From e47be0cfc06d7eb04e4252b5fba7242ea038652a Mon Sep 17 00:00:00 2001 From: Nicolas Lacasse Date: Wed, 1 Jun 2022 09:54:52 -0700 Subject: [PATCH] Move Send/RecvNotify calls outside of CopyIn/Out, due to lock order. We were calling Send/RecvNotify inside of endpoint.Send/RecvMsg, which is called as part of CopyIn/Out with mm.activeMu held. This leads to lock order violation because those Send/RecvNotify needs TaskSet.mu to send a signal to a task. This CL changes endpoint.Send/RecvMsg to return a notification callback which can be plumbed to the caller of CopyIn/Out and called without mm.activeMu held. Updated the lock documentation in mm.go to order TaskSet.mu > mm.activeMu. PiperOrigin-RevId: 452319981 --- pkg/sentry/mm/mm.go | 12 +++---- pkg/sentry/socket/netlink/socket_vfs2.go | 9 ++++-- pkg/sentry/socket/unix/io.go | 21 ++++++++++-- .../socket/unix/transport/connectioned.go | 4 +-- .../socket/unix/transport/connectionless.go | 9 +++--- pkg/sentry/socket/unix/transport/unix.go | 32 ++++++++++++------- pkg/sentry/socket/unix/unix.go | 30 ++++++++++++++--- pkg/sentry/socket/unix/unix_vfs2.go | 19 +++++++++-- 8 files changed, 100 insertions(+), 36 deletions(-) diff --git a/pkg/sentry/mm/mm.go b/pkg/sentry/mm/mm.go index c8682681e..c906c1b92 100644 --- a/pkg/sentry/mm/mm.go +++ b/pkg/sentry/mm/mm.go @@ -21,14 +21,14 @@ // mm.MemoryManager.metadataMu // mm.MemoryManager.mappingMu // Locks taken by memmap.Mappable methods other than Translate -// mm.MemoryManager.activeMu -// Locks taken by memmap.Mappable.Translate -// mm.privateRefs.mu -// platform.AddressSpace locks -// memmap.File locks +// kernel.TaskSet.mu +// mm.MemoryManager.activeMu +// Locks taken by memmap.Mappable.Translate +// mm.privateRefs.mu +// platform.AddressSpace locks +// memmap.File locks // mm.aioManager.mu // mm.AIOContext.mu -// kernel.TaskSet.mu // // Only mm.MemoryManager.Fork is permitted to lock mm.MemoryManager.activeMu in // multiple mm.MemoryManagers, as it does so in a well-defined order (forked diff --git a/pkg/sentry/socket/netlink/socket_vfs2.go b/pkg/sentry/socket/netlink/socket_vfs2.go index b746337fe..59583ec33 100644 --- a/pkg/sentry/socket/netlink/socket_vfs2.go +++ b/pkg/sentry/socket/netlink/socket_vfs2.go @@ -132,9 +132,14 @@ func (s *SocketVFS2) Read(ctx context.Context, dst usermem.IOSequence, opts vfs. if dst.NumBytes() == 0 { return 0, nil } - return dst.CopyOutFrom(ctx, &unix.EndpointReader{ + r := unix.EndpointReader{ Endpoint: s.ep, - }) + } + n, err := dst.CopyOutFrom(ctx, &r) + if r.Notify != nil { + r.Notify() + } + return n, err } // PWrite implements vfs.FileDescriptionImpl. diff --git a/pkg/sentry/socket/unix/io.go b/pkg/sentry/socket/unix/io.go index 129949990..619a4c64d 100644 --- a/pkg/sentry/socket/unix/io.go +++ b/pkg/sentry/socket/unix/io.go @@ -35,12 +35,18 @@ type EndpointWriter struct { // To is the endpoint to send to. May be nil. To transport.BoundEndpoint + + // Notify is the receiver.SendNotify notification callback that is set + // by WriteFromBlocks and should be called without mm.activeMu held + // (i.e. after CopyOut completes). + Notify func() } // WriteFromBlocks implements safemem.Writer.WriteFromBlocks. func (w *EndpointWriter) WriteFromBlocks(srcs safemem.BlockSeq) (uint64, error) { return safemem.FromVecWriterFunc{func(bufs [][]byte) (int64, error) { - n, err := w.Endpoint.SendMsg(w.Ctx, bufs, w.Control, w.To) + n, notify, err := w.Endpoint.SendMsg(w.Ctx, bufs, w.Control, w.To) + w.Notify = notify if err != nil { return int64(n), err.ToError() } @@ -81,15 +87,23 @@ type EndpointReader struct { // ControlTrunc indicates that SCM_RIGHTS FDs were discarded based on // the value of NumRights. ControlTrunc bool + + // Notify is the ConnectedEndpoint.RecvNotify callback that is set by + // ReadToBlocks and should be called without mm.activeMu held (i.e. + // after CopyIn completes). + Notify func() } // Truncate calls RecvMsg on the endpoint without writing to a destination. func (r *EndpointReader) Truncate() error { // Ignore bytes read since it will always be zero. - _, ms, c, ct, err := r.Endpoint.RecvMsg(r.Ctx, [][]byte{}, r.Creds, r.NumRights, r.Peek, r.From) + _, ms, c, ct, notify, err := r.Endpoint.RecvMsg(r.Ctx, [][]byte{}, r.Creds, r.NumRights, r.Peek, r.From) r.Control = c r.ControlTrunc = ct r.MsgSize = ms + if notify != nil { + notify() + } if err != nil { return err.ToError() } @@ -99,10 +113,11 @@ func (r *EndpointReader) Truncate() error { // ReadToBlocks implements safemem.Reader.ReadToBlocks. func (r *EndpointReader) ReadToBlocks(dsts safemem.BlockSeq) (uint64, error) { return safemem.FromVecReaderFunc{func(bufs [][]byte) (int64, error) { - n, ms, c, ct, err := r.Endpoint.RecvMsg(r.Ctx, bufs, r.Creds, r.NumRights, r.Peek, r.From) + n, ms, c, ct, notify, err := r.Endpoint.RecvMsg(r.Ctx, bufs, r.Creds, r.NumRights, r.Peek, r.From) r.Control = c r.ControlTrunc = ct r.MsgSize = ms + r.Notify = notify if err != nil { return int64(n), err.ToError() } diff --git a/pkg/sentry/socket/unix/transport/connectioned.go b/pkg/sentry/socket/unix/transport/connectioned.go index 88abc146a..ca8c7a61e 100644 --- a/pkg/sentry/socket/unix/transport/connectioned.go +++ b/pkg/sentry/socket/unix/transport/connectioned.go @@ -538,11 +538,11 @@ func (e *connectionedEndpoint) Bind(addr tcpip.FullAddress) *syserr.Error { // SendMsg writes data and a control message to the endpoint's peer. // This method does not block if the data cannot be written. -func (e *connectionedEndpoint) SendMsg(ctx context.Context, data [][]byte, c ControlMessages, to BoundEndpoint) (int64, *syserr.Error) { +func (e *connectionedEndpoint) SendMsg(ctx context.Context, data [][]byte, c ControlMessages, to BoundEndpoint) (int64, func(), *syserr.Error) { // Stream sockets do not support specifying the endpoint. Seqpacket // sockets ignore the passed endpoint. if e.stype == linux.SOCK_STREAM && to != nil { - return 0, syserr.ErrNotSupported + return 0, nil, syserr.ErrNotSupported } return e.baseEndpoint.SendMsg(ctx, data, c, to) } diff --git a/pkg/sentry/socket/unix/transport/connectionless.go b/pkg/sentry/socket/unix/transport/connectionless.go index 80929c0a9..6df52623c 100644 --- a/pkg/sentry/socket/unix/transport/connectionless.go +++ b/pkg/sentry/socket/unix/transport/connectionless.go @@ -103,14 +103,14 @@ func (e *connectionlessEndpoint) UnidirectionalConnect(ctx context.Context) (Con // SendMsg writes data and a control message to the specified endpoint. // This method does not block if the data cannot be written. -func (e *connectionlessEndpoint) SendMsg(ctx context.Context, data [][]byte, c ControlMessages, to BoundEndpoint) (int64, *syserr.Error) { +func (e *connectionlessEndpoint) SendMsg(ctx context.Context, data [][]byte, c ControlMessages, to BoundEndpoint) (int64, func(), *syserr.Error) { if to == nil { return e.baseEndpoint.SendMsg(ctx, data, c, nil) } connected, err := to.UnidirectionalConnect(ctx) if err != nil { - return 0, syserr.ErrInvalidEndpointState + return 0, nil, syserr.ErrInvalidEndpointState } defer connected.Release(ctx) @@ -118,11 +118,12 @@ func (e *connectionlessEndpoint) SendMsg(ctx context.Context, data [][]byte, c C n, notify, err := connected.Send(ctx, data, c, tcpip.FullAddress{Addr: tcpip.Address(e.path)}) e.Unlock() + var notifyFn func() if notify { - connected.SendNotify() + notifyFn = connected.SendNotify } - return n, err + return n, notifyFn, err } // Type implements Endpoint.Type. diff --git a/pkg/sentry/socket/unix/transport/unix.go b/pkg/sentry/socket/unix/transport/unix.go index 912605335..87bb7eb35 100644 --- a/pkg/sentry/socket/unix/transport/unix.go +++ b/pkg/sentry/socket/unix/transport/unix.go @@ -130,13 +130,19 @@ type Endpoint interface { // CMTruncated indicates that the numRights hint was used to receive fewer // than the total available SCM_RIGHTS FDs. Additional truncation may be // required by the caller. - RecvMsg(ctx context.Context, data [][]byte, creds bool, numRights int, peek bool, addr *tcpip.FullAddress) (recvLen, msgLen int64, cm ControlMessages, CMTruncated bool, err *syserr.Error) + // + // If set, notify is a callback that should be called after RecvMesg + // completes without mm.activeMu held. + RecvMsg(ctx context.Context, data [][]byte, creds bool, numRights int, peek bool, addr *tcpip.FullAddress) (recvLen, msgLen int64, cm ControlMessages, CMTruncated bool, notify func(), err *syserr.Error) // SendMsg writes data and a control message to the endpoint's peer. // This method does not block if the data cannot be written. // // SendMsg does not take ownership of any of its arguments on error. - SendMsg(context.Context, [][]byte, ControlMessages, BoundEndpoint) (int64, *syserr.Error) + // + // If set, notify is a callback that should be called after RecvMesg + // completes without mm.activeMu held. + SendMsg(context.Context, [][]byte, ControlMessages, BoundEndpoint) (int64, func(), *syserr.Error) // Connect connects this endpoint directly to another. // @@ -826,53 +832,55 @@ func (e *baseEndpoint) Connected() bool { } // RecvMsg reads data and a control message from the endpoint. -func (e *baseEndpoint) RecvMsg(ctx context.Context, data [][]byte, creds bool, numRights int, peek bool, addr *tcpip.FullAddress) (int64, int64, ControlMessages, bool, *syserr.Error) { +func (e *baseEndpoint) RecvMsg(ctx context.Context, data [][]byte, creds bool, numRights int, peek bool, addr *tcpip.FullAddress) (int64, int64, ControlMessages, bool, func(), *syserr.Error) { e.Lock() receiver := e.receiver if receiver == nil { e.Unlock() - return 0, 0, ControlMessages{}, false, syserr.ErrNotConnected + return 0, 0, ControlMessages{}, false, nil, syserr.ErrNotConnected } recvLen, msgLen, cms, cmt, a, notify, err := receiver.Recv(ctx, data, creds, numRights, peek) e.Unlock() if err != nil { - return 0, 0, ControlMessages{}, false, err + return 0, 0, ControlMessages{}, false, nil, err } + var notifyFn func() if notify { - receiver.RecvNotify() + notifyFn = receiver.RecvNotify } if addr != nil { *addr = a } - return recvLen, msgLen, cms, cmt, nil + return recvLen, msgLen, cms, cmt, notifyFn, nil } // SendMsg writes data and a control message to the endpoint's peer. // This method does not block if the data cannot be written. -func (e *baseEndpoint) SendMsg(ctx context.Context, data [][]byte, c ControlMessages, to BoundEndpoint) (int64, *syserr.Error) { +func (e *baseEndpoint) SendMsg(ctx context.Context, data [][]byte, c ControlMessages, to BoundEndpoint) (int64, func(), *syserr.Error) { e.Lock() if !e.Connected() { e.Unlock() - return 0, syserr.ErrNotConnected + return 0, nil, syserr.ErrNotConnected } if to != nil { e.Unlock() - return 0, syserr.ErrAlreadyConnected + return 0, nil, syserr.ErrAlreadyConnected } connected := e.connected n, notify, err := connected.Send(ctx, data, c, tcpip.FullAddress{Addr: tcpip.Address(e.path)}) e.Unlock() + var notifyFn func() if notify { - connected.SendNotify() + notifyFn = connected.SendNotify } - return n, err + return n, notifyFn, err } // SetSockOpt sets a socket option. diff --git a/pkg/sentry/socket/unix/unix.go b/pkg/sentry/socket/unix/unix.go index ff726e62d..aaccc8ac7 100644 --- a/pkg/sentry/socket/unix/unix.go +++ b/pkg/sentry/socket/unix/unix.go @@ -460,16 +460,25 @@ func (s *SocketOperations) Write(ctx context.Context, _ *fs.File, src usermem.IO ctrl := control.New(t, s.ep, nil) if src.NumBytes() == 0 { - nInt, err := s.ep.SendMsg(ctx, [][]byte{}, ctrl, nil) + nInt, notify, err := s.ep.SendMsg(ctx, [][]byte{}, ctrl, nil) + if notify != nil { + notify() + } return int64(nInt), err.ToError() } - return src.CopyInTo(ctx, &EndpointWriter{ + w := &EndpointWriter{ Ctx: ctx, Endpoint: s.ep, Control: ctrl, To: nil, - }) + } + + n, err := src.CopyInTo(ctx, w) + if w.Notify != nil { + w.Notify() + } + return n, err } // SendMsg implements the linux syscall sendmsg(2) for unix sockets backed by @@ -505,6 +514,9 @@ func (s *socketOpsCommon) SendMsg(t *kernel.Task, src usermem.IOSequence, to []b } n, err := src.CopyInTo(t, &w) + if w.Notify != nil { + w.Notify() + } if err != linuxerr.ErrWouldBlock || flags&linux.MSG_DONTWAIT != 0 { return int(n), syserr.FromError(err) } @@ -524,6 +536,9 @@ func (s *socketOpsCommon) SendMsg(t *kernel.Task, src usermem.IOSequence, to []b src = src.DropFirst64(n) n, err = src.CopyInTo(t, &w) + if w.Notify != nil { + w.Notify() + } total += n if err != linuxerr.ErrWouldBlock { break @@ -596,6 +611,9 @@ func (s *SocketOperations) Read(ctx context.Context, _ *fs.File, dst usermem.IOS From: nil, } n, err := dst.CopyOutFrom(ctx, r) + if r.Notify != nil { + r.Notify() + } // Drop control messages. r.Control.Release(ctx) return n, err @@ -641,7 +659,11 @@ func (s *socketOpsCommon) RecvMsg(t *kernel.Task, dst usermem.IOSequence, flags } doRead := func() (int64, error) { - return dst.CopyOutFrom(t, &r) + n, err := dst.CopyOutFrom(t, &r) + if r.Notify != nil { + r.Notify() + } + return n, err } // If MSG_TRUNC is set with a zero byte destination then we still need diff --git a/pkg/sentry/socket/unix/unix_vfs2.go b/pkg/sentry/socket/unix/unix_vfs2.go index 66fbad03b..556e535d1 100644 --- a/pkg/sentry/socket/unix/unix_vfs2.go +++ b/pkg/sentry/socket/unix/unix_vfs2.go @@ -287,6 +287,9 @@ func (s *SocketVFS2) Read(ctx context.Context, dst usermem.IOSequence, opts vfs. From: nil, } n, err := dst.CopyOutFrom(ctx, r) + if r.Notify != nil { + r.Notify() + } // Drop control messages. r.Control.Release(ctx) return n, err @@ -309,16 +312,26 @@ func (s *SocketVFS2) Write(ctx context.Context, src usermem.IOSequence, opts vfs ctrl := control.New(t, s.ep, nil) if src.NumBytes() == 0 { - nInt, err := s.ep.SendMsg(ctx, [][]byte{}, ctrl, nil) + nInt, notify, err := s.ep.SendMsg(ctx, [][]byte{}, ctrl, nil) + if notify != nil { + notify() + } return int64(nInt), err.ToError() } - return src.CopyInTo(ctx, &EndpointWriter{ + w := &EndpointWriter{ Ctx: ctx, Endpoint: s.ep, Control: ctrl, To: nil, - }) + } + + n, err := src.CopyInTo(ctx, w) + if w.Notify != nil { + w.Notify() + } + return n, err + } // Readiness implements waiter.Waitable.Readiness.