File size: 3,670 Bytes
30f011f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
//! 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(())
}