diff --git a/submitqueue/extension/changeprovider/github/provider.go b/submitqueue/extension/changeprovider/github/provider.go index 8ca69b5f1..748e6553f 100644 --- a/submitqueue/extension/changeprovider/github/provider.go +++ b/submitqueue/extension/changeprovider/github/provider.go @@ -106,10 +106,6 @@ func (p *provider) fetchAllPRs( return nil, fmt.Errorf("failed to fetch PR #%d: %w", cid.PRNumber, err) } - if err := validatePRStaleness(cid, prData); err != nil { - return nil, err - } - changeInfo := convertToChangeInfo(cid, prData) changeInfos = append(changeInfos, changeInfo) @@ -137,6 +133,10 @@ func (p *provider) fetchPullRequest(ctx context.Context, parsed entitygithub.Cha return nil, err } + if err := validatePRStaleness(parsed, data); err != nil { + return nil, err + } + if cursor == "" { prData = *data } diff --git a/submitqueue/extension/changeprovider/github/provider_test.go b/submitqueue/extension/changeprovider/github/provider_test.go index e7e968e8b..9ad844931 100644 --- a/submitqueue/extension/changeprovider/github/provider_test.go +++ b/submitqueue/extension/changeprovider/github/provider_test.go @@ -3,6 +3,7 @@ package github import ( "context" "encoding/json" + "fmt" "net/http" "net/http/httptest" "testing" @@ -160,6 +161,45 @@ func TestProvider_Get_Pagination(t *testing.T) { assert.Len(t, infos[0].Details.ChangedFiles, 2) } +func TestProvider_Get_PaginationRejectsStaleHead(t *testing.T) { + for _, tt := range []struct { + name string + headSHAs []string + wantCalls int + }{ + {name: "first page", headSHAs: []string{shaNew, shaXYZ}, wantCalls: 1}, + {name: "middle page", headSHAs: []string{shaXYZ, shaNew, shaXYZ}, wantCalls: 2}, + {name: "last page", headSHAs: []string{shaXYZ, shaNew}, wantCalls: 2}, + } { + t.Run(tt.name, func(t *testing.T) { + callCount := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.Less(t, callCount, len(tt.headSHAs)) + headSHA := tt.headSHAs[callCount] + callCount++ + servePR(t, w, pullRequestData{ + Number: 456, + HeadRefOid: headSHA, + Files: filesData{ + PageInfo: pageInfo{EndCursor: fmt.Sprintf("cursor%d", callCount), HasNextPage: callCount < len(tt.headSHAs)}, + Nodes: []fileNode{{Path: fmt.Sprintf("file%d.go", callCount)}}, + }, + }) + })) + t.Cleanup(server.Close) + + p := newTestProvider(t, server.URL) + infos, err := p.Get(context.Background(), entity.Request{Change: change.Change{ + URIs: []string{"github://github.example.com/uber/submitqueue/pull/456/" + shaXYZ}, + }}) + + require.Error(t, err) + assert.Empty(t, infos) + assert.Equal(t, tt.wantCalls, callCount) + }) + } +} + func TestProvider_Get_MultiplePRs(t *testing.T) { prData := []pullRequestData{ {Number: 123, HeadRefOid: shaA, Files: filesData{Nodes: []fileNode{{Path: "file1.go"}}}},