Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
126 changes: 97 additions & 29 deletions internal/usecase/devices/wsman/message.go
Original file line number Diff line number Diff line change
Expand Up @@ -141,8 +141,16 @@ func (g GoWSMANMessages) Worker() {
}
}

// setupResult carries the connection entry off the Worker goroutine. awaitAuth
// is non-nil when the entry still has to finish authenticating; the caller runs
// it on its own goroutine so the Worker is free for the next device.
type setupResult struct {
entry *ConnectionEntry
awaitAuth func(context.Context) (*ConnectionEntry, error)
}

func (g GoWSMANMessages) SetupWsmanClient(ctx context.Context, device entity.Device, isRedirection, logAMTMessages bool) (Management, error) {
resultChan := make(chan *ConnectionEntry, 1)
resultChan := make(chan setupResult, 1)
errChan := make(chan error, 1)
// Queue the request
requestQueue <- func() {
Expand Down Expand Up @@ -182,23 +190,35 @@ func (g GoWSMANMessages) SetupWsmanClient(ctx context.Context, device entity.Dev
}

connection.WsmanMessages = wsman.NewMessages(cp)
resultChan <- connection
} else {
resultChan <- g.setupWsmanClientInternal(device, isRedirection, logAMTMessages)
resultChan <- setupResult{entry: connection}

return
}

entry, awaitAuth := g.setupWsmanClientInternal(device, isRedirection, logAMTMessages)
resultChan <- setupResult{entry: entry, awaitAuth: awaitAuth}
}

select {
case err := <-errChan:
return nil, err
case result := <-resultChan:
return result, nil
if result.awaitAuth != nil {
entry, err := result.awaitAuth(ctx)
if err != nil {
return nil, err
}

return entry, nil
}

return result.entry, nil
case <-ctx.Done():
return nil, ErrCancelled.Wrap("SetupWsmanClient", "ctx.Done", ctx.Err())
}
}

func (g GoWSMANMessages) setupWsmanClientInternal(device entity.Device, isRedirection, logAMTMessages bool) *ConnectionEntry {
func (g GoWSMANMessages) setupWsmanClientInternal(device entity.Device, isRedirection, logAMTMessages bool) (conn *ConnectionEntry, awaitAuth func(context.Context) (*ConnectionEntry, error)) {
clientParams := client.Parameters{
Target: device.Hostname,
Username: device.Username,
Expand Down Expand Up @@ -226,34 +246,20 @@ func (g GoWSMANMessages) setupWsmanClientInternal(device entity.Device, isRedire
RemoveConnection(device.GUID)
})

return entry
return entry, nil
} else if entry.IsCIRA {
entry.WsmanMessages = wsman.NewMessages(clientParams)

return entry
return entry, nil
}

ticker := time.NewTicker(waitForAuthTickTime)

defer ticker.Stop()
// Hand the wait back to the caller instead of running it here. Capture
// the client now, on the Worker goroutine: entry.WsmanMessages is
// assigned without a lock elsewhere, so the caller must not re-read it.
pending := entry.WsmanMessages.Client

timeout := time.After(waitForAuth)

for {
select {
case <-ticker.C:
if entry.WsmanMessages.Client.IsAuthenticated() {
return entry
}
case <-timeout:
newEntry := &ConnectionEntry{
WsmanMessages: wsman.NewMessages(clientParams),
Timer: timer,
}
SetConnectionEntry(device.GUID, newEntry)

return newEntry
}
return entry, func(ctx context.Context) (*ConnectionEntry, error) {
return waitForAuthentication(ctx, device.GUID, entry, pending, clientParams, timer)
}
}

Expand All @@ -266,7 +272,51 @@ func (g GoWSMANMessages) setupWsmanClientInternal(device entity.Device, isRedire
newEntry.WsmanMessages.Client.IsAuthenticated()
SetConnectionEntry(device.GUID, newEntry)

return newEntry
return newEntry, nil
}

// waitForAuthentication polls until an in-flight request finishes the device
// handshake, falling back to a fresh client if it never does. It runs on the
// caller's goroutine rather than inside the queued closure, so a device that is
// slow to authenticate no longer delays every other device's setup (#1247).
func waitForAuthentication(ctx context.Context, guid string, entry *ConnectionEntry, pending client.WSMan, clientParams client.Parameters, timer *time.Timer) (*ConnectionEntry, error) {
ticker := time.NewTicker(waitForAuthTickTime)

defer ticker.Stop()

timeout := time.After(waitForAuth)

for {
select {
case <-ticker.C:
if pending.IsAuthenticated() {
// The cached entry keeps its own expiry timer; ours is surplus.
timer.Stop()

return entry, nil
}
case <-timeout:
newEntry := &ConnectionEntry{
WsmanMessages: wsman.NewMessages(clientParams),
Timer: timer,
}

cached := replaceConnectionEntry(guid, entry, newEntry)
if cached != newEntry {
// Another waiter replaced the entry first; drop ours so its
// expiry timer cannot evict the entry that is now cached.
timer.Stop()
}

return cached, nil
case <-ctx.Done():
Comment thread
amarnath-ac marked this conversation as resolved.
// Nothing takes ownership of the timer on this path; left armed it
// would evict whatever is cached for the device 90s from now.
timer.Stop()

return nil, ErrCancelled.Wrap("SetupWsmanClient", "ctx.Done", ctx.Err())
}
}
}

// RemoveConnection safely deletes a connection entry from the global map.
Expand All @@ -293,6 +343,24 @@ func SetConnectionEntry(guid string, entry *ConnectionEntry) {
connections[guid] = entry
}

// replaceConnectionEntry swaps stale for replacement under the map lock and
// returns the entry that is actually cached. Waiters on the same unauthenticated
// device all time out together; the first one installs its entry and the rest
// get it back, so a burst of requests shares one client instead of each caching
// its own.
func replaceConnectionEntry(guid string, stale, replacement *ConnectionEntry) *ConnectionEntry {
connectionsMu.Lock()
defer connectionsMu.Unlock()

if current, ok := connections[guid]; ok && current != stale {
return current
}

connections[guid] = replacement

return replacement
}

// HasConnections safely checks whether any connections exist.
func HasConnections() bool {
connectionsMu.RLock()
Expand Down
104 changes: 104 additions & 0 deletions internal/usecase/devices/wsman/message_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,14 +2,18 @@ package wsman

import (
"context"
"sync"
"testing"
"time"

"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"

gwmconfig "github.com/device-management-toolkit/go-wsman-messages/v2/pkg/config"
"github.com/device-management-toolkit/go-wsman-messages/v2/pkg/wsman"
"github.com/device-management-toolkit/go-wsman-messages/v2/pkg/wsman/client"

"github.com/device-management-toolkit/console/config"
"github.com/device-management-toolkit/console/internal/entity"
dto "github.com/device-management-toolkit/console/internal/entity/dto/v1"
"github.com/device-management-toolkit/console/pkg/logger"
Expand Down Expand Up @@ -173,3 +177,103 @@ func TestDestroyWsmanClient_MissingEntryIsNoop(t *testing.T) {
// Should not panic when the entry is absent.
g.DestroyWsmanClient(dto.Device{GUID: "destroy-missing-entry"})
}

// TestSetupWsmanClientSlowDeviceDoesNotDelayOthers is a regression test for
// #1247: a device that never authenticates made every other device's setup
// time out, because the wait for authentication ran inside the closure drained
// by the single Worker goroutine. The wait itself is unchanged; it now runs on
// the caller's goroutine, so the Worker stays free.
func TestSetupWsmanClientSlowDeviceDoesNotDelayOthers(t *testing.T) { //nolint:paralleltest // mutates package-level state (requestQueue, waitForAuth, connections)
const (
slowGUID = "issue-1247-slow-device"
healthyGUID = "issue-1247-healthy-device"
budget = 3 * time.Second
)

if config.ConsoleConfig == nil {
config.ConsoleConfig = &config.Config{}

t.Cleanup(func() { config.ConsoleConfig = nil })
}

// Keep the slow device waiting well past the budget without making the
// test sit through the production 30s.
origWait := waitForAuth
waitForAuth = time.Minute

t.Cleanup(func() { waitForAuth = origWait })
t.Cleanup(func() { RemoveConnection(slowGUID); RemoveConnection(healthyGUID) })

// A cached, unauthenticated entry pointing at a black-hole address: the
// authentication wait will run to its full timeout.
SetConnectionEntry(slowGUID, &ConnectionEntry{
WsmanMessages: wsman.NewMessages(client.Parameters{
Target: "192.0.2.1", Username: "u", Password: "p", UseDigest: true,
}),
Timer: time.AfterFunc(time.Hour, func() {}),
})
require.False(t, GetConnectionEntry(slowGUID).WsmanMessages.Client.IsAuthenticated(),
"the slow device must start out unauthenticated for this test to exercise the wait")

stopWorker := make(chan struct{})
workerDone := make(chan struct{})

go func() {
defer close(workerDone)

for {
select {
case <-stopWorker:
return
case req := <-requestQueue:
req()
// Stands in for the Worker's queueTickTime throttle without
// mutating that package-level var out from under it.
time.Sleep(time.Millisecond)
}
}
}()

t.Cleanup(func() {
close(stopWorker)

select {
case <-workerDone:
case <-time.After(time.Second):
}
})

g := NewGoWSMANMessages(logger.New("error"), passthroughCryptor{})

ctx, cancel := context.WithCancel(context.Background())
defer cancel()

var wg sync.WaitGroup

wg.Add(1)

go func() {
defer wg.Done()

_, _ = g.SetupWsmanClient(ctx, entity.Device{GUID: slowGUID, Hostname: "192.0.2.1", Username: "u", Password: "p"}, false, false)
}()

t.Cleanup(wg.Wait)

// Give the slow request time to reach the Worker and start waiting.
time.Sleep(50 * time.Millisecond)

done := make(chan error, 1)

go func() {
_, err := g.SetupWsmanClient(ctx, entity.Device{GUID: healthyGUID, Hostname: "192.0.2.2", Username: "u", Password: "p"}, false, false)
done <- err
}()

select {
case err := <-done:
require.NoError(t, err, "the healthy device must set up successfully")
case <-time.After(budget):
t.Fatal("a slow device is still delaying every other device's setup (#1247)")
}
}
Loading