use lsp_types::{Position, Range};
use rustledger_parser::ParseResult;
pub(crate) fn trim_span_end(source: &str, end: usize) -> usize {
let clamped = end.min(source.len());
source
.get(..clamped)
.map_or(clamped, |s| s.trim_end().len())
}
pub(crate) fn count_noun(count: usize, singular: &str) -> String {
if count == 1 {
format!("1 {singular}")
} else {
format!("{count} {singular}s")
}
}
#[must_use]
pub fn ranges_overlap(a: Range, b: Range) -> bool {
!(a.end <= b.start || b.end <= a.start)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PositionEncoding {
Utf8,
Utf16,
}
impl PositionEncoding {
#[must_use]
pub fn from_negotiated(negotiated: Option<&lsp_types::PositionEncodingKind>) -> Self {
match negotiated {
Some(kind) if *kind == lsp_types::PositionEncodingKind::UTF8 => Self::Utf8,
_ => Self::Utf16,
}
}
}
#[derive(Debug)]
pub struct LineIndex<'a> {
source: &'a str,
line_starts: Vec<usize>,
encoding: PositionEncoding,
utf16_rope: std::cell::OnceCell<ropey::Rope>,
}
impl<'a> LineIndex<'a> {
pub fn new(source: &'a str, encoding: PositionEncoding) -> Self {
let mut line_starts = vec![0];
for (i, ch) in source.char_indices() {
if ch == '\n' {
line_starts.push(i + 1); }
}
Self {
source,
line_starts,
encoding,
utf16_rope: std::cell::OnceCell::new(),
}
}
fn rope(&self) -> &ropey::Rope {
self.utf16_rope
.get_or_init(|| ropey::Rope::from_str(self.source))
}
fn byte_to_line(&self, byte: usize) -> usize {
match self.line_starts.binary_search(&byte) {
Ok(line) => line,
Err(line) => line.saturating_sub(1),
}
}
pub fn offset_to_position(&self, offset: usize) -> (u32, u32) {
let offset = offset.min(self.source.len());
let line = self.byte_to_line(offset);
let line_start = self.line_starts[line];
let col: u32 = match self.encoding {
PositionEncoding::Utf8 => (offset - line_start) as u32,
PositionEncoding::Utf16 => {
let rope = self.rope();
let char_at = rope.byte_to_char(offset);
let line_start_char = rope.byte_to_char(line_start);
(rope.char_to_utf16_cu(char_at) - rope.char_to_utf16_cu(line_start_char)) as u32
}
};
(line as u32, col)
}
pub fn position_to_offset(&self, line: u32, col: u32) -> Option<usize> {
let line_usize = line as usize;
if line_usize >= self.line_starts.len() {
return None;
}
let line_start = self.line_starts[line_usize];
let line_end_raw = self
.line_starts
.get(line_usize + 1)
.copied()
.unwrap_or(self.source.len());
let line_text_end = {
let bytes = self.source.as_bytes();
if line_end_raw > line_start && bytes.get(line_end_raw - 1) == Some(&b'\n') {
line_end_raw - 1
} else {
line_end_raw
}
};
match self.encoding {
PositionEncoding::Utf8 => {
let offset = line_start.checked_add(col as usize)?;
if offset > line_text_end {
return None;
}
if offset < self.source.len() && !self.source.is_char_boundary(offset) {
return None;
}
Some(offset)
}
PositionEncoding::Utf16 => {
let rope = self.rope();
let line_start_char = rope.byte_to_char(line_start);
let line_start_utf16 = rope.char_to_utf16_cu(line_start_char);
let line_text_end_char = rope.byte_to_char(line_text_end);
let line_text_end_utf16 = rope.char_to_utf16_cu(line_text_end_char);
let target_utf16 = line_start_utf16.checked_add(col as usize)?;
if target_utf16 > line_text_end_utf16 {
return None;
}
let char_idx = rope.utf16_cu_to_char(target_utf16);
if rope.char_to_utf16_cu(char_idx) != target_utf16 {
return None;
}
Some(rope.char_to_byte(char_idx))
}
}
}
pub fn line_count(&self) -> usize {
self.line_starts.len()
}
#[must_use]
pub fn line_start_byte(&self, line: u32) -> Option<usize> {
self.line_starts.get(line as usize).copied()
}
#[must_use]
pub fn byte_in_line_to_position(&self, line: u32, byte_in_line: usize) -> Option<Position> {
let line_start = self.line_start_byte(line)?;
let line_text = self.line_text(line)?;
if byte_in_line > line_text.len() {
return None;
}
let abs_byte = line_start.checked_add(byte_in_line)?;
let (l, c) = self.offset_to_position(abs_byte);
Some(Position::new(l, c))
}
#[must_use]
pub fn encoding(&self) -> PositionEncoding {
self.encoding
}
pub fn line_text(&self, line: u32) -> Option<&'a str> {
let line = line as usize;
let start = *self.line_starts.get(line)?;
let end = self
.line_starts
.get(line + 1)
.copied()
.unwrap_or(self.source.len());
Some(
self.source
.get(start..end)?
.trim_end_matches('\n')
.trim_end_matches('\r'),
)
}
}
pub fn get_word_at_position(
line: &str,
col: usize,
encoding: PositionEncoding,
) -> Option<(String, usize, usize)> {
let chars: Vec<char> = line.chars().collect();
let encoded_col_at_char = |char_idx: usize| -> usize {
chars
.iter()
.take(char_idx)
.map(|c| match encoding {
PositionEncoding::Utf8 => c.len_utf8(),
PositionEncoding::Utf16 => c.len_utf16(),
})
.sum()
};
let mut acc = 0usize;
let mut cursor_char_idx = 0usize;
for (i, c) in chars.iter().enumerate() {
if acc == col {
cursor_char_idx = i;
break;
}
let u = match encoding {
PositionEncoding::Utf8 => c.len_utf8(),
PositionEncoding::Utf16 => c.len_utf16(),
};
if acc + u > col {
return None;
}
acc += u;
cursor_char_idx = i + 1;
}
if acc < col {
return None;
}
let mut start = cursor_char_idx;
while start > 0 && is_word_char(chars[start - 1]) {
start -= 1;
}
let mut end = cursor_char_idx;
while end < chars.len() && is_word_char(chars[end]) {
end += 1;
}
if start == end {
return None;
}
let word: String = chars[start..end].iter().collect();
Some((word, encoded_col_at_char(start), encoded_col_at_char(end)))
}
pub fn get_word_at_source_position(
source: &str,
position: Position,
encoding: PositionEncoding,
) -> Option<String> {
let line = source.lines().nth(position.line as usize)?;
let (word, _, _) = get_word_at_position(line, position.character as usize, encoding)?;
Some(word)
}
pub fn is_word_char(c: char) -> bool {
c.is_alphanumeric() || c == ':' || c == '-' || c == '_'
}
pub fn is_account_like(s: &str) -> bool {
s.contains(':')
&& rustledger_core::ACCOUNT_TYPES
.iter()
.any(|t| s.starts_with(t))
}
#[must_use]
pub fn is_account_type(s: &str) -> bool {
rustledger_core::ACCOUNT_TYPES.contains(&s)
}
pub fn is_currency_like_simple(s: &str) -> bool {
s.len() >= 2
&& s.len() <= 5
&& s.chars()
.all(|c| c.is_ascii_uppercase() || c.is_ascii_digit())
}
#[must_use]
pub fn commodity_declaration_spans(
parse_result: &ParseResult,
) -> std::collections::HashSet<rustledger_parser::Span> {
parse_result
.directives
.iter()
.filter_map(|d| {
if !matches!(&d.value, rustledger_core::Directive::Commodity(_)) {
return None;
}
parse_result
.currency_occurrences
.iter()
.find(|o| o.span.start >= d.span.start && o.span.end <= d.span.end)
.map(|o| o.span)
})
.collect()
}
#[must_use]
pub fn account_declaration_spans(
parse_result: &ParseResult,
) -> std::collections::HashSet<rustledger_parser::Span> {
use rustledger_parser::SyntaxKind;
let bom_offset: usize = if parse_result.has_leading_bom { 3 } else { 0 };
let mut declarations = std::collections::HashSet::new();
for node in parse_result.syntax_node().descendants() {
let kind = node.kind();
if kind != SyntaxKind::OPEN_DIRECTIVE && kind != SyntaxKind::CLOSE_DIRECTIVE {
continue;
}
if node
.ancestors()
.skip(1)
.any(|a| a.kind() == SyntaxKind::ERROR_NODE)
{
continue;
}
let Some(account_token) = node
.descendants_with_tokens()
.filter_map(|n| n.into_token())
.find(|t| t.kind() == SyntaxKind::ACCOUNT)
else {
continue;
};
let range = account_token.text_range();
let start = u32::from(range.start()) as usize + bom_offset;
let end = u32::from(range.end()) as usize + bom_offset;
declarations.insert(rustledger_parser::Span::new(start, end));
}
declarations
}
pub fn is_currency_like(s: &str, parse_result: &ParseResult) -> bool {
if !s.chars().all(|c| c.is_uppercase() || c.is_numeric()) || s.len() < 2 || s.len() > 24 {
return false;
}
parse_result
.currency_occurrences
.iter()
.any(|occ| occ.value == s)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_count_noun_pluralization() {
assert_eq!(count_noun(0, "posting"), "0 postings");
assert_eq!(count_noun(1, "posting"), "1 posting");
assert_eq!(count_noun(2, "posting"), "2 postings");
assert_eq!(count_noun(1, "transaction"), "1 transaction");
assert_eq!(count_noun(3, "amount"), "3 amounts");
}
#[test]
fn test_line_index_basic() {
let source = "line1\nline2\nline3";
let index = LineIndex::new(source, PositionEncoding::Utf8);
assert_eq!(index.offset_to_position(0), (0, 0));
assert_eq!(index.offset_to_position(5), (0, 5));
assert_eq!(index.offset_to_position(6), (1, 0));
assert_eq!(index.offset_to_position(10), (1, 4));
assert_eq!(index.offset_to_position(12), (2, 0));
assert_eq!(index.offset_to_position(17), (2, 5));
assert_eq!(index.line_count(), 3);
}
#[test]
fn test_line_index_empty() {
let index = LineIndex::new("", PositionEncoding::Utf8);
assert_eq!(index.offset_to_position(0), (0, 0));
assert_eq!(index.line_count(), 1);
}
#[test]
fn test_line_index_single_line() {
let index = LineIndex::new("hello world", PositionEncoding::Utf8);
assert_eq!(index.offset_to_position(0), (0, 0));
assert_eq!(index.offset_to_position(5), (0, 5));
assert_eq!(index.offset_to_position(11), (0, 11));
assert_eq!(index.line_count(), 1);
}
#[test]
fn test_line_index_trailing_newline() {
let source = "line1\nline2\n";
let index = LineIndex::new(source, PositionEncoding::Utf8);
assert_eq!(index.offset_to_position(11), (1, 5));
assert_eq!(index.offset_to_position(12), (2, 0)); assert_eq!(index.line_count(), 3);
}
#[test]
fn test_line_index_position_to_offset() {
let source = "line1\nline2\nline3";
let index = LineIndex::new(source, PositionEncoding::Utf8);
assert_eq!(index.position_to_offset(0, 0), Some(0));
assert_eq!(index.position_to_offset(0, 5), Some(5));
assert_eq!(index.position_to_offset(1, 0), Some(6));
assert_eq!(index.position_to_offset(1, 4), Some(10));
assert_eq!(index.position_to_offset(2, 0), Some(12));
assert_eq!(index.position_to_offset(3, 0), None);
assert_eq!(index.position_to_offset(0, 100), None);
}
#[test]
fn test_line_index_utf8_matches_inline_walk() {
let source = "2024-01-01 open Assets:Bank USD\n2024-01-15 * \"Coffee\"\n Assets:Bank -5.00 USD\n Expenses:Food\n";
let index = LineIndex::new(source, PositionEncoding::Utf8);
for offset in 0..source.len() {
let mut line = 0u32;
let mut col = 0u32;
for (i, ch) in source.char_indices() {
if i >= offset {
break;
}
if ch == '\n' {
line += 1;
col = 0;
} else {
col += 1;
}
}
let indexed = index.offset_to_position(offset);
assert_eq!((line, col), indexed, "Mismatch at offset {}", offset);
}
}
#[test]
fn test_line_index_position_to_offset_overshoot_symmetric() {
let source = "line1\nline2\nline3";
let utf8 = LineIndex::new(source, PositionEncoding::Utf8);
assert_eq!(utf8.position_to_offset(0, 5), Some(5));
assert_eq!(utf8.position_to_offset(0, 6), None);
assert_eq!(utf8.position_to_offset(0, 100), None);
let utf16 = LineIndex::new(source, PositionEncoding::Utf16);
assert_eq!(utf16.position_to_offset(0, 5), Some(5));
assert_eq!(utf16.position_to_offset(0, 6), None);
assert_eq!(utf16.position_to_offset(0, 100), None);
}
#[test]
fn test_line_index_utf16_columns() {
let source = "💰 USD";
let after_emoji_byte = '💰'.len_utf8();
let utf8 = LineIndex::new(source, PositionEncoding::Utf8);
assert_eq!(utf8.offset_to_position(after_emoji_byte), (0, 4));
let utf16 = LineIndex::new(source, PositionEncoding::Utf16);
assert_eq!(utf16.offset_to_position(after_emoji_byte), (0, 2));
assert_eq!(utf16.position_to_offset(0, 2), Some(after_emoji_byte));
assert_eq!(utf16.offset_to_position(8), (0, 6));
}
#[test]
fn test_line_index_bare_cr_not_a_line_break() {
let source = "a\rb";
let index = LineIndex::new(source, PositionEncoding::Utf8);
assert_eq!(index.line_count(), 1);
assert_eq!(index.offset_to_position(2), (0, 2));
let crlf = LineIndex::new("a\r\nb", PositionEncoding::Utf8);
assert_eq!(crlf.line_count(), 2);
assert_eq!(crlf.offset_to_position(3), (1, 0));
}
#[test]
fn test_line_index_basic_offsets() {
let source = "line1\nline2\nline3";
let index = LineIndex::new(source, PositionEncoding::Utf8);
assert_eq!(index.offset_to_position(0), (0, 0));
assert_eq!(index.offset_to_position(5), (0, 5));
assert_eq!(index.offset_to_position(6), (1, 0));
assert_eq!(index.offset_to_position(10), (1, 4));
}
#[test]
fn test_byte_in_line_to_position_strict_overshoot() {
let source = "abc\ndefgh\nij";
let index = LineIndex::new(source, PositionEncoding::Utf8);
assert_eq!(
index.byte_in_line_to_position(0, 0),
Some(Position::new(0, 0))
);
assert_eq!(
index.byte_in_line_to_position(0, 3),
Some(Position::new(0, 3))
);
assert_eq!(index.byte_in_line_to_position(0, 4), None);
assert_eq!(index.byte_in_line_to_position(0, 100), None);
assert_eq!(
index.byte_in_line_to_position(1, 5),
Some(Position::new(1, 5))
);
assert_eq!(index.byte_in_line_to_position(1, 6), None);
assert_eq!(index.byte_in_line_to_position(99, 0), None);
}
#[test]
fn test_get_word_at_position() {
let line = " Assets:Bank -100.00 USD";
let result = get_word_at_position(line, 5, PositionEncoding::Utf8);
assert!(result.is_some());
let (word, start, end) = result.unwrap();
assert_eq!(word, "Assets:Bank");
assert_eq!(start, 2);
assert_eq!(end, 13);
let result = get_word_at_position(line, 24, PositionEncoding::Utf8);
assert!(result.is_some());
let (word, _, _) = result.unwrap();
assert_eq!(word, "USD");
}
#[test]
fn test_get_word_at_position_encoding_aware() {
let line = "Активы:Банк USD";
let (word, s, e) = get_word_at_position(line, 12, PositionEncoding::Utf16)
.expect("word at UTF-16 col 12 should resolve");
assert_eq!(word, "USD");
assert_eq!((s, e), (12, 15));
let (word, s, e) = get_word_at_position(line, 22, PositionEncoding::Utf8)
.expect("word at UTF-8 col 22 should resolve");
assert_eq!(word, "USD");
assert_eq!((s, e), (22, 25));
assert!(get_word_at_position(line, 1, PositionEncoding::Utf8).is_none());
}
#[test]
fn test_is_account_like() {
assert!(is_account_like("Assets:Bank"));
assert!(is_account_like("Expenses:Food:Groceries"));
assert!(!is_account_like("USD"));
assert!(!is_account_like("Bank"));
assert!(!is_account_like("Random:Thing"));
}
#[test]
fn test_is_account_type() {
assert!(is_account_type("Assets"));
assert!(is_account_type("Liabilities"));
assert!(is_account_type("Income"));
assert!(!is_account_type("Bank"));
assert!(!is_account_type("assets"));
}
#[test]
fn test_is_currency_like_simple() {
assert!(is_currency_like_simple("USD"));
assert!(is_currency_like_simple("EUR"));
assert!(is_currency_like_simple("BTC"));
assert!(!is_currency_like_simple("usd"));
assert!(!is_currency_like_simple("U"));
assert!(!is_currency_like_simple("TOOLONGCURRENCY"));
}
#[test]
fn test_is_currency_like() {
use rustledger_parser::parse;
let source = r#"2024-01-01 commodity USD
2024-01-01 open Assets:Bank USD
2024-01-15 * "Coffee"
Assets:Bank -5.00 USD
Expenses:Food 5.00 USD
2024-01-20 price GBP 1.27 USD
"#;
let parse_result = parse(source);
assert!(
!is_currency_like("usd", &parse_result),
"lowercase rejected"
);
assert!(!is_currency_like("U", &parse_result), "too short rejected");
assert!(
!is_currency_like("XYZ", &parse_result),
"unknown currency rejected"
);
assert!(is_currency_like("USD", &parse_result));
assert!(is_currency_like("GBP", &parse_result));
}
#[test]
fn test_is_word_char() {
assert!(is_word_char('a'));
assert!(is_word_char('Z'));
assert!(is_word_char('0'));
assert!(is_word_char(':'));
assert!(is_word_char('-'));
assert!(is_word_char('_'));
assert!(!is_word_char(' '));
assert!(!is_word_char('"'));
}
#[test]
fn account_declaration_spans_covers_open_and_close() {
use rustledger_parser::parse;
let source = "\
2024-01-01 open Assets:Bank USD
2024-06-15 * \"Coffee\"
Assets:Bank -5.00 USD
2024-08-01 balance Assets:Bank 95.00 USD
2024-12-31 close Assets:Bank
";
let result = parse(source);
assert!(result.errors.is_empty(), "{:?}", result.errors);
let decls = account_declaration_spans(&result);
assert_eq!(decls.len(), 2, "got {decls:?}");
let decl_occurrences: Vec<&rustledger_parser::Spanned<rustledger_core::Account>> = result
.account_occurrences
.iter()
.filter(|o| decls.contains(&o.span))
.collect();
assert_eq!(decl_occurrences.len(), 2);
let mut starts: Vec<usize> = decl_occurrences.iter().map(|o| o.span.start).collect();
starts.sort_unstable();
assert_eq!(
starts[0], 16,
"Open's Assets:Bank starts at byte 16 of the source",
);
let close_offset = source.find("close Assets:Bank").unwrap() + "close ".len();
assert_eq!(starts[1], close_offset);
}
#[test]
fn account_declaration_spans_handles_failed_open_conversion() {
use rustledger_core::Directive;
use rustledger_parser::parse;
let source = "2024-01-01 open Assets:Bank USD \"GARBAGE\"\n";
let result = parse(source);
assert!(
!result
.directives
.iter()
.any(|d| matches!(&d.value, Directive::Open(_))),
"expected the Open to be dropped from directives, got {:?}",
result.directives
);
let has_account = result
.account_occurrences
.iter()
.any(|o| o.value.as_str() == "Assets:Bank");
assert!(has_account, "{:?}", result.account_occurrences);
let decls = account_declaration_spans(&result);
assert_eq!(
decls.len(),
1,
"expected the failed-Open's ACCOUNT to still be a declaration; got {decls:?}",
);
}
#[test]
fn account_declaration_spans_skips_metadata_account_value() {
use rustledger_parser::parse;
let source = "\
2024-01-01 open Assets:Bank USD
payee_account: Assets:Other
";
let result = parse(source);
assert!(result.errors.is_empty(), "{:?}", result.errors);
let decls = account_declaration_spans(&result);
assert_eq!(decls.len(), 1, "got {decls:?}");
let decl = result
.account_occurrences
.iter()
.find(|o| decls.contains(&o.span))
.expect("at least one ACCOUNT occurrence is a declaration");
assert_eq!(
decl.value.as_str(),
"Assets:Bank",
"the declared account must be the directive header, not the metadata value",
);
}
#[test]
fn account_declaration_spans_excludes_posting_account() {
use rustledger_parser::parse;
let source = "\
2024-01-01 open Assets:Bank USD
2024-06-15 * \"Coffee\"
Assets:Bank -5.00 USD
Expenses:Food
";
let result = parse(source);
assert!(result.errors.is_empty(), "{:?}", result.errors);
let decls = account_declaration_spans(&result);
assert_eq!(decls.len(), 1, "got {decls:?}");
let posting_occurrences: Vec<_> = result
.account_occurrences
.iter()
.filter(|o| !decls.contains(&o.span))
.collect();
assert_eq!(
posting_occurrences.len(),
2,
"expected 2 non-declaration occurrences (the two postings); got {posting_occurrences:?}",
);
}
}