Files
trufflehog/pkg/sources/json_enumerator/json_enumerator.go
Brad Larsen f5370819b5 Improve the json-enumerator source (#5265)
* implement SourceUnitEnumerator and SourceUnitChunker
* natively support zstd-compressed inputs
* allow it to be used in multi-source config

Previously, multiple files would be processed sequentially, and zstd-compressed files would have to be decompressed externally.

These changes allow for slightly more efficient scanning of zstd-compressed NDJSON inputs (less CPU required), and enable greater CPU utilization on multicore systems when scanning multiple NDJSON files. With these changes, I see a 1.18x speedup in wall clock time in a simple experiment of scanning 10 zstd-compressed NDJSON files (~9GiB on disk) with verification disabled.
2026-09-03 13:44:35 -04:00

244 lines
6.2 KiB
Go

package json_enumerator
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"strings"
"unicode/utf8"
"github.com/klauspost/compress/zstd"
"google.golang.org/protobuf/proto"
"google.golang.org/protobuf/types/known/anypb"
"github.com/trufflesecurity/trufflehog/v3/pkg/common"
"github.com/trufflesecurity/trufflehog/v3/pkg/context"
"github.com/trufflesecurity/trufflehog/v3/pkg/handlers"
"github.com/trufflesecurity/trufflehog/v3/pkg/pb/source_metadatapb"
"github.com/trufflesecurity/trufflehog/v3/pkg/pb/sourcespb"
"github.com/trufflesecurity/trufflehog/v3/pkg/sources"
)
const SourceType = sourcespb.SourceType_SOURCE_TYPE_JSON_ENUMERATOR
type Source struct {
name string
sourceId sources.SourceID
jobId sources.JobID
verify bool
paths []string
sources.Progress
sources.CommonSourceUnitUnmarshaller
}
// Ensure the Source satisfies the interfaces at compile time
var _ sources.Source = (*Source)(nil)
var _ sources.SourceUnitUnmarshaller = (*Source)(nil)
var _ sources.SourceUnitEnumChunker = (*Source)(nil)
func (s *Source) Type() sourcespb.SourceType { return SourceType }
func (s *Source) SourceID() sources.SourceID { return s.sourceId }
func (s *Source) JobID() sources.JobID { return s.jobId }
func (s *Source) Init(
aCtx context.Context,
name string,
jobId sources.JobID,
sourceId sources.SourceID,
verify bool,
connection *anypb.Any,
concurrency int,
) error {
s.name = name
s.sourceId = sourceId
s.jobId = jobId
s.verify = verify
var conn sourcespb.JSONEnumerator
if err := anypb.UnmarshalTo(connection, &conn, proto.UnmarshalOptions{}); err != nil {
return fmt.Errorf("error unmarshalling connection: %w", err)
}
s.paths = conn.Paths
return nil
}
func (s *Source) Chunks(
ctx context.Context,
chunksChan chan *sources.Chunk,
_ ...sources.ChunkingTarget,
) error {
for i, path := range s.paths {
if common.IsDone(ctx) {
return nil
}
s.SetProgressComplete(i, len(s.paths), fmt.Sprintf("Path: %s", path), "")
if err := s.chunkJSONEnumerator(ctx, path, sources.ChanReporter{Ch: chunksChan}); err != nil {
ctx.Logger().Error(err, "error scanning JSON enumerator", "path", path)
}
}
return nil
}
// Enumerate implements the SourceUnitEnumerator interface, reporting each
// configured path as its own unit so paths can be chunked concurrently.
func (s *Source) Enumerate(ctx context.Context, reporter sources.UnitReporter) error {
for _, path := range s.paths {
f, err := os.Open(path)
if err != nil {
if err := reporter.UnitErr(ctx, err); err != nil {
return err
}
} else {
_ = f.Close()
if err := reporter.UnitOk(ctx, sources.CommonSourceUnit{ID: path}); err != nil {
return err
}
}
}
return nil
}
// ChunkUnit implements the SourceUnitChunker interface.
func (s *Source) ChunkUnit(ctx context.Context, unit sources.SourceUnit, reporter sources.ChunkReporter) error {
path, _ := unit.SourceUnitID()
if err := s.chunkJSONEnumerator(ctx, path, reporter); err != nil {
return reporter.ChunkErr(ctx, err)
}
return nil
}
type jsonEntry struct {
Metadata json.RawMessage
Data []byte
}
// jsonEntryAux is a helper struct to support marshalling the content to scan as either
// a UTF-8 string in `data` or a base64-encoded bytestring in `data_b64`.
type jsonEntryAux struct {
Metadata *json.RawMessage `json:"metadata"`
Data *string `json:"data,omitempty"`
DataB64 *[]byte `json:"data_b64,omitempty"`
}
// MarshalJSON implements custom JSON marshaling for jsonEntry.
// If Data is valid UTF-8, it's serialized as a `data` string field.
// If Data is not valid UTF-8, it's serialized as a `data_b64` base64-encoded string field.
func (e *jsonEntry) MarshalJSON() ([]byte, error) {
if utf8.Valid(e.Data) {
s := string(e.Data)
return json.Marshal(jsonEntryAux{
Metadata: &e.Metadata,
Data: &s,
})
} else {
return json.Marshal(jsonEntryAux{
Metadata: &e.Metadata,
DataB64: &e.Data,
})
}
}
// UnmarshalJSON implements custom JSON unmarshaling for jsonEntry.
func (e *jsonEntry) UnmarshalJSON(data []byte) error {
var aux jsonEntryAux
if err := json.Unmarshal(data, &aux); err != nil {
return err
}
if aux.Metadata == nil {
return fmt.Errorf("missing metadata")
}
if aux.Data == nil && aux.DataB64 == nil {
return fmt.Errorf("both data and data_b64 missing")
}
if aux.Data != nil && aux.DataB64 != nil {
return fmt.Errorf("both data and data_b64 present")
}
e.Metadata = *aux.Metadata
if aux.DataB64 != nil {
e.Data = *aux.DataB64
} else {
e.Data = []byte(*aux.Data)
}
return nil
}
func (s *Source) chunkJSONEnumeratorReader(
ctx context.Context,
input io.Reader,
reporter sources.ChunkReporter,
) error {
decoder := json.NewDecoder(input)
var entry jsonEntry
for {
if err := decoder.Decode(&entry); err != nil {
if errors.Is(err, io.EOF) {
// enumerator file is done
return nil
}
return err
}
metadataJSON, err := entry.Metadata.MarshalJSON()
if err != nil {
ctx.Logger().Error(err, "failed to convert metadata to JSON")
continue
}
chunkSkel := &sources.Chunk{
SourceType: s.Type(),
SourceName: s.name,
SourceID: s.SourceID(),
JobID: s.JobID(),
SourceVerify: s.verify,
SourceMetadata: &source_metadatapb.MetaData{
Data: &source_metadatapb.MetaData_JsonEnumerator{
JsonEnumerator: &source_metadatapb.JSONEnumerator{
Metadata: string(metadataJSON),
},
},
},
}
err = handlers.HandleFile(ctx, bytes.NewReader(entry.Data), chunkSkel, reporter)
if err != nil {
ctx.Logger().Error(err, "failed to scan data")
continue
}
}
}
func (s *Source) chunkJSONEnumerator(
ctx context.Context,
path string,
reporter sources.ChunkReporter,
) error {
ctx.Logger().V(3).Info("chunking JSON enumerator", "path", path)
enumeratorFile, err := os.Open(path)
if err != nil {
return fmt.Errorf("unable to open file: %w", err)
}
defer func() { _ = enumeratorFile.Close() }()
if strings.HasSuffix(path, ".zst") || strings.HasSuffix(path, ".zstd") {
r, err := zstd.NewReader(enumeratorFile)
if err != nil {
return fmt.Errorf("unable to open zstd reader: %w", err)
}
defer r.Close()
return s.chunkJSONEnumeratorReader(ctx, r, reporter)
} else {
return s.chunkJSONEnumeratorReader(ctx, enumeratorFile, reporter)
}
}