fix(gateway): enforce audited production traversal

This commit is contained in:
sechmachine
2026-07-30 04:48:02 +07:00
parent d3852d15f3
commit 70667d5aee
17 changed files with 1368 additions and 317 deletions
+61 -13
View File
@@ -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