diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index de88cd15..6eb8c295 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -1,7 +1,7 @@ -name: Release -permissions: - contents: write - +name: Release +permissions: + contents: write + on: workflow_dispatch: inputs: @@ -11,20 +11,20 @@ on: type: string push: tags: ["v*"] - -jobs: - prepare: - runs-on: ubuntu-latest - outputs: - should_run: ${{ steps.version.outputs.should_run }} - version: ${{ steps.version.outputs.version }} - is_prerelease: ${{ steps.version.outputs.is_prerelease }} - steps: - - name: Checkout - uses: actions/checkout@v4 - with: - fetch-depth: 0 - + +jobs: + prepare: + runs-on: ubuntu-latest + outputs: + should_run: ${{ steps.version.outputs.should_run }} + version: ${{ steps.version.outputs.version }} + is_prerelease: ${{ steps.version.outputs.is_prerelease }} + steps: + - name: Checkout + uses: actions/checkout@v4 + with: + fetch-depth: 0 + - name: Resolve version metadata id: version env: @@ -52,194 +52,320 @@ jobs: fi echo "should_run=$SHOULD_RUN" >> "$GITHUB_OUTPUT" - echo "version=$VERSION" >> "$GITHUB_OUTPUT" - if [[ "$VERSION" =~ ^v[0-9]+(\.[0-9]+)*$ ]]; then - echo "is_prerelease=false" >> "$GITHUB_OUTPUT" - else - echo "is_prerelease=true" >> "$GITHUB_OUTPUT" - fi - - build-frontend: - needs: prepare - if: needs.prepare.outputs.should_run == 'true' - runs-on: ubuntu-latest - steps: - - name: Checkout - uses: actions/checkout@v4 - with: - fetch-depth: 0 - - - name: Set up Node.js - uses: actions/setup-node@v4 - with: - node-version: 20 - - - name: Build Frontend - env: - CI: "" - VERSION: ${{ needs.prepare.outputs.version }} - run: | - cd openflare_server/web - corepack enable - pnpm install --frozen-lockfile - NEXT_PUBLIC_APP_VERSION="$VERSION" pnpm build - - - name: Upload Frontend Artifact - uses: actions/upload-artifact@v4 - with: - name: frontend-build - path: openflare_server/web/build - retention-days: 1 - - build-binaries: - needs: - - prepare - - build-frontend - if: needs.prepare.outputs.should_run == 'true' - runs-on: ubuntu-latest - strategy: - fail-fast: false - matrix: - include: - - goos: linux - goarch: amd64 - asset_name: openflare-server-linux-amd64 - - goos: linux - goarch: arm64 - asset_name: openflare-server-linux-arm64 - - goos: darwin - goarch: amd64 - asset_name: openflare-server-darwin-amd64 - - goos: darwin - goarch: arm64 - asset_name: openflare-server-darwin-arm64 - - goos: windows - goarch: amd64 - asset_name: openflare-server-windows-amd64.exe - - steps: - - name: Checkout - uses: actions/checkout@v4 - with: - fetch-depth: 0 - - - name: Download Frontend Artifact - uses: actions/download-artifact@v4 - with: - name: frontend-build - path: openflare_server/web/build - - - name: Set up Go - uses: actions/setup-go@v5 - with: - go-version-file: openflare_server/go.mod - - - name: Build Server - working-directory: openflare_server - env: - CGO_ENABLED: 0 - GOOS: ${{ matrix.goos }} - GOARCH: ${{ matrix.goarch }} - ASSET_NAME: ${{ matrix.asset_name }} - VERSION: ${{ needs.prepare.outputs.version }} - run: | - go mod download - mkdir -p ../dist - go build -trimpath -ldflags "-s -w -X 'openflare/common.Version=$VERSION'" -o "../dist/$ASSET_NAME" . - - - name: Upload Binary Artifact - uses: actions/upload-artifact@v4 - with: - name: server-${{ matrix.goos }}-${{ matrix.goarch }} - path: dist/${{ matrix.asset_name }} - retention-days: 1 - - build-agent-binaries: - needs: prepare - if: needs.prepare.outputs.should_run == 'true' - runs-on: ubuntu-latest - strategy: - fail-fast: false - matrix: - include: - - goos: linux - goarch: amd64 - asset_name: openflare-agent-linux-amd64 - - goos: linux - goarch: arm64 - asset_name: openflare-agent-linux-arm64 - - goos: darwin - goarch: amd64 - asset_name: openflare-agent-darwin-amd64 - - goos: darwin - goarch: arm64 - asset_name: openflare-agent-darwin-arm64 - - steps: - - name: Checkout - uses: actions/checkout@v4 - with: - fetch-depth: 0 - - - name: Set up Go - uses: actions/setup-go@v5 - with: - go-version-file: openflare_agent/go.mod - - - name: Build Agent - working-directory: openflare_agent - env: - CGO_ENABLED: 0 - GOOS: ${{ matrix.goos }} - GOARCH: ${{ matrix.goarch }} - ASSET_NAME: ${{ matrix.asset_name }} - VERSION: ${{ needs.prepare.outputs.version }} - run: | - go mod download - mkdir -p ../dist - go build -trimpath -ldflags "-s -w -X 'openflare-agent/internal/config.Version=$VERSION'" -o "../dist/$ASSET_NAME" ./cmd/agent + echo "version=$VERSION" >> "$GITHUB_OUTPUT" + if [[ "$VERSION" =~ ^v[0-9]+(\.[0-9]+)*$ ]]; then + echo "is_prerelease=false" >> "$GITHUB_OUTPUT" + else + echo "is_prerelease=true" >> "$GITHUB_OUTPUT" + fi + + build-frontend: + needs: prepare + if: needs.prepare.outputs.should_run == 'true' + runs-on: ubuntu-latest + steps: + - name: Checkout + uses: actions/checkout@v4 + with: + fetch-depth: 0 + + - name: Set up Node.js + uses: actions/setup-node@v4 + with: + node-version: 20 + + - name: Build Frontend + env: + CI: "" + VERSION: ${{ needs.prepare.outputs.version }} + run: | + cd openflare_server/web + corepack enable + pnpm install --frozen-lockfile + NEXT_PUBLIC_APP_VERSION="$VERSION" pnpm build + + - name: Upload Frontend Artifact + uses: actions/upload-artifact@v4 + with: + name: frontend-build + path: openflare_server/web/build + retention-days: 1 + + build-binaries: + needs: + - prepare + - build-frontend + if: needs.prepare.outputs.should_run == 'true' + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + include: + - goos: linux + goarch: amd64 + asset_name: openflare-server-linux-amd64 + - goos: linux + goarch: arm64 + asset_name: openflare-server-linux-arm64 + - goos: darwin + goarch: amd64 + asset_name: openflare-server-darwin-amd64 + - goos: darwin + goarch: arm64 + asset_name: openflare-server-darwin-arm64 + - goos: windows + goarch: amd64 + asset_name: openflare-server-windows-amd64.exe + + steps: + - name: Checkout + uses: actions/checkout@v4 + with: + fetch-depth: 0 + + - name: Download Frontend Artifact + uses: actions/download-artifact@v4 + with: + name: frontend-build + path: openflare_server/web/build + + - name: Set up Go + uses: actions/setup-go@v5 + with: + go-version-file: openflare_server/go.mod + + - name: Build Server + working-directory: openflare_server + env: + CGO_ENABLED: 0 + GOOS: ${{ matrix.goos }} + GOARCH: ${{ matrix.goarch }} + ASSET_NAME: ${{ matrix.asset_name }} + VERSION: ${{ needs.prepare.outputs.version }} + run: | + go mod download + mkdir -p ../dist + go build -trimpath -ldflags "-s -w -X 'openflare/common.Version=$VERSION'" -o "../dist/$ASSET_NAME" . + + - name: Upload Binary Artifact + uses: actions/upload-artifact@v4 + with: + name: server-${{ matrix.goos }}-${{ matrix.goarch }} + path: dist/${{ matrix.asset_name }} + retention-days: 1 + + build-agent-binaries: + needs: prepare + if: needs.prepare.outputs.should_run == 'true' + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + include: + - goos: linux + goarch: amd64 + asset_name: openflare-agent-linux-amd64 + - goos: linux + goarch: arm64 + asset_name: openflare-agent-linux-arm64 + - goos: darwin + goarch: amd64 + asset_name: openflare-agent-darwin-amd64 + - goos: darwin + goarch: arm64 + asset_name: openflare-agent-darwin-arm64 + + steps: + - name: Checkout + uses: actions/checkout@v4 + with: + fetch-depth: 0 + + - name: Set up Go + uses: actions/setup-go@v5 + with: + go-version-file: openflare_agent/go.mod + + - name: Build Agent + working-directory: openflare_agent + env: + CGO_ENABLED: 0 + GOOS: ${{ matrix.goos }} + GOARCH: ${{ matrix.goarch }} + ASSET_NAME: ${{ matrix.asset_name }} + VERSION: ${{ needs.prepare.outputs.version }} + run: | + go mod download + mkdir -p ../dist + go build -trimpath -ldflags "-s -w -X 'openflare-agent/internal/config.Version=$VERSION'" -o "../dist/$ASSET_NAME" ./cmd/agent (cd ../dist && sha256sum "$ASSET_NAME" > "$ASSET_NAME.sha256") - - - name: Upload Agent Artifact - uses: actions/upload-artifact@v4 - with: - name: agent-${{ matrix.goos }}-${{ matrix.goarch }} + + - name: Upload Agent Artifact + uses: actions/upload-artifact@v4 + with: + name: agent-${{ matrix.goos }}-${{ matrix.goarch }} path: | dist/${{ matrix.asset_name }} dist/${{ matrix.asset_name }}.sha256 - retention-days: 1 - - release: - needs: - - prepare - - build-binaries - - build-agent-binaries - if: needs.prepare.outputs.should_run == 'true' - runs-on: ubuntu-latest - steps: - - name: Download Server Artifacts - uses: actions/download-artifact@v4 - with: - pattern: "server-*" - path: dist - merge-multiple: true - - - name: Download Agent Artifacts - uses: actions/download-artifact@v4 - with: - pattern: "agent-*" - path: dist - merge-multiple: true - - - name: Release - uses: softprops/action-gh-release@v1 - with: - tag_name: ${{ needs.prepare.outputs.version }} - name: ${{ needs.prepare.outputs.version }} - target_commitish: ${{ github.sha }} - files: dist/* - draft: false - prerelease: ${{ needs.prepare.outputs.is_prerelease == 'true' }} - generate_release_notes: true - env: - GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} + retention-days: 1 + + build-relay-binaries: + needs: prepare + if: needs.prepare.outputs.should_run == 'true' + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + include: + - goos: linux + goarch: amd64 + asset_name: openflare-relay-linux-amd64 + - goos: linux + goarch: arm64 + asset_name: openflare-relay-linux-arm64 + - goos: darwin + goarch: amd64 + asset_name: openflare-relay-darwin-amd64 + - goos: darwin + goarch: arm64 + asset_name: openflare-relay-darwin-arm64 + + steps: + - name: Checkout + uses: actions/checkout@v4 + with: + fetch-depth: 0 + + - name: Set up Go + uses: actions/setup-go@v5 + with: + go-version-file: openflare_relay/go.mod + + - name: Build Relay + working-directory: openflare_relay + env: + CGO_ENABLED: 0 + GOOS: ${{ matrix.goos }} + GOARCH: ${{ matrix.goarch }} + ASSET_NAME: ${{ matrix.asset_name }} + VERSION: ${{ needs.prepare.outputs.version }} + run: | + go mod download + mkdir -p ../dist + go build -trimpath -ldflags "-s -w -X 'openflare-relay/internal/config.Version=$VERSION'" -o "../dist/$ASSET_NAME" ./cmd/relay + (cd ../dist && sha256sum "$ASSET_NAME" > "$ASSET_NAME.sha256") + + - name: Upload Relay Artifact + uses: actions/upload-artifact@v4 + with: + name: relay-${{ matrix.goos }}-${{ matrix.goarch }} + path: | + dist/${{ matrix.asset_name }} + dist/${{ matrix.asset_name }}.sha256 + retention-days: 1 + + build-flared-binaries: + needs: prepare + if: needs.prepare.outputs.should_run == 'true' + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + include: + - goos: linux + goarch: amd64 + asset_name: openflared-linux-amd64 + - goos: linux + goarch: arm64 + asset_name: openflared-linux-arm64 + - goos: darwin + goarch: amd64 + asset_name: openflared-darwin-amd64 + - goos: darwin + goarch: arm64 + asset_name: openflared-darwin-arm64 + + steps: + - name: Checkout + uses: actions/checkout@v4 + with: + fetch-depth: 0 + + - name: Set up Go + uses: actions/setup-go@v5 + with: + go-version-file: openflared/go.mod + + - name: Build Flared + working-directory: openflared + env: + CGO_ENABLED: 0 + GOOS: ${{ matrix.goos }} + GOARCH: ${{ matrix.goarch }} + ASSET_NAME: ${{ matrix.asset_name }} + VERSION: ${{ needs.prepare.outputs.version }} + run: | + go mod download + mkdir -p ../dist + go build -trimpath -ldflags "-s -w -X 'openflare-flared/internal/config.Version=$VERSION'" -o "../dist/$ASSET_NAME" ./cmd/flared + (cd ../dist && sha256sum "$ASSET_NAME" > "$ASSET_NAME.sha256") + + - name: Upload Flared Artifact + uses: actions/upload-artifact@v4 + with: + name: flared-${{ matrix.goos }}-${{ matrix.goarch }} + path: | + dist/${{ matrix.asset_name }} + dist/${{ matrix.asset_name }}.sha256 + retention-days: 1 + + release: + needs: + - prepare + - build-binaries + - build-agent-binaries + - build-relay-binaries + - build-flared-binaries + if: needs.prepare.outputs.should_run == 'true' + runs-on: ubuntu-latest + steps: + - name: Download Server Artifacts + uses: actions/download-artifact@v4 + with: + pattern: "server-*" + path: dist + merge-multiple: true + + - name: Download Agent Artifacts + uses: actions/download-artifact@v4 + with: + pattern: "agent-*" + path: dist + merge-multiple: true + + - name: Download Relay Artifacts + uses: actions/download-artifact@v4 + with: + pattern: "relay-*" + path: dist + merge-multiple: true + + - name: Download Flared Artifacts + uses: actions/download-artifact@v4 + with: + pattern: "flared-*" + path: dist + merge-multiple: true + + - name: Release + uses: softprops/action-gh-release@v1 + with: + tag_name: ${{ needs.prepare.outputs.version }} + name: ${{ needs.prepare.outputs.version }} + target_commitish: ${{ github.sha }} + files: dist/* + draft: false + prerelease: ${{ needs.prepare.outputs.is_prerelease == 'true' }} + generate_release_notes: true + env: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} diff --git a/openflare_relay/internal/heartbeat/service.go b/openflare_relay/internal/heartbeat/service.go index 3d93645e..1e4ceaf0 100644 --- a/openflare_relay/internal/heartbeat/service.go +++ b/openflare_relay/internal/heartbeat/service.go @@ -10,6 +10,7 @@ import ( "openflare-relay/internal/httpclient" "openflare-relay/internal/observability" "openflare-relay/internal/state" + "openflare-relay/internal/updater" "openflare/service" ) @@ -18,6 +19,7 @@ type Service struct { frpsManager *frps.Manager config *config.Config stateStore *state.Store + updater *updater.Service } func New(client *httpclient.Client, manager *frps.Manager, cfg *config.Config, stateStore *state.Store) *Service { @@ -26,6 +28,7 @@ func New(client *httpclient.Client, manager *frps.Manager, cfg *config.Config, s frpsManager: manager, config: cfg, stateStore: stateStore, + updater: updater.New(), } } @@ -72,4 +75,32 @@ func (s *Service) doHeartbeat(ctx context.Context) { // Update configs if changed s.frpsManager.UpdateConfig(resp.RelayConfig) + + if resp != nil && resp.RelaySettings != nil { + s.tryAutoUpdate(ctx, resp.RelaySettings) + } +} + +func (s *Service) tryAutoUpdate(ctx context.Context, settings *service.RelaySettings) { + if settings == nil || s.updater == nil { + return + } + force := settings.UpdateNow + shouldCheck := settings.AutoUpdate || force + if !shouldCheck || settings.UpdateRepo == "" { + return + } + channel := "stable" + if force && settings.UpdateChannel != "" { + channel = settings.UpdateChannel + } + slog.Info("checking for relay updates", "repo", settings.UpdateRepo, "channel", channel, "force", force) + err := s.updater.CheckAndUpdate(ctx, settings.UpdateRepo, updater.UpdateOptions{ + Channel: channel, + TagName: settings.UpdateTag, + Force: force, + }) + if err != nil { + slog.Error("relay update check failed", "error", err) + } } diff --git a/openflare_relay/internal/updater/restart_unix.go b/openflare_relay/internal/updater/restart_unix.go new file mode 100644 index 00000000..173af72b --- /dev/null +++ b/openflare_relay/internal/updater/restart_unix.go @@ -0,0 +1,51 @@ +//go:build !windows + +package updater + +import ( + "fmt" + "log/slog" + "os" + "syscall" +) + +func replaceAndRestart(execPath string, tmpPath string) error { + backupPath := execPath + ".bak" + if err := removeBackupBinary(backupPath); err != nil { + return err + } + if err := os.Rename(execPath, backupPath); err != nil { + renameErr := err + if err := os.Remove(tmpPath); err != nil && !os.IsNotExist(err) { + slog.Error("remove tmp binary failed", "path", tmpPath, "error", err) + return fmt.Errorf("backup current binary: %w; remove tmp binary: %v", renameErr, err) + } + return fmt.Errorf("backup current binary: %w", renameErr) + } + if err := os.Rename(tmpPath, execPath); err != nil { + replaceErr := err + if err := os.Rename(backupPath, execPath); err != nil { + slog.Error("restore backup binary failed", "path", backupPath, "error", err) + return fmt.Errorf("replace binary: %w; restore backup binary: %v", replaceErr, err) + } + return fmt.Errorf("replace binary: %w", replaceErr) + } + if err := removeBackupBinary(backupPath); err != nil { + return err + } + if err := syscall.Exec(execPath, os.Args, os.Environ()); err != nil { + return fmt.Errorf("exec restart: %w", err) + } + return fmt.Errorf("unreachable after exec") +} + +func removeBackupBinary(path string) error { + if err := os.Remove(path); err != nil { + if os.IsNotExist(err) { + return nil + } + slog.Error("remove backup binary failed", "path", path, "error", err) + return err + } + return nil +} diff --git a/openflare_relay/internal/updater/restart_windows.go b/openflare_relay/internal/updater/restart_windows.go new file mode 100644 index 00000000..66cb5e61 --- /dev/null +++ b/openflare_relay/internal/updater/restart_windows.go @@ -0,0 +1,53 @@ +//go:build windows + +package updater + +import ( + "fmt" + "os" + "os/exec" + "strings" +) + +func replaceAndRestart(execPath string, tmpPath string) error { + backupPath := execPath + ".bak" + scriptPath := execPath + ".update.cmd" + script := fmt.Sprintf(`@echo off +setlocal +:waitloop +move /Y "%s" "%s" >nul 2>nul +if errorlevel 1 ( + ping 127.0.0.1 -n 2 >nul + goto waitloop +) +move /Y "%s" "%s" >nul 2>nul +if errorlevel 1 exit /b 1 +start "" %s +del /Q "%s" >nul 2>nul +del /Q "%%~f0" >nul 2>nul +`, execPath, backupPath, tmpPath, execPath, buildWindowsCommandLine(execPath, os.Args[1:]), backupPath) + if err := os.WriteFile(scriptPath, []byte(script), 0o700); err != nil { + os.Remove(tmpPath) + return fmt.Errorf("write restart script: %w", err) + } + cmd := exec.Command("cmd", "/C", "start", "", scriptPath) + if err := cmd.Start(); err != nil { + os.Remove(scriptPath) + os.Remove(tmpPath) + return fmt.Errorf("schedule restart: %w", err) + } + os.Exit(0) + return nil +} + +func buildWindowsCommandLine(execPath string, args []string) string { + parts := []string{quoteWindowsArg(execPath)} + for _, arg := range args { + parts = append(parts, quoteWindowsArg(arg)) + } + return strings.Join(parts, " ") +} + +func quoteWindowsArg(value string) string { + return `"` + strings.ReplaceAll(value, `"`, `""`) + `"` +} diff --git a/openflare_relay/internal/updater/updater.go b/openflare_relay/internal/updater/updater.go new file mode 100644 index 00000000..f4e52a87 --- /dev/null +++ b/openflare_relay/internal/updater/updater.go @@ -0,0 +1,370 @@ +package updater + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "io" + "log/slog" + "net/http" + "openflare/utils" + "os" + "runtime" + "strings" + "time" + + "openflare-relay/internal/config" +) + +const maxChecksumAssetSize = 64 * 1024 + +var replaceAndRestartFunc = replaceAndRestart + +type Service struct { + httpClient *http.Client + lastCheckKey string +} + +func New() *Service { + return &Service{ + httpClient: &http.Client{Timeout: 30 * time.Second}, + } +} + +type githubRelease struct { + TagName string `json:"tag_name"` + Prerelease bool `json:"prerelease"` + Draft bool `json:"draft"` + Assets []githubAsset `json:"assets"` +} + +type githubAsset struct { + Name string `json:"name"` + BrowserDownloadURL string `json:"browser_download_url"` +} + +type UpdateOptions struct { + Channel string + TagName string + Force bool +} + +func (s *Service) CheckAndUpdate(ctx context.Context, repo string, options UpdateOptions) error { + release, err := s.getRelease(ctx, repo, options) + if err != nil { + return fmt.Errorf("check latest release: %w", err) + } + if release == nil || release.TagName == "" { + return nil + } + + remoteVersion := normalizeVersion(release.TagName) + localVersion := normalizeVersion(config.Version) + checkKey := buildReleaseCheckKey(options, remoteVersion) + + if remoteVersion == localVersion { + return nil + } + if !options.Force && checkKey != "" && checkKey == s.lastCheckKey { + return nil + } + if !isNewer(localVersion, remoteVersion) { + s.lastCheckKey = checkKey + return nil + } + + slog.Info("relay update available", "from", localVersion, "to", remoteVersion) + assetName := assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH) + checksumAssetName := assetName + ".sha256" + + var downloadURL string + var checksumURL string + for _, asset := range release.Assets { + switch asset.Name { + case assetName: + downloadURL = asset.BrowserDownloadURL + case checksumAssetName: + checksumURL = asset.BrowserDownloadURL + } + } + if downloadURL == "" { + s.lastCheckKey = checkKey + return fmt.Errorf("no matching asset %q in release %s", assetName, release.TagName) + } + if checksumURL == "" { + return fmt.Errorf("no matching checksum asset %q in release %s", checksumAssetName, release.TagName) + } + + expectedChecksum, err := s.downloadChecksum(ctx, checksumURL, assetName) + if err != nil { + return fmt.Errorf("download checksum: %w", err) + } + + execPath, err := os.Executable() + if err != nil { + return fmt.Errorf("get executable path: %w", err) + } + if err = s.downloadAndRestart(ctx, downloadURL, expectedChecksum, execPath); err != nil { + return fmt.Errorf("download and restart: %w", err) + } + s.lastCheckKey = checkKey + return nil +} + +func (s *Service) getRelease(ctx context.Context, repo string, options UpdateOptions) (*githubRelease, error) { + tagName := strings.TrimSpace(options.TagName) + if tagName != "" { + return s.getReleaseByTag(ctx, repo, tagName) + } + if strings.EqualFold(strings.TrimSpace(options.Channel), "preview") { + return s.getLatestPreviewRelease(ctx, repo) + } + return s.getLatestStableRelease(ctx, repo) +} + +func (s *Service) getLatestStableRelease(ctx context.Context, repo string) (*githubRelease, error) { + url := fmt.Sprintf("https://api.github.com/repos/%s/releases/latest", repo) + return s.fetchReleaseFromURL(ctx, url) +} + +func (s *Service) getLatestPreviewRelease(ctx context.Context, repo string) (*githubRelease, error) { + url := fmt.Sprintf("https://api.github.com/repos/%s/releases?per_page=20", repo) + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return nil, err + } + req.Header.Set("Accept", "application/vnd.github+json") + + resp, err := s.httpClient.Do(req) + if err != nil { + return nil, err + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("github api returned %s", resp.Status) + } + + var releases []githubRelease + if err = json.NewDecoder(resp.Body).Decode(&releases); err != nil { + return nil, err + } + for _, release := range releases { + if release.Draft || !release.Prerelease { + continue + } + releaseCopy := release + return &releaseCopy, nil + } + return nil, nil +} + +func (s *Service) getReleaseByTag(ctx context.Context, repo string, tag string) (*githubRelease, error) { + url := fmt.Sprintf("https://api.github.com/repos/%s/releases/tags/%s", repo, strings.TrimSpace(tag)) + return s.fetchReleaseFromURL(ctx, url) +} + +func (s *Service) fetchReleaseFromURL(ctx context.Context, url string) (*githubRelease, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return nil, err + } + req.Header.Set("Accept", "application/vnd.github+json") + + resp, err := s.httpClient.Do(req) + if err != nil { + return nil, err + } + defer func(Body io.ReadCloser) { + err := Body.Close() + if err != nil { + slog.Error("failed to close response body", "error", err) + } + }(resp.Body) + + if resp.StatusCode == http.StatusNotFound { + return nil, nil + } + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("github api returned %s", resp.Status) + } + + return decodeRelease(resp.Body) +} + +func decodeRelease(reader io.Reader) (*githubRelease, error) { + var release githubRelease + if err := json.NewDecoder(reader).Decode(&release); err != nil { + return nil, err + } + return &release, nil +} + +func (s *Service) downloadChecksum(ctx context.Context, url string, assetName string) (string, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return "", err + } + resp, err := s.httpClient.Do(req) + if err != nil { + return "", err + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return "", fmt.Errorf("checksum download returned %s", resp.Status) + } + + content, err := io.ReadAll(io.LimitReader(resp.Body, maxChecksumAssetSize+1)) + if err != nil { + return "", err + } + if len(content) > maxChecksumAssetSize { + return "", fmt.Errorf("checksum asset exceeds %d bytes", maxChecksumAssetSize) + } + checksum, err := parseSHA256Checksum(string(content), assetName) + if err != nil { + return "", err + } + return checksum, nil +} + +func parseSHA256Checksum(content string, assetName string) (string, error) { + assetName = strings.TrimSpace(assetName) + for _, line := range strings.Split(content, "\n") { + line = strings.TrimSpace(line) + if line == "" || strings.HasPrefix(line, "#") { + continue + } + if checksum, ok := parseSHA256Line(line, assetName); ok { + return checksum, nil + } + } + if assetName == "" { + return "", fmt.Errorf("checksum asset does not contain a valid sha256 digest") + } + return "", fmt.Errorf("checksum asset does not contain a sha256 digest for %q", assetName) +} + +func parseSHA256Line(line string, assetName string) (string, bool) { + fields := strings.Fields(line) + if len(fields) == 1 && isSHA256Hex(fields[0]) { + return strings.ToLower(fields[0]), true + } + if len(fields) >= 2 && isSHA256Hex(fields[0]) { + fileName := strings.TrimPrefix(strings.TrimSpace(fields[1]), "*") + if assetName == "" || fileName == assetName { + return strings.ToLower(fields[0]), true + } + } + + prefix := "SHA256(" + if strings.HasPrefix(line, prefix) { + closing := strings.Index(line, ")") + if closing > len(prefix) && closing+1 < len(line) { + fileName := strings.TrimSpace(line[len(prefix):closing]) + rest := strings.TrimSpace(line[closing+1:]) + rest = strings.TrimPrefix(rest, "=") + rest = strings.TrimSpace(rest) + if isSHA256Hex(rest) && (assetName == "" || fileName == assetName) { + return strings.ToLower(rest), true + } + } + } + return "", false +} + +func isSHA256Hex(value string) bool { + value = strings.TrimSpace(value) + if len(value) != sha256.Size*2 { + return false + } + _, err := hex.DecodeString(value) + return err == nil +} + +func (s *Service) downloadAndRestart(ctx context.Context, url string, expectedChecksum string, targetPath string) error { + expectedChecksum = strings.ToLower(strings.TrimSpace(expectedChecksum)) + if !isSHA256Hex(expectedChecksum) { + return fmt.Errorf("invalid expected sha256 checksum") + } + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return err + } + resp, err := s.httpClient.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("download returned %s", resp.Status) + } + + tmpPath := targetPath + ".update" + if runtime.GOOS == "windows" && !strings.HasSuffix(strings.ToLower(tmpPath), ".exe") { + tmpPath += ".exe" + } + tmpFile, err := os.OpenFile(tmpPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o600) + if err != nil { + return err + } + hasher := sha256.New() + if _, err = io.Copy(io.MultiWriter(tmpFile, hasher), resp.Body); err != nil { + tmpFile.Close() + os.Remove(tmpPath) + return err + } + if err = tmpFile.Close(); err != nil { + os.Remove(tmpPath) + return err + } + actualChecksum := hex.EncodeToString(hasher.Sum(nil)) + if actualChecksum != expectedChecksum { + os.Remove(tmpPath) + return fmt.Errorf("sha256 checksum mismatch: expected %s, got %s", expectedChecksum, actualChecksum) + } + if err = os.Chmod(tmpPath, 0o755); err != nil && runtime.GOOS != "windows" { + os.Remove(tmpPath) + return fmt.Errorf("set executable permission: %w", err) + } + + slog.Info("relay binary updated, restarting") + return replaceAndRestartFunc(targetPath, tmpPath) +} + +func assetNameForGOOSGOARCH(goos string, goarch string) string { + name := fmt.Sprintf("openflare-relay-%s-%s", goos, goarch) + if goos == "windows" { + return name + ".exe" + } + return name +} + +func normalizeVersion(v string) string { + v = strings.TrimSpace(v) + v = strings.TrimPrefix(v, "v") + return v +} + +func isNewer(local, remote string) bool { + return compareVersions(local, remote) < 0 +} + +func buildReleaseCheckKey(options UpdateOptions, remoteVersion string) string { + channel := strings.TrimSpace(options.Channel) + if channel == "" { + channel = "stable" + } + if tagName := strings.TrimSpace(options.TagName); tagName != "" { + return channel + ":" + tagName + } + return channel + ":" + remoteVersion +} + +func compareVersions(local string, remote string) int { + return utils.CompareVersions(local, remote) +} diff --git a/openflare_server/service/relay.go b/openflare_server/service/relay.go index 760f1c7c..dcdedcc0 100644 --- a/openflare_server/service/relay.go +++ b/openflare_server/service/relay.go @@ -40,8 +40,13 @@ type RelayConfig struct { // RelaySettings contains runtime settings for the Relay. type RelaySettings struct { - HeartbeatInterval int `json:"heartbeat_interval"` - WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"` + HeartbeatInterval int `json:"heartbeat_interval"` + WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"` + AutoUpdate bool `json:"auto_update"` + UpdateRepo string `json:"update_repo"` + UpdateNow bool `json:"update_now"` + UpdateChannel string `json:"update_channel"` + UpdateTag string `json:"update_tag"` } // RelayHeartbeatResponse is the response returned to the Relay from a heartbeat. @@ -63,6 +68,15 @@ func HeartbeatRelay(node *model.Node, payload RelayHeartbeatPayload) (*RelayHear payload.Name = strings.TrimSpace(payload.Name) payload.IP = strings.TrimSpace(payload.IP) + previous := *node + updateNow := node.UpdateRequested + updateChannel := normalizeReleaseChannel(node.UpdateChannel) + updateTag := strings.TrimSpace(node.UpdateTag) + + node.UpdateRequested = false + node.UpdateChannel = ReleaseChannelStable.String() + node.UpdateTag = "" + changes := make(map[string]any) appendRelayChange := func(key string, before any, after any) { if before != after { @@ -74,6 +88,16 @@ func HeartbeatRelay(node *model.Node, payload RelayHeartbeatPayload) (*RelayHear appendRelayChange("ext_version", node.ExtVersion, payload.ExtVersion) appendRelayChange("relay_status", node.RelayStatus, payload.RelayStatus) + if previous.UpdateRequested { + appendRelayChange("update_requested", previous.UpdateRequested, false) + } + if previous.UpdateChannel != ReleaseChannelStable.String() { + appendRelayChange("update_channel", previous.UpdateChannel, ReleaseChannelStable.String()) + } + if previous.UpdateTag != "" { + appendRelayChange("update_tag", previous.UpdateTag, "") + } + if payload.Name != "" && strings.TrimSpace(node.Name) == "" { appendRelayChange("name", node.Name, payload.Name) node.Name = payload.Name @@ -113,7 +137,7 @@ func HeartbeatRelay(node *model.Node, payload RelayHeartbeatPayload) (*RelayHear return &RelayHeartbeatResponse{ RelayConfig: buildRelayConfig(node), - RelaySettings: buildRelaySettings(), + RelaySettings: buildRelaySettings(node, updateNow, updateChannel.String(), updateTag), }, nil } @@ -145,10 +169,22 @@ func buildRelayConfig(node *model.Node) *RelayConfig { } } -func buildRelaySettings() *RelaySettings { +func buildRelaySettings(node *model.Node, updateNow bool, updateChannel string, updateTag string) *RelaySettings { + autoUpdate := false + if node != nil { + autoUpdate = node.AutoUpdateEnabled + } + if strings.TrimSpace(updateChannel) == "" { + updateChannel = ReleaseChannelStable.String() + } return &RelaySettings{ HeartbeatInterval: common.AgentHeartbeatInterval, WebsocketUpgradeEnabled: common.AgentWebsocketUpgradeEnabled, + AutoUpdate: autoUpdate, + UpdateRepo: common.AgentUpdateRepo, + UpdateNow: updateNow, + UpdateChannel: updateChannel, + UpdateTag: strings.TrimSpace(updateTag), } } @@ -236,6 +272,13 @@ func HeartbeatFlared(node *model.Node, payload FlaredHeartbeatPayload) (*FlaredH now := time.Now() previous := *node + updateNow := node.UpdateRequested + updateChannel := normalizeReleaseChannel(node.UpdateChannel) + updateTag := strings.TrimSpace(node.UpdateTag) + + node.UpdateRequested = false + node.UpdateChannel = ReleaseChannelStable.String() + node.UpdateTag = "" changes := make(map[string]any) if previous.Version != payload.ClientVersion { @@ -258,6 +301,16 @@ func HeartbeatFlared(node *model.Node, payload FlaredHeartbeatPayload) (*FlaredH node.LastSeenAt = now node.Status = NodeStatusOnline + if previous.UpdateRequested { + changes["update_requested"] = false + } + if previous.UpdateChannel != ReleaseChannelStable.String() { + changes["update_channel"] = ReleaseChannelStable.String() + } + if previous.UpdateTag != "" { + changes["update_tag"] = "" + } + if !node.IPManualOverride && payload.IP != "" && previous.IP != payload.IP { changes["ip"] = payload.IP node.IP = payload.IP @@ -290,7 +343,7 @@ func HeartbeatFlared(node *model.Node, payload FlaredHeartbeatPayload) (*FlaredH } return &FlaredHeartbeatResponse{ ActiveConfig: activeConfig, - TunnelSettings: buildRelaySettings(), + TunnelSettings: buildRelaySettings(node, updateNow, updateChannel.String(), updateTag), }, nil } diff --git a/openflare_server/service/relay_test.go b/openflare_server/service/relay_test.go index 897960b5..e9c41dc2 100644 --- a/openflare_server/service/relay_test.go +++ b/openflare_server/service/relay_test.go @@ -401,3 +401,127 @@ func TestTunnelRoutePublishAndFlaredConfigUseRelayPorts(t *testing.T) { t.Fatalf("unexpected proxy domains: %+v", proxy.CustomDomains) } } + +func TestHeartbeatRelaySelfUpdatePropagationAndReset(t *testing.T) { + setupServiceTestDB(t) + + node := &model.Node{ + NodeID: "relay-update-node", + Name: "relay-u", + IP: "1.1.1.1", + AccessToken: "relay-update-token", + Status: NodeStatusPending, + NodeType: "tunnel_relay", + AutoUpdateEnabled: true, + UpdateRequested: true, + UpdateChannel: "preview", + UpdateTag: "v1.2.3", + } + if err := node.Insert(); err != nil { + t.Fatalf("failed to seed relay node: %v", err) + } + + resp, err := HeartbeatRelay(node, RelayHeartbeatPayload{ + Version: "v1.0.0", + ExtVersion: "0.61.0", + RelayStatus: "healthy", + }) + if err != nil { + t.Fatalf("HeartbeatRelay failed: %v", err) + } + + if resp == nil || resp.RelaySettings == nil { + t.Fatal("expected non-nil response with RelaySettings") + } + + settings := resp.RelaySettings + if !settings.AutoUpdate { + t.Error("expected AutoUpdate to be true") + } + if !settings.UpdateNow { + t.Error("expected UpdateNow to be true") + } + if settings.UpdateChannel != "preview" { + t.Errorf("expected UpdateChannel to be preview, got %q", settings.UpdateChannel) + } + if settings.UpdateTag != "v1.2.3" { + t.Errorf("expected UpdateTag to be v1.2.3, got %q", settings.UpdateTag) + } + + // Verify that the requested update was cleared in the DB + updated, err := model.GetNodeByNodeID(node.NodeID) + if err != nil { + t.Fatalf("failed to reload node: %v", err) + } + if updated.UpdateRequested { + t.Error("expected UpdateRequested to be reset to false in the database") + } + if updated.UpdateChannel != "stable" { + t.Errorf("expected UpdateChannel to be reset to stable, got %q", updated.UpdateChannel) + } + if updated.UpdateTag != "" { + t.Errorf("expected UpdateTag to be reset to empty, got %q", updated.UpdateTag) + } +} + +func TestHeartbeatFlaredSelfUpdatePropagationAndReset(t *testing.T) { + setupServiceTestDB(t) + + node := &model.Node{ + NodeID: "flared-update-node", + Name: "flared-u", + IP: "1.1.1.2", + AccessToken: "flared-update-token", + Status: NodeStatusPending, + NodeType: "tunnel_client", + AutoUpdateEnabled: true, + UpdateRequested: true, + UpdateChannel: "stable", + UpdateTag: "v2.3.4", + } + if err := node.Insert(); err != nil { + t.Fatalf("failed to seed flared node: %v", err) + } + + resp, err := HeartbeatFlared(node, FlaredHeartbeatPayload{ + ClientVersion: "v1.0.0", + FrpVersion: "0.61.0", + TunnelStatus: "running", + }) + if err != nil { + t.Fatalf("HeartbeatFlared failed: %v", err) + } + + if resp == nil || resp.TunnelSettings == nil { + t.Fatal("expected non-nil response with TunnelSettings") + } + + settings := resp.TunnelSettings + if !settings.AutoUpdate { + t.Error("expected AutoUpdate to be true") + } + if !settings.UpdateNow { + t.Error("expected UpdateNow to be true") + } + if settings.UpdateChannel != "stable" { + t.Errorf("expected UpdateChannel to be stable, got %q", settings.UpdateChannel) + } + if settings.UpdateTag != "v2.3.4" { + t.Errorf("expected UpdateTag to be v2.3.4, got %q", settings.UpdateTag) + } + + // Verify that the requested update was cleared in the DB + updated, err := model.GetNodeByNodeID(node.NodeID) + if err != nil { + t.Fatalf("failed to reload node: %v", err) + } + if updated.UpdateRequested { + t.Error("expected UpdateRequested to be reset to false in the database") + } + if updated.UpdateChannel != "stable" { + t.Errorf("expected UpdateChannel to be reset to stable, got %q", updated.UpdateChannel) + } + if updated.UpdateTag != "" { + t.Errorf("expected UpdateTag to be reset to empty, got %q", updated.UpdateTag) + } +} diff --git a/openflared/internal/heartbeat/service.go b/openflared/internal/heartbeat/service.go index a565461b..85d3f5a3 100644 --- a/openflared/internal/heartbeat/service.go +++ b/openflared/internal/heartbeat/service.go @@ -9,6 +9,7 @@ import ( "openflare-flared/internal/config" "openflare-flared/internal/frpc" "openflare-flared/internal/httpclient" + "openflare-flared/internal/updater" "openflare/service" "openflare/utils/geoip" "openflare/utils/geoip/iputil" @@ -23,6 +24,7 @@ type Service struct { client *httpclient.Client frpcManager *frpc.Manager config *config.Config + updater *updater.Service } func New(client *httpclient.Client, manager *frpc.Manager, cfg *config.Config) *Service { @@ -30,6 +32,7 @@ func New(client *httpclient.Client, manager *frpc.Manager, cfg *config.Config) * client: client, frpcManager: manager, config: cfg, + updater: updater.New(), } } @@ -65,12 +68,40 @@ func (s *Service) doHeartbeat(ctx context.Context) { CurrentChecksum: s.frpcManager.GetCurrentConfigChecksum(), } - _, err := s.client.Heartbeat(ctx, payload) + resp, err := s.client.Heartbeat(ctx, payload) if err != nil { slog.Error("flared heartbeat failed", "error", err) return } slog.Debug("flared heartbeat succeeded") + + if resp != nil && resp.TunnelSettings != nil { + s.tryAutoUpdate(ctx, resp.TunnelSettings) + } +} + +func (s *Service) tryAutoUpdate(ctx context.Context, settings *service.RelaySettings) { + if settings == nil || s.updater == nil { + return + } + force := settings.UpdateNow + shouldCheck := settings.AutoUpdate || force + if !shouldCheck || settings.UpdateRepo == "" { + return + } + channel := "stable" + if force && settings.UpdateChannel != "" { + channel = settings.UpdateChannel + } + slog.Info("checking for client updates", "repo", settings.UpdateRepo, "channel", channel, "force", force) + err := s.updater.CheckAndUpdate(ctx, settings.UpdateRepo, updater.UpdateOptions{ + Channel: channel, + TagName: settings.UpdateTag, + Force: force, + }) + if err != nil { + slog.Error("client update check failed", "error", err) + } } func detectNodeIP() string { diff --git a/openflared/internal/updater/restart_unix.go b/openflared/internal/updater/restart_unix.go new file mode 100644 index 00000000..173af72b --- /dev/null +++ b/openflared/internal/updater/restart_unix.go @@ -0,0 +1,51 @@ +//go:build !windows + +package updater + +import ( + "fmt" + "log/slog" + "os" + "syscall" +) + +func replaceAndRestart(execPath string, tmpPath string) error { + backupPath := execPath + ".bak" + if err := removeBackupBinary(backupPath); err != nil { + return err + } + if err := os.Rename(execPath, backupPath); err != nil { + renameErr := err + if err := os.Remove(tmpPath); err != nil && !os.IsNotExist(err) { + slog.Error("remove tmp binary failed", "path", tmpPath, "error", err) + return fmt.Errorf("backup current binary: %w; remove tmp binary: %v", renameErr, err) + } + return fmt.Errorf("backup current binary: %w", renameErr) + } + if err := os.Rename(tmpPath, execPath); err != nil { + replaceErr := err + if err := os.Rename(backupPath, execPath); err != nil { + slog.Error("restore backup binary failed", "path", backupPath, "error", err) + return fmt.Errorf("replace binary: %w; restore backup binary: %v", replaceErr, err) + } + return fmt.Errorf("replace binary: %w", replaceErr) + } + if err := removeBackupBinary(backupPath); err != nil { + return err + } + if err := syscall.Exec(execPath, os.Args, os.Environ()); err != nil { + return fmt.Errorf("exec restart: %w", err) + } + return fmt.Errorf("unreachable after exec") +} + +func removeBackupBinary(path string) error { + if err := os.Remove(path); err != nil { + if os.IsNotExist(err) { + return nil + } + slog.Error("remove backup binary failed", "path", path, "error", err) + return err + } + return nil +} diff --git a/openflared/internal/updater/restart_windows.go b/openflared/internal/updater/restart_windows.go new file mode 100644 index 00000000..66cb5e61 --- /dev/null +++ b/openflared/internal/updater/restart_windows.go @@ -0,0 +1,53 @@ +//go:build windows + +package updater + +import ( + "fmt" + "os" + "os/exec" + "strings" +) + +func replaceAndRestart(execPath string, tmpPath string) error { + backupPath := execPath + ".bak" + scriptPath := execPath + ".update.cmd" + script := fmt.Sprintf(`@echo off +setlocal +:waitloop +move /Y "%s" "%s" >nul 2>nul +if errorlevel 1 ( + ping 127.0.0.1 -n 2 >nul + goto waitloop +) +move /Y "%s" "%s" >nul 2>nul +if errorlevel 1 exit /b 1 +start "" %s +del /Q "%s" >nul 2>nul +del /Q "%%~f0" >nul 2>nul +`, execPath, backupPath, tmpPath, execPath, buildWindowsCommandLine(execPath, os.Args[1:]), backupPath) + if err := os.WriteFile(scriptPath, []byte(script), 0o700); err != nil { + os.Remove(tmpPath) + return fmt.Errorf("write restart script: %w", err) + } + cmd := exec.Command("cmd", "/C", "start", "", scriptPath) + if err := cmd.Start(); err != nil { + os.Remove(scriptPath) + os.Remove(tmpPath) + return fmt.Errorf("schedule restart: %w", err) + } + os.Exit(0) + return nil +} + +func buildWindowsCommandLine(execPath string, args []string) string { + parts := []string{quoteWindowsArg(execPath)} + for _, arg := range args { + parts = append(parts, quoteWindowsArg(arg)) + } + return strings.Join(parts, " ") +} + +func quoteWindowsArg(value string) string { + return `"` + strings.ReplaceAll(value, `"`, `""`) + `"` +} diff --git a/openflared/internal/updater/updater.go b/openflared/internal/updater/updater.go new file mode 100644 index 00000000..8dbc6545 --- /dev/null +++ b/openflared/internal/updater/updater.go @@ -0,0 +1,370 @@ +package updater + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "io" + "log/slog" + "net/http" + "openflare/utils" + "os" + "runtime" + "strings" + "time" + + "openflare-flared/internal/config" +) + +const maxChecksumAssetSize = 64 * 1024 + +var replaceAndRestartFunc = replaceAndRestart + +type Service struct { + httpClient *http.Client + lastCheckKey string +} + +func New() *Service { + return &Service{ + httpClient: &http.Client{Timeout: 30 * time.Second}, + } +} + +type githubRelease struct { + TagName string `json:"tag_name"` + Prerelease bool `json:"prerelease"` + Draft bool `json:"draft"` + Assets []githubAsset `json:"assets"` +} + +type githubAsset struct { + Name string `json:"name"` + BrowserDownloadURL string `json:"browser_download_url"` +} + +type UpdateOptions struct { + Channel string + TagName string + Force bool +} + +func (s *Service) CheckAndUpdate(ctx context.Context, repo string, options UpdateOptions) error { + release, err := s.getRelease(ctx, repo, options) + if err != nil { + return fmt.Errorf("check latest release: %w", err) + } + if release == nil || release.TagName == "" { + return nil + } + + remoteVersion := normalizeVersion(release.TagName) + localVersion := normalizeVersion(config.Version) + checkKey := buildReleaseCheckKey(options, remoteVersion) + + if remoteVersion == localVersion { + return nil + } + if !options.Force && checkKey != "" && checkKey == s.lastCheckKey { + return nil + } + if !isNewer(localVersion, remoteVersion) { + s.lastCheckKey = checkKey + return nil + } + + slog.Info("flared update available", "from", localVersion, "to", remoteVersion) + assetName := assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH) + checksumAssetName := assetName + ".sha256" + + var downloadURL string + var checksumURL string + for _, asset := range release.Assets { + switch asset.Name { + case assetName: + downloadURL = asset.BrowserDownloadURL + case checksumAssetName: + checksumURL = asset.BrowserDownloadURL + } + } + if downloadURL == "" { + s.lastCheckKey = checkKey + return fmt.Errorf("no matching asset %q in release %s", assetName, release.TagName) + } + if checksumURL == "" { + return fmt.Errorf("no matching checksum asset %q in release %s", checksumAssetName, release.TagName) + } + + expectedChecksum, err := s.downloadChecksum(ctx, checksumURL, assetName) + if err != nil { + return fmt.Errorf("download checksum: %w", err) + } + + execPath, err := os.Executable() + if err != nil { + return fmt.Errorf("get executable path: %w", err) + } + if err = s.downloadAndRestart(ctx, downloadURL, expectedChecksum, execPath); err != nil { + return fmt.Errorf("download and restart: %w", err) + } + s.lastCheckKey = checkKey + return nil +} + +func (s *Service) getRelease(ctx context.Context, repo string, options UpdateOptions) (*githubRelease, error) { + tagName := strings.TrimSpace(options.TagName) + if tagName != "" { + return s.getReleaseByTag(ctx, repo, tagName) + } + if strings.EqualFold(strings.TrimSpace(options.Channel), "preview") { + return s.getLatestPreviewRelease(ctx, repo) + } + return s.getLatestStableRelease(ctx, repo) +} + +func (s *Service) getLatestStableRelease(ctx context.Context, repo string) (*githubRelease, error) { + url := fmt.Sprintf("https://api.github.com/repos/%s/releases/latest", repo) + return s.fetchReleaseFromURL(ctx, url) +} + +func (s *Service) getLatestPreviewRelease(ctx context.Context, repo string) (*githubRelease, error) { + url := fmt.Sprintf("https://api.github.com/repos/%s/releases?per_page=20", repo) + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return nil, err + } + req.Header.Set("Accept", "application/vnd.github+json") + + resp, err := s.httpClient.Do(req) + if err != nil { + return nil, err + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("github api returned %s", resp.Status) + } + + var releases []githubRelease + if err = json.NewDecoder(resp.Body).Decode(&releases); err != nil { + return nil, err + } + for _, release := range releases { + if release.Draft || !release.Prerelease { + continue + } + releaseCopy := release + return &releaseCopy, nil + } + return nil, nil +} + +func (s *Service) getReleaseByTag(ctx context.Context, repo string, tag string) (*githubRelease, error) { + url := fmt.Sprintf("https://api.github.com/repos/%s/releases/tags/%s", repo, strings.TrimSpace(tag)) + return s.fetchReleaseFromURL(ctx, url) +} + +func (s *Service) fetchReleaseFromURL(ctx context.Context, url string) (*githubRelease, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return nil, err + } + req.Header.Set("Accept", "application/vnd.github+json") + + resp, err := s.httpClient.Do(req) + if err != nil { + return nil, err + } + defer func(Body io.ReadCloser) { + err := Body.Close() + if err != nil { + slog.Error("failed to close response body", "error", err) + } + }(resp.Body) + + if resp.StatusCode == http.StatusNotFound { + return nil, nil + } + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("github api returned %s", resp.Status) + } + + return decodeRelease(resp.Body) +} + +func decodeRelease(reader io.Reader) (*githubRelease, error) { + var release githubRelease + if err := json.NewDecoder(reader).Decode(&release); err != nil { + return nil, err + } + return &release, nil +} + +func (s *Service) downloadChecksum(ctx context.Context, url string, assetName string) (string, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return "", err + } + resp, err := s.httpClient.Do(req) + if err != nil { + return "", err + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return "", fmt.Errorf("checksum download returned %s", resp.Status) + } + + content, err := io.ReadAll(io.LimitReader(resp.Body, maxChecksumAssetSize+1)) + if err != nil { + return "", err + } + if len(content) > maxChecksumAssetSize { + return "", fmt.Errorf("checksum asset exceeds %d bytes", maxChecksumAssetSize) + } + checksum, err := parseSHA256Checksum(string(content), assetName) + if err != nil { + return "", err + } + return checksum, nil +} + +func parseSHA256Checksum(content string, assetName string) (string, error) { + assetName = strings.TrimSpace(assetName) + for _, line := range strings.Split(content, "\n") { + line = strings.TrimSpace(line) + if line == "" || strings.HasPrefix(line, "#") { + continue + } + if checksum, ok := parseSHA256Line(line, assetName); ok { + return checksum, nil + } + } + if assetName == "" { + return "", fmt.Errorf("checksum asset does not contain a valid sha256 digest") + } + return "", fmt.Errorf("checksum asset does not contain a sha256 digest for %q", assetName) +} + +func parseSHA256Line(line string, assetName string) (string, bool) { + fields := strings.Fields(line) + if len(fields) == 1 && isSHA256Hex(fields[0]) { + return strings.ToLower(fields[0]), true + } + if len(fields) >= 2 && isSHA256Hex(fields[0]) { + fileName := strings.TrimPrefix(strings.TrimSpace(fields[1]), "*") + if assetName == "" || fileName == assetName { + return strings.ToLower(fields[0]), true + } + } + + prefix := "SHA256(" + if strings.HasPrefix(line, prefix) { + closing := strings.Index(line, ")") + if closing > len(prefix) && closing+1 < len(line) { + fileName := strings.TrimSpace(line[len(prefix):closing]) + rest := strings.TrimSpace(line[closing+1:]) + rest = strings.TrimPrefix(rest, "=") + rest = strings.TrimSpace(rest) + if isSHA256Hex(rest) && (assetName == "" || fileName == assetName) { + return strings.ToLower(rest), true + } + } + } + return "", false +} + +func isSHA256Hex(value string) bool { + value = strings.TrimSpace(value) + if len(value) != sha256.Size*2 { + return false + } + _, err := hex.DecodeString(value) + return err == nil +} + +func (s *Service) downloadAndRestart(ctx context.Context, url string, expectedChecksum string, targetPath string) error { + expectedChecksum = strings.ToLower(strings.TrimSpace(expectedChecksum)) + if !isSHA256Hex(expectedChecksum) { + return fmt.Errorf("invalid expected sha256 checksum") + } + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return err + } + resp, err := s.httpClient.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("download returned %s", resp.Status) + } + + tmpPath := targetPath + ".update" + if runtime.GOOS == "windows" && !strings.HasSuffix(strings.ToLower(tmpPath), ".exe") { + tmpPath += ".exe" + } + tmpFile, err := os.OpenFile(tmpPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o600) + if err != nil { + return err + } + hasher := sha256.New() + if _, err = io.Copy(io.MultiWriter(tmpFile, hasher), resp.Body); err != nil { + tmpFile.Close() + os.Remove(tmpPath) + return err + } + if err = tmpFile.Close(); err != nil { + os.Remove(tmpPath) + return err + } + actualChecksum := hex.EncodeToString(hasher.Sum(nil)) + if actualChecksum != expectedChecksum { + os.Remove(tmpPath) + return fmt.Errorf("sha256 checksum mismatch: expected %s, got %s", expectedChecksum, actualChecksum) + } + if err = os.Chmod(tmpPath, 0o755); err != nil && runtime.GOOS != "windows" { + os.Remove(tmpPath) + return fmt.Errorf("set executable permission: %w", err) + } + + slog.Info("flared binary updated, restarting") + return replaceAndRestartFunc(targetPath, tmpPath) +} + +func assetNameForGOOSGOARCH(goos string, goarch string) string { + name := fmt.Sprintf("openflared-%s-%s", goos, goarch) + if goos == "windows" { + return name + ".exe" + } + return name +} + +func normalizeVersion(v string) string { + v = strings.TrimSpace(v) + v = strings.TrimPrefix(v, "v") + return v +} + +func isNewer(local, remote string) bool { + return compareVersions(local, remote) < 0 +} + +func buildReleaseCheckKey(options UpdateOptions, remoteVersion string) string { + channel := strings.TrimSpace(options.Channel) + if channel == "" { + channel = "stable" + } + if tagName := strings.TrimSpace(options.TagName); tagName != "" { + return channel + ":" + tagName + } + return channel + ":" + remoteVersion +} + +func compareVersions(local string, remote string) int { + return utils.CompareVersions(local, remote) +}