From cf1a440641745f2d7ae602f21027abfaf236371e Mon Sep 17 00:00:00 2001 From: Chen Kai <281165273grape@gmail.com> Date: Thu, 16 Nov 2023 17:47:22 +0800 Subject: [PATCH] feat:handle offer extension Signed-off-by: Chen Kai <281165273grape@gmail.com> --- p2p/discover/portal_protocol.go | 120 +++++++++++++++++++++++--------- 1 file changed, 86 insertions(+), 34 deletions(-) diff --git a/p2p/discover/portal_protocol.go b/p2p/discover/portal_protocol.go index 1d073fd653..3cb4a7cf9b 100644 --- a/p2p/discover/portal_protocol.go +++ b/p2p/discover/portal_protocol.go @@ -171,19 +171,19 @@ func (p *PortalProtocol) setupUDPListening() (*net.UDPConn, error) { func(buf []byte, addr *net.UDPAddr) (int, error) { nodes := p.table.Nodes() var target *enode.Node - for _, node := range nodes { - if addr.Port != node.UDP() { + for _, n := range nodes { + if addr.Port != n.UDP() { continue } - if addr.IP != nil && addr.IP.To4().String() == node.IP().To4().String() { - target = node + if addr.IP != nil && addr.IP.To4().String() == n.IP().To4().String() { + target = n break } if addr.IP == nil { - nodeIp := node.IP().To4().String() + nodeIp := n.IP().To4().String() if nodeIp == "127.0.0.1" || nodeIp == "0.0.0.0" { - target = node + target = n break } } @@ -194,6 +194,7 @@ func (p *PortalProtocol) setupUDPListening() (*net.UDPConn, error) { return len(buf), err }) + // TODO: ZAP PRODUCTION LOG logger, err := zap.NewDevelopmentConfig().Build() if err != nil { return nil, err @@ -353,19 +354,19 @@ func (p *PortalProtocol) processContent(target *enode.Node, resp []byte) (byte, } p.log.Trace("Received content response", "id", target.ID(), "connIdMsg", connIdMsg) - rctx, rcancel := context.WithTimeout(context.Background(), defaultUTPConnectTimeout) + connctx, conncancel := context.WithTimeout(context.Background(), defaultUTPConnectTimeout) laddr := p.utp.Addr().(*utp.Addr) raddr := &utp.Addr{IP: target.IP(), Port: target.UDP()} connId := binary.BigEndian.Uint16(connIdMsg.Id[:]) - conn, err := utp.DialUTPOptions("utp", laddr, raddr, utp.WithContext(rctx), utp.WithSocketManager(p.utpSm), utp.WithConnId(uint32(connId))) + conn, err := utp.DialUTPOptions("utp", laddr, raddr, utp.WithContext(connctx), utp.WithSocketManager(p.utpSm), utp.WithConnId(uint32(connId))) if err != nil { - rcancel() + conncancel() return 0xff, nil, err } + conncancel() err = conn.SetReadDeadline(time.Now().Add(defaultUTPReadTimeout)) if err != nil { - rcancel() return 0xff, nil, err } // Read ALL the data from the connection until EOF and return it @@ -375,7 +376,6 @@ func (p *PortalProtocol) processContent(target *enode.Node, resp []byte) (byte, var n int n, err = conn.Read(buf) if err != nil { - rcancel() if errors.Is(err, io.EOF) { p.log.Trace("Received content response", "id", target.ID(), "data", data, "size", n) return resp[1], data, nil @@ -717,7 +717,7 @@ func (p *PortalProtocol) handleFindContent(id enode.ID, addr *net.UDPAddr, reque var conn *utp.Conn conn, err = p.utp.AcceptUTPContext(ctx, connIdSend) defer func(conn *utp.Conn) { - err := conn.Close() + err = conn.Close() if err != nil { p.log.Error("failed to close utp connection", "err", err) } @@ -767,13 +767,14 @@ func (p *PortalProtocol) handleFindContent(id enode.ID, addr *net.UDPAddr, reque } func (p *PortalProtocol) handleOffer(id enode.ID, addr *net.UDPAddr, request *portalwire.Offer) ([]byte, error) { + var err error contentKeyBitsets := bitfield.NewBitlist(uint64(len(request.ContentKeys))) contentKeys := make([][]byte, 0) for i, contentKey := range request.ContentKeys { contentId := p.storage.ContentId(contentKey) if contentId != nil { if p.inRange(p.Self().ID(), p.nodeRadius, contentId) { - if _, err := p.storage.Get(contentKey, contentId); err != nil { + if _, err = p.storage.Get(contentKey, contentId); err != nil { contentKeyBitsets.SetBitAt(uint64(i), true) contentKeys = append(contentKeys, contentKey) } @@ -783,30 +784,81 @@ func (p *PortalProtocol) handleOffer(id enode.ID, addr *net.UDPAddr, request *po } } - fmt.Println(contentKeys) - if contentKeyBitsets.Count() == 0 { - idBuffer := make([]byte, 2) + idBuffer := make([]byte, 2) + if contentKeyBitsets.Count() != 0 { + connIdGen := utp.NewConnIdGenerator() + connId := connIdGen.GenCid(id, false) + connIdSend := connId.SendId() + + go func() { + ctx, cancel := context.WithTimeout(context.Background(), defaultUTPConnectTimeout) + var conn *utp.Conn + conn, err = p.utp.AcceptUTPContext(ctx, connIdSend) + defer func(conn *utp.Conn) { + err = conn.Close() + if err != nil { + p.log.Error("failed to close utp connection", "err", err) + } + }(conn) + if err != nil { + p.log.Error("failed to accept utp connection", "connId", connIdSend, "err", err) + cancel() + return + } + cancel() + + err = conn.SetReadDeadline(time.Now().Add(defaultUTPReadTimeout)) + if err != nil { + p.log.Error("failed to set read deadline", "err", err) + return + } + // Read ALL the data from the connection until EOF and return it + data := make([]byte, 0) + buf := make([]byte, 1024) + for { + var n int + n, err = conn.Read(buf) + if err != nil { + if errors.Is(err, io.EOF) { + p.log.Trace("Received content response", "id", id, "data", data, "size", n) + break + } + + p.log.Error("failed to read from utp connection", "err", err) + return + } + data = append(data, buf[:n]...) + } + + p.handleOfferedContents(id, contentKeys, data) + }() + + binary.BigEndian.PutUint16(idBuffer, uint16(connIdSend)) + } else { binary.BigEndian.PutUint16(idBuffer, uint16(0)) - acceptMsg := &portalwire.Accept{ - ConnectionId: idBuffer, - ContentKeys: contentKeyBitsets.BytesNoTrim(), - } - - p.log.Trace("Sending accept response", "protocol", p.protocolId, "source", addr, "accept", acceptMsg) - var acceptMsgBytes []byte - acceptMsgBytes, err := acceptMsg.MarshalSSZ() - if err != nil { - return nil, err - } - - talkRespBytes := make([]byte, 0, len(acceptMsgBytes)+1) - talkRespBytes = append(talkRespBytes, portalwire.ACCEPT) - talkRespBytes = append(talkRespBytes, acceptMsgBytes...) - - return talkRespBytes, nil } - return nil, nil + acceptMsg := &portalwire.Accept{ + ConnectionId: idBuffer, + ContentKeys: contentKeyBitsets.BytesNoTrim(), + } + + p.log.Trace("Sending accept response", "protocol", p.protocolId, "source", addr, "accept", acceptMsg) + var acceptMsgBytes []byte + acceptMsgBytes, err = acceptMsg.MarshalSSZ() + if err != nil { + return nil, err + } + + talkRespBytes := make([]byte, 0, len(acceptMsgBytes)+1) + talkRespBytes = append(talkRespBytes, portalwire.ACCEPT) + talkRespBytes = append(talkRespBytes, acceptMsgBytes...) + + return talkRespBytes, nil +} + +func (p *PortalProtocol) handleOfferedContents(id enode.ID, keys [][]byte, data []byte) { + } func (p *PortalProtocol) Self() *enode.Node {