Files
zeromesh/sdk/agent.go

530 lines
11 KiB
Go

package sdk
import (
"context"
"crypto/cipher"
"crypto/sha256"
"encoding/binary"
"errors"
"fmt"
"log/slog"
"net"
"sync"
"time"
"golang.org/x/crypto/chacha20poly1305"
)
const (
vl1Version = 1
vl1HeaderSize = 8
pktHandshake = byte(1)
pktData = byte(2)
pktKeepalive = byte(3)
keepaliveInterval = 25 * time.Second
peerTimeout = 90 * time.Second
maintenanceTick = 10 * time.Second
macTableCap = 4096
macTableExpiry = 5 * time.Minute
)
type Agent struct {
cfg Config
client *Client
identity *Identity
udpConn *net.UDPConn
peers *peerManager
tapDev *tapInterface
log *slog.Logger
ctx context.Context
cancel context.CancelFunc
wg sync.WaitGroup
localMAC [6]byte
macTable *macTable
}
func NewAgent(client *Client) *Agent {
ctx, cancel := context.WithCancel(context.Background())
a := &Agent{
cfg: client.cfg,
client: client,
identity: client.identity,
log: slog.Default().With("component", "sdk-agent"),
ctx: ctx,
cancel: cancel,
}
if a.identity == nil {
a.identity = GenerateIdentity()
}
a.peers = newPeerManager()
a.macTable = newMACTable()
a.localMAC = generateMAC(a.cfg.NetworkID, a.identity.Address)
return a
}
func (a *Agent) Start() error {
port := a.cfg.ListenPort
if port == 0 {
port = 0 // OS auto-assign
}
udpAddr := &net.UDPAddr{Port: port}
conn, err := net.ListenUDP("udp", udpAddr)
if err != nil {
return fmt.Errorf("bind udp: %w", err)
}
a.udpConn = conn
a.log.Info("udp bound", "port", conn.LocalAddr().(*net.UDPAddr).Port)
tap, err := openTap(a.cfg.TapName, a.cfg.TapMTU)
if err != nil {
a.udpConn.Close()
return fmt.Errorf("open tap: %w", err)
}
a.tapDev = tap
a.log.Info("tap device opened", "name", tap.Name, "mtu", tap.MTU)
a.wg.Add(3)
go a.udpReadLoop()
go a.tapReadLoop()
go a.maintenanceLoop()
return nil
}
func (a *Agent) Stop() {
a.cancel()
if a.tapDev != nil {
a.tapDev.Close()
}
if a.udpConn != nil {
a.udpConn.Close()
}
a.wg.Wait()
a.log.Info("agent stopped")
}
func (a *Agent) LocalPort() int {
if a.udpConn == nil {
return 0
}
return a.udpConn.LocalAddr().(*net.UDPAddr).Port
}
func (a *Agent) udpReadLoop() {
defer a.wg.Done()
buf := make([]byte, 65535)
for {
select {
case <-a.ctx.Done():
return
default:
}
a.udpConn.SetReadDeadline(time.Now().Add(1 * time.Second))
n, addr, err := a.udpConn.ReadFromUDP(buf)
if err != nil {
if errors.Is(err, net.ErrClosed) {
return
}
continue
}
data := make([]byte, n)
copy(data, buf[:n])
a.handlePacket(data, addr)
}
}
func (a *Agent) handlePacket(data []byte, addr *net.UDPAddr) {
if len(data) < vl1HeaderSize {
return
}
pktType := data[1]
netID := binary.BigEndian.Uint32(data[2:6])
_ = netID
payloadLen := int(binary.BigEndian.Uint16(data[6:8]))
if vl1HeaderSize+payloadLen > len(data) {
return
}
payload := data[vl1HeaderSize : vl1HeaderSize+payloadLen]
switch pktType {
case pktHandshake:
a.handleHandshake(payload, addr)
case pktData:
a.handleData(payload, addr)
case pktKeepalive:
a.peers.touchByEndpoint(addr)
}
}
func (a *Agent) handleHandshake(payload []byte, addr *net.UDPAddr) {
if len(payload) < 32 {
return
}
var pubKey [32]byte
copy(pubKey[:], payload[:32])
peerAddr := AddressFromPublicKey(pubKey[:])
peer := a.peers.upsert(peerAddr, pubKey, addr)
peer.Touch()
a.log.Info("handshake from peer", "addr", peerAddr, "endpoint", addr)
sendKey, recvKey := deriveKeys([]byte(a.cfg.PSK), a.identity.PublicKey, pubKey[:])
c, err := chacha20poly1305.New(sendKey[:])
if err != nil {
return
}
_ = c
peer.sendCipher, err = chacha20poly1305.New(sendKey[:])
if err != nil {
return
}
peer.recvCipher, err = chacha20poly1305.New(recvKey[:])
if err != nil {
return
}
hello := make([]byte, 32)
copy(hello, a.identity.PublicKey)
a.sendPacket(pktHandshake, 0, hello, addr)
}
func (a *Agent) handleData(payload []byte, addr *net.UDPAddr) {
peer := a.peers.getByEndpoint(addr)
if peer == nil || peer.recvCipher == nil {
return
}
peer.Touch()
nonce := make([]byte, 12)
if len(payload) < 8 {
return
}
binary.BigEndian.PutUint64(nonce[4:], binary.BigEndian.Uint64(payload[:8]))
decrypted, err := peer.recvCipher.Open(nil, nonce, payload[8:], nil)
if err != nil {
return
}
frame := decrypted
if len(frame) < 14 {
return
}
a.macTable.learn(frame[6:12], peer.addr)
dstMAC := frame[:6]
srcMAC := frame[6:12]
_ = srcMAC
if isLocalMAC(dstMAC, a.localMAC[:]) {
if a.tapDev != nil {
a.tapDev.Write(frame)
}
return
}
if isBroadcastMAC(dstMAC) {
if a.tapDev != nil {
a.tapDev.Write(frame)
}
a.broadcastToPeers(frame, peer.addr)
return
}
dstAddr := a.macTable.lookup(string(dstMAC))
if dstAddr != nil {
if p := a.peers.getByAddr(*dstAddr); p != nil {
a.sendToPeer(frame, p)
} else {
a.broadcastToPeers(frame, peer.addr)
}
} else {
a.broadcastToPeers(frame, peer.addr)
}
}
func (a *Agent) tapReadLoop() {
defer a.wg.Done()
buf := make([]byte, a.cfg.TapMTU+14+64)
for {
select {
case <-a.ctx.Done():
return
default:
}
n, err := a.tapDev.Read(buf)
if err != nil {
continue
}
frame := make([]byte, n)
copy(frame, buf[:n])
if len(frame) < 14 {
continue
}
a.macTable.learn(frame[6:12], zeroAddr)
dstMAC := frame[:6]
if isBroadcastMAC(dstMAC) {
a.broadcastToPeers(frame, zeroAddr)
continue
}
dstAddr := a.macTable.lookup(string(dstMAC))
if dstAddr != nil {
if p := a.peers.getByAddr(*dstAddr); p != nil {
a.sendToPeer(frame, p)
} else {
a.broadcastToPeers(frame, zeroAddr)
}
} else {
a.broadcastToPeers(frame, zeroAddr)
}
}
}
func (a *Agent) maintenanceLoop() {
defer a.wg.Done()
ticker := time.NewTicker(maintenanceTick)
defer ticker.Stop()
for {
select {
case <-a.ctx.Done():
return
case <-ticker.C:
a.peers.cleanDead()
a.macTable.clean()
for _, p := range a.peers.all() {
if p.NeedsKeepalive() {
a.sendPacket(pktKeepalive, 0, nil, p.Endpoint)
p.LastSend = time.Now()
}
if !p.Connected && p.sendCipher != nil {
p.Connected = true
}
if !p.Connected {
a.sendHello(p)
}
}
}
}
}
func (a *Agent) sendPacket(pktType byte, netID uint32, payload []byte, addr *net.UDPAddr) {
if addr == nil {
return
}
total := vl1HeaderSize + len(payload)
buf := make([]byte, total)
buf[0] = vl1Version
buf[1] = pktType
binary.BigEndian.PutUint32(buf[2:6], netID)
binary.BigEndian.PutUint16(buf[6:8], uint16(len(payload)))
copy(buf[vl1HeaderSize:], payload)
a.udpConn.WriteTo(buf, addr)
}
func (a *Agent) sendToPeer(frame []byte, p *peer) {
if p == nil || p.sendCipher == nil {
return
}
p.LastSend = time.Now()
nonceBuf := make([]byte, 12)
binary.BigEndian.PutUint64(nonceBuf[4:], p.sendNonce)
p.sendNonce++
encrypted := p.sendCipher.Seal(nil, nonceBuf, frame, nil)
header := make([]byte, 8)
binary.BigEndian.PutUint64(header, binary.BigEndian.Uint64(nonceBuf[4:]))
payload := append(header, encrypted...)
a.sendPacket(pktData, a.cfg.NetworkID, payload, p.Endpoint)
}
func (a *Agent) broadcastToPeers(frame []byte, exclude Address) {
for _, p := range a.peers.all() {
if p.addr == exclude {
continue
}
if p.sendCipher == nil {
continue
}
a.sendToPeer(frame, p)
}
}
func (a *Agent) sendHello(p *peer) {
payload := make([]byte, 32)
copy(payload, a.identity.PublicKey)
a.sendPacket(pktHandshake, 0, payload, p.Endpoint)
}
var zeroAddr Address
type peer struct {
addr Address
pubKey [32]byte
Endpoint *net.UDPAddr
sendCipher cipher.AEAD
recvCipher cipher.AEAD
Connected bool
LastSeen time.Time
LastSend time.Time
sendNonce uint64
}
func (p *peer) Touch() { p.LastSeen = time.Now() }
func (p *peer) NeedsKeepalive() bool { return time.Since(p.LastSend) > keepaliveInterval }
func (p *peer) IsAlive() bool { return time.Since(p.LastSeen) < peerTimeout }
type peerManager struct {
mu sync.RWMutex
peers map[Address]*peer
epIdx map[string]*peer
}
func newPeerManager() *peerManager {
return &peerManager{
peers: make(map[Address]*peer),
epIdx: make(map[string]*peer),
}
}
func (pm *peerManager) upsert(addr Address, pubKey [32]byte, ep *net.UDPAddr) *peer {
pm.mu.Lock()
defer pm.mu.Unlock()
if p, ok := pm.peers[addr]; ok {
p.Endpoint = ep
pm.epIdx[ep.String()] = p
return p
}
p := &peer{addr: addr, pubKey: pubKey, Endpoint: ep}
pm.peers[addr] = p
pm.epIdx[ep.String()] = p
return p
}
func (pm *peerManager) getByEndpoint(ep *net.UDPAddr) *peer {
pm.mu.RLock()
defer pm.mu.RUnlock()
return pm.epIdx[ep.String()]
}
func (pm *peerManager) getByAddr(addr Address) *peer {
pm.mu.RLock()
defer pm.mu.RUnlock()
return pm.peers[addr]
}
func (pm *peerManager) touchByEndpoint(ep *net.UDPAddr) {
pm.mu.Lock()
defer pm.mu.Unlock()
if p, ok := pm.epIdx[ep.String()]; ok {
p.Touch()
}
}
func (pm *peerManager) all() []*peer {
pm.mu.RLock()
defer pm.mu.RUnlock()
out := make([]*peer, 0, len(pm.peers))
for _, p := range pm.peers {
out = append(out, p)
}
return out
}
func (pm *peerManager) cleanDead() {
pm.mu.Lock()
defer pm.mu.Unlock()
for addr, p := range pm.peers {
if !p.IsAlive() {
delete(pm.epIdx, p.Endpoint.String())
delete(pm.peers, addr)
}
}
}
type macEntry struct {
peerAddr Address
added time.Time
}
type macTable struct {
mu sync.RWMutex
tbl map[string]macEntry
}
func newMACTable() *macTable {
return &macTable{tbl: make(map[string]macEntry)}
}
func (mt *macTable) learn(mac []byte, peerAddr Address) {
mt.mu.Lock()
defer mt.mu.Unlock()
if len(mt.tbl) >= macTableCap {
for k, v := range mt.tbl {
if time.Since(v.added) > macTableExpiry {
delete(mt.tbl, k)
}
}
}
mt.tbl[string(mac)] = macEntry{peerAddr: peerAddr, added: time.Now()}
}
func (mt *macTable) lookup(mac string) *Address {
mt.mu.RLock()
defer mt.mu.RUnlock()
e, ok := mt.tbl[mac]
if !ok {
return nil
}
zero := Address{}
if e.peerAddr == zero {
return nil
}
return &e.peerAddr
}
func (mt *macTable) clean() {
mt.mu.Lock()
defer mt.mu.Unlock()
for k, v := range mt.tbl {
if time.Since(v.added) > macTableExpiry {
delete(mt.tbl, k)
}
}
}
func generateMAC(networkID uint32, addr Address) [6]byte {
var mac [6]byte
mac[0] = 0x02
binary.BigEndian.PutUint32(mac[1:5], networkID)
mac[5] = addr[4]
return mac
}
func isLocalMAC(dst, local []byte) bool {
return dst[0] == local[0] && dst[1] == local[1] && dst[2] == local[2] &&
dst[3] == local[3] && dst[4] == local[4] && dst[5] == local[5]
}
func isBroadcastMAC(dst []byte) bool {
for _, b := range dst {
if b != 0xff {
return false
}
}
return true
}
func deriveKeys(psk, localPub, remotePub []byte) ([32]byte, [32]byte) {
h := sha256.New()
h.Write(psk)
h.Write(localPub)
h.Write(remotePub)
var sendKey [32]byte
copy(sendKey[:], h.Sum(nil))
h2 := sha256.New()
h2.Write(sendKey[:])
h2.Write([]byte("reverse"))
var recvKey [32]byte
copy(recvKey[:], h2.Sum(nil))
return sendKey, recvKey
}