@@ -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
2894const (
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 !!!\n Remote key presented: %s\n This could be a man-in-the-middle attack" , host , presentedKey )
205+ return fmt .Errorf ("!!! HOST KEY MISMATCH FOR %s !!!\n Remote key presented: %s\n This 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 )
0 commit comments