From 0c0b58820c84280bb6f6d58163fa75c6b25b5292 Mon Sep 17 00:00:00 2001 From: GnomeZworc Date: Mon, 31 Aug 2026 18:35:43 +0200 Subject: [PATCH] f-46: dhcpapi: add the control socket, its protocol and its client #46 MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Le contrat et le listener vont dans internal/api/dhcp (package dhcpapi), sur la forme de internal/api/agent, et le client dans internal/client/dhcp. Chaînage des imports : statefile <- dhcpd <- dhcpapi <- dhcpclient, sans cycle. dhcpd parle net.IP et net.IPNet et garde ses structs disque privées ; dhcpapi parle chaînes JSON et convertit à la frontière. Un même type portait jusqu'ici le format du fil, la signature du Store et le format du .state — ce qui couplait le fichier au protocole alors que le ticket le décrit comme un détail interne. Le digest est calculé sur une forme canonique partagée par les deux côtés : MAC, IP et CIDR normalisés, hôtes triés, doublon de MAC refusé. Un écart de digest signale donc une vraie divergence, pas une différence d'écriture. Le listener pose un recover par connexion, plafonne les lignes à 64 Kio, refuse une ligne malformée sans fermer la connexion, écoute en 0600 et supprime une socket résiduelle avant le bind. Le client pose une deadline. 101 tests au total, -race propre, treize mutations toutes détectées. Signed-off-by: GnomeZworc --- internal/api/dhcp/convert.go | 83 ++++ internal/api/dhcp/digest.go | 116 ++++++ internal/api/dhcp/digest_test.go | 173 +++++++++ internal/api/dhcp/models.go | 58 +++ internal/api/dhcp/server.go | 190 +++++++++ internal/client/dhcp/client.go | 104 +++++ internal/client/dhcp/client_test.go | 364 ++++++++++++++++++ internal/dhcpd/engine.go | 16 +- internal/dhcpd/engine_test.go | 39 +- internal/dhcpd/reply_test.go | 23 +- internal/dhcpd/state.go | 285 -------------- internal/dhcpd/store.go | 255 ++++++++++++ .../dhcpd/{state_test.go => store_test.go} | 237 ++++++------ 13 files changed, 1490 insertions(+), 453 deletions(-) create mode 100644 internal/api/dhcp/convert.go create mode 100644 internal/api/dhcp/digest.go create mode 100644 internal/api/dhcp/digest_test.go create mode 100644 internal/api/dhcp/models.go create mode 100644 internal/api/dhcp/server.go create mode 100644 internal/client/dhcp/client.go create mode 100644 internal/client/dhcp/client_test.go delete mode 100644 internal/dhcpd/state.go create mode 100644 internal/dhcpd/store.go rename internal/dhcpd/{state_test.go => store_test.go} (55%) diff --git a/internal/api/dhcp/convert.go b/internal/api/dhcp/convert.go new file mode 100644 index 0000000..77dd7c4 --- /dev/null +++ b/internal/api/dhcp/convert.go @@ -0,0 +1,83 @@ +package dhcpapi + +import ( + "fmt" + "net" + + "git.g3e.fr/syonad/two/internal/dhcpd" +) + +func (s Subnet) toConfig() (dhcpd.SubnetConfig, error) { + _, network, err := net.ParseCIDR(s.Network) + if err != nil { + return dhcpd.SubnetConfig{}, fmt.Errorf("invalid network %q: %w", s.Network, err) + } + interfaceIP := net.ParseIP(s.InterfaceIP) + if interfaceIP == nil { + return dhcpd.SubnetConfig{}, fmt.Errorf("invalid interface ip %q", s.InterfaceIP) + } + + c := dhcpd.SubnetConfig{Network: network, InterfaceIP: interfaceIP} + + if s.VPCRoute != "" { + if _, c.VPCRoute, err = net.ParseCIDR(s.VPCRoute); err != nil { + return dhcpd.SubnetConfig{}, fmt.Errorf("invalid vpc route %q: %w", s.VPCRoute, err) + } + } + if s.DefaultGateway != "" { + if c.DefaultGateway = net.ParseIP(s.DefaultGateway); c.DefaultGateway == nil { + return dhcpd.SubnetConfig{}, fmt.Errorf("invalid default gateway %q", s.DefaultGateway) + } + } + return c, nil +} + +func (h Host) toHost() (dhcpd.Host, error) { + mac, err := net.ParseMAC(h.MAC) + if err != nil { + return dhcpd.Host{}, fmt.Errorf("invalid mac %q: %w", h.MAC, err) + } + ip := net.ParseIP(h.IP) + if ip == nil { + return dhcpd.Host{}, fmt.Errorf("invalid host ip %q", h.IP) + } + return dhcpd.Host{MAC: mac, IP: ip, VM: h.VM, DefaultRoute: h.DefaultRoute}, nil +} + +func subnetFromConfig(c dhcpd.SubnetConfig) Subnet { + s := Subnet{ + Network: c.Network.String(), + InterfaceIP: c.InterfaceIP.String(), + } + if c.VPCRoute != nil { + s.VPCRoute = c.VPCRoute.String() + } + if c.DefaultGateway != nil { + s.DefaultGateway = c.DefaultGateway.String() + } + return s +} + +func hostFromHost(h dhcpd.Host) Host { + return Host{ + MAC: h.MAC.String(), + IP: h.IP.String(), + VM: h.VM, + DefaultRoute: h.DefaultRoute, + } +} + +func stateFromStore(store *dhcpd.Store) State { + state := State{Hosts: make([]Host, 0)} + + if config, configured := store.Subnet(); configured { + subnet := subnetFromConfig(config) + state.Subnet = &subnet + } + for _, h := range store.Hosts() { + state.Hosts = append(state.Hosts, hostFromHost(h)) + } + SortHosts(state.Hosts) + + return state +} diff --git a/internal/api/dhcp/digest.go b/internal/api/dhcp/digest.go new file mode 100644 index 0000000..f311b22 --- /dev/null +++ b/internal/api/dhcp/digest.go @@ -0,0 +1,116 @@ +package dhcpapi + +import ( + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "net" + "sort" +) + +func canonicalMAC(s string) (string, error) { + mac, err := net.ParseMAC(s) + if err != nil { + return "", fmt.Errorf("invalid mac %q: %w", s, err) + } + return mac.String(), nil +} + +func canonicalIP(s string) (string, error) { + ip := net.ParseIP(s) + if ip == nil { + return "", fmt.Errorf("invalid ip %q", s) + } + return ip.String(), nil +} + +func canonicalCIDR(s string) (string, error) { + _, network, err := net.ParseCIDR(s) + if err != nil { + return "", fmt.Errorf("invalid cidr %q: %w", s, err) + } + return network.String(), nil +} + +func SortHosts(hosts []Host) { + sort.Slice(hosts, func(i, j int) bool { return hosts[i].MAC < hosts[j].MAC }) +} + +func CanonicalSubnet(s Subnet) (Subnet, error) { + network, err := canonicalCIDR(s.Network) + if err != nil { + return Subnet{}, err + } + interfaceIP, err := canonicalIP(s.InterfaceIP) + if err != nil { + return Subnet{}, err + } + + out := Subnet{Network: network, InterfaceIP: interfaceIP} + + if s.VPCRoute != "" { + if out.VPCRoute, err = canonicalCIDR(s.VPCRoute); err != nil { + return Subnet{}, err + } + } + if s.DefaultGateway != "" { + if out.DefaultGateway, err = canonicalIP(s.DefaultGateway); err != nil { + return Subnet{}, err + } + } + return out, nil +} + +func CanonicalHost(h Host) (Host, error) { + mac, err := canonicalMAC(h.MAC) + if err != nil { + return Host{}, err + } + ip, err := canonicalIP(h.IP) + if err != nil { + return Host{}, err + } + return Host{MAC: mac, IP: ip, VM: h.VM, DefaultRoute: h.DefaultRoute}, nil +} + +func Canonical(s State) (State, error) { + out := State{Hosts: make([]Host, 0, len(s.Hosts))} + + if s.Subnet != nil { + subnet, err := CanonicalSubnet(*s.Subnet) + if err != nil { + return State{}, err + } + out.Subnet = &subnet + } + + seen := make(map[string]struct{}, len(s.Hosts)) + for _, h := range s.Hosts { + host, err := CanonicalHost(h) + if err != nil { + return State{}, err + } + if _, dup := seen[host.MAC]; dup { + return State{}, fmt.Errorf("duplicate mac %s", host.MAC) + } + seen[host.MAC] = struct{}{} + out.Hosts = append(out.Hosts, host) + } + SortHosts(out.Hosts) + + return out, nil +} + +func Digest(s State) (string, error) { + canonical, err := Canonical(s) + if err != nil { + return "", err + } + raw, err := json.Marshal(canonical) + if err != nil { + return "", fmt.Errorf("encode state: %w", err) + } + sum := sha256.Sum256(raw) + return hex.EncodeToString(sum[:]), nil +} diff --git a/internal/api/dhcp/digest_test.go b/internal/api/dhcp/digest_test.go new file mode 100644 index 0000000..fe5dd29 --- /dev/null +++ b/internal/api/dhcp/digest_test.go @@ -0,0 +1,173 @@ +package dhcpapi + +import ( + "testing" +) + +func subnet() Subnet { + return Subnet{ + Network: "10.0.5.0/24", + InterfaceIP: "10.0.5.1", + VPCRoute: "10.0.0.0/16", + DefaultGateway: "10.0.5.254", + } +} + +func hosts() []Host { + return []Host{ + {MAC: "00:22:33:00:00:0a", IP: "10.0.5.10", VM: "vm-a", DefaultRoute: true}, + {MAC: "00:22:33:00:00:0b", IP: "10.0.5.11", VM: "vm-b"}, + } +} + +func digestOf(t *testing.T, s State) string { + t.Helper() + d, err := Digest(s) + if err != nil { + t.Fatalf("Digest: %v", err) + } + return d +} + +func TestDigest_IsStableAcrossHostOrder(t *testing.T) { + sub := subnet() + a := State{Subnet: &sub, Hosts: hosts()} + + reversed := hosts() + reversed[0], reversed[1] = reversed[1], reversed[0] + b := State{Subnet: &sub, Hosts: reversed} + + if digestOf(t, a) != digestOf(t, b) { + t.Error("host order must not change the digest: the watchdog would report a phantom drift") + } +} + +func TestDigest_IsStableAcrossMACCase(t *testing.T) { + sub := subnet() + a := State{Subnet: &sub, Hosts: hosts()} + + upper := hosts() + upper[0].MAC = "00:22:33:00:00:0A" + b := State{Subnet: &sub, Hosts: upper} + + if digestOf(t, a) != digestOf(t, b) { + t.Error("mac case must not change the digest") + } +} + +func TestDigest_IsStableAcrossIPv4InIPv6Notation(t *testing.T) { + sub := subnet() + a := State{Subnet: &sub, Hosts: hosts()} + + mapped := subnet() + mapped.InterfaceIP = "::ffff:10.0.5.1" + b := State{Subnet: &mapped, Hosts: hosts()} + + if digestOf(t, a) != digestOf(t, b) { + t.Error("the same address written in ipv4-mapped form must hash alike") + } +} + +func TestDigest_ChangesWhenAHostIPChanges(t *testing.T) { + sub := subnet() + a := State{Subnet: &sub, Hosts: hosts()} + + moved := hosts() + moved[0].IP = "10.0.5.99" + b := State{Subnet: &sub, Hosts: moved} + + if digestOf(t, a) == digestOf(t, b) { + t.Error("a changed reservation must change the digest") + } +} + +func TestDigest_ChangesWhenTheDefaultRouteFlagChanges(t *testing.T) { + sub := subnet() + a := State{Subnet: &sub, Hosts: hosts()} + + flipped := hosts() + flipped[0].DefaultRoute = false + b := State{Subnet: &sub, Hosts: flipped} + + if digestOf(t, a) == digestOf(t, b) { + t.Error("the default route flag is part of the served state") + } +} + +func TestDigest_ChangesWhenTheSubnetChanges(t *testing.T) { + sub := subnet() + a := State{Subnet: &sub, Hosts: hosts()} + + other := subnet() + other.DefaultGateway = "10.0.5.253" + b := State{Subnet: &other, Hosts: hosts()} + + if digestOf(t, a) == digestOf(t, b) { + t.Error("the subnet configuration is part of the served state") + } +} + +func TestDigest_DistinguishesNoSubnetFromAConfiguredOne(t *testing.T) { + sub := subnet() + configured := State{Subnet: &sub} + bare := State{} + + if digestOf(t, configured) == digestOf(t, bare) { + t.Error("an unconfigured subnet must not hash like a configured one") + } +} + +func TestDigest_EmptyAndNilHostsHashAlike(t *testing.T) { + if digestOf(t, State{Hosts: nil}) != digestOf(t, State{Hosts: []Host{}}) { + t.Error("nil and empty host lists describe the same state") + } +} + +func TestDigest_RejectsAnInvalidMAC(t *testing.T) { + if _, err := Digest(State{Hosts: []Host{{MAC: "nope", IP: "10.0.5.10"}}}); err == nil { + t.Fatal("an invalid mac must be reported, not hashed") + } +} + +func TestDigest_RejectsAnInvalidIP(t *testing.T) { + if _, err := Digest(State{Hosts: []Host{{MAC: "00:22:33:00:00:0a", IP: "10.0.5.300"}}}); err == nil { + t.Fatal("an invalid ip must be reported, not hashed") + } +} + +func TestCanonical_RejectsADuplicateMAC(t *testing.T) { + dup := []Host{ + {MAC: "00:22:33:00:00:0a", IP: "10.0.5.10"}, + {MAC: "00:22:33:00:00:0A", IP: "10.0.5.11"}, + } + if _, err := Canonical(State{Hosts: dup}); err == nil { + t.Fatal("the same mac twice is an inconsistent state, not something to hash") + } +} + +func TestCanonical_NormalizesTheNetworkToItsBaseAddress(t *testing.T) { + sub := subnet() + sub.Network = "10.0.5.42/24" + + got, err := CanonicalSubnet(sub) + if err != nil { + t.Fatalf("CanonicalSubnet: %v", err) + } + if got.Network != "10.0.5.0/24" { + t.Errorf("network = %s, want 10.0.5.0/24", got.Network) + } +} + +func TestCanonical_SortsHostsByMAC(t *testing.T) { + unsorted := []Host{ + {MAC: "00:22:33:00:00:0c", IP: "10.0.5.12"}, + {MAC: "00:22:33:00:00:0a", IP: "10.0.5.10"}, + } + got, err := Canonical(State{Hosts: unsorted}) + if err != nil { + t.Fatalf("Canonical: %v", err) + } + if got.Hosts[0].MAC != "00:22:33:00:00:0a" { + t.Errorf("hosts = %v, want sorted by mac", got.Hosts) + } +} diff --git a/internal/api/dhcp/models.go b/internal/api/dhcp/models.go new file mode 100644 index 0000000..b0268cf --- /dev/null +++ b/internal/api/dhcp/models.go @@ -0,0 +1,58 @@ +package dhcpapi + +type Verb string + +const ( + VerbSetSubnet Verb = "set-subnet" + VerbSetHost Verb = "set-host" + VerbDelHost Verb = "del-host" + VerbGetState Verb = "get-state" + VerbProbe Verb = "probe" +) + +const MaxMessageBytes = 64 * 1024 + +type Subnet struct { + Network string `json:"network"` + InterfaceIP string `json:"interface_ip"` + VPCRoute string `json:"vpc_route,omitempty"` + DefaultGateway string `json:"default_gateway,omitempty"` +} + +type Host struct { + MAC string `json:"mac"` + IP string `json:"ip"` + VM string `json:"vm,omitempty"` + DefaultRoute bool `json:"default_route"` +} + +type State struct { + Subnet *Subnet `json:"subnet,omitempty"` + Hosts []Host `json:"hosts"` +} + +type Lease struct { + MAC string `json:"mac"` + IP string `json:"ip"` + Netmask string `json:"netmask"` + Router string `json:"router,omitempty"` + DNS []string `json:"dns"` + Routes []string `json:"routes"` + LeaseSeconds uint32 `json:"lease_seconds"` +} + +type Request struct { + Verb Verb `json:"verb"` + Subnet *Subnet `json:"subnet,omitempty"` + Host *Host `json:"host,omitempty"` + MAC string `json:"mac,omitempty"` +} + +type Response struct { + OK bool `json:"ok"` + Error string `json:"error,omitempty"` + State *State `json:"state,omitempty"` + Digest string `json:"digest,omitempty"` + Lease *Lease `json:"lease,omitempty"` + Served bool `json:"served,omitempty"` +} diff --git a/internal/api/dhcp/server.go b/internal/api/dhcp/server.go new file mode 100644 index 0000000..879ac58 --- /dev/null +++ b/internal/api/dhcp/server.go @@ -0,0 +1,190 @@ +package dhcpapi + +import ( + "bufio" + "encoding/json" + "errors" + "fmt" + "log/slog" + "net" + "os" + "path/filepath" + + "git.g3e.fr/syonad/two/internal/dhcpd" + "git.g3e.fr/syonad/two/pkg/db/statefile" + + "github.com/insomniacslk/dhcp/dhcpv4" +) + +const SocketMode = 0o600 + +type Server struct { + store *dhcpd.Store + listener net.Listener + logger *slog.Logger +} + +func Listen(store *dhcpd.Store, path string, logger *slog.Logger) (*Server, error) { + dir := filepath.Dir(path) + if err := os.MkdirAll(dir, statefile.DirMode); err != nil { + return nil, fmt.Errorf("create %s: %w", dir, err) + } + if err := os.Remove(path); err != nil && !errors.Is(err, os.ErrNotExist) { + return nil, fmt.Errorf("remove stale socket %s: %w", path, err) + } + + listener, err := net.Listen("unix", path) + if err != nil { + return nil, fmt.Errorf("listen on %s: %w", path, err) + } + if err := os.Chmod(path, SocketMode); err != nil { + listener.Close() + return nil, fmt.Errorf("chmod %s: %w", path, err) + } + + return &Server{store: store, listener: listener, logger: logger}, nil +} + +func (s *Server) Addr() string { + return s.listener.Addr().String() +} + +func (s *Server) Close() error { + return s.listener.Close() +} + +func (s *Server) Serve() error { + for { + conn, err := s.listener.Accept() + if err != nil { + return err + } + go s.handleConn(conn) + } +} + +func (s *Server) handleConn(conn net.Conn) { + defer conn.Close() + defer func() { + if r := recover(); r != nil { + s.logger.Error("control connection panicked", "panic", r) + } + }() + + scanner := bufio.NewScanner(conn) + scanner.Buffer(make([]byte, 0, 4096), MaxMessageBytes) + encoder := json.NewEncoder(conn) + + for scanner.Scan() { + line := scanner.Bytes() + if len(line) == 0 { + continue + } + + var req Request + if err := json.Unmarshal(line, &req); err != nil { + if err := encoder.Encode(failure(fmt.Errorf("malformed request: %w", err))); err != nil { + return + } + continue + } + + if err := encoder.Encode(s.dispatch(req)); err != nil { + return + } + } + if err := scanner.Err(); err != nil { + s.logger.Error("control connection read failed", "error", err) + } +} + +func failure(err error) Response { + return Response{OK: false, Error: err.Error()} +} + +func (s *Server) dispatch(req Request) Response { + switch req.Verb { + case VerbSetSubnet: + if req.Subnet == nil { + return failure(errors.New("set-subnet requires a subnet")) + } + config, err := req.Subnet.toConfig() + if err != nil { + return failure(err) + } + if err := s.store.SetSubnet(config); err != nil { + return failure(err) + } + return Response{OK: true} + + case VerbSetHost: + if req.Host == nil { + return failure(errors.New("set-host requires a host")) + } + host, err := req.Host.toHost() + if err != nil { + return failure(err) + } + if err := s.store.SetHost(host); err != nil { + return failure(err) + } + return Response{OK: true} + + case VerbDelHost: + mac, err := net.ParseMAC(req.MAC) + if err != nil { + return failure(fmt.Errorf("invalid mac %q: %w", req.MAC, err)) + } + if err := s.store.DelHost(mac); err != nil { + return failure(err) + } + return Response{OK: true} + + case VerbGetState: + state := stateFromStore(s.store) + digest, err := Digest(state) + if err != nil { + return failure(err) + } + return Response{OK: true, State: &state, Digest: digest} + + case VerbProbe: + mac, err := net.ParseMAC(req.MAC) + if err != nil { + return failure(fmt.Errorf("invalid mac %q: %w", req.MAC, err)) + } + reply, err := s.store.Probe(mac) + if err != nil { + return failure(err) + } + if reply == nil { + return Response{OK: true, Served: false} + } + return Response{OK: true, Served: true, Lease: leaseFromReply(mac, reply)} + + default: + return failure(fmt.Errorf("unknown verb %q", req.Verb)) + } +} + +func leaseFromReply(mac net.HardwareAddr, reply *dhcpv4.DHCPv4) *Lease { + lease := &Lease{ + MAC: mac.String(), + IP: reply.YourIPAddr.String(), + Netmask: net.IP(reply.SubnetMask()).String(), + DNS: make([]string, 0, 2), + Routes: make([]string, 0, 3), + LeaseSeconds: uint32(reply.IPAddressLeaseTime(0).Seconds()), + } + + if routers := reply.Router(); len(routers) > 0 { + lease.Router = routers[0].String() + } + for _, dns := range reply.DNS() { + lease.DNS = append(lease.DNS, dns.String()) + } + for _, route := range reply.ClasslessStaticRoute() { + lease.Routes = append(lease.Routes, route.Dest.String()+" via "+route.Router.String()) + } + return lease +} diff --git a/internal/client/dhcp/client.go b/internal/client/dhcp/client.go new file mode 100644 index 0000000..d9133ae --- /dev/null +++ b/internal/client/dhcp/client.go @@ -0,0 +1,104 @@ +package dhcpclient + +import ( + "bufio" + "encoding/json" + "errors" + "fmt" + "net" + "time" + + dhcpapi "git.g3e.fr/syonad/two/internal/api/dhcp" +) + +const DefaultTimeout = 5 * time.Second + +var ErrNotServed = errors.New("mac is not served by this subnet") + +type Client struct { + path string + timeout time.Duration +} + +func New(path string) *Client { + return &Client{path: path, timeout: DefaultTimeout} +} + +func (c *Client) WithTimeout(d time.Duration) *Client { + return &Client{path: c.path, timeout: d} +} + +func (c *Client) call(req dhcpapi.Request) (dhcpapi.Response, error) { + conn, err := net.DialTimeout("unix", c.path, c.timeout) + if err != nil { + return dhcpapi.Response{}, fmt.Errorf("dial %s: %w", c.path, err) + } + defer conn.Close() + + if err := conn.SetDeadline(time.Now().Add(c.timeout)); err != nil { + return dhcpapi.Response{}, fmt.Errorf("set deadline on %s: %w", c.path, err) + } + + raw, err := json.Marshal(req) + if err != nil { + return dhcpapi.Response{}, fmt.Errorf("encode %s: %w", req.Verb, err) + } + if _, err := conn.Write(append(raw, '\n')); err != nil { + return dhcpapi.Response{}, fmt.Errorf("send %s: %w", req.Verb, err) + } + + scanner := bufio.NewScanner(conn) + scanner.Buffer(make([]byte, 0, 4096), dhcpapi.MaxMessageBytes) + if !scanner.Scan() { + if err := scanner.Err(); err != nil { + return dhcpapi.Response{}, fmt.Errorf("read reply to %s: %w", req.Verb, err) + } + return dhcpapi.Response{}, fmt.Errorf("no reply to %s", req.Verb) + } + + var resp dhcpapi.Response + if err := json.Unmarshal(scanner.Bytes(), &resp); err != nil { + return dhcpapi.Response{}, fmt.Errorf("parse reply to %s: %w", req.Verb, err) + } + if !resp.OK { + return resp, fmt.Errorf("%s refused: %s", req.Verb, resp.Error) + } + return resp, nil +} + +func (c *Client) SetSubnet(subnet dhcpapi.Subnet) error { + _, err := c.call(dhcpapi.Request{Verb: dhcpapi.VerbSetSubnet, Subnet: &subnet}) + return err +} + +func (c *Client) SetHost(host dhcpapi.Host) error { + _, err := c.call(dhcpapi.Request{Verb: dhcpapi.VerbSetHost, Host: &host}) + return err +} + +func (c *Client) DelHost(mac string) error { + _, err := c.call(dhcpapi.Request{Verb: dhcpapi.VerbDelHost, MAC: mac}) + return err +} + +func (c *Client) GetState() (dhcpapi.State, string, error) { + resp, err := c.call(dhcpapi.Request{Verb: dhcpapi.VerbGetState}) + if err != nil { + return dhcpapi.State{}, "", err + } + if resp.State == nil { + return dhcpapi.State{}, "", errors.New("get-state returned no state") + } + return *resp.State, resp.Digest, nil +} + +func (c *Client) Probe(mac string) (dhcpapi.Lease, error) { + resp, err := c.call(dhcpapi.Request{Verb: dhcpapi.VerbProbe, MAC: mac}) + if err != nil { + return dhcpapi.Lease{}, err + } + if !resp.Served || resp.Lease == nil { + return dhcpapi.Lease{}, ErrNotServed + } + return *resp.Lease, nil +} diff --git a/internal/client/dhcp/client_test.go b/internal/client/dhcp/client_test.go new file mode 100644 index 0000000..0bf3866 --- /dev/null +++ b/internal/client/dhcp/client_test.go @@ -0,0 +1,364 @@ +package dhcpclient + +import ( + "errors" + "io" + "log/slog" + "net" + "os" + "path/filepath" + "strings" + "testing" + "time" + + dhcpapi "git.g3e.fr/syonad/two/internal/api/dhcp" + "git.g3e.fr/syonad/two/internal/dhcpd" +) + +func shortTempDir(t *testing.T) string { + t.Helper() + dir, err := os.MkdirTemp("", "dhcpd") + if err != nil { + t.Fatalf("MkdirTemp: %v", err) + } + t.Cleanup(func() { os.RemoveAll(dir) }) + return dir +} + +func testSubnet() dhcpapi.Subnet { + return dhcpapi.Subnet{ + Network: "10.0.5.0/24", + InterfaceIP: "10.0.5.1", + VPCRoute: "10.0.0.0/16", + DefaultGateway: "10.0.5.254", + } +} + +func testHost() dhcpapi.Host { + return dhcpapi.Host{MAC: "00:22:33:00:00:0a", IP: "10.0.5.10", VM: "vm-test", DefaultRoute: true} +} + +func discardLogger() *slog.Logger { + return slog.New(slog.NewTextHandler(io.Discard, nil)) +} + +func serve(t *testing.T) (*Client, *dhcpd.Store, string) { + t.Helper() + + dir := shortTempDir(t) + store := dhcpd.NewStore(filepath.Join(dir, "s.state")) + if err := store.Load(); err != nil { + t.Fatalf("Load: %v", err) + } + + socketPath := filepath.Join(dir, "s.sock") + server, err := dhcpapi.Listen(store, socketPath, discardLogger()) + if err != nil { + t.Fatalf("Listen: %v", err) + } + go server.Serve() + t.Cleanup(func() { server.Close() }) + + return New(socketPath), store, socketPath +} + +func TestListen_SocketIsOwnerOnly(t *testing.T) { + _, _, socketPath := serve(t) + + info, err := os.Stat(socketPath) + if err != nil { + t.Fatalf("Stat: %v", err) + } + if got := info.Mode().Perm(); got != 0o600 { + t.Errorf("mode = %o, want 600: whoever reaches it rewrites the subnet addressing", got) + } +} + +func TestListen_ReplacesAStaleSocketFile(t *testing.T) { + dir := shortTempDir(t) + socketPath := filepath.Join(dir, "stale.sock") + if err := os.WriteFile(socketPath, []byte("leftover"), 0o600); err != nil { + t.Fatalf("WriteFile: %v", err) + } + + store := dhcpd.NewStore(filepath.Join(dir, "s.state")) + if err := store.Load(); err != nil { + t.Fatalf("Load: %v", err) + } + + server, err := dhcpapi.Listen(store, socketPath, discardLogger()) + if err != nil { + t.Fatalf("a socket left by an unclean stop must not block startup: %v", err) + } + server.Close() +} + +func TestSetSubnet_ReachesTheStore(t *testing.T) { + client, store, _ := serve(t) + + if err := client.SetSubnet(testSubnet()); err != nil { + t.Fatalf("SetSubnet: %v", err) + } + if _, configured := store.Subnet(); !configured { + t.Error("the subnet configuration did not reach the store") + } +} + +func TestSetSubnet_InvalidNetworkIsRefused(t *testing.T) { + client, _, _ := serve(t) + + subnet := testSubnet() + subnet.Network = "10.0.5.0" + err := client.SetSubnet(subnet) + if err == nil { + t.Fatal("an invalid network must be refused") + } + if !strings.Contains(err.Error(), "refused") { + t.Errorf("error = %v, want the server refusal to surface", err) + } +} + +func TestSetHost_IsIdempotent(t *testing.T) { + client, store, _ := serve(t) + + for range 3 { + if err := client.SetHost(testHost()); err != nil { + t.Fatalf("SetHost: %v", err) + } + } + if got := len(store.Hosts()); got != 1 { + t.Errorf("hosts = %d, want 1: set-host replaces the entry for that mac", got) + } +} + +func TestSetHost_InvalidMACIsRefused(t *testing.T) { + client, _, _ := serve(t) + + host := testHost() + host.MAC = "nope" + if err := client.SetHost(host); err == nil { + t.Fatal("an invalid mac must be refused") + } +} + +func TestDelHost_RemovesTheEntry(t *testing.T) { + client, store, _ := serve(t) + + if err := client.SetHost(testHost()); err != nil { + t.Fatalf("SetHost: %v", err) + } + if err := client.DelHost(testHost().MAC); err != nil { + t.Fatalf("DelHost: %v", err) + } + if got := len(store.Hosts()); got != 0 { + t.Errorf("hosts = %d, want 0", got) + } +} + +func TestDelHost_UnknownMACIsNotAnError(t *testing.T) { + client, _, _ := serve(t) + + if err := client.DelHost("00:22:33:ff:ff:ff"); err != nil { + t.Errorf("deleting an absent entry must be idempotent, got %v", err) + } +} + +func TestDelHost_InvalidMACIsRefused(t *testing.T) { + client, _, _ := serve(t) + + if err := client.DelHost("not-a-mac"); err == nil { + t.Fatal("an invalid mac must be refused") + } +} + +func TestGetState_ReturnsStateAndDigest(t *testing.T) { + client, _, _ := serve(t) + + if err := client.SetSubnet(testSubnet()); err != nil { + t.Fatalf("SetSubnet: %v", err) + } + if err := client.SetHost(testHost()); err != nil { + t.Fatalf("SetHost: %v", err) + } + + state, digest, err := client.GetState() + if err != nil { + t.Fatalf("GetState: %v", err) + } + if state.Subnet == nil || len(state.Hosts) != 1 { + t.Fatalf("state = %+v, want one subnet and one host", state) + } + if digest == "" { + t.Fatal("the digest is what the watchdog compares") + } + + local, err := dhcpapi.Digest(state) + if err != nil { + t.Fatalf("Digest: %v", err) + } + if local != digest { + t.Errorf("digest recomputed locally = %s, server said %s: the canonical form diverges", local, digest) + } +} + +func TestGetState_DigestFollowsTheState(t *testing.T) { + client, _, _ := serve(t) + + if err := client.SetSubnet(testSubnet()); err != nil { + t.Fatalf("SetSubnet: %v", err) + } + _, before, err := client.GetState() + if err != nil { + t.Fatalf("GetState: %v", err) + } + + if err := client.SetHost(testHost()); err != nil { + t.Fatalf("SetHost: %v", err) + } + _, after, err := client.GetState() + if err != nil { + t.Fatalf("GetState: %v", err) + } + + if before == after { + t.Error("adding a reservation must change the digest") + } +} + +func TestProbe_DescribesWhatWouldBeSent(t *testing.T) { + client, _, _ := serve(t) + + if err := client.SetSubnet(testSubnet()); err != nil { + t.Fatalf("SetSubnet: %v", err) + } + if err := client.SetHost(testHost()); err != nil { + t.Fatalf("SetHost: %v", err) + } + + lease, err := client.Probe("00:22:33:00:00:0A") + if err != nil { + t.Fatalf("Probe: %v", err) + } + if lease.IP != "10.0.5.10" { + t.Errorf("ip = %s, want 10.0.5.10", lease.IP) + } + if lease.Netmask != "255.255.255.0" { + t.Errorf("netmask = %s, want 255.255.255.0", lease.Netmask) + } + if lease.Router != "10.0.5.254" { + t.Errorf("router = %s, want 10.0.5.254", lease.Router) + } + if len(lease.DNS) != 2 { + t.Errorf("dns = %v, want two servers", lease.DNS) + } + if len(lease.Routes) != 3 { + t.Errorf("routes = %v, want metadata, vpc and default", lease.Routes) + } + if lease.LeaseSeconds != 43200 { + t.Errorf("lease = %ds, want 43200", lease.LeaseSeconds) + } + if lease.MAC != "00:22:33:00:00:0a" { + t.Errorf("mac = %s, want the normalized form", lease.MAC) + } +} + +func TestProbe_UnservedMACIsReported(t *testing.T) { + client, _, _ := serve(t) + + if err := client.SetSubnet(testSubnet()); err != nil { + t.Fatalf("SetSubnet: %v", err) + } + if _, err := client.Probe("00:22:33:ff:ff:ff"); !errors.Is(err, ErrNotServed) { + t.Fatalf("error = %v, want ErrNotServed", err) + } +} + +func TestProbe_WithoutSubnetConfigurationIsRefused(t *testing.T) { + client, _, _ := serve(t) + + if _, err := client.Probe("00:22:33:00:00:0a"); err == nil { + t.Fatal("probing an unconfigured subnet must be refused") + } +} + +func TestCall_UnknownVerbIsRefused(t *testing.T) { + _, _, socketPath := serve(t) + + conn, err := net.Dial("unix", socketPath) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer conn.Close() + + if _, err := conn.Write([]byte(`{"verb":"drop-everything"}` + "\n")); err != nil { + t.Fatalf("Write: %v", err) + } + + buf := make([]byte, 512) + n, err := conn.Read(buf) + if err != nil { + t.Fatalf("Read: %v", err) + } + if !strings.Contains(string(buf[:n]), "unknown verb") { + t.Errorf("reply = %s, want an unknown verb refusal", buf[:n]) + } +} + +func TestCall_MalformedLineIsRefusedWithoutClosingTheConnection(t *testing.T) { + _, _, socketPath := serve(t) + + conn, err := net.Dial("unix", socketPath) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer conn.Close() + + if _, err := conn.Write([]byte("{not json\n" + `{"verb":"get-state"}` + "\n")); err != nil { + t.Fatalf("Write: %v", err) + } + + buf := make([]byte, 4096) + n, err := conn.Read(buf) + if err != nil { + t.Fatalf("Read: %v", err) + } + if !strings.Contains(string(buf[:n]), "malformed request") { + t.Errorf("first reply = %s, want a malformed request refusal", buf[:n]) + } +} + +func TestCall_OnAnAbsentSocketFails(t *testing.T) { + client := New(filepath.Join(shortTempDir(t), "nothing.sock")) + + if err := client.DelHost("00:22:33:00:00:0a"); err == nil { + t.Fatal("an absent socket must be reported") + } +} + +func TestCall_HonoursItsTimeout(t *testing.T) { + socketPath := filepath.Join(shortTempDir(t), "mute.sock") + + listener, err := net.Listen("unix", socketPath) + if err != nil { + t.Fatalf("Listen: %v", err) + } + defer listener.Close() + + go func() { + conn, err := listener.Accept() + if err != nil { + return + } + defer conn.Close() + time.Sleep(3 * time.Second) + }() + + client := New(socketPath).WithTimeout(150 * time.Millisecond) + start := time.Now() + if err := client.DelHost("00:22:33:00:00:0a"); err == nil { + t.Fatal("a mute server must not hang the caller") + } + if elapsed := time.Since(start); elapsed > time.Second { + t.Errorf("returned after %s, want the 150ms deadline to apply", elapsed) + } +} diff --git a/internal/dhcpd/engine.go b/internal/dhcpd/engine.go index 15f81de..3c431da 100644 --- a/internal/dhcpd/engine.go +++ b/internal/dhcpd/engine.go @@ -36,10 +36,9 @@ func (s *Store) Handle(req *dhcpv4.DHCPv4) (*dhcpv4.DHCPv4, error) { return BuildReply(subnet, host, req) } -func (s *Store) Probe(mac string) (*dhcpv4.DHCPv4, error) { - key, err := normalizeMAC(mac) - if err != nil { - return nil, err +func (s *Store) Probe(mac net.HardwareAddr) (*dhcpv4.DHCPv4, error) { + if len(mac) == 0 { + return nil, ErrNoMAC } subnet, configured := s.Subnet() @@ -47,17 +46,12 @@ func (s *Store) Probe(mac string) (*dhcpv4.DHCPv4, error) { return nil, ErrNotConfigured } - parsed, err := net.ParseMAC(key) - if err != nil { - return nil, err - } - - host, known := s.Lookup(parsed) + host, known := s.Lookup(mac) if !known { return nil, nil } - req, err := dhcpv4.New(dhcpv4.WithMessageType(dhcpv4.MessageTypeRequest), dhcpv4.WithHwAddr(parsed)) + req, err := dhcpv4.New(dhcpv4.WithMessageType(dhcpv4.MessageTypeRequest), dhcpv4.WithHwAddr(mac)) if err != nil { return nil, err } diff --git a/internal/dhcpd/engine_test.go b/internal/dhcpd/engine_test.go index f035493..ae40cc3 100644 --- a/internal/dhcpd/engine_test.go +++ b/internal/dhcpd/engine_test.go @@ -11,24 +11,15 @@ import ( func configuredStore(t *testing.T) *Store { t.Helper() s, _ := loadedStore(t) - if err := s.SetSubnet(testSubnetSnapshot()); err != nil { + if err := s.SetSubnet(fullConfig(t)); err != nil { t.Fatalf("SetSubnet: %v", err) } - if err := s.SetHost(testHostSnapshot()); err != nil { + if err := s.SetHost(testHost(t)); err != nil { t.Fatalf("SetHost: %v", err) } return s } -func mac(t *testing.T, s string) net.HardwareAddr { - t.Helper() - m, err := net.ParseMAC(s) - if err != nil { - t.Fatalf("ParseMAC(%q): %v", s, err) - } - return m -} - func TestHandle_KnownMACGetsAnOfferOnDiscover(t *testing.T) { s := configuredStore(t) @@ -61,7 +52,7 @@ func TestHandle_UnknownMACIsAnsweredWithSilence(t *testing.T) { func TestHandle_UnconfiguredSubnetIsAnsweredWithSilence(t *testing.T) { s, _ := loadedStore(t) - if err := s.SetHost(testHostSnapshot()); err != nil { + if err := s.SetHost(testHost(t)); err != nil { t.Fatalf("SetHost: %v", err) } @@ -120,7 +111,7 @@ func TestHandle_NilRequestIsRejected(t *testing.T) { func TestHandle_DeletedHostStopsBeingAnswered(t *testing.T) { s := configuredStore(t) - if err := s.DelHost("00:22:33:00:00:0a"); err != nil { + if err := s.DelHost(mac(t, "00:22:33:00:00:0a")); err != nil { t.Fatalf("DelHost: %v", err) } @@ -136,7 +127,7 @@ func TestHandle_DeletedHostStopsBeingAnswered(t *testing.T) { func TestProbe_ReturnsWhatWouldBeSentToTheMAC(t *testing.T) { s := configuredStore(t) - reply, err := s.Probe("00:22:33:00:00:0A") + reply, err := s.Probe(mac(t, "00:22:33:00:00:0A")) if err != nil { t.Fatalf("Probe: %v", err) } @@ -154,7 +145,7 @@ func TestProbe_ReturnsWhatWouldBeSentToTheMAC(t *testing.T) { func TestProbe_UnknownMACReturnsNothing(t *testing.T) { s := configuredStore(t) - reply, err := s.Probe("00:22:33:ff:ff:ff") + reply, err := s.Probe(mac(t, "00:22:33:ff:ff:ff")) if err != nil { t.Fatalf("Probe: %v", err) } @@ -166,26 +157,26 @@ func TestProbe_UnknownMACReturnsNothing(t *testing.T) { func TestProbe_WithoutSubnetConfigurationIsRejected(t *testing.T) { s, _ := loadedStore(t) - if _, err := s.Probe("00:22:33:00:00:0a"); !errors.Is(err, ErrNotConfigured) { + if _, err := s.Probe(mac(t, "00:22:33:00:00:0a")); !errors.Is(err, ErrNotConfigured) { t.Fatalf("error = %v, want ErrNotConfigured", err) } } -func TestProbe_InvalidMACIsRejected(t *testing.T) { +func TestProbe_EmptyMACIsRejected(t *testing.T) { s := configuredStore(t) - if _, err := s.Probe("nope"); err == nil { - t.Fatal("an invalid mac must be reported") + if _, err := s.Probe(nil); !errors.Is(err, ErrNoMAC) { + t.Fatalf("error = %v, want ErrNoMAC", err) } } func TestHandle_SecondaryInterfaceGetsNoDefaultRoute(t *testing.T) { s := configuredStore(t) - snap := testHostSnapshot() - snap.MAC = "00:22:33:00:00:0b" - snap.IP = "10.0.5.11" - snap.DefaultRoute = false - if err := s.SetHost(snap); err != nil { + h := testHost(t) + h.MAC = mac(t, "00:22:33:00:00:0b") + h.IP = net.ParseIP("10.0.5.11") + h.DefaultRoute = false + if err := s.SetHost(h); err != nil { t.Fatalf("SetHost: %v", err) } diff --git a/internal/dhcpd/reply_test.go b/internal/dhcpd/reply_test.go index f96a126..702b430 100644 --- a/internal/dhcpd/reply_test.go +++ b/internal/dhcpd/reply_test.go @@ -28,13 +28,26 @@ func testConfig(t *testing.T) SubnetConfig { } } +func fullConfig(t *testing.T) SubnetConfig { + t.Helper() + c := testConfig(t) + c.VPCRoute = cidr(t, "10.0.0.0/16") + c.DefaultGateway = net.ParseIP("10.0.5.254") + return c +} + +func mac(t *testing.T, s string) net.HardwareAddr { + t.Helper() + m, err := net.ParseMAC(s) + if err != nil { + t.Fatalf("ParseMAC(%q): %v", s, err) + } + return m +} + func testHost(t *testing.T) Host { t.Helper() - mac, err := net.ParseMAC("00:22:33:00:00:0a") - if err != nil { - t.Fatalf("ParseMAC: %v", err) - } - return Host{MAC: mac, IP: net.ParseIP("10.0.5.10"), VM: "vm-test", DefaultRoute: true} + return Host{MAC: mac(t, "00:22:33:00:00:0a"), IP: net.ParseIP("10.0.5.10"), VM: "vm-test", DefaultRoute: true} } func request(t *testing.T, kind dhcpv4.MessageType, mac net.HardwareAddr) *dhcpv4.DHCPv4 { diff --git a/internal/dhcpd/state.go b/internal/dhcpd/state.go deleted file mode 100644 index 79f1905..0000000 --- a/internal/dhcpd/state.go +++ /dev/null @@ -1,285 +0,0 @@ -package dhcpd - -import ( - "encoding/json" - "errors" - "fmt" - "net" - "os" - "path/filepath" - "sort" - "sync" -) - -const ( - stateFileMode = 0o600 - stateDirMode = 0o700 -) - -var ( - ErrNoMAC = errors.New("host mac is required") - ErrNotConfigured = errors.New("subnet is not configured") -) - -type SubnetSnapshot struct { - Network string `json:"network"` - InterfaceIP string `json:"interface_ip"` - VPCRoute string `json:"vpc_route,omitempty"` - DefaultGateway string `json:"default_gateway,omitempty"` -} - -type HostSnapshot struct { - MAC string `json:"mac"` - IP string `json:"ip"` - VM string `json:"vm,omitempty"` - DefaultRoute bool `json:"default_route"` -} - -type Snapshot struct { - Subnet *SubnetSnapshot `json:"subnet,omitempty"` - Hosts []HostSnapshot `json:"hosts"` -} - -type Store struct { - mu sync.RWMutex - path string - subnet SubnetConfig - configured bool - hosts map[string]Host -} - -func NewStore(path string) *Store { - return &Store{path: path, hosts: make(map[string]Host)} -} - -func normalizeMAC(s string) (string, error) { - mac, err := net.ParseMAC(s) - if err != nil { - return "", fmt.Errorf("invalid mac %q: %w", s, err) - } - return mac.String(), nil -} - -func (s *Store) Load() error { - s.mu.Lock() - defer s.mu.Unlock() - - raw, err := os.ReadFile(s.path) - if errors.Is(err, os.ErrNotExist) { - return s.persist() - } - if err != nil { - return fmt.Errorf("read %s: %w", s.path, err) - } - - var snap Snapshot - if len(raw) > 0 { - if err := json.Unmarshal(raw, &snap); err != nil { - return fmt.Errorf("parse %s: %w", s.path, err) - } - } - return s.apply(snap) -} - -func (s *Store) apply(snap Snapshot) error { - subnet := SubnetConfig{} - configured := false - if snap.Subnet != nil { - parsed, err := parseSubnet(*snap.Subnet) - if err != nil { - return err - } - subnet = parsed - configured = true - } - - hosts := make(map[string]Host, len(snap.Hosts)) - for _, h := range snap.Hosts { - host, err := parseHost(h) - if err != nil { - return err - } - hosts[host.MAC.String()] = host - } - - s.subnet = subnet - s.configured = configured - s.hosts = hosts - return nil -} - -func parseSubnet(s SubnetSnapshot) (SubnetConfig, error) { - _, network, err := net.ParseCIDR(s.Network) - if err != nil { - return SubnetConfig{}, fmt.Errorf("invalid network %q: %w", s.Network, err) - } - - interfaceIP := net.ParseIP(s.InterfaceIP) - if interfaceIP == nil { - return SubnetConfig{}, ErrNoInterfaceIP - } - - c := SubnetConfig{Network: network, InterfaceIP: interfaceIP} - - if s.VPCRoute != "" { - _, vpcRoute, err := net.ParseCIDR(s.VPCRoute) - if err != nil { - return SubnetConfig{}, fmt.Errorf("invalid vpc route %q: %w", s.VPCRoute, err) - } - c.VPCRoute = vpcRoute - } - if s.DefaultGateway != "" { - gw := net.ParseIP(s.DefaultGateway) - if gw == nil { - return SubnetConfig{}, fmt.Errorf("invalid default gateway %q", s.DefaultGateway) - } - c.DefaultGateway = gw - } - return c, nil -} - -func parseHost(h HostSnapshot) (Host, error) { - mac, err := net.ParseMAC(h.MAC) - if err != nil { - return Host{}, fmt.Errorf("invalid mac %q: %w", h.MAC, err) - } - ip := net.ParseIP(h.IP) - if ip == nil { - return Host{}, fmt.Errorf("invalid host ip %q", h.IP) - } - return Host{MAC: mac, IP: ip, VM: h.VM, DefaultRoute: h.DefaultRoute}, nil -} - -func (s *Store) SetSubnet(snap SubnetSnapshot) error { - c, err := parseSubnet(snap) - if err != nil { - return err - } - - s.mu.Lock() - defer s.mu.Unlock() - - s.subnet = c - s.configured = true - return s.persist() -} - -func (s *Store) SetHost(snap HostSnapshot) error { - if snap.MAC == "" { - return ErrNoMAC - } - host, err := parseHost(snap) - if err != nil { - return err - } - - s.mu.Lock() - defer s.mu.Unlock() - - s.hosts[host.MAC.String()] = host - return s.persist() -} - -func (s *Store) DelHost(mac string) error { - key, err := normalizeMAC(mac) - if err != nil { - return err - } - - s.mu.Lock() - defer s.mu.Unlock() - - delete(s.hosts, key) - return s.persist() -} - -func (s *Store) Lookup(mac net.HardwareAddr) (Host, bool) { - s.mu.RLock() - defer s.mu.RUnlock() - - h, ok := s.hosts[mac.String()] - return h, ok -} - -func (s *Store) Subnet() (SubnetConfig, bool) { - s.mu.RLock() - defer s.mu.RUnlock() - - return s.subnet, s.configured -} - -func (s *Store) Snapshot() Snapshot { - s.mu.RLock() - defer s.mu.RUnlock() - - return s.snapshot() -} - -func (s *Store) snapshot() Snapshot { - snap := Snapshot{Hosts: make([]HostSnapshot, 0, len(s.hosts))} - - if s.configured { - sub := SubnetSnapshot{ - Network: s.subnet.Network.String(), - InterfaceIP: s.subnet.InterfaceIP.String(), - } - if s.subnet.VPCRoute != nil { - sub.VPCRoute = s.subnet.VPCRoute.String() - } - if s.subnet.DefaultGateway != nil { - sub.DefaultGateway = s.subnet.DefaultGateway.String() - } - snap.Subnet = &sub - } - - for _, h := range s.hosts { - snap.Hosts = append(snap.Hosts, HostSnapshot{ - MAC: h.MAC.String(), - IP: h.IP.String(), - VM: h.VM, - DefaultRoute: h.DefaultRoute, - }) - } - sort.Slice(snap.Hosts, func(i, j int) bool { return snap.Hosts[i].MAC < snap.Hosts[j].MAC }) - - return snap -} - -func (s *Store) persist() error { - raw, err := json.Marshal(s.snapshot()) - if err != nil { - return fmt.Errorf("encode state: %w", err) - } - - dir := filepath.Dir(s.path) - if err := os.MkdirAll(dir, stateDirMode); err != nil { - return fmt.Errorf("create %s: %w", dir, err) - } - - tmp, err := os.CreateTemp(dir, filepath.Base(s.path)+".tmp") - if err != nil { - return fmt.Errorf("create temp state in %s: %w", dir, err) - } - defer os.Remove(tmp.Name()) - - if err := tmp.Chmod(stateFileMode); err != nil { - tmp.Close() - return fmt.Errorf("chmod %s: %w", tmp.Name(), err) - } - if _, err := tmp.Write(raw); err != nil { - tmp.Close() - return fmt.Errorf("write %s: %w", tmp.Name(), err) - } - if err := tmp.Sync(); err != nil { - tmp.Close() - return fmt.Errorf("sync %s: %w", tmp.Name(), err) - } - if err := tmp.Close(); err != nil { - return fmt.Errorf("close %s: %w", tmp.Name(), err) - } - - if err := os.Rename(tmp.Name(), s.path); err != nil { - return fmt.Errorf("rename %s to %s: %w", tmp.Name(), s.path, err) - } - return nil -} diff --git a/internal/dhcpd/store.go b/internal/dhcpd/store.go new file mode 100644 index 0000000..86b756d --- /dev/null +++ b/internal/dhcpd/store.go @@ -0,0 +1,255 @@ +package dhcpd + +import ( + "errors" + "fmt" + "net" + "sort" + "sync" + + "git.g3e.fr/syonad/two/pkg/db/statefile" +) + +var ( + ErrNoMAC = errors.New("host mac is required") + ErrNotConfigured = errors.New("subnet is not configured") +) + +type diskSubnet struct { + Network string `json:"network"` + InterfaceIP string `json:"interface_ip"` + VPCRoute string `json:"vpc_route,omitempty"` + DefaultGateway string `json:"default_gateway,omitempty"` +} + +type diskHost struct { + MAC string `json:"mac"` + IP string `json:"ip"` + VM string `json:"vm,omitempty"` + DefaultRoute bool `json:"default_route"` +} + +type diskState struct { + Subnet *diskSubnet `json:"subnet,omitempty"` + Hosts []diskHost `json:"hosts"` +} + +type Store struct { + mu sync.RWMutex + file *statefile.File[diskState] + subnet SubnetConfig + configured bool + hosts map[string]Host +} + +func NewStore(path string) *Store { + return &Store{ + file: statefile.New[diskState](path), + hosts: make(map[string]Host), + } +} + +func (s *Store) Path() string { + return s.file.Path() +} + +func (s *Store) Load() error { + state, err := s.file.Load() + if err != nil { + return err + } + + subnet := SubnetConfig{} + configured := false + if state.Subnet != nil { + parsed, err := subnetFromDisk(*state.Subnet) + if err != nil { + return err + } + subnet = parsed + configured = true + } + + hosts := make(map[string]Host, len(state.Hosts)) + for _, h := range state.Hosts { + host, err := hostFromDisk(h) + if err != nil { + return err + } + hosts[host.MAC.String()] = host + } + + s.mu.Lock() + defer s.mu.Unlock() + + s.subnet = subnet + s.configured = configured + s.hosts = hosts + return nil +} + +func subnetFromDisk(d diskSubnet) (SubnetConfig, error) { + _, network, err := net.ParseCIDR(d.Network) + if err != nil { + return SubnetConfig{}, fmt.Errorf("invalid network %q: %w", d.Network, err) + } + interfaceIP := net.ParseIP(d.InterfaceIP) + if interfaceIP == nil { + return SubnetConfig{}, ErrNoInterfaceIP + } + + c := SubnetConfig{Network: network, InterfaceIP: interfaceIP} + + if d.VPCRoute != "" { + if _, c.VPCRoute, err = net.ParseCIDR(d.VPCRoute); err != nil { + return SubnetConfig{}, fmt.Errorf("invalid vpc route %q: %w", d.VPCRoute, err) + } + } + if d.DefaultGateway != "" { + if c.DefaultGateway = net.ParseIP(d.DefaultGateway); c.DefaultGateway == nil { + return SubnetConfig{}, fmt.Errorf("invalid default gateway %q", d.DefaultGateway) + } + } + return c, nil +} + +func hostFromDisk(d diskHost) (Host, error) { + mac, err := net.ParseMAC(d.MAC) + if err != nil { + return Host{}, fmt.Errorf("invalid mac %q: %w", d.MAC, err) + } + ip := net.ParseIP(d.IP) + if ip == nil { + return Host{}, fmt.Errorf("invalid host ip %q", d.IP) + } + return Host{MAC: mac, IP: ip, VM: d.VM, DefaultRoute: d.DefaultRoute}, nil +} + +func (s *Store) SetSubnet(c SubnetConfig) error { + if c.Network == nil { + return ErrNoNetwork + } + if c.InterfaceIP == nil { + return ErrNoInterfaceIP + } + + s.mu.Lock() + defer s.mu.Unlock() + + previous, wasConfigured := s.subnet, s.configured + s.subnet, s.configured = c, true + + if err := s.persist(); err != nil { + s.subnet, s.configured = previous, wasConfigured + return err + } + return nil +} + +func (s *Store) SetHost(h Host) error { + if len(h.MAC) == 0 { + return ErrNoMAC + } + if h.IP == nil { + return ErrNoHostIP + } + + key := h.MAC.String() + + s.mu.Lock() + defer s.mu.Unlock() + + previous, existed := s.hosts[key] + s.hosts[key] = h + + if err := s.persist(); err != nil { + if existed { + s.hosts[key] = previous + } else { + delete(s.hosts, key) + } + return err + } + return nil +} + +func (s *Store) DelHost(mac net.HardwareAddr) error { + if len(mac) == 0 { + return ErrNoMAC + } + key := mac.String() + + s.mu.Lock() + defer s.mu.Unlock() + + previous, existed := s.hosts[key] + delete(s.hosts, key) + + if err := s.persist(); err != nil { + if existed { + s.hosts[key] = previous + } + return err + } + return nil +} + +func (s *Store) Lookup(mac net.HardwareAddr) (Host, bool) { + s.mu.RLock() + defer s.mu.RUnlock() + + h, ok := s.hosts[mac.String()] + return h, ok +} + +func (s *Store) Subnet() (SubnetConfig, bool) { + s.mu.RLock() + defer s.mu.RUnlock() + + return s.subnet, s.configured +} + +func (s *Store) Hosts() []Host { + s.mu.RLock() + defer s.mu.RUnlock() + + return s.sortedHosts() +} + +func (s *Store) sortedHosts() []Host { + hosts := make([]Host, 0, len(s.hosts)) + for _, h := range s.hosts { + hosts = append(hosts, h) + } + sort.Slice(hosts, func(i, j int) bool { return hosts[i].MAC.String() < hosts[j].MAC.String() }) + return hosts +} + +func (s *Store) persist() error { + state := diskState{Hosts: make([]diskHost, 0, len(s.hosts))} + + if s.configured { + sub := diskSubnet{ + Network: s.subnet.Network.String(), + InterfaceIP: s.subnet.InterfaceIP.String(), + } + if s.subnet.VPCRoute != nil { + sub.VPCRoute = s.subnet.VPCRoute.String() + } + if s.subnet.DefaultGateway != nil { + sub.DefaultGateway = s.subnet.DefaultGateway.String() + } + state.Subnet = &sub + } + + for _, h := range s.sortedHosts() { + state.Hosts = append(state.Hosts, diskHost{ + MAC: h.MAC.String(), + IP: h.IP.String(), + VM: h.VM, + DefaultRoute: h.DefaultRoute, + }) + } + + return s.file.Save(state) +} diff --git a/internal/dhcpd/state_test.go b/internal/dhcpd/store_test.go similarity index 55% rename from internal/dhcpd/state_test.go rename to internal/dhcpd/store_test.go index 7240ddb..11a57d5 100644 --- a/internal/dhcpd/state_test.go +++ b/internal/dhcpd/store_test.go @@ -14,19 +14,6 @@ func statePath(t *testing.T) string { return filepath.Join(t.TempDir(), "vp-admin_br-000001.state") } -func testSubnetSnapshot() SubnetSnapshot { - return SubnetSnapshot{ - Network: "10.0.5.0/24", - InterfaceIP: "10.0.5.1", - VPCRoute: "10.0.0.0/16", - DefaultGateway: "10.0.5.254", - } -} - -func testHostSnapshot() HostSnapshot { - return HostSnapshot{MAC: "00:22:33:00:00:0a", IP: "10.0.5.10", VM: "vm-test", DefaultRoute: true} -} - func loadedStore(t *testing.T) (*Store, string) { t.Helper() path := statePath(t) @@ -49,28 +36,12 @@ func TestStore_LoadCreatesTheStateFileWhenAbsent(t *testing.T) { } } -func TestStore_PersistRestoresTheModeAfterAnExternalChmod(t *testing.T) { - s, path := loadedStore(t) - if err := os.Chmod(path, 0o644); err != nil { - t.Fatalf("Chmod: %v", err) - } - - if err := s.SetHost(testHostSnapshot()); err != nil { - t.Fatalf("SetHost: %v", err) - } - - info, err := os.Stat(path) - if err != nil { - t.Fatalf("Stat: %v", err) - } - if got := info.Mode().Perm(); got != 0o600 { - t.Errorf("mode = %o, want 600: each write must replace the file, not edit it in place", got) - } -} - func TestStore_LoadOnEmptyFileYieldsNoSubnet(t *testing.T) { path := statePath(t) - if err := os.WriteFile(path, nil, stateFileMode); err != nil { + if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil { + t.Fatalf("MkdirAll: %v", err) + } + if err := os.WriteFile(path, nil, 0o600); err != nil { t.Fatalf("WriteFile: %v", err) } @@ -85,7 +56,7 @@ func TestStore_LoadOnEmptyFileYieldsNoSubnet(t *testing.T) { func TestStore_LoadRejectsCorruptedState(t *testing.T) { path := statePath(t) - if err := os.WriteFile(path, []byte("{not json"), stateFileMode); err != nil { + if err := os.WriteFile(path, []byte("{not json"), 0o600); err != nil { t.Fatalf("WriteFile: %v", err) } @@ -94,9 +65,21 @@ func TestStore_LoadRejectsCorruptedState(t *testing.T) { } } +func TestStore_LoadRejectsAnInvalidStoredMAC(t *testing.T) { + path := statePath(t) + raw := []byte(`{"hosts":[{"mac":"nope","ip":"10.0.5.10","default_route":true}]}`) + if err := os.WriteFile(path, raw, 0o600); err != nil { + t.Fatalf("WriteFile: %v", err) + } + + if err := NewStore(path).Load(); err == nil { + t.Fatal("an unparseable stored mac must be reported") + } +} + func TestStore_SetSubnetIsPersisted(t *testing.T) { s, path := loadedStore(t) - if err := s.SetSubnet(testSubnetSnapshot()); err != nil { + if err := s.SetSubnet(fullConfig(t)); err != nil { t.Fatalf("SetSubnet: %v", err) } @@ -125,27 +108,27 @@ func TestStore_SetSubnetIsPersisted(t *testing.T) { func TestStore_SetSubnetRejectsAMissingInterfaceIP(t *testing.T) { s, _ := loadedStore(t) - snap := testSubnetSnapshot() - snap.InterfaceIP = "" + c := fullConfig(t) + c.InterfaceIP = nil - if err := s.SetSubnet(snap); !errors.Is(err, ErrNoInterfaceIP) { + if err := s.SetSubnet(c); !errors.Is(err, ErrNoInterfaceIP) { t.Fatalf("error = %v, want ErrNoInterfaceIP", err) } } -func TestStore_SetSubnetRejectsAnInvalidNetwork(t *testing.T) { +func TestStore_SetSubnetRejectsAMissingNetwork(t *testing.T) { s, _ := loadedStore(t) - snap := testSubnetSnapshot() - snap.Network = "10.0.5.0" + c := fullConfig(t) + c.Network = nil - if err := s.SetSubnet(snap); err == nil { - t.Fatal("a network without a prefix length must be rejected") + if err := s.SetSubnet(c); !errors.Is(err, ErrNoNetwork) { + t.Fatalf("error = %v, want ErrNoNetwork", err) } } func TestStore_SetHostIsPersistedAndFound(t *testing.T) { s, path := loadedStore(t) - if err := s.SetHost(testHostSnapshot()); err != nil { + if err := s.SetHost(testHost(t)); err != nil { t.Fatalf("SetHost: %v", err) } @@ -154,11 +137,7 @@ func TestStore_SetHostIsPersistedAndFound(t *testing.T) { t.Fatalf("Load: %v", err) } - mac, err := net.ParseMAC("00:22:33:00:00:0a") - if err != nil { - t.Fatalf("ParseMAC: %v", err) - } - h, known := reloaded.Lookup(mac) + h, known := reloaded.Lookup(mac(t, "00:22:33:00:00:0a")) if !known { t.Fatal("host lost across a restart") } @@ -175,52 +154,43 @@ func TestStore_SetHostIsPersistedAndFound(t *testing.T) { func TestStore_SetHostIsIdempotentOnTheSameMAC(t *testing.T) { s, _ := loadedStore(t) - snap := testHostSnapshot() - if err := s.SetHost(snap); err != nil { + h := testHost(t) + if err := s.SetHost(h); err != nil { t.Fatalf("SetHost: %v", err) } - snap.IP = "10.0.5.11" - snap.DefaultRoute = false - if err := s.SetHost(snap); err != nil { + h.IP = net.ParseIP("10.0.5.11") + h.DefaultRoute = false + if err := s.SetHost(h); err != nil { t.Fatalf("SetHost: %v", err) } - got := s.Snapshot() - if len(got.Hosts) != 1 { - t.Fatalf("hosts = %d, want 1: the mac is the key", len(got.Hosts)) + hosts := s.Hosts() + if len(hosts) != 1 { + t.Fatalf("hosts = %d, want 1: the mac is the key", len(hosts)) } - if got.Hosts[0].IP != "10.0.5.11" || got.Hosts[0].DefaultRoute { - t.Errorf("entry = %+v, want the second order to have replaced the first", got.Hosts[0]) + if !hosts[0].IP.Equal(net.ParseIP("10.0.5.11")) || hosts[0].DefaultRoute { + t.Errorf("entry = %+v, want the second order to have replaced the first", hosts[0]) } } -func TestStore_SetHostNormalizesTheMACCase(t *testing.T) { +func TestStore_LookupNormalizesTheMACCase(t *testing.T) { s, _ := loadedStore(t) - snap := testHostSnapshot() - snap.MAC = "00:22:33:AA:BB:CC" - if err := s.SetHost(snap); err != nil { + h := testHost(t) + h.MAC = mac(t, "00:22:33:AA:BB:CC") + if err := s.SetHost(h); err != nil { t.Fatalf("SetHost: %v", err) } - mac, err := net.ParseMAC("00:22:33:aa:bb:cc") - if err != nil { - t.Fatalf("ParseMAC: %v", err) - } - if _, known := s.Lookup(mac); !known { + if _, known := s.Lookup(mac(t, "00:22:33:aa:bb:cc")); !known { t.Error("an uppercase mac must be found in lowercase: the key would diverge") } } func TestStore_LoadNormalizesTheMACCase(t *testing.T) { path := statePath(t) - raw, err := json.Marshal(Snapshot{Hosts: []HostSnapshot{{ - MAC: "00:22:33:AA:BB:CC", IP: "10.0.5.12", VM: "vm-test", DefaultRoute: true, - }}}) - if err != nil { - t.Fatalf("Marshal: %v", err) - } - if err := os.WriteFile(path, raw, stateFileMode); err != nil { + raw := []byte(`{"hosts":[{"mac":"00:22:33:AA:BB:CC","ip":"10.0.5.12","default_route":true}]}`) + if err := os.WriteFile(path, raw, 0o600); err != nil { t.Fatalf("WriteFile: %v", err) } @@ -228,42 +198,37 @@ func TestStore_LoadNormalizesTheMACCase(t *testing.T) { if err := s.Load(); err != nil { t.Fatalf("Load: %v", err) } - - m, err := net.ParseMAC("00:22:33:aa:bb:cc") - if err != nil { - t.Fatalf("ParseMAC: %v", err) - } - if _, known := s.Lookup(m); !known { + if _, known := s.Lookup(mac(t, "00:22:33:aa:bb:cc")); !known { t.Error("a reloaded uppercase mac must be keyed in lowercase: the host would silently stop being served") } } func TestStore_SetHostRejectsAMissingMAC(t *testing.T) { s, _ := loadedStore(t) - snap := testHostSnapshot() - snap.MAC = "" + h := testHost(t) + h.MAC = nil - if err := s.SetHost(snap); !errors.Is(err, ErrNoMAC) { + if err := s.SetHost(h); !errors.Is(err, ErrNoMAC) { t.Fatalf("error = %v, want ErrNoMAC", err) } } -func TestStore_SetHostRejectsAnInvalidIP(t *testing.T) { +func TestStore_SetHostRejectsAMissingIP(t *testing.T) { s, _ := loadedStore(t) - snap := testHostSnapshot() - snap.IP = "10.0.5.300" + h := testHost(t) + h.IP = nil - if err := s.SetHost(snap); err == nil { - t.Fatal("an invalid host ip must be rejected") + if err := s.SetHost(h); !errors.Is(err, ErrNoHostIP) { + t.Fatalf("error = %v, want ErrNoHostIP", err) } } func TestStore_DelHostRemovesTheEntry(t *testing.T) { s, path := loadedStore(t) - if err := s.SetHost(testHostSnapshot()); err != nil { + if err := s.SetHost(testHost(t)); err != nil { t.Fatalf("SetHost: %v", err) } - if err := s.DelHost("00:22:33:00:00:0A"); err != nil { + if err := s.DelHost(mac(t, "00:22:33:00:00:0A")); err != nil { t.Fatalf("DelHost: %v", err) } @@ -271,76 +236,92 @@ func TestStore_DelHostRemovesTheEntry(t *testing.T) { if err := reloaded.Load(); err != nil { t.Fatalf("Load: %v", err) } - if got := len(reloaded.Snapshot().Hosts); got != 0 { + if got := len(reloaded.Hosts()); got != 0 { t.Errorf("hosts = %d, want 0 after deletion", got) } } func TestStore_DelHostOnAnUnknownMACIsNotAnError(t *testing.T) { s, _ := loadedStore(t) - if err := s.DelHost("00:22:33:ff:ff:ff"); err != nil { + if err := s.DelHost(mac(t, "00:22:33:ff:ff:ff")); err != nil { t.Errorf("deleting an absent entry must be idempotent, got %v", err) } } -func TestStore_DelHostRejectsAnInvalidMAC(t *testing.T) { +func TestStore_DelHostRejectsAnEmptyMAC(t *testing.T) { s, _ := loadedStore(t) - if err := s.DelHost("not-a-mac"); err == nil { - t.Fatal("an invalid mac must be reported") + if err := s.DelHost(nil); !errors.Is(err, ErrNoMAC) { + t.Fatalf("error = %v, want ErrNoMAC", err) } } -func TestStore_SnapshotSortsHostsByMAC(t *testing.T) { +func TestStore_HostsAreSortedByMAC(t *testing.T) { s, _ := loadedStore(t) - for _, mac := range []string{"00:22:33:00:00:0c", "00:22:33:00:00:0a", "00:22:33:00:00:0b"} { - snap := testHostSnapshot() - snap.MAC = mac - if err := s.SetHost(snap); err != nil { + for _, m := range []string{"00:22:33:00:00:0c", "00:22:33:00:00:0a", "00:22:33:00:00:0b"} { + h := testHost(t) + h.MAC = mac(t, m) + if err := s.SetHost(h); err != nil { t.Fatalf("SetHost: %v", err) } } - hosts := s.Snapshot().Hosts + hosts := s.Hosts() for i := 1; i < len(hosts); i++ { - if hosts[i-1].MAC >= hosts[i].MAC { + if hosts[i-1].MAC.String() >= hosts[i].MAC.String() { t.Fatalf("hosts are not sorted: %v", hosts) } } } -func TestStore_PersistLeavesNoTemporaryFileBehind(t *testing.T) { +func TestStore_PersistedStateIsSortedOnDisk(t *testing.T) { s, path := loadedStore(t) - if err := s.SetHost(testHostSnapshot()); err != nil { - t.Fatalf("SetHost: %v", err) - } - - entries, err := os.ReadDir(filepath.Dir(path)) - if err != nil { - t.Fatalf("ReadDir: %v", err) - } - if len(entries) != 1 { - t.Errorf("directory holds %d entries, want only the state file: %v", len(entries), entries) - } -} - -func TestStore_PersistedStateIsValidJSON(t *testing.T) { - s, path := loadedStore(t) - if err := s.SetSubnet(testSubnetSnapshot()); err != nil { - t.Fatalf("SetSubnet: %v", err) - } - if err := s.SetHost(testHostSnapshot()); err != nil { - t.Fatalf("SetHost: %v", err) + for _, m := range []string{"00:22:33:00:00:0c", "00:22:33:00:00:0a"} { + h := testHost(t) + h.MAC = mac(t, m) + if err := s.SetHost(h); err != nil { + t.Fatalf("SetHost: %v", err) + } } raw, err := os.ReadFile(path) if err != nil { t.Fatalf("ReadFile: %v", err) } - var snap Snapshot - if err := json.Unmarshal(raw, &snap); err != nil { + var state struct { + Hosts []struct { + MAC string `json:"mac"` + } `json:"hosts"` + } + if err := json.Unmarshal(raw, &state); err != nil { t.Fatalf("the state file must stay parseable: %v", err) } - if snap.Subnet == nil || len(snap.Hosts) != 1 { - t.Errorf("snapshot = %+v, want one subnet and one host", snap) + if len(state.Hosts) != 2 || state.Hosts[0].MAC != "00:22:33:00:00:0a" { + t.Errorf("hosts on disk = %+v, want sorted by mac", state.Hosts) + } +} + +func TestStore_PersistRestoresTheModeAfterAnExternalChmod(t *testing.T) { + s, path := loadedStore(t) + if err := os.Chmod(path, 0o644); err != nil { + t.Fatalf("Chmod: %v", err) + } + + if err := s.SetHost(testHost(t)); err != nil { + t.Fatalf("SetHost: %v", err) + } + + info, err := os.Stat(path) + if err != nil { + t.Fatalf("Stat: %v", err) + } + if got := info.Mode().Perm(); got != 0o600 { + t.Errorf("mode = %o, want 600: each write must replace the file, not edit it in place", got) + } +} + +func TestStore_PathReportsTheStateFile(t *testing.T) { + s, path := loadedStore(t) + if got := s.Path(); got != path { + t.Errorf("Path = %s, want %s", got, path) } }