From 91f58d2cc8d31d5e81e3bc2bf4b2e5edf7377408 Mon Sep 17 00:00:00 2001 From: Adin Scannell Date: Sat, 13 Nov 2021 12:51:42 -0800 Subject: [PATCH] Update Waitable API. Instead of passing the event mask at registratrion time, pass the mask as part of the waiter. This makes the mask immutable and simplifies the architecture of waiters. This is also necessary for a future fix that will allow the fdnotifier to keep persistent entries, as opposed to requiring constant updates. This change is intended to be a no-op in terms of function. The only exception is signalfd, where this mask was abused. To handle this case, the operation of signalfd changed to allow one layer of indirection. PiperOrigin-RevId: 409702998 --- pkg/sentry/devices/tundev/tundev.go | 4 +- pkg/sentry/fs/dev/net_tun.go | 4 +- pkg/sentry/fs/fdpipe/pipe.go | 4 +- pkg/sentry/fs/file.go | 4 +- pkg/sentry/fs/file_overlay.go | 6 +- pkg/sentry/fs/host/file.go | 4 +- pkg/sentry/fs/host/wait_test.go | 4 +- pkg/sentry/fs/lock/lock.go | 4 +- pkg/sentry/fs/timerfd/timerfd.go | 4 +- pkg/sentry/fs/tty/master.go | 4 +- pkg/sentry/fs/tty/replica.go | 4 +- pkg/sentry/fsimpl/devpts/devpts_test.go | 14 +- pkg/sentry/fsimpl/devpts/master.go | 4 +- pkg/sentry/fsimpl/devpts/replica.go | 4 +- pkg/sentry/fsimpl/eventfd/eventfd.go | 4 +- pkg/sentry/fsimpl/eventfd/eventfd_test.go | 4 +- pkg/sentry/fsimpl/fuse/dev.go | 4 +- pkg/sentry/fsimpl/fuse/dev_test.go | 4 +- pkg/sentry/fsimpl/gofer/special_file.go | 6 +- pkg/sentry/fsimpl/host/host.go | 4 +- pkg/sentry/fsimpl/mqfs/queue.go | 4 +- pkg/sentry/fsimpl/overlay/regular_file.go | 18 +- pkg/sentry/fsimpl/signalfd/signalfd.go | 42 ++-- pkg/sentry/fsimpl/timerfd/timerfd.go | 4 +- pkg/sentry/kernel/epoll/epoll.go | 22 ++- pkg/sentry/kernel/eventfd/eventfd.go | 4 +- pkg/sentry/kernel/eventfd/eventfd_test.go | 4 +- pkg/sentry/kernel/fasync/fasync.go | 13 +- pkg/sentry/kernel/mq/mq.go | 23 +-- pkg/sentry/kernel/msgqueue/msgqueue.go | 8 +- pkg/sentry/kernel/pipe/pipe_test.go | 8 +- pkg/sentry/kernel/pipe/vfs.go | 4 +- pkg/sentry/kernel/signalfd/signalfd.go | 46 +++-- pkg/sentry/kernel/task_exit.go | 4 +- pkg/sentry/kernel/task_signals.go | 4 +- pkg/sentry/kernel/time/time.go | 6 +- pkg/sentry/socket/hostinet/socket.go | 20 +- pkg/sentry/socket/hostinet/socket_vfs2.go | 4 +- pkg/sentry/socket/netlink/socket.go | 8 +- pkg/sentry/socket/netlink/socket_vfs2.go | 4 +- pkg/sentry/socket/netstack/netstack.go | 20 +- pkg/sentry/socket/netstack/netstack_vfs2.go | 4 +- pkg/sentry/socket/unix/transport/unix.go | 4 +- pkg/sentry/socket/unix/unix.go | 16 +- pkg/sentry/socket/unix/unix_vfs2.go | 8 +- pkg/sentry/syscalls/epoll.go | 4 +- pkg/sentry/syscalls/linux/sys_poll.go | 4 +- pkg/sentry/syscalls/linux/sys_read.go | 8 +- pkg/sentry/syscalls/linux/sys_splice.go | 16 +- pkg/sentry/syscalls/linux/sys_write.go | 8 +- pkg/sentry/syscalls/linux/vfs2/epoll.go | 4 +- pkg/sentry/syscalls/linux/vfs2/poll.go | 4 +- pkg/sentry/syscalls/linux/vfs2/read_write.go | 16 +- pkg/sentry/syscalls/linux/vfs2/splice.go | 8 +- pkg/sentry/vfs/epoll.go | 19 +- pkg/sentry/vfs/file_description.go | 4 +- pkg/sentry/vfs/file_description_impl_util.go | 2 +- pkg/sentry/vfs/inotify.go | 4 +- pkg/sentry/vfs/save_restore.go | 2 +- pkg/syncevent/broadcaster_test.go | 16 +- pkg/tcpip/adapters/gonet/gonet.go | 21 +- pkg/tcpip/adapters/gonet/gonet_test.go | 7 +- pkg/tcpip/network/ip_test.go | 4 +- pkg/tcpip/network/ipv4/ipv4_test.go | 4 +- pkg/tcpip/network/ipv6/ipv6_test.go | 12 +- pkg/tcpip/sample/tun_tcp_connect/main.go | 7 +- pkg/tcpip/sample/tun_tcp_echo/main.go | 9 +- pkg/tcpip/stack/ndp_test.go | 12 +- pkg/tcpip/stack/transport_demuxer_test.go | 4 +- pkg/tcpip/tests/integration/forward_test.go | 8 +- pkg/tcpip/tests/integration/iptables_test.go | 8 +- .../tests/integration/link_resolution_test.go | 80 ++++---- pkg/tcpip/tests/integration/loopback_test.go | 4 +- .../integration/multicast_broadcast_test.go | 4 +- pkg/tcpip/tests/integration/route_test.go | 12 +- pkg/tcpip/transport/tcp/dual_stack_test.go | 20 +- pkg/tcpip/transport/tcp/endpoint.go | 4 +- pkg/tcpip/transport/tcp/tcp_test.go | 181 +++++++++--------- pkg/tcpip/transport/tcp/tcp_timestamp_test.go | 8 +- .../transport/tcp/testing/context/context.go | 12 +- pkg/tcpip/transport/udp/udp_test.go | 4 +- pkg/waiter/waiter.go | 89 ++++++--- pkg/waiter/waiter_test.go | 37 +--- 83 files changed, 534 insertions(+), 532 deletions(-) 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.