use crate::executor::ExecCtx;
use nmbrs_workload::model::ScenarioNode;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ScheduleSpec {
pub levels: Vec<ConcurrencyLimit>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ConcurrencyLimit {
Bounded(u32),
Unlimited,
}
impl ScheduleSpec {
pub fn default_serial() -> Self {
Self {
levels: vec![ConcurrencyLimit::Bounded(1)],
}
}
pub fn parse(s: &str) -> Result<Self, String> {
let s = s.trim();
if s.is_empty() {
return Err("schedule spec is empty".into());
}
let mut levels = Vec::new();
for (i, part) in s.split('/').enumerate() {
let part = part.trim();
let limit = match part {
"*" => ConcurrencyLimit::Unlimited,
_ => {
let n: u32 = part.parse().map_err(|_| {
format!("schedule spec level {i}: '{part}' is not a number or '*'")
})?;
if n == 0 {
return Err(format!(
"schedule spec level {i}: 0 is not a valid concurrency limit (use '1' for serial)"
));
}
ConcurrencyLimit::Bounded(n)
}
};
levels.push(limit);
}
Ok(Self { levels })
}
pub fn limit_at(&self, depth: usize) -> ConcurrencyLimit {
if self.levels.is_empty() {
return ConcurrencyLimit::Bounded(1);
}
let idx = depth.min(self.levels.len() - 1);
self.levels[idx]
}
#[allow(dead_code)] pub fn is_serial(&self) -> bool {
self.levels
.iter()
.all(|l| matches!(l, ConcurrencyLimit::Bounded(1)))
}
}
pub trait PhaseScheduler {
fn run<'a>(
&'a self,
ctx: &'a mut ExecCtx,
nodes: &'a [ScenarioNode],
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<(), String>> + Send + 'a>>;
}
#[derive(Debug, Default)]
pub struct TreeScheduler;
impl PhaseScheduler for TreeScheduler {
fn run<'a>(
&'a self,
ctx: &'a mut ExecCtx,
nodes: &'a [ScenarioNode],
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<(), String>> + Send + 'a>> {
crate::executor::execute_tree(ctx, nodes)
}
}
pub fn build(_spec: &ScheduleSpec) -> Box<dyn PhaseScheduler> {
Box::new(TreeScheduler)
}
#[cfg(test)]
fn format_spec(spec: &ScheduleSpec) -> String {
let parts: Vec<String> = spec
.levels
.iter()
.map(|l| match l {
ConcurrencyLimit::Bounded(n) => n.to_string(),
ConcurrencyLimit::Unlimited => "*".into(),
})
.collect();
parts.join("/")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_serial() {
let s = ScheduleSpec::parse("1").unwrap();
assert_eq!(s.levels, vec![ConcurrencyLimit::Bounded(1)]);
assert!(s.is_serial());
}
#[test]
fn parse_unlimited() {
let s = ScheduleSpec::parse("*").unwrap();
assert_eq!(s.levels, vec![ConcurrencyLimit::Unlimited]);
assert!(!s.is_serial());
}
#[test]
fn parse_multilevel() {
let s = ScheduleSpec::parse("1/4/*").unwrap();
assert_eq!(
s.levels,
vec![
ConcurrencyLimit::Bounded(1),
ConcurrencyLimit::Bounded(4),
ConcurrencyLimit::Unlimited,
]
);
}
#[test]
fn limit_at_extends_trailing() {
let s = ScheduleSpec::parse("1/4").unwrap();
assert_eq!(s.limit_at(0), ConcurrencyLimit::Bounded(1));
assert_eq!(s.limit_at(1), ConcurrencyLimit::Bounded(4));
assert_eq!(s.limit_at(2), ConcurrencyLimit::Bounded(4));
assert_eq!(s.limit_at(99), ConcurrencyLimit::Bounded(4));
}
#[test]
fn rejects_zero() {
let err = ScheduleSpec::parse("0").unwrap_err();
assert!(
err.contains("0 is not a valid concurrency limit"),
"unexpected: {err}"
);
}
#[test]
fn rejects_garbage() {
let err = ScheduleSpec::parse("foo").unwrap_err();
assert!(err.contains("not a number"), "unexpected: {err}");
}
#[test]
fn rejects_empty() {
let err = ScheduleSpec::parse("").unwrap_err();
assert!(err.contains("empty"), "unexpected: {err}");
}
#[test]
fn n_one_is_serial() {
let s = ScheduleSpec::parse("1/1").unwrap();
assert!(s.is_serial());
}
#[test]
fn default_serial_matches_explicit_one() {
let default = ScheduleSpec::default_serial();
let parsed = ScheduleSpec::parse("1").unwrap();
assert_eq!(default, parsed);
}
#[test]
fn build_returns_tree_scheduler_for_any_spec() {
let _ = build(&ScheduleSpec::default_serial());
let _ = build(&ScheduleSpec::parse("1/4/*").unwrap());
}
#[test]
fn format_spec_round_trip() {
let s = ScheduleSpec::parse("1/4/*").unwrap();
assert_eq!(format_spec(&s), "1/4/*");
}
}