Ztorrent/internal/dht/dht.go

389 lines
8.9 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package dht
import (
"bytes"
"context"
"crypto/rand"
"encoding/hex"
"fmt"
"log"
"net"
"strings"
"sync"
"time"
"github.com/veggiedefender/torrent-client/internal/tracker"
)
const (
MaxNodes = 8
Port = 6881
ReadBufSize = 65536
Alpha = 5 // параллельных запросов
)
// Много bootstrap-нод для надёжности
var BootstrapNodes = []string{
"router.bittorrent.com:6881",
"dht.transmissionbt.com:6881",
"router.utorrent.com:6881",
"dht.aelitis.com:6881",
"bootstrap.jami.net:4222",
"router.silotis.us:6881",
"dht.libtorrent.org:25401",
}
type Server struct {
ID NodeID
conn *net.UDPConn
routingTable *RoutingTable
transactionsMu sync.Mutex
transactions map[string]chan Msg
PeersFound chan []tracker.Peer
}
func NewServer() *Server {
id := RandomNodeID()
return &Server{
ID: id,
routingTable: NewRoutingTable(id),
transactions: make(map[string]chan Msg),
PeersFound: make(chan []tracker.Peer, 100),
}
}
func (s *Server) Start(ctx context.Context, port int) error {
addr, err := net.ResolveUDPAddr("udp", fmt.Sprintf(":%d", port))
if err != nil {
return err
}
conn, err := net.ListenUDP("udp", addr)
if err != nil {
return err
}
s.conn = conn
go s.readLoop(ctx)
go s.bootstrap(ctx)
log.Printf("DHT Server listening on %s with ID %x", conn.LocalAddr(), s.ID[:8])
return nil
}
func (s *Server) readLoop(ctx context.Context) {
buf := make([]byte, ReadBufSize)
for {
select {
case <-ctx.Done():
s.conn.Close()
return
default:
}
s.conn.SetReadDeadline(time.Now().Add(1 * time.Second))
n, from, err := s.conn.ReadFromUDP(buf)
if err != nil {
if netErr, ok := err.(net.Error); ok && netErr.Timeout() {
continue
}
if strings.Contains(err.Error(), "use of closed network connection") {
return
}
continue
}
msg, err := DecodeMsg(buf[:n])
if err != nil {
continue
}
s.handleMsg(msg, from)
}
}
func (s *Server) handleMsg(msg Msg, from *net.UDPAddr) {
var senderID NodeID
var ok bool
if msg.Y == "q" && msg.A != nil {
if id, isStr := msg.A["id"].(string); isStr && len(id) == 20 {
copy(senderID[:], id)
ok = true
}
} else if msg.Y == "r" && msg.R != nil {
if id, isStr := msg.R["id"].(string); isStr && len(id) == 20 {
copy(senderID[:], id)
ok = true
}
}
if ok {
s.routingTable.AddNode(Node{ID: senderID, Addr: from})
}
if msg.Y == "r" || msg.Y == "e" {
s.transactionsMu.Lock()
ch, exists := s.transactions[msg.T]
s.transactionsMu.Unlock()
if exists {
select {
case ch <- msg:
default:
}
}
return
}
if msg.Y == "q" {
switch msg.Q {
case "ping":
s.sendResponse(from, msg.T, map[string]interface{}{
"id": string(s.ID[:]),
})
case "find_node":
targetStr, _ := msg.A["target"].(string)
if len(targetStr) == 20 {
var target NodeID
copy(target[:], targetStr)
nodes := s.routingTable.ClosestNodes(target, MaxNodes)
s.sendResponse(from, msg.T, map[string]interface{}{
"id": string(s.ID[:]),
"nodes": encodeNodes(nodes),
})
}
case "get_peers":
infoHashStr, _ := msg.A["info_hash"].(string)
if len(infoHashStr) == 20 {
var target NodeID
copy(target[:], infoHashStr)
nodes := s.routingTable.ClosestNodes(target, MaxNodes)
s.sendResponse(from, msg.T, map[string]interface{}{
"id": string(s.ID[:]),
"token": "token",
"nodes": encodeNodes(nodes),
})
}
}
}
}
func (s *Server) sendQuery(ctx context.Context, addr *net.UDPAddr, q string, a map[string]interface{}) (Msg, error) {
tidBytes := make([]byte, 2)
rand.Read(tidBytes)
tid := string(tidBytes)
a["id"] = string(s.ID[:])
msg := NewQuery(tid, q, a)
encoded, err := EncodeMsg(msg)
if err != nil {
return Msg{}, err
}
ch := make(chan Msg, 1)
s.transactionsMu.Lock()
s.transactions[tid] = ch
s.transactionsMu.Unlock()
defer func() {
s.transactionsMu.Lock()
delete(s.transactions, tid)
s.transactionsMu.Unlock()
}()
if _, err := s.conn.WriteToUDP(encoded, addr); err != nil {
return Msg{}, err
}
select {
case <-ctx.Done():
return Msg{}, ctx.Err()
case resp := <-ch:
if resp.Y == "e" {
return Msg{}, fmt.Errorf("KRPC error: %v", resp.E)
}
return resp, nil
case <-time.After(3 * time.Second):
return Msg{}, fmt.Errorf("timeout")
}
}
func (s *Server) sendResponse(addr *net.UDPAddr, tid string, r map[string]interface{}) {
msg := Msg{T: tid, Y: "r", R: r}
encoded, err := EncodeMsg(msg)
if err == nil {
s.conn.WriteToUDP(encoded, addr)
}
}
func (s *Server) bootstrap(ctx context.Context) {
var wg sync.WaitGroup
for _, addrStr := range BootstrapNodes {
addr, err := net.ResolveUDPAddr("udp", addrStr)
if err != nil {
continue
}
wg.Add(1)
go func(a *net.UDPAddr, name string) {
defer wg.Done()
qctx, cancel := context.WithTimeout(ctx, 5*time.Second)
defer cancel()
resp, err := s.sendQuery(qctx, a, "find_node", map[string]interface{}{
"target": string(s.ID[:]),
})
if err != nil {
log.Printf("DHT bootstrap %s failed: %v", name, err)
return
}
if resp.R != nil {
if nodesStr, ok := resp.R["nodes"].(string); ok {
count := len(nodesStr) / 26
log.Printf("DHT bootstrap %s: got %d nodes", name, count)
s.parseAndAddNodes(nodesStr)
}
}
}(addr, addrStr)
}
wg.Wait()
log.Printf("DHT bootstrap done, routing table: %d nodes", s.routingTable.Len())
}
// SearchForPeers непрерывно ищет пиров для данного info_hash.
// Не прекращает поиск, пока контекст не отменён.
func (s *Server) SearchForPeers(ctx context.Context, infoHash [20]byte) {
// Ждём bootstrap
for i := 0; i < 20; i++ {
if s.routingTable.Len() > 0 {
break
}
time.Sleep(300 * time.Millisecond)
}
log.Printf("DHT SearchForPeers starting, routing table: %d nodes", s.routingTable.Len())
targetID := NodeID(infoHash)
queried := make(map[string]bool) // ключ = IP:port строка
for {
select {
case <-ctx.Done():
return
default:
}
closest := s.routingTable.ClosestNodes(targetID, MaxNodes*4)
var toQuery []Node
for _, n := range closest {
key := n.Addr.String()
if !queried[key] {
toQuery = append(toQuery, n)
queried[key] = true
if len(toQuery) >= Alpha {
break
}
}
}
if len(toQuery) == 0 {
// Нет новых нод — сбрасываем карту уже запрошенных и пробуем снова
// (новые ноды могли добавиться в routing table)
log.Printf("DHT: no new nodes to query (%d total queried), resetting and retrying", len(queried))
queried = make(map[string]bool)
// Если routing table совсем пуста — делаем повторный bootstrap
if s.routingTable.Len() == 0 {
go s.bootstrap(ctx)
}
select {
case <-ctx.Done():
return
case <-time.After(5 * time.Second):
}
continue
}
var wg sync.WaitGroup
for _, n := range toQuery {
wg.Add(1)
go func(node Node) {
defer wg.Done()
qctx, cancel := context.WithTimeout(ctx, 5*time.Second)
defer cancel()
resp, err := s.sendQuery(qctx, node.Addr, "get_peers", map[string]interface{}{
"info_hash": string(infoHash[:]),
})
if err != nil {
return
}
if resp.R == nil {
return
}
// Добавляем новые ноды в routing table
if nodesStr, ok := resp.R["nodes"].(string); ok {
s.parseAndAddNodes(nodesStr)
}
// Пробуем извлечь пиров
if values, ok := resp.R["values"].([]interface{}); ok {
var allPeerData []byte
for _, v := range values {
if peerStr, ok := v.(string); ok {
allPeerData = append(allPeerData, []byte(peerStr)...)
}
}
if peers, err := tracker.ParsePeers(allPeerData); err == nil && len(peers) > 0 {
log.Printf("DHT: found %d peers from %s", len(peers), node.Addr)
select {
case s.PeersFound <- peers:
default:
// канал полный — пытаемся без блокировки
}
}
}
}(n)
}
wg.Wait()
}
}
func encodeNodes(nodes []Node) string {
var buf bytes.Buffer
for _, n := range nodes {
buf.Write(n.ID[:])
if v4 := n.Addr.IP.To4(); v4 != nil {
buf.Write(v4)
} else {
buf.Write(net.IPv4zero)
}
portBuf := make([]byte, 2)
portBuf[0] = byte(n.Addr.Port >> 8)
portBuf[1] = byte(n.Addr.Port)
buf.Write(portBuf)
}
return buf.String()
}
func (s *Server) parseAndAddNodes(nodesStr string) {
data := []byte(nodesStr)
for i := 0; i+26 <= len(data); i += 26 {
var id NodeID
copy(id[:], data[i:i+20])
ip := net.IP(data[i+20 : i+24])
port := int(data[i+24])<<8 | int(data[i+25])
if port == 0 {
continue
}
s.routingTable.AddNode(Node{
ID: id,
Addr: &net.UDPAddr{
IP: ip,
Port: port,
},
})
}
}
func HexToNodeID(h string) NodeID {
var id NodeID
b, _ := hex.DecodeString(h)
copy(id[:], b)
return id
}