diff --git a/cmd/devp2p/internal/ethtest/protocol.go b/cmd/devp2p/internal/ethtest/protocol.go index a21d1ca7a1..af76082318 100644 --- a/cmd/devp2p/internal/ethtest/protocol.go +++ b/cmd/devp2p/internal/ethtest/protocol.go @@ -86,3 +86,9 @@ func protoOffset(proto Proto) uint64 { panic("unhandled protocol") } } + +// msgTypePtr is the constraint for protocol message types. +type msgTypePtr[U any] interface { + *U + Kind() byte +} diff --git a/cmd/devp2p/internal/ethtest/suite.go b/cmd/devp2p/internal/ethtest/suite.go index f6de2a2db0..c23360bf82 100644 --- a/cmd/devp2p/internal/ethtest/suite.go +++ b/cmd/devp2p/internal/ethtest/suite.go @@ -237,39 +237,28 @@ concurrently, with different request IDs.`) // Wait for responses. // Note they can arrive in either order. - var resp1 *eth.BlockHeadersPacket - var resp2 *eth.BlockHeadersPacket - for range 2 { - r := new(eth.BlockHeadersPacket) - if err := conn.ReadMsg(ethProto, eth.BlockHeadersMsg, &r); err != nil { - t.Fatalf("error reading BlockHeaders msg: %v", err) - } - switch r.RequestId { - case req1.RequestId: - if resp1 != nil { - t.Fatalf("duplicate response for request with ID %v", req1.RequestId) - } - resp1 = r - case req2.RequestId: - if resp2 != nil { - t.Fatalf("duplicate response for request with ID %v", req2.RequestId) - } - resp2 = r - default: - t.Fatalf("unknown request ID in response: %v", r.RequestId) + resp, err := collectResponses(conn, 2, func(msg *eth.BlockHeadersPacket) uint64 { + if msg.RequestId != 111 && msg.RequestId != 222 { + t.Fatalf("response with unknown request ID: %v", msg.RequestId) } + return msg.RequestId + }) + if err != nil { + t.Fatal(err) } // Check if headers match. + resp1 := resp[111] if expected, err := s.chain.GetHeaders(req1); err != nil { t.Fatalf("failed to get expected headers for request 1: %v", err) } else if !headersMatch(expected, resp1.BlockHeadersRequest) { - t.Fatalf("header mismatch: \nexpected %v \ngot %v", expected, resp1) + t.Fatalf("header mismatch for request ID %v: \nexpected %v \ngot %v", 111, expected, resp1) } + resp2 := resp[222] if expected, err := s.chain.GetHeaders(req2); err != nil { t.Fatalf("failed to get expected headers for request 2: %v", err) } else if !headersMatch(expected, resp2.BlockHeadersRequest) { - t.Fatalf("header mismatch: \nexpected %v \ngot %v", expected, resp2) + t.Fatalf("header mismatch for request ID %v: \nexpected %v \ngot %v", 222, expected, resp2) } } @@ -314,42 +303,50 @@ same request ID. The node should handle the request by responding to both reques // Wait for the responses. They can arrive in either order, and we can't tell them // apart by their request ID, so use the number of headers instead. - var resp1 *eth.BlockHeadersPacket - var resp2 *eth.BlockHeadersPacket - for range 2 { - r := new(eth.BlockHeadersPacket) - if err := conn.ReadMsg(ethProto, eth.BlockHeadersMsg, &r); err != nil { - t.Fatalf("error reading BlockHeaders msg: %v", err) - } - switch uint64(len(r.BlockHeadersRequest)) { - case request1.GetBlockHeadersRequest.Amount: - if resp1 != nil { - t.Fatalf("duplicate response with %d headers", request1.GetBlockHeadersRequest.Amount) - } - resp1 = r - case request2.GetBlockHeadersRequest.Amount: - if resp2 != nil { - t.Fatalf("duplicate response with %d headers", request2.GetBlockHeadersRequest.Amount) - } - resp2 = r - default: - t.Fatalf("invalid number of items in response: %d", len(r.BlockHeadersRequest)) + resp, err := collectResponses(conn, 2, func(msg *eth.BlockHeadersPacket) uint64 { + id := uint64(len(msg.BlockHeadersRequest)) + if id != 2 && id != 3 { + t.Fatalf("invalid number of headers in response: %d", id) } + return id + }) + if err != nil { + t.Fatal(err) } // Check if headers match. + resp1 := resp[2] if expected, err := s.chain.GetHeaders(request1); err != nil { t.Fatalf("failed to get expected headers for request 1: %v", err) } else if !headersMatch(expected, resp1.BlockHeadersRequest) { - t.Fatalf("header mismatch: \nexpected %v \ngot %v", expected, resp1) + t.Fatalf("headers mismatch: \nexpected %v \ngot %v", expected, resp1) } + resp2 := resp[3] if expected, err := s.chain.GetHeaders(request2); err != nil { t.Fatalf("failed to get expected headers for request 2: %v", err) } else if !headersMatch(expected, resp2.BlockHeadersRequest) { - t.Fatalf("header mismatch: \nexpected %v \ngot %v", expected, resp2) + t.Fatalf("headers mismatch: \nexpected %v \ngot %v", expected, resp2) } } +// collectResponses waits for n messages of type T on the given connection. +// The messsages are collected according to the 'identity' function. +func collectResponses[T any, P msgTypePtr[T]](conn *Conn, n int, identity func(P) uint64) (map[uint64]P, error) { + resp := make(map[uint64]P, n) + for range n { + r := new(T) + if err := conn.ReadMsg(ethProto, eth.BlockHeadersMsg, r); err != nil { + return resp, fmt.Errorf("read error: %v", err) + } + id := identity(r) + if resp[id] != nil { + return resp, fmt.Errorf("duplicate response %v", r) + } + resp[id] = r + } + return resp, nil +} + func (s *Suite) TestZeroRequestID(t *utesting.T) { t.Log(`This test sends a GetBlockHeaders message with a request-id of zero, and expects a response.`)