feat(gateway): repair native Apollo provider path
This commit is contained in:
@@ -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]...)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
+535
-118
@@ -9,15 +9,20 @@ import (
|
||||
"crypto/x509"
|
||||
"encoding/binary"
|
||||
"encoding/hex"
|
||||
"encoding/xml"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
protocol "git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol/gen/go/protocol"
|
||||
)
|
||||
@@ -29,11 +34,11 @@ type NativeApolloBackend struct {
|
||||
Dialer *net.Dialer
|
||||
|
||||
mu sync.Mutex
|
||||
pending map[string]net.Conn
|
||||
pending map[string]*apolloRTSPSetup
|
||||
}
|
||||
|
||||
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) {
|
||||
@@ -53,7 +58,11 @@ func newPinnedApolloHTTPClient(work protocol.ProviderSessionWork) (*http.Client,
|
||||
if err != nil {
|
||||
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) {
|
||||
@@ -74,6 +83,88 @@ func apolloGet(ctx context.Context, client *http.Client, work protocol.ProviderS
|
||||
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) {
|
||||
identity, ok := providerIdentityFromKey(work.ProviderIdentity)
|
||||
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 {
|
||||
return nil, err
|
||||
}
|
||||
defer zeroApolloSecret(key)
|
||||
var keyID [4]byte
|
||||
if _, err := rand.Read(keyID[:]); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
launch, err := apolloGet(ctx, client, work, "/launch", url.Values{
|
||||
"uniqueid": {work.ClientID}, "appid": {work.ApplicationID}, "rikey": {hex.EncodeToString(key)},
|
||||
"rikeyid": {strconv.FormatUint(uint64(binary.BigEndian.Uint32(keyID[:])), 10)}, "localAudioPlayMode": {"0"},
|
||||
"corever": {"1"},
|
||||
})
|
||||
sessionPath, sessionValues := apolloSessionRequest(work, key, binary.BigEndian.Uint32(keyID[:]))
|
||||
launch, err := apolloGet(ctx, client, work, sessionPath, sessionValues)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -147,50 +236,25 @@ func (b *NativeApolloBackend) Setup(ctx context.Context, request LaunchRequest)
|
||||
if err != nil || streamPort != work.StreamPort {
|
||||
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 {
|
||||
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.pending[request.SessionID] = conn
|
||||
b.pending[request.SessionID] = setup
|
||||
b.mu.Unlock()
|
||||
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()
|
||||
conn, ok := b.pending[request.SessionID]
|
||||
setup, ok := b.pending[request.SessionID]
|
||||
delete(b.pending, request.SessionID)
|
||||
b.mu.Unlock()
|
||||
if !ok || conn == nil {
|
||||
if !ok || setup == nil {
|
||||
return nil, ErrProviderDisconnected
|
||||
}
|
||||
session := newNativeApolloSession(conn, request.SessionID)
|
||||
go session.readMedia()
|
||||
return session, nil
|
||||
return newNativeApolloProviderSession(ctx, setup)
|
||||
}
|
||||
|
||||
func readBounded(reader io.Reader, max int) ([]byte, error) {
|
||||
@@ -204,152 +268,503 @@ func readBounded(reader io.Reader, max int) ([]byte, error) {
|
||||
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 {
|
||||
conn net.Conn
|
||||
sessionID string
|
||||
video chan []byte
|
||||
audio chan []byte
|
||||
mu sync.Mutex
|
||||
state protocol.ProviderState
|
||||
closeOnce sync.Once
|
||||
done chan struct{}
|
||||
readDone chan struct{}
|
||||
enet *apolloENetPeer
|
||||
control *apolloControlCodec
|
||||
audioConn *net.UDPConn
|
||||
videoConn *net.UDPConn
|
||||
media *apolloMediaCodec
|
||||
videoFEC apolloVideoAssembler
|
||||
audioFEC apolloAudioAssembler
|
||||
audioPing []byte
|
||||
videoPing []byte
|
||||
sessionID string
|
||||
video chan []byte
|
||||
audio chan []byte
|
||||
events chan ProviderEvent
|
||||
mu sync.Mutex
|
||||
controlMu sync.Mutex
|
||||
state protocol.ProviderState
|
||||
pressed map[string]InputEvent
|
||||
closeOnce sync.Once
|
||||
channelsOnce sync.Once
|
||||
done 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 {
|
||||
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{})}
|
||||
func newNativeApolloSession(sessionID string) *nativeApolloSession {
|
||||
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 {
|
||||
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.state.State = ProviderStateReady
|
||||
s.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *nativeApolloSession) Video() <-chan []byte { return s.video }
|
||||
func (s *nativeApolloSession) Audio() <-chan []byte { return s.audio }
|
||||
func (s *nativeApolloSession) Video() <-chan []byte { return s.video }
|
||||
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 {
|
||||
payload, err := EncodeInputEvent(event)
|
||||
packet, err := encodeApolloInputEvent(event)
|
||||
if err != nil {
|
||||
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 {
|
||||
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 {
|
||||
return s.writeControl(ctx, ControlPacket{Kind: 4, Payload: []byte("RECN")})
|
||||
func (s *nativeApolloSession) ReadClipboard(ctx context.Context) (string, error) {
|
||||
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 {
|
||||
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 {
|
||||
_ = s.writeControl(ctx, ControlPacket{Kind: 5, Payload: []byte("TEAR")})
|
||||
var timedOut bool
|
||||
var cleanupErr error
|
||||
s.closeOnce.Do(func() {
|
||||
if err := s.ReleaseAll(ctx); err != nil {
|
||||
cleanupErr = err
|
||||
}
|
||||
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 {
|
||||
case <-s.readDone:
|
||||
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()
|
||||
if timedOut {
|
||||
cleanupErr = s.terminationErr
|
||||
if cleanupErr != nil {
|
||||
s.state.State = ProviderStateCleanup
|
||||
s.state.CleanupPending = true
|
||||
} else {
|
||||
s.state.State = ProviderStateTerminated
|
||||
}
|
||||
s.mu.Unlock()
|
||||
if timedOut {
|
||||
return ctx.Err()
|
||||
if cleanupErr != nil {
|
||||
return cleanupErr
|
||||
}
|
||||
s.control, s.media = nil, nil
|
||||
return nil
|
||||
}
|
||||
|
||||
func zeroApolloSecret(secret []byte) {
|
||||
for index := range secret {
|
||||
secret[index] = 0
|
||||
}
|
||||
}
|
||||
|
||||
func (s *nativeApolloSession) State() protocol.ProviderState {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.state
|
||||
}
|
||||
|
||||
func (s *nativeApolloSession) writeControl(ctx context.Context, packet ControlPacket) error {
|
||||
encoded, err := EncodeControlPacket(packet)
|
||||
func (s *nativeApolloSession) Telemetry() ProviderTelemetry {
|
||||
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 {
|
||||
return err
|
||||
}
|
||||
if deadline, ok := ctx.Deadline(); ok {
|
||||
_ = s.conn.SetWriteDeadline(deadline)
|
||||
if reliable {
|
||||
return s.enet.SendReliable(channel, encoded)
|
||||
}
|
||||
if _, err := s.conn.Write(encoded); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
return s.enet.SendUnsequenced(channel, encoded)
|
||||
}
|
||||
|
||||
func (s *nativeApolloSession) readMedia() {
|
||||
defer close(s.readDone)
|
||||
header := make([]byte, 4)
|
||||
func (s *nativeApolloSession) periodicApolloPing() {
|
||||
ticker := time.NewTicker(100 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
if _, err := io.ReadFull(s.conn, header); err != nil {
|
||||
select {
|
||||
case <-s.done:
|
||||
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)
|
||||
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
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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 {
|
||||
case channel <- payload:
|
||||
return false
|
||||
default:
|
||||
select {
|
||||
case <-channel:
|
||||
@@ -357,7 +772,9 @@ func pushLatest(channel chan []byte, payload []byte) {
|
||||
}
|
||||
select {
|
||||
case channel <- payload:
|
||||
return true
|
||||
default:
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+792
-27
@@ -2,22 +2,34 @@ package gateway
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/sha256"
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"encoding/binary"
|
||||
"encoding/hex"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
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) {
|
||||
serverTLS, clientTLS := testTLS(t)
|
||||
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]),
|
||||
ClientPrivateKeyPem: privateKeyPEM(t, clientTLS.Certificates[0]),
|
||||
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})
|
||||
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)
|
||||
streamListener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
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())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -72,40 +100,83 @@ func TestNativeApolloSetupLaunchesInventoriedApplicationWithEncryptedRTSP(t *tes
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
type fixtureKeyMaterial struct {
|
||||
key []byte
|
||||
keyID uint32
|
||||
}
|
||||
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)
|
||||
go func() {
|
||||
connection, acceptErr := streamListener.Accept()
|
||||
if acceptErr != nil {
|
||||
streamDone <- acceptErr
|
||||
return
|
||||
}
|
||||
defer connection.Close()
|
||||
key := <-keyReady
|
||||
header := make([]byte, encryptedRTSPHeaderSize)
|
||||
if _, readErr := io.ReadFull(connection, header); readErr != nil {
|
||||
streamDone <- readErr
|
||||
return
|
||||
}
|
||||
length := binary.BigEndian.Uint32(header[:4]) & 0x7fffffff
|
||||
frame := make([]byte, encryptedRTSPHeaderSize+int(length))
|
||||
copy(frame, header)
|
||||
if _, readErr := io.ReadFull(connection, frame[encryptedRTSPHeaderSize:]); readErr != nil {
|
||||
streamDone <- readErr
|
||||
return
|
||||
}
|
||||
codec, codecErr := newEncryptedRTSPCodec(key)
|
||||
if codecErr != nil {
|
||||
streamDone <- codecErr
|
||||
return
|
||||
}
|
||||
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" {
|
||||
streamDone <- ErrProviderMalformed
|
||||
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",
|
||||
}
|
||||
_, 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")))
|
||||
streamDone <- writeErr
|
||||
for index, method := range expectedMethods {
|
||||
connection, acceptErr := streamListener.Accept()
|
||||
if acceptErr != nil {
|
||||
streamDone <- acceptErr
|
||||
return
|
||||
}
|
||||
header := make([]byte, encryptedRTSPHeaderSize)
|
||||
if _, readErr := io.ReadFull(connection, header); readErr != nil {
|
||||
_ = connection.Close()
|
||||
streamDone <- readErr
|
||||
return
|
||||
}
|
||||
length := binary.BigEndian.Uint32(header[:4]) & 0x7fffffff
|
||||
frame := make([]byte, encryptedRTSPHeaderSize+int(length))
|
||||
copy(frame, header)
|
||||
if _, readErr := io.ReadFull(connection, frame[encryptedRTSPHeaderSize:]); readErr != nil {
|
||||
_ = connection.Close()
|
||||
streamDone <- readErr
|
||||
return
|
||||
}
|
||||
plaintext, decryptErr := codec.OpenClient(frame)
|
||||
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
|
||||
return
|
||||
}
|
||||
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
|
||||
return
|
||||
}
|
||||
}
|
||||
streamDone <- nil
|
||||
}()
|
||||
management := httptest.NewUnstartedServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) {
|
||||
if request.TLS == nil || len(request.TLS.PeerCertificates) != 1 {
|
||||
@@ -113,6 +184,8 @@ func TestNativeApolloSetupLaunchesInventoriedApplicationWithEncryptedRTSP(t *tes
|
||||
return
|
||||
}
|
||||
switch request.URL.Path {
|
||||
case "/serverinfo":
|
||||
_, _ = response.Write([]byte("<root><uniqueid>apollo-server</uniqueid></root>"))
|
||||
case "/applist":
|
||||
if request.URL.Query().Get("uniqueid") != "paired-client" {
|
||||
http.Error(response, "wrong client", http.StatusBadRequest)
|
||||
@@ -130,8 +203,38 @@ func TestNativeApolloSetupLaunchesInventoriedApplicationWithEncryptedRTSP(t *tes
|
||||
http.Error(response, "bad key", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
keyID, keyIDErr := strconv.ParseUint(query.Get("rikeyid"), 10, 32)
|
||||
if keyIDErr != nil {
|
||||
http.Error(response, "bad key id", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
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>"))
|
||||
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:
|
||||
http.NotFound(response, request)
|
||||
}
|
||||
@@ -154,19 +257,681 @@ func TestNativeApolloSetupLaunchesInventoriedApplicationWithEncryptedRTSP(t *tes
|
||||
ProviderIdentity: "apollo-server#sha256:" + hex.EncodeToString(pinned[:]), PolicyVersionID: "policy-1", ApplicationID: "42", ClientID: "paired-client",
|
||||
ManagementHost: managementHost, ManagementPort: managementPort, StreamHost: streamHost, StreamPort: streamPort,
|
||||
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()
|
||||
response, err := backend.Setup(context.Background(), LaunchRequest{SessionID: "session-1", ProviderProfile: ProviderProfileApollo, ProviderWork: work})
|
||||
if err != nil {
|
||||
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)
|
||||
}
|
||||
if err := <-streamDone; err != nil {
|
||||
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 {
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -97,6 +97,15 @@ func (c *ControlPlaneClient) ReportProviderState(ctx context.Context, state prot
|
||||
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) {
|
||||
if c == nil || c.HTTPClient == nil || c.BaseURL == "" {
|
||||
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 _ ClipboardAuditReporter = (*ControlPlaneClient)(nil)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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) {
|
||||
seed, _ := EncodeInputEvent(InputEvent{Sequence: 1, Device: "keyboard", Code: 7, Pressed: true})
|
||||
f.Add(seed)
|
||||
f.Add([]byte("INP1"))
|
||||
f.Add([]byte("VGI1"))
|
||||
f.Fuzz(func(t *testing.T, data []byte) {
|
||||
_, _ = DecodeInputEvent(data)
|
||||
})
|
||||
}
|
||||
|
||||
func TestInputEventUsesFixedProtocolVGI1Vector(t *testing.T) {
|
||||
encoded, err := EncodeInputEvent(InputEvent{Sequence: 7, Device: "keyboard", Code: 30, Pressed: true, Payload: []byte{2}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
const expected = "5647493101040102001e"
|
||||
if hex.EncodeToString(encoded) != expected {
|
||||
t.Fatalf("EncodeInputEvent() = %x, want %s", encoded, expected)
|
||||
}
|
||||
decoded, err := DecodeInputEvent(encoded)
|
||||
if err != nil || decoded.Sequence != 0 || decoded.Device != "keyboard" || decoded.Code != 30 || !decoded.Pressed || string(decoded.Payload) != string([]byte{2}) {
|
||||
t.Fatalf("DecodeInputEvent() = %#v, %v", decoded, err)
|
||||
}
|
||||
}
|
||||
|
||||
func 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) {
|
||||
capabilities := DefaultCapabilities()
|
||||
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) {
|
||||
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) {
|
||||
@@ -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()}
|
||||
admission := &oneTimeAdmission{authority: authority, released: make(chan struct{})}
|
||||
reporter := &recordingProviderStateReporter{}
|
||||
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: "gateway-1", Capabilities: DefaultCapabilities(), ProviderCapabilities: DefaultCapabilities(), Admission: admission, ProviderStateReporter: reporter, 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 {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -263,6 +255,44 @@ func TestAdmissionQUICMTLSRelayAndCleanup(t *testing.T) {
|
||||
t.Fatalf("unexpected media channel %d", frame.Channel)
|
||||
}
|
||||
}
|
||||
metrics := server.Metrics()
|
||||
if metrics.AdmittedSessions != 1 || metrics.MediaPackets < 2 || metrics.MediaBytes == 0 || metrics.ProcessingSamples < 2 || metrics.ProviderState != 2 {
|
||||
t.Fatalf("observed egress telemetry = %#v", metrics)
|
||||
}
|
||||
fakeSession, ok := fake.LastSession().(*fakeSession)
|
||||
if !ok {
|
||||
t.Fatal("fake provider session type")
|
||||
}
|
||||
fakeSession.mu.Lock()
|
||||
fakeSession.clipboard = "host clipboard"
|
||||
fakeSession.mu.Unlock()
|
||||
fakeSession.EmitEvent(ProviderEvent{Kind: ProviderEventRumble, Payload: []byte{1, 0x12, 0x34, 0x56, 0x78}})
|
||||
clipboardCtx, clipboardCancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
deliveredClipboard, clipboardErr := client.ReceiveClipboard(clipboardCtx)
|
||||
clipboardCancel()
|
||||
if clipboardErr != nil || deliveredClipboard.Direction != "provider_to_client" || deliveredClipboard.Text != "host clipboard" || deliveredClipboard.Encoding != "utf-8" {
|
||||
t.Fatalf("provider clipboard = %#v, %v", deliveredClipboard, clipboardErr)
|
||||
}
|
||||
if audits := reporter.Audits(); len(audits) == 0 || audits[0].Direction != "provider_to_client" || audits[0].Outcome != "forwarded" || audits[0].Reason != "forwarded" || audits[0].TextBytes != int64(len("host clipboard")) {
|
||||
t.Fatalf("clipboard audits = %#v", audits)
|
||||
}
|
||||
eventCtx, eventCancel := context.WithTimeout(context.Background(), time.Second)
|
||||
event, eventErr := client.ReceiveProviderEvent(eventCtx)
|
||||
eventCancel()
|
||||
if eventErr != nil || event.Kind != ProviderEventRumble || string(event.Payload) != string([]byte{1, 0x12, 0x34, 0x56, 0x78}) {
|
||||
t.Fatalf("provider event = %#v, %v", event, eventErr)
|
||||
}
|
||||
if err := client.SendClipboard(protocol.GatewayClipboardText{Direction: "client_to_provider", Text: "clipboard", Encoding: "utf-8", LoopToken: "abcdefghijklmnop"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
select {
|
||||
case value := <-fakeSession.clipboardWrites:
|
||||
if value != "clipboard" {
|
||||
t.Fatalf("provider clipboard = %q", value)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("gateway did not forward clipboard")
|
||||
}
|
||||
if err := client.SendInput(InputEvent{Sequence: 1, Device: "keyboard", Code: 7, Pressed: true}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -331,6 +361,7 @@ type oneTimeAdmission struct {
|
||||
type recordingProviderStateReporter struct {
|
||||
mu sync.Mutex
|
||||
states []protocol.ProviderState
|
||||
audits []protocol.GatewayClipboardAudit
|
||||
}
|
||||
|
||||
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...)
|
||||
}
|
||||
|
||||
func (r *recordingProviderStateReporter) ReportClipboardAudit(_ context.Context, audit protocol.GatewayClipboardAudit) error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.audits = append(r.audits, audit)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *recordingProviderStateReporter) Audits() []protocol.GatewayClipboardAudit {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
return append([]protocol.GatewayClipboardAudit(nil), r.audits...)
|
||||
}
|
||||
|
||||
func (a *oneTimeAdmission) Admit(context.Context, protocol.TunnelAdmissionRequest) (protocol.SessionAuthority, error) {
|
||||
if !a.used.CompareAndSwap(false, true) {
|
||||
return protocol.SessionAuthority{}, ErrAdmissionRejected
|
||||
@@ -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,
|
||||
StreamHost: "apollo.test", StreamPort: 47984, ClientCertificatePem: "certificate",
|
||||
ClientPrivateKeyPem: "private-key", ServerCertificatePem: "server-certificate",
|
||||
ClipboardPolicy: protocol.ClipboardPolicy{ClientToProviderEnabled: true, ProviderToClientEnabled: true, MaxTextBytes: 65536, MaxUpdatesPerMinute: 30},
|
||||
}, nil
|
||||
}
|
||||
|
||||
|
||||
+112
-18
@@ -3,36 +3,130 @@ package gateway
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
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) {
|
||||
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
|
||||
}
|
||||
encoded := make([]byte, gatewayInputHeaderSize+4)
|
||||
copy(encoded, "VGI1")
|
||||
encoded[4], encoded[5], encoded[6] = gatewayInputKeyboard, 4, 0
|
||||
if event.Pressed {
|
||||
encoded[6] = 1
|
||||
}
|
||||
if len(event.Payload) == 1 {
|
||||
encoded[7] = event.Payload[0]
|
||||
}
|
||||
binary.BigEndian.PutUint16(encoded[8:10], uint16(event.Code))
|
||||
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
|
||||
}
|
||||
encoded := make([]byte, 16+len(event.Device)+len(event.Payload))
|
||||
copy(encoded[:4], "INP1")
|
||||
binary.BigEndian.PutUint32(encoded[4:8], event.Sequence)
|
||||
binary.BigEndian.PutUint32(encoded[8:12], uint32(event.Code))
|
||||
if event.Pressed {
|
||||
encoded[12] = 1
|
||||
}
|
||||
encoded[13] = byte(len(event.Device))
|
||||
binary.BigEndian.PutUint16(encoded[14:16], uint16(len(event.Payload)))
|
||||
copy(encoded[16:16+len(event.Device)], event.Device)
|
||||
copy(encoded[16+len(event.Device):], event.Payload)
|
||||
return encoded, nil
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
deviceLength := int(data[13])
|
||||
payloadLength := int(binary.BigEndian.Uint16(data[14:16]))
|
||||
if deviceLength == 0 || deviceLength > 64 || payloadLength > 1024 || len(data) != 16+deviceLength+payloadLength {
|
||||
kind, body := data[4], data[gatewayInputHeaderSize:]
|
||||
switch kind {
|
||||
case gatewayInputKeyboard:
|
||||
if len(body) != 4 || body[0] > 1 || binary.BigEndian.Uint16(body[2:4]) == 0 {
|
||||
return InputEvent{}, ErrInputMalformed
|
||||
}
|
||||
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
|
||||
}
|
||||
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
|
||||
}
|
||||
|
||||
func anyNonzero(data []byte) bool {
|
||||
for _, value := range data {
|
||||
if value != 0 {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
+89
-54
@@ -2,7 +2,6 @@ package gateway
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"encoding/xml"
|
||||
"errors"
|
||||
"fmt"
|
||||
@@ -159,36 +158,6 @@ func ParseRTSPResponse(data []byte) (RTSPResponse, error) {
|
||||
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 {
|
||||
SessionID string
|
||||
Capabilities protocol.CapabilityProfile
|
||||
@@ -207,9 +176,35 @@ type InputEvent struct {
|
||||
|
||||
type Feedback struct {
|
||||
Sequence uint32
|
||||
Kind FeedbackKind
|
||||
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 {
|
||||
Start(context.Context, LaunchRequest) (ProviderSession, error)
|
||||
}
|
||||
@@ -218,9 +213,12 @@ type ProviderSession interface {
|
||||
Ready(context.Context) error
|
||||
Video() <-chan []byte
|
||||
Audio() <-chan []byte
|
||||
Events() <-chan ProviderEvent
|
||||
Input(context.Context, InputEvent) 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
|
||||
Terminate(context.Context) error
|
||||
State() protocol.ProviderState
|
||||
@@ -357,16 +355,18 @@ func (f *FakeApollo) Setup(context.Context, LaunchRequest) ([]byte, error) {
|
||||
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\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) {
|
||||
session := &fakeSession{
|
||||
failure: f.config.Failure,
|
||||
video: make(chan []byte, 16),
|
||||
audio: make(chan []byte, 16),
|
||||
state: protocol.ProviderState{Version: "1", SessionID: request.SessionID, State: ProviderStateStarting, Channels: []string{"video", "audio", "input", "feedback"}},
|
||||
pressed: make(map[string]struct{}),
|
||||
failure: f.config.Failure,
|
||||
video: 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"}},
|
||||
pressed: make(map[string]struct{}),
|
||||
}
|
||||
for _, payload := range f.config.Video {
|
||||
session.EmitVideo(payload)
|
||||
@@ -402,16 +402,19 @@ func (f *FakeApollo) DisconnectProvider() {
|
||||
}
|
||||
|
||||
type fakeSession struct {
|
||||
mu sync.Mutex
|
||||
failure FakeFailure
|
||||
video chan []byte
|
||||
audio chan []byte
|
||||
state protocol.ProviderState
|
||||
pressed map[string]struct{}
|
||||
inputs []InputEvent
|
||||
feedback []Feedback
|
||||
releaseAll int
|
||||
closeOnce sync.Once
|
||||
mu sync.Mutex
|
||||
failure FakeFailure
|
||||
video chan []byte
|
||||
audio chan []byte
|
||||
events chan ProviderEvent
|
||||
state protocol.ProviderState
|
||||
pressed map[string]struct{}
|
||||
inputs []InputEvent
|
||||
feedback []Feedback
|
||||
clipboard string
|
||||
clipboardWrites chan string
|
||||
releaseAll int
|
||||
closeOnce sync.Once
|
||||
}
|
||||
|
||||
func (s *fakeSession) Ready(ctx context.Context) error {
|
||||
@@ -428,8 +431,16 @@ func (s *fakeSession) Ready(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *fakeSession) Video() <-chan []byte { return s.video }
|
||||
func (s *fakeSession) Audio() <-chan []byte { return s.audio }
|
||||
func (s *fakeSession) Video() <-chan []byte { return s.video }
|
||||
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) {
|
||||
select {
|
||||
@@ -475,13 +486,33 @@ func (s *fakeSession) Feedback(_ context.Context, feedback Feedback) error {
|
||||
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()
|
||||
defer s.mu.Unlock()
|
||||
if s.state.State == ProviderStateTerminated {
|
||||
return ErrProviderTerminated
|
||||
if s.state.State != ProviderStateReady {
|
||||
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
|
||||
}
|
||||
|
||||
@@ -528,6 +559,10 @@ func (s *fakeSession) State() protocol.ProviderState {
|
||||
return s.state
|
||||
}
|
||||
|
||||
func (s *fakeSession) Telemetry() ProviderTelemetry {
|
||||
return ProviderTelemetry{State: s.State().State}
|
||||
}
|
||||
|
||||
func (s *fakeSession) Disconnect() {
|
||||
s.mu.Lock()
|
||||
s.state.State = ProviderStateDisconnected
|
||||
|
||||
+180
-15
@@ -2,33 +2,114 @@ package gateway
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Metrics struct {
|
||||
ActiveSessions atomic.Int64
|
||||
AdmissionRejects atomic.Uint64
|
||||
MediaDrops atomic.Uint64
|
||||
ProviderErrors atomic.Uint64
|
||||
InputRejected atomic.Uint64
|
||||
ActiveSessions atomic.Int64
|
||||
AdmittedSessions atomic.Uint64
|
||||
AdmissionRejects atomic.Uint64
|
||||
Reconnects atomic.Uint64
|
||||
DrainTransitions 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
|
||||
InputRejected atomic.Uint64
|
||||
ControlRTTNanos atomic.Uint64
|
||||
ControlJitterNanos atomic.Uint64
|
||||
ControlLossPPM atomic.Uint64
|
||||
PendingReliable atomic.Uint64
|
||||
ProviderState atomic.Uint64
|
||||
}
|
||||
|
||||
type MetricsSnapshot struct {
|
||||
ActiveSessions int64
|
||||
AdmissionRejects uint64
|
||||
MediaDrops uint64
|
||||
ProviderErrors uint64
|
||||
InputRejected uint64
|
||||
ActiveSessions int64
|
||||
AdmittedSessions uint64
|
||||
AdmissionRejects uint64
|
||||
Reconnects uint64
|
||||
DrainTransitions uint64
|
||||
MediaDrops uint64
|
||||
MediaPackets uint64
|
||||
MediaBytes uint64
|
||||
QueueDelayNanos uint64
|
||||
ProcessingDelayNanos uint64
|
||||
ProcessingSamples uint64
|
||||
PacingDelayNanos uint64
|
||||
ProviderErrors uint64
|
||||
InputRejected uint64
|
||||
ControlRTTNanos uint64
|
||||
ControlJitterNanos uint64
|
||||
ControlLossPPM uint64
|
||||
PendingReliable uint64
|
||||
ProviderState uint64
|
||||
}
|
||||
|
||||
func (m *Metrics) Snapshot() MetricsSnapshot {
|
||||
return MetricsSnapshot{
|
||||
ActiveSessions: m.ActiveSessions.Load(),
|
||||
AdmissionRejects: m.AdmissionRejects.Load(),
|
||||
MediaDrops: m.MediaDrops.Load(),
|
||||
ProviderErrors: m.ProviderErrors.Load(),
|
||||
InputRejected: m.InputRejected.Load(),
|
||||
ActiveSessions: m.ActiveSessions.Load(),
|
||||
AdmittedSessions: m.AdmittedSessions.Load(),
|
||||
AdmissionRejects: m.AdmissionRejects.Load(),
|
||||
Reconnects: m.Reconnects.Load(),
|
||||
DrainTransitions: m.DrainTransitions.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(),
|
||||
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
|
||||
}
|
||||
|
||||
// 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 {
|
||||
if kbps < 1 {
|
||||
return &Pacer{}
|
||||
|
||||
+1
-1
@@ -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
|
||||
|
||||
+393
-67
@@ -19,9 +19,10 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
defaultHelloLimit = 16 * 1024
|
||||
defaultControlLimit = 128 * 1024
|
||||
applicationError = quic.ApplicationErrorCode(0x100)
|
||||
defaultHelloLimit = 16 * 1024
|
||||
defaultControlLimit = 128 * 1024
|
||||
clientControlBacklog = 64
|
||||
applicationError = quic.ApplicationErrorCode(0x100)
|
||||
)
|
||||
|
||||
var (
|
||||
@@ -53,25 +54,31 @@ type ProviderStateReporter interface {
|
||||
ReportProviderState(context.Context, protocol.ProviderState) error
|
||||
}
|
||||
|
||||
type ClipboardAuditReporter interface {
|
||||
ReportClipboardAudit(context.Context, protocol.GatewayClipboardAudit) error
|
||||
}
|
||||
|
||||
type ServerConfig struct {
|
||||
ListenAddress string
|
||||
TLSConfig *tls.Config
|
||||
QUICConfig *quic.Config
|
||||
GatewayID string
|
||||
Capabilities protocol.CapabilityProfile
|
||||
ProviderCapabilities protocol.CapabilityProfile
|
||||
Admission Admission
|
||||
ProviderStateReporter ProviderStateReporter
|
||||
Provider Provider
|
||||
ProviderProfile string
|
||||
ProviderIdentity string
|
||||
PacerKbps int64
|
||||
ListenAddress string
|
||||
TLSConfig *tls.Config
|
||||
QUICConfig *quic.Config
|
||||
GatewayID string
|
||||
Capabilities protocol.CapabilityProfile
|
||||
ProviderCapabilities protocol.CapabilityProfile
|
||||
Admission Admission
|
||||
ProviderStateReporter ProviderStateReporter
|
||||
ClipboardAuditReporter ClipboardAuditReporter
|
||||
Provider Provider
|
||||
ProviderProfile string
|
||||
ProviderIdentity string
|
||||
PacerKbps int64
|
||||
}
|
||||
|
||||
type Server struct {
|
||||
listener *quic.Listener
|
||||
config ServerConfig
|
||||
metrics *Metrics
|
||||
pacer *fairPacer
|
||||
mu sync.Mutex
|
||||
sessions map[*gatewaySession]struct{}
|
||||
draining atomic.Bool
|
||||
@@ -115,7 +122,7 @@ func NewServer(config ServerConfig) (*Server, error) {
|
||||
if err != nil {
|
||||
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 {
|
||||
@@ -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) BeginDrain() {
|
||||
s.draining.Store(true)
|
||||
if s.draining.CompareAndSwap(false, true) {
|
||||
s.metrics.DrainTransitions.Add(1)
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
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 {
|
||||
_ = s.config.Admission.Release(context.Background(), authority)
|
||||
_ = 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)
|
||||
return
|
||||
}
|
||||
session := newGatewaySession(s, connection, stream, request, authority, providerSession)
|
||||
session := newGatewaySession(s, connection, stream, request, authority, providerSession, clipboard)
|
||||
s.addSession(session)
|
||||
s.metrics.ActiveSessions.Add(1)
|
||||
s.metrics.AdmittedSessions.Add(1)
|
||||
if authority.ReconnectSequence > 0 {
|
||||
s.metrics.Reconnects.Add(1)
|
||||
}
|
||||
defer func() {
|
||||
s.removeSession(session)
|
||||
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 {
|
||||
if s.config.ProviderStateReporter == nil {
|
||||
return nil
|
||||
}
|
||||
if err := state.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
s.metrics.observeProviderState(state.State)
|
||||
if s.config.ProviderStateReporter == nil {
|
||||
return nil
|
||||
}
|
||||
return s.config.ProviderStateReporter.ReportProviderState(ctx, state)
|
||||
}
|
||||
|
||||
@@ -315,25 +340,27 @@ func (s *Server) removeSession(session *gatewaySession) {
|
||||
}
|
||||
|
||||
type gatewaySession struct {
|
||||
server *Server
|
||||
connection *quic.Conn
|
||||
control *quic.Stream
|
||||
request protocol.TunnelAdmissionRequest
|
||||
authority protocol.SessionAuthority
|
||||
provider ProviderSession
|
||||
pacer *Pacer
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
cleanupOnce sync.Once
|
||||
inputMu sync.Mutex
|
||||
pressed map[string]struct{}
|
||||
sequence atomic.Uint32
|
||||
result chan error
|
||||
server *Server
|
||||
connection *quic.Conn
|
||||
control *quic.Stream
|
||||
request protocol.TunnelAdmissionRequest
|
||||
authority protocol.SessionAuthority
|
||||
provider ProviderSession
|
||||
clipboard *clipboardGate
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
cleanupOnce sync.Once
|
||||
inputMu sync.Mutex
|
||||
controlWriteMu sync.Mutex
|
||||
pressed map[string]struct{}
|
||||
sequence atomic.Uint32
|
||||
mediaDrops uint64
|
||||
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())
|
||||
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() {
|
||||
@@ -350,6 +377,11 @@ func (s *gatewaySession) run() {
|
||||
go s.controlLoop()
|
||||
go s.datagramLoop()
|
||||
go s.mediaLoop()
|
||||
go s.providerEventLoop()
|
||||
go s.providerTelemetryLoop()
|
||||
if s.clipboard != nil && s.clipboard.policy.ProviderToClientEnabled {
|
||||
go s.clipboardLoop()
|
||||
}
|
||||
select {
|
||||
case <-timer.C:
|
||||
s.server.metrics.InputRejected.Add(1)
|
||||
@@ -359,6 +391,96 @@ func (s *gatewaySession) run() {
|
||||
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() {
|
||||
for {
|
||||
data, err := readWire(s.control, defaultControlLimit)
|
||||
@@ -378,12 +500,32 @@ func (s *gatewaySession) controlLoop() {
|
||||
}
|
||||
switch frame.FlowID {
|
||||
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
|
||||
return
|
||||
}
|
||||
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
|
||||
return
|
||||
}
|
||||
@@ -408,19 +550,13 @@ func (s *gatewaySession) datagramLoop() {
|
||||
}
|
||||
switch frame.Channel {
|
||||
case ChannelInput:
|
||||
if err := s.handleInput(frame.Payload); err != nil {
|
||||
if err := s.handleInput(frame.Payload, frame.Sequence); err != nil {
|
||||
s.result <- err
|
||||
return
|
||||
}
|
||||
case ChannelText:
|
||||
if len(frame.Payload) > 4096 {
|
||||
s.result <- ErrFramePayloadLimit
|
||||
return
|
||||
}
|
||||
if err := s.provider.Feedback(s.ctx, Feedback{Sequence: frame.Sequence, Payload: append([]byte(nil), frame.Payload...)}); err != nil {
|
||||
s.result <- err
|
||||
return
|
||||
}
|
||||
s.result <- ErrFrameChannel
|
||||
return
|
||||
default:
|
||||
s.result <- ErrFrameChannel
|
||||
return
|
||||
@@ -458,6 +594,7 @@ func (s *gatewaySession) mediaLoop() {
|
||||
}
|
||||
|
||||
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)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -467,37 +604,122 @@ func (s *gatewaySession) sendMedia(channel byte, payload []byte) error {
|
||||
if err != nil {
|
||||
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
|
||||
}
|
||||
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 {
|
||||
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
|
||||
}
|
||||
|
||||
func (s *gatewaySession) handleControl(payload []byte) error {
|
||||
if len(payload) < 4 {
|
||||
return ErrProviderMalformed
|
||||
func (s *gatewaySession) sendControl(sequence uint32, payload []byte) error {
|
||||
if len(payload) > 1024 {
|
||||
return ErrFramePayloadLimit
|
||||
}
|
||||
switch string(payload[:4]) {
|
||||
case "TERM":
|
||||
return errors.New("client requested termination")
|
||||
case "RECN":
|
||||
return s.provider.Reconnect(s.ctx)
|
||||
case "FBRK":
|
||||
return s.provider.Feedback(s.ctx, Feedback{Payload: append([]byte(nil), 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)}
|
||||
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) 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:
|
||||
return ErrProviderMalformed
|
||||
return "provider"
|
||||
}
|
||||
}
|
||||
|
||||
func (s *gatewaySession) handleInput(payload []byte) error {
|
||||
func (s *gatewaySession) handleInput(payload []byte, sequence uint32) error {
|
||||
event, err := DecodeInputEvent(payload)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
event.Sequence = sequence
|
||||
if err := s.provider.Input(s.ctx, event); err != nil {
|
||||
s.server.metrics.InputRejected.Add(1)
|
||||
return err
|
||||
@@ -513,9 +735,17 @@ func (s *gatewaySession) handleInput(payload []byte) error {
|
||||
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() {
|
||||
s.cleanupOnce.Do(func() {
|
||||
s.cancel()
|
||||
s.server.pacer.remove(s.authority.SessionID)
|
||||
cleanupCtx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
defer cancel()
|
||||
releaseInputsErr := s.provider.ReleaseAll(cleanupCtx)
|
||||
@@ -613,9 +843,12 @@ func readWire(reader io.Reader, max int) ([]byte, error) {
|
||||
}
|
||||
|
||||
type Client struct {
|
||||
connection *quic.Conn
|
||||
control *quic.Stream
|
||||
Authority protocol.SessionAuthority
|
||||
connection *quic.Conn
|
||||
control *quic.Stream
|
||||
controlReadMu sync.Mutex
|
||||
controlWriteMu sync.Mutex
|
||||
pendingControl map[string][][]byte
|
||||
Authority protocol.SessionAuthority
|
||||
}
|
||||
|
||||
func Dial(ctx context.Context, address string, tlsConfig *tls.Config, request protocol.TunnelAdmissionRequest) (*Client, error) {
|
||||
@@ -658,7 +891,7 @@ func Dial(ctx context.Context, address string, tlsConfig *tls.Config, request pr
|
||||
}
|
||||
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 {
|
||||
@@ -671,15 +904,46 @@ func (c *Client) SendInput(event InputEvent) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeWire(c.control, encoded, defaultControlLimit)
|
||||
return c.writeControl(encoded)
|
||||
}
|
||||
|
||||
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)
|
||||
if err != nil {
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -691,6 +955,68 @@ func (c *Client) ReceiveFrame(ctx context.Context) (Frame, error) {
|
||||
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 {
|
||||
return c.connection.CloseWithError(applicationError, "client closed")
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user