use std::collections::HashMap;
use petgraph::graph::NodeIndex;
use super::rule_procedures::{
execute_cardinality_violation, execute_inverse_violation, execute_missing_required_edge,
execute_transitivity_violation, require_node_yield, type_indices,
};
use crate::datatypes::values::Value;
use crate::graph::languages::cypher::ast::YieldItem;
use crate::graph::languages::cypher::result::ResultRow;
use crate::graph::ontology::{OntologyStore, RelationshipDecl};
use crate::graph::schema::{DirGraph, InternedKey};
use crate::graph::storage::GraphRead;
fn accepted_types(store: &OntologyStore, class_or_type: &str) -> Vec<String> {
let mut out = vec![class_or_type.to_string()];
for name in store.classes.keys() {
if store
.ancestors(name)
.iter()
.any(|ancestor| ancestor == class_or_type)
{
out.push(name.clone());
}
}
out
}
pub(super) fn scan_endpoint_mismatch(
graph: &DirGraph,
edge_type: &str,
accepted: &[String],
check_source: bool,
) -> Vec<(NodeIndex, NodeIndex)> {
let key = InternedKey::from_str(edge_type);
let mut out = Vec::new();
for er in graph.graph.edge_references() {
if er.weight().connection_type != key {
continue;
}
let subject = if check_source {
er.source()
} else {
er.target()
};
let actual = match graph.graph.node_view(subject) {
Some(n) => n.node_type_str(&graph.interner).to_string(),
None => continue,
};
if !accepted.contains(&actual) {
out.push((er.source(), er.target()));
}
}
out
}
fn endpoint_rows(
pairs: Vec<(NodeIndex, NodeIndex)>,
src_var: &str,
tgt_var: &str,
) -> Vec<ResultRow> {
pairs
.into_iter()
.map(|(src, tgt)| {
let mut row = ResultRow::new();
row.node_bindings.insert(src_var.to_string(), src);
row.node_bindings.insert(tgt_var.to_string(), tgt);
row
})
.collect()
}
pub(super) fn stamp_rule(rows: &mut [ResultRow], yield_items: &[YieldItem], rule: &str) {
let Some(alias) = yield_items
.iter()
.find(|y| y.name == "rule")
.map(|y| y.alias.clone().unwrap_or_else(|| "rule".to_string()))
else {
return;
};
for row in rows {
row.projected
.insert(alias.clone(), Value::String(rule.to_string()));
}
}
fn live(graph: &DirGraph, node_type: &str) -> bool {
graph.type_indices.contains_key(node_type)
}
fn edge_type_exists(graph: &DirGraph, edge_type: &str) -> bool {
let key = InternedKey::from_str(edge_type);
graph
.graph
.edge_references()
.any(|er| er.weight().connection_type == key)
}
fn string_params(pairs: &[(&str, &str)]) -> HashMap<String, Value> {
pairs
.iter()
.map(|(k, v)| (k.to_string(), Value::String(v.to_string())))
.collect()
}
fn declared_checks(rel: &str, decl: &RelationshipDecl) -> Vec<DeclaredCheck> {
let mut out = Vec::new();
if decl.domain.is_some() {
out.push(DeclaredCheck::Domain);
}
if decl.range.is_some() {
out.push(DeclaredCheck::Range);
}
if decl.required {
out.push(DeclaredCheck::Required);
}
if decl.cardinality.is_some() {
out.push(DeclaredCheck::Cardinality);
}
if decl.inverse_name.is_some() {
out.push(DeclaredCheck::Inverse);
}
if decl.symmetric {
out.push(DeclaredCheck::Symmetric);
}
if decl.transitive {
out.push(DeclaredCheck::Transitive);
}
let _ = rel;
out
}
#[derive(Clone, Copy, PartialEq)]
enum DeclaredCheck {
Domain,
Range,
Required,
Cardinality,
Inverse,
Symmetric,
Transitive,
}
impl DeclaredCheck {
fn name(&self) -> &'static str {
match self {
DeclaredCheck::Domain => "domain",
DeclaredCheck::Range => "range",
DeclaredCheck::Required => "required",
DeclaredCheck::Cardinality => "cardinality",
DeclaredCheck::Inverse => "inverse",
DeclaredCheck::Symmetric => "symmetric",
DeclaredCheck::Transitive => "transitive",
}
}
}
fn check_rows(
graph: &DirGraph,
rel: &str,
decl: &RelationshipDecl,
check: DeclaredCheck,
yield_items: &[YieldItem],
) -> Result<Vec<ResultRow>, String> {
let store = &graph.ontology;
match check {
DeclaredCheck::Domain | DeclaredCheck::Range => {
let check_source = check == DeclaredCheck::Domain;
let endpoint = if check_source {
decl.domain.as_deref()
} else {
decl.range.as_deref()
}
.expect("declared_checks gated on presence");
let accepted = accepted_types(store, endpoint);
let proc = if check_source {
"type_domain_violation"
} else {
"type_range_violation"
};
let src_var = require_node_yield(yield_items, proc, "source")?;
let tgt_var = require_node_yield(yield_items, proc, "target")?;
let pairs = scan_endpoint_mismatch(graph, rel, &accepted, check_source);
Ok(endpoint_rows(pairs, &src_var, &tgt_var))
}
DeclaredCheck::Required => {
let Some(domain) = decl.domain.as_deref() else {
return Ok(Vec::new());
};
let mut rows = Vec::new();
for node_type in accepted_types(store, domain) {
if !live(graph, &node_type) {
continue;
}
let params = string_params(&[("type", &node_type), ("edge", rel)]);
rows.extend(execute_missing_required_edge(graph, ¶ms, yield_items)?);
}
Ok(rows)
}
DeclaredCheck::Cardinality => {
let Some(domain) = decl.domain.as_deref() else {
return Ok(Vec::new());
};
let card = decl.cardinality.expect("gated on presence");
let mut rows = Vec::new();
for node_type in accepted_types(store, domain) {
if !live(graph, &node_type) || !edge_type_exists(graph, rel) {
continue;
}
let mut params = string_params(&[("type", &node_type), ("edge", rel)]);
if let Some(min) = card.min {
params.insert("min".to_string(), Value::Int64(min as i64));
}
if let Some(max) = card.max {
params.insert("max".to_string(), Value::Int64(max as i64));
}
rows.extend(execute_cardinality_violation(graph, ¶ms, yield_items)?);
}
Ok(rows)
}
DeclaredCheck::Inverse | DeclaredCheck::Symmetric => {
let other = if check == DeclaredCheck::Symmetric {
rel
} else {
decl.inverse_name.as_deref().expect("gated on presence")
};
if !edge_type_exists(graph, rel) {
return Ok(Vec::new());
}
let params = string_params(&[("rel_a", rel), ("rel_b", other)]);
execute_inverse_violation(graph, ¶ms, yield_items)
}
DeclaredCheck::Transitive => {
if !edge_type_exists(graph, rel) {
return Ok(Vec::new());
}
let params = string_params(&[("rel", rel)]);
execute_transitivity_violation(graph, ¶ms, yield_items)
}
}
}
fn proc_check(proc_name: &str) -> Option<&'static [DeclaredCheck]> {
match proc_name {
"type_domain_violation" => Some(&[DeclaredCheck::Domain]),
"type_range_violation" => Some(&[DeclaredCheck::Range]),
"missing_required_edge" => Some(&[DeclaredCheck::Required]),
"cardinality_violation" => Some(&[DeclaredCheck::Cardinality]),
"inverse_violation" => Some(&[DeclaredCheck::Inverse, DeclaredCheck::Symmetric]),
"transitivity_violation" => Some(&[DeclaredCheck::Transitive]),
_ => None,
}
}
pub(super) fn no_arg_declaration_rows(
proc_name: &str,
graph: &DirGraph,
yield_items: &[YieldItem],
) -> Option<Result<Vec<ResultRow>, String>> {
let checks = proc_check(proc_name)?;
if graph.ontology.is_empty() {
return None;
}
let mut all = Vec::new();
for (rel, decl) in &graph.ontology.relationships {
for check in checks {
if !declared_checks(rel, decl).contains(check) {
continue;
}
match check_rows(graph, rel, decl, *check, yield_items) {
Ok(mut rows) => {
stamp_rule(&mut rows, yield_items, &format!("{rel}.{}", check.name()));
all.extend(rows);
}
Err(e) => return Some(Err(e)),
}
}
}
Some(Ok(all))
}
pub(super) fn execute_ontology_audit(
graph: &DirGraph,
params: &HashMap<String, Value>,
yield_items: &[YieldItem],
) -> Result<Vec<ResultRow>, String> {
if !params.is_empty() {
return Err("CALL ontology_audit takes no parameters".to_string());
}
let alias = |name: &str| {
yield_items
.iter()
.find(|y| y.name == name)
.map(|y| y.alias.clone().unwrap_or_else(|| name.to_string()))
};
let mut out = Vec::new();
for line in audit_counts(graph)? {
let mut row = ResultRow::new();
if let Some(a) = alias("rule") {
row.projected.insert(a, Value::String(line.rule));
}
if let Some(a) = alias("severity") {
row.projected
.insert(a, Value::String(line.severity.as_str().to_string()));
}
if let Some(a) = alias("violations") {
row.projected
.insert(a, Value::Int64(line.violations as i64));
}
if let Some(a) = alias("total") {
row.projected.insert(a, Value::Int64(line.total as i64));
}
if let Some(a) = alias("pct") {
row.projected.insert(a, Value::Float64(line.pct));
}
out.push(row);
}
Ok(out)
}
pub(crate) struct AuditLine {
pub(crate) rule: String,
pub(crate) severity: crate::graph::ontology::Enforcement,
pub(crate) violations: usize,
pub(crate) total: usize,
pub(crate) pct: f64,
}
pub(crate) fn audit_counts(graph: &DirGraph) -> Result<Vec<AuditLine>, String> {
if graph.ontology.is_empty() {
return Err(
"ontology_audit: no ontology declared — define one with define_ontology()".to_string(),
);
}
let mut out = Vec::new();
for (rel, decl) in &graph.ontology.relationships {
for check in declared_checks(rel, decl) {
let count_yield: Vec<YieldItem> = counting_yield(check);
let violations = check_rows(graph, rel, decl, check, &count_yield)?.len();
let total = check_total(graph, rel, decl, check);
let pct = if total == 0 {
0.0
} else {
((violations as f64 / total as f64 * 100.0) * 10.0).round() / 10.0
};
out.push(AuditLine {
rule: format!("{rel}.{}", check.name()),
severity: decl.enforcement,
violations,
total,
pct,
});
}
}
Ok(out)
}
fn counting_yield(check: DeclaredCheck) -> Vec<YieldItem> {
let names: &[&str] = match check {
DeclaredCheck::Domain | DeclaredCheck::Range => &["source", "target"],
DeclaredCheck::Required => &["node"],
DeclaredCheck::Cardinality => &["node", "count"],
DeclaredCheck::Inverse | DeclaredCheck::Symmetric => &["a", "b"],
DeclaredCheck::Transitive => &["a", "b", "c"],
};
names
.iter()
.map(|n| YieldItem {
name: n.to_string(),
alias: None,
})
.collect()
}
fn check_total(
graph: &DirGraph,
rel: &str,
decl: &RelationshipDecl,
check: DeclaredCheck,
) -> usize {
match check {
DeclaredCheck::Domain
| DeclaredCheck::Range
| DeclaredCheck::Inverse
| DeclaredCheck::Symmetric
| DeclaredCheck::Transitive => {
let key = InternedKey::from_str(rel);
graph
.graph
.edge_references()
.filter(|er| er.weight().connection_type == key)
.count()
}
DeclaredCheck::Required | DeclaredCheck::Cardinality => {
let Some(domain) = decl.domain.as_deref() else {
return 0;
};
accepted_types(&graph.ontology, domain)
.iter()
.filter_map(|t| type_indices(graph, t).ok())
.map(|nodes| nodes.iter().count())
.sum()
}
}
}