Ztorrent/internal/dht/dht.go

394 lines
8.8 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"
"net"
"strings"
"sync"
"time"
"github.com/veggiedefender/torrent-client/internal/logger"
"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)
logger.Info("DHT", "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) {
var tid string
s.transactionsMu.Lock()
for {
tidBytes := make([]byte, 2)
rand.Read(tidBytes)
tid = string(tidBytes)
if _, exists := s.transactions[tid]; !exists {
break
}
}
ch := make(chan Msg, 1)
s.transactions[tid] = ch
s.transactionsMu.Unlock()
a["id"] = string(s.ID[:])
msg := NewQuery(tid, q, a)
encoded, err := EncodeMsg(msg)
if err != nil {
return Msg{}, err
}
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 {
logger.Info("DHT", "DHT bootstrap %s failed: %v", name, err)
return
}
if resp.R != nil {
if nodesStr, ok := resp.R["nodes"].(string); ok {
count := len(nodesStr) / 26
logger.Info("DHT", "DHT bootstrap %s: got %d nodes", name, count)
s.parseAndAddNodes(nodesStr)
}
}
}(addr, addrStr)
}
wg.Wait()
logger.Info("DHT", "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)
}
logger.Info("DHT", "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 {
logger.Info("DHT", "DHT: no new nodes to query (%d total queried), waiting before retry", len(queried))
queried = make(map[string]bool)
if s.routingTable.Len() == 0 {
go s.bootstrap(ctx)
}
select {
case <-ctx.Done():
return
case <-time.After(30 * time.Second):
}
continue
}
for _, n := range toQuery {
go func(node Node) {
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 {
logger.Info("DHT", "DHT: found %d peers from %s", len(peers), node.Addr)
select {
case s.PeersFound <- peers:
case <-ctx.Done():
return
case <-time.After(2 * time.Second):
}
}
}
}(n)
}
select {
case <-ctx.Done():
return
case <-time.After(500 * time.Millisecond):
}
}
}
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
}