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<= 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 }