Files
VerseVDI-Data-Plane/gateway/apollo_enet_test.go
T

389 lines
14 KiB
Go

package gateway
import (
"context"
"encoding/binary"
"encoding/hex"
"errors"
"net"
"testing"
"time"
)
func TestApolloENetConnectWireVectorAndRTO(t *testing.T) {
now := time.UnixMilli(0x1234)
packet := apolloENetConnectPacket(now, 0x01020304, 0x12345678)
expected, err := hex.DecodeString("8fff123482ff00010000ffff00000578000080000000003000000000000000000000138800000002000000020102030412345678")
if err != nil {
t.Fatal(err)
}
if string(packet) != string(expected) {
t.Fatalf("CONNECT wire bytes = %x, want %x", packet, expected)
}
if got := apolloENetRetransmitTimeout(100*time.Millisecond, 30*time.Millisecond, 1); got != 200*time.Millisecond {
t.Fatalf("initial RTO = %s", got)
}
if got := apolloENetRetransmitTimeout(100*time.Millisecond, 30*time.Millisecond, 6); got != 1200*time.Millisecond {
t.Fatalf("bounded retry RTO = %s", got)
}
}
func TestApolloENetConnectVerifies48ChannelsAndFlushesAck(t *testing.T) {
server, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
if err != nil {
t.Fatal(err)
}
defer server.Close()
client, err := net.DialUDP("udp", nil, server.LocalAddr().(*net.UDPAddr))
if err != nil {
t.Fatal(err)
}
peer, err := newApolloENetPeer(client, time.Now)
if err != nil {
t.Fatal(err)
}
defer peer.close(nil)
received := make(chan string, 3)
peer.onPayload = func(_ uint8, _ bool, payload []byte) { received <- string(payload) }
serverDone := make(chan error, 1)
go func() {
buffer := make([]byte, apolloENetMaximumPacket)
count, remote, readErr := server.ReadFromUDP(buffer)
if readErr != nil {
serverDone <- readErr
return
}
packet := buffer[:count]
if len(packet) != 52 || packet[4] != apolloENetConnect|apolloENetAcknowledged || packet[5] != 0xff || binary.BigEndian.Uint32(packet[20:24]) != apolloENetChannels || binary.BigEndian.Uint32(packet[48:52]) != 0x12345678 {
serverDone <- ErrProviderMalformed
return
}
verify := make([]byte, 48)
apolloENetHeader(verify, 0, 0, time.Now())
verify[4] = apolloENetVerifyConnect | apolloENetAcknowledged
verify[5] = 0xff
binary.BigEndian.PutUint16(verify[6:8], 1)
binary.BigEndian.PutUint16(verify[8:10], 7)
verify[10], verify[11] = 2, 3
binary.BigEndian.PutUint32(verify[12:16], 1400)
binary.BigEndian.PutUint32(verify[16:20], 32768)
binary.BigEndian.PutUint32(verify[20:24], apolloENetChannels)
binary.BigEndian.PutUint32(verify[44:48], binary.BigEndian.Uint32(packet[44:48]))
if _, writeErr := server.WriteToUDP(verify, remote); writeErr != nil {
serverDone <- writeErr
return
}
count, _, readErr = server.ReadFromUDP(buffer)
if readErr != nil {
serverDone <- readErr
return
}
ack := buffer[:count]
if len(ack) != 12 || ack[4]&apolloENetCommandMask != 1 || binary.BigEndian.Uint16(ack[8:10]) != 1 {
serverDone <- ErrProviderMalformed
return
}
serverDone <- nil
}()
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := peer.Connect(ctx, 0x12345678); err != nil {
t.Fatalf("Connect() error = %v", err)
}
for _, command := range []struct {
sequence uint16
payload string
}{{1, "A"}, {3, "C"}, {2, "B"}} {
if err := peer.handleDatagram(sourceShapedENetReliablePacket(peer.peerID, peer.inboundSession, command.sequence, []byte(command.payload))); err != nil {
t.Fatalf("handleDatagram() error = %v", err)
}
}
if err := <-serverDone; err != nil {
t.Fatalf("ENet server error = %v", err)
}
ordered := ""
for range 3 {
select {
case payload := <-received:
ordered += payload
case <-time.After(time.Second):
t.Fatalf("reliable payload order = %q, want ABC", ordered)
}
}
if ordered != "ABC" {
t.Fatalf("reliable payload order = %q, want ABC", ordered)
}
}
func TestApolloENetFakeClockRetransmitsLostReliablePacketAtForkRTO(t *testing.T) {
server, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
if err != nil {
t.Fatal(err)
}
defer server.Close()
client, err := net.DialUDP("udp", nil, server.LocalAddr().(*net.UDPAddr))
if err != nil {
t.Fatal(err)
}
now := time.Unix(1, 0)
peer, err := newApolloENetPeer(client, func() time.Time { return now })
if err != nil {
t.Fatal(err)
}
defer peer.close(nil)
peer.state, peer.peerID, peer.outboundSession, peer.lastPing = apolloENetConnected, 1, 2, now
peer.rtt, peer.variance = 100*time.Millisecond, 30*time.Millisecond
if err := peer.SendReliable(apolloChannelKeyboard, []byte{0x01}); err != nil {
t.Fatal(err)
}
assertApolloENetReliablePacket(t, server, 1)
now = now.Add(199 * time.Millisecond)
if err := peer.maintain(); err != nil {
t.Fatal(err)
}
assertApolloENetPending(t, peer, 1, 200*time.Millisecond)
now = now.Add(time.Millisecond)
if err := peer.maintain(); err != nil {
t.Fatal(err)
}
assertApolloENetReliablePacket(t, server, 1)
assertApolloENetPending(t, peer, 2, 200*time.Millisecond)
now = now.Add(199 * time.Millisecond)
if err := peer.maintain(); err != nil {
t.Fatal(err)
}
assertApolloENetPending(t, peer, 2, 200*time.Millisecond)
now = now.Add(time.Millisecond)
if err := peer.maintain(); err != nil {
t.Fatal(err)
}
assertApolloENetReliablePacket(t, server, 1)
assertApolloENetPending(t, peer, 3, 400*time.Millisecond)
telemetry := peer.telemetry()
if telemetry.ControlRTT != 100*time.Millisecond || telemetry.ControlJitter != 30*time.Millisecond || telemetry.ReliableSent != 1 || telemetry.ReliableRetransmits != 2 || telemetry.PendingReliable != 1 {
t.Fatalf("ENet telemetry = %#v", telemetry)
}
}
func TestApolloENetFakeClockKeepsSessionAlivePastElevenVirtualSeconds(t *testing.T) {
server, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
if err != nil {
t.Fatal(err)
}
defer server.Close()
client, err := net.DialUDP("udp", nil, server.LocalAddr().(*net.UDPAddr))
if err != nil {
t.Fatal(err)
}
now := time.Unix(1, 0)
peer, err := newApolloENetPeer(client, func() time.Time { return now })
if err != nil {
t.Fatal(err)
}
defer peer.close(nil)
peer.state, peer.peerID, peer.inboundSession, peer.outboundSession, peer.lastPing, peer.lastReceive = apolloENetConnected, 1, 1, 2, now, now
for virtual := 500 * time.Millisecond; virtual <= 11*time.Second; virtual += 500 * time.Millisecond {
now = time.Unix(1, 0).Add(virtual)
if err := peer.maintain(); err != nil {
t.Fatalf("maintain() at virtual %s = %v", virtual, err)
}
assertApolloENetPingPacket(t, server)
peer.mu.Lock()
sequence := peer.channels[0].nextOutgoing
peer.mu.Unlock()
if err := peer.handleDatagram(sourceShapedENetAcknowledgePacket(peer.peerID, peer.inboundSession, 0, sequence)); err != nil {
t.Fatalf("handleDatagram() ACK at virtual %s = %v", virtual, err)
}
}
if peer.state != apolloENetConnected || len(peer.pending) != 0 {
t.Fatalf("virtual ENet session state = %v pending=%d", peer.state, len(peer.pending))
}
}
func TestApolloENetDisconnectCompletesOnProviderAcknowledgement(t *testing.T) {
server, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
if err != nil {
t.Fatal(err)
}
defer server.Close()
client, err := net.DialUDP("udp", nil, server.LocalAddr().(*net.UDPAddr))
if err != nil {
t.Fatal(err)
}
peer, err := newApolloENetPeer(client, time.Now)
if err != nil {
t.Fatal(err)
}
defer peer.close(nil)
peer.state, peer.peerID, peer.inboundSession, peer.outboundSession = apolloENetConnected, 1, 1, 2
go peer.run()
serverDone := make(chan error, 1)
go func() {
buffer := make([]byte, apolloENetMaximumPacket)
count, remote, readErr := server.ReadFromUDP(buffer)
if readErr != nil {
serverDone <- readErr
return
}
packet := buffer[:count]
if count != 12 || packet[4]&apolloENetCommandMask != apolloENetDisconnect || packet[5] != 0xff {
serverDone <- ErrProviderMalformed
return
}
_, writeErr := server.WriteToUDP(sourceShapedENetAcknowledgePacket(1, 1, 0xff, binary.BigEndian.Uint16(packet[6:8])), remote)
serverDone <- writeErr
}()
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
defer cancel()
if err := peer.Disconnect(ctx); err != nil {
t.Fatalf("Disconnect() = %v after provider ACK", err)
}
if err := <-serverDone; err != nil {
t.Fatalf("provider ACK = %v", err)
}
}
func assertApolloENetReliablePacket(t *testing.T, server *net.UDPConn, sequence uint16) {
t.Helper()
buffer := make([]byte, apolloENetMaximumPacket)
if err := server.SetReadDeadline(time.Now().Add(time.Second)); err != nil {
t.Fatal(err)
}
count, _, err := server.ReadFromUDP(buffer)
if err != nil {
t.Fatal(err)
}
if count != 11 || buffer[4]&apolloENetCommandMask != apolloENetSendReliable || binary.BigEndian.Uint16(buffer[6:8]) != sequence {
t.Fatalf("retransmitted packet = %x", buffer[:count])
}
}
func assertApolloENetPingPacket(t *testing.T, server *net.UDPConn) {
t.Helper()
buffer := make([]byte, apolloENetMaximumPacket)
if err := server.SetReadDeadline(time.Now().Add(time.Second)); err != nil {
t.Fatal(err)
}
count, _, err := server.ReadFromUDP(buffer)
if err != nil {
t.Fatal(err)
}
if count != 8 || buffer[4]&apolloENetCommandMask != apolloENetPing || buffer[4]&apolloENetAcknowledged == 0 {
t.Fatalf("ping packet = %x", buffer[:count])
}
}
func assertApolloENetPending(t *testing.T, peer *apolloENetPeer, attempts uint8, timeout time.Duration) {
t.Helper()
pending := peer.pending[apolloENetPendingKey{channel: apolloChannelKeyboard, sequence: 1}]
if pending == nil || pending.attempts != attempts || pending.timeout != timeout {
t.Fatalf("pending reliable state = %#v, want attempts=%d timeout=%s", pending, attempts, timeout)
}
}
func sourceShapedENetReliablePacket(peerID uint16, session uint8, sequence uint16, payload []byte) []byte {
return sourceShapedENetReliablePacketOn(peerID, session, apolloChannelKeyboard, sequence, payload)
}
func sourceShapedENetReliablePacketOn(peerID uint16, session, channel uint8, sequence uint16, payload []byte) []byte {
packet := make([]byte, 10+len(payload))
binary.BigEndian.PutUint16(packet[:2], peerID|(uint16(session&3)<<apolloENetSessionShift)|apolloENetSentTimeFlag)
binary.BigEndian.PutUint16(packet[2:4], 0x1234)
packet[4] = apolloENetSendReliable | apolloENetAcknowledged
packet[5] = channel
binary.BigEndian.PutUint16(packet[6:8], sequence)
binary.BigEndian.PutUint16(packet[8:10], uint16(len(payload)))
copy(packet[10:], payload)
return packet
}
func sourceShapedENetAcknowledgePacket(peerID uint16, session, channel uint8, sequence uint16) []byte {
packet := make([]byte, 12)
binary.BigEndian.PutUint16(packet[:2], peerID|(uint16(session&3)<<apolloENetSessionShift)|apolloENetSentTimeFlag)
packet[4] = 1
packet[5] = channel
binary.BigEndian.PutUint16(packet[8:10], sequence)
return packet
}
func TestApolloControlWireVectorAndTagFailure(t *testing.T) {
codec, err := newApolloControlCodec([]byte("0123456789abcdef"))
if err != nil {
t.Fatal(err)
}
packet, err := codec.SealClient(apolloControlTypePing, []byte{4, 0, 0, 0, 0, 0})
if err != nil {
t.Fatal(err)
}
const expected = "01001e0000000000c688b2e3cb8a5869293f01293a19486c692d0799eb8e220a74b4"
if string(packet) != string(mustDecodeHex(t, expected)) {
t.Fatalf("control wire bytes = %x", packet)
}
tampered := append([]byte(nil), packet...)
tampered[len(tampered)-1] ^= 1
if _, err := codec.OpenHost(tampered); err == nil {
t.Fatal("OpenHost() accepted a tag failure")
}
}
func TestApolloKeyboardInputWireVector(t *testing.T) {
packets, err := encodeApolloInputEvent(InputEvent{Device: "keyboard", Code: 30, Pressed: true, Payload: []byte{2}})
if err != nil {
t.Fatal(err)
}
const expected = "0000000a03000000001e00020000"
if len(packets) != 1 || string(packets[0].payload) != string(mustDecodeHex(t, expected)) {
t.Fatalf("keyboard packets = %#v, want %s", packets, expected)
}
}
func TestApolloAbsoluteAndScrollInputWireVectors(t *testing.T) {
// Independently implemented from the approved Apollo adc5c5a0 input.cpp
// consumer and its moonlight-common-c c999436 Input.h/InputStream.c pin.
absolute, err := encodeApolloInputEvent(InputEvent{Device: "mouse-absolute", Payload: []byte{0x04, 0xd2, 0x02, 0x37, 0x0a, 0x00, 0x05, 0xa0}})
if err != nil {
t.Fatal(err)
}
const absoluteExpected = "0000000e0500000004d20237000009ff059f"
if len(absolute) != 1 || absolute[0].channel != apolloChannelMouse || string(absolute[0].payload) != string(mustDecodeHex(t, absoluteExpected)) {
t.Fatalf("absolute packets = %#v, want %s", absolute, absoluteExpected)
}
scroll, err := encodeApolloInputEvent(InputEvent{Device: "mouse-scroll", Payload: []byte{0xff, 0x88, 0x00, 0x78}})
if err != nil {
t.Fatal(err)
}
const verticalExpected = "0000000a0a000000ff88ff880000"
const horizontalExpected = "00000006010000550078"
if len(scroll) != 2 || scroll[0].channel != apolloChannelMouse || scroll[1].channel != apolloChannelMouse ||
string(scroll[0].payload) != string(mustDecodeHex(t, verticalExpected)) || string(scroll[1].payload) != string(mustDecodeHex(t, horizontalExpected)) {
t.Fatalf("scroll packets = %#v", scroll)
}
zero, err := encodeApolloInputEvent(InputEvent{Device: "mouse-scroll", Payload: make([]byte, 4)})
if err != nil || len(zero) != 0 {
t.Fatalf("zero scroll packets = %#v, %v", zero, err)
}
for _, payload := range [][]byte{
{0, 0, 0, 0, 0, 1, 0, 2},
{0, 0, 0, 0, 0x80, 0, 0, 2},
{0, 0, 0, 0, 0, 2, 0x80, 0},
} {
if _, err := encodeApolloInputEvent(InputEvent{Device: "mouse-absolute", Payload: payload}); !errors.Is(err, ErrInputMalformed) {
t.Fatalf("Apollo accepted unrepresentable absolute payload %x: %v", payload, err)
}
}
}
func mustDecodeHex(t *testing.T, value string) []byte {
t.Helper()
decoded, err := hex.DecodeString(value)
if err != nil {
t.Fatal(err)
}
return decoded
}