package gateway import ( "context" "crypto/sha256" "crypto/tls" "crypto/x509" "encoding/binary" "encoding/hex" "encoding/pem" "io" "net" "net/http" "net/http/httptest" "strconv" "testing" protocol "git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol/gen/go/protocol" ) func TestNativeApolloManagementUsesSessionScopedMTLS(t *testing.T) { serverTLS, clientTLS := testTLS(t) server := httptest.NewUnstartedServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { if request.URL.Path != "/serverinfo" || request.TLS == nil || len(request.TLS.PeerCertificates) != 1 { http.Error(response, "mTLS required", http.StatusUnauthorized) return } _, _ = response.Write([]byte("apollo-server")) })) server.TLS = serverTLS server.StartTLS() defer server.Close() host, portText, err := net.SplitHostPort(server.Listener.Addr().String()) if err != nil { t.Fatal(err) } port, err := strconv.ParseInt(portText, 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: "1", ClientID: "paired-client", ManagementHost: host, ManagementPort: port, StreamHost: host, StreamPort: 47984, ClientCertificatePem: certificatePEM(t, clientTLS.Certificates[0]), ClientPrivateKeyPem: privateKeyPEM(t, clientTLS.Certificates[0]), ServerCertificatePem: certificatePEM(t, tls.Certificate{Certificate: [][]byte{serverTLS.Certificates[0].Certificate[1]}}), } data, err := NewNativeApolloBackend().Management(context.Background(), LaunchRequest{SessionID: "session-1", ProviderProfile: ProviderProfileApollo, ProviderWork: work}) if err != nil { t.Fatalf("Management() error = %v", err) } if info, err := ParseManagementXML(data); err != nil || info.Identity.UniqueID != "apollo-server" { t.Fatalf("ParseManagementXML() = %+v, %v", info, err) } } 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 { t.Fatal("expected certificate") } return string(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certificate.Certificate[0]})) } func privateKeyPEM(t *testing.T, certificate tls.Certificate) string { t.Helper() encoded, err := x509.MarshalPKCS8PrivateKey(certificate.PrivateKey) if err != nil { t.Fatal(err) } return string(pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: encoded})) }