diff --git a/pkg/sentry/devices/tundev/tundev.go b/pkg/sentry/devices/tundev/tundev.go index b4e2a6d91..d197c52d7 100644 --- a/pkg/sentry/devices/tundev/tundev.go +++ b/pkg/sentry/devices/tundev/tundev.go @@ -152,8 +152,8 @@ func (fd *tunFD) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements watier.Waitable.EventRegister. -func (fd *tunFD) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - fd.device.EventRegister(e, mask) +func (fd *tunFD) EventRegister(e *waiter.Entry) { + fd.device.EventRegister(e) } // EventUnregister implements watier.Waitable.EventUnregister. diff --git a/pkg/sentry/fs/dev/net_tun.go b/pkg/sentry/fs/dev/net_tun.go index 1abf11142..4b8a2d4f9 100644 --- a/pkg/sentry/fs/dev/net_tun.go +++ b/pkg/sentry/fs/dev/net_tun.go @@ -158,8 +158,8 @@ func (n *netTunFileOperations) Readiness(mask waiter.EventMask) waiter.EventMask } // EventRegister implements watier.Waitable.EventRegister. -func (n *netTunFileOperations) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - n.device.EventRegister(e, mask) +func (n *netTunFileOperations) EventRegister(e *waiter.Entry) { + n.device.EventRegister(e) } // EventUnregister implements watier.Waitable.EventUnregister. diff --git a/pkg/sentry/fs/fdpipe/pipe.go b/pkg/sentry/fs/fdpipe/pipe.go index d2eb03bb7..2d79f89d7 100644 --- a/pkg/sentry/fs/fdpipe/pipe.go +++ b/pkg/sentry/fs/fdpipe/pipe.go @@ -99,8 +99,8 @@ func (p *pipeOperations) init() error { } // EventRegister implements waiter.Waitable.EventRegister. -func (p *pipeOperations) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - p.Queue.EventRegister(e, mask) +func (p *pipeOperations) EventRegister(e *waiter.Entry) { + p.Queue.EventRegister(e) fdnotifier.UpdateFD(int32(p.file.FD())) } diff --git a/pkg/sentry/fs/file.go b/pkg/sentry/fs/file.go index df04f044d..e5d542419 100644 --- a/pkg/sentry/fs/file.go +++ b/pkg/sentry/fs/file.go @@ -182,8 +182,8 @@ func (f *File) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements waiter.Waitable.EventRegister. -func (f *File) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - f.FileOperations.EventRegister(e, mask) +func (f *File) EventRegister(e *waiter.Entry) { + f.FileOperations.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. diff --git a/pkg/sentry/fs/file_overlay.go b/pkg/sentry/fs/file_overlay.go index a27dd0b9a..ae2f959b2 100644 --- a/pkg/sentry/fs/file_overlay.go +++ b/pkg/sentry/fs/file_overlay.go @@ -100,14 +100,14 @@ func (f *overlayFileOperations) Release(ctx context.Context) { } // EventRegister implements FileOperations.EventRegister. -func (f *overlayFileOperations) EventRegister(we *waiter.Entry, mask waiter.EventMask) { +func (f *overlayFileOperations) EventRegister(we *waiter.Entry) { f.upperMu.Lock() defer f.upperMu.Unlock() if f.upper != nil { - f.upper.EventRegister(we, mask) + f.upper.EventRegister(we) return } - f.lower.EventRegister(we, mask) + f.lower.EventRegister(we) } // EventUnregister implements FileOperations.Unregister. diff --git a/pkg/sentry/fs/host/file.go b/pkg/sentry/fs/host/file.go index 1d0d95634..1d405c782 100644 --- a/pkg/sentry/fs/host/file.go +++ b/pkg/sentry/fs/host/file.go @@ -149,8 +149,8 @@ func newFile(ctx context.Context, dirent *fs.Dirent, flags fs.FileFlags, iops *i } // EventRegister implements waiter.Waitable.EventRegister. -func (f *fileOperations) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - f.iops.fileState.queue.EventRegister(e, mask) +func (f *fileOperations) EventRegister(e *waiter.Entry) { + f.iops.fileState.queue.EventRegister(e) fdnotifier.UpdateFD(int32(f.iops.fileState.FD())) } diff --git a/pkg/sentry/fs/host/wait_test.go b/pkg/sentry/fs/host/wait_test.go index bd6188e03..db4590ec6 100644 --- a/pkg/sentry/fs/host/wait_test.go +++ b/pkg/sentry/fs/host/wait_test.go @@ -46,8 +46,8 @@ func TestWait(t *testing.T) { t.Fatalf("File is ready for read when it shouldn't be.") } - e, ch := waiter.NewChannelEntry(nil) - file.EventRegister(&e, waiter.ReadableEvents) + e, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + file.EventRegister(&e) defer file.EventUnregister(&e) // Check that there are no notifications yet. diff --git a/pkg/sentry/fs/lock/lock.go b/pkg/sentry/fs/lock/lock.go index e39d340fe..3731a61b4 100644 --- a/pkg/sentry/fs/lock/lock.go +++ b/pkg/sentry/fs/lock/lock.go @@ -163,8 +163,8 @@ func (l *Locks) LockRegion(uid UniqueID, ownerPID int32, t LockType, r LockRange // continue blocking. res := l.locks.lock(uid, ownerPID, t, r) if !res && block != nil { - e, ch := waiter.NewChannelEntry(nil) - l.blockedQueue.EventRegister(&e, EventMaskAll) + e, ch := waiter.NewChannelEntry(EventMaskAll) + l.blockedQueue.EventRegister(&e) l.mu.Unlock() if err := block.Block(ch); err != nil { // We were interrupted, the caller can translate this to EINTR if applicable. diff --git a/pkg/sentry/fs/timerfd/timerfd.go b/pkg/sentry/fs/timerfd/timerfd.go index 457227954..66ae9debb 100644 --- a/pkg/sentry/fs/timerfd/timerfd.go +++ b/pkg/sentry/fs/timerfd/timerfd.go @@ -108,8 +108,8 @@ func (t *TimerOperations) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements waiter.Waitable.EventRegister. -func (t *TimerOperations) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - t.events.EventRegister(e, mask) +func (t *TimerOperations) EventRegister(e *waiter.Entry) { + t.events.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. diff --git a/pkg/sentry/fs/tty/master.go b/pkg/sentry/fs/tty/master.go index 88d6703a8..2003e658a 100644 --- a/pkg/sentry/fs/tty/master.go +++ b/pkg/sentry/fs/tty/master.go @@ -128,8 +128,8 @@ func (mf *masterFileOperations) Release(ctx context.Context) { } // EventRegister implements waiter.Waitable.EventRegister. -func (mf *masterFileOperations) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - mf.t.ld.masterWaiter.EventRegister(e, mask) +func (mf *masterFileOperations) EventRegister(e *waiter.Entry) { + mf.t.ld.masterWaiter.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. diff --git a/pkg/sentry/fs/tty/replica.go b/pkg/sentry/fs/tty/replica.go index ca5bc7535..272ce3ae9 100644 --- a/pkg/sentry/fs/tty/replica.go +++ b/pkg/sentry/fs/tty/replica.go @@ -113,8 +113,8 @@ func (sf *replicaFileOperations) Release(context.Context) { } // EventRegister implements waiter.Waitable.EventRegister. -func (sf *replicaFileOperations) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - sf.si.t.ld.replicaWaiter.EventRegister(e, mask) +func (sf *replicaFileOperations) EventRegister(e *waiter.Entry) { + sf.si.t.ld.replicaWaiter.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. diff --git a/pkg/sentry/fsimpl/devpts/devpts_test.go b/pkg/sentry/fsimpl/devpts/devpts_test.go index 1ef07d702..80e14b0f2 100644 --- a/pkg/sentry/fsimpl/devpts/devpts_test.go +++ b/pkg/sentry/fsimpl/devpts/devpts_test.go @@ -56,12 +56,6 @@ func TestSimpleMasterToReplica(t *testing.T) { } } -type callback func(*waiter.Entry, waiter.EventMask) - -func (cb callback) Callback(entry *waiter.Entry, mask waiter.EventMask) { - cb(entry, mask) -} - func TestEchoDeadlock(t *testing.T) { ctx := contexttest.Context(t) termios := linux.DefaultReplicaTermios @@ -69,11 +63,11 @@ func TestEchoDeadlock(t *testing.T) { ld := newLineDiscipline(termios) outBytes := make([]byte, 32) dst := usermem.BytesIOSequence(outBytes) - entry := &waiter.Entry{Callback: callback(func(*waiter.Entry, waiter.EventMask) { + entry := waiter.NewFunctionEntry(waiter.ReadableEvents, func(waiter.EventMask) { ld.inputQueueRead(ctx, dst) - })} - ld.masterWaiter.EventRegister(entry, waiter.ReadableEvents) - defer ld.masterWaiter.EventUnregister(entry) + }) + ld.masterWaiter.EventRegister(&entry) + defer ld.masterWaiter.EventUnregister(&entry) inBytes := []byte("hello, tty\n") n, err := ld.inputQueueWrite(ctx, usermem.BytesIOSequence(inBytes)) if err != nil { diff --git a/pkg/sentry/fsimpl/devpts/master.go b/pkg/sentry/fsimpl/devpts/master.go index 9a1a245dc..0ce84302f 100644 --- a/pkg/sentry/fsimpl/devpts/master.go +++ b/pkg/sentry/fsimpl/devpts/master.go @@ -103,8 +103,8 @@ func (mfd *masterFileDescription) Release(ctx context.Context) { } // EventRegister implements waiter.Waitable.EventRegister. -func (mfd *masterFileDescription) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - mfd.t.ld.masterWaiter.EventRegister(e, mask) +func (mfd *masterFileDescription) EventRegister(e *waiter.Entry) { + mfd.t.ld.masterWaiter.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. diff --git a/pkg/sentry/fsimpl/devpts/replica.go b/pkg/sentry/fsimpl/devpts/replica.go index e251897b4..5b6311d3c 100644 --- a/pkg/sentry/fsimpl/devpts/replica.go +++ b/pkg/sentry/fsimpl/devpts/replica.go @@ -112,8 +112,8 @@ var _ vfs.FileDescriptionImpl = (*replicaFileDescription)(nil) func (rfd *replicaFileDescription) Release(ctx context.Context) {} // EventRegister implements waiter.Waitable.EventRegister. -func (rfd *replicaFileDescription) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - rfd.inode.t.ld.replicaWaiter.EventRegister(e, mask) +func (rfd *replicaFileDescription) EventRegister(e *waiter.Entry) { + rfd.inode.t.ld.replicaWaiter.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. diff --git a/pkg/sentry/fsimpl/eventfd/eventfd.go b/pkg/sentry/fsimpl/eventfd/eventfd.go index af5ba5131..8ba81a74e 100644 --- a/pkg/sentry/fsimpl/eventfd/eventfd.go +++ b/pkg/sentry/fsimpl/eventfd/eventfd.go @@ -266,8 +266,8 @@ func (efd *EventFileDescription) Readiness(mask waiter.EventMask) waiter.EventMa } // EventRegister implements waiter.Waitable.EventRegister. -func (efd *EventFileDescription) EventRegister(entry *waiter.Entry, mask waiter.EventMask) { - efd.queue.EventRegister(entry, mask) +func (efd *EventFileDescription) EventRegister(entry *waiter.Entry) { + efd.queue.EventRegister(entry) efd.mu.Lock() defer efd.mu.Unlock() diff --git a/pkg/sentry/fsimpl/eventfd/eventfd_test.go b/pkg/sentry/fsimpl/eventfd/eventfd_test.go index 85718f813..fcf759b87 100644 --- a/pkg/sentry/fsimpl/eventfd/eventfd_test.go +++ b/pkg/sentry/fsimpl/eventfd/eventfd_test.go @@ -48,8 +48,8 @@ func TestEventFD(t *testing.T) { defer eventfd.DecRef(ctx) // Register a callback for a write event. - w, ch := waiter.NewChannelEntry(nil) - eventfd.EventRegister(&w, waiter.ReadableEvents) + w, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + eventfd.EventRegister(&w) defer eventfd.EventUnregister(&w) data := []byte("00000124") diff --git a/pkg/sentry/fsimpl/fuse/dev.go b/pkg/sentry/fsimpl/fuse/dev.go index 0f855ac59..f2c688219 100644 --- a/pkg/sentry/fsimpl/fuse/dev.go +++ b/pkg/sentry/fsimpl/fuse/dev.go @@ -378,8 +378,8 @@ func (fd *DeviceFD) readinessLocked(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements waiter.Waitable.EventRegister. -func (fd *DeviceFD) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - fd.waitQueue.EventRegister(e, mask) +func (fd *DeviceFD) EventRegister(e *waiter.Entry) { + fd.waitQueue.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. diff --git a/pkg/sentry/fsimpl/fuse/dev_test.go b/pkg/sentry/fsimpl/fuse/dev_test.go index 13b32fc7c..6e9dfe5ef 100644 --- a/pkg/sentry/fsimpl/fuse/dev_test.go +++ b/pkg/sentry/fsimpl/fuse/dev_test.go @@ -180,8 +180,8 @@ func ReadTest(serverTask *kernel.Task, fd *vfs.FileDescription, inIOseq usermem. dev := fd.Impl().(*DeviceFD) // Register for notifications. - w, ch := waiter.NewChannelEntry(nil) - dev.EventRegister(&w, waiter.ReadableEvents) + w, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + dev.EventRegister(&w) for { // Issue the request and break out if it completes with anything other than // "would block". diff --git a/pkg/sentry/fsimpl/gofer/special_file.go b/pkg/sentry/fsimpl/gofer/special_file.go index c568bbfd2..1536601b9 100644 --- a/pkg/sentry/fsimpl/gofer/special_file.go +++ b/pkg/sentry/fsimpl/gofer/special_file.go @@ -165,13 +165,13 @@ func (fd *specialFileFD) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements waiter.Waitable.EventRegister. -func (fd *specialFileFD) EventRegister(e *waiter.Entry, mask waiter.EventMask) { +func (fd *specialFileFD) EventRegister(e *waiter.Entry) { if fd.haveQueue { - fd.queue.EventRegister(e, mask) + fd.queue.EventRegister(e) fdnotifier.UpdateFD(fd.handle.fd) return } - fd.fileDescription.EventRegister(e, mask) + fd.fileDescription.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. diff --git a/pkg/sentry/fsimpl/host/host.go b/pkg/sentry/fsimpl/host/host.go index b970cef14..a27513103 100644 --- a/pkg/sentry/fsimpl/host/host.go +++ b/pkg/sentry/fsimpl/host/host.go @@ -892,8 +892,8 @@ func (f *fileDescription) ConfigureMMap(_ context.Context, opts *memmap.MMapOpts } // EventRegister implements waiter.Waitable.EventRegister. -func (f *fileDescription) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - f.inode.queue.EventRegister(e, mask) +func (f *fileDescription) EventRegister(e *waiter.Entry) { + f.inode.queue.EventRegister(e) if f.inode.mayBlock { fdnotifier.UpdateFD(int32(f.inode.hostFD)) } diff --git a/pkg/sentry/fsimpl/mqfs/queue.go b/pkg/sentry/fsimpl/mqfs/queue.go index 933dbc6ed..dae7208a9 100644 --- a/pkg/sentry/fsimpl/mqfs/queue.go +++ b/pkg/sentry/fsimpl/mqfs/queue.go @@ -135,8 +135,8 @@ func (fd *queueFD) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements Waitable.EventRegister. -func (fd *queueFD) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - fd.queue.EventRegister(e, mask) +func (fd *queueFD) EventRegister(e *waiter.Entry) { + fd.queue.EventRegister(e) } // EventUnregister implements Waitable.EventUnregister. diff --git a/pkg/sentry/fsimpl/overlay/regular_file.go b/pkg/sentry/fsimpl/overlay/regular_file.go index 156ffeaeb..719990346 100644 --- a/pkg/sentry/fsimpl/overlay/regular_file.go +++ b/pkg/sentry/fsimpl/overlay/regular_file.go @@ -63,7 +63,7 @@ type regularFileFD struct { // If copiedUp is false, lowerWaiters contains all waiter.Entries // registered with cachedFD. lowerWaiters is protected by mu. - lowerWaiters map[*waiter.Entry]waiter.EventMask + lowerWaiters map[*waiter.Entry]struct{} } func (fd *regularFileFD) getCurrentFD(ctx context.Context) (*vfs.FileDescription, error) { @@ -101,12 +101,10 @@ func (fd *regularFileFD) currentFDLocked(ctx context.Context) (*vfs.FileDescript } if len(fd.lowerWaiters) != 0 { ready := upperFD.Readiness(^waiter.EventMask(0)) - for e, mask := range fd.lowerWaiters { + for e := range fd.lowerWaiters { fd.cachedFD.EventUnregister(e) - upperFD.EventRegister(e, mask) - if m := ready & mask; m != 0 { - e.Callback.Callback(e, m) - } + upperFD.EventRegister(e) + e.NotifyEvent(ready) } } fd.cachedFD.DecRef(ctx) @@ -256,7 +254,7 @@ func (fd *regularFileFD) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements waiter.Waitable.EventRegister. -func (fd *regularFileFD) EventRegister(e *waiter.Entry, mask waiter.EventMask) { +func (fd *regularFileFD) EventRegister(e *waiter.Entry) { fd.mu.Lock() defer fd.mu.Unlock() wrappedFD, err := fd.currentFDLocked(context.Background()) @@ -267,12 +265,12 @@ func (fd *regularFileFD) EventRegister(e *waiter.Entry, mask waiter.EventMask) { log.Warningf("overlay.regularFileFD.EventRegister: currentFDLocked failed: %v", err) wrappedFD = fd.cachedFD } - wrappedFD.EventRegister(e, mask) + wrappedFD.EventRegister(e) if !fd.copiedUp { if fd.lowerWaiters == nil { - fd.lowerWaiters = make(map[*waiter.Entry]waiter.EventMask) + fd.lowerWaiters = make(map[*waiter.Entry]struct{}) } - fd.lowerWaiters[e] = mask + fd.lowerWaiters[e] = struct{}{} } } diff --git a/pkg/sentry/fsimpl/signalfd/signalfd.go b/pkg/sentry/fsimpl/signalfd/signalfd.go index bdb03ef96..746d205c9 100644 --- a/pkg/sentry/fsimpl/signalfd/signalfd.go +++ b/pkg/sentry/fsimpl/signalfd/signalfd.go @@ -44,11 +44,14 @@ type SignalFileDescription struct { // will undoubtedly become very complicated quickly. target *kernel.Task - // mu protects mask. + // queue is the queue for listeners. + queue waiter.Queue + + // mu protects entry. mu sync.Mutex `state:"nosave"` - // mask is the signal mask. Protected by mu. - mask linux.SignalSet + // entry is the entry in the task signal queue. + entry waiter.Entry } var _ vfs.FileDescriptionImpl = (*SignalFileDescription)(nil) @@ -59,13 +62,15 @@ func New(vfsObj *vfs.VirtualFilesystem, target *kernel.Task, mask linux.SignalSe defer vd.DecRef(target) sfd := &SignalFileDescription{ target: target, - mask: mask, } + sfd.entry.Init(sfd, waiter.EventMask(mask)) + sfd.target.SignalRegister(&sfd.entry) if err := sfd.vfsfd.Init(sfd, flags, vd.Mount(), vd.Dentry(), &vfs.FileDescriptionOptions{ UseDentryMetadata: true, DenyPRead: true, DenyPWrite: true, }); err != nil { + sfd.target.SignalUnregister(&sfd.entry) return nil, err } return &sfd.vfsfd, nil @@ -75,14 +80,16 @@ func New(vfsObj *vfs.VirtualFilesystem, target *kernel.Task, mask linux.SignalSe func (sfd *SignalFileDescription) Mask() linux.SignalSet { sfd.mu.Lock() defer sfd.mu.Unlock() - return sfd.mask + return linux.SignalSet(sfd.entry.Mask()) } // SetMask sets the signal mask. func (sfd *SignalFileDescription) SetMask(mask linux.SignalSet) { sfd.mu.Lock() defer sfd.mu.Unlock() - sfd.mask = mask + sfd.target.SignalUnregister(&sfd.entry) + sfd.entry.Init(sfd, waiter.EventMask(mask)) + sfd.target.SignalRegister(&sfd.entry) } // Read implements vfs.FileDescriptionImpl.Read. @@ -117,25 +124,28 @@ func (sfd *SignalFileDescription) Read(ctx context.Context, dst usermem.IOSequen func (sfd *SignalFileDescription) Readiness(mask waiter.EventMask) waiter.EventMask { sfd.mu.Lock() defer sfd.mu.Unlock() - if mask&waiter.ReadableEvents != 0 && sfd.target.PendingSignals()&sfd.mask != 0 { + if mask&waiter.ReadableEvents != 0 && sfd.target.PendingSignals()&linux.SignalSet(sfd.entry.Mask()) != 0 { return waiter.ReadableEvents // Pending signals. } return 0 } // EventRegister implements waiter.Waitable.EventRegister. -func (sfd *SignalFileDescription) EventRegister(entry *waiter.Entry, _ waiter.EventMask) { - sfd.mu.Lock() - defer sfd.mu.Unlock() - // Register for the signal set; ignore the passed events. - sfd.target.SignalRegister(entry, waiter.EventMask(sfd.mask)) +func (sfd *SignalFileDescription) EventRegister(e *waiter.Entry) { + sfd.queue.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. -func (sfd *SignalFileDescription) EventUnregister(entry *waiter.Entry) { - // Unregister the original entry. - sfd.target.SignalUnregister(entry) +func (sfd *SignalFileDescription) EventUnregister(e *waiter.Entry) { + sfd.queue.EventUnregister(e) +} + +// NotifyEvent implements waiter.EventListener.NotifyEvent. +func (sfd *SignalFileDescription) NotifyEvent(mask waiter.EventMask) { + sfd.queue.Notify(waiter.EventIn) // Always notify data available. } // Release implements vfs.FileDescriptionImpl.Release. -func (sfd *SignalFileDescription) Release(context.Context) {} +func (sfd *SignalFileDescription) Release(context.Context) { + sfd.target.SignalUnregister(&sfd.entry) +} diff --git a/pkg/sentry/fsimpl/timerfd/timerfd.go b/pkg/sentry/fsimpl/timerfd/timerfd.go index 565dfa6a0..1f07419bf 100644 --- a/pkg/sentry/fsimpl/timerfd/timerfd.go +++ b/pkg/sentry/fsimpl/timerfd/timerfd.go @@ -112,8 +112,8 @@ func (tfd *TimerFileDescription) Readiness(mask waiter.EventMask) waiter.EventMa } // EventRegister implements waiter.Waitable.EventRegister. -func (tfd *TimerFileDescription) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - tfd.events.EventRegister(e, mask) +func (tfd *TimerFileDescription) EventRegister(e *waiter.Entry) { + tfd.events.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. diff --git a/pkg/sentry/kernel/epoll/epoll.go b/pkg/sentry/kernel/epoll/epoll.go index 8d0a21baf..93accd7c9 100644 --- a/pkg/sentry/kernel/epoll/epoll.go +++ b/pkg/sentry/kernel/epoll/epoll.go @@ -113,7 +113,7 @@ type EventPoll struct { // different lock to avoid circular lock acquisition order involving // the wait queue mutexes and mu. The full order is mu, observed file // wait queue mutex, then listsMu; this allows listsMu to be acquired - // when (*pollEntry).Callback is called. + // when (*pollEntry).NotifyEvent is called. // // An entry is always in one of the following lists: // readyList -- when there's a chance that it's ready to have @@ -122,8 +122,8 @@ type EventPoll struct { // readEvents() functions always call the entry's file // Readiness() function to confirm it's ready. // waitingList -- when there's no chance that the entry is ready, - // so it's waiting for the (*pollEntry).Callback to be called - // on it before it gets moved to the readyList. + // so it's waiting for the (*pollEntry).NotifyEvent to be + // called on it before it gets moved to the readyList. // disabledList -- when the entry is disabled. This happens when // a one-shot entry gets delivered via readEvents(). listsMu sync.Mutex `state:"nosave"` @@ -275,11 +275,11 @@ func (e *EventPoll) ReadEvents(max int) []linux.EpollEvent { return ret } -// Callback implements waiter.EntryCallback.Callback. +// NotifyEvent implements waiter.EventListener.NotifyEvent. // -// Callback is called when one of the files we're polling becomes ready. It +// NotifyEvent is called when one of the files we're polling becomes ready. It // moves said file to the readyList if it's currently in the waiting list. -func (p *pollEntry) Callback(*waiter.Entry, waiter.EventMask) { +func (p *pollEntry) NotifyEvent(waiter.EventMask) { e := p.epoll e.listsMu.Lock() @@ -309,11 +309,12 @@ func (e *EventPoll) initEntryReadiness(entry *pollEntry) { // Register for event notifications. f := entry.id.File - f.EventRegister(&entry.waiter, entry.mask) + entry.waiter.Init(entry, entry.mask) + f.EventRegister(&entry.waiter) // Check if the file happens to already be in a ready state. if ready := f.Readiness(entry.mask) & entry.mask; ready != 0 { - entry.Callback(&entry.waiter, ready) + entry.NotifyEvent(ready) } } @@ -385,7 +386,7 @@ func (e *EventPoll) AddEntry(id FileIdentifier, flags EntryFlags, mask waiter.Ev flags: flags, mask: mask, } - entry.waiter.Callback = entry + entry.waiter.Init(entry, mask) e.files[id] = entry entry.file = refs.NewWeakRef(id.File, entry) @@ -408,7 +409,8 @@ func (e *EventPoll) UpdateEntry(id FileIdentifier, flags EntryFlags, mask waiter } // Unregister the old mask and remove entry from the list it's in, so - // (*pollEntry).Callback is guaranteed to not be called on this entry anymore. + // (*pollEntry).NotifyEvent is guaranteed to not be called on this + // entry anymore. entry.id.File.EventUnregister(&entry.waiter) // Remove entry from whatever list it's in. This ensure that no other diff --git a/pkg/sentry/kernel/eventfd/eventfd.go b/pkg/sentry/kernel/eventfd/eventfd.go index bf625dede..7bef7d41a 100644 --- a/pkg/sentry/kernel/eventfd/eventfd.go +++ b/pkg/sentry/kernel/eventfd/eventfd.go @@ -264,8 +264,8 @@ func (e *EventOperations) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements waiter.Waitable.EventRegister. -func (e *EventOperations) EventRegister(entry *waiter.Entry, mask waiter.EventMask) { - e.wq.EventRegister(entry, mask) +func (e *EventOperations) EventRegister(entry *waiter.Entry) { + e.wq.EventRegister(entry) e.mu.Lock() defer e.mu.Unlock() diff --git a/pkg/sentry/kernel/eventfd/eventfd_test.go b/pkg/sentry/kernel/eventfd/eventfd_test.go index 1b9e60b3a..cf3e47461 100644 --- a/pkg/sentry/kernel/eventfd/eventfd_test.go +++ b/pkg/sentry/kernel/eventfd/eventfd_test.go @@ -38,8 +38,8 @@ func TestEventfd(t *testing.T) { event := New(ctx, initVal, false) // Register a callback for a write event. - w, ch := waiter.NewChannelEntry(nil) - event.EventRegister(&w, waiter.ReadableEvents) + w, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + event.EventRegister(&w) defer event.EventUnregister(&w) data := []byte("00000124") diff --git a/pkg/sentry/kernel/fasync/fasync.go b/pkg/sentry/kernel/fasync/fasync.go index 473987a79..b4036cea7 100644 --- a/pkg/sentry/kernel/fasync/fasync.go +++ b/pkg/sentry/kernel/fasync/fasync.go @@ -95,8 +95,8 @@ type FileAsync struct { recipientT *kernel.Task } -// Callback sends a signal. -func (a *FileAsync) Callback(e *waiter.Entry, mask waiter.EventMask) { +// NotifyEvent implements waiter.EventListener.NotifyEvent. +func (a *FileAsync) NotifyEvent(mask waiter.EventMask) { a.mu.Lock() defer a.mu.Unlock() if !a.registered { @@ -149,19 +149,14 @@ func (a *FileAsync) Register(w waiter.Waitable) { a.regMu.Lock() defer a.regMu.Unlock() a.mu.Lock() - if a.registered { a.mu.Unlock() panic("registering already registered file") } - - if a.e.Callback == nil { - a.e.Callback = a - } + a.e.Init(a, waiter.ReadableEvents|waiter.WritableEvents|waiter.EventErr|waiter.EventHUp) a.registered = true - a.mu.Unlock() - w.EventRegister(&a.e, waiter.ReadableEvents|waiter.WritableEvents|waiter.EventErr|waiter.EventHUp) + w.EventRegister(&a.e) } // Unregister stops monitoring a file. diff --git a/pkg/sentry/kernel/mq/mq.go b/pkg/sentry/kernel/mq/mq.go index 7515a2772..e38a7fae1 100644 --- a/pkg/sentry/kernel/mq/mq.go +++ b/pkg/sentry/kernel/mq/mq.go @@ -263,13 +263,8 @@ type Queue struct { // mu protects all the fields below. mu sync.Mutex `state:"nosave"` - // senders is a queue of currently blocked senders. Senders are notified - // when space isi available in the queue for a new message. - senders waiter.Queue - - // receivers is a queue of currently blocked receivers. Receivers are - // notified when a new message is inserted in the queue. - receivers waiter.Queue + // queue is the queue of waiters. + queue waiter.Queue // messages is a list of messages currently in the queue. messages msgList @@ -423,25 +418,17 @@ func (q *Queue) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements Waitable.EventRegister. -func (q *Queue) EventRegister(e *waiter.Entry, mask waiter.EventMask) { +func (q *Queue) EventRegister(e *waiter.Entry) { q.mu.Lock() defer q.mu.Unlock() - - if mask&waiter.WritableEvents != 0 { - q.senders.EventRegister(e, waiter.EventOut) - } - if mask&waiter.ReadableEvents != 0 { - q.receivers.EventRegister(e, waiter.EventIn) - } + q.queue.EventRegister(e) } // EventUnregister implements Waitable.EventUnregister. func (q *Queue) EventUnregister(e *waiter.Entry) { q.mu.Lock() defer q.mu.Unlock() - - q.senders.EventUnregister(e) - q.receivers.EventUnregister(e) + q.queue.EventUnregister(e) } // HasPermissions returns true if the given credentials meet the access diff --git a/pkg/sentry/kernel/msgqueue/msgqueue.go b/pkg/sentry/kernel/msgqueue/msgqueue.go index c7c5e41fb..f891ac9eb 100644 --- a/pkg/sentry/kernel/msgqueue/msgqueue.go +++ b/pkg/sentry/kernel/msgqueue/msgqueue.go @@ -277,8 +277,8 @@ func (q *Queue) Send(ctx context.Context, m Message, b Blocker, wait bool, pid i // Slow path: at this point, the queue was found to be full, and we were // asked to block. - e, ch := waiter.NewChannelEntry(nil) - q.senders.EventRegister(&e, waiter.EventOut) + e, ch := waiter.NewChannelEntry(waiter.EventOut) + q.senders.EventRegister(&e) defer q.senders.EventUnregister(&e) // Note: we need to check again before blocking the first time since space @@ -373,8 +373,8 @@ func (q *Queue) Receive(ctx context.Context, b Blocker, mType int64, maxSize int // Slow path: at this point, the queue was found to be empty, and we were // asked to block. - e, ch := waiter.NewChannelEntry(nil) - q.receivers.EventRegister(&e, waiter.EventIn) + e, ch := waiter.NewChannelEntry(waiter.EventIn) + q.receivers.EventRegister(&e) defer q.receivers.EventUnregister(&e) // Note: we need to check again before blocking the first time since a diff --git a/pkg/sentry/kernel/pipe/pipe_test.go b/pkg/sentry/kernel/pipe/pipe_test.go index aa3ab305d..d76033195 100644 --- a/pkg/sentry/kernel/pipe/pipe_test.go +++ b/pkg/sentry/kernel/pipe/pipe_test.go @@ -96,8 +96,8 @@ func TestPipeWriteUntilEnd(t *testing.T) { ctx := contexttest.Context(t) buf := make([]byte, len(msg)+1) dst := usermem.BytesIOSequence(buf) - e, ch := waiter.NewChannelEntry(nil) - r.EventRegister(&e, waiter.ReadableEvents) + e, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + r.EventRegister(&e) defer r.EventUnregister(&e) for { n, err := r.Readv(ctx, dst) @@ -123,8 +123,8 @@ func TestPipeWriteUntilEnd(t *testing.T) { }() src := usermem.BytesIOSequence(msg) - e, ch := waiter.NewChannelEntry(nil) - w.EventRegister(&e, waiter.WritableEvents) + e, ch := waiter.NewChannelEntry(waiter.WritableEvents) + w.EventRegister(&e) defer w.EventUnregister(&e) for src.NumBytes() != 0 { n, err := w.Writev(ctx, src) diff --git a/pkg/sentry/kernel/pipe/vfs.go b/pkg/sentry/kernel/pipe/vfs.go index a6f1989f5..06325d99f 100644 --- a/pkg/sentry/kernel/pipe/vfs.go +++ b/pkg/sentry/kernel/pipe/vfs.go @@ -228,8 +228,8 @@ func (fd *VFSPipeFD) Allocate(ctx context.Context, mode, offset, length uint64) } // EventRegister implements waiter.Waitable.EventRegister. -func (fd *VFSPipeFD) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - fd.pipe.EventRegister(e, mask) +func (fd *VFSPipeFD) EventRegister(e *waiter.Entry) { + fd.pipe.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. diff --git a/pkg/sentry/kernel/signalfd/signalfd.go b/pkg/sentry/kernel/signalfd/signalfd.go index 9c5e6698c..0dca88579 100644 --- a/pkg/sentry/kernel/signalfd/signalfd.go +++ b/pkg/sentry/kernel/signalfd/signalfd.go @@ -32,7 +32,6 @@ import ( // // +stateify savable type SignalOperations struct { - fsutil.FileNoopRelease `state:"nosave"` fsutil.FilePipeSeek `state:"nosave"` fsutil.FileNotDirReaddir `state:"nosave"` fsutil.FileNoIoctl `state:"nosave"` @@ -52,11 +51,14 @@ type SignalOperations struct { // will undoubtedly become very complicated quickly. target *kernel.Task + // queue is the set of listeners. + queue waiter.Queue + // mu protects below. mu sync.Mutex `state:"nosave"` - // mask is the signal mask. Protected by mu. - mask linux.SignalSet + // entry is the entry reigstered with the target. + entry waiter.Entry } // New creates a new signalfd object with the supplied mask. @@ -68,28 +70,31 @@ func New(ctx context.Context, mask linux.SignalSet) (*fs.File, error) { } // name matches fs/signalfd.c:signalfd4. dirent := fs.NewDirent(ctx, anon.NewInode(ctx), "anon_inode:[signalfd]") - return fs.NewFile(ctx, dirent, fs.FileFlags{Read: true, Write: true}, &SignalOperations{ - target: t, - mask: mask, - }), nil + s := &SignalOperations{target: t} + s.entry.Init(s, waiter.EventMask(mask)) + s.target.SignalRegister(&s.entry) + return fs.NewFile(ctx, dirent, fs.FileFlags{Read: true, Write: true}, s), nil } // Release implements fs.FileOperations.Release. -func (s *SignalOperations) Release(context.Context) {} +func (s *SignalOperations) Release(context.Context) { + s.target.SignalUnregister(&s.entry) +} // Mask returns the signal mask. func (s *SignalOperations) Mask() linux.SignalSet { s.mu.Lock() - mask := s.mask - s.mu.Unlock() - return mask + defer s.mu.Unlock() + return linux.SignalSet(s.entry.Mask()) } // SetMask sets the signal mask. func (s *SignalOperations) SetMask(mask linux.SignalSet) { s.mu.Lock() - s.mask = mask - s.mu.Unlock() + defer s.mu.Unlock() + s.target.SignalUnregister(&s.entry) + s.entry.Init(s, waiter.EventMask(mask)) + s.target.SignalRegister(&s.entry) } // Read implements fs.FileOperations.Read. @@ -129,13 +134,16 @@ func (s *SignalOperations) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements waiter.Waitable.EventRegister. -func (s *SignalOperations) EventRegister(entry *waiter.Entry, _ waiter.EventMask) { - // Register for the signal set; ignore the passed events. - s.target.SignalRegister(entry, waiter.EventMask(s.Mask())) +func (s *SignalOperations) EventRegister(e *waiter.Entry) { + s.queue.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. -func (s *SignalOperations) EventUnregister(entry *waiter.Entry) { - // Unregister the original entry. - s.target.SignalUnregister(entry) +func (s *SignalOperations) EventUnregister(e *waiter.Entry) { + s.queue.EventUnregister(e) +} + +// NotifyEvent implements waiter.EventListener.NotifyEvent. +func (s *SignalOperations) NotifyEvent(mask waiter.EventMask) { + s.queue.Notify(waiter.EventIn) } diff --git a/pkg/sentry/kernel/task_exit.go b/pkg/sentry/kernel/task_exit.go index fa523c4db..dbd4c1068 100644 --- a/pkg/sentry/kernel/task_exit.go +++ b/pkg/sentry/kernel/task_exit.go @@ -855,8 +855,8 @@ func (t *Task) Wait(opts *WaitOptions) (*WaitResult, error) { if opts.BlockInterruptErr == nil { return t.waitOnce(opts) } - w, ch := waiter.NewChannelEntry(nil) - t.tg.eventQueue.EventRegister(&w, opts.Events) + w, ch := waiter.NewChannelEntry(opts.Events) + t.tg.eventQueue.EventRegister(&w) defer t.tg.eventQueue.EventUnregister(&w) for { wr, err := t.waitOnce(opts) diff --git a/pkg/sentry/kernel/task_signals.go b/pkg/sentry/kernel/task_signals.go index eeb3c5e69..e46c9c924 100644 --- a/pkg/sentry/kernel/task_signals.go +++ b/pkg/sentry/kernel/task_signals.go @@ -1094,9 +1094,9 @@ func (*runInterruptAfterSignalDeliveryStop) execute(t *Task) taskRunState { } // SignalRegister registers a waiter for pending signals. -func (t *Task) SignalRegister(e *waiter.Entry, mask waiter.EventMask) { +func (t *Task) SignalRegister(e *waiter.Entry) { t.tg.signalHandlers.mu.Lock() - t.signalQueue.EventRegister(e, mask) + t.signalQueue.EventRegister(e) t.tg.signalHandlers.mu.Unlock() } diff --git a/pkg/sentry/kernel/time/time.go b/pkg/sentry/kernel/time/time.go index 9f7201a7f..a92b4e23c 100644 --- a/pkg/sentry/kernel/time/time.go +++ b/pkg/sentry/kernel/time/time.go @@ -265,7 +265,7 @@ func (*NoClockEvents) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements waiter.Waitable.EventRegister. -func (*NoClockEvents) EventRegister(e *waiter.Entry, mask waiter.EventMask) { +func (*NoClockEvents) EventRegister(e *waiter.Entry) { } // EventUnregister implements waiter.Waitable.EventUnregister. @@ -473,8 +473,8 @@ func (t *Timer) init() { // If t.kicker is nil, the Timer goroutine can't be running, so we can't // race with it. t.kicker = time.NewTimer(0) - t.entry, t.events = waiter.NewChannelEntry(nil) - t.clock.EventRegister(&t.entry, timerTickEvents) + t.entry, t.events = waiter.NewChannelEntry(timerTickEvents) + t.clock.EventRegister(&t.entry) go t.runGoroutine() // S/R-SAFE: synchronized by t.mu } diff --git a/pkg/sentry/socket/hostinet/socket.go b/pkg/sentry/socket/hostinet/socket.go index a31f3ebec..c7f9ceafe 100644 --- a/pkg/sentry/socket/hostinet/socket.go +++ b/pkg/sentry/socket/hostinet/socket.go @@ -216,8 +216,8 @@ func (s *socketOpsCommon) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements waiter.Waitable.EventRegister. -func (s *socketOpsCommon) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - s.queue.EventRegister(e, mask) +func (s *socketOpsCommon) EventRegister(e *waiter.Entry) { + s.queue.EventRegister(e) _ = fdnotifier.UpdateFD(int32(s.fd)) } @@ -249,9 +249,9 @@ func (s *socketOpsCommon) Connect(t *kernel.Task, sockaddr []byte, blocking bool // level SOL-SOCKET to determine whether connect() completed successfully // (SO_ERROR is zero) or unsuccessfully (SO_ERROR is one of the usual error // codes listed here, explaining the reason for the failure)." - connect(2) - e, ch := waiter.NewChannelEntry(nil) writableMask := waiter.WritableEvents - s.EventRegister(&e, writableMask) + e, ch := waiter.NewChannelEntry(writableMask) + s.EventRegister(&e) defer s.EventUnregister(&e) if s.Readiness(writableMask)&writableMask == 0 { if err := t.Block(ch); err != nil { @@ -294,8 +294,8 @@ func (s *socketOpsCommon) Accept(t *kernel.Task, peerRequested bool, flags int, } } else { var e waiter.Entry - e, ch = waiter.NewChannelEntry(nil) - s.EventRegister(&e, waiter.ReadableEvents) + e, ch = waiter.NewChannelEntry(waiter.ReadableEvents) + s.EventRegister(&e) defer s.EventUnregister(&e) } fd, syscallErr = accept4(s.fd, peerAddrPtr, peerAddrlenPtr, unix.SOCK_NONBLOCK|unix.SOCK_CLOEXEC) @@ -546,8 +546,8 @@ func (s *socketOpsCommon) RecvMsg(t *kernel.Task, dst usermem.IOSequence, flags } } else { var e waiter.Entry - e, ch = waiter.NewChannelEntry(nil) - s.EventRegister(&e, waiter.ReadableEvents) + e, ch = waiter.NewChannelEntry(waiter.ReadableEvents) + s.EventRegister(&e) defer s.EventUnregister(&e) } n, err = copyToDst() @@ -721,8 +721,8 @@ func (s *socketOpsCommon) SendMsg(t *kernel.Task, src usermem.IOSequence, to []b } } else { var e waiter.Entry - e, ch = waiter.NewChannelEntry(nil) - s.EventRegister(&e, waiter.WritableEvents) + e, ch = waiter.NewChannelEntry(waiter.WritableEvents) + s.EventRegister(&e) defer s.EventUnregister(&e) } n, err = src.CopyInTo(t, sendmsgFromBlocks) diff --git a/pkg/sentry/socket/hostinet/socket_vfs2.go b/pkg/sentry/socket/hostinet/socket_vfs2.go index cd6e34ecc..7eeb17f71 100644 --- a/pkg/sentry/socket/hostinet/socket_vfs2.go +++ b/pkg/sentry/socket/hostinet/socket_vfs2.go @@ -89,8 +89,8 @@ func (s *socketVFS2) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements waiter.Waitable.EventRegister. -func (s *socketVFS2) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - s.socketOpsCommon.EventRegister(e, mask) +func (s *socketVFS2) EventRegister(e *waiter.Entry) { + s.socketOpsCommon.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. diff --git a/pkg/sentry/socket/netlink/socket.go b/pkg/sentry/socket/netlink/socket.go index 19c8f340d..22ceed20d 100644 --- a/pkg/sentry/socket/netlink/socket.go +++ b/pkg/sentry/socket/netlink/socket.go @@ -187,8 +187,8 @@ func (s *socketOpsCommon) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements waiter.Waitable.EventRegister. -func (s *socketOpsCommon) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - s.ep.EventRegister(e, mask) +func (s *socketOpsCommon) EventRegister(e *waiter.Entry) { + s.ep.EventRegister(e) // Writable readiness never changes, so no registration is needed. } @@ -542,8 +542,8 @@ func (s *socketOpsCommon) RecvMsg(t *kernel.Task, dst usermem.IOSequence, flags // We'll have to block. Register for notification and keep trying to // receive all the data. - e, ch := waiter.NewChannelEntry(nil) - s.EventRegister(&e, waiter.ReadableEvents) + e, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + s.EventRegister(&e) defer s.EventUnregister(&e) for { diff --git a/pkg/sentry/socket/netlink/socket_vfs2.go b/pkg/sentry/socket/netlink/socket_vfs2.go index 4d3cdea62..216b51550 100644 --- a/pkg/sentry/socket/netlink/socket_vfs2.go +++ b/pkg/sentry/socket/netlink/socket_vfs2.go @@ -96,8 +96,8 @@ func (s *SocketVFS2) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements waiter.Waitable.EventRegister. -func (s *SocketVFS2) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - s.socketOpsCommon.EventRegister(e, mask) +func (s *SocketVFS2) EventRegister(e *waiter.Entry) { + s.socketOpsCommon.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. diff --git a/pkg/sentry/socket/netstack/netstack.go b/pkg/sentry/socket/netstack/netstack.go index e38c4e5da..82b3d428f 100644 --- a/pkg/sentry/socket/netstack/netstack.go +++ b/pkg/sentry/socket/netstack/netstack.go @@ -436,8 +436,8 @@ func (s *socketOpsCommon) isPacketBased() bool { // Release implements fs.FileOperations.Release. func (s *socketOpsCommon) Release(ctx context.Context) { - e, ch := waiter.NewChannelEntry(nil) - s.EventRegister(&e, waiter.EventHUp|waiter.EventErr) + e, ch := waiter.NewChannelEntry(waiter.EventHUp | waiter.EventErr) + s.EventRegister(&e) defer s.EventUnregister(&e) s.Endpoint.Close() @@ -615,8 +615,8 @@ func (s *socketOpsCommon) Connect(t *kernel.Task, sockaddr []byte, blocking bool // Register for notification when the endpoint becomes writable, then // initiate the connection. - e, ch := waiter.NewChannelEntry(nil) - s.EventRegister(&e, waiter.WritableEvents) + e, ch := waiter.NewChannelEntry(waiter.WritableEvents) + s.EventRegister(&e) defer s.EventUnregister(&e) switch err := s.Endpoint.Connect(addr); err.(type) { @@ -712,8 +712,8 @@ func (s *socketOpsCommon) Listen(_ *kernel.Task, backlog int) *syserr.Error { // connections are ready to be accept, it will block until one becomes ready. func (s *socketOpsCommon) blockingAccept(t *kernel.Task, peerAddr *tcpip.FullAddress) (tcpip.Endpoint, *waiter.Queue, *syserr.Error) { // Register for notifications. - e, ch := waiter.NewChannelEntry(nil) - s.EventRegister(&e, waiter.ReadableEvents) + e, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + s.EventRegister(&e) defer s.EventUnregister(&e) // Try to accept the connection again; if it fails, then wait until we @@ -2858,8 +2858,8 @@ func (s *socketOpsCommon) RecvMsg(t *kernel.Task, dst usermem.IOSequence, flags // We'll have to block. Register for notifications and keep trying to // send all the data. - e, ch := waiter.NewChannelEntry(nil) - s.EventRegister(&e, waiter.ReadableEvents) + e, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + s.EventRegister(&e) defer s.EventUnregister(&e) for { @@ -2945,8 +2945,8 @@ func (s *socketOpsCommon) SendMsg(t *kernel.Task, src usermem.IOSequence, to []b if ch == nil { // We'll have to block. Register for notification and keep trying to // send all the data. - entry, ch = waiter.NewChannelEntry(nil) - s.EventRegister(&entry, waiter.WritableEvents) + entry, ch = waiter.NewChannelEntry(waiter.WritableEvents) + s.EventRegister(&entry) defer s.EventUnregister(&entry) } else { // Don't wait immediately after registration in case more data diff --git a/pkg/sentry/socket/netstack/netstack_vfs2.go b/pkg/sentry/socket/netstack/netstack_vfs2.go index ff10e159e..157ddd260 100644 --- a/pkg/sentry/socket/netstack/netstack_vfs2.go +++ b/pkg/sentry/socket/netstack/netstack_vfs2.go @@ -90,8 +90,8 @@ func (s *SocketVFS2) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements waiter.Waitable.EventRegister. -func (s *SocketVFS2) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - s.socketOpsCommon.EventRegister(e, mask) +func (s *SocketVFS2) EventRegister(e *waiter.Entry) { + s.socketOpsCommon.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. diff --git a/pkg/sentry/socket/unix/transport/unix.go b/pkg/sentry/socket/unix/transport/unix.go index 837ab4fde..ae30d02d5 100644 --- a/pkg/sentry/socket/unix/transport/unix.go +++ b/pkg/sentry/socket/unix/transport/unix.go @@ -774,8 +774,8 @@ type baseEndpoint struct { } // EventRegister implements waiter.Waitable.EventRegister. -func (e *baseEndpoint) EventRegister(we *waiter.Entry, mask waiter.EventMask) { - e.Queue.EventRegister(we, mask) +func (e *baseEndpoint) EventRegister(we *waiter.Entry) { + e.Queue.EventRegister(we) e.Lock() c := e.connected e.Unlock() diff --git a/pkg/sentry/socket/unix/unix.go b/pkg/sentry/socket/unix/unix.go index 032678032..c6e8f4aff 100644 --- a/pkg/sentry/socket/unix/unix.go +++ b/pkg/sentry/socket/unix/unix.go @@ -207,8 +207,8 @@ func (s *socketOpsCommon) Listen(t *kernel.Task, backlog int) *syserr.Error { // connections are ready to be accept, it will block until one becomes ready. func (s *SocketOperations) blockingAccept(t *kernel.Task, peerAddr *tcpip.FullAddress) (transport.Endpoint, *syserr.Error) { // Register for notifications. - e, ch := waiter.NewChannelEntry(nil) - s.EventRegister(&e, waiter.ReadableEvents) + e, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + s.EventRegister(&e) defer s.EventUnregister(&e) // Try to accept the connection; if it fails, then wait until we get a @@ -502,8 +502,8 @@ func (s *socketOpsCommon) SendMsg(t *kernel.Task, src usermem.IOSequence, to []b // We'll have to block. Register for notification and keep trying to // send all the data. - e, ch := waiter.NewChannelEntry(nil) - s.EventRegister(&e, waiter.WritableEvents) + e, ch := waiter.NewChannelEntry(waiter.WritableEvents) + s.EventRegister(&e) defer s.EventUnregister(&e) total := n @@ -544,8 +544,8 @@ func (s *socketOpsCommon) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements waiter.Waitable.EventRegister. -func (s *socketOpsCommon) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - s.ep.EventRegister(e, mask) +func (s *socketOpsCommon) EventRegister(e *waiter.Entry) { + s.ep.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. @@ -677,8 +677,8 @@ func (s *socketOpsCommon) RecvMsg(t *kernel.Task, dst usermem.IOSequence, flags // We'll have to block. Register for notification and keep trying to // send all the data. - e, ch := waiter.NewChannelEntry(nil) - s.EventRegister(&e, waiter.ReadableEvents) + e, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + s.EventRegister(&e) defer s.EventUnregister(&e) for { diff --git a/pkg/sentry/socket/unix/unix_vfs2.go b/pkg/sentry/socket/unix/unix_vfs2.go index b05233dfe..a3fb9d330 100644 --- a/pkg/sentry/socket/unix/unix_vfs2.go +++ b/pkg/sentry/socket/unix/unix_vfs2.go @@ -121,8 +121,8 @@ func (s *SocketVFS2) GetSockOpt(t *kernel.Task, level, name int, outPtr hostarch // connections are ready to be accept, it will block until one becomes ready. func (s *SocketVFS2) blockingAccept(t *kernel.Task, peerAddr *tcpip.FullAddress) (transport.Endpoint, *syserr.Error) { // Register for notifications. - e, ch := waiter.NewChannelEntry(nil) - s.socketOpsCommon.EventRegister(&e, waiter.ReadableEvents) + e, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + s.socketOpsCommon.EventRegister(&e) defer s.socketOpsCommon.EventUnregister(&e) // Try to accept the connection; if it fails, then wait until we get a @@ -315,8 +315,8 @@ func (s *SocketVFS2) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements waiter.Waitable.EventRegister. -func (s *SocketVFS2) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - s.socketOpsCommon.EventRegister(e, mask) +func (s *SocketVFS2) EventRegister(e *waiter.Entry) { + s.socketOpsCommon.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. diff --git a/pkg/sentry/syscalls/epoll.go b/pkg/sentry/syscalls/epoll.go index a69ed0746..01e5f991f 100644 --- a/pkg/sentry/syscalls/epoll.go +++ b/pkg/sentry/syscalls/epoll.go @@ -150,8 +150,8 @@ func WaitEpoll(t *kernel.Task, fd int32, max int, timeoutInNanos int64) ([]linux haveDeadline = true } - w, ch := waiter.NewChannelEntry(nil) - e.EventRegister(&w, waiter.ReadableEvents) + w, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + e.EventRegister(&w) defer e.EventUnregister(&w) // Try to read the events again until we succeed, timeout or get diff --git a/pkg/sentry/syscalls/linux/sys_poll.go b/pkg/sentry/syscalls/linux/sys_poll.go index 8cd318c8f..0bac94478 100644 --- a/pkg/sentry/syscalls/linux/sys_poll.go +++ b/pkg/sentry/syscalls/linux/sys_poll.go @@ -73,8 +73,8 @@ func initReadiness(t *kernel.Task, pfd *linux.PollFD, state *pollState, ch chan defer file.DecRef(t) } else { state.file = file - state.waiter, _ = waiter.NewChannelEntry(ch) - file.EventRegister(&state.waiter, waiter.EventMaskFromLinux(uint32(pfd.Events))) + state.waiter.Init(waiter.ChannelNotifier(ch), waiter.EventMaskFromLinux(uint32(pfd.Events))) + file.EventRegister(&state.waiter) } r := file.Readiness(waiter.EventMaskFromLinux(uint32(pfd.Events))) diff --git a/pkg/sentry/syscalls/linux/sys_read.go b/pkg/sentry/syscalls/linux/sys_read.go index 18ea23913..1b99aa208 100644 --- a/pkg/sentry/syscalls/linux/sys_read.go +++ b/pkg/sentry/syscalls/linux/sys_read.go @@ -313,8 +313,8 @@ func readv(t *kernel.Task, f *fs.File, dst usermem.IOSequence) (int64, error) { } // Register for notifications. - w, ch := waiter.NewChannelEntry(nil) - f.EventRegister(&w, EventMaskRead) + w, ch := waiter.NewChannelEntry(EventMaskRead) + f.EventRegister(&w) total := n for { @@ -359,8 +359,8 @@ func preadv(t *kernel.Task, f *fs.File, dst usermem.IOSequence, offset int64) (i } // Register for notifications. - w, ch := waiter.NewChannelEntry(nil) - f.EventRegister(&w, EventMaskRead) + w, ch := waiter.NewChannelEntry(EventMaskRead) + f.EventRegister(&w) total := n for { diff --git a/pkg/sentry/syscalls/linux/sys_splice.go b/pkg/sentry/syscalls/linux/sys_splice.go index 8c8847efa..eac50e69e 100644 --- a/pkg/sentry/syscalls/linux/sys_splice.go +++ b/pkg/sentry/syscalls/linux/sys_splice.go @@ -57,10 +57,10 @@ func doSplice(t *kernel.Task, outFile, inFile *fs.File, opts fs.SpliceOpts, nonB // these cases in turn before returning to the splice operation. if inFile.Readiness(EventMaskRead) == 0 { if inCh == nil { - inCh = make(chan struct{}, 1) - inW, _ := waiter.NewChannelEntry(inCh) - inFile.EventRegister(&inW, EventMaskRead) - defer inFile.EventUnregister(&inW) + var e waiter.Entry + e, inCh = waiter.NewChannelEntry(EventMaskRead) + inFile.EventRegister(&e) + defer inFile.EventUnregister(&e) // Need to refresh readiness. continue } @@ -73,10 +73,10 @@ func doSplice(t *kernel.Task, outFile, inFile *fs.File, opts fs.SpliceOpts, nonB // can be "ready" but will reject writes of certain sizes with // EWOULDBLOCK. if outCh == nil { - outCh = make(chan struct{}, 1) - outW, _ := waiter.NewChannelEntry(outCh) - outFile.EventRegister(&outW, EventMaskWrite) - defer outFile.EventUnregister(&outW) + var e waiter.Entry + e, outCh = waiter.NewChannelEntry(EventMaskWrite) + outFile.EventRegister(&e) + defer outFile.EventUnregister(&e) // We might be ready to write now. Try again before // blocking. continue diff --git a/pkg/sentry/syscalls/linux/sys_write.go b/pkg/sentry/syscalls/linux/sys_write.go index 4a4ef5046..5bd167689 100644 --- a/pkg/sentry/syscalls/linux/sys_write.go +++ b/pkg/sentry/syscalls/linux/sys_write.go @@ -283,8 +283,8 @@ func writev(t *kernel.Task, f *fs.File, src usermem.IOSequence) (int64, error) { } // Register for notifications. - w, ch := waiter.NewChannelEntry(nil) - f.EventRegister(&w, EventMaskWrite) + w, ch := waiter.NewChannelEntry(EventMaskWrite) + f.EventRegister(&w) total := n for { @@ -329,8 +329,8 @@ func pwritev(t *kernel.Task, f *fs.File, src usermem.IOSequence, offset int64) ( } // Register for notifications. - w, ch := waiter.NewChannelEntry(nil) - f.EventRegister(&w, EventMaskWrite) + w, ch := waiter.NewChannelEntry(EventMaskWrite) + f.EventRegister(&w) total := n for { diff --git a/pkg/sentry/syscalls/linux/vfs2/epoll.go b/pkg/sentry/syscalls/linux/vfs2/epoll.go index 84010db77..9a78a543b 100644 --- a/pkg/sentry/syscalls/linux/vfs2/epoll.go +++ b/pkg/sentry/syscalls/linux/vfs2/epoll.go @@ -163,8 +163,8 @@ func waitEpoll(t *kernel.Task, epfd int32, eventsAddr hostarch.Addr, maxEvents i // expires, or an interrupt arrives. if ch == nil { var w waiter.Entry - w, ch = waiter.NewChannelEntry(nil) - epfile.EventRegister(&w, waiter.ReadableEvents) + w, ch = waiter.NewChannelEntry(waiter.ReadableEvents) + epfile.EventRegister(&w) defer epfile.EventUnregister(&w) } else { // Set up the timer if a timeout was specified. diff --git a/pkg/sentry/syscalls/linux/vfs2/poll.go b/pkg/sentry/syscalls/linux/vfs2/poll.go index 1dcbe0a4d..30559c23c 100644 --- a/pkg/sentry/syscalls/linux/vfs2/poll.go +++ b/pkg/sentry/syscalls/linux/vfs2/poll.go @@ -76,8 +76,8 @@ func initReadiness(t *kernel.Task, pfd *linux.PollFD, state *pollState, ch chan defer file.DecRef(t) } else { state.file = file - state.waiter, _ = waiter.NewChannelEntry(ch) - file.EventRegister(&state.waiter, waiter.EventMaskFromLinux(uint32(pfd.Events))) + state.waiter.Init(waiter.ChannelNotifier(ch), waiter.EventMaskFromLinux(uint32(pfd.Events))) + file.EventRegister(&state.waiter) } r := file.Readiness(waiter.EventMaskFromLinux(uint32(pfd.Events))) diff --git a/pkg/sentry/syscalls/linux/vfs2/read_write.go b/pkg/sentry/syscalls/linux/vfs2/read_write.go index 4e7dc5080..06442b93d 100644 --- a/pkg/sentry/syscalls/linux/vfs2/read_write.go +++ b/pkg/sentry/syscalls/linux/vfs2/read_write.go @@ -102,8 +102,8 @@ func read(t *kernel.Task, file *vfs.FileDescription, dst usermem.IOSequence, opt } // Register for notifications. - w, ch := waiter.NewChannelEntry(nil) - file.EventRegister(&w, eventMaskRead) + w, ch := waiter.NewChannelEntry(eventMaskRead) + file.EventRegister(&w) total := n for { @@ -257,8 +257,8 @@ func pread(t *kernel.Task, file *vfs.FileDescription, dst usermem.IOSequence, of } // Register for notifications. - w, ch := waiter.NewChannelEntry(nil) - file.EventRegister(&w, eventMaskRead) + w, ch := waiter.NewChannelEntry(eventMaskRead) + file.EventRegister(&w) total := n for { @@ -353,8 +353,8 @@ func write(t *kernel.Task, file *vfs.FileDescription, src usermem.IOSequence, op } // Register for notifications. - w, ch := waiter.NewChannelEntry(nil) - file.EventRegister(&w, eventMaskWrite) + w, ch := waiter.NewChannelEntry(eventMaskWrite) + file.EventRegister(&w) total := n for { @@ -507,8 +507,8 @@ func pwrite(t *kernel.Task, file *vfs.FileDescription, src usermem.IOSequence, o } // Register for notifications. - w, ch := waiter.NewChannelEntry(nil) - file.EventRegister(&w, eventMaskWrite) + w, ch := waiter.NewChannelEntry(eventMaskWrite) + file.EventRegister(&w) total := n for { diff --git a/pkg/sentry/syscalls/linux/vfs2/splice.go b/pkg/sentry/syscalls/linux/vfs2/splice.go index 0205f09e0..5ba1207be 100644 --- a/pkg/sentry/syscalls/linux/vfs2/splice.go +++ b/pkg/sentry/syscalls/linux/vfs2/splice.go @@ -470,8 +470,8 @@ type dualWaiter struct { func (dw *dualWaiter) waitForBoth(t *kernel.Task) error { if dw.inFile.Readiness(eventMaskRead)&eventMaskRead == 0 { if dw.inCh == nil { - dw.inW, dw.inCh = waiter.NewChannelEntry(nil) - dw.inFile.EventRegister(&dw.inW, eventMaskRead) + dw.inW, dw.inCh = waiter.NewChannelEntry(eventMaskRead) + dw.inFile.EventRegister(&dw.inW) // We might be ready now. Try again before blocking. return nil } @@ -489,8 +489,8 @@ func (dw *dualWaiter) waitForOut(t *kernel.Task) error { // can be "ready" but will reject writes of certain sizes with // EWOULDBLOCK. See b/172075629, b/170743336. if dw.outCh == nil { - dw.outW, dw.outCh = waiter.NewChannelEntry(nil) - dw.outFile.EventRegister(&dw.outW, eventMaskWrite) + dw.outW, dw.outCh = waiter.NewChannelEntry(eventMaskWrite) + dw.outFile.EventRegister(&dw.outW) // We might be ready to write now. Try again before blocking. return nil } diff --git a/pkg/sentry/vfs/epoll.go b/pkg/sentry/vfs/epoll.go index fefd0fc9c..95132e1f8 100644 --- a/pkg/sentry/vfs/epoll.go +++ b/pkg/sentry/vfs/epoll.go @@ -151,8 +151,8 @@ func (ep *EpollInstance) Readiness(mask waiter.EventMask) waiter.EventMask { } // EventRegister implements waiter.Waitable.EventRegister. -func (ep *EpollInstance) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - ep.q.EventRegister(e, mask) +func (ep *EpollInstance) EventRegister(e *waiter.Entry) { + ep.q.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. @@ -202,14 +202,14 @@ func (ep *EpollInstance) AddInterest(file *FileDescription, num int32, event lin mask: mask, userData: event.Data, } - epi.waiter.Callback = epi ep.interest[key] = epi wmask := waiter.EventMaskFromLinux(mask) - file.EventRegister(&epi.waiter, wmask) + epi.waiter.Init(epi, wmask) + file.EventRegister(&epi.waiter) // Check if the file is already ready. if m := file.Readiness(wmask) & wmask; m != 0 { - epi.Callback(nil, m) + epi.NotifyEvent(m) } // Add epi to file.epolls so that it is removed when the last @@ -275,11 +275,12 @@ func (ep *EpollInstance) ModifyInterest(file *FileDescription, num int32, event // Re-register with the new mask. file.EventUnregister(&epi.waiter) wmask := waiter.EventMaskFromLinux(mask) - file.EventRegister(&epi.waiter, wmask) + epi.waiter.Init(epi, wmask) + file.EventRegister(&epi.waiter) // Check if the file is already ready with the new mask. if m := file.Readiness(wmask) & wmask; m != 0 { - epi.Callback(nil, m) + epi.NotifyEvent(m) } return nil @@ -314,8 +315,8 @@ func (ep *EpollInstance) DeleteInterest(file *FileDescription, num int32) error return nil } -// Callback implements waiter.EntryCallback.Callback. -func (epi *epollInterest) Callback(*waiter.Entry, waiter.EventMask) { +// NotifyEvent implements waiter.EventListener.NotifyEvent. +func (epi *epollInterest) NotifyEvent(waiter.EventMask) { newReady := false epi.epoll.mu.Lock() if !epi.ready { diff --git a/pkg/sentry/vfs/file_description.go b/pkg/sentry/vfs/file_description.go index ca3303dec..2a608feb8 100644 --- a/pkg/sentry/vfs/file_description.go +++ b/pkg/sentry/vfs/file_description.go @@ -586,8 +586,8 @@ func (fd *FileDescription) Readiness(mask waiter.EventMask) waiter.EventMask { // EventRegister implements waiter.Waitable.EventRegister. // // It registers e for I/O readiness events in mask. -func (fd *FileDescription) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - fd.impl.EventRegister(e, mask) +func (fd *FileDescription) EventRegister(e *waiter.Entry) { + fd.impl.EventRegister(e) } // EventUnregister implements waiter.Waitable.EventUnregister. diff --git a/pkg/sentry/vfs/file_description_impl_util.go b/pkg/sentry/vfs/file_description_impl_util.go index 452f5f1f9..ddb2b1658 100644 --- a/pkg/sentry/vfs/file_description_impl_util.go +++ b/pkg/sentry/vfs/file_description_impl_util.go @@ -78,7 +78,7 @@ func (FileDescriptionDefaultImpl) Readiness(mask waiter.EventMask) waiter.EventM // EventRegister implements waiter.Waitable.EventRegister analogously to // file_operations::poll == NULL in Linux. -func (FileDescriptionDefaultImpl) EventRegister(e *waiter.Entry, mask waiter.EventMask) { +func (FileDescriptionDefaultImpl) EventRegister(e *waiter.Entry) { } // EventUnregister implements waiter.Waitable.EventUnregister analogously to diff --git a/pkg/sentry/vfs/inotify.go b/pkg/sentry/vfs/inotify.go index 17d94b341..42d9f819e 100644 --- a/pkg/sentry/vfs/inotify.go +++ b/pkg/sentry/vfs/inotify.go @@ -157,8 +157,8 @@ func (i *Inotify) Allocate(ctx context.Context, mode, offset, length uint64) err } // EventRegister implements waiter.Waitable. -func (i *Inotify) EventRegister(e *waiter.Entry, mask waiter.EventMask) { - i.queue.EventRegister(e, mask) +func (i *Inotify) EventRegister(e *waiter.Entry) { + i.queue.EventRegister(e) } // EventUnregister implements waiter.Waitable. diff --git a/pkg/sentry/vfs/save_restore.go b/pkg/sentry/vfs/save_restore.go index 8a6ced365..d0dcad760 100644 --- a/pkg/sentry/vfs/save_restore.go +++ b/pkg/sentry/vfs/save_restore.go @@ -131,7 +131,7 @@ func (mnt *Mount) afterLoad() { func (epi *epollInterest) afterLoad() { // Mark all epollInterests as ready after restore so that the next call to // EpollInstance.ReadEvents() rechecks their readiness. - epi.Callback(nil, waiter.EventMaskFromLinux(epi.mask)) + epi.waiter.NotifyEvent(waiter.EventMaskFromLinux(epi.mask)) } // beforeSave is called by stateify. diff --git a/pkg/syncevent/broadcaster_test.go b/pkg/syncevent/broadcaster_test.go index e88779e23..cec3af601 100644 --- a/pkg/syncevent/broadcaster_test.go +++ b/pkg/syncevent/broadcaster_test.go @@ -136,11 +136,11 @@ func BenchmarkMapSubscribeUnsubscribe(b *testing.B) { func BenchmarkQueueSubscribeUnsubscribe(b *testing.B) { var q waiter.Queue - e, _ := waiter.NewChannelEntry(nil) + e, _ := waiter.NewChannelEntry(1) b.ResetTimer() for i := 0; i < b.N; i++ { - q.EventRegister(&e, 1) + q.EventRegister(&e) q.EventUnregister(&e) } } @@ -207,7 +207,7 @@ func BenchmarkQueueSubscribeUnsubscribeBatch(b *testing.B) { var q waiter.Queue es := make([]waiter.Entry, numBatchReceivers) for i := range es { - es[i], _ = waiter.NewChannelEntry(nil) + es[i], _ = waiter.NewChannelEntry(1) } // Generate a random order for unsubscriptions. @@ -216,7 +216,7 @@ func BenchmarkQueueSubscribeUnsubscribeBatch(b *testing.B) { b.ResetTimer() for i := 0; i < b.N/numBatchReceivers; i++ { for j := 0; j < numBatchReceivers; j++ { - q.EventRegister(&es[j], 1) + q.EventRegister(&es[j]) } for j := 0; j < numBatchReceivers; j++ { q.EventUnregister(&es[unsub[j]]) @@ -279,8 +279,8 @@ func BenchmarkQueueBroadcastRedundant(b *testing.B) { b.Run(fmt.Sprintf("%d", n), func(b *testing.B) { var q waiter.Queue for i := 0; i < n; i++ { - e, _ := waiter.NewChannelEntry(nil) - q.EventRegister(&e, 1) + e, _ := waiter.NewChannelEntry(1) + q.EventRegister(&e) } q.Notify(1) @@ -355,8 +355,8 @@ func BenchmarkQueueBroadcastAck(b *testing.B) { var q waiter.Queue chs := make([]chan struct{}, n) for i := range chs { - e, ch := waiter.NewChannelEntry(nil) - q.EventRegister(&e, 1) + e, ch := waiter.NewChannelEntry(1) + q.EventRegister(&e) chs[i] = ch } diff --git a/pkg/tcpip/adapters/gonet/gonet.go b/pkg/tcpip/adapters/gonet/gonet.go index 1f2bcaf65..6bfdb9e41 100644 --- a/pkg/tcpip/adapters/gonet/gonet.go +++ b/pkg/tcpip/adapters/gonet/gonet.go @@ -251,8 +251,8 @@ func (l *TCPListener) Accept() (net.Conn, error) { if _, ok := err.(*tcpip.ErrWouldBlock); ok { // Create wait queue entry that notifies a channel. - waitEntry, notifyCh := waiter.NewChannelEntry(nil) - l.wq.EventRegister(&waitEntry, waiter.ReadableEvents) + waitEntry, notifyCh := waiter.NewChannelEntry(waiter.ReadableEvents) + l.wq.EventRegister(&waitEntry) defer l.wq.EventUnregister(&waitEntry) for { @@ -301,8 +301,8 @@ func commonRead(b []byte, ep tcpip.Endpoint, wq *waiter.Queue, deadline <-chan s if _, ok := err.(*tcpip.ErrWouldBlock); ok { // Create wait queue entry that notifies a channel. - waitEntry, notifyCh := waiter.NewChannelEntry(nil) - wq.EventRegister(&waitEntry, waiter.ReadableEvents) + waitEntry, notifyCh := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&waitEntry) defer wq.EventUnregister(&waitEntry) for { res, err = ep.Read(&w, opts) @@ -381,9 +381,8 @@ func (c *TCPConn) Write(b []byte) (int, error) { case nil: case *tcpip.ErrWouldBlock: if ch == nil { - entry, ch = waiter.NewChannelEntry(nil) - - c.wq.EventRegister(&entry, waiter.WritableEvents) + entry, ch = waiter.NewChannelEntry(waiter.WritableEvents) + c.wq.EventRegister(&entry) defer c.wq.EventUnregister(&entry) } else { // Don't wait immediately after registration in case more data @@ -485,8 +484,8 @@ func DialTCPWithBind(ctx context.Context, s *stack.Stack, localAddr, remoteAddr // Create wait queue entry that notifies a channel. // // We do this unconditionally as Connect will always return an error. - waitEntry, notifyCh := waiter.NewChannelEntry(nil) - wq.EventRegister(&waitEntry, waiter.WritableEvents) + waitEntry, notifyCh := waiter.NewChannelEntry(waiter.WritableEvents) + wq.EventRegister(&waitEntry) defer wq.EventUnregister(&waitEntry) select { @@ -665,8 +664,8 @@ func (c *UDPConn) WriteTo(b []byte, addr net.Addr) (int, error) { n, err := c.ep.Write(&r, writeOptions) if _, ok := err.(*tcpip.ErrWouldBlock); ok { // Create wait queue entry that notifies a channel. - waitEntry, notifyCh := waiter.NewChannelEntry(nil) - c.wq.EventRegister(&waitEntry, waiter.WritableEvents) + waitEntry, notifyCh := waiter.NewChannelEntry(waiter.WritableEvents) + c.wq.EventRegister(&waitEntry) defer c.wq.EventUnregister(&waitEntry) for { select { diff --git a/pkg/tcpip/adapters/gonet/gonet_test.go b/pkg/tcpip/adapters/gonet/gonet_test.go index dcc9fff17..d1529e869 100644 --- a/pkg/tcpip/adapters/gonet/gonet_test.go +++ b/pkg/tcpip/adapters/gonet/gonet_test.go @@ -101,8 +101,8 @@ func connect(s *stack.Stack, addr tcpip.FullAddress) (*testConnection, tcpip.Err return nil, err } - entry, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&entry, waiter.WritableEvents) + entry, ch := waiter.NewChannelEntry(waiter.WritableEvents) + wq.EventRegister(&entry) err = ep.Connect(addr) if _, ok := err.(*tcpip.ErrConnectStarted); ok { @@ -114,7 +114,8 @@ func connect(s *stack.Stack, addr tcpip.FullAddress) (*testConnection, tcpip.Err } wq.EventUnregister(&entry) - wq.EventRegister(&entry, waiter.ReadableEvents) + entry, ch = waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&entry) return &testConnection{wq, &entry, ch, ep}, nil } diff --git a/pkg/tcpip/network/ip_test.go b/pkg/tcpip/network/ip_test.go index 87f650661..7e8cb262d 100644 --- a/pkg/tcpip/network/ip_test.go +++ b/pkg/tcpip/network/ip_test.go @@ -2104,8 +2104,8 @@ func TestSetNICIDBeforeDeliveringToRawEndpoint(t *testing.T) { }) var wq waiter.Queue - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) ep, err := s.NewRawEndpoint(udp.ProtocolNumber, test.proto, &wq, true /* associated */) if err != nil { t.Fatalf("NewEndpoint(%d, %d, _): %s", udp.ProtocolNumber, test.proto, err) diff --git a/pkg/tcpip/network/ipv4/ipv4_test.go b/pkg/tcpip/network/ipv4/ipv4_test.go index ef91245d7..50ef0d100 100644 --- a/pkg/tcpip/network/ipv4/ipv4_test.go +++ b/pkg/tcpip/network/ipv4/ipv4_test.go @@ -2724,8 +2724,8 @@ func TestReceiveFragments(t *testing.T) { } wq := waiter.Queue{} - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) defer close(ch) ep, err := s.NewEndpoint(udp.ProtocolNumber, header.IPv4ProtocolNumber, &wq) diff --git a/pkg/tcpip/network/ipv6/ipv6_test.go b/pkg/tcpip/network/ipv6/ipv6_test.go index 9e08b5318..e8f157dab 100644 --- a/pkg/tcpip/network/ipv6/ipv6_test.go +++ b/pkg/tcpip/network/ipv6/ipv6_test.go @@ -101,8 +101,8 @@ func testReceiveUDP(t *testing.T, s *stack.Stack, e *channel.Endpoint, src, dst t.Helper() wq := waiter.Queue{} - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) defer close(ch) @@ -927,8 +927,8 @@ func TestReceiveIPv6ExtHdrs(t *testing.T) { }) wq := waiter.Queue{} - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.WritableEvents) + we, ch := waiter.NewChannelEntry(waiter.WritableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) defer close(ch) ep, err := s.NewEndpoint(udp.ProtocolNumber, ProtocolNumber, &wq) @@ -2017,8 +2017,8 @@ func TestReceiveIPv6Fragments(t *testing.T) { } wq := waiter.Queue{} - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) defer close(ch) ep, err := s.NewEndpoint(udp.ProtocolNumber, ProtocolNumber, &wq) diff --git a/pkg/tcpip/sample/tun_tcp_connect/main.go b/pkg/tcpip/sample/tun_tcp_connect/main.go index 05b879543..9d5d5f576 100644 --- a/pkg/tcpip/sample/tun_tcp_connect/main.go +++ b/pkg/tcpip/sample/tun_tcp_connect/main.go @@ -177,8 +177,8 @@ func main() { } // Issue connect request and wait for it to complete. - waitEntry, notifyCh := waiter.NewChannelEntry(nil) - wq.EventRegister(&waitEntry, waiter.WritableEvents) + waitEntry, notifyCh := waiter.NewChannelEntry(waiter.WritableEvents) + wq.EventRegister(&waitEntry) terr := ep.Connect(remote) if _, ok := terr.(*tcpip.ErrConnectStarted); ok { fmt.Println("Connect is pending...") @@ -199,7 +199,8 @@ func main() { // Read data and write to standard output until the peer closes the // connection from its side. - wq.EventRegister(&waitEntry, waiter.ReadableEvents) + waitEntry, notifyCh = waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&waitEntry) for { _, err := ep.Read(os.Stdout, tcpip.ReadOptions{}) if err != nil { diff --git a/pkg/tcpip/sample/tun_tcp_echo/main.go b/pkg/tcpip/sample/tun_tcp_echo/main.go index a72afadda..7cc69b918 100644 --- a/pkg/tcpip/sample/tun_tcp_echo/main.go +++ b/pkg/tcpip/sample/tun_tcp_echo/main.go @@ -78,9 +78,8 @@ func echo(wq *waiter.Queue, ep tcpip.Endpoint) { defer ep.Close() // Create wait queue entry that notifies a channel. - waitEntry, notifyCh := waiter.NewChannelEntry(nil) - - wq.EventRegister(&waitEntry, waiter.ReadableEvents) + waitEntry, notifyCh := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&waitEntry) defer wq.EventUnregister(&waitEntry) w := endpointWriter{ @@ -215,8 +214,8 @@ func main() { } // Wait for connections to appear. - waitEntry, notifyCh := waiter.NewChannelEntry(nil) - wq.EventRegister(&waitEntry, waiter.ReadableEvents) + waitEntry, notifyCh := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&waitEntry) defer wq.EventUnregister(&waitEntry) for { diff --git a/pkg/tcpip/stack/ndp_test.go b/pkg/tcpip/stack/ndp_test.go index 40b33b6b5..41f4309ad 100644 --- a/pkg/tcpip/stack/ndp_test.go +++ b/pkg/tcpip/stack/ndp_test.go @@ -2964,8 +2964,8 @@ func addrForNewConnectionTo(t *testing.T, s *stack.Stack, addr tcpip.FullAddress t.Helper() wq := waiter.Queue{} - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) defer close(ch) ep, err := s.NewEndpoint(header.UDPProtocolNumber, header.IPv6ProtocolNumber, &wq) @@ -2998,8 +2998,8 @@ func addrForNewConnectionWithAddr(t *testing.T, s *stack.Stack, addr tcpip.FullA t.Helper() wq := waiter.Queue{} - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) defer close(ch) ep, err := s.NewEndpoint(header.UDPProtocolNumber, header.IPv6ProtocolNumber, &wq) @@ -3332,8 +3332,8 @@ func TestAutoGenAddrJobDeprecation(t *testing.T) { t.Fatal(err) } wq := waiter.Queue{} - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) defer close(ch) ep, err := s.NewEndpoint(header.UDPProtocolNumber, header.IPv6ProtocolNumber, &wq) diff --git a/pkg/tcpip/stack/transport_demuxer_test.go b/pkg/tcpip/stack/transport_demuxer_test.go index cd3a8c25a..bf31fc790 100644 --- a/pkg/tcpip/stack/transport_demuxer_test.go +++ b/pkg/tcpip/stack/transport_demuxer_test.go @@ -353,8 +353,8 @@ func TestBindToDeviceDistribution(t *testing.T) { for i, endpoint := range test.endpoints { // Try to receive the data. wq := waiter.Queue{} - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) t.Cleanup(func() { wq.EventUnregister(&we) close(ch) diff --git a/pkg/tcpip/tests/integration/forward_test.go b/pkg/tcpip/tests/integration/forward_test.go index 6e1d4720d..0bba30fff 100644 --- a/pkg/tcpip/tests/integration/forward_test.go +++ b/pkg/tcpip/tests/integration/forward_test.go @@ -84,8 +84,8 @@ func TestForwarding(t *testing.T) { newEP := func(t *testing.T, s *stack.Stack, transProto tcpip.TransportProtocolNumber, netProto tcpip.NetworkProtocolNumber) (tcpip.Endpoint, chan struct{}) { t.Helper() var wq waiter.Queue - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) ep, err := s.NewEndpoint(transProto, netProto, &wq) if err != nil { t.Fatalf("s.NewEndpoint(%d, %d, _): %s", transProto, netProto, err) @@ -220,8 +220,8 @@ func TestForwarding(t *testing.T) { t.Errorf("accepted address mismatch (-want +got):\n%s", diff) } - we, newCH := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, newCH := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) return newEP, newCH } }, diff --git a/pkg/tcpip/tests/integration/iptables_test.go b/pkg/tcpip/tests/integration/iptables_test.go index 51e7a130f..c6e8c08ad 100644 --- a/pkg/tcpip/tests/integration/iptables_test.go +++ b/pkg/tcpip/tests/integration/iptables_test.go @@ -1182,8 +1182,8 @@ func TestNAT(t *testing.T) { newEP := func(t *testing.T, s *stack.Stack, transProto tcpip.TransportProtocolNumber, netProto tcpip.NetworkProtocolNumber) (tcpip.Endpoint, chan struct{}) { t.Helper() var wq waiter.Queue - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) t.Cleanup(func() { wq.EventUnregister(&we) }) @@ -1647,8 +1647,8 @@ func TestNAT(t *testing.T) { t.Errorf("accepted address mismatch (-want +got):\n%s", diff) } - we, newCH := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, newCH := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) return newEP, newCH } }, diff --git a/pkg/tcpip/tests/integration/link_resolution_test.go b/pkg/tcpip/tests/integration/link_resolution_test.go index 95ddd8ec3..8b0968f20 100644 --- a/pkg/tcpip/tests/integration/link_resolution_test.go +++ b/pkg/tcpip/tests/integration/link_resolution_test.go @@ -154,8 +154,8 @@ func TestPing(t *testing.T) { host1Stack, _ := setupStack(t, stackOpts, host1NICID, host2NICID) var wq waiter.Queue - we, waiterCH := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, waiterCH := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) ep, err := host1Stack.NewEndpoint(test.transProto, test.netProto, &wq) if err != nil { t.Fatalf("host1Stack.NewEndpoint(%d, %d, _): %s", test.transProto, test.netProto, err) @@ -313,8 +313,8 @@ func TestTCPLinkResolutionFailure(t *testing.T) { } var clientWQ waiter.Queue - we, ch := waiter.NewChannelEntry(nil) - clientWQ.EventRegister(&we, waiter.WritableEvents|waiter.EventErr) + we, ch := waiter.NewChannelEntry(waiter.WritableEvents | waiter.EventErr) + clientWQ.EventRegister(&we) clientEP, err := host1Stack.NewEndpoint(tcp.ProtocolNumber, test.netProto, &clientWQ) if err != nil { t.Fatalf("host1Stack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, test.netProto, err) @@ -920,8 +920,8 @@ func TestWritePacketsLinkResolution(t *testing.T) { host1Stack, host2Stack := setupStack(t, stackOpts, host1NICID, host2NICID) var serverWQ waiter.Queue - serverWE, serverCH := waiter.NewChannelEntry(nil) - serverWQ.EventRegister(&serverWE, waiter.ReadableEvents) + serverWE, serverCH := waiter.NewChannelEntry(waiter.ReadableEvents) + serverWQ.EventRegister(&serverWE) serverEP, err := host2Stack.NewEndpoint(udp.ProtocolNumber, test.netProto, &serverWQ) if err != nil { t.Fatalf("host2Stack.NewEndpoint(%d, %d, _): %s", udp.ProtocolNumber, test.netProto, err) @@ -1099,8 +1099,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { neighborAddr: utils.RouterNIC1IPv4Addr.AddressWithPrefix.Address, getEndpoints: func(t *testing.T, host1Stack, _, host2Stack *stack.Stack) (tcpip.Endpoint, <-chan struct{}, tcpip.Endpoint, <-chan struct{}) { var listenerWQ waiter.Queue - listenerWE, listenerCH := waiter.NewChannelEntry(nil) - listenerWQ.EventRegister(&listenerWE, waiter.EventIn) + listenerWE, listenerCH := waiter.NewChannelEntry(waiter.EventIn) + listenerWQ.EventRegister(&listenerWE) listenerEP, err := host2Stack.NewEndpoint(tcp.ProtocolNumber, ipv4.ProtocolNumber, &listenerWQ) if err != nil { t.Fatalf("host2Stack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, ipv4.ProtocolNumber, err) @@ -1108,8 +1108,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { t.Cleanup(listenerEP.Close) var clientWQ waiter.Queue - clientWE, clientCH := waiter.NewChannelEntry(nil) - clientWQ.EventRegister(&clientWE, waiter.ReadableEvents|waiter.WritableEvents) + clientWE, clientCH := waiter.NewChannelEntry(waiter.ReadableEvents | waiter.WritableEvents) + clientWQ.EventRegister(&clientWE) clientEP, err := host1Stack.NewEndpoint(tcp.ProtocolNumber, ipv4.ProtocolNumber, &clientWQ) if err != nil { t.Fatalf("host1Stack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, ipv4.ProtocolNumber, err) @@ -1125,8 +1125,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { neighborAddr: utils.RouterNIC1IPv6Addr.AddressWithPrefix.Address, getEndpoints: func(t *testing.T, host1Stack, _, host2Stack *stack.Stack) (tcpip.Endpoint, <-chan struct{}, tcpip.Endpoint, <-chan struct{}) { var listenerWQ waiter.Queue - listenerWE, listenerCH := waiter.NewChannelEntry(nil) - listenerWQ.EventRegister(&listenerWE, waiter.EventIn) + listenerWE, listenerCH := waiter.NewChannelEntry(waiter.EventIn) + listenerWQ.EventRegister(&listenerWE) listenerEP, err := host2Stack.NewEndpoint(tcp.ProtocolNumber, ipv6.ProtocolNumber, &listenerWQ) if err != nil { t.Fatalf("host2Stack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, ipv6.ProtocolNumber, err) @@ -1134,8 +1134,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { t.Cleanup(listenerEP.Close) var clientWQ waiter.Queue - clientWE, clientCH := waiter.NewChannelEntry(nil) - clientWQ.EventRegister(&clientWE, waiter.ReadableEvents|waiter.WritableEvents) + clientWE, clientCH := waiter.NewChannelEntry(waiter.ReadableEvents | waiter.WritableEvents) + clientWQ.EventRegister(&clientWE) clientEP, err := host1Stack.NewEndpoint(tcp.ProtocolNumber, ipv6.ProtocolNumber, &clientWQ) if err != nil { t.Fatalf("host1Stack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, ipv6.ProtocolNumber, err) @@ -1151,8 +1151,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { neighborAddr: utils.RouterNIC1IPv4Addr.AddressWithPrefix.Address, getEndpoints: func(t *testing.T, host1Stack, routerStack, _ *stack.Stack) (tcpip.Endpoint, <-chan struct{}, tcpip.Endpoint, <-chan struct{}) { var listenerWQ waiter.Queue - listenerWE, listenerCH := waiter.NewChannelEntry(nil) - listenerWQ.EventRegister(&listenerWE, waiter.EventIn) + listenerWE, listenerCH := waiter.NewChannelEntry(waiter.EventIn) + listenerWQ.EventRegister(&listenerWE) listenerEP, err := routerStack.NewEndpoint(tcp.ProtocolNumber, ipv4.ProtocolNumber, &listenerWQ) if err != nil { t.Fatalf("routerStack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, ipv4.ProtocolNumber, err) @@ -1160,8 +1160,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { t.Cleanup(listenerEP.Close) var clientWQ waiter.Queue - clientWE, clientCH := waiter.NewChannelEntry(nil) - clientWQ.EventRegister(&clientWE, waiter.ReadableEvents|waiter.WritableEvents) + clientWE, clientCH := waiter.NewChannelEntry(waiter.ReadableEvents | waiter.WritableEvents) + clientWQ.EventRegister(&clientWE) clientEP, err := host1Stack.NewEndpoint(tcp.ProtocolNumber, ipv4.ProtocolNumber, &clientWQ) if err != nil { t.Fatalf("host1Stack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, ipv4.ProtocolNumber, err) @@ -1177,8 +1177,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { neighborAddr: utils.RouterNIC1IPv6Addr.AddressWithPrefix.Address, getEndpoints: func(t *testing.T, host1Stack, routerStack, _ *stack.Stack) (tcpip.Endpoint, <-chan struct{}, tcpip.Endpoint, <-chan struct{}) { var listenerWQ waiter.Queue - listenerWE, listenerCH := waiter.NewChannelEntry(nil) - listenerWQ.EventRegister(&listenerWE, waiter.EventIn) + listenerWE, listenerCH := waiter.NewChannelEntry(waiter.EventIn) + listenerWQ.EventRegister(&listenerWE) listenerEP, err := routerStack.NewEndpoint(tcp.ProtocolNumber, ipv6.ProtocolNumber, &listenerWQ) if err != nil { t.Fatalf("routerStack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, ipv6.ProtocolNumber, err) @@ -1186,8 +1186,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { t.Cleanup(listenerEP.Close) var clientWQ waiter.Queue - clientWE, clientCH := waiter.NewChannelEntry(nil) - clientWQ.EventRegister(&clientWE, waiter.ReadableEvents|waiter.WritableEvents) + clientWE, clientCH := waiter.NewChannelEntry(waiter.ReadableEvents | waiter.WritableEvents) + clientWQ.EventRegister(&clientWE) clientEP, err := host1Stack.NewEndpoint(tcp.ProtocolNumber, ipv6.ProtocolNumber, &clientWQ) if err != nil { t.Fatalf("host1Stack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, ipv6.ProtocolNumber, err) @@ -1203,8 +1203,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { neighborAddr: utils.RouterNIC1IPv4Addr.AddressWithPrefix.Address, getEndpoints: func(t *testing.T, host1Stack, routerStack, _ *stack.Stack) (tcpip.Endpoint, <-chan struct{}, tcpip.Endpoint, <-chan struct{}) { var listenerWQ waiter.Queue - listenerWE, listenerCH := waiter.NewChannelEntry(nil) - listenerWQ.EventRegister(&listenerWE, waiter.EventIn) + listenerWE, listenerCH := waiter.NewChannelEntry(waiter.EventIn) + listenerWQ.EventRegister(&listenerWE) listenerEP, err := host1Stack.NewEndpoint(tcp.ProtocolNumber, ipv4.ProtocolNumber, &listenerWQ) if err != nil { t.Fatalf("host1Stack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, ipv4.ProtocolNumber, err) @@ -1212,8 +1212,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { t.Cleanup(listenerEP.Close) var clientWQ waiter.Queue - clientWE, clientCH := waiter.NewChannelEntry(nil) - clientWQ.EventRegister(&clientWE, waiter.ReadableEvents|waiter.WritableEvents) + clientWE, clientCH := waiter.NewChannelEntry(waiter.ReadableEvents | waiter.WritableEvents) + clientWQ.EventRegister(&clientWE) clientEP, err := routerStack.NewEndpoint(tcp.ProtocolNumber, ipv4.ProtocolNumber, &clientWQ) if err != nil { t.Fatalf("routerStack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, ipv4.ProtocolNumber, err) @@ -1230,8 +1230,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { neighborAddr: utils.RouterNIC1IPv6Addr.AddressWithPrefix.Address, getEndpoints: func(t *testing.T, host1Stack, routerStack, _ *stack.Stack) (tcpip.Endpoint, <-chan struct{}, tcpip.Endpoint, <-chan struct{}) { var listenerWQ waiter.Queue - listenerWE, listenerCH := waiter.NewChannelEntry(nil) - listenerWQ.EventRegister(&listenerWE, waiter.EventIn) + listenerWE, listenerCH := waiter.NewChannelEntry(waiter.EventIn) + listenerWQ.EventRegister(&listenerWE) listenerEP, err := host1Stack.NewEndpoint(tcp.ProtocolNumber, ipv6.ProtocolNumber, &listenerWQ) if err != nil { t.Fatalf("host1Stack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, ipv6.ProtocolNumber, err) @@ -1239,8 +1239,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { t.Cleanup(listenerEP.Close) var clientWQ waiter.Queue - clientWE, clientCH := waiter.NewChannelEntry(nil) - clientWQ.EventRegister(&clientWE, waiter.ReadableEvents|waiter.WritableEvents) + clientWE, clientCH := waiter.NewChannelEntry(waiter.ReadableEvents | waiter.WritableEvents) + clientWQ.EventRegister(&clientWE) clientEP, err := routerStack.NewEndpoint(tcp.ProtocolNumber, ipv6.ProtocolNumber, &clientWQ) if err != nil { t.Fatalf("routerStack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, ipv6.ProtocolNumber, err) @@ -1257,8 +1257,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { neighborAddr: utils.RouterNIC1IPv4Addr.AddressWithPrefix.Address, getEndpoints: func(t *testing.T, host1Stack, _, host2Stack *stack.Stack) (tcpip.Endpoint, <-chan struct{}, tcpip.Endpoint, <-chan struct{}) { var listenerWQ waiter.Queue - listenerWE, listenerCH := waiter.NewChannelEntry(nil) - listenerWQ.EventRegister(&listenerWE, waiter.EventIn) + listenerWE, listenerCH := waiter.NewChannelEntry(waiter.EventIn) + listenerWQ.EventRegister(&listenerWE) listenerEP, err := host1Stack.NewEndpoint(tcp.ProtocolNumber, ipv4.ProtocolNumber, &listenerWQ) if err != nil { t.Fatalf("host1Stack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, ipv4.ProtocolNumber, err) @@ -1266,8 +1266,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { t.Cleanup(listenerEP.Close) var clientWQ waiter.Queue - clientWE, clientCH := waiter.NewChannelEntry(nil) - clientWQ.EventRegister(&clientWE, waiter.ReadableEvents|waiter.WritableEvents) + clientWE, clientCH := waiter.NewChannelEntry(waiter.ReadableEvents | waiter.WritableEvents) + clientWQ.EventRegister(&clientWE) clientEP, err := host2Stack.NewEndpoint(tcp.ProtocolNumber, ipv4.ProtocolNumber, &clientWQ) if err != nil { t.Fatalf("host2Stack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, ipv4.ProtocolNumber, err) @@ -1284,8 +1284,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { neighborAddr: utils.RouterNIC1IPv6Addr.AddressWithPrefix.Address, getEndpoints: func(t *testing.T, host1Stack, _, host2Stack *stack.Stack) (tcpip.Endpoint, <-chan struct{}, tcpip.Endpoint, <-chan struct{}) { var listenerWQ waiter.Queue - listenerWE, listenerCH := waiter.NewChannelEntry(nil) - listenerWQ.EventRegister(&listenerWE, waiter.EventIn) + listenerWE, listenerCH := waiter.NewChannelEntry(waiter.EventIn) + listenerWQ.EventRegister(&listenerWE) listenerEP, err := host1Stack.NewEndpoint(tcp.ProtocolNumber, ipv6.ProtocolNumber, &listenerWQ) if err != nil { t.Fatalf("host1Stack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, ipv6.ProtocolNumber, err) @@ -1293,8 +1293,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { t.Cleanup(listenerEP.Close) var clientWQ waiter.Queue - clientWE, clientCH := waiter.NewChannelEntry(nil) - clientWQ.EventRegister(&clientWE, waiter.ReadableEvents|waiter.WritableEvents) + clientWE, clientCH := waiter.NewChannelEntry(waiter.ReadableEvents | waiter.WritableEvents) + clientWQ.EventRegister(&clientWE) clientEP, err := host2Stack.NewEndpoint(tcp.ProtocolNumber, ipv6.ProtocolNumber, &clientWQ) if err != nil { t.Fatalf("host2Stack.NewEndpoint(%d, %d, _): %s", tcp.ProtocolNumber, ipv6.ProtocolNumber, err) @@ -1415,8 +1415,8 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { t.Fatalf("listenerEP.Accept(): %s", err) } defer peerEP.Close() - peerWE, peerCH := waiter.NewChannelEntry(nil) - peerWQ.EventRegister(&peerWE, waiter.ReadableEvents) + peerWE, peerCH := waiter.NewChannelEntry(waiter.ReadableEvents) + peerWQ.EventRegister(&peerWE) // Wait for the neighbor to be stale again then send data to the remote. // diff --git a/pkg/tcpip/tests/integration/loopback_test.go b/pkg/tcpip/tests/integration/loopback_test.go index f33223e79..244119f25 100644 --- a/pkg/tcpip/tests/integration/loopback_test.go +++ b/pkg/tcpip/tests/integration/loopback_test.go @@ -446,8 +446,8 @@ func TestLoopbackAcceptAllInSubnetTCP(t *testing.T) { }) var wq waiter.Queue - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) listeningEndpoint, err := s.NewEndpoint(tcp.ProtocolNumber, test.addAddress.Protocol, &wq) if err != nil { diff --git a/pkg/tcpip/tests/integration/multicast_broadcast_test.go b/pkg/tcpip/tests/integration/multicast_broadcast_test.go index 7753e7d6e..367636955 100644 --- a/pkg/tcpip/tests/integration/multicast_broadcast_test.go +++ b/pkg/tcpip/tests/integration/multicast_broadcast_test.go @@ -501,8 +501,8 @@ func TestReuseAddrAndBroadcast(t *testing.T) { // packet. for i := 0; i < 2; i++ { var wq waiter.Queue - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) ep, err := s.NewEndpoint(udp.ProtocolNumber, ipv4.ProtocolNumber, &wq) if err != nil { t.Fatalf("(eps[%d]) NewEndpoint(%d, %d, _): %s", len(eps), udp.ProtocolNumber, ipv4.ProtocolNumber, err) diff --git a/pkg/tcpip/tests/integration/route_test.go b/pkg/tcpip/tests/integration/route_test.go index 422eb8408..6d0b938d5 100644 --- a/pkg/tcpip/tests/integration/route_test.go +++ b/pkg/tcpip/tests/integration/route_test.go @@ -196,8 +196,8 @@ func TestLocalPing(t *testing.T) { } var wq waiter.Queue - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) ep, err := s.NewEndpoint(test.transProto, test.netProto, &wq) if err != nil { t.Fatalf("s.NewEndpoint(%d, %d, _): %s", test.transProto, test.netProto, err) @@ -319,8 +319,8 @@ func TestLocalUDP(t *testing.T) { } var serverWQ waiter.Queue - serverWE, serverCH := waiter.NewChannelEntry(nil) - serverWQ.EventRegister(&serverWE, waiter.ReadableEvents) + serverWE, serverCH := waiter.NewChannelEntry(waiter.ReadableEvents) + serverWQ.EventRegister(&serverWE) server, err := s.NewEndpoint(udp.ProtocolNumber, test.firstPrimaryAddr.Protocol, &serverWQ) if err != nil { t.Fatalf("s.NewEndpoint(%d, %d): %s", udp.ProtocolNumber, test.firstPrimaryAddr.Protocol, err) @@ -333,8 +333,8 @@ func TestLocalUDP(t *testing.T) { } var clientWQ waiter.Queue - clientWE, clientCH := waiter.NewChannelEntry(nil) - clientWQ.EventRegister(&clientWE, waiter.ReadableEvents) + clientWE, clientCH := waiter.NewChannelEntry(waiter.ReadableEvents) + clientWQ.EventRegister(&clientWE) client, err := s.NewEndpoint(udp.ProtocolNumber, test.firstPrimaryAddr.Protocol, &clientWQ) if err != nil { t.Fatalf("s.NewEndpoint(%d, %d): %s", udp.ProtocolNumber, test.firstPrimaryAddr.Protocol, err) diff --git a/pkg/tcpip/transport/tcp/dual_stack_test.go b/pkg/tcpip/transport/tcp/dual_stack_test.go index 5342aacfd..667fc3e2c 100644 --- a/pkg/tcpip/transport/tcp/dual_stack_test.go +++ b/pkg/tcpip/transport/tcp/dual_stack_test.go @@ -45,8 +45,8 @@ func TestV4MappedConnectOnV6Only(t *testing.T) { func testV4Connect(t *testing.T, c *context.Context, checkers ...checker.NetworkChecker) { // Start connection attempt. - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.WritableEvents) + we, ch := waiter.NewChannelEntry(waiter.WritableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) err := c.EP.Connect(tcpip.FullAddress{Addr: context.TestV4MappedAddr, Port: context.TestPort}) @@ -152,8 +152,8 @@ func TestV4ConnectWhenBoundToV4Mapped(t *testing.T) { func testV6Connect(t *testing.T, c *context.Context, checkers ...checker.NetworkChecker) { // Start connection attempt to IPv6 address. - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.WritableEvents) + we, ch := waiter.NewChannelEntry(waiter.WritableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) err := c.EP.Connect(tcpip.FullAddress{Addr: context.TestV6Addr, Port: context.TestPort}) @@ -387,8 +387,8 @@ func testV4Accept(t *testing.T, c *context.Context) { }) // Try to accept the connection. - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) nep, _, err := c.EP.Accept(nil) @@ -521,8 +521,8 @@ func TestV6AcceptOnV6(t *testing.T) { }) // Try to accept the connection. - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) var addr tcpip.FullAddress _, _, err := c.EP.Accept(&addr) @@ -608,8 +608,8 @@ func testV4ListenClose(t *testing.T, c *context.Context) { } // Try to accept the connection. - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) nep, _, err := c.EP.Accept(nil) if cmp.Equal(&tcpip.ErrWouldBlock{}, err) { diff --git a/pkg/tcpip/transport/tcp/endpoint.go b/pkg/tcpip/transport/tcp/endpoint.go index 066ffe051..4ceeb2e47 100644 --- a/pkg/tcpip/transport/tcp/endpoint.go +++ b/pkg/tcpip/transport/tcp/endpoint.go @@ -3053,8 +3053,8 @@ func (e *endpoint) Stats() tcpip.EndpointStats { // Wait implements stack.TransportEndpoint.Wait. func (e *endpoint) Wait() { - waitEntry, notifyCh := waiter.NewChannelEntry(nil) - e.waiterQueue.EventRegister(&waitEntry, waiter.EventHUp) + waitEntry, notifyCh := waiter.NewChannelEntry(waiter.EventHUp) + e.waiterQueue.EventRegister(&waitEntry) defer e.waiterQueue.EventUnregister(&waitEntry) for { e.LockUser() diff --git a/pkg/tcpip/transport/tcp/tcp_test.go b/pkg/tcpip/transport/tcp/tcp_test.go index 6f1ee3816..0566e53a3 100644 --- a/pkg/tcpip/transport/tcp/tcp_test.go +++ b/pkg/tcpip/transport/tcp/tcp_test.go @@ -125,8 +125,8 @@ func TestGiveUpConnect(t *testing.T) { } // Register for notification, then start connection attempt. - waitEntry, notifyCh := waiter.NewChannelEntry(nil) - wq.EventRegister(&waitEntry, waiter.EventHUp) + waitEntry, notifyCh := waiter.NewChannelEntry(waiter.EventHUp) + wq.EventRegister(&waitEntry) defer wq.EventUnregister(&waitEntry) { @@ -171,8 +171,8 @@ func TestConnectICMPError(t *testing.T) { t.Fatalf("NewEndpoint failed: %s", err) } - waitEntry, notifyCh := waiter.NewChannelEntry(nil) - wq.EventRegister(&waitEntry, waiter.EventHUp) + waitEntry, notifyCh := waiter.NewChannelEntry(waiter.EventHUp) + wq.EventRegister(&waitEntry) defer wq.EventUnregister(&waitEntry) { @@ -459,8 +459,8 @@ func TestTCPResetSentForACKWhenNotUsingSynCookies(t *testing.T) { c.SendPacket(nil, ackHeaders) // Try to accept the connection. - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) c.EP, _, err = ep.Accept(nil) @@ -880,8 +880,8 @@ func TestSimpleReceive(t *testing.T) { c.CreateConnected(context.TestInitialSequenceNumber, 30000, -1 /* epRcvBuf */) - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) ept := endpointTester{c.EP} @@ -1360,14 +1360,6 @@ func TestListenShutdown(t *testing.T) { )) } -var _ waiter.EntryCallback = (callback)(nil) - -type callback func(*waiter.Entry, waiter.EventMask) - -func (cb callback) Callback(entry *waiter.Entry, mask waiter.EventMask) { - cb(entry, mask) -} - func TestListenerReadinessOnEvent(t *testing.T) { s := stack.New(stack.Options{ TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol}, @@ -1423,15 +1415,15 @@ func TestListenerReadinessOnEvent(t *testing.T) { events := make(chan waiter.EventMask) // Scope `entry` to allow a binding of the same name below. { - entry := waiter.Entry{Callback: callback(func(_ *waiter.Entry, mask waiter.EventMask) { + entry := waiter.NewFunctionEntry(waiter.EventIn, func(mask waiter.EventMask) { events <- ep.Readiness(mask) - })} - wq.EventRegister(&entry, waiter.EventIn) + }) + wq.EventRegister(&entry) defer wq.EventUnregister(&entry) } - entry, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&entry, waiter.EventOut) + entry, ch := waiter.NewChannelEntry(waiter.EventOut) + wq.EventRegister(&entry) defer wq.EventUnregister(&entry) switch err := conn.Connect(address).(type) { @@ -1473,8 +1465,8 @@ func TestListenCloseWhileConnect(t *testing.T) { t.Fatal("Listen failed:", err) } - waitEntry, notifyCh := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&waitEntry, waiter.ReadableEvents) + waitEntry, notifyCh := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&waitEntry) defer c.WQ.EventUnregister(&waitEntry) executeHandshake(t, c, context.TestPort, false /* synCookiesInUse */) @@ -1613,8 +1605,8 @@ func TestConnectBindToDevice(t *testing.T) { t.Fatalf("c.EP.SetSockOpt(&%T(%d)): %s", test.device, test.device, err) } // Start connection attempt. - waitEntry, _ := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&waitEntry, waiter.WritableEvents) + waitEntry, _ := waiter.NewChannelEntry(waiter.WritableEvents) + c.WQ.EventRegister(&waitEntry) defer c.WQ.EventUnregister(&waitEntry) err := c.EP.Connect(tcpip.FullAddress{Addr: context.TestAddr, Port: context.TestPort}) @@ -1673,8 +1665,8 @@ func TestShutdownConnectingSocket(t *testing.T) { // the handshake process. c.Create(-1) - waitEntry, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&waitEntry, waiter.EventHUp) + waitEntry, ch := waiter.NewChannelEntry(waiter.EventHUp) + c.WQ.EventRegister(&waitEntry) defer c.WQ.EventUnregister(&waitEntry) // Start connection attempt. @@ -1738,8 +1730,8 @@ func TestSynSent(t *testing.T) { c.Create(-1) // Start connection attempt. - waitEntry, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&waitEntry, waiter.EventHUp) + waitEntry, ch := waiter.NewChannelEntry(waiter.EventHUp) + c.WQ.EventRegister(&waitEntry) defer c.WQ.EventUnregister(&waitEntry) addr := tcpip.FullAddress{Addr: context.TestAddr, Port: context.TestPort} @@ -1813,8 +1805,8 @@ func TestOutOfOrderReceive(t *testing.T) { c.CreateConnected(context.TestInitialSequenceNumber, 30000, -1 /* epRcvBuf */) - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) ept := endpointTester{c.EP} @@ -1955,8 +1947,8 @@ func TestRstOnCloseWithUnreadData(t *testing.T) { c.CreateConnected(context.TestInitialSequenceNumber, 30000, -1 /* epRcvBuf */) - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) ept := endpointTester{c.EP} @@ -2024,8 +2016,8 @@ func TestRstOnCloseWithUnreadDataFinConvertRst(t *testing.T) { c.CreateConnected(context.TestInitialSequenceNumber, 30000, -1 /* epRcvBuf */) - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) ept := endpointTester{c.EP} @@ -2136,8 +2128,8 @@ func TestFullWindowReceive(t *testing.T) { const rcvBufSz = 10 c.CreateConnected(context.TestInitialSequenceNumber, 30000, rcvBufSz) - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) ept := endpointTester{c.EP} @@ -2242,14 +2234,14 @@ func TestSmallReceiveBufferReadiness(t *testing.T) { }) } - listenerEntry, listenerCh := waiter.NewChannelEntry(nil) + listenerEntry, listenerCh := waiter.NewChannelEntry(waiter.ReadableEvents) var listenerWQ waiter.Queue listener, err := s.NewEndpoint(tcp.ProtocolNumber, ipv4.ProtocolNumber, &listenerWQ) if err != nil { t.Fatalf("NewEndpoint failed: %s", err) } defer listener.Close() - listenerWQ.EventRegister(&listenerEntry, waiter.ReadableEvents) + listenerWQ.EventRegister(&listenerEntry) defer listenerWQ.EventUnregister(&listenerEntry) if err := listener.Bind(tcpip.FullAddress{}); err != nil { @@ -2292,12 +2284,12 @@ func TestSmallReceiveBufferReadiness(t *testing.T) { // Send buffer size doesn't seem to affect this test. // server.SocketOptions().SetSendBufferSize(size, true) - clientEntry, clientCh := waiter.NewChannelEntry(nil) - clientWQ.EventRegister(&clientEntry, waiter.ReadableEvents) + clientEntry, clientCh := waiter.NewChannelEntry(waiter.ReadableEvents) + clientWQ.EventRegister(&clientEntry) defer clientWQ.EventUnregister(&clientEntry) - serverEntry, serverCh := waiter.NewChannelEntry(nil) - serverWQ.EventRegister(&serverEntry, waiter.WritableEvents) + serverEntry, serverCh := waiter.NewChannelEntry(waiter.WritableEvents) + serverWQ.EventRegister(&serverEntry) defer serverWQ.EventUnregister(&serverEntry) var total int64 @@ -2502,8 +2494,8 @@ func TestNoWindowShrinking(t *testing.T) { header.TCPOptionWS, 3, 0, header.TCPOptionNOP, }) - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) ept := endpointTester{c.EP} @@ -2820,8 +2812,8 @@ func TestScaledWindowAccept(t *testing.T) { c.PassiveConnectWithOptions(100, 3 /* wndScale */, header.TCPSynOptions{MSS: defaultIPv4MSS}, 0 /* delay */) // Try to accept the connection. - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) c.EP, _, err = ep.Accept(nil) @@ -2892,8 +2884,8 @@ func TestNonScaledWindowAccept(t *testing.T) { c.PassiveConnect(100, -1, header.TCPSynOptions{MSS: defaultIPv4MSS}) // Try to accept the connection. - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) c.EP, _, err = ep.Accept(nil) @@ -3484,8 +3476,8 @@ func TestPassiveSendMSSLessThanMTU(t *testing.T) { c.PassiveConnect(maxPayload, -1, header.TCPSynOptions{MSS: mtu - header.IPv4MinimumSize - header.TCPMinimumSize}) // Try to accept the connection. - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) c.EP, _, err = ep.Accept(nil) @@ -3538,8 +3530,8 @@ func TestSynCookiePassiveSendMSSLessThanMTU(t *testing.T) { c.PassiveConnect(maxPayload, -1, header.TCPSynOptions{MSS: mtu - header.IPv4MinimumSize - header.TCPMinimumSize}) // Try to accept the connection. - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) c.EP, _, err = ep.Accept(nil) @@ -3612,8 +3604,8 @@ func TestSynOptionsOnActiveConnect(t *testing.T) { c.EP.SocketOptions().SetReceiveBufferSize(rcvBufferSize*2, true /* notify */) // Start connection attempt. - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.WritableEvents) + we, ch := waiter.NewChannelEntry(waiter.WritableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) { @@ -3725,8 +3717,8 @@ func TestReceiveOnResetConnection(t *testing.T) { }) // Try to read. - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) loop: @@ -3822,8 +3814,8 @@ func TestMaxRetransmitsTimeout(t *testing.T) { } c.CreateConnected(context.TestInitialSequenceNumber, 30000 /* rcvWnd */, -1 /* epRcvBuf */) - waitEntry, notifyCh := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&waitEntry, waiter.EventHUp) + waitEntry, notifyCh := waiter.NewChannelEntry(waiter.EventHUp) + c.WQ.EventRegister(&waitEntry) defer c.WQ.EventUnregister(&waitEntry) var r bytes.Reader @@ -4706,8 +4698,8 @@ func TestReadAfterClosedState(t *testing.T) { c.CreateConnected(context.TestInitialSequenceNumber, 30000, -1 /* epRcvBuf */) - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) ept := endpointTester{c.EP} @@ -5081,8 +5073,8 @@ func TestSelfConnect(t *testing.T) { } // Register for notification, then start connection attempt. - waitEntry, notifyCh := waiter.NewChannelEntry(nil) - wq.EventRegister(&waitEntry, waiter.WritableEvents) + waitEntry, notifyCh := waiter.NewChannelEntry(waiter.WritableEvents) + wq.EventRegister(&waitEntry) defer wq.EventUnregister(&waitEntry) { @@ -5107,7 +5099,8 @@ func TestSelfConnect(t *testing.T) { // Read back what was written. wq.EventUnregister(&waitEntry) - wq.EventRegister(&waitEntry, waiter.ReadableEvents) + waitEntry, notifyCh = waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&waitEntry) ept := endpointTester{ep} rd := ept.CheckReadFull(t, len(data), notifyCh, 5*time.Second) @@ -5804,8 +5797,8 @@ func TestListenBacklogFull(t *testing.T) { c.CheckNoPacketTimeout("unexpected packet received", 50*time.Millisecond) // Try to accept the connections in the backlog. - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) for i := 0; i < listenBacklog; i++ { @@ -6134,8 +6127,8 @@ func TestListenSynRcvdQueueFull(t *testing.T) { }) // Verify if that is delivered to the accept queue. - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) <-ch @@ -6218,8 +6211,8 @@ func TestListenBacklogFullSynCookieInUse(t *testing.T) { c.CheckNoPacketTimeout("unexpected packet received", 50*time.Millisecond) // Verify that there is only one acceptable connection at this point. - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) _, _, err = c.EP.Accept(nil) @@ -6362,8 +6355,8 @@ func TestSynRcvdBadSeqNumber(t *testing.T) { // did not change the state from SYN-RCVD. // Get setup to be notified about connection establishment. - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) // Send ACK to move to ESTABLISHED state. @@ -6425,8 +6418,8 @@ func TestPassiveConnectionAttemptIncrement(t *testing.T) { srcPort := uint16(context.TestPort) executeHandshake(t, c, srcPort+1, false) - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) // Verify that there is only one acceptable connection at this point. @@ -6509,8 +6502,8 @@ func TestPassiveFailedConnectionAttemptIncrement(t *testing.T) { t.FailNow() } - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) // Now check that there is one acceptable connections. @@ -6605,8 +6598,8 @@ func TestEndpointBindListenAcceptState(t *testing.T) { c.PassiveConnectWithOptions(100, 5, header.TCPSynOptions{MSS: defaultIPv4MSS}, 0 /* delay */) // Try to accept the connection. - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) aep, _, err := ep.Accept(nil) @@ -7064,8 +7057,8 @@ func TestTCPTimeWaitRSTIgnored(t *testing.T) { c.SendPacket(nil, ackHeaders) // Try to accept the connection. - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) c.EP, _, err = ep.Accept(nil) @@ -7183,8 +7176,8 @@ func TestTCPTimeWaitOutOfOrder(t *testing.T) { c.SendPacket(nil, ackHeaders) // Try to accept the connection. - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) c.EP, _, err = ep.Accept(nil) @@ -7290,8 +7283,8 @@ func TestTCPTimeWaitNewSyn(t *testing.T) { c.SendPacket(nil, ackHeaders) // Try to accept the connection. - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) c.EP, _, err = ep.Accept(nil) @@ -7454,8 +7447,8 @@ func TestTCPTimeWaitDuplicateFINExtendsTimeWait(t *testing.T) { c.SendPacket(nil, ackHeaders) // Try to accept the connection. - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) c.EP, _, err = ep.Accept(nil) @@ -7604,8 +7597,8 @@ func TestTCPCloseWithData(t *testing.T) { c.SendPacket(nil, ackHeaders) // Try to accept the connection. - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) c.EP, _, err = ep.Accept(nil) @@ -7735,8 +7728,8 @@ func TestTCPUserTimeout(t *testing.T) { } c.CreateConnected(context.TestInitialSequenceNumber, 30000, -1 /* epRcvBuf */) - waitEntry, notifyCh := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&waitEntry, waiter.EventHUp) + waitEntry, notifyCh := waiter.NewChannelEntry(waiter.EventHUp) + c.WQ.EventRegister(&waitEntry) defer c.WQ.EventUnregister(&waitEntry) origEstablishedTimedout := c.Stack().Stats().TCP.EstablishedTimedout.Value() @@ -8456,8 +8449,8 @@ func TestSendBufferTuning(t *testing.T) { data[i] = byte(i) } - w, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&w, waiter.WritableEvents) + w, ch := waiter.NewChannelEntry(waiter.WritableEvents) + c.WQ.EventRegister(&w) defer c.WQ.EventUnregister(&w) bytesRead := 0 @@ -8566,8 +8559,8 @@ func TestTimestampSynCookies(t *testing.T) { }) c.EP, _, err = ep.Accept(nil) // Try to accept the connection. - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) if cmp.Equal(&tcpip.ErrWouldBlock{}, err) { // Wait for connection to be established. diff --git a/pkg/tcpip/transport/tcp/tcp_timestamp_test.go b/pkg/tcpip/transport/tcp/tcp_timestamp_test.go index 65925daa5..a18f6b2d7 100644 --- a/pkg/tcpip/transport/tcp/tcp_timestamp_test.go +++ b/pkg/tcpip/transport/tcp/tcp_timestamp_test.go @@ -45,8 +45,8 @@ func TestTimeStampEnabledConnect(t *testing.T) { rep := createConnectedWithTimestampOption(c) // Register for read and validate that we have data to read. - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) // The following tests ensure that TS option once enabled behaves @@ -272,8 +272,8 @@ func TestSegmentNotDroppedWhenTimestampMissing(t *testing.T) { rep := createConnectedWithTimestampOption(c) // Register for read. - we, ch := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.WQ.EventRegister(&we) defer c.WQ.EventUnregister(&we) droppedPacketsStat := c.Stack().Stats().DroppedPackets diff --git a/pkg/tcpip/transport/tcp/testing/context/context.go b/pkg/tcpip/transport/tcp/testing/context/context.go index fd0eca4bd..9b32044b9 100644 --- a/pkg/tcpip/transport/tcp/testing/context/context.go +++ b/pkg/tcpip/transport/tcp/testing/context/context.go @@ -692,8 +692,8 @@ func (c *Context) Connect(iss seqnum.Value, rcvWnd seqnum.Size, options []byte) c.t.Helper() // Start connection attempt. - waitEntry, notifyCh := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&waitEntry, waiter.WritableEvents) + waitEntry, notifyCh := waiter.NewChannelEntry(waiter.WritableEvents) + c.WQ.EventRegister(&waitEntry) defer c.WQ.EventUnregister(&waitEntry) err := c.EP.Connect(tcpip.FullAddress{Addr: TestAddr, Port: TestPort}) @@ -911,8 +911,8 @@ func (c *Context) CreateConnectedWithOptions(wantOptions header.TCPSynOptions, d } // Start connection attempt. - waitEntry, notifyCh := waiter.NewChannelEntry(nil) - c.WQ.EventRegister(&waitEntry, waiter.WritableEvents) + waitEntry, notifyCh := waiter.NewChannelEntry(waiter.WritableEvents) + c.WQ.EventRegister(&waitEntry) defer c.WQ.EventUnregister(&waitEntry) testFullAddr := tcpip.FullAddress{Addr: TestAddr, Port: TestPort} @@ -1082,8 +1082,8 @@ func (c *Context) AcceptWithOptions(wndScale int, synOptions header.TCPSynOption rep := c.PassiveConnectWithOptions(100, wndScale, synOptions, delay) // Try to accept the connection. - we, ch := waiter.NewChannelEntry(nil) - wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) defer wq.EventUnregister(&we) c.EP, _, err = ep.Accept(nil) diff --git a/pkg/tcpip/transport/udp/udp_test.go b/pkg/tcpip/transport/udp/udp_test.go index b3199489c..993acf109 100644 --- a/pkg/tcpip/transport/udp/udp_test.go +++ b/pkg/tcpip/transport/udp/udp_test.go @@ -608,8 +608,8 @@ func testReadInternal(c *testContext, flow testFlow, packetShouldBeDropped, expe c.injectPacket(flow, payload, false) // Try to receive the data. - we, ch := waiter.NewChannelEntry(nil) - c.wq.EventRegister(&we, waiter.ReadableEvents) + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + c.wq.EventRegister(&we) defer c.wq.EventUnregister(&we) // Take a snapshot of the stats to validate them at the end of the test. diff --git a/pkg/waiter/waiter.go b/pkg/waiter/waiter.go index 4ea067cb7..4d6dd4ce3 100644 --- a/pkg/waiter/waiter.go +++ b/pkg/waiter/waiter.go @@ -107,16 +107,16 @@ type Waitable interface { // EventRegister registers the given waiter entry to receive // notifications when an event occurs that makes the object ready for // at least one of the events in mask. - EventRegister(e *Entry, mask EventMask) + EventRegister(e *Entry) // EventUnregister unregisters a waiter entry previously registered with // EventRegister(). EventUnregister(e *Entry) } -// EntryCallback provides a notify callback. -type EntryCallback interface { - // Callback is the function to be called when the waiter entry is +// EventListener provides a notify callback. +type EventListener interface { + // NotifyEvent is the function to be called when the waiter entry is // notified. It is responsible for doing whatever is needed to wake up // the waiter. // @@ -126,7 +126,7 @@ type EntryCallback interface { // // The mask indicates the events that occurred and that the entry is // interested in. - Callback(e *Entry, mask EventMask) + NotifyEvent(mask EventMask) } // Entry represents a waiter that can be add to the a wait queue. It can @@ -135,21 +135,44 @@ type EntryCallback interface { // // +stateify savable type Entry struct { - Callback EntryCallback - - // The following fields are protected by the queue lock. - mask EventMask waiterEntry + + // eventListener receives the notification. + eventListener EventListener + + // mask should be immutable once queued. + mask EventMask } -type channelCallback struct { - ch chan struct{} +// Init initializes the Entry. +// +// This must only be called when unregistered. +func (e *Entry) Init(eventListener EventListener, mask EventMask) { + e.eventListener = eventListener + e.mask = mask } -// Callback implements EntryCallback.Callback. -func (c *channelCallback) Callback(*Entry, EventMask) { +// Mask returns the entry mask. +func (e *Entry) Mask() EventMask { + return e.mask +} + +// NotifyEvent notifies the event listener. +// +// Mask should be the full set of active events. +func (e *Entry) NotifyEvent(mask EventMask) { + if m := mask & e.mask; m != 0 { + e.eventListener.NotifyEvent(m) + } +} + +// ChannelNotifier is a simple channel-based notification. +type ChannelNotifier chan struct{} + +// NotifyEvent implements waiter.EventListener.NotifyEvent. +func (c ChannelNotifier) NotifyEvent(EventMask) { select { - case c.ch <- struct{}{}: + case c <- struct{}{}: default: } } @@ -157,15 +180,23 @@ func (c *channelCallback) Callback(*Entry, EventMask) { // NewChannelEntry initializes a new Entry that does a non-blocking write to a // struct{} channel when the callback is called. It returns the new Entry // instance and the channel being used. -// -// If a channel isn't specified (i.e., if "c" is nil), then NewChannelEntry -// allocates a new channel. -func NewChannelEntry(c chan struct{}) (Entry, chan struct{}) { - if c == nil { - c = make(chan struct{}, 1) - } +func NewChannelEntry(mask EventMask) (e Entry, ch chan struct{}) { + ch = make(chan struct{}, 1) + e.Init(ChannelNotifier(ch), mask) + return e, ch +} - return Entry{Callback: &channelCallback{ch: c}}, c +type functionNotifier func(EventMask) + +// NotifyEvent implements waiter.EventListener.NotifyEvent. +func (f functionNotifier) NotifyEvent(mask EventMask) { + f(mask) +} + +// NewFunctionEntry initializes a new Entry that calls the given function. +func NewFunctionEntry(mask EventMask, fn func(EventMask)) (e Entry) { + e.Init(functionNotifier(fn), mask) + return e } // Queue represents the wait queue where waiters can be added and @@ -179,11 +210,9 @@ type Queue struct { mu sync.RWMutex `state:"nosave"` } -// EventRegister adds a waiter to the wait queue; the waiter will be notified -// when at least one of the events specified in mask happens. -func (q *Queue) EventRegister(e *Entry, mask EventMask) { +// EventRegister adds a waiter to the wait queue. +func (q *Queue) EventRegister(e *Entry) { q.mu.Lock() - e.mask = mask q.list.PushBack(e) q.mu.Unlock() } @@ -200,9 +229,11 @@ func (q *Queue) EventUnregister(e *Entry) { func (q *Queue) Notify(mask EventMask) { q.mu.RLock() for e := q.list.Front(); e != nil; e = e.Next() { - if m := mask & e.mask; m != 0 { - e.Callback.Callback(e, m) + m := mask & e.mask + if m == 0 { + continue } + e.eventListener.NotifyEvent(m) // Skip intermediate call. } q.mu.RUnlock() } @@ -242,7 +273,7 @@ func (*AlwaysReady) Readiness(mask EventMask) EventMask { // EventRegister doesn't do anything because this object doesn't need to issue // notifications because its readiness never changes. -func (*AlwaysReady) EventRegister(*Entry, EventMask) { +func (*AlwaysReady) EventRegister(e *Entry) { } // EventUnregister doesn't do anything because this object doesn't need to issue diff --git a/pkg/waiter/waiter_test.go b/pkg/waiter/waiter_test.go index 6928f28b4..dbd127aa0 100644 --- a/pkg/waiter/waiter_test.go +++ b/pkg/waiter/waiter_test.go @@ -19,15 +19,6 @@ import ( "testing" ) -type callbackStub struct { - f func(e *Entry, m EventMask) -} - -// Callback implements EntryCallback.Callback. -func (c *callbackStub) Callback(e *Entry, m EventMask) { - c.f(e, m) -} - func TestEmptyQueue(t *testing.T) { var q Queue @@ -36,8 +27,8 @@ func TestEmptyQueue(t *testing.T) { // Register then unregister a waiter, then notify the queue. cnt := 0 - e := Entry{Callback: &callbackStub{func(*Entry, EventMask) { cnt++ }}} - q.EventRegister(&e, EventIn) + e := NewFunctionEntry(EventIn, func(EventMask) { cnt++ }) + q.EventRegister(&e) q.EventUnregister(&e) q.Notify(EventIn) if cnt != 0 { @@ -49,8 +40,8 @@ func TestMask(t *testing.T) { // Register a waiter. var q Queue var cnt int - e := Entry{Callback: &callbackStub{func(*Entry, EventMask) { cnt++ }}} - q.EventRegister(&e, EventIn|EventErr) + e := NewFunctionEntry(EventIn|EventErr, func(EventMask) { cnt++ }) + q.EventRegister(&e) // Notify with an overlapping mask. cnt = 0 @@ -100,20 +91,16 @@ func TestConcurrentRegistration(t *testing.T) { // Create goroutines that will all register/unregister concurrently. for i := 0; i < concurrency; i++ { go func() { - var e Entry - e.Callback = &callbackStub{func(entry *Entry, mask EventMask) { + e := NewFunctionEntry(EventIn|EventErr, func(mask EventMask) { cnt++ - if entry != &e { - t.Errorf("entry = %p, want %p", entry, &e) - } if mask != EventIn { t.Errorf("mask = %#x want %#x", mask, EventIn) } - }} + }) // Wait for notification, then register. <-ch1 - q.EventRegister(&e, EventIn|EventErr) + q.EventRegister(&e) // Tell main goroutine that we're done registering. ch2 <- struct{}{} @@ -160,18 +147,14 @@ func TestConcurrentNotification(t *testing.T) { // Register waiters. for i := 0; i < waiterCount; i++ { - var e Entry - e.Callback = &callbackStub{func(entry *Entry, mask EventMask) { + e := NewFunctionEntry(EventIn|EventErr, func(mask EventMask) { atomic.AddInt32(&cnt, 1) - if entry != &e { - t.Errorf("entry = %p, want %p", entry, &e) - } if mask != EventIn { t.Errorf("mask = %#x want %#x", mask, EventIn) } - }} + }) - q.EventRegister(&e, EventIn|EventErr) + q.EventRegister(&e) } // Launch notifiers.