use super::utils::{
LineIndex, PositionEncoding, account_declaration_spans, commodity_declaration_spans,
get_word_at_position, is_account_like, is_currency_like,
};
use lsp_types::{Location, Position, Range, ReferenceParams, Uri};
use rustledger_core::Directive;
use rustledger_parser::{ParseResult, parse};
#[derive(Clone, Copy, PartialEq, Eq)]
enum RefKind {
Account,
Currency,
Payee,
}
pub fn handle_references(
params: &ReferenceParams,
source: &str,
parse_result: &ParseResult,
other_files: &[(Uri, String)],
uri: &Uri,
encoding: PositionEncoding,
) -> Option<Vec<Location>> {
let position = params.text_document_position.position;
let include_declaration = params.context.include_declaration;
let line_idx = position.line as usize;
let lines: Vec<&str> = source.lines().collect();
let line = lines.get(line_idx)?;
let (word, _, _) = get_word_at_position(line, position.character as usize, encoding)?;
let kind = if is_account_like(&word) {
RefKind::Account
} else if is_currency_like(&word, parse_result) {
RefKind::Currency
} else if is_in_quotes(line, position.character as usize) {
RefKind::Payee
} else {
return None;
};
let collect_file = |src: &str, pr: &ParseResult, file_uri: &Uri| -> Vec<Location> {
let line_index = LineIndex::new(src, encoding);
let mut out = Vec::new();
match kind {
RefKind::Account => {
collect_account_references(
pr,
&line_index,
&word,
file_uri,
include_declaration,
&mut out,
);
}
RefKind::Currency => {
collect_currency_references(
pr,
&line_index,
&word,
file_uri,
include_declaration,
&mut out,
);
}
RefKind::Payee => {
collect_payee_references(src, pr, &line_index, &word, file_uri, &mut out);
}
}
out
};
let mut locations = collect_file(source, parse_result, uri);
for (f_uri, f_source) in other_files {
if f_uri == uri {
continue; }
let f_parse = parse(f_source);
locations.extend(collect_file(f_source, &f_parse, f_uri));
}
locations.sort_by(|a, b| {
a.uri
.as_str()
.cmp(b.uri.as_str())
.then(a.range.start.line.cmp(&b.range.start.line))
.then(a.range.start.character.cmp(&b.range.start.character))
});
locations.dedup_by(|a, b| a.uri == b.uri && a.range == b.range);
if locations.is_empty() {
None
} else {
Some(locations)
}
}
fn collect_account_references(
parse_result: &ParseResult,
line_index: &LineIndex,
account: &str,
uri: &Uri,
include_declaration: bool,
locations: &mut Vec<Location>,
) {
let declaration_spans = account_declaration_spans(parse_result);
for occurrence in &parse_result.account_occurrences {
if occurrence.value != account {
continue;
}
let is_declaration = declaration_spans.contains(&occurrence.span);
if is_declaration && !include_declaration {
continue;
}
let (start_line, start_col) = line_index.offset_to_position(occurrence.span.start);
let (end_line, end_col) = line_index.offset_to_position(occurrence.span.end);
locations.push(Location {
uri: uri.clone(),
range: Range {
start: Position::new(start_line, start_col),
end: Position::new(end_line, end_col),
},
});
}
locations.sort_by(|a, b| {
a.range
.start
.line
.cmp(&b.range.start.line)
.then(a.range.start.character.cmp(&b.range.start.character))
});
locations.dedup_by(|a, b| a.range == b.range);
}
fn collect_currency_references(
parse_result: &ParseResult,
line_index: &LineIndex,
currency: &str,
uri: &Uri,
include_declaration: bool,
locations: &mut Vec<Location>,
) {
let declaration_spans = commodity_declaration_spans(parse_result);
for occurrence in &parse_result.currency_occurrences {
if occurrence.value != currency {
continue;
}
let is_declaration = declaration_spans.contains(&occurrence.span);
if is_declaration && !include_declaration {
continue;
}
let (start_line, start_col) = line_index.offset_to_position(occurrence.span.start);
let (end_line, end_col) = line_index.offset_to_position(occurrence.span.end);
locations.push(Location {
uri: uri.clone(),
range: Range {
start: Position::new(start_line, start_col),
end: Position::new(end_line, end_col),
},
});
}
locations.sort_by(|a, b| {
a.range
.start
.line
.cmp(&b.range.start.line)
.then(a.range.start.character.cmp(&b.range.start.character))
});
locations.dedup_by(|a, b| a.range == b.range);
}
fn collect_payee_references(
source: &str,
parse_result: &ParseResult,
line_index: &LineIndex,
payee: &str,
uri: &Uri,
locations: &mut Vec<Location>,
) {
for spanned in &parse_result.directives {
if let Directive::Transaction(txn) = &spanned.value
&& let Some(ref txn_payee) = txn.payee
&& txn_payee.as_ref() == payee
{
let (line, _) = line_index.offset_to_position(spanned.span.start);
let line_text = source.lines().nth(line as usize).unwrap_or("");
if let Some(quote_byte) = line_text.find(&format!("\"{}\"", payee))
&& let Some(start) = line_index.byte_in_line_to_position(line, quote_byte + 1)
&& let Some(end) =
line_index.byte_in_line_to_position(line, quote_byte + 1 + payee.len())
{
locations.push(Location {
uri: uri.clone(),
range: Range { start, end },
});
}
}
}
}
fn is_in_quotes(line: &str, col: usize) -> bool {
let chars: Vec<char> = line.chars().collect();
let mut in_quotes = false;
for (i, c) in chars.iter().enumerate() {
if i >= col {
break;
}
if *c == '"' {
in_quotes = !in_quotes;
}
}
in_quotes
}
#[cfg(test)]
mod tests {
use super::*;
use rustledger_parser::parse;
#[test]
fn test_find_account_references() {
let source = r#"2024-01-01 open Assets:Bank USD
2024-01-15 * "Coffee"
Assets:Bank -5.00 USD
Expenses:Food
2024-01-31 balance Assets:Bank 100 USD
"#;
let result = parse(source);
let uri: Uri = "file:///test.beancount".parse().unwrap();
let params = ReferenceParams {
text_document_position: lsp_types::TextDocumentPositionParams {
text_document: lsp_types::TextDocumentIdentifier { uri: uri.clone() },
position: Position::new(0, 16), },
work_done_progress_params: Default::default(),
partial_result_params: Default::default(),
context: lsp_types::ReferenceContext {
include_declaration: true,
},
};
let refs = handle_references(¶ms, source, &result, &[], &uri, PositionEncoding::Utf16);
assert!(refs.is_some());
let refs = refs.unwrap();
assert_eq!(refs.len(), 3);
}
#[test]
fn test_find_currency_references() {
let source = r#"2024-01-01 open Assets:Bank USD
2024-01-15 * "Coffee"
Assets:Bank -5.00 USD
Expenses:Food 5.00 USD
"#;
let result = parse(source);
let uri: Uri = "file:///test.beancount".parse().unwrap();
let params = ReferenceParams {
text_document_position: lsp_types::TextDocumentPositionParams {
text_document: lsp_types::TextDocumentIdentifier { uri: uri.clone() },
position: Position::new(0, 28), },
work_done_progress_params: Default::default(),
partial_result_params: Default::default(),
context: lsp_types::ReferenceContext {
include_declaration: true,
},
};
let refs = handle_references(¶ms, source, &result, &[], &uri, PositionEncoding::Utf16);
assert!(refs.is_some());
let refs = refs.unwrap();
assert_eq!(refs.len(), 3);
}
#[test]
fn test_find_currency_references_no_false_positives() {
let source = r#"2024-01-01 open Assets:USD-Reserve
2024-01-01 commodity USD
2024-01-15 * "USD-to-EUR transfer"
Assets:USD-Reserve -100 USD
Assets:Bank 100 USD
"#;
let result = parse(source);
assert!(
result.errors.is_empty(),
"parse errors: {:?}",
result.errors
);
let uri: Uri = "file:///test.beancount".parse().unwrap();
let params = ReferenceParams {
text_document_position: lsp_types::TextDocumentPositionParams {
text_document: lsp_types::TextDocumentIdentifier { uri: uri.clone() },
position: Position::new(1, 21),
},
work_done_progress_params: Default::default(),
partial_result_params: Default::default(),
context: lsp_types::ReferenceContext {
include_declaration: true,
},
};
let refs = handle_references(¶ms, source, &result, &[], &uri, PositionEncoding::Utf16)
.expect("references returns Some");
assert_eq!(
refs.len(),
3,
"expected 3 currency references, got {}: {refs:#?}",
refs.len()
);
let params_no_decl = ReferenceParams {
context: lsp_types::ReferenceContext {
include_declaration: false,
},
..params
};
let refs_no_decl = handle_references(
¶ms_no_decl,
source,
&result,
&[],
&uri,
PositionEncoding::Utf16,
)
.expect("references returns Some");
assert_eq!(
refs_no_decl.len(),
2,
"expected 2 non-declaration references, got {}: {refs_no_decl:#?}",
refs_no_decl.len()
);
}
#[test]
fn test_find_account_references_no_false_positives() {
let source = r#"2024-01-01 open Assets:Bank USD
2024-01-15 * "Assets:Bank transfer note"
Assets:Bank -5.00 USD
memo: "moved Assets:Bank balance"
Expenses:Food
; rebalanced Assets:Bank yesterday
"#;
let result = parse(source);
assert!(
result.errors.is_empty(),
"parse errors: {:?}",
result.errors
);
let uri: Uri = "file:///test.beancount".parse().unwrap();
let params = ReferenceParams {
text_document_position: lsp_types::TextDocumentPositionParams {
text_document: lsp_types::TextDocumentIdentifier { uri: uri.clone() },
position: Position::new(0, 16), },
work_done_progress_params: Default::default(),
partial_result_params: Default::default(),
context: lsp_types::ReferenceContext {
include_declaration: true,
},
};
let refs = handle_references(¶ms, source, &result, &[], &uri, PositionEncoding::Utf16)
.expect("references returns Some");
assert_eq!(
refs.len(),
2,
"expected 2 account references, got {}: {refs:#?}",
refs.len()
);
let lines: Vec<u32> = refs.iter().map(|r| r.range.start.line).collect();
assert_eq!(
lines,
vec![0, 2],
"expected references on lines 0 (open) and 2 (posting), got {lines:?}"
);
for r in &refs {
assert_eq!(
r.range.end.character - r.range.start.character,
"Assets:Bank".len() as u32,
"reference range is wrong width: {r:?}"
);
}
let params_no_decl = ReferenceParams {
context: lsp_types::ReferenceContext {
include_declaration: false,
},
..params
};
let refs_no_decl = handle_references(
¶ms_no_decl,
source,
&result,
&[],
&uri,
PositionEncoding::Utf16,
)
.expect("references returns Some");
assert_eq!(
refs_no_decl.len(),
1,
"expected 1 non-declaration account reference, got {}: {refs_no_decl:#?}",
refs_no_decl.len()
);
assert_eq!(
refs_no_decl[0].range.start.line, 2,
"the surviving reference must be the posting on line 2"
);
}
#[test]
fn test_currency_in_commodity_metadata_is_not_a_declaration() {
let source = r#"2024-01-01 commodity USD
parent: USD
2024-01-15 * "Coffee"
Assets:Bank -5.00 USD
"#;
let result = parse(source);
assert!(
result.errors.is_empty(),
"parse errors: {:?}",
result.errors
);
let uri: Uri = "file:///test.beancount".parse().unwrap();
let params = ReferenceParams {
text_document_position: lsp_types::TextDocumentPositionParams {
text_document: lsp_types::TextDocumentIdentifier { uri: uri.clone() },
position: Position::new(0, 21), },
work_done_progress_params: Default::default(),
partial_result_params: Default::default(),
context: lsp_types::ReferenceContext {
include_declaration: false,
},
};
let refs = handle_references(¶ms, source, &result, &[], &uri, PositionEncoding::Utf16)
.expect("references returns Some");
assert_eq!(
refs.len(),
2,
"expected 2 references (metadata + posting); got {}: {refs:#?}",
refs.len()
);
}
#[test]
fn test_find_account_references_with_interleaved_metadata_1142() {
let source = "\
2024-01-01 open Assets:Bank USD
2024-01-15 * \"Test\"
Assets:Bank -5.00 USD
effective_date: 2024-01-20
Expenses:Food 5.00 USD
effective_date: 2024-01-21
";
let result = parse(source);
assert!(
result.errors.is_empty(),
"parse errors: {:?}",
result.errors
);
let uri: Uri = "file:///test.beancount".parse().unwrap();
let params = ReferenceParams {
text_document_position: lsp_types::TextDocumentPositionParams {
text_document: lsp_types::TextDocumentIdentifier { uri: uri.clone() },
position: Position::new(0, 16), },
work_done_progress_params: Default::default(),
partial_result_params: Default::default(),
context: lsp_types::ReferenceContext {
include_declaration: true,
},
};
let refs = handle_references(¶ms, source, &result, &[], &uri, PositionEncoding::Utf16)
.expect("at least the Open definition + 1 posting reference");
let metadata_lines = [3u32, 5u32];
for r in &refs {
assert!(
!metadata_lines.contains(&r.range.start.line),
"reference range landed on a metadata line: {r:?}"
);
}
let lines: Vec<u32> = refs.iter().map(|r| r.range.start.line).collect();
assert!(
lines.contains(&2),
"Assets:Bank posting on line 2 should appear in refs; got {lines:?}"
);
}
}