use std::path::Path;
use std::time::Instant;
use super::args::{EnrichArgs, EnrichOperation};
use super::events::*;
use super::extraction::{
call_body_enrich, call_body_extract, call_deep_research_synth, call_description_enrich,
call_domain_classify, call_entity_connect, call_entity_description, call_entity_type_validate,
call_graph_audit, call_memory_bindings, call_reembed, call_relation_reclassify,
call_weight_calibrate, take_last_openrouter_failure, EnrichItemResult,
};
use super::queue::{
dequeue_next_pending, heartbeat, item_type_for, open_queue_db, record_item_failure,
record_item_failure_typed, requeue_wrong_op, skip_wrong_type, validate_claim, ClaimCheck,
DequeueOutcome,
};
use super::prompts;
use super::DEFAULT_RATE_LIMIT_WAIT;
use crate::errors::AppError;
use crate::output::emit_json_line as emit_json;
use crate::paths::AppPaths;
use crate::storage::connection::open_rw;
#[derive(Default)]
pub(crate) struct DrainCounters {
pub completed: usize,
pub failed: usize,
pub skipped: usize,
pub cost_total: f64,
pub oauth_detected: bool,
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn drain_parallel(
args: &EnrichArgs,
paths: &AppPaths,
queue_path: &Path,
namespace: &str,
provider_binary: Option<&Path>,
provider_model: Option<&str>,
provider_timeout: u64,
op_label: &str,
backoff_clause: &str,
parallelism: usize,
total: usize,
llm_backend: crate::cli::LlmBackendChoice,
embedding_backend: crate::cli::EmbeddingBackendChoice,
counters: &mut DrainCounters,
) -> Result<(), AppError> {
let stdout_mu = parking_lot::Mutex::new(());
let budget = args.max_cost_usd;
let operation = args.operation().clone();
let mode = args.mode().clone();
let min_oc = args.min_output_chars;
let max_oc = args.max_output_chars;
let prompt_tpl = args.prompt_template.as_deref().map(|p| p.to_path_buf());
struct WorkerResult {
completed: usize,
failed: usize,
skipped: usize,
cost: f64,
oauth: bool,
db_busy: bool,
}
let results: Vec<WorkerResult> = std::thread::scope(|s| {
let handles: Vec<_> = (0..parallelism)
.map(|worker_id| {
let stdout_mu = &stdout_mu;
let paths = &paths;
let queue_path = &queue_path;
let namespace = &namespace;
let operation = &operation;
let mode = &mode;
let op_label = &op_label;
let expected_item_type = item_type_for(operation);
let prompt_tpl = prompt_tpl.as_deref();
s.spawn(move || {
let w_conn = match open_rw(&paths.db) {
Ok(c) => c,
Err(e) => {
tracing::error!(target: "enrich", worker = worker_id, error = %e, "worker failed to open DB");
return WorkerResult { completed: 0, failed: 0, skipped: 0, cost: 0.0, oauth: false, db_busy: false };
}
};
let w_queue = match open_queue_db(queue_path) {
Ok(c) => c,
Err(e) => {
tracing::error!(target: "enrich", worker = worker_id, error = %e, "worker failed to open queue DB");
return WorkerResult { completed: 0, failed: 0, skipped: 0, cost: 0.0, oauth: false, db_busy: false };
}
};
let mut w_completed = 0usize;
let mut w_failed = 0usize;
let mut w_skipped = 0usize;
let mut w_cost = 0.0f64;
let mut w_oauth = false;
let mut w_db_busy = false;
let mut w_backoff = DEFAULT_RATE_LIMIT_WAIT;
let w_deadline = std::time::Instant::now() + std::time::Duration::from_secs(3600);
let mut w_breaker = crate::retry::CircuitBreaker::new(
args.circuit_breaker_threshold.max(1),
std::time::Duration::from_secs(60),
);
loop {
if crate::shutdown_requested() {
tracing::info!(target: "enrich", "shutdown requested, worker stopping");
break;
}
if let Some(b) = budget {
if !w_oauth && w_cost >= b {
break;
}
}
let pending = match crate::storage::utils::with_busy_retry(|| {
dequeue_next_pending(&w_queue, op_label, backoff_clause)
}) {
Ok(DequeueOutcome::Claimed(p)) => Some(p),
Ok(DequeueOutcome::Empty) => None,
Err(AppError::DbBusy(msg)) => {
tracing::error!(target: "enrich", worker = worker_id, error = %msg, "SQLITE_BUSY exhausted bounded retries, worker aborting");
w_db_busy = true;
None
}
Err(e) => {
tracing::error!(target: "enrich", worker = worker_id, error = %e, "dequeue failed");
None
}
};
let claimed = match pending {
Some(p) => p,
None => break,
};
match validate_claim(&claimed, op_label, expected_item_type) {
ClaimCheck::Ok => {}
ClaimCheck::RequeueWrongOp => {
let _ = requeue_wrong_op(&w_queue, claimed.id);
continue;
}
ClaimCheck::SkipWrongType { reason } => {
let _ = skip_wrong_type(&w_queue, claimed.id, &reason);
w_skipped += 1;
continue;
}
}
let queue_id = claimed.id;
let item_key = claimed.item_key;
let attempt_current = claimed.attempt;
let _ = heartbeat(&w_queue, queue_id);
let item_started = Instant::now();
let current_index = w_completed + w_failed + w_skipped;
let provider_bin = provider_binary.unwrap_or_else(|| std::path::Path::new(""));
let call_result = match operation {
EnrichOperation::MemoryBindings | EnrichOperation::AugmentBindings => call_memory_bindings(&w_conn, namespace, &item_key, provider_bin, provider_model, provider_timeout, mode),
EnrichOperation::EntityDescriptions => call_entity_description(&w_conn, namespace, &item_key, provider_bin, provider_model, provider_timeout, mode, args.entity_description_grounding_threshold, &prompts::resolve_entity_description_domain(&args.entity_description_domain)),
EnrichOperation::BodyEnrich => call_body_enrich(&w_conn, namespace, &item_key, provider_bin, provider_model, provider_timeout, mode, min_oc, max_oc, prompt_tpl, args.preserve_threshold, paths, llm_backend, embedding_backend),
EnrichOperation::ReEmbed => call_reembed(&w_conn, namespace, &item_key, paths, llm_backend, embedding_backend),
EnrichOperation::WeightCalibrate => call_weight_calibrate(&w_conn, namespace, &item_key, provider_bin, provider_model, provider_timeout, mode),
EnrichOperation::RelationReclassify => call_relation_reclassify(&w_conn, namespace, &item_key, provider_bin, provider_model, provider_timeout, mode),
EnrichOperation::EntityConnect | EnrichOperation::CrossDomainBridges => call_entity_connect(&w_conn, namespace, &item_key, provider_bin, provider_model, provider_timeout, mode),
EnrichOperation::EntityTypeValidate => call_entity_type_validate(&w_conn, namespace, &item_key, provider_bin, provider_model, provider_timeout, mode),
EnrichOperation::DescriptionEnrich => call_description_enrich(&w_conn, namespace, &item_key, provider_bin, provider_model, provider_timeout, mode),
EnrichOperation::DomainClassify => call_domain_classify(&w_conn, namespace, &item_key, provider_bin, provider_model, provider_timeout, mode),
EnrichOperation::GraphAudit => call_graph_audit(&w_conn, namespace, &item_key, provider_bin, provider_model, provider_timeout, mode),
EnrichOperation::DeepResearchSynth => call_deep_research_synth(&w_conn, namespace, &item_key, provider_bin, provider_model, provider_timeout, mode),
EnrichOperation::BodyExtract => call_body_extract(&w_conn, namespace, &item_key, provider_bin, provider_model, provider_timeout, mode, args.body_extract_graph_only),
};
let openrouter_diag = take_last_openrouter_failure();
match call_result {
Ok(EnrichItemResult::Done { cost, is_oauth, memory_id, entity_id, entities, rels, chars_before, chars_after }) => {
if is_oauth { w_oauth = true; }
w_backoff = DEFAULT_RATE_LIMIT_WAIT;
let _ = w_queue.execute(
"UPDATE queue SET status='done', memory_id=?1, entity_id=?2, entities=?3, rels=?4, cost_usd=?5, elapsed_ms=?6, done_at=datetime('now') WHERE id=?7",
rusqlite::params![memory_id, entity_id, entities as i64, rels as i64, cost, item_started.elapsed().as_millis() as i64, queue_id],
);
w_completed += 1;
if !is_oauth { w_cost += cost; }
let _ = w_breaker
.record(crate::retry::AttemptOutcome::Success);
let _guard = stdout_mu.lock();
emit_json(&ItemEvent { item: &item_key, status: "done", memory_id, entity_id, entities: Some(entities), rels: Some(rels), chars_before, chars_after, cost_usd: if is_oauth { None } else { Some(cost) }, elapsed_ms: Some(item_started.elapsed().as_millis() as u64), error: None, index: current_index, total });
}
Ok(EnrichItemResult::Skipped { reason }) => {
w_skipped += 1;
let _ = w_queue.execute("UPDATE queue SET status='skipped', error=?1, done_at=datetime('now') WHERE id=?2", rusqlite::params![reason, queue_id]);
let _guard = stdout_mu.lock();
emit_json(&ItemEvent { item: &item_key, status: "skipped", memory_id: None, entity_id: None, entities: None, rels: None, chars_before: None, chars_after: None, cost_usd: None, elapsed_ms: Some(item_started.elapsed().as_millis() as u64), error: None, index: current_index, total });
}
Ok(EnrichItemResult::PreservationFailed { score, threshold, chars_before, chars_after }) => {
w_skipped += 1;
let reason = format!(
"preservation_failed: jaccard={score:.3} threshold={threshold:.3} (orig={chars_before} chars, new={chars_after} chars)"
);
let _ = w_queue.execute(
"UPDATE queue SET status='skipped', error=?1, done_at=datetime('now') WHERE id=?2",
rusqlite::params![reason, queue_id],
);
let _guard = stdout_mu.lock();
emit_json(&ItemEvent {
item: &item_key,
status: "preservation_failed",
memory_id: None,
entity_id: None,
entities: None,
rels: None,
chars_before: Some(chars_before),
chars_after: Some(chars_after),
cost_usd: None,
elapsed_ms: Some(item_started.elapsed().as_millis() as u64),
error: Some(reason),
index: current_index,
total,
});
}
Err(e) => {
let err_str = format!("{e}");
if matches!(e, AppError::RateLimited { .. }) {
if crate::retry::is_kill_switch_active() {
tracing::warn!(target: "enrich", "SQLITE_GRAPHRAG_DISABLE_RETRY=1, skipping rate-limit retry");
} else if std::time::Instant::now() >= w_deadline {
tracing::error!(target: "enrich", "rate-limit retry deadline (1h) exhausted in worker");
} else {
let half = w_backoff / 2;
let jitter = if half == 0 { 0 } else { fastrand::u64(0..half) };
let actual_wait = half + jitter;
tracing::warn!(target: "enrich", delay_secs = actual_wait, error_kind = "rate_limited", "rate limited in worker, backing off");
let _ = w_queue.execute("UPDATE queue SET status='pending' WHERE id=?1", rusqlite::params![queue_id]);
std::thread::sleep(std::time::Duration::from_secs(actual_wait));
w_backoff = (w_backoff * 2).min(900);
continue;
}
}
w_failed += 1;
let outcome = match openrouter_diag {
Some(diag) => record_item_failure_typed(
&w_queue,
queue_id,
attempt_current,
args.max_attempts,
diag.retry_class,
&err_str,
diag.finish_reason.as_deref(),
diag.prompt_tokens,
diag.completion_tokens,
),
None => record_item_failure(&w_queue, queue_id, attempt_current, args.max_attempts, &e),
};
let _guard = stdout_mu.lock();
emit_json(&ItemEvent { item: &item_key, status: "failed", memory_id: None, entity_id: None, entities: None, rels: None, chars_before: None, chars_after: None, cost_usd: None, elapsed_ms: Some(item_started.elapsed().as_millis() as u64), error: Some(err_str), index: current_index, total });
let breaker_opened = w_breaker.record(outcome);
if breaker_opened {
tracing::error!(target: "enrich",
consecutive_failures = w_breaker.consecutive_failures(),
"circuit breaker opened — aborting worker"
);
break;
}
}
}
}
WorkerResult { completed: w_completed, failed: w_failed, skipped: w_skipped, cost: w_cost, oauth: w_oauth, db_busy: w_db_busy }
})
})
.collect();
handles
.into_iter()
.map(|h| {
h.join().unwrap_or(WorkerResult {
completed: 0,
failed: 0,
skipped: 0,
cost: 0.0,
oauth: false,
db_busy: false,
})
})
.collect()
});
if results.iter().any(|r| r.db_busy) {
return Err(AppError::DbBusy(
"SQLITE_BUSY exhausted bounded retries while dequeuing (parallel worker)"
.into(),
));
}
for r in &results {
counters.completed += r.completed;
counters.failed += r.failed;
counters.skipped += r.skipped;
counters.cost_total += r.cost;
if r.oauth && !counters.oauth_detected {
counters.oauth_detected = true;
}
}
Ok(())
}