use crate::def::{default_max_edges, evaluate, NodeView, Predicate, RuleDef};
use core_storage::{list_tokens, Value, ValueKey};
use serde::Serialize;
use std::collections::{BTreeMap, BTreeSet};
use std::time::{Duration, Instant};
pub const DEFAULT_SEED: u64 = 0x4d75_7368_726f_6f6d;
pub const LOW_CARDINALITY_MAX: usize = 20;
pub const VECTOR_SIMILAR_MIN: f64 = 0.8;
pub const VECTOR_APPROX_THRESHOLD: usize = 2_000;
#[derive(Debug, Clone)]
pub struct SuggestConfig {
pub max_sample_nodes: usize,
pub max_sample_sources: usize,
pub max_examples: usize,
pub budget_ms: u64,
pub global_budget_ms: u64,
}
impl Default for SuggestConfig {
fn default() -> Self {
Self {
max_sample_nodes: 10_000,
max_sample_sources: 200,
max_examples: 3,
budget_ms: 250,
global_budget_ms: 5_000,
}
}
}
#[derive(Debug, Clone, Serialize)]
pub struct SuggestReport {
pub suggestions: Vec<RuleSuggestion>,
pub truncated: bool,
}
#[derive(Debug, Clone, Serialize)]
pub struct RuleSuggestion {
pub def: RuleDef,
pub est_edges: u64,
pub examples: Vec<(String, String, f64)>,
pub rationale: String,
}
#[inline]
fn lcg_step(state: &mut u64) -> u64 {
*state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
*state
}
fn sample_indices(n: usize, k: usize, seed: u64) -> Vec<usize> {
if n == 0 {
return Vec::new();
}
let take = k.min(n);
let mut indices: Vec<usize> = (0..n).collect();
let mut rng = seed;
for i in 0..take {
let r = lcg_step(&mut rng);
let j = i + (r as usize % (n - i));
indices.swap(i, j);
}
indices[..take].to_vec()
}
fn as_float_val(v: &Value) -> Option<f64> {
match v {
Value::Int(i) => Some(*i as f64),
Value::Float(f) if f.is_finite() => Some(*f),
_ => None,
}
}
fn as_float_list(v: &Value) -> Option<Vec<f64>> {
let Value::List(items) = v else {
return None;
};
if items.is_empty() {
return None;
}
items.iter().map(as_float_val).collect()
}
#[derive(Default)]
struct FieldProfile {
present: usize,
str_distinct: BTreeSet<String>,
numeric_vals: Vec<f64>,
list_tokens: Vec<(u32, BTreeSet<ValueKey>)>,
vec_entries: Vec<(u32, usize)>,
}
fn profile_label(
nodes: &[(u32, String)],
get_prop: &dyn Fn(u32, &str) -> Option<Value>,
all_fields: &[String],
max_sample: usize,
seed: u64,
) -> BTreeMap<String, FieldProfile> {
let sample = sample_indices(nodes.len(), max_sample, seed);
let mut profiles: BTreeMap<String, FieldProfile> = BTreeMap::new();
for si in sample {
let (node_id, _) = &nodes[si];
for field in all_fields {
let Some(val) = get_prop(*node_id, field) else {
continue;
};
let p = profiles.entry(field.clone()).or_default();
p.present += 1;
match &val {
Value::Str(s) => {
p.str_distinct.insert(s.clone());
}
Value::Int(_) | Value::Float(_) => {
if let Some(f) = as_float_val(&val) {
p.numeric_vals.push(f);
}
}
Value::List(_) => {
if let Some(fvec) = as_float_list(&val) {
p.vec_entries.push((*node_id, fvec.len()));
} else if let Some(toks) = list_tokens(&val) {
p.list_tokens.push((*node_id, toks));
}
}
_ => {}
}
}
}
profiles
}
fn dominant_dim(entries: &[(u32, usize)]) -> Option<usize> {
if entries.is_empty() {
return None;
}
let mut counts: BTreeMap<usize, usize> = BTreeMap::new();
for (_, dim) in entries {
*counts.entry(*dim).or_default() += 1;
}
let total = entries.len();
counts
.into_iter()
.find(|&(_, count)| count * 10 >= total * 8)
.map(|(dim, _)| dim)
}
fn is_covered(existing: &[RuleDef], src_label: &str, dst_label: &str, pred: &Predicate) -> bool {
existing.iter().any(|r| {
r.src_label == src_label
&& r.dst_label == dst_label
&& same_pred_kind_field(&r.predicate, pred)
})
}
fn same_pred_kind_field(a: &Predicate, b: &Predicate) -> bool {
match (a, b) {
(Predicate::KeyMatch { field: fa }, Predicate::KeyMatch { field: fb }) => fa == fb,
(Predicate::FieldEqual { field: fa }, Predicate::FieldEqual { field: fb }) => fa == fb,
(Predicate::Overlap { field: fa, .. }, Predicate::Overlap { field: fb, .. }) => fa == fb,
(
Predicate::NumericWithin { field: fa, .. },
Predicate::NumericWithin { field: fb, .. },
) => fa == fb,
(
Predicate::VectorSimilar { field: fa, .. },
Predicate::VectorSimilar { field: fb, .. },
) => fa == fb,
_ => false,
}
}
struct Preview {
est_edges: u64,
examples: Vec<(String, String, f64)>,
}
fn run_preview(
def: &RuleDef,
src_nodes: &[(u32, String)],
dst_nodes: &[(u32, String)],
get_prop: &dyn Fn(u32, &str) -> Option<Value>,
config: &SuggestConfig,
) -> Preview {
let src_n = src_nodes.len();
let dst_n = dst_nodes.len();
if src_n == 0 || dst_n == 0 {
return Preview {
est_edges: 0,
examples: Vec::new(),
};
}
let seed = def.name.bytes().fold(DEFAULT_SEED, |acc, b| {
acc.wrapping_mul(31).wrapping_add(b as u64)
});
let src_sample = sample_indices(src_n, config.max_sample_sources, seed);
let deadline = Instant::now() + Duration::from_millis(config.budget_ms);
let mut hit_edges = 0u64;
let mut examples: Vec<(String, String, f64)> = Vec::new();
let mut processed = 0usize;
'outer: for &si in &src_sample {
if Instant::now() >= deadline {
break;
}
let (src_id, src_key) = &src_nodes[si];
let sp = |f: &str| get_prop(*src_id, f);
let src_view = NodeView {
key: src_key.as_str(),
props: &sp,
};
let mut src_hits = 0u64;
for (dst_id, dst_key) in dst_nodes {
if src_key == dst_key {
continue; }
let dp = |f: &str| get_prop(*dst_id, f);
let dst_view = NodeView {
key: dst_key.as_str(),
props: &dp,
};
if let Some(score) = evaluate(&def.predicate, &src_view, &dst_view) {
src_hits += 1;
if examples.len() < config.max_examples {
examples.push((src_key.clone(), dst_key.clone(), score));
}
}
}
let kept = match def.max_edges {
Some(k) => src_hits.min(k),
None => src_hits,
};
hit_edges += kept;
processed += 1;
if Instant::now() >= deadline {
break 'outer;
}
}
let est_edges = if processed == 0 {
0
} else {
let avg_kept = hit_edges as f64 / processed as f64;
let raw = (avg_kept * src_n as f64).round() as u64;
match def.max_edges {
Some(k) => raw.min(k.saturating_mul(src_n as u64)),
None => raw,
}
};
Preview {
est_edges,
examples,
}
}
pub fn suggest_rules(
label_nodes: &BTreeMap<String, Vec<(u32, String)>>,
get_prop: &dyn Fn(u32, &str) -> Option<Value>,
all_fields: &[String],
existing: &[RuleDef],
config: &SuggestConfig,
seed: u64,
) -> SuggestReport {
if label_nodes.is_empty() || all_fields.is_empty() {
return SuggestReport {
suggestions: Vec::new(),
truncated: false,
};
}
let global_deadline = Instant::now() + Duration::from_millis(config.global_budget_ms);
let label_keys: BTreeMap<&str, BTreeSet<&str>> = label_nodes
.iter()
.map(|(label, nodes)| {
let keys: BTreeSet<&str> = nodes.iter().map(|(_, k)| k.as_str()).collect();
(label.as_str(), keys)
})
.collect();
let mut profiling_truncated = false;
let profiles: BTreeMap<String, BTreeMap<String, FieldProfile>> = label_nodes
.iter()
.enumerate()
.filter_map(|(i, (label, nodes))| {
if Instant::now() >= global_deadline {
profiling_truncated = true;
return None;
}
let label_seed = seed.wrapping_add(i as u64 ^ 0x9e37_79b9_7f4a_7c15);
let p = profile_label(
nodes,
get_prop,
all_fields,
config.max_sample_nodes,
label_seed,
);
Some((label.clone(), p))
})
.collect();
let labels: Vec<&str> = label_nodes.keys().map(String::as_str).collect();
let mut results: Vec<RuleSuggestion> = Vec::new();
let mut truncated = false;
'detect: {
for src_label in &labels {
let Some(src_profile) = profiles.get(*src_label) else {
continue;
};
let src_nodes = &label_nodes[*src_label];
for (field, fp) in src_profile {
if !field.ends_with("_id") || fp.str_distinct.is_empty() {
continue;
}
for dst_label in &labels {
let Some(dst_keys) = label_keys.get(dst_label) else {
continue;
};
let match_count = fp
.str_distinct
.iter()
.filter(|v| dst_keys.contains(v.as_str()))
.count();
if match_count == 0 {
continue;
}
let pred = Predicate::KeyMatch {
field: field.clone(),
};
if is_covered(existing, src_label, dst_label, &pred) {
continue;
}
if Instant::now() >= global_deadline {
truncated = true;
break 'detect;
}
let base = field.trim_end_matches("_id").to_uppercase();
let name = format!(
"suggest_km_{}_{}_{field}",
src_label.to_lowercase(),
dst_label.to_lowercase(),
);
let max_edges = Some(default_max_edges(&pred));
let def = RuleDef {
name,
src_label: src_label.to_string(),
dst_label: dst_label.to_string(),
predicate: pred,
edge_type: format!("{base}_OF"),
weight_prop: None,
max_edges,
approximate: false,
via_label: None,
via_edge: None,
via_dir: None,
};
let examples_preview: Vec<String> = fp
.str_distinct
.iter()
.filter(|v| dst_keys.contains(v.as_str()))
.take(3)
.cloned()
.collect();
let rationale = format!(
"Field '{field}' in {src_label} ends with '_id' and {match_count} \
sampled value(s) match keys in {dst_label} \
(e.g. {}). Suggests a foreign-key relationship.",
examples_preview.join(", ")
);
let preview =
run_preview(&def, src_nodes, &label_nodes[*dst_label], get_prop, config);
results.push(RuleSuggestion {
def,
est_edges: preview.est_edges,
examples: preview.examples,
rationale,
});
}
}
}
for (si, src_label) in labels.iter().enumerate() {
let Some(src_profile) = profiles.get(*src_label) else {
continue;
};
let src_nodes = &label_nodes[*src_label];
for (di, dst_label) in labels.iter().enumerate() {
if di < si {
continue; }
let Some(dst_profile) = profiles.get(*dst_label) else {
continue;
};
let dst_nodes = &label_nodes[*dst_label];
for field in all_fields {
let Some(src_fp) = src_profile.get(field) else {
continue;
};
let Some(dst_fp) = dst_profile.get(field) else {
continue;
};
if src_fp.list_tokens.is_empty() || dst_fp.list_tokens.is_empty() {
continue;
}
let n_src_toks = src_fp.list_tokens.len();
let n_dst_toks = dst_fp.list_tokens.len();
let n_pairs = 200.min(n_src_toks * n_dst_toks);
let mut rng = seed
.wrapping_add(0xAB_CD_EF_01u64)
.wrapping_add(si as u64 * 0x1111)
.wrapping_add(di as u64 * 0x2222)
.wrapping_add(field.len() as u64 * 0x3333);
let mut jaccards: Vec<f64> = Vec::with_capacity(n_pairs);
for _ in 0..n_pairs {
let si2 = lcg_step(&mut rng) as usize % n_src_toks;
let di2 = lcg_step(&mut rng) as usize % n_dst_toks;
let (_, src_toks) = &src_fp.list_tokens[si2];
let (_, dst_toks) = &dst_fp.list_tokens[di2];
let inter = src_toks.intersection(dst_toks).count();
let union = src_toks.union(dst_toks).count();
if union > 0 {
jaccards.push(inter as f64 / union as f64);
}
}
if jaccards.is_empty() {
continue;
}
jaccards.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let p50 = jaccards[jaccards.len() / 2];
if p50 <= 0.0 {
continue;
}
let min_val = ((p50 * 100.0).round() / 100.0).clamp(0.01, 1.0);
let pred = Predicate::Overlap {
field: field.clone(),
min: min_val,
};
if is_covered(existing, src_label, dst_label, &pred) {
continue;
}
if Instant::now() >= global_deadline {
truncated = true;
break 'detect;
}
let name = format!(
"suggest_ov_{}_{}_{field}",
src_label.to_lowercase(),
dst_label.to_lowercase(),
);
let max_edges = Some(default_max_edges(&pred));
let def = RuleDef {
name,
src_label: src_label.to_string(),
dst_label: dst_label.to_string(),
predicate: pred,
edge_type: format!("OVERLAPS_{}", field.to_uppercase()),
weight_prop: Some("score".into()),
max_edges,
approximate: false,
via_label: None,
via_edge: None,
via_dir: None,
};
let rationale = format!(
"Field '{field}' is a token list in both {src_label} and {dst_label}. \
Sampled Jaccard p50={p50:.2}; using that as the minimum threshold \
(min={min_val:.2}). Lists share common tokens suggesting semantic affinity."
);
let preview = run_preview(&def, src_nodes, dst_nodes, get_prop, config);
results.push(RuleSuggestion {
def,
est_edges: preview.est_edges,
examples: preview.examples,
rationale,
});
}
}
}
for (si, src_label) in labels.iter().enumerate() {
let Some(src_profile) = profiles.get(*src_label) else {
continue;
};
let src_nodes = &label_nodes[*src_label];
for (di, dst_label) in labels.iter().enumerate() {
if di < si {
continue;
}
let Some(dst_profile) = profiles.get(*dst_label) else {
continue;
};
let dst_nodes = &label_nodes[*dst_label];
for field in all_fields {
let Some(src_fp) = src_profile.get(field) else {
continue;
};
let Some(dst_fp) = dst_profile.get(field) else {
continue;
};
if src_fp.str_distinct.is_empty() || dst_fp.str_distinct.is_empty() {
continue;
}
if src_fp.str_distinct.len() > LOW_CARDINALITY_MAX
|| dst_fp.str_distinct.len() > LOW_CARDINALITY_MAX
{
continue;
}
let shared = src_fp
.str_distinct
.intersection(&dst_fp.str_distinct)
.count();
if shared == 0 {
continue;
}
let pred = Predicate::FieldEqual {
field: field.clone(),
};
if is_covered(existing, src_label, dst_label, &pred) {
continue;
}
if Instant::now() >= global_deadline {
truncated = true;
break 'detect;
}
let name = format!(
"suggest_fe_{}_{}_{field}",
src_label.to_lowercase(),
dst_label.to_lowercase(),
);
let max_edges = Some(default_max_edges(&pred));
let def = RuleDef {
name,
src_label: src_label.to_string(),
dst_label: dst_label.to_string(),
predicate: pred,
edge_type: format!("SAME_{}", field.to_uppercase()),
weight_prop: None,
max_edges,
approximate: false,
via_label: None,
via_edge: None,
via_dir: None,
};
let rationale = format!(
"Field '{field}' has low cardinality in {src_label} \
({} distinct value(s)) and {dst_label} ({} distinct value(s)), \
with {shared} shared value(s). Suggests a categorical grouping predicate.",
src_fp.str_distinct.len(),
dst_fp.str_distinct.len(),
);
let preview = run_preview(&def, src_nodes, dst_nodes, get_prop, config);
results.push(RuleSuggestion {
def,
est_edges: preview.est_edges,
examples: preview.examples,
rationale,
});
}
}
}
for (si, src_label) in labels.iter().enumerate() {
let Some(src_profile) = profiles.get(*src_label) else {
continue;
};
let src_nodes = &label_nodes[*src_label];
for (di, dst_label) in labels.iter().enumerate() {
if di < si {
continue;
}
let Some(dst_profile) = profiles.get(*dst_label) else {
continue;
};
let dst_nodes = &label_nodes[*dst_label];
for field in all_fields {
let Some(src_fp) = src_profile.get(field) else {
continue;
};
let Some(dst_fp) = dst_profile.get(field) else {
continue;
};
if src_fp.numeric_vals.is_empty() || dst_fp.numeric_vals.is_empty() {
continue;
}
let src_min = src_fp
.numeric_vals
.iter()
.cloned()
.fold(f64::INFINITY, f64::min);
let src_max = src_fp
.numeric_vals
.iter()
.cloned()
.fold(f64::NEG_INFINITY, f64::max);
let dst_min = dst_fp
.numeric_vals
.iter()
.cloned()
.fold(f64::INFINITY, f64::min);
let dst_max = dst_fp
.numeric_vals
.iter()
.cloned()
.fold(f64::NEG_INFINITY, f64::max);
if src_max < dst_min || dst_max < src_min {
continue;
}
let combined_min = src_min.min(dst_min);
let combined_max = src_max.max(dst_max);
let spread = combined_max - combined_min;
if !spread.is_finite() || spread <= 0.0 {
continue;
}
let tolerance = (spread / 4.0).max(1.0);
let pred = Predicate::NumericWithin {
field: field.clone(),
tolerance,
};
if is_covered(existing, src_label, dst_label, &pred) {
continue;
}
if Instant::now() >= global_deadline {
truncated = true;
break 'detect;
}
let name = format!(
"suggest_nw_{}_{}_{field}",
src_label.to_lowercase(),
dst_label.to_lowercase(),
);
let max_edges = Some(default_max_edges(&pred));
let def = RuleDef {
name,
src_label: src_label.to_string(),
dst_label: dst_label.to_string(),
predicate: pred,
edge_type: format!("NEAR_{}", field.to_uppercase()),
weight_prop: Some("score".into()),
max_edges,
approximate: false,
via_label: None,
via_edge: None,
via_dir: None,
};
let rationale = format!(
"Field '{field}' is numeric in {src_label} (range [{src_min:.2}, {src_max:.2}]) \
and {dst_label} (range [{dst_min:.2}, {dst_max:.2}]); ranges overlap. \
Tolerance {tolerance:.2} derived from combined spread {spread:.2}."
);
let preview = run_preview(&def, src_nodes, dst_nodes, get_prop, config);
results.push(RuleSuggestion {
def,
est_edges: preview.est_edges,
examples: preview.examples,
rationale,
});
}
}
}
for (si, src_label) in labels.iter().enumerate() {
let Some(src_profile) = profiles.get(*src_label) else {
continue;
};
let src_nodes = &label_nodes[*src_label];
for (di, dst_label) in labels.iter().enumerate() {
if di < si {
continue;
}
let Some(dst_profile) = profiles.get(*dst_label) else {
continue;
};
let dst_nodes = &label_nodes[*dst_label];
for field in all_fields {
let Some(src_fp) = src_profile.get(field) else {
continue;
};
let Some(dst_fp) = dst_profile.get(field) else {
continue;
};
if src_fp.vec_entries.is_empty() || dst_fp.vec_entries.is_empty() {
continue;
}
let src_dim = dominant_dim(&src_fp.vec_entries);
let dst_dim = dominant_dim(&dst_fp.vec_entries);
let (Some(sdim), Some(ddim)) = (src_dim, dst_dim) else {
continue;
};
if sdim != ddim || sdim == 0 {
continue;
}
let approximate = dst_nodes.len() > VECTOR_APPROX_THRESHOLD;
let pred = Predicate::VectorSimilar {
field: field.clone(),
min: VECTOR_SIMILAR_MIN,
};
if is_covered(existing, src_label, dst_label, &pred) {
continue;
}
if Instant::now() >= global_deadline {
truncated = true;
break 'detect;
}
let name = format!(
"suggest_vs_{}_{}_{field}",
src_label.to_lowercase(),
dst_label.to_lowercase(),
);
let max_edges = Some(default_max_edges(&pred));
let def = RuleDef {
name,
src_label: src_label.to_string(),
dst_label: dst_label.to_string(),
predicate: pred,
edge_type: format!("SIMILAR_{}", field.to_uppercase()),
weight_prop: Some("score".into()),
max_edges,
approximate,
via_label: None,
via_edge: None,
via_dir: None,
};
let rationale = format!(
"Field '{field}' is a float-array of dim {sdim} in both {src_label} \
and {dst_label}. Suggests embedding-based similarity (min={VECTOR_SIMILAR_MIN}){}.",
if approximate {
", approximate=true suggested (n>2000)"
} else {
""
}
);
let preview = run_preview(&def, src_nodes, dst_nodes, get_prop, config);
results.push(RuleSuggestion {
def,
est_edges: preview.est_edges,
examples: preview.examples,
rationale,
});
}
}
}
}
results.sort_by(|a, b| {
b.est_edges
.cmp(&a.est_edges)
.then(a.def.name.cmp(&b.def.name))
});
SuggestReport {
suggestions: results,
truncated: truncated || profiling_truncated,
}
}