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
}