Skip to content

Commit cfcad87

Browse files
committed
Only run accountant audit for NTT/WTT when enabled
1 parent 835d232 commit cfcad87

3 files changed

Lines changed: 88 additions & 15 deletions

File tree

node/pkg/accountant/accountant.go

Lines changed: 7 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -225,9 +225,13 @@ func (acct *Accountant) Start(ctx context.Context) error {
225225
return fmt.Errorf("failed to start watcher: %w", err)
226226
}
227227

228-
if err := supervisor.Run(ctx, "acctaudit", common.WrapWithScissors(acct.audit, "acctaudit")); err != nil {
229-
return fmt.Errorf("failed to start audit worker: %w", err)
230-
}
228+
}
229+
}
230+
231+
// Start the audit worker if not mocking/testing and either Global or NTT accountant are enabled
232+
if acct.env != common.AccountantMock && acct.env != common.GoTest && (acct.baseEnabled() || acct.nttEnabled()) {
233+
if err := supervisor.Run(ctx, "acctaudit", common.WrapWithScissors(acct.audit, "acctaudit")); err != nil {
234+
return fmt.Errorf("failed to start audit worker: %w", err)
231235
}
232236
}
233237

node/pkg/accountant/accountant_test.go

Lines changed: 68 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -102,10 +102,19 @@ type AuditMockWormchainConn struct {
102102
allPendingTransfersErr error
103103
batchTransferStatusResp []byte
104104
batchTransferStatusErr error
105+
106+
queryLock sync.Mutex
107+
queries []string
108+
contractAddresses []string
105109
}
106110

107111
func (c *AuditMockWormchainConn) SubmitQuery(ctx context.Context, contractAddress string, query []byte) ([]byte, error) {
108112
queryStr := string(query)
113+
c.queryLock.Lock()
114+
c.queries = append(c.queries, queryStr)
115+
c.contractAddresses = append(c.contractAddresses, contractAddress)
116+
c.queryLock.Unlock()
117+
109118
if strings.Contains(queryStr, "all_pending_transfers") {
110119
return c.allPendingTransfersResp, c.allPendingTransfersErr
111120
}
@@ -115,6 +124,12 @@ func (c *AuditMockWormchainConn) SubmitQuery(ctx context.Context, contractAddres
115124
return []byte{}, nil
116125
}
117126

127+
func (c *AuditMockWormchainConn) QueryCount() int {
128+
c.queryLock.Lock()
129+
defer c.queryLock.Unlock()
130+
return len(c.queries)
131+
}
132+
118133
func newAccountantForTest(
119134
t *testing.T,
120135
logger *zap.Logger,
@@ -739,6 +754,20 @@ func newAccountantForAuditTest(
739754
obsvReqWriteC chan *gossipv1.ObservationRequest,
740755
acctWriteC chan *common.MessagePublication,
741756
wormchainConn *AuditMockWormchainConn,
757+
) *Accountant {
758+
return newAccountantForAuditModeTest(t, logger, ctx, obsvReqWriteC, acctWriteC, "0xdeadbeef", wormchainConn, "", nil)
759+
}
760+
761+
func newAccountantForAuditModeTest(
762+
t *testing.T,
763+
logger *zap.Logger,
764+
ctx context.Context,
765+
obsvReqWriteC chan *gossipv1.ObservationRequest,
766+
acctWriteC chan *common.MessagePublication,
767+
contract string,
768+
wormchainConn *AuditMockWormchainConn,
769+
nttContract string,
770+
nttWormchainConn *AuditMockWormchainConn,
742771
) *Accountant {
743772
var db guardianDB.MockAccountantDB
744773

@@ -759,12 +788,12 @@ func newAccountantForAuditTest(
759788
logger,
760789
&db,
761790
obsvReqWriteC,
762-
"0xdeadbeef",
791+
contract,
763792
"none",
764793
wormchainConn,
765794
true, // enforceFlag
766-
"",
767-
nil,
795+
nttContract,
796+
nttWormchainConn,
768797
guardianSigner,
769798
gst,
770799
acctWriteC,
@@ -785,6 +814,42 @@ var (
785814
testTxHash = []byte{0x06, 0xf5, 0x41, 0xf5, 0xec, 0xfc, 0x43, 0x40, 0x7c, 0x31, 0x58, 0x7a, 0xa6, 0xac, 0x3a, 0x68, 0x9e, 0x89, 0x60, 0xf3, 0x6d, 0xc2, 0x3c, 0x33, 0x2d, 0xb5, 0x51, 0x0d, 0xfc, 0x6a, 0x40, 0x64}
786815
)
787816

817+
func TestRunAuditBaseOnlySkipsNttAudit(t *testing.T) {
818+
ctx := context.Background()
819+
logger := zaptest.NewLogger(t)
820+
obsvReqWriteC := make(chan *gossipv1.ObservationRequest, 10)
821+
acctChan := make(chan *common.MessagePublication, MsgChannelCapacity)
822+
wormchainConn := &AuditMockWormchainConn{
823+
allPendingTransfersResp: []byte(`{"pending":[]}`),
824+
}
825+
826+
acct := newAccountantForAuditModeTest(t, logger, ctx, obsvReqWriteC, acctChan, "base-contract", wormchainConn, "", nil)
827+
828+
require.NotPanics(t, func() {
829+
acct.runAudit(ctx)
830+
})
831+
assert.Equal(t, 1, wormchainConn.QueryCount(), "expected base audit query")
832+
assert.Equal(t, []string{"base-contract"}, wormchainConn.contractAddresses)
833+
}
834+
835+
func TestRunAuditNttOnlySkipsBaseAudit(t *testing.T) {
836+
ctx := context.Background()
837+
logger := zaptest.NewLogger(t)
838+
obsvReqWriteC := make(chan *gossipv1.ObservationRequest, 10)
839+
acctChan := make(chan *common.MessagePublication, MsgChannelCapacity)
840+
nttWormchainConn := &AuditMockWormchainConn{
841+
allPendingTransfersResp: []byte(`{"pending":[]}`),
842+
}
843+
844+
acct := newAccountantForAuditModeTest(t, logger, ctx, obsvReqWriteC, acctChan, "", nil, "ntt-contract", nttWormchainConn)
845+
846+
require.NotPanics(t, func() {
847+
acct.runAudit(ctx)
848+
})
849+
assert.Equal(t, 1, nttWormchainConn.QueryCount(), "expected NTT audit query")
850+
assert.Equal(t, []string{"ntt-contract"}, nttWormchainConn.contractAddresses)
851+
}
852+
788853
func TestPerformAuditResubmitsUnsignedTransfer(t *testing.T) {
789854
ctx := context.Background()
790855
logger := zaptest.NewLogger(t)

node/pkg/accountant/audit.go

Lines changed: 13 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -155,15 +155,19 @@ func (acct *Accountant) audit(ctx context.Context) error {
155155

156156
// runAudit is the entry point for the audit of the pending transfer map. It creates a temporary map of all pending transfers and invokes the main audit function.
157157
func (acct *Accountant) runAudit(ctx context.Context) {
158-
knownPendingTransferMap := acct.createAuditMap(false)
159-
acct.logger.Debug("in AuditPendingTransfers: starting base audit", zap.Int("numPending", numPendingEntries(knownPendingTransferMap)))
160-
acct.performAudit(ctx, knownPendingTransferMap, acct.wormchainConn, acct.contract)
161-
acct.logger.Debug("in AuditPendingTransfers: finished base audit")
162-
163-
knownPendingNttTransferMap := acct.createAuditMap(true)
164-
acct.logger.Debug("in AuditPendingTransfers: starting ntt audit", zap.Int("numPending", numPendingEntries(knownPendingNttTransferMap)))
165-
acct.performAudit(ctx, knownPendingNttTransferMap, acct.nttWormchainConn, acct.nttContract)
166-
acct.logger.Debug("in AuditPendingTransfers: finished ntt audit")
158+
if acct.baseEnabled() {
159+
knownPendingTransferMap := acct.createAuditMap(false)
160+
acct.logger.Debug("in AuditPendingTransfers: starting base audit", zap.Int("numPending", numPendingEntries(knownPendingTransferMap)))
161+
acct.performAudit(ctx, knownPendingTransferMap, acct.wormchainConn, acct.contract)
162+
acct.logger.Debug("in AuditPendingTransfers: finished base audit")
163+
}
164+
165+
if acct.nttEnabled() {
166+
knownPendingNttTransferMap := acct.createAuditMap(true)
167+
acct.logger.Debug("in AuditPendingTransfers: starting ntt audit", zap.Int("numPending", numPendingEntries(knownPendingNttTransferMap)))
168+
acct.performAudit(ctx, knownPendingNttTransferMap, acct.nttWormchainConn, acct.nttContract)
169+
acct.logger.Debug("in AuditPendingTransfers: finished ntt audit")
170+
}
167171
}
168172

169173
// numPendingEntries returns the total number of non-nil pending entries across all keys

0 commit comments

Comments
 (0)