[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:
Mustansir
2025-12-04 12:54:28 +05:00
committed by GitHub
co-authored by Amaan Ullah
parent 1935692d46
commit b6389e2419
8 changed files with 542 additions and 72 deletions
+27 -3
View File
@@ -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
}
+33 -3
View File
@@ -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
View File
@@ -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
}
+193
View File
@@ -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)
}
}
+19
View File
@@ -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")
}
+36
View File
@@ -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
}
+51
View File
@@ -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)
}
})
}
}
+30
View File
@@ -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)
}