use lsp_types::{
CallHierarchyIncomingCall, CallHierarchyIncomingCallsParams, CallHierarchyItem,
CallHierarchyOutgoingCall, CallHierarchyOutgoingCallsParams, CallHierarchyPrepareParams,
Position, Range, SymbolKind, Uri,
};
use rustledger_core::{Directive, SYNTHESIZED_FILE_ID};
use rustledger_parser::ParseResult;
use std::collections::HashMap;
use super::utils::{LineIndex, PositionEncoding, get_word_at_position, is_account_like};
pub fn handle_prepare_call_hierarchy(
params: &CallHierarchyPrepareParams,
source: &str,
parse_result: &ParseResult,
uri: &Uri,
encoding: PositionEncoding,
) -> Option<Vec<CallHierarchyItem>> {
let position = params.text_document_position_params.position;
let line_idx = position.line as usize;
let lines: Vec<&str> = source.lines().collect();
let line = lines.get(line_idx)?;
let (word, start, end) = get_word_at_position(line, position.character as usize, encoding)?;
if !is_account_like(&word) {
return None;
}
if !account_exists(&word, parse_result) {
return None;
}
let item = CallHierarchyItem {
name: word.clone(),
kind: SymbolKind::FUNCTION, tags: None,
detail: Some("Account".to_string()),
uri: uri.clone(),
range: Range {
start: Position::new(position.line, start as u32),
end: Position::new(position.line, end as u32),
},
selection_range: Range {
start: Position::new(position.line, start as u32),
end: Position::new(position.line, end as u32),
},
data: Some(serde_json::json!({ "account": word })),
};
Some(vec![item])
}
pub fn handle_incoming_calls(
params: &CallHierarchyIncomingCallsParams,
source: &str,
parse_result: &ParseResult,
uri: &Uri,
encoding: PositionEncoding,
) -> Option<Vec<CallHierarchyIncomingCall>> {
let account = params
.item
.data
.as_ref()
.and_then(|v| v.get("account"))
.and_then(|v| v.as_str())
.unwrap_or(¶ms.item.name);
let mut calls: Vec<CallHierarchyIncomingCall> = Vec::new();
let line_index = LineIndex::new(source, encoding);
for spanned in &parse_result.directives {
if let Directive::Transaction(txn) = &spanned.value {
let posting_indices: Vec<usize> = txn
.postings
.iter()
.enumerate()
.filter(|(_, p)| p.account.as_ref() == account)
.map(|(i, _)| i)
.collect();
if posting_indices.is_empty() {
continue;
}
let (txn_line, _) = line_index.offset_to_position(spanned.span.start);
let description = format!("{} {} \"{}\"", txn.date, txn.flag, txn.narration.as_ref());
let from_ranges: Vec<Range> = posting_indices
.iter()
.filter_map(|&idx| {
let sp = txn.postings.get(idx)?;
if sp.file_id == SYNTHESIZED_FILE_ID {
return None;
}
let (posting_line, _) = line_index.offset_to_position(sp.span.start);
let line_text = line_index.line_text(posting_line)?;
let col = line_text.find(account)?;
let start = line_index.byte_in_line_to_position(posting_line, col)?;
let end =
line_index.byte_in_line_to_position(posting_line, col + account.len())?;
Some(Range { start, end })
})
.collect();
if from_ranges.is_empty() {
continue;
}
let (txn_end_line, txn_end_col) = line_index.offset_to_position(spanned.span.end);
let normalized_end_line = if txn_end_col == 0 {
txn_end_line
} else {
txn_end_line.saturating_add(1)
};
let txn_item = CallHierarchyItem {
name: description,
kind: SymbolKind::EVENT, tags: None,
detail: Some(format!("{} postings", txn.postings.len())),
uri: uri.clone(),
range: Range {
start: Position::new(txn_line, 0),
end: Position::new(normalized_end_line, 0),
},
selection_range: Range {
start: Position::new(txn_line, 0),
end: Position::new(txn_line, 10), },
data: Some(serde_json::json!({
"type": "transaction",
"line": txn_line
})),
};
calls.push(CallHierarchyIncomingCall {
from: txn_item,
from_ranges,
});
}
}
if calls.is_empty() { None } else { Some(calls) }
}
pub fn handle_outgoing_calls(
params: &CallHierarchyOutgoingCallsParams,
source: &str,
parse_result: &ParseResult,
uri: &Uri,
encoding: PositionEncoding,
) -> Option<Vec<CallHierarchyOutgoingCall>> {
let data = params.item.data.as_ref()?;
let item_type = data.get("type").and_then(|v| v.as_str())?;
if item_type != "transaction" {
return None;
}
let txn_line = data.get("line").and_then(|v| v.as_u64())? as u32;
let line_index = LineIndex::new(source, encoding);
for spanned in &parse_result.directives {
if let Directive::Transaction(txn) = &spanned.value {
let (line, _) = line_index.offset_to_position(spanned.span.start);
if line != txn_line {
continue;
}
let mut account_postings: HashMap<String, Vec<usize>> = HashMap::new();
for (idx, posting) in txn.postings.iter().enumerate() {
let account = posting.account.to_string();
account_postings.entry(account).or_default().push(idx);
}
let calls: Vec<CallHierarchyOutgoingCall> = account_postings
.into_iter()
.filter_map(|(account, indices)| {
let account_location =
find_account_definition(parse_result, &line_index, &account);
let (_acc_line, acc_range) = match account_location {
Some(loc) => loc,
None => {
let sp = txn.postings.get(indices[0])?;
if sp.file_id == SYNTHESIZED_FILE_ID {
return None;
}
let (posting_line, _) = line_index.offset_to_position(sp.span.start);
let line_text = line_index.line_text(posting_line)?;
let col = line_text.find(&account)?;
let start = line_index.byte_in_line_to_position(posting_line, col)?;
let end = line_index
.byte_in_line_to_position(posting_line, col + account.len())?;
(posting_line, Range { start, end })
}
};
let from_ranges: Vec<Range> = indices
.iter()
.filter_map(|&idx| {
let sp = txn.postings.get(idx)?;
if sp.file_id == SYNTHESIZED_FILE_ID {
return None;
}
let (posting_line, _) = line_index.offset_to_position(sp.span.start);
let line_text = line_index.line_text(posting_line)?;
let col = line_text.find(&account)?;
let start = line_index.byte_in_line_to_position(posting_line, col)?;
let end = line_index
.byte_in_line_to_position(posting_line, col + account.len())?;
Some(Range { start, end })
})
.collect();
let account_item = CallHierarchyItem {
name: account.clone(),
kind: SymbolKind::FUNCTION,
tags: None,
detail: Some("Account".to_string()),
uri: uri.clone(),
range: acc_range,
selection_range: acc_range,
data: Some(serde_json::json!({ "account": account })),
};
Some(CallHierarchyOutgoingCall {
to: account_item,
from_ranges,
})
})
.collect();
return if calls.is_empty() { None } else { Some(calls) };
}
}
None
}
fn find_account_definition(
parse_result: &ParseResult,
line_index: &LineIndex<'_>,
account: &str,
) -> Option<(u32, Range)> {
for spanned in &parse_result.directives {
if let Directive::Open(open) = &spanned.value
&& open.account.as_ref() == account
{
let (line, _) = line_index.offset_to_position(spanned.span.start);
let line_text = line_index.line_text(line)?;
let col = line_text.find(account)?;
let start = line_index.byte_in_line_to_position(line, col)?;
let end = line_index.byte_in_line_to_position(line, col + account.len())?;
return Some((line, Range { start, end }));
}
}
None
}
fn account_exists(account: &str, parse_result: &ParseResult) -> bool {
for spanned in &parse_result.directives {
match &spanned.value {
Directive::Open(open) if open.account.as_ref() == account => return true,
Directive::Close(close) if close.account.as_ref() == account => return true,
Directive::Balance(bal) if bal.account.as_ref() == account => return true,
Directive::Transaction(txn)
if txn.postings.iter().any(|p| p.account.as_ref() == account) =>
{
return true;
}
_ => {}
}
}
false
}
#[cfg(test)]
mod tests {
use super::*;
use rustledger_parser::parse;
#[test]
fn test_prepare_call_hierarchy() {
let source = r#"2024-01-01 open Assets:Bank:Checking USD
2024-01-15 * "Coffee"
Assets:Bank:Checking -5.00 USD
Expenses:Food
"#;
let result = parse(source);
let uri: Uri = "file:///test.beancount".parse().unwrap();
let params = CallHierarchyPrepareParams {
text_document_position_params: lsp_types::TextDocumentPositionParams {
text_document: lsp_types::TextDocumentIdentifier { uri: uri.clone() },
position: Position::new(0, 20), },
work_done_progress_params: Default::default(),
};
let items =
handle_prepare_call_hierarchy(¶ms, source, &result, &uri, PositionEncoding::Utf16);
assert!(items.is_some());
let items = items.unwrap();
assert_eq!(items.len(), 1);
assert_eq!(items[0].name, "Assets:Bank:Checking");
assert_eq!(items[0].kind, SymbolKind::FUNCTION);
}
#[test]
fn test_incoming_calls() {
let source = r#"2024-01-01 open Assets:Bank:Checking USD
2024-01-15 * "Coffee"
Assets:Bank:Checking -5.00 USD
Expenses:Food
2024-01-16 * "Lunch"
Assets:Bank:Checking -10.00 USD
Expenses:Food
"#;
let result = parse(source);
let uri: Uri = "file:///test.beancount".parse().unwrap();
let item = CallHierarchyItem {
name: "Assets:Bank:Checking".to_string(),
kind: SymbolKind::FUNCTION,
tags: None,
detail: Some("Account".to_string()),
uri: uri.clone(),
range: Range::default(),
selection_range: Range::default(),
data: Some(serde_json::json!({ "account": "Assets:Bank:Checking" })),
};
let params = CallHierarchyIncomingCallsParams {
item,
work_done_progress_params: Default::default(),
partial_result_params: Default::default(),
};
let calls = handle_incoming_calls(¶ms, source, &result, &uri, PositionEncoding::Utf16);
assert!(calls.is_some());
let calls = calls.unwrap();
assert_eq!(calls.len(), 2); }
#[test]
fn test_outgoing_calls_from_transaction() {
let source = r#"2024-01-01 open Assets:Bank:Checking USD
2024-01-01 open Expenses:Food
2024-01-15 * "Coffee"
Assets:Bank:Checking -5.00 USD
Expenses:Food
"#;
let result = parse(source);
let uri: Uri = "file:///test.beancount".parse().unwrap();
let item = CallHierarchyItem {
name: "2024-01-15 * \"Coffee\"".to_string(),
kind: SymbolKind::EVENT,
tags: None,
detail: None,
uri: uri.clone(),
range: Range::default(),
selection_range: Range::default(),
data: Some(serde_json::json!({
"type": "transaction",
"line": 2
})),
};
let params = CallHierarchyOutgoingCallsParams {
item,
work_done_progress_params: Default::default(),
partial_result_params: Default::default(),
};
let calls = handle_outgoing_calls(¶ms, source, &result, &uri, PositionEncoding::Utf16);
assert!(calls.is_some());
let calls = calls.unwrap();
assert_eq!(calls.len(), 2); }
}