diff --git a/pkg/tui/commands/config.go b/pkg/tui/commands/config.go index 7e3b0d4..83a3d90 100644 --- a/pkg/tui/commands/config.go +++ b/pkg/tui/commands/config.go @@ -7,6 +7,7 @@ import ( "strings" "tusshi/pkg/config" + "tusshi/pkg/validation" ) const tabAll = "All" @@ -27,6 +28,11 @@ func AddConfig(mgr *config.Manager, parts []string) func(Context) { targetPath = filepath.Join(filepath.Dir(mgr.PrimaryPath), arg) } + if err := validation.ValidateConfigName(filepath.Base(targetPath)); err != nil { + ctx.SetError("Invalid config name: " + err.Error()) + return + } + if err := mgr.AddConfigFile(targetPath); err != nil { ctx.SetError("Add config error: " + err.Error()) } else { @@ -83,6 +89,11 @@ func RenameConfig(mgr *config.Manager, parts []string) func(Context) { newPath = filepath.Join(filepath.Dir(mgr.PrimaryPath), newName) } + if err := validation.ValidateConfigName(filepath.Base(newPath)); err != nil { + ctx.SetError("Invalid config name: " + err.Error()) + return + } + if err := mgr.RenameConfigFile(oldPath, newPath); err != nil { ctx.SetError("Rename config error: " + err.Error()) } else { diff --git a/pkg/tui/forms.go b/pkg/tui/forms.go index 7c59fb9..3903f6a 100644 --- a/pkg/tui/forms.go +++ b/pkg/tui/forms.go @@ -1,10 +1,10 @@ package tui import ( - "errors" "path/filepath" "tusshi/pkg/config" + "tusshi/pkg/validation" "github.com/charmbracelet/huh" ) @@ -65,12 +65,7 @@ func (m *Model) BuildHostForm(defaultFile string) *huh.Form { Description("What you will type to connect (e.g. prod-web-01)"). Placeholder("my-server"). Value(&m.FormHost.Alias). - Validate(func(str string) error { - if str == "" { - return errors.New("alias is required") - } - return nil - }), + Validate(validation.ValidateAlias), huh.NewInput(). Title("Server Address / HostName"). Description("Domain or IP address of the target server"). diff --git a/pkg/tui/tui_test.go b/pkg/tui/tui_test.go index 2a8b0b4..53ab153 100644 --- a/pkg/tui/tui_test.go +++ b/pkg/tui/tui_test.go @@ -110,6 +110,12 @@ func TestTUIConfigCommands(t *testing.T) { subPath := filepath.Join(tmpDir, "sub-config-tui") renamedPath := filepath.Join(tmpDir, "renamed-config-tui") + t.Run("add-config validation error", func(t *testing.T) { + m.ErrorText = "" + _, _ = m.executeCommand("add-config invalid*config") + assert.Contains(t, m.ErrorText, "Invalid config name") + }) + t.Run("add-config", func(t *testing.T) { _, _ = m.executeCommand("add-config " + subPath) assert.FileExists(t, subPath) @@ -117,6 +123,12 @@ func TestTUIConfigCommands(t *testing.T) { assert.Equal(t, subPath, m.ActiveTab) }) + t.Run("rename-config validation error", func(t *testing.T) { + m.ErrorText = "" + _, _ = m.executeCommand("rename-config invalid*rename") + assert.Contains(t, m.ErrorText, "Invalid config name") + }) + t.Run("rename-config", func(t *testing.T) { _, _ = m.executeCommand("rename-config " + renamedPath) assert.FileExists(t, renamedPath) diff --git a/pkg/validation/validate_alias_test.go b/pkg/validation/validate_alias_test.go new file mode 100644 index 0000000..8cdf36c --- /dev/null +++ b/pkg/validation/validate_alias_test.go @@ -0,0 +1,114 @@ +// Package validation contains testing suites for input validation functions. +package validation + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +// TestValidateAlias checks the SSH host connection alias validation rules. +func TestValidateAlias(t *testing.T) { + tests := []struct { + name string + alias string + wantErr bool + errMsg string + }{ + { + name: "empty alias", + alias: "", + wantErr: true, + errMsg: "alias is required", + }, + { + name: "alias with spaces", + alias: "my server", + wantErr: true, + errMsg: "alias cannot contain spaces", + }, + { + name: "alias with multiple spaces", + alias: "my server", + wantErr: true, + errMsg: "alias cannot contain spaces", + }, + { + name: "valid alias", + alias: "my-server-01", + wantErr: false, + }, + { + name: "forbidden less than", + alias: "server<01", + wantErr: true, + errMsg: "alias cannot contain forbidden characters", + }, + { + name: "forbidden greater than", + alias: "server>01", + wantErr: true, + errMsg: "alias cannot contain forbidden characters", + }, + { + name: "forbidden colon", + alias: "server:01", + wantErr: true, + errMsg: "alias cannot contain forbidden characters", + }, + { + name: "forbidden double quote", + alias: `server"01`, + wantErr: true, + errMsg: "alias cannot contain forbidden characters", + }, + { + name: "forbidden forward slash", + alias: "server/01", + wantErr: true, + errMsg: "alias cannot contain forbidden characters", + }, + { + name: "forbidden backslash", + alias: `server\01`, + wantErr: true, + errMsg: "alias cannot contain forbidden characters", + }, + { + name: "forbidden pipe", + alias: "server|01", + wantErr: true, + errMsg: "alias cannot contain forbidden characters", + }, + { + name: "forbidden question mark", + alias: "server?01", + wantErr: true, + errMsg: "alias cannot contain forbidden characters", + }, + { + name: "forbidden asterisk", + alias: "server*01", + wantErr: true, + errMsg: "alias cannot contain forbidden characters", + }, + { + name: "forbidden dollar", + alias: "server$01", + wantErr: true, + errMsg: "alias cannot contain forbidden characters", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := ValidateAlias(tt.alias) + if tt.wantErr { + assert.Error(t, err) + assert.Contains(t, err.Error(), tt.errMsg) + } else { + assert.NoError(t, err) + } + }) + } +} diff --git a/pkg/validation/validate_config_test.go b/pkg/validation/validate_config_test.go new file mode 100644 index 0000000..5e97051 --- /dev/null +++ b/pkg/validation/validate_config_test.go @@ -0,0 +1,119 @@ +// Package validation contains testing suites for input validation functions. +package validation + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +// TestValidateConfigName checks the SSH config file name validation rules. +func TestValidateConfigName(t *testing.T) { + tests := []struct { + name string + config string + wantErr bool + errMsg string + }{ + { + name: "empty config name", + config: "", + wantErr: true, + errMsg: "config name is required", + }, + { + name: "config name with spaces", + config: "my config", + wantErr: true, + errMsg: "config name cannot contain spaces", + }, + { + name: "config name with multiple spaces", + config: "my config", + wantErr: true, + errMsg: "config name cannot contain spaces", + }, + { + name: "valid config name", + config: "my-config", + wantErr: false, + }, + { + name: "valid config name with extension", + config: "config.txt", + wantErr: false, + }, + { + name: "forbidden less than", + config: "config<01", + wantErr: true, + errMsg: "config name cannot contain forbidden characters", + }, + { + name: "forbidden greater than", + config: "config>01", + wantErr: true, + errMsg: "config name cannot contain forbidden characters", + }, + { + name: "forbidden colon", + config: "config:01", + wantErr: true, + errMsg: "config name cannot contain forbidden characters", + }, + { + name: "forbidden double quote", + config: `config"01`, + wantErr: true, + errMsg: "config name cannot contain forbidden characters", + }, + { + name: "forbidden forward slash", + config: "config/01", + wantErr: true, + errMsg: "config name cannot contain forbidden characters", + }, + { + name: "forbidden backslash", + config: `config\01`, + wantErr: true, + errMsg: "config name cannot contain forbidden characters", + }, + { + name: "forbidden pipe", + config: "config|01", + wantErr: true, + errMsg: "config name cannot contain forbidden characters", + }, + { + name: "forbidden question mark", + config: "config?01", + wantErr: true, + errMsg: "config name cannot contain forbidden characters", + }, + { + name: "forbidden asterisk", + config: "config*01", + wantErr: true, + errMsg: "config name cannot contain forbidden characters", + }, + { + name: "forbidden hash", + config: "config#01", + wantErr: true, + errMsg: "config name cannot contain forbidden characters", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := ValidateConfigName(tt.config) + if tt.wantErr { + assert.Error(t, err) + assert.Contains(t, err.Error(), tt.errMsg) + } else { + assert.NoError(t, err) + } + }) + } +} diff --git a/pkg/validation/validation.go b/pkg/validation/validation.go new file mode 100644 index 0000000..9b0b2b8 --- /dev/null +++ b/pkg/validation/validation.go @@ -0,0 +1,40 @@ +// Package validation provides validation logic for connection aliases and config names. +package validation + +import ( + "errors" + "fmt" + "regexp" + "strings" +) + +var forbiddenAliasChars = regexp.MustCompile(`[<>:"/\\|?*$]`) +var forbiddenConfigNameChars = regexp.MustCompile(`[<>:"/\\|?*#]`) + +// ValidateAlias checks if the alias matches the validation rules. +func ValidateAlias(str string) error { + if str == "" { + return errors.New("alias is required") + } + if len(strings.Split(str, " ")) > 1 { + return errors.New("alias cannot contain spaces") + } + if forbiddenAliasChars.MatchString(str) { + return fmt.Errorf("alias cannot contain forbidden characters (%v)", strings.Split(forbiddenAliasChars.String(), "")) + } + return nil +} + +// ValidateConfigName checks if the configuration name matches the validation rules. +func ValidateConfigName(str string) error { + if str == "" { + return errors.New("config name is required") + } + if len(strings.Split(str, " ")) > 1 { + return errors.New("config name cannot contain spaces") + } + if forbiddenConfigNameChars.MatchString(str) { + return fmt.Errorf("config name cannot contain forbidden characters (%v)", strings.Split(forbiddenConfigNameChars.String(), "")) + } + return nil +}