competitive_library/string/
aho_corasick.rs1use 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}