use nmbrs_workload::model::{ParsedOp, WorkloadPhase};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct WrapperName(pub &'static str);
impl WrapperName {
pub const fn new(name: &'static str) -> Self {
Self(name)
}
pub fn as_str(&self) -> &'static str {
self.0
}
}
impl std::fmt::Display for WrapperName {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.0)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum WrapperLevel {
Op,
Stanza,
Phase,
Scenario,
Session,
}
#[derive(Clone, Copy)]
pub enum WrapperSubject<'a> {
Op(&'a ParsedOp),
Phase(&'a WorkloadPhase),
}
impl<'a> WrapperSubject<'a> {
pub fn level(&self) -> WrapperLevel {
match self {
WrapperSubject::Op(_) => WrapperLevel::Op,
WrapperSubject::Phase(_) => WrapperLevel::Phase,
}
}
pub fn op(&self) -> Option<&'a ParsedOp> {
match self {
WrapperSubject::Op(op) => Some(op),
_ => None,
}
}
pub fn phase(&self) -> Option<&'a WorkloadPhase> {
match self {
WrapperSubject::Phase(p) => Some(p),
_ => None,
}
}
pub fn has_owned_field(&self, field: &str) -> bool {
match self {
WrapperSubject::Op(op) => match field {
"if" => op.condition.is_some(),
"delay" => op.delay.is_some(),
"while" => op.while_cond.is_some(),
"rate" => op.rate.is_some(),
_ => op.params.contains_key(field),
},
WrapperSubject::Phase(p) => match field {
"interval" => p.interval.is_some(),
"repeat" => p.repeat.is_some(),
_ => false,
},
}
}
}
pub struct WrapperRegistration {
pub name: WrapperName,
pub owned_fields: &'static [&'static str],
pub triggers: fn(WrapperSubject) -> bool,
pub requires_inner: &'static [WrapperName],
pub forbids_outer: &'static [WrapperName],
pub mutually_exclusive_with: &'static [WrapperName],
pub describe_assignment: fn(WrapperSubject) -> Option<String>,
pub levels: &'static [WrapperLevel],
}
inventory::collect!(WrapperRegistration);
impl WrapperRegistration {
pub fn applies_at(&self, level: WrapperLevel) -> bool {
self.levels.contains(&level)
}
}
#[cfg(test)]
mod level_tests {
use super::*;
fn no_trigger(_: WrapperSubject) -> bool {
false
}
fn no_describe(_: WrapperSubject) -> Option<String> {
None
}
#[test]
fn applies_at_reads_declared_levels() {
let reg = WrapperRegistration {
name: WrapperName::new("t"),
owned_fields: &[],
triggers: no_trigger,
requires_inner: &[],
forbids_outer: &[],
mutually_exclusive_with: &[],
describe_assignment: no_describe,
levels: &[WrapperLevel::Op, WrapperLevel::Phase],
};
assert!(reg.applies_at(WrapperLevel::Op));
assert!(reg.applies_at(WrapperLevel::Phase));
assert!(!reg.applies_at(WrapperLevel::Session));
}
}
pub struct WrapperRegistry {
entries: Vec<&'static WrapperRegistration>,
}
impl WrapperRegistry {
pub fn from_inventory() -> Self {
let mut entries: Vec<&'static WrapperRegistration> =
inventory::iter::<WrapperRegistration>().collect();
entries.sort_by_key(|r| r.name);
Self { entries }
}
pub fn iter(&self) -> impl Iterator<Item = &'static WrapperRegistration> + '_ {
self.entries.iter().copied()
}
pub fn owns_field(&self, field: &str) -> bool {
self.iter().any(|reg| reg.owned_fields.contains(&field))
}
pub fn all_owned_fields(&self) -> std::collections::BTreeSet<&'static str> {
self.iter()
.flat_map(|reg| reg.owned_fields.iter().copied())
.collect()
}
pub fn get(&self, name: WrapperName) -> Option<&'static WrapperRegistration> {
self.entries.iter().copied().find(|r| r.name == name)
}
pub fn get_str(&self, name: &str) -> Option<&'static WrapperRegistration> {
self.entries
.iter()
.copied()
.find(|r| r.name.as_str() == name)
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn closest_match(&self, query: &str) -> Option<&'static str> {
closest_match(query, self.entries.iter().map(|r| r.name.as_str()))
}
pub fn misplaced_fields(&self, subject: WrapperSubject) -> Vec<(WrapperName, &'static str)> {
let mut out: Vec<(WrapperName, &'static str)> = Vec::new();
for reg in self.iter() {
if !reg.applies_at(subject.level()) || (reg.triggers)(subject) {
continue;
}
for &field in reg.owned_fields {
if subject.has_owned_field(field) {
out.push((reg.name, field));
}
}
}
out
}
}
pub fn closest_match<'a>(
query: &str,
candidates: impl IntoIterator<Item = &'a str>,
) -> Option<&'a str> {
let mut best: Option<(&str, usize)> = None;
for c in candidates {
let d = levenshtein(query, c);
match best {
Some((_, prev)) if d >= prev => {}
_ => best = Some((c, d)),
}
}
best.filter(|&(_, d)| d <= 3).map(|(s, _)| s)
}
fn levenshtein(a: &str, b: &str) -> usize {
let a: Vec<char> = a.chars().collect();
let b: Vec<char> = b.chars().collect();
let (n, m) = (a.len(), b.len());
if n == 0 {
return m;
}
if m == 0 {
return n;
}
let mut prev: Vec<usize> = (0..=m).collect();
let mut curr: Vec<usize> = vec![0; m + 1];
for i in 1..=n {
curr[0] = i;
for j in 1..=m {
let cost = if a[i - 1] == b[j - 1] { 0 } else { 1 };
curr[j] = (prev[j] + 1).min(curr[j - 1] + 1).min(prev[j - 1] + cost);
}
std::mem::swap(&mut prev, &mut curr);
}
prev[m]
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn levenshtein_basic() {
assert_eq!(levenshtein("", "abc"), 3);
assert_eq!(levenshtein("abc", ""), 3);
assert_eq!(levenshtein("kitten", "sitting"), 3);
assert_eq!(levenshtein("validate", "validatte"), 1);
}
#[test]
fn closest_match_finds_typo() {
let names = ["validate", "poll", "delay"];
assert_eq!(closest_match("validatte", names), Some("validate"));
assert_eq!(closest_match("plll", names), Some("poll"));
assert_eq!(closest_match("wildly_different", names), None);
}
}