Skip to content

Commit 1213b59

Browse files
author
AsciiMaik
committed
Implement host:port normalization and add tests for host/port parsing functions
1 parent 83e8178 commit 1213b59

3 files changed

Lines changed: 146 additions & 73 deletions

File tree

‎cmd/keymaster/main.go‎

Lines changed: 19 additions & 43 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,6 @@ import (
1616
"errors"
1717
"fmt"
1818
"log"
19-
"net"
2019
"os"
2120
"strings"
2221
"sync"
@@ -450,15 +449,17 @@ step before Keymaster can manage a new host.`,
450449
} else {
451450
hostname = target // Assume the whole string is the hostname if no '@'
452451
}
452+
// Always use host:port form with default :22
453+
canonicalHost := deploy.CanonicalizeHostPort(hostname)
453454

454-
fmt.Println(i18n.T("trust_host.retrieving_key", hostname))
455-
key, err := deploy.GetRemoteHostKey(hostname)
455+
fmt.Println(i18n.T("trust_host.retrieving_key", canonicalHost))
456+
key, err := deploy.GetRemoteHostKey(canonicalHost)
456457
if err != nil {
457458
log.Fatalf("%s", i18n.T("trust_host.error_get_key", err))
458459
}
459460

460461
fingerprint := ssh.FingerprintSHA256(key) // Use standard ssh package
461-
fmt.Printf("\n%s\n", i18n.T("trust_host.authenticity_warning_1", hostname))
462+
fmt.Printf("\n%s\n", i18n.T("trust_host.authenticity_warning_1", canonicalHost))
462463
fmt.Printf("%s\n", i18n.T("trust_host.authenticity_warning_2", key.Type(), fingerprint))
463464

464465
if warning := sshkey.CheckHostKeyAlgorithm(key); warning != "" {
@@ -473,12 +474,11 @@ step before Keymaster can manage a new host.`,
473474
}
474475

475476
keyStr := string(ssh.MarshalAuthorizedKey(key)) // Use standard ssh package
476-
normalized := normalizeKnownHostKeyName(hostname)
477-
if err := db.AddKnownHostKey(normalized, keyStr); err != nil {
477+
if err := db.AddKnownHostKey(canonicalHost, keyStr); err != nil {
478478
log.Fatalf("%s", i18n.T("trust_host.error_save_key", err))
479479
}
480480

481-
fmt.Printf("%s\n", i18n.T("trust_host.added_success", normalized, key.Type()))
481+
fmt.Printf("%s\n", i18n.T("trust_host.added_success", canonicalHost, key.Type()))
482482
},
483483
}
484484

@@ -564,13 +564,10 @@ Each account with a label will use the label as the Host alias.`,
564564
output.WriteString(fmt.Sprintf(" HostName %s\n", account.Hostname))
565565
output.WriteString(fmt.Sprintf(" User %s\n", account.Username))
566566

567-
// Parse hostname for port if it contains one
568-
_, port := account.Hostname, "22"
569-
if idx := strings.LastIndex(account.Hostname, ":"); idx > 0 {
570-
// Check if it's IPv6 by looking for multiple colons
571-
if strings.Count(account.Hostname, ":") == 1 {
572-
port = account.Hostname[idx+1:]
573-
}
567+
// Parse hostname for port (supports IPv4/IPv6 and names)
568+
_, port, _ := deploy.ParseHostPort(account.Hostname)
569+
if port == "" {
570+
port = "22"
574571
}
575572
if port != "22" {
576573
output.WriteString(fmt.Sprintf(" Port %s\n", port))
@@ -600,27 +597,6 @@ func promptForConfirmation(prompt string) string {
600597
return strings.TrimSpace(strings.ToLower(answer))
601598
}
602599

603-
// normalizeKnownHostKeyName normalizes a hostname for storage in the known hosts database
604-
// so that it matches lookups performed during SSH handshakes. It removes any port and
605-
// strips IPv6 brackets, returning just the host portion.
606-
func normalizeKnownHostKeyName(h string) string {
607-
h = strings.TrimSpace(h)
608-
if h == "" {
609-
return h
610-
}
611-
// If a port is present (e.g., "example.com:2222" or "[2001:db8::1]:2222"),
612-
// SplitHostPort returns the host without brackets for IPv6.
613-
if host, _, err := net.SplitHostPort(h); err == nil {
614-
return host
615-
}
616-
// If it's a bracketed IPv6 without port like "[2001:db8::1]", trim brackets.
617-
if strings.HasPrefix(h, "[") && strings.HasSuffix(h, "]") {
618-
return strings.TrimSuffix(strings.TrimPrefix(h, "["), "]")
619-
}
620-
// Otherwise, return as-is (covers plain hostnames or raw IPv6 without brackets/port).
621-
return h
622-
}
623-
624600
// runDeploymentForAccount is a simple wrapper for the CLI to match the
625601
// signature required by runParallelTasks. It calls the centralized
626602
// deployment logic from the deploy package.
@@ -656,7 +632,7 @@ Example (Full Restore):
656632

657633
backupData, err := readCompressedBackup(inputFile)
658634
if err != nil {
659-
log.Fatalf(i18n.T("restore.cli_error_read", err))
635+
log.Fatalf("%s", i18n.T("restore.cli_error_read", err))
660636
}
661637

662638
if fullRestore {
@@ -668,7 +644,7 @@ Example (Full Restore):
668644
}
669645

670646
if err != nil {
671-
log.Fatalf(i18n.T("restore.cli_error_import", err))
647+
log.Fatalf("%s", i18n.T("restore.cli_error_import", err))
672648
}
673649

674650
fmt.Println(i18n.T("restore.cli_success"))
@@ -734,11 +710,11 @@ Examples:
734710

735711
backupData, err := db.ExportDataForBackup()
736712
if err != nil {
737-
log.Fatalf(i18n.T("backup.cli_error_export", err))
713+
log.Fatalf("%s", i18n.T("backup.cli_error_export", err))
738714
}
739715

740716
if err := writeCompressedBackup(outputFile, backupData); err != nil {
741-
log.Fatalf(i18n.T("backup.cli_error_write", err))
717+
log.Fatalf("%s", i18n.T("backup.cli_error_write", err))
742718
}
743719

744720
fmt.Println(i18n.T("backup.cli_success", outputFile))
@@ -792,30 +768,30 @@ Example:
792768
targetDSN, _ := cmd.Flags().GetString("dsn")
793769

794770
if targetType == "" || targetDSN == "" {
795-
log.Fatalf(i18n.T("migrate.cli_error_flags"))
771+
log.Fatalf("%s", i18n.T("migrate.cli_error_flags"))
796772
}
797773

798774
// --- 1. Backup from source DB ---
799775
fmt.Println(i18n.T("migrate.cli_starting_backup"))
800776
backupData, err := db.ExportDataForBackup()
801777
if err != nil {
802-
log.Fatalf(i18n.T("migrate.cli_error_backup", err))
778+
log.Fatalf("%s", i18n.T("migrate.cli_error_backup", err))
803779
}
804780
fmt.Println(i18n.T("migrate.cli_backup_success"))
805781

806782
// --- 2. Connect to target DB and run migrations ---
807783
fmt.Println(i18n.T("migrate.cli_connecting_target", targetType))
808784
targetStore, err := initTargetDB(targetType, targetDSN)
809785
if err != nil {
810-
log.Fatalf(i18n.T("migrate.cli_error_connect", err))
786+
log.Fatalf("%s", i18n.T("migrate.cli_error_connect", err))
811787
}
812788
fmt.Println(i18n.T("migrate.cli_connect_success"))
813789

814790
// --- 3. Restore to target DB ---
815791
fmt.Println(i18n.T("migrate.cli_starting_restore"))
816792
// We call the method directly on our temporary store instance.
817793
if err := targetStore.ImportDataFromBackup(backupData); err != nil {
818-
log.Fatalf(i18n.T("migrate.cli_error_restore", err))
794+
log.Fatalf("%s", i18n.T("migrate.cli_error_restore", err))
819795
}
820796

821797
fmt.Println(i18n.T("migrate.cli_success"))

‎internal/deploy/ssh.go‎

Lines changed: 90 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@ import (
1616
"io"
1717
"net"
1818
"path"
19+
"regexp"
1920
"strings"
2021
"time"
2122

@@ -24,6 +25,71 @@ import (
2425
"golang.org/x/crypto/ssh"
2526
)
2627

28+
// CanonicalizeHostPort returns a normalized host:port string.
29+
// - If no port is provided, :22 is assumed.
30+
// - IPv6 literals will be bracketed as needed (e.g., [2001:db8::1]:22).
31+
// - If input is of the form user@host, the user part is discarded.
32+
// StripIPv6Brackets removes surrounding [ ] from an IPv6 literal if present.
33+
func StripIPv6Brackets(host string) string {
34+
if strings.HasPrefix(host, "[") && strings.HasSuffix(host, "]") {
35+
return strings.TrimSuffix(strings.TrimPrefix(host, "["), "]")
36+
}
37+
return host
38+
}
39+
40+
// ParseHostPort splits an address into host and port.
41+
// Behavior:
42+
// - Accepts host, host:port, [ipv6], [ipv6]:port, ipv6, ipv6:port
43+
// - Returns port "" if not specified
44+
// - Returns host without IPv6 brackets
45+
func ParseHostPort(addr string) (host string, port string, err error) {
46+
s := strings.TrimSpace(addr)
47+
if s == "" {
48+
return "", "", fmt.Errorf("empty address")
49+
}
50+
// Strip user@ part if present
51+
if at := strings.LastIndex(s, "@"); at != -1 {
52+
s = s[at+1:]
53+
}
54+
55+
// If bracketed IPv6
56+
if strings.HasPrefix(s, "[") {
57+
// regex: ^\[([^\]]+)\](?::(\d+))?$ -> capture host and optional port
58+
re := regexp.MustCompile(`^\[([^\]]+)\](?::(\d+))?$`)
59+
m := re.FindStringSubmatch(s)
60+
if m == nil {
61+
return "", "", fmt.Errorf("invalid bracketed IPv6: %s", s)
62+
}
63+
return m[1], m[2], nil
64+
}
65+
66+
// Try net.SplitHostPort for host:port or ipv6:port (unbracketed)
67+
if h, p, e := net.SplitHostPort(s); e == nil {
68+
return h, p, nil
69+
}
70+
71+
// No port specified, whole string is host (could be ipv4, name, or unbracketed ipv6)
72+
return s, "", nil
73+
}
74+
75+
// JoinHostPort joins host and port into a canonical host:port.
76+
// - If port is empty, defaultPort is used.
77+
// - IPv6 hosts will be bracketed.
78+
func JoinHostPort(host, port, defaultPort string) string {
79+
h := StripIPv6Brackets(strings.TrimSpace(host))
80+
p := strings.TrimSpace(port)
81+
if p == "" {
82+
p = defaultPort
83+
}
84+
return net.JoinHostPort(h, p)
85+
}
86+
87+
// CanonicalizeHostPort returns host:port form using default 22 if missing.
88+
func CanonicalizeHostPort(input string) string {
89+
host, port, _ := ParseHostPort(input)
90+
return JoinHostPort(host, port, "22")
91+
}
92+
2793
// Default timeout values for SSH operations
2894
const (
2995
// DefaultConnectionTimeout is the default timeout for establishing SSH connections
@@ -89,17 +155,13 @@ func newDeployerInternal(host, user, privateKey string, config *ConnectionConfig
89155
var hostKeyCallback ssh.HostKeyCallback
90156

91157
if isBootstrap {
92-
// For bootstrap, accept any host key and save it
158+
// For bootstrap, accept any host key and save it as canonical host:port
93159
hostKeyCallback = func(hostname string, remote net.Addr, key ssh.PublicKey) error {
94-
// Strip port if present
95-
hostOnly, _, err := net.SplitHostPort(hostname)
96-
if err != nil {
97-
hostOnly = hostname
98-
}
160+
canonical := CanonicalizeHostPort(hostname)
99161

100162
// Save the host key for future connections
101163
presentedKey := string(ssh.MarshalAuthorizedKey(key))
102-
if err := db.AddKnownHostKey(hostOnly, presentedKey); err != nil {
164+
if err := db.AddKnownHostKey(canonical, presentedKey); err != nil {
103165
// Log error but don't fail the connection
104166
// The host key can be added manually later
105167
}
@@ -109,42 +171,46 @@ func newDeployerInternal(host, user, privateKey string, config *ConnectionConfig
109171
} else {
110172
// Normal mode: verify host keys
111173
hostKeyCallback = func(hostname string, remote net.Addr, key ssh.PublicKey) error {
112-
// The hostname passed to the callback can include the port. We need to strip it
113-
// to ensure we're looking up the correct key in our database.
114-
host, _, err := net.SplitHostPort(hostname)
115-
if err != nil {
116-
// If SplitHostPort fails, it means there was no port, so we use the original string.
117-
host = hostname
118-
}
174+
// Always check canonical host:port first
175+
canonical := CanonicalizeHostPort(hostname)
119176

120177
// The key is presented in the format "ssh-ed25519 AAA..."
121178
presentedKey := string(ssh.MarshalAuthorizedKey(key))
122179

123-
// Check if we have a trusted key for this host in our database.
124-
knownKey, err := db.GetKnownHostKey(host)
180+
// Check if we have a trusted key for this canonical host:port in our database.
181+
knownKey, err := db.GetKnownHostKey(canonical)
125182
if err != nil {
126183
return fmt.Errorf("failed to query known_hosts database: %w", err)
127184
}
128185

129186
// If we don't have a key, this is the first connection.
130187
if knownKey == "" {
131-
return fmt.Errorf("unknown host key for %s. run 'keymaster trust-host' to add it", host)
188+
// Backward compatibility: try legacy host-only key (without port)
189+
if hostOnly, _, err := net.SplitHostPort(canonical); err == nil {
190+
legacyKey, lerr := db.GetKnownHostKey(hostOnly)
191+
if lerr != nil {
192+
return fmt.Errorf("failed to query known_hosts database: %w", lerr)
193+
}
194+
if legacyKey != "" {
195+
knownKey = legacyKey
196+
}
197+
}
198+
if knownKey == "" {
199+
return fmt.Errorf("unknown host key for %s. run 'keymaster trust-host' to add it", canonical)
200+
}
132201
}
133202

134203
// If the key exists, it must match exactly.
135204
if knownKey != presentedKey {
136-
return fmt.Errorf("!!! HOST KEY MISMATCH FOR %s !!!\nRemote key presented: %s\nThis could be a man-in-the-middle attack", host, presentedKey)
205+
return fmt.Errorf("!!! HOST KEY MISMATCH FOR %s !!!\nRemote key presented: %s\nThis could be a man-in-the-middle attack", canonical, presentedKey)
137206
}
138207

139208
return nil // Host key is trusted.
140209
}
141210
}
142211

143212
// Add port 22 if not specified.
144-
addr := host
145-
if _, _, err := net.SplitHostPort(host); err != nil {
146-
addr = net.JoinHostPort(host, "22")
147-
}
213+
addr := CanonicalizeHostPort(host)
148214
var client *ssh.Client
149215

150216
// If a private key is provided, use it exclusively. This is the standard path
@@ -234,10 +300,7 @@ func newDeployerWithExpectedHostKey(host, user, privateKey string, config *Conne
234300
}
235301

236302
// Add port 22 if not specified
237-
addr := host
238-
if _, _, err := net.SplitHostPort(host); err != nil {
239-
addr = net.JoinHostPort(host, "22")
240-
}
303+
addr := CanonicalizeHostPort(host)
241304

242305
// Parse the private key
243306
signer, err := ssh.ParsePrivateKey([]byte(privateKey))
@@ -459,10 +522,7 @@ func GetRemoteHostKeyWithTimeout(host string, timeout time.Duration) (ssh.Public
459522
Timeout: timeout,
460523
}
461524

462-
addr := host
463-
if _, _, err := net.SplitHostPort(host); err != nil {
464-
addr = net.JoinHostPort(host, "22")
465-
}
525+
addr := CanonicalizeHostPort(host)
466526

467527
// We expect ssh.Dial to fail with our specific error.
468528
_, err := ssh.Dial("tcp", addr, config)

‎internal/deploy/ssh_test.go‎

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -167,3 +167,40 @@ func stringContains(s, substr string) bool {
167167
}
168168
return false
169169
}
170+
171+
func TestHostPortHelpers(t *testing.T) {
172+
cases := []struct {
173+
in string
174+
host string
175+
port string
176+
canon string
177+
}{
178+
{"example.com", "example.com", "", "example.com:22"},
179+
{"example.com:2222", "example.com", "2222", "example.com:2222"},
180+
{"192.168.1.10", "192.168.1.10", "", "192.168.1.10:22"},
181+
{"192.168.1.10:2200", "192.168.1.10", "2200", "192.168.1.10:2200"},
182+
{"[2001:db8::1]", "2001:db8::1", "", "[2001:db8::1]:22"},
183+
{"[2001:db8::1]:2200", "2001:db8::1", "2200", "[2001:db8::1]:2200"},
184+
{"2001:db8::1", "2001:db8::1", "", "[2001:db8::1]:22"},
185+
{"user@example.com", "example.com", "", "example.com:22"},
186+
{"user@[2001:db8::1]:2222", "2001:db8::1", "2222", "[2001:db8::1]:2222"},
187+
}
188+
for _, c := range cases {
189+
h, p, err := ParseHostPort(c.in)
190+
if err != nil {
191+
t.Fatalf("unexpected error parsing %q: %v", c.in, err)
192+
}
193+
if h != c.host || p != c.port {
194+
t.Errorf("ParseHostPort(%q) => host=%q port=%q; want host=%q port=%q", c.in, h, p, c.host, c.port)
195+
}
196+
canon := CanonicalizeHostPort(c.in)
197+
if canon != c.canon {
198+
t.Errorf("CanonicalizeHostPort(%q) => %q; want %q", c.in, canon, c.canon)
199+
}
200+
// Join should reconstruct canon from components
201+
joined := JoinHostPort(h, p, "22")
202+
if joined != c.canon {
203+
t.Errorf("JoinHostPort(%q,%q,22) => %q; want %q", h, p, joined, c.canon)
204+
}
205+
}
206+
}

0 commit comments

Comments
 (0)