bert-agent / daemon /src /main.rs
SNAPKITTYWEST's picture
push from SNAPKITTYWEST/bert-agent
30f011f verified
Raw
History Blame Contribute Delete
3.67 kB
//! bert-daemon β€” Sovereign Cross-Encoder Entailment Inference Daemon
//!
//! Startup sequence:
//! 1. Load DaemonConfig from config/daemon.json
//! 2. Build TensorRT ORT session (loads cached .plan or compiles ~5 min)
//! 3. Load tokenizer from tokenizer/
//! 4. Spawn: WORM ledger worker
//! 5. Spawn: inference daemon (dual-trigger continuous batching)
//! 6. Bind Axum HTTP server on configured port
//!
//! All entailment decisions are sealed with BLAKE3 and appended to the
//! WORM audit chain before the response is returned to the caller.
mod inference;
mod ledger;
mod session;
mod server;
mod types;
use std::sync::Arc;
use clap::Parser;
use tokio::sync::mpsc;
use crate::ledger::run_ledger_worker;
use crate::inference::run_inference_daemon;
use crate::session::build_trt_session;
use crate::server::{AppState, build_router};
use crate::types::DaemonConfig;
#[derive(Parser, Debug)]
#[command(name = "bert-daemon", about = "Cross-Encoder entailment inference daemon")]
struct Cli {
#[arg(short, long, default_value = "config/daemon.json")]
config: String,
}
#[tokio::main]
async fn main() -> anyhow::Result<()> {
env_logger::init();
let cli = Cli::parse();
// 1. Load config
let cfg: DaemonConfig = {
let raw = std::fs::read_to_string(&cli.config)
.unwrap_or_else(|_| {
log::warn!("config not found at {}, using defaults", cli.config);
serde_json::to_string(&DaemonConfig::default()).unwrap()
});
serde_json::from_str(&raw)?
};
let cfg = Arc::new(cfg);
log::info!("[main] config loaded: model={} threshold={}", cfg.model_path, cfg.threshold);
// 2. Build TRT session
log::info!("[main] initialising TensorRT session...");
let session = Arc::new(build_trt_session(&cfg)?);
log::info!("[main] TRT session ready");
// 3. Load tokenizer
let tokenizer = Arc::new(
tokenizers::Tokenizer::from_pretrained(
"microsoft/deberta-v3-base",
None,
).expect("tokenizer load failed β€” run training first or place tokenizer/ in working dir"),
);
// 4. MPSC channels
// inference_tx/rx: HTTP handlers β†’ inference daemon
let (inference_tx, inference_rx) = mpsc::channel::<crate::types::VerifyRequest>(10_000);
// ledger_tx/rx: inference daemon β†’ WORM ledger worker
let (ledger_tx, ledger_rx) = mpsc::channel::<([u8; 32], Vec<u8>)>(10_000);
// 5. Spawn WORM ledger worker
let ledger_path = cfg.ledger_path.clone();
tokio::spawn(async move {
run_ledger_worker(ledger_rx, ledger_path).await;
});
log::info!("[main] WORM ledger worker spawned");
// 6. Spawn inference daemon
{
let session = Arc::clone(&session);
let cfg_clone = Arc::clone(&cfg);
let ledger_tx = ledger_tx.clone();
tokio::spawn(async move {
run_inference_daemon(inference_rx, session, cfg_clone, ledger_tx).await;
});
}
log::info!("[main] inference daemon spawned β€” batch={} flush={}ms",
cfg.max_batch_size, cfg.flush_interval_ms);
// 7. Bind HTTP server
let app_state = Arc::new(AppState {
tx: inference_tx,
cfg: Arc::clone(&cfg),
tokenizer,
});
let router = build_router(app_state);
let addr = format!("0.0.0.0:{}", cfg.http_port);
log::info!("[main] HTTP server β†’ {}", addr);
let listener = tokio::net::TcpListener::bind(&addr).await?;
axum::serve(listener, router).await?;
Ok(())
}