competitive_library/string/
aho_corasick.rs

1use core::hash::Hash;
2use std::{
3    collections::{HashMap, VecDeque},
4    rc::Rc,
5};
6
7#[derive(Debug)]
8struct Node<T>
9where
10    T: Eq + Hash + Copy,
11{
12    to: HashMap<T, usize>,
13    keyword: Option<Rc<Vec<T>>>,
14    failure: usize,
15}
16impl<T> Node<T>
17where
18    T: Eq + Hash + Copy,
19{
20    pub fn new() -> Self {
21        Self {
22            to: HashMap::new(),
23            failure: 0,
24            keyword: None,
25        }
26    }
27}
28
29#[derive(Default, Debug)]
30pub struct AhoCorasick<T>
31where
32    T: Eq + Hash + Copy,
33{
34    nodes: Vec<Node<T>>,
35    prepared: bool,
36}
37impl<T> AhoCorasick<T>
38where
39    T: Eq + Hash + Copy,
40{
41    pub fn new() -> Self {
42        Self {
43            nodes: vec![Node::new()],
44            prepared: false,
45        }
46    }
47    pub fn add(&mut self, keyword: &[T]) {
48        assert!(!self.prepared);
49
50        let mut index = 0;
51
52        for value in keyword {
53            let buff = index;
54            index = match self.nodes[buff].to.get(value) {
55                Some(&index) => index,
56                None => {
57                    self.nodes.push(Node::new());
58                    let index = self.nodes.len() - 1;
59
60                    self.nodes[buff].to.insert(*value, index);
61                    index
62                }
63            };
64        }
65
66        self.nodes[index].keyword = Some(Rc::new(keyword.to_vec()));
67    }
68
69    pub fn make_failure_link(&mut self) {
70        self.prepared = true;
71        let mut queue = VecDeque::new();
72
73        queue.push_back(0);
74
75        while let Some(i) = queue.pop_front() {
76            for (value, &index) in self.nodes[i].to.clone().iter() {
77                if i != 0 {
78                    let mut buff = self.nodes[i].failure;
79                    self.nodes[index].failure = loop {
80                        if let Some(&next_index) = self.get_to(buff).get(value) {
81                            break next_index;
82                        } else {
83                            if buff == 0 {
84                                break 0;
85                            }
86                            buff = self.nodes[buff].failure
87                        };
88                    }
89                }
90                queue.push_back(index);
91            }
92        }
93    }
94    pub fn create_matcher<'a, 'b>(&'a self, target: &'b [T]) -> Matcher<'a, 'b, T> {
95        assert!(self.prepared);
96
97        Matcher::new(self, target)
98    }
99
100    pub(crate) fn get_to(&self, index: usize) -> &HashMap<T, usize> {
101        &self.nodes[index].to
102    }
103    pub(crate) fn get_failure(&self, index: usize) -> usize {
104        self.nodes[index].failure
105    }
106    pub(crate) fn get_keyword(&self, index: usize) -> Option<Rc<Vec<T>>> {
107        self.nodes[index].keyword.clone()
108    }
109}
110
111impl<T> Iterator for AhoCorasick<T>
112where
113    T: Eq + Hash + Copy,
114{
115    type Item = (Vec<T>, usize);
116    fn next(&mut self) -> Option<Self::Item> {
117        None
118    }
119}
120
121pub struct Matcher<'a, 'b, T>
122where
123    T: Eq + Hash + Copy,
124{
125    aho: &'a AhoCorasick<T>,
126    target: &'b [T],
127    target_index: usize,
128    aho_index: usize,
129    failure_index: usize,
130}
131impl<'a, 'b, T> Matcher<'a, 'b, T>
132where
133    T: Eq + Hash + Copy,
134{
135    pub fn new(aho: &'a AhoCorasick<T>, target: &'b [T]) -> Self {
136        Self {
137            aho,
138            target,
139            target_index: 0,
140            aho_index: 0,
141            failure_index: 0,
142        }
143    }
144}
145
146impl<'a, 'b, T> Iterator for Matcher<'a, 'b, T>
147where
148    T: Eq + Hash + Copy,
149{
150    type Item = (Rc<Vec<T>>, usize);
151    fn next(&mut self) -> Option<Self::Item> {
152        while self.target_index < self.target.len() {
153            while self.failure_index != 0 {
154                if let Some(keyword) = self.aho.get_keyword(self.failure_index) {
155                    self.failure_index = self.aho.get_failure(self.failure_index);
156                    return Some((keyword.clone(), self.target_index - keyword.len()));
157                }
158                self.failure_index = self.aho.get_failure(self.failure_index);
159            }
160
161            match self
162                .aho
163                .get_to(self.aho_index)
164                .get(&self.target[self.target_index])
165            {
166                Some(&index) => {
167                    self.target_index += 1;
168                    self.failure_index = self.aho.get_failure(index);
169                    self.aho_index = index;
170                }
171                None => {
172                    let buff = self.aho.get_failure(self.aho_index);
173                    if buff == self.aho_index {
174                        self.target_index += 1;
175                        continue;
176                    }
177                    self.aho_index = buff;
178                    continue;
179                }
180            };
181
182            if let Some(keyword) = self.aho.get_keyword(self.aho_index) {
183                return Some((keyword.clone(), self.target_index - keyword.len()));
184            }
185        }
186        while self.aho_index != 0 {
187            self.aho_index = self.aho.get_failure(self.aho_index);
188            if let Some(keyword) = self.aho.get_keyword(self.aho_index) {
189                return Some((keyword.clone(), self.target_index - keyword.len()));
190            }
191        }
192        None
193    }
194}
195
196#[cfg(test)]
197mod tests {
198    use std::ops::Deref;
199
200    use super::*;
201
202    #[test]
203    fn t() {
204        let mut aho = AhoCorasick::new();
205
206        aho.add(&"aho".to_string().chars().collect::<Vec<char>>());
207        aho.add(&"corasick".to_string().chars().collect::<Vec<char>>());
208        aho.add(&"aho-corasick".to_string().chars().collect::<Vec<char>>());
209
210        let s = "aho-corasick".to_string().chars().collect::<Vec<char>>();
211        aho.make_failure_link();
212        let mut m = aho.create_matcher(&s);
213        dbg!(&aho);
214        assert_eq!(m.next().unwrap().0.deref(), &['a', 'h', 'o']);
215
216        assert_eq!(
217            m.next().unwrap().0.deref(),
218            &['a', 'h', 'o', '-', 'c', 'o', 'r', 'a', 's', 'i', 'c', 'k']
219        );
220        assert_eq!(
221            m.next().unwrap().0.deref(),
222            &['c', 'o', 'r', 'a', 's', 'i', 'c', 'k']
223        );
224        assert_eq!(m.next(), None);
225    }
226
227    #[test]
228    fn aaa() {
229        let mut aho = AhoCorasick::new();
230
231        aho.add(&"a".to_string().chars().collect::<Vec<char>>());
232        aho.make_failure_link();
233        let s = "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"
234            .to_string()
235            .chars()
236            .collect::<Vec<char>>();
237        for (i, (keyword, index)) in aho.create_matcher(&s).enumerate() {
238            assert_eq!(keyword.deref(), &['a']);
239            assert_eq!(index, i);
240        }
241    }
242    #[test]
243    fn abc() {
244        let mut aho = AhoCorasick::new();
245
246        let s = "abcdefghijklmnopqrstuvwxyz"
247            .to_string()
248            .chars()
249            .collect::<Vec<char>>();
250
251        for i in 0..s.len() {
252            aho.add(&s.iter().skip(i).copied().collect::<Vec<char>>());
253        }
254
255        aho.make_failure_link();
256        for (i, (keyword, index)) in aho.create_matcher(&s).enumerate() {
257            assert_eq!(index, i);
258            assert_eq!(
259                keyword.deref(),
260                &s.iter().skip(i).copied().collect::<Vec<char>>()
261            )
262        }
263    }
264    #[test]
265    fn add() {
266        let mut aho = AhoCorasick::new();
267
268        let s = "zabcdefghijklmn".to_string().chars().collect::<Vec<char>>();
269
270        aho.add(&"abcd".to_string().chars().collect::<Vec<char>>());
271        aho.add(&"ijk".to_string().chars().collect::<Vec<char>>());
272        aho.add(&"ghi".to_string().chars().collect::<Vec<char>>());
273        aho.make_failure_link();
274        let mut m = aho.create_matcher(&s);
275
276        assert_eq!(m.next().unwrap().0.deref(), &['a', 'b', 'c', 'd']);
277        assert_eq!(m.next().unwrap().0.deref(), &['g', 'h', 'i']);
278        assert_eq!(m.next().unwrap().0.deref(), &['i', 'j', 'k']);
279        assert_eq!(m.next(), None);
280    }
281    #[test]
282    fn xbabcdex() {
283        let mut aho = AhoCorasick::new();
284
285        let s = "xbabcdex".to_string().chars().collect::<Vec<char>>();
286
287        aho.add(&"ab".to_string().chars().collect::<Vec<char>>());
288        aho.add(&"bc".to_string().chars().collect::<Vec<char>>());
289        aho.add(&"bab".to_string().chars().collect::<Vec<char>>());
290        aho.add(&"d".to_string().chars().collect::<Vec<char>>());
291        aho.add(&"abcde".to_string().chars().collect::<Vec<char>>());
292        aho.make_failure_link();
293        let mut m = aho.create_matcher(&s);
294
295        assert_eq!(m.next().unwrap().0.deref(), &['b', 'a', 'b']);
296        assert_eq!(m.next().unwrap().0.deref(), &['a', 'b']);
297
298        assert_eq!(m.next().unwrap().0.deref(), &['b', 'c']);
299        assert_eq!(m.next().unwrap().0.deref(), &['d']);
300        assert_eq!(m.next().unwrap().0.deref(), &['a', 'b', 'c', 'd', 'e']);
301        assert_eq!(m.next(), None);
302    }
303}