修复: SFTP连接首拉不显示与传输挂起
- watcher握手期伪触发致幻影本地列表,根因修复 - 连接探测2.5s超时竞速,握手10s预算 - 30s保活自动掐死半开连接,取连接时剔除已断开 - 副本命名与进度计数抽公共,下载缓存清理去重
This commit is contained in:
+100
-4
@@ -11,11 +11,22 @@ import (
|
||||
"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
|
||||
}
|
||||
|
||||
@@ -73,6 +84,11 @@ func (m *Manager) Disconnect(host string, port int) {
|
||||
}
|
||||
}
|
||||
|
||||
// Evict 从连接池剔除连接(底层已断开,仅移除不重复关闭)
|
||||
func (m *Manager) Evict(connID string) {
|
||||
m.clients.Delete(connID)
|
||||
}
|
||||
|
||||
// Shutdown 关闭所有连接
|
||||
func (m *Manager) Shutdown() {
|
||||
m.clients.Range(func(key, value any) bool {
|
||||
@@ -84,7 +100,18 @@ func (m *Manager) Shutdown() {
|
||||
|
||||
// --- 内部 ---
|
||||
|
||||
// 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{
|
||||
@@ -115,8 +142,8 @@ func newClient(config *Config) (*Client, error) {
|
||||
}
|
||||
sshConfig.Auth = []ssh.AuthMethod{ssh.PublicKeys(signer)}
|
||||
} else if config.Password != "" {
|
||||
pw := config.Password
|
||||
sshConfig.Auth = []ssh.AuthMethod{
|
||||
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))
|
||||
@@ -131,10 +158,13 @@ func newClient(config *Config) (*Client, error) {
|
||||
}
|
||||
|
||||
addr := fmt.Sprintf("%s:%d", config.Host, config.Port)
|
||||
sshConn, err := net.DialTimeout("tcp", addr, config.Timeout)
|
||||
// 拨号与认证、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 {
|
||||
@@ -143,7 +173,9 @@ func newClient(config *Config) (*Client, error) {
|
||||
}
|
||||
|
||||
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}
|
||||
@@ -156,6 +188,62 @@ func newClient(config *Config) (*Client, error) {
|
||||
}, 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()
|
||||
@@ -204,7 +292,8 @@ func (c *Client) WithRetry(fn func(*sftp.Client) error) error {
|
||||
}
|
||||
|
||||
func (c *Client) reconnect() error {
|
||||
nc, err := newClient(c.config)
|
||||
// 用 buildClient 而非 newClient,避免保活循环绑定到即将丢弃的临时容器
|
||||
nc, err := buildClient(c.config)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -213,6 +302,8 @@ func (c *Client) reconnect() error {
|
||||
c.closeLocked()
|
||||
c.client = nc.client
|
||||
c.sshClient = nc.sshClient
|
||||
c.closed = false
|
||||
c.startKeepalive()
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -223,6 +314,10 @@ func (c *Client) Close() {
|
||||
}
|
||||
|
||||
func (c *Client) closeLocked() {
|
||||
if c.stopKeep != nil {
|
||||
close(c.stopKeep)
|
||||
c.stopKeep = nil
|
||||
}
|
||||
if c.client != nil {
|
||||
c.client.Close()
|
||||
c.client = nil
|
||||
@@ -231,6 +326,7 @@ func (c *Client) closeLocked() {
|
||||
c.sshClient.Close()
|
||||
c.sshClient = nil
|
||||
}
|
||||
c.closed = true
|
||||
}
|
||||
|
||||
// RunCommand 通过 SSH Session 执行远程命令,返回 stdout
|
||||
|
||||
+285
-16
@@ -1,8 +1,12 @@
|
||||
package sftp
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"log/slog"
|
||||
"io"
|
||||
"io/fs"
|
||||
"os"
|
||||
@@ -12,6 +16,7 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"u-desk/internal/common"
|
||||
"u-desk/internal/filesystem"
|
||||
"u-desk/internal/storage"
|
||||
|
||||
@@ -84,7 +89,7 @@ func (s *Service) ReadFile(connID string, filePath string) (string, error) {
|
||||
fi, e := sc.Stat(filePath)
|
||||
if e != nil { return e }
|
||||
if fi.Size() > maxSize {
|
||||
return fmt.Errorf("文件过大 (%s),超过 %d 限制", filesystem.FormatBytes(fi.Size()), maxSize)
|
||||
return fmt.Errorf("文件过大 (%s),超过 %d 限制", common.FormatBytesInt(fi.Size()), maxSize)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
@@ -105,7 +110,7 @@ func (s *Service) ReadFile(connID string, filePath string) (string, error) {
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("读取文件失败: %w", err)
|
||||
}
|
||||
return filesystem.BytesToString(data), nil
|
||||
return common.BytesToString(data), nil
|
||||
}
|
||||
|
||||
func (s *Service) WriteFile(connID string, filePath string, content string) error {
|
||||
@@ -168,7 +173,7 @@ func (s *Service) GetFileInfo(connID string, filePath string) (map[string]interf
|
||||
"name": info.Name(),
|
||||
"path": toUnixPath(filePath),
|
||||
"size": info.Size(),
|
||||
"size_str": filesystem.FormatBytes(info.Size()),
|
||||
"size_str": common.FormatBytesInt(info.Size()),
|
||||
"is_dir": info.IsDir(),
|
||||
"mod_time": info.ModTime().Format("2006-01-02 15:04:05"),
|
||||
"mode": info.Mode().String(),
|
||||
@@ -209,7 +214,10 @@ func (s *Service) CreateFile(connID string, filePath string) (*filesystem.FileOp
|
||||
return nil, fmt.Errorf("创建文件失败: %w", err)
|
||||
}
|
||||
|
||||
infoMap, _ := s.GetFileInfo(connID, filePath)
|
||||
infoMap, err := s.GetFileInfo(connID, filePath)
|
||||
if err != nil {
|
||||
slog.Warn("创建后获取详情失败", "path", filePath, "error", err)
|
||||
}
|
||||
return toFileOperationResult(infoMap, false), nil
|
||||
}
|
||||
|
||||
@@ -291,7 +299,7 @@ func (s *Service) downloadToTempDirect(connID string, remotePath string) (string
|
||||
fi, e := sc.Stat(remotePath)
|
||||
if e != nil { return e }
|
||||
if fi.Size() > maxPreviewSize {
|
||||
return fmt.Errorf("预览文件过大: %s", filesystem.FormatBytes(fi.Size()))
|
||||
return fmt.Errorf("预览文件过大: %s", common.FormatBytesInt(fi.Size()))
|
||||
}
|
||||
return nil
|
||||
})
|
||||
@@ -345,11 +353,20 @@ func (s *Service) DownloadSiteForPreview(connID string, remotePath string) (stri
|
||||
return "", err
|
||||
}
|
||||
|
||||
// 1. 创建临时目录
|
||||
tmpDir, err := os.MkdirTemp("", "udesk-sftp-site-*")
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("创建临时目录失败: %w", err)
|
||||
// 1. 缓存检查:用确定性路径,命中直接返回
|
||||
htmlInfo, _ := s.GetFileInfo(connID, remotePath)
|
||||
htmlSize, _ := htmlInfo["size"].(int64)
|
||||
htmlModTime, _ := htmlInfo["mod_time"].(string)
|
||||
cacheDir := common.SiteCacheDir("sftp", connID, remotePath, htmlSize, htmlModTime)
|
||||
htmlCachePath := filepath.Join(cacheDir, filepath.FromSlash(path.Dir(remotePath)), path.Base(remotePath))
|
||||
if _, err := os.Stat(htmlCachePath); err == nil {
|
||||
return htmlCachePath, nil // 缓存命中
|
||||
}
|
||||
os.RemoveAll(cacheDir) // 清理旧缓存
|
||||
if err := os.MkdirAll(cacheDir, 0755); err != nil {
|
||||
return "", fmt.Errorf("创建缓存目录失败: %w", err)
|
||||
}
|
||||
tmpDir := cacheDir
|
||||
|
||||
// 2. 确定远程网站根目录(从 HTML 路径推断)
|
||||
keyDir := path.Dir(remotePath)
|
||||
@@ -379,7 +396,7 @@ func (s *Service) DownloadSiteForPreview(connID string, remotePath string) (stri
|
||||
if err != nil {
|
||||
return htmlLocalPath, nil
|
||||
}
|
||||
resources := filesystem.ExtractHtmlResources(string(htmlContent))
|
||||
resources := common.ExtractHtmlResources(string(htmlContent))
|
||||
|
||||
// 5. 下载静态引用资源(嗅探网站根)
|
||||
htmlRemoteDir := keyDir
|
||||
@@ -399,8 +416,11 @@ func (s *Service) DownloadSiteForPreview(connID string, remotePath string) (stri
|
||||
}
|
||||
}
|
||||
|
||||
// 绝对路径资源:串行嗅探 siteRoot,相对路径:并行下载
|
||||
type relativeTask struct{ remoteKey, localPath string }
|
||||
var relativeTasks []relativeTask
|
||||
for _, resPath := range resources {
|
||||
if filesystem.ShouldSkipResource(resPath) {
|
||||
if common.ShouldSkipResource(resPath) {
|
||||
continue
|
||||
}
|
||||
isAbsolute := strings.HasPrefix(resPath, "/")
|
||||
@@ -418,12 +438,20 @@ func (s *Service) DownloadSiteForPreview(connID string, remotePath string) (stri
|
||||
}
|
||||
} else {
|
||||
remoteKey := path.Join(htmlRemoteDir, cleanPath)
|
||||
if s.sftpTryDownload(c, remoteKey, localPath) {
|
||||
recordDir(remoteKey)
|
||||
}
|
||||
relativeTasks = append(relativeTasks, relativeTask{remoteKey, localPath})
|
||||
}
|
||||
}
|
||||
|
||||
// 相对路径资源并行下载(recordDir 需加锁,多 goroutine 并发回调)
|
||||
var mu sync.Mutex
|
||||
common.RunConcurrent(relativeTasks, common.SiteDownloadConcurrency, func(t relativeTask) {
|
||||
if s.sftpTryDownload(c, t.remoteKey, t.localPath) {
|
||||
mu.Lock()
|
||||
recordDir(t.remoteKey)
|
||||
mu.Unlock()
|
||||
}
|
||||
})
|
||||
|
||||
// 6. 补充下载已发现目录中的剩余文件(覆盖动态 chunk 等)
|
||||
for _, dir := range discoveredDirs {
|
||||
sftpSupplementDir(s, c, dir, tmpDir, siteRoot)
|
||||
@@ -494,6 +522,7 @@ func sftpResolveAndDownload(s *Service, c *Client, htmlDir string, cleanPath str
|
||||
}
|
||||
|
||||
// sftpSupplementDir 补充下载远程目录中尚未下载的文件(只处理已知资源所在目录)
|
||||
// 数量与单文件大小上限见 common.MaxSupplementFiles / common.MaxSupplementFileSize
|
||||
func sftpSupplementDir(s *Service, c *Client, remoteDir string, tmpDir string, siteRoot string) {
|
||||
var entries []fs.FileInfo
|
||||
err := c.WithRetry(func(sc *sftpclient.Client) error {
|
||||
@@ -504,19 +533,254 @@ func sftpSupplementDir(s *Service, c *Client, remoteDir string, tmpDir string, s
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
count := 0
|
||||
for _, entry := range entries {
|
||||
if entry.IsDir() || entry.Size() == 0 {
|
||||
if entry.IsDir() || entry.Size() == 0 || entry.Size() > common.MaxSupplementFileSize {
|
||||
continue
|
||||
}
|
||||
if count >= common.MaxSupplementFiles {
|
||||
log.Printf("[站点下载] 补充扫描达到上限: count=%d", common.MaxSupplementFiles)
|
||||
break
|
||||
}
|
||||
fullPath := path.Join(remoteDir, entry.Name())
|
||||
localPath := filepath.Join(tmpDir, filepath.FromSlash(fullPath))
|
||||
if _, err := os.Stat(localPath); err == nil {
|
||||
continue
|
||||
}
|
||||
s.sftpTryDownload(c, fullPath, localPath)
|
||||
count++
|
||||
}
|
||||
}
|
||||
|
||||
// DownloadToFile 下载远程文件/目录到用户指定的本地目录(流式,无大小上限),返回实际落地路径
|
||||
// 目标重名自动追加「 - 副本」后缀;onProgress 按累计字节回调(name 为当前文件名)
|
||||
func (s *Service) DownloadToFile(connID, remotePath, localDstDir string, onProgress func(fileName string, copied, total int64)) (string, error) {
|
||||
c, err := s.getClient(connID)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
var info fs.FileInfo
|
||||
err = c.WithRetry(func(sc *sftpclient.Client) error {
|
||||
var e error
|
||||
info, e = sc.Stat(remotePath)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("获取远程文件信息失败: %w", err)
|
||||
}
|
||||
|
||||
// 单文件:直接流式落地
|
||||
if !info.IsDir() {
|
||||
dstPath := filesystem.UniquePath(filepath.Join(localDstDir, path.Base(remotePath)))
|
||||
var copied int64
|
||||
if err := s.downloadFileTo(c, remotePath, dstPath, &copied, info.Size(), onProgress); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return dstPath, nil
|
||||
}
|
||||
|
||||
// 目录:先收集远端文件清单并统计总大小
|
||||
type remoteEntry struct {
|
||||
remote, local string
|
||||
size int64
|
||||
}
|
||||
var entries []remoteEntry
|
||||
var total int64
|
||||
var walk func(dir, localDir string) error
|
||||
walk = func(dir, localDir string) error {
|
||||
var list []fs.FileInfo
|
||||
err := c.WithRetry(func(sc *sftpclient.Client) error {
|
||||
var e error
|
||||
list, e = sc.ReadDir(dir)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, fi := range list {
|
||||
rp := path.Join(dir, fi.Name())
|
||||
lp := filepath.Join(localDir, fi.Name())
|
||||
if fi.IsDir() {
|
||||
if err := walk(rp, lp); err != nil {
|
||||
return err
|
||||
}
|
||||
} else if fi.Mode().IsRegular() {
|
||||
entries = append(entries, remoteEntry{rp, lp, fi.Size()})
|
||||
total += fi.Size()
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
dstDir := filesystem.UniquePath(filepath.Join(localDstDir, path.Base(remotePath)))
|
||||
if err := walk(remotePath, dstDir); err != nil {
|
||||
return "", fmt.Errorf("遍历远程目录失败: %w", err)
|
||||
}
|
||||
|
||||
var copied int64
|
||||
for _, e := range entries {
|
||||
if err := s.downloadFileTo(c, e.remote, e.local, &copied, total, onProgress); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
return dstDir, nil
|
||||
}
|
||||
|
||||
// downloadFileTo 流式下载单个远程文件,进度按累计字节回调
|
||||
func (s *Service) downloadFileTo(c *Client, remotePath, localPath string, copied *int64, total int64, onProgress func(fileName string, copied, total int64)) error {
|
||||
if err := os.MkdirAll(filepath.Dir(localPath), 0755); err != nil {
|
||||
return fmt.Errorf("创建本地目录失败: %w", err)
|
||||
}
|
||||
return c.WithRetry(func(sc *sftpclient.Client) error {
|
||||
src, err := sc.Open(remotePath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer src.Close()
|
||||
dst, err := os.Create(localPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer dst.Close()
|
||||
var reader io.Reader = src
|
||||
if onProgress != nil {
|
||||
reader = &common.CountingReader{R: src, OnN: func(n int64) {
|
||||
*copied += n
|
||||
onProgress(path.Base(remotePath), *copied, total)
|
||||
}}
|
||||
}
|
||||
_, err = io.Copy(dst, reader)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
// UploadFromFile 上传本地文件/目录到远程目录(流式),远端重名自动追加副本后缀
|
||||
func (s *Service) UploadFromFile(connID, localPath, remoteDir string, onProgress func(fileName string, copied, total int64)) error {
|
||||
c, err := s.getClient(connID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
info, err := os.Stat(localPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("本地路径不存在: %w", err)
|
||||
}
|
||||
|
||||
// 单文件:直接流式上传
|
||||
if !info.IsDir() {
|
||||
remotePath, err := s.uniqueRemotePath(c, path.Join(remoteDir, filepath.Base(localPath)))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var copied int64
|
||||
return s.uploadFileTo(c, localPath, remotePath, &copied, info.Size(), onProgress)
|
||||
}
|
||||
|
||||
// 目录:本地清单 + 总大小(相对根路径,远端落地根确定后再拼完整路径)
|
||||
type localEntry struct {
|
||||
local, rel string
|
||||
size int64
|
||||
}
|
||||
var files []localEntry
|
||||
var total int64
|
||||
err = filepath.WalkDir(localPath, func(p string, d fs.DirEntry, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !d.Type().IsRegular() {
|
||||
return nil
|
||||
}
|
||||
fi, e := d.Info()
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
rel, e := filepath.Rel(localPath, p)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
files = append(files, localEntry{p, filepath.ToSlash(rel), fi.Size()})
|
||||
total += fi.Size()
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("遍历本地目录失败: %w", err)
|
||||
}
|
||||
|
||||
remoteRoot, err := s.uniqueRemotePath(c, path.Join(remoteDir, filepath.Base(localPath)))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var copied int64
|
||||
for _, f := range files {
|
||||
if err := s.uploadFileTo(c, f.local, path.Join(remoteRoot, f.rel), &copied, total, onProgress); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// uploadFileTo 流式上传单个本地文件(自动创建远端父目录),进度按累计字节回调
|
||||
func (s *Service) uploadFileTo(c *Client, localPath, remotePath string, copied *int64, total int64, onProgress func(fileName string, copied, total int64)) error {
|
||||
return c.WithRetry(func(sc *sftpclient.Client) error {
|
||||
if err := sc.MkdirAll(path.Dir(remotePath)); err != nil {
|
||||
return err
|
||||
}
|
||||
src, err := os.Open(localPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer src.Close()
|
||||
dst, err := sc.Create(remotePath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer dst.Close()
|
||||
var reader io.Reader = src
|
||||
if onProgress != nil {
|
||||
reader = &common.CountingReader{R: src, OnN: func(n int64) {
|
||||
*copied += n
|
||||
onProgress(filepath.Base(localPath), *copied, total)
|
||||
}}
|
||||
}
|
||||
_, err = io.Copy(dst, reader)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
// uniqueRemotePath 远端重名时追加「 - 副本」「 - 副本 (2)」后缀(扩展名前插入)
|
||||
// 命名循环复用 common.UniqueName(与本地 UniquePath 同源)
|
||||
func (s *Service) uniqueRemotePath(c *Client, remotePath string) (string, error) {
|
||||
dir := path.Dir(remotePath)
|
||||
unique, err := common.UniqueName(path.Base(remotePath), func(candidate string) (bool, error) {
|
||||
return s.remoteExists(c, path.Join(dir, candidate))
|
||||
})
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return path.Join(dir, unique), nil
|
||||
}
|
||||
|
||||
// remoteExists 判断远程路径是否存在(区分不存在与网络错误)
|
||||
func (s *Service) remoteExists(c *Client, p string) (bool, error) {
|
||||
var notExist bool
|
||||
err := c.WithRetry(func(sc *sftpclient.Client) error {
|
||||
_, e := sc.Stat(p)
|
||||
if e == nil {
|
||||
return nil
|
||||
}
|
||||
if errors.Is(e, fs.ErrNotExist) {
|
||||
notExist = true
|
||||
return nil
|
||||
}
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return !notExist, nil
|
||||
}
|
||||
|
||||
// GetCommonPaths 返回 SFTP 远程主机常用路径
|
||||
func (s *Service) GetCommonPaths(connID string) (map[string]string, error) {
|
||||
c := s.manager.GetClient(connID)
|
||||
@@ -676,6 +940,11 @@ func (s *Service) getClient(connID string) (*Client, error) {
|
||||
if c == nil {
|
||||
return nil, fmt.Errorf("SFTP 连接不存在: %s", connID)
|
||||
}
|
||||
if c.IsClosed() {
|
||||
// 保活失败或显式断开置位,剔除死连接避免后续操作继续挂在上面
|
||||
s.manager.Evict(connID)
|
||||
return nil, fmt.Errorf("SFTP 连接已断开,请重新连接: %s", connID)
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
|
||||
@@ -694,7 +963,7 @@ func toFileOperationResult(m map[string]interface{}, isDir bool) *filesystem.Fil
|
||||
Path: p,
|
||||
Name: name,
|
||||
Size: size,
|
||||
SizeStr: filesystem.FormatBytes(size),
|
||||
SizeStr: common.FormatBytesInt(size),
|
||||
IsDir: isDir,
|
||||
ModTime: modTime,
|
||||
Mode: mode,
|
||||
|
||||
Reference in New Issue
Block a user