From 3fcebc1c8079166f246d1ea8421cd7d5f86d4b87 Mon Sep 17 00:00:00 2001 From: Vlad Date: Tue, 12 Mar 2019 20:49:14 +0400 Subject: [PATCH] swarm/pss: eliminated the global variables --- swarm/pss/prox_test.go | 227 ++++++++++++++++++++--------------------- 1 file changed, 108 insertions(+), 119 deletions(-) diff --git a/swarm/pss/prox_test.go b/swarm/pss/prox_test.go index 2beb5ab815..e87d1aa6ee 100644 --- a/swarm/pss/prox_test.go +++ b/swarm/pss/prox_test.go @@ -30,7 +30,7 @@ import ( ) // needed to make the enode id of the receiving node available to the handler for triggers -type handlerContextFunc func(*adapters.NodeConfig) *handler +type handlerContextFunc func(*testData, *adapters.NodeConfig) *handler // struct to notify reception of messages to simulation driver // TODO To make code cleaner: @@ -41,81 +41,55 @@ type handlerNotification struct { serial uint64 } -var ( - pof = pot.DefaultPof(256) // generate messages and index them - topic = BytesToTopic([]byte{0x00, 0x00, 0x06, 0x82}) - mu sync.Mutex // keeps handlerDone in sync - sim *simulation.Simulation - - mx sync.Mutex // prevents data race for the test variables - handlerDone bool // set to true on termination of the simulation run +type testData struct { + mu sync.Mutex + sim *simulation.Simulation + handlerDone bool // set to true on termination of the simulation run requiredMessages int allowedMessages int messageCount int + kademlias map[enode.ID]*network.Kademlia + nodeAddrs map[enode.ID][]byte // make predictable overlay addresses from the generated random enode ids + recipients map[int][]enode.ID // for logging output only + allowed map[int][]enode.ID // allowed recipients + expectedMsgs 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 + senders map[int]enode.ID // originating nodes of the messages (intention is to choose as far as possible from the receiving neighborhood) + handlerC chan handlerNotification // passes message from pss message handler to simulation driver + 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 + msgs [][]byte // recipient addresses of messages +} - kademlias map[enode.ID]*network.Kademlia - nodeAddrs map[enode.ID][]byte // make predictable overlay addresses from the generated random enode ids - recipients map[int][]enode.ID // for logging output only - allowed map[int][]enode.ID // allowed recipients - expectedMsgs 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 - senders map[int]enode.ID // originating nodes of the messages (intention is to choose as far as possible from the receiving neighborhood) - handlerC chan handlerNotification // passes message from pss message handler to simulation driver - 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 - msgs [][]byte // recipient addresses of messages +var ( + pof = pot.DefaultPof(256) // generate messages and index them + topic = BytesToTopic([]byte{0x00, 0x00, 0x06, 0x82}) ) -func resetTestVariables() { - handlerDone = false - requiredMessages = 0 - allowedMessages = 0 - msgs = nil - sim = nil - resetMsgCount() - kademlias = make(map[enode.ID]*network.Kademlia) - nodeAddrs = make(map[enode.ID][]byte) - recipients = make(map[int][]enode.ID) - allowed = make(map[int][]enode.ID) - expectedMsgs = make(map[enode.ID][]uint64) - allowedMsgs = make(map[enode.ID][]uint64) - senders = make(map[int]enode.ID) - handlerC = make(chan handlerNotification) - doneC = make(chan struct{}) - errC = make(chan error) - msgC = make(chan handlerNotification) +func (d *testData) getMsgCount() int { + d.mu.Lock() + defer d.mu.Unlock() + return d.messageCount } -func getMsgCount() int { - mx.Lock() - defer mx.Unlock() - return messageCount +func (d *testData) incrementMsgCount() int { + d.mu.Lock() + defer d.mu.Unlock() + d.messageCount++ + return d.messageCount } -func resetMsgCount() { - mx.Lock() - messageCount = 0 - mx.Unlock() +func (d *testData) isDone() bool { + d.mu.Lock() + defer d.mu.Unlock() + return d.handlerDone } -func incrementMsgCount() int { - mx.Lock() - defer mx.Unlock() - messageCount++ - return messageCount -} - -func isDone() bool { - mu.Lock() - defer mu.Unlock() - return handlerDone -} - -func setDone() { - mu.Lock() - defer mu.Unlock() - handlerDone = true +func (d *testData) setDone() { + d.mu.Lock() + defer d.mu.Unlock() + d.handlerDone = true } func getCmdParams(t *testing.T) (int, int) { @@ -149,23 +123,34 @@ func readSnapshot(t *testing.T, nodeCount int) simulations.Snapshot { return snap } -func assignTestVariables(sim *simulation.Simulation, msgCount int) { +func initializeTestData(d *testData, msgCount int) { log.Debug("TestProxNetwork start") - for _, nodeId := range sim.NodeIDs() { - nodeAddrs[nodeId] = nodeIDToAddr(nodeId) + d.nodeAddrs = make(map[enode.ID][]byte) + d.recipients = make(map[int][]enode.ID) + d.allowed = make(map[int][]enode.ID) + d.expectedMsgs = make(map[enode.ID][]uint64) + d.allowedMsgs = make(map[enode.ID][]uint64) + d.senders = make(map[int]enode.ID) + d.handlerC = make(chan handlerNotification) + d.doneC = make(chan struct{}) + d.errC = make(chan error) + d.msgC = make(chan handlerNotification) + + for _, nodeId := range d.sim.NodeIDs() { + d.nodeAddrs[nodeId] = nodeIDToAddr(nodeId) } for i := 0; i < int(msgCount); i++ { msgAddr := pot.RandomAddress() // we choose message addresses randomly - msgs = append(msgs, msgAddr.Bytes()) + d.msgs = append(d.msgs, msgAddr.Bytes()) smallestPo := 256 var targets []enode.ID var closestPO int // loop through all nodes and add the message to recipient indices - for _, nod := range sim.Net.GetNodes() { - po, _ := pof(msgs[i], nodeAddrs[nod.ID()], 0) - depth := kademlias[nod.ID()].NeighbourhoodDepth() + for _, nod := range d.sim.Net.GetNodes() { + po, _ := pof(d.msgs[i], d.nodeAddrs[nod.ID()], 0) + depth := d.kademlias[nod.ID()].NeighbourhoodDepth() // only nodes with closest IDs (wrt msg) will receive the msg if po > closestPO { @@ -177,27 +162,27 @@ func assignTestVariables(sim *simulation.Simulation, msgCount int) { } if po >= depth { - allowedMessages++ - allowed[i] = append(allowed[i], nod.ID()) - allowedMsgs[nod.ID()] = append(allowedMsgs[nod.ID()], uint64(i)) + d.allowedMessages++ + d.allowed[i] = append(d.allowed[i], nod.ID()) + d.allowedMsgs[nod.ID()] = append(d.allowedMsgs[nod.ID()], uint64(i)) } // a node with the smallest PO (wrt msg) will be the sender if po < smallestPo { smallestPo = po - senders[i] = nod.ID() + d.senders[i] = nod.ID() } } - requiredMessages += len(targets) + d.requiredMessages += len(targets) for _, id := range targets { - recipients[i] = append(recipients[i], id) - expectedMsgs[id] = append(expectedMsgs[id], uint64(i)) + d.recipients[i] = append(d.recipients[i], id) + d.expectedMsgs[id] = append(d.expectedMsgs[id], uint64(i)) } - log.Debug("nn for msg", "targets", len(recipients[i]), "msgidx", i, "msg", common.Bytes2Hex(msgAddr[:8]), "sender", senders[i], "senderpo", smallestPo) + log.Debug("nn for msg", "targets", len(d.recipients[i]), "msgidx", i, "msg", common.Bytes2Hex(msgAddr[:8]), "sender", d.senders[i], "senderpo", smallestPo) } - log.Debug("msgs to receive", "count", requiredMessages) + log.Debug("msgs to receive", "count", d.requiredMessages) } func TestProxNetwork(t *testing.T) { @@ -221,33 +206,37 @@ func TestProxNetworkLong(t *testing.T) { // Upon sending the messages, it verifies that the respective message is passed to the message handlers of these recipients. // It will fail if a recipient handles a message it should not, or if after propagation not all expected messages are handled (timeout) func testProxNetwork(t *testing.T) { - resetTestVariables() + var tstdata testData msgCount, nodeCount := getCmdParams(t) handlerContextFuncs := make(map[Topic]handlerContextFunc) handlerContextFuncs[topic] = nodeMsgHandler - services := newProxServices(true, handlerContextFuncs, kademlias) - sim = simulation.New(services) - defer sim.Close() - err := sim.UploadSnapshot(fmt.Sprintf("testdata/snapshot_%d.json", nodeCount)) + tstdata.kademlias = make(map[enode.ID]*network.Kademlia) + services := newProxServices(&tstdata, true, handlerContextFuncs, tstdata.kademlias) + tstdata.sim = simulation.New(services) + defer tstdata.sim.Close() + err := tstdata.sim.UploadSnapshot(fmt.Sprintf("testdata/snapshot_%d.json", nodeCount)) if err != nil { t.Fatal(err) } - ctx, cancel := context.WithTimeout(context.Background(), time.Second*20) // todo: review + ctx, cancel := context.WithTimeout(context.Background(), time.Second*3) // todo: increase before commit defer cancel() snap := readSnapshot(t, nodeCount) - err = sim.WaitTillSnapshotRecreated(ctx, snap) + err = tstdata.sim.WaitTillSnapshotRecreated(ctx, snap) if err != nil { t.Fatalf("failed to recreate snapshot: %s", err) } - assignTestVariables(sim, msgCount) - result := sim.Run(ctx, runFunc) + initializeTestData(&tstdata, msgCount) + wrapper := func(c context.Context, s *simulation.Simulation) error { + return runFunc(&tstdata, c, s) + } + result := tstdata.sim.Run(ctx, wrapper) if result.Error != nil { // context deadline exceeded // however, it might just mean that not all possible messages are received // now we must check if all required messages are received - cnt := getMsgCount() + cnt := tstdata.getMsgCount() log.Debug("TestProxNetwork finnished", "rcv", cnt) - if cnt < requiredMessages { + if cnt < tstdata.requiredMessages { t.Fatal(result.Error) } } @@ -268,22 +257,22 @@ func sendAllMsgs(sim *simulation.Simulation, msgs [][]byte, senders map[int]enod log.Debug("all messages sent") } -func runFunc(ctx context.Context, sim *simulation.Simulation) error { - go handlerChannelListener(ctx) - go sendAllMsgs(sim, msgs, senders) +func runFunc(tstdata *testData, ctx context.Context, sim *simulation.Simulation) error { + go handlerChannelListener(tstdata, ctx) + go sendAllMsgs(sim, tstdata.msgs, tstdata.senders) received := 0 // collect incoming messages and terminate with corresponding status when message handler listener ends for { select { - case err := <-errC: + case err := <-tstdata.errC: return err - case hn := <-msgC: + case hn := <-tstdata.msgC: received++ - log.Debug("msg received", "msgs_received", received, "total_expected", requiredMessages, "id", hn.id, "serial", hn.serial) - if received == allowedMessages { - doneC <- struct{}{} - close(doneC) + log.Debug("msg received", "msgs_received", received, "total_expected", tstdata.requiredMessages, "id", hn.id, "serial", hn.serial) + if received == tstdata.allowedMessages { + tstdata.doneC <- struct{}{} + close(tstdata.doneC) return nil } } @@ -291,26 +280,26 @@ func runFunc(ctx context.Context, sim *simulation.Simulation) error { return nil } -func handlerChannelListener(ctx context.Context) { +func handlerChannelListener(tstdata *testData, ctx context.Context) { for { select { - case <-doneC: // graceful exit - setDone() - errC <- nil + case <-tstdata.doneC: // graceful exit + tstdata.setDone() + tstdata.errC <- nil return case <-ctx.Done(): // timeout or cancel - setDone() - errC <- ctx.Err() + tstdata.setDone() + tstdata.errC <- ctx.Err() return // incoming message from pss message handler - case handlerNotification := <-handlerC: + case handlerNotification := <-tstdata.handlerC: // check if recipient has already received all its messages and notify to fail the test if so - aMsgs := allowedMsgs[handlerNotification.id] + aMsgs := tstdata.allowedMsgs[handlerNotification.id] if len(aMsgs) == 0 { - setDone() - errC <- fmt.Errorf("too many messages received by recipient %x", handlerNotification.id) + tstdata.setDone() + tstdata.errC <- fmt.Errorf("too many messages received by recipient %x", handlerNotification.id) return } @@ -323,23 +312,23 @@ func handlerChannelListener(ctx context.Context) { } } if idx == -1 { - setDone() - errC <- fmt.Errorf("message %d received by wrong recipient %v", handlerNotification.serial, handlerNotification.id) + tstdata.setDone() + tstdata.errC <- fmt.Errorf("message %d received by wrong recipient %v", handlerNotification.serial, handlerNotification.id) return } // message is ok, so remove that message serial from the recipient expectation array and notify the main sim thread aMsgs[idx] = aMsgs[len(aMsgs)-1] aMsgs = aMsgs[:len(aMsgs)-1] - msgC <- handlerNotification + tstdata.msgC <- handlerNotification } } } -func nodeMsgHandler(config *adapters.NodeConfig) *handler { +func nodeMsgHandler(tstdata *testData, config *adapters.NodeConfig) *handler { return &handler{ f: func(msg []byte, p *p2p.Peer, asymmetric bool, keyid string) error { - cnt := incrementMsgCount() + cnt := tstdata.incrementMsgCount() log.Debug("nodeMsgHandler rcv", "cnt", cnt) // using simple serial in message body, makes it easy to keep track of who's getting what @@ -348,12 +337,12 @@ func nodeMsgHandler(config *adapters.NodeConfig) *handler { log.Crit(fmt.Sprintf("corrupt message received by %x (uvarint parse returned %d)", config.ID, c)) } - if isDone() { + if tstdata.isDone() { return errors.New("handlers aborted") // terminate if simulation is over } // pass message context to the listener in the simulation - handlerC <- handlerNotification{ + tstdata.handlerC <- handlerNotification{ id: config.ID, serial: serial, } @@ -368,7 +357,7 @@ func nodeMsgHandler(config *adapters.NodeConfig) *handler { // an adaptation of the same services setup as in pss_test.go // replaces pss_test.go when those tests are rewritten to the new swarm/network/simulation package -func newProxServices(allowRaw bool, handlerContextFuncs map[Topic]handlerContextFunc, kademlias map[enode.ID]*network.Kademlia) map[string]simulation.ServiceFunc { +func newProxServices(tstdata *testData, allowRaw bool, handlerContextFuncs map[Topic]handlerContextFunc, kademlias map[enode.ID]*network.Kademlia) map[string]simulation.ServiceFunc { stateStore := state.NewInmemoryStore() kademlia := func(id enode.ID) *network.Kademlia { if k, ok := kademlias[id]; ok { @@ -422,7 +411,7 @@ func newProxServices(allowRaw bool, handlerContextFuncs map[Topic]handlerContext // register the handlers we've been passed var deregisters []func() for tpc, hndlrFunc := range handlerContextFuncs { - deregisters = append(deregisters, ps.Register(&tpc, hndlrFunc(ctx.Config))) + deregisters = append(deregisters, ps.Register(&tpc, hndlrFunc(tstdata, ctx.Config))) } // if handshake mode is set, add the controller