Skip to main content

luna_core/
pattern.rs

1//! Lua pattern matching engine — a port of lstrlib.c's matcher.
2//! Pure functions over byte slices (stone candidate: no runtime types).
3//!
4//! The matcher is the 5.2+ one (explicit pattern end, `matchdepth` bound);
5//! `Flavor` carries the few places where older dialects differ. 5.1's
6//! NUL-terminated patterns are the caller's business: it passes the pattern
7//! cut at the first NUL.
8
9const MAX_CAPTURES: usize = 32;
10/// PUC `MAXCCALLS`: nested `match` calls allowed before "pattern too complex".
11const MAXCCALLS: u32 = 200;
12/// 5.1 has no `matchdepth`; its matcher recurses until the C stack runs out.
13/// This bound stands in for that stack so a runaway pattern raises instead
14/// of overflowing ours.
15const MAXCCALLS_51: u32 = 5000;
16
17/// One capture produced by a successful pattern match.
18#[derive(Clone, Copy, PartialEq, Eq, Debug)]
19pub enum Cap {
20    /// captured span [start, end) in source bytes
21    Span(usize, usize),
22    /// position capture `()` — byte offset (0-based; callers add 1)
23    Pos(usize),
24}
25
26/// Error returned by the pattern matcher (malformed pattern, runaway depth,
27/// invalid `%f` frontier, etc.).
28#[derive(Debug)]
29pub struct PatError(
30    /// Human-readable message describing the malformation.
31    pub String,
32);
33
34/// A successful match against a Lua pattern, with the captures it produced.
35pub struct Match {
36    /// whole-match span [start, end)
37    pub start: usize,
38    /// End offset of the whole match (exclusive).
39    pub end: usize,
40    /// Captured spans / positions, in pattern order.
41    pub caps: Vec<Cap>,
42}
43
44/// Where the dialects' matchers differ.
45#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Debug)]
46pub(crate) enum Flavor {
47    /// No `%g` class, capture-index errors without the index, "unbalanced
48    /// pattern" for a short `%b`, no `matchdepth`.
49    Lua51,
50    /// Numbered capture-index errors while matching, unnumbered ones when a
51    /// capture is fetched.
52    Lua52,
53    /// 5.3 onwards.
54    Lua53,
55}
56
57const CAP_UNFINISHED: isize = -1;
58const CAP_POSITION: isize = -2;
59
60/// A capture as `get_onecapture` sees it once a match succeeded.
61#[derive(Clone, Copy, PartialEq, Eq, Debug)]
62pub(crate) enum CapValue {
63    Span(usize, usize),
64    Pos(usize),
65}
66
67/// PUC `MatchState`: the subject, the pattern, and the captures of the match
68/// in progress. `try_at` is `reprepstate` + `match`.
69pub(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    /// Match the whole pattern at exactly `s`; `Some(end)` on success.
99    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    /// Number of captures of the last successful match.
105    pub(crate) fn level(&self) -> usize {
106        self.level
107    }
108
109    /// PUC `get_onecapture`: capture `i` of a match spanning `[s, e)`; with
110    /// no captures, index 0 is the whole match.
111    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                        // PUC reads the subject's terminating NUL at its end
174                        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            // a single-char class plus an optional suffix
196            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                // do-while: the first byte is consumed before looking for
238                // ']', so "[]" and "[^]" start a set containing ']'
239                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    /// `[set]` test; `p` is the '[' and `ec` the closing ']'.
276    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    /// Back-reference `%d`. A position capture passes the index check and
407    /// then never matches: its length is the (huge) `CAP_POSITION` cast.
408    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
427/// Split a leading `^` anchor from the pattern body. The caller decides what
428/// the anchor means (find/match scan at most once; gsub stops after the first
429/// position).
430pub fn anchor_split(pat: &[u8]) -> (bool, &[u8]) {
431    match pat.first() {
432        Some(b'^') => (true, &pat[1..]),
433        _ => (false, pat),
434    }
435}
436
437/// Try to match `pat_body` (already `^`-stripped) at exactly position `s`,
438/// with no forward scan. Returns the Match (whose `start == s`) or None; a
439/// capture left open by a successful match is an error.
440pub 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
460/// Scan from `init` for the first match (PUC str_find_aux without the plain
461/// fast path). A leading `^` anchors the search to `init`.
462pub 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
479/// Whether the pattern contains a byte from PUC's `SPECIALS`; a pattern
480/// without one is searched for as plain text.
481pub 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
490/// Plain substring search (find with plain=true).
491pub 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        // position capture
546        let mm = find(b"abc", b"a()b", 0).unwrap().unwrap();
547        assert_eq!(mm.caps[0], Cap::Pos(1));
548        // balanced
549        assert_eq!(m("(foo(bar))baz", "%b()"), Some((0, 10)));
550        // frontier
551        assert_eq!(m("THE (quick) fox", "%f[%a]%a+"), Some((0, 3)));
552        // back-reference
553        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()); // unfinished capture
562        assert!(find(b"x", b"%1", 0).is_err());
563    }
564}