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
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:<alias> or raw:<target>)", "")
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:<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.displayName, "display-name", "", "Display name")
addCmd.Flags().StringVar(&addFlags.notes, "notes", "", "Notes")

View File

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

View File

@ -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 {
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 {
var secret []byte
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))
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"
}
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 password: %w", err)
return fmt.Errorf("read %s: %w", label, err)
}
if len(pw) > 0 {
secret = string(pw)
if len(secret) == 0 {
return fmt.Errorf("%s cannot be empty", label)
}
} 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)
defer func() {
for i := range secret {
secret[i] = 0
}
if len(pw) > 0 {
secret = string(pw)
}()
}
}
if err := syncServerSecrets(v, alias, server, secret); 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 {
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 &copyServer
}
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:<alias>, raw:<target>, 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")
}

View File

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

View File

@ -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,31 +18,30 @@ 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 := 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("Mode: %s\n", server.Route.RouteMode())
fmt.Printf("ProxyJump: %s\n", server.Route.ProxyJumpString())
if server.Route.HasProfileLinks() {
fmt.Printf("Spec: %s\n", model.FormatRouteSpec(server.Route))
fmt.Println("Hops:")
for _, h := range server.Route.Hops {
if h.IsProfile {
fmt.Printf(" - %s (profile)\n", h.Alias)
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(" - %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
},
}
@ -52,47 +49,38 @@ var routeShowCmd = &cobra.Command{
var routeSetCmd = &cobra.Command{
Use: "set <alias>",
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 {
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.TrimSpace(jumps) == "" {
return fmt.Errorf("--jumps is required unless --mode=direct/clear")
}
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})
route, err := parseRouteSpec(jumps)
if err != nil {
return err
}
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 {
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:<alias> or raw:<target> for explicit type")
routeCmd.AddCommand(routeShowCmd)
routeCmd.AddCommand(routeSetCmd)
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)
}
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
}
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) {
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.")

View File

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

View File

@ -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"`
}

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 {
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

View File

@ -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":

View File

@ -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:<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 {
if fm.err != nil {
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")
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") {

View File

@ -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) runSave() tea.Cmd {
return func() tea.Msg {
func (fm *forwardFormModel) buildForwardFromForm() (*model.Forward, error) {
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")}
return nil, fmt.Errorf("name is required")
}
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{
forward := &model.Forward{
ID: fm.editID,
ServerID: fm.serverID,
Name: name,
Description: desc,
Description: strings.TrimSpace(fm.descInput.Value()),
Type: fm.currentType,
LocalAddr: localAddr,
LocalPort: localPort,
RemoteAddr: remoteAddr,
RemotePort: remotePort,
Enabled: true,
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 {
forward, err := fm.buildForwardFromForm()
if err != nil {
return saveDoneMsg{err: err}
}
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)
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"},
},

View File

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

View File

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