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

990 lines
28 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
)
// 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
}
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)
}
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) (*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 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
}
var readers sync.WaitGroup
readers.Add(2)
read := func(conn *net.UDPConn, output chan ProviderMedia, video bool) {
defer readers.Done()
buffer := make([]byte, apolloMediaMaximumPacket+1)
for {
if s.mediaQuiesced.Load() {
return
}
if err := conn.SetReadDeadline(time.Now().Add(250 * time.Millisecond)); err != nil {
return
}
count, err := conn.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)
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
}
var evicted bool
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(output, payload, receivedAt)
}
}
}
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[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)