Add missing pagination on github calls (#30)

* Add missing pagination on github calls

Includes some refactoring to improve readability and code reuse.

* Close response body and handle rate limit

* Re-include support for including users as repos to github scans

* Fix gist test to match new func signature

* Add current test name to logging

* Support username as org use case

* Also include no-auth user as org

Co-authored-by: Bill Rich <[email protected]>
This commit is contained in:
Bill Rich
2022-02-15 18:54:47 -08:00
committed by Dustin Decker
co-authored by Bill Rich
parent 6b183424f5
commit 1fb767247f
5 changed files with 246 additions and 93 deletions
+4 -4
View File
@@ -21,16 +21,16 @@ check:
go vet $(shell go list ./... | grep -v /vendor/) go vet $(shell go list ./... | grep -v /vendor/)
test-failing: test-failing:
CGO_ENABLED=0 go test -timeout=30s $(shell go list ./... | grep -v /vendor/) | grep FAIL CGO_ENABLED=0 go test -timeout=5m $(shell go list ./... | grep -v /vendor/) | grep FAIL
test: test:
CGO_ENABLED=0 go test -timeout=30s $(shell go list ./... | grep -v /vendor/ | grep -v /pkg/detectors) CGO_ENABLED=0 go test -timeout=5m $(shell go list ./... | grep -v /vendor/ | grep -v /pkg/detectors)
test-race: test-race:
CGO_ENABLED=1 go test -timeout=30s -race $(shell go list ./... | grep -v /vendor/ | grep -v /pkg/detectors) CGO_ENABLED=1 go test -timeout=5m -race $(shell go list ./... | grep -v /vendor/ | grep -v /pkg/detectors)
test-detectors: test-detectors:
CGO_ENABLED=0 go test -timeout=30s $(shell go list ./... | grep /pkg/detectors) CGO_ENABLED=0 go test -timeout=5m $(shell go list ./... | grep /pkg/detectors)
bench: bench:
CGO_ENABLED=0 go test $(shell go list ./pkg/secrets/... | grep -v /vendor/) -benchmem -run=xxx -bench . CGO_ENABLED=0 go test $(shell go list ./pkg/secrets/... | grep -v /vendor/) -benchmem -run=xxx -bench .
+19
View File
@@ -0,0 +1,19 @@
package common
func AddStringSliceItem(item string, slice *[]string) {
for _, i := range *slice {
if i == item {
return
}
}
*slice = append(*slice, item)
}
func RemoveStringSliceItem(item string, slice *[]string) {
for i, listItem := range *slice {
if item == listItem {
(*slice)[i] = (*slice)[len(*slice)-1]
*slice = (*slice)[:len(*slice)-1]
}
}
}
+76
View File
@@ -0,0 +1,76 @@
package common
import (
"reflect"
"testing"
)
func TestAddItem(t *testing.T) {
type Case struct {
Slice []string
Modifier []string
Expected []string
}
tests := map[string]Case{
"newItem": {
Slice: []string{"a", "b", "c"},
Modifier: []string{"d"},
Expected: []string{"a", "b", "c", "d"},
},
"newDuplicate": {
Slice: []string{"a", "b", "c"},
Modifier: []string{"c"},
Expected: []string{"a", "b", "c"},
},
}
for name, test := range tests {
for _, item := range test.Modifier {
AddStringSliceItem(item, &test.Slice)
}
if !reflect.DeepEqual(test.Slice, test.Expected) {
t.Errorf("%s: expected:%v, got:%v", name, test.Expected, test.Slice)
}
}
}
func TestRemoveItem(t *testing.T) {
type Case struct {
Slice []string
Modifier []string
Expected []string
}
tests := map[string]Case{
"existingItemEnd": {
Slice: []string{"a", "b", "c"},
Modifier: []string{"c"},
Expected: []string{"a", "b"},
},
"existingItemMiddle": {
Slice: []string{"a", "b", "c"},
Modifier: []string{"b"},
Expected: []string{"a", "c"},
},
"existingItemBeginning": {
Slice: []string{"a", "b", "c"},
Modifier: []string{"a"},
Expected: []string{"c", "b"},
},
"nonExistingItem": {
Slice: []string{"a", "b", "c"},
Modifier: []string{"d"},
Expected: []string{"a", "b", "c"},
},
}
for name, test := range tests {
for _, item := range test.Modifier {
RemoveStringSliceItem(item, &test.Slice)
}
if !reflect.DeepEqual(test.Slice, test.Expected) {
t.Errorf("%s: expected:%v, got:%v", name, test.Expected, test.Slice)
}
}
}
+87 -73
View File
@@ -141,30 +141,18 @@ func (s *Source) Chunks(ctx context.Context, chunksChan chan *sources.Chunk) err
s.repos = s.conn.Repositories s.repos = s.conn.Repositories
s.orgs = s.conn.Organizations s.orgs = s.conn.Organizations
var apiClient *github.Client
switch cred := s.conn.GetCredential().(type) { switch cred := s.conn.GetCredential().(type) {
case *sourcespb.GitHub_Unauthenticated: case *sourcespb.GitHub_Unauthenticated:
apiClient := github.NewClient(s.httpClient) apiClient = github.NewClient(s.httpClient)
if len(s.orgs) > 30 { if len(s.orgs) > 30 {
log.Warn("You may experience rate limiting when using the unauthenticated GitHub api. Consider using an authenticated scan instead.") log.Warn("You may experience rate limiting when using the unauthenticated GitHub api. Consider using an authenticated scan instead.")
} }
if len(s.repos) > 0 {
for i, repo := range s.repos {
if !strings.HasSuffix(repo, ".git") {
if repo, err := giturl.NormalizeGithubRepo(repo); err != nil {
// This wasn't formatted as expected, let the user know why that might be.
log.WithError(err).Warnf("Repo not in expected format, attempting to paginate repos instead.")
} else {
s.repos[i] = repo
}
s.paginateRepos(ctx, apiClient, repo)
}
}
}
if len(s.orgs) > 0 { if len(s.orgs) > 0 {
for _, org := range s.orgs { for _, org := range s.orgs {
s.paginateRepos(ctx, apiClient, org) s.addReposByOrg(ctx, apiClient, org)
s.addReposByUser(ctx, apiClient, org)
} }
} }
case *sourcespb.GitHub_Token: case *sourcespb.GitHub_Token:
@@ -177,7 +165,6 @@ func (s *Source) Chunks(ctx context.Context, chunksChan chan *sources.Chunk) err
) )
tc := oauth2.NewClient(context.TODO(), ts) tc := oauth2.NewClient(context.TODO(), ts)
var apiClient *github.Client
var err error var err error
// If we're using public github, make a regular client. // If we're using public github, make a regular client.
// Otherwise make an enterprise client // Otherwise make an enterprise client
@@ -196,24 +183,14 @@ func (s *Source) Chunks(ctx context.Context, chunksChan chan *sources.Chunk) err
if len(s.repos) > 0 { if len(s.repos) > 0 {
specificScope = true specificScope = true
for i, repo := range s.repos {
if !strings.HasSuffix(repo, ".git") {
if repo, err := giturl.NormalizeGithubRepo(repo); err != nil {
// This wasn't formatted as expected, let the user know why that might be.
log.WithError(err).Warnf("Repo not in expected format, attempting to paginate repos instead.")
} else {
s.repos[i] = repo
}
s.paginateRepos(ctx, apiClient, repo)
}
}
} }
if len(s.orgs) > 0 { if len(s.orgs) > 0 {
specificScope = true specificScope = true
for _, org := range s.orgs { for _, org := range s.orgs {
if !strings.HasSuffix(org, ".git") { if !strings.HasSuffix(org, ".git") {
s.paginateRepos(ctx, apiClient, org) s.addReposByOrg(ctx, apiClient, org)
s.addReposByUser(ctx, apiClient, org)
} }
} }
} }
@@ -222,14 +199,21 @@ func (s *Source) Chunks(ctx context.Context, chunksChan chan *sources.Chunk) err
if err != nil { if err != nil {
return errors.New(err) return errors.New(err)
} }
// TODO: this should enumerate an organizations gists too...
s.paginateGists(ctx, user.GetLogin(), chunksChan)
if !specificScope { if !specificScope {
s.paginateRepos(ctx, apiClient, user.GetLogin()) s.addReposByUser(ctx, apiClient, user.GetLogin())
// Scan for orgs is default with a token. GitHub App enumerates the repositories // Scan for orgs is default with a token. GitHub App enumerates the repositories
// that were assigned to it in GitHub App settings. // that were assigned to it in GitHub App settings.
s.paginateOrgs(ctx, apiClient, *user.Name) s.addOrgsByUser(ctx, apiClient, user.GetLogin())
for _, org := range s.orgs {
s.addReposByOrg(ctx, apiClient, org)
}
}
s.addGistsByUser(ctx, apiClient, user.GetLogin())
for _, org := range s.orgs {
// TODO: Test it actually works to list org gists like this.
s.addGistsByUser(ctx, apiClient, org)
} }
case *sourcespb.GitHub_GithubApp: case *sourcespb.GitHub_GithubApp:
installationID, err := strconv.ParseInt(cred.GithubApp.InstallationId, 10, 64) installationID, err := strconv.ParseInt(cred.GithubApp.InstallationId, 10, 64)
@@ -252,7 +236,7 @@ func (s *Source) Chunks(ctx context.Context, chunksChan chan *sources.Chunk) err
return errors.New(err) return errors.New(err)
} }
itr.BaseURL = apiEndpoint itr.BaseURL = apiEndpoint
apiClient, err := github.NewEnterpriseClient(apiEndpoint, apiEndpoint, &http.Client{Transport: itr}) apiClient, err = github.NewEnterpriseClient(apiEndpoint, apiEndpoint, &http.Client{Transport: itr})
if err != nil { if err != nil {
return errors.New(err) return errors.New(err)
} }
@@ -272,22 +256,21 @@ func (s *Source) Chunks(ctx context.Context, chunksChan chan *sources.Chunk) err
return errors.New(err) return errors.New(err)
} }
err = s.paginateApp(ctx, apiClient) err = s.addReposByApp(ctx, apiClient)
if err != nil { if err != nil {
return err return err
} }
//check if we need to find user repos //check if we need to find user repos
if s.conn.ScanUsers { if s.conn.ScanUsers {
err := s.paginateMembers(ctx, installationClient, apiClient) err := s.addMembersByApp(ctx, installationClient, apiClient)
if err != nil { if err != nil {
return err return err
} }
log.Infof("Scanning repos from %v organization members.", len(s.members)) log.Infof("Scanning repos from %v organization members.", len(s.members))
for _, member := range s.members { for _, member := range s.members {
//all org member's gists s.addGistsByUser(ctx, apiClient, member)
s.paginateGists(ctx, member, chunksChan) s.addReposByUser(ctx, apiClient, member)
s.paginateRepos(ctx, apiClient, member)
} }
} }
@@ -295,6 +278,8 @@ func (s *Source) Chunks(ctx context.Context, chunksChan chan *sources.Chunk) err
return errors.Errorf("Invalid configuration given for source. Name: %s, Type: %s", s.name, s.Type()) return errors.Errorf("Invalid configuration given for source. Name: %s, Type: %s", s.name, s.Type())
} }
s.normalizeRepos(ctx, apiClient)
if _, ok := os.LookupEnv("DO_NOT_RANDOMIZE"); !ok { if _, ok := os.LookupEnv("DO_NOT_RANDOMIZE"); !ok {
//Randomize channel scan order on each scan //Randomize channel scan order on each scan
rand.Seed(time.Now().UnixNano()) rand.Seed(time.Now().UnixNano())
@@ -341,7 +326,7 @@ func (s *Source) Chunks(ctx context.Context, chunksChan chan *sources.Chunk) err
defer os.RemoveAll(path) defer os.RemoveAll(path)
if err != nil { if err != nil {
log.WithError(err).Errorf("unable to clone repo, continuing") log.WithError(err).Errorf("unable to clone repo (%s), continuing", repoURL)
return return
} }
err = s.git.ScanRepo(ctx, repo, git.NewScanOptions(), chunksChan) err = s.git.ScanRepo(ctx, repo, git.NewScanOptions(), chunksChan)
@@ -377,7 +362,7 @@ func handleRateLimit(err error) bool {
return true return true
} }
func (s *Source) paginateReposByOrg(ctx context.Context, apiClient *github.Client, org string) { func (s *Source) addReposByOrg(ctx context.Context, apiClient *github.Client, org string) {
opts := &github.RepositoryListByOrgOptions{ opts := &github.RepositoryListByOrgOptions{
ListOptions: github.ListOptions{ ListOptions: github.ListOptions{
PerPage: 100, PerPage: 100,
@@ -395,7 +380,7 @@ func (s *Source) paginateReposByOrg(ctx context.Context, apiClient *github.Clien
break break
} }
for _, r := range someRepos { for _, r := range someRepos {
s.repos = append(s.repos, r.GetCloneURL()) common.AddStringSliceItem(r.GetCloneURL(), &s.repos)
} }
if res.NextPage == 0 { if res.NextPage == 0 {
break break
@@ -404,9 +389,8 @@ func (s *Source) paginateReposByOrg(ctx context.Context, apiClient *github.Clien
} }
} }
func (s *Source) paginateRepos(ctx context.Context, apiClient *github.Client, user string) { func (s *Source) addReposByUser(ctx context.Context, apiClient *github.Client, user string) {
opts := &github.RepositoryListOptions{ opts := &github.RepositoryListOptions{
// Visibility: "all",
ListOptions: github.ListOptions{ ListOptions: github.ListOptions{
PerPage: 50, PerPage: 50,
}, },
@@ -423,7 +407,7 @@ func (s *Source) paginateRepos(ctx context.Context, apiClient *github.Client, us
break break
} }
for _, r := range someRepos { for _, r := range someRepos {
s.repos = append(s.repos, r.GetCloneURL()) common.AddStringSliceItem(r.GetCloneURL(), &s.repos)
} }
if res.NextPage == 0 { if res.NextPage == 0 {
break break
@@ -432,34 +416,31 @@ func (s *Source) paginateRepos(ctx context.Context, apiClient *github.Client, us
} }
} }
func (s *Source) paginateGists(ctx context.Context, user string, chunksChan chan *sources.Chunk) { func (s *Source) addGistsByUser(ctx context.Context, apiClient *github.Client, user string) {
apiClient := github.NewClient(s.httpClient) gistOpts := &github.GistListOptions{}
gists, _, err := apiClient.Gists.List(ctx, user, &github.GistListOptions{}) for {
gists, resp, err := apiClient.Gists.List(ctx, user, gistOpts)
if err == nil {
defer resp.Body.Close()
}
if handled := handleRateLimit(err); handled {
continue
}
if err != nil { if err != nil {
log.WithError(err).Warnf("Could not get gists for user %s", user) log.WithError(err).Warnf("Could not get gists for user %s", user)
return
} }
for _, gist := range gists { for _, gist := range gists {
path, repo, err := git.CloneRepoUsingUnauthenticated(*gist.GitPullURL) common.AddStringSliceItem(gist.GetGitPullURL(), &s.repos)
defer os.RemoveAll(path)
if err != nil {
log.WithError(err).Warnf("Could not get gist %s from user %s", *gist.HTMLURL, user)
continue
} }
s.log.WithField("repo", *gist.HTMLURL).Debugf("attempting to clone gist from user %s", user) if resp == nil || resp.NextPage == 0 {
break
scanCtx := context.Background()
err = s.git.ScanRepo(scanCtx, repo, git.NewScanOptions(), chunksChan)
if err != nil {
log.WithError(err).Warnf("Could not scan after clone: %s", *gist.HTMLURL)
continue
} }
gistOpts.Page = resp.NextPage
} }
return
} }
func (s *Source) paginateMembers(ctx context.Context, installationClient *github.Client, apiClient *github.Client) error { func (s *Source) addMembersByApp(ctx context.Context, installationClient *github.Client, apiClient *github.Client) error {
opts := &github.ListOptions{ opts := &github.ListOptions{
PerPage: 500, PerPage: 500,
@@ -493,20 +474,19 @@ func (s *Source) paginateMembers(ctx context.Context, installationClient *github
if usr == nil || *usr == "" { if usr == nil || *usr == "" {
continue continue
} }
s.members = append(s.members, *usr) common.AddStringSliceItem(*usr, &s.members)
} }
if res.NextPage == 0 { if res.NextPage == 0 {
break break
} }
opts.Page = res.NextPage opts.Page = res.NextPage
} }
} }
return nil return nil
} }
func (s *Source) paginateApp(ctx context.Context, apiClient *github.Client) error { func (s *Source) addReposByApp(ctx context.Context, apiClient *github.Client) error {
// Authenticated enumeration of repos // Authenticated enumeration of repos
opts := &github.ListOptions{ opts := &github.ListOptions{
PerPage: 100, PerPage: 100,
@@ -523,7 +503,7 @@ func (s *Source) paginateApp(ctx context.Context, apiClient *github.Client) erro
return errors.WrapPrefix(err, "unable to list repositories", 0) return errors.WrapPrefix(err, "unable to list repositories", 0)
} }
for _, r := range someRepos.Repositories { for _, r := range someRepos.Repositories {
s.repos = append(s.repos, r.GetCloneURL()) common.AddStringSliceItem(r.GetCloneURL(), &s.repos)
} }
if res.NextPage == 0 { if res.NextPage == 0 {
break break
@@ -533,9 +513,18 @@ func (s *Source) paginateApp(ctx context.Context, apiClient *github.Client) erro
return nil return nil
} }
func (s *Source) paginateOrgs(ctx context.Context, apiClient *github.Client, user string) { func (s *Source) addOrgsByUser(ctx context.Context, apiClient *github.Client, user string) {
orgOpts := &github.ListOptions{} orgOpts := &github.ListOptions{
orgs, _, err := apiClient.Organizations.List(ctx, "", orgOpts) PerPage: 100,
}
for {
orgs, resp, err := apiClient.Organizations.List(ctx, "", orgOpts)
if err == nil {
defer resp.Body.Close()
}
if handled := handleRateLimit(err); handled {
continue
}
if err != nil { if err != nil {
log.WithError(err).Errorf("Could not list organizations for %s", user) log.WithError(err).Errorf("Could not list organizations for %s", user)
return return
@@ -549,7 +538,32 @@ func (s *Source) paginateOrgs(ctx context.Context, apiClient *github.Client, use
} else { } else {
continue continue
} }
s.paginateReposByOrg(ctx, apiClient, name) common.AddStringSliceItem(name, &s.orgs)
}
if resp.NextPage == 0 {
break
}
orgOpts.Page = resp.NextPage
}
}
func (s *Source) normalizeRepos(ctx context.Context, apiClient *github.Client) {
// TODO: Add check/fix for repos that are missing scheme
var newRepoList []string
for _, repo := range s.repos {
if parts := strings.Split(repo, "/"); len(parts) == 1 {
origSources := len(s.repos)
s.addGistsByUser(ctx, apiClient, repo)
s.addReposByUser(ctx, apiClient, repo)
if origSources != len(s.repos) {
common.RemoveStringSliceItem(repo, &s.repos)
continue
}
}
repoNormalized, err := giturl.NormalizeGithubRepo(repo)
if err != nil {
log.WithError(err).Warnf("Repo not in expected format: %s", repo)
}
newRepoList = append(newRepoList, repoNormalized)
} }
} }
+46 -2
View File
@@ -5,6 +5,7 @@ import (
"encoding/base64" "encoding/base64"
"errors" "errors"
"fmt" "fmt"
"github.com/google/go-github/v41/github"
"os" "os"
"testing" "testing"
"time" "time"
@@ -378,6 +379,7 @@ func TestSource_paginateGists(t *testing.T) {
wantChunk *sources.Chunk wantChunk *sources.Chunk
wantErr bool wantErr bool
user string user string
minRepos int
}{ }{
{ {
name: "get gist secret", name: "get gist secret",
@@ -406,6 +408,26 @@ func TestSource_paginateGists(t *testing.T) {
}, },
wantErr: false, wantErr: false,
user: "dustin-decker", user: "dustin-decker",
minRepos: 1,
},
{
name: "get multiple pages of gists",
init: init{
name: "test source",
connection: &sourcespb.GitHub{
Credential: &sourcespb.GitHub_GithubApp{
GithubApp: &credentialspb.GitHubApp{
PrivateKey: githubPrivateKeyNew,
InstallationId: githubInstallationIDNew,
AppId: githubAppIDNew,
},
},
},
},
wantChunk: nil,
wantErr: false,
user: "andrew",
minRepos: 101,
}, },
/* { /* {
name: "get multiple pages of gists", name: "get multiple pages of gists",
@@ -459,15 +481,37 @@ func TestSource_paginateGists(t *testing.T) {
} }
chunksCh := make(chan *sources.Chunk, 5) chunksCh := make(chan *sources.Chunk, 5)
go func() { go func() {
s.paginateGists(ctx, tt.user, chunksCh) s.addGistsByUser(ctx, github.NewClient(s.httpClient), tt.user)
chunksCh <- &sources.Chunk{}
}() }()
if err = handleChannel(chunksCh, basicCheckFunc(0, 0, tt.wantChunk, &s)); err != nil { var wantedRepo string
if tt.wantChunk != nil {
wantedRepo = tt.wantChunk.SourceMetadata.GetGithub().Repository
}
if err = handleChannel(chunksCh, gistsCheckFunc(wantedRepo, tt.minRepos, &s)); err != nil {
t.Error(err) t.Error(err)
} }
}) })
} }
} }
func gistsCheckFunc(expected string, minRepos int, s *Source) chunkFunc {
return func(chunk *sources.Chunk) error {
if minRepos != 0 && minRepos > len(s.repos) {
return fmt.Errorf("didn't find enough repos. expected: %d, got :%d", minRepos, len(s.repos))
}
if expected != "" {
for _, repo := range s.repos {
if repo == expected {
return nil
}
}
return fmt.Errorf("expected repo not included: %s", expected)
}
return nil
}
}
func basicCheckFunc(minOrg, minRepo int, wantChunk *sources.Chunk, s *Source) chunkFunc { func basicCheckFunc(minOrg, minRepo int, wantChunk *sources.Chunk, s *Source) chunkFunc {
return func(chunk *sources.Chunk) error { return func(chunk *sources.Chunk) error {
if minOrg != 0 && minOrg > len(s.orgs) { if minOrg != 0 && minOrg > len(s.orgs) {