refactor: normalize server routes and stable identities

This commit is contained in:
mirivlad 2026-09-05 18:55:42 +08:00
parent a19b3deb24
commit 3cc21a7b22
9 changed files with 1315 additions and 337 deletions

View File

@ -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 {
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)
data, err := v.Get(oldID)
if err == nil {
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)
}
}
}
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 {
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
}
}

View File

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

View File

@ -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
servers = append(servers, s)
}
if lastTest.Valid {
s.LastTestAt = &lastTest.Time
if err := rows.Err(); err != nil {
return nil, err
}
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)
for _, s := range servers {
s.Route, err = db.loadRoute(s.ID)
if err != nil {
return nil, err
}
s.Tags = tags
servers = append(servers, &s)
s.ProxyJump = s.Route.ProxyJumpString()
s.Tags, err = db.GetServerTags(s.ID)
if err != nil {
return nil, err
}
return servers, rows.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
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 ?
)
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)
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)
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()
}

View File

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

159
internal/db/v040.go Normal file
View File

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

View File

@ -32,8 +32,10 @@ type Server struct {
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"`
@ -143,18 +145,39 @@ func (f *Forward) ForwardTarget() string {
type Tag struct {
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"`

View File

@ -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:<alias> requires an existing sshkeeper profile; raw:<target> 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, ",")
}

View File

@ -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 runPrepared(cfg, args, server, getVault)
}
return cmd.Wait()
}
func Connect(cfg *config.Config, server *model.Server, getVault VaultFunc) error {
return ConnectResolved(cfg, server, nil, getVault)
}
func RunCommand(cfg *config.Config, server *model.Server, getVault VaultFunc, command string) error {
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")
invocation, err := PrepareSSHInvocation(server, nil, false, resolve)
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()
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
}
return runPrepared(cfg, args, server, vaultFunc)
}
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()
}
}
func BuildForwardArgs(forwards []*model.Forward, exitOnForwardFailure bool) []string {
var args []string
for _, f := range forwards {

225
internal/ssh/planner.go Normal file
View File

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