diff --git a/.github/workflows/autoupgrade-test.yml b/.github/workflows/autoupgrade-test.yml index 03091ef23..7c69e01a5 100644 --- a/.github/workflows/autoupgrade-test.yml +++ b/.github/workflows/autoupgrade-test.yml @@ -1,19 +1,376 @@ -# Stub: registers this workflow on the default branch so it can be dispatched -# to v2 via `gh workflow run autoupgrade-test.yml --ref v2`. -# The real implementation lives on the v2 branch. -name: Auto-upgrade canary test - on: workflow_dispatch: - + inputs: + dryrun: + description: 'Skip PagerDuty alert on failure' + type: boolean + default: true +name: Auto-upgrade test permissions: contents: read - jobs: - notice: + resolve-versions: runs-on: ubuntu-latest + outputs: + latest: ${{ steps.versions.outputs.latest }} steps: - - run: | - echo "This workflow should be dispatched with --ref v2." - echo "The real implementation lives on the v2 branch." + - name: Fetch latest release version + id: versions + run: | + for i in 1 2 3; do + LATEST=$(curl -sSL -H "Authorization: Bearer $GITHUB_TOKEN" \ + "https://api.github.com/repos/stripe/stripe-cli/releases/latest" | \ + sed -n 's/.*"tag_name": *"v\([^"]*\)".*/\1/p') + if [ -n "$LATEST" ]; then + echo "latest=$LATEST" >> "$GITHUB_OUTPUT" + echo "Latest: $LATEST" + exit 0 + fi + echo "Attempt $i: failed to fetch version, retrying in 10s..." + sleep 10 + done + echo "Error: could not determine latest version after 3 attempts" exit 1 + env: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} + + upgrade-macos: + runs-on: macos-latest + needs: [resolve-versions] + outputs: + failed_step: ${{ steps.result.outputs.failed_step }} + steps: + - name: Checkout code + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6.1.0 + with: + persist-credentials: false + + - name: Install Go + uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0 + with: + go-version: '1.26.0' + + - name: Build binary with fake old version + id: build + env: + STRIPE_NO_AUTO_UPDATE: "1" + run: | + go generate ./... + CGO_ENABLED=0 go build -ldflags="-s -w -X github.com/stripe/stripe-cli/pkg/version.Version=1.0.0" \ + -o "$HOME/.stripe/bin/stripe" ./cmd/stripe + chmod +x "$HOME/.stripe/bin/stripe" + export PATH="$HOME/.stripe/bin:$PATH" + INSTALLED=$(stripe --version | grep -oE '[0-9]+\.[0-9]+\.[0-9]+') + if [ "$INSTALLED" != "1.0.0" ]; then + echo "Error: expected version 1.0.0, got $INSTALLED" + exit 1 + fi + echo "Built binary with version v$INSTALLED" + + - name: Stage update marker + id: update-check + run: | + LATEST="${{ needs.resolve-versions.outputs.latest }}" + ARCH=$(uname -m) + case "$ARCH" in + x86_64) ARCH_NAME="x86_64" ;; + arm64|aarch64) ARCH_NAME="arm64" ;; + *) ARCH_NAME="$ARCH" ;; + esac + URL="https://github.com/stripe/stripe-cli/releases/download/v${LATEST}/stripe_${LATEST}_mac-os_${ARCH_NAME}.tar.gz" + + # Must match autoupdate.UpdateMarker's JSON encoding in pkg/autoupdate/checker.go. + mkdir -p "$HOME/.stripe/state" + printf '{"version":"%s","download_url":"%s","checksum":"","release_notes":""}' \ + "$LATEST" "$URL" > "$HOME/.stripe/state/update-available" + echo "Wrote update marker:" + cat "$HOME/.stripe/state/update-available" + + - name: Verify auto-upgrade + id: upgrade + env: + STRIPE_INSTALL_METHOD: curl + run: | + export PATH="$HOME/.stripe/bin:$PATH" + OUTPUT=$(stripe --version 2>&1) + echo "$OUTPUT" + UPGRADED=$(echo "$OUTPUT" | grep -oE '[0-9]+\.[0-9]+\.[0-9]+' | tail -1) + if [ "$UPGRADED" != "${{ needs.resolve-versions.outputs.latest }}" ]; then + echo "Error: expected upgrade to v${{ needs.resolve-versions.outputs.latest }}, got v$UPGRADED" + exit 1 + fi + echo "Successfully upgraded to v$UPGRADED" + + - name: Record failed step + id: result + if: failure() + run: | + if [ "${{ steps.build.outcome }}" = "failure" ]; then + echo "failed_step=build" >> "$GITHUB_OUTPUT" + elif [ "${{ steps.update-check.outcome }}" = "failure" ]; then + echo "failed_step=update-check" >> "$GITHUB_OUTPUT" + elif [ "${{ steps.upgrade.outcome }}" = "failure" ]; then + echo "failed_step=auto-upgrade" >> "$GITHUB_OUTPUT" + fi + + upgrade-linux: + runs-on: ubuntu-latest + needs: [resolve-versions] + outputs: + failed_step: ${{ steps.result.outputs.failed_step }} + steps: + - name: Checkout code + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6.1.0 + with: + persist-credentials: false + + - name: Install Go + uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0 + with: + go-version: '1.26.0' + + - name: Build binary with fake old version + id: build + env: + STRIPE_NO_AUTO_UPDATE: "1" + run: | + go generate ./... + CGO_ENABLED=0 go build -ldflags="-s -w -X github.com/stripe/stripe-cli/pkg/version.Version=1.0.0" \ + -o "$HOME/.stripe/bin/stripe" ./cmd/stripe + chmod +x "$HOME/.stripe/bin/stripe" + export PATH="$HOME/.stripe/bin:$PATH" + INSTALLED=$(stripe --version | grep -oE '[0-9]+\.[0-9]+\.[0-9]+') + if [ "$INSTALLED" != "1.0.0" ]; then + echo "Error: expected version 1.0.0, got $INSTALLED" + exit 1 + fi + echo "Built binary with version v$INSTALLED" + + - name: Stage update marker + id: update-check + run: | + LATEST="${{ needs.resolve-versions.outputs.latest }}" + ARCH=$(uname -m) + case "$ARCH" in + x86_64) ARCH_NAME="x86_64" ;; + arm64|aarch64) ARCH_NAME="arm64" ;; + *) ARCH_NAME="$ARCH" ;; + esac + URL="https://github.com/stripe/stripe-cli/releases/download/v${LATEST}/stripe_${LATEST}_linux_${ARCH_NAME}.tar.gz" + + # Must match autoupdate.UpdateMarker's JSON encoding in pkg/autoupdate/checker.go. + mkdir -p "$HOME/.stripe/state" + printf '{"version":"%s","download_url":"%s","checksum":"","release_notes":""}' \ + "$LATEST" "$URL" > "$HOME/.stripe/state/update-available" + echo "Wrote update marker:" + cat "$HOME/.stripe/state/update-available" + + - name: Verify auto-upgrade + id: upgrade + env: + STRIPE_INSTALL_METHOD: curl + run: | + export PATH="$HOME/.stripe/bin:$PATH" + OUTPUT=$(stripe --version 2>&1) + echo "$OUTPUT" + UPGRADED=$(echo "$OUTPUT" | grep -oE '[0-9]+\.[0-9]+\.[0-9]+' | tail -1) + if [ "$UPGRADED" != "${{ needs.resolve-versions.outputs.latest }}" ]; then + echo "Error: expected upgrade to v${{ needs.resolve-versions.outputs.latest }}, got v$UPGRADED" + exit 1 + fi + echo "Successfully upgraded to v$UPGRADED" + + - name: Record failed step + id: result + if: failure() + run: | + if [ "${{ steps.build.outcome }}" = "failure" ]; then + echo "failed_step=build" >> "$GITHUB_OUTPUT" + elif [ "${{ steps.update-check.outcome }}" = "failure" ]; then + echo "failed_step=update-check" >> "$GITHUB_OUTPUT" + elif [ "${{ steps.upgrade.outcome }}" = "failure" ]; then + echo "failed_step=auto-upgrade" >> "$GITHUB_OUTPUT" + fi + + # Windows earns a few steps the other two do not need. The outgoing binary + # cannot be deleted while a process is running from it, so an update there moves + # it aside instead and a later invocation clears it — which is the part of the + # sequence that leaves a trace on disk, and so the part worth asserting on. + upgrade-windows: + runs-on: windows-latest + needs: [resolve-versions] + defaults: + run: + shell: bash + outputs: + failed_step: ${{ steps.result.outputs.failed_step }} + steps: + - name: Checkout code + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6.1.0 + with: + persist-credentials: false + + - name: Install Go + uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0 + with: + go-version: '1.26.0' + + - name: Build binary with fake old version + id: build + env: + STRIPE_NO_AUTO_UPDATE: "1" + run: | + go generate ./... + CGO_ENABLED=0 go build -ldflags="-s -w -X github.com/stripe/stripe-cli/pkg/version.Version=1.0.0" \ + -o "$HOME/.stripe/bin/stripe.exe" ./cmd/stripe + INSTALLED=$("$HOME/.stripe/bin/stripe.exe" --version | grep -oE '[0-9]+\.[0-9]+\.[0-9]+') + if [ "$INSTALLED" != "1.0.0" ]; then + echo "Error: expected version 1.0.0, got $INSTALLED" + exit 1 + fi + echo "Built binary with version v$INSTALLED" + + - name: Stage update marker + id: update-check + run: | + LATEST="${{ needs.resolve-versions.outputs.latest }}" + # .goreleaser/windows.yml publishes x86_64 and i386 only; a Windows ARM + # machine runs the x64 archive under emulation, and the runners are x64. + URL="https://github.com/stripe/stripe-cli/releases/download/v${LATEST}/stripe_${LATEST}_windows_x86_64.zip" + + # Must match autoupdate.UpdateMarker's JSON encoding in pkg/autoupdate/checker.go. + mkdir -p "$HOME/.stripe/state" + printf '{"version":"%s","download_url":"%s","checksum":"","release_notes":""}' \ + "$LATEST" "$URL" > "$HOME/.stripe/state/update-available" + echo "Wrote update marker:" + cat "$HOME/.stripe/state/update-available" + + - name: Verify auto-upgrade + id: upgrade + env: + STRIPE_INSTALL_METHOD: curl + run: | + OUTPUT=$("$HOME/.stripe/bin/stripe.exe" --version 2>&1) + echo "$OUTPUT" + UPGRADED=$(echo "$OUTPUT" | grep -oE '[0-9]+\.[0-9]+\.[0-9]+' | tail -1) + if [ "$UPGRADED" != "${{ needs.resolve-versions.outputs.latest }}" ]; then + echo "Error: expected upgrade to v${{ needs.resolve-versions.outputs.latest }}, got v$UPGRADED" + exit 1 + fi + echo "Successfully upgraded to v$UPGRADED" + + - name: Verify the moved-aside binary is cleaned up + id: cleanup + env: + STRIPE_INSTALL_METHOD: curl + run: | + STRIPE="$HOME/.stripe/bin/stripe.exe" + + # Windows could not delete the moved-aside image: the process doing the + # update was still running from it, and the one it re-exec'd could not + # delete it either while its parent lived. + if [ ! -f "$STRIPE.old" ]; then + echo "Error: expected $STRIPE.old to be left behind by the update" + ls -la "$HOME/.stripe/bin" + exit 1 + fi + + # A later invocation is what clears it. The release just installed + # predates the cleanup, so put a build that has it back in place. + CGO_ENABLED=0 go build -ldflags="-s -w -X github.com/stripe/stripe-cli/pkg/version.Version=1.0.0" \ + -o "$STRIPE" ./cmd/stripe + + # Opted out, so that this also covers the cleanup running ahead of the + # opt-out check: turning auto-update off must not park the outgoing + # binary next to the new one forever. + STRIPE_NO_AUTO_UPDATE=1 "$STRIPE" --version + + if [ -f "$STRIPE.old" ]; then + echo "Error: $STRIPE.old was not cleaned up" + ls -la "$HOME/.stripe/bin" + exit 1 + fi + echo "Moved-aside binary was cleaned up" + + - name: Record failed step + id: result + if: failure() + run: | + if [ "${{ steps.build.outcome }}" = "failure" ]; then + echo "failed_step=build" >> "$GITHUB_OUTPUT" + elif [ "${{ steps.update-check.outcome }}" = "failure" ]; then + echo "failed_step=update-check" >> "$GITHUB_OUTPUT" + elif [ "${{ steps.upgrade.outcome }}" = "failure" ]; then + echo "failed_step=auto-upgrade" >> "$GITHUB_OUTPUT" + elif [ "${{ steps.cleanup.outcome }}" = "failure" ]; then + echo "failed_step=cleanup" >> "$GITHUB_OUTPUT" + fi + + notify: + runs-on: ubuntu-latest + needs: [resolve-versions, upgrade-macos, upgrade-linux, upgrade-windows] + if: always() + steps: + - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6.1.0 + if: needs.resolve-versions.result == 'success' + with: + persist-credentials: false + sparse-checkout: | + scripts/notify.sh + sparse-checkout-cone-mode: false + + - name: Build failure summary + id: summary + if: needs.resolve-versions.result == 'success' && contains(needs.*.result, 'failure') + run: | + PLATFORMS=(macOS Linux Windows) + STEPS=( + "${{ needs.upgrade-macos.outputs.failed_step }}" + "${{ needs.upgrade-linux.outputs.failed_step }}" + "${{ needs.upgrade-windows.outputs.failed_step }}" + ) + + # One clause per failing step, naming every platform that failed in it, so + # that the same breakage everywhere stays a single clause. sort -u to keep + # the order the message is built in independent of which platform failed. + MSG="" + for STEP in $(printf '%s\n' "${STEPS[@]}" | sed '/^$/d' | sort -u); do + WHERE="" + for i in "${!PLATFORMS[@]}"; do + if [ "${STEPS[$i]}" = "$STEP" ]; then + WHERE="${WHERE:+$WHERE, }${PLATFORMS[$i]}" + fi + done + MSG="${MSG:+$MSG; }Step '${STEP}' failed on ${WHERE}" + done + + echo "msg=${MSG:-Unknown failure}" >> "$GITHUB_OUTPUT" + + - name: Trigger PagerDuty alert + if: needs.resolve-versions.result == 'success' && contains(needs.*.result, 'failure') + run: bash scripts/notify.sh + env: + OVERALL_RESULT: failure + PAGERDUTY_INTEGRATION_KEY: ${{ secrets.PAGERDUTY_INTEGRATION_KEY }} + SLACK_WEBHOOK_URL: ${{ secrets.SLACK_WEBHOOK_URL }} + PAGERDUTY_DEDUP_KEY: gh-actions-stripe-cli-autoupgrade-test + PAGERDUTY_SUMMARY: "Auto-upgrade test failed: ${{ steps.summary.outputs.msg }}. Investigate: https://github.com/stripe/stripe-cli/actions/workflows/autoupgrade-test.yml" + PAGERDUTY_RESOLVE_SUMMARY: "Auto-upgrade test is passing again" + PAGERDUTY_SEVERITY: critical + DRYRUN: "true" # TODO: switch to ${{ inputs.dryrun }} once confirmed stable + + - name: Resolve PagerDuty alert + if: needs.upgrade-macos.result == 'success' && needs.upgrade-linux.result == 'success' && needs.upgrade-windows.result == 'success' + run: bash scripts/notify.sh + env: + OVERALL_RESULT: success + PAGERDUTY_INTEGRATION_KEY: ${{ secrets.PAGERDUTY_INTEGRATION_KEY }} + PAGERDUTY_DEDUP_KEY: gh-actions-stripe-cli-autoupgrade-test + PAGERDUTY_SUMMARY: "unused" + PAGERDUTY_RESOLVE_SUMMARY: "Auto-upgrade test is passing again" + PAGERDUTY_SEVERITY: critical + DRYRUN: "true" # TODO: switch to ${{ inputs.dryrun }} once confirmed stable + + - name: Skip notification (setup failure) + if: needs.resolve-versions.result == 'failure' + run: echo "Setup job (version resolution) failed — likely transient GitHub API issue. Skipping PagerDuty alert." diff --git a/.github/workflows/install-test.yml b/.github/workflows/install-test.yml index ad0273cc9..6330e900b 100644 --- a/.github/workflows/install-test.yml +++ b/.github/workflows/install-test.yml @@ -91,7 +91,44 @@ jobs: - id: install_test shell: powershell run: bash scripts/install-test.sh scoop - winget: + # winget install failures are often specific to the runner's local WinGet + # source/cache state, so a failure retries on a brand-new windows-latest + # runner rather than just rerunning the install inside the same VM. + winget-attempt-1: + runs-on: windows-latest + outputs: + install_test_result: ${{ steps.install_test.outcome }} + steps: + - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6.1.0 + with: + sparse-checkout: | + scripts/install-test.sh + sparse-checkout-cone-mode: false + persist-credentials: false + - id: install_test + shell: powershell + run: bash scripts/install-test.sh winget + + winget-attempt-2: + needs: winget-attempt-1 + if: needs.winget-attempt-1.outputs.install_test_result == 'failure' + runs-on: windows-latest + outputs: + install_test_result: ${{ steps.install_test.outcome }} + steps: + - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6.1.0 + with: + sparse-checkout: | + scripts/install-test.sh + sparse-checkout-cone-mode: false + persist-credentials: false + - id: install_test + shell: powershell + run: bash scripts/install-test.sh winget + + winget-attempt-3: + needs: winget-attempt-2 + if: always() && needs.winget-attempt-2.outputs.install_test_result == 'failure' runs-on: windows-latest outputs: install_test_result: ${{ steps.install_test.outcome }} @@ -105,6 +142,24 @@ jobs: - id: install_test shell: powershell run: bash scripts/install-test.sh winget + + winget: + needs: [winget-attempt-1, winget-attempt-2, winget-attempt-3] + if: always() + runs-on: ubuntu-latest + outputs: + install_test_result: ${{ steps.result.outputs.result }} + steps: + - id: result + run: | + if [ "${{ needs.winget-attempt-1.outputs.install_test_result }}" = "success" ] || \ + [ "${{ needs.winget-attempt-2.outputs.install_test_result }}" = "success" ] || \ + [ "${{ needs.winget-attempt-3.outputs.install_test_result }}" = "success" ]; then + echo "result=success" >> "$GITHUB_OUTPUT" + else + echo "result=failure" >> "$GITHUB_OUTPUT" + fi + docker: runs-on: ubuntu-latest # Deliberately unpinned: this canary exists to verify that the *currently diff --git a/.github/workflows/notify-developer-products-review.yml b/.github/workflows/notify-developer-products-review.yml new file mode 100644 index 000000000..8be646587 --- /dev/null +++ b/.github/workflows/notify-developer-products-review.yml @@ -0,0 +1,57 @@ +name: Notify developer-products review + +on: + pull_request_target: + types: [review_requested] + +permissions: {} + +jobs: + notify: + if: >- + github.event.requested_team.slug == 'developer-products' && + github.event.pull_request.draft == false + runs-on: ubuntu-latest + steps: + - name: Notify Slack + env: + PR_AUTHOR: ${{ github.event.pull_request.user.login }} + PR_NUMBER: ${{ github.event.pull_request.number }} + PR_TITLE: ${{ github.event.pull_request.title }} + PR_URL: ${{ github.event.pull_request.html_url }} + REPOSITORY: ${{ github.repository }} + SLACK_GROUP_ID: ${{ vars.SLACK_STRIPE_CLI_ENG_GROUP_ID }} + SLACK_WEBHOOK_URL: ${{ secrets.SLACK_STRIPE_CLI_PRS_WEBHOOK_URL }} + shell: bash + run: | + payload="$( + jq --null-input \ + --arg author "$PR_AUTHOR" \ + --arg group_id "$SLACK_GROUP_ID" \ + --arg number "$PR_NUMBER" \ + --arg repository "$REPOSITORY" \ + --arg title "$PR_TITLE" \ + --arg url "$PR_URL" \ + ' + def slack_text: + gsub("&"; "&") + | gsub("<"; "<") + | gsub(">"; ">") + | gsub("[\\r\\n]+"; " "); + + { + text: ( + " Developer Products review requested.\n" + + "*" + ($repository | slack_text) + "#" + $number + "* — " + + "<" + $url + "|" + ($title | slack_text) + ">\n" + + "Author: `" + ($author | slack_text) + "`" + ) + } + ' + )" + + curl --fail-with-body --silent --show-error \ + --request POST \ + --header 'Content-Type: application/json' \ + --data "$payload" \ + "$SLACK_WEBHOOK_URL" diff --git a/canary/listen_test.go b/canary/listen_test.go index 2cf3b6f1e..82736e0e8 100644 --- a/canary/listen_test.go +++ b/canary/listen_test.go @@ -114,7 +114,9 @@ func TestAPIListenForwardTo(t *testing.T) { "STRIPE_API_KEY": testutil.GetAPIKey(), })) - listen, err := runner.RunBackground("listen", "--forward-to", server.URL) + // --forward-to needs an explicit subscription; --all-snapshot is what + // forwarding every snapshot event is spelled as now. + listen, err := runner.RunBackground("listen", "--all-snapshot", "--forward-to", server.URL) if err != nil { fatalf(t, "Failed to start listen: %v", err) } diff --git a/cmd/stripe/main.go b/cmd/stripe/main.go index 542741210..7755db100 100644 --- a/cmd/stripe/main.go +++ b/cmd/stripe/main.go @@ -9,6 +9,7 @@ import ( goversion "github.com/hashicorp/go-version" + "github.com/stripe/stripe-cli/pkg/autoupdate" "github.com/stripe/stripe-cli/pkg/cmd" "github.com/stripe/stripe-cli/pkg/reporting" "github.com/stripe/stripe-cli/pkg/stripe" @@ -18,6 +19,16 @@ import ( const sentryDSN = "https://0e1c83fa780a5946e14bfc0f6d0a7ddd@errors.stripe.com/11762" func main() { + // Apply a pending update before anything else: this re-execs the new binary, + // so the command the user typed runs on the version they are being moved to. + autoupdate.ApplyIfPending() + + // Check for the next update after the command has run, so the network call + // never delays it. Deferred rather than called at the end of main because + // every branch below returns early, and the check has to happen on all of + // them. This still does not cover cmd.Execute's os.Exit on command failure. + defer autoupdate.CheckForUpdate() + ctx := context.Background() if stripe.TelemetryOptedOut(os.Getenv("STRIPE_CLI_TELEMETRY_OPTOUT")) || stripe.TelemetryOptedOut(os.Getenv("DO_NOT_TRACK")) { diff --git a/pkg/agentsetup/codex.go b/pkg/agentsetup/codex.go index d797e5b53..f063a2d21 100644 --- a/pkg/agentsetup/codex.go +++ b/pkg/agentsetup/codex.go @@ -12,24 +12,27 @@ import ( ) const ( - ClientCodex = "codex" - CodexBinaryName = "codex" - CodexPluginName = "stripe" - CodexMarketplace = "openai-curated" - TargetCodexPlugin = "stripe@openai-curated" - CodexDisplayName = "Codex CLI" + ClientCodex = "codex" + CodexBinaryName = "codex" + CodexPluginName = "stripe" + CodexDisplayName = "Codex CLI" codexListTimeout = 5 * time.Second ) +// Official Codex marketplaces for ChatGPT and API-key users, respectively: +// https://github.com/openai/plugins/blob/main/.agents/plugins/marketplace.json#L2 +// https://github.com/openai/plugins/blob/main/.agents/plugins/api_marketplace.json +var codexMarketplaces = [...]string{"openai-curated", "openai-api-curated"} + // RunOutputFunc runs a command and returns its standard output. It exists so // Codex detection (which shells out to `codex plugin list --json`) is testable. type RunOutputFunc func(context.Context, string, ...string) ([]byte, error) // CodexProvider detects and installs the Stripe plugin for Codex CLI. // -// Codex has a real plugin CLI, so detection runs `codex plugin list --json` and -// installation runs `codex plugin add stripe@openai-curated`. +// Detection selects the first available supported marketplace, and installation +// runs `codex plugin add stripe@` using that selection. type CodexProvider struct { Scanner Scanner RunCommand RunCommandFunc @@ -67,17 +70,24 @@ func (p CodexProvider) Detect() Status { status.ExecutablePath = binPath status.Status = StatusMissing + marketplace, err := p.marketplace(context.Background()) + if err != nil { + status.Status = StatusError + status.Error = err.Error() + return status + } + status.Plugin.ID = CodexPluginName + "@" + marketplace + ctx, cancel := context.WithTimeout(context.Background(), codexListTimeout) defer cancel() - version, ok, supportsPlugins := p.stripePluginStatus(ctx) + version, ok, supportsPlugins := p.stripePluginStatus(ctx, marketplace) if !supportsPlugins { status.Error = "upgrade Codex to enable plugin support" return status } if ok { status.Plugin.Installed = true - status.Plugin.ID = TargetCodexPlugin status.Plugin.Version = version status.Plugin.Scope = "user" status.Status = StatusInstalled @@ -86,25 +96,57 @@ func (p CodexProvider) Detect() Status { return status } +// marketplace selects the first available supported marketplace. +func (p CodexProvider) marketplace(ctx context.Context) (string, error) { + ctx, cancel := context.WithTimeout(ctx, codexListTimeout) + defer cancel() + runOutput := p.RunOutput + if runOutput == nil { + runOutput = runCommandOutput + } + out, err := runOutput(ctx, CodexBinaryName, "plugin", "marketplace", "list", "--json") + var list struct { + Marketplaces []struct { + Name string `json:"name"` + } `json:"marketplaces"` + } + if err == nil { + err = json.Unmarshal(out, &list) + } + if err != nil { + return "", errorcategory.Errorf(errorcategory.Internal, "listing Codex marketplaces: %w", err) + } + + for _, marketplace := range codexMarketplaces { + for _, available := range list.Marketplaces { + if available.Name == marketplace { + return marketplace, nil + } + } + } + return "", errorcategory.Errorf(errorcategory.Internal, "no supported Codex marketplace is available; expected %s", strings.Join(codexMarketplaces[:], " or ")) +} + // stripePluginStatus runs `codex plugin list --json` and reports whether (1) // the command is supported (supportsPlugins), and if so (2) whether the Stripe // plugin is installed and its version. When the command fails (e.g. old Codex // version without plugin support), supportsPlugins is false. -func (p CodexProvider) stripePluginStatus(ctx context.Context) (version string, installed bool, supportsPlugins bool) { +func (p CodexProvider) stripePluginStatus(ctx context.Context, marketplace string) (version string, installed bool, supportsPlugins bool) { runOutput := p.RunOutput if runOutput == nil { runOutput = runCommandOutput } - out, err := runOutput(ctx, CodexBinaryName, "plugin", "list", "--json") + // The unfiltered list can omit locally installed curated plugins. + out, err := runOutput(ctx, CodexBinaryName, "plugin", "list", "--marketplace", marketplace, "--json") if err != nil { return "", false, false } - v, ok := findCodexStripePlugin(out) + v, ok := findCodexStripePlugin(out, marketplace) return v, ok, true } func (p CodexProvider) Plan(status Status, force bool) Plan { - command := []string{CodexBinaryName, "plugin", "add", TargetCodexPlugin} + command := []string{CodexBinaryName, "plugin", "add", status.Plugin.ID} switch { case status.Status == StatusError: @@ -131,16 +173,18 @@ func (p CodexProvider) Apply(ctx context.Context, _ io.Writer, plan Plan) error if runCommand == nil { runCommand = RunCommand } + pluginID := plan.Command[len(plan.Command)-1] + _, marketplace, _ := strings.Cut(pluginID, "@") if err := runCommand(ctx, plan.Command[0], plan.Command[1:]...); err != nil { - return err + return errorcategory.Errorf(errorcategory.Internal, "could not install the Stripe plugin from %s: %w", marketplace, err) } // `codex plugin add` exits 0 even when it fails (e.g. the marketplace is not // configured), so the exit code cannot be trusted. Confirm the plugin is // actually installed before reporting success. - if _, installed, _ := p.stripePluginStatus(ctx); !installed { + if _, installed, _ := p.stripePluginStatus(ctx, marketplace); !installed { return errorcategory.Errorf(errorcategory.Internal, "codex reported success but %s is not installed; run `%s` to see the underlying error", - TargetCodexPlugin, strings.Join(plan.Command, " ")) + pluginID, strings.Join(plan.Command, " ")) } return nil } @@ -163,26 +207,26 @@ type codexInstalledPlugin struct { // findCodexStripePlugin reports whether the Stripe plugin appears in the // installed list and returns its version when available. -func findCodexStripePlugin(listJSON []byte) (string, bool) { +func findCodexStripePlugin(listJSON []byte, marketplace string) (string, bool) { var list codexPluginList if err := json.Unmarshal(listJSON, &list); err != nil { return "", false } for _, plugin := range list.Installed { - if codexPluginIsStripe(plugin) { + if codexPluginIsStripe(plugin, marketplace) { return plugin.Version, true } } return "", false } -func codexPluginIsStripe(plugin codexInstalledPlugin) bool { - if strings.EqualFold(plugin.PluginID, TargetCodexPlugin) { +func codexPluginIsStripe(plugin codexInstalledPlugin, marketplace string) bool { + if strings.EqualFold(plugin.PluginID, CodexPluginName+"@"+marketplace) { return true } return strings.EqualFold(plugin.Name, CodexPluginName) && - strings.EqualFold(plugin.Marketplace, CodexMarketplace) + strings.EqualFold(plugin.Marketplace, marketplace) } func runCommandOutput(ctx context.Context, name string, args ...string) ([]byte, error) { diff --git a/pkg/agentsetup/codex_test.go b/pkg/agentsetup/codex_test.go index 7a64f1f47..ce13ce660 100644 --- a/pkg/agentsetup/codex_test.go +++ b/pkg/agentsetup/codex_test.go @@ -3,11 +3,14 @@ package agentsetup import ( "context" "errors" + "fmt" "testing" "github.com/stretchr/testify/require" ) +const codexTestMarketplaceList = `{"marketplaces":[{"name":"openai-curated"},{"name":"openai-api-curated"}]}` + func TestScanCodex_NotDetected(t *testing.T) { provider := CodexProvider{ Scanner: Scanner{LookPath: func(string) (string, error) { return "", errors.New("missing") }}, @@ -32,7 +35,8 @@ func TestScanCodex_PluginMissing(t *testing.T) { require.True(t, status.Detected) require.Equal(t, StatusMissing, status.Status) require.False(t, status.Plugin.Installed) - require.Equal(t, Plan{Action: ActionInstall, Command: []string{"codex", "plugin", "add", TargetCodexPlugin}}, provider.Plan(status, false)) + require.Equal(t, "stripe@openai-curated", status.Plugin.ID) + require.Equal(t, Plan{Action: ActionInstall, Command: []string{"codex", "plugin", "add", "stripe@openai-curated"}}, provider.Plan(status, false)) } func TestScanCodex_PluginInstalled(t *testing.T) { @@ -43,7 +47,7 @@ func TestScanCodex_PluginInstalled(t *testing.T) { require.Equal(t, StatusInstalled, status.Status) require.True(t, status.Plugin.Installed) - require.Equal(t, TargetCodexPlugin, status.Plugin.ID) + require.Equal(t, "stripe@openai-curated", status.Plugin.ID) require.Equal(t, "3fdeeb49", status.Plugin.Version) require.Equal(t, Plan{Action: ActionNone}, provider.Plan(status, false)) } @@ -58,72 +62,230 @@ func TestScanCodex_PluginInstalledByNameAndMarketplace(t *testing.T) { require.Equal(t, "2.0.0", status.Plugin.Version) } +func TestScanCodex_APIPluginInstalled(t *testing.T) { + provider := codexTestProvider(`{"installed":[{"pluginId":"stripe@openai-api-curated","version":"1.0.0"}]}`, nil, nil) + runOutput := provider.RunOutput + provider.RunOutput = func(ctx context.Context, name string, args ...string) ([]byte, error) { + if args[1] == "marketplace" { + return []byte(`{"marketplaces":[{"name":"openai-api-curated"}]}`), nil + } + return runOutput(ctx, name, args...) + } + status := provider.Detect() + + require.Equal(t, StatusInstalled, status.Status) + require.True(t, status.Plugin.Installed) + require.Equal(t, "stripe@openai-api-curated", status.Plugin.ID) + require.Equal(t, "1.0.0", status.Plugin.Version) + require.Equal(t, Plan{Action: ActionNone}, provider.Plan(status, false)) + require.Equal(t, Plan{Action: ActionReinstall, Command: []string{"codex", "plugin", "add", "stripe@openai-api-curated"}}, provider.Plan(status, true)) +} + +func TestScanCodex_DoesNotFallBackAfterLookupError(t *testing.T) { + provider := codexTestProvider("", nil, nil) + var marketplaces []string + provider.RunOutput = func(_ context.Context, _ string, args ...string) ([]byte, error) { + if args[1] == "marketplace" { + return []byte(codexTestMarketplaceList), nil + } + marketplace := args[3] + marketplaces = append(marketplaces, marketplace) + if marketplace == "openai-curated" { + return nil, errors.New("marketplace unavailable") + } + return []byte(`{"installed":[{"pluginId":"stripe@openai-api-curated","version":"1.0.0"}]}`), nil + } + + status := provider.Detect() + + require.Equal(t, []string{"openai-curated"}, marketplaces) + require.Equal(t, StatusMissing, status.Status) + require.Equal(t, "stripe@openai-curated", status.Plugin.ID) + require.False(t, status.Plugin.Installed) + require.Contains(t, status.Error, "upgrade Codex") +} + func TestScanCodex_OldVersionWithoutPluginSupport(t *testing.T) { provider := codexTestProvider("", errors.New("unrecognized subcommand 'plugin'"), nil) status := provider.Detect() - // Old Codex shows as detected but with an error hint — the TUI renders - // it as disabled (visible but not selectable). require.True(t, status.Detected) - require.Equal(t, StatusMissing, status.Status) - require.Contains(t, status.Error, "upgrade Codex") + require.Equal(t, StatusError, status.Status) + require.Contains(t, status.Error, "listing Codex marketplaces") + require.Equal(t, Plan{Action: ActionNone}, provider.Plan(status, true)) } -func TestCodexApply_RunsAddCommandAndVerifies(t *testing.T) { - var gotName string - var gotArgs []string - installed := false +// An install failure, including exit zero without installing, must not cause +// an attempt from a different marketplace, even when both are available. +func TestCodexApply_DoesNotFallBack(t *testing.T) { + for _, installErr := range []error{nil, errors.New("marketplace unavailable")} { + var attempts []string + provider := codexTestProvider(`{"installed":[]}`, nil, func(_ context.Context, _ string, args ...string) error { + attempts = append(attempts, args[2]) + return installErr + }) - provider := CodexProvider{ - Scanner: Scanner{LookPath: func(string) (string, error) { return "/usr/local/bin/codex", nil }}, - RunCommand: func(_ context.Context, name string, args ...string) error { - gotName = name - gotArgs = args - installed = true // simulate a successful add - return nil + status := provider.Detect() + plan := provider.Plan(status, false) + err := provider.Apply(context.Background(), nil, plan) + + require.Error(t, err) + require.Equal(t, []string{"stripe@openai-curated"}, attempts) + require.Contains(t, err.Error(), "openai-curated") + require.NotContains(t, err.Error(), "openai-api-curated") + if installErr != nil { + require.ErrorIs(t, err, installErr) + } else { + require.ErrorContains(t, err, "codex reported success but stripe@openai-curated is not installed") + } + } +} + +func TestCodexSetup_AvailableMarketplaces(t *testing.T) { + for _, tt := range []struct { + name string + listOutput string + marketplace string + }{ + { + name: "supported order, ignoring unrelated and duplicate entries", + listOutput: `{"marketplaces":[{"name":"other"},{"name":"openai-api-curated"},{"name":"openai-curated"},{"name":"openai-curated"}]}`, + marketplace: "openai-curated", }, - RunOutput: func(context.Context, string, ...string) ([]byte, error) { - if installed { - return []byte(`{"installed":[{"pluginId":"stripe@openai-curated","name":"stripe","marketplaceName":"openai-curated","version":"1.0.0"}]}`), nil - } - return []byte(`{"installed":[]}`), nil + { + name: "ChatGPT marketplace only", + listOutput: `{"marketplaces":[{"name":"openai-curated"}]}`, + marketplace: "openai-curated", }, - } + { + name: "API marketplace only", + listOutput: `{"marketplaces":[{"name":"openai-api-curated"}]}`, + marketplace: "openai-api-curated", + }, + } { + t.Run(tt.name, func(t *testing.T) { + var queried, attempts []string + discoveries := 0 + installed := "" + provider := codexTestProvider("", nil, func(_ context.Context, name string, args ...string) error { + require.Equal(t, "codex", name) + require.Equal(t, []string{"plugin", "add"}, args[:2]) + attempts = append(attempts, args[2]) + installed = args[2] + return nil + }) + provider.RunOutput = func(ctx context.Context, name string, args ...string) ([]byte, error) { + require.Equal(t, "codex", name) + require.NoError(t, ctx.Err()) + if args[1] == "marketplace" { + require.Equal(t, []string{"plugin", "marketplace", "list", "--json"}, args) + discoveries++ + return []byte(tt.listOutput), nil + } + require.Equal(t, []string{"plugin", "list", "--marketplace"}, args[:3]) + require.Equal(t, "--json", args[4]) + queried = append(queried, args[3]) + if installed == "stripe@"+args[3] { + return []byte(fmt.Sprintf(`{"installed":[{"pluginId":%q,"version":"1.0.0"}]}`, installed)), nil + } + return []byte(`{"installed":[]}`), nil + } - status := provider.Detect() - plan := provider.Plan(status, false) - err := provider.Apply(context.Background(), nil, plan) + status := provider.Detect() + require.Equal(t, StatusMissing, status.Status) + require.Empty(t, status.Error) + require.Equal(t, []string{tt.marketplace}, queried) + pluginID := "stripe@" + tt.marketplace + require.Equal(t, pluginID, status.Plugin.ID) + plan := provider.Plan(status, false) + require.Equal(t, Plan{Action: ActionInstall, Command: []string{"codex", "plugin", "add", pluginID}}, plan) + require.NoError(t, provider.Apply(context.Background(), nil, plan)) + require.Equal(t, []string{pluginID}, attempts) + require.Equal(t, 1, discoveries, "Plan and Apply must reuse the selection from Detect") - require.NoError(t, err) - require.Equal(t, "codex", gotName) - require.Equal(t, []string{"plugin", "add", TargetCodexPlugin}, gotArgs) + status = provider.Detect() + require.Equal(t, StatusInstalled, status.Status) + require.Equal(t, pluginID, status.Plugin.ID) + require.NoError(t, provider.Apply(context.Background(), nil, provider.Plan(status, false))) + require.Equal(t, []string{pluginID}, attempts, "skip an installed plugin") + require.NoError(t, provider.Apply(context.Background(), nil, provider.Plan(status, true))) + require.Equal(t, []string{pluginID, pluginID}, attempts, "force reinstalls from the detected marketplace") + require.Equal(t, 2, discoveries, "reinstall must also reuse the selection from Detect") + }) + } } -// TestCodexApply_FailsWhenExitZeroButNotInstalled covers the real-world case -// where `codex plugin add` prints an error but exits 0. Apply must not report -// success when the plugin is still not present afterward. -func TestCodexApply_FailsWhenExitZeroButNotInstalled(t *testing.T) { - provider := codexTestProvider(`{"installed":[]}`, nil, func(context.Context, string, ...string) error { - return nil // add "succeeds" (exit 0) but installs nothing - }) - - status := provider.Detect() - plan := provider.Plan(status, false) - err := provider.Apply(context.Background(), nil, plan) +func TestScanCodex_MarketplaceDiscoveryFailures(t *testing.T) { + for _, tt := range []struct { + name string + listOutput string + listErr error + wantError string + }{ + { + name: "no marketplaces", + listOutput: `{"marketplaces":[]}`, + wantError: "no supported Codex marketplace is available", + }, + { + name: "only unrelated marketplaces", + listOutput: `{"marketplaces":[{"name":"other"}]}`, + wantError: "no supported Codex marketplace is available", + }, + { + name: "listing fails", + listErr: errors.New("marketplace listing failed"), + wantError: "listing Codex marketplaces: marketplace listing failed", + }, + { + name: "listing times out", + listErr: context.DeadlineExceeded, + wantError: "listing Codex marketplaces: context deadline exceeded", + }, + { + name: "invalid JSON", + listOutput: "not JSON", + wantError: "listing Codex marketplaces: invalid character", + }, + } { + t.Run(tt.name, func(t *testing.T) { + provider := codexTestProvider("", nil, func(context.Context, string, ...string) error { + t.Fatal("must not install without a detected marketplace") + return nil + }) + discoveries := 0 + provider.RunOutput = func(_ context.Context, _ string, args ...string) ([]byte, error) { + require.Equal(t, []string{"plugin", "marketplace", "list", "--json"}, args) + discoveries++ + return []byte(tt.listOutput), tt.listErr + } - require.Error(t, err) - require.Contains(t, err.Error(), "is not installed") + status := provider.Detect() + require.Equal(t, StatusError, status.Status) + require.Contains(t, status.Error, tt.wantError) + require.Empty(t, status.Plugin.ID) + for _, force := range []bool{false, true} { + plan := provider.Plan(status, force) + require.Equal(t, Plan{Action: ActionNone}, plan) + require.NoError(t, provider.Apply(context.Background(), nil, plan)) + } + require.Equal(t, 1, discoveries) + }) + } } func codexTestProvider(listOutput string, listErr error, runCommand RunCommandFunc) CodexProvider { return CodexProvider{ Scanner: Scanner{LookPath: func(string) (string, error) { return "/usr/local/bin/codex", nil }}, RunCommand: runCommand, - RunOutput: func(context.Context, string, ...string) ([]byte, error) { + RunOutput: func(_ context.Context, _ string, args ...string) ([]byte, error) { if listErr != nil { return nil, listErr } + if args[1] == "marketplace" { + return []byte(codexTestMarketplaceList), nil + } return []byte(listOutput), nil }, } diff --git a/pkg/agentsetup/grok.go b/pkg/agentsetup/grok.go new file mode 100644 index 000000000..795f4df42 --- /dev/null +++ b/pkg/agentsetup/grok.go @@ -0,0 +1,157 @@ +package agentsetup + +import ( + "context" + "encoding/json" + "io" + "strings" + "time" + + "github.com/stripe/stripe-cli/pkg/errorcategory" +) + +const ( + ClientGrok = "grok" + GrokBinaryName = "grok" + GrokPluginName = "stripe" + GrokDisplayName = "Grok" + + grokListTimeout = 5 * time.Second +) + +// GrokProvider detects and installs the Stripe plugin for Grok Build (xAI). +type GrokProvider struct { + Scanner Scanner + RunCommand RunCommandFunc + RunOutput RunOutputFunc +} + +// NewGrokProvider returns a Grok Build setup provider. +func NewGrokProvider(scanner Scanner, runCommand RunCommandFunc) Provider { + if runCommand == nil { + runCommand = RunCommand + } + return GrokProvider{ + Scanner: scanner, + RunCommand: runCommand, + RunOutput: runCommandOutput, + } +} + +func (p GrokProvider) ID() string { return ClientGrok } + +func (p GrokProvider) Detect() Status { + s := p.Scanner.withDefaults() + + status := Status{ + Client: ClientGrok, + DisplayName: GrokDisplayName, + Status: StatusNotDetected, + } + + binPath, err := s.LookPath(GrokBinaryName) + if err != nil { + return status + } + status.Detected = true + status.ExecutablePath = binPath + status.Status = StatusMissing + + ctx, cancel := context.WithTimeout(context.Background(), grokListTimeout) + defer cancel() + + plugin, ok, supportsPlugins := p.stripePluginStatus(ctx) + if !supportsPlugins { + status.Error = "upgrade Grok Build to enable plugin support" + return status + } + if ok { + status.Plugin.Installed = true + status.Plugin.ID = plugin.Name + status.Plugin.Version = plugin.Version + status.Plugin.StatePath = plugin.Path + status.Status = StatusInstalled + } + + return status +} + +// stripePluginStatus runs `grok plugin list --json` and reports whether the +// Stripe plugin is installed. When the command fails (e.g. an old Grok +// version without plugin support), supportsPlugins is false. +func (p GrokProvider) stripePluginStatus(ctx context.Context) (plugin grokInstalledPlugin, installed bool, supportsPlugins bool) { + runOutput := p.RunOutput + if runOutput == nil { + runOutput = runCommandOutput + } + out, err := runOutput(ctx, GrokBinaryName, "plugin", "list", "--json") + if err != nil { + return grokInstalledPlugin{}, false, false + } + plugin, ok := findGrokStripePlugin(out) + return plugin, ok, true +} + +// grokInstalledPlugin is an entry in `grok plugin list --json` output +type grokInstalledPlugin struct { + Status string `json:"status"` + Name string `json:"name"` + RepoKey string `json:"repo_key"` + Version string `json:"version"` + Path string `json:"path"` + Source string `json:"source"` + Marketplace string `json:"marketplace"` +} + +// findGrokStripePlugin reports whether the Stripe plugin appears in the +// output of `grok plugin list --json`. +func findGrokStripePlugin(listJSON []byte) (grokInstalledPlugin, bool) { + var plugins []grokInstalledPlugin + if err := json.Unmarshal(listJSON, &plugins); err != nil { + return grokInstalledPlugin{}, false + } + + for _, plugin := range plugins { + if grokPluginIsStripe(plugin) { + return plugin, true + } + } + return grokInstalledPlugin{}, false +} + +func grokPluginIsStripe(plugin grokInstalledPlugin) bool { + return strings.EqualFold(plugin.Name, GrokPluginName) && strings.EqualFold(plugin.Status, "installed") +} + +func (p GrokProvider) Plan(status Status, force bool) Plan { + switch { + case status.Status == StatusError: + return Plan{Action: ActionNone} + case !status.Detected: + return Plan{Action: ActionNone} + case status.Plugin.Installed && force: + // `grok plugin install` is idempotent when already installed — it + // prints "Plugin stripe is already installed ... Run `grok plugin + // update stripe` to update it" rather than reinstalling, so a forced + // refresh has to go through `update` instead. + return Plan{Action: ActionReinstall, Command: []string{GrokBinaryName, "plugin", "update", GrokPluginName}} + case status.Plugin.Installed: + return Plan{Action: ActionNone} + default: + return Plan{Action: ActionInstall, Command: []string{GrokBinaryName, "plugin", "install", GrokPluginName, "--trust"}} + } +} + +// Apply installs (or updates) the Stripe Grok plugin. Unlike Codex's +// `plugin add` (see codex.go), `grok plugin install`/`update` exit non-zero +// on failure and 0 on success, so the exit code can be trusted without a +// post-install re-check. +func (p GrokProvider) Apply(ctx context.Context, _ io.Writer, plan Plan) error { + if plan.Action == ActionNone { + return nil + } + if len(plan.Command) == 0 { + return errorcategory.Errorf(errorcategory.Internal, "missing command for %s action", plan.Action) + } + return p.RunCommand(ctx, plan.Command[0], plan.Command[1:]...) +} diff --git a/pkg/agentsetup/grok_test.go b/pkg/agentsetup/grok_test.go new file mode 100644 index 000000000..71a808573 --- /dev/null +++ b/pkg/agentsetup/grok_test.go @@ -0,0 +1,122 @@ +package agentsetup + +import ( + "context" + "errors" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestGrok_NotDetected(t *testing.T) { + provider := GrokProvider{ + Scanner: Scanner{LookPath: func(string) (string, error) { return "", errors.New("missing") }}, + RunOutput: func(context.Context, string, ...string) ([]byte, error) { + t.Fatal("plugin list should not run when Grok is not detected") + return nil, nil + }, + } + + status := provider.Detect() + + require.Equal(t, ClientGrok, status.Client) + require.Equal(t, "Grok", status.DisplayName) + require.False(t, status.Detected) + require.Equal(t, StatusNotDetected, status.Status) +} + +func TestGrok_PluginMissing(t *testing.T) { + provider := grokTestProvider(`[]`, nil, nil) + + status := provider.Detect() + + require.True(t, status.Detected) + require.Equal(t, StatusMissing, status.Status) + require.False(t, status.Plugin.Installed) + require.Equal(t, + Plan{Action: ActionInstall, Command: []string{"grok", "plugin", "install", GrokPluginName, "--trust"}}, + provider.Plan(status, false)) +} + +func TestGrok_PluginInstalled(t *testing.T) { + provider := grokTestProvider(`[ + {"status":"installed","name":"stripe","repo_key":"plugin-760cfec9","version":"0.7.1", + "path":"/Users/x/.grok/installed-plugins/plugin-760cfec9", + "source":"https://github.com/stripe/ai.git","marketplace":"xAI Official"} + ]`, nil, nil) + + status := provider.Detect() + + require.Equal(t, StatusInstalled, status.Status) + require.True(t, status.Plugin.Installed) + require.Equal(t, "stripe", status.Plugin.ID) + require.Equal(t, "0.7.1", status.Plugin.Version) + require.Equal(t, "/Users/x/.grok/installed-plugins/plugin-760cfec9", status.Plugin.StatePath) + require.Equal(t, Plan{Action: ActionNone}, provider.Plan(status, false)) +} + +func TestGrok_OldVersionWithoutPluginSupport(t *testing.T) { + provider := grokTestProvider("", errors.New("unrecognized subcommand 'plugin'"), nil) + + status := provider.Detect() + + require.True(t, status.Detected) + require.Equal(t, StatusMissing, status.Status) + require.Contains(t, status.Error, "upgrade Grok Build") +} + +func TestGrok_PlanReinstallWhenForced(t *testing.T) { + status := Status{Detected: true, Plugin: PluginStatus{Installed: true}} + provider := GrokProvider{} + + plan := provider.Plan(status, true) + + require.Equal(t, ActionReinstall, plan.Action) + require.Equal(t, []string{"grok", "plugin", "update", GrokPluginName}, plan.Command) +} + +func TestGrokApply_RunsInstallCommand(t *testing.T) { + var gotName string + var gotArgs []string + + provider := GrokProvider{ + RunCommand: func(_ context.Context, name string, args ...string) error { + gotName = name + gotArgs = args + return nil + }, + } + + plan := Plan{Action: ActionInstall, Command: []string{"grok", "plugin", "install", GrokPluginName, "--trust"}} + err := provider.Apply(context.Background(), nil, plan) + + require.NoError(t, err) + require.Equal(t, "grok", gotName) + require.Equal(t, []string{"plugin", "install", GrokPluginName, "--trust"}, gotArgs) +} + +func TestGrokApply_NoneIsNoop(t *testing.T) { + provider := GrokProvider{ + RunCommand: func(context.Context, string, ...string) error { + t.Fatal("RunCommand should not run for ActionNone") + return nil + }, + } + + err := provider.Apply(context.Background(), nil, Plan{Action: ActionNone}) + + require.NoError(t, err) +} + +func grokTestProvider(listOutput string, listErr error, runCommand RunCommandFunc) GrokProvider { + return GrokProvider{ + Scanner: Scanner{LookPath: func(string) (string, error) { return "/usr/local/bin/grok", nil }}, + RunCommand: runCommand, + RunOutput: func(context.Context, string, ...string) ([]byte, error) { + if listErr != nil { + return nil, listErr + } + return []byte(listOutput), nil + }, + } +} diff --git a/pkg/agentsetup/provider.go b/pkg/agentsetup/provider.go index b06f5a82c..a1297e151 100644 --- a/pkg/agentsetup/provider.go +++ b/pkg/agentsetup/provider.go @@ -36,10 +36,12 @@ func DefaultProviders() map[string]Provider { claude := NewClaudeProvider(scanner, RunCommand) cursor := NewCursorProvider(scanner, RunCommand) codex := NewCodexProvider(scanner, RunCommand) + grok := NewGrokProvider(scanner, RunCommand) return map[string]Provider{ claude.ID(): claude, cursor.ID(): cursor, codex.ID(): codex, + grok.ID(): grok, } } diff --git a/pkg/autoupdate/checker.go b/pkg/autoupdate/checker.go new file mode 100644 index 000000000..a79e1814d --- /dev/null +++ b/pkg/autoupdate/checker.go @@ -0,0 +1,328 @@ +package autoupdate + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "io" + "net/http" + "os" + "path/filepath" + "runtime" + "strconv" + "strings" + "time" + + "github.com/google/go-github/v72/github" + semver "github.com/hashicorp/go-version" + log "github.com/sirupsen/logrus" + + "github.com/stripe/stripe-cli/pkg/version" +) + +const httpTimeout = 10 * time.Second + +const checkInterval = 24 * time.Hour + +// UpdateMarker represents a staged update ready to be applied. +type UpdateMarker struct { + Version string `json:"version"` + DownloadURL string `json:"download_url"` + Checksum string `json:"checksum"` + ReleaseNotes string `json:"release_notes"` +} + +// CheckForUpdate checks for a newer CLI version and writes a marker file +// if an update is available. Skips major version changes. This is called +// synchronously after command execution, rate-limited to once per day. +func CheckForUpdate() { + defer func() { + if r := recover(); r != nil { + log.Debugf("autoupdate check panicked: %v", r) + } + }() + + if !shouldCheck() { + return + } + + latest, url, checksum, releaseNotes := fetchLatestRelease() + if latest == "" { + return + } + + current := strings.TrimPrefix(version.Version, "v") + latestClean := strings.TrimPrefix(latest, "v") + + if current == latestClean { + recordLastCheck() + return + } + + if isMajorVersionChange(current, latestClean) { + log.Debugf("autoupdate: skipping major version change %s → %s", current, latestClean) + recordLastCheck() + return + } + + WriteMarker(UpdateMarker{ + Version: latestClean, + DownloadURL: url, + Checksum: checksum, + ReleaseNotes: releaseNotes, + }) + sendTelemetryEvent("Auto-Update Available", fmt.Sprintf("from=%s to=%s", current, latestClean)) +} + +func isMajorVersionChange(current, latest string) bool { + cur, err := semver.NewVersion(current) + if err != nil { + return false + } + lat, err := semver.NewVersion(latest) + if err != nil { + return false + } + return cur.Segments()[0] != lat.Segments()[0] +} + +func shouldCheck() bool { + if version.Version == "master" { + return false + } + if IsOptedOut() { + return false + } + if !IsCurlInstall() { + return false + } + + stateDir := GetStateDir() + if stateDir == "" { + return false + } + + lastCheckFile := filepath.Join(stateDir, "last_update_check") + data, err := os.ReadFile(lastCheckFile) + if err != nil { + return true + } + + ts, err := strconv.ParseInt(strings.TrimSpace(string(data)), 10, 64) + if err != nil { + return true + } + + return time.Since(time.Unix(ts, 0)) >= checkInterval +} + +func fetchLatestRelease() (ver string, downloadURL string, checksum string, releaseNotes string) { + ctx, cancel := context.WithTimeout(context.Background(), httpTimeout) + defer cancel() + + client := github.NewClient(nil) + release, _, err := client.Repositories.GetLatestRelease(ctx, "stripe", "stripe-cli") + if err != nil { + log.Debug("autoupdate: failed to fetch latest release: ", err) + return "", "", "", "" + } + + ver = release.GetTagName() + releaseNotes = release.GetBody() + assetName := binaryAssetName(strings.TrimPrefix(ver, "v")) + checksumAsset := checksumAssetName() + + var binaryURL, checksumURL string + for _, asset := range release.Assets { + name := asset.GetName() + if name == assetName { + binaryURL = asset.GetBrowserDownloadURL() + } + if name == checksumAsset { + checksumURL = asset.GetBrowserDownloadURL() + } + } + + if binaryURL == "" { + log.Debug("autoupdate: binary asset not found: ", assetName) + return "", "", "", "" + } + + if checksumURL != "" { + checksum = fetchChecksumForAsset(checksumURL, assetName) + } + + return ver, binaryURL, checksum, releaseNotes +} + +func fetchChecksumForAsset(checksumURL, assetName string) string { + ctx, cancel := context.WithTimeout(context.Background(), httpTimeout) + defer cancel() + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, checksumURL, nil) + if err != nil { + log.Debug("autoupdate: failed to create checksum request: ", err) + return "" + } + + resp, err := http.DefaultClient.Do(req) + if err != nil { + log.Debug("autoupdate: failed to fetch checksums: ", err) + return "" + } + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + if err != nil { + return "" + } + + for _, line := range strings.Split(string(body), "\n") { + parts := strings.Fields(line) + if len(parts) == 2 && parts[1] == assetName { + return parts[0] + } + } + return "" +} + +func binaryAssetName(ver string) string { + return binaryAssetNameFor(ver, runtime.GOOS, runtime.GOARCH) +} + +// binaryAssetNameFor is the release archive published for a platform. +// +// The names come from the archive templates in .goreleaser/, which do not use +// the Go names for either half: darwin is published as "mac-os" and amd64 as +// "x86_64". Passing runtime.GOOS straight through asks for an asset that does +// not exist, and a missing asset stops the update silently. +func binaryAssetNameFor(ver, goos, goarch string) string { + osLabel := goos + ext := "tar.gz" + + switch goos { + case "darwin": + osLabel = "mac-os" + case "windows": + ext = "zip" + } + + return fmt.Sprintf("stripe_%s_%s_%s.%s", ver, osLabel, archAssetLabel(goos, goarch), ext) +} + +func archAssetLabel(goos, goarch string) string { + switch goarch { + case "amd64": + return "x86_64" + case "386": + return "i386" + case "arm64": + // .goreleaser/windows.yml builds amd64 and 386 only. Windows on ARM runs + // the x64 binary under emulation, so that is the archive to fetch. + if goos == "windows" { + return "x86_64" + } + + return "arm64" + default: + return goarch + } +} + +func checksumAssetName() string { + switch runtime.GOOS { + case "darwin": + return "stripe-mac-checksums.txt" + case "linux": + return "stripe-linux-checksums.txt" + case "windows": + return "stripe-windows-checksums.txt" + default: + return "" + } +} + +// WriteMarker writes an update marker to the state directory. +func WriteMarker(m UpdateMarker) { + stateDir := GetStateDir() + if stateDir == "" { + return + } + + if err := os.MkdirAll(stateDir, 0755); err != nil { + return + } + + content, err := json.Marshal(m) + if err != nil { + return + } + + markerPath := filepath.Join(stateDir, "update-available") + _ = os.WriteFile(markerPath, content, 0644) + + recordLastCheck() +} + +func recordLastCheck() { + stateDir := GetStateDir() + if stateDir == "" { + return + } + if err := os.MkdirAll(stateDir, 0755); err != nil { + return + } + now := strconv.FormatInt(time.Now().Unix(), 10) + _ = os.WriteFile(filepath.Join(stateDir, "last_update_check"), []byte(now), 0644) +} + +// ReadMarker reads a pending update marker, or returns nil if none exists. +func ReadMarker() *UpdateMarker { + stateDir := GetStateDir() + if stateDir == "" { + return nil + } + + data, err := os.ReadFile(filepath.Join(stateDir, "update-available")) + if err != nil { + return nil + } + + var m UpdateMarker + if err := json.Unmarshal(data, &m); err != nil { + return nil + } + return &m +} + +// ClearMarker removes the pending update marker. +func ClearMarker() { + stateDir := GetStateDir() + if stateDir == "" { + return + } + _ = os.Remove(filepath.Join(stateDir, "update-available")) +} + +// VerifyChecksum verifies the SHA256 checksum of a file. +func VerifyChecksum(filePath, expected string) bool { + if expected == "" { + return true + } + + f, err := os.Open(filePath) + if err != nil { + return false + } + defer f.Close() + + h := sha256.New() + if _, err := io.Copy(h, f); err != nil { + return false + } + + actual := hex.EncodeToString(h.Sum(nil)) + return strings.EqualFold(actual, expected) +} diff --git a/pkg/autoupdate/checker_test.go b/pkg/autoupdate/checker_test.go new file mode 100644 index 000000000..3ebc8d7ee --- /dev/null +++ b/pkg/autoupdate/checker_test.go @@ -0,0 +1,163 @@ +package autoupdate + +import ( + "os" + "path/filepath" + "runtime" + "strconv" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestIsMajorVersionChange(t *testing.T) { + tests := []struct { + current string + latest string + expected bool + }{ + {"1.23.0", "1.24.0", false}, + {"1.23.0", "1.23.1", false}, + {"1.99.0", "2.0.0", true}, + {"2.0.0", "1.99.0", true}, + {"1.0.0", "1.0.0", false}, + {"1.0.0-beta", "1.0.0", false}, + {"invalid", "1.0.0", false}, + } + + for _, tt := range tests { + t.Run(tt.current+"→"+tt.latest, func(t *testing.T) { + assert.Equal(t, tt.expected, isMajorVersionChange(tt.current, tt.latest)) + }) + } +} + +// The expected names are the ones the .goreleaser archive templates produce. An +// asset name that does not match one of them is not a download that fails +// loudly — fetchLatestRelease finds no matching asset and gives up silently, so +// auto-update simply never happens on that platform. +func TestBinaryAssetNameFor(t *testing.T) { + tests := []struct { + goos string + goarch string + expected string + }{ + {"darwin", "amd64", "stripe_1.24.0_mac-os_x86_64.tar.gz"}, + {"darwin", "arm64", "stripe_1.24.0_mac-os_arm64.tar.gz"}, + {"linux", "amd64", "stripe_1.24.0_linux_x86_64.tar.gz"}, + {"linux", "arm64", "stripe_1.24.0_linux_arm64.tar.gz"}, + {"linux", "386", "stripe_1.24.0_linux_i386.tar.gz"}, + {"windows", "amd64", "stripe_1.24.0_windows_x86_64.zip"}, + {"windows", "386", "stripe_1.24.0_windows_i386.zip"}, + // No arm64 Windows build is published; that machine runs the x64 one. + {"windows", "arm64", "stripe_1.24.0_windows_x86_64.zip"}, + } + + for _, tt := range tests { + t.Run(tt.goos+"/"+tt.goarch, func(t *testing.T) { + assert.Equal(t, tt.expected, binaryAssetNameFor("1.24.0", tt.goos, tt.goarch)) + }) + } +} + +func TestBinaryAssetNameUsesTheRunningPlatform(t *testing.T) { + assert.Equal(t, binaryAssetNameFor("1.24.0", runtime.GOOS, runtime.GOARCH), binaryAssetName("1.24.0")) +} + +func TestChecksumAssetName(t *testing.T) { + // The checksums file the release publishes for this platform, which is where + // the archive's expected digest is read from. + expected := map[string]string{ + "darwin": "stripe-mac-checksums.txt", + "linux": "stripe-linux-checksums.txt", + "windows": "stripe-windows-checksums.txt", + }[runtime.GOOS] + + assert.Equal(t, expected, checksumAssetName()) +} + +func TestMarkerReadWrite(t *testing.T) { + tmpDir := t.TempDir() + original := GetStateDirFn + defer func() { GetStateDirFn = original }() + GetStateDirFn = func() string { return tmpDir } + + m := UpdateMarker{ + Version: "1.24.0", + DownloadURL: "https://example.com/stripe.tar.gz", + Checksum: "abc123", + } + + WriteMarker(m) + + got := ReadMarker() + require.NotNil(t, got) + assert.Equal(t, "1.24.0", got.Version) + assert.Equal(t, "https://example.com/stripe.tar.gz", got.DownloadURL) + assert.Equal(t, "abc123", got.Checksum) + + ClearMarker() + assert.Nil(t, ReadMarker()) +} + +func TestMarkerReadWriteWithReleaseNotes(t *testing.T) { + tmpDir := t.TempDir() + original := GetStateDirFn + defer func() { GetStateDirFn = original }() + GetStateDirFn = func() string { return tmpDir } + + notes := "## Changes\n\n- Added one thing\n- Fixed another thing" + WriteMarker(UpdateMarker{ + Version: "1.24.0", + DownloadURL: "https://example.com/stripe.tar.gz", + Checksum: "abc123", + ReleaseNotes: notes, + }) + + got := ReadMarker() + require.NotNil(t, got) + assert.Equal(t, notes, got.ReleaseNotes) +} + +func TestReadMarkerWithoutReleaseNotes(t *testing.T) { + tmpDir := t.TempDir() + original := GetStateDirFn + defer func() { GetStateDirFn = original }() + GetStateDirFn = func() string { return tmpDir } + + marker := `{"version":"1.24.0","download_url":"https://example.com/stripe.tar.gz","checksum":"abc123"}` + require.NoError(t, os.WriteFile(filepath.Join(tmpDir, "update-available"), []byte(marker), 0644)) + + got := ReadMarker() + require.NotNil(t, got) + assert.Empty(t, got.ReleaseNotes) +} + +func TestRecordLastCheck(t *testing.T) { + tmpDir := t.TempDir() + original := GetStateDirFn + defer func() { GetStateDirFn = original }() + GetStateDirFn = func() string { return tmpDir } + + recordLastCheck() + + data, err := os.ReadFile(filepath.Join(tmpDir, "last_update_check")) + require.NoError(t, err) + + ts, err := strconv.ParseInt(string(data), 10, 64) + require.NoError(t, err) + assert.WithinDuration(t, time.Now(), time.Unix(ts, 0), 2*time.Second) +} + +func TestVerifyChecksum(t *testing.T) { + tmpFile := filepath.Join(t.TempDir(), "testfile") + os.WriteFile(tmpFile, []byte("hello world\n"), 0644) + + // sha256 of "hello world\n" + assert.True(t, VerifyChecksum(tmpFile, "a948904f2f0f479b8f8197694b30184b0d2ed1c1cd2a1ec0fb85d299a192a447")) + assert.False(t, VerifyChecksum(tmpFile, "0000000000000000000000000000000000000000000000000000000000000000")) + // Empty expected = skip verification + assert.True(t, VerifyChecksum(tmpFile, "")) +} diff --git a/pkg/autoupdate/config.go b/pkg/autoupdate/config.go new file mode 100644 index 000000000..3b11307ea --- /dev/null +++ b/pkg/autoupdate/config.go @@ -0,0 +1,101 @@ +// Package autoupdate implements automatic version updates for curl-installed Stripe CLI binaries. +package autoupdate + +import ( + "os" + "path/filepath" + "strings" + + "github.com/mitchellh/go-homedir" + "github.com/spf13/viper" +) + +// IsOptedOut reports whether the user has disabled auto-update. +func IsOptedOut() bool { + if os.Getenv("STRIPE_NO_AUTO_UPDATE") != "" { + return true + } + + configFolder := getConfigFolder() + configFile := filepath.Join(configFolder, "config.toml") + + v := viper.New() + v.SetConfigType("toml") + v.SetConfigFile(configFile) + + if err := v.ReadInConfig(); err != nil { + return false + } + + // Both spellings count. The install script tells users to write a [settings] + // table; the message the CLI prints when it updates tells them to write a + // top-level auto_update, which is also where every other CLI setting lives in + // config.toml. Someone who followed either has said what they want. + for _, key := range []string{"auto_update", "settings.auto_update"} { + if v.IsSet(key) && !v.GetBool(key) { + return true + } + } + + return false +} + +// IsCurlInstall reports whether the current binary was installed via curl (lives in ~/.stripe/bin/). +func IsCurlInstall() bool { + if method := os.Getenv("STRIPE_INSTALL_METHOD"); method != "" { + return method == "curl" + } + + exe, err := os.Executable() + if err != nil { + return false + } + + exe, err = filepath.EvalSymlinks(exe) + if err != nil { + return false + } + + home, err := homedir.Dir() + if err != nil { + return false + } + + stripeBinDir := filepath.Join(home, ".stripe", "bin") + stripeBinDir, err = filepath.EvalSymlinks(stripeBinDir) + if err != nil { + return false + } + + exeLower := strings.ToLower(filepath.ToSlash(exe)) + expectedLower := strings.ToLower(filepath.ToSlash(stripeBinDir)) + + return strings.HasPrefix(exeLower, expectedLower) +} + +func getConfigFolder() string { + if xdg := os.Getenv("XDG_CONFIG_HOME"); xdg != "" { + return filepath.Join(xdg, "stripe") + } + home, err := homedir.Dir() + if err != nil { + return "" + } + return filepath.Join(home, ".config", "stripe") +} + +// GetStateDirFn is the active implementation; tests can override it. +var GetStateDirFn = getStateDirDefault + +// GetStateDir returns the path to the autoupdate state directory. +func GetStateDir() string { + return GetStateDirFn() +} + +func getStateDirDefault() string { + home, err := homedir.Dir() + if err != nil { + return "" + } + return filepath.Join(home, ".stripe", "state") +} diff --git a/pkg/autoupdate/config_test.go b/pkg/autoupdate/config_test.go new file mode 100644 index 000000000..cfe9aec73 --- /dev/null +++ b/pkg/autoupdate/config_test.go @@ -0,0 +1,65 @@ +package autoupdate + +import ( + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestIsOptedOut_EnvVar(t *testing.T) { + t.Setenv("STRIPE_NO_AUTO_UPDATE", "1") + assert.True(t, IsOptedOut()) +} + +func TestIsOptedOut_NoConfig(t *testing.T) { + t.Setenv("STRIPE_NO_AUTO_UPDATE", "") + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + assert.False(t, IsOptedOut()) +} + +// scripts/install.sh tells users to write a [settings] table; the message the CLI +// prints when it updates tells them to write a top-level auto_update. Whichever +// set of instructions a user followed has to opt them out. +func TestIsOptedOut_Config(t *testing.T) { + tests := []struct { + name string + config string + optedOut bool + }{ + {"top-level false", "auto_update = false\n", true}, + {"top-level true", "auto_update = true\n", false}, + {"settings table false", "[settings]\nauto_update = false\n", true}, + {"settings table true", "[settings]\nauto_update = true\n", false}, + { + // Where it lands in a config.toml that already has profiles in it: with + // the other top-level settings, ahead of the first table. + "alongside the other top-level settings", + "color = ''\nauto_update = false\n\n[default]\ndevice_name = 'laptop'\n", + true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Setenv("STRIPE_NO_AUTO_UPDATE", "") + configDir := t.TempDir() + stripeDir := filepath.Join(configDir, "stripe") + require.NoError(t, os.MkdirAll(stripeDir, 0755)) + require.NoError(t, os.WriteFile(filepath.Join(stripeDir, "config.toml"), []byte(tt.config), 0644)) + t.Setenv("XDG_CONFIG_HOME", configDir) + + assert.Equal(t, tt.optedOut, IsOptedOut()) + }) + } +} + +func TestIsCurlInstall_EnvOverride(t *testing.T) { + t.Setenv("STRIPE_INSTALL_METHOD", "curl") + assert.True(t, IsCurlInstall()) + + t.Setenv("STRIPE_INSTALL_METHOD", "homebrew") + assert.False(t, IsCurlInstall()) +} diff --git a/pkg/autoupdate/reexec_unix.go b/pkg/autoupdate/reexec_unix.go new file mode 100644 index 000000000..bf0914ea9 --- /dev/null +++ b/pkg/autoupdate/reexec_unix.go @@ -0,0 +1,19 @@ +//go:build !windows + +package autoupdate + +import ( + "os" + "syscall" + + log "github.com/sirupsen/logrus" +) + +// reexec hands the invocation to the binary that was just installed, so the +// command the user typed runs on the new version. +func reexec(exe string) { + err := syscall.Exec(exe, os.Args, os.Environ()) + if err != nil { + log.Debug("autoupdate: re-exec failed: ", err) + } +} diff --git a/pkg/autoupdate/reexec_windows.go b/pkg/autoupdate/reexec_windows.go new file mode 100644 index 000000000..457f361fc --- /dev/null +++ b/pkg/autoupdate/reexec_windows.go @@ -0,0 +1,44 @@ +//go:build windows + +package autoupdate + +import ( + "errors" + "os" + "os/exec" + "os/signal" + + log "github.com/sirupsen/logrus" +) + +// reexec runs the binary that was just installed as a child process and exits +// with its status. +// +// Windows has no execve. syscall.Exec is a stub there that always fails, so the +// only way to hand the invocation to the new binary is to start it and forward +// what it did. The child inherits this process's standard streams and console, so +// an interactive command such as `stripe login` still works. +// +// Returning instead of exiting means the update landed but this invocation +// carries on running the old image, which is a worse outcome than re-execing and +// a better one than failing the command outright. +func reexec(exe string) { + cmd := exec.Command(exe, os.Args[1:]...) //nolint:gosec // exe is this process's own path + cmd.Stdin, cmd.Stdout, cmd.Stderr = os.Stdin, os.Stdout, os.Stderr + + // Ctrl-C is delivered to every process attached to the console, this one + // included. Ignore it here so that the child, which is the process actually + // doing the work now, decides what a Ctrl-C means. + signal.Ignore(os.Interrupt) + + err := cmd.Run() + + var exitErr *exec.ExitError + if err != nil && !errors.As(err, &exitErr) { + log.Debug("autoupdate: could not run the updated binary: ", err) + + return + } + + os.Exit(cmd.ProcessState.ExitCode()) +} diff --git a/pkg/autoupdate/replace.go b/pkg/autoupdate/replace.go new file mode 100644 index 000000000..b185c0c79 --- /dev/null +++ b/pkg/autoupdate/replace.go @@ -0,0 +1,111 @@ +package autoupdate + +import ( + "fmt" + "os" + "path/filepath" + "strings" +) + +// oldSuffix names the outgoing binary while it is being replaced. +const oldSuffix = ".old" + +// rename is indirected so tests can fail the second move and check that the +// working binary comes back. +var rename = os.Rename + +// replaceBinary moves staged into place at dst, which is the image of the process +// calling it. +// +// The live binary is renamed aside first rather than overwritten. On Windows that +// is the only thing that works: the OS locks the image of a running executable +// against writes and deletes, but a rename only touches the directory entry, so +// it is allowed. Unix would tolerate a plain rename over the target, but taking +// the same path on both platforms means the interesting case is the one that runs +// everywhere, including in the unit tests. +func replaceBinary(dst, staged string) error { + if err := os.Chmod(staged, 0755); err != nil { + return fmt.Errorf("chmod failed: %w", err) + } + + aside, err := asideName(dst) + if err != nil { + return fmt.Errorf("cannot make room next to %s: %w", dst, err) + } + + movedAside := false + + if _, err := os.Stat(dst); err == nil { + if err := rename(dst, aside); err != nil { + return fmt.Errorf("cannot move %s aside: %w", dst, err) + } + + movedAside = true + } + + if err := rename(staged, dst); err != nil { + if movedAside { + // Put the working binary back. Failing to install an update is + // recoverable; leaving no stripe on PATH at all is not. + _ = rename(aside, dst) + } + + return fmt.Errorf("cannot replace binary: %w", err) + } + + // Best effort: on Windows this fails while the calling process is still running + // from the old image. removeSupersededBinary clears it on a later invocation. + _ = os.Remove(aside) + + return nil +} + +// asideName returns a free path to move the outgoing binary to. +// +// The usual name is dst+".old", left over from a previous update or not there at +// all. When it is there and cannot be deleted, the move has to go somewhere else: +// Windows refuses to delete or replace a file while any process is running from +// it, and after a concurrent update this process may be running from that very +// file. Reusing the name in that case fails the whole update, so fall back to a +// unique one, which a later invocation cleans up along with the rest. +func asideName(dst string) (string, error) { + aside := dst + oldSuffix + if err := os.Remove(aside); err == nil || os.IsNotExist(err) { + return aside, nil + } + + // Created rather than just named so that two updates racing here cannot pick + // the same path. Renaming over the empty placeholder is what claims it. + placeholder, err := os.CreateTemp(filepath.Dir(dst), filepath.Base(dst)+oldSuffix+".*") + if err != nil { + return "", err + } + + name := placeholder.Name() + _ = placeholder.Close() + + return name, nil +} + +// removeSupersededBinary deletes the outgoing binaries earlier updates parked +// next to the new one. +// +// On Windows those deletes could not happen at update time, because a process was +// still running from each file. Any later invocation is a process that is not, so +// this runs on every invocation and almost always finds nothing. It scans rather +// than deleting one known name because asideName falls back to a suffixed name +// when the usual one is occupied. +func removeSupersededBinary(exe string) { + dir, base := filepath.Split(exe) + + entries, err := os.ReadDir(dir) + if err != nil { + return + } + + for _, entry := range entries { + if strings.HasPrefix(entry.Name(), base+oldSuffix) { + _ = os.Remove(filepath.Join(dir, entry.Name())) + } + } +} diff --git a/pkg/autoupdate/replace_test.go b/pkg/autoupdate/replace_test.go new file mode 100644 index 000000000..f3c548ebc --- /dev/null +++ b/pkg/autoupdate/replace_test.go @@ -0,0 +1,135 @@ +package autoupdate + +import ( + "errors" + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestReplaceBinaryInstallsOverALiveBinary(t *testing.T) { + dir := t.TempDir() + dst := filepath.Join(dir, binaryName()) + staged := filepath.Join(dir, "staged") + + require.NoError(t, os.WriteFile(dst, []byte("old"), 0755)) + require.NoError(t, os.WriteFile(staged, []byte("new"), 0644)) + + require.NoError(t, replaceBinary(dst, staged)) + + installed, err := os.ReadFile(dst) + require.NoError(t, err) + assert.Equal(t, "new", string(installed)) + + // Unix can unlink the moved-aside file immediately. Windows cannot while a + // process is running from the image, but nothing is running this one. + assert.NoFileExists(t, dst+oldSuffix) + assert.NoFileExists(t, staged) +} + +func TestReplaceBinaryInstallsWhenNothingIsThere(t *testing.T) { + dir := t.TempDir() + dst := filepath.Join(dir, binaryName()) + staged := filepath.Join(dir, "staged") + + require.NoError(t, os.WriteFile(staged, []byte("new"), 0644)) + + require.NoError(t, replaceBinary(dst, staged)) + + installed, err := os.ReadFile(dst) + require.NoError(t, err) + assert.Equal(t, "new", string(installed)) +} + +func TestReplaceBinaryClearsAStaleOldFile(t *testing.T) { + dir := t.TempDir() + dst := filepath.Join(dir, binaryName()) + staged := filepath.Join(dir, "staged") + + require.NoError(t, os.WriteFile(dst, []byte("old"), 0755)) + require.NoError(t, os.WriteFile(dst+oldSuffix, []byte("older"), 0755)) + require.NoError(t, os.WriteFile(staged, []byte("new"), 0644)) + + require.NoError(t, replaceBinary(dst, staged)) + + assert.NoFileExists(t, dst+oldSuffix) +} + +func TestReplaceBinaryWhenTheOldNameCannotBeCleared(t *testing.T) { + dir := t.TempDir() + dst := filepath.Join(dir, binaryName()) + staged := filepath.Join(dir, "staged") + + require.NoError(t, os.WriteFile(dst, []byte("old"), 0755)) + require.NoError(t, os.WriteFile(staged, []byte("new"), 0644)) + + // Windows refuses to delete or replace a file while a process is running from + // it, which is the state of the .old file left behind by a concurrent update. A + // non-empty directory is refused by os.Remove and os.Rename on every platform, + // so it stands in for that obstacle in a test that can run anywhere. + require.NoError(t, os.MkdirAll(filepath.Join(dst+oldSuffix, "occupied"), 0755)) + + require.NoError(t, replaceBinary(dst, staged), + "an unusable .old name must not fail the update") + + installed, err := os.ReadFile(dst) + require.NoError(t, err) + assert.Equal(t, "new", string(installed)) + + // The outgoing binary went to a suffixed name, and whatever was holding the + // usual one was left alone rather than clobbered. + assert.DirExists(t, dst+oldSuffix) +} + +func TestReplaceBinaryRestoresTheOldBinaryOnFailure(t *testing.T) { + dir := t.TempDir() + dst := filepath.Join(dir, binaryName()) + staged := filepath.Join(dir, "staged") + + require.NoError(t, os.WriteFile(dst, []byte("old"), 0755)) + require.NoError(t, os.WriteFile(staged, []byte("new"), 0644)) + + // Fail the move of the new binary into place, after the live one has been moved + // aside: a full disk, a revoked permission, an antivirus hold on the download. + original := rename + rename = func(from, to string) error { + if from == staged { + return errors.New("cannot move the new binary into place") + } + + return original(from, to) + } + + t.Cleanup(func() { rename = original }) + + err := replaceBinary(dst, staged) + require.Error(t, err) + + restored, readErr := os.ReadFile(dst) + require.NoError(t, readErr) + assert.Equal(t, "old", string(restored), "a failed update has to leave a working binary behind") +} + +func TestRemoveSupersededBinary(t *testing.T) { + dir := t.TempDir() + exe := filepath.Join(dir, binaryName()) + + require.NoError(t, os.WriteFile(exe, []byte("current"), 0755)) + require.NoError(t, os.WriteFile(exe+oldSuffix, []byte("previous"), 0755)) + // What asideName falls back to when .old itself is occupied. + require.NoError(t, os.WriteFile(exe+oldSuffix+".2971828", []byte("older"), 0755)) + require.NoError(t, os.WriteFile(filepath.Join(dir, "unrelated"), []byte("keep"), 0644)) + + removeSupersededBinary(exe) + + assert.NoFileExists(t, exe+oldSuffix) + assert.NoFileExists(t, exe+oldSuffix+".2971828") + assert.FileExists(t, exe) + assert.FileExists(t, filepath.Join(dir, "unrelated")) + + // Nothing to remove is the common case and must not be an error. + removeSupersededBinary(exe) +} diff --git a/pkg/autoupdate/telemetry.go b/pkg/autoupdate/telemetry.go new file mode 100644 index 000000000..d0aaa28b0 --- /dev/null +++ b/pkg/autoupdate/telemetry.go @@ -0,0 +1,68 @@ +package autoupdate + +import ( + "fmt" + "net/http" + "net/url" + "os" + "runtime" + "strings" + "time" + + "github.com/google/uuid" + log "github.com/sirupsen/logrus" + + "github.com/stripe/stripe-cli/pkg/version" +) + +var telemetryEndpoint = "https://r.stripe.com/0" + +func init() { + if raw := os.Getenv("STRIPE_TELEMETRY_URL"); raw != "" { + telemetryEndpoint = raw + } +} + +func sendTelemetryEvent(eventName, eventValue string) { + if isTelemetryOptedOut() { + return + } + + data := url.Values{} + data.Set("client_id", "stripe-cli") + data.Set("event_id", uuid.NewString()) + data.Set("event_name", eventName) + data.Set("event_value", eventValue) + data.Set("created", fmt.Sprint(time.Now().Unix())) + data.Set("cli_version", version.Version) + data.Set("os", runtime.GOOS) + data.Set("arch", runtime.GOARCH) + data.Set("install_method", "curl") + + req, err := http.NewRequest(http.MethodPost, telemetryEndpoint, strings.NewReader(data.Encode())) + if err != nil { + log.Debug("autoupdate telemetry: failed to create request: ", err) + return + } + + req.Header.Set("origin", "stripe-cli") + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + + client := &http.Client{Timeout: 3 * time.Second} + resp, err := client.Do(req) + if err != nil { + log.Debug("autoupdate telemetry: failed to send event: ", err) + return + } + resp.Body.Close() +} + +func isTelemetryOptedOut() bool { + for _, key := range []string{"STRIPE_CLI_TELEMETRY_OPTOUT", "DO_NOT_TRACK"} { + val := strings.ToLower(os.Getenv(key)) + if val == "1" || val == "true" { + return true + } + } + return false +} diff --git a/pkg/autoupdate/telemetry_test.go b/pkg/autoupdate/telemetry_test.go new file mode 100644 index 000000000..ab234e51d --- /dev/null +++ b/pkg/autoupdate/telemetry_test.go @@ -0,0 +1,80 @@ +package autoupdate + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestSendTelemetryEvent(t *testing.T) { + var received bool + var gotBody string + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + received = true + assert.Equal(t, "stripe-cli", r.Header.Get("origin")) + assert.Equal(t, "application/x-www-form-urlencoded", r.Header.Get("Content-Type")) + assert.Equal(t, http.MethodPost, r.Method) + + err := r.ParseForm() + assert.NoError(t, err) + gotBody = r.Form.Get("event_name") + + assert.Equal(t, "stripe-cli", r.Form.Get("client_id")) + assert.Equal(t, "curl", r.Form.Get("install_method")) + assert.NotEmpty(t, r.Form.Get("event_id")) + assert.NotEmpty(t, r.Form.Get("created")) + + w.WriteHeader(http.StatusOK) + })) + defer server.Close() + + original := telemetryEndpoint + telemetryEndpoint = server.URL + defer func() { telemetryEndpoint = original }() + + sendTelemetryEvent("Auto-Update Succeeded", "from=1.0.0 to=1.1.0") + + assert.True(t, received) + assert.Equal(t, "Auto-Update Succeeded", gotBody) +} + +func TestSendTelemetryEvent_OptedOut(t *testing.T) { + var received bool + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + received = true + w.WriteHeader(http.StatusOK) + })) + defer server.Close() + + original := telemetryEndpoint + telemetryEndpoint = server.URL + defer func() { telemetryEndpoint = original }() + + t.Setenv("STRIPE_CLI_TELEMETRY_OPTOUT", "1") + + sendTelemetryEvent("Auto-Update Succeeded", "from=1.0.0 to=1.1.0") + + assert.False(t, received) +} + +func TestSendTelemetryEvent_DoNotTrack(t *testing.T) { + var received bool + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + received = true + w.WriteHeader(http.StatusOK) + })) + defer server.Close() + + original := telemetryEndpoint + telemetryEndpoint = server.URL + defer func() { telemetryEndpoint = original }() + + t.Setenv("DO_NOT_TRACK", "true") + + sendTelemetryEvent("Auto-Update Succeeded", "from=1.0.0 to=1.1.0") + + assert.False(t, received) +} diff --git a/pkg/autoupdate/updater.go b/pkg/autoupdate/updater.go new file mode 100644 index 000000000..81fbfefa1 --- /dev/null +++ b/pkg/autoupdate/updater.go @@ -0,0 +1,258 @@ +package autoupdate + +import ( + "archive/tar" + "archive/zip" + "bytes" + "compress/gzip" + "fmt" + "io" + "net/http" + "os" + "path" + "path/filepath" + "runtime" + "strings" + + log "github.com/sirupsen/logrus" + + "github.com/stripe/stripe-cli/pkg/errorcategory" + "github.com/stripe/stripe-cli/pkg/version" +) + +// ApplyIfPending checks for a pending update marker and applies it. +// If an update is applied, it re-execs the current process with the new binary. +// This function only returns if no update was applied — or, on Windows, if the +// new binary could not be started, in which case this invocation carries on +// running the image it already has. +func ApplyIfPending() { + if version.Version == "master" { + return + } + if !IsCurlInstall() { + return + } + + exe, err := resolvedExecutable() + if err != nil { + log.Debug("autoupdate: cannot determine executable path: ", err) + return + } + + // Ahead of the opt-out check, so that turning auto-update off does not leave + // the outgoing binary from an earlier update parked next to the new one + // forever. On Windows the update that installed it could not delete it: a + // process was still running from that file. This one is not, so it can. + removeSupersededBinary(exe) + + if IsOptedOut() { + return + } + + marker := ReadMarker() + if marker == nil { + return + } + + current := strings.TrimPrefix(version.Version, "v") + target := strings.TrimPrefix(marker.Version, "v") + if current == target { + ClearMarker() + return + } + + fmt.Fprintf(os.Stderr, "Automatically updating Stripe CLI from %s to %s.\n", current, target) + fmt.Fprintf(os.Stderr, "To disable auto-update, set STRIPE_NO_AUTO_UPDATE=1 or add auto_update = false to ~/.config/stripe/config.toml\n") + + if err := downloadAndReplace(marker, exe); err != nil { + fmt.Fprintf(os.Stderr, "Auto-update failed: %v. Continuing with current version.\n", err) + sendTelemetryEvent("Auto-Update Failed", fmt.Sprintf("from=%s to=%s error=%s", current, target, err.Error())) + ClearMarker() + return + } + + ClearMarker() + fmt.Fprintf(os.Stderr, "Updated successfully ✓\n") + fmt.Fprintf(os.Stderr, "Run 'stripe version --notes' to see what's new.\n") + sendTelemetryEvent("Auto-Update Succeeded", fmt.Sprintf("from=%s to=%s", current, target)) + + reexec(exe) +} + +func downloadAndReplace(marker *UpdateMarker, exePath string) error { + resp, err := http.Get(marker.DownloadURL) //nolint:gosec + if err != nil { + return fmt.Errorf("download failed: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return errorcategory.Errorf(errorcategory.Network, "download returned status %d", resp.StatusCode) + } + + tmpArchive, err := os.CreateTemp(filepath.Dir(exePath), "stripe-update-archive-*") + if err != nil { + return fmt.Errorf("cannot create temp file: %w", err) + } + tmpArchivePath := tmpArchive.Name() + defer os.Remove(tmpArchivePath) + + if _, err := io.Copy(tmpArchive, resp.Body); err != nil { + tmpArchive.Close() + return fmt.Errorf("download interrupted: %w", err) + } + tmpArchive.Close() + + if marker.Checksum != "" && !VerifyChecksum(tmpArchivePath, marker.Checksum) { + return errorcategory.Errorf(errorcategory.Network, "checksum verification failed") + } + + tmpBinary, err := os.CreateTemp(filepath.Dir(exePath), "stripe-update-*") + if err != nil { + return fmt.Errorf("cannot create temp binary: %w", err) + } + tmpBinaryPath := tmpBinary.Name() + tmpBinary.Close() + + if err := extractBinary(tmpArchivePath, tmpBinaryPath); err != nil { + os.Remove(tmpBinaryPath) + return fmt.Errorf("extraction failed: %w", err) + } + + if err := replaceBinary(exePath, tmpBinaryPath); err != nil { + os.Remove(tmpBinaryPath) + return err + } + + return nil +} + +// resolvedExecutable is the path of the running binary with every symlink +// resolved. A user may have symlinked stripe onto their PATH, or a component of +// the install path may itself be a link, and the update has to land on the real +// file. +func resolvedExecutable() (string, error) { + exe, err := os.Executable() + if err != nil { + return "", err + } + + return filepath.EvalSymlinks(exe) +} + +// binaryName is the name of the executable inside the release archive, which is +// also the name it is installed under. goreleaser appends .exe to Windows builds. +func binaryName() string { + if runtime.GOOS == "windows" { + return "stripe.exe" + } + + return "stripe" +} + +// extractBinary pulls the stripe executable out of a release archive. +// +// Which format that is is read off the file rather than its name: the archive was +// downloaded to a temporary path whose name says nothing about its contents. Mac +// and Linux releases ship a .tar.gz, Windows a .zip. +func extractBinary(archivePath, destPath string) error { + zipped, err := isZipArchive(archivePath) + if err != nil { + return err + } + + if zipped { + return extractFromZip(archivePath, destPath) + } + + return extractFromTarGz(archivePath, destPath) +} + +func isZipArchive(archivePath string) (bool, error) { + f, err := os.Open(archivePath) + if err != nil { + return false, err + } + defer f.Close() + + magic := make([]byte, 4) + if _, err := io.ReadFull(f, magic); err != nil { + // Too short to be either format. Let the tar reader say so, so that the + // error names what the file was expected to be. + return false, nil + } + + return bytes.Equal(magic, []byte("PK\x03\x04")), nil +} + +func extractFromTarGz(archivePath, destPath string) error { + f, err := os.Open(archivePath) + if err != nil { + return err + } + defer f.Close() + + gz, err := gzip.NewReader(f) + if err != nil { + return err + } + defer gz.Close() + + tr := tar.NewReader(gz) + for { + hdr, err := tr.Next() + if err == io.EOF { + break + } + if err != nil { + return err + } + + if path.Base(hdr.Name) == binaryName() && hdr.Typeflag == tar.TypeReg { + out, err := os.Create(destPath) + if err != nil { + return err + } + if _, err := io.Copy(out, tr); err != nil { + out.Close() + return err + } + return out.Close() + } + } + return errorcategory.Errorf(errorcategory.Internal, "stripe binary not found in archive") +} + +func extractFromZip(archivePath, destPath string) error { + r, err := zip.OpenReader(archivePath) + if err != nil { + return err + } + defer r.Close() + + for _, entry := range r.File { + // Only this one entry is extracted, to a path this function chose, so a + // crafted archive cannot write anywhere else. + if entry.FileInfo().IsDir() || path.Base(entry.Name) != binaryName() { + continue + } + + in, err := entry.Open() + if err != nil { + return err + } + defer in.Close() + + out, err := os.Create(destPath) + if err != nil { + return err + } + if _, err := io.Copy(out, in); err != nil { + out.Close() + return err + } + return out.Close() + } + + return errorcategory.Errorf(errorcategory.Internal, "stripe binary not found in archive") +} diff --git a/pkg/autoupdate/updater_test.go b/pkg/autoupdate/updater_test.go new file mode 100644 index 000000000..27df2da91 --- /dev/null +++ b/pkg/autoupdate/updater_test.go @@ -0,0 +1,321 @@ +package autoupdate + +import ( + "archive/tar" + "archive/zip" + "compress/gzip" + "crypto/sha256" + "encoding/hex" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "runtime" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func createTestTarGz(t *testing.T, filename string, content []byte) string { + t.Helper() + path := filepath.Join(t.TempDir(), "test.tar.gz") + f, err := os.Create(path) + require.NoError(t, err) + + gw := gzip.NewWriter(f) + tw := tar.NewWriter(gw) + + require.NoError(t, tw.WriteHeader(&tar.Header{ + Name: filename, + Size: int64(len(content)), + Mode: 0755, + Typeflag: tar.TypeReg, + })) + _, err = tw.Write(content) + require.NoError(t, err) + + require.NoError(t, tw.Close()) + require.NoError(t, gw.Close()) + require.NoError(t, f.Close()) + return path +} + +// createTestZip writes the archive format the Windows release publishes. +func createTestZip(t *testing.T, filename string, content []byte) string { + t.Helper() + path := filepath.Join(t.TempDir(), "test.zip") + f, err := os.Create(path) + require.NoError(t, err) + + zw := zip.NewWriter(f) + entry, err := zw.Create(filename) + require.NoError(t, err) + _, err = entry.Write(content) + require.NoError(t, err) + + require.NoError(t, zw.Close()) + require.NoError(t, f.Close()) + return path +} + +func sha256sum(path string) string { + data, _ := os.ReadFile(path) + h := sha256.Sum256(data) + return hex.EncodeToString(h[:]) +} + +func TestExtractFromTarGz(t *testing.T) { + content := []byte("#!/bin/sh\necho hello\n") + archivePath := createTestTarGz(t, binaryName(), content) + + destPath := filepath.Join(t.TempDir(), "stripe") + err := extractFromTarGz(archivePath, destPath) + require.NoError(t, err) + + got, err := os.ReadFile(destPath) + require.NoError(t, err) + assert.Equal(t, content, got) +} + +func TestExtractFromTarGz_NestedPath(t *testing.T) { + content := []byte("binary content") + archivePath := createTestTarGz(t, "stripe_1.43.8_linux_arm64/"+binaryName(), content) + + destPath := filepath.Join(t.TempDir(), "stripe") + err := extractFromTarGz(archivePath, destPath) + require.NoError(t, err) + + got, err := os.ReadFile(destPath) + require.NoError(t, err) + assert.Equal(t, content, got) +} + +func TestExtractFromTarGz_NoBinary(t *testing.T) { + archivePath := createTestTarGz(t, "not-stripe", []byte("nope")) + + destPath := filepath.Join(t.TempDir(), "stripe") + err := extractFromTarGz(archivePath, destPath) + assert.Error(t, err) + assert.Contains(t, err.Error(), "not found in archive") +} + +func TestExtractFromZip(t *testing.T) { + content := []byte("MZ windows binary") + archivePath := createTestZip(t, binaryName(), content) + + destPath := filepath.Join(t.TempDir(), binaryName()) + err := extractFromZip(archivePath, destPath) + require.NoError(t, err) + + got, err := os.ReadFile(destPath) + require.NoError(t, err) + assert.Equal(t, content, got) +} + +func TestExtractFromZip_NestedPath(t *testing.T) { + content := []byte("MZ windows binary") + // Zip entry names use forward slashes whatever the platform reading them. + archivePath := createTestZip(t, "stripe_1.43.8_windows_x86_64/"+binaryName(), content) + + destPath := filepath.Join(t.TempDir(), binaryName()) + err := extractFromZip(archivePath, destPath) + require.NoError(t, err) + + got, err := os.ReadFile(destPath) + require.NoError(t, err) + assert.Equal(t, content, got) +} + +func TestExtractFromZip_NoBinary(t *testing.T) { + archivePath := createTestZip(t, "not-stripe", []byte("nope")) + + destPath := filepath.Join(t.TempDir(), binaryName()) + err := extractFromZip(archivePath, destPath) + assert.Error(t, err) + assert.Contains(t, err.Error(), "not found in archive") +} + +// The archive is downloaded to a temporary name, so its format has to be read off +// the bytes rather than an extension. +func TestExtractBinary_PicksTheFormatFromTheContents(t *testing.T) { + for _, tt := range []struct { + format string + archive func(*testing.T, string, []byte) string + }{ + {"tar.gz", createTestTarGz}, + {"zip", createTestZip}, + } { + t.Run(tt.format, func(t *testing.T) { + content := []byte("binary for " + tt.format) + // A name that gives nothing away, as os.CreateTemp produces. + archivePath := tt.archive(t, binaryName(), content) + anonymous := filepath.Join(t.TempDir(), "stripe-update-archive-1234") + require.NoError(t, os.Rename(archivePath, anonymous)) + + destPath := filepath.Join(t.TempDir(), binaryName()) + require.NoError(t, extractBinary(anonymous, destPath)) + + got, err := os.ReadFile(destPath) + require.NoError(t, err) + assert.Equal(t, content, got) + }) + } +} + +func TestExtractBinary_ShortFileIsAnError(t *testing.T) { + archivePath := filepath.Join(t.TempDir(), "stripe-update-archive-1234") + require.NoError(t, os.WriteFile(archivePath, []byte("PK"), 0644)) + + assert.Error(t, extractBinary(archivePath, filepath.Join(t.TempDir(), binaryName()))) +} + +func TestDownloadAndReplace(t *testing.T) { + content := []byte("#!/bin/sh\necho updated\n") + archivePath := createTestTarGz(t, binaryName(), content) + archiveData, err := os.ReadFile(archivePath) + require.NoError(t, err) + + checksum := sha256sum(archivePath) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Write(archiveData) + })) + defer server.Close() + + dir := t.TempDir() + exePath := filepath.Join(dir, binaryName()) + require.NoError(t, os.WriteFile(exePath, []byte("old binary"), 0755)) + + marker := &UpdateMarker{ + Version: "1.43.8", + DownloadURL: server.URL + "/stripe.tar.gz", + Checksum: checksum, + } + + err = downloadAndReplace(marker, exePath) + require.NoError(t, err) + + got, err := os.ReadFile(exePath) + require.NoError(t, err) + assert.Equal(t, content, got) + + if runtime.GOOS != "windows" { + info, err := os.Stat(exePath) + require.NoError(t, err) + assert.Equal(t, os.FileMode(0755), info.Mode().Perm()) + } + + // Nothing staged is left in the install directory. + assert.NoFileExists(t, exePath+oldSuffix) + entries, err := os.ReadDir(dir) + require.NoError(t, err) + assert.Len(t, entries, 1) +} + +// The Windows release ships a zip rather than a tar.gz, and the binary inside it +// is named stripe.exe. +func TestDownloadAndReplace_Zip(t *testing.T) { + content := []byte("MZ updated windows binary") + archivePath := createTestZip(t, binaryName(), content) + archiveData, err := os.ReadFile(archivePath) + require.NoError(t, err) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Write(archiveData) + })) + defer server.Close() + + dir := t.TempDir() + exePath := filepath.Join(dir, binaryName()) + require.NoError(t, os.WriteFile(exePath, []byte("old binary"), 0755)) + + marker := &UpdateMarker{ + Version: "1.43.8", + DownloadURL: server.URL + "/stripe.zip", + Checksum: sha256sum(archivePath), + } + + require.NoError(t, downloadAndReplace(marker, exePath)) + + got, err := os.ReadFile(exePath) + require.NoError(t, err) + assert.Equal(t, content, got) +} + +func TestDownloadAndReplace_BadChecksum(t *testing.T) { + content := []byte("#!/bin/sh\necho updated\n") + archivePath := createTestTarGz(t, binaryName(), content) + archiveData, err := os.ReadFile(archivePath) + require.NoError(t, err) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Write(archiveData) + })) + defer server.Close() + + dir := t.TempDir() + exePath := filepath.Join(dir, binaryName()) + require.NoError(t, os.WriteFile(exePath, []byte("old binary"), 0755)) + + marker := &UpdateMarker{ + Version: "1.43.8", + DownloadURL: server.URL + "/stripe.tar.gz", + Checksum: "0000000000000000000000000000000000000000000000000000000000000000", + } + + err = downloadAndReplace(marker, exePath) + assert.Error(t, err) + assert.Contains(t, err.Error(), "checksum verification failed") + + got, _ := os.ReadFile(exePath) + assert.Equal(t, []byte("old binary"), got, "original binary should be unchanged") +} + +func TestDownloadAndReplace_ServerError(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + defer server.Close() + + dir := t.TempDir() + exePath := filepath.Join(dir, binaryName()) + require.NoError(t, os.WriteFile(exePath, []byte("old binary"), 0755)) + + marker := &UpdateMarker{ + Version: "1.43.8", + DownloadURL: server.URL + "/stripe.tar.gz", + } + + err := downloadAndReplace(marker, exePath) + assert.Error(t, err) + assert.Contains(t, err.Error(), "status 500") +} + +func TestApplyIfPending_NoMarker(t *testing.T) { + tmpDir := t.TempDir() + original := GetStateDirFn + defer func() { GetStateDirFn = original }() + GetStateDirFn = func() string { return tmpDir } + + // Should return without panic when no marker exists + ApplyIfPending() +} + +func TestApplyIfPending_SameVersion(t *testing.T) { + tmpDir := t.TempDir() + original := GetStateDirFn + defer func() { GetStateDirFn = original }() + GetStateDirFn = func() string { return tmpDir } + + // Write a marker with the current version — should be cleared without action + WriteMarker(UpdateMarker{ + Version: "master", + DownloadURL: "https://example.com/stripe.tar.gz", + }) + + ApplyIfPending() + + // Marker should still exist since version.Version is "master" and we return early + // (the "master" check happens before reading the marker) +} diff --git a/pkg/cmd/agent.go b/pkg/cmd/agent.go index be105fafa..de42d59c2 100644 --- a/pkg/cmd/agent.go +++ b/pkg/cmd/agent.go @@ -32,6 +32,7 @@ var agentClientID = map[string]string{ "claude_code": agentsetup.ClientClaudeCode, "codex_cli": agentsetup.ClientCodex, "cursor": agentsetup.ClientCursor, + "grok": agentsetup.ClientGrok, } // providerOrder is the canonical display order for known clients. Providers not @@ -40,6 +41,7 @@ var providerOrder = []string{ agentsetup.ClientClaudeCode, agentsetup.ClientCodex, agentsetup.ClientCursor, + agentsetup.ClientGrok, } type agentCmd struct { @@ -410,6 +412,12 @@ func (asc *agentSetupCmd) install(ctx context.Context, out io.Writer, providers plan := provider.Plan(status, asc.force) fmt.Fprintf(out, "\n %s\n", status.DisplayName) + if status.Status == agentsetup.StatusError { + fmt.Fprintf(out, " %s error: %s\n", cross, status.Error) + sendAgentEvent(ctx, "Agent Setup: Plugin Install", status.Client+":error") + errCount++ + continue + } switch plan.Action { case agentsetup.ActionNone: fmt.Fprintln(out, " already set up") @@ -858,6 +866,7 @@ Supported clients for automatic setup: • Claude Code https://claude.ai/code • Cursor https://cursor.com • Codex CLI https://openai.com/codex/ + • Grok Build https://x.ai/build You can still install Stripe skills. `) @@ -870,6 +879,7 @@ Supported clients for automatic setup: • Claude Code https://claude.ai/code • Cursor https://cursor.com • Codex CLI https://openai.com/codex/ + • Grok Build https://x.ai/build Once a client is installed, re-run: stripe agent setup `) diff --git a/pkg/cmd/agent_env_test.go b/pkg/cmd/agent_env_test.go index fccc074e6..3e2789b39 100644 --- a/pkg/cmd/agent_env_test.go +++ b/pkg/cmd/agent_env_test.go @@ -14,6 +14,8 @@ func clearAgentEnv(t *testing.T) { t.Helper() for _, key := range []string{ + "AI_AGENT", + "AGENT", "ANTIGRAVITY_CLI_ALIAS", "CLAUDECODE", "CLAUDE_CODE_ENTRYPOINT", @@ -25,6 +27,9 @@ func clearAgentEnv(t *testing.T) { "CODEX_THREAD_ID", "CURSOR_AGENT", "GEMINI_CLI", + "GROK_AGENT", + "GROK_SESSION_ID", + "HERMES_AGENT", "OPENCLAW_SHELL", "OPENCODE", } { diff --git a/pkg/cmd/agent_test.go b/pkg/cmd/agent_test.go index 305b3e55b..285cfafba 100644 --- a/pkg/cmd/agent_test.go +++ b/pkg/cmd/agent_test.go @@ -215,6 +215,39 @@ func TestAgentSetupJSONShowsUpgradeHintWhenPluginCommandFails(t *testing.T) { require.Contains(t, result.Clients[0].Error, "upgrade Claude Code") } +func TestAgentSetupCodexDiscoveryFailureReportsError(t *testing.T) { + for _, tt := range []struct { + name string + listErr error + wantError string + }{ + {"listing fails", errors.New("listing unavailable"), "listing Codex marketplaces: listing unavailable"}, + {"no supported marketplace", nil, "no supported Codex marketplace is available"}, + } { + t.Run(tt.name, func(t *testing.T) { + codex := codexMissingProvider(func(context.Context, string, ...string) error { + t.Fatal("must not install without a detected marketplace") + return nil + }) + codex.RunOutput = func(_ context.Context, _ string, args ...string) ([]byte, error) { + require.Equal(t, []string{"plugin", "marketplace", "list", "--json"}, args) + return []byte(`{"marketplaces":[]}`), tt.listErr + } + setup := testAgentSetupCmd() + setup.providers = map[string]agentsetup.Provider{codex.ID(): codex} + setup.callingAgent = func() string { return "" } + setup.cmd.SetContext(context.Background()) + + output, err := executeCommand(setup.cmd, "--client", "codex", "--yes") + + require.ErrorContains(t, err, "1 item(s) failed to set up") + require.Contains(t, output, tt.wantError) + require.Contains(t, output, "0 installed, 0 updated, 0 skipped, 1 errors") + require.NotContains(t, output, "already set up") + }) + } +} + func TestAgentSetupForceYesInvokesInstallerWhenInstalled(t *testing.T) { var called bool setup := newTestAgentSetupCmdInstalled(t, func(ctx context.Context, name string, args ...string) error { @@ -365,28 +398,52 @@ func TestAgentSetupClientFlagDoesNotCheckSkills(t *testing.T) { } func TestAgentSetupAutoInstallsForCallingAgent(t *testing.T) { - var installed []string - record := func(_ context.Context, name string, args ...string) error { - installed = append(installed, name) - return nil + callingAgents := []struct { + name string + displayName string + agent string + makeProvider func(record agentsetup.RunCommandFunc) agentsetup.Provider + }{ + { + name: "codex_cli", + displayName: "Codex CLI", + agent: "codex", + makeProvider: func(record agentsetup.RunCommandFunc) agentsetup.Provider { return codexMissingProvider(record) }, + }, + { + name: "grok", + displayName: "Grok", + agent: "grok", + makeProvider: func(record agentsetup.RunCommandFunc) agentsetup.Provider { return grokMissingProvider(record) }, + }, } - claude := agentsetup.NewClaudeProvider(claudeMissingPluginScanner(t), record) - codex := codexMissingProvider(record) + for _, agent := range callingAgents { + t.Run(agent.name, func(t *testing.T) { + var installedAgents []string + record := func(_ context.Context, name string, args ...string) error { + installedAgents = append(installedAgents, name) + return nil + } - setup := testAgentSetupCmd() - setup.providers = map[string]agentsetup.Provider{claude.ID(): claude, codex.ID(): codex} - // Simulate being invoked by Codex CLI — only its plugin should install, - // even though Claude is also detected, and with no --client flag. - setup.callingAgent = func() string { return "codex_cli" } - setup.cmd.SetContext(context.Background()) + claude := agentsetup.NewClaudeProvider(claudeMissingPluginScanner(t), record) + provider := agent.makeProvider(record) - output, err := executeCommand(setup.cmd) + setup := testAgentSetupCmd() + setup.providers = map[string]agentsetup.Provider{claude.ID(): claude, provider.ID(): provider} + // Simulate being invoked by the calling agent — only its plugin should + // install, even though Claude is also detected, and with no --client flag. + setup.callingAgent = func() string { return agent.name } + setup.cmd.SetContext(context.Background()) - require.NoError(t, err) - require.Equal(t, []string{"codex"}, installed) // Claude NOT installed - require.Contains(t, output, "Detected Codex CLI — setting up its Stripe plugin.") - require.Contains(t, output, "1 installed, 0 updated, 0 skipped, 0 errors") + output, err := executeCommand(setup.cmd) + + require.NoError(t, err) + require.Equal(t, []string{agent.agent}, installedAgents) // Claude NOT installed + require.Contains(t, output, fmt.Sprintf("Detected %s — setting up its Stripe plugin.", agent.displayName)) + require.Contains(t, output, "1 installed, 0 updated, 0 skipped, 0 errors") + }) + } } func TestAgentSetupCallingAgentDoesNotCheckSkills(t *testing.T) { @@ -430,7 +487,10 @@ func codexMissingProvider(record agentsetup.RunCommandFunc) agentsetup.CodexProv } return nil }, - RunOutput: func(context.Context, string, ...string) ([]byte, error) { + RunOutput: func(_ context.Context, _ string, args ...string) ([]byte, error) { + if args[1] == "marketplace" { + return []byte(`{"marketplaces":[{"name":"openai-curated"}]}`), nil + } if installed { return []byte(`{"installed":[{"pluginId":"stripe@openai-curated","name":"stripe","marketplaceName":"openai-curated","version":"1.0.0"}]}`), nil } @@ -439,6 +499,18 @@ func codexMissingProvider(record agentsetup.RunCommandFunc) agentsetup.CodexProv } } +// grokMissingProvider returns a Grok provider that detects the binary, starts +// with the Stripe plugin not installed, and records install commands via record. +func grokMissingProvider(record agentsetup.RunCommandFunc) agentsetup.GrokProvider { + return agentsetup.GrokProvider{ + Scanner: agentsetup.Scanner{LookPath: func(string) (string, error) { return "/usr/local/bin/grok", nil }}, + RunCommand: record, + RunOutput: func(context.Context, string, ...string) ([]byte, error) { + return []byte(`[]`), nil + }, + } +} + func TestAgentSetupUnsupportedAgentInstallsSkillsToLocal(t *testing.T) { var gotDir string claude := agentsetup.NewClaudeProvider(claudeMissingPluginScanner(t), func(context.Context, string, ...string) error { diff --git a/pkg/cmd/config_migration.go b/pkg/cmd/config_migration.go index 58a59969f..4dfac142d 100644 --- a/pkg/cmd/config_migration.go +++ b/pkg/cmd/config_migration.go @@ -11,6 +11,7 @@ import ( "github.com/spf13/afero" "github.com/spf13/cobra" + "github.com/stripe/stripe-cli/pkg/ansi" "github.com/stripe/stripe-cli/pkg/config" "github.com/stripe/stripe-cli/pkg/errorcategory" "github.com/stripe/stripe-cli/pkg/plugins" @@ -22,15 +23,15 @@ import ( // about it. Every dependency is a field so the policy can be tested without a // terminal or an installed plugin. type configMigration struct { - profilesFile string - needsMigration func() bool - pluginsReady func() bool - incompatibilities func() ([]plugins.ConfigV2Incompatibility, error) - installedPluginCount func() int - upgradePlugin func(plugins.ConfigV2Incompatibility) (string, error) - migrate func(path string) (bool, error) - reload func() error - out io.Writer + profilesFile string + needsMigration func() bool + pluginsReady func() bool + incompatibilities func() ([]plugins.ConfigV2Incompatibility, error) + upgradePlugin func(plugins.ConfigV2Incompatibility) (string, error) + migrate func(path string) (bool, error) + stampNew func(path string) error + reload func() error + out io.Writer } func newConfigMigration(cfg *config.Config, ctx context.Context) configMigration { @@ -45,20 +46,13 @@ func newConfigMigration(cfg *config.Config, ctx context.Context) configMigration incompatibilities: func() ([]plugins.ConfigV2Incompatibility, error) { return plugins.ConfigV2Incompatibilities(cfg, fs) }, - installedPluginCount: func() int { - names, err := plugins.GetInstalledPluginNames(cfg, fs) - if err != nil { - return 0 - } - - return len(names) - }, upgradePlugin: func(incompatibility plugins.ConfigV2Incompatibility) (string, error) { return upgradePluginForConfigV2(ctx, cfg, fs, incompatibility) }, - migrate: config.MigrateConfigFile, - reload: config.ReloadConfigFile, - out: os.Stderr, + migrate: config.MigrateConfigFile, + stampNew: config.StampNewConfigFile, + reload: config.ReloadConfigFile, + out: os.Stderr, } } @@ -73,10 +67,10 @@ func migrateConfigIfNeeded(cmd *cobra.Command) { newConfigMigration(&Config, cmd.Context()).run() } -// migrationSafeCommand reports whether it is acceptable to write status lines -// and to rewrite the config file while running this command. Help and shell -// completion output gets read by other programs, and a status line in the -// middle of it would be worse than a config file left in the old layout. +// migrationSafeCommand reports whether it is acceptable to rewrite the config +// file, and to say anything at all, while running this command. Help and shell +// completion output gets read by other programs, and a plugin upgrade notice in +// the middle of it would be worse than a config file left in the old layout. func migrationSafeCommand(cmd *cobra.Command) bool { for c := cmd; c != nil; c = c.Parent() { switch c.Name() { @@ -103,9 +97,9 @@ func (m configMigration) run() { } if _, err := os.Stat(m.profilesFile); err != nil { - // No config file yet, so there is nothing to move. The first write picks - // the layout. - logger.Debugf("Skipping the config migration: %s", err) + // No config file yet, so there is nothing to move -- but there is a layout + // to choose for whatever gets written first. + m.stampNewConfigFile(logger, err) return } @@ -118,59 +112,80 @@ func (m configMigration) run() { return } - m.migrateAndReload() + m.migrateAndReload(logger) +} + +// stampNewConfigFile records the v2 layout in a config file that does not exist +// yet, so that the first write into it -- usually a login -- lands in the new +// layout directly. +// +// Without this, the first command writes the flat layout and the *next* command +// migrates it. That shows a brand-new user a migration notice before they have +// logged in, and leaves a backup file holding a copy of the credentials the +// previous command just wrote. +// +// Gated on plugins for the same reason the migration is: once a profile is written +// under the profiles table, a plugin too old to look there cannot find it. Silent, +// though. There is nothing to migrate and nothing at risk, so an incompatible +// plugin simply means the file keeps the flat layout until that plugin is +// upgraded -- and a status line here would land in front of a user who has not +// run anything yet. +func (m configMigration) stampNewConfigFile(logger *log.Entry, statErr error) { + logger.Debugf("No config file to migrate: %s", statErr) + + if m.stampNew == nil || !m.pluginsReady() { + return + } + + incompatibilities, err := m.incompatibilities() + if err != nil { + logger.Debugf("Not recording the new config format: could not check installed plugins: %s", err) + return + } + + if len(incompatibilities) > 0 { + logger.Debugf("Not recording the new config format: %s", incompatibilities[0].Error()) + return + } + + if err := m.stampNew(m.profilesFile); err != nil { + logger.Debugf("Could not record the new config format in %s: %s", m.profilesFile, err) + return + } + + if err := m.reload(); err != nil { + logger.Debugf("Recorded the new config format in %s but could not re-read it: %s", m.profilesFile, err) + } } // ensurePluginsCompatible upgrades any installed plugin that cannot read the v2 // layout. It returns false when an upgrade fails, in which case the config file // is left alone. func (m configMigration) ensurePluginsCompatible(logger *log.Entry) bool { - fmt.Fprint(m.out, "checking installed plugins...") - incompatibilities, err := m.incompatibilities() if err != nil { - fmt.Fprintln(m.out) logger.Debugf("Skipping the config migration: could not check installed plugins: %s", err) return false } - if len(incompatibilities) == 0 { - switch n := m.pluginCount(); n { - case 0: - fmt.Fprintln(m.out, " none installed.") - case 1: - fmt.Fprintln(m.out, " 1 is compatible.") - default: - fmt.Fprintf(m.out, " all %d are compatible.\n", n) - } - - return true - } - - fmt.Fprintln(m.out) - for _, incompatibility := range incompatibilities { + color := ansi.Color(m.out) + fmt.Fprintln(m.out, color.Faint(fmt.Sprintf( + "Upgrading the %s plugin so it can read the updated config file.", incompatibility.Plugin, + )).String()) + newVersion, err := m.upgradeOne(incompatibility) if err != nil { - logger.Debugf("could not upgrade %s: %s", incompatibility.Plugin, err) - m.reportUpgradeFailure(incompatibility) + logger.Debugf("Skipping the config migration: could not upgrade %s (%s): %s", incompatibility.Plugin, incompatibility.Error(), err) return false } - fmt.Fprintf(m.out, "✔ upgraded %s from v%s to v%s.\n", incompatibility.Plugin, incompatibility.InstalledVersion, newVersion) + logger.Debugf("Upgraded %s from v%s to v%s", incompatibility.Plugin, incompatibility.InstalledVersion, newVersion) } return true } -func (m configMigration) pluginCount() int { - if m.installedPluginCount == nil { - return 0 - } - - return m.installedPluginCount() -} - func (m configMigration) upgradeOne(incompatibility plugins.ConfigV2Incompatibility) (string, error) { if m.upgradePlugin == nil { return "", errorcategory.New(errorcategory.Internal, "plugin upgrade is not configured") @@ -179,23 +194,12 @@ func (m configMigration) upgradeOne(incompatibility plugins.ConfigV2Incompatibil return m.upgradePlugin(incompatibility) } -func (m configMigration) reportUpgradeFailure(incompatibility plugins.ConfigV2Incompatibility) { - if incompatibility.MinimumVersion != "" { - fmt.Fprintf(m.out, "! could not upgrade %s to the minimum required version (%s).\n", incompatibility.Plugin, incompatibility.MinimumVersion) - } else { - fmt.Fprintf(m.out, "! could not upgrade %s to a version that reads the new config format.\n", incompatibility.Plugin) - } - - fmt.Fprintf(m.out, " run `%s`, then try again.\n", incompatibility.UpgradeCommand()) - fmt.Fprintln(m.out, "your config file was not changed.") -} - -func (m configMigration) migrateAndReload() { +func (m configMigration) migrateAndReload(logger *log.Entry) { changed, err := m.migrate(m.profilesFile) if err != nil { - fmt.Fprintf(m.out, "Could not update %s to the new config format: %s\n", m.profilesFile, err) - fmt.Fprintln(m.out, "The file was left as it was, and the CLI still reads it.") - + // MigrateConfigFile leaves the original in place when it cannot finish, and + // this CLI reads that layout, so the command the user typed is unaffected. + logger.Debugf("Could not update %s to the new config format, leaving it as it was: %s", m.profilesFile, err) return } @@ -204,11 +208,11 @@ func (m configMigration) migrateAndReload() { } if err := m.reload(); err != nil { - fmt.Fprintf(m.out, "Updated %s to the new config format, but could not re-read it: %s\n", m.profilesFile, err) + logger.Debugf("Updated %s to the new config format but could not re-read it: %s", m.profilesFile, err) return } - fmt.Fprintf(m.out, "✔ updated %s to the new config format (backup saved to %s)\n", m.profilesFile, filepath.Base(m.profilesFile+config.ConfigBackupSuffix)) + logger.Debugf("Updated %s to the new config format, backup saved to %s", m.profilesFile, filepath.Base(m.profilesFile+config.ConfigBackupSuffix)) } // upgradePluginForConfigV2 installs the latest release of a plugin that is too diff --git a/pkg/cmd/config_migration_test.go b/pkg/cmd/config_migration_test.go index 999f7da82..0228582bb 100644 --- a/pkg/cmd/config_migration_test.go +++ b/pkg/cmd/config_migration_test.go @@ -22,6 +22,8 @@ type migrationHarness struct { reloaded bool migratedPath string upgrades []string + stamped bool + stampedPath string } func newMigrationHarness(t *testing.T) *migrationHarness { @@ -32,11 +34,10 @@ func newMigrationHarness(t *testing.T) *migrationHarness { h := &migrationHarness{out: &bytes.Buffer{}} h.migration = configMigration{ - profilesFile: profilesFile, - needsMigration: func() bool { return true }, - pluginsReady: func() bool { return true }, - incompatibilities: func() ([]plugins.ConfigV2Incompatibility, error) { return nil, nil }, - installedPluginCount: func() int { return 3 }, + profilesFile: profilesFile, + needsMigration: func() bool { return true }, + pluginsReady: func() bool { return true }, + incompatibilities: func() ([]plugins.ConfigV2Incompatibility, error) { return nil, nil }, upgradePlugin: func(incompatibility plugins.ConfigV2Incompatibility) (string, error) { h.upgrades = append(h.upgrades, incompatibility.Plugin) return "1.2.0", nil @@ -47,6 +48,12 @@ func newMigrationHarness(t *testing.T) *migrationHarness { return true, nil }, + stampNew: func(path string) error { + h.stamped = true + h.stampedPath = path + + return nil + }, reload: func() error { h.reloaded = true @@ -58,6 +65,9 @@ func newMigrationHarness(t *testing.T) *migrationHarness { return h } +// The common case, and the whole point of the quiet contract: the file is +// reorganized and the user sees nothing about it. They asked to run a command, +// nothing about it changed, and there is nothing for them to decide. func TestConfigMigrationRunsWhenNeeded(t *testing.T) { h := newMigrationHarness(t) @@ -66,9 +76,7 @@ func TestConfigMigrationRunsWhenNeeded(t *testing.T) { require.True(t, h.migrated) require.True(t, h.reloaded) require.Equal(t, h.migration.profilesFile, h.migratedPath) - require.Contains(t, h.out.String(), "checking installed plugins... all 3 are compatible.") - require.Contains(t, h.out.String(), "✔ updated "+h.migration.profilesFile+" to the new config format") - require.Contains(t, h.out.String(), "backup saved to config.toml"+config.ConfigBackupSuffix) + require.Empty(t, h.out.String()) } func TestNewConfigMigrationUsesEffectiveConfigPath(t *testing.T) { @@ -126,9 +134,10 @@ func TestConfigMigrationUpgradesAPluginThatIsTooOld(t *testing.T) { require.Equal(t, []string{"projects"}, h.upgrades) require.True(t, h.migrated) - require.Contains(t, h.out.String(), "checking installed plugins...") - require.Contains(t, h.out.String(), "✔ upgraded projects from v0.8.2 to v1.2.0.") - require.Contains(t, h.out.String(), "✔ updated "+h.migration.profilesFile+" to the new config format") + + // The one thing worth saying out loud, and it is said before the download rather + // than reported after it, because the point is to explain the wait. + require.Equal(t, "Upgrading the projects plugin so it can read the updated config file.\n", h.out.String()) } func TestConfigMigrationDoesNotMigrateWhenPluginUpgradeFails(t *testing.T) { @@ -147,9 +156,7 @@ func TestConfigMigrationDoesNotMigrateWhenPluginUpgradeFails(t *testing.T) { h.migration.run() require.False(t, h.migrated) - require.Contains(t, h.out.String(), "! could not upgrade projects to the minimum required version (1.2.0).") - require.Contains(t, h.out.String(), "run `stripe plugin upgrade projects`, then try again.") - require.Contains(t, h.out.String(), "your config file was not changed.") + require.Equal(t, "Upgrading the projects plugin so it can read the updated config file.\n", h.out.String()) } func TestConfigMigrationSkipsUntilPluginVersionsAreKnown(t *testing.T) { @@ -163,8 +170,8 @@ func TestConfigMigrationSkipsUntilPluginVersionsAreKnown(t *testing.T) { } // A migration that fails has already restored the original file, and the command -// the user asked for still runs. -func TestConfigMigrationReportsFailureWithoutFailingTheCommand(t *testing.T) { +// the user asked for still runs against it -- so there is nothing to report. +func TestConfigMigrationStaysQuietWhenTheMigrationFails(t *testing.T) { h := newMigrationHarness(t) h.migration.migrate = func(string) (bool, error) { return false, os.ErrPermission @@ -173,8 +180,7 @@ func TestConfigMigrationReportsFailureWithoutFailingTheCommand(t *testing.T) { h.migration.run() require.False(t, h.reloaded) - require.Contains(t, h.out.String(), "Could not update") - require.Contains(t, h.out.String(), "still reads it") + require.Empty(t, h.out.String()) } // Help and completion output is read by other programs, so a status line in the @@ -195,3 +201,68 @@ func TestMigrationSafeCommand(t *testing.T) { require.False(t, migrationSafeCommand(completionZsh)) require.False(t, migrationSafeCommand(help)) } + +// A config file that does not exist yet has nothing to migrate, but it still has a +// layout to choose. Recording v2 now means the first write -- usually a login -- +// lands in the new layout, instead of writing the flat layout and migrating it on +// the next command, which would leave a backup holding a copy of the credentials +// just written. +func TestRunStampsAConfigFileThatDoesNotExistYet(t *testing.T) { + h := newMigrationHarness(t) + h.migration.profilesFile = filepath.Join(t.TempDir(), "absent", "config.toml") + + h.migration.run() + + require.True(t, h.stamped, "a new config file should record the v2 layout") + require.Equal(t, h.migration.profilesFile, h.stampedPath) + require.True(t, h.reloaded, "viper has to see the stamped file") + require.False(t, h.migrated, "there is nothing to migrate") + + // A brand-new user has not run anything yet; a status line here would be the + // first thing they ever see from the CLI. + require.Empty(t, h.out.String()) +} + +// Gated for the same reason the migration is: once a profile is written under the +// profiles table, a plugin too old to look there cannot find it. Nothing is at +// risk, so the file just keeps the flat layout. +func TestRunDoesNotStampWhenAPluginCannotReadV2(t *testing.T) { + h := newMigrationHarness(t) + h.migration.profilesFile = filepath.Join(t.TempDir(), "absent", "config.toml") + h.migration.incompatibilities = func() ([]plugins.ConfigV2Incompatibility, error) { + return []plugins.ConfigV2Incompatibility{{ + Plugin: "apps", + InstalledVersion: "1.19.0", + }}, nil + } + + h.migration.run() + + require.False(t, h.stamped) + require.Empty(t, h.out.String(), "nothing to upgrade for, so say nothing") +} + +func TestRunDoesNotStampBeforeAnyPluginReleaseIsKnown(t *testing.T) { + h := newMigrationHarness(t) + h.migration.profilesFile = filepath.Join(t.TempDir(), "absent", "config.toml") + h.migration.pluginsReady = func() bool { return false } + + h.migration.run() + + require.False(t, h.stamped) +} + +func TestRunWritesNothingInV2WhileTheMinimumVersionMapIsEmpty(t *testing.T) { + profilesFile := filepath.Join(t.TempDir(), "absent", "config.toml") + migration := newConfigMigration(&config.Config{ProfilesFile: profilesFile}, t.Context()) + + require.False(t, migration.pluginsReady(), "configV2MinimumVersions is empty, so the gate has to be closed") + + // The only stub, and it opens a gate rather than closing one: there is no config + // file here, so this is what gets run() as far as the stamp. + migration.needsMigration = func() bool { return true } + + migration.run() + + require.NoFileExists(t, profilesFile) +} diff --git a/pkg/cmd/plugin/auto_update.go b/pkg/cmd/plugin/auto_update.go index 258de3c84..aa641af40 100644 --- a/pkg/cmd/plugin/auto_update.go +++ b/pkg/cmd/plugin/auto_update.go @@ -1,7 +1,9 @@ package plugin import ( + "fmt" "slices" + "strings" "github.com/spf13/cobra" @@ -65,5 +67,36 @@ func (ac *AutoUpdateCmd) run(cmd *cobra.Command, args []string) error { value = config.PluginConfigOn } - return ac.cfg.WriteConfigField(config.PluginConfigKey(scope, config.PluginConfigUpdatesField), value) + if err := ac.cfg.WriteConfigField(config.PluginConfigKey(scope, config.PluginConfigUpdatesField), value); err != nil { + return err + } + + ac.printSettings(cmd, scope) + return nil +} + +func (ac *AutoUpdateCmd) printSettings(cmd *cobra.Command, scope string) { + action := "Enable" + state := "disabled" + if ac.enable { + action = "Disable" + state = "enabled" + } + + out := cmd.OutOrStdout() + if scope == config.PluginConfigGlobalScope { + fmt.Fprintf(out, "Automatic updates are %s for all plugins\n\n", state) + fmt.Fprintf(out, "%s them with 'stripe plugin auto-update --%s'\n", action, strings.ToLower(action)) + return + } + + pluginName := strings.ToUpper(scope[:1]) + scope[1:] + globalState := "disabled" + if config.PluginUpdatesEnabled("") { + globalState = "enabled" + } + + fmt.Fprintf(out, "Automatic updates are %s for the %s plugin\n\n", state, pluginName) + fmt.Fprintf(out, "%s it with 'stripe plugin auto-update %s --%s'\n", action, scope, strings.ToLower(action)) + fmt.Fprintf(out, "Follow the global setting with 'stripe plugin auto-update %s --unset' (current: %s)\n", scope, globalState) } diff --git a/pkg/cmd/plugin/auto_update_test.go b/pkg/cmd/plugin/auto_update_test.go index c7f4c71e7..28c19af67 100644 --- a/pkg/cmd/plugin/auto_update_test.go +++ b/pkg/cmd/plugin/auto_update_test.go @@ -1,6 +1,7 @@ package plugin import ( + "bytes" "path/filepath" "testing" @@ -38,10 +39,13 @@ func TestGlobalEnable(t *testing.T) { ac := NewAutoUpdateCmd(cfg) ac.enable = true + var output bytes.Buffer + ac.Cmd.SetOut(&output) err := ac.run(ac.Cmd, []string{}) require.NoError(t, err) assert.Equal(t, "on", viper.GetString(config.PluginConfigKey(config.PluginConfigGlobalScope, config.PluginConfigUpdatesField))) + assert.Equal(t, "Automatic updates are enabled for all plugins\n\nDisable them with 'stripe plugin auto-update --disable'\n", output.String()) } // -- global --disable ------------------------------------------------------- @@ -52,10 +56,13 @@ func TestGlobalDisable(t *testing.T) { ac := NewAutoUpdateCmd(cfg) ac.disable = true + var output bytes.Buffer + ac.Cmd.SetOut(&output) err := ac.run(ac.Cmd, []string{}) require.NoError(t, err) assert.Equal(t, "off", viper.GetString(config.PluginConfigKey(config.PluginConfigGlobalScope, config.PluginConfigUpdatesField))) + assert.Equal(t, "Automatic updates are disabled for all plugins\n\nEnable them with 'stripe plugin auto-update --enable'\n", output.String()) } // -- no flags → help -------------------------------------------------------- @@ -78,13 +85,17 @@ func TestPluginEnable(t *testing.T) { defer cleanup() require.NoError(t, cfg.WriteConfigField("installed_plugins", []string{"apps"})) + require.NoError(t, cfg.WriteConfigField(config.PluginConfigKey(config.PluginConfigGlobalScope, config.PluginConfigUpdatesField), config.PluginConfigOn)) ac := NewAutoUpdateCmd(cfg) ac.enable = true + var output bytes.Buffer + ac.Cmd.SetOut(&output) err := ac.run(ac.Cmd, []string{"apps"}) require.NoError(t, err) assert.Equal(t, "on", viper.GetString(config.PluginConfigKey("apps", config.PluginConfigUpdatesField))) + assert.Equal(t, "Automatic updates are enabled for the Apps plugin\n\nDisable it with 'stripe plugin auto-update apps --disable'\nFollow the global setting with 'stripe plugin auto-update apps --unset' (current: enabled)\n", output.String()) } // -- per-plugin --disable --------------------------------------------------- @@ -97,10 +108,13 @@ func TestPluginDisable(t *testing.T) { ac := NewAutoUpdateCmd(cfg) ac.disable = true + var output bytes.Buffer + ac.Cmd.SetOut(&output) err := ac.run(ac.Cmd, []string{"apps"}) require.NoError(t, err) assert.Equal(t, "off", viper.GetString(config.PluginConfigKey("apps", config.PluginConfigUpdatesField))) + assert.Equal(t, "Automatic updates are disabled for the Apps plugin\n\nEnable it with 'stripe plugin auto-update apps --enable'\nFollow the global setting with 'stripe plugin auto-update apps --unset' (current: disabled)\n", output.String()) } // -- per-plugin not installed ----------------------------------------------- diff --git a/pkg/cmd/plugin_cmds.go b/pkg/cmd/plugin_cmds.go index b47a5d03b..84acbee22 100644 --- a/pkg/cmd/plugin_cmds.go +++ b/pkg/cmd/plugin_cmds.go @@ -35,7 +35,10 @@ type pluginTemplateCmd struct { fs afero.Fs ParsedArgs []string - runPluginCmdFn func(cmd *cobra.Command, args []string) error + // runPluginCmdFn hands off to the plugin binary. skipAutoUpgrade leaves out the + // pre-run upgrade check, for a handoff that has nothing to gain from it; see + // runPluginCmd. + runPluginCmdFn func(cmd *cobra.Command, args []string, skipAutoUpgrade bool) error } // newPluginTemplateCmd is a generic plugin command template to dynamically use @@ -53,7 +56,7 @@ func newPluginTemplateCmd(config *config.Config, plugin *plugins.Plugin) *plugin RunE: func(cmd *cobra.Command, args []string) error { // "stripe [host_flags...] plugin_name [plugin_subcommands...] [plugin_flags...]" => "[plugin_subcommands...] [plugin_flags...]" pluginArgs := cmdutil.ArgsAfter(os.Args, cmd.Name()) - return ptc.runPluginCmdFn(cmd, pluginArgs) + return ptc.runPluginCmdFn(cmd, pluginArgs, false) }, Annotations: map[string]string{"scope": "plugin"}, FParseErrWhitelist: cobra.FParseErrWhitelist{ @@ -73,7 +76,10 @@ func newPluginTemplateCmd(config *config.Config, plugin *plugins.Plugin) *plugin // "stripe plugin_name [plugin_subcommands...] --help" => "[plugin_subcommands...] --help" args = cmdutil.ArgsAfter(s, c.Name()) } - ptc.runPluginCmdFn(c, args) + // Asking what a command does should not install software, hence the skip. Cobra + // passes this func down to every subcommand stub too, so `stripe plugin_name + // subcommand --help` lands here as well and is covered by the same decision. + ptc.runPluginCmdFn(c, args, true) }) // Add subcommand stubs from manifest metadata so they appear in --map and help @@ -95,7 +101,7 @@ func addPluginSubcommandStubs(parent *cobra.Command, commands []plugins.CommandI Short: ci.Desc, RunE: func(cmd *cobra.Command, args []string) error { pluginArgs := cmdutil.ArgsAfter(os.Args, ptc.cmd.Name()) - return ptc.runPluginCmdFn(cmd, pluginArgs) + return ptc.runPluginCmdFn(cmd, pluginArgs, false) }, Annotations: map[string]string{"scope": "plugin"}, FParseErrWhitelist: cobra.FParseErrWhitelist{ @@ -123,6 +129,13 @@ func explicitFlagValue(cmd *cobra.Command, name, value string) string { // user's original command. It reuses the template command so a just-installed // plugin gets the same execution, version-check, and exit-code handling as one // that was already present when the CLI started. +// +// The one thing it does differently is skip the upgrade check, because the install it +// follows resolved the newest release moments ago -- there is nothing newer to find, and +// this holds whether the plugin was installed to run a command or to print its own help. +// It is an invariant of this function's contract rather than of its callers: anything +// reaching a plugin that was not just installed should go through the template command's +// normal path instead. func runPluginByName(cmd *cobra.Command, name string, args []string) error { fs := afero.NewOsFs() @@ -131,11 +144,14 @@ func runPluginByName(cmd *cobra.Command, name string, args []string) error { return err } - return newPluginTemplateCmd(&Config, &plugin).runPluginCmd(cmd, args) + return newPluginTemplateCmd(&Config, &plugin).runPluginCmd(cmd, args, true) } -// runPluginCmd hands off to the plugin itself to take over -func (ptc *pluginTemplateCmd) runPluginCmd(cmd *cobra.Command, args []string) error { +// runPluginCmd hands off to the plugin itself to take over. +// +// skipAutoUpgrade leaves out the pre-run upgrade check; see +// Plugin.RunWithoutAutoUpgrade for which handoffs want that and why. +func (ptc *pluginTemplateCmd) runPluginCmd(cmd *cobra.Command, args []string, skipAutoUpgrade bool) error { ctx := withSIGTERMCancel(commandContextOrBackground(cmd), func() { log.WithFields(log.Fields{ "prefix": "cmd.pluginCmd.runPluginCmd", @@ -182,7 +198,12 @@ func (ptc *pluginTemplateCmd) runPluginCmd(cmd *cobra.Command, args []string) er return err } - err = plugin.Run(ctx, ptc.cfg, fs, ptc.ParsedArgs, "", "", + run := plugin.Run + if skipAutoUpgrade { + run = plugin.RunWithoutAutoUpgrade + } + + err = run(ctx, ptc.cfg, fs, ptc.ParsedArgs, "", "", explicitFlagValue(cmd, "api-base", apiBaseURL), explicitFlagValue(cmd, "dashboard-base", rawDashboardBaseURL), explicitFlagValue(cmd, "access-base", accessBaseURL)) @@ -201,8 +222,18 @@ func (ptc *pluginTemplateCmd) runPluginCmd(cmd *cobra.Command, args []string) er "prefix": "pluginTemplateCmd.runPluginCmd", }).Debug(fmt.Sprintf("Plugin command '%s' exited with error: %s", plugin.Shortname, err)) - // We can't return err because the plugin will have already printed the error message at - // this point, and we can't return nil because the host will exit with code 0. + // A plugin that started and then failed has already printed why, so printing + // it again would duplicate it. Anything else failed before the plugin was + // ever launched -- an install that failed, a plugin too old to read the + // config file, a handshake that never completed -- and nothing has printed + // it, so exiting silently here is the difference between an actionable + // message and a bare exit code 1. + if !plugins.PluginAlreadyReported(err) { + fmt.Fprintln(os.Stderr, err) + } + + // We can't return err because it is either already printed or printed just + // above, and we can't return nil because the host would exit with code 0. os.Exit(1) } diff --git a/pkg/cmd/plugin_cmds_test.go b/pkg/cmd/plugin_cmds_test.go index 3c257e22e..0e7ed9ca6 100644 --- a/pkg/cmd/plugin_cmds_test.go +++ b/pkg/cmd/plugin_cmds_test.go @@ -186,9 +186,11 @@ func TestGeneratedPluginSubcommandStubUsesExecutingCommand(t *testing.T) { var capturedCmd *cobra.Command var capturedArgs []string - ptc.runPluginCmdFn = func(cmd *cobra.Command, args []string) error { + var capturedSkipAutoUpgrade bool + ptc.runPluginCmdFn = func(cmd *cobra.Command, args []string, skipAutoUpgrade bool) error { capturedCmd = cmd capturedArgs = append([]string(nil), args...) + capturedSkipAutoUpgrade = skipAutoUpgrade return sentinelErr } @@ -207,6 +209,91 @@ func TestGeneratedPluginSubcommandStubUsesExecutingCommand(t *testing.T) { assert.Equal(t, "catalog", capturedCmd.Name()) assert.Equal(t, []string{"catalog"}, capturedArgs) assert.Equal(t, "sentinel", capturedCmd.Context().Value(ctxKey{})) + assert.False(t, capturedSkipAutoUpgrade, "a subcommand the user asked to run gets the upgrade check") +} + +// TestPluginHelpDoesNotRunAsACommand covers every way the user can ask a plugin for help. +// All of them run the plugin binary, because the plugin owns its help text, and none of +// them should be treated as the user asking the plugin to do work -- most concretely, +// none should trigger an auto-upgrade. See Plugin.RunWithoutAutoUpgrade. +func TestPluginHelpDoesNotRunAsACommand(t *testing.T) { + tests := []struct { + name string + argv []string + wantArgs []string + }{ + { + name: "--help on the plugin", + argv: []string{"stripe", "projects", "--help"}, + wantArgs: []string{"--help"}, + }, + { + name: "-h on the plugin", + argv: []string{"stripe", "projects", "-h"}, + wantArgs: []string{"-h"}, + }, + { + // Cobra hands its help func down to subcommands, so the stubs are covered + // by the same opt-out rather than needing one of their own. + // + // The subcommand name being dropped here is a separate, pre-existing bug -- + // the help func slices argv after the command it was called on, which for a + // stub is the subcommand itself -- so the plugin prints its top-level help. + // Pinned as-is rather than fixed, since what this test is about is that the + // handoff happens without an upgrade. + name: "--help on a subcommand stub", + argv: []string{"stripe", "projects", "catalog", "--help"}, + wantArgs: []string{"--help"}, + }, + { + // The other spelling, which arrives through cobra's own help command with + // no args of its own, so the help func rebuilds them from os.Args. + name: "the help command", + argv: []string{"stripe", "help", "projects"}, + wantArgs: []string{"--help"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + plugin := plugins.Plugin{ + Shortname: "projects", + Shortdesc: "Projects plugin", + Binary: "stripe-cli-projects", + MagicCookieValue: "magic", + Commands: []plugins.CommandInfo{ + {Name: "catalog", Desc: "Browse the projects catalog"}, + }, + } + + ptc := newPluginTemplateCmd(&Config, &plugin) + + var calls int + var capturedArgs []string + var capturedSkipAutoUpgrade bool + ptc.runPluginCmdFn = func(cmd *cobra.Command, args []string, skipAutoUpgrade bool) error { + calls++ + capturedArgs = append([]string(nil), args...) + capturedSkipAutoUpgrade = skipAutoUpgrade + return nil + } + + root := &cobra.Command{Use: "stripe"} + root.AddCommand(ptc.cmd) + + oldArgs := os.Args + os.Args = tt.argv + defer func() { os.Args = oldArgs }() + + root.SetArgs(tt.argv[1:]) + require.NoError(t, root.ExecuteContext(context.Background())) + + require.Equal(t, 1, calls, "the plugin should be asked for its help text exactly once") + assert.Equal(t, tt.wantArgs, capturedArgs) + require.True(t, capturedSkipAutoUpgrade, + "asking what a command does would auto-upgrade the plugin, and could install a new binary") + }) + } } func TestCommandContextOrBackgroundUsesCommandContext(t *testing.T) { diff --git a/pkg/cmd/root.go b/pkg/cmd/root.go index f4f0d76a5..fa4633c52 100644 --- a/pkg/cmd/root.go +++ b/pkg/cmd/root.go @@ -60,7 +60,6 @@ var rootCmd = &cobra.Command{ "trigger": "webhooks", "listen": "webhooks", "logs": "stripe", - "status": "stripe", "resources": "resources", AIAgentHelpAnnotationKey: " If you do not have an account, run `stripe sandbox create` (provisions a claimable sandbox without a browser).\n" + " Visit https://docs.stripe.com/llms.txt?utm_source=cli for latest guidance on how to integrate correctly.\n" + @@ -362,6 +361,7 @@ func init() { rootCmd.AddCommand(newResourcesCmd().cmd) rootCmd.AddCommand(newSamplesCmd()) rootCmd.AddCommand(newServeCmd()) + rootCmd.AddCommand(newStatusCmd()) rootCmd.AddCommand(newSwitchCmd().cmd) rootCmd.AddCommand(newTriggerCmd().cmd) rootCmd.AddCommand(newVersionCmd().cmd) diff --git a/pkg/cmd/samples.go b/pkg/cmd/samples.go index b878ef57d..3856fcac6 100644 --- a/pkg/cmd/samples.go +++ b/pkg/cmd/samples.go @@ -8,14 +8,14 @@ import ( "github.com/stripe/stripe-cli/pkg/errorcategory" ) -// errCommandRemoved is returned by the shims for commands removed in v1.60.0. +// errCommandRemoved is returned by the shims for commands removed in v1.51.0. // root.go recognizes this sentinel to suppress duplicate error output and error // reporting while still exiting non-zero, so callers that scripted a removed // command see a failure rather than a silent no-op. var errCommandRemoved = errorcategory.New(errorcategory.UserInput, "command removed") func deprecatedCommandMessage(command string) string { - return fmt.Sprintf("The `%s` command is no longer available in Stripe CLI v1.60.0 and later. To use it, install a version earlier than v1.60.0.", command) + return fmt.Sprintf("The `%s` command is no longer available in Stripe CLI v1.51.0 and later. To use it, install a version earlier than v1.51.0.", command) } func newDeprecatedCommand(use, command string) *cobra.Command { @@ -50,3 +50,7 @@ func newServeCmd() *cobra.Command { cmd.Aliases = []string{"srv"} return cmd } + +func newStatusCmd() *cobra.Command { + return newDeprecatedCommand("status", "stripe status") +} diff --git a/pkg/cmd/samples_test.go b/pkg/cmd/samples_test.go index f468f4a7e..435b943c7 100644 --- a/pkg/cmd/samples_test.go +++ b/pkg/cmd/samples_test.go @@ -14,7 +14,7 @@ func newRootWithDeprecatedCmds() *cobra.Command { // Mirror the real root's silencing so the shims' output is what a user // actually sees; without it cobra appends its own error and usage block. root := &cobra.Command{Use: "stripe", SilenceUsage: true, SilenceErrors: true} - root.AddCommand(newSamplesCmd(), newServeCmd()) + root.AddCommand(newSamplesCmd(), newServeCmd(), newStatusCmd()) return root } @@ -27,17 +27,22 @@ func TestDeprecatedCommands(t *testing.T) { { name: "samples", args: []string{"samples", "create", "checkout", "--force", "--integration", "react"}, - expected: "The `stripe samples` command is no longer available in Stripe CLI v1.60.0 and later. To use it, install a version earlier than v1.60.0.\n", + expected: "The `stripe samples` command is no longer available in Stripe CLI v1.51.0 and later. To use it, install a version earlier than v1.51.0.\n", }, { name: "serve", args: []string{"serve", ".", "--port", "8080"}, - expected: "The `stripe serve` command is no longer available in Stripe CLI v1.60.0 and later. To use it, install a version earlier than v1.60.0.\n", + expected: "The `stripe serve` command is no longer available in Stripe CLI v1.51.0 and later. To use it, install a version earlier than v1.51.0.\n", }, { name: "serve alias", args: []string{"srv", "."}, - expected: "The `stripe serve` command is no longer available in Stripe CLI v1.60.0 and later. To use it, install a version earlier than v1.60.0.\n", + expected: "The `stripe serve` command is no longer available in Stripe CLI v1.51.0 and later. To use it, install a version earlier than v1.51.0.\n", + }, + { + name: "status", + args: []string{"status"}, + expected: "The `stripe status` command is no longer available in Stripe CLI v1.51.0 and later. To use it, install a version earlier than v1.51.0.\n", }, } @@ -62,27 +67,37 @@ func TestDeprecatedCommandHelp(t *testing.T) { { name: "samples", args: []string{"help", "samples"}, - expected: "The `stripe samples` command is no longer available in Stripe CLI v1.60.0 and later. To use it, install a version earlier than v1.60.0.\n", + expected: "The `stripe samples` command is no longer available in Stripe CLI v1.51.0 and later. To use it, install a version earlier than v1.51.0.\n", }, { name: "samples help flag", args: []string{"samples", "--help"}, - expected: "The `stripe samples` command is no longer available in Stripe CLI v1.60.0 and later. To use it, install a version earlier than v1.60.0.\n", + expected: "The `stripe samples` command is no longer available in Stripe CLI v1.51.0 and later. To use it, install a version earlier than v1.51.0.\n", }, { name: "samples help shorthand", args: []string{"samples", "-h"}, - expected: "The `stripe samples` command is no longer available in Stripe CLI v1.60.0 and later. To use it, install a version earlier than v1.60.0.\n", + expected: "The `stripe samples` command is no longer available in Stripe CLI v1.51.0 and later. To use it, install a version earlier than v1.51.0.\n", }, { name: "serve", args: []string{"help", "serve"}, - expected: "The `stripe serve` command is no longer available in Stripe CLI v1.60.0 and later. To use it, install a version earlier than v1.60.0.\n", + expected: "The `stripe serve` command is no longer available in Stripe CLI v1.51.0 and later. To use it, install a version earlier than v1.51.0.\n", }, { name: "serve help flag", args: []string{"serve", "--help"}, - expected: "The `stripe serve` command is no longer available in Stripe CLI v1.60.0 and later. To use it, install a version earlier than v1.60.0.\n", + expected: "The `stripe serve` command is no longer available in Stripe CLI v1.51.0 and later. To use it, install a version earlier than v1.51.0.\n", + }, + { + name: "status", + args: []string{"help", "status"}, + expected: "The `stripe status` command is no longer available in Stripe CLI v1.51.0 and later. To use it, install a version earlier than v1.51.0.\n", + }, + { + name: "status help flag", + args: []string{"status", "--help"}, + expected: "The `stripe status` command is no longer available in Stripe CLI v1.51.0 and later. To use it, install a version earlier than v1.51.0.\n", }, } @@ -104,6 +119,7 @@ func TestDeprecatedCommandsHidden(t *testing.T) { require.NoError(t, err) require.NotContains(t, rootHelp, "samples") require.NotContains(t, rootHelp, "serve") + require.NotContains(t, rootHelp, "status") } // The tests above build their own root, so they can't catch the shims falling @@ -113,6 +129,7 @@ func TestDeprecatedCommandsRegisteredOnRoot(t *testing.T) { {"samples"}, {"serve"}, {"srv"}, + {"status"}, } for _, path := range paths { diff --git a/pkg/config/migrate.go b/pkg/config/migrate.go index 7dff0b4e2..f55c030e5 100644 --- a/pkg/config/migrate.go +++ b/pkg/config/migrate.go @@ -91,7 +91,8 @@ func NeedsMigration() bool { // migration itself always agree on whether there is work to do. func hasFlatProfileTable(v *viper.Viper) bool { for key, value := range v.AllSettings() { - if reservedTopLevelKeys[key] { + if key == ProfilesTableName { + // The v2 container: everything inside it is already migrated. continue } @@ -203,6 +204,8 @@ func planMigration(contents []byte) (*migrationPlan, bool, error) { alreadyV2 := version == ConfigVersionV2 // Start from the profiles already in the v2 table, if there are any. + // nestedIsContainer tells the v2 container apart from a v1 profile named + // "profiles"; only the container holds profiles that are already in place. nested, nestedIsContainer := profilesContainer(doc, alreadyV2) if nestedIsContainer { for name, value := range nested { @@ -250,15 +253,16 @@ func planMigration(contents []byte) (*migrationPlan, bool, error) { // reservedProfileNameException reports whether a top-level key must be treated // as a setting rather than as a profile. Only the profiles table itself // qualifies, and only when it is acting as the v2 container. +// +// Every other key is judged on contents rather than on its name: a table holding +// profile fields is a profile even when a settings key already owns that name, and +// separating the two is what the v2 layout is for. func reservedProfileNameException(key string, nestedIsContainer bool) bool { - if key == ProfilesTableName { - return nestedIsContainer - } - - return reservedTopLevelKeys[key] + return key == ProfilesTableName && nestedIsContainer } -// profilesContainer returns the v2 profiles table, if the document has one. +// profilesContainer returns the v2 profiles table and whether the document has +// one at all. // // A v1 file can contain a profile literally named "profiles", which occupies the // same key as the v2 container: nothing stops `stripe login --project-name @@ -480,3 +484,36 @@ func writeAndSync(file *os.File, contents []byte) error { return file.Sync() } + +// StampNewConfigFile creates a config file that records the v2 layout, so that the +// first write into it nests profiles under the profiles table instead of writing +// the flat layout and needing a migration on the next command. +// +// The document is the one encodePlan produces for an empty plan, so a file created +// here is indistinguishable from one the migration would have produced. +// +// It refuses to touch a file that already exists: picking the layout of a file that +// already holds something is MigrateConfigFile's job, and it has a backup and a +// verification pass for exactly that reason. +func StampNewConfigFile(path string) error { + if _, err := os.Stat(path); err == nil { + return errorcategory.Errorf(errorcategory.Filesystem, + "refusing to stamp %s: the file already exists", path) + } else if !os.IsNotExist(err) { + return err + } + + if err := makePath(path); err != nil { + return err + } + + contents, err := encodePlan(&migrationPlan{ + profiles: make(map[string]map[string]interface{}), + settings: make(map[string]interface{}), + }) + if err != nil { + return err + } + + return writeFileSync(path, contents) +} diff --git a/pkg/config/migrate_test.go b/pkg/config/migrate_test.go index 1c9ba3882..33fdfe8b6 100644 --- a/pkg/config/migrate_test.go +++ b/pkg/config/migrate_test.go @@ -396,3 +396,121 @@ func TestVerifyPlanCatchesADroppedSetting(t *testing.T) { err := verifyPlan(plan, []byte("config_version = 2\n\n[profiles]\n")) require.ErrorContains(t, err, "is missing from the migrated config") } + +func TestStampNewConfigFileWritesTheV2Layout(t *testing.T) { + path := filepath.Join(t.TempDir(), "nested", "config.toml") + + require.NoError(t, StampNewConfigFile(path)) + + v := viper.New() + v.SetConfigFile(path) + require.NoError(t, v.ReadInConfig()) + require.Equal(t, ConfigVersionV2, v.GetInt(ConfigVersionName)) + require.True(t, isMigrated(v), "a write into this file has to nest") + + // Windows has no Unix permission bits, so the mode the file was created with + // does not survive a Stat there. + if runtime.GOOS != "windows" { + info, err := os.Stat(path) + require.NoError(t, err) + require.Equal(t, os.FileMode(0600), info.Mode().Perm()) + } +} + +// Choosing the layout of a file that already holds something is +// MigrateConfigFile's job: it takes a backup and verifies the result first. +func TestStampNewConfigFileRefusesAnExistingFile(t *testing.T) { + path := writeConfigFileForMigration(t, "[default]\n display_name = 'Acme'\n") + + require.Error(t, StampNewConfigFile(path)) + require.Contains(t, string(helperLoadBytes(t, path)), "display_name = 'Acme'") +} + +// `stripe login --project-name installed_plugins` is allowed, so a profile can sit +// on a settings key. Moving it out is what keeps the next plugin install from +// overwriting it. +func TestMigrateConfigFileMovesProfileNamedAfterAReservedKey(t *testing.T) { + for _, name := range []string{"installed_plugins", "plugin_configs", "color", "user_info", "project-name"} { + t.Run(name, func(t *testing.T) { + path := writeConfigFileForMigration(t, "["+name+"]\n"+ + " display_name = 'Collided Account'\n"+ + " test_mode_api_key = 'sk_test_collided_key'\n") + + changed, err := MigrateConfigFile(path) + require.NoError(t, err) + require.True(t, changed) + + v := viper.New() + v.SetConfigFile(path) + require.NoError(t, v.ReadInConfig()) + + require.Equal(t, "sk_test_collided_key", + v.GetString(ProfilesTableName+"."+name+".test_mode_api_key")) + require.False(t, v.IsSet(name+".test_mode_api_key")) + }) + } +} + +// The flip side: a settings key holding its real value has no profile field in it, +// so it stays at the top level. +func TestMigrateConfigFileLeavesRealReservedSettingsAlone(t *testing.T) { + path := writeConfigFileForMigration(t, `installed_plugins = ['apps'] +machine_uuid = 'uuid-reserved' + +[plugin_configs.__global] + updates = 'on' + +[user_info] + compartments = [] + +[default] + display_name = 'Acme' + test_mode_api_key = 'sk_test_acme_key' +`) + + changed, err := MigrateConfigFile(path) + require.NoError(t, err) + require.True(t, changed) + + v := viper.New() + v.SetConfigFile(path) + require.NoError(t, v.ReadInConfig()) + + require.Equal(t, []string{"apps"}, v.GetStringSlice("installed_plugins")) + require.Equal(t, "on", v.GetString("plugin_configs.__global.updates")) + require.True(t, v.IsSet("user_info")) + require.False(t, v.IsSet(ProfilesTableName+".plugin_configs")) + require.False(t, v.IsSet(ProfilesTableName+".user_info")) + require.Equal(t, "sk_test_acme_key", v.GetString(ProfilesTableName+".default.test_mode_api_key")) +} + +// NeedsMigration has to agree with the migration about what counts as work, or the +// file never converges. +func TestNeedsMigrationSeesProfileNamedAfterAReservedKey(t *testing.T) { + setupProfileConfig(t, `config_version = 2 + +[profiles.default] + display_name = 'Acme' + +[installed_plugins] + display_name = 'Collided Account' + test_mode_api_key = 'sk_test_collided_key' +`) + + require.True(t, NeedsMigration()) +} + +// ...while a fully migrated file with only real settings at the top level is done. +func TestNeedsMigrationIsFalseForAMigratedFile(t *testing.T) { + setupProfileConfig(t, `config_version = 2 +installed_plugins = ['apps'] + +[plugin_configs.__global] + updates = 'on' + +[profiles.default] + display_name = 'Acme' +`) + + require.False(t, NeedsMigration()) +} diff --git a/pkg/login/login.go b/pkg/login/login.go index 6912f3acb..b2eab39de 100644 --- a/pkg/login/login.go +++ b/pkg/login/login.go @@ -90,31 +90,103 @@ func InitiateLogin(ctx context.Context, baseURL, accessBaseURL string, cfg *conf return nil } -// initiateOAuthDeviceLogin calls the device authorization endpoint, saves the -// pending state to disk, and prints the JSON session output. -func initiateOAuthDeviceLogin(ctx context.Context, accessBaseURL string) error { +// OAuthLoginSession describes a non-interactive OAuth device-code login that's ready for the +// user to complete out-of-band, whether just minted or resumed from a still-valid pending one. +type OAuthLoginSession struct { + BrowserURL string + VerificationCode string + ExpiresIn int // seconds remaining until the device code expires +} + +// InitiateOAuthLogin starts (or resumes) a non-interactive OAuth device-code login for +// accessBaseURL, without printing anything. Unlike `stripe login --non-interactive` (see +// initiateOAuthDeviceLogin), which always mints a fresh device code, this resumes a +// still-valid pending one instead of minting a new one - see initiateOrResumeOAuthDeviceLogin. +// This is the resume behavior exposed to plugins via the OAuthInitiateLogin RPC. +func InitiateOAuthLogin(ctx context.Context, accessBaseURL string) (*OAuthLoginSession, error) { + return initiateOrResumeOAuthDeviceLogin(ctx, accessBaseURL) +} + +// FindPendingOAuthLogin looks for an OAuth device-code login already in progress for +// accessBaseURL (started by this process or another one, e.g. `stripe login +// --non-interactive`), without starting a new one. Returns nil, nil if there is no pending +// login attempt, it's for a different accessBaseURL, or it has expired. +func FindPendingOAuthLogin(accessBaseURL string) (*OAuthLoginSession, error) { + return findPendingOAuthLoginSession(accessBaseURL), nil +} + +func findPendingOAuthLoginSession(accessBaseURL string) *OAuthLoginSession { + cont, err := loadPendingDeviceAuth() + if err != nil || cont.AccessBaseURL != accessBaseURL { + return nil + } + remaining := time.Until(cont.deadline()) + if remaining <= 0 { + return nil + } + return &OAuthLoginSession{ + BrowserURL: cont.VerificationURI, + VerificationCode: cont.UserCode, + ExpiresIn: int(remaining.Seconds()), + } +} + +// initiateOrResumeOAuthDeviceLogin returns the still-valid pending device code for +// accessBaseURL if one exists, instead of minting a new one - so repeated calls to the +// OAuthInitiateLogin RPC (e.g. a retrying plugin) converge on one browser_url/ +// verification_code rather than orphaning the previous one every retry. `stripe login +// --non-interactive` does not use this - see mintOAuthDeviceLogin. +func initiateOrResumeOAuthDeviceLogin(ctx context.Context, accessBaseURL string) (*OAuthLoginSession, error) { + if session := findPendingOAuthLoginSession(accessBaseURL); session != nil { + return session, nil + } + return mintOAuthDeviceLogin(ctx, accessBaseURL) +} + +// mintOAuthDeviceLogin always requests a fresh device code from accessBaseURL and saves it as +// the new pending state, overwriting any still-valid pending login that may already exist. +func mintOAuthDeviceLogin(ctx context.Context, accessBaseURL string) (*OAuthLoginSession, error) { clientID := clientIDForAccessBaseURL(accessBaseURL) authResp, err := RequestDeviceCode(ctx, accessBaseURL, clientID) if err != nil { - return fmt.Errorf("failed to request device code: %w", err) + return nil, fmt.Errorf("failed to request device code: %w", err) } if err := validateBrowserURL(authResp.VerificationURI, accessBaseURL); err != nil { - return err + return nil, err } cont := &oauthContinuation{ - DeviceCode: authResp.DeviceCode, - Interval: authResp.Interval, - ExpiresIn: authResp.ExpiresIn, - AccessBaseURL: accessBaseURL, + DeviceCode: authResp.DeviceCode, + Interval: authResp.Interval, + ExpiresIn: authResp.ExpiresIn, + AccessBaseURL: accessBaseURL, + VerificationURI: authResp.VerificationURI, + UserCode: authResp.UserCode, + IssuedAt: time.Now(), } if err := savePendingDeviceAuth(cont); err != nil { - return fmt.Errorf("failed to save pending auth state: %w", err) + return nil, fmt.Errorf("failed to save pending auth state: %w", err) } - out := loginSessionOutput{ + return &OAuthLoginSession{ BrowserURL: authResp.VerificationURI, VerificationCode: authResp.UserCode, + ExpiresIn: authResp.ExpiresIn, + }, nil +} + +// initiateOAuthDeviceLogin always mints a fresh device code (see mintOAuthDeviceLogin - unlike +// the OAuthInitiateLogin RPC, this does not resume a still-valid pending login), saves it as +// the pending state, and prints the JSON session output. +func initiateOAuthDeviceLogin(ctx context.Context, accessBaseURL string) error { + session, err := mintOAuthDeviceLogin(ctx, accessBaseURL) + if err != nil { + return err + } + + out := loginSessionOutput{ + BrowserURL: session.BrowserURL, + VerificationCode: session.VerificationCode, NextStep: "stripe login --complete-device", } b, err := json.MarshalIndent(out, "", " ") @@ -153,43 +225,121 @@ func PollForLogin(ctx context.Context, pollURL string, cfg *config.Config) error return nil } -// PollPendingDeviceAuth loads the OAuth device auth state saved by -// InitiateLogin and polls the token endpoint until the user approves. -func PollPendingDeviceAuth(ctx context.Context, cfg *config.Config) error { +// loadValidatedPendingDeviceAuth loads the pending device auth state and checks that its +// AccessBaseURL is one ValidateAccessBaseURL accepts, since requests built from it carry OAuth +// bearer/refresh tokens. +func loadValidatedPendingDeviceAuth() (*oauthContinuation, error) { cont, err := loadPendingDeviceAuth() if err != nil { - return err + return nil, err } - // Remove the pending file whether polling succeeds or fails. - defer clearPendingDeviceAuth() - if err := ValidateAccessBaseURL(cont.AccessBaseURL); err != nil { - return err + return nil, err + } + return cont, nil +} + +// PollPendingOAuthLogin waits for the login started by InitiateOAuthLogin (or InitiateLogin's +// OAuth path) to complete. Returns (nil, nil) if ctx is done before the user completes +// authentication and the device code's own expiry hasn't been reached yet - callers should +// call this again to keep waiting. Returns an error and clears the pending state if the +// device code has actually expired or the server returned a terminal OAuth error (e.g. +// access_denied); on success, also clears the pending state. +func PollPendingOAuthLogin(ctx context.Context, cfg *config.Config) (*DeviceCodeLoginResult, error) { + cont, err := loadValidatedPendingDeviceAuth() + if err != nil { + return nil, err } clientID := clientIDForAccessBaseURL(cont.AccessBaseURL) interval := max(time.Duration(cont.Interval)*time.Second, 5*time.Second) - expiresIn := max(time.Duration(cont.ExpiresIn)*time.Second, 10*time.Minute) + deadline := cont.deadline() - pollCtx, cancel := context.WithTimeout(ctx, expiresIn) + pollCtx, cancel := context.WithDeadline(ctx, deadline) defer cancel() - waitCtx, stop := signal.NotifyContext(pollCtx, os.Interrupt) + + result, err := PollAndSaveDeviceCredentials(pollCtx, cont.AccessBaseURL, clientID, cont.DeviceCode, interval, cfg) + if err != nil { + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + if time.Now().After(deadline) { + clearPendingDeviceAuth() + return nil, errorcategory.Errorf(errorcategory.Auth, "device code expired; please run 'stripe login --non-interactive' again") + } + // The caller's own ctx stopped waiting this time (e.g. ^C or its own timeout); + // the device code is still good, so leave the pending state for a later retry. + return nil, nil + } + // Terminal OAuth error (access_denied, expired_token, ...): the device code is dead. + clearPendingDeviceAuth() + return nil, err + } + + clearPendingDeviceAuth() + return result, nil +} + +// CheckPendingOAuthLogin makes a single, non-blocking check on the login started by +// InitiateOAuthLogin (or InitiateLogin's OAuth path). Unlike PollPendingOAuthLogin, it never +// waits for the user - it makes one request and returns immediately, so callers that want to +// wait (e.g. a plugin driving its own retry loop) should call this repeatedly on their own +// schedule instead. Returns (nil, nil) if the device code is still valid but the user hasn't +// completed authentication yet. Returns an error and clears the pending state if the device +// code has actually expired, the server returned a terminal OAuth error (e.g. access_denied), +// or the token exchange succeeded but saving the resulting credentials failed (the device code +// is already consumed at that point, so it can't be retried); on success, also clears the +// pending state. +func CheckPendingOAuthLogin(ctx context.Context, cfg *config.Config) (*DeviceCodeLoginResult, error) { + cont, err := loadValidatedPendingDeviceAuth() + if err != nil { + return nil, err + } + + if time.Now().After(cont.deadline()) { + clearPendingDeviceAuth() + return nil, errorcategory.Errorf(errorcategory.Auth, "device code expired; please run 'stripe login --non-interactive' again") + } + + clientID := clientIDForAccessBaseURL(cont.AccessBaseURL) + tokenResp, err := CheckDeviceToken(ctx, cont.AccessBaseURL, clientID, cont.DeviceCode) + if err != nil { + var oauthErr *OAuthError + if errors.As(err, &oauthErr) && (oauthErr.Code == "authorization_pending" || oauthErr.Code == "slow_down") { + // The user hasn't completed authentication yet; the device code is still good. + return nil, nil + } + // Terminal OAuth error (access_denied, expired_token, ...): the device code is dead. + clearPendingDeviceAuth() + return nil, err + } + + result, err := saveDeviceCredentials(context.WithoutCancel(ctx), cont.AccessBaseURL, tokenResp, cfg) + if err != nil { + // The device code has already been redeemed for tokenResp, so it can't be used again - + // retrying with the same pending state would just fail with a confusing OAuth error + // instead of surfacing this one. Matches PollPendingOAuthLogin's behavior. + clearPendingDeviceAuth() + return nil, err + } + clearPendingDeviceAuth() + return result, nil +} + +// PollPendingDeviceAuth loads the OAuth device auth state saved by InitiateLogin and polls +// the token endpoint until the user approves. +func PollPendingDeviceAuth(ctx context.Context, cfg *config.Config) error { + waitCtx, stop := signal.NotifyContext(ctx, os.Interrupt) defer stop() s := ansi.StartNewSpinner("Waiting for confirmation...", os.Stdout) - result, err := PollAndSaveDeviceCredentials(waitCtx, cont.AccessBaseURL, clientID, cont.DeviceCode, interval, cfg) + result, err := PollPendingOAuthLogin(waitCtx, cfg) ansi.StopSpinner(s, "", os.Stdout) if err != nil { - switch { - case errors.Is(err, context.Canceled): - ansi.ClearLine(os.Stdout) - fmt.Println("Canceled. Run 'stripe login --non-interactive' again to try again.") - return nil - case pollCtx.Err() != nil: - return errorcategory.Errorf(errorcategory.Auth, "device code expired; please run 'stripe login --non-interactive' again") - default: - return err - } + return err + } + if result == nil { + ansi.ClearLine(os.Stdout) + fmt.Println("Canceled. Run 'stripe login --non-interactive' again to try again.") + return nil } printAuthorizedSummary(result.Accounts, result.ActiveAccountID, result.ActiveLivemode) diff --git a/pkg/login/login_test.go b/pkg/login/login_test.go new file mode 100644 index 000000000..71d352139 --- /dev/null +++ b/pkg/login/login_test.go @@ -0,0 +1,559 @@ +package login + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "net/url" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/stripe/stripe-cli/pkg/config" +) + +// rewriteHostTransport redirects every request to target, regardless of the URL it was built +// with. This lets tests exercise code paths gated by ValidateAccessBaseURL (which only accepts +// the real access.stripe.com/qa-access.stripe.com origins) against a local httptest server, +// without weakening that validation. +type rewriteHostTransport struct { + target *url.URL + base http.RoundTripper +} + +func (t rewriteHostTransport) RoundTrip(req *http.Request) (*http.Response, error) { + req = req.Clone(req.Context()) + req.URL.Scheme = t.target.Scheme + req.URL.Host = t.target.Host + return t.base.RoundTrip(req) +} + +// stubAccessSrv points accessSrvHTTPClient at ts for the duration of the test. +func stubAccessSrv(t *testing.T, ts *httptest.Server) { + t.Helper() + target, err := url.Parse(ts.URL) + require.NoError(t, err) + + orig := accessSrvHTTPClient + accessSrvHTTPClient = &http.Client{ + CheckRedirect: orig.CheckRedirect, + Transport: rewriteHostTransport{target: target, base: http.DefaultTransport}, + } + t.Cleanup(func() { accessSrvHTTPClient = orig }) +} + +func TestInitiateOrResumeOAuthDeviceLogin_MintsWhenNoPending(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + + var mintCount atomic.Int32 + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + mintCount.Add(1) + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(DeviceAuthResponse{ //nolint:errcheck + DeviceCode: "device-code", + UserCode: "ABCD-EFGH", + VerificationURI: "https://access.stripe.com/verify", + ExpiresIn: 300, + Interval: 5, + }) + })) + defer ts.Close() + + session, err := initiateOrResumeOAuthDeviceLogin(context.Background(), ts.URL) + require.NoError(t, err) + assert.Equal(t, "https://access.stripe.com/verify", session.BrowserURL) + assert.Equal(t, "ABCD-EFGH", session.VerificationCode) + assert.Equal(t, int32(1), mintCount.Load()) +} + +func TestInitiateOrResumeOAuthDeviceLogin_ResumesStillValidPending(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + + var mintCount atomic.Int32 + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + mintCount.Add(1) + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(DeviceAuthResponse{ //nolint:errcheck + DeviceCode: "device-code", + UserCode: "ABCD-EFGH", + VerificationURI: "https://access.stripe.com/verify", + ExpiresIn: 300, + Interval: 5, + }) + })) + defer ts.Close() + + first, err := initiateOrResumeOAuthDeviceLogin(context.Background(), ts.URL) + require.NoError(t, err) + + second, err := initiateOrResumeOAuthDeviceLogin(context.Background(), ts.URL) + require.NoError(t, err) + + assert.Equal(t, int32(1), mintCount.Load(), "a second call within the expiry window must not mint a new device code") + assert.Equal(t, first.BrowserURL, second.BrowserURL) + assert.Equal(t, first.VerificationCode, second.VerificationCode) +} + +func TestInitiateOrResumeOAuthDeviceLogin_MintsFreshWhenPendingExpired(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + + var mintCount atomic.Int32 + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + mintCount.Add(1) + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(DeviceAuthResponse{ //nolint:errcheck + DeviceCode: "new-device-code", + UserCode: "NEW1-CODE", + VerificationURI: "https://access.stripe.com/verify", + ExpiresIn: 300, + Interval: 5, + }) + })) + defer ts.Close() + + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{ + DeviceCode: "stale-device-code", + AccessBaseURL: ts.URL, + VerificationURI: "https://access.stripe.com/stale", + UserCode: "STALE-CODE", + ExpiresIn: 1, + IssuedAt: time.Now().Add(-11 * time.Minute), // past the 10-minute deadline floor + })) + + session, err := initiateOrResumeOAuthDeviceLogin(context.Background(), ts.URL) + require.NoError(t, err) + assert.Equal(t, int32(1), mintCount.Load()) + assert.Equal(t, "NEW1-CODE", session.VerificationCode) +} + +func TestInitiateOrResumeOAuthDeviceLogin_MintsFreshForDifferentAccessBaseURL(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + + var mintCount atomic.Int32 + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + mintCount.Add(1) + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(DeviceAuthResponse{ //nolint:errcheck + DeviceCode: "new-device-code", + UserCode: "NEW1-CODE", + VerificationURI: "https://access.stripe.com/verify", + ExpiresIn: 300, + Interval: 5, + }) + })) + defer ts.Close() + + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{ + DeviceCode: "other-env-device-code", + AccessBaseURL: "https://a-different-env.example.com", + VerificationURI: "https://access.stripe.com/other", + UserCode: "OTHER-CODE", + ExpiresIn: 300, + IssuedAt: time.Now(), + })) + + session, err := initiateOrResumeOAuthDeviceLogin(context.Background(), ts.URL) + require.NoError(t, err) + assert.Equal(t, int32(1), mintCount.Load()) + assert.Equal(t, "NEW1-CODE", session.VerificationCode) +} + +func TestMintOAuthDeviceLogin_AlwaysMintsFreshEvenWithValidPending(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + + var mintCount atomic.Int32 + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + mintCount.Add(1) + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(DeviceAuthResponse{ //nolint:errcheck + DeviceCode: "new-device-code", + UserCode: "NEW1-CODE", + VerificationURI: "https://access.stripe.com/verify", + ExpiresIn: 300, + Interval: 5, + }) + })) + defer ts.Close() + + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{ + DeviceCode: "still-valid-device-code", + AccessBaseURL: ts.URL, + VerificationURI: "https://access.stripe.com/still-valid", + UserCode: "STILL-VALID", + ExpiresIn: 300, + IssuedAt: time.Now(), + })) + + // `stripe login --non-interactive` (initiateOAuthDeviceLogin) calls mintOAuthDeviceLogin + // directly, not initiateOrResumeOAuthDeviceLogin, so it must mint a new device code even + // though the pending one saved above is still well within its expiry window. + session, err := mintOAuthDeviceLogin(context.Background(), ts.URL) + require.NoError(t, err) + assert.Equal(t, int32(1), mintCount.Load()) + assert.Equal(t, "NEW1-CODE", session.VerificationCode) + + cont, err := loadPendingDeviceAuth() + require.NoError(t, err) + assert.Equal(t, "new-device-code", cont.DeviceCode, "the fresh device code must overwrite the still-valid pending one") +} + +func TestFindPendingOAuthLogin_NoneExists(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + + session, err := FindPendingOAuthLogin("https://example.com") + require.NoError(t, err) + assert.Nil(t, session) +} + +func TestFindPendingOAuthLogin_FindsStillValidPending(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{ + DeviceCode: "device-code", + AccessBaseURL: "https://example.com", + VerificationURI: "https://access.stripe.com/verify", + UserCode: "ABCD-EFGH", + ExpiresIn: 300, + IssuedAt: time.Now(), + })) + + session, err := FindPendingOAuthLogin("https://example.com") + require.NoError(t, err) + require.NotNil(t, session) + assert.Equal(t, "https://access.stripe.com/verify", session.BrowserURL) + assert.Equal(t, "ABCD-EFGH", session.VerificationCode) +} + +func TestFindPendingOAuthLogin_IgnoresExpiredPending(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{ + DeviceCode: "device-code", + AccessBaseURL: "https://example.com", + ExpiresIn: 1, + IssuedAt: time.Now().Add(-11 * time.Minute), // past the 10-minute deadline floor + })) + + session, err := FindPendingOAuthLogin("https://example.com") + require.NoError(t, err) + assert.Nil(t, session) +} + +func TestFindPendingOAuthLogin_IgnoresDifferentAccessBaseURL(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{ + DeviceCode: "device-code", + AccessBaseURL: "https://a-different-env.example.com", + ExpiresIn: 300, + IssuedAt: time.Now(), + })) + + session, err := FindPendingOAuthLogin("https://example.com") + require.NoError(t, err) + assert.Nil(t, session) +} + +func TestPollPendingOAuthLogin_CallerTimeoutLeavesPendingState(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + cfg, cleanup := setupOAuthTestConfig(t) + defer cleanup() + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + json.NewEncoder(w).Encode(tokenErrorResponse{Error: "authorization_pending"}) //nolint:errcheck + })) + defer ts.Close() + stubAccessSrv(t, ts) + + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{ + DeviceCode: "device-code", + Interval: 5, + ExpiresIn: 300, + AccessBaseURL: QAAccessBaseURL, + IssuedAt: time.Now(), + })) + + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + + result, err := PollPendingOAuthLogin(ctx, cfg) + require.NoError(t, err) + assert.Nil(t, result) + + _, loadErr := loadPendingDeviceAuth() + assert.NoError(t, loadErr, "pending state must survive a caller-side timeout so a later poll can resume") +} + +func TestPollPendingOAuthLogin_DeviceCodeExpiryClearsPendingState(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + cfg, cleanup := setupOAuthTestConfig(t) + defer cleanup() + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + json.NewEncoder(w).Encode(tokenErrorResponse{Error: "authorization_pending"}) //nolint:errcheck + })) + defer ts.Close() + stubAccessSrv(t, ts) + + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{ + DeviceCode: "device-code", + Interval: 5, + ExpiresIn: 1, + AccessBaseURL: QAAccessBaseURL, + IssuedAt: time.Now().Add(-11 * time.Minute), // past the 10-minute deadline floor + })) + + _, err := PollPendingOAuthLogin(context.Background(), cfg) + require.Error(t, err) + assert.Contains(t, err.Error(), "expired") + + _, loadErr := loadPendingDeviceAuth() + assert.Error(t, loadErr, "pending state must be cleared once the device code has actually expired") +} + +func TestPollPendingOAuthLogin_TerminalOAuthErrorClearsPendingState(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + cfg, cleanup := setupOAuthTestConfig(t) + defer cleanup() + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + json.NewEncoder(w).Encode(tokenErrorResponse{Error: "access_denied"}) //nolint:errcheck + })) + defer ts.Close() + stubAccessSrv(t, ts) + + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{ + DeviceCode: "device-code", + Interval: 5, + ExpiresIn: 300, + AccessBaseURL: QAAccessBaseURL, + IssuedAt: time.Now(), + })) + + _, err := PollPendingOAuthLogin(context.Background(), cfg) + require.Error(t, err) + assert.Contains(t, err.Error(), "access_denied") + + _, loadErr := loadPendingDeviceAuth() + assert.Error(t, loadErr, "pending state must be cleared once the server returns a terminal error") +} + +func TestPollPendingOAuthLogin_SuccessClearsPendingStateAndReturnsResult(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + cfg, cleanup := setupOAuthTestConfig(t) + defer cleanup() + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + switch r.URL.Path { + case "/stripecli/oauth2/token": + json.NewEncoder(w).Encode(OAuthTokenResponse{ //nolint:errcheck + AccessToken: "oaac_test_access", + RefreshToken: "oart_test_refresh", + TokenType: "Bearer", + ExpiresIn: 3600, + }) + case "/stripecli/oauth2/token/accounts": + json.NewEncoder(w).Encode(listAccountsResponse{Accounts: []config.AuthorizedAccount{ //nolint:errcheck + {ID: "acct_123", Name: "Test Account", Modes: []string{"test", "live"}}, + }}) + default: + http.NotFound(w, r) + } + })) + defer ts.Close() + stubAccessSrv(t, ts) + + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{ + DeviceCode: "device-code", + Interval: 1, + ExpiresIn: 300, + AccessBaseURL: QAAccessBaseURL, + IssuedAt: time.Now(), + })) + + result, err := PollPendingOAuthLogin(context.Background(), cfg) + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, "acct_123", result.ActiveAccountID) + + _, loadErr := loadPendingDeviceAuth() + assert.Error(t, loadErr, "pending state must be cleared once login succeeds") +} + +func TestCheckPendingOAuthLogin_NotYetLoggedInLeavesPendingState(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + cfg, cleanup := setupOAuthTestConfig(t) + defer cleanup() + + var callCount atomic.Int32 + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + callCount.Add(1) + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + json.NewEncoder(w).Encode(tokenErrorResponse{Error: "authorization_pending"}) //nolint:errcheck + })) + defer ts.Close() + stubAccessSrv(t, ts) + + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{ + DeviceCode: "device-code", + Interval: 5, + ExpiresIn: 300, + AccessBaseURL: QAAccessBaseURL, + IssuedAt: time.Now(), + })) + + result, err := CheckPendingOAuthLogin(context.Background(), cfg) + require.NoError(t, err) + assert.Nil(t, result) + assert.Equal(t, int32(1), callCount.Load(), "a single check must make exactly one request, never loop or sleep") + + _, loadErr := loadPendingDeviceAuth() + assert.NoError(t, loadErr, "pending state must survive a not-yet-completed check") +} + +func TestCheckPendingOAuthLogin_DeviceCodeExpiryClearsPendingState(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + cfg, cleanup := setupOAuthTestConfig(t) + defer cleanup() + + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{ + DeviceCode: "device-code", + Interval: 5, + ExpiresIn: 1, + AccessBaseURL: QAAccessBaseURL, + IssuedAt: time.Now().Add(-11 * time.Minute), // past the 10-minute deadline floor + })) + + _, err := CheckPendingOAuthLogin(context.Background(), cfg) + require.Error(t, err) + assert.Contains(t, err.Error(), "expired") + + _, loadErr := loadPendingDeviceAuth() + assert.Error(t, loadErr, "pending state must be cleared once the device code has actually expired") +} + +func TestCheckPendingOAuthLogin_TerminalOAuthErrorClearsPendingState(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + cfg, cleanup := setupOAuthTestConfig(t) + defer cleanup() + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + json.NewEncoder(w).Encode(tokenErrorResponse{Error: "access_denied"}) //nolint:errcheck + })) + defer ts.Close() + stubAccessSrv(t, ts) + + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{ + DeviceCode: "device-code", + Interval: 5, + ExpiresIn: 300, + AccessBaseURL: QAAccessBaseURL, + IssuedAt: time.Now(), + })) + + _, err := CheckPendingOAuthLogin(context.Background(), cfg) + require.Error(t, err) + assert.Contains(t, err.Error(), "access_denied") + + _, loadErr := loadPendingDeviceAuth() + assert.Error(t, loadErr, "pending state must be cleared once the server returns a terminal error") +} + +func TestCheckPendingOAuthLogin_SuccessClearsPendingStateAndReturnsResult(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + cfg, cleanup := setupOAuthTestConfig(t) + defer cleanup() + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + switch r.URL.Path { + case "/stripecli/oauth2/token": + json.NewEncoder(w).Encode(OAuthTokenResponse{ //nolint:errcheck + AccessToken: "oaac_test_access", + RefreshToken: "oart_test_refresh", + TokenType: "Bearer", + ExpiresIn: 3600, + }) + case "/stripecli/oauth2/token/accounts": + json.NewEncoder(w).Encode(listAccountsResponse{Accounts: []config.AuthorizedAccount{ //nolint:errcheck + {ID: "acct_123", Name: "Test Account", Modes: []string{"test", "live"}}, + }}) + default: + http.NotFound(w, r) + } + })) + defer ts.Close() + stubAccessSrv(t, ts) + + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{ + DeviceCode: "device-code", + Interval: 1, + ExpiresIn: 300, + AccessBaseURL: QAAccessBaseURL, + IssuedAt: time.Now(), + })) + + result, err := CheckPendingOAuthLogin(context.Background(), cfg) + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, "acct_123", result.ActiveAccountID) + + _, loadErr := loadPendingDeviceAuth() + assert.Error(t, loadErr, "pending state must be cleared once login succeeds") +} + +func TestCheckPendingOAuthLogin_SaveFailureAfterTokenExchangeClearsPendingState(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + cfg, cleanup := setupOAuthTestConfig(t) + defer cleanup() + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + switch r.URL.Path { + case "/stripecli/oauth2/token": + json.NewEncoder(w).Encode(OAuthTokenResponse{ //nolint:errcheck + AccessToken: "oaac_test_access", + RefreshToken: "oart_test_refresh", + TokenType: "Bearer", + ExpiresIn: 3600, + }) + case "/stripecli/oauth2/token/accounts": + // The token exchange succeeded (the device code is now consumed), but fetching + // account info fails. + w.WriteHeader(http.StatusInternalServerError) + default: + http.NotFound(w, r) + } + })) + defer ts.Close() + stubAccessSrv(t, ts) + + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{ + DeviceCode: "device-code", + Interval: 1, + ExpiresIn: 300, + AccessBaseURL: QAAccessBaseURL, + IssuedAt: time.Now(), + })) + + result, err := CheckPendingOAuthLogin(context.Background(), cfg) + require.Error(t, err) + assert.Nil(t, result) + + _, loadErr := loadPendingDeviceAuth() + assert.Error(t, loadErr, "pending state must be cleared once the (single-use) device code has been consumed, even if saving the resulting credentials failed") +} diff --git a/pkg/login/oauth_accounts.go b/pkg/login/oauth_accounts.go index 52a2a5fc9..7932bfbfc 100644 --- a/pkg/login/oauth_accounts.go +++ b/pkg/login/oauth_accounts.go @@ -7,6 +7,7 @@ import ( "io" "net/http" "os" + "strings" "github.com/stripe/stripe-cli/pkg/ansi" "github.com/stripe/stripe-cli/pkg/config" @@ -52,6 +53,39 @@ func fetchAuthorizedAccounts(ctx context.Context, accessBaseURL, accessToken str return result.Accounts, nil } +// AuthorizedAccountsResult bundles the accounts a user has authorized with which account and +// mode are currently active. +type AuthorizedAccountsResult struct { + Accounts []config.AuthorizedAccount + ActiveAccountID string + ActiveLivemode bool +} + +// ListAuthorizedAccountsForActiveSession fetches the accounts authorized for cfg's stored OAuth +// session, along with which account and mode are currently active. Returns an error if cfg +// isn't logged in via OAuth. +func ListAuthorizedAccountsForActiveSession(ctx context.Context, accessBaseURL string, cfg *config.Config) (*AuthorizedAccountsResult, error) { + uat, err := cfg.Profile.GetUAT() + if err != nil { + return nil, err + } + if !strings.HasPrefix(uat, "oak_") { + return nil, errorcategory.Errorf(errorcategory.Auth, "not logged in; run 'stripe login' first") + } + + accounts, err := ListAuthorizedAccounts(ctx, accessBaseURL, uat) + if err != nil { + return nil, fmt.Errorf("failed to fetch authorized accounts: %w", err) + } + + result := &AuthorizedAccountsResult{Accounts: accounts} + if ac, _ := config.GetActiveContext(); ac != nil { + result.ActiveAccountID = ac.AccountID + result.ActiveLivemode = ac.Livemode + } + return result, nil +} + // PrintAuthorizedContexts fetches the authorized accounts for accessToken and // prints them as a formatted list, marking the active context. func PrintAuthorizedContexts(ctx context.Context, accessBaseURL, accessToken string) error { diff --git a/pkg/login/oauth_accounts_test.go b/pkg/login/oauth_accounts_test.go index 02656073e..57bb7c382 100644 --- a/pkg/login/oauth_accounts_test.go +++ b/pkg/login/oauth_accounts_test.go @@ -56,3 +56,35 @@ func TestListAuthorizedAccounts_returnsRealDataWhenAvailable(t *testing.T) { require.NoError(t, err) assert.Equal(t, want, got) } + +func TestListAuthorizedAccountsForActiveSession_NotLoggedIn(t *testing.T) { + cfg, cleanup := setupOAuthTestConfig(t) + defer cleanup() + + _, err := ListAuthorizedAccountsForActiveSession(context.Background(), "https://access.example", cfg) + assert.ErrorContains(t, err, "not logged in") +} + +func TestListAuthorizedAccountsForActiveSession_ReturnsAccountsAndActiveContext(t *testing.T) { + cfg, cleanup := setupOAuthTestConfig(t) + defer cleanup() + + want := []config.AuthorizedAccount{ + {ID: "acct_123", Name: "Test Co", Modes: []string{"test"}}, + {ID: "acct_456", Name: "Live Co", Modes: []string{"live", "test"}}, + } + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(listAccountsResponse{Accounts: want}) //nolint:errcheck + })) + defer srv.Close() + + require.NoError(t, config.KeyRing.Set(config.UATKeychainItemKey, []byte("oak_test"), "")) + require.NoError(t, config.SaveActiveContext("acct_456", true)) + + result, err := ListAuthorizedAccountsForActiveSession(context.Background(), srv.URL, cfg) + require.NoError(t, err) + assert.Equal(t, want, result.Accounts) + assert.Equal(t, "acct_456", result.ActiveAccountID) + assert.True(t, result.ActiveLivemode) +} diff --git a/pkg/login/oauth_device.go b/pkg/login/oauth_device.go index 553a562bd..d46c9ed9e 100644 --- a/pkg/login/oauth_device.go +++ b/pkg/login/oauth_device.go @@ -97,18 +97,50 @@ func RequestDeviceCode(ctx context.Context, accessBaseURL, clientID string) (*De return &authResp, nil } -// PollDeviceToken polls the token endpoint until the user approves, ctx is -// canceled or times out, or a terminal error is returned. -// -// Callers should create ctx with a deadline matching DeviceAuthResponse.ExpiresIn -// to automatically stop polling when the device code expires. -func PollDeviceToken(ctx context.Context, accessBaseURL, clientID, deviceCode string, interval time.Duration) (*OAuthTokenResponse, error) { +// CheckDeviceToken makes a single, non-blocking request to the token endpoint and returns +// immediately with whatever the server reports right now - unlike PollDeviceToken, it never +// loops or sleeps waiting for the user to complete authentication. A non-nil *OAuthError with +// Code "authorization_pending" or "slow_down" means the user hasn't completed authentication +// yet, not that the device code is dead. +func CheckDeviceToken(ctx context.Context, accessBaseURL, clientID, deviceCode string) (*OAuthTokenResponse, error) { endpoint := accessBaseURL + accessAPNPath + "/token" data := url.Values{} data.Set("grant_type", "urn:ietf:params:oauth:grant-type:device_code") data.Set("client_id", clientID) data.Set("device_code", deviceCode) + resp, err := doPostForm(ctx, endpoint, data) + if err != nil { + return nil, err + } + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + if err != nil { + return nil, err + } + + if resp.StatusCode == http.StatusOK { + var tokenResp OAuthTokenResponse + if err := json.Unmarshal(body, &tokenResp); err != nil { + return nil, fmt.Errorf("failed to parse token response: %w", err) + } + return &tokenResp, nil + } + + var errResp tokenErrorResponse + if jsonErr := json.Unmarshal(body, &errResp); jsonErr != nil || errResp.Error == "" { + return nil, errorcategory.Errorf(errorcategory.Auth, "token request failed (status %d): %s", resp.StatusCode, string(body)) + } + return nil, &OAuthError{Code: errResp.Error, Description: errResp.ErrorDescription, HTTPStatus: resp.StatusCode} +} + +// PollDeviceToken polls the token endpoint until the user approves, ctx is +// canceled or times out, or a terminal error is returned. +// +// Callers should create ctx with a deadline matching DeviceAuthResponse.ExpiresIn +// to automatically stop polling when the device code expires. +func PollDeviceToken(ctx context.Context, accessBaseURL, clientID, deviceCode string, interval time.Duration) (*OAuthTokenResponse, error) { for { select { case <-ctx.Done(): @@ -116,34 +148,18 @@ func PollDeviceToken(ctx context.Context, accessBaseURL, clientID, deviceCode st default: } - resp, err := doPostForm(ctx, endpoint, data) - if err != nil { - return nil, err + tokenResp, err := CheckDeviceToken(ctx, accessBaseURL, clientID, deviceCode) + if err == nil { + return tokenResp, nil } - body, err := io.ReadAll(resp.Body) - resp.Body.Close() - if err != nil { + var oauthErr *OAuthError + if !errors.As(err, &oauthErr) { return nil, err } - if resp.StatusCode == http.StatusOK { - var tokenResp OAuthTokenResponse - if err := json.Unmarshal(body, &tokenResp); err != nil { - return nil, fmt.Errorf("failed to parse token response: %w", err) - } - return &tokenResp, nil - } - - var errResp tokenErrorResponse - if jsonErr := json.Unmarshal(body, &errResp); jsonErr != nil || errResp.Error == "" { - return nil, errorcategory.Errorf(errorcategory.Auth, "token request failed (status %d): %s", resp.StatusCode, string(body)) - } - - oauthErr := &OAuthError{Code: errResp.Error, Description: errResp.ErrorDescription, HTTPStatus: resp.StatusCode} - var wait time.Duration - switch errResp.Error { + switch oauthErr.Code { case "authorization_pending": wait = interval case "slow_down": @@ -203,8 +219,13 @@ func PollAndSaveDeviceCredentials(ctx context.Context, accessBaseURL, clientID, // cancellation/deadline: a caller-side timeout (or the natural device-code expiry) firing at // this exact moment shouldn't leave a valid token saved but the account list and active // context unpopulated. - ctx = context.WithoutCancel(ctx) + return saveDeviceCredentials(context.WithoutCancel(ctx), accessBaseURL, tokenResp, cfg) +} +// saveDeviceCredentials persists a token endpoint response as the active credentials: it clears +// any stale credentials, saves the new OAuth tokens, fetches the authorized accounts, and +// populates cfg's profile with the active account/mode. +func saveDeviceCredentials(ctx context.Context, accessBaseURL string, tokenResp *OAuthTokenResponse, cfg *config.Config) (*DeviceCodeLoginResult, error) { // Clear all stale credentials before saving new ones, so this succeeds even if a // previously stored credential is expired or revoked. _ = cfg.RemoveAuthFields(cfg.Profile.ProfileName) diff --git a/pkg/login/oauth_pending.go b/pkg/login/oauth_pending.go index 85fbca6f5..ca0cf57a9 100644 --- a/pkg/login/oauth_pending.go +++ b/pkg/login/oauth_pending.go @@ -6,34 +6,69 @@ import ( "fmt" "os" "path/filepath" + "time" "github.com/stripe/stripe-cli/pkg/config" "github.com/stripe/stripe-cli/pkg/errorcategory" ) -// oauthContinuation holds the data needed to poll for an OAuth device token. -// It is written to disk by InitiateLogin and read by PollPendingDeviceAuth. +// oauthContinuation holds the data needed to poll for an OAuth device token, plus enough of +// the original device-authorization response to resume (rather than restart) a login attempt. +// It is written to disk by InitiateLogin/InitiateOAuthLogin and read by PollPendingDeviceAuth/ +// PollPendingOAuthLogin. type oauthContinuation struct { - DeviceCode string `json:"device_code"` - Interval int `json:"interval"` - ExpiresIn int `json:"expires_in"` - AccessBaseURL string `json:"access_base"` + DeviceCode string `json:"device_code"` + Interval int `json:"interval"` + ExpiresIn int `json:"expires_in"` + AccessBaseURL string `json:"access_base"` + VerificationURI string `json:"verification_uri"` + UserCode string `json:"user_code"` + IssuedAt time.Time `json:"issued_at"` +} + +// deadline returns the absolute time at which this device code expires. It's computed from +// the persisted IssuedAt rather than "now", so it stays correct across a login attempt that's +// resumed (via InitiateOAuthLogin) or polled (via PollPendingOAuthLogin) well after it started. +func (c *oauthContinuation) deadline() time.Time { + return c.IssuedAt.Add(max(time.Duration(c.ExpiresIn)*time.Second, 10*time.Minute)) } func pendingDeviceAuthPath() string { return filepath.Join(filepath.Dir(config.CredentialsFilePath()), "oauth_pending.json") } +// savePendingDeviceAuth writes cont via a temp file + rename, so a concurrent read (or a +// concurrent write from another process racing to mint its own device code) never observes a +// partially-written file. It does not otherwise coordinate between concurrent writers: if two +// processes mint at once, the second write wins and the first process's caller ends up holding +// a browser_url/verification_code for a device code no longer on disk. That's a narrow, +// single-user-CLI race that isn't worth solving with cross-process locking. func savePendingDeviceAuth(cont *oauthContinuation) error { path := pendingDeviceAuthPath() - if err := os.MkdirAll(filepath.Dir(path), 0700); err != nil { + dir := filepath.Dir(path) + if err := os.MkdirAll(dir, 0700); err != nil { return err } data, err := json.Marshal(cont) if err != nil { return err } - return os.WriteFile(path, data, 0600) + tmp, err := os.CreateTemp(dir, ".oauth_pending-*.tmp") + if err != nil { + return err + } + defer os.Remove(tmp.Name()) + if _, err := tmp.Write(data); err != nil { + tmp.Close() + return err + } + if err := tmp.Close(); err != nil { + return err + } + if err := os.Chmod(tmp.Name(), 0600); err != nil { + return err + } + return os.Rename(tmp.Name(), path) } func loadPendingDeviceAuth() (*oauthContinuation, error) { diff --git a/pkg/login/oauth_pending_test.go b/pkg/login/oauth_pending_test.go new file mode 100644 index 000000000..5756d4d02 --- /dev/null +++ b/pkg/login/oauth_pending_test.go @@ -0,0 +1,59 @@ +package login + +import ( + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestOauthContinuationDeadline(t *testing.T) { + issuedAt := time.Now().Add(-1 * time.Minute) + cont := &oauthContinuation{IssuedAt: issuedAt, ExpiresIn: 1800} + assert.WithinDuration(t, issuedAt.Add(1800*time.Second), cont.deadline(), time.Second) +} + +func TestOauthContinuationDeadline_MinimumTenMinutes(t *testing.T) { + issuedAt := time.Now() + cont := &oauthContinuation{IssuedAt: issuedAt, ExpiresIn: 30} + assert.WithinDuration(t, issuedAt.Add(10*time.Minute), cont.deadline(), time.Second) +} + +func TestSaveAndLoadPendingDeviceAuth(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + + issuedAt := time.Now().Truncate(time.Second) + cont := &oauthContinuation{ + DeviceCode: "device-code", + Interval: 5, + ExpiresIn: 300, + AccessBaseURL: QAAccessBaseURL, + VerificationURI: "https://qa-access.stripe.com/verify", + UserCode: "ABCD-EFGH", + IssuedAt: issuedAt, + } + require.NoError(t, savePendingDeviceAuth(cont)) + + loaded, err := loadPendingDeviceAuth() + require.NoError(t, err) + assert.Equal(t, cont.DeviceCode, loaded.DeviceCode) + assert.Equal(t, cont.VerificationURI, loaded.VerificationURI) + assert.Equal(t, cont.UserCode, loaded.UserCode) + assert.True(t, cont.IssuedAt.Equal(loaded.IssuedAt)) + + clearPendingDeviceAuth() + _, err = loadPendingDeviceAuth() + require.Error(t, err) +} + +func TestSavePendingDeviceAuth_OverwritesPrevious(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{DeviceCode: "first", AccessBaseURL: QAAccessBaseURL, IssuedAt: time.Now(), ExpiresIn: 300})) + require.NoError(t, savePendingDeviceAuth(&oauthContinuation{DeviceCode: "second", AccessBaseURL: QAAccessBaseURL, IssuedAt: time.Now(), ExpiresIn: 300})) + + loaded, err := loadPendingDeviceAuth() + require.NoError(t, err) + assert.Equal(t, "second", loaded.DeviceCode) +} diff --git a/pkg/plugins/auto_upgrade.go b/pkg/plugins/auto_upgrade.go index d1a74a5a0..89d6b2749 100644 --- a/pkg/plugins/auto_upgrade.go +++ b/pkg/plugins/auto_upgrade.go @@ -4,6 +4,8 @@ import ( "context" "fmt" "os" + "path/filepath" + "strings" "time" log "github.com/sirupsen/logrus" @@ -27,23 +29,143 @@ import ( // command, but a large binary on a slow connection is not a failure to cut short. var autoUpgradeResolveTimeout = 3 * time.Second +// autoUpgradeCheckInterval is how long one upgrade check stands before another is +// worth making. +// +// Without a floor on how often it runs, an opted-in plugin spends a request on every +// single command -- and, on a machine that cannot reach the endpoint, the whole +// timeout above on every single command. What that buys is a slightly sooner upgrade, +// which is not something the user is waiting for; they are waiting for the command +// they typed. +// +// A few hours rather than a day: this should still land an upgrade the same working +// session it ships, and one check per plugin per morning is already close to free. +// +// Best-effort, not a guarantee. Two CLIs launched close enough together can both find +// the stamp expired before either has written it, and both check. Bounding it properly +// would mean a lock file and a policy for when to steal one from a process that died +// holding it -- machinery whose failure mode is "never checks again", to save at most +// one duplicate request per concurrent invocation. The measure that matters is +// requests per command, and that is already one per interval for anyone not running +// two plugins at the same instant. +// +// What does deserve locking is the install itself, which is a separate and older +// problem: cleanUpPluginPath deletes every version directory but the one it just +// wrote, so any two processes installing at once can pull a binary out from under a +// third, whether or not an auto-upgrade check is what set them off. +var autoUpgradeCheckInterval = 4 * time.Hour + // Swappable for test injection. These are every effect maybeAutoUpgrade has outside // its own package: the setting (global config state), the lookup (network), the -// download (network and disk), and the plugin's own PostInstall hook (a subprocess). +// download (network and disk), the plugin's own PostInstall hook (a subprocess), and +// the clock the check interval is measured against. var ( pluginUpdatesEnabled = config.PluginUpdatesEnabled autoUpgradeResolver = ResolvePluginForUpgrade autoUpgradePostInstall = runPostInstallHook + autoUpgradeNow = time.Now autoUpgradeInstaller = func(ctx context.Context, resolved *ResolvedPluginVersion, cfg config.IConfig, fs afero.Fs, apiBaseURL, dashboardBaseURL string) error { return resolved.Install(ctx, cfg, fs, apiBaseURL, dashboardBaseURL) } ) +// autoUpgradeCheckStampPath returns the file recording when this plugin was last +// checked for an upgrade. +// +// It sits beside the plugin's local metadata, which is already the directory for +// per-plugin state the CLI keeps for itself, and is per-plugin because the setting is +// too -- one plugin's check should not silence another's. Deliberately not a config +// field: this is written on plugin commands the user runs all day, and viper rewrites +// the entire config file per field. +// +// The extension keeps it out of getLocalPluginMetadataNames, which reads that +// directory to enumerate installed plugins and counts only `.toml` entries. +func autoUpgradeCheckStampPath(cfg config.IConfig, pluginName string) (string, error) { + if err := ValidatePluginShortname(pluginName); err != nil { + return "", err + } + + return filepath.Join(getLocalPluginMetadataDir(cfg), pluginName+".last-upgrade-check"), nil +} + +// removeAutoUpgradeCheckStamp deletes a plugin's check stamp, for an uninstall that +// should not leave anything of the plugin behind. +// +// A missing stamp is not an error: most uninstalls are of plugins that never had one, +// because the setting is off by default. +func removeAutoUpgradeCheckStamp(cfg config.IConfig, fs afero.Fs, pluginName string) error { + path, err := autoUpgradeCheckStampPath(cfg, pluginName) + if err != nil { + return err + } + + if err := fs.Remove(path); err != nil && !os.IsNotExist(err) { + return err + } + + return nil +} + +// autoUpgradeCheckDue reports whether enough time has passed since the last upgrade +// check to be worth spending another request on. See autoUpgradeCheckInterval. +// +// Anything unreadable counts as due. A missing stamp is a first run, and a corrupt one +// is not worth refusing to upgrade over: being wrong in this direction costs the one +// request the stamp exists to save, while being wrong in the other direction means +// never upgrading again. +func autoUpgradeCheckDue(cfg config.IConfig, fs afero.Fs, pluginName string) bool { + path, err := autoUpgradeCheckStampPath(cfg, pluginName) + if err != nil { + return true + } + + body, err := afero.ReadFile(fs, path) + if err != nil { + return true + } + + lastCheck, err := time.Parse(time.RFC3339, strings.TrimSpace(string(body))) + if err != nil { + return true + } + + // A stamp in the future is a clock that has moved backwards, most often a machine + // correcting its time. Waiting for the future to arrive could park the check for + // years, so treat it as due and let the next check overwrite it. + if now := autoUpgradeNow(); lastCheck.After(now) { + return true + } else if now.Sub(lastCheck) < autoUpgradeCheckInterval { + return false + } + + return true +} + +// recordAutoUpgradeCheck claims the current interval for a check about to be made. +// +// Stored as text rather than leaned on the file's mtime, so that `cat`-ing it while +// working out why a plugin did or did not upgrade answers the question. +func recordAutoUpgradeCheck(cfg config.IConfig, fs afero.Fs, pluginName string) error { + path, err := autoUpgradeCheckStampPath(cfg, pluginName) + if err != nil { + return err + } + + if err := fs.MkdirAll(filepath.Dir(path), 0755); err != nil { + return err + } + + return afero.WriteFile(fs, path, []byte(autoUpgradeNow().UTC().Format(time.RFC3339)+"\n"), 0644) +} + // maybeAutoUpgrade upgrades a plugin to the newest release available to this CLI // before it runs, when the user turned `stripe plugin auto-update` on for it. It // returns the plugin and version to run: the newly installed pair when it upgraded, // and the pair it was given every other time. // +// Most calls return without looking anything up. It runs on every invocation of an +// opted-in plugin, but only actually checks once per autoUpgradeCheckInterval. +// // It returns no error, by design. The user asked to run a plugin, not to upgrade // one, so every way this can come up short -- a setting that is off, an endpoint // that will not answer, a download that breaks -- is handled the same way: leave @@ -66,9 +188,21 @@ func maybeAutoUpgrade(ctx context.Context, cfg *config.Config, fs afero.Fs, p *P } switch { - case PluginsPath != "": - // A plugin loaded from a local path is not something the metadata endpoint - // knows about, and overwriting it would throw away what was built there. + case pluginsDirOverride() != "": + // A plugin loaded from a directory the user pointed the CLI at is not something + // the metadata endpoint knows about, and installing over it would throw away + // whatever was built there -- not just overwrite the one version, since the + // install then deletes every other version directory beside it. + // + // Asked of both ways to point the CLI somewhere else, not just the compiled-in + // one. A plugin developer working under STRIPE_PLUGINS_PATH is the likeliest + // person here, and their build is the likeliest thing to lose. + // + // This is broader than it strictly has to be: someone using that variable to + // relocate ordinary installs, rather than to develop a plugin, gives up + // automatic upgrades for them. That is the safe direction -- they still get the + // upgrade hint, and `stripe plugin install` still upgrades on request, whereas + // guessing wrong the other way destroys work with no way to get it back. return p, installedVersion case installedVersion == "": // Nothing to upgrade. Run's auto-install handles a missing binary before @@ -80,12 +214,36 @@ func maybeAutoUpgrade(ctx context.Context, cfg *config.Config, fs afero.Fs, p *P // Read before the lookup below so a user who left this off pays nothing for // the feature, not even one request per plugin command. return p, installedVersion + case !autoUpgradeCheckDue(cfg, fs, p.Shortname): + // Checked recently enough. Ordered after the setting because that read is free + // and this one touches the disk. See autoUpgradeCheckInterval. + logger.Debug("skipping auto-upgrade, checked for one recently") + return p, installedVersion } // Filled in here because the base URLs handed to Run carry only what the user // explicitly passed, and a metadata request has to name a real host. installAPIBaseURL, installDashboardBaseURL := resolveInstallBaseURLs(apiBaseURL, dashboardBaseURL) + // Stamped before the lookup rather than after it, and regardless of how it turns + // out. The request is the cost being rationed, so a lookup that fails or finds + // nothing has to count; claiming the interval up front is what keeps two CLIs + // started at once from both reading an expired stamp and both making the request, + // which stamping afterwards would leave a whole lookup's worth of room for. + // + // It narrows that window rather than closing it -- read and write are still two + // operations. See autoUpgradeCheckInterval for why that is where this stops. + // + // The cost of claiming first is that a process killed between here and the answer + // defers the check by an interval, having learned nothing. For an upgrade the user + // did not ask for, and can wait a few hours for, that is the safe direction to err + // in -- interrupting a plugin download is a supported thing to do. + if stampErr := recordAutoUpgradeCheck(cfg, fs, p.Shortname); stampErr != nil { + // Nothing to do about it beyond checking again next time, which is what the + // feature did before there was a stamp at all. + logger.Debugf("could not record the upgrade check: %s", stampErr) + } + resolveCtx, cancel := context.WithTimeout(ctx, autoUpgradeResolveTimeout) defer cancel() @@ -120,10 +278,17 @@ func maybeAutoUpgrade(ctx context.Context, cfg *config.Config, fs afero.Fs, p *P // first request. // // The cached-metadata fallback is what usually lands here: it can name a version - // but never a binary URL. Skipping it costs nothing, because auto-upgrade runs on - // every invocation and the next one starts a fresh budget. It also means an - // auto-upgrade only ever installs a release a live metadata response just - // offered, which is the guarantee ErrPluginRequiresNewerCLI's doc relies on. + // but never a binary URL. Skipping it defers the upgrade by an + // autoUpgradeCheckInterval, since the check it just declined is the one that got + // stamped -- acceptable for something the user did not ask for, and the price of + // not letting a machine that cannot reach the endpoint retry on every command. It + // also means an auto-upgrade only ever installs a release a live metadata response + // just offered, which is the guarantee ErrPluginRequiresNewerCLI's doc relies on. + // + // Nothing tells the user about the version named here, because the upgrade hint is + // suppressed for a plugin that auto-updates; see CheckLatestPluginVersion. A plugin + // that keeps landing here therefore stays quietly behind, which is worth reporting + // from this side one day rather than by putting the hint's request back. if resolved.BinaryURL == "" { logger.Debugf("skipping auto-upgrade to v%s, the lookup returned no binary URL", resolved.Version) return p, installedVersion diff --git a/pkg/plugins/auto_upgrade_test.go b/pkg/plugins/auto_upgrade_test.go index 8366fabe2..345e6777c 100644 --- a/pkg/plugins/auto_upgrade_test.go +++ b/pkg/plugins/auto_upgrade_test.go @@ -3,7 +3,9 @@ package plugins import ( "context" "errors" + "path/filepath" "runtime" + "strings" "testing" "time" @@ -48,6 +50,9 @@ type autoUpgradeStubs struct { // blockUntilCanceled makes the resolver wait for its context instead of // answering, so a test can prove the lookup is actually bounded. blockUntilCanceled bool + // now is what maybeAutoUpgrade reads the clock as, pinned so a test can place a + // check stamp at an exact age instead of depending on the wall clock. + now time.Time // Recorded calls. settingReads []string @@ -59,26 +64,35 @@ type autoUpgradeStubs struct { func stubAutoUpgrade(t *testing.T) *autoUpgradeStubs { t.Helper() - stubs := &autoUpgradeStubs{updatesEnabled: true} + stubs := &autoUpgradeStubs{ + updatesEnabled: true, + now: time.Date(2026, 4, 1, 12, 0, 0, 0, time.UTC), + } origUpdatesEnabled := pluginUpdatesEnabled origResolver := autoUpgradeResolver origInstaller := autoUpgradeInstaller origPostInstall := autoUpgradePostInstall + origNow := autoUpgradeNow origPluginsPath := PluginsPath // Every test here runs as a normal, non-local-dev install unless it says otherwise. - // Left set by another test in this package, it would skip auto-upgrade outright and - // every assertion below would pass for the wrong reason. + // Left set by another test in this package, either of these would skip auto-upgrade + // outright and every assertion below would pass for the wrong reason. Both spellings, + // because the guard now asks about both. PluginsPath = "" + t.Setenv("STRIPE_PLUGINS_PATH", "") t.Cleanup(func() { pluginUpdatesEnabled = origUpdatesEnabled autoUpgradeResolver = origResolver autoUpgradeInstaller = origInstaller autoUpgradePostInstall = origPostInstall + autoUpgradeNow = origNow PluginsPath = origPluginsPath }) + autoUpgradeNow = func() time.Time { return stubs.now } + pluginUpdatesEnabled = func(pluginName string) bool { stubs.settingReads = append(stubs.settingReads, pluginName) return stubs.updatesEnabled @@ -170,6 +184,45 @@ func autoUpgradeTestConfig() *cfgpkg.Config { return &cfg.Config } +// writeAutoUpgradeCheckStamp records a check as having happened at the given time, in +// the format maybeAutoUpgrade writes rather than by calling the writer, so a test that +// breaks the reader cannot be rescued by a matching break in the writer. +func writeAutoUpgradeCheckStamp(t *testing.T, cfg cfgpkg.IConfig, fs afero.Fs, pluginName string, at time.Time) { + t.Helper() + + path, err := autoUpgradeCheckStampPath(cfg, pluginName) + require.NoError(t, err) + require.NoError(t, fs.MkdirAll(filepath.Dir(path), 0755)) + require.NoError(t, afero.WriteFile(fs, path, []byte(at.UTC().Format(time.RFC3339)+"\n"), 0644)) +} + +// readAutoUpgradeCheckStamp returns the recorded check time, and fails the test if +// there is no stamp to read. +func readAutoUpgradeCheckStamp(t *testing.T, cfg cfgpkg.IConfig, fs afero.Fs, pluginName string) time.Time { + t.Helper() + + path, err := autoUpgradeCheckStampPath(cfg, pluginName) + require.NoError(t, err) + body, err := afero.ReadFile(fs, path) + require.NoError(t, err) + + at, err := time.Parse(time.RFC3339, strings.TrimSpace(string(body))) + require.NoError(t, err) + + return at +} + +func autoUpgradeCheckStampExists(t *testing.T, cfg cfgpkg.IConfig, fs afero.Fs, pluginName string) bool { + t.Helper() + + path, err := autoUpgradeCheckStampPath(cfg, pluginName) + require.NoError(t, err) + exists, err := afero.Exists(fs, path) + require.NoError(t, err) + + return exists +} + func TestMaybeAutoUpgradeInstallsNewerRelease(t *testing.T) { stubs := stubAutoUpgrade(t) stubs.resolved = autoUpgradeResolvedPlugin("1.3.0") @@ -282,10 +335,14 @@ func TestMaybeAutoUpgradeSkips(t *testing.T) { tests := []struct { name string pluginsPath string + pluginsPathEnv string installedVersion string updatesDisabled bool resolved *ResolvedPluginVersion resolveErr error + // lastCheckedAgo places a check stamp that far in the past. Zero leaves the + // plugin unstamped, which is how every case but the throttle one runs. + lastCheckedAgo time.Duration // wantSettingRead is false for the checks that come before it, which is the // point of ordering them that way. wantSettingRead bool @@ -298,6 +355,24 @@ func TestMaybeAutoUpgradeSkips(t *testing.T) { pluginsPath: "/some/local/dev/path", installedVersion: "1.2.0", }, + { + // The same thing said the other way. A plugin developer is far more likely to + // point the CLI at their build with this than to compile a path into it, and + // the check used to miss them entirely -- installing over the directory, and + // deleting every other version in it on the way out. + name: "a plugin loaded from a local path set in the environment", + pluginsPathEnv: "/some/local/dev/path", + installedVersion: "1.2.0", + }, + { + // Whichever way it is set, before the setting is read: it costs nothing, and + // a developer who once turned updates on for a plugin they now have a build of + // should not have that decision reach it. + name: "a local path with updates turned on", + pluginsPathEnv: "/some/local/dev/path", + installedVersion: "1.2.0", + resolved: autoUpgradeResolvedPlugin("1.3.0"), + }, { // Run's auto-install already handles a missing binary, and resolves the // newest release itself while doing so. @@ -316,6 +391,15 @@ func TestMaybeAutoUpgradeSkips(t *testing.T) { updatesDisabled: true, wantSettingRead: true, }, + { + // The other half of what the feature costs someone who left it on: one + // config read and one stat, for all but the first command in a few hours. + name: "checked for an upgrade recently", + installedVersion: "1.2.0", + lastCheckedAgo: autoUpgradeCheckInterval - time.Minute, + resolved: autoUpgradeResolvedPlugin("1.3.0"), + wantSettingRead: true, + }, { name: "the lookup failed", installedVersion: "1.2.0", @@ -392,13 +476,22 @@ func TestMaybeAutoUpgradeSkips(t *testing.T) { stubs.resolved = tt.resolved stubs.resolveErr = tt.resolveErr PluginsPath = tt.pluginsPath + if tt.pluginsPathEnv != "" { + t.Setenv("STRIPE_PLUGINS_PATH", tt.pluginsPathEnv) + } + + cfg := autoUpgradeTestConfig() + fs := afero.NewMemMapFs() + if tt.lastCheckedAgo != 0 { + writeAutoUpgradeCheckStamp(t, cfg, fs, "apps", stubs.now.Add(-tt.lastCheckedAgo)) + } installed := autoUpgradeTestPlugin("1.2.0") var gotPlugin *Plugin var gotVersion string output := captureStderr(t, func() { - gotPlugin, gotVersion = maybeAutoUpgrade(context.Background(), autoUpgradeTestConfig(), afero.NewMemMapFs(), + gotPlugin, gotVersion = maybeAutoUpgrade(context.Background(), cfg, fs, installed, tt.installedVersion, "", "", "") }) @@ -413,6 +506,11 @@ func TestMaybeAutoUpgradeSkips(t *testing.T) { require.Equal(t, tt.wantSettingRead, len(stubs.settingReads) > 0) require.Equal(t, tt.wantLookup, len(stubs.resolveCalls) > 0) + + // A check that spent a request is stamped however it turned out; one that + // bailed before the lookup leaves the next command free to make it. + require.Equal(t, tt.wantLookup || tt.lastCheckedAgo != 0, + autoUpgradeCheckStampExists(t, cfg, fs, "apps")) }) } } @@ -482,3 +580,177 @@ func TestMaybeAutoUpgradeWithNilContext(t *testing.T) { require.Len(t, stubs.resolveCalls, 1) require.True(t, stubs.resolveCalls[0].hasDeadline) } + +// The throttle is meant to be forgotten about: once its interval is up, the check +// happens exactly as it would have without one. +func TestMaybeAutoUpgradeChecksAgainOnceTheIntervalIsUp(t *testing.T) { + stubs := stubAutoUpgrade(t) + stubs.resolved = autoUpgradeResolvedPlugin("1.3.0") + + cfg := autoUpgradeTestConfig() + fs := afero.NewMemMapFs() + lastCheck := stubs.now.Add(-autoUpgradeCheckInterval) + writeAutoUpgradeCheckStamp(t, cfg, fs, "apps", lastCheck) + + var gotVersion string + captureStderr(t, func() { + _, gotVersion = maybeAutoUpgrade(context.Background(), cfg, fs, autoUpgradeTestPlugin("1.2.0"), "1.2.0", "", "", "") + }) + + require.Equal(t, "1.3.0", gotVersion) + require.Len(t, stubs.resolveCalls, 1) + + // Moved forward, not left at the previous check. A stamp that never advanced would + // leave the plugin permanently due and the throttle would do nothing at all. + require.Equal(t, stubs.now, readAutoUpgradeCheckStamp(t, cfg, fs, "apps").UTC()) +} + +// The case the throttle exists for. Without it, a machine that cannot reach the +// metadata endpoint pays the full lookup timeout in front of every plugin command. +func TestMaybeAutoUpgradeThrottlesAFailedCheck(t *testing.T) { + stubs := stubAutoUpgrade(t) + stubs.resolveErr = errors.New("metadata endpoint unreachable") + + cfg := autoUpgradeTestConfig() + fs := afero.NewMemMapFs() + + captureStderr(t, func() { + for range 3 { + maybeAutoUpgrade(context.Background(), cfg, fs, autoUpgradeTestPlugin("1.2.0"), "1.2.0", "", "", "") + } + }) + + require.Len(t, stubs.resolveCalls, 1, "a failed check should be throttled like any other") + require.Equal(t, stubs.now, readAutoUpgradeCheckStamp(t, cfg, fs, "apps").UTC()) +} + +// A successful upgrade is throttled the same way, so that the command right after one +// does not go straight back to the endpoint to be told it is up to date. +func TestMaybeAutoUpgradeThrottlesAfterUpgrading(t *testing.T) { + stubs := stubAutoUpgrade(t) + stubs.resolved = autoUpgradeResolvedPlugin("1.3.0") + + cfg := autoUpgradeTestConfig() + fs := afero.NewMemMapFs() + + var secondVersion string + captureStderr(t, func() { + maybeAutoUpgrade(context.Background(), cfg, fs, autoUpgradeTestPlugin("1.2.0"), "1.2.0", "", "", "") + // As Run would call it next time: the upgraded version is what is installed now. + _, secondVersion = maybeAutoUpgrade(context.Background(), cfg, fs, stubs.resolved.Plugin, "1.3.0", "", "", "") + }) + + require.Len(t, stubs.resolveCalls, 1) + require.Len(t, stubs.installCalls, 1) + require.Equal(t, "1.3.0", secondVersion) +} + +// The stamp is state on disk that anything could have written, so every way of +// reading it wrong has to fall the same way: check, rather than never check again. +func TestMaybeAutoUpgradeChecksWhenTheStampIsUnusable(t *testing.T) { + tests := []struct { + name string + body string + }{ + {name: "empty", body: ""}, + {name: "not a timestamp", body: "yesterday\n"}, + {name: "truncated", body: "2026-04-0"}, + // A clock corrected backwards. Waiting for the recorded time to arrive could + // park the check for as long as the correction was large. + {name: "in the future", body: "2027-04-01T12:00:00Z\n"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + stubs := stubAutoUpgrade(t) + stubs.resolved = autoUpgradeResolvedPlugin("1.3.0") + + cfg := autoUpgradeTestConfig() + fs := afero.NewMemMapFs() + path, err := autoUpgradeCheckStampPath(cfg, "apps") + require.NoError(t, err) + require.NoError(t, fs.MkdirAll(filepath.Dir(path), 0755)) + require.NoError(t, afero.WriteFile(fs, path, []byte(tt.body), 0644)) + + var gotVersion string + captureStderr(t, func() { + _, gotVersion = maybeAutoUpgrade(context.Background(), cfg, fs, autoUpgradeTestPlugin("1.2.0"), "1.2.0", "", "", "") + }) + + require.Equal(t, "1.3.0", gotVersion) + require.Len(t, stubs.resolveCalls, 1) + // Overwritten with something readable, so the throttle works from here on. + require.Equal(t, stubs.now, readAutoUpgradeCheckStamp(t, cfg, fs, "apps").UTC()) + }) + } +} + +// A stamp that cannot be written is the throttle failing, not the upgrade failing. +func TestMaybeAutoUpgradeUpgradesWhenTheStampCannotBeWritten(t *testing.T) { + stubs := stubAutoUpgrade(t) + stubs.resolved = autoUpgradeResolvedPlugin("1.3.0") + + cfg := autoUpgradeTestConfig() + fs := afero.NewReadOnlyFs(afero.NewMemMapFs()) + + var gotVersion string + captureStderr(t, func() { + _, gotVersion = maybeAutoUpgrade(context.Background(), cfg, fs, autoUpgradeTestPlugin("1.2.0"), "1.2.0", "", "", "") + }) + + require.Equal(t, "1.3.0", gotVersion) + require.Len(t, stubs.installCalls, 1) +} + +// The setting is per-plugin, so the throttle has to be too: one plugin's check must +// not stand in for another's. +func TestMaybeAutoUpgradeThrottlesEachPluginSeparately(t *testing.T) { + stubs := stubAutoUpgrade(t) + stubs.resolved = autoUpgradeResolvedPlugin("1.3.0") + + cfg := autoUpgradeTestConfig() + fs := afero.NewMemMapFs() + writeAutoUpgradeCheckStamp(t, cfg, fs, "apps", stubs.now) + + other := autoUpgradeTestPlugin("1.2.0") + other.Shortname = "projects" + + captureStderr(t, func() { + maybeAutoUpgrade(context.Background(), cfg, fs, autoUpgradeTestPlugin("1.2.0"), "1.2.0", "", "", "") + maybeAutoUpgrade(context.Background(), cfg, fs, other, "1.2.0", "", "", "") + }) + + require.Len(t, stubs.resolveCalls, 1) + require.Equal(t, "projects", stubs.resolveCalls[0].pluginName) + require.True(t, autoUpgradeCheckStampExists(t, cfg, fs, "projects")) +} + +// The interval has to be claimed before the request goes out rather than when it comes +// back. The gap between the two is what a concurrently starting CLI slips through, and +// it is as wide as the lookup -- up to the whole resolve timeout. +func TestMaybeAutoUpgradeClaimsTheIntervalBeforeLookingUp(t *testing.T) { + stubs := stubAutoUpgrade(t) + stubs.resolved = autoUpgradeResolvedPlugin("1.3.0") + + cfg := autoUpgradeTestConfig() + fs := afero.NewMemMapFs() + + // Wrapping the stub rather than replacing it keeps its call recording intact. + // stubAutoUpgrade's cleanup restores this along with everything else. + recording := autoUpgradeResolver + var claimedDuringLookup bool + autoUpgradeResolver = func(ctx context.Context, c cfgpkg.IConfig, f afero.Fs, pluginName, apiBaseURL, dashboardBaseURL string) (*ResolvedPluginVersion, error) { + claimedDuringLookup = autoUpgradeCheckStampExists(t, cfg, fs, pluginName) + return recording(ctx, c, f, pluginName, apiBaseURL, dashboardBaseURL) + } + + var gotVersion string + captureStderr(t, func() { + _, gotVersion = maybeAutoUpgrade(context.Background(), cfg, fs, autoUpgradeTestPlugin("1.2.0"), "1.2.0", "", "", "") + }) + + require.Len(t, stubs.resolveCalls, 1) + require.Equal(t, "1.3.0", gotVersion, "claiming the interval must not cost the upgrade") + require.True(t, claimedDuringLookup, + "a second CLI starting while this lookup was in flight would have made the same request") +} diff --git a/pkg/plugins/core_cli_helper.go b/pkg/plugins/core_cli_helper.go index 1a992a787..7e3868181 100644 --- a/pkg/plugins/core_cli_helper.go +++ b/pkg/plugins/core_cli_helper.go @@ -35,6 +35,9 @@ type CoreCLIHelper interface { // `stripe switch context` does. If accountID is empty, shows an interactive picker; // switched is false if the user cancels it, in which case the other return values are empty. SwitchContext(accountID string, livemode bool) (resultAccountID string, accountName string, resultLivemode bool, switched bool, err error) + // ListAuthorizedAccounts returns the accounts the current session is authorized for, along + // with which account and mode are currently active. Returns an error if not logged in. + ListAuthorizedAccounts() (accounts []config.AuthorizedAccount, activeAccountID string, activeLivemode bool, err error) // Login starts a Stripe CLI login, the same way `stripe login --new-session` does when run // interactively: it revokes any existing OAuth session first (so this works even if the // stored credential is expired or revoked), then runs the normal login flow, printing the @@ -45,6 +48,32 @@ type CoreCLIHelper interface { // Calling Login again starts a brand new login attempt (a new device code and browser URL), // not a resumption of this one. Login(timeoutSeconds int32) (accountID string, accountName string, livemode bool, loggedIn bool, err error) + // The OAuth-prefixed methods below (OAuthInitiateLogin, OAuthFindPendingLogin, + // OAuthCheckLoginStatus) are low-level building blocks for a non-interactive, resumable + // login flow. Prefer Login unless you specifically need non-blocking, resumable behavior + // (e.g. driving your own retry loop). + // OAuthInitiateLogin starts a non-interactive OAuth device-code login: it returns + // immediately with a browser URL and verification code for the plugin to present to the + // user (or open itself), instead of blocking until the user completes it like Login does. + // Calling it again before the previous attempt is completed or has expired returns the + // same browser URL and verification code rather than minting a new device code, so a + // plugin (or an agent driving it) can safely retry without orphaning an in-flight login - + // unlike `stripe login --non-interactive`, which always mints a fresh device code and does + // not resume. Call OAuthCheckLoginStatus to check whether the user has completed it. + OAuthInitiateLogin() (browserURL string, verificationCode string, expiresIn int32, err error) + // OAuthFindPendingLogin looks for an OAuth device-code login already in progress - started + // by this call chain, another plugin, or `stripe login --non-interactive` - without + // starting a new one. found is false if there is no pending login attempt or it has + // expired, in which case the other return values are empty. + OAuthFindPendingLogin() (found bool, browserURL string, verificationCode string, expiresIn int32, err error) + // OAuthCheckLoginStatus makes a single, non-blocking check on whether the login started by + // OAuthInitiateLogin has completed - it does not wait for the user, so callers that want to + // wait should call this repeatedly on their own schedule. loggedIn is false without error + // if the user hasn't completed authentication yet, in which case the other return values + // are empty and callers should call this again later to keep checking. Returns an error if + // the device code expired or the user denied authorization, in which case a new + // OAuthInitiateLogin call is required to try again. + OAuthCheckLoginStatus() (accountID string, accountName string, livemode bool, loggedIn bool, err error) } type CoreCLIHelperClient struct { @@ -132,6 +161,18 @@ func (c *CoreCLIHelperClient) SwitchContext(accountID string, livemode bool) (st return resp.AccountId, resp.AccountName, resp.Livemode, resp.Switched, nil } +func (c *CoreCLIHelperClient) ListAuthorizedAccounts() ([]config.AuthorizedAccount, string, bool, error) { + resp, err := c.client.ListAuthorizedAccounts(context.Background(), &proto.ListAuthorizedAccountsRequest{}) + if err != nil { + return nil, "", false, err + } + accounts := make([]config.AuthorizedAccount, len(resp.Accounts)) + for i, a := range resp.Accounts { + accounts[i] = config.AuthorizedAccount{ID: a.Id, Name: a.Name, Modes: a.Modes} + } + return accounts, resp.ActiveAccountId, resp.ActiveLivemode, nil +} + func (c *CoreCLIHelperClient) Login(timeoutSeconds int32) (string, string, bool, bool, error) { resp, err := c.client.Login(context.Background(), &proto.LoginRequest{TimeoutSeconds: timeoutSeconds}) if err != nil { @@ -140,6 +181,30 @@ func (c *CoreCLIHelperClient) Login(timeoutSeconds int32) (string, string, bool, return resp.AccountId, resp.AccountName, resp.Livemode, resp.LoggedIn, nil } +func (c *CoreCLIHelperClient) OAuthInitiateLogin() (string, string, int32, error) { + resp, err := c.client.OAuthInitiateLogin(context.Background(), &proto.OAuthInitiateLoginRequest{}) + if err != nil { + return "", "", 0, err + } + return resp.BrowserUrl, resp.VerificationCode, resp.ExpiresIn, nil +} + +func (c *CoreCLIHelperClient) OAuthFindPendingLogin() (bool, string, string, int32, error) { + resp, err := c.client.OAuthFindPendingLogin(context.Background(), &proto.OAuthFindPendingLoginRequest{}) + if err != nil { + return false, "", "", 0, err + } + return resp.Found, resp.BrowserUrl, resp.VerificationCode, resp.ExpiresIn, nil +} + +func (c *CoreCLIHelperClient) OAuthCheckLoginStatus() (string, string, bool, bool, error) { + resp, err := c.client.OAuthCheckLoginStatus(context.Background(), &proto.OAuthCheckLoginStatusRequest{}) + if err != nil { + return "", "", false, false, err + } + return resp.AccountId, resp.AccountName, resp.Livemode, resp.LoggedIn, nil +} + type CoreCLIHelperServer struct { proto.CoreCLIHelperServer Impl CoreCLIHelper @@ -225,6 +290,18 @@ func (s *CoreCLIHelperServer) SwitchContext(ctx context.Context, req *proto.Swit return &proto.SwitchContextResponse{AccountId: accountID, AccountName: accountName, Livemode: livemode, Switched: switched}, nil } +func (s *CoreCLIHelperServer) ListAuthorizedAccounts(ctx context.Context, req *proto.ListAuthorizedAccountsRequest) (*proto.ListAuthorizedAccountsResponse, error) { + accounts, activeAccountID, activeLivemode, err := s.Impl.ListAuthorizedAccounts() + if err != nil { + return nil, err + } + protoAccounts := make([]*proto.AuthorizedAccount, len(accounts)) + for i, a := range accounts { + protoAccounts[i] = &proto.AuthorizedAccount{Id: a.ID, Name: a.Name, Modes: a.Modes} + } + return &proto.ListAuthorizedAccountsResponse{Accounts: protoAccounts, ActiveAccountId: activeAccountID, ActiveLivemode: activeLivemode}, nil +} + func (s *CoreCLIHelperServer) Login(ctx context.Context, req *proto.LoginRequest) (*proto.LoginResponse, error) { accountID, accountName, livemode, loggedIn, err := s.Impl.Login(req.TimeoutSeconds) if err != nil { @@ -233,6 +310,30 @@ func (s *CoreCLIHelperServer) Login(ctx context.Context, req *proto.LoginRequest return &proto.LoginResponse{AccountId: accountID, AccountName: accountName, Livemode: livemode, LoggedIn: loggedIn}, nil } +func (s *CoreCLIHelperServer) OAuthInitiateLogin(ctx context.Context, req *proto.OAuthInitiateLoginRequest) (*proto.OAuthInitiateLoginResponse, error) { + browserURL, verificationCode, expiresIn, err := s.Impl.OAuthInitiateLogin() + if err != nil { + return nil, err + } + return &proto.OAuthInitiateLoginResponse{BrowserUrl: browserURL, VerificationCode: verificationCode, ExpiresIn: expiresIn}, nil +} + +func (s *CoreCLIHelperServer) OAuthFindPendingLogin(ctx context.Context, req *proto.OAuthFindPendingLoginRequest) (*proto.OAuthFindPendingLoginResponse, error) { + found, browserURL, verificationCode, expiresIn, err := s.Impl.OAuthFindPendingLogin() + if err != nil { + return nil, err + } + return &proto.OAuthFindPendingLoginResponse{Found: found, BrowserUrl: browserURL, VerificationCode: verificationCode, ExpiresIn: expiresIn}, nil +} + +func (s *CoreCLIHelperServer) OAuthCheckLoginStatus(ctx context.Context, req *proto.OAuthCheckLoginStatusRequest) (*proto.OAuthCheckLoginStatusResponse, error) { + accountID, accountName, livemode, loggedIn, err := s.Impl.OAuthCheckLoginStatus() + if err != nil { + return nil, err + } + return &proto.OAuthCheckLoginStatusResponse{AccountId: accountID, AccountName: accountName, Livemode: livemode, LoggedIn: loggedIn}, nil +} + // coreCLIHelper is the real implementation of the CoreCLIHelper interface. type coreCLIHelper struct { ctx context.Context @@ -318,6 +419,10 @@ func clearPendingKeychainValue(key string) { // made by coreCLIHelper.SwitchContext. var loginSwitchContext = login.SwitchContext +// loginListAuthorizedAccountsForActiveSession is a package variable so tests can stub out the +// network call made by coreCLIHelper.ListAuthorizedAccounts. +var loginListAuthorizedAccountsForActiveSession = login.ListAuthorizedAccountsForActiveSession + // loginRevokeToken and loginLogin are package variables so tests can stub out the network/ // keychain calls made by coreCLIHelper.Login. var ( @@ -325,6 +430,16 @@ var ( loginLogin = login.Login ) +// loginInitiateOAuthLogin, loginFindPendingOAuthLogin, and loginCheckPendingOAuthLogin are +// package variables so tests can stub out the network/keychain calls made by +// coreCLIHelper.OAuthInitiateLogin, coreCLIHelper.OAuthFindPendingLogin, and +// coreCLIHelper.OAuthCheckLoginStatus. +var ( + loginInitiateOAuthLogin = login.InitiateOAuthLogin + loginFindPendingOAuthLogin = login.FindPendingOAuthLogin + loginCheckPendingOAuthLogin = login.CheckPendingOAuthLogin +) + // NewCoreCLIHelper creates a new CoreCLIHelper with the given context, config, and filesystem. // apiBaseURL, dashboardBaseURL, and accessBaseURL should be empty unless the user explicitly // passed --api-base/--dashboard-base/--access-base to the CLI. @@ -485,6 +600,25 @@ func (h *coreCLIHelper) SwitchContext(accountID string, livemode bool) (string, return result.Account.ID, result.Account.Name, result.Mode == "live", true, nil } +// ListAuthorizedAccounts returns the accounts the current session is authorized for, along +// with which account and mode are currently active. +func (h *coreCLIHelper) ListAuthorizedAccounts() ([]config.AuthorizedAccount, string, bool, error) { + cfg, ok := h.config.(*config.Config) + if !ok { + return nil, "", false, errorcategory.Errorf(errorcategory.Internal, "could not list authorized accounts: config type mismatch") + } + accessBaseURL := h.accessBaseURL + if accessBaseURL == "" { + accessBaseURL = login.DefaultAccessBaseURL + } + + result, err := loginListAuthorizedAccountsForActiveSession(h.ctx, accessBaseURL, cfg) + if err != nil { + return nil, "", false, err + } + return result.Accounts, result.ActiveAccountID, result.ActiveLivemode, nil +} + // Login starts a Stripe CLI login, the same way `stripe login --new-session` does when run // interactively. func (h *coreCLIHelper) Login(timeoutSeconds int32) (string, string, bool, bool, error) { @@ -529,3 +663,55 @@ func (h *coreCLIHelper) Login(timeoutSeconds int32) (string, string, bool, bool, } return cfg.Profile.AccountID, cfg.Profile.DisplayName, livemode, true, nil } + +// OAuthInitiateLogin starts (or resumes) a non-interactive OAuth device-code login, the same +// way `stripe login --non-interactive` does, without printing anything or blocking on +// completion. +func (h *coreCLIHelper) OAuthInitiateLogin() (string, string, int32, error) { + accessBaseURL := h.accessBaseURL + if accessBaseURL == "" { + accessBaseURL = login.DefaultAccessBaseURL + } + + session, err := loginInitiateOAuthLogin(h.ctx, accessBaseURL) + if err != nil { + return "", "", 0, err + } + return session.BrowserURL, session.VerificationCode, int32(session.ExpiresIn), nil +} + +// OAuthFindPendingLogin looks for an OAuth device-code login already in progress, without +// starting a new one. +func (h *coreCLIHelper) OAuthFindPendingLogin() (bool, string, string, int32, error) { + accessBaseURL := h.accessBaseURL + if accessBaseURL == "" { + accessBaseURL = login.DefaultAccessBaseURL + } + + session, err := loginFindPendingOAuthLogin(accessBaseURL) + if err != nil { + return false, "", "", 0, err + } + if session == nil { + return false, "", "", 0, nil + } + return true, session.BrowserURL, session.VerificationCode, int32(session.ExpiresIn), nil +} + +// OAuthCheckLoginStatus makes a single, non-blocking check on whether the login started by +// OAuthInitiateLogin has completed. +func (h *coreCLIHelper) OAuthCheckLoginStatus() (string, string, bool, bool, error) { + cfg, ok := h.config.(*config.Config) + if !ok { + return "", "", false, false, errorcategory.Errorf(errorcategory.Internal, "could not log in: config type mismatch") + } + + result, err := loginCheckPendingOAuthLogin(h.ctx, cfg) + if err != nil { + return "", "", false, false, err + } + if result == nil { + return "", "", false, false, nil + } + return result.ActiveAccountID, result.ActiveDisplayName, result.ActiveLivemode, true, nil +} diff --git a/pkg/plugins/core_cli_helper_test.go b/pkg/plugins/core_cli_helper_test.go index 97b47e71b..ea76106d0 100644 --- a/pkg/plugins/core_cli_helper_test.go +++ b/pkg/plugins/core_cli_helper_test.go @@ -526,6 +526,54 @@ func TestSwitchContextPropagatesError(t *testing.T) { require.Empty(t, accountName) } +func TestListAuthorizedAccountsReturnsConfigTypeMismatchError(t *testing.T) { + coreCLIHelper := NewCoreCLIHelper(context.Background(), nil, afero.NewMemMapFs(), "", "", "") + accounts, activeAccountID, _, err := coreCLIHelper.ListAuthorizedAccounts() + require.Error(t, err) + require.Empty(t, accounts) + require.Empty(t, activeAccountID) +} + +func TestListAuthorizedAccountsSuccess(t *testing.T) { + original := loginListAuthorizedAccountsForActiveSession + t.Cleanup(func() { loginListAuthorizedAccountsForActiveSession = original }) + + want := []config.AuthorizedAccount{ + {ID: "acct_123", Name: "Test Co", Modes: []string{"test"}}, + {ID: "acct_456", Name: "Live Co", Modes: []string{"live", "test"}}, + } + loginListAuthorizedAccountsForActiveSession = func(ctx context.Context, accessBaseURL string, cfg *config.Config) (*login.AuthorizedAccountsResult, error) { + return &login.AuthorizedAccountsResult{ + Accounts: want, + ActiveAccountID: "acct_456", + ActiveLivemode: true, + }, nil + } + + coreCLIHelper := NewCoreCLIHelper(context.Background(), &config.Config{}, afero.NewMemMapFs(), "", "", "") + accounts, activeAccountID, activeLivemode, err := coreCLIHelper.ListAuthorizedAccounts() + require.NoError(t, err) + require.Equal(t, want, accounts) + require.Equal(t, "acct_456", activeAccountID) + require.True(t, activeLivemode) +} + +func TestListAuthorizedAccountsPropagatesError(t *testing.T) { + original := loginListAuthorizedAccountsForActiveSession + t.Cleanup(func() { loginListAuthorizedAccountsForActiveSession = original }) + + expectedErr := errors.New("boom") + loginListAuthorizedAccountsForActiveSession = func(ctx context.Context, accessBaseURL string, cfg *config.Config) (*login.AuthorizedAccountsResult, error) { + return nil, expectedErr + } + + coreCLIHelper := NewCoreCLIHelper(context.Background(), &config.Config{}, afero.NewMemMapFs(), "", "", "") + accounts, activeAccountID, _, err := coreCLIHelper.ListAuthorizedAccounts() + require.ErrorIs(t, err, expectedErr) + require.Empty(t, accounts) + require.Empty(t, activeAccountID) +} + func TestLoginReturnsConfigTypeMismatchError(t *testing.T) { coreCLIHelper := NewCoreCLIHelper(context.Background(), nil, afero.NewMemMapFs(), "", "", "") accountID, accountName, _, loggedIn, err := coreCLIHelper.Login(0) @@ -676,6 +724,164 @@ func TestLoginPropagatesError(t *testing.T) { require.Empty(t, accountName) } +func TestOAuthInitiateLoginSuccess(t *testing.T) { + originalInitiate := loginInitiateOAuthLogin + t.Cleanup(func() { loginInitiateOAuthLogin = originalInitiate }) + + loginInitiateOAuthLogin = func(ctx context.Context, accessBaseURL string) (*login.OAuthLoginSession, error) { + require.Equal(t, login.DefaultAccessBaseURL, accessBaseURL) + return &login.OAuthLoginSession{ + BrowserURL: "https://access.stripe.com/verify", + VerificationCode: "ABCD-EFGH", + ExpiresIn: 300, + }, nil + } + + coreCLIHelper := NewCoreCLIHelper(context.Background(), &config.Config{}, afero.NewMemMapFs(), "", "", "") + browserURL, verificationCode, expiresIn, err := coreCLIHelper.OAuthInitiateLogin() + require.NoError(t, err) + require.Equal(t, "https://access.stripe.com/verify", browserURL) + require.Equal(t, "ABCD-EFGH", verificationCode) + require.Equal(t, int32(300), expiresIn) +} + +func TestOAuthInitiateLoginPropagatesError(t *testing.T) { + originalInitiate := loginInitiateOAuthLogin + t.Cleanup(func() { loginInitiateOAuthLogin = originalInitiate }) + + expectedErr := errors.New("boom") + loginInitiateOAuthLogin = func(ctx context.Context, accessBaseURL string) (*login.OAuthLoginSession, error) { + return nil, expectedErr + } + + coreCLIHelper := NewCoreCLIHelper(context.Background(), &config.Config{}, afero.NewMemMapFs(), "", "", "") + browserURL, verificationCode, expiresIn, err := coreCLIHelper.OAuthInitiateLogin() + require.ErrorIs(t, err, expectedErr) + require.Empty(t, browserURL) + require.Empty(t, verificationCode) + require.Zero(t, expiresIn) +} + +func TestOAuthFindPendingLoginFound(t *testing.T) { + originalFind := loginFindPendingOAuthLogin + t.Cleanup(func() { loginFindPendingOAuthLogin = originalFind }) + + loginFindPendingOAuthLogin = func(accessBaseURL string) (*login.OAuthLoginSession, error) { + require.Equal(t, login.DefaultAccessBaseURL, accessBaseURL) + return &login.OAuthLoginSession{ + BrowserURL: "https://access.stripe.com/verify", + VerificationCode: "ABCD-EFGH", + ExpiresIn: 300, + }, nil + } + + coreCLIHelper := NewCoreCLIHelper(context.Background(), &config.Config{}, afero.NewMemMapFs(), "", "", "") + found, browserURL, verificationCode, expiresIn, err := coreCLIHelper.OAuthFindPendingLogin() + require.NoError(t, err) + require.True(t, found) + require.Equal(t, "https://access.stripe.com/verify", browserURL) + require.Equal(t, "ABCD-EFGH", verificationCode) + require.Equal(t, int32(300), expiresIn) +} + +func TestOAuthFindPendingLoginNotFound(t *testing.T) { + originalFind := loginFindPendingOAuthLogin + t.Cleanup(func() { loginFindPendingOAuthLogin = originalFind }) + + loginFindPendingOAuthLogin = func(accessBaseURL string) (*login.OAuthLoginSession, error) { + return nil, nil + } + + coreCLIHelper := NewCoreCLIHelper(context.Background(), &config.Config{}, afero.NewMemMapFs(), "", "", "") + found, browserURL, verificationCode, expiresIn, err := coreCLIHelper.OAuthFindPendingLogin() + require.NoError(t, err) + require.False(t, found) + require.Empty(t, browserURL) + require.Empty(t, verificationCode) + require.Zero(t, expiresIn) +} + +func TestOAuthFindPendingLoginPropagatesError(t *testing.T) { + originalFind := loginFindPendingOAuthLogin + t.Cleanup(func() { loginFindPendingOAuthLogin = originalFind }) + + expectedErr := errors.New("boom") + loginFindPendingOAuthLogin = func(accessBaseURL string) (*login.OAuthLoginSession, error) { + return nil, expectedErr + } + + coreCLIHelper := NewCoreCLIHelper(context.Background(), &config.Config{}, afero.NewMemMapFs(), "", "", "") + found, browserURL, verificationCode, expiresIn, err := coreCLIHelper.OAuthFindPendingLogin() + require.ErrorIs(t, err, expectedErr) + require.False(t, found) + require.Empty(t, browserURL) + require.Empty(t, verificationCode) + require.Zero(t, expiresIn) +} + +func TestOAuthCheckLoginStatusReturnsConfigTypeMismatchError(t *testing.T) { + coreCLIHelper := NewCoreCLIHelper(context.Background(), nil, afero.NewMemMapFs(), "", "", "") + accountID, accountName, _, loggedIn, err := coreCLIHelper.OAuthCheckLoginStatus() + require.Error(t, err) + require.False(t, loggedIn) + require.Empty(t, accountID) + require.Empty(t, accountName) +} + +func TestOAuthCheckLoginStatusSuccess(t *testing.T) { + originalCheck := loginCheckPendingOAuthLogin + t.Cleanup(func() { loginCheckPendingOAuthLogin = originalCheck }) + + loginCheckPendingOAuthLogin = func(ctx context.Context, cfg *config.Config) (*login.DeviceCodeLoginResult, error) { + return &login.DeviceCodeLoginResult{ + ActiveAccountID: "acct_123", + ActiveDisplayName: "Acme Inc", + ActiveLivemode: true, + }, nil + } + + coreCLIHelper := NewCoreCLIHelper(context.Background(), &config.Config{}, afero.NewMemMapFs(), "", "", "") + accountID, accountName, livemode, loggedIn, err := coreCLIHelper.OAuthCheckLoginStatus() + require.NoError(t, err) + require.True(t, loggedIn) + require.Equal(t, "acct_123", accountID) + require.Equal(t, "Acme Inc", accountName) + require.True(t, livemode) +} + +func TestOAuthCheckLoginStatusNotYetLoggedIn(t *testing.T) { + originalCheck := loginCheckPendingOAuthLogin + t.Cleanup(func() { loginCheckPendingOAuthLogin = originalCheck }) + + loginCheckPendingOAuthLogin = func(ctx context.Context, cfg *config.Config) (*login.DeviceCodeLoginResult, error) { + return nil, nil + } + + coreCLIHelper := NewCoreCLIHelper(context.Background(), &config.Config{}, afero.NewMemMapFs(), "", "", "") + accountID, accountName, _, loggedIn, err := coreCLIHelper.OAuthCheckLoginStatus() + require.NoError(t, err) + require.False(t, loggedIn) + require.Empty(t, accountID) + require.Empty(t, accountName) +} + +func TestOAuthCheckLoginStatusPropagatesError(t *testing.T) { + originalCheck := loginCheckPendingOAuthLogin + t.Cleanup(func() { loginCheckPendingOAuthLogin = originalCheck }) + + expectedErr := errors.New("boom") + loginCheckPendingOAuthLogin = func(ctx context.Context, cfg *config.Config) (*login.DeviceCodeLoginResult, error) { + return nil, expectedErr + } + + coreCLIHelper := NewCoreCLIHelper(context.Background(), &config.Config{}, afero.NewMemMapFs(), "", "", "") + accountID, accountName, _, loggedIn, err := coreCLIHelper.OAuthCheckLoginStatus() + require.ErrorIs(t, err, expectedErr) + require.False(t, loggedIn) + require.Empty(t, accountID) + require.Empty(t, accountName) +} + func TestSendAnalyticsWithTelemetryClient(t *testing.T) { // Test with a NoOp telemetry client ctx := context.Background() diff --git a/pkg/plugins/plugin.go b/pkg/plugins/plugin.go index 7e76c11a9..f3f6b0c01 100644 --- a/pkg/plugins/plugin.go +++ b/pkg/plugins/plugin.go @@ -443,6 +443,18 @@ func (p *Plugin) Uninstall(ctx context.Context, config config.IConfig, fs afero. return err } + // Last, and deliberately not part of the rollback above. The stamp only rations how + // often the CLI asks about upgrades, so an uninstall that has already removed the + // binary and the metadata has succeeded whether or not this cache goes with it -- + // and putting a whole uninstall back because a timestamp would not delete would be + // far worse than leaving the timestamp. + if err := removeAutoUpgradeCheckStamp(config, fs, p.Shortname); err != nil { + log.WithFields(log.Fields{ + "prefix": "plugins.plugin.Uninstall", + "plugin": p.Shortname, + }).Debugf("could not remove the upgrade check stamp: %s", err) + } + return nil } @@ -650,9 +662,32 @@ func (p *Plugin) Run(ctx context.Context, config *config.Config, fs afero.Fs, ar return p.run(ctx, config, fs, args, cwd, versionOverride, apiBaseURL, dashboardBaseURL, accessBaseURL, true) } -// run is Run with the auto-upgrade check made optional, so that the one caller who -// reaches a plugin without the user having asked for it can leave it out. See -// CoreCLIHelper.RunPeerPlugin. +// RunWithoutAutoUpgrade is Run for a handoff that should not spend a metadata request +// on the auto-upgrade check. Two callers want this, for different reasons: +// +// Printing the plugin's own help, which the CLI hands to the plugin because the plugin +// owns that text rather than the manifest. `--help` is a question about a command, and +// answering it should not download and install software: someone reading help is usually +// deciding whether to run something, or has just been told they got the flags wrong, and +// neither is a moment for an upgrade they did not ask for. +// +// Running a plugin that was installed earlier in this same invocation, where the install +// already resolved the newest release -- so a check here would spend a second request to +// be told what the first one just said. This mirrors what run's own auto-install branch +// does when the install happens inside it. +// +// Either way the upgrade is deferred, not lost: the next command that does real work on +// an already-installed plugin makes the check. +// +// An install still happens here when the binary is missing, since there is nowhere else +// for the plugin or its help text to come from. +func (p *Plugin) RunWithoutAutoUpgrade(ctx context.Context, config *config.Config, fs afero.Fs, args []string, cwd string, versionOverride string, apiBaseURL, dashboardBaseURL, accessBaseURL string) error { + return p.run(ctx, config, fs, args, cwd, versionOverride, apiBaseURL, dashboardBaseURL, accessBaseURL, false) +} + +// run is Run with the auto-upgrade check made optional, so that callers with nothing to +// gain from it can leave it out. See CoreCLIHelper.RunPeerPlugin and +// RunWithoutAutoUpgrade. func (p *Plugin) run(ctx context.Context, config *config.Config, fs afero.Fs, args []string, cwd string, versionOverride string, apiBaseURL, dashboardBaseURL, accessBaseURL string, allowAutoUpgrade bool) error { logger := log.WithFields(log.Fields{ "prefix": "plugins.plugin.Run", @@ -726,17 +761,17 @@ func (p *Plugin) run(ctx context.Context, config *config.Config, fs afero.Fs, ar case Dispatcher: logger.Debug("negotiated net/rpc with plugin process") if _, err = d.RunCommand(args); err != nil { - return err + return pluginReportedError{err} } case DispatcherGRPC: logger.Debug("negotiated gRPC with plugin process") if err = d.RunCommand(buildAdditionalInfo(logger, apiBaseURL, dashboardBaseURL, accessBaseURL), args); err != nil { - return err + return pluginReportedError{err} } case DispatcherV3: logger.Debug("negotiated gRPC with plugin process (v3)") if err = d.RunCommand(buildAdditionalInfo(logger, apiBaseURL, dashboardBaseURL, accessBaseURL), args, NewCoreCLIHelper(ctx, config, fs, apiBaseURL, dashboardBaseURL, accessBaseURL)); err != nil { - return err + return pluginReportedError{err} } default: return errorcategory.New(errorcategory.Internal, "dispensed an unknown plugin interface") diff --git a/pkg/plugins/plugin_test.go b/pkg/plugins/plugin_test.go index eb8066789..a45d7097e 100644 --- a/pkg/plugins/plugin_test.go +++ b/pkg/plugins/plugin_test.go @@ -13,6 +13,7 @@ import ( "runtime" "strings" "testing" + "time" "github.com/BurntSushi/toml" "github.com/spf13/afero" @@ -862,9 +863,15 @@ func setUpRunAutoUpgrade(t *testing.T, installedVersion string) (*autoUpgradeStu cfg := &TestConfig{} cfg.InitConfig() - t.Setenv("STRIPE_PLUGINS_PATH", "/plugins") + // Run takes a *config.Config, so TestConfig's "/" config folder does not apply inside + // it and the plugins directory has to be moved onto the memory filesystem some other + // way. XDG_CONFIG_HOME rather than STRIPE_PLUGINS_PATH: the latter is an overridden + // plugins directory, which maybeAutoUpgrade now refuses to install into, so every + // test here would pass for the wrong reason. This moves the whole config folder, + // which is what a real machine does too. + t.Setenv("XDG_CONFIG_HOME", "/xdg") - installDir := filepath.Join("/plugins/appA", installedVersion) + installDir := filepath.Join(getPluginsDir(&cfg.Config), "appA", installedVersion) require.NoError(t, fs.MkdirAll(installDir, 0755)) require.NoError(t, afero.WriteFile(fs, filepath.Join(installDir, "stripe-cli-app-a"+GetBinaryExtension()), []byte("bin"), 0755)) @@ -936,6 +943,22 @@ func TestRunSkipsAutoUpgradeForLocalDevelopmentBuild(t *testing.T) { require.Empty(t, stubs.installCalls) } +func TestRunWithoutAutoUpgradeSkipsTheCheck(t *testing.T) { + stubs, cfg, fs := setUpRunAutoUpgrade(t, "1.0.1") + + plugin, err := LookUpPlugin(context.Background(), cfg, fs, "appA") + require.NoError(t, err) + + require.Error(t, plugin.RunWithoutAutoUpgrade(context.Background(), &cfg.Config, fs, []string{"--help"}, "", "", "", "", "")) + + // Not even the setting is read. Both callers of this reach a plugin with nothing to + // gain from the check -- printing help, or running something whose install just + // resolved the newest release -- so neither should pay a request for it. + require.Empty(t, stubs.settingReads) + require.Empty(t, stubs.resolveCalls) + require.Empty(t, stubs.installCalls) +} + func TestRunPeerPluginSkipsAutoUpgrade(t *testing.T) { stubs, cfg, fs := setUpRunAutoUpgrade(t, "1.0.1") @@ -1194,6 +1217,53 @@ func TestUninstallSucceedsWithLocalMetadataOnly(t *testing.T) { require.Equal(t, 0, len(config.GetInstalledPlugins())) } +// The check stamp is the one piece of per-plugin state that does not live in the +// metadata file or the plugin directory, so nothing else in Uninstall reaches it. +func TestUninstallRemovesTheAutoUpgradeCheckStamp(t *testing.T) { + fs := afero.NewMemMapFs() + config := &TestConfig{} + config.InitConfig() + plugin := Plugin{ + Shortname: "sample-plugin", + Binary: "stripe-cli-sample-plugin", + MagicCookieValue: "SAMPLE-COOKIE", + Releases: []Release{ + {Arch: runtime.GOARCH, OS: runtime.GOOS, Version: "1.0.0", Sum: "abc123"}, + }, + } + + require.NoError(t, writeLocalPluginMetadata(config, fs, plugin)) + require.NoError(t, fs.MkdirAll("/plugins/sample-plugin/1.0.0", 0755)) + writeAutoUpgradeCheckStamp(t, config, fs, "sample-plugin", time.Now()) + require.True(t, autoUpgradeCheckStampExists(t, config, fs, "sample-plugin")) + + require.NoError(t, plugin.Uninstall(context.Background(), config, fs)) + + require.False(t, autoUpgradeCheckStampExists(t, config, fs, "sample-plugin")) +} + +// The ordinary case, since auto-update is off by default: there is no stamp to remove, +// and an uninstall must not report that as a problem. +func TestUninstallSucceedsWithoutAnAutoUpgradeCheckStamp(t *testing.T) { + fs := afero.NewMemMapFs() + config := &TestConfig{} + config.InitConfig() + plugin := Plugin{ + Shortname: "sample-plugin", + Binary: "stripe-cli-sample-plugin", + MagicCookieValue: "SAMPLE-COOKIE", + Releases: []Release{ + {Arch: runtime.GOARCH, OS: runtime.GOOS, Version: "1.0.0", Sum: "abc123"}, + }, + } + + require.NoError(t, writeLocalPluginMetadata(config, fs, plugin)) + require.NoError(t, fs.MkdirAll("/plugins/sample-plugin/1.0.0", 0755)) + require.False(t, autoUpgradeCheckStampExists(t, config, fs, "sample-plugin")) + + require.NoError(t, plugin.Uninstall(context.Background(), config, fs)) +} + func TestUninstallRejectsInvalidPluginShortnames(t *testing.T) { tests := []string{"../victim", "..\\victim"} diff --git a/pkg/plugins/proto/main.pb.go b/pkg/plugins/proto/main.pb.go index b2af5d019..7c486e58f 100644 --- a/pkg/plugins/proto/main.pb.go +++ b/pkg/plugins/proto/main.pb.go @@ -1367,6 +1367,165 @@ func (x *SwitchContextResponse) GetSwitched() bool { return false } +type ListAuthorizedAccountsRequest struct { + state protoimpl.MessageState `protogen:"open.v1"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *ListAuthorizedAccountsRequest) Reset() { + *x = ListAuthorizedAccountsRequest{} + mi := &file_pkg_plugins_proto_main_proto_msgTypes[27] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *ListAuthorizedAccountsRequest) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*ListAuthorizedAccountsRequest) ProtoMessage() {} + +func (x *ListAuthorizedAccountsRequest) ProtoReflect() protoreflect.Message { + mi := &file_pkg_plugins_proto_main_proto_msgTypes[27] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use ListAuthorizedAccountsRequest.ProtoReflect.Descriptor instead. +func (*ListAuthorizedAccountsRequest) Descriptor() ([]byte, []int) { + return file_pkg_plugins_proto_main_proto_rawDescGZIP(), []int{27} +} + +type AuthorizedAccount struct { + state protoimpl.MessageState `protogen:"open.v1"` + Id string `protobuf:"bytes,1,opt,name=id,proto3" json:"id,omitempty"` + Name string `protobuf:"bytes,2,opt,name=name,proto3" json:"name,omitempty"` + // modes lists the API mode(s) ("test", "live") this account grants access to. + Modes []string `protobuf:"bytes,3,rep,name=modes,proto3" json:"modes,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *AuthorizedAccount) Reset() { + *x = AuthorizedAccount{} + mi := &file_pkg_plugins_proto_main_proto_msgTypes[28] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *AuthorizedAccount) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*AuthorizedAccount) ProtoMessage() {} + +func (x *AuthorizedAccount) ProtoReflect() protoreflect.Message { + mi := &file_pkg_plugins_proto_main_proto_msgTypes[28] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use AuthorizedAccount.ProtoReflect.Descriptor instead. +func (*AuthorizedAccount) Descriptor() ([]byte, []int) { + return file_pkg_plugins_proto_main_proto_rawDescGZIP(), []int{28} +} + +func (x *AuthorizedAccount) GetId() string { + if x != nil { + return x.Id + } + return "" +} + +func (x *AuthorizedAccount) GetName() string { + if x != nil { + return x.Name + } + return "" +} + +func (x *AuthorizedAccount) GetModes() []string { + if x != nil { + return x.Modes + } + return nil +} + +type ListAuthorizedAccountsResponse struct { + state protoimpl.MessageState `protogen:"open.v1"` + Accounts []*AuthorizedAccount `protobuf:"bytes,1,rep,name=accounts,proto3" json:"accounts,omitempty"` + // active_account_id and active_livemode identify which authorized account and mode are + // currently active; empty/false if none is active yet. + ActiveAccountId string `protobuf:"bytes,2,opt,name=active_account_id,json=activeAccountId,proto3" json:"active_account_id,omitempty"` + ActiveLivemode bool `protobuf:"varint,3,opt,name=active_livemode,json=activeLivemode,proto3" json:"active_livemode,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *ListAuthorizedAccountsResponse) Reset() { + *x = ListAuthorizedAccountsResponse{} + mi := &file_pkg_plugins_proto_main_proto_msgTypes[29] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *ListAuthorizedAccountsResponse) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*ListAuthorizedAccountsResponse) ProtoMessage() {} + +func (x *ListAuthorizedAccountsResponse) ProtoReflect() protoreflect.Message { + mi := &file_pkg_plugins_proto_main_proto_msgTypes[29] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use ListAuthorizedAccountsResponse.ProtoReflect.Descriptor instead. +func (*ListAuthorizedAccountsResponse) Descriptor() ([]byte, []int) { + return file_pkg_plugins_proto_main_proto_rawDescGZIP(), []int{29} +} + +func (x *ListAuthorizedAccountsResponse) GetAccounts() []*AuthorizedAccount { + if x != nil { + return x.Accounts + } + return nil +} + +func (x *ListAuthorizedAccountsResponse) GetActiveAccountId() string { + if x != nil { + return x.ActiveAccountId + } + return "" +} + +func (x *ListAuthorizedAccountsResponse) GetActiveLivemode() bool { + if x != nil { + return x.ActiveLivemode + } + return false +} + type LoginRequest struct { state protoimpl.MessageState `protogen:"open.v1"` // timeout_seconds bounds how long Login waits for the user to complete @@ -1379,7 +1538,7 @@ type LoginRequest struct { func (x *LoginRequest) Reset() { *x = LoginRequest{} - mi := &file_pkg_plugins_proto_main_proto_msgTypes[27] + mi := &file_pkg_plugins_proto_main_proto_msgTypes[30] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -1391,7 +1550,7 @@ func (x *LoginRequest) String() string { func (*LoginRequest) ProtoMessage() {} func (x *LoginRequest) ProtoReflect() protoreflect.Message { - mi := &file_pkg_plugins_proto_main_proto_msgTypes[27] + mi := &file_pkg_plugins_proto_main_proto_msgTypes[30] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -1404,7 +1563,7 @@ func (x *LoginRequest) ProtoReflect() protoreflect.Message { // Deprecated: Use LoginRequest.ProtoReflect.Descriptor instead. func (*LoginRequest) Descriptor() ([]byte, []int) { - return file_pkg_plugins_proto_main_proto_rawDescGZIP(), []int{27} + return file_pkg_plugins_proto_main_proto_rawDescGZIP(), []int{30} } func (x *LoginRequest) GetTimeoutSeconds() int32 { @@ -1430,7 +1589,7 @@ type LoginResponse struct { func (x *LoginResponse) Reset() { *x = LoginResponse{} - mi := &file_pkg_plugins_proto_main_proto_msgTypes[28] + mi := &file_pkg_plugins_proto_main_proto_msgTypes[31] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -1442,7 +1601,7 @@ func (x *LoginResponse) String() string { func (*LoginResponse) ProtoMessage() {} func (x *LoginResponse) ProtoReflect() protoreflect.Message { - mi := &file_pkg_plugins_proto_main_proto_msgTypes[28] + mi := &file_pkg_plugins_proto_main_proto_msgTypes[31] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -1455,7 +1614,7 @@ func (x *LoginResponse) ProtoReflect() protoreflect.Message { // Deprecated: Use LoginResponse.ProtoReflect.Descriptor instead. func (*LoginResponse) Descriptor() ([]byte, []int) { - return file_pkg_plugins_proto_main_proto_rawDescGZIP(), []int{28} + return file_pkg_plugins_proto_main_proto_rawDescGZIP(), []int{31} } func (x *LoginResponse) GetAccountId() string { @@ -1486,6 +1645,319 @@ func (x *LoginResponse) GetLoggedIn() bool { return false } +type OAuthInitiateLoginRequest struct { + state protoimpl.MessageState `protogen:"open.v1"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *OAuthInitiateLoginRequest) Reset() { + *x = OAuthInitiateLoginRequest{} + mi := &file_pkg_plugins_proto_main_proto_msgTypes[32] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *OAuthInitiateLoginRequest) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*OAuthInitiateLoginRequest) ProtoMessage() {} + +func (x *OAuthInitiateLoginRequest) ProtoReflect() protoreflect.Message { + mi := &file_pkg_plugins_proto_main_proto_msgTypes[32] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use OAuthInitiateLoginRequest.ProtoReflect.Descriptor instead. +func (*OAuthInitiateLoginRequest) Descriptor() ([]byte, []int) { + return file_pkg_plugins_proto_main_proto_rawDescGZIP(), []int{32} +} + +type OAuthInitiateLoginResponse struct { + state protoimpl.MessageState `protogen:"open.v1"` + BrowserUrl string `protobuf:"bytes,1,opt,name=browser_url,json=browserUrl,proto3" json:"browser_url,omitempty"` + VerificationCode string `protobuf:"bytes,2,opt,name=verification_code,json=verificationCode,proto3" json:"verification_code,omitempty"` + // expires_in is the number of seconds remaining before the browser_url/ + // verification_code expire and a new OAuthInitiateLogin call is required. + ExpiresIn int32 `protobuf:"varint,3,opt,name=expires_in,json=expiresIn,proto3" json:"expires_in,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *OAuthInitiateLoginResponse) Reset() { + *x = OAuthInitiateLoginResponse{} + mi := &file_pkg_plugins_proto_main_proto_msgTypes[33] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *OAuthInitiateLoginResponse) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*OAuthInitiateLoginResponse) ProtoMessage() {} + +func (x *OAuthInitiateLoginResponse) ProtoReflect() protoreflect.Message { + mi := &file_pkg_plugins_proto_main_proto_msgTypes[33] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use OAuthInitiateLoginResponse.ProtoReflect.Descriptor instead. +func (*OAuthInitiateLoginResponse) Descriptor() ([]byte, []int) { + return file_pkg_plugins_proto_main_proto_rawDescGZIP(), []int{33} +} + +func (x *OAuthInitiateLoginResponse) GetBrowserUrl() string { + if x != nil { + return x.BrowserUrl + } + return "" +} + +func (x *OAuthInitiateLoginResponse) GetVerificationCode() string { + if x != nil { + return x.VerificationCode + } + return "" +} + +func (x *OAuthInitiateLoginResponse) GetExpiresIn() int32 { + if x != nil { + return x.ExpiresIn + } + return 0 +} + +type OAuthFindPendingLoginRequest struct { + state protoimpl.MessageState `protogen:"open.v1"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *OAuthFindPendingLoginRequest) Reset() { + *x = OAuthFindPendingLoginRequest{} + mi := &file_pkg_plugins_proto_main_proto_msgTypes[34] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *OAuthFindPendingLoginRequest) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*OAuthFindPendingLoginRequest) ProtoMessage() {} + +func (x *OAuthFindPendingLoginRequest) ProtoReflect() protoreflect.Message { + mi := &file_pkg_plugins_proto_main_proto_msgTypes[34] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use OAuthFindPendingLoginRequest.ProtoReflect.Descriptor instead. +func (*OAuthFindPendingLoginRequest) Descriptor() ([]byte, []int) { + return file_pkg_plugins_proto_main_proto_rawDescGZIP(), []int{34} +} + +type OAuthFindPendingLoginResponse struct { + state protoimpl.MessageState `protogen:"open.v1"` + // found is false if there is no pending login attempt, or it has + // expired; in that case the other fields are empty. + Found bool `protobuf:"varint,1,opt,name=found,proto3" json:"found,omitempty"` + BrowserUrl string `protobuf:"bytes,2,opt,name=browser_url,json=browserUrl,proto3" json:"browser_url,omitempty"` + VerificationCode string `protobuf:"bytes,3,opt,name=verification_code,json=verificationCode,proto3" json:"verification_code,omitempty"` + // expires_in is the number of seconds remaining before the browser_url/ + // verification_code expire and a new OAuthInitiateLogin call is required. + ExpiresIn int32 `protobuf:"varint,4,opt,name=expires_in,json=expiresIn,proto3" json:"expires_in,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *OAuthFindPendingLoginResponse) Reset() { + *x = OAuthFindPendingLoginResponse{} + mi := &file_pkg_plugins_proto_main_proto_msgTypes[35] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *OAuthFindPendingLoginResponse) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*OAuthFindPendingLoginResponse) ProtoMessage() {} + +func (x *OAuthFindPendingLoginResponse) ProtoReflect() protoreflect.Message { + mi := &file_pkg_plugins_proto_main_proto_msgTypes[35] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use OAuthFindPendingLoginResponse.ProtoReflect.Descriptor instead. +func (*OAuthFindPendingLoginResponse) Descriptor() ([]byte, []int) { + return file_pkg_plugins_proto_main_proto_rawDescGZIP(), []int{35} +} + +func (x *OAuthFindPendingLoginResponse) GetFound() bool { + if x != nil { + return x.Found + } + return false +} + +func (x *OAuthFindPendingLoginResponse) GetBrowserUrl() string { + if x != nil { + return x.BrowserUrl + } + return "" +} + +func (x *OAuthFindPendingLoginResponse) GetVerificationCode() string { + if x != nil { + return x.VerificationCode + } + return "" +} + +func (x *OAuthFindPendingLoginResponse) GetExpiresIn() int32 { + if x != nil { + return x.ExpiresIn + } + return 0 +} + +type OAuthCheckLoginStatusRequest struct { + state protoimpl.MessageState `protogen:"open.v1"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *OAuthCheckLoginStatusRequest) Reset() { + *x = OAuthCheckLoginStatusRequest{} + mi := &file_pkg_plugins_proto_main_proto_msgTypes[36] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *OAuthCheckLoginStatusRequest) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*OAuthCheckLoginStatusRequest) ProtoMessage() {} + +func (x *OAuthCheckLoginStatusRequest) ProtoReflect() protoreflect.Message { + mi := &file_pkg_plugins_proto_main_proto_msgTypes[36] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use OAuthCheckLoginStatusRequest.ProtoReflect.Descriptor instead. +func (*OAuthCheckLoginStatusRequest) Descriptor() ([]byte, []int) { + return file_pkg_plugins_proto_main_proto_rawDescGZIP(), []int{36} +} + +type OAuthCheckLoginStatusResponse struct { + state protoimpl.MessageState `protogen:"open.v1"` + AccountId string `protobuf:"bytes,1,opt,name=account_id,json=accountId,proto3" json:"account_id,omitempty"` + AccountName string `protobuf:"bytes,2,opt,name=account_name,json=accountName,proto3" json:"account_name,omitempty"` + Livemode bool `protobuf:"varint,3,opt,name=livemode,proto3" json:"livemode,omitempty"` + // logged_in is false if the user hasn't completed authentication yet; in + // that case the other fields are empty and callers should call + // OAuthCheckLoginStatus again later to keep checking. + LoggedIn bool `protobuf:"varint,4,opt,name=logged_in,json=loggedIn,proto3" json:"logged_in,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *OAuthCheckLoginStatusResponse) Reset() { + *x = OAuthCheckLoginStatusResponse{} + mi := &file_pkg_plugins_proto_main_proto_msgTypes[37] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *OAuthCheckLoginStatusResponse) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*OAuthCheckLoginStatusResponse) ProtoMessage() {} + +func (x *OAuthCheckLoginStatusResponse) ProtoReflect() protoreflect.Message { + mi := &file_pkg_plugins_proto_main_proto_msgTypes[37] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use OAuthCheckLoginStatusResponse.ProtoReflect.Descriptor instead. +func (*OAuthCheckLoginStatusResponse) Descriptor() ([]byte, []int) { + return file_pkg_plugins_proto_main_proto_rawDescGZIP(), []int{37} +} + +func (x *OAuthCheckLoginStatusResponse) GetAccountId() string { + if x != nil { + return x.AccountId + } + return "" +} + +func (x *OAuthCheckLoginStatusResponse) GetAccountName() string { + if x != nil { + return x.AccountName + } + return "" +} + +func (x *OAuthCheckLoginStatusResponse) GetLivemode() bool { + if x != nil { + return x.Livemode + } + return false +} + +func (x *OAuthCheckLoginStatusResponse) GetLoggedIn() bool { + if x != nil { + return x.LoggedIn + } + return false +} + var File_pkg_plugins_proto_main_proto protoreflect.FileDescriptor const file_pkg_plugins_proto_main_proto_rawDesc = "" + @@ -1570,7 +2042,16 @@ const file_pkg_plugins_proto_main_proto_rawDesc = "" + "account_id\x18\x01 \x01(\tR\taccountId\x12!\n" + "\faccount_name\x18\x02 \x01(\tR\vaccountName\x12\x1a\n" + "\blivemode\x18\x03 \x01(\bR\blivemode\x12\x1a\n" + - "\bswitched\x18\x04 \x01(\bR\bswitched\"7\n" + + "\bswitched\x18\x04 \x01(\bR\bswitched\"\x1f\n" + + "\x1dListAuthorizedAccountsRequest\"M\n" + + "\x11AuthorizedAccount\x12\x0e\n" + + "\x02id\x18\x01 \x01(\tR\x02id\x12\x12\n" + + "\x04name\x18\x02 \x01(\tR\x04name\x12\x14\n" + + "\x05modes\x18\x03 \x03(\tR\x05modes\"\xab\x01\n" + + "\x1eListAuthorizedAccountsResponse\x124\n" + + "\baccounts\x18\x01 \x03(\v2\x18.proto.AuthorizedAccountR\baccounts\x12*\n" + + "\x11active_account_id\x18\x02 \x01(\tR\x0factiveAccountId\x12'\n" + + "\x0factive_livemode\x18\x03 \x01(\bR\x0eactiveLivemode\"7\n" + "\fLoginRequest\x12'\n" + "\x0ftimeout_seconds\x18\x01 \x01(\x05R\x0etimeoutSeconds\"\x8a\x01\n" + "\rLoginResponse\x12\x1d\n" + @@ -1578,12 +2059,35 @@ const file_pkg_plugins_proto_main_proto_rawDesc = "" + "account_id\x18\x01 \x01(\tR\taccountId\x12!\n" + "\faccount_name\x18\x02 \x01(\tR\vaccountName\x12\x1a\n" + "\blivemode\x18\x03 \x01(\bR\blivemode\x12\x1b\n" + + "\tlogged_in\x18\x04 \x01(\bR\bloggedIn\"\x1b\n" + + "\x19OAuthInitiateLoginRequest\"\x89\x01\n" + + "\x1aOAuthInitiateLoginResponse\x12\x1f\n" + + "\vbrowser_url\x18\x01 \x01(\tR\n" + + "browserUrl\x12+\n" + + "\x11verification_code\x18\x02 \x01(\tR\x10verificationCode\x12\x1d\n" + + "\n" + + "expires_in\x18\x03 \x01(\x05R\texpiresIn\"\x1e\n" + + "\x1cOAuthFindPendingLoginRequest\"\xa2\x01\n" + + "\x1dOAuthFindPendingLoginResponse\x12\x14\n" + + "\x05found\x18\x01 \x01(\bR\x05found\x12\x1f\n" + + "\vbrowser_url\x18\x02 \x01(\tR\n" + + "browserUrl\x12+\n" + + "\x11verification_code\x18\x03 \x01(\tR\x10verificationCode\x12\x1d\n" + + "\n" + + "expires_in\x18\x04 \x01(\x05R\texpiresIn\"\x1e\n" + + "\x1cOAuthCheckLoginStatusRequest\"\x9a\x01\n" + + "\x1dOAuthCheckLoginStatusResponse\x12\x1d\n" + + "\n" + + "account_id\x18\x01 \x01(\tR\taccountId\x12!\n" + + "\faccount_name\x18\x02 \x01(\tR\vaccountName\x12\x1a\n" + + "\blivemode\x18\x03 \x01(\bR\blivemode\x12\x1b\n" + "\tlogged_in\x18\x04 \x01(\bR\bloggedIn2\xd8\x01\n" + "\x04Main\x12A\n" + "\n" + "RunCommand\x12\x18.proto.RunCommandRequest\x1a\x19.proto.RunCommandResponse\x12D\n" + "\vPostInstall\x12\x19.proto.PostInstallRequest\x1a\x1a.proto.PostInstallResponse\x12G\n" + - "\fPreUninstall\x12\x1a.proto.PreUninstallRequest\x1a\x1b.proto.PreUninstallResponse2\xaa\a\n" + + "\fPreUninstall\x12\x1a.proto.PreUninstallRequest\x1a\x1b.proto.PreUninstallResponse2\xb4\n" + + "\n" + "\rCoreCLIHelper\x12/\n" + "\x04Echo\x12\x12.proto.EchoRequest\x1a\x13.proto.EchoResponse\x12J\n" + "\rSendAnalytics\x12\x1b.proto.SendAnalyticsRequest\x1a\x1c.proto.SendAnalyticsResponse\x12\\\n" + @@ -1594,8 +2098,12 @@ const file_pkg_plugins_proto_main_proto_rawDesc = "" + "\rRunPeerPlugin\x12\x1b.proto.RunPeerPluginRequest\x1a\x1c.proto.RunPeerPluginResponse\x12Y\n" + "\x12ResolveCredentials\x12 .proto.ResolveCredentialsRequest\x1a!.proto.ResolveCredentialsResponse\x12c\n" + "\x1cResolveCredentialsForAnyMode\x12 .proto.ResolveCredentialsRequest\x1a!.proto.ResolveCredentialsResponse\x12J\n" + - "\rSwitchContext\x12\x1b.proto.SwitchContextRequest\x1a\x1c.proto.SwitchContextResponse\x122\n" + - "\x05Login\x12\x13.proto.LoginRequest\x1a\x14.proto.LoginResponseB,Z*github.com/stripe/stripe-cli/plugins/protob\x06proto3" + "\rSwitchContext\x12\x1b.proto.SwitchContextRequest\x1a\x1c.proto.SwitchContextResponse\x12e\n" + + "\x16ListAuthorizedAccounts\x12$.proto.ListAuthorizedAccountsRequest\x1a%.proto.ListAuthorizedAccountsResponse\x122\n" + + "\x05Login\x12\x13.proto.LoginRequest\x1a\x14.proto.LoginResponse\x12Y\n" + + "\x12OAuthInitiateLogin\x12 .proto.OAuthInitiateLoginRequest\x1a!.proto.OAuthInitiateLoginResponse\x12b\n" + + "\x15OAuthFindPendingLogin\x12#.proto.OAuthFindPendingLoginRequest\x1a$.proto.OAuthFindPendingLoginResponse\x12b\n" + + "\x15OAuthCheckLoginStatus\x12#.proto.OAuthCheckLoginStatusRequest\x1a$.proto.OAuthCheckLoginStatusResponseB,Z*github.com/stripe/stripe-cli/plugins/protob\x06proto3" var ( file_pkg_plugins_proto_main_proto_rawDescOnce sync.Once @@ -1609,7 +2117,7 @@ func file_pkg_plugins_proto_main_proto_rawDescGZIP() []byte { return file_pkg_plugins_proto_main_proto_rawDescData } -var file_pkg_plugins_proto_main_proto_msgTypes = make([]protoimpl.MessageInfo, 29) +var file_pkg_plugins_proto_main_proto_msgTypes = make([]protoimpl.MessageInfo, 38) var file_pkg_plugins_proto_main_proto_goTypes = []any{ (*RunCommandRequest)(nil), // 0: proto.RunCommandRequest (*RunCommandResponse)(nil), // 1: proto.RunCommandResponse @@ -1638,8 +2146,17 @@ var file_pkg_plugins_proto_main_proto_goTypes = []any{ (*ResolveCredentialsResponse)(nil), // 24: proto.ResolveCredentialsResponse (*SwitchContextRequest)(nil), // 25: proto.SwitchContextRequest (*SwitchContextResponse)(nil), // 26: proto.SwitchContextResponse - (*LoginRequest)(nil), // 27: proto.LoginRequest - (*LoginResponse)(nil), // 28: proto.LoginResponse + (*ListAuthorizedAccountsRequest)(nil), // 27: proto.ListAuthorizedAccountsRequest + (*AuthorizedAccount)(nil), // 28: proto.AuthorizedAccount + (*ListAuthorizedAccountsResponse)(nil), // 29: proto.ListAuthorizedAccountsResponse + (*LoginRequest)(nil), // 30: proto.LoginRequest + (*LoginResponse)(nil), // 31: proto.LoginResponse + (*OAuthInitiateLoginRequest)(nil), // 32: proto.OAuthInitiateLoginRequest + (*OAuthInitiateLoginResponse)(nil), // 33: proto.OAuthInitiateLoginResponse + (*OAuthFindPendingLoginRequest)(nil), // 34: proto.OAuthFindPendingLoginRequest + (*OAuthFindPendingLoginResponse)(nil), // 35: proto.OAuthFindPendingLoginResponse + (*OAuthCheckLoginStatusRequest)(nil), // 36: proto.OAuthCheckLoginStatusRequest + (*OAuthCheckLoginStatusResponse)(nil), // 37: proto.OAuthCheckLoginStatusResponse } var file_pkg_plugins_proto_main_proto_depIdxs = []int32{ 6, // 0: proto.RunCommandRequest.additional_info:type_name -> proto.AdditionalInfo @@ -1647,39 +2164,48 @@ var file_pkg_plugins_proto_main_proto_depIdxs = []int32{ 6, // 2: proto.PreUninstallRequest.additional_info:type_name -> proto.AdditionalInfo 7, // 3: proto.AdditionalInfo.is_terminal:type_name -> proto.IsTerminal 8, // 4: proto.AdditionalInfo.terminal_dimensions:type_name -> proto.TerminalDimensions - 0, // 5: proto.Main.RunCommand:input_type -> proto.RunCommandRequest - 2, // 6: proto.Main.PostInstall:input_type -> proto.PostInstallRequest - 4, // 7: proto.Main.PreUninstall:input_type -> proto.PreUninstallRequest - 9, // 8: proto.CoreCLIHelper.Echo:input_type -> proto.EchoRequest - 11, // 9: proto.CoreCLIHelper.SendAnalytics:input_type -> proto.SendAnalyticsRequest - 13, // 10: proto.CoreCLIHelper.KeychainGetPassword:input_type -> proto.KeychainGetPasswordRequest - 15, // 11: proto.CoreCLIHelper.KeychainSetPassword:input_type -> proto.KeychainSetPasswordRequest - 17, // 12: proto.CoreCLIHelper.KeychainDeletePassword:input_type -> proto.KeychainDeletePasswordRequest - 19, // 13: proto.CoreCLIHelper.KeychainFindCredentials:input_type -> proto.KeychainFindCredentialsRequest - 21, // 14: proto.CoreCLIHelper.RunPeerPlugin:input_type -> proto.RunPeerPluginRequest - 23, // 15: proto.CoreCLIHelper.ResolveCredentials:input_type -> proto.ResolveCredentialsRequest - 23, // 16: proto.CoreCLIHelper.ResolveCredentialsForAnyMode:input_type -> proto.ResolveCredentialsRequest - 25, // 17: proto.CoreCLIHelper.SwitchContext:input_type -> proto.SwitchContextRequest - 27, // 18: proto.CoreCLIHelper.Login:input_type -> proto.LoginRequest - 1, // 19: proto.Main.RunCommand:output_type -> proto.RunCommandResponse - 3, // 20: proto.Main.PostInstall:output_type -> proto.PostInstallResponse - 5, // 21: proto.Main.PreUninstall:output_type -> proto.PreUninstallResponse - 10, // 22: proto.CoreCLIHelper.Echo:output_type -> proto.EchoResponse - 12, // 23: proto.CoreCLIHelper.SendAnalytics:output_type -> proto.SendAnalyticsResponse - 14, // 24: proto.CoreCLIHelper.KeychainGetPassword:output_type -> proto.KeychainGetPasswordResponse - 16, // 25: proto.CoreCLIHelper.KeychainSetPassword:output_type -> proto.KeychainSetPasswordResponse - 18, // 26: proto.CoreCLIHelper.KeychainDeletePassword:output_type -> proto.KeychainDeletePasswordResponse - 20, // 27: proto.CoreCLIHelper.KeychainFindCredentials:output_type -> proto.KeychainFindCredentialsResponse - 22, // 28: proto.CoreCLIHelper.RunPeerPlugin:output_type -> proto.RunPeerPluginResponse - 24, // 29: proto.CoreCLIHelper.ResolveCredentials:output_type -> proto.ResolveCredentialsResponse - 24, // 30: proto.CoreCLIHelper.ResolveCredentialsForAnyMode:output_type -> proto.ResolveCredentialsResponse - 26, // 31: proto.CoreCLIHelper.SwitchContext:output_type -> proto.SwitchContextResponse - 28, // 32: proto.CoreCLIHelper.Login:output_type -> proto.LoginResponse - 19, // [19:33] is the sub-list for method output_type - 5, // [5:19] is the sub-list for method input_type - 5, // [5:5] is the sub-list for extension type_name - 5, // [5:5] is the sub-list for extension extendee - 0, // [0:5] is the sub-list for field type_name + 28, // 5: proto.ListAuthorizedAccountsResponse.accounts:type_name -> proto.AuthorizedAccount + 0, // 6: proto.Main.RunCommand:input_type -> proto.RunCommandRequest + 2, // 7: proto.Main.PostInstall:input_type -> proto.PostInstallRequest + 4, // 8: proto.Main.PreUninstall:input_type -> proto.PreUninstallRequest + 9, // 9: proto.CoreCLIHelper.Echo:input_type -> proto.EchoRequest + 11, // 10: proto.CoreCLIHelper.SendAnalytics:input_type -> proto.SendAnalyticsRequest + 13, // 11: proto.CoreCLIHelper.KeychainGetPassword:input_type -> proto.KeychainGetPasswordRequest + 15, // 12: proto.CoreCLIHelper.KeychainSetPassword:input_type -> proto.KeychainSetPasswordRequest + 17, // 13: proto.CoreCLIHelper.KeychainDeletePassword:input_type -> proto.KeychainDeletePasswordRequest + 19, // 14: proto.CoreCLIHelper.KeychainFindCredentials:input_type -> proto.KeychainFindCredentialsRequest + 21, // 15: proto.CoreCLIHelper.RunPeerPlugin:input_type -> proto.RunPeerPluginRequest + 23, // 16: proto.CoreCLIHelper.ResolveCredentials:input_type -> proto.ResolveCredentialsRequest + 23, // 17: proto.CoreCLIHelper.ResolveCredentialsForAnyMode:input_type -> proto.ResolveCredentialsRequest + 25, // 18: proto.CoreCLIHelper.SwitchContext:input_type -> proto.SwitchContextRequest + 27, // 19: proto.CoreCLIHelper.ListAuthorizedAccounts:input_type -> proto.ListAuthorizedAccountsRequest + 30, // 20: proto.CoreCLIHelper.Login:input_type -> proto.LoginRequest + 32, // 21: proto.CoreCLIHelper.OAuthInitiateLogin:input_type -> proto.OAuthInitiateLoginRequest + 34, // 22: proto.CoreCLIHelper.OAuthFindPendingLogin:input_type -> proto.OAuthFindPendingLoginRequest + 36, // 23: proto.CoreCLIHelper.OAuthCheckLoginStatus:input_type -> proto.OAuthCheckLoginStatusRequest + 1, // 24: proto.Main.RunCommand:output_type -> proto.RunCommandResponse + 3, // 25: proto.Main.PostInstall:output_type -> proto.PostInstallResponse + 5, // 26: proto.Main.PreUninstall:output_type -> proto.PreUninstallResponse + 10, // 27: proto.CoreCLIHelper.Echo:output_type -> proto.EchoResponse + 12, // 28: proto.CoreCLIHelper.SendAnalytics:output_type -> proto.SendAnalyticsResponse + 14, // 29: proto.CoreCLIHelper.KeychainGetPassword:output_type -> proto.KeychainGetPasswordResponse + 16, // 30: proto.CoreCLIHelper.KeychainSetPassword:output_type -> proto.KeychainSetPasswordResponse + 18, // 31: proto.CoreCLIHelper.KeychainDeletePassword:output_type -> proto.KeychainDeletePasswordResponse + 20, // 32: proto.CoreCLIHelper.KeychainFindCredentials:output_type -> proto.KeychainFindCredentialsResponse + 22, // 33: proto.CoreCLIHelper.RunPeerPlugin:output_type -> proto.RunPeerPluginResponse + 24, // 34: proto.CoreCLIHelper.ResolveCredentials:output_type -> proto.ResolveCredentialsResponse + 24, // 35: proto.CoreCLIHelper.ResolveCredentialsForAnyMode:output_type -> proto.ResolveCredentialsResponse + 26, // 36: proto.CoreCLIHelper.SwitchContext:output_type -> proto.SwitchContextResponse + 29, // 37: proto.CoreCLIHelper.ListAuthorizedAccounts:output_type -> proto.ListAuthorizedAccountsResponse + 31, // 38: proto.CoreCLIHelper.Login:output_type -> proto.LoginResponse + 33, // 39: proto.CoreCLIHelper.OAuthInitiateLogin:output_type -> proto.OAuthInitiateLoginResponse + 35, // 40: proto.CoreCLIHelper.OAuthFindPendingLogin:output_type -> proto.OAuthFindPendingLoginResponse + 37, // 41: proto.CoreCLIHelper.OAuthCheckLoginStatus:output_type -> proto.OAuthCheckLoginStatusResponse + 24, // [24:42] is the sub-list for method output_type + 6, // [6:24] is the sub-list for method input_type + 6, // [6:6] is the sub-list for extension type_name + 6, // [6:6] is the sub-list for extension extendee + 0, // [0:6] is the sub-list for field type_name } func init() { file_pkg_plugins_proto_main_proto_init() } @@ -1693,7 +2219,7 @@ func file_pkg_plugins_proto_main_proto_init() { GoPackagePath: reflect.TypeOf(x{}).PkgPath(), RawDescriptor: unsafe.Slice(unsafe.StringData(file_pkg_plugins_proto_main_proto_rawDesc), len(file_pkg_plugins_proto_main_proto_rawDesc)), NumEnums: 0, - NumMessages: 29, + NumMessages: 38, NumExtensions: 0, NumServices: 2, }, diff --git a/pkg/plugins/proto/main.proto b/pkg/plugins/proto/main.proto index ba4e8585f..9bbc7b117 100644 --- a/pkg/plugins/proto/main.proto +++ b/pkg/plugins/proto/main.proto @@ -86,12 +86,41 @@ service CoreCLIHelper { // same way `stripe switch context` does. If account_id is empty, shows an // interactive picker. rpc SwitchContext(SwitchContextRequest) returns (SwitchContextResponse); + // ListAuthorizedAccounts returns the Stripe accounts the current session is + // authorized for, and which account and mode are currently active. + rpc ListAuthorizedAccounts(ListAuthorizedAccountsRequest) returns (ListAuthorizedAccountsResponse); // Login starts a Stripe CLI login, the same way `stripe login --new-session` // does when run interactively: it revokes any existing OAuth session first // (so this works even if the stored credential is expired or revoked), // then runs the normal login flow, printing the same output and opening // the browser only after the user presses enter. rpc Login(LoginRequest) returns (LoginResponse); + // The OAuth-prefixed RPCs below (OAuthInitiateLogin, OAuthFindPendingLogin, + // OAuthCheckLoginStatus) are low-level building blocks for a non-interactive, + // resumable login flow. Prefer Login unless you specifically need + // non-blocking, resumable behavior (e.g. driving your own retry loop). + // OAuthInitiateLogin starts (or resumes) a non-interactive OAuth + // device-code login: it returns immediately with a browser URL and + // verification code for the plugin to present, instead of blocking until + // the user completes it like Login does. Calling it again before the + // previous attempt completes or expires returns the same browser URL and + // verification code rather than minting a new device code, so a plugin (or + // an agent driving it) can safely retry without orphaning an in-flight + // login - unlike `stripe login --non-interactive`, which always mints a + // fresh device code. Call OAuthCheckLoginStatus to check whether the user + // has completed it. + rpc OAuthInitiateLogin(OAuthInitiateLoginRequest) returns (OAuthInitiateLoginResponse); + // OAuthFindPendingLogin looks for an OAuth device-code login already in + // progress - started by this call chain, another plugin, or + // `stripe login --non-interactive` - without starting a new one. Useful + // for a non-interactive caller to check whether it should resume an + // existing attempt instead of calling OAuthInitiateLogin. + rpc OAuthFindPendingLogin(OAuthFindPendingLoginRequest) returns (OAuthFindPendingLoginResponse); + // OAuthCheckLoginStatus makes a single, non-blocking check on whether the + // login started by OAuthInitiateLogin has completed. It does not wait for + // the user; call it again later (e.g. on your own poll loop) to keep + // checking. + rpc OAuthCheckLoginStatus(OAuthCheckLoginStatusRequest) returns (OAuthCheckLoginStatusResponse); } message EchoRequest { @@ -179,6 +208,24 @@ message SwitchContextResponse { bool switched = 4; } +message ListAuthorizedAccountsRequest { +} + +message AuthorizedAccount { + string id = 1; + string name = 2; + // modes lists the API mode(s) ("test", "live") this account grants access to. + repeated string modes = 3; +} + +message ListAuthorizedAccountsResponse { + repeated AuthorizedAccount accounts = 1; + // active_account_id and active_livemode identify which authorized account and mode are + // currently active; empty/false if none is active yet. + string active_account_id = 2; + bool active_livemode = 3; +} + message LoginRequest { // timeout_seconds bounds how long Login waits for the user to complete // authentication. If unset or 0, it waits indefinitely, matching @@ -196,3 +243,41 @@ message LoginResponse { // URL), not a resumption of this one. bool logged_in = 4; } + +message OAuthInitiateLoginRequest { +} + +message OAuthInitiateLoginResponse { + string browser_url = 1; + string verification_code = 2; + // expires_in is the number of seconds remaining before the browser_url/ + // verification_code expire and a new OAuthInitiateLogin call is required. + int32 expires_in = 3; +} + +message OAuthFindPendingLoginRequest { +} + +message OAuthFindPendingLoginResponse { + // found is false if there is no pending login attempt, or it has + // expired; in that case the other fields are empty. + bool found = 1; + string browser_url = 2; + string verification_code = 3; + // expires_in is the number of seconds remaining before the browser_url/ + // verification_code expire and a new OAuthInitiateLogin call is required. + int32 expires_in = 4; +} + +message OAuthCheckLoginStatusRequest { +} + +message OAuthCheckLoginStatusResponse { + string account_id = 1; + string account_name = 2; + bool livemode = 3; + // logged_in is false if the user hasn't completed authentication yet; in + // that case the other fields are empty and callers should call + // OAuthCheckLoginStatus again later to keep checking. + bool logged_in = 4; +} diff --git a/pkg/plugins/proto/main_grpc.pb.go b/pkg/plugins/proto/main_grpc.pb.go index ea1c77929..e5961bb4d 100644 --- a/pkg/plugins/proto/main_grpc.pb.go +++ b/pkg/plugins/proto/main_grpc.pb.go @@ -215,7 +215,11 @@ const ( CoreCLIHelper_ResolveCredentials_FullMethodName = "/proto.CoreCLIHelper/ResolveCredentials" CoreCLIHelper_ResolveCredentialsForAnyMode_FullMethodName = "/proto.CoreCLIHelper/ResolveCredentialsForAnyMode" CoreCLIHelper_SwitchContext_FullMethodName = "/proto.CoreCLIHelper/SwitchContext" + CoreCLIHelper_ListAuthorizedAccounts_FullMethodName = "/proto.CoreCLIHelper/ListAuthorizedAccounts" CoreCLIHelper_Login_FullMethodName = "/proto.CoreCLIHelper/Login" + CoreCLIHelper_OAuthInitiateLogin_FullMethodName = "/proto.CoreCLIHelper/OAuthInitiateLogin" + CoreCLIHelper_OAuthFindPendingLogin_FullMethodName = "/proto.CoreCLIHelper/OAuthFindPendingLogin" + CoreCLIHelper_OAuthCheckLoginStatus_FullMethodName = "/proto.CoreCLIHelper/OAuthCheckLoginStatus" ) // CoreCLIHelperClient is the client API for CoreCLIHelper service. @@ -240,12 +244,41 @@ type CoreCLIHelperClient interface { // same way `stripe switch context` does. If account_id is empty, shows an // interactive picker. SwitchContext(ctx context.Context, in *SwitchContextRequest, opts ...grpc.CallOption) (*SwitchContextResponse, error) + // ListAuthorizedAccounts returns the Stripe accounts the current session is + // authorized for, and which account and mode are currently active. + ListAuthorizedAccounts(ctx context.Context, in *ListAuthorizedAccountsRequest, opts ...grpc.CallOption) (*ListAuthorizedAccountsResponse, error) // Login starts a Stripe CLI login, the same way `stripe login --new-session` // does when run interactively: it revokes any existing OAuth session first // (so this works even if the stored credential is expired or revoked), // then runs the normal login flow, printing the same output and opening // the browser only after the user presses enter. Login(ctx context.Context, in *LoginRequest, opts ...grpc.CallOption) (*LoginResponse, error) + // The OAuth-prefixed RPCs below (OAuthInitiateLogin, OAuthFindPendingLogin, + // OAuthCheckLoginStatus) are low-level building blocks for a non-interactive, + // resumable login flow. Prefer Login unless you specifically need + // non-blocking, resumable behavior (e.g. driving your own retry loop). + // OAuthInitiateLogin starts (or resumes) a non-interactive OAuth + // device-code login: it returns immediately with a browser URL and + // verification code for the plugin to present, instead of blocking until + // the user completes it like Login does. Calling it again before the + // previous attempt completes or expires returns the same browser URL and + // verification code rather than minting a new device code, so a plugin (or + // an agent driving it) can safely retry without orphaning an in-flight + // login - unlike `stripe login --non-interactive`, which always mints a + // fresh device code. Call OAuthCheckLoginStatus to check whether the user + // has completed it. + OAuthInitiateLogin(ctx context.Context, in *OAuthInitiateLoginRequest, opts ...grpc.CallOption) (*OAuthInitiateLoginResponse, error) + // OAuthFindPendingLogin looks for an OAuth device-code login already in + // progress - started by this call chain, another plugin, or + // `stripe login --non-interactive` - without starting a new one. Useful + // for a non-interactive caller to check whether it should resume an + // existing attempt instead of calling OAuthInitiateLogin. + OAuthFindPendingLogin(ctx context.Context, in *OAuthFindPendingLoginRequest, opts ...grpc.CallOption) (*OAuthFindPendingLoginResponse, error) + // OAuthCheckLoginStatus makes a single, non-blocking check on whether the + // login started by OAuthInitiateLogin has completed. It does not wait for + // the user; call it again later (e.g. on your own poll loop) to keep + // checking. + OAuthCheckLoginStatus(ctx context.Context, in *OAuthCheckLoginStatusRequest, opts ...grpc.CallOption) (*OAuthCheckLoginStatusResponse, error) } type coreCLIHelperClient struct { @@ -357,6 +390,16 @@ func (c *coreCLIHelperClient) SwitchContext(ctx context.Context, in *SwitchConte return out, nil } +func (c *coreCLIHelperClient) ListAuthorizedAccounts(ctx context.Context, in *ListAuthorizedAccountsRequest, opts ...grpc.CallOption) (*ListAuthorizedAccountsResponse, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(ListAuthorizedAccountsResponse) + err := c.cc.Invoke(ctx, CoreCLIHelper_ListAuthorizedAccounts_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + func (c *coreCLIHelperClient) Login(ctx context.Context, in *LoginRequest, opts ...grpc.CallOption) (*LoginResponse, error) { cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) out := new(LoginResponse) @@ -367,6 +410,36 @@ func (c *coreCLIHelperClient) Login(ctx context.Context, in *LoginRequest, opts return out, nil } +func (c *coreCLIHelperClient) OAuthInitiateLogin(ctx context.Context, in *OAuthInitiateLoginRequest, opts ...grpc.CallOption) (*OAuthInitiateLoginResponse, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(OAuthInitiateLoginResponse) + err := c.cc.Invoke(ctx, CoreCLIHelper_OAuthInitiateLogin_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *coreCLIHelperClient) OAuthFindPendingLogin(ctx context.Context, in *OAuthFindPendingLoginRequest, opts ...grpc.CallOption) (*OAuthFindPendingLoginResponse, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(OAuthFindPendingLoginResponse) + err := c.cc.Invoke(ctx, CoreCLIHelper_OAuthFindPendingLogin_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *coreCLIHelperClient) OAuthCheckLoginStatus(ctx context.Context, in *OAuthCheckLoginStatusRequest, opts ...grpc.CallOption) (*OAuthCheckLoginStatusResponse, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(OAuthCheckLoginStatusResponse) + err := c.cc.Invoke(ctx, CoreCLIHelper_OAuthCheckLoginStatus_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + // CoreCLIHelperServer is the server API for CoreCLIHelper service. // All implementations must embed UnimplementedCoreCLIHelperServer // for forward compatibility. @@ -389,12 +462,41 @@ type CoreCLIHelperServer interface { // same way `stripe switch context` does. If account_id is empty, shows an // interactive picker. SwitchContext(context.Context, *SwitchContextRequest) (*SwitchContextResponse, error) + // ListAuthorizedAccounts returns the Stripe accounts the current session is + // authorized for, and which account and mode are currently active. + ListAuthorizedAccounts(context.Context, *ListAuthorizedAccountsRequest) (*ListAuthorizedAccountsResponse, error) // Login starts a Stripe CLI login, the same way `stripe login --new-session` // does when run interactively: it revokes any existing OAuth session first // (so this works even if the stored credential is expired or revoked), // then runs the normal login flow, printing the same output and opening // the browser only after the user presses enter. Login(context.Context, *LoginRequest) (*LoginResponse, error) + // The OAuth-prefixed RPCs below (OAuthInitiateLogin, OAuthFindPendingLogin, + // OAuthCheckLoginStatus) are low-level building blocks for a non-interactive, + // resumable login flow. Prefer Login unless you specifically need + // non-blocking, resumable behavior (e.g. driving your own retry loop). + // OAuthInitiateLogin starts (or resumes) a non-interactive OAuth + // device-code login: it returns immediately with a browser URL and + // verification code for the plugin to present, instead of blocking until + // the user completes it like Login does. Calling it again before the + // previous attempt completes or expires returns the same browser URL and + // verification code rather than minting a new device code, so a plugin (or + // an agent driving it) can safely retry without orphaning an in-flight + // login - unlike `stripe login --non-interactive`, which always mints a + // fresh device code. Call OAuthCheckLoginStatus to check whether the user + // has completed it. + OAuthInitiateLogin(context.Context, *OAuthInitiateLoginRequest) (*OAuthInitiateLoginResponse, error) + // OAuthFindPendingLogin looks for an OAuth device-code login already in + // progress - started by this call chain, another plugin, or + // `stripe login --non-interactive` - without starting a new one. Useful + // for a non-interactive caller to check whether it should resume an + // existing attempt instead of calling OAuthInitiateLogin. + OAuthFindPendingLogin(context.Context, *OAuthFindPendingLoginRequest) (*OAuthFindPendingLoginResponse, error) + // OAuthCheckLoginStatus makes a single, non-blocking check on whether the + // login started by OAuthInitiateLogin has completed. It does not wait for + // the user; call it again later (e.g. on your own poll loop) to keep + // checking. + OAuthCheckLoginStatus(context.Context, *OAuthCheckLoginStatusRequest) (*OAuthCheckLoginStatusResponse, error) mustEmbedUnimplementedCoreCLIHelperServer() } @@ -435,9 +537,21 @@ func (UnimplementedCoreCLIHelperServer) ResolveCredentialsForAnyMode(context.Con func (UnimplementedCoreCLIHelperServer) SwitchContext(context.Context, *SwitchContextRequest) (*SwitchContextResponse, error) { return nil, status.Error(codes.Unimplemented, "method SwitchContext not implemented") } +func (UnimplementedCoreCLIHelperServer) ListAuthorizedAccounts(context.Context, *ListAuthorizedAccountsRequest) (*ListAuthorizedAccountsResponse, error) { + return nil, status.Error(codes.Unimplemented, "method ListAuthorizedAccounts not implemented") +} func (UnimplementedCoreCLIHelperServer) Login(context.Context, *LoginRequest) (*LoginResponse, error) { return nil, status.Error(codes.Unimplemented, "method Login not implemented") } +func (UnimplementedCoreCLIHelperServer) OAuthInitiateLogin(context.Context, *OAuthInitiateLoginRequest) (*OAuthInitiateLoginResponse, error) { + return nil, status.Error(codes.Unimplemented, "method OAuthInitiateLogin not implemented") +} +func (UnimplementedCoreCLIHelperServer) OAuthFindPendingLogin(context.Context, *OAuthFindPendingLoginRequest) (*OAuthFindPendingLoginResponse, error) { + return nil, status.Error(codes.Unimplemented, "method OAuthFindPendingLogin not implemented") +} +func (UnimplementedCoreCLIHelperServer) OAuthCheckLoginStatus(context.Context, *OAuthCheckLoginStatusRequest) (*OAuthCheckLoginStatusResponse, error) { + return nil, status.Error(codes.Unimplemented, "method OAuthCheckLoginStatus not implemented") +} func (UnimplementedCoreCLIHelperServer) mustEmbedUnimplementedCoreCLIHelperServer() {} func (UnimplementedCoreCLIHelperServer) testEmbeddedByValue() {} @@ -639,6 +753,24 @@ func _CoreCLIHelper_SwitchContext_Handler(srv interface{}, ctx context.Context, return interceptor(ctx, in, info, handler) } +func _CoreCLIHelper_ListAuthorizedAccounts_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(ListAuthorizedAccountsRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(CoreCLIHelperServer).ListAuthorizedAccounts(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: CoreCLIHelper_ListAuthorizedAccounts_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(CoreCLIHelperServer).ListAuthorizedAccounts(ctx, req.(*ListAuthorizedAccountsRequest)) + } + return interceptor(ctx, in, info, handler) +} + func _CoreCLIHelper_Login_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { in := new(LoginRequest) if err := dec(in); err != nil { @@ -657,6 +789,60 @@ func _CoreCLIHelper_Login_Handler(srv interface{}, ctx context.Context, dec func return interceptor(ctx, in, info, handler) } +func _CoreCLIHelper_OAuthInitiateLogin_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(OAuthInitiateLoginRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(CoreCLIHelperServer).OAuthInitiateLogin(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: CoreCLIHelper_OAuthInitiateLogin_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(CoreCLIHelperServer).OAuthInitiateLogin(ctx, req.(*OAuthInitiateLoginRequest)) + } + return interceptor(ctx, in, info, handler) +} + +func _CoreCLIHelper_OAuthFindPendingLogin_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(OAuthFindPendingLoginRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(CoreCLIHelperServer).OAuthFindPendingLogin(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: CoreCLIHelper_OAuthFindPendingLogin_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(CoreCLIHelperServer).OAuthFindPendingLogin(ctx, req.(*OAuthFindPendingLoginRequest)) + } + return interceptor(ctx, in, info, handler) +} + +func _CoreCLIHelper_OAuthCheckLoginStatus_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(OAuthCheckLoginStatusRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(CoreCLIHelperServer).OAuthCheckLoginStatus(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: CoreCLIHelper_OAuthCheckLoginStatus_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(CoreCLIHelperServer).OAuthCheckLoginStatus(ctx, req.(*OAuthCheckLoginStatusRequest)) + } + return interceptor(ctx, in, info, handler) +} + // CoreCLIHelper_ServiceDesc is the grpc.ServiceDesc for CoreCLIHelper service. // It's only intended for direct use with grpc.RegisterService, // and not to be introspected or modified (even as a copy) @@ -704,10 +890,26 @@ var CoreCLIHelper_ServiceDesc = grpc.ServiceDesc{ MethodName: "SwitchContext", Handler: _CoreCLIHelper_SwitchContext_Handler, }, + { + MethodName: "ListAuthorizedAccounts", + Handler: _CoreCLIHelper_ListAuthorizedAccounts_Handler, + }, { MethodName: "Login", Handler: _CoreCLIHelper_Login_Handler, }, + { + MethodName: "OAuthInitiateLogin", + Handler: _CoreCLIHelper_OAuthInitiateLogin_Handler, + }, + { + MethodName: "OAuthFindPendingLogin", + Handler: _CoreCLIHelper_OAuthFindPendingLogin_Handler, + }, + { + MethodName: "OAuthCheckLoginStatus", + Handler: _CoreCLIHelper_OAuthCheckLoginStatus_Handler, + }, }, Streams: []grpc.StreamDesc{}, Metadata: "pkg/plugins/proto/main.proto", diff --git a/pkg/plugins/reported_error.go b/pkg/plugins/reported_error.go new file mode 100644 index 000000000..b472ff226 --- /dev/null +++ b/pkg/plugins/reported_error.go @@ -0,0 +1,34 @@ +package plugins + +import "errors" + +// pluginReportedError marks an error that the plugin process itself produced after +// it had started, and so has already written to the terminal. +// +// Run returns errors from two distinct situations, and the caller has to tell them +// apart. A plugin that started and then failed has printed its own message, so +// printing it again would duplicate it. Everything else -- an install that failed, +// a plugin refused for being too old to read the config file, a handshake that +// never completed -- happens before the plugin is launched, so nothing has been +// printed and the caller is the only one that can say what went wrong. +type pluginReportedError struct { + err error +} + +func (e pluginReportedError) Error() string { + return e.err.Error() +} + +func (e pluginReportedError) Unwrap() error { + return e.err +} + +// PluginAlreadyReported reports whether err came from a plugin process that had +// already started, meaning the plugin has printed the message itself. +// +// A caller that exits on error must print any error for which this is false, or +// the command fails with no output at all. +func PluginAlreadyReported(err error) bool { + var reported pluginReportedError + return errors.As(err, &reported) +} diff --git a/pkg/plugins/reported_error_test.go b/pkg/plugins/reported_error_test.go new file mode 100644 index 000000000..0f88fac6b --- /dev/null +++ b/pkg/plugins/reported_error_test.go @@ -0,0 +1,41 @@ +package plugins + +import ( + "errors" + "testing" + + "github.com/spf13/viper" + "github.com/stretchr/testify/require" + + "github.com/stripe/stripe-cli/pkg/config" +) + +// The caller exits on any error from Run, so it has to print the ones nobody else +// has. Refusing a plugin for being too old happens before the plugin is launched, +// which is what made that refusal exit 1 with no output at all. +func TestPluginAlreadyReportedIsFalseForARefusal(t *testing.T) { + t.Cleanup(viper.Reset) + viper.Set(config.ConfigVersionName, config.ConfigVersionV2) + withConfigV2MinimumVersions(t, map[string]string{"apps": "2.0.0"}) + + plugin := Plugin{Shortname: "apps"} + err := plugin.refuseIfConfigTooNew("1.0.0") + + require.Error(t, err) + require.False(t, PluginAlreadyReported(err), + "a refusal happens before the plugin launches, so the caller must print it") + require.Contains(t, err.Error(), "stripe plugin upgrade apps") +} + +// An error the plugin process produced after starting is already on screen. +func TestPluginAlreadyReportedIsTrueForPluginOutput(t *testing.T) { + inner := errors.New("the plugin printed this itself") + + require.True(t, PluginAlreadyReported(pluginReportedError{inner})) + require.False(t, PluginAlreadyReported(inner)) + require.False(t, PluginAlreadyReported(nil)) + + // Unwrapping has to keep working, so callers can still match on the cause. + require.True(t, errors.Is(pluginReportedError{inner}, inner)) + require.Equal(t, inner.Error(), pluginReportedError{inner}.Error()) +} diff --git a/pkg/plugins/utilities.go b/pkg/plugins/utilities.go index 4b86c8b0b..58c79da93 100644 --- a/pkg/plugins/utilities.go +++ b/pkg/plugins/utilities.go @@ -108,22 +108,34 @@ func GetBinaryExtension() string { return "" } +// pluginsDirOverride returns the directory plugins have been pointed at instead of the +// CLI's own, or "" when they have not been. +// +// There are two ways to do that -- the STRIPE_PLUGINS_PATH environment variable, and +// PluginsPath compiled in by a `localdev` build -- and anything deciding what the CLI may +// do to a plugin directory has to ask about both. Checking only PluginsPath is what let +// auto-upgrade overwrite a plugin under STRIPE_PLUGINS_PATH: the same directory, with +// none of the protection, because the guard knew only the other spelling of it. +// +// The env var wins where both are set, matching the order these have always resolved in: +// a variable set for one invocation is a narrower statement than one baked into a binary. +func pluginsDirOverride() string { + if envPluginsPath := os.Getenv("STRIPE_PLUGINS_PATH"); envPluginsPath != "" { + return envPluginsPath + } + + return PluginsPath +} + // getPluginsDir computes where plugins are installed locally func getPluginsDir(config config.IConfig) string { - var pluginsDir string - tempEnvPluginsPath := os.Getenv("STRIPE_PLUGINS_PATH") - - switch { - case tempEnvPluginsPath != "": - pluginsDir = tempEnvPluginsPath - case PluginsPath != "": - pluginsDir = PluginsPath - default: - configPath := config.GetConfigFolder(os.Getenv("XDG_CONFIG_HOME")) - pluginsDir = filepath.Join(configPath, "plugins") + if override := pluginsDirOverride(); override != "" { + return override } - return pluginsDir + configPath := config.GetConfigFolder(os.Getenv("XDG_CONFIG_HOME")) + + return filepath.Join(configPath, "plugins") } func getLocalPluginMetadataDir(config config.IConfig) string { @@ -1166,22 +1178,44 @@ func FetchRemoteResource(ctx context.Context, url string) ([]byte, error) { // CheckLatestPluginVersion prints an upgrade hint to stderr if live metadata // has a newer version of the plugin than what is currently installed. // -// It stays quiet for a plugin that auto-updates. The upgrade check already ran before -// the command, so this would be a second lookup of the same thing on the same -// invocation, and the user would wait out both timeouts to be advised of an upgrade -// the CLI already made for them. +// It stays quiet for a plugin that auto-updates, whose owner asked the CLI to handle +// upgrades rather than be told about them -- for as long as the CLI can actually handle +// them, which the guard below is about. Where maybeAutoUpgrade already ran this +// invocation, this is a second lookup of the same thing, ending in advice about an +// upgrade the CLI just made. Where it did not run -- which is most invocations, since +// it checks at most once per autoUpgradeCheckInterval -- this would spend exactly the +// per-command request that throttle exists to avoid, and undo it from the other end. // -// What that gives up: when the pre-run check resolves a newer version from cached -// metadata, it declines to upgrade, but this would still have named that version. Such -// a run now says nothing. Keeping the hint for it would charge every auto-updating -// command an extra request, and a check that keeps declining is better reported by the -// check itself than inferred from a hint here. +// What that gives up is every case where the pre-run check knows a newer version exists +// but declines to install it: a resolution from cached metadata, or the whole interval +// after such a decline. Those runs now say nothing at all. The alternative is charging +// every auto-updating command a request to say it, and a check that keeps declining is +// better reported by the check itself than inferred from a hint here. func CheckLatestPluginVersion(ctx context.Context, config config.IConfig, fs afero.Fs, plugin Plugin, apiBaseURL, dashboardBaseURL string) { + // PluginsPath alone, deliberately narrower than the same-looking guard in + // maybeAutoUpgrade: a `localdev` build has no published release to be behind, but + // someone who merely relocated their plugins with STRIPE_PLUGINS_PATH still wants to + // hear about upgrades. Printing a line can only be wrong; installing over the + // directory can delete a build, which is why that side asks the broader question. if PluginsPath != "" { return } - if pluginUpdatesEnabled(plugin.Shortname) { + // Handing the job to the pre-run check, but only where that check will take it. + // maybeAutoUpgrade refuses a plugins directory the user pointed the CLI at, so + // deferring to it there would leave a plugin that auto-updates under + // STRIPE_PLUGINS_PATH with no upgrade and no word that one exists -- silently behind, + // on the strength of a setting asking for the opposite. + // + // Not the same as deferring across the throttle, which this still does: that decline + // is for the current invocation and some later one will upgrade, so the silence costs + // a few hours. An overridden directory is refused on every invocation there will ever + // be, so nothing arrives to break it. + // + // Asked of pluginsDirOverride rather than the environment directly, even though the + // guard above has already returned for the compiled-in half of it, so that this and + // maybeAutoUpgrade keep reading the same answer from the same place. + if pluginsDirOverride() == "" && pluginUpdatesEnabled(plugin.Shortname) { return } diff --git a/pkg/plugins/utilities_test.go b/pkg/plugins/utilities_test.go index a17dc5986..e1fc9c986 100644 --- a/pkg/plugins/utilities_test.go +++ b/pkg/plugins/utilities_test.go @@ -1157,6 +1157,107 @@ func TestCheckLatestPluginVersionSilentWhenLookupTimesOut(t *testing.T) { } } +// TestCheckLatestPluginVersionStillHintsUnderAnEnvironmentPluginsPath pins the one place +// the two plugins-path guards deliberately disagree. maybeAutoUpgrade refuses to install +// into a directory the user pointed the CLI at, whichever way they pointed it; the hint +// only goes quiet for a localdev build, which has no published release to be behind. +// Someone who relocated ordinary installs with the environment variable still wants to +// hear that an upgrade exists -- all the more so now that they will not get it silently. +func TestCheckLatestPluginVersionStillHintsUnderAnEnvironmentPluginsPath(t *testing.T) { + origPluginsPath := PluginsPath + origResolver := checkLatestPluginVersionResolver + PluginsPath = "" + t.Setenv("STRIPE_PLUGINS_PATH", "/somewhere/else") + checkLatestPluginVersionResolver = func(ctx context.Context, cfg cfgpkg.IConfig, fs afero.Fs, pluginName, apiBaseURL, dashboardBaseURL string) (*ResolvedPluginVersion, error) { + return &ResolvedPluginVersion{ + Plugin: &Plugin{ + Shortname: "myplugin", + Releases: []Release{ + {Arch: runtime.GOARCH, OS: runtime.GOOS, Version: "1.1.0", Sum: "abc123"}, + }, + }, + Version: "1.1.0", + }, nil + } + defer func() { + PluginsPath = origPluginsPath + checkLatestPluginVersionResolver = origResolver + }() + + fs := afero.NewMemMapFs() + config := &TestConfig{} + + plugin := Plugin{ + Shortname: "myplugin", + Binary: "stripe-cli-myplugin", + MagicCookieValue: "MY-COOKIE", + } + + pluginBinaryPath := fmt.Sprintf("/somewhere/else/myplugin/1.0.0/stripe-cli-myplugin%s", GetBinaryExtension()) + require.NoError(t, fs.MkdirAll(filepath.Dir(pluginBinaryPath), 0755)) + require.NoError(t, afero.WriteFile(fs, pluginBinaryPath, []byte("binary"), 0755)) + + output := captureStderr(t, func() { + CheckLatestPluginVersion(context.Background(), config, fs, plugin, stripe.DefaultAPIBaseURL, "") + }) + + require.Contains(t, output, "A newer version of the myplugin plugin is available") +} + +func TestGetPluginsDirOverrides(t *testing.T) { + // TestConfig's config folder is "/", which is why every other test in this package + // finds plugins at /plugins without arranging anything. Joined rather than written + // out because this is the one case getPluginsDir builds a path for, and Windows + // builds it with the other separator. The overrides below are handed back verbatim, + // so they are the same string everywhere. + defaultPluginsDir := filepath.Join("/", "plugins") + + tests := []struct { + name string + pluginsPathEnv string + pluginsPath string + want string + }{ + { + name: "neither, so the CLI's own config folder", + want: defaultPluginsDir, + }, + { + name: "the environment variable", + pluginsPathEnv: "/from/the/environment", + want: "/from/the/environment", + }, + { + name: "a path compiled into a localdev build", + pluginsPath: "/compiled/in", + want: "/compiled/in", + }, + { + // The order these have always resolved in, kept because a variable set for + // one invocation is a narrower statement than one baked into a binary. + name: "both, so the environment variable", + pluginsPathEnv: "/from/the/environment", + pluginsPath: "/compiled/in", + want: "/from/the/environment", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + origPluginsPath := PluginsPath + PluginsPath = tt.pluginsPath + t.Setenv("STRIPE_PLUGINS_PATH", tt.pluginsPathEnv) + defer func() { PluginsPath = origPluginsPath }() + + require.Equal(t, tt.want, getPluginsDir(&TestConfig{})) + + // What the auto-upgrade guard reads. Anything but the config folder is a + // directory the CLI was pointed at and must not install over. + require.Equal(t, tt.want != defaultPluginsDir, pluginsDirOverride() != "") + }) + } +} + func TestCheckLatestPluginVersionSilentInDevMode(t *testing.T) { origPluginsPath := PluginsPath origResolver := checkLatestPluginVersionResolver @@ -1196,6 +1297,11 @@ func TestCheckLatestPluginVersionSilentWhenPluginAutoUpdates(t *testing.T) { origUpdatesEnabled := pluginUpdatesEnabled origResolver := checkLatestPluginVersionResolver PluginsPath = "" + // A plugins directory the CLI has not been pointed at, which is what makes deferring + // to the pre-run check the right thing to do here. Pinned rather than assumed: the + // suppression this asserts is now conditional on it, so a stray variable in the + // environment running the tests would turn the whole test into its own opposite. + t.Setenv("STRIPE_PLUGINS_PATH", "") var settingReads []string pluginUpdatesEnabled = func(pluginName string) bool { @@ -1247,11 +1353,80 @@ func TestCheckLatestPluginVersionSilentWhenPluginAutoUpdates(t *testing.T) { require.Equal(t, []string{"myplugin"}, settingReads) require.Empty(t, output) - // The point is the request, not just the message: the pre-run upgrade check already - // asked this question on this invocation. + // The point is the request, not just the message. A hint here would put a lookup on + // every command of an auto-updating plugin, which is the cost + // autoUpgradeCheckInterval exists to keep maybeAutoUpgrade from imposing. require.Zero(t, resolveCalls) } +// TestCheckLatestPluginVersionHintsWhenAutoUpgradeWillNotRun covers the one state where +// both halves of the feature could go quiet at once: auto-update is on for the plugin, so +// the hint would hand the job to the pre-run check, while the plugins directory is +// overridden, so that check refuses it outright. Deferring to something that never runs +// leaves the plugin silently out of date -- the single outcome neither guard is willing to +// own, and the reason the suppression above asks whether the upgrade can happen at all. +// +// Distinct from the throttle, which is also a decline: that one is for this invocation and +// the next one may well upgrade, so staying quiet costs nothing but a few hours. An +// overridden directory is refused on every invocation, forever. +func TestCheckLatestPluginVersionHintsWhenAutoUpgradeWillNotRun(t *testing.T) { + origPluginsPath := PluginsPath + origUpdatesEnabled := pluginUpdatesEnabled + origResolver := checkLatestPluginVersionResolver + PluginsPath = "" + t.Setenv("STRIPE_PLUGINS_PATH", "/somewhere/else") + + var settingReads []string + pluginUpdatesEnabled = func(pluginName string) bool { + settingReads = append(settingReads, pluginName) + return true + } + checkLatestPluginVersionResolver = func(ctx context.Context, cfg cfgpkg.IConfig, fs afero.Fs, pluginName, apiBaseURL, dashboardBaseURL string) (*ResolvedPluginVersion, error) { + return &ResolvedPluginVersion{ + Plugin: &Plugin{ + Shortname: "myplugin", + Releases: []Release{ + {Arch: runtime.GOARCH, OS: runtime.GOOS, Version: "1.1.0", Sum: "abc123"}, + }, + }, + Version: "1.1.0", + }, nil + } + defer func() { + PluginsPath = origPluginsPath + pluginUpdatesEnabled = origUpdatesEnabled + checkLatestPluginVersionResolver = origResolver + }() + + fs := afero.NewMemMapFs() + config := &TestConfig{} + + plugin := Plugin{ + Shortname: "myplugin", + Binary: "stripe-cli-myplugin", + MagicCookieValue: "MY-COOKIE", + Releases: []Release{ + {Arch: runtime.GOARCH, OS: runtime.GOOS, Version: "1.0.0", Sum: "abc123"}, + }, + } + + pluginBinaryPath := fmt.Sprintf("/somewhere/else/myplugin/1.0.0/stripe-cli-myplugin%s", GetBinaryExtension()) + require.NoError(t, fs.MkdirAll(filepath.Dir(pluginBinaryPath), 0755)) + require.NoError(t, afero.WriteFile(fs, pluginBinaryPath, []byte("binary"), 0755)) + + output := captureStderr(t, func() { + CheckLatestPluginVersion(context.Background(), config, fs, plugin, stripe.DefaultAPIBaseURL, "") + }) + + require.Contains(t, output, "A newer version of the myplugin plugin is available") + + // The setting is not read at all. Under an overridden directory it has nothing left to + // decide, and asserting that rules out passing for the neighboring reason -- a hint + // printed because the setting happened to be off rather than because the override + // took precedence over it. + require.Empty(t, settingReads) +} + func TestIsPluginCommand(t *testing.T) { pluginCmd := &cobra.Command{ Annotations: map[string]string{"scope": "plugin"}, diff --git a/pkg/useragent/useragent.go b/pkg/useragent/useragent.go index 055eb7b4d..f2446da01 100644 --- a/pkg/useragent/useragent.go +++ b/pkg/useragent/useragent.go @@ -70,7 +70,8 @@ func DetectTerminalProgram(getEnv func(string) string) string { // It accepts an environment getter function to allow testing without modifying the actual environment. // // Agent-specific variables are checked first. When none match it falls back to the two host -// variables DetectAgentHost reads, which identify an agent surface and so imply the agent. +// variables DetectAgentHost reads, which identify an agent surface and so imply the agent, and +// finally to AI_AGENT and then AGENT, which some agents report through directly. func DetectAIAgent(getEnv func(string) string) string { if getEnv("ANTIGRAVITY_CLI_ALIAS") != "" { return "antigravity" @@ -90,6 +91,12 @@ func DetectAIAgent(getEnv func(string) string) string { if getEnv("GEMINI_CLI") != "" { return "gemini_cli" } + if getEnv("HERMES_AGENT") != "" { + return "hermes" + } + if getEnv("GROK_AGENT") != "" || getEnv("GROK_SESSION_ID") != "" { + return "grok" + } if getEnv("OPENCODE") != "" { return "open_code" } @@ -114,6 +121,17 @@ func DetectAIAgent(getEnv func(string) string) string { return "codex_cli" } + // Last resort: AI_AGENT is a convention some agents report themselves through + // (see DetectAgentVersion), so a value here is itself evidence of an agent even + // when none of the specific variables above matched. + if aiAgent := strings.TrimSpace(getEnv("AI_AGENT")); aiAgent != "" { + return aiAgent + } + // AGENT is the same convention under the name Goose, Amp and Bun use. + if agent := strings.TrimSpace(getEnv("AGENT")); agent != "" { + return agent + } + return "" } @@ -130,9 +148,14 @@ func DetectAIAgent(getEnv func(string) string) string { // from a vendor's source, which matters because mapping one otherwise costs a code // change, a release, and users upgrading before it is even visible. // -// One host is inferred rather than reported: Codex names every surface except the terminal, -// so a detected Codex agent with no host is a terminal one. See the empty-host branch below. +// Two hosts are inferred rather than reported: Codex names every surface except the +// terminal, and Grok Build names neither of its two surfaces, so a detected agent of +// either with no host is inferred from other signals. See the empty-host branch below. func DetectAgentHost(getEnv func(string) string) (kind string, raw string) { + if getEnv("HERMES_DESKTOP") != "" { + return "desktop", "hermes" + } + host := getEnv("CLAUDE_CODE_ENTRYPOINT") if host == "" { // Codex Desktop sets this alongside the generic Codex signals, and it is @@ -157,6 +180,17 @@ func DetectAgentHost(getEnv func(string) string) (kind string, raw string) { return "terminal", "codex-cli" } + // Grok Build also has an ACP surface for IDE integration (`grok agent + // stdio`), and neither surface sets a dedicated host variable the way + // Claude/Codex do. The only signal separating them is that GROK_AGENT is set for + // the terminal TUI and absent (with only GROK_SESSION_ID present) under ACP. + if DetectAIAgent(getEnv) == "grok" { + if getEnv("GROK_AGENT") != "" { + return "terminal", "grok-cli" + } + return "ide", "grok-acp" + } + return "", "" } diff --git a/pkg/useragent/useragent_test.go b/pkg/useragent/useragent_test.go index aa6aab045..1d0396368 100644 --- a/pkg/useragent/useragent_test.go +++ b/pkg/useragent/useragent_test.go @@ -78,6 +78,8 @@ func TestDetectAgentHost(t *testing.T) { kind string raw string }{ + {"hermes desktop", map[string]string{"HERMES_DESKTOP": "true"}, "desktop", "hermes"}, + {"hermes desktop, any non-empty value counts", map[string]string{"HERMES_DESKTOP": "1"}, "desktop", "hermes"}, {"claude desktop", map[string]string{"CLAUDE_CODE_ENTRYPOINT": "claude-desktop"}, "desktop", "claude-desktop"}, // Both are desktop, and raw is the only thing that tells them apart. {"claude desktop 3p", map[string]string{"CLAUDE_CODE_ENTRYPOINT": "claude-desktop-3p"}, "desktop", "claude-desktop-3p"}, @@ -112,6 +114,10 @@ func TestDetectAgentHost(t *testing.T) { // originator is a terminal one. This is the only inferred host. {"codex terminal inferred from sandbox signal", map[string]string{"CODEX_SANDBOX": "1"}, "terminal", "codex-cli"}, {"codex terminal inferred from thread signal", map[string]string{"CODEX_THREAD_ID": "thread-abc"}, "terminal", "codex-cli"}, + // Grok Build sets GROK_AGENT for its terminal TUI but not for its ACP/IDE + // surface, which sets only GROK_SESSION_ID + {"grok terminal inferred from agent signal", map[string]string{"GROK_AGENT": "1"}, "terminal", "grok-cli"}, + {"grok acp inferred from session signal without agent", map[string]string{"GROK_SESSION_ID": "01a0b015-bf40-7673-a055-afee8019dc33"}, "ide", "grok-acp"}, // The inference is gated on the agent, so it does not fire for anyone else. Claude // Code without an entrypoint has no host, rather than a guessed terminal. {"claude code without entrypoint stays hostless", map[string]string{"CLAUDECODE": "1"}, "", ""}, @@ -179,6 +185,49 @@ func TestDetectAIAgent_InferredFromHost(t *testing.T) { } } +func TestDetectAIAgent_Hermes(t *testing.T) { + require.Equal(t, "hermes", DetectAIAgent(mapEnv(map[string]string{"HERMES_AGENT": "1"}))) +} + +// TestDetectAIAgent_AIAgentFallback covers the last-resort fallback: when no agent-specific +// variable or inherited host names the agent, AI_AGENT and then AGENT are reported directly. +func TestDetectAIAgent_AIAgentFallback(t *testing.T) { + tests := []struct { + name string + envs map[string]string + expected string + description string + }{ + {"reported when nothing else matches", map[string]string{"AI_AGENT": "goose_1-2-3"}, "goose_1-2-3", ""}, + {"whitespace trimmed", map[string]string{"AI_AGENT": " goose "}, "goose", ""}, + { + name: "specific agent variable wins over AI_AGENT", + envs: map[string]string{"CURSOR_AGENT": "1", "AI_AGENT": "goose"}, + expected: "cursor", + description: "AI_AGENT is checked last, so a direct signal still takes priority", + }, + {"blank AI_AGENT reports nothing", map[string]string{"AI_AGENT": " "}, "", ""}, + { + name: "AGENT used when AI_AGENT is absent", + envs: map[string]string{"AGENT": "amp"}, + expected: "amp", + description: "AGENT is the same convention under the name Goose, Amp and Bun use", + }, + { + name: "AI_AGENT wins over AGENT", + envs: map[string]string{"AI_AGENT": "goose", "AGENT": "amp"}, + expected: "goose", + }, + {"blank AGENT reports nothing", map[string]string{"AGENT": " "}, "", ""}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.expected, DetectAIAgent(mapEnv(tt.envs)), tt.description) + }) + } +} + func TestDetectAgentVersion(t *testing.T) { tests := []struct { name string @@ -303,6 +352,19 @@ func TestObservedAgentSessions(t *testing.T) { description: "Codex sets no originator from a terminal, so the terminal host is " + "inferred from its absence; every other Codex surface names itself", }, + { + name: "grok build", + envs: map[string]string{ + "GROK_AGENT": "1", + "GROK_SESSION_ID": sensitiveSessionID, + }, + agent: "grok", + hostKind: "terminal", + hostRaw: "grok-cli", + version: "", + description: "Grok Build's terminal TUI sets GROK_AGENT alongside GROK_SESSION_ID, and " + + "reports no version through the AI_AGENT/AGENT convention", + }, } for _, tt := range tests { @@ -342,6 +404,8 @@ func TestObservedAgentSessions_NoSensitiveValuesReported(t *testing.T) { "CODEX_THREAD_ID": sensitiveThreadID, "CODEX_INTERNAL_ORIGINATOR_OVERRIDE": "Codex Desktop", "CODEX_PERMISSION_PROFILE": ":read-only", + "GROK_AGENT": "1", + "GROK_SESSION_ID": sensitiveSessionID, } getEnv := mapEnv(envs) diff --git a/scripts/install.sh b/scripts/install.sh index c1118f132..8a8765da1 100755 --- a/scripts/install.sh +++ b/scripts/install.sh @@ -4,6 +4,8 @@ set -eu INSTALL_DIR="${STRIPE_INSTALL_DIR:-$HOME/.stripe/bin}" GITHUB_REPO="stripe/stripe-cli" NEEDS_SOURCE=false +TELEMETRY_URL="${STRIPE_TELEMETRY_URL:-https://r.stripe.com/0}" +INSTALL_SUCCESS=false main() { detect_platform @@ -11,6 +13,8 @@ main() { download_and_verify install_binary setup_path + INSTALL_SUCCESS=true + send_telemetry "Install Succeeded" "version=$VERSION" print_success } @@ -103,7 +107,7 @@ download_and_verify() { BASE_URL="https://github.com/${GITHUB_REPO}/releases/download/v${VERSION}" TMP_DIR=$(mktemp -d) - trap 'rm -rf "$TMP_DIR"' EXIT + trap 'rm -rf "$TMP_DIR"; if [ "$INSTALL_SUCCESS" = "false" ]; then send_telemetry "Install Failed" "version=${VERSION:-unknown}"; fi' EXIT echo "Downloading stripe v${VERSION}..." http_download "$BASE_URL/$ARCHIVE" "$TMP_DIR/$ARCHIVE" @@ -199,6 +203,25 @@ version_lt() { [ "$1" != "$2" ] && [ "$(printf '%s\n%s' "$1" "$2" | sort -V | head -n1)" = "$1" ] } +send_telemetry() { + event_name="$1" + event_value="$2" + + case "${STRIPE_CLI_TELEMETRY_OPTOUT:-}${DO_NOT_TRACK:-}" in + *1*|*true*|*TRUE*) return ;; + esac + + # install_method matches what pkg/installmethod reports for a script install, + # so events from the installer and from the CLI group together. + telemetry_data="client_id=stripe-cli&event_name=${event_name}&event_value=${event_value}&os=${OS:-unknown}&arch=${ARCH_LABEL:-unknown}&cli_version=${VERSION:-unknown}&install_method=script" + + if command -v curl >/dev/null 2>&1; then + curl -sS --max-time 3 -X POST -H "origin: stripe-cli" -H "Content-Type: application/x-www-form-urlencoded" -d "$telemetry_data" "$TELEMETRY_URL" >/dev/null 2>&1 || true + elif command -v wget >/dev/null 2>&1; then + wget -q --timeout=3 -O /dev/null --post-data="$telemetry_data" --header="origin: stripe-cli" --header="Content-Type: application/x-www-form-urlencoded" "$TELEMETRY_URL" 2>/dev/null || true + fi +} + print_success() { echo "" echo "stripe v${VERSION} installed to $INSTALL_DIR/stripe"