package sftp import ( "fmt" "net" "os" "strconv" "sync" "time" "github.com/pkg/sftp" "golang.org/x/crypto/ssh" ) // 连接建立全流程(拨号+认证+SFTP 子系统)的总超时预算 const connectTimeout = 10 * time.Second // 保活参数:周期发送 keepalive 并等待回复,超过 keepaliveReplyWait 未回复判定连接已死 const ( keepaliveInterval = 30 * time.Second keepaliveReplyWait = 30 * time.Second ) // Client SFTP 客户端封装(单连接) type Client struct { config *Config client *sftp.Client sshClient *ssh.Client stopKeep chan struct{} // 当前 SSH 连接保活循环的停止信号 closed bool // 保活失败或显式关闭后置位,供取连接时剔除 mu sync.Mutex } // Manager 全局 SFTP 连接管理器(以 host:port 为 key 的连接池) type Manager struct { clients sync.Map // map[string]*Client mu sync.Mutex } var globalManager = &Manager{} // GetManager 获取全局连接管理器 func GetManager() *Manager { return globalManager } // Connect 创建或复用 SFTP 连接 func (m *Manager) Connect(config *Config) (*Client, error) { key := fmt.Sprintf("%s:%d", config.Host, config.Port) m.mu.Lock() defer m.mu.Unlock() if existing, ok := m.clients.Load(key); ok { c := existing.(*Client) if c.IsHealthy() { return c, nil } c.Close() m.clients.Delete(key) } c, err := newClient(config) if err != nil { return nil, err } m.clients.Store(key, c) return c, nil } // GetClient 获取已有连接(不复用也不新建) func (m *Manager) GetClient(connID string) *Client { if val, ok := m.clients.Load(connID); ok { return val.(*Client) } return nil } // Disconnect 关闭并移除指定连接 func (m *Manager) Disconnect(host string, port int) { key := fmt.Sprintf("%s:%d", host, port) if val, ok := m.clients.LoadAndDelete(key); ok { val.(*Client).Close() } } // Evict 从连接池剔除连接(底层已断开,仅移除不重复关闭) func (m *Manager) Evict(connID string) { m.clients.Delete(connID) } // Shutdown 关闭所有连接 func (m *Manager) Shutdown() { m.clients.Range(func(key, value any) bool { value.(*Client).Close() m.clients.Delete(key) return true }) } // --- 内部 --- // newClient 建立连接并启动保活循环 func newClient(config *Config) (*Client, error) { c, err := buildClient(config) if err != nil { return nil, err } c.startKeepalive() return c, nil } // buildClient 拨号+认证+SFTP 子系统建立(全程受 connectTimeout 约束),不启动保活 func buildClient(config *Config) (*Client, error) { sshConfig := &ssh.ClientConfig{ Config: ssh.Config{ KeyExchanges: []string{ "curve25519-sha256", "curve25519-sha256@libssh.org", "ecdh-sha2-nistp256", "ecdh-sha2-nistp384", "diffie-hellman-group14-sha256", "diffie-hellman-group14-sha1", }, }, User: config.Username, HostKeyCallback: ssh.InsecureIgnoreHostKey(), Timeout: config.Timeout, } // 认证方式选择 if config.KeyPath != "" { key, err := os.ReadFile(config.KeyPath) if err != nil { return nil, &ConnectionError{Op: "auth", Err: fmt.Errorf("读取密钥文件失败: %w", err)} } var signer ssh.Signer if config.KeyPassphrase != "" { signer, err = ssh.ParsePrivateKeyWithPassphrase(key, []byte(config.KeyPassphrase)) } else { signer, err = ssh.ParsePrivateKey(key) } if err != nil { return nil, &ConnectionError{Op: "auth", Err: fmt.Errorf("解析密钥失败: %w", err)} } sshConfig.Auth = []ssh.AuthMethod{ssh.PublicKeys(signer)} } else if config.Password != "" { pw := config.Password sshConfig.Auth = []ssh.AuthMethod{ ssh.Password(pw), ssh.KeyboardInteractive(func(user, instruction string, questions []string, echos []bool) ([]string, error) { answers := make([]string, len(questions)) for i := range questions { answers[i] = pw } return answers, nil }), } } else { return nil, &ConnectionError{Op: "auth", Err: fmt.Errorf("必须提供密码或密钥文件")} } // JoinHostPort 保证 IPv6 主机地址被方括号包裹,fmt 拼接对 IPv6 产生非法地址 addr := net.JoinHostPort(config.Host, strconv.Itoa(config.Port)) // 拨号与认证、SFTP 子系统握手共享同一总超时预算,认证阶段无界挂起由 deadline 掐断 deadline := time.Now().Add(connectTimeout) sshConn, err := net.DialTimeout("tcp", addr, connectTimeout) if err != nil { return nil, &ConnectionError{Op: "dial", Err: err} } sshConn.SetDeadline(deadline) sshConnConn, chans, reqs, err := ssh.NewClientConn(sshConn, addr, sshConfig) if err != nil { sshConn.Close() return nil, &ConnectionError{Op: "handshake", Err: err} } sshClient := ssh.NewClient(sshConnConn, chans, reqs) sftpClient, err := sftp.NewClient(sshClient) sshConn.SetDeadline(time.Time{}) if err != nil { sshClient.Close() return nil, &ConnectionError{Op: "sftp_init", Err: err} } return &Client{ config: config, client: sftpClient, sshClient: sshClient, }, nil } // startKeepalive 启动当前 SSH 连接的保活循环(须持有 c.mu 或独占 Client 时调用) // 防止空闲连接被 NAT/防火墙杀掉;探测失败时关闭连接并置位 closed,唤醒所有阻塞中的操作 func (c *Client) startKeepalive() { stop := make(chan struct{}) sshClient := c.sshClient c.stopKeep = stop go func() { t := time.NewTicker(keepaliveInterval) defer t.Stop() for { select { case <-stop: return case <-t.C: } if !sshKeepalive(sshClient, keepaliveReplyWait) { c.markDead(sshClient) return } } }() } // sshKeepalive 发送一次 keepalive 请求并限时等待回复,超时或出错视作连接已死 func sshKeepalive(sshClient *ssh.Client, wait time.Duration) bool { done := make(chan error, 1) go func() { _, _, err := sshClient.SendRequest("keepalive@openssh.com", true, nil) done <- err }() select { case err := <-done: return err == nil case <-time.After(wait): return false } } // markDead 标记连接死亡并关闭底层资源 // 仅当容器仍持有该 SSH 连接时生效,避免误杀 reconnect 换上的新连接 func (c *Client) markDead(sshClient *ssh.Client) { c.mu.Lock() defer c.mu.Unlock() if c.sshClient != sshClient { return } c.closeLocked() } // IsClosed 连接是否已断开(保活失败或显式关闭) func (c *Client) IsClosed() bool { c.mu.Lock() defer c.mu.Unlock() return c.closed } // IsHealthy 检查连接是否健康(先取引用再解锁,避免持锁做 I/O) func (c *Client) IsHealthy() bool { c.mu.Lock() client := c.client c.mu.Unlock() if client == nil { return false } _, err := client.Stat("/") return err == nil } // WithRetry 带重试的操作执行(自动处理断线重连) func (c *Client) WithRetry(fn func(*sftp.Client) error) error { const maxRetries = 3 var lastErr error for attempt := 0; attempt < maxRetries; attempt++ { if attempt > 0 { time.Sleep(time.Duration((attempt+1)*2) * time.Second) if reconnectErr := c.reconnect(); reconnectErr != nil { lastErr = reconnectErr continue } } c.mu.Lock() client := c.client c.mu.Unlock() if client == nil { lastErr = fmt.Errorf("SFTP 客户端未初始化") continue } if err := fn(client); err != nil { if isNetworkError(err) { lastErr = err continue } return err } return nil } return fmt.Errorf("操作失败(已重试 %d 次): %w", maxRetries, lastErr) } func (c *Client) reconnect() error { // 用 buildClient 而非 newClient,避免保活循环绑定到即将丢弃的临时容器 nc, err := buildClient(c.config) if err != nil { return err } c.mu.Lock() defer c.mu.Unlock() c.closeLocked() c.client = nc.client c.sshClient = nc.sshClient c.closed = false c.startKeepalive() return nil } func (c *Client) Close() { c.mu.Lock() defer c.mu.Unlock() c.closeLocked() } func (c *Client) closeLocked() { if c.stopKeep != nil { close(c.stopKeep) c.stopKeep = nil } if c.client != nil { c.client.Close() c.client = nil } if c.sshClient != nil { c.sshClient.Close() c.sshClient = nil } c.closed = true } // RunCommand 通过 SSH Session 执行远程命令,返回 stdout func (c *Client) RunCommand(cmd string) (string, error) { c.mu.Lock() sshClient := c.sshClient c.mu.Unlock() if sshClient == nil { return "", fmt.Errorf("SSH 客户端未初始化") } session, err := sshClient.NewSession() if err != nil { return "", fmt.Errorf("创建 SSH 会话失败: %w", err) } defer session.Close() out, err := session.CombinedOutput(cmd) if err != nil { return "", fmt.Errorf("执行命令失败 [%s]: %w", cmd, err) } return string(out), nil } // RawClient 获取底层 sftp.Client(高级用法) func (c *Client) RawClient() *sftp.Client { c.mu.Lock() defer c.mu.Unlock() return c.client }