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