fix: harden ONNX model wrapper
- Store ONNX Session in Option<Mutex<Session>> 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 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.7
parent
fc34e52911
commit
36e4a49ca7
@@ -5,6 +5,7 @@ use ort::Session;
|
|||||||
use serde::Serialize;
|
use serde::Serialize;
|
||||||
use std::collections::BTreeMap;
|
use std::collections::BTreeMap;
|
||||||
use std::path::Path;
|
use std::path::Path;
|
||||||
|
use std::sync::Mutex;
|
||||||
|
|
||||||
/// Prediction result containing the top label, confidence, and all probabilities.
|
/// Prediction result containing the top label, confidence, and all probabilities.
|
||||||
#[derive(Debug, Clone, Serialize)]
|
#[derive(Debug, Clone, Serialize)]
|
||||||
@@ -16,12 +17,13 @@ pub struct Prediction {
|
|||||||
|
|
||||||
/// Service for running ONNX model inference.
|
/// 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.
|
/// 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 {
|
pub struct ModelService {
|
||||||
model_path: std::path::PathBuf,
|
model_path: std::path::PathBuf,
|
||||||
input_size: u32,
|
input_size: u32,
|
||||||
session: Option<Session>,
|
session: Option<Mutex<Session>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl ModelService {
|
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.
|
/// 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.
|
/// 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 {
|
pub fn new(model_path: &Path, input_size: u32) -> Self {
|
||||||
let session = Session::builder()
|
let session = Session::builder()
|
||||||
.ok()
|
.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 {
|
Self {
|
||||||
model_path: model_path.to_path_buf(),
|
model_path: model_path.to_path_buf(),
|
||||||
@@ -59,15 +63,20 @@ impl ModelService {
|
|||||||
/// Runs inference on the given input array.
|
/// Runs inference on the given input array.
|
||||||
///
|
///
|
||||||
/// Returns ModelUnavailable if the model is not loaded.
|
/// 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<f32>) -> Result<Prediction, ServiceError> {
|
pub fn predict(&self, input: Array4<f32>) -> Result<Prediction, ServiceError> {
|
||||||
let session = self
|
let session = self
|
||||||
.session
|
.session
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.ok_or_else(|| ServiceError::ModelUnavailable("Model is not loaded".to_string()))?;
|
.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
|
// Run inference
|
||||||
let outputs = session
|
let outputs = session_guard
|
||||||
.run(ort::inputs![input]?)
|
.run(ort::inputs![input]?)
|
||||||
.map_err(|_| ServiceError::PredictionFailed("Prediction failed".to_string()))?;
|
.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.
|
/// Maps a probability vector to a Prediction with label and all probabilities.
|
||||||
///
|
///
|
||||||
/// Expects a vector of length 4 (one per label in LABELS).
|
/// 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<Prediction, ServiceError> {
|
pub fn prediction_from_probabilities(probs: &[f32]) -> Result<Prediction, ServiceError> {
|
||||||
if probs.len() != LABELS.len() {
|
if probs.len() != LABELS.len() {
|
||||||
return Err(ServiceError::PredictionFailed("Prediction failed".to_string()));
|
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
|
// Find the index with the highest probability
|
||||||
let top_idx = probs
|
let top_idx = probs
|
||||||
.iter()
|
.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]
|
#[test]
|
||||||
fn missing_model_file_creates_unloaded_service() {
|
fn missing_model_file_creates_unloaded_service() {
|
||||||
let model_path = Path::new("/nonexistent/model.onnx");
|
let model_path = Path::new("/nonexistent/model.onnx");
|
||||||
|
|||||||
Reference in New Issue
Block a user