//! inference.rs — continuous batching loop with dual-trigger flush //! //! Architecture: //! - Tokio MPSC channel receives VerifyRequests from HTTP handlers //! - tokio::select! races: MAX_BATCH_SIZE trigger vs 5 ms timer //! - execute_batch: dynamic pad → ndarray → TRT forward → softmax → BLAKE3 seal //! - Attestations dispatched to background WORM ledger worker //! - Results routed back through oneshot channels (no polling) use std::sync::Arc; use std::time::SystemTime; use ndarray::{s, Array2}; use ort::Session; use tokio::sync::mpsc; use tokio::time::{interval, Duration}; use crate::types::{ DaemonConfig, EntailmentAttestation, Verdict, VerifyRequest, VerifyResponse, }; // ── Softmax ───────────────────────────────────────────────────────────────── fn softmax_entailment(logits: &[f32]) -> f32 { // logits: [contradiction, neutral, entailment] let max_l = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max); let exps: Vec = logits.iter().map(|l| (l - max_l).exp()).collect(); let sum: f32 = exps.iter().sum(); exps[2] / sum // P(Entailment) } // ── Batch execution ────────────────────────────────────────────────────────── async fn execute_batch( batch: &mut Vec, session: &Arc, cfg: &DaemonConfig, ledger_tx: &mpsc::Sender<([u8; 32], Vec)>, ) { if batch.is_empty() { return; } let batch_size = batch.len(); // 1. Dynamic padding: pad to the longest sequence in THIS batch (not global max). // Avoids wasting compute padding short sequences to 512. let max_len = batch .iter() .map(|r| r.input_ids.len()) .max() .unwrap_or(0); let mut input_ids_arr = Array2::::zeros((batch_size, max_len)); let mut attention_mask_arr = Array2::::zeros((batch_size, max_len)); let mut token_types_arr = Array2::::zeros((batch_size, max_len)); for (i, req) in batch.iter().enumerate() { let len = req.input_ids.len().min(max_len); input_ids_arr .slice_mut(s![i, ..len]) .assign(&ndarray::ArrayView::from(&req.input_ids[..len])); attention_mask_arr .slice_mut(s![i, ..len]) .assign(&ndarray::ArrayView::from(&req.attention_mask[..len])); token_types_arr .slice_mut(s![i, ..len]) .assign(&ndarray::ArrayView::from(&req.token_type_ids[..len])); } // 2. Run TRT forward pass on a blocking thread (keeps async reactor free). let session_clone = Arc::clone(session); let outputs = tokio::task::spawn_blocking(move || { let inputs = ort::inputs![ "input_ids" => input_ids_arr, "attention_mask" => attention_mask_arr, "token_type_ids" => token_types_arr, ] .expect("input construction failed"); session_clone.run(inputs).expect("TRT inference failed") }) .await .expect("spawn_blocking panicked"); // 3. Extract logits (B, 3) → entailment scores. let logits_tensor = outputs["logits"] .extract_tensor::() .expect("logits extraction failed"); let logits_view = logits_tensor.view(); // shape (B, 3) let timestamp_ns = SystemTime::now() .duration_since(SystemTime::UNIX_EPOCH) .unwrap() .as_nanos() as u64; // 4. For each request: compute score → BLAKE3 seal → route response. for (i, req) in batch.drain(..).enumerate() { let row: Vec = logits_view .row(i) .iter() .cloned() .collect(); let score = softmax_entailment(&row); let verdict = Verdict::from_score(score, cfg.threshold); // 5. Build deterministic attestation and seal with BLAKE3. let attestation = EntailmentAttestation { timestamp_ns, chunk_id: req.chunk_id.clone(), claim_text: req.claim_text.clone(), entailment_score: score, verdict: format!("{:?}", verdict), model_signature: cfg.model_signature.clone(), threshold: cfg.threshold, }; let (hash, payload) = attestation.seal(); // 6. Dispatch attestation to WORM ledger (non-blocking). let _ = ledger_tx.try_send((hash, payload)); // 7. Return result to caller through oneshot channel. let _ = req.responder.send(VerifyResponse { entailment_score: score, label: verdict, attestation_hash: hash, }); } } // ── Continuous batching event loop ─────────────────────────────────────────── /// Dual-trigger: flush when MAX_BATCH_SIZE is reached OR every flush_interval_ms. /// Guarantees maximum latency = flush_interval_ms (default 5 ms). pub async fn run_inference_daemon( mut rx: mpsc::Receiver, session: Arc, cfg: Arc, ledger_tx: mpsc::Sender<([u8; 32], Vec)>, ) { let mut batch: Vec = Vec::with_capacity(cfg.max_batch_size); let mut flush_timer = interval(Duration::from_millis(cfg.flush_interval_ms)); log::info!( "[daemon] running — max_batch={} flush_interval={}ms threshold={}", cfg.max_batch_size, cfg.flush_interval_ms, cfg.threshold ); loop { tokio::select! { // New request arrived Some(req) = rx.recv() => { batch.push(req); if batch.len() >= cfg.max_batch_size { execute_batch(&mut batch, &session, &cfg, &ledger_tx).await; } } // Flush timer fired — process whatever is in the queue _ = flush_timer.tick() => { if !batch.is_empty() { execute_batch(&mut batch, &session, &cfg, &ledger_tx).await; } } // Channel closed — drain remaining requests and exit else => { if !batch.is_empty() { execute_batch(&mut batch, &session, &cfg, &ledger_tx).await; } log::info!("[daemon] channel closed, shutting down"); break; } } } }