use super::fan_out::{chunk_ranges, fan_out_chunk, reassemble_ordered};
use crate::embedder::{
is_openrouter_initialized, shared_runtime, LlmBackendKind, OPENROUTER_CLIENT,
};
use crate::errors::AppError;
use std::path::Path;
use std::sync::Arc;
use tokio::task::JoinSet;
type EmbedChunkResult = (usize, Result<Vec<Vec<f32>>, AppError>);
#[deprecated(
since = "1.2.3",
note = "clones the whole corpus; use `embed_passages_parallel_shared`, which takes an \
`Arc<[String]>` and is the real implementation this delegates to"
)]
pub fn embed_passages_parallel_with_embedding_choice(
models_dir: &Path,
texts: &[String],
parallelism: usize,
local_batch_size: usize,
embedding_backend: crate::cli::EmbeddingBackendChoice,
llm_backend: crate::cli::LlmBackendChoice,
) -> Result<Vec<Vec<f32>>, AppError> {
embed_passages_parallel_shared(
models_dir,
Arc::from(texts.to_vec()),
parallelism,
local_batch_size,
embedding_backend,
llm_backend,
)
}
pub(crate) fn embed_passages_parallel_shared(
_models_dir: &Path,
texts: Arc<[String]>,
parallelism: usize,
_local_batch_size: usize,
embedding_backend: crate::cli::EmbeddingBackendChoice,
llm_backend: crate::cli::LlmBackendChoice,
) -> Result<Vec<Vec<f32>>, AppError> {
let texts: &Arc<[String]> = &texts;
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)
.min(crate::constants::joint_parallelism_ceiling())
.max(1);
let chunk = fan_out_chunk();
if texts.len() <= chunk || 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, range) in chunk_ranges(texts.len(), chunk).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 shared = Arc::clone(texts);
set.spawn(async move {
let refs: Vec<&str> =
shared[range.clone()].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 {
Err(AppError::Embedding(
crate::i18n::validation::embedding_openrouter_client_not_initialised(),
))
}
}