From 714d871faa41f5bc6a8051d67b685dd507db01eb Mon Sep 17 00:00:00 2001 From: Janos Guljas Date: Tue, 12 Feb 2019 15:59:51 +0100 Subject: [PATCH] swarm/network/stream: add concurrent safe bool for watchDisconnections --- swarm/network/stream/common_test.go | 28 ++++++++++++++++++++-- swarm/network/stream/delivery_test.go | 12 ++++------ swarm/network/stream/intervals_test.go | 6 ++--- swarm/network/stream/snapshot_sync_test.go | 6 ++--- swarm/network/stream/syncer_test.go | 6 ++--- 5 files changed, 36 insertions(+), 22 deletions(-) diff --git a/swarm/network/stream/common_test.go b/swarm/network/stream/common_test.go index 7b9a62c704..afd08d2754 100644 --- a/swarm/network/stream/common_test.go +++ b/swarm/network/stream/common_test.go @@ -323,13 +323,14 @@ func createTestLocalStorageForID(id enode.ID, addr *network.BzzAddr) (storage.Ch // watchDisconnections receives simulation peer events in a new goroutine and sets atomic value // disconnected to true in case of a disconnect event. -func watchDisconnections(ctx context.Context, sim *simulation.Simulation) (disconnected atomic.Value) { +func watchDisconnections(ctx context.Context, sim *simulation.Simulation) (disconnected *boolean) { log.Debug("Watching for disconnections") disconnections := sim.PeerEvents( ctx, sim.NodeIDs(), simulation.NewPeerEventsFilter().Drop(), ) + disconnected = new(boolean) go func() { for { select { @@ -341,9 +342,32 @@ func watchDisconnections(ctx context.Context, sim *simulation.Simulation) (disco } else { log.Error("peer drop", "node", d.NodeID, "peer", d.PeerID) } - disconnected.Store(true) + disconnected.set(true) } } }() return disconnected } + +// boolean is used to concurrently set +// and read a boolean value. +type boolean struct { + v bool + mu sync.RWMutex +} + +// set sets the value. +func (b *boolean) set(v bool) { + b.mu.Lock() + defer b.mu.Unlock() + + b.v = v +} + +// bool reads the value. +func (b *boolean) bool() bool { + b.mu.RLock() + defer b.mu.RUnlock() + + return b.v +} diff --git a/swarm/network/stream/delivery_test.go b/swarm/network/stream/delivery_test.go index ce118e3254..e5821df4f9 100644 --- a/swarm/network/stream/delivery_test.go +++ b/swarm/network/stream/delivery_test.go @@ -550,10 +550,8 @@ func testDeliveryFromNodes(t *testing.T, nodes, chunkCount int, skipCheck bool) disconnected := watchDisconnections(ctx, sim) defer func() { - if err != nil { - if yes, ok := disconnected.Load().(bool); ok && yes { - err = errors.New("disconnect events received") - } + if err != nil && disconnected.bool() { + err = errors.New("disconnect events received") } }() @@ -667,10 +665,8 @@ func benchmarkDeliveryFromNodes(b *testing.B, nodes, chunkCount int, skipCheck b disconnected := watchDisconnections(ctx, sim) defer func() { - if err != nil { - if yes, ok := disconnected.Load().(bool); ok && yes { - err = errors.New("disconnect events received") - } + if err != nil && disconnected.bool() { + err = errors.New("disconnect events received") } }() // benchmark loop diff --git a/swarm/network/stream/intervals_test.go b/swarm/network/stream/intervals_test.go index 682cb9ddb5..009a941ef4 100644 --- a/swarm/network/stream/intervals_test.go +++ b/swarm/network/stream/intervals_test.go @@ -140,10 +140,8 @@ func testIntervals(t *testing.T, live bool, history *Range, skipCheck bool) { disconnected := watchDisconnections(ctx, sim) defer func() { - if err != nil { - if yes, ok := disconnected.Load().(bool); ok && yes { - err = errors.New("disconnect events received") - } + if err != nil && disconnected.bool() { + err = errors.New("disconnect events received") } }() diff --git a/swarm/network/stream/snapshot_sync_test.go b/swarm/network/stream/snapshot_sync_test.go index 8d9c5a879a..b45d0aed50 100644 --- a/swarm/network/stream/snapshot_sync_test.go +++ b/swarm/network/stream/snapshot_sync_test.go @@ -195,10 +195,8 @@ func runSim(conf *synctestConfig, ctx context.Context, sim *simulation.Simulatio return sim.Run(ctx, func(ctx context.Context, sim *simulation.Simulation) (err error) { disconnected := watchDisconnections(ctx, sim) defer func() { - if err != nil { - if yes, ok := disconnected.Load().(bool); ok && yes { - err = errors.New("disconnect events received") - } + if err != nil && disconnected.bool() { + err = errors.New("disconnect events received") } }() diff --git a/swarm/network/stream/syncer_test.go b/swarm/network/stream/syncer_test.go index 5d2d27120c..be0752a9d0 100644 --- a/swarm/network/stream/syncer_test.go +++ b/swarm/network/stream/syncer_test.go @@ -141,10 +141,8 @@ func testSyncBetweenNodes(t *testing.T, nodes, chunkCount int, skipCheck bool, p disconnected := watchDisconnections(ctx, sim) defer func() { - if err != nil { - if yes, ok := disconnected.Load().(bool); ok && yes { - err = errors.New("disconnect events received") - } + if err != nil && disconnected.bool() { + err = errors.New("disconnect events received") } }()