From 7a7235f83b13daae62e738609c76e23999ee7fcc Mon Sep 17 00:00:00 2001 From: Scott Feldman Date: Fri, 1 Dec 2023 22:30:46 -0800 Subject: [PATCH] move netdev interface to use net/netip for IP addr/ports --- espat/espat.go | 42 ++++++++++++----------- netdev/netdev.go | 15 ++++---- rtl8720dn/rtl8720dn.go | 77 +++++++++++++++++++++++------------------- wifinina/wifinina.go | 67 ++++++++++++++++++------------------ 4 files changed, 108 insertions(+), 93 deletions(-) diff --git a/espat/espat.go b/espat/espat.go index aca6d7a..e571cb0 100644 --- a/espat/espat.go +++ b/espat/espat.go @@ -25,6 +25,7 @@ import ( "fmt" "machine" "net" + "net/netip" "strconv" "strings" "sync" @@ -44,8 +45,8 @@ type Config struct { type socket struct { inUse bool protocol int - lip net.IP - lport int + lip netip.Addr + lport uint16 } type Device struct { @@ -106,7 +107,7 @@ func (d *Device) NetConnect(params *netlink.ConnectParams) error { fmt.Printf("CONNECTED\r\n") - ip, err := d.GetIPAddr() + ip, err := d.Addr() if err != nil { return err } @@ -125,28 +126,31 @@ func (d *Device) NetNotify(cb func(netlink.Event)) { // Not supported } -func (d *Device) GetHostByName(name string) (net.IP, error) { +func (d *Device) GetHostByName(name string) (netip.Addr, error) { ip, err := d.GetDNS(name) - return net.ParseIP(ip), err + if err != nil { + return netip.Addr{}, err + } + return netip.ParseAddr(ip) } func (d *Device) GetHardwareAddr() (net.HardwareAddr, error) { return net.HardwareAddr{}, netlink.ErrNotSupported } -func (d *Device) GetIPAddr() (net.IP, error) { +func (d *Device) Addr() (netip.Addr, error) { resp, err := d.GetClientIP() if err != nil { - return net.IP{}, err + return netip.Addr{}, err } prefix := "+CIPSTA:ip:" for _, line := range strings.Split(resp, "\n") { if ok := strings.HasPrefix(line, prefix); ok { ip := line[len(prefix)+1 : len(line)-2] - return net.ParseIP(ip), nil + return netip.ParseAddr(ip) } } - return net.IP{}, fmt.Errorf("Error getting IP address") + return netip.Addr{}, fmt.Errorf("Error getting IP address") } func (d *Device) Socket(domain int, stype int, protocol int) (int, error) { @@ -175,17 +179,17 @@ func (d *Device) Socket(domain int, stype int, protocol int) (int, error) { return 0, nil } -func (d *Device) Bind(sockfd int, ip net.IP, port int) error { - d.socket.lip = ip - d.socket.lport = port +func (d *Device) Bind(sockfd int, ip netip.AddrPort) error { + d.socket.lip = ip.Addr() + d.socket.lport = ip.Port() return nil } -func (d *Device) Connect(sockfd int, host string, ip net.IP, port int) error { +func (d *Device) Connect(sockfd int, host string, ip netip.AddrPort) error { var err error - var addr = ip.String() - var rport = strconv.Itoa(port) - var lport = strconv.Itoa(d.socket.lport) + var addr = ip.Addr().String() + var rport = strconv.Itoa(int(ip.Port())) + var lport = strconv.Itoa(int(d.socket.lport)) switch d.socket.protocol { case netdev.IPPROTO_TCP: @@ -198,9 +202,9 @@ func (d *Device) Connect(sockfd int, host string, ip net.IP, port int) error { if err != nil { if host == "" { - return fmt.Errorf("Connect to %s:%d timed out", ip, port) + return fmt.Errorf("Connect to %s timed out", ip) } else { - return fmt.Errorf("Connect to %s:%d timed out", host, port) + return fmt.Errorf("Connect to %s:%d timed out", host, ip.Port()) } } @@ -216,7 +220,7 @@ func (d *Device) Listen(sockfd int, backlog int) error { return nil } -func (d *Device) Accept(sockfd int, ip net.IP, port int) (int, error) { +func (d *Device) Accept(sockfd int, ip netip.AddrPort) (int, error) { return -1, netdev.ErrNotSupported } diff --git a/netdev/netdev.go b/netdev/netdev.go index cf07c07..1709270 100644 --- a/netdev/netdev.go +++ b/netdev/netdev.go @@ -4,7 +4,7 @@ package netdev import ( "errors" - "net" + "net/netip" "time" _ "unsafe" // to use go:linkname ) @@ -28,6 +28,7 @@ const ( // GethostByName() errors var ( ErrHostUnknown = errors.New("Host unknown") + ErrMalAddr = errors.New("Malformed address") ) // Socket errors @@ -70,18 +71,18 @@ type Netdever interface { // GetHostByName returns the IP address of either a hostname or IPv4 // address in standard dot notation - GetHostByName(name string) (net.IP, error) + GetHostByName(name string) (netip.Addr, error) - // GetIPAddr returns IP address assigned to the interface, either by + // Addr returns IP address assigned to the interface, either by // DHCP or statically - GetIPAddr() (net.IP, error) + Addr() (netip.Addr, error) // Berkely Sockets-like interface, Go-ified. See man page for socket(2), etc. Socket(domain int, stype int, protocol int) (int, error) - Bind(sockfd int, ip net.IP, port int) error - Connect(sockfd int, host string, ip net.IP, port int) error + Bind(sockfd int, ip netip.AddrPort) error + Connect(sockfd int, host string, ip netip.AddrPort) error Listen(sockfd int, backlog int) error - Accept(sockfd int, ip net.IP, port int) (int, error) + Accept(sockfd int, ip netip.AddrPort) (int, error) Send(sockfd int, buf []byte, flags int, deadline time.Time) (int, error) Recv(sockfd int, buf []byte, flags int, deadline time.Time) (int, error) Close(sockfd int) error diff --git a/rtl8720dn/rtl8720dn.go b/rtl8720dn/rtl8720dn.go index fa0f384..75cd8c8 100644 --- a/rtl8720dn/rtl8720dn.go +++ b/rtl8720dn/rtl8720dn.go @@ -12,6 +12,7 @@ import ( "io" "machine" "net" + "net/netip" "strings" "sync" "time" @@ -195,9 +196,9 @@ func (r *rtl8720dn) showIP() { if debugging(debugBasic) { ip, subnet, gateway, _ := r.getIP() fmt.Printf("\r\n") - fmt.Printf("DHCP-assigned IP : %s\r\n", ip.String()) - fmt.Printf("DHCP-assigned subnet : %s\r\n", subnet.String()) - fmt.Printf("DHCP-assigned gateway : %s\r\n", gateway.String()) + fmt.Printf("DHCP-assigned IP : %s\r\n", ip) + fmt.Printf("DHCP-assigned subnet : %s\r\n", subnet) + fmt.Printf("DHCP-assigned gateway : %s\r\n", gateway) fmt.Printf("\r\n") } } @@ -315,7 +316,7 @@ func (r *rtl8720dn) NetNotify(cb func(netlink.Event)) { r.notifyCb = cb } -func (r *rtl8720dn) GetHostByName(name string) (net.IP, error) { +func (r *rtl8720dn) GetHostByName(name string) (netip.Addr, error) { if debugging(debugNetdev) { fmt.Printf("[GetHostByName] name: %s\r\n", name) @@ -327,10 +328,15 @@ func (r *rtl8720dn) GetHostByName(name string) (net.IP, error) { var ip [4]byte result := r.rpc_netconn_gethostbyname(name, ip[:]) if result == -1 { - return net.IP{}, netdev.ErrHostUnknown + return netip.Addr{}, netdev.ErrHostUnknown } - return net.IP(ip[:]), nil + addr, ok := netip.AddrFromSlice(ip[:]) + if !ok { + return netip.Addr{}, netdev.ErrMalAddr + } + + return addr, nil } func (r *rtl8720dn) GetHardwareAddr() (net.HardwareAddr, error) { @@ -348,7 +354,7 @@ func (r *rtl8720dn) GetHardwareAddr() (net.HardwareAddr, error) { return net.HardwareAddr(addr), err } -func (r *rtl8720dn) GetIPAddr() (net.IP, error) { +func (r *rtl8720dn) Addr() (netip.Addr, error) { if debugging(debugNetdev) { fmt.Printf("[GetIPAddr]\r\n") @@ -359,7 +365,7 @@ func (r *rtl8720dn) GetIPAddr() (net.IP, error) { ip, _, _, err := r.getIP() - return net.IP(ip), err + return ip, err } func (r *rtl8720dn) clientTLS() uint32 { @@ -415,26 +421,26 @@ func (r *rtl8720dn) Socket(domain int, stype int, protocol int) (int, error) { return int(newSock), nil } -func addrToName(ip net.IP, port int) []byte { +func ipToName(ip netip.AddrPort) []byte { name := make([]byte, 16) name[0] = 0x00 name[1] = netdev.AF_INET - name[2] = byte(port >> 8) - name[3] = byte(port) - if len(ip) == 4 { - name[4] = byte(ip[0]) - name[5] = byte(ip[1]) - name[6] = byte(ip[2]) - name[7] = byte(ip[3]) + name[2] = byte(ip.Port() >> 8) + name[3] = byte(ip.Port()) + if ip.Addr().Is4() { + addr := ip.Addr().As4() + name[4] = byte(addr[0]) + name[5] = byte(addr[1]) + name[6] = byte(addr[2]) + name[7] = byte(addr[3]) } - return name } -func (r *rtl8720dn) Bind(sockfd int, ip net.IP, port int) error { +func (r *rtl8720dn) Bind(sockfd int, ip netip.AddrPort) error { if debugging(debugNetdev) { - fmt.Printf("[Bind] sockfd: %d, addr: %s:%d\r\n", sockfd, ip, port) + fmt.Printf("[Bind] sockfd: %d, addr: %s\r\n", sockfd, ip) } r.mu.Lock() @@ -442,13 +448,13 @@ func (r *rtl8720dn) Bind(sockfd int, ip net.IP, port int) error { var sock = sock(sockfd) var socket = r.sockets[sock] - var name = addrToName(ip, port) + var name = ipToName(ip) switch socket.protocol { case netdev.IPPROTO_TCP, netdev.IPPROTO_UDP: result := r.rpc_lwip_bind(int32(sock), name, uint32(len(name))) if result == -1 { - return fmt.Errorf("Bind to %s:%d failed", ip, port) + return fmt.Errorf("Bind to %s failed", ip) } default: return netdev.ErrProtocolNotSupported @@ -457,11 +463,13 @@ func (r *rtl8720dn) Bind(sockfd int, ip net.IP, port int) error { return nil } -func (r *rtl8720dn) Connect(sockfd int, host string, ip net.IP, port int) error { +func (r *rtl8720dn) Connect(sockfd int, host string, ip netip.AddrPort) error { + + port := ip.Port() if debugging(debugNetdev) { if host == "" { - fmt.Printf("[Connect] sockfd: %d, addr: %s:%d\r\n", sockfd, ip, port) + fmt.Printf("[Connect] sockfd: %d, addr: %s\r\n", sockfd, ip) } else { fmt.Printf("[Connect] sockfd: %d, host: %s:%d\r\n", sockfd, host, port) } @@ -472,14 +480,14 @@ func (r *rtl8720dn) Connect(sockfd int, host string, ip net.IP, port int) error var sock = sock(sockfd) var socket = r.sockets[sock] - var name = addrToName(ip, port) + var name = ipToName(ip) // Start the connection switch socket.protocol { case netdev.IPPROTO_TCP, netdev.IPPROTO_UDP: result := r.rpc_lwip_connect(int32(sock), name, uint32(len(name))) if result == -1 { - return fmt.Errorf("Connect to %s:%d failed", ip, port) + return fmt.Errorf("Connect to %s failed", ip) } case netdev.IPPROTO_TLS: result := r.rpc_wifi_start_ssl_client(uint32(sock), @@ -526,10 +534,10 @@ func (r *rtl8720dn) Listen(sockfd int, backlog int) error { return nil } -func (r *rtl8720dn) Accept(sockfd int, ip net.IP, port int) (int, error) { +func (r *rtl8720dn) Accept(sockfd int, ip netip.AddrPort) (int, error) { if debugging(debugNetdev) { - fmt.Printf("[Accept] sockfd: %d, peer: %s:%d\r\n", sockfd, ip, port) + fmt.Printf("[Accept] sockfd: %d, peer: %s\r\n", sockfd, ip) } r.mu.Lock() @@ -538,7 +546,7 @@ func (r *rtl8720dn) Accept(sockfd int, ip net.IP, port int) (int, error) { var newSock int32 var lsock = sock(sockfd) var socket = r.sockets[lsock] - var addr = addrToName(ip, port) + var name = ipToName(ip) switch socket.protocol { case netdev.IPPROTO_TCP: @@ -554,8 +562,8 @@ func (r *rtl8720dn) Accept(sockfd int, ip net.IP, port int) (int, error) { r.mu.Lock() // Check if a client connected. O_NONBLOCK is set on lsock. - addrlen := uint32(len(addr)) - newSock = r.rpc_lwip_accept(int32(lsock), addr, &addrlen) + namelen := uint32(len(name)) + newSock = r.rpc_lwip_accept(int32(lsock), name, &namelen) if newSock == -1 { // No new client time.Sleep(100 * time.Millisecond) @@ -761,16 +769,15 @@ func (r *rtl8720dn) getMACAddr() string { return string(mac[:]) } -func (r *rtl8720dn) getIP() (ip, subnet, gateway net.IP, err error) { +func (r *rtl8720dn) getIP() (ip, subnet, gateway netip.Addr, err error) { var ip_info [12]byte result := r.rpc_tcpip_adapter_get_ip_info(0, ip_info[:]) if result == -1 { err = fmt.Errorf("Get IP info failed") return } - ip, subnet, gateway = make([]byte, 4), make([]byte, 4), make([]byte, 4) - copy(ip[:], ip_info[0:4]) - copy(subnet[:], ip_info[4:8]) - copy(gateway[:], ip_info[8:12]) + ip, _ = netip.AddrFromSlice(ip_info[0:4]) + subnet, _ = netip.AddrFromSlice(ip_info[4:8]) + gateway, _ = netip.AddrFromSlice(ip_info[8:12]) return } diff --git a/wifinina/wifinina.go b/wifinina/wifinina.go index efcc765..1f2d724 100644 --- a/wifinina/wifinina.go +++ b/wifinina/wifinina.go @@ -17,6 +17,7 @@ import ( "machine" "math/bits" "net" + "net/netip" "sync" "time" @@ -164,8 +165,7 @@ type hwerr uint8 type socket struct { protocol int - ip net.IP - port int + ip netip.AddrPort inuse bool } @@ -375,9 +375,9 @@ func (w *wifinina) showIP() { if debugging(debugBasic) { ip, subnet, gateway := w.getIP() fmt.Printf("\r\n") - fmt.Printf("DHCP-assigned IP : %s\r\n", ip.String()) - fmt.Printf("DHCP-assigned subnet : %s\r\n", subnet.String()) - fmt.Printf("DHCP-assigned gateway : %s\r\n", gateway.String()) + fmt.Printf("DHCP-assigned IP : %s\r\n", ip) + fmt.Printf("DHCP-assigned subnet : %s\r\n", subnet) + fmt.Printf("DHCP-assigned gateway : %s\r\n", gateway) fmt.Printf("\r\n") } } @@ -498,7 +498,7 @@ func (w *wifinina) NetNotify(cb func(netlink.Event)) { w.notifyCb = cb } -func (w *wifinina) GetHostByName(name string) (net.IP, error) { +func (w *wifinina) GetHostByName(name string) (netip.Addr, error) { if debugging(debugNetdev) { fmt.Printf("[GetHostByName] name: %s\r\n", name) @@ -509,10 +509,15 @@ func (w *wifinina) GetHostByName(name string) (net.IP, error) { ip := w.getHostByName(name) if ip == "" { - return net.IP{}, netdev.ErrHostUnknown + return netip.Addr{}, netdev.ErrHostUnknown } - return net.IP([]byte(ip)), nil + addr, ok := netip.AddrFromSlice([]byte(ip)) + if !ok { + return netip.Addr{}, netdev.ErrMalAddr + } + + return addr, nil } func (w *wifinina) GetHardwareAddr() (net.HardwareAddr, error) { @@ -527,7 +532,7 @@ func (w *wifinina) GetHardwareAddr() (net.HardwareAddr, error) { return w.getMACAddr(), nil } -func (w *wifinina) GetIPAddr() (net.IP, error) { +func (w *wifinina) Addr() (netip.Addr, error) { if debugging(debugNetdev) { fmt.Printf("[GetIPAddr]\r\n") @@ -538,7 +543,7 @@ func (w *wifinina) GetIPAddr() (net.IP, error) { ip, _, _ := w.getIP() - return net.IP(ip), nil + return ip, nil } // See man socket(2) for standard Berkely sockets for Socket, Bind, etc. @@ -579,10 +584,10 @@ func (w *wifinina) Socket(domain int, stype int, protocol int) (int, error) { return int(sock), nil } -func (w *wifinina) Bind(sockfd int, ip net.IP, port int) error { +func (w *wifinina) Bind(sockfd int, ip netip.AddrPort) error { if debugging(debugNetdev) { - fmt.Printf("[Bind] sockfd: %d, addr: %s:%d\r\n", sockfd, ip, port) + fmt.Printf("[Bind] sockfd: %d, addr: %s\r\n", sockfd, ip) } w.mu.Lock() @@ -595,29 +600,28 @@ func (w *wifinina) Bind(sockfd int, ip net.IP, port int) error { case netdev.IPPROTO_TCP: case netdev.IPPROTO_TLS: case netdev.IPPROTO_UDP: - w.startServer(sock, uint16(port), protoModeUDP) + w.startServer(sock, ip.Port(), protoModeUDP) } - socket.ip, socket.port = ip, port + socket.ip = ip return nil } -func toUint32(ip net.IP) uint32 { - ip = ip.To4() +func toUint32(ip [4]byte) uint32 { return uint32(ip[0])<<24 | uint32(ip[1])<<16 | uint32(ip[2])<<8 | uint32(ip[3]) } -func (w *wifinina) Connect(sockfd int, host string, ip net.IP, port int) error { +func (w *wifinina) Connect(sockfd int, host string, ip netip.AddrPort) error { if debugging(debugNetdev) { if host == "" { - fmt.Printf("[Connect] sockfd: %d, addr: %s:%d\r\n", sockfd, ip, port) + fmt.Printf("[Connect] sockfd: %d, addr: %s\r\n", sockfd, ip) } else { - fmt.Printf("[Connect] sockfd: %d, host: %s:%d\r\n", sockfd, host, port) + fmt.Printf("[Connect] sockfd: %d, host: %s:%d\r\n", sockfd, host, ip.Port()) } } @@ -630,11 +634,11 @@ func (w *wifinina) Connect(sockfd int, host string, ip net.IP, port int) error { // Start the connection switch socket.protocol { case netdev.IPPROTO_TCP: - w.startClient(sock, "", toUint32(ip), uint16(port), protoModeTCP) + w.startClient(sock, "", toUint32(ip.Addr().As4()), ip.Port(), protoModeTCP) case netdev.IPPROTO_TLS: - w.startClient(sock, host, 0, uint16(port), protoModeTLS) + w.startClient(sock, host, 0, ip.Port(), protoModeTLS) case netdev.IPPROTO_UDP: - w.startClient(sock, "", toUint32(ip), uint16(port), protoModeUDP) + w.startClient(sock, "", toUint32(ip.Addr().As4()), ip.Port(), protoModeUDP) return nil } @@ -643,9 +647,9 @@ func (w *wifinina) Connect(sockfd int, host string, ip net.IP, port int) error { } if host == "" { - return fmt.Errorf("Connect to %s:%d failed", ip, port) + return fmt.Errorf("Connect to %s failed", ip) } else { - return fmt.Errorf("Connect to %s:%d failed", host, port) + return fmt.Errorf("Connect to %s:%d failed", host, ip.Port()) } } @@ -663,7 +667,7 @@ func (w *wifinina) Listen(sockfd int, backlog int) error { switch socket.protocol { case netdev.IPPROTO_TCP: - w.startServer(sock, uint16(socket.port), protoModeTCP) + w.startServer(sock, socket.ip.Port(), protoModeTCP) case netdev.IPPROTO_UDP: default: return netdev.ErrProtocolNotSupported @@ -672,10 +676,10 @@ func (w *wifinina) Listen(sockfd int, backlog int) error { return nil } -func (w *wifinina) Accept(sockfd int, ip net.IP, port int) (int, error) { +func (w *wifinina) Accept(sockfd int, ip netip.AddrPort) (int, error) { if debugging(debugNetdev) { - fmt.Printf("[Accept] sockfd: %d, peer: %s:%d\r\n", sockfd, ip, port) + fmt.Printf("[Accept] sockfd: %d, peer: %s\r\n", sockfd, ip) } w.mu.Lock() @@ -1210,7 +1214,7 @@ func (w *wifinina) faultf(f string, args ...any) { } } -func (w *wifinina) getIP() (ip, subnet, gateway net.IP) { +func (w *wifinina) getIP() (ip, subnet, gateway netip.Addr) { if debugging(debugCmd) { fmt.Printf(" [cmdGetIPAddr]\r\n") } @@ -1220,10 +1224,9 @@ func (w *wifinina) getIP() (ip, subnet, gateway net.IP) { w.faultf("getIP wanted l=3, got l=%d", l) return } - ip, subnet, gateway = make([]byte, 4), make([]byte, 4), make([]byte, 4) - copy(ip[:], sl[0]) - copy(subnet[:], sl[1]) - copy(gateway[:], sl[2]) + ip, _ = netip.AddrFromSlice([]byte(sl[0])[:4]) + subnet, _ = netip.AddrFromSlice([]byte(sl[1])[:4]) + gateway, _ = netip.AddrFromSlice([]byte(sl[2])[:4]) return }