use super::QueueWorker;
use crate::error::DbError;
use crate::server::llm_client::LLMClient;
use crate::storage::collection::vector::{pending_embed_count, release_pending_embed};
use crate::storage::index::{extract_field_value, VectorIndexConfig};
use crate::storage::Collection;
const EMBED_BATCH: usize = 128;
const ERROR_BACKOFF_SECS: u64 = 60;
fn now_secs() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0)
}
impl QueueWorker {
pub(crate) async fn check_embeddings(&self) {
let claimed = pending_embed_count();
if claimed == 0 {
return;
}
let _lock = match self.claiming_lock.try_lock() {
Ok(l) => l,
Err(_) => return,
};
let grouped = self.storage.collections_grouped();
let mut saw_pending = false;
let now = now_secs();
if let Ok(mut backoff) = self.embed_backoff.lock() {
backoff.retain(|_, until| *until > now);
}
for (db_name, coll_names) in grouped {
let db = match self.storage.get_database(&db_name) {
Ok(d) => d,
Err(_) => continue,
};
for coll_name in coll_names {
let coll = match db.system_collection(&coll_name) {
Ok(c) => c,
Err(_) => continue,
};
let configs = coll.get_all_vector_index_configs();
for config in configs {
if config.embedding_source.is_none() {
continue;
}
let backoff_key = format!("{}\u{0}{}\u{0}{}", db_name, coll_name, config.name);
let backing_off = self
.embed_backoff
.lock()
.map(|b| b.get(&backoff_key).is_some_and(|until| *until > now))
.unwrap_or(false);
if backing_off {
saw_pending = true;
continue;
}
let pending = coll.list_embed_pending(&config.name, EMBED_BATCH);
if pending.is_empty() {
continue;
}
saw_pending = true;
if let Err(e) = self
.embed_pending_batch(&db_name, &coll, &config, &pending)
.await
{
tracing::warn!(
"Auto-embed worker: {}/{} index '{}' failed: {} (backing off {}s)",
db_name,
coll_name,
config.name,
e,
ERROR_BACKOFF_SECS
);
if let Ok(mut backoff) = self.embed_backoff.lock() {
backoff.insert(backoff_key, now_secs() + ERROR_BACKOFF_SECS);
}
}
}
}
}
if !saw_pending {
release_pending_embed(claimed);
}
}
async fn embed_pending_batch(
&self,
db_name: &str,
coll: &Collection,
config: &VectorIndexConfig,
doc_keys: &[String],
) -> Result<(), DbError> {
let source_field = config.embedding_source.as_deref().unwrap_or_default();
let mut keys: Vec<String> = Vec::new();
let mut texts: Vec<String> = Vec::new();
for dk in doc_keys {
let doc = match coll.get(dk) {
Ok(d) => d,
Err(_) => {
coll.clear_embed_pending(&config.name, dk);
continue;
}
};
let value = doc.to_value();
match extract_field_value(&value, source_field).as_str() {
Some(t) if !t.trim().is_empty() => {
keys.push(dk.clone());
texts.push(t.to_string());
}
_ => coll.clear_embed_pending(&config.name, dk),
}
}
if keys.is_empty() {
return Ok(());
}
let provider = config
.embedding_provider
.clone()
.unwrap_or_else(|| "openai".to_string());
let client = LLMClient::from_storage(
&self.storage,
db_name,
Some(&provider),
config.embedding_model.clone(),
)?;
let text_refs: Vec<&str> = texts.iter().map(|s| s.as_str()).collect();
let embeddings = client.embed_batch(&text_refs).await?;
for (dk, emb) in keys.iter().zip(embeddings) {
if emb.len() != config.dimension {
tracing::warn!(
"Auto-embed worker: dim mismatch for '{}' index '{}' (got {}, expected {}); dropping marker",
dk,
config.name,
emb.len(),
config.dimension
);
coll.clear_embed_pending(&config.name, dk);
continue;
}
let doc = match coll.get(dk) {
Ok(d) => d,
Err(_) => {
coll.clear_embed_pending(&config.name, dk);
continue;
}
};
let mut value = doc.to_value();
if let Some(obj) = value.as_object_mut() {
obj.insert(config.field.clone(), serde_json::json!(emb));
if let Err(e) = coll.update(dk, value) {
tracing::warn!(
"Auto-embed worker: failed to persist vector for '{}': {}",
dk,
e
);
}
}
}
Ok(())
}
}