From e08f204299dfcd6a93fde73375933cfa5f017740 Mon Sep 17 00:00:00 2001 From: Andrei Vagin Date: Fri, 20 Jan 2023 17:08:27 -0800 Subject: [PATCH] inet: each socket has to hold a reference to its network namespace Otherwise a network namespace can be destroyed before sockets. Reported-by: syzbot+78dcf6a117cd41dcb84e@syzkaller.appspotmail.com PiperOrigin-RevId: 503552997 --- pkg/sentry/socket/netstack/netstack.go | 118 ++++++++++--------- pkg/sentry/socket/netstack/netstack_state.go | 4 +- 2 files changed, 63 insertions(+), 59 deletions(-) diff --git a/pkg/sentry/socket/netstack/netstack.go b/pkg/sentry/socket/netstack/netstack.go index 6db78c8c4..a032e0215 100644 --- a/pkg/sentry/socket/netstack/netstack.go +++ b/pkg/sentry/socket/netstack/netstack.go @@ -342,11 +342,11 @@ type commonEndpoint interface { SocketOptions() *tcpip.SocketOptions } -// Socket encapsulates all the state needed to represent a network stack +// sock encapsulates all the state needed to represent a network stack // endpoint in the kernel context. // // +stateify savable -type Socket struct { +type sock struct { vfsfd vfs.FileDescription vfs.FileDescriptionDefaultImpl vfs.DentryMetadataFileDescriptionImpl @@ -359,6 +359,8 @@ type Socket struct { skType linux.SockType protocol int + namespace *inet.Namespace + // readMu protects access to the below fields. readMu sync.Mutex `state:"nosave"` @@ -379,7 +381,7 @@ type Socket struct { sockOptInq bool } -var _ = socket.Socket(&Socket{}) +var _ = socket.Socket(&sock{}) // New creates a new endpoint socket. func New(t *kernel.Task, family int, skType linux.SockType, protocol int, queue *waiter.Queue, endpoint tcpip.Endpoint) (*vfs.FileDescription, *syserr.Error) { @@ -391,12 +393,14 @@ func New(t *kernel.Task, family int, skType linux.SockType, protocol int, queue d := sockfs.NewDentry(t, mnt) defer d.DecRef(t) - s := &Socket{ - Queue: queue, - family: family, - Endpoint: endpoint, - skType: skType, - protocol: protocol, + namespace := t.NetworkNamespace() + s := &sock{ + Queue: queue, + family: family, + Endpoint: endpoint, + skType: skType, + protocol: protocol, + namespace: namespace, } s.LockFD.Init(&vfs.FileLocks{}) vfsfd := &s.vfsfd @@ -407,11 +411,12 @@ func New(t *kernel.Task, family int, skType linux.SockType, protocol int, queue }); err != nil { return nil, syserr.FromError(err) } + namespace.IncRef() return vfsfd, nil } // Release implements vfs.FileDescriptionImpl.Release. -func (s *Socket) Release(ctx context.Context) { +func (s *sock) Release(ctx context.Context) { kernel.KernelFromContext(ctx).DeleteSocket(&s.vfsfd) e, ch := waiter.NewChannelEntry(waiter.EventHUp | waiter.EventErr) s.EventRegister(&e) @@ -421,31 +426,30 @@ func (s *Socket) Release(ctx context.Context) { // SO_LINGER option is valid only for TCP. For other socket types // return after endpoint close. - if family, skType, _ := s.Type(); skType != linux.SOCK_STREAM || (family != linux.AF_INET && family != linux.AF_INET6) { - return - } - - v := s.Endpoint.SocketOptions().GetLinger() - // The case for zero timeout is handled in tcp endpoint close function. - // Close is blocked until either: - // 1. The endpoint state is not in any of the states: FIN-WAIT1, - // CLOSING and LAST_ACK. - // 2. Timeout is reached. - if v.Enabled && v.Timeout != 0 { - t := kernel.TaskFromContext(ctx) - start := t.Kernel().MonotonicClock().Now() - deadline := start.Add(v.Timeout) - _ = t.BlockWithDeadline(ch, true, deadline) + if family, skType, _ := s.Type(); skType == linux.SOCK_STREAM && (family == linux.AF_INET || family == linux.AF_INET6) { + v := s.Endpoint.SocketOptions().GetLinger() + // The case for zero timeout is handled in tcp endpoint close function. + // Close is blocked until either: + // 1. The endpoint state is not in any of the states: FIN-WAIT1, + // CLOSING and LAST_ACK. + // 2. Timeout is reached. + if v.Enabled && v.Timeout != 0 { + t := kernel.TaskFromContext(ctx) + start := t.Kernel().MonotonicClock().Now() + deadline := start.Add(v.Timeout) + _ = t.BlockWithDeadline(ch, true, deadline) + } } + s.namespace.DecRef() } // Epollable implements FileDescriptionImpl.Epollable. -func (s *Socket) Epollable() bool { +func (s *sock) Epollable() bool { return true } // Read implements vfs.FileDescriptionImpl. -func (s *Socket) Read(ctx context.Context, dst usermem.IOSequence, opts vfs.ReadOptions) (int64, error) { +func (s *sock) Read(ctx context.Context, dst usermem.IOSequence, opts vfs.ReadOptions) (int64, error) { // All flags other than RWF_NOWAIT should be ignored. // TODO(gvisor.dev/issue/2601): Support RWF_NOWAIT. if opts.Flags != 0 { @@ -466,7 +470,7 @@ func (s *Socket) Read(ctx context.Context, dst usermem.IOSequence, opts vfs.Read } // Write implements vfs.FileDescriptionImpl. -func (s *Socket) Write(ctx context.Context, src usermem.IOSequence, opts vfs.WriteOptions) (int64, error) { +func (s *sock) Write(ctx context.Context, src usermem.IOSequence, opts vfs.WriteOptions) (int64, error) { // All flags other than RWF_NOWAIT should be ignored. // TODO(gvisor.dev/issue/2601): Support RWF_NOWAIT. if opts.Flags != 0 { @@ -491,7 +495,7 @@ func (s *Socket) Write(ctx context.Context, src usermem.IOSequence, opts vfs.Wri // Accept implements the linux syscall accept(2) for sockets backed by // tcpip.Endpoint. -func (s *Socket) Accept(t *kernel.Task, peerRequested bool, flags int, blocking bool) (int32, linux.SockAddr, uint32, *syserr.Error) { +func (s *sock) Accept(t *kernel.Task, peerRequested bool, flags int, blocking bool) (int32, linux.SockAddr, uint32, *syserr.Error) { // Issue the accept request to get the new endpoint. var peerAddr *tcpip.FullAddress if peerRequested { @@ -538,7 +542,7 @@ func (s *Socket) Accept(t *kernel.Task, peerRequested bool, flags int, blocking // GetSockOpt implements the linux syscall getsockopt(2) for sockets backed by // tcpip.Endpoint. -func (s *Socket) GetSockOpt(t *kernel.Task, level, name int, outPtr hostarch.Addr, outLen int) (marshal.Marshallable, *syserr.Error) { +func (s *sock) GetSockOpt(t *kernel.Task, level, name int, outPtr hostarch.Addr, outLen int) (marshal.Marshallable, *syserr.Error) { // TODO(b/78348848): Unlike other socket options, SO_TIMESTAMP is // implemented specifically for netstack.Socket rather than // commonEndpoint. commonEndpoint should be extended to support socket @@ -574,7 +578,7 @@ func (s *Socket) GetSockOpt(t *kernel.Task, level, name int, outPtr hostarch.Add // SetSockOpt implements the linux syscall setsockopt(2) for sockets backed by // tcpip.Endpoint. -func (s *Socket) SetSockOpt(t *kernel.Task, level int, name int, optVal []byte) *syserr.Error { +func (s *sock) SetSockOpt(t *kernel.Task, level int, name int, optVal []byte) *syserr.Error { // TODO(b/78348848): Unlike other socket options, SO_TIMESTAMP is // implemented specifically for netstack.Socket rather than // commonEndpoint. commonEndpoint should be extended to support socket @@ -608,7 +612,7 @@ var sockAddrLinkSize = (*linux.SockAddrLink)(nil).SizeBytes() // minSockAddrLen returns the minimum length in bytes of a socket address for // the socket's family. -func (s *Socket) minSockAddrLen() int { +func (s *sock) minSockAddrLen() int { const addressFamilySize = 2 switch s.family { @@ -627,12 +631,12 @@ func (s *Socket) minSockAddrLen() int { } } -func (s *Socket) isPacketBased() bool { +func (s *sock) isPacketBased() bool { return s.skType == linux.SOCK_DGRAM || s.skType == linux.SOCK_SEQPACKET || s.skType == linux.SOCK_RDM || s.skType == linux.SOCK_RAW } // Readiness returns a mask of ready events for socket s. -func (s *Socket) Readiness(mask waiter.EventMask) waiter.EventMask { +func (s *sock) Readiness(mask waiter.EventMask) waiter.EventMask { return s.Endpoint.Readiness(mask) } @@ -641,7 +645,7 @@ func (s *Socket) Readiness(mask waiter.EventMask) waiter.EventMask { // // If exact is true, then the specified address family must be an exact match // with the socket's family. -func (s *Socket) checkFamily(family uint16, exact bool) bool { +func (s *sock) checkFamily(family uint16, exact bool) bool { if family == uint16(s.family) { return true } @@ -660,7 +664,7 @@ func (s *Socket) checkFamily(family uint16, exact bool) bool { // represented by the empty string. // // TODO(gvisor.dev/issue/1556): remove this function. -func (s *Socket) mapFamily(addr tcpip.FullAddress, family uint16) tcpip.FullAddress { +func (s *sock) mapFamily(addr tcpip.FullAddress, family uint16) tcpip.FullAddress { if len(addr.Addr) == 0 && s.family == linux.AF_INET6 && family == linux.AF_INET { addr.Addr = "\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\xff\xff\x00\x00\x00\x00" } @@ -669,7 +673,7 @@ func (s *Socket) mapFamily(addr tcpip.FullAddress, family uint16) tcpip.FullAddr // Connect implements the linux syscall connect(2) for sockets backed by // tpcip.Endpoint. -func (s *Socket) Connect(t *kernel.Task, sockaddr []byte, blocking bool) *syserr.Error { +func (s *sock) Connect(t *kernel.Task, sockaddr []byte, blocking bool) *syserr.Error { addr, family, err := socket.AddressAndFamily(sockaddr) if err != nil { return err @@ -724,7 +728,7 @@ func (s *Socket) Connect(t *kernel.Task, sockaddr []byte, blocking bool) *syserr // Bind implements the linux syscall bind(2) for sockets backed by // tcpip.Endpoint. -func (s *Socket) Bind(_ *kernel.Task, sockaddr []byte) *syserr.Error { +func (s *sock) Bind(_ *kernel.Task, sockaddr []byte) *syserr.Error { if len(sockaddr) < 2 { return syserr.ErrInvalidArgument } @@ -784,13 +788,13 @@ func (s *Socket) Bind(_ *kernel.Task, sockaddr []byte) *syserr.Error { // Listen implements the linux syscall listen(2) for sockets backed by // tcpip.Endpoint. -func (s *Socket) Listen(_ *kernel.Task, backlog int) *syserr.Error { +func (s *sock) Listen(_ *kernel.Task, backlog int) *syserr.Error { return syserr.TranslateNetstackError(s.Endpoint.Listen(backlog)) } // blockingAccept implements a blocking version of accept(2), that is, if no // connections are ready to be accept, it will block until one becomes ready. -func (s *Socket) blockingAccept(t *kernel.Task, peerAddr *tcpip.FullAddress) (tcpip.Endpoint, *waiter.Queue, *syserr.Error) { +func (s *sock) blockingAccept(t *kernel.Task, peerAddr *tcpip.FullAddress) (tcpip.Endpoint, *waiter.Queue, *syserr.Error) { // Register for notifications. e, ch := waiter.NewChannelEntry(waiter.ReadableEvents) s.EventRegister(&e) @@ -828,7 +832,7 @@ func ConvertShutdown(how int) (tcpip.ShutdownFlags, *syserr.Error) { // Shutdown implements the linux syscall shutdown(2) for sockets backed by // tcpip.Endpoint. -func (s *Socket) Shutdown(_ *kernel.Task, how int) *syserr.Error { +func (s *sock) Shutdown(_ *kernel.Task, how int) *syserr.Error { f, err := ConvertShutdown(how) if err != nil { return err @@ -2579,7 +2583,7 @@ func setSockOptIP(t *kernel.Task, s socket.Socket, ep commonEndpoint, name int, // GetSockName implements the linux syscall getsockname(2) for sockets backed by // tcpip.Endpoint. -func (s *Socket) GetSockName(*kernel.Task) (linux.SockAddr, uint32, *syserr.Error) { +func (s *sock) GetSockName(*kernel.Task) (linux.SockAddr, uint32, *syserr.Error) { addr, err := s.Endpoint.GetLocalAddress() if err != nil { return nil, 0, syserr.TranslateNetstackError(err) @@ -2591,7 +2595,7 @@ func (s *Socket) GetSockName(*kernel.Task) (linux.SockAddr, uint32, *syserr.Erro // GetPeerName implements the linux syscall getpeername(2) for sockets backed by // tcpip.Endpoint. -func (s *Socket) GetPeerName(*kernel.Task) (linux.SockAddr, uint32, *syserr.Error) { +func (s *sock) GetPeerName(*kernel.Task) (linux.SockAddr, uint32, *syserr.Error) { addr, err := s.Endpoint.GetRemoteAddress() if err != nil { return nil, 0, syserr.TranslateNetstackError(err) @@ -2601,7 +2605,7 @@ func (s *Socket) GetPeerName(*kernel.Task) (linux.SockAddr, uint32, *syserr.Erro return a, l, nil } -func (s *Socket) fillCmsgInq(cmsg *socket.ControlMessages) { +func (s *sock) fillCmsgInq(cmsg *socket.ControlMessages) { if !s.sockOptInq { return } @@ -2633,7 +2637,7 @@ func toLinuxPacketType(pktType tcpip.PacketType) uint8 { // nonBlockingRead issues a non-blocking read. // // TODO(b/78348848): Support timestamps for stream sockets. -func (s *Socket) nonBlockingRead(ctx context.Context, dst usermem.IOSequence, peek, trunc, senderRequested bool) (int, int, linux.SockAddr, uint32, socket.ControlMessages, *syserr.Error) { +func (s *sock) nonBlockingRead(ctx context.Context, dst usermem.IOSequence, peek, trunc, senderRequested bool) (int, int, linux.SockAddr, uint32, socket.ControlMessages, *syserr.Error) { isPacket := s.isPacketBased() readOptions := tcpip.ReadOptions{ @@ -2722,7 +2726,7 @@ func (s *Socket) nonBlockingRead(ctx context.Context, dst usermem.IOSequence, pe return res.Count, 0, nil, 0, cmsg, syserr.TranslateNetstackError(err) } -func (s *Socket) netstackToLinuxControlMessages(cm tcpip.ReceivableControlMessages) socket.ControlMessages { +func (s *sock) netstackToLinuxControlMessages(cm tcpip.ReceivableControlMessages) socket.ControlMessages { readCM := socket.NewIPControlMessages(s.family, cm) return socket.ControlMessages{ IP: socket.IPControlMessages{ @@ -2748,7 +2752,7 @@ func (s *Socket) netstackToLinuxControlMessages(cm tcpip.ReceivableControlMessag } } -func (s *Socket) linuxToNetstackControlMessages(cm socket.ControlMessages) tcpip.SendableControlMessages { +func (s *sock) linuxToNetstackControlMessages(cm socket.ControlMessages) tcpip.SendableControlMessages { return tcpip.SendableControlMessages{ HasTTL: cm.IP.HasTTL, TTL: uint8(cm.IP.TTL), @@ -2761,7 +2765,7 @@ func (s *Socket) linuxToNetstackControlMessages(cm socket.ControlMessages) tcpip // successfully writing packet data out to userspace. // // Precondition: s.readMu must be locked. -func (s *Socket) updateTimestamp(cm tcpip.ReceivableControlMessages) { +func (s *sock) updateTimestamp(cm tcpip.ReceivableControlMessages) { // Save the SIOCGSTAMP timestamp only if SO_TIMESTAMP is disabled. if !s.sockOptTimestamp { s.timestampValid = true @@ -2770,7 +2774,7 @@ func (s *Socket) updateTimestamp(cm tcpip.ReceivableControlMessages) { } // dequeueErr is analogous to net/core/skbuff.c:sock_dequeue_err_skb(). -func (s *Socket) dequeueErr() *tcpip.SockError { +func (s *sock) dequeueErr() *tcpip.SockError { so := s.Endpoint.SocketOptions() err := so.DequeueErr() if err == nil { @@ -2801,7 +2805,7 @@ func addrFamilyFromNetProto(net tcpip.NetworkProtocolNumber) int { // recvErr handles MSG_ERRQUEUE for recvmsg(2). // This is analogous to net/ipv4/ip_sockglue.c:ip_recv_error(). -func (s *Socket) recvErr(t *kernel.Task, dst usermem.IOSequence) (int, int, linux.SockAddr, uint32, socket.ControlMessages, *syserr.Error) { +func (s *sock) recvErr(t *kernel.Task, dst usermem.IOSequence) (int, int, linux.SockAddr, uint32, socket.ControlMessages, *syserr.Error) { sockErr := s.dequeueErr() if sockErr == nil { return 0, 0, nil, 0, socket.ControlMessages{}, syserr.ErrTryAgain @@ -2827,7 +2831,7 @@ func (s *Socket) recvErr(t *kernel.Task, dst usermem.IOSequence) (int, int, linu // RecvMsg implements the linux syscall recvmsg(2) for sockets backed by // tcpip.Endpoint. -func (s *Socket) RecvMsg(t *kernel.Task, dst usermem.IOSequence, flags int, haveDeadline bool, deadline ktime.Time, senderRequested bool, _ uint64) (n int, msgFlags int, senderAddr linux.SockAddr, senderAddrLen uint32, controlMessages socket.ControlMessages, err *syserr.Error) { +func (s *sock) RecvMsg(t *kernel.Task, dst usermem.IOSequence, flags int, haveDeadline bool, deadline ktime.Time, senderRequested bool, _ uint64) (n int, msgFlags int, senderAddr linux.SockAddr, senderAddrLen uint32, controlMessages socket.ControlMessages, err *syserr.Error) { if flags&linux.MSG_ERRQUEUE != 0 { return s.recvErr(t, dst) } @@ -2899,7 +2903,7 @@ func (s *Socket) RecvMsg(t *kernel.Task, dst usermem.IOSequence, flags int, have // SendMsg implements the linux syscall sendmsg(2) for sockets backed by // tcpip.Endpoint. -func (s *Socket) SendMsg(t *kernel.Task, src usermem.IOSequence, to []byte, flags int, haveDeadline bool, deadline ktime.Time, controlMessages socket.ControlMessages) (int, *syserr.Error) { +func (s *sock) SendMsg(t *kernel.Task, src usermem.IOSequence, to []byte, flags int, haveDeadline bool, deadline ktime.Time, controlMessages socket.ControlMessages) (int, *syserr.Error) { // Reject Unix control messages. if !controlMessages.Unix.Empty() { return 0, syserr.ErrInvalidArgument @@ -2972,7 +2976,7 @@ func (s *Socket) SendMsg(t *kernel.Task, src usermem.IOSequence, to []byte, flag } // Ioctl implements vfs.FileDescriptionImpl. -func (s *Socket) Ioctl(ctx context.Context, uio usermem.IO, args arch.SyscallArguments) (uintptr, error) { +func (s *sock) Ioctl(ctx context.Context, uio usermem.IO, args arch.SyscallArguments) (uintptr, error) { t := kernel.TaskFromContext(ctx) if t == nil { panic("ioctl(2) may only be called from a task goroutine") @@ -3329,7 +3333,7 @@ func isICMPSocket(skType linux.SockType, skProto int) bool { // State implements socket.Socket.State. State translates the internal state // returned by netstack to values defined by Linux. -func (s *Socket) State() uint32 { +func (s *sock) State() uint32 { if s.family != linux.AF_INET && s.family != linux.AF_INET6 { // States not implemented for this socket's family. return 0 @@ -3389,17 +3393,17 @@ func (s *Socket) State() uint32 { } // Type implements socket.Socket.Type. -func (s *Socket) Type() (family int, skType linux.SockType, protocol int) { +func (s *sock) Type() (family int, skType linux.SockType, protocol int) { return s.family, s.skType, s.protocol } // EventRegister implements waiter.Waitable. -func (s *Socket) EventRegister(e *waiter.Entry) error { +func (s *sock) EventRegister(e *waiter.Entry) error { s.Queue.EventRegister(e) return nil } // EventUnregister implements waiter.Waitable.EventUnregister. -func (s *Socket) EventUnregister(e *waiter.Entry) { +func (s *sock) EventUnregister(e *waiter.Entry) { s.Queue.EventUnregister(e) } diff --git a/pkg/sentry/socket/netstack/netstack_state.go b/pkg/sentry/socket/netstack/netstack_state.go index a3cf1fa92..d995cc679 100644 --- a/pkg/sentry/socket/netstack/netstack_state.go +++ b/pkg/sentry/socket/netstack/netstack_state.go @@ -18,13 +18,13 @@ import ( "time" ) -func (s *Socket) saveTimestamp() int64 { +func (s *sock) saveTimestamp() int64 { s.readMu.Lock() defer s.readMu.Unlock() return s.timestamp.UnixNano() } -func (s *Socket) loadTimestamp(nsec int64) { +func (s *sock) loadTimestamp(nsec int64) { s.readMu.Lock() defer s.readMu.Unlock() s.timestamp = time.Unix(0, nsec)