1080 lines
30 KiB
Go
1080 lines
30 KiB
Go
package gateway
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"crypto/rand"
|
|
"crypto/sha256"
|
|
"crypto/tls"
|
|
"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"
|
|
)
|
|
|
|
const (
|
|
nativeApolloVideoQueuePackets = 16
|
|
nativeApolloVideoQueueBytes = 4 << 20
|
|
nativeApolloVideoQueueLatency = 250 * time.Millisecond
|
|
nativeApolloAudioQueuePackets = 16
|
|
nativeApolloEventQueuePackets = 16
|
|
nativeApolloVideoIngressSlots = 2048
|
|
nativeApolloVideoPacketBytes = apolloVideoHeaderSize + apolloVideoRawPacketSize
|
|
nativeApolloVideoReadBuffer = nativeApolloVideoIngressSlots * nativeApolloVideoPacketBytes
|
|
)
|
|
|
|
// NativeApolloBackend keeps provider sockets inside the gateway process. The
|
|
// session-scoped Server work is the sole source of provider endpoint and mTLS
|
|
// material; it is never serialized into a client manifest or authority.
|
|
type NativeApolloBackend struct {
|
|
Dialer *net.Dialer
|
|
|
|
mu sync.Mutex
|
|
pending map[string]*apolloRTSPSetup
|
|
configureMedia func(*apolloMediaCodec)
|
|
}
|
|
|
|
func NewNativeApolloBackend() *NativeApolloBackend {
|
|
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) {
|
|
work := request.ProviderWork
|
|
if err := work.Validate(); err != nil || validateApolloStreamPolicy(work.StreamPolicy) != nil ||
|
|
request.SessionID == "" || request.SessionID != work.SessionID || work.ProviderProfile != ProviderProfileApollo {
|
|
return nil, ErrProviderMalformed
|
|
}
|
|
client, err := newPinnedApolloHTTPClient(work)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return apolloGet(ctx, client, work, "/serverinfo", nil)
|
|
}
|
|
|
|
func newPinnedApolloHTTPClient(work protocol.ProviderSessionWork) (*http.Client, error) {
|
|
tlsConfig, err := pinnedApolloTLSConfig(work)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
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) {
|
|
endpoint := url.URL{Scheme: "https", Host: net.JoinHostPort(work.ManagementHost, strconv.FormatInt(work.ManagementPort, 10)), Path: path}
|
|
endpoint.RawQuery = values.Encode()
|
|
request, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint.String(), nil)
|
|
if err != nil {
|
|
return nil, ErrProviderMalformed
|
|
}
|
|
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 management status %d", response.StatusCode)
|
|
}
|
|
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:") {
|
|
return nil, ErrProviderIdentity
|
|
}
|
|
pinned, err := hex.DecodeString(strings.TrimPrefix(identity.Fingerprint, "sha256:"))
|
|
if err != nil || len(pinned) != sha256.Size {
|
|
return nil, ErrProviderIdentity
|
|
}
|
|
certificate, err := tls.X509KeyPair([]byte(work.ClientCertificatePem), []byte(work.ClientPrivateKeyPem))
|
|
if err != nil {
|
|
return nil, ErrProviderIdentity
|
|
}
|
|
trust := x509.NewCertPool()
|
|
if !trust.AppendCertsFromPEM([]byte(work.ServerCertificatePem)) {
|
|
return nil, ErrProviderIdentity
|
|
}
|
|
return &tls.Config{
|
|
MinVersion: tls.VersionTLS13, Certificates: []tls.Certificate{certificate}, RootCAs: trust,
|
|
VerifyPeerCertificate: func(rawCertificates [][]byte, _ [][]*x509.Certificate) error {
|
|
if len(rawCertificates) == 0 {
|
|
return ErrProviderIdentity
|
|
}
|
|
digest := sha256.Sum256(rawCertificates[0])
|
|
if !bytes.Equal(digest[:], pinned) {
|
|
return ErrProviderIdentity
|
|
}
|
|
return nil
|
|
},
|
|
}, nil
|
|
}
|
|
|
|
func (b *NativeApolloBackend) Setup(ctx context.Context, request LaunchRequest, management []byte) ([]byte, error) {
|
|
work := request.ProviderWork
|
|
if err := work.Validate(); err != nil || validateApolloStreamPolicy(work.StreamPolicy) != nil ||
|
|
request.SessionID == "" || request.SessionID != work.SessionID || work.ProviderProfile != ProviderProfileApollo {
|
|
return nil, ErrProviderMalformed
|
|
}
|
|
info, err := ParseManagementXML(management)
|
|
if err != nil || validateApolloProviderStreamPolicy(info, work.StreamPolicy) != nil {
|
|
return nil, ErrProviderMalformed
|
|
}
|
|
client, err := newPinnedApolloHTTPClient(work)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
inventory, err := apolloGet(ctx, client, work, "/applist", url.Values{"uniqueid": []string{work.ClientID}})
|
|
if err != nil || !apolloInventoryContains(inventory, work.ApplicationID) {
|
|
return nil, ErrProviderMalformed
|
|
}
|
|
key := make([]byte, 16)
|
|
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
|
|
}
|
|
sessionPath, sessionValues := apolloSessionRequest(work, key, binary.BigEndian.Uint32(keyID[:]))
|
|
launch, err := apolloGet(ctx, client, work, sessionPath, sessionValues)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
launchResponse, err := parseApolloLaunchResponse(launch)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
streamURL, err := url.Parse(launchResponse.SessionURL)
|
|
if err != nil || streamURL.Scheme != "rtspenc" || streamURL.Hostname() != work.StreamHost {
|
|
return nil, ErrProviderMalformed
|
|
}
|
|
streamPort, err := strconv.ParseInt(streamURL.Port(), 10, 64)
|
|
if err != nil || streamPort != work.StreamPort {
|
|
return nil, ErrProviderMalformed
|
|
}
|
|
setup, response, err := b.performRTSPHandshake(ctx, work, key, binary.BigEndian.Uint32(keyID[:]), streamURL)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
b.mu.Lock()
|
|
b.pending[request.SessionID] = setup
|
|
b.mu.Unlock()
|
|
return response, nil
|
|
}
|
|
|
|
func (b *NativeApolloBackend) Open(ctx context.Context, request LaunchRequest, _ RTSPResponse) (ProviderSession, error) {
|
|
b.mu.Lock()
|
|
setup, ok := b.pending[request.SessionID]
|
|
delete(b.pending, request.SessionID)
|
|
b.mu.Unlock()
|
|
if !ok || setup == nil {
|
|
return nil, ErrProviderDisconnected
|
|
}
|
|
return newNativeApolloProviderSession(ctx, setup, b.configureMedia)
|
|
}
|
|
|
|
func readBounded(reader io.Reader, max int) ([]byte, error) {
|
|
data, err := io.ReadAll(io.LimitReader(reader, int64(max)+1))
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if len(data) > max {
|
|
return nil, ErrProviderMalformed
|
|
}
|
|
return data, nil
|
|
}
|
|
|
|
type nativeApolloSession 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 ProviderMedia
|
|
audio chan ProviderMedia
|
|
events chan ProviderEvent
|
|
mu sync.Mutex
|
|
eventMu sync.Mutex
|
|
mediaMu sync.Mutex
|
|
controlMu sync.Mutex
|
|
state protocol.ProviderState
|
|
pressed map[string]InputEvent
|
|
closeOnce sync.Once
|
|
disconnectOnce 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
|
|
mediaQuiesced atomic.Bool
|
|
mediaIngress atomic.Uint64
|
|
mediaRecovered atomic.Uint64
|
|
mediaEnqueued atomic.Uint64
|
|
mediaQueueMaximum atomic.Uint64
|
|
mediaQueueBytes atomic.Int64
|
|
mediaQueueMaximumBytes atomic.Uint64
|
|
mediaQueueSequence atomic.Uint64
|
|
}
|
|
|
|
func newNativeApolloSession(sessionID string) *nativeApolloSession {
|
|
return &nativeApolloSession{
|
|
sessionID: sessionID,
|
|
video: make(chan ProviderMedia, nativeApolloVideoQueuePackets),
|
|
audio: make(chan ProviderMedia, nativeApolloAudioQueuePackets),
|
|
events: make(chan ProviderEvent, nativeApolloEventQueuePackets),
|
|
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, configureMedia func(*apolloMediaCodec)) (*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
|
|
}
|
|
if configureMedia != nil {
|
|
configureMedia(media)
|
|
}
|
|
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 := videoConn.SetReadBuffer(nativeApolloVideoReadBuffer); err != nil {
|
|
_ = audioConn.Close()
|
|
_ = videoConn.Close()
|
|
peer.close(err)
|
|
return nil, fmt.Errorf("set Apollo video receive buffer: %w", 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 ProviderMedia { return s.video }
|
|
func (s *nativeApolloSession) Audio() <-chan ProviderMedia { return s.audio }
|
|
func (s *nativeApolloSession) Events() <-chan ProviderEvent { return s.events }
|
|
|
|
func (s *nativeApolloSession) Input(ctx context.Context, event InputEvent) error {
|
|
packet, err := encodeApolloInputEvent(event)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
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 {
|
|
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) 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 {
|
|
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 {
|
|
var cleanupErr error
|
|
s.mu.Lock()
|
|
disconnected := s.state.State == ProviderStateDisconnected
|
|
s.mu.Unlock()
|
|
s.closeOnce.Do(func() {
|
|
if err := s.ReleaseAll(ctx); err != nil {
|
|
cleanupErr = err
|
|
}
|
|
close(s.done)
|
|
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():
|
|
if cleanupErr == nil {
|
|
cleanupErr = ctx.Err()
|
|
}
|
|
}
|
|
if cleanupErr == nil {
|
|
s.closeMediaChannels()
|
|
if s.allowApplicationTermination && !disconnected {
|
|
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()
|
|
cleanupErr = s.terminationErr
|
|
if cleanupErr != nil {
|
|
s.state.State = ProviderStateCleanup
|
|
s.state.CleanupPending = true
|
|
} else if disconnected {
|
|
s.state.State = ProviderStateDisconnected
|
|
} else {
|
|
s.state.State = ProviderStateTerminated
|
|
}
|
|
s.mu.Unlock()
|
|
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) 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 reliable {
|
|
return s.enet.SendReliable(channel, encoded)
|
|
}
|
|
return s.enet.SendUnsequenced(channel, encoded)
|
|
}
|
|
|
|
func (s *nativeApolloSession) periodicApolloPing() {
|
|
ticker := time.NewTicker(100 * time.Millisecond)
|
|
defer ticker.Stop()
|
|
for {
|
|
select {
|
|
case <-s.done:
|
|
return
|
|
case <-ticker.C:
|
|
if err := s.writeApolloControl(apolloChannelGeneric, true, apolloControlTypePing, []byte{4, 0, 0, 0, 0, 0, 0, 0}); err != nil {
|
|
s.handleApolloDisconnect(err)
|
|
return
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
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.quiesceMedia()
|
|
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) {
|
|
if !s.enqueueProviderEvent(event) {
|
|
s.handleApolloDisconnect(ErrProviderMalformed)
|
|
}
|
|
}
|
|
|
|
func (s *nativeApolloSession) enqueueProviderEvent(event ProviderEvent) bool {
|
|
s.eventMu.Lock()
|
|
defer s.eventMu.Unlock()
|
|
select {
|
|
case s.events <- event:
|
|
return true
|
|
default:
|
|
}
|
|
if event.Kind != ProviderEventTerminated && event.Kind != ProviderEventDisconnected {
|
|
return false
|
|
}
|
|
select {
|
|
case <-s.events:
|
|
default:
|
|
return false
|
|
}
|
|
select {
|
|
case s.events <- event:
|
|
return true
|
|
default:
|
|
return false
|
|
}
|
|
}
|
|
|
|
func (s *nativeApolloSession) handleApolloDisconnect(err error) {
|
|
if err == nil {
|
|
return
|
|
}
|
|
s.quiesceMedia()
|
|
s.disconnectOnce.Do(func() {
|
|
s.mu.Lock()
|
|
if s.state.State == ProviderStateTerminated {
|
|
s.mu.Unlock()
|
|
return
|
|
}
|
|
s.state.State = ProviderStateDisconnected
|
|
s.mu.Unlock()
|
|
s.enqueueProviderEvent(ProviderEvent{Kind: ProviderEventDisconnected})
|
|
})
|
|
}
|
|
|
|
func (s *nativeApolloSession) quiesceMedia() {
|
|
if !s.mediaQuiesced.CompareAndSwap(false, true) {
|
|
return
|
|
}
|
|
if s.audioConn != nil {
|
|
_ = s.audioConn.Close()
|
|
}
|
|
if s.videoConn != nil {
|
|
_ = s.videoConn.Close()
|
|
}
|
|
}
|
|
|
|
func (s *nativeApolloSession) closeMediaChannels() {
|
|
s.channelsOnce.Do(func() {
|
|
s.mediaQuiesced.Store(true)
|
|
s.mediaMu.Lock()
|
|
defer s.mediaMu.Unlock()
|
|
for {
|
|
select {
|
|
case media := <-s.video:
|
|
if media.expiry != nil {
|
|
media.expiry.Stop()
|
|
}
|
|
media.releaseQueue()
|
|
default:
|
|
close(s.video)
|
|
close(s.audio)
|
|
return
|
|
}
|
|
}
|
|
})
|
|
}
|
|
|
|
func (s *nativeApolloSession) enqueueMedia(output chan ProviderMedia, payload []byte, receivedAt time.Time) bool {
|
|
s.mediaMu.Lock()
|
|
defer s.mediaMu.Unlock()
|
|
if len(payload) == 0 || s.mediaQuiesced.Load() {
|
|
return false
|
|
}
|
|
if output == s.video && len(payload) > maxCompleteFrameBytes {
|
|
s.mediaDrops.Add(1)
|
|
return false
|
|
}
|
|
s.mediaRecovered.Add(1)
|
|
media := ProviderMedia{Payload: payload, ReceivedAt: receivedAt, EnqueuedAt: time.Now()}
|
|
if output == s.video {
|
|
dropped := uint64(0)
|
|
for s.mediaQueueBytes.Load()+int64(len(payload)) > nativeApolloVideoQueueBytes {
|
|
select {
|
|
case replaced := <-output:
|
|
if replaced.expiry != nil {
|
|
replaced.expiry.Stop()
|
|
}
|
|
replaced.releaseQueue()
|
|
dropped++
|
|
default:
|
|
s.mediaDrops.Add(dropped + 1)
|
|
return false
|
|
}
|
|
}
|
|
media.queueID = s.mediaQueueSequence.Add(1)
|
|
media.accounting = &providerMediaQueueAccounting{
|
|
bytes: int64(len(payload)), total: &s.mediaQueueBytes,
|
|
}
|
|
currentBytes := uint64(s.mediaQueueBytes.Add(int64(len(payload))))
|
|
media.expiry = time.AfterFunc(nativeApolloVideoQueueLatency, func() {
|
|
s.expireVideo(media.queueID)
|
|
media.releaseQueue()
|
|
})
|
|
for maximum := s.mediaQueueMaximumBytes.Load(); currentBytes > maximum && !s.mediaQueueMaximumBytes.CompareAndSwap(maximum, currentBytes); maximum = s.mediaQueueMaximumBytes.Load() {
|
|
}
|
|
if dropped > 0 {
|
|
s.mediaDrops.Add(dropped)
|
|
}
|
|
}
|
|
dropped := false
|
|
select {
|
|
case output <- media:
|
|
default:
|
|
select {
|
|
case replaced := <-output:
|
|
if replaced.expiry != nil {
|
|
replaced.expiry.Stop()
|
|
}
|
|
replaced.releaseQueue()
|
|
default:
|
|
}
|
|
select {
|
|
case output <- media:
|
|
dropped = true
|
|
default:
|
|
if media.expiry != nil {
|
|
media.expiry.Stop()
|
|
}
|
|
media.releaseQueue()
|
|
dropped = true
|
|
}
|
|
}
|
|
if dropped {
|
|
s.mediaDrops.Add(1)
|
|
}
|
|
s.mediaEnqueued.Add(1)
|
|
depth := uint64(len(output))
|
|
for maximum := s.mediaQueueMaximum.Load(); depth > maximum && !s.mediaQueueMaximum.CompareAndSwap(maximum, depth); maximum = s.mediaQueueMaximum.Load() {
|
|
}
|
|
return true
|
|
}
|
|
|
|
func (s *nativeApolloSession) expireVideo(queueID uint64) {
|
|
s.mediaMu.Lock()
|
|
defer s.mediaMu.Unlock()
|
|
if s.mediaQuiesced.Load() {
|
|
return
|
|
}
|
|
retained := make([]ProviderMedia, 0, cap(s.video))
|
|
removed := false
|
|
for {
|
|
select {
|
|
case media := <-s.video:
|
|
if media.queueID == queueID {
|
|
removed = true
|
|
media.releaseQueue()
|
|
continue
|
|
}
|
|
retained = append(retained, media)
|
|
default:
|
|
for _, media := range retained {
|
|
s.video <- media
|
|
}
|
|
if removed {
|
|
s.mediaDrops.Add(1)
|
|
}
|
|
return
|
|
}
|
|
}
|
|
}
|
|
|
|
func (s *nativeApolloSession) readUDPMedia() {
|
|
if s.media == nil {
|
|
close(s.readDone)
|
|
s.closeMediaChannels()
|
|
return
|
|
}
|
|
videoIngress := newNativeApolloVideoIngress()
|
|
var workers sync.WaitGroup
|
|
workers.Add(3)
|
|
go func() {
|
|
defer workers.Done()
|
|
buffer := make([]byte, apolloMediaMaximumPacket+1)
|
|
for {
|
|
if s.mediaQuiesced.Load() {
|
|
return
|
|
}
|
|
if err := s.audioConn.SetReadDeadline(time.Now().Add(250 * time.Millisecond)); err != nil {
|
|
return
|
|
}
|
|
count, err := s.audioConn.Read(buffer)
|
|
receivedAt := time.Now()
|
|
if err != nil {
|
|
if networkErr, ok := err.(net.Error); ok && networkErr.Timeout() {
|
|
select {
|
|
case <-s.done:
|
|
return
|
|
default:
|
|
continue
|
|
}
|
|
}
|
|
return
|
|
}
|
|
if count > apolloMediaMaximumPacket {
|
|
continue
|
|
}
|
|
if s.mediaQuiesced.Load() {
|
|
return
|
|
}
|
|
s.mediaIngress.Add(1)
|
|
shard, openErr := s.media.OpenAudio(buffer[:count])
|
|
if openErr != nil {
|
|
continue
|
|
}
|
|
payloads, evicted, err := s.audioFEC.Add(s.media, shard)
|
|
if evicted {
|
|
s.mediaDrops.Add(1)
|
|
}
|
|
if err != nil {
|
|
continue
|
|
}
|
|
for _, payload := range payloads {
|
|
s.enqueueMedia(s.audio, payload, receivedAt)
|
|
}
|
|
}
|
|
}()
|
|
go func() {
|
|
defer workers.Done()
|
|
s.drainApolloVideo(videoIngress)
|
|
}()
|
|
go func() {
|
|
defer workers.Done()
|
|
s.processApolloVideo(videoIngress)
|
|
}()
|
|
go func() {
|
|
workers.Wait()
|
|
close(s.readDone)
|
|
s.closeMediaChannels()
|
|
}()
|
|
}
|
|
|
|
type nativeApolloVideoIngressSlot struct {
|
|
packet [apolloMediaMaximumPacket + 1]byte
|
|
count int
|
|
receivedAt time.Time
|
|
}
|
|
|
|
type nativeApolloVideoIngress struct {
|
|
slots [nativeApolloVideoIngressSlots]nativeApolloVideoIngressSlot
|
|
free chan uint16
|
|
ready chan uint16
|
|
scratch [apolloMediaMaximumPacket + 1]byte
|
|
}
|
|
|
|
func newNativeApolloVideoIngress() *nativeApolloVideoIngress {
|
|
ingress := &nativeApolloVideoIngress{
|
|
free: make(chan uint16, nativeApolloVideoIngressSlots),
|
|
ready: make(chan uint16, nativeApolloVideoIngressSlots),
|
|
}
|
|
for index := range nativeApolloVideoIngressSlots {
|
|
ingress.free <- uint16(index)
|
|
}
|
|
return ingress
|
|
}
|
|
|
|
func (s *nativeApolloSession) drainApolloVideo(ingress *nativeApolloVideoIngress) {
|
|
defer close(ingress.ready)
|
|
for {
|
|
if s.mediaQuiesced.Load() {
|
|
return
|
|
}
|
|
select {
|
|
case index := <-ingress.free:
|
|
slot := &ingress.slots[index]
|
|
count, err := s.videoConn.Read(slot.packet[:])
|
|
receivedAt := time.Now()
|
|
if err != nil {
|
|
ingress.free <- index
|
|
return
|
|
}
|
|
if count > apolloMediaMaximumPacket {
|
|
ingress.free <- index
|
|
continue
|
|
}
|
|
if s.mediaQuiesced.Load() {
|
|
ingress.free <- index
|
|
return
|
|
}
|
|
s.mediaIngress.Add(1)
|
|
slot.count, slot.receivedAt = count, receivedAt
|
|
ingress.ready <- index
|
|
default:
|
|
count, err := s.videoConn.Read(ingress.scratch[:])
|
|
if err != nil {
|
|
return
|
|
}
|
|
if count > apolloMediaMaximumPacket {
|
|
continue
|
|
}
|
|
if s.mediaQuiesced.Load() {
|
|
return
|
|
}
|
|
s.mediaIngress.Add(1)
|
|
s.mediaDrops.Add(1)
|
|
}
|
|
}
|
|
}
|
|
|
|
func (s *nativeApolloSession) processApolloVideo(ingress *nativeApolloVideoIngress) {
|
|
for index := range ingress.ready {
|
|
slot := &ingress.slots[index]
|
|
if !s.mediaQuiesced.Load() {
|
|
shard, err := s.media.OpenVideo(slot.packet[:slot.count])
|
|
if err == nil {
|
|
payload, fecErr := s.videoFEC.Add(shard)
|
|
if fecErr == nil && len(payload) > 0 {
|
|
s.enqueueMedia(s.video, payload, slot.receivedAt)
|
|
}
|
|
}
|
|
}
|
|
slot.count, slot.receivedAt = 0, time.Time{}
|
|
ingress.free <- index
|
|
}
|
|
}
|
|
|
|
func pushLatest[T any](channel chan T, payload T) bool {
|
|
select {
|
|
case channel <- payload:
|
|
return false
|
|
default:
|
|
select {
|
|
case <-channel:
|
|
default:
|
|
}
|
|
select {
|
|
case channel <- payload:
|
|
return true
|
|
default:
|
|
return true
|
|
}
|
|
}
|
|
}
|
|
|
|
var _ ApolloBackend = (*NativeApolloBackend)(nil)
|
|
var _ ProviderSession = (*nativeApolloSession)(nil)
|