#![cfg(feature = "ort")]
use std::path::{Path, PathBuf};
use std::sync::Mutex;
use ort::execution_providers::CPUExecutionProvider;
use ort::inputs;
use ort::session::{builder::GraphOptimizationLevel, Session};
use ort::value::Value;
use tokenizers::Tokenizer;
use crate::engine::EngineError;
const MAX_SEQ_LEN: usize = 512;
pub struct InjectionClassifier {
session: Mutex<Session>,
tokenizer: Tokenizer,
injection_idx: usize,
#[allow(dead_code)]
model_dir: PathBuf,
}
impl std::fmt::Debug for InjectionClassifier {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("InjectionClassifier")
.field("model_dir", &self.model_dir)
.field("injection_idx", &self.injection_idx)
.finish_non_exhaustive()
}
}
impl InjectionClassifier {
pub fn from_model_dir(dir: &Path) -> Result<Self, EngineError> {
let dir_str = dir.to_string_lossy().into_owned();
let model_dir = dir.to_path_buf();
let onnx_path = model_dir.join("model.onnx");
let tok_path = model_dir.join("tokenizer.json");
let cfg_path = model_dir.join("config.json");
for p in [&onnx_path, &tok_path, &cfg_path] {
if !p.exists() {
return Err(EngineError::ModelNotFound {
dir: dir_str.clone(),
});
}
}
let tokenizer = Tokenizer::from_file(&tok_path)
.map_err(|e| EngineError::TokenizerLoad(e.to_string()))?;
let injection_idx = parse_injection_idx(&cfg_path)?;
let _ = ort::init()
.with_name("vigil-injection-classifier")
.with_execution_providers([CPUExecutionProvider::default().build()])
.commit();
let session = Session::builder()
.map_err(|e| EngineError::SessionInit(e.to_string()))?
.with_optimization_level(GraphOptimizationLevel::Level1)
.map_err(|e| EngineError::SessionInit(e.to_string()))?
.with_intra_threads(4)
.map_err(|e| EngineError::SessionInit(e.to_string()))?
.commit_from_file(&onnx_path)
.map_err(|e| EngineError::SessionInit(e.to_string()))?;
Ok(Self {
session: Mutex::new(session),
tokenizer,
injection_idx,
model_dir,
})
}
pub fn classify(&self, text: &str) -> Result<f32, EngineError> {
let enc = self
.tokenizer
.encode(text, true)
.map_err(|e| EngineError::InferRun(e.to_string()))?;
let mut ids: Vec<i64> = enc.get_ids().iter().map(|&i| i as i64).collect();
let mut mask: Vec<i64> = enc.get_attention_mask().iter().map(|&m| m as i64).collect();
if ids.len() > MAX_SEQ_LEN {
ids.truncate(MAX_SEQ_LEN);
mask.truncate(MAX_SEQ_LEN);
}
let seq_len = ids.len();
if seq_len == 0 {
return Ok(0.0);
}
let input_ids_val = Value::from_array((vec![1i64, seq_len as i64], ids))
.map_err(|e| EngineError::DecodeShape(e.to_string()))?;
let mask_val = Value::from_array((vec![1i64, seq_len as i64], mask))
.map_err(|e| EngineError::DecodeShape(e.to_string()))?;
let (shape, data): (Vec<i64>, Vec<f32>) = {
let mut session = self
.session
.lock()
.map_err(|e| EngineError::Internal(format!("session mutex poisoned: {e}")))?;
let outputs = session
.run(inputs![
"input_ids" => input_ids_val,
"attention_mask" => mask_val,
])
.map_err(|e| EngineError::InferRun(e.to_string()))?;
let (_name, logits_val) = outputs
.iter()
.next()
.ok_or_else(|| EngineError::DecodeShape("no output tensor".to_string()))?;
let (raw_shape, raw_data) = logits_val
.try_extract_tensor::<f32>()
.map_err(|e| EngineError::DecodeShape(e.to_string()))?;
(raw_shape.to_vec(), raw_data.to_vec())
};
let num_labels = match shape.as_slice() {
[1, n] => *n as usize,
[n] => *n as usize,
_ => {
return Err(EngineError::DecodeShape(format!(
"unexpected logits shape (want [1,2] or [2]): {shape:?}"
)));
}
};
if num_labels < 2 || self.injection_idx >= num_labels || data.len() < num_labels {
return Err(EngineError::DecodeShape(format!(
"logits len/labels mismatch: shape={shape:?} data_len={}",
data.len()
)));
}
let logits = &data[..num_labels];
let max_logit = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let sum_exp: f32 = logits.iter().map(|&v| (v - max_logit).exp()).sum();
if sum_exp <= 0.0 {
return Ok(0.0);
}
let p_injection = (logits[self.injection_idx] - max_logit).exp() / sum_exp;
Ok(p_injection)
}
pub fn warmup(&self) -> Result<(), EngineError> {
let _ = self.classify("a")?;
Ok(())
}
}
fn parse_injection_idx(cfg_path: &Path) -> Result<usize, EngineError> {
let raw = std::fs::read_to_string(cfg_path)
.map_err(|e| EngineError::Internal(format!("read config.json: {e}")))?;
let cfg: serde_json::Value = serde_json::from_str(&raw)
.map_err(|e| EngineError::Internal(format!("parse config.json: {e}")))?;
if let Some(idx) = cfg
.get("label2id")
.and_then(|v| v.get("INJECTION"))
.and_then(|v| v.as_u64())
{
return Ok(idx as usize);
}
if let Some(obj) = cfg.get("id2label").and_then(|v| v.as_object()) {
for (k, v) in obj {
if v.as_str() == Some("INJECTION") {
if let Ok(idx) = k.parse::<usize>() {
return Ok(idx);
}
}
}
}
Err(EngineError::Internal(
"config.json missing INJECTION label in label2id/id2label".to_string(),
))
}
#[cfg(test)]
mod injection_static_assertions {
use super::*;
fn _assert_send_sync<T: Send + Sync>() {}
#[allow(dead_code)]
fn _check() {
_assert_send_sync::<InjectionClassifier>();
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
use std::io::Write;
#[test]
fn from_model_dir_missing_returns_modelnotfound() {
let r = InjectionClassifier::from_model_dir(Path::new("/nonexistent/vigil/deberta"));
assert!(
matches!(r, Err(EngineError::ModelNotFound { .. })),
"缺三件套应返 ModelNotFound,实际: {:?}",
r.map(|_| "Ok")
);
}
#[test]
fn parse_injection_idx_from_label2id() {
let dir = tempfile::tempdir().unwrap();
let cfg = dir.path().join("config.json");
let mut f = std::fs::File::create(&cfg).unwrap();
write!(
f,
r#"{{"id2label":{{"0":"SAFE","1":"INJECTION"}},"label2id":{{"SAFE":0,"INJECTION":1}}}}"#
)
.unwrap();
assert_eq!(parse_injection_idx(&cfg).unwrap(), 1);
}
#[test]
fn parse_injection_idx_fallback_id2label() {
let dir = tempfile::tempdir().unwrap();
let cfg = dir.path().join("config.json");
let mut f = std::fs::File::create(&cfg).unwrap();
write!(f, r#"{{"id2label":{{"0":"INJECTION","1":"SAFE"}}}}"#).unwrap();
assert_eq!(parse_injection_idx(&cfg).unwrap(), 0);
}
#[test]
fn parse_injection_idx_missing_label_fails() {
let dir = tempfile::tempdir().unwrap();
let cfg = dir.path().join("config.json");
let mut f = std::fs::File::create(&cfg).unwrap();
write!(f, r#"{{"id2label":{{"0":"FOO","1":"BAR"}}}}"#).unwrap();
assert!(matches!(
parse_injection_idx(&cfg),
Err(EngineError::Internal(_))
));
}
}