use std::collections::{BTreeMap, BTreeSet};
use crate::component_category::CategoryPredicates;
use crate::{ComponentGraph, Edge, Error, Node};
use super::predicates::{
is_measurable_component, parent_meters, reached_only_through, reaches_any_below,
};
#[derive(Debug, PartialEq)]
pub(super) enum Measurement {
Single(u64),
Diamond {
components: Vec<u64>,
meters: Vec<u64>,
},
Subtraction {
parent_meters: Vec<u64>,
subtracted: Vec<u64>,
components: Vec<u64>,
},
}
fn drop_subsumed(points: &mut Vec<Measurement>, subsumed: &[u64]) {
points
.retain(|point| !matches!(point, Measurement::Single(single) if subsumed.contains(single)));
}
pub(super) fn measurement_points<N: Node, E: Edge>(
graph: &ComponentGraph<N, E>,
targets: &BTreeSet<u64>,
) -> Result<Vec<Measurement>, Error> {
let mut remaining = targets.clone();
let mut points = Vec::new();
let mut claimed = BTreeSet::new();
let mut groups = BTreeMap::new();
while let Some(id) = remaining.pop_first() {
let group = match classify(graph, id, targets, &mut groups)? {
Some(group)
if group.subtracted.is_none()
|| group.meters.iter().all(|meter| !claimed.contains(meter)) =>
{
group
}
_ => {
if claimed.insert(id) {
points.push(Measurement::Single(id));
}
continue;
}
};
for &component in &group.components {
remaining.remove(&component);
}
drop_subsumed(&mut points, &group.components);
if let Some(subtracted) = group.subtracted {
claimed.extend(&group.meters);
points.push(Measurement::Subtraction {
parent_meters: group.meters,
subtracted,
components: group.components,
});
} else if group.meters.len() > 1 {
claimed.extend(&group.meters);
points.push(Measurement::Diamond {
components: group.components,
meters: group.meters,
});
} else {
for meter in group.meters {
if claimed.insert(meter) {
points.push(Measurement::Single(meter));
}
}
}
}
Ok(points)
}
fn group_components(seed: u64, siblings: Vec<u64>) -> Vec<u64> {
let mut components: Vec<u64> = std::iter::once(seed).chain(siblings).collect();
components.sort_unstable();
components
}
#[derive(Clone)]
struct Group {
meters: Vec<u64>,
components: Vec<u64>,
subtracted: Option<Vec<u64>>,
}
fn classify<N: Node, E: Edge>(
graph: &ComponentGraph<N, E>,
seed: u64,
targets: &BTreeSet<u64>,
cache: &mut BTreeMap<BTreeSet<u64>, Option<Group>>,
) -> Result<Option<Group>, Error> {
let Some(meters) = parent_meters(graph, seed)? else {
return Ok(None);
};
if let Some(group) = cache.get(&meters) {
return Ok(group.clone());
}
let group = classify_for_parents(graph, seed, targets, &meters)?;
cache.insert(meters, group.clone());
Ok(group)
}
fn classify_for_parents<N: Node, E: Edge>(
graph: &ComponentGraph<N, E>,
seed: u64,
targets: &BTreeSet<u64>,
meters: &BTreeSet<u64>,
) -> Result<Option<Group>, Error> {
let mut covered = Vec::new();
let mut covered_has_meter = false;
let mut subtracted = Vec::new();
for sibling in graph.siblings_from_predecessors(seed)? {
let sibling_id = sibling.component_id();
if targets.contains(&sibling_id) {
covered_has_meter |= sibling.is_meter();
covered.push(sibling_id);
} else if (sibling.is_meter() || is_measurable_component(sibling, &graph.config))
&& sibling.provides_telemetry()
{
subtracted.push(sibling_id);
} else {
return Ok(None);
}
}
let is_subtraction = !subtracted.is_empty();
if !is_subtraction {
if covered.iter().any(|sibling| meters.contains(sibling)) {
return Ok(None);
}
} else {
if covered_has_meter {
return Ok(None);
}
for &other in &subtracted {
if !reached_only_through(graph, other, meters)? {
return Ok(None);
}
if reaches_any_below(graph, other, targets)? {
return Ok(None);
}
}
let subtracted_set: BTreeSet<u64> = subtracted.iter().copied().collect();
let mut redundant = Vec::new();
for &other in &subtracted {
if !reaches_any_below(graph, other, &subtracted_set)? {
continue;
}
if graph.has_successors(other)?
&& graph
.successors(other)?
.all(|child| subtracted_set.contains(&child.component_id()))
{
redundant.push(other);
} else {
return Ok(None);
}
}
subtracted.retain(|other| !redundant.contains(other));
subtracted.sort_unstable();
}
if !group_reached_only_through(graph, seed, &covered, meters)? {
return Ok(None);
}
Ok(Some(Group {
meters: meters.iter().copied().collect(),
components: group_components(seed, covered),
subtracted: is_subtraction.then_some(subtracted),
}))
}
fn group_reached_only_through<N: Node, E: Edge>(
graph: &ComponentGraph<N, E>,
id: u64,
siblings: &[u64],
meters: &BTreeSet<u64>,
) -> Result<bool, Error> {
for &component in std::iter::once(&id).chain(siblings) {
if !reached_only_through(graph, component, meters)? {
return Ok(false);
}
}
Ok(true)
}