Skip to content

Commit 1e87163

Browse files
committed
Abort mux srflx gather on close
1 parent 9651b02 commit 1e87163

6 files changed

Lines changed: 344 additions & 33 deletions

File tree

agent_test.go

Lines changed: 128 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -376,6 +376,134 @@ func TestAgentCloseAbortsBlockedUDPMuxWrite(t *testing.T) {
376376
require.NoError(t, <-runErr)
377377
}
378378

379+
func TestAgentCloseAbortsBlockedUDPMuxSrflxGatherWrite(t *testing.T) {
380+
defer test.CheckRoutines(t)()
381+
defer test.TimeOut(2 * time.Second).Stop()
382+
383+
udpConn := newDeadlineBlockingPacketConn()
384+
udpMux := NewUniversalUDPMuxDefault(UniversalUDPMuxParams{UDPConn: udpConn})
385+
defer func() {
386+
_ = udpMux.Close()
387+
}()
388+
389+
agent, err := NewAgent(&AgentConfig{
390+
NetworkTypes: []NetworkType{NetworkTypeUDP4},
391+
CandidateTypes: []CandidateType{CandidateTypeServerReflexive},
392+
Urls: []*stun.URI{{
393+
Scheme: stun.SchemeTypeSTUN,
394+
Host: "192.0.2.2",
395+
Port: 3478,
396+
}},
397+
UDPMuxSrflx: udpMux,
398+
})
399+
require.NoError(t, err)
400+
401+
require.NoError(t, agent.OnCandidate(func(Candidate) {}))
402+
require.NoError(t, agent.GatherCandidates())
403+
404+
select {
405+
case <-udpConn.writeStarted:
406+
case <-time.After(time.Second):
407+
require.FailNow(t, "timed out waiting for UDP mux srflx gather write to block")
408+
}
409+
410+
closeErr := make(chan error, 1)
411+
go func() {
412+
closeErr <- agent.Close()
413+
}()
414+
415+
select {
416+
case err := <-closeErr:
417+
require.NoError(t, err)
418+
case <-time.After(time.Second):
419+
require.FailNow(t, "agent close did not abort blocked UDP mux srflx gather write")
420+
}
421+
}
422+
423+
func TestAgentCloseDoesNotAbortOtherAgentUDPMuxSrflxGatherWrite(t *testing.T) { //nolint:cyclop
424+
defer test.CheckRoutines(t)()
425+
defer test.TimeOut(2 * time.Second).Stop()
426+
427+
udpConn := newDeadlineBlockingPacketConn()
428+
udpMux := NewUniversalUDPMuxDefault(UniversalUDPMuxParams{UDPConn: udpConn})
429+
defer func() {
430+
_ = udpMux.Close()
431+
}()
432+
433+
newSrflxAgent := func(t *testing.T) *Agent {
434+
t.Helper()
435+
436+
agent, err := NewAgent(&AgentConfig{
437+
NetworkTypes: []NetworkType{NetworkTypeUDP4},
438+
CandidateTypes: []CandidateType{CandidateTypeServerReflexive},
439+
Urls: []*stun.URI{{
440+
Scheme: stun.SchemeTypeSTUN,
441+
Host: "192.0.2.2",
442+
Port: 3478,
443+
}},
444+
UDPMuxSrflx: udpMux,
445+
})
446+
require.NoError(t, err)
447+
448+
return agent
449+
}
450+
451+
agent1 := newSrflxAgent(t)
452+
defer func() {
453+
_ = agent1.Close()
454+
}()
455+
456+
agent2 := newSrflxAgent(t)
457+
defer func() {
458+
_ = agent2.Close()
459+
}()
460+
461+
require.NoError(t, agent2.OnCandidate(func(Candidate) {}))
462+
require.NoError(t, agent2.GatherCandidates())
463+
464+
select {
465+
case <-udpConn.writeStarted:
466+
case <-time.After(time.Second):
467+
require.FailNow(t, "timed out waiting for second agent UDP mux srflx gather write to block")
468+
}
469+
470+
closeErr := make(chan error, 1)
471+
go func() {
472+
closeErr <- agent1.Close()
473+
}()
474+
475+
select {
476+
case err := <-closeErr:
477+
require.NoError(t, err)
478+
case <-time.After(time.Second):
479+
require.FailNow(t, "first agent close blocked")
480+
}
481+
482+
select {
483+
case <-udpConn.writeDeadlineSet:
484+
require.FailNow(t, "first agent close aborted another agent's UDP mux srflx gather write")
485+
case <-time.After(200 * time.Millisecond):
486+
}
487+
488+
secondCloseErr := make(chan error, 1)
489+
go func() {
490+
secondCloseErr <- agent2.Close()
491+
}()
492+
493+
select {
494+
case <-udpConn.writeDeadlineSet:
495+
case <-time.After(time.Second):
496+
require.FailNow(t, "second agent close did not abort its own UDP mux srflx gather write")
497+
}
498+
499+
select {
500+
case err := <-secondCloseErr:
501+
require.NoError(t, err)
502+
case <-time.After(time.Second):
503+
require.FailNow(t, "second agent close blocked")
504+
}
505+
}
506+
379507
func TestAgentCloseClearsSharedUDPMuxAbortDeadlineForOtherAgent(t *testing.T) { //nolint:cyclop
380508
defer test.CheckRoutines(t)()
381509
defer test.TimeOut(2 * time.Second).Stop()

gather.go

Lines changed: 22 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -760,7 +760,7 @@ func (a *Agent) gatherCandidatesSrflxUDPMux(ctx context.Context, urls []*stun.UR
760760
return
761761
}
762762

763-
xorAddr, err := a.udpMuxSrflx.GetXORMappedAddr(serverAddr, a.stunGatherTimeout)
763+
xorAddr, err := getXORMappedAddr(ctx, a.udpMuxSrflx, serverAddr, a.stunGatherTimeout)
764764
if err != nil {
765765
a.log.Warnf("Failed get server reflexive address %s %s: %v", network, url, err)
766766

@@ -809,6 +809,27 @@ func (a *Agent) gatherCandidatesSrflxUDPMux(ctx context.Context, urls []*stun.UR
809809
}
810810
}
811811

812+
type contextXORMappedAddrGetter interface {
813+
GetXORMappedAddrContext(context.Context, net.Addr, time.Duration) (*stun.XORMappedAddress, error)
814+
}
815+
816+
func getXORMappedAddr(
817+
ctx context.Context,
818+
mux UniversalUDPMux,
819+
serverAddr net.Addr,
820+
deadline time.Duration,
821+
) (*stun.XORMappedAddress, error) {
822+
if muxWithContext, ok := mux.(contextXORMappedAddrGetter); ok {
823+
return muxWithContext.GetXORMappedAddrContext(ctx, serverAddr, deadline)
824+
}
825+
826+
if err := ctx.Err(); err != nil {
827+
return nil, err
828+
}
829+
830+
return mux.GetXORMappedAddr(serverAddr, deadline)
831+
}
832+
812833
//nolint:cyclop,gocognit
813834
func (a *Agent) gatherCandidatesSrflx(ctx context.Context, urls []*stun.URI, networkTypes []NetworkType) {
814835
var wg sync.WaitGroup

gather_test.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1522,7 +1522,7 @@ func TestGatherCandidatesRelayRespectsInterfaceFilter(t *testing.T) {
15221522
},
15231523
}),
15241524
WithInterfaceFilter(func(iface string) bool {
1525-
return iface == "eth0"
1525+
return iface == "eth0" //nolint:goconst
15261526
}),
15271527
WithIncludeLoopback(),
15281528
)

udp_mux.go

Lines changed: 58 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
package ice
55

66
import (
7+
"context"
78
"errors"
89
"io"
910
"net"
@@ -50,7 +51,11 @@ type UDPMuxDefault struct {
5051
// whether the UDP connection listens on an unspecified address
5152
isUnspecified bool
5253

53-
// writeState packs write in-flight count.
54+
// writeState coordinates context cancellation for WriteTo calls.
55+
// Low bits count writes currently inside UDPConn.WriteTo. blocked means an
56+
// abort arming the shared write deadline, so new writes wait.
57+
// deadline means SetWriteDeadline(time.Now()) succeeded and the last
58+
// in-flight writer must clear it before new writes can enter.
5459
writeState atomic.Uint64
5560
}
5661

@@ -269,13 +274,53 @@ func (m *UDPMuxDefault) Close() error {
269274
}
270275

271276
func (m *UDPMuxDefault) writeTo(buf []byte, rAddr net.Addr) (n int, err error) {
272-
m.startWrite()
277+
return m.writeToContext(context.Background(), buf, rAddr)
278+
}
279+
280+
func (m *UDPMuxDefault) writeToContext(ctx context.Context, buf []byte, rAddr net.Addr) (n int, err error) {
281+
if err = m.startWriteContext(ctx); err != nil {
282+
return 0, err
283+
}
273284

274285
defer func() {
275286
err = m.finishWrite(err)
276287
}()
277288

278-
return m.params.UDPConn.WriteTo(buf, rAddr)
289+
if err = ctx.Err(); err != nil {
290+
return 0, err
291+
}
292+
293+
if done := ctx.Done(); done != nil {
294+
// net.PacketConn writes cannot be canceled directly. If ctx is
295+
// canceled while WriteTo is blocked, abortWrite interrupts it by
296+
// temporarily setting the shared socket write deadline to now.
297+
stopAbort := make(chan struct{})
298+
var stopped atomic.Bool
299+
defer func() {
300+
stopped.Store(true)
301+
close(stopAbort)
302+
}()
303+
go func() {
304+
select {
305+
case <-done:
306+
if !stopped.Load() {
307+
if abortErr := m.abortWrite(); abortErr != nil {
308+
m.params.Logger.Warnf("Failed to abort UDP write: %v", abortErr)
309+
}
310+
}
311+
case <-stopAbort:
312+
}
313+
}()
314+
}
315+
316+
n, err = m.params.UDPConn.WriteTo(buf, rAddr)
317+
if err != nil {
318+
if ctxErr := ctx.Err(); ctxErr != nil {
319+
return n, ctxErr
320+
}
321+
}
322+
323+
return n, err
279324
}
280325

281326
func (m *UDPMuxDefault) abortWrite() error {
@@ -289,6 +334,8 @@ func (m *UDPMuxDefault) abortWrite() error {
289334
continue
290335
}
291336

337+
// The deadline applies to the shared UDPConn, so blocked stays set
338+
// until the final in-flight writer clears the deadline in finishWrite.
292339
if err := m.params.UDPConn.SetWriteDeadline(time.Now()); err != nil {
293340
m.clearWriteAbortState()
294341

@@ -301,8 +348,12 @@ func (m *UDPMuxDefault) abortWrite() error {
301348
}
302349
}
303350

304-
func (m *UDPMuxDefault) startWrite() {
351+
func (m *UDPMuxDefault) startWriteContext(ctx context.Context) error {
305352
for {
353+
if err := ctx.Err(); err != nil {
354+
return err
355+
}
356+
306357
state := m.writeState.Load()
307358
if state&udpMuxWriteBlockedBit != 0 {
308359
runtime.Gosched()
@@ -311,7 +362,7 @@ func (m *UDPMuxDefault) startWrite() {
311362
}
312363

313364
if m.writeState.CompareAndSwap(state, state+1) {
314-
return
365+
return nil
315366
}
316367
}
317368
}
@@ -357,6 +408,8 @@ func (m *UDPMuxDefault) clearWriteDeadlineAfterAbort(writeErr error) error {
357408
return writeErr
358409
}
359410
if state&udpMuxWriteDeadlineBit == 0 {
411+
// The last writer can race with abortWrite after blocked is set but
412+
// before SetWriteDeadline returns.
360413
runtime.Gosched()
361414

362415
continue

0 commit comments

Comments
 (0)