Convert Recv and RecvMsg arguments/output into structs.

...they were getting out of control.

The two interface methods now take the same RecvArgs inputs, and return
RecvOutput.

I left `data [][]byte` as an explicit argument since it's technically an
argument (passed into Recv method), but acts like an output in that data is
written to it.

PiperOrigin-RevId: 578625002
This commit is contained in:
Nicolas Lacasse
2023-11-01 13:37:10 -07:00
committed by gVisor bot
parent ea234b5da4
commit 9c0d595c8f
4 changed files with 144 additions and 102 deletions
+26 -15
View File
@@ -77,8 +77,8 @@ type EndpointReader struct {
// sockets, it is the amount read.
MsgSize int64
// From, if not nil, will be set with the address read from.
From *transport.Address
// From will be set with the address read from.
From transport.Address
// Control contains the received control messages.
Control transport.ControlMessages
@@ -99,12 +99,17 @@ 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, 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
args := transport.RecvArgs{
Creds: r.Creds,
NumRights: r.NumRights,
Peek: r.Peek,
}
out, notify, err := r.Endpoint.RecvMsg(r.Ctx, [][]byte{}, args)
r.MsgSize = out.MsgLen
r.Control = out.Control
r.ControlTrunc = out.ControlTrunc
r.UnusedRights = out.UnusedRights
r.From = out.Source
if notify != nil {
notify()
}
@@ -117,15 +122,21 @@ 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, 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
args := transport.RecvArgs{
Creds: r.Creds,
NumRights: r.NumRights,
Peek: r.Peek,
}
out, notify, err := r.Endpoint.RecvMsg(r.Ctx, bufs, args)
r.MsgSize = out.MsgLen
r.Control = out.Control
r.ControlTrunc = out.ControlTrunc
r.UnusedRights = out.UnusedRights
r.From = out.Source
r.Notify = notify
if err != nil {
return int64(n), err.ToError()
return int64(out.RecvLen), err.ToError()
}
return int64(n), nil
return int64(out.RecvLen), nil
}}.ReadToBlocks(dsts)
}
+18 -12
View File
@@ -229,50 +229,56 @@ 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, []RightsControlMessage, bool, Address, bool, *syserr.Error) {
func (c *HostConnectedEndpoint) Recv(ctx context.Context, data [][]byte, args RecvArgs) (RecvOutput, bool, *syserr.Error) {
c.mu.RLock()
defer c.mu.RUnlock()
var cm unet.ControlMessage
if numRights > 0 {
cm.EnableFDs(int(numRights))
if args.NumRights > 0 {
cm.EnableFDs(int(args.NumRights))
}
// N.B. Unix sockets don't have a receive buffer, the send buffer
// serves both purposes.
rl, ml, cl, cTrunc, err := fdReadVec(c.fd, data, []byte(cm), peek, c.RecvMaxQueueSize())
if rl > 0 && err != nil {
out := RecvOutput{Source: Address{Addr: c.addr}}
var err error
var controlLen uint64
out.RecvLen, out.MsgLen, controlLen, out.ControlTrunc, err = fdReadVec(c.fd, data, []byte(cm), args.Peek, c.RecvMaxQueueSize())
if out.RecvLen > 0 && err != nil {
// We got some data, so all we need to do on error is return
// the data that we got. Short reads are fine, no need to
// block.
err = nil
}
if err != nil {
return 0, 0, ControlMessages{}, nil, false, Address{}, false, syserr.FromError(err)
return RecvOutput{}, false, syserr.FromError(err)
}
// There is no need for the callee to call RecvNotify because fdReadVec uses
// the host's recvmsg(2) and the host kernel's queue.
// Trim the control data if we received less than the full amount.
if cl < uint64(len(cm)) {
cm = cm[:cl]
if controlLen < uint64(len(cm)) {
cm = cm[:controlLen]
}
// Avoid extra allocations in the case where there isn't any control data.
if len(cm) == 0 {
return rl, ml, ControlMessages{}, nil, cTrunc, Address{Addr: c.addr}, false, nil
return out, false, nil
}
fds, err := cm.ExtractFDs()
if err != nil {
return 0, 0, ControlMessages{}, nil, false, Address{}, false, syserr.FromError(err)
return RecvOutput{}, false, syserr.FromError(err)
}
if len(fds) == 0 {
return rl, ml, ControlMessages{}, nil, cTrunc, Address{Addr: c.addr}, false, nil
return out, false, nil
}
return rl, ml, ControlMessages{Rights: &SCMRights{fds}}, nil, cTrunc, Address{Addr: c.addr}, false, nil
out.Control = ControlMessages{
Rights: &SCMRights{fds},
}
return out, false, nil
}
// RecvNotify implements Receiver.RecvNotify.
+96 -67
View File
@@ -87,6 +87,56 @@ func (c *ControlMessages) Release(ctx context.Context) {
*c = ControlMessages{}
}
// RecvArgs are the arguments to Endpoint.RecvMsg and Receiver.Recv.
type RecvArgs struct {
// Creds indicates if credential control messages are requested by the
// caller. This is useful for determining if control messages can be
// coalesced. Creds is a hint and can be safely ignored by the
// implementation if no coalescing is possible. It is fine to return
// credential control messages when none were requested or to not
// return credential control messages when they were requested.
Creds bool
// NumRights is the number of SCM_RIGHTS FDs requested by the caller.
// This is useful if one must allocate a buffer to receive a SCM_RIGHTS
// message or determine if control messages can be coalesced. numRights
// is a hint and can be safely ignored by the implementation if the
// number of available SCM_RIGHTS FDs is known and no coalescing is
// possible. It is fine for the returned number of SCM_RIGHTS FDs to be
// either higher or lower than the requested number.
NumRights int
// If Peek is true, no data should be consumed from the Endpoint. Any and
// all data returned from a peek should be available in the next call to
// Recv or RecvMsg.
Peek bool
}
// RecvOutput is the output from Endpoint.RecvMsg and Receiver.Recv.
type RecvOutput struct {
// RecvLen is the number of bytes copied into RecvArgs.Data.
RecvLen int64
// MsgLen is the length of the read message consumed for datagram Endpoints.
// MsgLen is always the same as RecvLen for stream Endpoints.
MsgLen int64
// Source is the source address we received from.
Source Address
// Control is the ControlMessages read.
Control ControlMessages
// ControlTrunc 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.
ControlTrunc bool
// UnusedRights is a slice of unused RightsControlMessage which should
// be Release()d.
UnusedRights []RightsControlMessage
}
// Endpoint is the interface implemented by Unix transport protocol
// implementations that expose functionality like sendmsg, recvmsg, connect,
// etc. to Unix socket implementations.
@@ -101,40 +151,8 @@ type Endpoint interface {
// RecvMsg reads data and a control message from the endpoint. This method
// does not block if there is no data pending.
//
// creds indicates if credential control messages are requested by the
// caller. This is useful for determining if control messages can be
// coalesced. creds is a hint and can be safely ignored by the
// implementation if no coalescing is possible. It is fine to return
// credential control messages when none were requested or to not return
// credential control messages when they were requested.
//
// numRights is the number of SCM_RIGHTS FDs requested by the caller. This
// is useful if one must allocate a buffer to receive a SCM_RIGHTS message
// or determine if control messages can be coalesced. numRights is a hint
// and can be safely ignored by the implementation if the number of
// available SCM_RIGHTS FDs is known and no coalescing is possible. It is
// fine for the returned number of SCM_RIGHTS FDs to be either higher or
// lower than the requested number.
//
// If peek is true, no data should be consumed from the Endpoint. Any and
// all data returned from a peek should be available in the next call to
// RecvMsg.
//
// recvLen is the number of bytes copied into data.
//
// 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
// 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 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)
// The returned callback should be called if not nil.
RecvMsg(ctx context.Context, data [][]byte, args RecvArgs) (RecvOutput, func(), *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,10 +362,8 @@ func (m *message) Truncate(n int64) {
type Receiver interface {
// Recv receives a single message. This method does not block.
//
// 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, unusedRights []RightsControlMessage, cmTruncated bool, source Address, notify bool, err *syserr.Error)
Recv(ctx context.Context, data [][]byte, args RecvArgs) (out RecvOutput, notify bool, err *syserr.Error)
// RecvNotify notifies the Receiver of a successful Recv. This must not be
// called while holding any endpoint locks.
@@ -394,17 +410,17 @@ 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, []RightsControlMessage, bool, Address, bool, *syserr.Error) {
func (q *queueReceiver) Recv(ctx context.Context, data [][]byte, args RecvArgs) (RecvOutput, bool, *syserr.Error) {
var m *message
var notify bool
var err *syserr.Error
if peek {
if args.Peek {
m, err = q.readQueue.Peek()
} else {
m, notify, err = q.readQueue.Dequeue()
}
if err != nil {
return 0, 0, ControlMessages{}, nil, false, Address{}, false, err
return RecvOutput{}, false, err
}
src := []byte(m.Data)
var copied int64
@@ -413,7 +429,13 @@ 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, nil, false, m.Address, notify, nil
out := RecvOutput{
RecvLen: copied,
MsgLen: int64(len(m.Data)),
Control: m.Control,
Source: m.Address,
}
return out, notify, nil
}
// RecvNotify implements Receiver.RecvNotify.
@@ -506,12 +528,11 @@ 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, []RightsControlMessage, bool, Address, bool, *syserr.Error) {
func (q *streamQueueReceiver) Recv(ctx context.Context, data [][]byte, args RecvArgs) (RecvOutput, 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 {
@@ -520,7 +541,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{}, unusedRights, false, Address{}, false, err
return RecvOutput{}, false, err
}
notify = n
q.buffer = []byte(m.Data)
@@ -529,14 +550,20 @@ func (q *streamQueueReceiver) Recv(ctx context.Context, data [][]byte, wantCreds
}
var copied int64
if peek {
if args.Peek {
// Don't consume control message if we are peeking.
c := q.control.Clone()
// Don't consume data since we are peeking.
copied, _, _ = vecCopy(data, q.buffer)
return copied, copied, c, unusedRights, false, q.addr, notify, nil
out := RecvOutput{
RecvLen: copied,
MsgLen: copied,
Control: c,
Source: q.addr,
}
return out, notify, nil
}
// Consume data and control message since we are not peeking.
@@ -547,15 +574,16 @@ func (q *streamQueueReceiver) Recv(ctx context.Context, data [][]byte, wantCreds
// Remove rights from q.control and leave behind just the creds.
q.control.Rights = nil
if !wantCreds {
if !args.Creds {
c.Credentials = nil
}
var cmTruncated bool
if c.Rights != nil && numRights == 0 {
unusedRights = append(unusedRights, c.Rights)
var out RecvOutput
if c.Rights != nil && args.NumRights == 0 {
// We won't use these rights.
out.UnusedRights = append(out.UnusedRights, c.Rights)
c.Rights = nil
cmTruncated = true
out.ControlTrunc = true
}
haveRights := c.Rights != nil
@@ -578,7 +606,7 @@ func (q *streamQueueReceiver) Recv(ctx context.Context, data [][]byte, wantCreds
q.control = m.Control
q.addr = m.Address
if wantCreds {
if args.Creds {
if (q.control.Credentials == nil) != (c.Credentials == nil) {
// One message has credentials, the other does not.
break
@@ -590,7 +618,7 @@ func (q *streamQueueReceiver) Recv(ctx context.Context, data [][]byte, wantCreds
}
}
if numRights != 0 && c.Rights != nil && q.control.Rights != nil {
if args.NumRights != 0 && c.Rights != nil && q.control.Rights != nil {
// Both messages have rights.
break
}
@@ -606,9 +634,9 @@ func (q *streamQueueReceiver) Recv(ctx context.Context, data [][]byte, wantCreds
if q.control.Rights != nil {
// Consume rights.
if numRights == 0 {
cmTruncated = true
unusedRights = append(unusedRights, q.control.Rights)
if args.NumRights == 0 {
out.ControlTrunc = true
out.UnusedRights = append(out.UnusedRights, q.control.Rights)
} else {
c.Rights = q.control.Rights
haveRights = true
@@ -616,7 +644,12 @@ func (q *streamQueueReceiver) Recv(ctx context.Context, data [][]byte, wantCreds
q.control.Rights = nil
}
}
return copied, copied, c, unusedRights, cmTruncated, q.addr, notify, nil
out.MsgLen = copied
out.RecvLen = copied
out.Source = q.addr
out.Control = c
return out, notify, nil
}
// Release implements Receiver.Release.
@@ -863,29 +896,25 @@ 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, []RightsControlMessage, bool, func(), *syserr.Error) {
func (e *baseEndpoint) RecvMsg(ctx context.Context, data [][]byte, args RecvArgs) (RecvOutput, func(), *syserr.Error) {
e.Lock()
receiver := e.receiver
e.Unlock()
if receiver == nil {
return 0, 0, ControlMessages{}, nil, false, nil, syserr.ErrNotConnected
return RecvOutput{}, nil, syserr.ErrNotConnected
}
recvLen, msgLen, cms, unusedRights, cmt, a, notify, err := receiver.Recv(ctx, data, creds, numRights, peek)
out, notify, err := receiver.Recv(ctx, data, args)
if err != nil {
return 0, 0, ControlMessages{}, unusedRights, false, nil, err
return RecvOutput{}, nil, err
}
var notifyFn func()
if notify {
notifyFn = receiver.RecvNotify
return out, receiver.RecvNotify, nil
}
if addr != nil {
*addr = a
}
return recvLen, msgLen, cms, unusedRights, cmt, notifyFn, nil
return out, nil, nil
}
// SendMsg writes data and a control message to the endpoint's peer.
+4 -8
View File
@@ -302,7 +302,6 @@ func (s *Socket) Read(ctx context.Context, dst usermem.IOSequence, opts vfs.Read
Endpoint: s.ep,
NumRights: 0,
Peek: false,
From: nil,
}
n, err := dst.CopyOutFrom(ctx, r)
if r.Notify != nil {
@@ -741,9 +740,6 @@ func (s *Socket) RecvMsg(t *kernel.Task, dst usermem.IOSequence, flags int, have
NumRights: numRights,
Peek: peek,
}
if senderRequested {
r.From = &transport.Address{}
}
doRead := func() (int64, error) {
n, err := dst.CopyOutFrom(t, &r)
@@ -778,8 +774,8 @@ func (s *Socket) RecvMsg(t *kernel.Task, dst usermem.IOSequence, flags int, have
if n, err := doRead(); err != linuxerr.ErrWouldBlock || dontWait {
var from linux.SockAddr
var fromLen uint32
if r.From != nil && len([]byte(r.From.Addr)) != 0 {
from, fromLen = convertAddress(*r.From)
if senderRequested && len([]byte(r.From.Addr)) != 0 {
from, fromLen = convertAddress(r.From)
}
if r.ControlTrunc {
@@ -813,8 +809,8 @@ func (s *Socket) RecvMsg(t *kernel.Task, dst usermem.IOSequence, flags int, have
if n, err := doRead(); err != linuxerr.ErrWouldBlock {
var from linux.SockAddr
var fromLen uint32
if r.From != nil {
from, fromLen = convertAddress(*r.From)
if senderRequested {
from, fromLen = convertAddress(r.From)
}
if r.ControlTrunc {