feat: unify route workflows and forwarding UX
This commit is contained in:
parent
3cc21a7b22
commit
f05e8e8e84
79
cmd/add.go
79
cmd/add.go
|
|
@ -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")
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
|
||||||
148
cmd/edit.go
148
cmd/edit.go
|
|
@ -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 {
|
||||||
server.Tags = strings.Split(parsedTags, ",")
|
if strings.TrimSpace(parsedTags) == "" {
|
||||||
|
server.Tags = nil
|
||||||
|
} else {
|
||||||
|
server.Tags = strings.Split(parsedTags, ",")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := model.ValidateServerBasics(server); err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
if parsedAuth != "" && oldAuthMethod != server.AuthMethod {
|
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 {
|
||||||
fmt.Println()
|
case model.AuthPassword, model.AuthKeyPassphrase:
|
||||||
if err != nil {
|
label := "password"
|
||||||
return fmt.Errorf("read password: %w", err)
|
if server.AuthMethod == model.AuthKeyPassphrase {
|
||||||
}
|
label = "key passphrase"
|
||||||
if len(pw) > 0 {
|
|
||||||
secret = string(pw)
|
|
||||||
}
|
|
||||||
} else if server.AuthMethod == model.AuthKeyPassphrase {
|
|
||||||
fmt.Print("Enter key passphrase (stored in vault, input hidden): ")
|
|
||||||
pw, err := term.ReadPassword(int(syscall.Stdin))
|
|
||||||
fmt.Println()
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("read passphrase: %w", err)
|
|
||||||
}
|
|
||||||
if len(pw) > 0 {
|
|
||||||
secret = string(pw)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
fmt.Printf("Enter new %s (stored in vault, input hidden): ", label)
|
||||||
if err := syncServerSecrets(v, alias, server, secret); err != nil {
|
secret, err = term.ReadPassword(int(syscall.Stdin))
|
||||||
return fmt.Errorf("sync vault secrets: %w", err)
|
fmt.Println()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("read %s: %w", label, err)
|
||||||
}
|
}
|
||||||
if err := v.Save(); err != nil {
|
if len(secret) == 0 {
|
||||||
return fmt.Errorf("save vault: %w", err)
|
return fmt.Errorf("%s cannot be empty", label)
|
||||||
}
|
}
|
||||||
|
defer func() {
|
||||||
|
for i := range secret {
|
||||||
|
secret[i] = 0
|
||||||
|
}
|
||||||
|
}()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := appDB.UpdateServer(server); err != nil {
|
if err := appDB.UpdateServerByAlias(alias, server); err != nil {
|
||||||
return fmt.Errorf("update server: %w", err)
|
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 ©Server
|
||||||
|
}
|
||||||
|
|
||||||
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")
|
||||||
}
|
}
|
||||||
|
|
|
||||||
70
cmd/extra.go
70
cmd/extra.go
|
|
@ -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
|
|
||||||
}
|
}
|
||||||
|
|
|
||||||
100
cmd/route.go
100
cmd/route.go
|
|
@ -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,30 +18,29 @@ 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 := fmt.Sprintf("%s@%s:%d", server.User, server.Host, server.Port)
|
target := server.Host
|
||||||
if len(server.Route.Hops) > 0 {
|
if server.User != "" {
|
||||||
fmt.Printf("Route: %s\n", server.Route.DisplaySummary(target))
|
target = server.User + "@" + server.Host
|
||||||
fmt.Printf("Mode: %s\n", server.Route.RouteMode())
|
}
|
||||||
fmt.Printf("ProxyJump: %s\n", server.Route.ProxyJumpString())
|
target = fmt.Sprintf("%s:%d", target, server.Port)
|
||||||
if server.Route.HasProfileLinks() {
|
if len(server.Route.Hops) == 0 {
|
||||||
fmt.Println("Hops:")
|
|
||||||
for _, h := range server.Route.Hops {
|
|
||||||
if h.IsProfile {
|
|
||||||
fmt.Printf(" - %s (profile)\n", h.Alias)
|
|
||||||
} else {
|
|
||||||
fmt.Printf(" - %s (raw)\n", h.Raw)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
} else if server.ProxyJump != "" {
|
|
||||||
fmt.Printf("ProxyJump: %s\n", server.ProxyJump)
|
|
||||||
} else {
|
|
||||||
fmt.Println("Direct connection (no route)")
|
fmt.Println("Direct connection (no route)")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
fmt.Printf("Route: %s\n", server.Route.DisplaySummary(target))
|
||||||
|
fmt.Printf("Mode: %s\n", server.Route.RouteMode())
|
||||||
|
fmt.Printf("Spec: %s\n", model.FormatRouteSpec(server.Route))
|
||||||
|
fmt.Println("Hops:")
|
||||||
|
for index, hop := range server.Route.Hops {
|
||||||
|
if hop.Profile() {
|
||||||
|
fmt.Printf(" %d. %s (sshkeeper profile #%d)\n", index+1, hop.Alias, hop.ServerID)
|
||||||
|
} else {
|
||||||
|
fmt.Printf(" %d. %s (raw OpenSSH target)\n", index+1, hop.Raw)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
return nil
|
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, ":") {
|
|
||||||
hops = append(hops, model.RouteHop{Raw: p, IsProfile: false})
|
|
||||||
} else {
|
|
||||||
hops = append(hops, model.RouteHop{Alias: p, IsProfile: true})
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
server.Route = model.Route{Hops: hops}
|
route, err := parseRouteSpec(jumps)
|
||||||
server.ProxyJump = server.Route.ProxyJumpString()
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
server.Route = route
|
||||||
|
server.ProxyJump = 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)
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
}
|
||||||
87
cmd/tui.go
87
cmd/tui.go
|
|
@ -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 {
|
||||||
|
return err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
if err := appDB.CreateServer(server); err != nil {
|
if err := appDB.SetServerTags(server.ID, server.Tags); err != nil {
|
||||||
|
if original != nil {
|
||||||
|
_ = appDB.UpdateServerByAlias(server.Alias, original)
|
||||||
|
_ = appDB.SetServerTags(original.ID, original.Tags)
|
||||||
|
} else {
|
||||||
|
_ = appDB.DeleteServer(server.Alias)
|
||||||
|
}
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return appDB.SetServerTags(server.ID, server.Tags)
|
|
||||||
|
v := getOrCreateVault()
|
||||||
|
if !v.IsUnlocked() {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if err := syncServerSecrets(v, oldAlias, server, password); err != nil {
|
||||||
|
rollbackSavedServer(server, original)
|
||||||
|
return fmt.Errorf("sync vault secrets: %w", err)
|
||||||
|
}
|
||||||
|
if err := v.Save(); err != nil {
|
||||||
|
rollbackSavedServer(server, original)
|
||||||
|
return fmt.Errorf("save vault: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
tui.GetGroups = func() ([]string, error) {
|
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.")
|
||||||
|
|
|
||||||
|
|
@ -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))
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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"`
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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":
|
||||||
|
|
|
||||||
|
|
@ -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))
|
||||||
|
|
|
||||||
|
|
@ -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") {
|
||||||
|
|
|
||||||
|
|
@ -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) buildForwardFromForm() (*model.Forward, error) {
|
||||||
|
name := strings.TrimSpace(fm.nameInput.Value())
|
||||||
|
if name == "" {
|
||||||
|
return nil, fmt.Errorf("name is required")
|
||||||
|
}
|
||||||
|
forward := &model.Forward{
|
||||||
|
ID: fm.editID,
|
||||||
|
ServerID: fm.serverID,
|
||||||
|
Name: name,
|
||||||
|
Description: strings.TrimSpace(fm.descInput.Value()),
|
||||||
|
Type: fm.currentType,
|
||||||
|
Enabled: fm.enabled,
|
||||||
|
}
|
||||||
|
var err error
|
||||||
|
switch fm.currentType {
|
||||||
|
case model.ForwardLocal:
|
||||||
|
forward.LocalAddr = strings.TrimSpace(fm.inputs[0].Value())
|
||||||
|
if forward.LocalAddr == "" {
|
||||||
|
forward.LocalAddr = "127.0.0.1"
|
||||||
|
}
|
||||||
|
forward.LocalPort, err = parseNamedPort("Listen port", fm.inputs[1].Value())
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
forward.RemoteAddr = strings.TrimSpace(fm.inputs[2].Value())
|
||||||
|
if forward.RemoteAddr == "" {
|
||||||
|
return nil, fmt.Errorf("target host is required for local forward")
|
||||||
|
}
|
||||||
|
forward.RemotePort, err = parseNamedPort("Target port", fm.inputs[3].Value())
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
case model.ForwardRemote:
|
||||||
|
forward.RemoteAddr = strings.TrimSpace(fm.inputs[0].Value())
|
||||||
|
if forward.RemoteAddr == "" {
|
||||||
|
return nil, fmt.Errorf("remote listen address is required")
|
||||||
|
}
|
||||||
|
forward.RemotePort, err = parseNamedPort("Remote listen port", fm.inputs[1].Value())
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
forward.LocalAddr = strings.TrimSpace(fm.inputs[2].Value())
|
||||||
|
if forward.LocalAddr == "" {
|
||||||
|
forward.LocalAddr = "127.0.0.1"
|
||||||
|
}
|
||||||
|
forward.LocalPort, err = parseNamedPort("Local target port", fm.inputs[3].Value())
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
case model.ForwardDynamic:
|
||||||
|
forward.LocalAddr = strings.TrimSpace(fm.inputs[0].Value())
|
||||||
|
if forward.LocalAddr == "" {
|
||||||
|
forward.LocalAddr = "127.0.0.1"
|
||||||
|
}
|
||||||
|
forward.LocalPort, err = parseNamedPort("Listen port", fm.inputs[1].Value())
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return nil, fmt.Errorf("unsupported forward type: %s", fm.currentType)
|
||||||
|
}
|
||||||
|
return forward, nil
|
||||||
|
}
|
||||||
|
|
||||||
func (fm *forwardFormModel) runSave() tea.Cmd {
|
func (fm *forwardFormModel) runSave() tea.Cmd {
|
||||||
return func() tea.Msg {
|
return func() tea.Msg {
|
||||||
name := strings.TrimSpace(fm.nameInput.Value())
|
forward, err := fm.buildForwardFromForm()
|
||||||
desc := strings.TrimSpace(fm.descInput.Value())
|
if err != nil {
|
||||||
localAddr, remoteAddr := "", ""
|
return saveDoneMsg{err: err}
|
||||||
localPort, remotePort := 0, 0
|
|
||||||
var err error
|
|
||||||
|
|
||||||
if name == "" {
|
|
||||||
return saveDoneMsg{err: 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{
|
|
||||||
ServerID: fm.serverID,
|
|
||||||
Name: name,
|
|
||||||
Description: desc,
|
|
||||||
Type: fm.currentType,
|
|
||||||
LocalAddr: localAddr,
|
|
||||||
LocalPort: localPort,
|
|
||||||
RemoteAddr: remoteAddr,
|
|
||||||
RemotePort: remotePort,
|
|
||||||
Enabled: true,
|
|
||||||
}
|
|
||||||
|
|
||||||
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)
|
preview := "Preview ssh " + strings.Join(fwd.ForwardSSHArgs(), " ") + " -o ExitOnForwardFailure=yes"
|
||||||
fmt.Sscanf(fm.inputs[3].Value(), "%d", &fwd.RemotePort)
|
lines = append(lines, wrapCells(preview, contentWidth)...)
|
||||||
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 ]"
|
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"},
|
||||||
},
|
},
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue