diff --git a/docs/code/maad-conformance.md b/docs/code/maad-conformance.md index df36305c..00977579 100644 --- a/docs/code/maad-conformance.md +++ b/docs/code/maad-conformance.md @@ -115,6 +115,15 @@ Add `--ipv6` to compare IPv6 files. Add another file as another positional different tolerance. A passing case prints one compact summary; command, input, JSON, row-count, metadata, or numeric mismatches return nonzero. +Pass `--weighted` to compare measure-weighted MAAD. Each case is then an +`ADDR,MEASURE` CSV with one row per distinct address and a finite, positive +measure: the oracle keeps only the first row for a repeated address, while Rust +sums them, so the validator rejects duplicates. Rust runs `netflow-db maad +--weighted` and the oracle runs with `--csv --meas-col 1`. Only structure and +dimensions are compared, because the weighted estimator does not emit a +spectrum. Combine it with `--ipv6` for IPv6 CSVs. Every summary line reports +`max_abs_dtau` and `max_abs_ddim`. + ## Known edge cases The Haskell executable cannot produce JSON for an empty set, a singleton, or a diff --git a/scripts/local/validate_maad.py b/scripts/local/validate_maad.py index 63c85925..6e04d626 100755 --- a/scripts/local/validate_maad.py +++ b/scripts/local/validate_maad.py @@ -47,35 +47,68 @@ def parse_case_spec(raw: str) -> Case: return Case(name=name, path=path) -def validate_input(path: str, ipv6: bool) -> None: - """Check each non-empty line without rewriting or materializing the file.""" - +def parse_address( + value: str, path: str, line_number: int, ipv6: bool +) -> ipaddress.IPv4Address | ipaddress.IPv6Address: family = "IPv6" if ipv6 else "IPv4" expected_type = ipaddress.IPv6Address if ipv6 else ipaddress.IPv4Address + try: + address = ipaddress.ip_address(value) + except ValueError as error: + raise CaseError( + f"{path}: line {line_number}: invalid {family} address {value!r}" + ) from error + if not isinstance(address, expected_type): + raise CaseError(f"{path}: line {line_number}: expected {family}, got {value!r}") + if ipv6 and "." in value: + hextets = ":".join( + f"{(int(address) >> shift) & 0xFFFF:x}" for shift in range(112, -1, -16) + ) + raise CaseError( + f"{path}: line {line_number}: the oracle misreads embedded IPv4 " + f"notation; write {value!r} as {hextets!r}" + ) + return address + + +def parse_measure(value: str, path: str, line_number: int) -> float: + try: + measure = float(value) + except ValueError as error: + raise CaseError(f"{path}: line {line_number}: invalid measure {value!r}") from error + if not math.isfinite(measure) or measure <= 0: + raise CaseError(f"{path}: line {line_number}: measure must be finite and positive") + return measure + + +def validate_input(path: str, ipv6: bool, weighted: bool) -> None: + """Check each non-empty line without rewriting the file. + + Weighted inputs must hold one ADDR,MEASURE row per distinct address: the oracle keeps + only the first row for a repeated address while Rust sums them. + """ + + seen: set[int] = set() try: with open(path, "r", encoding="ascii", errors="strict", newline="") as stream: for line_number, raw_line in enumerate(stream, start=1): value = raw_line.strip() if not value: + family = "IPv6" if ipv6 else "IPv4" raise CaseError(f"{path}: line {line_number}: expected an {family} address") - try: - address = ipaddress.ip_address(value) - except ValueError as error: - raise CaseError( - f"{path}: line {line_number}: invalid {family} address {value!r}" - ) from error - if not isinstance(address, expected_type): - raise CaseError( - f"{path}: line {line_number}: expected {family}, got {value!r}" - ) - if ipv6 and "." in value: - hextets = ":".join( - f"{(int(address) >> shift) & 0xFFFF:x}" for shift in range(112, -1, -16) - ) + if not weighted: + parse_address(value, path, line_number, ipv6) + continue + fields = value.split(",") + if len(fields) != 2: + raise CaseError(f"{path}: line {line_number}: expected ADDR,MEASURE") + address = int(parse_address(fields[0].strip(), path, line_number, ipv6)) + parse_measure(fields[1].strip(), path, line_number) + if address in seen: raise CaseError( - f"{path}: line {line_number}: the oracle misreads embedded IPv4 " - f"notation; write {value!r} as {hextets!r}" + f"{path}: line {line_number}: duplicate address {fields[0]!r}" ) + seen.add(address) except CaseError: raise except (OSError, UnicodeError) as error: @@ -208,11 +241,30 @@ def compare_rows( return errors +def max_deviation( + rust: dict[str, Any], haskell: dict[str, Any], section: str, field: str +) -> float: + rust_rows = rust.get(section) + haskell_rows = haskell.get(section) + if not isinstance(rust_rows, list) or not isinstance(haskell_rows, list): + return math.nan + deviation = 0.0 + for index, (rust_row, haskell_row) in enumerate(zip(rust_rows, haskell_rows)): + try: + rust_value = numeric(rust_row.get(field), f"rust {section}[{index}].{field}") + haskell_value = numeric(haskell_row.get(field), f"haskell {section}[{index}].{field}") + except (AttributeError, CaseError): + return math.nan + deviation = max(deviation, abs(rust_value - haskell_value)) + return deviation + + def compare_results( rust: dict[str, Any], haskell: dict[str, Any], absolute: float, relative: float, + weighted: bool, ) -> tuple[list[str], tuple[int, int, int, int, int]]: errors: list[str] = [] @@ -242,6 +294,8 @@ def compare_results( "spectrum": ("alpha", "f"), "dimensions": ("q", "dim", "sd"), } + if weighted: + del section_fields["spectrum"] for section, fields in section_fields.items(): errors.extend(compare_rows(rust, haskell, section, fields, absolute, relative)) if len(errors) >= MAX_ERROR_DETAILS: @@ -278,7 +332,10 @@ def case_specs(parser: argparse.ArgumentParser, args: argparse.Namespace) -> lis def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser( - description="Compare Rust and Haskell MAAD JSON for local IPv4 or IPv6 address files.", + description=( + "Compare Rust and Haskell MAAD JSON for local IPv4 or IPv6 address files, or for " + "ADDR,MEASURE CSV files with --weighted." + ), epilog=( "Cases are NAME=PATH; paths are passed unchanged to both binaries. " "Example: %(prog)s --rust target/release/netflow-db " @@ -287,12 +344,6 @@ def build_parser() -> argparse.ArgumentParser: ) parser.add_argument("--rust", required=True, help="Rust netflow-db binary") parser.add_argument("--haskell", required=True, help="Haskell MAAD binary") - parser.add_argument( - "-6", - "--ipv6", - action="store_true", - help="every case contains IPv6 addresses (default: IPv4)", - ) parser.add_argument( "--abs-tol", "--absolute-tolerance", @@ -309,6 +360,20 @@ def build_parser() -> argparse.ArgumentParser: default=DEFAULT_TOLERANCE, help=f"relative numeric tolerance (default: {DEFAULT_TOLERANCE:g})", ) + parser.add_argument( + "-6", + "--ipv6", + action="store_true", + help="every case contains IPv6 addresses (default: IPv4)", + ) + parser.add_argument( + "--weighted", + action="store_true", + help=( + "cases are ADDR,MEASURE CSV files with one row per distinct address; compare " + "measure-weighted structure and dimensions (no spectrum)" + ), + ) parser.add_argument( "--case", dest="named_cases", @@ -328,34 +393,44 @@ def run_case( absolute: float, relative: float, ipv6: bool, + weighted: bool, ) -> bool: family_args = ["--ipv6"] if ipv6 else [] + rust_command = [ + rust_binary, + "maad", + *family_args, + *(["--weighted"] if weighted else []), + case.path, + ] + haskell_command = [ + haskell_binary, + *family_args, + "--input", + case.path, + "--output", + "-", + "--format", + "json", + "--structure", + "--dimensions", + *(["--csv", "--meas-col", "1"] if weighted else ["--spectrum"]), + ] try: - validate_input(case.path, ipv6) - rust = run_json([rust_binary, "maad", *family_args, case.path], "Rust") - haskell = run_json( - [ - haskell_binary, - *family_args, - "--input", - case.path, - "--output", - "-", - "--format", - "json", - "--structure", - "--spectrum", - "--dimensions", - ], - "Haskell", - ) - errors, counts = compare_results(rust, haskell, absolute, relative) + validate_input(case.path, ipv6, weighted) + rust = run_json(rust_command, "Rust") + haskell = run_json(haskell_command, "Haskell") + errors, counts = compare_results(rust, haskell, absolute, relative, weighted) except CaseError as error: print(f"{case.name}: FAIL {error}") return False + deviations = ( + f"max_abs_dtau={max_deviation(rust, haskell, 'structure', 'tauTilde'):.3g} " + f"max_abs_ddim={max_deviation(rust, haskell, 'dimensions', 'dim'):.3g}" + ) if errors: - print(f"{case.name}: FAIL") + print(f"{case.name}: FAIL {deviations}") for error in errors: print(f" {error}") return False @@ -363,7 +438,8 @@ def run_case( total, structure_rows, spectrum_rows, dimension_rows, prefix_count = counts print( f"{case.name}: PASS total={total} prefixes={prefix_count} " - f"rows=structure:{structure_rows},spectrum:{spectrum_rows},dimensions:{dimension_rows}" + f"rows=structure:{structure_rows},spectrum:{spectrum_rows},dimensions:{dimension_rows} " + f"{deviations}" ) return True @@ -382,6 +458,7 @@ def main(argv: Iterable[str] | None = None) -> int: args.absolute_tolerance, args.relative_tolerance, args.ipv6, + args.weighted, ) and all_passed return 0 if all_passed else 1 diff --git a/tools/netflow-db/src/maad.rs b/tools/netflow-db/src/maad.rs index 07378f3a..5cf5aeb3 100644 --- a/tools/netflow-db/src/maad.rs +++ b/tools/netflow-db/src/maad.rs @@ -1,8 +1,9 @@ //! In-process MAAD-compatible multifractal analysis for IPv4 and IPv6 address sets. +use rayon::prelude::*; use serde::Serialize; use std::io::Write; -use std::net::{Ipv4Addr, Ipv6Addr}; +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; const MIN_MAAD_ADDRESSES: usize = 2; const SCHEMA_VERSION: u32 = 3; @@ -61,7 +62,7 @@ impl PrefixBits for u128 { } /// An address family MAAD can analyze, with its upstream default prefix range. -pub trait MaadAddress: Copy { +pub trait MaadAddress: Copy + Into { type Bits: PrefixBits; fn bits(self) -> Self::Bits; @@ -168,6 +169,10 @@ pub enum MaadError { }, #[error("full_threshold must be finite and in [0, 1) (full_threshold={full_threshold})")] InvalidFullThreshold { full_threshold: f64 }, + #[error("weight for {address} must be finite and positive (weight={weight})")] + InvalidWeight { address: IpAddr, weight: f64 }, + #[error("summed weights must be finite (total={total})")] + NonFiniteTotalWeight { total: f64 }, } #[derive(Clone, Debug, PartialEq, Serialize)] @@ -212,9 +217,9 @@ pub struct DimensionRow { } #[derive(Clone, Debug)] -struct PreparedMoment { - parent_counts: Vec, - child_counts: Vec>, +struct PreparedMoment { + parent_masses: Vec, + child_masses: Vec<(M, Option)>, } /// Compute MAAD-compatible output from one address family's address set. @@ -235,28 +240,102 @@ pub fn compute_with_config( if addresses.len() < MIN_MAAD_ADDRESSES { return Ok(empty_result(addresses.len())); } - let counts = build_prefix_counts(&addresses, config.max_prefix_length + 1); - let prepared = prepare_valid_moments(&counts, &config); - if prepared.is_empty() { - return Ok(empty_result(addresses.len())); - } - let (prefix_lengths, prepared): (Vec<_>, Vec<_>) = prepared.into_iter().unzip(); - let structure = compute_structure(&prepared, &q_values); - let spectrum = compute_spectrum(&structure, config.q_step); - let dimensions = compute_dimensions(&counts, &prefix_lengths, &structure, addresses.len()); - Ok(MaadResult { - schema_version: SCHEMA_VERSION, - metadata: MaadMetadata { - input: "-", - min_prefix_length: prefix_lengths.first().copied(), - max_prefix_length: prefix_lengths.last().copied(), - prefix_lengths, - total_addrs: addresses.len(), - }, - structure, - spectrum, - dimensions, - }) + Ok(analyze(&addresses, &[], &[], &config).address_result(&q_values, config.q_step)) +} + +/// Compute measure-weighted MAAD structure and dimensions from `(address, weight)` pairs. +/// +/// Weights of duplicate addresses are summed. Prefix validity and path pruning use distinct +/// address counts; moments and entropy use summed weights. The spectrum is always empty. +pub fn compute_weighted( + entries: impl IntoIterator, +) -> Result { + compute_weighted_with_config(entries, A::default_config()) +} + +/// Compute measure-weighted MAAD output using an explicitly validated configuration. +pub fn compute_weighted_with_config( + entries: impl IntoIterator, + config: MaadConfig, +) -> Result { + let (analysis, q_values) = analyze_measures( + entries + .into_iter() + .map(|(address, weight)| (address, [weight])), + &config, + )?; + Ok(analysis.weighted_result(0, &q_values)) +} + +/// Compute the distinct-address result and one measure-weighted result per weight column. +/// +/// Every measure shares one sort and one walk over the prefix levels, because prefix validity +/// and path pruning depend only on the distinct addresses. Duplicate addresses sum their +/// weights in input order. Each result equals the one [`compute`] or [`compute_weighted`] +/// returns for the same addresses and weights in the same order. +pub fn compute_measures( + entries: impl IntoIterator, +) -> Result<(MaadResult, [MaadResult; K]), MaadError> { + compute_measures_with_config(entries, A::default_config()) +} + +/// Compute every measure using an explicitly validated configuration. +pub fn compute_measures_with_config( + entries: impl IntoIterator, + config: MaadConfig, +) -> Result<(MaadResult, [MaadResult; K]), MaadError> { + let (analysis, q_values) = analyze_measures(entries, &config)?; + Ok(( + analysis.address_result(&q_values, config.q_step), + std::array::from_fn(|k| analysis.weighted_result(k, &q_values)), + )) +} + +/// Validate, sort and sum weighted entries, then walk their prefix levels once. +fn analyze_measures( + entries: impl IntoIterator, + config: &MaadConfig, +) -> Result<(Analysis, Vec), MaadError> { + let q_values = validate_config(config, A::Bits::WIDTH)?; + let mut entries = entries + .into_iter() + .map(|(address, weights)| { + match weights + .iter() + .find(|weight| !(weight.is_finite() && **weight > 0.0)) + { + Some(&weight) => Err(MaadError::InvalidWeight { + address: address.into(), + weight, + }), + None => Ok((address.bits(), weights)), + } + }) + .collect::, _>>()?; + entries.sort_by_key(|&(address, _)| address); + entries.dedup_by(|next, kept| { + let duplicate = next.0 == kept.0; + if duplicate { + for (kept, next) in kept.1.iter_mut().zip(next.1) { + *kept += next; + } + } + duplicate + }); + let totals: [f64; K] = + std::array::from_fn(|k| entries.iter().map(|(_, weights)| weights[k]).sum::()); + if let Some(&total) = totals.iter().find(|total| !total.is_finite()) { + return Err(MaadError::NonFiniteTotalWeight { total }); + } + if entries.len() < MIN_MAAD_ADDRESSES { + return Ok((Analysis::empty(entries.len(), K), q_values)); + } + let addresses: Vec<_> = entries.iter().map(|&(address, _)| address).collect(); + let columns: Vec> = (0..K) + .map(|k| entries.iter().map(|(_, weights)| weights[k]).collect()) + .collect(); + drop(entries); + Ok((analyze(&addresses, &columns, &totals, config), q_values)) } /// Serialize a computed MAAD result using the established JSON field names. @@ -265,6 +344,88 @@ pub fn write_json(result: &MaadResult, mut output: W) -> Result<(), se output.write_all(b"\n").map_err(serde_json::Error::io) } +/// Prepared moments and level entropies for one measure. +struct MeasureAnalysis { + moments: Vec>, + entropies: Vec, +} + +impl MeasureAnalysis { + const fn new() -> Self { + Self { + moments: Vec::new(), + entropies: Vec::new(), + } + } +} + +/// Every measure's moments over the prefix levels that have a valid parent. +struct Analysis { + total_addrs: usize, + prefix_lengths: Vec, + counts: MeasureAnalysis, + weighted: Vec>, +} + +impl Analysis { + fn empty(total_addrs: usize, measures: usize) -> Self { + Self { + total_addrs, + prefix_lengths: Vec::new(), + counts: MeasureAnalysis::new(), + weighted: (0..measures).map(|_| MeasureAnalysis::new()).collect(), + } + } + + fn address_result(&self, q_values: &[f64], q_step: f64) -> MaadResult { + if self.prefix_lengths.is_empty() { + return empty_result(self.total_addrs); + } + let structure = compute_structure(&self.counts.moments, q_values); + let spectrum = compute_spectrum(&structure, q_step); + let dimensions = self.dimensions(&structure, &self.counts.entropies); + self.result(structure, spectrum, dimensions) + } + + fn weighted_result(&self, measure: usize, q_values: &[f64]) -> MaadResult { + if self.prefix_lengths.is_empty() { + return empty_result(self.total_addrs); + } + let weighted = &self.weighted[measure]; + let structure = compute_weighted_structure(&weighted.moments, q_values); + let dimensions = self.dimensions(&structure, &weighted.entropies); + self.result(structure, Vec::new(), dimensions) + } + + fn dimensions(&self, structure: &[StructureRow], entropies: &[f64]) -> Vec { + if structure.is_empty() { + return Vec::new(); + } + dimension_rows(structure, info_dimension(&self.prefix_lengths, entropies)) + } + + fn result( + &self, + structure: Vec, + spectrum: Vec, + dimensions: Vec, + ) -> MaadResult { + MaadResult { + schema_version: SCHEMA_VERSION, + metadata: MaadMetadata { + input: "-", + min_prefix_length: self.prefix_lengths.first().copied(), + max_prefix_length: self.prefix_lengths.last().copied(), + prefix_lengths: self.prefix_lengths.clone(), + total_addrs: self.total_addrs, + }, + structure, + spectrum, + dimensions, + } + } +} + fn empty_result(total_addrs: usize) -> MaadResult { MaadResult { schema_version: SCHEMA_VERSION, @@ -281,26 +442,155 @@ fn empty_result(total_addrs: usize) -> MaadResult { } } -fn build_prefix_counts( +/// Walk the prefix levels of sorted, distinct addresses from the root down. +/// +/// Only a parent level and its child level are held at once, so memory stays linear in the +/// address count. Validity and path pruning use distinct-address counts; each weight column +/// contributes its own masses and entropies over the same selected parents. +fn analyze( addresses: &[B], - max_prefix_length: u8, -) -> Vec> { - let mut counts = Vec::with_capacity(usize::from(max_prefix_length) + 1); - for prefix_length in 0..=max_prefix_length { - let mut prefixes = Vec::new(); - for &address in addresses { - let prefix = address.prefix(prefix_length); - if let Some((last_prefix, count)) = prefixes.last_mut() - && *last_prefix == prefix - { - *count += 1; - } else { - prefixes.push((prefix, 1)); + weights: &[Vec], + totals: &[f64], + config: &MaadConfig, +) -> Analysis { + let total_addrs = addresses.len(); + let mut analysis = Analysis { + total_addrs, + prefix_lengths: Vec::new(), + counts: MeasureAnalysis::new(), + weighted: weights.iter().map(|_| MeasureAnalysis::new()).collect(), + }; + let mut parents = prefix_level(addresses, 0); + let mut parent_masses: Vec<_> = weights + .iter() + .map(|column| weight_level(addresses, column, 0)) + .collect(); + let mut path_allowed = vec![true; parents.len()]; + + for prefix_length in 0..=config.max_prefix_length { + let children = prefix_level(addresses, prefix_length + 1); + let child_masses: Vec<_> = weights + .iter() + .map(|column| weight_level(addresses, column, prefix_length + 1)) + .collect(); + let measured = prefix_length >= config.min_prefix_length; + let propagate = prefix_length < config.max_prefix_length; + let mut selected = Vec::new(); + let mut child_path_allowed = Vec::with_capacity(if propagate { children.len() } else { 0 }); + let mut next_child = 0; + + for (parent_index, &(prefix, count)) in parents.iter().enumerate() { + let first_child = prefix.child(false); + let last_child = prefix.child(true); + while next_child < children.len() && children[next_child].0 < first_child { + next_child += 1; + } + let child_start = next_child; + while next_child < children.len() && children[next_child].0 <= last_child { + next_child += 1; + } + let valid = is_valid_parent::(count, prefix_length, config.full_threshold); + if measured && path_allowed[parent_index] && valid && child_start < next_child { + selected.push((parent_index, child_start, next_child)); + } + if propagate { + let is_branch = next_child - child_start == 2; + let allowed = path_allowed[parent_index] && (!is_branch || valid); + child_path_allowed.extend(std::iter::repeat_n(allowed, next_child - child_start)); + } + } + + if !selected.is_empty() { + analysis.prefix_lengths.push(prefix_length); + analysis.counts.moments.push(select_moment( + &selected, + |index| parents[index].1, + |index| children[index].1, + )); + analysis.counts.entropies.push(level_entropy( + parents.iter().map(|&(_, count)| count as f64), + total_addrs as f64, + )); + for (measure, weighted) in analysis.weighted.iter_mut().enumerate() { + let parents = &parent_masses[measure]; + let children = &child_masses[measure]; + weighted.moments.push(select_moment( + &selected, + |index| parents[index], + |index| children[index], + )); + weighted + .entropies + .push(level_entropy(parents.iter().copied(), totals[measure])); } } - counts.push(prefixes); + + debug_assert!(!propagate || child_path_allowed.len() == children.len()); + path_allowed = child_path_allowed; + parents = children; + parent_masses = child_masses; + } + + analysis +} + +fn prefix_level(addresses: &[B], prefix_length: u8) -> Vec<(B, usize)> { + let mut prefixes: Vec<(B, usize)> = Vec::new(); + for &address in addresses { + let prefix = address.prefix(prefix_length); + if let Some((last_prefix, count)) = prefixes.last_mut() + && *last_prefix == prefix + { + *count += 1; + } else { + prefixes.push((prefix, 1)); + } + } + prefixes +} + +fn weight_level(addresses: &[B], weights: &[f64], prefix_length: u8) -> Vec { + let mut level = Vec::new(); + let mut last_prefix = None; + for (&address, &weight) in addresses.iter().zip(weights) { + let prefix = address.prefix(prefix_length); + match level.last_mut() { + Some(total) if last_prefix == Some(prefix) => *total += weight, + _ => level.push(weight), + } + last_prefix = Some(prefix); } - counts + level +} + +fn select_moment( + selected: &[(usize, usize, usize)], + parent: impl Fn(usize) -> M, + child: impl Fn(usize) -> M, +) -> PreparedMoment { + PreparedMoment { + parent_masses: selected + .iter() + .map(|&(parent_index, _, _)| parent(parent_index)) + .collect(), + child_masses: selected + .iter() + .map(|&(_, child_start, child_end)| { + ( + child(child_start), + (child_end - child_start == 2).then(|| child(child_start + 1)), + ) + }) + .collect(), + } +} + +fn level_entropy(masses: impl Iterator, total: f64) -> f64 { + masses + .map(|mass| mass / total) + .filter(|&probability| probability > 0.0) + .map(|probability| probability * probability.log2()) + .sum::() } fn validate_config(config: &MaadConfig, address_bits: u8) -> Result, MaadError> { @@ -381,128 +671,26 @@ fn is_valid_parent(count: usize, prefix_length: u8, full_threshol count > 1 && (count as f64).log2() / f64::from(B::WIDTH - prefix_length) < 1.0 - full_threshold } -fn prepare_valid_moments( - counts: &[Vec<(B, usize)>], - config: &MaadConfig, -) -> Vec<(u8, PreparedMoment)> { - let mut prepared = Vec::new(); - let mut path_allowed = vec![true; counts[0].len()]; - - for prefix_length in 0..=config.max_prefix_length { - let parents = &counts[usize::from(prefix_length)]; - let children = &counts[usize::from(prefix_length) + 1]; - - if prefix_length >= config.min_prefix_length { - let moment = prepare_moment_at_length( - parents, - children, - &path_allowed, - prefix_length, - config.full_threshold, - ); - if !moment.parent_counts.is_empty() { - prepared.push((prefix_length, moment)); - } - } - - if prefix_length < config.max_prefix_length { - path_allowed = propagate_allowed_paths( - parents, - children, - &path_allowed, - prefix_length, - config.full_threshold, - ); - } - } - - prepared -} - -fn prepare_moment_at_length( - parents: &[(B, usize)], - children: &[(B, usize)], - path_allowed: &[bool], - prefix_length: u8, - full_threshold: f64, -) -> PreparedMoment { - let mut parent_counts = Vec::new(); - let mut child_counts = Vec::new(); - let mut next_child = 0; - - for (parent_index, &(prefix, count)) in parents.iter().enumerate() { - let first_child = prefix.child(false); - let last_child = prefix.child(true); - while next_child < children.len() && children[next_child].0 < first_child { - next_child += 1; - } - let child_start = next_child; - while next_child < children.len() && children[next_child].0 <= last_child { - next_child += 1; - } - if path_allowed[parent_index] - && is_valid_parent::(count, prefix_length, full_threshold) - && child_start < next_child - { - parent_counts.push(count); - child_counts.push( - children[child_start..next_child] - .iter() - .map(|&(_, child_count)| child_count) - .collect(), - ); - } - } - - PreparedMoment { - parent_counts, - child_counts, - } -} - -fn propagate_allowed_paths( - parents: &[(B, usize)], - children: &[(B, usize)], - path_allowed: &[bool], - prefix_length: u8, - full_threshold: f64, -) -> Vec { - let mut child_path_allowed = Vec::with_capacity(children.len()); - let mut next_child = 0; - - for (parent_index, &(prefix, count)) in parents.iter().enumerate() { - let first_child = prefix.child(false); - let last_child = prefix.child(true); - while next_child < children.len() && children[next_child].0 < first_child { - next_child += 1; - } - let child_start = next_child; - while next_child < children.len() && children[next_child].0 <= last_child { - next_child += 1; - } - let is_branch = next_child - child_start == 2; - let allowed = path_allowed[parent_index] - && (!is_branch || is_valid_parent::(count, prefix_length, full_threshold)); - child_path_allowed.extend(std::iter::repeat_n(allowed, next_child - child_start)); - } - - debug_assert_eq!(child_path_allowed.len(), children.len()); - child_path_allowed -} - -fn one_moment(prepared: &PreparedMoment, powers: &[f64]) -> (f64, f64) { - if prepared.parent_counts.is_empty() { +fn one_moment( + prepared: &PreparedMoment, + parent_power: impl Fn(M) -> f64, + child_power: impl Fn(M) -> f64, +) -> (f64, f64) { + if prepared.parent_masses.is_empty() { return (0.0, 0.0); } let parent_powers: Vec<_> = prepared - .parent_counts + .parent_masses .iter() - .map(|&count| powers[count]) + .map(|&mass| parent_power(mass)) .collect(); let child_power_sums: Vec<_> = prepared - .child_counts + .child_masses .iter() - .map(|children| children.iter().map(|&count| powers[count]).sum::()) + .map(|&(first, second)| match second { + Some(second) => child_power(first) + child_power(second), + None => child_power(first), + }) .collect(); let this_z: f64 = parent_powers.iter().sum(); let next_z: f64 = child_power_sums.iter().sum(); @@ -517,44 +705,130 @@ fn one_moment(prepared: &PreparedMoment, powers: &[f64]) -> (f64, f64) { (this_z.log2() - next_z.log2(), d2) } -fn compute_structure(prepared: &[PreparedMoment], q_values: &[f64]) -> Vec { +fn moment_masses(moment: &PreparedMoment) -> impl Iterator + '_ { + moment.parent_masses.iter().copied().chain( + moment + .child_masses + .iter() + .flat_map(|&(first, second)| std::iter::once(first).chain(second)), + ) +} + +fn compute_structure(prepared: &[PreparedMoment], q_values: &[f64]) -> Vec { if prepared.is_empty() { return Vec::new(); } let max_count = prepared .iter() - .flat_map(|moment| { - moment - .parent_counts - .iter() - .chain(moment.child_counts.iter().flatten()) - }) - .copied() + .flat_map(moment_masses) .max() .unwrap_or_default(); q_values - .iter() - .copied() - .map(|q| { + .par_iter() + .map(|&q| { let powers: Vec<_> = (0..=max_count) .map(|count| (count as f64).powf(q)) .collect(); - let (tau_sum, d2_sum) = prepared + structure_row( + q, + prepared + .iter() + .map(|moment| one_moment(moment, |count| powers[count], |count| powers[count])), + ) + }) + .collect() +} + +/// Weighted masses repeat heavily, so each q raises only the distinct masses to the power +/// and every moment looks its masses up by index. A q whose raw powers leave the normal +/// range recomputes each moment from masses rescaled within that moment. +fn compute_weighted_structure( + prepared: &[PreparedMoment], + q_values: &[f64], +) -> Vec { + let mut distinct: Vec = prepared.iter().flat_map(moment_masses).collect(); + distinct.sort_unstable_by(f64::total_cmp); + distinct.dedup_by(|next, kept| next.total_cmp(kept).is_eq()); + let index = |mass: f64| { + distinct + .binary_search_by(|probe| probe.total_cmp(&mass)) + .expect("every mass is in the distinct table") + }; + let indexed: Vec> = prepared + .iter() + .map(|moment| PreparedMoment { + parent_masses: moment + .parent_masses + .iter() + .map(|&mass| index(mass)) + .collect(), + child_masses: moment + .child_masses .iter() - .map(|moment| one_moment(moment, &powers)) - .fold((0.0, 0.0), |(tau_sum, d2_sum), (tau, d2)| { - (tau_sum + tau, d2_sum + d2) - }); - let count = prepared.len() as f64; - StructureRow { + .map(|&(first, second)| (index(first), second.map(index))) + .collect(), + }) + .collect(); + q_values + .par_iter() + .map(|&q| { + let powers: Vec<_> = distinct.iter().map(|mass| mass.powf(q)).collect(); + let raw_powers_normal = powers.iter().all(|power| power.is_normal()); + structure_row( q, - tau_tilde: tau_sum / count, - sd: d2_sum.sqrt() / count, - } + indexed.iter().zip(prepared).map(|(moment, masses)| { + raw_powers_normal + .then(|| one_moment(moment, |index| powers[index], |index| powers[index])) + .filter(|(tau, d2)| tau.is_finite() && d2.is_finite()) + .unwrap_or_else(|| rescaled_moment(masses, q)) + }), + ) }) .collect() } +/// One weighted moment with parent and child masses divided by the mass that dominates their +/// partition sum at `q`, so the sums stay near one at any weight scale. +fn rescaled_moment(prepared: &PreparedMoment, q: f64) -> (f64, f64) { + let dominant: fn(f64, f64) -> f64 = if q >= 0.0 { f64::max } else { f64::min }; + let parent_scale = prepared + .parent_masses + .iter() + .copied() + .reduce(dominant) + .unwrap_or(1.0); + let child_scale = prepared + .child_masses + .iter() + .flat_map(|&(first, second)| std::iter::once(first).chain(second)) + .reduce(dominant) + .unwrap_or(1.0); + let (tau, d2) = one_moment( + prepared, + |mass| (mass / parent_scale).powf(q), + |mass| (mass / child_scale).powf(q), + ); + let scale_ratio = parent_scale / child_scale; + let log_scale_ratio = if scale_ratio.is_finite() { + scale_ratio.log2() + } else { + parent_scale.log2() - child_scale.log2() + }; + (tau + q * log_scale_ratio, d2) +} + +fn structure_row(q: f64, moments: impl ExactSizeIterator) -> StructureRow { + let count = moments.len() as f64; + let (tau_sum, d2_sum) = moments.fold((0.0, 0.0), |(tau_sum, d2_sum), (tau, d2)| { + (tau_sum + tau, d2_sum + d2) + }); + StructureRow { + q, + tau_tilde: tau_sum / count, + sd: d2_sum.sqrt() / count, + } +} + /// Central-difference `(q, alpha, f)` estimates for every interior structure row. fn spectrum_estimates(structure: &[StructureRow], q_step: f64) -> Vec<(f64, SpectrumRow)> { structure @@ -615,21 +889,6 @@ fn compute_spectrum(structure: &[StructureRow], q_step: f64) -> Vec rows } -fn compute_dimensions( - counts: &[Vec<(B, usize)>], - prefix_lengths: &[u8], - structure: &[StructureRow], - total_addresses: usize, -) -> Vec { - if prefix_lengths.is_empty() || structure.is_empty() { - return Vec::new(); - } - dimension_rows( - structure, - info_dimension(counts, prefix_lengths, total_addresses), - ) -} - fn dimension_rows(structure: &[StructureRow], information_dimension: f64) -> Vec { let from_structure = |q: f64| { structure @@ -653,24 +912,11 @@ fn dimension_rows(structure: &[StructureRow], information_dimension: f64) -> Vec ] } -fn info_dimension( - counts: &[Vec<(B, usize)>], - prefix_lengths: &[u8], - total_addresses: usize, -) -> f64 { - let total = total_addresses as f64; +fn info_dimension(prefix_lengths: &[u8], entropies: &[f64]) -> f64 { let points: Vec<_> = prefix_lengths .iter() - .map(|&prefix_length| { - let entropy = counts[usize::from(prefix_length)] - .iter() - .map(|&(_, count)| { - let probability = count as f64 / total; - probability * probability.log2() - }) - .sum::(); - (-(f64::from(prefix_length)), entropy) - }) + .zip(entropies) + .map(|(&prefix_length, &entropy)| (-(f64::from(prefix_length)), entropy)) .collect(); let point_count = points.len() as f64; let mean_x = points.iter().map(|point| point.0).sum::() / point_count; @@ -705,25 +951,90 @@ mod tests { where A::Bits: Into, { - let width = A::Bits::WIDTH; - let addresses: BTreeSet = addresses + let masses: BTreeMap = addresses .into_iter() - .map(|address| address.bits().into()) + .map(|address| (address.bits().into(), 1.0)) .collect(); - if addresses.len() < MIN_MAAD_ADDRESSES { - return empty_result(addresses.len()); + reference_measure(&masses, A::Bits::WIDTH, config, true) + } + + fn reference_weighted_compute( + entries: &[(A, f64)], + config: MaadConfig, + ) -> MaadResult + where + A::Bits: Into, + { + let mut masses = BTreeMap::::new(); + for &(address, weight) in entries { + *masses.entry(address.bits().into()).or_insert(0.0) += weight; } + reference_measure(&masses, A::Bits::WIDTH, config, false) + } + + /// Ordered-map MAAD over per-address masses. Validity uses distinct counts; + /// moments and entropy use the summed masses. + fn reference_measure( + masses: &BTreeMap, + width: u8, + config: MaadConfig, + with_spectrum: bool, + ) -> MaadResult { + if masses.len() < MIN_MAAD_ADDRESSES { + return empty_result(masses.len()); + } + let addresses: BTreeSet = masses.keys().copied().collect(); let counts = reference_prefix_counts(&addresses, width); - let prepared = reference_prepare_valid_moments(&counts, &config, width); - if prepared.is_empty() { - return empty_result(addresses.len()); + let level_masses: Vec> = (0..=width) + .map(|prefix_length| { + let mut level = BTreeMap::new(); + for (&address, &mass) in masses { + let prefix = if prefix_length == 0 { + 0 + } else { + address >> (width - prefix_length) + }; + *level.entry(prefix).or_insert(0.0) += mass; + } + level + }) + .collect(); + let selected = reference_selected_parents(&counts, &config, width); + if selected.is_empty() { + return empty_result(masses.len()); } - let (prefix_lengths, prepared): (Vec<_>, Vec<_>) = prepared.into_iter().unzip(); + let prefix_lengths: Vec = selected.iter().map(|(length, _)| *length).collect(); + let prepared: Vec> = selected + .iter() + .map(|(prefix_length, parents)| { + let level = usize::from(*prefix_length); + PreparedMoment { + parent_masses: parents + .iter() + .map(|parent| level_masses[level][parent]) + .collect(), + child_masses: parents + .iter() + .map(|&parent| { + let children: Vec<_> = [parent << 1, (parent << 1) | 1] + .into_iter() + .filter_map(|child| level_masses[level + 1].get(&child).copied()) + .collect(); + (children[0], children.get(1).copied()) + }) + .collect(), + } + }) + .collect(); let q_values = reference_q_values(&config); let structure = reference_structure(&prepared, &q_values); - let spectrum = compute_spectrum(&structure, config.q_step); - let dimensions = - reference_dimensions(&counts, &prefix_lengths, &structure, addresses.len()); + let spectrum = if with_spectrum { + compute_spectrum(&structure, config.q_step) + } else { + Vec::new() + }; + let total = masses.values().sum::(); + let dimensions = reference_dimensions(&level_masses, &prefix_lengths, &structure, total); MaadResult { schema_version: SCHEMA_VERSION, metadata: MaadMetadata { @@ -731,7 +1042,7 @@ mod tests { min_prefix_length: prefix_lengths.first().copied(), max_prefix_length: prefix_lengths.last().copied(), prefix_lengths, - total_addrs: addresses.len(), + total_addrs: masses.len(), }, structure, spectrum, @@ -759,12 +1070,12 @@ mod tests { counts } - fn reference_prepare_valid_moments( + fn reference_selected_parents( counts: &[BTreeMap], config: &MaadConfig, width: u8, - ) -> Vec<(u8, PreparedMoment)> { - let mut prepared = Vec::new(); + ) -> Vec<(u8, Vec)> { + let mut selected = Vec::new(); let mut path_allowed = BTreeMap::from([(0, true)]); for prefix_length in 0..=config.max_prefix_length { @@ -772,36 +1083,24 @@ mod tests { let children = &counts[usize::from(prefix_length) + 1]; if prefix_length >= config.min_prefix_length { - let mut parent_counts = Vec::new(); - let mut child_counts = Vec::new(); - for (&prefix, &count) in parents { - if !path_allowed[&prefix] - || !reference_valid_parent( - count, - prefix_length, - config.full_threshold, - width, - ) - { - continue; - } - let child_counts_for_parent: Vec<_> = [prefix << 1, (prefix << 1) | 1] - .into_iter() - .filter_map(|child| children.get(&child).copied()) - .collect(); - if !child_counts_for_parent.is_empty() { - parent_counts.push(count); - child_counts.push(child_counts_for_parent); - } - } - if !parent_counts.is_empty() { - prepared.push(( - prefix_length, - PreparedMoment { - parent_counts, - child_counts, - }, - )); + let level: Vec = parents + .iter() + .filter(|&(prefix, &count)| { + path_allowed[prefix] + && reference_valid_parent( + count, + prefix_length, + config.full_threshold, + width, + ) + && [prefix << 1, (prefix << 1) | 1] + .iter() + .any(|child| children.contains_key(child)) + }) + .map(|(&prefix, _)| prefix) + .collect(); + if !level.is_empty() { + selected.push((prefix_length, level)); } } @@ -827,7 +1126,7 @@ mod tests { } } - prepared + selected } fn reference_valid_parent( @@ -855,7 +1154,10 @@ mod tests { .collect() } - fn reference_structure(prepared: &[PreparedMoment], q_values: &[f64]) -> Vec { + fn reference_structure( + prepared: &[PreparedMoment], + q_values: &[f64], + ) -> Vec { q_values .iter() .copied() @@ -864,17 +1166,17 @@ mod tests { .iter() .map(|moment| { let parent_powers: Vec<_> = moment - .parent_counts + .parent_masses .iter() - .map(|&count| (count as f64).powf(q)) + .map(|&mass| mass.powf(q)) .collect(); let child_power_sums: Vec<_> = moment - .child_counts + .child_masses .iter() - .map(|children| { - children - .iter() - .map(|&count| (count as f64).powf(q)) + .map(|&(first, second)| { + std::iter::once(first) + .chain(second) + .map(|mass| mass.powf(q)) .sum::() }) .collect(); @@ -904,19 +1206,18 @@ mod tests { } fn reference_dimensions( - counts: &[BTreeMap], + level_masses: &[BTreeMap], prefix_lengths: &[u8], structure: &[StructureRow], - total_addresses: usize, + total: f64, ) -> Vec { - let total = total_addresses as f64; let points: Vec<_> = prefix_lengths .iter() .map(|&prefix_length| { - let entropy = counts[usize::from(prefix_length)] + let entropy = level_masses[usize::from(prefix_length)] .values() - .map(|&count| { - let probability = count as f64 / total; + .map(|&mass| { + let probability = mass / total; probability * probability.log2() }) .sum::(); @@ -1459,4 +1760,332 @@ mod tests { close(q2.sd, (2.0 / 49.0_f64).sqrt() / 2.0); } + + fn pseudo_random_addresses(count: usize) -> Vec { + let mut state = 0x2545_f491_u32; + (0..count) + .map(|_| { + state = state.wrapping_mul(1_664_525).wrapping_add(1_013_904_223); + Ipv4Addr::from(state) + }) + .collect() + } + + #[test] + fn unit_weights_reproduce_the_unweighted_result_exactly() { + let dense: Vec<_> = (0..2) + .flat_map(|third| (0..=255).map(move |fourth| Ipv4Addr::new(10, 0, third, fourth))) + .chain([Ipv4Addr::new(192, 0, 2, 1)]) + .collect(); + let clustered: Vec<_> = (0..64) + .map(|index| Ipv4Addr::new(10, 1, index / 4, index * 3)) + .chain(pseudo_random_addresses(512)) + .collect(); + + for addresses in [dense, clustered, pseudo_random_addresses(1024)] { + let unweighted = compute(addresses.iter().copied()); + let weighted = + compute_weighted(addresses.iter().map(|&address| (address, 1.0))).unwrap(); + + assert!(!unweighted.structure.is_empty()); + assert_eq!(weighted.metadata, unweighted.metadata); + assert_eq!(weighted.structure, unweighted.structure); + assert_eq!(weighted.dimensions, unweighted.dimensions); + assert!(weighted.spectrum.is_empty()); + } + } + + #[test] + fn ipv6_unit_weights_reproduce_the_unweighted_result_exactly() { + let addresses: Vec<_> = (0..512_u128) + .map(|index| documentation_ipv6((index % 37) << 72 | (index * 7919) << 40)) + .collect(); + let unweighted = compute(addresses.iter().copied()); + let weighted = compute_weighted(addresses.iter().map(|&address| (address, 1.0))).unwrap(); + + assert!(!unweighted.structure.is_empty()); + assert_eq!(unweighted.metadata.min_prefix_length, Some(23)); + assert_eq!(weighted.metadata, unweighted.metadata); + assert_eq!(weighted.structure, unweighted.structure); + assert_eq!(weighted.dimensions, unweighted.dimensions); + assert!(weighted.spectrum.is_empty()); + } + + #[test] + fn shared_measures_equal_separate_computations() { + let v4: Vec<_> = pseudo_random_addresses(2_000) + .into_iter() + .chain((0..=255).map(|last| Ipv4Addr::new(198, 51, 100, last))) + .enumerate() + .map(|(index, address)| { + let index = index as u32; + ( + address, + [f64::from(index % 7 + 1), f64::from(index % 5 * 1_500 + 40)], + ) + }) + .collect(); + let (addresses, [packets, bytes]) = compute_measures(v4.iter().copied()).unwrap(); + assert!(!addresses.spectrum.is_empty()); + assert_eq!(addresses, compute(v4.iter().map(|&(address, _)| address))); + assert_eq!( + packets, + compute_weighted(v4.iter().map(|&(address, weights)| (address, weights[0]))).unwrap() + ); + assert_eq!( + bytes, + compute_weighted(v4.iter().map(|&(address, weights)| (address, weights[1]))).unwrap() + ); + + let v6: Vec<_> = (0..1_024_u128) + .map(|index| { + ( + documentation_ipv6((index % 37) << 72 | (index * 7919) << 40), + [(index % 3 + 1) as f64 * 0.5, 1e12 + index as f64], + ) + }) + .collect(); + let (addresses, [first, second]) = compute_measures(v6.iter().copied()).unwrap(); + assert!(!addresses.structure.is_empty()); + assert_eq!(addresses, compute(v6.iter().map(|&(address, _)| address))); + assert_eq!( + first, + compute_weighted(v6.iter().map(|&(address, weights)| (address, weights[0]))).unwrap() + ); + assert_eq!( + second, + compute_weighted(v6.iter().map(|&(address, weights)| (address, weights[1]))).unwrap() + ); + } + + fn assert_weighted_matches_reference(entries: &[(A, f64)]) + where + A::Bits: Into, + { + let config = A::default_config(); + let result = compute_weighted_with_config(entries.iter().copied(), config).unwrap(); + let reference = reference_weighted_compute(entries, config); + assert!(!result.structure.is_empty()); + assert_eq!(result.metadata, reference.metadata); + assert!(result.spectrum.is_empty() && reference.spectrum.is_empty()); + assert_eq!(result.structure.len(), reference.structure.len()); + assert_eq!(result.dimensions.len(), reference.dimensions.len()); + for (actual, expected) in result.structure.iter().zip(&reference.structure) { + close(actual.q, expected.q); + close(actual.tau_tilde, expected.tau_tilde); + close(actual.sd, expected.sd); + } + for (actual, expected) in result.dimensions.iter().zip(&reference.dimensions) { + close(actual.q, expected.q); + close(actual.dim, expected.dim); + close(actual.sd, expected.sd); + } + } + + #[test] + fn shared_measures_sum_repeated_addresses_like_separate_computations() { + let repeated = [ + (Ipv4Addr::new(10, 0, 0, 1), [1e16, 1.0]), + (Ipv4Addr::new(10, 0, 0, 1), [1.0, 1e16]), + (Ipv4Addr::new(10, 0, 0, 1), [1.0, 3.0]), + ]; + let entries: Vec<_> = pseudo_random_addresses(400) + .into_iter() + .map(|address| (address, [2.0, 5.0])) + .chain(repeated) + .collect(); + + let (_, [first, second]) = compute_measures(entries.iter().copied()).unwrap(); + + assert_eq!( + first, + compute_weighted( + entries + .iter() + .map(|&(address, weights)| (address, weights[0])) + ) + .unwrap() + ); + assert_eq!( + second, + compute_weighted( + entries + .iter() + .map(|&(address, weights)| (address, weights[1])) + ) + .unwrap() + ); + } + + #[test] + fn weighted_path_matches_the_ordered_map_reference() { + let v4: Vec<_> = pseudo_random_addresses(600) + .into_iter() + .chain((0..=255).map(|last| Ipv4Addr::new(198, 51, 100, last))) + .enumerate() + .map(|(index, address)| (address, f64::from((index % 11) as u32 * 37 + 1))) + .collect(); + let duplicates: Vec<_> = v4 + .iter() + .take(50) + .map(|&(address, weight)| (address, weight * 0.5 + 0.25)) + .collect(); + assert_weighted_matches_reference(&[v4, duplicates].concat()); + + let v6: Vec<_> = (0..1_024_u128) + .map(|index| { + ( + documentation_ipv6((index % 37) << 72 | (index * 7919) << 40), + (index % 5 + 1) as f64 * 1_500.0, + ) + }) + .collect(); + assert_weighted_matches_reference(&v6); + } + + #[test] + fn weighted_result_ignores_order_and_sums_duplicate_addresses() { + let entries: Vec<_> = pseudo_random_addresses(300) + .into_iter() + .chain((0..40).map(|last| Ipv4Addr::new(198, 51, 100, last))) + .enumerate() + .map(|(index, address)| (address, f64::from((index % 13) as u32 * 4 + 4))) + .collect(); + let expected = compute_weighted(entries.iter().copied()).unwrap(); + + let mut reversed = entries.clone(); + reversed.reverse(); + let mut rotated = entries.clone(); + rotated.rotate_left(117); + let split: Vec<_> = entries + .iter() + .flat_map(|&(address, weight)| [(address, weight / 4.0), (address, weight * 3.0 / 4.0)]) + .rev() + .collect(); + + assert!(!expected.structure.is_empty()); + assert_eq!(compute_weighted(reversed).unwrap(), expected); + assert_eq!(compute_weighted(rotated).unwrap(), expected); + assert_eq!(compute_weighted(split).unwrap(), expected); + } + + #[test] + fn weighted_moments_and_entropy_match_a_hand_computed_case() { + let entries = [ + (Ipv4Addr::new(10, 0, 0, 0), 1.0), + (Ipv4Addr::new(10, 0, 0, 2), 3.0), + (Ipv4Addr::new(10, 0, 0, 4), 4.0), + ]; + let config = MaadConfig { + q_min: 0.0, + q_max: 2.0, + q_step: 1.0, + min_prefix_length: 29, + max_prefix_length: 30, + ..MaadConfig::default() + }; + + let result = compute_weighted_with_config(entries, config).unwrap(); + + assert_eq!(result.metadata.total_addrs, 3); + assert_eq!(result.metadata.prefix_lengths, vec![29, 30]); + assert!(result.spectrum.is_empty()); + let tau = |q: f64| ((q - 1.0) + (2.0 * q - (1.0 + 3.0_f64.powf(q)).log2())) / 2.0; + assert_eq!(result.structure.len(), 3); + for row in &result.structure { + close(row.tau_tilde, tau(row.q)); + close(row.sd, 0.0); + } + let dimensions: Vec<_> = result + .dimensions + .iter() + .map(|row| (row.q, row.dim, row.sd)) + .collect(); + assert_eq!(dimensions.len(), 3); + close(dimensions[0].0, 0.0); + close(dimensions[0].1, 1.0); + close(dimensions[1].0, 1.0); + close(dimensions[1].1, 1.0); + close(dimensions[1].2, 0.0); + close(dimensions[2].0, 2.0); + close(dimensions[2].1, (5.0 - 10.0_f64.log2()) / 2.0); + } + + #[test] + fn weighted_moments_are_invariant_to_the_weight_scale() { + let pair = [Ipv4Addr::new(10, 0, 0, 0), Ipv4Addr::new(10, 0, 0, 128)]; + for scale in [1.0, 1e160, 1e-200] { + let result = compute_weighted(pair.map(|address| (address, scale))).unwrap(); + let d2 = result.dimensions.iter().find(|row| row.q == 2.0).unwrap(); + close(d2.dim, 1.0 / 17.0); + assert!(result.structure.iter().all(|row| row.tau_tilde.is_finite())); + } + + let entries: Vec<_> = pseudo_random_addresses(600) + .into_iter() + .chain((0..=255).map(|last| Ipv4Addr::new(198, 51, 100, last))) + .enumerate() + .map(|(index, address)| (address, f64::from((index % 11) as u32 * 37 + 1))) + .collect(); + let expected = compute_weighted(entries.iter().copied()).unwrap(); + for scale in [1e300, 1e160, 1e-200, 1e-300] { + let scaled = compute_weighted( + entries + .iter() + .map(|&(address, weight)| (address, weight * scale)), + ) + .unwrap(); + assert_eq!(scaled.metadata, expected.metadata); + for (actual, expected) in scaled.structure.iter().zip(&expected.structure) { + close(actual.tau_tilde, expected.tau_tilde); + close(actual.sd, expected.sd); + } + for (actual, expected) in scaled.dimensions.iter().zip(&expected.dimensions) { + close(actual.dim, expected.dim); + close(actual.sd, expected.sd); + } + } + } + + #[test] + fn weighted_entropy_treats_masses_that_normalize_to_zero_as_absent() { + let entries = [ + (Ipv4Addr::new(10, 0, 0, 0), 1e-300), + (Ipv4Addr::new(10, 0, 0, 128), 1e-300), + (Ipv4Addr::new(192, 0, 2, 1), 1e50), + ]; + + let result = compute_weighted(entries).unwrap(); + + let d1 = result.dimensions.iter().find(|row| row.q == 1.0).unwrap(); + close(d1.dim, 0.0); + assert!( + result + .structure + .iter() + .all(|row| row.tau_tilde.is_finite() && row.sd.is_finite()) + ); + assert!(result.dimensions.iter().all(|row| row.dim.is_finite())); + } + + #[test] + fn weighted_input_rejects_non_positive_and_non_finite_weights() { + let valid = (Ipv4Addr::new(192, 0, 2, 1), 1.0); + for weight in [0.0, -1.0, f64::NAN, f64::INFINITY, f64::NEG_INFINITY] { + let address = Ipv4Addr::new(192, 0, 2, 2); + let error = compute_weighted([valid, (address, weight)]).unwrap_err(); + assert!(matches!( + error, + MaadError::InvalidWeight { address: rejected, .. } if rejected == IpAddr::V4(address) + )); + } + assert!(matches!( + compute_weighted([ + valid, + (Ipv4Addr::new(192, 0, 2, 2), f64::MAX), + (valid.0, f64::MAX) + ]), + Err(MaadError::NonFiniteTotalWeight { .. }) + )); + } } diff --git a/tools/netflow-db/src/main.rs b/tools/netflow-db/src/main.rs index 15c9d547..030e9b26 100644 --- a/tools/netflow-db/src/main.rs +++ b/tools/netflow-db/src/main.rs @@ -249,6 +249,9 @@ struct MaadArgs { /// Read IPv6 addresses and use the IPv6 prefix range (/23-/64). #[arg(short = '6', long)] ipv6: bool, + /// Read `ADDR,MEASURE` rows and emit measure-weighted structure and dimensions. + #[arg(long)] + weighted: bool, } #[derive(Debug, Args)] @@ -604,16 +607,52 @@ fn run_web_verify(args: WebVerifyArgs) -> Result<()> { } fn run_maad(args: MaadArgs) -> Result<()> { - let result = if args.ipv6 { - maad::compute(read_address_lines::(args.input, "IPv6")?) - } else { - maad::compute(read_address_lines::(args.input, "IPv4")?) + let result = match (args.ipv6, args.weighted) { + (false, false) => maad::compute(read_address_lines::(args.input, "IPv4")?), + (true, false) => maad::compute(read_address_lines::(args.input, "IPv6")?), + (false, true) => { + maad::compute_weighted(read_weighted_lines::(args.input, "IPv4")?)? + } + (true, true) => { + maad::compute_weighted(read_weighted_lines::(args.input, "IPv6")?)? + } }; maad::write_json(&result, io::stdout().lock())?; io::stdout().flush()?; Ok(()) } +/// Read one address family's `ADDR,MEASURE` rows from a file or standard input. +fn read_weighted_lines(input: Option, family: &str) -> Result> +where + A: std::str::FromStr, + A::Err: std::error::Error + Send + Sync + 'static, +{ + let input = open_input(input)?; + let mut entries = Vec::new(); + for line in input.lines() { + let line = line?; + let value = line.trim(); + if value.is_empty() { + continue; + } + let (address, measure) = value + .split_once(',') + .with_context(|| format!("expected ADDR,MEASURE, got {value:?}"))?; + entries.push(( + address + .trim() + .parse::() + .with_context(|| format!("invalid {family} address {address:?}"))?, + measure + .trim() + .parse::() + .with_context(|| format!("invalid measure {measure:?}"))?, + )); + } + Ok(entries) +} + fn run_singularity(args: SingularityArgs) -> Result<()> { let addresses = read_address_lines::(args.input, "IPv4")?; singularity::write_csv(&singularity::score(addresses), io::stdout().lock())?; @@ -621,18 +660,22 @@ fn run_singularity(args: SingularityArgs) -> Result<()> { Ok(()) } +fn open_input(input: Option) -> Result> { + Ok(match input { + Some(path) => Box::new(BufReader::new( + File::open(&path).with_context(|| format!("unable to open {}", path.display()))?, + )), + None => Box::new(BufReader::new(io::stdin())), + }) +} + /// Read one address family's addresses, one per line, from a file or standard input. fn read_address_lines(input: Option, family: &str) -> Result> where A: std::str::FromStr, A::Err: std::error::Error + Send + Sync + 'static, { - let input: Box = match input { - Some(path) => Box::new(BufReader::new( - File::open(&path).with_context(|| format!("unable to open {}", path.display()))?, - )), - None => Box::new(BufReader::new(io::stdin())), - }; + let input = open_input(input)?; let mut addresses = Vec::new(); for line in input.lines() { let line = line?;