initial: ZeroTier-like P2P mesh VPN server with multi-tenant Web UI
This commit is contained in:
324
internal/agent/agent.go
Normal file
324
internal/agent/agent.go
Normal file
@@ -0,0 +1,324 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"zeromesh/internal/config"
|
||||
"zeromesh/internal/identity"
|
||||
"zeromesh/internal/tap"
|
||||
"zeromesh/internal/vl1"
|
||||
"zeromesh/internal/vl2"
|
||||
)
|
||||
|
||||
type Agent struct {
|
||||
cfg config.AgentConfig
|
||||
identity *identity.Identity
|
||||
transport *vl1.Transport
|
||||
peers *vl1.PeerManager
|
||||
network *vl2.Network
|
||||
tapDev *tap.Interface
|
||||
log *slog.Logger
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
func New(cfg config.AgentConfig, log *slog.Logger) (*Agent, error) {
|
||||
id, err := identity.LoadOrGenerate("./data/agent.identity")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("load identity: %w", err)
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
return &Agent{
|
||||
cfg: cfg,
|
||||
identity: id,
|
||||
peers: vl1.NewPeerManager(log),
|
||||
log: log.With("component", "agent"),
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (a *Agent) Start() error {
|
||||
transport, err := vl1.NewTransport(a.cfg.ListenPort, a.log)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
a.transport = transport
|
||||
|
||||
netConfig := vl2.NetworkConfig{
|
||||
ID: a.cfg.NetworkID,
|
||||
Name: "default",
|
||||
MTU: a.cfg.TAPMTU,
|
||||
Multicast: true,
|
||||
}
|
||||
a.network = vl2.NewNetwork(netConfig, a.identity.Address, a, a.log)
|
||||
|
||||
if a.cfg.ControllerURL != "" {
|
||||
a.log.Info("agent started (controller mode)",
|
||||
"address", a.identity.Address,
|
||||
"port", a.transport.Port(),
|
||||
"controller", a.cfg.ControllerURL)
|
||||
}
|
||||
|
||||
if a.cfg.TAPName != "" {
|
||||
tapName := a.cfg.TAPName
|
||||
if tapName == "" {
|
||||
tapName = fmt.Sprintf("zm%x", a.identity.Address[:4])
|
||||
}
|
||||
tapDev, err := tap.Open(tapName, a.cfg.TAPMTU)
|
||||
if err != nil {
|
||||
a.log.Warn("TAP device not available (run with sufficient privileges)", "err", err)
|
||||
} else {
|
||||
a.tapDev = tapDev
|
||||
a.log.Info("TAP device created", "name", tapName, "mtu", a.cfg.TAPMTU)
|
||||
a.wg.Add(1)
|
||||
go a.tapReadLoop()
|
||||
}
|
||||
}
|
||||
|
||||
a.wg.Add(1)
|
||||
go a.udpReadLoop()
|
||||
|
||||
a.wg.Add(1)
|
||||
go a.maintenanceLoop()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *Agent) Stop() {
|
||||
a.cancel()
|
||||
if a.transport != nil {
|
||||
a.transport.Close()
|
||||
}
|
||||
if a.tapDev != nil {
|
||||
a.tapDev.Close()
|
||||
}
|
||||
a.wg.Wait()
|
||||
}
|
||||
|
||||
func (a *Agent) Identity() *identity.Identity {
|
||||
return a.identity
|
||||
}
|
||||
|
||||
func (a *Agent) Peers() *vl1.PeerManager {
|
||||
return a.peers
|
||||
}
|
||||
|
||||
func (a *Agent) Network() *vl2.Network {
|
||||
return a.network
|
||||
}
|
||||
|
||||
func (a *Agent) SendToPeer(peerAddr identity.Address, networkID uint32, frame []byte) error {
|
||||
peer := a.peers.GetPeer(peerAddr)
|
||||
if peer == nil {
|
||||
return fmt.Errorf("unknown peer: %s", peerAddr)
|
||||
}
|
||||
if !peer.IsConnected() {
|
||||
return fmt.Errorf("peer not connected: %s", peerAddr)
|
||||
}
|
||||
|
||||
pkt := vl1.NewDataPacket(networkID, frame)
|
||||
encoded := pkt.Encode()
|
||||
|
||||
// Encrypt payload portion (after header)
|
||||
encrypted, err := peer.Encrypt(encoded[vl1.HeaderSize:])
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
encPacket := vl1.NewDataPacket(networkID, encrypted)
|
||||
if peer.Endpoint == nil {
|
||||
return fmt.Errorf("peer %s: no endpoint", peerAddr)
|
||||
}
|
||||
return a.transport.SendPacket(&encPacket, peer.Endpoint)
|
||||
}
|
||||
|
||||
func (a *Agent) BroadcastToPeers(networkID uint32, frame []byte, excludePeer identity.Address) error {
|
||||
for _, peer := range a.peers.ConnectedPeers() {
|
||||
if peer.Address == excludePeer {
|
||||
continue
|
||||
}
|
||||
pkt := vl1.NewDataPacket(networkID, frame)
|
||||
encoded := pkt.Encode()
|
||||
encrypted, err := peer.Encrypt(encoded[vl1.HeaderSize:])
|
||||
if err != nil {
|
||||
a.log.Debug("encrypt for broadcast", "peer", peer.Address, "err", err)
|
||||
continue
|
||||
}
|
||||
encPacket := vl1.NewDataPacket(networkID, encrypted)
|
||||
if peer.Endpoint != nil {
|
||||
if err := a.transport.SendPacket(&encPacket, peer.Endpoint); err != nil {
|
||||
a.log.Debug("broadcast send", "peer", peer.Address, "err", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *Agent) udpReadLoop() {
|
||||
defer a.wg.Done()
|
||||
buf := make([]byte, vl1.MaxPacketSize)
|
||||
for {
|
||||
select {
|
||||
case <-a.ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
n, remoteAddr, err := a.transport.ReadFrom(buf)
|
||||
if err != nil {
|
||||
if a.ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
a.log.Error("UDP read error", "err", err)
|
||||
time.Sleep(time.Millisecond)
|
||||
continue
|
||||
}
|
||||
a.handlePacket(buf[:n], remoteAddr)
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) handlePacket(data []byte, from *net.UDPAddr) {
|
||||
pkt, err := vl1.DecodePacket(data)
|
||||
if err != nil {
|
||||
a.log.Debug("decode packet", "err", err)
|
||||
return
|
||||
}
|
||||
|
||||
switch pkt.Header.Type {
|
||||
case vl1.PacketTypeHandshake:
|
||||
a.handleHandshake(pkt.Payload, from)
|
||||
case vl1.PacketTypeData:
|
||||
a.handleDataPacket(pkt, from)
|
||||
case vl1.PacketTypeKeepalive:
|
||||
if peer := a.peers.GetPeerByEndpoint(from); peer != nil {
|
||||
peer.Touch()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) handleHandshake(payload []byte, from *net.UDPAddr) {
|
||||
var pubKey [32]byte
|
||||
if len(payload) < 32 {
|
||||
return
|
||||
}
|
||||
copy(pubKey[:], payload[:32])
|
||||
addr := identity.AddressFromPublicKey(pubKey[:])
|
||||
|
||||
peer := a.peers.GetPeer(addr)
|
||||
if peer != nil {
|
||||
a.peers.UpdatePeerEndpoint(addr, from)
|
||||
peer.Touch()
|
||||
if !peer.IsConnected() {
|
||||
sendKey, recvKey := vl1.DeriveKeysFromPSK(a.cfg.PSK, a.identity.PublicKey, pubKey[:])
|
||||
peer.SetCipher(vl1.NewNoiseCipher(sendKey, recvKey))
|
||||
a.log.Info("peer connected", "peer", peer.Address)
|
||||
}
|
||||
return
|
||||
}
|
||||
peer = a.peers.AddPeer(addr, pubKey, from)
|
||||
sendKey, recvKey := vl1.DeriveKeysFromPSK(a.cfg.PSK, a.identity.PublicKey, pubKey[:])
|
||||
peer.SetCipher(vl1.NewNoiseCipher(sendKey, recvKey))
|
||||
a.log.Info("new peer", "peer", peer.Address)
|
||||
a.sendHello(peer)
|
||||
}
|
||||
|
||||
func (a *Agent) handleDataPacket(pkt *vl1.Packet, from *net.UDPAddr) {
|
||||
peer := a.peers.GetPeerByEndpoint(from)
|
||||
if peer == nil {
|
||||
return
|
||||
}
|
||||
peer.Touch()
|
||||
|
||||
plaintext, err := peer.Decrypt(pkt.Payload)
|
||||
if err != nil {
|
||||
a.log.Debug("decrypt failed", "peer", peer.Address, "err", err)
|
||||
return
|
||||
}
|
||||
|
||||
if a.network == nil {
|
||||
return
|
||||
}
|
||||
|
||||
frameToInject, err := a.network.Switch.HandleRemoteFrame(peer.Address, plaintext)
|
||||
if err != nil {
|
||||
a.log.Debug("switch handle remote frame", "err", err)
|
||||
return
|
||||
}
|
||||
if frameToInject != nil {
|
||||
if a.tapDev != nil {
|
||||
if _, err := a.tapDev.Write(frameToInject); err != nil {
|
||||
a.log.Debug("TAP write", "err", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) tapReadLoop() {
|
||||
defer a.wg.Done()
|
||||
buf := make([]byte, a.cfg.TAPMTU+64)
|
||||
for {
|
||||
select {
|
||||
case <-a.ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
n, err := a.tapDev.Read(buf)
|
||||
if err != nil {
|
||||
if a.ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
a.log.Error("TAP read error", "err", err)
|
||||
time.Sleep(time.Millisecond)
|
||||
continue
|
||||
}
|
||||
frame := make([]byte, n)
|
||||
copy(frame, buf[:n])
|
||||
if a.network != nil {
|
||||
if err := a.network.Switch.HandleLocalFrame(frame); err != nil {
|
||||
a.log.Debug("switch handle local", "err", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) sendHello(peer *vl1.Peer) {
|
||||
pkt := vl1.NewHandshakePacket(a.identity.PublicKey[:])
|
||||
if peer.Endpoint != nil {
|
||||
a.transport.SendPacket(&pkt, peer.Endpoint)
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Agent) maintenanceLoop() {
|
||||
defer a.wg.Done()
|
||||
ticker := time.NewTicker(10 * time.Second)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-a.ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
for _, peer := range a.peers.ConnectedPeers() {
|
||||
if peer.NeedsKeepalive() {
|
||||
pkt := vl1.NewKeepalivePacket()
|
||||
if peer.Endpoint != nil {
|
||||
a.transport.SendPacket(&pkt, peer.Endpoint)
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, peer := range a.peers.AllPeers() {
|
||||
if !peer.IsConnected() {
|
||||
a.sendHello(peer)
|
||||
}
|
||||
}
|
||||
a.peers.CleanDead()
|
||||
if a.network != nil {
|
||||
a.network.Switch.CleanExpired()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user