mod rules;
pub use rules::{
AttributeProtoValidityRule, DuplicateValueNameRule, FunctionProtoValidityRule,
GraphAcyclicRule, InitializerTypeMatchesDeclaredRule, InputOutputDeclaredRule,
IrVersionFeatureRule, IrVersionSupportedRule, MetadataKeysUniqueRule, MissingOpsetImportRule,
MultiDeviceConfigurationRule, NoUnconnectedNodesRule, ProtoTypeValidityRule,
SchemaNodeConformsRule, SparseTensorValidityRule, TensorPayloadValidityRule,
TypeConstraintSatisfiedRule,
};
use crate::model::Model;
use crate::schema::SchemaRegistry;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Severity {
Error,
Warning,
Info,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ViolationLocation {
Model,
Graph { graph_name: String },
Node {
graph_name: String,
node_name: String,
},
Value { value_name: String },
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Violation {
pub rule_id: String,
pub severity: Severity,
pub message: String,
pub location: ViolationLocation,
}
#[derive(Clone, Debug)]
pub struct ValidationContext {
schemas: std::sync::Arc<SchemaRegistry>,
}
impl ValidationContext {
pub fn new(schemas: SchemaRegistry) -> Self {
Self {
schemas: std::sync::Arc::new(schemas),
}
}
pub fn schemas(&self) -> &SchemaRegistry {
&self.schemas
}
}
impl Default for ValidationContext {
fn default() -> Self {
Self::new(SchemaRegistry::builtins())
}
}
pub trait ValidationRule: Send + Sync {
fn id(&self) -> &str;
fn severity(&self) -> Severity;
fn check(&self, model: &Model, ctx: &ValidationContext) -> Vec<Violation>;
}
#[derive(Clone, Debug, Default)]
pub struct ValidationResult {
pub violations: Vec<Violation>,
pub errors: usize,
pub warnings: usize,
}
impl ValidationResult {
pub fn is_valid(&self) -> bool {
self.errors == 0
}
}
pub struct OnnxChecker {
rules: Vec<Box<dyn ValidationRule>>,
disabled: std::collections::HashSet<String>,
context: ValidationContext,
}
impl OnnxChecker {
pub fn empty() -> Self {
Self {
rules: Vec::new(),
disabled: std::collections::HashSet::new(),
context: ValidationContext::default(),
}
}
pub fn new() -> Self {
let mut checker = Self::empty();
checker.add_rule(MissingOpsetImportRule);
checker.add_rule(DuplicateValueNameRule);
checker.add_rule(GraphAcyclicRule);
checker.add_rule(SchemaNodeConformsRule);
checker.add_rule(InputOutputDeclaredRule);
checker.add_rule(NoUnconnectedNodesRule);
checker.add_rule(TypeConstraintSatisfiedRule);
checker.add_rule(InitializerTypeMatchesDeclaredRule);
checker.add_rule(IrVersionSupportedRule);
checker.add_rule(IrVersionFeatureRule);
checker.add_rule(FunctionProtoValidityRule);
checker.add_rule(MetadataKeysUniqueRule);
checker.add_rule(AttributeProtoValidityRule);
checker.add_rule(ProtoTypeValidityRule);
checker.add_rule(TensorPayloadValidityRule);
checker.add_rule(SparseTensorValidityRule);
checker.add_rule(MultiDeviceConfigurationRule);
checker
}
pub fn with_schema_registry(schemas: SchemaRegistry) -> Self {
let mut checker = Self::new();
checker.context = ValidationContext::new(schemas);
checker
}
pub fn add_rule<R: ValidationRule + 'static>(&mut self, rule: R) {
self.rules.push(Box::new(rule));
}
pub fn disable_rule(&mut self, rule_id: &str) {
self.disabled.insert(rule_id.to_string());
}
pub fn rule_ids(&self) -> Vec<&str> {
self.rules.iter().map(|r| r.id()).collect()
}
pub fn check(&self, model: &Model) -> ValidationResult {
let mut result = ValidationResult::default();
for rule in &self.rules {
if self.disabled.contains(rule.id()) {
continue;
}
for violation in rule.check(model, &self.context) {
match violation.severity {
Severity::Error => result.errors += 1,
Severity::Warning => result.warnings += 1,
Severity::Info => {}
}
result.violations.push(violation);
}
}
result
}
}
impl Default for OnnxChecker {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
struct AlwaysFails;
impl ValidationRule for AlwaysFails {
fn id(&self) -> &str {
"test.always_fails"
}
fn severity(&self) -> Severity {
Severity::Error
}
fn check(&self, _model: &Model, _ctx: &ValidationContext) -> Vec<Violation> {
vec![Violation {
rule_id: self.id().to_string(),
severity: Severity::Error,
message: "boom".to_string(),
location: ViolationLocation::Model,
}]
}
}
fn empty_model() -> Model {
let mut g = onnx_runtime_ir::Graph::new();
g.opset_imports.insert(String::new(), 21);
Model::new(g)
}
#[test]
fn empty_checker_reports_no_violations() {
let result = OnnxChecker::empty().check(&empty_model());
assert!(result.is_valid());
assert!(result.violations.is_empty());
}
#[test]
fn custom_rule_runs_and_counts_errors() {
let mut checker = OnnxChecker::empty();
checker.add_rule(AlwaysFails);
let result = checker.check(&empty_model());
assert_eq!(result.errors, 1);
assert!(!result.is_valid());
}
#[test]
fn disabled_rule_is_skipped() {
let mut checker = OnnxChecker::empty();
checker.add_rule(AlwaysFails);
checker.disable_rule("test.always_fails");
let result = checker.check(&empty_model());
assert!(result.is_valid());
}
#[test]
fn default_checker_has_builtin_rules() {
let checker = OnnxChecker::new();
let ids = checker.rule_ids();
assert!(ids.contains(&"ir.opset_import_present"));
assert!(ids.contains(&"structure.duplicate_value_name"));
assert!(ids.contains(&"structure.graph_acyclic"));
assert!(ids.contains(&"schema.node_conforms"));
assert!(ids.contains(&"structure.input_output_declared"));
assert!(ids.contains(&"structure.no_unconnected_nodes"));
assert!(ids.contains(&"schema.type_constraint_satisfied"));
assert!(ids.contains(&"type.initializer_matches_declared"));
assert!(ids.contains(&"ir.version_supported"));
assert!(ids.contains(&"ir.version_gated_features"));
assert!(ids.contains(&"proto.function_valid"));
assert!(ids.contains(&"proto.metadata_keys_unique"));
assert!(ids.contains(&"proto.attribute_valid"));
assert!(ids.contains(&"proto.type_valid"));
assert!(ids.contains(&"proto.tensor_payload_valid"));
assert!(ids.contains(&"proto.sparse_tensor_valid"));
assert!(ids.contains(&"multidevice.configuration_valid"));
}
#[test]
fn validation_context_exposes_custom_registry() {
let context = ValidationContext::new(SchemaRegistry::new());
assert_eq!(context.schemas().iter().count(), 0);
let checker = OnnxChecker::with_schema_registry(SchemaRegistry::new());
assert!(checker.rule_ids().contains(&"schema.node_conforms"));
}
}