fix(gateway): secure control and terminal ownership

This commit is contained in:
sechmachine
2026-07-30 11:15:48 +07:00
parent df75b1d250
commit baf4073f68
16 changed files with 488 additions and 36 deletions
+162 -8
View File
@@ -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)