[INS-104] Support units in S3 source (#4560)
* implemented Source unit for S3. Implemented. integration test * use bucket as source unit * remove code duplication, reuse from Chunks * remove unnecessary change * remove unused functions * revisit tests * revert unnecessary change * change SourceUnitKind to s3_bucket * handle nil objectCount inside scanBucket * handle nil objectCount outside loop * add bucket to resume log * add bucket and role to error log, remove enumerating log * implement sub unit resumption * add comment to checkpointer for unit scans * implement SourceUnitUnmarshaller on source with the new S3SourceUnit, add test to test resumption on multiple buckets with concurrent ChunkUnit processing * add role to SourceUnitID * Revert "add role to SourceUnitID" This reverts commit 549e6bede9fe4ed276f7a23f5a5b2bf037133264. * add role to source unit ID, keep track of resumption using source unit ID instead of just bucket name * rename bucket -> unitID in UnmarshalSourceUnit --------- Co-authored-by: Amaan Ullah <[email protected]>
This commit is contained in:
@@ -11,6 +11,8 @@ import (
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/sources"
|
||||
)
|
||||
|
||||
// TODO [INS-207] Add role to legacy scan resumption info
|
||||
|
||||
// Checkpointer maintains resumption state for S3 bucket scanning,
|
||||
// enabling resumable scans by tracking which objects have been successfully processed.
|
||||
// It provides checkpoints that can be used to resume interrupted scans without missing objects.
|
||||
@@ -33,6 +35,10 @@ import (
|
||||
// resuming from the correct bucket. The scan will continue from the last checkpointed object
|
||||
// in that bucket.
|
||||
//
|
||||
// Unit scans are also supported. The encoded resume info in this case tracks the last processed object
|
||||
// for each unit separately by using the SetEncodedResumeInfoFor method on Progress. To use the
|
||||
// checkpointer for unit scans, call SetIsUnitScan(true) before starting the scan.
|
||||
//
|
||||
// For example, if scanning is interrupted after processing 1500 objects across 2 pages:
|
||||
// Page 1 (objects 0-999): Fully processed, checkpoint saved at object 999
|
||||
// Page 2 (objects 1000-1999): Partially processed through 1600, but only consecutive through 1499
|
||||
@@ -56,6 +62,8 @@ type Checkpointer struct {
|
||||
// progress holds the scan's overall progress state and enables persistence.
|
||||
// The EncodedResumeInfo field stores the JSON-encoded ResumeInfo checkpoint.
|
||||
progress *sources.Progress // Reference to source's Progress
|
||||
|
||||
isUnitScan bool // Indicates if scanning is done in unit scan mode
|
||||
}
|
||||
|
||||
const defaultMaxObjectsPerPage = 1000
|
||||
@@ -153,9 +161,10 @@ func (p *Checkpointer) UpdateObjectCompletion(
|
||||
ctx context.Context,
|
||||
completedIdx int,
|
||||
bucket string,
|
||||
role string,
|
||||
pageContents []s3types.Object,
|
||||
) error {
|
||||
ctx = context.WithValues(ctx, "bucket", bucket, "completedIdx", completedIdx)
|
||||
ctx = context.WithValues(ctx, "bucket", bucket, "role", role, "completedIdx", completedIdx)
|
||||
ctx.Logger().V(5).Info("Updating progress")
|
||||
|
||||
if completedIdx >= len(p.completedObjects) {
|
||||
@@ -184,7 +193,7 @@ func (p *Checkpointer) UpdateObjectCompletion(
|
||||
}
|
||||
obj := pageContents[checkpointIdx]
|
||||
|
||||
return p.updateCheckpoint(bucket, *obj.Key)
|
||||
return p.updateCheckpoint(bucket, role, *obj.Key)
|
||||
}
|
||||
|
||||
// advanceLowestIncompleteIdx moves the lowest incomplete index forward to the next incomplete object.
|
||||
@@ -198,7 +207,14 @@ func (p *Checkpointer) advanceLowestIncompleteIdx() {
|
||||
|
||||
// updateCheckpoint persists the current resumption state.
|
||||
// Must be called with lock held.
|
||||
func (p *Checkpointer) updateCheckpoint(bucket string, lastKey string) error {
|
||||
func (p *Checkpointer) updateCheckpoint(bucket string, role string, lastKey string) error {
|
||||
if p.isUnitScan {
|
||||
unitID := constructS3SourceUnitID(bucket, role)
|
||||
// track sub-unit resumption state
|
||||
p.progress.SetEncodedResumeInfoFor(unitID, lastKey)
|
||||
return nil
|
||||
}
|
||||
|
||||
encoded, err := json.Marshal(&ResumeInfo{CurrentBucket: bucket, StartAfter: lastKey})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to encode resume info: %w", err)
|
||||
@@ -212,3 +228,11 @@ func (p *Checkpointer) updateCheckpoint(bucket string, lastKey string) error {
|
||||
)
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetIsUnitScan sets whether the checkpointer is operating in unit scan mode.
|
||||
func (p *Checkpointer) SetIsUnitScan(isUnitScan bool) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
|
||||
p.isUnitScan = isUnitScan
|
||||
}
|
||||
|
||||
@@ -31,7 +31,7 @@ func TestCheckpointerResumption(t *testing.T) {
|
||||
|
||||
// Process first 6 objects.
|
||||
for i := range 6 {
|
||||
err := tracker.UpdateObjectCompletion(ctx, i, "test-bucket", firstPage.Contents)
|
||||
err := tracker.UpdateObjectCompletion(ctx, i, "test-bucket", "", firstPage.Contents)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
@@ -50,7 +50,7 @@ func TestCheckpointerResumption(t *testing.T) {
|
||||
|
||||
// Process remaining objects.
|
||||
for i := range len(resumePage.Contents) {
|
||||
err := resumeTracker.UpdateObjectCompletion(ctx, i, "test-bucket", resumePage.Contents)
|
||||
err := resumeTracker.UpdateObjectCompletion(ctx, i, "test-bucket", "", resumePage.Contents)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
@@ -244,7 +244,7 @@ func TestCheckpointerUpdate(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
err := tracker.UpdateObjectCompletion(ctx, tt.completedIdx, "test-bucket", page.Contents)
|
||||
err := tracker.UpdateObjectCompletion(ctx, tt.completedIdx, "test-bucket", "", page.Contents)
|
||||
assert.NoError(t, err, "Unexpected error updating progress")
|
||||
|
||||
var info ResumeInfo
|
||||
@@ -258,6 +258,36 @@ func TestCheckpointerUpdate(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckpointerUpdateUnitScan(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
progress := new(sources.Progress)
|
||||
tracker := NewCheckpointer(ctx, progress)
|
||||
tracker.SetIsUnitScan(true)
|
||||
|
||||
page := &s3.ListObjectsV2Output{
|
||||
Contents: make([]s3types.Object, 3),
|
||||
}
|
||||
for i := range 3 {
|
||||
key := fmt.Sprintf("key-%d", i)
|
||||
page.Contents[i] = s3types.Object{Key: &key}
|
||||
}
|
||||
|
||||
// Complete first object.
|
||||
err := tracker.UpdateObjectCompletion(ctx, 0, "test-bucket", "test-role", page.Contents)
|
||||
assert.NoError(t, err, "Unexpected error updating progress")
|
||||
|
||||
var info map[string]string
|
||||
err = json.Unmarshal([]byte(progress.EncodedResumeInfo), &info)
|
||||
var gotUnitID, gotStartAfter string
|
||||
for k, v := range info {
|
||||
gotUnitID = k
|
||||
gotStartAfter = v
|
||||
}
|
||||
assert.NoError(t, err, "Failed to decode resume info")
|
||||
assert.Equal(t, "test-role|test-bucket", gotUnitID, "Incorrect unit ID")
|
||||
assert.Equal(t, "key-0", gotStartAfter, "Incorrect resume point")
|
||||
}
|
||||
|
||||
func TestComplete(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
|
||||
+153
-66
@@ -1,6 +1,7 @@
|
||||
package s3
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"slices"
|
||||
"strings"
|
||||
@@ -54,14 +55,13 @@ type Source struct {
|
||||
errorCount *sync.Map
|
||||
jobPool *errgroup.Group
|
||||
maxObjectSize int64
|
||||
|
||||
sources.CommonSourceUnitUnmarshaller
|
||||
}
|
||||
|
||||
// Ensure the Source satisfies the interfaces at compile time
|
||||
var _ sources.Source = (*Source)(nil)
|
||||
var _ sources.SourceUnitUnmarshaller = (*Source)(nil)
|
||||
var _ sources.Validator = (*Source)(nil)
|
||||
var _ sources.SourceUnitEnumChunker = (*Source)(nil)
|
||||
|
||||
// Type returns the type of source
|
||||
func (s *Source) Type() sourcespb.SourceType { return SourceType }
|
||||
@@ -225,6 +225,7 @@ func (s *Source) getBucketsToScan(ctx context.Context, client *s3.Client) ([]str
|
||||
// pageMetadata contains metadata about a single page of S3 objects being scanned.
|
||||
type pageMetadata struct {
|
||||
bucket string // The name of the S3 bucket being scanned
|
||||
role string // The AWS role ARN used for scanning
|
||||
pageNumber int // Current page number in the pagination sequence
|
||||
client *s3.Client // AWS S3 client configured for the appropriate region
|
||||
page *s3.ListObjectsV2Output // Contains the list of S3 objects in this page
|
||||
@@ -294,7 +295,7 @@ func (s *Source) scanBuckets(
|
||||
if role != "" {
|
||||
ctx = context.WithValue(ctx, "role", role)
|
||||
}
|
||||
var objectCount uint64
|
||||
var totalObjectCount uint64
|
||||
|
||||
pos := determineResumePosition(ctx, s.checkpointer, bucketsToScan)
|
||||
switch {
|
||||
@@ -316,16 +317,7 @@ func (s *Source) scanBuckets(
|
||||
|
||||
bucketsToScanCount := len(bucketsToScan)
|
||||
for bucketIdx := pos.index; bucketIdx < bucketsToScanCount; bucketIdx++ {
|
||||
s.metricsCollector.RecordBucketForRole(role)
|
||||
bucket := bucketsToScan[bucketIdx]
|
||||
ctx := context.WithValue(ctx, "bucket", bucket)
|
||||
|
||||
if common.IsDone(ctx) {
|
||||
ctx.Logger().Error(ctx.Err(), "context done, while scanning bucket")
|
||||
return
|
||||
}
|
||||
|
||||
ctx.Logger().V(3).Info("Scanning bucket")
|
||||
|
||||
s.SetProgressComplete(
|
||||
bucketIdx,
|
||||
@@ -334,63 +326,95 @@ func (s *Source) scanBuckets(
|
||||
s.Progress.EncodedResumeInfo,
|
||||
)
|
||||
|
||||
regionalClient, err := s.getRegionalClientForBucket(ctx, client, role, bucket)
|
||||
if err != nil {
|
||||
ctx.Logger().Error(err, "could not get regional client for bucket")
|
||||
continue
|
||||
}
|
||||
|
||||
errorCount := sync.Map{}
|
||||
|
||||
input := &s3.ListObjectsV2Input{Bucket: &bucket}
|
||||
var startAfter *string
|
||||
if bucket == pos.bucket && pos.startAfter != "" {
|
||||
input.StartAfter = &pos.startAfter
|
||||
startAfter = &pos.startAfter
|
||||
ctx.Logger().V(3).Info(
|
||||
"Resuming bucket scan",
|
||||
"start_after", pos.startAfter,
|
||||
"bucket", bucket,
|
||||
)
|
||||
}
|
||||
|
||||
pageNumber := 1
|
||||
paginator := s3.NewListObjectsV2Paginator(regionalClient, input)
|
||||
for paginator.HasMorePages() {
|
||||
output, err := paginator.NextPage(ctx)
|
||||
if err != nil {
|
||||
if role == "" {
|
||||
ctx.Logger().Error(err, "could not list objects in bucket")
|
||||
} else {
|
||||
// Our documentation blesses specifying a role to assume without specifying buckets to scan, which will
|
||||
// often cause this to happen a lot (because in that case the scanner tries to scan every bucket in the
|
||||
// account, but the role probably doesn't have access to all of them). This makes it expected behavior
|
||||
// and therefore not an error.
|
||||
ctx.Logger().V(3).Info("could not list objects in bucket", "err", err)
|
||||
}
|
||||
break
|
||||
}
|
||||
pageMetadata := pageMetadata{
|
||||
bucket: bucket,
|
||||
pageNumber: pageNumber,
|
||||
client: regionalClient,
|
||||
page: output,
|
||||
}
|
||||
processingState := processingState{
|
||||
errorCount: &errorCount,
|
||||
objectCount: &objectCount,
|
||||
}
|
||||
s.pageChunker(ctx, pageMetadata, processingState, chunksChan)
|
||||
|
||||
pageNumber++
|
||||
}
|
||||
objectCount := s.scanBucket(ctx, client, role, bucket, sources.ChanReporter{Ch: chunksChan}, startAfter)
|
||||
totalObjectCount += objectCount
|
||||
}
|
||||
|
||||
s.SetProgressComplete(
|
||||
len(bucketsToScan),
|
||||
len(bucketsToScan),
|
||||
fmt.Sprintf("Completed scanning source %s. %d objects scanned.", s.name, objectCount),
|
||||
fmt.Sprintf("Completed scanning source %s. %d objects scanned.", s.name, totalObjectCount),
|
||||
"",
|
||||
)
|
||||
}
|
||||
|
||||
func (s *Source) scanBucket(
|
||||
ctx context.Context,
|
||||
client *s3.Client,
|
||||
role string,
|
||||
bucket string,
|
||||
reporter sources.ChunkReporter,
|
||||
startAfter *string,
|
||||
) uint64 {
|
||||
s.metricsCollector.RecordBucketForRole(role)
|
||||
|
||||
ctx = context.WithValue(ctx, "bucket", bucket)
|
||||
|
||||
if common.IsDone(ctx) {
|
||||
ctx.Logger().Error(ctx.Err(), "context done, while scanning bucket")
|
||||
return 0
|
||||
}
|
||||
|
||||
ctx.Logger().V(3).Info("Scanning bucket")
|
||||
|
||||
regionalClient, err := s.getRegionalClientForBucket(ctx, client, role, bucket)
|
||||
if err != nil {
|
||||
ctx.Logger().Error(err, "could not get regional client for bucket")
|
||||
return 0
|
||||
}
|
||||
|
||||
errorCount := sync.Map{}
|
||||
|
||||
input := &s3.ListObjectsV2Input{Bucket: &bucket}
|
||||
if startAfter != nil {
|
||||
input.StartAfter = startAfter
|
||||
}
|
||||
|
||||
pageNumber := 1
|
||||
paginator := s3.NewListObjectsV2Paginator(regionalClient, input)
|
||||
var objectCount uint64
|
||||
for paginator.HasMorePages() {
|
||||
output, err := paginator.NextPage(ctx)
|
||||
if err != nil {
|
||||
if role == "" {
|
||||
ctx.Logger().Error(err, "could not list objects in bucket")
|
||||
} else {
|
||||
// Our documentation blesses specifying a role to assume without specifying buckets to scan, which will
|
||||
// often cause this to happen a lot (because in that case the scanner tries to scan every bucket in the
|
||||
// account, but the role probably doesn't have access to all of them). This makes it expected behavior
|
||||
// and therefore not an error.
|
||||
ctx.Logger().V(3).Info("could not list objects in bucket", "err", err)
|
||||
}
|
||||
break
|
||||
}
|
||||
pageMetadata := pageMetadata{
|
||||
bucket: bucket,
|
||||
role: role,
|
||||
pageNumber: pageNumber,
|
||||
client: regionalClient,
|
||||
page: output,
|
||||
}
|
||||
processingState := processingState{
|
||||
errorCount: &errorCount,
|
||||
objectCount: &objectCount,
|
||||
}
|
||||
s.pageChunker(ctx, pageMetadata, processingState, reporter)
|
||||
|
||||
pageNumber++
|
||||
}
|
||||
return objectCount
|
||||
}
|
||||
|
||||
// Chunks emits chunks of bytes over a channel.
|
||||
func (s *Source) Chunks(ctx context.Context, chunksChan chan *sources.Chunk, _ ...sources.ChunkingTarget) error {
|
||||
visitor := func(c context.Context, defaultRegionClient *s3.Client, roleArn string, buckets []string) error {
|
||||
@@ -429,14 +453,12 @@ func (s *Source) pageChunker(
|
||||
ctx context.Context,
|
||||
metadata pageMetadata,
|
||||
state processingState,
|
||||
chunksChan chan *sources.Chunk,
|
||||
reporter sources.ChunkReporter,
|
||||
) {
|
||||
s.checkpointer.Reset() // Reset the checkpointer for each PAGE
|
||||
ctx = context.WithValues(ctx, "bucket", metadata.bucket, "page_number", metadata.pageNumber)
|
||||
|
||||
for objIdx, obj := range metadata.page.Contents {
|
||||
ctx = context.WithValues(ctx, "key", *obj.Key, "size", *obj.Size)
|
||||
|
||||
if common.IsDone(ctx) {
|
||||
return
|
||||
}
|
||||
@@ -445,7 +467,7 @@ func (s *Source) pageChunker(
|
||||
if obj.StorageClass == s3types.ObjectStorageClassGlacier || obj.StorageClass == s3types.ObjectStorageClassGlacierIr {
|
||||
ctx.Logger().V(5).Info("Skipping object in storage class", "storage_class", obj.StorageClass)
|
||||
s.metricsCollector.RecordObjectSkipped(metadata.bucket, "storage_class", float64(*obj.Size))
|
||||
if err := s.checkpointer.UpdateObjectCompletion(ctx, objIdx, metadata.bucket, metadata.page.Contents); err != nil {
|
||||
if err := s.checkpointer.UpdateObjectCompletion(ctx, objIdx, metadata.bucket, metadata.role, metadata.page.Contents); err != nil {
|
||||
ctx.Logger().Error(err, "could not update progress for glacier object")
|
||||
}
|
||||
continue
|
||||
@@ -455,7 +477,7 @@ func (s *Source) pageChunker(
|
||||
if *obj.Size > s.maxObjectSize {
|
||||
ctx.Logger().V(5).Info("Skipping large file", "max_object_size", s.maxObjectSize)
|
||||
s.metricsCollector.RecordObjectSkipped(metadata.bucket, "size_limit", float64(*obj.Size))
|
||||
if err := s.checkpointer.UpdateObjectCompletion(ctx, objIdx, metadata.bucket, metadata.page.Contents); err != nil {
|
||||
if err := s.checkpointer.UpdateObjectCompletion(ctx, objIdx, metadata.bucket, metadata.role, metadata.page.Contents); err != nil {
|
||||
ctx.Logger().Error(err, "could not update progress for large file")
|
||||
}
|
||||
continue
|
||||
@@ -465,7 +487,7 @@ func (s *Source) pageChunker(
|
||||
if *obj.Size == 0 {
|
||||
ctx.Logger().V(5).Info("Skipping empty file")
|
||||
s.metricsCollector.RecordObjectSkipped(metadata.bucket, "empty_file", 0)
|
||||
if err := s.checkpointer.UpdateObjectCompletion(ctx, objIdx, metadata.bucket, metadata.page.Contents); err != nil {
|
||||
if err := s.checkpointer.UpdateObjectCompletion(ctx, objIdx, metadata.bucket, metadata.role, metadata.page.Contents); err != nil {
|
||||
ctx.Logger().Error(err, "could not update progress for empty file")
|
||||
}
|
||||
continue
|
||||
@@ -475,7 +497,7 @@ func (s *Source) pageChunker(
|
||||
if common.SkipFile(*obj.Key) {
|
||||
ctx.Logger().V(5).Info("Skipping file with incompatible extension")
|
||||
s.metricsCollector.RecordObjectSkipped(metadata.bucket, "incompatible_extension", float64(*obj.Size))
|
||||
if err := s.checkpointer.UpdateObjectCompletion(ctx, objIdx, metadata.bucket, metadata.page.Contents); err != nil {
|
||||
if err := s.checkpointer.UpdateObjectCompletion(ctx, objIdx, metadata.bucket, metadata.role, metadata.page.Contents); err != nil {
|
||||
ctx.Logger().Error(err, "could not update progress for incompatible file")
|
||||
}
|
||||
continue
|
||||
@@ -572,12 +594,11 @@ func (s *Source) pageChunker(
|
||||
Verify: s.verify,
|
||||
}
|
||||
|
||||
if err := handlers.HandleFile(ctx, res.Body, chunkSkel, sources.ChanReporter{Ch: chunksChan}); err != nil {
|
||||
if err := handlers.HandleFile(ctx, res.Body, chunkSkel, reporter); err != nil {
|
||||
ctx.Logger().Error(err, "error handling file")
|
||||
s.metricsCollector.RecordObjectError(metadata.bucket)
|
||||
return nil
|
||||
}
|
||||
|
||||
atomic.AddUint64(state.objectCount, 1)
|
||||
ctx.Logger().V(5).Info("S3 object scanned.", "object_count", state.objectCount)
|
||||
nErr, ok = state.errorCount.Load(prefix)
|
||||
@@ -587,17 +608,14 @@ func (s *Source) pageChunker(
|
||||
if nErr.(int) > 0 {
|
||||
state.errorCount.Store(prefix, 0)
|
||||
}
|
||||
|
||||
// Update progress after successful processing.
|
||||
if err := s.checkpointer.UpdateObjectCompletion(ctx, objIdx, metadata.bucket, metadata.page.Contents); err != nil {
|
||||
if err := s.checkpointer.UpdateObjectCompletion(ctx, objIdx, metadata.bucket, metadata.role, metadata.page.Contents); err != nil {
|
||||
ctx.Logger().Error(err, "could not update progress for scanned object")
|
||||
}
|
||||
s.metricsCollector.RecordObjectScanned(metadata.bucket, float64(*obj.Size))
|
||||
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
_ = s.jobPool.Wait()
|
||||
}
|
||||
|
||||
@@ -681,3 +699,72 @@ func (s *Source) visitRoles(
|
||||
func makeS3Link(bucket, region, key string) string {
|
||||
return fmt.Sprintf("https://%s.s3.%s.amazonaws.com/%s", bucket, region, key)
|
||||
}
|
||||
|
||||
// Enumerate implements SourceUnitEnumerator interface. This implementation visits
|
||||
// each configured role and passes each s3 bucket as a source unit
|
||||
func (s *Source) Enumerate(ctx context.Context, reporter sources.UnitReporter) error {
|
||||
visitor := func(c context.Context, defaultRegionClient *s3.Client, roleArn string, buckets []string) error {
|
||||
for _, bucket := range buckets {
|
||||
if common.IsDone(ctx) {
|
||||
return ctx.Err()
|
||||
}
|
||||
|
||||
unit := S3SourceUnit{
|
||||
Bucket: bucket,
|
||||
Role: roleArn,
|
||||
}
|
||||
|
||||
if err := reporter.UnitOk(ctx, unit); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
return s.visitRoles(ctx, visitor)
|
||||
}
|
||||
|
||||
// ChunkUnit implements SourceUnitChunker interface. This implementation scans
|
||||
// the given S3 bucket source unit and emits chunks for each object found.
|
||||
// It supports sub-unit resumption by utilizing the checkpointer to track progress.
|
||||
func (s *Source) ChunkUnit(ctx context.Context, unit sources.SourceUnit, reporter sources.ChunkReporter) error {
|
||||
|
||||
s3unit, ok := unit.(S3SourceUnit)
|
||||
if !ok {
|
||||
return fmt.Errorf("expected *S3SourceUnit, got %T", unit)
|
||||
}
|
||||
// unitID is a combination of bucket name and role ARN
|
||||
unitID, _ := s3unit.SourceUnitID()
|
||||
defaultClient, err := s.newClient(ctx, defaultAWSRegion, s3unit.Role)
|
||||
if err != nil {
|
||||
return fmt.Errorf("could not create s3 client for bucket %s and role %s: %w", s3unit.Bucket, s3unit.Role, err)
|
||||
}
|
||||
|
||||
s.checkpointer.SetIsUnitScan(true)
|
||||
|
||||
var startAfterPtr *string
|
||||
startAfter := s.Progress.GetEncodedResumeInfoFor(unitID)
|
||||
if startAfter != "" {
|
||||
ctx.Logger().V(3).Info(
|
||||
"Resuming unit scan",
|
||||
"start_after", startAfter,
|
||||
"unitID", unitID,
|
||||
)
|
||||
startAfterPtr = &startAfter
|
||||
}
|
||||
defer s.Progress.ClearEncodedResumeInfoFor(unitID)
|
||||
s.scanBucket(ctx, defaultClient, s3unit.Role, s3unit.Bucket, reporter, startAfterPtr)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Source) UnmarshalSourceUnit(data []byte) (sources.SourceUnit, error) {
|
||||
var unit S3SourceUnit
|
||||
if err := json.Unmarshal(data, &unit); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
unitID, kind := unit.SourceUnitID()
|
||||
if unitID == "" || kind != SourceUnitKindBucket {
|
||||
return nil, fmt.Errorf("not an S3SourceUnit")
|
||||
}
|
||||
return unit, nil
|
||||
}
|
||||
|
||||
@@ -16,8 +16,10 @@ import (
|
||||
"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/source_metadatapb"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/pb/sourcespb"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/sources"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/sourcestest"
|
||||
)
|
||||
|
||||
func TestSource_ChunksCount(t *testing.T) {
|
||||
@@ -391,3 +393,194 @@ func TestSourceChunksResumptionMultipleBucketsIgnoredBucket(t *testing.T) {
|
||||
|
||||
assert.Equal(t, 103, count, "Should have processed all remaining data on resume")
|
||||
}
|
||||
|
||||
func TestSource_Enumerate(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
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")
|
||||
|
||||
connection := &sourcespb.S3{
|
||||
Credential: &sourcespb.S3_AccessKey{
|
||||
AccessKey: &credentialspb.KeySecret{
|
||||
Key: s3key,
|
||||
Secret: s3secret,
|
||||
},
|
||||
},
|
||||
Buckets: []string{"truffletestbucket"},
|
||||
}
|
||||
|
||||
conn, err := anypb.New(connection)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
s := Source{}
|
||||
err = s.Init(ctx, "test enumerate", 0, 0, false, conn, 1)
|
||||
assert.NoError(t, err)
|
||||
|
||||
reporter := sourcestest.TestReporter{}
|
||||
err = s.Enumerate(ctx, &reporter)
|
||||
assert.NoError(t, err)
|
||||
|
||||
assert.Equal(t, len(reporter.Units), 1)
|
||||
assert.Equal(t, 0, len(reporter.UnitErrs), "Expected no errors during enumeration")
|
||||
|
||||
for _, unit := range reporter.Units {
|
||||
id, _ := unit.SourceUnitID()
|
||||
assert.NotEmpty(t, id, "Unit ID should not be empty")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSource_ChunkUnit(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
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")
|
||||
|
||||
connection := &sourcespb.S3{
|
||||
Credential: &sourcespb.S3_AccessKey{
|
||||
AccessKey: &credentialspb.KeySecret{
|
||||
Key: s3key,
|
||||
Secret: s3secret,
|
||||
},
|
||||
},
|
||||
Buckets: []string{"truffletestbucket"},
|
||||
}
|
||||
|
||||
conn, err := anypb.New(connection)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
s := Source{}
|
||||
err = s.Init(ctx, "test enumerate", 0, 0, false, conn, 1)
|
||||
assert.NoError(t, err)
|
||||
|
||||
reporter := sourcestest.TestReporter{}
|
||||
err = s.Enumerate(ctx, &reporter)
|
||||
assert.NoError(t, err)
|
||||
|
||||
for _, unit := range reporter.Units {
|
||||
err = s.ChunkUnit(ctx, unit, &reporter)
|
||||
assert.NoError(t, err, "Expected no error during ChunkUnit")
|
||||
}
|
||||
|
||||
assert.Equal(t, 103, len(reporter.Chunks))
|
||||
assert.Equal(t, 0, len(reporter.ChunkErrs))
|
||||
}
|
||||
|
||||
func TestSource_ChunkUnit_Resumption(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second*30)
|
||||
defer cancel()
|
||||
|
||||
s := new(Source)
|
||||
s.Progress = sources.Progress{
|
||||
Message: "Bucket: integration-resumption-tests",
|
||||
EncodedResumeInfo: "{\"integration-resumption-tests\":\"test-dir/\"}",
|
||||
SectionsCompleted: 0,
|
||||
SectionsRemaining: 1,
|
||||
}
|
||||
connection := &sourcespb.S3{
|
||||
Credential: &sourcespb.S3_Unauthenticated{},
|
||||
Buckets: []string{"integration-resumption-tests"},
|
||||
EnableResumption: true,
|
||||
}
|
||||
conn, err := anypb.New(connection)
|
||||
require.NoError(t, err)
|
||||
|
||||
err = s.Init(ctx, "test name", 0, 0, false, conn, 2)
|
||||
require.NoError(t, err)
|
||||
|
||||
reporter := sourcestest.TestReporter{}
|
||||
err = s.Enumerate(ctx, &reporter)
|
||||
assert.NoError(t, err)
|
||||
|
||||
for _, unit := range reporter.Units {
|
||||
err = s.ChunkUnit(ctx, unit, &reporter)
|
||||
assert.NoError(t, err, "Expected no error during ChunkUnit")
|
||||
}
|
||||
|
||||
// Verify that we processed all remaining data on resume.
|
||||
assert.Equal(t, 9638, len(reporter.Chunks), "Should have processed all remaining data on resume")
|
||||
}
|
||||
|
||||
// TestSource_ChunkUnit_Resumption_MultipleBucketsConcurrent tests resumption across multiple buckets
|
||||
// with concurrent ChunkUnit processing.
|
||||
func TestSource_ChunkUnit_Resumption_MultipleBucketsConcurrent(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
src := new(Source)
|
||||
connection := &sourcespb.S3{
|
||||
Credential: &sourcespb.S3_Unauthenticated{},
|
||||
Buckets: []string{"integration-resumption-tests", "trufflesec-ahrav-test-2"},
|
||||
EnableResumption: true,
|
||||
}
|
||||
conn, err := anypb.New(connection)
|
||||
require.NoError(t, err)
|
||||
err = src.Init(ctx, "test name", 0, 0, false, conn, 2)
|
||||
require.NoError(t, err)
|
||||
|
||||
src.Progress = sources.Progress{
|
||||
Message: "Buckets: [integration-resumption-tests trufflesec-ahrav-test-2]",
|
||||
EncodedResumeInfo: "{\"integration-resumption-tests\":\"test-dir\", \"trufflesec-ahrav-test-2\":\"test-dir/smmed_random_data.json.zip\"}",
|
||||
SectionsCompleted: 0,
|
||||
SectionsRemaining: 2,
|
||||
}
|
||||
|
||||
reporter := sourcestest.SafeTestReporter{}
|
||||
err = src.Enumerate(ctx, &reporter)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 2, len(reporter.Units), "Expected two source units from enumeration")
|
||||
|
||||
var wg sync.WaitGroup
|
||||
|
||||
for _, unit := range reporter.Units {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
err = src.ChunkUnit(ctx, unit, &reporter)
|
||||
assert.NoError(t, err, "Expected no error during ChunkUnit")
|
||||
}()
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
|
||||
bucketChunkCounts := map[string]int{
|
||||
"integration-resumption-tests": 9638,
|
||||
"trufflesec-ahrav-test-2": 2,
|
||||
}
|
||||
actualBucketChunkCounts := make(map[string]int)
|
||||
|
||||
for _, chunk := range reporter.Chunks {
|
||||
metadata, _ := chunk.SourceMetadata.Data.(*source_metadatapb.MetaData_S3)
|
||||
actualBucketChunkCounts[metadata.S3.Bucket]++
|
||||
}
|
||||
|
||||
for bucket, wantCount := range bucketChunkCounts {
|
||||
gotCount, ok := actualBucketChunkCounts[bucket]
|
||||
require.True(t, ok, "Expected chunks for bucket %s", bucket)
|
||||
assert.Equal(t, wantCount, gotCount, "Chunk count mismatch for bucket %s", bucket)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
|
||||
"github.com/kylelemons/godebug/pretty"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/protobuf/types/known/anypb"
|
||||
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/common"
|
||||
@@ -163,3 +164,21 @@ func TestSource_Chunks(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSource_UnmarshalSourceUnit(t *testing.T) {
|
||||
s := Source{}
|
||||
|
||||
unitJSON := `{
|
||||
"Bucket": "my-test-bucket",
|
||||
"Role": "my-test-role"
|
||||
}`
|
||||
|
||||
unit, err := s.UnmarshalSourceUnit([]byte(unitJSON))
|
||||
require.NoError(t, err, "UnmarshalSourceUnit should not return an error")
|
||||
|
||||
s3Unit, ok := unit.(S3SourceUnit)
|
||||
require.True(t, ok, "Unmarshaled unit should be of type S3SourceUnit")
|
||||
|
||||
assert.Equal(t, "my-test-bucket", s3Unit.Bucket, "Bucket field should match")
|
||||
assert.Equal(t, "my-test-role", s3Unit.Role, "Role field should match")
|
||||
}
|
||||
|
||||
@@ -0,0 +1,36 @@
|
||||
package s3
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/sources"
|
||||
)
|
||||
|
||||
const SourceUnitKindBucket sources.SourceUnitKind = "bucket"
|
||||
|
||||
type S3SourceUnit struct {
|
||||
Bucket string
|
||||
Role string
|
||||
}
|
||||
|
||||
var _ sources.SourceUnit = S3SourceUnit{}
|
||||
|
||||
func (s S3SourceUnit) SourceUnitID() (string, sources.SourceUnitKind) {
|
||||
// ID is a combination of bucket and role (if any)
|
||||
return constructS3SourceUnitID(s.Bucket, s.Role), SourceUnitKindBucket
|
||||
}
|
||||
|
||||
func (s S3SourceUnit) Display() string {
|
||||
if s.Role != "" {
|
||||
return fmt.Sprintf("Role=%s Bucket=%s", s.Role, s.Bucket)
|
||||
}
|
||||
return s.Bucket
|
||||
}
|
||||
|
||||
func constructS3SourceUnitID(bucket string, role string) string {
|
||||
unitID := ""
|
||||
if role != "" {
|
||||
unitID += role + "|"
|
||||
}
|
||||
return unitID + bucket
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
package s3
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestS3Unit(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
bucket string
|
||||
role string
|
||||
wantID string
|
||||
wantDisplay string
|
||||
}{
|
||||
{
|
||||
name: "Bucket with role",
|
||||
bucket: "my-bucket",
|
||||
role: "arn:aws:iam::123456789012:role/MyRole",
|
||||
wantID: "arn:aws:iam::123456789012:role/MyRole|my-bucket",
|
||||
wantDisplay: "Role=arn:aws:iam::123456789012:role/MyRole Bucket=my-bucket",
|
||||
},
|
||||
{
|
||||
name: "Bucket without role",
|
||||
bucket: "my-bucket",
|
||||
wantID: "my-bucket",
|
||||
wantDisplay: "my-bucket",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
unit := S3SourceUnit{
|
||||
Bucket: tt.bucket,
|
||||
Role: tt.role,
|
||||
}
|
||||
|
||||
gotID, gotKind := unit.SourceUnitID()
|
||||
gotDisplay := unit.Display()
|
||||
|
||||
if gotKind != SourceUnitKindBucket {
|
||||
t.Errorf("SourceUnitID() got kind = %v, want %v", gotKind, SourceUnitKindBucket)
|
||||
}
|
||||
if gotID != tt.wantID {
|
||||
t.Errorf("SourceUnitID() got id= %v, want %v", gotID, tt.wantID)
|
||||
}
|
||||
if gotDisplay != tt.wantDisplay {
|
||||
t.Errorf("Display() = %v, want %v", gotDisplay, tt.wantDisplay)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -2,6 +2,7 @@ package sourcestest
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/context"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/sources"
|
||||
@@ -59,3 +60,32 @@ func (ErrReporter) ChunkOk(context.Context, sources.Chunk) error {
|
||||
func (ErrReporter) ChunkErr(context.Context, error) error {
|
||||
return fmt.Errorf("ErrReporter: ChunkErr error")
|
||||
}
|
||||
|
||||
// SafeTestReporter is a helper struct that implements both UnitReporter and
|
||||
// ChunkReporter by recording the values passed in the methods with thread safety.
|
||||
type SafeTestReporter struct {
|
||||
TestReporter
|
||||
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func (t *SafeTestReporter) UnitOk(_ context.Context, unit sources.SourceUnit) error {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
return t.TestReporter.UnitOk(nil, unit)
|
||||
}
|
||||
func (t *SafeTestReporter) UnitErr(_ context.Context, err error) error {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
return t.TestReporter.UnitErr(nil, err)
|
||||
}
|
||||
func (t *SafeTestReporter) ChunkOk(_ context.Context, chunk sources.Chunk) error {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
return t.TestReporter.ChunkOk(nil, chunk)
|
||||
}
|
||||
func (t *SafeTestReporter) ChunkErr(_ context.Context, err error) error {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
return t.TestReporter.ChunkErr(nil, err)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user