diff --git a/internal/desktopruntime/tunnel_config_windows.go b/internal/desktopruntime/tunnel_config_windows.go index d412d4b0..015a5d8e 100644 --- a/internal/desktopruntime/tunnel_config_windows.go +++ b/internal/desktopruntime/tunnel_config_windows.go @@ -18,8 +18,66 @@ func platformConfigureTunnel(ctx context.Context, request TunnelConfigureRequest if err := ensureDesktopCredentials(runtime.root); err != nil { return err } + + // A configure request updates several files plus the startup entry before the + // replacement tunnel is proven ready. Keep the previously committed state so + // any later failure can restore it as one transaction. + snapshotPaths := []string{ + runtime.files.manifest, + runtime.files.mode, + runtime.files.serverURL, + runtime.files.namedServerURL, + runtime.files.quickURL, + runtime.files.token, + } + snapshots := make([]fileSnapshot, 0, len(snapshotPaths)) + for _, path := range snapshotPaths { + snapshot, snapshotErr := snapshotFile(path) + if snapshotErr != nil { + return fmt.Errorf("备份 Tunnel 配置失败: %w", snapshotErr) + } + snapshots = append(snapshots, snapshot) + } + oldAutostart, err := tunnelAutostartEnabled(runtime.manifest) + if err != nil { + return fmt.Errorf("读取 Tunnel 开机启动状态失败: %w", err) + } + + tunnelStopped := false + rollback := func(cause error) error { + var restoreErr error + if tunnelStopped { + if err := stopTunnel(ctx, runtime); err != nil { + restoreErr = errors.Join(restoreErr, err) + } + } + restoreErr = errors.Join(restoreErr, restoreSnapshots(snapshots)) + if tunnelStopped { + if err := platformSetTunnelAutostart(ctx, runtime.root, oldAutostart); err != nil { + restoreErr = errors.Join(restoreErr, err) + } + oldRuntime, loadErr := loadTunnelRuntime(request.RuntimeRoot) + if loadErr != nil { + restoreErr = errors.Join(restoreErr, loadErr) + } else { + if err := platformServiceAction(ctx, oldRuntime.root, "restart"); err != nil { + restoreErr = errors.Join(restoreErr, err) + } + if oldRuntime.mode != "none" { + if err := startTunnel(ctx, oldRuntime); err != nil { + restoreErr = errors.Join(restoreErr, err) + } + } + } + } + if restoreErr != nil { + return fmt.Errorf("%w;同时恢复 Tunnel 配置失败: %v", cause, restoreErr) + } + return cause + } + if err := preserveNamedServerURL(runtime); err != nil { - return err + return rollback(err) } namedServerURL := "" @@ -28,96 +86,106 @@ func platformConfigureTunnel(ctx context.Context, request TunnelConfigureRequest if candidate == "" { candidate, err = readTrimmedText(runtime.files.namedServerURL) if err != nil { - return err + return rollback(err) } } namedServerURL, err = normalizeHTTPSOrigin(candidate) if err != nil { - return err + return rollback(err) } providedToken, err := readSecretFile(request.TokenFile) if err != nil { - return err + return rollback(err) } if providedToken != "" { if err := writeProtectedText(runtime.files.token, providedToken, tunnelTokenEntropy); err != nil { - return fmt.Errorf("保存 Cloudflare Tunnel Token 失败: %w", err) + return rollback(fmt.Errorf("保存 Cloudflare Tunnel Token 失败: %w", err)) } } storedToken, err := readProtectedText(runtime.files.token, tunnelTokenEntropy) if err != nil { if errors.Is(err, os.ErrNotExist) { - return errors.New("固定域名模式需要 Cloudflare Tunnel Token") + return rollback(errors.New("固定域名模式需要 Cloudflare Tunnel Token")) } - return fmt.Errorf("读取 Cloudflare Tunnel Token 失败: %w", err) + return rollback(fmt.Errorf("读取 Cloudflare Tunnel Token 失败: %w", err)) } if strings.TrimSpace(storedToken) == "" { - return errors.New("固定域名模式需要 Cloudflare Tunnel Token") + return rollback(errors.New("固定域名模式需要 Cloudflare Tunnel Token")) } } if err := stopTunnel(ctx, runtime); err != nil { - return err + return rollback(err) } + tunnelStopped = true switch request.Mode { case "none": if err := writeRuntimeText(runtime.files.mode, "none"); err != nil { - return err + return rollback(err) } if err := clearActivePublicURL(runtime.files); err != nil { - return err + return rollback(err) } if err := runtime.updateManifest("none", ""); err != nil { - return err + return rollback(err) } if err := platformSetTunnelAutostart(ctx, runtime.root, false); err != nil { - return err + return rollback(err) + } + if err := platformServiceAction(ctx, runtime.root, "restart"); err != nil { + return rollback(err) } - return platformServiceAction(ctx, runtime.root, "restart") + return nil case "quick": if err := writeRuntimeText(runtime.files.mode, "quick"); err != nil { - return err + return rollback(err) } if err := clearActivePublicURL(runtime.files); err != nil { - return err + return rollback(err) } if err := runtime.updateManifest("none", ""); err != nil { - return err + return rollback(err) } if err := platformSetTunnelAutostart(ctx, runtime.root, true); err != nil { - return err + return rollback(err) } if err := platformServiceAction(ctx, runtime.root, "restart"); err != nil { - return err + return rollback(err) } runtime.mode = "quick" - return startTunnel(ctx, runtime) + if err := startTunnel(ctx, runtime); err != nil { + return rollback(err) + } + return nil case "named": if err := writeRuntimeText(runtime.files.namedServerURL, namedServerURL); err != nil { - return err + return rollback(err) } if err := writeRuntimeText(runtime.files.serverURL, namedServerURL); err != nil { - return err + return rollback(err) } if err := writeRuntimeText(runtime.files.mode, "named"); err != nil { - return err + return rollback(err) } if err := os.Remove(runtime.files.quickURL); err != nil && !errors.Is(err, os.ErrNotExist) { - return fmt.Errorf("删除 Quick Tunnel ready 文件失败: %w", err) + return rollback(fmt.Errorf("删除 Quick Tunnel ready 文件失败: %w", err)) } if err := runtime.updateManifest("named", namedServerURL); err != nil { - return err + return rollback(err) } if err := platformSetTunnelAutostart(ctx, runtime.root, true); err != nil { - return err + return rollback(err) } if err := platformServiceAction(ctx, runtime.root, "restart"); err != nil { - return err + return rollback(err) } runtime.mode = "named" - return startTunnel(ctx, runtime) + if err := startTunnel(ctx, runtime); err != nil { + return rollback(err) + } + return nil default: - return fmt.Errorf("不支持的公网模式:%s", request.Mode) + return rollback(fmt.Errorf("不支持的公网模式:%s", request.Mode)) } } diff --git a/internal/desktopruntime/tunnel_config_windows_test.go b/internal/desktopruntime/tunnel_config_windows_test.go new file mode 100644 index 00000000..afebec1e --- /dev/null +++ b/internal/desktopruntime/tunnel_config_windows_test.go @@ -0,0 +1,116 @@ +//go:build windows + +package desktopruntime + +import ( + "context" + "errors" + "net" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" +) + +func TestPlatformConfigureTunnelRollsBackFailedNamedStart(t *testing.T) { + root := t.TempDir() + newURL := "https://new.example.test" + oldToken := "old-stable-token" + newToken := "new-replacement-token" + + healthServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + })) + defer healthServer.Close() + addr := healthServer.Listener.Addr().(*net.TCPAddr) + + runValueName := "AgentDockTunnelRollbackTest-" + filepath.Base(root) + defer func() { _ = removeRunValue(runValueName) }() + + manifest := Manifest{ + SchemaVersion: SchemaVersion, + InstallRoot: root, + AgentDockBinary: filepath.Join(root, "bin", "agentdock.exe"), + CloudflaredBinary: filepath.Join(root, "bin", "missing-cloudflared.exe"), + CloudflaredStartupValueName: runValueName, + Host: addr.IP.String(), + Port: addr.Port, + LocalMCPURL: "http://127.0.0.1:8765/mcp", + TunnelMode: "none", + InstallChannel: "test", + PrivilegeMode: "standard", + } + manifestPath := filepath.Join(root, "runtime.json") + if err := Save(manifestPath, manifest); err != nil { + t.Fatal(err) + } + modePath := filepath.Join(root, "cloudflared-mode.txt") + if err := writeRuntimeText(modePath, "none"); err != nil { + t.Fatal(err) + } + serverURLPath := filepath.Join(root, "server-url.txt") + if err := writeRuntimeText(serverURLPath, ""); err != nil { + t.Fatal(err) + } + tokenPath := filepath.Join(root, "cloudflared-token.dpapi") + if err := writeProtectedText(tokenPath, oldToken, tunnelTokenEntropy); err != nil { + t.Fatal(err) + } + + tokenFile := filepath.Join(root, "replacement-token.txt") + if err := os.WriteFile(tokenFile, []byte(newToken), 0o600); err != nil { + t.Fatal(err) + } + + err := platformConfigureTunnel(context.Background(), TunnelConfigureRequest{ + RuntimeRoot: root, + Mode: "named", + ServerURL: newURL, + TokenFile: tokenFile, + }) + if err == nil || !strings.Contains(err.Error(), "找不到 cloudflared.exe") { + t.Fatalf("expected cloudflared start failure, got %v", err) + } + + storedToken, err := readProtectedText(tokenPath, tunnelTokenEntropy) + if err != nil { + t.Fatal(err) + } + if storedToken != oldToken { + t.Fatalf("stored token = %q, want previous token", storedToken) + } + namedURLPath := filepath.Join(root, "named-server-url.txt") + if _, err := os.Stat(namedURLPath); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("named URL was not rolled back: %v", err) + } + serverURL, err := readTrimmedText(serverURLPath) + if err != nil { + t.Fatal(err) + } + if serverURL != "" { + t.Fatalf("server URL = %q, want empty", serverURL) + } + mode, err := readTrimmedText(modePath) + if err != nil { + t.Fatal(err) + } + if mode != "none" { + t.Fatalf("mode = %q, want none", mode) + } + restoredManifest, err := Load(manifestPath) + if err != nil { + t.Fatal(err) + } + if restoredManifest.TunnelMode != "none" || restoredManifest.PublicURL != "" { + t.Fatalf("manifest was not rolled back: mode=%q public_url=%q", restoredManifest.TunnelMode, restoredManifest.PublicURL) + } + startupEnabled, err := tunnelAutostartEnabled(restoredManifest) + if err != nil { + t.Fatal(err) + } + if startupEnabled { + t.Fatal("tunnel autostart was not rolled back") + } +}