Skip to content
Draft
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
24 changes: 16 additions & 8 deletions pkg/capabilities/registry/atomic.go
Original file line number Diff line number Diff line change
Expand Up @@ -100,17 +100,21 @@ func (a *atomicTriggerCapability) Load() *capabilities.TriggerCapability {
}

func (a *atomicTriggerCapability) RegisterTrigger(ctx context.Context, request capabilities.TriggerRegistrationRequest) (<-chan capabilities.TriggerResponse, error) {
a.mu.Lock()
defer a.mu.Unlock()
// A read lock suffices because registrations is thread-safe; it allows multiple triggers to
// be registered concurrently, while Update rebinds registrations under the exclusive write lock.
a.mu.RLock()
defer a.mu.RUnlock()
if a.cap == nil {
return nil, errors.New("capability unavailable")
}
return a.registrations.register(ctx, a.cap, request)
}

func (a *atomicTriggerCapability) UnregisterTrigger(ctx context.Context, request capabilities.TriggerRegistrationRequest) error {
a.mu.Lock()
defer a.mu.Unlock()
// A read lock suffices because registrations is thread-safe; it allows triggers to be
// unregistered concurrently, while Update rebinds registrations under the exclusive write lock.
a.mu.RLock()
defer a.mu.RUnlock()
if a.cap == nil {
return errors.New("capability unavailable")
}
Expand Down Expand Up @@ -277,17 +281,21 @@ func (a *atomicExecuteAndTriggerCapability) Load() *capabilities.ExecutableAndTr
}

func (a *atomicExecuteAndTriggerCapability) RegisterTrigger(ctx context.Context, request capabilities.TriggerRegistrationRequest) (<-chan capabilities.TriggerResponse, error) {
a.mu.Lock()
defer a.mu.Unlock()
// A read lock suffices because registrations is thread-safe; it allows multiple triggers to
// be registered concurrently, while Update rebinds registrations under the exclusive write lock.
a.mu.RLock()
defer a.mu.RUnlock()
if a.cap == nil {
return nil, errors.New("capability unavailable")
}
return a.registrations.register(ctx, a.cap, request)
}

func (a *atomicExecuteAndTriggerCapability) UnregisterTrigger(ctx context.Context, request capabilities.TriggerRegistrationRequest) error {
a.mu.Lock()
defer a.mu.Unlock()
// A read lock suffices because registrations is thread-safe; it allows triggers to be
// unregistered concurrently, while Update rebinds registrations under the exclusive write lock.
a.mu.RLock()
defer a.mu.RUnlock()
if a.cap == nil {
return errors.New("capability unavailable")
}
Expand Down
223 changes: 222 additions & 1 deletion pkg/capabilities/registry/base_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import (
"sync"
"sync/atomic"
"testing"
"time"

"github.com/google/uuid"
"github.com/stretchr/testify/assert"
Expand Down Expand Up @@ -280,7 +281,10 @@ type testTrigger struct {
mu sync.Mutex
registrations map[string]chan capabilities.TriggerResponse
registerCount atomic.Int32
failAfter int32 // fail RegisterTrigger after this many successful calls (-1 = never fail)
inFlight atomic.Int32
maxInFlight atomic.Int32
failAfter int32 // fail RegisterTrigger after this many successful calls (-1 = never fail)
registerDelay time.Duration // delay each RegisterTrigger call to simulate a slow RPC
}

func newTestTrigger(name string) *testTrigger {
Expand All @@ -302,6 +306,21 @@ func newTestTriggerWithFailures(name string, failAfter int32) *testTrigger {

func (t *testTrigger) RegisterTrigger(_ context.Context, req capabilities.TriggerRegistrationRequest) (<-chan capabilities.TriggerResponse, error) {
count := t.registerCount.Add(1)

// track how many RegisterTrigger calls are in flight concurrently
inFlight := t.inFlight.Add(1)
for {
cur := t.maxInFlight.Load()
if inFlight <= cur || t.maxInFlight.CompareAndSwap(cur, inFlight) {
break
}
}
defer t.inFlight.Add(-1)

if t.registerDelay > 0 {
time.Sleep(t.registerDelay)
}

t.mu.Lock()
defer t.mu.Unlock()
if t.failAfter >= 0 && count > t.failAfter {
Expand Down Expand Up @@ -345,6 +364,11 @@ func (t *testTrigger) GetRegistrationCount() int32 {
return t.registerCount.Load()
}

// GetMaxInFlight returns the highest number of RegisterTrigger calls observed concurrently.
func (t *testTrigger) GetMaxInFlight() int32 {
return t.maxInFlight.Load()
}

func (t *testTrigger) GetState() connectivity.State {
return connectivity.Shutdown
}
Expand Down Expand Up @@ -375,3 +399,200 @@ func TestAtomicTrigger_RegistrationsReplayed(t *testing.T) {
resp1 := <-outCh
assert.Equal(t, "event1", resp1.Event.ID)
}

func TestAtomicTrigger_ConcurrentRegistrations(t *testing.T) {
ctx := t.Context()
r := registry.NewBaseRegistry(logger.Test(t))

const (
nTriggers = 10
delay = 100 * time.Millisecond
)

trigger := newTestTrigger("trigger")
trigger.registerDelay = delay
require.NoError(t, r.Add(ctx, trigger))

tc, err := r.GetTrigger(ctx, "trigger@1.0.0")
require.NoError(t, err)

outChs := make([]<-chan capabilities.TriggerResponse, nTriggers)
start := time.Now()
var wg sync.WaitGroup
for i := range nTriggers {
wg.Add(1)
go func() {
defer wg.Done()
outCh, err := tc.RegisterTrigger(ctx, capabilities.TriggerRegistrationRequest{TriggerID: fmt.Sprintf("reg%d", i)})
assert.NoError(t, err)
outChs[i] = outCh
}()
}
wg.Wait()
elapsed := time.Since(start)

require.Equal(t, int32(nTriggers), trigger.GetRegistrationCount())
// Registrations of different triggers must not serialize on a single lock: each call to the
// underlying capability takes delay, so a serialized implementation would need at least
// nTriggers*delay in total.
assert.GreaterOrEqual(t, trigger.GetMaxInFlight(), int32(2), "registrations should run concurrently in the underlying capability")
assert.Less(t, elapsed, time.Duration(nTriggers)*delay, "concurrent registrations should complete faster than serial ones")

// events flow through every registration's channel
for i := range nTriggers {
require.NotNil(t, outChs[i])
require.True(t, trigger.SendEvent(fmt.Sprintf("reg%d", i), capabilities.TriggerResponse{Event: capabilities.TriggerEvent{ID: fmt.Sprintf("event%d", i)}}))
}
for i := range nTriggers {
select {
case resp := <-outChs[i]:
assert.Equal(t, fmt.Sprintf("event%d", i), resp.Event.ID)
case <-time.After(2 * time.Second):
t.Fatalf("timed out waiting for event on registration %d", i)
}
}
}

func TestAtomicTrigger_ConcurrentSameTriggerRegistrations(t *testing.T) {
ctx := t.Context()
r := registry.NewBaseRegistry(logger.Test(t))

trigger := newTestTrigger("trigger")
trigger.registerDelay = 10 * time.Millisecond
require.NoError(t, r.Add(ctx, trigger))

tc, err := r.GetTrigger(ctx, "trigger@1.0.0")
require.NoError(t, err)

const n = 5
outChs := make([]<-chan capabilities.TriggerResponse, n)
var wg sync.WaitGroup
for i := range n {
wg.Add(1)
go func() {
defer wg.Done()
outCh, err := tc.RegisterTrigger(ctx, capabilities.TriggerRegistrationRequest{TriggerID: "same-reg"})
assert.NoError(t, err)
outChs[i] = outCh
}()
}
wg.Wait()

// every call registers with the underlying capability and all callers of the same trigger
// share one channel, matching the sequential behaviour
require.Equal(t, int32(n), trigger.GetRegistrationCount())
for i := range n {
require.NotNil(t, outChs[i])
assert.Equal(t, outChs[0], outChs[i])
}

require.True(t, trigger.SendEvent("same-reg", capabilities.TriggerResponse{Event: capabilities.TriggerEvent{ID: "event1"}}))
select {
case resp := <-outChs[0]:
assert.Equal(t, "event1", resp.Event.ID)
case <-time.After(2 * time.Second):
t.Fatal("timed out waiting for event")
}
}

func TestAtomicTrigger_Unregister(t *testing.T) {
ctx := t.Context()
r := registry.NewBaseRegistry(logger.Test(t))

trigger := newTestTrigger("trigger")
require.NoError(t, r.Add(ctx, trigger))

tc, err := r.GetTrigger(ctx, "trigger@1.0.0")
require.NoError(t, err)

outCh, err := tc.RegisterTrigger(ctx, capabilities.TriggerRegistrationRequest{TriggerID: "reg1"})
require.NoError(t, err)

// unregistering an unknown trigger is a no-op on the cache and still hits the underlying capability
require.NoError(t, tc.UnregisterTrigger(ctx, capabilities.TriggerRegistrationRequest{TriggerID: "unknown"}))

require.NoError(t, tc.UnregisterTrigger(ctx, capabilities.TriggerRegistrationRequest{TriggerID: "reg1"}))

// the caller's channel is closed and the registration is not replayed on rebind
select {
case _, ok := <-outCh:
require.False(t, ok, "expected the registration channel to be closed")
case <-time.After(2 * time.Second):
t.Fatal("expected the registration channel to be closed")
}

trigger2 := newTestTrigger("trigger")
require.NoError(t, r.Add(ctx, trigger2))
assert.Equal(t, int32(0), trigger2.GetRegistrationCount())

// re-registering after an unregister works and gets a fresh channel
outCh2, err := tc.RegisterTrigger(ctx, capabilities.TriggerRegistrationRequest{TriggerID: "reg1"})
require.NoError(t, err)
require.NotNil(t, outCh2)
assert.NotEqual(t, outCh, outCh2)

require.True(t, trigger2.SendEvent("reg1", capabilities.TriggerResponse{Event: capabilities.TriggerEvent{ID: "event2"}}))
select {
case resp := <-outCh2:
assert.Equal(t, "event2", resp.Event.ID)
case <-time.After(2 * time.Second):
t.Fatal("timed out waiting for event")
}
}

func TestAtomicTrigger_ConcurrentRegisterUnregisterUpdate(t *testing.T) {
ctx := t.Context()
r := registry.NewBaseRegistry(logger.Test(t))

trigger := newTestTrigger("trigger")
trigger.registerDelay = time.Millisecond
require.NoError(t, r.Add(ctx, trigger))

tc, err := r.GetTrigger(ctx, "trigger@1.0.0")
require.NoError(t, err)

const (
nGoroutines = 8
nIterations = 25
nTriggers = 5
)

var wg sync.WaitGroup
for g := range nGoroutines {
wg.Add(1)
go func() {
defer wg.Done()
for it := range nIterations {
id := fmt.Sprintf("reg%d", (g+it)%nTriggers)
switch it % 3 {
case 0, 1:
outCh, regErr := tc.RegisterTrigger(ctx, capabilities.TriggerRegistrationRequest{TriggerID: id})
assert.NoError(t, regErr)
if outCh != nil {
// the channel may already be closed by a concurrent unregister
select {
case _, ok := <-outCh:
assert.False(t, ok)
default:
}
}
case 2:
assert.NoError(t, tc.UnregisterTrigger(ctx, capabilities.TriggerRegistrationRequest{TriggerID: id}))
}
if it%7 == 0 {
// swap the underlying capability to exercise rebind concurrently
assert.NoError(t, r.Add(ctx, newTestTrigger("trigger")))
}
}
}()
}
wg.Wait()

// after the storm, unregister everything and swap in a fresh capability: nothing should be replayed
for i := range nTriggers {
require.NoError(t, tc.UnregisterTrigger(ctx, capabilities.TriggerRegistrationRequest{TriggerID: fmt.Sprintf("reg%d", i)}))
}
final := newTestTrigger("trigger")
require.NoError(t, r.Add(ctx, final))
assert.Equal(t, int32(0), final.GetRegistrationCount())
}
Loading
Loading