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()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
142
internal/config/config.go
Normal file
142
internal/config/config.go
Normal file
@@ -0,0 +1,142 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
Server ServerConfig `yaml:"server"`
|
||||
Database DatabaseConfig `yaml:"database"`
|
||||
JWT JWTConfig `yaml:"jwt"`
|
||||
Controller ControllerConfig `yaml:"controller"`
|
||||
Agent AgentConfig `yaml:"agent"`
|
||||
Relay RelayConfig `yaml:"relay"`
|
||||
CORS CORSConfig `yaml:"cors"`
|
||||
RateLimit RateLimitConfig `yaml:"rate_limit"`
|
||||
Logging LoggingConfig `yaml:"logging"`
|
||||
Env EnvConfig `yaml:"env"`
|
||||
}
|
||||
|
||||
type ServerConfig struct {
|
||||
Port int `yaml:"port"`
|
||||
}
|
||||
|
||||
func (s *ServerConfig) Addr() string {
|
||||
return fmt.Sprintf(":%d", s.Port)
|
||||
}
|
||||
|
||||
type DatabaseConfig struct {
|
||||
Driver string `yaml:"driver"`
|
||||
AutoMigrate bool `yaml:"auto_migrate"`
|
||||
InitData bool `yaml:"init_data"`
|
||||
SQLite SQLiteConfig `yaml:"sqlite"`
|
||||
MySQL MySQLConfig `yaml:"mysql"`
|
||||
}
|
||||
|
||||
type SQLiteConfig struct {
|
||||
Path string `yaml:"path"`
|
||||
}
|
||||
|
||||
type MySQLConfig struct {
|
||||
Host string `yaml:"host"`
|
||||
Port int `yaml:"port"`
|
||||
Username string `yaml:"username"`
|
||||
Password string `yaml:"password"`
|
||||
Database string `yaml:"database"`
|
||||
Charset string `yaml:"charset"`
|
||||
MaxIdleConns int `yaml:"max_idle_conns"`
|
||||
MaxOpenConns int `yaml:"max_open_conns"`
|
||||
}
|
||||
|
||||
type JWTConfig struct {
|
||||
Secret string `yaml:"secret"`
|
||||
ExpireHours int `yaml:"expire_hours"`
|
||||
}
|
||||
|
||||
type ControllerConfig struct {
|
||||
ListenPort int `yaml:"listen_port"`
|
||||
PublicEndpoint string `yaml:"public_endpoint"`
|
||||
PlanetID string `yaml:"planet_id"`
|
||||
}
|
||||
|
||||
type AgentConfig struct {
|
||||
ControllerURL string `yaml:"controller_url"`
|
||||
ListenPort int `yaml:"listen_port"`
|
||||
TAPName string `yaml:"tap_name"`
|
||||
TAPMTU int `yaml:"tap_mtu"`
|
||||
NetworkID uint32 `yaml:"network_id"`
|
||||
PSK string `yaml:"psk"`
|
||||
StaticPeers []string `yaml:"static_peers"`
|
||||
}
|
||||
|
||||
type RelayConfig struct {
|
||||
ListenPort int `yaml:"listen_port"`
|
||||
Realm string `yaml:"realm"`
|
||||
}
|
||||
|
||||
type CORSConfig struct {
|
||||
AllowedOrigins []string `yaml:"allowed_origins"`
|
||||
}
|
||||
|
||||
type RateLimitConfig struct {
|
||||
Enabled bool `yaml:"enabled"`
|
||||
RequestsPerSecond int `yaml:"requests_per_second"`
|
||||
BurstSize int `yaml:"burst_size"`
|
||||
}
|
||||
|
||||
type LoggingConfig struct {
|
||||
Level string `yaml:"level"`
|
||||
Format string `yaml:"format"`
|
||||
}
|
||||
|
||||
type EnvConfig struct {
|
||||
Name string `yaml:"name"`
|
||||
}
|
||||
|
||||
func Load(path string) (*Config, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read config: %w", err)
|
||||
}
|
||||
var cfg Config
|
||||
if err := yaml.Unmarshal(data, &cfg); err != nil {
|
||||
return nil, fmt.Errorf("parse config: %w", err)
|
||||
}
|
||||
setDefaults(&cfg)
|
||||
return &cfg, nil
|
||||
}
|
||||
|
||||
func setDefaults(cfg *Config) {
|
||||
if cfg.Server.Port == 0 {
|
||||
cfg.Server.Port = 10001
|
||||
}
|
||||
if cfg.Controller.ListenPort == 0 {
|
||||
cfg.Controller.ListenPort = 9993
|
||||
}
|
||||
if cfg.Agent.TAPName == "" {
|
||||
cfg.Agent.TAPName = "zeromesh0"
|
||||
}
|
||||
if cfg.Agent.TAPMTU == 0 {
|
||||
cfg.Agent.TAPMTU = 2800
|
||||
}
|
||||
if cfg.JWT.ExpireHours == 0 {
|
||||
cfg.JWT.ExpireHours = 720
|
||||
}
|
||||
if cfg.Logging.Level == "" {
|
||||
cfg.Logging.Level = "info"
|
||||
}
|
||||
if cfg.Env.Name == "" {
|
||||
cfg.Env.Name = "dev"
|
||||
}
|
||||
}
|
||||
|
||||
func (d *DatabaseConfig) BuildDSN() string {
|
||||
if d.Driver == "mysql" {
|
||||
return fmt.Sprintf("%s:%s@tcp(%s:%d)/%s?charset=%s&parseTime=True&loc=Local",
|
||||
d.MySQL.Username, d.MySQL.Password, d.MySQL.Host, d.MySQL.Port, d.MySQL.Database, d.MySQL.Charset)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
192
internal/controller/controller.go
Normal file
192
internal/controller/controller.go
Normal file
@@ -0,0 +1,192 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"zeromesh/internal/config"
|
||||
"zeromesh/internal/database"
|
||||
"zeromesh/internal/identity"
|
||||
"zeromesh/internal/model"
|
||||
"zeromesh/internal/repository"
|
||||
"zeromesh/internal/vl1"
|
||||
)
|
||||
|
||||
type Controller struct {
|
||||
cfg *config.Config
|
||||
db *database.DB
|
||||
identity *identity.Identity
|
||||
peers *vl1.PeerManager
|
||||
transport *vl1.Transport
|
||||
netRepo *repository.NetworkRepo
|
||||
nodeRepo *repository.NodeRepo
|
||||
wsClients map[string]net.Conn
|
||||
mu sync.RWMutex
|
||||
log *slog.Logger
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
func New(cfg *config.Config, db *database.DB, log *slog.Logger) (*Controller, error) {
|
||||
id, err := identity.LoadOrGenerate("./data/controller.identity")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("load controller identity: %w", err)
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
return &Controller{
|
||||
cfg: cfg,
|
||||
db: db,
|
||||
identity: id,
|
||||
peers: vl1.NewPeerManager(log),
|
||||
netRepo: repository.NewNetworkRepo(db),
|
||||
nodeRepo: repository.NewNodeRepo(db),
|
||||
wsClients: make(map[string]net.Conn),
|
||||
log: log.With("component", "controller"),
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (c *Controller) Start() error {
|
||||
transport, err := vl1.NewTransport(c.cfg.Controller.ListenPort, c.log)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
c.transport = transport
|
||||
c.log.Info("controller started",
|
||||
"address", c.identity.Address,
|
||||
"port", c.transport.Port(),
|
||||
)
|
||||
|
||||
c.wg.Add(1)
|
||||
go c.udpReadLoop()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Controller) Stop() {
|
||||
c.cancel()
|
||||
if c.transport != nil {
|
||||
c.transport.Close()
|
||||
}
|
||||
c.wg.Wait()
|
||||
}
|
||||
|
||||
func (c *Controller) udpReadLoop() {
|
||||
defer c.wg.Done()
|
||||
buf := make([]byte, vl1.MaxPacketSize)
|
||||
for {
|
||||
select {
|
||||
case <-c.ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
n, remoteAddr, err := c.transport.ReadFrom(buf)
|
||||
if err != nil {
|
||||
if c.ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
c.log.Error("UDP read error", "err", err)
|
||||
time.Sleep(time.Millisecond)
|
||||
continue
|
||||
}
|
||||
c.handlePacket(buf[:n], remoteAddr)
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Controller) handlePacket(data []byte, from *net.UDPAddr) {
|
||||
pkt, err := vl1.DecodePacket(data)
|
||||
if err != nil {
|
||||
c.log.Debug("decode packet", "err", err)
|
||||
return
|
||||
}
|
||||
switch pkt.Header.Type {
|
||||
case vl1.PacketTypeHandshake:
|
||||
c.handleHandshake(pkt.Payload, from)
|
||||
case vl1.PacketTypeKeepalive:
|
||||
if peer := c.peers.GetPeerByEndpoint(from); peer != nil {
|
||||
peer.Touch()
|
||||
nodeID := peer.Address.String()
|
||||
_ = c.nodeRepo.SetOnline(nodeID, true)
|
||||
}
|
||||
case vl1.PacketTypeData:
|
||||
// Controller does not forward data; data is P2P between agents
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Controller) handleHandshake(payload []byte, from *net.UDPAddr) {
|
||||
if len(payload) < 32 {
|
||||
return
|
||||
}
|
||||
var pubKey [32]byte
|
||||
copy(pubKey[:], payload[:32])
|
||||
addr := identity.AddressFromPublicKey(pubKey[:])
|
||||
nodeID := addr.String()
|
||||
|
||||
peer := c.peers.GetPeer(addr)
|
||||
if peer == nil {
|
||||
peer = c.peers.AddPeer(addr, pubKey, from)
|
||||
c.log.Info("new peer connected", "node_id", nodeID, "addr", from)
|
||||
} else {
|
||||
c.peers.UpdatePeerEndpoint(addr, from)
|
||||
peer.Touch()
|
||||
}
|
||||
_ = c.HandleNodeHello(nodeID, fmt.Sprintf("%x", pubKey[:]), from)
|
||||
}
|
||||
|
||||
func (c *Controller) Identity() *identity.Identity {
|
||||
return c.identity
|
||||
}
|
||||
|
||||
func (c *Controller) HandleNodeHello(nodeID string, publicKey string, addr *net.UDPAddr) error {
|
||||
node, err := c.nodeRepo.FindByNodeID(nodeID)
|
||||
if err != nil {
|
||||
// New node: register it
|
||||
node = &model.Node{
|
||||
NodeID: nodeID,
|
||||
PublicKey: publicKey,
|
||||
IPAddress: addr.IP.String(),
|
||||
Port: addr.Port,
|
||||
Online: true,
|
||||
Version: "1.0.0",
|
||||
}
|
||||
if err := c.nodeRepo.Create(node); err != nil {
|
||||
return err
|
||||
}
|
||||
c.log.Info("new node registered", "node_id", nodeID, "addr", addr)
|
||||
} else {
|
||||
c.nodeRepo.SetOnline(nodeID, true)
|
||||
c.log.Debug("node hello", "node_id", nodeID, "addr", addr)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Controller) GetNodeList() ([]model.Node, error) {
|
||||
return c.nodeRepo.ListOnline()
|
||||
}
|
||||
|
||||
func (c *Controller) GetNetworkConfig(networkID uint32) (*model.Network, error) {
|
||||
return c.netRepo.FindByNetworkID(networkID)
|
||||
}
|
||||
|
||||
func (c *Controller) AuthorizeNode(networkID uint32, nodeID string) error {
|
||||
_, err := c.netRepo.FindByNetworkID(networkID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = c.nodeRepo.FindByNodeID(nodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return c.netRepo.AddMember(&model.NetworkMember{
|
||||
NetworkID: networkID,
|
||||
NodeID: nodeID,
|
||||
Authorized: true,
|
||||
})
|
||||
}
|
||||
84
internal/database/database.go
Normal file
84
internal/database/database.go
Normal file
@@ -0,0 +1,84 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
|
||||
"gorm.io/driver/mysql"
|
||||
"gorm.io/driver/sqlite"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"zeromesh/internal/config"
|
||||
"zeromesh/internal/model"
|
||||
)
|
||||
|
||||
type DB struct {
|
||||
*gorm.DB
|
||||
}
|
||||
|
||||
func Init(cfg *config.DatabaseConfig) (*DB, error) {
|
||||
var db *gorm.DB
|
||||
var err error
|
||||
|
||||
switch cfg.Driver {
|
||||
case "mysql":
|
||||
dsn := cfg.BuildDSN()
|
||||
db, err = gorm.Open(mysql.Open(dsn), &gorm.Config{})
|
||||
case "sqlite":
|
||||
fallthrough
|
||||
default:
|
||||
db, err = gorm.Open(sqlite.Open(cfg.SQLite.Path), &gorm.Config{})
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open database: %w", err)
|
||||
}
|
||||
|
||||
sqlDB, err := db.DB()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get sql db: %w", err)
|
||||
}
|
||||
|
||||
if cfg.Driver == "mysql" {
|
||||
sqlDB.SetMaxIdleConns(cfg.MySQL.MaxIdleConns)
|
||||
sqlDB.SetMaxOpenConns(cfg.MySQL.MaxOpenConns)
|
||||
}
|
||||
|
||||
if cfg.AutoMigrate {
|
||||
if err := db.AutoMigrate(
|
||||
&model.User{},
|
||||
&model.Network{},
|
||||
&model.NetworkMember{},
|
||||
&model.Node{},
|
||||
&model.Config{},
|
||||
&model.ApiLog{},
|
||||
&model.ErrorLog{},
|
||||
); err != nil {
|
||||
return nil, fmt.Errorf("auto migrate: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if cfg.InitData {
|
||||
seedData(db)
|
||||
}
|
||||
|
||||
slog.Info("database initialized", "driver", cfg.Driver)
|
||||
return &DB{db}, nil
|
||||
}
|
||||
|
||||
func seedData(db *gorm.DB) {
|
||||
var count int64
|
||||
db.Model(&model.Config{}).Count(&count)
|
||||
if count > 0 {
|
||||
return
|
||||
}
|
||||
|
||||
defaults := []model.Config{
|
||||
{Key: "allow_register", Value: "true"},
|
||||
{Key: "default_network_pool", Value: "10.147.0.0/16"},
|
||||
}
|
||||
for _, c := range defaults {
|
||||
db.Create(&c)
|
||||
}
|
||||
slog.Info("seed data inserted")
|
||||
}
|
||||
62
internal/handler/admin.go
Normal file
62
internal/handler/admin.go
Normal file
@@ -0,0 +1,62 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"zeromesh/internal/database"
|
||||
"zeromesh/internal/model"
|
||||
"zeromesh/internal/service"
|
||||
)
|
||||
|
||||
type AdminHandler struct {
|
||||
authService *service.AuthService
|
||||
db *database.DB
|
||||
nodeService *service.NodeService
|
||||
netService *service.NetworkService
|
||||
}
|
||||
|
||||
func NewAdminHandler(authService *service.AuthService, db *database.DB, nodeService *service.NodeService, netService *service.NetworkService) *AdminHandler {
|
||||
return &AdminHandler{
|
||||
authService: authService,
|
||||
db: db,
|
||||
nodeService: nodeService,
|
||||
netService: netService,
|
||||
}
|
||||
}
|
||||
|
||||
func (h *AdminHandler) Dashboard(c *gin.Context) {
|
||||
uid := userID(c)
|
||||
|
||||
var authorizedMembers int64
|
||||
if isAdmin(c) {
|
||||
h.db.Model(&model.NetworkMember{}).Where("authorized = ?", true).Count(&authorizedMembers)
|
||||
} else {
|
||||
networks, _ := h.netService.ListNetworks(uid)
|
||||
var ids []uint32
|
||||
for _, n := range networks {
|
||||
ids = append(ids, n.NetworkID)
|
||||
}
|
||||
if len(ids) > 0 {
|
||||
h.db.Model(&model.NetworkMember{}).Where("authorized = ? AND network_id IN ?", true, ids).Count(&authorizedMembers)
|
||||
}
|
||||
}
|
||||
|
||||
nodes, _ := h.nodeService.ListNode(uid)
|
||||
var online int64
|
||||
for _, n := range nodes {
|
||||
if n.Online {
|
||||
online++
|
||||
}
|
||||
}
|
||||
|
||||
networks, _ := h.netService.ListNetworks(uid)
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"authorized_members": authorizedMembers,
|
||||
"nodes_total": len(nodes),
|
||||
"nodes_online": online,
|
||||
"networks_count": len(networks),
|
||||
})
|
||||
}
|
||||
87
internal/handler/auth.go
Normal file
87
internal/handler/auth.go
Normal file
@@ -0,0 +1,87 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"zeromesh/internal/service"
|
||||
)
|
||||
|
||||
type AuthHandler struct {
|
||||
authService *service.AuthService
|
||||
}
|
||||
|
||||
func NewAuthHandler(authService *service.AuthService) *AuthHandler {
|
||||
return &AuthHandler{authService: authService}
|
||||
}
|
||||
|
||||
type loginReq struct {
|
||||
Username string `json:"username" binding:"required"`
|
||||
Password string `json:"password" binding:"required"`
|
||||
}
|
||||
|
||||
type registerReq struct {
|
||||
Username string `json:"username" binding:"required"`
|
||||
Password string `json:"password" binding:"required"`
|
||||
}
|
||||
|
||||
func (h *AuthHandler) Login(c *gin.Context) {
|
||||
var req loginReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
token, user, err := h.authService.Login(req.Username, req.Password)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"token": token,
|
||||
"user": user,
|
||||
})
|
||||
}
|
||||
|
||||
func (h *AuthHandler) Register(c *gin.Context) {
|
||||
var req registerReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
user, token, err := h.authService.Register(req.Username, req.Password)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"token": token,
|
||||
"user": user,
|
||||
})
|
||||
}
|
||||
|
||||
func (h *AuthHandler) InitAdmin(c *gin.Context) {
|
||||
var req registerReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
user, token, err := h.authService.InitAdmin(req.Username, req.Password)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"token": token,
|
||||
"user": user,
|
||||
})
|
||||
}
|
||||
|
||||
func (h *AuthHandler) CheckAdmin(c *gin.Context) {
|
||||
exists, err := h.authService.CheckAdmin()
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"admin_exists": exists})
|
||||
}
|
||||
23
internal/handler/context.go
Normal file
23
internal/handler/context.go
Normal file
@@ -0,0 +1,23 @@
|
||||
package handler
|
||||
|
||||
import "github.com/gin-gonic/gin"
|
||||
|
||||
func userID(c *gin.Context) uint {
|
||||
id, exists := c.Get("user_id")
|
||||
if !exists {
|
||||
return 0
|
||||
}
|
||||
uid, ok := id.(uint)
|
||||
if !ok {
|
||||
return 0
|
||||
}
|
||||
return uid
|
||||
}
|
||||
|
||||
func isAdmin(c *gin.Context) bool {
|
||||
role, exists := c.Get("role")
|
||||
if !exists {
|
||||
return false
|
||||
}
|
||||
return role.(string) == "admin"
|
||||
}
|
||||
75
internal/handler/log.go
Normal file
75
internal/handler/log.go
Normal file
@@ -0,0 +1,75 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"zeromesh/internal/service"
|
||||
)
|
||||
|
||||
type LogHandler struct {
|
||||
logService *service.LogService
|
||||
}
|
||||
|
||||
func NewLogHandler(logService *service.LogService) *LogHandler {
|
||||
return &LogHandler{logService: logService}
|
||||
}
|
||||
|
||||
func (h *LogHandler) ListApiLogs(c *gin.Context) {
|
||||
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
||||
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20"))
|
||||
ip := c.Query("ip")
|
||||
path := c.Query("path")
|
||||
method := c.Query("method")
|
||||
statusCode, _ := strconv.Atoi(c.Query("status_code"))
|
||||
|
||||
logs, total, err := h.logService.ListApiLogs(page, pageSize, ip, path, method, statusCode)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"data": logs, "total": total, "page": page, "page_size": pageSize})
|
||||
}
|
||||
|
||||
func (h *LogHandler) ListErrorLogs(c *gin.Context) {
|
||||
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
||||
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20"))
|
||||
ip := c.Query("ip")
|
||||
path := c.Query("path")
|
||||
method := c.Query("method")
|
||||
statusCode, _ := strconv.Atoi(c.Query("status_code"))
|
||||
|
||||
logs, total, err := h.logService.ListErrorLogs(page, pageSize, ip, path, method, statusCode)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"data": logs, "total": total, "page": page, "page_size": pageSize})
|
||||
}
|
||||
|
||||
func (h *LogHandler) ListApiLogGrouped(c *gin.Context) {
|
||||
groups, err := h.logService.ListApiLogGrouped()
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"data": groups})
|
||||
}
|
||||
|
||||
func (h *LogHandler) ClearApiLogs(c *gin.Context) {
|
||||
if err := h.logService.ClearApiLogs(); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"message": "ok"})
|
||||
}
|
||||
|
||||
func (h *LogHandler) ClearErrorLogs(c *gin.Context) {
|
||||
if err := h.logService.ClearErrorLogs(); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"message": "ok"})
|
||||
}
|
||||
143
internal/handler/network.go
Normal file
143
internal/handler/network.go
Normal file
@@ -0,0 +1,143 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"zeromesh/internal/service"
|
||||
)
|
||||
|
||||
type NetworkHandler struct {
|
||||
networkService *service.NetworkService
|
||||
}
|
||||
|
||||
func NewNetworkHandler(networkService *service.NetworkService) *NetworkHandler {
|
||||
return &NetworkHandler{networkService: networkService}
|
||||
}
|
||||
|
||||
type createNetworkReq struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
IPRange string `json:"ip_range"`
|
||||
}
|
||||
|
||||
func (h *NetworkHandler) Create(c *gin.Context) {
|
||||
var req createNetworkReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
uid := userID(c)
|
||||
network, err := h.networkService.CreateNetwork(uid, req.Name, req.IPRange)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"network": network})
|
||||
}
|
||||
|
||||
func (h *NetworkHandler) List(c *gin.Context) {
|
||||
uid := userID(c)
|
||||
if !isAdmin(c) {
|
||||
networks, err := h.networkService.ListNetworks(uid)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"networks": networks})
|
||||
return
|
||||
}
|
||||
networks, err := h.networkService.ListNetworks(0)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"networks": networks})
|
||||
}
|
||||
|
||||
func (h *NetworkHandler) Get(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 32)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid network id"})
|
||||
return
|
||||
}
|
||||
uid := userID(c)
|
||||
network, err := h.networkService.GetNetwork(uint32(id), uid)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "network not found"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"network": network})
|
||||
}
|
||||
|
||||
func (h *NetworkHandler) Delete(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 32)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid network id"})
|
||||
return
|
||||
}
|
||||
uid := userID(c)
|
||||
if err := h.networkService.DeleteNetwork(uint32(id), uid); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"message": "deleted"})
|
||||
}
|
||||
|
||||
func (h *NetworkHandler) Members(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 32)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid network id"})
|
||||
return
|
||||
}
|
||||
uid := userID(c)
|
||||
members, err := h.networkService.ListMembers(uint32(id), uid)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"members": members})
|
||||
}
|
||||
|
||||
type authorizeReq struct {
|
||||
NodeID string `json:"node_id" binding:"required"`
|
||||
}
|
||||
|
||||
func (h *NetworkHandler) Authorize(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 32)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid network id"})
|
||||
return
|
||||
}
|
||||
var req authorizeReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
uid := userID(c)
|
||||
if err := h.networkService.AuthorizeMember(uint32(id), req.NodeID, uid); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"message": "member authorized"})
|
||||
}
|
||||
|
||||
func (h *NetworkHandler) Deauthorize(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 32)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid network id"})
|
||||
return
|
||||
}
|
||||
var req authorizeReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
uid := userID(c)
|
||||
if err := h.networkService.DeauthorizeMember(uint32(id), req.NodeID, uid); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"message": "member deauthorized"})
|
||||
}
|
||||
81
internal/handler/node.go
Normal file
81
internal/handler/node.go
Normal file
@@ -0,0 +1,81 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"zeromesh/internal/service"
|
||||
)
|
||||
|
||||
type NodeHandler struct {
|
||||
nodeService *service.NodeService
|
||||
}
|
||||
|
||||
func NewNodeHandler(nodeService *service.NodeService) *NodeHandler {
|
||||
return &NodeHandler{nodeService: nodeService}
|
||||
}
|
||||
|
||||
type registerNodeReq struct {
|
||||
NodeID string `json:"node_id" binding:"required"`
|
||||
PublicKey string `json:"public_key" binding:"required"`
|
||||
Name string `json:"name"`
|
||||
IPAddress string `json:"ip_address"`
|
||||
Port int `json:"port"`
|
||||
Version string `json:"version"`
|
||||
}
|
||||
|
||||
func (h *NodeHandler) Register(c *gin.Context) {
|
||||
var req registerNodeReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
uid := userID(c)
|
||||
node, err := h.nodeService.Register(req.NodeID, req.PublicKey, req.Name, req.IPAddress, req.Port, req.Version, uid)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"node": node})
|
||||
}
|
||||
|
||||
func (h *NodeHandler) List(c *gin.Context) {
|
||||
uid := userID(c)
|
||||
if !isAdmin(c) {
|
||||
nodes, err := h.nodeService.ListNode(uid)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"nodes": nodes})
|
||||
return
|
||||
}
|
||||
nodes, err := h.nodeService.ListNode(0)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"nodes": nodes})
|
||||
}
|
||||
|
||||
func (h *NodeHandler) ListOnline(c *gin.Context) {
|
||||
uid := userID(c)
|
||||
nodes, err := h.nodeService.ListOnline(uid)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"nodes": nodes})
|
||||
}
|
||||
|
||||
func (h *NodeHandler) Heartbeat(c *gin.Context) {
|
||||
nodeID := c.Param("node_id")
|
||||
if err := h.nodeService.SetOnline(nodeID, true); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
time.Sleep(5 * time.Second)
|
||||
c.JSON(http.StatusOK, gin.H{"status": "ok"})
|
||||
}
|
||||
142
internal/identity/identity.go
Normal file
142
internal/identity/identity.go
Normal file
@@ -0,0 +1,142 @@
|
||||
package identity
|
||||
|
||||
import (
|
||||
"crypto/ed25519"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type Identity struct {
|
||||
PublicKey ed25519.PublicKey `json:"public_key"`
|
||||
PrivateKey ed25519.PrivateKey `json:"private_key"`
|
||||
Address Address `json:"address"`
|
||||
}
|
||||
|
||||
type Address [5]byte
|
||||
|
||||
func (a Address) String() string {
|
||||
return hex.EncodeToString(a[:])
|
||||
}
|
||||
|
||||
func (a Address) MarshalText() ([]byte, error) {
|
||||
return []byte(a.String()), nil
|
||||
}
|
||||
|
||||
func (a *Address) UnmarshalText(text []byte) error {
|
||||
decoded, err := hex.DecodeString(string(text))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(decoded) != 5 {
|
||||
return fmt.Errorf("address must be 5 bytes")
|
||||
}
|
||||
copy(a[:], decoded)
|
||||
return nil
|
||||
}
|
||||
|
||||
func AddressFromPublicKey(pubKey []byte) Address {
|
||||
var addr Address
|
||||
h := hashBytes(pubKey)
|
||||
copy(addr[:], h[:5])
|
||||
return addr
|
||||
}
|
||||
|
||||
func Generate() *Identity {
|
||||
pub, priv, err := ed25519.GenerateKey(rand.Reader)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return &Identity{
|
||||
PublicKey: pub,
|
||||
PrivateKey: priv,
|
||||
Address: AddressFromPublicKey(pub),
|
||||
}
|
||||
}
|
||||
|
||||
func (id *Identity) PublicKeyHex() string {
|
||||
return hex.EncodeToString(id.PublicKey)
|
||||
}
|
||||
|
||||
func (id *Identity) PrivateKeyHex() string {
|
||||
return hex.EncodeToString(id.PrivateKey)
|
||||
}
|
||||
|
||||
func (id *Identity) Sign(data []byte) []byte {
|
||||
return ed25519.Sign(id.PrivateKey, data)
|
||||
}
|
||||
|
||||
func (id *Identity) Verify(data, sig []byte) bool {
|
||||
return ed25519.Verify(id.PublicKey, data, sig)
|
||||
}
|
||||
|
||||
func LoadOrGenerate(path string) (*Identity, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err == nil {
|
||||
return Parse(strings.TrimSpace(string(data)))
|
||||
}
|
||||
if !os.IsNotExist(err) {
|
||||
return nil, err
|
||||
}
|
||||
id := Generate()
|
||||
encoded, err := id.Serialize()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := os.MkdirAll(dir(path), 0755); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := os.WriteFile(path, []byte(encoded), 0600); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
func (id *Identity) Serialize() (string, error) {
|
||||
data, err := json.Marshal(id)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(data), nil
|
||||
}
|
||||
|
||||
func Parse(data string) (*Identity, error) {
|
||||
var id Identity
|
||||
if err := json.Unmarshal([]byte(data), &id); err != nil {
|
||||
return nil, fmt.Errorf("parse identity: %w", err)
|
||||
}
|
||||
if len(id.PrivateKey) == 0 {
|
||||
return nil, fmt.Errorf("invalid identity: no private key")
|
||||
}
|
||||
id.Address = AddressFromPublicKey(id.PublicKey)
|
||||
return &id, nil
|
||||
}
|
||||
|
||||
func dir(path string) string {
|
||||
idx := strings.LastIndex(path, "/")
|
||||
if idx == -1 {
|
||||
idx = strings.LastIndex(path, "\\")
|
||||
}
|
||||
if idx == -1 {
|
||||
return "."
|
||||
}
|
||||
return path[:idx]
|
||||
}
|
||||
|
||||
func hashBytes(data []byte) []byte {
|
||||
h := make([]byte, 32)
|
||||
for i, b := range data {
|
||||
h[i%32] ^= b
|
||||
}
|
||||
// simple hash expansion
|
||||
for round := 0; round < 3; round++ {
|
||||
for i := 0; i < 32; i++ {
|
||||
h[i] = h[i] ^ h[(i+1)%32] ^ h[(i+7)%32]
|
||||
h[i] = (h[i] << 3) | (h[i] >> 5)
|
||||
}
|
||||
}
|
||||
return h
|
||||
}
|
||||
29
internal/middleware/cors.go
Normal file
29
internal/middleware/cors.go
Normal file
@@ -0,0 +1,29 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func CORS(allowedOrigins []string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
origin := c.GetHeader("Origin")
|
||||
for _, allowed := range allowedOrigins {
|
||||
if allowed == "*" || allowed == origin {
|
||||
c.Header("Access-Control-Allow-Origin", origin)
|
||||
break
|
||||
}
|
||||
}
|
||||
if origin == "" {
|
||||
c.Header("Access-Control-Allow-Origin", "*")
|
||||
}
|
||||
c.Header("Access-Control-Allow-Methods", "GET, POST, PUT, DELETE, OPTIONS")
|
||||
c.Header("Access-Control-Allow-Headers", "Content-Type, Authorization")
|
||||
c.Header("Access-Control-Allow-Credentials", "true")
|
||||
|
||||
if c.Request.Method == "OPTIONS" {
|
||||
c.AbortWithStatus(204)
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
58
internal/middleware/jwt.go
Normal file
58
internal/middleware/jwt.go
Normal file
@@ -0,0 +1,58 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
|
||||
type JWTClaims struct {
|
||||
UserID uint `json:"user_id"`
|
||||
Username string `json:"username"`
|
||||
Role string `json:"role"`
|
||||
jwt.RegisteredClaims
|
||||
}
|
||||
|
||||
func JWTAuth(secret string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
auth := c.GetHeader("Authorization")
|
||||
if auth == "" {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "missing authorization"})
|
||||
return
|
||||
}
|
||||
parts := strings.SplitN(auth, " ", 2)
|
||||
if len(parts) != 2 || parts[0] != "Bearer" {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "invalid authorization format"})
|
||||
return
|
||||
}
|
||||
claims := &JWTClaims{}
|
||||
token, err := jwt.ParseWithClaims(parts[1], claims, func(t *jwt.Token) (interface{}, error) {
|
||||
return []byte(secret), nil
|
||||
})
|
||||
if err != nil || !token.Valid {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "invalid token"})
|
||||
return
|
||||
}
|
||||
c.Set("user_id", claims.UserID)
|
||||
c.Set("username", claims.Username)
|
||||
c.Set("role", claims.Role)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
func GenerateToken(secret string, userID uint, username, role string, expireHours int) (string, error) {
|
||||
claims := JWTClaims{
|
||||
UserID: userID,
|
||||
Username: username,
|
||||
Role: role,
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Duration(expireHours) * time.Hour)),
|
||||
IssuedAt: jwt.NewNumericDate(time.Now()),
|
||||
},
|
||||
}
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
return token.SignedString([]byte(secret))
|
||||
}
|
||||
97
internal/middleware/logger.go
Normal file
97
internal/middleware/logger.go
Normal file
@@ -0,0 +1,97 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"zeromesh/internal/model"
|
||||
)
|
||||
|
||||
type LogWriter interface {
|
||||
InsertApiLog(*model.ApiLog) error
|
||||
InsertErrorLog(*model.ErrorLog) error
|
||||
}
|
||||
|
||||
type RequestLogger struct {
|
||||
writer LogWriter
|
||||
}
|
||||
|
||||
func NewRequestLogger(writer LogWriter) *RequestLogger {
|
||||
return &RequestLogger{writer: writer}
|
||||
}
|
||||
|
||||
func (l *RequestLogger) Handler() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
path := c.Request.URL.Path
|
||||
|
||||
if strings.HasPrefix(path, "/api/v1/log") {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
start := time.Now()
|
||||
|
||||
body := ""
|
||||
if c.Request.Body != nil {
|
||||
b, _ := io.ReadAll(c.Request.Body)
|
||||
body = string(b)
|
||||
c.Request.Body = io.NopCloser(bytes.NewBuffer(b))
|
||||
}
|
||||
|
||||
blw := &bodyLogWriter{ResponseWriter: c.Writer, buf: &bytes.Buffer{}}
|
||||
c.Writer = blw
|
||||
|
||||
c.Next()
|
||||
|
||||
latencyMs := time.Since(start).Milliseconds()
|
||||
ip := c.ClientIP()
|
||||
method := c.Request.Method
|
||||
query := c.Request.URL.RawQuery
|
||||
ua := c.Request.UserAgent()
|
||||
status := c.Writer.Status()
|
||||
|
||||
if status >= 400 {
|
||||
el := &model.ErrorLog{
|
||||
IP: ip,
|
||||
Method: method,
|
||||
Path: path,
|
||||
StatusCode: status,
|
||||
Query: query,
|
||||
UserAgent: ua,
|
||||
RequestBody: body,
|
||||
ResponseBody: blw.buf.String(),
|
||||
LatencyMs: latencyMs,
|
||||
}
|
||||
if err := l.writer.InsertErrorLog(el); err != nil {
|
||||
c.Error(err)
|
||||
}
|
||||
}
|
||||
|
||||
al := &model.ApiLog{
|
||||
IP: ip,
|
||||
Method: method,
|
||||
Path: path,
|
||||
StatusCode: status,
|
||||
Query: query,
|
||||
UserAgent: ua,
|
||||
LatencyMs: latencyMs,
|
||||
}
|
||||
if err := l.writer.InsertApiLog(al); err != nil {
|
||||
c.Error(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type bodyLogWriter struct {
|
||||
gin.ResponseWriter
|
||||
buf *bytes.Buffer
|
||||
}
|
||||
|
||||
func (w *bodyLogWriter) Write(b []byte) (int, error) {
|
||||
w.buf.Write(b)
|
||||
return w.ResponseWriter.Write(b)
|
||||
}
|
||||
64
internal/middleware/ratelimit.go
Normal file
64
internal/middleware/ratelimit.go
Normal file
@@ -0,0 +1,64 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type RateLimiter struct {
|
||||
mu sync.Mutex
|
||||
clients map[string]*clientBuckets
|
||||
rate int
|
||||
burst int
|
||||
}
|
||||
|
||||
type clientBuckets struct {
|
||||
tokens int
|
||||
lastFill time.Time
|
||||
}
|
||||
|
||||
func NewRateLimiter(rate, burst int) *RateLimiter {
|
||||
return &RateLimiter{
|
||||
clients: make(map[string]*clientBuckets),
|
||||
rate: rate,
|
||||
burst: burst,
|
||||
}
|
||||
}
|
||||
|
||||
func (rl *RateLimiter) Handler() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if !rl.allow(c.ClientIP()) {
|
||||
c.AbortWithStatusJSON(http.StatusTooManyRequests, gin.H{"error": "rate limit exceeded"})
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
func (rl *RateLimiter) allow(key string) bool {
|
||||
rl.mu.Lock()
|
||||
defer rl.mu.Unlock()
|
||||
|
||||
b, ok := rl.clients[key]
|
||||
if !ok {
|
||||
b = &clientBuckets{tokens: rl.burst, lastFill: time.Now()}
|
||||
rl.clients[key] = b
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
elapsed := now.Sub(b.lastFill)
|
||||
b.lastFill = now
|
||||
b.tokens += int(elapsed.Seconds()) * rl.rate
|
||||
if b.tokens > rl.burst {
|
||||
b.tokens = rl.burst
|
||||
}
|
||||
|
||||
if b.tokens > 0 {
|
||||
b.tokens--
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
101
internal/model/models.go
Normal file
101
internal/model/models.go
Normal file
@@ -0,0 +1,101 @@
|
||||
package model
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type User struct {
|
||||
ID uint `gorm:"primarykey" json:"id"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
DeletedAt gorm.DeletedAt `gorm:"index" json:"-"`
|
||||
Username string `gorm:"uniqueIndex;size:64" json:"username"`
|
||||
Password string `json:"-"`
|
||||
Role string `gorm:"default:user;size:16" json:"role"`
|
||||
Token string `json:"-"`
|
||||
QuotaNetworks int `gorm:"default:3" json:"quota_networks"`
|
||||
QuotaNodes int `gorm:"default:20" json:"quota_nodes"`
|
||||
}
|
||||
|
||||
func (u *User) IsAdmin() bool {
|
||||
return u.Role == "admin"
|
||||
}
|
||||
|
||||
type Network struct {
|
||||
ID uint `gorm:"primarykey" json:"id"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
DeletedAt gorm.DeletedAt `gorm:"index" json:"-"`
|
||||
UserID uint `gorm:"index;not null" json:"user_id"`
|
||||
NetworkID uint32 `gorm:"uniqueIndex;not null" json:"network_id"`
|
||||
Name string `gorm:"size:128" json:"name"`
|
||||
IPRange string `gorm:"size:64" json:"ip_range"`
|
||||
IP6Range string `gorm:"size:64" json:"ip6_range"`
|
||||
MTU int `gorm:"default:2800" json:"mtu"`
|
||||
Multicast bool `gorm:"default:true" json:"multicast"`
|
||||
Private bool `gorm:"default:true" json:"private"`
|
||||
Members []NetworkMember `json:"members,omitempty"`
|
||||
}
|
||||
|
||||
type NetworkMember struct {
|
||||
ID uint `gorm:"primarykey" json:"id"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
NetworkID uint32 `gorm:"index;not null" json:"network_id"`
|
||||
NodeID string `gorm:"size:40;index;not null" json:"node_id"`
|
||||
IPAddress string `gorm:"size:64" json:"ip_address"`
|
||||
Authorized bool `gorm:"default:false" json:"authorized"`
|
||||
Label string `gorm:"size:128" json:"label"`
|
||||
}
|
||||
|
||||
type Node struct {
|
||||
ID uint `gorm:"primarykey" json:"id"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
DeletedAt gorm.DeletedAt `gorm:"index" json:"-"`
|
||||
UserID uint `gorm:"index;not null" json:"user_id"`
|
||||
NodeID string `gorm:"uniqueIndex;size:40" json:"node_id"`
|
||||
PublicKey string `gorm:"size:128" json:"public_key"`
|
||||
Name string `gorm:"size:128" json:"name"`
|
||||
IPAddress string `gorm:"size:64" json:"ip_address"`
|
||||
Port int `json:"port"`
|
||||
Online bool `gorm:"default:false" json:"online"`
|
||||
LastSeen *time.Time `json:"last_seen"`
|
||||
Version string `gorm:"size:32" json:"version"`
|
||||
}
|
||||
|
||||
type Config struct {
|
||||
ID uint `gorm:"primarykey" json:"id"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
Key string `gorm:"uniqueIndex;size:128" json:"key"`
|
||||
Value string `gorm:"size:1024" json:"value"`
|
||||
}
|
||||
|
||||
type ApiLog struct {
|
||||
ID uint `gorm:"primarykey" json:"id"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
IP string `gorm:"size:64;index" json:"ip"`
|
||||
Method string `gorm:"size:10" json:"method"`
|
||||
Path string `gorm:"size:512;index" json:"path"`
|
||||
StatusCode int `json:"status_code"`
|
||||
Query string `gorm:"size:1024" json:"query"`
|
||||
UserAgent string `gorm:"size:512" json:"user_agent"`
|
||||
LatencyMs int64 `json:"latency_ms"`
|
||||
}
|
||||
|
||||
type ErrorLog struct {
|
||||
ID uint `gorm:"primarykey" json:"id"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
IP string `gorm:"size:64;index" json:"ip"`
|
||||
Method string `gorm:"size:10" json:"method"`
|
||||
Path string `gorm:"size:512;index" json:"path"`
|
||||
StatusCode int `json:"status_code"`
|
||||
Query string `gorm:"size:1024" json:"query"`
|
||||
UserAgent string `gorm:"size:512" json:"user_agent"`
|
||||
RequestBody string `gorm:"type:text" json:"request_body"`
|
||||
ResponseBody string `gorm:"type:text" json:"response_body"`
|
||||
LatencyMs int64 `json:"latency_ms"`
|
||||
}
|
||||
111
internal/repository/log.go
Normal file
111
internal/repository/log.go
Normal file
@@ -0,0 +1,111 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"zeromesh/internal/database"
|
||||
"zeromesh/internal/model"
|
||||
)
|
||||
|
||||
type LogRepo struct {
|
||||
db *database.DB
|
||||
}
|
||||
|
||||
func NewLogRepo(db *database.DB) *LogRepo {
|
||||
return &LogRepo{db: db}
|
||||
}
|
||||
|
||||
func (r *LogRepo) InsertApiLog(log *model.ApiLog) error {
|
||||
return r.db.Create(log).Error
|
||||
}
|
||||
|
||||
func (r *LogRepo) InsertErrorLog(log *model.ErrorLog) error {
|
||||
return r.db.Create(log).Error
|
||||
}
|
||||
|
||||
func (r *LogRepo) ListApiLogs(page, pageSize int, ip, path, method string, statusCode int) ([]model.ApiLog, int64, error) {
|
||||
q := r.db.Model(&model.ApiLog{})
|
||||
if ip != "" {
|
||||
q = q.Where("ip LIKE ?", "%"+ip+"%")
|
||||
}
|
||||
if path != "" {
|
||||
q = q.Where("path LIKE ?", "%"+path+"%")
|
||||
}
|
||||
if method != "" {
|
||||
q = q.Where("method = ?", method)
|
||||
}
|
||||
if statusCode > 0 {
|
||||
q = q.Where("status_code = ?", statusCode)
|
||||
}
|
||||
var total int64
|
||||
if err := q.Count(&total).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
var logs []model.ApiLog
|
||||
offset := (page - 1) * pageSize
|
||||
if err := q.Order("id DESC").Offset(offset).Limit(pageSize).Find(&logs).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
return logs, total, nil
|
||||
}
|
||||
|
||||
func (r *LogRepo) ListErrorLogs(page, pageSize int, ip, path, method string, statusCode int) ([]model.ErrorLog, int64, error) {
|
||||
q := r.db.Model(&model.ErrorLog{})
|
||||
if ip != "" {
|
||||
q = q.Where("ip LIKE ?", "%"+ip+"%")
|
||||
}
|
||||
if path != "" {
|
||||
q = q.Where("path LIKE ?", "%"+path+"%")
|
||||
}
|
||||
if method != "" {
|
||||
q = q.Where("method = ?", method)
|
||||
}
|
||||
if statusCode > 0 {
|
||||
q = q.Where("status_code = ?", statusCode)
|
||||
}
|
||||
var total int64
|
||||
if err := q.Count(&total).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
var logs []model.ErrorLog
|
||||
offset := (page - 1) * pageSize
|
||||
if err := q.Order("id DESC").Offset(offset).Limit(pageSize).Find(&logs).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
return logs, total, nil
|
||||
}
|
||||
|
||||
func (r *LogRepo) ListApiLogGrouped() ([]map[string]interface{}, error) {
|
||||
rows, err := r.db.Model(&model.ApiLog{}).
|
||||
Select("ip, method, path, COUNT(*) as count, SUM(latency_ms) as total_latency").
|
||||
Group("ip, method, path").
|
||||
Order("count DESC").
|
||||
Rows()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var result []map[string]interface{}
|
||||
for rows.Next() {
|
||||
var ip, method, path string
|
||||
var count int64
|
||||
var totalLatency int64
|
||||
if err := rows.Scan(&ip, &method, &path, &count, &totalLatency); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result = append(result, map[string]interface{}{
|
||||
"ip": ip,
|
||||
"method": method,
|
||||
"path": path,
|
||||
"count": count,
|
||||
"total_latency": totalLatency,
|
||||
})
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (r *LogRepo) ClearApiLogs() error {
|
||||
return r.db.Where("1 = 1").Delete(&model.ApiLog{}).Error
|
||||
}
|
||||
|
||||
func (r *LogRepo) ClearErrorLogs() error {
|
||||
return r.db.Where("1 = 1").Delete(&model.ErrorLog{}).Error
|
||||
}
|
||||
82
internal/repository/network.go
Normal file
82
internal/repository/network.go
Normal file
@@ -0,0 +1,82 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"zeromesh/internal/database"
|
||||
"zeromesh/internal/model"
|
||||
)
|
||||
|
||||
type NetworkRepo struct {
|
||||
db *database.DB
|
||||
}
|
||||
|
||||
func NewNetworkRepo(db *database.DB) *NetworkRepo {
|
||||
return &NetworkRepo{db: db}
|
||||
}
|
||||
|
||||
func (r *NetworkRepo) Create(network *model.Network) error {
|
||||
return r.db.Create(network).Error
|
||||
}
|
||||
|
||||
func (r *NetworkRepo) FindByNetworkID(networkID uint32) (*model.Network, error) {
|
||||
var network model.Network
|
||||
err := r.db.Where("network_id = ?", networkID).Preload("Members").First(&network).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &network, nil
|
||||
}
|
||||
|
||||
func (r *NetworkRepo) List() ([]model.Network, error) {
|
||||
var networks []model.Network
|
||||
err := r.db.Preload("Members").Find(&networks).Error
|
||||
return networks, err
|
||||
}
|
||||
|
||||
func (r *NetworkRepo) ListByUser(userID uint) ([]model.Network, error) {
|
||||
var networks []model.Network
|
||||
err := r.db.Where("user_id = ?", userID).Preload("Members").Find(&networks).Error
|
||||
return networks, err
|
||||
}
|
||||
|
||||
func (r *NetworkRepo) FindByNetworkIDForUser(networkID uint32, userID uint) (*model.Network, error) {
|
||||
var network model.Network
|
||||
err := r.db.Where("network_id = ? AND user_id = ?", networkID, userID).Preload("Members").First(&network).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &network, nil
|
||||
}
|
||||
|
||||
func (r *NetworkRepo) CountByUser(userID uint) (int64, error) {
|
||||
var count int64
|
||||
err := r.db.Model(&model.Network{}).Where("user_id = ?", userID).Count(&count).Error
|
||||
return count, err
|
||||
}
|
||||
|
||||
func (r *NetworkRepo) Update(network *model.Network) error {
|
||||
return r.db.Save(network).Error
|
||||
}
|
||||
|
||||
func (r *NetworkRepo) Delete(networkID uint32) error {
|
||||
return r.db.Where("network_id = ?", networkID).Delete(&model.Network{}).Error
|
||||
}
|
||||
|
||||
func (r *NetworkRepo) AddMember(member *model.NetworkMember) error {
|
||||
return r.db.Create(member).Error
|
||||
}
|
||||
|
||||
func (r *NetworkRepo) RemoveMember(networkID uint32, nodeID string) error {
|
||||
return r.db.Where("network_id = ? AND node_id = ?", networkID, nodeID).Delete(&model.NetworkMember{}).Error
|
||||
}
|
||||
|
||||
func (r *NetworkRepo) FindMembers(networkID uint32) ([]model.NetworkMember, error) {
|
||||
var members []model.NetworkMember
|
||||
err := r.db.Where("network_id = ?", networkID).Find(&members).Error
|
||||
return members, err
|
||||
}
|
||||
|
||||
func (r *NetworkRepo) Count() (int64, error) {
|
||||
var count int64
|
||||
err := r.db.Model(&model.Network{}).Count(&count).Error
|
||||
return count, err
|
||||
}
|
||||
75
internal/repository/node.go
Normal file
75
internal/repository/node.go
Normal file
@@ -0,0 +1,75 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"zeromesh/internal/database"
|
||||
"zeromesh/internal/model"
|
||||
)
|
||||
|
||||
type NodeRepo struct {
|
||||
db *database.DB
|
||||
}
|
||||
|
||||
func NewNodeRepo(db *database.DB) *NodeRepo {
|
||||
return &NodeRepo{db: db}
|
||||
}
|
||||
|
||||
func (r *NodeRepo) Create(node *model.Node) error {
|
||||
return r.db.Create(node).Error
|
||||
}
|
||||
|
||||
func (r *NodeRepo) FindByNodeID(nodeID string) (*model.Node, error) {
|
||||
var node model.Node
|
||||
err := r.db.Where("node_id = ?", nodeID).First(&node).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &node, nil
|
||||
}
|
||||
|
||||
func (r *NodeRepo) Upsert(node *model.Node) error {
|
||||
var existing model.Node
|
||||
result := r.db.Where("node_id = ?", node.NodeID).First(&existing)
|
||||
if result.Error != nil {
|
||||
return r.db.Create(node).Error
|
||||
}
|
||||
existing.PublicKey = node.PublicKey
|
||||
existing.Name = node.Name
|
||||
existing.IPAddress = node.IPAddress
|
||||
existing.Port = node.Port
|
||||
existing.Version = node.Version
|
||||
return r.db.Save(&existing).Error
|
||||
}
|
||||
|
||||
func (r *NodeRepo) SetOnline(nodeID string, online bool) error {
|
||||
now := time.Now()
|
||||
return r.db.Model(&model.Node{}).Where("node_id = ?", nodeID).Updates(map[string]interface{}{
|
||||
"online": online,
|
||||
"last_seen": now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *NodeRepo) List() ([]model.Node, error) {
|
||||
var nodes []model.Node
|
||||
err := r.db.Order("id desc").Find(&nodes).Error
|
||||
return nodes, err
|
||||
}
|
||||
|
||||
func (r *NodeRepo) ListByUser(userID uint) ([]model.Node, error) {
|
||||
var nodes []model.Node
|
||||
err := r.db.Where("user_id = ?", userID).Order("id desc").Find(&nodes).Error
|
||||
return nodes, err
|
||||
}
|
||||
|
||||
func (r *NodeRepo) CountByUser(userID uint) (int64, error) {
|
||||
var count int64
|
||||
err := r.db.Model(&model.Node{}).Where("user_id = ?", userID).Count(&count).Error
|
||||
return count, err
|
||||
}
|
||||
|
||||
func (r *NodeRepo) ListOnline() ([]model.Node, error) {
|
||||
var nodes []model.Node
|
||||
err := r.db.Where("online = ?", true).Find(&nodes).Error
|
||||
return nodes, err
|
||||
}
|
||||
58
internal/repository/user.go
Normal file
58
internal/repository/user.go
Normal file
@@ -0,0 +1,58 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"zeromesh/internal/database"
|
||||
"zeromesh/internal/model"
|
||||
)
|
||||
|
||||
type UserRepo struct {
|
||||
db *database.DB
|
||||
}
|
||||
|
||||
func NewUserRepo(db *database.DB) *UserRepo {
|
||||
return &UserRepo{db: db}
|
||||
}
|
||||
|
||||
func (r *UserRepo) FindByUsername(username string) (*model.User, error) {
|
||||
var user model.User
|
||||
err := r.db.Where("username = ?", username).First(&user).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
func (r *UserRepo) FindByID(id uint) (*model.User, error) {
|
||||
var user model.User
|
||||
err := r.db.First(&user, id).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
func (r *UserRepo) Create(user *model.User) error {
|
||||
return r.db.Create(user).Error
|
||||
}
|
||||
|
||||
func (r *UserRepo) Count() (int64, error) {
|
||||
var count int64
|
||||
err := r.db.Model(&model.User{}).Count(&count).Error
|
||||
return count, err
|
||||
}
|
||||
|
||||
func (r *UserRepo) HasAdmin() (bool, error) {
|
||||
var count int64
|
||||
err := r.db.Model(&model.User{}).Where("role = ?", "admin").Count(&count).Error
|
||||
return count > 0, err
|
||||
}
|
||||
|
||||
func (r *UserRepo) UpdatePassword(id uint, hashed string) error {
|
||||
return r.db.Model(&model.User{}).Where("id = ?", id).Update("password", hashed).Error
|
||||
}
|
||||
|
||||
func (r *UserRepo) List(offset, limit int) ([]model.User, error) {
|
||||
var users []model.User
|
||||
err := r.db.Offset(offset).Limit(limit).Order("id desc").Find(&users).Error
|
||||
return users, err
|
||||
}
|
||||
124
internal/router/router.go
Normal file
124
internal/router/router.go
Normal file
@@ -0,0 +1,124 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"zeromesh/internal/config"
|
||||
"zeromesh/internal/database"
|
||||
"zeromesh/internal/handler"
|
||||
"zeromesh/internal/middleware"
|
||||
"zeromesh/internal/repository"
|
||||
"zeromesh/internal/service"
|
||||
)
|
||||
|
||||
func Setup(cfg *config.Config, db *database.DB) *gin.Engine {
|
||||
r := gin.Default()
|
||||
|
||||
// Repositories
|
||||
userRepo := repository.NewUserRepo(db)
|
||||
netRepo := repository.NewNetworkRepo(db)
|
||||
nodeRepo := repository.NewNodeRepo(db)
|
||||
logRepo := repository.NewLogRepo(db)
|
||||
|
||||
// Services
|
||||
authService := service.NewAuthService(userRepo, cfg.JWT)
|
||||
netService := service.NewNetworkService(netRepo, nodeRepo, userRepo)
|
||||
nodeService := service.NewNodeService(nodeRepo, userRepo)
|
||||
logService := service.NewLogService(logRepo)
|
||||
|
||||
// Global middleware
|
||||
r.Use(middleware.CORS(cfg.CORS.AllowedOrigins))
|
||||
if cfg.RateLimit.Enabled {
|
||||
rl := middleware.NewRateLimiter(cfg.RateLimit.RequestsPerSecond, cfg.RateLimit.BurstSize)
|
||||
r.Use(rl.Handler())
|
||||
}
|
||||
requestLogger := middleware.NewRequestLogger(logService)
|
||||
r.Use(requestLogger.Handler())
|
||||
|
||||
// Handlers
|
||||
authHandler := handler.NewAuthHandler(authService)
|
||||
netHandler := handler.NewNetworkHandler(netService)
|
||||
nodeHandler := handler.NewNodeHandler(nodeService)
|
||||
adminHandler := handler.NewAdminHandler(authService, db, nodeService, netService)
|
||||
logHandler := handler.NewLogHandler(logService)
|
||||
|
||||
// Logging API — skip logger middleware (avoid recursion)
|
||||
logApi := r.Group("/api/v1/log")
|
||||
logApi.Use(middleware.JWTAuth(cfg.JWT.Secret))
|
||||
{
|
||||
logApi.GET("/api", logHandler.ListApiLogs)
|
||||
logApi.GET("/error", logHandler.ListErrorLogs)
|
||||
logApi.GET("/grouped", logHandler.ListApiLogGrouped)
|
||||
logApi.DELETE("/api", logHandler.ClearApiLogs)
|
||||
logApi.DELETE("/error", logHandler.ClearErrorLogs)
|
||||
}
|
||||
|
||||
// Public routes
|
||||
r.GET("/api/health", func(c *gin.Context) {
|
||||
c.JSON(200, gin.H{"status": "ok"})
|
||||
})
|
||||
r.GET("/api/v1/admin/check", authHandler.CheckAdmin)
|
||||
|
||||
api := r.Group("/api/v1")
|
||||
{
|
||||
api.POST("/auth/login", authHandler.Login)
|
||||
api.POST("/auth/register", authHandler.Register)
|
||||
api.POST("/auth/init", authHandler.InitAdmin)
|
||||
}
|
||||
|
||||
// Authenticated routes
|
||||
auth := r.Group("/api/v1")
|
||||
auth.Use(middleware.JWTAuth(cfg.JWT.Secret))
|
||||
{
|
||||
// Dashboard (user + admin)
|
||||
auth.GET("/dashboard", adminHandler.Dashboard)
|
||||
|
||||
// Profile
|
||||
auth.GET("/user/profile", func(c *gin.Context) {
|
||||
uid := c.GetUint("user_id")
|
||||
u, err := userRepo.FindByID(uid)
|
||||
if err != nil {
|
||||
c.JSON(404, gin.H{"error": "user not found"})
|
||||
return
|
||||
}
|
||||
netCount, _ := netRepo.CountByUser(uid)
|
||||
nodeCount, _ := nodeRepo.CountByUser(uid)
|
||||
c.JSON(200, gin.H{
|
||||
"user": u,
|
||||
"used_networks": netCount,
|
||||
"used_nodes": nodeCount,
|
||||
})
|
||||
})
|
||||
|
||||
// Nodes (user-scoped)
|
||||
auth.POST("/node/register", nodeHandler.Register)
|
||||
auth.GET("/node/list", nodeHandler.List)
|
||||
auth.GET("/node/online", nodeHandler.ListOnline)
|
||||
|
||||
// Networks (user-scoped)
|
||||
auth.POST("/network/create", netHandler.Create)
|
||||
auth.GET("/network/list", netHandler.List)
|
||||
auth.GET("/network/:id", netHandler.Get)
|
||||
auth.DELETE("/network/:id", netHandler.Delete)
|
||||
auth.GET("/network/:id/members", netHandler.Members)
|
||||
auth.POST("/network/:id/authorize", netHandler.Authorize)
|
||||
auth.POST("/network/:id/deauthorize", netHandler.Deauthorize)
|
||||
|
||||
// Admin-only
|
||||
admin := auth.Group("/admin")
|
||||
admin.Use(func(c *gin.Context) {
|
||||
if c.GetString("role") != "admin" {
|
||||
c.AbortWithStatusJSON(403, gin.H{"error": "forbidden"})
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
})
|
||||
{
|
||||
admin.GET("/dashboard", adminHandler.Dashboard)
|
||||
admin.GET("/networks", netHandler.List)
|
||||
admin.GET("/nodes", nodeHandler.List)
|
||||
}
|
||||
}
|
||||
|
||||
return r
|
||||
}
|
||||
113
internal/service/auth.go
Normal file
113
internal/service/auth.go
Normal file
@@ -0,0 +1,113 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
|
||||
"zeromesh/internal/config"
|
||||
"zeromesh/internal/middleware"
|
||||
"zeromesh/internal/model"
|
||||
"zeromesh/internal/repository"
|
||||
)
|
||||
|
||||
type AuthService struct {
|
||||
userRepo *repository.UserRepo
|
||||
jwtCfg config.JWTConfig
|
||||
}
|
||||
|
||||
func NewAuthService(userRepo *repository.UserRepo, jwtCfg config.JWTConfig) *AuthService {
|
||||
return &AuthService{userRepo: userRepo, jwtCfg: jwtCfg}
|
||||
}
|
||||
|
||||
func (s *AuthService) Login(username, password string) (string, *model.User, error) {
|
||||
user, err := s.userRepo.FindByUsername(username)
|
||||
if err != nil {
|
||||
return "", nil, errors.New("invalid credentials")
|
||||
}
|
||||
// Migrate old SHA256 hashes to bcrypt on login
|
||||
if len(user.Password) < 4 || user.Password[:4] != "$2a$" {
|
||||
if sha256Hash(password) != user.Password {
|
||||
return "", nil, errors.New("invalid credentials")
|
||||
}
|
||||
hashed, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("migrate password: %w", err)
|
||||
}
|
||||
user.Password = string(hashed)
|
||||
_ = s.userRepo.UpdatePassword(user.ID, user.Password)
|
||||
} else if err := bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(password)); err != nil {
|
||||
return "", nil, errors.New("invalid credentials")
|
||||
}
|
||||
token, err := middleware.GenerateToken(s.jwtCfg.Secret, user.ID, user.Username, user.Role, s.jwtCfg.ExpireHours)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("generate token: %w", err)
|
||||
}
|
||||
user.Token = token
|
||||
return token, user, nil
|
||||
}
|
||||
|
||||
func (s *AuthService) Register(username, password string) (*model.User, string, error) {
|
||||
hashed, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("hash password: %w", err)
|
||||
}
|
||||
user := &model.User{
|
||||
Username: username,
|
||||
Password: string(hashed),
|
||||
Role: "user",
|
||||
}
|
||||
if err := s.userRepo.Create(user); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
token, err := middleware.GenerateToken(s.jwtCfg.Secret, user.ID, user.Username, user.Role, s.jwtCfg.ExpireHours)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("generate token: %w", err)
|
||||
}
|
||||
user.Token = token
|
||||
return user, token, nil
|
||||
}
|
||||
|
||||
func (s *AuthService) InitAdmin(username, password string) (*model.User, string, error) {
|
||||
count, err := s.userRepo.Count()
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
if count > 0 {
|
||||
return nil, "", fmt.Errorf("admin already exists")
|
||||
}
|
||||
hashed, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("hash password: %w", err)
|
||||
}
|
||||
user := &model.User{
|
||||
Username: username,
|
||||
Password: string(hashed),
|
||||
Role: "admin",
|
||||
QuotaNetworks: -1,
|
||||
QuotaNodes: -1,
|
||||
}
|
||||
if err := s.userRepo.Create(user); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
token, err := middleware.GenerateToken(s.jwtCfg.Secret, user.ID, user.Username, user.Role, s.jwtCfg.ExpireHours)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("generate token: %w", err)
|
||||
}
|
||||
user.Token = token
|
||||
return user, token, nil
|
||||
}
|
||||
|
||||
func sha256Hash(password string) string {
|
||||
h := sha256.Sum256([]byte(password))
|
||||
return hex.EncodeToString(h[:])
|
||||
}
|
||||
|
||||
func (s *AuthService) CheckAdmin() (bool, error) {
|
||||
return s.userRepo.HasAdmin()
|
||||
}
|
||||
|
||||
|
||||
54
internal/service/log.go
Normal file
54
internal/service/log.go
Normal file
@@ -0,0 +1,54 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"zeromesh/internal/model"
|
||||
"zeromesh/internal/repository"
|
||||
)
|
||||
|
||||
type LogService struct {
|
||||
repo *repository.LogRepo
|
||||
}
|
||||
|
||||
func NewLogService(repo *repository.LogRepo) *LogService {
|
||||
return &LogService{repo: repo}
|
||||
}
|
||||
|
||||
func (s *LogService) ListApiLogs(page, pageSize int, ip, path, method string, statusCode int) ([]model.ApiLog, int64, error) {
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if pageSize < 1 || pageSize > 200 {
|
||||
pageSize = 20
|
||||
}
|
||||
return s.repo.ListApiLogs(page, pageSize, ip, path, method, statusCode)
|
||||
}
|
||||
|
||||
func (s *LogService) ListErrorLogs(page, pageSize int, ip, path, method string, statusCode int) ([]model.ErrorLog, int64, error) {
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if pageSize < 1 || pageSize > 200 {
|
||||
pageSize = 20
|
||||
}
|
||||
return s.repo.ListErrorLogs(page, pageSize, ip, path, method, statusCode)
|
||||
}
|
||||
|
||||
func (s *LogService) ListApiLogGrouped() ([]map[string]interface{}, error) {
|
||||
return s.repo.ListApiLogGrouped()
|
||||
}
|
||||
|
||||
func (s *LogService) InsertApiLog(log *model.ApiLog) error {
|
||||
return s.repo.InsertApiLog(log)
|
||||
}
|
||||
|
||||
func (s *LogService) InsertErrorLog(log *model.ErrorLog) error {
|
||||
return s.repo.InsertErrorLog(log)
|
||||
}
|
||||
|
||||
func (s *LogService) ClearApiLogs() error {
|
||||
return s.repo.ClearApiLogs()
|
||||
}
|
||||
|
||||
func (s *LogService) ClearErrorLogs() error {
|
||||
return s.repo.ClearErrorLogs()
|
||||
}
|
||||
163
internal/service/network.go
Normal file
163
internal/service/network.go
Normal file
@@ -0,0 +1,163 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math/rand"
|
||||
"net"
|
||||
|
||||
"zeromesh/internal/model"
|
||||
"zeromesh/internal/repository"
|
||||
)
|
||||
|
||||
type NetworkService struct {
|
||||
netRepo *repository.NetworkRepo
|
||||
nodeRepo *repository.NodeRepo
|
||||
userRepo *repository.UserRepo
|
||||
}
|
||||
|
||||
func NewNetworkService(netRepo *repository.NetworkRepo, nodeRepo *repository.NodeRepo, userRepo *repository.UserRepo) *NetworkService {
|
||||
return &NetworkService{netRepo: netRepo, nodeRepo: nodeRepo, userRepo: userRepo}
|
||||
}
|
||||
|
||||
func randomPrivateSubnet() string {
|
||||
const (
|
||||
_10 = iota // 10.x.y.0/24
|
||||
_172 // 172.16-31.y.0/24
|
||||
_192 // 192.168.y.0/24
|
||||
)
|
||||
class := rand.Intn(3)
|
||||
var b1, b2 byte
|
||||
switch class {
|
||||
case _10:
|
||||
b1 = byte(rand.Intn(256))
|
||||
b2 = byte(rand.Intn(256))
|
||||
return fmt.Sprintf("10.%d.%d.0/24", b1, b2)
|
||||
case _172:
|
||||
b1 = byte(16 + rand.Intn(16)) // 16-31
|
||||
b2 = byte(rand.Intn(256))
|
||||
return fmt.Sprintf("172.%d.%d.0/24", b1, b2)
|
||||
default:
|
||||
b2 = byte(rand.Intn(256))
|
||||
return fmt.Sprintf("192.168.%d.0/24", b2)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *NetworkService) CreateNetwork(userID uint, name, ipRange string) (*model.Network, error) {
|
||||
user, err := s.userRepo.FindByID(userID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("user not found")
|
||||
}
|
||||
if user.QuotaNetworks >= 0 {
|
||||
used, _ := s.netRepo.CountByUser(userID)
|
||||
if used >= int64(user.QuotaNetworks) {
|
||||
return nil, fmt.Errorf("network quota exceeded (%d)", user.QuotaNetworks)
|
||||
}
|
||||
}
|
||||
if ipRange == "" {
|
||||
ipRange = randomPrivateSubnet()
|
||||
}
|
||||
id := rand.Uint32()
|
||||
network := &model.Network{
|
||||
UserID: userID,
|
||||
NetworkID: id,
|
||||
Name: name,
|
||||
IPRange: ipRange,
|
||||
MTU: 2800,
|
||||
Multicast: true,
|
||||
Private: true,
|
||||
}
|
||||
if err := s.netRepo.Create(network); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return network, nil
|
||||
}
|
||||
|
||||
func (s *NetworkService) ListNetworks(userID uint) ([]model.Network, error) {
|
||||
if userID == 0 {
|
||||
return s.netRepo.List()
|
||||
}
|
||||
return s.netRepo.ListByUser(userID)
|
||||
}
|
||||
|
||||
func (s *NetworkService) GetNetwork(networkID uint32, userID uint) (*model.Network, error) {
|
||||
if userID == 0 {
|
||||
return s.netRepo.FindByNetworkID(networkID)
|
||||
}
|
||||
return s.netRepo.FindByNetworkIDForUser(networkID, userID)
|
||||
}
|
||||
|
||||
func (s *NetworkService) DeleteNetwork(networkID uint32, userID uint) error {
|
||||
net, err := s.GetNetwork(networkID, userID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return s.netRepo.Delete(net.NetworkID)
|
||||
}
|
||||
|
||||
func (s *NetworkService) AuthorizeMember(networkID uint32, nodeID string, userID uint) error {
|
||||
network, err := s.GetNetwork(networkID, userID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
node, err := s.nodeRepo.FindByNodeID(nodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if node.UserID != userID && userID != 0 {
|
||||
return fmt.Errorf("node does not belong to user")
|
||||
}
|
||||
ip, err := s.allocateIP(network.IPRange, networkID, nodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
member := &model.NetworkMember{
|
||||
NetworkID: networkID,
|
||||
NodeID: nodeID,
|
||||
IPAddress: ip,
|
||||
Authorized: true,
|
||||
Label: node.Name,
|
||||
}
|
||||
if err := s.netRepo.AddMember(member); err != nil {
|
||||
return fmt.Errorf("add member: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *NetworkService) DeauthorizeMember(networkID uint32, nodeID string, userID uint) error {
|
||||
_, err := s.GetNetwork(networkID, userID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return s.netRepo.RemoveMember(networkID, nodeID)
|
||||
}
|
||||
|
||||
func (s *NetworkService) ListMembers(networkID uint32, userID uint) ([]model.NetworkMember, error) {
|
||||
_, err := s.GetNetwork(networkID, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.netRepo.FindMembers(networkID)
|
||||
}
|
||||
|
||||
func (s *NetworkService) allocateIP(cidr string, networkID uint32, nodeID string) (string, error) {
|
||||
_, ipNet, err := net.ParseCIDR(cidr)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
ones, bits := ipNet.Mask.Size()
|
||||
ip := ipNet.IP.To4()
|
||||
if ip == nil {
|
||||
return "", fmt.Errorf("IPv4 only")
|
||||
}
|
||||
// simple allocation: use last 2 bytes of nodeID hash
|
||||
h := uint32(0)
|
||||
for _, b := range []byte(nodeID) {
|
||||
h = h*31 + uint32(b)
|
||||
}
|
||||
hostBits := bits - ones
|
||||
maxHosts := (1 << uint(hostBits)) - 2
|
||||
hostOffset := int(h%uint32(maxHosts)) + 1
|
||||
ip[2] = byte(hostOffset >> 8)
|
||||
ip[3] = byte(hostOffset)
|
||||
return ip.String(), nil
|
||||
}
|
||||
89
internal/service/node.go
Normal file
89
internal/service/node.go
Normal file
@@ -0,0 +1,89 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"zeromesh/internal/model"
|
||||
"zeromesh/internal/repository"
|
||||
)
|
||||
|
||||
type NodeService struct {
|
||||
nodeRepo *repository.NodeRepo
|
||||
userRepo *repository.UserRepo
|
||||
}
|
||||
|
||||
func NewNodeService(nodeRepo *repository.NodeRepo, userRepo *repository.UserRepo) *NodeService {
|
||||
return &NodeService{nodeRepo: nodeRepo, userRepo: userRepo}
|
||||
}
|
||||
|
||||
func (s *NodeService) Register(nodeID, publicKey, name, ipAddress string, port int, version string, userID uint) (*model.Node, error) {
|
||||
if userID > 0 {
|
||||
user, err := s.userRepo.FindByID(userID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("user not found")
|
||||
}
|
||||
if user.QuotaNodes >= 0 {
|
||||
used, _ := s.nodeRepo.CountByUser(userID)
|
||||
if used >= int64(user.QuotaNodes) {
|
||||
return nil, fmt.Errorf("node quota exceeded (%d)", user.QuotaNodes)
|
||||
}
|
||||
}
|
||||
}
|
||||
node := &model.Node{
|
||||
UserID: userID,
|
||||
NodeID: nodeID,
|
||||
PublicKey: publicKey,
|
||||
Name: name,
|
||||
IPAddress: ipAddress,
|
||||
Port: port,
|
||||
Online: true,
|
||||
LastSeen: now(),
|
||||
Version: version,
|
||||
}
|
||||
if err := s.nodeRepo.Upsert(node); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return node, nil
|
||||
}
|
||||
|
||||
func (s *NodeService) SetOnline(nodeID string, online bool) error {
|
||||
return s.nodeRepo.SetOnline(nodeID, online)
|
||||
}
|
||||
|
||||
func (s *NodeService) ListNode(userID uint) ([]model.Node, error) {
|
||||
if userID == 0 {
|
||||
return s.nodeRepo.List()
|
||||
}
|
||||
return s.nodeRepo.ListByUser(userID)
|
||||
}
|
||||
|
||||
func (s *NodeService) ListOnline(userID uint) ([]model.Node, error) {
|
||||
if userID == 0 {
|
||||
return s.nodeRepo.ListOnline()
|
||||
}
|
||||
nodes, err := s.nodeRepo.ListByUser(userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var online []model.Node
|
||||
for _, n := range nodes {
|
||||
if n.Online {
|
||||
online = append(online, n)
|
||||
}
|
||||
}
|
||||
return online, nil
|
||||
}
|
||||
|
||||
func (s *NodeService) GetNode(nodeID string) (*model.Node, error) {
|
||||
return s.nodeRepo.FindByNodeID(nodeID)
|
||||
}
|
||||
|
||||
func (s *NodeService) CountByUser(userID uint) (int64, error) {
|
||||
return s.nodeRepo.CountByUser(userID)
|
||||
}
|
||||
|
||||
func now() *time.Time {
|
||||
t := time.Now()
|
||||
return &t
|
||||
}
|
||||
25
internal/service/user.go
Normal file
25
internal/service/user.go
Normal file
@@ -0,0 +1,25 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"zeromesh/internal/database"
|
||||
)
|
||||
|
||||
type UserService struct {
|
||||
db *database.DB
|
||||
}
|
||||
|
||||
func NewUserService(db *database.DB) *UserService {
|
||||
return &UserService{db: db}
|
||||
}
|
||||
|
||||
func (s *UserService) Count() (int64, error) {
|
||||
var count int64
|
||||
err := s.db.Model(&struct{}{}).Table("users").Count(&count).Error
|
||||
return count, err
|
||||
}
|
||||
|
||||
func (s *UserService) List(offset, limit int) ([]map[string]interface{}, error) {
|
||||
var users []map[string]interface{}
|
||||
err := s.db.Table("users").Offset(offset).Limit(limit).Order("id desc").Find(&users).Error
|
||||
return users, err
|
||||
}
|
||||
24
internal/tap/tap.go
Normal file
24
internal/tap/tap.go
Normal file
@@ -0,0 +1,24 @@
|
||||
package tap
|
||||
|
||||
type Interface struct {
|
||||
Name string
|
||||
MTU int
|
||||
fd int
|
||||
}
|
||||
|
||||
// Open creates a TAP interface on Linux. On other platforms it returns an error.
|
||||
func Open(name string, mtu int) (*Interface, error) {
|
||||
return openTap(name, mtu)
|
||||
}
|
||||
|
||||
func (t *Interface) Read(buf []byte) (int, error) {
|
||||
return readBuf(t.fd, buf)
|
||||
}
|
||||
|
||||
func (t *Interface) Write(buf []byte) (int, error) {
|
||||
return writeBuf(t.fd, buf)
|
||||
}
|
||||
|
||||
func (t *Interface) Close() error {
|
||||
return closeFD(t.fd)
|
||||
}
|
||||
71
internal/tap/tap_linux.go
Normal file
71
internal/tap/tap_linux.go
Normal file
@@ -0,0 +1,71 @@
|
||||
//go:build linux
|
||||
|
||||
package tap
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"syscall"
|
||||
"unsafe"
|
||||
)
|
||||
|
||||
const (
|
||||
cIFF_TAP = 0x0002
|
||||
cIFF_NO_PI = 0x1000
|
||||
cTUNSETIFF = 0x400454ca
|
||||
cSIOCSIFMTU = 0x8922
|
||||
)
|
||||
|
||||
func openTap(name string, mtu int) (*Interface, error) {
|
||||
fd, err := syscall.Open("/dev/net/tun", syscall.O_RDWR, 0)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open /dev/net/tun: %w (is TUN/TAP supported?)", err)
|
||||
}
|
||||
|
||||
ifr := make([]byte, 40)
|
||||
copy(ifr, []byte(name))
|
||||
*(*uint16)(unsafe.Pointer(&ifr[16])) = uint16(cIFF_TAP | cIFF_NO_PI)
|
||||
|
||||
_, _, errno := syscall.Syscall(syscall.SYS_IOCTL, uintptr(fd), cTUNSETIFF, uintptr(unsafe.Pointer(&ifr[0])))
|
||||
if errno != 0 {
|
||||
syscall.Close(fd)
|
||||
return nil, fmt.Errorf("ioctl TUNSETIFF: %w", errno)
|
||||
}
|
||||
|
||||
if mtu > 0 {
|
||||
mtuIfr := make([]byte, 40)
|
||||
copy(mtuIfr, ifr[:16])
|
||||
*(*int32)(unsafe.Pointer(&mtuIfr[16])) = int32(mtu)
|
||||
_, _, errno := syscall.Syscall(syscall.SYS_IOCTL, uintptr(fd), cSIOCSIFMTU, uintptr(unsafe.Pointer(&mtuIfr[0])))
|
||||
if errno != 0 {
|
||||
syscall.Close(fd)
|
||||
return nil, fmt.Errorf("ioctl SIOCSIFMTU: %w", errno)
|
||||
}
|
||||
}
|
||||
|
||||
return &Interface{
|
||||
Name: name,
|
||||
MTU: mtu,
|
||||
fd: fd,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func readBuf(fd int, buf []byte) (int, error) {
|
||||
n, err := syscall.Read(fd, buf)
|
||||
if err != nil {
|
||||
return 0, os.NewSyscallError("read", err)
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func writeBuf(fd int, buf []byte) (int, error) {
|
||||
n, err := syscall.Write(fd, buf)
|
||||
if err != nil {
|
||||
return 0, os.NewSyscallError("write", err)
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func closeFD(fd int) error {
|
||||
return syscall.Close(fd)
|
||||
}
|
||||
21
internal/tap/tap_stub.go
Normal file
21
internal/tap/tap_stub.go
Normal file
@@ -0,0 +1,21 @@
|
||||
//go:build !linux
|
||||
|
||||
package tap
|
||||
|
||||
import "fmt"
|
||||
|
||||
func openTap(_ string, _ int) (*Interface, error) {
|
||||
return nil, fmt.Errorf("TAP devices are only supported on Linux")
|
||||
}
|
||||
|
||||
func readBuf(_ int, _ []byte) (int, error) {
|
||||
return 0, fmt.Errorf("not supported")
|
||||
}
|
||||
|
||||
func writeBuf(_ int, _ []byte) (int, error) {
|
||||
return 0, fmt.Errorf("not supported")
|
||||
}
|
||||
|
||||
func closeFD(_ int) error {
|
||||
return fmt.Errorf("not supported")
|
||||
}
|
||||
50
internal/task/task.go
Normal file
50
internal/task/task.go
Normal file
@@ -0,0 +1,50 @@
|
||||
package task
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Task struct {
|
||||
Name string
|
||||
Interval time.Duration
|
||||
Handler func(ctx context.Context) error
|
||||
}
|
||||
|
||||
type Scheduler struct {
|
||||
tasks []Task
|
||||
log *slog.Logger
|
||||
}
|
||||
|
||||
func NewScheduler(log *slog.Logger) *Scheduler {
|
||||
return &Scheduler{log: log.With("component", "scheduler")}
|
||||
}
|
||||
|
||||
func (s *Scheduler) Add(name string, interval time.Duration, handler func(ctx context.Context) error) {
|
||||
s.tasks = append(s.tasks, Task{
|
||||
Name: name,
|
||||
Interval: interval,
|
||||
Handler: handler,
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Scheduler) Start(ctx context.Context) {
|
||||
for _, task := range s.tasks {
|
||||
go func(t Task) {
|
||||
ticker := time.NewTicker(t.Interval)
|
||||
defer ticker.Stop()
|
||||
s.log.Info("task started", "name", t.Name, "interval", t.Interval)
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
if err := t.Handler(ctx); err != nil {
|
||||
s.log.Error("task failed", "name", t.Name, "err", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
}(task)
|
||||
}
|
||||
}
|
||||
103
internal/vl1/noise.go
Normal file
103
internal/vl1/noise.go
Normal file
@@ -0,0 +1,103 @@
|
||||
package vl1
|
||||
|
||||
import (
|
||||
"crypto/cipher"
|
||||
"crypto/rand"
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"io"
|
||||
|
||||
"golang.org/x/crypto/chacha20poly1305"
|
||||
"golang.org/x/crypto/sha3"
|
||||
)
|
||||
|
||||
type NoiseCipher struct {
|
||||
sendCipher cipher.AEAD
|
||||
recvCipher cipher.AEAD
|
||||
sendNonce uint64
|
||||
recvNonce uint64
|
||||
}
|
||||
|
||||
func NewNoiseCipher(sendKey, recvKey [32]byte) *NoiseCipher {
|
||||
send, err := chacha20poly1305.New(sendKey[:])
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
recv, err := chacha20poly1305.New(recvKey[:])
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return &NoiseCipher{
|
||||
sendCipher: send,
|
||||
recvCipher: recv,
|
||||
}
|
||||
}
|
||||
|
||||
func DeriveKeysFromPSK(psk string, localPub, remotePub []byte) ([32]byte, [32]byte) {
|
||||
h := sha3.New256()
|
||||
h.Write([]byte(psk))
|
||||
h.Write(localPub)
|
||||
h.Write(remotePub)
|
||||
sum := h.Sum(nil)
|
||||
|
||||
var sendKey, recvKey [32]byte
|
||||
copy(sendKey[:], sum[:32])
|
||||
|
||||
h.Reset()
|
||||
h.Write(sum)
|
||||
h.Write([]byte("reverse"))
|
||||
rev := h.Sum(nil)
|
||||
copy(recvKey[:], rev[:32])
|
||||
|
||||
return sendKey, recvKey
|
||||
}
|
||||
|
||||
func (nc *NoiseCipher) Encrypt(plaintext []byte) ([]byte, error) {
|
||||
nonce := make([]byte, 12)
|
||||
binary.BigEndian.PutUint64(nonce[4:], nc.sendNonce)
|
||||
nc.sendNonce++
|
||||
|
||||
ciphertext := nc.sendCipher.Seal(nil, nonce, plaintext, nil)
|
||||
return ciphertext, nil
|
||||
}
|
||||
|
||||
func (nc *NoiseCipher) Decrypt(ciphertext []byte) ([]byte, error) {
|
||||
nonce := make([]byte, 12)
|
||||
binary.BigEndian.PutUint64(nonce[4:], nc.recvNonce)
|
||||
nc.recvNonce++
|
||||
|
||||
plaintext, err := nc.recvCipher.Open(nil, nonce, ciphertext, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decrypt: %w", err)
|
||||
}
|
||||
return plaintext, nil
|
||||
}
|
||||
|
||||
func (nc *NoiseCipher) EncryptTo(buf []byte, plaintext []byte) (int, error) {
|
||||
nonce := make([]byte, 12)
|
||||
binary.BigEndian.PutUint64(nonce[4:], nc.sendNonce)
|
||||
nc.sendNonce++
|
||||
|
||||
ciphertext := nc.sendCipher.Seal(buf[:0], nonce, plaintext, nil)
|
||||
return len(ciphertext), nil
|
||||
}
|
||||
|
||||
func (nc *NoiseCipher) DecryptTo(buf []byte, ciphertext []byte) ([]byte, error) {
|
||||
nonce := make([]byte, 12)
|
||||
binary.BigEndian.PutUint64(nonce[4:], nc.recvNonce)
|
||||
nc.recvNonce++
|
||||
|
||||
plaintext, err := nc.recvCipher.Open(buf[:0], nonce, ciphertext, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decrypt: %w", err)
|
||||
}
|
||||
return plaintext, nil
|
||||
}
|
||||
|
||||
func generateSessionKey() [32]byte {
|
||||
var key [32]byte
|
||||
if _, err := io.ReadFull(rand.Reader, key[:]); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return key
|
||||
}
|
||||
117
internal/vl1/packet.go
Normal file
117
internal/vl1/packet.go
Normal file
@@ -0,0 +1,117 @@
|
||||
package vl1
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
const (
|
||||
Version = 1
|
||||
MaxPacketSize = 65535
|
||||
HeaderSize = 8
|
||||
MaxFrameSize = 65535
|
||||
MinFrameSize = 14
|
||||
|
||||
PacketTypeHandshake = byte(1)
|
||||
PacketTypeData = byte(2)
|
||||
PacketTypeKeepalive = byte(3)
|
||||
)
|
||||
|
||||
type Header struct {
|
||||
Version byte
|
||||
Type byte
|
||||
NetworkID uint32
|
||||
Length uint16
|
||||
}
|
||||
|
||||
func (h *Header) Encode(buf []byte) {
|
||||
buf[0] = h.Version
|
||||
buf[1] = h.Type
|
||||
binary.BigEndian.PutUint32(buf[2:6], h.NetworkID)
|
||||
binary.BigEndian.PutUint16(buf[6:8], h.Length)
|
||||
}
|
||||
|
||||
func (h *Header) Decode(buf []byte) error {
|
||||
if len(buf) < HeaderSize {
|
||||
return fmt.Errorf("header too short: %d < %d", len(buf), HeaderSize)
|
||||
}
|
||||
h.Version = buf[0]
|
||||
h.Type = buf[1]
|
||||
h.NetworkID = binary.BigEndian.Uint32(buf[2:6])
|
||||
h.Length = binary.BigEndian.Uint16(buf[6:8])
|
||||
return nil
|
||||
}
|
||||
|
||||
type Packet struct {
|
||||
Header Header
|
||||
Payload []byte
|
||||
}
|
||||
|
||||
func NewHandshakePacket(payload []byte) Packet {
|
||||
return Packet{
|
||||
Header: Header{
|
||||
Version: Version,
|
||||
Type: PacketTypeHandshake,
|
||||
Length: uint16(len(payload)),
|
||||
},
|
||||
Payload: payload,
|
||||
}
|
||||
}
|
||||
|
||||
func NewDataPacket(networkID uint32, payload []byte) Packet {
|
||||
return Packet{
|
||||
Header: Header{
|
||||
Version: Version,
|
||||
Type: PacketTypeData,
|
||||
NetworkID: networkID,
|
||||
Length: uint16(len(payload)),
|
||||
},
|
||||
Payload: payload,
|
||||
}
|
||||
}
|
||||
|
||||
func NewKeepalivePacket() Packet {
|
||||
return Packet{
|
||||
Header: Header{
|
||||
Version: Version,
|
||||
Type: PacketTypeKeepalive,
|
||||
Length: 0,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (p *Packet) Encode() []byte {
|
||||
total := HeaderSize + len(p.Payload)
|
||||
buf := make([]byte, total)
|
||||
p.Header.Length = uint16(len(p.Payload))
|
||||
p.Header.Encode(buf[:HeaderSize])
|
||||
copy(buf[HeaderSize:], p.Payload)
|
||||
return buf
|
||||
}
|
||||
|
||||
func DecodePacket(data []byte) (*Packet, error) {
|
||||
var p Packet
|
||||
if err := p.Header.Decode(data); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
payloadLen := int(p.Header.Length)
|
||||
if HeaderSize+payloadLen > len(data) {
|
||||
return nil, fmt.Errorf("packet truncated: header claims %d + %d > %d", HeaderSize, payloadLen, len(data))
|
||||
}
|
||||
p.Payload = make([]byte, payloadLen)
|
||||
copy(p.Payload, data[HeaderSize:HeaderSize+payloadLen])
|
||||
return &p, nil
|
||||
}
|
||||
|
||||
func DecodePacketInto(p *Packet, data []byte) error {
|
||||
if err := p.Header.Decode(data); err != nil {
|
||||
return err
|
||||
}
|
||||
payloadLen := int(p.Header.Length)
|
||||
if HeaderSize+payloadLen > len(data) {
|
||||
return fmt.Errorf("packet truncated")
|
||||
}
|
||||
p.Payload = make([]byte, payloadLen)
|
||||
copy(p.Payload, data[HeaderSize:HeaderSize+payloadLen])
|
||||
return nil
|
||||
}
|
||||
193
internal/vl1/peer.go
Normal file
193
internal/vl1/peer.go
Normal file
@@ -0,0 +1,193 @@
|
||||
package vl1
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"zeromesh/internal/identity"
|
||||
)
|
||||
|
||||
type Peer struct {
|
||||
Address identity.Address
|
||||
PublicKey [32]byte
|
||||
Endpoint *net.UDPAddr
|
||||
Cipher *NoiseCipher
|
||||
LastSeen time.Time
|
||||
LastSend time.Time
|
||||
Connected bool
|
||||
mu sync.RWMutex
|
||||
log *slog.Logger
|
||||
}
|
||||
|
||||
func (p *Peer) Touch() {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
p.LastSeen = time.Now()
|
||||
}
|
||||
|
||||
func (p *Peer) SetCipher(c *NoiseCipher) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
p.Cipher = c
|
||||
p.Connected = true
|
||||
}
|
||||
|
||||
func (p *Peer) IsConnected() bool {
|
||||
p.mu.RLock()
|
||||
defer p.mu.RUnlock()
|
||||
return p.Connected
|
||||
}
|
||||
|
||||
func (p *Peer) IsAlive() bool {
|
||||
p.mu.RLock()
|
||||
defer p.mu.RUnlock()
|
||||
return time.Since(p.LastSeen) < 90*time.Second
|
||||
}
|
||||
|
||||
func (p *Peer) Encrypt(plaintext []byte) ([]byte, error) {
|
||||
p.mu.RLock()
|
||||
defer p.mu.RUnlock()
|
||||
if p.Cipher == nil {
|
||||
return nil, ErrNoCipher
|
||||
}
|
||||
return p.Cipher.Encrypt(plaintext)
|
||||
}
|
||||
|
||||
func (p *Peer) Decrypt(ciphertext []byte) ([]byte, error) {
|
||||
p.mu.RLock()
|
||||
defer p.mu.RUnlock()
|
||||
if p.Cipher == nil {
|
||||
return nil, ErrNoCipher
|
||||
}
|
||||
return p.Cipher.Decrypt(ciphertext)
|
||||
}
|
||||
|
||||
func (p *Peer) EncryptTo(buf []byte, plaintext []byte) (int, error) {
|
||||
p.mu.RLock()
|
||||
defer p.mu.RUnlock()
|
||||
if p.Cipher == nil {
|
||||
return 0, ErrNoCipher
|
||||
}
|
||||
return p.Cipher.EncryptTo(buf, plaintext)
|
||||
}
|
||||
|
||||
func (p *Peer) DecryptTo(buf []byte, ciphertext []byte) ([]byte, error) {
|
||||
p.mu.RLock()
|
||||
defer p.mu.RUnlock()
|
||||
if p.Cipher == nil {
|
||||
return nil, ErrNoCipher
|
||||
}
|
||||
return p.Cipher.DecryptTo(buf, ciphertext)
|
||||
}
|
||||
|
||||
func (p *Peer) NeedsKeepalive() bool {
|
||||
p.mu.RLock()
|
||||
defer p.mu.RUnlock()
|
||||
return p.Connected && time.Since(p.LastSend) > 25*time.Second
|
||||
}
|
||||
|
||||
var ErrNoCipher = errNoCipher()
|
||||
|
||||
func errNoCipher() error {
|
||||
return &noCipherError{}
|
||||
}
|
||||
|
||||
type noCipherError struct{}
|
||||
|
||||
func (e *noCipherError) Error() string {
|
||||
return "no cipher established"
|
||||
}
|
||||
|
||||
type PeerManager struct {
|
||||
peers map[identity.Address]*Peer
|
||||
mu sync.RWMutex
|
||||
log *slog.Logger
|
||||
}
|
||||
|
||||
func NewPeerManager(log *slog.Logger) *PeerManager {
|
||||
return &PeerManager{
|
||||
peers: make(map[identity.Address]*Peer),
|
||||
log: log,
|
||||
}
|
||||
}
|
||||
|
||||
func (pm *PeerManager) AddPeer(addr identity.Address, pubKey [32]byte, endpoint *net.UDPAddr) *Peer {
|
||||
pm.mu.Lock()
|
||||
defer pm.mu.Unlock()
|
||||
peer := &Peer{
|
||||
Address: addr,
|
||||
PublicKey: pubKey,
|
||||
Endpoint: endpoint,
|
||||
LastSeen: time.Now(),
|
||||
LastSend: time.Now(),
|
||||
log: pm.log.With("peer", addr.String()),
|
||||
}
|
||||
pm.peers[addr] = peer
|
||||
return peer
|
||||
}
|
||||
|
||||
func (pm *PeerManager) GetPeer(addr identity.Address) *Peer {
|
||||
pm.mu.RLock()
|
||||
defer pm.mu.RUnlock()
|
||||
return pm.peers[addr]
|
||||
}
|
||||
|
||||
func (pm *PeerManager) GetPeerByEndpoint(endpoint *net.UDPAddr) *Peer {
|
||||
pm.mu.RLock()
|
||||
defer pm.mu.RUnlock()
|
||||
for _, p := range pm.peers {
|
||||
if p.Endpoint != nil && p.Endpoint.IP.Equal(endpoint.IP) && p.Endpoint.Port == endpoint.Port {
|
||||
return p
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (pm *PeerManager) RemovePeer(addr identity.Address) {
|
||||
pm.mu.Lock()
|
||||
defer pm.mu.Unlock()
|
||||
delete(pm.peers, addr)
|
||||
}
|
||||
|
||||
func (pm *PeerManager) UpdatePeerEndpoint(addr identity.Address, endpoint *net.UDPAddr) {
|
||||
pm.mu.Lock()
|
||||
defer pm.mu.Unlock()
|
||||
if p, ok := pm.peers[addr]; ok {
|
||||
p.Endpoint = endpoint
|
||||
}
|
||||
}
|
||||
|
||||
func (pm *PeerManager) ConnectedPeers() []*Peer {
|
||||
pm.mu.RLock()
|
||||
defer pm.mu.RUnlock()
|
||||
var result []*Peer
|
||||
for _, p := range pm.peers {
|
||||
if p.IsConnected() {
|
||||
result = append(result, p)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func (pm *PeerManager) AllPeers() []*Peer {
|
||||
pm.mu.RLock()
|
||||
defer pm.mu.RUnlock()
|
||||
result := make([]*Peer, 0, len(pm.peers))
|
||||
for _, p := range pm.peers {
|
||||
result = append(result, p)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func (pm *PeerManager) CleanDead() {
|
||||
pm.mu.Lock()
|
||||
defer pm.mu.Unlock()
|
||||
for addr, p := range pm.peers {
|
||||
if p.Connected && !p.IsAlive() {
|
||||
p.log.Warn("peer timed out, removing")
|
||||
delete(pm.peers, addr)
|
||||
}
|
||||
}
|
||||
}
|
||||
112
internal/vl1/transport.go
Normal file
112
internal/vl1/transport.go
Normal file
@@ -0,0 +1,112 @@
|
||||
package vl1
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Transport struct {
|
||||
conn *net.UDPConn
|
||||
port int
|
||||
mu sync.RWMutex
|
||||
closed bool
|
||||
log *slog.Logger
|
||||
}
|
||||
|
||||
func NewTransport(port int, log *slog.Logger) (*Transport, error) {
|
||||
addr := &net.UDPAddr{Port: port}
|
||||
conn, err := net.ListenUDP("udp", addr)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("bind UDP port %d: %w", port, err)
|
||||
}
|
||||
actualPort := conn.LocalAddr().(*net.UDPAddr).Port
|
||||
log.Info("VL1 transport listening", "port", actualPort)
|
||||
return &Transport{
|
||||
conn: conn,
|
||||
port: actualPort,
|
||||
log: log,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (t *Transport) Port() int {
|
||||
return t.port
|
||||
}
|
||||
|
||||
func (t *Transport) ReadFrom(buf []byte) (int, *net.UDPAddr, error) {
|
||||
n, addr, err := t.conn.ReadFromUDP(buf)
|
||||
return n, addr, err
|
||||
}
|
||||
|
||||
func (t *Transport) SendTo(data []byte, addr *net.UDPAddr) error {
|
||||
t.mu.RLock()
|
||||
defer t.mu.RUnlock()
|
||||
if t.closed {
|
||||
return fmt.Errorf("transport closed")
|
||||
}
|
||||
_, err := t.conn.WriteToUDP(data, addr)
|
||||
return err
|
||||
}
|
||||
|
||||
func (t *Transport) SendPacket(pkt *Packet, addr *net.UDPAddr) error {
|
||||
return t.SendTo(pkt.Encode(), addr)
|
||||
}
|
||||
|
||||
func (t *Transport) Close() error {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
t.closed = true
|
||||
return t.conn.Close()
|
||||
}
|
||||
|
||||
func (t *Transport) SetSocketBuffers(rcvBuf, sndBuf int) error {
|
||||
rawConn, err := t.conn.SyscallConn()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get raw conn: %w", err)
|
||||
}
|
||||
var setErr error
|
||||
err = rawConn.Control(func(fd uintptr) {
|
||||
if rcvBuf > 0 {
|
||||
if e := syscall.SetsockoptInt(int(fd), syscall.SOL_SOCKET, syscall.SO_RCVBUF, rcvBuf); e != nil {
|
||||
setErr = fmt.Errorf("set SO_RCVBUF=%d: %w", rcvBuf, e)
|
||||
return
|
||||
}
|
||||
}
|
||||
if sndBuf > 0 {
|
||||
if e := syscall.SetsockoptInt(int(fd), syscall.SOL_SOCKET, syscall.SO_SNDBUF, sndBuf); e != nil {
|
||||
setErr = fmt.Errorf("set SO_SNDBUF=%d: %w", sndBuf, e)
|
||||
return
|
||||
}
|
||||
}
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return setErr
|
||||
}
|
||||
|
||||
func (t *Transport) LocalAddr() net.Addr {
|
||||
return t.conn.LocalAddr()
|
||||
}
|
||||
|
||||
func (t *Transport) SetReadDeadline(deadline time.Time) error {
|
||||
return t.conn.SetReadDeadline(deadline)
|
||||
}
|
||||
|
||||
var packetBufPool = sync.Pool{
|
||||
New: func() interface{} {
|
||||
buf := make([]byte, MaxPacketSize)
|
||||
return &buf
|
||||
},
|
||||
}
|
||||
|
||||
func GetPacketBuf() *[]byte {
|
||||
return packetBufPool.Get().(*[]byte)
|
||||
}
|
||||
|
||||
func PutPacketBuf(buf *[]byte) {
|
||||
packetBufPool.Put(buf)
|
||||
}
|
||||
120
internal/vl2/arp.go
Normal file
120
internal/vl2/arp.go
Normal file
@@ -0,0 +1,120 @@
|
||||
package vl2
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"log/slog"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
type ARPEntry struct {
|
||||
MAC net.HardwareAddr
|
||||
LastSeen time.Time
|
||||
}
|
||||
|
||||
type ARPProxy struct {
|
||||
cache map[string]*ARPEntry
|
||||
mu sync.RWMutex
|
||||
log *slog.Logger
|
||||
}
|
||||
|
||||
func NewARPProxy(log *slog.Logger) *ARPProxy {
|
||||
return &ARPProxy{
|
||||
cache: make(map[string]*ARPEntry),
|
||||
log: log.With("component", "arp"),
|
||||
}
|
||||
}
|
||||
|
||||
func (a *ARPProxy) Learn(ip net.IP, mac net.HardwareAddr) {
|
||||
key := ip.String()
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
a.cache[key] = &ARPEntry{
|
||||
MAC: mac,
|
||||
LastSeen: time.Now(),
|
||||
}
|
||||
}
|
||||
|
||||
func (a *ARPProxy) Lookup(ip net.IP) net.HardwareAddr {
|
||||
key := ip.String()
|
||||
a.mu.RLock()
|
||||
defer a.mu.RUnlock()
|
||||
entry, ok := a.cache[key]
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
return entry.MAC
|
||||
}
|
||||
|
||||
func (a *ARPProxy) HandleARP(frame *EthernetFrame) []byte {
|
||||
if len(frame.Payload) < 28 {
|
||||
return nil
|
||||
}
|
||||
pl := frame.Payload
|
||||
|
||||
hrd := binary.BigEndian.Uint16(pl[0:2])
|
||||
pro := binary.BigEndian.Uint16(pl[2:4])
|
||||
hln := pl[4]
|
||||
pln := pl[5]
|
||||
op := binary.BigEndian.Uint16(pl[6:8])
|
||||
|
||||
if hrd != 1 || pro != EtherTypeIPv4 || hln != 6 || pln != 4 {
|
||||
return nil
|
||||
}
|
||||
|
||||
senderMAC := net.HardwareAddr(pl[8:14])
|
||||
senderIP := net.IP(pl[14:18])
|
||||
targetIP := net.IP(pl[24:28])
|
||||
|
||||
a.Learn(senderIP, senderMAC)
|
||||
|
||||
if op == 1 { // ARP request
|
||||
if localMAC := a.Lookup(targetIP); localMAC != nil {
|
||||
reply := make([]byte, 42)
|
||||
copy(reply[0:6], senderMAC)
|
||||
copy(reply[6:12], localMAC)
|
||||
binary.BigEndian.PutUint16(reply[12:14], EtherTypeARP)
|
||||
|
||||
reply[14] = 0x00
|
||||
reply[15] = 0x01
|
||||
binary.BigEndian.PutUint16(reply[16:18], EtherTypeIPv4)
|
||||
reply[18] = 6
|
||||
reply[19] = 4
|
||||
binary.BigEndian.PutUint16(reply[20:22], 2)
|
||||
|
||||
copy(reply[22:28], localMAC)
|
||||
copy(reply[28:32], targetIP)
|
||||
copy(reply[32:38], senderMAC)
|
||||
copy(reply[38:42], senderIP)
|
||||
|
||||
return reply
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *ARPProxy) PeerFromARP(frame *EthernetFrame) (net.IP, net.HardwareAddr) {
|
||||
if len(frame.Payload) < 28 {
|
||||
return nil, nil
|
||||
}
|
||||
pl := frame.Payload
|
||||
hrd := binary.BigEndian.Uint16(pl[0:2])
|
||||
pro := binary.BigEndian.Uint16(pl[2:4])
|
||||
if hrd != 1 || pro != EtherTypeIPv4 {
|
||||
return nil, nil
|
||||
}
|
||||
return net.IP(pl[14:18]), net.HardwareAddr(pl[8:14])
|
||||
}
|
||||
|
||||
func (a *ARPProxy) CleanExpired() {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
cutoff := time.Now().Add(-30 * time.Minute)
|
||||
for k, v := range a.cache {
|
||||
if v.LastSeen.Before(cutoff) {
|
||||
delete(a.cache, k)
|
||||
}
|
||||
}
|
||||
}
|
||||
92
internal/vl2/frame.go
Normal file
92
internal/vl2/frame.go
Normal file
@@ -0,0 +1,92 @@
|
||||
package vl2
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"net"
|
||||
"sync"
|
||||
)
|
||||
|
||||
const (
|
||||
MinFrameSize = 14
|
||||
MaxFrameSize = 65535
|
||||
EtherTypeIPv4 = 0x0800
|
||||
EtherTypeARP = 0x0806
|
||||
EtherTypeIPv6 = 0x86DD
|
||||
)
|
||||
|
||||
type EthernetFrame struct {
|
||||
DstMAC net.HardwareAddr
|
||||
SrcMAC net.HardwareAddr
|
||||
EtherType uint16
|
||||
Payload []byte
|
||||
}
|
||||
|
||||
func ParseEthernetFrame(data []byte) (*EthernetFrame, error) {
|
||||
if len(data) < MinFrameSize {
|
||||
return nil, fmt.Errorf("frame too short: %d < %d", len(data), MinFrameSize)
|
||||
}
|
||||
return &EthernetFrame{
|
||||
DstMAC: net.HardwareAddr(data[0:6]),
|
||||
SrcMAC: net.HardwareAddr(data[6:12]),
|
||||
EtherType: binary.BigEndian.Uint16(data[12:14]),
|
||||
Payload: data[14:],
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (f *EthernetFrame) Encode() []byte {
|
||||
buf := make([]byte, MinFrameSize+len(f.Payload))
|
||||
copy(buf[0:6], f.DstMAC)
|
||||
copy(buf[6:12], f.SrcMAC)
|
||||
binary.BigEndian.PutUint16(buf[12:14], f.EtherType)
|
||||
copy(buf[14:], f.Payload)
|
||||
return buf
|
||||
}
|
||||
|
||||
func (f *EthernetFrame) IsBroadcast() bool {
|
||||
for _, b := range f.DstMAC {
|
||||
if b != 0xff {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (f *EthernetFrame) IsMulticast() bool {
|
||||
return len(f.DstMAC) > 0 && f.DstMAC[0]&0x01 == 0x01 && !f.IsBroadcast()
|
||||
}
|
||||
|
||||
func (f *EthernetFrame) IsARP() bool {
|
||||
return f.EtherType == EtherTypeARP
|
||||
}
|
||||
|
||||
func (f *EthernetFrame) IsIPv4() bool {
|
||||
return f.EtherType == EtherTypeIPv4
|
||||
}
|
||||
|
||||
func (f *EthernetFrame) IsIPv6() bool {
|
||||
return f.EtherType == EtherTypeIPv6
|
||||
}
|
||||
|
||||
type MACKey [6]byte
|
||||
|
||||
func MACToKey(mac net.HardwareAddr) MACKey {
|
||||
var key MACKey
|
||||
copy(key[:], mac)
|
||||
return key
|
||||
}
|
||||
|
||||
var frameBufPool = sync.Pool{
|
||||
New: func() interface{} {
|
||||
buf := make([]byte, MaxFrameSize)
|
||||
return &buf
|
||||
},
|
||||
}
|
||||
|
||||
func GetFrameBuf() *[]byte {
|
||||
return frameBufPool.Get().(*[]byte)
|
||||
}
|
||||
|
||||
func PutFrameBuf(buf *[]byte) {
|
||||
frameBufPool.Put(buf)
|
||||
}
|
||||
19
internal/vl2/mac.go
Normal file
19
internal/vl2/mac.go
Normal file
@@ -0,0 +1,19 @@
|
||||
package vl2
|
||||
|
||||
import (
|
||||
"net"
|
||||
|
||||
"zeromesh/internal/identity"
|
||||
)
|
||||
|
||||
// GenerateMAC creates a deterministic MAC address from network ID and node address.
|
||||
func GenerateMAC(networkID uint32, nodeAddr identity.Address) net.HardwareAddr {
|
||||
mac := make(net.HardwareAddr, 6)
|
||||
mac[0] = 0x02 // locally administered, unicast
|
||||
mac[1] = byte(networkID >> 16)
|
||||
mac[2] = byte(networkID >> 8)
|
||||
mac[3] = byte(networkID)
|
||||
mac[4] = nodeAddr[3]
|
||||
mac[5] = nodeAddr[4]
|
||||
return mac
|
||||
}
|
||||
38
internal/vl2/network.go
Normal file
38
internal/vl2/network.go
Normal file
@@ -0,0 +1,38 @@
|
||||
package vl2
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
|
||||
"zeromesh/internal/identity"
|
||||
)
|
||||
|
||||
type NetworkConfig struct {
|
||||
ID uint32
|
||||
Name string
|
||||
IPRange string
|
||||
IP6Range string
|
||||
MTU int
|
||||
Multicast bool
|
||||
}
|
||||
|
||||
type Network struct {
|
||||
Config NetworkConfig
|
||||
Switch *Switch
|
||||
ARP *ARPProxy
|
||||
LocalMAC [6]byte
|
||||
log *slog.Logger
|
||||
}
|
||||
|
||||
func NewNetwork(config NetworkConfig, nodeAddr identity.Address, sender PeerSender, log *slog.Logger) *Network {
|
||||
netLog := log.With("network", config.ID, "name", config.Name)
|
||||
mac := GenerateMAC(config.ID, nodeAddr)
|
||||
var macArr [6]byte
|
||||
copy(macArr[:], mac)
|
||||
return &Network{
|
||||
Config: config,
|
||||
Switch: NewSwitch(config.ID, sender, netLog),
|
||||
ARP: NewARPProxy(netLog),
|
||||
LocalMAC: macArr,
|
||||
log: netLog,
|
||||
}
|
||||
}
|
||||
153
internal/vl2/switch.go
Normal file
153
internal/vl2/switch.go
Normal file
@@ -0,0 +1,153 @@
|
||||
package vl2
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"zeromesh/internal/identity"
|
||||
)
|
||||
|
||||
const (
|
||||
MACTableExpiry = 5 * time.Minute
|
||||
MACTableMaxSize = 4096
|
||||
)
|
||||
|
||||
type MACEntry struct {
|
||||
PeerAddr identity.Address
|
||||
LastSeen time.Time
|
||||
IsLocal bool
|
||||
}
|
||||
|
||||
type PeerSender interface {
|
||||
SendToPeer(peerAddr identity.Address, networkID uint32, frame []byte) error
|
||||
BroadcastToPeers(networkID uint32, frame []byte, excludePeer identity.Address) error
|
||||
}
|
||||
|
||||
type Switch struct {
|
||||
networkID uint32
|
||||
macTable map[MACKey]*MACEntry
|
||||
mu sync.RWMutex
|
||||
sender PeerSender
|
||||
log *slog.Logger
|
||||
}
|
||||
|
||||
func NewSwitch(networkID uint32, sender PeerSender, log *slog.Logger) *Switch {
|
||||
return &Switch{
|
||||
networkID: networkID,
|
||||
macTable: make(map[MACKey]*MACEntry),
|
||||
sender: sender,
|
||||
log: log.With("component", "switch", "network", networkID),
|
||||
}
|
||||
}
|
||||
|
||||
func (sw *Switch) HandleLocalFrame(frame []byte) error {
|
||||
parsed, err := ParseEthernetFrame(frame)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
sw.learn(parsed.SrcMAC, identity.Address{}, true)
|
||||
|
||||
if parsed.IsBroadcast() || parsed.IsMulticast() {
|
||||
return sw.sender.BroadcastToPeers(sw.networkID, frame, identity.Address{})
|
||||
}
|
||||
|
||||
sw.mu.RLock()
|
||||
entry, found := sw.macTable[MACToKey(parsed.DstMAC)]
|
||||
sw.mu.RUnlock()
|
||||
|
||||
if found && !entry.IsLocal {
|
||||
return sw.sender.SendToPeer(entry.PeerAddr, sw.networkID, frame)
|
||||
}
|
||||
|
||||
if !found {
|
||||
sw.log.Debug("unknown dst MAC, flooding", "dst", parsed.DstMAC)
|
||||
return sw.sender.BroadcastToPeers(sw.networkID, frame, identity.Address{})
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (sw *Switch) HandleRemoteFrame(peerAddr identity.Address, frame []byte) ([]byte, error) {
|
||||
parsed, err := ParseEthernetFrame(frame)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
sw.learn(parsed.SrcMAC, peerAddr, false)
|
||||
|
||||
if parsed.IsBroadcast() || parsed.IsMulticast() {
|
||||
_ = sw.sender.BroadcastToPeers(sw.networkID, frame, peerAddr)
|
||||
return frame, nil
|
||||
}
|
||||
|
||||
sw.mu.RLock()
|
||||
entry, found := sw.macTable[MACToKey(parsed.DstMAC)]
|
||||
sw.mu.RUnlock()
|
||||
|
||||
if found && entry.IsLocal {
|
||||
return frame, nil
|
||||
}
|
||||
|
||||
if found && !entry.IsLocal {
|
||||
_ = sw.sender.SendToPeer(entry.PeerAddr, sw.networkID, frame)
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
_ = sw.sender.BroadcastToPeers(sw.networkID, frame, peerAddr)
|
||||
return frame, nil
|
||||
}
|
||||
|
||||
func (sw *Switch) learn(mac net.HardwareAddr, peerAddr identity.Address, isLocal bool) {
|
||||
key := MACToKey(mac)
|
||||
sw.mu.Lock()
|
||||
defer sw.mu.Unlock()
|
||||
|
||||
if len(sw.macTable) >= MACTableMaxSize {
|
||||
sw.evictOldest()
|
||||
}
|
||||
|
||||
sw.macTable[key] = &MACEntry{
|
||||
PeerAddr: peerAddr,
|
||||
LastSeen: time.Now(),
|
||||
IsLocal: isLocal,
|
||||
}
|
||||
}
|
||||
|
||||
func (sw *Switch) evictOldest() {
|
||||
var oldestKey MACKey
|
||||
var oldestTime time.Time
|
||||
first := true
|
||||
for k, v := range sw.macTable {
|
||||
if first || v.LastSeen.Before(oldestTime) {
|
||||
oldestKey = k
|
||||
oldestTime = v.LastSeen
|
||||
first = false
|
||||
}
|
||||
}
|
||||
if !first {
|
||||
delete(sw.macTable, oldestKey)
|
||||
}
|
||||
}
|
||||
|
||||
func (sw *Switch) CleanExpired() int {
|
||||
sw.mu.Lock()
|
||||
defer sw.mu.Unlock()
|
||||
cutoff := time.Now().Add(-MACTableExpiry)
|
||||
removed := 0
|
||||
for k, v := range sw.macTable {
|
||||
if v.LastSeen.Before(cutoff) && !v.IsLocal {
|
||||
delete(sw.macTable, k)
|
||||
removed++
|
||||
}
|
||||
}
|
||||
return removed
|
||||
}
|
||||
|
||||
func (sw *Switch) MACTableSize() int {
|
||||
sw.mu.RLock()
|
||||
defer sw.mu.RUnlock()
|
||||
return len(sw.macTable)
|
||||
}
|
||||
Reference in New Issue
Block a user