refactor: normalize server routes and stable identities
This commit is contained in:
parent
a19b3deb24
commit
3cc21a7b22
170
cmd/secrets.go
170
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 {
|
||||
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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
WHERE st.server_id = servers.id AND t.name LIKE ?
|
||||
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 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 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)
|
||||
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()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
@ -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"`
|
||||
|
|
|
|||
|
|
@ -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, ",")
|
||||
}
|
||||
|
|
@ -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")
|
||||
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
|
||||
}
|
||||
|
||||
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 {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
Loading…
Reference in New Issue