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