From f05e8e8e8426b8f4ffee5c46ea70570d3741653a Mon Sep 17 00:00:00 2001 From: mirivlad Date: Sat, 5 Sep 2026 19:10:02 +0800 Subject: [PATCH] feat: unify route workflows and forwarding UX --- cmd/add.go | 79 ++++-- cmd/connect.go | 53 +--- cmd/edit.go | 148 +++++++---- cmd/extra.go | 70 +++-- cmd/route.go | 100 +++---- cmd/runtime_helpers.go | 27 ++ cmd/tui.go | 87 +++--- cmd/tunnel.go | 18 +- internal/model/server.go | 1 + internal/ssh/command.go | 11 +- internal/tui/app.go | 20 ++ internal/tui/form.go | 378 +++++++++++++++++++++++---- internal/tui/form_validation_test.go | 4 +- internal/tui/forward.go | 204 ++++++++------- internal/tui/forward_test.go | 36 ++- internal/tunnel/manager.go | 26 +- 16 files changed, 846 insertions(+), 416 deletions(-) create mode 100644 cmd/runtime_helpers.go diff --git a/cmd/add.go b/cmd/add.go index 6645cc8..3d3a538 100644 --- a/cmd/add.go +++ b/cmd/add.go @@ -21,6 +21,7 @@ var addFlags struct { authMethod string identityFile string proxyJump string + route string groupName string displayName string notes string @@ -50,6 +51,14 @@ func addInteractive() error { } func addNonInteractive(alias string) error { + routeSpec := strings.TrimSpace(addFlags.route) + if routeSpec == "" { + routeSpec = strings.TrimSpace(addFlags.proxyJump) + } + route, err := parseRouteSpec(routeSpec) + if err != nil { + return fmt.Errorf("route: %w", err) + } server := &model.Server{ Alias: alias, DisplayName: addFlags.displayName, @@ -58,7 +67,8 @@ func addNonInteractive(alias string) error { User: addFlags.user, AuthMethod: model.AuthMethod(addFlags.authMethod), IdentityFile: addFlags.identityFile, - ProxyJump: addFlags.proxyJump, + Route: route, + ProxyJump: route.ProxyJumpString(), GroupName: addFlags.groupName, Notes: addFlags.notes, StartupCommand: addFlags.startup, @@ -78,13 +88,30 @@ func addNonInteractive(alias string) error { } func saveServerWithOptionalSecret(server *model.Server) error { - // Handle password/passphrase auth — request interactively, never via argv - if server.AuthMethod == model.AuthPassword || server.AuthMethod == model.AuthKeyPassphrase { + if len(server.Route.Hops) == 0 && strings.TrimSpace(server.ProxyJump) != "" { + route, err := parseRouteSpec(server.ProxyJump) + if err != nil { + return fmt.Errorf("route: %w", err) + } + server.Route = route + } + server.ProxyJump = server.Route.ProxyJumpString() + if err := model.ValidateServerBasics(server); err != nil { + return err + } + + if addFlags.tags != "" { + server.Tags = strings.Split(addFlags.tags, ",") + } + + var secret []byte + var v = getOrCreateVault() + needsSecret := server.AuthMethod == model.AuthPassword || server.AuthMethod == model.AuthKeyPassphrase + if needsSecret { secretType := "password" if server.AuthMethod == model.AuthKeyPassphrase { secretType = "passphrase" } - fmt.Printf("Enter %s (will be stored in vault, input hidden): ", secretType) password, err := term.ReadPassword(int(syscall.Stdin)) fmt.Println() @@ -94,40 +121,41 @@ func saveServerWithOptionalSecret(server *model.Server) error { if len(password) == 0 { return fmt.Errorf("%s cannot be empty", secretType) } - - v := getOrCreateVault() + secret = password + defer func() { + for i := range secret { + secret[i] = 0 + } + }() if err := unlockVaultForCommand(v); err != nil { return err } - - vaultKey := fmt.Sprintf("server:%s:ssh_password", server.Alias) - vaultType := "ssh_password" - if server.AuthMethod == model.AuthKeyPassphrase { - vaultKey = fmt.Sprintf("server:%s:key_passphrase", server.Alias) - vaultType = "key_passphrase" - } - - if err := v.Put(vaultKey, vaultType, password); err != nil { - return fmt.Errorf("store %s in vault: %w", secretType, err) - } - if err := v.Save(); err != nil { - return fmt.Errorf("save vault: %w", err) - } } if err := appDB.CreateServer(server); err != nil { return fmt.Errorf("create server: %w", err) } + rollbackDB := func() { _ = appDB.DeleteServer(server.Alias) } - if addFlags.tags != "" { - server.Tags = strings.Split(addFlags.tags, ",") - } if len(server.Tags) > 0 { if err := appDB.SetServerTags(server.ID, server.Tags); err != nil { + rollbackDB() return fmt.Errorf("set tags: %w", err) } } + if needsSecret { + if err := syncServerSecrets(v, "", server, string(secret)); err != nil { + rollbackDB() + return fmt.Errorf("store secret in vault: %w", err) + } + if err := v.Save(); err != nil { + cleanupServerSecretsForServer(v, server) + rollbackDB() + return fmt.Errorf("save vault: %w", err) + } + } + fmt.Println("Saved.") return nil } @@ -171,7 +199,7 @@ func promptServerForAdd(in io.Reader, out io.Writer) (*model.Server, error) { if err != nil { return nil, err } - proxyJump, err := promptOptional(reader, out, "ProxyJump", "") + proxyJump, err := promptOptional(reader, out, "Route / ProxyJump (profile: or raw:)", "") if err != nil { return nil, err } @@ -253,7 +281,8 @@ func init() { addCmd.Flags().StringVar(&addFlags.user, "user", "", "SSH username") addCmd.Flags().StringVar(&addFlags.authMethod, "auth", "key", "Auth method: password, key, key_passphrase, agent") addCmd.Flags().StringVar(&addFlags.identityFile, "identity-file", "", "Path to SSH private key") - addCmd.Flags().StringVar(&addFlags.proxyJump, "proxy-jump", "", "ProxyJump host") + addCmd.Flags().StringVar(&addFlags.route, "route", "", "Route hops: profile:, raw:, comma-separated") + addCmd.Flags().StringVar(&addFlags.proxyJump, "proxy-jump", "", "Compatibility alias for --route") addCmd.Flags().StringVar(&addFlags.groupName, "group", "", "Server group") addCmd.Flags().StringVar(&addFlags.displayName, "display-name", "", "Display name") addCmd.Flags().StringVar(&addFlags.notes, "notes", "", "Notes") diff --git a/cmd/connect.go b/cmd/connect.go index 97c0776..d4e0988 100644 --- a/cmd/connect.go +++ b/cmd/connect.go @@ -19,33 +19,9 @@ var connectCmd = &cobra.Command{ if err != nil { return fmt.Errorf("server not found: %s", alias) } - - v := getOrCreateVault() - vaultFunc := func(serverAlias string, secretType string) (string, error) { - if !v.IsUnlocked() { - return "", fmt.Errorf("%s", vaultLockedProcessMessage()) - } - key := fmt.Sprintf("server:%s:%s", serverAlias, secretType) - data, err := v.Get(key) - if err != nil { - return "", err - } - return string(data), nil - } - - if err := ssh.Connect(cfg, &model.Server{ - Alias: server.Alias, - Host: server.Host, - Port: server.Port, - User: server.User, - AuthMethod: server.AuthMethod, - IdentityFile: server.IdentityFile, - ProxyJump: server.ProxyJump, - Route: server.Route, - }, vaultFunc); err != nil { + if err := ssh.ConnectResolved(cfg, server, dbProfileResolver, serverVaultFunc(server)); err != nil { return err } - appDB.UpdateLastConnected(alias) return nil }, @@ -61,31 +37,7 @@ var testCmd = &cobra.Command{ if err != nil { return fmt.Errorf("server not found: %s", alias) } - - v := getOrCreateVault() - vaultFunc := func(serverAlias string, secretType string) (string, error) { - if !v.IsUnlocked() { - return "", fmt.Errorf("%s", vaultLockedProcessMessage()) - } - key := fmt.Sprintf("server:%s:%s", serverAlias, secretType) - data, err := v.Get(key) - if err != nil { - return "", err - } - return string(data), nil - } - - ok, testErr := ssh.Test(cfg, &model.Server{ - Alias: server.Alias, - Host: server.Host, - Port: server.Port, - User: server.User, - AuthMethod: server.AuthMethod, - IdentityFile: server.IdentityFile, - ProxyJump: server.ProxyJump, - Route: server.Route, - }, vaultFunc) - + ok, testErr := ssh.TestResolved(cfg, server, dbProfileResolver, serverVaultFunc(server)) if ok { fmt.Println("Connection OK.") appDB.UpdateTestResult(alias, model.TestOK, "") @@ -93,7 +45,6 @@ var testCmd = &cobra.Command{ fmt.Printf("Connection failed:\n%s\n", testErr) appDB.UpdateTestResult(alias, model.TestFailed, testErr) } - return nil }, } diff --git a/cmd/edit.go b/cmd/edit.go index 527e103..ecdf583 100644 --- a/cmd/edit.go +++ b/cmd/edit.go @@ -20,99 +20,138 @@ var editCmd = &cobra.Command{ if err != nil { return fmt.Errorf("server not found: %s", alias) } + original := cloneServer(server) - oldAuthMethod := server.AuthMethod - - if parsedHost != "" { + if cmd.Flags().Changed("host") { server.Host = parsedHost } - if parsedPort != 0 { + if cmd.Flags().Changed("port") { + if parsedPort < 1 || parsedPort > 65535 { + return fmt.Errorf("port must be between 1 and 65535") + } server.Port = parsedPort } - if parsedUser != "" { + if cmd.Flags().Changed("user") { server.User = parsedUser } - if parsedAuth != "" { + authChanged := cmd.Flags().Changed("auth") + if authChanged { server.AuthMethod = model.AuthMethod(parsedAuth) } - if parsedIdentity != "" { + if cmd.Flags().Changed("identity-file") { server.IdentityFile = parsedIdentity } - if parsedProxyJump != "" { - server.ProxyJump = parsedProxyJump - } - if parsedGroup != "" { + if cmd.Flags().Changed("group") { server.GroupName = parsedGroup } - if parsedDisplayName != "" { + if cmd.Flags().Changed("display-name") { server.DisplayName = parsedDisplayName } - if parsedNotes != "" { + if cmd.Flags().Changed("notes") { server.Notes = parsedNotes } - if parsedStartup != "" { + if cmd.Flags().Changed("startup-command") { server.StartupCommand = parsedStartup } + routeChanged := cmd.Flags().Changed("route") || cmd.Flags().Changed("proxy-jump") + if routeChanged { + if cmd.Flags().Changed("route") && cmd.Flags().Changed("proxy-jump") { + return fmt.Errorf("use either --route or --proxy-jump, not both") + } + spec := parsedRoute + if cmd.Flags().Changed("proxy-jump") { + spec = parsedProxyJump + } + route, err := parseRouteSpec(spec) + if err != nil { + return fmt.Errorf("route: %w", err) + } + server.Route = route + server.ProxyJump = route.ProxyJumpString() + } tagsChanged := cmd.Flags().Changed("tags") if tagsChanged { - server.Tags = strings.Split(parsedTags, ",") + if strings.TrimSpace(parsedTags) == "" { + server.Tags = nil + } else { + server.Tags = strings.Split(parsedTags, ",") + } + } + if err := model.ValidateServerBasics(server); err != nil { + return err } - if parsedAuth != "" && oldAuthMethod != server.AuthMethod { - v := getOrCreateVault() - if v.IsUnlocked() { - var secret string - if server.AuthMethod == model.AuthPassword { - fmt.Print("Enter new password (stored in vault, input hidden): ") - pw, err := term.ReadPassword(int(syscall.Stdin)) - fmt.Println() - if err != nil { - return fmt.Errorf("read password: %w", err) - } - if len(pw) > 0 { - secret = string(pw) - } - } else if server.AuthMethod == model.AuthKeyPassphrase { - fmt.Print("Enter key passphrase (stored in vault, input hidden): ") - pw, err := term.ReadPassword(int(syscall.Stdin)) - fmt.Println() - if err != nil { - return fmt.Errorf("read passphrase: %w", err) - } - if len(pw) > 0 { - secret = string(pw) - } + var secret []byte + v := getOrCreateVault() + if authChanged { + if err := unlockVaultForCommand(v); err != nil { + return err + } + switch server.AuthMethod { + case model.AuthPassword, model.AuthKeyPassphrase: + label := "password" + if server.AuthMethod == model.AuthKeyPassphrase { + label = "key passphrase" } - - if err := syncServerSecrets(v, alias, server, secret); err != nil { - return fmt.Errorf("sync vault secrets: %w", err) + fmt.Printf("Enter new %s (stored in vault, input hidden): ", label) + secret, err = term.ReadPassword(int(syscall.Stdin)) + fmt.Println() + if err != nil { + return fmt.Errorf("read %s: %w", label, err) } - if err := v.Save(); err != nil { - return fmt.Errorf("save vault: %w", err) + if len(secret) == 0 { + return fmt.Errorf("%s cannot be empty", label) } + defer func() { + for i := range secret { + secret[i] = 0 + } + }() } } - if err := appDB.UpdateServer(server); err != nil { + if err := appDB.UpdateServerByAlias(alias, server); err != nil { return fmt.Errorf("update server: %w", err) } + rollback := func() { + _ = appDB.UpdateServerByAlias(server.Alias, original) + _ = appDB.SetServerTags(original.ID, original.Tags) + } if tagsChanged { if err := appDB.SetServerTags(server.ID, server.Tags); err != nil { + rollback() return fmt.Errorf("set tags: %w", err) } } - + if authChanged { + if err := syncServerSecrets(v, alias, server, string(secret)); err != nil { + rollback() + return fmt.Errorf("sync vault secrets: %w", err) + } + if err := v.Save(); err != nil { + rollback() + return fmt.Errorf("save vault: %w", err) + } + } fmt.Println("Saved.") return nil }, } +func cloneServer(server *model.Server) *model.Server { + copyServer := *server + copyServer.Route.Hops = append([]model.RouteHop(nil), server.Route.Hops...) + copyServer.Tags = append([]string(nil), server.Tags...) + return ©Server +} + var ( parsedHost string parsedPort int parsedUser string parsedAuth string parsedIdentity string + parsedRoute string parsedProxyJump string parsedGroup string parsedDisplayName string @@ -124,13 +163,14 @@ var ( func init() { editCmd.Flags().StringVar(&parsedHost, "host", "", "Server hostname or IP") editCmd.Flags().IntVar(&parsedPort, "port", 0, "SSH port") - editCmd.Flags().StringVar(&parsedUser, "user", "", "SSH username") + editCmd.Flags().StringVar(&parsedUser, "user", "", "SSH username; empty lets OpenSSH choose") editCmd.Flags().StringVar(&parsedAuth, "auth", "", "Auth method") - editCmd.Flags().StringVar(&parsedIdentity, "identity-file", "", "Path to SSH private key") - editCmd.Flags().StringVar(&parsedProxyJump, "proxy-jump", "", "ProxyJump host") - editCmd.Flags().StringVar(&parsedGroup, "group", "", "Server group") - editCmd.Flags().StringVar(&parsedDisplayName, "display-name", "", "Display name") - editCmd.Flags().StringVar(&parsedNotes, "notes", "", "Notes") - editCmd.Flags().StringVar(&parsedStartup, "startup-command", "", "Command to run after connecting") - editCmd.Flags().StringVar(&parsedTags, "tags", "", "Comma-separated tags") + editCmd.Flags().StringVar(&parsedIdentity, "identity-file", "", "Path to SSH private key; empty clears it") + editCmd.Flags().StringVar(&parsedRoute, "route", "", "Route hops: profile:, raw:, comma-separated; empty means direct") + editCmd.Flags().StringVar(&parsedProxyJump, "proxy-jump", "", "Compatibility alias for --route") + editCmd.Flags().StringVar(&parsedGroup, "group", "", "Server group; empty clears it") + editCmd.Flags().StringVar(&parsedDisplayName, "display-name", "", "Display name; empty clears it") + editCmd.Flags().StringVar(&parsedNotes, "notes", "", "Notes; empty clears them") + editCmd.Flags().StringVar(&parsedStartup, "startup-command", "", "Startup command; empty clears it") + editCmd.Flags().StringVar(&parsedTags, "tags", "", "Comma-separated tags; empty clears all") } diff --git a/cmd/extra.go b/cmd/extra.go index 55278ff..e425bae 100644 --- a/cmd/extra.go +++ b/cmd/extra.go @@ -43,7 +43,6 @@ func importServersFromSSHConfig(report func(format string, args ...interface{})) if err != nil { return 0, fmt.Errorf("import: %w", err) } - if len(servers) == 0 { if report != nil { report("No servers found in ~/.ssh/config") @@ -51,27 +50,63 @@ func importServersFromSSHConfig(report func(format string, args ...interface{})) return 0, nil } + // First pass creates every profile without routes. That makes ProxyJump + // aliases resolvable to stable sshkeeper IDs in the second pass, regardless + // of declaration order in ~/.ssh/config. + type pendingRoute struct { + server *model.Server + spec string + } + pending := make([]pendingRoute, 0, len(servers)) imported := 0 - for _, s := range servers { - existing, _ := appDB.GetServer(s.Alias) - if existing != nil { + for _, server := range servers { + if existing, _ := appDB.GetServer(server.Alias); existing != nil { if report != nil { - report(" skip (exists): %s", s.Alias) + report(" skip (exists): %s", server.Alias) } continue } - if err := appDB.CreateServer(s); err != nil { + spec := strings.TrimSpace(server.ProxyJump) + server.ProxyJump = "" + server.Route = model.Route{} + if err := appDB.CreateServer(server); err != nil { if report != nil { - report(" error: %s: %v", s.Alias, err) + report(" error: %s: %v", server.Alias, err) } continue } - if report != nil { - report(" imported: %s (%s@%s:%d)", s.Alias, s.User, s.Host, s.Port) - } + pending = append(pending, pendingRoute{server: server, spec: spec}) imported++ } + for _, item := range pending { + if item.spec != "" { + route, err := parseRouteSpec(item.spec) + if err != nil { + if report != nil { + report(" warning: %s imported direct; route %q could not be parsed: %v", item.server.Alias, item.spec, err) + } + continue + } + item.server.Route = route + item.server.ProxyJump = route.ProxyJumpString() + if err := appDB.UpdateServer(item.server); err != nil { + item.server.Route = model.Route{} + item.server.ProxyJump = "" + if report != nil { + report(" warning: %s imported direct; route could not be saved: %v", item.server.Alias, err) + } + continue + } + } + if report != nil { + routeSuffix := "" + if len(item.server.Route.Hops) > 0 { + routeSuffix = " via " + model.FormatRouteSpec(item.server.Route) + } + report(" imported: %s (%s@%s:%d)%s", item.server.Alias, item.server.User, item.server.Host, item.server.Port, routeSuffix) + } + } return imported, nil } @@ -101,18 +136,5 @@ var runCmd = &cobra.Command{ } func runCommandOnServer(server *model.Server, command string) error { - return ssh.RunCommand(cfg, server, commandVaultFunc, command) -} - -func commandVaultFunc(serverAlias string, secretType string) (string, error) { - v := getOrCreateVault() - if !v.IsUnlocked() { - return "", fmt.Errorf("%s", vaultLockedProcessMessage()) - } - vaultKey := fmt.Sprintf("server:%s:%s", serverAlias, secretType) - data, err := v.Get(vaultKey) - if err != nil { - return "", err - } - return string(data), nil + return ssh.RunCommandResolved(cfg, server, dbProfileResolver, serverVaultFunc(server), command) } diff --git a/cmd/route.go b/cmd/route.go index fd892af..5552d81 100644 --- a/cmd/route.go +++ b/cmd/route.go @@ -8,11 +8,9 @@ import ( "github.com/spf13/cobra" ) -// --- Route commands --- - var routeCmd = &cobra.Command{ Use: "route", - Short: "Manage server routes (ProxyJump)", + Short: "Manage server routes (bastions / ProxyJump)", } var routeShowCmd = &cobra.Command{ @@ -20,30 +18,29 @@ var routeShowCmd = &cobra.Command{ Short: "Show route for a server", Args: cobra.ExactArgs(1), RunE: func(cmd *cobra.Command, args []string) error { - alias := args[0] - server, err := appDB.GetServer(alias) + server, err := appDB.GetServer(args[0]) if err != nil { - return fmt.Errorf("server not found: %s", alias) + return fmt.Errorf("server not found: %s", args[0]) } - target := fmt.Sprintf("%s@%s:%d", server.User, server.Host, server.Port) - if len(server.Route.Hops) > 0 { - fmt.Printf("Route: %s\n", server.Route.DisplaySummary(target)) - fmt.Printf("Mode: %s\n", server.Route.RouteMode()) - fmt.Printf("ProxyJump: %s\n", server.Route.ProxyJumpString()) - if server.Route.HasProfileLinks() { - fmt.Println("Hops:") - for _, h := range server.Route.Hops { - if h.IsProfile { - fmt.Printf(" - %s (profile)\n", h.Alias) - } else { - fmt.Printf(" - %s (raw)\n", h.Raw) - } - } - } - } else if server.ProxyJump != "" { - fmt.Printf("ProxyJump: %s\n", server.ProxyJump) - } else { + target := server.Host + if server.User != "" { + target = server.User + "@" + server.Host + } + target = fmt.Sprintf("%s:%d", target, server.Port) + if len(server.Route.Hops) == 0 { fmt.Println("Direct connection (no route)") + return nil + } + fmt.Printf("Route: %s\n", server.Route.DisplaySummary(target)) + fmt.Printf("Mode: %s\n", server.Route.RouteMode()) + fmt.Printf("Spec: %s\n", model.FormatRouteSpec(server.Route)) + fmt.Println("Hops:") + for index, hop := range server.Route.Hops { + if hop.Profile() { + fmt.Printf(" %d. %s (sshkeeper profile #%d)\n", index+1, hop.Alias, hop.ServerID) + } else { + fmt.Printf(" %d. %s (raw OpenSSH target)\n", index+1, hop.Raw) + } } return nil }, @@ -52,47 +49,38 @@ var routeShowCmd = &cobra.Command{ var routeSetCmd = &cobra.Command{ Use: "set ", Short: "Set route for a server", - Args: cobra.MinimumNArgs(1), + Long: `Set an ordered route. Known aliases are resolved to stable sshkeeper profile IDs. +Use profile: to require a profile reference and raw: to force a literal OpenSSH target.`, + Args: cobra.ExactArgs(1), RunE: func(cmd *cobra.Command, args []string) error { - alias := args[0] - server, err := appDB.GetServer(alias) + server, err := appDB.GetServer(args[0]) if err != nil { - return fmt.Errorf("server not found: %s", alias) + return fmt.Errorf("server not found: %s", args[0]) } - mode, _ := cmd.Flags().GetString("mode") jumps, _ := cmd.Flags().GetString("jumps") - - if mode == "clear" || jumps == "" { + mode = strings.ToLower(strings.TrimSpace(mode)) + if mode == "clear" || mode == "direct" { server.Route = model.Route{} server.ProxyJump = "" } else { - parts := strings.Split(jumps, ",") - hops := make([]model.RouteHop, 0, len(parts)) - for _, p := range parts { - p = strings.TrimSpace(p) - if p == "" { - continue - } - if strings.Contains(p, "@") || strings.Contains(p, ":") { - hops = append(hops, model.RouteHop{Raw: p, IsProfile: false}) - } else { - hops = append(hops, model.RouteHop{Alias: p, IsProfile: true}) - } + if strings.TrimSpace(jumps) == "" { + return fmt.Errorf("--jumps is required unless --mode=direct/clear") } - server.Route = model.Route{Hops: hops} - server.ProxyJump = server.Route.ProxyJumpString() + route, err := parseRouteSpec(jumps) + if err != nil { + return err + } + server.Route = route + server.ProxyJump = route.ProxyJumpString() } - if err := appDB.UpdateServer(server); err != nil { return fmt.Errorf("update route: %w", err) } - - target := fmt.Sprintf("%s@%s:%d", server.User, server.Host, server.Port) - if len(server.Route.Hops) > 0 { - fmt.Printf("✓ Route set: %s\n", server.Route.DisplaySummary(target)) - } else { + if len(server.Route.Hops) == 0 { fmt.Println("✓ Route cleared (direct connection)") + } else { + fmt.Printf("✓ Route set: %s\n", model.FormatRouteSpec(server.Route)) } return nil }, @@ -103,10 +91,9 @@ var routeClearCmd = &cobra.Command{ Short: "Clear route for a server (set direct)", Args: cobra.ExactArgs(1), RunE: func(cmd *cobra.Command, args []string) error { - alias := args[0] - server, err := appDB.GetServer(alias) + server, err := appDB.GetServer(args[0]) if err != nil { - return fmt.Errorf("server not found: %s", alias) + return fmt.Errorf("server not found: %s", args[0]) } server.Route = model.Route{} server.ProxyJump = "" @@ -119,9 +106,8 @@ var routeClearCmd = &cobra.Command{ } func init() { - routeSetCmd.Flags().String("mode", "via", "Route mode: direct, via, chain, or clear") - routeSetCmd.Flags().String("jumps", "", "Comma-separated jump hosts (aliases or raw addresses)") - + routeSetCmd.Flags().String("mode", "via", "Route mode: via, chain, direct, or clear") + routeSetCmd.Flags().String("jumps", "", "Comma-separated hops; use profile: or raw: for explicit type") routeCmd.AddCommand(routeShowCmd) routeCmd.AddCommand(routeSetCmd) routeCmd.AddCommand(routeClearCmd) diff --git a/cmd/runtime_helpers.go b/cmd/runtime_helpers.go new file mode 100644 index 0000000..afe765a --- /dev/null +++ b/cmd/runtime_helpers.go @@ -0,0 +1,27 @@ +package cmd + +import ( + "github.com/mirivlad/sshkeeper/internal/model" + "github.com/mirivlad/sshkeeper/internal/ssh" +) + +func dbProfileResolver(serverID int64) (*model.Server, error) { + return appDB.GetServerByID(serverID) +} + +func parseRouteSpec(spec string) (model.Route, error) { + return model.ParseRouteSpec(spec, appDB.ResolveAlias) +} + +func serverVaultFunc(server *model.Server) ssh.VaultFunc { + return vaultFuncForServer(getOrCreateVault(), server) +} + +func rollbackSavedServer(server, original *model.Server) { + if original != nil { + _ = appDB.UpdateServerByAlias(server.Alias, original) + _ = appDB.SetServerTags(original.ID, original.Tags) + return + } + _ = appDB.DeleteServer(server.Alias) +} diff --git a/cmd/tui.go b/cmd/tui.go index e4ee95d..87b26d0 100644 --- a/cmd/tui.go +++ b/cmd/tui.go @@ -17,19 +17,6 @@ func runTUI() error { return fmt.Errorf("load servers: %w", err) } - vaultFunc := func(sa string, st string) (string, error) { - v := getOrCreateVault() - if !v.IsUnlocked() { - return "", fmt.Errorf("vault is locked") - } - key := fmt.Sprintf("server:%s:%s", sa, st) - data, err := v.Get(key) - if err != nil { - return "", err - } - return string(data), nil - } - tui.ListServers = func() ([]*model.Server, error) { return appDB.ListServers() } @@ -37,56 +24,81 @@ func runTUI() error { return appDB.SearchServers(query) } tui.DeleteServer = func(alias string) error { + server, err := appDB.GetServer(alias) + if err != nil { + return err + } if err := appDB.DeleteServer(alias); err != nil { return err } v := getOrCreateVault() if v.IsUnlocked() { - cleanupServerSecrets(v, alias) + cleanupServerSecretsForServer(v, server) if err := v.Save(); err != nil { + // The vault file is unchanged on save failure; restore the DB profile. + _ = appDB.CreateServer(server) + _ = appDB.SetServerTags(server.ID, server.Tags) return fmt.Errorf("save vault after cleanup: %w", err) } } return nil } tui.TestConnection = func(server *model.Server) (bool, string) { - return ssh.Test(cfg, server, vaultFunc) + return ssh.TestResolved(cfg, server, dbProfileResolver, serverVaultFunc(server)) } tui.TestConnectionWithPassword = func(server *model.Server, password string) (bool, string) { - return ssh.Test(cfg, server, formTestVaultFunc(vaultFunc, server, password)) + base := serverVaultFunc(server) + return ssh.TestResolved(cfg, server, dbProfileResolver, formTestVaultFunc(base, server, password)) } tui.SaveServer = func(server *model.Server, password string, oldAlias string) error { - v := getOrCreateVault() - if v.IsUnlocked() { - if err := syncServerSecrets(v, oldAlias, server, password); err != nil { - return fmt.Errorf("sync vault secrets: %w", err) - } - if err := v.Save(); err != nil { - return fmt.Errorf("save vault: %w", err) - } - } - lookupAlias := server.Alias if oldAlias != "" { lookupAlias = oldAlias } existing, _ := appDB.GetServer(lookupAlias) + var original *model.Server if existing != nil { + original = cloneServer(existing) server.ID = existing.ID if err := appDB.UpdateServerByAlias(existing.Alias, server); err != nil { return err } - return appDB.SetServerTags(existing.ID, server.Tags) + } else { + if err := appDB.CreateServer(server); err != nil { + return err + } } - if err := appDB.CreateServer(server); err != nil { + if err := appDB.SetServerTags(server.ID, server.Tags); err != nil { + if original != nil { + _ = appDB.UpdateServerByAlias(server.Alias, original) + _ = appDB.SetServerTags(original.ID, original.Tags) + } else { + _ = appDB.DeleteServer(server.Alias) + } return err } - return appDB.SetServerTags(server.ID, server.Tags) + + v := getOrCreateVault() + if !v.IsUnlocked() { + return nil + } + if err := syncServerSecrets(v, oldAlias, server, password); err != nil { + rollbackSavedServer(server, original) + return fmt.Errorf("sync vault secrets: %w", err) + } + if err := v.Save(); err != nil { + rollbackSavedServer(server, original) + return fmt.Errorf("save vault: %w", err) + } + return nil } tui.GetGroups = func() ([]string, error) { return appDB.GetGroups() } + tui.ResolveRouteAlias = func(alias string) (int64, bool) { + return appDB.ResolveAlias(alias) + } tui.RenameGroup = func(oldName, newName string) error { return appDB.RenameGroup(oldName, newName) } @@ -123,7 +135,7 @@ func runTUI() error { if err != nil { return "", err } - return ssh.RunCommandOutput(cfg, fresh, vaultFunc, command) + return ssh.RunCommandOutputResolved(cfg, fresh, dbProfileResolver, serverVaultFunc(fresh), command) } tui.ListForwards = func(serverID int64) ([]*model.Forward, error) { return appDB.GetForwards(serverID) @@ -157,7 +169,11 @@ func runTUI() error { if !v.IsUnlocked() { return false } - return v.HasSecret(serverSecretID(alias, secretType)) + server, err := appDB.GetServer(alias) + if err != nil { + return false + } + return hasServerSecret(v, server, secretType) } // Run TUI in a loop — if user requests connect, handle it and restart TUI @@ -185,7 +201,7 @@ func runTUI() error { fmt.Printf("Connecting to %s@%s:%d...\n", fresh.User, fresh.Host, fresh.Port) - if err := ssh.Connect(cfg, fresh, vaultFunc); err != nil { + if err := ssh.ConnectResolved(cfg, fresh, dbProfileResolver, serverVaultFunc(fresh)); err != nil { fmt.Fprintf(os.Stderr, "Connection error: %v\n", err) } else { fmt.Println("Connection closed.") @@ -210,7 +226,7 @@ func runTUI() error { continue } fmt.Printf("Running template %q on %s...\n", result.TemplateName, fresh.Alias) - if err := ssh.RunCommand(cfg, fresh, vaultFunc, result.Command); err != nil { + if err := ssh.RunCommandResolved(cfg, fresh, dbProfileResolver, serverVaultFunc(fresh), result.Command); err != nil { fmt.Fprintf(os.Stderr, "Command error on %s: %v\n", fresh.Alias, err) } appDB.UpdateLastConnected(fresh.Alias) @@ -279,7 +295,7 @@ func runTUI() error { servers, _ = appDB.ListServers() continue } - state, err := tunnelpkg.Start(cfg, fresh, forwards, forwardOnly) + state, err := tunnelpkg.StartResolved(cfg, fresh, forwards, forwardOnly, dbProfileResolver) if err != nil { fmt.Fprintf(os.Stderr, "Start tunnel: %v\n", err) } else { @@ -295,8 +311,7 @@ func runTUI() error { fmt.Printf("Starting session to %s...\n", fresh.Alias) } - sshArgs := ssh.BuildSSHArgs(fresh, forwards, forwardOnly) - if err := ssh.ConnectWithArgs(cfg, sshArgs, vaultFunc, fresh); err != nil { + if err := ssh.ConnectWithForwardsResolved(cfg, fresh, forwards, forwardOnly, dbProfileResolver, serverVaultFunc(fresh)); err != nil { fmt.Fprintf(os.Stderr, "Tunnel error: %v\n", err) } else { fmt.Println("Tunnel closed.") diff --git a/cmd/tunnel.go b/cmd/tunnel.go index a9534b1..770c57d 100644 --- a/cmd/tunnel.go +++ b/cmd/tunnel.go @@ -34,7 +34,7 @@ var tunnelCmd = &cobra.Command{ if err := validateBackgroundTunnel(server, forwards); err != nil { return err } - state, err := tunnelpkg.Start(cfg, server, forwards, true) + state, err := tunnelpkg.StartResolved(cfg, server, forwards, true, dbProfileResolver) if err != nil { return err } @@ -46,31 +46,17 @@ var tunnelCmd = &cobra.Command{ return fmt.Errorf("no forwards configured for %s", alias) } - v := getOrCreateVault() - vaultFunc := func(serverAlias string, secretType string) (string, error) { - if !v.IsUnlocked() { - return "", fmt.Errorf("%s", vaultLockedProcessMessage()) - } - key := fmt.Sprintf("server:%s:%s", serverAlias, secretType) - data, err := v.Get(key) - if err != nil { - return "", err - } - return string(data), nil - } - if len(forwards) > 0 { fmt.Printf("Starting tunnel to %s with %d forward(s)...\n", alias, len(forwards)) } else { fmt.Printf("Starting session to %s...\n", alias) } - sshArgs := ssh.BuildSSHArgs(server, forwards, forwardsOnly) if forwardsOnly { fmt.Printf("Tunnel mode (ssh -N). Press Ctrl+C to exit.\n") } - return ssh.ConnectWithArgs(cfg, sshArgs, vaultFunc, server) + return ssh.ConnectWithForwardsResolved(cfg, server, forwards, forwardsOnly, dbProfileResolver, serverVaultFunc(server)) }, } diff --git a/internal/model/server.go b/internal/model/server.go index c2cd368..f1f411a 100644 --- a/internal/model/server.go +++ b/internal/model/server.go @@ -245,6 +245,7 @@ type TunnelState struct { Name string `json:"name"` PID int `json:"pid"` ForwardIDs []int64 `json:"forward_ids"` + ConfigPath string `json:"config_path,omitempty"` StartedAt time.Time `json:"started_at"` LastError string `json:"last_error"` } diff --git a/internal/ssh/command.go b/internal/ssh/command.go index faf3a6d..e209e7b 100644 --- a/internal/ssh/command.go +++ b/internal/ssh/command.go @@ -266,7 +266,11 @@ func BuildForwardArgs(forwards []*model.Forward, exitOnForwardFailure bool) []st func BuildSSHArgs(server *model.Server, forwards []*model.Forward, forwardOnly bool) []string { var args []string - args = append(args, "-p", fmt.Sprintf("%d", server.Port)) + port := server.Port + if port == 0 { + port = 22 + } + args = append(args, "-p", fmt.Sprintf("%d", port)) if server.IdentityFile != "" { args = append(args, "-i", server.IdentityFile) @@ -291,7 +295,10 @@ func BuildSSHArgs(server *model.Server, forwards []*model.Forward, forwardOnly b args = append(args, "-N") } - target := fmt.Sprintf("%s@%s", server.User, server.Host) + target := server.Host + if strings.TrimSpace(server.User) != "" { + target = fmt.Sprintf("%s@%s", server.User, server.Host) + } args = append(args, target) return args diff --git a/internal/tui/app.go b/internal/tui/app.go index 8133dd7..669ab1e 100644 --- a/internal/tui/app.go +++ b/internal/tui/app.go @@ -180,6 +180,7 @@ var ( UpdateTestResult func(alias string, status model.TestStatus, testErr string) error HasSecret func(alias string, secretType string) bool GetGroups func() ([]string, error) + ResolveRouteAlias func(alias string) (int64, bool) RenameGroup func(oldName, newName string) error DeleteGroup func(name string) error ListTags func() ([]string, error) @@ -769,12 +770,14 @@ func (m *tuiModel) updateList(msg tea.KeyMsg) (tea.Model, tea.Cmd) { case tea.KeyCtrlA: m.form = newFormModel(m.width, m.height) + m.form.setRouteProfiles(m.servers) m.screen = screenForm return m, nil case tea.KeyCtrlE: if item, ok := m.list.SelectedItem().(serverItem); ok { m.form = newEditFormModel(item.server, m.width, m.height) + m.form.setRouteProfiles(m.servers) m.screen = screenForm } return m, nil @@ -1355,6 +1358,7 @@ func (m *tuiModel) updateActionMenu(msg tea.KeyMsg) (tea.Model, tea.Cmd) { case "route": if item, ok := m.list.SelectedItem().(serverItem); ok { m.form = newEditFormModel(item.server, m.width, m.height) + m.form.setRouteProfiles(m.servers) m.form.focusIdx = 7 m.form.updateFocus() m.screen = screenForm @@ -1372,6 +1376,7 @@ func (m *tuiModel) updateActionMenu(msg tea.KeyMsg) (tea.Model, tea.Cmd) { case "edit": if item, ok := m.list.SelectedItem().(serverItem); ok { m.form = newEditFormModel(item.server, m.width, m.height) + m.form.setRouteProfiles(m.servers) m.screen = screenForm m.actionMenu = nil return m, nil @@ -1445,6 +1450,21 @@ func (m *tuiModel) updateForwardList(msg tea.KeyMsg) (tea.Model, tea.Cmd) { if m.forwardScreen != nil { return m, m.forwardScreen.editSelected() } + case tea.KeySpace: + if m.forwardScreen != nil && m.forwardScreen.selected >= 0 && m.forwardScreen.selected < len(m.forwardScreen.list) { + selected := *m.forwardScreen.list[m.forwardScreen.selected] + selected.Enabled = !selected.Enabled + return m, func() tea.Msg { + if UpdateForward == nil { + return forwardsLoadedMsg{err: fmt.Errorf("forward update is unavailable")} + } + if err := UpdateForward(&selected); err != nil { + return forwardsLoadedMsg{err: err} + } + forwards, err := ListForwards(m.forwardScreen.serverID) + return forwardsLoadedMsg{forwards: forwards, err: err} + } + } case tea.KeyRunes: switch msg.String() { case "a", "A": diff --git a/internal/tui/form.go b/internal/tui/form.go index 2147aa2..14bcd28 100644 --- a/internal/tui/form.go +++ b/internal/tui/form.go @@ -23,6 +23,22 @@ func (i groupItem) Title() string { return i.name } func (i groupItem) Description() string { return "" } func (i groupItem) FilterValue() string { return i.name } +type routeProfileItem struct { + server *model.Server +} + +func (i routeProfileItem) Title() string { return i.server.Alias } +func (i routeProfileItem) Description() string { + target := i.server.Host + if i.server.User != "" { + target = i.server.User + "@" + i.server.Host + } + return fmt.Sprintf("%s:%d", target, i.server.Port) +} +func (i routeProfileItem) FilterValue() string { + return strings.Join([]string{i.server.Alias, i.server.DisplayName, i.server.Host, i.server.User, i.server.GroupName}, " ") +} + func newStringList(values []string, title string, width, height int) list.Model { items := make([]list.Item, len(values)) for i, value := range values { @@ -37,6 +53,28 @@ func newStringList(values []string, title string, width, height int) list.Model return l } +func (fm *formModel) setRouteProfiles(servers []*model.Server) { + fm.routeProfiles = nil + items := []list.Item{} + for _, server := range servers { + if server == nil { + continue + } + if fm.server != nil && ((fm.server.ID > 0 && server.ID == fm.server.ID) || server.Alias == fm.server.Alias) { + continue + } + fm.routeProfiles = append(fm.routeProfiles, server) + items = append(items, routeProfileItem{server: server}) + } + l := list.New(items, list.NewDefaultDelegate(), 44, 14) + l.Title = "Available server profiles" + l.SetShowStatusBar(false) + l.SetShowHelp(false) + l.SetFilteringEnabled(true) + l.Styles.Title = titleStyle + fm.routeList = l +} + // --- Form model --- type formModel struct { @@ -63,6 +101,11 @@ type formModel struct { showGroupList bool authList list.Model showAuthList bool + routeProfiles []*model.Server + routeList list.Model + showRouteList bool + routePane int // 0=current route, 1=available profiles + routeCursor int initial formSnapshot } @@ -81,7 +124,7 @@ func newFormModel(w, h int) *formModel { "User", "Auth Method (password/key/key_passphrase/agent)", "Identity File", - "Route hops (comma-separated, or pick from profiles)", + "Route (direct / ordered bastions)", "Group (type new or pick from list)", "Notes", "Startup Command", @@ -128,7 +171,6 @@ func newFormModel(w, h int) *formModel { fm.groupList = newStringList(groups, "Select group", 30, 8) } } - fm.updateFocus() fm.initial = fm.snapshot() return fm @@ -150,8 +192,8 @@ func placeholderForLabel(label string) string { return "key" case "Identity File": return "~/.ssh/id_ed25519" - case "Route hops (comma-separated, or pick from profiles)": - return "bastion, dmz-gw" + case "Route (direct / ordered bastions)": + return "profile:bastion, raw:user@gw.example" case "Group (type new or pick from list)": return "KP" case "Notes": @@ -177,17 +219,9 @@ func newEditFormModel(s *model.Server, w, h int) *formModel { fm.inputs[5].SetValue(string(s.AuthMethod)) fm.inputs[6].SetValue(s.IdentityFile) - // Populate Route hops + // Store an unambiguous route spec in the editable field. if len(s.Route.Hops) > 0 { - hopStrs := make([]string, len(s.Route.Hops)) - for i, h := range s.Route.Hops { - if h.IsProfile { - hopStrs[i] = h.Alias - } else { - hopStrs[i] = h.Raw - } - } - fm.inputs[7].SetValue(strings.Join(hopStrs, ", ")) + fm.inputs[7].SetValue(model.FormatRouteSpec(s.Route)) } else if s.ProxyJump != "" { fm.inputs[7].SetValue(s.ProxyJump) } @@ -236,6 +270,136 @@ func (fm *formModel) Dirty() bool { return false } +func (fm *formModel) resolveRouteAlias(alias string) (int64, bool) { + alias = strings.TrimSpace(alias) + for _, server := range fm.routeProfiles { + if server != nil && server.Alias == alias && server.ID > 0 { + return server.ID, true + } + } + if ResolveRouteAlias != nil { + return ResolveRouteAlias(alias) + } + return 0, false +} + +func (fm *formModel) parseRouteInput() (model.Route, error) { + return model.ParseRouteSpec(fm.inputs[7].Value(), fm.resolveRouteAlias) +} + +func (fm *formModel) currentRoute() model.Route { + route, err := fm.parseRouteInput() + if err != nil { + return model.Route{} + } + return route +} + +func (fm *formModel) setCurrentRoute(route model.Route) { + fm.inputs[7].SetValue(model.FormatRouteSpec(route)) + if len(route.Hops) == 0 { + fm.routeCursor = 0 + } else if fm.routeCursor >= len(route.Hops) { + fm.routeCursor = len(route.Hops) - 1 + } +} + +func (fm *formModel) routeContainsProfile(route model.Route, serverID int64) bool { + for _, hop := range route.Hops { + if hop.Profile() && hop.ServerID == serverID { + return true + } + } + return false +} + +func (fm *formModel) updateRouteEditor(msg tea.Msg) (tea.Model, tea.Cmd) { + key, ok := msg.(tea.KeyMsg) + if !ok { + if fm.routePane == 1 { + var cmd tea.Cmd + fm.routeList, cmd = fm.routeList.Update(msg) + return fm, cmd + } + return fm, nil + } + if key.Type == tea.KeyEsc { + fm.showRouteList = false + fm.inputs[7].Focus() + return fm, nil + } + if key.Type == tea.KeyTab || key.Type == tea.KeyShiftTab { + if fm.routePane == 0 { + fm.routePane = 1 + } else { + fm.routePane = 0 + } + return fm, nil + } + + route := fm.currentRoute() + if fm.routePane == 0 { + switch key.Type { + case tea.KeyUp: + if fm.routeCursor > 0 { + fm.routeCursor-- + } + return fm, nil + case tea.KeyDown: + if fm.routeCursor+1 < len(route.Hops) { + fm.routeCursor++ + } + return fm, nil + case tea.KeyBackspace, tea.KeyDelete: + if len(route.Hops) > 0 && fm.routeCursor < len(route.Hops) { + route.Hops = append(route.Hops[:fm.routeCursor], route.Hops[fm.routeCursor+1:]...) + fm.setCurrentRoute(route) + } + return fm, nil + case tea.KeyRunes: + switch key.String() { + case "x", "X": + if len(route.Hops) > 0 && fm.routeCursor < len(route.Hops) { + route.Hops = append(route.Hops[:fm.routeCursor], route.Hops[fm.routeCursor+1:]...) + fm.setCurrentRoute(route) + } + return fm, nil + case "[": + if fm.routeCursor > 0 && fm.routeCursor < len(route.Hops) { + i := fm.routeCursor + route.Hops[i-1], route.Hops[i] = route.Hops[i], route.Hops[i-1] + fm.routeCursor-- + fm.setCurrentRoute(route) + } + return fm, nil + case "]": + if fm.routeCursor >= 0 && fm.routeCursor+1 < len(route.Hops) { + i := fm.routeCursor + route.Hops[i], route.Hops[i+1] = route.Hops[i+1], route.Hops[i] + fm.routeCursor++ + fm.setCurrentRoute(route) + } + return fm, nil + } + } + return fm, nil + } + + if key.Type == tea.KeyEnter { + if item, ok := fm.routeList.SelectedItem().(routeProfileItem); ok && item.server != nil { + if !fm.routeContainsProfile(route, item.server.ID) { + route.Hops = append(route.Hops, model.RouteHop{ServerID: item.server.ID, Alias: item.server.Alias, IsProfile: true}) + fm.routeCursor = len(route.Hops) - 1 + fm.setCurrentRoute(route) + } + return fm, nil + } + } + var cmd tea.Cmd + fm.routeList, cmd = fm.routeList.Update(msg) + return fm, cmd +} + func (fm *formModel) Init() tea.Cmd { return nil } @@ -276,6 +440,10 @@ func (fm *formModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) { return fm, cmd } + if fm.showRouteList { + return fm.updateRouteEditor(msg) + } + if fm.showGroupList { switch msg := msg.(type) { case tea.KeyMsg: @@ -342,6 +510,15 @@ func (fm *formModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) { fm.showAuthList = true return fm, nil } + if len(msg.Runes) == 1 && msg.Runes[0] == '/' && !msg.Alt && fm.focusIdx == 7 { + fm.showRouteList = true + fm.routePane = 1 + route := fm.currentRoute() + if len(route.Hops) > 0 { + fm.routeCursor = len(route.Hops) - 1 + } + return fm, nil + } if len(msg.Runes) == 1 && msg.Runes[0] == '/' && !msg.Alt && fm.focusIdx == 8 && len(fm.groups) > 0 { fm.showGroupList = true return fm, nil @@ -414,6 +591,8 @@ func (fm *formModel) applySaveError(err error) { fm.focusIdx = 2 case strings.Contains(message, "port"): fm.focusIdx = 3 + case strings.Contains(message, "route"): + fm.focusIdx = 7 default: return } @@ -450,6 +629,9 @@ func (fm *formModel) labelAt(index int) string { if index == 5 { return "Auth Method (/ pick)" } + if index == 7 { + return "Route (/ edit)" + } if index == 8 { if len(fm.groups) > 0 { return "Group (/ pick)" @@ -471,7 +653,11 @@ func (fm *formModel) runTest() tea.Cmd { fm.testing = false return func() tea.Msg { return testDoneMsg{ok: false, err: err.Error()} } } - s := fm.buildServer() + s, err := fm.buildServerValidated() + if err != nil { + fm.testing = false + return func() tea.Msg { return testDoneMsg{ok: false, err: err.Error()} } + } pw := fm.password.Value() return tea.Batch( @@ -499,7 +685,10 @@ func (fm *formModel) runSave() tea.Cmd { if _, err := parsePort(fm.inputs[3].Value()); err != nil { return func() tea.Msg { return saveDoneMsg{err: err} } } - s := fm.buildServer() + s, err := fm.buildServerValidated() + if err != nil { + return func() tea.Msg { return saveDoneMsg{err: err} } + } pw := fm.password.Value() return tea.Batch( @@ -521,55 +710,54 @@ func (fm *formModel) runSave() tea.Cmd { ) } -// parseRouteHops parses the route hops input string into a model.Route. -// Format: comma-separated list of aliases or raw addresses. -func parseRouteHops(input string) model.Route { - input = strings.TrimSpace(input) - if input == "" { - return model.Route{} - } - parts := strings.Split(input, ",") - hops := make([]model.RouteHop, 0, len(parts)) - for _, p := range parts { - p = strings.TrimSpace(p) - if p == "" { - continue - } - // Heuristic: if it contains @ or :, treat as raw address - if strings.Contains(p, "@") || strings.Contains(p, ":") { - hops = append(hops, model.RouteHop{Raw: p, IsProfile: false}) - } else { - // Treat as profile alias - hops = append(hops, model.RouteHop{Alias: p, IsProfile: true}) - } - } - return model.Route{Hops: hops} +// parseRouteHops parses an explicit route spec. Exact known aliases become +// stable profile references; unknown unprefixed values remain raw OpenSSH targets. +func parseRouteHops(input string) (model.Route, error) { + return model.ParseRouteSpec(input, ResolveRouteAlias) } -func (fm *formModel) buildServer() *model.Server { - port, _ := parsePort(fm.inputs[3].Value()) - authMethod := model.AuthMethod(fm.inputs[5].Value()) +func (fm *formModel) buildServerValidated() (*model.Server, error) { + port, err := parsePort(fm.inputs[3].Value()) + if err != nil { + return nil, err + } + authMethod := model.AuthMethod(strings.TrimSpace(fm.inputs[5].Value())) if authMethod == "" { authMethod = model.AuthKey } - - route := parseRouteHops(fm.inputs[7].Value()) - - return &model.Server{ - Alias: fm.inputs[0].Value(), - DisplayName: fm.inputs[1].Value(), - Host: fm.inputs[2].Value(), + route, err := fm.parseRouteInput() + if err != nil { + return nil, fmt.Errorf("route: %w", err) + } + server := &model.Server{ + Alias: strings.TrimSpace(fm.inputs[0].Value()), + DisplayName: strings.TrimSpace(fm.inputs[1].Value()), + Host: strings.TrimSpace(fm.inputs[2].Value()), Port: port, - User: fm.inputs[4].Value(), + User: strings.TrimSpace(fm.inputs[4].Value()), AuthMethod: authMethod, - IdentityFile: fm.inputs[6].Value(), + IdentityFile: strings.TrimSpace(fm.inputs[6].Value()), ProxyJump: route.ProxyJumpString(), Route: route, - GroupName: fm.inputs[8].Value(), + GroupName: strings.TrimSpace(fm.inputs[8].Value()), Notes: fm.inputs[9].Value(), StartupCommand: fm.inputs[10].Value(), Tags: splitCSV(fm.inputs[11].Value()), } + if fm.edit && fm.server != nil { + server.ID = fm.server.ID + } + if err := model.ValidateServerBasics(server); err != nil { + return nil, err + } + return server, nil +} + +// buildServer is kept as a convenience for view/tests that already populate +// valid fields. Save/Test paths use buildServerValidated and surface errors. +func (fm *formModel) buildServer() *model.Server { + server, _ := fm.buildServerValidated() + return server } func parsePort(value string) (int, error) { @@ -589,6 +777,9 @@ func (fm *formModel) View() string { if fm.edit { title = "Edit Server: " + fm.server.Alias } + if fm.showRouteList { + return fm.routeEditorView(title) + } if fm.showAuthList || fm.showGroupList { var dropdown list.Model fieldIndex := 8 @@ -684,6 +875,89 @@ func (fm *formModel) View() string { }) } +func (fm *formModel) routeEditorView(title string) string { + route := fm.currentRoute() + body := func(width, height int) string { + lines := []string{dashboardSection("Current route")} + if len(route.Hops) == 0 { + line := " Direct connection" + if fm.routePane == 0 { + line = selectedRowStyle.Render("> Direct connection") + } + lines = append(lines, line) + } else { + for index, hop := range route.Hops { + kind := "raw" + name := hop.Raw + if hop.Profile() { + kind = "profile" + name = hop.Alias + } + prefix := " " + line := fmt.Sprintf("%s%d. %-24s [%s]", prefix, index+1, name, kind) + if fm.routePane == 0 && index == fm.routeCursor { + line = selectedRowStyle.Render(fmt.Sprintf("> %d. %-24s [%s]", index+1, name, kind)) + } + lines = append(lines, line) + } + } + target := strings.TrimSpace(fm.inputs[2].Value()) + if target == "" { + target = "target" + } + lines = append(lines, dashboardHelp("Preview: "+route.DisplaySummary(target)), "", dashboardSection("Available server profiles")) + + if len(fm.routeList.Items()) == 0 { + lines = append(lines, dashboardHelp("No other server profiles are available.")) + } else { + capacity := max(1, height-len(lines)-4) + start, end := visibleServerRange(len(fm.routeList.Items()), fm.routeList.Index(), capacity) + for index := start; index < end; index++ { + item, ok := fm.routeList.Items()[index].(routeProfileItem) + if !ok || item.server == nil { + continue + } + prefix := " " + mark := " " + if fm.routeContainsProfile(route, item.server.ID) { + mark = "✓" + } + target := item.server.Host + if item.server.User != "" { + target = item.server.User + "@" + item.server.Host + } + line := fmt.Sprintf("%s%s %-20s %s:%d", prefix, mark, item.server.Alias, target, item.server.Port) + if fm.routePane == 1 && index == fm.routeList.Index() { + line = selectedRowStyle.Render(fmt.Sprintf("> %s %-20s %s:%d", mark, item.server.Alias, target, item.server.Port)) + } + lines = append(lines, fitLine(line, max(1, width-4))) + } + } + lines = append(lines, "", dashboardHelp("Need a host that is not a sshkeeper profile? Esc and type raw: in the Route field.")) + return renderPaddedPanel(width, height, lines) + } + pane := "profiles" + if fm.routePane == 0 { + pane = "current route" + } + return renderScreenShell(screenShell{ + breadcrumb: title + " / Route Editor", + status: "Editing " + pane, + width: fm.width, + height: fm.height, + body: body, + footer: []helpItem{ + {Key: "Tab", Action: "switch pane"}, + {Key: "↑/↓", Action: "move"}, + {Key: "Enter", Action: "add profile"}, + {Key: "x/Del", Action: "remove hop"}, + {Key: "[/]", Action: "reorder hop"}, + {Key: "/", Action: "filter profiles"}, + {Key: "Esc", Action: "done"}, + }, + }) +} + func (fm *formModel) formStatusLine() string { if fm.err != nil { return errorStyle.Render(fmt.Sprintf("✗ Error: %v", fm.err)) diff --git a/internal/tui/form_validation_test.go b/internal/tui/form_validation_test.go index 7a5b877..a0cb73f 100644 --- a/internal/tui/form_validation_test.go +++ b/internal/tui/form_validation_test.go @@ -81,8 +81,8 @@ func TestForwardValidationMovesFocusToInvalidPort(t *testing.T) { fm.inputs[3].SetValue("5432") updated, _ := fm.Update(fm.runSave()()) fm = updated.(*forwardFormModel) - if fm.focusIdx != 6 { - t.Fatalf("invalid listen port focus = %d, want 6", fm.focusIdx) + if fm.focusIdx != 7 { + t.Fatalf("invalid listen port focus = %d, want 7", fm.focusIdx) } view := fm.View() if !strings.Contains(view, "Listen Port") || !strings.Contains(view, "must be a number") { diff --git a/internal/tui/forward.go b/internal/tui/forward.go index 98dc70f..490174e 100644 --- a/internal/tui/forward.go +++ b/internal/tui/forward.go @@ -112,6 +112,7 @@ func (m *forwardScreenModel) View() string { footer: []helpItem{ {Key: "Ctrl+A (a)", Action: "add"}, {Key: "Ctrl+E/Enter", Action: "edit"}, + {Key: "Space", Action: "enable/disable"}, {Key: "Ctrl+D (d)", Action: "delete"}, {Key: "Ctrl+H", Action: "help"}, {Key: "Esc", Action: "back"}, @@ -221,6 +222,7 @@ type forwardFormModel struct { nameInput textinput.Model descInput textinput.Model typeIdx int // 0=local, 1=remote, 2=socks + enabled bool width int height int initial forwardFormSnapshot @@ -231,6 +233,7 @@ type forwardFormSnapshot struct { description string values []string forwardType model.ForwardType + enabled bool } var forwardTypes = []forwardTypeItem{ @@ -262,6 +265,7 @@ func newForwardFormModel(serverID int64, w, h int) *forwardFormModel { focusIdx: 0, currentType: model.ForwardLocal, typeIdx: 0, + enabled: true, nameInput: nameInput, descInput: descInput, width: w, @@ -280,6 +284,7 @@ func newForwardEditModel(serverID int64, fwd *model.Forward, w, h int) *forwardF fm.descInput.SetValue(fwd.Description) fm.currentType = fwd.Type fm.typeIdx = typeIndex(fwd.Type) + fm.enabled = fwd.Enabled if fwd.Type == model.ForwardRemote { fm.inputs[0].SetValue(fwd.RemoteAddr) fm.inputs[1].SetValue(strconv.Itoa(fwd.RemotePort)) @@ -306,12 +311,13 @@ func (fm *forwardFormModel) snapshot() forwardFormSnapshot { description: fm.descInput.Value(), values: values, forwardType: fm.currentType, + enabled: fm.enabled, } } func (fm *forwardFormModel) Dirty() bool { current := fm.snapshot() - if current.name != fm.initial.name || current.description != fm.initial.description || current.forwardType != fm.initial.forwardType || len(current.values) != len(fm.initial.values) { + if current.name != fm.initial.name || current.description != fm.initial.description || current.forwardType != fm.initial.forwardType || current.enabled != fm.initial.enabled || len(current.values) != len(fm.initial.values) { return true } for i := range current.values { @@ -382,7 +388,7 @@ func (fm *forwardFormModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) { switch msg.Type { case tea.KeyTab: fm.focusIdx++ - total := 2 + 3 + len(fm.visibleFields()) + 1 // name + desc + type(3) + fields + save + total := 2 + 3 + 1 + len(fm.visibleFields()) + 1 // name + desc + type(3) + fields + save if fm.focusIdx >= total { fm.focusIdx = 0 } @@ -391,7 +397,7 @@ func (fm *forwardFormModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) { case tea.KeyShiftTab: fm.focusIdx-- if fm.focusIdx < 0 { - total := 2 + 3 + len(fm.visibleFields()) + 1 + total := 2 + 3 + 1 + len(fm.visibleFields()) + 1 fm.focusIdx = total - 1 } fm.updateFocus() @@ -405,7 +411,11 @@ func (fm *forwardFormModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) { fm.updateFocus() return fm, nil } - if fm.focusIdx == 2+3+len(fm.visibleFields()) { + if fm.focusIdx == 2+3 { + fm.enabled = !fm.enabled + return fm, nil + } + if fm.focusIdx == 2+3+1+len(fm.visibleFields()) { return fm, fm.runSave() } fm.focusIdx++ @@ -415,7 +425,7 @@ func (fm *forwardFormModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) { return fm, nil case tea.KeyDown: fm.focusIdx++ - total := 2 + 3 + len(fm.visibleFields()) + 1 + total := 2 + 3 + 1 + len(fm.visibleFields()) + 1 if fm.focusIdx >= total { fm.focusIdx = 0 } @@ -424,12 +434,16 @@ func (fm *forwardFormModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) { case tea.KeyUp: fm.focusIdx-- if fm.focusIdx < 0 { - total := 2 + 3 + len(fm.visibleFields()) + 1 + total := 2 + 3 + 1 + len(fm.visibleFields()) + 1 fm.focusIdx = total - 1 } fm.updateFocus() return fm, nil case tea.KeyRunes: + if fm.focusIdx == 2+3 && msg.String() == " " { + fm.enabled = !fm.enabled + return fm, nil + } // Direct number keys select a type only while the type selector has focus. if fm.focusIdx >= 2 && fm.focusIdx < 2+len(forwardTypes) && len(msg.Runes) == 1 { switch msg.Runes[0] { @@ -465,8 +479,8 @@ func (fm *forwardFormModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) { return fm, cmd } visible := fm.visibleFields() - if fm.focusIdx >= 2+3 && fm.focusIdx < 2+3+len(visible) { - fieldIdx := visible[fm.focusIdx-(2+3)] + if fm.focusIdx >= 2+3+1 && fm.focusIdx < 2+3+1+len(visible) { + fieldIdx := visible[fm.focusIdx-(2+3+1)] var cmd tea.Cmd fm.inputs[fieldIdx], cmd = fm.inputs[fieldIdx].Update(msg) return fm, cmd @@ -485,7 +499,7 @@ func (fm *forwardFormModel) updateFocus() { fm.inputs[i].Prompt = blurredStyle.Render(fm.labelForField(i) + ": ") } - total := 2 + 3 + len(fm.visibleFields()) + 1 + total := 2 + 3 + 1 + len(fm.visibleFields()) + 1 switch { case fm.focusIdx == 0: fm.nameInput.Focus() @@ -495,98 +509,96 @@ func (fm *forwardFormModel) updateFocus() { fm.descInput.Prompt = focusedStyle.Render("Description> ") case fm.focusIdx >= 2 && fm.focusIdx < 2+3: // Type selector focused — no input to focus - case fm.focusIdx >= 2+3 && fm.focusIdx < total-1: + case fm.focusIdx == 2+3: + // Enabled toggle focused. + case fm.focusIdx >= 2+3+1 && fm.focusIdx < total-1: visible := fm.visibleFields() - fieldIdx := visible[fm.focusIdx-(2+3)] + fieldIdx := visible[fm.focusIdx-(2+3+1)] fm.inputs[fieldIdx].Focus() fm.inputs[fieldIdx].Prompt = focusedStyle.Render(fm.labelForField(fieldIdx) + "> ") } } +func (fm *forwardFormModel) buildForwardFromForm() (*model.Forward, error) { + name := strings.TrimSpace(fm.nameInput.Value()) + if name == "" { + return nil, fmt.Errorf("name is required") + } + forward := &model.Forward{ + ID: fm.editID, + ServerID: fm.serverID, + Name: name, + Description: strings.TrimSpace(fm.descInput.Value()), + Type: fm.currentType, + Enabled: fm.enabled, + } + var err error + switch fm.currentType { + case model.ForwardLocal: + forward.LocalAddr = strings.TrimSpace(fm.inputs[0].Value()) + if forward.LocalAddr == "" { + forward.LocalAddr = "127.0.0.1" + } + forward.LocalPort, err = parseNamedPort("Listen port", fm.inputs[1].Value()) + if err != nil { + return nil, err + } + forward.RemoteAddr = strings.TrimSpace(fm.inputs[2].Value()) + if forward.RemoteAddr == "" { + return nil, fmt.Errorf("target host is required for local forward") + } + forward.RemotePort, err = parseNamedPort("Target port", fm.inputs[3].Value()) + if err != nil { + return nil, err + } + case model.ForwardRemote: + forward.RemoteAddr = strings.TrimSpace(fm.inputs[0].Value()) + if forward.RemoteAddr == "" { + return nil, fmt.Errorf("remote listen address is required") + } + forward.RemotePort, err = parseNamedPort("Remote listen port", fm.inputs[1].Value()) + if err != nil { + return nil, err + } + forward.LocalAddr = strings.TrimSpace(fm.inputs[2].Value()) + if forward.LocalAddr == "" { + forward.LocalAddr = "127.0.0.1" + } + forward.LocalPort, err = parseNamedPort("Local target port", fm.inputs[3].Value()) + if err != nil { + return nil, err + } + case model.ForwardDynamic: + forward.LocalAddr = strings.TrimSpace(fm.inputs[0].Value()) + if forward.LocalAddr == "" { + forward.LocalAddr = "127.0.0.1" + } + forward.LocalPort, err = parseNamedPort("Listen port", fm.inputs[1].Value()) + if err != nil { + return nil, err + } + default: + return nil, fmt.Errorf("unsupported forward type: %s", fm.currentType) + } + return forward, nil +} + func (fm *forwardFormModel) runSave() tea.Cmd { return func() tea.Msg { - name := strings.TrimSpace(fm.nameInput.Value()) - desc := strings.TrimSpace(fm.descInput.Value()) - localAddr, remoteAddr := "", "" - localPort, remotePort := 0, 0 - var err error - - if name == "" { - return saveDoneMsg{err: fmt.Errorf("name is required")} + forward, err := fm.buildForwardFromForm() + if err != nil { + return saveDoneMsg{err: err} } - switch fm.currentType { - case model.ForwardLocal: - localAddr = strings.TrimSpace(fm.inputs[0].Value()) - if localAddr == "" { - localAddr = "127.0.0.1" - } - localPort, err = parseNamedPort("Listen port", fm.inputs[1].Value()) - if err != nil { - return saveDoneMsg{err: err} - } - remoteAddr = strings.TrimSpace(fm.inputs[2].Value()) - if remoteAddr == "" { - return saveDoneMsg{err: fmt.Errorf("target host is required for local forward")} - } - remotePort, err = parseNamedPort("Target port", fm.inputs[3].Value()) - if err != nil { - return saveDoneMsg{err: err} - } - case model.ForwardRemote: - remoteAddr = strings.TrimSpace(fm.inputs[0].Value()) - if remoteAddr == "" { - return saveDoneMsg{err: fmt.Errorf("remote listen address is required")} - } - remotePort, err = parseNamedPort("Remote listen port", fm.inputs[1].Value()) - if err != nil { - return saveDoneMsg{err: err} - } - localAddr = strings.TrimSpace(fm.inputs[2].Value()) - if localAddr == "" { - localAddr = "127.0.0.1" - } - localPort, err = parseNamedPort("Local target port", fm.inputs[3].Value()) - if err != nil { - return saveDoneMsg{err: err} - } - case model.ForwardDynamic: - localAddr = strings.TrimSpace(fm.inputs[0].Value()) - if localAddr == "" { - localAddr = "127.0.0.1" - } - localPort, err = parseNamedPort("Listen port", fm.inputs[1].Value()) - if err != nil { - return saveDoneMsg{err: err} - } - remoteAddr = "" - remotePort = 0 - } - - fwd := &model.Forward{ - ServerID: fm.serverID, - Name: name, - Description: desc, - Type: fm.currentType, - LocalAddr: localAddr, - LocalPort: localPort, - RemoteAddr: remoteAddr, - RemotePort: remotePort, - Enabled: true, - } - if fm.editMode { - fwd.ID = fm.editID if UpdateForward == nil { return saveDoneMsg{err: fmt.Errorf("update not available")} } - return saveDoneMsg{err: UpdateForward(fwd)} + return saveDoneMsg{err: UpdateForward(forward)} } - if SaveForward == nil { return saveDoneMsg{err: fmt.Errorf("forward storage is unavailable")} } - err = SaveForward(fwd) - return saveDoneMsg{err: err} + return saveDoneMsg{err: SaveForward(forward)} } } @@ -612,7 +624,7 @@ func (fm *forwardFormModel) applySaveError(err error) { fieldIndex = 3 } if fieldIndex >= 0 { - fm.focusIdx = 2 + len(forwardTypes) + fieldIndex + fm.focusIdx = 2 + len(forwardTypes) + 1 + fieldIndex fm.updateFocus() } } @@ -658,6 +670,15 @@ func (fm *forwardFormModel) View() string { if width >= 100 { lines = append(lines, helpStyle.Copy().MarginLeft(0).Render(forwardTypes[fm.typeIdx].description)) } + enabledMark := "[ ]" + if fm.enabled { + enabledMark = "[x]" + } + enabledLine := " Enabled " + enabledMark + if fm.focusIdx == 2+3 { + enabledLine = selectedStyle.Render("> Enabled " + enabledMark + " Enter/Space toggles") + } + lines = append(lines, enabledLine) visible := fm.visibleFields() for _, idx := range visible { lines = append(lines, fm.inputs[idx].View()) @@ -666,13 +687,12 @@ func (fm *forwardFormModel) View() string { lines = append(lines, helpStyle.Copy().MarginLeft(0).Render("⚠ This port will be accessible from the network.")) } if width >= 70 && fm.currentType != "" && fm.inputs[1].Value() != "" { - fwd := &model.Forward{Type: fm.currentType, LocalAddr: fm.inputs[0].Value(), RemoteAddr: fm.inputs[2].Value()} - fmt.Sscanf(fm.inputs[1].Value(), "%d", &fwd.LocalPort) - fmt.Sscanf(fm.inputs[3].Value(), "%d", &fwd.RemotePort) - preview := "Preview ssh " + strings.Join(fwd.ForwardSSHArgs(), " ") + " -o ExitOnForwardFailure=yes" - lines = append(lines, wrapCells(preview, contentWidth)...) + if fwd, err := fm.buildForwardFromForm(); err == nil { + preview := "Preview ssh " + strings.Join(fwd.ForwardSSHArgs(), " ") + " -o ExitOnForwardFailure=yes" + lines = append(lines, wrapCells(preview, contentWidth)...) + } } - total := 2 + 3 + len(visible) + 1 + total := 2 + 3 + 1 + len(visible) + 1 button := " [ Save ]" if fm.focusIdx == total-1 { button = selectedStyle.Render("> [ Save ]") @@ -691,7 +711,7 @@ func (fm *forwardFormModel) View() string { {Key: "Tab/↓", Action: "next"}, {Key: "↑", Action: "prev"}, {Key: "1/2/3", Action: "select type"}, - {Key: "Enter", Action: "save"}, + {Key: "Enter/Space", Action: "toggle/save"}, {Key: "Ctrl+H", Action: "help"}, {Key: "Esc", Action: "back"}, }, diff --git a/internal/tui/forward_test.go b/internal/tui/forward_test.go index 556e6d0..6f09101 100644 --- a/internal/tui/forward_test.go +++ b/internal/tui/forward_test.go @@ -11,7 +11,7 @@ import ( func TestForwardFormDigitsReachFocusedInput(t *testing.T) { fm := newForwardFormModel(1, 100, 30) - fm.focusIdx = 2 + len(forwardTypes) + 1 + fm.focusIdx = 2 + len(forwardTypes) + 2 fm.updateFocus() for _, digit := range []rune{'1', '2', '3'} { @@ -72,6 +72,40 @@ func TestRemoteForwardEditPopulatesSemanticFields(t *testing.T) { } } +func TestForwardEditPreservesDisabledState(t *testing.T) { + forward := &model.Forward{ID: 11, ServerID: 7, Name: "db", Type: model.ForwardLocal, LocalAddr: "127.0.0.1", LocalPort: 15432, RemoteAddr: "127.0.0.1", RemotePort: 5432, Enabled: false} + fm := newForwardEditModel(7, forward, 80, 24) + if fm.enabled { + t.Fatal("disabled forward became enabled in edit form") + } + built, err := fm.buildForwardFromForm() + if err != nil { + t.Fatalf("build forward: %v", err) + } + if built.Enabled { + t.Fatal("disabled forward would be saved as enabled") + } +} + +func TestRemoteForwardPreviewUsesSavedEndpointMapping(t *testing.T) { + fm := newForwardFormModel(7, 100, 30) + fm.currentType = model.ForwardRemote + fm.typeIdx = typeIndex(model.ForwardRemote) + fm.nameInput.SetValue("remote web") + fm.inputs[0].SetValue("0.0.0.0") + fm.inputs[1].SetValue("18080") + fm.inputs[2].SetValue("127.0.0.1") + fm.inputs[3].SetValue("8080") + built, err := fm.buildForwardFromForm() + if err != nil { + t.Fatalf("build preview forward: %v", err) + } + want := []string{"-R", "0.0.0.0:18080:127.0.0.1:8080"} + if got := built.ForwardSSHArgs(); !reflect.DeepEqual(got, want) { + t.Fatalf("preview args = %#v, want %#v", got, want) + } +} + func TestForwardFormDigitShortcutsWorkOnTypeSelector(t *testing.T) { tests := []struct { digit rune diff --git a/internal/tunnel/manager.go b/internal/tunnel/manager.go index 3f7f96e..2b25359 100644 --- a/internal/tunnel/manager.go +++ b/internal/tunnel/manager.go @@ -88,8 +88,14 @@ func Get(id int64) *model.TunnelState { return states[id] } -// Start starts a tunnel for the given server with its forwards. +// Start starts a tunnel without profile resolution (legacy/direct routes). func Start(cfg *config.Config, server *model.Server, forwards []*model.Forward, forwardOnly bool) (*model.TunnelState, error) { + return StartResolved(cfg, server, forwards, forwardOnly, nil) +} + +// StartResolved starts a background tunnel and retains any generated OpenSSH +// config until the tunnel is stopped. +func StartResolved(cfg *config.Config, server *model.Server, forwards []*model.Forward, forwardOnly bool, resolve ssh.ProfileResolver) (*model.TunnelState, error) { mu.Lock() defer mu.Unlock() @@ -105,9 +111,11 @@ func Start(cfg *config.Config, server *model.Server, forwards []*model.Forward, } } - sshArgs := ssh.BuildSSHArgs(server, active, forwardOnly) - args := make([]string, len(sshArgs)) - copy(args, sshArgs) + invocation, err := ssh.PrepareSSHInvocation(server, active, forwardOnly, resolve) + if err != nil { + return nil, err + } + args := append([]string(nil), invocation.Args...) cmd := exec.Command(cfg.SSH.Binary, args...) cmd.Env = os.Environ() @@ -116,6 +124,7 @@ func Start(cfg *config.Config, server *model.Server, forwards []*model.Forward, cmd.Stderr = nil if err := cmd.Start(); err != nil { + invocation.Cleanup() return nil, fmt.Errorf("start tunnel: %w", err) } @@ -132,6 +141,7 @@ func Start(cfg *config.Config, server *model.Server, forwards []*model.Forward, Name: fmt.Sprintf("Tunnel to %s", server.Alias), PID: cmd.Process.Pid, ForwardIDs: forwardIDs, + ConfigPath: invocation.ConfigPath, StartedAt: time.Now(), } @@ -139,10 +149,12 @@ func Start(cfg *config.Config, server *model.Server, forwards []*model.Forward, if err := saveStates(); err != nil { delete(states, id) _ = cmd.Process.Kill() + invocation.Cleanup() return nil, fmt.Errorf("save tunnel state: %w", err) } if err := cmd.Process.Release(); err != nil { delete(states, id) + invocation.Cleanup() return nil, fmt.Errorf("release tunnel process: %w", err) } @@ -166,6 +178,9 @@ func Stop(id int64) error { } } + if state.ConfigPath != "" { + _ = os.Remove(state.ConfigPath) + } delete(states, id) return saveStates() } @@ -182,6 +197,9 @@ func StopAll() error { proc.Kill() } } + if state.ConfigPath != "" { + _ = os.Remove(state.ConfigPath) + } delete(states, id) } return saveStates()