diff --git a/p2p/subscriber.go b/p2p/subscriber.go index 7ecfde1d..58635b33 100644 --- a/p2p/subscriber.go +++ b/p2p/subscriber.go @@ -138,10 +138,18 @@ func (s *Subscriber[H]) Broadcast(ctx context.Context, header H, opts ...pubsub. if err != nil { return err } + + opts = append(opts, pubsub.WithValidatorData(header)) return s.topic.Publish(ctx, bin, opts...) } func (s *Subscriber[H]) verifyMessage(ctx context.Context, p peer.ID, msg *pubsub.Message) (res pubsub.ValidationResult) { + if msg.ValidatorData != nil { + // means the message is local and was already validated + // so simply accept it + return pubsub.ValidationAccept + } + defer func() { err := recover() if err != nil { diff --git a/p2p/subscription_test.go b/p2p/subscription_test.go index 9911503a..4a5b6d77 100644 --- a/p2p/subscription_test.go +++ b/p2p/subscription_test.go @@ -14,7 +14,7 @@ import ( "github.com/celestiaorg/go-header/headertest" ) -// TestSubscriber tests the header Module's implementation of Subscriber. +// TestSubscriber a simple test to check if the subscriber can receive headers from the network. func TestSubscriber(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), time.Second*15) defer cancel() @@ -73,23 +73,24 @@ func TestSubscriber(t *testing.T) { } // subscribe - _, err = p2pSub2.Subscribe() + senderSubscription, err := p2pSub2.Subscribe() require.NoError(t, err) subscription, err := p2pSub1.Subscribe() require.NoError(t, err) expectedHeader := suite.GenDummyHeaders(1)[0] - bin, err := expectedHeader.MarshalBinary() - require.NoError(t, err) - - err = p2pSub2.topic.Publish(ctx, bin, pubsub.WithReadiness(pubsub.MinTopicSize(1))) + err = p2pSub2.Broadcast(ctx, expectedHeader, pubsub.WithReadiness(pubsub.MinTopicSize(1))) require.NoError(t, err) // get next Header from network header, err := subscription.NextHeader(ctx) require.NoError(t, err) + assert.Equal(t, expectedHeader.Height(), header.Height()) + assert.Equal(t, expectedHeader.Hash(), header.Hash()) + header, err = senderSubscription.NextHeader(ctx) + require.NoError(t, err) assert.Equal(t, expectedHeader.Height(), header.Height()) assert.Equal(t, expectedHeader.Hash(), header.Hash()) }