1use crate::{Domain, Known, Rule, Test, Then};
2use crate::trace::Firing;
3
4/// Where a node stands, given what is known.
5#[derive(Clone, Copy, Debug, PartialEq, Eq)]
6pub enum State {
7    /// Every test up to here holds.
8    Holds,
9    /// A test up to here is known to fail. Nothing downstream can fire.
10    Fails,
11    /// No test has failed, and at least one is on a fact not yet known.
12    Waits,
13}
14
15impl State {
16    fn and(self, other: State) -> State {
17        match (self, other) {
18            (State::Fails, _) | (_, State::Fails) => State::Fails,
19            (State::Holds, State::Holds) => State::Holds,
20            _ => State::Waits,
21        }
22    }
23}
24
25/// A join node: everything `left` requires, and one more test. `left` is
26/// another join, or `None` when this is a rule's first test.
27#[derive(Clone, Copy, Debug, PartialEq, Eq)]
28pub struct Join {
29    pub left: Option<usize>,
30    /// Index into [`Network::alphas`].
31    pub alpha: usize,
32    /// How many tests are joined here: 1 for a rule's first.
33    pub depth: usize,
34}
35
36/// A rule's place in the network: the join that holds when the rule does.
37#[derive(Clone, Debug)]
38pub struct Terminal<D: Domain> {
39    pub name: String,
40    pub join: usize,
41    pub then: Then<D>,
42}
43
44/// The rules as a Rete network.
45#[derive(Clone, Debug)]
46pub struct Network<D: Domain> {
47    /// One per distinct test, however many rules use it.
48    pub alphas: Vec<Test<D>>,
49    /// Shared by every rule that begins with the same tests.
50    pub joins: Vec<Join>,
51    /// In rule order, which is priority.
52    pub terminals: Vec<Terminal<D>>,
53}
54
55/// What to do next.
56#[derive(Clone, Debug, PartialEq, Eq)]
57pub enum Next<D: Domain> {
58    /// Ask the sources for all of these, in one request.
59    Ask(Vec<D::Fact>),
60    Do(D::Effect),
61    /// A rule ends the query outright.
62    End(D::End),
63    /// No rule has anything left to do and nothing more can be learned: the
64    /// query is over, and what to show is in the notes of the rules that hold.
65    Done,
66}
67
68impl<D: Domain> Network<D> {
69    /// Merges the rules: one alpha per distinct test, and for each rule a
70    /// chain of joins, shared by rules that begin with the same tests.
71    ///
72    /// # Panics
73    ///
74    /// On a rule with no tests, which holds always and so cannot be placed
75    /// in a network of tests. [`crate::load`] refuses one before it gets here.
76    pub fn compile(rules: &[Rule<'_, D>]) -> Self {
77        let mut network = Network { alphas: Vec::new(), joins: Vec::new(), terminals: Vec::new() };
78        for rule in rules {
79            let mut left = None;
80            for (depth, test) in rule.when.iter().enumerate() {
81                let alpha = network.alphas.iter().position(|known| known == test).unwrap_or_else(|| {
82                    network.alphas.push(*test);
83                    network.alphas.len() - 1
84                });
85                let join = Join { left, alpha, depth: depth + 1 };
86                left = Some(network.joins.iter().position(|known| *known == join).unwrap_or_else(|| {
87                    network.joins.push(join);
88                    network.joins.len() - 1
89                }));
90            }
91            let join = left.expect("a rule has at least one test");
92            network.terminals.push(Terminal { name: rule.name.to_owned(), join, then: rule.then });
93        }
94        network
95    }
96
97    pub fn alpha(&self, alpha: usize, known: &Known<D>) -> State {
98        let test = self.alphas[alpha];
99        match (test, known.get(test.fact())) {
100            (_, None) => State::Waits,
101            (Test::Known(_), Some(_)) => State::Holds,
102            (Test::Is(_, wanted), Some(value)) if wanted == value => State::Holds,
103            (Test::Is(..), Some(_)) => State::Fails,
104        }
105    }
106
107    pub fn join(&self, join: usize, known: &Known<D>) -> State {
108        let Join { left, alpha, .. } = self.joins[join];
109        let here = self.alpha(alpha, known);
110        left.map_or(here, |left| self.join(left, known).and(here))
111    }
112
113    /// The facts a join's tests are on, first test first.
114    pub(crate) fn facts(&self, join: usize) -> Vec<D::Fact> {
115        let Join { left, alpha, .. } = self.joins[join];
116        let mut facts = left.map_or_else(Vec::new, |left| self.facts(left));
117        facts.push(self.alphas[alpha].fact());
118        facts
119    }
120
121    /// What to do, given what is known, and the rule that says so (its index
122    /// in [`Network::terminals`]), when one does.
123    ///
124    /// A rule that holds decides, in rule order; an effect whose fact is
125    /// already known has been done and is passed over. If no rule decides,
126    /// every fact a source supplies that a rule not yet failed is waiting on
127    /// is wanted, and they are asked for together.
128    pub fn decide(&self, known: &Known<D>) -> (Option<usize>, Next<D>) {
129        for (index, terminal) in self.terminals.iter().enumerate() {
130            if self.join(terminal.join, known) != State::Holds {
131                continue;
132            }
133            match terminal.then {
134                Then::End(end) => return (Some(index), Next::End(end)),
135                Then::Do(effect) if known.get(D::teaches(effect)).is_none() => {
136                    return (Some(index), Next::Do(effect));
137                }
138                Then::Do(_) | Then::Note(_) => {}
139            }
140        }
141        let mut wanted: Vec<D::Fact> = self
142            .terminals
143            .iter()
144            .filter(|terminal| self.join(terminal.join, known) == State::Waits)
145            .flat_map(|terminal| self.facts(terminal.join))
146            .filter(|fact| D::asked_for(*fact) && known.get(*fact).is_none())
147            .collect();
148        wanted.sort();
149        wanted.dedup();
150        (None, if wanted.is_empty() { Next::Done } else { Next::Ask(wanted) })
151    }
152
153    /// What to do, given what is known.
154    pub fn next(&self, known: &Known<D>) -> Next<D> {
155        self.decide(known).1
156    }
157
158    /// The rules that hold, in rule order.
159    pub fn held<'a>(&'a self, known: &'a Known<D>) -> impl Iterator<Item = &'a Terminal<D>> {
160        self.terminals.iter().filter(move |terminal| self.join(terminal.join, known) == State::Holds)
161    }
162
163    /// The names of the rules that hold, in rule order.
164    pub fn holding<'a>(&'a self, known: &'a Known<D>) -> Vec<&'a str> {
165        self.held(known).map(|terminal| terminal.name.as_str()).collect()
166    }
167
168    /// The joins and alphas that a rule that holds stands on: what, of
169    /// everything known, turned out to matter.
170    pub fn used(&self, known: &Known<D>) -> (Vec<bool>, Vec<bool>) {
171        let mut joins = vec![false; self.joins.len()];
172        let mut alphas = vec![false; self.alphas.len()];
173        for terminal in &self.terminals {
174            if self.join(terminal.join, known) != State::Holds {
175                continue;
176            }
177            let mut at = Some(terminal.join);
178            while let Some(join) = at {
179                joins[join] = true;
180                alphas[self.joins[join].alpha] = true;
181                at = self.joins[join].left;
182            }
183        }
184        (joins, alphas)
185    }
186
187    /// What rule `rule` (an index into [`Network::terminals`]) stands on,
188    /// given what is known: for a trace.
189    pub fn explain(&self, rule: usize, known: &Known<D>) -> Firing<D> {
190        let terminal = &self.terminals[rule];
191        let on = self.facts(terminal.join).into_iter().filter_map(|fact| known.get(fact).map(|value| (fact, value))).collect();
192        Firing { rule, name: terminal.name.clone(), then: terminal.then, on }
193    }
194}