fix(dht/routing) make GetProviders respect context
This commit makes GetProviders (sync) respect the request context. It also amends all of GetProviders' callsites to pass a context in. This meant changing the signature of the dht's handlerfunc. I think I'll start referring to the request context as Vito Corleone. cc @whyrusleeping @jbenet License: MIT Signed-off-by: Brian Tiger Chow <brian@perfmode.com>
Brian Tiger Chow committed
Dec 2, 2014 at 00:40 UTC
de374226be0229a72e337b486588e792233ec80c
5 files changed
+22
-19
routing/dht/dht.go
+1
-1
@@ -161,7 +161,7 @@ func (dht *IpfsDHT) HandleMessage(ctx context.Context, mes msg.NetMessage) msg.N
161
}
162
163
// dispatch handler.
164
- rpmes, err := handler(mPeer, pmes)
164
+ rpmes, err := handler(ctx, mPeer, pmes)
165
if err != nil {
166
log.Errorf("handle message error: %s", err)
167
return nil
routing/dht/handlers.go
+12
-13
@@ -5,20 +5,19 @@ import (
5
"fmt"
6
"time"
7
8
- "github.com/jbenet/go-ipfs/Godeps/_workspace/src/code.google.com/p/goprotobuf/proto"
9
-
8
+ context "github.com/jbenet/go-ipfs/Godeps/_workspace/src/code.google.com/p/go.net/context"
9
+ proto "github.com/jbenet/go-ipfs/Godeps/_workspace/src/code.google.com/p/goprotobuf/proto"
10
+ ds "github.com/jbenet/go-ipfs/Godeps/_workspace/src/github.com/jbenet/go-datastore"
11
peer "github.com/jbenet/go-ipfs/peer"
12
pb "github.com/jbenet/go-ipfs/routing/dht/pb"
13
u "github.com/jbenet/go-ipfs/util"
13
-
14
- ds "github.com/jbenet/go-ipfs/Godeps/_workspace/src/github.com/jbenet/go-datastore"
14
)
15
16
// The number of closer peers to send on requests.
17
var CloserPeerCount = 4
18
19
// dhthandler specifies the signature of functions that handle DHT messages.
21
-type dhtHandler func(peer.Peer, *pb.Message) (*pb.Message, error)
20
+type dhtHandler func(context.Context, peer.Peer, *pb.Message) (*pb.Message, error)
21
22
func (dht *IpfsDHT) handlerForMsgType(t pb.Message_MessageType) dhtHandler {
23
switch t {
@@ -39,7 +38,7 @@ func (dht *IpfsDHT) handlerForMsgType(t pb.Message_MessageType) dhtHandler {
38
}
39
}
40
42
-func (dht *IpfsDHT) handleGetValue(p peer.Peer, pmes *pb.Message) (*pb.Message, error) {
41
+func (dht *IpfsDHT) handleGetValue(ctx context.Context, p peer.Peer, pmes *pb.Message) (*pb.Message, error) {
42
log.Debugf("%s handleGetValue for key: %s\n", dht.self, pmes.GetKey())
43
44
// setup response
@@ -85,7 +84,7 @@ func (dht *IpfsDHT) handleGetValue(p peer.Peer, pmes *pb.Message) (*pb.Message,
84
}
85
86
// if we know any providers for the requested value, return those.
88
- provs := dht.providers.GetProviders(u.Key(pmes.GetKey()))
87
+ provs := dht.providers.GetProviders(ctx, u.Key(pmes.GetKey()))
88
if len(provs) > 0 {
89
log.Debugf("handleGetValue returning %d provider[s]", len(provs))
90
resp.ProviderPeers = pb.PeersToPBPeers(provs)
@@ -107,7 +106,7 @@ func (dht *IpfsDHT) handleGetValue(p peer.Peer, pmes *pb.Message) (*pb.Message,
106
}
107
108
// Store a value in this peer local storage
110
-func (dht *IpfsDHT) handlePutValue(p peer.Peer, pmes *pb.Message) (*pb.Message, error) {
109
+func (dht *IpfsDHT) handlePutValue(ctx context.Context, p peer.Peer, pmes *pb.Message) (*pb.Message, error) {
110
dht.dslock.Lock()
111
defer dht.dslock.Unlock()
112
dskey := u.Key(pmes.GetKey()).DsKey()
@@ -129,12 +128,12 @@ func (dht *IpfsDHT) handlePutValue(p peer.Peer, pmes *pb.Message) (*pb.Message,
128
return pmes, err
129
}
130
132
-func (dht *IpfsDHT) handlePing(p peer.Peer, pmes *pb.Message) (*pb.Message, error) {
131
+func (dht *IpfsDHT) handlePing(_ context.Context, p peer.Peer, pmes *pb.Message) (*pb.Message, error) {
132
log.Debugf("%s Responding to ping from %s!\n", dht.self, p)
133
return pmes, nil
134
}
135
137
-func (dht *IpfsDHT) handleFindPeer(p peer.Peer, pmes *pb.Message) (*pb.Message, error) {
136
+func (dht *IpfsDHT) handleFindPeer(ctx context.Context, p peer.Peer, pmes *pb.Message) (*pb.Message, error) {
137
resp := pb.NewMessage(pmes.GetType(), "", pmes.GetClusterLevel())
138
var closest []peer.Peer
139
@@ -164,7 +163,7 @@ func (dht *IpfsDHT) handleFindPeer(p peer.Peer, pmes *pb.Message) (*pb.Message,
163
return resp, nil
164
}
165
167
-func (dht *IpfsDHT) handleGetProviders(p peer.Peer, pmes *pb.Message) (*pb.Message, error) {
166
+func (dht *IpfsDHT) handleGetProviders(ctx context.Context, p peer.Peer, pmes *pb.Message) (*pb.Message, error) {
167
resp := pb.NewMessage(pmes.GetType(), pmes.GetKey(), pmes.GetClusterLevel())
168
169
// check if we have this value, to add ourselves as provider.
@@ -177,7 +176,7 @@ func (dht *IpfsDHT) handleGetProviders(p peer.Peer, pmes *pb.Message) (*pb.Messa
176
}
177
178
// setup providers
180
- providers := dht.providers.GetProviders(u.Key(pmes.GetKey()))
179
+ providers := dht.providers.GetProviders(ctx, u.Key(pmes.GetKey()))
180
if has {
181
providers = append(providers, dht.self)
182
}
@@ -201,7 +200,7 @@ type providerInfo struct {
200
Value peer.Peer
201
}
202
204
-func (dht *IpfsDHT) handleAddProvider(p peer.Peer, pmes *pb.Message) (*pb.Message, error) {
203
+func (dht *IpfsDHT) handleAddProvider(ctx context.Context, p peer.Peer, pmes *pb.Message) (*pb.Message, error) {
204
key := u.Key(pmes.GetKey())
205
206
log.Debugf("%s adding %s as a provider for '%s'\n", dht.self, p, peer.ID(key))
routing/dht/providers.go
+7
-3
@@ -101,12 +101,16 @@ func (pm *ProviderManager) AddProvider(k u.Key, val peer.Peer) {
101
}
102
}
103
104
-func (pm *ProviderManager) GetProviders(k u.Key) []peer.Peer {
104
+func (pm *ProviderManager) GetProviders(ctx context.Context, k u.Key) []peer.Peer {
105
gp := new(getProv)
106
gp.k = k
107
gp.resp = make(chan []peer.Peer)
108
- pm.getprovs <- gp
109
- return <-gp.resp
108
+ select {
109
+ case pm.getprovs <- gp:
110
+ return <-gp.resp
111
+ case <-ctx.Done():
112
+ return nil
113
+ }
114
}
115
116
func (pm *ProviderManager) GetLocal() []u.Key {
routing/dht/providers_test.go
+1
-1
@@ -15,7 +15,7 @@ func TestProviderManager(t *testing.T) {
15
p := NewProviderManager(ctx, mid)
16
a := u.Key("test")
17
p.AddProvider(a, peer.WithIDString("testingprovider"))
18
- resp := p.GetProviders(a)
18
+ resp := p.GetProviders(ctx, a)
19
if len(resp) != 1 {
20
t.Fatal("Could not retrieve provider.")
21
}
routing/dht/routing.go
+1
-1
@@ -133,7 +133,7 @@ func (dht *IpfsDHT) FindProvidersAsync(ctx context.Context, key u.Key, count int
133
134
ps := newPeerSet()
135
// TODO may want to make this function async to hide latency
136
- provs := dht.providers.GetProviders(key)
136
+ provs := dht.providers.GetProviders(ctx, key)
137
for _, p := range provs {
138
count--
139
// NOTE: assuming that this list of peers is unique