diff --git a/pkg/handlers/archive.go b/pkg/handlers/archive.go index 0b4a300fc..b93c3bdf5 100644 --- a/pkg/handlers/archive.go +++ b/pkg/handlers/archive.go @@ -14,7 +14,6 @@ import ( "time" "github.com/gabriel-vasile/mimetype" - "github.com/h2non/filetype" "github.com/mholt/archiver/v4" "github.com/trufflesecurity/trufflehog/v3/pkg/common" @@ -122,15 +121,16 @@ func (a *Archive) openArchive(ctx logContext.Context, depth int, reader io.Reade const mimeTypeBufferSize = 512 func (a *Archive) handleNonArchiveContent(ctx logContext.Context, reader io.Reader, archiveChan chan []byte) error { - bufReader := bufio.NewReaderSize(reader, mimeTypeBufferSize) - // A buffer of 512 bytes is used since many file formats store their magic numbers within the first 512 bytes. - // If fewer bytes are read, MIME type detection may still succeed. - buffer, err := bufReader.Peek(mimeTypeBufferSize) - if err != nil && !errors.Is(err, io.EOF) { - return fmt.Errorf("unable to read file for MIME type detection: %w", err) + mime, reader, err := determineMimeType(reader) + if err != nil { + return err + } + + // If the MIME type is nil, the reader is non-nil it is likely a directory. + if mime == nil && reader != nil { + return nil } - mime := mimetype.Detect(buffer) mimeT := mimeType(mime.String()) if common.SkipFile(mime.Extension()) { @@ -146,7 +146,7 @@ func (a *Archive) handleNonArchiveContent(ctx logContext.Context, reader io.Read } chunkReader := sources.NewChunkReader() - chunkResChan := chunkReader(ctx, bufReader) + chunkResChan := chunkReader(ctx, reader) for data := range chunkResChan { if err := data.Error(); err != nil { ctx.Logger().Error(err, "error reading chunk") @@ -303,11 +303,12 @@ func ensureToolsForMimeType(mimeType mimeType) error { // If the file is specialized, the returned boolean is true with no error. // The caller is responsible for closing the returned reader. func (a *Archive) HandleSpecialized(ctx logContext.Context, reader io.Reader) (io.Reader, bool, error) { - mimeType, reader, err := determineMimeType(reader) + mime, reader, err := determineMimeType(reader) if err != nil { return nil, false, err } + mimeType := mimeType(mime.String()) switch mimeType { case arMimeType: // includes .deb files if err := ensureToolsForMimeType(mimeType); err != nil { @@ -413,10 +414,11 @@ func (a *Archive) handleNestedFileMIME(ctx logContext.Context, tempEnv tempEnv, } defer nestedFile.Close() - mimeType, reader, err := determineMimeType(nestedFile) + mime, reader, err := determineMimeType(nestedFile) if err != nil { return "", fmt.Errorf("unable to determine MIME type of nested filename: %s, %w", nestedFile.Name(), err) } + mimeType := mimeType(mime.String()) switch mimeType { case arMimeType, rpmMimeType: @@ -434,25 +436,26 @@ func (a *Archive) handleNestedFileMIME(ctx logContext.Context, tempEnv tempEnv, // determineMimeType reads from the provided reader to detect the MIME type. // It returns the detected MIME type and a new reader that includes the read portion. -func determineMimeType(reader io.Reader) (mimeType, io.Reader, error) { +func determineMimeType(reader io.Reader) (*mimetype.MIME, io.Reader, error) { + bufReader := bufio.NewReaderSize(reader, mimeTypeBufferSize) // A buffer of 512 bytes is used since many file formats store their magic numbers within the first 512 bytes. // If fewer bytes are read, MIME type detection may still succeed. - buffer := make([]byte, 512) - n, err := reader.Read(buffer) + buffer, err := bufReader.Peek(mimeTypeBufferSize) if err != nil && !errors.Is(err, io.EOF) { - return "", nil, fmt.Errorf("unable to read file for MIME type detection: %w", err) + return nil, nil, fmt.Errorf("unable to read file for MIME type detection: %w", err) } - // Create a new reader that starts with the buffer we just read - // and continues with the rest of the original reader. - reader = io.MultiReader(bytes.NewReader(buffer[:n]), reader) - - kind, err := filetype.Match(buffer) - if err != nil { - return "", nil, fmt.Errorf("unable to determine file type: %w", err) + if len(buffer) == 0 { + return nil, bufReader, nil } - return mimeType(kind.MIME.Value), reader, nil + mime := mimetype.Detect(buffer) + mimeT := mimeType(mime.String()) + if mimeT == "" { + return nil, bufReader, fmt.Errorf("unable to determine MIME type") + } + + return mime, bufReader, nil } // handleExtractedFiles processes each file in the extracted directory using a provided handler function. diff --git a/pkg/handlers/archive_test.go b/pkg/handlers/archive_test.go index 6fe77f340..b5099ce27 100644 --- a/pkg/handlers/archive_test.go +++ b/pkg/handlers/archive_test.go @@ -383,7 +383,7 @@ func TestDetermineMimeType(t *testing.T) { { name: "Text file", input: bytes.NewReader(textBytes), - expected: mimeType("text/plain"), + expected: mimeType("text/plain; charset=utf-8"), shouldFail: false, }, { @@ -416,7 +416,7 @@ func TestDetermineMimeType(t *testing.T) { } if !tt.shouldFail { - assert.Equal(t, tt.expected, mime) + assert.Equal(t, tt.expected, mimeType(mime.String())) } // Ensure the reader still contains all the original data.