325 lines
7.2 KiB
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()
|
|
}
|
|
}
|
|
}
|
|
}
|