Skip to content
Open
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
40 changes: 38 additions & 2 deletions channel/persistence/keyvalue/persistrestorer_internal_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,9 @@ import (
"github.com/stretchr/testify/require"

_ "perun.network/go-perun/backend/sim" // backend init
"perun.network/go-perun/channel/persistence/test"
"perun.network/go-perun/channel"
ptest "perun.network/go-perun/channel/persistence/test"
wiretest "perun.network/go-perun/wire/test"
"polycry.pt/poly-go/sortedkv"
"polycry.pt/poly-go/sortedkv/leveldb"
"polycry.pt/poly-go/sortedkv/memorydb"
Expand All @@ -44,7 +46,7 @@ func TestPersistRestorer_Generic(t *testing.T) {
defer func() { require.NoError(t, db.Close()) }()
pr := NewPersistRestorer(db)
rng := pkgtest.Prng(t, i)
test.GenericPersistRestorerTest(
ptest.GenericPersistRestorerTest(
context.Background(),
t,
rng,
Expand All @@ -63,3 +65,37 @@ func TestChannelIterator_Next_Empty(t *testing.T) {
assert.False(t, success)
require.NoError(t, it.err)
}

func TestPersistRestorer_RestoreChannelRejectsUnexpectedFirstKey(t *testing.T) {
db := memorydb.NewDatabase()
pr := NewPersistRestorer(db)
defer func() { require.NoError(t, pr.Close()) }()

var id channel.ID
require.NoError(t, pr.channelDB(id).Put("bogus", "x"))

ch, err := pr.RestoreChannel(context.Background(), id)
require.Nil(t, ch)
require.ErrorContains(t, err, `unexpected iterator key`)
require.ErrorContains(t, err, `:bogus`)
require.ErrorContains(t, err, `expected suffix "current"`)
}

func TestPersistRestorer_RestoreChannelRejectsTrailingBytesInCurrent(t *testing.T) {
db := memorydb.NewDatabase()
pr := NewPersistRestorer(db)
defer func() { require.NoError(t, pr.Close()) }()

rng := pkgtest.Prng(t)
client := ptest.NewClient(context.Background(), t, rng, pr)
peer := wiretest.NewRandomAddress(rng)
ch := client.NewChannel(t, peer, nil)

current, err := pr.channelDB(ch.ID()).GetBytes("current")
require.NoError(t, err)
require.NoError(t, pr.channelDB(ch.ID()).Put("current", string(append(current, 0xFF))))

restored, err := pr.RestoreChannel(context.Background(), ch.ID())
require.Nil(t, restored)
require.ErrorContains(t, err, "decoding current incomplete")
}
65 changes: 52 additions & 13 deletions channel/persistence/keyvalue/restorer.go
Original file line number Diff line number Diff line change
Expand Up @@ -153,17 +153,32 @@ func (i *ChannelIterator) Next(context.Context) bool {
}

i.ch = persistence.NewChannel()
if !i.decodeNext("current", &i.ch.CurrentTXV, allowEnd) ||
current, ok := i.readNextValue("current", allowEnd)
if !ok ||
!i.decodeNext("index", &i.ch.IdxV, noOpts) ||
!i.decodeNext("params", i.ch.ParamsV, noOpts) ||
!i.decodeNext("parent", optChannelIDDec{&i.ch.Parent}, noOpts) ||
!i.decodeNext("peers", (*wire.AddressMapArray)(&i.ch.PeersV), noOpts) ||
!i.decodeNext("phase", &i.ch.PhaseV, noOpts) {
return false
}
currentTXDec := channel.TransactionDec{Tx: &i.ch.CurrentTXV, Parts: i.ch.ParamsV.Parts}
if err := currentTXDec.Decode(current); err != nil {
i.err = errors.WithMessage(err, "decoding current")
return false
}
if current.Len() != 0 {
i.err = errors.Errorf("decoding current incomplete (%d bytes left)", current.Len())
return false
}
i.ch.StagingTXV.Sigs = make([]wallet.Sig, len(i.ch.ParamsV.Parts))
for idx, key := range sigKeys(len(i.ch.ParamsV.Parts)) {
i.decodeNext(key, wallet.SigDec{Sig: &i.ch.StagingTXV.Sigs[idx]}, allowEmpty)
if !i.decodeNext(key, wallet.SigDec{
Sig: &i.ch.StagingTXV.Sigs[idx],
BackendID: participantBackendID(i.ch.ParamsV.Parts[idx]),
}, allowEmpty) {
return false
}
}

return i.decodeNext("staging:state", &PersistedState{&i.ch.StagingTXV.State}, allowEmpty)
Expand Down Expand Up @@ -210,28 +225,52 @@ func (i *ChannelIterator) recoverFromEmptyIterator(key string, allowedToEnd decO
// an iterator ends in the middle of decoding a channel, then the channel
// iterator's error is set. Returns whether a value was decoded without error.
func (i *ChannelIterator) decodeNext(key string, v interface{}, opts decOpts) bool {
buf, ok := i.readNextValue(key, opts)
if !ok {
return false
}
if buf == nil {
return true
}
origLen := buf.Len()
i.err = errors.WithMessage(perunio.Decode(buf, v), "decoding "+key)
if i.err != nil {
i.err = errors.WithMessagef(i.err, "value length %d", origLen)
return false
}
if buf.Len() != 0 {
i.err = errors.Errorf("decoding %s incomplete (%d bytes left)", key, buf.Len())
}

return i.err == nil
}

func (i *ChannelIterator) readNextValue(key string, opts decOpts) (*bytes.Buffer, bool) {
for !i.its[0].Next() {
if !i.recoverFromEmptyIterator(key, opts) {
return false
return nil, false
}
}
if actual := i.its[0].Key(); !strings.HasSuffix(actual, key) {
i.err = errors.Errorf("unexpected iterator key %q, expected suffix %q", actual, key)
Comment on lines +254 to +255

Copilot AI Apr 8, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

readNextValue uses strings.HasSuffix(actual, key) to validate the iterator key. This can incorrectly accept unexpected keys that merely end with the expected key (e.g., ...:mycurrent would pass for key="current"), weakening the new guard.

Consider additionally checking that the byte immediately preceding the suffix is the channel DB separator (":"), or otherwise validating the exact key format produced by channelDB(...)+":"+key.

Suggested change
if actual := i.its[0].Key(); !strings.HasSuffix(actual, key) {
i.err = errors.Errorf("unexpected iterator key %q, expected suffix %q", actual, key)
expectedSuffix := ":" + key
if actual := i.its[0].Key(); !strings.HasSuffix(actual, expectedSuffix) {
i.err = errors.Errorf("unexpected iterator key %q, expected suffix %q", actual, expectedSuffix)

Copilot uses AI. Check for mistakes.
return nil, false
}

buf := bytes.NewBuffer(i.its[0].ValueBytes())
if buf.Len() == 0 {
if allowEmpty.isSetIn(opts) {
return true
return nil, true
}
i.err = errors.Errorf("unexpected empty value")
return false
return nil, false
}
return buf, true
}

i.err = errors.WithMessage(perunio.Decode(buf, v), "decoding "+key)
if i.err != nil {
return false
}
if buf.Len() != 0 {
i.err = errors.Errorf("decoding %s incomplete (%d bytes left)", key, buf.Len())
func participantBackendID(part map[wallet.BackendID]wallet.Address) *wallet.BackendID {
backendID, ok := wallet.SingleBackendID(part)
if !ok {
return nil
}

return i.err == nil
return &backendID
}
37 changes: 30 additions & 7 deletions channel/transaction.go
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,17 @@ type Transaction struct {
Sigs []wallet.Sig
}

var _ perunio.Serializer = (*Transaction)(nil)
// TransactionDec is a helper type to decode a transaction while optionally
// using participant backend IDs for signature decoding.
type TransactionDec struct {
Tx *Transaction
Parts []map[wallet.BackendID]wallet.Address
}

var (
_ perunio.Serializer = (*Transaction)(nil)
_ perunio.Decoder = TransactionDec{}
)

// Clone returns a deep copy of Transaction.
func (t Transaction) Clone() Transaction {
Expand All @@ -57,6 +67,12 @@ func (t Transaction) Encode(w io.Writer) error {

// Decode decodes a transaction from an `io.Reader` or returns an `error`.
func (t *Transaction) Decode(r io.Reader) error {
return TransactionDec{Tx: t}.Decode(r)
}

// Decode decodes a transaction and uses participant backend IDs for signature
// decoding when they are known.
func (d TransactionDec) Decode(r io.Reader) error {
// Decode stateSet
var stateSet uint8
if err := perunio.Decode(r, &stateSet); err != nil {
Expand All @@ -65,20 +81,27 @@ func (t *Transaction) Decode(r io.Reader) error {

switch stateSet {
case 0:
t.State = nil
d.Tx.State = nil
return nil
case 1:
default:
return errors.Errorf("unknown stateSet value: %v", stateSet)
}

// Decode State
t.State = new(State)
if err := perunio.Decode(r, t.State); err != nil {
d.Tx.State = new(State)
if err := perunio.Decode(r, d.Tx.State); err != nil {
return errors.WithMessage(err, "decoding state")
}

t.Sigs = make([]wallet.Sig, t.NumParts())

return wallet.DecodeSparseSigs(r, &t.Sigs)
d.Tx.Sigs = make([]wallet.Sig, d.Tx.NumParts())
// A nil or empty Parts slice means that no backend context is available for
// signature decoding, so decoding falls back to the global wallet decoder.
if len(d.Parts) == 0 {
return wallet.DecodeSparseSigs(r, &d.Tx.Sigs)
}
if len(d.Parts) != d.Tx.NumParts() {
return errors.Errorf("participant count mismatch: state has %d participants, params have %d", d.Tx.NumParts(), len(d.Parts))
}
return wallet.DecodeSparseSigsForParts(r, &d.Tx.Sigs, d.Parts)
}
51 changes: 51 additions & 0 deletions channel/transaction_dec_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,51 @@
// Copyright 2026 - See NOTICE file for copyright holders.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.

package channel_test

import (
"bytes"
"testing"

"github.com/stretchr/testify/require"

_ "perun.network/go-perun/backend/sim/wallet"
"perun.network/go-perun/channel"
ctest "perun.network/go-perun/channel/test"
"perun.network/go-perun/wallet"
wallettest "perun.network/go-perun/wallet/test"
pkgtest "polycry.pt/poly-go/test"
)

func TestTransactionDecDecodeUsesParticipantBackends(t *testing.T) {
rng := pkgtest.Prng(t)
accs, addrs := wallettest.NewRandomAccounts(rng, 2, channel.TestBackendID)
params := ctest.NewRandomParams(rng, ctest.WithParts(addrs))
state := ctest.NewRandomState(rng, ctest.WithID(params.ID()), ctest.WithNumParts(len(addrs)))

sigs := make([]wallet.Sig, len(addrs))
for i := range addrs {
sig, err := channel.Sign(accs[i][channel.TestBackendID], state, channel.TestBackendID)
require.NoError(t, err)
sigs[i] = sig
}

var encoded bytes.Buffer
original := channel.Transaction{State: state, Sigs: sigs}
require.NoError(t, original.Encode(&encoded))

var decoded channel.Transaction
require.NoError(t, (channel.TransactionDec{Tx: &decoded, Parts: addrs}).Decode(&encoded))
require.Equal(t, original, decoded)
}
12 changes: 12 additions & 0 deletions wallet/address.go
Original file line number Diff line number Diff line change
Expand Up @@ -83,6 +83,18 @@ func IndexOfAddrs(addrs []map[BackendID]Address, addr map[BackendID]Address) int
return -1
}

// SingleBackendID returns the single backend ID contained in the participant
// map. It returns false if the participant exposes zero or multiple backends.
func SingleBackendID(part map[BackendID]Address) (BackendID, bool) {
if len(part) != 1 {
return 0, false
}
for backendID := range part {
return backendID, true
}
return 0, false
}

// CloneAddress returns a clone of an Address using its binary marshaling
// implementation. It panics if an error occurs during binary (un)marshaling.
func CloneAddress(a Address) Address {
Expand Down
10 changes: 10 additions & 0 deletions wallet/backend.go
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,16 @@ func NewAddress(id BackendID) Address {
return backend[id].NewAddress()
}

// decodeSigForBackend calls DecodeSig of the given backend and returns an error
// if no backend is registered for the id.
func decodeSigForBackend(r io.Reader, id BackendID) (Sig, error) {
b := backend[id]
if b == nil {
return nil, fmt.Errorf("no wallet backend registered for id %d", id)
}
return b.DecodeSig(r)
}

// DecodeSig calls DecodeSig of all Backends and returns an error if none return a valid signature.
func DecodeSig(r io.Reader) (Sig, error) {
var err error
Expand Down
44 changes: 36 additions & 8 deletions wallet/sig.go
Original file line number Diff line number Diff line change
Expand Up @@ -47,13 +47,19 @@ const bitsPerByte = 8

// SigDec is a helper type to decode signatures.
type SigDec struct {
Sig *Sig
BackendID int
Sig *Sig
// BackendID optionally selects the backend-specific signature decoder. If it
// is nil, decoding falls back to the global multi-backend decoder.
BackendID *BackendID
}
Comment on lines 49 to 54

Copilot AI Apr 8, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

SigDec is an exported type, and changing BackendID from an int/value field to a *BackendID is a breaking API change for downstream users constructing wallet.SigDec{...}.

If preserving backwards compatibility matters, consider keeping the old field (deprecated) or adding a new field (e.g., BackendIDPtr *BackendID) / constructor helper, and translating internally.

Copilot uses AI. Check for mistakes.

// Decode decodes a single signature.
func (s SigDec) Decode(r io.Reader) (err error) {
*s.Sig, err = DecodeSig(r)
if s.BackendID == nil {
*s.Sig, err = DecodeSig(r)
return err
}
*s.Sig, err = decodeSigForBackend(r, *s.BackendID)
return err
}

Expand Down Expand Up @@ -86,6 +92,28 @@ func EncodeSparseSigs(w io.Writer, sigs []Sig) error {

// DecodeSparseSigs decodes a collection of signatures in the form (mask, sig, ...).
func DecodeSparseSigs(r io.Reader, sigs *[]Sig) (err error) {
return decodeSparseSigs(r, sigs, func(r io.Reader, _ int) (Sig, error) {
return DecodeSig(r)
})
}

// DecodeSparseSigsForParts decodes a sparse signature collection using the
// participant backend IDs when they are known. If a participant exposes zero or
// multiple backend IDs, it falls back to the global decoder.
func DecodeSparseSigsForParts(r io.Reader, sigs *[]Sig, parts []map[BackendID]Address) (err error) {
if len(*sigs) != len(parts) {
return errors.Errorf("signature/participant count mismatch: %d != %d", len(*sigs), len(parts))
}

return decodeSparseSigs(r, sigs, func(r io.Reader, sigIdx int) (Sig, error) {
if backendID, ok := SingleBackendID(parts[sigIdx]); ok {
return decodeSigForBackend(r, backendID)
}
return DecodeSig(r)
})
}

func decodeSparseSigs(r io.Reader, sigs *[]Sig, decoder func(io.Reader, int) (Sig, error)) (err error) {
masklen := int(math.Ceil(float64(len(*sigs)) / float64(bitsPerByte)))
mask := make([]uint8, masklen)

Expand All @@ -100,11 +128,11 @@ func DecodeSparseSigs(r io.Reader, sigs *[]Sig) (err error) {
for bitIdx := 0; bitIdx < bitsPerByte && sigIdx < len(*sigs); bitIdx, sigIdx = bitIdx+1, sigIdx+1 {
if ((mask[maskIdx] >> bitIdx) % binaryModulo) == 0 {
(*sigs)[sigIdx] = nil
} else {
(*sigs)[sigIdx], err = DecodeSig(r)
if err != nil {
return errors.WithMessagef(err, "decoding signature %d", sigIdx)
}
continue
}
(*sigs)[sigIdx], err = decoder(r, sigIdx)
if err != nil {
return errors.WithMessagef(err, "decoding signature %d", sigIdx)
}
}
}
Expand Down
Loading
Loading