188 lines
7.2 KiB
Go
188 lines
7.2 KiB
Go
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("<root><uniqueid>apollo-server</uniqueid></root>"))
|
|
}))
|
|
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("<root><App><ID>42</ID></App></root>"))
|
|
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("<root status_code=\"200\"><sessionUrl0>rtspenc://" + streamListener.Addr().String() + "</sessionUrl0></root>"))
|
|
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}))
|
|
}
|