diff --git a/internal/machine/cluster/cluster.go b/internal/machine/cluster/cluster.go index 3db9e760..87ba85be 100644 --- a/internal/machine/cluster/cluster.go +++ b/internal/machine/cluster/cluster.go @@ -14,6 +14,7 @@ import ( "github.com/psviderski/uncloud/internal/machine/network" "github.com/psviderski/uncloud/internal/machine/store" "github.com/psviderski/uncloud/internal/secret" + "github.com/psviderski/uncloud/pkg/api" "google.golang.org/grpc/codes" "google.golang.org/grpc/status" "google.golang.org/protobuf/types/known/emptypb" @@ -87,6 +88,11 @@ func (c *Cluster) AddMachine(ctx context.Context, req *pb.AddMachineRequest) (*p func (c *Cluster) AddMachineWithoutReadyCheck( ctx context.Context, req *pb.AddMachineRequest, ) (*pb.AddMachineResponse, error) { + if req.Name != "" { + if err := api.ValidateMachineName(req.Name); err != nil { + return nil, status.Error(codes.InvalidArgument, err.Error()) + } + } if req.Network == nil { return nil, status.Error(codes.InvalidArgument, "network not set") } diff --git a/internal/machine/cluster/cluster_test.go b/internal/machine/cluster/cluster_test.go new file mode 100644 index 00000000..a798f60a --- /dev/null +++ b/internal/machine/cluster/cluster_test.go @@ -0,0 +1,28 @@ +package cluster + +import ( + "context" + "testing" + + "github.com/psviderski/uncloud/api/pb" + "github.com/stretchr/testify/require" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" +) + +func TestAddMachine_InvalidName(t *testing.T) { + t.Parallel() + + ready := make(chan struct{}) + close(ready) + // Invalid names must be rejected before any store access. + c := NewCluster(nil, nil, nil, ready) + for _, name := range []string{"VPS1", "worker.example", "rr", "nearest", "c337f00600de51ef4375c9a9a267dba5"} { + t.Run("add", func(t *testing.T) { + resp, err := c.AddMachine(context.Background(), &pb.AddMachineRequest{Name: name}) + require.Nil(t, resp) + require.Equal(t, codes.InvalidArgument, status.Code(err)) + require.ErrorContains(t, err, "invalid machine name") + }) + } +} diff --git a/internal/machine/cluster/machine.go b/internal/machine/cluster/machine.go index 78fbe420..3626acf8 100644 --- a/internal/machine/cluster/machine.go +++ b/internal/machine/cluster/machine.go @@ -6,6 +6,7 @@ import ( "strings" "github.com/psviderski/uncloud/internal/secret" + "github.com/psviderski/uncloud/pkg/api" ) // NewMachineID generates a new unique machine ID. @@ -27,7 +28,7 @@ func NewRandomMachineName() (string, error) { // a numeric suffix ("-1", "-2", etc.) if needed. func DefaultMachineName(hostname string, existing []string) (string, error) { name := machineNameFromHostname(hostname) - if name == "" { + if api.ValidateMachineName(name) != nil { var err error if name, err = NewRandomMachineName(); err != nil { return "", err @@ -38,7 +39,9 @@ func DefaultMachineName(hostname string, existing []string) (string, error) { return name, nil } for i := 1; ; i++ { - candidate := fmt.Sprintf("%s-%d", name, i) + suffix := fmt.Sprintf("-%d", i) + base := strings.TrimRight(name[:min(len(name), 63-len(suffix))], "-") + candidate := base + suffix if !slices.Contains(existing, candidate) { return candidate, nil } @@ -55,12 +58,13 @@ func machineNameFromHostname(hostname string) string { var b strings.Builder for _, r := range label { switch { - case r >= 'a' && r <= 'z', r >= '0' && r <= '9', r == '-', r == '_': + case r >= 'a' && r <= 'z', r >= '0' && r <= '9', r == '-': b.WriteRune(r) default: b.WriteRune('-') } } // Trim leading and trailing hyphens that may result from the sanitisation above. - return strings.Trim(b.String(), "-") + name := strings.Trim(b.String(), "-") + return strings.TrimRight(name[:min(len(name), 63)], "-") } diff --git a/internal/machine/cluster/machine_test.go b/internal/machine/cluster/machine_test.go index 1129c78f..93d9e323 100644 --- a/internal/machine/cluster/machine_test.go +++ b/internal/machine/cluster/machine_test.go @@ -1,8 +1,10 @@ package cluster import ( + "strings" "testing" + "github.com/psviderski/uncloud/pkg/api" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -18,9 +20,11 @@ func TestMachineNameFromHostname(t *testing.T) { {"simple", "web", "web"}, {"fqdn uses first label", "web-1.example.com", "web-1"}, {"uppercase lowercased", "Web-Server", "web-server"}, - {"invalid chars replaced", "host_name@1", "host_name-1"}, - {"trim surrounding hyphens", "-_host_-", "_host_"}, + {"invalid chars replaced", "host_name@1", "host-name-1"}, + {"trim surrounding hyphens", "-_host_-", "host"}, {"whitespace trimmed", " myhost ", "myhost"}, + {"long hostname truncated", strings.Repeat("a", 64), strings.Repeat("a", 63)}, + {"truncation trims trailing hyphen", strings.Repeat("a", 62) + "-b", strings.Repeat("a", 62)}, {"empty", "", ""}, {"only invalid chars", "@#", ""}, {"dot only", ".example.com", ""}, @@ -37,10 +41,12 @@ func TestDefaultMachineName(t *testing.T) { t.Parallel() t.Run("falls back to random", func(t *testing.T) { - // An empty or fully invalid hostname falls back to a random "machine-xxxx" name. - got, err := DefaultMachineName("***", nil) - require.NoError(t, err) - assert.Regexp(t, `^machine-[a-zA-Z0-9]{4}$`, got) + for _, hostname := range []string{"***", "rr", "nearest", strings.Repeat("a", 32)} { + got, err := DefaultMachineName(hostname, nil) + require.NoError(t, err) + assert.Regexp(t, `^machine-[a-z0-9]{4}$`, got) + require.NoError(t, api.ValidateMachineName(got)) + } }) tests := []struct { @@ -49,10 +55,45 @@ func TestDefaultMachineName(t *testing.T) { existing []string want string }{ - {"from hostname", "web-1.example.com", nil, "web-1"}, - {"dedup against existing", "web", []string{"web"}, "web-1"}, - {"dedup multiple", "web", []string{"web", "web-1"}, "web-2"}, - {"sanitized", "My_Host", nil, "my_host"}, + { + name: "from hostname", + hostname: "web-1.example.com", + want: "web-1", + }, + { + name: "dedup against existing", + hostname: "web", + existing: []string{"web"}, + want: "web-1", + }, + { + name: "dedup multiple", + hostname: "web", + existing: []string{"web", "web-1"}, + want: "web-2", + }, + { + name: "sanitized", + hostname: "My_Host", + want: "my-host", + }, + { + name: "long hostname", + hostname: strings.Repeat("a", 64), + want: strings.Repeat("a", 63), + }, + { + name: "dedup maximum length", + hostname: strings.Repeat("a", 63), + existing: []string{strings.Repeat("a", 63)}, + want: strings.Repeat("a", 61) + "-1", + }, + { + name: "dedup trims trailing hyphen", + hostname: strings.Repeat("a", 60) + "-bb", + existing: []string{strings.Repeat("a", 60) + "-bb"}, + want: strings.Repeat("a", 60) + "-1", + }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { @@ -60,6 +101,7 @@ func TestDefaultMachineName(t *testing.T) { got, err := DefaultMachineName(tt.hostname, tt.existing) require.NoError(t, err) assert.Equal(t, tt.want, got) + require.NoError(t, api.ValidateMachineName(got)) }) } } diff --git a/internal/machine/dns/resolver.go b/internal/machine/dns/resolver.go index a4f8074e..0e486658 100644 --- a/internal/machine/dns/resolver.go +++ b/internal/machine/dns/resolver.go @@ -10,7 +10,7 @@ import ( "sync" "time" - "github.com/psviderski/uncloud/internal/machine/api/pb" + "github.com/psviderski/uncloud/api/pb" "github.com/psviderski/uncloud/internal/machine/network" "github.com/psviderski/uncloud/internal/machine/store" ) diff --git a/internal/machine/dns/resolver_test.go b/internal/machine/dns/resolver_test.go index 11a46931..4e856801 100644 --- a/internal/machine/dns/resolver_test.go +++ b/internal/machine/dns/resolver_test.go @@ -7,7 +7,7 @@ import ( "github.com/docker/docker/api/types/container" "github.com/docker/docker/api/types/network" - "github.com/psviderski/uncloud/internal/machine/api/pb" + "github.com/psviderski/uncloud/api/pb" "github.com/psviderski/uncloud/internal/machine/store" "github.com/psviderski/uncloud/pkg/api" "github.com/stretchr/testify/assert" diff --git a/internal/machine/machine.go b/internal/machine/machine.go index 2b37998b..1680c6af 100644 --- a/internal/machine/machine.go +++ b/internal/machine/machine.go @@ -836,6 +836,11 @@ func (m *Machine) InitCluster(ctx context.Context, req *pb.InitClusterRequest) ( if m.Initialised() { return nil, status.Error(codes.FailedPrecondition, "machine is already configured as a cluster member") } + if req.MachineName != "" { + if err := api.ValidateMachineName(req.MachineName); err != nil { + return nil, status.Error(codes.InvalidArgument, err.Error()) + } + } clusterNetwork, err := req.Network.ToPrefix() if err != nil { @@ -1220,8 +1225,8 @@ func (m *Machine) applyMachineUpdate(ctx context.Context, req *pb.UpdateMachineR defer m.state.mu.Unlock() if req.Name != nil { - if *req.Name == "" { - return status.Error(codes.InvalidArgument, "machine name cannot be empty") + if err := api.ValidateMachineName(*req.Name); err != nil { + return status.Error(codes.InvalidArgument, err.Error()) } // Check for duplicate names across the cluster, excluding this machine. if *req.Name != m.state.Name { diff --git a/internal/machine/machine_test.go b/internal/machine/machine_test.go new file mode 100644 index 00000000..04128c91 --- /dev/null +++ b/internal/machine/machine_test.go @@ -0,0 +1,38 @@ +package machine + +import ( + "context" + "testing" + + "github.com/psviderski/uncloud/api/pb" + "github.com/stretchr/testify/require" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" +) + +func TestInitCluster_InvalidMachineName(t *testing.T) { + t.Parallel() + + for _, name := range []string{"VPS1", "rr", "nearest", "c337f00600de51ef4375c9a9a267dba5"} { + // Invalid names must be rejected before initializing cluster state. + m := &Machine{state: &State{}} + + resp, err := m.InitCluster(context.Background(), &pb.InitClusterRequest{MachineName: name}) + require.Nil(t, resp) + require.Equal(t, codes.InvalidArgument, status.Code(err)) + require.ErrorContains(t, err, "invalid machine name") + } +} + +func TestUpdateMachine_InvalidName(t *testing.T) { + t.Parallel() + + for _, name := range []string{"", "VPS1", "worker.example", "rr", "nearest", "c337f00600de51ef4375c9a9a267dba5"} { + m := &Machine{state: &State{ID: "machine-id", Name: "worker"}} + resp, err := m.UpdateMachine(context.Background(), &pb.UpdateMachineRequest{Name: &name}) + require.Nil(t, resp) + require.Equal(t, codes.InvalidArgument, status.Code(err)) + require.ErrorContains(t, err, "invalid machine name") + require.Equal(t, "worker", m.state.Name) + } +} diff --git a/pkg/api/machine.go b/pkg/api/machine.go index d5f95740..78b36059 100644 --- a/pkg/api/machine.go +++ b/pkg/api/machine.go @@ -1,12 +1,33 @@ package api import ( + "fmt" "net/netip" "strings" "github.com/psviderski/uncloud/api/pb" ) +// ValidateMachineName checks that a machine name is a lowercase DNS label and doesn't conflict with +// internal DNS query modes or machine IDs. +func ValidateMachineName(name string) error { + if !DNSLabelRegex.MatchString(name) { + return fmt.Errorf("invalid machine name %q: must be 1-63 characters, lowercase letters, numbers, "+ + "and hyphens only, starting and ending with a letter or number", name) + } + + switch name { + case "rr", "nearest": + return fmt.Errorf("invalid machine name %q: reserved for internal DNS query modes", name) + } + if IDRegex.MatchString(name) { + return fmt.Errorf( + "invalid machine name %q: must not match the machine ID format (32 hexadecimal characters)", name) + } + + return nil +} + // MachineFilter defines criteria to filter machines in ListMachines. type MachineFilter struct { // Available filters machines that are not DOWN. diff --git a/pkg/api/machine_test.go b/pkg/api/machine_test.go index fe268bcb..5d3d130a 100644 --- a/pkg/api/machine_test.go +++ b/pkg/api/machine_test.go @@ -3,6 +3,7 @@ package api import ( "encoding/json" "net/netip" + "strings" "testing" "github.com/psviderski/uncloud/api/pb" @@ -10,6 +11,52 @@ import ( "github.com/stretchr/testify/require" ) +func TestValidateMachineName(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + input string + wantErr string + }{ + {name: "single letter", input: "a"}, + {name: "single digit", input: "1"}, + {name: "two characters", input: "a1"}, + {name: "generated name", input: "machine-ab12"}, + {name: "hyphens and digits", input: "1-worker-2"}, + {name: "maximum length", input: strings.Repeat("a", 63)}, + {name: "maximum length with hyphens", input: "a" + strings.Repeat("-", 61) + "1"}, + {name: "machine namespace label", input: "m"}, + {name: "mode prefix", input: "nearest-worker"}, + {name: "short hexadecimal name", input: strings.Repeat("a", 31)}, + {name: "long hexadecimal name", input: strings.Repeat("a", 33)}, + {name: "non-hexadecimal 32 characters", input: strings.Repeat("g", 32)}, + {name: "empty", wantErr: "must be 1-63 characters"}, + {name: "too long", input: strings.Repeat("a", 64), wantErr: "must be 1-63 characters"}, + {name: "uppercase", input: "VPS1", wantErr: "lowercase letters"}, + {name: "leading hyphen", input: "-worker", wantErr: "starting and ending"}, + {name: "trailing hyphen", input: "worker-", wantErr: "starting and ending"}, + {name: "underscore", input: "worker_1", wantErr: "hyphens only"}, + {name: "dot", input: "worker.example", wantErr: "hyphens only"}, + {name: "space", input: "worker 1", wantErr: "hyphens only"}, + {name: "leading whitespace", input: " worker", wantErr: "hyphens only"}, + {name: "trailing whitespace", input: "worker\t", wantErr: "hyphens only"}, + {name: "non-ASCII", input: "wörker", wantErr: "lowercase letters"}, + {name: "round-robin mode", input: "rr", wantErr: "reserved for internal DNS query modes"}, + {name: "nearest mode", input: "nearest", wantErr: "reserved for internal DNS query modes"}, + {name: "machine ID", input: "c337f00600de51ef4375c9a9a267dba5", wantErr: "machine ID format"}, + } + + for _, tt := range tests { + err := ValidateMachineName(tt.input) + if tt.wantErr == "" { + require.NoError(t, err) + } else { + require.ErrorContains(t, err, tt.wantErr) + } + } +} + func TestMachineMembersList_Info(t *testing.T) { t.Parallel() diff --git a/pkg/api/service.go b/pkg/api/service.go index 77375c70..5d9ff21f 100644 --- a/pkg/api/service.go +++ b/pkg/api/service.go @@ -41,12 +41,30 @@ const ( ) var ( - serviceIDRegexp = regexp.MustCompile("^[0-9a-f]{32}$") - dnsLabelRegexp = regexp.MustCompile(`^[a-z0-9]([-a-z0-9]*[a-z0-9])?$`) + IDRegex = regexp.MustCompile("^[0-9a-f]{32}$") + DNSLabelRegex = regexp.MustCompile(`^[a-z0-9]([-a-z0-9]{0,61}[a-z0-9])?$`) ) func ValidateServiceID(id string) bool { - return serviceIDRegexp.MatchString(id) + return IDRegex.MatchString(id) +} + +// ValidateServiceName checks that a service name is a lowercase DNS label and doesn't conflict with +// the machine DNS namespace or service IDs. +func ValidateServiceName(name string) error { + if !DNSLabelRegex.MatchString(name) { + return fmt.Errorf("invalid service name %q: must be 1-63 characters, lowercase letters, numbers, "+ + "and hyphens only, starting and ending with a letter or number", name) + } + if name == "m" { + return fmt.Errorf("invalid service name %q: reserved for the machine DNS namespace", name) + } + if IDRegex.MatchString(name) { + return fmt.Errorf( + "invalid service name %q: must not match the service ID format (32 hexadecimal characters)", name) + } + + return nil } // ServiceSpec defines the desired state of a service. @@ -147,12 +165,8 @@ func (s *ServiceSpec) Validate() error { } if s.Name != "" { - if len(s.Name) > 63 { - return fmt.Errorf("service name too long (max 63 characters): %q", s.Name) - } - if !dnsLabelRegexp.MatchString(s.Name) { - return fmt.Errorf("invalid service name: %q. must be 1-63 characters, lowercase letters, numbers, "+ - "and dashes only; must start and end with a letter or number", s.Name) + if err := ValidateServiceName(s.Name); err != nil { + return err } } diff --git a/pkg/api/service_test.go b/pkg/api/service_test.go index 1b067ab5..637eb770 100644 --- a/pkg/api/service_test.go +++ b/pkg/api/service_test.go @@ -2,6 +2,7 @@ package api import ( "os" + "strings" "testing" "github.com/docker/docker/api/types/container" @@ -9,6 +10,59 @@ import ( "github.com/stretchr/testify/require" ) +func TestServiceSpec_Validate_Name(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + input string + wantErr string + }{ + {name: "empty allows generated name"}, + {name: "single character", input: "a"}, + {name: "single digit", input: "1"}, + {name: "two characters", input: "a1"}, + {name: "hyphens and digits", input: "1-web-2"}, + {name: "consecutive hyphens", input: "web--1"}, + {name: "maximum length", input: strings.Repeat("a", 63)}, + {name: "maximum length with hyphens", input: "a" + strings.Repeat("-", 61) + "1"}, + {name: "round-robin mode is a valid service name", input: "rr"}, + {name: "nearest mode is a valid service name", input: "nearest"}, + {name: "machine namespace prefix", input: "m-service"}, + {name: "short hexadecimal name", input: strings.Repeat("a", 31)}, + {name: "long hexadecimal name", input: strings.Repeat("a", 33)}, + {name: "non-hexadecimal 32 characters", input: strings.Repeat("g", 32)}, + {name: "too long", input: strings.Repeat("a", 64), wantErr: "must be 1-63 characters"}, + {name: "uppercase", input: "WEB1", wantErr: "lowercase letters"}, + {name: "leading hyphen", input: "-web", wantErr: "starting and ending"}, + {name: "trailing hyphen", input: "web-", wantErr: "starting and ending"}, + {name: "hyphen only", input: "-", wantErr: "starting and ending"}, + {name: "underscore", input: "web_1", wantErr: "hyphens only"}, + {name: "dot", input: "web.example", wantErr: "hyphens only"}, + {name: "slash", input: "web/api", wantErr: "hyphens only"}, + {name: "space", input: "web 1", wantErr: "hyphens only"}, + {name: "leading whitespace", input: " web", wantErr: "hyphens only"}, + {name: "trailing whitespace", input: "web ", wantErr: "hyphens only"}, + {name: "tab", input: "web\t1", wantErr: "hyphens only"}, + {name: "newline", input: "web\n", wantErr: "hyphens only"}, + {name: "non-ASCII", input: "wéb", wantErr: "lowercase letters"}, + {name: "machine DNS namespace", input: "m", wantErr: "reserved for the machine DNS namespace"}, + {name: "service ID", input: "c337f00600de51ef4375c9a9a267dba5", wantErr: "service ID format"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + spec := ServiceSpec{Name: tt.input, Container: ContainerSpec{Image: "nginx:latest"}} + err := spec.Validate() + if tt.wantErr != "" { + require.ErrorContains(t, err, tt.wantErr) + } else { + require.NoError(t, err) + } + }) + } +} + func TestServiceSpec_Validate_CaddyAndPorts(t *testing.T) { tests := []struct { name string