From 67f28cf23a8ae59a38e0128390ccdad27b5526be Mon Sep 17 00:00:00 2001 From: Steffen Vogel Date: Thu, 9 Feb 2023 21:23:59 +0100 Subject: [PATCH] Move util.go to more appropriately named files Move util.go to more appropriately named files --- addr.go | 48 +++++++++++++++++ agent.go | 7 +-- errors.go | 1 - gather.go | 3 +- internal/stun/stun.go | 69 ++++++++++++++++++++++++ util.go => net.go | 103 ------------------------------------ util_test.go => net_test.go | 0 stun.go | 24 --------- 8 files changed, 123 insertions(+), 132 deletions(-) create mode 100644 addr.go create mode 100644 internal/stun/stun.go rename util.go => net.go (53%) rename util_test.go => net_test.go (100%) delete mode 100644 stun.go diff --git a/addr.go b/addr.go new file mode 100644 index 0000000..cbcdc4c --- /dev/null +++ b/addr.go @@ -0,0 +1,48 @@ +package ice + +import "net" + +func parseMulticastAnswerAddr(in net.Addr) (net.IP, bool) { + switch addr := in.(type) { + case *net.IPAddr: + return addr.IP, true + case *net.UDPAddr: + return addr.IP, true + case *net.TCPAddr: + return addr.IP, true + } + return nil, false +} + +func parseAddr(in net.Addr) (net.IP, int, NetworkType, bool) { + switch addr := in.(type) { + case *net.UDPAddr: + return addr.IP, addr.Port, NetworkTypeUDP4, true + case *net.TCPAddr: + return addr.IP, addr.Port, NetworkTypeTCP4, true + } + return nil, 0, 0, false +} + +func createAddr(network NetworkType, ip net.IP, port int) net.Addr { + switch { + case network.IsTCP(): + return &net.TCPAddr{IP: ip, Port: port} + default: + return &net.UDPAddr{IP: ip, Port: port} + } +} + +func addrEqual(a, b net.Addr) bool { + aIP, aPort, aType, aOk := parseAddr(a) + if !aOk { + return false + } + + bIP, bPort, bType, bOk := parseAddr(b) + if !bOk { + return false + } + + return aType == bType && aIP.Equal(bIP) && aPort == bPort +} diff --git a/agent.go b/agent.go index b1017d3..3497a44 100644 --- a/agent.go +++ b/agent.go @@ -12,6 +12,7 @@ import ( "time" atomicx "github.com/pion/ice/v2/internal/atomic" + stunx "github.com/pion/ice/v2/internal/stun" "github.com/pion/logging" "github.com/pion/mdns" "github.com/pion/stun" @@ -1075,7 +1076,7 @@ func (a *Agent) handleInbound(m *stun.Message, local Candidate, remote net.Addr) remoteCandidate := a.findRemoteCandidate(local.NetworkType(), remote) if m.Type.Class == stun.ClassSuccessResponse { - if err = assertInboundMessageIntegrity(m, []byte(a.remotePwd)); err != nil { + if err = stun.MessageIntegrity([]byte(a.remotePwd)).Check(m); err != nil { a.log.Warnf("discard message from (%s), %v", remote, err) return } @@ -1087,10 +1088,10 @@ func (a *Agent) handleInbound(m *stun.Message, local Candidate, remote net.Addr) a.selector.HandleSuccessResponse(m, local, remoteCandidate, remote) } else if m.Type.Class == stun.ClassRequest { - if err = assertInboundUsername(m, a.localUfrag+":"+a.remoteUfrag); err != nil { + if err = stunx.AssertUsername(m, a.localUfrag+":"+a.remoteUfrag); err != nil { a.log.Warnf("discard message from (%s), %v", remote, err) return - } else if err = assertInboundMessageIntegrity(m, []byte(a.localPwd)); err != nil { + } else if err = stun.MessageIntegrity([]byte(a.localPwd)).Check(m); err != nil { a.log.Warnf("discard message from (%s), %v", remote, err) return } diff --git a/errors.go b/errors.go index 647b2be..93b9a64 100644 --- a/errors.go +++ b/errors.go @@ -128,7 +128,6 @@ var ( errTooManyColonsAddr = errors.New("too many colons in address") errRead = errors.New("unexpected error trying to read") errUnknownRole = errors.New("unknown role") - errMismatchUsername = errors.New("username mismatch") errICEWriteSTUNMessage = errors.New("the ICE conn can't write STUN messages") errUDPMuxDisabled = errors.New("UDPMux is not enabled") errNoXorAddrMapping = errors.New("no address mapping") diff --git a/gather.go b/gather.go index 9c90eaa..bba3b4f 100644 --- a/gather.go +++ b/gather.go @@ -13,6 +13,7 @@ import ( "github.com/pion/dtls/v2" "github.com/pion/ice/v2/internal/fakenet" + stunx "github.com/pion/ice/v2/internal/stun" "github.com/pion/logging" "github.com/pion/turn/v2" ) @@ -466,7 +467,7 @@ func (a *Agent) gatherCandidatesSrflx(ctx context.Context, urls []*URL, networkT } }() - xorAddr, err := getXORMappedAddr(conn, serverAddr, stunGatherTimeout) + xorAddr, err := stunx.GetXORMappedAddr(conn, serverAddr, stunGatherTimeout) if err != nil { closeConnAndLog(conn, a.log, fmt.Sprintf("could not get server reflexive address %s %s: %v", network, url, err)) return diff --git a/internal/stun/stun.go b/internal/stun/stun.go new file mode 100644 index 0000000..0298abd --- /dev/null +++ b/internal/stun/stun.go @@ -0,0 +1,69 @@ +// Package stun contains ICE specific STUN code +package stun + +import ( + "errors" + "fmt" + "net" + "time" + + "github.com/pion/stun" +) + +var ( + errGetXorMappedAddrResponse = errors.New("failed to get XOR-MAPPED-ADDRESS response") + errMismatchUsername = errors.New("username mismatch") +) + +// GetXORMappedAddr initiates a stun requests to serverAddr using conn, reads the response and returns +// the XORMappedAddress returned by the STUN server. +func GetXORMappedAddr(conn net.PacketConn, serverAddr net.Addr, timeout time.Duration) (*stun.XORMappedAddress, error) { + if timeout > 0 { + if err := conn.SetReadDeadline(time.Now().Add(timeout)); err != nil { + return nil, err + } + + // Reset timeout after completion + defer conn.SetReadDeadline(time.Time{}) //nolint:errcheck + } + + req, err := stun.Build(stun.BindingRequest, stun.TransactionID) + if err != nil { + return nil, err + } + + if _, err = conn.WriteTo(req.Raw, serverAddr); err != nil { + return nil, err + } + + const maxMessageSize = 1280 + buf := make([]byte, maxMessageSize) + n, _, err := conn.ReadFrom(buf) + if err != nil { + return nil, err + } + + res := &stun.Message{Raw: buf[:n]} + if err = res.Decode(); err != nil { + return nil, err + } + + var addr stun.XORMappedAddress + if err = addr.GetFrom(res); err != nil { + return nil, fmt.Errorf("%w: %v", errGetXorMappedAddrResponse, err) + } + + return &addr, nil +} + +// AssertUsername checks that the given STUN message m has a USERNAME attribute with a given value +func AssertUsername(m *stun.Message, expectedUsername string) error { + var username stun.Username + if err := username.GetFrom(m); err != nil { + return err + } else if string(username) != expectedUsername { + return fmt.Errorf("%w expected(%x) actual(%x)", errMismatchUsername, expectedUsername, string(username)) + } + + return nil +} diff --git a/util.go b/net.go similarity index 53% rename from util.go rename to net.go index ec18691..d28970c 100644 --- a/util.go +++ b/net.go @@ -1,12 +1,9 @@ package ice import ( - "fmt" "net" - "time" "github.com/pion/logging" - "github.com/pion/stun" "github.com/pion/transport/v2" ) @@ -32,106 +29,6 @@ func isZeros(ip net.IP) bool { return true } -func parseMulticastAnswerAddr(in net.Addr) (net.IP, bool) { - switch addr := in.(type) { - case *net.IPAddr: - return addr.IP, true - case *net.UDPAddr: - return addr.IP, true - case *net.TCPAddr: - return addr.IP, true - } - return nil, false -} - -func parseAddr(in net.Addr) (net.IP, int, NetworkType, bool) { - switch addr := in.(type) { - case *net.UDPAddr: - return addr.IP, addr.Port, NetworkTypeUDP4, true - case *net.TCPAddr: - return addr.IP, addr.Port, NetworkTypeTCP4, true - } - return nil, 0, 0, false -} - -func createAddr(network NetworkType, ip net.IP, port int) net.Addr { - switch { - case network.IsTCP(): - return &net.TCPAddr{IP: ip, Port: port} - default: - return &net.UDPAddr{IP: ip, Port: port} - } -} - -func addrEqual(a, b net.Addr) bool { - aIP, aPort, aType, aOk := parseAddr(a) - if !aOk { - return false - } - - bIP, bPort, bType, bOk := parseAddr(b) - if !bOk { - return false - } - - return aType == bType && aIP.Equal(bIP) && aPort == bPort -} - -// getXORMappedAddr initiates a stun requests to serverAddr using conn, reads the response and returns -// the XORMappedAddress returned by the stun server. -// -// Adapted from stun v0.2. -func getXORMappedAddr(conn net.PacketConn, serverAddr net.Addr, deadline time.Duration) (*stun.XORMappedAddress, error) { - if deadline > 0 { - if err := conn.SetReadDeadline(time.Now().Add(deadline)); err != nil { - return nil, err - } - } - defer func() { - if deadline > 0 { - _ = conn.SetReadDeadline(time.Time{}) - } - }() - resp, err := stunRequest( - func(p []byte) (int, error) { - n, _, err := conn.ReadFrom(p) - return n, err - }, - func(b []byte) (int, error) { - return conn.WriteTo(b, serverAddr) - }, - ) - if err != nil { - return nil, err - } - var addr stun.XORMappedAddress - if err = addr.GetFrom(resp); err != nil { - return nil, fmt.Errorf("%w: %v", errGetXorMappedAddrResponse, err) - } - return &addr, nil -} - -func stunRequest(read func([]byte) (int, error), write func([]byte) (int, error)) (*stun.Message, error) { - req, err := stun.Build(stun.BindingRequest, stun.TransactionID) - if err != nil { - return nil, err - } - if _, err = write(req.Raw); err != nil { - return nil, err - } - const maxMessageSize = 1280 - bs := make([]byte, maxMessageSize) - n, err := read(bs) - if err != nil { - return nil, err - } - res := &stun.Message{Raw: bs[:n]} - if err := res.Decode(); err != nil { - return nil, err - } - return res, nil -} - func localInterfaces(n transport.Net, interfaceFilter func(string) bool, ipFilter func(net.IP) bool, networkTypes []NetworkType, includeLoopback bool) ([]net.IP, error) { //nolint:gocognit ips := []net.IP{} ifaces, err := n.Interfaces() diff --git a/util_test.go b/net_test.go similarity index 100% rename from util_test.go rename to net_test.go diff --git a/stun.go b/stun.go deleted file mode 100644 index bef7c87..0000000 --- a/stun.go +++ /dev/null @@ -1,24 +0,0 @@ -package ice - -import ( - "fmt" - - "github.com/pion/stun" -) - -func assertInboundUsername(m *stun.Message, expectedUsername string) error { - var username stun.Username - if err := username.GetFrom(m); err != nil { - return err - } - if string(username) != expectedUsername { - return fmt.Errorf("%w expected(%x) actual(%x)", errMismatchUsername, expectedUsername, string(username)) - } - - return nil -} - -func assertInboundMessageIntegrity(m *stun.Message, key []byte) error { - messageIntegrityAttr := stun.MessageIntegrity(key) - return messageIntegrityAttr.Check(m) -}