895 lines
22 KiB
Go
895 lines
22 KiB
Go
package torrent
|
||
|
||
import (
|
||
"bytes"
|
||
"context"
|
||
"encoding/binary"
|
||
"errors"
|
||
"fmt"
|
||
"io"
|
||
"log"
|
||
"net"
|
||
"sync"
|
||
"sync/atomic"
|
||
"time"
|
||
|
||
"github.com/jackpal/bencode-go"
|
||
"golang.org/x/time/rate"
|
||
|
||
"github.com/veggiedefender/torrent-client/internal/torrentfile"
|
||
"github.com/veggiedefender/torrent-client/internal/tracker"
|
||
)
|
||
|
||
const (
|
||
wireProtocolString = "BitTorrent protocol"
|
||
|
||
msgChoke = 0
|
||
msgUnchoke = 1
|
||
msgInterested = 2
|
||
msgNotInterested = 3
|
||
msgHave = 4
|
||
msgBitfield = 5
|
||
msgRequest = 6
|
||
msgPiece = 7
|
||
msgCancel = 8
|
||
msgExtended = 20
|
||
|
||
maxWireMessageSize = 2 * 1024 * 1024
|
||
requestBlockSize = 16 * 1024
|
||
|
||
// Adaptive pipeline bounds
|
||
minPipelineDepth = 4
|
||
maxPipelineDepth = 64
|
||
initPipelineDepth = 8
|
||
|
||
keepaliveInterval = 90 * time.Second
|
||
|
||
peerConnectTimeout = 5 * time.Second
|
||
peerReadTimeout = 15 * time.Second
|
||
peerWriteTimeout = 10 * time.Second
|
||
|
||
peerSocketReadBuffer = 512 * 1024
|
||
peerSocketWriteBuffer = 512 * 1024
|
||
)
|
||
|
||
type extendedHandshake struct {
|
||
M map[string]int `bencode:"m"`
|
||
MetadataSize int `bencode:"metadata_size"`
|
||
}
|
||
|
||
type pexMessage struct {
|
||
Added string `bencode:"added,omitempty"`
|
||
Added6 string `bencode:"added6,omitempty"`
|
||
}
|
||
|
||
type wireMessage struct {
|
||
ID int
|
||
Payload []byte
|
||
}
|
||
|
||
type pieceTransferStats struct {
|
||
DownloadedBytes int64
|
||
UploadedBytes int64
|
||
AvgBlockLatency time.Duration
|
||
Blocks int
|
||
Duration time.Duration
|
||
}
|
||
|
||
type measuredConn struct {
|
||
net.Conn
|
||
readBytes atomic.Int64
|
||
writeBytes atomic.Int64
|
||
}
|
||
|
||
func (c *measuredConn) Read(p []byte) (int, error) {
|
||
n, err := c.Conn.Read(p)
|
||
if n > 0 {
|
||
c.readBytes.Add(int64(n))
|
||
}
|
||
return n, err
|
||
}
|
||
|
||
func (c *measuredConn) Write(p []byte) (int, error) {
|
||
n, err := c.Conn.Write(p)
|
||
if n > 0 {
|
||
c.writeBytes.Add(int64(n))
|
||
}
|
||
return n, err
|
||
}
|
||
|
||
func (c *measuredConn) Snapshot() (readBytes int64, writeBytes int64) {
|
||
return c.readBytes.Load(), c.writeBytes.Load()
|
||
}
|
||
|
||
// rateLimitedConn оборачивает measuredConn и применяет token-bucket rate limiting.
|
||
type rateLimitedConn struct {
|
||
*measuredConn
|
||
downLimiter *rate.Limiter
|
||
upLimiter *rate.Limiter
|
||
mu sync.RWMutex
|
||
}
|
||
|
||
func newRateLimitedConn(inner *measuredConn) *rateLimitedConn {
|
||
return &rateLimitedConn{
|
||
measuredConn: inner,
|
||
downLimiter: rate.NewLimiter(rate.Inf, 0),
|
||
upLimiter: rate.NewLimiter(rate.Inf, 0),
|
||
}
|
||
}
|
||
|
||
func (c *rateLimitedConn) SetDownloadLimit(bps int64) {
|
||
c.mu.Lock()
|
||
defer c.mu.Unlock()
|
||
if bps <= 0 {
|
||
c.downLimiter.SetLimit(rate.Inf)
|
||
c.downLimiter.SetBurst(0)
|
||
} else {
|
||
c.downLimiter.SetLimit(rate.Limit(bps))
|
||
c.downLimiter.SetBurst(int(bps))
|
||
}
|
||
}
|
||
|
||
func (c *rateLimitedConn) SetUploadLimit(bps int64) {
|
||
c.mu.Lock()
|
||
defer c.mu.Unlock()
|
||
if bps <= 0 {
|
||
c.upLimiter.SetLimit(rate.Inf)
|
||
c.upLimiter.SetBurst(0)
|
||
} else {
|
||
c.upLimiter.SetLimit(rate.Limit(bps))
|
||
c.upLimiter.SetBurst(int(bps))
|
||
}
|
||
}
|
||
|
||
func (c *rateLimitedConn) Read(p []byte) (int, error) {
|
||
n, err := c.measuredConn.Read(p)
|
||
if n > 0 {
|
||
c.mu.RLock()
|
||
lim := c.downLimiter
|
||
c.mu.RUnlock()
|
||
if lim.Limit() != rate.Inf {
|
||
_ = lim.WaitN(context.Background(), min(n, lim.Burst()))
|
||
}
|
||
}
|
||
return n, err
|
||
}
|
||
|
||
func (c *rateLimitedConn) Write(p []byte) (int, error) {
|
||
if len(p) > 0 {
|
||
c.mu.RLock()
|
||
lim := c.upLimiter
|
||
c.mu.RUnlock()
|
||
if lim.Limit() != rate.Inf {
|
||
_ = lim.WaitN(context.Background(), min(len(p), lim.Burst()))
|
||
}
|
||
}
|
||
return c.measuredConn.Write(p)
|
||
}
|
||
|
||
func min(a, b int) int {
|
||
if a < b {
|
||
return a
|
||
}
|
||
return b
|
||
}
|
||
|
||
type peerClient struct {
|
||
conn *measuredConn
|
||
rlConn *rateLimitedConn
|
||
|
||
have []bool
|
||
hasPieceInfo bool
|
||
peerIsChoked bool
|
||
|
||
supportsExtensions bool
|
||
peerUtMetadataID int
|
||
metadataSize int
|
||
peerUtPexID int
|
||
OnPex func([]tracker.Peer)
|
||
|
||
// keepalive
|
||
kaStop chan struct{}
|
||
kaOnce sync.Once
|
||
}
|
||
|
||
|
||
func newPeerClient(ctx context.Context, addr string, infoHash [20]byte, peerID [20]byte, pieceCount int) (*peerClient, error) {
|
||
dialer := net.Dialer{Timeout: peerConnectTimeout}
|
||
rawConn, err := dialer.DialContext(ctx, "tcp", addr)
|
||
if err != nil {
|
||
log.Printf("peer dial failed %s: %v", addr, err)
|
||
return nil, err
|
||
}
|
||
log.Printf("connected to peer %s", addr)
|
||
|
||
if tcpConn, ok := rawConn.(*net.TCPConn); ok {
|
||
_ = tcpConn.SetNoDelay(true)
|
||
_ = tcpConn.SetReadBuffer(peerSocketReadBuffer)
|
||
_ = tcpConn.SetWriteBuffer(peerSocketWriteBuffer)
|
||
}
|
||
|
||
measured := &measuredConn{Conn: rawConn}
|
||
rlConn := newRateLimitedConn(measured)
|
||
pc := &peerClient{
|
||
conn: measured,
|
||
rlConn: rlConn,
|
||
have: make([]bool, pieceCount),
|
||
peerIsChoked: true,
|
||
kaStop: make(chan struct{}),
|
||
}
|
||
|
||
if err := pc.sendHandshake(ctx, infoHash, peerID); err != nil {
|
||
measured.Close()
|
||
return nil, err
|
||
}
|
||
if err := pc.readHandshake(ctx, infoHash); err != nil {
|
||
measured.Close()
|
||
return nil, err
|
||
}
|
||
|
||
if pc.supportsExtensions {
|
||
if err := pc.sendExtendedHandshake(ctx); err != nil {
|
||
log.Printf("failed to send extended handshake to %s: %v", addr, err)
|
||
}
|
||
}
|
||
|
||
if err := pc.sendMessage(ctx, msgInterested, nil); err != nil {
|
||
measured.Close()
|
||
return nil, err
|
||
}
|
||
debugf("sent interested to %s", addr)
|
||
|
||
if err := pc.readInitialMessages(ctx); err != nil {
|
||
log.Printf("peer %s initial message read failed: %v", addr, err)
|
||
measured.Close()
|
||
return nil, err
|
||
}
|
||
|
||
go pc.keepaliveLoop()
|
||
return pc, nil
|
||
}
|
||
|
||
|
||
func (pc *peerClient) PeerUtPexID() int {
|
||
return pc.peerUtPexID
|
||
}
|
||
|
||
func newIncomingPeerClient(ctx context.Context, rawConn net.Conn, extensions bool, pieceCount int) (*peerClient, error) {
|
||
measured := &measuredConn{Conn: rawConn}
|
||
rlConn := newRateLimitedConn(measured)
|
||
pc := &peerClient{
|
||
conn: measured,
|
||
rlConn: rlConn,
|
||
have: make([]bool, pieceCount),
|
||
peerIsChoked: true,
|
||
kaStop: make(chan struct{}),
|
||
supportsExtensions: extensions,
|
||
}
|
||
|
||
if pc.supportsExtensions {
|
||
if err := pc.sendExtendedHandshake(ctx); err != nil {
|
||
log.Printf("failed to send extended handshake to incoming peer: %v", err)
|
||
}
|
||
}
|
||
|
||
if err := pc.sendMessage(ctx, msgInterested, nil); err != nil {
|
||
measured.Close()
|
||
return nil, err
|
||
}
|
||
|
||
if err := pc.readInitialMessages(ctx); err != nil {
|
||
measured.Close()
|
||
return nil, err
|
||
}
|
||
|
||
go pc.keepaliveLoop()
|
||
return pc, nil
|
||
}
|
||
|
||
func (pc *peerClient) Close() error {
|
||
pc.kaOnce.Do(func() { close(pc.kaStop) })
|
||
return pc.conn.Close()
|
||
}
|
||
|
||
// keepaliveLoop отправляет keepalive (каждые 90с) пока соединение активно.
|
||
func (pc *peerClient) keepaliveLoop() {
|
||
ticker := time.NewTicker(keepaliveInterval)
|
||
defer ticker.Stop()
|
||
for {
|
||
select {
|
||
case <-pc.kaStop:
|
||
return
|
||
case <-ticker.C:
|
||
// keepalive: send length-prefix 0 (no message ID)
|
||
_ = pc.conn.SetWriteDeadline(time.Now().Add(peerWriteTimeout))
|
||
_, _ = pc.conn.Write([]byte{0, 0, 0, 0})
|
||
debugf("sent keepalive")
|
||
}
|
||
}
|
||
}
|
||
|
||
|
||
func (pc *peerClient) PieceAvailability() ([]bool, bool) {
|
||
have := make([]bool, len(pc.have))
|
||
copy(have, pc.have)
|
||
return have, pc.hasPieceInfo
|
||
}
|
||
|
||
func (pc *peerClient) DownloadPiece(ctx context.Context, pieceIndex int, pieceLength int) ([]byte, pieceTransferStats, error) {
|
||
var transfer pieceTransferStats
|
||
if pieceLength <= 0 {
|
||
return nil, transfer, fmt.Errorf("invalid piece length %d", pieceLength)
|
||
}
|
||
|
||
if err := pc.waitForUnchoke(ctx); err != nil {
|
||
return nil, transfer, err
|
||
}
|
||
|
||
startReadBytes, startWriteBytes := pc.conn.Snapshot()
|
||
startedAt := time.Now()
|
||
|
||
piece := make([]byte, pieceLength)
|
||
type pendingRequest struct {
|
||
length int
|
||
requestedAt time.Time
|
||
}
|
||
|
||
pending := make(map[int]pendingRequest, maxPipelineDepth)
|
||
offset := 0
|
||
received := 0
|
||
var latencySum time.Duration
|
||
blocksCompleted := 0
|
||
depth := initPipelineDepth // adaptive, пересчитывается каждые 4 блока
|
||
|
||
for received < pieceLength {
|
||
for len(pending) < depth && offset < pieceLength {
|
||
blockLength := requestBlockSize
|
||
if remaining := pieceLength - offset; remaining < blockLength {
|
||
blockLength = remaining
|
||
}
|
||
|
||
if err := pc.sendRequest(ctx, pieceIndex, offset, blockLength); err != nil {
|
||
return nil, transfer, err
|
||
}
|
||
pending[offset] = pendingRequest{
|
||
length: blockLength,
|
||
requestedAt: time.Now(),
|
||
}
|
||
offset += blockLength
|
||
}
|
||
|
||
gotIndex, gotBegin, block, err := pc.readPieceMessage(ctx)
|
||
if err != nil {
|
||
return nil, transfer, err
|
||
}
|
||
if gotIndex != pieceIndex {
|
||
continue
|
||
}
|
||
|
||
req, ok := pending[gotBegin]
|
||
if !ok {
|
||
continue
|
||
}
|
||
if len(block) != req.length {
|
||
return nil, transfer, fmt.Errorf("unexpected block size for piece=%d begin=%d: got=%d want=%d", pieceIndex, gotBegin, len(block), req.length)
|
||
}
|
||
if gotBegin < 0 || gotBegin+len(block) > len(piece) {
|
||
return nil, transfer, fmt.Errorf("piece block bounds invalid for piece=%d begin=%d block=%d", pieceIndex, gotBegin, len(block))
|
||
}
|
||
|
||
copy(piece[gotBegin:gotBegin+len(block)], block)
|
||
delete(pending, gotBegin)
|
||
received += len(block)
|
||
|
||
blocksCompleted++
|
||
blockLatency := time.Since(req.requestedAt)
|
||
latencySum += blockLatency
|
||
|
||
// Пересчёт adaptive depth каждые 4 блока
|
||
if blocksCompleted%4 == 0 && blockLatency > 0 {
|
||
elapsedSec := blockLatency.Seconds()
|
||
curReadBytes, _ := pc.conn.Snapshot()
|
||
bytesSoFar := curReadBytes - startReadBytes
|
||
if elapsedSec > 0 && bytesSoFar > 0 && time.Since(startedAt).Seconds() > 0 {
|
||
bwBps := float64(bytesSoFar) / time.Since(startedAt).Seconds()
|
||
newDepth := int(bwBps * elapsedSec / float64(requestBlockSize))
|
||
if newDepth < minPipelineDepth {
|
||
newDepth = minPipelineDepth
|
||
}
|
||
if newDepth > maxPipelineDepth {
|
||
newDepth = maxPipelineDepth
|
||
}
|
||
depth = newDepth
|
||
}
|
||
}
|
||
|
||
}
|
||
|
||
endReadBytes, endWriteBytes := pc.conn.Snapshot()
|
||
transfer.DownloadedBytes = maxInt64(0, endReadBytes-startReadBytes)
|
||
transfer.UploadedBytes = maxInt64(0, endWriteBytes-startWriteBytes)
|
||
transfer.Blocks = blocksCompleted
|
||
transfer.Duration = time.Since(startedAt)
|
||
if blocksCompleted > 0 {
|
||
transfer.AvgBlockLatency = latencySum / time.Duration(blocksCompleted)
|
||
}
|
||
|
||
return piece, transfer, nil
|
||
}
|
||
|
||
// SendCancel отправляет сообщение cancel пиру (используется в endgame-режиме).
|
||
func (pc *peerClient) SendCancel(pieceIndex, begin, length int) {
|
||
payload := make([]byte, 12)
|
||
binary.BigEndian.PutUint32(payload[0:4], uint32(pieceIndex))
|
||
binary.BigEndian.PutUint32(payload[4:8], uint32(begin))
|
||
binary.BigEndian.PutUint32(payload[8:12], uint32(length))
|
||
_ = pc.conn.SetWriteDeadline(time.Now().Add(peerWriteTimeout))
|
||
_ = pc.sendMessageDirect(msgCancel, payload)
|
||
}
|
||
|
||
// sendMessageDirect — упрощенная версия sendMessage без ctx (deadline уже выставлен).
|
||
func (pc *peerClient) sendMessageDirect(msgID int, payload []byte) error {
|
||
length := uint32(1 + len(payload))
|
||
buf := make([]byte, 4+length)
|
||
binary.BigEndian.PutUint32(buf[0:4], length)
|
||
buf[4] = byte(msgID)
|
||
copy(buf[5:], payload)
|
||
_, err := pc.conn.Write(buf)
|
||
return err
|
||
}
|
||
|
||
|
||
func (pc *peerClient) readInitialMessages(ctx context.Context) error {
|
||
if err := pc.setReadDeadlineFromContext(ctx, 2*time.Second); err != nil {
|
||
return err
|
||
}
|
||
defer pc.conn.SetReadDeadline(time.Time{})
|
||
|
||
for {
|
||
msg, err := readWireMessage(pc.conn)
|
||
if err != nil {
|
||
if isTimeout(err) {
|
||
debugf("peer initial message window finished")
|
||
return nil
|
||
}
|
||
return err
|
||
}
|
||
pc.consumeMessage(msg)
|
||
}
|
||
}
|
||
|
||
func (pc *peerClient) waitForUnchoke(ctx context.Context) error {
|
||
if !pc.peerIsChoked {
|
||
return nil
|
||
}
|
||
|
||
deadline := time.Now().Add(12 * time.Second)
|
||
for {
|
||
if err := ctx.Err(); err != nil {
|
||
return err
|
||
}
|
||
if time.Now().After(deadline) {
|
||
return errors.New("peer did not unchoke")
|
||
}
|
||
|
||
if err := pc.setReadDeadlineFromContext(ctx, peerReadTimeout); err != nil {
|
||
return err
|
||
}
|
||
msg, err := readWireMessage(pc.conn)
|
||
if err != nil {
|
||
if isTimeout(err) {
|
||
continue
|
||
}
|
||
return err
|
||
}
|
||
pc.consumeMessage(msg)
|
||
if msg.ID == msgUnchoke {
|
||
debugf("peer unchoked")
|
||
return nil
|
||
}
|
||
}
|
||
}
|
||
|
||
func (pc *peerClient) sendRequest(ctx context.Context, pieceIndex, begin, length int) error {
|
||
debugf("request piece %d offset %d length %d", pieceIndex, begin, length)
|
||
payload := make([]byte, 12)
|
||
binary.BigEndian.PutUint32(payload[0:4], uint32(pieceIndex))
|
||
binary.BigEndian.PutUint32(payload[4:8], uint32(begin))
|
||
binary.BigEndian.PutUint32(payload[8:12], uint32(length))
|
||
return pc.sendMessage(ctx, msgRequest, payload)
|
||
}
|
||
|
||
func (pc *peerClient) readPieceMessage(ctx context.Context) (pieceIndex int, begin int, block []byte, err error) {
|
||
for {
|
||
if err := ctx.Err(); err != nil {
|
||
return 0, 0, nil, err
|
||
}
|
||
if err := pc.setReadDeadlineFromContext(ctx, peerReadTimeout); err != nil {
|
||
return 0, 0, nil, err
|
||
}
|
||
|
||
msg, err := readWireMessage(pc.conn)
|
||
if err != nil {
|
||
if isTimeout(err) {
|
||
continue
|
||
}
|
||
return 0, 0, nil, err
|
||
}
|
||
|
||
switch msg.ID {
|
||
case msgPiece:
|
||
if len(msg.Payload) < 8 {
|
||
continue
|
||
}
|
||
gotIndex := int(binary.BigEndian.Uint32(msg.Payload[0:4]))
|
||
gotBegin := int(binary.BigEndian.Uint32(msg.Payload[4:8]))
|
||
data := msg.Payload[8:]
|
||
if len(data) == 0 {
|
||
continue
|
||
}
|
||
debugf("received piece %d offset %d block=%d", gotIndex, gotBegin, len(data))
|
||
return gotIndex, gotBegin, data, nil
|
||
case msgChoke:
|
||
pc.consumeMessage(msg)
|
||
return 0, 0, nil, errors.New("peer choked")
|
||
default:
|
||
pc.consumeMessage(msg)
|
||
}
|
||
}
|
||
}
|
||
|
||
func (pc *peerClient) consumeMessage(msg wireMessage) {
|
||
switch msg.ID {
|
||
case msgChoke:
|
||
pc.peerIsChoked = true
|
||
debugf("peer sent choke")
|
||
case msgUnchoke:
|
||
pc.peerIsChoked = false
|
||
debugf("peer sent unchoke")
|
||
case msgHave:
|
||
if len(msg.Payload) < 4 {
|
||
return
|
||
}
|
||
idx := int(binary.BigEndian.Uint32(msg.Payload[:4]))
|
||
if idx >= 0 && idx < len(pc.have) {
|
||
pc.have[idx] = true
|
||
pc.hasPieceInfo = true
|
||
debugf("received have piece %d", idx)
|
||
}
|
||
case msgBitfield:
|
||
pc.hasPieceInfo = true
|
||
debugf("received bitfield (%d bytes)", len(msg.Payload))
|
||
for i := range pc.have {
|
||
pc.have[i] = bitfieldHasPiece(msg.Payload, i)
|
||
}
|
||
case msgExtended:
|
||
if len(msg.Payload) == 0 {
|
||
return
|
||
}
|
||
extID := int(msg.Payload[0])
|
||
if extID == 0 {
|
||
var extMsg extendedHandshake
|
||
if err := bencode.Unmarshal(bytes.NewReader(msg.Payload[1:]), &extMsg); err == nil {
|
||
if id, ok := extMsg.M["ut_metadata"]; ok {
|
||
pc.peerUtMetadataID = id
|
||
}
|
||
if id, ok := extMsg.M["ut_pex"]; ok {
|
||
pc.peerUtPexID = id
|
||
}
|
||
if extMsg.MetadataSize > 0 {
|
||
pc.metadataSize = extMsg.MetadataSize
|
||
}
|
||
debugf("received extended handshake, ut_metadata=%d, ut_pex=%d, size=%d", pc.peerUtMetadataID, pc.peerUtPexID, pc.metadataSize)
|
||
}
|
||
} else if extID == pc.peerUtPexID && pc.peerUtPexID != 0 {
|
||
if pc.OnPex != nil {
|
||
if peers := ParsePexPayload(msg.Payload[1:]); len(peers) > 0 {
|
||
pc.OnPex(peers)
|
||
}
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
func (pc *peerClient) sendExtendedHandshake(ctx context.Context) error {
|
||
debugf("sending extended handshake")
|
||
msg := extendedHandshake{
|
||
M: map[string]int{
|
||
"ut_metadata": 1,
|
||
"ut_pex": 2,
|
||
},
|
||
}
|
||
var buf bytes.Buffer
|
||
buf.WriteByte(0) // Extended message ID 0 for handshake
|
||
if err := bencode.Marshal(&buf, msg); err != nil {
|
||
return err
|
||
}
|
||
return pc.sendMessage(ctx, msgExtended, buf.Bytes())
|
||
}
|
||
|
||
func (pc *peerClient) sendHandshake(ctx context.Context, infoHash [20]byte, peerID [20]byte) error {
|
||
debugf("sending handshake")
|
||
payload := make([]byte, 49+len(wireProtocolString))
|
||
payload[0] = byte(len(wireProtocolString))
|
||
copy(payload[1:1+len(wireProtocolString)], wireProtocolString)
|
||
payload[25] |= 0x10 // Set extension protocol bit
|
||
copy(payload[1+len(wireProtocolString)+8:1+len(wireProtocolString)+8+20], infoHash[:])
|
||
copy(payload[1+len(wireProtocolString)+8+20:], peerID[:])
|
||
|
||
if err := pc.setWriteDeadlineFromContext(ctx, peerWriteTimeout); err != nil {
|
||
return err
|
||
}
|
||
_, err := pc.conn.Write(payload)
|
||
return err
|
||
}
|
||
|
||
func (pc *peerClient) readHandshake(ctx context.Context, expectedInfoHash [20]byte) error {
|
||
head := make([]byte, 1)
|
||
if err := pc.setReadDeadlineFromContext(ctx, peerReadTimeout); err != nil {
|
||
return err
|
||
}
|
||
if _, err := io.ReadFull(pc.conn, head); err != nil {
|
||
return err
|
||
}
|
||
|
||
pstrlen := int(head[0])
|
||
if pstrlen <= 0 || pstrlen > 64 {
|
||
return fmt.Errorf("invalid handshake pstrlen %d", pstrlen)
|
||
}
|
||
|
||
rest := make([]byte, pstrlen+48)
|
||
if _, err := io.ReadFull(pc.conn, rest); err != nil {
|
||
return err
|
||
}
|
||
|
||
if string(rest[:pstrlen]) != wireProtocolString {
|
||
return errors.New("invalid peer protocol string")
|
||
}
|
||
|
||
infoHashOffset := pstrlen + 8
|
||
if !bytes.Equal(rest[infoHashOffset:infoHashOffset+20], expectedInfoHash[:]) {
|
||
return errors.New("peer info_hash mismatch")
|
||
}
|
||
pc.supportsExtensions = (rest[24] & 0x10) != 0
|
||
debugf("handshake complete (extensions: %v)", pc.supportsExtensions)
|
||
|
||
return nil
|
||
}
|
||
|
||
func (pc *peerClient) SendUnchoke(ctx context.Context) error {
|
||
return pc.sendMessage(ctx, msgUnchoke, nil)
|
||
}
|
||
|
||
func (pc *peerClient) SendChoke(ctx context.Context) error {
|
||
return pc.sendMessage(ctx, msgChoke, nil)
|
||
}
|
||
|
||
func (pc *peerClient) SendBitfield(ctx context.Context, bitfield []byte) error {
|
||
return pc.sendMessage(ctx, msgBitfield, bitfield)
|
||
}
|
||
|
||
func (pc *peerClient) SendPiece(ctx context.Context, index, begin int, data []byte) error {
|
||
payload := make([]byte, 8+len(data))
|
||
binary.BigEndian.PutUint32(payload[0:4], uint32(index))
|
||
binary.BigEndian.PutUint32(payload[4:8], uint32(begin))
|
||
copy(payload[8:], data)
|
||
return pc.sendMessage(ctx, msgPiece, payload)
|
||
}
|
||
|
||
func ParsePexPayload(payload []byte) []tracker.Peer {
|
||
var msg pexMessage
|
||
if err := bencode.Unmarshal(bytes.NewReader(payload), &msg); err != nil {
|
||
return nil
|
||
}
|
||
|
||
var allPeers []tracker.Peer
|
||
if len(msg.Added) > 0 {
|
||
if peers, err := tracker.ParsePeers([]byte(msg.Added)); err == nil {
|
||
allPeers = append(allPeers, peers...)
|
||
}
|
||
}
|
||
if len(msg.Added6) > 0 {
|
||
if peers6, err := tracker.ParsePeers6([]byte(msg.Added6)); err == nil {
|
||
allPeers = append(allPeers, peers6...)
|
||
}
|
||
}
|
||
return allPeers
|
||
}
|
||
|
||
func (pc *peerClient) SendPex(ctx context.Context, added []tracker.Peer) error {
|
||
if !pc.supportsExtensions || pc.peerUtPexID == 0 {
|
||
return errors.New("peer does not support ut_pex")
|
||
}
|
||
|
||
var addedBuf bytes.Buffer
|
||
for _, p := range added {
|
||
if v4 := p.IP.To4(); v4 != nil {
|
||
addedBuf.Write(v4)
|
||
var portBuf [2]byte
|
||
binary.BigEndian.PutUint16(portBuf[:], p.Port)
|
||
addedBuf.Write(portBuf[:])
|
||
}
|
||
}
|
||
|
||
msg := pexMessage{
|
||
Added: addedBuf.String(),
|
||
}
|
||
|
||
var buf bytes.Buffer
|
||
buf.WriteByte(byte(pc.peerUtPexID))
|
||
if err := bencode.Marshal(&buf, msg); err != nil {
|
||
return err
|
||
}
|
||
|
||
return pc.sendMessage(ctx, msgExtended, buf.Bytes())
|
||
}
|
||
|
||
func (pc *peerClient) ReadMessage(ctx context.Context) (wireMessage, error) {
|
||
if err := pc.setReadDeadlineFromContext(ctx, peerReadTimeout); err != nil {
|
||
return wireMessage{}, err
|
||
}
|
||
return readWireMessage(pc.conn)
|
||
}
|
||
|
||
func (pc *peerClient) sendMessage(ctx context.Context, msgID int, payload []byte) error {
|
||
if err := pc.setWriteDeadlineFromContext(ctx, peerWriteTimeout); err != nil {
|
||
return err
|
||
}
|
||
|
||
length := uint32(1 + len(payload))
|
||
buf := make([]byte, 4+length)
|
||
binary.BigEndian.PutUint32(buf[0:4], length)
|
||
buf[4] = byte(msgID)
|
||
copy(buf[5:], payload)
|
||
|
||
_, err := pc.conn.Write(buf)
|
||
return err
|
||
}
|
||
|
||
func readWireMessage(r io.Reader) (wireMessage, error) {
|
||
var lengthBuf [4]byte
|
||
if _, err := io.ReadFull(r, lengthBuf[:]); err != nil {
|
||
return wireMessage{}, err
|
||
}
|
||
|
||
length := binary.BigEndian.Uint32(lengthBuf[:])
|
||
if length == 0 {
|
||
return wireMessage{ID: -1}, nil
|
||
}
|
||
if length > maxWireMessageSize {
|
||
return wireMessage{}, fmt.Errorf("wire message too large: %d", length)
|
||
}
|
||
|
||
msg := make([]byte, length)
|
||
if _, err := io.ReadFull(r, msg); err != nil {
|
||
return wireMessage{}, err
|
||
}
|
||
return wireMessage{ID: int(msg[0]), Payload: msg[1:]}, nil
|
||
}
|
||
|
||
func bitfieldHasPiece(bitfield []byte, index int) bool {
|
||
byteIndex := index / 8
|
||
if byteIndex < 0 || byteIndex >= len(bitfield) {
|
||
return false
|
||
}
|
||
bitOffset := 7 - (index % 8)
|
||
return bitfield[byteIndex]&(1<<bitOffset) != 0
|
||
}
|
||
|
||
func pieceSizeForIndex(tf *torrentfile.TorrentFile, pieceIndex int) int {
|
||
if pieceIndex < 0 || pieceIndex >= len(tf.PieceHashes) {
|
||
return 0
|
||
}
|
||
if pieceIndex == len(tf.PieceHashes)-1 {
|
||
lastPieceSize := tf.Length % tf.PieceLength
|
||
if lastPieceSize == 0 {
|
||
return tf.PieceLength
|
||
}
|
||
return lastPieceSize
|
||
}
|
||
return tf.PieceLength
|
||
}
|
||
|
||
func (pc *peerClient) PeerUtMetadataID() int {
|
||
return pc.peerUtMetadataID
|
||
}
|
||
|
||
func (pc *peerClient) MetadataSize() int {
|
||
return pc.metadataSize
|
||
}
|
||
|
||
func (pc *peerClient) SendMetadataRequest(ctx context.Context, piece int) error {
|
||
if pc.peerUtMetadataID == 0 {
|
||
return errors.New("peer does not support ut_metadata")
|
||
}
|
||
msg := map[string]int{
|
||
"msg_type": 0, // request
|
||
"piece": piece,
|
||
}
|
||
var buf bytes.Buffer
|
||
buf.WriteByte(byte(pc.peerUtMetadataID))
|
||
if err := bencode.Marshal(&buf, msg); err != nil {
|
||
return err
|
||
}
|
||
return pc.sendMessage(ctx, msgExtended, buf.Bytes())
|
||
}
|
||
|
||
func (pc *peerClient) ReadMetadataMessage(ctx context.Context) (piece int, data []byte, reject bool, err error) {
|
||
for {
|
||
if err := ctx.Err(); err != nil {
|
||
return 0, nil, false, err
|
||
}
|
||
if err := pc.setReadDeadlineFromContext(ctx, peerReadTimeout); err != nil {
|
||
return 0, nil, false, err
|
||
}
|
||
msg, err := readWireMessage(pc.conn)
|
||
if err != nil {
|
||
if isTimeout(err) {
|
||
continue
|
||
}
|
||
return 0, nil, false, err
|
||
}
|
||
if msg.ID == msgExtended {
|
||
if len(msg.Payload) == 0 {
|
||
continue
|
||
}
|
||
extID := int(msg.Payload[0])
|
||
if extID == 1 { // our ut_metadata ID is 1
|
||
reader := bytes.NewReader(msg.Payload[1:])
|
||
var dict map[string]int
|
||
if err := bencode.Unmarshal(reader, &dict); err != nil {
|
||
continue
|
||
}
|
||
msgType, ok := dict["msg_type"]
|
||
if !ok {
|
||
continue
|
||
}
|
||
pieceIdx := dict["piece"]
|
||
|
||
if msgType == 2 {
|
||
return pieceIdx, nil, true, nil
|
||
}
|
||
if msgType == 1 {
|
||
bytesRead := len(msg.Payload[1:]) - reader.Len()
|
||
data := msg.Payload[1+bytesRead:]
|
||
return pieceIdx, data, false, nil
|
||
}
|
||
} else if extID == 0 {
|
||
pc.consumeMessage(msg)
|
||
}
|
||
} else {
|
||
pc.consumeMessage(msg)
|
||
}
|
||
}
|
||
}
|
||
|
||
func (pc *peerClient) setReadDeadlineFromContext(ctx context.Context, fallback time.Duration) error {
|
||
deadline := time.Now().Add(fallback)
|
||
if ctxDeadline, ok := ctx.Deadline(); ok && ctxDeadline.Before(deadline) {
|
||
deadline = ctxDeadline
|
||
}
|
||
return pc.conn.SetReadDeadline(deadline)
|
||
}
|
||
|
||
func (pc *peerClient) setWriteDeadlineFromContext(ctx context.Context, fallback time.Duration) error {
|
||
deadline := time.Now().Add(fallback)
|
||
if ctxDeadline, ok := ctx.Deadline(); ok && ctxDeadline.Before(deadline) {
|
||
deadline = ctxDeadline
|
||
}
|
||
return pc.conn.SetWriteDeadline(deadline)
|
||
}
|
||
|
||
func isTimeout(err error) bool {
|
||
if err == nil {
|
||
return false
|
||
}
|
||
var netErr net.Error
|
||
return errors.As(err, &netErr) && netErr.Timeout()
|
||
}
|
||
|
||
func maxInt64(a, b int64) int64 {
|
||
if a >= b {
|
||
return a
|
||
}
|
||
return b
|
||
}
|