use super::{
clone_client, embed_passage, embed_passages_controlled, get_embedder,
is_openrouter_initialized, shared_runtime, LlmBackendKind, OPENROUTER_CLIENT,
};
use crate::errors::AppError;
use crate::extract::llm_embedding::LlmEmbedding;
use parking_lot::Mutex;
use std::path::Path;
use std::sync::Arc;
use std::sync::OnceLock;
use tokio::sync::{mpsc, Semaphore};
use tokio::task::JoinSet;
use tokio_util::sync::CancellationToken;
pub const CHUNK_EMBED_BATCH_SIZE: usize = 8;
pub const ENTITY_EMBED_BATCH_SIZE: usize = 25;
pub const EMBED_BATCH_CALIBRATION_DIM: usize = 64;
pub(crate) fn adaptive_batch_for_dim(base: usize, dim: usize) -> usize {
let base = base.max(1);
(base * EMBED_BATCH_CALIBRATION_DIM / dim.max(1)).clamp(1, base)
}
pub fn chunk_embed_batch_size() -> usize {
let dim = crate::constants::embedding_dim();
let batch = adaptive_batch_for_dim(CHUNK_EMBED_BATCH_SIZE, dim);
tracing::debug!(
dim,
base = CHUNK_EMBED_BATCH_SIZE,
batch,
"adaptive chunk batch size (G44)"
);
batch
}
pub fn entity_embed_batch_size() -> usize {
let dim = crate::constants::embedding_dim();
let batch = adaptive_batch_for_dim(ENTITY_EMBED_BATCH_SIZE, dim);
tracing::debug!(
dim,
base = ENTITY_EMBED_BATCH_SIZE,
batch,
"adaptive entity batch size (G44)"
);
batch
}
pub fn embed_passages_controlled_local(
models_dir: &Path,
texts: &[&str],
token_counts: &[usize],
) -> Result<Vec<Vec<f32>>, AppError> {
let embedder = get_embedder(models_dir)?;
embed_passages_controlled(embedder, texts, token_counts)
}
pub fn embed_passages_parallel_local(
models_dir: &Path,
texts: &[String],
parallelism: usize,
batch_size: usize,
) -> Result<Vec<Vec<f32>>, AppError> {
let embedder = get_embedder(models_dir)?;
embed_texts_parallel(embedder, texts, parallelism, batch_size)
}
type EmbedChunkResult = (usize, Result<Vec<Vec<f32>>, AppError>);
pub(crate) fn reassemble_ordered(mut parts: Vec<(usize, Vec<Vec<f32>>)>) -> Vec<Vec<f32>> {
parts.sort_by_key(|(idx, _)| *idx);
parts.into_iter().flat_map(|(_, v)| v).collect()
}
pub fn embed_passages_parallel_with_embedding_choice(
models_dir: &Path,
texts: &[String],
parallelism: usize,
batch_size: usize,
embedding_backend: crate::cli::EmbeddingBackendChoice,
llm_backend: crate::cli::LlmBackendChoice,
) -> Result<Vec<Vec<f32>>, AppError> {
let chain = embedding_backend.to_chain(llm_backend);
if chain.first() == Some(&LlmBackendKind::OpenRouter) && is_openrouter_initialized() {
let client = OPENROUTER_CLIENT.get().ok_or_else(|| {
AppError::Embedding(
crate::i18n::validation::embedding_openrouter_client_not_initialised(),
)
})?;
let k = parallelism.clamp(1, 16);
if texts.len() <= 32 || k == 1 {
let refs: Vec<&str> = texts.iter().map(|s| s.as_str()).collect();
let vecs = match tokio::runtime::Handle::try_current() {
Ok(handle) => tokio::task::block_in_place(|| {
handle.block_on(client.embed_batch(&refs, client.default_input_type()))
})?,
Err(_) => shared_runtime()?
.block_on(client.embed_batch(&refs, client.default_input_type()))?,
};
return Ok(vecs);
}
let fan_out = async move {
let mut set: JoinSet<EmbedChunkResult> = JoinSet::new();
let mut parts: Vec<(usize, Vec<Vec<f32>>)> = Vec::new();
for (idx, chunk) in texts.chunks(32).enumerate() {
if set.len() >= k {
if let Some(joined) = set.join_next().await {
let (cidx, res) = joined.map_err(|e| {
AppError::Embedding(
crate::i18n::validation::embedding_task_join_error(e),
)
})?;
parts.push((cidx, res?));
}
}
let owned: Vec<String> = chunk.to_vec();
set.spawn(async move {
let refs: Vec<&str> = owned.iter().map(|s| s.as_str()).collect();
let r = client
.embed_batch(&refs, client.default_input_type())
.await
.map_err(AppError::from);
(idx, r)
});
}
while let Some(joined) = set.join_next().await {
let (cidx, res) = joined.map_err(|e| {
AppError::Embedding(crate::i18n::validation::embedding_task_join_error(e))
})?;
parts.push((cidx, res?));
}
Ok::<Vec<Vec<f32>>, AppError>(reassemble_ordered(parts))
};
let vecs = match tokio::runtime::Handle::try_current() {
Ok(handle) => tokio::task::block_in_place(|| handle.block_on(fan_out))?,
Err(_) => shared_runtime()?.block_on(fan_out)?,
};
Ok(vecs)
} else {
embed_passages_parallel_local(models_dir, texts, parallelism, batch_size)
}
}
type EntityEmbedCacheMap = std::collections::HashMap<u64, Arc<Vec<f32>>>;
static ENTITY_EMBED_CACHE: OnceLock<parking_lot::Mutex<EntityEmbedCacheMap>> = OnceLock::new();
pub(crate) fn entity_embed_cache() -> &'static parking_lot::Mutex<EntityEmbedCacheMap> {
ENTITY_EMBED_CACHE.get_or_init(|| parking_lot::Mutex::new(std::collections::HashMap::new()))
}
pub(crate) fn entity_cache_key(model: &str, text: &str) -> u64 {
let mut hasher = blake3::Hasher::new();
hasher.update(model.as_bytes());
hasher.update(b"\0");
hasher.update(text.as_bytes());
let h = hasher.finalize();
let bytes = h.as_bytes();
u64::from_le_bytes([
bytes[0], bytes[1], bytes[2], bytes[3], bytes[4], bytes[5], bytes[6], bytes[7],
])
}
pub fn embed_entity_texts_cached(
models_dir: &Path,
texts: &[String],
parallelism: usize,
embedding_backend: crate::cli::EmbeddingBackendChoice,
llm_backend: crate::cli::LlmBackendChoice,
) -> Result<(Vec<Vec<f32>>, EmbedCacheStats), AppError> {
if texts.is_empty() {
return Ok((Vec::new(), EmbedCacheStats::default()));
}
let chain = embedding_backend.to_chain(llm_backend);
if chain.as_slice() == [LlmBackendKind::None] {
let out: Vec<Vec<f32>> = texts.iter().map(|_| Vec::new()).collect();
return Ok((
out,
EmbedCacheStats {
requested: texts.len(),
hits: 0,
misses: texts.len(),
},
));
}
let routed_openrouter =
chain.first() == Some(&LlmBackendKind::OpenRouter) && is_openrouter_initialized();
let model = if routed_openrouter {
format!("openrouter:{}", crate::constants::embedding_dim())
} else {
get_embedder(models_dir)?.lock().model_label()
};
let cache = entity_embed_cache();
let mut hits: Vec<Option<Arc<Vec<f32>>>> = vec![None; texts.len()];
let mut miss_indices: Vec<usize> = Vec::with_capacity(texts.len());
{
let guard = cache.lock();
for (i, text) in texts.iter().enumerate() {
let key = entity_cache_key(&model, text);
if let Some(v) = guard.get(&key) {
hits[i] = Some(Arc::clone(v));
} else {
miss_indices.push(i);
}
}
}
let miss_count = miss_indices.len();
if miss_count > 0 {
let miss_texts: Vec<String> = miss_indices.iter().map(|&i| texts[i].clone()).collect();
let miss_vecs = embed_passages_parallel_with_embedding_choice(
models_dir,
&miss_texts,
parallelism,
entity_embed_batch_size(),
embedding_backend,
llm_backend,
)?;
let mut guard = cache.lock();
for (slot, &orig_idx) in miss_indices.iter().enumerate() {
let vec = Arc::new(miss_vecs[slot].clone());
let key = entity_cache_key(&model, &texts[orig_idx]);
guard.insert(key, Arc::clone(&vec));
hits[orig_idx] = Some(vec);
}
}
let mut out = Vec::with_capacity(texts.len());
for hit in hits.into_iter() {
let v = hit.ok_or_else(|| {
AppError::Embedding(crate::i18n::validation::embedding_entity_cache_null())
})?;
out.push((*v).clone());
}
Ok((
out,
EmbedCacheStats {
requested: texts.len(),
hits: texts.len() - miss_count,
misses: miss_count,
},
))
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, serde::Serialize)]
pub struct EmbedCacheStats {
pub requested: usize,
pub hits: usize,
pub misses: usize,
}
impl EmbedCacheStats {
pub fn hit_rate(&self) -> f64 {
if self.requested == 0 {
0.0
} else {
self.hits as f64 / self.requested as f64
}
}
}
pub fn embed_texts_parallel(
embedder: &Mutex<LlmEmbedding>,
texts: &[String],
parallelism: usize,
batch_size: usize,
) -> Result<Vec<Vec<f32>>, AppError> {
let mut slots: Vec<Option<Vec<f32>>> = vec![None; texts.len()];
embed_texts_parallel_with(embedder, texts, parallelism, batch_size, |idx, v| {
slots[idx] = Some(v.to_vec());
Ok(())
})?;
let mut out = Vec::with_capacity(slots.len());
for (idx, slot) in slots.into_iter().enumerate() {
out.push(slot.ok_or_else(|| {
AppError::Embedding(crate::i18n::validation::embedding_fanout_lost_index(idx))
})?);
}
Ok(out)
}
pub fn embed_texts_parallel_with(
embedder: &Mutex<LlmEmbedding>,
texts: &[String],
parallelism: usize,
batch_size: usize,
mut on_result: impl FnMut(usize, &[f32]) -> Result<(), AppError>,
) -> Result<(), AppError> {
if texts.is_empty() {
return Ok(());
}
let dim = crate::constants::embedding_dim();
if texts.len() == 1 {
let v = embed_passage(embedder, &texts[0])?;
return on_result(0, &v);
}
let client = clone_client(embedder);
let permits = effective_permits(parallelism);
let batches = build_batches(texts, batch_size.max(1));
let token = crate::cancel_token().clone();
let work = move |batch: Vec<(usize, String)>| {
let client = client.clone();
async move {
client
.embed_batch_async(crate::constants::PASSAGE_PREFIX, &batch)
.await
}
};
let fan_out = run_bounded(batches, permits, dim, token, work, &mut on_result);
match tokio::runtime::Handle::try_current() {
Ok(handle) => tokio::task::block_in_place(|| handle.block_on(fan_out)),
Err(_) => shared_runtime()?.block_on(fan_out),
}
}
pub(crate) fn build_batches(texts: &[String], batch_size: usize) -> Vec<Vec<(usize, String)>> {
texts
.iter()
.cloned()
.enumerate()
.collect::<Vec<_>>()
.chunks(batch_size)
.map(|c| c.to_vec())
.collect()
}
pub fn effective_permits(requested: usize) -> usize {
let cpus = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(4);
let by_ram = ((crate::memory_guard::available_memory_mb() / 2)
/ crate::constants::LLM_WORKER_RSS_MB)
.max(1) as usize;
requested.clamp(1, 32).min(cpus).min(by_ram).max(1)
}
pub(crate) async fn run_bounded<F, Fut>(
batches: Vec<Vec<(usize, String)>>,
permits: usize,
dim: usize,
token: CancellationToken,
work: F,
on_result: &mut impl FnMut(usize, &[f32]) -> Result<(), AppError>,
) -> Result<(), AppError>
where
F: Fn(Vec<(usize, String)>) -> Fut + Clone + Send + 'static,
Fut: std::future::Future<Output = Result<Vec<(usize, Vec<f32>)>, AppError>> + Send,
{
let total_batches = batches.len();
let semaphore = Arc::new(Semaphore::new(permits));
let (tx, mut rx) = mpsc::channel::<Result<Vec<(usize, Vec<f32>)>, AppError>>(permits * 2);
let mut set: JoinSet<()> = JoinSet::new();
for (batch_idx, batch) in batches.into_iter().enumerate() {
let sem = Arc::clone(&semaphore);
let token = token.clone();
let tx = tx.clone();
let work = work.clone();
set.spawn(async move {
let wait_start = std::time::Instant::now();
let Ok(_permit) = sem.acquire_owned().await else {
let _ = tx
.send(Err(AppError::Embedding(
crate::i18n::validation::embedding_semaphore_closed(),
)))
.await;
return;
};
let permit_wait_ms = wait_start.elapsed().as_millis() as u64;
let work_start = std::time::Instant::now();
let outcome = if crate::should_obey_shutdown() {
tokio::select! {
res = work(batch) => res,
_ = token.cancelled() => Err(AppError::Embedding(
crate::i18n::validation::embedding_cancelled_by_shutdown(),
)),
}
} else {
work(batch).await
};
tracing::debug!(
target: "embedding",
batch_idx,
permit_wait_ms,
work_ms = work_start.elapsed().as_millis() as u64,
ok = outcome.is_ok(),
"embedding batch finished"
);
let _ = tx.send(outcome).await;
});
}
drop(tx);
let mut completed = 0usize;
let mut failed = 0usize;
let mut cancelled = 0usize;
let mut first_error: Option<AppError> = None;
while let Some(message) = rx.recv().await {
match message {
Ok(items) => {
completed += 1;
if first_error.is_none() {
for (idx, v) in items {
if v.len() != dim {
first_error = Some(AppError::Embedding(
crate::i18n::validation::embedding_llm_item_dims(
v.len(),
idx,
dim,
),
));
break;
}
if let Err(e) = on_result(idx, &v) {
first_error = Some(e);
break;
}
}
if first_error.is_some() {
set.shutdown().await;
}
}
}
Err(e) => {
if matches!(&e, AppError::Embedding(msg) if msg.contains("cancelled")) {
cancelled += 1;
} else {
failed += 1;
}
if first_error.is_none() {
first_error = Some(e);
set.shutdown().await;
}
}
}
}
while let Some(join_result) = set.join_next().await {
if let Err(join_err) = join_result {
if join_err.is_panic() {
failed += 1;
if first_error.is_none() {
first_error = Some(AppError::Embedding(
crate::i18n::validation::embedding_task_panicked(join_err),
));
}
} else {
cancelled += 1;
}
}
}
tracing::debug!(
target: "embedding",
total_batches,
completed,
failed,
cancelled,
"embedding fan-out finished"
);
match first_error {
Some(e) => Err(e),
None => Ok(()),
}
}