Skip to main content

harn_kernel/pure/
regex.rs

1use std::cell::RefCell;
2use std::collections::{BTreeMap, HashMap};
3use std::rc::Rc;
4
5const REGEX_CACHE_LIMIT: usize = 128;
6pub const MAX_REGEX_PATTERN_BYTES: usize = 64 * 1024;
7
8thread_local! {
9    static REGEX_CACHE: RefCell<HashMap<String, Rc<regex::Regex>>> = RefCell::new(HashMap::new());
10    static LAST_REGEX: RefCell<Option<(String, String, Rc<regex::Regex>)>> =
11        const { RefCell::new(None) };
12}
13
14#[derive(Debug, Clone, PartialEq, Eq)]
15pub struct RegexCapture {
16    pub full_match: String,
17    pub groups: Vec<Option<String>>,
18    pub start: usize,
19    pub end: usize,
20    pub line: usize,
21    pub named: BTreeMap<String, String>,
22}
23
24fn compiled(pattern: &str, flags: &str) -> Result<Rc<regex::Regex>, String> {
25    if pattern.len() > MAX_REGEX_PATTERN_BYTES {
26        return Err(format!(
27            "regex pattern exceeds the {MAX_REGEX_PATTERN_BYTES}-byte limit"
28        ));
29    }
30    if let Some(regex) = LAST_REGEX.with(|slot| {
31        slot.borrow()
32            .as_ref()
33            .filter(|(cached_pattern, cached_flags, _)| {
34                cached_pattern == pattern && cached_flags == flags
35            })
36            .map(|(_, _, regex)| Rc::clone(regex))
37    }) {
38        return Ok(regex);
39    }
40
41    let regex = REGEX_CACHE.with(|cache| -> Result<Rc<regex::Regex>, String> {
42        let key = format!("{flags}\0{pattern}");
43        let mut cache = cache.borrow_mut();
44        if let Some(regex) = cache.get(&key) {
45            return Ok(Rc::clone(regex));
46        }
47        let mut builder = regex::RegexBuilder::new(pattern);
48        for flag in flags.chars() {
49            match flag {
50                'i' => builder.case_insensitive(true),
51                'm' => builder.multi_line(true),
52                's' => builder.dot_matches_new_line(true),
53                'x' => builder.ignore_whitespace(true),
54                _ => {
55                    return Err(format!(
56                        "unsupported regex flag '{flag}', expected one of i/m/s/x"
57                    ));
58                }
59            };
60        }
61        let regex = Rc::new(builder.build().map_err(|error| error.to_string())?);
62        if cache.len() >= REGEX_CACHE_LIMIT {
63            cache.clear();
64        }
65        cache.insert(key, Rc::clone(&regex));
66        Ok(regex)
67    })?;
68
69    LAST_REGEX.with(|slot| {
70        *slot.borrow_mut() = Some((pattern.to_string(), flags.to_string(), Rc::clone(&regex)));
71    });
72    Ok(regex)
73}
74
75pub fn regex_matches(pattern: &str, text: &str, flags: &str) -> Result<Vec<String>, String> {
76    Ok(compiled(pattern, flags)?
77        .find_iter(text)
78        .map(|matched| matched.as_str().to_string())
79        .collect())
80}
81
82pub fn regex_replace(
83    pattern: &str,
84    replacement: &str,
85    text: &str,
86    flags: &str,
87) -> Result<String, String> {
88    Ok(compiled(pattern, flags)?
89        .replace_all(text, replacement)
90        .into_owned())
91}
92
93pub fn regex_split(pattern: &str, text: &str, flags: &str) -> Result<Vec<String>, String> {
94    Ok(compiled(pattern, flags)?
95        .split(text)
96        .map(str::to_string)
97        .collect())
98}
99
100/// Keys every `regex_captures` result dict carries. Named groups are merged
101/// into the same dict, so a group may not reuse one of these names.
102pub const REGEX_CAPTURES_RESERVED_KEYS: [&str; 5] = ["match", "groups", "start", "end", "line"];
103
104pub fn regex_captures(pattern: &str, text: &str, flags: &str) -> Result<Vec<RegexCapture>, String> {
105    let regex = compiled(pattern, flags)?;
106    let names = regex
107        .capture_names()
108        .flatten()
109        .map(str::to_string)
110        .collect::<Vec<_>>();
111    if let Some(name) = names
112        .iter()
113        .find(|name| REGEX_CAPTURES_RESERVED_KEYS.contains(&name.as_str()))
114    {
115        return Err(format!(
116            "named group `{name}` collides with a reserved regex_captures key \
117             (match, groups, start, end, line); rename the group"
118        ));
119    }
120    let mut scanned_byte = 0;
121    let mut chars_before = 0;
122    let mut newlines_before = 0;
123    let mut results = Vec::new();
124
125    for captures in regex.captures_iter(text) {
126        let whole = captures
127            .get(0)
128            .expect("regex capture always includes the full match");
129        #[expect(
130            clippy::string_slice,
131            reason = "regex match bounds are char boundaries"
132        )]
133        let gap = &text[scanned_byte..whole.start()];
134        chars_before += gap.chars().count();
135        newlines_before += gap.bytes().filter(|byte| *byte == b'\n').count();
136        let start = chars_before;
137        let line = newlines_before + 1;
138        let matched = whole.as_str();
139        chars_before += matched.chars().count();
140        newlines_before += matched.bytes().filter(|byte| *byte == b'\n').count();
141        scanned_byte = whole.end();
142
143        let groups = (1..captures.len())
144            .map(|index| captures.get(index).map(|value| value.as_str().to_string()))
145            .collect();
146        let named = names
147            .iter()
148            .filter_map(|name| {
149                captures
150                    .name(name)
151                    .map(|value| (name.clone(), value.as_str().to_string()))
152            })
153            .collect();
154        results.push(RegexCapture {
155            full_match: matched.to_string(),
156            groups,
157            start,
158            end: chars_before,
159            line,
160            named,
161        });
162    }
163    Ok(results)
164}
165
166#[cfg(test)]
167mod tests {
168    use super::*;
169
170    #[test]
171    fn captures_use_character_offsets_names_and_lines() {
172        let captures = regex_captures(r"(?m)^(?<word>\w+)-(\d+)$", "λ-1\nHarn-42", "").unwrap();
173        assert_eq!(captures.len(), 2);
174        assert_eq!(captures[0].start, 0);
175        assert_eq!(captures[0].end, 3);
176        assert_eq!(captures[1].line, 2);
177        assert_eq!(captures[1].named["word"], "Harn");
178        assert_eq!(captures[1].groups[1].as_deref(), Some("42"));
179    }
180
181    #[test]
182    fn reserved_named_groups_are_rejected() {
183        for key in REGEX_CAPTURES_RESERVED_KEYS {
184            let error = regex_captures(&format!(r"(?P<{key}>\d+)"), "10", "").unwrap_err();
185            assert!(error.contains(&format!("`{key}`")), "{error}");
186        }
187        // Near-miss names stay valid.
188        let captures = regex_captures(r"(?P<starts>\d+)-(?P<line_no>\d+)", "10-20", "").unwrap();
189        assert_eq!(captures[0].start, 0);
190        assert_eq!(captures[0].named["starts"], "10");
191        assert_eq!(captures[0].named["line_no"], "20");
192    }
193
194    #[test]
195    fn flags_and_pattern_size_are_bounded() {
196        assert!(regex_matches("harn", "HARN", "i").unwrap().len() == 1);
197        assert!(regex_matches("harn", "harn", "q").is_err());
198        assert!(regex_matches(&"x".repeat(MAX_REGEX_PATTERN_BYTES + 1), "x", "").is_err());
199    }
200}