Files
u-desk/internal/sftp/client.go
T
lxy 6ea9d9ac99 修复: SFTP连接首拉不显示与传输挂起
- watcher握手期伪触发致幻影本地列表,根因修复
- 连接探测2.5s超时竞速,握手10s预算
- 30s保活自动掐死半开连接,取连接时剔除已断开
- 副本命名与进度计数抽公共,下载缓存清理去重
2026-09-15 22:26:22 +08:00

361 lines
8.6 KiB
Go

package sftp
import (
"fmt"
"net"
"os"
"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("必须提供密码或密钥文件")}
}
addr := fmt.Sprintf("%s:%d", config.Host, 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
}