mirror of
https://github.com/tinygo-org/drivers.git
synced 2026-08-20 14:39:00 +00:00
wifinina: fix concurrency issues with multiple sockets
wifinina driver was not handling concurrent socket connections
correctly. When trying to make multiple socket connections, the sockfd
returned by the first Socket() call would be the same sockfd returned by
subsequent Socket() calls. The problem is the sockfd returned by
Socket() isn't used until Connect() (or Appect()). This would result in
Connect() trying to use the same sockfd for multiple connections,
causing wifinina fw to lock up.
The solution in this PR is to create a new sockfd space managed by the
driver that gives the app a unique, safe sockfd for each connection.
The real underlying sock fd returned by fw is set on Connect() (or
Accept()), and mapped back to the apps sockfd using a map:
sockets map[int]*Socket // keyed by sockfd
Where Socket has a reference to the fw sock:
type Socket struct {
protocol int
clientConnected bool
laddr netip.AddrPort // Set in Bind()
raddr netip.AddrPort // Set in Connect()
sock // Device socket, as returned from w.getSocket()
}
This commit is contained in:
@@ -39,6 +39,7 @@ var (
|
|||||||
ErrNoMoreSockets = errors.New("No more sockets")
|
ErrNoMoreSockets = errors.New("No more sockets")
|
||||||
ErrClosingSocket = errors.New("Error closing socket")
|
ErrClosingSocket = errors.New("Error closing socket")
|
||||||
ErrNotSupported = errors.New("Not supported")
|
ErrNotSupported = errors.New("Not supported")
|
||||||
|
ErrInvalidSocketFd = errors.New("Invalid socket fd")
|
||||||
)
|
)
|
||||||
|
|
||||||
// Duplicate of non-exported net.errTimeout
|
// Duplicate of non-exported net.errTimeout
|
||||||
|
|||||||
+132
-93
@@ -163,11 +163,12 @@ type encryptionType uint8
|
|||||||
type sock uint8
|
type sock uint8
|
||||||
type hwerr uint8
|
type hwerr uint8
|
||||||
|
|
||||||
type socket struct {
|
type Socket struct {
|
||||||
protocol int
|
protocol int
|
||||||
laddr netip.AddrPort // Set in Bind()
|
clientConnected bool
|
||||||
raddr netip.AddrPort // Set in Connect()
|
laddr netip.AddrPort // Set in Bind()
|
||||||
inuse bool
|
raddr netip.AddrPort // Set in Connect()
|
||||||
|
sock // Device socket, as returned from w.getSocket()
|
||||||
}
|
}
|
||||||
|
|
||||||
type Config struct {
|
type Config struct {
|
||||||
@@ -214,17 +215,13 @@ type wifinina struct {
|
|||||||
killWatchdog chan bool
|
killWatchdog chan bool
|
||||||
fault error
|
fault error
|
||||||
|
|
||||||
sockets map[sock]*socket // keyed by sock as returned by getSocket()
|
sockets map[int]*Socket // keyed by sockfd
|
||||||
}
|
|
||||||
|
|
||||||
func newSocket(protocol int) *socket {
|
|
||||||
return &socket{protocol: protocol, inuse: true}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func New(cfg *Config) *wifinina {
|
func New(cfg *Config) *wifinina {
|
||||||
w := wifinina{
|
w := wifinina{
|
||||||
cfg: cfg,
|
cfg: cfg,
|
||||||
sockets: make(map[sock]*socket),
|
sockets: make(map[int]*Socket),
|
||||||
killWatchdog: make(chan bool),
|
killWatchdog: make(chan bool),
|
||||||
cs: cfg.Cs,
|
cs: cfg.Cs,
|
||||||
ack: cfg.Ack,
|
ack: cfg.Ack,
|
||||||
@@ -505,6 +502,12 @@ func (w *wifinina) GetHostByName(name string) (netip.Addr, error) {
|
|||||||
fmt.Printf("[GetHostByName] name: %s\r\n", name)
|
fmt.Printf("[GetHostByName] name: %s\r\n", name)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// If it's already in dotted-decimal notation, return a copy
|
||||||
|
// per gethostbyname(3).
|
||||||
|
if ip, err := netip.ParseAddr(name); err == nil {
|
||||||
|
return ip, nil
|
||||||
|
}
|
||||||
|
|
||||||
w.mu.Lock()
|
w.mu.Lock()
|
||||||
defer w.mu.Unlock()
|
defer w.mu.Unlock()
|
||||||
|
|
||||||
@@ -547,6 +550,20 @@ func (w *wifinina) Addr() (netip.Addr, error) {
|
|||||||
return ip, nil
|
return ip, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// newSockfd returns the next available sockfd, or -1 if none available
|
||||||
|
func (w *wifinina) newSockfd() int {
|
||||||
|
if len(w.sockets) >= maxNetworks {
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
// Search for the next available sockfd starting at 0
|
||||||
|
for sockfd := 0; ; sockfd++ {
|
||||||
|
if _, ok := w.sockets[sockfd]; !ok {
|
||||||
|
return sockfd
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
|
||||||
// See man socket(2) for standard Berkely sockets for Socket, Bind, etc.
|
// See man socket(2) for standard Berkely sockets for Socket, Bind, etc.
|
||||||
// The driver strives to meet the function and semantics of socket(2).
|
// The driver strives to meet the function and semantics of socket(2).
|
||||||
|
|
||||||
@@ -574,34 +591,46 @@ func (w *wifinina) Socket(domain int, stype int, protocol int) (int, error) {
|
|||||||
w.mu.Lock()
|
w.mu.Lock()
|
||||||
defer w.mu.Unlock()
|
defer w.mu.Unlock()
|
||||||
|
|
||||||
sock := w.getSocket()
|
sockfd := w.newSockfd()
|
||||||
if sock == noSocketAvail {
|
if sockfd == -1 {
|
||||||
return -1, netdev.ErrNoMoreSockets
|
return -1, netdev.ErrNoMoreSockets
|
||||||
}
|
}
|
||||||
|
|
||||||
socket := newSocket(protocol)
|
w.sockets[sockfd] = &Socket{
|
||||||
w.sockets[sock] = socket
|
protocol: protocol,
|
||||||
|
sock: noSocketAvail,
|
||||||
|
}
|
||||||
|
|
||||||
return int(sock), nil
|
if debugging(debugNetdev) {
|
||||||
|
fmt.Printf("[Socket] <-- sockfd %d\r\n", sockfd)
|
||||||
|
}
|
||||||
|
|
||||||
|
return sockfd, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (w *wifinina) Bind(sockfd int, ip netip.AddrPort) error {
|
func (w *wifinina) Bind(sockfd int, ip netip.AddrPort) error {
|
||||||
|
|
||||||
if debugging(debugNetdev) {
|
if debugging(debugNetdev) {
|
||||||
fmt.Printf("[Bind] sockfd: %d, addr: %s\r\n", sockfd, ip)
|
fmt.Printf("[Bind] sockfd: %d, addr: %s:%d\r\n", sockfd, ip.Addr(), ip.Port())
|
||||||
}
|
}
|
||||||
|
|
||||||
w.mu.Lock()
|
w.mu.Lock()
|
||||||
defer w.mu.Unlock()
|
defer w.mu.Unlock()
|
||||||
|
|
||||||
var sock = sock(sockfd)
|
socket, ok := w.sockets[sockfd]
|
||||||
var socket = w.sockets[sock]
|
if !ok {
|
||||||
|
return netdev.ErrInvalidSocketFd
|
||||||
|
}
|
||||||
|
|
||||||
switch socket.protocol {
|
switch socket.protocol {
|
||||||
case netdev.IPPROTO_TCP:
|
case netdev.IPPROTO_TCP:
|
||||||
case netdev.IPPROTO_TLS:
|
case netdev.IPPROTO_TLS:
|
||||||
case netdev.IPPROTO_UDP:
|
case netdev.IPPROTO_UDP:
|
||||||
w.startServer(sock, ip.Port(), protoModeUDP)
|
socket.sock = w.getSocket()
|
||||||
|
if socket.sock == noSocketAvail {
|
||||||
|
return netdev.ErrNoMoreSockets
|
||||||
|
}
|
||||||
|
w.startServer(socket.sock, ip.Port(), protoModeUDP)
|
||||||
}
|
}
|
||||||
|
|
||||||
socket.laddr = ip
|
socket.laddr = ip
|
||||||
@@ -629,22 +658,40 @@ func (w *wifinina) Connect(sockfd int, host string, ip netip.AddrPort) error {
|
|||||||
w.mu.Lock()
|
w.mu.Lock()
|
||||||
defer w.mu.Unlock()
|
defer w.mu.Unlock()
|
||||||
|
|
||||||
var sock = sock(sockfd)
|
socket, ok := w.sockets[sockfd]
|
||||||
var socket = w.sockets[sock]
|
if !ok {
|
||||||
|
return netdev.ErrInvalidSocketFd
|
||||||
|
}
|
||||||
|
|
||||||
// Start the connection
|
// Start the connection
|
||||||
switch socket.protocol {
|
switch socket.protocol {
|
||||||
|
|
||||||
case netdev.IPPROTO_TCP:
|
case netdev.IPPROTO_TCP:
|
||||||
w.startClient(sock, "", toUint32(ip.Addr().As4()), ip.Port(), protoModeTCP)
|
socket.sock = w.getSocket()
|
||||||
|
if socket.sock == noSocketAvail {
|
||||||
|
return netdev.ErrNoMoreSockets
|
||||||
|
}
|
||||||
|
w.startClient(socket.sock, "", toUint32(ip.Addr().As4()), ip.Port(), protoModeTCP)
|
||||||
|
|
||||||
case netdev.IPPROTO_TLS:
|
case netdev.IPPROTO_TLS:
|
||||||
w.startClient(sock, host, 0, ip.Port(), protoModeTLS)
|
socket.sock = w.getSocket()
|
||||||
|
if socket.sock == noSocketAvail {
|
||||||
|
return netdev.ErrNoMoreSockets
|
||||||
|
}
|
||||||
|
w.startClient(socket.sock, host, 0, ip.Port(), protoModeTLS)
|
||||||
|
|
||||||
case netdev.IPPROTO_UDP:
|
case netdev.IPPROTO_UDP:
|
||||||
|
if socket.sock == noSocketAvail {
|
||||||
|
return fmt.Errorf("Must Bind before Connecting")
|
||||||
|
}
|
||||||
// See start in sendUDP()
|
// See start in sendUDP()
|
||||||
socket.raddr = ip
|
socket.raddr = ip
|
||||||
|
socket.clientConnected = true
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
if w.getClientState(sock) == tcpStateEstablished {
|
if w.getClientState(socket.sock) == tcpStateEstablished {
|
||||||
|
socket.clientConnected = true
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -664,12 +711,18 @@ func (w *wifinina) Listen(sockfd int, backlog int) error {
|
|||||||
w.mu.Lock()
|
w.mu.Lock()
|
||||||
defer w.mu.Unlock()
|
defer w.mu.Unlock()
|
||||||
|
|
||||||
var sock = sock(sockfd)
|
socket, ok := w.sockets[sockfd]
|
||||||
var socket = w.sockets[sock]
|
if !ok {
|
||||||
|
return netdev.ErrInvalidSocketFd
|
||||||
|
}
|
||||||
|
|
||||||
switch socket.protocol {
|
switch socket.protocol {
|
||||||
case netdev.IPPROTO_TCP:
|
case netdev.IPPROTO_TCP:
|
||||||
w.startServer(sock, socket.laddr.Port(), protoModeTCP)
|
socket.sock = w.getSocket()
|
||||||
|
if socket.sock == noSocketAvail {
|
||||||
|
return netdev.ErrNoMoreSockets
|
||||||
|
}
|
||||||
|
w.startServer(socket.sock, socket.laddr.Port(), protoModeTCP)
|
||||||
case netdev.IPPROTO_UDP:
|
case netdev.IPPROTO_UDP:
|
||||||
default:
|
default:
|
||||||
return netdev.ErrProtocolNotSupported
|
return netdev.ErrProtocolNotSupported
|
||||||
@@ -687,9 +740,10 @@ func (w *wifinina) Accept(sockfd int) (int, netip.AddrPort, error) {
|
|||||||
w.mu.Lock()
|
w.mu.Lock()
|
||||||
defer w.mu.Unlock()
|
defer w.mu.Unlock()
|
||||||
|
|
||||||
var client sock
|
socket, ok := w.sockets[sockfd]
|
||||||
var sock = sock(sockfd)
|
if !ok {
|
||||||
var socket = w.sockets[sock]
|
return -1, netip.AddrPort{}, netdev.ErrInvalidSocketFd
|
||||||
|
}
|
||||||
|
|
||||||
switch socket.protocol {
|
switch socket.protocol {
|
||||||
case netdev.IPPROTO_TCP:
|
case netdev.IPPROTO_TCP:
|
||||||
@@ -697,8 +751,9 @@ func (w *wifinina) Accept(sockfd int) (int, netip.AddrPort, error) {
|
|||||||
return -1, netip.AddrPort{}, netdev.ErrProtocolNotSupported
|
return -1, netip.AddrPort{}, netdev.ErrProtocolNotSupported
|
||||||
}
|
}
|
||||||
|
|
||||||
|
skip:
|
||||||
for {
|
for {
|
||||||
// Accept() will be sleeping most of the time, checking for a
|
// Accept() will be sleeping most of the time, checking for
|
||||||
// new clients every 1/10 sec.
|
// new clients every 1/10 sec.
|
||||||
w.mu.Unlock()
|
w.mu.Unlock()
|
||||||
time.Sleep(100 * time.Millisecond)
|
time.Sleep(100 * time.Millisecond)
|
||||||
@@ -709,48 +764,43 @@ func (w *wifinina) Accept(sockfd int) (int, netip.AddrPort, error) {
|
|||||||
return -1, netip.AddrPort{}, w.fault
|
return -1, netip.AddrPort{}, w.fault
|
||||||
}
|
}
|
||||||
|
|
||||||
// TODO: BUG: Currently, a sock that is 100% busy will always be
|
|
||||||
// TODO: returned by w.accept(sock), starving other socks
|
|
||||||
// TODO: from begin serviced. Need to figure out how to
|
|
||||||
// TODO: service socks fairly (round-robin?) so no one sock
|
|
||||||
// TODO: can dominate.
|
|
||||||
|
|
||||||
// Check if a client has data
|
// Check if a client has data
|
||||||
client = w.accept(sock)
|
var client sock = w.accept(socket.sock)
|
||||||
if client == noSocketAvail {
|
if client == noSocketAvail {
|
||||||
// None ready
|
// None ready
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
raddr := w.getRemoteData(client)
|
// If we already have a socket for the client, skip
|
||||||
|
for _, s := range w.sockets {
|
||||||
// If we've already seen this socket, we can reuse
|
if s.sock == client {
|
||||||
// the socket and return it. But, only if the socket
|
continue skip
|
||||||
// is closed. If it's not closed, we'll just come back
|
|
||||||
// later to reuse it.
|
|
||||||
|
|
||||||
clientSocket, ok := w.sockets[client]
|
|
||||||
if ok {
|
|
||||||
// Wait for client to Close
|
|
||||||
if clientSocket.inuse {
|
|
||||||
continue
|
|
||||||
}
|
}
|
||||||
// Reuse client socket
|
|
||||||
return int(client), raddr, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Create new socket for client and return fd
|
// Otherwise, create a new socket
|
||||||
w.sockets[client] = newSocket(socket.protocol)
|
clientfd := w.newSockfd()
|
||||||
return int(client), raddr, nil
|
if clientfd == -1 {
|
||||||
|
return -1, netip.AddrPort{}, netdev.ErrNoMoreSockets
|
||||||
|
}
|
||||||
|
|
||||||
|
w.sockets[clientfd] = &Socket{
|
||||||
|
protocol: netdev.IPPROTO_TCP,
|
||||||
|
sock: client,
|
||||||
|
clientConnected: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
raddr := w.getRemoteData(client)
|
||||||
|
|
||||||
|
return clientfd, raddr, nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (w *wifinina) sockDown(sock sock) bool {
|
func (w *wifinina) sockDown(socket *Socket) bool {
|
||||||
var socket = w.sockets[sock]
|
|
||||||
if socket.protocol == netdev.IPPROTO_UDP {
|
if socket.protocol == netdev.IPPROTO_UDP {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
return w.getClientState(sock) != tcpStateEstablished
|
return w.getClientState(socket.sock) != tcpStateEstablished
|
||||||
}
|
}
|
||||||
|
|
||||||
func (w *wifinina) sendTCP(sock sock, buf []byte, deadline time.Time) (int, error) {
|
func (w *wifinina) sendTCP(sock sock, buf []byte, deadline time.Time) (int, error) {
|
||||||
@@ -778,7 +828,7 @@ func (w *wifinina) sendTCP(sock sock, buf []byte, deadline time.Time) (int, erro
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Check if socket went down
|
// Check if socket went down
|
||||||
if w.sockDown(sock) {
|
if w.getClientState(sock) != tcpStateEstablished {
|
||||||
return -1, io.EOF
|
return -1, io.EOF
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -817,8 +867,10 @@ func (w *wifinina) sendUDP(sock sock, raddr netip.AddrPort, buf []byte, deadline
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (w *wifinina) sendChunk(sockfd int, buf []byte, deadline time.Time) (int, error) {
|
func (w *wifinina) sendChunk(sockfd int, buf []byte, deadline time.Time) (int, error) {
|
||||||
var sock = sock(sockfd)
|
socket, ok := w.sockets[sockfd]
|
||||||
var socket = w.sockets[sock]
|
if !ok {
|
||||||
|
return -1, netdev.ErrInvalidSocketFd
|
||||||
|
}
|
||||||
|
|
||||||
// Check if we've timed out
|
// Check if we've timed out
|
||||||
if !deadline.IsZero() {
|
if !deadline.IsZero() {
|
||||||
@@ -829,9 +881,9 @@ func (w *wifinina) sendChunk(sockfd int, buf []byte, deadline time.Time) (int, e
|
|||||||
|
|
||||||
switch socket.protocol {
|
switch socket.protocol {
|
||||||
case netdev.IPPROTO_TCP, netdev.IPPROTO_TLS:
|
case netdev.IPPROTO_TCP, netdev.IPPROTO_TLS:
|
||||||
return w.sendTCP(sock, buf, deadline)
|
return w.sendTCP(socket.sock, buf, deadline)
|
||||||
case netdev.IPPROTO_UDP:
|
case netdev.IPPROTO_UDP:
|
||||||
return w.sendUDP(sock, socket.raddr, buf, deadline)
|
return w.sendUDP(socket.sock, socket.raddr, buf, deadline)
|
||||||
}
|
}
|
||||||
|
|
||||||
return -1, netdev.ErrProtocolNotSupported
|
return -1, netdev.ErrProtocolNotSupported
|
||||||
@@ -876,7 +928,10 @@ func (w *wifinina) Recv(sockfd int, buf []byte, flags int,
|
|||||||
w.mu.Lock()
|
w.mu.Lock()
|
||||||
defer w.mu.Unlock()
|
defer w.mu.Unlock()
|
||||||
|
|
||||||
var sock = sock(sockfd)
|
socket, ok := w.sockets[sockfd]
|
||||||
|
if !ok {
|
||||||
|
return -1, netdev.ErrInvalidSocketFd
|
||||||
|
}
|
||||||
|
|
||||||
// Limit max read size to chunk large read requests
|
// Limit max read size to chunk large read requests
|
||||||
var max = len(buf)
|
var max = len(buf)
|
||||||
@@ -897,22 +952,22 @@ func (w *wifinina) Recv(sockfd int, buf []byte, flags int,
|
|||||||
// doesn't return unless there is data, even a single byte, or
|
// doesn't return unless there is data, even a single byte, or
|
||||||
// on error such as timeout or EOF.
|
// on error such as timeout or EOF.
|
||||||
|
|
||||||
n := int(w.getDataBuf(sock, buf[:max]))
|
n := int(w.getDataBuf(socket.sock, buf[:max]))
|
||||||
if n > 0 {
|
if n > 0 {
|
||||||
if debugging(debugNetdev) {
|
if debugging(debugNetdev) {
|
||||||
fmt.Printf("[<--Recv] sockfd: %d, n: %d\r\n",
|
fmt.Printf("[<--Recv] sockfd: %d, n: %d\r\n",
|
||||||
sock, n)
|
sockfd, n)
|
||||||
}
|
}
|
||||||
return n, nil
|
return n, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check if socket went down
|
// Check if socket went down
|
||||||
if w.sockDown(sock) {
|
if w.sockDown(socket) {
|
||||||
// Get any last bytes
|
// Get any last bytes
|
||||||
n = int(w.getDataBuf(sock, buf[:max]))
|
n = int(w.getDataBuf(socket.sock, buf[:max]))
|
||||||
if debugging(debugNetdev) {
|
if debugging(debugNetdev) {
|
||||||
fmt.Printf("[<--Recv] sockfd: %d, n: %d, EOF\r\n",
|
fmt.Printf("[<--Recv] sockfd: %d, n: %d, EOF\r\n",
|
||||||
sock, n)
|
sockfd, n)
|
||||||
}
|
}
|
||||||
if n > 0 {
|
if n > 0 {
|
||||||
return n, io.EOF
|
return n, io.EOF
|
||||||
@@ -941,34 +996,18 @@ func (w *wifinina) Close(sockfd int) error {
|
|||||||
w.mu.Lock()
|
w.mu.Lock()
|
||||||
defer w.mu.Unlock()
|
defer w.mu.Unlock()
|
||||||
|
|
||||||
var sock = sock(sockfd)
|
socket, ok := w.sockets[sockfd]
|
||||||
var socket = w.sockets[sock]
|
if !ok {
|
||||||
|
return netdev.ErrInvalidSocketFd
|
||||||
if !socket.inuse {
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
w.stopClient(sock)
|
if socket.clientConnected {
|
||||||
|
w.stopClient(socket.sock)
|
||||||
if socket.protocol == netdev.IPPROTO_UDP {
|
|
||||||
socket.inuse = false
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
start := time.Now()
|
delete(w.sockets, sockfd)
|
||||||
for time.Since(start) < 5*time.Second {
|
|
||||||
|
|
||||||
if w.getClientState(sock) == tcpStateClosed {
|
return nil
|
||||||
socket.inuse = false
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
w.mu.Unlock()
|
|
||||||
time.Sleep(100 * time.Millisecond)
|
|
||||||
w.mu.Lock()
|
|
||||||
}
|
|
||||||
|
|
||||||
return netdev.ErrClosingSocket
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (w *wifinina) SetSockOpt(sockfd int, level int, opt int, value interface{}) error {
|
func (w *wifinina) SetSockOpt(sockfd int, level int, opt int, value interface{}) error {
|
||||||
|
|||||||
Reference in New Issue
Block a user