1const MAX_CAPTURES: usize = 32;
10const MAXCCALLS: u32 = 200;
12const MAXCCALLS_51: u32 = 5000;
16
17#[derive(Clone, Copy, PartialEq, Eq, Debug)]
19pub enum Cap {
20 Span(usize, usize),
22 Pos(usize),
24}
25
26#[derive(Debug)]
29pub struct PatError(
30 pub String,
32);
33
34pub struct Match {
36 pub start: usize,
38 pub end: usize,
40 pub caps: Vec<Cap>,
42}
43
44#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Debug)]
46pub(crate) enum Flavor {
47 Lua51,
50 Lua52,
53 Lua53,
55}
56
57const CAP_UNFINISHED: isize = -1;
58const CAP_POSITION: isize = -2;
59
60#[derive(Clone, Copy, PartialEq, Eq, Debug)]
62pub(crate) enum CapValue {
63 Span(usize, usize),
64 Pos(usize),
65}
66
67pub(crate) struct MatchState<'a> {
70 src: &'a [u8],
71 pat: &'a [u8],
72 flavor: Flavor,
73 level: usize,
74 capture: [(usize, isize); MAX_CAPTURES],
75 matchdepth: u32,
76}
77
78fn err<T>(msg: &str) -> Result<T, PatError> {
79 Err(PatError(msg.to_string()))
80}
81
82impl<'a> MatchState<'a> {
83 pub(crate) fn new(src: &'a [u8], pat: &'a [u8], flavor: Flavor) -> Self {
84 MatchState {
85 src,
86 pat,
87 flavor,
88 level: 0,
89 capture: [(0, 0); MAX_CAPTURES],
90 matchdepth: if flavor == Flavor::Lua51 {
91 MAXCCALLS_51
92 } else {
93 MAXCCALLS
94 },
95 }
96 }
97
98 pub(crate) fn try_at(&mut self, s: usize) -> Result<Option<usize>, PatError> {
100 self.level = 0;
101 self.do_match(s, 0)
102 }
103
104 pub(crate) fn level(&self) -> usize {
106 self.level
107 }
108
109 pub(crate) fn get_capture(&self, i: usize, s: usize, e: usize) -> Result<CapValue, PatError> {
112 if i >= self.level {
113 if i != 0 {
114 return match self.flavor {
115 Flavor::Lua53 => Err(PatError(format!("invalid capture index %{}", i + 1))),
116 _ => err("invalid capture index"),
117 };
118 }
119 return Ok(CapValue::Span(s, e));
120 }
121 let (init, len) = self.capture[i];
122 match len {
123 CAP_UNFINISHED => err("unfinished capture"),
124 CAP_POSITION => Ok(CapValue::Pos(init)),
125 _ => Ok(CapValue::Span(init, init + len as usize)),
126 }
127 }
128
129 fn do_match(&mut self, s: usize, p: usize) -> Result<Option<usize>, PatError> {
130 if self.matchdepth == 0 {
131 return err("pattern too complex");
132 }
133 self.matchdepth -= 1;
134 let r = self.match_body(s, p);
135 self.matchdepth += 1;
136 r
137 }
138
139 fn match_body(&mut self, mut s: usize, mut p: usize) -> Result<Option<usize>, PatError> {
140 let pat = self.pat;
141 loop {
142 if p == pat.len() {
143 return Ok(Some(s));
144 }
145 match pat[p] {
146 b'(' => {
147 return if pat.get(p + 1) == Some(&b')') {
148 self.start_capture(s, p + 2, CAP_POSITION)
149 } else {
150 self.start_capture(s, p + 1, CAP_UNFINISHED)
151 };
152 }
153 b')' => return self.end_capture(s, p + 1),
154 b'$' if p + 1 == pat.len() => {
155 return Ok((s == self.src.len()).then_some(s));
156 }
157 b'%' => match pat.get(p + 1) {
158 Some(b'b') => match self.match_balance(s, p + 2)? {
159 Some(ns) => {
160 s = ns;
161 p += 4;
162 continue;
163 }
164 None => return Ok(None),
165 },
166 Some(b'f') => {
167 p += 2;
168 if pat.get(p) != Some(&b'[') {
169 return err("missing '[' after '%f' in pattern");
170 }
171 let ep = self.class_end(p)?;
172 let prev = if s == 0 { 0 } else { self.src[s - 1] };
173 let cur = self.src.get(s).copied().unwrap_or(0);
175 if !self.match_bracket(prev, p, ep - 1)
176 && self.match_bracket(cur, p, ep - 1)
177 {
178 p = ep;
179 continue;
180 }
181 return Ok(None);
182 }
183 Some(&d) if d.is_ascii_digit() => match self.match_capture(s, d)? {
184 Some(ns) => {
185 s = ns;
186 p += 2;
187 continue;
188 }
189 None => return Ok(None),
190 },
191 _ => {}
192 },
193 _ => {}
194 }
195 let ep = self.class_end(p)?;
197 let suffix = pat.get(ep).copied();
198 if !self.single_match(s, p, ep) {
199 if matches!(suffix, Some(b'*' | b'?' | b'-')) {
200 p = ep + 1;
201 continue;
202 }
203 return Ok(None);
204 }
205 match suffix {
206 Some(b'?') => {
207 if let Some(r) = self.do_match(s + 1, ep + 1)? {
208 return Ok(Some(r));
209 }
210 p = ep + 1;
211 }
212 Some(b'+') => return self.max_expand(s + 1, p, ep),
213 Some(b'*') => return self.max_expand(s, p, ep),
214 Some(b'-') => return self.min_expand(s, p, ep),
215 _ => {
216 s += 1;
217 p = ep;
218 }
219 }
220 }
221 }
222
223 fn class_end(&self, p: usize) -> Result<usize, PatError> {
224 let pat = self.pat;
225 match pat[p] {
226 b'%' => {
227 if p + 1 == pat.len() {
228 return err("malformed pattern (ends with '%')");
229 }
230 Ok(p + 2)
231 }
232 b'[' => {
233 let mut q = p + 1;
234 if pat.get(q) == Some(&b'^') {
235 q += 1;
236 }
237 loop {
240 if q == pat.len() {
241 return err("malformed pattern (missing ']')");
242 }
243 let c = pat[q];
244 q += 1;
245 if c == b'%' && q < pat.len() {
246 q += 1;
247 }
248 if pat.get(q) == Some(&b']') {
249 return Ok(q + 1);
250 }
251 }
252 }
253 _ => Ok(p + 1),
254 }
255 }
256
257 fn match_class(&self, c: u8, cl: u8) -> bool {
258 let res = match cl.to_ascii_lowercase() {
259 b'a' => c.is_ascii_alphabetic(),
260 b'c' => c.is_ascii_control(),
261 b'd' => c.is_ascii_digit(),
262 b'g' if self.flavor >= Flavor::Lua52 => c.is_ascii_graphic(),
263 b'l' => c.is_ascii_lowercase(),
264 b'p' => c.is_ascii_punctuation(),
265 b's' => matches!(c, b' ' | b'\t' | b'\n' | 0x0B | 0x0C | b'\r'),
266 b'u' => c.is_ascii_uppercase(),
267 b'w' => c.is_ascii_alphanumeric(),
268 b'x' => c.is_ascii_hexdigit(),
269 b'z' => c == 0,
270 _ => return cl == c,
271 };
272 if cl.is_ascii_lowercase() { res } else { !res }
273 }
274
275 fn match_bracket(&self, c: u8, mut p: usize, ec: usize) -> bool {
277 let pat = self.pat;
278 let mut sig = true;
279 if pat[p + 1] == b'^' {
280 sig = false;
281 p += 1;
282 }
283 loop {
284 p += 1;
285 if p >= ec {
286 return !sig;
287 }
288 if pat[p] == b'%' {
289 p += 1;
290 if self.match_class(c, pat[p]) {
291 return sig;
292 }
293 } else if pat[p + 1] == b'-' && p + 2 < ec {
294 p += 2;
295 if pat[p - 2] <= c && c <= pat[p] {
296 return sig;
297 }
298 } else if pat[p] == c {
299 return sig;
300 }
301 }
302 }
303
304 fn single_match(&self, s: usize, p: usize, ep: usize) -> bool {
305 let Some(&c) = self.src.get(s) else {
306 return false;
307 };
308 match self.pat[p] {
309 b'.' => true,
310 b'%' => self.match_class(c, self.pat[p + 1]),
311 b'[' => self.match_bracket(c, p, ep - 1),
312 pc => pc == c,
313 }
314 }
315
316 fn match_balance(&self, s: usize, p: usize) -> Result<Option<usize>, PatError> {
317 if p + 1 >= self.pat.len() {
318 return if self.flavor == Flavor::Lua51 {
319 err("unbalanced pattern")
320 } else {
321 err("malformed pattern (missing arguments to '%b')")
322 };
323 }
324 let (b, e) = (self.pat[p], self.pat[p + 1]);
325 if self.src.get(s) != Some(&b) {
326 return Ok(None);
327 }
328 let mut cont = 1;
329 for (i, &c) in self.src.iter().enumerate().skip(s + 1) {
330 if c == e {
331 cont -= 1;
332 if cont == 0 {
333 return Ok(Some(i + 1));
334 }
335 } else if c == b {
336 cont += 1;
337 }
338 }
339 Ok(None)
340 }
341
342 fn max_expand(&mut self, s: usize, p: usize, ep: usize) -> Result<Option<usize>, PatError> {
343 let mut i = 0;
344 while self.single_match(s + i, p, ep) {
345 i += 1;
346 }
347 loop {
348 if let Some(r) = self.do_match(s + i, ep + 1)? {
349 return Ok(Some(r));
350 }
351 if i == 0 {
352 return Ok(None);
353 }
354 i -= 1;
355 }
356 }
357
358 fn min_expand(&mut self, mut s: usize, p: usize, ep: usize) -> Result<Option<usize>, PatError> {
359 loop {
360 if let Some(r) = self.do_match(s, ep + 1)? {
361 return Ok(Some(r));
362 }
363 if self.single_match(s, p, ep) {
364 s += 1;
365 } else {
366 return Ok(None);
367 }
368 }
369 }
370
371 fn start_capture(
372 &mut self,
373 s: usize,
374 p: usize,
375 what: isize,
376 ) -> Result<Option<usize>, PatError> {
377 if self.level >= MAX_CAPTURES {
378 return err("too many captures");
379 }
380 self.capture[self.level] = (s, what);
381 self.level += 1;
382 let r = self.do_match(s, p)?;
383 if r.is_none() {
384 self.level -= 1;
385 }
386 Ok(r)
387 }
388
389 fn end_capture(&mut self, s: usize, p: usize) -> Result<Option<usize>, PatError> {
390 let l = self.capture_to_close()?;
391 self.capture[l].1 = (s - self.capture[l].0) as isize;
392 let r = self.do_match(s, p)?;
393 if r.is_none() {
394 self.capture[l].1 = CAP_UNFINISHED;
395 }
396 Ok(r)
397 }
398
399 fn capture_to_close(&self) -> Result<usize, PatError> {
400 (0..self.level)
401 .rev()
402 .find(|&l| self.capture[l].1 == CAP_UNFINISHED)
403 .map_or_else(|| err("invalid pattern capture"), Ok)
404 }
405
406 fn match_capture(&self, s: usize, d: u8) -> Result<Option<usize>, PatError> {
409 let l = d as isize - b'1' as isize;
410 if l < 0 || l as usize >= self.level || self.capture[l as usize].1 == CAP_UNFINISHED {
411 return if self.flavor == Flavor::Lua51 {
412 err("invalid capture index")
413 } else {
414 Err(PatError(format!("invalid capture index %{}", l + 1)))
415 };
416 }
417 let (init, len) = self.capture[l as usize];
418 let len = len as usize;
419 if self.src.len() - s >= len && self.src[init..init + len] == self.src[s..s + len] {
420 Ok(Some(s + len))
421 } else {
422 Ok(None)
423 }
424 }
425}
426
427pub fn anchor_split(pat: &[u8]) -> (bool, &[u8]) {
431 match pat.first() {
432 Some(b'^') => (true, &pat[1..]),
433 _ => (false, pat),
434 }
435}
436
437pub fn match_at(src: &[u8], pat_body: &[u8], s: usize) -> Result<Option<Match>, PatError> {
441 let mut ms = MatchState::new(src, pat_body, Flavor::Lua53);
442 let Some(e) = ms.try_at(s)? else {
443 return Ok(None);
444 };
445 let caps = (0..ms.level())
446 .map(|i| {
447 ms.get_capture(i, s, e).map(|c| match c {
448 CapValue::Span(a, b) => Cap::Span(a, b),
449 CapValue::Pos(p) => Cap::Pos(p),
450 })
451 })
452 .collect::<Result<Vec<_>, _>>()?;
453 Ok(Some(Match {
454 start: s,
455 end: e,
456 caps,
457 }))
458}
459
460pub fn find(src: &[u8], pat: &[u8], init: usize) -> Result<Option<Match>, PatError> {
463 if init > src.len() {
464 return Ok(None);
465 }
466 let (anchor, pat_body) = anchor_split(pat);
467 let mut s = init;
468 loop {
469 if let Some(m) = match_at(src, pat_body, s)? {
470 return Ok(Some(m));
471 }
472 if anchor || s >= src.len() {
473 return Ok(None);
474 }
475 s += 1;
476 }
477}
478
479pub fn has_specials(pat: &[u8]) -> bool {
482 pat.iter().any(|c| {
483 matches!(
484 c,
485 b'^' | b'$' | b'*' | b'+' | b'?' | b'.' | b'(' | b'[' | b'%' | b'-'
486 )
487 })
488}
489
490pub fn plain_find(hay: &[u8], needle: &[u8], init: usize) -> Option<usize> {
492 if init > hay.len() {
493 return None;
494 }
495 if needle.is_empty() {
496 return Some(init);
497 }
498 hay[init..]
499 .windows(needle.len())
500 .position(|w| w == needle)
501 .map(|i| i + init)
502}
503
504#[cfg(test)]
505mod tests {
506 use super::*;
507
508 fn m(src: &str, pat: &str) -> Option<(usize, usize)> {
509 find(src.as_bytes(), pat.as_bytes(), 0)
510 .unwrap()
511 .map(|m| (m.start, m.end))
512 }
513
514 #[test]
515 fn basics() {
516 assert_eq!(m("hello", "l+"), Some((2, 4)));
517 assert_eq!(m("hello", "^h"), Some((0, 1)));
518 assert_eq!(m("hello", "^e"), None);
519 assert_eq!(m("hello", "o$"), Some((4, 5)));
520 assert_eq!(m("hello", "%a+"), Some((0, 5)));
521 assert_eq!(m("a1b2", "%d"), Some((1, 2)));
522 assert_eq!(m("abc", "a.c"), Some((0, 3)));
523 assert_eq!(m("", ".*"), Some((0, 0)));
524 assert_eq!(m("abc", "x*"), Some((0, 0)));
525 }
526
527 #[test]
528 fn sets_and_quantifiers() {
529 assert_eq!(m("hello world", "[aeiou]"), Some((1, 2)));
530 assert_eq!(m("hello", "[^aeiou]+"), Some((0, 1)));
531 assert_eq!(m("x123y", "[0-9]+"), Some((1, 4)));
532 assert_eq!(m("aaa", "a-"), Some((0, 0)));
533 assert_eq!(m("<a><b>", "<.->"), Some((0, 3)));
534 assert_eq!(m("<a><b>", "<.*>"), Some((0, 6)));
535 assert_eq!(m("abc", "ab?c"), Some((0, 3)));
536 assert_eq!(m("ac", "ab?c"), Some((0, 2)));
537 }
538
539 #[test]
540 fn captures_and_specials() {
541 let mm = find(b"key=value", b"(%w+)=(%w+)", 0).unwrap().unwrap();
542 assert_eq!(mm.caps.len(), 2);
543 assert_eq!(mm.caps[0], Cap::Span(0, 3));
544 assert_eq!(mm.caps[1], Cap::Span(4, 9));
545 let mm = find(b"abc", b"a()b", 0).unwrap().unwrap();
547 assert_eq!(mm.caps[0], Cap::Pos(1));
548 assert_eq!(m("(foo(bar))baz", "%b()"), Some((0, 10)));
550 assert_eq!(m("THE (quick) fox", "%f[%a]%a+"), Some((0, 3)));
552 assert_eq!(m("abcabc", "(abc)%1"), Some((0, 6)));
554 assert_eq!(m("abcabd", "(abc)%1"), None);
555 }
556
557 #[test]
558 fn errors() {
559 assert!(find(b"x", b"%", 0).is_err());
560 assert!(find(b"x", b"[abc", 0).is_err());
561 assert!(find(b"a", b"(a", 0).is_err()); assert!(find(b"x", b"%1", 0).is_err());
563 }
564}