diff --git a/cmd/publisher/auth/github-at.go b/cmd/publisher/auth/github-at.go index 0da68a4c6..fef3774da 100644 --- a/cmd/publisher/auth/github-at.go +++ b/cmd/publisher/auth/github-at.go @@ -26,6 +26,20 @@ const ( // maximum to the implementation; this prevents a misbehaving auth server // from growing the interval unboundedly. maxPollInterval = 60 + // defaultExpiresIn is the device-code lifetime in seconds used when the + // device-code response omits expires_in. + defaultExpiresIn = 900 + // maxExpiresIn caps the device-code lifetime a response may claim. It sits + // well past GitHub's documented 15 minutes; the cap stops a misbehaving or + // hostile response from pushing the deadline (and the polling loop) far + // into the future. + maxExpiresIn = 3600 + // incorrectDeviceCodeGraceRetries is the total number of extra polls a + // single login may spend on incorrect_device_code before giving up — a + // budget for the whole login, not per occurrence. cli/cli#9302 reports the + // same error for a code that had just been issued; a genuinely invalid code + // still fails, only a few seconds later. + incorrectDeviceCodeGraceRetries = 2 ) // DeviceCodeResponse represents the response from GitHub's device code endpoint @@ -43,6 +57,38 @@ type AccessTokenResponse struct { TokenType string `json:"token_type"` Scope string `json:"scope"` Error string `json:"error,omitempty"` + // ErrorDescription and ErrorURI accompany Error in GitHub's responses; + // surfacing them turns opaque failures like "incorrect_device_code" into + // diagnosable reports. + ErrorDescription string `json:"error_description,omitempty"` + ErrorURI string `json:"error_uri,omitempty"` +} + +// errorDetail renders Error together with any error_description and error_uri +// GitHub returned, so failures reach the user diagnosable rather than opaque. +func (r AccessTokenResponse) errorDetail() string { + detail := r.Error + if r.ErrorDescription != "" { + detail += ": " + r.ErrorDescription + } + if r.ErrorURI != "" { + detail += " (" + r.ErrorURI + ")" + } + return detail +} + +func decodeAccessTokenResponse(statusCode int, body []byte) (AccessTokenResponse, error) { + var tokenResp AccessTokenResponse + // OAuth errors may use HTTP 400 (RFC 6749 §5.2), not just HTTP 200. + if statusCode != http.StatusOK && statusCode != http.StatusBadRequest { + return tokenResp, fmt.Errorf("token endpoint returned %d: %s", statusCode, body) + } + + err := json.Unmarshal(body, &tokenResp) + if statusCode == http.StatusBadRequest && (err != nil || tokenResp.Error == "") { + return tokenResp, fmt.Errorf("token endpoint returned %d: %s", statusCode, body) + } + return tokenResp, err } // RegistryTokenResponse represents the response from registry's token exchange endpoint @@ -61,9 +107,17 @@ type GitHubATProvider struct { // accessTokenURL is the GitHub access-token polling endpoint. It is a field // (rather than the package constant) so tests can point it at a mock server. accessTokenURL string + // deviceCodeURL is the GitHub device-code endpoint, a field for the same + // reason as accessTokenURL. + deviceCodeURL string // pollInterval is the initial polling interval in seconds. Defaults to - // defaultPollInterval; overridable in tests to avoid real delays. + // defaultPollInterval; overridable in tests to avoid real delays. Updated + // from the device-code response's interval when GitHub returns one. pollInterval int + // expiresIn is the device-code lifetime in seconds. Defaults to + // defaultExpiresIn; updated from the device-code response's expires_in + // when GitHub returns one. + expiresIn int // sleep abstracts time.Sleep so tests can run without real delays and // assert the back-off sequence. Defaults to time.Sleep. sleep func(time.Duration) @@ -86,7 +140,9 @@ func NewGitHubATProvider(registryURL, token string) Provider { registryURL: registryURL, providedToken: token, accessTokenURL: GitHubAccessTokenURL, + deviceCodeURL: GitHubDeviceCodeURL, pollInterval: defaultPollInterval, + expiresIn: defaultExpiresIn, sleep: time.Sleep, } } @@ -173,7 +229,7 @@ func (g *GitHubATProvider) requestDeviceCode(ctx context.Context) (string, strin return "", "", "", err } - req, err := http.NewRequestWithContext(ctx, http.MethodPost, GitHubDeviceCodeURL, bytes.NewBuffer(jsonData)) + req, err := http.NewRequestWithContext(ctx, http.MethodPost, g.deviceCodeURL, bytes.NewBuffer(jsonData)) if err != nil { return "", "", "", err } @@ -202,6 +258,21 @@ func (g *GitHubATProvider) requestDeviceCode(ctx context.Context) (string, strin return "", "", "", err } + // Per RFC 8628 §3.2 interval and expires_in are bound to this device code, + // so reset any pacing left over from a previous code before adopting them. + // Clamp both: a hostile or buggy response must not grow the interval past + // maxPollInterval (the invariant maxPollInterval exists for) nor push the + // deadline unboundedly into the future, and clamping also keeps the + // derived time.Duration well clear of int64 overflow. + g.pollInterval = defaultPollInterval + g.expiresIn = defaultExpiresIn + if deviceCodeResp.Interval > 0 { + g.pollInterval = min(deviceCodeResp.Interval, maxPollInterval) + } + if deviceCodeResp.ExpiresIn > 0 { + g.expiresIn = min(deviceCodeResp.ExpiresIn, maxExpiresIn) + } + return deviceCodeResp.DeviceCode, deviceCodeResp.UserCode, deviceCodeResp.VerificationURI, nil } @@ -222,11 +293,26 @@ func (g *GitHubATProvider) pollForToken(ctx context.Context, deviceCode string) return "", err } - // Default polling interval and expiration time + // Pacing comes from the device-code response (via requestDeviceCode), + // falling back to GitHub's documented defaults. interval := g.pollInterval - expiresIn := 900 // 15 minutes + expiresIn := g.expiresIn + // Zero-value fallback: providers built directly rather than through + // requestDeviceCode (the test seam) carry no lifetime. Keep this guard. + if expiresIn <= 0 { + expiresIn = defaultExpiresIn + } deadline := time.Now().Add(time.Duration(expiresIn) * time.Second) + graceRetries := 0 + // serverErrorInterval grows the wait between consecutive 5xx retries and is + // reset by any other response. + serverErrorInterval := interval + // lastErr holds the most recent retryable failure so a deadline reached + // mid-retry reports that diagnostic instead of a bare "timed out". It is + // cleared whenever polling recovers, so a genuine user timeout after a + // recovered blip still reports plainly. + var lastErr error for time.Now().Before(deadline) { req, err := http.NewRequestWithContext(ctx, http.MethodPost, g.accessTokenURL, bytes.NewBuffer(jsonData)) if err != nil { @@ -247,8 +333,20 @@ func (g *GitHubATProvider) pollForToken(ctx context.Context, deviceCode string) return "", err } - var tokenResp AccessTokenResponse - err = json.Unmarshal(body, &tokenResp) + // A transient 5xx must not abandon a login the user may already have + // approved in the browser. Each consecutive 5xx grows the wait by 5s up + // to maxPollInterval, so a sustained outage costs a couple of dozen + // requests rather than hundreds; the deadline bounds the number of + // these retries. + if resp.StatusCode >= http.StatusInternalServerError { + lastErr = fmt.Errorf("token endpoint returned %d", resp.StatusCode) + serverErrorInterval = min(serverErrorInterval+5, maxPollInterval) + g.sleep(time.Duration(serverErrorInterval) * time.Second) + continue + } + serverErrorInterval = interval + + tokenResp, err := decodeAccessTokenResponse(resp.StatusCode, body) if err != nil { return "", err } @@ -263,12 +361,26 @@ func (g *GitHubATProvider) pollForToken(ctx context.Context, deviceCode string) interval = maxPollInterval } } + lastErr = nil g.sleep(time.Duration(interval) * time.Second) continue } if tokenResp.Error != "" { - return "", fmt.Errorf("token request failed: %s", tokenResp.Error) + failure := fmt.Errorf("token request failed: %s", tokenResp.errorDetail()) + + // incorrect_device_code has been reported for codes that were in + // fact valid (cli/cli#9302; cause never confirmed), so spend a + // couple of grace polls before abandoning a login the user may + // have approved. + if tokenResp.Error == "incorrect_device_code" && graceRetries < incorrectDeviceCodeGraceRetries { + graceRetries++ + lastErr = failure + g.sleep(time.Duration(interval) * time.Second) + continue + } + + return "", failure } if tokenResp.AccessToken != "" { @@ -279,6 +391,9 @@ func (g *GitHubATProvider) pollForToken(ctx context.Context, deviceCode string) return "", fmt.Errorf("failed to obtain access token") } + if lastErr != nil { + return "", fmt.Errorf("device code authorization timed out; last response: %w", lastErr) + } return "", fmt.Errorf("device code authorization timed out") } diff --git a/cmd/publisher/auth/github_at_device_flow_internal_test.go b/cmd/publisher/auth/github_at_device_flow_internal_test.go new file mode 100644 index 000000000..b1e13accc --- /dev/null +++ b/cmd/publisher/auth/github_at_device_flow_internal_test.go @@ -0,0 +1,532 @@ +package auth + +import ( + "context" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// seqTokenServer serves a fixed sequence of (status, body) responses, one per +// request, reusing the last one once the sequence is exhausted. Unlike +// newMockTokenServer it can serve non-200 statuses and non-JSON bodies, which +// the transient-5xx path needs. +type seqTokenServer struct { + srv *httptest.Server + mu sync.Mutex + seen int + items []seqResp +} + +type seqResp struct { + status int + body string +} + +func newSeqTokenServer(t *testing.T, items ...seqResp) *seqTokenServer { + t.Helper() + s := &seqTokenServer{items: items} + s.srv = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + s.mu.Lock() + item := s.items[len(s.items)-1] + if s.seen < len(s.items) { + item = s.items[s.seen] + } + s.seen++ + s.mu.Unlock() + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(item.status) + _, _ = w.Write([]byte(item.body)) + })) + t.Cleanup(s.srv.Close) + return s +} + +func (s *seqTokenServer) requests() int { + s.mu.Lock() + defer s.mu.Unlock() + return s.seen +} + +func okResp(body string) seqResp { return seqResp{status: http.StatusOK, body: body} } + +// serverErrorResp is a 502 carrying an HTML body rather than JSON, the shape a +// transient edge failure takes. +func serverErrorResp() seqResp { + return seqResp{status: http.StatusBadGateway, body: `Server Error`} +} + +const ( + testClientID = "test-client-id" + + pendingBody = `{"error":"authorization_pending"}` + incorrectBody = `{"error":"incorrect_device_code"}` + // The full shape GitHub sends for this error, description and URI included. + incorrectVerboseBody = `{"error":"incorrect_device_code","error_description":"The device_code provided is not valid.","error_uri":"https://docs.github.com/developers/apps/authorizing-oauth-apps#error-codes-for-the-device-flow"}` + tokenBody = `{"access_token":"gho_test_token","token_type":"bearer","scope":"read:org,read:user"}` // #nosec G101 -- test fixture, not a real secret +) + +// TestPollForToken_IncorrectDeviceCodeIsTerminalAfterGrace confirms a code that +// stays invalid past the grace polls still fails, and fails with the error +// GitHub reported rather than a rewritten one. +func TestPollForToken_IncorrectDeviceCodeIsTerminalAfterGrace(t *testing.T) { + gh := newSeqTokenServer(t, okResp(pendingBody), okResp(pendingBody), okResp(incorrectBody)) + p, _ := newPollTestProvider(gh.srv.URL) + + _, err := p.pollForToken(context.Background(), "issued-device-code") + require.Error(t, err) + assert.Equal(t, "token request failed: incorrect_device_code", err.Error()) + // 2 pending + 1 incorrect + 2 grace polls that also answer incorrect. + assert.Equal(t, 5, gh.requests()) +} + +// TestPollForToken_TransientIncorrectDeviceCodeRecovered covers the case the +// grace retries exist for: one incorrect_device_code between a pending poll and +// a successful one must not discard a token the user already authorised. +func TestPollForToken_TransientIncorrectDeviceCodeRecovered(t *testing.T) { + gh := newSeqTokenServer(t, okResp(pendingBody), okResp(incorrectBody), okResp(tokenBody)) + p, _ := newPollTestProvider(gh.srv.URL) + + token, err := p.pollForToken(context.Background(), "issued-device-code") + require.NoError(t, err, "one transient incorrect_device_code must not be fatal") + assert.Equal(t, "gho_test_token", token) + assert.Equal(t, 3, gh.requests()) +} + +func TestPollForToken_OAuthErrorsRetryOnBadRequest(t *testing.T) { + for _, status := range []int{http.StatusOK, http.StatusBadRequest} { + for _, code := range []string{"authorization_pending", "slow_down", "incorrect_device_code"} { + t.Run(fmt.Sprintf("%d/%s", status, code), func(t *testing.T) { + gh := newSeqTokenServer(t, + seqResp{status: status, body: fmt.Sprintf(`{"error":%q}`, code)}, + okResp(tokenBody), + ) + p, slept := newPollTestProvider(gh.srv.URL) + p.pollInterval = defaultPollInterval + + token, err := p.pollForToken(context.Background(), "issued-device-code") + require.NoError(t, err) + assert.Equal(t, "gho_test_token", token) + assert.Equal(t, 2, gh.requests()) + wait := time.Duration(defaultPollInterval) * time.Second + if code == "slow_down" { + wait += 5 * time.Second + } + assert.Equal(t, []time.Duration{wait}, *slept) + }) + } + } +} + +func TestPollForToken_BadRequestOAuthErrorsRemainTerminal(t *testing.T) { + for _, code := range []string{"access_denied", "expired_token", "unknown_error", "incorrect_device_code"} { + t.Run(code, func(t *testing.T) { + gh := newSeqTokenServer(t, seqResp{ + status: http.StatusBadRequest, + body: fmt.Sprintf(`{"error":%q,"error_description":"Request rejected","error_uri":"https://github.com/error"}`, code), + }) + p, slept := newPollTestProvider(gh.srv.URL) + + token, err := p.pollForToken(context.Background(), "issued-device-code") + require.EqualError(t, err, "token request failed: "+code+": Request rejected (https://github.com/error)") + assert.Empty(t, token) + retries := 0 + if code == "incorrect_device_code" { + retries = incorrectDeviceCodeGraceRetries + } + assert.Equal(t, 1+retries, gh.requests()) + assert.Len(t, *slept, retries) + }) + } +} + +// TestPollForToken_ErrorDescriptionAndURISurfaced confirms the grace budget is +// spent in full against a persistently invalid code, and that GitHub's +// error_description and error_uri reach the user instead of being discarded. +func TestPollForToken_ErrorDescriptionAndURISurfaced(t *testing.T) { + gh := newSeqTokenServer(t, okResp(incorrectVerboseBody)) + p, _ := newPollTestProvider(gh.srv.URL) + + _, err := p.pollForToken(context.Background(), "issued-device-code") + require.Error(t, err) + assert.Contains(t, err.Error(), "incorrect_device_code") + assert.Contains(t, err.Error(), "The device_code provided is not valid.") + assert.Contains(t, err.Error(), "error-codes-for-the-device-flow") + assert.Equal(t, 1+incorrectDeviceCodeGraceRetries, gh.requests()) +} + +// TestPollForToken_TransientServerErrorRetried confirms a 5xx (which carries an +// HTML body, not JSON) is retried rather than ending the login. +func TestPollForToken_TransientServerErrorRetried(t *testing.T) { + gh := newSeqTokenServer(t, + serverErrorResp(), + okResp(tokenBody), + ) + p, slept := newPollTestProvider(gh.srv.URL) + + token, err := p.pollForToken(context.Background(), "issued-device-code") + require.NoError(t, err) + assert.Equal(t, "gho_test_token", token) + assert.Equal(t, 2, gh.requests()) + assert.Len(t, *slept, 1, "the 5xx must be followed by one back-off") +} + +// TestPollForToken_FirstPollIsImmediate pins the polling schedule: the first +// request goes out without a preceding sleep, so a user who authorises quickly +// is not made to wait an interval for nothing. +func TestPollForToken_FirstPollIsImmediate(t *testing.T) { + gh := newSeqTokenServer(t, okResp(tokenBody)) + p, slept := newPollTestProvider(gh.srv.URL) + + _, err := p.pollForToken(context.Background(), "issued-device-code") + require.NoError(t, err) + assert.Equal(t, 1, gh.requests()) + assert.Empty(t, *slept, "an immediately successful poll must not sleep") +} + +// TestRequestDeviceCode_HonoursServerIntervalAndExpiry confirms the interval and +// expires_in in the device-code response are adopted rather than dropped. +func TestRequestDeviceCode_HonoursServerIntervalAndExpiry(t *testing.T) { + device := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"device_code":"dc-1","user_code":"AAAA-BBBB","verification_uri":"https://github.com/login/device","expires_in":60,"interval":7}`)) + })) + t.Cleanup(device.Close) + + p := &GitHubATProvider{ + clientID: testClientID, + deviceCodeURL: device.URL, + pollInterval: defaultPollInterval, + expiresIn: defaultExpiresIn, + } + + deviceCode, userCode, uri, err := p.requestDeviceCode(context.Background()) + require.NoError(t, err) + assert.Equal(t, "dc-1", deviceCode) + assert.Equal(t, "AAAA-BBBB", userCode) + assert.Equal(t, "https://github.com/login/device", uri) + assert.Equal(t, 7, p.pollInterval, "interval from the response must be adopted") + assert.Equal(t, 60, p.expiresIn, "expires_in from the response must be adopted") +} + +// TestRequestDeviceCode_KeepsDefaultsWhenOmitted confirms a response without +// interval or expires_in leaves the documented defaults in place. +func TestRequestDeviceCode_KeepsDefaultsWhenOmitted(t *testing.T) { + device := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"device_code":"dc-2","user_code":"EEEE-FFFF","verification_uri":"https://github.com/login/device"}`)) + })) + t.Cleanup(device.Close) + + p := &GitHubATProvider{ + clientID: testClientID, + deviceCodeURL: device.URL, + pollInterval: defaultPollInterval, + expiresIn: defaultExpiresIn, + } + + _, _, _, err := p.requestDeviceCode(context.Background()) + require.NoError(t, err) + assert.Equal(t, defaultPollInterval, p.pollInterval) + assert.Equal(t, defaultExpiresIn, p.expiresIn) +} + +// deviceCodeServerReturning serves a single device-code response with the given +// raw JSON body, for exercising interval/expires_in handling. +func deviceCodeServerReturning(t *testing.T, body string) string { + t.Helper() + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(body)) + })) + t.Cleanup(srv.Close) + return srv.URL +} + +// TestRequestDeviceCode_ClampsExcessiveInterval confirms an interval far above +// maxPollInterval is clamped, protecting the same invariant maxPollInterval +// enforces on the slow_down path. +func TestRequestDeviceCode_ClampsExcessiveInterval(t *testing.T) { + url := deviceCodeServerReturning(t, + `{"device_code":"dc","user_code":"AAAA-BBBB","verification_uri":"https://github.com/login/device","expires_in":900,"interval":86400}`) + + p := &GitHubATProvider{clientID: testClientID, deviceCodeURL: url, pollInterval: defaultPollInterval, expiresIn: defaultExpiresIn} + _, _, _, err := p.requestDeviceCode(context.Background()) + require.NoError(t, err) + assert.Equal(t, maxPollInterval, p.pollInterval, "interval above the cap must be clamped") +} + +// TestRequestDeviceCode_ClampsOverflowValues covers the overflow pair: values +// large enough that time.Duration(n)*time.Second wraps negative would make +// sleep return instantly (spinning the loop) and put the deadline in the past +// (yielding zero polls). Clamping keeps both within safe bounds. +func TestRequestDeviceCode_ClampsOverflowValues(t *testing.T) { + url := deviceCodeServerReturning(t, + `{"device_code":"dc","user_code":"AAAA-BBBB","verification_uri":"https://github.com/login/device","expires_in":10000000000,"interval":10000000000}`) + + p := &GitHubATProvider{clientID: testClientID, deviceCodeURL: url, pollInterval: defaultPollInterval, expiresIn: defaultExpiresIn} + _, _, _, err := p.requestDeviceCode(context.Background()) + require.NoError(t, err) + + assert.Equal(t, maxPollInterval, p.pollInterval) + assert.Equal(t, maxExpiresIn, p.expiresIn) +} + +// TestRequestDeviceCode_IgnoresNegativeValues confirms negative interval and +// expires_in are rejected in favour of the defaults, so neither can produce a +// negative sleep or a deadline already in the past. +func TestRequestDeviceCode_IgnoresNegativeValues(t *testing.T) { + url := deviceCodeServerReturning(t, + `{"device_code":"dc","user_code":"AAAA-BBBB","verification_uri":"https://github.com/login/device","expires_in":-1,"interval":-1}`) + + p := &GitHubATProvider{clientID: testClientID, deviceCodeURL: url, pollInterval: defaultPollInterval, expiresIn: defaultExpiresIn} + _, _, _, err := p.requestDeviceCode(context.Background()) + require.NoError(t, err) + + assert.Equal(t, defaultPollInterval, p.pollInterval, "negative interval must fall back to the default") + assert.Equal(t, defaultExpiresIn, p.expiresIn, "negative expires_in must fall back to the default") +} + +// TestPollForToken_TimeoutSurfacesLastError covers F2: when the deadline lands +// mid-retry, the timeout must carry the last retryable diagnostic instead of a +// bare "timed out". The grace poll records the full-shape incorrect_device_code +// and its following sleep crosses the (1-second) deadline. +func TestPollForToken_TimeoutSurfacesLastError(t *testing.T) { + gh := newSeqTokenServer(t, okResp(incorrectVerboseBody)) + p := &GitHubATProvider{ + clientID: testClientID, + accessTokenURL: gh.srv.URL, + pollInterval: 1, + expiresIn: 1, // deadline ~1s out + sleep: func(time.Duration) { time.Sleep(1100 * time.Millisecond) }, + } + + _, err := p.pollForToken(context.Background(), "issued-device-code") + require.Error(t, err) + assert.Contains(t, err.Error(), "timed out") + assert.Contains(t, err.Error(), "The device_code provided is not valid.", + "the last retryable diagnostic must survive the timeout") +} + +// TestPollForToken_TimeoutAfterRecoveryIsBare confirms lastErr is cleared when +// polling recovers: one transient 502 followed by healthy pending polls and +// then a genuine user timeout must report plainly, not blame the stale 502. +func TestPollForToken_TimeoutAfterRecoveryIsBare(t *testing.T) { + gh := newSeqTokenServer(t, + serverErrorResp(), + okResp(pendingBody), + ) + p := &GitHubATProvider{ + clientID: testClientID, + accessTokenURL: gh.srv.URL, + pollInterval: 1, + expiresIn: 1, // deadline ~1s out + sleep: func(time.Duration) { time.Sleep(400 * time.Millisecond) }, + } + + _, err := p.pollForToken(context.Background(), "issued-device-code") + require.Error(t, err) + assert.Equal(t, "device code authorization timed out", err.Error(), + "a recovered 502 must not be reported as the cause of a user timeout") +} + +// Non-OAuth failures retain HTTP diagnostics rather than JSON parsing errors. +func TestPollForToken_NonOKStatusSurfacesStatusAndBody(t *testing.T) { + for _, tt := range []struct { + name string + resp seqResp + }{ + {"rate_limited", seqResp{status: http.StatusTooManyRequests, body: "rate limited"}}, + {"bad_request_html", seqResp{status: http.StatusBadRequest, body: "bad request"}}, + {"bad_request_invalid_json", seqResp{status: http.StatusBadRequest, body: `{"error":`}}, + {"bad_request_empty_error", seqResp{status: http.StatusBadRequest, body: `{"error":""}`}}, + {"bad_request_token", seqResp{status: http.StatusBadRequest, body: tokenBody}}, + {"pending_on_other_status", seqResp{status: http.StatusForbidden, body: pendingBody}}, + } { + t.Run(tt.name, func(t *testing.T) { + gh := newSeqTokenServer(t, tt.resp) + p, slept := newPollTestProvider(gh.srv.URL) + + token, err := p.pollForToken(context.Background(), "issued-device-code") + require.EqualError(t, err, fmt.Sprintf("token endpoint returned %d: %s", tt.resp.status, tt.resp.body)) + assert.Empty(t, token) + assert.Equal(t, 1, gh.requests(), "a non-retryable status must stop polling") + assert.Empty(t, *slept) + }) + } +} + +// TestPollForToken_ConsecutiveServerErrorsBackOff confirms the 5xx wait grows by +// 5s per consecutive failure and is capped, so a sustained outage costs a +// bounded number of requests rather than hammering the endpoint. +func TestPollForToken_ConsecutiveServerErrorsBackOff(t *testing.T) { + responses := make([]seqResp, 0, 4) + for range [3]struct{}{} { + responses = append(responses, serverErrorResp()) + } + responses = append(responses, okResp(tokenBody)) + + gh := newSeqTokenServer(t, responses...) + p, slept := newPollTestProvider(gh.srv.URL) + + token, err := p.pollForToken(context.Background(), "issued-device-code") + require.NoError(t, err) + assert.Equal(t, "gho_test_token", token) + + // Base interval is 0 in the test provider, so the waits are 5, 10, 15. + require.Len(t, *slept, 3) + assert.Equal(t, 5*time.Second, (*slept)[0]) + assert.Equal(t, 10*time.Second, (*slept)[1]) + assert.Equal(t, 15*time.Second, (*slept)[2]) +} + +// TestPollForToken_ServerErrorBackOffResetsOnRecovery confirms the 5xx growth is +// reset by any non-5xx response, so an intermittent endpoint does not +// accumulate an ever-longer wait across separate blips. +func TestPollForToken_ServerErrorBackOffResetsOnRecovery(t *testing.T) { + serverError := serverErrorResp() + gh := newSeqTokenServer(t, + serverError, + okResp(pendingBody), + serverError, + okResp(tokenBody), + ) + p, slept := newPollTestProvider(gh.srv.URL) + + token, err := p.pollForToken(context.Background(), "issued-device-code") + require.NoError(t, err) + assert.Equal(t, "gho_test_token", token) + + // 5s for the first 5xx, 0s for the pending poll (base interval), then 5s + // again — not 10s — because the pending response reset the growth. + require.Len(t, *slept, 3) + assert.Equal(t, 5*time.Second, (*slept)[0]) + assert.Equal(t, 0*time.Second, (*slept)[1]) + assert.Equal(t, 5*time.Second, (*slept)[2]) +} + +// TestRequestDeviceCode_PacingDoesNotLeakAcrossCodes covers F4: interval and +// expires_in are bound to the code they were issued with, so a second request +// that omits them must fall back to defaults rather than inherit the first. +func TestRequestDeviceCode_PacingDoesNotLeakAcrossCodes(t *testing.T) { + var call int + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + if call == 0 { + _, _ = w.Write([]byte(`{"device_code":"dc1","user_code":"AAAA-BBBB","verification_uri":"https://github.com/login/device","expires_in":120,"interval":55}`)) + } else { + _, _ = w.Write([]byte(`{"device_code":"dc2","user_code":"CCCC-DDDD","verification_uri":"https://github.com/login/device"}`)) + } + call++ + })) + t.Cleanup(srv.Close) + + p := &GitHubATProvider{clientID: testClientID, deviceCodeURL: srv.URL, pollInterval: defaultPollInterval, expiresIn: defaultExpiresIn} + + _, _, _, err := p.requestDeviceCode(context.Background()) + require.NoError(t, err) + assert.Equal(t, 55, p.pollInterval) + assert.Equal(t, 120, p.expiresIn) + + _, _, _, err = p.requestDeviceCode(context.Background()) + require.NoError(t, err) + assert.Equal(t, defaultPollInterval, p.pollInterval, "second code must not inherit the first's interval") + assert.Equal(t, defaultExpiresIn, p.expiresIn, "second code must not inherit the first's expires_in") +} + +// TestLogin_ClientIDStableAcrossDeviceCodeAndPolls covers the whole device flow +// end to end and pins the one invariant that makes incorrect_device_code +// meaningful: the health endpoint is read once, and the client_id the device +// code was issued under is the one every poll carries. The mock token endpoint +// answers incorrect_device_code for any other pairing. +func TestLogin_ClientIDStableAcrossDeviceCodeAndPolls(t *testing.T) { + var mu sync.Mutex + healthCalls := 0 + deviceClientID := "" + pollClientIDs := []string{} + pollCount := 0 + const issued = "dc-login" + + registry := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/v0/health" { + http.NotFound(w, r) + return + } + mu.Lock() + healthCalls++ + id := fmt.Sprintf("Iv23liCLIENT%04d", healthCalls) + mu.Unlock() + w.Header().Set("Content-Type", "application/json") + _, _ = fmt.Fprintf(w, `{"status":"ok","github_client_id":"%s"}`, id) + })) + t.Cleanup(registry.Close) + + device := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var req map[string]string + // assert, not require: t.FailNow from a handler goroutine is unsupported. + if !assert.NoError(t, json.NewDecoder(r.Body).Decode(&req)) { + http.Error(w, "bad request body", http.StatusBadRequest) + return + } + mu.Lock() + deviceClientID = req["client_id"] + mu.Unlock() + w.Header().Set("Content-Type", "application/json") + _, _ = fmt.Fprintf(w, `{"device_code":"%s","user_code":"CCCC-DDDD","verification_uri":"https://github.com/login/device","expires_in":900,"interval":5}`, issued) + })) + t.Cleanup(device.Close) + + tokenSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var req map[string]string + // assert, not require: t.FailNow from a handler goroutine is unsupported. + if !assert.NoError(t, json.NewDecoder(r.Body).Decode(&req)) { + http.Error(w, "bad request body", http.StatusBadRequest) + return + } + mu.Lock() + pollClientIDs = append(pollClientIDs, req["client_id"]) + bound := req["client_id"] == deviceClientID && req["device_code"] == issued + pollCount++ + first := pollCount == 1 + mu.Unlock() + + w.Header().Set("Content-Type", "application/json") + switch { + case !bound: + _, _ = w.Write([]byte(incorrectBody)) + case first: + _, _ = w.Write([]byte(pendingBody)) + default: + _, _ = w.Write([]byte(tokenBody)) + } + })) + t.Cleanup(tokenSrv.Close) + + p := &GitHubATProvider{ + registryURL: registry.URL, + deviceCodeURL: device.URL, + accessTokenURL: tokenSrv.URL, + pollInterval: defaultPollInterval, + expiresIn: defaultExpiresIn, + sleep: func(time.Duration) {}, + } + + require.NoError(t, p.Login(context.Background())) + assert.Equal(t, "gho_test_token", p.githubToken) + + mu.Lock() + defer mu.Unlock() + assert.Equal(t, 1, healthCalls, "Login must read the health endpoint exactly once") + assert.Equal(t, "Iv23liCLIENT0001", deviceClientID) + for _, id := range pollClientIDs { + assert.Equal(t, deviceClientID, id, + "every poll must carry the client_id the device code was issued under") + } +} diff --git a/docs/reference/cli/commands.md b/docs/reference/cli/commands.md index c8d761579..690a7f43d 100644 --- a/docs/reference/cli/commands.md +++ b/docs/reference/cli/commands.md @@ -70,6 +70,10 @@ mcp-publisher login github [--token=PAT] [--registry=URL] [publishing from GitHub Actions](../../modelcontextprotocol-io/github-actions.mdx) authenticates without a browser. This flag is accepted by `login github` only. +Interactive login handles retryable OAuth errors in both HTTP 200 and HTTP 400 responses. +Access denial, expired codes, and unrecognized OAuth errors still stop polling; an +`incorrect_device_code` error receives at most two additional attempts per login. + #### GitHub OIDC (CI/CD) ```bash mcp-publisher login github-oidc [--registry=URL]