From f13f56d2a93c8f0735524f9ffdfcfa88f5175ef1 Mon Sep 17 00:00:00 2001 From: Claude Date: Sat, 28 Feb 2026 10:44:13 +0000 Subject: [PATCH 1/2] Introduce named ID types and return value types from public API Add defined types (LedgerID, SubledgerID, AccountID, TransactionID, EntryID, HoldID) so map keys and method parameters clearly communicate what kind of identifier they expect. These are defined types (not aliases), giving compile-time safety against mixing e.g. a HoldID where an AccountID is expected. Change all public Service methods to return value types instead of pointers to internal state. This prevents callers from accidentally mutating the service's data. Transaction returns use a deep copy helper (copyTransaction) that copies the Entries slice and Metadata map. GetSnapshot now returns ErrSnapshotNotFound instead of nil when no snapshot exists for the given parameters. https://claude.ai/code/session_01947Cs6jqVPDDzmAqQMoB35 --- errors.go | 4 + service.go | 235 ++++++++++++++++++++++++++---------------------- service_test.go | 27 +++--- types.go | 36 +++++--- 4 files changed, 166 insertions(+), 136 deletions(-) diff --git a/errors.go b/errors.go index 037d3e04..044ff036 100644 --- a/errors.go +++ b/errors.go @@ -64,6 +64,10 @@ var ( // (debit/credit) determines the sign of the balance impact. ErrInvalidAmount = errors.New("amount must be positive") + // ErrSnapshotNotFound is returned when no end-of-day snapshot + // exists for the given account and date. + ErrSnapshotNotFound = errors.New("snapshot not found") + // ErrInsufficientBalance is returned when a hold or transaction // would cause the available balance to go below zero for account // types where that is not permitted. Note: this is only enforced diff --git a/service.go b/service.go index a0edf764..456fb7c6 100644 --- a/service.go +++ b/service.go @@ -30,23 +30,23 @@ type Service struct { mu sync.RWMutex // Entity storage, keyed by ID. - ledgers map[string]*Ledger - subledgers map[string]*Subledger - accounts map[string]*Account - transactions map[string]*Transaction - holds map[string]*Hold + ledgers map[LedgerID]*Ledger + subledgers map[SubledgerID]*Subledger + accounts map[AccountID]*Account + transactions map[TransactionID]*Transaction + holds map[HoldID]*Hold // idempotencyIndex maps idempotency keys to transaction IDs. // This allows the system to detect and reject duplicate postings. - idempotencyIndex map[string]string + idempotencyIndex map[string]TransactionID // accountHolds maps account IDs to their active hold IDs. // This index enables efficient lookup of holds per account. - accountHolds map[string][]string + accountHolds map[AccountID][]HoldID // snapshots stores end-of-day balance snapshots. // Structure: accountID -> dateKey -> snapshot. - snapshots map[string]map[string]*BalanceSnapshot + snapshots map[AccountID]map[string]*BalanceSnapshot // auditLog is an append-only log of all mutations. Once appended, // entries are never modified or removed. @@ -69,14 +69,14 @@ type Service struct { // acct, _ := svc.CreateAccount(sl.ID, "Customer A", ledger.Asset) func NewService() *Service { return &Service{ - ledgers: make(map[string]*Ledger), - subledgers: make(map[string]*Subledger), - accounts: make(map[string]*Account), - transactions: make(map[string]*Transaction), - holds: make(map[string]*Hold), - idempotencyIndex: make(map[string]string), - accountHolds: make(map[string][]string), - snapshots: make(map[string]map[string]*BalanceSnapshot), + ledgers: make(map[LedgerID]*Ledger), + subledgers: make(map[SubledgerID]*Subledger), + accounts: make(map[AccountID]*Account), + transactions: make(map[TransactionID]*Transaction), + holds: make(map[HoldID]*Hold), + idempotencyIndex: make(map[string]TransactionID), + accountHolds: make(map[AccountID][]HoldID), + snapshots: make(map[AccountID]map[string]*BalanceSnapshot), clock: time.Now, } } @@ -112,31 +112,31 @@ func (s *Service) appendAudit(eventType AuditEventType, entityID string, payload // a book of accounts (e.g., "General Ledger", "Trading Book"). // // Returns the created ledger. -func (s *Service) CreateLedger(name string) (*Ledger, error) { +func (s *Service) CreateLedger(name string) (Ledger, error) { s.mu.Lock() defer s.mu.Unlock() l := &Ledger{ - ID: s.nextID("ldg"), + ID: LedgerID(s.nextID("ldg")), Name: name, CreatedAt: s.now(), } s.ledgers[l.ID] = l - s.appendAudit(EventLedgerCreated, l.ID, l) - return l, nil + s.appendAudit(EventLedgerCreated, string(l.ID), l) + return *l, nil } // GetLedger retrieves a ledger by its ID. // Returns ErrLedgerNotFound if the ledger does not exist. -func (s *Service) GetLedger(id string) (*Ledger, error) { +func (s *Service) GetLedger(id LedgerID) (Ledger, error) { s.mu.RLock() defer s.mu.RUnlock() l, ok := s.ledgers[id] if !ok { - return nil, ErrLedgerNotFound + return Ledger{}, ErrLedgerNotFound } - return l, nil + return *l, nil } // CreateSubledger creates a new subledger under an existing ledger. @@ -144,36 +144,36 @@ func (s *Service) GetLedger(id string) (*Ledger, error) { // (e.g., "Accounts Receivable", "Checking Accounts", "Loan Portfolio"). // // Returns ErrLedgerNotFound if the parent ledger does not exist. -func (s *Service) CreateSubledger(ledgerID, name string) (*Subledger, error) { +func (s *Service) CreateSubledger(ledgerID LedgerID, name string) (Subledger, error) { s.mu.Lock() defer s.mu.Unlock() if _, ok := s.ledgers[ledgerID]; !ok { - return nil, ErrLedgerNotFound + return Subledger{}, ErrLedgerNotFound } sl := &Subledger{ - ID: s.nextID("sub"), + ID: SubledgerID(s.nextID("sub")), LedgerID: ledgerID, Name: name, CreatedAt: s.now(), } s.subledgers[sl.ID] = sl - s.appendAudit(EventSubledgerCreated, sl.ID, sl) - return sl, nil + s.appendAudit(EventSubledgerCreated, string(sl.ID), sl) + return *sl, nil } // GetSubledger retrieves a subledger by its ID. // Returns ErrSubledgerNotFound if the subledger does not exist. -func (s *Service) GetSubledger(id string) (*Subledger, error) { +func (s *Service) GetSubledger(id SubledgerID) (Subledger, error) { s.mu.RLock() defer s.mu.RUnlock() sl, ok := s.subledgers[id] if !ok { - return nil, ErrSubledgerNotFound + return Subledger{}, ErrSubledgerNotFound } - return sl, nil + return *sl, nil } // --------------------------------------------------------------------------- @@ -190,37 +190,37 @@ func (s *Service) GetSubledger(id string) (*Subledger, error) { // The account starts with a zero balance. // // Returns ErrSubledgerNotFound if the parent subledger does not exist. -func (s *Service) CreateAccount(subledgerID, name string, accountType AccountType) (*Account, error) { +func (s *Service) CreateAccount(subledgerID SubledgerID, name string, accountType AccountType) (Account, error) { s.mu.Lock() defer s.mu.Unlock() if _, ok := s.subledgers[subledgerID]; !ok { - return nil, ErrSubledgerNotFound + return Account{}, ErrSubledgerNotFound } acct := &Account{ - ID: s.nextID("acct"), + ID: AccountID(s.nextID("acct")), SubledgerID: subledgerID, Name: name, Type: accountType, CreatedAt: s.now(), } s.accounts[acct.ID] = acct - s.appendAudit(EventAccountCreated, acct.ID, acct) - return acct, nil + s.appendAudit(EventAccountCreated, string(acct.ID), acct) + return *acct, nil } // GetAccount retrieves an account by its ID. // Returns ErrAccountNotFound if the account does not exist. -func (s *Service) GetAccount(id string) (*Account, error) { +func (s *Service) GetAccount(id AccountID) (Account, error) { s.mu.RLock() defer s.mu.RUnlock() acct, ok := s.accounts[id] if !ok { - return nil, ErrAccountNotFound + return Account{}, ErrAccountNotFound } - return acct, nil + return *acct, nil } // --------------------------------------------------------------------------- @@ -290,39 +290,39 @@ type PostTransactionRequest struct { // // Internally, balances are stored as signed values where positive means // a balance in the account's normal direction. -func (s *Service) PostTransaction(req PostTransactionRequest) (*Transaction, error) { +func (s *Service) PostTransaction(req PostTransactionRequest) (Transaction, error) { s.mu.Lock() defer s.mu.Unlock() // Validate: non-empty entries. if len(req.Entries) == 0 { - return nil, ErrEmptyTransaction + return Transaction{}, ErrEmptyTransaction } // Validate: amounts. for _, e := range req.Entries { if e.Amount <= 0 { - return nil, ErrInvalidAmount + return Transaction{}, ErrInvalidAmount } } // Validate: all referenced accounts exist. for _, e := range req.Entries { if _, ok := s.accounts[e.AccountID]; !ok { - return nil, ErrAccountNotFound + return Transaction{}, ErrAccountNotFound } } // Validate: idempotency key. if req.IdempotencyKey != "" { if _, ok := s.idempotencyIndex[req.IdempotencyKey]; ok { - return nil, ErrDuplicateIdempotencyKey + return Transaction{}, ErrDuplicateIdempotencyKey } } // Validate: balanced. if err := validateBalance(req.Entries); err != nil { - return nil, err + return Transaction{}, err } // Set defaults for dates. @@ -339,12 +339,12 @@ func (s *Service) PostTransaction(req PostTransactionRequest) (*Transaction, err // Assign IDs to entries. entries := make([]Entry, len(req.Entries)) for i, e := range req.Entries { - e.ID = s.nextID("ent") + e.ID = EntryID(s.nextID("ent")) entries[i] = e } tx := &Transaction{ - ID: s.nextID("tx"), + ID: TransactionID(s.nextID("tx")), IdempotencyKey: req.IdempotencyKey, Entries: entries, BookingDate: bookingDate, @@ -360,8 +360,23 @@ func (s *Service) PostTransaction(req PostTransactionRequest) (*Transaction, err s.idempotencyIndex[req.IdempotencyKey] = tx.ID } - s.appendAudit(EventTransactionPosted, tx.ID, tx) - return tx, nil + s.appendAudit(EventTransactionPosted, string(tx.ID), tx) + return copyTransaction(tx), nil +} + +// copyTransaction returns a deep copy of a Transaction, including its +// Entries slice and Metadata map. +func copyTransaction(tx *Transaction) Transaction { + cp := *tx + cp.Entries = make([]Entry, len(tx.Entries)) + copy(cp.Entries, tx.Entries) + if tx.Metadata != nil { + cp.Metadata = make(map[string]string, len(tx.Metadata)) + for k, v := range tx.Metadata { + cp.Metadata[k] = v + } + } + return cp } // validateBalance checks that total debits equal total credits. @@ -386,28 +401,28 @@ func validateBalance(entries []Entry) error { // GetTransaction retrieves a transaction by its ID. // Returns ErrTransactionNotFound if the transaction does not exist. -func (s *Service) GetTransaction(id string) (*Transaction, error) { +func (s *Service) GetTransaction(id TransactionID) (Transaction, error) { s.mu.RLock() defer s.mu.RUnlock() tx, ok := s.transactions[id] if !ok { - return nil, ErrTransactionNotFound + return Transaction{}, ErrTransactionNotFound } - return tx, nil + return copyTransaction(tx), nil } // GetTransactionByIdempotencyKey retrieves a transaction by its idempotency key. // Returns ErrTransactionNotFound if no transaction with that key exists. -func (s *Service) GetTransactionByIdempotencyKey(key string) (*Transaction, error) { +func (s *Service) GetTransactionByIdempotencyKey(key string) (Transaction, error) { s.mu.RLock() defer s.mu.RUnlock() txID, ok := s.idempotencyIndex[key] if !ok { - return nil, ErrTransactionNotFound + return Transaction{}, ErrTransactionNotFound } - return s.transactions[txID], nil + return copyTransaction(s.transactions[txID]), nil } // --------------------------------------------------------------------------- @@ -432,16 +447,16 @@ func (s *Service) GetTransactionByIdempotencyKey(key string) (*Transaction, erro // Returns: // - ErrTransactionNotFound if the original does not exist. // - ErrTransactionAlreadyReversed if the original was already reversed. -func (s *Service) ReverseTransaction(txID, description string) (*Transaction, error) { +func (s *Service) ReverseTransaction(txID TransactionID, description string) (Transaction, error) { s.mu.Lock() defer s.mu.Unlock() original, ok := s.transactions[txID] if !ok { - return nil, ErrTransactionNotFound + return Transaction{}, ErrTransactionNotFound } if original.Status == Reversed { - return nil, ErrTransactionAlreadyReversed + return Transaction{}, ErrTransactionAlreadyReversed } // Build reversal entries: flip every direction. @@ -449,7 +464,7 @@ func (s *Service) ReverseTransaction(txID, description string) (*Transaction, er entries := make([]Entry, len(original.Entries)) for i, e := range original.Entries { entries[i] = Entry{ - ID: s.nextID("ent"), + ID: EntryID(s.nextID("ent")), AccountID: e.AccountID, Amount: e.Amount, Direction: e.Direction.Opposite(), @@ -457,7 +472,7 @@ func (s *Service) ReverseTransaction(txID, description string) (*Transaction, er } reversal := &Transaction{ - ID: s.nextID("tx"), + ID: TransactionID(s.nextID("tx")), Entries: entries, BookingDate: now, ValueDate: original.ValueDate, @@ -470,12 +485,12 @@ func (s *Service) ReverseTransaction(txID, description string) (*Transaction, er original.Status = Reversed s.transactions[reversal.ID] = reversal - s.appendAudit(EventTransactionReversed, original.ID, map[string]string{ - "original_id": original.ID, - "reversal_id": reversal.ID, + s.appendAudit(EventTransactionReversed, string(original.ID), map[string]string{ + "original_id": string(original.ID), + "reversal_id": string(reversal.ID), }) - return reversal, nil + return copyTransaction(reversal), nil } // --------------------------------------------------------------------------- @@ -485,7 +500,7 @@ func (s *Service) ReverseTransaction(txID, description string) (*Transaction, er // CreateHoldRequest contains all the parameters needed to create a hold. type CreateHoldRequest struct { // AccountID is the account whose available balance will be reduced. - AccountID string + AccountID AccountID // Amount is the positive hold amount in minor currency units. Amount Amount @@ -512,20 +527,20 @@ type CreateHoldRequest struct { // Returns: // - ErrAccountNotFound if the account does not exist. // - ErrInvalidAmount if the amount is not positive. -func (s *Service) CreateHold(req CreateHoldRequest) (*Hold, error) { +func (s *Service) CreateHold(req CreateHoldRequest) (Hold, error) { s.mu.Lock() defer s.mu.Unlock() if _, ok := s.accounts[req.AccountID]; !ok { - return nil, ErrAccountNotFound + return Hold{}, ErrAccountNotFound } if req.Amount <= 0 { - return nil, ErrInvalidAmount + return Hold{}, ErrInvalidAmount } now := s.now() h := &Hold{ - ID: s.nextID("hld"), + ID: HoldID(s.nextID("hld")), AccountID: req.AccountID, Amount: req.Amount, ExpiresAt: req.ExpiresAt, @@ -536,8 +551,8 @@ func (s *Service) CreateHold(req CreateHoldRequest) (*Hold, error) { s.holds[h.ID] = h s.accountHolds[req.AccountID] = append(s.accountHolds[req.AccountID], h.ID) - s.appendAudit(EventHoldCreated, h.ID, h) - return h, nil + s.appendAudit(EventHoldCreated, string(h.ID), h) + return *h, nil } // ReleaseHold cancels an active hold, restoring the available balance. @@ -549,7 +564,7 @@ func (s *Service) CreateHold(req CreateHoldRequest) (*Hold, error) { // Returns: // - ErrHoldNotFound if the hold does not exist. // - ErrHoldNotActive if the hold has already been released or captured. -func (s *Service) ReleaseHold(holdID string) error { +func (s *Service) ReleaseHold(holdID HoldID) error { s.mu.Lock() defer s.mu.Unlock() @@ -562,7 +577,7 @@ func (s *Service) ReleaseHold(holdID string) error { } h.Status = HoldReleased - s.appendAudit(EventHoldReleased, h.ID, h) + s.appendAudit(EventHoldReleased, string(h.ID), h) return nil } @@ -590,20 +605,20 @@ func (s *Service) ReleaseHold(holdID string) error { // - ErrHoldNotFound if the hold does not exist. // - ErrHoldNotActive if the hold has already been released or captured. // - ErrAccountNotFound if the counterparty account does not exist. -func (s *Service) CaptureHold(holdID, counterpartyAccountID string, captureAmount Amount, description string) (*Transaction, error) { +func (s *Service) CaptureHold(holdID HoldID, counterpartyAccountID AccountID, captureAmount Amount, description string) (Transaction, error) { s.mu.Lock() defer s.mu.Unlock() h, ok := s.holds[holdID] if !ok { - return nil, ErrHoldNotFound + return Transaction{}, ErrHoldNotFound } if h.Status != HoldActive { - return nil, ErrHoldNotActive + return Transaction{}, ErrHoldNotActive } if _, ok := s.accounts[counterpartyAccountID]; !ok { - return nil, ErrAccountNotFound + return Transaction{}, ErrAccountNotFound } if captureAmount <= 0 { @@ -620,13 +635,13 @@ func (s *Service) CaptureHold(holdID, counterpartyAccountID string, captureAmoun now := s.now() entries := []Entry{ { - ID: s.nextID("ent"), + ID: EntryID(s.nextID("ent")), AccountID: h.AccountID, Amount: captureAmount, Direction: holdDirection, }, { - ID: s.nextID("ent"), + ID: EntryID(s.nextID("ent")), AccountID: counterpartyAccountID, Amount: captureAmount, Direction: counterDirection, @@ -634,7 +649,7 @@ func (s *Service) CaptureHold(holdID, counterpartyAccountID string, captureAmoun } tx := &Transaction{ - ID: s.nextID("tx"), + ID: TransactionID(s.nextID("tx")), Entries: entries, BookingDate: now, ValueDate: now, @@ -646,26 +661,26 @@ func (s *Service) CaptureHold(holdID, counterpartyAccountID string, captureAmoun h.Status = HoldCaptured s.transactions[tx.ID] = tx - s.appendAudit(EventHoldCaptured, h.ID, map[string]string{ - "hold_id": h.ID, - "transaction_id": tx.ID, + s.appendAudit(EventHoldCaptured, string(h.ID), map[string]string{ + "hold_id": string(h.ID), + "transaction_id": string(tx.ID), }) - s.appendAudit(EventTransactionPosted, tx.ID, tx) + s.appendAudit(EventTransactionPosted, string(tx.ID), tx) - return tx, nil + return copyTransaction(tx), nil } // GetHold retrieves a hold by its ID. // Returns ErrHoldNotFound if the hold does not exist. -func (s *Service) GetHold(id string) (*Hold, error) { +func (s *Service) GetHold(id HoldID) (Hold, error) { s.mu.RLock() defer s.mu.RUnlock() h, ok := s.holds[id] if !ok { - return nil, ErrHoldNotFound + return Hold{}, ErrHoldNotFound } - return h, nil + return *h, nil } // --------------------------------------------------------------------------- @@ -687,19 +702,19 @@ func (s *Service) GetHold(id string) (*Hold, error) { // the amount that can actually be used for new transactions. // // Returns ErrAccountNotFound if the account does not exist. -func (s *Service) GetBalance(accountID string) (*Balance, error) { +func (s *Service) GetBalance(accountID AccountID) (Balance, error) { s.mu.RLock() defer s.mu.RUnlock() acct, ok := s.accounts[accountID] if !ok { - return nil, ErrAccountNotFound + return Balance{}, ErrAccountNotFound } book := s.computeBookBalance(accountID, acct.Type) holds := s.computeActiveHolds(accountID) - return &Balance{ + return Balance{ Book: book, Holds: holds, Available: book - holds, @@ -722,7 +737,7 @@ func (s *Service) GetBalance(accountID string) (*Balance, error) { // The Reversed status is informational — the corresponding reversal // transaction's entries are what actually cancel out the original's // balance impact. This preserves the full audit trail. -func (s *Service) computeBookBalance(accountID string, accountType AccountType) Amount { +func (s *Service) computeBookBalance(accountID AccountID, accountType AccountType) Amount { var balance Amount normal := accountType.NormalBalance() @@ -743,7 +758,7 @@ func (s *Service) computeBookBalance(accountID string, accountType AccountType) } // computeActiveHolds sums all active, non-expired holds for an account. -func (s *Service) computeActiveHolds(accountID string) Amount { +func (s *Service) computeActiveHolds(accountID AccountID) Amount { var total Amount now := s.now() @@ -786,13 +801,13 @@ func (s *Service) computeActiveHolds(accountID string) Amount { // overwritten (useful for end-of-day recalculation after late postings). // // Returns ErrAccountNotFound if the account does not exist. -func (s *Service) TakeEndOfDaySnapshot(accountID string, date time.Time) (*BalanceSnapshot, error) { +func (s *Service) TakeEndOfDaySnapshot(accountID AccountID, date time.Time) (BalanceSnapshot, error) { s.mu.Lock() defer s.mu.Unlock() acct, ok := s.accounts[accountID] if !ok { - return nil, ErrAccountNotFound + return BalanceSnapshot{}, ErrAccountNotFound } book := s.computeBookBalance(accountID, acct.Type) @@ -816,31 +831,31 @@ func (s *Service) TakeEndOfDaySnapshot(accountID string, date time.Time) (*Balan dateKey := date.Format("2006-01-02") s.snapshots[accountID][dateKey] = snap - s.appendAudit(EventSnapshotTaken, accountID, snap) - return snap, nil + s.appendAudit(EventSnapshotTaken, string(accountID), snap) + return *snap, nil } // GetSnapshot retrieves an end-of-day balance snapshot for an account // and business date. // -// Returns nil if no snapshot exists for the given parameters. +// Returns ErrSnapshotNotFound if no snapshot exists for the given parameters. // Returns ErrAccountNotFound if the account does not exist. -func (s *Service) GetSnapshot(accountID string, date time.Time) (*BalanceSnapshot, error) { +func (s *Service) GetSnapshot(accountID AccountID, date time.Time) (BalanceSnapshot, error) { s.mu.RLock() defer s.mu.RUnlock() if _, ok := s.accounts[accountID]; !ok { - return nil, ErrAccountNotFound + return BalanceSnapshot{}, ErrAccountNotFound } dateKey := date.Format("2006-01-02") if byAccount, ok := s.snapshots[accountID]; ok { if snap, ok := byAccount[dateKey]; ok { - return snap, nil + return *snap, nil } } - return nil, nil + return BalanceSnapshot{}, ErrSnapshotNotFound } // --------------------------------------------------------------------------- @@ -859,26 +874,28 @@ func (s *Service) GetSnapshot(accountID string, date time.Time) (*BalanceSnapsho // // In a production system, the audit log would typically be stored in a // separate, write-once data store with strict access controls. -func (s *Service) GetAuditLog() []*AuditEvent { +func (s *Service) GetAuditLog() []AuditEvent { s.mu.RLock() defer s.mu.RUnlock() - // Return a copy to prevent external mutation. - result := make([]*AuditEvent, len(s.auditLog)) - copy(result, s.auditLog) + // Return copies to prevent external mutation. + result := make([]AuditEvent, len(s.auditLog)) + for i, e := range s.auditLog { + result[i] = *e + } return result } // GetAuditLogForEntity returns all audit events related to a specific // entity, identified by its ID. -func (s *Service) GetAuditLogForEntity(entityID string) []*AuditEvent { +func (s *Service) GetAuditLogForEntity(entityID string) []AuditEvent { s.mu.RLock() defer s.mu.RUnlock() - var result []*AuditEvent + var result []AuditEvent for _, e := range s.auditLog { if e.EntityID == entityID { - result = append(result, e) + result = append(result, *e) } } return result diff --git a/service_test.go b/service_test.go index 1e0af39f..5fbe794e 100644 --- a/service_test.go +++ b/service_test.go @@ -30,7 +30,7 @@ func testService(t *testing.T) *Service { // │ └── Cash (Asset) // └── Revenue (subledger) // └── Fee Income (Revenue) -func setupChartOfAccounts(t *testing.T, svc *Service) (alice, bob, cash, feeIncome *Account) { +func setupChartOfAccounts(t *testing.T, svc *Service) (alice, bob, cash, feeIncome Account) { t.Helper() gl, err := svc.CreateLedger("General Ledger") @@ -822,13 +822,10 @@ func TestEndOfDaySnapshot(t *testing.T) { assertNoError(t, err) assertEqual(t, "retrieved book", got.Balance.Book, Amount(10000)) - // Non-existent snapshot returns nil. + // Non-existent snapshot returns ErrSnapshotNotFound. otherDate := time.Date(2025, 1, 16, 0, 0, 0, 0, time.UTC) - missing, err := svc.GetSnapshot(alice.ID, otherDate) - assertNoError(t, err) - if missing != nil { - t.Fatal("expected nil snapshot for non-existent date") - } + _, err = svc.GetSnapshot(alice.ID, otherDate) + assertError(t, err, ErrSnapshotNotFound) } func TestEndOfDaySnapshot_AccountNotFound(t *testing.T) { @@ -923,12 +920,12 @@ func TestAuditLogForEntity(t *testing.T) { }) // Get events for Alice's account. - aliceEvents := svc.GetAuditLogForEntity(alice.ID) + aliceEvents := svc.GetAuditLogForEntity(string(alice.ID)) assertEqual(t, "alice events", len(aliceEvents), 1) assertEqual(t, "event type", aliceEvents[0].Type, EventAccountCreated) // Get events for the transaction. - txEvents := svc.GetAuditLogForEntity(tx.ID) + txEvents := svc.GetAuditLogForEntity(string(tx.ID)) assertEqual(t, "tx events", len(txEvents), 1) assertEqual(t, "event type", txEvents[0].Type, EventTransactionPosted) } @@ -941,8 +938,8 @@ func TestAuditLog_ImmutableCopy(t *testing.T) { log2 := svc.GetAuditLog() // Modifying the returned slice should not affect the internal log. - log1[0] = nil - if log2[0] == nil { + log1[0].Type = "tampered" + if log2[0].Type == "tampered" { t.Fatal("audit log returned mutable reference") } } @@ -986,7 +983,7 @@ func TestGetBalance_AllAccountTypes(t *testing.T) { revenue, _ := svc.CreateAccount(sl.ID, "Revenue", Revenue) expense, _ := svc.CreateAccount(sl.ID, "Expense", Expense) - accounts := []*Account{asset, liability, equity, revenue, expense} + accounts := []Account{asset, liability, equity, revenue, expense} // Post a debit of 100 and credit of 100 between pairs. // Debit asset, credit liability. @@ -1162,13 +1159,13 @@ func TestFullBankingWorkflow(t *testing.T) { // Helper functions for tests // --------------------------------------------------------------------------- -func findAccountByName(t *testing.T, svc *Service, name string) *Account { +func findAccountByName(t *testing.T, svc *Service, name string) Account { t.Helper() for _, acct := range svc.accounts { if acct.Name == name { - return acct + return *acct } } t.Fatalf("account %q not found", name) - return nil + return Account{} } diff --git a/types.go b/types.go index 0fcdbe8f..d80efe80 100644 --- a/types.go +++ b/types.go @@ -2,6 +2,18 @@ package ledger import "time" +// ID types for each entity. These are defined types (not aliases) so the +// compiler prevents accidentally passing e.g. a HoldID where an AccountID +// is expected. +type ( + LedgerID string + SubledgerID string + AccountID string + TransactionID string + EntryID string + HoldID string +) + // Amount represents a monetary value in the smallest unit of the currency // (e.g., cents for USD, pence for GBP). This is the standard approach // used by most payment systems and banks. @@ -70,23 +82,23 @@ func (d Direction) Opposite() Direction { // Ledger is a top-level grouping for accounts (e.g., "General Ledger"). type Ledger struct { - ID string + ID LedgerID Name string CreatedAt time.Time } // Subledger is a subdivision of a ledger (e.g., "Accounts Receivable"). type Subledger struct { - ID string - LedgerID string + ID SubledgerID + LedgerID LedgerID Name string CreatedAt time.Time } // Account is a financial account within a subledger. type Account struct { - ID string - SubledgerID string + ID AccountID + SubledgerID SubledgerID Name string Type AccountType CreatedAt time.Time @@ -95,8 +107,8 @@ type Account struct { // Entry is a single leg of a transaction, representing a debit or credit // to an account. type Entry struct { - ID string - AccountID string + ID EntryID + AccountID AccountID Amount Amount Direction Direction } @@ -119,7 +131,7 @@ func (s TransactionStatus) String() string { // Transaction is a multi-legged accounting entry. All entries within a // transaction must balance (total debits = total credits). type Transaction struct { - ID string + ID TransactionID IdempotencyKey string Entries []Entry BookingDate time.Time // When the transaction was recorded in the system @@ -130,7 +142,7 @@ type Transaction struct { CreatedAt time.Time // ReversalOf is set when this transaction is a reversal of another. - ReversalOf string + ReversalOf TransactionID } // HoldStatus tracks the lifecycle of a hold. @@ -158,8 +170,8 @@ func (s HoldStatus) String() string { // Hold represents a pending authorization that reduces the available // balance of an account without affecting the book balance. type Hold struct { - ID string - AccountID string + ID HoldID + AccountID AccountID Amount Amount ExpiresAt time.Time Description string @@ -177,7 +189,7 @@ type Balance struct { // BalanceSnapshot is a point-in-time record of an account's balance, // taken at end-of-day for a given value date. type BalanceSnapshot struct { - AccountID string + AccountID AccountID Date time.Time // The business day this snapshot represents Balance Balance TakenAt time.Time // When the snapshot was actually taken From 261ee9437259ad3e4f4206080d49c414c97757a8 Mon Sep 17 00:00:00 2001 From: Claude Date: Sat, 28 Feb 2026 10:51:07 +0000 Subject: [PATCH 2/2] Enforce ErrInsufficientBalance for Asset and Expense accounts PostTransaction and CreateHold now check that the available balance (book minus active holds) would not go negative for Asset and Expense accounts. Liability, Equity, and Revenue accounts are not checked. https://claude.ai/code/session_01947Cs6jqVPDDzmAqQMoB35 --- service.go | 54 +++++++++++++++- service_test.go | 163 ++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 216 insertions(+), 1 deletion(-) diff --git a/service.go b/service.go index 456fb7c6..8ebc5434 100644 --- a/service.go +++ b/service.go @@ -274,6 +274,7 @@ type PostTransactionRequest struct { // 3. All referenced accounts must exist. // 4. If an idempotency key is provided, it must not already be used. // 5. Total debits must equal total credits. +// 6. Asset and Expense accounts must have sufficient available balance. // // If all validations pass, the entries are atomically applied to the // account balances and the transaction is recorded. @@ -325,6 +326,11 @@ func (s *Service) PostTransaction(req PostTransactionRequest) (Transaction, erro return Transaction{}, err } + // Validate: sufficient balance for Asset and Expense accounts. + if err := s.checkSufficientBalance(req.Entries); err != nil { + return Transaction{}, err + } + // Set defaults for dates. now := s.now() bookingDate := req.BookingDate @@ -399,6 +405,40 @@ func validateBalance(entries []Entry) error { return nil } +// checkSufficientBalance verifies that the entries would not cause any +// Asset or Expense account's available balance to go below zero. +// Liability, Equity, and Revenue accounts are not checked. +func (s *Service) checkSufficientBalance(entries []Entry) error { + // Compute the net balance impact per account. + impact := make(map[AccountID]Amount) + for _, e := range entries { + acct := s.accounts[e.AccountID] + if e.Direction == acct.Type.NormalBalance() { + impact[e.AccountID] += e.Amount + } else { + impact[e.AccountID] -= e.Amount + } + } + + for accountID, delta := range impact { + acct := s.accounts[accountID] + if acct.Type != Asset && acct.Type != Expense { + continue + } + // Only check when the transaction decreases the balance. + if delta >= 0 { + continue + } + book := s.computeBookBalance(accountID, acct.Type) + holds := s.computeActiveHolds(accountID) + available := book - holds + if available+delta < 0 { + return ErrInsufficientBalance + } + } + return nil +} + // GetTransaction retrieves a transaction by its ID. // Returns ErrTransactionNotFound if the transaction does not exist. func (s *Service) GetTransaction(id TransactionID) (Transaction, error) { @@ -527,17 +567,29 @@ type CreateHoldRequest struct { // Returns: // - ErrAccountNotFound if the account does not exist. // - ErrInvalidAmount if the amount is not positive. +// - ErrInsufficientBalance if the hold would overdraw an Asset or Expense account. func (s *Service) CreateHold(req CreateHoldRequest) (Hold, error) { s.mu.Lock() defer s.mu.Unlock() - if _, ok := s.accounts[req.AccountID]; !ok { + acct, ok := s.accounts[req.AccountID] + if !ok { return Hold{}, ErrAccountNotFound } if req.Amount <= 0 { return Hold{}, ErrInvalidAmount } + // Check sufficient balance for Asset and Expense accounts. + if acct.Type == Asset || acct.Type == Expense { + book := s.computeBookBalance(req.AccountID, acct.Type) + holds := s.computeActiveHolds(req.AccountID) + available := book - holds + if available-req.Amount < 0 { + return Hold{}, ErrInsufficientBalance + } + } + now := s.now() h := &Hold{ ID: HoldID(s.nextID("hld")), diff --git a/service_test.go b/service_test.go index 5fbe794e..ad8c7c25 100644 --- a/service_test.go +++ b/service_test.go @@ -944,6 +944,169 @@ func TestAuditLog_ImmutableCopy(t *testing.T) { } } +// --------------------------------------------------------------------------- +// Insufficient Balance Tests +// --------------------------------------------------------------------------- + +// TestPostTransaction_InsufficientBalance_Asset tests that a transaction +// is rejected when it would cause an Asset account's available balance +// to go below zero. +func TestPostTransaction_InsufficientBalance_Asset(t *testing.T) { + svc := testService(t) + alice, _, cash, _ := setupChartOfAccounts(t, svc) + + // Fund cash account with $100. + svc.PostTransaction(PostTransactionRequest{ + Description: "Initial deposit", + Entries: []Entry{ + {AccountID: cash.ID, Amount: 10000, Direction: Debit}, + {AccountID: alice.ID, Amount: 10000, Direction: Credit}, + }, + }) + + // Try to withdraw more cash than available (credit cash $150). + _, err := svc.PostTransaction(PostTransactionRequest{ + Description: "Overdraw cash", + Entries: []Entry{ + {AccountID: alice.ID, Amount: 15000, Direction: Debit}, + {AccountID: cash.ID, Amount: 15000, Direction: Credit}, + }, + }) + assertError(t, err, ErrInsufficientBalance) + + // Cash balance should be unchanged. + bal, _ := svc.GetBalance(cash.ID) + assertEqual(t, "cash balance unchanged", bal.Book, Amount(10000)) +} + +// TestPostTransaction_InsufficientBalance_WithHolds tests that holds +// are considered when checking available balance for transactions. +func TestPostTransaction_InsufficientBalance_WithHolds(t *testing.T) { + svc := testService(t) + alice, _, _, _ := setupChartOfAccounts(t, svc) + + l, _ := svc.CreateLedger("Test") + sl, _ := svc.CreateSubledger(l.ID, "Test") + assetAcct, _ := svc.CreateAccount(sl.ID, "Test Asset", Asset) + + // Fund asset account with $100 using a Liability counterparty. + svc.PostTransaction(PostTransactionRequest{ + Description: "Fund", + Entries: []Entry{ + {AccountID: assetAcct.ID, Amount: 10000, Direction: Debit}, + {AccountID: alice.ID, Amount: 10000, Direction: Credit}, + }, + }) + + // Place hold for $60. + svc.CreateHold(CreateHoldRequest{ + AccountID: assetAcct.ID, + Amount: 6000, + }) + + // Try to withdraw $50 — book is $100, holds $60, available $40. + _, err := svc.PostTransaction(PostTransactionRequest{ + Description: "Exceeds available", + Entries: []Entry{ + {AccountID: assetAcct.ID, Amount: 5000, Direction: Credit}, + {AccountID: alice.ID, Amount: 5000, Direction: Debit}, + }, + }) + assertError(t, err, ErrInsufficientBalance) + + // A smaller withdrawal should succeed. + _, err = svc.PostTransaction(PostTransactionRequest{ + Description: "Within available", + Entries: []Entry{ + {AccountID: assetAcct.ID, Amount: 4000, Direction: Credit}, + {AccountID: alice.ID, Amount: 4000, Direction: Debit}, + }, + }) + assertNoError(t, err) +} + +// TestPostTransaction_InsufficientBalance_LiabilityNotChecked tests that +// Liability accounts are not subject to balance checking. +func TestPostTransaction_InsufficientBalance_LiabilityNotChecked(t *testing.T) { + svc := testService(t) + alice, _, cash, _ := setupChartOfAccounts(t, svc) + + // Debit Alice (Liability) without any prior credit — this should succeed + // because Liability accounts are not checked for insufficient balance. + _, err := svc.PostTransaction(PostTransactionRequest{ + Description: "Debit unfunded liability", + Entries: []Entry{ + {AccountID: alice.ID, Amount: 5000, Direction: Debit}, + {AccountID: cash.ID, Amount: 5000, Direction: Debit}, + }, + }) + // This will fail with ErrUnbalancedTransaction since both are debits, + // but let's do a proper balanced test. + assertError(t, err, ErrUnbalancedTransaction) + + // Proper test: debit a Liability with no prior balance. + l, _ := svc.CreateLedger("Test") + sl, _ := svc.CreateSubledger(l.ID, "Test") + liab, _ := svc.CreateAccount(sl.ID, "Test Liability", Liability) + + _, err = svc.PostTransaction(PostTransactionRequest{ + Description: "Debit unfunded liability", + Entries: []Entry{ + {AccountID: liab.ID, Amount: 5000, Direction: Debit}, + {AccountID: alice.ID, Amount: 5000, Direction: Credit}, + }, + }) + assertNoError(t, err) +} + +// TestCreateHold_InsufficientBalance tests that a hold is rejected when +// it would cause an Asset account's available balance to go below zero. +func TestCreateHold_InsufficientBalance(t *testing.T) { + svc := testService(t) + alice, _, _, _ := setupChartOfAccounts(t, svc) + + l, _ := svc.CreateLedger("Test") + sl, _ := svc.CreateSubledger(l.ID, "Test") + assetAcct, _ := svc.CreateAccount(sl.ID, "Test Asset", Asset) + + // Fund asset account with $100 using a Liability counterparty. + svc.PostTransaction(PostTransactionRequest{ + Description: "Fund", + Entries: []Entry{ + {AccountID: assetAcct.ID, Amount: 10000, Direction: Debit}, + {AccountID: alice.ID, Amount: 10000, Direction: Credit}, + }, + }) + + // Hold for $100 should succeed (exactly available). + _, err := svc.CreateHold(CreateHoldRequest{ + AccountID: assetAcct.ID, + Amount: 10000, + }) + assertNoError(t, err) + + // Another hold should fail — available is now $0. + _, err = svc.CreateHold(CreateHoldRequest{ + AccountID: assetAcct.ID, + Amount: 1, + }) + assertError(t, err, ErrInsufficientBalance) +} + +// TestCreateHold_LiabilityNotChecked tests that holds on Liability +// accounts are not subject to balance checking. +func TestCreateHold_LiabilityNotChecked(t *testing.T) { + svc := testService(t) + alice, _, _, _ := setupChartOfAccounts(t, svc) + + // Alice (Liability) has $0 balance, but hold should succeed. + _, err := svc.CreateHold(CreateHoldRequest{ + AccountID: alice.ID, + Amount: 5000, + }) + assertNoError(t, err) +} + // --------------------------------------------------------------------------- // Balance Edge Cases // ---------------------------------------------------------------------------