Files
trufflehog/pkg/sources/s3/s3_test.go
ahrav e4956615ad
Lint / golangci-lint (push) Waiting to run
Lint / semgrep (push) Waiting to run
Release / Release (push) Waiting to run
Scan for secrets / test (push) Waiting to run
Test / test (push) Waiting to run
Test / test-community (push) Waiting to run
[feat] - Support S3 Source Resumption (#3570)
* add config option for s3 resumption

* updates

* initial progress tracking logic

* more testing

* revert s3 source file

* UpdateScanProgress tests

* adjust

* updates

* invert

* updates

* updates

* fix

* update

* adjust test

* fix

* remove progress tracking

* cleanup

* cleanup

* remove dupe

* remove context cancellation logic

* fix comment format

* make resumption logic more clear

* rename

* fixes

* update

* add edge case test

* remove dupe mu

* add comment

* fix comment
2024-11-22 13:33:34 -08:00

143 lines
4.0 KiB
Go

package s3
import (
"encoding/base64"
"fmt"
"os"
"sync"
"testing"
"time"
"github.com/kylelemons/godebug/pretty"
"github.com/stretchr/testify/assert"
"google.golang.org/protobuf/types/known/anypb"
"github.com/trufflesecurity/trufflehog/v3/pkg/common"
"github.com/trufflesecurity/trufflehog/v3/pkg/context"
"github.com/trufflesecurity/trufflehog/v3/pkg/pb/credentialspb"
"github.com/trufflesecurity/trufflehog/v3/pkg/pb/sourcespb"
"github.com/trufflesecurity/trufflehog/v3/pkg/sources"
)
func TestSource_Init_IncludeAndIgnoreBucketsError(t *testing.T) {
conn, err := anypb.New(&sourcespb.S3{
Credential: &sourcespb.S3_AccessKey{
AccessKey: &credentialspb.KeySecret{
Key: "ignored for test",
Secret: "ignore for test",
},
},
Buckets: []string{"a"},
IgnoreBuckets: []string{"b"},
})
assert.NoError(t, err)
s := Source{}
err = s.Init(context.Background(), "s3 test source", 0, 0, false, conn, 1)
assert.Error(t, err)
}
func TestSource_Chunks(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), time.Second*30)
defer cancel()
secret, err := common.GetTestSecret(ctx)
if err != nil {
t.Fatal(fmt.Errorf("failed to access secret: %v", err))
}
s3key := secret.MustGetField("AWS_S3_KEY")
s3secret := secret.MustGetField("AWS_S3_SECRET")
type init struct {
name string
verify bool
connection *sourcespb.S3
setEnv map[string]string
}
tests := []struct {
name string
init init
wantErr bool
wantChunkData string
}{
{
name: "gets chunks",
init: init{
connection: &sourcespb.S3{
Credential: &sourcespb.S3_AccessKey{
AccessKey: &credentialspb.KeySecret{
Key: s3key,
Secret: s3secret,
},
},
Buckets: []string{"truffletestbucket-s3-tests"},
},
},
wantErr: false,
wantChunkData: `W2RlZmF1bHRdCmF3c19hY2Nlc3Nfa2V5X2lkID0gQUtJQTM1T0hYMkRTT1pHNjQ3TkgKYXdzX3NlY3JldF9hY2Nlc3Nfa2V5ID0gUXk5OVMrWkIvQ1dsRk50eFBBaWQ3Z0d6dnNyWGhCQjd1ckFDQUxwWgpvdXRwdXQgPSBqc29uCnJlZ2lvbiA9IHVzLWVhc3QtMg==`,
},
{
name: "gets chunks after assuming role",
// This test will attempt to scan every bucket in the account, but the role policy blocks access to every
// bucket except the one we want. This (expected behavior) causes errors in the test log output, but these
// errors shouldn't actually cause test failures.
init: init{
connection: &sourcespb.S3{
Roles: []string{"arn:aws:iam::619888638459:role/s3-test-assume-role"},
},
setEnv: map[string]string{
"AWS_ACCESS_KEY_ID": s3key,
"AWS_SECRET_ACCESS_KEY": s3secret,
},
},
wantErr: false,
wantChunkData: `W2RlZmF1bHRdCmF3c19zZWNyZXRfYWNjZXNzX2tleSA9IFF5OTlTK1pCL0NXbEZOdHhQQWlkN2dHenZzclhoQkI3dXJBQ0FMcFoKYXdzX2FjY2Vzc19rZXlfaWQgPSBBS0lBMzVPSFgyRFNPWkc2NDdOSApvdXRwdXQgPSBqc29uCnJlZ2lvbiA9IHVzLWVhc3QtMg==`,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), time.Second*30)
var cancelOnce sync.Once
defer cancelOnce.Do(cancel)
for k, v := range tt.init.setEnv {
t.Setenv(k, v)
}
s := Source{}
conn, err := anypb.New(tt.init.connection)
if err != nil {
t.Fatal(err)
}
err = s.Init(ctx, tt.init.name, 0, 0, tt.init.verify, conn, 8)
if (err != nil) != tt.wantErr {
t.Errorf("Source.Init() error = %v, wantErr %v", err, tt.wantErr)
return
}
chunksCh := make(chan *sources.Chunk)
var wg sync.WaitGroup
wg.Add(1)
go func() {
defer wg.Done()
err = s.Chunks(ctx, chunksCh)
if (err != nil) != tt.wantErr {
t.Errorf("Source.Chunks() error = %v, wantErr %v", err, tt.wantErr)
os.Exit(1)
}
}()
gotChunk := <-chunksCh
wantData, _ := base64.StdEncoding.DecodeString(tt.wantChunkData)
if diff := pretty.Compare(gotChunk.Data, wantData); diff != "" {
t.Errorf("%s: Source.Chunks() diff: (-got +want)\n%s", tt.name, diff)
}
wg.Wait()
assert.Equal(t, "", s.GetProgress().EncodedResumeInfo)
assert.Equal(t, int64(100), s.GetProgress().PercentComplete)
})
}
}