use std::cell::RefCell;
use std::collections::{HashSet, VecDeque};
use std::rc::Rc;
use std::sync::Arc;
use std::time::{Duration, Instant};
use crate::testing::fdr::config::{Failure, FdrConfig, Trace};
use crate::testing::fdr::explorer::{MemoizationCache, RefinementChecker};
use crate::testing::specs::csp::{Action, Event, Process, State};
struct TimeoutChecker {
start: Instant,
timeout: Duration,
}
impl TimeoutChecker {
fn new(timeout_ms: u64) -> Self {
Self { start: Instant::now(), timeout: Duration::from_millis(timeout_ms) }
}
fn is_expired(&self) -> bool {
self.start.elapsed() >= self.timeout
}
}
pub struct DefaultRefinementChecker<'a, M>
where
M: MemoizationCache,
{
config: Arc<FdrConfig>,
process: &'a Process,
cache: Rc<RefCell<M>>,
}
impl<'a, M> DefaultRefinementChecker<'a, M>
where
M: MemoizationCache,
{
pub fn new(process: &'a Process, config: Arc<FdrConfig>, cache: Rc<RefCell<M>>) -> Self {
Self { config, process, cache }
}
pub fn config(&self) -> &FdrConfig {
&self.config
}
pub fn process(&self) -> &Process {
self.process
}
fn max_traces() -> usize {
<Self as RefinementChecker>::MAX_TRACES
}
fn max_queue_size() -> usize {
<Self as RefinementChecker>::MAX_QUEUE_SIZE
}
fn max_visited() -> usize {
<Self as RefinementChecker>::MAX_VISITED
}
fn check_limits(queue_len: usize, visited_len: usize, traces_count: usize) -> bool {
queue_len >= Self::max_queue_size() || visited_len >= Self::max_visited() || traces_count >= Self::max_traces()
}
fn is_stable_state(process: &Process, state: State) -> bool {
let enabled_actions = process.enabled(state);
!enabled_actions.iter().any(|action| process.hidden.contains(&action.event))
}
fn process_transition(process: &Process, event: &Event, trace: Trace, depth: usize) -> (Trace, usize) {
if process.hidden.contains(event) {
(trace, depth)
} else {
let mut new_trace = trace;
new_trace.push(*event);
(new_trace, depth + 1)
}
}
fn extract_linear_trace(process: &Process, max_depth: usize) -> Option<Trace> {
let mut trace = Vec::new();
let mut current_state = process.initial;
let mut visited_states = HashSet::new();
visited_states.insert(current_state);
loop {
if process.terminal.contains(¤t_state) {
return Some(trace);
}
if trace.len() >= max_depth {
return None; }
if visited_states.len() > 1000 {
return None; }
let mut has_tau_transitions = true;
while has_tau_transitions {
has_tau_transitions = false;
let current_enabled = process.enabled(current_state);
for action in ¤t_enabled {
if process.hidden.contains(&action.event) {
let next_states = process.step(current_state, &action.event);
if next_states.len() != 1 {
return None; }
current_state = next_states[0];
if !visited_states.insert(current_state) {
return None; }
has_tau_transitions = true;
break; }
}
if process.terminal.contains(¤t_state) {
return Some(trace);
}
}
let stable_enabled = process.enabled(current_state);
let observable_actions: Vec<_> = stable_enabled
.iter()
.filter(|action| !process.hidden.contains(&action.event))
.collect();
if observable_actions.len() > 1 {
return None; }
if let Some(action) = observable_actions.first() {
let next_states = process.step(current_state, &action.event);
if next_states.len() != 1 {
return None; }
trace.push(action.event);
current_state = next_states[0];
if !visited_states.insert(current_state) {
return None; }
} else {
if process.terminal.contains(¤t_state) {
return Some(trace);
}
return None; }
}
}
fn bfs_with_callbacks<T, FState, FTransition>(
&self,
process: &Process,
max_depth: usize,
mut data: T,
mut on_state: FState,
mut on_transition: FTransition,
) -> T
where
FState: FnMut(&mut T, State, &Trace, usize) -> bool,
FTransition: FnMut(&mut T, &mut VecDeque<(State, Trace, usize)>, State, Trace, usize, &Event, State),
{
let mut queue = VecDeque::new();
let mut visited = HashSet::new();
queue.push_back((process.initial, Vec::new(), 0usize));
while let Some((state, trace, depth)) = queue.pop_front() {
if queue.len() >= Self::max_queue_size() || visited.len() >= Self::max_visited() {
break;
}
if trace.len() >= max_depth {
continue;
}
let visit_key = (state, trace.clone());
if visited.contains(&visit_key) {
continue;
}
visited.insert(visit_key);
let skip_transitions = on_state(&mut data, state, &trace, depth);
if skip_transitions {
continue;
}
for action in process.enabled(state) {
for next_state in process.step(state, &action.event) {
on_transition(&mut data, &mut queue, state, trace.clone(), depth, &action.event, next_state);
}
}
}
data
}
fn has_tau_cycle(&self, tau_states_seen: &HashSet<(State, Trace)>, next_state: State, trace: &Trace) -> bool {
let next_key = (next_state, trace.clone());
tau_states_seen.contains(&next_key)
}
#[allow(clippy::too_many_arguments)]
fn process_observable_events(
spec: &Process,
state: State,
next_event: &Event,
enabled_actions: &[Action],
queue: &mut VecDeque<(State, usize)>,
visited: &mut HashSet<(State, usize)>,
trace_idx: usize,
max_queue_size: usize,
max_visited: usize,
timeout_checker: &TimeoutChecker,
) -> bool {
for action in enabled_actions {
if timeout_checker.is_expired() {
return false;
}
if &action.event == next_event {
for next_state in spec.step(state, &action.event) {
if timeout_checker.is_expired() {
return false;
}
if queue.len() >= max_queue_size || visited.len() >= max_visited {
break;
}
queue.push_back((next_state, trace_idx + 1));
}
}
}
true
}
#[allow(clippy::too_many_arguments)]
fn process_tau_transitions(
spec: &Process,
state: State,
enabled_actions: &[Action],
queue: &mut VecDeque<(State, usize)>,
visited: &mut HashSet<(State, usize)>,
trace_idx: usize,
max_queue_size: usize,
max_visited: usize,
timeout_checker: &TimeoutChecker,
) -> bool {
let mut tau_count = 0;
const MAX_TAU_PER_STATE: usize = 10;
for action in enabled_actions {
if timeout_checker.is_expired() {
return false;
}
if spec.hidden.contains(&action.event) {
if tau_count >= MAX_TAU_PER_STATE {
break; }
tau_count += 1;
for next_state in spec.step(state, &action.event) {
if timeout_checker.is_expired() {
return false;
}
if queue.len() >= max_queue_size || visited.len() >= max_visited {
break;
}
queue.push_back((next_state, trace_idx));
}
}
}
true
}
fn failure_exists_in_spec(spec_failures: &[Failure], impl_trace: &Trace, impl_refusal: &HashSet<Event>) -> bool {
for (spec_trace, spec_refusal) in spec_failures {
if spec_trace == impl_trace && impl_refusal.is_subset(spec_refusal) {
return true;
}
}
false
}
fn trace_exists_in_spec(
spec: &Process,
target_trace: &Trace,
max_depth: usize,
max_queue_size: usize,
max_visited: usize,
timeout_ms: u64,
) -> bool {
if target_trace.len() > max_depth {
return false;
}
let timeout_checker = TimeoutChecker::new(timeout_ms);
let mut queue = VecDeque::new();
let mut visited = HashSet::new();
queue.push_back((spec.initial, 0usize));
while let Some((state, trace_idx)) = queue.pop_front() {
if timeout_checker.is_expired() {
return false;
}
if queue.len() >= max_queue_size || visited.len() >= max_visited {
return false;
}
let visit_key = (state, trace_idx);
if visited.contains(&visit_key) {
continue;
}
visited.insert(visit_key);
if trace_idx >= target_trace.len() {
return true;
}
let next_event = &target_trace[trace_idx];
let enabled_actions = spec.enabled(state);
if !Self::process_observable_events(
spec,
state,
next_event,
&enabled_actions,
&mut queue,
&mut visited,
trace_idx,
max_queue_size,
max_visited,
&timeout_checker,
) {
return false;
}
if !Self::process_tau_transitions(
spec,
state,
&enabled_actions,
&mut queue,
&mut visited,
trace_idx,
max_queue_size,
max_visited,
&timeout_checker,
) {
return false;
}
}
false
}
}
impl<'a, M> RefinementChecker for DefaultRefinementChecker<'a, M>
where
M: MemoizationCache,
{
fn check_trace_refinement(&mut self, spec: &Process, impl_process: &Process) -> (bool, Option<Trace>) {
let impl_traces = self.compute_traces(impl_process, self.config.max_depth);
let longest_trace = impl_traces.iter().max_by_key(|t| t.len()).cloned();
if let Some(full_trace) = longest_trace {
let max_queue = Self::max_queue_size();
let max_visited = Self::max_visited();
if !Self::trace_exists_in_spec(
spec,
&full_trace,
self.config.max_depth,
max_queue,
max_visited,
self.config.timeout_ms,
) {
return (false, Some(full_trace));
}
(true, None)
} else {
(false, Some(Vec::new()))
}
}
fn check_failures_refinement(&mut self, spec: &Process, impl_process: &Process) -> (bool, Option<Failure>) {
let spec_failures = self.compute_failures(spec, self.config.max_depth);
let impl_failures = self.compute_failures(impl_process, self.config.max_depth);
for (impl_trace, impl_refusal) in &impl_failures {
if !<DefaultRefinementChecker<'a, M>>::failure_exists_in_spec(&spec_failures, impl_trace, impl_refusal) {
return (false, Some((impl_trace.clone(), impl_refusal.clone())));
}
}
(true, None)
}
fn check_divergence_refinement(&mut self, spec: &Process, impl_process: &Process) -> (bool, Option<Trace>) {
let spec_divergences = self.compute_divergences(spec, self.config.max_depth);
let impl_divergences = self.compute_divergences(impl_process, self.config.max_depth);
for impl_div in &impl_divergences {
if !spec_divergences.contains(impl_div) {
return (false, Some(impl_div.clone()));
}
}
(true, None)
}
fn compute_traces(&mut self, process: &Process, max_depth: usize) -> HashSet<Trace> {
if let Some(cached) = self.cache.borrow().get_cached_traces(process.name) {
return cached.into_iter().collect();
}
if let Some(linear_trace) = Self::extract_linear_trace(process, max_depth) {
let mut traces = HashSet::new();
traces.insert(linear_trace);
return traces;
}
let timeout_checker = TimeoutChecker::new(self.config.timeout_ms);
let mut traces = HashSet::new();
traces.insert(Vec::new());
let mut queue = VecDeque::new();
let mut visited = HashSet::new();
queue.push_back((process.initial, Vec::new(), 0usize));
while let Some((state, trace, depth)) = queue.pop_front() {
if timeout_checker.is_expired() {
break;
}
if Self::check_limits(queue.len(), visited.len(), traces.len()) {
break;
}
if trace.len() >= max_depth {
continue;
}
let visit_key = (state, trace.clone());
if visited.contains(&visit_key) {
continue;
}
visited.insert(visit_key);
let enabled_actions = process.enabled(state);
for action in enabled_actions {
let next_states = process.step(state, &action.event);
for next_state in next_states {
if Self::check_limits(queue.len(), visited.len(), traces.len()) {
break;
}
let (new_trace, new_depth) = Self::process_transition(process, &action.event, trace.clone(), depth);
if !process.hidden.contains(&action.event) {
traces.insert(new_trace.clone());
}
queue.push_back((next_state, new_trace, new_depth));
}
}
}
let traces_vec: Vec<Trace> = traces.iter().cloned().collect();
self.cache.borrow_mut().cache_traces(process.name.to_string(), traces_vec);
traces
}
fn compute_failures(&mut self, process: &Process, max_depth: usize) -> Vec<Failure> {
if let Some(cached) = self.cache.borrow().get_cached_failures(process.name) {
return cached;
}
let failures = Vec::new();
let visited = HashSet::new();
let data = (failures, visited);
let (failures, _) = self.bfs_with_callbacks(
process,
max_depth,
data,
|(failures, visited), state, trace, _depth| {
let visit_key = (state, trace.clone());
if visited.contains(&visit_key) {
return true;
}
visited.insert(visit_key);
if Self::is_stable_state(process, state) {
let refusals = self.compute_refusals(process, state);
let failure = (trace.clone(), refusals);
if !failures.contains(&failure) {
failures.push(failure);
}
}
false
},
|(_failures, _visited), queue, _state, trace, depth, event, next_state| {
let (new_trace, new_depth) = Self::process_transition(process, event, trace, depth);
queue.push_back((next_state, new_trace, new_depth));
},
);
self.cache
.borrow_mut()
.cache_failures(process.name.to_string(), failures.clone());
failures
}
fn compute_divergences(&mut self, process: &Process, max_depth: usize) -> HashSet<Trace> {
if let Some(cached) = self.cache.borrow().get_cached_divergences(process.name) {
return cached.into_iter().collect();
}
let mut divergences = HashSet::new();
let mut queue = VecDeque::new();
let mut initial_tau_seen = HashSet::new();
initial_tau_seen.insert((process.initial, Vec::new()));
queue.push_back((process.initial, Vec::new(), initial_tau_seen));
let mut global_visited = HashSet::new();
while let Some((state, trace, tau_states_seen)) = queue.pop_front() {
if trace.len() >= max_depth {
continue;
}
let visit_key = (state, trace.clone());
if global_visited.contains(&visit_key) {
continue;
}
for action in process.enabled(state) {
for next_state in process.step(state, &action.event) {
if process.hidden.contains(&action.event) {
if self.has_tau_cycle(&tau_states_seen, next_state, &trace) {
divergences.insert(trace.clone());
continue;
}
let mut new_tau_seen = tau_states_seen.clone();
new_tau_seen.insert((next_state, trace.clone()));
queue.push_back((next_state, trace.clone(), new_tau_seen));
} else {
let mut new_trace = trace.clone();
new_trace.push(action.event);
let mut new_tau_seen = HashSet::new();
new_tau_seen.insert((next_state, new_trace.clone()));
queue.push_back((next_state, new_trace, new_tau_seen));
}
}
}
global_visited.insert(visit_key);
}
let divergences_vec: Vec<Trace> = divergences.iter().cloned().collect();
self.cache
.borrow_mut()
.cache_divergences(process.name.to_string(), divergences_vec);
divergences
}
fn compute_refusals(&self, process: &Process, state: State) -> HashSet<Event> {
let enabled_events: HashSet<Event> = process
.enabled(state)
.iter()
.filter_map(|action| {
if !process.hidden.contains(&action.event) {
Some(action.event)
} else {
None
}
})
.collect();
process
.observable
.iter()
.filter(|&event| !enabled_events.contains(event))
.cloned()
.collect()
}
}