Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
126 changes: 97 additions & 29 deletions internal/desktopruntime/tunnel_config_windows.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 := ""
Expand All @@ -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))
}
}
116 changes: 116 additions & 0 deletions internal/desktopruntime/tunnel_config_windows_test.go
Original file line number Diff line number Diff line change
@@ -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")
}
}