From 3cc21a7b22653fcf701493acd0217759a715c2f5 Mon Sep 17 00:00:00 2001 From: mirivlad Date: Sat, 5 Sep 2026 18:55:42 +0800 Subject: [PATCH] refactor: normalize server routes and stable identities --- cmd/secrets.go | 182 ++++++++-- internal/db/db.go | 7 + internal/db/servers.go | 631 +++++++++++++++++++++++++---------- internal/db/servers_test.go | 8 +- internal/db/v040.go | 159 +++++++++ internal/model/server.go | 66 ++-- internal/model/validation.go | 147 ++++++++ internal/ssh/command.go | 227 ++++++------- internal/ssh/planner.go | 225 +++++++++++++ 9 files changed, 1315 insertions(+), 337 deletions(-) create mode 100644 internal/db/v040.go create mode 100644 internal/model/validation.go create mode 100644 internal/ssh/planner.go diff --git a/cmd/secrets.go b/cmd/secrets.go index 18352f5..9cf6f49 100644 --- a/cmd/secrets.go +++ b/cmd/secrets.go @@ -20,58 +20,187 @@ var serverSecretTypes = []string{ secretSudoPassword, } +// serverSecretID is the legacy alias-based key kept for migration/tests. func serverSecretID(alias, secretType string) string { return fmt.Sprintf("server:%s:%s", alias, secretType) } +func stableServerSecretID(serverID int64, secretType string) string { + return fmt.Sprintf("server-id:%d:%s", serverID, secretType) +} + +func getServerSecret(v *vault.Vault, server *model.Server, secretType string) ([]byte, error) { + if server == nil { + return nil, fmt.Errorf("server is required") + } + if server.ID > 0 { + stableID := stableServerSecretID(server.ID, secretType) + if data, err := v.Get(stableID); err == nil { + return data, nil + } + } + legacyID := serverSecretID(server.Alias, secretType) + data, err := v.Get(legacyID) + if err != nil { + return nil, err + } + if server.ID > 0 { + if err := v.Put(stableServerSecretID(server.ID, secretType), secretType, data); err != nil { + return nil, err + } + v.Delete(legacyID) + if err := v.Save(); err != nil { + return nil, fmt.Errorf("save migrated vault secret: %w", err) + } + } + return data, nil +} + +func hasServerSecret(v *vault.Vault, server *model.Server, secretType string) bool { + if server == nil { + return false + } + if server.ID > 0 && v.HasSecret(stableServerSecretID(server.ID, secretType)) { + return true + } + return v.HasSecret(serverSecretID(server.Alias, secretType)) +} + +func cleanupServerSecretsForServer(v *vault.Vault, server *model.Server, legacyAliases ...string) { + if server == nil { + return + } + aliases := append([]string{server.Alias}, legacyAliases...) + for _, secretType := range serverSecretTypes { + if server.ID > 0 { + v.Delete(stableServerSecretID(server.ID, secretType)) + } + for _, alias := range aliases { + if alias != "" { + v.Delete(serverSecretID(alias, secretType)) + } + } + } +} + +// syncServerSecrets writes credentials only under stable identity after the DB +// save has succeeded. Existing alias keys are migrated without depending on a +// rename operation, so a failed DB rename cannot orphan credentials. + +// cleanupServerSecrets keeps the legacy helper surface for CLI/tests and also +// removes stable-ID records when the server still exists. func cleanupServerSecrets(v *vault.Vault, alias string) { + if appDB != nil { + server, _ := appDB.GetServer(alias) + if server != nil { + cleanupServerSecretsForServer(v, server) + return + } + } for _, secretType := range serverSecretTypes { v.Delete(serverSecretID(alias, secretType)) } } func syncServerSecrets(v *vault.Vault, oldAlias string, server *model.Server, secret string) error { - if oldAlias == "" { - oldAlias = server.Alias + if server == nil { + return fmt.Errorf("server is required") } - if oldAlias != server.Alias { - for _, secretType := range serverSecretTypes { - oldID := serverSecretID(oldAlias, secretType) - data, err := v.Get(oldID) - if err == nil { - if err := v.Put(serverSecretID(server.Alias, secretType), secretType, data); err != nil { - return err + if server.ID <= 0 { + // Compatibility path for pre-persistence callers/tests. Real saves assign + // Server.ID before this function is called. Keep old alias-based vaults + // working and complete alias renames atomically in memory. + if oldAlias != "" && oldAlias != server.Alias { + for _, secretType := range serverSecretTypes { + oldID := serverSecretID(oldAlias, secretType) + if data, err := v.Get(oldID); err == nil { + if err := v.Put(serverSecretID(server.Alias, secretType), secretType, data); err != nil { + return err + } + v.Delete(oldID) } } - v.Delete(oldID) + } + key := func(secretType string) string { return serverSecretID(server.Alias, secretType) } + switch server.AuthMethod { + case model.AuthPassword: + v.Delete(key(secretKeyPassphrase)) + if secret != "" { + return v.Put(key(secretSSHPassword), secretSSHPassword, []byte(secret)) + } + case model.AuthKeyPassphrase: + v.Delete(key(secretSSHPassword)) + if secret != "" { + return v.Put(key(secretKeyPassphrase), secretKeyPassphrase, []byte(secret)) + } + default: + v.Delete(key(secretSSHPassword)) + v.Delete(key(secretKeyPassphrase)) + } + return nil + } + + aliases := []string{server.Alias} + if oldAlias != "" && oldAlias != server.Alias { + aliases = append(aliases, oldAlias) + } + for _, secretType := range serverSecretTypes { + stableID := stableServerSecretID(server.ID, secretType) + if !v.HasSecret(stableID) { + for _, alias := range aliases { + legacyID := serverSecretID(alias, secretType) + if data, err := v.Get(legacyID); err == nil { + if err := v.Put(stableID, secretType, data); err != nil { + return err + } + break + } + } + } + for _, alias := range aliases { + v.Delete(serverSecretID(alias, secretType)) } } switch server.AuthMethod { case model.AuthPassword: - v.Delete(serverSecretID(server.Alias, secretKeyPassphrase)) + v.Delete(stableServerSecretID(server.ID, secretKeyPassphrase)) if secret != "" { - return v.Put(serverSecretID(server.Alias, secretSSHPassword), secretSSHPassword, []byte(secret)) + return v.Put(stableServerSecretID(server.ID, secretSSHPassword), secretSSHPassword, []byte(secret)) } case model.AuthKeyPassphrase: - v.Delete(serverSecretID(server.Alias, secretSSHPassword)) + v.Delete(stableServerSecretID(server.ID, secretSSHPassword)) if secret != "" { - return v.Put(serverSecretID(server.Alias, secretKeyPassphrase), secretKeyPassphrase, []byte(secret)) + return v.Put(stableServerSecretID(server.ID, secretKeyPassphrase), secretKeyPassphrase, []byte(secret)) } default: - v.Delete(serverSecretID(server.Alias, secretSSHPassword)) - v.Delete(serverSecretID(server.Alias, secretKeyPassphrase)) + v.Delete(stableServerSecretID(server.ID, secretSSHPassword)) + v.Delete(stableServerSecretID(server.ID, secretKeyPassphrase)) } - return nil } func deleteVaultSecrets(v *vault.Vault, alias string, secretType string) error { - if secretType != "" { - v.Delete(serverSecretID(alias, secretType)) + var server *model.Server + if appDB != nil { + server, _ = appDB.GetServer(alias) + } + if server == nil { + if secretType != "" { + v.Delete(serverSecretID(alias, secretType)) + } else { + for _, t := range serverSecretTypes { + v.Delete(serverSecretID(alias, t)) + } + } return nil } - cleanupServerSecrets(v, alias) + if secretType != "" { + v.Delete(stableServerSecretID(server.ID, secretType)) + v.Delete(serverSecretID(server.Alias, secretType)) + return nil + } + cleanupServerSecretsForServer(v, server) return nil } @@ -83,3 +212,16 @@ func formTestVaultFunc(getVault ssh.VaultFunc, server *model.Server, formSecret return getVault(serverAlias, secretType) } } + +func vaultFuncForServer(v *vault.Vault, server *model.Server) ssh.VaultFunc { + return func(_ string, secretType string) (string, error) { + if !v.IsUnlocked() { + return "", fmt.Errorf("%s", vaultLockedProcessMessage()) + } + data, err := getServerSecret(v, server, secretType) + if err != nil { + return "", err + } + return string(data), nil + } +} diff --git a/internal/db/db.go b/internal/db/db.go index be85e0f..55a28e8 100644 --- a/internal/db/db.go +++ b/internal/db/db.go @@ -28,6 +28,9 @@ func Open(dataDir string) (*DB, error) { if err := conn.Ping(); err != nil { return nil, fmt.Errorf("ping database: %w", err) } + if _, err := conn.Exec("PRAGMA foreign_keys = ON"); err != nil { + return nil, fmt.Errorf("enable foreign keys: %w", err) + } db := &DB{conn: conn} @@ -69,6 +72,10 @@ func (db *DB) ensureSchema() error { } } + if err := db.ensureV040Schema(); err != nil { + return err + } + // Add forwards name/description/enabled columns for _, col := range []struct { name string diff --git a/internal/db/servers.go b/internal/db/servers.go index 955506d..e045635 100644 --- a/internal/db/servers.go +++ b/internal/db/servers.go @@ -3,6 +3,7 @@ package db import ( "database/sql" "encoding/json" + "fmt" "sort" "strings" "time" @@ -41,60 +42,25 @@ func unmarshalRoute(s string) model.Route { // --- Server CRUD --- -func (db *DB) CreateServer(s *model.Server) error { - result, err := db.conn.Exec(` - INSERT INTO servers (alias, display_name, host, port, user, auth_method, identity_file, proxy_jump, route_hops, group_name, notes, startup_command) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, - s.Alias, s.DisplayName, s.Host, s.Port, s.User, s.AuthMethod, s.IdentityFile, s.ProxyJump, marshalRoute(s.Route), s.GroupName, s.Notes, s.StartupCommand) - if err != nil { - return err - } - s.ID, _ = result.LastInsertId() - return nil +type rowScanner interface { + Scan(dest ...any) error } -func (db *DB) UpdateServer(s *model.Server) error { - _, err := db.conn.Exec(` - UPDATE servers SET - display_name=?, host=?, port=?, user=?, auth_method=?, - identity_file=?, proxy_jump=?, route_hops=?, group_name=?, notes=?, startup_command=?, updated_at=CURRENT_TIMESTAMP - WHERE alias=?`, - s.DisplayName, s.Host, s.Port, s.User, s.AuthMethod, - s.IdentityFile, s.ProxyJump, marshalRoute(s.Route), s.GroupName, s.Notes, s.StartupCommand, s.Alias) - return err -} +const serverSelectColumns = ` + id, alias, display_name, host, port, user, auth_method, + identity_file, proxy_jump, route_hops, COALESCE(group_id, 0), group_name, + notes, startup_command, created_at, updated_at, last_connected_at, + last_test_at, last_test_status, last_test_error` -func (db *DB) UpdateServerByAlias(oldAlias string, s *model.Server) error { - _, err := db.conn.Exec(` - UPDATE servers SET - alias=?, display_name=?, host=?, port=?, user=?, auth_method=?, - identity_file=?, proxy_jump=?, route_hops=?, group_name=?, notes=?, startup_command=?, updated_at=CURRENT_TIMESTAMP - WHERE alias=?`, - s.Alias, s.DisplayName, s.Host, s.Port, s.User, s.AuthMethod, - s.IdentityFile, s.ProxyJump, marshalRoute(s.Route), s.GroupName, s.Notes, s.StartupCommand, oldAlias) - return err -} - -func (db *DB) DeleteServer(alias string) error { - _, err := db.conn.Exec("DELETE FROM servers WHERE alias=?", alias) - return err -} - -func (db *DB) GetServer(alias string) (*model.Server, error) { +func scanServerBase(row rowScanner) (*model.Server, error) { var s model.Server var lastConnected, lastTest sql.NullTime - var routeHops sql.NullString - err := db.conn.QueryRow(` - SELECT id, alias, display_name, host, port, user, auth_method, - identity_file, proxy_jump, route_hops, group_name, notes, startup_command, - created_at, updated_at, last_connected_at, - last_test_at, last_test_status, last_test_error - FROM servers WHERE alias=?`, alias).Scan( + var legacyRoute sql.NullString + if err := row.Scan( &s.ID, &s.Alias, &s.DisplayName, &s.Host, &s.Port, &s.User, &s.AuthMethod, - &s.IdentityFile, &s.ProxyJump, &routeHops, &s.GroupName, &s.Notes, &s.StartupCommand, - &s.CreatedAt, &s.UpdatedAt, &lastConnected, - &lastTest, &s.LastTestStatus, &s.LastTestError) - if err != nil { + &s.IdentityFile, &s.ProxyJump, &legacyRoute, &s.GroupID, &s.GroupName, + &s.Notes, &s.StartupCommand, &s.CreatedAt, &s.UpdatedAt, &lastConnected, + &lastTest, &s.LastTestStatus, &s.LastTestError); err != nil { return nil, err } if lastConnected.Valid { @@ -103,134 +69,402 @@ func (db *DB) GetServer(alias string) (*model.Server, error) { if lastTest.Valid { s.LastTestAt = &lastTest.Time } - if routeHops.Valid && routeHops.String != "" { - s.Route = unmarshalRoute(routeHops.String) - } - if len(s.Route.Hops) == 0 && s.ProxyJump != "" { - s.Route = unmarshalRoute(s.ProxyJump) - } - tags, err := db.GetServerTags(s.ID) - if err != nil { - return nil, err - } - s.Tags = tags return &s, nil } -func (db *DB) ListServers() ([]*model.Server, error) { +func (db *DB) ResolveAlias(alias string) (int64, bool) { + var id int64 + if err := db.conn.QueryRow(`SELECT id FROM servers WHERE alias=?`, strings.TrimSpace(alias)).Scan(&id); err != nil { + return 0, false + } + return id, true +} + +func (db *DB) normalizeRoute(route model.Route) (model.Route, error) { + resolved := model.Route{Hops: make([]model.RouteHop, 0, len(route.Hops))} + for _, hop := range route.Hops { + if hop.Profile() { + var id int64 + alias := strings.TrimSpace(hop.Alias) + if hop.ServerID > 0 { + id = hop.ServerID + if err := db.conn.QueryRow(`SELECT alias FROM servers WHERE id=?`, id).Scan(&alias); err != nil { + if err == sql.ErrNoRows { + return model.Route{}, fmt.Errorf("route profile #%d not found", id) + } + return model.Route{}, err + } + } else { + if alias == "" { + return model.Route{}, fmt.Errorf("route profile alias is empty") + } + if err := db.conn.QueryRow(`SELECT id FROM servers WHERE alias=?`, alias).Scan(&id); err != nil { + if err == sql.ErrNoRows { + return model.Route{}, fmt.Errorf("route profile not found: %s", alias) + } + return model.Route{}, err + } + } + resolved.Hops = append(resolved.Hops, model.RouteHop{ServerID: id, Alias: alias, IsProfile: true}) + continue + } + raw := strings.TrimSpace(hop.Raw) + if raw == "" { + return model.Route{}, fmt.Errorf("raw route hop is empty") + } + resolved.Hops = append(resolved.Hops, model.RouteHop{Raw: raw}) + } + return resolved, nil +} + +func (db *DB) ValidateRoute(targetID int64, route model.Route) error { + resolved, err := db.normalizeRoute(route) + if err != nil { + return err + } + if err := model.ValidateRouteShape(targetID, resolved); err != nil { + return err + } + for _, hop := range resolved.Hops { + if !hop.Profile() { + continue + } + reaches, err := db.routeReaches(hop.ServerID, targetID, map[int64]bool{}) + if err != nil { + return err + } + if reaches { + return fmt.Errorf("route cycle detected through %s", hop.Alias) + } + } + return nil +} + +func (db *DB) routeReaches(startID, targetID int64, visiting map[int64]bool) (bool, error) { + if targetID > 0 && startID == targetID { + return true, nil + } + if visiting[startID] { + return false, fmt.Errorf("existing route cycle detected at server #%d", startID) + } + visiting[startID] = true + defer delete(visiting, startID) + route, err := db.loadRoute(startID) + if err != nil { + return false, err + } + for _, hop := range route.Hops { + if !hop.Profile() { + continue + } + reaches, err := db.routeReaches(hop.ServerID, targetID, visiting) + if err != nil || reaches { + return reaches, err + } + } + return false, nil +} + +func insertRouteTx(tx *sql.Tx, targetID int64, route model.Route) error { + if _, err := tx.Exec(`DELETE FROM server_route_hops WHERE target_server_id=?`, targetID); err != nil { + return err + } + for pos, hop := range route.Hops { + if hop.Profile() { + if _, err := tx.Exec(`INSERT INTO server_route_hops(target_server_id, position, hop_server_id) VALUES(?,?,?)`, targetID, pos, hop.ServerID); err != nil { + return err + } + } else { + if _, err := tx.Exec(`INSERT INTO server_route_hops(target_server_id, position, raw_target) VALUES(?,?,?)`, targetID, pos, hop.Raw); err != nil { + return err + } + } + } + return nil +} + +func (db *DB) loadRoute(targetID int64) (model.Route, error) { rows, err := db.conn.Query(` - SELECT id, alias, display_name, host, port, user, auth_method, - identity_file, proxy_jump, route_hops, group_name, notes, startup_command, - created_at, updated_at, last_connected_at, - last_test_at, last_test_status, last_test_error - FROM servers ORDER BY alias`) + SELECT h.hop_server_id, h.raw_target, COALESCE(s.alias, '') + FROM server_route_hops h + LEFT JOIN servers s ON s.id=h.hop_server_id + WHERE h.target_server_id=? ORDER BY h.position`, targetID) + if err != nil { + return model.Route{}, err + } + defer rows.Close() + route := model.Route{} + for rows.Next() { + var hopID sql.NullInt64 + var raw sql.NullString + var alias string + if err := rows.Scan(&hopID, &raw, &alias); err != nil { + return model.Route{}, err + } + if hopID.Valid { + route.Hops = append(route.Hops, model.RouteHop{ServerID: hopID.Int64, Alias: alias, IsProfile: true}) + } else if raw.Valid { + route.Hops = append(route.Hops, model.RouteHop{Raw: raw.String}) + } + } + return route, rows.Err() +} + +func (db *DB) routeDependents(serverID int64) ([]string, error) { + rows, err := db.conn.Query(` + SELECT DISTINCT s.alias FROM server_route_hops h + JOIN servers s ON s.id=h.target_server_id + WHERE h.hop_server_id=? ORDER BY s.alias`, serverID) if err != nil { return nil, err } defer rows.Close() + var aliases []string + for rows.Next() { + var alias string + if err := rows.Scan(&alias); err != nil { + return nil, err + } + aliases = append(aliases, alias) + } + return aliases, rows.Err() +} +func (db *DB) refreshRouteCompatibility() error { + rows, err := db.conn.Query(`SELECT id FROM servers ORDER BY id`) + if err != nil { + return err + } + var ids []int64 + for rows.Next() { + var id int64 + if err := rows.Scan(&id); err != nil { + rows.Close() + return err + } + ids = append(ids, id) + } + rows.Close() + for _, id := range ids { + route, err := db.loadRoute(id) + if err != nil { + return err + } + proxy, legacy := routeCompatibilityProjection(route) + if _, err := db.conn.Exec(`UPDATE servers SET proxy_jump=?, route_hops=? WHERE id=?`, proxy, legacy, id); err != nil { + return err + } + } + return nil +} + +func (db *DB) CreateServer(s *model.Server) error { + resolved, err := db.normalizeRoute(s.Route) + if err != nil { + return err + } + s.Route = resolved + s.ProxyJump = s.Route.ProxyJumpString() + if err := model.ValidateServerBasics(s); err != nil { + return err + } + + tx, err := db.conn.Begin() + if err != nil { + return err + } + defer tx.Rollback() + groupID, err := ensureGroupTx(tx, s.GroupName) + if err != nil { + return err + } + proxy, legacy := routeCompatibilityProjection(s.Route) + result, err := tx.Exec(` + INSERT INTO servers (alias, display_name, host, port, user, auth_method, identity_file, + proxy_jump, route_hops, group_id, group_name, notes, startup_command) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + s.Alias, s.DisplayName, s.Host, s.Port, s.User, s.AuthMethod, s.IdentityFile, + proxy, legacy, nullGroupID(groupID), strings.TrimSpace(s.GroupName), s.Notes, s.StartupCommand) + if err != nil { + return err + } + s.ID, err = result.LastInsertId() + if err != nil { + return err + } + s.GroupID = groupID + if err := model.ValidateRouteShape(s.ID, s.Route); err != nil { + return err + } + if err := insertRouteTx(tx, s.ID, s.Route); err != nil { + return err + } + return tx.Commit() +} + +func nullGroupID(id int64) any { + if id == 0 { + return nil + } + return id +} + +func (db *DB) UpdateServer(s *model.Server) error { + return db.UpdateServerByAlias(s.Alias, s) +} + +func (db *DB) UpdateServerByAlias(oldAlias string, s *model.Server) error { + var id int64 + if err := db.conn.QueryRow(`SELECT id FROM servers WHERE alias=?`, oldAlias).Scan(&id); err != nil { + return err + } + resolved, err := db.normalizeRoute(s.Route) + if err != nil { + return err + } + s.ID = id + s.Route = resolved + s.ProxyJump = s.Route.ProxyJumpString() + if err := model.ValidateServerBasics(s); err != nil { + return err + } + if err := db.ValidateRoute(id, s.Route); err != nil { + return err + } + + tx, err := db.conn.Begin() + if err != nil { + return err + } + defer tx.Rollback() + groupID, err := ensureGroupTx(tx, s.GroupName) + if err != nil { + return err + } + proxy, legacy := routeCompatibilityProjection(s.Route) + result, err := tx.Exec(` + UPDATE servers SET alias=?, display_name=?, host=?, port=?, user=?, auth_method=?, + identity_file=?, proxy_jump=?, route_hops=?, group_id=?, group_name=?, notes=?, startup_command=?, updated_at=CURRENT_TIMESTAMP + WHERE id=?`, + s.Alias, s.DisplayName, s.Host, s.Port, s.User, s.AuthMethod, + s.IdentityFile, proxy, legacy, nullGroupID(groupID), strings.TrimSpace(s.GroupName), s.Notes, s.StartupCommand, id) + if err != nil { + return err + } + if affected, _ := result.RowsAffected(); affected != 1 { + return fmt.Errorf("server not found: %s", oldAlias) + } + if err := insertRouteTx(tx, id, s.Route); err != nil { + return err + } + if err := tx.Commit(); err != nil { + return err + } + s.GroupID = groupID + return db.refreshRouteCompatibility() +} + +func (db *DB) DeleteServer(alias string) error { + var id int64 + if err := db.conn.QueryRow(`SELECT id FROM servers WHERE alias=?`, alias).Scan(&id); err != nil { + return err + } + dependents, err := db.routeDependents(id) + if err != nil { + return err + } + if len(dependents) > 0 { + return fmt.Errorf("server %q is used as a route hop by: %s", alias, strings.Join(dependents, ", ")) + } + _, err = db.conn.Exec(`DELETE FROM servers WHERE id=?`, id) + return err +} + +func (db *DB) loadServerByQuery(query string, arg any) (*model.Server, error) { + s, err := scanServerBase(db.conn.QueryRow(`SELECT `+serverSelectColumns+` FROM servers WHERE `+query, arg)) + if err != nil { + return nil, err + } + s.Route, err = db.loadRoute(s.ID) + if err != nil { + return nil, err + } + s.ProxyJump = s.Route.ProxyJumpString() + s.Tags, err = db.GetServerTags(s.ID) + if err != nil { + return nil, err + } + return s, nil +} + +func (db *DB) GetServer(alias string) (*model.Server, error) { + return db.loadServerByQuery(`alias=?`, alias) +} + +func (db *DB) GetServerByID(id int64) (*model.Server, error) { + return db.loadServerByQuery(`id=?`, id) +} + +func (db *DB) listServersQuery(query string, args ...any) ([]*model.Server, error) { + rows, err := db.conn.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() var servers []*model.Server for rows.Next() { - var s model.Server - var lastConnected, lastTest sql.NullTime - var routeHops sql.NullString - err := rows.Scan( - &s.ID, &s.Alias, &s.DisplayName, &s.Host, &s.Port, &s.User, &s.AuthMethod, - &s.IdentityFile, &s.ProxyJump, &routeHops, &s.GroupName, &s.Notes, &s.StartupCommand, - &s.CreatedAt, &s.UpdatedAt, &lastConnected, - &lastTest, &s.LastTestStatus, &s.LastTestError) + s, err := scanServerBase(rows) if err != nil { return nil, err } - if lastConnected.Valid { - s.LastConnectedAt = &lastConnected.Time - } - if lastTest.Valid { - s.LastTestAt = &lastTest.Time - } - if routeHops.Valid && routeHops.String != "" { - s.Route = unmarshalRoute(routeHops.String) - } - if len(s.Route.Hops) == 0 && s.ProxyJump != "" { - s.Route = unmarshalRoute(s.ProxyJump) - } - tags, err := db.GetServerTags(s.ID) - if err != nil { - return nil, err - } - s.Tags = tags - servers = append(servers, &s) + servers = append(servers, s) } - return servers, rows.Err() + if err := rows.Err(); err != nil { + return nil, err + } + for _, s := range servers { + s.Route, err = db.loadRoute(s.ID) + if err != nil { + return nil, err + } + s.ProxyJump = s.Route.ProxyJumpString() + s.Tags, err = db.GetServerTags(s.ID) + if err != nil { + return nil, err + } + } + return servers, nil +} + +func (db *DB) ListServers() ([]*model.Server, error) { + return db.listServersQuery(`SELECT ` + serverSelectColumns + ` FROM servers ORDER BY alias`) } func (db *DB) SearchServers(query string) ([]*model.Server, error) { pattern := "%" + query + "%" - rows, err := db.conn.Query(` - SELECT id, alias, display_name, host, port, user, auth_method, - identity_file, proxy_jump, route_hops, group_name, notes, startup_command, - created_at, updated_at, last_connected_at, - last_test_at, last_test_status, last_test_error - FROM servers + return db.listServersQuery(` + SELECT `+serverSelectColumns+` FROM servers WHERE alias LIKE ? OR display_name LIKE ? OR host LIKE ? OR user LIKE ? - OR group_name LIKE ? OR notes LIKE ? OR proxy_jump LIKE ? OR route_hops LIKE ? - OR EXISTS ( - SELECT 1 FROM server_tags st - JOIN tags t ON t.id = st.tag_id - WHERE st.server_id = servers.id AND t.name LIKE ? - ) - OR EXISTS ( - SELECT 1 FROM forwards f - WHERE f.server_id = servers.id - AND ( - f.name LIKE ? OR f.description LIKE ? - OR f.local_addr LIKE ? OR f.remote_addr LIKE ? - OR CAST(f.local_port AS TEXT) LIKE ? - OR CAST(f.remote_port AS TEXT) LIKE ? - ) - ) + OR group_name LIKE ? OR notes LIKE ? OR proxy_jump LIKE ? OR route_hops LIKE ? + OR EXISTS ( + SELECT 1 FROM server_route_hops rh + LEFT JOIN servers hs ON hs.id=rh.hop_server_id + WHERE rh.target_server_id=servers.id + AND (hs.alias LIKE ? OR rh.raw_target LIKE ?) + ) + OR EXISTS ( + SELECT 1 FROM server_tags st JOIN tags t ON t.id=st.tag_id + WHERE st.server_id=servers.id AND t.name LIKE ? + ) + OR EXISTS ( + SELECT 1 FROM forwards f WHERE f.server_id=servers.id + AND (f.name LIKE ? OR f.description LIKE ? OR f.local_addr LIKE ? OR f.remote_addr LIKE ? + OR CAST(f.local_port AS TEXT) LIKE ? OR CAST(f.remote_port AS TEXT) LIKE ?) + ) ORDER BY alias`, pattern, pattern, pattern, pattern, pattern, pattern, pattern, pattern, - pattern, + pattern, pattern, pattern, pattern, pattern, pattern, pattern, pattern, pattern) - if err != nil { - return nil, err - } - defer rows.Close() - - var servers []*model.Server - for rows.Next() { - var s model.Server - var lastConnected, lastTest sql.NullTime - var routeHops sql.NullString - err := rows.Scan( - &s.ID, &s.Alias, &s.DisplayName, &s.Host, &s.Port, &s.User, &s.AuthMethod, - &s.IdentityFile, &s.ProxyJump, &routeHops, &s.GroupName, &s.Notes, &s.StartupCommand, - &s.CreatedAt, &s.UpdatedAt, &lastConnected, - &lastTest, &s.LastTestStatus, &s.LastTestError) - if err != nil { - return nil, err - } - if lastConnected.Valid { - s.LastConnectedAt = &lastConnected.Time - } - if lastTest.Valid { - s.LastTestAt = &lastTest.Time - } - if routeHops.Valid && routeHops.String != "" { - s.Route = unmarshalRoute(routeHops.String) - } - if len(s.Route.Hops) == 0 && s.ProxyJump != "" { - s.Route = unmarshalRoute(s.ProxyJump) - } - tags, err := db.GetServerTags(s.ID) - if err != nil { - return nil, err - } - s.Tags = tags - servers = append(servers, &s) - } - return servers, rows.Err() } func (db *DB) UpdateTestResult(alias string, status model.TestStatus, testErr string) error { @@ -474,38 +708,93 @@ func uniqueCleanStrings(values []string) []string { // --- Group methods --- -func (db *DB) GetGroups() ([]string, error) { +func (db *DB) CreateGroup(name string) error { + name = strings.TrimSpace(name) + if name == "" { + return fmt.Errorf("group name is required") + } + _, err := db.conn.Exec(`INSERT INTO groups(name) VALUES(?)`, name) + return err +} + +func (db *DB) ListGroups() ([]*model.Group, error) { rows, err := db.conn.Query(` - SELECT group_name FROM servers - WHERE group_name != '' - GROUP BY group_name - ORDER BY group_name`) + SELECT g.id, g.name, count(s.id) + FROM groups g LEFT JOIN servers s ON s.group_id=g.id + GROUP BY g.id, g.name ORDER BY g.name`) if err != nil { return nil, err } defer rows.Close() - - var groups []string + var groups []*model.Group for rows.Next() { - var name string - if err := rows.Scan(&name); err != nil { + var group model.Group + if err := rows.Scan(&group.ID, &group.Name, &group.ServerCount); err != nil { return nil, err } - groups = append(groups, name) + groups = append(groups, &group) } return groups, rows.Err() } +func (db *DB) GetGroups() ([]string, error) { + groups, err := db.ListGroups() + if err != nil { + return nil, err + } + names := make([]string, len(groups)) + for i, group := range groups { + names[i] = group.Name + } + return names, nil +} + func (db *DB) RenameGroup(oldName, newName string) error { - _, err := db.conn.Exec( - "UPDATE servers SET group_name = ?, updated_at = CURRENT_TIMESTAMP WHERE group_name = ?", - newName, oldName) - return err + oldName = strings.TrimSpace(oldName) + newName = strings.TrimSpace(newName) + if oldName == "" || newName == "" { + return fmt.Errorf("group name is required") + } + tx, err := db.conn.Begin() + if err != nil { + return err + } + defer tx.Rollback() + var id int64 + if err := tx.QueryRow(`SELECT id FROM groups WHERE name=?`, oldName).Scan(&id); err != nil { + return err + } + if _, err := tx.Exec(`UPDATE groups SET name=? WHERE id=?`, newName, id); err != nil { + return err + } + if _, err := tx.Exec(`UPDATE servers SET group_name=?, updated_at=CURRENT_TIMESTAMP WHERE group_id=?`, newName, id); err != nil { + return err + } + return tx.Commit() } func (db *DB) DeleteGroup(name string) error { - _, err := db.conn.Exec( - "UPDATE servers SET group_name = '', updated_at = CURRENT_TIMESTAMP WHERE group_name = ?", - name) - return err + name = strings.TrimSpace(name) + if name == "" { + return nil + } + tx, err := db.conn.Begin() + if err != nil { + return err + } + defer tx.Rollback() + var id int64 + if err := tx.QueryRow(`SELECT id FROM groups WHERE name=?`, name).Scan(&id); err != nil { + if err == sql.ErrNoRows { + return nil + } + return err + } + if _, err := tx.Exec(`UPDATE servers SET group_id=NULL, group_name='', updated_at=CURRENT_TIMESTAMP WHERE group_id=?`, id); err != nil { + return err + } + if _, err := tx.Exec(`DELETE FROM groups WHERE id=?`, id); err != nil { + return err + } + return tx.Commit() } diff --git a/internal/db/servers_test.go b/internal/db/servers_test.go index 91d74ba..cf405cb 100644 --- a/internal/db/servers_test.go +++ b/internal/db/servers_test.go @@ -240,6 +240,10 @@ func TestSearchServersMatchesTagsRoutesAndForwardPorts(t *testing.T) { } defer db.Close() + bastion := &model.Server{Alias: "bastion", Host: "bastion.internal", Port: 22, User: "root", AuthMethod: model.AuthKey} + if err := db.CreateServer(bastion); err != nil { + t.Fatalf("create bastion: %v", err) + } server := &model.Server{ Alias: "db", Host: "db.internal", @@ -247,8 +251,8 @@ func TestSearchServersMatchesTagsRoutesAndForwardPorts(t *testing.T) { User: "postgres", AuthMethod: model.AuthKey, Route: model.Route{Hops: []model.RouteHop{ - {Alias: "bastion", IsProfile: true}, - {Raw: "dmz.example.org", IsProfile: false}, + {ServerID: bastion.ID, Alias: "bastion", IsProfile: true}, + {Raw: "dmz.example.org"}, }}, } if err := db.CreateServer(server); err != nil { diff --git a/internal/db/v040.go b/internal/db/v040.go new file mode 100644 index 0000000..19e29ea --- /dev/null +++ b/internal/db/v040.go @@ -0,0 +1,159 @@ +package db + +import ( + "database/sql" + "fmt" + "strings" + + "github.com/mirivlad/sshkeeper/internal/model" +) + +func (db *DB) ensureV040Schema() error { + if _, err := db.conn.Exec(` + CREATE TABLE IF NOT EXISTS groups ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL UNIQUE + )`); err != nil { + return fmt.Errorf("create groups: %w", err) + } + + hasGroupID, err := db.hasColumn("servers", "group_id") + if err != nil { + return err + } + if !hasGroupID { + if _, err := db.conn.Exec("ALTER TABLE servers ADD COLUMN group_id INTEGER"); err != nil { + return fmt.Errorf("add servers.group_id: %w", err) + } + } + + if _, err := db.conn.Exec(` + CREATE TABLE IF NOT EXISTS server_route_hops ( + target_server_id INTEGER NOT NULL REFERENCES servers(id) ON DELETE CASCADE, + position INTEGER NOT NULL, + hop_server_id INTEGER REFERENCES servers(id) ON DELETE RESTRICT, + raw_target TEXT, + PRIMARY KEY (target_server_id, position), + CHECK ((hop_server_id IS NOT NULL AND raw_target IS NULL) OR + (hop_server_id IS NULL AND raw_target IS NOT NULL AND length(trim(raw_target)) > 0)) + )`); err != nil { + return fmt.Errorf("create server_route_hops: %w", err) + } + if _, err := db.conn.Exec(`CREATE INDEX IF NOT EXISTS idx_route_hop_profile ON server_route_hops(hop_server_id)`); err != nil { + return fmt.Errorf("index server_route_hops: %w", err) + } + + if err := db.migrateLegacyGroups(); err != nil { + return err + } + if err := db.migrateLegacyRoutes(); err != nil { + return err + } + return nil +} + +func (db *DB) migrateLegacyGroups() error { + if _, err := db.conn.Exec(` + INSERT OR IGNORE INTO groups(name) + SELECT DISTINCT trim(group_name) FROM servers WHERE trim(group_name) != ''`); err != nil { + return fmt.Errorf("migrate groups: %w", err) + } + if _, err := db.conn.Exec(` + UPDATE servers + SET group_id = (SELECT id FROM groups WHERE groups.name = servers.group_name) + WHERE group_id IS NULL AND trim(group_name) != ''`); err != nil { + return fmt.Errorf("link migrated groups: %w", err) + } + return nil +} + +func (db *DB) migrateLegacyRoutes() error { + rows, err := db.conn.Query(`SELECT id, proxy_jump, route_hops FROM servers ORDER BY id`) + if err != nil { + return fmt.Errorf("read legacy routes: %w", err) + } + defer rows.Close() + + type legacyServer struct { + id int64 + proxyJump string + routeHops string + } + var legacy []legacyServer + for rows.Next() { + var item legacyServer + if err := rows.Scan(&item.id, &item.proxyJump, &item.routeHops); err != nil { + return err + } + legacy = append(legacy, item) + } + if err := rows.Err(); err != nil { + return err + } + + for _, item := range legacy { + var count int + if err := db.conn.QueryRow(`SELECT count(*) FROM server_route_hops WHERE target_server_id=?`, item.id).Scan(&count); err != nil { + return err + } + if count > 0 { + continue + } + source := strings.TrimSpace(item.routeHops) + if source == "" { + source = strings.TrimSpace(item.proxyJump) + } + if source == "" { + continue + } + route := unmarshalRoute(source) + for pos, hop := range route.Hops { + hopID := hop.ServerID + candidate := strings.TrimSpace(hop.Alias) + if candidate == "" { + candidate = strings.TrimSpace(hop.Raw) + } + if hopID == 0 && candidate != "" { + _ = db.conn.QueryRow(`SELECT id FROM servers WHERE alias=?`, candidate).Scan(&hopID) + } + if hopID > 0 && hopID != item.id { + if _, err := db.conn.Exec(`INSERT INTO server_route_hops(target_server_id, position, hop_server_id) VALUES(?,?,?)`, item.id, pos, hopID); err != nil { + return fmt.Errorf("migrate route hop: %w", err) + } + continue + } + raw := candidate + if raw == "" { + raw = strings.TrimSpace(hop.Raw) + } + if raw == "" { + continue + } + if _, err := db.conn.Exec(`INSERT INTO server_route_hops(target_server_id, position, raw_target) VALUES(?,?,?)`, item.id, pos, raw); err != nil { + return fmt.Errorf("migrate raw route hop: %w", err) + } + } + } + return nil +} + +func ensureGroupTx(tx *sql.Tx, name string) (int64, error) { + name = strings.TrimSpace(name) + if name == "" { + return 0, nil + } + if _, err := tx.Exec(`INSERT OR IGNORE INTO groups(name) VALUES(?)`, name); err != nil { + return 0, err + } + var id int64 + if err := tx.QueryRow(`SELECT id FROM groups WHERE name=?`, name).Scan(&id); err != nil { + return 0, err + } + return id, nil +} + +func routeCompatibilityProjection(route model.Route) (proxyJump, routeJSON string) { + proxyJump = route.ProxyJumpString() + routeJSON = marshalRoute(route) + return proxyJump, routeJSON +} diff --git a/internal/model/server.go b/internal/model/server.go index 95c7ace..c2cd368 100644 --- a/internal/model/server.go +++ b/internal/model/server.go @@ -24,16 +24,18 @@ const ( ) type Server struct { - ID int64 `json:"id"` - Alias string `json:"alias"` - DisplayName string `json:"display_name"` - Host string `json:"host"` - Port int `json:"port"` - User string `json:"user"` - AuthMethod AuthMethod `json:"auth_method"` - IdentityFile string `json:"identity_file"` + ID int64 `json:"id"` + Alias string `json:"alias"` + DisplayName string `json:"display_name"` + Host string `json:"host"` + Port int `json:"port"` + User string `json:"user"` + AuthMethod AuthMethod `json:"auth_method"` + IdentityFile string `json:"identity_file"` + // ProxyJump is a deprecated compatibility projection of Route. ProxyJump string `json:"proxy_jump"` Route Route `json:"route"` + GroupID int64 `json:"group_id"` GroupName string `json:"group_name"` Notes string `json:"notes"` StartupCommand string `json:"startup_command"` @@ -141,20 +143,41 @@ func (f *Forward) ForwardTarget() string { } type Tag struct { - ID int64 `json:"id"` - Name string `json:"name"` + ID int64 `json:"id"` + Name string `json:"name"` + ServerCount int `json:"server_count,omitempty"` +} + +type Group struct { + ID int64 `json:"id"` + Name string `json:"name"` + ServerCount int `json:"server_count,omitempty"` } // --- Route --- -// RouteHop represents a single jump host in a route. -// IsProfile: true = use Alias (references a sshkeeper profile), false = use Raw (literal address). +// RouteHop is either a stable reference to another sshkeeper profile or a raw +// OpenSSH jump target. Alias is a display/backward-compatibility cache only; +// ServerID is the identity for profile hops. type RouteHop struct { - Alias string `json:"alias"` - Raw string `json:"raw"` + ServerID int64 `json:"server_id,omitempty"` + Alias string `json:"alias,omitempty"` + Raw string `json:"raw,omitempty"` IsProfile bool `json:"is_profile"` } +func (h RouteHop) Profile() bool { return h.IsProfile || h.ServerID != 0 } + +func (h RouteHop) DisplayName() string { + if h.Profile() { + if h.Alias != "" { + return h.Alias + } + return fmt.Sprintf("profile#%d", h.ServerID) + } + return h.Raw +} + // Route represents the SSH jump route for a server. // Mode is computed from Hops length: 0=direct, 1=via, 2+=chain type Route struct { @@ -177,11 +200,7 @@ func (r Route) RouteMode() string { func (r Route) ProxyJumpString() string { parts := make([]string, len(r.Hops)) for i, h := range r.Hops { - if h.IsProfile { - parts[i] = h.Alias - } else { - parts[i] = h.Raw - } + parts[i] = h.DisplayName() } return strings.Join(parts, ",") } @@ -194,11 +213,7 @@ func (r Route) DisplaySummary(target string) string { } names := make([]string, len(r.Hops)) for i, h := range r.Hops { - if h.IsProfile { - names[i] = h.Alias - } else { - names[i] = h.Raw - } + names[i] = h.DisplayName() } return strings.Join(names, " → ") + " → " + target } @@ -206,7 +221,7 @@ func (r Route) DisplaySummary(target string) string { // HasProfileLinks returns true if any hop references a known profile. func (r Route) HasProfileLinks() bool { for _, h := range r.Hops { - if h.IsProfile { + if h.Profile() { return true } } @@ -215,7 +230,6 @@ func (r Route) HasProfileLinks() bool { type CommandTemplate struct { ID int64 `json:"id"` - ServerID int64 `json:"server_id"` Name string `json:"name"` Command string `json:"command"` Description string `json:"description"` diff --git a/internal/model/validation.go b/internal/model/validation.go new file mode 100644 index 0000000..e303cc4 --- /dev/null +++ b/internal/model/validation.go @@ -0,0 +1,147 @@ +package model + +import ( + "fmt" + "strconv" + "strings" +) + +// IsSupportedAuthMethod reports whether method is a supported sshkeeper auth mode. +func IsSupportedAuthMethod(method AuthMethod) bool { + switch method { + case AuthPassword, AuthKey, AuthKeyPassphrase, AuthAgent: + return true + default: + return false + } +} + +// ValidateServerBasics validates fields that do not require database access. +func ValidateServerBasics(s *Server) error { + if s == nil { + return fmt.Errorf("server is required") + } + if strings.TrimSpace(s.Alias) == "" { + return fmt.Errorf("alias is required") + } + if strings.TrimSpace(s.Host) == "" { + return fmt.Errorf("host is required") + } + if s.Port < 1 || s.Port > 65535 { + return fmt.Errorf("port must be between 1 and 65535") + } + if s.AuthMethod == "" { + s.AuthMethod = AuthKey + } + if !IsSupportedAuthMethod(s.AuthMethod) { + return fmt.Errorf("unsupported auth method: %s", s.AuthMethod) + } + if (s.AuthMethod == AuthKey || s.AuthMethod == AuthKeyPassphrase) && strings.TrimSpace(s.IdentityFile) == "" { + // OpenSSH may still find a default key, so this is intentionally allowed. + } + return ValidateRouteShape(s.ID, s.Route) +} + +// ValidateRouteShape validates a route without resolving external references. +func ValidateRouteShape(targetID int64, route Route) error { + seenProfiles := map[int64]bool{} + seenRaw := map[string]bool{} + for i, hop := range route.Hops { + if hop.Profile() { + if hop.ServerID <= 0 && strings.TrimSpace(hop.Alias) == "" { + return fmt.Errorf("route hop %d has no profile reference", i+1) + } + if targetID > 0 && hop.ServerID == targetID { + return fmt.Errorf("route cannot use the target server itself as a hop") + } + if hop.ServerID > 0 { + if seenProfiles[hop.ServerID] { + return fmt.Errorf("route contains duplicate profile hop %s", hop.DisplayName()) + } + seenProfiles[hop.ServerID] = true + } + continue + } + raw := strings.TrimSpace(hop.Raw) + if raw == "" { + return fmt.Errorf("route hop %d is empty", i+1) + } + if seenRaw[raw] { + return fmt.Errorf("route contains duplicate raw hop %q", raw) + } + seenRaw[raw] = true + } + return nil +} + +// AliasResolver resolves an sshkeeper alias to a stable server ID. +type AliasResolver func(alias string) (int64, bool) + +// ParseRouteSpec parses CLI/legacy route syntax into an explicit Route. +// profile: requires an existing sshkeeper profile; raw: is always +// a literal OpenSSH target. Unprefixed entries are backward-compatible: an +// exact known alias becomes a profile reference, otherwise the entry is raw. +func ParseRouteSpec(input string, resolve AliasResolver) (Route, error) { + input = strings.TrimSpace(input) + if input == "" { + return Route{}, nil + } + parts := strings.Split(input, ",") + route := Route{Hops: make([]RouteHop, 0, len(parts))} + for _, part := range parts { + part = strings.TrimSpace(part) + if part == "" { + continue + } + switch { + case strings.HasPrefix(part, "profile:"): + alias := strings.TrimSpace(strings.TrimPrefix(part, "profile:")) + if alias == "" { + return Route{}, fmt.Errorf("empty profile route hop") + } + if resolve == nil { + return Route{}, fmt.Errorf("cannot resolve profile route hop %q", alias) + } + id, ok := resolve(alias) + if !ok || id <= 0 { + return Route{}, fmt.Errorf("route profile not found: %s", alias) + } + route.Hops = append(route.Hops, RouteHop{ServerID: id, Alias: alias, IsProfile: true}) + case strings.HasPrefix(part, "raw:"): + raw := strings.TrimSpace(strings.TrimPrefix(part, "raw:")) + if raw == "" { + return Route{}, fmt.Errorf("empty raw route hop") + } + route.Hops = append(route.Hops, RouteHop{Raw: raw}) + default: + if resolve != nil { + if id, ok := resolve(part); ok && id > 0 { + route.Hops = append(route.Hops, RouteHop{ServerID: id, Alias: part, IsProfile: true}) + continue + } + } + route.Hops = append(route.Hops, RouteHop{Raw: part}) + } + } + if err := ValidateRouteShape(0, route); err != nil { + return Route{}, err + } + return route, nil +} + +// FormatRouteSpec returns an unambiguous CLI representation. +func FormatRouteSpec(route Route) string { + parts := make([]string, 0, len(route.Hops)) + for _, hop := range route.Hops { + if hop.Profile() { + name := hop.Alias + if name == "" { + name = strconv.FormatInt(hop.ServerID, 10) + } + parts = append(parts, "profile:"+name) + } else { + parts = append(parts, "raw:"+hop.Raw) + } + } + return strings.Join(parts, ",") +} diff --git a/internal/ssh/command.go b/internal/ssh/command.go index f7743f0..faf3a6d 100644 --- a/internal/ssh/command.go +++ b/internal/ssh/command.go @@ -36,90 +36,95 @@ func validateSSHBinaryForOS(goos string, binary string, lookPath func(string) (s return nil } -func Connect(cfg *config.Config, server *model.Server, getVault VaultFunc) error { +func runPrepared(cfg *config.Config, args []string, server *model.Server, getVault VaultFunc) error { + switch server.AuthMethod { + case model.AuthPassword: + password, err := getVault(server.Alias, "ssh_password") + if err != nil { + return fmt.Errorf("get password from vault: %w", err) + } + return ConnectWithPassword(cfg.SSH.Binary, args, password) + case model.AuthKeyPassphrase: + passphrase, err := getVault(server.Alias, "key_passphrase") + if err != nil { + return fmt.Errorf("get key passphrase from vault: %w", err) + } + return ConnectWithPassword(cfg.SSH.Binary, args, passphrase) + default: + cmd := exec.Command(cfg.SSH.Binary, args...) + cmd.Stdin = os.Stdin + cmd.Stdout = os.Stdout + cmd.Stderr = os.Stderr + if err := cmd.Start(); err != nil { + return fmt.Errorf("start ssh: %w", err) + } + return cmd.Wait() + } +} + +func insertBeforeTarget(args []string, values ...string) []string { + if len(args) == 0 { + return append([]string(nil), values...) + } + result := make([]string, 0, len(args)+len(values)) + result = append(result, args[:len(args)-1]...) + result = append(result, values...) + result = append(result, args[len(args)-1]) + return result +} + +func ConnectResolved(cfg *config.Config, server *model.Server, resolve ProfileResolver, getVault VaultFunc) error { if err := EnsureSSHBinary(cfg.SSH.Binary); err != nil { return err } - - args := BuildSSHArgsSimple(server) + invocation, err := PrepareSSHInvocation(server, nil, false, resolve) + if err != nil { + return err + } + defer invocation.Cleanup() + args := append([]string(nil), invocation.Args...) if strings.TrimSpace(server.StartupCommand) != "" { args = append(args, server.StartupCommand) } - - switch server.AuthMethod { - case model.AuthPassword: - password, err := getVault(server.Alias, "ssh_password") - if err != nil { - return fmt.Errorf("get password from vault: %w", err) - } - return ConnectWithPassword(cfg.SSH.Binary, args, password) - - case model.AuthKeyPassphrase: - passphrase, err := getVault(server.Alias, "key_passphrase") - if err != nil { - return fmt.Errorf("get key passphrase from vault: %w", err) - } - return ConnectWithPassword(cfg.SSH.Binary, args, passphrase) - - default: - // key and agent auth use direct OpenSSH execution. - cmd := exec.Command(cfg.SSH.Binary, args...) - cmd.Stdin = os.Stdin - cmd.Stdout = os.Stdout - cmd.Stderr = os.Stderr - - if err := cmd.Start(); err != nil { - return fmt.Errorf("start ssh: %w", err) - } - - return cmd.Wait() - } + return runPrepared(cfg, args, server, getVault) } -func RunCommand(cfg *config.Config, server *model.Server, getVault VaultFunc, command string) error { +func Connect(cfg *config.Config, server *model.Server, getVault VaultFunc) error { + return ConnectResolved(cfg, server, nil, getVault) +} + +func RunCommandResolved(cfg *config.Config, server *model.Server, resolve ProfileResolver, getVault VaultFunc, command string) error { if err := EnsureSSHBinary(cfg.SSH.Binary); err != nil { return err } - - args := BuildSSHArgsSimple(server) - args = append(args, command) - - switch server.AuthMethod { - case model.AuthPassword: - password, err := getVault(server.Alias, "ssh_password") - if err != nil { - return fmt.Errorf("get password from vault: %w", err) - } - return ConnectWithPassword(cfg.SSH.Binary, args, password) - case model.AuthKeyPassphrase: - passphrase, err := getVault(server.Alias, "key_passphrase") - if err != nil { - return fmt.Errorf("get key passphrase from vault: %w", err) - } - return ConnectWithPassword(cfg.SSH.Binary, args, passphrase) - default: - cmd := exec.Command(cfg.SSH.Binary, args...) - cmd.Stdin = os.Stdin - cmd.Stdout = os.Stdout - cmd.Stderr = os.Stderr - if err := cmd.Start(); err != nil { - return fmt.Errorf("start ssh: %w", err) - } - return cmd.Wait() + invocation, err := PrepareSSHInvocation(server, nil, false, resolve) + if err != nil { + return err } + defer invocation.Cleanup() + args := append(append([]string(nil), invocation.Args...), command) + return runPrepared(cfg, args, server, getVault) } -func RunCommandOutput(cfg *config.Config, server *model.Server, getVault VaultFunc, command string) (string, error) { +func RunCommand(cfg *config.Config, server *model.Server, getVault VaultFunc, command string) error { + return RunCommandResolved(cfg, server, nil, getVault, command) +} + +func RunCommandOutputResolved(cfg *config.Config, server *model.Server, resolve ProfileResolver, getVault VaultFunc, command string) (string, error) { if err := EnsureSSHBinary(cfg.SSH.Binary); err != nil { return "", err } - - args := BuildSSHArgsSimple(server) - args = append(args, "-o", fmt.Sprintf("ConnectTimeout=%d", cfg.SSH.ConnectTimeoutSec)) + invocation, err := PrepareSSHInvocation(server, nil, false, resolve) + if err != nil { + return "", err + } + defer invocation.Cleanup() + args := insertBeforeTarget(invocation.Args, "-o", fmt.Sprintf("ConnectTimeout=%d", cfg.SSH.ConnectTimeoutSec)) switch server.AuthMethod { case model.AuthPassword: - args = append(args, "-o", "NumberOfPasswordPrompts=1", command) + args = insertBeforeTarget(args, "-o", "NumberOfPasswordPrompts=1") + args = append(args, command) password, err := getVault(server.Alias, "ssh_password") if err != nil { return "", fmt.Errorf("get password from vault: %w", err) @@ -130,7 +135,8 @@ func RunCommandOutput(cfg *config.Config, server *model.Server, getVault VaultFu } return output, nil case model.AuthKeyPassphrase: - args = append(args, "-o", "NumberOfPasswordPrompts=1", command) + args = insertBeforeTarget(args, "-o", "NumberOfPasswordPrompts=1") + args = append(args, command) passphrase, err := getVault(server.Alias, "key_passphrase") if err != nil { return "", fmt.Errorf("get key passphrase from vault: %w", err) @@ -141,7 +147,8 @@ func RunCommandOutput(cfg *config.Config, server *model.Server, getVault VaultFu } return output, nil default: - args = append(args, "-o", "BatchMode=yes", command) + args = insertBeforeTarget(args, "-o", "BatchMode=yes") + args = append(args, command) cmd := exec.Command(cfg.SSH.Binary, args...) output, err := cmd.CombinedOutput() if err != nil { @@ -151,104 +158,88 @@ func RunCommandOutput(cfg *config.Config, server *model.Server, getVault VaultFu } } -func Test(cfg *config.Config, server *model.Server, getVault VaultFunc) (bool, string) { +func RunCommandOutput(cfg *config.Config, server *model.Server, getVault VaultFunc, command string) (string, error) { + return RunCommandOutputResolved(cfg, server, nil, getVault, command) +} + +func TestResolved(cfg *config.Config, server *model.Server, resolve ProfileResolver, getVault VaultFunc) (bool, string) { if err := EnsureSSHBinary(cfg.SSH.Binary); err != nil { return false, err.Error() } - - args := BuildSSHArgsSimple(server) - args = append(args, "-o", fmt.Sprintf("ConnectTimeout=%d", cfg.SSH.ConnectTimeoutSec)) + invocation, err := PrepareSSHInvocation(server, nil, false, resolve) + if err != nil { + return false, err.Error() + } + defer invocation.Cleanup() + args := insertBeforeTarget(invocation.Args, "-o", fmt.Sprintf("ConnectTimeout=%d", cfg.SSH.ConnectTimeoutSec)) switch server.AuthMethod { case model.AuthPassword: - args = append(args, "-o", "NumberOfPasswordPrompts=1") + args = insertBeforeTarget(args, "-o", "NumberOfPasswordPrompts=1") password, err := getVault(server.Alias, "ssh_password") if err != nil { return false, fmt.Sprintf("vault error: %v", err) } - return testWithPassword(cfg, args, password) - + return testWithPassword(cfg, append(args, cfg.SSH.TestCommand), password) case model.AuthKeyPassphrase: - args = append(args, "-o", "NumberOfPasswordPrompts=1") + args = insertBeforeTarget(args, "-o", "NumberOfPasswordPrompts=1") passphrase, err := getVault(server.Alias, "key_passphrase") if err != nil { return false, fmt.Sprintf("vault error: %v", err) } - return testWithPassword(cfg, args, passphrase) - + return testWithPassword(cfg, append(args, cfg.SSH.TestCommand), passphrase) default: - // key and agent auth should not prompt during tests. - args = append(args, "-o", "BatchMode=yes") + args = insertBeforeTarget(args, "-o", "BatchMode=yes") args = append(args, cfg.SSH.TestCommand) - cmd := exec.Command(cfg.SSH.Binary, args...) cmd.Stdin = nil - output, err := cmd.CombinedOutput() if err != nil { return false, strings.TrimSpace(string(output)) } - result := strings.TrimSpace(string(output)) - if result == "SSHKEEPER_OK" { + if result == "SSHKEEPER_OK" || strings.Contains(result, "SSHKEEPER_OK") { return true, "" } return false, result } } -// testWithPassword tests SSH connection with password auth via PTY-wrapper. -// It connects, sends the password, runs the test command, and checks the output. -func testWithPassword(cfg *config.Config, args []string, password string) (bool, string) { - args = append(args, cfg.SSH.TestCommand) +func Test(cfg *config.Config, server *model.Server, getVault VaultFunc) (bool, string) { + return TestResolved(cfg, server, nil, getVault) +} +func testWithPassword(cfg *config.Config, args []string, password string) (bool, string) { ok, output := connectWithPasswordAndRead(cfg.SSH.Binary, args, password, cfg.SSH.ConnectTimeoutSec) if !ok { return false, output } - result := strings.TrimSpace(output) - if result == "SSHKEEPER_OK" { - return true, "" - } - // The output might have the test command echo before the result - if strings.Contains(result, "SSHKEEPER_OK") { + if result == "SSHKEEPER_OK" || strings.Contains(result, "SSHKEEPER_OK") { return true, "" } return false, result } +func ConnectWithForwardsResolved(cfg *config.Config, server *model.Server, forwards []*model.Forward, forwardOnly bool, resolve ProfileResolver, getVault VaultFunc) error { + if err := EnsureSSHBinary(cfg.SSH.Binary); err != nil { + return err + } + invocation, err := PrepareSSHInvocation(server, forwards, forwardOnly, resolve) + if err != nil { + return err + } + defer invocation.Cleanup() + return runPrepared(cfg, invocation.Args, server, getVault) +} + func ConnectWithArgs(cfg *config.Config, args []string, vaultFunc VaultFunc, server *model.Server) error { if err := EnsureSSHBinary(cfg.SSH.Binary); err != nil { return err } - - switch server.AuthMethod { - case model.AuthPassword: - password, err := vaultFunc(server.Alias, "ssh_password") - if err != nil { - return fmt.Errorf("get password from vault: %w", err) - } - return ConnectWithPassword(cfg.SSH.Binary, args, password) - - case model.AuthKeyPassphrase: - passphrase, err := vaultFunc(server.Alias, "key_passphrase") - if err != nil { - return fmt.Errorf("get key passphrase from vault: %w", err) - } - return ConnectWithPassword(cfg.SSH.Binary, args, passphrase) - - default: - cmd := exec.Command(cfg.SSH.Binary, args...) - cmd.Stdin = os.Stdin - cmd.Stdout = os.Stdout - cmd.Stderr = os.Stderr - if err := cmd.Start(); err != nil { - return fmt.Errorf("start ssh: %w", err) - } - return cmd.Wait() - } + return runPrepared(cfg, args, server, vaultFunc) } + func BuildForwardArgs(forwards []*model.Forward, exitOnForwardFailure bool) []string { var args []string for _, f := range forwards { diff --git a/internal/ssh/planner.go b/internal/ssh/planner.go new file mode 100644 index 0000000..5af14d2 --- /dev/null +++ b/internal/ssh/planner.go @@ -0,0 +1,225 @@ +package ssh + +import ( + "fmt" + "os" + "sort" + "strings" + + "github.com/mirivlad/sshkeeper/internal/model" +) + +// ProfileResolver resolves a stable sshkeeper server ID for route planning. +type ProfileResolver func(serverID int64) (*model.Server, error) + +type PlannedHop struct { + Server *model.Server + Raw string +} + +type ConnectionPlan struct { + Target *model.Server + Hops []PlannedHop + UsesProfiles bool +} + +func PlanConnection(target *model.Server, resolve ProfileResolver) (*ConnectionPlan, error) { + if target == nil { + return nil, fmt.Errorf("target server is required") + } + plan := &ConnectionPlan{Target: target} + stack := map[int64]bool{} + if target.ID > 0 { + stack[target.ID] = true + } + seenProfiles := map[int64]bool{} + seenRaw := map[string]bool{} + hops, err := flattenRoute(target, resolve, stack, seenProfiles, seenRaw) + if err != nil { + return nil, err + } + plan.Hops = hops + for _, hop := range hops { + if hop.Server != nil { + plan.UsesProfiles = true + break + } + } + return plan, nil +} + +func flattenRoute(owner *model.Server, resolve ProfileResolver, stack, seenProfiles map[int64]bool, seenRaw map[string]bool) ([]PlannedHop, error) { + var result []PlannedHop + for _, hop := range owner.Route.Hops { + if hop.Profile() { + if hop.ServerID <= 0 { + return nil, fmt.Errorf("route profile %q has no stable ID; edit and re-save the route", hop.Alias) + } + if resolve == nil { + return nil, fmt.Errorf("route profile %s requires sshkeeper profile resolution", hop.DisplayName()) + } + if stack[hop.ServerID] { + return nil, fmt.Errorf("route cycle detected at %s", hop.DisplayName()) + } + server, err := resolve(hop.ServerID) + if err != nil { + return nil, fmt.Errorf("resolve route profile %s: %w", hop.DisplayName(), err) + } + if server.AuthMethod == model.AuthPassword || server.AuthMethod == model.AuthKeyPassphrase { + return nil, fmt.Errorf("jump profile %s uses %s authentication; password/passphrase jump profiles are not supported by the current OpenSSH vault flow", server.Alias, server.AuthMethod) + } + stack[server.ID] = true + nested, err := flattenRoute(server, resolve, stack, seenProfiles, seenRaw) + delete(stack, server.ID) + if err != nil { + return nil, err + } + result = append(result, nested...) + if seenProfiles[server.ID] { + return nil, fmt.Errorf("route resolves to duplicate profile hop %s", server.Alias) + } + seenProfiles[server.ID] = true + result = append(result, PlannedHop{Server: server}) + continue + } + raw := strings.TrimSpace(hop.Raw) + if raw == "" { + return nil, fmt.Errorf("route contains an empty raw hop") + } + if seenRaw[raw] { + return nil, fmt.Errorf("route resolves to duplicate raw hop %q", raw) + } + seenRaw[raw] = true + result = append(result, PlannedHop{Raw: raw}) + } + return result, nil +} + +func syntheticProfileHost(id int64) string { + return fmt.Sprintf("sshkeeper-profile-%d", id) +} + +func syntheticTargetHost(target *model.Server) string { + if target.ID > 0 { + return fmt.Sprintf("sshkeeper-target-%d", target.ID) + } + return "sshkeeper-target-unsaved" +} + +func appendHostBlock(b *strings.Builder, hostAlias string, server *model.Server) { + fmt.Fprintf(b, "Host %s\n", hostAlias) + fmt.Fprintf(b, " HostName %s\n", server.Host) + port := server.Port + if port == 0 { + port = 22 + } + fmt.Fprintf(b, " Port %d\n", port) + if server.User != "" { + fmt.Fprintf(b, " User %s\n", server.User) + } + if server.IdentityFile != "" && server.AuthMethod != model.AuthPassword && server.AuthMethod != model.AuthAgent { + fmt.Fprintf(b, " IdentityFile %s\n", server.IdentityFile) + } + fmt.Fprintln(b, " StrictHostKeyChecking accept-new") + fmt.Fprintln(b) +} + +// OpenSSHConfig renders the deterministic temporary config used when a route +// references sshkeeper profiles. The user's normal config is included first so +// raw OpenSSH jump targets keep their existing configuration. +func (p *ConnectionPlan) OpenSSHConfig() string { + var b strings.Builder + b.WriteString("# Temporary config generated by sshkeeper\n") + b.WriteString("Include ~/.ssh/config\n\n") + profiles := map[int64]*model.Server{} + for _, hop := range p.Hops { + if hop.Server != nil { + profiles[hop.Server.ID] = hop.Server + } + } + ids := make([]int64, 0, len(profiles)) + for id := range profiles { + ids = append(ids, id) + } + sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] }) + for _, id := range ids { + appendHostBlock(&b, syntheticProfileHost(id), profiles[id]) + } + appendHostBlock(&b, syntheticTargetHost(p.Target), p.Target) + if len(p.Hops) > 0 { + parts := make([]string, 0, len(p.Hops)) + for _, hop := range p.Hops { + if hop.Server != nil { + parts = append(parts, syntheticProfileHost(hop.Server.ID)) + } else { + parts = append(parts, hop.Raw) + } + } + // OpenSSH uses the first value it obtains. Append a target-specific + // stanza after the generic one with ProxyJump before any competing rule. + fmt.Fprintf(&b, "Host %s\n", syntheticTargetHost(p.Target)) + fmt.Fprintf(&b, " ProxyJump %s\n\n", strings.Join(parts, ",")) + } + return b.String() +} + +type PreparedInvocation struct { + Args []string + ConfigPath string +} + +func (p *PreparedInvocation) Cleanup() { + if p != nil && p.ConfigPath != "" { + _ = os.Remove(p.ConfigPath) + p.ConfigPath = "" + } +} + +func enabledForwards(forwards []*model.Forward) []*model.Forward { + result := make([]*model.Forward, 0, len(forwards)) + for _, forward := range forwards { + if forward != nil && forward.Enabled { + result = append(result, forward) + } + } + return result +} + +func PrepareSSHInvocation(server *model.Server, forwards []*model.Forward, forwardOnly bool, resolve ProfileResolver) (*PreparedInvocation, error) { + plan, err := PlanConnection(server, resolve) + if err != nil { + return nil, err + } + active := enabledForwards(forwards) + if !plan.UsesProfiles { + return &PreparedInvocation{Args: BuildSSHArgs(server, active, forwardOnly)}, nil + } + file, err := os.CreateTemp("", "sshkeeper-*.conf") + if err != nil { + return nil, fmt.Errorf("create temporary ssh config: %w", err) + } + path := file.Name() + if err := file.Chmod(0600); err != nil { + file.Close() + os.Remove(path) + return nil, err + } + if _, err := file.WriteString(plan.OpenSSHConfig()); err != nil { + file.Close() + os.Remove(path) + return nil, err + } + if err := file.Close(); err != nil { + os.Remove(path) + return nil, err + } + args := []string{"-F", path} + if len(active) > 0 { + args = append(args, BuildForwardArgs(active, true)...) + } + if forwardOnly { + args = append(args, "-N") + } + args = append(args, syntheticTargetHost(server)) + return &PreparedInvocation{Args: args, ConfigPath: path}, nil +}