diff --git a/agent.go b/agent.go index 36d85c82..7e9aa5c9 100644 --- a/agent.go +++ b/agent.go @@ -1014,6 +1014,19 @@ func (a *Agent) AddRemoteCandidate(cand Candidate) error { return nil } +// AddVirtualCandidate registers a virtual candidate as a local ICE candidate, +// using the supplied packet connection to send and receive packets. +func (a *Agent) AddVirtualCandidate(cand Candidate, candidateConn net.PacketConn) error { + if cand == nil { + return nil + } + if candidateConn == nil { + return errCandidatePacketConnNil + } + + return a.addCandidate(a.loop, cand, candidateConn, true) +} + func isMulticastDNSCandidate(cand Candidate) bool { return cand.Type() == CandidateTypeHost && strings.HasSuffix(cand.Address(), ".local") } @@ -1397,15 +1410,27 @@ func (a *Agent) shouldAcceptRemoteCandidate(cand Candidate) bool { return true } -func (a *Agent) addCandidate(ctx context.Context, cand Candidate, candidateConn net.PacketConn) error { +func (a *Agent) addCandidate( + ctx context.Context, + cand Candidate, + candidateConn net.PacketConn, + errorOnDuplicate bool, +) error { if err := ctx.Err(); err != nil { return err } - return a.loop.Run(ctx, func(context.Context) { + var addErr error + err := a.loop.Run(ctx, func(context.Context) { set := a.localCandidates[cand.NetworkType()] for _, candidate := range set { if candidate.Equal(cand) { + if errorOnDuplicate { + addErr = errDuplicateCandidate + + return + } + a.log.Debugf("Ignore duplicate candidate: %s", cand) if err := cand.close(); err != nil { a.log.Warnf("Failed to close duplicate candidate: %v", err) @@ -1425,10 +1450,8 @@ func (a *Agent) addCandidate(ctx context.Context, cand Candidate, candidateConn set = append(set, cand) a.localCandidates[cand.NetworkType()] = set - if remoteCandidates, ok := a.remoteCandidates[cand.NetworkType()]; ok { - for _, remoteCandidate := range remoteCandidates { - a.addPair(cand, remoteCandidate) - } + for _, remoteCandidate := range a.remoteCandidates[cand.NetworkType()] { + a.addPair(cand, remoteCandidate) } a.requestConnectivityCheck() @@ -1437,6 +1460,11 @@ func (a *Agent) addCandidate(ctx context.Context, cand Candidate, candidateConn a.candidateNotifier.EnqueueCandidate(cand) } }) + if err != nil { + return err + } + + return addErr } func (a *Agent) setCandidateExtensions(cand Candidate) { diff --git a/agent_test.go b/agent_test.go index 086bfbbc..ebb2b273 100644 --- a/agent_test.go +++ b/agent_test.go @@ -51,6 +51,30 @@ type blockingWritePacketConn struct { closeOnce sync.Once } +type localCandidatePacketConn struct { + addr net.Addr + closeCount atomic.Int32 +} + +func (c *localCandidatePacketConn) ReadFrom([]byte) (int, net.Addr, error) { + return 0, c.addr, io.EOF +} + +func (c *localCandidatePacketConn) WriteTo(payload []byte, _ net.Addr) (int, error) { + return len(payload), nil +} + +func (c *localCandidatePacketConn) Close() error { + c.closeCount.Add(1) + + return nil +} + +func (c *localCandidatePacketConn) LocalAddr() net.Addr { return c.addr } +func (c *localCandidatePacketConn) SetDeadline(time.Time) error { return nil } +func (c *localCandidatePacketConn) SetReadDeadline(time.Time) error { return nil } +func (c *localCandidatePacketConn) SetWriteDeadline(time.Time) error { return nil } + func newBlockingWritePacketConn() *blockingWritePacketConn { return &blockingWritePacketConn{ writeStarted: make(chan struct{}), @@ -2782,6 +2806,143 @@ func TestAddRemoteCandidateHonorsRemoteIPFilter(t *testing.T) { }, time.Second, 10*time.Millisecond) } +func TestAddVirtualCandidateRegistersExternalRelay(t *testing.T) { + agent, err := NewAgentWithOptions( + WithCandidateTypes([]CandidateType{}), + WithMulticastDNSMode(MulticastDNSModeDisabled), + ) + require.NoError(t, err) + defer func() { require.NoError(t, agent.Close()) }() + + gatheringComplete := make(chan struct{}) + candidates := make(chan Candidate, 1) + require.NoError(t, agent.OnCandidate(func(candidate Candidate) { + if candidate == nil { + close(gatheringComplete) + + return + } + candidates <- candidate + })) + + candidate, err := NewCandidateRelay(&CandidateRelayConfig{ + Network: NetworkTypeUDP4.String(), + Address: "192.0.2.10", + Port: 5000, + Component: ComponentRTP, + RelayProtocol: "custom", + }) + require.NoError(t, err) + packetConn := &localCandidatePacketConn{ + addr: &net.UDPAddr{IP: net.IPv4(192, 0, 2, 10), Port: 5000}, + } + + require.NoError(t, agent.AddVirtualCandidate(candidate, packetConn)) + localCandidates, err := agent.GetLocalCandidates() + require.NoError(t, err) + require.Len(t, localCandidates, 1) + relayCandidate, ok := localCandidates[0].(*CandidateRelay) + require.True(t, ok) + require.Equal(t, "custom", relayCandidate.RelayProtocol()) + + select { + case got := <-candidates: + require.Equal(t, candidate, got) + case <-time.After(time.Second): + require.FailNow(t, "timed out waiting for local candidate callback") + } + + require.NoError(t, agent.GatherCandidates()) + select { + case <-gatheringComplete: + case <-time.After(time.Second): + require.FailNow(t, "timed out waiting for gathering completion") + } +} + +func TestAddVirtualCandidateRejectsNilPacketConn(t *testing.T) { + agent, err := NewAgentWithOptions(WithMulticastDNSMode(MulticastDNSModeDisabled)) + require.NoError(t, err) + defer func() { require.NoError(t, agent.Close()) }() + + candidate, err := NewCandidateRelay(&CandidateRelayConfig{ + Network: NetworkTypeUDP4.String(), + Address: "192.0.2.11", + Port: 5001, + Component: ComponentRTP, + }) + require.NoError(t, err) + require.ErrorIs(t, agent.AddVirtualCandidate(candidate, nil), errCandidatePacketConnNil) +} + +func TestAddVirtualCandidateRejectsDuplicateWithoutClosing(t *testing.T) { + agent, err := NewAgentWithOptions(WithMulticastDNSMode(MulticastDNSModeDisabled)) + require.NoError(t, err) + defer func() { require.NoError(t, agent.Close()) }() + + require.NoError(t, agent.OnCandidate(func(Candidate) {})) + var candidateCloseCount atomic.Int32 + candidate, err := NewCandidateRelay(&CandidateRelayConfig{ + Network: NetworkTypeUDP4.String(), + Address: "192.0.2.12", + Port: 5002, + Component: ComponentRTP, + OnClose: func() error { + candidateCloseCount.Add(1) + + return nil + }, + }) + require.NoError(t, err) + packetConn := &localCandidatePacketConn{ + addr: &net.UDPAddr{IP: net.IPv4(192, 0, 2, 12), Port: 5002}, + } + + require.NoError(t, agent.AddVirtualCandidate(candidate, packetConn)) + require.ErrorIs(t, agent.AddVirtualCandidate(candidate, packetConn), errDuplicateCandidate) + require.Zero(t, candidateCloseCount.Load()) + require.Zero(t, packetConn.closeCount.Load()) + + localCandidates, err := agent.GetLocalCandidates() + require.NoError(t, err) + require.Equal(t, []Candidate{candidate}, localCandidates) +} + +func TestAddCandidateClosesDuplicate(t *testing.T) { + agent, err := NewAgentWithOptions(WithMulticastDNSMode(MulticastDNSModeDisabled)) + require.NoError(t, err) + defer func() { require.NoError(t, agent.Close()) }() + + require.NoError(t, agent.OnCandidate(func(Candidate) {})) + config := CandidateRelayConfig{ + Network: NetworkTypeUDP4.String(), + Address: "192.0.2.13", + Port: 5003, + Component: ComponentRTP, + } + first, err := NewCandidateRelay(&config) + require.NoError(t, err) + var duplicateCloseCount atomic.Int32 + config.OnClose = func() error { + duplicateCloseCount.Add(1) + + return nil + } + duplicate, err := NewCandidateRelay(&config) + require.NoError(t, err) + firstConn := &localCandidatePacketConn{ + addr: &net.UDPAddr{IP: net.IPv4(192, 0, 2, 13), Port: 5003}, + } + duplicateConn := &localCandidatePacketConn{ + addr: &net.UDPAddr{IP: net.IPv4(192, 0, 2, 13), Port: 5003}, + } + + require.NoError(t, agent.addCandidate(context.Background(), first, firstConn, false)) + require.NoError(t, agent.addCandidate(context.Background(), duplicate, duplicateConn, false)) + require.Equal(t, int32(1), duplicateCloseCount.Load()) + require.Equal(t, int32(1), duplicateConn.closeCount.Load()) +} + func TestGetLocalCandidates(t *testing.T) { var config AgentConfig @@ -2807,7 +2968,7 @@ func TestGetLocalCandidates(t *testing.T) { expectedCandidates = append(expectedCandidates, cand) - err = agent.addCandidate(context.Background(), cand, dummyConn) + err = agent.addCandidate(context.Background(), cand, dummyConn, false) require.NoError(t, err) } @@ -3480,7 +3641,7 @@ func TestSetCandidatesUfrag(t *testing.T) { cand, errCand := NewCandidateHost(&cfg) require.NoError(t, errCand) - err = agent.addCandidate(context.Background(), cand, dummyConn) + err = agent.addCandidate(context.Background(), cand, dummyConn, false) require.NoError(t, err) } diff --git a/errors.go b/errors.go index 99578525..aaf76be7 100644 --- a/errors.go +++ b/errors.go @@ -177,6 +177,8 @@ var ( ErrAgentOptionNotUpdatable = errors.New("option can only be set during agent construction") errAttributeTooShortICECandidate = errors.New("attribute not long enough to be ICE candidate") + errCandidatePacketConnNil = errors.New("candidate packet connection is nil") + errDuplicateCandidate = errors.New("candidate already added") errClosingConnection = errors.New("failed to close connection") errConnectionAddrAlreadyExist = errors.New("connection with same remote address already exists") errInvalidAddress = errors.New("invalid address") diff --git a/gather.go b/gather.go index 42390b49..1b6eb29d 100644 --- a/gather.go +++ b/gather.go @@ -485,7 +485,7 @@ func (a *Agent) gatherCandidatesLocal(ctx context.Context, networkTypes []Networ continue } - if err := a.addCandidate(ctx, candidateHost, connAndPort.conn); err != nil { + if err := a.addCandidate(ctx, candidateHost, connAndPort.conn, false); err != nil { if closeErr := candidateHost.close(); closeErr != nil { a.log.Warnf("Failed to close candidate: %v", closeErr) } @@ -594,7 +594,7 @@ func (a *Agent) gatherCandidatesLocalUDPMux(ctx context.Context) error { //nolin continue } - if err := a.addCandidate(ctx, c, conn); err != nil { + if err := a.addCandidate(ctx, c, conn, false); err != nil { if closeErr := c.close(); closeErr != nil { a.log.Warnf("Failed to close candidate: %v", closeErr) } @@ -709,7 +709,7 @@ func (a *Agent) gatherCandidatesSrflxMapped(ctx context.Context, networkTypes [] continue } - if err := a.addCandidate(ctx, c, currentConn); err != nil { + if err := a.addCandidate(ctx, c, currentConn, false); err != nil { if closeErr := c.close(); closeErr != nil { a.log.Warnf("Failed to close candidate: %v", closeErr) } @@ -798,7 +798,7 @@ func (a *Agent) gatherCandidatesSrflxUDPMux(ctx context.Context, urls []*stun.UR return } - if err := a.addCandidate(ctx, c, conn); err != nil { + if err := a.addCandidate(ctx, c, conn, false); err != nil { if closeErr := c.close(); closeErr != nil { a.log.Warnf("Failed to close candidate: %v", closeErr) } @@ -929,7 +929,7 @@ func (a *Agent) gatherCandidatesSrflx(ctx context.Context, urls []*stun.URI, net return } - if err := a.addCandidate(ctx, c, conn); err != nil { + if err := a.addCandidate(ctx, c, conn, false); err != nil { if closeErr := c.close(); closeErr != nil { a.log.Warnf("Failed to close candidate: %v", closeErr) } @@ -1379,7 +1379,7 @@ func (a *Agent) createRelayCandidate(ctx context.Context, ep relayEndpoint, ip n return err } - if err := a.addCandidate(ctx, candidate, ep.conn); err != nil { + if err := a.addCandidate(ctx, candidate, ep.conn, false); err != nil { if closeErr := candidate.close(); closeErr != nil { a.log.Warnf("Failed to close candidate: %v", closeErr) }