1use 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
41pub(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 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 _ => {}
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;
106const 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
156pub(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 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 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 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#[async_trait]
393pub trait EmbedderProvider: Send + Sync {
394 fn name(&self) -> &str;
401
402 fn dimensions(&self) -> usize;
407
408 async fn build(&self) -> RuntimeResult<Arc<dyn EmbeddingService>>;
414}
415
416pub(crate) struct EmbedderEntry {
419 provider: Arc<dyn EmbedderProvider>,
420 cell: Arc<OnceCell<Arc<dyn EmbeddingService>>>,
421 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#[derive(Clone, Default)]
443pub struct EmbedderRegistry {
444 entries: HashMap<String, EmbedderEntry>,
445}
446
447impl EmbedderRegistry {
448 pub fn new() -> Self {
450 Self {
451 entries: HashMap::new(),
452 }
453 }
454
455 pub fn register<P: EmbedderProvider + 'static>(&mut self, provider: P) {
464 self.insert(provider, false);
465 }
466
467 pub(crate) fn register_builtin(&mut self, provider: LatticeEmbedderProvider) {
470 self.insert(provider, true);
471 }
472
473 #[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 pub fn get_provider(&self, name: &str) -> Option<&dyn EmbedderProvider> {
503 self.entries.get(name).map(|e| e.provider.as_ref())
504 }
505
506 pub fn contains(&self, name: &str) -> bool {
508 self.entries.contains_key(name)
509 }
510
511 pub fn names(&self) -> Vec<String> {
513 self.entries.keys().cloned().collect()
514 }
515
516 pub(crate) fn get_entry(&self, name: &str) -> Option<EmbedderEntry> {
522 self.entries.get(name).cloned()
523 }
524
525 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#[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 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
632pub struct LatticeEmbedderProvider {
642 model: EmbeddingModel,
643 name: String,
645}
646
647impl LatticeEmbedderProvider {
648 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
672fn 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#[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 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(®);
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}