1975 lines
75 KiB
Go
1975 lines
75 KiB
Go
package gateway
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"crypto/ecdsa"
|
|
"crypto/elliptic"
|
|
"crypto/rand"
|
|
"crypto/tls"
|
|
"crypto/x509"
|
|
"crypto/x509/pkix"
|
|
"encoding/base64"
|
|
"encoding/binary"
|
|
"encoding/hex"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"math/big"
|
|
"net"
|
|
"os"
|
|
"reflect"
|
|
"strings"
|
|
"sync"
|
|
"sync/atomic"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/quic-go/quic-go"
|
|
|
|
protocol "git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol/gen/go/protocol"
|
|
)
|
|
|
|
const protocolTerminalReceiptVector = "VGF1\x00\x03\x00\x00"
|
|
|
|
func TestFrameValidationAndFragmentation(t *testing.T) {
|
|
frames, err := FragmentPayload(ChannelVideo, 7, 11, make([]byte, 1180))
|
|
if err != nil || len(frames) != 2 || len(frames[0].Payload) != 1177 || len(frames[1].Payload) != 3 {
|
|
t.Fatalf("fragmentation = %#v, err = %v", frames, err)
|
|
}
|
|
encoded, err := EncodeFrame(frames[0])
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
decoded, err := DecodeFrame(encoded)
|
|
if err != nil || string(decoded.Payload) != string(frames[0].Payload) {
|
|
t.Fatalf("decoded = %#v, err = %v", decoded, err)
|
|
}
|
|
for _, raw := range [][]byte{
|
|
{0x56, 0x44},
|
|
append([]byte(nil), encoded[:len(encoded)-1]...),
|
|
append(append([]byte(nil), encoded...), 0),
|
|
} {
|
|
if err := ValidateFrame(raw); err == nil {
|
|
t.Fatalf("accepted malformed frame %x", raw)
|
|
}
|
|
}
|
|
}
|
|
|
|
func FuzzDecodeFrame(f *testing.F) {
|
|
seed, _ := hex.DecodeString("5644010a0000000000000000000000000000010000")
|
|
f.Add(seed)
|
|
f.Add([]byte("not-a-frame"))
|
|
f.Fuzz(func(t *testing.T, data []byte) {
|
|
_, _ = DecodeFrame(data)
|
|
})
|
|
}
|
|
|
|
func FuzzDecodeInputEvent(f *testing.F) {
|
|
seed, _ := EncodeInputEvent(InputEvent{Sequence: 1, Device: "keyboard", Code: 7, Pressed: true})
|
|
f.Add(seed)
|
|
absolute, _ := EncodeInputEvent(InputEvent{Device: "mouse-absolute", Payload: []byte{0, 1, 0, 1, 0, 2, 0, 2}})
|
|
f.Add(absolute)
|
|
scroll, _ := EncodeInputEvent(InputEvent{Device: "mouse-scroll", Payload: []byte{0xff, 0x88, 0, 0x78}})
|
|
f.Add(scroll)
|
|
f.Add([]byte("VGI1"))
|
|
f.Fuzz(func(t *testing.T, data []byte) {
|
|
_, _ = DecodeInputEvent(data)
|
|
})
|
|
}
|
|
|
|
func TestInputEventUsesFixedProtocolVGI1Vector(t *testing.T) {
|
|
encoded, err := EncodeInputEvent(InputEvent{Sequence: 7, Device: "keyboard", Code: 30, Pressed: true, Payload: []byte{2}})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
const expected = "5647493101040102001e"
|
|
if hex.EncodeToString(encoded) != expected {
|
|
t.Fatalf("EncodeInputEvent() = %x, want %s", encoded, expected)
|
|
}
|
|
decoded, err := DecodeInputEvent(encoded)
|
|
if err != nil || decoded.Sequence != 0 || decoded.Device != "keyboard" || decoded.Code != 30 || !decoded.Pressed || string(decoded.Payload) != string([]byte{2}) {
|
|
t.Fatalf("DecodeInputEvent() = %#v, %v", decoded, err)
|
|
}
|
|
}
|
|
|
|
func TestInputEventUsesFixedAbsoluteAndScrollVectors(t *testing.T) {
|
|
tests := []struct {
|
|
name string
|
|
event InputEvent
|
|
expected string
|
|
}{
|
|
{
|
|
name: "absolute",
|
|
event: InputEvent{Device: "mouse-absolute", Payload: []byte{0x04, 0xd2, 0x02, 0x37, 0x0a, 0x00, 0x05, 0xa0}},
|
|
expected: "56474931060804d202370a0005a0",
|
|
},
|
|
{
|
|
name: "scroll",
|
|
event: InputEvent{Device: "mouse-scroll", Payload: []byte{0xff, 0x88, 0x00, 0x78}},
|
|
expected: "564749310704ff880078",
|
|
},
|
|
}
|
|
for _, test := range tests {
|
|
t.Run(test.name, func(t *testing.T) {
|
|
encoded, err := EncodeInputEvent(test.event)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if hex.EncodeToString(encoded) != test.expected {
|
|
t.Fatalf("EncodeInputEvent() = %x, want %s", encoded, test.expected)
|
|
}
|
|
decoded, err := DecodeInputEvent(encoded)
|
|
if err != nil || decoded.Device != test.event.Device || decoded.Code != 0 || decoded.Pressed || !bytes.Equal(decoded.Payload, test.event.Payload) {
|
|
t.Fatalf("DecodeInputEvent() = %#v, %v", decoded, err)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestInputEventRejectsMalformedAbsoluteAndScroll(t *testing.T) {
|
|
for _, event := range []InputEvent{
|
|
{Device: "mouse-absolute", Pressed: true, Payload: make([]byte, 8)},
|
|
{Device: "mouse-absolute", Payload: []byte{0, 0, 0, 0, 0, 0, 0, 1}},
|
|
{Device: "mouse-absolute", Payload: []byte{0, 2, 0, 0, 0, 2, 0, 1}},
|
|
{Device: "mouse-absolute", Payload: []byte{0, 0, 0, 1, 0, 2, 0, 1}},
|
|
{Device: "mouse-scroll", Pressed: true, Payload: make([]byte, 4)},
|
|
{Device: "mouse-scroll", Payload: make([]byte, 3)},
|
|
} {
|
|
if _, err := EncodeInputEvent(event); !errors.Is(err, ErrInputMalformed) {
|
|
t.Fatalf("EncodeInputEvent(%#v) = %v", event, err)
|
|
}
|
|
}
|
|
for _, value := range []string{
|
|
"56474931060800000000000005a0",
|
|
"5647493106080a0000000a0005a0",
|
|
"564749310608000005a00a0005a0",
|
|
"56474931060700000000010001",
|
|
"5647493107020000",
|
|
} {
|
|
raw, err := hex.DecodeString(value)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, err := DecodeInputEvent(raw); !errors.Is(err, ErrInputMalformed) {
|
|
t.Fatalf("DecodeInputEvent(%s) = %v", value, err)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestGatewayRejectsUnadvertisedInputBeforeProviderTranslation(t *testing.T) {
|
|
provider := &fakeSession{
|
|
state: protocol.ProviderState{State: ProviderStateReady},
|
|
pressed: make(map[string]struct{}),
|
|
}
|
|
session := &gatewaySession{
|
|
server: &Server{config: ServerConfig{}},
|
|
provider: provider,
|
|
ctx: context.Background(),
|
|
pressed: make(map[string]struct{}),
|
|
}
|
|
absolute, err := EncodeInputEvent(InputEvent{Device: "mouse-absolute", Payload: []byte{0, 1, 0, 1, 0, 2, 0, 2}})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := session.handleInput(absolute, 1); !errors.Is(err, ErrInputMalformed) {
|
|
t.Fatalf("unadvertised absolute input = %v", err)
|
|
}
|
|
if len(provider.inputs) != 0 {
|
|
t.Fatalf("unadvertised absolute input reached provider: %#v", provider.inputs)
|
|
}
|
|
session.server.config.Features = []string{"input.absolute.v1"}
|
|
if err := session.handleInput(absolute, 2); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
scroll, err := EncodeInputEvent(InputEvent{Device: "mouse-scroll", Payload: []byte{0, 1, 0, 1}})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := session.handleInput(scroll, 3); !errors.Is(err, ErrInputMalformed) {
|
|
t.Fatalf("unadvertised scroll input = %v", err)
|
|
}
|
|
if len(provider.inputs) != 1 || len(session.pressed) != 0 {
|
|
t.Fatalf("provider inputs=%#v pressed=%#v", provider.inputs, session.pressed)
|
|
}
|
|
}
|
|
|
|
func TestClientFeedbackUsesFixedProtocolVGFVector(t *testing.T) {
|
|
feedback := Feedback{Sequence: 9, Kind: FeedbackFEC, Payload: []byte{0, 0, 0, 42, 0, 5, 0, 3, 0, 2, 0, 10, 0, 2, 0, 8, 0, 2, 20, 0, 1}}
|
|
encoded, err := EncodeClientFeedback(feedback)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
const expected = "56474631000200150000002a000500030002000a000200080002140001"
|
|
if hex.EncodeToString(encoded) != expected {
|
|
t.Fatalf("EncodeClientFeedback() = %x, want %s", encoded, expected)
|
|
}
|
|
decoded, err := DecodeClientFeedback(encoded)
|
|
if err != nil || decoded.Kind != FeedbackFEC || string(decoded.Payload) != string(feedback.Payload) {
|
|
t.Fatalf("DecodeClientFeedback() = %#v, %v", decoded, err)
|
|
}
|
|
if _, err := DecodeClientFeedback([]byte{'F', 'B', 'R', 'K', 0}); !errors.Is(err, ErrProviderMalformed) {
|
|
t.Fatalf("legacy feedback accepted: %v", err)
|
|
}
|
|
disconnected, err := EncodeProviderEvent(ProviderEvent{Kind: ProviderEventDisconnected})
|
|
if err != nil || hex.EncodeToString(disconnected) != "5647463101130000" {
|
|
t.Fatalf("disconnected event vector = %x, %v", disconnected, err)
|
|
}
|
|
if event, err := DecodeProviderEvent(disconnected); err != nil || event.Kind != ProviderEventDisconnected {
|
|
t.Fatalf("decoded disconnected event = %#v, %v", event, err)
|
|
}
|
|
receipt, err := EncodeClientFeedback(Feedback{Kind: FeedbackTerminalReceipt})
|
|
if err != nil || string(receipt) != protocolTerminalReceiptVector {
|
|
t.Fatalf("terminal receipt vector = %x, %v", receipt, err)
|
|
}
|
|
}
|
|
|
|
func TestCapabilityIntersectionAndBoundedQueue(t *testing.T) {
|
|
capabilities := DefaultCapabilities()
|
|
if _, err := IntersectCapabilities(capabilities, capabilities); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
other := capabilities
|
|
other.Audio = "different"
|
|
if !errors.Is(func() error { _, err := IntersectCapabilities(capabilities, other); return err }(), ErrNoCapabilityOverlap) {
|
|
t.Fatal("capability mismatch was accepted")
|
|
}
|
|
queue := NewBoundedQueue[int](2)
|
|
_ = queue.PushLatest(1)
|
|
_ = queue.PushLatest(2)
|
|
_ = queue.PushLatest(3)
|
|
if queue.Dropped() != 1 || queue.Len() != 2 {
|
|
t.Fatalf("queue length=%d dropped=%d", queue.Len(), queue.Dropped())
|
|
}
|
|
ctx := context.Background()
|
|
first, _ := queue.Pop(ctx)
|
|
second, _ := queue.Pop(ctx)
|
|
if first != 2 || second != 3 {
|
|
t.Fatalf("queue values=%d,%d", first, second)
|
|
}
|
|
}
|
|
|
|
func TestNewServerRejectsPartiallyConfiguredCapabilities(t *testing.T) {
|
|
serverTLS, _ := testTLS(t)
|
|
fake := NewFakeApollo(FakeApolloConfig{Now: time.Now()})
|
|
server, err := NewServer(ServerConfig{
|
|
TLSConfig: serverTLS, GatewayID: "gateway-1",
|
|
Capabilities: protocol.CapabilityProfile{SourceRateControl: "server"},
|
|
ProviderCapabilities: DefaultCapabilities(),
|
|
Admission: &oneTimeAdmission{},
|
|
Provider: fake,
|
|
})
|
|
if server != nil {
|
|
_ = server.Close()
|
|
}
|
|
if err == nil {
|
|
t.Fatal("partial capability configuration was silently replaced with defaults")
|
|
}
|
|
}
|
|
|
|
func TestSyntheticImpairmentPacingAndResourceBounds(t *testing.T) {
|
|
payload := make([]byte, maxCompleteFrameBytes+1)
|
|
if _, err := FragmentPayload(ChannelVideo, 1, 0, payload); !errors.Is(err, ErrFrameFragmentedLimit) {
|
|
t.Fatalf("oversized media payload accepted: %v", err)
|
|
}
|
|
frames, err := FragmentPayload(ChannelVideo, 1, 0, bytesRepeat(0x5a, frameV2PayloadSize*4))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var delivered int
|
|
var deliveredBytes int
|
|
for index, frame := range frames {
|
|
if (index+1)%3 == 0 { // deterministic synthetic loss profile: every third frame.
|
|
continue
|
|
}
|
|
delivered++
|
|
encoded, encodeErr := EncodeFrame(frame)
|
|
if encodeErr != nil {
|
|
t.Fatal(encodeErr)
|
|
}
|
|
decoded, decodeErr := DecodeFrame(encoded)
|
|
if decodeErr != nil {
|
|
t.Fatal(decodeErr)
|
|
}
|
|
deliveredBytes += len(decoded.Payload)
|
|
}
|
|
if delivered != 3 || deliveredBytes != frameV2PayloadSize*3 {
|
|
t.Fatalf("synthetic impairment delivered=%d bytes=%d", delivered, deliveredBytes)
|
|
}
|
|
}
|
|
|
|
func TestFragmentPayloadCarriesCompleteEncodedFrame(t *testing.T) {
|
|
payload := bytesRepeat(0x5a, 256*1024)
|
|
frames, err := FragmentPayload(ChannelVideo, 9, 11, payload)
|
|
if errors.Is(err, ErrFrameFragmentedLimit) {
|
|
t.Fatalf("complete encoded frame rejected at legacy fragment ceiling: %v", err)
|
|
}
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var recovered []byte
|
|
for _, frame := range frames {
|
|
recovered = append(recovered, frame.Payload...)
|
|
}
|
|
if !bytes.Equal(recovered, payload) {
|
|
t.Fatal("complete encoded frame payload changed")
|
|
}
|
|
}
|
|
|
|
func TestApolloFixturesAndLifecycle(t *testing.T) {
|
|
management, err := os.ReadFile("testdata/apollo-management.xml")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
info, err := ParseManagementXML(management)
|
|
if err != nil || info.Identity.UniqueID != "apollo-fixture-1" {
|
|
t.Fatalf("management = %#v, err = %v", info, err)
|
|
}
|
|
rtspText, err := os.ReadFile("testdata/rtsp-setup-response.txt")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
rtsp, err := ParseRTSPResponse([]byte(strings.ReplaceAll(string(rtspText), `\r\n`, "\r\n")))
|
|
if err != nil || rtsp.StatusCode != 200 {
|
|
t.Fatalf("RTSP = %#v, err = %v", rtsp, err)
|
|
}
|
|
video, _ := hex.DecodeString(strings.TrimSpace(string(mustRead(t, "testdata/encoded-video.hex"))))
|
|
audio, _ := hex.DecodeString(strings.TrimSpace(string(mustRead(t, "testdata/encoded-audio.hex"))))
|
|
now := time.Date(2026, time.January, 1, 0, 0, 0, 0, time.UTC)
|
|
identity := ProviderIdentity{UniqueID: "apollo-fixture-1", Fingerprint: "sha256:fixture-apollo-1"}
|
|
fake := NewFakeApollo(FakeApolloConfig{Identity: identity, Now: now, Video: [][]byte{video}, Audio: [][]byte{audio}})
|
|
session, err := fake.Start(context.Background(), LaunchRequest{SessionID: "session-1", ProviderProfile: ProviderProfileApollo, ProviderIdentity: identity.Key(), Capabilities: DefaultCapabilities()})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if got := <-session.Video(); string(got.Payload) != string(video) {
|
|
t.Fatalf("video changed: %x", got.Payload)
|
|
}
|
|
if got := <-session.Audio(); string(got.Payload) != string(audio) {
|
|
t.Fatalf("audio changed: %x", got.Payload)
|
|
}
|
|
if err := session.Input(context.Background(), InputEvent{Sequence: 1, Device: "keyboard", Code: 7, Pressed: true}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := session.ReleaseAll(context.Background()); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := session.Terminate(context.Background()); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if state := session.State(); state.State != ProviderStateTerminated || state.CleanupPending {
|
|
t.Fatalf("state = %#v", state)
|
|
}
|
|
|
|
identityFailure := NewFakeApollo(FakeApolloConfig{Identity: identity, Now: now, Failure: FakeFailureIdentity})
|
|
if _, err := identityFailure.Start(context.Background(), LaunchRequest{SessionID: "session-2", ProviderProfile: ProviderProfileApollo, ProviderIdentity: identity.Key(), Capabilities: DefaultCapabilities()}); !errors.Is(err, ErrProviderIdentity) {
|
|
t.Fatalf("identity failure = %v", err)
|
|
}
|
|
}
|
|
|
|
func TestProviderTimeoutAndBoundedInput(t *testing.T) {
|
|
now := time.Date(2026, time.January, 1, 0, 0, 0, 0, time.UTC)
|
|
fake := NewFakeApollo(FakeApolloConfig{Now: now, Failure: FakeFailureReadinessTimeout})
|
|
ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond)
|
|
defer cancel()
|
|
started := time.Now()
|
|
_, err := fake.Start(ctx, LaunchRequest{SessionID: "session-timeout", ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key(), Capabilities: DefaultCapabilities()})
|
|
if !errors.Is(err, ErrProviderTimeout) || time.Since(started) > time.Second {
|
|
t.Fatalf("readiness timeout = %v after %s", err, time.Since(started))
|
|
}
|
|
if _, err := EncodeInputEvent(InputEvent{Device: strings.Repeat("d", 65)}); !errors.Is(err, ErrInputMalformed) {
|
|
t.Fatalf("oversized input accepted: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestAdmissionQUICMTLSRelayAndCleanup(t *testing.T) {
|
|
serverTLS, clientTLS := testTLS(t)
|
|
fake := NewFakeApollo(FakeApolloConfig{Now: time.Date(2026, time.January, 1, 0, 0, 0, 0, time.UTC)})
|
|
authority := protocol.SessionAuthority{Version: "1", SessionID: "session-1", GatewayID: "gateway-1", Audience: "versevdi-gateway", ReconnectSequence: 0, ExpiresAt: time.Now().Add(5 * time.Second).UTC().Format(time.RFC3339Nano), Capabilities: DefaultCapabilities(), ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key()}
|
|
admission := &oneTimeAdmission{authority: authority, released: make(chan struct{})}
|
|
reporter := &recordingProviderStateReporter{}
|
|
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: "gateway-1", Capabilities: DefaultCapabilities(), ProviderCapabilities: DefaultCapabilities(), Admission: admission, ProviderStateReporter: reporter, ClipboardAuditReporter: reporter, Provider: fake})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
defer cancel()
|
|
serveDone := make(chan error, 1)
|
|
go func() { serveDone <- server.Serve(ctx) }()
|
|
request := protocol.TunnelAdmissionRequest{Version: "1", SessionID: "session-1", GatewayID: "gateway-1", Audience: "versevdi-gateway", Grant: strings.Repeat("g", 64), ReconnectSequence: 0, ClientNonce: "nonce-0000000001", DeviceSignature: strings.Repeat("s", 86), Capabilities: DefaultCapabilities()}
|
|
client, err := Dial(context.Background(), server.Addr().String(), clientTLS, request)
|
|
if err != nil {
|
|
_ = server.Close()
|
|
t.Fatal(err)
|
|
}
|
|
for i := 0; i < 2; i++ {
|
|
frame, receiveErr := client.ReceiveFrame(context.Background())
|
|
if receiveErr != nil {
|
|
t.Fatal(receiveErr)
|
|
}
|
|
if frame.Channel != ChannelVideo && frame.Channel != ChannelAudio {
|
|
t.Fatalf("unexpected media channel %d", frame.Channel)
|
|
}
|
|
}
|
|
metrics := server.Metrics()
|
|
if metrics.AdmittedSessions != 1 || metrics.MediaPackets < 2 || metrics.MediaBytes == 0 || metrics.ProcessingSamples < 2 || metrics.ProviderState != 2 {
|
|
t.Fatalf("observed egress telemetry = %#v", metrics)
|
|
}
|
|
fakeSession, ok := fake.LastSession().(*fakeSession)
|
|
if !ok {
|
|
t.Fatal("fake provider session type")
|
|
}
|
|
fakeSession.mu.Lock()
|
|
fakeSession.clipboard = "host clipboard"
|
|
fakeSession.mu.Unlock()
|
|
fakeSession.EmitEvent(ProviderEvent{Kind: ProviderEventRumble, Payload: []byte{1, 0x12, 0x34, 0x56, 0x78}})
|
|
clipboardCtx, clipboardCancel := context.WithTimeout(context.Background(), 2*time.Second)
|
|
deliveredClipboard, clipboardErr := client.ReceiveClipboard(clipboardCtx)
|
|
clipboardCancel()
|
|
if clipboardErr != nil || deliveredClipboard.Direction != "provider_to_client" || deliveredClipboard.Text != "host clipboard" || deliveredClipboard.Encoding != "utf-8" {
|
|
t.Fatalf("provider clipboard = %#v, %v", deliveredClipboard, clipboardErr)
|
|
}
|
|
if audits := reporter.Audits(); len(audits) == 0 || audits[0].Direction != "provider_to_client" || audits[0].Outcome != "forwarded" || audits[0].Reason != "forwarded" || audits[0].TextBytes != int64(len("host clipboard")) {
|
|
t.Fatalf("clipboard audits = %#v", audits)
|
|
}
|
|
eventCtx, eventCancel := context.WithTimeout(context.Background(), time.Second)
|
|
event, eventErr := client.ReceiveProviderEvent(eventCtx)
|
|
eventCancel()
|
|
if eventErr != nil || event.Kind != ProviderEventRumble || string(event.Payload) != string([]byte{1, 0x12, 0x34, 0x56, 0x78}) {
|
|
t.Fatalf("provider event = %#v, %v", event, eventErr)
|
|
}
|
|
if err := client.SendClipboard(protocol.GatewayClipboardText{Direction: "client_to_provider", Text: "clipboard", Encoding: "utf-8", LoopToken: "abcdefghijklmnop"}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
select {
|
|
case value := <-fakeSession.clipboardWrites:
|
|
if value != "clipboard" {
|
|
t.Fatalf("provider clipboard = %q", value)
|
|
}
|
|
case <-time.After(time.Second):
|
|
t.Fatal("gateway did not forward clipboard")
|
|
}
|
|
if err := client.SendInput(InputEvent{Sequence: 1, Device: "keyboard", Code: 7, Pressed: true}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
_ = client.Close()
|
|
deadline := time.NewTimer(2 * time.Second)
|
|
select {
|
|
case <-admission.released:
|
|
case <-deadline.C:
|
|
t.Fatal("gateway did not release admission")
|
|
}
|
|
deadline.Stop()
|
|
if fake.LastSession().State().State != ProviderStateTerminated {
|
|
t.Fatalf("provider state = %#v", fake.LastSession().State())
|
|
}
|
|
if _, err := Dial(context.Background(), server.Addr().String(), clientTLS, request); err == nil {
|
|
t.Fatal("replayed grant was accepted")
|
|
}
|
|
_ = server.Close()
|
|
if err := <-serveDone; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
states := reporter.States()
|
|
if len(states) != 3 || states[0].State != ProviderStateStarting || states[1].State != ProviderStateReady || states[2].State != ProviderStateTerminated {
|
|
t.Fatalf("provider states = %#v", states)
|
|
}
|
|
}
|
|
|
|
func TestStableErrorFramePrecedesConnectionTeardown(t *testing.T) {
|
|
for _, retryable := range []bool{false, true} {
|
|
t.Run(fmt.Sprintf("retryable=%t", retryable), func(t *testing.T) {
|
|
serverTLS, clientTLS := testTLS(t)
|
|
clientTLS.NextProtos = []string{"versevdi-gateway-v1"}
|
|
fake := NewFakeApollo(FakeApolloConfig{Now: time.Now()})
|
|
authority := protocol.SessionAuthority{
|
|
Version: "1", SessionID: "session-1", GatewayID: "gateway-1", Audience: "versevdi-gateway",
|
|
ExpiresAt: time.Now().Add(time.Minute).UTC().Format(time.RFC3339Nano), Capabilities: DefaultCapabilities(),
|
|
ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key(),
|
|
}
|
|
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID, Admission: &oneTimeAdmission{authority: authority, released: make(chan struct{})}, Provider: fake})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if retryable {
|
|
server.BeginDrain()
|
|
}
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
defer cancel()
|
|
go func() { _ = server.Serve(ctx) }()
|
|
connection, err := quic.DialAddr(context.Background(), server.Addr().String(), clientTLS, &quic.Config{EnableDatagrams: true})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
stream, err := connection.OpenStreamSync(context.Background())
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
gatewayID := authority.GatewayID
|
|
if !retryable {
|
|
gatewayID = "wrong-gateway"
|
|
}
|
|
request := protocol.TunnelAdmissionRequest{Version: "1", SessionID: authority.SessionID, GatewayID: gatewayID, Audience: authority.Audience, Grant: strings.Repeat("g", 64), ClientNonce: "nonce-0000000001", DeviceSignature: strings.Repeat("s", 86), Capabilities: DefaultCapabilities()}
|
|
payload, err := protocol.EncodeTunnelAdmissionRequest(request)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := writeWire(stream, payload, defaultHelloLimit); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
response, err := readWire(stream, defaultHelloLimit)
|
|
if err != nil {
|
|
t.Fatalf("stable response lost before connection teardown: %v", err)
|
|
}
|
|
stable, err := protocol.DecodeStableError(response)
|
|
if err != nil || stable.Retryable != retryable {
|
|
t.Fatalf("stable response = %#v, %v", stable, err)
|
|
}
|
|
select {
|
|
case <-connection.Context().Done():
|
|
case <-time.After(2 * time.Second):
|
|
t.Fatal("server retained rejected connection")
|
|
}
|
|
_ = server.Close()
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestStableErrorDrainOutlivesExpiredHandlerContext(t *testing.T) {
|
|
serverTLS, clientTLS := testTLS(t)
|
|
clientTLS.NextProtos = []string{"versevdi-gateway-v1"}
|
|
fake := NewFakeApollo(FakeApolloConfig{Now: time.Now()})
|
|
authority := protocol.SessionAuthority{Version: "1", SessionID: "session-1", GatewayID: "gateway-1", Audience: "versevdi-gateway", ExpiresAt: time.Now().Add(time.Minute).UTC().Format(time.RFC3339Nano), Capabilities: DefaultCapabilities(), ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key()}
|
|
admission := AdmissionFunc(func(ctx context.Context, _ protocol.TunnelAdmissionRequest) (protocol.SessionAuthority, error) {
|
|
<-ctx.Done()
|
|
return protocol.SessionAuthority{}, ctx.Err()
|
|
})
|
|
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID, Admission: admission, Provider: fake})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
server.helloTimeout = 20 * time.Millisecond
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
go func() { _ = server.Serve(ctx) }()
|
|
connection, err := quic.DialAddr(context.Background(), server.Addr().String(), clientTLS, &quic.Config{EnableDatagrams: true})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
stream, err := connection.OpenStreamSync(context.Background())
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
request := protocol.TunnelAdmissionRequest{Version: "1", SessionID: authority.SessionID, GatewayID: authority.GatewayID, Audience: authority.Audience, Grant: strings.Repeat("g", 64), ClientNonce: "nonce-0000000001", DeviceSignature: strings.Repeat("s", 86), Capabilities: DefaultCapabilities()}
|
|
payload, err := protocol.EncodeTunnelAdmissionRequest(request)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := writeWire(stream, payload, defaultHelloLimit); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
response, err := readWire(stream, defaultHelloLimit)
|
|
if err != nil {
|
|
t.Fatalf("stable response lost with expired handler context: %v", err)
|
|
}
|
|
if stable, err := protocol.DecodeStableError(response); err != nil || stable.Code != "admission_rejected" {
|
|
t.Fatalf("stable response = %#v, %v", stable, err)
|
|
}
|
|
_ = connection.CloseWithError(applicationError, "done")
|
|
cancel()
|
|
_ = server.Close()
|
|
}
|
|
|
|
func TestServerCloseInterruptsPartialHelloClient(t *testing.T) {
|
|
serverTLS, clientTLS := testTLS(t)
|
|
clientTLS.NextProtos = []string{"versevdi-gateway-v1"}
|
|
fake := NewFakeApollo(FakeApolloConfig{Now: time.Now()})
|
|
authority := protocol.SessionAuthority{Version: "1", SessionID: "session-1", GatewayID: "gateway-1", Audience: "versevdi-gateway", ExpiresAt: time.Now().Add(time.Minute).UTC().Format(time.RFC3339Nano), Capabilities: DefaultCapabilities(), ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key()}
|
|
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID, Admission: &oneTimeAdmission{authority: authority, released: make(chan struct{})}, Provider: fake})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
serveDone := make(chan error, 1)
|
|
go func() { serveDone <- server.Serve(ctx) }()
|
|
connection, err := quic.DialAddr(context.Background(), server.Addr().String(), clientTLS, &quic.Config{EnableDatagrams: true})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
stream, err := connection.OpenStreamSync(context.Background())
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, err := stream.Write([]byte{0, 0}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
closed := make(chan error, 1)
|
|
go func() { closed <- server.Close() }()
|
|
select {
|
|
case err := <-closed:
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
case <-time.After(time.Second):
|
|
t.Fatal("Server.Close blocked on partial hello")
|
|
}
|
|
select {
|
|
case <-connection.Context().Done():
|
|
case <-time.After(time.Second):
|
|
t.Fatal("partial hello connection remained open")
|
|
}
|
|
cancel()
|
|
if err := <-serveDone; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
|
|
func TestHelloUsesOneAbsoluteDeadlineAgainstTrickle(t *testing.T) {
|
|
serverTLS, clientTLS := testTLS(t)
|
|
clientTLS.NextProtos = []string{"versevdi-gateway-v1"}
|
|
fake := NewFakeApollo(FakeApolloConfig{Now: time.Now()})
|
|
authority := protocol.SessionAuthority{Version: "1", SessionID: "session-1", GatewayID: "gateway-1", Audience: "versevdi-gateway", ExpiresAt: time.Now().Add(time.Minute).UTC().Format(time.RFC3339Nano), Capabilities: DefaultCapabilities(), ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key()}
|
|
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID, Admission: &oneTimeAdmission{authority: authority, released: make(chan struct{})}, Provider: fake})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
server.helloTimeout = 40 * time.Millisecond
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
defer cancel()
|
|
go func() { _ = server.Serve(ctx) }()
|
|
connection, err := quic.DialAddr(context.Background(), server.Addr().String(), clientTLS, &quic.Config{EnableDatagrams: true})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
stream, err := connection.OpenStreamSync(context.Background())
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, err := stream.Write([]byte{0}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
started := time.Now()
|
|
time.Sleep(25 * time.Millisecond)
|
|
_, _ = stream.Write([]byte{0})
|
|
response, err := readWire(stream, defaultHelloLimit)
|
|
if err != nil {
|
|
t.Fatalf("absolute hello deadline did not produce a stable error: %v", err)
|
|
}
|
|
stable, err := protocol.DecodeStableError(response)
|
|
if err != nil || stable.Code != "invalid_hello" {
|
|
t.Fatalf("stable response = %#v, %v", stable, err)
|
|
}
|
|
if time.Since(started) > 250*time.Millisecond {
|
|
t.Fatal("trickling hello reset or escaped its absolute deadline")
|
|
}
|
|
_ = connection.CloseWithError(applicationError, "done")
|
|
_ = server.Close()
|
|
}
|
|
|
|
func TestHelloDeadlineIsClearedAfterAdmission(t *testing.T) {
|
|
serverTLS, clientTLS := testTLS(t)
|
|
fake := NewFakeApollo(FakeApolloConfig{Now: time.Now()})
|
|
authority := protocol.SessionAuthority{Version: "1", SessionID: "session-1", GatewayID: "gateway-1", Audience: "versevdi-gateway", ExpiresAt: time.Now().Add(time.Minute).UTC().Format(time.RFC3339Nano), Capabilities: DefaultCapabilities(), ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key()}
|
|
admission := &oneTimeAdmission{authority: authority, released: make(chan struct{}), disableClipboard: true}
|
|
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID, Admission: admission, Provider: fake})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
server.helloTimeout = 40 * time.Millisecond
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
defer cancel()
|
|
go func() { _ = server.Serve(ctx) }()
|
|
request := protocol.TunnelAdmissionRequest{Version: "1", SessionID: authority.SessionID, GatewayID: authority.GatewayID, Audience: authority.Audience, Grant: strings.Repeat("g", 64), ClientNonce: "nonce-0000000001", DeviceSignature: strings.Repeat("s", 86), Capabilities: DefaultCapabilities()}
|
|
client, err := Dial(context.Background(), server.Addr().String(), clientTLS, request)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
time.Sleep(80 * time.Millisecond)
|
|
select {
|
|
case <-client.connection.Context().Done():
|
|
t.Fatal("hello deadline leaked into the admitted session")
|
|
default:
|
|
}
|
|
_ = client.Close()
|
|
_ = server.Close()
|
|
}
|
|
|
|
func TestStableErrorNeverLeaksInternalProviderDetails(t *testing.T) {
|
|
var wire bytes.Buffer
|
|
if err := writeStableError(&wire, "provider_unavailable", errors.New("https://provider.invalid/launch?rikey=secret-sentinel"), true); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
response, err := readWire(&wire, defaultHelloLimit)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
stable, err := protocol.DecodeStableError(response)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if strings.Contains(stable.Message, "provider.invalid") || strings.Contains(stable.Message, "secret-sentinel") {
|
|
t.Fatalf("stable error leaked internal details: %q", stable.Message)
|
|
}
|
|
}
|
|
|
|
func TestPostAdmissionProviderFailureIsTerminalEvenWhenReleaseFails(t *testing.T) {
|
|
serverTLS, clientTLS := testTLS(t)
|
|
clientTLS.NextProtos = []string{"versevdi-gateway-v1"}
|
|
fake := NewFakeApollo(FakeApolloConfig{Now: time.Now()})
|
|
authority := protocol.SessionAuthority{Version: "1", SessionID: "session-1", GatewayID: "gateway-1", Audience: "versevdi-gateway", ExpiresAt: time.Now().Add(time.Minute).UTC().Format(time.RFC3339Nano), Capabilities: DefaultCapabilities(), ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key()}
|
|
admission := &oneTimeAdmission{authority: authority, released: make(chan struct{}), disableClipboard: true, releaseErr: errors.New("release failed")}
|
|
provider := providerStartFunc(func(context.Context, LaunchRequest) (ProviderSession, error) {
|
|
return nil, context.DeadlineExceeded
|
|
})
|
|
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID, Admission: admission, Provider: provider})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
defer cancel()
|
|
go func() { _ = server.Serve(ctx) }()
|
|
connection, err := quic.DialAddr(context.Background(), server.Addr().String(), clientTLS, &quic.Config{EnableDatagrams: true})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
stream, err := connection.OpenStreamSync(context.Background())
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
request := protocol.TunnelAdmissionRequest{Version: "1", SessionID: authority.SessionID, GatewayID: authority.GatewayID, Audience: authority.Audience, Grant: strings.Repeat("g", 64), ClientNonce: "nonce-0000000001", DeviceSignature: strings.Repeat("s", 86), Capabilities: DefaultCapabilities()}
|
|
payload, err := protocol.EncodeTunnelAdmissionRequest(request)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := writeWire(stream, payload, defaultHelloLimit); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
response, err := readWire(stream, defaultHelloLimit)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
stable, err := protocol.DecodeStableError(response)
|
|
if err != nil || stable.Code != "provider_timeout" || stable.Retryable {
|
|
t.Fatalf("post-admission stable response = %#v, %v", stable, err)
|
|
}
|
|
select {
|
|
case <-admission.released:
|
|
case <-time.After(time.Second):
|
|
t.Fatal("post-admission failure did not attempt release")
|
|
}
|
|
if got := admission.releases.Load(); got != 1 {
|
|
t.Fatalf("release attempts = %d, want 1", got)
|
|
}
|
|
_ = connection.CloseWithError(applicationError, "done")
|
|
_ = server.Close()
|
|
}
|
|
|
|
func TestGatewayTelemetrySeparatesQueueProcessingAndPacing(t *testing.T) {
|
|
serverTLS, clientTLS := testTLS(t)
|
|
session := &fakeSession{
|
|
video: make(chan ProviderMedia, 1),
|
|
audio: make(chan ProviderMedia),
|
|
events: make(chan ProviderEvent, 1),
|
|
clipboardWrites: make(chan string, 1),
|
|
state: protocol.ProviderState{
|
|
Version: "1", SessionID: "session-timing", State: ProviderStateStarting,
|
|
Channels: []string{"video", "audio", "input", "feedback"},
|
|
},
|
|
pressed: make(map[string]struct{}),
|
|
}
|
|
provider := providerStartFunc(func(context.Context, LaunchRequest) (ProviderSession, error) {
|
|
enqueuedAt := time.Now()
|
|
session.video <- ProviderMedia{Payload: bytesRepeat(0x5a, 2000), ReceivedAt: enqueuedAt, EnqueuedAt: enqueuedAt}
|
|
time.Sleep(60 * time.Millisecond)
|
|
session.mu.Lock()
|
|
session.state.State = ProviderStateReady
|
|
session.mu.Unlock()
|
|
return session, nil
|
|
})
|
|
authority := protocol.SessionAuthority{
|
|
Version: "1", SessionID: "session-timing", GatewayID: "gateway-1", Audience: "versevdi-gateway",
|
|
ExpiresAt: time.Now().Add(5 * time.Second).UTC().Format(time.RFC3339Nano),
|
|
Capabilities: DefaultCapabilities(), ProviderProfile: ProviderProfileApollo,
|
|
ProviderIdentity: "apollo-fixture-1#sha256:fixture-apollo-1",
|
|
}
|
|
admission := &oneTimeAdmission{authority: authority, released: make(chan struct{}), disableClipboard: true}
|
|
server, err := NewServer(ServerConfig{
|
|
ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID,
|
|
Capabilities: DefaultCapabilities(), ProviderCapabilities: DefaultCapabilities(),
|
|
Admission: admission, Provider: provider, PacerKbps: 24,
|
|
})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
serveDone := make(chan error, 1)
|
|
go func() { serveDone <- server.Serve(ctx) }()
|
|
request := protocol.TunnelAdmissionRequest{
|
|
Version: "1", SessionID: authority.SessionID, GatewayID: authority.GatewayID, Audience: authority.Audience,
|
|
Grant: strings.Repeat("g", 64), ClientNonce: "nonce-0000000001", DeviceSignature: strings.Repeat("s", 86),
|
|
Capabilities: DefaultCapabilities(),
|
|
}
|
|
client, err := Dial(context.Background(), server.Addr().String(), clientTLS, request)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
receiveCtx, receiveCancel := context.WithTimeout(context.Background(), 2*time.Second)
|
|
for index := uint16(0); index < 2; index++ {
|
|
frame, err := client.ReceiveFrame(receiveCtx)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if frame.FragmentIndex != index || frame.FragmentCount != 2 {
|
|
t.Fatalf("timing frame = %#v", frame)
|
|
}
|
|
}
|
|
receiveCancel()
|
|
metrics := server.Metrics()
|
|
if metrics.ProcessingSamples != 1 {
|
|
t.Fatalf("timing samples = %d, want one provider unit", metrics.ProcessingSamples)
|
|
}
|
|
if metrics.MediaPackets != 2 {
|
|
t.Fatalf("media packets = %d, want two fragments", metrics.MediaPackets)
|
|
}
|
|
queue, processing, pacing := time.Duration(metrics.QueueDelayNanos), time.Duration(metrics.ProcessingDelayNanos), time.Duration(metrics.PacingDelayNanos)
|
|
if queue < 40*time.Millisecond || queue > 150*time.Millisecond {
|
|
t.Fatalf("queue residence = %s, want the controlled 60ms provider queue wait", queue)
|
|
}
|
|
if processing >= 100*time.Millisecond {
|
|
t.Fatalf("processing = %s, pacing leaked into gateway processing", processing)
|
|
}
|
|
if pacing < 500*time.Millisecond {
|
|
t.Fatalf("pacing = %s, want the controlled scheduler wait", pacing)
|
|
}
|
|
_ = client.Close()
|
|
cancel()
|
|
_ = server.Close()
|
|
if err := <-serveDone; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
|
|
func TestGatewayRejectsProviderWorkOutsideNegotiatedDecodeProfile(t *testing.T) {
|
|
serverTLS, clientTLS := testTLS(t)
|
|
fake := NewFakeApollo(FakeApolloConfig{Now: time.Date(2026, time.January, 1, 0, 0, 0, 0, time.UTC)})
|
|
capabilities := DefaultCapabilities()
|
|
capabilities.ClientDecode = []string{"h264-opus"}
|
|
authority := protocol.SessionAuthority{
|
|
Version: "1", SessionID: "session-policy", GatewayID: "gateway-1", Audience: "versevdi-gateway",
|
|
ExpiresAt: time.Now().Add(5 * time.Second).UTC().Format(time.RFC3339Nano),
|
|
Capabilities: capabilities, ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key(),
|
|
}
|
|
admission := &oneTimeAdmission{
|
|
authority: authority, released: make(chan struct{}),
|
|
streamPolicy: protocol.ProviderStreamPolicy{
|
|
ResolutionWidth: 2560, ResolutionHeight: 1440, Fps: 120,
|
|
Codec: "HEVC", BitrateKbps: 40000, AudioEnabled: true,
|
|
},
|
|
}
|
|
server, err := NewServer(ServerConfig{
|
|
ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: "gateway-1",
|
|
Capabilities: capabilities, ProviderCapabilities: capabilities, Admission: admission, Provider: fake,
|
|
})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
defer cancel()
|
|
serveDone := make(chan error, 1)
|
|
go func() { serveDone <- server.Serve(ctx) }()
|
|
request := protocol.TunnelAdmissionRequest{
|
|
Version: "1", SessionID: authority.SessionID, GatewayID: authority.GatewayID, Audience: authority.Audience,
|
|
Grant: strings.Repeat("g", 64), ClientNonce: "nonce-0000000001", DeviceSignature: strings.Repeat("s", 86),
|
|
Capabilities: capabilities,
|
|
}
|
|
if _, err := Dial(context.Background(), server.Addr().String(), clientTLS, request); err == nil {
|
|
t.Fatal("gateway accepted HEVC provider work for an H.264-only negotiated profile")
|
|
}
|
|
select {
|
|
case <-admission.released:
|
|
case <-time.After(time.Second):
|
|
t.Fatal("gateway did not release rejected provider work")
|
|
}
|
|
if session, _ := fake.LastSession().(*fakeSession); session != nil {
|
|
t.Fatal("provider started before policy/capability rejection")
|
|
}
|
|
_ = server.Close()
|
|
if err := <-serveDone; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
|
|
func TestGatewayNegotiatesRegisteredProfilesWithIndependentClient(t *testing.T) {
|
|
for _, test := range []struct {
|
|
name string
|
|
clientProfiles []string
|
|
selected string
|
|
codec string
|
|
}{
|
|
{name: "h264-only", clientProfiles: []string{"h264-opus"}, selected: "h264-opus", codec: "H264"},
|
|
{name: "hevc-only", clientProfiles: []string{"hevc-opus"}, selected: "hevc-opus", codec: "HEVC"},
|
|
{name: "policy-selects-hevc", clientProfiles: []string{"h264-opus", "hevc-opus"}, selected: "hevc-opus", codec: "HEVC"},
|
|
} {
|
|
t.Run(test.name, func(t *testing.T) {
|
|
serverTLS, clientTLS := testTLS(t)
|
|
fake := NewFakeApollo(FakeApolloConfig{Now: time.Date(2026, time.January, 1, 0, 0, 0, 0, time.UTC)})
|
|
clientCapabilities := DefaultCapabilities()
|
|
clientCapabilities.ClientDecode = test.clientProfiles
|
|
authority := protocol.SessionAuthority{
|
|
Version: "1", SessionID: "session-profile-" + test.name, GatewayID: "gateway-1", Audience: "versevdi-gateway",
|
|
ExpiresAt: time.Now().Add(5 * time.Second).UTC().Format(time.RFC3339Nano),
|
|
Capabilities: clientCapabilities, ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key(),
|
|
}
|
|
admission := &oneTimeAdmission{
|
|
authority: authority, released: make(chan struct{}), disableClipboard: true,
|
|
streamPolicy: protocol.ProviderStreamPolicy{
|
|
ResolutionWidth: 1920, ResolutionHeight: 1080, Fps: 60,
|
|
Codec: test.codec, BitrateKbps: 8000, AudioEnabled: true,
|
|
},
|
|
}
|
|
started := make(chan LaunchRequest, 1)
|
|
provider := providerStartFunc(func(ctx context.Context, request LaunchRequest) (ProviderSession, error) {
|
|
started <- request
|
|
return fake.Start(ctx, request)
|
|
})
|
|
server, err := NewServer(ServerConfig{
|
|
ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID,
|
|
Capabilities: DefaultCapabilities(), ProviderCapabilities: DefaultCapabilities(),
|
|
Admission: admission, Provider: provider,
|
|
})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
serveDone := make(chan error, 1)
|
|
go func() { serveDone <- server.Serve(ctx) }()
|
|
request := protocol.TunnelAdmissionRequest{
|
|
Version: "1", SessionID: authority.SessionID, GatewayID: authority.GatewayID, Audience: authority.Audience,
|
|
Grant: strings.Repeat("g", 64), ClientNonce: "nonce-0000000001", DeviceSignature: strings.Repeat("s", 86),
|
|
Capabilities: clientCapabilities,
|
|
}
|
|
client, err := Dial(context.Background(), server.Addr().String(), clientTLS, request)
|
|
if err != nil {
|
|
cancel()
|
|
_ = server.Close()
|
|
t.Fatal(err)
|
|
}
|
|
launch := <-started
|
|
if !reflect.DeepEqual(launch.Capabilities.ClientDecode, []string{test.selected}) {
|
|
t.Fatalf("provider selected profiles = %v", launch.Capabilities.ClientDecode)
|
|
}
|
|
_ = client.Close()
|
|
cancel()
|
|
_ = server.Close()
|
|
if err := <-serveDone; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestRegisteredChannelFramesTraversePublicTransport(t *testing.T) {
|
|
h := newGatewayTransportHarness(t)
|
|
|
|
input, err := EncodeInputEvent(InputEvent{Sequence: 7, Device: "keyboard", Code: 7, Pressed: true})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
absolute, err := EncodeInputEvent(InputEvent{Sequence: 10, Device: "mouse-absolute", Payload: []byte{0, 1, 0, 1, 0, 2, 0, 2}})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
scroll, err := EncodeInputEvent(InputEvent{Sequence: 11, Device: "mouse-scroll", Payload: []byte{0, 1, 0, 1}})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
feedback, err := EncodeClientFeedback(Feedback{Sequence: 8, Kind: FeedbackIDR})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
clipboard, err := protocol.EncodeGatewayClipboardText(protocol.GatewayClipboardText{Direction: "client_to_provider", Text: "registered clipboard", Encoding: "utf-8", LoopToken: "abcdefghijklmnop"})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
for _, frame := range []protocol.ChannelFrame{
|
|
testChannelFrame("input.sequenced.v1", 7, input),
|
|
testChannelFrame("control.ack.v1", 8, feedback),
|
|
testChannelFrame("clipboard.text.v1", 9, clipboard),
|
|
testChannelFrame("input.sequenced.v1", 10, absolute),
|
|
testChannelFrame("input.sequenced.v1", 11, scroll),
|
|
} {
|
|
encoded, encodeErr := protocol.EncodeChannelFrame(frame)
|
|
if encodeErr != nil {
|
|
t.Fatal(encodeErr)
|
|
}
|
|
if writeErr := h.client.writeControl(encoded); writeErr != nil {
|
|
t.Fatal(writeErr)
|
|
}
|
|
}
|
|
|
|
select {
|
|
case value := <-h.session.clipboardWrites:
|
|
if value != "registered clipboard" {
|
|
t.Fatalf("provider clipboard = %q", value)
|
|
}
|
|
case <-time.After(time.Second):
|
|
t.Fatal("registered clipboard flow was not forwarded")
|
|
}
|
|
deadline := time.Now().Add(time.Second)
|
|
for {
|
|
h.session.mu.Lock()
|
|
inputs := append([]InputEvent(nil), h.session.inputs...)
|
|
feedbacks := append([]Feedback(nil), h.session.feedback...)
|
|
h.session.mu.Unlock()
|
|
if len(inputs) == 3 && inputs[0].Sequence == 7 && inputs[1].Sequence == 10 && inputs[1].Device == "mouse-absolute" && inputs[2].Sequence == 11 && inputs[2].Device == "mouse-scroll" && len(feedbacks) == 1 && feedbacks[0].Sequence == 8 && feedbacks[0].Kind == FeedbackIDR {
|
|
break
|
|
}
|
|
if time.Now().After(deadline) {
|
|
t.Fatalf("registered flows inputs=%#v feedback=%#v", inputs, feedbacks)
|
|
}
|
|
time.Sleep(time.Millisecond)
|
|
}
|
|
|
|
h.session.EmitEvent(ProviderEvent{Kind: ProviderEventRumble, Payload: []byte{1, 2, 3, 4, 5}})
|
|
eventCtx, cancelEvent := context.WithTimeout(context.Background(), time.Second)
|
|
event, err := h.client.ReceiveProviderEvent(eventCtx)
|
|
cancelEvent()
|
|
if err != nil || event.Kind != ProviderEventRumble || string(event.Payload) != string([]byte{1, 2, 3, 4, 5}) {
|
|
t.Fatalf("registered provider feedback = %#v, %v", event, err)
|
|
}
|
|
|
|
h.session.mu.Lock()
|
|
h.session.clipboard = "registered provider clipboard"
|
|
h.session.mu.Unlock()
|
|
clipboardCtx, cancelClipboard := context.WithTimeout(context.Background(), time.Second)
|
|
value, err := h.client.ReceiveClipboard(clipboardCtx)
|
|
cancelClipboard()
|
|
if err != nil || value.Text != "registered provider clipboard" {
|
|
t.Fatalf("registered provider clipboard = %#v, %v", value, err)
|
|
}
|
|
}
|
|
|
|
func TestPrivateChannelAliasesAreRejectedOnPublicTransport(t *testing.T) {
|
|
for _, test := range []struct {
|
|
alias string
|
|
payload func(t *testing.T) []byte
|
|
}{
|
|
{"control", func(t *testing.T) []byte {
|
|
t.Helper()
|
|
value, err := EncodeClientFeedback(Feedback{Kind: FeedbackIDR})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return value
|
|
}},
|
|
{"input", func(t *testing.T) []byte {
|
|
t.Helper()
|
|
value, err := EncodeInputEvent(InputEvent{Device: "keyboard", Code: 7, Pressed: true})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return value
|
|
}},
|
|
{"clipboard", func(t *testing.T) []byte {
|
|
t.Helper()
|
|
value, err := protocol.EncodeGatewayClipboardText(protocol.GatewayClipboardText{Direction: "client_to_provider", Text: "alias", Encoding: "utf-8", LoopToken: "abcdefghijklmnop"})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return value
|
|
}},
|
|
} {
|
|
t.Run(test.alias, func(t *testing.T) {
|
|
h := newGatewayTransportHarness(t)
|
|
encoded, err := protocol.EncodeChannelFrame(testChannelFrame(test.alias, 1, test.payload(t)))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := h.client.writeControl(encoded); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
h.waitReleased(t)
|
|
h.session.mu.Lock()
|
|
defer h.session.mu.Unlock()
|
|
if len(h.session.inputs) != 0 || len(h.session.feedback) != 0 {
|
|
t.Fatalf("alias reached provider: inputs=%#v feedback=%#v", h.session.inputs, h.session.feedback)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestProviderClipboardAuditWaitsForPublicTransportDelivery(t *testing.T) {
|
|
h := newGatewayTransportHarness(t)
|
|
h.client.control.CancelRead(0)
|
|
h.session.mu.Lock()
|
|
h.session.clipboard = "undeliverable clipboard"
|
|
h.session.mu.Unlock()
|
|
|
|
h.waitReleased(t)
|
|
for _, audit := range h.reporter.Audits() {
|
|
if audit.Direction == "provider_to_client" && audit.Outcome == "forwarded" {
|
|
t.Fatalf("failed delivery was audited as forwarded: %#v", audit)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestProviderTerminationEndsPublicGatewaySession(t *testing.T) {
|
|
h := newGatewayTransportHarnessWithoutClipboard(t)
|
|
h.drainInitialMedia(t)
|
|
h.session.EmitEvent(ProviderEvent{Kind: ProviderEventTerminated, Payload: []byte{1, 2, 3, 4}})
|
|
|
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
|
event, err := h.client.ReceiveProviderEvent(ctx)
|
|
cancel()
|
|
if err != nil || event.Kind != ProviderEventTerminated {
|
|
t.Fatalf("provider termination event = %#v, %v", event, err)
|
|
}
|
|
h.waitReleased(t)
|
|
if states := h.reporter.States(); len(states) == 0 || states[len(states)-1].State != ProviderStateTerminated {
|
|
t.Fatalf("provider states = %#v", states)
|
|
}
|
|
h.session.mu.Lock()
|
|
if len(h.session.feedback) != 0 {
|
|
t.Fatalf("terminal receipt reached provider: %#v", h.session.feedback)
|
|
}
|
|
h.session.mu.Unlock()
|
|
h.assertMediaClosed(t)
|
|
}
|
|
|
|
func TestEncryptedNativeHostTerminationEndsPublicGatewaySession(t *testing.T) {
|
|
h := newNativeGatewayLifecycleHarness(t, "session-native-terminal")
|
|
h.native.handleApolloControlPayload(apolloChannelGeneric, true, sourceSealHostControl(t, h.key, 0, apolloControlTypeTerm, []byte{1, 2, 3, 4}))
|
|
tryQueueNativeMedia(h.native, h.native.video, []byte("queued-video"))
|
|
tryQueueNativeMedia(h.native, h.native.audio, []byte("queued-audio"))
|
|
|
|
eventCtx, eventCancel := context.WithTimeout(context.Background(), time.Second)
|
|
event, err := h.client.ReceiveProviderEvent(eventCtx)
|
|
eventCancel()
|
|
if err != nil || event.Kind != ProviderEventTerminated {
|
|
t.Fatalf("native provider termination = %#v, %v", event, err)
|
|
}
|
|
h.client.waitClosed(t)
|
|
tryQueueNativeMedia(h.native, h.native.video, []byte("new-video"))
|
|
tryQueueNativeMedia(h.native, h.native.audio, []byte("new-audio"))
|
|
h.assertNoMedia(t)
|
|
h.waitReleased(t)
|
|
if states := h.reporter.States(); len(states) == 0 || states[len(states)-1].State != ProviderStateTerminated {
|
|
t.Fatalf("provider states = %#v", states)
|
|
}
|
|
}
|
|
|
|
func TestNativeENetDisconnectQuiescesPublicGatewaySession(t *testing.T) {
|
|
h := newNativeGatewayLifecycleHarness(t, "session-native-disconnect")
|
|
h.native.handleApolloDisconnect(ErrProviderDisconnected)
|
|
tryQueueNativeMedia(h.native, h.native.video, []byte("queued-video"))
|
|
tryQueueNativeMedia(h.native, h.native.audio, []byte("queued-audio"))
|
|
|
|
eventCtx, eventCancel := context.WithTimeout(context.Background(), time.Second)
|
|
event, err := h.client.ReceiveProviderEvent(eventCtx)
|
|
eventCancel()
|
|
if err != nil || event.Kind != ProviderEventDisconnected {
|
|
t.Fatalf("native provider disconnect = %#v, %v", event, err)
|
|
}
|
|
h.client.waitClosed(t)
|
|
tryQueueNativeMedia(h.native, h.native.video, []byte("new-video"))
|
|
tryQueueNativeMedia(h.native, h.native.audio, []byte("new-audio"))
|
|
h.assertNoMedia(t)
|
|
h.waitReleased(t)
|
|
if states := h.reporter.States(); len(states) == 0 || states[len(states)-1].State != ProviderStateDisconnected || states[len(states)-1].CleanupPending {
|
|
t.Fatalf("provider states = %#v", states)
|
|
}
|
|
}
|
|
|
|
func TestTerminalTunnelClosesWhenIndependentClientDoesNotAcknowledge(t *testing.T) {
|
|
h := newNativeGatewayLifecycleHarness(t, "session-native-no-receipt")
|
|
h.native.handleApolloDisconnect(ErrProviderDisconnected)
|
|
|
|
eventCtx, eventCancel := context.WithTimeout(context.Background(), time.Second)
|
|
event, err := h.client.receiveProviderEvent(eventCtx, false)
|
|
eventCancel()
|
|
if err != nil || event.Kind != ProviderEventDisconnected {
|
|
t.Fatalf("native provider disconnect = %#v, %v", event, err)
|
|
}
|
|
h.client.waitClosed(t)
|
|
h.waitReleased(t)
|
|
}
|
|
|
|
type nativeGatewayLifecycleHarness struct {
|
|
native *nativeApolloSession
|
|
key []byte
|
|
client *independentGatewayClient
|
|
admission *oneTimeAdmission
|
|
reporter *recordingProviderStateReporter
|
|
}
|
|
|
|
func newNativeGatewayLifecycleHarness(t *testing.T, sessionID string) nativeGatewayLifecycleHarness {
|
|
t.Helper()
|
|
serverTLS, clientTLS := testTLS(t)
|
|
key := []byte("0123456789abcdef")
|
|
native := newNativeApolloSession(sessionID)
|
|
control, err := newApolloControlCodec(key)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
native.control = control
|
|
close(native.readDone)
|
|
session := nativeLifecycleSession{native}
|
|
provider := providerStartFunc(func(context.Context, LaunchRequest) (ProviderSession, error) {
|
|
if err := native.Ready(context.Background()); err != nil {
|
|
return nil, err
|
|
}
|
|
return session, nil
|
|
})
|
|
authority := protocol.SessionAuthority{
|
|
Version: "1", SessionID: native.sessionID, GatewayID: "gateway-1", Audience: "versevdi-gateway",
|
|
ExpiresAt: time.Now().Add(5 * time.Second).UTC().Format(time.RFC3339Nano),
|
|
Capabilities: DefaultCapabilities(), ProviderProfile: ProviderProfileApollo, ProviderIdentity: "apollo-fixture-1#sha256:fixture-apollo-1",
|
|
}
|
|
admission := &oneTimeAdmission{authority: authority, released: make(chan struct{}), disableClipboard: true}
|
|
reporter := &recordingProviderStateReporter{}
|
|
server, err := NewServer(ServerConfig{
|
|
ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID,
|
|
Capabilities: DefaultCapabilities(), ProviderCapabilities: DefaultCapabilities(),
|
|
Admission: admission, ProviderStateReporter: reporter, Provider: provider,
|
|
})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
serveDone := make(chan error, 1)
|
|
go func() { serveDone <- server.Serve(ctx) }()
|
|
request := protocol.TunnelAdmissionRequest{
|
|
Version: "1", SessionID: authority.SessionID, GatewayID: authority.GatewayID, Audience: authority.Audience,
|
|
Grant: strings.Repeat("g", 64), ClientNonce: "nonce-0000000001", DeviceSignature: strings.Repeat("s", 86),
|
|
Capabilities: DefaultCapabilities(),
|
|
}
|
|
client, err := dialIndependentGateway(context.Background(), server.Addr().String(), clientTLS, request)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
t.Cleanup(func() {
|
|
_ = client.Close()
|
|
cancel()
|
|
_ = server.Close()
|
|
if err := <-serveDone; err != nil {
|
|
t.Errorf("serve: %v", err)
|
|
}
|
|
})
|
|
return nativeGatewayLifecycleHarness{native: native, key: key, client: client, admission: admission, reporter: reporter}
|
|
}
|
|
|
|
func TestGatewayEgressReturnsProviderFreeClientAuthority(t *testing.T) {
|
|
h := newNativeGatewayLifecycleHarness(t, "session-client-authority")
|
|
var got map[string]any
|
|
if err := json.Unmarshal(h.client.authorityRaw, &got); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
want := map[string]any{
|
|
"version": "1",
|
|
"session_id": "session-client-authority",
|
|
"gateway_id": "gateway-1",
|
|
"audience": "versevdi-gateway",
|
|
"reconnect_sequence": float64(0),
|
|
"expires_at": h.admission.authority.ExpiresAt,
|
|
"capabilities": map[string]any{
|
|
"transport": "quic-tls13",
|
|
"framing": "datagram-v2",
|
|
"media": "encoded",
|
|
"audio": "encoded",
|
|
"source_rate_control": "server",
|
|
"client_decode": []any{"h264-opus"},
|
|
},
|
|
}
|
|
if !reflect.DeepEqual(got, want) {
|
|
t.Fatalf("raw client authority = %s", h.client.authorityRaw)
|
|
}
|
|
if err := h.client.Close(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
h.waitReleased(t)
|
|
if !reflect.DeepEqual(h.admission.releaseAuthority, h.admission.authority) {
|
|
t.Fatalf("release authority = %#v, want provider-bearing %#v", h.admission.releaseAuthority, h.admission.authority)
|
|
}
|
|
}
|
|
|
|
func TestClientSessionAuthorityRejectsProviderBearingRC3(t *testing.T) {
|
|
var _ protocol.ClientSessionAuthority = Client{}.Authority
|
|
raw, err := protocol.EncodeSessionAuthority(protocol.SessionAuthority{
|
|
Version: "1", SessionID: "session-rc3", GatewayID: "gateway-1", Audience: "versevdi-gateway",
|
|
ExpiresAt: time.Now().Add(time.Minute).UTC().Format(time.RFC3339Nano), Capabilities: DefaultCapabilities(),
|
|
ProviderProfile: ProviderProfileApollo, ProviderIdentity: "apollo-fixture-1#sha256:fixture-apollo-1",
|
|
})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, err := protocol.DecodeClientSessionAuthority(raw); err == nil {
|
|
t.Fatalf("accepted provider-bearing RC3 authority: %s", raw)
|
|
}
|
|
}
|
|
|
|
type independentGatewayClient struct {
|
|
connection *quic.Conn
|
|
control *quic.Stream
|
|
authority protocol.ClientSessionAuthority
|
|
authorityRaw []byte
|
|
media independentMediaReassembler
|
|
}
|
|
|
|
func dialIndependentGateway(ctx context.Context, address string, tlsConfig *tls.Config, request protocol.TunnelAdmissionRequest) (*independentGatewayClient, error) {
|
|
config := tlsConfig.Clone()
|
|
config.NextProtos = []string{"versevdi-gateway-v1"}
|
|
connection, err := quic.DialAddr(ctx, address, config, &quic.Config{EnableDatagrams: true, MaxIdleTimeout: 30 * time.Second})
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
stream, err := connection.OpenStreamSync(ctx)
|
|
if err != nil {
|
|
_ = connection.CloseWithError(applicationError, "independent client setup failed")
|
|
return nil, err
|
|
}
|
|
encoded, err := protocol.EncodeTunnelAdmissionRequest(request)
|
|
if err == nil {
|
|
err = independentWriteWire(stream, encoded)
|
|
}
|
|
if err == nil {
|
|
encoded, err = independentReadWire(stream, defaultHelloLimit)
|
|
}
|
|
var authority protocol.ClientSessionAuthority
|
|
if err == nil {
|
|
authority, err = protocol.DecodeClientSessionAuthority(encoded)
|
|
}
|
|
if err != nil {
|
|
_ = connection.CloseWithError(applicationError, "independent client admission failed")
|
|
return nil, err
|
|
}
|
|
return &independentGatewayClient{
|
|
connection: connection,
|
|
control: stream,
|
|
authority: authority,
|
|
authorityRaw: encoded,
|
|
media: independentMediaReassembler{incomplete: make(map[independentMediaKey]*independentMediaUnit)},
|
|
}, nil
|
|
}
|
|
|
|
func (c *independentGatewayClient) ReceiveProviderEvent(ctx context.Context) (ProviderEvent, error) {
|
|
return c.receiveProviderEvent(ctx, true)
|
|
}
|
|
|
|
func (c *independentGatewayClient) receiveProviderEvent(ctx context.Context, acknowledge bool) (ProviderEvent, error) {
|
|
if deadline, ok := ctx.Deadline(); ok {
|
|
if err := c.control.SetReadDeadline(deadline); err != nil {
|
|
return ProviderEvent{}, err
|
|
}
|
|
defer c.control.SetReadDeadline(time.Time{})
|
|
}
|
|
for {
|
|
encoded, err := independentReadWire(c.control, defaultControlLimit)
|
|
if err != nil {
|
|
return ProviderEvent{}, err
|
|
}
|
|
frame, err := protocol.DecodeChannelFrame(encoded)
|
|
if err != nil || frame.FlowID != "control.ack.v1" {
|
|
return ProviderEvent{}, ErrProviderMalformed
|
|
}
|
|
payload, err := base64.StdEncoding.DecodeString(frame.Payload)
|
|
if err != nil {
|
|
return ProviderEvent{}, err
|
|
}
|
|
event, err := DecodeProviderEvent(payload)
|
|
if acknowledge && err == nil && (event.Kind == ProviderEventTerminated || event.Kind == ProviderEventDisconnected) {
|
|
ack, encodeErr := protocol.EncodeChannelFrame(testChannelFrame("control.ack.v1", 0, []byte(protocolTerminalReceiptVector)))
|
|
if encodeErr != nil {
|
|
return ProviderEvent{}, encodeErr
|
|
}
|
|
if encodeErr = independentWriteWire(c.control, ack); encodeErr != nil {
|
|
return ProviderEvent{}, encodeErr
|
|
}
|
|
}
|
|
return event, err
|
|
}
|
|
}
|
|
|
|
func (c *independentGatewayClient) ReceiveFrame(ctx context.Context) (Frame, error) {
|
|
data, err := c.connection.ReceiveDatagram(ctx)
|
|
if err != nil {
|
|
return Frame{}, err
|
|
}
|
|
return independentDecodeFrame(data)
|
|
}
|
|
|
|
func (c *independentGatewayClient) ReceiveMedia(ctx context.Context) ([]byte, error) {
|
|
return c.receiveMedia(ctx, nil)
|
|
}
|
|
|
|
func (c *independentGatewayClient) receiveMedia(ctx context.Context, observe func(time.Time, int)) ([]byte, error) {
|
|
for {
|
|
data, err := c.connection.ReceiveDatagram(ctx)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
receivedAt := time.Now()
|
|
if observe != nil {
|
|
observe(receivedAt, len(data))
|
|
}
|
|
payload, complete, err := c.media.Add(data, receivedAt)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if complete {
|
|
return payload, nil
|
|
}
|
|
}
|
|
}
|
|
|
|
type independentMediaKey struct {
|
|
channel byte
|
|
sequence uint32
|
|
}
|
|
|
|
type independentMediaUnit struct {
|
|
started time.Time
|
|
timestamp uint64
|
|
fragments [][]byte
|
|
received []bool
|
|
bytes int
|
|
}
|
|
|
|
type independentMediaReassembler struct {
|
|
incomplete map[independentMediaKey]*independentMediaUnit
|
|
}
|
|
|
|
func independentDecodeFrame(data []byte) (Frame, error) {
|
|
const (
|
|
v1Header, v2Header = 21, 23
|
|
v1Payload, v2Payload = 1179, 1177
|
|
)
|
|
if len(data) < 3 {
|
|
return Frame{}, ErrFrameTruncated
|
|
}
|
|
if data[0] != 'V' || data[1] != 'D' {
|
|
return Frame{}, ErrFrameMagic
|
|
}
|
|
version := data[2]
|
|
headerSize, payloadLimit := v1Header, v1Payload
|
|
if version == 2 {
|
|
headerSize, payloadLimit = v2Header, v2Payload
|
|
} else if version != 1 {
|
|
return Frame{}, ErrFrameVersion
|
|
}
|
|
if len(data) < headerSize || version == 2 && len(data) > 1200 {
|
|
return Frame{}, ErrFrameSize
|
|
}
|
|
frame := Frame{
|
|
Version: version,
|
|
Channel: data[3],
|
|
Flags: data[4],
|
|
Sequence: binary.BigEndian.Uint32(data[5:9]),
|
|
TimestampMS: binary.BigEndian.Uint64(data[9:17]),
|
|
}
|
|
payloadLengthOffset := 19
|
|
if version == 1 {
|
|
frame.FragmentIndex = uint16(data[17])
|
|
frame.FragmentCount = uint16(data[18])
|
|
} else {
|
|
frame.FragmentIndex = binary.BigEndian.Uint16(data[17:19])
|
|
frame.FragmentCount = binary.BigEndian.Uint16(data[19:21])
|
|
payloadLengthOffset = 21
|
|
}
|
|
if (frame.Channel != ChannelVideo && frame.Channel != ChannelAudio) || frame.Flags != 0 ||
|
|
frame.FragmentCount == 0 || version == 1 && frame.FragmentCount > 16 ||
|
|
version == 2 && frame.FragmentCount > 891 || frame.FragmentIndex >= frame.FragmentCount {
|
|
return Frame{}, ErrFrameFragment
|
|
}
|
|
payloadLength := int(binary.BigEndian.Uint16(data[payloadLengthOffset : payloadLengthOffset+2]))
|
|
if payloadLength > payloadLimit || len(data) != headerSize+payloadLength {
|
|
return Frame{}, ErrFrameLength
|
|
}
|
|
frame.Payload = append([]byte(nil), data[headerSize:]...)
|
|
return frame, nil
|
|
}
|
|
|
|
func (r *independentMediaReassembler) Add(data []byte, now time.Time) ([]byte, bool, error) {
|
|
frame, err := independentDecodeFrame(data)
|
|
if err != nil {
|
|
return nil, false, err
|
|
}
|
|
for key, unit := range r.incomplete {
|
|
if now.Sub(unit.started) > 250*time.Millisecond {
|
|
delete(r.incomplete, key)
|
|
}
|
|
}
|
|
key := independentMediaKey{channel: frame.Channel, sequence: frame.Sequence}
|
|
unit := r.incomplete[key]
|
|
if unit == nil {
|
|
if len(r.incomplete) == 4 {
|
|
var oldestKey independentMediaKey
|
|
var oldest time.Time
|
|
for candidate, current := range r.incomplete {
|
|
if oldest.IsZero() || current.started.Before(oldest) {
|
|
oldestKey, oldest = candidate, current.started
|
|
}
|
|
}
|
|
delete(r.incomplete, oldestKey)
|
|
}
|
|
unit = &independentMediaUnit{
|
|
started: now, timestamp: frame.TimestampMS,
|
|
fragments: make([][]byte, frame.FragmentCount), received: make([]bool, frame.FragmentCount),
|
|
}
|
|
r.incomplete[key] = unit
|
|
}
|
|
if len(unit.fragments) != int(frame.FragmentCount) || unit.timestamp != frame.TimestampMS {
|
|
return nil, false, ErrFrameFragment
|
|
}
|
|
index := int(frame.FragmentIndex)
|
|
if unit.received[index] {
|
|
if !bytes.Equal(unit.fragments[index], frame.Payload) {
|
|
return nil, false, ErrFrameFragment
|
|
}
|
|
return nil, false, nil
|
|
}
|
|
if unit.bytes+len(frame.Payload) > 1<<20 {
|
|
delete(r.incomplete, key)
|
|
return nil, false, ErrFrameSize
|
|
}
|
|
unit.fragments[index] = frame.Payload
|
|
unit.received[index] = true
|
|
unit.bytes += len(frame.Payload)
|
|
for _, received := range unit.received {
|
|
if !received {
|
|
return nil, false, nil
|
|
}
|
|
}
|
|
payload := make([]byte, 0, unit.bytes)
|
|
for _, fragment := range unit.fragments {
|
|
payload = append(payload, fragment...)
|
|
}
|
|
delete(r.incomplete, key)
|
|
return payload, true, nil
|
|
}
|
|
|
|
func TestIndependentClientReassemblesProtocolDatagramV2(t *testing.T) {
|
|
fixed, err := hex.DecodeString("5644020a00000000010000000000000002000000010003010203")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
frame, err := independentDecodeFrame(fixed)
|
|
if err != nil || frame.Version != 2 || frame.Channel != ChannelVideo ||
|
|
frame.Sequence != 1 || frame.TimestampMS != 2 || frame.FragmentCount != 1 ||
|
|
!bytes.Equal(frame.Payload, []byte{1, 2, 3}) {
|
|
t.Fatalf("fixed Protocol v2 frame = %#v, %v", frame, err)
|
|
}
|
|
|
|
payload := bytes.Repeat([]byte("frame-boundary-"), 300)
|
|
fragments, err := FragmentPayload(ChannelVideo, 7, 11, payload)
|
|
if err != nil || len(fragments) < 3 {
|
|
t.Fatalf("fragments = %d, %v", len(fragments), err)
|
|
}
|
|
encoded := make([][]byte, len(fragments))
|
|
for index, fragment := range fragments {
|
|
encoded[index], err = EncodeFrame(fragment)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
reassembler := independentMediaReassembler{incomplete: make(map[independentMediaKey]*independentMediaUnit)}
|
|
now := time.Unix(0, 0)
|
|
order := []int{2, 0, 0, 1, 3}
|
|
var recovered []byte
|
|
for _, index := range order {
|
|
var complete bool
|
|
recovered, complete, err = reassembler.Add(encoded[index], now)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if index != 3 && complete {
|
|
t.Fatalf("unit completed at fragment %d", index)
|
|
}
|
|
}
|
|
if !bytes.Equal(recovered, payload) {
|
|
t.Fatal("independent client changed the complete encoded frame")
|
|
}
|
|
|
|
if _, _, err := reassembler.Add(encoded[0], now); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
conflict := append([]byte(nil), encoded[0]...)
|
|
conflict[len(conflict)-1] ^= 0xff
|
|
if _, _, err := reassembler.Add(conflict, now); !errors.Is(err, ErrFrameFragment) {
|
|
t.Fatalf("conflicting duplicate = %v", err)
|
|
}
|
|
reassembler = independentMediaReassembler{incomplete: make(map[independentMediaKey]*independentMediaUnit)}
|
|
if _, _, err := reassembler.Add(encoded[0], now); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
expired := append([]byte(nil), encoded[0]...)
|
|
binary.BigEndian.PutUint32(expired[5:9], 8)
|
|
if _, _, err := reassembler.Add(expired, now.Add(251*time.Millisecond)); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if len(reassembler.incomplete) != 1 {
|
|
t.Fatalf("expired incomplete units = %d, want 1", len(reassembler.incomplete))
|
|
}
|
|
|
|
reassembler = independentMediaReassembler{incomplete: make(map[independentMediaKey]*independentMediaUnit)}
|
|
for sequence := uint32(1); sequence <= 5; sequence++ {
|
|
partial, encodeErr := EncodeFrame(Frame{
|
|
Version: 2, Channel: ChannelVideo, Sequence: sequence, TimestampMS: 1,
|
|
FragmentIndex: 0, FragmentCount: 2, Payload: []byte{byte(sequence)},
|
|
})
|
|
if encodeErr != nil {
|
|
t.Fatal(encodeErr)
|
|
}
|
|
if _, _, err := reassembler.Add(partial, now.Add(time.Duration(sequence)*time.Nanosecond)); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
if len(reassembler.incomplete) != 4 || reassembler.incomplete[independentMediaKey{channel: ChannelVideo, sequence: 1}] != nil {
|
|
t.Fatalf("fifth unit did not evict the oldest: %#v", reassembler.incomplete)
|
|
}
|
|
|
|
reassembler = independentMediaReassembler{incomplete: make(map[independentMediaKey]*independentMediaUnit)}
|
|
for index := uint16(0); index < 891; index++ {
|
|
fragment, encodeErr := EncodeFrame(Frame{
|
|
Version: 2, Channel: ChannelVideo, Sequence: 99, TimestampMS: 1,
|
|
FragmentIndex: index, FragmentCount: 891, Payload: make([]byte, 1177),
|
|
})
|
|
if encodeErr != nil {
|
|
t.Fatal(encodeErr)
|
|
}
|
|
_, _, err = reassembler.Add(fragment, now)
|
|
}
|
|
if !errors.Is(err, ErrFrameSize) || len(reassembler.incomplete) != 0 {
|
|
t.Fatalf("oversized reassembly = %v, incomplete=%d", err, len(reassembler.incomplete))
|
|
}
|
|
}
|
|
|
|
func (c *independentGatewayClient) waitClosed(t *testing.T) {
|
|
t.Helper()
|
|
select {
|
|
case <-c.connection.Context().Done():
|
|
case <-time.After(terminalAckTimeout + time.Second):
|
|
t.Fatal("gateway did not close the terminal QUIC connection")
|
|
}
|
|
}
|
|
|
|
func (c *independentGatewayClient) Close() error {
|
|
return c.connection.CloseWithError(applicationError, "independent client closed")
|
|
}
|
|
|
|
func independentWriteWire(writer io.Writer, payload []byte) error {
|
|
var header [4]byte
|
|
binary.BigEndian.PutUint32(header[:], uint32(len(payload)))
|
|
if _, err := writer.Write(header[:]); err != nil {
|
|
return err
|
|
}
|
|
_, err := writer.Write(payload)
|
|
return err
|
|
}
|
|
|
|
func independentReadWire(reader io.Reader, limit int) ([]byte, error) {
|
|
var header [4]byte
|
|
if _, err := io.ReadFull(reader, header[:]); err != nil {
|
|
return nil, err
|
|
}
|
|
length := binary.BigEndian.Uint32(header[:])
|
|
if length > uint32(limit) {
|
|
return nil, ErrFrameSize
|
|
}
|
|
payload := make([]byte, length)
|
|
_, err := io.ReadFull(reader, payload)
|
|
return payload, err
|
|
}
|
|
|
|
func (h nativeGatewayLifecycleHarness) waitReleased(t *testing.T) {
|
|
t.Helper()
|
|
select {
|
|
case <-h.admission.released:
|
|
case <-time.After(2 * time.Second):
|
|
t.Fatal("native terminal state did not release admission")
|
|
}
|
|
}
|
|
|
|
func (h nativeGatewayLifecycleHarness) assertNoMedia(t *testing.T) {
|
|
t.Helper()
|
|
ctx, cancel := context.WithTimeout(context.Background(), 150*time.Millisecond)
|
|
defer cancel()
|
|
if frame, err := h.client.ReceiveFrame(ctx); err == nil {
|
|
t.Fatalf("media crossed after native terminal signal: channel=%d payload=%q", frame.Channel, frame.Payload)
|
|
}
|
|
}
|
|
|
|
func tryQueueNativeMedia(session *nativeApolloSession, channel chan ProviderMedia, payload []byte) {
|
|
session.enqueueMedia(channel, payload, time.Now())
|
|
}
|
|
|
|
func TestProviderDisconnectEndsPublicGatewaySessionReconnectable(t *testing.T) {
|
|
h := newGatewayTransportHarnessWithoutClipboard(t)
|
|
h.drainInitialMedia(t)
|
|
h.session.Disconnect()
|
|
|
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
|
if event, err := h.client.ReceiveProviderEvent(ctx); err != nil || event.Kind != ProviderEventDisconnected {
|
|
t.Fatalf("provider disconnect event = %#v, %v", event, err)
|
|
}
|
|
cancel()
|
|
h.waitReleased(t)
|
|
if states := h.reporter.States(); len(states) == 0 || states[len(states)-1].State != ProviderStateDisconnected || states[len(states)-1].CleanupPending {
|
|
t.Fatalf("provider states = %#v", states)
|
|
}
|
|
h.assertMediaClosed(t)
|
|
}
|
|
|
|
func TestProviderTerminalCleanupFailureReportsCleanupPending(t *testing.T) {
|
|
h := newGatewayTransportHarnessWithoutClipboard(t)
|
|
h.drainInitialMedia(t)
|
|
h.session.mu.Lock()
|
|
h.session.failure = FakeFailureTerminationTimeout
|
|
h.session.mu.Unlock()
|
|
h.session.EmitEvent(ProviderEvent{Kind: ProviderEventTerminated, Payload: []byte{1, 2, 3, 4}})
|
|
|
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
|
if event, err := h.client.ReceiveProviderEvent(ctx); err != nil || event.Kind != ProviderEventTerminated {
|
|
t.Fatalf("provider termination event = %#v, %v", event, err)
|
|
}
|
|
cancel()
|
|
h.waitReleased(t)
|
|
if states := h.reporter.States(); len(states) == 0 || states[len(states)-1].State != ProviderStateCleanup || !states[len(states)-1].CleanupPending {
|
|
t.Fatalf("provider states = %#v", states)
|
|
}
|
|
h.assertMediaClosed(t)
|
|
}
|
|
|
|
func testChannelFrame(flowID string, sequence int64, payload []byte) protocol.ChannelFrame {
|
|
return protocol.ChannelFrame{Version: "1", FlowID: flowID, Sequence: sequence, Flags: 0, FragmentIndex: 0, FragmentCount: 1, TimestampMs: time.Now().UnixMilli(), Payload: base64.StdEncoding.EncodeToString(payload)}
|
|
}
|
|
|
|
type gatewayTransportHarness struct {
|
|
client *Client
|
|
session *fakeSession
|
|
admission *oneTimeAdmission
|
|
reporter *recordingProviderStateReporter
|
|
server *Server
|
|
}
|
|
|
|
type providerStartFunc func(context.Context, LaunchRequest) (ProviderSession, error)
|
|
|
|
func (fn providerStartFunc) Start(ctx context.Context, request LaunchRequest) (ProviderSession, error) {
|
|
return fn(ctx, request)
|
|
}
|
|
|
|
type nativeLifecycleSession struct{ *nativeApolloSession }
|
|
|
|
func (s nativeLifecycleSession) Telemetry() ProviderTelemetry {
|
|
return ProviderTelemetry{State: s.State().State}
|
|
}
|
|
|
|
func newGatewayTransportHarness(t *testing.T) gatewayTransportHarness {
|
|
return newGatewayTransportHarnessWithClipboard(t, true)
|
|
}
|
|
|
|
func newGatewayTransportHarnessWithoutClipboard(t *testing.T) gatewayTransportHarness {
|
|
return newGatewayTransportHarnessWithClipboard(t, false)
|
|
}
|
|
|
|
func newGatewayTransportHarnessWithClipboard(t *testing.T, clipboardEnabled bool) gatewayTransportHarness {
|
|
t.Helper()
|
|
serverTLS, clientTLS := testTLS(t)
|
|
fake := NewFakeApollo(FakeApolloConfig{Now: time.Date(2026, time.January, 1, 0, 0, 0, 0, time.UTC)})
|
|
authority := protocol.SessionAuthority{Version: "1", SessionID: "session-transport", GatewayID: "gateway-1", Audience: "versevdi-gateway", ReconnectSequence: 0, ExpiresAt: time.Now().Add(5 * time.Second).UTC().Format(time.RFC3339Nano), Capabilities: DefaultCapabilities(), ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key()}
|
|
admission := &oneTimeAdmission{authority: authority, released: make(chan struct{}), disableClipboard: !clipboardEnabled}
|
|
reporter := &recordingProviderStateReporter{}
|
|
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: "gateway-1", Features: DefaultFeatures(), Capabilities: DefaultCapabilities(), ProviderCapabilities: DefaultCapabilities(), Admission: admission, ProviderStateReporter: reporter, ClipboardAuditReporter: reporter, Provider: fake})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
serveDone := make(chan error, 1)
|
|
go func() { serveDone <- server.Serve(ctx) }()
|
|
request := protocol.TunnelAdmissionRequest{Version: "1", SessionID: authority.SessionID, GatewayID: authority.GatewayID, Audience: authority.Audience, Grant: strings.Repeat("g", 64), ReconnectSequence: 0, ClientNonce: "nonce-0000000001", DeviceSignature: strings.Repeat("s", 86), Capabilities: DefaultCapabilities()}
|
|
client, err := Dial(context.Background(), server.Addr().String(), clientTLS, request)
|
|
if err != nil {
|
|
cancel()
|
|
_ = server.Close()
|
|
t.Fatal(err)
|
|
}
|
|
session, ok := fake.LastSession().(*fakeSession)
|
|
if !ok {
|
|
t.Fatal("fake provider session type")
|
|
}
|
|
t.Cleanup(func() {
|
|
_ = client.Close()
|
|
cancel()
|
|
_ = server.Close()
|
|
if err := <-serveDone; err != nil {
|
|
t.Errorf("serve: %v", err)
|
|
}
|
|
})
|
|
return gatewayTransportHarness{client: client, session: session, admission: admission, reporter: reporter, server: server}
|
|
}
|
|
|
|
func (h gatewayTransportHarness) waitReleased(t *testing.T) {
|
|
t.Helper()
|
|
select {
|
|
case <-h.admission.released:
|
|
case <-time.After(2 * time.Second):
|
|
t.Fatal("gateway did not reject the channel")
|
|
}
|
|
}
|
|
|
|
func (h gatewayTransportHarness) drainInitialMedia(t *testing.T) {
|
|
t.Helper()
|
|
for range 2 {
|
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
|
_, err := h.client.ReceiveFrame(ctx)
|
|
cancel()
|
|
if err != nil {
|
|
t.Fatalf("drain initial media: %v", err)
|
|
}
|
|
}
|
|
}
|
|
|
|
func (h gatewayTransportHarness) assertMediaClosed(t *testing.T) {
|
|
t.Helper()
|
|
h.session.EmitVideo([]byte("must-not-forward"))
|
|
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
|
|
defer cancel()
|
|
if frame, err := h.client.ReceiveFrame(ctx); err == nil {
|
|
t.Fatalf("media remained open after provider terminal state: %#v", frame)
|
|
}
|
|
}
|
|
|
|
func testTLS(t *testing.T) (*tls.Config, *tls.Config) {
|
|
t.Helper()
|
|
caKey, _ := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
|
caTemplate := &x509.Certificate{SerialNumber: big.NewInt(1), Subject: pkix.Name{CommonName: "Verse Test CA"}, NotBefore: time.Now().Add(-time.Hour), NotAfter: time.Now().Add(time.Hour), IsCA: true, BasicConstraintsValid: true, KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageDigitalSignature}
|
|
caDER, err := x509.CreateCertificate(rand.Reader, caTemplate, caTemplate, &caKey.PublicKey, caKey)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
caCert, err := x509.ParseCertificate(caDER)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
makeLeaf := func(serial int64, dns string, usage x509.ExtKeyUsage) tls.Certificate {
|
|
key, keyErr := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
|
if keyErr != nil {
|
|
t.Fatal(keyErr)
|
|
}
|
|
template := &x509.Certificate{SerialNumber: big.NewInt(serial), Subject: pkix.Name{CommonName: dns}, DNSNames: []string{dns}, IPAddresses: []net.IP{net.ParseIP("127.0.0.1")}, NotBefore: time.Now().Add(-time.Hour), NotAfter: time.Now().Add(time.Hour), ExtKeyUsage: []x509.ExtKeyUsage{usage}, KeyUsage: x509.KeyUsageDigitalSignature}
|
|
der, createErr := x509.CreateCertificate(rand.Reader, template, caCert, &key.PublicKey, caKey)
|
|
if createErr != nil {
|
|
t.Fatal(createErr)
|
|
}
|
|
return tls.Certificate{Certificate: [][]byte{der, caDER}, PrivateKey: key}
|
|
}
|
|
serverCert := makeLeaf(2, "gateway.test", x509.ExtKeyUsageServerAuth)
|
|
clientCert := makeLeaf(3, "client.test", x509.ExtKeyUsageClientAuth)
|
|
pool := x509.NewCertPool()
|
|
pool.AddCert(caCert)
|
|
return &tls.Config{MinVersion: tls.VersionTLS13, Certificates: []tls.Certificate{serverCert}, ClientAuth: tls.RequireAndVerifyClientCert, ClientCAs: pool}, &tls.Config{MinVersion: tls.VersionTLS13, Certificates: []tls.Certificate{clientCert}, RootCAs: pool, ServerName: "gateway.test"}
|
|
}
|
|
|
|
type oneTimeAdmission struct {
|
|
used atomic.Bool
|
|
authority protocol.SessionAuthority
|
|
releaseAuthority protocol.SessionAuthority
|
|
releases atomic.Int64
|
|
released chan struct{}
|
|
streamPolicy protocol.ProviderStreamPolicy
|
|
providerWork *protocol.ProviderSessionWork
|
|
disableClipboard bool
|
|
releaseErr error
|
|
}
|
|
|
|
type recordingProviderStateReporter struct {
|
|
mu sync.Mutex
|
|
states []protocol.ProviderState
|
|
audits []protocol.GatewayClipboardAudit
|
|
}
|
|
|
|
func (r *recordingProviderStateReporter) ReportProviderState(_ context.Context, state protocol.ProviderState) error {
|
|
r.mu.Lock()
|
|
defer r.mu.Unlock()
|
|
r.states = append(r.states, state)
|
|
return nil
|
|
}
|
|
|
|
func (r *recordingProviderStateReporter) States() []protocol.ProviderState {
|
|
r.mu.Lock()
|
|
defer r.mu.Unlock()
|
|
return append([]protocol.ProviderState(nil), r.states...)
|
|
}
|
|
|
|
func (r *recordingProviderStateReporter) ReportClipboardAudit(_ context.Context, audit protocol.GatewayClipboardAudit) error {
|
|
r.mu.Lock()
|
|
defer r.mu.Unlock()
|
|
r.audits = append(r.audits, audit)
|
|
return nil
|
|
}
|
|
|
|
func (r *recordingProviderStateReporter) Audits() []protocol.GatewayClipboardAudit {
|
|
r.mu.Lock()
|
|
defer r.mu.Unlock()
|
|
return append([]protocol.GatewayClipboardAudit(nil), r.audits...)
|
|
}
|
|
|
|
func (a *oneTimeAdmission) Admit(context.Context, protocol.TunnelAdmissionRequest) (protocol.SessionAuthority, error) {
|
|
if !a.used.CompareAndSwap(false, true) {
|
|
return protocol.SessionAuthority{}, ErrAdmissionRejected
|
|
}
|
|
return a.authority, nil
|
|
}
|
|
|
|
func (a *oneTimeAdmission) ProviderWork(_ context.Context, authority protocol.SessionAuthority) (protocol.ProviderSessionWork, error) {
|
|
if !reflect.DeepEqual(authority, a.authority) {
|
|
return protocol.ProviderSessionWork{}, ErrAdmissionRejected
|
|
}
|
|
if a.providerWork != nil {
|
|
return *a.providerWork, nil
|
|
}
|
|
streamPolicy := a.streamPolicy
|
|
if streamPolicy == (protocol.ProviderStreamPolicy{}) {
|
|
streamPolicy = protocol.ProviderStreamPolicy{ResolutionWidth: 1920, ResolutionHeight: 1080, Fps: 60, Codec: "H264", BitrateKbps: 8000, AudioEnabled: true}
|
|
}
|
|
clipboardPolicy := protocol.ClipboardPolicy{MaxTextBytes: 65536, MaxUpdatesPerMinute: 30}
|
|
if !a.disableClipboard {
|
|
clipboardPolicy.ClientToProviderEnabled = true
|
|
clipboardPolicy.ProviderToClientEnabled = true
|
|
}
|
|
return protocol.ProviderSessionWork{
|
|
Version: "1", SessionID: authority.SessionID, GatewayID: authority.GatewayID,
|
|
ReconnectSequence: authority.ReconnectSequence, ExpiresAt: authority.ExpiresAt,
|
|
ProviderProfile: ProviderProfileApollo, ProviderIdentity: authority.ProviderIdentity,
|
|
PolicyVersionID: "policy-1", ApplicationID: "1", ClientID: "paired-client", ManagementHost: "apollo.test", ManagementPort: 47990,
|
|
StreamPolicy: streamPolicy,
|
|
StreamHost: "apollo.test", StreamPort: 47984, ClientCertificatePem: "certificate",
|
|
ClientPrivateKeyPem: "private-key", ServerCertificatePem: "server-certificate",
|
|
ClipboardPolicy: clipboardPolicy,
|
|
}, nil
|
|
}
|
|
|
|
func (a *oneTimeAdmission) Release(_ context.Context, authority protocol.SessionAuthority) error {
|
|
if a.releases.Add(1) == 1 {
|
|
a.releaseAuthority = authority
|
|
close(a.released)
|
|
}
|
|
return a.releaseErr
|
|
}
|
|
|
|
func mustRead(t *testing.T, path string) []byte {
|
|
t.Helper()
|
|
data, err := os.ReadFile(path)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return data
|
|
}
|
|
|
|
func bytesRepeat(value byte, count int) []byte {
|
|
data := make([]byte, count)
|
|
for index := range data {
|
|
data[index] = value
|
|
}
|
|
return data
|
|
}
|