Download daemon/src/server.rs from Snapkitty/bert-agent: direct link, hf CLI and curl.
- Browser
- Download file 5.42 kB
-
https://huggingface.co/Snapkitty/bert-agent/resolve/main/daemon/src/server.rs
- Command line
-
hf download hf://Snapkitty/bert-agent/daemon/src/server.rs
-
curl -L -o server.rs https://huggingface.co/Snapkitty/bert-agent/resolve/main/daemon/src/server.rs
5.42 kB
| //! server.rs β Axum HTTP server | |
| //! | |
| //! POST /verify | |
| //! Body: { "premise": "...", "hypothesis": "...", "chunk_id": "..." } | |
| //! Response: { "score": 0.94, "verdict": "Entailment", "hash": "abc123..." } | |
| //! | |
| //! The handler tokenises the (premise, hypothesis) pair using the same | |
| //! cross-encoder format as training: [CLS] premise [SEP] hypothesis [SEP]. | |
| //! It then sends a VerifyRequest through the MPSC channel to the inference | |
| //! daemon and awaits the oneshot response. | |
| use std::sync::Arc; | |
| use axum::{ | |
| extract::State, | |
| http::StatusCode, | |
| response::IntoResponse, | |
| routing::post, | |
| Json, Router, | |
| }; | |
| use serde::{Deserialize, Serialize}; | |
| use tokio::sync::{mpsc, oneshot}; | |
| use crate::types::{DaemonConfig, VerifyRequest, VerifyResponse, Verdict}; | |
| // ββ Request / Response DTOs βββββββββββββββββββββββββββββββββββββββββββββββββ | |
| pub struct VerifyBody { | |
| pub premise: String, // retrieved source chunk | |
| pub hypothesis: String, // LLM generated claim | |
| pub chunk_id: String, | |
| } | |
| pub struct VerifyReply { | |
| pub score: f32, | |
| pub verdict: String, | |
| pub hash: String, // BLAKE3 hex for the audit ledger | |
| } | |
| // ββ Shared state βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| pub struct AppState { | |
| pub tx: mpsc::Sender<VerifyRequest>, | |
| pub cfg: Arc<DaemonConfig>, | |
| pub tokenizer: Arc<tokenizers::Tokenizer>, | |
| } | |
| // ββ Tokenisation βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| /// Encode (premise, hypothesis) as a cross-encoder input: | |
| /// [CLS] premise_tokens [SEP] hypothesis_tokens [SEP] | |
| fn encode_pair( | |
| tokenizer: &tokenizers::Tokenizer, | |
| premise: &str, | |
| hypothesis: &str, | |
| max_length: usize, | |
| ) -> (Vec<i64>, Vec<i64>, Vec<i64>) { | |
| use tokenizers::EncodeInput; | |
| let encoding = tokenizer | |
| .encode( | |
| EncodeInput::Dual( | |
| tokenizers::InputSequence::Raw(premise.into()), | |
| tokenizers::InputSequence::Raw(hypothesis.into()), | |
| ), | |
| true, | |
| ) | |
| .expect("tokenisation failed"); | |
| let ids: Vec<i64> = encoding.get_ids().iter().map(|&x| x as i64).collect(); | |
| let mask: Vec<i64> = encoding.get_attention_mask().iter().map(|&x| x as i64).collect(); | |
| let types: Vec<i64> = encoding.get_type_ids().iter().map(|&x| x as i64).collect(); | |
| // Truncate to max_length | |
| let trunc = |v: Vec<i64>| v.into_iter().take(max_length).collect::<Vec<_>>(); | |
| (trunc(ids), trunc(mask), trunc(types)) | |
| } | |
| // ββ Handler ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| async fn verify_handler( | |
| State(state): State<Arc<AppState>>, | |
| Json(body): Json<VerifyBody>, | |
| ) -> impl IntoResponse { | |
| let (input_ids, attention_mask, token_type_ids) = encode_pair( | |
| &state.tokenizer, | |
| &body.premise, | |
| &body.hypothesis, | |
| 512, | |
| ); | |
| let (resp_tx, resp_rx) = oneshot::channel::<VerifyResponse>(); | |
| let request = VerifyRequest { | |
| input_ids, | |
| attention_mask, | |
| token_type_ids, | |
| chunk_id: body.chunk_id, | |
| claim_text: body.hypothesis, | |
| responder: resp_tx, | |
| }; | |
| if state.tx.send(request).await.is_err() { | |
| return ( | |
| StatusCode::SERVICE_UNAVAILABLE, | |
| Json(serde_json::json!({"error": "inference daemon unavailable"})), | |
| ); | |
| } | |
| match resp_rx.await { | |
| Ok(resp) => { | |
| let hash_hex = resp.attestation_hash | |
| .iter() | |
| .map(|b| format!("{:02x}", b)) | |
| .collect::<String>(); | |
| ( | |
| StatusCode::OK, | |
| Json(serde_json::json!({ | |
| "score": resp.entailment_score, | |
| "verdict": format!("{:?}", resp.label), | |
| "hash": hash_hex, | |
| })), | |
| ) | |
| } | |
| Err(_) => ( | |
| StatusCode::INTERNAL_SERVER_ERROR, | |
| Json(serde_json::json!({"error": "inference worker dropped"})), | |
| ), | |
| } | |
| } | |
| // ββ Health check βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| async fn health_handler() -> impl IntoResponse { | |
| Json(serde_json::json!({"status": "ok"})) | |
| } | |
| // ββ Router builder βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| pub fn build_router(state: Arc<AppState>) -> Router { | |
| Router::new() | |
| .route("/verify", post(verify_handler)) | |
| .route("/health", axum::routing::get(health_handler)) | |
| .with_state(state) | |
| } | |