use regex::Regex;
#[derive(Debug, Clone)]
pub struct FindPattern {
re: Regex,
case_sensitive: bool,
has_captures: bool,
}
impl FindPattern {
pub fn compile(query: &str) -> Result<Self, regex::Error> {
let case_sensitive = query.chars().any(char::is_uppercase);
let re = if case_sensitive {
Regex::new(query)?
} else {
Regex::new(&format!("(?i){query}"))?
};
let has_captures = re.captures_len() > 1;
Ok(Self {
re,
case_sensitive,
has_captures,
})
}
pub fn as_regex(&self) -> &Regex {
&self.re
}
pub fn case_sensitive(&self) -> bool {
self.case_sensitive
}
pub fn has_captures(&self) -> bool {
self.has_captures
}
pub fn count_matches<S: AsRef<str>>(&self, rows: impl Iterator<Item = S>) -> usize {
rows.map(|row| self.re.find_iter(row.as_ref()).count())
.sum()
}
pub fn match_spans<S: AsRef<str>>(
&self,
rows: impl Iterator<Item = S>,
) -> Vec<(usize, usize, usize)> {
let mut out = Vec::new();
for (row, line) in rows.enumerate() {
let line = line.as_ref();
for m in self.re.find_iter(line) {
let start = line[..m.start()].chars().count();
let end = start + line[m.range()].chars().count();
out.push((row, start, end));
}
}
out
}
pub fn expand(&self, caps: ®ex::Captures<'_>, replacement: &str) -> String {
if !self.has_captures {
return replacement.to_string();
}
let mut out = String::new();
let mut rest = replacement;
while let Some(dollar) = rest.find('$') {
out.push_str(&rest[..dollar]);
let tail = &rest[dollar..];
if let Some(after) = tail.strip_prefix("$$") {
out.push('$');
rest = after;
continue;
}
let end = reference_end(tail);
let reference = &tail[..end];
let mut expanded = String::new();
caps.expand(reference, &mut expanded);
if expanded.is_empty() && !group_exists(caps, reference) {
out.push_str(reference);
} else {
out.push_str(&expanded);
}
rest = &tail[end..];
}
out.push_str(rest);
out
}
}
fn reference_end(s: &str) -> usize {
debug_assert!(s.starts_with('$'));
if let Some(rest) = s.strip_prefix("${") {
return match rest.find('}') {
Some(close) => 2 + close + 1,
None => s.len(),
};
}
let name_len = s[1..]
.find(|c: char| !c.is_ascii_alphanumeric() && c != '_')
.unwrap_or(s.len() - 1);
1 + name_len
}
fn group_exists(caps: ®ex::Captures<'_>, reference: &str) -> bool {
let name = reference
.trim_start_matches('$')
.trim_start_matches('{')
.trim_end_matches('}');
if name.is_empty() {
return false;
}
match name.parse::<usize>() {
Ok(index) => index < caps.len(),
Err(_) => caps.name(name).is_some(),
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PreviewSpan {
pub row: usize,
pub start: usize,
pub end: usize,
pub is_current: bool,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Preview {
pub lines: Vec<String>,
pub spans: Vec<PreviewSpan>,
}
pub fn build_preview<S: AsRef<str>>(
pattern: &FindPattern,
rows: impl Iterator<Item = S>,
replacement: &str,
current: Option<(usize, usize)>,
) -> Preview {
let mut out_lines = Vec::new();
let mut spans = Vec::new();
for (row, line) in rows.enumerate() {
let line = line.as_ref();
let mut rebuilt = String::with_capacity(line.len());
let mut last_byte = 0usize;
let mut out_chars = 0usize;
for caps in pattern.as_regex().captures_iter(line) {
let m = caps.get(0).expect("group 0 always exists");
let gap = &line[last_byte..m.start()];
rebuilt.push_str(gap);
out_chars += gap.chars().count();
let expanded = pattern.expand(&caps, replacement);
let expanded_chars = expanded.chars().count();
let match_start_chars = line[..m.start()].chars().count();
let is_current = current == Some((row, match_start_chars));
rebuilt.push_str(&expanded);
spans.push(PreviewSpan {
row,
start: out_chars,
end: out_chars + expanded_chars,
is_current,
});
out_chars += expanded_chars;
last_byte = m.end();
if m.start() == m.end() && m.end() == last_byte {
continue;
}
}
rebuilt.push_str(&line[last_byte..]);
out_lines.push(rebuilt);
}
Preview {
lines: out_lines,
spans,
}
}
pub fn replace_all(
pattern: &FindPattern,
lines: &[String],
replacement: &str,
) -> Option<(Vec<String>, usize)> {
let preview = build_preview(pattern, lines.iter(), replacement, None);
let count = preview.spans.len();
if count == 0 {
return None;
}
Some((preview.lines, count))
}
#[cfg(test)]
mod tests {
use super::*;
fn lines(v: &[&str]) -> Vec<String> {
v.iter().map(|s| s.to_string()).collect()
}
#[test]
fn lowercase_pattern_matches_any_case() {
let p = FindPattern::compile("todo").unwrap();
assert!(!p.case_sensitive());
assert_eq!(p.count_matches(lines(&["todo Todo TODO"]).iter()), 3);
}
#[test]
fn pattern_with_uppercase_is_exact() {
let p = FindPattern::compile("Todo").unwrap();
assert!(p.case_sensitive());
assert_eq!(p.count_matches(lines(&["todo Todo TODO"]).iter()), 1);
}
#[test]
fn user_written_inline_flag_overrides_smartcase() {
let p = FindPattern::compile("(?-i)todo").unwrap();
assert_eq!(p.count_matches(lines(&["todo Todo TODO"]).iter()), 1);
}
#[test]
fn dollar_is_literal_when_pattern_does_not_capture() {
let p = FindPattern::compile("price").unwrap();
assert!(!p.has_captures());
let (out, n) = replace_all(&p, &lines(&["the price here"]), "$5").unwrap();
assert_eq!(n, 1);
assert_eq!(out, lines(&["the $5 here"]));
}
#[test]
fn dollar_expands_when_pattern_captures() {
let p = FindPattern::compile(r"(\w+)-(\w+)").unwrap();
assert!(p.has_captures());
let (out, _) = replace_all(&p, &lines(&["alpha-beta"]), "$2 $1").unwrap();
assert_eq!(out, lines(&["beta alpha"]));
}
#[test]
fn a_dollar_naming_no_group_stays_literal_even_when_the_pattern_captures() {
let p = FindPattern::compile(r"(Total)").unwrap();
let (out, _) = replace_all(&p, &lines(&["Total: 5 due"]), "$1 cost $100").unwrap();
assert_eq!(out, lines(&["Total cost $100: 5 due"]));
}
#[test]
fn inline_latex_survives_a_capturing_pattern() {
let p = FindPattern::compile(r"(area)").unwrap();
let (out, _) = replace_all(&p, &lines(&["the area"]), "$1 $x^2$").unwrap();
assert_eq!(out, lines(&["the area $x^2$"]));
}
#[test]
fn braced_and_named_references_still_expand() {
let p = FindPattern::compile(r"(?<word>\w+)-(\d+)").unwrap();
let (out, _) = replace_all(&p, &lines(&["ab-12"]), "${word}/$2").unwrap();
assert_eq!(out, lines(&["ab/12"]));
}
#[test]
fn double_dollar_is_still_an_escape() {
let p = FindPattern::compile(r"(x)").unwrap();
let (out, _) = replace_all(&p, &lines(&["x"]), "$$1").unwrap();
assert_eq!(out, lines(&["$1"]));
}
#[test]
fn a_group_that_matched_nothing_expands_to_nothing() {
let p = FindPattern::compile(r"a(z*)").unwrap();
let (out, _) = replace_all(&p, &lines(&["a"]), "[$1]").unwrap();
assert_eq!(out, lines(&["[]"]));
}
#[test]
fn non_capturing_group_does_not_enable_expansion() {
let p = FindPattern::compile(r"(?:foo)").unwrap();
assert!(!p.has_captures());
let (out, _) = replace_all(&p, &lines(&["foo"]), "$1").unwrap();
assert_eq!(out, lines(&["$1"]));
}
#[test]
fn replace_all_rewrites_every_line() {
let p = FindPattern::compile("a").unwrap();
let (out, n) = replace_all(&p, &lines(&["aa", "b", "a"]), "x").unwrap();
assert_eq!(n, 3);
assert_eq!(out, lines(&["xx", "b", "x"]));
}
#[test]
fn replace_all_reports_none_when_nothing_matches() {
let p = FindPattern::compile("zzz").unwrap();
assert!(replace_all(&p, &lines(&["abc"]), "x").is_none());
}
#[test]
fn empty_replacement_deletes_matches() {
let p = FindPattern::compile("todo ").unwrap();
let (out, n) = replace_all(&p, &lines(&["todo todo done"]), "").unwrap();
assert_eq!(n, 2);
assert_eq!(out, lines(&["done"]));
}
#[test]
fn replace_all_never_changes_the_line_count() {
let p = FindPattern::compile("x").unwrap();
let src = lines(&["x", "", "xx", "y"]);
let (out, _) = replace_all(&p, &src, "longer").unwrap();
assert_eq!(out.len(), src.len());
}
#[test]
fn preview_spans_are_in_preview_coordinates_not_buffer_ones() {
let p = FindPattern::compile("ab").unwrap();
let pv = build_preview(&p, lines(&["ab-ab"]).iter(), "XYZW", None);
assert_eq!(pv.lines, lines(&["XYZW-XYZW"]));
assert_eq!(pv.spans[0].start, 0);
assert_eq!(pv.spans[0].end, 4);
assert_eq!(pv.spans[1].start, 5);
assert_eq!(pv.spans[1].end, 9);
}
#[test]
fn preview_flags_the_current_match_by_buffer_position() {
let p = FindPattern::compile("ab").unwrap();
let pv = build_preview(&p, lines(&["ab-ab"]).iter(), "X", Some((0, 3)));
assert_eq!(
pv.spans.iter().map(|s| s.is_current).collect::<Vec<_>>(),
vec![false, true]
);
}
#[test]
fn preview_handles_multibyte_content() {
let p = FindPattern::compile("é").unwrap();
let pv = build_preview(&p, lines(&["aéb"]).iter(), "ü", None);
assert_eq!(pv.lines, lines(&["aüb"]));
assert_eq!(pv.spans[0].start, 1);
assert_eq!(pv.spans[0].end, 2);
}
#[test]
fn preview_with_captures_differs_per_match() {
let p = FindPattern::compile(r"(\w)(\d)").unwrap();
let pv = build_preview(&p, lines(&["a1 b2"]).iter(), "$2$1", None);
assert_eq!(pv.lines, lines(&["1a 2b"]));
}
#[test]
fn zero_width_pattern_terminates() {
let p = FindPattern::compile(r"\b").unwrap();
let pv = build_preview(&p, lines(&["hi there"]).iter(), "|", None);
assert_eq!(pv.lines, lines(&["|hi| |there|"]));
}
#[test]
fn empty_lines_survive_preview() {
let p = FindPattern::compile("x").unwrap();
let pv = build_preview(&p, lines(&["", "x", ""]).iter(), "y", None);
assert_eq!(pv.lines, lines(&["", "y", ""]));
}
#[test]
fn invalid_pattern_reports_error() {
assert!(FindPattern::compile("[").is_err());
}
}