diff --git a/pkg/broadcast/broadcast.go b/pkg/broadcast/broadcast.go index 6f0fce6..c1573c7 100644 --- a/pkg/broadcast/broadcast.go +++ b/pkg/broadcast/broadcast.go @@ -22,86 +22,82 @@ package broadcast import ( - "context" + "sync" "github.com/americanexpress/earlybird/v4/pkg/scan" ) type BroadcastServer interface { - Subscribe() <-chan scan.Hit CancelSubscription(<-chan scan.Hit) + GetListeners() []chan scan.Hit } type broadcastServer struct { - source <-chan scan.Hit - listeners []chan scan.Hit - addListener chan chan scan.Hit - removeListener chan (<-chan scan.Hit) + source <-chan scan.Hit + listeners []chan scan.Hit + wg *sync.WaitGroup } -// Subscribe() creates a subcribtion on broadcastServer. +// Subscribe() creates a subscription on broadcastServer. func (s *broadcastServer) Subscribe() <-chan scan.Hit { newListener := make(chan scan.Hit) - s.addListener <- newListener + s.listeners = append(s.listeners, newListener) + return newListener } -// CancelSubscription() cancel a subcribtion on broadcastServer. +// CancelSubscription() cancel a subscription on broadcastServer. func (s *broadcastServer) CancelSubscription(channel <-chan scan.Hit) { - s.removeListener <- channel + for i, ch := range s.listeners { + if ch == channel { + s.wg.Done() + s.listeners[i] = s.listeners[len(s.listeners)-1] + s.listeners = s.listeners[:len(s.listeners)-1] + close(ch) + break + } + } +} + +func (s *broadcastServer) GetListeners() []chan scan.Hit { + return s.listeners +} + +func (s *broadcastServer) AddSubscriber(count int, wg *sync.WaitGroup) { + for i := 0; i < count; i++ { + s.wg.Add(1) + s.Subscribe() + } +} + +func (s *broadcastServer) CloseBroadcast() { + for _, listener := range s.listeners { + if listener != nil { + close(listener) + } + } } // NewBroadcastServer() create a broadcast server and starts new routine. -func NewBroadcastServer(ctx context.Context, source <-chan scan.Hit) BroadcastServer { +func NewBroadcastServer(source <-chan scan.Hit, count int, wg *sync.WaitGroup) BroadcastServer { service := &broadcastServer{ - source: source, - listeners: make([]chan scan.Hit, 0), - addListener: make(chan chan scan.Hit, 10), - removeListener: make(chan (<-chan scan.Hit)), + source: source, + listeners: make([]chan scan.Hit, 0), + wg: wg, } - go service.serve(ctx) + service.AddSubscriber(count, wg) + go service.broadCastData() return service } -// serve() run the server and manages listener counts. -func (s *broadcastServer) serve(ctx context.Context) { - defer func() { +// broadCastData() run the server and manages listener counts. +func (s *broadcastServer) broadCastData() { + for val := range s.source { for _, listener := range s.listeners { if listener != nil { - close(listener) - } - } - }() - - for { - select { - case <-ctx.Done(): - return - case newListener := <-s.addListener: - s.listeners = append(s.listeners, newListener) - case listenerToRemove := <-s.removeListener: - for i, ch := range s.listeners { - if ch == listenerToRemove { - s.listeners[i] = s.listeners[len(s.listeners)-1] - s.listeners = s.listeners[:len(s.listeners)-1] - close(ch) - break - } - } - case val, ok := <-s.source: - if !ok { - return - } - for _, listener := range s.listeners { - if listener != nil { - select { - case listener <- val: - case <-ctx.Done(): - return - } - - } + listener <- val } } } + s.CloseBroadcast() } diff --git a/pkg/core/core.go b/pkg/core/core.go index ce7934a..5ebbbac 100644 --- a/pkg/core/core.go +++ b/pkg/core/core.go @@ -18,7 +18,6 @@ package core import ( "bufio" - "context" "flag" "fmt" "log" @@ -293,11 +292,13 @@ func (eb *EarlybirdCfg) Scan() { if err != nil { log.Fatal("Failed to get FileContext: ", err) } + var wg sync.WaitGroup HitChannel := make(chan scan.Hit) - go scan.SearchFiles(&eb.Config, fileContext.Files, fileContext.CompressPaths, fileContext.ConvertPaths, HitChannel) - // Send output to a writer - eb.WriteResults(start, HitChannel, fileContext) + eb.WriteResults(start, HitChannel, fileContext, &wg) // Registering the hit receiver. + scan.SearchFiles(&eb.Config, fileContext.Files, fileContext.CompressPaths, fileContext.ConvertPaths, HitChannel) // sending the hits to the channel from the worker threads running on go-routine. + + wg.Wait() // this wait ensures that all writers goroutine are finished utils.DeleteGit(eb.Config.Gitrepo, eb.Config.SearchDir) if eb.Config.FailScan { @@ -336,41 +337,44 @@ func (eb *EarlybirdCfg) FileContext() (fileContext file.Context, err error) { } // WriteResults reads hits from the channel to the console or target file -func (eb *EarlybirdCfg) WriteResults(start time.Time, HitChannel chan scan.Hit, fileContext file.Context) { +func (eb *EarlybirdCfg) WriteResults(start time.Time, HitChannel chan scan.Hit, fileContext file.Context, wg *sync.WaitGroup) { // Send output to a writer var err error - + // if eb.Config.WithConsole && eb.Config.OutputFormat == "json" { - var wg sync.WaitGroup - wg.Add(2) - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - broadcaster := broadcast.NewBroadcastServer(ctx, HitChannel) - listener1 := broadcaster.Subscribe() - listener2 := broadcaster.Subscribe() + // initializing the broadcaster with two listeners, and starting the broadcast server + broadcaster := broadcast.NewBroadcastServer(HitChannel, 2, wg) + listener := broadcaster.GetListeners() go func() { defer wg.Done() - err = writers.WriteConsole(listener1, "", eb.Config.ShowFullLine) + err = writers.WriteConsole(listener[0], "", eb.Config.ShowFullLine) log.Printf("\n%d files scanned in %s", len(fileContext.Files), time.Since(start)) - log.Printf("\n%d rules observed\n", len(scan.CombinedRules)) + printError(err) }() go func() { defer wg.Done() - err = writers.WriteJSON(listener2, eb.Config, fileContext, eb.Config.OutputFile) + err = writers.WriteJSON(listener[1], eb.Config, fileContext, eb.Config.OutputFile) + printError(err) }() - wg.Wait() } else { - switch { - case eb.Config.OutputFormat == "json": - err = writers.WriteJSON(HitChannel, eb.Config, fileContext, eb.Config.OutputFile) - case eb.Config.OutputFormat == "csv": - err = writers.WriteCSV(HitChannel, eb.Config.OutputFile) - default: - err = writers.WriteConsole(HitChannel, eb.Config.OutputFile, eb.Config.ShowFullLine) - log.Printf("\n%d files scanned in %s", len(fileContext.Files), time.Since(start)) - log.Printf("\n%d rules observed\n", len(scan.CombinedRules)) - } + wg.Add(1) + go func() { + defer wg.Done() + switch { + case eb.Config.OutputFormat == "json": + err = writers.WriteJSON(HitChannel, eb.Config, fileContext, eb.Config.OutputFile) + case eb.Config.OutputFormat == "csv": + err = writers.WriteCSV(HitChannel, eb.Config.OutputFile) + default: + err = writers.WriteConsole(HitChannel, eb.Config.OutputFile, eb.Config.ShowFullLine) + log.Printf("\n%d files scanned in %s", len(fileContext.Files), time.Since(start)) + } + printError(err) + }() } +} + +func printError(err error) { if err != nil { log.Println("Writing Results failed:", err) } diff --git a/pkg/writers/consoleout.go b/pkg/writers/consoleout.go index 3ad9273..6f5b96b 100644 --- a/pkg/writers/consoleout.go +++ b/pkg/writers/consoleout.go @@ -35,7 +35,7 @@ type issue struct { var issues = make(map[string]int) -//WriteConsole streams hits from the result channel to the command line or target file +// WriteConsole streams hits from the result channel to the command line or target file func WriteConsole(hits <-chan scan.Hit, fileName string, showFullLine bool) error { // If no filename was passed in, just print to stdout if fileName == "" { @@ -53,6 +53,7 @@ func WriteConsole(hits <-chan scan.Hit, fileName string, showFullLine bool) erro } } displayIssues() + log.Printf("\n%d rules observed\n", len(scan.CombinedRules)) return nil } diff --git a/pkg/writers/jsonout.go b/pkg/writers/jsonout.go index d1e89f8..8be81b5 100755 --- a/pkg/writers/jsonout.go +++ b/pkg/writers/jsonout.go @@ -19,11 +19,12 @@ package writers import ( "encoding/json" "fmt" + "os" + "time" + cfgReader "github.com/americanexpress/earlybird/v4/pkg/config" "github.com/americanexpress/earlybird/v4/pkg/file" "github.com/americanexpress/earlybird/v4/pkg/scan" - "os" - "time" ) // WriteJSON takes the hits, converts them into JSON report and passing report to reportToJSONWriter(). @@ -52,7 +53,7 @@ func WriteJSON(hits <-chan scan.Hit, config cfgReader.EarlybirdConfig, fileConte return err } -//reportToJSONWriter Outputs an object as a JSON blob to an output file or console +// reportToJSONWriter Outputs an object as a JSON blob to an output file or console func reportToJSONWriter(v interface{}, fileName string) (s string, err error) { b, err := json.MarshalIndent(v, "", "\t") if err != nil {