Skip to content
Merged
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
32 changes: 30 additions & 2 deletions block/manager.go
Original file line number Diff line number Diff line change
Expand Up @@ -603,7 +603,7 @@ func (m *Manager) publishBlockInternal(ctx context.Context) error {
}
}

signature, err = m.getSignature(header.Header)
signature, err = m.getHeaderSignature(header.Header)
if err != nil {
return err
}
Expand Down Expand Up @@ -906,7 +906,7 @@ func bytesToBatchData(data []byte) ([][]byte, error) {
return result, nil
}

func (m *Manager) getSignature(header types.Header) (types.Signature, error) {
func (m *Manager) getHeaderSignature(header types.Header) (types.Signature, error) {
b, err := header.MarshalBinary()
if err != nil {
return nil, err
Expand All @@ -917,6 +917,17 @@ func (m *Manager) getSignature(header types.Header) (types.Signature, error) {
return m.signer.Sign(b)
}

func (m *Manager) getDataSignature(data *types.Data) (types.Signature, error) {
dataBz, err := data.MarshalBinary()
if err != nil {
return nil, err
}
if m.signer == nil {
return nil, fmt.Errorf("signer is nil; cannot sign data")
}
return m.signer.Sign(dataBz)
}

// NotifyNewTransactions signals that new transactions are available for processing
// This method will be called by the Reaper when it receives new transactions
func (m *Manager) NotifyNewTransactions() {
Expand Down Expand Up @@ -977,3 +988,20 @@ func (m *Manager) SaveCache() error {
}
return nil
}

// isValidSignedData returns true if the data signature is valid for the expected sequencer.
func (m *Manager) isValidSignedData(signedData *types.SignedData) bool {
if signedData == nil || signedData.Txs == nil {
return false
}
if !bytes.Equal(signedData.Signer.Address, m.genesis.ProposerAddress) {
return false
}
dataBytes, err := signedData.Data.MarshalBinary()
if err != nil {
return false
}

valid, err := signedData.Signer.PubKey.Verify(dataBytes, signedData.Signature)
return err == nil && valid
}
135 changes: 135 additions & 0 deletions block/manager_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -365,3 +365,138 @@ func TestBytesToBatchData(t *testing.T) {
assert.Error(err)
assert.Contains(err.Error(), "corrupted data")
}

// TestGetDataSignature_Success ensures a valid signature is returned when the signer is set.
func TestGetDataSignature_Success(t *testing.T) {
require := require.New(t)
mockDAC := mocks.NewDA(t)
m, _ := getManager(t, mockDAC, -1, -1)

privKey, _, err := crypto.GenerateKeyPair(crypto.Ed25519, 256)
require.NoError(err)
signer, err := noopsigner.NewNoopSigner(privKey)
require.NoError(err)
m.signer = signer
_, data := types.GetRandomBlock(1, 2, "TestGetDataSignature")
sig, err := m.getDataSignature(data)
require.NoError(err)
require.NotEmpty(sig)
}

// TestGetDataSignature_NilSigner ensures the correct error is returned when the signer is nil.
func TestGetDataSignature_NilSigner(t *testing.T) {
require := require.New(t)
mockDAC := mocks.NewDA(t)
m, _ := getManager(t, mockDAC, -1, -1)

privKey, _, err := crypto.GenerateKeyPair(crypto.Ed25519, 256)
require.NoError(err)
signer, err := noopsigner.NewNoopSigner(privKey)
require.NoError(err)
m.signer = signer
_, data := types.GetRandomBlock(1, 2, "TestGetDataSignature")

m.signer = nil
_, err = m.getDataSignature(data)
require.ErrorContains(err, "signer is nil; cannot sign data")
}

// TestIsValidSignedData covers valid, nil, wrong proposer, and invalid signature cases for isValidSignedData.
func TestIsValidSignedData(t *testing.T) {
require := require.New(t)
privKey, _, err := crypto.GenerateKeyPair(crypto.Ed25519, 256)
require.NoError(err)
testSigner, err := noopsigner.NewNoopSigner(privKey)
require.NoError(err)
proposerAddr, err := testSigner.GetAddress()
require.NoError(err)
gen := genesispkg.NewGenesis(
"testchain",
1,
time.Now(),
proposerAddr,
)
m := &Manager{
signer: testSigner,
genesis: gen,
}

t.Run("valid signed data", func(t *testing.T) {
batch := &types.Data{
Txs: types.Txs{types.Tx("tx1"), types.Tx("tx2")},
}
sig, err := m.getDataSignature(batch)
require.NoError(err)
pubKey, err := m.signer.GetPublic()
require.NoError(err)
signedData := &types.SignedData{
Data: *batch,
Signature: sig,
Signer: types.Signer{
PubKey: pubKey,
Address: proposerAddr,
},
}
assert.True(t, m.isValidSignedData(signedData))
})

t.Run("nil signed data", func(t *testing.T) {
assert.False(t, m.isValidSignedData(nil))
})

t.Run("nil Txs", func(t *testing.T) {
signedData := &types.SignedData{
Data: types.Data{},
Signer: types.Signer{
Address: proposerAddr,
},
}
signedData.Txs = nil
assert.False(t, m.isValidSignedData(signedData))
})

t.Run("wrong proposer address", func(t *testing.T) {
batch := &types.Data{
Txs: types.Txs{types.Tx("tx1")},
}
sig, err := m.getDataSignature(batch)
require.NoError(err)
pubKey, err := m.signer.GetPublic()
require.NoError(err)
wrongAddr := make([]byte, len(proposerAddr))
copy(wrongAddr, proposerAddr)
wrongAddr[0] ^= 0xFF // flip a bit
signedData := &types.SignedData{
Data: *batch,
Signature: sig,
Signer: types.Signer{
PubKey: pubKey,
Address: wrongAddr,
},
}
assert.False(t, m.isValidSignedData(signedData))
})

t.Run("invalid signature", func(t *testing.T) {
batch := &types.Data{
Txs: types.Txs{types.Tx("tx1")},
}
sig, err := m.getDataSignature(batch)
require.NoError(err)
pubKey, err := m.signer.GetPublic()
require.NoError(err)
// Corrupt the signature
badSig := make([]byte, len(sig))
copy(badSig, sig)
badSig[0] ^= 0xFF
signedData := &types.SignedData{
Data: *batch,
Signature: badSig,
Signer: types.Signer{
PubKey: pubKey,
Address: proposerAddr,
},
}
assert.False(t, m.isValidSignedData(signedData))
})
}
40 changes: 23 additions & 17 deletions block/retriever.go
Original file line number Diff line number Diff line change
Expand Up @@ -85,7 +85,7 @@ func (m *Manager) processNextDAHeaderAndData(ctx context.Context) error {
if m.handlePotentialHeader(ctx, bz, daHeight) {
continue
}
m.handlePotentialBatch(ctx, bz, daHeight)
m.handlePotentialData(ctx, bz, daHeight)
}
return nil
}
Expand Down Expand Up @@ -117,6 +117,11 @@ func (m *Manager) handlePotentialHeader(ctx context.Context, bz []byte, daHeight
m.logger.Debug("failed to decode unmarshalled header", "error", err)
return true
}
// Stronger validation: check for obviously invalid headers using ValidateBasic
if err := header.ValidateBasic(); err != nil {
m.logger.Debug("blob does not look like a valid header", "daHeight", daHeight, "error", err)
return false
}
// early validation to reject junk headers
if !m.isUsingExpectedSingleSequencer(header) {
m.logger.Debug("skipping header from unexpected sequencer",
Expand All @@ -140,36 +145,37 @@ func (m *Manager) handlePotentialHeader(ctx context.Context, bz []byte, daHeight
return true
}

// handlePotentialBatch tries to decode and process a batch. No return value.
func (m *Manager) handlePotentialBatch(ctx context.Context, bz []byte, daHeight uint64) {
var batchPb pb.Batch
err := proto.Unmarshal(bz, &batchPb)
// handlePotentialData tries to decode and process a data. No return value.
func (m *Manager) handlePotentialData(ctx context.Context, bz []byte, daHeight uint64) {
var signedData types.SignedData
err := signedData.UnmarshalBinary(bz)
if err != nil {
m.logger.Debug("failed to unmarshal batch", "error", err)
m.logger.Debug("failed to unmarshal signed data", "error", err)
return
}
if len(batchPb.Txs) == 0 {
m.logger.Debug("ignoring empty batch", "daHeight", daHeight)
if len(signedData.Txs) == 0 {
m.logger.Debug("ignoring empty signed data", "daHeight", daHeight)
return
}
data := &types.Data{
Txs: make(types.Txs, len(batchPb.Txs)),
}
for i, tx := range batchPb.Txs {
data.Txs[i] = types.Tx(tx)

// Early validation to reject junk data
if !m.isValidSignedData(&signedData) {
m.logger.Debug("invalid data signature", "daHeight", daHeight)
return
}
dataHashStr := data.DACommitment().String()

dataHashStr := signedData.Data.DACommitment().String()
m.dataCache.SetDAIncluded(dataHashStr)
m.sendNonBlockingSignalToDAIncluderCh()
m.logger.Info("batch marked as DA included", "batchHash", dataHashStr, "daHeight", daHeight)
m.logger.Info("signed data marked as DA included", "dataHash", dataHashStr, "daHeight", daHeight)
if !m.dataCache.IsSeen(dataHashStr) {
select {
case <-ctx.Done():
return
default:
m.logger.Warn("dataInCh backlog full, dropping batch", "daHeight", daHeight)
m.logger.Warn("dataInCh backlog full, dropping signed data", "daHeight", daHeight)
}
m.dataInCh <- NewDataEvent{data, daHeight}
m.dataInCh <- NewDataEvent{&signedData.Data, daHeight}
}
}

Expand Down
Loading