From 8826f5c80442211f7403b5308283a61bbe3c65e0 Mon Sep 17 00:00:00 2001 From: sechmachine <97589681+sechmachine727@users.noreply.github.com> Date: Wed, 29 Jul 2026 22:13:32 +0700 Subject: [PATCH] fix(gateway): bound fake media shutdown --- gateway/gateway_test.go | 3 ++- gateway/provider.go | 32 ++++++++++++++++++++++------ gateway/resource_test.go | 45 ++++++++++++++++++++++++++++++++++++++++ 3 files changed, 73 insertions(+), 7 deletions(-) create mode 100644 gateway/resource_test.go diff --git a/gateway/gateway_test.go b/gateway/gateway_test.go index f222f52..82e428f 100644 --- a/gateway/gateway_test.go +++ b/gateway/gateway_test.go @@ -465,6 +465,7 @@ type gatewayTransportHarness struct { session *fakeSession admission *oneTimeAdmission reporter *recordingProviderStateReporter + server *Server } func newGatewayTransportHarness(t *testing.T) gatewayTransportHarness { @@ -500,7 +501,7 @@ func newGatewayTransportHarness(t *testing.T) gatewayTransportHarness { 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) { diff --git a/gateway/provider.go b/gateway/provider.go index ea2b4e2..cd2f3dd 100644 --- a/gateway/provider.go +++ b/gateway/provider.go @@ -443,20 +443,42 @@ func (s *fakeSession) EmitEvent(event ProviderEvent) { } 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 { case s.video <- append([]byte(nil), payload...): default: - <-s.video - s.video <- append([]byte(nil), payload...) + select { + case <-s.video: + default: + } + select { + case s.video <- append([]byte(nil), payload...): + default: + } } } 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 { case s.audio <- append([]byte(nil), payload...): default: - <-s.audio - s.audio <- append([]byte(nil), payload...) + select { + 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 } s.state.State = ProviderStateTerminating - s.mu.Unlock() s.closeOnce.Do(func() { close(s.video) close(s.audio) }) - s.mu.Lock() s.state.State = ProviderStateTerminated s.mu.Unlock() return nil diff --git a/gateway/resource_test.go b/gateway/resource_test.go new file mode 100644 index 0000000..cc6310b --- /dev/null +++ b/gateway/resource_test.go @@ -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) + } +}