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..8ebc5434 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 } // --------------------------------------------------------------------------- @@ -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. @@ -290,39 +291,44 @@ 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 + } + + // Validate: sufficient balance for Asset and Expense accounts. + if err := s.checkSufficientBalance(req.Entries); err != nil { + return Transaction{}, err } // Set defaults for dates. @@ -339,12 +345,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 +366,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. @@ -384,30 +405,64 @@ 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 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 +487,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 +504,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 +512,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 +525,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 +540,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 +567,32 @@ 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) { +// - 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 { - return nil, ErrAccountNotFound + acct, ok := s.accounts[req.AccountID] + if !ok { + return Hold{}, ErrAccountNotFound } if req.Amount <= 0 { - return nil, ErrInvalidAmount + 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: s.nextID("hld"), + ID: HoldID(s.nextID("hld")), AccountID: req.AccountID, Amount: req.Amount, ExpiresAt: req.ExpiresAt, @@ -536,8 +603,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 +616,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 +629,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 +657,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 +687,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 +701,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 +713,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 +754,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 +789,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 +810,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 +853,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 +883,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 +926,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..ad8c7c25 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,12 +938,175 @@ 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") } } +// --------------------------------------------------------------------------- +// 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 // --------------------------------------------------------------------------- @@ -986,7 +1146,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 +1322,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