diff --git a/client.go b/client.go index b89dd31a..b1104542 100644 --- a/client.go +++ b/client.go @@ -5,6 +5,7 @@ package turn import ( b64 "encoding/base64" + "errors" "fmt" "math" "net" @@ -23,6 +24,10 @@ const ( defaultRTO = 200 * time.Millisecond maxRtxCount = 7 // Total 7 requests (Rc) maxDataBufferSize = math.MaxUint16 // Message size limit for Chromium + // while there is no limit in the RFC (RFC 5389/8489), + // this caps how many consecutive ALTERNATE-SERVER redirects + // we have a generous limit of 64 to prevent infinite loops and abuse. + maxAltSrvRedirects = 64 ) // interval [msec] @@ -67,7 +72,7 @@ type Client struct { conn net.PacketConn // Read-only net transport.Net // Read-only stunServerAddr net.Addr // Read-only - turnServerAddr net.Addr // Read-only + turnServerAddr net.Addr username stun.Username // Read-only password string // Read-only @@ -98,6 +103,19 @@ type Client struct { bindingCheckInterval time.Duration } +type errAlternateServerError struct { + code stun.ErrorCodeAttribute + alternate *net.UDPAddr +} + +func (e *errAlternateServerError) Error() string { + if e.alternate != nil { + return fmt.Sprintf("turn: %s (alternate %s)", e.code.String(), e.alternate.String()) + } + + return fmt.Sprintf("turn: %s (alternate server missing)", e.code.String()) +} + // inferAddressFamilyFromConn attempts to determine the address // family (IPv4 or IPv6) from a PacketConn's local address. // Returns an error if the address type is not IP-based. @@ -328,7 +346,80 @@ func (c *Client) SendBindingRequest() (net.Addr, error) { return c.SendBindingRequestTo(c.stunServerAddr) } -func (c *Client) sendAllocateRequest(protocol proto.Protocol) ( //nolint:cyclop +func parseAlternateServer(msg *stun.Message) (*net.UDPAddr, error) { + var alt stun.AlternateServer + if err := alt.GetFrom(msg); err != nil { + return nil, err + } + + return &net.UDPAddr{ + IP: alt.IP, + Port: alt.Port, + }, nil +} + +// handleAllocateError classifies Allocate error responses. +// If allowAuthRetry is true, 401/438 are treated as continuable (so caller can +// extract nonce/realm). Otherwise they are returned as errors. +// Returns proceed=true when the caller should continue. +// Returns an error when the caller should stop (redirect or terminal error). +func handleAllocateError(res *stun.Message, allowAuthRetry bool, allowAlternate bool) (proceed bool, err error) { + if res.Type.Class != stun.ClassErrorResponse { + return true, nil + } + + var code stun.ErrorCodeAttribute + if err = code.GetFrom(res); err != nil { + return false, err + } + + switch code.Code { + case stun.CodeTryAlternate: + if !allowAlternate { + return false, &stun.TurnError{StunMessageType: res.Type, ErrorCodeAttr: code} + } + + if alt, altErr := parseAlternateServer(res); altErr == nil { + return false, &errAlternateServerError{code: code, alternate: alt} + } + + return false, &errAlternateServerError{code: code} + case stun.CodeUnauthorized, stun.CodeStaleNonce: + if allowAuthRetry { + // Continue so nonce/realm can be extracted. + return true, nil + } + + return false, &stun.TurnError{StunMessageType: res.Type, ErrorCodeAttr: code} + default: + return false, &stun.TurnError{StunMessageType: res.Type, ErrorCodeAttr: code} + } +} + +// parseAuthResponse processes an unauthenticated Allocate error response, +// returning nonce and realm when the flow should continue with auth. +func parseAuthResponse(res *stun.Message, allowAlternate bool) (stun.Nonce, stun.Realm, error) { + var nonce stun.Nonce + var realm stun.Realm + + if proceed, err := handleAllocateError(res, true, allowAlternate); err != nil || !proceed { + return nonce, realm, err + } + + if err := nonce.GetFrom(res); err != nil { + return nonce, realm, err + } + if err := realm.GetFrom(res); err != nil { + return nonce, realm, err + } + + // Copy realm to detach from the STUN message buffer. + realm = append([]byte(nil), realm...) + + return nonce, realm, nil +} + +func (c *Client) sendAllocateRequest(protocol proto.Protocol, turnAddr net.Addr, allowAlternate bool) ( //nolint:cyclop relayed proto.RelayedAddress, lifetime proto.Lifetime, nonce stun.Nonce, @@ -356,21 +447,16 @@ func (c *Client) sendAllocateRequest(protocol proto.Protocol) ( //nolint:cyclop return relayed, lifetime, nonce, reservationToken, err } - trRes, err := c.PerformTransaction(msg, c.turnServerAddr, false) + trRes, err := c.PerformTransaction(msg, turnAddr, false) if err != nil { return relayed, lifetime, nonce, reservationToken, err } res := trRes.Msg - // Anonymous allocate failed, trying to authenticate. - if err = nonce.GetFrom(res); err != nil { - return relayed, lifetime, nonce, reservationToken, err - } - if err = c.realm.GetFrom(res); err != nil { + if nonce, c.realm, err = parseAuthResponse(res, allowAlternate); err != nil { return relayed, lifetime, nonce, reservationToken, err } - c.realm = append([]byte(nil), c.realm...) c.integrity = stun.NewLongTermIntegrity( c.username.String(), c.realm.String(), c.password, ) @@ -400,24 +486,14 @@ func (c *Client) sendAllocateRequest(protocol proto.Protocol) ( //nolint:cyclop return relayed, lifetime, nonce, reservationToken, err } - trRes, err = c.PerformTransaction(msg, c.turnServerAddr, false) + trRes, err = c.PerformTransaction(msg, turnAddr, false) if err != nil { return relayed, lifetime, nonce, reservationToken, err } res = trRes.Msg - if res.Type.Class == stun.ClassErrorResponse { - var code stun.ErrorCodeAttribute - if err = code.GetFrom(res); err == nil { - turnError := &stun.TurnError{ - StunMessageType: res.Type, - ErrorCodeAttr: code, - } - - return relayed, lifetime, nonce, reservationToken, turnError - } - - return relayed, lifetime, nonce, reservationToken, fmt.Errorf("%s", res.Type) //nolint:err113 + if proceed, err := handleAllocateError(res, false, allowAlternate); err != nil || !proceed { + return relayed, lifetime, nonce, reservationToken, err } // Getting relayed addresses from response. @@ -440,6 +516,49 @@ func (c *Client) sendAllocateRequest(protocol proto.Protocol) ( //nolint:cyclop return relayed, lifetime, nonce, reservationToken, nil } +// sendAllocateWithRedirect wraps sendAllocateRequest and follows TURN +// ALTERNATE-SERVER (error code 300). +func (c *Client) sendAllocateWithRedirect(protocol proto.Protocol) ( //nolint:cyclop + relayed proto.RelayedAddress, + lifetime proto.Lifetime, + nonce stun.Nonce, + reservationToken proto.ReservationToken, + serverAddr net.Addr, + err error, +) { + currentTurn := c.turnServerAddr + redirects := 0 + visited := map[string]struct{}{} + + for { + relayed, lifetime, nonce, reservationToken, err = c.sendAllocateRequest(protocol, currentTurn, true) + if err == nil { + return relayed, lifetime, nonce, reservationToken, currentTurn, nil + } + + var altErr *errAlternateServerError + if errors.As(err, &altErr) && altErr.alternate != nil { + redirects++ + if redirects > maxAltSrvRedirects { + return relayed, lifetime, nonce, reservationToken, currentTurn, errAlternateServerRedirects + } + + key := altErr.alternate.String() + if _, ok := visited[key]; ok { + return relayed, lifetime, nonce, reservationToken, currentTurn, errAlternateServerLoop + } + + visited[key] = struct{}{} + + currentTurn = altErr.alternate + + continue + } + + return relayed, lifetime, nonce, reservationToken, currentTurn, err + } +} + // Allocate sends a TURN allocation request to the given transport address. func (c *Client) Allocate() (net.PacketConn, error) { if err := c.allocTryLock.Lock(); err != nil { @@ -452,7 +571,7 @@ func (c *Client) Allocate() (net.PacketConn, error) { return nil, fmt.Errorf("%w: %s", errAlreadyAllocated, relayedConn.LocalAddr().String()) } - relayed, lifetime, nonce, reservationToken, err := c.sendAllocateRequest(proto.ProtoUDP) + relayed, lifetime, nonce, reservationToken, serverAddr, err := c.sendAllocateWithRedirect(proto.ProtoUDP) if err != nil { return nil, err } @@ -465,7 +584,7 @@ func (c *Client) Allocate() (net.PacketConn, error) { relayedConn = client.NewUDPConn(&client.AllocationConfig{ Client: c, RelayedAddr: relayedAddr, - ServerAddr: c.turnServerAddr, + ServerAddr: serverAddr, Realm: c.realm, Username: c.username, Integrity: c.integrity, @@ -495,7 +614,7 @@ func (c *Client) AllocateTCP() (*client.TCPAllocation, error) { return nil, fmt.Errorf("%w: %s", errAlreadyAllocated, allocation.Addr()) } - relayed, lifetime, nonce, reservationToken, err := c.sendAllocateRequest(proto.ProtoTCP) + relayed, lifetime, nonce, reservationToken, err := c.sendAllocateRequest(proto.ProtoTCP, c.turnServerAddr, false) if err != nil { return nil, err } diff --git a/client_test.go b/client_test.go index 1ec77531..cd0d6a02 100644 --- a/client_test.go +++ b/client_test.go @@ -8,6 +8,7 @@ package turn import ( "context" + "errors" "io" "net" "runtime" @@ -23,6 +24,57 @@ import ( "github.com/stretchr/testify/require" ) +type allocateRedirectStub struct { + t *testing.T + client *Client + withAlternate bool + redirects int +} + +func (s *allocateRedirectStub) ReadFrom(_ []byte) (int, net.Addr, error) { + return 0, nil, errors.New("not implemented") //nolint:err113 +} + +func (s *allocateRedirectStub) WriteTo(packet []byte, addr net.Addr) (int, error) { + s.redirects++ + + var req stun.Message + req.Raw = append(req.Raw[:0], packet...) + require.NoError(s.t, req.Decode()) + + setters := []stun.Setter{ + &stun.Message{TransactionID: req.TransactionID}, + stun.NewType(stun.MethodAllocate, stun.ClassErrorResponse), + stun.ErrorCodeAttribute{Code: stun.CodeTryAlternate, Reason: []byte("Try Alternate")}, + } + + if s.withAlternate { + setters = append(setters, &stun.AlternateServer{ + IP: net.ParseIP("127.0.0.1"), + Port: 20000 + s.redirects, + }) + } + + setters = append(setters, stun.Fingerprint) + + resp, err := stun.Build(setters...) + require.NoError(s.t, err) + + go func(raw []byte, a net.Addr) { + err := s.client.handleSTUNMessage(raw, a) + + assert.NoError(s.t, err) + }(append([]byte(nil), resp.Raw...), addr) + + return len(packet), nil +} + +func (s *allocateRedirectStub) Close() error { return nil } +func (s *allocateRedirectStub) LocalAddr() net.Addr { return &net.UDPAddr{} } +func (s *allocateRedirectStub) SetDeadline(_ time.Time) error { return nil } +func (s *allocateRedirectStub) SetReadDeadline(_ time.Time) error { return nil } +func (s *allocateRedirectStub) SetWriteDeadline(_ time.Time) error { return nil } + func buildMsg( transactionID [stun.TransactionIDSize]byte, msgType stun.MessageType, @@ -383,6 +435,309 @@ func TestClientReadTimout(t *testing.T) { assert.Contains(t, err.Error(), "use of closed network connection") } +func TestClientAllocateFollowsAlternateServer(t *testing.T) { + loggerFactory := logging.NewDefaultLoggerFactory() + + altListener, err := net.ListenPacket("udp4", "127.0.0.1:0") // nolint: noctx + require.NoError(t, err) + + altServer, err := NewServer(ServerConfig{ + AuthHandler: func(ra *RequestAttributes) (userID string, key []byte, ok bool) { + return ra.Username, GenerateAuthKey(ra.Username, ra.Realm, "pass"), true + }, + PacketConnConfigs: []PacketConnConfig{ + { + PacketConn: altListener, + RelayAddressGenerator: &RelayAddressGeneratorStatic{ + RelayAddress: net.ParseIP("127.0.0.1"), + Address: "0.0.0.0", + }, + }, + }, + Realm: "pion.ly", + }) + require.NoError(t, err) + + redirectListener, err := net.ListenPacket("udp4", "127.0.0.1:0") // nolint: noctx + require.NoError(t, err) + + redirectDone := make(chan struct{}) + go func() { //nolint:dupl + defer close(redirectDone) + buf := make([]byte, 1500) + for { + n, from, readErr := redirectListener.ReadFrom(buf) + if readErr != nil { + return + } + + var req stun.Message + req.Raw = append(req.Raw[:0], buf[:n]...) + if decodeErr := req.Decode(); decodeErr != nil { + continue + } + + altAddr, ok := altListener.LocalAddr().(*net.UDPAddr) + assert.True(t, ok) + resp, buildErr := stun.Build( + &stun.Message{TransactionID: req.TransactionID}, + stun.NewType(stun.MethodAllocate, stun.ClassErrorResponse), + stun.ErrorCodeAttribute{Code: stun.CodeTryAlternate, Reason: []byte("Try Alternate")}, + &stun.AlternateServer{IP: altAddr.IP, Port: altAddr.Port}, + stun.Fingerprint, + ) + if buildErr != nil { + continue + } + + _, _ = redirectListener.WriteTo(resp.Raw, from) + } + }() + + clientConn, err := net.ListenPacket("udp4", "0.0.0.0:0") // nolint: noctx + require.NoError(t, err) + + redirectAddr := redirectListener.LocalAddr().String() + + client, err := NewClient(&ClientConfig{ + Conn: clientConn, + STUNServerAddr: redirectAddr, + TURNServerAddr: redirectAddr, + Username: "foo", + Password: "pass", + Realm: "pion.ly", + LoggerFactory: loggerFactory, + }) + require.NoError(t, err) + require.NoError(t, client.Listen()) + origServer := client.TURNServerAddr().String() + + allocation, err := client.Allocate() + require.NoError(t, err) + require.NotNil(t, allocation) + require.Equal(t, origServer, client.TURNServerAddr().String()) + + require.NoError(t, client.CreatePermission(&net.UDPAddr{ + IP: net.ParseIP("127.0.0.1"), + Port: 30000, + })) + + assert.NoError(t, allocation.Close()) + assert.NoError(t, clientConn.Close()) + assert.NoError(t, redirectListener.Close()) + assert.NoError(t, altServer.Close()) + <-redirectDone +} + +func TestClientAlternateServerLoop(t *testing.T) { + loggerFactory := logging.NewDefaultLoggerFactory() + + redirectListener, err := net.ListenPacket("udp4", "127.0.0.1:0") // nolint: noctx + require.NoError(t, err) + + redirectDone := make(chan struct{}) + go func() { //nolint:dupl + defer close(redirectDone) + buf := make([]byte, 1500) + for { + n, from, readErr := redirectListener.ReadFrom(buf) + if readErr != nil { + return + } + + var req stun.Message + req.Raw = append(req.Raw[:0], buf[:n]...) + if decodeErr := req.Decode(); decodeErr != nil { + continue + } + + selfAddr, ok := redirectListener.LocalAddr().(*net.UDPAddr) + assert.True(t, ok) + resp, buildErr := stun.Build( + &stun.Message{TransactionID: req.TransactionID}, + stun.NewType(stun.MethodAllocate, stun.ClassErrorResponse), + stun.ErrorCodeAttribute{Code: stun.CodeTryAlternate, Reason: []byte("Try Alternate")}, + &stun.AlternateServer{IP: selfAddr.IP, Port: selfAddr.Port}, + stun.Fingerprint, + ) + if buildErr != nil { + continue + } + + _, _ = redirectListener.WriteTo(resp.Raw, from) + } + }() + + clientConn, err := net.ListenPacket("udp4", "0.0.0.0:0") // nolint: noctx + require.NoError(t, err) + + redirectAddr := redirectListener.LocalAddr().String() + + client, err := NewClient(&ClientConfig{ + Conn: clientConn, + STUNServerAddr: redirectAddr, + TURNServerAddr: redirectAddr, + Username: "foo", + Password: "pass", + Realm: "pion.ly", + LoggerFactory: loggerFactory, + }) + require.NoError(t, err) + require.NoError(t, client.Listen()) + + _, err = client.Allocate() + require.Error(t, err) + assert.True(t, errors.Is(err, errAlternateServerLoop)) + + assert.NoError(t, clientConn.Close()) + assert.NoError(t, redirectListener.Close()) + <-redirectDone +} + +func TestClientAlternateServerRedirectLimit(t *testing.T) { + loggerFactory := logging.NewDefaultLoggerFactory() + + stub := &allocateRedirectStub{t: t, withAlternate: true} + + client, err := NewClient(&ClientConfig{ + Conn: stub, + STUNServerAddr: "127.0.0.1:3478", + TURNServerAddr: "127.0.0.1:3478", + Username: "foo", + Password: "pass", + Realm: "pion.ly", + LoggerFactory: loggerFactory, + }) + require.NoError(t, err) + stub.client = client + + _, err = client.Allocate() + require.Error(t, err) + assert.True(t, errors.Is(err, errAlternateServerRedirects)) + assert.Equal(t, maxAltSrvRedirects+1, stub.redirects) +} + +func TestClientAlternateServerMissingAttribute(t *testing.T) { + loggerFactory := logging.NewDefaultLoggerFactory() + + stub := &allocateRedirectStub{t: t, withAlternate: false} + + client, err := NewClient(&ClientConfig{ + Conn: stub, + STUNServerAddr: "127.0.0.1:3478", + TURNServerAddr: "127.0.0.1:3478", + Username: "foo", + Password: "pass", + Realm: "pion.ly", + LoggerFactory: loggerFactory, + }) + require.NoError(t, err) + stub.client = client + + _, err = client.Allocate() + require.Error(t, err) + + var altErr *errAlternateServerError + require.True(t, errors.As(err, &altErr)) + assert.Nil(t, altErr.alternate) +} + +func TestClientTCPAllocateFollowsAlternateServer(t *testing.T) { + loggerFactory := logging.NewDefaultLoggerFactory() + + altListener, err := net.Listen("tcp4", "127.0.0.1:0") // nolint: gosec,noctx + require.NoError(t, err) + + altServer, err := NewServer(ServerConfig{ + AuthHandler: func(ra *RequestAttributes) (userID string, key []byte, ok bool) { + return ra.Username, GenerateAuthKey(ra.Username, ra.Realm, "pass"), true + }, + ListenerConfigs: []ListenerConfig{ + { + Listener: altListener, + RelayAddressGenerator: &RelayAddressGeneratorStatic{ + RelayAddress: net.ParseIP("127.0.0.1"), + Address: "0.0.0.0", + }, + }, + }, + Realm: "pion.ly", + }) + require.NoError(t, err) + + redirectListener, err := net.Listen("tcp4", "127.0.0.1:0") // nolint: gosec,noctx + require.NoError(t, err) + + redirectDone := make(chan struct{}) + go func() { + defer close(redirectDone) + + conn, acceptErr := redirectListener.Accept() + if acceptErr != nil { + return + } + defer conn.Close() //nolint:errcheck + + buf := make([]byte, 1500) + for { + n, readErr := conn.Read(buf) + if readErr != nil { + return + } + + var req stun.Message + req.Raw = append(req.Raw[:0], buf[:n]...) + if decodeErr := req.Decode(); decodeErr != nil { + continue + } + + altAddr, ok := altListener.Addr().(*net.TCPAddr) + assert.True(t, ok) + resp, buildErr := stun.Build( + &stun.Message{TransactionID: req.TransactionID}, + stun.NewType(stun.MethodAllocate, stun.ClassErrorResponse), + stun.ErrorCodeAttribute{Code: stun.CodeTryAlternate, Reason: []byte("Try Alternate")}, + &stun.AlternateServer{IP: altAddr.IP, Port: altAddr.Port}, + stun.Fingerprint, + ) + if buildErr != nil { + continue + } + + _, _ = conn.Write(resp.Raw) + } + }() + + clientConn, err := net.Dial("tcp4", redirectListener.Addr().String()) // nolint: gosec,noctx + require.NoError(t, err) + + redirectAddr := redirectListener.Addr().String() + + client, err := NewClient(&ClientConfig{ + Conn: NewSTUNConn(clientConn), + STUNServerAddr: redirectAddr, + TURNServerAddr: redirectAddr, + Username: "foo", + Password: "pass", + Realm: "pion.ly", + LoggerFactory: loggerFactory, + }) + require.NoError(t, err) + require.NoError(t, client.Listen()) + + allocation, err := client.AllocateTCP() + require.Error(t, err) + require.Nil(t, allocation) + var turnErr *stun.TurnError + require.True(t, errors.As(err, &turnErr)) + require.Equal(t, stun.CodeTryAlternate, turnErr.ErrorCodeAttr.Code) + + assert.NoError(t, clientConn.Close()) + assert.NoError(t, redirectListener.Close()) + assert.NoError(t, altServer.Close()) + <-redirectDone +} + func TestTCPClientDial(t *testing.T) { tcpListener, err := net.Listen("tcp4", "0.0.0.0:3478") //nolint: gosec,noctx require.NoError(t, err) diff --git a/errors.go b/errors.go index 8e7a6ede..e5622ddc 100644 --- a/errors.go +++ b/errors.go @@ -28,4 +28,6 @@ var ( errFailedToDecodeSTUN = errors.New("failed to decode STUN message") errUnexpectedSTUNRequestMessage = errors.New("unexpected STUN request message") errRelayAddressGeneratorNil = errors.New("RelayAddressGenerator is nil") + errAlternateServerRedirects = errors.New("turn: too many ALTERNATE-SERVER redirects") + errAlternateServerLoop = errors.New("turn: ALTERNATE-SERVER redirect loop detected") )