feat: unify route workflows and forwarding UX

This commit is contained in:
mirivlad 2026-09-05 19:10:02 +08:00
parent 3cc21a7b22
commit f05e8e8e84
16 changed files with 846 additions and 416 deletions

View File

@ -21,6 +21,7 @@ var addFlags struct {
authMethod string authMethod string
identityFile string identityFile string
proxyJump string proxyJump string
route string
groupName string groupName string
displayName string displayName string
notes string notes string
@ -50,6 +51,14 @@ func addInteractive() error {
} }
func addNonInteractive(alias string) 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{ server := &model.Server{
Alias: alias, Alias: alias,
DisplayName: addFlags.displayName, DisplayName: addFlags.displayName,
@ -58,7 +67,8 @@ func addNonInteractive(alias string) error {
User: addFlags.user, User: addFlags.user,
AuthMethod: model.AuthMethod(addFlags.authMethod), AuthMethod: model.AuthMethod(addFlags.authMethod),
IdentityFile: addFlags.identityFile, IdentityFile: addFlags.identityFile,
ProxyJump: addFlags.proxyJump, Route: route,
ProxyJump: route.ProxyJumpString(),
GroupName: addFlags.groupName, GroupName: addFlags.groupName,
Notes: addFlags.notes, Notes: addFlags.notes,
StartupCommand: addFlags.startup, StartupCommand: addFlags.startup,
@ -78,13 +88,30 @@ func addNonInteractive(alias string) error {
} }
func saveServerWithOptionalSecret(server *model.Server) error { func saveServerWithOptionalSecret(server *model.Server) error {
// Handle password/passphrase auth — request interactively, never via argv if len(server.Route.Hops) == 0 && strings.TrimSpace(server.ProxyJump) != "" {
if server.AuthMethod == model.AuthPassword || server.AuthMethod == model.AuthKeyPassphrase { 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" secretType := "password"
if server.AuthMethod == model.AuthKeyPassphrase { if server.AuthMethod == model.AuthKeyPassphrase {
secretType = "passphrase" secretType = "passphrase"
} }
fmt.Printf("Enter %s (will be stored in vault, input hidden): ", secretType) fmt.Printf("Enter %s (will be stored in vault, input hidden): ", secretType)
password, err := term.ReadPassword(int(syscall.Stdin)) password, err := term.ReadPassword(int(syscall.Stdin))
fmt.Println() fmt.Println()
@ -94,40 +121,41 @@ func saveServerWithOptionalSecret(server *model.Server) error {
if len(password) == 0 { if len(password) == 0 {
return fmt.Errorf("%s cannot be empty", secretType) return fmt.Errorf("%s cannot be empty", secretType)
} }
secret = password
v := getOrCreateVault() defer func() {
for i := range secret {
secret[i] = 0
}
}()
if err := unlockVaultForCommand(v); err != nil { if err := unlockVaultForCommand(v); err != nil {
return err 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 { if err := appDB.CreateServer(server); err != nil {
return fmt.Errorf("create server: %w", err) 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 len(server.Tags) > 0 {
if err := appDB.SetServerTags(server.ID, server.Tags); err != nil { if err := appDB.SetServerTags(server.ID, server.Tags); err != nil {
rollbackDB()
return fmt.Errorf("set tags: %w", err) 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.") fmt.Println("Saved.")
return nil return nil
} }
@ -171,7 +199,7 @@ func promptServerForAdd(in io.Reader, out io.Writer) (*model.Server, error) {
if err != nil { if err != nil {
return nil, err return nil, err
} }
proxyJump, err := promptOptional(reader, out, "ProxyJump", "") proxyJump, err := promptOptional(reader, out, "Route / ProxyJump (profile:<alias> or raw:<target>)", "")
if err != nil { if err != nil {
return nil, err return nil, err
} }
@ -253,7 +281,8 @@ func init() {
addCmd.Flags().StringVar(&addFlags.user, "user", "", "SSH username") 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.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.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:<alias>, raw:<target>, 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.groupName, "group", "", "Server group")
addCmd.Flags().StringVar(&addFlags.displayName, "display-name", "", "Display name") addCmd.Flags().StringVar(&addFlags.displayName, "display-name", "", "Display name")
addCmd.Flags().StringVar(&addFlags.notes, "notes", "", "Notes") addCmd.Flags().StringVar(&addFlags.notes, "notes", "", "Notes")

View File

@ -19,33 +19,9 @@ var connectCmd = &cobra.Command{
if err != nil { if err != nil {
return fmt.Errorf("server not found: %s", alias) return fmt.Errorf("server not found: %s", alias)
} }
if err := ssh.ConnectResolved(cfg, server, dbProfileResolver, serverVaultFunc(server)); err != nil {
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 {
return err return err
} }
appDB.UpdateLastConnected(alias) appDB.UpdateLastConnected(alias)
return nil return nil
}, },
@ -61,31 +37,7 @@ var testCmd = &cobra.Command{
if err != nil { if err != nil {
return fmt.Errorf("server not found: %s", alias) return fmt.Errorf("server not found: %s", alias)
} }
ok, testErr := ssh.TestResolved(cfg, server, dbProfileResolver, serverVaultFunc(server))
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)
if ok { if ok {
fmt.Println("Connection OK.") fmt.Println("Connection OK.")
appDB.UpdateTestResult(alias, model.TestOK, "") appDB.UpdateTestResult(alias, model.TestOK, "")
@ -93,7 +45,6 @@ var testCmd = &cobra.Command{
fmt.Printf("Connection failed:\n%s\n", testErr) fmt.Printf("Connection failed:\n%s\n", testErr)
appDB.UpdateTestResult(alias, model.TestFailed, testErr) appDB.UpdateTestResult(alias, model.TestFailed, testErr)
} }
return nil return nil
}, },
} }

View File

@ -20,99 +20,138 @@ var editCmd = &cobra.Command{
if err != nil { if err != nil {
return fmt.Errorf("server not found: %s", alias) return fmt.Errorf("server not found: %s", alias)
} }
original := cloneServer(server)
oldAuthMethod := server.AuthMethod if cmd.Flags().Changed("host") {
if parsedHost != "" {
server.Host = parsedHost 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 server.Port = parsedPort
} }
if parsedUser != "" { if cmd.Flags().Changed("user") {
server.User = parsedUser server.User = parsedUser
} }
if parsedAuth != "" { authChanged := cmd.Flags().Changed("auth")
if authChanged {
server.AuthMethod = model.AuthMethod(parsedAuth) server.AuthMethod = model.AuthMethod(parsedAuth)
} }
if parsedIdentity != "" { if cmd.Flags().Changed("identity-file") {
server.IdentityFile = parsedIdentity server.IdentityFile = parsedIdentity
} }
if parsedProxyJump != "" { if cmd.Flags().Changed("group") {
server.ProxyJump = parsedProxyJump
}
if parsedGroup != "" {
server.GroupName = parsedGroup server.GroupName = parsedGroup
} }
if parsedDisplayName != "" { if cmd.Flags().Changed("display-name") {
server.DisplayName = parsedDisplayName server.DisplayName = parsedDisplayName
} }
if parsedNotes != "" { if cmd.Flags().Changed("notes") {
server.Notes = parsedNotes server.Notes = parsedNotes
} }
if parsedStartup != "" { if cmd.Flags().Changed("startup-command") {
server.StartupCommand = parsedStartup 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") tagsChanged := cmd.Flags().Changed("tags")
if tagsChanged { if tagsChanged {
if strings.TrimSpace(parsedTags) == "" {
server.Tags = nil
} else {
server.Tags = strings.Split(parsedTags, ",") server.Tags = strings.Split(parsedTags, ",")
} }
}
if err := model.ValidateServerBasics(server); err != nil {
return err
}
if parsedAuth != "" && oldAuthMethod != server.AuthMethod { var secret []byte
v := getOrCreateVault() v := getOrCreateVault()
if v.IsUnlocked() { if authChanged {
var secret string if err := unlockVaultForCommand(v); err != nil {
if server.AuthMethod == model.AuthPassword { return err
fmt.Print("Enter new password (stored in vault, input hidden): ") }
pw, err := term.ReadPassword(int(syscall.Stdin)) switch server.AuthMethod {
case model.AuthPassword, model.AuthKeyPassphrase:
label := "password"
if server.AuthMethod == model.AuthKeyPassphrase {
label = "key passphrase"
}
fmt.Printf("Enter new %s (stored in vault, input hidden): ", label)
secret, err = term.ReadPassword(int(syscall.Stdin))
fmt.Println() fmt.Println()
if err != nil { if err != nil {
return fmt.Errorf("read password: %w", err) return fmt.Errorf("read %s: %w", label, err)
} }
if len(pw) > 0 { if len(secret) == 0 {
secret = string(pw) return fmt.Errorf("%s cannot be empty", label)
} }
} else if server.AuthMethod == model.AuthKeyPassphrase { defer func() {
fmt.Print("Enter key passphrase (stored in vault, input hidden): ") for i := range secret {
pw, err := term.ReadPassword(int(syscall.Stdin)) secret[i] = 0
fmt.Println()
if err != nil {
return fmt.Errorf("read passphrase: %w", err)
} }
if len(pw) > 0 { }()
secret = string(pw)
} }
} }
if err := syncServerSecrets(v, alias, server, secret); err != nil { if err := appDB.UpdateServerByAlias(alias, server); err != nil {
return fmt.Errorf("sync vault secrets: %w", err)
}
if err := v.Save(); err != nil {
return fmt.Errorf("save vault: %w", err)
}
}
}
if err := appDB.UpdateServer(server); err != nil {
return fmt.Errorf("update server: %w", err) return fmt.Errorf("update server: %w", err)
} }
rollback := func() {
_ = appDB.UpdateServerByAlias(server.Alias, original)
_ = appDB.SetServerTags(original.ID, original.Tags)
}
if tagsChanged { if tagsChanged {
if err := appDB.SetServerTags(server.ID, server.Tags); err != nil { if err := appDB.SetServerTags(server.ID, server.Tags); err != nil {
rollback()
return fmt.Errorf("set tags: %w", err) 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.") fmt.Println("Saved.")
return nil 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 &copyServer
}
var ( var (
parsedHost string parsedHost string
parsedPort int parsedPort int
parsedUser string parsedUser string
parsedAuth string parsedAuth string
parsedIdentity string parsedIdentity string
parsedRoute string
parsedProxyJump string parsedProxyJump string
parsedGroup string parsedGroup string
parsedDisplayName string parsedDisplayName string
@ -124,13 +163,14 @@ var (
func init() { func init() {
editCmd.Flags().StringVar(&parsedHost, "host", "", "Server hostname or IP") editCmd.Flags().StringVar(&parsedHost, "host", "", "Server hostname or IP")
editCmd.Flags().IntVar(&parsedPort, "port", 0, "SSH port") 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(&parsedAuth, "auth", "", "Auth method")
editCmd.Flags().StringVar(&parsedIdentity, "identity-file", "", "Path to SSH private key") editCmd.Flags().StringVar(&parsedIdentity, "identity-file", "", "Path to SSH private key; empty clears it")
editCmd.Flags().StringVar(&parsedProxyJump, "proxy-jump", "", "ProxyJump host") editCmd.Flags().StringVar(&parsedRoute, "route", "", "Route hops: profile:<alias>, raw:<target>, comma-separated; empty means direct")
editCmd.Flags().StringVar(&parsedGroup, "group", "", "Server group") editCmd.Flags().StringVar(&parsedProxyJump, "proxy-jump", "", "Compatibility alias for --route")
editCmd.Flags().StringVar(&parsedDisplayName, "display-name", "", "Display name") editCmd.Flags().StringVar(&parsedGroup, "group", "", "Server group; empty clears it")
editCmd.Flags().StringVar(&parsedNotes, "notes", "", "Notes") editCmd.Flags().StringVar(&parsedDisplayName, "display-name", "", "Display name; empty clears it")
editCmd.Flags().StringVar(&parsedStartup, "startup-command", "", "Command to run after connecting") editCmd.Flags().StringVar(&parsedNotes, "notes", "", "Notes; empty clears them")
editCmd.Flags().StringVar(&parsedTags, "tags", "", "Comma-separated tags") editCmd.Flags().StringVar(&parsedStartup, "startup-command", "", "Startup command; empty clears it")
editCmd.Flags().StringVar(&parsedTags, "tags", "", "Comma-separated tags; empty clears all")
} }

View File

@ -43,7 +43,6 @@ func importServersFromSSHConfig(report func(format string, args ...interface{}))
if err != nil { if err != nil {
return 0, fmt.Errorf("import: %w", err) return 0, fmt.Errorf("import: %w", err)
} }
if len(servers) == 0 { if len(servers) == 0 {
if report != nil { if report != nil {
report("No servers found in ~/.ssh/config") report("No servers found in ~/.ssh/config")
@ -51,27 +50,63 @@ func importServersFromSSHConfig(report func(format string, args ...interface{}))
return 0, nil 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 imported := 0
for _, s := range servers { for _, server := range servers {
existing, _ := appDB.GetServer(s.Alias) if existing, _ := appDB.GetServer(server.Alias); existing != nil {
if existing != nil {
if report != nil { if report != nil {
report(" skip (exists): %s", s.Alias) report(" skip (exists): %s", server.Alias)
} }
continue 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 { if report != nil {
report(" error: %s: %v", s.Alias, err) report(" error: %s: %v", server.Alias, err)
} }
continue continue
} }
if report != nil { pending = append(pending, pendingRoute{server: server, spec: spec})
report(" imported: %s (%s@%s:%d)", s.Alias, s.User, s.Host, s.Port)
}
imported++ 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 return imported, nil
} }
@ -101,18 +136,5 @@ var runCmd = &cobra.Command{
} }
func runCommandOnServer(server *model.Server, command string) error { func runCommandOnServer(server *model.Server, command string) error {
return ssh.RunCommand(cfg, server, commandVaultFunc, command) return ssh.RunCommandResolved(cfg, server, dbProfileResolver, serverVaultFunc(server), 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
} }

View File

@ -8,11 +8,9 @@ import (
"github.com/spf13/cobra" "github.com/spf13/cobra"
) )
// --- Route commands ---
var routeCmd = &cobra.Command{ var routeCmd = &cobra.Command{
Use: "route", Use: "route",
Short: "Manage server routes (ProxyJump)", Short: "Manage server routes (bastions / ProxyJump)",
} }
var routeShowCmd = &cobra.Command{ var routeShowCmd = &cobra.Command{
@ -20,31 +18,30 @@ var routeShowCmd = &cobra.Command{
Short: "Show route for a server", Short: "Show route for a server",
Args: cobra.ExactArgs(1), Args: cobra.ExactArgs(1),
RunE: func(cmd *cobra.Command, args []string) error { RunE: func(cmd *cobra.Command, args []string) error {
alias := args[0] server, err := appDB.GetServer(args[0])
server, err := appDB.GetServer(alias)
if err != nil { if err != nil {
return fmt.Errorf("server not found: %s", alias) return fmt.Errorf("server not found: %s", args[0])
}
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
} }
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("Route: %s\n", server.Route.DisplaySummary(target))
fmt.Printf("Mode: %s\n", server.Route.RouteMode()) fmt.Printf("Mode: %s\n", server.Route.RouteMode())
fmt.Printf("ProxyJump: %s\n", server.Route.ProxyJumpString()) fmt.Printf("Spec: %s\n", model.FormatRouteSpec(server.Route))
if server.Route.HasProfileLinks() {
fmt.Println("Hops:") fmt.Println("Hops:")
for _, h := range server.Route.Hops { for index, hop := range server.Route.Hops {
if h.IsProfile { if hop.Profile() {
fmt.Printf(" - %s (profile)\n", h.Alias) fmt.Printf(" %d. %s (sshkeeper profile #%d)\n", index+1, hop.Alias, hop.ServerID)
} else { } else {
fmt.Printf(" - %s (raw)\n", h.Raw) fmt.Printf(" %d. %s (raw OpenSSH target)\n", index+1, hop.Raw)
} }
} }
}
} else if server.ProxyJump != "" {
fmt.Printf("ProxyJump: %s\n", server.ProxyJump)
} else {
fmt.Println("Direct connection (no route)")
}
return nil return nil
}, },
} }
@ -52,47 +49,38 @@ var routeShowCmd = &cobra.Command{
var routeSetCmd = &cobra.Command{ var routeSetCmd = &cobra.Command{
Use: "set <alias>", Use: "set <alias>",
Short: "Set route for a server", 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:<alias> to require a profile reference and raw:<target> to force a literal OpenSSH target.`,
Args: cobra.ExactArgs(1),
RunE: func(cmd *cobra.Command, args []string) error { RunE: func(cmd *cobra.Command, args []string) error {
alias := args[0] server, err := appDB.GetServer(args[0])
server, err := appDB.GetServer(alias)
if err != nil { 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") mode, _ := cmd.Flags().GetString("mode")
jumps, _ := cmd.Flags().GetString("jumps") jumps, _ := cmd.Flags().GetString("jumps")
mode = strings.ToLower(strings.TrimSpace(mode))
if mode == "clear" || jumps == "" { if mode == "clear" || mode == "direct" {
server.Route = model.Route{} server.Route = model.Route{}
server.ProxyJump = "" server.ProxyJump = ""
} else { } else {
parts := strings.Split(jumps, ",") if strings.TrimSpace(jumps) == "" {
hops := make([]model.RouteHop, 0, len(parts)) return fmt.Errorf("--jumps is required unless --mode=direct/clear")
for _, p := range parts {
p = strings.TrimSpace(p)
if p == "" {
continue
} }
if strings.Contains(p, "@") || strings.Contains(p, ":") { route, err := parseRouteSpec(jumps)
hops = append(hops, model.RouteHop{Raw: p, IsProfile: false}) if err != nil {
} else { return err
hops = append(hops, model.RouteHop{Alias: p, IsProfile: true})
} }
server.Route = route
server.ProxyJump = route.ProxyJumpString()
} }
server.Route = model.Route{Hops: hops}
server.ProxyJump = server.Route.ProxyJumpString()
}
if err := appDB.UpdateServer(server); err != nil { if err := appDB.UpdateServer(server); err != nil {
return fmt.Errorf("update route: %w", err) return fmt.Errorf("update route: %w", err)
} }
if len(server.Route.Hops) == 0 {
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 {
fmt.Println("✓ Route cleared (direct connection)") fmt.Println("✓ Route cleared (direct connection)")
} else {
fmt.Printf("✓ Route set: %s\n", model.FormatRouteSpec(server.Route))
} }
return nil return nil
}, },
@ -103,10 +91,9 @@ var routeClearCmd = &cobra.Command{
Short: "Clear route for a server (set direct)", Short: "Clear route for a server (set direct)",
Args: cobra.ExactArgs(1), Args: cobra.ExactArgs(1),
RunE: func(cmd *cobra.Command, args []string) error { RunE: func(cmd *cobra.Command, args []string) error {
alias := args[0] server, err := appDB.GetServer(args[0])
server, err := appDB.GetServer(alias)
if err != nil { 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.Route = model.Route{}
server.ProxyJump = "" server.ProxyJump = ""
@ -119,9 +106,8 @@ var routeClearCmd = &cobra.Command{
} }
func init() { func init() {
routeSetCmd.Flags().String("mode", "via", "Route mode: direct, via, chain, or clear") routeSetCmd.Flags().String("mode", "via", "Route mode: via, chain, direct, or clear")
routeSetCmd.Flags().String("jumps", "", "Comma-separated jump hosts (aliases or raw addresses)") routeSetCmd.Flags().String("jumps", "", "Comma-separated hops; use profile:<alias> or raw:<target> for explicit type")
routeCmd.AddCommand(routeShowCmd) routeCmd.AddCommand(routeShowCmd)
routeCmd.AddCommand(routeSetCmd) routeCmd.AddCommand(routeSetCmd)
routeCmd.AddCommand(routeClearCmd) routeCmd.AddCommand(routeClearCmd)

27
cmd/runtime_helpers.go Normal file
View File

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

View File

@ -17,19 +17,6 @@ func runTUI() error {
return fmt.Errorf("load servers: %w", err) 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) { tui.ListServers = func() ([]*model.Server, error) {
return appDB.ListServers() return appDB.ListServers()
} }
@ -37,56 +24,81 @@ func runTUI() error {
return appDB.SearchServers(query) return appDB.SearchServers(query)
} }
tui.DeleteServer = func(alias string) error { tui.DeleteServer = func(alias string) error {
server, err := appDB.GetServer(alias)
if err != nil {
return err
}
if err := appDB.DeleteServer(alias); err != nil { if err := appDB.DeleteServer(alias); err != nil {
return err return err
} }
v := getOrCreateVault() v := getOrCreateVault()
if v.IsUnlocked() { if v.IsUnlocked() {
cleanupServerSecrets(v, alias) cleanupServerSecretsForServer(v, server)
if err := v.Save(); err != nil { 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 fmt.Errorf("save vault after cleanup: %w", err)
} }
} }
return nil return nil
} }
tui.TestConnection = func(server *model.Server) (bool, string) { 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) { 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 { 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 lookupAlias := server.Alias
if oldAlias != "" { if oldAlias != "" {
lookupAlias = oldAlias lookupAlias = oldAlias
} }
existing, _ := appDB.GetServer(lookupAlias) existing, _ := appDB.GetServer(lookupAlias)
var original *model.Server
if existing != nil { if existing != nil {
original = cloneServer(existing)
server.ID = existing.ID server.ID = existing.ID
if err := appDB.UpdateServerByAlias(existing.Alias, server); err != nil { if err := appDB.UpdateServerByAlias(existing.Alias, server); err != nil {
return err return err
} }
return appDB.SetServerTags(existing.ID, server.Tags) } else {
}
if err := appDB.CreateServer(server); err != nil { if err := appDB.CreateServer(server); err != nil {
return err return err
} }
return appDB.SetServerTags(server.ID, server.Tags) }
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
}
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) { tui.GetGroups = func() ([]string, error) {
return appDB.GetGroups() return appDB.GetGroups()
} }
tui.ResolveRouteAlias = func(alias string) (int64, bool) {
return appDB.ResolveAlias(alias)
}
tui.RenameGroup = func(oldName, newName string) error { tui.RenameGroup = func(oldName, newName string) error {
return appDB.RenameGroup(oldName, newName) return appDB.RenameGroup(oldName, newName)
} }
@ -123,7 +135,7 @@ func runTUI() error {
if err != nil { if err != nil {
return "", err 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) { tui.ListForwards = func(serverID int64) ([]*model.Forward, error) {
return appDB.GetForwards(serverID) return appDB.GetForwards(serverID)
@ -157,7 +169,11 @@ func runTUI() error {
if !v.IsUnlocked() { if !v.IsUnlocked() {
return false 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 // 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) 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) fmt.Fprintf(os.Stderr, "Connection error: %v\n", err)
} else { } else {
fmt.Println("Connection closed.") fmt.Println("Connection closed.")
@ -210,7 +226,7 @@ func runTUI() error {
continue continue
} }
fmt.Printf("Running template %q on %s...\n", result.TemplateName, fresh.Alias) 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) fmt.Fprintf(os.Stderr, "Command error on %s: %v\n", fresh.Alias, err)
} }
appDB.UpdateLastConnected(fresh.Alias) appDB.UpdateLastConnected(fresh.Alias)
@ -279,7 +295,7 @@ func runTUI() error {
servers, _ = appDB.ListServers() servers, _ = appDB.ListServers()
continue continue
} }
state, err := tunnelpkg.Start(cfg, fresh, forwards, forwardOnly) state, err := tunnelpkg.StartResolved(cfg, fresh, forwards, forwardOnly, dbProfileResolver)
if err != nil { if err != nil {
fmt.Fprintf(os.Stderr, "Start tunnel: %v\n", err) fmt.Fprintf(os.Stderr, "Start tunnel: %v\n", err)
} else { } else {
@ -295,8 +311,7 @@ func runTUI() error {
fmt.Printf("Starting session to %s...\n", fresh.Alias) fmt.Printf("Starting session to %s...\n", fresh.Alias)
} }
sshArgs := ssh.BuildSSHArgs(fresh, forwards, forwardOnly) if err := ssh.ConnectWithForwardsResolved(cfg, fresh, forwards, forwardOnly, dbProfileResolver, serverVaultFunc(fresh)); err != nil {
if err := ssh.ConnectWithArgs(cfg, sshArgs, vaultFunc, fresh); err != nil {
fmt.Fprintf(os.Stderr, "Tunnel error: %v\n", err) fmt.Fprintf(os.Stderr, "Tunnel error: %v\n", err)
} else { } else {
fmt.Println("Tunnel closed.") fmt.Println("Tunnel closed.")

View File

@ -34,7 +34,7 @@ var tunnelCmd = &cobra.Command{
if err := validateBackgroundTunnel(server, forwards); err != nil { if err := validateBackgroundTunnel(server, forwards); err != nil {
return err return err
} }
state, err := tunnelpkg.Start(cfg, server, forwards, true) state, err := tunnelpkg.StartResolved(cfg, server, forwards, true, dbProfileResolver)
if err != nil { if err != nil {
return err return err
} }
@ -46,31 +46,17 @@ var tunnelCmd = &cobra.Command{
return fmt.Errorf("no forwards configured for %s", alias) 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 { if len(forwards) > 0 {
fmt.Printf("Starting tunnel to %s with %d forward(s)...\n", alias, len(forwards)) fmt.Printf("Starting tunnel to %s with %d forward(s)...\n", alias, len(forwards))
} else { } else {
fmt.Printf("Starting session to %s...\n", alias) fmt.Printf("Starting session to %s...\n", alias)
} }
sshArgs := ssh.BuildSSHArgs(server, forwards, forwardsOnly)
if forwardsOnly { if forwardsOnly {
fmt.Printf("Tunnel mode (ssh -N). Press Ctrl+C to exit.\n") 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))
}, },
} }

View File

@ -245,6 +245,7 @@ type TunnelState struct {
Name string `json:"name"` Name string `json:"name"`
PID int `json:"pid"` PID int `json:"pid"`
ForwardIDs []int64 `json:"forward_ids"` ForwardIDs []int64 `json:"forward_ids"`
ConfigPath string `json:"config_path,omitempty"`
StartedAt time.Time `json:"started_at"` StartedAt time.Time `json:"started_at"`
LastError string `json:"last_error"` LastError string `json:"last_error"`
} }

View File

@ -266,7 +266,11 @@ func BuildForwardArgs(forwards []*model.Forward, exitOnForwardFailure bool) []st
func BuildSSHArgs(server *model.Server, forwards []*model.Forward, forwardOnly bool) []string { func BuildSSHArgs(server *model.Server, forwards []*model.Forward, forwardOnly bool) []string {
var args []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 != "" { if server.IdentityFile != "" {
args = append(args, "-i", 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") 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) args = append(args, target)
return args return args

View File

@ -180,6 +180,7 @@ var (
UpdateTestResult func(alias string, status model.TestStatus, testErr string) error UpdateTestResult func(alias string, status model.TestStatus, testErr string) error
HasSecret func(alias string, secretType string) bool HasSecret func(alias string, secretType string) bool
GetGroups func() ([]string, error) GetGroups func() ([]string, error)
ResolveRouteAlias func(alias string) (int64, bool)
RenameGroup func(oldName, newName string) error RenameGroup func(oldName, newName string) error
DeleteGroup func(name string) error DeleteGroup func(name string) error
ListTags func() ([]string, error) ListTags func() ([]string, error)
@ -769,12 +770,14 @@ func (m *tuiModel) updateList(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
case tea.KeyCtrlA: case tea.KeyCtrlA:
m.form = newFormModel(m.width, m.height) m.form = newFormModel(m.width, m.height)
m.form.setRouteProfiles(m.servers)
m.screen = screenForm m.screen = screenForm
return m, nil return m, nil
case tea.KeyCtrlE: case tea.KeyCtrlE:
if item, ok := m.list.SelectedItem().(serverItem); ok { if item, ok := m.list.SelectedItem().(serverItem); ok {
m.form = newEditFormModel(item.server, m.width, m.height) m.form = newEditFormModel(item.server, m.width, m.height)
m.form.setRouteProfiles(m.servers)
m.screen = screenForm m.screen = screenForm
} }
return m, nil return m, nil
@ -1355,6 +1358,7 @@ func (m *tuiModel) updateActionMenu(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
case "route": case "route":
if item, ok := m.list.SelectedItem().(serverItem); ok { if item, ok := m.list.SelectedItem().(serverItem); ok {
m.form = newEditFormModel(item.server, m.width, m.height) m.form = newEditFormModel(item.server, m.width, m.height)
m.form.setRouteProfiles(m.servers)
m.form.focusIdx = 7 m.form.focusIdx = 7
m.form.updateFocus() m.form.updateFocus()
m.screen = screenForm m.screen = screenForm
@ -1372,6 +1376,7 @@ func (m *tuiModel) updateActionMenu(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
case "edit": case "edit":
if item, ok := m.list.SelectedItem().(serverItem); ok { if item, ok := m.list.SelectedItem().(serverItem); ok {
m.form = newEditFormModel(item.server, m.width, m.height) m.form = newEditFormModel(item.server, m.width, m.height)
m.form.setRouteProfiles(m.servers)
m.screen = screenForm m.screen = screenForm
m.actionMenu = nil m.actionMenu = nil
return m, nil return m, nil
@ -1445,6 +1450,21 @@ func (m *tuiModel) updateForwardList(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
if m.forwardScreen != nil { if m.forwardScreen != nil {
return m, m.forwardScreen.editSelected() 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: case tea.KeyRunes:
switch msg.String() { switch msg.String() {
case "a", "A": case "a", "A":

View File

@ -23,6 +23,22 @@ func (i groupItem) Title() string { return i.name }
func (i groupItem) Description() string { return "" } func (i groupItem) Description() string { return "" }
func (i groupItem) FilterValue() string { return i.name } 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 { func newStringList(values []string, title string, width, height int) list.Model {
items := make([]list.Item, len(values)) items := make([]list.Item, len(values))
for i, value := range values { for i, value := range values {
@ -37,6 +53,28 @@ func newStringList(values []string, title string, width, height int) list.Model
return l 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 --- // --- Form model ---
type formModel struct { type formModel struct {
@ -63,6 +101,11 @@ type formModel struct {
showGroupList bool showGroupList bool
authList list.Model authList list.Model
showAuthList bool showAuthList bool
routeProfiles []*model.Server
routeList list.Model
showRouteList bool
routePane int // 0=current route, 1=available profiles
routeCursor int
initial formSnapshot initial formSnapshot
} }
@ -81,7 +124,7 @@ func newFormModel(w, h int) *formModel {
"User", "User",
"Auth Method (password/key/key_passphrase/agent)", "Auth Method (password/key/key_passphrase/agent)",
"Identity File", "Identity File",
"Route hops (comma-separated, or pick from profiles)", "Route (direct / ordered bastions)",
"Group (type new or pick from list)", "Group (type new or pick from list)",
"Notes", "Notes",
"Startup Command", "Startup Command",
@ -128,7 +171,6 @@ func newFormModel(w, h int) *formModel {
fm.groupList = newStringList(groups, "Select group", 30, 8) fm.groupList = newStringList(groups, "Select group", 30, 8)
} }
} }
fm.updateFocus() fm.updateFocus()
fm.initial = fm.snapshot() fm.initial = fm.snapshot()
return fm return fm
@ -150,8 +192,8 @@ func placeholderForLabel(label string) string {
return "key" return "key"
case "Identity File": case "Identity File":
return "~/.ssh/id_ed25519" return "~/.ssh/id_ed25519"
case "Route hops (comma-separated, or pick from profiles)": case "Route (direct / ordered bastions)":
return "bastion, dmz-gw" return "profile:bastion, raw:user@gw.example"
case "Group (type new or pick from list)": case "Group (type new or pick from list)":
return "KP" return "KP"
case "Notes": case "Notes":
@ -177,17 +219,9 @@ func newEditFormModel(s *model.Server, w, h int) *formModel {
fm.inputs[5].SetValue(string(s.AuthMethod)) fm.inputs[5].SetValue(string(s.AuthMethod))
fm.inputs[6].SetValue(s.IdentityFile) fm.inputs[6].SetValue(s.IdentityFile)
// Populate Route hops // Store an unambiguous route spec in the editable field.
if len(s.Route.Hops) > 0 { if len(s.Route.Hops) > 0 {
hopStrs := make([]string, len(s.Route.Hops)) fm.inputs[7].SetValue(model.FormatRouteSpec(s.Route))
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, ", "))
} else if s.ProxyJump != "" { } else if s.ProxyJump != "" {
fm.inputs[7].SetValue(s.ProxyJump) fm.inputs[7].SetValue(s.ProxyJump)
} }
@ -236,6 +270,136 @@ func (fm *formModel) Dirty() bool {
return false 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 { func (fm *formModel) Init() tea.Cmd {
return nil return nil
} }
@ -276,6 +440,10 @@ func (fm *formModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
return fm, cmd return fm, cmd
} }
if fm.showRouteList {
return fm.updateRouteEditor(msg)
}
if fm.showGroupList { if fm.showGroupList {
switch msg := msg.(type) { switch msg := msg.(type) {
case tea.KeyMsg: case tea.KeyMsg:
@ -342,6 +510,15 @@ func (fm *formModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
fm.showAuthList = true fm.showAuthList = true
return fm, nil 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 { if len(msg.Runes) == 1 && msg.Runes[0] == '/' && !msg.Alt && fm.focusIdx == 8 && len(fm.groups) > 0 {
fm.showGroupList = true fm.showGroupList = true
return fm, nil return fm, nil
@ -414,6 +591,8 @@ func (fm *formModel) applySaveError(err error) {
fm.focusIdx = 2 fm.focusIdx = 2
case strings.Contains(message, "port"): case strings.Contains(message, "port"):
fm.focusIdx = 3 fm.focusIdx = 3
case strings.Contains(message, "route"):
fm.focusIdx = 7
default: default:
return return
} }
@ -450,6 +629,9 @@ func (fm *formModel) labelAt(index int) string {
if index == 5 { if index == 5 {
return "Auth Method (/ pick)" return "Auth Method (/ pick)"
} }
if index == 7 {
return "Route (/ edit)"
}
if index == 8 { if index == 8 {
if len(fm.groups) > 0 { if len(fm.groups) > 0 {
return "Group (/ pick)" return "Group (/ pick)"
@ -471,7 +653,11 @@ func (fm *formModel) runTest() tea.Cmd {
fm.testing = false fm.testing = false
return func() tea.Msg { return testDoneMsg{ok: false, err: err.Error()} } 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() pw := fm.password.Value()
return tea.Batch( return tea.Batch(
@ -499,7 +685,10 @@ func (fm *formModel) runSave() tea.Cmd {
if _, err := parsePort(fm.inputs[3].Value()); err != nil { if _, err := parsePort(fm.inputs[3].Value()); err != nil {
return func() tea.Msg { return saveDoneMsg{err: err} } 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() pw := fm.password.Value()
return tea.Batch( return tea.Batch(
@ -521,55 +710,54 @@ func (fm *formModel) runSave() tea.Cmd {
) )
} }
// parseRouteHops parses the route hops input string into a model.Route. // parseRouteHops parses an explicit route spec. Exact known aliases become
// Format: comma-separated list of aliases or raw addresses. // stable profile references; unknown unprefixed values remain raw OpenSSH targets.
func parseRouteHops(input string) model.Route { func parseRouteHops(input string) (model.Route, error) {
input = strings.TrimSpace(input) return model.ParseRouteSpec(input, ResolveRouteAlias)
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}
} }
func (fm *formModel) buildServer() *model.Server { func (fm *formModel) buildServerValidated() (*model.Server, error) {
port, _ := parsePort(fm.inputs[3].Value()) port, err := parsePort(fm.inputs[3].Value())
authMethod := model.AuthMethod(fm.inputs[5].Value()) if err != nil {
return nil, err
}
authMethod := model.AuthMethod(strings.TrimSpace(fm.inputs[5].Value()))
if authMethod == "" { if authMethod == "" {
authMethod = model.AuthKey authMethod = model.AuthKey
} }
route, err := fm.parseRouteInput()
route := parseRouteHops(fm.inputs[7].Value()) if err != nil {
return nil, fmt.Errorf("route: %w", err)
return &model.Server{ }
Alias: fm.inputs[0].Value(), server := &model.Server{
DisplayName: fm.inputs[1].Value(), Alias: strings.TrimSpace(fm.inputs[0].Value()),
Host: fm.inputs[2].Value(), DisplayName: strings.TrimSpace(fm.inputs[1].Value()),
Host: strings.TrimSpace(fm.inputs[2].Value()),
Port: port, Port: port,
User: fm.inputs[4].Value(), User: strings.TrimSpace(fm.inputs[4].Value()),
AuthMethod: authMethod, AuthMethod: authMethod,
IdentityFile: fm.inputs[6].Value(), IdentityFile: strings.TrimSpace(fm.inputs[6].Value()),
ProxyJump: route.ProxyJumpString(), ProxyJump: route.ProxyJumpString(),
Route: route, Route: route,
GroupName: fm.inputs[8].Value(), GroupName: strings.TrimSpace(fm.inputs[8].Value()),
Notes: fm.inputs[9].Value(), Notes: fm.inputs[9].Value(),
StartupCommand: fm.inputs[10].Value(), StartupCommand: fm.inputs[10].Value(),
Tags: splitCSV(fm.inputs[11].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) { func parsePort(value string) (int, error) {
@ -589,6 +777,9 @@ func (fm *formModel) View() string {
if fm.edit { if fm.edit {
title = "Edit Server: " + fm.server.Alias title = "Edit Server: " + fm.server.Alias
} }
if fm.showRouteList {
return fm.routeEditorView(title)
}
if fm.showAuthList || fm.showGroupList { if fm.showAuthList || fm.showGroupList {
var dropdown list.Model var dropdown list.Model
fieldIndex := 8 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:<user@host:port> 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 { func (fm *formModel) formStatusLine() string {
if fm.err != nil { if fm.err != nil {
return errorStyle.Render(fmt.Sprintf("✗ Error: %v", fm.err)) return errorStyle.Render(fmt.Sprintf("✗ Error: %v", fm.err))

View File

@ -81,8 +81,8 @@ func TestForwardValidationMovesFocusToInvalidPort(t *testing.T) {
fm.inputs[3].SetValue("5432") fm.inputs[3].SetValue("5432")
updated, _ := fm.Update(fm.runSave()()) updated, _ := fm.Update(fm.runSave()())
fm = updated.(*forwardFormModel) fm = updated.(*forwardFormModel)
if fm.focusIdx != 6 { if fm.focusIdx != 7 {
t.Fatalf("invalid listen port focus = %d, want 6", fm.focusIdx) t.Fatalf("invalid listen port focus = %d, want 7", fm.focusIdx)
} }
view := fm.View() view := fm.View()
if !strings.Contains(view, "Listen Port") || !strings.Contains(view, "must be a number") { if !strings.Contains(view, "Listen Port") || !strings.Contains(view, "must be a number") {

View File

@ -112,6 +112,7 @@ func (m *forwardScreenModel) View() string {
footer: []helpItem{ footer: []helpItem{
{Key: "Ctrl+A (a)", Action: "add"}, {Key: "Ctrl+A (a)", Action: "add"},
{Key: "Ctrl+E/Enter", Action: "edit"}, {Key: "Ctrl+E/Enter", Action: "edit"},
{Key: "Space", Action: "enable/disable"},
{Key: "Ctrl+D (d)", Action: "delete"}, {Key: "Ctrl+D (d)", Action: "delete"},
{Key: "Ctrl+H", Action: "help"}, {Key: "Ctrl+H", Action: "help"},
{Key: "Esc", Action: "back"}, {Key: "Esc", Action: "back"},
@ -221,6 +222,7 @@ type forwardFormModel struct {
nameInput textinput.Model nameInput textinput.Model
descInput textinput.Model descInput textinput.Model
typeIdx int // 0=local, 1=remote, 2=socks typeIdx int // 0=local, 1=remote, 2=socks
enabled bool
width int width int
height int height int
initial forwardFormSnapshot initial forwardFormSnapshot
@ -231,6 +233,7 @@ type forwardFormSnapshot struct {
description string description string
values []string values []string
forwardType model.ForwardType forwardType model.ForwardType
enabled bool
} }
var forwardTypes = []forwardTypeItem{ var forwardTypes = []forwardTypeItem{
@ -262,6 +265,7 @@ func newForwardFormModel(serverID int64, w, h int) *forwardFormModel {
focusIdx: 0, focusIdx: 0,
currentType: model.ForwardLocal, currentType: model.ForwardLocal,
typeIdx: 0, typeIdx: 0,
enabled: true,
nameInput: nameInput, nameInput: nameInput,
descInput: descInput, descInput: descInput,
width: w, width: w,
@ -280,6 +284,7 @@ func newForwardEditModel(serverID int64, fwd *model.Forward, w, h int) *forwardF
fm.descInput.SetValue(fwd.Description) fm.descInput.SetValue(fwd.Description)
fm.currentType = fwd.Type fm.currentType = fwd.Type
fm.typeIdx = typeIndex(fwd.Type) fm.typeIdx = typeIndex(fwd.Type)
fm.enabled = fwd.Enabled
if fwd.Type == model.ForwardRemote { if fwd.Type == model.ForwardRemote {
fm.inputs[0].SetValue(fwd.RemoteAddr) fm.inputs[0].SetValue(fwd.RemoteAddr)
fm.inputs[1].SetValue(strconv.Itoa(fwd.RemotePort)) fm.inputs[1].SetValue(strconv.Itoa(fwd.RemotePort))
@ -306,12 +311,13 @@ func (fm *forwardFormModel) snapshot() forwardFormSnapshot {
description: fm.descInput.Value(), description: fm.descInput.Value(),
values: values, values: values,
forwardType: fm.currentType, forwardType: fm.currentType,
enabled: fm.enabled,
} }
} }
func (fm *forwardFormModel) Dirty() bool { func (fm *forwardFormModel) Dirty() bool {
current := fm.snapshot() 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 return true
} }
for i := range current.values { for i := range current.values {
@ -382,7 +388,7 @@ func (fm *forwardFormModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
switch msg.Type { switch msg.Type {
case tea.KeyTab: case tea.KeyTab:
fm.focusIdx++ 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 { if fm.focusIdx >= total {
fm.focusIdx = 0 fm.focusIdx = 0
} }
@ -391,7 +397,7 @@ func (fm *forwardFormModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
case tea.KeyShiftTab: case tea.KeyShiftTab:
fm.focusIdx-- fm.focusIdx--
if fm.focusIdx < 0 { if fm.focusIdx < 0 {
total := 2 + 3 + len(fm.visibleFields()) + 1 total := 2 + 3 + 1 + len(fm.visibleFields()) + 1
fm.focusIdx = total - 1 fm.focusIdx = total - 1
} }
fm.updateFocus() fm.updateFocus()
@ -405,7 +411,11 @@ func (fm *forwardFormModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
fm.updateFocus() fm.updateFocus()
return fm, nil 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() return fm, fm.runSave()
} }
fm.focusIdx++ fm.focusIdx++
@ -415,7 +425,7 @@ func (fm *forwardFormModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
return fm, nil return fm, nil
case tea.KeyDown: case tea.KeyDown:
fm.focusIdx++ fm.focusIdx++
total := 2 + 3 + len(fm.visibleFields()) + 1 total := 2 + 3 + 1 + len(fm.visibleFields()) + 1
if fm.focusIdx >= total { if fm.focusIdx >= total {
fm.focusIdx = 0 fm.focusIdx = 0
} }
@ -424,12 +434,16 @@ func (fm *forwardFormModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
case tea.KeyUp: case tea.KeyUp:
fm.focusIdx-- fm.focusIdx--
if fm.focusIdx < 0 { if fm.focusIdx < 0 {
total := 2 + 3 + len(fm.visibleFields()) + 1 total := 2 + 3 + 1 + len(fm.visibleFields()) + 1
fm.focusIdx = total - 1 fm.focusIdx = total - 1
} }
fm.updateFocus() fm.updateFocus()
return fm, nil return fm, nil
case tea.KeyRunes: 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. // 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 { if fm.focusIdx >= 2 && fm.focusIdx < 2+len(forwardTypes) && len(msg.Runes) == 1 {
switch msg.Runes[0] { switch msg.Runes[0] {
@ -465,8 +479,8 @@ func (fm *forwardFormModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
return fm, cmd return fm, cmd
} }
visible := fm.visibleFields() visible := fm.visibleFields()
if fm.focusIdx >= 2+3 && fm.focusIdx < 2+3+len(visible) { if fm.focusIdx >= 2+3+1 && fm.focusIdx < 2+3+1+len(visible) {
fieldIdx := visible[fm.focusIdx-(2+3)] fieldIdx := visible[fm.focusIdx-(2+3+1)]
var cmd tea.Cmd var cmd tea.Cmd
fm.inputs[fieldIdx], cmd = fm.inputs[fieldIdx].Update(msg) fm.inputs[fieldIdx], cmd = fm.inputs[fieldIdx].Update(msg)
return fm, cmd return fm, cmd
@ -485,7 +499,7 @@ func (fm *forwardFormModel) updateFocus() {
fm.inputs[i].Prompt = blurredStyle.Render(fm.labelForField(i) + ": ") 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 { switch {
case fm.focusIdx == 0: case fm.focusIdx == 0:
fm.nameInput.Focus() fm.nameInput.Focus()
@ -495,98 +509,96 @@ func (fm *forwardFormModel) updateFocus() {
fm.descInput.Prompt = focusedStyle.Render("Description> ") fm.descInput.Prompt = focusedStyle.Render("Description> ")
case fm.focusIdx >= 2 && fm.focusIdx < 2+3: case fm.focusIdx >= 2 && fm.focusIdx < 2+3:
// Type selector focused — no input to focus // 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() visible := fm.visibleFields()
fieldIdx := visible[fm.focusIdx-(2+3)] fieldIdx := visible[fm.focusIdx-(2+3+1)]
fm.inputs[fieldIdx].Focus() fm.inputs[fieldIdx].Focus()
fm.inputs[fieldIdx].Prompt = focusedStyle.Render(fm.labelForField(fieldIdx) + "> ") fm.inputs[fieldIdx].Prompt = focusedStyle.Render(fm.labelForField(fieldIdx) + "> ")
} }
} }
func (fm *forwardFormModel) runSave() tea.Cmd { func (fm *forwardFormModel) buildForwardFromForm() (*model.Forward, error) {
return func() tea.Msg {
name := strings.TrimSpace(fm.nameInput.Value()) name := strings.TrimSpace(fm.nameInput.Value())
desc := strings.TrimSpace(fm.descInput.Value())
localAddr, remoteAddr := "", ""
localPort, remotePort := 0, 0
var err error
if name == "" { if name == "" {
return saveDoneMsg{err: fmt.Errorf("name is required")} return nil, fmt.Errorf("name is required")
} }
switch fm.currentType { forward := &model.Forward{
case model.ForwardLocal: ID: fm.editID,
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, ServerID: fm.serverID,
Name: name, Name: name,
Description: desc, Description: strings.TrimSpace(fm.descInput.Value()),
Type: fm.currentType, Type: fm.currentType,
LocalAddr: localAddr, Enabled: fm.enabled,
LocalPort: localPort, }
RemoteAddr: remoteAddr, var err error
RemotePort: remotePort, switch fm.currentType {
Enabled: true, 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 {
forward, err := fm.buildForwardFromForm()
if err != nil {
return saveDoneMsg{err: err}
}
if fm.editMode { if fm.editMode {
fwd.ID = fm.editID
if UpdateForward == nil { if UpdateForward == nil {
return saveDoneMsg{err: fmt.Errorf("update not available")} return saveDoneMsg{err: fmt.Errorf("update not available")}
} }
return saveDoneMsg{err: UpdateForward(fwd)} return saveDoneMsg{err: UpdateForward(forward)}
} }
if SaveForward == nil { if SaveForward == nil {
return saveDoneMsg{err: fmt.Errorf("forward storage is unavailable")} return saveDoneMsg{err: fmt.Errorf("forward storage is unavailable")}
} }
err = SaveForward(fwd) return saveDoneMsg{err: SaveForward(forward)}
return saveDoneMsg{err: err}
} }
} }
@ -612,7 +624,7 @@ func (fm *forwardFormModel) applySaveError(err error) {
fieldIndex = 3 fieldIndex = 3
} }
if fieldIndex >= 0 { if fieldIndex >= 0 {
fm.focusIdx = 2 + len(forwardTypes) + fieldIndex fm.focusIdx = 2 + len(forwardTypes) + 1 + fieldIndex
fm.updateFocus() fm.updateFocus()
} }
} }
@ -658,6 +670,15 @@ func (fm *forwardFormModel) View() string {
if width >= 100 { if width >= 100 {
lines = append(lines, helpStyle.Copy().MarginLeft(0).Render(forwardTypes[fm.typeIdx].description)) 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() visible := fm.visibleFields()
for _, idx := range visible { for _, idx := range visible {
lines = append(lines, fm.inputs[idx].View()) 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.")) 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() != "" { 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()} if fwd, err := fm.buildForwardFromForm(); err == nil {
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" preview := "Preview ssh " + strings.Join(fwd.ForwardSSHArgs(), " ") + " -o ExitOnForwardFailure=yes"
lines = append(lines, wrapCells(preview, contentWidth)...) lines = append(lines, wrapCells(preview, contentWidth)...)
} }
total := 2 + 3 + len(visible) + 1 }
total := 2 + 3 + 1 + len(visible) + 1
button := " [ Save ]" button := " [ Save ]"
if fm.focusIdx == total-1 { if fm.focusIdx == total-1 {
button = selectedStyle.Render("> [ Save ]") button = selectedStyle.Render("> [ Save ]")
@ -691,7 +711,7 @@ func (fm *forwardFormModel) View() string {
{Key: "Tab/↓", Action: "next"}, {Key: "Tab/↓", Action: "next"},
{Key: "↑", Action: "prev"}, {Key: "↑", Action: "prev"},
{Key: "1/2/3", Action: "select type"}, {Key: "1/2/3", Action: "select type"},
{Key: "Enter", Action: "save"}, {Key: "Enter/Space", Action: "toggle/save"},
{Key: "Ctrl+H", Action: "help"}, {Key: "Ctrl+H", Action: "help"},
{Key: "Esc", Action: "back"}, {Key: "Esc", Action: "back"},
}, },

View File

@ -11,7 +11,7 @@ import (
func TestForwardFormDigitsReachFocusedInput(t *testing.T) { func TestForwardFormDigitsReachFocusedInput(t *testing.T) {
fm := newForwardFormModel(1, 100, 30) fm := newForwardFormModel(1, 100, 30)
fm.focusIdx = 2 + len(forwardTypes) + 1 fm.focusIdx = 2 + len(forwardTypes) + 2
fm.updateFocus() fm.updateFocus()
for _, digit := range []rune{'1', '2', '3'} { 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) { func TestForwardFormDigitShortcutsWorkOnTypeSelector(t *testing.T) {
tests := []struct { tests := []struct {
digit rune digit rune

View File

@ -88,8 +88,14 @@ func Get(id int64) *model.TunnelState {
return states[id] 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) { 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() mu.Lock()
defer mu.Unlock() defer mu.Unlock()
@ -105,9 +111,11 @@ func Start(cfg *config.Config, server *model.Server, forwards []*model.Forward,
} }
} }
sshArgs := ssh.BuildSSHArgs(server, active, forwardOnly) invocation, err := ssh.PrepareSSHInvocation(server, active, forwardOnly, resolve)
args := make([]string, len(sshArgs)) if err != nil {
copy(args, sshArgs) return nil, err
}
args := append([]string(nil), invocation.Args...)
cmd := exec.Command(cfg.SSH.Binary, args...) cmd := exec.Command(cfg.SSH.Binary, args...)
cmd.Env = os.Environ() cmd.Env = os.Environ()
@ -116,6 +124,7 @@ func Start(cfg *config.Config, server *model.Server, forwards []*model.Forward,
cmd.Stderr = nil cmd.Stderr = nil
if err := cmd.Start(); err != nil { if err := cmd.Start(); err != nil {
invocation.Cleanup()
return nil, fmt.Errorf("start tunnel: %w", err) 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), Name: fmt.Sprintf("Tunnel to %s", server.Alias),
PID: cmd.Process.Pid, PID: cmd.Process.Pid,
ForwardIDs: forwardIDs, ForwardIDs: forwardIDs,
ConfigPath: invocation.ConfigPath,
StartedAt: time.Now(), StartedAt: time.Now(),
} }
@ -139,10 +149,12 @@ func Start(cfg *config.Config, server *model.Server, forwards []*model.Forward,
if err := saveStates(); err != nil { if err := saveStates(); err != nil {
delete(states, id) delete(states, id)
_ = cmd.Process.Kill() _ = cmd.Process.Kill()
invocation.Cleanup()
return nil, fmt.Errorf("save tunnel state: %w", err) return nil, fmt.Errorf("save tunnel state: %w", err)
} }
if err := cmd.Process.Release(); err != nil { if err := cmd.Process.Release(); err != nil {
delete(states, id) delete(states, id)
invocation.Cleanup()
return nil, fmt.Errorf("release tunnel process: %w", err) 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) delete(states, id)
return saveStates() return saveStates()
} }
@ -182,6 +197,9 @@ func StopAll() error {
proc.Kill() proc.Kill()
} }
} }
if state.ConfigPath != "" {
_ = os.Remove(state.ConfigPath)
}
delete(states, id) delete(states, id)
} }
return saveStates() return saveStates()