mirror of
https://github.com/soypat/lneto.git
synced 2026-08-24 00:19:03 +00:00
mega refactor internet.TCPConn->tcp.Conn
This commit is contained in:
@@ -0,0 +1,149 @@
|
||||
package internet
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net"
|
||||
"time"
|
||||
|
||||
"github.com/soypat/lneto"
|
||||
"github.com/soypat/lneto/internal"
|
||||
"github.com/soypat/lneto/tcp"
|
||||
)
|
||||
|
||||
var _ StackNode = (*NodeTCPListener)(nil)
|
||||
|
||||
type NodeTCPListener struct {
|
||||
connID uint64
|
||||
conns []tcp.Conn
|
||||
accepted []bool
|
||||
port uint16
|
||||
getISS func() uint32
|
||||
}
|
||||
|
||||
func (listener *NodeTCPListener) AcceptRaw() (*tcp.Conn, error) {
|
||||
connid := listener.connID
|
||||
for {
|
||||
if listener.isClosed() || connid != listener.connID {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
for i := range listener.conns {
|
||||
isAvailable := listener.connReceivedSyn(i) && !listener.connAccepted(i)
|
||||
if !isAvailable {
|
||||
continue
|
||||
}
|
||||
// Connection received as SYN and is not yet accepted.
|
||||
listener.accepted[i] = true
|
||||
return &listener.conns[i], nil
|
||||
}
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
}
|
||||
panic("unreachable")
|
||||
}
|
||||
|
||||
func (listener *NodeTCPListener) Close() error {
|
||||
if listener.isClosed() {
|
||||
return errors.New("already closed")
|
||||
}
|
||||
listener.connID++
|
||||
listener.port = 0
|
||||
return nil
|
||||
}
|
||||
|
||||
func (listener *NodeTCPListener) LocalPort() uint16 { return listener.port }
|
||||
|
||||
func (listener *NodeTCPListener) ConnectionID() *uint64 { return &listener.connID }
|
||||
|
||||
func (listener *NodeTCPListener) Protocol() uint64 { return uint64(lneto.IPProtoTCP) }
|
||||
|
||||
func (listener *NodeTCPListener) Encapsulate(carrierData []byte, tcpFrameOffset int) (int, error) {
|
||||
if listener.isClosed() {
|
||||
return 0, net.ErrClosed
|
||||
}
|
||||
for i := range listener.conns {
|
||||
conn := &listener.conns[i]
|
||||
if conn.State().IsClosed() {
|
||||
continue
|
||||
}
|
||||
n, err := conn.Encapsulate(carrierData, tcpFrameOffset)
|
||||
if err != nil {
|
||||
listener.maintainConn(i, err)
|
||||
}
|
||||
if n == 0 {
|
||||
continue
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
func (listener *NodeTCPListener) Demux(carrierData []byte, tcpFrameOffset int) error {
|
||||
if listener.isClosed() {
|
||||
return net.ErrClosed
|
||||
}
|
||||
tfrm, err := tcp.NewFrame(carrierData[tcpFrameOffset:])
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
addr, _, err := internal.GetIPSourceAddr(carrierData)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
dst := tfrm.DestinationPort()
|
||||
if dst != listener.port {
|
||||
return errors.New("not our port")
|
||||
}
|
||||
src := tfrm.DestinationPort()
|
||||
_, flags := tfrm.OffsetAndFlags()
|
||||
for i := range listener.conns {
|
||||
if listener.conns[i].RemotePort() != src || !bytes.Equal(listener.conns[i].RemoteAddr(), addr) {
|
||||
continue
|
||||
}
|
||||
conn := &listener.conns[i]
|
||||
err := conn.Demux(carrierData, tcpFrameOffset)
|
||||
if err != nil {
|
||||
listener.maintainConn(i, err)
|
||||
}
|
||||
return err
|
||||
}
|
||||
if !flags.HasAll(tcp.FlagSYN) {
|
||||
return nil // Not a synchronizing packet, drop it.
|
||||
}
|
||||
// New connection must be assigned.
|
||||
for i := range listener.conns {
|
||||
conn := &listener.conns[i]
|
||||
isOpen := !conn.State().IsClosed()
|
||||
if isOpen {
|
||||
continue
|
||||
}
|
||||
if conn.State() == tcp.StateTimeWait {
|
||||
conn.Abort()
|
||||
}
|
||||
|
||||
err = conn.OpenListen(dst, tcp.Value(listener.getISS()))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return conn.Demux(carrierData, tcpFrameOffset)
|
||||
}
|
||||
slog.Error("tcpListener:no-free-conn")
|
||||
return nil
|
||||
}
|
||||
|
||||
func (listener *NodeTCPListener) maintainConn(connIdx int, err error) {
|
||||
if err == net.ErrClosed {
|
||||
listener.conns[connIdx].Abort()
|
||||
}
|
||||
}
|
||||
|
||||
func (listener *NodeTCPListener) isClosed() bool {
|
||||
return listener.port == 0
|
||||
}
|
||||
|
||||
func (listener *NodeTCPListener) connReceivedSyn(idx int) bool {
|
||||
return listener.conns[idx].RemotePort() != 0
|
||||
}
|
||||
func (listener *NodeTCPListener) connAccepted(idx int) bool {
|
||||
return listener.accepted[idx]
|
||||
}
|
||||
@@ -164,7 +164,7 @@ func (sb *StackIP) Register(h StackNode) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (sb *StackIP) RegisterTCPConn(conn *TCPConn) error {
|
||||
func (sb *StackIP) RegisterTCPConn(conn *tcp.Conn) error {
|
||||
if conn.LocalPort() == 0 {
|
||||
return errZeroPort
|
||||
}
|
||||
|
||||
@@ -11,7 +11,7 @@ import (
|
||||
func TestBasicStack(t *testing.T) {
|
||||
rng := rand.New(rand.NewSource(1))
|
||||
var sbCl, sbSv StackIP
|
||||
var connCl, connSv TCPConn
|
||||
var connCl, connSv tcp.Conn
|
||||
setupClientServer(t, rng, &sbCl, &sbSv, &connCl, &connSv)
|
||||
var buf [2048]byte
|
||||
nextToSend := &sbCl
|
||||
@@ -37,7 +37,7 @@ func TestBasicStack(t *testing.T) {
|
||||
func TestBasicStack2(t *testing.T) {
|
||||
rng := rand.New(rand.NewSource(1))
|
||||
var sbCl, sbSv StackIP
|
||||
var connCl, connSv TCPConn
|
||||
var connCl, connSv tcp.Conn
|
||||
setupClientServerEstablished(t, rng, &sbCl, &sbSv, &connCl, &connSv)
|
||||
|
||||
}
|
||||
@@ -57,7 +57,7 @@ func expectExchange(t *testing.T, from, to *StackIP, buf []byte) {
|
||||
}
|
||||
}
|
||||
|
||||
func setupClientServerEstablished(t *testing.T, rng *rand.Rand, client, server *StackIP, connClient, connServer *TCPConn) {
|
||||
func setupClientServerEstablished(t *testing.T, rng *rand.Rand, client, server *StackIP, connClient, connServer *tcp.Conn) {
|
||||
t.Helper()
|
||||
setupClientServer(t, rng, client, server, connClient, connServer)
|
||||
var buf [2048]byte
|
||||
@@ -85,7 +85,7 @@ func setupClientServerEstablished(t *testing.T, rng *rand.Rand, client, server *
|
||||
}
|
||||
}
|
||||
|
||||
func setupClientServer(t *testing.T, rng *rand.Rand, client, server *StackIP, connClient, connServer *TCPConn) {
|
||||
func setupClientServer(t *testing.T, rng *rand.Rand, client, server *StackIP, connClient, connServer *tcp.Conn) {
|
||||
bufsize := 2048
|
||||
// Ensure buffer sizes are OK with reused buffers.
|
||||
svip := netip.AddrPortFrom(netip.AddrFrom4([4]byte{192, 168, 1, 0}), 80)
|
||||
@@ -93,7 +93,7 @@ func setupClientServer(t *testing.T, rng *rand.Rand, client, server *StackIP, co
|
||||
server.SetAddr(svip.Addr())
|
||||
client.SetAddr(clip.Addr())
|
||||
|
||||
err := connServer.Configure(&TCPConnConfig{
|
||||
err := connServer.Configure(&tcp.ConnConfig{
|
||||
RxBuf: make([]byte, bufsize),
|
||||
TxBuf: make([]byte, bufsize),
|
||||
TxPacketQueueSize: 3,
|
||||
@@ -102,7 +102,7 @@ func setupClientServer(t *testing.T, rng *rand.Rand, client, server *StackIP, co
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err = connClient.Configure(&TCPConnConfig{
|
||||
err = connClient.Configure(&tcp.ConnConfig{
|
||||
RxBuf: make([]byte, bufsize),
|
||||
TxBuf: make([]byte, bufsize),
|
||||
TxPacketQueueSize: 3,
|
||||
|
||||
@@ -1,329 +0,0 @@
|
||||
package internet
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"runtime"
|
||||
"time"
|
||||
|
||||
"github.com/soypat/lneto"
|
||||
"github.com/soypat/lneto/internal"
|
||||
"github.com/soypat/lneto/ipv4"
|
||||
"github.com/soypat/lneto/ipv6"
|
||||
"github.com/soypat/lneto/tcp"
|
||||
)
|
||||
|
||||
var (
|
||||
errDeadlineExceeded = os.ErrDeadlineExceeded
|
||||
)
|
||||
|
||||
type TCPConn struct {
|
||||
h tcp.Handler
|
||||
remoteAddr []byte
|
||||
|
||||
rdead time.Time
|
||||
wdead time.Time
|
||||
lastTx time.Time
|
||||
lastRx time.Time
|
||||
|
||||
ipID uint16
|
||||
abortErr error
|
||||
logger
|
||||
}
|
||||
|
||||
type TCPConnConfig struct {
|
||||
RxBuf []byte
|
||||
TxBuf []byte
|
||||
TxPacketQueueSize int
|
||||
Logger *slog.Logger
|
||||
}
|
||||
|
||||
func (conn *TCPConn) Configure(config *TCPConnConfig) (err error) {
|
||||
err = conn.h.SetBuffers(config.TxBuf, config.RxBuf, config.TxPacketQueueSize)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
conn.logger.log = config.Logger
|
||||
return nil
|
||||
}
|
||||
|
||||
// LocalPort returns the local port on which the socket is listening or connected to.
|
||||
func (conn *TCPConn) LocalPort() uint16 { return conn.h.LocalPort() }
|
||||
|
||||
// RemotePort returns the port of the incoming remote connection. Is non-zero if connection is established.
|
||||
func (conn *TCPConn) RemotePort() uint16 { return conn.h.RemotePort() }
|
||||
|
||||
// State returns the TCP state of the socket.
|
||||
func (conn *TCPConn) State() tcp.State { return conn.h.State() }
|
||||
|
||||
// BufferedInput returns the number of bytes in the socket's receive/input buffer.
|
||||
func (conn *TCPConn) BufferedInput() int { return conn.h.BufferedInput() }
|
||||
|
||||
// OpenActive opens a connection to a remote peer with a known IP address and port combination.
|
||||
// iss is the initial send sequence number which is ideally a random number which is far away from the last sequence number used on a connection to the same host.
|
||||
func (conn *TCPConn) OpenActive(remote netip.AddrPort, localPort uint16, iss tcp.Value) error {
|
||||
err := conn.h.OpenActive(localPort, remote.Port(), iss)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
conn.reset(conn.h)
|
||||
raddr := remote.Addr()
|
||||
if raddr.Is4() {
|
||||
addr4 := raddr.As4()
|
||||
conn.remoteAddr = append(conn.remoteAddr[:0], addr4[:]...)
|
||||
} else if raddr.Is6() {
|
||||
addr6 := raddr.As16()
|
||||
conn.remoteAddr = append(conn.remoteAddr[:0], addr6[:]...)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// OpenListen opens a passive connection which listens for the first SYN packet to be received on a local port.
|
||||
// iss is the initial send sequence number which is usually a randomly chosen number.
|
||||
func (conn *TCPConn) OpenListen(localPort uint16, iss tcp.Value) error {
|
||||
err := conn.h.OpenListen(localPort, iss)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
conn.reset(conn.h)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (conn *TCPConn) Close() error {
|
||||
conn.trace("TCPConn.Close")
|
||||
return conn.h.Close()
|
||||
}
|
||||
|
||||
func (conn *TCPConn) Demux(buf []byte, off int) (err error) {
|
||||
conn.trace("tcpconn.Recv:start")
|
||||
if off >= len(buf) {
|
||||
return errors.New("bad offset in TCPConn.Recv")
|
||||
}
|
||||
raddr, id, err := getIPAddr(buf[:off])
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if conn.isRaddrSet() && !bytes.Equal(conn.remoteAddr, raddr) {
|
||||
return errors.New("IP addr mismatch on TCPConn")
|
||||
}
|
||||
err = conn.h.Recv(buf[off:])
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !conn.isRaddrSet() && conn.h.RemotePort() != 0 {
|
||||
conn.remoteAddr = append(conn.remoteAddr[:0], raddr...)
|
||||
conn.ipID = ^(id - 1)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Write writes argument data to the TCPConns's output buffer which is queued to be sent.
|
||||
func (conn *TCPConn) Write(b []byte) (int, error) {
|
||||
err := conn.checkPipeOpen()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
plen := len(b)
|
||||
conn.trace("TCPConn.Write:start")
|
||||
connid := conn.h.ConnectionID()
|
||||
if conn.deadlineExceeded(conn.wdead) {
|
||||
return 0, errDeadlineExceeded
|
||||
} else if plen == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
backoff := internal.NewBackoff(internal.BackoffTCPConn)
|
||||
n := 0
|
||||
for {
|
||||
if conn.abortErr != nil {
|
||||
return n, conn.abortErr
|
||||
} else if connid != conn.h.ConnectionID() {
|
||||
return n, net.ErrClosed
|
||||
}
|
||||
ngot, _ := conn.h.Write(b)
|
||||
n += ngot
|
||||
b = b[ngot:]
|
||||
if n == plen {
|
||||
break
|
||||
} else if ngot > 0 {
|
||||
backoff.Hit()
|
||||
runtime.Gosched() // Do a little yield since we won't have data for sure otherwise.
|
||||
} else {
|
||||
backoff.Miss()
|
||||
}
|
||||
conn.trace("TCPConn.Write:insuf-buf", slog.Int("missing", plen-n))
|
||||
if conn.deadlineExceeded(conn.wdead) {
|
||||
return n, errDeadlineExceeded
|
||||
}
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// Read reads data from the socket's input buffer. If the buffer is empty,
|
||||
// Read will block until data is available or connection closes.
|
||||
func (conn *TCPConn) Read(b []byte) (int, error) {
|
||||
err := conn.checkPipeOpen()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
conn.trace("TCPConn.Read:start")
|
||||
connid := conn.h.ConnectionID()
|
||||
backoff := internal.NewBackoff(internal.BackoffTCPConn)
|
||||
for conn.h.BufferedInput() == 0 && conn.State() == tcp.StateEstablished {
|
||||
if conn.abortErr != nil {
|
||||
return 0, conn.abortErr
|
||||
} else if connid != conn.h.ConnectionID() {
|
||||
return 0, net.ErrClosed
|
||||
}
|
||||
if conn.deadlineExceeded(conn.rdead) {
|
||||
return 0, errDeadlineExceeded
|
||||
}
|
||||
backoff.Miss()
|
||||
}
|
||||
n, err := conn.h.Read(b)
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (conn *TCPConn) checkPipeOpen() error {
|
||||
if conn.abortErr != nil {
|
||||
return conn.abortErr
|
||||
}
|
||||
state := conn.State()
|
||||
if state.IsClosed() {
|
||||
return net.ErrClosed
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (conn *TCPConn) Encapsulate(buf []byte, off int) (n int, err error) {
|
||||
if len(conn.remoteAddr) == 0 {
|
||||
return 0, errors.New("unset IP address")
|
||||
}
|
||||
raddr, _, err := getIPAddr(buf[:off])
|
||||
if err != nil {
|
||||
return 0, err
|
||||
} else if len(raddr) != len(conn.remoteAddr) {
|
||||
return 0, errors.New("mismatched IP version")
|
||||
}
|
||||
n, err = conn.h.Send(buf[off:])
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
err = setDstAddr(buf[:off], conn.ipID, conn.remoteAddr)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
conn.ipID++
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func (conn *TCPConn) Protocol() uint64 {
|
||||
return uint64(lneto.IPProtoTCP)
|
||||
}
|
||||
|
||||
func getIPAddr(buf []byte) (addr []byte, id uint16, err error) {
|
||||
switch buf[0] >> 4 {
|
||||
case 4:
|
||||
ifrm4, err := ipv4.NewFrame(buf)
|
||||
if err != nil {
|
||||
return addr, 0, err
|
||||
}
|
||||
addr = ifrm4.SourceAddr()[:]
|
||||
id = ifrm4.ID()
|
||||
case 6:
|
||||
ifrm6, err := ipv6.NewFrame(buf)
|
||||
if err != nil {
|
||||
return addr, 0, err
|
||||
}
|
||||
addr = ifrm6.SourceAddr()[:]
|
||||
default:
|
||||
err = errors.New("unsupported IP version")
|
||||
}
|
||||
return addr, id, err
|
||||
}
|
||||
|
||||
func setDstAddr(buf []byte, id uint16, addr []byte) (err error) {
|
||||
var dstaddr []byte
|
||||
switch buf[0] >> 4 {
|
||||
case 4:
|
||||
ifrm4, err := ipv4.NewFrame(buf)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
dstaddr = ifrm4.DestinationAddr()[:]
|
||||
ifrm4.SetID(id)
|
||||
case 6:
|
||||
ifrm6, err := ipv6.NewFrame(buf)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
dstaddr = ifrm6.DestinationAddr()[:]
|
||||
default:
|
||||
err = errors.New("unsupported IP version")
|
||||
}
|
||||
if err == nil && len(dstaddr) != len(addr) {
|
||||
return errors.New("invalid ip version to setDstAddr")
|
||||
}
|
||||
copy(dstaddr, addr)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (conn *TCPConn) isRaddrSet() bool {
|
||||
return len(conn.remoteAddr) != 0
|
||||
}
|
||||
|
||||
func (conn *TCPConn) reset(h tcp.Handler) {
|
||||
*conn = TCPConn{
|
||||
h: h,
|
||||
remoteAddr: conn.remoteAddr[:0],
|
||||
logger: conn.logger,
|
||||
}
|
||||
}
|
||||
|
||||
// SetDeadline sets the read and write deadlines associated
|
||||
// with the connection. It is equivalent to calling both
|
||||
// SetReadDeadline and SetWriteDeadline. Implements [net.Conn].
|
||||
func (conn *TCPConn) SetDeadline(t time.Time) error {
|
||||
err := conn.SetReadDeadline(t)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return conn.SetWriteDeadline(t)
|
||||
}
|
||||
|
||||
// SetReadDeadline sets the deadline for future Read calls
|
||||
// and any currently-blocked Read call. A zero value for t means Read will not time out.
|
||||
func (conn *TCPConn) SetReadDeadline(t time.Time) error {
|
||||
conn.trace("TCPConn.SetReadDeadline:start")
|
||||
err := conn.checkPipeOpen()
|
||||
if err == nil {
|
||||
conn.rdead = t
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// SetWriteDeadline sets the deadline for future Write calls
|
||||
// and any currently-blocked Write call.
|
||||
// Even if write times out, it may return n > 0, indicating that
|
||||
// some of the data was successfully written.
|
||||
// A zero value for t means Write will not time out.
|
||||
func (conn *TCPConn) SetWriteDeadline(t time.Time) error {
|
||||
conn.trace("TCPConn.SetWriteDeadline:start")
|
||||
err := conn.checkPipeOpen()
|
||||
if err == nil {
|
||||
conn.wdead = t
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func (conn *TCPConn) deadlineExceeded(deadline time.Time) bool {
|
||||
return !deadline.IsZero() && time.Since(deadline) > 0
|
||||
}
|
||||
|
||||
func (conn *TCPConn) ConnectionID() *uint64 {
|
||||
return conn.h.ConnectionID()
|
||||
}
|
||||
Reference in New Issue
Block a user