diff --git a/README.md b/README.md index 1c70ee56e..115e46125 100644 --- a/README.md +++ b/README.md @@ -418,6 +418,7 @@ TruffleHog has a sub-command for each source of data that you may want to scan: - jenkins - elasticsearch - stdin +- multi-scan Each subcommand can have options that you can see with the `--help` flag provided to the sub command: @@ -481,6 +482,33 @@ For example, to scan a `git` repository, start with trufflehog git https://github.com/trufflesecurity/trufflehog.git ``` +## Configuration + +TruffleHog supports defining [custom regex detectors](#regex-detector-alpha) +and multiple sources in a configuration file provided via the `--config` flag. +The regex detectors can be used with any subcommand, while the sources defined +in configuration are only for the `multi-scan` subcommand. + +The configuration format for sources can be found on Truffle Security's +[source configuration documentation page](https://docs.trufflesecurity.com/scan-data-for-secrets). + +Example GitHub source configuration and [options reference](https://docs.trufflesecurity.com/github#Fvm1I): + +```yaml +sources: +- connection: + '@type': type.googleapis.com/sources.GitHub + repositories: + - https://github.com/trufflesecurity/test_keys.git + unauthenticated: {} + name: example config scan + type: SOURCE_TYPE_GITHUB + verify: true +``` + +You may define multiple connections under the `sources` key (see above), and +TruffleHog will scan all of the sources concurrently. + ## S3 The S3 source supports assuming IAM roles for scanning in addition to IAM users. This makes it easier for users to scan multiple AWS accounts without needing to rely on hardcoded credentials for each account. diff --git a/main.go b/main.go index ec477c61b..e406702a7 100644 --- a/main.go +++ b/main.go @@ -255,6 +255,7 @@ var ( huggingfaceIncludePrs = huggingfaceScan.Flag("include-prs", "Include pull requests in scan.").Bool() stdinInputScan = cli.Command("stdin", "Find credentials from stdin.") + multiScanScan = cli.Command("multi-scan", "Find credentials in multiple sources defined in configuration.") analyzeCmd = analyzer.Command(cli) usingTUI = false @@ -515,7 +516,8 @@ func run(state overseer.State) { verificationCacheMetrics := verificationcache.InMemoryMetrics{} engConf := engine.Config{ - Concurrency: *concurrency, + Concurrency: *concurrency, + ConfiguredSources: conf.Sources, // The engine must always be configured with the list of // default detectors, which can be further filtered by the // user. The filters are applied by the engine and are only @@ -540,6 +542,16 @@ func run(state overseer.State) { engConf.VerificationResultCache = simple.NewCache[detectors.Result]() } + // Check that there are no sources defined for non-scan subcommands. If + // there are, return an error as it is ambiguous what the user is + // trying to do. + if cmd != multiScanScan.FullCommand() && len(conf.Sources) > 0 { + logFatal( + fmt.Errorf("ambiguous configuration"), + "sources should only be defined in configuration for the 'multi-scan' command", + ) + } + if *compareDetectionStrategies { if err := compareScans(ctx, cmd, engConf); err != nil { logFatal(err, "error comparing detection strategies") @@ -702,7 +714,7 @@ func runSingleScan(ctx context.Context, cmd string, cfg engine.Config) (metrics, } }() - var ref sources.JobProgressRef + var refs []sources.JobProgressRef switch cmd { case gitScan.FullCommand(): gitCfg := sources.GitConfig{ @@ -715,8 +727,10 @@ func runSingleScan(ctx context.Context, cmd string, cfg engine.Config) (metrics, Bare: *gitScanBare, ExcludeGlobs: *gitScanExcludeGlobs, } - if ref, err = eng.ScanGit(ctx, gitCfg); err != nil { + if ref, err := eng.ScanGit(ctx, gitCfg); err != nil { return scanMetrics, fmt.Errorf("failed to scan Git: %v", err) + } else { + refs = []sources.JobProgressRef{ref} } case githubScan.FullCommand(): filter, err := common.FilterFromFiles(*githubScanIncludePaths, *githubScanExcludePaths) @@ -745,8 +759,10 @@ func runSingleScan(ctx context.Context, cmd string, cfg engine.Config) (metrics, Filter: filter, AuthInUrl: *githubAuthInUrl, } - if ref, err = eng.ScanGitHub(ctx, cfg); err != nil { + if ref, err := eng.ScanGitHub(ctx, cfg); err != nil { return scanMetrics, fmt.Errorf("failed to scan Github: %v", err) + } else { + refs = []sources.JobProgressRef{ref} } case githubExperimentalScan.FullCommand(): cfg := sources.GitHubExperimentalConfig{ @@ -756,8 +772,10 @@ func runSingleScan(ctx context.Context, cmd string, cfg engine.Config) (metrics, CollisionThreshold: *githubExperimentalCollisionThreshold, DeleteCachedData: *githubExperimentalDeleteCache, } - if ref, err = eng.ScanGitHubExperimental(ctx, cfg); err != nil { + if ref, err := eng.ScanGitHubExperimental(ctx, cfg); err != nil { return scanMetrics, fmt.Errorf("failed to scan using Github Experimental: %v", err) + } else { + refs = []sources.JobProgressRef{ref} } case gitlabScan.FullCommand(): filter, err := common.FilterFromFiles(*gitlabScanIncludePaths, *gitlabScanExcludePaths) @@ -774,8 +792,10 @@ func runSingleScan(ctx context.Context, cmd string, cfg engine.Config) (metrics, Filter: filter, AuthInUrl: *gitlabAuthInUrl, } - if ref, err = eng.ScanGitLab(ctx, cfg); err != nil { + if ref, err := eng.ScanGitLab(ctx, cfg); err != nil { return scanMetrics, fmt.Errorf("failed to scan GitLab: %v", err) + } else { + refs = []sources.JobProgressRef{ref} } case filesystemScan.FullCommand(): if len(*filesystemDirectories) > 0 { @@ -789,8 +809,10 @@ func runSingleScan(ctx context.Context, cmd string, cfg engine.Config) (metrics, IncludePathsFile: *filesystemScanIncludePaths, ExcludePathsFile: *filesystemScanExcludePaths, } - if ref, err = eng.ScanFileSystem(ctx, cfg); err != nil { + if ref, err := eng.ScanFileSystem(ctx, cfg); err != nil { return scanMetrics, fmt.Errorf("failed to scan filesystem: %v", err) + } else { + refs = []sources.JobProgressRef{ref} } case s3Scan.FullCommand(): cfg := sources.S3Config{ @@ -803,8 +825,10 @@ func runSingleScan(ctx context.Context, cmd string, cfg engine.Config) (metrics, CloudCred: *s3ScanCloudEnv, MaxObjectSize: int64(*s3ScanMaxObjectSize), } - if ref, err = eng.ScanS3(ctx, cfg); err != nil { + if ref, err := eng.ScanS3(ctx, cfg); err != nil { return scanMetrics, fmt.Errorf("failed to scan S3: %v", err) + } else { + refs = []sources.JobProgressRef{ref} } case syslogScan.FullCommand(): cfg := sources.SyslogConfig{ @@ -815,16 +839,22 @@ func runSingleScan(ctx context.Context, cmd string, cfg engine.Config) (metrics, KeyPath: *syslogTLSKey, Concurrency: *concurrency, } - if ref, err = eng.ScanSyslog(ctx, cfg); err != nil { + if ref, err := eng.ScanSyslog(ctx, cfg); err != nil { return scanMetrics, fmt.Errorf("failed to scan syslog: %v", err) + } else { + refs = []sources.JobProgressRef{ref} } case circleCiScan.FullCommand(): - if ref, err = eng.ScanCircleCI(ctx, *circleCiScanToken); err != nil { + if ref, err := eng.ScanCircleCI(ctx, *circleCiScanToken); err != nil { return scanMetrics, fmt.Errorf("failed to scan CircleCI: %v", err) + } else { + refs = []sources.JobProgressRef{ref} } case travisCiScan.FullCommand(): - if ref, err = eng.ScanTravisCI(ctx, *travisCiScanToken); err != nil { + if ref, err := eng.ScanTravisCI(ctx, *travisCiScanToken); err != nil { return scanMetrics, fmt.Errorf("failed to scan TravisCI: %v", err) + } else { + refs = []sources.JobProgressRef{ref} } case gcsScan.FullCommand(): cfg := sources.GCSConfig{ @@ -840,8 +870,10 @@ func runSingleScan(ctx context.Context, cmd string, cfg engine.Config) (metrics, Concurrency: *concurrency, MaxObjectSize: int64(*gcsMaxObjectSize), } - if ref, err = eng.ScanGCS(ctx, cfg); err != nil { + if ref, err := eng.ScanGCS(ctx, cfg); err != nil { return scanMetrics, fmt.Errorf("failed to scan GCS: %v", err) + } else { + refs = []sources.JobProgressRef{ref} } case dockerScan.FullCommand(): cfg := sources.DockerConfig{ @@ -849,8 +881,10 @@ func runSingleScan(ctx context.Context, cmd string, cfg engine.Config) (metrics, Images: *dockerScanImages, UseDockerKeychain: *dockerScanToken == "", } - if ref, err = eng.ScanDocker(ctx, cfg); err != nil { + if ref, err := eng.ScanDocker(ctx, cfg); err != nil { return scanMetrics, fmt.Errorf("failed to scan Docker: %v", err) + } else { + refs = []sources.JobProgressRef{ref} } case postmanScan.FullCommand(): // handle deprecated flag @@ -886,8 +920,10 @@ func runSingleScan(ctx context.Context, cmd string, cfg engine.Config) (metrics, WorkspacePaths: *postmanWorkspacePaths, EnvironmentPaths: *postmanEnvironmentPaths, } - if ref, err = eng.ScanPostman(ctx, cfg); err != nil { + if ref, err := eng.ScanPostman(ctx, cfg); err != nil { return scanMetrics, fmt.Errorf("failed to scan Postman: %v", err) + } else { + refs = []sources.JobProgressRef{ref} } case elasticsearchScan.FullCommand(): cfg := sources.ElasticsearchConfig{ @@ -902,8 +938,10 @@ func runSingleScan(ctx context.Context, cmd string, cfg engine.Config) (metrics, SinceTimestamp: *elasticsearchSinceTimestamp, BestEffortScan: *elasticsearchBestEffortScan, } - if ref, err = eng.ScanElasticsearch(ctx, cfg); err != nil { + if ref, err := eng.ScanElasticsearch(ctx, cfg); err != nil { return scanMetrics, fmt.Errorf("failed to scan Elasticsearch: %v", err) + } else { + refs = []sources.JobProgressRef{ref} } case jenkinsScan.FullCommand(): cfg := engine.JenkinsConfig{ @@ -912,8 +950,10 @@ func runSingleScan(ctx context.Context, cmd string, cfg engine.Config) (metrics, Username: *jenkinsUsername, Password: *jenkinsPassword, } - if ref, err = eng.ScanJenkins(ctx, cfg); err != nil { + if ref, err := eng.ScanJenkins(ctx, cfg); err != nil { return scanMetrics, fmt.Errorf("failed to scan Jenkins: %v", err) + } else { + refs = []sources.JobProgressRef{ref} } case huggingfaceScan.FullCommand(): if *huggingfaceEndpoint != "" { @@ -945,13 +985,26 @@ func runSingleScan(ctx context.Context, cmd string, cfg engine.Config) (metrics, IncludePrs: *huggingfaceIncludePrs, Concurrency: *concurrency, } - if ref, err = eng.ScanHuggingface(ctx, cfg); err != nil { + if ref, err := eng.ScanHuggingface(ctx, cfg); err != nil { return scanMetrics, fmt.Errorf("failed to scan HuggingFace: %v", err) + } else { + refs = []sources.JobProgressRef{ref} + } + case multiScanScan.FullCommand(): + if *configFilename == "" { + return scanMetrics, fmt.Errorf("missing required flag: --config") + } + if rs, err := eng.ScanConfig(ctx, cfg.ConfiguredSources...); err != nil { + return scanMetrics, fmt.Errorf("failed to scan via config: %w", err) + } else { + refs = rs } case stdinInputScan.FullCommand(): cfg := sources.StdinConfig{} - if ref, err = eng.ScanStdinInput(ctx, cfg); err != nil { + if ref, err := eng.ScanStdinInput(ctx, cfg); err != nil { return scanMetrics, fmt.Errorf("failed to scan stdin input: %v", err) + } else { + refs = []sources.JobProgressRef{ref} } default: return scanMetrics, fmt.Errorf("invalid command: %s", cmd) @@ -962,13 +1015,19 @@ func runSingleScan(ctx context.Context, cmd string, cfg engine.Config) (metrics, return scanMetrics, fmt.Errorf("engine failed to finish execution: %v", err) } - // Print any errors reported during the scan. - if errs := ref.Snapshot().Errors; len(errs) > 0 { - errMsgs := make([]string, len(errs)) - for i := 0; i < len(errs); i++ { - errMsgs[i] = errs[i].Error() + // Print any non-fatal errors reported during the scan. + for _, ref := range refs { + if errs := ref.Snapshot().Errors; len(errs) > 0 { + errMsgs := make([]string, len(errs)) + for i := 0; i < len(errs); i++ { + errMsgs[i] = errs[i].Error() + } + ctx.Logger().Error(nil, "encountered errors during scan", + "job", ref.JobID, + "source_name", ref.SourceName, + "errors", errMsgs, + ) } - ctx.Logger().Error(nil, "encountered errors during scan", "errors", errMsgs) } if *printAvgDetectorTime { diff --git a/pkg/config/config.go b/pkg/config/config.go index e37a7a6db..47fe3868c 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -1,16 +1,29 @@ package config import ( + "fmt" "os" "github.com/trufflesecurity/trufflehog/v3/pkg/custom_detectors" "github.com/trufflesecurity/trufflehog/v3/pkg/detectors" - "github.com/trufflesecurity/trufflehog/v3/pkg/pb/custom_detectorspb" + "github.com/trufflesecurity/trufflehog/v3/pkg/pb/configpb" + "github.com/trufflesecurity/trufflehog/v3/pkg/pb/sourcespb" "github.com/trufflesecurity/trufflehog/v3/pkg/protoyaml" + "github.com/trufflesecurity/trufflehog/v3/pkg/sources" + "github.com/trufflesecurity/trufflehog/v3/pkg/sources/docker" + "github.com/trufflesecurity/trufflehog/v3/pkg/sources/filesystem" + "github.com/trufflesecurity/trufflehog/v3/pkg/sources/gcs" + "github.com/trufflesecurity/trufflehog/v3/pkg/sources/git" + "github.com/trufflesecurity/trufflehog/v3/pkg/sources/github" + "github.com/trufflesecurity/trufflehog/v3/pkg/sources/gitlab" + "github.com/trufflesecurity/trufflehog/v3/pkg/sources/jenkins" + "github.com/trufflesecurity/trufflehog/v3/pkg/sources/postman" + "github.com/trufflesecurity/trufflehog/v3/pkg/sources/s3" ) // Config holds user supplied configuration. type Config struct { + Sources []sources.ConfiguredSource Detectors []detectors.Detector } @@ -25,21 +38,74 @@ func Read(filename string) (*Config, error) { // NewYAML parses the given YAML data into a Config. func NewYAML(input []byte) (*Config, error) { + var inputYAML configpb.Config // Parse the raw YAML into a structure. - var messages custom_detectorspb.CustomDetectors - if err := protoyaml.UnmarshalStrict(input, &messages); err != nil { + if err := protoyaml.UnmarshalStrict(input, &inputYAML); err != nil { return nil, err } - // Convert the structured YAML into detectors. - var d []detectors.Detector - for _, detectorConfig := range messages.Detectors { + + // Convert to detectors. + var detectorConfigs []detectors.Detector + for _, detectorConfig := range inputYAML.Detectors { detector, err := custom_detectors.NewWebhookCustomRegex(detectorConfig) if err != nil { return nil, err } - d = append(d, detector) + detectorConfigs = append(detectorConfigs, detector) } + + // Convert to configured sources. + var sourceConfigs []sources.ConfiguredSource + for _, pbSource := range inputYAML.Sources { + s, err := instantiateSourceFromType(pbSource.GetType()) + if err != nil { + return nil, err + } + src := sources.NewConfiguredSource(s, pbSource) + + sourceConfigs = append(sourceConfigs, src) + } + return &Config{ - Detectors: d, + Detectors: detectorConfigs, + Sources: sourceConfigs, }, nil } + +// instantiateSourceFromType creates a concrete implementation of +// sources.Source for the provided type. +func instantiateSourceFromType(sourceType string) (sources.Source, error) { + var source sources.Source + switch sourceType { + case sourcespb.SourceType_SOURCE_TYPE_GIT.String(): + source = new(git.Source) + case sourcespb.SourceType_SOURCE_TYPE_GITHUB.String(): + source = new(github.Source) + case sourcespb.SourceType_SOURCE_TYPE_GITHUB_UNAUTHENTICATED_ORG.String(): + source = new(github.Source) + case sourcespb.SourceType_SOURCE_TYPE_PUBLIC_GIT.String(): + source = new(git.Source) + case sourcespb.SourceType_SOURCE_TYPE_GITLAB.String(): + source = new(gitlab.Source) + case sourcespb.SourceType_SOURCE_TYPE_POSTMAN.String(): + source = new(postman.Source) + case sourcespb.SourceType_SOURCE_TYPE_S3.String(): + source = new(s3.Source) + case sourcespb.SourceType_SOURCE_TYPE_S3_UNAUTHED.String(): + source = new(s3.Source) + case sourcespb.SourceType_SOURCE_TYPE_FILESYSTEM.String(): + source = new(filesystem.Source) + case sourcespb.SourceType_SOURCE_TYPE_JENKINS.String(): + source = new(jenkins.Source) + case sourcespb.SourceType_SOURCE_TYPE_GCS.String(): + source = new(gcs.Source) + case sourcespb.SourceType_SOURCE_TYPE_GCS_UNAUTHED.String(): + source = new(gcs.Source) + case sourcespb.SourceType_SOURCE_TYPE_DOCKER.String(): + source = new(docker.Source) + default: + return nil, fmt.Errorf("got unexpected source type: %q", sourceType) + } + + return source, nil +} diff --git a/pkg/engine/engine.go b/pkg/engine/engine.go index 67c76370f..287a19188 100644 --- a/pkg/engine/engine.go +++ b/pkg/engine/engine.go @@ -99,6 +99,7 @@ type Config struct { // also serves as a multiplier for other worker types (e.g., detector workers, notifier workers) Concurrency int + ConfiguredSources []sources.ConfiguredSource Decoders []decoders.Decoder Detectors []detectors.Detector DetectorVerificationOverrides map[config.DetectorID]bool @@ -546,6 +547,17 @@ func (r *verificationOverlapTracker) increment() { const ignoreTag = "trufflehog:ignore" +// AhoCorasickCoreKeywords returns a set of keywords that the engine's +// AhoCorasickCore is using. +func (e *Engine) AhoCorasickCoreKeywords() map[string]struct{} { + // Turn AhoCorasick keywordsToDetectors into a map of keywords + keywords := make(map[string]struct{}) + for key := range e.AhoCorasickCore.KeywordsToDetectors() { + keywords[key] = struct{}{} + } + return keywords +} + // HasFoundResults returns true if any results are found. func (e *Engine) HasFoundResults() bool { return atomic.LoadUint32(&e.numFoundResults) > 0 diff --git a/pkg/engine/postman.go b/pkg/engine/postman.go index ac6b7af62..fc5e6ae08 100644 --- a/pkg/engine/postman.go +++ b/pkg/engine/postman.go @@ -39,12 +39,6 @@ func (e *Engine) ScanPostman(ctx context.Context, c sources.PostmanConfig) (sour return sources.JobProgressRef{}, errors.New("no path to locally exported data or API token provided") } - // Turn AhoCorasick keywordsToDetectors into a map of keywords - keywords := make(map[string]struct{}) - for key := range e.AhoCorasickCore.KeywordsToDetectors() { - keywords[key] = struct{}{} - } - var conn anypb.Any err := anypb.MarshalFrom(&conn, &connection, proto.MarshalOptions{}) if err != nil { @@ -56,8 +50,9 @@ func (e *Engine) ScanPostman(ctx context.Context, c sources.PostmanConfig) (sour sourceID, jobID, _ := e.sourceManager.GetIDs(ctx, sourceName, postman.SourceType) postmanSource := &postman.Source{ - DetectorKeywords: keywords, + DetectorKeywords: e.AhoCorasickCoreKeywords(), } + if err := postmanSource.Init(ctx, sourceName, jobID, sourceID, true, &conn, c.Concurrency); err != nil { return sources.JobProgressRef{}, err } diff --git a/pkg/engine/scan.go b/pkg/engine/scan.go new file mode 100644 index 000000000..167204e2a --- /dev/null +++ b/pkg/engine/scan.go @@ -0,0 +1,36 @@ +package engine + +import ( + "github.com/trufflesecurity/trufflehog/v3/pkg/context" + "github.com/trufflesecurity/trufflehog/v3/pkg/sources" + "github.com/trufflesecurity/trufflehog/v3/pkg/sources/postman" +) + +// ScanConfig starts a scan of all of the configured (but not initialized) +// sources and returns their job references. If there is an error during +// initialization or starting of the scan, an error is returned along with the +// references that successfully started up to that point. +func (e *Engine) ScanConfig(ctx context.Context, configuredSources ...sources.ConfiguredSource) ([]sources.JobProgressRef, error) { + var refs []sources.JobProgressRef + for _, configuredSource := range configuredSources { + sourceID, jobID, _ := e.sourceManager.GetIDs(ctx, configuredSource.Name, configuredSource.SourceType()) + + source, err := configuredSource.Init(ctx, sourceID, jobID) + if err != nil { + return refs, err + } + // Postman needs special initialization to set Keywords from + // the engine. + if postmanSource, ok := source.(*postman.Source); ok { + postmanSource.DetectorKeywords = e.AhoCorasickCoreKeywords() + } + + // Start the scan. + ref, err := e.sourceManager.EnumerateAndScan(ctx, configuredSource.Name, source) + if err != nil { + return refs, err + } + refs = append(refs, ref) + } + return refs, nil +} diff --git a/pkg/pb/configpb/config.pb.go b/pkg/pb/configpb/config.pb.go new file mode 100644 index 000000000..cc760fcba --- /dev/null +++ b/pkg/pb/configpb/config.pb.go @@ -0,0 +1,166 @@ +// Code generated by protoc-gen-go. DO NOT EDIT. +// versions: +// protoc-gen-go v1.33.0 +// protoc v4.25.3 +// source: config.proto + +package configpb + +import ( + custom_detectorspb "github.com/trufflesecurity/trufflehog/v3/pkg/pb/custom_detectorspb" + sourcespb "github.com/trufflesecurity/trufflehog/v3/pkg/pb/sourcespb" + protoreflect "google.golang.org/protobuf/reflect/protoreflect" + protoimpl "google.golang.org/protobuf/runtime/protoimpl" + reflect "reflect" + sync "sync" +) + +const ( + // Verify that this generated code is sufficiently up-to-date. + _ = protoimpl.EnforceVersion(20 - protoimpl.MinVersion) + // Verify that runtime/protoimpl is sufficiently up-to-date. + _ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20) +) + +type Config struct { + state protoimpl.MessageState + sizeCache protoimpl.SizeCache + unknownFields protoimpl.UnknownFields + + Sources []*sourcespb.LocalSource `protobuf:"bytes,9,rep,name=sources,proto3" json:"sources,omitempty"` + Detectors []*custom_detectorspb.CustomRegex `protobuf:"bytes,13,rep,name=detectors,proto3" json:"detectors,omitempty"` +} + +func (x *Config) Reset() { + *x = Config{} + if protoimpl.UnsafeEnabled { + mi := &file_config_proto_msgTypes[0] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) + } +} + +func (x *Config) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*Config) ProtoMessage() {} + +func (x *Config) ProtoReflect() protoreflect.Message { + mi := &file_config_proto_msgTypes[0] + if protoimpl.UnsafeEnabled && x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use Config.ProtoReflect.Descriptor instead. +func (*Config) Descriptor() ([]byte, []int) { + return file_config_proto_rawDescGZIP(), []int{0} +} + +func (x *Config) GetSources() []*sourcespb.LocalSource { + if x != nil { + return x.Sources + } + return nil +} + +func (x *Config) GetDetectors() []*custom_detectorspb.CustomRegex { + if x != nil { + return x.Detectors + } + return nil +} + +var File_config_proto protoreflect.FileDescriptor + +var file_config_proto_rawDesc = []byte{ + 0x0a, 0x0c, 0x63, 0x6f, 0x6e, 0x66, 0x69, 0x67, 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x12, 0x06, + 0x63, 0x6f, 0x6e, 0x66, 0x69, 0x67, 0x1a, 0x0d, 0x73, 0x6f, 0x75, 0x72, 0x63, 0x65, 0x73, 0x2e, + 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x1a, 0x16, 0x63, 0x75, 0x73, 0x74, 0x6f, 0x6d, 0x5f, 0x64, 0x65, + 0x74, 0x65, 0x63, 0x74, 0x6f, 0x72, 0x73, 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x22, 0x75, 0x0a, + 0x06, 0x43, 0x6f, 0x6e, 0x66, 0x69, 0x67, 0x12, 0x2e, 0x0a, 0x07, 0x73, 0x6f, 0x75, 0x72, 0x63, + 0x65, 0x73, 0x18, 0x09, 0x20, 0x03, 0x28, 0x0b, 0x32, 0x14, 0x2e, 0x73, 0x6f, 0x75, 0x72, 0x63, + 0x65, 0x73, 0x2e, 0x4c, 0x6f, 0x63, 0x61, 0x6c, 0x53, 0x6f, 0x75, 0x72, 0x63, 0x65, 0x52, 0x07, + 0x73, 0x6f, 0x75, 0x72, 0x63, 0x65, 0x73, 0x12, 0x3b, 0x0a, 0x09, 0x64, 0x65, 0x74, 0x65, 0x63, + 0x74, 0x6f, 0x72, 0x73, 0x18, 0x0d, 0x20, 0x03, 0x28, 0x0b, 0x32, 0x1d, 0x2e, 0x63, 0x75, 0x73, + 0x74, 0x6f, 0x6d, 0x5f, 0x64, 0x65, 0x74, 0x65, 0x63, 0x74, 0x6f, 0x72, 0x73, 0x2e, 0x43, 0x75, + 0x73, 0x74, 0x6f, 0x6d, 0x52, 0x65, 0x67, 0x65, 0x78, 0x52, 0x09, 0x64, 0x65, 0x74, 0x65, 0x63, + 0x74, 0x6f, 0x72, 0x73, 0x42, 0x3a, 0x5a, 0x38, 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, 0x63, 0x6f, 0x6e, 0x66, 0x69, 0x67, 0x70, 0x62, + 0x62, 0x06, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x33, +} + +var ( + file_config_proto_rawDescOnce sync.Once + file_config_proto_rawDescData = file_config_proto_rawDesc +) + +func file_config_proto_rawDescGZIP() []byte { + file_config_proto_rawDescOnce.Do(func() { + file_config_proto_rawDescData = protoimpl.X.CompressGZIP(file_config_proto_rawDescData) + }) + return file_config_proto_rawDescData +} + +var file_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1) +var file_config_proto_goTypes = []interface{}{ + (*Config)(nil), // 0: config.Config + (*sourcespb.LocalSource)(nil), // 1: sources.LocalSource + (*custom_detectorspb.CustomRegex)(nil), // 2: custom_detectors.CustomRegex +} +var file_config_proto_depIdxs = []int32{ + 1, // 0: config.Config.sources:type_name -> sources.LocalSource + 2, // 1: config.Config.detectors:type_name -> custom_detectors.CustomRegex + 2, // [2:2] is the sub-list for method output_type + 2, // [2:2] is the sub-list for method input_type + 2, // [2:2] is the sub-list for extension type_name + 2, // [2:2] is the sub-list for extension extendee + 0, // [0:2] is the sub-list for field type_name +} + +func init() { file_config_proto_init() } +func file_config_proto_init() { + if File_config_proto != nil { + return + } + if !protoimpl.UnsafeEnabled { + file_config_proto_msgTypes[0].Exporter = func(v interface{}, i int) interface{} { + switch v := v.(*Config); i { + case 0: + return &v.state + case 1: + return &v.sizeCache + case 2: + return &v.unknownFields + default: + return nil + } + } + } + type x struct{} + out := protoimpl.TypeBuilder{ + File: protoimpl.DescBuilder{ + GoPackagePath: reflect.TypeOf(x{}).PkgPath(), + RawDescriptor: file_config_proto_rawDesc, + NumEnums: 0, + NumMessages: 1, + NumExtensions: 0, + NumServices: 0, + }, + GoTypes: file_config_proto_goTypes, + DependencyIndexes: file_config_proto_depIdxs, + MessageInfos: file_config_proto_msgTypes, + }.Build() + File_config_proto = out.File + file_config_proto_rawDesc = nil + file_config_proto_goTypes = nil + file_config_proto_depIdxs = nil +} diff --git a/pkg/pb/configpb/config.pb.validate.go b/pkg/pb/configpb/config.pb.validate.go new file mode 100644 index 000000000..e02545a95 --- /dev/null +++ b/pkg/pb/configpb/config.pb.validate.go @@ -0,0 +1,202 @@ +// Code generated by protoc-gen-validate. DO NOT EDIT. +// source: config.proto + +package configpb + +import ( + "bytes" + "errors" + "fmt" + "net" + "net/mail" + "net/url" + "regexp" + "sort" + "strings" + "time" + "unicode/utf8" + + "google.golang.org/protobuf/types/known/anypb" +) + +// ensure the imports are used +var ( + _ = bytes.MinRead + _ = errors.New("") + _ = fmt.Print + _ = utf8.UTFMax + _ = (*regexp.Regexp)(nil) + _ = (*strings.Reader)(nil) + _ = net.IPv4len + _ = time.Duration(0) + _ = (*url.URL)(nil) + _ = (*mail.Address)(nil) + _ = anypb.Any{} + _ = sort.Sort +) + +// Validate checks the field values on Config with the rules defined in the +// proto definition for this message. If any rules are violated, the first +// error encountered is returned, or nil if there are no violations. +func (m *Config) Validate() error { + return m.validate(false) +} + +// ValidateAll checks the field values on Config with the rules defined in the +// proto definition for this message. If any rules are violated, the result is +// a list of violation errors wrapped in ConfigMultiError, or nil if none found. +func (m *Config) ValidateAll() error { + return m.validate(true) +} + +func (m *Config) validate(all bool) error { + if m == nil { + return nil + } + + var errors []error + + for idx, item := range m.GetSources() { + _, _ = idx, item + + if all { + switch v := interface{}(item).(type) { + case interface{ ValidateAll() error }: + if err := v.ValidateAll(); err != nil { + errors = append(errors, ConfigValidationError{ + field: fmt.Sprintf("Sources[%v]", idx), + reason: "embedded message failed validation", + cause: err, + }) + } + case interface{ Validate() error }: + if err := v.Validate(); err != nil { + errors = append(errors, ConfigValidationError{ + field: fmt.Sprintf("Sources[%v]", idx), + reason: "embedded message failed validation", + cause: err, + }) + } + } + } else if v, ok := interface{}(item).(interface{ Validate() error }); ok { + if err := v.Validate(); err != nil { + return ConfigValidationError{ + field: fmt.Sprintf("Sources[%v]", idx), + reason: "embedded message failed validation", + cause: err, + } + } + } + + } + + for idx, item := range m.GetDetectors() { + _, _ = idx, item + + if all { + switch v := interface{}(item).(type) { + case interface{ ValidateAll() error }: + if err := v.ValidateAll(); err != nil { + errors = append(errors, ConfigValidationError{ + field: fmt.Sprintf("Detectors[%v]", idx), + reason: "embedded message failed validation", + cause: err, + }) + } + case interface{ Validate() error }: + if err := v.Validate(); err != nil { + errors = append(errors, ConfigValidationError{ + field: fmt.Sprintf("Detectors[%v]", idx), + reason: "embedded message failed validation", + cause: err, + }) + } + } + } else if v, ok := interface{}(item).(interface{ Validate() error }); ok { + if err := v.Validate(); err != nil { + return ConfigValidationError{ + field: fmt.Sprintf("Detectors[%v]", idx), + reason: "embedded message failed validation", + cause: err, + } + } + } + + } + + if len(errors) > 0 { + return ConfigMultiError(errors) + } + + return nil +} + +// ConfigMultiError is an error wrapping multiple validation errors returned by +// Config.ValidateAll() if the designated constraints aren't met. +type ConfigMultiError []error + +// Error returns a concatenation of all the error messages it wraps. +func (m ConfigMultiError) Error() string { + var msgs []string + for _, err := range m { + msgs = append(msgs, err.Error()) + } + return strings.Join(msgs, "; ") +} + +// AllErrors returns a list of validation violation errors. +func (m ConfigMultiError) AllErrors() []error { return m } + +// ConfigValidationError is the validation error returned by Config.Validate if +// the designated constraints aren't met. +type ConfigValidationError struct { + field string + reason string + cause error + key bool +} + +// Field function returns field value. +func (e ConfigValidationError) Field() string { return e.field } + +// Reason function returns reason value. +func (e ConfigValidationError) Reason() string { return e.reason } + +// Cause function returns cause value. +func (e ConfigValidationError) Cause() error { return e.cause } + +// Key function returns key value. +func (e ConfigValidationError) Key() bool { return e.key } + +// ErrorName returns error name. +func (e ConfigValidationError) ErrorName() string { return "ConfigValidationError" } + +// Error satisfies the builtin error interface +func (e ConfigValidationError) Error() string { + cause := "" + if e.cause != nil { + cause = fmt.Sprintf(" | caused by: %v", e.cause) + } + + key := "" + if e.key { + key = "key for " + } + + return fmt.Sprintf( + "invalid %sConfig.%s: %s%s", + key, + e.field, + e.reason, + cause) +} + +var _ error = ConfigValidationError{} + +var _ interface { + Field() string + Reason() string + Key() bool + Cause() error + ErrorName() string +} = ConfigValidationError{} diff --git a/pkg/sources/sources.go b/pkg/sources/sources.go index d4763715e..32815e124 100644 --- a/pkg/sources/sources.go +++ b/pkg/sources/sources.go @@ -1,6 +1,8 @@ package sources import ( + "errors" + "runtime" "sync" "google.golang.org/protobuf/types/known/anypb" @@ -65,7 +67,7 @@ type Source interface { SourceID() SourceID // JobID returns the initialized job ID used for tracking relationships in the DB. JobID() JobID - // Init initializes the source. + // Init initializes the source. Calling this method more than once is undefined behavior. Init(aCtx context.Context, name string, jobId JobID, sourceId SourceID, verify bool, connection *anypb.Any, concurrency int) error // Chunks emits data over a channel which is then decoded and scanned for secrets. // By default, data is obtained indiscriminately. However, by providing one or more @@ -103,6 +105,56 @@ type SourceUnitEnumerator interface { Enumerate(ctx context.Context, reporter UnitReporter) error } +// ConfiguredSource is a Source with most of its initialization values +// pre-configured from a [sourcespb.LocalSource] configuration struct. It +// exposes a simplified Init() method and can be only initialized once. This +// struct is not necessary for running sources, but it helps simplify gathering +// all of the necessary information to call the [Source.Init] method. +type ConfiguredSource struct { + Name string + source Source + initParams struct { + verify bool + conn *anypb.Any + concurrency int + } +} + +// NewConfiguredSource pre-configures an instantiated Source object with the +// provided protobuf configuration. +func NewConfiguredSource(s Source, config *sourcespb.LocalSource) ConfiguredSource { + return ConfiguredSource{ + Name: config.GetName(), + source: s, + initParams: struct { + verify bool + conn *anypb.Any + concurrency int + }{ + verify: config.GetVerify(), + conn: config.GetConnection(), + concurrency: runtime.NumCPU(), + }, + } +} + +// SourceType exposes the underlying source type. +func (c *ConfiguredSource) SourceType() sourcespb.SourceType { + return c.source.Type() +} + +// Init returns the initialized Source. The ConfiguredSource is unusable after +// calling this method because initializing a [Source] more than once is undefined. +func (c *ConfiguredSource) Init(ctx context.Context, sourceID SourceID, jobID JobID) (Source, error) { + if c.source == nil { + return nil, errors.New("source already initialized") + } + src := c.source + err := src.Init(ctx, c.Name, jobID, sourceID, c.initParams.verify, c.initParams.conn, c.initParams.concurrency) + c.source = nil + return src, err +} + // BaseUnitReporter is a helper struct that implements the UnitReporter interface // and includes a JobProgress reference. type baseUnitReporter struct { diff --git a/proto/config.proto b/proto/config.proto new file mode 100644 index 000000000..b74bbd541 --- /dev/null +++ b/proto/config.proto @@ -0,0 +1,13 @@ +syntax = "proto3"; + +package config; + +option go_package = "github.com/trufflesecurity/trufflehog/v3/pkg/pb/configpb"; + +import "sources.proto"; +import "custom_detectors.proto"; + +message Config { + repeated sources.LocalSource sources = 9; + repeated custom_detectors.CustomRegex detectors = 13; +}