hyprspace/node/node.go
2026-06-19 00:10:04 +02:00

522 lines
13 KiB
Go

package node
import (
"context"
"encoding/binary"
"errors"
"fmt"
"io/fs"
"net"
"net/http"
"os"
"path/filepath"
"sync"
"time"
"github.com/hyprspace/hyprspace/config"
hsdns "github.com/hyprspace/hyprspace/dns"
"github.com/hyprspace/hyprspace/p2p"
hsrpc "github.com/hyprspace/hyprspace/rpc"
"github.com/hyprspace/hyprspace/svc"
"github.com/hyprspace/hyprspace/tun"
"github.com/ipfs/go-log/v2"
"github.com/libp2p/go-cidranger"
dht "github.com/libp2p/go-libp2p-kad-dht"
"github.com/libp2p/go-libp2p/core/connmgr"
"github.com/libp2p/go-libp2p/core/event"
"github.com/libp2p/go-libp2p/core/host"
"github.com/libp2p/go-libp2p/core/network"
"github.com/libp2p/go-libp2p/core/peer"
"github.com/multiformats/go-multiaddr"
"github.com/prometheus/client_golang/prometheus/promhttp"
"go.uber.org/zap"
)
var logger = log.Logger("hyprspace/node")
type SharedStream struct {
Stream *network.Stream
Lock *sync.Mutex
}
type Node struct {
cfg *config.Config
p2p host.Host
dht *dht.IpfsDHT
tunDev *tun.TUN
activeStreams map[peer.ID]SharedStream
activeStreamsLock sync.RWMutex
ctx context.Context
cancel func()
lockPath string
configPath string
interfaceName string
wg *sync.WaitGroup
}
func New(ctx context.Context, configPath string, ifName string) Node {
innerCtx, ctxCancel := context.WithCancel(ctx)
return Node{
cfg: &config.Config{},
p2p: nil,
tunDev: &tun.TUN{},
ctx: innerCtx,
cancel: ctxCancel,
configPath: configPath,
interfaceName: ifName,
}
}
func (node *Node) Run() error {
// Read in configuration from file.
cfg2, err := config.Read(node.configPath)
if err != nil {
logger.With(err).Error("Failed to read config")
return err
}
cfg2.Interface = node.interfaceName
node.cfg = cfg2
logger.Info("Creating TUN Device")
// Create new TUN device
node.tunDev, err = tun.New(
node.cfg.Interface,
tun.Address(node.cfg.BuiltinAddr4.String()+"/32"),
tun.Address(node.cfg.BuiltinAddr6.String()+"/128"),
tun.MTU(1420),
)
if err != nil {
logger.With(err).Error("Failed to create TUN Device")
return err
}
allRoutes4, err := node.cfg.PeerLookup.ByRoute.CoveredNetworks(*cidranger.AllIPv4)
if err != nil {
logger.With(err).Error("Failed to lookup IPv4 peer-routes")
return err
}
allRoutes6, err := node.cfg.PeerLookup.ByRoute.CoveredNetworks(*cidranger.AllIPv6)
if err != nil {
logger.With(err).Error("Failed to lookup IPv6 peer-routes")
return err
}
var routeOpts []tun.Option
for _, r := range allRoutes4 {
routeOpts = append(routeOpts, tun.Route(r.Network()))
}
for _, r := range allRoutes6 {
routeOpts = append(routeOpts, tun.Route(r.Network()))
}
recursionGater := p2p.NewRecursionGater(node.cfg)
var gater connmgr.ConnectionGater
if node.cfg.FilterPrivateAddresses {
gater = p2p.NewMultiGater(
recursionGater,
p2p.NewFilterGater(
// IPv4 local
parseCIDR("10.0.0.0/8"),
parseCIDR("172.16.0.0/12"),
parseCIDR("192.168.0.0/16"),
// IPv4 link-local
parseCIDR("169.254.0.0/16"),
// IPv4 loopback
parseCIDR("127.0.0.0/8"),
// IPv6 link-local
parseCIDR("fe80::/10"),
// IPv6 loopback
parseCIDR("::1/128"),
),
)
} else {
gater = recursionGater
}
logger.Info("Creating LibP2P node")
// Create P2P Node
node.p2p, node.dht, err = p2p.CreateNode(
node.ctx,
node.cfg.PrivateKey,
node.cfg.ListenAddresses,
node.streamHandler,
p2p.NewClosedCircuitRelayFilter(node.cfg.Peers),
gater,
node.cfg.Peers,
)
if err != nil {
logger.With(err).Error("Failed to create Libp2p node")
return err
}
node.p2p.SetStreamHandler(p2p.PeXProtocol, p2p.NewPeXStreamHandler(node.p2p, node.cfg))
for _, p := range node.cfg.Peers {
node.p2p.ConnManager().Protect(p.ID, "/hyprspace/peer")
}
node.wg = &sync.WaitGroup{}
logger.Debug("Setting up Node discovery via DHT")
// Setup P2P Discovery
go p2p.Discover(node.ctx, node.wg, node.p2p, node.dht, node.cfg.Peers)
// Configure path for lock
node.lockPath = filepath.Join(filepath.Dir(node.cfg.Path), node.cfg.Interface+".lock")
logger.Debug("Starting Peer-Exchange service")
// PeX
go p2p.PeXService(node.ctx, node.wg, node.p2p, node.cfg)
logger.Debug("Starting Route Metrics service")
// Route metrics and latency
go p2p.RouteMetricsService(node.ctx, node.wg, node.p2p, node.cfg)
// Log about various events
err = node.eventLogger(node.ctx, node.p2p)
if err != nil {
logger.With(err).Error("Failed to subscribe to EventBus")
return err
}
logger.Debug("Starting RPC server")
// RPC server
go hsrpc.RpcServer(node.ctx, node.wg, multiaddr.StringCast(fmt.Sprintf("/unix/run/hyprspace-rpc.%s.sock", node.cfg.Interface)), node.p2p, *node.cfg, *node.tunDev)
logger.Debug("Starting DNS server")
// Magic DNS server
go hsdns.MagicDnsServer(node.ctx, node.wg, *node.cfg, node.p2p)
// metrics endpoint
metricsPort, ok := os.LookupEnv("HYPRSPACE_METRICS_PORT")
if ok {
metricsTuple := fmt.Sprintf("127.0.0.1:%s", metricsPort)
http.Handle("/metrics", promhttp.Handler())
go func() {
logger.Debug("Starting metrics API server")
http.ListenAndServe(metricsTuple, nil)
}()
logger.Info(fmt.Sprintf("Listening for metrics scrape requests on http://%s/metrics", metricsTuple))
}
serviceNet := svc.NewServiceNetwork(node.p2p, node.cfg, node.tunDev)
for name, service := range node.cfg.Services {
proxy, err := svc.ProxyTo(service.Target)
if err != nil {
return err
}
serviceNet.Register(
name,
proxy,
)
}
var svcNetIds [][4]byte
for _, p := range node.cfg.Peers {
svcNetIds = append(svcNetIds, config.MkNetID(p.ID))
}
svcNetIds = append(svcNetIds, config.MkNetID(node.p2p.ID()))
for _, netId := range svcNetIds {
addr := make([]byte, 16)
copy(addr, serviceNet.NetworkRange.IP)
copy(addr[10:], netId[:])
mask1, mask0 := serviceNet.NetworkRange.Mask.Size()
routeOpts = append(routeOpts, tun.Route(net.IPNet{
IP: addr,
Mask: net.CIDRMask(mask1+32, mask0),
}))
}
// Write lock to filesystem to indicate an existing running daemon.
err = os.WriteFile(node.lockPath, fmt.Append(nil, os.Getpid()), os.ModePerm)
if err != nil {
return err
}
logger.Debug("Bringing up TUN device")
// Bring Up TUN Device
err = node.tunDev.Up()
if err != nil {
logger.With(err).Error("Failed to bring TUN device up")
return errors.New("unable to bring up tun device: " + err.Error())
}
err = node.tunDev.Apply(routeOpts...)
if err != nil {
return errors.New("unable to apply routing options: " + err.Error())
}
logger.Info("Network setup complete")
// Initialize active streams map and packet byte array.
node.activeStreams = make(map[peer.ID]SharedStream)
go func() {
for {
var packet = make([]byte, 1420)
// Read in a packet from the tun device.
plen, err := node.tunDev.Iface.Read(packet)
if errors.Is(err, fs.ErrClosed) {
logger.Warn("Interface closed")
<-node.ctx.Done()
time.Sleep(1 * time.Second)
return
} else if err != nil {
logger.With(err).Error("Failed to read from interface")
continue
}
var dstIP net.IP
proto := packet[0] & 0xf0
if proto == 0x40 {
dstIP = net.IP(packet[16:20])
if node.cfg.BuiltinAddr4.Equal(dstIP) {
continue
}
} else if proto == 0x60 {
dstIP = net.IP(packet[24:40])
if node.cfg.BuiltinAddr6.Equal(dstIP) {
continue
} else if serviceNet.NetworkRange.Contains(dstIP) {
// Are you TCP because your protocol is 6, or is your protocol 6 because you are TCP?
if packet[6] == 0x06 {
port := uint16(packet[42])*256 + uint16(packet[43])
if serviceNet.EnsureListener([16]byte(packet[24:40]), port) {
count, err := (*serviceNet.Tun).Write([][]byte{packet}, 0)
if count == 0 || err != nil {
logger.With(err).Error("Error writing to service-network tunnel")
}
}
}
continue
}
} else {
continue
}
var dst peer.ID
// Check route table for destination address.
route, found := node.cfg.FindRouteForIP(dstIP)
if found {
dst = route.Target.ID
go node.sendPacket(dst, packet, plen)
}
}
}()
return nil
}
func (node *Node) getActiveStream(pid peer.ID) (SharedStream, bool) {
node.activeStreamsLock.RLock()
defer node.activeStreamsLock.RUnlock()
s, ok := node.activeStreams[pid]
return s, ok
}
func (node *Node) insertActiveStream(pid peer.ID, ss SharedStream) bool {
node.activeStreamsLock.Lock()
if _, exists := node.activeStreams[pid]; exists {
node.activeStreamsLock.Unlock()
return false
}
node.activeStreams[pid] = ss
node.activeStreamsLock.Unlock()
return true
}
func (node *Node) expireActiveStream(pid peer.ID) {
node.activeStreamsLock.Lock()
delete(node.activeStreams, pid)
node.activeStreamsLock.Unlock()
}
func (node *Node) streamHandler(stream network.Stream) {
remotePeerID := stream.Conn().RemotePeer()
// If the remote node ID isn't in the list of known nodes don't respond.
if _, ok := config.FindPeer(node.cfg.Peers, remotePeerID); !ok {
stream.Reset()
return
}
var streamLock = new(sync.Mutex)
// Version 0 nodes don't read from this stream, so we can't reuse it.
if stream.Protocol() != p2p.ProtocolV0 {
inserted := node.insertActiveStream(remotePeerID, SharedStream{
Stream: &stream,
Lock: streamLock,
})
if inserted {
defer node.expireActiveStream(remotePeerID)
}
}
var packet = make([]byte, 1420)
var packetSize = make([]byte, 2)
for {
// Read the incoming packet's size as a binary value.
_, err := stream.Read(packetSize)
if err != nil {
stream.Close()
return
}
// Decode the incoming packet's size from binary.
size := binary.LittleEndian.Uint16(packetSize)
// Read in the packet until completion.
var plen uint16 = 0
for plen < size {
tmp, err := stream.Read(packet[plen:size])
plen += uint16(tmp)
if err != nil {
stream.Close()
return
}
}
streamLock.Lock()
err = stream.SetWriteDeadline(time.Now().Add(25 * time.Second))
streamLock.Unlock()
if err != nil {
logger.With(err).Error("Failed to set write deadline")
stream.Close()
return
}
_, _ = node.tunDev.Iface.Write(packet[:size])
}
}
func (node *Node) sendPacket(dst peer.ID, packet []byte, plen int) {
// Check if we already have an open connection to the destination peer.
ms, ok := node.getActiveStream(dst)
if ok {
if func() bool {
ms.Lock.Lock()
defer ms.Lock.Unlock()
// Write out the packet's length to the libp2p stream to ensure
// we know the full size of the packet at the other end.
err := binary.Write(*ms.Stream, binary.LittleEndian, uint16(plen))
if err == nil {
// Write the packet out to the libp2p stream.
// If everyting succeeds continue on to the next packet.
_, err = (*ms.Stream).Write(packet[:plen])
if err == nil {
err := (*ms.Stream).SetWriteDeadline(time.Now().Add(25 * time.Second))
if err == nil {
return true
}
}
}
// If we encounter an error when writing to a stream we should
// close that stream and delete it from the active stream map.
(*ms.Stream).Close()
node.expireActiveStream(dst)
return false
}() {
return
}
}
stream, err := node.p2p.NewStream(node.ctx, dst, p2p.Protocols...)
if err != nil {
logger.With(zap.String("destination", dst.String()), zap.Error(err)).Error("Failed to open stream")
go p2p.Rediscover()
return
}
err = stream.SetWriteDeadline(time.Now().Add(25 * time.Second))
if err != nil {
logger.With(err).Error("Failed to set write deadline")
stream.Close()
return
}
// Write packet length
err = binary.Write(stream, binary.LittleEndian, uint16(plen))
if err != nil {
stream.Close()
return
}
// Write the packet
_, err = stream.Write(packet[:plen])
if err != nil {
stream.Close()
return
}
go node.streamHandler(stream)
}
func (node *Node) eventLogger(ctx context.Context, host host.Host) error {
subCon, err := host.EventBus().Subscribe(new(event.EvtPeerConnectednessChanged))
if err != nil {
return err
}
go func() {
for {
select {
case <-ctx.Done():
node.wg.Done()
return
case ev := <-subCon.Out():
evt := ev.(event.EvtPeerConnectednessChanged)
for _, vpnPeer := range node.cfg.Peers {
if vpnPeer.ID == evt.Peer {
switch evt.Connectedness {
case network.Connected:
for _, c := range host.Network().ConnsToPeer(evt.Peer) {
logger.Info(fmt.Sprintf("Connected to %s/p2p/%s", c.RemoteMultiaddr().String(), evt.Peer.String()))
}
case network.NotConnected:
logger.Info(fmt.Sprintf("Disconnected from %s", evt.Peer.String()))
}
break
}
}
}
}
}()
return nil
}
func (node *Node) Rebootstrap() {
node.p2p.ConnManager().TrimOpenConns(context.Background())
<-node.dht.ForceRefresh()
p2p.Rediscover()
}
func (node *Node) Stop() error {
err := node.p2p.Close()
if err != nil {
return err
}
err = os.Remove(node.lockPath)
if err != nil {
return err
}
logger.Info("Received signal, shutting down...")
err = node.tunDev.Down()
if err != nil {
return err
}
node.tunDev.Iface.Close()
node.cancel()
node.wg.Wait()
return nil
}
func parseCIDR(s string) net.IPNet {
_, n, err := net.ParseCIDR(s)
if err != nil {
panic("invalid CIDR: " + s)
}
return *n
}