From e1e7edcc1b676018f6c577fe8f0767fb91a6326f Mon Sep 17 00:00:00 2001 From: Nicolas Lacasse Date: Wed, 1 Nov 2023 11:32:45 -0700 Subject: [PATCH] Don't release unused RightsControlMessage in Recv. This Recv method is called with mm.activeRWMutex in mm.CopyOutFrom, which can lead to lock order violatations because Release()ing a RightsControlMessage can trigger filesystem operations (like inotify) which may need filesystem locks. Instead, we plumb these "unusedRights" messages allll the way up to the EndpointReader, where they can be safely released outside of mm.activeRWMutex. PiperOrigin-RevId: 578584323 --- pkg/sentry/socket/unix/io.go | 10 ++++- pkg/sentry/socket/unix/transport/host.go | 12 +++--- pkg/sentry/socket/unix/transport/unix.go | 49 +++++++++++------------- pkg/sentry/socket/unix/unix.go | 11 ++++++ 4 files changed, 47 insertions(+), 35 deletions(-) diff --git a/pkg/sentry/socket/unix/io.go b/pkg/sentry/socket/unix/io.go index 9cb477de7..22fd935c7 100644 --- a/pkg/sentry/socket/unix/io.go +++ b/pkg/sentry/socket/unix/io.go @@ -83,6 +83,10 @@ type EndpointReader struct { // Control contains the received control messages. Control transport.ControlMessages + // UnusedRights is a slice of unused RightsControlMessage that must be + // Release()d before this EndpointReader is discarded. + UnusedRights []transport.RightsControlMessage + // ControlTrunc indicates that SCM_RIGHTS FDs were discarded based on // the value of NumRights. ControlTrunc bool @@ -96,8 +100,9 @@ type EndpointReader struct { // 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, notify, err := r.Endpoint.RecvMsg(r.Ctx, [][]byte{}, r.Creds, r.NumRights, r.Peek, r.From) + _, ms, c, unusedRights, ct, notify, err := r.Endpoint.RecvMsg(r.Ctx, [][]byte{}, r.Creds, r.NumRights, r.Peek, r.From) r.Control = c + r.UnusedRights = unusedRights r.ControlTrunc = ct r.MsgSize = ms if notify != nil { @@ -112,8 +117,9 @@ 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, notify, err := r.Endpoint.RecvMsg(r.Ctx, bufs, r.Creds, r.NumRights, r.Peek, r.From) + n, ms, c, unusedRights, ct, notify, err := r.Endpoint.RecvMsg(r.Ctx, bufs, r.Creds, r.NumRights, r.Peek, r.From) r.Control = c + r.UnusedRights = unusedRights r.ControlTrunc = ct r.MsgSize = ms r.Notify = notify diff --git a/pkg/sentry/socket/unix/transport/host.go b/pkg/sentry/socket/unix/transport/host.go index 69be4a6f9..79e5a991e 100644 --- a/pkg/sentry/socket/unix/transport/host.go +++ b/pkg/sentry/socket/unix/transport/host.go @@ -229,7 +229,7 @@ func (c *HostConnectedEndpoint) EventUpdate() error { } // Recv implements Receiver.Recv. -func (c *HostConnectedEndpoint) Recv(ctx context.Context, data [][]byte, creds bool, numRights int, peek bool) (int64, int64, ControlMessages, bool, Address, bool, *syserr.Error) { +func (c *HostConnectedEndpoint) Recv(ctx context.Context, data [][]byte, creds bool, numRights int, peek bool) (int64, int64, ControlMessages, []RightsControlMessage, bool, Address, bool, *syserr.Error) { c.mu.RLock() defer c.mu.RUnlock() @@ -248,7 +248,7 @@ func (c *HostConnectedEndpoint) Recv(ctx context.Context, data [][]byte, creds b err = nil } if err != nil { - return 0, 0, ControlMessages{}, false, Address{}, false, syserr.FromError(err) + return 0, 0, ControlMessages{}, nil, false, Address{}, false, syserr.FromError(err) } // There is no need for the callee to call RecvNotify because fdReadVec uses @@ -261,18 +261,18 @@ func (c *HostConnectedEndpoint) Recv(ctx context.Context, data [][]byte, creds b // Avoid extra allocations in the case where there isn't any control data. if len(cm) == 0 { - return rl, ml, ControlMessages{}, cTrunc, Address{Addr: c.addr}, false, nil + return rl, ml, ControlMessages{}, nil, cTrunc, Address{Addr: c.addr}, false, nil } fds, err := cm.ExtractFDs() if err != nil { - return 0, 0, ControlMessages{}, false, Address{}, false, syserr.FromError(err) + return 0, 0, ControlMessages{}, nil, false, Address{}, false, syserr.FromError(err) } if len(fds) == 0 { - return rl, ml, ControlMessages{}, cTrunc, Address{Addr: c.addr}, false, nil + return rl, ml, ControlMessages{}, nil, cTrunc, Address{Addr: c.addr}, false, nil } - return rl, ml, ControlMessages{Rights: &SCMRights{fds}}, cTrunc, Address{Addr: c.addr}, false, nil + return rl, ml, ControlMessages{Rights: &SCMRights{fds}}, nil, cTrunc, Address{Addr: c.addr}, false, nil } // RecvNotify implements Receiver.RecvNotify. diff --git a/pkg/sentry/socket/unix/transport/unix.go b/pkg/sentry/socket/unix/transport/unix.go index 26528b91f..e5003c87a 100644 --- a/pkg/sentry/socket/unix/transport/unix.go +++ b/pkg/sentry/socket/unix/transport/unix.go @@ -125,13 +125,16 @@ type Endpoint interface { // msgLen is the length of the read message consumed for datagram Endpoints. // msgLen is always the same as recvLen for stream Endpoints. // - // CMTruncated indicates that the numRights hint was used to receive fewer + // 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. // + // unusedRights is a slice of unused RightsControlMessage which should + // be Release()d. + // // 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 *Address) (recvLen, msgLen int64, cm ControlMessages, CMTruncated bool, notify func(), err *syserr.Error) + // completes without mm.activeMu hed. + RecvMsg(ctx context.Context, data [][]byte, creds bool, numRights int, peek bool, addr *Address) (recvLen, msgLen int64, cm ControlMessages, unusedRights []RightsControlMessage, 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. @@ -344,7 +347,7 @@ type Receiver interface { // See Endpoint.RecvMsg for documentation on shared arguments. // // notify indicates if RecvNotify should be called. - Recv(ctx context.Context, data [][]byte, creds bool, numRights int, peek bool) (recvLen, msgLen int64, cm ControlMessages, CMTruncated bool, source Address, notify bool, err *syserr.Error) + Recv(ctx context.Context, data [][]byte, creds bool, numRights int, peek bool) (recvLen, msgLen int64, cm ControlMessages, unusedRights []RightsControlMessage, cmTruncated bool, source Address, notify bool, err *syserr.Error) // RecvNotify notifies the Receiver of a successful Recv. This must not be // called while holding any endpoint locks. @@ -391,7 +394,7 @@ type queueReceiver struct { } // Recv implements Receiver.Recv. -func (q *queueReceiver) Recv(ctx context.Context, data [][]byte, creds bool, numRights int, peek bool) (int64, int64, ControlMessages, bool, Address, bool, *syserr.Error) { +func (q *queueReceiver) Recv(ctx context.Context, data [][]byte, creds bool, numRights int, peek bool) (int64, int64, ControlMessages, []RightsControlMessage, bool, Address, bool, *syserr.Error) { var m *message var notify bool var err *syserr.Error @@ -401,7 +404,7 @@ func (q *queueReceiver) Recv(ctx context.Context, data [][]byte, creds bool, num m, notify, err = q.readQueue.Dequeue() } if err != nil { - return 0, 0, ControlMessages{}, false, Address{}, false, err + return 0, 0, ControlMessages{}, nil, false, Address{}, false, err } src := []byte(m.Data) var copied int64 @@ -410,7 +413,7 @@ func (q *queueReceiver) Recv(ctx context.Context, data [][]byte, creds bool, num copied += int64(n) src = src[n:] } - return copied, int64(len(m.Data)), m.Control, false, m.Address, notify, nil + return copied, int64(len(m.Data)), m.Control, nil, false, m.Address, notify, nil } // RecvNotify implements Receiver.RecvNotify. @@ -503,20 +506,12 @@ func (q *streamQueueReceiver) RecvMaxQueueSize() int64 { } // Recv implements Receiver.Recv. -func (q *streamQueueReceiver) Recv(ctx context.Context, data [][]byte, wantCreds bool, numRights int, peek bool) (int64, int64, ControlMessages, bool, Address, bool, *syserr.Error) { - // RightsControlMessages must be released without q.mu held. We do this in a - // defer to simplify control flow logic. - var rightsToRelease []RightsControlMessage - defer func() { - for _, rcm := range rightsToRelease { - rcm.Release(ctx) - } - }() - +func (q *streamQueueReceiver) Recv(ctx context.Context, data [][]byte, wantCreds bool, numRights int, peek bool) (int64, int64, ControlMessages, []RightsControlMessage, bool, Address, bool, *syserr.Error) { q.mu.Lock() defer q.mu.Unlock() var notify bool + var unusedRights []RightsControlMessage // If we have no data in the endpoint, we need to get some. if len(q.buffer) == 0 { @@ -525,7 +520,7 @@ func (q *streamQueueReceiver) Recv(ctx context.Context, data [][]byte, wantCreds // the next time Recv() is called. m, n, err := q.readQueue.Dequeue() if err != nil { - return 0, 0, ControlMessages{}, false, Address{}, false, err + return 0, 0, ControlMessages{}, unusedRights, false, Address{}, false, err } notify = n q.buffer = []byte(m.Data) @@ -541,7 +536,7 @@ func (q *streamQueueReceiver) Recv(ctx context.Context, data [][]byte, wantCreds // Don't consume data since we are peeking. copied, _, _ = vecCopy(data, q.buffer) - return copied, copied, c, false, q.addr, notify, nil + return copied, copied, c, unusedRights, false, q.addr, notify, nil } // Consume data and control message since we are not peeking. @@ -558,7 +553,7 @@ func (q *streamQueueReceiver) Recv(ctx context.Context, data [][]byte, wantCreds var cmTruncated bool if c.Rights != nil && numRights == 0 { - rightsToRelease = append(rightsToRelease, c.Rights) + unusedRights = append(unusedRights, c.Rights) c.Rights = nil cmTruncated = true } @@ -613,7 +608,7 @@ func (q *streamQueueReceiver) Recv(ctx context.Context, data [][]byte, wantCreds // Consume rights. if numRights == 0 { cmTruncated = true - rightsToRelease = append(rightsToRelease, q.control.Rights) + unusedRights = append(unusedRights, q.control.Rights) } else { c.Rights = q.control.Rights haveRights = true @@ -621,7 +616,7 @@ func (q *streamQueueReceiver) Recv(ctx context.Context, data [][]byte, wantCreds q.control.Rights = nil } } - return copied, copied, c, cmTruncated, q.addr, notify, nil + return copied, copied, c, unusedRights, cmTruncated, q.addr, notify, nil } // Release implements Receiver.Release. @@ -868,18 +863,18 @@ 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 *Address) (int64, int64, ControlMessages, bool, func(), *syserr.Error) { +func (e *baseEndpoint) RecvMsg(ctx context.Context, data [][]byte, creds bool, numRights int, peek bool, addr *Address) (int64, int64, ControlMessages, []RightsControlMessage, bool, func(), *syserr.Error) { e.Lock() receiver := e.receiver e.Unlock() if receiver == nil { - return 0, 0, ControlMessages{}, false, nil, syserr.ErrNotConnected + return 0, 0, ControlMessages{}, nil, false, nil, syserr.ErrNotConnected } - recvLen, msgLen, cms, cmt, a, notify, err := receiver.Recv(ctx, data, creds, numRights, peek) + recvLen, msgLen, cms, unusedRights, cmt, a, notify, err := receiver.Recv(ctx, data, creds, numRights, peek) if err != nil { - return 0, 0, ControlMessages{}, false, nil, err + return 0, 0, ControlMessages{}, unusedRights, false, nil, err } var notifyFn func() @@ -890,7 +885,7 @@ func (e *baseEndpoint) RecvMsg(ctx context.Context, data [][]byte, creds bool, n if addr != nil { *addr = a } - return recvLen, msgLen, cms, cmt, notifyFn, nil + return recvLen, msgLen, cms, unusedRights, cmt, notifyFn, nil } // SendMsg writes data and a control message to the endpoint's peer. diff --git a/pkg/sentry/socket/unix/unix.go b/pkg/sentry/socket/unix/unix.go index a212e8fce..ea4b5e9cd 100644 --- a/pkg/sentry/socket/unix/unix.go +++ b/pkg/sentry/socket/unix/unix.go @@ -308,6 +308,10 @@ func (s *Socket) Read(ctx context.Context, dst usermem.IOSequence, opts vfs.Read if r.Notify != nil { r.Notify() } + // Drop any unused rights messages. + for _, rm := range r.UnusedRights { + rm.Release(ctx) + } // Drop control messages. r.Control.Release(ctx) return n, err @@ -749,6 +753,13 @@ func (s *Socket) RecvMsg(t *kernel.Task, dst usermem.IOSequence, flags int, have return n, err } + // Drop any unused rights messages after reading. + defer func() { + for _, rm := range r.UnusedRights { + rm.Release(t) + } + }() + // If MSG_TRUNC is set with a zero byte destination then we still need // to read the message and discard it, or in the case where MSG_PEEK is // set, leave it be. In both cases the full message length must be