use regex::Regex;
use crate::types::{Cursor, Query, Search};
use hjkl_vim_types::Operator;
#[derive(Debug, Clone)]
pub struct SearchPrompt {
pub text: String,
pub cursor: usize,
pub forward: bool,
pub operator: Option<(Operator, usize, (usize, usize))>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CaseMode {
Sensitive,
Insensitive,
Smart,
}
impl CaseMode {
pub fn from_options(ignorecase: bool, smartcase: bool) -> Self {
if !ignorecase {
CaseMode::Sensitive
} else if smartcase {
CaseMode::Smart
} else {
CaseMode::Insensitive
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum MagicLevel {
VeryMagic,
Magic,
NoMagic,
VeryNoMagic,
}
fn very_magic_special(ch: char) -> bool {
matches!(
ch,
'(' | ')' | '+' | '?' | '|' | '{' | '}' | '=' | '<' | '>'
)
}
fn magic_special(ch: char) -> bool {
matches!(ch, '.' | '*' | '[' | ']' | '~')
}
fn nomagic_special(ch: char) -> bool {
matches!(ch, '^' | '$')
}
fn regex_meta(ch: char) -> bool {
matches!(
ch,
'\\' | '.' | '+' | '*' | '?' | '(' | ')' | '|' | '[' | ']' | '{' | '}' | '^' | '$'
)
}
fn is_special_unescaped(ch: char, level: MagicLevel) -> bool {
if very_magic_special(ch) {
level == MagicLevel::VeryMagic
} else if magic_special(ch) {
matches!(level, MagicLevel::VeryMagic | MagicLevel::Magic)
} else if nomagic_special(ch) {
level != MagicLevel::VeryNoMagic
} else {
false
}
}
fn emit_special(
out: &mut String,
ch: char,
chars: &mut std::iter::Peekable<std::str::Chars>,
last_sub: &str,
) {
match ch {
'(' => out.push('('),
')' => out.push(')'),
'+' => out.push('+'),
'?' => out.push('?'),
'=' => out.push('?'), '|' => out.push('|'),
'<' | '>' => out.push_str(r"\b"),
'{' => {
out.push('{');
emit_counted_repeat(out, chars);
}
'}' => out.push('}'), '.' => out.push('.'),
'*' => out.push('*'),
']' => out.push_str(r"\]"), '~' => out.push_str(last_sub),
'^' => out.push('^'),
'$' => out.push('$'),
_ => out.push(ch),
}
}
fn emit_counted_repeat(out: &mut String, chars: &mut std::iter::Peekable<std::str::Chars>) {
loop {
match chars.next() {
Some('\\') => {
if chars.peek() == Some(&'}') {
chars.next();
out.push('}');
return;
} else if let Some(c2) = chars.next() {
out.push(c2);
} else {
return;
}
}
Some('}') => {
out.push('}');
return;
}
Some(c2) => out.push(c2),
None => return,
}
}
}
fn emit_literal(out: &mut String, ch: char) {
if regex_meta(ch) {
out.push('\\');
}
out.push(ch);
}
fn translate_pattern(pat: &str, last_sub: &str) -> (String, Option<bool>) {
let mut out = String::with_capacity(pat.len());
let mut level = MagicLevel::Magic;
let mut override_mode: Option<bool> = None;
let mut chars = pat.chars().peekable();
let mut in_bracket = false;
while let Some(ch) = chars.next() {
if in_bracket {
out.push(ch);
if ch == ']' {
in_bracket = false;
}
continue;
}
if ch == '\\' {
match chars.next() {
Some('c') => override_mode = Some(true), Some('C') => override_mode = Some(false), Some('v') => level = MagicLevel::VeryMagic,
Some('V') => level = MagicLevel::VeryNoMagic,
Some('m') => level = MagicLevel::Magic,
Some('M') => level = MagicLevel::NoMagic,
Some(d @ '0'..='9') => {
out.push('\\');
out.push(d);
}
Some(c2) if very_magic_special(c2) || magic_special(c2) || nomagic_special(c2) => {
if is_special_unescaped(c2, level) {
emit_literal(&mut out, c2);
} else if c2 == '[' {
out.push('[');
in_bracket = true;
} else {
emit_special(&mut out, c2, &mut chars, last_sub);
}
}
Some(other) => {
out.push('\\');
out.push(other);
}
None => out.push('\\'),
}
continue;
}
if is_special_unescaped(ch, level) {
if ch == '[' {
out.push('[');
in_bracket = true;
} else {
emit_special(&mut out, ch, &mut chars, last_sub);
}
} else {
emit_literal(&mut out, ch);
}
}
(out, override_mode)
}
pub fn resolve_case_mode(pat: &str, base: CaseMode, last_sub: &str) -> (String, CaseMode) {
let (out, override_mode) = translate_pattern(pat, last_sub);
let resolved = match override_mode {
Some(true) => CaseMode::Insensitive,
Some(false) => CaseMode::Sensitive,
None => match base {
CaseMode::Smart => {
if out.chars().any(|c| c.is_uppercase()) {
CaseMode::Sensitive
} else {
CaseMode::Insensitive
}
}
other => other,
},
};
(out, resolved)
}
pub fn vim_to_rust_regex(pat: &str) -> String {
resolve_case_mode(pat, CaseMode::Sensitive, "").0
}
#[derive(Debug, Clone, Default)]
pub struct SearchState {
pub pattern: Option<Regex>,
pub forward: bool,
pub matches: Vec<Vec<(usize, usize)>>,
pub generations: Vec<u64>,
pub wrap_around: bool,
}
impl SearchState {
pub fn new() -> Self {
Self {
pattern: None,
forward: true,
matches: Vec::new(),
generations: Vec::new(),
wrap_around: true,
}
}
pub fn set_pattern(&mut self, re: Option<Regex>) {
self.pattern = re;
self.matches.clear();
self.generations.clear();
}
pub fn matches_for(&mut self, row: usize, line: &str, dirty_gen: u64) -> &[(usize, usize)] {
let Some(ref re) = self.pattern else {
return &[];
};
if self.matches.len() <= row {
self.matches.resize_with(row + 1, Vec::new);
self.generations.resize(row + 1, u64::MAX);
}
if self.generations[row] != dirty_gen {
self.matches[row] = hjkl_buffer::search_match_ranges(re, line);
self.generations[row] = dirty_gen;
}
&self.matches[row]
}
}
pub fn search_forward<B: Cursor + Query + Search>(
buf: &mut B,
state: &mut SearchState,
skip_current: bool,
) -> bool {
let Some(re) = state.pattern.clone() else {
return false;
};
let cursor = buf.cursor();
let total = buf.line_count();
if total == 0 {
return false;
}
let from = if skip_current {
let from_byte = buf.byte_offset(cursor);
buf.pos_at_byte(from_byte.saturating_add(1))
} else {
cursor
};
if let Some(range) = buf.find_next(from, &re) {
if !state.wrap_around && range.start.line < cursor.line {
return false;
}
Cursor::set_cursor(buf, range.start);
return true;
}
false
}
pub fn search_backward<B: Cursor + Query + Search>(
buf: &mut B,
state: &mut SearchState,
skip_current: bool,
) -> bool {
let Some(re) = state.pattern.clone() else {
return false;
};
let cursor = buf.cursor();
let total = buf.line_count();
if total == 0 {
return false;
}
let initial = buf.find_prev(cursor, &re);
let range = if skip_current {
match initial {
Some(m) if m.start == cursor => {
let cb = buf.byte_offset(m.start);
if cb == 0 {
None
} else {
let anchor = buf.pos_at_byte(cb.saturating_sub(1));
buf.find_prev(anchor, &re)
}
}
other => other,
}
} else {
initial
};
if let Some(range) = range {
if !state.wrap_around && range.start.line > cursor.line {
return false;
}
Cursor::set_cursor(buf, range.start);
return true;
}
false
}
pub fn search_matches<B: Query>(
buf: &B,
state: &mut SearchState,
dirty_gen: u64,
row: usize,
) -> Vec<(usize, usize)> {
if state.pattern.is_none() {
return Vec::new();
}
let line_count = buf.line_count() as usize;
if row >= line_count {
return Vec::new();
}
let line = buf.line(row as u32);
state.matches_for(row, &line, dirty_gen).to_vec()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::Pos;
use hjkl_buffer::View;
fn re(pat: &str) -> Regex {
Regex::new(pat).unwrap()
}
fn vim_re(pat: &str) -> Regex {
Regex::new(&vim_to_rust_regex(pat)).unwrap()
}
#[test]
fn vim_boundary_rewrites_to_b() {
assert_eq!(vim_to_rust_regex(r"\<foo\>"), r"\bfoo\b");
assert_eq!(vim_to_rust_regex(r"\<"), r"\b");
assert_eq!(vim_to_rust_regex(r"\>"), r"\b");
}
#[test]
fn escaped_backslash_left_alone() {
let input = r"\\<";
let output = vim_to_rust_regex(input);
assert_eq!(output, r"\\<");
}
#[test]
fn other_escapes_unchanged() {
assert_eq!(vim_to_rust_regex(r"\b"), r"\b");
assert_eq!(vim_to_rust_regex(r"\B"), r"\B");
assert_eq!(vim_to_rust_regex(r"\d\+"), r"\d+");
assert_eq!(vim_to_rust_regex(r"^\w\+$"), r"^\w+$");
}
#[test]
fn mixed_boundary_and_word_class() {
assert_eq!(vim_to_rust_regex(r"\<\w\+\>"), r"\b\w+\b");
}
#[test]
fn vim_boundary_matches_standalone_word_not_suffix() {
let re = vim_re(r"foo\<bar\>");
assert!(!re.is_match("foobar"));
let re2 = vim_re(r"\<bar\>");
assert!(re2.is_match("foo bar baz"));
assert!(!re2.is_match("foobar"));
}
#[test]
fn vim_boundary_start_only() {
let re = vim_re(r"\<word");
assert!(re.is_match("word here"));
assert!(re.is_match("some word here"));
assert!(!re.is_match("sword"));
assert!(!re.is_match("aword"));
}
#[test]
fn vim_boundary_end_only() {
let re = vim_re(r"word\>");
assert!(re.is_match("some word"));
assert!(re.is_match("word"));
assert!(!re.is_match("words"));
assert!(!re.is_match("wordsmith"));
}
#[test]
fn existing_b_boundary_unchanged() {
let re = vim_re(r"\bfoo\b");
assert!(re.is_match("foo"));
assert!(re.is_match("a foo b"));
assert!(!re.is_match("foobar"));
assert!(!re.is_match("afoo"));
}
#[test]
fn vim_whole_word_pattern() {
let re = vim_re(r"\<\w\+\>");
let matches: Vec<_> = re.find_iter("foo bar baz").map(|m| m.as_str()).collect();
assert_eq!(matches, vec!["foo", "bar", "baz"]);
}
#[test]
fn empty_state_no_match() {
let mut b = View::from_str("anything");
let mut s = SearchState::new();
assert!(!search_forward(&mut b, &mut s, false));
assert!(!search_backward(&mut b, &mut s, false));
}
#[test]
fn default_magic_groups_and_backref_replacement_side() {
assert_eq!(
vim_to_rust_regex(r"\(hello\) \(world\)"),
r"(hello) (world)"
);
}
#[test]
fn default_magic_quantifiers_and_alternation() {
assert_eq!(vim_to_rust_regex(r"a\+"), r"a+");
assert_eq!(vim_to_rust_regex(r"a\?"), r"a?");
assert_eq!(vim_to_rust_regex(r"a\="), r"a?");
assert_eq!(vim_to_rust_regex(r"a\|b"), r"a|b");
}
#[test]
fn default_magic_counted_repeat_bare_close() {
assert_eq!(vim_to_rust_regex(r"a\{1,2}"), r"a{1,2}");
assert_eq!(vim_to_rust_regex(r"a\{1,2\}"), r"a{1,2}");
}
#[test]
fn default_magic_unescaped_group_chars_are_literal() {
assert_eq!(vim_to_rust_regex("(a)"), r"\(a\)");
assert_eq!(vim_to_rust_regex("a+b"), r"a\+b");
assert_eq!(vim_to_rust_regex("a|b"), r"a\|b");
assert_eq!(vim_to_rust_regex("a?b"), r"a\?b");
}
#[test]
fn default_magic_dot_star_bracket_caret_dollar_stay_magic() {
assert_eq!(vim_to_rust_regex("a.b"), "a.b");
assert_eq!(vim_to_rust_regex("a*"), "a*");
assert_eq!(vim_to_rust_regex("[0-9]"), "[0-9]");
assert_eq!(vim_to_rust_regex("^foo$"), "^foo$");
}
#[test]
fn magic_tilde_expands_to_last_sub_empty_via_wrapper() {
assert_eq!(vim_to_rust_regex("a~b"), "ab");
assert_eq!(vim_to_rust_regex(r"a\~b"), "a~b");
}
#[test]
fn magic_tilde_expands_to_last_sub() {
let (out, _) = resolve_case_mode("~", CaseMode::Sensitive, "BAR");
assert_eq!(out, "BAR");
let (out, _) = resolve_case_mode("x~y", CaseMode::Sensitive, "BAR");
assert_eq!(out, "xBARy");
}
#[test]
fn escaped_tilde_stays_literal_and_does_not_expand() {
let (out, _) = resolve_case_mode(r"\~", CaseMode::Sensitive, "BAR");
assert_eq!(out, "~");
let re = Regex::new(&out).unwrap();
assert!(re.is_match("a~b"));
assert!(!re.is_match("BAR"));
}
#[test]
fn tilde_in_bracket_class_is_literal() {
let (out, _) = resolve_case_mode("[~]", CaseMode::Sensitive, "BAR");
assert_eq!(out, "[~]");
}
#[test]
fn magic_tilde_no_previous_sub_expands_empty() {
let (out, _) = resolve_case_mode("a~b", CaseMode::Sensitive, "");
assert_eq!(out, "ab");
}
#[test]
fn very_magic_mode_switch_at_start() {
assert_eq!(vim_to_rust_regex(r"\v(\w+) (\w+)"), r"(\w+) (\w+)");
assert_eq!(vim_to_rust_regex(r"\v\d+"), r"\d+");
assert_eq!(vim_to_rust_regex(r"\v<foo>"), r"\bfoo\b");
assert_eq!(vim_to_rust_regex(r"\va=b"), r"a?b");
}
#[test]
fn very_magic_mode_escaped_chars_are_literal() {
assert_eq!(vim_to_rust_regex(r"\v\(a\)"), r"\(a\)");
assert_eq!(vim_to_rust_regex(r"\va\+b"), r"a\+b");
}
#[test]
fn very_nomagic_mode_is_all_literal_except_backslash() {
assert_eq!(vim_to_rust_regex(r"\Va.b"), r"a\.b");
assert_eq!(vim_to_rust_regex(r"\V(a)"), r"\(a\)");
assert_eq!(vim_to_rust_regex(r"\Va\.b"), r"a.b");
}
#[test]
fn nomagic_mode_only_caret_dollar_special() {
assert_eq!(vim_to_rust_regex(r"\M^a.b$"), r"^a\.b$");
assert_eq!(vim_to_rust_regex(r"\Ma\.b"), r"a.b");
}
#[test]
fn mode_switch_mid_pattern() {
assert_eq!(vim_to_rust_regex(r"(a)\v(b)"), r"\(a\)(b)");
assert_eq!(vim_to_rust_regex(r"\va\mb+"), r"ab\+");
}
#[test]
fn backreference_in_pattern_passes_through_unchanged() {
assert_eq!(vim_to_rust_regex(r"\(a\)\1"), r"(a)\1");
}
#[test]
fn character_class_contents_not_translated() {
assert_eq!(vim_to_rust_regex("[()]"), "[()]");
}
#[test]
fn search_forward_reveals_fold() {
use hjkl_buffer::View;
let mut buf = View::from_str("header\nneedle\nfooter");
buf.add_fold(0, 2, true);
assert!(buf.is_row_hidden(1), "row 1 must be hidden before search");
let mut state = SearchState::new();
state.set_pattern(Some(re("needle")));
let found = search_forward(&mut buf, &mut state, false);
assert!(found, "search_forward must find 'needle'");
let row = crate::types::Cursor::cursor(&buf).line as usize;
buf.reveal_row(row);
assert!(
!buf.is_row_hidden(1),
"row 1 must be revealed after search finds it there"
);
}
#[test]
fn search_backward_reveals_fold() {
use hjkl_buffer::View;
let mut buf = View::from_str("footer\nneedle\nheader");
buf.add_fold(0, 2, true);
crate::types::Cursor::set_cursor(&mut buf, crate::types::Pos::new(2, 0));
assert!(buf.is_row_hidden(1), "row 1 must be hidden before search");
let mut state = SearchState::new();
state.set_pattern(Some(re("needle")));
let found = search_backward(&mut buf, &mut state, false);
assert!(found, "search_backward must find 'needle'");
let row = crate::types::Cursor::cursor(&buf).line as usize;
buf.reveal_row(row);
assert!(
!buf.is_row_hidden(1),
"row 1 must be revealed after backward search finds it"
);
}
#[test]
fn forward_finds_first_match() {
let mut b = View::from_str("foo bar foo baz");
let mut s = SearchState::new();
s.set_pattern(Some(re("foo")));
assert!(search_forward(&mut b, &mut s, false));
assert_eq!(Cursor::cursor(&b), Pos::new(0, 0));
}
#[test]
fn forward_skip_current_walks_past() {
let mut b = View::from_str("foo bar foo baz");
let mut s = SearchState::new();
s.set_pattern(Some(re("foo")));
search_forward(&mut b, &mut s, false);
search_forward(&mut b, &mut s, true);
assert_eq!(Cursor::cursor(&b), Pos::new(0, 8));
}
#[test]
fn forward_wraps_to_top() {
let mut b = View::from_str("zzz\nfoo");
Cursor::set_cursor(&mut b, Pos::new(1, 2));
let mut s = SearchState::new();
s.set_pattern(Some(re("zzz")));
s.wrap_around = true;
assert!(search_forward(&mut b, &mut s, true));
assert_eq!(Cursor::cursor(&b), Pos::new(0, 0));
}
#[test]
fn search_matches_caches_against_dirty_gen() {
let b = View::from_str("foo bar");
let mut s = SearchState::new();
s.set_pattern(Some(re("bar")));
let dgen = b.dirty_gen();
let initial = search_matches(&b, &mut s, dgen, 0);
assert_eq!(initial, vec![(4, 7)]);
}
#[test]
fn case_mode_from_options_matrix() {
assert_eq!(CaseMode::from_options(false, false), CaseMode::Sensitive);
assert_eq!(CaseMode::from_options(false, true), CaseMode::Sensitive);
assert_eq!(CaseMode::from_options(true, false), CaseMode::Insensitive);
assert_eq!(CaseMode::from_options(true, true), CaseMode::Smart);
}
#[test]
fn resolve_case_mode_no_override_smart_lowercase() {
let (stripped, mode) = resolve_case_mode("foo", CaseMode::Smart, "");
assert_eq!(stripped, "foo");
assert_eq!(mode, CaseMode::Insensitive);
}
#[test]
fn resolve_case_mode_no_override_smart_uppercase() {
let (stripped, mode) = resolve_case_mode("Foo", CaseMode::Smart, "");
assert_eq!(stripped, "Foo");
assert_eq!(mode, CaseMode::Sensitive);
}
#[test]
fn resolve_case_mode_lower_c_override() {
let (stripped, mode) = resolve_case_mode(r"\cFoo", CaseMode::Sensitive, "");
assert_eq!(stripped, "Foo");
assert_eq!(mode, CaseMode::Insensitive);
}
#[test]
fn resolve_case_mode_upper_c_override() {
let (stripped, mode) = resolve_case_mode(r"foo\C", CaseMode::Smart, "");
assert_eq!(stripped, "foo");
assert_eq!(mode, CaseMode::Sensitive);
}
#[test]
fn resolve_case_mode_last_wins() {
let (stripped, mode) = resolve_case_mode(r"\cfoo\C", CaseMode::Smart, "");
assert_eq!(stripped, "foo");
assert_eq!(mode, CaseMode::Sensitive);
}
fn build_regex_from(pat: &str, ic: bool, smart: bool) -> Regex {
let base = CaseMode::from_options(ic, smart);
let (stripped, mode) = resolve_case_mode(pat, base, "");
let src = if mode == CaseMode::Insensitive {
format!("(?i){stripped}")
} else {
stripped
};
Regex::new(&src).unwrap()
}
#[test]
fn search_finds_capital_with_smartcase_lowercase_pattern() {
let re = build_regex_from("foo", true, true);
assert!(re.is_match("FOO"), "expected match on 'FOO'");
assert!(re.is_match("foo"), "expected match on 'foo'");
}
#[test]
fn search_skips_capital_with_smartcase_mixed_pattern() {
let re = build_regex_from("Foo", true, true);
assert!(!re.is_match("FOO"), "must not match 'FOO' (case-sensitive)");
assert!(re.is_match("Foo"), "must match exact 'Foo'");
}
#[test]
fn search_lower_c_override_finds_capital() {
let re = build_regex_from(r"\cFoo", false, false);
assert!(re.is_match("FOO"), "\\c override must match 'FOO'");
assert!(re.is_match("foo"), "\\c override must match 'foo'");
}
#[test]
fn vim_to_rust_regex_strips_case_overrides() {
assert_eq!(vim_to_rust_regex(r"\cfoo"), "foo");
assert_eq!(vim_to_rust_regex(r"foo\C"), "foo");
assert_eq!(vim_to_rust_regex(r"\<bar\>"), r"\bbar\b");
}
#[test]
fn star_search_finds_lowercase_when_smartcase_lower_word() {
let pat = r"\bfoo\b";
let re = build_regex_from(pat, true, true);
let text = "FOO foo Foo";
let hits: Vec<_> = re.find_iter(text).map(|m| m.as_str()).collect();
assert!(
hits.contains(&"FOO"),
"smartcase lower-word * must match FOO: {hits:?}"
);
assert!(
hits.contains(&"foo"),
"smartcase lower-word * must match foo: {hits:?}"
);
}
}