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 }