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