mirror of
https://github.com/netbirdio/gvisor.git
synced 2026-05-22 17:12:49 -07:00
Track and export socket state.
This is necessary for implementing network diagnostic interfaces like
/proc/net/{tcp,udp,unix} and sock_diag(7).
For pass-through endpoints such as hostinet, we obtain the socket
state from the backend. For netstack, we add explicit tracking of TCP
states.
PiperOrigin-RevId: 251934850
This commit is contained in:
@@ -200,6 +200,22 @@ const (
|
||||
SS_DISCONNECTING = 4 // In process of disconnecting.
|
||||
)
|
||||
|
||||
// TCP protocol states, from include/net/tcp_states.h.
|
||||
const (
|
||||
TCP_ESTABLISHED uint32 = iota + 1
|
||||
TCP_SYN_SENT
|
||||
TCP_SYN_RECV
|
||||
TCP_FIN_WAIT1
|
||||
TCP_FIN_WAIT2
|
||||
TCP_TIME_WAIT
|
||||
TCP_CLOSE
|
||||
TCP_CLOSE_WAIT
|
||||
TCP_LAST_ACK
|
||||
TCP_LISTEN
|
||||
TCP_CLOSING
|
||||
TCP_NEW_SYN_RECV
|
||||
)
|
||||
|
||||
// SockAddrMax is the maximum size of a struct sockaddr, from
|
||||
// uapi/linux/socket.h.
|
||||
const SockAddrMax = 128
|
||||
|
||||
@@ -240,24 +240,6 @@ func (n *netUnix) ReadSeqFileData(ctx context.Context, h seqfile.SeqHandle) ([]s
|
||||
}
|
||||
}
|
||||
|
||||
var sockState int
|
||||
switch sops.Endpoint().Type() {
|
||||
case linux.SOCK_DGRAM:
|
||||
sockState = linux.SS_CONNECTING
|
||||
// Unlike Linux, we don't have unbound connection-less sockets,
|
||||
// so no SS_DISCONNECTING.
|
||||
|
||||
case linux.SOCK_SEQPACKET:
|
||||
fallthrough
|
||||
case linux.SOCK_STREAM:
|
||||
// Connectioned.
|
||||
if sops.Endpoint().(transport.ConnectingEndpoint).Connected() {
|
||||
sockState = linux.SS_CONNECTED
|
||||
} else {
|
||||
sockState = linux.SS_UNCONNECTED
|
||||
}
|
||||
}
|
||||
|
||||
// In the socket entry below, the value for the 'Num' field requires
|
||||
// some consideration. Linux prints the address to the struct
|
||||
// unix_sock representing a socket in the kernel, but may redact the
|
||||
@@ -282,7 +264,7 @@ func (n *netUnix) ReadSeqFileData(ctx context.Context, h seqfile.SeqHandle) ([]s
|
||||
0, // Protocol, always 0 for UDS.
|
||||
sockFlags, // Flags.
|
||||
sops.Endpoint().Type(), // Type.
|
||||
sockState, // State.
|
||||
sops.State(), // State.
|
||||
sfile.InodeID(), // Inode.
|
||||
)
|
||||
|
||||
|
||||
@@ -52,6 +52,7 @@ import (
|
||||
"gvisor.googlesource.com/gvisor/pkg/tcpip"
|
||||
"gvisor.googlesource.com/gvisor/pkg/tcpip/buffer"
|
||||
"gvisor.googlesource.com/gvisor/pkg/tcpip/stack"
|
||||
"gvisor.googlesource.com/gvisor/pkg/tcpip/transport/tcp"
|
||||
"gvisor.googlesource.com/gvisor/pkg/waiter"
|
||||
)
|
||||
|
||||
@@ -2281,3 +2282,46 @@ func nicStateFlagsToLinux(f stack.NICStateFlags) uint32 {
|
||||
}
|
||||
return rv
|
||||
}
|
||||
|
||||
// State implements socket.Socket.State. State translates the internal state
|
||||
// returned by netstack to values defined by Linux.
|
||||
func (s *SocketOperations) State() uint32 {
|
||||
if s.family != linux.AF_INET && s.family != linux.AF_INET6 {
|
||||
// States not implemented for this socket's family.
|
||||
return 0
|
||||
}
|
||||
|
||||
if !s.isPacketBased() {
|
||||
// TCP socket.
|
||||
switch tcp.EndpointState(s.Endpoint.State()) {
|
||||
case tcp.StateEstablished:
|
||||
return linux.TCP_ESTABLISHED
|
||||
case tcp.StateSynSent:
|
||||
return linux.TCP_SYN_SENT
|
||||
case tcp.StateSynRecv:
|
||||
return linux.TCP_SYN_RECV
|
||||
case tcp.StateFinWait1:
|
||||
return linux.TCP_FIN_WAIT1
|
||||
case tcp.StateFinWait2:
|
||||
return linux.TCP_FIN_WAIT2
|
||||
case tcp.StateTimeWait:
|
||||
return linux.TCP_TIME_WAIT
|
||||
case tcp.StateClose, tcp.StateInitial, tcp.StateBound, tcp.StateConnecting, tcp.StateError:
|
||||
return linux.TCP_CLOSE
|
||||
case tcp.StateCloseWait:
|
||||
return linux.TCP_CLOSE_WAIT
|
||||
case tcp.StateLastAck:
|
||||
return linux.TCP_LAST_ACK
|
||||
case tcp.StateListen:
|
||||
return linux.TCP_LISTEN
|
||||
case tcp.StateClosing:
|
||||
return linux.TCP_CLOSING
|
||||
default:
|
||||
// Internal or unknown state.
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
// TODO(b/112063468): Export states for UDP, ICMP, and raw sockets.
|
||||
return 0
|
||||
}
|
||||
|
||||
@@ -19,7 +19,9 @@ import (
|
||||
"syscall"
|
||||
|
||||
"gvisor.googlesource.com/gvisor/pkg/abi/linux"
|
||||
"gvisor.googlesource.com/gvisor/pkg/binary"
|
||||
"gvisor.googlesource.com/gvisor/pkg/fdnotifier"
|
||||
"gvisor.googlesource.com/gvisor/pkg/log"
|
||||
"gvisor.googlesource.com/gvisor/pkg/sentry/context"
|
||||
"gvisor.googlesource.com/gvisor/pkg/sentry/fs"
|
||||
"gvisor.googlesource.com/gvisor/pkg/sentry/fs/fsutil"
|
||||
@@ -519,6 +521,28 @@ func translateIOSyscallError(err error) error {
|
||||
return err
|
||||
}
|
||||
|
||||
// State implements socket.Socket.State.
|
||||
func (s *socketOperations) State() uint32 {
|
||||
info := linux.TCPInfo{}
|
||||
buf, err := getsockopt(s.fd, syscall.SOL_TCP, syscall.TCP_INFO, linux.SizeOfTCPInfo)
|
||||
if err != nil {
|
||||
if err != syscall.ENOPROTOOPT {
|
||||
log.Warningf("Failed to get TCP socket info from %+v: %v", s, err)
|
||||
}
|
||||
// For non-TCP sockets, silently ignore the failure.
|
||||
return 0
|
||||
}
|
||||
if len(buf) != linux.SizeOfTCPInfo {
|
||||
// Unmarshal below will panic if getsockopt returns a buffer of
|
||||
// unexpected size.
|
||||
log.Warningf("Failed to get TCP socket info from %+v: getsockopt(2) returned %d bytes, expecting %d bytes.", s, len(buf), linux.SizeOfTCPInfo)
|
||||
return 0
|
||||
}
|
||||
|
||||
binary.Unmarshal(buf, usermem.ByteOrder, &info)
|
||||
return uint32(info.State)
|
||||
}
|
||||
|
||||
type socketProvider struct {
|
||||
family int
|
||||
}
|
||||
|
||||
@@ -616,3 +616,8 @@ func (s *Socket) Write(ctx context.Context, _ *fs.File, src usermem.IOSequence,
|
||||
n, err := s.sendMsg(ctx, src, nil, 0, socket.ControlMessages{})
|
||||
return int64(n), err.ToError()
|
||||
}
|
||||
|
||||
// State implements socket.Socket.State.
|
||||
func (s *Socket) State() uint32 {
|
||||
return s.ep.State()
|
||||
}
|
||||
|
||||
@@ -830,6 +830,12 @@ func (s *socketOperations) SendMsg(t *kernel.Task, src usermem.IOSequence, to []
|
||||
}
|
||||
}
|
||||
|
||||
// State implements socket.Socket.State.
|
||||
func (s *socketOperations) State() uint32 {
|
||||
// TODO(b/127845868): Define a new rpc to query the socket state.
|
||||
return 0
|
||||
}
|
||||
|
||||
type socketProvider struct {
|
||||
family int
|
||||
}
|
||||
|
||||
@@ -116,6 +116,10 @@ type Socket interface {
|
||||
// SendTimeout gets the current timeout (in ns) for send operations. Zero
|
||||
// means no timeout, and negative means DONTWAIT.
|
||||
SendTimeout() int64
|
||||
|
||||
// State returns the current state of the socket, as represented by Linux in
|
||||
// procfs. The returned state value is protocol-specific.
|
||||
State() uint32
|
||||
}
|
||||
|
||||
// Provider is the interface implemented by providers of sockets for specific
|
||||
|
||||
@@ -28,6 +28,7 @@ go_library(
|
||||
importpath = "gvisor.googlesource.com/gvisor/pkg/sentry/socket/unix/transport",
|
||||
visibility = ["//:sandbox"],
|
||||
deps = [
|
||||
"//pkg/abi/linux",
|
||||
"//pkg/ilist",
|
||||
"//pkg/refs",
|
||||
"//pkg/syserr",
|
||||
|
||||
@@ -17,6 +17,7 @@ package transport
|
||||
import (
|
||||
"sync"
|
||||
|
||||
"gvisor.googlesource.com/gvisor/pkg/abi/linux"
|
||||
"gvisor.googlesource.com/gvisor/pkg/syserr"
|
||||
"gvisor.googlesource.com/gvisor/pkg/tcpip"
|
||||
"gvisor.googlesource.com/gvisor/pkg/waiter"
|
||||
@@ -458,3 +459,11 @@ func (e *connectionedEndpoint) Readiness(mask waiter.EventMask) waiter.EventMask
|
||||
|
||||
return ready
|
||||
}
|
||||
|
||||
// State implements socket.Socket.State.
|
||||
func (e *connectionedEndpoint) State() uint32 {
|
||||
if e.Connected() {
|
||||
return linux.SS_CONNECTED
|
||||
}
|
||||
return linux.SS_UNCONNECTED
|
||||
}
|
||||
|
||||
@@ -15,6 +15,7 @@
|
||||
package transport
|
||||
|
||||
import (
|
||||
"gvisor.googlesource.com/gvisor/pkg/abi/linux"
|
||||
"gvisor.googlesource.com/gvisor/pkg/syserr"
|
||||
"gvisor.googlesource.com/gvisor/pkg/tcpip"
|
||||
"gvisor.googlesource.com/gvisor/pkg/waiter"
|
||||
@@ -194,3 +195,18 @@ func (e *connectionlessEndpoint) Readiness(mask waiter.EventMask) waiter.EventMa
|
||||
|
||||
return ready
|
||||
}
|
||||
|
||||
// State implements socket.Socket.State.
|
||||
func (e *connectionlessEndpoint) State() uint32 {
|
||||
e.Lock()
|
||||
defer e.Unlock()
|
||||
|
||||
switch {
|
||||
case e.isBound():
|
||||
return linux.SS_UNCONNECTED
|
||||
case e.Connected():
|
||||
return linux.SS_CONNECTING
|
||||
default:
|
||||
return linux.SS_DISCONNECTING
|
||||
}
|
||||
}
|
||||
|
||||
@@ -191,6 +191,10 @@ type Endpoint interface {
|
||||
// GetSockOpt gets a socket option. opt should be a pointer to one of the
|
||||
// tcpip.*Option types.
|
||||
GetSockOpt(opt interface{}) *tcpip.Error
|
||||
|
||||
// State returns the current state of the socket, as represented by Linux in
|
||||
// procfs.
|
||||
State() uint32
|
||||
}
|
||||
|
||||
// A Credentialer is a socket or endpoint that supports the SO_PASSCRED socket
|
||||
|
||||
@@ -596,6 +596,11 @@ func (s *SocketOperations) RecvMsg(t *kernel.Task, dst usermem.IOSequence, flags
|
||||
}
|
||||
}
|
||||
|
||||
// State implements socket.Socket.State.
|
||||
func (s *SocketOperations) State() uint32 {
|
||||
return s.ep.State()
|
||||
}
|
||||
|
||||
// provider is a unix domain socket provider.
|
||||
type provider struct{}
|
||||
|
||||
|
||||
@@ -188,6 +188,10 @@ func (f *fakeTransportEndpoint) HandleControlPacket(stack.TransportEndpointID, s
|
||||
f.proto.controlCount++
|
||||
}
|
||||
|
||||
func (f *fakeTransportEndpoint) State() uint32 {
|
||||
return 0
|
||||
}
|
||||
|
||||
type fakeTransportGoodOption bool
|
||||
|
||||
type fakeTransportBadOption bool
|
||||
|
||||
@@ -377,6 +377,10 @@ type Endpoint interface {
|
||||
// GetSockOpt gets a socket option. opt should be a pointer to one of the
|
||||
// *Option types.
|
||||
GetSockOpt(opt interface{}) *Error
|
||||
|
||||
// State returns a socket's lifecycle state. The returned value is
|
||||
// protocol-specific and is primarily used for diagnostics.
|
||||
State() uint32
|
||||
}
|
||||
|
||||
// WriteOptions contains options for Endpoint.Write.
|
||||
|
||||
@@ -33,6 +33,7 @@ go_library(
|
||||
"//pkg/tcpip/header",
|
||||
"//pkg/tcpip/stack",
|
||||
"//pkg/tcpip/transport/raw",
|
||||
"//pkg/tcpip/transport/tcp",
|
||||
"//pkg/waiter",
|
||||
],
|
||||
)
|
||||
|
||||
@@ -708,3 +708,9 @@ func (e *endpoint) HandlePacket(r *stack.Route, id stack.TransportEndpointID, vv
|
||||
// HandleControlPacket implements stack.TransportEndpoint.HandleControlPacket.
|
||||
func (e *endpoint) HandleControlPacket(id stack.TransportEndpointID, typ stack.ControlType, extra uint32, vv buffer.VectorisedView) {
|
||||
}
|
||||
|
||||
// State implements tcpip.Endpoint.State. The ICMP endpoint currently doesn't
|
||||
// expose internal socket state.
|
||||
func (e *endpoint) State() uint32 {
|
||||
return 0
|
||||
}
|
||||
|
||||
@@ -519,3 +519,8 @@ func (ep *endpoint) HandlePacket(route *stack.Route, netHeader buffer.View, vv b
|
||||
ep.waiterQueue.Notify(waiter.EventIn)
|
||||
}
|
||||
}
|
||||
|
||||
// State implements socket.Socket.State.
|
||||
func (ep *endpoint) State() uint32 {
|
||||
return 0
|
||||
}
|
||||
|
||||
@@ -226,7 +226,6 @@ func (l *listenContext) createConnectingEndpoint(s *segment, iss seqnum.Value, i
|
||||
}
|
||||
|
||||
n.isRegistered = true
|
||||
n.state = stateConnecting
|
||||
|
||||
// Create sender and receiver.
|
||||
//
|
||||
@@ -258,8 +257,9 @@ func (l *listenContext) createEndpointAndPerformHandshake(s *segment, opts *head
|
||||
ep.Close()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
ep.state = stateConnected
|
||||
ep.mu.Lock()
|
||||
ep.state = StateEstablished
|
||||
ep.mu.Unlock()
|
||||
|
||||
// Update the receive window scaling. We can't do it before the
|
||||
// handshake because it's possible that the peer doesn't support window
|
||||
@@ -276,7 +276,7 @@ func (e *endpoint) deliverAccepted(n *endpoint) {
|
||||
e.mu.RLock()
|
||||
state := e.state
|
||||
e.mu.RUnlock()
|
||||
if state == stateListen {
|
||||
if state == StateListen {
|
||||
e.acceptedChan <- n
|
||||
e.waiterQueue.Notify(waiter.EventIn)
|
||||
} else {
|
||||
@@ -406,7 +406,7 @@ func (e *endpoint) handleListenSegment(ctx *listenContext, s *segment) {
|
||||
n.tsOffset = 0
|
||||
|
||||
// Switch state to connected.
|
||||
n.state = stateConnected
|
||||
n.state = StateEstablished
|
||||
|
||||
// Do the delivery in a separate goroutine so
|
||||
// that we don't block the listen loop in case
|
||||
@@ -429,7 +429,7 @@ func (e *endpoint) protocolListenLoop(rcvWnd seqnum.Size) *tcpip.Error {
|
||||
// handleSynSegment() from attempting to queue new connections
|
||||
// to the endpoint.
|
||||
e.mu.Lock()
|
||||
e.state = stateClosed
|
||||
e.state = StateClose
|
||||
|
||||
// Do cleanup if needed.
|
||||
e.completeWorkerLocked()
|
||||
|
||||
@@ -151,6 +151,9 @@ func (h *handshake) resetToSynRcvd(iss seqnum.Value, irs seqnum.Value, opts *hea
|
||||
h.mss = opts.MSS
|
||||
h.sndWndScale = opts.WS
|
||||
h.listenEP = listenEP
|
||||
h.ep.mu.Lock()
|
||||
h.ep.state = StateSynRecv
|
||||
h.ep.mu.Unlock()
|
||||
}
|
||||
|
||||
// checkAck checks if the ACK number, if present, of a segment received during
|
||||
@@ -219,6 +222,9 @@ func (h *handshake) synSentState(s *segment) *tcpip.Error {
|
||||
// but resend our own SYN and wait for it to be acknowledged in the
|
||||
// SYN-RCVD state.
|
||||
h.state = handshakeSynRcvd
|
||||
h.ep.mu.Lock()
|
||||
h.ep.state = StateSynRecv
|
||||
h.ep.mu.Unlock()
|
||||
synOpts := header.TCPSynOptions{
|
||||
WS: h.rcvWndScale,
|
||||
TS: rcvSynOpts.TS,
|
||||
@@ -668,7 +674,7 @@ func (e *endpoint) makeOptions(sackBlocks []header.SACKBlock) []byte {
|
||||
// sendRaw sends a TCP segment to the endpoint's peer.
|
||||
func (e *endpoint) sendRaw(data buffer.VectorisedView, flags byte, seq, ack seqnum.Value, rcvWnd seqnum.Size) *tcpip.Error {
|
||||
var sackBlocks []header.SACKBlock
|
||||
if e.state == stateConnected && e.rcv.pendingBufSize > 0 && (flags&header.TCPFlagAck != 0) {
|
||||
if e.state == StateEstablished && e.rcv.pendingBufSize > 0 && (flags&header.TCPFlagAck != 0) {
|
||||
sackBlocks = e.sack.Blocks[:e.sack.NumBlocks]
|
||||
}
|
||||
options := e.makeOptions(sackBlocks)
|
||||
@@ -719,8 +725,7 @@ func (e *endpoint) handleClose() *tcpip.Error {
|
||||
// protocol goroutine.
|
||||
func (e *endpoint) resetConnectionLocked(err *tcpip.Error) {
|
||||
e.sendRaw(buffer.VectorisedView{}, header.TCPFlagAck|header.TCPFlagRst, e.snd.sndUna, e.rcv.rcvNxt, 0)
|
||||
|
||||
e.state = stateError
|
||||
e.state = StateError
|
||||
e.hardError = err
|
||||
}
|
||||
|
||||
@@ -876,14 +881,19 @@ func (e *endpoint) protocolMainLoop(handshake bool) *tcpip.Error {
|
||||
// handshake, and then inform potential waiters about its
|
||||
// completion.
|
||||
h := newHandshake(e, seqnum.Size(e.receiveBufferAvailable()))
|
||||
e.mu.Lock()
|
||||
h.ep.state = StateSynSent
|
||||
e.mu.Unlock()
|
||||
|
||||
if err := h.execute(); err != nil {
|
||||
e.lastErrorMu.Lock()
|
||||
e.lastError = err
|
||||
e.lastErrorMu.Unlock()
|
||||
|
||||
e.mu.Lock()
|
||||
e.state = stateError
|
||||
e.state = StateError
|
||||
e.hardError = err
|
||||
|
||||
// Lock released below.
|
||||
epilogue()
|
||||
|
||||
@@ -905,7 +915,7 @@ func (e *endpoint) protocolMainLoop(handshake bool) *tcpip.Error {
|
||||
|
||||
// Tell waiters that the endpoint is connected and writable.
|
||||
e.mu.Lock()
|
||||
e.state = stateConnected
|
||||
e.state = StateEstablished
|
||||
drained := e.drainDone != nil
|
||||
e.mu.Unlock()
|
||||
if drained {
|
||||
@@ -1005,7 +1015,7 @@ func (e *endpoint) protocolMainLoop(handshake bool) *tcpip.Error {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if e.state != stateError {
|
||||
if e.state != StateError {
|
||||
close(e.drainDone)
|
||||
<-e.undrain
|
||||
}
|
||||
@@ -1061,8 +1071,8 @@ func (e *endpoint) protocolMainLoop(handshake bool) *tcpip.Error {
|
||||
|
||||
// Mark endpoint as closed.
|
||||
e.mu.Lock()
|
||||
if e.state != stateError {
|
||||
e.state = stateClosed
|
||||
if e.state != StateError {
|
||||
e.state = StateClose
|
||||
}
|
||||
// Lock released below.
|
||||
epilogue()
|
||||
|
||||
@@ -32,18 +32,81 @@ import (
|
||||
"gvisor.googlesource.com/gvisor/pkg/waiter"
|
||||
)
|
||||
|
||||
type endpointState int
|
||||
// EndpointState represents the state of a TCP endpoint.
|
||||
type EndpointState uint32
|
||||
|
||||
// Endpoint states. Note that are represented in a netstack-specific manner and
|
||||
// may not be meaningful externally. Specifically, they need to be translated to
|
||||
// Linux's representation for these states if presented to userspace.
|
||||
const (
|
||||
stateInitial endpointState = iota
|
||||
stateBound
|
||||
stateListen
|
||||
stateConnecting
|
||||
stateConnected
|
||||
stateClosed
|
||||
stateError
|
||||
// Endpoint states internal to netstack. These map to the TCP state CLOSED.
|
||||
StateInitial EndpointState = iota
|
||||
StateBound
|
||||
StateConnecting // Connect() called, but the initial SYN hasn't been sent.
|
||||
StateError
|
||||
|
||||
// TCP protocol states.
|
||||
StateEstablished
|
||||
StateSynSent
|
||||
StateSynRecv
|
||||
StateFinWait1
|
||||
StateFinWait2
|
||||
StateTimeWait
|
||||
StateClose
|
||||
StateCloseWait
|
||||
StateLastAck
|
||||
StateListen
|
||||
StateClosing
|
||||
)
|
||||
|
||||
// connected is the set of states where an endpoint is connected to a peer.
|
||||
func (s EndpointState) connected() bool {
|
||||
switch s {
|
||||
case StateEstablished, StateFinWait1, StateFinWait2, StateTimeWait, StateCloseWait, StateLastAck, StateClosing:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// String implements fmt.Stringer.String.
|
||||
func (s EndpointState) String() string {
|
||||
switch s {
|
||||
case StateInitial:
|
||||
return "INITIAL"
|
||||
case StateBound:
|
||||
return "BOUND"
|
||||
case StateConnecting:
|
||||
return "CONNECTING"
|
||||
case StateError:
|
||||
return "ERROR"
|
||||
case StateEstablished:
|
||||
return "ESTABLISHED"
|
||||
case StateSynSent:
|
||||
return "SYN-SENT"
|
||||
case StateSynRecv:
|
||||
return "SYN-RCVD"
|
||||
case StateFinWait1:
|
||||
return "FIN-WAIT1"
|
||||
case StateFinWait2:
|
||||
return "FIN-WAIT2"
|
||||
case StateTimeWait:
|
||||
return "TIME-WAIT"
|
||||
case StateClose:
|
||||
return "CLOSED"
|
||||
case StateCloseWait:
|
||||
return "CLOSE-WAIT"
|
||||
case StateLastAck:
|
||||
return "LAST-ACK"
|
||||
case StateListen:
|
||||
return "LISTEN"
|
||||
case StateClosing:
|
||||
return "CLOSING"
|
||||
default:
|
||||
panic("unreachable")
|
||||
}
|
||||
}
|
||||
|
||||
// Reasons for notifying the protocol goroutine.
|
||||
const (
|
||||
notifyNonZeroReceiveWindow = 1 << iota
|
||||
@@ -108,10 +171,14 @@ type endpoint struct {
|
||||
rcvBufUsed int
|
||||
|
||||
// The following fields are protected by the mutex.
|
||||
mu sync.RWMutex `state:"nosave"`
|
||||
id stack.TransportEndpointID
|
||||
state endpointState `state:".(endpointState)"`
|
||||
isPortReserved bool `state:"manual"`
|
||||
mu sync.RWMutex `state:"nosave"`
|
||||
id stack.TransportEndpointID
|
||||
|
||||
// state endpointState `state:".(endpointState)"`
|
||||
// pState ProtocolState
|
||||
state EndpointState `state:".(EndpointState)"`
|
||||
|
||||
isPortReserved bool `state:"manual"`
|
||||
isRegistered bool
|
||||
boundNICID tcpip.NICID `state:"manual"`
|
||||
route stack.Route `state:"manual"`
|
||||
@@ -304,6 +371,7 @@ func newEndpoint(stack *stack.Stack, netProto tcpip.NetworkProtocolNumber, waite
|
||||
stack: stack,
|
||||
netProto: netProto,
|
||||
waiterQueue: waiterQueue,
|
||||
state: StateInitial,
|
||||
rcvBufSize: DefaultBufferSize,
|
||||
sndBufSize: DefaultBufferSize,
|
||||
sndMTU: int(math.MaxInt32),
|
||||
@@ -351,14 +419,14 @@ func (e *endpoint) Readiness(mask waiter.EventMask) waiter.EventMask {
|
||||
defer e.mu.RUnlock()
|
||||
|
||||
switch e.state {
|
||||
case stateInitial, stateBound, stateConnecting:
|
||||
case StateInitial, StateBound, StateConnecting, StateSynSent, StateSynRecv:
|
||||
// Ready for nothing.
|
||||
|
||||
case stateClosed, stateError:
|
||||
case StateClose, StateError:
|
||||
// Ready for anything.
|
||||
result = mask
|
||||
|
||||
case stateListen:
|
||||
case StateListen:
|
||||
// Check if there's anything in the accepted channel.
|
||||
if (mask & waiter.EventIn) != 0 {
|
||||
if len(e.acceptedChan) > 0 {
|
||||
@@ -366,7 +434,7 @@ func (e *endpoint) Readiness(mask waiter.EventMask) waiter.EventMask {
|
||||
}
|
||||
}
|
||||
|
||||
case stateConnected:
|
||||
case StateEstablished, StateFinWait1, StateFinWait2, StateTimeWait, StateCloseWait, StateLastAck, StateClosing:
|
||||
// Determine if the endpoint is writable if requested.
|
||||
if (mask & waiter.EventOut) != 0 {
|
||||
e.sndBufMu.Lock()
|
||||
@@ -427,7 +495,7 @@ func (e *endpoint) Close() {
|
||||
// are immediately available for reuse after Close() is called. If also
|
||||
// registered, we unregister as well otherwise the next user would fail
|
||||
// in Listen() when trying to register.
|
||||
if e.state == stateListen && e.isPortReserved {
|
||||
if e.state == StateListen && e.isPortReserved {
|
||||
if e.isRegistered {
|
||||
e.stack.UnregisterTransportEndpoint(e.boundNICID, e.effectiveNetProtos, ProtocolNumber, e.id, e)
|
||||
e.isRegistered = false
|
||||
@@ -487,15 +555,15 @@ func (e *endpoint) Read(*tcpip.FullAddress) (buffer.View, tcpip.ControlMessages,
|
||||
e.mu.RLock()
|
||||
// The endpoint can be read if it's connected, or if it's already closed
|
||||
// but has some pending unread data. Also note that a RST being received
|
||||
// would cause the state to become stateError so we should allow the
|
||||
// would cause the state to become StateError so we should allow the
|
||||
// reads to proceed before returning a ECONNRESET.
|
||||
e.rcvListMu.Lock()
|
||||
bufUsed := e.rcvBufUsed
|
||||
if s := e.state; s != stateConnected && s != stateClosed && bufUsed == 0 {
|
||||
if s := e.state; !s.connected() && s != StateClose && bufUsed == 0 {
|
||||
e.rcvListMu.Unlock()
|
||||
he := e.hardError
|
||||
e.mu.RUnlock()
|
||||
if s == stateError {
|
||||
if s == StateError {
|
||||
return buffer.View{}, tcpip.ControlMessages{}, he
|
||||
}
|
||||
return buffer.View{}, tcpip.ControlMessages{}, tcpip.ErrInvalidEndpointState
|
||||
@@ -511,7 +579,7 @@ func (e *endpoint) Read(*tcpip.FullAddress) (buffer.View, tcpip.ControlMessages,
|
||||
|
||||
func (e *endpoint) readLocked() (buffer.View, *tcpip.Error) {
|
||||
if e.rcvBufUsed == 0 {
|
||||
if e.rcvClosed || e.state != stateConnected {
|
||||
if e.rcvClosed || !e.state.connected() {
|
||||
return buffer.View{}, tcpip.ErrClosedForReceive
|
||||
}
|
||||
return buffer.View{}, tcpip.ErrWouldBlock
|
||||
@@ -547,9 +615,9 @@ func (e *endpoint) Write(p tcpip.Payload, opts tcpip.WriteOptions) (uintptr, <-c
|
||||
defer e.mu.RUnlock()
|
||||
|
||||
// The endpoint cannot be written to if it's not connected.
|
||||
if e.state != stateConnected {
|
||||
if !e.state.connected() {
|
||||
switch e.state {
|
||||
case stateError:
|
||||
case StateError:
|
||||
return 0, nil, e.hardError
|
||||
default:
|
||||
return 0, nil, tcpip.ErrClosedForSend
|
||||
@@ -612,8 +680,8 @@ func (e *endpoint) Peek(vec [][]byte) (uintptr, tcpip.ControlMessages, *tcpip.Er
|
||||
|
||||
// The endpoint can be read if it's connected, or if it's already closed
|
||||
// but has some pending unread data.
|
||||
if s := e.state; s != stateConnected && s != stateClosed {
|
||||
if s == stateError {
|
||||
if s := e.state; !s.connected() && s != StateClose {
|
||||
if s == StateError {
|
||||
return 0, tcpip.ControlMessages{}, e.hardError
|
||||
}
|
||||
return 0, tcpip.ControlMessages{}, tcpip.ErrInvalidEndpointState
|
||||
@@ -623,7 +691,7 @@ func (e *endpoint) Peek(vec [][]byte) (uintptr, tcpip.ControlMessages, *tcpip.Er
|
||||
defer e.rcvListMu.Unlock()
|
||||
|
||||
if e.rcvBufUsed == 0 {
|
||||
if e.rcvClosed || e.state != stateConnected {
|
||||
if e.rcvClosed || !e.state.connected() {
|
||||
return 0, tcpip.ControlMessages{}, tcpip.ErrClosedForReceive
|
||||
}
|
||||
return 0, tcpip.ControlMessages{}, tcpip.ErrWouldBlock
|
||||
@@ -789,7 +857,7 @@ func (e *endpoint) SetSockOpt(opt interface{}) *tcpip.Error {
|
||||
defer e.mu.Unlock()
|
||||
|
||||
// We only allow this to be set when we're in the initial state.
|
||||
if e.state != stateInitial {
|
||||
if e.state != StateInitial {
|
||||
return tcpip.ErrInvalidEndpointState
|
||||
}
|
||||
|
||||
@@ -841,7 +909,7 @@ func (e *endpoint) readyReceiveSize() (int, *tcpip.Error) {
|
||||
defer e.mu.RUnlock()
|
||||
|
||||
// The endpoint cannot be in listen state.
|
||||
if e.state == stateListen {
|
||||
if e.state == StateListen {
|
||||
return 0, tcpip.ErrInvalidEndpointState
|
||||
}
|
||||
|
||||
@@ -1057,7 +1125,7 @@ func (e *endpoint) connect(addr tcpip.FullAddress, handshake bool, run bool) (er
|
||||
|
||||
nicid := addr.NIC
|
||||
switch e.state {
|
||||
case stateBound:
|
||||
case StateBound:
|
||||
// If we're already bound to a NIC but the caller is requesting
|
||||
// that we use a different one now, we cannot proceed.
|
||||
if e.boundNICID == 0 {
|
||||
@@ -1070,16 +1138,16 @@ func (e *endpoint) connect(addr tcpip.FullAddress, handshake bool, run bool) (er
|
||||
|
||||
nicid = e.boundNICID
|
||||
|
||||
case stateInitial:
|
||||
// Nothing to do. We'll eventually fill-in the gaps in the ID
|
||||
// (if any) when we find a route.
|
||||
case StateInitial:
|
||||
// Nothing to do. We'll eventually fill-in the gaps in the ID (if any)
|
||||
// when we find a route.
|
||||
|
||||
case stateConnecting:
|
||||
// A connection request has already been issued but hasn't
|
||||
// completed yet.
|
||||
case StateConnecting, StateSynSent, StateSynRecv:
|
||||
// A connection request has already been issued but hasn't completed
|
||||
// yet.
|
||||
return tcpip.ErrAlreadyConnecting
|
||||
|
||||
case stateConnected:
|
||||
case StateEstablished:
|
||||
// The endpoint is already connected. If caller hasn't been notified yet, return success.
|
||||
if !e.isConnectNotified {
|
||||
e.isConnectNotified = true
|
||||
@@ -1088,7 +1156,7 @@ func (e *endpoint) connect(addr tcpip.FullAddress, handshake bool, run bool) (er
|
||||
// Otherwise return that it's already connected.
|
||||
return tcpip.ErrAlreadyConnected
|
||||
|
||||
case stateError:
|
||||
case StateError:
|
||||
return e.hardError
|
||||
|
||||
default:
|
||||
@@ -1154,7 +1222,7 @@ func (e *endpoint) connect(addr tcpip.FullAddress, handshake bool, run bool) (er
|
||||
}
|
||||
|
||||
e.isRegistered = true
|
||||
e.state = stateConnecting
|
||||
e.state = StateConnecting
|
||||
e.route = r.Clone()
|
||||
e.boundNICID = nicid
|
||||
e.effectiveNetProtos = netProtos
|
||||
@@ -1175,7 +1243,7 @@ func (e *endpoint) connect(addr tcpip.FullAddress, handshake bool, run bool) (er
|
||||
}
|
||||
e.segmentQueue.mu.Unlock()
|
||||
e.snd.updateMaxPayloadSize(int(e.route.MTU()), 0)
|
||||
e.state = stateConnected
|
||||
e.state = StateEstablished
|
||||
}
|
||||
|
||||
if run {
|
||||
@@ -1199,8 +1267,8 @@ func (e *endpoint) Shutdown(flags tcpip.ShutdownFlags) *tcpip.Error {
|
||||
defer e.mu.Unlock()
|
||||
e.shutdownFlags |= flags
|
||||
|
||||
switch e.state {
|
||||
case stateConnected:
|
||||
switch {
|
||||
case e.state.connected():
|
||||
// Close for read.
|
||||
if (e.shutdownFlags & tcpip.ShutdownRead) != 0 {
|
||||
// Mark read side as closed.
|
||||
@@ -1241,7 +1309,7 @@ func (e *endpoint) Shutdown(flags tcpip.ShutdownFlags) *tcpip.Error {
|
||||
e.sndCloseWaker.Assert()
|
||||
}
|
||||
|
||||
case stateListen:
|
||||
case e.state == StateListen:
|
||||
// Tell protocolListenLoop to stop.
|
||||
if flags&tcpip.ShutdownRead != 0 {
|
||||
e.notifyProtocolGoroutine(notifyClose)
|
||||
@@ -1269,7 +1337,7 @@ func (e *endpoint) Listen(backlog int) (err *tcpip.Error) {
|
||||
// When the endpoint shuts down, it sets workerCleanup to true, and from
|
||||
// that point onward, acceptedChan is the responsibility of the cleanup()
|
||||
// method (and should not be touched anywhere else, including here).
|
||||
if e.state == stateListen && !e.workerCleanup {
|
||||
if e.state == StateListen && !e.workerCleanup {
|
||||
// Adjust the size of the channel iff we can fix existing
|
||||
// pending connections into the new one.
|
||||
if len(e.acceptedChan) > backlog {
|
||||
@@ -1288,7 +1356,7 @@ func (e *endpoint) Listen(backlog int) (err *tcpip.Error) {
|
||||
}
|
||||
|
||||
// Endpoint must be bound before it can transition to listen mode.
|
||||
if e.state != stateBound {
|
||||
if e.state != StateBound {
|
||||
return tcpip.ErrInvalidEndpointState
|
||||
}
|
||||
|
||||
@@ -1298,7 +1366,7 @@ func (e *endpoint) Listen(backlog int) (err *tcpip.Error) {
|
||||
}
|
||||
|
||||
e.isRegistered = true
|
||||
e.state = stateListen
|
||||
e.state = StateListen
|
||||
if e.acceptedChan == nil {
|
||||
e.acceptedChan = make(chan *endpoint, backlog)
|
||||
}
|
||||
@@ -1325,7 +1393,7 @@ func (e *endpoint) Accept() (tcpip.Endpoint, *waiter.Queue, *tcpip.Error) {
|
||||
defer e.mu.RUnlock()
|
||||
|
||||
// Endpoint must be in listen state before it can accept connections.
|
||||
if e.state != stateListen {
|
||||
if e.state != StateListen {
|
||||
return nil, nil, tcpip.ErrInvalidEndpointState
|
||||
}
|
||||
|
||||
@@ -1353,7 +1421,7 @@ func (e *endpoint) Bind(addr tcpip.FullAddress) (err *tcpip.Error) {
|
||||
// Don't allow binding once endpoint is not in the initial state
|
||||
// anymore. This is because once the endpoint goes into a connected or
|
||||
// listen state, it is already bound.
|
||||
if e.state != stateInitial {
|
||||
if e.state != StateInitial {
|
||||
return tcpip.ErrAlreadyBound
|
||||
}
|
||||
|
||||
@@ -1408,7 +1476,7 @@ func (e *endpoint) Bind(addr tcpip.FullAddress) (err *tcpip.Error) {
|
||||
}
|
||||
|
||||
// Mark endpoint as bound.
|
||||
e.state = stateBound
|
||||
e.state = StateBound
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -1430,7 +1498,7 @@ func (e *endpoint) GetRemoteAddress() (tcpip.FullAddress, *tcpip.Error) {
|
||||
e.mu.RLock()
|
||||
defer e.mu.RUnlock()
|
||||
|
||||
if e.state != stateConnected {
|
||||
if !e.state.connected() {
|
||||
return tcpip.FullAddress{}, tcpip.ErrNotConnected
|
||||
}
|
||||
|
||||
@@ -1739,3 +1807,11 @@ func (e *endpoint) initGSO() {
|
||||
gso.MaxSize = e.route.GSOMaxSize()
|
||||
e.gso = gso
|
||||
}
|
||||
|
||||
// State implements tcpip.Endpoint.State. It exports the endpoint's protocol
|
||||
// state for diagnostics.
|
||||
func (e *endpoint) State() uint32 {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
return uint32(e.state)
|
||||
}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user