package ssh import ( "errors" "fmt" "net" "os" cryptossh "golang.org/x/crypto/ssh" "golang.org/x/crypto/ssh/knownhosts" ) func hostKeyCallback(options Options) (cryptossh.HostKeyCallback, error) { switch options.HostKeyPolicy { case HostKeyPolicyOff: return cryptossh.InsecureIgnoreHostKey(), nil case HostKeyPolicyStrict: if options.KnownHosts == "" { return nil, fmt.Errorf("known_hosts is required for strict host key checking") } callback, err := knownhosts.New(options.KnownHosts) if err != nil { return nil, fmt.Errorf("load known_hosts %q: %w", options.KnownHosts, err) } return callback, nil case HostKeyPolicyAcceptNew: return acceptNewHostKeyCallback(options) default: return nil, fmt.Errorf("host_key_policy must be strict, accept-new, or off") } } func acceptNewHostKeyCallback(options Options) (cryptossh.HostKeyCallback, error) { var checker cryptossh.HostKeyCallback if options.KnownHosts != "" { loaded, err := knownhosts.New(options.KnownHosts) if err == nil { checker = loaded } else if !errors.Is(err, os.ErrNotExist) { return nil, fmt.Errorf("load known_hosts %q: %w", options.KnownHosts, err) } } return func(hostname string, remote net.Addr, key cryptossh.PublicKey) error { if checker != nil { err := checker(hostname, remote, key) if err == nil { return nil } var keyErr *knownhosts.KeyError if !errors.As(err, &keyErr) { return err } if len(keyErr.Want) > 0 { return fmt.Errorf("host key for %s has changed: %w", hostname, err) } } if options.KnownHosts == "" { return fmt.Errorf("host key for %s is unknown and no writable known_hosts path is available", hostname) } if err := appendKnownHost(options.KnownHosts, hostname, key); err != nil { return err } loaded, err := knownhosts.New(options.KnownHosts) if err == nil { checker = loaded } return nil }, nil } func appendKnownHost(path, host string, key cryptossh.PublicKey) error { file, err := os.OpenFile(path, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o600) if err != nil { return fmt.Errorf("persist accepted host key to known_hosts %q: %w", path, err) } defer file.Close() if _, err := fmt.Fprintln(file, knownhosts.Line([]string{knownhosts.Normalize(host)}, key)); err != nil { return fmt.Errorf("persist accepted host key to known_hosts %q: %w", path, err) } return nil }