use lsp_types::{
CodeAction, CodeActionKind, CodeActionParams, CodeActionResponse, Position, Range, TextEdit,
Uri, WorkspaceEdit,
};
use rustledger_core::Directive;
use rustledger_parser::{ParseErrorKind, ParseResult};
use std::collections::{BTreeSet, HashMap, HashSet};
use super::utils::{LineIndex, PositionEncoding, ranges_overlap};
pub fn handle_code_actions(
params: &CodeActionParams,
source: &str,
parse_result: &ParseResult,
encoding: PositionEncoding,
) -> Option<CodeActionResponse> {
let mut actions = Vec::new();
let range = params.range;
let uri = params.text_document.uri.clone();
let line_index = LineIndex::new(source, encoding);
let defined_accounts = collect_defined_accounts(parse_result);
let used_accounts = collect_used_accounts(parse_result);
let undefined_accounts: Vec<_> = used_accounts
.difference(&defined_accounts)
.cloned()
.collect();
for account in undefined_accounts {
if is_account_in_range(&line_index, &account, range, parse_result) {
let action = create_open_directive_action(&uri, &account);
actions.push(action);
}
}
if let Some(action) = check_unbalanced_transactions(params, &line_index, parse_result) {
actions.push(action);
}
actions.extend(bom_removal_actions(
parse_result,
source,
&uri,
range,
encoding,
));
actions.extend(super::import::import_code_actions(
&parse_result.directives,
source,
range,
encoding,
));
if actions.is_empty() {
None
} else {
Some(actions.into_iter().map(|a| a.into()).collect())
}
}
fn bom_removal_actions(
parse_result: &ParseResult,
source: &str,
uri: &Uri,
request_range: Range,
encoding: PositionEncoding,
) -> Vec<CodeAction> {
let line_index = LineIndex::new(source, encoding);
let mut bom_offsets: BTreeSet<usize> = BTreeSet::new();
for err in &parse_result.errors {
if !matches!(err.kind, ParseErrorKind::BomInDirectiveBody) {
continue;
}
let span_text = source.get(err.span.start..err.span.end).unwrap_or("");
for (offset_in_span, _) in span_text.match_indices('\u{FEFF}') {
bom_offsets.insert(err.span.start + offset_in_span);
}
}
let bom_byte_len = '\u{FEFF}'.len_utf8();
let edits: Vec<TextEdit> = bom_offsets
.into_iter()
.filter_map(|bom_start| {
let bom_end = bom_start + bom_byte_len;
let (sl, sc) = line_index.offset_to_position(bom_start);
let (el, ec) = line_index.offset_to_position(bom_end);
let bom_range = Range::new(Position::new(sl, sc), Position::new(el, ec));
if !ranges_overlap(bom_range, request_range) {
return None;
}
Some(TextEdit {
range: bom_range,
new_text: String::new(),
})
})
.collect();
if edits.is_empty() {
return Vec::new();
}
let title = if edits.len() == 1 {
"Remove BOM (U+FEFF)".to_string()
} else {
format!("Remove {} BOM bytes (U+FEFF)", edits.len())
};
#[allow(clippy::mutable_key_type)]
let mut changes = HashMap::new();
changes.insert(uri.clone(), edits);
vec![CodeAction {
title,
kind: Some(CodeActionKind::QUICKFIX),
edit: Some(WorkspaceEdit {
changes: Some(changes),
document_changes: None,
change_annotations: None,
}),
..CodeAction::default()
}]
}
fn collect_defined_accounts(parse_result: &ParseResult) -> HashSet<String> {
let mut accounts = HashSet::new();
for spanned in &parse_result.directives {
if let Directive::Open(open) = &spanned.value {
accounts.insert(open.account.to_string());
}
}
accounts
}
fn collect_used_accounts(parse_result: &ParseResult) -> HashSet<String> {
let mut accounts = HashSet::new();
for spanned in &parse_result.directives {
match &spanned.value {
Directive::Transaction(txn) => {
for posting in &txn.postings {
accounts.insert(posting.account.to_string());
}
}
Directive::Balance(bal) => {
accounts.insert(bal.account.to_string());
}
Directive::Pad(pad) => {
accounts.insert(pad.account.to_string());
accounts.insert(pad.source_account.to_string());
}
Directive::Note(note) => {
accounts.insert(note.account.to_string());
}
Directive::Document(doc) => {
accounts.insert(doc.account.to_string());
}
Directive::Close(close) => {
accounts.insert(close.account.to_string());
}
_ => {}
}
}
accounts
}
fn is_account_in_range(
line_index: &LineIndex<'_>,
account: &str,
range: Range,
parse_result: &ParseResult,
) -> bool {
let start_line = range.start.line;
let window_start = start_line.saturating_sub(3);
let window_end = start_line.saturating_add(10);
for line_idx in window_start..=window_end {
if let Some(line) = line_index.line_text(line_idx)
&& line.contains(account)
{
return true;
}
}
for spanned in &parse_result.directives {
if let Directive::Transaction(txn) = &spanned.value {
let (dir_line, _) = line_index.offset_to_position(spanned.span.start);
let (end_line, _) = line_index.offset_to_position(spanned.span.end);
if (range.start.line <= end_line) && (range.end.line >= dir_line) {
for posting in &txn.postings {
if posting.account.as_ref() == account {
return true;
}
}
}
}
}
false
}
fn create_open_directive_action(uri: &Uri, account: &str) -> CodeAction {
let data = serde_json::json!({
"kind": "add_open_directive",
"account": account,
"uri": uri.as_str(),
});
CodeAction {
title: format!("Add 'open {}' directive", account),
kind: Some(CodeActionKind::QUICKFIX),
diagnostics: None,
edit: None, command: None,
is_preferred: Some(true),
disabled: None,
data: Some(data),
}
}
#[allow(clippy::mutable_key_type)] pub fn handle_code_action_resolve(
action: CodeAction,
source: &str,
parse_result: &ParseResult,
uri: &Uri,
encoding: PositionEncoding,
) -> CodeAction {
let mut resolved = action.clone();
if let Some(data) = &action.data
&& data.get("kind").and_then(|v| v.as_str()) == Some("add_open_directive")
&& let Some(account) = data.get("account").and_then(|v| v.as_str())
{
let line_index = LineIndex::new(source, encoding);
resolved.edit = Some(compute_open_directive_edit(
uri,
&line_index,
account,
parse_result,
));
}
resolved
}
#[allow(clippy::mutable_key_type)] fn compute_open_directive_edit(
uri: &Uri,
line_index: &LineIndex,
account: &str,
parse_result: &ParseResult,
) -> WorkspaceEdit {
let earliest_date =
find_earliest_date(parse_result).unwrap_or_else(|| "2000-01-01".to_string());
let insert_position = find_open_directive_position(line_index, parse_result);
let new_text = format!("{} open {}\n", earliest_date, account);
let mut changes = HashMap::new();
changes.insert(
uri.clone(),
vec![TextEdit {
range: Range {
start: insert_position,
end: insert_position,
},
new_text,
}],
);
WorkspaceEdit {
changes: Some(changes),
document_changes: None,
change_annotations: None,
}
}
fn find_earliest_date(parse_result: &ParseResult) -> Option<String> {
let mut earliest: Option<rustledger_core::NaiveDate> = None;
for spanned in &parse_result.directives {
let date = match &spanned.value {
Directive::Transaction(t) => Some(t.date),
Directive::Open(o) => Some(o.date),
Directive::Close(c) => Some(c.date),
Directive::Balance(b) => Some(b.date),
Directive::Pad(p) => Some(p.date),
Directive::Commodity(c) => Some(c.date),
Directive::Event(e) => Some(e.date),
Directive::Note(n) => Some(n.date),
Directive::Document(d) => Some(d.date),
Directive::Price(p) => Some(p.date),
Directive::Query(q) => Some(q.date),
Directive::Custom(c) => Some(c.date),
};
if let Some(d) = date {
earliest = Some(earliest.map_or(d, |e| e.min(d)));
}
}
earliest.map(|d| d.to_string())
}
fn find_open_directive_position(line_index: &LineIndex, parse_result: &ParseResult) -> Position {
let mut last_open_end: Option<usize> = None;
for spanned in &parse_result.directives {
if matches!(&spanned.value, Directive::Open(_)) {
last_open_end = Some(spanned.span.end);
}
}
if let Some(offset) = last_open_end {
let (line, _) = line_index.offset_to_position(offset);
Position::new(line + 1, 0)
} else {
Position::new(0, 0)
}
}
fn check_unbalanced_transactions(
params: &CodeActionParams,
line_index: &LineIndex,
parse_result: &ParseResult,
) -> Option<CodeAction> {
let range = params.range;
for spanned in &parse_result.directives {
if let Directive::Transaction(txn) = &spanned.value {
let (start_line, _) = line_index.offset_to_position(spanned.span.start);
let (end_line, _) = line_index.offset_to_position(spanned.span.end);
if range.start.line >= start_line && range.start.line <= end_line {
let postings_without_amount =
txn.postings.iter().filter(|p| p.units.is_none()).count();
let postings_with_amount =
txn.postings.iter().filter(|p| p.units.is_some()).count();
if postings_without_amount == 1 && postings_with_amount >= 1 {
continue;
}
if postings_without_amount == 0 && postings_with_amount >= 2 {
continue;
}
}
}
}
None
}
#[cfg(test)]
mod tests {
use super::*;
use rustledger_parser::parse;
#[test]
fn test_collect_accounts() {
let source = r#"
2024-01-01 open Assets:Bank USD
2024-01-15 * "Coffee Shop"
Assets:Bank -5.00 USD
Expenses:Food
"#;
let result = parse(source);
let defined = collect_defined_accounts(&result);
assert!(defined.contains("Assets:Bank"));
assert!(!defined.contains("Expenses:Food"));
let used = collect_used_accounts(&result);
assert!(used.contains("Assets:Bank"));
assert!(used.contains("Expenses:Food"));
}
#[test]
#[allow(clippy::mutable_key_type)]
fn test_bom_removal_code_action_for_mid_file_bom() {
use lsp_types::{
CodeActionContext, PartialResultParams, TextDocumentIdentifier, Uri,
WorkDoneProgressParams,
};
let source = "2024-01-01 open Assets:Bank USD\n\u{FEFF}2024-01-02 open Assets:Cash USD\n";
let result = parse(source);
assert!(
result
.errors
.iter()
.any(|e| matches!(e.kind, ParseErrorKind::BomInDirectiveBody)),
"expected a BomInDirectiveBody error in parse output"
);
let uri: Uri = "file:///test.bean".parse().unwrap();
let params = CodeActionParams {
text_document: TextDocumentIdentifier { uri: uri.clone() },
range: Range::new(Position::new(0, 0), Position::new(100, 0)),
context: CodeActionContext::default(),
work_done_progress_params: WorkDoneProgressParams::default(),
partial_result_params: PartialResultParams::default(),
};
let response = handle_code_actions(¶ms, source, &result, PositionEncoding::Utf16)
.expect("expected at least one code action");
let bom_action = response
.into_iter()
.find_map(|a| match a {
lsp_types::CodeActionOrCommand::CodeAction(action) => {
if action.title.contains("Remove BOM") {
Some(action)
} else {
None
}
}
lsp_types::CodeActionOrCommand::Command(_) => None,
})
.expect("expected a 'Remove BOM' quick-fix action");
assert_eq!(bom_action.kind, Some(CodeActionKind::QUICKFIX));
let edit = bom_action.edit.expect("action must carry a WorkspaceEdit");
let changes = edit.changes.expect("edit must include changes");
let text_edits = &changes[&uri];
assert_eq!(text_edits.len(), 1, "expected exactly one TextEdit");
assert!(
text_edits[0].new_text.is_empty(),
"BOM removal must be a deletion (empty new_text)"
);
assert_eq!(
text_edits[0].range,
Range::new(Position::new(1, 0), Position::new(1, 1)),
"BOM range must be in UTF-16 code units (BOM = 1 UTF-16 unit); \
a (1,0)..(1,3) range here would indicate byte columns leaking through"
);
}
#[test]
#[allow(clippy::mutable_key_type)]
fn test_bom_removal_action_covers_all_boms_in_span() {
use lsp_types::{
CodeActionContext, PartialResultParams, TextDocumentIdentifier, Uri,
WorkDoneProgressParams,
};
let source = "2024-01-01 open Assets:Bank \u{FEFF}USD \u{FEFF}EUR\n";
let result = parse(source);
let bom_errs = result
.errors
.iter()
.filter(|e| matches!(e.kind, ParseErrorKind::BomInDirectiveBody))
.count();
assert!(
bom_errs >= 1,
"expected at least one BomInDirectiveBody error in parse output, got: {:?}",
result.errors
);
let uri: Uri = "file:///test.bean".parse().unwrap();
let params = CodeActionParams {
text_document: TextDocumentIdentifier { uri: uri.clone() },
range: Range::new(Position::new(0, 0), Position::new(100, 0)),
context: CodeActionContext::default(),
work_done_progress_params: WorkDoneProgressParams::default(),
partial_result_params: PartialResultParams::default(),
};
let response = handle_code_actions(¶ms, source, &result, PositionEncoding::Utf16)
.expect("expected at least one code action");
let bom_actions: Vec<_> = response
.into_iter()
.filter_map(|a| match a {
lsp_types::CodeActionOrCommand::CodeAction(action)
if action.title.contains("Remove") && action.title.contains("BOM") =>
{
Some(action)
}
_ => None,
})
.collect();
assert!(
!bom_actions.is_empty(),
"expected at least one 'Remove ... BOM' quick-fix action"
);
let mut all_edits: Vec<TextEdit> = Vec::new();
for action in bom_actions {
let edit = action.edit.expect("action must carry a WorkspaceEdit");
let changes = edit.changes.expect("edit must include changes");
all_edits.extend(changes[&uri].iter().cloned());
}
assert_eq!(
all_edits.len(),
2,
"expected one TextEdit per BOM occurrence across all actions; got {} edits",
all_edits.len()
);
for edit in &all_edits {
assert!(
edit.new_text.is_empty(),
"each BOM removal is a deletion (empty new_text)"
);
}
assert_ne!(
all_edits[0].range, all_edits[1].range,
"multi-BOM edits must target distinct positions, not duplicates"
);
}
#[test]
#[allow(clippy::mutable_key_type)]
fn test_bom_removal_action_utf8_encoding_emits_byte_columns() {
use lsp_types::{
CodeActionContext, PartialResultParams, TextDocumentIdentifier, Uri,
WorkDoneProgressParams,
};
let source = "2024-01-01 open Assets:Bank USD\n\u{FEFF}2024-01-02 open Assets:Cash USD\n";
let result = parse(source);
let uri: Uri = "file:///test.bean".parse().unwrap();
let params = CodeActionParams {
text_document: TextDocumentIdentifier { uri: uri.clone() },
range: Range::new(Position::new(0, 0), Position::new(100, 0)),
context: CodeActionContext::default(),
work_done_progress_params: WorkDoneProgressParams::default(),
partial_result_params: PartialResultParams::default(),
};
let response = handle_code_actions(¶ms, source, &result, PositionEncoding::Utf8)
.expect("expected at least one code action");
let bom_action = response
.into_iter()
.find_map(|a| match a {
lsp_types::CodeActionOrCommand::CodeAction(action) => {
if action.title.contains("Remove BOM") {
Some(action)
} else {
None
}
}
lsp_types::CodeActionOrCommand::Command(_) => None,
})
.expect("expected a 'Remove BOM' quick-fix action");
let edit = bom_action.edit.expect("action must carry a WorkspaceEdit");
let changes = edit.changes.expect("edit must include changes");
let text_edits = &changes[&uri];
assert_eq!(
text_edits[0].range,
Range::new(Position::new(1, 0), Position::new(1, 3)),
"UTF-8 encoding must emit byte columns: BOM is 3 bytes wide"
);
}
#[test]
fn test_find_earliest_date() {
let source = r#"
2024-06-15 open Assets:Bank
2024-01-01 open Assets:Cash
2024-03-01 * "Test"
Assets:Bank -10 USD
Assets:Cash
"#;
let result = parse(source);
let earliest = find_earliest_date(&result);
assert_eq!(earliest, Some("2024-01-01".to_string()));
}
#[test]
#[allow(clippy::mutable_key_type)] fn test_code_action_resolve() {
let source = r#"
2024-01-01 open Assets:Bank USD
2024-01-15 * "Coffee"
Assets:Bank -5.00 USD
Expenses:Food
"#;
let result = parse(source);
let uri: Uri = "file:///test.beancount".parse().unwrap();
let action = CodeAction {
title: "Add 'open Expenses:Food' directive".to_string(),
kind: Some(CodeActionKind::QUICKFIX),
diagnostics: None,
edit: None, command: None,
is_preferred: Some(true),
disabled: None,
data: Some(serde_json::json!({
"kind": "add_open_directive",
"account": "Expenses:Food",
"uri": uri.as_str(),
})),
};
let resolved =
handle_code_action_resolve(action, source, &result, &uri, PositionEncoding::Utf16);
assert!(resolved.edit.is_some());
let edit = resolved.edit.unwrap();
let changes = edit.changes.unwrap();
let edits = changes.get(&uri).unwrap();
assert_eq!(edits.len(), 1);
assert!(edits[0].new_text.contains("open Expenses:Food"));
assert!(edits[0].new_text.contains("2024-01-01")); }
}