From 36e4a49ca7f1268870cc3f674ada2fb598f081b0 Mon Sep 17 00:00:00 2001 From: Asep Haryana Saputra <90584806+MythEclipse@users.noreply.github.com> Date: Sat, 23 May 2026 09:44:25 +0000 Subject: [PATCH] fix: harden ONNX model wrapper - Store ONNX Session in Option> for thread-safe concurrent access - Lock session before run() in predict(), map lock poisoning to PredictionFailed - Preserve ModelUnavailable when session is None - Reject non-finite values (NaN, +inf, -inf) in prediction_from_probabilities - Add unit tests for NaN, positive infinity, and negative infinity rejection - Keep all existing tests passing Co-Authored-By: Claude Opus 4.7 --- apps/ml-service/src/model.rs | 69 ++++++++++++++++++++++++++++++++---- 1 file changed, 63 insertions(+), 6 deletions(-) diff --git a/apps/ml-service/src/model.rs b/apps/ml-service/src/model.rs index 20b5742..7049140 100644 --- a/apps/ml-service/src/model.rs +++ b/apps/ml-service/src/model.rs @@ -5,6 +5,7 @@ use ort::Session; use serde::Serialize; use std::collections::BTreeMap; use std::path::Path; +use std::sync::Mutex; /// Prediction result containing the top label, confidence, and all probabilities. #[derive(Debug, Clone, Serialize)] @@ -16,12 +17,13 @@ pub struct Prediction { /// Service for running ONNX model inference. /// -/// Stores the model path, input size, and an optional ONNX session. +/// Stores the model path, input size, and an optional thread-safe ONNX session. /// If the model fails to load, the session remains None and predictions will fail. +/// The session is wrapped in a Mutex to ensure thread-safe access from concurrent Axum requests. pub struct ModelService { model_path: std::path::PathBuf, input_size: u32, - session: Option, + session: Option>, } impl ModelService { @@ -29,10 +31,12 @@ impl ModelService { /// /// If the model file does not exist or fails to load, the session is stored as None. /// This allows the service to report unloaded state via health checks. + /// The session is wrapped in a Mutex for thread-safe concurrent access. pub fn new(model_path: &Path, input_size: u32) -> Self { let session = Session::builder() .ok() - .and_then(|builder| builder.commit_from_file(model_path).ok()); + .and_then(|builder| builder.commit_from_file(model_path).ok()) + .map(Mutex::new); Self { model_path: model_path.to_path_buf(), @@ -59,15 +63,20 @@ impl ModelService { /// Runs inference on the given input array. /// /// Returns ModelUnavailable if the model is not loaded. - /// Returns PredictionFailed if inference fails or output format is invalid. + /// Returns PredictionFailed if inference fails, lock is poisoned, or output format is invalid. pub fn predict(&self, input: Array4) -> Result { let session = self .session .as_ref() .ok_or_else(|| ServiceError::ModelUnavailable("Model is not loaded".to_string()))?; + // Lock the session for thread-safe access + let session_guard = session + .lock() + .map_err(|_| ServiceError::PredictionFailed("Prediction failed".to_string()))?; + // Run inference - let outputs = session + let outputs = session_guard .run(ort::inputs![input]?) .map_err(|_| ServiceError::PredictionFailed("Prediction failed".to_string()))?; @@ -84,12 +93,18 @@ impl ModelService { /// Maps a probability vector to a Prediction with label and all probabilities. /// /// Expects a vector of length 4 (one per label in LABELS). - /// Returns PredictionFailed if the length is incorrect. + /// Rejects non-finite values (NaN, +inf, -inf) to prevent invalid predictions. + /// Returns PredictionFailed if the length is incorrect or any value is non-finite. pub fn prediction_from_probabilities(probs: &[f32]) -> Result { if probs.len() != LABELS.len() { return Err(ServiceError::PredictionFailed("Prediction failed".to_string())); } + // Reject non-finite values (NaN, +inf, -inf) + if probs.iter().any(|p| !p.is_finite()) { + return Err(ServiceError::PredictionFailed("Prediction failed".to_string())); + } + // Find the index with the highest probability let top_idx = probs .iter() @@ -153,6 +168,48 @@ mod tests { } } + #[test] + fn prediction_mapping_rejects_nan_values() { + let probs = [0.1, f32::NAN, 0.6, 0.1]; + let result = ModelService::prediction_from_probabilities(&probs); + + assert!(result.is_err()); + match result.unwrap_err() { + ServiceError::PredictionFailed(msg) => { + assert_eq!(msg, "Prediction failed"); + } + _ => panic!("expected PredictionFailed error"), + } + } + + #[test] + fn prediction_mapping_rejects_positive_infinity() { + let probs = [0.1, 0.2, f32::INFINITY, 0.1]; + let result = ModelService::prediction_from_probabilities(&probs); + + assert!(result.is_err()); + match result.unwrap_err() { + ServiceError::PredictionFailed(msg) => { + assert_eq!(msg, "Prediction failed"); + } + _ => panic!("expected PredictionFailed error"), + } + } + + #[test] + fn prediction_mapping_rejects_negative_infinity() { + let probs = [0.1, 0.2, 0.6, f32::NEG_INFINITY]; + let result = ModelService::prediction_from_probabilities(&probs); + + assert!(result.is_err()); + match result.unwrap_err() { + ServiceError::PredictionFailed(msg) => { + assert_eq!(msg, "Prediction failed"); + } + _ => panic!("expected PredictionFailed error"), + } + } + #[test] fn missing_model_file_creates_unloaded_service() { let model_path = Path::new("/nonexistent/model.onnx");