1use super::{
9 clone_client, embed_passage, embed_passages_controlled, get_embedder,
10 is_openrouter_initialized, shared_runtime, LlmBackendKind, OPENROUTER_CLIENT,
11};
12use crate::errors::AppError;
13use crate::extract::llm_embedding::LlmEmbedding;
14use parking_lot::Mutex;
15use std::path::Path;
16use std::sync::Arc;
17use std::sync::OnceLock;
18use tokio::sync::{mpsc, Semaphore};
19use tokio::task::JoinSet;
20use tokio_util::sync::CancellationToken;
21
22pub const CHUNK_EMBED_BATCH_SIZE: usize = 8;
26
27pub const ENTITY_EMBED_BATCH_SIZE: usize = 25;
31
32pub const EMBED_BATCH_CALIBRATION_DIM: usize = 64;
34
35pub(crate) fn adaptive_batch_for_dim(base: usize, dim: usize) -> usize {
43 let base = base.max(1);
44 (base * EMBED_BATCH_CALIBRATION_DIM / dim.max(1)).clamp(1, base)
45}
46
47pub fn chunk_embed_batch_size() -> usize {
49 let dim = crate::constants::embedding_dim();
50 let batch = adaptive_batch_for_dim(CHUNK_EMBED_BATCH_SIZE, dim);
51 tracing::debug!(
52 dim,
53 base = CHUNK_EMBED_BATCH_SIZE,
54 batch,
55 "adaptive chunk batch size (G44)"
56 );
57 batch
58}
59
60pub fn entity_embed_batch_size() -> usize {
62 let dim = crate::constants::embedding_dim();
63 let batch = adaptive_batch_for_dim(ENTITY_EMBED_BATCH_SIZE, dim);
64 tracing::debug!(
65 dim,
66 base = ENTITY_EMBED_BATCH_SIZE,
67 batch,
68 "adaptive entity batch size (G44)"
69 );
70 batch
71}
72
73pub fn embed_passages_controlled_local(
75 models_dir: &Path,
76 texts: &[&str],
77 token_counts: &[usize],
78) -> Result<Vec<Vec<f32>>, AppError> {
79 let embedder = get_embedder(models_dir)?;
80 embed_passages_controlled(embedder, texts, token_counts)
81}
82
83pub fn embed_passages_parallel_local(
86 models_dir: &Path,
87 texts: &[String],
88 parallelism: usize,
89 batch_size: usize,
90) -> Result<Vec<Vec<f32>>, AppError> {
91 let embedder = get_embedder(models_dir)?;
92 embed_texts_parallel(embedder, texts, parallelism, batch_size)
93}
94
95type EmbedChunkResult = (usize, Result<Vec<Vec<f32>>, AppError>);
99
100pub(crate) fn reassemble_ordered(mut parts: Vec<(usize, Vec<Vec<f32>>)>) -> Vec<Vec<f32>> {
105 parts.sort_by_key(|(idx, _)| *idx);
106 parts.into_iter().flat_map(|(_, v)| v).collect()
107}
108
109pub fn embed_passages_parallel_with_embedding_choice(
116 models_dir: &Path,
117 texts: &[String],
118 parallelism: usize,
119 batch_size: usize,
120 embedding_backend: crate::cli::EmbeddingBackendChoice,
121 llm_backend: crate::cli::LlmBackendChoice,
122) -> Result<Vec<Vec<f32>>, AppError> {
123 let chain = embedding_backend.to_chain(llm_backend);
124 if chain.first() == Some(&LlmBackendKind::OpenRouter) && is_openrouter_initialized() {
125 let client = OPENROUTER_CLIENT.get().ok_or_else(|| {
126 AppError::Embedding(
127 crate::i18n::validation::embedding_openrouter_client_not_initialised(),
128 )
129 })?;
130
131 let k = parallelism.clamp(1, 16);
136 if texts.len() <= 32 || k == 1 {
137 let refs: Vec<&str> = texts.iter().map(|s| s.as_str()).collect();
138 let vecs = match tokio::runtime::Handle::try_current() {
140 Ok(handle) => tokio::task::block_in_place(|| {
141 handle.block_on(client.embed_batch(&refs, client.default_input_type()))
142 })?,
143 Err(_) => shared_runtime()?
144 .block_on(client.embed_batch(&refs, client.default_input_type()))?,
145 };
146 return Ok(vecs);
147 }
148
149 let fan_out = async move {
158 let mut set: JoinSet<EmbedChunkResult> = JoinSet::new();
159 let mut parts: Vec<(usize, Vec<Vec<f32>>)> = Vec::new();
160
161 for (idx, chunk) in texts.chunks(32).enumerate() {
162 if set.len() >= k {
163 if let Some(joined) = set.join_next().await {
164 let (cidx, res) = joined.map_err(|e| {
165 AppError::Embedding(
166 crate::i18n::validation::embedding_task_join_error(e),
167 )
168 })?;
169 parts.push((cidx, res?));
170 }
171 }
172 let owned: Vec<String> = chunk.to_vec();
173 set.spawn(async move {
174 let refs: Vec<&str> = owned.iter().map(|s| s.as_str()).collect();
175 let r = client
179 .embed_batch(&refs, client.default_input_type())
180 .await
181 .map_err(AppError::from);
182 (idx, r)
183 });
184 }
185
186 while let Some(joined) = set.join_next().await {
187 let (cidx, res) = joined.map_err(|e| {
188 AppError::Embedding(crate::i18n::validation::embedding_task_join_error(e))
189 })?;
190 parts.push((cidx, res?));
191 }
192
193 Ok::<Vec<Vec<f32>>, AppError>(reassemble_ordered(parts))
194 };
195 let vecs = match tokio::runtime::Handle::try_current() {
196 Ok(handle) => tokio::task::block_in_place(|| handle.block_on(fan_out))?,
197 Err(_) => shared_runtime()?.block_on(fan_out)?,
198 };
199 Ok(vecs)
200 } else {
201 embed_passages_parallel_local(models_dir, texts, parallelism, batch_size)
202 }
203}
204
205type EntityEmbedCacheMap = std::collections::HashMap<u64, Arc<Vec<f32>>>;
217
218static ENTITY_EMBED_CACHE: OnceLock<parking_lot::Mutex<EntityEmbedCacheMap>> = OnceLock::new();
219
220pub(crate) fn entity_embed_cache() -> &'static parking_lot::Mutex<EntityEmbedCacheMap> {
221 ENTITY_EMBED_CACHE.get_or_init(|| parking_lot::Mutex::new(std::collections::HashMap::new()))
222}
223
224pub(crate) fn entity_cache_key(model: &str, text: &str) -> u64 {
225 let mut hasher = blake3::Hasher::new();
226 hasher.update(model.as_bytes());
227 hasher.update(b"\0");
228 hasher.update(text.as_bytes());
229 let h = hasher.finalize();
230 let bytes = h.as_bytes();
231 u64::from_le_bytes([
232 bytes[0], bytes[1], bytes[2], bytes[3], bytes[4], bytes[5], bytes[6], bytes[7],
233 ])
234}
235
236pub fn embed_entity_texts_cached(
246 models_dir: &Path,
247 texts: &[String],
248 parallelism: usize,
249 embedding_backend: crate::cli::EmbeddingBackendChoice,
250 llm_backend: crate::cli::LlmBackendChoice,
251) -> Result<(Vec<Vec<f32>>, EmbedCacheStats), AppError> {
252 if texts.is_empty() {
253 return Ok((Vec::new(), EmbedCacheStats::default()));
254 }
255 let chain = embedding_backend.to_chain(llm_backend);
259
260 if chain.as_slice() == [LlmBackendKind::None] {
266 let out: Vec<Vec<f32>> = texts.iter().map(|_| Vec::new()).collect();
267 return Ok((
268 out,
269 EmbedCacheStats {
270 requested: texts.len(),
271 hits: 0,
272 misses: texts.len(),
273 },
274 ));
275 }
276
277 let routed_openrouter =
283 chain.first() == Some(&LlmBackendKind::OpenRouter) && is_openrouter_initialized();
284 let model = if routed_openrouter {
285 format!("openrouter:{}", crate::constants::embedding_dim())
286 } else {
287 get_embedder(models_dir)?.lock().model_label()
288 };
289 let cache = entity_embed_cache();
290 let mut hits: Vec<Option<Arc<Vec<f32>>>> = vec![None; texts.len()];
291 let mut miss_indices: Vec<usize> = Vec::with_capacity(texts.len());
292 {
293 let guard = cache.lock();
294 for (i, text) in texts.iter().enumerate() {
295 let key = entity_cache_key(&model, text);
296 if let Some(v) = guard.get(&key) {
297 hits[i] = Some(Arc::clone(v));
298 } else {
299 miss_indices.push(i);
300 }
301 }
302 }
303 let miss_count = miss_indices.len();
304 if miss_count > 0 {
305 let miss_texts: Vec<String> = miss_indices.iter().map(|&i| texts[i].clone()).collect();
306 let miss_vecs = embed_passages_parallel_with_embedding_choice(
310 models_dir,
311 &miss_texts,
312 parallelism,
313 entity_embed_batch_size(),
314 embedding_backend,
315 llm_backend,
316 )?;
317 let mut guard = cache.lock();
318 for (slot, &orig_idx) in miss_indices.iter().enumerate() {
319 let vec = Arc::new(miss_vecs[slot].clone());
320 let key = entity_cache_key(&model, &texts[orig_idx]);
321 guard.insert(key, Arc::clone(&vec));
322 hits[orig_idx] = Some(vec);
323 }
324 }
325 let mut out = Vec::with_capacity(texts.len());
326 for hit in hits.into_iter() {
327 let v = hit.ok_or_else(|| {
328 AppError::Embedding(crate::i18n::validation::embedding_entity_cache_null())
329 })?;
330 out.push((*v).clone());
331 }
332 Ok((
333 out,
334 EmbedCacheStats {
335 requested: texts.len(),
336 hits: texts.len() - miss_count,
337 misses: miss_count,
338 },
339 ))
340}
341
342#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, serde::Serialize)]
344pub struct EmbedCacheStats {
345 pub requested: usize,
347 pub hits: usize,
349 pub misses: usize,
351}
352
353impl EmbedCacheStats {
354 pub fn hit_rate(&self) -> f64 {
356 if self.requested == 0 {
357 0.0
358 } else {
359 self.hits as f64 / self.requested as f64
360 }
361 }
362}
363
364pub fn embed_texts_parallel(
377 embedder: &Mutex<LlmEmbedding>,
378 texts: &[String],
379 parallelism: usize,
380 batch_size: usize,
381) -> Result<Vec<Vec<f32>>, AppError> {
382 let mut slots: Vec<Option<Vec<f32>>> = vec![None; texts.len()];
383 embed_texts_parallel_with(embedder, texts, parallelism, batch_size, |idx, v| {
384 slots[idx] = Some(v.to_vec());
385 Ok(())
386 })?;
387 let mut out = Vec::with_capacity(slots.len());
388 for (idx, slot) in slots.into_iter().enumerate() {
389 out.push(slot.ok_or_else(|| {
390 AppError::Embedding(crate::i18n::validation::embedding_fanout_lost_index(idx))
391 })?);
392 }
393 Ok(out)
394}
395
396pub fn embed_texts_parallel_with(
400 embedder: &Mutex<LlmEmbedding>,
401 texts: &[String],
402 parallelism: usize,
403 batch_size: usize,
404 mut on_result: impl FnMut(usize, &[f32]) -> Result<(), AppError>,
405) -> Result<(), AppError> {
406 if texts.is_empty() {
407 return Ok(());
408 }
409 let dim = crate::constants::embedding_dim();
410 if texts.len() == 1 {
411 let v = embed_passage(embedder, &texts[0])?;
412 return on_result(0, &v);
413 }
414
415 let client = clone_client(embedder);
416 let permits = effective_permits(parallelism);
417 let batches = build_batches(texts, batch_size.max(1));
418 let token = crate::cancel_token().clone();
419
420 let work = move |batch: Vec<(usize, String)>| {
421 let client = client.clone();
422 async move {
423 client
424 .embed_batch_async(crate::constants::PASSAGE_PREFIX, &batch)
425 .await
426 }
427 };
428
429 let fan_out = run_bounded(batches, permits, dim, token, work, &mut on_result);
430 match tokio::runtime::Handle::try_current() {
431 Ok(handle) => tokio::task::block_in_place(|| handle.block_on(fan_out)),
432 Err(_) => shared_runtime()?.block_on(fan_out),
433 }
434}
435
436pub(crate) fn build_batches(texts: &[String], batch_size: usize) -> Vec<Vec<(usize, String)>> {
438 texts
439 .iter()
440 .cloned()
441 .enumerate()
442 .collect::<Vec<_>>()
443 .chunks(batch_size)
444 .map(|c| c.to_vec())
445 .collect()
446}
447
448pub fn effective_permits(requested: usize) -> usize {
453 let cpus = std::thread::available_parallelism()
454 .map(|n| n.get())
455 .unwrap_or(4);
456 let by_ram = ((crate::memory_guard::available_memory_mb() / 2)
457 / crate::constants::LLM_WORKER_RSS_MB)
458 .max(1) as usize;
459 requested.clamp(1, 32).min(cpus).min(by_ram).max(1)
460}
461
462pub(crate) async fn run_bounded<F, Fut>(
472 batches: Vec<Vec<(usize, String)>>,
473 permits: usize,
474 dim: usize,
475 token: CancellationToken,
476 work: F,
477 on_result: &mut impl FnMut(usize, &[f32]) -> Result<(), AppError>,
478) -> Result<(), AppError>
479where
480 F: Fn(Vec<(usize, String)>) -> Fut + Clone + Send + 'static,
481 Fut: std::future::Future<Output = Result<Vec<(usize, Vec<f32>)>, AppError>> + Send,
482{
483 let total_batches = batches.len();
484 let semaphore = Arc::new(Semaphore::new(permits));
485 let (tx, mut rx) = mpsc::channel::<Result<Vec<(usize, Vec<f32>)>, AppError>>(permits * 2);
488 let mut set: JoinSet<()> = JoinSet::new();
489
490 for (batch_idx, batch) in batches.into_iter().enumerate() {
491 let sem = Arc::clone(&semaphore);
492 let token = token.clone();
493 let tx = tx.clone();
494 let work = work.clone();
495 set.spawn(async move {
496 let wait_start = std::time::Instant::now();
497 let Ok(_permit) = sem.acquire_owned().await else {
500 let _ = tx
501 .send(Err(AppError::Embedding(
502 crate::i18n::validation::embedding_semaphore_closed(),
503 )))
504 .await;
505 return;
506 };
507 let permit_wait_ms = wait_start.elapsed().as_millis() as u64;
508 let work_start = std::time::Instant::now();
509 let outcome = if crate::should_obey_shutdown() {
515 tokio::select! {
516 res = work(batch) => res,
517 _ = token.cancelled() => Err(AppError::Embedding(
518 crate::i18n::validation::embedding_cancelled_by_shutdown(),
519 )),
520 }
521 } else {
522 work(batch).await
523 };
524 tracing::debug!(
526 target: "embedding",
527 batch_idx,
528 permit_wait_ms,
529 work_ms = work_start.elapsed().as_millis() as u64,
530 ok = outcome.is_ok(),
531 "embedding batch finished"
532 );
533 let _ = tx.send(outcome).await;
534 });
535 }
536 drop(tx);
537
538 let mut completed = 0usize;
539 let mut failed = 0usize;
540 let mut cancelled = 0usize;
541 let mut first_error: Option<AppError> = None;
542
543 while let Some(message) = rx.recv().await {
544 match message {
545 Ok(items) => {
546 completed += 1;
547 if first_error.is_none() {
548 for (idx, v) in items {
549 if v.len() != dim {
550 first_error = Some(AppError::Embedding(
551 crate::i18n::validation::embedding_llm_item_dims(
552 v.len(),
553 idx,
554 dim,
555 ),
556 ));
557 break;
558 }
559 if let Err(e) = on_result(idx, &v) {
560 first_error = Some(e);
561 break;
562 }
563 }
564 if first_error.is_some() {
565 set.shutdown().await;
568 }
569 }
570 }
571 Err(e) => {
572 if matches!(&e, AppError::Embedding(msg) if msg.contains("cancelled")) {
573 cancelled += 1;
574 } else {
575 failed += 1;
576 }
577 if first_error.is_none() {
578 first_error = Some(e);
579 set.shutdown().await;
580 }
581 }
582 }
583 }
584
585 while let Some(join_result) = set.join_next().await {
588 if let Err(join_err) = join_result {
589 if join_err.is_panic() {
590 failed += 1;
591 if first_error.is_none() {
592 first_error = Some(AppError::Embedding(
593 crate::i18n::validation::embedding_task_panicked(join_err),
594 ));
595 }
596 } else {
597 cancelled += 1;
598 }
599 }
600 }
601
602 tracing::debug!(
612 target: "embedding",
613 total_batches,
614 completed,
615 failed,
616 cancelled,
617 "embedding fan-out finished"
618 );
619
620 match first_error {
621 Some(e) => Err(e),
622 None => Ok(()),
623 }
624}