Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions go/.changes/unreleased/fixed-20261002-124000.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
kind: fixed
body: Fixed refunds for persisted server-signed batch channels after a client restart.
time: 2026-10-02T12:40:00.000000Z
34 changes: 34 additions & 0 deletions go/mechanisms/svm/batch-settlement/client/harness_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,8 @@ func boolPtr(value bool) *bool { return &value }
type memoryStorage struct {
mu sync.Mutex
records map[string]BatchClientChannelRecord
gets []string
sets []string
}

func newMemoryStorage() *memoryStorage {
Expand All @@ -47,6 +49,7 @@ func newMemoryStorage() *memoryStorage {
func (s *memoryStorage) Get(key string) (*BatchClientChannelRecord, error) {
s.mu.Lock()
defer s.mu.Unlock()
s.gets = append(s.gets, key)
record, ok := s.records[key]
if !ok {
return nil, nil
Expand All @@ -58,6 +61,7 @@ func (s *memoryStorage) Get(key string) (*BatchClientChannelRecord, error) {
func (s *memoryStorage) Set(key string, record BatchClientChannelRecord) error {
s.mu.Lock()
defer s.mu.Unlock()
s.sets = append(s.sets, key)
s.records[key] = record
return nil
}
Expand All @@ -78,6 +82,36 @@ func (s *memoryStorage) only() BatchClientChannelRecord {
return BatchClientChannelRecord{}
}

func (s *memoryStorage) got(key string) bool {
s.mu.Lock()
defer s.mu.Unlock()
for _, got := range s.gets {
if got == key {
return true
}
}
return false
}

func (s *memoryStorage) setCount() int {
s.mu.Lock()
defer s.mu.Unlock()
return len(s.sets)
}

func (s *memoryStorage) size() int {
s.mu.Lock()
defer s.mu.Unlock()
return len(s.records)
}

func (s *memoryStorage) resetCalls() {
s.mu.Lock()
defer s.mu.Unlock()
s.gets = nil
s.sets = nil
}

type rpcStub struct {
mu sync.Mutex
owner solana.PublicKey
Expand Down
61 changes: 61 additions & 0 deletions go/mechanisms/svm/batch-settlement/client/lifecycle_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@ func TestBatchClientLifecycle(t *testing.T) {
t.Run("validates client terms and configuration boundaries", testValidatesTerms)
t.Run("builds a refund from a cached channel and rejects a missing one", testRefundFromCache)
t.Run("refunds a client-signed channel when the probe lists server-signed first", testRefundServerFirst)
t.Run("refunds a persisted server-signed channel after restart without discovery", testRefundPersistedServerSignedAfterRestart)
t.Run("reports no channel when the probe is server-signed and nothing is open", testRefundNoChannelServerProbe)
}

Expand Down Expand Up @@ -561,6 +562,66 @@ func testRefundServerFirst(t *testing.T) {
require.Equal(t, "1000", nestedString(t, cooperative.Payload, "voucher", "maxClaimableAmount"))
}

func testRefundPersistedServerSignedAfterRestart(t *testing.T) {
h := newHarness(t)
operator := newKey(t)
storage := newMemoryStorage()
serverRequirements := h.requirements("", map[string]any{
batchsettlement.ExtraOperator: operator.Address().String(),
batchsettlement.ExtraVoucherSigner: batchsettlement.VoucherSignerServer,
})
keyed := h.scheme(t, &BatchSvmClientConfig{
ChannelStorage: storage,
DiscoverChannels: boolPtr(false),
ServerSignedChannelsPolicy: &BatchServerSignedChannelsPolicy{
AllowedOperators: []string{operator.Address().String()},
},
})
key := keyed.channelKey(serverRequirements, h.feePayer.String(), 900)
storage.records[key] = BatchClientChannelRecord{
ChannelConfig: batchsettlement.BatchChannelConfig{
OpenSlot: 123,
Payer: h.payer.Address().String(),
PayerAuthorizer: operator.Address().String(),
Receiver: svm.USDCMainnetAddress,
ReceiverAuthorizer: h.receiverAuthorizer.String(),
Salt: "0",
Token: testMint,
VoucherSigner: batchsettlement.VoucherSignerServer,
WithdrawDelay: 900,
},
ChannelID: svm.USDCMainnetAddress,
ChargedCumulativeAmount: "1000",
Deposit: "5000",
}

untrusted := h.scheme(t, &BatchSvmClientConfig{
ChannelStorage: storage,
DiscoverChannels: boolPtr(false),
})
_, err := untrusted.CreateRefundPayload(context.Background(), 2, serverRequirements, RefundPayloadOptions{})
require.ErrorIs(t, err, ErrNoBatchChannelToRefund)
require.False(t, storage.got(key))
storage.resetCalls()

restarted := h.scheme(t, &BatchSvmClientConfig{
ChannelStorage: storage,
DiscoverChannels: boolPtr(false),
ServerSignedChannelsPolicy: &BatchServerSignedChannelsPolicy{
AllowedOperators: []string{operator.Address().String()},
},
})
cooperative, err := restarted.CreateRefundPayload(context.Background(), 2, serverRequirements, RefundPayloadOptions{})
require.NoError(t, err)
require.Equal(t, 2, cooperative.X402Version)
require.Equal(t, "refund", cooperative.Payload["type"])
require.Equal(t, "0", nestedString(t, cooperative.Payload, "authorization", "authorizedAmount"))
require.Equal(t, svm.USDCMainnetAddress, nestedString(t, cooperative.Payload, "authorization", "channelId"))
require.True(t, storage.got(key))
require.Zero(t, storage.setCount())
require.Equal(t, 1, storage.size())
}

func testRefundNoChannelServerProbe(t *testing.T) {
h := newHarness(t)
operator := newKey(t)
Expand Down
14 changes: 7 additions & 7 deletions go/mechanisms/svm/batch-settlement/client/scheme_state.go
Original file line number Diff line number Diff line change
Expand Up @@ -204,17 +204,17 @@ func (s *BatchSvmScheme) locateRefundChannel(
}
return nil, lookup, resolvedTerms{}, serverErr
}
discovered, discoverErr := s.discoverChannel(ctx, requirements, serverTerms)
if discoverErr != nil {
return nil, lookup, resolvedTerms{}, discoverErr
serverChannel, loadErr := s.loadRefundChannel(ctx, requirements, serverTerms, nil)
if loadErr != nil {
return nil, lookup, resolvedTerms{}, loadErr
}
if discovered != nil {
lookup = AlignRefundRequirements(requirements, discovered.tracker.ChannelConfig)
terms, err = s.resolveRefundTerms(ctx, lookup, discovered)
if serverChannel != nil {
lookup = AlignRefundRequirements(requirements, serverChannel.tracker.ChannelConfig)
terms, err = s.resolveRefundTerms(ctx, lookup, serverChannel)
if err != nil {
return nil, lookup, resolvedTerms{}, err
}
return discovered, lookup, terms, nil
return serverChannel, lookup, terms, nil
}
}

Expand Down
Loading