diff --git a/cmd/verse-gateway/main.go b/cmd/verse-gateway/main.go index 0f56834..8eab735 100644 --- a/cmd/verse-gateway/main.go +++ b/cmd/verse-gateway/main.go @@ -57,7 +57,7 @@ func run() error { controlPlaneClient := gateway.NewControlPlaneClient(controlPlane, &http.Client{Transport: transport, Timeout: 5 * time.Second}) provider := gateway.NewApolloAdapter(gateway.NewNativeApolloBackend(), gateway.ProviderIdentity{}) capabilities := gateway.DefaultCapabilities() - server, err := gateway.NewServer(gateway.ServerConfig{ListenAddress: listen, TLSConfig: serverTLS, GatewayID: gatewayID, Capabilities: capabilities, ProviderCapabilities: capabilities, Admission: controlPlaneClient, ProviderStateReporter: controlPlaneClient, Provider: provider, PacerKbps: 100000}) + server, err := gateway.NewServer(gateway.ServerConfig{ListenAddress: listen, TLSConfig: serverTLS, GatewayID: gatewayID, Capabilities: capabilities, ProviderCapabilities: capabilities, Admission: controlPlaneClient, ProviderStateReporter: controlPlaneClient, ClipboardAuditReporter: controlPlaneClient, Provider: provider, PacerKbps: 100000}) if err != nil { return err } diff --git a/gateway/apollo_audio_fec.go b/gateway/apollo_audio_fec.go new file mode 100644 index 0000000..7f33493 --- /dev/null +++ b/gateway/apollo_audio_fec.go @@ -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]...) +} diff --git a/gateway/apollo_control.go b/gateway/apollo_control.go new file mode 100644 index 0000000..09db479 --- /dev/null +++ b/gateway/apollo_control.go @@ -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 +} diff --git a/gateway/apollo_enet.go b/gateway/apollo_enet.go new file mode 100644 index 0000000..fd99565 --- /dev/null +++ b/gateway/apollo_enet.go @@ -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) + } + }) +} diff --git a/gateway/apollo_enet_test.go b/gateway/apollo_enet_test.go new file mode 100644 index 0000000..9d49fa6 --- /dev/null +++ b/gateway/apollo_enet_test.go @@ -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)< 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 + } +} diff --git a/gateway/apollo_media.go b/gateway/apollo_media.go new file mode 100644 index 0000000..9e08cd3 --- /dev/null +++ b/gateway/apollo_media.go @@ -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 +} diff --git a/gateway/apollo_native.go b/gateway/apollo_native.go index 338157e..dab57ae 100644 --- a/gateway/apollo_native.go +++ b/gateway/apollo_native.go @@ -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 } } } diff --git a/gateway/apollo_native_test.go b/gateway/apollo_native_test.go index bcb3717..3c00824 100644 --- a/gateway/apollo_native_test.go +++ b/gateway/apollo_native_test.go @@ -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("apollo-server")) 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("rtspenc://" + streamListener.Addr().String() + "")) + case "/resume": + _, _ = response.Write([]byte("")) + case "/cancel": + if request.Method != http.MethodGet { + http.Error(response, "cancel must be GET", http.StatusMethodNotAllowed) + return + } + cancelCalls <- struct{}{} + _, _ = response.Write([]byte("1")) + 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 { diff --git a/gateway/apollo_parser_fuzz_test.go b/gateway/apollo_parser_fuzz_test.go new file mode 100644 index 0000000..ad68b6b --- /dev/null +++ b/gateway/apollo_parser_fuzz_test.go @@ -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) + }) +} diff --git a/gateway/apollo_rtsp_handshake.go b/gateway/apollo_rtsp_handshake.go new file mode 100644 index 0000000..b820914 --- /dev/null +++ b/gateway/apollo_rtsp_handshake.go @@ -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 +} diff --git a/gateway/apollo_video_fec.go b/gateway/apollo_video_fec.go new file mode 100644 index 0000000..6e995fa --- /dev/null +++ b/gateway/apollo_video_fec.go @@ -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 +} diff --git a/gateway/clipboard.go b/gateway/clipboard.go new file mode 100644 index 0000000..f929e4d --- /dev/null +++ b/gateway/clipboard.go @@ -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 +} diff --git a/gateway/clipboard_test.go b/gateway/clipboard_test.go new file mode 100644 index 0000000..6f63b91 --- /dev/null +++ b/gateway/clipboard_test.go @@ -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) + } +} diff --git a/gateway/control_plane.go b/gateway/control_plane.go index dc048d2..ecdf5f1 100644 --- a/gateway/control_plane.go +++ b/gateway/control_plane.go @@ -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) diff --git a/gateway/fair_pacer_test.go b/gateway/fair_pacer_test.go new file mode 100644 index 0000000..c4a1409 --- /dev/null +++ b/gateway/fair_pacer_test.go @@ -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 + } +} diff --git a/gateway/feedback.go b/gateway/feedback.go new file mode 100644 index 0000000..6c8bd0b --- /dev/null +++ b/gateway/feedback.go @@ -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 +} diff --git a/gateway/gateway_test.go b/gateway/gateway_test.go index fdf4168..b3c52ac 100644 --- a/gateway/gateway_test.go +++ b/gateway/gateway_test.go @@ -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 } diff --git a/gateway/input.go b/gateway/input.go index fc91c66..023cb71 100644 --- a/gateway/input.go +++ b/gateway/input.go @@ -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 } diff --git a/gateway/provider.go b/gateway/provider.go index 28eda6b..ea2b4e2 100644 --- a/gateway/provider.go +++ b/gateway/provider.go @@ -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 diff --git a/gateway/telemetry.go b/gateway/telemetry.go index a61825e..105449e 100644 --- a/gateway/telemetry.go +++ b/gateway/telemetry.go @@ -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{} diff --git a/gateway/testdata/rtsp-setup-response.txt b/gateway/testdata/rtsp-setup-response.txt index 8fc30e2..35c5624 100644 --- a/gateway/testdata/rtsp-setup-response.txt +++ b/gateway/testdata/rtsp-setup-response.txt @@ -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 diff --git a/gateway/transport.go b/gateway/transport.go index 4c11175..a0e5580 100644 --- a/gateway/transport.go +++ b/gateway/transport.go @@ -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") }