Files
zeromesh/internal/agent/agent.go

325 lines
7.2 KiB
Go

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()
}
}
}
}