diff --git a/comm.go b/comm.go index 6efc933d..a0f51b5d 100644 --- a/comm.go +++ b/comm.go @@ -1,19 +1,10 @@ package pubsub import ( - "context" - "encoding/binary" - "io" - "time" - - pool "github.com/libp2p/go-buffer-pool" - "github.com/multiformats/go-varint" - "google.golang.org/protobuf/proto" - "github.com/libp2p/go-libp2p/core/network" - "github.com/libp2p/go-libp2p/core/peer" - "github.com/libp2p/go-msgio" + "google.golang.org/protobuf/proto" + "github.com/libp2p/go-libp2p-pubsub/internal/peercomm" pb "github.com/libp2p/go-libp2p-pubsub/pb" ) @@ -52,226 +43,42 @@ func (p *PubSub) getHelloPacket() *RPC { return &rpc } -func (p *PubSub) handleNewStream(s network.Stream) { - peer := s.Conn().RemotePeer() - done := make(chan struct{}) - sentNewStream := false - - defer func() { - p.inboundStreamsMx.Lock() - if p.inboundStreams[peer].s == s { - delete(p.inboundStreams, peer) - } - p.inboundStreamsMx.Unlock() - - if sentNewStream { - select { - case p.incoming <- incomingUnion{kind: incomingKindClosedStream, s: s}: - case <-p.ctx.Done(): - } - } - - close(done) - }() - - p.inboundStreamsMx.Lock() - prev, hasPrev := p.inboundStreams[peer] - p.inboundStreams[peer] = inboundHandler{s: s, done: done} - p.inboundStreamsMx.Unlock() - - if hasPrev { - p.logger.Debug("duplicate inbound stream; replacing handler", "peer", peer) - prev.s.Reset() - select { - case <-prev.done: - case <-p.ctx.Done(): - return - } - } - +func (p *PubSub) enqueuePeerEvent(event incomingUnion) { select { - case p.incoming <- incomingUnion{kind: incomingKindNewStream, s: s}: - sentNewStream = true + case p.incoming <- event: case <-p.ctx.Done(): - // Close is useless because the other side isn't reading. - s.Reset() - return - } - - r := msgio.NewVarintReaderSize(s, p.maxMessageSize) - for { - // Peek at the message length to know when we should mark the start time - // for measuring how long it took to receive a message. - _, _ = r.NextMsgLen() - start := time.Now() - msgbytes, err := r.ReadMsg() - if err != nil { - r.ReleaseMsg(msgbytes) - if err != io.EOF { - s.Reset() - p.rpcLogger.Debug("error reading rpc", "from", s.Conn().RemotePeer(), "err", err) - } else { - // Just be nice. They probably won't read this - // but it doesn't hurt to send it. - s.Close() - } - - return - } - if len(msgbytes) == 0 { - continue - } - - err = pb.ValidateRawRPCControlMessageSize(msgbytes, p.maxControlMessageSize) - if err != nil { - r.ReleaseMsg(msgbytes) - s.Reset() - - p.rpcLogger.Warn("RPC's control message is too large. Disconnecting. If this is a mistake, you should increase `WithMaxControlMessageSize`", "peer", s.Conn().RemotePeer(), "err", err) - return - } - - rpc := new(RPC) - err = proto.Unmarshal(msgbytes, &rpc.RPC) - r.ReleaseMsg(msgbytes) - if err != nil { - s.Reset() - - p.rpcLogger.Warn("bogus rpc from", "peer", s.Conn().RemotePeer(), "err", err) - return - } - - timeToReceive := time.Since(start) - p.rpcLogger.Debug("received", "peer", s.Conn().RemotePeer(), "duration_s", timeToReceive.Seconds(), "rpc", rpc) - - rpc.from = peer - select { - case p.incoming <- incomingUnion{ - kind: incomingKindRPC, - rpc: rpc, - }: - case <-p.ctx.Done(): - // Close is useless because the other side isn't reading. - s.Reset() - return - } - } -} - -func (p *PubSub) notifyPeerDead(pid peer.ID) { - p.peerDeadPrioLk.RLock() - p.peerDeadMx.Lock() - p.peerDeadPend[pid] = struct{}{} - p.peerDeadMx.Unlock() - p.peerDeadPrioLk.RUnlock() - - select { - case p.peerDead <- struct{}{}: - default: - } -} - -func (p *PubSub) handleNewPeer(ctx context.Context, pid peer.ID, outgoing *rpcQueue) { - s, err := p.host.NewStream(ctx, pid, p.rt.Protocols()...) - if err != nil { - p.logger.Debug("error opening new stream to peer", "err", err, "peer", pid) - - select { - case p.newPeerError <- pid: - case <-ctx.Done(): - } - - return - } - - firstMessage := make(chan *RPC, 1) - sCtx, cancel := context.WithCancel(ctx) - go p.handleSendingMessages(sCtx, s, outgoing, firstMessage) - go p.handlePeerDead(s) - select { - case p.newPeerStream <- peerOutgoingStream{Stream: s, FirstMessage: firstMessage, Cancel: cancel}: - case <-ctx.Done(): - cancel() } } -func (p *PubSub) handleNewPeerWithBackoff(ctx context.Context, pid peer.ID, backoff time.Duration, outgoing *rpcQueue) { +func (p *PubSub) handleNewStream(s network.Stream) { + response := make(chan *peercomm.Actor, 1) select { - case <-time.After(backoff): - p.handleNewPeer(ctx, pid, outgoing) - case <-ctx.Done(): - return - } -} - -func (p *PubSub) handlePeerDead(s network.Stream) { - pid := s.Conn().RemotePeer() - - _, err := s.Read([]byte{0}) - if err == nil { - p.logger.Debug("unexpected message from peer", "peer", pid) - } - - s.Reset() - p.notifyPeerDead(pid) -} - -func (p *PubSub) handleSendingMessages(ctx context.Context, s network.Stream, outgoing *rpcQueue, firstMessage chan *RPC) { - writeRpc := func(rpc *RPC) error { - size := uint64(proto.Size(&rpc.RPC)) - - buf := pool.Get(varint.UvarintSize(size) + int(size)) - defer pool.Put(buf) - - n := binary.PutUvarint(buf, size) - out, err := proto.MarshalOptions{}.MarshalAppend(buf[:n], &rpc.RPC) - if err != nil { - return err - } - - if err := s.SetWriteDeadline(time.Now().Add(time.Second * 30)); err != nil { - p.rpcLogger.Debug("failed to set write deadline", "peer", s.Conn().RemotePeer(), "err", err) - return err + case p.eval <- func() { + pid := s.Conn().RemotePeer() + if p.blacklist.Contains(pid) { + response <- nil + return } - - _, err = s.Write(out) - if err != nil { - p.rpcLogger.Debug("failed to send message", "peer", s.Conn().RemotePeer(), "rpc", rpc, "err", err) - return err + a, ok := p.peerComm.Lookup(pid) + if !ok { + a = p.peerComm.GetOrCreate(pid) + _ = a.Start(p.openRequest(0)) } - p.rpcLogger.Debug("sent", "peer", s.Conn().RemotePeer(), "rpc", rpc) - return nil - } - - select { - case rpc := <-firstMessage: - if proto.Size(&rpc.RPC) > 0 { - err := writeRpc(rpc) - if err != nil { - s.Reset() - p.logger.Debug("error writing message to peer", "peer", s.Conn().RemotePeer(), "err", err) - return - } - } - case <-ctx.Done(): - s.Reset() + response <- a + }: + case <-p.ctx.Done(): + _ = s.Reset() return } - - defer s.Close() - for ctx.Err() == nil { - rpc, err := outgoing.Pop(ctx) - if err != nil { - p.logger.Debug("error popping message from the queue to send to peer", "peer", s.Conn().RemotePeer(), "err", err) - return - } - - err = writeRpc(rpc) - if err != nil { - s.Reset() - p.logger.Debug("error writing message to peer", "peer", s.Conn().RemotePeer(), "err", err) + select { + case a := <-response: + if a == nil { + _ = s.Reset() return } + a.HandleInbound(s) + case <-p.ctx.Done(): + _ = s.Reset() } } diff --git a/extensions.go b/extensions.go index 6b921391..dc97c0bb 100644 --- a/extensions.go +++ b/extensions.go @@ -86,6 +86,7 @@ type extensionsState struct { myExtensions PeerExtensions peerExtensions map[peer.ID]PeerExtensions // peer's extensions sentExtensions map[peer.ID]struct{} + activeExtensions map[peer.ID]PeerExtensions reportMisbehavior func(peer.ID) sendRPC func(p peer.ID, r *RPC, urgent bool) testExtension *testExtension @@ -98,6 +99,7 @@ func newExtensionsState(myExtensions PeerExtensions, reportMisbehavior func(peer myExtensions: myExtensions, peerExtensions: make(map[peer.ID]PeerExtensions), sentExtensions: make(map[peer.ID]struct{}), + activeExtensions: make(map[peer.ID]PeerExtensions), reportMisbehavior: reportMisbehavior, sendRPC: sendRPC, testExtension: nil, @@ -109,11 +111,7 @@ func (es *extensionsState) HandleRPC(rpc *RPC) error { // We know this is the first message because we didn't have extensions // for this peer, and we always set extensions on the first rpc. es.peerExtensions[rpc.from] = peerExtensionsFromRPC(rpc) - if _, ok := es.sentExtensions[rpc.from]; ok { - // We just finished both sending and receiving the extensions - // control message. - es.extensionsOnNewOutboundStream(rpc.from) - } + es.activatePeerExtensions(rpc.from) } else { // We already have an extension for this peer. If they send us another // extensions control message, that is a protocol error. We should @@ -130,6 +128,7 @@ func (es *extensionsState) OnNewIncomingStream(peer.ID, protocol.ID) { } func (es *extensionsState) OnClosedIncomingStream(id peer.ID, _ protocol.ID) { + es.deactivatePeerExtensions(id) delete(es.peerExtensions, id) if len(es.peerExtensions) == 0 { es.peerExtensions = make(map[peer.ID]PeerExtensions) @@ -141,48 +140,58 @@ func (es *extensionsState) OnNewOutboundStream(id peer.ID, helloPacket *RPC) *RP helloPacket = es.myExtensions.ExtendRPC(helloPacket) es.sentExtensions[id] = struct{}{} - if _, ok := es.peerExtensions[id]; ok { - // We've just finished sending and receiving the extensions control - // message. - es.extensionsOnNewOutboundStream(id) - } + es.activatePeerExtensions(id) return helloPacket } func (es *extensionsState) OnClosedOutboundStream(id peer.ID) { - _, recvdExt := es.peerExtensions[id] - _, sentExt := es.sentExtensions[id] - if recvdExt && sentExt { - // Add peer was previously called, so we need to call remove peer - es.extensionsOnClosedOutboundStream(id) - } + es.deactivatePeerExtensions(id) delete(es.sentExtensions, id) if len(es.sentExtensions) == 0 { es.sentExtensions = make(map[peer.ID]struct{}) } } -// extensionsOnNewOutboundStream is only called once we've both sent and received the -// extensions control message. -func (es *extensionsState) extensionsOnNewOutboundStream(id peer.ID) { - if es.myExtensions.TestExtension && es.peerExtensions[id].TestExtension { +func (es *extensionsState) activatePeerExtensions(id peer.ID) { + peerExtensions, received := es.peerExtensions[id] + _, sent := es.sentExtensions[id] + if !received || !sent { + return + } + + active := PeerExtensions{ + TestExtension: es.myExtensions.TestExtension && peerExtensions.TestExtension, + PartialMessages: es.myExtensions.PartialMessages && peerExtensions.PartialMessages, + } + es.activeExtensions[id] = active + + if active.TestExtension && es.testExtension != nil { es.testExtension.OnNewOutboundStream(id) } } -// extensionsOnClosedOutboundStream is always called after extensionsOnNewOutboundStream. -func (es *extensionsState) extensionsOnClosedOutboundStream(id peer.ID) { - if es.myExtensions.PartialMessages && es.peerExtensions[id].PartialMessages { +func (es *extensionsState) deactivatePeerExtensions(id peer.ID) { + active, ok := es.activeExtensions[id] + if !ok { + return + } + delete(es.activeExtensions, id) + if len(es.activeExtensions) == 0 { + es.activeExtensions = make(map[peer.ID]PeerExtensions) + } + + if active.PartialMessages && es.partialMessagesExtension != nil { es.partialMessagesExtension.OnClosedOutboundStream(id) } } func (es *extensionsState) extensionsHandleRPC(rpc *RPC) error { - if es.myExtensions.TestExtension && es.peerExtensions[rpc.from].TestExtension { + active := es.activeExtensions[rpc.from] + if active.TestExtension && es.testExtension != nil { es.testExtension.HandleRPC(rpc.from, rpc.TestExtension) } - if es.myExtensions.PartialMessages && es.peerExtensions[rpc.from].PartialMessages && rpc.Partial != nil { + if active.PartialMessages && rpc.Partial != nil && es.partialMessagesExtension != nil { err := es.partialMessagesExtension.HandleRPC(rpc.from, rpc.Partial) if err != nil { return err @@ -277,8 +286,9 @@ func (r partialMessageRouter) MeshPeers(topic string) iter.Seq[peer.ID] { } for peer := range peerSet { - if r.gs.extensions.peerExtensions[peer].PartialMessages && - (r.gs.iRequestPartial(topic) && r.gs.peerSupportsSendingPartial(peer, topic)) || (r.gs.iSupportSendingPartial(topic) && r.gs.peerRequestsPartial(peer, topic)) { + if r.gs.extensions.activeExtensions[peer].PartialMessages && + ((r.gs.iRequestPartial(topic) && r.gs.peerSupportsSendingPartial(peer, topic)) || + (r.gs.iSupportSendingPartial(topic) && r.gs.peerRequestsPartial(peer, topic))) { if !yield(peer) { return } diff --git a/extensions_test.go b/extensions_test.go new file mode 100644 index 00000000..46de2107 --- /dev/null +++ b/extensions_test.go @@ -0,0 +1,139 @@ +package pubsub + +import ( + "testing" + + pubsub_pb "github.com/libp2p/go-libp2p-pubsub/pb" + "github.com/libp2p/go-libp2p/core/peer" +) + +type lifecyclePartialMessages struct { + closed []peer.ID +} + +func (m *lifecyclePartialMessages) OnClosedOutboundStream(id peer.ID) { + m.closed = append(m.closed, id) +} + +func (*lifecyclePartialMessages) HandleRPC(peer.ID, *pubsub_pb.PartialMessagesExtension) error { + return nil +} + +func (*lifecyclePartialMessages) Heartbeat() {} + +func (*lifecyclePartialMessages) EmitGossip(string, []peer.ID) {} + +func newPartialLifecycleState(cleanup *lifecyclePartialMessages) *extensionsState { + es := newExtensionsState(PeerExtensions{PartialMessages: true}, nil, nil) + es.partialMessagesExtension = cleanup + return es +} + +func partialExtensionsHello(id peer.ID) *RPC { + enabled := true + return &RPC{ + RPC: pubsub_pb.RPC{Control: &pubsub_pb.ControlMessage{ + Extensions: &pubsub_pb.ControlExtensions{PartialMessages: &enabled}, + }}, + from: id, + } +} + +func activatePartialExtensions(t *testing.T, es *extensionsState, id peer.ID) { + t.Helper() + es.OnNewOutboundStream(id, &RPC{}) + if err := es.HandleRPC(partialExtensionsHello(id)); err != nil { + t.Fatalf("handle extensions hello: %v", err) + } + if !es.activeExtensions[id].PartialMessages { + t.Fatal("partial messages extension was not activated") + } +} + +func TestExtensionsDeactivateOnEitherHalfClosing(t *testing.T) { + for _, test := range []struct { + name string + close func(*extensionsState, peer.ID) + }{ + {"incoming", func(es *extensionsState, id peer.ID) { es.OnClosedIncomingStream(id, "") }}, + {"outbound", func(es *extensionsState, id peer.ID) { es.OnClosedOutboundStream(id) }}, + } { + t.Run(test.name, func(t *testing.T) { + cleanup := new(lifecyclePartialMessages) + es := newPartialLifecycleState(cleanup) + id := peer.ID("peer") + activatePartialExtensions(t, es, id) + + test.close(es, id) + if es.activeExtensions[id].PartialMessages { + t.Fatal("extension remained active after stream closure") + } + if len(cleanup.closed) != 1 || cleanup.closed[0] != id { + t.Fatalf("expected one cleanup for %q, got %v", id, cleanup.closed) + } + + test.close(es, id) + if len(cleanup.closed) != 1 { + t.Fatalf("duplicate closure triggered %d cleanups", len(cleanup.closed)) + } + }) + } +} + +func TestExtensionsReplacementHalfReactivates(t *testing.T) { + for _, test := range []struct { + name string + closeHalf func(*extensionsState, peer.ID) + replace func(*testing.T, *extensionsState, peer.ID) + }{ + { + name: "incoming", + closeHalf: func(es *extensionsState, id peer.ID) { es.OnClosedIncomingStream(id, "") }, + replace: func(t *testing.T, es *extensionsState, id peer.ID) { + t.Helper() + if err := es.HandleRPC(partialExtensionsHello(id)); err != nil { + t.Fatalf("handle replacement hello: %v", err) + } + }, + }, + { + name: "outbound", + closeHalf: func(es *extensionsState, id peer.ID) { es.OnClosedOutboundStream(id) }, + replace: func(t *testing.T, es *extensionsState, id peer.ID) { + t.Helper() + es.OnNewOutboundStream(id, &RPC{}) + }, + }, + } { + t.Run(test.name, func(t *testing.T) { + cleanup := new(lifecyclePartialMessages) + es := newPartialLifecycleState(cleanup) + id := peer.ID("peer") + activatePartialExtensions(t, es, id) + + test.closeHalf(es, id) + test.replace(t, es, id) + if !es.activeExtensions[id].PartialMessages { + t.Fatal("replacement half did not reactivate extension") + } + if len(cleanup.closed) != 1 { + t.Fatalf("expected one cleanup before reactivation, got %d", len(cleanup.closed)) + } + }) + } +} + +func TestExtensionsPartialCleanupUsesActiveSnapshot(t *testing.T) { + cleanup := new(lifecyclePartialMessages) + es := newPartialLifecycleState(cleanup) + id := peer.ID("peer") + activatePartialExtensions(t, es, id) + + es.peerExtensions[id] = PeerExtensions{} + es.myExtensions.PartialMessages = false + es.OnClosedOutboundStream(id) + + if len(cleanup.closed) != 1 || cleanup.closed[0] != id { + t.Fatalf("expected cleanup from negotiated snapshot, got %v", cleanup.closed) + } +} diff --git a/floodsub.go b/floodsub.go index 07a52c1c..e035becf 100644 --- a/floodsub.go +++ b/floodsub.go @@ -89,12 +89,12 @@ func (fs *FloodSubRouter) Publish(msg *Message) { continue } - q, ok := fs.p.peers[pid] + q, ok := fs.p.peerComm.Lookup(pid) if !ok { continue } - err := q.Push(out, false) + err := q.Send(&out.RPC, false) if err != nil { fs.p.logger.Info("dropping message to peer: queue full", "peer", pid) fs.tracer.DropRPC(out, pid) diff --git a/floodsub_test.go b/floodsub_test.go index d667d5d6..1f2e30f3 100644 --- a/floodsub_test.go +++ b/floodsub_test.go @@ -1284,6 +1284,13 @@ func TestDedupInboundStreams(t *testing.T) { h1 := hosts[0] h2 := hosts[1] + // Keep h1's outbound stream open so this test only exercises inbound + // stream replacement, not initial outbound-open failure retirement. + h2.SetStreamHandler(FloodSubID, func(s network.Stream) { + <-ctx.Done() + _ = s.Reset() + }) + _, err := NewFloodSub(ctx, h1) if err != nil { t.Fatal(err) diff --git a/gossipsub.go b/gossipsub.go index 37bbdb14..b76f987d 100644 --- a/gossipsub.go +++ b/gossipsub.go @@ -12,6 +12,7 @@ import ( "sort" "time" + "github.com/libp2p/go-libp2p-pubsub/internal/peercomm" pb "github.com/libp2p/go-libp2p-pubsub/pb" "github.com/libp2p/go-libp2p/core/event" @@ -1534,7 +1535,7 @@ func (gs *GossipSubRouter) sendPrune(p peer.ID, topic string, isUnsubscribe bool } func (gs *GossipSubRouter) sendRPC(p peer.ID, out *RPC, urgent bool) { - q, ok := gs.p.peers[p] + q, ok := gs.p.peerComm.Lookup(p) if !ok { // No queue to send to this peer. Nothing to do. gs.doDropRPC(out, p, "No send queue for peer. Can't send RPC") @@ -1602,13 +1603,8 @@ func (gs *GossipSubRouter) doDropRPC(rpc *RPC, p peer.ID, reason string) { } } -func (gs *GossipSubRouter) doSendRPC(rpc *RPC, p peer.ID, q *rpcQueue, urgent bool) { - var err error - if urgent { - err = q.UrgentPush(rpc, false) - } else { - err = q.Push(rpc, false) - } +func (gs *GossipSubRouter) doSendRPC(rpc *RPC, p peer.ID, q *peercomm.Actor, urgent bool) { + err := q.Send(&rpc.RPC, urgent) if err != nil { gs.doDropRPC(rpc, p, "queue full") return diff --git a/gossipsub_peer_lifecycle_test.go b/gossipsub_peer_lifecycle_test.go index 82ad1f3b..74589578 100644 --- a/gossipsub_peer_lifecycle_test.go +++ b/gossipsub_peer_lifecycle_test.go @@ -2,11 +2,14 @@ package pubsub import ( "context" + "errors" "sync" + "sync/atomic" "testing" "time" pb "github.com/libp2p/go-libp2p-pubsub/pb" + "github.com/libp2p/go-libp2p/core/host" "github.com/libp2p/go-libp2p/core/network" "github.com/libp2p/go-libp2p/core/peer" "github.com/libp2p/go-libp2p/core/protocol" @@ -14,6 +17,201 @@ import ( "google.golang.org/protobuf/proto" ) +type outboundOpenFailureHost struct { + host.Host + + started chan struct{} + release chan struct{} + err error + once sync.Once + calls atomic.Int32 +} + +func (h *outboundOpenFailureHost) NewStream(ctx context.Context, p peer.ID, pids ...protocol.ID) (network.Stream, error) { + h.calls.Add(1) + h.once.Do(func() { close(h.started) }) + select { + case <-h.release: + return nil, h.err + case <-ctx.Done(): + return nil, ctx.Err() + } +} + +type outboundCloseCountingTracer struct { + mockRawTracer + closed atomic.Int32 +} + +func (t *outboundCloseCountingTracer) OnClosedOutboundStream(peer.ID) { + t.closed.Add(1) +} + +func waitForLifecycleCondition(t *testing.T, ps *PubSub, desc string, condition func() bool) { + t.Helper() + deadline := time.Now().Add(5 * time.Second) + for { + result := make(chan bool, 1) + ps.eval <- func() { result <- condition() } + if <-result { + return + } + if time.Now().After(deadline) { + t.Fatalf("timed out waiting for %s", desc) + } + time.Sleep(10 * time.Millisecond) + } +} + +func writeLifecycleSubscription(t *testing.T, stream network.Stream, topic string) { + t.Helper() + rpc := &pb.RPC{Subscriptions: []*pb.RPC_SubOpts{{ + Topicid: proto.String(topic), + Subscribe: proto.Bool(true), + }}} + b, err := proto.Marshal(rpc) + if err != nil { + t.Fatal(err) + } + if err := msgio.NewVarintWriter(stream).WriteMsg(b); err != nil { + t.Fatal(err) + } +} + +func newOutboundOpenFailureTest(t *testing.T, ctx context.Context) (*PubSub, *outboundOpenFailureHost, host.Host, *outboundCloseCountingTracer) { + t.Helper() + hosts := getDefaultHosts(t, 2) + local := &outboundOpenFailureHost{ + Host: hosts[0], + started: make(chan struct{}), + release: make(chan struct{}), + err: errors.New("controlled outbound open failure"), + } + tracer := &outboundCloseCountingTracer{} + ps := getGossipsub(ctx, local, WithRawTracer(tracer), WithMessageSignaturePolicy(StrictNoSign)) + hosts[1].SetStreamHandler(GossipSubID_v13, func(s network.Stream) { _ = s.Reset() }) + connect(t, hosts[1], local) + select { + case <-local.started: + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for initial outbound open") + } + return ps, local, hosts[1], tracer +} + +func TestInitialOutboundOpenFailureRetiresInboundPeer(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + ps, local, remote, tracer := newOutboundOpenFailureTest(t, ctx) + const topicID = "outbound-open-failure" + topic, err := ps.Join(topicID) + if err != nil { + t.Fatal(err) + } + events, err := topic.EventHandler() + if err != nil { + t.Fatal(err) + } + defer events.Cancel() + + inbound, err := remote.NewStream(ctx, local.ID(), GossipSubID_v13) + if err != nil { + t.Fatal(err) + } + defer inbound.Close() + writeLifecycleSubscription(t, inbound, topicID) + + eventCtx, eventCancel := context.WithTimeout(ctx, 5*time.Second) + join, err := events.NextPeerEvent(eventCtx) + eventCancel() + if err != nil { + t.Fatal(err) + } + if join.Type != PeerJoin || join.Peer != remote.ID() { + t.Fatalf("first event = %+v, want PeerJoin for %s", join, remote.ID()) + } + + reset := make(chan error, 1) + go func() { + var b [1]byte + _, err := inbound.Read(b[:]) + reset <- err + }() + close(local.release) + + waitForLifecycleCondition(t, ps, "peer retirement and topic cleanup", func() bool { + _, inRegistry := ps.peerComm.Lookup(remote.ID()) + _, inTopic := ps.topics[topicID][remote.ID()] + return !inRegistry && !inTopic + }) + select { + case err := <-reset: + if !errors.Is(err, network.ErrReset) { + t.Fatalf("inbound read error = %v, want stream reset", err) + } + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for inbound stream reset") + } + + eventCtx, eventCancel = context.WithTimeout(ctx, 5*time.Second) + leave, err := events.NextPeerEvent(eventCtx) + eventCancel() + if err != nil { + t.Fatal(err) + } + if leave.Type != PeerLeave || leave.Peer != join.Peer { + t.Fatalf("second event = %+v, want matching PeerLeave for %s", leave, join.Peer) + } + eventCtx, eventCancel = context.WithTimeout(ctx, 100*time.Millisecond) + _, err = events.NextPeerEvent(eventCtx) + eventCancel() + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("unexpected additional peer event: %v", err) + } + + if got := tracer.closed.Load(); got != 0 { + t.Fatalf("OnClosedOutboundStream calls = %d, want 0", got) + } + if got := local.calls.Load(); got != 1 { + t.Fatalf("outbound open attempts = %d, want 1", got) + } + if connected := local.Network().Connectedness(remote.ID()); connected != network.Connected { + t.Fatalf("host connectedness = %s, want connected", connected) + } +} + +func TestInitialOutboundOpenFailureWithoutInboundSubscriptionEmitsNoPeerLeave(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + ps, local, remote, _ := newOutboundOpenFailureTest(t, ctx) + topic, err := ps.Join("outbound-open-failure-no-subscription") + if err != nil { + t.Fatal(err) + } + events, err := topic.EventHandler() + if err != nil { + t.Fatal(err) + } + defer events.Cancel() + + close(local.release) + waitForLifecycleCondition(t, ps, "peer retirement", func() bool { + _, inRegistry := ps.peerComm.Lookup(remote.ID()) + return !inRegistry + }) + if got := local.calls.Load(); got != 1 { + t.Fatalf("outbound open attempts = %d, want 1", got) + } + + eventCtx, eventCancel := context.WithTimeout(ctx, 100*time.Millisecond) + defer eventCancel() + if event, err := events.NextPeerEvent(eventCtx); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("unexpected peer event %+v: %v", event, err) + } +} + // lifecycleSkeletonGossipsub is a minimal gossipsub peer for testing peer // lifecycle edge cases. It manages two streams: // @@ -192,7 +390,7 @@ func TestNoLeakFromDisconnectedPeer(t *testing.T) { } } waitForCond("peer added", func() bool { - _, inPeers := ps.peers[remoteHost.ID()] + _, inPeers := ps.peerComm.Lookup(remoteHost.ID()) _, inRouter := gs.peers[remoteHost.ID()] return inPeers && inRouter }) @@ -204,9 +402,9 @@ func TestNoLeakFromDisconnectedPeer(t *testing.T) { // Wait for the local to fully remove the peer. handleDeadPeers removes // it, tries to reconnect (fails because handler is gone), and newPeerError - // cleans up ps.peers. + // cleans up the peer registry. waitForCond("peer removed", func() bool { - _, inPeers := ps.peers[remoteHost.ID()] + _, inPeers := ps.peerComm.Lookup(remoteHost.ID()) _, inGossipsub := gs.peers[remoteHost.ID()] return !inPeers && !inGossipsub }) @@ -233,15 +431,15 @@ func TestNoLeakFromDisconnectedPeer(t *testing.T) { time.Sleep(time.Second) // --- Step 4: Observe leaked state. --- - // The peer is not in ps.peers (it was removed and the reconnect failed), + // The peer is not in the registry (it was removed and the reconnect failed), // but the stale RPC re-added it to ps.topics. Since the peer is not in - // ps.peers, handleDeadPeers will never clean it up — it checks ps.peers + // the registry, handleDeadPeers will never clean it up — it checks registry membership // first and skips unknown peers. This is permanently leaked state. checkDone := make(chan struct{}) ps.eval <- func() { defer close(checkDone) - _, inPeers := ps.peers[remoteHost.ID()] + _, inPeers := ps.peerComm.Lookup(remoteHost.ID()) _, inTopics := ps.topics[topic][remoteHost.ID()] if inPeers || inTopics { diff --git a/internal/peercomm/peercomm.go b/internal/peercomm/peercomm.go new file mode 100644 index 00000000..032b1c06 --- /dev/null +++ b/internal/peercomm/peercomm.go @@ -0,0 +1,693 @@ +// Package peercomm owns pubsub control transport communication for each peer. +package peercomm + +import ( + "context" + "encoding/binary" + "errors" + "io" + "iter" + "sync" + "time" + + pool "github.com/libp2p/go-buffer-pool" + "github.com/libp2p/go-libp2p/core/network" + "github.com/libp2p/go-libp2p/core/peer" + "github.com/libp2p/go-libp2p/core/protocol" + "github.com/libp2p/go-msgio" + "github.com/multiformats/go-varint" + "google.golang.org/protobuf/proto" + + pb "github.com/libp2p/go-libp2p-pubsub/pb" +) + +const WriteTimeout = 30 * time.Second + +var ( + ErrQueueFull = errors.New("peercomm: outbound queue full") + ErrQueueClosed = errors.New("peercomm: outbound queue closed") + ErrActorRetired = errors.New("peercomm: actor retired") + ErrNoProtocols = errors.New("peercomm: no protocols available") +) + +// StreamOpener is the subset of host.Host needed by peer communication. +type StreamOpener interface { + NewStream(context.Context, peer.ID, ...protocol.ID) (network.Stream, error) +} + +// Transport identifies the control or topic stream carrying an inbound RPC. +type Transport uint8 + +const ( + TransportControl Transport = iota + TransportTopic +) + +// Hooks emits transport events without calling back into root PubSub for policy decisions. +type Hooks struct { + InboundOpened func(*Actor, network.Stream) + InboundRPC func(*Actor, network.Stream, Transport, *pb.RPC) + InboundClosed func(*Actor, network.Stream) + OutboundReady func(*Actor, network.Stream) + OutboundSent func(*Actor, network.Stream, *pb.RPC) + OutboundSendFailed func(*Actor, network.Stream, *pb.RPC, error) + OutboundOpenFailed func(*Actor, error) + OutboundDead func(*Actor, network.Stream, error) +} + +// Config configures all actors in a Registry. +type Config struct { + Host StreamOpener + Hooks Hooks + QueueSize int + MaxMessageSize int + MaxControlMessageSize int +} + +// Registry owns at most one Actor for each peer. +type Registry struct { + ctx context.Context + cancel context.CancelFunc + config Config + + mu sync.Mutex + actors map[peer.ID]*Actor + stopped bool +} + +func NewRegistry(ctx context.Context, config Config) (*Registry, error) { + if ctx == nil { + return nil, errors.New("peercomm: nil context") + } + if config.Host == nil { + return nil, errors.New("peercomm: nil host") + } + if config.QueueSize <= 0 || config.MaxMessageSize <= 0 || config.MaxControlMessageSize <= 0 { + return nil, errors.New("peercomm: queue and message limits must be positive") + } + ctx, cancel := context.WithCancel(ctx) + return &Registry{ctx: ctx, cancel: cancel, config: config, actors: make(map[peer.ID]*Actor)}, nil +} + +// GetOrCreate returns the current actor or atomically creates its successor. +func (r *Registry) GetOrCreate(p peer.ID) *Actor { + r.mu.Lock() + defer r.mu.Unlock() + if a := r.actors[p]; a != nil { + return a + } + a := newActor(r, p) + if !r.stopped { + r.actors[p] = a + } + return a +} + +// Lookup returns the current actor without creating it. +func (r *Registry) Lookup(p peer.ID) (*Actor, bool) { + r.mu.Lock() + defer r.mu.Unlock() + a, ok := r.actors[p] + return a, ok +} + +// IsCurrent reports whether actor is the authoritative generation for its peer. +func (r *Registry) IsCurrent(actor *Actor) bool { + if actor == nil { + return false + } + r.mu.Lock() + defer r.mu.Unlock() + return r.actors[actor.peer] == actor +} + +// All returns an iterator over a stable snapshot of the current actors. +func (r *Registry) All() iter.Seq2[peer.ID, *Actor] { + r.mu.Lock() + actors := make(map[peer.ID]*Actor, len(r.actors)) + for id, actor := range r.actors { + actors[id] = actor + } + r.mu.Unlock() + return func(yield func(peer.ID, *Actor) bool) { + for id, actor := range actors { + if !yield(id, actor) { + return + } + } + } +} + +// Retire removes and stops p's current actor. A later Actor call creates a fresh one. +func (r *Registry) Retire(p peer.ID) { + r.mu.Lock() + a := r.actors[p] + if a != nil { + delete(r.actors, p) + } + r.mu.Unlock() + if a != nil { + a.Retire() + } +} + +// Retirement describes logical transport teardown claimed while retiring an actor. +type Retirement struct { + InboundProtocol protocol.ID + HadInbound bool +} + +// RetireActor removes and stops actor only if it is still current. It claims +// logical inbound teardown before the actor can be replaced. +func (r *Registry) RetireActor(actor *Actor) (Retirement, bool) { + if actor == nil { + return Retirement{}, false + } + r.mu.Lock() + if r.actors[actor.peer] != actor { + r.mu.Unlock() + return Retirement{}, false + } + retirement := actor.claimCurrentInboundClose() + delete(r.actors, actor.peer) + r.mu.Unlock() + actor.Retire() + return retirement, true +} + +// Stop retires every actor and prevents ongoing communication. +func (r *Registry) Stop() { + r.cancel() + r.mu.Lock() + r.stopped = true + actors := r.actors + r.actors = make(map[peer.ID]*Actor) + r.mu.Unlock() + for _, a := range actors { + a.Retire() + } +} + +type commandKind uint8 + +const ( + commandStart commandKind = iota + commandOpenResult + commandOutboundClosed + commandActivate +) + +type command struct { + kind commandKind + generation uint64 + backoff time.Duration + stream network.Stream + err error + hello *pb.RPC + protocols []protocol.ID +} + +// OpenRequest is an immutable outbound-open command prepared by PubSub. +type OpenRequest struct { + Protocols []protocol.ID + Backoff time.Duration +} + +// Actor owns a peer's outbound queue, stream generations, reconnect timer, and +// inbound replacement state. +type Actor struct { + registry *Registry + peer peer.ID + ctx context.Context + cancel context.CancelFunc + queue *rpcQueue + commands chan command + done chan struct{} + + inboundMu sync.Mutex + currentInbound *inboundRun + notifiedInbound *inboundRun + retire sync.Once +} + +type inboundRun struct { + stream network.Stream + done chan struct{} + opened bool + closeClaimed bool +} + +func newActor(r *Registry, p peer.ID) *Actor { + ctx, cancel := context.WithCancel(r.ctx) + a := &Actor{ + registry: r, peer: p, ctx: ctx, cancel: cancel, + queue: newRPCQueue(r.config.QueueSize), commands: make(chan command, 16), done: make(chan struct{}), + } + go a.run() + return a +} + +func (a *Actor) Peer() peer.ID { return a.peer } +func (a *Actor) Done() <-chan struct{} { return a.done } + +// Start submits a copied outbound-open request. Starting while a stream is live is a no-op. +func (a *Actor) Start(request OpenRequest) error { + if len(request.Protocols) == 0 { + return ErrNoProtocols + } + protocols := append([]protocol.ID(nil), request.Protocols...) + return a.command(command{kind: commandStart, backoff: request.Backoff, protocols: protocols}) +} + +// Activate accepts a ready stream and supplies the prepared immutable hello. +func (a *Actor) Activate(s network.Stream, hello *pb.RPC) error { + if hello != nil { + hello = proto.Clone(hello).(*pb.RPC) + } + return a.command(command{kind: commandActivate, stream: s, hello: hello}) +} + +// Send enqueues an RPC without blocking. Urgent RPCs are drained before normal +// RPCs while preserving FIFO order within each class. A successful Send snapshots +// rpc, so callers may inspect or mutate the original afterward, and actors do not +// share protobuf runtime state. If Send fails, it imposes no asynchronous ownership +// or lifetime obligation on rpc. +func (a *Actor) Send(rpc *pb.RPC, urgent bool) error { + if rpc == nil { + return errors.New("peercomm: nil rpc") + } + select { + case <-a.ctx.Done(): + return ErrActorRetired + default: + } + err := a.queue.push(rpc, urgent) + if errors.Is(err, ErrQueueClosed) { + return ErrActorRetired + } + return err +} + +// Retire permanently cancels this actor and closes its queue. +func (a *Actor) Retire() { + a.retire.Do(func() { + a.cancel() + a.queue.close() + }) +} + +func (a *Actor) command(c command) error { + select { + case <-a.ctx.Done(): + return ErrActorRetired + default: + } + select { + case a.commands <- c: + return nil + case <-a.ctx.Done(): + return ErrActorRetired + case <-a.done: + return ErrActorRetired + } +} + +// HandleInbound authenticates the remote peer, atomically replaces an older +// inbound stream, emits close before open for duplicates, and reads RPCs until +// terminal input. EOF closes politely; every other terminal condition resets. +func (a *Actor) HandleInbound(s network.Stream) { + if s == nil || s.Conn().RemotePeer() != a.peer { + if s != nil { + _ = s.Reset() + } + return + } + + run := &inboundRun{stream: s, done: make(chan struct{})} + a.inboundMu.Lock() + previous := a.currentInbound + a.currentInbound = run + a.inboundMu.Unlock() + + if previous != nil { + _ = previous.stream.Reset() + select { + case <-previous.done: + case <-a.ctx.Done(): + _ = s.Reset() + a.finishInbound(run) + return + } + } + + a.inboundMu.Lock() + if a.currentInbound != run || a.ctx.Err() != nil { + a.inboundMu.Unlock() + _ = s.Reset() + a.finishInbound(run) + return + } + run.opened = true + a.notifiedInbound = run + a.inboundMu.Unlock() + + if h := a.registry.config.Hooks.InboundOpened; h != nil { + h(a, s) + } + a.readInbound(run) +} + +func (a *Actor) readInbound(run *inboundRun) { + s := run.stream + defer a.finishInbound(run) + + r := msgio.NewVarintReaderSize(s, a.registry.config.MaxMessageSize) + for { + _, _ = r.NextMsgLen() + b, err := r.ReadMsg() + if err != nil { + r.ReleaseMsg(b) + if errors.Is(err, io.EOF) { + _ = s.Close() + } else { + _ = s.Reset() + } + return + } + if len(b) == 0 { + r.ReleaseMsg(b) + continue + } + if err = pb.ValidateRawRPCControlMessageSize(b, a.registry.config.MaxControlMessageSize); err != nil { + r.ReleaseMsg(b) + _ = s.Reset() + return + } + rpc := new(pb.RPC) + err = proto.Unmarshal(b, rpc) + r.ReleaseMsg(b) + if err != nil { + _ = s.Reset() + return + } + if h := a.registry.config.Hooks.InboundRPC; h != nil { + h(a, s, TransportControl, rpc) + } + select { + case <-a.ctx.Done(): + _ = s.Reset() + return + default: + } + } +} + +func (a *Actor) finishInbound(run *inboundRun) { + a.inboundMu.Lock() + if a.currentInbound == run { + a.currentInbound = nil + } + _, notify := a.claimInboundCloseLocked(run) + a.inboundMu.Unlock() + if notify { + if h := a.registry.config.Hooks.InboundClosed; h != nil { + h(a, run.stream) + } + } + close(run.done) +} + +func (a *Actor) claimCurrentInboundClose() Retirement { + a.inboundMu.Lock() + defer a.inboundMu.Unlock() + proto, ok := a.claimInboundCloseLocked(a.notifiedInbound) + return Retirement{InboundProtocol: proto, HadInbound: ok} +} + +func (a *Actor) claimInboundCloseLocked(run *inboundRun) (protocol.ID, bool) { + if run == nil || !run.opened || run.closeClaimed { + return "", false + } + run.closeClaimed = true + if a.notifiedInbound == run { + a.notifiedInbound = nil + } + return run.stream.Protocol(), true +} + +func (a *Actor) run() { + defer close(a.done) + var generation uint64 + var current network.Stream + var pending network.Stream + var streamCancel context.CancelFunc + var timerC <-chan time.Time + opening := false + var protocols []protocol.ID + + terminate := func(notify bool, err error) { + generation++ + opening = false + timerC = nil + if streamCancel != nil { + streamCancel() + streamCancel = nil + } + if pending != nil { + _ = pending.Reset() + pending = nil + } + if current != nil { + s := current + current = nil + _ = s.Reset() + if notify && a.registry.config.Hooks.OutboundDead != nil { + a.registry.config.Hooks.OutboundDead(a, s, err) + } + } + } + schedule := func(backoff time.Duration) { + if backoff < 0 { + backoff = 0 + } + timerC = time.After(backoff) + } + + for { + select { + case <-a.ctx.Done(): + terminate(false, a.ctx.Err()) + a.closeInbound() + return + case <-timerC: + timerC = nil + if opening || current != nil { + continue + } + generation++ + gen := generation + opening = true + go a.open(gen, protocols) + case c := <-a.commands: + switch c.kind { + case commandStart: + if current == nil && pending == nil && !opening { + protocols = c.protocols + schedule(c.backoff) + } + case commandOpenResult: + if c.generation != generation || !opening { + if c.stream != nil { + _ = c.stream.Reset() + } + continue + } + opening = false + if c.err != nil { + if h := a.registry.config.Hooks.OutboundOpenFailed; h != nil { + h(a, c.err) + } + continue + } + pending = c.stream + if h := a.registry.config.Hooks.OutboundReady; h != nil { + h(a, pending) + } + case commandActivate: + if pending == nil || c.stream != pending { + continue + } + current = pending + pending = nil + streamCtx, cancel := context.WithCancel(a.ctx) + streamCancel = cancel + go a.writeLoop(streamCtx, generation, current, c.hello) + go a.watchDeath(streamCtx, generation, current) + case commandOutboundClosed: + if c.generation != generation || current != c.stream { + continue + } + terminate(true, c.err) + } + } + } +} + +func (a *Actor) open(generation uint64, protocols []protocol.ID) { + s, err := a.registry.config.Host.NewStream(a.ctx, a.peer, protocols...) + if a.ctx.Err() != nil { + if s != nil { + _ = s.Reset() + } + return + } + select { + case a.commands <- command{kind: commandOpenResult, generation: generation, stream: s, err: err}: + case <-a.ctx.Done(): + if s != nil { + _ = s.Reset() + } + } +} + +func (a *Actor) writeLoop(ctx context.Context, generation uint64, s network.Stream, hello *pb.RPC) { + if hello != nil && proto.Size(hello) > 0 { + if err := a.writeRPC(s, hello); err != nil { + _ = a.command(command{kind: commandOutboundClosed, generation: generation, stream: s, err: err}) + return + } + } + for { + rpc, err := a.queue.pop(ctx) + if err != nil { + return + } + if err = a.writeRPC(s, rpc); err != nil { + _ = a.command(command{kind: commandOutboundClosed, generation: generation, stream: s, err: err}) + return + } + } +} + +func (a *Actor) writeRPC(s network.Stream, rpc *pb.RPC) error { + err := writeRPC(s, rpc) + if err != nil { + if h := a.registry.config.Hooks.OutboundSendFailed; h != nil { + h(a, s, rpc, err) + } + return err + } + if h := a.registry.config.Hooks.OutboundSent; h != nil { + h(a, s, rpc) + } + return nil +} + +func (a *Actor) watchDeath(ctx context.Context, generation uint64, s network.Stream) { + one := []byte{0} + _, err := s.Read(one) + select { + case <-ctx.Done(): + return + default: + } + if err == nil { + err = errors.New("peercomm: unexpected data on outbound stream") + } + _ = a.command(command{kind: commandOutboundClosed, generation: generation, stream: s, err: err}) +} + +func (a *Actor) closeInbound() { + a.inboundMu.Lock() + if a.currentInbound != nil { + _ = a.currentInbound.stream.Reset() + } + a.inboundMu.Unlock() +} + +func writeRPC(s network.Stream, rpc *pb.RPC) error { + size := uint64(proto.Size(rpc)) + buf := pool.Get(varint.UvarintSize(size) + int(size)) + defer pool.Put(buf) + n := binary.PutUvarint(buf, size) + out, err := proto.MarshalOptions{}.MarshalAppend(buf[:n], rpc) + if err != nil { + return err + } + if err = s.SetWriteDeadline(time.Now().Add(WriteTimeout)); err != nil { + return err + } + written, err := s.Write(out) + if err == nil && written != len(out) { + return io.ErrShortWrite + } + return err +} + +type rpcQueue struct { + mu sync.Mutex + available *sync.Cond + urgent []*pb.RPC + normal []*pb.RPC + capacity int + closed bool +} + +func newRPCQueue(capacity int) *rpcQueue { + q := &rpcQueue{capacity: capacity} + q.available = sync.NewCond(&q.mu) + return q +} + +func (q *rpcQueue) push(rpc *pb.RPC, urgent bool) error { + q.mu.Lock() + defer q.mu.Unlock() + if q.closed { + return ErrQueueClosed + } + if len(q.urgent)+len(q.normal) >= q.capacity { + return ErrQueueFull + } + rpc = proto.Clone(rpc).(*pb.RPC) + if urgent { + q.urgent = append(q.urgent, rpc) + } else { + q.normal = append(q.normal, rpc) + } + q.available.Signal() + return nil +} + +func (q *rpcQueue) pop(ctx context.Context) (*pb.RPC, error) { + q.mu.Lock() + defer q.mu.Unlock() + stop := context.AfterFunc(ctx, q.available.Broadcast) + defer stop() + for len(q.urgent)+len(q.normal) == 0 && !q.closed && ctx.Err() == nil { + q.available.Wait() + } + if ctx.Err() != nil { + return nil, ctx.Err() + } + if q.closed { + return nil, ErrQueueClosed + } + var rpc *pb.RPC + if len(q.urgent) > 0 { + rpc = q.urgent[0] + q.urgent[0] = nil + q.urgent = q.urgent[1:] + } else { + rpc = q.normal[0] + q.normal[0] = nil + q.normal = q.normal[1:] + } + return rpc, nil +} + +func (q *rpcQueue) close() { + q.mu.Lock() + q.closed = true + q.available.Broadcast() + q.mu.Unlock() +} diff --git a/internal/peercomm/peercomm_test.go b/internal/peercomm/peercomm_test.go new file mode 100644 index 00000000..c0e7beef --- /dev/null +++ b/internal/peercomm/peercomm_test.go @@ -0,0 +1,844 @@ +package peercomm + +import ( + "context" + "encoding/binary" + "errors" + "io" + "sync" + "testing" + "time" + + "github.com/libp2p/go-libp2p/core/network" + "github.com/libp2p/go-libp2p/core/peer" + "github.com/libp2p/go-libp2p/core/protocol" + "google.golang.org/protobuf/proto" + + pb "github.com/libp2p/go-libp2p-pubsub/pb" +) + +func TestQueuePriorityFullAndClose(t *testing.T) { + q := newRPCQueue(3) + n1 := &pb.RPC{Publish: []*pb.Message{{Data: []byte("normal-1")}}} + u1 := &pb.RPC{Publish: []*pb.Message{{Data: []byte("urgent-1")}}} + u2 := &pb.RPC{Publish: []*pb.Message{{Data: []byte("urgent-2")}}} + if err := q.push(n1, false); err != nil { + t.Fatal(err) + } + if err := q.push(u1, true); err != nil { + t.Fatal(err) + } + if err := q.push(u2, true); err != nil { + t.Fatal(err) + } + if err := q.push(&pb.RPC{}, false); !errors.Is(err, ErrQueueFull) { + t.Fatalf("full push: %v", err) + } + + ctx := context.Background() + for i, want := range []*pb.RPC{u1, u2, n1} { + got, err := q.pop(ctx) + if err != nil { + t.Fatalf("pop %d: %v", i, err) + } + if !proto.Equal(got, want) { + t.Fatalf("pop %d returned wrong priority: got %v, want %v", i, got, want) + } + } + q.close() + if _, err := q.pop(ctx); !errors.Is(err, ErrQueueClosed) { + t.Fatalf("closed pop: %v", err) + } + if err := q.push(&pb.RPC{}, false); !errors.Is(err, ErrQueueClosed) { + t.Fatalf("closed push: %v", err) + } +} + +func TestQueuePopCancellation(t *testing.T) { + q := newRPCQueue(1) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if _, err := q.pop(ctx); !errors.Is(err, context.Canceled) { + t.Fatalf("pop: %v", err) + } +} + +func TestActorSendSnapshotsRPC(t *testing.T) { + r, err := NewRegistry(context.Background(), testConfig(&failingHost{}, Hooks{})) + if err != nil { + t.Fatal(err) + } + defer r.Stop() + + a := r.GetOrCreate(peer.ID("peer")) + rpc := &pb.RPC{Publish: []*pb.Message{{Data: []byte("original")}}} + if err := a.Send(rpc, false); err != nil { + t.Fatal(err) + } + got, err := a.queue.pop(context.Background()) + if err != nil { + t.Fatal(err) + } + if got == rpc { + t.Fatal("Send enqueued the supplied RPC pointer") + } + if !proto.Equal(got, rpc) { + t.Fatalf("queued RPC differs from original: got %v, want %v", got, rpc) + } + + rpc.Publish[0].Data[0] = 'O' + if string(got.Publish[0].Data) != "original" { + t.Fatalf("queued nested data changed with original: %q", got.Publish[0].Data) + } +} + +type failingHost struct { + mu sync.Mutex + calls int + err error +} + +func (h *failingHost) NewStream(context.Context, peer.ID, ...protocol.ID) (network.Stream, error) { + h.mu.Lock() + h.calls++ + h.mu.Unlock() + return nil, h.err +} + +func testConfig(h StreamOpener, hooks Hooks) Config { + return Config{Host: h, Hooks: hooks, QueueSize: 4, MaxMessageSize: 1 << 20, MaxControlMessageSize: 1 << 20} +} + +func TestRegistryReturnsOneActorAndRetires(t *testing.T) { + h := &failingHost{err: errors.New("open failed")} + r, err := NewRegistry(context.Background(), testConfig(h, Hooks{})) + if err != nil { + t.Fatal(err) + } + p := peer.ID("peer") + a := r.GetOrCreate(p) + if got := r.GetOrCreate(p); got != a { + t.Fatal("registry created a duplicate actor") + } + r.Retire(p) + if err := a.Send(&pb.RPC{}, false); !errors.Is(err, ErrActorRetired) { + t.Fatalf("retired send: %v", err) + } + if got := r.GetOrCreate(p); got == a { + t.Fatal("registry reused retired actor") + } + r.Stop() +} + +func TestRegistryRejectsStaleActorRetirement(t *testing.T) { + r, err := NewRegistry(context.Background(), testConfig(&failingHost{}, Hooks{})) + if err != nil { + t.Fatal(err) + } + defer r.Stop() + p := peer.ID("peer") + stale := r.GetOrCreate(p) + if _, ok := r.RetireActor(stale); !ok { + t.Fatal("current actor was not retired") + } + current := r.GetOrCreate(p) + if _, ok := r.RetireActor(stale); ok { + t.Fatal("stale actor retired current generation") + } + if !r.IsCurrent(current) { + t.Fatal("replacement actor is not current") + } +} + +func TestStartCoalescesWhileOpening(t *testing.T) { + ctxSeen := make(chan struct{}) + h := &blockingHost{started: ctxSeen} + r, err := NewRegistry(context.Background(), testConfig(h, Hooks{})) + if err != nil { + t.Fatal(err) + } + a := r.GetOrCreate(peer.ID("peer")) + if err := a.Start(testOpen(0)); err != nil { + t.Fatal(err) + } + select { + case <-ctxSeen: + case <-time.After(time.Second): + t.Fatal("open did not start") + } + for i := 0; i < 10; i++ { + if err := a.Start(testOpen(0)); err != nil { + t.Fatal(err) + } + } + time.Sleep(20 * time.Millisecond) + if got := h.count(); got != 1 { + t.Fatalf("open calls = %d, want 1", got) + } + r.Stop() +} + +type blockingHost struct { + mu sync.Mutex + calls int + started chan struct{} + once sync.Once +} + +func (h *blockingHost) NewStream(ctx context.Context, _ peer.ID, _ ...protocol.ID) (network.Stream, error) { + h.mu.Lock() + h.calls++ + h.mu.Unlock() + h.once.Do(func() { close(h.started) }) + <-ctx.Done() + return nil, ctx.Err() +} +func (h *blockingHost) count() int { h.mu.Lock(); defer h.mu.Unlock(); return h.calls } + +type streamRead struct { + data []byte + err error +} + +type testConn struct { + network.Conn + remote peer.ID +} + +func (c *testConn) RemotePeer() peer.ID { return c.remote } + +type testStream struct { + network.Stream + conn network.Conn + + reads chan streamRead + writes chan []byte + + mu sync.Mutex + pending []byte + terminalErr error + writeErr error + readStarted chan struct{} + writeStarted chan struct{} + readGate <-chan struct{} + writeGate <-chan struct{} + readOnce sync.Once + writeOnce sync.Once + resetCount int + closeCount int + reset chan struct{} + resetOnce sync.Once + protocol protocol.ID +} + +func newTestStream(p peer.ID) *testStream { + return &testStream{ + conn: &testConn{remote: p}, + reads: make(chan streamRead, 8), + writes: make(chan []byte, 8), + reset: make(chan struct{}), + } +} + +func (s *testStream) Conn() network.Conn { return s.conn } +func (s *testStream) Protocol() protocol.ID { return s.protocol } +func (s *testStream) SetProtocol(id protocol.ID) error { s.protocol = id; return nil } + +func (s *testStream) Read(p []byte) (int, error) { + if s.readStarted != nil { + s.readOnce.Do(func() { close(s.readStarted) }) + } + if s.readGate != nil { + select { + case <-s.readGate: + case <-s.reset: + return 0, network.ErrReset + } + } + for { + s.mu.Lock() + if len(s.pending) > 0 { + n := copy(p, s.pending) + s.pending = s.pending[n:] + s.mu.Unlock() + return n, nil + } + if s.terminalErr != nil { + err := s.terminalErr + s.mu.Unlock() + return 0, err + } + s.mu.Unlock() + select { + case r := <-s.reads: + if len(r.data) == 0 { + s.mu.Lock() + s.terminalErr = r.err + s.mu.Unlock() + return 0, r.err + } + s.mu.Lock() + s.pending = append(s.pending, r.data...) + s.mu.Unlock() + case <-s.reset: + return 0, network.ErrReset + } + } +} + +func (s *testStream) Write(p []byte) (int, error) { + if s.writeStarted != nil { + s.writeOnce.Do(func() { close(s.writeStarted) }) + } + if s.writeGate != nil { + select { + case <-s.writeGate: + case <-s.reset: + return 0, network.ErrReset + } + } + s.mu.Lock() + err := s.writeErr + s.mu.Unlock() + if err != nil { + return 0, err + } + b := append([]byte(nil), p...) + s.writes <- b + return len(p), nil +} + +func (s *testStream) Close() error { + s.mu.Lock() + s.closeCount++ + s.mu.Unlock() + s.resetOnce.Do(func() { close(s.reset) }) + return nil +} + +func (s *testStream) Reset() error { + s.mu.Lock() + s.resetCount++ + s.mu.Unlock() + s.resetOnce.Do(func() { close(s.reset) }) + return nil +} + +func (s *testStream) SetWriteDeadline(time.Time) error { return nil } + +func (s *testStream) counts() (closed, reset int) { + s.mu.Lock() + defer s.mu.Unlock() + return s.closeCount, s.resetCount +} + +func frameProto(t *testing.T, rpc proto.Message) []byte { + t.Helper() + payload, err := proto.Marshal(rpc) + if err != nil { + t.Fatal(err) + } + buf := make([]byte, binary.MaxVarintLen64+len(payload)) + n := binary.PutUvarint(buf, uint64(len(payload))) + copy(buf[n:], payload) + return buf[:n+len(payload)] +} + +func decodeFrameInto(t *testing.T, frame []byte, message proto.Message) { + t.Helper() + size, n := binary.Uvarint(frame) + if n <= 0 || int(size) != len(frame)-n { + t.Fatalf("invalid frame: size=%d prefix=%d bytes=%d", size, n, len(frame)) + } + if err := proto.Unmarshal(frame[n:], message); err != nil { + t.Fatal(err) + } +} + +func decodeFrame(t *testing.T, frame []byte) *pb.RPC { + rpc := new(pb.RPC) + decodeFrameInto(t, frame, rpc) + return rpc +} + +func receive[T any](t *testing.T, ch <-chan T, what string) T { + t.Helper() + select { + case v := <-ch: + return v + case <-time.After(2 * time.Second): + t.Fatalf("timed out waiting for %s", what) + var zero T + return zero + } +} + +func waitDone(t *testing.T, ch <-chan struct{}, what string) { + t.Helper() + receive(t, ch, what) +} + +func TestInboundReplacementCallbackOrdering(t *testing.T) { + p := peer.ID("peer") + events := make(chan string, 4) + oldClosing := make(chan struct{}) + releaseClose := make(chan struct{}) + old := newTestStream(p) + fresh := newTestStream(p) + + r, err := NewRegistry(context.Background(), testConfig(&failingHost{}, Hooks{ + InboundOpened: func(_ *Actor, s network.Stream) { + if s == old { + events <- "open-old" + } else { + events <- "open-new" + } + }, + InboundClosed: func(_ *Actor, s network.Stream) { + if s == old { + close(oldClosing) + <-releaseClose + events <- "close-old" + } else { + events <- "close-new" + } + }, + })) + if err != nil { + t.Fatal(err) + } + a := r.GetOrCreate(p) + oldDone := make(chan struct{}) + go func() { a.HandleInbound(old); close(oldDone) }() + if got := receive(t, events, "old open"); got != "open-old" { + t.Fatalf("first event = %q", got) + } + newDone := make(chan struct{}) + go func() { a.HandleInbound(fresh); close(newDone) }() + waitDone(t, oldClosing, "old close callback") + select { + case got := <-events: + t.Fatalf("replacement opened before old close completed: %q", got) + default: + } + close(releaseClose) + if got := receive(t, events, "old close"); got != "close-old" { + t.Fatalf("second event = %q", got) + } + if got := receive(t, events, "new open"); got != "open-new" { + t.Fatalf("third event = %q", got) + } + waitDone(t, oldDone, "old handler") + + // Completion of the stale old run must not clear the replacement. + a.inboundMu.Lock() + if a.currentInbound == nil || a.currentInbound.stream != fresh { + t.Fatal("old close cleared replacement state") + } + a.inboundMu.Unlock() + r.Stop() + waitDone(t, newDone, "new handler") +} + +func TestInboundCloseClaimedExactlyOnce(t *testing.T) { + p := peer.ID("peer") + const protoID = protocol.ID("/inbound/1") + opened := make(chan struct{}) + closed := make(chan struct{}, 1) + r, err := NewRegistry(context.Background(), testConfig(&failingHost{}, Hooks{ + InboundOpened: func(*Actor, network.Stream) { close(opened) }, + InboundClosed: func(*Actor, network.Stream) { closed <- struct{}{} }, + })) + if err != nil { + t.Fatal(err) + } + defer r.Stop() + a := r.GetOrCreate(p) + s := newTestStream(p) + s.protocol = protoID + done := make(chan struct{}) + go func() { a.HandleInbound(s); close(done) }() + waitDone(t, opened, "inbound open") + s.reads <- streamRead{err: io.EOF} + waitDone(t, closed, "inbound close") + waitDone(t, done, "inbound handler") + retirement, ok := r.RetireActor(a) + if !ok { + t.Fatal("actor was not retired") + } + if retirement.HadInbound { + t.Fatalf("retirement reclaimed closed inbound protocol %q", retirement.InboundProtocol) + } +} + +func TestRetirementClaimsInboundAndSuppressesCloseHook(t *testing.T) { + p := peer.ID("peer") + const protoID = protocol.ID("/inbound/1") + opened := make(chan struct{}) + closed := make(chan struct{}, 1) + r, err := NewRegistry(context.Background(), testConfig(&failingHost{}, Hooks{ + InboundOpened: func(*Actor, network.Stream) { close(opened) }, + InboundClosed: func(*Actor, network.Stream) { closed <- struct{}{} }, + })) + if err != nil { + t.Fatal(err) + } + defer r.Stop() + a := r.GetOrCreate(p) + s := newTestStream(p) + s.protocol = protoID + done := make(chan struct{}) + go func() { a.HandleInbound(s); close(done) }() + waitDone(t, opened, "inbound open") + retirement, ok := r.RetireActor(a) + if !ok { + t.Fatal("actor was not retired") + } + if !retirement.HadInbound || retirement.InboundProtocol != protoID { + t.Fatalf("retirement inbound = %q/%t, want %q/true", retirement.InboundProtocol, retirement.HadInbound, protoID) + } + waitDone(t, done, "retired inbound handler") + select { + case <-closed: + t.Fatal("retired inbound emitted duplicate close hook") + default: + } +} + +func TestInboundValidFrameAuthenticatesPeer(t *testing.T) { + p := peer.ID("authenticated-peer") + rpcs := make(chan *pb.RPC, 1) + actors := make(chan peer.ID, 1) + r, err := NewRegistry(context.Background(), testConfig(&failingHost{}, Hooks{ + InboundRPC: func(a *Actor, _ network.Stream, transport Transport, rpc *pb.RPC) { + if transport != TransportControl { + t.Errorf("transport = %v, want control", transport) + } + actors <- a.Peer() + rpcs <- rpc + }, + })) + if err != nil { + t.Fatal(err) + } + defer r.Stop() + a := r.GetOrCreate(p) + s := newTestStream(p) + want := &pb.RPC{Subscriptions: []*pb.RPC_SubOpts{{Topicid: proto.String("topic"), Subscribe: proto.Bool(true)}}} + s.reads <- streamRead{data: frameProto(t, want)} + s.reads <- streamRead{err: io.EOF} + done := make(chan struct{}) + go func() { a.HandleInbound(s); close(done) }() + got := receive(t, rpcs, "inbound RPC") + if !proto.Equal(got, want) { + t.Fatalf("RPC mismatch: got %v want %v", got, want) + } + if gotPeer := receive(t, actors, "authenticated actor"); gotPeer != p { + t.Fatalf("actor peer = %q, want %q", gotPeer, p) + } + waitDone(t, done, "inbound EOF") + closed, reset := s.counts() + if closed != 1 || reset != 0 { + t.Fatalf("EOF close/reset = %d/%d, want 1/0", closed, reset) + } +} + +func TestInboundFramingFailuresReset(t *testing.T) { + tests := []struct { + name string + config func(Config) Config + input func(*testing.T) []byte + }{ + {name: "malformed protobuf", input: func(*testing.T) []byte { return []byte{1, 0xff} }}, + {name: "oversized envelope", config: func(c Config) Config { c.MaxMessageSize = 2; return c }, input: func(*testing.T) []byte { return []byte{3} }}, + {name: "oversized control", config: func(c Config) Config { c.MaxControlMessageSize = 1; return c }, input: func(t *testing.T) []byte { + return frameProto(t, &pb.RPC{Control: &pb.ControlMessage{Ihave: []*pb.ControlIHave{{TopicID: proto.String("too-large")}}}}) + }}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + p := peer.ID("peer") + cfg := testConfig(&failingHost{}, Hooks{}) + if tt.config != nil { + cfg = tt.config(cfg) + } + r, err := NewRegistry(context.Background(), cfg) + if err != nil { + t.Fatal(err) + } + defer r.Stop() + s := newTestStream(p) + s.reads <- streamRead{data: tt.input(t)} + done := make(chan struct{}) + go func() { r.GetOrCreate(p).HandleInbound(s); close(done) }() + waitDone(t, done, "framing failure") + closed, reset := s.counts() + if closed != 0 || reset != 1 { + t.Fatalf("close/reset = %d/%d, want 0/1", closed, reset) + } + }) + } +} + +type openRequest struct { + ctx context.Context + result chan openResult +} + +type openResult struct { + stream network.Stream + err error +} + +type controlledHost struct{ calls chan openRequest } + +func (h *controlledHost) NewStream(ctx context.Context, _ peer.ID, _ ...protocol.ID) (network.Stream, error) { + req := openRequest{ctx: ctx, result: make(chan openResult, 1)} + h.calls <- req + result := <-req.result + return result.stream, result.err +} + +func outboundConfig(h StreamOpener, hooks Hooks) Config { + if hooks.OutboundReady == nil { + hooks.OutboundReady = func(a *Actor, stream network.Stream) { _ = a.Activate(stream, nil) } + } + return testConfig(h, hooks) +} + +func testOpen(backoff time.Duration) OpenRequest { + return OpenRequest{Protocols: []protocol.ID{"/test/1"}, Backoff: backoff} +} + +func TestHelloWrittenBeforeQueuedRPC(t *testing.T) { + h := &controlledHost{calls: make(chan openRequest, 1)} + hello := &pb.RPC{Subscriptions: []*pb.RPC_SubOpts{{Topicid: proto.String("hello")}}} + queued := &pb.RPC{Subscriptions: []*pb.RPC_SubOpts{{Topicid: proto.String("queued")}}} + r, err := NewRegistry(context.Background(), outboundConfig(h, Hooks{OutboundReady: func(a *Actor, stream network.Stream) { _ = a.Activate(stream, hello) }})) + if err != nil { + t.Fatal(err) + } + defer r.Stop() + a := r.GetOrCreate(peer.ID("peer")) + if err := a.Send(queued, false); err != nil { + t.Fatal(err) + } + if err := a.Start(testOpen(0)); err != nil { + t.Fatal(err) + } + req := receive(t, h.calls, "open request") + s := newTestStream(a.Peer()) + req.result <- openResult{stream: s} + if got := decodeFrame(t, receive(t, s.writes, "hello write")); !proto.Equal(got, hello) { + t.Fatalf("first write = %v", got) + } + if got := decodeFrame(t, receive(t, s.writes, "queued write")); !proto.Equal(got, queued) { + t.Fatalf("second write = %v", got) + } +} + +func TestWriteAndDeathRaceNotifiesTerminalOnce(t *testing.T) { + h := &controlledHost{calls: make(chan openRequest, 1)} + writeErr := errors.New("write failed") + dead := make(chan error, 2) + failed := make(chan error, 2) + r, err := NewRegistry(context.Background(), outboundConfig(h, Hooks{ + OutboundDead: func(_ *Actor, _ network.Stream, err error) { dead <- err }, + OutboundSendFailed: func(_ *Actor, _ network.Stream, _ *pb.RPC, err error) { failed <- err }, + })) + if err != nil { + t.Fatal(err) + } + defer r.Stop() + a := r.GetOrCreate(peer.ID("peer")) + if err := a.Send(&pb.RPC{}, false); err != nil { + t.Fatal(err) + } + if err := a.Start(testOpen(0)); err != nil { + t.Fatal(err) + } + req := receive(t, h.calls, "open request") + s := newTestStream(a.Peer()) + readGate := make(chan struct{}) + writeGate := make(chan struct{}) + s.readStarted = make(chan struct{}) + s.writeStarted = make(chan struct{}) + s.readGate = readGate + s.writeGate = writeGate + s.mu.Lock() + s.writeErr = writeErr + s.terminalErr = errors.New("read failed") + s.mu.Unlock() + req.result <- openResult{stream: s} + waitDone(t, s.readStarted, "death watcher read") + waitDone(t, s.writeStarted, "writer") + close(readGate) + close(writeGate) + if !errors.Is(receive(t, failed, "send failure"), writeErr) { + t.Fatal("wrong send failure") + } + receive(t, dead, "terminal notification") + select { + case err := <-dead: + t.Fatalf("duplicate terminal notification: %v", err) + case <-time.After(50 * time.Millisecond): + } +} + +func TestRetirementDuringOpenAndSendRace(t *testing.T) { + h := &controlledHost{calls: make(chan openRequest, 1)} + r, err := NewRegistry(context.Background(), outboundConfig(h, Hooks{})) + if err != nil { + t.Fatal(err) + } + a := r.GetOrCreate(peer.ID("peer")) + if err := a.Start(testOpen(0)); err != nil { + t.Fatal(err) + } + req := receive(t, h.calls, "open request") + + const senders = 32 + start := make(chan struct{}) + errs := make(chan error, senders) + for i := 0; i < senders; i++ { + go func() { <-start; errs <- a.Send(&pb.RPC{}, false) }() + } + close(start) + r.Retire(a.Peer()) + for i := 0; i < senders; i++ { + err := receive(t, errs, "racing send") + if err != nil && !errors.Is(err, ErrActorRetired) && !errors.Is(err, ErrQueueFull) { + t.Fatalf("send error = %v", err) + } + } + stale := newTestStream(a.Peer()) + req.result <- openResult{stream: stale} + waitDone(t, a.Done(), "actor retirement") + if err := a.Send(&pb.RPC{}, false); !errors.Is(err, ErrActorRetired) { + t.Fatalf("post-retirement send = %v", err) + } + waitDone(t, stale.reset, "retired open result reset") + _, resets := stale.counts() + if resets != 1 { + t.Fatalf("retired open result resets = %d, want 1", resets) + } +} + +func TestStaleActorCallbacksAndRegistryShutdown(t *testing.T) { + h := &failingHost{err: errors.New("open failed")} + r, err := NewRegistry(context.Background(), testConfig(h, Hooks{})) + if err != nil { + t.Fatal(err) + } + p := peer.ID("peer") + old := r.GetOrCreate(p) + r.Retire(p) + fresh := r.GetOrCreate(p) + if fresh == old { + t.Fatal("actor was not replaced") + } + if _, ok := r.RetireActor(old); ok { + t.Fatal("stale actor callback retired replacement") + } + if got, ok := r.Lookup(p); !ok || got != fresh { + t.Fatal("replacement missing after stale callback") + } + + r.Stop() + waitDone(t, fresh.Done(), "registry shutdown") + if _, ok := r.Lookup(p); ok { + t.Fatal("registry retained actor after shutdown") + } + postStop := r.GetOrCreate(p) + waitDone(t, postStop.Done(), "post-stop actor") + if err := postStop.Start(testOpen(0)); !errors.Is(err, ErrActorRetired) { + t.Fatalf("post-stop start = %v", err) + } + if err := postStop.Send(&pb.RPC{}, false); !errors.Is(err, ErrActorRetired) { + t.Fatalf("post-stop send = %v", err) + } +} + +func TestRetirementDuringBackoffReadAndWrite(t *testing.T) { + t.Run("backoff", func(t *testing.T) { + h := &controlledHost{calls: make(chan openRequest, 1)} + r, err := NewRegistry(context.Background(), outboundConfig(h, Hooks{})) + if err != nil { + t.Fatal(err) + } + a := r.GetOrCreate(peer.ID("peer")) + if err := a.Start(testOpen(time.Hour)); err != nil { + t.Fatal(err) + } + a.Retire() + waitDone(t, a.Done(), "backoff retirement") + select { + case <-h.calls: + t.Fatal("open started after retirement during backoff") + default: + } + }) + + t.Run("inbound read", func(t *testing.T) { + p := peer.ID("peer") + closed := make(chan struct{}) + r, err := NewRegistry(context.Background(), testConfig(&failingHost{}, Hooks{InboundClosed: func(*Actor, network.Stream) { close(closed) }})) + if err != nil { + t.Fatal(err) + } + a := r.GetOrCreate(p) + s := newTestStream(p) + started := make(chan struct{}) + s.readStarted = started + done := make(chan struct{}) + go func() { a.HandleInbound(s); close(done) }() + waitDone(t, started, "inbound read") + a.Retire() + waitDone(t, closed, "inbound close callback") + waitDone(t, done, "inbound read retirement") + _, resets := s.counts() + if resets == 0 { + t.Fatal("retirement did not reset inbound read") + } + }) + + t.Run("outbound write", func(t *testing.T) { + h := &controlledHost{calls: make(chan openRequest, 1)} + r, err := NewRegistry(context.Background(), outboundConfig(h, Hooks{})) + if err != nil { + t.Fatal(err) + } + a := r.GetOrCreate(peer.ID("peer")) + if err := a.Send(&pb.RPC{}, false); err != nil { + t.Fatal(err) + } + if err := a.Start(testOpen(0)); err != nil { + t.Fatal(err) + } + req := receive(t, h.calls, "open request") + s := newTestStream(a.Peer()) + s.writeStarted = make(chan struct{}) + s.writeGate = make(chan struct{}) + req.result <- openResult{stream: s} + waitDone(t, s.writeStarted, "outbound write") + a.Retire() + waitDone(t, a.Done(), "outbound write retirement") + waitDone(t, s.reset, "outbound write reset") + }) +} + +func TestInboundRejectsMismatchedPeer(t *testing.T) { + r, err := NewRegistry(context.Background(), testConfig(&failingHost{}, Hooks{})) + if err != nil { + t.Fatal(err) + } + defer r.Stop() + s := newTestStream(peer.ID("impostor")) + done := make(chan struct{}) + go func() { r.GetOrCreate(peer.ID("expected")).HandleInbound(s); close(done) }() + waitDone(t, done, "peer mismatch rejection") + _, resets := s.counts() + if resets != 1 { + t.Fatalf("mismatched peer resets = %d, want 1", resets) + } +} diff --git a/pubsub.go b/pubsub.go index fce5c544..f80a3ab1 100644 --- a/pubsub.go +++ b/pubsub.go @@ -14,6 +14,7 @@ import ( "time" "github.com/libp2p/go-libp2p-pubsub/internal/gologshim" + "github.com/libp2p/go-libp2p-pubsub/internal/peercomm" pb "github.com/libp2p/go-libp2p-pubsub/pb" "github.com/libp2p/go-libp2p-pubsub/timecache" @@ -54,12 +55,6 @@ type peerTopicState struct { supportsPartial bool } -type peerOutgoingStream struct { - network.Stream - FirstMessage chan *RPC - Cancel context.CancelFunc -} - // PubSub is the implementation of the pubsub system. type PubSub struct { // atomic counter for seqnos @@ -128,17 +123,6 @@ type PubSub struct { newPeersMx sync.Mutex newPeersPend map[peer.ID]struct{} - // a notification channel for new outoging peer streams - newPeerStream chan peerOutgoingStream - - // a notification channel for errors opening new peer streams - newPeerError chan peer.ID - - // a notification channel for when our peers die - peerDead chan struct{} - peerDeadPrioLk sync.RWMutex - peerDeadMx sync.Mutex - peerDeadPend map[peer.ID]struct{} // backoff for retrying new connections to dead peers deadPeerBackoff *backoff @@ -173,10 +157,7 @@ type PubSub struct { blacklist Blacklist blacklistPeer chan peer.ID - peers map[peer.ID]*rpcQueue - - inboundStreamsMx sync.Mutex - inboundStreams map[peer.ID]inboundHandler + peerComm *peercomm.Registry seenMessages timecache.TimeCache seenMsgTTL time.Duration @@ -281,20 +262,15 @@ func (m *Message) GetFrom() peer.ID { return peer.ID(m.Message.GetFrom()) } -// inboundHandler tracks an active inbound stream handler. The done channel is -// closed when the handler exits, allowing successive handlers for the same peer -// to serialize their newStream/closedStream notifications. -type inboundHandler struct { - s network.Stream - done chan struct{} -} - type incomingKind uint8 const ( incomingKindRPC = iota incomingKindNewStream incomingKindClosedStream + incomingKindOutboundOpenFailed + incomingKindOutboundDead + incomingKindOutboundReady ) // incomingUnion wraps the different messages the incoming stream handler can @@ -302,15 +278,24 @@ const ( type incomingUnion struct { rpc *RPC // only set when kind == RPC // s is only set when kind == NewStream or kind == ClosedStream - s network.Stream - kind incomingKind + s network.Stream + actor *peercomm.Actor + err error + kind incomingKind } type RPC struct { pb.RPC // unexported on purpose, not sending this over the wire - from peer.ID + from peer.ID + transport peercomm.Transport +} + +func wrapInboundRPC(rpc *pb.RPC, from peer.ID, transport peercomm.Transport) *RPC { + wrapped := &RPC{from: from, transport: transport} + proto.Merge(&wrapped.RPC, rpc) + return wrapped } func (rpc *RPC) From() peer.ID { @@ -636,10 +621,6 @@ func NewPubSub(ctx context.Context, h host.Host, rt PubSubRouter, opts ...Option incoming: make(chan incomingUnion, 32), newPeers: make(chan struct{}, 1), newPeersPend: make(map[peer.ID]struct{}), - newPeerStream: make(chan peerOutgoingStream), - newPeerError: make(chan peer.ID), - peerDead: make(chan struct{}, 1), - peerDeadPend: make(map[peer.ID]struct{}), deadPeerBackoff: newBackoff(ctx, 1000, BackoffCleanupInterval, MaxBackoffAttempts), cancelCh: make(chan *Subscription), getPeers: make(chan *listPeerReq), @@ -658,8 +639,6 @@ func NewPubSub(ctx context.Context, h host.Host, rt PubSubRouter, opts ...Option mySubs: make(map[string]map[*Subscription]struct{}), myRelays: make(map[string]int), topics: make(map[string]map[peer.ID]peerTopicState), - peers: make(map[peer.ID]*rpcQueue), - inboundStreams: make(map[peer.ID]inboundHandler), blacklist: NewMapBlacklist(), blacklistPeer: make(chan peer.ID), seenMsgTTL: TimeCacheDuration, @@ -698,6 +677,41 @@ func NewPubSub(ctx context.Context, h host.Host, rt PubSubRouter, opts ...Option rt.Attach(ps) + var err error + ps.peerComm, err = peercomm.NewRegistry(ctx, peercomm.Config{ + Host: h, QueueSize: ps.peerOutboundQueueSize, + MaxMessageSize: ps.maxMessageSize, MaxControlMessageSize: ps.maxControlMessageSize, + Hooks: peercomm.Hooks{ + InboundOpened: func(a *peercomm.Actor, s network.Stream) { + ps.enqueuePeerEvent(incomingUnion{kind: incomingKindNewStream, actor: a, s: s}) + }, + InboundRPC: func(a *peercomm.Actor, s network.Stream, transport peercomm.Transport, rpc *pb.RPC) { + ps.enqueuePeerEvent(incomingUnion{kind: incomingKindRPC, actor: a, s: s, rpc: wrapInboundRPC(rpc, a.Peer(), transport)}) + }, + InboundClosed: func(a *peercomm.Actor, s network.Stream) { + ps.enqueuePeerEvent(incomingUnion{kind: incomingKindClosedStream, actor: a, s: s}) + }, + OutboundReady: func(a *peercomm.Actor, s network.Stream) { + ps.enqueuePeerEvent(incomingUnion{kind: incomingKindOutboundReady, actor: a, s: s}) + }, + OutboundOpenFailed: func(a *peercomm.Actor, err error) { + ps.enqueuePeerEvent(incomingUnion{kind: incomingKindOutboundOpenFailed, actor: a, err: err}) + }, + OutboundSent: func(a *peercomm.Actor, s network.Stream, rpc *pb.RPC) { + ps.rpcLogger.Debug("sent", "peer", a.Peer(), "rpc", rpc) + }, + OutboundSendFailed: func(a *peercomm.Actor, s network.Stream, rpc *pb.RPC, err error) { + ps.rpcLogger.Debug("failed to send message", "peer", a.Peer(), "rpc", rpc, "err", err) + }, + OutboundDead: func(a *peercomm.Actor, s network.Stream, err error) { + ps.enqueuePeerEvent(incomingUnion{kind: incomingKindOutboundDead, actor: a, s: s, err: err}) + }, + }, + }) + if err != nil { + return nil, err + } + for _, id := range rt.Protocols() { if ps.protoMatchFunc != nil { h.SetStreamHandlerMatch(id, ps.protoMatchFunc(id), ps.handleNewStream) @@ -961,11 +975,7 @@ func WithAppSpecificRpcInspector(inspector func(peer.ID, *RPC) error) Option { // processLoop handles all inputs arriving on the channels func (p *PubSub) processLoop(ctx context.Context) { defer func() { - // Clean up go routines. - for _, queue := range p.peers { - queue.Close() - } - p.peers = nil + p.peerComm.Stop() p.topics = nil p.seenMessages.Done() p.deliveredMessages.Done() @@ -976,36 +986,6 @@ func (p *PubSub) processLoop(ctx context.Context) { case <-p.newPeers: p.handlePendingPeers() - case s := <-p.newPeerStream: - pid := s.Conn().RemotePeer() - - q, ok := p.peers[pid] - if !ok { - p.logger.Warn("new stream for unknown peer", "peer", pid) - s.Cancel() - s.Reset() - continue - } - - if p.blacklist.Contains(pid) { - p.logger.Warn("closing stream for blacklisted peer", "peer", pid) - q.Close() - delete(p.peers, pid) - s.Cancel() - s.Reset() - continue - } - - helloPacket := p.getHelloPacket() - helloPacket = p.rt.OnNewOutboundStream(pid, s.Protocol(), helloPacket) - s.FirstMessage <- helloPacket - - case pid := <-p.newPeerError: - delete(p.peers, pid) - - case <-p.peerDead: - p.handleDeadPeers() - case treq := <-p.getTopics: var out []string for t := range p.mySubs { @@ -1031,26 +1011,42 @@ func (p *PubSub) processLoop(ctx context.Context) { continue } var peers []peer.ID - for p := range p.peers { + for pid := range p.peerComm.All() { if preq.topic != "" { - _, ok := tmap[p] - if !ok { + if _, ok := tmap[pid]; !ok { continue } } - peers = append(peers, p) + peers = append(peers, pid) } preq.resp <- peers case in := <-p.incoming: + if in.actor == nil { + continue + } + if !p.peerComm.IsCurrent(in.actor) { + if in.kind == incomingKindNewStream || in.kind == incomingKindRPC || in.kind == incomingKindClosedStream { + in.s.Reset() + } + continue + } switch in.kind { case incomingKindRPC: p.handleIncomingRPC(in.rpc) case incomingKindNewStream: - p.rt.OnNewIncomingStream( - in.s.Conn().RemotePeer(), in.s.Protocol()) + p.rt.OnNewIncomingStream(in.actor.Peer(), in.s.Protocol()) case incomingKindClosedStream: - p.onClosedIncomingStream( - in.s.Conn().RemotePeer(), in.s.Protocol()) + p.onClosedIncomingStream(in.actor.Peer(), in.s.Protocol()) + case incomingKindOutboundReady: + hello := p.rt.OnNewOutboundStream(in.actor.Peer(), in.s.Protocol(), p.getHelloPacket()) + if hello == nil || in.actor.Activate(in.s, &hello.RPC) != nil { + _ = in.s.Reset() + } + case incomingKindOutboundOpenFailed: + p.logger.Debug("error opening new stream to peer", "err", in.err, "peer", in.actor.Peer()) + p.retirePeer(in.actor, false) + case incomingKindOutboundDead: + p.replaceDeadActor(in.actor, true) } case msg := <-p.sendMsg: p.publishMessage(msg) @@ -1071,12 +1067,8 @@ func (p *PubSub) processLoop(ctx context.Context) { p.logger.Info("Blacklisting peer", "peer", pid) p.blacklist.Add(pid) - q, ok := p.peers[pid] - if ok { - q.Close() - delete(p.peers, pid) - p.clearPeerFromTopicsState(pid) - p.rt.OnClosedOutboundStream(pid) + if actor, ok := p.peerComm.Lookup(pid); ok { + p.retirePeer(actor, true) } case <-ctx.Done(): @@ -1086,6 +1078,10 @@ func (p *PubSub) processLoop(ctx context.Context) { } } +func (p *PubSub) openRequest(backoff time.Duration) peercomm.OpenRequest { + return peercomm.OpenRequest{Protocols: append([]protocol.ID(nil), p.rt.Protocols()...), Backoff: backoff} +} + func (p *PubSub) handlePendingPeers() { p.newPeersPrioLk.Lock() @@ -1105,7 +1101,7 @@ func (p *PubSub) handlePendingPeers() { continue } - if _, ok := p.peers[pid]; ok { + if _, ok := p.peerComm.Lookup(pid); ok { p.logger.Debug("already have connection to peer", "peer", pid) continue } @@ -1115,9 +1111,8 @@ func (p *PubSub) handlePendingPeers() { continue } - rpcQueue := newRpcQueue(p.peerOutboundQueueSize) - p.peers[pid] = rpcQueue - go p.handleNewPeer(p.ctx, pid, rpcQueue) + actor := p.peerComm.GetOrCreate(pid) + _ = actor.Start(p.openRequest(0)) } } @@ -1140,45 +1135,37 @@ func (p *PubSub) clearPeerFromTopicsState(pid peer.ID) { } } -func (p *PubSub) handleDeadPeers() { - p.peerDeadPrioLk.Lock() - - if len(p.peerDeadPend) == 0 { - p.peerDeadPrioLk.Unlock() - return +func (p *PubSub) retirePeer(actor *peercomm.Actor, notifyOutbound bool) bool { + pid := actor.Peer() + retirement, ok := p.peerComm.RetireActor(actor) + if !ok { + return false } - - deadPeers := p.peerDeadPend - p.peerDeadPend = make(map[peer.ID]struct{}) - p.peerDeadPrioLk.Unlock() - - for pid := range deadPeers { - q, ok := p.peers[pid] - if !ok { - continue - } - - q.Close() - delete(p.peers, pid) - - p.clearPeerFromTopicsState(pid) + if retirement.HadInbound { + p.onClosedIncomingStream(pid, retirement.InboundProtocol) + } + p.clearPeerFromTopicsState(pid) + if notifyOutbound { p.rt.OnClosedOutboundStream(pid) + } + return true +} - if p.host.Network().Connectedness(pid) == network.Connected { - backoffDelay, err := p.deadPeerBackoff.updateAndGet(pid) - if err != nil { - p.logger.Debug("error updating backoff", "err", err, "peer", pid) - continue - } - - // still connected, must be a duplicate connection being closed. - // we respawn the writer as we need to ensure there is a stream active - p.logger.Debug("peer declared dead but still connected; respawning writer", "peer", pid) - rpcQueue := newRpcQueue(p.peerOutboundQueueSize) - p.peers[pid] = rpcQueue - go p.handleNewPeerWithBackoff(p.ctx, pid, backoffDelay, rpcQueue) - } +func (p *PubSub) replaceDeadActor(actor *peercomm.Actor, outboundOpened bool) { + pid := actor.Peer() + if !p.retirePeer(actor, outboundOpened) { + return + } + if p.host.Network().Connectedness(pid) != network.Connected || p.blacklist.Contains(pid) { + return + } + backoffDelay, err := p.deadPeerBackoff.updateAndGet(pid) + if err != nil { + p.logger.Debug("error updating backoff", "err", err, "peer", pid) + return } + replacement := p.peerComm.GetOrCreate(pid) + _ = replacement.Start(p.openRequest(backoffDelay)) } // handleAddTopic adds a tracker for a particular topic. @@ -1358,12 +1345,14 @@ func (p *PubSub) announce(topic string, sub bool) { } out := rpcWithSubs(subopt) - for pid, peer := range p.peers { - err := peer.Push(out, false) + for pid, actor := range p.peerComm.All() { + err := actor.Send(&out.RPC, false) if err != nil { - p.logger.Info("Can't send announce message to peer: queue full; scheduling retry", "peer", pid) p.tracer.DropRPC(out, pid) - go p.announceRetry(pid, topic, sub) + if errors.Is(err, peercomm.ErrQueueFull) { + p.logger.Info("Can't send announce message to peer: queue full; scheduling retry", "peer", pid) + go p.announceRetry(pid, topic, sub) + } continue } p.tracer.SendRPC(out, pid) @@ -1391,7 +1380,7 @@ func (p *PubSub) announceRetry(pid peer.ID, topic string, sub bool) { } func (p *PubSub) doAnnounceRetry(pid peer.ID, topic string, sub bool) { - peer, ok := p.peers[pid] + peer, ok := p.peerComm.Lookup(pid) if !ok { return } @@ -1412,11 +1401,13 @@ func (p *PubSub) doAnnounceRetry(pid peer.ID, topic string, sub bool) { } out := rpcWithSubs(subopt) - err := peer.Push(out, false) + err := peer.Send(&out.RPC, false) if err != nil { - p.logger.Info("Can't send announce message to peer: queue full; scheduling retry", "peer", pid) p.tracer.DropRPC(out, pid) - go p.announceRetry(pid, topic, sub) + if errors.Is(err, peercomm.ErrQueueFull) { + p.logger.Info("Can't send announce message to peer: queue full; scheduling retry", "peer", pid) + go p.announceRetry(pid, topic, sub) + } return } p.tracer.SendRPC(out, pid) diff --git a/pubsub_test.go b/pubsub_test.go index 5899b5f9..916fdefd 100644 --- a/pubsub_test.go +++ b/pubsub_test.go @@ -1,17 +1,20 @@ package pubsub import ( + "bytes" "context" "runtime" "testing" "testing/synctest" "time" + "github.com/libp2p/go-libp2p-pubsub/internal/peercomm" pb "github.com/libp2p/go-libp2p-pubsub/pb" "github.com/libp2p/go-libp2p/core/host" "github.com/libp2p/go-libp2p/core/peer" "github.com/libp2p/go-libp2p/x/simlibp2p" "github.com/marcopolo/simnet" + "google.golang.org/protobuf/encoding/protowire" ) // synctestTest wraps synctest.Test with GOMAXPROCS(1) to work around a Go @@ -25,6 +28,28 @@ func synctestTest(t *testing.T, f func(t *testing.T)) { synctest.Test(t, f) } +func TestWrapInboundRPCPreservesMetadataAndCopiesProto(t *testing.T) { + topic := "topic" + source := &pb.RPC{Subscriptions: []*pb.RPC_SubOpts{{Topicid: &topic}}} + unknown := protowire.AppendTag(nil, 100, protowire.BytesType) + unknown = protowire.AppendBytes(unknown, []byte("extension")) + source.ProtoReflect().SetUnknown(unknown) + + from := peer.ID("peer-a") + wrapped := wrapInboundRPC(source, from, peercomm.TransportTopic) + if wrapped.from != from || wrapped.transport != peercomm.TransportTopic { + t.Fatalf("metadata = (%q, %v), want (%q, %v)", wrapped.from, wrapped.transport, from, peercomm.TransportTopic) + } + if wrapped.Subscriptions[0].GetTopicid() != topic || !bytes.Equal(wrapped.ProtoReflect().GetUnknown(), unknown) { + t.Fatal("protobuf fields were not preserved") + } + source.Subscriptions[0].Topicid = nil + source.ProtoReflect().SetUnknown(nil) + if wrapped.Subscriptions[0].GetTopicid() != topic || !bytes.Equal(wrapped.ProtoReflect().GetUnknown(), unknown) { + t.Fatal("wrapped RPC shares mutable protobuf state with source") + } +} + func TestClearPeerFromTopicsStateRemovesEmptyTopicMap(t *testing.T) { pid := peer.ID("peer-a") other := peer.ID("peer-b") diff --git a/randomsub.go b/randomsub.go index b5ca3034..ce11e17a 100644 --- a/randomsub.go +++ b/randomsub.go @@ -150,12 +150,12 @@ func (rs *RandomSubRouter) Publish(msg *Message) { out := rpcWithMessages(msg.Message) for p := range tosend { - q, ok := rs.p.peers[p] + q, ok := rs.p.peerComm.Lookup(p) if !ok { continue } - err := q.Push(out, false) + err := q.Send(&out.RPC, false) if err != nil { rs.p.logger.Info("dropping message to peer: queue full", "peer", p) rs.tracer.DropRPC(out, p) diff --git a/rpc_queue.go b/rpc_queue.go deleted file mode 100644 index e5c22935..00000000 --- a/rpc_queue.go +++ /dev/null @@ -1,147 +0,0 @@ -package pubsub - -import ( - "context" - "errors" - "sync" -) - -var ( - ErrQueueCancelled = errors.New("rpc queue operation cancelled") - ErrQueueClosed = errors.New("rpc queue closed") - ErrQueueFull = errors.New("rpc queue full") - ErrQueuePushOnClosed = errors.New("push on closed rpc queue") -) - -type priorityQueue struct { - normal []*RPC - priority []*RPC -} - -func (q *priorityQueue) Len() int { - return len(q.normal) + len(q.priority) -} - -func (q *priorityQueue) NormalPush(rpc *RPC) { - q.normal = append(q.normal, rpc) -} - -func (q *priorityQueue) PriorityPush(rpc *RPC) { - q.priority = append(q.priority, rpc) -} - -func (q *priorityQueue) Pop() *RPC { - var rpc *RPC - - if len(q.priority) > 0 { - rpc = q.priority[0] - q.priority[0] = nil - q.priority = q.priority[1:] - } else if len(q.normal) > 0 { - rpc = q.normal[0] - q.normal[0] = nil - q.normal = q.normal[1:] - } - - return rpc -} - -type rpcQueue struct { - dataAvailable sync.Cond - spaceAvailable sync.Cond - // Mutex used to access queue - queueMu sync.Mutex - queue priorityQueue - - closed bool - maxSize int -} - -func newRpcQueue(maxSize int) *rpcQueue { - q := &rpcQueue{maxSize: maxSize} - q.dataAvailable.L = &q.queueMu - q.spaceAvailable.L = &q.queueMu - return q -} - -func (q *rpcQueue) Push(rpc *RPC, block bool) error { - return q.push(rpc, false, block) -} - -func (q *rpcQueue) UrgentPush(rpc *RPC, block bool) error { - return q.push(rpc, true, block) -} - -func (q *rpcQueue) push(rpc *RPC, urgent bool, block bool) error { - q.queueMu.Lock() - defer q.queueMu.Unlock() - - if q.closed { - panic(ErrQueuePushOnClosed) - } - - for q.queue.Len() == q.maxSize { - if block { - q.spaceAvailable.Wait() - // It can receive a signal because the queue is closed. - if q.closed { - panic(ErrQueuePushOnClosed) - } - } else { - return ErrQueueFull - } - } - if urgent { - q.queue.PriorityPush(rpc) - } else { - q.queue.NormalPush(rpc) - } - - q.dataAvailable.Signal() - return nil -} - -// Note that, when the queue is empty and there are two blocked Pop calls, it -// doesn't mean that the first Pop will get the item from the next Push. The -// second Pop will probably get it instead. -func (q *rpcQueue) Pop(ctx context.Context) (*RPC, error) { - q.queueMu.Lock() - defer q.queueMu.Unlock() - - if q.closed { - return nil, ErrQueueClosed - } - - unregisterAfterFunc := context.AfterFunc(ctx, func() { - // Wake up all the waiting routines. The only routine that correponds - // to this Pop call will return from the function. Note that this can - // be expensive, if there are too many waiting routines. - q.dataAvailable.Broadcast() - }) - defer unregisterAfterFunc() - - for q.queue.Len() == 0 { - select { - case <-ctx.Done(): - return nil, ErrQueueCancelled - default: - } - q.dataAvailable.Wait() - // It can receive a signal because the queue is closed. - if q.closed { - return nil, ErrQueueClosed - } - } - rpc := q.queue.Pop() - q.spaceAvailable.Signal() - return rpc, nil -} - -func (q *rpcQueue) Close() { - q.queueMu.Lock() - defer q.queueMu.Unlock() - - q.closed = true - q.dataAvailable.Broadcast() - q.spaceAvailable.Broadcast() -} diff --git a/rpc_queue_test.go b/rpc_queue_test.go deleted file mode 100644 index 6380d383..00000000 --- a/rpc_queue_test.go +++ /dev/null @@ -1,256 +0,0 @@ -package pubsub - -import ( - "context" - "testing" - "time" -) - -func TestNewRpcQueue(t *testing.T) { - synctestTest(t, func(t *testing.T) { - maxSize := 32 - q := newRpcQueue(maxSize) - if q.maxSize != maxSize { - t.Fatalf("rpc queue has wrong max size, expected %d but got %d", maxSize, q.maxSize) - } - if q.dataAvailable.L != &q.queueMu { - t.Fatalf("the dataAvailable field of rpc queue has an incorrect mutex") - } - if q.spaceAvailable.L != &q.queueMu { - t.Fatalf("the spaceAvailable field of rpc queue has an incorrect mutex") - } - }) -} - -func TestRpcQueueUrgentPush(t *testing.T) { - synctestTest(t, func(t *testing.T) { - maxSize := 32 - q := newRpcQueue(maxSize) - - rpc1 := &RPC{} - rpc2 := &RPC{} - rpc3 := &RPC{} - rpc4 := &RPC{} - q.Push(rpc1, true) - q.UrgentPush(rpc2, true) - q.Push(rpc3, true) - q.UrgentPush(rpc4, true) - pop1, err := q.Pop(context.Background()) - if err != nil { - t.Fatal(err) - } - pop2, err := q.Pop(context.Background()) - if err != nil { - t.Fatal(err) - } - pop3, err := q.Pop(context.Background()) - if err != nil { - t.Fatal(err) - } - pop4, err := q.Pop(context.Background()) - if err != nil { - t.Fatal(err) - } - if pop1 != rpc2 { - t.Fatalf("get wrong item from rpc queue Pop") - } - if pop2 != rpc4 { - t.Fatalf("get wrong item from rpc queue Pop") - } - if pop3 != rpc1 { - t.Fatalf("get wrong item from rpc queue Pop") - } - if pop4 != rpc3 { - t.Fatalf("get wrong item from rpc queue Pop") - } - }) -} - -func TestRpcQueuePushThenPop(t *testing.T) { - synctestTest(t, func(t *testing.T) { - maxSize := 32 - q := newRpcQueue(maxSize) - - rpc1 := &RPC{} - rpc2 := &RPC{} - q.Push(rpc1, true) - q.Push(rpc2, true) - pop1, err := q.Pop(context.Background()) - if err != nil { - t.Fatal(err) - } - pop2, err := q.Pop(context.Background()) - if err != nil { - t.Fatal(err) - } - if pop1 != rpc1 { - t.Fatalf("get wrong item from rpc queue Pop") - } - if pop2 != rpc2 { - t.Fatalf("get wrong item from rpc queue Pop") - } - }) -} - -func TestRpcQueuePopThenPush(t *testing.T) { - synctestTest(t, func(t *testing.T) { - maxSize := 32 - q := newRpcQueue(maxSize) - - rpc1 := &RPC{} - rpc2 := &RPC{} - go func() { - // Wait to make sure the main goroutine is blocked. - time.Sleep(1 * time.Millisecond) - q.Push(rpc1, true) - q.Push(rpc2, true) - }() - pop1, err := q.Pop(context.Background()) - if err != nil { - t.Fatal(err) - } - pop2, err := q.Pop(context.Background()) - if err != nil { - t.Fatal(err) - } - if pop1 != rpc1 { - t.Fatalf("get wrong item from rpc queue Pop") - } - if pop2 != rpc2 { - t.Fatalf("get wrong item from rpc queue Pop") - } - }) -} - -func TestRpcQueueBlockPushWhenFull(t *testing.T) { - synctestTest(t, func(t *testing.T) { - maxSize := 1 - q := newRpcQueue(maxSize) - - finished := make(chan struct{}) - q.Push(&RPC{}, true) - go func() { - defer func() { recover() }() // recover from close - q.Push(&RPC{}, true) - finished <- struct{}{} - }() - // Wait to make sure the goroutine is blocked. - time.Sleep(1 * time.Millisecond) - select { - case <-finished: - t.Fatalf("blocking rpc queue Push is not blocked when it is full") - default: - } - // Unblock the goroutine so synctest can exit cleanly - q.Close() - }) -} - -func TestRpcQueueNonblockPushWhenFull(t *testing.T) { - synctestTest(t, func(t *testing.T) { - maxSize := 1 - q := newRpcQueue(maxSize) - - q.Push(&RPC{}, true) - err := q.Push(&RPC{}, false) - if err != ErrQueueFull { - t.Fatalf("non-blocking rpc queue Push returns wrong error when it is full") - } - }) -} - -func TestRpcQueuePushAfterClose(t *testing.T) { - synctestTest(t, func(t *testing.T) { - maxSize := 32 - q := newRpcQueue(maxSize) - q.Close() - - defer func() { - if r := recover(); r == nil { - t.Fatalf("rpc queue Push does not panick after closed") - } - }() - q.Push(&RPC{}, true) - }) -} - -func TestRpcQueuePopAfterClose(t *testing.T) { - synctestTest(t, func(t *testing.T) { - maxSize := 32 - q := newRpcQueue(maxSize) - q.Close() - _, err := q.Pop(context.Background()) - if err != ErrQueueClosed { - t.Fatalf("rpc queue Pop returns wrong error after closed") - } - }) -} - -func TestRpcQueueCloseWhilePush(t *testing.T) { - synctestTest(t, func(t *testing.T) { - maxSize := 1 - q := newRpcQueue(maxSize) - q.Push(&RPC{}, true) - - defer func() { - if r := recover(); r == nil { - t.Fatalf("rpc queue Push does not panick when it's closed on the fly") - } - }() - - go func() { - // Wait to make sure the main goroutine is blocked. - time.Sleep(1 * time.Millisecond) - q.Close() - }() - q.Push(&RPC{}, true) - }) -} - -func TestRpcQueueCloseWhilePop(t *testing.T) { - synctestTest(t, func(t *testing.T) { - maxSize := 32 - q := newRpcQueue(maxSize) - go func() { - // Wait to make sure the main goroutine is blocked. - time.Sleep(1 * time.Millisecond) - q.Close() - }() - _, err := q.Pop(context.Background()) - if err != ErrQueueClosed { - t.Fatalf("rpc queue Pop returns wrong error when it's closed on the fly") - } - }) -} - -func TestRpcQueuePushWhenFullThenPop(t *testing.T) { - synctestTest(t, func(t *testing.T) { - maxSize := 1 - q := newRpcQueue(maxSize) - - q.Push(&RPC{}, true) - go func() { - // Wait to make sure the main goroutine is blocked. - time.Sleep(1 * time.Millisecond) - q.Pop(context.Background()) - }() - q.Push(&RPC{}, true) - }) -} - -func TestRpcQueueCancelPop(t *testing.T) { - synctestTest(t, func(t *testing.T) { - maxSize := 32 - q := newRpcQueue(maxSize) - ctx, cancel := context.WithCancel(context.Background()) - go func() { - // Wait to make sure the main goroutine is blocked. - time.Sleep(1 * time.Millisecond) - cancel() - }() - _, err := q.Pop(ctx) - if err != ErrQueueCancelled { - t.Fatalf("rpc queue Pop returns wrong error when it's cancelled") - } - }) -}