diff --git a/agent.go b/agent.go index 09124b4..e8c2272 100644 --- a/agent.go +++ b/agent.go @@ -702,13 +702,33 @@ func (a *Agent) GetLocalCandidates() ([]Candidate, error) { } // GetLocalUserCredentials returns the local user credentials -func (a *Agent) GetLocalUserCredentials() (frag string, pwd string) { - return a.localUfrag, a.localPwd +func (a *Agent) GetLocalUserCredentials() (frag string, pwd string, err error) { + valSet := make(chan struct{}) + err = a.run(func(agent *Agent) { + frag = agent.localUfrag + pwd = agent.localPwd + close(valSet) + }, nil) + + if err == nil { + <-valSet + } + return } // GetRemoteUserCredentials returns the remote user credentials -func (a *Agent) GetRemoteUserCredentials() (frag string, pwd string) { - return a.remoteUfrag, a.remotePwd +func (a *Agent) GetRemoteUserCredentials() (frag string, pwd string, err error) { + valSet := make(chan struct{}) + err = a.run(func(agent *Agent) { + frag = agent.remoteUfrag + pwd = agent.remotePwd + close(valSet) + }, nil) + + if err == nil { + <-valSet + } + return } // Close cleans up the Agent diff --git a/agent_test.go b/agent_test.go index 67f687e..2b1602c 100644 --- a/agent_test.go +++ b/agent_test.go @@ -416,8 +416,12 @@ func TestConnectivityOnStartup(t *testing.T) { aConn, bConn := func(aAgent, bAgent *Agent) (*Conn, *Conn) { // Manual signaling - aUfrag, aPwd := aAgent.GetLocalUserCredentials() - bUfrag, bPwd := bAgent.GetLocalUserCredentials() + aUfrag, aPwd, err := aAgent.GetLocalUserCredentials() + assert.NoError(t, err) + + bUfrag, bPwd, err := bAgent.GetLocalUserCredentials() + assert.NoError(t, err) + gatherAndExchangeCandidates(aAgent, bAgent) accepted := make(chan struct{}) @@ -1439,8 +1443,14 @@ func TestAgentRestart(t *testing.T) { assert.NoError(t, connB.agent.Restart("", "")) // Exchange Candidates and Credentials - assert.NoError(t, connA.agent.SetRemoteCredentials(connB.agent.GetLocalUserCredentials())) - assert.NoError(t, connB.agent.SetRemoteCredentials(connA.agent.GetLocalUserCredentials())) + ufrag, pwd, err := connB.agent.GetLocalUserCredentials() + assert.NoError(t, err) + assert.NoError(t, connA.agent.SetRemoteCredentials(ufrag, pwd)) + + ufrag, pwd, err = connA.agent.GetLocalUserCredentials() + assert.NoError(t, err) + assert.NoError(t, connB.agent.SetRemoteCredentials(ufrag, pwd)) + gatherAndExchangeCandidates(connA.agent, connB.agent) // Wait until both have gone back to connected @@ -1457,9 +1467,11 @@ func TestAgentRestart(t *testing.T) { } func TestGetRemoteCredentials(t *testing.T) { - a := Agent{remoteUfrag: "remoteUfrag", remotePwd: "remotePwd"} + a := Agent{remoteUfrag: "remoteUfrag", remotePwd: "remotePwd", muChan: make(chan struct{}, 1)} + + actualUfrag, actualPwd, err := a.GetRemoteUserCredentials() + assert.NoError(t, err) - actualUfrag, actualPwd := a.GetRemoteUserCredentials() assert.Equal(t, actualUfrag, a.remoteUfrag) assert.Equal(t, actualPwd, a.remotePwd) } diff --git a/connectivity_vnet_test.go b/connectivity_vnet_test.go index cc40e18..efbf2e5 100644 --- a/connectivity_vnet_test.go +++ b/connectivity_vnet_test.go @@ -163,8 +163,11 @@ func buildVNet(natType0, natType1 *vnet.NATType) (*virtualNet, error) { func connectWithVNet(aAgent, bAgent *Agent) (*Conn, *Conn) { // Manual signaling - aUfrag, aPwd := aAgent.GetLocalUserCredentials() - bUfrag, bPwd := bAgent.GetLocalUserCredentials() + aUfrag, aPwd, err := aAgent.GetLocalUserCredentials() + check(err) + + bUfrag, bPwd, err := bAgent.GetLocalUserCredentials() + check(err) gatherAndExchangeCandidates(aAgent, bAgent) @@ -585,8 +588,11 @@ func TestWriteUseValidPair(t *testing.T) { gatherAndExchangeCandidates(controllingAgent, controlledAgent) - controllingUfrag, controllingPwd := controllingAgent.GetLocalUserCredentials() - controlledUfrag, controlledPwd := controlledAgent.GetLocalUserCredentials() + controllingUfrag, controllingPwd, err := controllingAgent.GetLocalUserCredentials() + assert.NoError(t, err) + + controlledUfrag, controlledPwd, err := controlledAgent.GetLocalUserCredentials() + assert.NoError(t, err) assert.NoError(t, controllingAgent.startConnectivityChecks(true, controlledUfrag, controlledPwd)) assert.NoError(t, controlledAgent.startConnectivityChecks(false, controllingUfrag, controllingPwd)) diff --git a/transport_test.go b/transport_test.go index 62329eb..f301e5f 100644 --- a/transport_test.go +++ b/transport_test.go @@ -226,13 +226,15 @@ func connect(aAgent, bAgent *Agent) (*Conn, *Conn) { go func() { var acceptErr error - bUfrag, bPwd := bAgent.GetLocalUserCredentials() + bUfrag, bPwd, acceptErr := bAgent.GetLocalUserCredentials() + check(acceptErr) aConn, acceptErr = aAgent.Accept(context.TODO(), bUfrag, bPwd) check(acceptErr) close(accepted) }() - aUfrag, aPwd := aAgent.GetLocalUserCredentials() + aUfrag, aPwd, err := aAgent.GetLocalUserCredentials() + check(err) bConn, err := bAgent.Dial(context.TODO(), aUfrag, aPwd) check(err)