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::{AtomicU8, AtomicUsize, Ordering};
9use std::sync::{Arc, OnceLock};
10
11use async_trait::async_trait;
12use lattice_embed::{
13    CachedEmbeddingService, EmbeddingModel, EmbeddingRole, EmbeddingService,
14    NativeEmbeddingService, DEFAULT_MAX_BATCH_SIZE, MAX_TEXT_BYTES,
15};
16use tokio::sync::{mpsc, Notify, OnceCell};
17
18use crate::error::{RuntimeError, RuntimeResult};
19
20const ADMISSION_WAITING: u8 = 1;
21const ADMISSION_ACCEPTED: u8 = 2;
22
23struct EmbeddingAdmission {
24    deadline: khive_storage::RequestReadDeadline,
25    request: khive_storage::RequestReadContext,
26    state: AtomicU8,
27    changed: Notify,
28}
29
30impl EmbeddingAdmission {
31    fn expired(&self) -> bool {
32        tokio::time::Instant::now() >= self.deadline.async_at()
33            || self.request.stop_reason().is_some()
34    }
35}
36
37tokio::task_local! {
38    static EMBEDDING_ADMISSION: Arc<EmbeddingAdmission>;
39}
40
41/// Bound only the built-in worker's pre-admission wait, above the lattice error
42/// boundary. Providers that never enter that wait retain their own semantics.
43pub(crate) async fn with_embedding_admission<T>(
44    future: impl std::future::Future<Output = lattice_embed::Result<T>>,
45) -> RuntimeResult<T> {
46    let deadline = khive_storage::effective_request_read_deadline(
47        khive_storage::RequestReadDeadline::after(khive_storage::request_read_timeout_from_env()),
48    );
49    let timeout = deadline
50        .async_at()
51        .saturating_duration_since(tokio::time::Instant::now());
52    let admission = Arc::new(EmbeddingAdmission {
53        deadline,
54        request: khive_storage::capture_request_read_context(),
55        state: AtomicU8::new(0),
56        changed: Notify::new(),
57    });
58    EMBEDDING_ADMISSION
59        .scope(Arc::clone(&admission), async move {
60            tokio::pin!(future);
61            let stopped = async {
62                tokio::select! {
63                    _ = tokio::time::sleep_until(deadline.async_at()) => {},
64                    _ = admission.request.clone().wait_for_stop() => {},
65                }
66            };
67            tokio::pin!(stopped);
68            let mut bound_expired = false;
69            loop {
70                tokio::select! {
71                    // A ready cache hit precedes admission, including an expired bound.
72                    biased;
73                    result = &mut future => return result.map_err(RuntimeError::from),
74                    _ = admission.changed.notified() => {},
75                    _ = &mut stopped, if !bound_expired => bound_expired = true,
76                }
77                match admission.state.load(Ordering::Acquire) {
78                    ADMISSION_ACCEPTED => return future.await.map_err(RuntimeError::from),
79                    ADMISSION_WAITING if bound_expired => {
80                        return Err(RuntimeError::Storage(
81                            khive_storage::StorageError::AdmissionTimeout {
82                                operation: "embedding admission".into(),
83                                timeout_ms: u64::try_from(timeout.as_millis()).unwrap_or(u64::MAX),
84                                pool_identity: None,
85                            },
86                        ));
87                    }
88                    // A provider may enter the owned adapter later. Keep the
89                    // original expired bound without timing out its other work.
90                    _ => {}
91                }
92            }
93        })
94        .await
95}
96
97#[derive(Clone, Copy)]
98enum EmbeddingCall {
99    Generic,
100    Query,
101    Passage,
102}
103
104const EMBEDDING_QUEUE_CAPACITY: usize = 32;
105const EMBEDDING_MAX_JOB_BYTES: usize = DEFAULT_MAX_BATCH_SIZE * MAX_TEXT_BYTES;
106// 32 queue slots × the normal 128-text batch × 32 KiB/text = 128 MiB in flight.
107const EMBEDDING_QUEUE_BYTE_BUDGET: usize = EMBEDDING_QUEUE_CAPACITY * 128 * MAX_TEXT_BYTES;
108
109struct InFlightBytes {
110    counter: Arc<AtomicUsize>,
111    bytes: usize,
112}
113
114impl InFlightBytes {
115    fn reserve(
116        counter: Arc<AtomicUsize>,
117        byte_budget: usize,
118        bytes: usize,
119    ) -> lattice_embed::Result<Self> {
120        let mut current = counter.load(Ordering::Acquire);
121        loop {
122            let Some(next) = current.checked_add(bytes) else {
123                return Err(lattice_embed::EmbedError::Internal(format!(
124                    "embedding worker byte budget exceeded: in-flight byte count overflowed the {byte_budget}-byte budget"
125                )));
126            };
127            if next > byte_budget {
128                return Err(lattice_embed::EmbedError::Internal(format!(
129                    "embedding worker byte budget exceeded: {current} in flight + {bytes} job bytes > {byte_budget}"
130                )));
131            }
132            match counter.compare_exchange_weak(current, next, Ordering::AcqRel, Ordering::Acquire)
133            {
134                Ok(_) => return Ok(Self { counter, bytes }),
135                Err(observed) => current = observed,
136            }
137        }
138    }
139}
140
141impl Drop for InFlightBytes {
142    fn drop(&mut self) {
143        let previous = self.counter.fetch_sub(self.bytes, Ordering::AcqRel);
144        debug_assert!(previous >= self.bytes, "embedding byte counter underflow");
145    }
146}
147
148struct EmbeddingJob {
149    texts: Vec<String>,
150    model: EmbeddingModel,
151    call: EmbeddingCall,
152    reply: tokio::sync::oneshot::Sender<lattice_embed::Result<Vec<Vec<f32>>>>,
153    _in_flight: InFlightBytes,
154}
155
156/// Bounds non-cancellable native inference to one worker and a fixed queue.
157/// Callers can detach safely because closed queued jobs are skipped before inference.
158pub(crate) struct BlockingEmbeddingService<S> {
159    inner: Arc<S>,
160    worker: OnceLock<Result<mpsc::Sender<EmbeddingJob>, String>>,
161    in_flight_bytes: Arc<AtomicUsize>,
162    byte_budget: usize,
163}
164
165impl<S> BlockingEmbeddingService<S> {
166    pub(crate) fn new(inner: Arc<S>) -> Self {
167        Self {
168            inner,
169            worker: OnceLock::new(),
170            in_flight_bytes: Arc::new(AtomicUsize::new(0)),
171            byte_budget: EMBEDDING_QUEUE_BYTE_BUDGET,
172        }
173    }
174
175    #[cfg(test)]
176    fn with_byte_budget(inner: Arc<S>, byte_budget: usize) -> Self {
177        Self {
178            inner,
179            worker: OnceLock::new(),
180            in_flight_bytes: Arc::new(AtomicUsize::new(0)),
181            byte_budget,
182        }
183    }
184}
185
186impl<S: EmbeddingService + 'static> BlockingEmbeddingService<S> {
187    fn input_bytes(texts: &[String]) -> lattice_embed::Result<usize> {
188        if texts.is_empty() {
189            return Err(lattice_embed::EmbedError::InvalidInput(
190                "no texts provided".to_owned(),
191            ));
192        }
193        let input_bytes = texts.iter().try_fold(0usize, |total, text| {
194            total.checked_add(text.len()).ok_or_else(|| {
195                lattice_embed::EmbedError::InvalidInput(format!(
196                    "embedding job input exceeds the {EMBEDDING_MAX_JOB_BYTES}-byte maximum"
197                ))
198            })
199        })?;
200        if input_bytes > EMBEDDING_MAX_JOB_BYTES {
201            return Err(lattice_embed::EmbedError::InvalidInput(format!(
202                "embedding job input is {input_bytes} bytes; maximum is {EMBEDDING_MAX_JOB_BYTES} bytes"
203            )));
204        }
205        if texts.len() > DEFAULT_MAX_BATCH_SIZE {
206            return Err(lattice_embed::EmbedError::InvalidInput(format!(
207                "batch size {} exceeds maximum {DEFAULT_MAX_BATCH_SIZE}",
208                texts.len()
209            )));
210        }
211        if let Some(text) = texts.iter().find(|text| text.len() > MAX_TEXT_BYTES) {
212            return Err(lattice_embed::EmbedError::TextTooLong {
213                length: text.len(),
214                max: MAX_TEXT_BYTES,
215            });
216        }
217        Ok(input_bytes)
218    }
219
220    fn worker(&self) -> lattice_embed::Result<&mpsc::Sender<EmbeddingJob>> {
221        self.worker
222            .get_or_init(|| {
223                let (sender, receiver) = mpsc::channel(EMBEDDING_QUEUE_CAPACITY);
224                let inner = Arc::clone(&self.inner);
225                let runtime = tokio::runtime::Handle::current();
226                std::thread::Builder::new()
227                    .name("khive-embedding".to_owned())
228                    .spawn(move || Self::run_worker(inner, runtime, receiver))
229                    .map(|_| sender)
230                    .map_err(|error| error.to_string())
231            })
232            .as_ref()
233            .map_err(|error| lattice_embed::EmbedError::Internal(error.clone()))
234    }
235
236    fn run_worker(
237        inner: Arc<S>,
238        runtime: tokio::runtime::Handle,
239        mut receiver: mpsc::Receiver<EmbeddingJob>,
240    ) {
241        while let Some(job) = receiver.blocking_recv() {
242            if job.reply.is_closed() {
243                continue;
244            }
245            let result = runtime.block_on(async {
246                match job.call {
247                    EmbeddingCall::Generic => inner.embed(&job.texts, job.model).await,
248                    EmbeddingCall::Query => inner.embed_query(&job.texts, job.model).await,
249                    EmbeddingCall::Passage => inner.embed_passage(&job.texts, job.model).await,
250                }
251            });
252            let _ = job.reply.send(result);
253        }
254    }
255
256    async fn run(
257        &self,
258        texts: &[String],
259        model: EmbeddingModel,
260        call: EmbeddingCall,
261    ) -> lattice_embed::Result<Vec<Vec<f32>>> {
262        let input_bytes = Self::input_bytes(texts)?;
263        let sender = self.worker()?;
264        let admission = EMBEDDING_ADMISSION.try_with(Arc::clone).ok();
265        let permit = if let Some(admission) = &admission {
266            admission.state.store(ADMISSION_WAITING, Ordering::Release);
267            admission.changed.notify_one();
268            if admission.expired() {
269                return std::future::pending().await;
270            }
271            sender.reserve().await.map_err(|_| {
272                lattice_embed::EmbedError::Internal(
273                    "embedding worker channel is disconnected".to_owned(),
274                )
275            })?
276        } else {
277            // Direct trait callers have no RuntimeResult boundary at which to
278            // report a typed admission timeout, so retain finite fail-fast admission.
279            sender.try_reserve().map_err(|error| match error {
280                mpsc::error::TrySendError::Full(_) => {
281                    lattice_embed::EmbedError::Internal("embedding worker queue is full".to_owned())
282                }
283                mpsc::error::TrySendError::Closed(_) => lattice_embed::EmbedError::Internal(
284                    "embedding worker channel is disconnected".to_owned(),
285                ),
286            })?
287        };
288        // Capacity and expiry can become ready in one poll. Do not publish a
289        // job after the original bound, even when reserve() returned a permit.
290        if admission
291            .as_ref()
292            .is_some_and(|admission| admission.expired())
293        {
294            return std::future::pending().await;
295        }
296        let in_flight = InFlightBytes::reserve(
297            Arc::clone(&self.in_flight_bytes),
298            self.byte_budget,
299            input_bytes,
300        )?;
301        let (reply, receiver) = tokio::sync::oneshot::channel();
302        let job = EmbeddingJob {
303            texts: texts.to_vec(),
304            model,
305            call,
306            reply,
307            _in_flight: in_flight,
308        };
309        if let Some(admission) = &admission {
310            if admission.expired() {
311                return std::future::pending().await;
312            }
313            admission.state.store(ADMISSION_ACCEPTED, Ordering::Release);
314            admission.changed.notify_one();
315        }
316        permit.send(job);
317        receiver
318            .await
319            .map_err(|error| lattice_embed::EmbedError::Internal(error.to_string()))?
320    }
321}
322
323#[async_trait]
324impl<S: EmbeddingService + 'static> EmbeddingService for BlockingEmbeddingService<S> {
325    async fn embed(
326        &self,
327        texts: &[String],
328        model: EmbeddingModel,
329    ) -> lattice_embed::Result<Vec<Vec<f32>>> {
330        self.run(texts, model, EmbeddingCall::Generic).await
331    }
332
333    async fn embed_with_role(
334        &self,
335        texts: &[String],
336        model: EmbeddingModel,
337        role: EmbeddingRole,
338    ) -> lattice_embed::Result<Vec<Vec<f32>>> {
339        // The outer cache delegates raw caller text here; preparation belongs
340        // to the native service, after the caller-input admission checks.
341        let call = match role {
342            EmbeddingRole::Generic => EmbeddingCall::Generic,
343            EmbeddingRole::Query => EmbeddingCall::Query,
344            EmbeddingRole::Passage => EmbeddingCall::Passage,
345            _ => {
346                return Err(lattice_embed::EmbedError::InvalidInput(
347                    "unsupported embedding role".to_owned(),
348                ))
349            }
350        };
351        self.run(texts, model, call).await
352    }
353
354    async fn embed_query(
355        &self,
356        texts: &[String],
357        model: EmbeddingModel,
358    ) -> lattice_embed::Result<Vec<Vec<f32>>> {
359        self.run(texts, model, EmbeddingCall::Query).await
360    }
361
362    async fn embed_passage(
363        &self,
364        texts: &[String],
365        model: EmbeddingModel,
366    ) -> lattice_embed::Result<Vec<Vec<f32>>> {
367        self.run(texts, model, EmbeddingCall::Passage).await
368    }
369
370    fn model_config(&self, model: EmbeddingModel) -> lattice_embed::ModelConfig {
371        self.inner.model_config(model)
372    }
373
374    fn supports_model(&self, model: EmbeddingModel) -> bool {
375        self.inner.supports_model(model)
376    }
377
378    fn name(&self) -> &'static str {
379        self.inner.name()
380    }
381}
382
383/// A source that can produce an [`EmbeddingService`] by name.
384///
385/// Packs implement this trait to register custom embedding backends.
386/// The runtime calls [`build`](EmbedderProvider::build) lazily — once per
387/// process per model — and caches the result. Subsequent calls to
388/// `KhiveRuntime::embedder(name)` are cheap.
389///
390/// Built-in lattice models are registered automatically via
391/// [`LatticeEmbedderProvider`]; packs need not re-register them.
392#[async_trait]
393pub trait EmbedderProvider: Send + Sync {
394    /// Stable, case-sensitive name for this embedder.
395    ///
396    /// Must be unique across all registered providers. The name is used as
397    /// the key in `KhiveRuntime::embedder(name)` lookups and as the storage
398    /// table suffix for vector indices. Use the model's canonical short form
399    /// (e.g. `"all-minilm-l6-v2"`, `"my-custom-encoder"`).
400    fn name(&self) -> &str;
401
402    /// Output vector dimension for this embedder.
403    ///
404    /// Must be consistent with what [`build`](Self::build) produces.
405    /// The runtime uses this to pre-register the vector store columns.
406    fn dimensions(&self) -> usize;
407
408    /// Construct the underlying [`EmbeddingService`].
409    ///
410    /// Called at most once per process. The result is cached in a
411    /// [`OnceCell`]; concurrent callers block on the first call and share
412    /// the result thereafter.
413    async fn build(&self) -> RuntimeResult<Arc<dyn EmbeddingService>>;
414}
415
416/// An entry in the [`EmbedderRegistry`] combining a provider with its
417/// lazy-initialized service.
418pub(crate) struct EmbedderEntry {
419    provider: Arc<dyn EmbedderProvider>,
420    cell: Arc<OnceCell<Arc<dyn EmbeddingService>>>,
421    /// Only the runtime's built-in lattice provider has an audited document
422    /// preparation path. Pack replacements, even under a built-in name, do not.
423    audited_document_preparation: bool,
424}
425
426impl Clone for EmbedderEntry {
427    fn clone(&self) -> Self {
428        Self {
429            provider: Arc::clone(&self.provider),
430            cell: Arc::clone(&self.cell),
431            audited_document_preparation: self.audited_document_preparation,
432        }
433    }
434}
435
436/// Registry of named [`EmbedderProvider`] instances.
437///
438/// Built during `KhiveRuntime` construction and optionally extended by packs
439/// via [`crate::KhiveRuntime::register_embedder`]. The registry is internally
440/// reference-counted so `KhiveRuntime::clone()` shares the same providers
441/// and cached service instances.
442#[derive(Clone, Default)]
443pub struct EmbedderRegistry {
444    entries: HashMap<String, EmbedderEntry>,
445}
446
447impl EmbedderRegistry {
448    /// Create an empty registry.
449    pub fn new() -> Self {
450        Self {
451            entries: HashMap::new(),
452        }
453    }
454
455    /// Register a provider.
456    ///
457    /// If a provider with the same [`name`](EmbedderProvider::name) already
458    /// exists, it is replaced (last-writer wins) and any cached service is
459    /// discarded, since pack registration order is not guaranteed and packs
460    /// may legitimately override a default model under the same name.
461    /// Callers needing strict collision detection should check
462    /// [`names`](Self::names) before registering.
463    pub fn register<P: EmbedderProvider + 'static>(&mut self, provider: P) {
464        self.insert(provider, false);
465    }
466
467    /// Register the runtime-owned lattice adapter whose passage preparation is
468    /// audited. This is deliberately not available to pack providers.
469    pub(crate) fn register_builtin(&mut self, provider: LatticeEmbedderProvider) {
470        self.insert(provider, true);
471    }
472
473    /// Test-only attested provider. The wrapper below owns passage preparation,
474    /// so a fake backend cannot change the bytes whose digest is recorded.
475    #[cfg(feature = "test-internals")]
476    pub fn register_test_audited<P: EmbedderProvider + 'static>(
477        &mut self,
478        model: EmbeddingModel,
479        provider: P,
480    ) {
481        assert_eq!(provider.name(), model.to_string());
482        self.insert(TestAuditedProvider { provider }, true);
483    }
484
485    fn insert<P: EmbedderProvider + 'static>(
486        &mut self,
487        provider: P,
488        audited_document_preparation: bool,
489    ) {
490        let name = provider.name().to_owned();
491        self.entries.insert(
492            name,
493            EmbedderEntry {
494                provider: Arc::new(provider),
495                cell: Arc::new(OnceCell::new()),
496                audited_document_preparation,
497            },
498        );
499    }
500
501    /// Look up a provider by name.
502    pub fn get_provider(&self, name: &str) -> Option<&dyn EmbedderProvider> {
503        self.entries.get(name).map(|e| e.provider.as_ref())
504    }
505
506    /// Returns `true` if a provider with this name is registered.
507    pub fn contains(&self, name: &str) -> bool {
508        self.entries.contains_key(name)
509    }
510
511    /// Names of all registered providers, in unspecified order.
512    pub fn names(&self) -> Vec<String> {
513        self.entries.keys().cloned().collect()
514    }
515
516    /// Return a cloned entry for `name` without holding any lock.
517    ///
518    /// The caller can then call [`EmbedderEntry::resolve`] without holding
519    /// a lock — this avoids holding a `RwLockGuard` across `await` points.
520    /// Returns `None` if `name` is not registered.
521    pub(crate) fn get_entry(&self, name: &str) -> Option<EmbedderEntry> {
522        self.entries.get(name).cloned()
523    }
524
525    /// Lazily resolve a registered provider to its live [`EmbeddingService`].
526    ///
527    /// Returns [`RuntimeError::UnknownModel`] if `name` is not registered.
528    /// The first call for a given name triggers [`EmbedderProvider::build`];
529    /// subsequent calls return the cached `Arc`.
530    ///
531    /// Prefer [`crate::KhiveRuntime::embedder`] over calling this directly from pack
532    /// handlers — the runtime method handles alias resolution and error mapping.
533    pub async fn get_service(&self, name: &str) -> RuntimeResult<Arc<dyn EmbeddingService>> {
534        let entry = self
535            .entries
536            .get(name)
537            .ok_or_else(|| RuntimeError::UnknownModel(name.to_string()))?
538            .clone();
539
540        Ok(entry.resolve().await?.0)
541    }
542}
543
544#[cfg(feature = "test-internals")]
545struct TestAuditedProvider<P> {
546    provider: P,
547}
548
549#[cfg(feature = "test-internals")]
550#[async_trait]
551impl<P: EmbedderProvider> EmbedderProvider for TestAuditedProvider<P> {
552    fn name(&self) -> &str {
553        self.provider.name()
554    }
555
556    fn dimensions(&self) -> usize {
557        self.provider.dimensions()
558    }
559
560    async fn build(&self) -> RuntimeResult<Arc<dyn EmbeddingService>> {
561        Ok(Arc::new(TestAuditedService(self.provider.build().await?)))
562    }
563}
564
565/// Deliberately does not delegate `embed_passage`: the lattice trait default
566/// applies the model's document instruction before calling `embed`, exactly
567/// as the audited built-in path does. Test providers supply only fake vectors.
568#[cfg(feature = "test-internals")]
569struct TestAuditedService(Arc<dyn EmbeddingService>);
570
571#[cfg(feature = "test-internals")]
572#[async_trait]
573impl EmbeddingService for TestAuditedService {
574    async fn embed(
575        &self,
576        texts: &[String],
577        model: EmbeddingModel,
578    ) -> lattice_embed::Result<Vec<Vec<f32>>> {
579        self.0.embed(texts, model).await
580    }
581
582    fn supports_model(&self, model: EmbeddingModel) -> bool {
583        self.0.supports_model(model)
584    }
585
586    fn name(&self) -> &'static str {
587        self.0.name()
588    }
589}
590
591impl EmbedderEntry {
592    pub(crate) fn cached_service(&self) -> Option<Arc<dyn EmbeddingService>> {
593        self.cell.get().map(Arc::clone)
594    }
595
596    pub(crate) fn has_audited_document_preparation(&self) -> bool {
597        self.audited_document_preparation
598    }
599
600    /// Lazily initialise and return the embedding service for this entry.
601    ///
602    /// `OnceCell::get_or_try_init` single-flights concurrent cold callers: only
603    /// one of them runs [`EmbedderProvider::build`], and every other waiter
604    /// receives the same result once it completes. Only the caller whose task
605    /// actually ran `build()` gets back `Some(duration)`; every other caller
606    /// (cache hit or wait-for-in-flight-build) gets `None`.
607    ///
608    /// Returns `RuntimeError` if `build()` fails, rather than panicking. On
609    /// failure the cell stays unset, so a later call retries the build.
610    pub(crate) async fn resolve(self) -> RuntimeResult<(Arc<dyn EmbeddingService>, Option<i64>)> {
611        let mut own_init_duration_us: Option<i64> = None;
612        let provider = Arc::clone(&self.provider);
613        let init_duration_us = &mut own_init_duration_us;
614        let svc = self
615            .cell
616            .get_or_try_init(|| async move {
617                let init_start = std::time::Instant::now();
618                let svc = provider.build().await.map_err(|e| {
619                    crate::error::RuntimeError::Internal(format!(
620                        "EmbedderProvider '{}' build() failed: {e}",
621                        provider.name()
622                    ))
623                })?;
624                *init_duration_us = Some(init_start.elapsed().as_micros() as i64);
625                Ok::<_, RuntimeError>(svc)
626            })
627            .await?;
628        Ok((Arc::clone(svc), own_init_duration_us))
629    }
630}
631
632// ── LatticeEmbedderProvider ───────────────────────────────────────────────────
633
634/// Adapter that wraps a [`lattice_embed::EmbeddingModel`] as an
635/// [`EmbedderProvider`].
636///
637/// All built-in models (MiniLM, paraphrase-multilingual, BGE variants, etc.)
638/// are registered as `LatticeEmbedderProvider` instances during
639/// `KhiveRuntime` construction. External callers do not need to use this type
640/// unless they are constructing a custom registry from scratch.
641pub struct LatticeEmbedderProvider {
642    model: EmbeddingModel,
643    /// Cached `to_string()` result so `name()` can return `&str`.
644    name: String,
645}
646
647impl LatticeEmbedderProvider {
648    /// Create a new provider wrapping the given lattice model.
649    pub fn new(model: EmbeddingModel) -> Self {
650        let name = model.to_string();
651        Self { model, name }
652    }
653}
654
655#[async_trait]
656impl EmbedderProvider for LatticeEmbedderProvider {
657    fn name(&self) -> &str {
658        &self.name
659    }
660
661    fn dimensions(&self) -> usize {
662        self.model.dimensions()
663    }
664
665    async fn build(&self) -> RuntimeResult<Arc<dyn EmbeddingService>> {
666        let native = Arc::new(NativeEmbeddingService::with_model(self.model));
667        native.ensure_loaded().await?;
668        Ok(cached_blocking_service(native))
669    }
670}
671
672/// Keep result-cache lookup outside worker admission for every retrieval role.
673fn cached_blocking_service<S: EmbeddingService + 'static>(
674    inner: Arc<S>,
675) -> Arc<dyn EmbeddingService> {
676    let blocking = Arc::new(BlockingEmbeddingService::new(inner));
677    Arc::new(CachedEmbeddingService::with_default_cache(blocking))
678}
679
680#[cfg(test)]
681#[path = "embedding_cache_admission_tests.rs"]
682mod cache_admission_tests;
683
684// ── Unit tests ────────────────────────────────────────────────────────────────
685
686#[cfg(test)]
687mod tests {
688    use super::*;
689    use std::collections::HashSet;
690    use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
691    use std::sync::{Condvar, Mutex};
692    use std::time::Duration;
693    use tokio::sync::Notify;
694
695    struct ConstVecProvider {
696        name: String,
697        dims: usize,
698        build_calls: Arc<AtomicUsize>,
699    }
700
701    impl ConstVecProvider {
702        fn new(name: &str, dims: usize) -> Self {
703            Self {
704                name: name.to_owned(),
705                dims,
706                build_calls: Arc::new(AtomicUsize::new(0)),
707            }
708        }
709    }
710
711    /// A trivial embedding service that returns a constant vector of `1.0`s.
712    /// The `model` parameter is ignored — this service always returns the
713    /// same synthetic vector regardless of which model is requested.
714    pub(super) struct ConstVecService {
715        pub(super) dims: usize,
716    }
717
718    #[async_trait]
719    impl EmbeddingService for ConstVecService {
720        async fn embed(
721            &self,
722            texts: &[String],
723            _model: EmbeddingModel,
724        ) -> std::result::Result<Vec<Vec<f32>>, lattice_embed::EmbedError> {
725            Ok(texts.iter().map(|_| vec![1.0_f32; self.dims]).collect())
726        }
727
728        fn supports_model(&self, _model: EmbeddingModel) -> bool {
729            true
730        }
731
732        fn name(&self) -> &'static str {
733            "const-vec-service"
734        }
735    }
736
737    #[async_trait]
738    impl EmbedderProvider for ConstVecProvider {
739        fn name(&self) -> &str {
740            &self.name
741        }
742
743        fn dimensions(&self) -> usize {
744            self.dims
745        }
746
747        async fn build(&self) -> RuntimeResult<Arc<dyn EmbeddingService>> {
748            self.build_calls.fetch_add(1, Ordering::SeqCst);
749            Ok(Arc::new(ConstVecService { dims: self.dims }))
750        }
751    }
752
753    #[test]
754    fn builtin_input_attestation_stays_with_cloned_entry_after_canonical_name_override() {
755        let model = EmbeddingModel::MultilingualE5Small;
756        let name = model.to_string();
757        let mut registry = EmbedderRegistry::new();
758        registry.register_builtin(LatticeEmbedderProvider::new(model));
759        let builtin_entry = registry.get_entry(&name).expect("builtin entry");
760        assert!(builtin_entry.has_audited_document_preparation());
761
762        registry.register(ConstVecProvider::new(&name, model.dimensions()));
763        let replacement_entry = registry.get_entry(&name).expect("replacement entry");
764        assert!(builtin_entry.has_audited_document_preparation());
765        assert!(!replacement_entry.has_audited_document_preparation());
766    }
767
768    struct FirstLoadBlockingService {
769        loaded: AtomicBool,
770    }
771
772    #[async_trait]
773    impl EmbeddingService for FirstLoadBlockingService {
774        async fn embed(
775            &self,
776            texts: &[String],
777            _model: EmbeddingModel,
778        ) -> lattice_embed::Result<Vec<Vec<f32>>> {
779            if !self.loaded.swap(true, Ordering::SeqCst) {
780                tokio::task::spawn_blocking(|| {})
781                    .await
782                    .map_err(|error| lattice_embed::EmbedError::Internal(error.to_string()))?;
783            }
784            Ok(texts.iter().map(|_| vec![1.0]).collect())
785        }
786
787        fn supports_model(&self, _model: EmbeddingModel) -> bool {
788            true
789        }
790
791        fn name(&self) -> &'static str {
792            "first-load-blocking-service"
793        }
794    }
795
796    pub(super) struct BlockingTestService {
797        pub(super) calls: Mutex<Vec<String>>,
798        pub(super) entered: AtomicUsize,
799        release: (Mutex<bool>, Condvar),
800        pub(super) thread_ids: Mutex<HashSet<std::thread::ThreadId>>,
801    }
802
803    impl BlockingTestService {
804        pub(super) fn new() -> Self {
805            Self {
806                calls: Mutex::new(Vec::new()),
807                entered: AtomicUsize::new(0),
808                release: (Mutex::new(false), Condvar::new()),
809                thread_ids: Mutex::new(HashSet::new()),
810            }
811        }
812
813        pub(super) fn release(&self) {
814            *self
815                .release
816                .0
817                .lock()
818                .expect("release lock must not be poisoned") = true;
819            self.release.1.notify_all();
820        }
821    }
822
823    #[async_trait]
824    impl EmbeddingService for BlockingTestService {
825        async fn embed(
826            &self,
827            texts: &[String],
828            _model: EmbeddingModel,
829        ) -> lattice_embed::Result<Vec<Vec<f32>>> {
830            let text = texts.first().cloned().unwrap_or_default();
831            self.thread_ids
832                .lock()
833                .expect("thread id lock must not be poisoned")
834                .insert(std::thread::current().id());
835            self.calls
836                .lock()
837                .expect("call lock must not be poisoned")
838                .push(text.clone());
839            self.entered.fetch_add(1, Ordering::Release);
840
841            if text != "later" {
842                let (released, wake) = &self.release;
843                let guard = released.lock().expect("release lock must not be poisoned");
844                let _guard = wake
845                    .wait_while(guard, |released| !*released)
846                    .expect("release lock must not be poisoned");
847            }
848
849            Ok(texts.iter().map(|_| vec![1.0]).collect())
850        }
851
852        fn supports_model(&self, _model: EmbeddingModel) -> bool {
853            true
854        }
855
856        fn name(&self) -> &'static str {
857            "blocking-test-service"
858        }
859    }
860
861    #[test]
862    fn blocking_adapter_first_use_completes_with_single_blocking_thread() {
863        let runtime = tokio::runtime::Builder::new_current_thread()
864            .enable_time()
865            .max_blocking_threads(1)
866            .build()
867            .expect("current-thread runtime must build");
868        let service = BlockingEmbeddingService::new(Arc::new(FirstLoadBlockingService {
869            loaded: AtomicBool::new(false),
870        }));
871
872        let result = runtime.block_on(async {
873            tokio::time::timeout(
874                Duration::from_secs(5),
875                service.embed(&["first use".to_owned()], EmbeddingModel::default()),
876            )
877            .await
878        });
879        runtime.shutdown_timeout(Duration::from_secs(1));
880
881        let embeddings = result
882            .expect("first-use embedding must not exhaust the blocking pool")
883            .expect("first-use embedding must succeed");
884        assert_eq!(embeddings, vec![vec![1.0]]);
885    }
886
887    #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
888    async fn blocking_adapter_uses_one_worker_for_concurrent_calls() {
889        const CALL_COUNT: usize = 32;
890        let inner = Arc::new(BlockingTestService::new());
891        let service = Arc::new(BlockingEmbeddingService::new(Arc::clone(&inner)));
892        let mut calls = Vec::with_capacity(CALL_COUNT);
893
894        for index in 0..CALL_COUNT {
895            let service = Arc::clone(&service);
896            calls.push(tokio::spawn(async move {
897                service
898                    .embed(&[format!("request-{index}")], EmbeddingModel::default())
899                    .await
900            }));
901        }
902
903        let _ = tokio::time::timeout(Duration::from_millis(250), async {
904            while inner.entered.load(Ordering::Acquire) < 2 {
905                tokio::task::yield_now().await;
906            }
907        })
908        .await;
909        inner.release();
910
911        for call in calls {
912            call.await
913                .expect("embedding task must not panic")
914                .expect("embedding call must succeed");
915        }
916        assert_eq!(
917            inner
918                .thread_ids
919                .lock()
920                .expect("thread id lock must not be poisoned")
921                .len(),
922            1,
923            "concurrent calls must share one native worker thread"
924        );
925    }
926
927    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
928    async fn blocking_adapter_skips_timed_out_call_and_serves_later_call() {
929        let inner = Arc::new(BlockingTestService::new());
930        let service = Arc::new(BlockingEmbeddingService::new(Arc::clone(&inner)));
931
932        let first_service = Arc::clone(&service);
933        let first = tokio::spawn(async move {
934            first_service
935                .embed(&["first".to_owned()], EmbeddingModel::default())
936                .await
937        });
938        tokio::time::timeout(Duration::from_secs(1), async {
939            while inner.entered.load(Ordering::Acquire) == 0 {
940                tokio::task::yield_now().await;
941            }
942        })
943        .await
944        .expect("first embedding call must enter the native service");
945
946        let abandoned = tokio::time::timeout(
947            Duration::from_millis(50),
948            service.embed(&["abandoned".to_owned()], EmbeddingModel::default()),
949        )
950        .await;
951        let later_service = Arc::clone(&service);
952        let later = tokio::spawn(async move {
953            later_service
954                .embed(&["later".to_owned()], EmbeddingModel::default())
955                .await
956        });
957        inner.release();
958
959        first
960            .await
961            .expect("first embedding task must not panic")
962            .expect("first embedding call must succeed");
963        let later_result = tokio::time::timeout(Duration::from_secs(1), later)
964            .await
965            .expect("later embedding call must be served")
966            .expect("later embedding task must not panic")
967            .expect("later embedding call must succeed");
968
969        assert!(abandoned.is_err(), "queued embedding call must time out");
970        assert_eq!(later_result, vec![vec![1.0]]);
971        assert_eq!(
972            *inner.calls.lock().expect("call lock must not be poisoned"),
973            vec!["first".to_owned(), "later".to_owned()],
974            "the worker must skip a queued call whose receiver is closed"
975        );
976    }
977
978    pub(super) struct ReleaseWorkerOnDrop(pub(super) Arc<BlockingTestService>);
979
980    impl Drop for ReleaseWorkerOnDrop {
981        fn drop(&mut self) {
982            self.0.release();
983        }
984    }
985
986    struct ServiceProvider(Arc<dyn EmbeddingService>);
987
988    #[async_trait]
989    impl EmbedderProvider for ServiceProvider {
990        fn name(&self) -> &str {
991            "queue-test"
992        }
993        fn dimensions(&self) -> usize {
994            1
995        }
996        async fn build(&self) -> RuntimeResult<Arc<dyn EmbeddingService>> {
997            Ok(Arc::clone(&self.0))
998        }
999    }
1000
1001    pub(super) fn queue_runtime(service: Arc<dyn EmbeddingService>) -> crate::KhiveRuntime {
1002        let runtime = crate::KhiveRuntime::memory().expect("memory runtime");
1003        runtime.register_embedder(ServiceProvider(service));
1004        runtime
1005    }
1006
1007    pub(super) async fn poll_once<F: std::future::Future + ?Sized>(
1008        mut future: std::pin::Pin<&mut F>,
1009    ) -> std::task::Poll<F::Output> {
1010        std::future::poll_fn(|cx| std::task::Poll::Ready(future.as_mut().poll(cx))).await
1011    }
1012
1013    pub(super) async fn wait_for_entered(inner: &BlockingTestService, count: usize) {
1014        let watchdog = std::time::Instant::now() + Duration::from_secs(2);
1015        while inner.entered.load(Ordering::Acquire) < count {
1016            assert!(
1017                std::time::Instant::now() < watchdog,
1018                "setup watchdog: native worker did not enter"
1019            );
1020            tokio::task::yield_now().await;
1021        }
1022    }
1023
1024    async fn drive_until_entered<F: std::future::Future + ?Sized>(
1025        mut future: std::pin::Pin<&mut F>,
1026        inner: &BlockingTestService,
1027        count: usize,
1028    ) {
1029        let watchdog = std::time::Instant::now() + Duration::from_secs(2);
1030        while inner.entered.load(Ordering::Acquire) < count {
1031            assert!(
1032                std::time::Instant::now() < watchdog,
1033                "setup watchdog: driven runtime call did not reach the native worker"
1034            );
1035            assert!(
1036                poll_once(future.as_mut()).await.is_pending(),
1037                "held inference must remain pending while the runtime call is driven"
1038            );
1039            tokio::task::yield_now().await;
1040        }
1041    }
1042
1043    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
1044    async fn runtime_embedding_waits_for_capacity_and_drains_n_plus_one_calls() {
1045        let inner = Arc::new(BlockingTestService::new());
1046        let _release = ReleaseWorkerOnDrop(Arc::clone(&inner));
1047        let service = Arc::new(BlockingEmbeddingService::new(Arc::clone(&inner)));
1048        let runtime = queue_runtime(service.clone());
1049        runtime
1050            .embedder("queue-test")
1051            .await
1052            .expect("resolve provider before the admission fixture");
1053        let first_text = vec!["first".to_owned()];
1054        let mut first = Box::pin(runtime.embed_batch_with_model("queue-test", &first_text));
1055        assert!(poll_once(first.as_mut()).await.is_pending());
1056        drive_until_entered(first.as_mut(), &inner, 1).await;
1057
1058        let texts: Vec<_> = (0..EMBEDDING_QUEUE_CAPACITY)
1059            .map(|index| vec![format!("queued-{index}")])
1060            .collect();
1061        let mut queued = Vec::new();
1062        for text in &texts {
1063            let mut call = Box::pin(runtime.embed_batch_with_model("queue-test", text));
1064            assert!(
1065                poll_once(call.as_mut()).await.is_pending(),
1066                "the bounded queue must admit its N actual runtime calls"
1067            );
1068            queued.push(call);
1069        }
1070        assert_eq!(
1071            service.in_flight_bytes.load(Ordering::Acquire),
1072            first_text.iter().map(String::len).sum::<usize>()
1073                + texts.iter().flatten().map(String::len).sum::<usize>(),
1074            "the held worker and N queued runtime jobs must own their exact input bytes"
1075        );
1076        let overflow_text = vec!["overflow".to_owned()];
1077        let mut overflow = Box::pin(runtime.embed_batch_with_model("queue-test", &overflow_text));
1078        let overflow_before_release = poll_once(overflow.as_mut()).await;
1079        inner.release();
1080        assert!(
1081            overflow_before_release.is_pending(),
1082            "request N+1 must wait for a slot instead of failing: {overflow_before_release:?}"
1083        );
1084        assert_eq!(first.await.unwrap(), vec![vec![1.0]]);
1085        for call in queued {
1086            assert_eq!(call.await.unwrap(), vec![vec![1.0]]);
1087        }
1088        assert_eq!(overflow.await.unwrap(), vec![vec![1.0]]);
1089        assert_eq!(
1090            inner.entered.load(Ordering::Acquire),
1091            EMBEDDING_QUEUE_CAPACITY + 2
1092        );
1093        assert_eq!(inner.thread_ids.lock().unwrap().len(), 1);
1094    }
1095
1096    #[tokio::test(start_paused = true)]
1097    async fn runtime_embedding_expired_admission_is_retryable_and_never_enqueued() {
1098        let inner = Arc::new(BlockingTestService::new());
1099        let _release = ReleaseWorkerOnDrop(Arc::clone(&inner));
1100        let service = Arc::new(BlockingEmbeddingService::new(Arc::clone(&inner)));
1101        let runtime = queue_runtime(service.clone());
1102        runtime
1103            .embedder("queue-test")
1104            .await
1105            .expect("resolve provider before the admission fixture");
1106        let first_text = vec!["first".to_owned()];
1107        let mut first = Box::pin(runtime.embed_batch_with_model("queue-test", &first_text));
1108        assert!(poll_once(first.as_mut()).await.is_pending());
1109        drive_until_entered(first.as_mut(), &inner, 1).await;
1110        let texts: Vec<_> = (0..EMBEDDING_QUEUE_CAPACITY)
1111            .map(|index| vec![format!("queued-{index}")])
1112            .collect();
1113        let mut queued = Vec::new();
1114        for text in &texts {
1115            let mut call = Box::pin(runtime.embed_batch_with_model("queue-test", text));
1116            assert!(poll_once(call.as_mut()).await.is_pending());
1117            queued.push(call);
1118        }
1119        assert_eq!(
1120            service.in_flight_bytes.load(Ordering::Acquire),
1121            first_text.iter().map(String::len).sum::<usize>()
1122                + texts.iter().flatten().map(String::len).sum::<usize>(),
1123            "the held worker and N queued runtime jobs must own their exact input bytes"
1124        );
1125        let deadline = khive_storage::RequestReadDeadline::after(Duration::from_millis(100));
1126        let mut expired = Box::pin(khive_storage::scope_request_read_deadline_at(
1127            deadline,
1128            runtime.embed_with_model("queue-test", "expired"),
1129        ));
1130        assert!(
1131            poll_once(expired.as_mut()).await.is_pending(),
1132            "an unexpired full-queue call must wait"
1133        );
1134        tokio::time::advance(Duration::from_millis(100)).await;
1135        let result = poll_once(expired.as_mut()).await;
1136        inner.release();
1137        let error = match result {
1138            std::task::Poll::Ready(Err(error)) => error,
1139            other => panic!("expired admission must produce a typed runtime error: {other:?}"),
1140        };
1141        assert!(
1142            matches!(&error, RuntimeError::Storage(khive_storage::StorageError::AdmissionTimeout {
1143            operation, timeout_ms: 100, pool_identity: None,
1144        }) if operation == "embedding admission"),
1145            "{error:?}"
1146        );
1147        assert!(
1148            error.retryable_failure_context().is_some(),
1149            "pre-admission expiry must use the existing retryable classification"
1150        );
1151        let projected = crate::error_projection::runtime_error_value(
1152            error,
1153            crate::DomainDisposition::NotCommitted,
1154        );
1155        assert_eq!(projected["retryable"], true);
1156        assert_eq!(projected["operation"], "embedding admission");
1157        assert_eq!(first.await.unwrap(), vec![vec![1.0]]);
1158        for call in queued {
1159            assert_eq!(call.await.unwrap(), vec![vec![1.0]]);
1160        }
1161        assert!(
1162            !inner
1163                .calls
1164                .lock()
1165                .unwrap()
1166                .iter()
1167                .any(|text| text == "expired"),
1168            "an expired waiter must never reach inference"
1169        );
1170    }
1171
1172    #[tokio::test(start_paused = true)]
1173    async fn runtime_embedding_earlier_absolute_bound_is_not_renewed_when_slot_is_ready() {
1174        let inner = Arc::new(BlockingTestService::new());
1175        let _release = ReleaseWorkerOnDrop(Arc::clone(&inner));
1176        let runtime = queue_runtime(Arc::new(BlockingEmbeddingService::new(Arc::clone(&inner))));
1177        runtime
1178            .embedder("queue-test")
1179            .await
1180            .expect("resolve provider before the admission deadline");
1181        let deadline = khive_storage::RequestReadDeadline::after(Duration::from_millis(100));
1182        tokio::time::advance(Duration::from_millis(100)).await;
1183        let mut call = Box::pin(khive_storage::scope_request_read_deadline_at(
1184            deadline,
1185            runtime.embed_with_model("queue-test", "expired-ready-slot"),
1186        ));
1187        let result = poll_once(call.as_mut()).await;
1188        inner.release();
1189        assert!(
1190            matches!(
1191                result,
1192                std::task::Poll::Ready(Err(RuntimeError::Storage(
1193                    khive_storage::StorageError::AdmissionTimeout { timeout_ms: 0, .. }
1194                )))
1195            ),
1196            "an already-expired original bound must refuse even an available slot: {result:?}"
1197        );
1198        assert_eq!(
1199            inner.entered.load(Ordering::Acquire),
1200            0,
1201            "capacity becoming ready must not bypass expiry"
1202        );
1203    }
1204
1205    #[tokio::test(start_paused = true)]
1206    async fn runtime_embedding_all_six_call_families_refuse_expired_admission() {
1207        let inner = Arc::new(BlockingTestService::new());
1208        let _release = ReleaseWorkerOnDrop(Arc::clone(&inner));
1209        let runtime = queue_runtime(Arc::new(BlockingEmbeddingService::new(Arc::clone(&inner))));
1210        let texts = vec!["later".to_owned()];
1211        for family in 0..6 {
1212            let result = khive_storage::scope_request_read_deadline(Duration::ZERO, async {
1213                match family {
1214                    0 => runtime.embed_with_model("queue-test", "later").await,
1215                    1 => runtime
1216                        .embed_document_with_model_outcome("queue-test", "later")
1217                        .await
1218                        .map(|outcome| outcome.vector),
1219                    2 => runtime.embed_query_with_model("queue-test", "later").await,
1220                    3 => runtime
1221                        .embed_batch_with_model("queue-test", &texts)
1222                        .await
1223                        .map(|vectors| vectors[0].clone()),
1224                    4 => runtime
1225                        .embed_document_batch_with_model("queue-test", &texts)
1226                        .await
1227                        .map(|vectors| vectors[0].clone()),
1228                    _ => runtime
1229                        .embed_query_batch_with_model("queue-test", &texts)
1230                        .await
1231                        .map(|vectors| vectors[0].clone()),
1232                }
1233            })
1234            .await;
1235            assert!(
1236                matches!(
1237                    result,
1238                    Err(RuntimeError::Storage(
1239                        khive_storage::StorageError::AdmissionTimeout { .. }
1240                    ))
1241                ),
1242                "family {family} must reach the actual runtime admission boundary: {result:?}"
1243            );
1244        }
1245        assert_eq!(inner.entered.load(Ordering::Acquire), 0);
1246    }
1247
1248    #[tokio::test(start_paused = true)]
1249    async fn runtime_embedding_admitted_inference_is_not_reported_as_admission_timeout() {
1250        let inner = Arc::new(BlockingTestService::new());
1251        let _release = ReleaseWorkerOnDrop(Arc::clone(&inner));
1252        let runtime = queue_runtime(Arc::new(BlockingEmbeddingService::new(Arc::clone(&inner))));
1253        runtime
1254            .embedder("queue-test")
1255            .await
1256            .expect("resolve provider before the admission fixture");
1257        let mut call = Box::pin(khive_storage::scope_request_read_deadline(
1258            Duration::from_millis(100),
1259            runtime.embed_with_model("queue-test", "first"),
1260        ));
1261        assert!(poll_once(call.as_mut()).await.is_pending());
1262        drive_until_entered(call.as_mut(), &inner, 1).await;
1263        tokio::time::advance(Duration::from_millis(100)).await;
1264        let after_deadline = poll_once(call.as_mut()).await;
1265        inner.release();
1266        assert!(after_deadline.is_pending(),
1267            "already-running inference cannot be called a pre-admission refusal: {after_deadline:?}");
1268        assert_eq!(call.await.unwrap(), vec![1.0]);
1269    }
1270
1271    struct CustomWaitingService {
1272        release: Notify,
1273    }
1274
1275    #[async_trait]
1276    impl EmbeddingService for CustomWaitingService {
1277        async fn embed(
1278            &self,
1279            texts: &[String],
1280            _model: EmbeddingModel,
1281        ) -> lattice_embed::Result<Vec<Vec<f32>>> {
1282            self.release.notified().await;
1283            Ok(texts.iter().map(|_| vec![1.0]).collect())
1284        }
1285        fn supports_model(&self, _model: EmbeddingModel) -> bool {
1286            true
1287        }
1288        fn name(&self) -> &'static str {
1289            "custom-waiting"
1290        }
1291    }
1292
1293    #[tokio::test(start_paused = true)]
1294    async fn runtime_custom_inference_is_not_reported_as_builtin_admission_timeout() {
1295        let inner = Arc::new(CustomWaitingService {
1296            release: Notify::new(),
1297        });
1298        let runtime = queue_runtime(inner.clone());
1299        let mut call = Box::pin(khive_storage::scope_request_read_deadline(
1300            Duration::from_millis(100),
1301            runtime.embed_with_model("queue-test", "custom"),
1302        ));
1303        assert!(poll_once(call.as_mut()).await.is_pending());
1304        tokio::time::advance(Duration::from_millis(100)).await;
1305        let after_deadline = poll_once(call.as_mut()).await;
1306        inner.release.notify_one();
1307        assert!(
1308            after_deadline.is_pending(),
1309            "custom inference did not enter the built-in pre-admission wait: {after_deadline:?}"
1310        );
1311        assert_eq!(call.await.unwrap(), vec![1.0]);
1312    }
1313
1314    #[tokio::test]
1315    async fn blocking_adapter_rejects_oversized_job_before_enqueue() {
1316        let service = BlockingEmbeddingService::new(Arc::new(ConstVecService { dims: 1 }));
1317        let oversized =
1318            "x".repeat(lattice_embed::DEFAULT_MAX_BATCH_SIZE * lattice_embed::MAX_TEXT_BYTES + 1);
1319
1320        let error = service
1321            .embed(&[oversized], EmbeddingModel::default())
1322            .await
1323            .expect_err("an oversized embedding job must be rejected");
1324
1325        assert!(
1326            error.to_string().contains("embedding job input"),
1327            "oversized admission must use the embedding error path: {error}"
1328        );
1329        assert!(
1330            service.worker.get().is_none(),
1331            "oversized work must be rejected before the worker queue is initialized"
1332        );
1333    }
1334
1335    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
1336    async fn blocking_adapter_byte_budget_rejects_excess_and_queued_jobs_complete() {
1337        const ADMITTED_JOBS: usize = 4;
1338        let byte_budget = ADMITTED_JOBS * EMBEDDING_MAX_JOB_BYTES;
1339        let inner = Arc::new(BlockingTestService::new());
1340        let service = Arc::new(BlockingEmbeddingService::with_byte_budget(
1341            Arc::clone(&inner),
1342            byte_budget,
1343        ));
1344        let max_texts = Arc::new(vec![
1345            "x".repeat(lattice_embed::MAX_TEXT_BYTES);
1346            lattice_embed::DEFAULT_MAX_BATCH_SIZE
1347        ]);
1348        let mut admitted = Vec::with_capacity(ADMITTED_JOBS);
1349
1350        for _ in 0..ADMITTED_JOBS {
1351            let service = Arc::clone(&service);
1352            let texts = Arc::clone(&max_texts);
1353            admitted.push(tokio::spawn(async move {
1354                service.embed(&texts, EmbeddingModel::default()).await
1355            }));
1356        }
1357        tokio::time::timeout(Duration::from_secs(1), async {
1358            while service.in_flight_bytes.load(Ordering::Acquire) < byte_budget {
1359                tokio::task::yield_now().await;
1360            }
1361        })
1362        .await
1363        .expect("all jobs within the byte budget must be admitted");
1364
1365        let overflow = tokio::time::timeout(
1366            Duration::from_millis(100),
1367            service.embed(&max_texts, EmbeddingModel::default()),
1368        )
1369        .await
1370        .expect("a byte-budget overflow must fail without waiting")
1371        .expect_err("a byte-budget overflow must return an embedding error");
1372
1373        inner.release();
1374        for call in admitted {
1375            call.await
1376                .expect("admitted embedding task must not panic")
1377                .expect("admitted embedding job must complete");
1378        }
1379        assert!(
1380            overflow.to_string().contains("byte budget"),
1381            "byte saturation must use the embedding failure path: {overflow}"
1382        );
1383    }
1384
1385    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
1386    async fn blocking_adapter_releases_byte_budget_after_completion_and_skipped_job() {
1387        let byte_budget = "first".len() + "abandoned".len();
1388        let inner = Arc::new(BlockingTestService::new());
1389        let service = Arc::new(BlockingEmbeddingService::with_byte_budget(
1390            Arc::clone(&inner),
1391            byte_budget,
1392        ));
1393
1394        let first_service = Arc::clone(&service);
1395        let first = tokio::spawn(async move {
1396            first_service
1397                .embed(&["first".to_owned()], EmbeddingModel::default())
1398                .await
1399        });
1400        tokio::time::timeout(Duration::from_secs(1), async {
1401            while inner.entered.load(Ordering::Acquire) == 0 {
1402                tokio::task::yield_now().await;
1403            }
1404        })
1405        .await
1406        .expect("first embedding call must occupy the native worker");
1407
1408        let abandoned = tokio::time::timeout(
1409            Duration::from_millis(50),
1410            service.embed(&["abandoned".to_owned()], EmbeddingModel::default()),
1411        )
1412        .await;
1413        assert!(abandoned.is_err(), "queued embedding call must time out");
1414        assert_eq!(
1415            service.in_flight_bytes.load(Ordering::Acquire),
1416            byte_budget,
1417            "running and queued jobs must both consume the byte budget"
1418        );
1419
1420        inner.release();
1421        first
1422            .await
1423            .expect("first embedding task must not panic")
1424            .expect("first embedding call must succeed");
1425        tokio::time::timeout(Duration::from_secs(1), async {
1426            while service.in_flight_bytes.load(Ordering::Acquire) != 0 {
1427                tokio::task::yield_now().await;
1428            }
1429        })
1430        .await
1431        .expect("completed and skipped jobs must release their byte reservations");
1432
1433        let later = service
1434            .embed(&["later".to_owned()], EmbeddingModel::default())
1435            .await
1436            .expect("a later call must succeed after the byte budget is released");
1437        assert_eq!(later, vec![vec![1.0]]);
1438    }
1439
1440    #[test]
1441    fn register_and_get_provider_round_trip() {
1442        let mut reg = EmbedderRegistry::new();
1443        reg.register(ConstVecProvider::new("mock-384", 384));
1444
1445        assert!(reg.contains("mock-384"), "registered name must be present");
1446        let provider = reg.get_provider("mock-384").expect("provider must exist");
1447        assert_eq!(provider.name(), "mock-384");
1448        assert_eq!(provider.dimensions(), 384);
1449    }
1450
1451    #[test]
1452    fn duplicate_name_last_wins() {
1453        let mut reg = EmbedderRegistry::new();
1454        reg.register(ConstVecProvider::new("shared", 128));
1455        reg.register(ConstVecProvider::new("shared", 256));
1456
1457        let provider = reg.get_provider("shared").expect("provider must exist");
1458        assert_eq!(
1459            provider.dimensions(),
1460            256,
1461            "last registration must win; expected dims=256"
1462        );
1463    }
1464
1465    #[test]
1466    fn names_returns_all_registered() {
1467        let mut reg = EmbedderRegistry::new();
1468        reg.register(ConstVecProvider::new("model-a", 64));
1469        reg.register(ConstVecProvider::new("model-b", 128));
1470        reg.register(ConstVecProvider::new("model-c", 256));
1471
1472        let mut names = reg.names();
1473        names.sort();
1474        assert_eq!(names, vec!["model-a", "model-b", "model-c"]);
1475    }
1476
1477    #[tokio::test]
1478    async fn get_service_unknown_name_returns_error() {
1479        let reg = EmbedderRegistry::new();
1480        let result = reg.get_service("does-not-exist").await;
1481        let err = result.err().expect("expected Err for unknown name, got Ok");
1482        assert!(
1483            matches!(err, RuntimeError::UnknownModel(ref n) if n == "does-not-exist"),
1484            "expected UnknownModel, got {err:?}"
1485        );
1486    }
1487
1488    #[tokio::test]
1489    async fn get_service_calls_build_once() {
1490        let counter = Arc::new(AtomicUsize::new(0));
1491        let provider = ConstVecProvider {
1492            name: "cached-model".to_owned(),
1493            dims: 32,
1494            build_calls: Arc::clone(&counter),
1495        };
1496        let mut reg = EmbedderRegistry::new();
1497        reg.register(provider);
1498
1499        let _ = reg.get_service("cached-model").await.unwrap();
1500        let _ = reg.get_service("cached-model").await.unwrap();
1501        let _ = reg.get_service("cached-model").await.unwrap();
1502
1503        assert_eq!(
1504            counter.load(Ordering::SeqCst),
1505            1,
1506            "build must be called exactly once regardless of get_service call count"
1507        );
1508    }
1509
1510    struct SlowBuildProvider {
1511        name: String,
1512        dims: usize,
1513        build_calls: Arc<AtomicUsize>,
1514    }
1515
1516    #[async_trait]
1517    impl EmbedderProvider for SlowBuildProvider {
1518        fn name(&self) -> &str {
1519            &self.name
1520        }
1521
1522        fn dimensions(&self) -> usize {
1523            self.dims
1524        }
1525
1526        async fn build(&self) -> RuntimeResult<Arc<dyn EmbeddingService>> {
1527            self.build_calls.fetch_add(1, Ordering::SeqCst);
1528            tokio::time::sleep(Duration::from_millis(50)).await;
1529            Ok(Arc::new(ConstVecService { dims: self.dims }))
1530        }
1531    }
1532
1533    #[tokio::test(flavor = "multi_thread", worker_threads = 8)]
1534    async fn concurrent_cold_resolutions_single_flight_one_build() {
1535        const CALLERS: usize = 16;
1536        let counter = Arc::new(AtomicUsize::new(0));
1537        let mut reg = EmbedderRegistry::new();
1538        reg.register(SlowBuildProvider {
1539            name: "cold-model".to_owned(),
1540            dims: 8,
1541            build_calls: Arc::clone(&counter),
1542        });
1543        let reg = Arc::new(reg);
1544
1545        let mut callers = Vec::with_capacity(CALLERS);
1546        for _ in 0..CALLERS {
1547            let reg = Arc::clone(&reg);
1548            callers.push(tokio::spawn(
1549                async move { reg.get_service("cold-model").await },
1550            ));
1551        }
1552
1553        for caller in callers {
1554            let service = caller
1555                .await
1556                .expect("resolution task must not panic")
1557                .expect("every concurrent cold resolution must receive a working service");
1558            assert_eq!(service.name(), "const-vec-service");
1559        }
1560
1561        assert_eq!(
1562            counter.load(Ordering::SeqCst),
1563            1,
1564            "concurrent cold resolutions must share a single in-flight build()"
1565        );
1566    }
1567}