dht_net.go 5.64 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
	for {
		// receive msg
		pmes := new(pb.Message)
		if err := r.ReadMsg(pmes); err != nil {
38
			s.Reset()
39 40 41 42 43 44 45 46 47 48
			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 {
49
			s.Reset()
50 51 52 53 54 55 56
			log.Debug("got back nil handler from handlerForMsgType")
			return
		}

		// dispatch handler.
		rpmes, err := handler(ctx, mPeer, pmes)
		if err != nil {
57
			s.Reset()
58 59 60 61 62 63
			log.Debugf("handle message error: %s", err)
			return
		}

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

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

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

81
	ms := dht.messageSenderForPeer(p)
82 83 84

	start := time.Now()

85 86
	rpmes, err := ms.SendRequest(ctx, pmes)
	if err != nil {
87 88 89
		return nil, err
	}

90 91 92
	// update the peer (on valid msgs only)
	dht.updateFromMessage(ctx, p, rpmes)

93
	dht.peerstore.RecordLatency(p, time.Since(start))
94 95 96
	log.Event(ctx, "dhtReceivedMessage", dht.self, p, rpmes)
	return rpmes, nil
}
97 98

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

101
	ms := dht.messageSenderForPeer(p)
102

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

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

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
135 136

	singleMes int
137 138 139 140 141 142 143 144 145 146 147
}

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
148
	nstr, err := ms.dht.host.NewStream(ms.dht.ctx, ms.p, ProtocolDHT, ProtocolDHTOld)
149 150 151 152
	if err != nil {
		return err
	}

153 154 155 156 157 158 159
	ms.r = ggio.NewDelimitedReader(nstr, inet.MessageSizeMax)
	ms.w = ggio.NewDelimitedWriter(nstr)
	ms.s = nstr

	return nil
}

160 161 162 163
// 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
164

165 166 167
func (ms *messageSender) SendMessage(ctx context.Context, pmes *pb.Message) error {
	ms.lk.Lock()
	defer ms.lk.Unlock()
168 169
	retry := false
	for {
Jeromy's avatar
Jeromy committed
170 171 172 173 174
		if err := ms.prep(); err != nil {
			return err
		}

		if err := ms.w.WriteMsg(pmes); err != nil {
175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194
			ms.s.Reset()
			ms.s = nil

			if retry {
				log.Info("error writing message, bailing: ", err)
				return err
			} else {
				log.Info("error writing message, trying again: ", err)
				retry = true
				continue
			}
		}

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

		if ms.singleMes > streamReuseTries {
			ms.s.Close()
			ms.s = nil
		} else if retry {
			ms.singleMes++
Jeromy's avatar
Jeromy committed
195 196
		}

197
		return nil
198
	}
199
}
200 201 202 203

func (ms *messageSender) SendRequest(ctx context.Context, pmes *pb.Message) (*pb.Message, error) {
	ms.lk.Lock()
	defer ms.lk.Unlock()
204 205 206 207 208
	retry := false
	for {
		if err := ms.prep(); err != nil {
			return nil, err
		}
209

210 211 212 213 214 215 216 217 218 219 220 221 222
		if err := ms.w.WriteMsg(pmes); err != nil {
			ms.s.Reset()
			ms.s = nil

			if retry {
				log.Info("error writing message, bailing: ", err)
				return nil, err
			} else {
				log.Info("error writing message, trying again: ", err)
				retry = true
				continue
			}
		}
223

224 225 226 227 228 229 230 231 232 233 234 235 236 237
		mes := new(pb.Message)
		if err := ms.ctxReadMsg(ctx, mes); err != nil {
			ms.s.Reset()
			ms.s = nil

			if retry {
				log.Info("error reading message, bailing: ", err)
				return nil, err
			} else {
				log.Info("error reading message, trying again: ", err)
				retry = true
				continue
			}
		}
238

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

241 242 243 244 245 246
		if ms.singleMes > streamReuseTries {
			ms.s.Close()
			ms.s = nil
		} else if retry {
			ms.singleMes++
		}
Jeromy's avatar
Jeromy committed
247

248 249
		return mes, nil
	}
250
}
251 252 253

func (ms *messageSender) ctxReadMsg(ctx context.Context, mes *pb.Message) error {
	errc := make(chan error, 1)
254 255 256
	go func(r ggio.ReadCloser) {
		errc <- r.ReadMsg(mes)
	}(ms.r)
257

258 259 260
	t := time.NewTimer(dhtReadMessageTimeout)
	defer t.Stop()

261 262 263 264 265
	select {
	case err := <-errc:
		return err
	case <-ctx.Done():
		return ctx.Err()
266 267
	case <-t.C:
		return ErrReadTimeout
268 269
	}
}