Files
MyGoNavi/internal/ssh/ssh.go

558 lines
16 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package ssh
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"io"
"net"
"os"
"strconv"
"strings"
"sync"
"time"
"GoNavi-Wails/internal/connection"
"GoNavi-Wails/internal/logger"
"github.com/go-sql-driver/mysql"
"golang.org/x/crypto/ssh"
"golang.org/x/sync/singleflight"
)
// ViaSSHDialer registers a custom network for MySQL that proxies through SSH
type ViaSSHDialer struct {
sshClient *ssh.Client
}
func (d *ViaSSHDialer) Dial(ctx context.Context, addr string) (net.Conn, error) {
return dialContext(ctx, d.sshClient, "tcp", addr)
}
func dialContext(ctx context.Context, client *ssh.Client, network, addr string) (net.Conn, error) {
type result struct {
conn net.Conn
err error
}
ch := make(chan result, 1)
go func() {
c, err := client.Dial(network, addr)
ch <- result{conn: c, err: err}
}()
select {
case <-ctx.Done():
go func() {
r := <-ch
if r.conn != nil {
_ = r.conn.Close()
}
}()
return nil, ctx.Err()
case r := <-ch:
return r.conn, r.err
}
}
// connectSSH establishes an SSH connection and returns a Dialer
func connectSSH(config connection.SSHConfig) (*ssh.Client, error) {
logger.Infof("开始建立 SSH 连接:地址=%s:%d 用户=%s", config.Host, config.Port, config.User)
authMethods := []ssh.AuthMethod{}
if keyPath := strings.TrimSpace(config.KeyPath); keyPath != "" {
key, err := os.ReadFile(keyPath)
if err != nil {
logger.Warnf("读取 SSH 私钥失败:路径=%s原因%v", keyPath, err)
return nil, fmt.Errorf("failed to read SSH private key %s: %w", keyPath, err)
}
signer, err := ssh.ParsePrivateKey(key)
if err != nil {
logger.Warnf("解析 SSH 私钥失败:路径=%s原因%v", keyPath, err)
var passphraseErr *ssh.PassphraseMissingError
if errors.As(err, &passphraseErr) {
return nil, fmt.Errorf("SSH private key %s is encrypted with a passphrase; passphrase-protected keys are not supported", keyPath)
}
return nil, fmt.Errorf("failed to parse SSH private key %s: %w", keyPath, err)
}
authMethods = append(authMethods, ssh.PublicKeys(signer))
}
if config.Password != "" {
authMethods = append(authMethods, ssh.Password(config.Password))
}
if len(authMethods) == 0 {
logger.Warnf("SSH 未配置认证方式(密码或私钥)")
}
sshConfig := &ssh.ClientConfig{
User: config.User,
Auth: authMethods,
HostKeyCallback: ssh.InsecureIgnoreHostKey(), // Use strict checking in production!
Timeout: 5 * time.Second,
}
addr := fmt.Sprintf("%s:%d", config.Host, config.Port)
client, err := ssh.Dial("tcp", addr, sshConfig)
if err != nil {
logger.Error(err, "SSH 连接建立失败:地址=%s 用户=%s", addr, config.User)
return nil, err
}
logger.Infof("SSH 连接建立成功:地址=%s 用户=%s", addr, config.User)
return client, nil
}
// sshNetworkName 按 SSH 目标确定性派生 go-sql-driver 的自定义 network 名。
//
// 必须是确定性的mysql.RegisterDialContext 写入驱动内一张永不回收的全局 map
// DeregisterDialContext 在本仓库无任何调用点)。若每次调用都用时间戳生成新名字,
// 每次(重)连接都会新增一条永久条目并钉住其闭包捕获的 ssh.Client形成随重连线性增长的
// SSH 连接与 goroutine 泄漏。相同目标复用同名注册后map 大小收敛为 SSH 目标个数。
//
// 用 %q 做字段分隔以保证单射host/user 里的引号会被转义),并只取短哈希,避免在
// network 名与日志中泄露认证指纹明文。
func sshNetworkName(key sshClientCacheKey) string {
sum := sha256.Sum256([]byte(fmt.Sprintf("%q %d %q %q", key.host, key.port, key.user, key.auth)))
return "ssh_" + hex.EncodeToString(sum[:8])
}
// RegisterSSHNetwork registers a network name for a specific SSH tunnel
// Returns the network name to use in DSN
func RegisterSSHNetwork(sshConfig connection.SSHConfig) (string, error) {
// 走缓存创建客户端,使其进入 sshClientCache从而能被 CloseAllSSHClients 统一回收;
// 直接调 connectSSH 会产出一个既不入缓存、也无人关闭的孤立客户端。
if _, err := GetOrCreateSSHClient(sshConfig); err != nil {
return "", err
}
netName := sshNetworkName(newSSHClientCacheKey(sshConfig))
logger.Infof("注册 SSH 网络:%s地址=%s:%d 用户=%s", netName, sshConfig.Host, sshConfig.Port, sshConfig.User)
// 闭包在拨号时才取客户端不捕获固定实例GetOrCreateSSHClient 会探测存活并在断开后重建,
// 因此这条注册项对同一目标可长期复用,也不会把一个已死的 client 永久钉在驱动的全局 map 里。
mysql.RegisterDialContext(netName, func(ctx context.Context, addr string) (net.Conn, error) {
client, err := GetOrCreateSSHClient(sshConfig)
if err != nil {
return nil, err
}
return dialContext(ctx, client, "tcp", addr)
})
return netName, nil
}
// DialContextThroughSSH creates a context-aware connection through an SSH tunnel.
func DialContextThroughSSH(ctx context.Context, config connection.SSHConfig, network, address string) (net.Conn, error) {
client, err := GetOrCreateSSHClient(config)
if err != nil {
return nil, fmt.Errorf("failed to establish SSH connection: %w", err)
}
conn, err := dialContext(ctx, client, network, address)
if err != nil {
return nil, fmt.Errorf("failed to connect to %s through SSH tunnel: %w", address, err)
}
logger.Infof("已通过 SSH 隧道连接到:%s", address)
return conn, nil
}
// sshClientCache stores SSH clients to avoid creating multiple connections
var (
sshClientCache = make(map[sshClientCacheKey]*ssh.Client)
sshClientCacheMu sync.RWMutex
sshClientFlights singleflight.Group
connectSSHClient = connectSSH
localForwarders = make(map[forwarderCacheKey]*LocalForwarder)
forwarderMu sync.RWMutex
)
type sshClientCacheKey struct {
host string
port int
user string
auth string
}
type forwarderCacheKey struct {
ssh sshClientCacheKey
remoteHost string
remotePort int
}
func sshAuthFingerprint(config connection.SSHConfig) string {
hasher := sha256.New()
_, _ = hasher.Write([]byte(config.Password))
_, _ = hasher.Write([]byte{0})
_, _ = hasher.Write([]byte(config.KeyPath))
if config.KeyPath != "" {
if st, err := os.Stat(config.KeyPath); err == nil {
_, _ = hasher.Write([]byte{0})
_, _ = hasher.Write([]byte(st.ModTime().UTC().Format(time.RFC3339Nano)))
_, _ = hasher.Write([]byte{0})
_, _ = hasher.Write([]byte(strconv.FormatInt(st.Size(), 10)))
} else {
_, _ = hasher.Write([]byte{0})
_, _ = hasher.Write([]byte("stat_err"))
}
}
sum := hasher.Sum(nil)
return hex.EncodeToString(sum[:8])
}
func newSSHClientCacheKey(config connection.SSHConfig) sshClientCacheKey {
return sshClientCacheKey{
host: config.Host,
port: config.Port,
user: config.User,
auth: sshAuthFingerprint(config),
}
}
func formatSSHClientKeyForLog(key sshClientCacheKey) string {
return fmt.Sprintf("%s:%d 用户=%s", key.host, key.port, key.user)
}
// LocalForwarder represents a local port forwarder through SSH
type LocalForwarder struct {
LocalAddr string
RemoteAddr string
SSHClient *ssh.Client
listener net.Listener
closeChan chan struct{}
closeOnce sync.Once
closed bool
closedMu sync.RWMutex
// shared/cacheKey identify a lease returned by AcquireLocalForwarder.
// The cached forwarder itself keeps shared nil and owns the listener.
shared *LocalForwarder
cacheKey forwarderCacheKey
leaseOnce sync.Once
refCount int // guarded by forwarderMu; meaningful only on the cached forwarder
}
// NewLocalForwarder creates a new local port forwarder
// It listens on a random local port and forwards all connections through SSH tunnel
func NewLocalForwarder(sshConfig connection.SSHConfig, remoteHost string, remotePort int) (*LocalForwarder, error) {
client, err := GetOrCreateSSHClient(sshConfig)
if err != nil {
return nil, fmt.Errorf("failed to establish SSH connection: %w", err)
}
// Listen on localhost with a random port
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
return nil, fmt.Errorf("failed to create local listener: %w", err)
}
localAddr := listener.Addr().String()
remoteAddr := fmt.Sprintf("%s:%d", remoteHost, remotePort)
forwarder := &LocalForwarder{
LocalAddr: localAddr,
RemoteAddr: remoteAddr,
SSHClient: client,
listener: listener,
closeChan: make(chan struct{}),
}
// Start forwarding in background
go forwarder.forward()
logger.Infof("已创建 SSH 端口转发:本地 %s -> 远程 %s", localAddr, remoteAddr)
return forwarder, nil
}
// forward handles the port forwarding
func (f *LocalForwarder) forward() {
for {
localConn, err := f.listener.Accept()
if err != nil {
// Check if we're shutting down
select {
case <-f.closeChan:
return
default:
logger.Warnf("接受本地连接失败:%v", err)
// listener可能已关闭,退出循环
return
}
}
go f.handleConnection(localConn)
}
}
// handleConnection handles a single connection
func (f *LocalForwarder) handleConnection(localConn net.Conn) {
defer localConn.Close()
// Connect to remote through SSH with timeout
remoteConn, err := f.SSHClient.Dial("tcp", f.RemoteAddr)
if err != nil {
logger.Warnf("通过 SSH 连接到远程 %s 失败:%v", f.RemoteAddr, err)
return
}
defer remoteConn.Close()
// Bidirectional copy with error channel
errc := make(chan error, 2)
// Copy from local to remote
go func() {
_, err := io.Copy(remoteConn, localConn)
if err != nil {
logger.Warnf("本地->远程数据复制错误:%v", err)
}
errc <- err
}()
// Copy from remote to local
go func() {
_, err := io.Copy(localConn, remoteConn)
if err != nil {
logger.Warnf("远程->本地数据复制错误:%v", err)
}
errc <- err
}()
// Wait for BOTH goroutines to complete
<-errc
<-errc
}
// Close releases a cached lease, or closes a standalone forwarder created by
// NewLocalForwarder. It is thread-safe and can be called multiple times.
func (f *LocalForwarder) Close() error {
if f == nil {
return nil
}
if f.shared != nil {
var err error
f.leaseOnce.Do(func() {
err = releaseLocalForwarder(f.cacheKey, f.shared)
})
return err
}
return f.closeUnderlying()
}
// Release releases this acquisition. It is an explicit lifecycle alias for
// callers that obtained the forwarder through AcquireLocalForwarder.
func (f *LocalForwarder) Release() error {
return f.Close()
}
func (f *LocalForwarder) closeUnderlying() error {
var err error
f.closeOnce.Do(func() {
f.closedMu.Lock()
f.closed = true
f.closedMu.Unlock()
close(f.closeChan)
err = f.listener.Close()
if err != nil {
logger.Warnf("关闭端口转发监听器失败:%v", err)
}
})
return err
}
// IsClosed returns whether the forwarder is closed
func (f *LocalForwarder) IsClosed() bool {
if f == nil {
return true
}
if f.shared != nil {
return f.shared.IsClosed()
}
f.closedMu.RLock()
defer f.closedMu.RUnlock()
return f.closed
}
// AcquireLocalForwarder acquires a lease on a cached forwarder or creates one.
// Each successful call must be paired with Release. The shared listener is
// closed and evicted only after the last lease is released.
func AcquireLocalForwarder(sshConfig connection.SSHConfig, remoteHost string, remotePort int) (*LocalForwarder, error) {
key := forwarderCacheKey{
ssh: newSSHClientCacheKey(sshConfig),
remoteHost: remoteHost,
remotePort: remotePort,
}
logKey := fmt.Sprintf("%s:%d:%s->%s:%d",
sshConfig.Host, sshConfig.Port, sshConfig.User, remoteHost, remotePort)
forwarderMu.Lock()
if forwarder := localForwarders[key]; forwarder != nil && !forwarder.IsClosed() {
lease := acquireForwarderLeaseLocked(key, forwarder)
forwarderMu.Unlock()
logger.Infof("复用已有端口转发:%s", logKey)
return lease, nil
}
delete(localForwarders, key)
forwarderMu.Unlock()
forwarder, err := NewLocalForwarder(sshConfig, remoteHost, remotePort)
if err != nil {
return nil, err
}
forwarderMu.Lock()
if existing := localForwarders[key]; existing != nil && !existing.IsClosed() {
lease := acquireForwarderLeaseLocked(key, existing)
forwarderMu.Unlock()
_ = forwarder.closeUnderlying()
logger.Infof("复用已有端口转发:%s", logKey)
return lease, nil
}
delete(localForwarders, key)
localForwarders[key] = forwarder
lease := acquireForwarderLeaseLocked(key, forwarder)
forwarderMu.Unlock()
return lease, nil
}
// GetOrCreateLocalForwarder is kept for internal compatibility. New callers
// should use AcquireLocalForwarder so the lease ownership is explicit.
func GetOrCreateLocalForwarder(sshConfig connection.SSHConfig, remoteHost string, remotePort int) (*LocalForwarder, error) {
return AcquireLocalForwarder(sshConfig, remoteHost, remotePort)
}
func acquireForwarderLeaseLocked(key forwarderCacheKey, shared *LocalForwarder) *LocalForwarder {
shared.refCount++
return &LocalForwarder{
LocalAddr: shared.LocalAddr,
RemoteAddr: shared.RemoteAddr,
SSHClient: shared.SSHClient,
shared: shared,
cacheKey: key,
}
}
func releaseLocalForwarder(key forwarderCacheKey, shared *LocalForwarder) error {
forwarderMu.Lock()
if shared.refCount > 0 {
shared.refCount--
}
if shared.refCount > 0 {
forwarderMu.Unlock()
return nil
}
if localForwarders[key] == shared {
delete(localForwarders, key)
}
forwarderMu.Unlock()
return shared.closeUnderlying()
}
// CloseAllForwarders force-closes all cached local forwarders regardless of
// active leases.
func CloseAllForwarders() {
forwarderMu.Lock()
defer forwarderMu.Unlock()
for _, forwarder := range localForwarders {
if forwarder != nil {
forwarder.refCount = 0
_ = forwarder.closeUnderlying()
logger.Infof("已关闭端口转发:本地 %s -> 远程 %s", forwarder.LocalAddr, forwarder.RemoteAddr)
}
}
localForwarders = make(map[forwarderCacheKey]*LocalForwarder)
}
// GetOrCreateSSHClient returns a cached SSH client or creates a new one
func GetOrCreateSSHClient(config connection.SSHConfig) (*ssh.Client, error) {
key := newSSHClientCacheKey(config)
value, err, _ := sshClientFlights.Do(sshClientFlightKey(key), func() (interface{}, error) {
return getOrCreateSSHClient(config, key)
})
if err != nil {
return nil, err
}
client, ok := value.(*ssh.Client)
if !ok || client == nil {
return nil, fmt.Errorf("SSH client creation returned an invalid result")
}
return client, nil
}
func getOrCreateSSHClient(config connection.SSHConfig, key sshClientCacheKey) (*ssh.Client, error) {
sshClientCacheMu.RLock()
client, exists := sshClientCache[key]
sshClientCacheMu.RUnlock()
if exists && client != nil {
// Test if connection is still alive by creating a test session
session, err := client.NewSession()
if err == nil {
session.Close()
logger.Infof("复用已有 SSH 连接:%s", formatSSHClientKeyForLog(key))
return client, nil
}
// Connection is dead, remove from cache
logger.Warnf("SSH 连接已断开,重新建立:%s (错误: %v)", formatSSHClientKeyForLog(key), err)
sshClientCacheMu.Lock()
delete(sshClientCache, key)
sshClientCacheMu.Unlock()
// Try to close the dead client
_ = client.Close()
}
// Create new SSH client
client, err := connectSSHClient(config)
if err != nil {
return nil, err
}
// Cache the client
sshClientCacheMu.Lock()
sshClientCache[key] = client
sshClientCacheMu.Unlock()
logger.Infof("已缓存 SSH 连接:%s", formatSSHClientKeyForLog(key))
return client, nil
}
func sshClientFlightKey(key sshClientCacheKey) string {
return fmt.Sprintf("%q\x00%d\x00%q\x00%s", key.host, key.port, key.user, key.auth)
}
// DialThroughSSH creates a connection through SSH tunnel
// This is a generic dialer that can be used by any database driver
func DialThroughSSH(config connection.SSHConfig, network, address string) (net.Conn, error) {
client, err := GetOrCreateSSHClient(config)
if err != nil {
return nil, fmt.Errorf("failed to establish SSH connection: %w", err)
}
conn, err := client.Dial(network, address)
if err != nil {
return nil, fmt.Errorf("failed to connect to %s through SSH tunnel: %w", address, err)
}
logger.Infof("已通过 SSH 隧道连接到:%s", address)
return conn, nil
}
// CloseAllSSHClients closes all cached SSH clients
func CloseAllSSHClients() {
sshClientCacheMu.Lock()
defer sshClientCacheMu.Unlock()
for key, client := range sshClientCache {
if client != nil {
_ = client.Close()
logger.Infof("已关闭 SSH 连接:%s", formatSSHClientKeyForLog(key))
}
}
sshClientCache = make(map[sshClientCacheKey]*ssh.Client)
}