Files
trufflehog/pkg/sources/s3/s3.go
ahrav 7b3d98d133 [feat] - S3 metrics (#3577)
* 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

* add metrics to s3 scan

* make collector a singleton

* address comments

* fix

* remove
2024-11-26 10:52:02 -08:00

691 lines
22 KiB
Go

package s3
import (
"fmt"
"slices"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/aws/aws-sdk-go/aws"
"github.com/aws/aws-sdk-go/aws/credentials"
"github.com/aws/aws-sdk-go/aws/credentials/stscreds"
"github.com/aws/aws-sdk-go/aws/session"
"github.com/aws/aws-sdk-go/service/s3"
"github.com/aws/aws-sdk-go/service/s3/s3manager"
"github.com/aws/aws-sdk-go/service/sts"
"github.com/go-errors/errors"
"golang.org/x/sync/errgroup"
"google.golang.org/protobuf/proto"
"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/handlers"
"github.com/trufflesecurity/trufflehog/v3/pkg/log"
"github.com/trufflesecurity/trufflehog/v3/pkg/pb/source_metadatapb"
"github.com/trufflesecurity/trufflehog/v3/pkg/pb/sourcespb"
"github.com/trufflesecurity/trufflehog/v3/pkg/sanitizer"
"github.com/trufflesecurity/trufflehog/v3/pkg/sources"
)
const (
SourceType = sourcespb.SourceType_SOURCE_TYPE_S3
defaultAWSRegion = "us-east-1"
defaultMaxObjectSize = 250 * 1024 * 1024 // 250 MiB
maxObjectSizeLimit = 250 * 1024 * 1024 // 250 MiB
)
type Source struct {
name string
sourceID sources.SourceID
jobID sources.JobID
verify bool
concurrency int
conn *sourcespb.S3
checkpointer *Checkpointer
sources.Progress
metricsCollector metricsCollector
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)
// Type returns the type of source
func (s *Source) Type() sourcespb.SourceType { return SourceType }
func (s *Source) SourceID() sources.SourceID { return s.sourceID }
func (s *Source) JobID() sources.JobID { return s.jobID }
// Init returns an initialized AWS source
func (s *Source) Init(
ctx context.Context,
name string,
jobID sources.JobID,
sourceID sources.SourceID,
verify bool,
connection *anypb.Any,
concurrency int,
) error {
s.name = name
s.sourceID = sourceID
s.jobID = jobID
s.verify = verify
s.concurrency = concurrency
s.errorCount = &sync.Map{}
s.jobPool = &errgroup.Group{}
s.jobPool.SetLimit(concurrency)
var conn sourcespb.S3
if err := anypb.UnmarshalTo(connection, &conn, proto.UnmarshalOptions{}); err != nil {
return fmt.Errorf("error unmarshalling connection: %w", err)
}
s.conn = &conn
s.checkpointer = NewCheckpointer(ctx, conn.GetEnableResumption(), &s.Progress)
s.metricsCollector = metricsInstance
s.setMaxObjectSize(conn.GetMaxObjectSize())
if len(conn.GetBuckets()) > 0 && len(conn.GetIgnoreBuckets()) > 0 {
return errors.New("either a bucket include list or a bucket ignore list can be specified, but not both")
}
return nil
}
func (s *Source) Validate(ctx context.Context) []error {
var errs []error
visitor := func(c context.Context, defaultRegionClient *s3.S3, roleArn string, buckets []string) error {
roleErrs := s.validateBucketAccess(c, defaultRegionClient, roleArn, buckets)
if len(roleErrs) > 0 {
errs = append(errs, roleErrs...)
}
return nil
}
if err := s.visitRoles(ctx, visitor); err != nil {
errs = append(errs, err)
}
return errs
}
// setMaxObjectSize sets the maximum size of objects that will be scanned. If
// not set, set to a negative number, or set larger than the
// maxObjectSizeLimit, the defaultMaxObjectSizeLimit will be used.
func (s *Source) setMaxObjectSize(maxObjectSize int64) {
if maxObjectSize <= 0 || maxObjectSize > maxObjectSizeLimit {
s.maxObjectSize = defaultMaxObjectSize
} else {
s.maxObjectSize = maxObjectSize
}
}
func (s *Source) newClient(region, roleArn string) (*s3.S3, error) {
cfg := aws.NewConfig()
cfg.CredentialsChainVerboseErrors = aws.Bool(true)
cfg.Region = aws.String(region)
switch cred := s.conn.GetCredential().(type) {
case *sourcespb.S3_SessionToken:
cfg.Credentials = credentials.NewStaticCredentials(
cred.SessionToken.GetKey(),
cred.SessionToken.GetSecret(),
cred.SessionToken.GetSessionToken(),
)
log.RedactGlobally(cred.SessionToken.GetSecret())
log.RedactGlobally(cred.SessionToken.GetSessionToken())
case *sourcespb.S3_AccessKey:
cfg.Credentials = credentials.NewStaticCredentials(cred.AccessKey.GetKey(), cred.AccessKey.GetSecret(), "")
log.RedactGlobally(cred.AccessKey.GetSecret())
case *sourcespb.S3_Unauthenticated:
cfg.Credentials = credentials.AnonymousCredentials
default:
// In all other cases, the AWS SDK will follow its normal waterfall logic to pick up credentials (i.e. they can
// come from the environment or the credentials file or whatever else AWS gets up to).
}
if roleArn != "" {
sess, err := session.NewSession(cfg)
if err != nil {
return nil, err
}
stsClient := sts.New(sess)
cfg.Credentials = stscreds.NewCredentialsWithClient(stsClient, roleArn, func(p *stscreds.AssumeRoleProvider) {
p.RoleSessionName = "trufflehog"
})
}
sess, err := session.NewSessionWithOptions(session.Options{
SharedConfigState: session.SharedConfigEnable,
Config: *cfg,
})
if err != nil {
return nil, err
}
return s3.New(sess), nil
}
// getBucketsToScan returns a list of S3 buckets to scan.
// If the connection has a list of buckets specified, those are returned.
// Otherwise, it lists all buckets the client has access to and filters out the ignored ones.
// The list of buckets is sorted lexicographically to ensure consistent ordering,
// which allows resuming scanning from the same place if the scan is interrupted.
//
// Note: The IAM identity needs the s3:ListBuckets permission.
func (s *Source) getBucketsToScan(client *s3.S3) ([]string, error) {
if buckets := s.conn.GetBuckets(); len(buckets) > 0 {
slices.Sort(buckets)
return buckets, nil
}
ignore := make(map[string]struct{}, len(s.conn.GetIgnoreBuckets()))
for _, bucket := range s.conn.GetIgnoreBuckets() {
ignore[bucket] = struct{}{}
}
res, err := client.ListBuckets(&s3.ListBucketsInput{})
if err != nil {
return nil, err
}
var bucketsToScan []string
for _, bucket := range res.Buckets {
name := *bucket.Name
if _, ignored := ignore[name]; !ignored {
bucketsToScan = append(bucketsToScan, name)
}
}
slices.Sort(bucketsToScan)
return bucketsToScan, nil
}
// 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
pageNumber int // Current page number in the pagination sequence
client *s3.S3 // AWS S3 client configured for the appropriate region
page *s3.ListObjectsV2Output // Contains the list of S3 objects in this page
}
// processingState tracks the state of concurrent S3 object processing.
type processingState struct {
errorCount *sync.Map // Thread-safe map tracking errors per prefix
objectCount *uint64 // Total number of objects processed
}
// resumePosition tracks where to restart scanning S3 buckets and objects after an interruption.
// It encapsulates all the information needed to resume a scan from its last known position.
type resumePosition struct {
bucket string // The bucket name we were processing
index int // Index in the buckets slice where we should resume
startAfter string // The last processed object key within the bucket
isNewScan bool // True if we're starting a fresh scan
exactMatch bool // True if we found the exact bucket we were previously processing
}
// determineResumePosition calculates where to resume scanning from based on the last saved checkpoint
// and the current list of available buckets to scan. It handles several scenarios:
//
// 1. If getting the resume point fails or there is no previous bucket saved (CurrentBucket is empty),
// we start a new scan from the beginning, this is the safest option.
//
// 2. If the previous bucket exists in our current scan list (exactMatch=true),
// we resume from that exact position and use the StartAfter value
// to continue from the last processed object within that bucket.
//
// 3. If the previous bucket is not found in our current scan list (exactMatch=false), this typically means:
// - The bucket was deleted since our last scan
// - The bucket was explicitly excluded from this scan's configuration
// - The IAM role no longer has access to the bucket
// - The bucket name changed due to a configuration update
// In this case, we use binary search to find the closest position where the bucket would have been,
// allowing us to resume from the nearest available point in our sorted bucket list rather than
// restarting the entire scan.
func determineResumePosition(ctx context.Context, tracker *Checkpointer, buckets []string) resumePosition {
resumePoint, err := tracker.ResumePoint(ctx)
if err != nil {
ctx.Logger().Error(err, "failed to get resume point; starting from the beginning")
return resumePosition{isNewScan: true}
}
if resumePoint.CurrentBucket == "" {
return resumePosition{isNewScan: true}
}
startIdx, found := slices.BinarySearch(buckets, resumePoint.CurrentBucket)
return resumePosition{
bucket: resumePoint.CurrentBucket,
startAfter: resumePoint.StartAfter,
index: startIdx,
exactMatch: found,
}
}
func (s *Source) scanBuckets(
ctx context.Context,
client *s3.S3,
role string,
bucketsToScan []string,
chunksChan chan *sources.Chunk,
) {
if role != "" {
ctx = context.WithValue(ctx, "role", role)
}
var objectCount uint64
pos := determineResumePosition(ctx, s.checkpointer, bucketsToScan)
switch {
case pos.isNewScan:
ctx.Logger().Info("Starting new scan from beginning")
case !pos.exactMatch:
ctx.Logger().Info(
"Resume bucket no longer available, starting from closest position",
"original_bucket", pos.bucket,
"position", pos.index,
)
default:
ctx.Logger().Info(
"Resuming scan from previous scan's bucket",
"bucket", pos.bucket,
"position", pos.index,
)
}
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,
len(bucketsToScan),
fmt.Sprintf("Bucket: %s", bucket),
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}
if bucket == pos.bucket && pos.startAfter != "" {
input.StartAfter = &pos.startAfter
ctx.Logger().V(3).Info(
"Resuming bucket scan",
"start_after", pos.startAfter,
)
}
pageNumber := 1
err = regionalClient.ListObjectsV2PagesWithContext(
ctx,
input,
func(page *s3.ListObjectsV2Output, _ bool) bool {
pageMetadata := pageMetadata{
bucket: bucket,
pageNumber: pageNumber,
client: regionalClient,
page: page,
}
processingState := processingState{
errorCount: &errorCount,
objectCount: &objectCount,
}
s.pageChunker(ctx, pageMetadata, processingState, chunksChan)
pageNumber++
return true
})
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)
}
}
}
s.SetProgressComplete(
len(bucketsToScan),
len(bucketsToScan),
fmt.Sprintf("Completed scanning source %s. %d objects scanned.", s.name, 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.S3, roleArn string, buckets []string) error {
s.scanBuckets(c, defaultRegionClient, roleArn, buckets, chunksChan)
return nil
}
return s.visitRoles(ctx, visitor)
}
func (s *Source) getRegionalClientForBucket(
ctx context.Context,
defaultRegionClient *s3.S3,
role string,
bucket string,
) (*s3.S3, error) {
region, err := s3manager.GetBucketRegionWithClient(ctx, defaultRegionClient, bucket)
if err != nil {
return nil, fmt.Errorf("could not get s3 region for bucket: %s", bucket)
}
if region == defaultAWSRegion {
return defaultRegionClient, nil
}
regionalClient, err := s.newClient(region, role)
if err != nil {
return nil, fmt.Errorf("could not create regional s3 client for bucket %s: %w", bucket, err)
}
return regionalClient, nil
}
// pageChunker emits chunks onto the given channel from a page.
func (s *Source) pageChunker(
ctx context.Context,
metadata pageMetadata,
state processingState,
chunksChan chan *sources.Chunk,
) {
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 {
if obj == nil {
s.metricsCollector.RecordObjectSkipped(metadata.bucket, "nil_object")
if err := s.checkpointer.UpdateObjectCompletion(ctx, objIdx, metadata.bucket, metadata.page.Contents); err != nil {
ctx.Logger().Error(err, "could not update progress for nil object")
}
continue
}
ctx = context.WithValues(ctx, "key", *obj.Key, "size", *obj.Size)
if common.IsDone(ctx) {
return
}
// Skip GLACIER and GLACIER_IR objects.
if obj.StorageClass == nil || strings.Contains(*obj.StorageClass, "GLACIER") {
ctx.Logger().V(5).Info("Skipping object in storage class", "storage_class", *obj.StorageClass)
s.metricsCollector.RecordObjectSkipped(metadata.bucket, "storage_class")
if err := s.checkpointer.UpdateObjectCompletion(ctx, objIdx, metadata.bucket, metadata.page.Contents); err != nil {
ctx.Logger().Error(err, "could not update progress for glacier object")
}
continue
}
// Ignore large files.
if *obj.Size > s.maxObjectSize {
ctx.Logger().V(5).Info("Skipping %d byte file (over maxObjectSize limit)")
s.metricsCollector.RecordObjectSkipped(metadata.bucket, "size_limit")
if err := s.checkpointer.UpdateObjectCompletion(ctx, objIdx, metadata.bucket, metadata.page.Contents); err != nil {
ctx.Logger().Error(err, "could not update progress for large file")
}
continue
}
// File empty file.
if *obj.Size == 0 {
ctx.Logger().V(5).Info("Skipping empty file")
s.metricsCollector.RecordObjectSkipped(metadata.bucket, "empty_file")
if err := s.checkpointer.UpdateObjectCompletion(ctx, objIdx, metadata.bucket, metadata.page.Contents); err != nil {
ctx.Logger().Error(err, "could not update progress for empty file")
}
continue
}
// Skip incompatible extensions.
if common.SkipFile(*obj.Key) {
ctx.Logger().V(5).Info("Skipping file with incompatible extension")
s.metricsCollector.RecordObjectSkipped(metadata.bucket, "incompatible_extension")
if err := s.checkpointer.UpdateObjectCompletion(ctx, objIdx, metadata.bucket, metadata.page.Contents); err != nil {
ctx.Logger().Error(err, "could not update progress for incompatible file")
}
continue
}
s.jobPool.Go(func() error {
defer common.RecoverWithExit(ctx)
if common.IsDone(ctx) {
return ctx.Err()
}
if strings.HasSuffix(*obj.Key, "/") {
ctx.Logger().V(5).Info("Skipping directory")
s.metricsCollector.RecordObjectSkipped(metadata.bucket, "directory")
return nil
}
path := strings.Split(*obj.Key, "/")
prefix := strings.Join(path[:len(path)-1], "/")
nErr, ok := state.errorCount.Load(prefix)
if !ok {
nErr = 0
}
if nErr.(int) > 3 {
ctx.Logger().V(2).Info("Skipped due to excessive errors")
return nil
}
// Make sure we use a separate context for the GetObjectWithContext call.
// This ensures that the timeout is isolated and does not affect any downstream operations. (e.g. HandleFile)
const getObjectTimeout = 30 * time.Second
objCtx, cancel := context.WithTimeout(ctx, getObjectTimeout)
defer cancel()
res, err := metadata.client.GetObjectWithContext(objCtx, &s3.GetObjectInput{
Bucket: &metadata.bucket,
Key: obj.Key,
})
if err != nil {
if strings.Contains(err.Error(), "AccessDenied") {
ctx.Logger().Error(err, "could not get S3 object; access denied")
s.metricsCollector.RecordObjectSkipped(metadata.bucket, "access_denied")
} else {
ctx.Logger().Error(err, "could not get S3 object")
s.metricsCollector.RecordObjectError(metadata.bucket)
}
// According to the documentation for GetObjectWithContext,
// the response can be non-nil even if there was an error.
// It's uncertain if the body will be nil in such cases,
// but we'll close it if it's not.
if res != nil && res.Body != nil {
res.Body.Close()
}
nErr, ok := state.errorCount.Load(prefix)
if !ok {
nErr = 0
}
if nErr.(int) > 3 {
ctx.Logger().V(3).Info("Skipped due to excessive errors")
return nil
}
nErr = nErr.(int) + 1
state.errorCount.Store(prefix, nErr)
// too many consecutive errors on this page
if nErr.(int) > 3 {
ctx.Logger().V(2).Info("Too many consecutive errors, excluding prefix", "prefix", prefix)
}
return nil
}
defer res.Body.Close()
email := "Unknown"
if obj.Owner != nil {
email = *obj.Owner.DisplayName
}
modified := obj.LastModified.String()
chunkSkel := &sources.Chunk{
SourceType: s.Type(),
SourceName: s.name,
SourceID: s.SourceID(),
JobID: s.JobID(),
SourceMetadata: &source_metadatapb.MetaData{
Data: &source_metadatapb.MetaData_S3{
S3: &source_metadatapb.S3{
Bucket: metadata.bucket,
File: sanitizer.UTF8(*obj.Key),
Link: sanitizer.UTF8(makeS3Link(metadata.bucket, *metadata.client.Config.Region, *obj.Key)),
Email: sanitizer.UTF8(email),
Timestamp: sanitizer.UTF8(modified),
},
},
},
Verify: s.verify,
}
if err := handlers.HandleFile(ctx, res.Body, chunkSkel, sources.ChanReporter{Ch: chunksChan}); 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)
if !ok {
nErr = 0
}
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 {
ctx.Logger().Error(err, "could not update progress for scanned object")
}
s.metricsCollector.RecordObjectScanned(metadata.bucket)
return nil
})
}
_ = s.jobPool.Wait()
}
func (s *Source) validateBucketAccess(ctx context.Context, client *s3.S3, roleArn string, buckets []string) []error {
shouldHaveAccessToAllBuckets := roleArn == ""
wasAbleToListAnyBucket := false
var errs []error
for _, bucket := range buckets {
if common.IsDone(ctx) {
return append(errs, ctx.Err())
}
regionalClient, err := s.getRegionalClientForBucket(ctx, client, roleArn, bucket)
if err != nil {
errs = append(errs, fmt.Errorf("could not get regional client for bucket %q: %w", bucket, err))
continue
}
_, err = regionalClient.ListObjectsV2(&s3.ListObjectsV2Input{Bucket: &bucket})
if err == nil {
wasAbleToListAnyBucket = true
} else if shouldHaveAccessToAllBuckets {
errs = append(errs, fmt.Errorf("could not list objects in bucket %q: %w", bucket, err))
}
}
if !wasAbleToListAnyBucket {
if roleArn == "" {
errs = append(errs, errors.New("could not list objects in any bucket"))
} else {
errs = append(errs, fmt.Errorf("role %q could not list objects in any bucket", roleArn))
}
}
return errs
}
// visitRoles iterates over the configured AWS roles and calls the provided function
// for each role, passing in the default S3 client, the role ARN, and the list of
// buckets to scan.
//
// The provided function parameter typically implements the core scanning logic
// and must handle context cancellation appropriately.
//
// If no roles are configured, it will call the function with an empty role ARN.
func (s *Source) visitRoles(
ctx context.Context,
f func(c context.Context, defaultRegionClient *s3.S3, roleArn string, buckets []string) error,
) error {
roles := s.conn.GetRoles()
if len(roles) == 0 {
roles = []string{""}
}
for _, role := range roles {
s.metricsCollector.RecordRoleScanned(role)
client, err := s.newClient(defaultAWSRegion, role)
if err != nil {
return fmt.Errorf("could not create s3 client: %w", err)
}
bucketsToScan, err := s.getBucketsToScan(client)
if err != nil {
return fmt.Errorf("role %q could not list any s3 buckets for scanning: %w", role, err)
}
if err := f(ctx, client, role, bucketsToScan); err != nil {
return err
}
}
return nil
}
// S3 links currently have the general format of:
// https://[bucket].s3[.region unless us-east-1].amazonaws.com/[key]
func makeS3Link(bucket, region, key string) string {
if region == defaultAWSRegion {
region = ""
} else {
region = "." + region
}
return fmt.Sprintf("https://%s.s3%s.amazonaws.com/%s", bucket, region, key)
}