fix(gateway): secure control and terminal ownership
This commit is contained in:
+162
-8
@@ -9,8 +9,10 @@ import (
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"encoding/base64"
|
||||
"encoding/binary"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"io"
|
||||
"math/big"
|
||||
"net"
|
||||
"os"
|
||||
@@ -21,6 +23,8 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/quic-go/quic-go"
|
||||
|
||||
protocol "git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol/gen/go/protocol"
|
||||
)
|
||||
|
||||
@@ -105,6 +109,10 @@ func TestClientFeedbackUsesFixedProtocolVGFVector(t *testing.T) {
|
||||
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 || hex.EncodeToString(receipt) != "5647463100030000" {
|
||||
t.Fatalf("terminal receipt vector = %x, %v", receipt, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCapabilityIntersectionAndBoundedQueue(t *testing.T) {
|
||||
@@ -179,12 +187,6 @@ func TestSyntheticImpairmentPacingAndResourceBounds(t *testing.T) {
|
||||
if delivered != 3 || deliveredBytes != 1179*3 {
|
||||
t.Fatalf("synthetic impairment delivered=%d bytes=%d", delivered, deliveredBytes)
|
||||
}
|
||||
pacer := NewPacer(1)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Millisecond)
|
||||
defer cancel()
|
||||
if err := pacer.Wait(ctx, 100); !errors.Is(err, context.DeadlineExceeded) {
|
||||
t.Fatalf("pacer ignored bounded context: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApolloFixturesAndLifecycle(t *testing.T) {
|
||||
@@ -701,6 +703,11 @@ func TestProviderTerminationEndsPublicGatewaySession(t *testing.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)
|
||||
}
|
||||
|
||||
@@ -716,6 +723,7 @@ func TestEncryptedNativeHostTerminationEndsPublicGatewaySession(t *testing.T) {
|
||||
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)
|
||||
@@ -737,6 +745,7 @@ func TestNativeENetDisconnectQuiescesPublicGatewaySession(t *testing.T) {
|
||||
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)
|
||||
@@ -746,10 +755,24 @@ func TestNativeENetDisconnectQuiescesPublicGatewaySession(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
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 *Client
|
||||
client *independentGatewayClient
|
||||
admission *oneTimeAdmission
|
||||
reporter *recordingProviderStateReporter
|
||||
}
|
||||
@@ -795,7 +818,7 @@ func newNativeGatewayLifecycleHarness(t *testing.T, sessionID string) nativeGate
|
||||
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)
|
||||
client, err := dialIndependentGateway(context.Background(), server.Addr().String(), clientTLS, request)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -810,6 +833,127 @@ func newNativeGatewayLifecycleHarness(t *testing.T, sessionID string) nativeGate
|
||||
return nativeGatewayLifecycleHarness{native: native, key: key, client: client, admission: admission, reporter: reporter}
|
||||
}
|
||||
|
||||
type independentGatewayClient struct {
|
||||
connection *quic.Conn
|
||||
control *quic.Stream
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
if err == nil {
|
||||
_, err = protocol.DecodeSessionAuthority(encoded)
|
||||
}
|
||||
if err != nil {
|
||||
_ = connection.CloseWithError(applicationError, "independent client admission failed")
|
||||
return nil, err
|
||||
}
|
||||
return &independentGatewayClient{connection: connection, control: stream}, 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) {
|
||||
receipt, encodeErr := EncodeClientFeedback(Feedback{Kind: FeedbackTerminalReceipt})
|
||||
if encodeErr != nil {
|
||||
return ProviderEvent{}, encodeErr
|
||||
}
|
||||
ack, encodeErr := protocol.EncodeChannelFrame(testChannelFrame("control.ack.v1", 0, receipt))
|
||||
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 DecodeFrame(data)
|
||||
}
|
||||
|
||||
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 {
|
||||
@@ -837,6 +981,11 @@ func TestProviderDisconnectEndsPublicGatewaySessionReconnectable(t *testing.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)
|
||||
@@ -852,6 +1001,11 @@ func TestProviderTerminalCleanupFailureReportsCleanupPending(t *testing.T) {
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user