dht_net.go 5.44 KB
Newer Older
1 2 3
package dht

import (
Jeromy's avatar
Jeromy committed
4
	"context"
5
	"fmt"
6
	"sync"
7 8
	"time"

9 10
	ggio "github.com/gogo/protobuf/io"
	ctxio "github.com/jbenet/go-context/io"
George Antoniadis's avatar
George Antoniadis committed
11
	pb "github.com/libp2p/go-libp2p-kad-dht/pb"
12 13
	inet "github.com/libp2p/go-libp2p-net"
	peer "github.com/libp2p/go-libp2p-peer"
14 15
)

16 17 18
var dhtReadMessageTimeout = time.Minute
var ErrReadTimeout = fmt.Errorf("timed out reading response")

19 20 21 22 23 24 25 26 27
// handleNewStream implements the inet.StreamHandler
func (dht *IpfsDHT) handleNewStream(s inet.Stream) {
	go dht.handleNewMessage(s)
}

func (dht *IpfsDHT) handleNewMessage(s inet.Stream) {
	defer s.Close()

	ctx := dht.Context()
28 29
	cr := ctxio.NewReader(ctx, s) // ok to use. we defer close stream in this func
	cw := ctxio.NewWriter(ctx, s) // ok to use. we defer close stream in this func
30 31
	r := ggio.NewDelimitedReader(cr, inet.MessageSizeMax)
	w := ggio.NewDelimitedWriter(cw)
32 33
	mPeer := s.Conn().RemotePeer()

34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60
	for {
		// receive msg
		pmes := new(pb.Message)
		if err := r.ReadMsg(pmes); err != nil {
			log.Debugf("Error unmarshaling data: %s", err)
			return
		}

		// update the peer (on valid msgs only)
		dht.updateFromMessage(ctx, mPeer, pmes)

		// get handler for this msg type.
		handler := dht.handlerForMsgType(pmes.GetType())
		if handler == nil {
			log.Debug("got back nil handler from handlerForMsgType")
			return
		}

		// dispatch handler.
		rpmes, err := handler(ctx, mPeer, pmes)
		if err != nil {
			log.Debugf("handle message error: %s", err)
			return
		}

		// if nil response, return it before serializing
		if rpmes == nil {
61
			log.Debug("got back nil response from request")
62 63 64 65 66 67 68 69
			continue
		}

		// send out response msg
		if err := w.WriteMsg(rpmes); err != nil {
			log.Debugf("send response error: %s", err)
			return
		}
70 71 72 73 74
	}
}

// sendRequest sends out a request, but also makes sure to
// measure the RTT for latency measurements.
75
func (dht *IpfsDHT) sendRequest(ctx context.Context, p peer.ID, pmes *pb.Message) (*pb.Message, error) {
76

77
	ms := dht.messageSenderForPeer(p)
78 79 80

	start := time.Now()

81 82
	rpmes, err := ms.SendRequest(ctx, pmes)
	if err != nil {
83 84 85
		return nil, err
	}

86 87 88
	// update the peer (on valid msgs only)
	dht.updateFromMessage(ctx, p, rpmes)

89
	dht.peerstore.RecordLatency(p, time.Since(start))
90 91 92
	log.Event(ctx, "dhtReceivedMessage", dht.self, p, rpmes)
	return rpmes, nil
}
93 94

// sendMessage sends out a message
95
func (dht *IpfsDHT) sendMessage(ctx context.Context, p peer.ID, pmes *pb.Message) error {
96

97
	ms := dht.messageSenderForPeer(p)
98

99
	if err := ms.SendMessage(ctx, pmes); err != nil {
100 101 102 103 104
		return err
	}
	log.Event(ctx, "dhtSentMessage", dht.self, p, pmes)
	return nil
}
105 106 107 108 109

func (dht *IpfsDHT) updateFromMessage(ctx context.Context, p peer.ID, mes *pb.Message) error {
	dht.Update(ctx, p)
	return nil
}
110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130

func (dht *IpfsDHT) messageSenderForPeer(p peer.ID) *messageSender {
	dht.smlk.Lock()
	defer dht.smlk.Unlock()

	ms, ok := dht.strmap[p]
	if !ok {
		ms = dht.newMessageSender(p)
		dht.strmap[p] = ms
	}

	return ms
}

type messageSender struct {
	s   inet.Stream
	r   ggio.ReadCloser
	w   ggio.WriteCloser
	lk  sync.Mutex
	p   peer.ID
	dht *IpfsDHT
Jeromy's avatar
Jeromy committed
131 132

	singleMes int
133 134 135 136 137 138 139 140 141 142 143
}

func (dht *IpfsDHT) newMessageSender(p peer.ID) *messageSender {
	return &messageSender{p: p, dht: dht}
}

func (ms *messageSender) prep() error {
	if ms.s != nil {
		return nil
	}

Jeromy's avatar
Jeromy committed
144
	nstr, err := ms.dht.host.NewStream(ms.dht.ctx, ms.p, ProtocolDHT, ProtocolDHTOld)
145 146 147 148
	if err != nil {
		return err
	}

149 150 151 152 153 154 155
	ms.r = ggio.NewDelimitedReader(nstr, inet.MessageSizeMax)
	ms.w = ggio.NewDelimitedWriter(nstr)
	ms.s = nstr

	return nil
}

156 157 158 159
// streamReuseTries is the number of times we will try to reuse a stream to a
// given peer before giving up and reverting to the old one-message-per-stream
// behaviour.
const streamReuseTries = 3
160

161 162 163 164
func (ms *messageSender) SendMessage(ctx context.Context, pmes *pb.Message) error {
	ms.lk.Lock()
	defer ms.lk.Unlock()
	if err := ms.prep(); err != nil {
165 166
		return err
	}
167

Jeromy's avatar
Jeromy committed
168 169 170 171
	if err := ms.writeMessage(pmes); err != nil {
		return err
	}

172
	if ms.singleMes > streamReuseTries {
Jeromy's avatar
Jeromy committed
173 174 175 176
		ms.s.Close()
		ms.s = nil
	}

177 178
	return nil
}
179

Jeromy's avatar
Jeromy committed
180
func (ms *messageSender) writeMessage(pmes *pb.Message) error {
181 182
	err := ms.w.WriteMsg(pmes)
	if err != nil {
Jeromy's avatar
Jeromy committed
183 184 185 186 187
		// If the other side isnt expecting us to be reusing streams, we're gonna
		// end up erroring here. To make sure things work seamlessly, lets retry once
		// before continuing

		log.Infof("error writing message: ", err)
188 189
		ms.s.Close()
		ms.s = nil
Jeromy's avatar
Jeromy committed
190 191 192 193 194 195 196 197 198 199 200
		if err := ms.prep(); err != nil {
			return err
		}

		if err := ms.w.WriteMsg(pmes); err != nil {
			return err
		}

		// keep track of this happening. If it happens a few times, its
		// likely we can assume the otherside will never support stream reuse
		ms.singleMes++
201
	}
202 203
	return nil
}
204 205 206 207 208 209 210 211

func (ms *messageSender) SendRequest(ctx context.Context, pmes *pb.Message) (*pb.Message, error) {
	ms.lk.Lock()
	defer ms.lk.Unlock()
	if err := ms.prep(); err != nil {
		return nil, err
	}

Jeromy's avatar
Jeromy committed
212
	if err := ms.writeMessage(pmes); err != nil {
213 214 215 216 217 218
		return nil, err
	}

	log.Event(ctx, "dhtSentMessage", ms.dht.self, ms.p, pmes)

	mes := new(pb.Message)
219
	if err := ms.ctxReadMsg(ctx, mes); err != nil {
Steven Allen's avatar
Steven Allen committed
220
		ms.s.Reset()
221 222 223 224
		ms.s = nil
		return nil, err
	}

225
	if ms.singleMes > streamReuseTries {
Jeromy's avatar
Jeromy committed
226 227 228 229
		ms.s.Close()
		ms.s = nil
	}

230 231
	return mes, nil
}
232 233 234

func (ms *messageSender) ctxReadMsg(ctx context.Context, mes *pb.Message) error {
	errc := make(chan error, 1)
235 236 237
	go func(r ggio.ReadCloser) {
		errc <- r.ReadMsg(mes)
	}(ms.r)
238

239 240 241
	t := time.NewTimer(dhtReadMessageTimeout)
	defer t.Stop()

242 243 244 245 246
	select {
	case err := <-errc:
		return err
	case <-ctx.Done():
		return ctx.Err()
247 248
	case <-t.C:
		return ErrReadTimeout
249 250
	}
}