fix(gateway): enforce audited production traversal
This commit is contained in:
+61
-13
@@ -198,12 +198,16 @@ func pinnedApolloTLSConfig(work protocol.ProviderSessionWork) (*tls.Config, erro
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *NativeApolloBackend) Setup(ctx context.Context, request LaunchRequest) ([]byte, error) {
|
||||
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
|
||||
@@ -281,10 +285,11 @@ type nativeApolloSession struct {
|
||||
audioPing []byte
|
||||
videoPing []byte
|
||||
sessionID string
|
||||
video chan []byte
|
||||
audio chan []byte
|
||||
video chan ProviderMedia
|
||||
audio chan ProviderMedia
|
||||
events chan ProviderEvent
|
||||
mu sync.Mutex
|
||||
mediaMu sync.Mutex
|
||||
controlMu sync.Mutex
|
||||
state protocol.ProviderState
|
||||
pressed map[string]InputEvent
|
||||
@@ -299,10 +304,15 @@ type nativeApolloSession struct {
|
||||
allowApplicationTermination bool
|
||||
terminationErr error
|
||||
mediaDrops atomic.Uint64
|
||||
mediaQuiesced atomic.Bool
|
||||
mediaIngress atomic.Uint64
|
||||
mediaRecovered atomic.Uint64
|
||||
mediaEnqueued atomic.Uint64
|
||||
mediaQueueMaximum atomic.Uint64
|
||||
}
|
||||
|
||||
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{})}
|
||||
return &nativeApolloSession{sessionID: sessionID, video: make(chan ProviderMedia, 16), audio: make(chan ProviderMedia, 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) {
|
||||
@@ -405,8 +415,8 @@ func (s *nativeApolloSession) Ready(context.Context) error {
|
||||
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 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 {
|
||||
@@ -646,6 +656,7 @@ func (s *nativeApolloSession) handleApolloControlPayload(_ uint8, _ bool, payloa
|
||||
s.handleApolloDisconnect(ErrProviderMalformed)
|
||||
return
|
||||
}
|
||||
s.quiesceMedia()
|
||||
s.mu.Lock()
|
||||
s.state.State = ProviderStateTerminated
|
||||
s.mu.Unlock()
|
||||
@@ -686,6 +697,7 @@ func (s *nativeApolloSession) handleApolloDisconnect(err error) {
|
||||
if err == nil {
|
||||
return
|
||||
}
|
||||
s.quiesceMedia()
|
||||
s.disconnectOnce.Do(func() {
|
||||
s.mu.Lock()
|
||||
if s.state.State == ProviderStateTerminated {
|
||||
@@ -701,13 +713,45 @@ func (s *nativeApolloSession) handleApolloDisconnect(err error) {
|
||||
})
|
||||
}
|
||||
|
||||
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.mediaMu.Lock()
|
||||
defer s.mediaMu.Unlock()
|
||||
close(s.video)
|
||||
close(s.audio)
|
||||
})
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
s.mediaRecovered.Add(1)
|
||||
media := ProviderMedia{Payload: payload, ReceivedAt: receivedAt, EnqueuedAt: time.Now()}
|
||||
if pushLatest(output, media) {
|
||||
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) readUDPMedia() {
|
||||
if s.media == nil {
|
||||
close(s.readDone)
|
||||
@@ -716,14 +760,18 @@ func (s *nativeApolloSession) readUDPMedia() {
|
||||
}
|
||||
var readers sync.WaitGroup
|
||||
readers.Add(2)
|
||||
read := func(conn *net.UDPConn, output chan []byte, video bool) {
|
||||
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 {
|
||||
@@ -738,6 +786,10 @@ func (s *nativeApolloSession) readUDPMedia() {
|
||||
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])
|
||||
@@ -764,11 +816,7 @@ func (s *nativeApolloSession) readUDPMedia() {
|
||||
continue
|
||||
}
|
||||
for _, payload := range payloads {
|
||||
if len(payload) != 0 {
|
||||
if pushLatest(output, payload) {
|
||||
s.mediaDrops.Add(1)
|
||||
}
|
||||
}
|
||||
s.enqueueMedia(output, payload, receivedAt)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -781,7 +829,7 @@ func (s *nativeApolloSession) readUDPMedia() {
|
||||
}()
|
||||
}
|
||||
|
||||
func pushLatest(channel chan []byte, payload []byte) bool {
|
||||
func pushLatest[T any](channel chan T, payload T) bool {
|
||||
select {
|
||||
case channel <- payload:
|
||||
return false
|
||||
|
||||
Reference in New Issue
Block a user