diff --git a/swarm/pss/prox_test.go b/swarm/pss/prox_test.go index 1b3e1511dd..b007fc18e0 100644 --- a/swarm/pss/prox_test.go +++ b/swarm/pss/prox_test.go @@ -5,7 +5,6 @@ import ( "crypto/ecdsa" "encoding/binary" "fmt" - "strings" "sync" "testing" "time" @@ -37,14 +36,13 @@ type handlerNotification struct { } type testData struct { - sim *simulation.Simulation - kademlias map[enode.ID]*network.Kademlia - nodeAddrs map[enode.ID][]byte // make predictable overlay addresses from the generated random enode ids - senders map[int]enode.ID // originating nodes of the messages (intention is to choose as far as possible from the receiving neighborhood) - msgs [][]byte // recipient addresses of messages + sim *simulation.Simulation + kademlias map[enode.ID]*network.Kademlia + nodeAddresses map[enode.ID][]byte // make predictable overlay addresses from the generated random enode ids + senders map[int]enode.ID // originating nodes of the messages (intention is to choose as far as possible from the receiving neighborhood) + recipientAddresses [][]byte requiredMsgCount int - allowedMsgCount int requiredMsgs map[enode.ID][]uint64 // message serials we expect respective nodes to receive allowedMsgs map[enode.ID][]uint64 // message serials we expect respective nodes to receive @@ -52,10 +50,6 @@ type testData struct { totalMsgCount int handlerDone bool // set to true on termination of the simulation run mu sync.Mutex - - doneC chan struct{} // terminates the handler channel listener - errC chan error // error to pass to main sim thread - msgC chan handlerNotification // message receipt notification to main sim thread } var ( @@ -63,63 +57,60 @@ var ( topic = BytesToTopic([]byte{0xf3, 0x9e, 0x06, 0x82}) ) -func (d *testData) pushNotification(val handlerNotification) { - d.mu.Lock() - d.notifications = append(d.notifications, val) - d.mu.Unlock() +func (td *testData) pushNotification(val handlerNotification) { + td.mu.Lock() + td.notifications = append(td.notifications, val) + td.mu.Unlock() } -func (d *testData) popNotification() (ret handlerNotification, found bool) { - d.mu.Lock() - if len(d.notifications) > 0 { - found = true - ret = d.notifications[0] - d.notifications = d.notifications[1:] +func (td *testData) popNotification() (first handlerNotification, ok bool) { + td.mu.Lock() + if len(td.notifications) > 0 { + ok = true + first = td.notifications[0] + td.notifications = td.notifications[1:] } - d.mu.Unlock() - return ret, found + td.mu.Unlock() + return first, ok } -func (d *testData) getMsgCount() int { - d.mu.Lock() - defer d.mu.Unlock() - return d.totalMsgCount +func (td *testData) getMsgCount() int { + td.mu.Lock() + defer td.mu.Unlock() + return td.totalMsgCount } -func (d *testData) incrementMsgCount() int { - d.mu.Lock() - defer d.mu.Unlock() - d.totalMsgCount++ - return d.totalMsgCount +func (td *testData) incrementMsgCount() int { + td.mu.Lock() + defer td.mu.Unlock() + td.totalMsgCount++ + return td.totalMsgCount } -func (d *testData) isDone() bool { - d.mu.Lock() - defer d.mu.Unlock() - return d.handlerDone +func (td *testData) isDone() bool { + td.mu.Lock() + defer td.mu.Unlock() + return td.handlerDone } -func (d *testData) setDone() { - d.mu.Lock() - defer d.mu.Unlock() - d.handlerDone = true +func (td *testData) setDone() { + td.mu.Lock() + defer td.mu.Unlock() + td.handlerDone = true } func newTestData() *testData { return &testData{ - kademlias: make(map[enode.ID]*network.Kademlia), - nodeAddrs: make(map[enode.ID][]byte), - requiredMsgs: make(map[enode.ID][]uint64), - allowedMsgs: make(map[enode.ID][]uint64), - senders: make(map[int]enode.ID), - doneC: make(chan struct{}), - errC: make(chan error), - msgC: make(chan handlerNotification), + kademlias: make(map[enode.ID]*network.Kademlia), + nodeAddresses: make(map[enode.ID][]byte), + requiredMsgs: make(map[enode.ID][]uint64), + allowedMsgs: make(map[enode.ID][]uint64), + senders: make(map[int]enode.ID), } } -func (d *testData) getKademlia(nodeId *enode.ID) (*network.Kademlia, error) { - kadif, ok := d.sim.NodeItem(*nodeId, simulation.BucketKeyKademlia) +func (td *testData) getKademlia(nodeId *enode.ID) (*network.Kademlia, error) { + kadif, ok := td.sim.NodeItem(*nodeId, simulation.BucketKeyKademlia) if !ok { return nil, fmt.Errorf("no kademlia entry for %v", nodeId) } @@ -130,29 +121,29 @@ func (d *testData) getKademlia(nodeId *enode.ID) (*network.Kademlia, error) { return kad, nil } -func (d *testData) init(msgCount int) error { +func (td *testData) init(msgCount int) error { log.Debug("TestProxNetwork start") - for _, nodeId := range d.sim.NodeIDs() { - kad, err := d.getKademlia(&nodeId) + for _, nodeId := range td.sim.NodeIDs() { + kad, err := td.getKademlia(&nodeId) if err != nil { return err } - d.nodeAddrs[nodeId] = kad.BaseAddr() + td.nodeAddresses[nodeId] = kad.BaseAddr() } for i := 0; i < int(msgCount); i++ { msgAddr := pot.RandomAddress() // we choose message addresses randomly - d.msgs = append(d.msgs, msgAddr.Bytes()) + td.recipientAddresses = append(td.recipientAddresses, msgAddr.Bytes()) smallestPo := 256 var targets []enode.ID var closestPO int // loop through all nodes and find the required and allowed recipients of each message // (for more information, please see the comment to the main test function) - for _, nod := range d.sim.Net.GetNodes() { - po, _ := pof(d.msgs[i], d.nodeAddrs[nod.ID()], 0) - depth := d.kademlias[nod.ID()].NeighbourhoodDepth() + for _, nod := range td.sim.Net.GetNodes() { + po, _ := pof(td.recipientAddresses[i], td.nodeAddresses[nod.ID()], 0) + depth := td.kademlias[nod.ID()].NeighbourhoodDepth() // only nodes with closest IDs (wrt the msg address) will be required recipients if po > closestPO { @@ -164,26 +155,25 @@ func (d *testData) init(msgCount int) error { } if po >= depth { - d.allowedMsgCount++ - d.allowedMsgs[nod.ID()] = append(d.allowedMsgs[nod.ID()], uint64(i)) + td.allowedMsgs[nod.ID()] = append(td.allowedMsgs[nod.ID()], uint64(i)) } // a node with the smallest PO (wrt msg) will be the sender, // in order to increase the distance the msg must travel if po < smallestPo { smallestPo = po - d.senders[i] = nod.ID() + td.senders[i] = nod.ID() } } - d.requiredMsgCount += len(targets) + td.requiredMsgCount += len(targets) for _, id := range targets { - d.requiredMsgs[id] = append(d.requiredMsgs[id], uint64(i)) + td.requiredMsgs[id] = append(td.requiredMsgs[id], uint64(i)) } - log.Debug("nn for msg", "targets", len(targets), "msgidx", i, "msg", common.Bytes2Hex(msgAddr[:8]), "sender", d.senders[i], "senderpo", smallestPo) + log.Debug("nn for msg", "targets", len(targets), "msgidx", i, "msg", common.Bytes2Hex(msgAddr[:8]), "sender", td.senders[i], "senderpo", smallestPo) } - log.Debug("msgs to receive", "count", d.requiredMsgCount) + log.Debug("recipientAddresses to receive", "count", td.requiredMsgCount) return nil } @@ -207,11 +197,10 @@ func (d *testData) init(msgCount int) error { // whereas nodes X, Y and Z will be allowed recipients. func TestProxNetwork(t *testing.T) { t.Run("16_nodes,_16_messages,_16_seconds", func(t *testing.T) { - testProxNetwork(t, 16, 16, 3*time.Second) + testProxNetwork(t, 16, 16, 16*time.Second) }) } -// params in run name: nodes/msgs func TestProxNetworkLong(t *testing.T) { if !*longrunning { t.Skip("run with --longrunning flag to run extensive network tests") @@ -256,7 +245,7 @@ func testProxNetwork(t *testing.T, nodeCount int, msgCount int, timeout time.Dur } result := td.sim.Run(ctx, wrapper) // call the main test function if result.Error != nil { - timedOut := strings.Compare(result.Error.Error(), "context deadline exceeded") == 0 + timedOut := result.Error == context.DeadlineExceeded if !timedOut || td.getMsgCount() < td.requiredMsgCount { t.Fatal(result.Error) } @@ -265,7 +254,7 @@ func testProxNetwork(t *testing.T, nodeCount int, msgCount int, timeout time.Dur func (td *testData) sendAllMsgs() error { nodes := make(map[int]*rpc.Client) - for i := range td.msgs { + for i := range td.recipientAddresses { nodeClient, err := td.sim.Net.GetNode(td.senders[i]).Client() if err != nil { return err @@ -273,7 +262,7 @@ func (td *testData) sendAllMsgs() error { nodes[i] = nodeClient } - for i, msg := range td.msgs { + for i, msg := range td.recipientAddresses { log.Debug("sending msg", "idx", i, "from", td.senders[i]) nodeClient := nodes[i] var uvarByte [8]byte @@ -285,60 +274,64 @@ func (td *testData) sendAllMsgs() error { // testRoutine is the main test function, called by Simulation.Run() func testRoutine(td *testData, ctx context.Context) error { - go handlerChannelListener(td) - res := td.sendAllMsgs() - if res != nil { - return res + if err := td.sendAllMsgs(); err != nil { + return err } - received := 0 - - for { + isMoreTimeLeft := func() bool { select { - case <-ctx.Done(): // timeout or cancel - td.setDone() - if td.getMsgCount() < td.requiredMsgCount { - res = ctx.Err() - } - case err := <-td.errC: - // only the first error matters - if res == nil && err != nil { - res = err - td.setDone() - } - case hn := <-td.msgC: - received++ - log.Debug("msg received", "msgs_received", received, "total_expected", td.requiredMsgCount, "id", hn.id, "serial", hn.serial) - case <-td.doneC: - return res + case <-ctx.Done(): + return false + default: + return true } } -} -func handlerChannelListener(d *testData) { - for !d.isDone() { - notification, exist := d.popNotification() - if exist { - if d.isAllowedMessage(notification) { - d.msgC <- notification //notify the main sim thread + hasMoreRound := func(err error, hadMessage bool) bool { + return err == nil && (hadMessage || isMoreTimeLeft()) + } + + var err error + received := 0 + hadMessage := false + for oneMoreRound := true; oneMoreRound; oneMoreRound = hasMoreRound(err, hadMessage) { + message, hadMessage := td.popNotification() + + if !isMoreTimeLeft() { + // Stop handlers from sending more messages. + // Note: only best effort, race is possible. + td.setDone() + } + + if hadMessage { + if td.isAllowedMessage(message) { + received++ + log.Debug("msg received", "msgs_received", received, "total_expected", td.requiredMsgCount, "id", message.id, "serial", message.serial) } else { - log.Error("message received by wrong recipient", "num", notification.serial) - d.errC <- fmt.Errorf("message %d received by wrong recipient %v", notification.serial, notification.id) - break + err = fmt.Errorf("message %d received by wrong recipient %v", message.serial, message.id) } } else { - time.Sleep(time.Millisecond * 32) + time.Sleep(32 * time.Millisecond) } + + oneMoreRound = err == nil && (hadMessage || isMoreTimeLeft()) } - close(d.doneC) - close(d.errC) - close(d.msgC) + + if err != nil { + return err + } + + if td.getMsgCount() < td.requiredMsgCount { + return ctx.Err() + } + return nil + } -func (d *testData) isAllowedMessage(n handlerNotification) bool { +func (td *testData) isAllowedMessage(n handlerNotification) bool { // check if message serial is in expected messages for this recipient - for _, s := range d.allowedMsgs[n.id] { + for _, s := range td.allowedMsgs[n.id] { if n.serial == s { return true } @@ -346,10 +339,10 @@ func (d *testData) isAllowedMessage(n handlerNotification) bool { return false } -func (d *testData) removeAllowedMessage(id enode.ID, index int) { - last := len(d.allowedMsgs[id]) - 1 - d.allowedMsgs[id][index] = d.allowedMsgs[id][last] - d.allowedMsgs[id] = d.allowedMsgs[id][:last] +func (td *testData) removeAllowedMessage(id enode.ID, index int) { + last := len(td.allowedMsgs[id]) - 1 + td.allowedMsgs[id][index] = td.allowedMsgs[id][last] + td.allowedMsgs[id] = td.allowedMsgs[id][:last] } func nodeMsgHandler(td *testData, config *adapters.NodeConfig) *handler { diff --git a/swarm/pss/pss_test.go b/swarm/pss/pss_test.go index ea7a591b1e..9884ffbe94 100644 --- a/swarm/pss/pss_test.go +++ b/swarm/pss/pss_test.go @@ -1364,7 +1364,7 @@ func TestNetwork(t *testing.T) { } // params in run name: -// nodes/msgs/addrbytes/adaptertype +// nodes/recipientAddresses/addrbytes/adaptertype // if adaptertype is exec uses execadapter, simadapter otherwise func TestNetwork2000(t *testing.T) { if !*longrunning {