fix(gateway): bound fake media shutdown

This commit is contained in:
sechmachine
2026-07-29 22:13:32 +07:00
parent 994d1f00d1
commit 8826f5c804
3 changed files with 73 additions and 7 deletions
+2 -1
View File
@@ -465,6 +465,7 @@ type gatewayTransportHarness struct {
session *fakeSession session *fakeSession
admission *oneTimeAdmission admission *oneTimeAdmission
reporter *recordingProviderStateReporter reporter *recordingProviderStateReporter
server *Server
} }
func newGatewayTransportHarness(t *testing.T) gatewayTransportHarness { func newGatewayTransportHarness(t *testing.T) gatewayTransportHarness {
@@ -500,7 +501,7 @@ func newGatewayTransportHarness(t *testing.T) gatewayTransportHarness {
t.Errorf("serve: %v", err) t.Errorf("serve: %v", err)
} }
}) })
return gatewayTransportHarness{client: client, session: session, admission: admission, reporter: reporter} return gatewayTransportHarness{client: client, session: session, admission: admission, reporter: reporter, server: server}
} }
func (h gatewayTransportHarness) waitReleased(t *testing.T) { func (h gatewayTransportHarness) waitReleased(t *testing.T) {
+26 -6
View File
@@ -443,20 +443,42 @@ func (s *fakeSession) EmitEvent(event ProviderEvent) {
} }
func (s *fakeSession) EmitVideo(payload []byte) { func (s *fakeSession) EmitVideo(payload []byte) {
s.mu.Lock()
defer s.mu.Unlock()
if s.state.State == ProviderStateTerminating || s.state.State == ProviderStateTerminated || s.state.State == ProviderStateDisconnected {
return
}
select { select {
case s.video <- append([]byte(nil), payload...): case s.video <- append([]byte(nil), payload...):
default: default:
<-s.video select {
s.video <- append([]byte(nil), payload...) case <-s.video:
default:
}
select {
case s.video <- append([]byte(nil), payload...):
default:
}
} }
} }
func (s *fakeSession) EmitAudio(payload []byte) { func (s *fakeSession) EmitAudio(payload []byte) {
s.mu.Lock()
defer s.mu.Unlock()
if s.state.State == ProviderStateTerminating || s.state.State == ProviderStateTerminated || s.state.State == ProviderStateDisconnected {
return
}
select { select {
case s.audio <- append([]byte(nil), payload...): case s.audio <- append([]byte(nil), payload...):
default: default:
<-s.audio select {
s.audio <- append([]byte(nil), payload...) case <-s.audio:
default:
}
select {
case s.audio <- append([]byte(nil), payload...):
default:
}
} }
} }
@@ -542,12 +564,10 @@ func (s *fakeSession) Terminate(ctx context.Context) error {
return nil return nil
} }
s.state.State = ProviderStateTerminating s.state.State = ProviderStateTerminating
s.mu.Unlock()
s.closeOnce.Do(func() { s.closeOnce.Do(func() {
close(s.video) close(s.video)
close(s.audio) close(s.audio)
}) })
s.mu.Lock()
s.state.State = ProviderStateTerminated s.state.State = ProviderStateTerminated
s.mu.Unlock() s.mu.Unlock()
return nil return nil
+45
View File
@@ -0,0 +1,45 @@
package gateway
import (
"net"
"testing"
"time"
)
func TestGatewaySlowReaderStillCleansUpWithinBound(t *testing.T) {
harness := newGatewayTransportHarness(t)
for sequence := 0; sequence < 10_000; sequence++ {
harness.session.EmitVideo([]byte{byte(sequence)})
}
if err := harness.client.Close(); err != nil {
t.Fatal(err)
}
harness.waitReleased(t)
}
func TestGatewayMalformedUDPDoesNotAmplify(t *testing.T) {
harness := newGatewayTransportHarness(t)
connection, err := net.DialUDP("udp", nil, harness.server.Addr().(*net.UDPAddr))
if err != nil {
t.Fatal(err)
}
defer connection.Close()
request := []byte("invalid")
if _, err := connection.Write(request); err != nil {
t.Fatal(err)
}
if err := connection.SetReadDeadline(time.Now().Add(100 * time.Millisecond)); err != nil {
t.Fatal(err)
}
response := make([]byte, len(request)*3+1)
count, _, err := connection.ReadFromUDP(response)
if err != nil {
if timeout, ok := err.(net.Error); ok && timeout.Timeout() {
return
}
t.Fatal(err)
}
if count > len(request)*3 {
t.Fatalf("malformed UDP amplified %d bytes to %d", len(request), count)
}
}