diff --git a/atsf_server/service/update.go b/atsf_server/service/update.go index 666870da..81d7c574 100644 --- a/atsf_server/service/update.go +++ b/atsf_server/service/update.go @@ -161,7 +161,15 @@ func UploadManualServerBinary(ctx context.Context, fileName string, reader io.Re return nil, fmt.Errorf("缺少上传文件内容") } - tempPath, err := persistUploadedServerBinary(fileName, reader) + execPath, err := os.Executable() + if err != nil { + return nil, fmt.Errorf("获取当前服务程序路径失败: %v", err) + } + if err = verifyExecutableDirectoryWritable(execPath); err != nil { + return nil, err + } + + tempPath, err := persistUploadedServerBinary(filepath.Dir(execPath), fileName, reader) if err != nil { return nil, err } @@ -180,16 +188,6 @@ func UploadManualServerBinary(ctx context.Context, fileName string, reader io.Re return info, nil } - execPath, err := os.Executable() - if err != nil { - _ = os.Remove(tempPath) - return nil, fmt.Errorf("获取当前服务程序路径失败: %v", err) - } - if err = verifyExecutableDirectoryWritable(execPath); err != nil { - _ = os.Remove(tempPath) - return nil, err - } - uploadToken, err := newUpgradeToken() if err != nil { _ = os.Remove(tempPath) @@ -712,12 +710,16 @@ func isManualServerUpgradeSupported(currentVersion string) bool { return normalized != "" && !strings.EqualFold(normalized, "dev") } -func persistUploadedServerBinary(fileName string, reader io.Reader) (string, error) { +func persistUploadedServerBinary(tempDir string, fileName string, reader io.Reader) (string, error) { suffix := filepath.Ext(strings.TrimSpace(fileName)) if runtime.GOOS == "windows" && suffix == "" { suffix = ".exe" } - tempFile, err := os.CreateTemp("", "atsflare-server-manual-upgrade-*"+suffix) + tempDir = strings.TrimSpace(tempDir) + if tempDir == "" { + tempDir = os.TempDir() + } + tempFile, err := os.CreateTemp(tempDir, "atsflare-server-manual-upgrade-*"+suffix) if err != nil { return "", fmt.Errorf("创建临时升级文件失败: %v", err) } diff --git a/atsf_server/service/update_restart_unix.go b/atsf_server/service/update_restart_unix.go index 96093a6a..20dd4a21 100644 --- a/atsf_server/service/update_restart_unix.go +++ b/atsf_server/service/update_restart_unix.go @@ -4,19 +4,22 @@ package service import ( "fmt" + "io" "os" "syscall" ) +var unixRename = os.Rename + func replaceAndRestartServer(execPath string, tmpPath string) error { backupPath := execPath + ".bak" _ = os.Remove(backupPath) - if err := os.Rename(execPath, backupPath); err != nil { + if err := unixRename(execPath, backupPath); err != nil { _ = os.Remove(tmpPath) return fmt.Errorf("备份当前服务端二进制失败: %w", err) } - if err := os.Rename(tmpPath, execPath); err != nil { - _ = os.Rename(backupPath, execPath) + if err := replaceFileUnix(tmpPath, execPath); err != nil { + _ = unixRename(backupPath, execPath) return fmt.Errorf("替换服务端二进制失败: %w", err) } _ = os.Remove(backupPath) @@ -25,3 +28,46 @@ func replaceAndRestartServer(execPath string, tmpPath string) error { } return fmt.Errorf("unreachable after exec") } + +func replaceFileUnix(srcPath string, dstPath string) error { + if err := unixRename(srcPath, dstPath); err == nil { + return nil + } else if linkErr, ok := err.(*os.LinkError); !ok || linkErr.Err != syscall.EXDEV { + return err + } + + sourceFile, err := os.Open(srcPath) + if err != nil { + return err + } + defer sourceFile.Close() + + info, err := sourceFile.Stat() + if err != nil { + return err + } + + destinationFile, err := os.OpenFile(dstPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, info.Mode().Perm()) + if err != nil { + return err + } + + copyErr := func() error { + defer destinationFile.Close() + if _, err = io.Copy(destinationFile, sourceFile); err != nil { + return err + } + if err = destinationFile.Sync(); err != nil { + return err + } + return nil + }() + if copyErr != nil { + return copyErr + } + + if err = os.Chmod(dstPath, info.Mode().Perm()); err != nil { + return err + } + return os.Remove(srcPath) +} diff --git a/atsf_server/service/update_restart_unix_test.go b/atsf_server/service/update_restart_unix_test.go new file mode 100644 index 00000000..b71cc268 --- /dev/null +++ b/atsf_server/service/update_restart_unix_test.go @@ -0,0 +1,50 @@ +//go:build !windows + +package service + +import ( + "errors" + "os" + "path/filepath" + "syscall" + "testing" +) + +func TestReplaceFileUnixFallsBackOnCrossDeviceRename(t *testing.T) { + tempDir := t.TempDir() + srcPath := filepath.Join(tempDir, "source.bin") + dstPath := filepath.Join(tempDir, "target.bin") + + if err := os.WriteFile(srcPath, []byte("new-binary"), 0o755); err != nil { + t.Fatalf("failed to write source file: %v", err) + } + if err := os.WriteFile(dstPath, []byte("old-binary"), 0o755); err != nil { + t.Fatalf("failed to write target file: %v", err) + } + + originalRename := unixRename + unixRename = func(oldPath string, newPath string) error { + if oldPath == srcPath && newPath == dstPath { + return &os.LinkError{Op: "rename", Old: oldPath, New: newPath, Err: syscall.EXDEV} + } + return os.Rename(oldPath, newPath) + } + t.Cleanup(func() { + unixRename = originalRename + }) + + if err := replaceFileUnix(srcPath, dstPath); err != nil { + t.Fatalf("expected cross-device fallback to succeed: %v", err) + } + + content, err := os.ReadFile(dstPath) + if err != nil { + t.Fatalf("failed to read target file: %v", err) + } + if string(content) != "new-binary" { + t.Fatalf("unexpected target content: %s", string(content)) + } + if _, err = os.Stat(srcPath); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("expected source file to be removed, got err=%v", err) + } +} diff --git a/atsf_server/service/update_test.go b/atsf_server/service/update_test.go index 5f473ed9..340d4866 100644 --- a/atsf_server/service/update_test.go +++ b/atsf_server/service/update_test.go @@ -5,6 +5,7 @@ import ( "bytes" "context" "os" + "path/filepath" "runtime" "testing" "time" @@ -174,6 +175,13 @@ func TestUploadManualServerBinary(t *testing.T) { if candidate.UploadToken != info.UploadToken { t.Fatalf("unexpected stored upload token: %s", candidate.UploadToken) } + execPath, err := os.Executable() + if err != nil { + t.Fatalf("failed to get executable path: %v", err) + } + if filepath.Dir(candidate.TempPath) != filepath.Dir(execPath) { + t.Fatalf("expected temporary binary in executable dir, got %s want %s", filepath.Dir(candidate.TempPath), filepath.Dir(execPath)) + } } func TestUploadManualServerBinaryRejectsSameVersion(t *testing.T) { diff --git a/atsf_server/web/components/layout/dashboard-topbar.tsx b/atsf_server/web/components/layout/dashboard-topbar.tsx index 27cdb9b1..7b65dcc7 100644 --- a/atsf_server/web/components/layout/dashboard-topbar.tsx +++ b/atsf_server/web/components/layout/dashboard-topbar.tsx @@ -47,6 +47,7 @@ export function DashboardTopbar() { useState(null); const menuRef = useRef(null); const isRoot = (user?.role ?? 0) >= 100; + const upgradeStatusPollInterval = 3000; const publicStatusQuery = useQuery({ queryKey: ['public-status'], @@ -57,13 +58,26 @@ export function DashboardTopbar() { queryKey: ['update', 'latest-release', 'stable'], queryFn: () => getLatestRelease('stable'), enabled: isRoot, - refetchInterval: 60 * 60 * 1000, + refetchInterval: (query) => { + const release = query.state.data; + if (isVersionModalOpen && release?.in_progress) { + return upgradeStatusPollInterval; + } + return 60 * 60 * 1000; + }, }); const previewReleaseQuery = useQuery({ queryKey: ['update', 'latest-release', 'preview'], queryFn: () => getLatestRelease('preview'), enabled: false, + refetchInterval: (query) => { + const release = query.state.data; + if (isVersionModalOpen && release?.in_progress) { + return upgradeStatusPollInterval; + } + return false; + }, }); const upgradeMutation = useMutation({