feat(gateway): repair native Apollo provider path

This commit is contained in:
sechmachine
2026-07-29 17:24:34 +07:00
parent 67eeb9e898
commit 236c7c2973
23 changed files with 5088 additions and 345 deletions
+1 -1
View File
@@ -57,7 +57,7 @@ func run() error {
controlPlaneClient := gateway.NewControlPlaneClient(controlPlane, &http.Client{Transport: transport, Timeout: 5 * time.Second}) controlPlaneClient := gateway.NewControlPlaneClient(controlPlane, &http.Client{Transport: transport, Timeout: 5 * time.Second})
provider := gateway.NewApolloAdapter(gateway.NewNativeApolloBackend(), gateway.ProviderIdentity{}) provider := gateway.NewApolloAdapter(gateway.NewNativeApolloBackend(), gateway.ProviderIdentity{})
capabilities := gateway.DefaultCapabilities() capabilities := gateway.DefaultCapabilities()
server, err := gateway.NewServer(gateway.ServerConfig{ListenAddress: listen, TLSConfig: serverTLS, GatewayID: gatewayID, Capabilities: capabilities, ProviderCapabilities: capabilities, Admission: controlPlaneClient, ProviderStateReporter: controlPlaneClient, Provider: provider, PacerKbps: 100000}) server, err := gateway.NewServer(gateway.ServerConfig{ListenAddress: listen, TLSConfig: serverTLS, GatewayID: gatewayID, Capabilities: capabilities, ProviderCapabilities: capabilities, Admission: controlPlaneClient, ProviderStateReporter: controlPlaneClient, ClipboardAuditReporter: controlPlaneClient, Provider: provider, PacerKbps: 100000})
if err != nil { if err != nil {
return err return err
} }
+144
View File
@@ -0,0 +1,144 @@
package gateway
const (
apolloAudioDataShards = 4
apolloAudioParityShards = 2
apolloAudioTotalShards = apolloAudioDataShards + apolloAudioParityShards
apolloAudioMaximumBlocks = 4
)
type apolloAudioFECBlock struct {
base uint16
timestamp uint32
ssrc uint32
haveFEC bool
size int
shards [apolloAudioTotalShards][]byte
received [apolloAudioTotalShards]bool
count int
}
type apolloAudioAssembler struct {
blocks map[uint16]*apolloAudioFECBlock
}
func (a *apolloAudioAssembler) Add(codec *apolloMediaCodec, shard apolloAudioShard) ([][]byte, error) {
if codec == nil || len(shard.payload) == 0 || len(shard.payload) > 1408 || len(shard.payload)%16 != 0 {
return nil, errApolloMedia
}
if a.blocks == nil {
a.blocks = make(map[uint16]*apolloAudioFECBlock)
}
base := shard.base
if base&3 != 0 {
return nil, errApolloMedia
}
block := a.blocks[base]
if block == nil {
if len(a.blocks) >= apolloAudioMaximumBlocks {
return nil, errApolloMedia
}
block = &apolloAudioFECBlock{base: base}
a.blocks[base] = block
}
if block.size == 0 {
block.size = len(shard.payload)
} else if block.size != len(shard.payload) {
return nil, errApolloMedia
}
index := 0
if shard.parity {
if shard.parityIndex >= apolloAudioParityShards {
return nil, errApolloMedia
}
index = apolloAudioDataShards + int(shard.parityIndex)
if block.haveFEC && (block.timestamp != shard.timestamp || block.ssrc != shard.ssrc) {
return nil, errApolloMedia
}
block.timestamp, block.ssrc, block.haveFEC = shard.timestamp, shard.ssrc, true
} else {
index = int(uint16(shard.sequence - base))
if index >= apolloAudioDataShards {
return nil, errApolloMedia
}
if block.haveFEC && (shard.timestamp != block.timestamp+uint32(index*5) || shard.ssrc != block.ssrc) {
return nil, errApolloMedia
}
}
if block.received[index] {
return nil, errApolloMedia
}
block.shards[index] = append([]byte(nil), shard.payload...)
block.received[index] = true
block.count++
if block.count < apolloAudioDataShards {
return nil, nil
}
if err := reconstructApolloAudioBlock(block); err != nil {
return nil, err
}
output := make([][]byte, apolloAudioDataShards)
for index := range output {
payload, err := codec.openApolloAudioCipher(base+uint16(index), block.shards[index])
if err != nil {
return nil, err
}
output[index] = payload
}
delete(a.blocks, base)
return output, nil
}
func reconstructApolloAudioBlock(block *apolloAudioFECBlock) error {
if block == nil || block.count < apolloAudioDataShards || block.size == 0 {
return errApolloMedia
}
missing := false
for index := 0; index < apolloAudioDataShards; index++ {
if !block.received[index] {
missing = true
block.shards[index] = make([]byte, block.size)
}
}
if !missing {
return nil
}
if !block.haveFEC {
return errApolloMedia
}
rows := make([][]byte, 0, apolloAudioDataShards)
shards := make([][]byte, 0, apolloAudioDataShards)
for index, received := range block.received {
if !received {
continue
}
rows = append(rows, apolloAudioFECRow(index))
shards = append(shards, block.shards[index])
if len(rows) == apolloAudioDataShards {
break
}
}
inverse, ok := apolloGFInvert(rows)
if !ok {
return errApolloMedia
}
for index := 0; index < apolloAudioDataShards; index++ {
if block.received[index] {
continue
}
for source, coefficient := range inverse[index] {
apolloGFAXPY(block.shards[index], shards[source], coefficient)
}
}
return nil
}
func apolloAudioFECRow(index int) []byte {
if index < apolloAudioDataShards {
row := make([]byte, apolloAudioDataShards)
row[index] = 1
return row
}
parity := [8]byte{0x77, 0x40, 0x38, 0x0e, 0xc7, 0xa7, 0x0d, 0x6c}
return append([]byte(nil), parity[(index-apolloAudioDataShards)*apolloAudioDataShards:(index-apolloAudioDataShards+1)*apolloAudioDataShards]...)
}
+106
View File
@@ -0,0 +1,106 @@
package gateway
import (
"crypto/aes"
"crypto/cipher"
"encoding/binary"
"errors"
)
const (
apolloControlOuterType = 0x0001
apolloControlHeaderSize = 8
apolloControlTagSize = 16
apolloControlInnerSize = 4
apolloControlMaximumPlain = 2048
apolloControlTypeStart = 0x0307
apolloControlTypeIDR = 0x0302
apolloControlTypePing = 0x0200
apolloControlTypeInput = 0x0206
apolloControlTypeFEC = 0x5502
apolloControlTypeRumble = 0x010b
apolloControlTypeHDR = 0x010e
apolloControlTypeTerm = 0x0109
)
var errApolloControl = errors.New("apollo control malformed")
type apolloControlMessage struct {
typeID uint16
payload []byte
}
type apolloControlCodec struct {
aead cipher.AEAD
nextClient uint32
lastHost uint32
hostSeen bool
}
func newApolloControlCodec(key []byte) (*apolloControlCodec, error) {
block, err := aes.NewCipher(key)
if err != nil {
return nil, err
}
aead, err := cipher.NewGCM(block)
if err != nil {
return nil, err
}
return &apolloControlCodec{aead: aead}, nil
}
func (c *apolloControlCodec) SealClient(typeID uint16, payload []byte) ([]byte, error) {
if c == nil || c.aead == nil || typeID == 0 || len(payload) > apolloControlMaximumPlain || c.nextClient == ^uint32(0) {
return nil, errApolloControl
}
sequence := c.nextClient
c.nextClient++
inner := make([]byte, apolloControlInnerSize+len(payload))
binary.LittleEndian.PutUint16(inner[:2], typeID)
binary.LittleEndian.PutUint16(inner[2:4], uint16(len(payload)))
copy(inner[4:], payload)
nonce := apolloControlNonce(sequence, 'C')
sealed := c.aead.Seal(nil, nonce[:], inner, nil)
packet := make([]byte, apolloControlHeaderSize+len(sealed))
binary.LittleEndian.PutUint16(packet[:2], apolloControlOuterType)
binary.LittleEndian.PutUint16(packet[2:4], uint16(4+len(sealed)))
binary.LittleEndian.PutUint32(packet[4:8], sequence)
copy(packet[8:24], sealed[len(inner):])
copy(packet[24:], sealed[:len(inner)])
return packet, nil
}
func (c *apolloControlCodec) OpenHost(packet []byte) (apolloControlMessage, error) {
if c == nil || c.aead == nil || len(packet) < apolloControlHeaderSize+apolloControlTagSize+apolloControlInnerSize || len(packet) > apolloControlHeaderSize+apolloControlTagSize+apolloControlMaximumPlain {
return apolloControlMessage{}, errApolloControl
}
if binary.LittleEndian.Uint16(packet[:2]) != apolloControlOuterType || int(binary.LittleEndian.Uint16(packet[2:4])) != len(packet)-4 {
return apolloControlMessage{}, errApolloControl
}
sequence := binary.LittleEndian.Uint32(packet[4:8])
if c.hostSeen && sequence <= c.lastHost {
return apolloControlMessage{}, errApolloControl
}
nonce := apolloControlNonce(sequence, 'H')
sealed := make([]byte, len(packet)-apolloControlHeaderSize)
copy(sealed, packet[24:])
copy(sealed[len(packet)-24:], packet[8:24])
plaintext, err := c.aead.Open(nil, nonce[:], sealed, nil)
if err != nil || len(plaintext) < apolloControlInnerSize {
return apolloControlMessage{}, errApolloControl
}
length := int(binary.LittleEndian.Uint16(plaintext[2:4]))
if length != len(plaintext)-apolloControlInnerSize || length > apolloControlMaximumPlain {
return apolloControlMessage{}, errApolloControl
}
c.lastHost, c.hostSeen = sequence, true
return apolloControlMessage{typeID: binary.LittleEndian.Uint16(plaintext[:2]), payload: append([]byte(nil), plaintext[4:]...)}, nil
}
func apolloControlNonce(sequence uint32, origin byte) [12]byte {
var nonce [12]byte
binary.LittleEndian.PutUint32(nonce[:4], sequence)
nonce[10], nonce[11] = origin, 'C'
return nonce
}
+709
View File
@@ -0,0 +1,709 @@
package gateway
import (
"context"
"crypto/rand"
"encoding/binary"
"errors"
"net"
"sync"
"time"
)
const (
apolloENetChannels = 48
apolloENetMaximumPacket = 4096
apolloENetMaximumPayload = 2048
apolloENetMaximumCommands = 32
apolloENetMaximumPending = 128
apolloENetMaximumReorder = 64
// These are the pinned ENet fork defaults. The connection remains subject
// to the stricter ten-second no-receive peer deadline below.
apolloENetTimeoutLimit = 32
apolloENetTimeoutMinimum = 5 * time.Second
apolloENetTimeoutMaximum = 30 * time.Second
apolloENetPeerTimeout = 10 * time.Second
apolloENetPeerIDMask = 0x0fff
apolloENetSentTimeFlag = 0x8000
apolloENetCompressedFlag = 0x4000
apolloENetSessionMask = 0x3000
apolloENetSessionShift = 12
apolloENetCommandMask = 0x0f
apolloENetAcknowledged = 0x80
apolloENetUnsequenced = 0x40
apolloENetConnect = 2
apolloENetVerifyConnect = 3
apolloENetDisconnect = 4
apolloENetPing = 5
apolloENetSendReliable = 6
apolloENetSendUnsequenced = 9
apolloENetBandwidthLimit = 10
apolloENetThrottleConfig = 11
)
var errApolloENet = errors.New("apollo ENet malformed")
type apolloENetState uint8
const (
apolloENetConnecting apolloENetState = iota
apolloENetConnected
apolloENetDisconnecting
apolloENetClosed
)
type apolloENetChannel struct {
nextOutgoing uint16
lastIncoming uint16
hasIncoming bool
incoming map[uint16][]byte
}
type apolloENetPendingKey struct {
channel uint8
sequence uint16
}
type apolloENetPending struct {
packet []byte
firstSent time.Time
sentTime time.Time
timeout time.Duration
attempts uint8
}
// apolloENetPeer is deliberately scoped to the Apollo adapter. It implements
// one negotiated ENet peer over one connected UDP socket and exposes no
// reusable transport abstraction.
type apolloENetPeer struct {
conn *net.UDPConn
now func() time.Time
mu sync.Mutex
state apolloENetState
peerID uint16
inboundSession uint8
outboundSession uint8
connectID uint32
channels [apolloENetChannels]apolloENetChannel
pending map[apolloENetPendingKey]*apolloENetPending
unsequenced uint16
rtt time.Duration
variance time.Duration
reliableSent uint64
retransmits uint64
lastReceive time.Time
lastSend time.Time
lastPing time.Time
disconnectAck chan struct{}
disconnectSeq uint16
onPayload func(uint8, bool, []byte)
onDisconnect func(error)
done chan struct{}
closeOnce sync.Once
}
func newApolloENetPeer(conn *net.UDPConn, now func() time.Time) (*apolloENetPeer, error) {
if conn == nil {
return nil, ErrProviderMalformed
}
if now == nil {
now = time.Now
}
return &apolloENetPeer{
conn: conn, now: now, state: apolloENetConnecting, pending: make(map[apolloENetPendingKey]*apolloENetPending),
rtt: 500 * time.Millisecond, variance: time.Millisecond, done: make(chan struct{}),
}, nil
}
func (p *apolloENetPeer) Connect(ctx context.Context, connectData uint32) error {
if p == nil || connectData == 0 {
return ErrProviderMalformed
}
var id [4]byte
if _, err := rand.Read(id[:]); err != nil {
return err
}
p.mu.Lock()
if p.state != apolloENetConnecting {
p.mu.Unlock()
return ErrProviderMalformed
}
p.connectID = binary.BigEndian.Uint32(id[:])
now := p.now()
packet := apolloENetConnectPacket(now, p.connectID, connectData)
p.pending[apolloENetPendingKey{channel: 0xff, sequence: 1}] = &apolloENetPending{packet: append([]byte(nil), packet...), firstSent: now, sentTime: now, timeout: apolloENetRetransmitTimeout(p.rtt, p.variance, 1), attempts: 1}
err := p.writeLocked(packet)
p.mu.Unlock()
if err != nil {
p.close(err)
return err
}
for {
if err := p.readOnce(ctx); err != nil {
p.close(err)
return err
}
p.mu.Lock()
connected := p.state == apolloENetConnected
p.mu.Unlock()
if connected {
go p.run()
return nil
}
}
}
func apolloENetConnectPacket(now time.Time, connectID, data uint32) []byte {
packet := make([]byte, 52)
apolloENetHeader(packet, apolloENetPeerIDMask, 0, now)
packet[4] = apolloENetConnect | apolloENetAcknowledged
packet[5] = 0xff
binary.BigEndian.PutUint16(packet[6:8], 1)
binary.BigEndian.PutUint16(packet[8:10], 0)
packet[10], packet[11] = 0xff, 0xff
binary.BigEndian.PutUint32(packet[12:16], 1400)
binary.BigEndian.PutUint32(packet[16:20], 32768)
binary.BigEndian.PutUint32(packet[20:24], apolloENetChannels)
binary.BigEndian.PutUint32(packet[24:28], 0)
binary.BigEndian.PutUint32(packet[28:32], 0)
binary.BigEndian.PutUint32(packet[32:36], 5000)
binary.BigEndian.PutUint32(packet[36:40], 2)
binary.BigEndian.PutUint32(packet[40:44], 2)
binary.BigEndian.PutUint32(packet[44:48], connectID)
binary.BigEndian.PutUint32(packet[48:52], data)
return packet
}
func apolloENetHeader(packet []byte, peerID uint16, session uint8, now time.Time) {
value := peerID&apolloENetPeerIDMask | (uint16(session&3) << apolloENetSessionShift) | apolloENetSentTimeFlag
binary.BigEndian.PutUint16(packet[:2], value)
binary.BigEndian.PutUint16(packet[2:4], uint16(now.UnixMilli()))
}
func (p *apolloENetPeer) run() {
ticker := time.NewTicker(25 * time.Millisecond)
defer ticker.Stop()
for {
select {
case <-p.done:
return
case <-ticker.C:
if err := p.maintain(); err != nil {
p.close(err)
return
}
default:
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond)
err := p.readOnce(ctx)
cancel()
if err != nil && !errors.Is(err, context.DeadlineExceeded) && !isApolloENetTimeout(err) {
p.close(err)
return
}
}
}
}
func (p *apolloENetPeer) readOnce(ctx context.Context) error {
if p == nil || p.conn == nil {
return ErrProviderDisconnected
}
deadline := time.Now().Add(100 * time.Millisecond)
if contextDeadline, ok := ctx.Deadline(); ok && contextDeadline.Before(deadline) {
deadline = contextDeadline
}
if err := p.conn.SetReadDeadline(deadline); err != nil {
return err
}
buffer := make([]byte, apolloENetMaximumPacket+1)
count, err := p.conn.Read(buffer)
if err != nil {
if networkErr, ok := err.(net.Error); ok && networkErr.Timeout() {
return context.DeadlineExceeded
}
return err
}
if count < 4 || count > apolloENetMaximumPacket {
return errApolloENet
}
return p.handleDatagram(buffer[:count])
}
func isApolloENetTimeout(err error) bool {
return errors.Is(err, context.DeadlineExceeded)
}
func (p *apolloENetPeer) handleDatagram(packet []byte) error {
if len(packet) < 4 || len(packet) > apolloENetMaximumPacket {
return errApolloENet
}
header := binary.BigEndian.Uint16(packet[:2])
if header&apolloENetCompressedFlag != 0 {
return errApolloENet
}
peerID := header & apolloENetPeerIDMask
session := uint8((header & apolloENetSessionMask) >> apolloENetSessionShift)
offset := 2
sentTime := uint16(0)
if header&apolloENetSentTimeFlag != 0 {
if len(packet) < 4 {
return errApolloENet
}
sentTime = binary.BigEndian.Uint16(packet[2:4])
offset = 4
}
p.mu.Lock()
state := p.state
if state == apolloENetClosed || (state == apolloENetConnected && (peerID != p.peerID || session != p.inboundSession)) {
p.mu.Unlock()
return errApolloENet
}
p.lastReceive = p.now()
p.mu.Unlock()
commands := 0
for offset < len(packet) {
commands++
if commands > apolloENetMaximumCommands || len(packet)-offset < 4 {
return errApolloENet
}
command := packet[offset] & apolloENetCommandMask
flags := packet[offset]
channel := packet[offset+1]
sequence := binary.BigEndian.Uint16(packet[offset+2 : offset+4])
consumed, err := p.handleCommand(command, flags, channel, sequence, sentTime, packet[offset:])
if err != nil || consumed < 4 || consumed > len(packet)-offset {
return errApolloENet
}
offset += consumed
}
return nil
}
func (p *apolloENetPeer) handleCommand(command, flags, channel uint8, sequence, sentTime uint16, data []byte) (int, error) {
switch command {
case 1:
if len(data) < 8 {
return 0, errApolloENet
}
return 8, p.acknowledge(channel, binary.BigEndian.Uint16(data[4:6]), binary.BigEndian.Uint16(data[6:8]))
case apolloENetVerifyConnect:
if len(data) < 44 {
return 0, errApolloENet
}
return 44, p.verifyConnect(sequence, sentTime, data[:44])
case apolloENetDisconnect:
if len(data) < 8 {
return 0, errApolloENet
}
if flags&apolloENetAcknowledged != 0 {
if err := p.sendAcknowledge(channel, sequence, sentTime); err != nil {
return 0, err
}
}
return 8, ErrProviderDisconnected
case apolloENetPing:
if flags&apolloENetAcknowledged != 0 {
if err := p.sendAcknowledge(channel, sequence, sentTime); err != nil {
return 0, err
}
}
return 4, nil
case apolloENetSendReliable:
if len(data) < 6 {
return 0, errApolloENet
}
length := int(binary.BigEndian.Uint16(data[4:6]))
if length > apolloENetMaximumPayload || len(data) < 6+length {
return 0, errApolloENet
}
if flags&apolloENetAcknowledged == 0 || channel >= apolloENetChannels {
return 0, errApolloENet
}
if err := p.sendAcknowledge(channel, sequence, sentTime); err != nil {
return 0, err
}
deliver, err := p.acceptReliable(channel, sequence, data[6:6+length])
if err != nil {
return 0, err
}
for _, message := range deliver {
p.deliver(channel, true, message)
}
return 6 + length, nil
case apolloENetSendUnsequenced:
if len(data) < 8 {
return 0, errApolloENet
}
length := int(binary.BigEndian.Uint16(data[6:8]))
if flags&apolloENetUnsequenced == 0 || channel >= apolloENetChannels || length > apolloENetMaximumPayload || len(data) < 8+length {
return 0, errApolloENet
}
p.deliver(channel, false, data[8:8+length])
return 8 + length, nil
case apolloENetBandwidthLimit:
if len(data) < 12 {
return 0, errApolloENet
}
return 12, nil
case apolloENetThrottleConfig:
if len(data) < 16 {
return 0, errApolloENet
}
return 16, nil
case apolloENetConnect, 7, 8, 12:
return 0, errApolloENet
default:
return 0, errApolloENet
}
}
func (p *apolloENetPeer) verifyConnect(sequence, sentTime uint16, data []byte) error {
p.mu.Lock()
defer p.mu.Unlock()
if p.state != apolloENetConnecting || binary.BigEndian.Uint32(data[40:44]) != p.connectID || binary.BigEndian.Uint32(data[16:20]) != apolloENetChannels {
return errApolloENet
}
p.peerID = binary.BigEndian.Uint16(data[4:6])
p.inboundSession = data[6]
p.outboundSession = data[7]
p.state = apolloENetConnected
delete(p.pending, apolloENetPendingKey{channel: 0xff, sequence: 1})
return p.sendAcknowledgeLocked(0xff, sequence, sentTime)
}
func (p *apolloENetPeer) acknowledge(channel uint8, sequence, sentTime uint16) error {
p.mu.Lock()
if p.state == apolloENetDisconnecting && channel == 0xff && sequence == p.disconnectSeq && p.disconnectAck != nil {
close(p.disconnectAck)
p.disconnectAck = nil
p.mu.Unlock()
return nil
}
key := apolloENetPendingKey{channel: channel, sequence: sequence}
pending, ok := p.pending[key]
if !ok {
p.mu.Unlock()
return nil
}
delete(p.pending, key)
measured := p.now().Sub(pending.sentTime)
if measured < 0 || measured > 30*time.Second {
p.mu.Unlock()
return errApolloENet
}
delta := durationAbs(p.rtt - measured)
p.variance += (delta - p.variance) / 4
p.rtt += (measured - p.rtt) / 8
p.mu.Unlock()
_ = sentTime
return nil
}
func durationAbs(value time.Duration) time.Duration {
if value < 0 {
return -value
}
return value
}
func (p *apolloENetPeer) acceptReliable(channel uint8, sequence uint16, payload []byte) ([][]byte, error) {
p.mu.Lock()
defer p.mu.Unlock()
state := &p.channels[channel]
if !state.hasIncoming {
state.hasIncoming = true
state.lastIncoming = sequence
return [][]byte{append([]byte(nil), payload...)}, nil
}
if sequence == state.lastIncoming+1 {
state.lastIncoming = sequence
deliver := [][]byte{append([]byte(nil), payload...)}
for {
next := state.lastIncoming + 1
queued, ok := state.incoming[next]
if !ok {
return deliver, nil
}
delete(state.incoming, next)
state.lastIncoming = next
deliver = append(deliver, queued)
}
}
if apolloENetSequenceGreater(sequence, state.lastIncoming) {
if uint16(sequence-state.lastIncoming) > 1024 || len(state.incoming) >= apolloENetMaximumReorder {
return nil, errApolloENet
}
if state.incoming == nil {
state.incoming = make(map[uint16][]byte)
}
if _, duplicate := state.incoming[sequence]; !duplicate {
state.incoming[sequence] = append([]byte(nil), payload...)
}
}
return nil, nil
}
func apolloENetSequenceGreater(first, second uint16) bool {
return (first > second && first-second <= 32768) || (first < second && second-first > 32768)
}
func (p *apolloENetPeer) deliver(channel uint8, reliable bool, payload []byte) {
p.mu.Lock()
callback := p.onPayload
p.mu.Unlock()
if callback != nil {
callback(channel, reliable, append([]byte(nil), payload...))
}
}
func (p *apolloENetPeer) SendReliable(channel uint8, payload []byte) error {
if channel >= apolloENetChannels || len(payload) == 0 || len(payload) > apolloENetMaximumPayload {
return ErrProviderMalformed
}
p.mu.Lock()
defer p.mu.Unlock()
if p.state != apolloENetConnected || len(p.pending) >= apolloENetMaximumPending {
return ErrProviderDisconnected
}
state := &p.channels[channel]
state.nextOutgoing++
if state.nextOutgoing == 0 {
state.nextOutgoing++
}
now := p.now()
packet := make([]byte, 10+len(payload))
apolloENetHeader(packet, p.peerID, p.outboundSession, now)
packet[4] = apolloENetSendReliable | apolloENetAcknowledged
packet[5] = channel
binary.BigEndian.PutUint16(packet[6:8], state.nextOutgoing)
binary.BigEndian.PutUint16(packet[8:10], uint16(len(payload)))
copy(packet[10:], payload)
key := apolloENetPendingKey{channel: channel, sequence: state.nextOutgoing}
p.pending[key] = &apolloENetPending{packet: append([]byte(nil), packet...), firstSent: now, sentTime: now, timeout: apolloENetRetransmitTimeout(p.rtt, p.variance, 1), attempts: 1}
if err := p.writeLocked(packet); err != nil {
delete(p.pending, key)
return err
}
p.reliableSent++
return nil
}
func (p *apolloENetPeer) SendUnsequenced(channel uint8, payload []byte) error {
if channel >= apolloENetChannels || len(payload) == 0 || len(payload) > apolloENetMaximumPayload {
return ErrProviderMalformed
}
p.mu.Lock()
defer p.mu.Unlock()
if p.state != apolloENetConnected {
return ErrProviderDisconnected
}
p.unsequenced++
packet := make([]byte, 12+len(payload))
apolloENetHeader(packet, p.peerID, p.outboundSession, p.now())
packet[4] = apolloENetSendUnsequenced | apolloENetUnsequenced
packet[5] = channel
binary.BigEndian.PutUint16(packet[6:8], 0)
binary.BigEndian.PutUint16(packet[8:10], p.unsequenced)
binary.BigEndian.PutUint16(packet[10:12], uint16(len(payload)))
copy(packet[12:], payload)
return p.writeLocked(packet)
}
func (p *apolloENetPeer) sendAcknowledge(channel uint8, sequence, sentTime uint16) error {
p.mu.Lock()
defer p.mu.Unlock()
return p.sendAcknowledgeLocked(channel, sequence, sentTime)
}
func (p *apolloENetPeer) sendAcknowledgeLocked(channel uint8, sequence, sentTime uint16) error {
if p.state == apolloENetClosed {
return ErrProviderDisconnected
}
packet := make([]byte, 12)
peerID := p.peerID
session := p.outboundSession
if p.state == apolloENetConnecting {
peerID, session = apolloENetPeerIDMask, 0
}
apolloENetHeader(packet, peerID, session, p.now())
packet[4] = 1
packet[5] = channel
binary.BigEndian.PutUint16(packet[6:8], 0)
binary.BigEndian.PutUint16(packet[8:10], sequence)
binary.BigEndian.PutUint16(packet[10:12], sentTime)
return p.writeLocked(packet)
}
func (p *apolloENetPeer) maintain() error {
p.mu.Lock()
defer p.mu.Unlock()
if p.state != apolloENetConnected && p.state != apolloENetDisconnecting {
return nil
}
now := p.now()
if !p.lastReceive.IsZero() && now.Sub(p.lastReceive) > apolloENetPeerTimeout {
return ErrProviderTimeout
}
for key, pending := range p.pending {
if now.Sub(pending.sentTime) < pending.timeout {
continue
}
if now.Sub(pending.firstSent) >= apolloENetTimeoutMaximum || (apolloENetExceededTimeoutLimit(pending.attempts) && now.Sub(pending.firstSent) >= apolloENetTimeoutMinimum) {
return ErrProviderTimeout
}
pending.timeout = apolloENetRetransmitTimeout(p.rtt, p.variance, pending.attempts)
pending.attempts++
p.retransmits++
pending.sentTime = now
binary.BigEndian.PutUint16(pending.packet[2:4], uint16(now.UnixMilli()))
if err := p.writeLocked(pending.packet); err != nil {
return err
}
p.pending[key] = pending
}
if now.Sub(p.lastPing) >= 500*time.Millisecond {
p.lastPing = now
return p.sendPingLocked()
}
return nil
}
func apolloENetRetransmitTimeout(rtt, variance time.Duration, attempts uint8) time.Duration {
if rtt < time.Millisecond {
rtt = time.Millisecond
}
if variance < time.Millisecond {
variance = time.Millisecond
}
base := rtt + minDuration(rtt, 4*variance)
if base > apolloENetTimeoutMaximum/5 {
base = apolloENetTimeoutMaximum / 5
}
if attempts == 0 {
attempts = 1
}
if attempts > apolloENetTimeoutLimit {
attempts = apolloENetTimeoutLimit
}
return base * time.Duration(attempts)
}
func apolloENetExceededTimeoutLimit(attempts uint8) bool {
return attempts >= 6 // 1 << (attempts - 1) reaches the fork's limit of 32.
}
func minDuration(first, second time.Duration) time.Duration {
if first < second {
return first
}
return second
}
func (p *apolloENetPeer) sendPingLocked() error {
if len(p.pending) >= apolloENetMaximumPending {
return ErrProviderTimeout
}
state := &p.channels[0]
state.nextOutgoing++
if state.nextOutgoing == 0 {
state.nextOutgoing++
}
now := p.now()
packet := make([]byte, 8)
apolloENetHeader(packet, p.peerID, p.outboundSession, now)
packet[4] = apolloENetPing | apolloENetAcknowledged
packet[5] = 0
binary.BigEndian.PutUint16(packet[6:8], state.nextOutgoing)
p.pending[apolloENetPendingKey{channel: 0, sequence: state.nextOutgoing}] = &apolloENetPending{packet: append([]byte(nil), packet...), firstSent: now, sentTime: now, timeout: apolloENetRetransmitTimeout(p.rtt, p.variance, 1), attempts: 1}
if err := p.writeLocked(packet); err != nil {
delete(p.pending, apolloENetPendingKey{channel: 0, sequence: state.nextOutgoing})
return err
}
p.reliableSent++
return nil
}
func (p *apolloENetPeer) telemetry() ProviderTelemetry {
if p == nil {
return ProviderTelemetry{}
}
p.mu.Lock()
defer p.mu.Unlock()
return ProviderTelemetry{
ControlRTT: p.rtt, ControlJitter: p.variance, ReliableSent: p.reliableSent,
ReliableRetransmits: p.retransmits, PendingReliable: uint64(len(p.pending)),
}
}
func (p *apolloENetPeer) writeLocked(packet []byte) error {
if len(packet) < 4 || len(packet) > apolloENetMaximumPacket {
return ErrProviderMalformed
}
count, err := p.conn.Write(packet)
if err != nil || count != len(packet) {
if err != nil {
return err
}
return ErrProviderDisconnected
}
p.lastSend = p.now()
return nil
}
func (p *apolloENetPeer) Disconnect(ctx context.Context) error {
if p == nil {
return nil
}
p.mu.Lock()
if p.state == apolloENetClosed {
p.mu.Unlock()
return nil
}
p.state = apolloENetDisconnecting
ack := make(chan struct{})
p.disconnectAck, p.disconnectSeq = ack, 1
packet := make([]byte, 12)
apolloENetHeader(packet, p.peerID, p.outboundSession, p.now())
packet[4] = apolloENetDisconnect | apolloENetAcknowledged
packet[5] = 0xff
binary.BigEndian.PutUint16(packet[6:8], 1)
if err := p.writeLocked(packet); err != nil {
p.mu.Unlock()
p.close(err)
return err
}
p.mu.Unlock()
deadline := time.NewTimer(2 * time.Second)
defer deadline.Stop()
select {
case <-ack:
p.close(nil)
return nil
case <-p.done:
return nil
case <-ctx.Done():
p.close(ctx.Err())
return ctx.Err()
case <-deadline.C:
p.close(ErrProviderTimeout)
return ErrProviderTimeout
}
}
func (p *apolloENetPeer) close(err error) {
if p == nil {
return
}
p.closeOnce.Do(func() {
p.mu.Lock()
p.state = apolloENetClosed
callback := p.onDisconnect
p.mu.Unlock()
close(p.done)
_ = p.conn.Close()
if callback != nil && err != nil {
callback(err)
}
})
}
+348
View File
@@ -0,0 +1,348 @@
package gateway
import (
"context"
"encoding/binary"
"encoding/hex"
"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) {
packet, err := encodeApolloInputEvent(InputEvent{Device: "keyboard", Code: 30, Pressed: true, Payload: []byte{2}})
if err != nil {
t.Fatal(err)
}
const expected = "0000000a03000000001e00020000"
if string(packet.payload) != string(mustDecodeHex(t, expected)) {
t.Fatalf("keyboard packet = %x, want %s", packet.payload, expected)
}
}
func mustDecodeHex(t *testing.T, value string) []byte {
t.Helper()
decoded, err := hex.DecodeString(value)
if err != nil {
t.Fatal(err)
}
return decoded
}
+90
View File
@@ -0,0 +1,90 @@
package gateway
import (
"encoding/binary"
"unicode/utf8"
)
const (
apolloChannelGeneric = 0
apolloChannelUrgent = 1
apolloChannelKeyboard = 2
apolloChannelMouse = 3
apolloChannelUTF8 = 6
apolloChannelGamepad = 16
)
type apolloInputPacket struct {
channel uint8
payload []byte
}
func encodeApolloInputEvent(event InputEvent) (apolloInputPacket, error) {
switch event.Device {
case "keyboard":
if event.Code < 0 || event.Code > 0xffff || len(event.Payload) > 1 {
return apolloInputPacket{}, ErrInputMalformed
}
packet := make([]byte, 14)
binary.BigEndian.PutUint32(packet[:4], 10)
if event.Pressed {
binary.LittleEndian.PutUint32(packet[4:8], 3)
} else {
binary.LittleEndian.PutUint32(packet[4:8], 4)
}
binary.LittleEndian.PutUint16(packet[9:11], uint16(event.Code))
if len(event.Payload) == 1 {
packet[11] = event.Payload[0]
}
return apolloInputPacket{channel: apolloChannelKeyboard, payload: packet}, nil
case "mouse-button":
if event.Code < 1 || event.Code > 8 || len(event.Payload) != 0 {
return apolloInputPacket{}, ErrInputMalformed
}
packet := make([]byte, 9)
binary.BigEndian.PutUint32(packet[:4], 5)
magic := uint32(8)
if !event.Pressed {
magic = 9
}
binary.LittleEndian.PutUint32(packet[4:8], magic)
packet[8] = byte(event.Code)
return apolloInputPacket{channel: apolloChannelMouse, payload: packet}, nil
case "mouse-relative":
if event.Pressed || len(event.Payload) != 4 {
return apolloInputPacket{}, ErrInputMalformed
}
packet := make([]byte, 12)
binary.BigEndian.PutUint32(packet[:4], 8)
binary.LittleEndian.PutUint32(packet[4:8], 7)
copy(packet[8:], event.Payload)
return apolloInputPacket{channel: apolloChannelMouse, payload: packet}, nil
case "utf8":
if event.Pressed || len(event.Payload) == 0 || len(event.Payload) > utf8.UTFMax || !utf8.Valid(event.Payload) || utf8.RuneCount(event.Payload) != 1 {
return apolloInputPacket{}, ErrInputMalformed
}
packet := make([]byte, 8+len(event.Payload))
binary.BigEndian.PutUint32(packet[:4], uint32(4+len(event.Payload)))
binary.LittleEndian.PutUint32(packet[4:8], 0x17)
copy(packet[8:], event.Payload)
return apolloInputPacket{channel: apolloChannelUTF8, payload: packet}, nil
case "controller":
if event.Code < 0 || event.Code > 15 || len(event.Payload) != 16 {
return apolloInputPacket{}, ErrInputMalformed
}
packet := make([]byte, 34)
binary.BigEndian.PutUint32(packet[:4], 30)
binary.LittleEndian.PutUint32(packet[4:8], 0x0c)
binary.LittleEndian.PutUint16(packet[8:10], 0x1a)
binary.LittleEndian.PutUint16(packet[10:12], uint16(event.Code))
copy(packet[12:14], event.Payload[:2])
binary.LittleEndian.PutUint16(packet[14:16], 0x14)
copy(packet[16:28], event.Payload[2:14])
binary.LittleEndian.PutUint16(packet[28:30], 0x9c)
copy(packet[30:32], event.Payload[14:16])
binary.LittleEndian.PutUint16(packet[32:34], 0x55)
return apolloInputPacket{channel: apolloChannelGamepad + uint8(event.Code), payload: packet}, nil
default:
return apolloInputPacket{}, ErrInputMalformed
}
}
+181
View File
@@ -0,0 +1,181 @@
package gateway
import (
"crypto/aes"
"crypto/cipher"
"encoding/binary"
"errors"
)
const (
// Apollo limits clear audio payloads to 1400 bytes and uses AES-CBC with
// PKCS#7 padding. FEC adds its fixed 12-byte header outside that ciphertext.
// This is the largest accepted UDP datagram, not an allocation hint.
apolloMediaMaximumPacket = 12 + 12 + 1408
apolloVideoHeaderSize = 32
apolloRTPHeaderSize = 12
apolloVideoNVHeaderSize = 16
apolloVideoRawPacketSize = 1024 + 16
)
var (
errApolloMedia = errors.New("apollo media malformed")
errApolloMediaParity = errors.New("apollo media parity packet")
)
type apolloMediaCodec struct {
block cipher.Block
aead cipher.AEAD
keyID uint32
}
type apolloRTPPacket struct {
extension bool
payloadType byte
sequence uint16
payload []byte
}
type apolloAudioShard struct {
sequence uint16
timestamp uint32
ssrc uint32
base uint16
parityIndex uint8
parity bool
payload []byte
}
func newApolloMediaCodec(key []byte, keyID uint32) (*apolloMediaCodec, error) {
block, err := aes.NewCipher(key)
if err != nil {
return nil, err
}
aead, err := cipher.NewGCM(block)
if err != nil {
return nil, err
}
return &apolloMediaCodec{block: block, aead: aead, keyID: keyID}, nil
}
func apolloMediaPing(payload []byte, sequence uint32) []byte {
if len(payload) != 16 {
return nil
}
ping := make([]byte, 20)
copy(ping, payload)
binary.BigEndian.PutUint32(ping[16:], sequence)
return ping
}
func (c *apolloMediaCodec) OpenVideo(packet []byte) (apolloVideoShard, error) {
if c == nil || c.aead == nil || len(packet) != apolloVideoHeaderSize+apolloVideoRawPacketSize {
return apolloVideoShard{}, errApolloMedia
}
sealed := make([]byte, len(packet)-apolloVideoHeaderSize+apolloControlTagSize)
copy(sealed, packet[apolloVideoHeaderSize:])
copy(sealed[len(packet)-apolloVideoHeaderSize:], packet[16:apolloVideoHeaderSize])
plaintext, err := c.aead.Open(nil, packet[:12], sealed, nil)
if err != nil {
return apolloVideoShard{}, errApolloMedia
}
rtp, err := parseApolloRTP(plaintext)
if err != nil || !rtp.extension || rtp.payloadType != 0 || len(rtp.payload) != 1024 {
return apolloVideoShard{}, errApolloMedia
}
nv := rtp.payload[:apolloVideoNVHeaderSize]
flags := nv[8]
fecInfo := binary.LittleEndian.Uint32(nv[12:16])
dataPackets := int(fecInfo >> 22)
fecIndex := int((fecInfo >> 12) & 0x03ff)
fecPercent := int((fecInfo >> 4) & 0xff)
if dataPackets < 1 || dataPackets > apolloVideoMaximumDataShards {
return apolloVideoShard{}, errApolloMedia
}
parityPackets := (dataPackets*fecPercent + 99) / 100
if dataPackets+parityPackets > 255 || fecIndex >= dataPackets+parityPackets {
return apolloVideoShard{}, errApolloMedia
}
block := (nv[11] >> 4) & 0x03
lastBlock := (nv[11] >> 6) & 0x03
if block > lastBlock {
return apolloVideoShard{}, errApolloMedia
}
return apolloVideoShard{
frame: binary.LittleEndian.Uint32(nv[4:8]),
block: block,
lastBlock: lastBlock,
dataPackets: dataPackets,
parity: parityPackets,
index: fecIndex,
sequence: rtp.sequence,
streamIndex: binary.LittleEndian.Uint32(nv[:4]) >> 8,
flags: flags,
payload: append([]byte(nil), rtp.payload[apolloVideoNVHeaderSize:]...),
}, nil
}
func (c *apolloMediaCodec) OpenAudio(packet []byte) (apolloAudioShard, error) {
if c == nil || c.block == nil || len(packet) <= apolloRTPHeaderSize || len(packet) > apolloMediaMaximumPacket {
return apolloAudioShard{}, errApolloMedia
}
rtp, err := parseApolloRTP(packet)
if err != nil || rtp.extension || len(rtp.payload) == 0 {
return apolloAudioShard{}, errApolloMedia
}
if rtp.payloadType == 97 && len(rtp.payload)%aes.BlockSize == 0 {
return apolloAudioShard{
sequence: rtp.sequence, timestamp: binary.BigEndian.Uint32(packet[4:8]), ssrc: binary.BigEndian.Uint32(packet[8:12]),
base: rtp.sequence &^ 3, payload: append([]byte(nil), rtp.payload...),
}, nil
}
if rtp.payloadType != 127 || len(rtp.payload) <= 12 || len(rtp.payload)-12 > 1408 || rtp.payload[0] > 1 || rtp.payload[1] != 97 {
return apolloAudioShard{}, errApolloMedia
}
base := binary.BigEndian.Uint16(rtp.payload[2:4])
if base&3 != 0 || len(rtp.payload[12:])%aes.BlockSize != 0 {
return apolloAudioShard{}, errApolloMedia
}
return apolloAudioShard{
sequence: rtp.sequence, timestamp: binary.BigEndian.Uint32(rtp.payload[4:8]), ssrc: binary.BigEndian.Uint32(rtp.payload[8:12]),
base: base, parityIndex: rtp.payload[0], parity: true, payload: append([]byte(nil), rtp.payload[12:]...),
}, nil
}
func (c *apolloMediaCodec) openApolloAudioCipher(sequence uint16, payload []byte) ([]byte, error) {
if c == nil || c.block == nil || len(payload) == 0 || len(payload) > 1408 || len(payload)%aes.BlockSize != 0 {
return nil, errApolloMedia
}
plaintext := append([]byte(nil), payload...)
iv := make([]byte, aes.BlockSize)
binary.BigEndian.PutUint32(iv, c.keyID+uint32(sequence))
cipher.NewCBCDecrypter(c.block, iv).CryptBlocks(plaintext, plaintext)
padding := int(plaintext[len(plaintext)-1])
if padding == 0 || padding > aes.BlockSize || padding > len(plaintext) {
return nil, errApolloMedia
}
for _, value := range plaintext[len(plaintext)-padding:] {
if int(value) != padding {
return nil, errApolloMedia
}
}
return append([]byte(nil), plaintext[:len(plaintext)-padding]...), nil
}
func parseApolloRTP(packet []byte) (apolloRTPPacket, error) {
if len(packet) < apolloRTPHeaderSize || packet[0]>>6 != 2 || packet[0]&0x2f != 0 {
return apolloRTPPacket{}, errApolloMedia
}
offset := apolloRTPHeaderSize
extension := packet[0]&0x10 != 0
if extension {
if len(packet) < offset+4 || binary.BigEndian.Uint16(packet[14:16]) != 0 {
return apolloRTPPacket{}, errApolloMedia
}
offset += 4
}
if offset >= len(packet) {
return apolloRTPPacket{}, errApolloMedia
}
return apolloRTPPacket{extension: extension, payloadType: packet[1] & 0x7f, sequence: binary.BigEndian.Uint16(packet[2:4]), payload: packet[offset:]}, nil
}
+524 -107
View File
@@ -9,15 +9,20 @@ import (
"crypto/x509" "crypto/x509"
"encoding/binary" "encoding/binary"
"encoding/hex" "encoding/hex"
"encoding/xml"
"errors"
"fmt" "fmt"
"io" "io"
"net" "net"
"net/http" "net/http"
"net/url" "net/url"
"sort"
"strconv" "strconv"
"strings" "strings"
"sync" "sync"
"sync/atomic"
"time" "time"
"unicode/utf8"
protocol "git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol/gen/go/protocol" protocol "git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol/gen/go/protocol"
) )
@@ -29,11 +34,11 @@ type NativeApolloBackend struct {
Dialer *net.Dialer Dialer *net.Dialer
mu sync.Mutex mu sync.Mutex
pending map[string]net.Conn pending map[string]*apolloRTSPSetup
} }
func NewNativeApolloBackend() *NativeApolloBackend { func NewNativeApolloBackend() *NativeApolloBackend {
return &NativeApolloBackend{Dialer: &net.Dialer{Timeout: 5 * time.Second}, pending: make(map[string]net.Conn)} return &NativeApolloBackend{Dialer: &net.Dialer{Timeout: 5 * time.Second}, pending: make(map[string]*apolloRTSPSetup)}
} }
func (b *NativeApolloBackend) Management(ctx context.Context, request LaunchRequest) ([]byte, error) { func (b *NativeApolloBackend) Management(ctx context.Context, request LaunchRequest) ([]byte, error) {
@@ -53,7 +58,11 @@ func newPinnedApolloHTTPClient(work protocol.ProviderSessionWork) (*http.Client,
if err != nil { if err != nil {
return nil, err return nil, err
} }
return &http.Client{Transport: &http.Transport{TLSClientConfig: tlsConfig}, Timeout: 5 * time.Second}, nil return &http.Client{
Transport: &http.Transport{TLSClientConfig: tlsConfig},
Timeout: 5 * time.Second,
CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse },
}, nil
} }
func apolloGet(ctx context.Context, client *http.Client, work protocol.ProviderSessionWork, path string, values url.Values) ([]byte, error) { func apolloGet(ctx context.Context, client *http.Client, work protocol.ProviderSessionWork, path string, values url.Values) ([]byte, error) {
@@ -74,6 +83,88 @@ func apolloGet(ctx context.Context, client *http.Client, work protocol.ProviderS
return readBounded(response.Body, 64*1024) return readBounded(response.Body, 64*1024)
} }
func apolloSessionRequest(work protocol.ProviderSessionWork, key []byte, keyID uint32) (string, url.Values) {
values := url.Values{
"rikey": {hex.EncodeToString(key)},
"rikeyid": {strconv.FormatUint(uint64(keyID), 10)},
"localAudioPlayMode": {"0"},
}
if work.ReconnectSequence > 0 {
return "/resume", values
}
values.Set("uniqueid", work.ClientID)
values.Set("appid", work.ApplicationID)
values.Set("corever", "1")
return "/launch", values
}
func apolloClipboardRequest(ctx context.Context, client *http.Client, host string, port int64, method, text string) ([]byte, error) {
if client == nil || host == "" || port < 1 || port > maxApolloRTSPPort || (method != http.MethodGet && method != http.MethodPost) {
return nil, ErrProviderMalformed
}
endpoint := url.URL{Scheme: "https", Host: net.JoinHostPort(host, strconv.FormatInt(port, 10)), Path: "/actions/clipboard", RawQuery: url.Values{"type": {"text"}}.Encode()}
var body io.Reader
if method == http.MethodPost {
if !utf8.ValidString(text) || len(text) > 65536 {
return nil, ErrProviderMalformed
}
body = bytes.NewReader([]byte(text))
}
request, err := http.NewRequestWithContext(ctx, method, endpoint.String(), body)
if err != nil {
return nil, ErrProviderMalformed
}
if method == http.MethodPost {
request.Header.Set("Content-Type", "text/plain; charset=utf-8")
}
response, err := client.Do(request)
if err != nil {
return nil, err
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
return nil, fmt.Errorf("provider clipboard status %d", response.StatusCode)
}
return readBounded(response.Body, 65536)
}
func apolloCancelRequest(ctx context.Context, client *http.Client, host string, port int64) error {
if client == nil || host == "" || port < 1 || port > maxApolloRTSPPort {
return ErrProviderMalformed
}
endpoint := url.URL{Scheme: "https", Host: net.JoinHostPort(host, strconv.FormatInt(port, 10)), Path: "/cancel"}
request, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint.String(), nil)
if err != nil {
return ErrProviderMalformed
}
response, err := client.Do(request)
if err != nil {
return err
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
return fmt.Errorf("provider cancel status %d", response.StatusCode)
}
body, err := readBounded(response.Body, 64*1024)
if err != nil {
return err
}
var result struct {
XMLName xml.Name `xml:"root"`
StatusCode int `xml:"status_code,attr"`
Cancel int `xml:"cancel"`
}
decoder := xml.NewDecoder(bytes.NewReader(body))
decoder.Strict = true
if err := decoder.Decode(&result); err != nil || result.XMLName.Local != "root" || result.StatusCode != http.StatusOK || result.Cancel != 1 {
return ErrProviderMalformed
}
if err := decoder.Decode(&struct{}{}); !errors.Is(err, io.EOF) {
return ErrProviderMalformed
}
return nil
}
func pinnedApolloTLSConfig(work protocol.ProviderSessionWork) (*tls.Config, error) { func pinnedApolloTLSConfig(work protocol.ProviderSessionWork) (*tls.Config, error) {
identity, ok := providerIdentityFromKey(work.ProviderIdentity) identity, ok := providerIdentityFromKey(work.ProviderIdentity)
if !ok || !strings.HasPrefix(identity.Fingerprint, "sha256:") { if !ok || !strings.HasPrefix(identity.Fingerprint, "sha256:") {
@@ -123,15 +214,13 @@ func (b *NativeApolloBackend) Setup(ctx context.Context, request LaunchRequest)
if _, err := rand.Read(key); err != nil { if _, err := rand.Read(key); err != nil {
return nil, err return nil, err
} }
defer zeroApolloSecret(key)
var keyID [4]byte var keyID [4]byte
if _, err := rand.Read(keyID[:]); err != nil { if _, err := rand.Read(keyID[:]); err != nil {
return nil, err return nil, err
} }
launch, err := apolloGet(ctx, client, work, "/launch", url.Values{ sessionPath, sessionValues := apolloSessionRequest(work, key, binary.BigEndian.Uint32(keyID[:]))
"uniqueid": {work.ClientID}, "appid": {work.ApplicationID}, "rikey": {hex.EncodeToString(key)}, launch, err := apolloGet(ctx, client, work, sessionPath, sessionValues)
"rikeyid": {strconv.FormatUint(uint64(binary.BigEndian.Uint32(keyID[:])), 10)}, "localAudioPlayMode": {"0"},
"corever": {"1"},
})
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -147,50 +236,25 @@ func (b *NativeApolloBackend) Setup(ctx context.Context, request LaunchRequest)
if err != nil || streamPort != work.StreamPort { if err != nil || streamPort != work.StreamPort {
return nil, ErrProviderMalformed return nil, ErrProviderMalformed
} }
conn, err := b.Dialer.DialContext(ctx, "tcp", net.JoinHostPort(work.StreamHost, strconv.FormatInt(work.StreamPort, 10))) setup, response, err := b.performRTSPHandshake(ctx, work, key, binary.BigEndian.Uint32(keyID[:]), streamURL)
if err != nil { if err != nil {
return nil, err return nil, err
} }
if deadline, ok := ctx.Deadline(); ok {
_ = conn.SetDeadline(deadline)
}
codec, err := newEncryptedRTSPCodec(key)
if err != nil {
_ = conn.Close()
return nil, err
}
requestText := "SETUP rtsp://" + work.StreamHost + "/streamid=video/0/0 RTSP/1.0\r\nCSeq: 1\r\nTransport: RTP/AVP/TCP;interleaved=0-1\r\n\r\n"
encoded, err := codec.SealClient([]byte(requestText))
if err != nil {
_ = conn.Close()
return nil, err
}
if _, err := conn.Write(encoded); err != nil {
_ = conn.Close()
return nil, err
}
response, err := readEncryptedRTSPHeaders(conn, codec)
if err != nil {
_ = conn.Close()
return nil, err
}
b.mu.Lock() b.mu.Lock()
b.pending[request.SessionID] = conn b.pending[request.SessionID] = setup
b.mu.Unlock() b.mu.Unlock()
return response, nil return response, nil
} }
func (b *NativeApolloBackend) Open(_ context.Context, request LaunchRequest, _ RTSPResponse) (ProviderSession, error) { func (b *NativeApolloBackend) Open(ctx context.Context, request LaunchRequest, _ RTSPResponse) (ProviderSession, error) {
b.mu.Lock() b.mu.Lock()
conn, ok := b.pending[request.SessionID] setup, ok := b.pending[request.SessionID]
delete(b.pending, request.SessionID) delete(b.pending, request.SessionID)
b.mu.Unlock() b.mu.Unlock()
if !ok || conn == nil { if !ok || setup == nil {
return nil, ErrProviderDisconnected return nil, ErrProviderDisconnected
} }
session := newNativeApolloSession(conn, request.SessionID) return newNativeApolloProviderSession(ctx, setup)
go session.readMedia()
return session, nil
} }
func readBounded(reader io.Reader, max int) ([]byte, error) { func readBounded(reader io.Reader, max int) ([]byte, error) {
@@ -204,44 +268,134 @@ func readBounded(reader io.Reader, max int) ([]byte, error) {
return data, nil return data, nil
} }
func readEncryptedRTSPHeaders(conn net.Conn, codec *encryptedRTSPCodec) ([]byte, error) {
header := make([]byte, encryptedRTSPHeaderSize)
if _, err := io.ReadFull(conn, header); err != nil {
return nil, err
}
length := binary.BigEndian.Uint32(header[:4]) & 0x7fffffff
if length == 0 || length > encryptedRTSPMaxPayload {
return nil, ErrProviderMalformed
}
frame := make([]byte, encryptedRTSPHeaderSize+int(length))
copy(frame, header)
if _, err := io.ReadFull(conn, frame[encryptedRTSPHeaderSize:]); err != nil {
return nil, err
}
plaintext, err := codec.OpenHost(frame)
if err != nil || len(plaintext) > 16*1024 || !strings.HasSuffix(string(plaintext), "\r\n\r\n") {
return nil, ErrProviderMalformed
}
return plaintext, nil
}
type nativeApolloSession struct { type nativeApolloSession struct {
conn net.Conn enet *apolloENetPeer
control *apolloControlCodec
audioConn *net.UDPConn
videoConn *net.UDPConn
media *apolloMediaCodec
videoFEC apolloVideoAssembler
audioFEC apolloAudioAssembler
audioPing []byte
videoPing []byte
sessionID string sessionID string
video chan []byte video chan []byte
audio chan []byte audio chan []byte
events chan ProviderEvent
mu sync.Mutex mu sync.Mutex
controlMu sync.Mutex
state protocol.ProviderState state protocol.ProviderState
pressed map[string]InputEvent
closeOnce sync.Once closeOnce sync.Once
channelsOnce sync.Once
done chan struct{} done chan struct{}
readDone chan struct{} readDone chan struct{}
managementClient *http.Client
managementHost string
managementPort int64
allowApplicationTermination bool
terminationErr error
mediaDrops atomic.Uint64
} }
func newNativeApolloSession(conn net.Conn, sessionID string) *nativeApolloSession { func newNativeApolloSession(sessionID string) *nativeApolloSession {
return &nativeApolloSession{conn: conn, sessionID: sessionID, video: make(chan []byte, 16), audio: make(chan []byte, 16), state: protocol.ProviderState{Version: "1", SessionID: sessionID, State: ProviderStateStarting, Channels: []string{"video", "audio", "input", "feedback"}}, done: make(chan struct{}), readDone: make(chan struct{})} return &nativeApolloSession{sessionID: sessionID, video: make(chan []byte, 16), audio: make(chan []byte, 16), events: make(chan ProviderEvent, 16), state: protocol.ProviderState{Version: "1", SessionID: sessionID, State: ProviderStateStarting, Channels: []string{"video", "audio", "input", "feedback"}}, pressed: make(map[string]InputEvent), done: make(chan struct{}), readDone: make(chan struct{})}
}
func newNativeApolloProviderSession(ctx context.Context, setup *apolloRTSPSetup) (*nativeApolloSession, error) {
if setup == nil || len(setup.streamKey) != 16 || setup.controlPort == 0 || setup.audioPort == 0 || setup.videoPort == 0 {
return nil, ErrProviderMalformed
}
defer zeroApolloSecret(setup.streamKey)
controlConn, err := dialApolloUDP(setup.streamHost, setup.controlPort)
if err != nil {
return nil, err
}
peer, err := newApolloENetPeer(controlConn, time.Now)
if err != nil {
_ = controlConn.Close()
return nil, err
}
codec, err := newApolloControlCodec(setup.streamKey)
if err != nil {
peer.close(err)
return nil, err
}
media, err := newApolloMediaCodec(setup.streamKey, setup.streamKeyID)
if err != nil {
peer.close(err)
return nil, err
}
session := newNativeApolloSession(setup.sessionID)
managementClient, err := newPinnedApolloHTTPClient(setup.providerWork)
if err != nil {
peer.close(err)
return nil, err
}
session.enet, session.control, session.media = peer, codec, media
session.managementClient = managementClient
session.managementHost, session.managementPort = setup.providerWork.ManagementHost, setup.providerWork.ManagementPort
session.allowApplicationTermination = setup.providerWork.ProviderApplicationTerminationAllowed
session.audioPing = append([]byte(nil), setup.audioPing...)
session.videoPing = append([]byte(nil), setup.videoPing...)
peer.onPayload = session.handleApolloControlPayload
peer.onDisconnect = session.handleApolloDisconnect
connectCtx, cancel := context.WithTimeout(ctx, 10*time.Second)
defer cancel()
if err := peer.Connect(connectCtx, setup.controlConnect); err != nil {
peer.close(err)
return nil, err
}
audioConn, err := dialApolloUDP(setup.streamHost, setup.audioPort)
if err != nil {
peer.close(err)
return nil, err
}
videoConn, err := dialApolloUDP(setup.streamHost, setup.videoPort)
if err != nil {
_ = audioConn.Close()
peer.close(err)
return nil, err
}
if _, err := audioConn.Write(apolloMediaPing(setup.audioPing, 1)); err != nil {
_ = audioConn.Close()
_ = videoConn.Close()
peer.close(err)
return nil, err
}
if _, err := videoConn.Write(apolloMediaPing(setup.videoPing, 1)); err != nil {
_ = audioConn.Close()
_ = videoConn.Close()
peer.close(err)
return nil, err
}
session.audioConn, session.videoConn = audioConn, videoConn
go session.readUDPMedia()
go session.periodicApolloMediaPing()
return session, nil
}
func dialApolloUDP(host string, port int) (*net.UDPConn, error) {
if host == "" || port < 1 || port > maxApolloRTSPPort {
return nil, ErrProviderMalformed
}
remote, err := net.ResolveUDPAddr("udp", net.JoinHostPort(host, strconv.Itoa(port)))
if err != nil {
return nil, err
}
return net.DialUDP("udp", nil, remote)
} }
func (s *nativeApolloSession) Ready(context.Context) error { func (s *nativeApolloSession) Ready(context.Context) error {
if s.enet != nil {
if err := s.writeApolloControl(apolloChannelGeneric, true, apolloControlTypeIDR, []byte{0, 0}); err != nil {
return err
}
if err := s.writeApolloControl(apolloChannelGeneric, true, apolloControlTypeStart, []byte{0}); err != nil {
return err
}
go s.periodicApolloPing()
}
s.mu.Lock() s.mu.Lock()
s.state.State = ProviderStateReady s.state.State = ProviderStateReady
s.mu.Unlock() s.mu.Unlock()
@@ -250,106 +404,367 @@ func (s *nativeApolloSession) Ready(context.Context) error {
func (s *nativeApolloSession) Video() <-chan []byte { return s.video } func (s *nativeApolloSession) Video() <-chan []byte { return s.video }
func (s *nativeApolloSession) Audio() <-chan []byte { return s.audio } func (s *nativeApolloSession) Audio() <-chan []byte { return s.audio }
func (s *nativeApolloSession) Events() <-chan ProviderEvent { return s.events }
func (s *nativeApolloSession) Input(ctx context.Context, event InputEvent) error { func (s *nativeApolloSession) Input(ctx context.Context, event InputEvent) error {
payload, err := EncodeInputEvent(event) packet, err := encodeApolloInputEvent(event)
if err != nil { if err != nil {
return err return err
} }
return s.writeControl(ctx, ControlPacket{Kind: 1, Sequence: event.Sequence, Payload: payload}) if err := ctx.Err(); err != nil {
return err
}
if err := s.writeApolloControl(packet.channel, true, apolloControlTypeInput, packet.payload); err != nil {
return err
}
if event.Device == "keyboard" || event.Device == "mouse-button" || event.Device == "controller" {
key := fmt.Sprintf("%s:%d", event.Device, event.Code)
s.mu.Lock()
if event.Pressed {
s.pressed[key] = event
} else {
delete(s.pressed, key)
}
s.mu.Unlock()
}
return nil
} }
func (s *nativeApolloSession) Feedback(ctx context.Context, feedback Feedback) error { func (s *nativeApolloSession) Feedback(ctx context.Context, feedback Feedback) error {
return s.writeControl(ctx, ControlPacket{Kind: 3, Sequence: feedback.Sequence, Payload: feedback.Payload}) if err := ctx.Err(); err != nil {
return err
}
switch feedback.Kind {
case FeedbackIDR:
if len(feedback.Payload) != 0 {
return ErrProviderMalformed
}
return s.writeApolloControl(apolloChannelUrgent, true, apolloControlTypeIDR, []byte{0, 0})
case FeedbackFEC:
if !validGatewayFECStatus(feedback.Payload) {
return ErrProviderMalformed
}
return s.writeApolloControl(apolloChannelGeneric, false, apolloControlTypeFEC, feedback.Payload)
default:
return ErrProviderMalformed
}
} }
func (s *nativeApolloSession) Reconnect(ctx context.Context) error { func (s *nativeApolloSession) ReadClipboard(ctx context.Context) (string, error) {
return s.writeControl(ctx, ControlPacket{Kind: 4, Payload: []byte("RECN")}) response, err := apolloClipboardRequest(ctx, s.managementClient, s.managementHost, s.managementPort, http.MethodGet, "")
if err != nil {
return "", err
}
if !utf8.Valid(response) {
return "", ErrProviderMalformed
}
return string(response), nil
}
func (s *nativeApolloSession) WriteClipboard(ctx context.Context, text string) error {
_, err := apolloClipboardRequest(ctx, s.managementClient, s.managementHost, s.managementPort, http.MethodPost, text)
return err
} }
func (s *nativeApolloSession) ReleaseAll(ctx context.Context) error { func (s *nativeApolloSession) ReleaseAll(ctx context.Context) error {
return s.writeControl(ctx, ControlPacket{Kind: 2, Payload: []byte("RELEASE_ALL")}) s.mu.Lock()
pressed := make([]InputEvent, 0, len(s.pressed))
for _, event := range s.pressed {
pressed = append(pressed, event)
}
s.mu.Unlock()
sort.Slice(pressed, func(first, second int) bool {
if pressed[first].Device != pressed[second].Device {
return pressed[first].Device < pressed[second].Device
}
return pressed[first].Code < pressed[second].Code
})
for _, event := range pressed {
event.Pressed = false
if event.Device == "controller" {
event.Payload = make([]byte, len(event.Payload))
}
if err := s.Input(ctx, event); err != nil {
return err
}
}
return nil
} }
func (s *nativeApolloSession) Terminate(ctx context.Context) error { func (s *nativeApolloSession) Terminate(ctx context.Context) error {
_ = s.writeControl(ctx, ControlPacket{Kind: 5, Payload: []byte("TEAR")}) var cleanupErr error
var timedOut bool
s.closeOnce.Do(func() { s.closeOnce.Do(func() {
if err := s.ReleaseAll(ctx); err != nil {
cleanupErr = err
}
close(s.done) close(s.done)
_ = s.conn.Close() if s.enet != nil {
if err := s.enet.Disconnect(ctx); err != nil && cleanupErr == nil {
cleanupErr = err
}
}
if s.audioConn != nil {
_ = s.audioConn.Close()
}
if s.videoConn != nil {
_ = s.videoConn.Close()
}
select { select {
case <-s.readDone: case <-s.readDone:
case <-ctx.Done(): case <-ctx.Done():
timedOut = true if cleanupErr == nil {
cleanupErr = ctx.Err()
} }
if !timedOut {
close(s.video)
close(s.audio)
} }
if cleanupErr == nil {
s.closeMediaChannels()
if s.allowApplicationTermination {
if err := apolloCancelRequest(ctx, s.managementClient, s.managementHost, s.managementPort); err != nil {
cleanupErr = err
}
}
}
if s.managementClient != nil {
s.managementClient.CloseIdleConnections()
}
s.mu.Lock()
s.terminationErr = cleanupErr
s.mu.Unlock()
}) })
s.mu.Lock() s.mu.Lock()
if timedOut { cleanupErr = s.terminationErr
if cleanupErr != nil {
s.state.State = ProviderStateCleanup s.state.State = ProviderStateCleanup
s.state.CleanupPending = true s.state.CleanupPending = true
} else { } else {
s.state.State = ProviderStateTerminated s.state.State = ProviderStateTerminated
} }
s.mu.Unlock() s.mu.Unlock()
if timedOut { if cleanupErr != nil {
return ctx.Err() return cleanupErr
} }
s.control, s.media = nil, nil
return nil return nil
} }
func zeroApolloSecret(secret []byte) {
for index := range secret {
secret[index] = 0
}
}
func (s *nativeApolloSession) State() protocol.ProviderState { func (s *nativeApolloSession) State() protocol.ProviderState {
s.mu.Lock() s.mu.Lock()
defer s.mu.Unlock() defer s.mu.Unlock()
return s.state return s.state
} }
func (s *nativeApolloSession) writeControl(ctx context.Context, packet ControlPacket) error { func (s *nativeApolloSession) Telemetry() ProviderTelemetry {
encoded, err := EncodeControlPacket(packet) telemetry := s.enet.telemetry()
telemetry.State = s.State().State
telemetry.MediaDrops = s.mediaDrops.Load()
return telemetry
}
func (s *nativeApolloSession) writeApolloControl(channel uint8, reliable bool, typeID uint16, payload []byte) error {
if s.enet == nil || s.control == nil {
return ErrProviderDisconnected
}
s.controlMu.Lock()
encoded, err := s.control.SealClient(typeID, payload)
s.controlMu.Unlock()
if err != nil { if err != nil {
return err return err
} }
if deadline, ok := ctx.Deadline(); ok { if reliable {
_ = s.conn.SetWriteDeadline(deadline) return s.enet.SendReliable(channel, encoded)
} }
if _, err := s.conn.Write(encoded); err != nil { return s.enet.SendUnsequenced(channel, encoded)
return err
}
return nil
} }
func (s *nativeApolloSession) readMedia() { func (s *nativeApolloSession) periodicApolloPing() {
defer close(s.readDone) ticker := time.NewTicker(100 * time.Millisecond)
header := make([]byte, 4) defer ticker.Stop()
for { for {
if _, err := io.ReadFull(s.conn, header); err != nil { select {
case <-s.done:
return
case <-ticker.C:
if err := s.writeApolloControl(apolloChannelGeneric, true, apolloControlTypePing, []byte{4, 0, 0, 0, 0, 0, 0, 0}); err != nil {
s.handleApolloDisconnect(err)
return return
} }
if header[0] != '$' || (header[1] != 0 && header[1] != 1) {
return
}
length := int(header[2])<<8 | int(header[3])
if length > 65536 {
return
}
payload := make([]byte, length)
if _, err := io.ReadFull(s.conn, payload); err != nil {
return
}
if header[1] == 0 {
pushLatest(s.video, payload)
} else {
pushLatest(s.audio, payload)
} }
} }
} }
func pushLatest(channel chan []byte, payload []byte) { func (s *nativeApolloSession) periodicApolloMediaPing() {
ticker := time.NewTicker(500 * time.Millisecond)
defer ticker.Stop()
sequence := uint32(2)
for {
select {
case <-s.done:
return
case <-ticker.C:
if s.audioConn == nil || s.videoConn == nil {
s.handleApolloDisconnect(ErrProviderDisconnected)
return
}
if _, err := s.audioConn.Write(apolloMediaPing(s.audioPing, sequence)); err != nil {
s.handleApolloDisconnect(err)
return
}
if _, err := s.videoConn.Write(apolloMediaPing(s.videoPing, sequence)); err != nil {
s.handleApolloDisconnect(err)
return
}
sequence++
}
}
}
func (s *nativeApolloSession) handleApolloControlPayload(_ uint8, _ bool, payload []byte) {
s.controlMu.Lock()
message, err := s.control.OpenHost(payload)
s.controlMu.Unlock()
if err != nil {
s.handleApolloDisconnect(err)
return
}
switch message.typeID {
case apolloControlTypeTerm:
if len(message.payload) != 4 {
s.handleApolloDisconnect(ErrProviderMalformed)
return
}
s.mu.Lock()
s.state.State = ProviderStateTerminated
s.mu.Unlock()
s.emitProviderEvent(ProviderEvent{Kind: ProviderEventTerminated, Payload: message.payload})
case apolloControlTypeRumble:
if len(message.payload) != 10 {
s.handleApolloDisconnect(ErrProviderMalformed)
return
}
controller := binary.LittleEndian.Uint16(message.payload[4:6])
if controller > 15 {
s.handleApolloDisconnect(ErrProviderMalformed)
return
}
payload := make([]byte, 5)
payload[0] = byte(controller)
binary.BigEndian.PutUint16(payload[1:3], binary.LittleEndian.Uint16(message.payload[6:8]))
binary.BigEndian.PutUint16(payload[3:5], binary.LittleEndian.Uint16(message.payload[8:10]))
s.emitProviderEvent(ProviderEvent{Kind: ProviderEventRumble, Payload: payload})
case apolloControlTypeHDR:
if len(message.payload) != 27 || message.payload[0] > 1 {
s.handleApolloDisconnect(ErrProviderMalformed)
return
}
s.emitProviderEvent(ProviderEvent{Kind: ProviderEventHDR, Payload: []byte{message.payload[0]}})
}
}
func (s *nativeApolloSession) emitProviderEvent(event ProviderEvent) {
select {
case s.events <- event:
default:
s.handleApolloDisconnect(ErrProviderMalformed)
}
}
func (s *nativeApolloSession) handleApolloDisconnect(err error) {
if err == nil {
return
}
s.mu.Lock()
if s.state.State != ProviderStateTerminated {
s.state.State = ProviderStateDisconnected
}
s.mu.Unlock()
}
func (s *nativeApolloSession) closeMediaChannels() {
s.channelsOnce.Do(func() {
close(s.video)
close(s.audio)
})
}
func (s *nativeApolloSession) readUDPMedia() {
if s.media == nil {
close(s.readDone)
s.closeMediaChannels()
return
}
var readers sync.WaitGroup
readers.Add(2)
read := func(conn *net.UDPConn, output chan []byte, video bool) {
defer readers.Done()
buffer := make([]byte, apolloMediaMaximumPacket+1)
for {
if err := conn.SetReadDeadline(time.Now().Add(250 * time.Millisecond)); err != nil {
return
}
count, err := conn.Read(buffer)
if err != nil {
if networkErr, ok := err.(net.Error); ok && networkErr.Timeout() {
select {
case <-s.done:
return
default:
continue
}
}
return
}
if count > apolloMediaMaximumPacket {
continue
}
var payloads [][]byte
if video {
shard, openErr := s.media.OpenVideo(buffer[:count])
if openErr != nil {
continue
}
payload, err := s.videoFEC.Add(shard)
if err != nil || len(payload) == 0 {
continue
}
payloads = [][]byte{payload}
} else {
shard, openErr := s.media.OpenAudio(buffer[:count])
if openErr != nil {
continue
}
payloads, err = s.audioFEC.Add(s.media, shard)
}
if err != nil {
continue
}
for _, payload := range payloads {
if len(payload) != 0 {
if pushLatest(output, payload) {
s.mediaDrops.Add(1)
}
}
}
}
}
go read(s.audioConn, s.audio, false)
go read(s.videoConn, s.video, true)
go func() {
readers.Wait()
close(s.readDone)
s.closeMediaChannels()
}()
}
func pushLatest(channel chan []byte, payload []byte) bool {
select { select {
case channel <- payload: case channel <- payload:
return false
default: default:
select { select {
case <-channel: case <-channel:
@@ -357,7 +772,9 @@ func pushLatest(channel chan []byte, payload []byte) {
} }
select { select {
case channel <- payload: case channel <- payload:
return true
default: default:
return true
} }
} }
} }
+776 -11
View File
@@ -2,22 +2,34 @@ package gateway
import ( import (
"context" "context"
"crypto/aes"
"crypto/cipher"
"crypto/sha256" "crypto/sha256"
"crypto/tls" "crypto/tls"
"crypto/x509" "crypto/x509"
"encoding/binary" "encoding/binary"
"encoding/hex" "encoding/hex"
"encoding/pem" "encoding/pem"
"errors"
"fmt"
"io" "io"
"net" "net"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"strconv" "strconv"
"strings"
"testing" "testing"
"time"
protocol "git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol/gen/go/protocol" protocol "git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol/gen/go/protocol"
) )
type nativeRoundTripFunc func(*http.Request) (*http.Response, error)
func (fn nativeRoundTripFunc) RoundTrip(request *http.Request) (*http.Response, error) {
return fn(request)
}
func TestNativeApolloManagementUsesSessionScopedMTLS(t *testing.T) { func TestNativeApolloManagementUsesSessionScopedMTLS(t *testing.T) {
serverTLS, clientTLS := testTLS(t) serverTLS, clientTLS := testTLS(t)
server := httptest.NewUnstartedServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { server := httptest.NewUnstartedServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) {
@@ -47,6 +59,7 @@ func TestNativeApolloManagementUsesSessionScopedMTLS(t *testing.T) {
ClientCertificatePem: certificatePEM(t, clientTLS.Certificates[0]), ClientCertificatePem: certificatePEM(t, clientTLS.Certificates[0]),
ClientPrivateKeyPem: privateKeyPEM(t, clientTLS.Certificates[0]), ClientPrivateKeyPem: privateKeyPEM(t, clientTLS.Certificates[0]),
ServerCertificatePem: certificatePEM(t, tls.Certificate{Certificate: [][]byte{serverTLS.Certificates[0].Certificate[1]}}), ServerCertificatePem: certificatePEM(t, tls.Certificate{Certificate: [][]byte{serverTLS.Certificates[0].Certificate[1]}}),
ClipboardPolicy: protocol.ClipboardPolicy{MaxTextBytes: 65536, MaxUpdatesPerMinute: 30},
} }
data, err := NewNativeApolloBackend().Management(context.Background(), LaunchRequest{SessionID: "session-1", ProviderProfile: ProviderProfileApollo, ProviderWork: work}) data, err := NewNativeApolloBackend().Management(context.Background(), LaunchRequest{SessionID: "session-1", ProviderProfile: ProviderProfileApollo, ProviderWork: work})
if err != nil { if err != nil {
@@ -57,13 +70,28 @@ func TestNativeApolloManagementUsesSessionScopedMTLS(t *testing.T) {
} }
} }
func TestNativeApolloSetupLaunchesInventoriedApplicationWithEncryptedRTSP(t *testing.T) { func TestNativeApolloSetupRequiresModernEncryptedRTSPOrder(t *testing.T) {
serverTLS, clientTLS := testTLS(t) serverTLS, clientTLS := testTLS(t)
streamListener, err := net.Listen("tcp", "127.0.0.1:0") streamListener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
defer streamListener.Close() defer streamListener.Close()
controlServer, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
if err != nil {
t.Fatal(err)
}
defer controlServer.Close()
audioServer, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
if err != nil {
t.Fatal(err)
}
defer audioServer.Close()
videoServer, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
if err != nil {
t.Fatal(err)
}
defer videoServer.Close()
streamHost, streamPortText, err := net.SplitHostPort(streamListener.Addr().String()) streamHost, streamPortText, err := net.SplitHostPort(streamListener.Addr().String())
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
@@ -72,18 +100,43 @@ func TestNativeApolloSetupLaunchesInventoriedApplicationWithEncryptedRTSP(t *tes
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
type fixtureKeyMaterial struct {
key []byte
keyID uint32
}
keyReady := make(chan []byte, 1) keyReady := make(chan []byte, 1)
keyMaterialReady := make(chan fixtureKeyMaterial, 1)
clipboardWrites := make(chan string, 1)
cancelCalls := make(chan struct{}, 1)
streamDone := make(chan error, 1) streamDone := make(chan error, 1)
go func() { go func() {
key := <-keyReady
codec, codecErr := newEncryptedRTSPCodec(key)
if codecErr != nil {
streamDone <- codecErr
return
}
expectedMethods := []string{"OPTIONS", "DESCRIBE", "SETUP", "SETUP", "SETUP", "ANNOUNCE", "PLAY"}
expectedTargets := []string{"rtspenc://" + streamListener.Addr().String(), "rtspenc://" + streamListener.Addr().String(), "streamid=audio/0/0", "streamid=video/0/0", "streamid=control/13/0", "streamid=control/13/0", "/"}
describeBody := "a=x-ss-general.featureFlags:1\r\na=x-ss-general.encryptionSupported:7\r\na=x-ss-general.encryptionRequested:1\r\na=fmtp:97 surround-params=21101\r\n"
responses := []string{
"RTSP/1.0 200 OK\r\nCSeq: 1\r\n\r\n",
fmt.Sprintf("RTSP/1.0 200 OK\r\nCSeq: 2\r\nContent-Type: application/sdp\r\nContent-Length: %d\r\n\r\n%s", len(describeBody), describeBody),
fmt.Sprintf("RTSP/1.0 200 OK\r\nCSeq: 3\r\nSession: fixture-session;timeout=90\r\nTransport: unicast;server_port=%d\r\nX-SS-Ping-Payload: 0123456789abcdef\r\n\r\n", audioServer.LocalAddr().(*net.UDPAddr).Port),
fmt.Sprintf("RTSP/1.0 200 OK\r\nCSeq: 4\r\nSession: fixture-session\r\nTransport: unicast;server_port=%d\r\nX-SS-Ping-Payload: fedcba9876543210\r\n\r\n", videoServer.LocalAddr().(*net.UDPAddr).Port),
fmt.Sprintf("RTSP/1.0 200 OK\r\nCSeq: 5\r\nSession: fixture-session\r\nTransport: unicast;server_port=%d\r\nX-SS-Connect-Data: 305419896\r\n\r\n", controlServer.LocalAddr().(*net.UDPAddr).Port),
"RTSP/1.0 200 OK\r\nCSeq: 6\r\nSession: fixture-session\r\n\r\n",
"RTSP/1.0 200 OK\r\nCSeq: 7\r\nSession: fixture-session\r\n\r\n",
}
for index, method := range expectedMethods {
connection, acceptErr := streamListener.Accept() connection, acceptErr := streamListener.Accept()
if acceptErr != nil { if acceptErr != nil {
streamDone <- acceptErr streamDone <- acceptErr
return return
} }
defer connection.Close()
key := <-keyReady
header := make([]byte, encryptedRTSPHeaderSize) header := make([]byte, encryptedRTSPHeaderSize)
if _, readErr := io.ReadFull(connection, header); readErr != nil { if _, readErr := io.ReadFull(connection, header); readErr != nil {
_ = connection.Close()
streamDone <- readErr streamDone <- readErr
return return
} }
@@ -91,21 +144,39 @@ func TestNativeApolloSetupLaunchesInventoriedApplicationWithEncryptedRTSP(t *tes
frame := make([]byte, encryptedRTSPHeaderSize+int(length)) frame := make([]byte, encryptedRTSPHeaderSize+int(length))
copy(frame, header) copy(frame, header)
if _, readErr := io.ReadFull(connection, frame[encryptedRTSPHeaderSize:]); readErr != nil { if _, readErr := io.ReadFull(connection, frame[encryptedRTSPHeaderSize:]); readErr != nil {
_ = connection.Close()
streamDone <- readErr streamDone <- readErr
return return
} }
codec, codecErr := newEncryptedRTSPCodec(key)
if codecErr != nil {
streamDone <- codecErr
return
}
plaintext, decryptErr := codec.OpenClient(frame) plaintext, decryptErr := codec.OpenClient(frame)
if decryptErr != nil || string(plaintext) != "SETUP rtsp://"+streamHost+"/streamid=video/0/0 RTSP/1.0\r\nCSeq: 1\r\nTransport: RTP/AVP/TCP;interleaved=0-1\r\n\r\n" { firstLine, _, _ := strings.Cut(string(plaintext), "\r\n")
if decryptErr != nil || firstLine != method+" "+expectedTargets[index]+" RTSP/1.0" || !strings.Contains(string(plaintext), "CSeq: "+strconv.Itoa(index+1)+"\r\n") {
_, _ = connection.Write(hostEncryptedRTSPFrame(t, key, uint32(index+1), []byte("RTSP/1.0 400 Bad Request\r\nCSeq: "+strconv.Itoa(index+1)+"\r\n\r\n")))
_ = connection.Close()
streamDone <- ErrProviderMalformed streamDone <- ErrProviderMalformed
return return
} }
_, writeErr := connection.Write(hostEncryptedRTSPFrame(t, key, 1, []byte("RTSP/1.0 200 OK\r\nSession: host-session\r\nTransport: RTP/AVP/TCP;interleaved=0-1\r\n\r\n"))) if method == "ANNOUNCE" {
for _, required := range []string{
"a=x-nv-video[0].clientViewportWd:1920", "a=x-nv-video[0].clientViewportHt:1080", "a=x-nv-video[0].maxFPS:60",
"a=x-nv-video[0].packetSize:1024", "a=x-nv-vqos[0].bw.maximumBitrateKbps:8000", "a=x-nv-audio.surround.numChannels:2",
"a=x-nv-general.useReliableUdp:13", "a=x-ss-general.encryptionEnabled:7",
} {
if !strings.Contains(string(plaintext), required+"\r\n") {
_ = connection.Close()
streamDone <- ErrProviderMalformed
return
}
}
}
_, writeErr := connection.Write(hostEncryptedRTSPFrame(t, key, uint32(index+1), []byte(responses[index])))
_ = connection.Close()
if writeErr != nil {
streamDone <- writeErr streamDone <- writeErr
return
}
}
streamDone <- nil
}() }()
management := httptest.NewUnstartedServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { management := httptest.NewUnstartedServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) {
if request.TLS == nil || len(request.TLS.PeerCertificates) != 1 { if request.TLS == nil || len(request.TLS.PeerCertificates) != 1 {
@@ -113,6 +184,8 @@ func TestNativeApolloSetupLaunchesInventoriedApplicationWithEncryptedRTSP(t *tes
return return
} }
switch request.URL.Path { switch request.URL.Path {
case "/serverinfo":
_, _ = response.Write([]byte("<root><uniqueid>apollo-server</uniqueid></root>"))
case "/applist": case "/applist":
if request.URL.Query().Get("uniqueid") != "paired-client" { if request.URL.Query().Get("uniqueid") != "paired-client" {
http.Error(response, "wrong client", http.StatusBadRequest) http.Error(response, "wrong client", http.StatusBadRequest)
@@ -130,8 +203,38 @@ func TestNativeApolloSetupLaunchesInventoriedApplicationWithEncryptedRTSP(t *tes
http.Error(response, "bad key", http.StatusBadRequest) http.Error(response, "bad key", http.StatusBadRequest)
return return
} }
keyID, keyIDErr := strconv.ParseUint(query.Get("rikeyid"), 10, 32)
if keyIDErr != nil {
http.Error(response, "bad key id", http.StatusBadRequest)
return
}
keyReady <- key keyReady <- key
keyMaterialReady <- fixtureKeyMaterial{key: append([]byte(nil), key...), keyID: uint32(keyID)}
_, _ = response.Write([]byte("<root status_code=\"200\"><sessionUrl0>rtspenc://" + streamListener.Addr().String() + "</sessionUrl0></root>")) _, _ = response.Write([]byte("<root status_code=\"200\"><sessionUrl0>rtspenc://" + streamListener.Addr().String() + "</sessionUrl0></root>"))
case "/resume":
_, _ = response.Write([]byte("<root status_code=\"200\"/>"))
case "/cancel":
if request.Method != http.MethodGet {
http.Error(response, "cancel must be GET", http.StatusMethodNotAllowed)
return
}
cancelCalls <- struct{}{}
_, _ = response.Write([]byte("<root status_code=\"200\"><cancel>1</cancel></root>"))
case "/actions/clipboard":
if request.URL.Query().Get("type") != "text" || (request.Method != http.MethodGet && request.Method != http.MethodPost) {
http.Error(response, "bad clipboard request", http.StatusBadRequest)
return
}
if request.Method == http.MethodPost {
body, readErr := io.ReadAll(request.Body)
if readErr != nil || string(body) != "client clipboard" {
http.Error(response, "bad clipboard body", http.StatusBadRequest)
return
}
clipboardWrites <- string(body)
return
}
_, _ = response.Write([]byte("fixture clipboard"))
default: default:
http.NotFound(response, request) http.NotFound(response, request)
} }
@@ -155,18 +258,680 @@ func TestNativeApolloSetupLaunchesInventoriedApplicationWithEncryptedRTSP(t *tes
ManagementHost: managementHost, ManagementPort: managementPort, StreamHost: streamHost, StreamPort: streamPort, ManagementHost: managementHost, ManagementPort: managementPort, StreamHost: streamHost, StreamPort: streamPort,
ClientCertificatePem: certificatePEM(t, clientTLS.Certificates[0]), ClientPrivateKeyPem: privateKeyPEM(t, clientTLS.Certificates[0]), ClientCertificatePem: certificatePEM(t, clientTLS.Certificates[0]), ClientPrivateKeyPem: privateKeyPEM(t, clientTLS.Certificates[0]),
ServerCertificatePem: certificatePEM(t, tls.Certificate{Certificate: [][]byte{serverTLS.Certificates[0].Certificate[1]}}), ServerCertificatePem: certificatePEM(t, tls.Certificate{Certificate: [][]byte{serverTLS.Certificates[0].Certificate[1]}}),
ClipboardPolicy: protocol.ClipboardPolicy{MaxTextBytes: 65536, MaxUpdatesPerMinute: 30},
ProviderApplicationTerminationAllowed: true,
} }
backend := NewNativeApolloBackend() backend := NewNativeApolloBackend()
response, err := backend.Setup(context.Background(), LaunchRequest{SessionID: "session-1", ProviderProfile: ProviderProfileApollo, ProviderWork: work}) response, err := backend.Setup(context.Background(), LaunchRequest{SessionID: "session-1", ProviderProfile: ProviderProfileApollo, ProviderWork: work})
if err != nil { if err != nil {
t.Fatalf("Setup() error = %v", err) t.Fatalf("Setup() error = %v", err)
} }
if parsed, err := ParseRTSPResponse(response); err != nil || parsed.Session != "host-session" { if parsed, err := ParseRTSPResponse(response); err != nil || parsed.Session != "fixture-session" {
t.Fatalf("ParseRTSPResponse() = %+v, %v", parsed, err) t.Fatalf("ParseRTSPResponse() = %+v, %v", parsed, err)
} }
if err := <-streamDone; err != nil { if err := <-streamDone; err != nil {
t.Fatalf("stream transaction error = %v", err) t.Fatalf("stream transaction error = %v", err)
} }
material := <-keyMaterialReady
videoData, videoParity := sourceShapedEncryptedVideoFEC(t, material.key)
audioPackets := make([][]byte, 0, apolloAudioDataShards)
for index, payload := range [][]byte{{'A'}, {'B'}, {'C'}, {'D'}} {
audioPackets = append(audioPackets, sourceShapedEncryptedAudioPacket(t, material.key, material.keyID, uint16(index), payload))
}
mediaDone := make(chan error, 2)
serveMedia := func(server *net.UDPConn, ping string, packets [][]byte) {
buffer := make([]byte, apolloMediaMaximumPacket)
count, remote, readErr := server.ReadFromUDP(buffer)
if readErr != nil || count != 20 || string(buffer[:16]) != ping {
if readErr != nil {
mediaDone <- readErr
} else {
mediaDone <- ErrProviderMalformed
}
return
}
for _, packet := range packets {
if _, writeErr := server.WriteToUDP(packet, remote); writeErr != nil {
mediaDone <- writeErr
return
}
}
mediaDone <- nil
}
go serveMedia(audioServer, "0123456789abcdef", audioPackets)
go serveMedia(videoServer, "fedcba9876543210", [][]byte{videoData, videoParity})
type observedControl struct {
typeID uint16
reliable bool
}
controls := make(chan observedControl, 16)
controlRemote := make(chan *net.UDPAddr, 1)
controlDone := make(chan error, 1)
go func() {
buffer := make([]byte, apolloENetMaximumPacket)
count, remote, readErr := controlServer.ReadFromUDP(buffer)
if readErr != nil {
controlDone <- readErr
return
}
connect := buffer[:count]
if count != 52 || connect[4] != apolloENetConnect|apolloENetAcknowledged || connect[5] != 0xff || binary.BigEndian.Uint32(connect[20:24]) != apolloENetChannels || binary.BigEndian.Uint32(connect[48:52]) != 305419896 {
controlDone <- 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(connect[44:48]))
if _, writeErr := controlServer.WriteToUDP(verify, remote); writeErr != nil {
controlDone <- writeErr
return
}
controlRemote <- remote
for {
count, _, readErr = controlServer.ReadFromUDP(buffer)
if readErr != nil {
controlDone <- readErr
return
}
packet := buffer[:count]
if len(packet) < 8 {
controlDone <- ErrProviderMalformed
return
}
command, channel := packet[4]&apolloENetCommandMask, packet[5]
sequence := binary.BigEndian.Uint16(packet[6:8])
switch command {
case 1:
continue
case apolloENetSendReliable:
if len(packet) < 10 || int(binary.BigEndian.Uint16(packet[8:10])) != len(packet)-10 {
controlDone <- ErrProviderMalformed
return
}
typeID, _ := sourceOpenClientControl(t, material.key, packet[10:])
controls <- observedControl{typeID: typeID, reliable: true}
if _, writeErr := controlServer.WriteToUDP(sourceShapedENetAcknowledgePacket(7, 2, channel, sequence), remote); writeErr != nil {
controlDone <- writeErr
return
}
case apolloENetSendUnsequenced:
if len(packet) < 12 || int(binary.BigEndian.Uint16(packet[10:12])) != len(packet)-12 {
controlDone <- ErrProviderMalformed
return
}
typeID, _ := sourceOpenClientControl(t, material.key, packet[12:])
controls <- observedControl{typeID: typeID}
case apolloENetPing:
if _, writeErr := controlServer.WriteToUDP(sourceShapedENetAcknowledgePacket(7, 2, channel, sequence), remote); writeErr != nil {
controlDone <- writeErr
return
}
case apolloENetDisconnect:
_, writeErr := controlServer.WriteToUDP(sourceShapedENetAcknowledgePacket(7, 2, channel, sequence), remote)
controlDone <- writeErr
return
default:
controlDone <- ErrProviderMalformed
return
}
}
}()
request := LaunchRequest{SessionID: "session-1", ProviderProfile: ProviderProfileApollo, ProviderWork: work}
parsed, err := ParseRTSPResponse(response)
if err != nil {
t.Fatal(err)
}
session, err := backend.Open(context.Background(), request, parsed)
if err != nil {
t.Fatalf("Open() error = %v", err)
}
if err := session.Ready(context.Background()); err != nil {
t.Fatalf("Ready() error = %v", err)
}
if clipboard, clipboardErr := session.ReadClipboard(context.Background()); clipboardErr != nil || clipboard != "fixture clipboard" {
t.Fatalf("ReadClipboard() = %q, %v", clipboard, clipboardErr)
}
if err := session.WriteClipboard(context.Background(), "client clipboard"); err != nil {
t.Fatalf("WriteClipboard() error = %v", err)
}
select {
case <-clipboardWrites:
case <-time.After(time.Second):
t.Fatal("provider did not receive authenticated clipboard write")
}
awaitControl := func(want uint16, reliable bool) {
deadline := time.NewTimer(time.Second)
defer deadline.Stop()
for {
select {
case observed := <-controls:
if observed.typeID == want && observed.reliable == reliable {
return
}
case <-deadline.C:
t.Fatalf("provider did not receive control %#x", want)
}
}
}
awaitControl(apolloControlTypeIDR, true)
awaitControl(apolloControlTypeStart, true)
if err := session.Input(context.Background(), InputEvent{Device: "keyboard", Code: 7, Pressed: true}); err != nil {
t.Fatal(err)
}
awaitControl(apolloControlTypeInput, true)
if err := session.Feedback(context.Background(), Feedback{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}}); err != nil {
t.Fatal(err)
}
awaitControl(apolloControlTypeFEC, false)
remote := <-controlRemote
hostTermination := sourceSealHostControl(t, material.key, 0, apolloControlTypeTerm, []byte{1, 2, 3, 4})
if _, err := controlServer.WriteToUDP(sourceShapedENetReliablePacketOn(7, 2, apolloChannelGeneric, 1, hostTermination), remote); err != nil {
t.Fatal(err)
}
select {
case event := <-session.Events():
if event.Kind != ProviderEventTerminated || string(event.Payload) != string([]byte{1, 2, 3, 4}) {
t.Fatalf("provider termination event = %#v", event)
}
case <-time.After(time.Second):
t.Fatal("encrypted host termination was not forwarded")
}
select {
case payload := <-session.Video():
if len(payload) != 1001 || payload[0] != 'A' || payload[1000] != 'B' {
t.Fatalf("source-shaped video relay = %x", payload)
}
case <-time.After(time.Second):
t.Fatal("source-shaped video was not relayed")
}
select {
case payload := <-session.Audio():
if string(payload) != "A" {
t.Fatalf("source-shaped audio relay = %x", payload)
}
case <-time.After(time.Second):
t.Fatal("source-shaped audio was not relayed")
}
terminateCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := session.Terminate(terminateCtx); err != nil {
t.Fatalf("Terminate() error = %v", err)
}
select {
case <-cancelCalls:
case <-time.After(time.Second):
t.Fatal("authorized provider application cancellation was not sent")
}
if err := <-controlDone; err != nil {
t.Fatalf("control fake = %v", err)
}
for range 2 {
if err := <-mediaDone; err != nil {
t.Fatalf("media fake = %v", err)
}
}
}
func TestValidateApolloDescribeRejectsUnknownLines(t *testing.T) {
message := apolloRTSPMessage{
headers: map[string]string{"content-type": "application/sdp"},
body: []byte("a=x-ss-general.featureFlags:1\r\n" +
"a=x-ss-general.encryptionSupported:7\r\n" +
"a=x-ss-general.encryptionRequested:1\r\n" +
"a=fmtp:97 surround-params=21101\r\n" +
"unexpected-source-line\r\n"),
}
if err := validateApolloDescribe(message); err == nil {
t.Fatal("validateApolloDescribe() accepted an unknown source line")
}
}
func TestNativeApolloTerminateSkipsProviderCancelWithoutServerPolicy(t *testing.T) {
session := newNativeApolloSession("session-1")
close(session.readDone)
session.managementClient = &http.Client{Transport: nativeRoundTripFunc(func(request *http.Request) (*http.Response, error) {
t.Fatalf("unauthorized provider cancellation request: %s %s", request.Method, request.URL)
return nil, fmt.Errorf("unexpected provider cancellation")
})}
session.managementHost = "provider.invalid"
session.managementPort = 47984
if err := session.Terminate(context.Background()); err != nil {
t.Fatalf("Terminate() error = %v", err)
}
}
func TestNativeApolloTerminateRetainsCleanupPendingAfterFailure(t *testing.T) {
session := newNativeApolloSession("session-1")
failed, cancel := context.WithCancel(context.Background())
cancel()
if err := session.Terminate(failed); !errors.Is(err, context.Canceled) {
t.Fatalf("first Terminate() error = %v", err)
}
if state := session.State(); state.State != ProviderStateCleanup || !state.CleanupPending {
t.Fatalf("failed cleanup state = %#v", state)
}
if err := session.Terminate(context.Background()); !errors.Is(err, context.Canceled) {
t.Fatalf("second Terminate() error = %v", err)
}
if state := session.State(); state.State != ProviderStateCleanup || !state.CleanupPending {
t.Fatalf("repeated cleanup state = %#v", state)
}
}
func TestApolloReconnectUsesResumeWithFreshRIK(t *testing.T) {
work := protocol.ProviderSessionWork{ReconnectSequence: 1, ApplicationID: "42", ClientID: "paired-client"}
path, values := apolloSessionRequest(work, []byte("0123456789abcdef"), 0x01020304)
if path != "/resume" || values.Get("rikey") != "30313233343536373839616263646566" || values.Get("rikeyid") != "16909060" || values.Get("localAudioPlayMode") != "0" {
t.Fatalf("resume request = path %q values %#v", path, values)
}
if values.Get("appid") != "" || values.Get("uniqueid") != "" || values.Get("corever") != "" {
t.Fatalf("resume request leaked launch-only values %#v", values)
}
}
func TestNativeApolloSessionRelaysOnlyAuthenticatedEncodedUDPMedia(t *testing.T) {
key := []byte("0123456789abcdef")
audioServer, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
if err != nil {
t.Fatal(err)
}
defer audioServer.Close()
videoServer, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
if err != nil {
t.Fatal(err)
}
defer videoServer.Close()
audioClient, err := net.DialUDP("udp", nil, audioServer.LocalAddr().(*net.UDPAddr))
if err != nil {
t.Fatal(err)
}
videoClient, err := net.DialUDP("udp", nil, videoServer.LocalAddr().(*net.UDPAddr))
if err != nil {
t.Fatal(err)
}
session := newNativeApolloSession("session-media")
const keyID = 0x01020304
session.media, err = newApolloMediaCodec(key, keyID)
if err != nil {
t.Fatal(err)
}
session.audioConn, session.videoConn = audioClient, videoClient
go session.readUDPMedia()
videoPacket := sourceShapedEncryptedVideoPacket(t, key, []byte{0x01, 0x02, 0x03})
if _, err := videoServer.WriteToUDP(videoPacket, videoClient.LocalAddr().(*net.UDPAddr)); err != nil {
t.Fatal(err)
}
wantAudio := [][]byte{{0xf8, 0x08}, {0xf8, 0x09}, {0xf8, 0x0a}, {0xf8, 0x0b}}
for index, payload := range wantAudio {
packet := sourceShapedEncryptedAudioPacket(t, key, keyID, uint16(8+index), payload)
if _, err := audioServer.WriteToUDP(packet, audioClient.LocalAddr().(*net.UDPAddr)); err != nil {
t.Fatal(err)
}
}
select {
case payload := <-session.Video():
if string(payload) != string([]byte{0x01, 0x02, 0x03}) {
t.Fatalf("video relay = %x, want encoded payload", payload)
}
case <-time.After(time.Second):
t.Fatal("encrypted video was not relayed")
}
for _, want := range wantAudio {
select {
case payload := <-session.Audio():
if string(payload) != string(want) {
t.Fatalf("audio relay = %x, want %x", payload, want)
}
case <-time.After(time.Second):
t.Fatal("encrypted audio was not relayed")
}
}
dataPacket, parityPacket := sourceShapedEncryptedVideoFEC(t, key)
for _, packet := range [][]byte{dataPacket, parityPacket} {
if _, err := videoServer.WriteToUDP(packet, videoClient.LocalAddr().(*net.UDPAddr)); err != nil {
t.Fatal(err)
}
}
select {
case payload := <-session.Video():
if len(payload) != 1001 || payload[0] != 'A' || payload[999] != 'A' || payload[1000] != 'B' {
t.Fatalf("FEC video relay = %x", payload)
}
case <-time.After(time.Second):
t.Fatal("encrypted FEC video was not recovered")
}
for _, packet := range sourceShapedEncryptedAudioFEC(t, key, keyID) {
if _, err := audioServer.WriteToUDP(packet, audioClient.LocalAddr().(*net.UDPAddr)); err != nil {
t.Fatal(err)
}
}
for _, want := range [][]byte{{'A'}, {'B'}, {'C'}, {'D'}} {
select {
case payload := <-session.Audio():
if string(payload) != string(want) {
t.Fatalf("FEC audio relay = %x, want %x", payload, want)
}
case <-time.After(time.Second):
t.Fatal("encrypted FEC audio was not recovered")
}
}
terminateCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := session.Terminate(terminateCtx); err != nil {
t.Fatalf("Terminate() error = %v", err)
}
}
func TestNativeApolloTerminateReleasesPressedProviderInput(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)
}
peer.state, peer.peerID, peer.outboundSession = apolloENetConnected, 1, 1
codecKey := []byte("0123456789abcdef")
codec, err := newApolloControlCodec(codecKey)
if err != nil {
t.Fatal(err)
}
session := newNativeApolloSession("session-terminate")
session.enet, session.control = peer, codec
close(session.readDone)
packets := make(chan []byte, 3)
go func() {
defer close(packets)
buffer := make([]byte, apolloENetMaximumPacket)
for {
_ = server.SetReadDeadline(time.Now().Add(time.Second))
count, _, readErr := server.ReadFromUDP(buffer)
if readErr != nil {
return
}
packet := append([]byte(nil), buffer[:count]...)
packets <- packet
if len(packet) >= 5 && packet[4]&apolloENetCommandMask == apolloENetDisconnect {
peer.close(nil)
return
}
}
}()
if err := session.Input(context.Background(), InputEvent{Device: "keyboard", Code: 7, Pressed: true}); err != nil {
t.Fatal(err)
}
terminateCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := session.Terminate(terminateCtx); err != nil {
t.Fatalf("Terminate() error = %v", err)
}
var observed [][]byte
for packet := range packets {
observed = append(observed, packet)
}
if len(observed) != 3 {
t.Fatalf("provider packets = %d, want pressed input, release, disconnect", len(observed))
}
typeID, payload := sourceOpenClientControl(t, codecKey, observed[1][10:])
if typeID != apolloControlTypeInput || len(payload) < 8 || binary.LittleEndian.Uint32(payload[4:8]) != 4 {
t.Fatalf("release packet = type %#x payload %x", typeID, payload)
}
}
func TestNativeApolloSessionForwardsEncryptedHostFeedback(t *testing.T) {
key := []byte("0123456789abcdef")
codec, err := newApolloControlCodec(key)
if err != nil {
t.Fatal(err)
}
session := newNativeApolloSession("session-feedback")
session.control = codec
messages := []struct {
typeID uint16
payload []byte
want ProviderEvent
}{
{apolloControlTypeRumble, []byte{0, 0, 0, 0, 1, 0, 0x34, 0x12, 0x78, 0x56}, ProviderEvent{Kind: ProviderEventRumble, Payload: []byte{1, 0x12, 0x34, 0x56, 0x78}}},
{apolloControlTypeHDR, append([]byte{1}, make([]byte, 26)...), ProviderEvent{Kind: ProviderEventHDR, Payload: []byte{1}}},
{apolloControlTypeTerm, []byte{1, 2, 3, 4}, ProviderEvent{Kind: ProviderEventTerminated, Payload: []byte{1, 2, 3, 4}}},
}
for sequence, message := range messages {
session.handleApolloControlPayload(apolloChannelGeneric, true, sourceSealHostControl(t, key, uint32(sequence), message.typeID, message.payload))
select {
case event := <-session.Events():
if event.Kind != message.want.Kind || string(event.Payload) != string(message.want.Payload) {
t.Fatalf("provider event = %#v, want %#v", event, message.want)
}
case <-time.After(time.Second):
t.Fatalf("host event %#x was not forwarded", message.typeID)
}
}
}
func TestPushLatestDropsExactlyOneOldPayload(t *testing.T) {
queue := make(chan []byte, 1)
if dropped := pushLatest(queue, []byte("old")); dropped {
t.Fatal("first media payload dropped")
}
if dropped := pushLatest(queue, []byte("new")); !dropped {
t.Fatal("bounded media queue did not report a drop")
}
if got := string(<-queue); got != "new" {
t.Fatalf("bounded media queue payload = %q", got)
}
}
func sourceShapedEncryptedVideoPacket(t *testing.T, key, encoded []byte) []byte {
t.Helper()
payload := make([]byte, apolloVideoShardPayloadSize)
payload[0], payload[3] = 0x01, 0x01
binary.LittleEndian.PutUint16(payload[4:6], uint16(8+len(encoded)))
copy(payload[8:], encoded)
plaintext := sourceShapedVideoRaw(42, 7, 1, 0x07, 1, 0, 0, payload)
return sourceEncryptVideoRaw(t, key, plaintext, "0123456789aV")
}
func sourceShapedEncryptedVideoFEC(t *testing.T, key []byte) ([]byte, []byte) {
t.Helper()
firstPayload := make([]byte, apolloVideoShardPayloadSize)
firstPayload[0], firstPayload[3] = 0x01, 0x01
binary.LittleEndian.PutUint16(firstPayload[4:6], 1)
for index := 8; index < len(firstPayload); index++ {
firstPayload[index] = 'A'
}
secondPayload := make([]byte, apolloVideoShardPayloadSize)
secondPayload[0] = 'B'
first := sourceShapedVideoRaw(43, 100, 100, 0x05, 2, 50, 0, firstPayload)
second := sourceShapedVideoRaw(43, 101, 101, 0x03, 2, 50, 1, secondPayload)
parity := make([]byte, len(first))
for index := range parity {
parity[index] = first[index] ^ sourceGFMultiply(second[index], 142)
}
sourceConfigureVideoShard(first, 43, 100, 0, 2, 50, 0)
sourceConfigureVideoShard(second, 43, 101, 1, 2, 50, 1)
sourceConfigureVideoShard(parity, 43, 102, 2, 2, 50, 2)
return sourceEncryptVideoRaw(t, key, second, "0123456789bV"), sourceEncryptVideoRaw(t, key, parity, "0123456789cV")
}
func sourceShapedVideoRaw(frame uint32, sequence uint16, streamIndex uint32, flags byte, dataShards, percentage, shardIndex int, payload []byte) []byte {
raw := make([]byte, apolloVideoRawPacketSize)
binary.LittleEndian.PutUint32(raw[16:20], streamIndex<<8)
binary.LittleEndian.PutUint32(raw[20:24], frame)
raw[24], raw[26] = flags, 0x10
binary.LittleEndian.PutUint32(raw[28:32], uint32(shardIndex<<12|dataShards<<22|percentage<<4))
copy(raw[32:], payload)
sourceConfigureVideoShard(raw, frame, sequence, streamIndex, dataShards, percentage, shardIndex)
return raw
}
func sourceConfigureVideoShard(raw []byte, frame uint32, sequence uint16, streamIndex uint32, dataShards, percentage, shardIndex int) {
raw[0] = 0x90
binary.BigEndian.PutUint16(raw[2:4], sequence)
binary.BigEndian.PutUint32(raw[4:8], 99)
binary.LittleEndian.PutUint32(raw[20:24], frame)
raw[27] = 0
binary.LittleEndian.PutUint32(raw[28:32], uint32(shardIndex<<12|dataShards<<22|percentage<<4))
}
func sourceEncryptVideoRaw(t *testing.T, key, plaintext []byte, iv string) []byte {
t.Helper()
if len(plaintext) != apolloVideoRawPacketSize || len(iv) != 12 {
t.Fatal("invalid source video fixture")
}
block, err := aes.NewCipher(key)
if err != nil {
t.Fatal(err)
}
aead, err := cipher.NewGCM(block)
if err != nil {
t.Fatal(err)
}
sealed := aead.Seal(nil, []byte(iv), plaintext, nil)
packet := make([]byte, 32+len(plaintext))
copy(packet[:12], iv)
binary.LittleEndian.PutUint32(packet[12:16], binary.LittleEndian.Uint32(plaintext[20:24]))
copy(packet[16:32], sealed[len(plaintext):])
copy(packet[32:], sealed[:len(plaintext)])
return packet
}
func sourceOpenClientControl(t *testing.T, key, packet []byte) (uint16, []byte) {
t.Helper()
if len(packet) < apolloControlHeaderSize+apolloControlTagSize+apolloControlInnerSize || binary.LittleEndian.Uint16(packet[:2]) != apolloControlOuterType || int(binary.LittleEndian.Uint16(packet[2:4])) != len(packet)-4 {
t.Fatalf("control packet = %x", packet)
}
block, err := aes.NewCipher(key)
if err != nil {
t.Fatal(err)
}
aead, err := cipher.NewGCM(block)
if err != nil {
t.Fatal(err)
}
nonce := make([]byte, 12)
binary.LittleEndian.PutUint32(nonce, binary.LittleEndian.Uint32(packet[4:8]))
nonce[10], nonce[11] = 'C', 'C'
sealed := append([]byte(nil), packet[24:]...)
sealed = append(sealed, packet[8:24]...)
plaintext, err := aead.Open(nil, nonce, sealed, nil)
if err != nil || len(plaintext) < apolloControlInnerSize || int(binary.LittleEndian.Uint16(plaintext[2:4])) != len(plaintext)-apolloControlInnerSize {
t.Fatalf("control open = %x, %v", packet, err)
}
return binary.LittleEndian.Uint16(plaintext[:2]), append([]byte(nil), plaintext[4:]...)
}
func sourceSealHostControl(t *testing.T, key []byte, sequence uint32, typeID uint16, payload []byte) []byte {
t.Helper()
block, err := aes.NewCipher(key)
if err != nil {
t.Fatal(err)
}
aead, err := cipher.NewGCM(block)
if err != nil {
t.Fatal(err)
}
inner := make([]byte, apolloControlInnerSize+len(payload))
binary.LittleEndian.PutUint16(inner[:2], typeID)
binary.LittleEndian.PutUint16(inner[2:4], uint16(len(payload)))
copy(inner[4:], payload)
nonce := make([]byte, 12)
binary.LittleEndian.PutUint32(nonce, sequence)
nonce[10], nonce[11] = 'H', 'C'
sealed := aead.Seal(nil, nonce, inner, nil)
packet := make([]byte, apolloControlHeaderSize+len(sealed))
binary.LittleEndian.PutUint16(packet[:2], apolloControlOuterType)
binary.LittleEndian.PutUint16(packet[2:4], uint16(4+len(sealed)))
binary.LittleEndian.PutUint32(packet[4:8], sequence)
copy(packet[8:24], sealed[len(inner):])
copy(packet[24:], sealed[:len(inner)])
return packet
}
func sourceGFMultiply(first, second byte) byte {
var product byte
for second != 0 {
if second&1 != 0 {
product ^= first
}
high := first & 0x80
first <<= 1
if high != 0 {
first ^= 0x1d
}
second >>= 1
}
return product
}
func sourceShapedEncryptedAudioPacket(t *testing.T, key []byte, keyID uint32, sequence uint16, encoded []byte) []byte {
return sourceShapedEncryptedAudioPacketWithHeaders(t, key, keyID, sequence, 99, 1, encoded)
}
func sourceShapedEncryptedAudioPacketWithHeaders(t *testing.T, key []byte, keyID uint32, sequence uint16, timestamp, ssrc uint32, encoded []byte) []byte {
t.Helper()
packet := make([]byte, 12)
packet[0], packet[1] = 0x80, 97
binary.BigEndian.PutUint16(packet[2:4], sequence)
binary.BigEndian.PutUint32(packet[4:8], timestamp)
binary.BigEndian.PutUint32(packet[8:12], ssrc)
padded := append([]byte(nil), encoded...)
padding := aes.BlockSize - len(padded)%aes.BlockSize
for range padding {
padded = append(padded, byte(padding))
}
iv := make([]byte, aes.BlockSize)
binary.BigEndian.PutUint32(iv, keyID+uint32(sequence))
block, err := aes.NewCipher(key)
if err != nil {
t.Fatal(err)
}
cipher.NewCBCEncrypter(block, iv).CryptBlocks(padded, padded)
return append(packet, padded...)
}
func sourceShapedEncryptedAudioFEC(t *testing.T, key []byte, keyID uint32) [][]byte {
t.Helper()
const base = uint16(12)
const timestamp = uint32(100)
data := make([][]byte, apolloAudioDataShards)
for index, payload := range [][]byte{{'A'}, {'B'}, {'C'}, {'D'}} {
data[index] = sourceShapedEncryptedAudioPacketWithHeaders(t, key, keyID, base+uint16(index), timestamp+uint32(index*5), 0, payload)
}
parity := make([]byte, len(data[0])-apolloRTPHeaderSize)
for index, coefficient := range sourceAudioFECRow() {
for offset, value := range data[index][apolloRTPHeaderSize:] {
parity[offset] ^= sourceGFMultiply(value, coefficient)
}
}
fec := make([]byte, apolloRTPHeaderSize+12+len(parity))
fec[0], fec[1] = 0x80, 127
binary.BigEndian.PutUint16(fec[2:4], base+apolloAudioDataShards)
fec[apolloRTPHeaderSize+1] = 97
binary.BigEndian.PutUint16(fec[apolloRTPHeaderSize+2:apolloRTPHeaderSize+4], base)
binary.BigEndian.PutUint32(fec[apolloRTPHeaderSize+4:apolloRTPHeaderSize+8], timestamp)
binary.BigEndian.PutUint32(fec[apolloRTPHeaderSize+8:apolloRTPHeaderSize+12], 0)
copy(fec[apolloRTPHeaderSize+12:], parity)
return [][]byte{data[1], data[2], data[3], fec}
}
func sourceAudioFECRow() []byte {
return []byte{0x77, 0x40, 0x38, 0x0e}
} }
func certificatePEM(t *testing.T, certificate tls.Certificate) string { func certificatePEM(t *testing.T, certificate tls.Certificate) string {
+36
View File
@@ -0,0 +1,36 @@
package gateway
import (
"net"
"testing"
"time"
)
func FuzzApolloENetAndRTSPParsersStayBounded(f *testing.F) {
server, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
if err != nil {
f.Fatal(err)
}
defer server.Close()
client, err := net.DialUDP("udp", nil, server.LocalAddr().(*net.UDPAddr))
if err != nil {
f.Fatal(err)
}
defer client.Close()
f.Add([]byte{0x80, 0x01, 0, 0, apolloENetPing | apolloENetAcknowledged, 0, 0, 1})
f.Add([]byte("RTSP/1.0 200 OK\r\nCSeq: 1\r\n\r\n"))
f.Fuzz(func(t *testing.T, data []byte) {
if len(data) > apolloENetMaximumPacket+1 {
data = data[:apolloENetMaximumPacket+1]
}
peer := &apolloENetPeer{
conn: client, now: time.Now, state: apolloENetConnected, peerID: 1, inboundSession: 0,
pending: make(map[apolloENetPendingKey]*apolloENetPending), rtt: time.Millisecond, variance: time.Millisecond,
}
_ = peer.handleDatagram(data)
_, _ = parseApolloRTSPMessage(data)
_, _ = ParseRTSPResponse(data)
_ = validateApolloDescribe(apolloRTSPMessage{headers: map[string]string{"content-type": "application/sdp"}, body: data})
_, _ = parseApolloRTP(data)
})
}
+509
View File
@@ -0,0 +1,509 @@
package gateway
import (
"context"
"fmt"
"io"
"net"
"net/url"
"strconv"
"strings"
"time"
protocol "git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol/gen/go/protocol"
)
const (
maxApolloRTSPHeaders = 16 << 10
maxApolloRTSPBody = 48 << 10
maxApolloRTSPPort = 65535
apolloEncryptionAll = 0x07
)
type apolloRTSPMessage struct {
status int
cseq uint32
headers map[string]string
body []byte
raw []byte
}
type apolloRTSPSetup struct {
sessionID string
audioPort int
videoPort int
controlPort int
audioPing []byte
videoPing []byte
controlConnect uint32
streamHost string
streamPort int64
providerWork protocol.ProviderSessionWork
streamKey []byte
streamKeyID uint32
}
func (b *NativeApolloBackend) performRTSPHandshake(ctx context.Context, work protocol.ProviderSessionWork, key []byte, keyID uint32, streamURL *url.URL) (*apolloRTSPSetup, []byte, error) {
if streamURL == nil || streamURL.Scheme != "rtspenc" || streamURL.Hostname() != work.StreamHost || streamURL.Port() != strconv.FormatInt(work.StreamPort, 10) || streamURL.User != nil || streamURL.RawQuery != "" || streamURL.Fragment != "" {
return nil, nil, ErrProviderMalformed
}
codec, err := newEncryptedRTSPCodec(key)
if err != nil {
return nil, nil, err
}
request := func(method, target, session string, headers []apolloRTSPHeader, body []byte, sequence uint32) (apolloRTSPMessage, error) {
return b.encryptedRTSPRequest(ctx, work, codec, method, target, session, headers, body, sequence)
}
options, err := request("OPTIONS", streamURL.String(), "", nil, nil, 1)
if err != nil {
return nil, nil, err
}
describe, err := request("DESCRIBE", streamURL.String(), "", []apolloRTSPHeader{{"Accept", "application/sdp"}, {"If-Modified-Since", "Thu, 01 Jan 1970 00:00:00 GMT"}}, nil, 2)
if err != nil {
return nil, nil, err
}
if err := validateApolloDescribe(describe); err != nil {
return nil, nil, err
}
setupHeaders := []apolloRTSPHeader{{"Transport", "unicast;X-GS-ClientPort=50000-50001"}, {"If-Modified-Since", "Thu, 01 Jan 1970 00:00:00 GMT"}}
audio, err := request("SETUP", "streamid=audio/0/0", "", setupHeaders, nil, 3)
if err != nil {
return nil, nil, err
}
sessionID, err := apolloRTSPSession(audio)
if err != nil {
return nil, nil, err
}
audioPort, err := apolloRTSPServerPort(audio)
if err != nil {
return nil, nil, err
}
audioPing, err := apolloRTSPPingPayload(audio)
if err != nil {
return nil, nil, err
}
video, err := request("SETUP", "streamid=video/0/0", sessionID, setupHeaders, nil, 4)
if err != nil {
return nil, nil, err
}
if err := apolloRTSPMatchSession(video, sessionID); err != nil {
return nil, nil, err
}
videoPort, err := apolloRTSPServerPort(video)
if err != nil {
return nil, nil, err
}
videoPing, err := apolloRTSPPingPayload(video)
if err != nil {
return nil, nil, err
}
control, err := request("SETUP", "streamid=control/13/0", sessionID, setupHeaders, nil, 5)
if err != nil {
return nil, nil, err
}
if err := apolloRTSPMatchSession(control, sessionID); err != nil {
return nil, nil, err
}
controlPort, err := apolloRTSPServerPort(control)
if err != nil {
return nil, nil, err
}
connectData, err := apolloRTSPConnectData(control)
if err != nil {
return nil, nil, err
}
announceBody := apolloAnnounceProfile()
announce, err := request("ANNOUNCE", "streamid=control/13/0", sessionID, []apolloRTSPHeader{{"Content-Type", "application/sdp"}}, announceBody, 6)
if err != nil {
return nil, nil, err
}
if err := apolloRTSPMatchSession(announce, sessionID); err != nil {
return nil, nil, err
}
play, err := request("PLAY", "/", sessionID, nil, nil, 7)
if err != nil {
return nil, nil, err
}
if err := apolloRTSPMatchSession(play, sessionID); err != nil {
return nil, nil, err
}
_ = options
setup := &apolloRTSPSetup{
sessionID: sessionID, audioPort: audioPort, videoPort: videoPort, controlPort: controlPort,
audioPing: audioPing, videoPing: videoPing, controlConnect: connectData, streamHost: work.StreamHost,
streamPort: work.StreamPort, providerWork: work, streamKey: append([]byte(nil), key...), streamKeyID: keyID,
}
return setup, append([]byte(nil), control.raw...), nil
}
type apolloRTSPHeader struct{ key, value string }
func (b *NativeApolloBackend) encryptedRTSPRequest(ctx context.Context, work protocol.ProviderSessionWork, codec *encryptedRTSPCodec, method, target, session string, headers []apolloRTSPHeader, body []byte, sequence uint32) (apolloRTSPMessage, error) {
if b == nil || b.Dialer == nil || codec == nil || sequence == 0 || len(body) > maxApolloRTSPBody || method == "" || target == "" {
return apolloRTSPMessage{}, ErrProviderMalformed
}
conn, err := b.Dialer.DialContext(ctx, "tcp", net.JoinHostPort(work.StreamHost, strconv.FormatInt(work.StreamPort, 10)))
if err != nil {
return apolloRTSPMessage{}, err
}
defer conn.Close()
deadline := time.Now().Add(5 * time.Second)
if contextDeadline, ok := ctx.Deadline(); ok && contextDeadline.Before(deadline) {
deadline = contextDeadline
}
if err := conn.SetDeadline(deadline); err != nil {
return apolloRTSPMessage{}, err
}
plaintext, err := buildApolloRTSPRequest(method, target, session, headers, body, sequence)
if err != nil {
return apolloRTSPMessage{}, err
}
frame, err := codec.SealClient(plaintext)
if err != nil {
return apolloRTSPMessage{}, err
}
if _, err := conn.Write(frame); err != nil {
return apolloRTSPMessage{}, err
}
response, err := readEncryptedRTSPMessage(conn, codec)
if err != nil {
return apolloRTSPMessage{}, err
}
if response.status != 200 || response.cseq != sequence {
return apolloRTSPMessage{}, ErrProviderMalformed
}
return response, nil
}
func buildApolloRTSPRequest(method, target, session string, headers []apolloRTSPHeader, body []byte, sequence uint32) ([]byte, error) {
if method == "" || target == "" || strings.ContainsAny(method, "\r\n ") || strings.ContainsAny(target, "\r\n") || sequence == 0 {
return nil, ErrProviderMalformed
}
var builder strings.Builder
builder.Grow(256 + len(body))
fmt.Fprintf(&builder, "%s %s RTSP/1.0\r\nCSeq: %d\r\n", method, target, sequence)
if session != "" {
if !validApolloRTSPToken(session) {
return nil, ErrProviderMalformed
}
fmt.Fprintf(&builder, "Session: %s\r\n", session)
}
seen := map[string]struct{}{"cseq": {}, "session": {}}
for _, header := range headers {
key := strings.ToLower(header.key)
if !validApolloRTSPToken(header.key) || header.value == "" || len(header.value) > 1024 || strings.ContainsAny(header.value, "\r\n") {
return nil, ErrProviderMalformed
}
if _, ok := seen[key]; ok {
return nil, ErrProviderMalformed
}
seen[key] = struct{}{}
fmt.Fprintf(&builder, "%s: %s\r\n", header.key, header.value)
}
if len(body) != 0 {
if _, ok := seen["content-length"]; ok {
return nil, ErrProviderMalformed
}
fmt.Fprintf(&builder, "Content-Length: %d\r\n", len(body))
}
builder.WriteString("\r\n")
builder.Write(body)
return []byte(builder.String()), nil
}
func readEncryptedRTSPMessage(conn net.Conn, codec *encryptedRTSPCodec) (apolloRTSPMessage, error) {
header := make([]byte, encryptedRTSPHeaderSize)
if _, err := io.ReadFull(conn, header); err != nil {
return apolloRTSPMessage{}, err
}
length := uint64(header[0]&0x7f)<<24 | uint64(header[1])<<16 | uint64(header[2])<<8 | uint64(header[3])
if header[0]&0x80 == 0 || length == 0 || length > encryptedRTSPMaxPayload {
return apolloRTSPMessage{}, ErrProviderMalformed
}
frame := make([]byte, encryptedRTSPHeaderSize+int(length))
copy(frame, header)
if _, err := io.ReadFull(conn, frame[encryptedRTSPHeaderSize:]); err != nil {
return apolloRTSPMessage{}, err
}
plaintext, err := codec.OpenHost(frame)
if err != nil {
return apolloRTSPMessage{}, ErrProviderMalformed
}
return parseApolloRTSPMessage(plaintext)
}
func parseApolloRTSPMessage(data []byte) (apolloRTSPMessage, error) {
if len(data) == 0 || len(data) > encryptedRTSPMaxPayload {
return apolloRTSPMessage{}, ErrProviderMalformed
}
headerEnd := strings.Index(string(data), "\r\n\r\n")
if headerEnd < 0 || headerEnd+4 > maxApolloRTSPHeaders {
return apolloRTSPMessage{}, ErrProviderMalformed
}
lines := strings.Split(string(data[:headerEnd]), "\r\n")
if len(lines) < 1 {
return apolloRTSPMessage{}, ErrProviderMalformed
}
parts := strings.SplitN(lines[0], " ", 3)
if len(parts) != 3 || parts[0] != "RTSP/1.0" || len(parts[2]) == 0 || len(parts[2]) > 128 {
return apolloRTSPMessage{}, ErrProviderMalformed
}
status, err := strconv.Atoi(parts[1])
if err != nil || status < 100 || status > 599 {
return apolloRTSPMessage{}, ErrProviderMalformed
}
message := apolloRTSPMessage{status: status, headers: make(map[string]string), raw: append([]byte(nil), data...)}
for _, line := range lines[1:] {
key, value, ok := strings.Cut(line, ":")
key = strings.ToLower(strings.TrimSpace(key))
value = strings.TrimSpace(value)
if !ok || !validApolloRTSPToken(key) || value == "" || len(value) > 1024 {
return apolloRTSPMessage{}, ErrProviderMalformed
}
if _, duplicate := message.headers[key]; duplicate {
return apolloRTSPMessage{}, ErrProviderMalformed
}
message.headers[key] = value
}
cseq, ok := message.headers["cseq"]
if !ok {
return apolloRTSPMessage{}, ErrProviderMalformed
}
parsedCSeq, err := strconv.ParseUint(cseq, 10, 32)
if err != nil || parsedCSeq == 0 {
return apolloRTSPMessage{}, ErrProviderMalformed
}
message.cseq = uint32(parsedCSeq)
message.body = append([]byte(nil), data[headerEnd+4:]...)
if len(message.body) > maxApolloRTSPBody {
return apolloRTSPMessage{}, ErrProviderMalformed
}
if length, hasLength := message.headers["content-length"]; hasLength {
declared, err := strconv.ParseUint(length, 10, 16)
if err != nil || int(declared) != len(message.body) {
return apolloRTSPMessage{}, ErrProviderMalformed
}
} else if len(message.body) != 0 {
return apolloRTSPMessage{}, ErrProviderMalformed
}
return message, nil
}
func validateApolloDescribe(message apolloRTSPMessage) error {
if message.headers["content-type"] != "application/sdp" || len(message.body) == 0 {
return ErrProviderMalformed
}
body := string(message.body)
lineEnding := "\n"
if strings.Contains(body, "\r\n") {
if strings.Contains(strings.ReplaceAll(body, "\r\n", ""), "\n") {
return ErrProviderMalformed
}
lineEnding = "\r\n"
}
if !strings.HasSuffix(body, lineEnding) {
return ErrProviderMalformed
}
attributes := map[string]string{}
seen := map[string]struct{}{}
stage := 0
stereo := false
for _, line := range strings.Split(strings.TrimSuffix(body, lineEnding), lineEnding) {
if line == "" {
return ErrProviderMalformed
}
if !strings.HasPrefix(line, "a=") {
if line == "sprop-parameter-sets=AAAAAU" && stage >= 3 && stage <= 4 {
stage = 5
continue
}
return ErrProviderMalformed
}
key, value, ok := strings.Cut(strings.TrimPrefix(line, "a="), ":")
if !ok || key == "" || value == "" || len(key) > 128 || len(value) > 256 {
return ErrProviderMalformed
}
if key == "fmtp" {
if stage < 3 || !validApolloSurroundParameters(value) {
return ErrProviderMalformed
}
if _, duplicate := seen["fmtp:"+value]; duplicate {
return ErrProviderMalformed
}
seen["fmtp:"+value] = struct{}{}
if value == "97 surround-params=21101" {
stereo = true
}
stage = 6
continue
}
if key == "rtpmap" && value == "98 AV1/90000" && stage >= 3 && stage <= 5 {
if _, duplicate := seen[key]; duplicate {
return ErrProviderMalformed
}
seen[key] = struct{}{}
stage = 6
continue
}
if key == "x-nv-video[0].refPicInvalidation" && value == "1" && stage == 3 {
if _, duplicate := seen[key]; duplicate {
return ErrProviderMalformed
}
seen[key] = struct{}{}
stage = 4
continue
}
if key != "x-ss-general.featureFlags" && key != "x-ss-general.encryptionSupported" && key != "x-ss-general.encryptionRequested" {
return ErrProviderMalformed
}
if _, duplicate := attributes[key]; duplicate || (key == "x-ss-general.featureFlags" && stage != 0) || (key == "x-ss-general.encryptionSupported" && stage != 1) || (key == "x-ss-general.encryptionRequested" && stage != 2) {
return ErrProviderMalformed
}
attributes[key] = value
stage++
}
featureFlags, featureFlagsOK := attributes["x-ss-general.featureFlags"]
supported, supportedOK := attributes["x-ss-general.encryptionSupported"]
requested, requestedOK := attributes["x-ss-general.encryptionRequested"]
if !featureFlagsOK || !supportedOK || !requestedOK || !stereo {
return ErrProviderMalformed
}
if _, err := strconv.ParseUint(featureFlags, 10, 32); err != nil {
return ErrProviderMalformed
}
supportedFlags, err := strconv.ParseUint(supported, 10, 32)
if err != nil || supportedFlags&apolloEncryptionAll != apolloEncryptionAll {
return ErrProviderMalformed
}
requestedFlags, err := strconv.ParseUint(requested, 10, 32)
if err != nil || requestedFlags&^supportedFlags != 0 || requestedFlags&1 == 0 {
return ErrProviderMalformed
}
return nil
}
func validApolloSurroundParameters(value string) bool {
const prefix = "97 surround-params="
if !strings.HasPrefix(value, prefix) {
return false
}
parameters := strings.TrimPrefix(value, prefix)
if len(parameters) < 5 || len(parameters) > 11 {
return false
}
channels := int(parameters[0] - '0')
streams := int(parameters[1] - '0')
coupled := int(parameters[2] - '0')
if (channels != 2 && channels != 6 && channels != 8) || len(parameters) != channels+3 || streams+coupled != channels || streams == 0 {
return false
}
used := [8]bool{}
for _, character := range parameters[3:] {
if character < '0' || int(character-'0') >= channels || used[character-'0'] {
return false
}
used[character-'0'] = true
}
return true
}
func apolloRTSPSession(message apolloRTSPMessage) (string, error) {
value, ok := message.headers["session"]
if !ok {
return "", ErrProviderMalformed
}
token, _, _ := strings.Cut(value, ";")
token = strings.TrimSpace(token)
if !validApolloRTSPToken(token) || len(token) > 256 {
return "", ErrProviderMalformed
}
return token, nil
}
func apolloRTSPMatchSession(message apolloRTSPMessage, expected string) error {
actual, err := apolloRTSPSession(message)
if err != nil || actual != expected {
return ErrProviderMalformed
}
return nil
}
func apolloRTSPServerPort(message apolloRTSPMessage) (int, error) {
transport, ok := message.headers["transport"]
if !ok {
return 0, ErrProviderMalformed
}
parts := strings.Split(transport, ";")
if len(parts) != 2 || parts[0] != "unicast" {
return 0, ErrProviderMalformed
}
key, value, ok := strings.Cut(parts[1], "=")
port, err := strconv.Atoi(value)
if !ok || key != "server_port" || err != nil || port < 1 || port > maxApolloRTSPPort {
return 0, ErrProviderMalformed
}
return port, nil
}
func apolloRTSPPingPayload(message apolloRTSPMessage) ([]byte, error) {
payload, ok := message.headers["x-ss-ping-payload"]
if !ok || len(payload) != 16 || !validApolloRTSPToken(payload) {
return nil, ErrProviderMalformed
}
return []byte(payload), nil
}
func apolloRTSPConnectData(message apolloRTSPMessage) (uint32, error) {
value, ok := message.headers["x-ss-connect-data"]
if !ok {
return 0, ErrProviderMalformed
}
if value == "" || strings.Trim(value, "0123456789") != "" {
return 0, ErrProviderMalformed
}
parsed, err := strconv.ParseUint(value, 10, 32)
if err != nil || parsed == 0 {
return 0, ErrProviderMalformed
}
return uint32(parsed), nil
}
func apolloAnnounceProfile() []byte {
return []byte("v=0\r\n" +
"o=android 0 0 IN IP4 0.0.0.0\r\n" +
"s=NVIDIA Streaming Client\r\n" +
"a=x-nv-video[0].clientViewportWd:1920\r\n" +
"a=x-nv-video[0].clientViewportHt:1080\r\n" +
"a=x-nv-video[0].maxFPS:60\r\n" +
"a=x-nv-video[0].packetSize:1024\r\n" +
"a=x-nv-video[0].videoEncoderSlicesPerFrame:1\r\n" +
"a=x-nv-video[0].maxNumReferenceFrames:0\r\n" +
"a=x-nv-vqos[0].bitStreamFormat:0\r\n" +
"a=x-nv-vqos[0].bw.maximumBitrateKbps:8000\r\n" +
"a=x-nv-vqos[0].fec.minRequiredFecPackets:2\r\n" +
"a=x-nv-vqos[0].qosTrafficType:5\r\n" +
"a=x-nv-audio.surround.numChannels:2\r\n" +
"a=x-nv-audio.surround.channelMask:3\r\n" +
"a=x-nv-audio.surround.AudioQuality:0\r\n" +
"a=x-nv-aqos.packetDuration:5\r\n" +
"a=x-nv-aqos.qosTrafficType:4\r\n" +
"a=x-nv-general.useReliableUdp:13\r\n" +
"a=x-nv-general.featureFlags:167\r\n" +
"a=x-ml-general.featureFlags:0\r\n" +
"a=x-ml-video.configuredBitrateKbps:8000\r\n" +
"a=x-ss-general.encryptionEnabled:7\r\n" +
"a=x-ss-video[0].chromaSamplingType:0\r\n" +
"a=x-ss-video[0].intraRefresh:0\r\n")
}
func validApolloRTSPToken(value string) bool {
if value == "" || len(value) > 128 {
return false
}
for _, character := range value {
if character <= 0x20 || character >= 0x7f || strings.ContainsRune("()<>@,;:\\\"/[]?={}", character) {
return false
}
}
return true
}
+279
View File
@@ -0,0 +1,279 @@
package gateway
import (
"bytes"
"encoding/binary"
)
const (
apolloVideoMaximumDataShards = 255
apolloVideoMaximumBlocks = 4
apolloVideoShardPayloadSize = apolloVideoRawPacketSize - apolloRTPHeaderSize - 4 - apolloVideoNVHeaderSize
)
type apolloVideoShard struct {
frame uint32
block uint8
lastBlock uint8
dataPackets int
parity int
index int
sequence uint16
streamIndex uint32
flags byte
payload []byte
}
type apolloVideoFECBlock struct {
dataPackets int
parity int
firstSeq uint16
streamBase uint32
haveBase bool
shards [][]byte
received []bool
count int
complete bool
}
type apolloVideoAssembler struct {
haveFrame bool
frame uint32
lastBlock uint8
blocks [apolloVideoMaximumBlocks]*apolloVideoFECBlock
}
func (a *apolloVideoAssembler) Add(shard apolloVideoShard) ([]byte, error) {
if len(shard.payload) != apolloVideoShardPayloadSize || shard.dataPackets < 1 || shard.dataPackets > apolloVideoMaximumDataShards || shard.parity < 0 || shard.dataPackets+shard.parity > 255 || shard.block > shard.lastBlock || shard.lastBlock >= apolloVideoMaximumBlocks || shard.index >= shard.dataPackets+shard.parity {
return nil, errApolloMedia
}
if !a.haveFrame || apolloFrameNewer(shard.frame, a.frame) {
*a = apolloVideoAssembler{haveFrame: true, frame: shard.frame, lastBlock: shard.lastBlock}
} else if shard.frame != a.frame {
return nil, nil
}
if shard.lastBlock != a.lastBlock {
return nil, errApolloMedia
}
block := a.blocks[shard.block]
if block == nil {
block = &apolloVideoFECBlock{
dataPackets: shard.dataPackets,
parity: shard.parity,
shards: make([][]byte, shard.dataPackets+shard.parity),
received: make([]bool, shard.dataPackets+shard.parity),
}
a.blocks[shard.block] = block
} else if block.dataPackets != shard.dataPackets || block.parity != shard.parity {
return nil, errApolloMedia
}
if shard.index < shard.dataPackets {
if shard.index == 0 && shard.flags&0x04 == 0 || shard.index == shard.dataPackets-1 && shard.flags&0x02 == 0 || shard.flags&^byte(0x07) != 0 {
return nil, errApolloMedia
}
base := shard.sequence - uint16(shard.index)
streamBase := shard.streamIndex - uint32(shard.index)
if !block.haveBase {
block.firstSeq, block.streamBase, block.haveBase = base, streamBase, true
} else if block.firstSeq != base || block.streamBase != streamBase {
return nil, errApolloMedia
}
}
if block.received[shard.index] {
if !bytes.Equal(block.shards[shard.index], shard.payload) {
return nil, errApolloMedia
}
return nil, nil
}
block.shards[shard.index] = append([]byte(nil), shard.payload...)
block.received[shard.index] = true
block.count++
if block.count < block.dataPackets {
return nil, nil
}
if !block.complete {
if err := reconstructApolloVideoBlock(block); err != nil {
return nil, err
}
block.complete = true
}
for index := uint8(0); index <= a.lastBlock; index++ {
if a.blocks[index] == nil || !a.blocks[index].complete {
return nil, nil
}
}
capacity := 0
for blockIndex := uint8(0); blockIndex <= a.lastBlock; blockIndex++ {
capacity += a.blocks[blockIndex].dataPackets * apolloVideoShardPayloadSize
}
frame := make([]byte, 0, capacity)
for blockIndex := uint8(0); blockIndex <= a.lastBlock; blockIndex++ {
for shardIndex := 0; shardIndex < a.blocks[blockIndex].dataPackets; shardIndex++ {
frame = append(frame, a.blocks[blockIndex].shards[shardIndex]...)
}
}
if len(frame) < 8 || frame[0] != 0x01 {
return nil, errApolloMedia
}
lastPayloadLength := int(binary.LittleEndian.Uint16(frame[4:6]))
end := len(frame) - apolloVideoShardPayloadSize + lastPayloadLength
if lastPayloadLength == 0 || lastPayloadLength > apolloVideoShardPayloadSize || end <= 8 || end > len(frame) {
return nil, errApolloMedia
}
*a = apolloVideoAssembler{}
return append([]byte(nil), frame[8:end]...), nil
}
func apolloFrameNewer(first, second uint32) bool {
return int32(first-second) > 0
}
func reconstructApolloVideoBlock(block *apolloVideoFECBlock) error {
if block == nil || block.count < block.dataPackets {
return errApolloMedia
}
missing := false
for index := 0; index < block.dataPackets; index++ {
if !block.received[index] {
missing = true
block.shards[index] = make([]byte, apolloVideoShardPayloadSize)
}
}
if !missing {
return nil
}
selectedRows := make([][]byte, 0, block.dataPackets)
selectedShards := make([][]byte, 0, block.dataPackets)
for index, received := range block.received {
if !received {
continue
}
selectedRows = append(selectedRows, apolloVideoFECRow(index, block.dataPackets, block.parity))
selectedShards = append(selectedShards, block.shards[index])
if len(selectedRows) == block.dataPackets {
break
}
}
if len(selectedRows) != block.dataPackets {
return errApolloMedia
}
inverse, ok := apolloGFInvert(selectedRows)
if !ok {
return errApolloMedia
}
for index := 0; index < block.dataPackets; index++ {
if block.received[index] {
continue
}
for source, coefficient := range inverse[index] {
apolloGFAXPY(block.shards[index], selectedShards[source], coefficient)
}
block.received[index] = true
}
return nil
}
func apolloVideoFECRow(index, dataPackets, parity int) []byte {
row := make([]byte, dataPackets)
if index < dataPackets {
row[index] = 1
return row
}
parityIndex := index - dataPackets
for dataIndex := range row {
row[dataIndex] = apolloGFInverse(byte((parity + dataIndex) ^ parityIndex))
}
return row
}
func apolloGFInvert(matrix [][]byte) ([][]byte, bool) {
size := len(matrix)
if size == 0 {
return nil, false
}
work := make([][]byte, size)
inverse := make([][]byte, size)
for row := range matrix {
if len(matrix[row]) != size {
return nil, false
}
work[row] = append([]byte(nil), matrix[row]...)
inverse[row] = make([]byte, size)
inverse[row][row] = 1
}
for column := 0; column < size; column++ {
pivot := column
for pivot < size && work[pivot][column] == 0 {
pivot++
}
if pivot == size {
return nil, false
}
work[column], work[pivot] = work[pivot], work[column]
inverse[column], inverse[pivot] = inverse[pivot], inverse[column]
factor := apolloGFInverse(work[column][column])
for index := column; index < size; index++ {
work[column][index] = apolloGFMultiply(work[column][index], factor)
}
for index := range inverse[column] {
inverse[column][index] = apolloGFMultiply(inverse[column][index], factor)
}
for row := 0; row < size; row++ {
if row == column || work[row][column] == 0 {
continue
}
factor = work[row][column]
for index := column; index < size; index++ {
work[row][index] ^= apolloGFMultiply(work[column][index], factor)
}
for index := range inverse[row] {
inverse[row][index] ^= apolloGFMultiply(inverse[column][index], factor)
}
}
}
return inverse, true
}
func apolloGFAXPY(destination, source []byte, coefficient byte) {
if coefficient == 0 {
return
}
for index := range destination {
destination[index] ^= apolloGFMultiply(source[index], coefficient)
}
}
func apolloGFInverse(value byte) byte {
if value == 0 {
return 0
}
return apolloGFPow(value, 254)
}
func apolloGFPow(value byte, exponent uint8) byte {
result := byte(1)
for exponent != 0 {
if exponent&1 != 0 {
result = apolloGFMultiply(result, value)
}
value = apolloGFMultiply(value, value)
exponent >>= 1
}
return result
}
func apolloGFMultiply(first, second byte) byte {
var product byte
for second != 0 {
if second&1 != 0 {
product ^= first
}
high := first & 0x80
first <<= 1
if high != 0 {
first ^= 0x1d
}
second >>= 1
}
return product
}
+166
View File
@@ -0,0 +1,166 @@
package gateway
import (
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"errors"
"sync"
"time"
"unicode/utf8"
protocol "git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol/gen/go/protocol"
)
var (
ErrClipboardDenied = errors.New("clipboard policy denied")
ErrClipboardRate = errors.New("clipboard rate limited")
)
const clipboardRetention = time.Minute
type clipboardRecord struct {
digest [sha256.Size]byte
direction string
at time.Time
}
type clipboardGate struct {
policy protocol.ClipboardPolicy
now func() time.Time
mu sync.Mutex
updates []time.Time
seen map[string]clipboardRecord
}
func newClipboardGate(policy protocol.ClipboardPolicy, now func() time.Time) (*clipboardGate, error) {
if err := policy.Validate(); err != nil || now == nil {
return nil, ErrProviderMalformed
}
return &clipboardGate{policy: policy, now: now, seen: make(map[string]clipboardRecord, policy.MaxUpdatesPerMinute)}, nil
}
// ValidateGatewayClipboard applies the authenticated Server policy before any
// clipboard value can reach a provider or Verse client.
func ValidateGatewayClipboard(policy protocol.ClipboardPolicy, value protocol.GatewayClipboardText) error {
if err := policy.Validate(); err != nil || value.Validate() != nil || !utf8.ValidString(value.Text) {
return ErrProviderMalformed
}
if len(value.Text) > int(policy.MaxTextBytes) || len(value.LoopToken) < 16 || len(value.LoopToken) > 128 {
return ErrProviderMalformed
}
if _, err := base64.RawURLEncoding.DecodeString(value.LoopToken); err != nil {
return ErrProviderMalformed
}
switch value.Direction {
case "client_to_provider":
if !policy.ClientToProviderEnabled {
return ErrClipboardDenied
}
case "provider_to_client":
if !policy.ProviderToClientEnabled {
return ErrClipboardDenied
}
default:
return ErrProviderMalformed
}
return nil
}
func (g *clipboardGate) fromClient(value protocol.GatewayClipboardText) (bool, error) {
if g == nil {
return false, ErrClipboardDenied
}
if err := ValidateGatewayClipboard(g.policy, value); err != nil {
return false, err
}
now := g.now()
digest := sha256.Sum256([]byte(value.Text))
g.mu.Lock()
defer g.mu.Unlock()
g.pruneLocked(now)
if record, ok := g.seen[value.LoopToken]; ok {
if record.direction == "provider_to_client" && record.digest == digest {
return true, nil
}
return false, ErrClipboardDenied
}
if !g.allowUpdateLocked(now) {
return false, ErrClipboardRate
}
g.seen[value.LoopToken] = clipboardRecord{digest: digest, direction: value.Direction, at: now}
return false, nil
}
func (g *clipboardGate) fromProvider(text string) (protocol.GatewayClipboardText, bool, error) {
if g == nil || !g.policy.ProviderToClientEnabled {
return protocol.GatewayClipboardText{}, false, ErrClipboardDenied
}
if !utf8.ValidString(text) || len(text) > int(g.policy.MaxTextBytes) {
return protocol.GatewayClipboardText{}, false, ErrProviderMalformed
}
now := g.now()
digest := sha256.Sum256([]byte(text))
g.mu.Lock()
defer g.mu.Unlock()
g.pruneLocked(now)
for _, record := range g.seen {
if record.digest == digest {
return protocol.GatewayClipboardText{}, true, nil
}
}
if !g.allowUpdateLocked(now) {
return protocol.GatewayClipboardText{}, false, ErrClipboardRate
}
for attempts := 0; attempts < 3; attempts++ {
var raw [24]byte
if _, err := rand.Read(raw[:]); err != nil {
return protocol.GatewayClipboardText{}, false, err
}
token := base64.RawURLEncoding.EncodeToString(raw[:])
if _, exists := g.seen[token]; exists {
continue
}
value := protocol.GatewayClipboardText{Direction: "provider_to_client", Text: text, Encoding: "utf-8", LoopToken: token}
g.seen[token] = clipboardRecord{digest: digest, direction: value.Direction, at: now}
return value, false, nil
}
return protocol.GatewayClipboardText{}, false, ErrClipboardDenied
}
func (g *clipboardGate) retractClient(value protocol.GatewayClipboardText) {
if g == nil {
return
}
digest := sha256.Sum256([]byte(value.Text))
g.mu.Lock()
defer g.mu.Unlock()
if record, ok := g.seen[value.LoopToken]; ok && record.direction == "client_to_provider" && record.digest == digest {
delete(g.seen, value.LoopToken)
}
}
func (g *clipboardGate) pruneLocked(now time.Time) {
minimum := now.Add(-clipboardRetention)
index := 0
for _, update := range g.updates {
if update.After(minimum) {
g.updates[index] = update
index++
}
}
g.updates = g.updates[:index]
for token, record := range g.seen {
if !record.at.After(minimum) {
delete(g.seen, token)
}
}
}
func (g *clipboardGate) allowUpdateLocked(now time.Time) bool {
if len(g.updates) >= int(g.policy.MaxUpdatesPerMinute) {
return false
}
g.updates = append(g.updates, now)
return true
}
+75
View File
@@ -0,0 +1,75 @@
package gateway
import (
"errors"
"strings"
"testing"
"time"
protocol "git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol/gen/go/protocol"
)
func TestValidateGatewayClipboardEnforcesServerOwnedPolicy(t *testing.T) {
policy := protocol.ClipboardPolicy{ClientToProviderEnabled: true, ProviderToClientEnabled: true, MaxTextBytes: 5, MaxUpdatesPerMinute: 2}
valid := protocol.GatewayClipboardText{Direction: "client_to_provider", Text: "hello", Encoding: "utf-8", LoopToken: "abcdefghijklmnop"}
if err := ValidateGatewayClipboard(policy, valid); err != nil {
t.Fatalf("ValidateGatewayClipboard() valid text = %v", err)
}
if err := ValidateGatewayClipboard(policy, protocol.GatewayClipboardText{Direction: "client_to_provider", Text: "hello", Encoding: "utf-8", LoopToken: "not base64url!"}); err == nil {
t.Fatal("ValidateGatewayClipboard() accepted malformed loop token")
}
if err := ValidateGatewayClipboard(policy, protocol.GatewayClipboardText{Direction: "client_to_provider", Text: strings.Repeat("x", 6), Encoding: "utf-8", LoopToken: "abcdefghijklmnop"}); err == nil {
t.Fatal("ValidateGatewayClipboard() accepted oversized text")
}
disabled := policy
disabled.ClientToProviderEnabled = false
if err := ValidateGatewayClipboard(disabled, valid); err == nil {
t.Fatal("ValidateGatewayClipboard() accepted disabled direction")
}
}
func TestClipboardGateSuppressesReflectionsAndBoundsRate(t *testing.T) {
now := time.Date(2026, time.January, 1, 0, 0, 0, 0, time.UTC)
policy := protocol.ClipboardPolicy{ClientToProviderEnabled: true, ProviderToClientEnabled: true, MaxTextBytes: 64, MaxUpdatesPerMinute: 2}
gate, err := newClipboardGate(policy, func() time.Time { return now })
if err != nil {
t.Fatal(err)
}
client := protocol.GatewayClipboardText{Direction: "client_to_provider", Text: "client", Encoding: "utf-8", LoopToken: "abcdefghijklmnop"}
if suppress, err := gate.fromClient(client); err != nil || suppress {
t.Fatalf("fromClient() = suppress %t, err %v", suppress, err)
}
if _, suppress, err := gate.fromProvider("client"); err != nil || !suppress {
t.Fatalf("fromProvider() reflection = suppress %t, err %v", suppress, err)
}
host, suppress, err := gate.fromProvider("host")
if err != nil || suppress || host.Direction != "provider_to_client" {
t.Fatalf("fromProvider() host = %#v, suppress %t, err %v", host, suppress, err)
}
if suppress, err := gate.fromClient(host); err != nil || !suppress {
t.Fatalf("fromClient() host reflection = suppress %t, err %v", suppress, err)
}
if _, err := gate.fromClient(protocol.GatewayClipboardText{Direction: "client_to_provider", Text: "third", Encoding: "utf-8", LoopToken: "qrstuvwxyzABCDEF"}); !errors.Is(err, ErrClipboardRate) {
t.Fatalf("fromClient() rate error = %v, want ErrClipboardRate", err)
}
now = now.Add(time.Minute)
if suppress, err := gate.fromClient(protocol.GatewayClipboardText{Direction: "client_to_provider", Text: "after-window", Encoding: "utf-8", LoopToken: "0123456789abcdef"}); err != nil || suppress {
t.Fatalf("fromClient() after window = suppress %t, err %v", suppress, err)
}
}
func TestClipboardGatePermitsRetryAfterProviderWriteFailure(t *testing.T) {
policy := protocol.ClipboardPolicy{ClientToProviderEnabled: true, MaxTextBytes: 64, MaxUpdatesPerMinute: 2}
gate, err := newClipboardGate(policy, time.Now)
if err != nil {
t.Fatal(err)
}
value := protocol.GatewayClipboardText{Direction: "client_to_provider", Text: "retry", Encoding: "utf-8", LoopToken: "abcdefghijklmnop"}
if suppress, err := gate.fromClient(value); err != nil || suppress {
t.Fatalf("fromClient() = suppress %t, err %v", suppress, err)
}
gate.retractClient(value)
if suppress, err := gate.fromClient(value); err != nil || suppress {
t.Fatalf("fromClient() retry = suppress %t, err %v", suppress, err)
}
}
+10
View File
@@ -97,6 +97,15 @@ func (c *ControlPlaneClient) ReportProviderState(ctx context.Context, state prot
return err return err
} }
func (c *ControlPlaneClient) ReportClipboardAudit(ctx context.Context, audit protocol.GatewayClipboardAudit) error {
payload, err := protocol.EncodeGatewayClipboardAudit(audit)
if err != nil {
return err
}
_, err = c.post(ctx, "/api/v1/gateway/clipboard-audit", payload)
return err
}
func (c *ControlPlaneClient) post(ctx context.Context, path string, payload []byte) ([]byte, error) { func (c *ControlPlaneClient) post(ctx context.Context, path string, payload []byte) ([]byte, error) {
if c == nil || c.HTTPClient == nil || c.BaseURL == "" { if c == nil || c.HTTPClient == nil || c.BaseURL == "" {
return nil, errors.New("control-plane client is not configured") return nil, errors.New("control-plane client is not configured")
@@ -129,3 +138,4 @@ func (c *ControlPlaneClient) post(ctx context.Context, path string, payload []by
} }
var _ Admission = (*ControlPlaneClient)(nil) var _ Admission = (*ControlPlaneClient)(nil)
var _ ClipboardAuditReporter = (*ControlPlaneClient)(nil)
+94
View File
@@ -0,0 +1,94 @@
package gateway
import (
"sort"
"testing"
"time"
)
type syntheticPacerDelivery struct {
at time.Time
flow string
bytes int64
}
func TestFairPacerEightFlowSharesAndCapacitySteps(t *testing.T) {
start := time.Date(2026, time.January, 1, 0, 0, 0, 0, time.UTC)
flows := []string{"one", "two", "three", "four", "five", "six", "seven", "eight"}
pacer := newFairPacer(8000)
next := make(map[string]time.Time, len(flows))
baseline := runSyntheticPacer(pacer, start, start.Add(60*time.Second), flows, next)
assertSyntheticFairness(t, baseline, flows)
assertSyntheticCap(t, baseline, 1_000_000)
pacer.setKbps(6000)
quarter := runSyntheticPacer(pacer, start.Add(60*time.Second), start.Add(70*time.Second), flows, next)
assertSyntheticFairness(t, quarter, flows)
assertSyntheticCap(t, quarter, 750_000)
pacer.setKbps(4000)
half := runSyntheticPacer(pacer, start.Add(70*time.Second), start.Add(80*time.Second), flows, next)
assertSyntheticFairness(t, half, flows)
assertSyntheticCap(t, half, 500_000)
}
func runSyntheticPacer(pacer *fairPacer, start, end time.Time, flows []string, next map[string]time.Time) []syntheticPacerDelivery {
const packetBytes = 1000
for _, flow := range flows {
if next[flow].IsZero() {
next[flow] = pacer.reserveAt(start, flow, packetBytes)
}
}
var deliveries []syntheticPacerDelivery
for {
flow := ""
at := end.Add(time.Nanosecond)
for _, candidate := range flows {
if next[candidate].Before(at) {
flow, at = candidate, next[candidate]
}
}
if at.After(end) {
return deliveries
}
deliveries = append(deliveries, syntheticPacerDelivery{at: at, flow: flow, bytes: packetBytes})
next[flow] = pacer.reserveAt(at, flow, packetBytes)
}
}
func assertSyntheticFairness(t *testing.T, deliveries []syntheticPacerDelivery, flows []string) {
t.Helper()
counts := make(map[string]int64, len(flows))
for _, delivery := range deliveries {
counts[delivery.flow] += delivery.bytes
}
total := int64(0)
for _, flow := range flows {
total += counts[flow]
}
target := total / int64(len(flows))
for _, flow := range flows {
delta := counts[flow] - target
if delta < 0 {
delta = -delta
}
if target == 0 || float64(delta)/float64(target) > 0.10 {
t.Fatalf("flow %s share=%d target=%d", flow, counts[flow], target)
}
}
}
func assertSyntheticCap(t *testing.T, deliveries []syntheticPacerDelivery, bytesPerSecond int64) {
t.Helper()
sort.Slice(deliveries, func(first, second int) bool { return deliveries[first].at.Before(deliveries[second].at) })
for first, total, last := 0, int64(0), 0; first < len(deliveries); first++ {
for last < len(deliveries) && deliveries[last].at.Sub(deliveries[first].at) <= 5*time.Second {
total += deliveries[last].bytes
last++
}
if total > bytesPerSecond*5*105/100 {
t.Fatalf("five-second egress=%d exceeds cap=%d", total, bytesPerSecond*5)
}
total -= deliveries[first].bytes
}
}
+149
View File
@@ -0,0 +1,149 @@
package gateway
import "encoding/binary"
type FeedbackKind uint8
const (
FeedbackIDR FeedbackKind = iota + 1
FeedbackFEC
)
const (
gatewayFeedbackHeaderSize = 8
gatewayFeedbackClient = 0
gatewayFeedbackGateway = 1
gatewayFeedbackIDR = 1
gatewayFeedbackFEC = 2
gatewayFeedbackTerminated = 0x10
gatewayFeedbackRumble = 0x11
gatewayFeedbackHDR = 0x12
)
type gatewayFeedbackMessage struct {
direction byte
kind byte
payload []byte
}
func EncodeProviderEvent(event ProviderEvent) ([]byte, error) {
switch event.Kind {
case ProviderEventTerminated:
if len(event.Payload) != 4 {
return nil, ErrProviderMalformed
}
return encodeGatewayFeedback(gatewayFeedbackGateway, gatewayFeedbackTerminated, event.Payload)
case ProviderEventRumble:
if len(event.Payload) != 5 || event.Payload[0] > 15 {
return nil, ErrProviderMalformed
}
return encodeGatewayFeedback(gatewayFeedbackGateway, gatewayFeedbackRumble, event.Payload)
case ProviderEventHDR:
if len(event.Payload) != 1 || event.Payload[0] > 1 {
return nil, ErrProviderMalformed
}
return encodeGatewayFeedback(gatewayFeedbackGateway, gatewayFeedbackHDR, event.Payload)
default:
return nil, ErrProviderMalformed
}
}
func EncodeClientFeedback(feedback Feedback) ([]byte, error) {
var kind byte
switch feedback.Kind {
case FeedbackIDR:
kind = gatewayFeedbackIDR
if len(feedback.Payload) != 0 {
return nil, ErrProviderMalformed
}
case FeedbackFEC:
kind = gatewayFeedbackFEC
if !validGatewayFECStatus(feedback.Payload) {
return nil, ErrProviderMalformed
}
default:
return nil, ErrProviderMalformed
}
return encodeGatewayFeedback(gatewayFeedbackClient, kind, feedback.Payload)
}
func DecodeClientFeedback(data []byte) (Feedback, error) {
message, err := decodeGatewayFeedback(data)
if err != nil || message.direction != gatewayFeedbackClient {
return Feedback{}, ErrProviderMalformed
}
switch message.kind {
case gatewayFeedbackIDR:
if len(message.payload) != 0 {
return Feedback{}, ErrProviderMalformed
}
return Feedback{Kind: FeedbackIDR}, nil
case gatewayFeedbackFEC:
if !validGatewayFECStatus(message.payload) {
return Feedback{}, ErrProviderMalformed
}
return Feedback{Kind: FeedbackFEC, Payload: message.payload}, nil
default:
return Feedback{}, ErrProviderMalformed
}
}
func validGatewayFECStatus(payload []byte) bool {
if len(payload) != 21 {
return false
}
totalData := binary.BigEndian.Uint16(payload[10:12])
totalParity := binary.BigEndian.Uint16(payload[12:14])
receivedData := binary.BigEndian.Uint16(payload[14:16])
receivedParity := binary.BigEndian.Uint16(payload[16:18])
return totalData > 0 && receivedData <= totalData && receivedParity <= totalParity && payload[18] <= 100 && payload[20] > 0 && payload[19] < payload[20]
}
func DecodeProviderEvent(data []byte) (ProviderEvent, error) {
message, err := decodeGatewayFeedback(data)
if err != nil || message.direction != gatewayFeedbackGateway {
return ProviderEvent{}, ErrProviderMalformed
}
switch message.kind {
case gatewayFeedbackTerminated:
if len(message.payload) != 4 {
return ProviderEvent{}, ErrProviderMalformed
}
return ProviderEvent{Kind: ProviderEventTerminated, Payload: message.payload}, nil
case gatewayFeedbackRumble:
if len(message.payload) != 5 || message.payload[0] > 15 {
return ProviderEvent{}, ErrProviderMalformed
}
return ProviderEvent{Kind: ProviderEventRumble, Payload: message.payload}, nil
case gatewayFeedbackHDR:
if len(message.payload) != 1 || message.payload[0] > 1 {
return ProviderEvent{}, ErrProviderMalformed
}
return ProviderEvent{Kind: ProviderEventHDR, Payload: message.payload}, nil
default:
return ProviderEvent{}, ErrProviderMalformed
}
}
func encodeGatewayFeedback(direction, kind byte, payload []byte) ([]byte, error) {
if len(payload) > 1016 {
return nil, ErrProviderMalformed
}
encoded := make([]byte, gatewayFeedbackHeaderSize+len(payload))
copy(encoded, "VGF1")
encoded[4], encoded[5] = direction, kind
binary.BigEndian.PutUint16(encoded[6:8], uint16(len(payload)))
copy(encoded[8:], payload)
return encoded, nil
}
func decodeGatewayFeedback(data []byte) (gatewayFeedbackMessage, error) {
if len(data) < gatewayFeedbackHeaderSize || len(data) > 1024 || string(data[:4]) != "VGF1" || len(data) != gatewayFeedbackHeaderSize+int(binary.BigEndian.Uint16(data[6:8])) {
return gatewayFeedbackMessage{}, ErrProviderMalformed
}
message := gatewayFeedbackMessage{direction: data[4], kind: data[5], payload: append([]byte(nil), data[8:]...)}
if message.direction != gatewayFeedbackClient && message.direction != gatewayFeedbackGateway {
return gatewayFeedbackMessage{}, ErrProviderMalformed
}
return message, nil
}
+89 -44
View File
@@ -55,24 +55,49 @@ func FuzzDecodeFrame(f *testing.F) {
}) })
} }
func FuzzDecodeControlPacket(f *testing.F) {
seed, _ := EncodeControlPacket(ControlPacket{Kind: 1, Sequence: 2, Payload: []byte("fixture")})
f.Add(seed)
f.Add([]byte("APC1"))
f.Fuzz(func(t *testing.T, data []byte) {
_, _ = DecodeControlPacket(data)
})
}
func FuzzDecodeInputEvent(f *testing.F) { func FuzzDecodeInputEvent(f *testing.F) {
seed, _ := EncodeInputEvent(InputEvent{Sequence: 1, Device: "keyboard", Code: 7, Pressed: true}) seed, _ := EncodeInputEvent(InputEvent{Sequence: 1, Device: "keyboard", Code: 7, Pressed: true})
f.Add(seed) f.Add(seed)
f.Add([]byte("INP1")) f.Add([]byte("VGI1"))
f.Fuzz(func(t *testing.T, data []byte) { f.Fuzz(func(t *testing.T, data []byte) {
_, _ = DecodeInputEvent(data) _, _ = 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 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)
}
}
func TestCapabilityIntersectionAndBoundedQueue(t *testing.T) { func TestCapabilityIntersectionAndBoundedQueue(t *testing.T) {
capabilities := DefaultCapabilities() capabilities := DefaultCapabilities()
if _, err := IntersectCapabilities(capabilities, capabilities); err != nil { if _, err := IntersectCapabilities(capabilities, capabilities); err != nil {
@@ -199,39 +224,6 @@ func TestProviderTimeoutAndBoundedInput(t *testing.T) {
if _, err := EncodeInputEvent(InputEvent{Device: strings.Repeat("d", 65)}); !errors.Is(err, ErrInputMalformed) { if _, err := EncodeInputEvent(InputEvent{Device: strings.Repeat("d", 65)}); !errors.Is(err, ErrInputMalformed) {
t.Fatalf("oversized input accepted: %v", err) t.Fatalf("oversized input accepted: %v", err)
} }
if _, err := DecodeControlPacket([]byte("APC1")); !errors.Is(err, ErrProviderMalformed) {
t.Fatalf("truncated control accepted: %v", err)
}
}
func TestNativeApolloEncodedRelay(t *testing.T) {
provider, peer := net.Pipe()
session := newNativeApolloSession(provider, "session-native")
go session.readMedia()
go func() {
_, _ = peer.Write([]byte{'$', 0, 0, 3, 1, 2, 3})
_, _ = peer.Write([]byte{'$', 1, 0, 2, 4, 5})
}()
select {
case payload := <-session.Video():
if string(payload) != string([]byte{1, 2, 3}) {
t.Fatalf("video payload changed: %x", payload)
}
case <-time.After(time.Second):
t.Fatal("video payload not relayed")
}
select {
case payload := <-session.Audio():
if string(payload) != string([]byte{4, 5}) {
t.Fatalf("audio payload changed: %x", payload)
}
case <-time.After(time.Second):
t.Fatal("audio payload not relayed")
}
terminateCtx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond)
defer cancel()
_ = session.Terminate(terminateCtx)
_ = peer.Close()
} }
func TestAdmissionQUICMTLSRelayAndCleanup(t *testing.T) { func TestAdmissionQUICMTLSRelayAndCleanup(t *testing.T) {
@@ -240,7 +232,7 @@ func TestAdmissionQUICMTLSRelayAndCleanup(t *testing.T) {
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()} 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{})} admission := &oneTimeAdmission{authority: authority, released: make(chan struct{})}
reporter := &recordingProviderStateReporter{} 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, Provider: fake}) 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 { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -263,6 +255,44 @@ func TestAdmissionQUICMTLSRelayAndCleanup(t *testing.T) {
t.Fatalf("unexpected media channel %d", frame.Channel) 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 { if err := client.SendInput(InputEvent{Sequence: 1, Device: "keyboard", Code: 7, Pressed: true}); err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -331,6 +361,7 @@ type oneTimeAdmission struct {
type recordingProviderStateReporter struct { type recordingProviderStateReporter struct {
mu sync.Mutex mu sync.Mutex
states []protocol.ProviderState states []protocol.ProviderState
audits []protocol.GatewayClipboardAudit
} }
func (r *recordingProviderStateReporter) ReportProviderState(_ context.Context, state protocol.ProviderState) error { func (r *recordingProviderStateReporter) ReportProviderState(_ context.Context, state protocol.ProviderState) error {
@@ -346,6 +377,19 @@ func (r *recordingProviderStateReporter) States() []protocol.ProviderState {
return append([]protocol.ProviderState(nil), r.states...) 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) { func (a *oneTimeAdmission) Admit(context.Context, protocol.TunnelAdmissionRequest) (protocol.SessionAuthority, error) {
if !a.used.CompareAndSwap(false, true) { if !a.used.CompareAndSwap(false, true) {
return protocol.SessionAuthority{}, ErrAdmissionRejected return protocol.SessionAuthority{}, ErrAdmissionRejected
@@ -364,6 +408,7 @@ func (a *oneTimeAdmission) ProviderWork(_ context.Context, authority protocol.Se
PolicyVersionID: "policy-1", ApplicationID: "1", ClientID: "paired-client", ManagementHost: "apollo.test", ManagementPort: 47990, PolicyVersionID: "policy-1", ApplicationID: "1", ClientID: "paired-client", ManagementHost: "apollo.test", ManagementPort: 47990,
StreamHost: "apollo.test", StreamPort: 47984, ClientCertificatePem: "certificate", StreamHost: "apollo.test", StreamPort: 47984, ClientCertificatePem: "certificate",
ClientPrivateKeyPem: "private-key", ServerCertificatePem: "server-certificate", ClientPrivateKeyPem: "private-key", ServerCertificatePem: "server-certificate",
ClipboardPolicy: protocol.ClipboardPolicy{ClientToProviderEnabled: true, ProviderToClientEnabled: true, MaxTextBytes: 65536, MaxUpdatesPerMinute: 30},
}, nil }, nil
} }
+109 -15
View File
@@ -3,36 +3,130 @@ package gateway
import ( import (
"encoding/binary" "encoding/binary"
"errors" "errors"
"unicode/utf8"
) )
var ErrInputMalformed = errors.New("input event malformed") var ErrInputMalformed = errors.New("input event malformed")
const (
gatewayInputHeaderSize = 6
gatewayInputKeyboard = 1
gatewayInputMouse = 2
gatewayInputRelative = 3
gatewayInputUTF8 = 4
gatewayInputController = 5
)
func EncodeInputEvent(event InputEvent) ([]byte, error) { func EncodeInputEvent(event InputEvent) ([]byte, error) {
if len(event.Device) == 0 || len(event.Device) > 64 || len(event.Payload) > 1024 { switch event.Device {
case "keyboard":
if event.Code < 1 || event.Code > 0xffff || len(event.Payload) > 1 {
return nil, ErrInputMalformed return nil, ErrInputMalformed
} }
encoded := make([]byte, 16+len(event.Device)+len(event.Payload)) encoded := make([]byte, gatewayInputHeaderSize+4)
copy(encoded[:4], "INP1") copy(encoded, "VGI1")
binary.BigEndian.PutUint32(encoded[4:8], event.Sequence) encoded[4], encoded[5], encoded[6] = gatewayInputKeyboard, 4, 0
binary.BigEndian.PutUint32(encoded[8:12], uint32(event.Code))
if event.Pressed { if event.Pressed {
encoded[12] = 1 encoded[6] = 1
} }
encoded[13] = byte(len(event.Device)) if len(event.Payload) == 1 {
binary.BigEndian.PutUint16(encoded[14:16], uint16(len(event.Payload))) encoded[7] = event.Payload[0]
copy(encoded[16:16+len(event.Device)], event.Device) }
copy(encoded[16+len(event.Device):], event.Payload) binary.BigEndian.PutUint16(encoded[8:10], uint16(event.Code))
return encoded, nil return encoded, nil
case "mouse-button":
if event.Code < 1 || event.Code > 5 || len(event.Payload) != 0 {
return nil, ErrInputMalformed
}
encoded := make([]byte, gatewayInputHeaderSize+3)
copy(encoded, "VGI1")
encoded[4], encoded[5], encoded[7] = gatewayInputMouse, 3, byte(event.Code)
if event.Pressed {
encoded[6] = 1
}
return encoded, nil
case "mouse-relative":
if event.Pressed || event.Code != 0 || len(event.Payload) != 4 {
return nil, ErrInputMalformed
}
encoded := make([]byte, gatewayInputHeaderSize+4)
copy(encoded, "VGI1")
encoded[4], encoded[5] = gatewayInputRelative, 4
copy(encoded[6:], event.Payload)
return encoded, nil
case "utf8":
if event.Pressed || event.Code != 0 || len(event.Payload) == 0 || len(event.Payload) > utf8.UTFMax || !utf8.Valid(event.Payload) || utf8.RuneCount(event.Payload) != 1 {
return nil, ErrInputMalformed
}
encoded := make([]byte, gatewayInputHeaderSize+len(event.Payload))
copy(encoded, "VGI1")
encoded[4], encoded[5] = gatewayInputUTF8, byte(len(event.Payload))
copy(encoded[6:], event.Payload)
return encoded, nil
case "controller":
if event.Code < 0 || event.Code > 15 || len(event.Payload) != 16 {
return nil, ErrInputMalformed
}
active := binary.BigEndian.Uint16(event.Payload[:2])
if (!event.Pressed && anyNonzero(event.Payload)) || (event.Pressed && active == 0) {
return nil, ErrInputMalformed
}
encoded := make([]byte, gatewayInputHeaderSize+17)
copy(encoded, "VGI1")
encoded[4], encoded[5], encoded[6] = gatewayInputController, 17, byte(event.Code)
copy(encoded[7:], event.Payload)
return encoded, nil
default:
return nil, ErrInputMalformed
}
} }
func DecodeInputEvent(data []byte) (InputEvent, error) { func DecodeInputEvent(data []byte) (InputEvent, error) {
if len(data) < 16 || len(data) > 1179 || string(data[:4]) != "INP1" || (data[12] != 0 && data[12] != 1) { if len(data) < gatewayInputHeaderSize || len(data) > 1179 || string(data[:4]) != "VGI1" || len(data) != gatewayInputHeaderSize+int(data[5]) {
return InputEvent{}, ErrInputMalformed return InputEvent{}, ErrInputMalformed
} }
deviceLength := int(data[13]) kind, body := data[4], data[gatewayInputHeaderSize:]
payloadLength := int(binary.BigEndian.Uint16(data[14:16])) switch kind {
if deviceLength == 0 || deviceLength > 64 || payloadLength > 1024 || len(data) != 16+deviceLength+payloadLength { case gatewayInputKeyboard:
if len(body) != 4 || body[0] > 1 || binary.BigEndian.Uint16(body[2:4]) == 0 {
return InputEvent{}, ErrInputMalformed return InputEvent{}, ErrInputMalformed
} }
return InputEvent{Sequence: binary.BigEndian.Uint32(data[4:8]), Code: int32(binary.BigEndian.Uint32(data[8:12])), Pressed: data[12] == 1, Device: string(data[16 : 16+deviceLength]), Payload: append([]byte(nil), data[16+deviceLength:]...)}, nil return InputEvent{Device: "keyboard", Code: int32(binary.BigEndian.Uint16(body[2:4])), Pressed: body[0] == 1, Payload: []byte{body[1]}}, nil
case gatewayInputMouse:
if len(body) != 3 || body[0] > 1 || body[1] < 1 || body[1] > 5 || body[2] != 0 {
return InputEvent{}, ErrInputMalformed
}
return InputEvent{Device: "mouse-button", Code: int32(body[1]), Pressed: body[0] == 1}, nil
case gatewayInputRelative:
if len(body) != 4 {
return InputEvent{}, ErrInputMalformed
}
return InputEvent{Device: "mouse-relative", Payload: append([]byte(nil), body...)}, nil
case gatewayInputUTF8:
if len(body) == 0 || len(body) > utf8.UTFMax || !utf8.Valid(body) || utf8.RuneCount(body) != 1 {
return InputEvent{}, ErrInputMalformed
}
return InputEvent{Device: "utf8", Payload: append([]byte(nil), body...)}, nil
case gatewayInputController:
if len(body) != 17 || body[0] > 15 {
return InputEvent{}, ErrInputMalformed
}
payload := append([]byte(nil), body[1:]...)
active := binary.BigEndian.Uint16(payload[:2])
if active == 0 && anyNonzero(payload[2:]) {
return InputEvent{}, ErrInputMalformed
}
return InputEvent{Device: "controller", Code: int32(body[0]), Pressed: active != 0, Payload: payload}, nil
default:
return InputEvent{}, ErrInputMalformed
}
}
func anyNonzero(data []byte) bool {
for _, value := range data {
if value != 0 {
return true
}
}
return false
} }
+72 -37
View File
@@ -2,7 +2,6 @@ package gateway
import ( import (
"context" "context"
"encoding/binary"
"encoding/xml" "encoding/xml"
"errors" "errors"
"fmt" "fmt"
@@ -159,36 +158,6 @@ func ParseRTSPResponse(data []byte) (RTSPResponse, error) {
return response, nil return response, nil
} }
type ControlPacket struct {
Kind byte
Sequence uint32
Payload []byte
}
func EncodeControlPacket(packet ControlPacket) ([]byte, error) {
if len(packet.Payload) > 4096 {
return nil, ErrProviderMalformed
}
encoded := make([]byte, 11+len(packet.Payload))
copy(encoded[:4], "APC1")
encoded[4] = packet.Kind
binary.BigEndian.PutUint32(encoded[5:9], packet.Sequence)
binary.BigEndian.PutUint16(encoded[9:11], uint16(len(packet.Payload)))
copy(encoded[11:], packet.Payload)
return encoded, nil
}
func DecodeControlPacket(data []byte) (ControlPacket, error) {
if len(data) < 11 || len(data) > 4107 || string(data[:4]) != "APC1" {
return ControlPacket{}, ErrProviderMalformed
}
length := int(binary.BigEndian.Uint16(data[9:11]))
if length > 4096 || len(data) != 11+length {
return ControlPacket{}, ErrProviderMalformed
}
return ControlPacket{Kind: data[4], Sequence: binary.BigEndian.Uint32(data[5:9]), Payload: append([]byte(nil), data[11:]...)}, nil
}
type LaunchRequest struct { type LaunchRequest struct {
SessionID string SessionID string
Capabilities protocol.CapabilityProfile Capabilities protocol.CapabilityProfile
@@ -207,9 +176,35 @@ type InputEvent struct {
type Feedback struct { type Feedback struct {
Sequence uint32 Sequence uint32
Kind FeedbackKind
Payload []byte Payload []byte
} }
type ProviderEventKind uint8
const (
ProviderEventTerminated ProviderEventKind = iota + 1
ProviderEventRumble
ProviderEventHDR
)
type ProviderEvent struct {
Kind ProviderEventKind
Payload []byte
}
// ProviderTelemetry holds measured provider-channel state only; it never
// contains provider routes, credentials, or payload bytes.
type ProviderTelemetry struct {
State string
ControlRTT time.Duration
ControlJitter time.Duration
ReliableSent uint64
ReliableRetransmits uint64
PendingReliable uint64
MediaDrops uint64
}
type Provider interface { type Provider interface {
Start(context.Context, LaunchRequest) (ProviderSession, error) Start(context.Context, LaunchRequest) (ProviderSession, error)
} }
@@ -218,9 +213,12 @@ type ProviderSession interface {
Ready(context.Context) error Ready(context.Context) error
Video() <-chan []byte Video() <-chan []byte
Audio() <-chan []byte Audio() <-chan []byte
Events() <-chan ProviderEvent
Input(context.Context, InputEvent) error Input(context.Context, InputEvent) error
Feedback(context.Context, Feedback) error Feedback(context.Context, Feedback) error
Reconnect(context.Context) error ReadClipboard(context.Context) (string, error)
WriteClipboard(context.Context, string) error
Telemetry() ProviderTelemetry
ReleaseAll(context.Context) error ReleaseAll(context.Context) error
Terminate(context.Context) error Terminate(context.Context) error
State() protocol.ProviderState State() protocol.ProviderState
@@ -357,7 +355,7 @@ func (f *FakeApollo) Setup(context.Context, LaunchRequest) ([]byte, error) {
if f.config.Failure == FakeFailureMalformed { if f.config.Failure == FakeFailureMalformed {
return []byte("RTSP/1.0 200 OK\r\n\r\n"), nil return []byte("RTSP/1.0 200 OK\r\n\r\n"), nil
} }
return []byte("RTSP/1.0 200 OK\r\nSession: fixture-session\r\nTransport: RTP/AVP/TCP;interleaved=0-1\r\n\r\n"), nil return []byte("RTSP/1.0 200 OK\r\nSession: fixture-session\r\nTransport: unicast;server_port=43000\r\n\r\n"), nil
} }
func (f *FakeApollo) Open(_ context.Context, request LaunchRequest, _ RTSPResponse) (ProviderSession, error) { func (f *FakeApollo) Open(_ context.Context, request LaunchRequest, _ RTSPResponse) (ProviderSession, error) {
@@ -365,6 +363,8 @@ func (f *FakeApollo) Open(_ context.Context, request LaunchRequest, _ RTSPRespon
failure: f.config.Failure, failure: f.config.Failure,
video: make(chan []byte, 16), video: make(chan []byte, 16),
audio: make(chan []byte, 16), audio: make(chan []byte, 16),
events: make(chan ProviderEvent, 16),
clipboardWrites: make(chan string, 1),
state: protocol.ProviderState{Version: "1", SessionID: request.SessionID, State: ProviderStateStarting, Channels: []string{"video", "audio", "input", "feedback"}}, state: protocol.ProviderState{Version: "1", SessionID: request.SessionID, State: ProviderStateStarting, Channels: []string{"video", "audio", "input", "feedback"}},
pressed: make(map[string]struct{}), pressed: make(map[string]struct{}),
} }
@@ -406,10 +406,13 @@ type fakeSession struct {
failure FakeFailure failure FakeFailure
video chan []byte video chan []byte
audio chan []byte audio chan []byte
events chan ProviderEvent
state protocol.ProviderState state protocol.ProviderState
pressed map[string]struct{} pressed map[string]struct{}
inputs []InputEvent inputs []InputEvent
feedback []Feedback feedback []Feedback
clipboard string
clipboardWrites chan string
releaseAll int releaseAll int
closeOnce sync.Once closeOnce sync.Once
} }
@@ -430,6 +433,14 @@ func (s *fakeSession) Ready(ctx context.Context) error {
func (s *fakeSession) Video() <-chan []byte { return s.video } func (s *fakeSession) Video() <-chan []byte { return s.video }
func (s *fakeSession) Audio() <-chan []byte { return s.audio } func (s *fakeSession) Audio() <-chan []byte { return s.audio }
func (s *fakeSession) Events() <-chan ProviderEvent { return s.events }
func (s *fakeSession) EmitEvent(event ProviderEvent) {
select {
case s.events <- ProviderEvent{Kind: event.Kind, Payload: append([]byte(nil), event.Payload...)}:
default:
}
}
func (s *fakeSession) EmitVideo(payload []byte) { func (s *fakeSession) EmitVideo(payload []byte) {
select { select {
@@ -475,13 +486,33 @@ func (s *fakeSession) Feedback(_ context.Context, feedback Feedback) error {
return nil return nil
} }
func (s *fakeSession) Reconnect(_ context.Context) error { func (s *fakeSession) ReadClipboard(ctx context.Context) (string, error) {
if err := ctx.Err(); err != nil {
return "", err
}
s.mu.Lock() s.mu.Lock()
defer s.mu.Unlock() defer s.mu.Unlock()
if s.state.State == ProviderStateTerminated { if s.state.State != ProviderStateReady {
return ErrProviderTerminated return "", ErrProviderDisconnected
}
return s.clipboard, nil
}
func (s *fakeSession) WriteClipboard(ctx context.Context, value string) error {
if err := ctx.Err(); err != nil {
return err
}
s.mu.Lock()
if s.state.State != ProviderStateReady {
s.mu.Unlock()
return ErrProviderDisconnected
}
s.clipboard = value
s.mu.Unlock()
select {
case s.clipboardWrites <- value:
default:
} }
s.state.State = ProviderStateReady
return nil return nil
} }
@@ -528,6 +559,10 @@ func (s *fakeSession) State() protocol.ProviderState {
return s.state return s.state
} }
func (s *fakeSession) Telemetry() ProviderTelemetry {
return ProviderTelemetry{State: s.State().State}
}
func (s *fakeSession) Disconnect() { func (s *fakeSession) Disconnect() {
s.mu.Lock() s.mu.Lock()
s.state.State = ProviderStateDisconnected s.state.State = ProviderStateDisconnected
+165
View File
@@ -2,33 +2,114 @@ package gateway
import ( import (
"context" "context"
"sync"
"sync/atomic" "sync/atomic"
"time" "time"
) )
type Metrics struct { type Metrics struct {
ActiveSessions atomic.Int64 ActiveSessions atomic.Int64
AdmittedSessions atomic.Uint64
AdmissionRejects atomic.Uint64 AdmissionRejects atomic.Uint64
Reconnects atomic.Uint64
DrainTransitions atomic.Uint64
MediaDrops atomic.Uint64 MediaDrops atomic.Uint64
MediaPackets atomic.Uint64
MediaBytes atomic.Uint64
QueueDelayNanos atomic.Uint64
ProcessingDelayNanos atomic.Uint64
ProcessingSamples atomic.Uint64
PacingDelayNanos atomic.Uint64
ProviderErrors atomic.Uint64 ProviderErrors atomic.Uint64
InputRejected atomic.Uint64 InputRejected atomic.Uint64
ControlRTTNanos atomic.Uint64
ControlJitterNanos atomic.Uint64
ControlLossPPM atomic.Uint64
PendingReliable atomic.Uint64
ProviderState atomic.Uint64
} }
type MetricsSnapshot struct { type MetricsSnapshot struct {
ActiveSessions int64 ActiveSessions int64
AdmittedSessions uint64
AdmissionRejects uint64 AdmissionRejects uint64
Reconnects uint64
DrainTransitions uint64
MediaDrops uint64 MediaDrops uint64
MediaPackets uint64
MediaBytes uint64
QueueDelayNanos uint64
ProcessingDelayNanos uint64
ProcessingSamples uint64
PacingDelayNanos uint64
ProviderErrors uint64 ProviderErrors uint64
InputRejected uint64 InputRejected uint64
ControlRTTNanos uint64
ControlJitterNanos uint64
ControlLossPPM uint64
PendingReliable uint64
ProviderState uint64
} }
func (m *Metrics) Snapshot() MetricsSnapshot { func (m *Metrics) Snapshot() MetricsSnapshot {
return MetricsSnapshot{ return MetricsSnapshot{
ActiveSessions: m.ActiveSessions.Load(), ActiveSessions: m.ActiveSessions.Load(),
AdmittedSessions: m.AdmittedSessions.Load(),
AdmissionRejects: m.AdmissionRejects.Load(), AdmissionRejects: m.AdmissionRejects.Load(),
Reconnects: m.Reconnects.Load(),
DrainTransitions: m.DrainTransitions.Load(),
MediaDrops: m.MediaDrops.Load(), MediaDrops: m.MediaDrops.Load(),
MediaPackets: m.MediaPackets.Load(),
MediaBytes: m.MediaBytes.Load(),
QueueDelayNanos: m.QueueDelayNanos.Load(),
ProcessingDelayNanos: m.ProcessingDelayNanos.Load(),
ProcessingSamples: m.ProcessingSamples.Load(),
PacingDelayNanos: m.PacingDelayNanos.Load(),
ProviderErrors: m.ProviderErrors.Load(), ProviderErrors: m.ProviderErrors.Load(),
InputRejected: m.InputRejected.Load(), InputRejected: m.InputRejected.Load(),
ControlRTTNanos: m.ControlRTTNanos.Load(),
ControlJitterNanos: m.ControlJitterNanos.Load(),
ControlLossPPM: m.ControlLossPPM.Load(),
PendingReliable: m.PendingReliable.Load(),
ProviderState: m.ProviderState.Load(),
}
}
func (m *Metrics) observeProviderTelemetry(telemetry ProviderTelemetry) {
if m == nil {
return
}
m.ControlRTTNanos.Store(uint64(telemetry.ControlRTT))
m.ControlJitterNanos.Store(uint64(telemetry.ControlJitter))
m.PendingReliable.Store(telemetry.PendingReliable)
if telemetry.ReliableSent == 0 {
m.ControlLossPPM.Store(0)
} else {
m.ControlLossPPM.Store(telemetry.ReliableRetransmits * 1_000_000 / telemetry.ReliableSent)
}
}
func (m *Metrics) observeProviderState(state string) {
if m == nil {
return
}
switch state {
case ProviderStateStarting:
m.ProviderState.Store(1)
case ProviderStateReady:
m.ProviderState.Store(2)
case ProviderStateDisconnected:
m.ProviderState.Store(3)
case ProviderStateTerminating:
m.ProviderState.Store(4)
case ProviderStateTerminated:
m.ProviderState.Store(5)
case ProviderStateCleanup:
m.ProviderState.Store(6)
case ProviderStateFailed:
m.ProviderState.Store(7)
default:
m.ProviderState.Store(0)
} }
} }
@@ -37,6 +118,90 @@ type Pacer struct {
last time.Time last time.Time
} }
// fairPacer is the gateway's one shared, equal-tier media scheduler. Each
// session can hold only its existing bounded provider media channel while it
// waits for the next reservation, so a slow client cannot grow a global queue.
type fairPacer struct {
mu sync.Mutex
bytesPerSecond int64
flows map[string]fairPacerFlow
}
type fairPacerFlow struct {
next time.Time
lastSeen time.Time
}
func newFairPacer(kbps int64) *fairPacer {
pacer := &fairPacer{flows: make(map[string]fairPacerFlow)}
pacer.setKbps(kbps)
return pacer
}
func (p *fairPacer) setKbps(kbps int64) {
if p == nil {
return
}
p.mu.Lock()
if kbps > 0 {
p.bytesPerSecond = kbps * 1000 / 8
} else {
p.bytesPerSecond = 0
}
p.mu.Unlock()
}
func (p *fairPacer) remove(flow string) {
if p == nil || flow == "" {
return
}
p.mu.Lock()
delete(p.flows, flow)
p.mu.Unlock()
}
func (p *fairPacer) reserveAt(now time.Time, flow string, bytes int) time.Time {
if p == nil || flow == "" || bytes < 1 {
return now
}
p.mu.Lock()
defer p.mu.Unlock()
if p.bytesPerSecond < 1 {
return now
}
for key, state := range p.flows {
if now.Sub(state.lastSeen) > time.Second {
delete(p.flows, key)
}
}
state := p.flows[flow]
state.lastSeen = now
p.flows[flow] = state
base := now
if state.next.After(base) {
base = state.next
}
numerator := int64(bytes) * int64(len(p.flows)) * int64(time.Second)
delay := time.Duration((numerator + p.bytesPerSecond - 1) / p.bytesPerSecond)
state.next = base.Add(delay)
p.flows[flow] = state
return state.next
}
func (p *fairPacer) wait(ctx context.Context, flow string, bytes int) error {
target := p.reserveAt(time.Now(), flow, bytes)
if delay := time.Until(target); delay > 0 {
timer := time.NewTimer(delay)
defer timer.Stop()
select {
case <-ctx.Done():
return ctx.Err()
case <-timer.C:
}
}
return nil
}
func NewPacer(kbps int64) *Pacer { func NewPacer(kbps int64) *Pacer {
if kbps < 1 { if kbps < 1 {
return &Pacer{} return &Pacer{}
+1 -1
View File
@@ -1 +1 @@
RTSP/1.0 200 OK\r\nSession: fixture-session\r\nTransport: RTP/AVP/TCP;interleaved=0-1\r\n\r\n RTSP/1.0 200 OK\r\nSession: fixture-session\r\nTransport: unicast;server_port=43000\r\n\r\n
+361 -35
View File
@@ -21,6 +21,7 @@ import (
const ( const (
defaultHelloLimit = 16 * 1024 defaultHelloLimit = 16 * 1024
defaultControlLimit = 128 * 1024 defaultControlLimit = 128 * 1024
clientControlBacklog = 64
applicationError = quic.ApplicationErrorCode(0x100) applicationError = quic.ApplicationErrorCode(0x100)
) )
@@ -53,6 +54,10 @@ type ProviderStateReporter interface {
ReportProviderState(context.Context, protocol.ProviderState) error ReportProviderState(context.Context, protocol.ProviderState) error
} }
type ClipboardAuditReporter interface {
ReportClipboardAudit(context.Context, protocol.GatewayClipboardAudit) error
}
type ServerConfig struct { type ServerConfig struct {
ListenAddress string ListenAddress string
TLSConfig *tls.Config TLSConfig *tls.Config
@@ -62,6 +67,7 @@ type ServerConfig struct {
ProviderCapabilities protocol.CapabilityProfile ProviderCapabilities protocol.CapabilityProfile
Admission Admission Admission Admission
ProviderStateReporter ProviderStateReporter ProviderStateReporter ProviderStateReporter
ClipboardAuditReporter ClipboardAuditReporter
Provider Provider Provider Provider
ProviderProfile string ProviderProfile string
ProviderIdentity string ProviderIdentity string
@@ -72,6 +78,7 @@ type Server struct {
listener *quic.Listener listener *quic.Listener
config ServerConfig config ServerConfig
metrics *Metrics metrics *Metrics
pacer *fairPacer
mu sync.Mutex mu sync.Mutex
sessions map[*gatewaySession]struct{} sessions map[*gatewaySession]struct{}
draining atomic.Bool draining atomic.Bool
@@ -115,7 +122,7 @@ func NewServer(config ServerConfig) (*Server, error) {
if err != nil { if err != nil {
return nil, err return nil, err
} }
return &Server{listener: listener, config: config, metrics: &Metrics{}, sessions: make(map[*gatewaySession]struct{})}, nil return &Server{listener: listener, config: config, metrics: &Metrics{}, pacer: newFairPacer(config.PacerKbps), sessions: make(map[*gatewaySession]struct{})}, nil
} }
func validateServerTLS(config *tls.Config) error { func validateServerTLS(config *tls.Config) error {
@@ -130,7 +137,9 @@ func (s *Server) Metrics() MetricsSnapshot { return s.metrics.Snapshot() }
func (s *Server) Draining() bool { return s.draining.Load() } func (s *Server) Draining() bool { return s.draining.Load() }
func (s *Server) BeginDrain() { func (s *Server) BeginDrain() {
s.draining.Store(true) if s.draining.CompareAndSwap(false, true) {
s.metrics.DrainTransitions.Add(1)
}
} }
func (s *Server) Serve(ctx context.Context) error { func (s *Server) Serve(ctx context.Context) error {
@@ -228,6 +237,17 @@ func (s *Server) handleConnection(parent context.Context, connection *quic.Conn)
_ = writeStableError(stream, "no_capability_overlap", err, false) _ = writeStableError(stream, "no_capability_overlap", err, false)
return return
} }
clipboard, err := newClipboardGate(work.ClipboardPolicy, time.Now)
if err != nil {
_ = s.config.Admission.Release(context.Background(), authority)
_ = writeStableError(stream, "provider_work_unavailable", ErrAdmissionRejected, false)
return
}
if (work.ClipboardPolicy.ClientToProviderEnabled || work.ClipboardPolicy.ProviderToClientEnabled) && s.config.ClipboardAuditReporter == nil {
_ = s.config.Admission.Release(context.Background(), authority)
_ = writeStableError(stream, "clipboard_audit_unavailable", ErrAdmissionRejected, true)
return
}
if err := s.reportProviderState(ctx, protocol.ProviderState{Version: "1", SessionID: request.SessionID, State: ProviderStateStarting, CleanupPending: false, Channels: []string{"video", "audio", "input", "feedback"}}); err != nil { if err := s.reportProviderState(ctx, protocol.ProviderState{Version: "1", SessionID: request.SessionID, State: ProviderStateStarting, CleanupPending: false, Channels: []string{"video", "audio", "input", "feedback"}}); err != nil {
_ = s.config.Admission.Release(context.Background(), authority) _ = s.config.Admission.Release(context.Background(), authority)
_ = writeStableError(stream, "provider_state_unavailable", err, true) _ = writeStableError(stream, "provider_state_unavailable", err, true)
@@ -256,9 +276,13 @@ func (s *Server) handleConnection(parent context.Context, connection *quic.Conn)
_ = s.config.Admission.Release(context.Background(), authority) _ = s.config.Admission.Release(context.Background(), authority)
return return
} }
session := newGatewaySession(s, connection, stream, request, authority, providerSession) session := newGatewaySession(s, connection, stream, request, authority, providerSession, clipboard)
s.addSession(session) s.addSession(session)
s.metrics.ActiveSessions.Add(1) s.metrics.ActiveSessions.Add(1)
s.metrics.AdmittedSessions.Add(1)
if authority.ReconnectSequence > 0 {
s.metrics.Reconnects.Add(1)
}
defer func() { defer func() {
s.removeSession(session) s.removeSession(session)
s.metrics.ActiveSessions.Add(-1) s.metrics.ActiveSessions.Add(-1)
@@ -267,12 +291,13 @@ func (s *Server) handleConnection(parent context.Context, connection *quic.Conn)
} }
func (s *Server) reportProviderState(ctx context.Context, state protocol.ProviderState) error { func (s *Server) reportProviderState(ctx context.Context, state protocol.ProviderState) error {
if s.config.ProviderStateReporter == nil {
return nil
}
if err := state.Validate(); err != nil { if err := state.Validate(); err != nil {
return err return err
} }
s.metrics.observeProviderState(state.State)
if s.config.ProviderStateReporter == nil {
return nil
}
return s.config.ProviderStateReporter.ReportProviderState(ctx, state) return s.config.ProviderStateReporter.ReportProviderState(ctx, state)
} }
@@ -321,19 +346,21 @@ type gatewaySession struct {
request protocol.TunnelAdmissionRequest request protocol.TunnelAdmissionRequest
authority protocol.SessionAuthority authority protocol.SessionAuthority
provider ProviderSession provider ProviderSession
pacer *Pacer clipboard *clipboardGate
ctx context.Context ctx context.Context
cancel context.CancelFunc cancel context.CancelFunc
cleanupOnce sync.Once cleanupOnce sync.Once
inputMu sync.Mutex inputMu sync.Mutex
controlWriteMu sync.Mutex
pressed map[string]struct{} pressed map[string]struct{}
sequence atomic.Uint32 sequence atomic.Uint32
mediaDrops uint64
result chan error result chan error
} }
func newGatewaySession(server *Server, connection *quic.Conn, control *quic.Stream, request protocol.TunnelAdmissionRequest, authority protocol.SessionAuthority, provider ProviderSession) *gatewaySession { func newGatewaySession(server *Server, connection *quic.Conn, control *quic.Stream, request protocol.TunnelAdmissionRequest, authority protocol.SessionAuthority, provider ProviderSession, clipboard *clipboardGate) *gatewaySession {
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())
return &gatewaySession{server: server, connection: connection, control: control, request: request, authority: authority, provider: provider, pacer: NewPacer(server.config.PacerKbps), ctx: ctx, cancel: cancel, pressed: make(map[string]struct{}), result: make(chan error, 3)} return &gatewaySession{server: server, connection: connection, control: control, request: request, authority: authority, provider: provider, clipboard: clipboard, ctx: ctx, cancel: cancel, pressed: make(map[string]struct{}), result: make(chan error, 3)}
} }
func (s *gatewaySession) run() { func (s *gatewaySession) run() {
@@ -350,6 +377,11 @@ func (s *gatewaySession) run() {
go s.controlLoop() go s.controlLoop()
go s.datagramLoop() go s.datagramLoop()
go s.mediaLoop() go s.mediaLoop()
go s.providerEventLoop()
go s.providerTelemetryLoop()
if s.clipboard != nil && s.clipboard.policy.ProviderToClientEnabled {
go s.clipboardLoop()
}
select { select {
case <-timer.C: case <-timer.C:
s.server.metrics.InputRejected.Add(1) s.server.metrics.InputRejected.Add(1)
@@ -359,6 +391,96 @@ func (s *gatewaySession) run() {
s.cancel() s.cancel()
} }
func (s *gatewaySession) providerEventLoop() {
events := s.provider.Events()
for events != nil {
select {
case <-s.ctx.Done():
return
case event, ok := <-events:
if !ok {
return
}
payload, err := EncodeProviderEvent(event)
if err == nil {
err = s.sendControl(s.sequence.Add(1), payload)
}
if err != nil {
s.result <- err
return
}
}
}
}
func (s *gatewaySession) clipboardLoop() {
ticker := time.NewTicker(500 * time.Millisecond)
defer ticker.Stop()
for {
select {
case <-s.ctx.Done():
return
case <-ticker.C:
text, err := s.provider.ReadClipboard(s.ctx)
if err != nil {
if auditErr := s.reportClipboardAudit("provider_to_client", "rejected", clipboardAuditTextBytes(text), clipboardAuditReason(err)); auditErr != nil {
s.result <- auditErr
return
}
s.result <- err
return
}
value, suppress, err := s.clipboard.fromProvider(text)
if err != nil {
if auditErr := s.reportClipboardAudit("provider_to_client", "rejected", clipboardAuditTextBytes(text), clipboardAuditReason(err)); auditErr != nil {
s.result <- auditErr
return
}
s.result <- err
return
}
if suppress {
if err := s.reportClipboardAudit("provider_to_client", "suppressed", clipboardAuditTextBytes(text), "loop"); err != nil {
s.result <- err
return
}
continue
}
if err := s.reportClipboardAudit(value.Direction, "forwarded", clipboardAuditTextBytes(value.Text), "forwarded"); err != nil {
s.result <- err
return
}
if err := s.sendClipboard(value); err != nil {
s.result <- err
return
}
}
}
}
func (s *gatewaySession) providerTelemetryLoop() {
s.observeProviderTelemetry()
ticker := time.NewTicker(500 * time.Millisecond)
defer ticker.Stop()
for {
select {
case <-s.ctx.Done():
return
case <-ticker.C:
s.observeProviderTelemetry()
}
}
}
func (s *gatewaySession) observeProviderTelemetry() {
telemetry := s.provider.Telemetry()
if telemetry.MediaDrops >= s.mediaDrops {
s.server.metrics.MediaDrops.Add(telemetry.MediaDrops - s.mediaDrops)
s.mediaDrops = telemetry.MediaDrops
}
s.server.metrics.observeProviderTelemetry(telemetry)
}
func (s *gatewaySession) controlLoop() { func (s *gatewaySession) controlLoop() {
for { for {
data, err := readWire(s.control, defaultControlLimit) data, err := readWire(s.control, defaultControlLimit)
@@ -378,12 +500,32 @@ func (s *gatewaySession) controlLoop() {
} }
switch frame.FlowID { switch frame.FlowID {
case "control": case "control":
if err := s.handleControl(payload); err != nil { sequence, sequenceErr := channelSequence(frame.Sequence)
if sequenceErr != nil {
s.result <- sequenceErr
return
}
if err := s.handleControl(payload, sequence); err != nil {
s.result <- err s.result <- err
return return
} }
case "input": case "input":
if err := s.handleInput(payload); err != nil { sequence, sequenceErr := channelSequence(frame.Sequence)
if sequenceErr != nil {
s.result <- sequenceErr
return
}
if err := s.handleInput(payload, sequence); err != nil {
s.result <- err
return
}
case "clipboard":
value, decodeErr := protocol.DecodeGatewayClipboardText(payload)
if decodeErr != nil {
s.result <- ErrProviderMalformed
return
}
if err := s.handleClipboard(value); err != nil {
s.result <- err s.result <- err
return return
} }
@@ -408,19 +550,13 @@ func (s *gatewaySession) datagramLoop() {
} }
switch frame.Channel { switch frame.Channel {
case ChannelInput: case ChannelInput:
if err := s.handleInput(frame.Payload); err != nil { if err := s.handleInput(frame.Payload, frame.Sequence); err != nil {
s.result <- err s.result <- err
return return
} }
case ChannelText: case ChannelText:
if len(frame.Payload) > 4096 { s.result <- ErrFrameChannel
s.result <- ErrFramePayloadLimit
return return
}
if err := s.provider.Feedback(s.ctx, Feedback{Sequence: frame.Sequence, Payload: append([]byte(nil), frame.Payload...)}); err != nil {
s.result <- err
return
}
default: default:
s.result <- ErrFrameChannel s.result <- ErrFrameChannel
return return
@@ -458,6 +594,7 @@ func (s *gatewaySession) mediaLoop() {
} }
func (s *gatewaySession) sendMedia(channel byte, payload []byte) error { func (s *gatewaySession) sendMedia(channel byte, payload []byte) error {
processingStarted := time.Now()
frames, err := FragmentPayload(channel, s.sequence.Add(1), uint64(time.Now().UnixMilli()), payload) frames, err := FragmentPayload(channel, s.sequence.Add(1), uint64(time.Now().UnixMilli()), payload)
if err != nil { if err != nil {
return err return err
@@ -467,37 +604,122 @@ func (s *gatewaySession) sendMedia(channel byte, payload []byte) error {
if err != nil { if err != nil {
return err return err
} }
if err := s.pacer.Wait(s.ctx, len(encoded)); err != nil { pacingStarted := time.Now()
if err := s.server.pacer.wait(s.ctx, s.authority.SessionID, len(encoded)); err != nil {
return err return err
} }
s.server.metrics.PacingDelayNanos.Add(uint64(time.Since(pacingStarted)))
s.server.metrics.QueueDelayNanos.Add(uint64(time.Since(pacingStarted)))
if err := s.connection.SendDatagram(encoded); err != nil { if err := s.connection.SendDatagram(encoded); err != nil {
return err return err
} }
s.server.metrics.MediaPackets.Add(1)
s.server.metrics.MediaBytes.Add(uint64(len(encoded)))
s.server.metrics.ProcessingDelayNanos.Add(uint64(time.Since(processingStarted)))
s.server.metrics.ProcessingSamples.Add(1)
} }
return nil return nil
} }
func (s *gatewaySession) handleControl(payload []byte) error { func (s *gatewaySession) sendControl(sequence uint32, payload []byte) error {
if len(payload) < 4 { if len(payload) > 1024 {
return ErrProviderMalformed return ErrFramePayloadLimit
} }
switch string(payload[:4]) { frame := protocol.ChannelFrame{Version: "1", FlowID: "control", Sequence: int64(sequence), Flags: 0, FragmentIndex: 0, FragmentCount: 1, TimestampMs: time.Now().UnixMilli(), Payload: base64.StdEncoding.EncodeToString(payload)}
case "TERM": encoded, err := protocol.EncodeChannelFrame(frame)
return errors.New("client requested termination") if err != nil {
case "RECN": return err
return s.provider.Reconnect(s.ctx) }
case "FBRK": s.controlWriteMu.Lock()
return s.provider.Feedback(s.ctx, Feedback{Payload: append([]byte(nil), payload[4:]...)}) defer s.controlWriteMu.Unlock()
return writeWire(s.control, encoded, defaultControlLimit)
}
func (s *gatewaySession) sendClipboard(value protocol.GatewayClipboardText) error {
payload, err := protocol.EncodeGatewayClipboardText(value)
if err != nil {
return err
}
frame := protocol.ChannelFrame{Version: "1", FlowID: "clipboard", Sequence: int64(s.sequence.Add(1)), Flags: 0, FragmentIndex: 0, FragmentCount: 1, TimestampMs: time.Now().UnixMilli(), Payload: base64.StdEncoding.EncodeToString(payload)}
encoded, err := protocol.EncodeChannelFrame(frame)
if err != nil {
return err
}
s.controlWriteMu.Lock()
defer s.controlWriteMu.Unlock()
return writeWire(s.control, encoded, defaultControlLimit)
}
func (s *gatewaySession) handleControl(payload []byte, sequence uint32) error {
feedback, err := DecodeClientFeedback(payload)
if err != nil {
return err
}
feedback.Sequence = sequence
return s.provider.Feedback(s.ctx, feedback)
}
func (s *gatewaySession) handleClipboard(value protocol.GatewayClipboardText) error {
if s.clipboard == nil {
return ErrClipboardDenied
}
suppress, err := s.clipboard.fromClient(value)
if err != nil {
if auditErr := s.reportClipboardAudit(value.Direction, "rejected", clipboardAuditTextBytes(value.Text), clipboardAuditReason(err)); auditErr != nil {
return auditErr
}
return err
}
if suppress {
return s.reportClipboardAudit(value.Direction, "suppressed", clipboardAuditTextBytes(value.Text), "loop")
}
if err := s.provider.WriteClipboard(s.ctx, value.Text); err != nil {
s.clipboard.retractClient(value)
if auditErr := s.reportClipboardAudit(value.Direction, "rejected", clipboardAuditTextBytes(value.Text), "provider"); auditErr != nil {
return auditErr
}
return err
}
return s.reportClipboardAudit(value.Direction, "forwarded", clipboardAuditTextBytes(value.Text), "forwarded")
}
func (s *gatewaySession) reportClipboardAudit(direction, outcome string, textBytes int, reason string) error {
if s.server.config.ClipboardAuditReporter == nil {
return ErrClipboardDenied
}
ctx, cancel := context.WithTimeout(s.ctx, 5*time.Second)
defer cancel()
return s.server.config.ClipboardAuditReporter.ReportClipboardAudit(ctx, protocol.GatewayClipboardAudit{
Version: "1", SessionID: s.authority.SessionID, Direction: direction, Outcome: outcome, TextBytes: int64(textBytes), Reason: reason,
})
}
func clipboardAuditTextBytes(text string) int {
if len(text) > 65536 {
return 65536
}
return len(text)
}
func clipboardAuditReason(err error) string {
switch {
case errors.Is(err, ErrClipboardRate):
return "rate"
case errors.Is(err, ErrProviderMalformed):
return "malformed"
case errors.Is(err, ErrClipboardDenied):
return "policy"
default: default:
return ErrProviderMalformed return "provider"
} }
} }
func (s *gatewaySession) handleInput(payload []byte) error { func (s *gatewaySession) handleInput(payload []byte, sequence uint32) error {
event, err := DecodeInputEvent(payload) event, err := DecodeInputEvent(payload)
if err != nil { if err != nil {
return err return err
} }
event.Sequence = sequence
if err := s.provider.Input(s.ctx, event); err != nil { if err := s.provider.Input(s.ctx, event); err != nil {
s.server.metrics.InputRejected.Add(1) s.server.metrics.InputRejected.Add(1)
return err return err
@@ -513,9 +735,17 @@ func (s *gatewaySession) handleInput(payload []byte) error {
return nil return nil
} }
func channelSequence(sequence int64) (uint32, error) {
if sequence < 0 || sequence > int64(^uint32(0)) {
return 0, ErrProviderMalformed
}
return uint32(sequence), nil
}
func (s *gatewaySession) cleanup() { func (s *gatewaySession) cleanup() {
s.cleanupOnce.Do(func() { s.cleanupOnce.Do(func() {
s.cancel() s.cancel()
s.server.pacer.remove(s.authority.SessionID)
cleanupCtx, cancel := context.WithTimeout(context.Background(), time.Second) cleanupCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel() defer cancel()
releaseInputsErr := s.provider.ReleaseAll(cleanupCtx) releaseInputsErr := s.provider.ReleaseAll(cleanupCtx)
@@ -615,6 +845,9 @@ func readWire(reader io.Reader, max int) ([]byte, error) {
type Client struct { type Client struct {
connection *quic.Conn connection *quic.Conn
control *quic.Stream control *quic.Stream
controlReadMu sync.Mutex
controlWriteMu sync.Mutex
pendingControl map[string][][]byte
Authority protocol.SessionAuthority Authority protocol.SessionAuthority
} }
@@ -658,7 +891,7 @@ func Dial(ctx context.Context, address string, tlsConfig *tls.Config, request pr
} }
return nil, authorityErr return nil, authorityErr
} }
return &Client{connection: connection, control: stream, Authority: authority}, nil return &Client{connection: connection, control: stream, pendingControl: make(map[string][][]byte), Authority: authority}, nil
} }
func (c *Client) SendInput(event InputEvent) error { func (c *Client) SendInput(event InputEvent) error {
@@ -671,15 +904,46 @@ func (c *Client) SendInput(event InputEvent) error {
if err != nil { if err != nil {
return err return err
} }
return writeWire(c.control, encoded, defaultControlLimit) return c.writeControl(encoded)
} }
func (c *Client) SendControl(payload []byte) error { func (c *Client) SendControl(payload []byte) error {
frame := protocol.ChannelFrame{Version: "1", FlowID: "control", Sequence: 0, Flags: 0, FragmentIndex: 0, FragmentCount: 1, TimestampMs: time.Now().UnixMilli(), Payload: base64.StdEncoding.EncodeToString(payload)} return c.sendControl(0, payload)
}
func (c *Client) SendFeedback(feedback Feedback) error {
payload, err := EncodeClientFeedback(feedback)
if err != nil {
return err
}
return c.sendControl(feedback.Sequence, payload)
}
func (c *Client) SendClipboard(value protocol.GatewayClipboardText) error {
payload, err := protocol.EncodeGatewayClipboardText(value)
if err != nil {
return err
}
frame := protocol.ChannelFrame{Version: "1", FlowID: "clipboard", Sequence: 0, Flags: 0, FragmentIndex: 0, FragmentCount: 1, TimestampMs: time.Now().UnixMilli(), Payload: base64.StdEncoding.EncodeToString(payload)}
encoded, err := protocol.EncodeChannelFrame(frame) encoded, err := protocol.EncodeChannelFrame(frame)
if err != nil { if err != nil {
return err return err
} }
return c.writeControl(encoded)
}
func (c *Client) sendControl(sequence uint32, payload []byte) error {
frame := protocol.ChannelFrame{Version: "1", FlowID: "control", Sequence: int64(sequence), Flags: 0, FragmentIndex: 0, FragmentCount: 1, TimestampMs: time.Now().UnixMilli(), Payload: base64.StdEncoding.EncodeToString(payload)}
encoded, err := protocol.EncodeChannelFrame(frame)
if err != nil {
return err
}
return c.writeControl(encoded)
}
func (c *Client) writeControl(encoded []byte) error {
c.controlWriteMu.Lock()
defer c.controlWriteMu.Unlock()
return writeWire(c.control, encoded, defaultControlLimit) return writeWire(c.control, encoded, defaultControlLimit)
} }
@@ -691,6 +955,68 @@ func (c *Client) ReceiveFrame(ctx context.Context) (Frame, error) {
return DecodeFrame(data) return DecodeFrame(data)
} }
func (c *Client) ReceiveProviderEvent(ctx context.Context) (ProviderEvent, error) {
payload, err := c.receiveControlPayload(ctx, "control")
if err != nil {
return ProviderEvent{}, err
}
if len(payload) > 1024 {
return ProviderEvent{}, ErrProviderMalformed
}
return DecodeProviderEvent(payload)
}
func (c *Client) ReceiveClipboard(ctx context.Context) (protocol.GatewayClipboardText, error) {
payload, err := c.receiveControlPayload(ctx, "clipboard")
if err != nil {
return protocol.GatewayClipboardText{}, err
}
return protocol.DecodeGatewayClipboardText(payload)
}
func (c *Client) receiveControlPayload(ctx context.Context, flowID string) ([]byte, error) {
if c == nil || c.control == nil || (flowID != "control" && flowID != "clipboard") {
return nil, ErrProviderMalformed
}
c.controlReadMu.Lock()
defer c.controlReadMu.Unlock()
if err := ctx.Err(); err != nil {
return nil, err
}
if queued := c.pendingControl[flowID]; len(queued) > 0 {
payload := queued[0]
c.pendingControl[flowID] = queued[1:]
return payload, nil
}
if deadline, ok := ctx.Deadline(); ok {
if err := c.control.SetReadDeadline(deadline); err != nil {
return nil, err
}
defer c.control.SetReadDeadline(time.Time{})
}
for {
data, err := readWire(c.control, defaultControlLimit)
if err != nil {
return nil, err
}
frame, err := protocol.DecodeChannelFrame(data)
if err != nil || (frame.FlowID != "control" && frame.FlowID != "clipboard") {
return nil, ErrProviderMalformed
}
payload, err := base64.StdEncoding.DecodeString(frame.Payload)
if err != nil || len(payload) > maxFrameSize {
return nil, ErrProviderMalformed
}
if frame.FlowID == flowID {
return payload, nil
}
if len(c.pendingControl[frame.FlowID]) >= clientControlBacklog {
return nil, ErrFramePayloadLimit
}
c.pendingControl[frame.FlowID] = append(c.pendingControl[frame.FlowID], payload)
}
}
func (c *Client) Close() error { func (c *Client) Close() error {
return c.connection.CloseWithError(applicationError, "client closed") return c.connection.CloseWithError(applicationError, "client closed")
} }