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