all repos

clerk @ c958816

missing tooling for ledger/hledger

clerk/internal/linter/rule_unbalanced_transaction.go (view raw)

Oleksandr Smirnov Oleksandr Smirnov
olexsmir@gmail.com
ast: store postings as values instead of pointers, 1 month ago
1
package linter
2
3
import (
4
	"fmt"
5
6
	"olexsmir.xyz/clerk/internal/analyzer"
7
	"olexsmir.xyz/clerk/internal/decimal"
8
	"olexsmir.xyz/clerk/journal/ast"
9
	"olexsmir.xyz/clerk/journal/token"
10
)
11
12
const UnbalancedTransactionID RuleID = "unbalanced-transaction"
13
14
// UnbalancedTransaction flags transactions whose postings don't balance to zero.
15
type UnbalancedTransaction struct{}
16
17
func (UnbalancedTransaction) ID() RuleID { return UnbalancedTransactionID }
18
func (u *UnbalancedTransaction) CheckJournal(an *analyzer.Analysis) []Find {
19
	var finds []Find
20
	for _, txn := range an.Transactions {
21
		finds = append(finds, u.check(txn.Postings, txn.Span)...)
22
	}
23
	for _, ptx := range an.PeriodicTransactions {
24
		finds = append(finds, u.check(ptx.Postings, ptx.Span)...)
25
	}
26
	for _, atx := range an.AutomatedTransactions {
27
		finds = append(finds, u.check(atx.Postings, atx.Span)...)
28
	}
29
	return finds
30
}
31
32
func (u *UnbalancedTransaction) check(postings []ast.Posting, span token.Span) []Find {
33
	var hasExpr, hasCost bool
34
	var realPostings []*ast.Posting
35
	autoBalancingPostings := 0
36
37
	for i := range postings {
38
		posting := &postings[i]
39
		if posting.Type != ast.PostingReal {
40
			continue
41
		}
42
		realPostings = append(realPostings, posting)
43
		if posting.Amount == nil {
44
			autoBalancingPostings++
45
		} else {
46
			if posting.Amount.IsExpr {
47
				hasExpr = true
48
			}
49
			if posting.Cost != nil {
50
				hasCost = true
51
			}
52
		}
53
	}
54
55
	if len(realPostings) == 0 ||
56
		autoBalancingPostings == 1 ||
57
		autoBalancingPostings > 1 ||
58
		hasExpr || hasCost {
59
		return nil
60
	}
61
62
	sums := make(map[string]decimal.Decimal)
63
	for _, posting := range realPostings {
64
		if posting.Amount == nil {
65
			continue
66
		}
67
		sums[posting.Amount.Commodity] = sums[posting.Amount.Commodity].Add(posting.Amount.Quantity)
68
	}
69
70
	var finds []Find
71
	for commodity, sum := range sums {
72
		if !sum.IsZero() {
73
			var msg string
74
			if commodity != "" {
75
				msg = fmt.Sprintf("transaction is unbalanced; %s balance is %s", commodity, sum.String())
76
			} else {
77
				msg = fmt.Sprintf("transaction is unbalanced; net balance is %s", sum.String())
78
			}
79
			finds = append(finds, Find{
80
				Code:    u.ID(),
81
				Span:    span,
82
				Message: msg,
83
			})
84
		}
85
	}
86
	return finds
87
}