fixed muxer errors
Juan Batiz-Benet committed
Sep 26, 2014 at 03:57 UTC
c1219303a092161960094dbb666946f1f28d14d0
2 files changed
+36
-20
net/mux/mux.go
+30
-14
@@ -2,6 +2,7 @@ package mux
2
3
import (
4
"errors"
5
+ "sync"
6
7
msg "github.com/jbenet/go-ipfs/net/message"
8
u "github.com/jbenet/go-ipfs/util"
@@ -30,6 +31,8 @@ type Muxer struct {
31
32
// cancel is the function to stop the Muxer
33
cancel context.CancelFunc
34
+ ctx context.Context
35
+ wg sync.WaitGroup
36
37
*msg.Pipe
38
}
@@ -58,11 +61,14 @@ func (m *Muxer) Start(ctx context.Context) error {
61
}
62
63
// make a cancellable context.
61
- ctx, m.cancel = context.WithCancel(ctx)
64
+ m.ctx, m.cancel = context.WithCancel(ctx)
65
+ m.wg = sync.WaitGroup{}
66
63
- go m.handleIncomingMessages(ctx)
67
+ m.wg.Add(1)
68
+ go m.handleIncomingMessages()
69
for pid, proto := range m.Protocols {
65
- go m.handleOutgoingMessages(ctx, pid, proto)
70
+ m.wg.Add(1)
71
+ go m.handleOutgoingMessages(pid, proto)
72
}
73
74
return nil
@@ -70,8 +76,15 @@ func (m *Muxer) Start(ctx context.Context) error {
76
77
// Stop stops muxer activity.
78
func (m *Muxer) Stop() {
79
+ if m.cancel == nil {
80
+ panic("muxer stopped twice.")
81
+ }
82
+ // issue cancel, and wipe func.
83
m.cancel()
84
m.cancel = context.CancelFunc(nil)
85
+
86
+ // wait for everything to wind down.
87
+ m.wg.Wait()
88
}
89
90
// AddProtocol adds a Protocol with given ProtocolID to the Muxer.
@@ -86,7 +99,8 @@ func (m *Muxer) AddProtocol(p Protocol, pid ProtocolID) error {
99
100
// handleIncoming consumes the messages on the m.Incoming channel and
101
// routes them appropriately (to the protocols).
89
-func (m *Muxer) handleIncomingMessages(ctx context.Context) {
102
+func (m *Muxer) handleIncomingMessages() {
103
+ defer m.wg.Done()
104
105
for {
106
if m == nil {
@@ -98,16 +112,16 @@ func (m *Muxer) handleIncomingMessages(ctx context.Context) {
112
if !more {
113
return
114
}
101
- go m.handleIncomingMessage(ctx, msg)
115
+ go m.handleIncomingMessage(msg)
116
103
- case <-ctx.Done():
117
+ case <-m.ctx.Done():
118
return
119
}
120
}
121
}
122
123
// handleIncomingMessage routes message to the appropriate protocol.
110
-func (m *Muxer) handleIncomingMessage(ctx context.Context, m1 msg.NetMessage) {
124
+func (m *Muxer) handleIncomingMessage(m1 msg.NetMessage) {
125
126
data, pid, err := unwrapData(m1.Data())
127
if err != nil {
@@ -124,31 +138,33 @@ func (m *Muxer) handleIncomingMessage(ctx context.Context, m1 msg.NetMessage) {
138
139
select {
140
case proto.GetPipe().Incoming <- m2:
127
- case <-ctx.Done():
128
- u.PErr("%v\n", ctx.Err())
141
+ case <-m.ctx.Done():
142
+ u.PErr("%v\n", m.ctx.Err())
143
return
144
}
145
}
146
147
// handleOutgoingMessages consumes the messages on the proto.Outgoing channel,
148
// wraps them and sends them out.
135
-func (m *Muxer) handleOutgoingMessages(ctx context.Context, pid ProtocolID, proto Protocol) {
149
+func (m *Muxer) handleOutgoingMessages(pid ProtocolID, proto Protocol) {
150
+ defer m.wg.Done()
151
+
152
for {
153
select {
154
case msg, more := <-proto.GetPipe().Outgoing:
155
if !more {
156
return
157
}
142
- go m.handleOutgoingMessage(ctx, pid, msg)
158
+ go m.handleOutgoingMessage(pid, msg)
159
144
- case <-ctx.Done():
160
+ case <-m.ctx.Done():
161
return
162
}
163
}
164
}
165
166
// handleOutgoingMessage wraps out a message and sends it out the
151
-func (m *Muxer) handleOutgoingMessage(ctx context.Context, pid ProtocolID, m1 msg.NetMessage) {
167
+func (m *Muxer) handleOutgoingMessage(pid ProtocolID, m1 msg.NetMessage) {
168
data, err := wrapData(m1.Data(), pid)
169
if err != nil {
170
u.PErr("muxer serializing error: %v\n", err)
@@ -158,7 +174,7 @@ func (m *Muxer) handleOutgoingMessage(ctx context.Context, pid ProtocolID, m1 ms
174
m2 := msg.New(m1.Peer(), data)
175
select {
176
case m.GetPipe().Outgoing <- m2:
161
- case <-ctx.Done():
177
+ case <-m.ctx.Done():
178
return
179
}
180
}
net/mux/mux_test.go
+6
-6
@@ -229,13 +229,13 @@ func TestStopping(t *testing.T) {
229
mux1.Start(context.Background())
230
231
// test outgoing p1
232
- for _, s := range []string{"foo", "bar", "baz"} {
232
+ for _, s := range []string{"foo1", "bar1", "baz1"} {
233
p1.Outgoing <- msg.New(peer1, []byte(s))
234
testWrappedMsg(t, <-mux1.Outgoing, pid1, []byte(s))
235
}
236
237
// test incoming p1
238
- for _, s := range []string{"foo", "bar", "baz"} {
238
+ for _, s := range []string{"foo2", "bar2", "baz2"} {
239
d, err := wrapData([]byte(s), pid1)
240
if err != nil {
241
t.Error(err)
@@ -250,17 +250,17 @@ func TestStopping(t *testing.T) {
250
}
251
252
// test outgoing p1
253
- for _, s := range []string{"foo", "bar", "baz"} {
253
+ for _, s := range []string{"foo3", "bar3", "baz3"} {
254
p1.Outgoing <- msg.New(peer1, []byte(s))
255
select {
256
- case <-mux1.Outgoing:
257
- t.Error("should not have received anything.")
256
+ case m := <-mux1.Outgoing:
257
+ t.Errorf("should not have received anything. Got: %v", string(m.Data()))
258
case <-time.After(time.Millisecond):
259
}
260
}
261
262
// test incoming p1
263
- for _, s := range []string{"foo", "bar", "baz"} {
263
+ for _, s := range []string{"foo4", "bar4", "baz4"} {
264
d, err := wrapData([]byte(s), pid1)
265
if err != nil {
266
t.Error(err)