389 lines
8.9 KiB
Go
389 lines
8.9 KiB
Go
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
|
||
}
|