Merge branch 'archive-entry-paths' of https://github.com/trufflesecurity/trufflehog into archive-entry-paths
This commit is contained in:
@@ -88,7 +88,7 @@ require (
|
||||
github.com/sergi/go-diff v1.3.2-0.20230802210424-5b0b94c5c0d3
|
||||
github.com/shuheiktgw/go-travis v0.3.1
|
||||
github.com/snowflakedb/gosnowflake v1.12.0
|
||||
github.com/stretchr/testify v1.9.0
|
||||
github.com/stretchr/testify v1.10.0
|
||||
github.com/tailscale/depaware v0.0.0-20241028160002-3d7f3b30ed0e
|
||||
github.com/testcontainers/testcontainers-go v0.34.0
|
||||
github.com/testcontainers/testcontainers-go/modules/elasticsearch v0.34.0
|
||||
@@ -110,7 +110,7 @@ require (
|
||||
golang.org/x/oauth2 v0.24.0
|
||||
golang.org/x/sync v0.9.0
|
||||
golang.org/x/text v0.20.0
|
||||
google.golang.org/api v0.208.0
|
||||
google.golang.org/api v0.209.0
|
||||
google.golang.org/protobuf v1.35.2
|
||||
gopkg.in/h2non/gock.v1 v1.1.2
|
||||
gopkg.in/yaml.v2 v2.4.0
|
||||
|
||||
@@ -712,6 +712,8 @@ github.com/stretchr/testify v1.8.2/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o
|
||||
github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo=
|
||||
github.com/stretchr/testify v1.9.0 h1:HtqpIVDClZ4nwg75+f6Lvsy/wHu+3BoSGCbBAcpTsTg=
|
||||
github.com/stretchr/testify v1.9.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
|
||||
github.com/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOfJA=
|
||||
github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
|
||||
github.com/tailscale/depaware v0.0.0-20241028160002-3d7f3b30ed0e h1:Hb50wYyy5VblH5zpKkoy49TrJy3pxVWOaRSOEdzTWKA=
|
||||
github.com/tailscale/depaware v0.0.0-20241028160002-3d7f3b30ed0e/go.mod h1:p9lPsd+cx33L3H9nNoecRRxPssFKUwwI50I3pZ0yT+8=
|
||||
github.com/testcontainers/testcontainers-go v0.34.0 h1:5fbgF0vIN5u+nD3IWabQwRybuB4GY8G2HHgCkbMzMHo=
|
||||
@@ -1067,6 +1069,8 @@ google.golang.org/api v0.207.0 h1:Fvt6IGCYjf7YLcQ+GCegeAI2QSQCfIWhRkmrMPj3JRM=
|
||||
google.golang.org/api v0.207.0/go.mod h1:I53S168Yr/PNDNMi5yPnDc0/LGRZO6o7PoEbl/HY3CM=
|
||||
google.golang.org/api v0.208.0 h1:8Y62MUGRviQnnP9/41/bYAGySPKAN9iwzV96ZvhwyVE=
|
||||
google.golang.org/api v0.208.0/go.mod h1:I53S168Yr/PNDNMi5yPnDc0/LGRZO6o7PoEbl/HY3CM=
|
||||
google.golang.org/api v0.209.0 h1:Ja2OXNlyRlWCWu8o+GgI4yUn/wz9h/5ZfFbKz+dQX+w=
|
||||
google.golang.org/api v0.209.0/go.mod h1:I53S168Yr/PNDNMi5yPnDc0/LGRZO6o7PoEbl/HY3CM=
|
||||
google.golang.org/appengine v1.1.0/go.mod h1:EbEs0AVv82hx2wNQdGPgUI5lhzA/G0D9YwlJXL52JkM=
|
||||
google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4=
|
||||
google.golang.org/appengine v1.5.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4=
|
||||
|
||||
@@ -138,8 +138,9 @@ func verifyMatch(ctx context.Context, appId, apiKey string) (bool, map[string]st
|
||||
case http.StatusUnauthorized:
|
||||
return false, nil, nil
|
||||
case http.StatusForbidden:
|
||||
// Key is valid but lacks permissions.
|
||||
return true, nil, nil
|
||||
// Invalidated key.
|
||||
// {"message":"Invalid Application-ID or API key","status":403}
|
||||
return false, nil, nil
|
||||
default:
|
||||
return false, nil, fmt.Errorf("unexpected HTTP response status %d", res.StatusCode)
|
||||
}
|
||||
|
||||
@@ -2,7 +2,9 @@ package buildkite
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
@@ -15,6 +17,10 @@ import (
|
||||
|
||||
type Scanner struct{}
|
||||
|
||||
type APIResponse struct {
|
||||
Scopes []string `json:"scopes"`
|
||||
}
|
||||
|
||||
func (s Scanner) Version() int { return 1 }
|
||||
|
||||
// Ensure the Scanner satisfies the interface at compile time.
|
||||
@@ -49,21 +55,15 @@ func (s Scanner) FromData(ctx context.Context, verify bool, data []byte) (result
|
||||
s1 := detectors.Result{
|
||||
DetectorType: detectorspb.DetectorType_Buildkite,
|
||||
Raw: []byte(resMatch),
|
||||
ExtraData: make(map[string]string),
|
||||
}
|
||||
|
||||
if verify {
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", "https://api.buildkite.com/v2/access-token", nil)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
req.Header.Add("Authorization", fmt.Sprintf("Bearer %s", resMatch))
|
||||
res, err := client.Do(req)
|
||||
if err == nil {
|
||||
defer res.Body.Close()
|
||||
if res.StatusCode >= 200 && res.StatusCode < 300 {
|
||||
s1.Verified = true
|
||||
}
|
||||
}
|
||||
extraData, isVerified, verificationErr := VerifyBuildKite(ctx, client, resMatch)
|
||||
s1.Verified = isVerified
|
||||
s1.SetVerificationError(verificationErr, resMatch)
|
||||
|
||||
s1.ExtraData = extraData
|
||||
}
|
||||
|
||||
results = append(results, s1)
|
||||
@@ -79,3 +79,42 @@ func (s Scanner) Type() detectorspb.DetectorType {
|
||||
func (s Scanner) Description() string {
|
||||
return "Buildkite is a platform for running fast, secure, and scalable continuous integration pipelines. Buildkite API tokens can be used to access and modify pipeline data and configurations."
|
||||
}
|
||||
|
||||
func VerifyBuildKite(ctx context.Context, client *http.Client, secret string) (map[string]string, bool, error) {
|
||||
// create a request
|
||||
// api doc: https://buildkite.com/docs/apis/rest-api/access-token#get-the-current-token
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", "https://api.buildkite.com/v2/access-token", nil)
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
|
||||
// add authorization header
|
||||
req.Header.Add("Authorization", fmt.Sprintf("Bearer %s", secret))
|
||||
|
||||
res, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
defer func() {
|
||||
_, _ = io.Copy(io.Discard, res.Body)
|
||||
_ = res.Body.Close()
|
||||
}()
|
||||
|
||||
switch res.StatusCode {
|
||||
case http.StatusOK:
|
||||
var response APIResponse
|
||||
|
||||
if err := json.NewDecoder(res.Body).Decode(&response); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
|
||||
extraData := make(map[string]string)
|
||||
|
||||
extraData["scopes"] = strings.Join(response.Scopes, ", ")
|
||||
return extraData, true, nil
|
||||
case http.StatusUnauthorized:
|
||||
return nil, false, nil
|
||||
default:
|
||||
return nil, false, fmt.Errorf("unexpected HTTP response status %d", res.StatusCode)
|
||||
}
|
||||
}
|
||||
@@ -2,14 +2,13 @@ package buildkitev2
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
regexp "github.com/wasilibs/go-re2"
|
||||
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/common"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors"
|
||||
v1 "github.com/trufflesecurity/trufflehog/v3/pkg/detectors/buildkite/v1"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/pb/detectorspb"
|
||||
)
|
||||
|
||||
@@ -52,18 +51,10 @@ func (s Scanner) FromData(ctx context.Context, verify bool, data []byte) (result
|
||||
}
|
||||
|
||||
if verify {
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", "https://api.buildkite.com/v2/access-token", nil)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
req.Header.Add("Authorization", fmt.Sprintf("Bearer %s", resMatch))
|
||||
res, err := client.Do(req)
|
||||
if err == nil {
|
||||
defer res.Body.Close()
|
||||
if res.StatusCode >= 200 && res.StatusCode < 300 {
|
||||
s1.Verified = true
|
||||
}
|
||||
}
|
||||
extraData, isVerified, verificationErr := v1.VerifyBuildKite(ctx, client, resMatch)
|
||||
s1.Verified = isVerified
|
||||
s1.SetVerificationError(verificationErr, resMatch)
|
||||
s1.ExtraData = extraData
|
||||
}
|
||||
|
||||
results = append(results, s1)
|
||||
@@ -0,0 +1,140 @@
|
||||
package godaddy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
|
||||
regexp "github.com/wasilibs/go-re2"
|
||||
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/common"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/pb/detectorspb"
|
||||
)
|
||||
|
||||
type Scanner struct {
|
||||
client *http.Client
|
||||
}
|
||||
|
||||
var (
|
||||
// ensure the scanner satisfies the interface at compile time.
|
||||
_ detectors.Detector = (*Scanner)(nil)
|
||||
_ detectors.Versioner = (*Scanner)(nil)
|
||||
|
||||
defaultClient = common.SaneHttpClient()
|
||||
|
||||
// the key for the GoDaddy OTE environment is a 37-character alphanumeric string that may include underscores.
|
||||
keyPattern = regexp.MustCompile(detectors.PrefixRegex([]string{"godaddy", "ote"}) + common.BuildRegex("a-zA-Z0-9", "_", 37))
|
||||
// the secret for the GoDaddy OTE environment is a 22-character alphanumeric string.
|
||||
secretPattern = regexp.MustCompile(detectors.PrefixRegex([]string{"godaddy", "ote"}) + common.BuildRegex("a-zA-Z0-9", "", 22))
|
||||
|
||||
// ote environment
|
||||
ote = "api.ote-godaddy.com"
|
||||
)
|
||||
|
||||
func (s *Scanner) getClient() *http.Client {
|
||||
if s.client != nil {
|
||||
return s.client
|
||||
}
|
||||
|
||||
return defaultClient
|
||||
}
|
||||
|
||||
func (s *Scanner) Version() int { return 1 }
|
||||
|
||||
// Keywords are used for efficiently pre-filtering chunks.
|
||||
// Use identifiers in the secret preferably, or the provider name.
|
||||
func (s Scanner) Keywords() []string {
|
||||
return []string{"godaddy", "ote"}
|
||||
}
|
||||
|
||||
func (s Scanner) Description() string {
|
||||
return "GoDaddy offers website building, hosting and security tools and services to construct, expand and protect the online presence." +
|
||||
"GoDaddy provides applications and access to relevant third-party products and platforms to connect their customers"
|
||||
}
|
||||
|
||||
func (s Scanner) Type() detectorspb.DetectorType {
|
||||
return detectorspb.DetectorType_GoDaddy
|
||||
}
|
||||
|
||||
// FromData will find and optionally verify GoDaddy API Key and secrets in a given set of bytes.
|
||||
func (s Scanner) FromData(ctx context.Context, verify bool, data []byte) (results []detectors.Result, err error) {
|
||||
// convert the data to string
|
||||
dataStr := string(data)
|
||||
|
||||
// find all the matching keys and secret in data and make a unique maps of both keys and secret.
|
||||
uniqueKeys, uniqueSecrets := make(map[string]struct{}), make(map[string]struct{})
|
||||
|
||||
for _, foundKey := range keyPattern.FindAllStringSubmatch(dataStr, -1) {
|
||||
uniqueKeys[foundKey[1]] = struct{}{}
|
||||
}
|
||||
|
||||
for _, foundSecret := range secretPattern.FindAllStringSubmatch(dataStr, -1) {
|
||||
uniqueSecrets[foundSecret[1]] = struct{}{}
|
||||
}
|
||||
|
||||
for key := range uniqueKeys {
|
||||
for secret := range uniqueSecrets {
|
||||
result := detectors.Result{
|
||||
DetectorType: detectorspb.DetectorType_GoDaddy,
|
||||
Raw: []byte(key),
|
||||
ExtraData: make(map[string]string),
|
||||
}
|
||||
|
||||
if verify {
|
||||
isVerified, verificationErr := VerifyGoDaddySecret(ctx, s.getClient(), ote, MakeAuthHeaderValue(key, secret))
|
||||
|
||||
result.Verified = isVerified
|
||||
result.SetVerificationError(verificationErr, secret)
|
||||
|
||||
// in case of successful verification add the enviorement name in extradata to let user know which env this secret belong to.
|
||||
if isVerified {
|
||||
result.ExtraData["Environment"] = "OTE"
|
||||
}
|
||||
}
|
||||
|
||||
results = append(results, result)
|
||||
}
|
||||
}
|
||||
|
||||
return results, nil
|
||||
|
||||
}
|
||||
|
||||
// VerifyGoDaddySecret make a call to godaddy api with given secret to check if secret is valid or not.
|
||||
func VerifyGoDaddySecret(ctx context.Context, client *http.Client, environment, secret string) (bool, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", fmt.Sprintf("https://%s/v1/domains/available?domain=example.com", environment), http.NoBody)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
// set the required auth header
|
||||
req.Header.Set("Authorization", secret)
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
defer func() {
|
||||
_, _ = io.Copy(io.Discard, resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
}()
|
||||
|
||||
switch resp.StatusCode {
|
||||
case http.StatusOK:
|
||||
return true, nil
|
||||
case http.StatusUnauthorized:
|
||||
return false, nil
|
||||
case http.StatusForbidden:
|
||||
// as per documentation in case of 403 the token is actually verified but it does not have access.
|
||||
return true, nil
|
||||
default:
|
||||
return false, fmt.Errorf("unexpected status code: %d", resp.StatusCode)
|
||||
}
|
||||
}
|
||||
|
||||
// MakeAuthHeaderValue return a value made from key and secret that can be used as authorization header value for godaddy API's.
|
||||
func MakeAuthHeaderValue(key, secret string) string {
|
||||
return fmt.Sprintf("sso-key %s:%s", key, secret)
|
||||
}
|
||||
@@ -0,0 +1,120 @@
|
||||
//go:build detectors
|
||||
// +build detectors
|
||||
|
||||
package godaddy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/kylelemons/godebug/pretty"
|
||||
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/common"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/pb/detectorspb"
|
||||
)
|
||||
|
||||
func TestGoDaddy_FromChunk(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second*5)
|
||||
defer cancel()
|
||||
testSecrets, err := common.GetSecret(ctx, "trufflehog-testing", "detectors5")
|
||||
if err != nil {
|
||||
t.Fatalf("could not get test secrets from GCP: %s", err)
|
||||
}
|
||||
secret := testSecrets.MustGetField("GODADDY_OTE")
|
||||
inactiveSecret := testSecrets.MustGetField("GODADDY_OTE_INACTIVE")
|
||||
|
||||
type args struct {
|
||||
ctx context.Context
|
||||
data []byte
|
||||
verify bool
|
||||
}
|
||||
tests := []struct {
|
||||
name string
|
||||
s Scanner
|
||||
args args
|
||||
want []detectors.Result
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "found, verified",
|
||||
s: Scanner{},
|
||||
args: args{
|
||||
ctx: context.Background(),
|
||||
data: []byte(fmt.Sprintf("You can find a godaddy secret %s within", secret)),
|
||||
verify: true,
|
||||
},
|
||||
want: []detectors.Result{
|
||||
{
|
||||
DetectorType: detectorspb.DetectorType_GoDaddy,
|
||||
Verified: true,
|
||||
},
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "found, unverified",
|
||||
s: Scanner{},
|
||||
args: args{
|
||||
ctx: context.Background(),
|
||||
data: []byte(fmt.Sprintf("You can find a godaddy secret %s within but not valid", inactiveSecret)), // the secret would satisfy the regex but not pass validation
|
||||
verify: true,
|
||||
},
|
||||
want: []detectors.Result{
|
||||
{
|
||||
DetectorType: detectorspb.DetectorType_GoDaddy,
|
||||
Verified: false,
|
||||
},
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "not found",
|
||||
s: Scanner{},
|
||||
args: args{
|
||||
ctx: context.Background(),
|
||||
data: []byte("You cannot find the secret within"),
|
||||
verify: true,
|
||||
},
|
||||
want: nil,
|
||||
wantErr: false,
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
s := Scanner{}
|
||||
got, err := s.FromData(tt.args.ctx, tt.args.verify, tt.args.data)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("GoDaddy.FromData() error = %v, wantErr %v", err, tt.wantErr)
|
||||
return
|
||||
}
|
||||
for i := range got {
|
||||
if len(got[i].Raw) == 0 {
|
||||
t.Fatalf("no raw secret present: \n %+v", got[i])
|
||||
}
|
||||
got[i].Raw = nil
|
||||
}
|
||||
if diff := pretty.Compare(got, tt.want); diff != "" {
|
||||
t.Errorf("GoDaddy.FromData() %s diff: (-got +want)\n%s", tt.name, diff)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkFromData(benchmark *testing.B) {
|
||||
ctx := context.Background()
|
||||
s := Scanner{}
|
||||
for name, data := range detectors.MustGetBenchmarkData() {
|
||||
benchmark.Run(name, func(b *testing.B) {
|
||||
b.ResetTimer()
|
||||
for n := 0; n < b.N; n++ {
|
||||
_, err := s.FromData(ctx, false, data)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,90 @@
|
||||
package godaddy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/google/go-cmp/cmp"
|
||||
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/engine/ahocorasick"
|
||||
)
|
||||
|
||||
var (
|
||||
validPattern = `[{
|
||||
"_id": "1a8d0cca-e1a9-4318-bc2f-f5658ab2dcb5",
|
||||
"name": "GoDaddy",
|
||||
"type": "Detector",
|
||||
"api": true,
|
||||
"authentication_type": "",
|
||||
"verification_url": "https://api.example.com/example",
|
||||
"test_secrets": {
|
||||
"godaddyKey": "2TM44WqB21o4zH_3xM44WkB21i4zHHhXSoHjO",
|
||||
"godaddySecret": "3xM44WkB21i4zHHhXSoHjO",
|
||||
"not_godaddySecret": "2TM44WqB21o4zH$3xM44WkB21i4zHHhXSoHjO"
|
||||
},
|
||||
"expected_response": "200",
|
||||
"method": "GET",
|
||||
"deprecated": false
|
||||
}]`
|
||||
secret = "2TM44WqB21o4zH_3xM44WkB21i4zHHhXSoHjO"
|
||||
)
|
||||
|
||||
func TestGoDaddy_Pattern(t *testing.T) {
|
||||
d := Scanner{}
|
||||
ahoCorasickCore := ahocorasick.NewAhoCorasickCore([]detectors.Detector{d})
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
want []string
|
||||
}{
|
||||
{
|
||||
name: "valid pattern",
|
||||
input: validPattern,
|
||||
want: []string{secret},
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
matchedDetectors := ahoCorasickCore.FindDetectorMatches([]byte(test.input))
|
||||
if len(matchedDetectors) == 0 {
|
||||
t.Errorf("keywords '%v' not matched by: %s", d.Keywords(), test.input)
|
||||
return
|
||||
}
|
||||
|
||||
results, err := d.FromData(context.Background(), false, []byte(test.input))
|
||||
if err != nil {
|
||||
t.Errorf("error = %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
if len(results) != len(test.want) {
|
||||
if len(results) == 0 {
|
||||
t.Errorf("did not receive result")
|
||||
} else {
|
||||
t.Errorf("expected %d results, only received %d", len(test.want), len(results))
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
actual := make(map[string]struct{}, len(results))
|
||||
for _, r := range results {
|
||||
if len(r.RawV2) > 0 {
|
||||
actual[string(r.RawV2)] = struct{}{}
|
||||
} else {
|
||||
actual[string(r.Raw)] = struct{}{}
|
||||
}
|
||||
}
|
||||
expected := make(map[string]struct{}, len(test.want))
|
||||
for _, v := range test.want {
|
||||
expected[v] = struct{}{}
|
||||
}
|
||||
|
||||
if diff := cmp.Diff(expected, actual); diff != "" {
|
||||
t.Errorf("%s diff: (-want +got)\n%s", test.name, diff)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,136 @@
|
||||
package godaddy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
|
||||
regexp "github.com/wasilibs/go-re2"
|
||||
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/common"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors"
|
||||
v1 "github.com/trufflesecurity/trufflehog/v3/pkg/detectors/godaddy/v1"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/pb/detectorspb"
|
||||
)
|
||||
|
||||
type Scanner struct {
|
||||
client *http.Client
|
||||
}
|
||||
|
||||
var (
|
||||
// ensure the scanner satisfies the interface at compile time.
|
||||
_ detectors.Detector = (*Scanner)(nil)
|
||||
_ detectors.Versioner = (*Scanner)(nil)
|
||||
|
||||
defaultClient = common.SaneHttpClient()
|
||||
|
||||
// the key for the GoDaddy Prod environment is a 35-character alphanumeric string that may include underscores.
|
||||
keyPattern = regexp.MustCompile(detectors.PrefixRegex([]string{"godaddy"}) + common.BuildRegex("a-zA-Z0-9", "_", 35))
|
||||
// the secret for the GoDaddy Prod environment is a 22-character alphanumeric string.
|
||||
secretPattern = regexp.MustCompile(detectors.PrefixRegex([]string{"godaddy"}) + common.BuildRegex("a-zA-Z0-9", "", 22))
|
||||
|
||||
// prod environment
|
||||
prod = "api.godaddy.com"
|
||||
)
|
||||
|
||||
func (s *Scanner) getClient() *http.Client {
|
||||
if s.client != nil {
|
||||
return s.client
|
||||
}
|
||||
|
||||
return defaultClient
|
||||
}
|
||||
|
||||
func (s *Scanner) Version() int { return 2 }
|
||||
|
||||
// Keywords are used for efficiently pre-filtering chunks.
|
||||
// Use identifiers in the secret preferably, or the provider name.
|
||||
func (s Scanner) Keywords() []string {
|
||||
return []string{"godaddy"}
|
||||
}
|
||||
|
||||
func (s Scanner) Description() string {
|
||||
return "GoDaddy offers website building, hosting and security tools and services to construct, expand and protect the online presence." +
|
||||
"GoDaddy provides applications and access to relevant third-party products and platforms to connect their customers"
|
||||
}
|
||||
|
||||
func (s Scanner) Type() detectorspb.DetectorType {
|
||||
return detectorspb.DetectorType_GoDaddy
|
||||
}
|
||||
|
||||
// FromData will find and optionally verify GoDaddy API Key and secrets in a given set of bytes.
|
||||
func (s Scanner) FromData(ctx context.Context, verify bool, data []byte) (results []detectors.Result, err error) {
|
||||
// convert the data to string
|
||||
dataStr := string(data)
|
||||
|
||||
// find all the matching keys and secret in data and make a unique maps of both keys and secret.
|
||||
uniqueKeys, uniqueSecrets := make(map[string]struct{}), make(map[string]struct{})
|
||||
|
||||
for _, foundKey := range keyPattern.FindAllStringSubmatch(dataStr, -1) {
|
||||
uniqueKeys[foundKey[1]] = struct{}{}
|
||||
}
|
||||
|
||||
for _, foundSecret := range secretPattern.FindAllStringSubmatch(dataStr, -1) {
|
||||
uniqueSecrets[foundSecret[1]] = struct{}{}
|
||||
}
|
||||
|
||||
for key := range uniqueKeys {
|
||||
for secret := range uniqueSecrets {
|
||||
result := detectors.Result{
|
||||
DetectorType: detectorspb.DetectorType_GoDaddy,
|
||||
Raw: []byte(key),
|
||||
ExtraData: make(map[string]string),
|
||||
}
|
||||
|
||||
if verify {
|
||||
isVerified, verificationErr := VerifyGoDaddySecret(ctx, s.getClient(), prod, v1.MakeAuthHeaderValue(key, secret))
|
||||
|
||||
result.Verified = isVerified
|
||||
result.SetVerificationError(verificationErr, secret)
|
||||
|
||||
// in case of successful verification add the enviorement name in extradata to let user know which env this secret belong to.
|
||||
if isVerified {
|
||||
result.ExtraData["Environment"] = "Prod"
|
||||
}
|
||||
}
|
||||
|
||||
results = append(results, result)
|
||||
}
|
||||
}
|
||||
|
||||
return results, nil
|
||||
|
||||
}
|
||||
|
||||
// VerifyGoDaddySecret make a call to godaddy api with given secret to check if secret is valid or not.
|
||||
func VerifyGoDaddySecret(ctx context.Context, client *http.Client, environment, secret string) (bool, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", fmt.Sprintf("https://%s/v1/domains/available?domain=example.com", environment), http.NoBody)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
// set the required auth header
|
||||
req.Header.Set("Authorization", secret)
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
defer func() {
|
||||
_, _ = io.Copy(io.Discard, resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
}()
|
||||
|
||||
switch resp.StatusCode {
|
||||
case http.StatusOK:
|
||||
return true, nil
|
||||
case http.StatusUnauthorized:
|
||||
return false, nil
|
||||
case http.StatusForbidden:
|
||||
// as per documentation in case of 403 the token is actually verified but it does not have access.
|
||||
return true, nil
|
||||
default:
|
||||
return false, fmt.Errorf("unexpected status code: %d", resp.StatusCode)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,120 @@
|
||||
//go:build detectors
|
||||
// +build detectors
|
||||
|
||||
package godaddy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/kylelemons/godebug/pretty"
|
||||
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/common"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/pb/detectorspb"
|
||||
)
|
||||
|
||||
func TestGoDaddy_FromChunk(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second*5)
|
||||
defer cancel()
|
||||
testSecrets, err := common.GetSecret(ctx, "trufflehog-testing", "detectors5")
|
||||
if err != nil {
|
||||
t.Fatalf("could not get test secrets from GCP: %s", err)
|
||||
}
|
||||
secret := testSecrets.MustGetField("GODADDY_PROD")
|
||||
inactiveSecret := testSecrets.MustGetField("GODADDY_PROD_INACTIVE")
|
||||
|
||||
type args struct {
|
||||
ctx context.Context
|
||||
data []byte
|
||||
verify bool
|
||||
}
|
||||
tests := []struct {
|
||||
name string
|
||||
s Scanner
|
||||
args args
|
||||
want []detectors.Result
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "found, verified",
|
||||
s: Scanner{},
|
||||
args: args{
|
||||
ctx: context.Background(),
|
||||
data: []byte(fmt.Sprintf("You can find a godaddy secret %s within", secret)),
|
||||
verify: true,
|
||||
},
|
||||
want: []detectors.Result{
|
||||
{
|
||||
DetectorType: detectorspb.DetectorType_GoDaddy,
|
||||
Verified: true,
|
||||
},
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "found, unverified",
|
||||
s: Scanner{},
|
||||
args: args{
|
||||
ctx: context.Background(),
|
||||
data: []byte(fmt.Sprintf("You can find a godaddy secret %s within but not valid", inactiveSecret)), // the secret would satisfy the regex but not pass validation
|
||||
verify: true,
|
||||
},
|
||||
want: []detectors.Result{
|
||||
{
|
||||
DetectorType: detectorspb.DetectorType_GoDaddy,
|
||||
Verified: false,
|
||||
},
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "not found",
|
||||
s: Scanner{},
|
||||
args: args{
|
||||
ctx: context.Background(),
|
||||
data: []byte("You cannot find the secret within"),
|
||||
verify: true,
|
||||
},
|
||||
want: nil,
|
||||
wantErr: false,
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
s := Scanner{}
|
||||
got, err := s.FromData(tt.args.ctx, tt.args.verify, tt.args.data)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("GoDaddy.FromData() error = %v, wantErr %v", err, tt.wantErr)
|
||||
return
|
||||
}
|
||||
for i := range got {
|
||||
if len(got[i].Raw) == 0 {
|
||||
t.Fatalf("no raw secret present: \n %+v", got[i])
|
||||
}
|
||||
got[i].Raw = nil
|
||||
}
|
||||
if diff := pretty.Compare(got, tt.want); diff != "" {
|
||||
t.Errorf("GoDaddy.FromData() %s diff: (-got +want)\n%s", tt.name, diff)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkFromData(benchmark *testing.B) {
|
||||
ctx := context.Background()
|
||||
s := Scanner{}
|
||||
for name, data := range detectors.MustGetBenchmarkData() {
|
||||
benchmark.Run(name, func(b *testing.B) {
|
||||
b.ResetTimer()
|
||||
for n := 0; n < b.N; n++ {
|
||||
_, err := s.FromData(ctx, false, data)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,90 @@
|
||||
package godaddy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/google/go-cmp/cmp"
|
||||
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/engine/ahocorasick"
|
||||
)
|
||||
|
||||
var (
|
||||
validPattern = `[{
|
||||
"_id": "1a8d0cca-e1a9-4318-bc2f-f5658ab2dcb5",
|
||||
"name": "GoDaddy",
|
||||
"type": "Detector",
|
||||
"api": true,
|
||||
"authentication_type": "",
|
||||
"verification_url": "https://api.example.com/example",
|
||||
"test_secrets": {
|
||||
"godaddyKey": "2TM44WqB21o4zH_3xM44WkB21i4zHHhXSoH",
|
||||
"godaddySecret": "3xM44WkB21i4zHHhXSoHjO",
|
||||
"not_godaddySecret": "2TM44WqB21o4zH@3xM44WkB21i4zHHhXSoH"
|
||||
},
|
||||
"expected_response": "200",
|
||||
"method": "GET",
|
||||
"deprecated": false
|
||||
}]`
|
||||
secret = "2TM44WqB21o4zH_3xM44WkB21i4zHHhXSoH"
|
||||
)
|
||||
|
||||
func TestGoDaddy_Pattern(t *testing.T) {
|
||||
d := Scanner{}
|
||||
ahoCorasickCore := ahocorasick.NewAhoCorasickCore([]detectors.Detector{d})
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
want []string
|
||||
}{
|
||||
{
|
||||
name: "valid pattern",
|
||||
input: validPattern,
|
||||
want: []string{secret},
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
matchedDetectors := ahoCorasickCore.FindDetectorMatches([]byte(test.input))
|
||||
if len(matchedDetectors) == 0 {
|
||||
t.Errorf("keywords '%v' not matched by: %s", d.Keywords(), test.input)
|
||||
return
|
||||
}
|
||||
|
||||
results, err := d.FromData(context.Background(), false, []byte(test.input))
|
||||
if err != nil {
|
||||
t.Errorf("error = %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
if len(results) != len(test.want) {
|
||||
if len(results) == 0 {
|
||||
t.Errorf("did not receive result")
|
||||
} else {
|
||||
t.Errorf("expected %d results, only received %d", len(test.want), len(results))
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
actual := make(map[string]struct{}, len(results))
|
||||
for _, r := range results {
|
||||
if len(r.RawV2) > 0 {
|
||||
actual[string(r.RawV2)] = struct{}{}
|
||||
} else {
|
||||
actual[string(r.Raw)] = struct{}{}
|
||||
}
|
||||
}
|
||||
expected := make(map[string]struct{}, len(test.want))
|
||||
for _, v := range test.want {
|
||||
expected[v] = struct{}{}
|
||||
}
|
||||
|
||||
if diff := cmp.Diff(expected, actual); diff != "" {
|
||||
t.Errorf("%s diff: (-want +got)\n%s", test.name, diff)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,7 +1,7 @@
|
||||
//go:build detectors
|
||||
// +build detectors
|
||||
|
||||
package hubspotapikey
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
@@ -34,5 +34,5 @@ func (e *Engine) ScanCircleCI(ctx context.Context, token string) (sources.JobPro
|
||||
if err := circleSource.Init(ctx, "trufflehog - Circle CI", jobID, sourceID, true, &conn, runtime.NumCPU()); err != nil {
|
||||
return sources.JobProgressRef{}, err
|
||||
}
|
||||
return e.sourceManager.Run(ctx, sourceName, circleSource)
|
||||
return e.sourceManager.EnumerateAndScan(ctx, sourceName, circleSource)
|
||||
}
|
||||
|
||||
@@ -109,8 +109,8 @@ import (
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors/budibase"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors/bugherd"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors/bugsnag"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors/buildkite"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors/buildkitev2"
|
||||
buildKitev1 "github.com/trufflesecurity/trufflehog/v3/pkg/detectors/buildkite/v1"
|
||||
buildKitev2 "github.com/trufflesecurity/trufflehog/v3/pkg/detectors/buildkite/v2"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors/bulbul"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors/bulksms"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors/buttercms"
|
||||
@@ -321,6 +321,8 @@ import (
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors/glassnode"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors/gocanvas"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors/gocardless"
|
||||
godaddyv1 "github.com/trufflesecurity/trufflehog/v3/pkg/detectors/godaddy/v1"
|
||||
godaddyv2 "github.com/trufflesecurity/trufflehog/v3/pkg/detectors/godaddy/v2"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors/goodday"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors/googleoauth2"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/detectors/grafana"
|
||||
@@ -932,8 +934,8 @@ func buildDetectorList() []detectors.Detector {
|
||||
&budibase.Scanner{},
|
||||
&bugherd.Scanner{},
|
||||
&bugsnag.Scanner{},
|
||||
&buildkite.Scanner{},
|
||||
&buildkitev2.Scanner{},
|
||||
&buildKitev1.Scanner{},
|
||||
&buildKitev2.Scanner{},
|
||||
&bulbul.Scanner{},
|
||||
&bulksms.Scanner{},
|
||||
&buttercms.Scanner{},
|
||||
@@ -1149,6 +1151,8 @@ func buildDetectorList() []detectors.Detector {
|
||||
&glassnode.Scanner{},
|
||||
&gocanvas.Scanner{},
|
||||
&gocardless.Scanner{},
|
||||
&godaddyv1.Scanner{},
|
||||
&godaddyv2.Scanner{},
|
||||
&goodday.Scanner{},
|
||||
&googleoauth2.Scanner{},
|
||||
&grafana.Scanner{},
|
||||
|
||||
@@ -39,5 +39,5 @@ func (e *Engine) ScanDocker(ctx context.Context, c sources.DockerConfig) (source
|
||||
if err := dockerSource.Init(ctx, sourceName, jobID, sourceID, true, &conn, runtime.NumCPU()); err != nil {
|
||||
return sources.JobProgressRef{}, err
|
||||
}
|
||||
return e.sourceManager.Run(ctx, sourceName, dockerSource)
|
||||
return e.sourceManager.EnumerateAndScan(ctx, sourceName, dockerSource)
|
||||
}
|
||||
|
||||
@@ -41,5 +41,5 @@ func (e *Engine) ScanElasticsearch(ctx context.Context, c sources.ElasticsearchC
|
||||
if err := elasticsearchSource.Init(ctx, sourceName, jobID, sourceID, true, &conn, runtime.NumCPU()); err != nil {
|
||||
return sources.JobProgressRef{}, err
|
||||
}
|
||||
return e.sourceManager.Run(ctx, sourceName, elasticsearchSource)
|
||||
return e.sourceManager.EnumerateAndScan(ctx, sourceName, elasticsearchSource)
|
||||
}
|
||||
|
||||
@@ -33,5 +33,5 @@ func (e *Engine) ScanFileSystem(ctx context.Context, c sources.FilesystemConfig)
|
||||
if err := fileSystemSource.Init(ctx, sourceName, jobID, sourceID, true, &conn, runtime.NumCPU()); err != nil {
|
||||
return sources.JobProgressRef{}, err
|
||||
}
|
||||
return e.sourceManager.Run(ctx, sourceName, fileSystemSource)
|
||||
return e.sourceManager.EnumerateAndScan(ctx, sourceName, fileSystemSource)
|
||||
}
|
||||
|
||||
+1
-1
@@ -51,7 +51,7 @@ func (e *Engine) ScanGCS(ctx context.Context, c sources.GCSConfig) (sources.JobP
|
||||
if err := gcsSource.Init(ctx, sourceName, jobID, sourceID, true, &conn, int(c.Concurrency)); err != nil {
|
||||
return sources.JobProgressRef{}, err
|
||||
}
|
||||
return e.sourceManager.Run(ctx, sourceName, gcsSource)
|
||||
return e.sourceManager.EnumerateAndScan(ctx, sourceName, gcsSource)
|
||||
}
|
||||
|
||||
func isAuthValid(ctx context.Context, c sources.GCSConfig, connection *sourcespb.GCS) bool {
|
||||
|
||||
+1
-1
@@ -39,5 +39,5 @@ func (e *Engine) ScanGit(ctx context.Context, c sources.GitConfig) (sources.JobP
|
||||
return sources.JobProgressRef{}, err
|
||||
}
|
||||
|
||||
return e.sourceManager.Run(ctx, sourceName, gitSource)
|
||||
return e.sourceManager.EnumerateAndScan(ctx, sourceName, gitSource)
|
||||
}
|
||||
|
||||
@@ -59,5 +59,5 @@ func (e *Engine) ScanGitHub(ctx context.Context, c sources.GithubConfig) (source
|
||||
return sources.JobProgressRef{}, err
|
||||
}
|
||||
githubSource.WithScanOptions(scanOptions)
|
||||
return e.sourceManager.Run(ctx, sourceName, githubSource)
|
||||
return e.sourceManager.EnumerateAndScan(ctx, sourceName, githubSource)
|
||||
}
|
||||
|
||||
@@ -60,5 +60,5 @@ func (e *Engine) ScanGitHubExperimental(ctx context.Context, c sources.GitHubExp
|
||||
return sources.JobProgressRef{}, err
|
||||
}
|
||||
githubExperimentalSource.WithScanOptions(scanOptions)
|
||||
return e.sourceManager.Run(ctx, sourceName, githubExperimentalSource)
|
||||
return e.sourceManager.EnumerateAndScan(ctx, sourceName, githubExperimentalSource)
|
||||
}
|
||||
|
||||
@@ -66,5 +66,5 @@ func (e *Engine) ScanGitLab(ctx context.Context, c sources.GitlabConfig) (source
|
||||
return sources.JobProgressRef{}, err
|
||||
}
|
||||
gitlabSource.WithScanOptions(scanOptions)
|
||||
return e.sourceManager.Run(ctx, sourceName, gitlabSource)
|
||||
return e.sourceManager.EnumerateAndScan(ctx, sourceName, gitlabSource)
|
||||
}
|
||||
|
||||
@@ -76,5 +76,5 @@ func (e *Engine) ScanHuggingface(ctx context.Context, c HuggingfaceConfig) (sour
|
||||
if err := huggingfaceSource.Init(ctx, sourceName, jobID, sourceID, true, &conn, c.Concurrency); err != nil {
|
||||
return sources.JobProgressRef{}, err
|
||||
}
|
||||
return e.sourceManager.Run(ctx, sourceName, huggingfaceSource)
|
||||
return e.sourceManager.EnumerateAndScan(ctx, sourceName, huggingfaceSource)
|
||||
}
|
||||
|
||||
@@ -77,5 +77,5 @@ func (e *Engine) ScanJenkins(ctx context.Context, jenkinsConfig JenkinsConfig) (
|
||||
if err := jenkinsSource.Init(ctx, "trufflehog - Jenkins", jobID, sourceID, true, &conn, runtime.NumCPU()); err != nil {
|
||||
return sources.JobProgressRef{}, err
|
||||
}
|
||||
return e.sourceManager.Run(ctx, sourceName, jenkinsSource)
|
||||
return e.sourceManager.EnumerateAndScan(ctx, sourceName, jenkinsSource)
|
||||
}
|
||||
|
||||
@@ -61,5 +61,5 @@ func (e *Engine) ScanPostman(ctx context.Context, c sources.PostmanConfig) (sour
|
||||
if err := postmanSource.Init(ctx, sourceName, jobID, sourceID, true, &conn, c.Concurrency); err != nil {
|
||||
return sources.JobProgressRef{}, err
|
||||
}
|
||||
return e.sourceManager.Run(ctx, sourceName, postmanSource)
|
||||
return e.sourceManager.EnumerateAndScan(ctx, sourceName, postmanSource)
|
||||
}
|
||||
|
||||
+1
-1
@@ -68,5 +68,5 @@ func (e *Engine) ScanS3(ctx context.Context, c sources.S3Config) (sources.JobPro
|
||||
if err := s3Source.Init(ctx, sourceName, jobID, sourceID, true, &conn, runtime.NumCPU()); err != nil {
|
||||
return sources.JobProgressRef{}, err
|
||||
}
|
||||
return e.sourceManager.Run(ctx, sourceName, s3Source)
|
||||
return e.sourceManager.EnumerateAndScan(ctx, sourceName, s3Source)
|
||||
}
|
||||
|
||||
@@ -49,5 +49,5 @@ func (e *Engine) ScanSyslog(ctx context.Context, c sources.SyslogConfig) (source
|
||||
}
|
||||
syslogSource.InjectConnection(connection)
|
||||
|
||||
return e.sourceManager.Run(ctx, sourceName, syslogSource)
|
||||
return e.sourceManager.EnumerateAndScan(ctx, sourceName, syslogSource)
|
||||
}
|
||||
|
||||
@@ -34,5 +34,5 @@ func (e *Engine) ScanTravisCI(ctx context.Context, token string) (sources.JobPro
|
||||
if err := travisSource.Init(ctx, sourceName, jobID, sourceID, true, &conn, runtime.NumCPU()); err != nil {
|
||||
return sources.JobProgressRef{}, err
|
||||
}
|
||||
return e.sourceManager.Run(ctx, sourceName, travisSource)
|
||||
return e.sourceManager.EnumerateAndScan(ctx, sourceName, travisSource)
|
||||
}
|
||||
|
||||
@@ -1110,6 +1110,7 @@ const (
|
||||
DetectorType_WeightsAndBiases DetectorType = 1005
|
||||
DetectorType_ZohoCRM DetectorType = 1006
|
||||
DetectorType_AzureOpenAI DetectorType = 1007
|
||||
DetectorType_GoDaddy DetectorType = 1008
|
||||
)
|
||||
|
||||
// Enum value maps for DetectorType.
|
||||
@@ -2119,6 +2120,7 @@ var (
|
||||
1005: "WeightsAndBiases",
|
||||
1006: "ZohoCRM",
|
||||
1007: "AzureOpenAI",
|
||||
1008: "GoDaddy",
|
||||
}
|
||||
DetectorType_value = map[string]int32{
|
||||
"Alibaba": 0,
|
||||
@@ -3125,6 +3127,7 @@ var (
|
||||
"WeightsAndBiases": 1005,
|
||||
"ZohoCRM": 1006,
|
||||
"AzureOpenAI": 1007,
|
||||
"GoDaddy": 1008,
|
||||
}
|
||||
)
|
||||
|
||||
@@ -3578,7 +3581,7 @@ var file_detectors_proto_rawDesc = []byte{
|
||||
0x4c, 0x41, 0x49, 0x4e, 0x10, 0x01, 0x12, 0x0a, 0x0a, 0x06, 0x42, 0x41, 0x53, 0x45, 0x36, 0x34,
|
||||
0x10, 0x02, 0x12, 0x09, 0x0a, 0x05, 0x55, 0x54, 0x46, 0x31, 0x36, 0x10, 0x03, 0x12, 0x13, 0x0a,
|
||||
0x0f, 0x45, 0x53, 0x43, 0x41, 0x50, 0x45, 0x44, 0x5f, 0x55, 0x4e, 0x49, 0x43, 0x4f, 0x44, 0x45,
|
||||
0x10, 0x04, 0x2a, 0xdf, 0x80, 0x01, 0x0a, 0x0c, 0x44, 0x65, 0x74, 0x65, 0x63, 0x74, 0x6f, 0x72,
|
||||
0x10, 0x04, 0x2a, 0xed, 0x80, 0x01, 0x0a, 0x0c, 0x44, 0x65, 0x74, 0x65, 0x63, 0x74, 0x6f, 0x72,
|
||||
0x54, 0x79, 0x70, 0x65, 0x12, 0x0b, 0x0a, 0x07, 0x41, 0x6c, 0x69, 0x62, 0x61, 0x62, 0x61, 0x10,
|
||||
0x00, 0x12, 0x08, 0x0a, 0x04, 0x41, 0x4d, 0x51, 0x50, 0x10, 0x01, 0x12, 0x07, 0x0a, 0x03, 0x41,
|
||||
0x57, 0x53, 0x10, 0x02, 0x12, 0x09, 0x0a, 0x05, 0x41, 0x7a, 0x75, 0x72, 0x65, 0x10, 0x03, 0x12,
|
||||
@@ -4608,11 +4611,12 @@ var file_detectors_proto_rawDesc = []byte{
|
||||
0x0a, 0x10, 0x57, 0x65, 0x69, 0x67, 0x68, 0x74, 0x73, 0x41, 0x6e, 0x64, 0x42, 0x69, 0x61, 0x73,
|
||||
0x65, 0x73, 0x10, 0xed, 0x07, 0x12, 0x0c, 0x0a, 0x07, 0x5a, 0x6f, 0x68, 0x6f, 0x43, 0x52, 0x4d,
|
||||
0x10, 0xee, 0x07, 0x12, 0x10, 0x0a, 0x0b, 0x41, 0x7a, 0x75, 0x72, 0x65, 0x4f, 0x70, 0x65, 0x6e,
|
||||
0x41, 0x49, 0x10, 0xef, 0x07, 0x42, 0x3d, 0x5a, 0x3b, 0x67, 0x69, 0x74, 0x68, 0x75, 0x62, 0x2e,
|
||||
0x63, 0x6f, 0x6d, 0x2f, 0x74, 0x72, 0x75, 0x66, 0x66, 0x6c, 0x65, 0x73, 0x65, 0x63, 0x75, 0x72,
|
||||
0x69, 0x74, 0x79, 0x2f, 0x74, 0x72, 0x75, 0x66, 0x66, 0x6c, 0x65, 0x68, 0x6f, 0x67, 0x2f, 0x76,
|
||||
0x33, 0x2f, 0x70, 0x6b, 0x67, 0x2f, 0x70, 0x62, 0x2f, 0x64, 0x65, 0x74, 0x65, 0x63, 0x74, 0x6f,
|
||||
0x72, 0x73, 0x70, 0x62, 0x62, 0x06, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x33,
|
||||
0x41, 0x49, 0x10, 0xef, 0x07, 0x12, 0x0c, 0x0a, 0x07, 0x47, 0x6f, 0x44, 0x61, 0x64, 0x64, 0x79,
|
||||
0x10, 0xf0, 0x07, 0x42, 0x3d, 0x5a, 0x3b, 0x67, 0x69, 0x74, 0x68, 0x75, 0x62, 0x2e, 0x63, 0x6f,
|
||||
0x6d, 0x2f, 0x74, 0x72, 0x75, 0x66, 0x66, 0x6c, 0x65, 0x73, 0x65, 0x63, 0x75, 0x72, 0x69, 0x74,
|
||||
0x79, 0x2f, 0x74, 0x72, 0x75, 0x66, 0x66, 0x6c, 0x65, 0x68, 0x6f, 0x67, 0x2f, 0x76, 0x33, 0x2f,
|
||||
0x70, 0x6b, 0x67, 0x2f, 0x70, 0x62, 0x2f, 0x64, 0x65, 0x74, 0x65, 0x63, 0x74, 0x6f, 0x72, 0x73,
|
||||
0x70, 0x62, 0x62, 0x06, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x33,
|
||||
}
|
||||
|
||||
var (
|
||||
|
||||
@@ -3,6 +3,7 @@ package sources
|
||||
import (
|
||||
"bufio"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/context"
|
||||
@@ -99,6 +100,22 @@ func readInChunks(ctx context.Context, reader io.Reader, config *chunkReaderConf
|
||||
go func() {
|
||||
defer close(chunkResultChan)
|
||||
|
||||
// Defer a panic recovery to handle any panics that occur while reading, which can sometimes unavoidably happen
|
||||
// due to third-party library bugs.
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
var panicErr error
|
||||
if e, ok := r.(error); ok {
|
||||
panicErr = e
|
||||
} else {
|
||||
panicErr = fmt.Errorf("panic occurred: %v", r)
|
||||
}
|
||||
chunkResultChan <- ChunkResult{
|
||||
err: fmt.Errorf("panic error: %w", panicErr),
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
for {
|
||||
chunkRes := ChunkResult{}
|
||||
chunkBytes := make([]byte, config.totalSize)
|
||||
|
||||
@@ -2,6 +2,7 @@ package sources
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"math/rand"
|
||||
"runtime"
|
||||
"strings"
|
||||
@@ -10,6 +11,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/context"
|
||||
)
|
||||
@@ -121,6 +123,21 @@ func TestNewChunkedReader(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
type panicReader struct{}
|
||||
|
||||
var _ io.Reader = (*panicReader)(nil)
|
||||
|
||||
func (_ panicReader) Read([]byte) (int, error) {
|
||||
panic("panic for testing")
|
||||
}
|
||||
|
||||
func TestChunkReader_UnderlyingReaderPanics_DoesNotPanic(t *testing.T) {
|
||||
require.NotPanics(t, func() {
|
||||
for range NewChunkReader()(context.Background(), &panicReader{}) {
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func BenchmarkChunkReader(b *testing.B) {
|
||||
var bigChunk = make([]byte, 1<<24) // 16MB
|
||||
|
||||
|
||||
@@ -958,8 +958,11 @@ func (s *Source) addMembersByOrg(ctx context.Context, org string) error {
|
||||
if s.handleRateLimit(ctx, err) {
|
||||
continue
|
||||
}
|
||||
if err != nil || len(members) == 0 {
|
||||
return fmt.Errorf("could not list organization members: account may not have access to list organization members %w", err)
|
||||
if err != nil {
|
||||
return fmt.Errorf("could not list organization (%q) members: account may not have access to list organization members: %w", org, err)
|
||||
}
|
||||
if len(members) == 0 {
|
||||
return fmt.Errorf("organization (%q) had 0 members: account may not have access to list organization members", org)
|
||||
}
|
||||
|
||||
logger.V(2).Info("Listed members", "page", opts.Page, "last_page", res.LastPage)
|
||||
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
"reflect"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -207,6 +208,41 @@ func TestAddMembersByOrg(t *testing.T) {
|
||||
assert.True(t, gock.IsDone())
|
||||
}
|
||||
|
||||
func TestAddMembersByOrg_AuthFailure(t *testing.T) {
|
||||
defer gock.Off()
|
||||
|
||||
gock.New("https://api.github.com").
|
||||
Get("/orgs/org1/members").
|
||||
Reply(401).
|
||||
JSON([]map[string]string{{
|
||||
"message": "Bad credentials",
|
||||
"documentation_url": "https://docs.github.com/rest",
|
||||
"status": "401",
|
||||
}})
|
||||
|
||||
s := initTestSource(&sourcespb.GitHub{Credential: &sourcespb.GitHub_Unauthenticated{}})
|
||||
err := s.addMembersByOrg(context.Background(), "org1")
|
||||
assert.True(t, strings.HasPrefix(err.Error(), "could not list organization"))
|
||||
assert.False(t, gock.HasUnmatchedRequest())
|
||||
assert.True(t, gock.IsDone())
|
||||
}
|
||||
|
||||
func TestAddMembersByOrg_NoMembers(t *testing.T) {
|
||||
defer gock.Off()
|
||||
|
||||
gock.New("https://api.github.com").
|
||||
Get("/orgs/org1/members").
|
||||
Reply(200).
|
||||
JSON([]map[string]string{})
|
||||
|
||||
s := initTestSource(&sourcespb.GitHub{Credential: &sourcespb.GitHub_Unauthenticated{}})
|
||||
err := s.addMembersByOrg(context.Background(), "org1")
|
||||
|
||||
assert.Equal(t, fmt.Sprintf("organization (%q) had 0 members: account may not have access to list organization members", "org1"), err.Error())
|
||||
assert.False(t, gock.HasUnmatchedRequest())
|
||||
assert.True(t, gock.IsDone())
|
||||
}
|
||||
|
||||
func TestAddMembersByApp(t *testing.T) {
|
||||
defer gock.Off()
|
||||
|
||||
|
||||
@@ -11,16 +11,12 @@ import (
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/sources"
|
||||
)
|
||||
|
||||
// ProgressTracker maintains scan progress state for S3 bucket scanning,
|
||||
// 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.
|
||||
//
|
||||
// S3 buckets are organized as flat namespaces of objects identified by unique keys.
|
||||
// When listing objects, S3 returns paginated results with a maximum of 1000 objects per page.
|
||||
// The ListObjectsV2 API accepts a 'StartAfter' parameter that allows resuming the listing
|
||||
// from a specific object key.
|
||||
//
|
||||
// The tracker maintains state for the current page of objects (up to 1000) using a boolean array
|
||||
// The checkpointer maintains state for the current page of objects (up to 1000) using a boolean array
|
||||
// to track completion status and an ordered list to record the sequence of completions.
|
||||
// This enables finding the highest consecutive completed index as a "low water mark".
|
||||
//
|
||||
@@ -41,14 +37,24 @@ import (
|
||||
// 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
|
||||
// On resume: StartAfter=object1499 in saved bucket, scanning continues from object 1500
|
||||
type ProgressTracker struct {
|
||||
//
|
||||
// Important constraints:
|
||||
// - Only tracks completion state for a single page of objects (up to 1000)
|
||||
// - Supports concurrent object processing within a page
|
||||
// - Does NOT support concurrent page processing
|
||||
// - Must be Reset() between pages
|
||||
type Checkpointer struct {
|
||||
enabled bool
|
||||
|
||||
// completedObjects tracks which indices in the current page have been processed.
|
||||
sync.Mutex
|
||||
mu sync.Mutex // protects concurrent access to completion state.
|
||||
completedObjects []bool
|
||||
completionOrder []int // Track the order in which objects complete
|
||||
|
||||
// lowestIncompleteIdx tracks the first index that hasn't been completed.
|
||||
// This optimizes checkpoint creation by avoiding recalculation.
|
||||
lowestIncompleteIdx int
|
||||
|
||||
// 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
|
||||
@@ -56,13 +62,13 @@ type ProgressTracker struct {
|
||||
|
||||
const defaultMaxObjectsPerPage = 1000
|
||||
|
||||
// NewProgressTracker creates a new progress tracker for S3 scanning operations.
|
||||
// The enabled parameter determines if progress tracking is active, and progress
|
||||
// NewCheckpointer creates a new checkpointer for S3 scanning operations.
|
||||
// The enabled parameter determines if checkpointing is active, and progress
|
||||
// provides the underlying mechanism for persisting scan state.
|
||||
func NewProgressTracker(ctx context.Context, enabled bool, progress *sources.Progress) *ProgressTracker {
|
||||
ctx.Logger().Info("Creating progress tracker")
|
||||
func NewCheckpointer(ctx context.Context, enabled bool, progress *sources.Progress) *Checkpointer {
|
||||
ctx.Logger().Info("Creating checkpointer")
|
||||
|
||||
return &ProgressTracker{
|
||||
return &Checkpointer{
|
||||
// We are resuming if we have completed objects from a previous scan.
|
||||
completedObjects: make([]bool, defaultMaxObjectsPerPage),
|
||||
completionOrder: make([]int, 0, defaultMaxObjectsPerPage),
|
||||
@@ -72,16 +78,18 @@ func NewProgressTracker(ctx context.Context, enabled bool, progress *sources.Pro
|
||||
}
|
||||
|
||||
// Reset prepares the tracker for a new page of objects by clearing the completion state.
|
||||
func (p *ProgressTracker) Reset() {
|
||||
// Must be called before processing each new page of objects.
|
||||
func (p *Checkpointer) Reset() {
|
||||
if !p.enabled {
|
||||
return
|
||||
}
|
||||
|
||||
p.Lock()
|
||||
defer p.Unlock()
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
// Store the current completed count before moving to next page.
|
||||
p.completedObjects = make([]bool, defaultMaxObjectsPerPage)
|
||||
p.completionOrder = make([]int, 0, defaultMaxObjectsPerPage)
|
||||
p.lowestIncompleteIdx = 0
|
||||
}
|
||||
|
||||
// ResumeInfo represents the state needed to resume an interrupted operation.
|
||||
@@ -92,11 +100,11 @@ type ResumeInfo struct {
|
||||
StartAfter string `json:"start_after"` // Last processed object key
|
||||
}
|
||||
|
||||
// GetResumePoint retrieves the last saved checkpoint state if one exists.
|
||||
// ResumePoint retrieves the last saved checkpoint state if one exists.
|
||||
// It returns nil if progress tracking is disabled or no resume state exists.
|
||||
// This method decodes the stored resume information and validates it contains
|
||||
// the minimum required data to enable resumption.
|
||||
func (p *ProgressTracker) GetResumePoint(ctx context.Context) (ResumeInfo, error) {
|
||||
func (p *Checkpointer) ResumePoint(ctx context.Context) (ResumeInfo, error) {
|
||||
resume := ResumeInfo{}
|
||||
|
||||
if !p.enabled || p.progress.EncodedResumeInfo == "" {
|
||||
@@ -118,7 +126,7 @@ func (p *ProgressTracker) GetResumePoint(ctx context.Context) (ResumeInfo, error
|
||||
|
||||
// Complete marks the entire scanning operation as finished and clears the resume state.
|
||||
// This should only be called once all scanning operations are complete.
|
||||
func (p *ProgressTracker) Complete(_ context.Context, message string) error {
|
||||
func (p *Checkpointer) Complete(_ context.Context, message string) error {
|
||||
// Preserve existing progress counters while clearing resume state.
|
||||
p.progress.SetProgressComplete(
|
||||
int(p.progress.SectionsCompleted),
|
||||
@@ -129,14 +137,11 @@ func (p *ProgressTracker) Complete(_ context.Context, message string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdateObjectProgress records successfully processed objects within the current page
|
||||
// UpdateObjectCompletion records successfully processed objects within the current page
|
||||
// and maintains fine-grained resumption checkpoints. It uses a conservative tracking
|
||||
// strategy that ensures no objects are missed by only checkpointing consecutively
|
||||
// completed objects.
|
||||
//
|
||||
// This method manages the detailed object-level progress tracking and creates
|
||||
// checkpoints that enable resumption of interrupted scans.
|
||||
//
|
||||
// This approach ensures scan reliability by only checkpointing consecutively completed
|
||||
// objects. While this may result in re-scanning some objects when resuming, it guarantees
|
||||
// no objects are missed in case of interruption.
|
||||
@@ -146,10 +151,13 @@ func (p *ProgressTracker) Complete(_ context.Context, message string) error {
|
||||
// - Objects completed: [0,1,2,3,4,5,7,8]
|
||||
// - The checkpoint will only include objects 0-5 since they are consecutive
|
||||
// - If scanning is interrupted and resumed:
|
||||
// - Scan resumes after object 5 (the last checkpoint)
|
||||
// - Objects 7-8 will be re-scanned even though they completed before
|
||||
// - This ensures object 6 is not missed
|
||||
func (p *ProgressTracker) UpdateObjectProgress(
|
||||
// -- Scan resumes after object 5 (the last checkpoint)
|
||||
// -- Objects 7-8 will be re-scanned even though they completed before
|
||||
// -- This ensures object 6 is not missed
|
||||
//
|
||||
// Thread-safe for concurrent object processing within a single page.
|
||||
// WARNING: Not safe for concurrent page processing.
|
||||
func (p *Checkpointer) UpdateObjectCompletion(
|
||||
ctx context.Context,
|
||||
completedIdx int,
|
||||
bucket string,
|
||||
@@ -166,46 +174,48 @@ func (p *ProgressTracker) UpdateObjectProgress(
|
||||
return fmt.Errorf("completed index %d exceeds maximum page size", completedIdx)
|
||||
}
|
||||
|
||||
p.Lock()
|
||||
defer p.Unlock()
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
|
||||
// Only track completion if this is the first time this index is marked complete.
|
||||
// Only process if this is the first time this index is marked complete.
|
||||
if !p.completedObjects[completedIdx] {
|
||||
p.completedObjects[completedIdx] = true
|
||||
p.completionOrder = append(p.completionOrder, completedIdx)
|
||||
}
|
||||
|
||||
// Find the highest safe checkpoint we can create.
|
||||
lastSafeIdx := -1
|
||||
var safeIndices [defaultMaxObjectsPerPage]bool
|
||||
|
||||
// Mark all completed indices.
|
||||
for _, idx := range p.completionOrder {
|
||||
safeIndices[idx] = true
|
||||
}
|
||||
|
||||
// Find the highest consecutive completed index.
|
||||
for i := range len(p.completedObjects) {
|
||||
if !safeIndices[i] {
|
||||
break
|
||||
// If we completed the lowest incomplete index, scan forward to find the new lowest.
|
||||
if completedIdx == p.lowestIncompleteIdx {
|
||||
p.advanceLowestIncompleteIdx()
|
||||
}
|
||||
lastSafeIdx = i
|
||||
}
|
||||
|
||||
// Update progress if we have at least one completed object.
|
||||
if lastSafeIdx < 0 {
|
||||
return nil
|
||||
// lowestIncompleteIdx points to first incomplete object, so everything before
|
||||
// it is complete. We want to checkpoint at the last complete object.
|
||||
checkpointIdx := p.lowestIncompleteIdx - 1
|
||||
if checkpointIdx < 0 {
|
||||
return nil // No completed objects yet
|
||||
}
|
||||
obj := pageContents[checkpointIdx]
|
||||
|
||||
obj := pageContents[lastSafeIdx]
|
||||
info := &ResumeInfo{CurrentBucket: bucket, StartAfter: *obj.Key}
|
||||
encoded, err := json.Marshal(info)
|
||||
return p.updateCheckpoint(bucket, *obj.Key)
|
||||
}
|
||||
|
||||
// advanceLowestIncompleteIdx moves the lowest incomplete index forward to the next incomplete object.
|
||||
// Must be called with lock held.
|
||||
func (p *Checkpointer) advanceLowestIncompleteIdx() {
|
||||
for p.lowestIncompleteIdx < len(p.completedObjects) &&
|
||||
p.completedObjects[p.lowestIncompleteIdx] {
|
||||
p.lowestIncompleteIdx++
|
||||
}
|
||||
}
|
||||
|
||||
// updateCheckpoint persists the current resumption state.
|
||||
// Must be called with lock held.
|
||||
func (p *Checkpointer) updateCheckpoint(bucket string, lastKey string) error {
|
||||
encoded, err := json.Marshal(&ResumeInfo{CurrentBucket: bucket, StartAfter: lastKey})
|
||||
if err != nil {
|
||||
return err
|
||||
return fmt.Errorf("failed to encode resume info: %w", err)
|
||||
}
|
||||
|
||||
// Purposefully avoid updating any progress counts.
|
||||
// Only update resume info.
|
||||
p.progress.SetProgressComplete(
|
||||
int(p.progress.SectionsCompleted),
|
||||
int(p.progress.SectionsRemaining),
|
||||
@@ -13,12 +13,12 @@ import (
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/sources"
|
||||
)
|
||||
|
||||
func TestProgressTrackerResumption(t *testing.T) {
|
||||
func TestCheckpointerResumption(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
// First scan - process 6 objects then interrupt.
|
||||
initialProgress := &sources.Progress{}
|
||||
tracker := NewProgressTracker(ctx, true, initialProgress)
|
||||
tracker := NewCheckpointer(ctx, true, initialProgress)
|
||||
|
||||
firstPage := &s3.ListObjectsV2Output{
|
||||
Contents: make([]*s3.Object, 12), // Total of 12 objects
|
||||
@@ -30,18 +30,18 @@ func TestProgressTrackerResumption(t *testing.T) {
|
||||
|
||||
// Process first 6 objects.
|
||||
for i := range 6 {
|
||||
err := tracker.UpdateObjectProgress(ctx, i, "test-bucket", firstPage.Contents)
|
||||
err := tracker.UpdateObjectCompletion(ctx, i, "test-bucket", firstPage.Contents)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
// Verify resume info is set correctly.
|
||||
resumeInfo, err := tracker.GetResumePoint(ctx)
|
||||
resumeInfo, err := tracker.ResumePoint(ctx)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "test-bucket", resumeInfo.CurrentBucket)
|
||||
assert.Equal(t, "key-5", resumeInfo.StartAfter)
|
||||
|
||||
// Resume scan with existing progress.
|
||||
resumeTracker := NewProgressTracker(ctx, true, initialProgress)
|
||||
resumeTracker := NewCheckpointer(ctx, true, initialProgress)
|
||||
|
||||
resumePage := &s3.ListObjectsV2Output{
|
||||
Contents: firstPage.Contents[6:], // Remaining 6 objects
|
||||
@@ -49,18 +49,18 @@ func TestProgressTrackerResumption(t *testing.T) {
|
||||
|
||||
// Process remaining objects.
|
||||
for i := range len(resumePage.Contents) {
|
||||
err := resumeTracker.UpdateObjectProgress(ctx, i, "test-bucket", resumePage.Contents)
|
||||
err := resumeTracker.UpdateObjectCompletion(ctx, i, "test-bucket", resumePage.Contents)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
// Verify final resume info.
|
||||
finalResumeInfo, err := resumeTracker.GetResumePoint(ctx)
|
||||
finalResumeInfo, err := resumeTracker.ResumePoint(ctx)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "test-bucket", finalResumeInfo.CurrentBucket)
|
||||
assert.Equal(t, "key-11", finalResumeInfo.StartAfter)
|
||||
}
|
||||
|
||||
func TestProgressTrackerReset(t *testing.T) {
|
||||
func TestCheckpointerReset(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
enabled bool
|
||||
@@ -75,7 +75,7 @@ func TestProgressTrackerReset(t *testing.T) {
|
||||
|
||||
ctx := context.Background()
|
||||
progress := new(sources.Progress)
|
||||
tracker := NewProgressTracker(ctx, tt.enabled, progress)
|
||||
tracker := NewCheckpointer(ctx, tt.enabled, progress)
|
||||
|
||||
tracker.completedObjects[1] = true
|
||||
tracker.completedObjects[2] = true
|
||||
@@ -150,9 +150,9 @@ func TestGetResumePoint(t *testing.T) {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tracker := &ProgressTracker{enabled: tt.enabled, progress: tt.progress}
|
||||
tracker := &Checkpointer{enabled: tt.enabled, progress: tt.progress}
|
||||
|
||||
resumePoint, err := tracker.GetResumePoint(context.Background())
|
||||
resumePoint, err := tracker.ResumePoint(context.Background())
|
||||
if tt.expectError {
|
||||
assert.Error(t, err, "Expected an error decoding resume info")
|
||||
} else {
|
||||
@@ -164,37 +164,50 @@ func TestGetResumePoint(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestProgressTrackerUpdateProgress(t *testing.T) {
|
||||
func TestCheckpointerUpdate(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
description string
|
||||
completedIdx int
|
||||
pageSize int
|
||||
preCompleted []int
|
||||
expectedKey string
|
||||
name string
|
||||
description string
|
||||
completedIdx int
|
||||
pageSize int
|
||||
preCompleted []int
|
||||
expectedKey string
|
||||
expectedLowestIncomplete int
|
||||
}{
|
||||
{
|
||||
name: "first object completed",
|
||||
description: "Basic case - completing first object",
|
||||
completedIdx: 0,
|
||||
pageSize: 3,
|
||||
expectedKey: "key-0",
|
||||
name: "first object completed",
|
||||
description: "Basic case - completing first object",
|
||||
completedIdx: 0,
|
||||
pageSize: 3,
|
||||
expectedKey: "key-0",
|
||||
expectedLowestIncomplete: 1,
|
||||
},
|
||||
{
|
||||
name: "completing missing middle",
|
||||
description: "Completing object when previous is done",
|
||||
completedIdx: 1,
|
||||
pageSize: 3,
|
||||
preCompleted: []int{0},
|
||||
expectedKey: "key-1",
|
||||
name: "completing missing middle",
|
||||
description: "Completing object when previous is done",
|
||||
completedIdx: 1,
|
||||
pageSize: 3,
|
||||
preCompleted: []int{0},
|
||||
expectedKey: "key-1",
|
||||
expectedLowestIncomplete: 2,
|
||||
},
|
||||
{
|
||||
name: "all objects completed in order",
|
||||
description: "Completing final object in sequence",
|
||||
completedIdx: 2,
|
||||
pageSize: 3,
|
||||
preCompleted: []int{0, 1},
|
||||
expectedKey: "key-2",
|
||||
name: "all objects completed in order",
|
||||
description: "Completing final object in sequence",
|
||||
completedIdx: 2,
|
||||
pageSize: 3,
|
||||
preCompleted: []int{0, 1},
|
||||
expectedKey: "key-2",
|
||||
expectedLowestIncomplete: 3,
|
||||
},
|
||||
{
|
||||
name: "out of order completion before lowest",
|
||||
description: "Completing object before current lowest incomplete - should not affect checkpoint",
|
||||
completedIdx: 1,
|
||||
pageSize: 4,
|
||||
preCompleted: []int{0, 2, 3},
|
||||
expectedKey: "key-3",
|
||||
expectedLowestIncomplete: 4,
|
||||
},
|
||||
{
|
||||
name: "last index in max page",
|
||||
@@ -203,12 +216,13 @@ func TestProgressTrackerUpdateProgress(t *testing.T) {
|
||||
pageSize: 1000,
|
||||
preCompleted: func() []int {
|
||||
indices := make([]int, 999)
|
||||
for i := range 999 {
|
||||
for i := range indices {
|
||||
indices[i] = i
|
||||
}
|
||||
return indices
|
||||
}(),
|
||||
expectedKey: "key-999",
|
||||
expectedKey: "key-999",
|
||||
expectedLowestIncomplete: 1000,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -218,11 +232,12 @@ func TestProgressTrackerUpdateProgress(t *testing.T) {
|
||||
|
||||
ctx := context.Background()
|
||||
progress := new(sources.Progress)
|
||||
tracker := &ProgressTracker{
|
||||
enabled: true,
|
||||
progress: progress,
|
||||
completedObjects: make([]bool, tt.pageSize),
|
||||
completionOrder: make([]int, 0, tt.pageSize),
|
||||
tracker := &Checkpointer{
|
||||
enabled: true,
|
||||
progress: progress,
|
||||
completedObjects: make([]bool, tt.pageSize),
|
||||
completionOrder: make([]int, 0, tt.pageSize),
|
||||
lowestIncompleteIdx: 0,
|
||||
}
|
||||
|
||||
page := &s3.ListObjectsV2Output{Contents: make([]*s3.Object, tt.pageSize)}
|
||||
@@ -231,21 +246,30 @@ func TestProgressTrackerUpdateProgress(t *testing.T) {
|
||||
page.Contents[i] = &s3.Object{Key: &key}
|
||||
}
|
||||
|
||||
// Apply pre-completed indices in order.
|
||||
if tt.preCompleted != nil {
|
||||
for _, idx := range tt.preCompleted {
|
||||
tracker.completedObjects[idx] = true
|
||||
tracker.completionOrder = append(tracker.completionOrder, idx)
|
||||
// Setup pre-completed objects.
|
||||
for _, idx := range tt.preCompleted {
|
||||
tracker.completedObjects[idx] = true
|
||||
tracker.completionOrder = append(tracker.completionOrder, idx)
|
||||
}
|
||||
|
||||
// Find the correct lowest incomplete index after pre-completion.
|
||||
for i := range tt.pageSize {
|
||||
if !tracker.completedObjects[i] {
|
||||
tracker.lowestIncompleteIdx = i
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
err := tracker.UpdateObjectProgress(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
|
||||
err = json.Unmarshal([]byte(progress.EncodedResumeInfo), &info)
|
||||
assert.NoError(t, err, "Failed to decode resume info")
|
||||
assert.Equal(t, tt.expectedKey, info.StartAfter, "Incorrect resume point")
|
||||
|
||||
assert.Equal(t, tt.expectedLowestIncomplete, tracker.lowestIncompleteIdx,
|
||||
"Incorrect lowest incomplete index")
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -313,7 +337,7 @@ func TestComplete(t *testing.T) {
|
||||
EncodedResumeInfo: tt.initialState.resumeInfo,
|
||||
Message: tt.initialState.message,
|
||||
}
|
||||
tracker := NewProgressTracker(ctx, tt.enabled, progress)
|
||||
tracker := NewCheckpointer(ctx, tt.enabled, progress)
|
||||
|
||||
err := tracker.Complete(ctx, tt.completeMessage)
|
||||
assert.NoError(t, err)
|
||||
+180
-35
@@ -2,6 +2,7 @@ package s3
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
@@ -43,8 +44,10 @@ type Source struct {
|
||||
jobID sources.JobID
|
||||
verify bool
|
||||
concurrency int
|
||||
conn *sourcespb.S3
|
||||
|
||||
checkpointer *Checkpointer
|
||||
sources.Progress
|
||||
conn *sourcespb.S3
|
||||
|
||||
errorCount *sync.Map
|
||||
jobPool *errgroup.Group
|
||||
@@ -67,7 +70,7 @@ func (s *Source) JobID() sources.JobID { return s.jobID }
|
||||
|
||||
// Init returns an initialized AWS source
|
||||
func (s *Source) Init(
|
||||
_ context.Context,
|
||||
ctx context.Context,
|
||||
name string,
|
||||
jobID sources.JobID,
|
||||
sourceID sources.SourceID,
|
||||
@@ -90,6 +93,8 @@ func (s *Source) Init(
|
||||
}
|
||||
s.conn = &conn
|
||||
|
||||
s.checkpointer = NewCheckpointer(ctx, conn.GetEnableResumption(), &s.Progress)
|
||||
|
||||
s.setMaxObjectSize(conn.GetMaxObjectSize())
|
||||
|
||||
if len(conn.GetBuckets()) > 0 && len(conn.GetIgnoreBuckets()) > 0 {
|
||||
@@ -173,9 +178,16 @@ func (s *Source) newClient(region, roleArn string) (*s3.S3, error) {
|
||||
return s3.New(sess), nil
|
||||
}
|
||||
|
||||
// IAM identity needs s3:ListBuckets permission
|
||||
// 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
|
||||
}
|
||||
|
||||
@@ -196,9 +208,73 @@ func (s *Source) getBucketsToScan(client *s3.S3) ([]string, error) {
|
||||
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,
|
||||
@@ -206,22 +282,48 @@ func (s *Source) scanBuckets(
|
||||
bucketsToScan []string,
|
||||
chunksChan chan *sources.Chunk,
|
||||
) {
|
||||
var objectCount uint64
|
||||
|
||||
if role != "" {
|
||||
ctx = context.WithValue(ctx, "role", role)
|
||||
}
|
||||
var objectCount uint64
|
||||
|
||||
for i, bucket := range bucketsToScan {
|
||||
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++ {
|
||||
bucket := bucketsToScan[bucketIdx]
|
||||
ctx := context.WithValue(ctx, "bucket", bucket)
|
||||
|
||||
if common.IsDone(ctx) {
|
||||
ctx.Logger().Error(ctx.Err(), "context done, while scanning bucket")
|
||||
return
|
||||
}
|
||||
|
||||
s.SetProgressComplete(i, len(bucketsToScan), fmt.Sprintf("Bucket: %s", bucket), "")
|
||||
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")
|
||||
@@ -230,10 +332,33 @@ func (s *Source) scanBuckets(
|
||||
|
||||
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, &s3.ListObjectsV2Input{Bucket: &bucket},
|
||||
ctx,
|
||||
input,
|
||||
func(page *s3.ListObjectsV2Output, _ bool) bool {
|
||||
s.pageChunker(ctx, regionalClient, chunksChan, bucket, page, &errorCount, i+1, &objectCount)
|
||||
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
|
||||
})
|
||||
|
||||
@@ -249,6 +374,7 @@ func (s *Source) scanBuckets(
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
s.SetProgressComplete(
|
||||
len(bucketsToScan),
|
||||
len(bucketsToScan),
|
||||
@@ -289,29 +415,25 @@ func (s *Source) getRegionalClientForBucket(
|
||||
return regionalClient, nil
|
||||
}
|
||||
|
||||
// pageChunker emits chunks onto the given channel from a page
|
||||
// pageChunker emits chunks onto the given channel from a page.
|
||||
func (s *Source) pageChunker(
|
||||
ctx context.Context,
|
||||
client *s3.S3,
|
||||
metadata pageMetadata,
|
||||
state processingState,
|
||||
chunksChan chan *sources.Chunk,
|
||||
bucket string,
|
||||
page *s3.ListObjectsV2Output,
|
||||
errorCount *sync.Map,
|
||||
pageNumber int,
|
||||
objectCount *uint64,
|
||||
) {
|
||||
for _, obj := range page.Contents {
|
||||
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 {
|
||||
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,
|
||||
"bucket", bucket,
|
||||
"page", pageNumber,
|
||||
"size", *obj.Size,
|
||||
)
|
||||
ctx = context.WithValues(ctx, "key", *obj.Key, "size", *obj.Size)
|
||||
|
||||
if common.IsDone(ctx) {
|
||||
return
|
||||
@@ -320,29 +442,44 @@ func (s *Source) pageChunker(
|
||||
// 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)
|
||||
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)")
|
||||
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")
|
||||
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")
|
||||
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")
|
||||
@@ -352,7 +489,7 @@ func (s *Source) pageChunker(
|
||||
path := strings.Split(*obj.Key, "/")
|
||||
prefix := strings.Join(path[:len(path)-1], "/")
|
||||
|
||||
nErr, ok := errorCount.Load(prefix)
|
||||
nErr, ok := state.errorCount.Load(prefix)
|
||||
if !ok {
|
||||
nErr = 0
|
||||
}
|
||||
@@ -366,8 +503,8 @@ func (s *Source) pageChunker(
|
||||
objCtx, cancel := context.WithTimeout(ctx, getObjectTimeout)
|
||||
defer cancel()
|
||||
|
||||
res, err := client.GetObjectWithContext(objCtx, &s3.GetObjectInput{
|
||||
Bucket: &bucket,
|
||||
res, err := metadata.client.GetObjectWithContext(objCtx, &s3.GetObjectInput{
|
||||
Bucket: &metadata.bucket,
|
||||
Key: obj.Key,
|
||||
})
|
||||
if err != nil {
|
||||
@@ -382,7 +519,7 @@ func (s *Source) pageChunker(
|
||||
res.Body.Close()
|
||||
}
|
||||
|
||||
nErr, ok := errorCount.Load(prefix)
|
||||
nErr, ok := state.errorCount.Load(prefix)
|
||||
if !ok {
|
||||
nErr = 0
|
||||
}
|
||||
@@ -391,7 +528,7 @@ func (s *Source) pageChunker(
|
||||
return nil
|
||||
}
|
||||
nErr = nErr.(int) + 1
|
||||
errorCount.Store(prefix, nErr)
|
||||
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)
|
||||
@@ -413,9 +550,9 @@ func (s *Source) pageChunker(
|
||||
SourceMetadata: &source_metadatapb.MetaData{
|
||||
Data: &source_metadatapb.MetaData_S3{
|
||||
S3: &source_metadatapb.S3{
|
||||
Bucket: bucket,
|
||||
Bucket: metadata.bucket,
|
||||
File: sanitizer.UTF8(*obj.Key),
|
||||
Link: sanitizer.UTF8(makeS3Link(bucket, *client.Config.Region, *obj.Key)),
|
||||
Link: sanitizer.UTF8(makeS3Link(metadata.bucket, *metadata.client.Config.Region, *obj.Key)),
|
||||
Email: sanitizer.UTF8(email),
|
||||
Timestamp: sanitizer.UTF8(modified),
|
||||
},
|
||||
@@ -429,14 +566,19 @@ func (s *Source) pageChunker(
|
||||
return nil
|
||||
}
|
||||
|
||||
atomic.AddUint64(objectCount, 1)
|
||||
ctx.Logger().V(5).Info("S3 object scanned.", "object_count", objectCount)
|
||||
nErr, ok = errorCount.Load(prefix)
|
||||
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 {
|
||||
errorCount.Store(prefix, 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")
|
||||
}
|
||||
|
||||
return nil
|
||||
@@ -485,6 +627,9 @@ func (s *Source) validateBucketAccess(ctx context.Context, client *s3.S3, roleAr
|
||||
// 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,
|
||||
|
||||
@@ -10,11 +10,12 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/protobuf/types/known/anypb"
|
||||
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/common"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/pb/credentialspb"
|
||||
"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"
|
||||
)
|
||||
@@ -215,3 +216,162 @@ func TestSource_Validate(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSourceChunksNoResumption(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second*30)
|
||||
defer cancel()
|
||||
|
||||
s := Source{}
|
||||
connection := &sourcespb.S3{
|
||||
Credential: &sourcespb.S3_Unauthenticated{},
|
||||
Buckets: []string{"integration-resumption-tests"},
|
||||
}
|
||||
conn, err := anypb.New(connection)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
err = s.Init(ctx, "test name", 0, 0, false, conn, 1)
|
||||
chunksCh := make(chan *sources.Chunk)
|
||||
go func() {
|
||||
defer close(chunksCh)
|
||||
err = s.Chunks(ctx, chunksCh)
|
||||
assert.Nil(t, err)
|
||||
}()
|
||||
|
||||
wantChunkCount := 19787
|
||||
got := 0
|
||||
|
||||
for range chunksCh {
|
||||
got++
|
||||
}
|
||||
assert.Equal(t, wantChunkCount, got)
|
||||
}
|
||||
|
||||
func TestSourceChunksResumption(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
src := new(Source)
|
||||
src.Progress = sources.Progress{
|
||||
Message: "Bucket: integration-resumption-tests",
|
||||
EncodedResumeInfo: "{\"current_bucket\":\"integration-resumption-tests\",\"start_after\":\"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 = src.Init(ctx, "test name", 0, 0, false, conn, 2)
|
||||
require.NoError(t, err)
|
||||
|
||||
chunksCh := make(chan *sources.Chunk)
|
||||
var count int
|
||||
|
||||
cancelCtx, ctxCancel := context.WithCancel(ctx)
|
||||
defer ctxCancel()
|
||||
|
||||
go func() {
|
||||
defer close(chunksCh)
|
||||
err = src.Chunks(cancelCtx, chunksCh)
|
||||
assert.NoError(t, err, "Should not error during scan")
|
||||
}()
|
||||
|
||||
for range chunksCh {
|
||||
count++
|
||||
}
|
||||
|
||||
// Verify that we processed all remaining data on resume.
|
||||
// Also verify that we processed less than the total number of chunks for the source.
|
||||
sourceTotalChunkCount := 19787
|
||||
assert.Equal(t, 9638, count, "Should have processed all remaining data on resume")
|
||||
assert.Less(t, count, sourceTotalChunkCount, "Should have processed less than total chunks on resume")
|
||||
}
|
||||
|
||||
func TestSourceChunksNoResumptionMultipleBuckets(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second*30)
|
||||
defer cancel()
|
||||
|
||||
s := Source{}
|
||||
connection := &sourcespb.S3{
|
||||
Credential: &sourcespb.S3_Unauthenticated{},
|
||||
Buckets: []string{"integration-resumption-tests", "truffletestbucket"},
|
||||
}
|
||||
conn, err := anypb.New(connection)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
err = s.Init(ctx, "test name", 0, 0, false, conn, 1)
|
||||
chunksCh := make(chan *sources.Chunk)
|
||||
go func() {
|
||||
defer close(chunksCh)
|
||||
err = s.Chunks(ctx, chunksCh)
|
||||
assert.Nil(t, err)
|
||||
}()
|
||||
|
||||
wantChunkCount := 19890
|
||||
got := 0
|
||||
|
||||
for range chunksCh {
|
||||
got++
|
||||
}
|
||||
assert.Equal(t, wantChunkCount, got)
|
||||
}
|
||||
|
||||
func TestSourceChunksResumptionMultipleBucketsIgnoredBucket(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
src := new(Source)
|
||||
|
||||
// The bucket stored in EncodedResumeInfo is NOT in the list of buckets to scan.
|
||||
// Therefore, resume from the other provided bucket (truffletestbucket).
|
||||
src.Progress = sources.Progress{
|
||||
Message: "Bucket: integration-resumption-tests",
|
||||
EncodedResumeInfo: "{\"current_bucket\":\"integration-resumption-tests\",\"start_after\":\"test-dir/\"}",
|
||||
SectionsCompleted: 0,
|
||||
SectionsRemaining: 1,
|
||||
}
|
||||
connection := &sourcespb.S3{
|
||||
Credential: &sourcespb.S3_Unauthenticated{},
|
||||
Buckets: []string{"truffletestbucket"},
|
||||
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)
|
||||
|
||||
chunksCh := make(chan *sources.Chunk)
|
||||
var count int
|
||||
|
||||
cancelCtx, ctxCancel := context.WithCancel(ctx)
|
||||
defer ctxCancel()
|
||||
|
||||
go func() {
|
||||
defer close(chunksCh)
|
||||
err = src.Chunks(cancelCtx, chunksCh)
|
||||
assert.NoError(t, err, "Should not error during scan")
|
||||
}()
|
||||
|
||||
for range chunksCh {
|
||||
count++
|
||||
}
|
||||
|
||||
assert.Equal(t, 103, count, "Should have processed all remaining data on resume")
|
||||
}
|
||||
|
||||
@@ -10,12 +10,13 @@ import (
|
||||
|
||||
"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"
|
||||
"google.golang.org/protobuf/types/known/anypb"
|
||||
)
|
||||
|
||||
func TestSource_Init_IncludeAndIgnoreBucketsError(t *testing.T) {
|
||||
|
||||
@@ -123,10 +123,10 @@ func (s *SourceManager) GetIDs(ctx context.Context, sourceName string, kind sour
|
||||
return s.api.GetIDs(ctx, sourceName, kind)
|
||||
}
|
||||
|
||||
// Run blocks until a resource is available to run the source, then
|
||||
// EnumerateAndScan blocks until a resource is available to run the source, then
|
||||
// asynchronously runs it. Error information is stored and accessible via the
|
||||
// JobProgressRef as it becomes available.
|
||||
func (s *SourceManager) Run(ctx context.Context, sourceName string, source Source, targets ...ChunkingTarget) (JobProgressRef, error) {
|
||||
func (s *SourceManager) EnumerateAndScan(ctx context.Context, sourceName string, source Source, targets ...ChunkingTarget) (JobProgressRef, error) {
|
||||
sourceID, jobID := source.SourceID(), source.JobID()
|
||||
// Do preflight checks before waiting on the pool.
|
||||
if err := s.preflightChecks(ctx); err != nil {
|
||||
@@ -169,6 +169,98 @@ func (s *SourceManager) Run(ctx context.Context, sourceName string, source Sourc
|
||||
return progress.Ref(), nil
|
||||
}
|
||||
|
||||
func (s *SourceManager) Enumerate(ctx context.Context, sourceName string, source Source, reporter UnitReporter) (JobProgressRef, error) {
|
||||
sourceID, jobID := source.SourceID(), source.JobID()
|
||||
// Do preflight checks before waiting on the pool.
|
||||
if err := s.preflightChecks(ctx); err != nil {
|
||||
return JobProgressRef{
|
||||
SourceName: sourceName,
|
||||
SourceID: sourceID,
|
||||
JobID: jobID,
|
||||
}, err
|
||||
}
|
||||
|
||||
// Create a JobProgress object for tracking progress.
|
||||
sem := s.sem
|
||||
ctx, cancel := context.WithCancelCause(ctx)
|
||||
progress := NewJobProgress(jobID, sourceID, sourceName, WithHooks(s.hooks...), WithCancel(cancel))
|
||||
if err := sem.Acquire(ctx, 1); err != nil {
|
||||
// Context cancelled.
|
||||
progress.ReportError(Fatal{err})
|
||||
return progress.Ref(), Fatal{err}
|
||||
}
|
||||
|
||||
// Wrap the passed in reporter so we update the progress information.
|
||||
reporter = baseUnitReporter{
|
||||
child: reporter,
|
||||
progress: progress,
|
||||
}
|
||||
|
||||
s.wg.Add(1)
|
||||
go func() {
|
||||
// Call Finish after the semaphore has been released.
|
||||
defer progress.Finish()
|
||||
defer sem.Release(1)
|
||||
defer s.wg.Done()
|
||||
ctx := context.WithValues(ctx,
|
||||
"source_manager_worker_id", common.RandomID(5),
|
||||
)
|
||||
defer common.Recover(ctx)
|
||||
defer cancel(nil)
|
||||
if err := s.enumerate(ctx, source, progress, reporter); err != nil {
|
||||
select {
|
||||
case s.firstErr <- err:
|
||||
default:
|
||||
}
|
||||
progress.ReportError(Fatal{err})
|
||||
}
|
||||
}()
|
||||
return progress.Ref(), nil
|
||||
}
|
||||
|
||||
// Scan blocks until a resource is available to run the source against a single
|
||||
// SourceUnit, then asynchronously runs it. Error information is stored and
|
||||
// accessible via the JobProgressRef as it becomes available.
|
||||
func (s *SourceManager) Scan(ctx context.Context, sourceName string, source Source, unit SourceUnit) (JobProgressRef, error) {
|
||||
sourceID, jobID := source.SourceID(), source.JobID()
|
||||
// Do preflight checks before waiting on the pool.
|
||||
if err := s.preflightChecks(ctx); err != nil {
|
||||
return JobProgressRef{
|
||||
SourceName: sourceName,
|
||||
SourceID: sourceID,
|
||||
JobID: jobID,
|
||||
}, err
|
||||
}
|
||||
// Create a JobProgress object for tracking progress.
|
||||
ctx, cancel := context.WithCancelCause(ctx)
|
||||
progress := NewJobProgress(jobID, sourceID, sourceName, WithHooks(s.hooks...), WithCancel(cancel))
|
||||
if err := s.sem.Acquire(ctx, 1); err != nil {
|
||||
// Context cancelled.
|
||||
progress.ReportError(Fatal{err})
|
||||
return progress.Ref(), Fatal{err}
|
||||
}
|
||||
s.wg.Add(1)
|
||||
go func() {
|
||||
// Call Finish after the semaphore has been released.
|
||||
defer progress.Finish()
|
||||
defer s.sem.Release(1)
|
||||
defer s.wg.Done()
|
||||
ctx := context.WithValues(ctx,
|
||||
"source_manager_worker_id", common.RandomID(5),
|
||||
)
|
||||
defer common.Recover(ctx)
|
||||
defer cancel(nil)
|
||||
if err := s.scan(ctx, source, progress, unit); err != nil {
|
||||
select {
|
||||
case s.firstErr <- err:
|
||||
default:
|
||||
}
|
||||
progress.ReportError(Fatal{err})
|
||||
}
|
||||
}()
|
||||
return progress.Ref(), nil
|
||||
}
|
||||
|
||||
// Chunks returns the read only channel of all the chunks produced by all of
|
||||
// the sources managed by this manager.
|
||||
func (s *SourceManager) Chunks() <-chan *Chunk {
|
||||
@@ -272,7 +364,7 @@ func (s *SourceManager) run(ctx context.Context, source Source, report *JobProgr
|
||||
ctx = context.WithValue(ctx, "source_type", source.Type().String())
|
||||
}
|
||||
|
||||
// Check for the preferred method of tracking source units.
|
||||
// Check if source units are supported and configured.
|
||||
canUseSourceUnits := len(targets) == 0 && s.useSourceUnitsFunc != nil
|
||||
if enumChunker, ok := source.(SourceUnitEnumChunker); ok && canUseSourceUnits && s.useSourceUnitsFunc() {
|
||||
ctx.Logger().Info("running source",
|
||||
@@ -286,6 +378,111 @@ func (s *SourceManager) run(ctx context.Context, source Source, report *JobProgr
|
||||
return s.runWithoutUnits(ctx, source, report, targets...)
|
||||
}
|
||||
|
||||
// enumerate is a helper method to enumerate a Source.
|
||||
func (s *SourceManager) enumerate(ctx context.Context, source Source, report *JobProgress, reporter UnitReporter) error {
|
||||
report.Start(time.Now())
|
||||
defer func() { report.End(time.Now()) }()
|
||||
|
||||
defer func() {
|
||||
if err := context.Cause(ctx); err != nil {
|
||||
report.ReportError(Fatal{err})
|
||||
}
|
||||
}()
|
||||
|
||||
report.TrackProgress(source.GetProgress())
|
||||
if ctx.Value("job_id") == "" {
|
||||
ctx = context.WithValue(ctx, "job_id", report.JobID)
|
||||
}
|
||||
if ctx.Value("source_id") == "" {
|
||||
ctx = context.WithValue(ctx, "source_id", report.SourceID)
|
||||
}
|
||||
if ctx.Value("source_name") == "" {
|
||||
ctx = context.WithValue(ctx, "source_name", report.SourceName)
|
||||
}
|
||||
if ctx.Value("source_type") == "" {
|
||||
ctx = context.WithValue(ctx, "source_type", source.Type().String())
|
||||
}
|
||||
|
||||
// Check if source units are supported and configured.
|
||||
canUseSourceUnits := s.useSourceUnitsFunc != nil
|
||||
if enumChunker, ok := source.(SourceUnitEnumerator); ok && canUseSourceUnits && s.useSourceUnitsFunc() {
|
||||
ctx.Logger().Info("running source",
|
||||
"with_units", true)
|
||||
return s.enumerateWithUnits(ctx, enumChunker, report, reporter)
|
||||
}
|
||||
return fmt.Errorf("Enumeration not supported or configured for source: %s", source.Type().String())
|
||||
}
|
||||
|
||||
// scan runs a scan against a single SourceUnit as its only job. This method
|
||||
// manages the lifecycle of the provided report.
|
||||
func (s *SourceManager) scan(ctx context.Context, source Source, report *JobProgress, unit SourceUnit) error {
|
||||
report.Start(time.Now())
|
||||
defer func() { report.End(time.Now()) }()
|
||||
|
||||
defer func() {
|
||||
if err := context.Cause(ctx); err != nil {
|
||||
report.ReportError(Fatal{err})
|
||||
}
|
||||
}()
|
||||
|
||||
report.TrackProgress(source.GetProgress())
|
||||
if ctx.Value("job_id") == "" {
|
||||
ctx = context.WithValue(ctx, "job_id", report.JobID)
|
||||
}
|
||||
if ctx.Value("source_id") == "" {
|
||||
ctx = context.WithValue(ctx, "source_id", report.SourceID)
|
||||
}
|
||||
if ctx.Value("source_name") == "" {
|
||||
ctx = context.WithValue(ctx, "source_name", report.SourceName)
|
||||
}
|
||||
if ctx.Value("source_type") == "" {
|
||||
ctx = context.WithValue(ctx, "source_type", source.Type().String())
|
||||
}
|
||||
|
||||
// Check if source units are supported and configured.
|
||||
canUseSourceUnits := s.useSourceUnitsFunc != nil
|
||||
if unitChunker, ok := source.(SourceUnitChunker); ok && canUseSourceUnits && s.useSourceUnitsFunc() {
|
||||
ctx.Logger().Info("running source",
|
||||
"with_units", true)
|
||||
return s.scanWithUnit(ctx, unitChunker, report, unit)
|
||||
}
|
||||
return fmt.Errorf("source units not supported or configured for source: %s (%s)", report.SourceName, source.Type().String())
|
||||
}
|
||||
|
||||
// enumerateWithUnits is a helper method to enumerate a Source that is also a
|
||||
// SourceUnitEnumerator. This allows better introspection of what is getting
|
||||
// enumerated and any errors encountered.
|
||||
func (s *SourceManager) enumerateWithUnits(ctx context.Context, source SourceUnitEnumerator, report *JobProgress, reporter UnitReporter) error {
|
||||
// Create a function that will save the first error encountered (if
|
||||
// any) and discard the rest.
|
||||
fatalErr := make(chan error, 1)
|
||||
catchFirstFatal := func(err error) {
|
||||
select {
|
||||
case fatalErr <- err:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
// Produce units.
|
||||
func() {
|
||||
// TODO: Catch panics and add to report.
|
||||
report.StartEnumerating(time.Now())
|
||||
defer func() { report.EndEnumerating(time.Now()) }()
|
||||
ctx.Logger().V(2).Info("enumerating source with units")
|
||||
if err := source.Enumerate(ctx, reporter); err != nil {
|
||||
report.ReportError(Fatal{err})
|
||||
catchFirstFatal(Fatal{err})
|
||||
}
|
||||
}()
|
||||
|
||||
select {
|
||||
case err := <-fatalErr:
|
||||
return err
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// runWithoutUnits is a helper method to run a Source. It has coarse-grained
|
||||
// job reporting.
|
||||
func (s *SourceManager) runWithoutUnits(ctx context.Context, source Source, report *JobProgress, targets ...ChunkingTarget) error {
|
||||
@@ -302,6 +499,7 @@ func (s *SourceManager) runWithoutUnits(ctx context.Context, source Source, repo
|
||||
s.outputChunks <- chunk
|
||||
}
|
||||
}()
|
||||
|
||||
// Don't return from this function until the goroutine has finished
|
||||
// outputting chunks to the downstream channel. Closing the channel
|
||||
// will stop the goroutine, so that needs to happen first in the defer
|
||||
@@ -393,6 +591,44 @@ func (s *SourceManager) runWithUnits(ctx context.Context, source SourceUnitEnumC
|
||||
}
|
||||
}
|
||||
|
||||
// scanWithUnit produces chunks from a single SourceUnit.
|
||||
func (s *SourceManager) scanWithUnit(ctx context.Context, source SourceUnitChunker, report *JobProgress, unit SourceUnit) error {
|
||||
// Create a function that will save the first error encountered (if
|
||||
// any) and discard the rest.
|
||||
chunkReporter := &mgrChunkReporter{
|
||||
unit: unit,
|
||||
chunkCh: make(chan *Chunk, defaultChannelSize),
|
||||
report: report,
|
||||
}
|
||||
// Produce chunks from the given unit.
|
||||
var chunkErr error
|
||||
go func() {
|
||||
report.StartUnitChunking(unit, time.Now())
|
||||
// TODO: Catch panics and add to report.
|
||||
defer close(chunkReporter.chunkCh)
|
||||
id, kind := unit.SourceUnitID()
|
||||
ctx := context.WithValues(ctx, "unit_kind", kind, "unit", id)
|
||||
ctx.Logger().V(3).Info("chunking unit")
|
||||
if err := source.ChunkUnit(ctx, unit, chunkReporter); err != nil {
|
||||
report.ReportError(Fatal{ChunkError{Unit: unit, Err: err}})
|
||||
chunkErr = Fatal{err}
|
||||
}
|
||||
}()
|
||||
// Consume chunks and export chunks.
|
||||
// This anonymous function blocks until the chunkReporter.chunkCh is
|
||||
// closed in the above goroutine.
|
||||
func() {
|
||||
defer func() { report.EndUnitChunking(unit, time.Now()) }()
|
||||
for chunk := range chunkReporter.chunkCh {
|
||||
if src, ok := source.(Source); ok {
|
||||
chunk.JobID = src.JobID()
|
||||
}
|
||||
s.outputChunks <- chunk
|
||||
}
|
||||
}()
|
||||
return chunkErr
|
||||
}
|
||||
|
||||
// headlessAPI implements the apiClient interface locally.
|
||||
type headlessAPI struct {
|
||||
// Counters for assigning source and job IDs.
|
||||
|
||||
@@ -114,7 +114,7 @@ func TestSourceManagerRun(t *testing.T) {
|
||||
source, err := buildDummy(&counterChunker{count: 1})
|
||||
assert.NoError(t, err)
|
||||
for i := 0; i < 3; i++ {
|
||||
ref, err := mgr.Run(context.Background(), "dummy", source)
|
||||
ref, err := mgr.EnumerateAndScan(context.Background(), "dummy", source)
|
||||
<-ref.Done()
|
||||
assert.NoError(t, err)
|
||||
assert.NoError(t, ref.Snapshot().FatalError())
|
||||
@@ -132,7 +132,7 @@ func TestSourceManagerWait(t *testing.T) {
|
||||
source, err := buildDummy(&counterChunker{count: 1})
|
||||
assert.NoError(t, err)
|
||||
// Asynchronously run the source.
|
||||
_, err = mgr.Run(context.Background(), "dummy", source)
|
||||
_, err = mgr.EnumerateAndScan(context.Background(), "dummy", source)
|
||||
assert.NoError(t, err)
|
||||
// Read the 1 chunk we're expecting so Waiting completes.
|
||||
<-mgr.Chunks()
|
||||
@@ -141,7 +141,7 @@ func TestSourceManagerWait(t *testing.T) {
|
||||
// Run should return an error now.
|
||||
_, err = buildDummy(&counterChunker{count: 1})
|
||||
assert.NoError(t, err)
|
||||
_, err = mgr.Run(context.Background(), "dummy", source)
|
||||
_, err = mgr.EnumerateAndScan(context.Background(), "dummy", source)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
@@ -149,7 +149,7 @@ func TestSourceManagerError(t *testing.T) {
|
||||
mgr := NewManager()
|
||||
source, err := buildDummy(errorChunker{fmt.Errorf("oops")})
|
||||
assert.NoError(t, err)
|
||||
ref, err := mgr.Run(context.Background(), "dummy", source)
|
||||
ref, err := mgr.EnumerateAndScan(context.Background(), "dummy", source)
|
||||
assert.NoError(t, err)
|
||||
<-ref.Done()
|
||||
assert.Error(t, ref.Snapshot().FatalError())
|
||||
@@ -165,7 +165,7 @@ func TestSourceManagerReport(t *testing.T) {
|
||||
mgr := NewManager(opts...)
|
||||
source, err := buildDummy(&counterChunker{count: 4})
|
||||
assert.NoError(t, err)
|
||||
ref, err := mgr.Run(context.Background(), "dummy", source)
|
||||
ref, err := mgr.EnumerateAndScan(context.Background(), "dummy", source)
|
||||
assert.NoError(t, err)
|
||||
<-ref.Done()
|
||||
assert.Equal(t, 0, len(ref.Snapshot().Errors))
|
||||
@@ -173,6 +173,56 @@ func TestSourceManagerReport(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestSourceManagerEnumerate(t *testing.T) {
|
||||
mgr := NewManager(WithBufferedOutput(8), WithSourceUnits())
|
||||
source, err := buildDummy(&counterChunker{count: 1})
|
||||
assert.NoError(t, err)
|
||||
var enumeratedUnits []SourceUnit
|
||||
reporter := visitorUnitReporter{
|
||||
ok: func(_ context.Context, unit SourceUnit) error {
|
||||
enumeratedUnits = append(enumeratedUnits, unit)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
for i := 0; i < 3; i++ {
|
||||
ref, err := mgr.Enumerate(context.Background(), "dummy", source, reporter)
|
||||
<-ref.Done()
|
||||
assert.NoError(t, err)
|
||||
assert.NoError(t, ref.Snapshot().FatalError())
|
||||
// The Chunks channel should be empty because we only enumerated.
|
||||
_, err = tryRead(mgr.Chunks())
|
||||
assert.Error(t, err)
|
||||
// Each time the loop iterates, we add 1 unit to the slice.
|
||||
assert.Equal(t, i+1, len(enumeratedUnits), ref.Snapshot())
|
||||
}
|
||||
}
|
||||
|
||||
func TestSourceManagerScan(t *testing.T) {
|
||||
mgr := NewManager(WithBufferedOutput(8), WithSourceUnits())
|
||||
source, err := buildDummy(&counterChunker{count: 1})
|
||||
assert.NoError(t, err)
|
||||
for i := 0; i < 3; i++ {
|
||||
ref, err := mgr.Scan(context.Background(), "dummy", source, countChunk(123))
|
||||
<-ref.Done()
|
||||
assert.NoError(t, err)
|
||||
assert.NoError(t, ref.Snapshot().FatalError())
|
||||
chunk, err := tryRead(mgr.Chunks())
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, []byte{123}, chunk.Data)
|
||||
// The Chunks channel should be empty now.
|
||||
_, err = tryRead(mgr.Chunks())
|
||||
assert.Error(t, err)
|
||||
}
|
||||
}
|
||||
|
||||
type visitorUnitReporter struct {
|
||||
ok func(context.Context, SourceUnit) error
|
||||
err func(context.Context, error) error
|
||||
}
|
||||
|
||||
func (v visitorUnitReporter) UnitOk(ctx context.Context, u SourceUnit) error { return v.ok(ctx, u) }
|
||||
func (v visitorUnitReporter) UnitErr(ctx context.Context, err error) error { return v.err(ctx, err) }
|
||||
|
||||
type unitChunk struct {
|
||||
unit string
|
||||
output string
|
||||
@@ -230,7 +280,7 @@ func TestSourceManagerNonFatalError(t *testing.T) {
|
||||
mgr := NewManager(WithBufferedOutput(8), WithSourceUnits())
|
||||
source, err := buildDummy(&unitChunker{input})
|
||||
assert.NoError(t, err)
|
||||
ref, err := mgr.Run(context.Background(), "dummy", source)
|
||||
ref, err := mgr.EnumerateAndScan(context.Background(), "dummy", source)
|
||||
assert.NoError(t, err)
|
||||
<-ref.Done()
|
||||
report := ref.Snapshot()
|
||||
@@ -247,7 +297,7 @@ func TestSourceManagerContextCancelled(t *testing.T) {
|
||||
assert.NoError(t, err)
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
ref, err := mgr.Run(ctx, "dummy", source)
|
||||
ref, err := mgr.EnumerateAndScan(ctx, "dummy", source)
|
||||
assert.NoError(t, err)
|
||||
|
||||
cancel()
|
||||
@@ -291,7 +341,7 @@ func TestSourceManagerCancelRun(t *testing.T) {
|
||||
}})
|
||||
assert.NoError(t, err)
|
||||
|
||||
ref, err := mgr.Run(context.Background(), "dummy", source)
|
||||
ref, err := mgr.EnumerateAndScan(context.Background(), "dummy", source)
|
||||
assert.NoError(t, err)
|
||||
|
||||
cancelErr := fmt.Errorf("abort! abort!")
|
||||
@@ -313,7 +363,7 @@ func TestSourceManagerAvailableCapacity(t *testing.T) {
|
||||
assert.NoError(t, err)
|
||||
|
||||
assert.Equal(t, 1337, mgr.AvailableCapacity())
|
||||
ref, err := mgr.Run(context.Background(), "dummy", source)
|
||||
ref, err := mgr.EnumerateAndScan(context.Background(), "dummy", source)
|
||||
assert.NoError(t, err)
|
||||
|
||||
<-start // Wait for start signal.
|
||||
@@ -338,7 +388,7 @@ func TestSourceManagerUnitHook(t *testing.T) {
|
||||
)
|
||||
source, err := buildDummy(&unitChunker{input})
|
||||
assert.NoError(t, err)
|
||||
ref, err := mgr.Run(context.Background(), "dummy", source)
|
||||
ref, err := mgr.EnumerateAndScan(context.Background(), "dummy", source)
|
||||
assert.NoError(t, err)
|
||||
<-ref.Done()
|
||||
assert.NoError(t, mgr.Wait())
|
||||
@@ -399,7 +449,7 @@ func TestSourceManagerUnitHookBackPressure(t *testing.T) {
|
||||
)
|
||||
source, err := buildDummy(&unitChunker{input})
|
||||
assert.NoError(t, err)
|
||||
ref, err := mgr.Run(context.Background(), "dummy", source)
|
||||
ref, err := mgr.EnumerateAndScan(context.Background(), "dummy", source)
|
||||
assert.NoError(t, err)
|
||||
|
||||
var metrics []UnitMetrics
|
||||
@@ -428,7 +478,7 @@ func TestSourceManagerUnitHookNoUnits(t *testing.T) {
|
||||
source, err := buildDummy(&counterChunker{count: 5})
|
||||
assert.NoError(t, err)
|
||||
|
||||
ref, err := mgr.Run(context.Background(), "dummy", source)
|
||||
ref, err := mgr.EnumerateAndScan(context.Background(), "dummy", source)
|
||||
assert.NoError(t, err)
|
||||
<-ref.Done()
|
||||
assert.NoError(t, mgr.Wait())
|
||||
|
||||
@@ -105,6 +105,30 @@ type SourceUnitEnumerator interface {
|
||||
Enumerate(ctx context.Context, reporter UnitReporter) error
|
||||
}
|
||||
|
||||
// BaseUnitReporter is a helper struct that implements the UnitReporter interface
|
||||
// and includes a JobProgress reference.
|
||||
type baseUnitReporter struct {
|
||||
child UnitReporter
|
||||
progress *JobProgress
|
||||
}
|
||||
|
||||
func (b baseUnitReporter) UnitOk(ctx context.Context, unit SourceUnit) error {
|
||||
b.progress.ReportUnit(unit)
|
||||
if b.child != nil {
|
||||
return b.child.UnitOk(ctx, unit)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (b baseUnitReporter) UnitErr(ctx context.Context, err error) error {
|
||||
b.progress.ReportError(err)
|
||||
if b.child != nil {
|
||||
return b.child.UnitErr(ctx, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
// UnitReporter defines the interface a source will use to report whether a
|
||||
// unit was found during enumeration. Either method may be called any number of
|
||||
// times. Implementors of this interface should allow for concurrent calls.
|
||||
|
||||
@@ -1017,6 +1017,7 @@ enum DetectorType {
|
||||
WeightsAndBiases = 1005;
|
||||
ZohoCRM = 1006;
|
||||
AzureOpenAI = 1007;
|
||||
GoDaddy = 1008;
|
||||
}
|
||||
|
||||
message Result {
|
||||
|
||||
Reference in New Issue
Block a user