diff --git a/gateway/apollo_native.go b/gateway/apollo_native.go index 57e9c45..338157e 100644 --- a/gateway/apollo_native.go +++ b/gateway/apollo_native.go @@ -1,17 +1,19 @@ package gateway import ( - "bufio" "bytes" "context" + "crypto/rand" "crypto/sha256" "crypto/tls" "crypto/x509" + "encoding/binary" "encoding/hex" "fmt" "io" "net" "net/http" + "net/url" "strconv" "strings" "sync" @@ -36,20 +38,32 @@ func NewNativeApolloBackend() *NativeApolloBackend { func (b *NativeApolloBackend) Management(ctx context.Context, request LaunchRequest) ([]byte, error) { work := request.ProviderWork - if err := work.Validate(); err != nil || work.ProviderProfile != ProviderProfileApollo { + if err := work.Validate(); err != nil || request.SessionID == "" || request.SessionID != work.SessionID || work.ProviderProfile != ProviderProfileApollo { return nil, ErrProviderMalformed } + client, err := newPinnedApolloHTTPClient(work) + if err != nil { + return nil, err + } + return apolloGet(ctx, client, work, "/serverinfo", nil) +} + +func newPinnedApolloHTTPClient(work protocol.ProviderSessionWork) (*http.Client, error) { tlsConfig, err := pinnedApolloTLSConfig(work) if err != nil { return nil, err } - client := &http.Client{Transport: &http.Transport{TLSClientConfig: tlsConfig}, Timeout: 5 * time.Second} - managementURL := "https://" + net.JoinHostPort(work.ManagementHost, strconv.FormatInt(work.ManagementPort, 10)) + "/serverinfo" - httpRequest, err := http.NewRequestWithContext(ctx, http.MethodGet, managementURL, nil) + return &http.Client{Transport: &http.Transport{TLSClientConfig: tlsConfig}, Timeout: 5 * time.Second}, nil +} + +func apolloGet(ctx context.Context, client *http.Client, work protocol.ProviderSessionWork, path string, values url.Values) ([]byte, error) { + endpoint := url.URL{Scheme: "https", Host: net.JoinHostPort(work.ManagementHost, strconv.FormatInt(work.ManagementPort, 10)), Path: path} + endpoint.RawQuery = values.Encode() + request, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint.String(), nil) if err != nil { return nil, ErrProviderMalformed } - response, err := client.Do(httpRequest) + response, err := client.Do(request) if err != nil { return nil, err } @@ -94,7 +108,43 @@ func pinnedApolloTLSConfig(work protocol.ProviderSessionWork) (*tls.Config, erro func (b *NativeApolloBackend) Setup(ctx context.Context, request LaunchRequest) ([]byte, error) { work := request.ProviderWork - if err := work.Validate(); err != nil || request.SessionID == "" { + if err := work.Validate(); err != nil || request.SessionID == "" || request.SessionID != work.SessionID || work.ProviderProfile != ProviderProfileApollo { + return nil, ErrProviderMalformed + } + client, err := newPinnedApolloHTTPClient(work) + if err != nil { + return nil, err + } + inventory, err := apolloGet(ctx, client, work, "/applist", url.Values{"uniqueid": []string{work.ClientID}}) + if err != nil || !apolloInventoryContains(inventory, work.ApplicationID) { + return nil, ErrProviderMalformed + } + key := make([]byte, 16) + if _, err := rand.Read(key); err != nil { + return nil, err + } + var keyID [4]byte + if _, err := rand.Read(keyID[:]); err != nil { + return nil, err + } + launch, err := apolloGet(ctx, client, work, "/launch", url.Values{ + "uniqueid": {work.ClientID}, "appid": {work.ApplicationID}, "rikey": {hex.EncodeToString(key)}, + "rikeyid": {strconv.FormatUint(uint64(binary.BigEndian.Uint32(keyID[:])), 10)}, "localAudioPlayMode": {"0"}, + "corever": {"1"}, + }) + if err != nil { + return nil, err + } + launchResponse, err := parseApolloLaunchResponse(launch) + if err != nil { + return nil, err + } + streamURL, err := url.Parse(launchResponse.SessionURL) + if err != nil || streamURL.Scheme != "rtspenc" || streamURL.Hostname() != work.StreamHost { + return nil, ErrProviderMalformed + } + streamPort, err := strconv.ParseInt(streamURL.Port(), 10, 64) + if err != nil || streamPort != work.StreamPort { return nil, ErrProviderMalformed } conn, err := b.Dialer.DialContext(ctx, "tcp", net.JoinHostPort(work.StreamHost, strconv.FormatInt(work.StreamPort, 10))) @@ -104,12 +154,22 @@ func (b *NativeApolloBackend) Setup(ctx context.Context, request LaunchRequest) if deadline, ok := ctx.Deadline(); ok { _ = conn.SetDeadline(deadline) } - requestText := "SETUP rtsp://" + work.StreamHost + "/streamid=video/0/0 RTSP/1.0\r\nCSeq: 1\r\nTransport: RTP/AVP/TCP;interleaved=0-1\r\nSession: " + request.SessionID + "\r\n\r\n" - if _, err := io.WriteString(conn, requestText); err != nil { + codec, err := newEncryptedRTSPCodec(key) + if err != nil { _ = conn.Close() return nil, err } - response, err := readRTSPHeaders(conn, 16*1024) + requestText := "SETUP rtsp://" + work.StreamHost + "/streamid=video/0/0 RTSP/1.0\r\nCSeq: 1\r\nTransport: RTP/AVP/TCP;interleaved=0-1\r\n\r\n" + encoded, err := codec.SealClient([]byte(requestText)) + if err != nil { + _ = conn.Close() + return nil, err + } + if _, err := conn.Write(encoded); err != nil { + _ = conn.Close() + return nil, err + } + response, err := readEncryptedRTSPHeaders(conn, codec) if err != nil { _ = conn.Close() return nil, err @@ -122,10 +182,10 @@ func (b *NativeApolloBackend) Setup(ctx context.Context, request LaunchRequest) func (b *NativeApolloBackend) Open(_ context.Context, request LaunchRequest, _ RTSPResponse) (ProviderSession, error) { b.mu.Lock() - conn := b.pending[request.SessionID] + conn, ok := b.pending[request.SessionID] delete(b.pending, request.SessionID) b.mu.Unlock() - if conn == nil { + if !ok || conn == nil { return nil, ErrProviderDisconnected } session := newNativeApolloSession(conn, request.SessionID) @@ -144,20 +204,25 @@ func readBounded(reader io.Reader, max int) ([]byte, error) { return data, nil } -func readRTSPHeaders(conn net.Conn, max int) ([]byte, error) { - reader := bufio.NewReaderSize(conn, 4096) - var response []byte - for len(response) < max { - line, err := reader.ReadBytes('\n') - if err != nil { - return nil, err - } - response = append(response, line...) - if strings.HasSuffix(string(response), "\r\n\r\n") { - return response, nil - } +func readEncryptedRTSPHeaders(conn net.Conn, codec *encryptedRTSPCodec) ([]byte, error) { + header := make([]byte, encryptedRTSPHeaderSize) + if _, err := io.ReadFull(conn, header); err != nil { + return nil, err } - return nil, ErrProviderMalformed + length := binary.BigEndian.Uint32(header[:4]) & 0x7fffffff + if length == 0 || length > encryptedRTSPMaxPayload { + return nil, ErrProviderMalformed + } + frame := make([]byte, encryptedRTSPHeaderSize+int(length)) + copy(frame, header) + if _, err := io.ReadFull(conn, frame[encryptedRTSPHeaderSize:]); err != nil { + return nil, err + } + plaintext, err := codec.OpenHost(frame) + if err != nil || len(plaintext) > 16*1024 || !strings.HasSuffix(string(plaintext), "\r\n\r\n") { + return nil, ErrProviderMalformed + } + return plaintext, nil } type nativeApolloSession struct { diff --git a/gateway/apollo_native_test.go b/gateway/apollo_native_test.go index 75e1c52..bcb3717 100644 --- a/gateway/apollo_native_test.go +++ b/gateway/apollo_native_test.go @@ -5,8 +5,10 @@ import ( "crypto/sha256" "crypto/tls" "crypto/x509" + "encoding/binary" "encoding/hex" "encoding/pem" + "io" "net" "net/http" "net/http/httptest" @@ -55,6 +57,118 @@ func TestNativeApolloManagementUsesSessionScopedMTLS(t *testing.T) { } } +func TestNativeApolloSetupLaunchesInventoriedApplicationWithEncryptedRTSP(t *testing.T) { + serverTLS, clientTLS := testTLS(t) + streamListener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer streamListener.Close() + streamHost, streamPortText, err := net.SplitHostPort(streamListener.Addr().String()) + if err != nil { + t.Fatal(err) + } + streamPort, err := strconv.ParseInt(streamPortText, 10, 64) + if err != nil { + t.Fatal(err) + } + keyReady := make(chan []byte, 1) + streamDone := make(chan error, 1) + go func() { + connection, acceptErr := streamListener.Accept() + if acceptErr != nil { + streamDone <- acceptErr + return + } + defer connection.Close() + key := <-keyReady + header := make([]byte, encryptedRTSPHeaderSize) + if _, readErr := io.ReadFull(connection, header); readErr != nil { + streamDone <- readErr + return + } + length := binary.BigEndian.Uint32(header[:4]) & 0x7fffffff + frame := make([]byte, encryptedRTSPHeaderSize+int(length)) + copy(frame, header) + if _, readErr := io.ReadFull(connection, frame[encryptedRTSPHeaderSize:]); readErr != nil { + streamDone <- readErr + return + } + codec, codecErr := newEncryptedRTSPCodec(key) + if codecErr != nil { + streamDone <- codecErr + return + } + plaintext, decryptErr := codec.OpenClient(frame) + if decryptErr != nil || string(plaintext) != "SETUP rtsp://"+streamHost+"/streamid=video/0/0 RTSP/1.0\r\nCSeq: 1\r\nTransport: RTP/AVP/TCP;interleaved=0-1\r\n\r\n" { + streamDone <- ErrProviderMalformed + return + } + _, writeErr := connection.Write(hostEncryptedRTSPFrame(t, key, 1, []byte("RTSP/1.0 200 OK\r\nSession: host-session\r\nTransport: RTP/AVP/TCP;interleaved=0-1\r\n\r\n"))) + streamDone <- writeErr + }() + management := httptest.NewUnstartedServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + if request.TLS == nil || len(request.TLS.PeerCertificates) != 1 { + http.Error(response, "mTLS required", http.StatusUnauthorized) + return + } + switch request.URL.Path { + case "/applist": + if request.URL.Query().Get("uniqueid") != "paired-client" { + http.Error(response, "wrong client", http.StatusBadRequest) + return + } + _, _ = response.Write([]byte("42")) + case "/launch": + query := request.URL.Query() + if query.Get("uniqueid") != "paired-client" || query.Get("appid") != "42" || query.Get("corever") != "1" || query.Get("localAudioPlayMode") != "0" { + http.Error(response, "wrong launch", http.StatusBadRequest) + return + } + key, decodeErr := hex.DecodeString(query.Get("rikey")) + if decodeErr != nil || len(key) != 16 || query.Get("rikeyid") == "" { + http.Error(response, "bad key", http.StatusBadRequest) + return + } + keyReady <- key + _, _ = response.Write([]byte("rtspenc://" + streamListener.Addr().String() + "")) + default: + http.NotFound(response, request) + } + })) + management.TLS = serverTLS + management.StartTLS() + defer management.Close() + managementHost, managementPortText, err := net.SplitHostPort(management.Listener.Addr().String()) + if err != nil { + t.Fatal(err) + } + managementPort, err := strconv.ParseInt(managementPortText, 10, 64) + if err != nil { + t.Fatal(err) + } + pinned := sha256.Sum256(serverTLS.Certificates[0].Certificate[0]) + work := protocol.ProviderSessionWork{ + Version: "1", SessionID: "session-1", GatewayID: "gateway-1", ReconnectSequence: 0, + ExpiresAt: "2099-01-01T00:00:00Z", ProviderProfile: ProviderProfileApollo, + ProviderIdentity: "apollo-server#sha256:" + hex.EncodeToString(pinned[:]), PolicyVersionID: "policy-1", ApplicationID: "42", ClientID: "paired-client", + ManagementHost: managementHost, ManagementPort: managementPort, StreamHost: streamHost, StreamPort: streamPort, + ClientCertificatePem: certificatePEM(t, clientTLS.Certificates[0]), ClientPrivateKeyPem: privateKeyPEM(t, clientTLS.Certificates[0]), + ServerCertificatePem: certificatePEM(t, tls.Certificate{Certificate: [][]byte{serverTLS.Certificates[0].Certificate[1]}}), + } + backend := NewNativeApolloBackend() + response, err := backend.Setup(context.Background(), LaunchRequest{SessionID: "session-1", ProviderProfile: ProviderProfileApollo, ProviderWork: work}) + if err != nil { + t.Fatalf("Setup() error = %v", err) + } + if parsed, err := ParseRTSPResponse(response); err != nil || parsed.Session != "host-session" { + t.Fatalf("ParseRTSPResponse() = %+v, %v", parsed, err) + } + if err := <-streamDone; err != nil { + t.Fatalf("stream transaction error = %v", err) + } +} + func certificatePEM(t *testing.T, certificate tls.Certificate) string { t.Helper() if len(certificate.Certificate) == 0 { diff --git a/gateway/apollo_rtsp.go b/gateway/apollo_rtsp.go index 5d6c5a8..a832fb0 100644 --- a/gateway/apollo_rtsp.go +++ b/gateway/apollo_rtsp.go @@ -18,10 +18,12 @@ var errEncryptedRTSPFrame = errors.New("invalid encrypted RTSP frame") // only strictly increasing host sequence numbers, so a replay cannot be fed // into the RTSP parser after it has already authenticated once. type encryptedRTSPCodec struct { - aead cipher.AEAD - nextClient uint32 - lastHost uint32 - hostReceived bool + aead cipher.AEAD + nextClient uint32 + lastClient uint32 + clientReceived bool + lastHost uint32 + hostReceived bool } func newEncryptedRTSPCodec(key []byte) (*encryptedRTSPCodec, error) { @@ -54,6 +56,17 @@ func (codec *encryptedRTSPCodec) SealClient(plaintext []byte) ([]byte, error) { } func (codec *encryptedRTSPCodec) OpenHost(frame []byte) ([]byte, error) { + return codec.open(frame, 'H', 'R', &codec.lastHost, &codec.hostReceived) +} + +// OpenClient validates the client-originated direction. The native gateway +// only sends this direction, but retaining the inverse lets a bounded fake +// provider verify the negotiated session key and frame shape. +func (codec *encryptedRTSPCodec) OpenClient(frame []byte) ([]byte, error) { + return codec.open(frame, 'C', 'R', &codec.lastClient, &codec.clientReceived) +} + +func (codec *encryptedRTSPCodec) open(frame []byte, origin, protocol byte, last *uint32, received *bool) ([]byte, error) { if codec == nil || codec.aead == nil || len(frame) < encryptedRTSPHeaderSize { return nil, errEncryptedRTSPFrame } @@ -62,10 +75,10 @@ func (codec *encryptedRTSPCodec) OpenHost(frame []byte) ([]byte, error) { return nil, errEncryptedRTSPFrame } sequence := binary.BigEndian.Uint32(frame[4:8]) - if sequence == 0 || (codec.hostReceived && sequence <= codec.lastHost) { + if sequence == 0 || (*received && sequence <= *last) { return nil, errEncryptedRTSPFrame } - nonce := encryptedRTSPNonce(sequence, 'H', 'R') + nonce := encryptedRTSPNonce(sequence, origin, protocol) sealed := make([]byte, int(length&0x7fffffff)+codec.aead.Overhead()) copy(sealed, frame[24:]) copy(sealed[length&0x7fffffff:], frame[8:24]) @@ -73,7 +86,7 @@ func (codec *encryptedRTSPCodec) OpenHost(frame []byte) ([]byte, error) { if err != nil { return nil, errEncryptedRTSPFrame } - codec.lastHost, codec.hostReceived = sequence, true + *last, *received = sequence, true return plaintext, nil }