From 849951be1d1a7ee9f9302006ccb187bf5b4e36f3 Mon Sep 17 00:00:00 2001 From: Paul Buetow Date: Wed, 22 Jul 2026 23:51:18 +0300 Subject: =?UTF-8?q?feat:=20DTail=20fork=20=E2=80=94=20server/client=20feat?= =?UTF-8?q?ure=20development?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Squashed development of the snonux/dtail fork's product code (internal/, cmd/) since diverging from mimecast/dtail. Major areas: - Read/output path: the former "turbo" channel-less path is now the single, default server-side read/output path for cat/grep/tail and MapReduce; the old channel-based path and its config/env toggles were removed. - MapReduce: single aggregate implementation (server + serverless) fed directly by a processor pipeline, with input-exhausted finalization via the shutdown coordinator; high-concurrency and data-race fixes. - Journal source reads (journal:unit.service) via journalctl, Linux-gated behind a journal-v1 capability. - Auth-key fast reconnect: in-memory per-user public-key cache with TTL/max-keys, registered over an authenticated session (AUTHKEY), checked before authorized_keys. - Interactive query reload (--interactive-query) with SESSION START/UPDATE generation boundaries and capability negotiation. - Client-side deadlines: --timeout / --shutdownAfter as context deadlines; follow shutdown handling. - Client logging: diagnostics-only daily log by default, opt-in payload tee via --log-payload. - Numerous correctness fixes (buffer-pool double-recycle races, EOF-sentinel leaks, glob-expansion cap, TOCTOU in CSV parsing) with accompanying unit tests. Co-Authored-By: Claude Opus 4.8 --- internal/mapr/aggregateset.go | 22 +- internal/mapr/client/aggregate.go | 73 ++- internal/mapr/client/aggregate_test.go | 96 +++ internal/mapr/client/session_state.go | 95 +++ internal/mapr/client/session_state_test.go | 81 +++ internal/mapr/funcs/function.go | 64 +- internal/mapr/funcs/function_test.go | 114 ++-- internal/mapr/globalgroupset.go | 11 +- internal/mapr/globalgroupset_test.go | 95 +++ internal/mapr/groupset.go | 133 +++- internal/mapr/groupset_avg_nan_test.go | 79 +++ internal/mapr/groupset_ordering_test.go | 125 ++++ internal/mapr/groupset_percentage_test.go | 143 +++++ internal/mapr/groupsetresult.go | 108 ++-- internal/mapr/groupsetresult_renderer_test.go | 113 ++++ internal/mapr/logformat/csv.go | 97 ++- internal/mapr/logformat/csv_test.go | 170 +++++- internal/mapr/logformat/custom1.go | 5 +- internal/mapr/logformat/custom2.go | 5 +- internal/mapr/logformat/default.go | 239 ++++++-- internal/mapr/logformat/default_benchmark_test.go | 44 ++ internal/mapr/logformat/default_test.go | 42 +- internal/mapr/logformat/delimited.go | 12 + internal/mapr/logformat/generic.go | 15 +- internal/mapr/logformat/generickv.go | 36 +- internal/mapr/logformat/mimecast.go | 5 +- internal/mapr/logformat/parser.go | 123 +++- internal/mapr/logformat/parser_test.go | 69 +++ internal/mapr/logformat/variables.go | 107 ++++ internal/mapr/logformat/variables_test.go | 145 +++++ internal/mapr/parserfieldplan.go | 81 +++ internal/mapr/parserfieldplan_test.go | 32 + internal/mapr/query.go | 29 +- internal/mapr/query_test.go | 83 +++ internal/mapr/queryvariables.go | 92 +++ internal/mapr/queryvariables_test.go | 78 +++ internal/mapr/result_renderer.go | 34 ++ internal/mapr/safe_aggregateset.go | 72 +++ internal/mapr/safe_aggregateset_test.go | 147 +++++ internal/mapr/selectcondition.go | 6 + internal/mapr/server/aggregate.go | 620 +++++++++++++------ internal/mapr/server/aggregate_test.go | 714 ++++++++++++++++++++++ internal/mapr/server/groupkey.go | 31 + internal/mapr/server/parsername.go | 10 + internal/mapr/server/parsername_test.go | 62 ++ internal/mapr/token.go | 35 +- internal/mapr/token_test.go | 77 +++ 47 files changed, 4168 insertions(+), 501 deletions(-) create mode 100644 internal/mapr/client/aggregate_test.go create mode 100644 internal/mapr/client/session_state.go create mode 100644 internal/mapr/client/session_state_test.go create mode 100644 internal/mapr/globalgroupset_test.go create mode 100644 internal/mapr/groupset_avg_nan_test.go create mode 100644 internal/mapr/groupset_ordering_test.go create mode 100644 internal/mapr/groupset_percentage_test.go create mode 100644 internal/mapr/groupsetresult_renderer_test.go create mode 100644 internal/mapr/logformat/default_benchmark_test.go create mode 100644 internal/mapr/logformat/delimited.go create mode 100644 internal/mapr/logformat/parser_test.go create mode 100644 internal/mapr/logformat/variables.go create mode 100644 internal/mapr/logformat/variables_test.go create mode 100644 internal/mapr/parserfieldplan.go create mode 100644 internal/mapr/parserfieldplan_test.go create mode 100644 internal/mapr/queryvariables.go create mode 100644 internal/mapr/queryvariables_test.go create mode 100644 internal/mapr/result_renderer.go create mode 100644 internal/mapr/safe_aggregateset.go create mode 100644 internal/mapr/safe_aggregateset_test.go create mode 100644 internal/mapr/server/aggregate_test.go create mode 100644 internal/mapr/server/groupkey.go create mode 100644 internal/mapr/server/parsername.go create mode 100644 internal/mapr/server/parsername_test.go create mode 100644 internal/mapr/token_test.go (limited to 'internal/mapr') diff --git a/internal/mapr/aggregateset.go b/internal/mapr/aggregateset.go index c50c7a1..3281353 100644 --- a/internal/mapr/aggregateset.go +++ b/internal/mapr/aggregateset.go @@ -46,6 +46,10 @@ func (s *AggregateSet) Merge(query *Query, set *AggregateSet) error { case Sum: fallthrough case Avg: + fallthrough + case Percentage: + fallthrough + case Percentile: value := set.FValues[storage] s.addFloat(storage, value) case Min: @@ -67,21 +71,25 @@ func (s *AggregateSet) Merge(query *Query, set *AggregateSet) error { return nil } -// Serialize the aggregate set so it can be sent over the wire. -func (s *AggregateSet) Serialize(ctx context.Context, groupKey string, ch chan<- string) { +// Serialize the aggregate set so it can be sent over the wire. Returns true +// when the serialized message was successfully sent, and false when the +// context was cancelled before the send completed. Callers that own the +// source state (e.g. Aggregate) must re-merge unsent sets so data is +// not silently lost. +func (s *AggregateSet) Serialize(ctx context.Context, groupKey string, ch chan<- string) bool { dlog.Common.Trace("Serialising mapr.AggregateSet", s) sb := pool.BuilderBuffer.Get().(*strings.Builder) defer pool.RecycleBuilderBuffer(sb) sb.WriteString(groupKey) sb.WriteString(protocol.AggregateDelimiter) - sb.WriteString(fmt.Sprintf("%d", s.Samples)) + sb.WriteString(strconv.Itoa(s.Samples)) sb.WriteString(protocol.AggregateDelimiter) for k, v := range s.FValues { sb.WriteString(k) sb.WriteString(protocol.AggregateKVDelimiter) - sb.WriteString(fmt.Sprintf("%v", v)) + sb.WriteString(strconv.FormatFloat(v, 'f', -1, 64)) sb.WriteString(protocol.AggregateDelimiter) } @@ -94,7 +102,9 @@ func (s *AggregateSet) Serialize(ctx context.Context, groupKey string, ch chan<- select { case ch <- sb.String(): + return true case <-ctx.Done(): + return false } } @@ -177,6 +187,10 @@ func (s *AggregateSet) Aggregate(key string, agg AggregateOperation, value strin case Sum: fallthrough case Avg: + fallthrough + case Percentage: + fallthrough + case Percentile: s.addFloat(key, f) case Min: s.addFloatMin(key, f) diff --git a/internal/mapr/client/aggregate.go b/internal/mapr/client/aggregate.go index 2e9b61a..9989e8f 100644 --- a/internal/mapr/client/aggregate.go +++ b/internal/mapr/client/aggregate.go @@ -12,30 +12,46 @@ import ( // Aggregate mapreduce data on the DTail client side. type Aggregate struct { - // This is the mapr query specified on the command line. - query *mapr.Query // This represents aggregated data of a single remote server. group *mapr.GroupSet - // This represents the merged aggregated data of all servers. - globalGroup *mapr.GlobalGroupSet + // Shared per-client session state. + session *SessionState + // The currently tracked shared generation. + generation uint64 // The server we aggregate the data for (logging and debugging purposes only) server string } // NewAggregate create new client aggregator. -func NewAggregate(server string, query *mapr.Query, - globalGroup *mapr.GlobalGroupSet) *Aggregate { +func NewAggregate(server string, session *SessionState) *Aggregate { + generation := uint64(0) + if session != nil { + generation = session.Snapshot().Generation + } return &Aggregate{ - query: query, - group: mapr.NewGroupSet(), - globalGroup: globalGroup, - server: server, + group: mapr.NewGroupSet(), + session: session, + generation: generation, + server: server, } } // Aggregate data from mapr log line into local (and global) group sets. func (a *Aggregate) Aggregate(message string) error { + if a.session == nil { + return fmt.Errorf("missing client mapreduce session state") + } + + snapshot := a.session.Snapshot() + if snapshot.Query == nil || snapshot.GlobalGroup == nil { + return fmt.Errorf("missing client mapreduce query state") + } + if snapshot.Generation != a.generation { + a.group.InitSet() + a.generation = snapshot.Generation + } + parts := strings.Split(message, protocol.AggregateDelimiter) if len(parts) < 4 { return fmt.Errorf("aggregate message without any real data") @@ -51,7 +67,7 @@ func (a *Aggregate) Aggregate(message string) error { set := a.group.GetSet(groupKey) var addedSamples bool - for _, sc := range a.query.Select { + for _, sc := range snapshot.Query.Select { if val, ok := fields[sc.FieldStorage]; ok { if err := set.Aggregate(sc.FieldStorage, sc.Operation, val, true); err != nil { dlog.Client.Error(err) @@ -65,9 +81,9 @@ func (a *Aggregate) Aggregate(message string) error { } // Merge data from group into global group. - isMerged, err := a.globalGroup.MergeNoblock(a.query, a.group) + isMerged, err := snapshot.GlobalGroup.MergeNoblock(snapshot.Query, a.group) if err != nil { - panic(err) + return fmt.Errorf("unable to merge aggregate data for server %s: %w", a.server, err) } if isMerged { // Re-init local group (make it empty again). @@ -76,15 +92,40 @@ func (a *Aggregate) Aggregate(message string) error { return nil } +// Flush merges any pending per-server aggregate state into the shared global group. +// The normal hot path uses MergeNoblock to avoid stalling on the global merge lock. +// During shutdown we need a blocking flush so the last local batch is not lost. +func (a *Aggregate) Flush() error { + if a.session == nil { + return fmt.Errorf("missing client mapreduce session state") + } + + snapshot := a.session.Snapshot() + if snapshot.Query == nil || snapshot.GlobalGroup == nil { + return nil + } + if snapshot.Generation != a.generation { + a.group.InitSet() + a.generation = snapshot.Generation + return nil + } + + if err := snapshot.GlobalGroup.Merge(snapshot.Query, a.group); err != nil { + return fmt.Errorf("unable to flush aggregate data for server %s: %w", a.server, err) + } + a.group.InitSet() + return nil +} + // Create a map of key-value pairs from a part list such as ["foo=bar", "bar=baz"]. func (a *Aggregate) makeFields(parts []string) map[string]string { fields := make(map[string]string, len(parts)) for _, part := range parts { - kv := strings.SplitN(part, protocol.AggregateKVDelimiter, 2) - if len(kv) != 2 { + key, value, ok := strings.Cut(part, protocol.AggregateKVDelimiter) + if !ok { continue } - fields[kv[0]] = kv[1] + fields[key] = value } return fields } diff --git a/internal/mapr/client/aggregate_test.go b/internal/mapr/client/aggregate_test.go new file mode 100644 index 0000000..3387a63 --- /dev/null +++ b/internal/mapr/client/aggregate_test.go @@ -0,0 +1,96 @@ +package client + +import ( + "strings" + "testing" + + "github.com/mimecast/dtail/internal/mapr" + "github.com/mimecast/dtail/internal/protocol" +) + +func TestAggregateResetsPendingLocalStateOnGenerationChange(t *testing.T) { + query := mustSessionStateQuery(t, "select status,count(status) from stats group by status") + state := NewSessionState(query) + aggregate := NewAggregate("srv1", state) + countStorage := aggregateCountStorage(t, query) + + oldSet := aggregate.group.GetSet("ERROR") + oldSet.Samples = 1 + oldSet.FValues[countStorage] = 1 + + rawQuery := "select status,count(status) from warnings group by status" + if _, err := state.CommitQuery(rawQuery, 2); err != nil { + t.Fatalf("CommitQuery() error = %v", err) + } + + snapshot := state.Snapshot() + message := strings.Join([]string{ + "WARN", + "1", + aggregateCountStorage(t, snapshot.Query) + protocol.AggregateKVDelimiter + "1", + "", + }, protocol.AggregateDelimiter) + + if err := aggregate.Aggregate(message); err != nil { + t.Fatalf("Aggregate() error = %v", err) + } + + result, numRows, err := snapshot.GlobalGroup.Result(snapshot.Query, 10, nil) + if err != nil { + t.Fatalf("Result() error = %v", err) + } + if numRows != 1 { + t.Fatalf("numRows = %d, want 1", numRows) + } + if !strings.Contains(result, "1") { + t.Fatalf("expected one new-generation aggregate row, got %q", result) + } +} + +func TestAggregateRejectsMalformedMessage(t *testing.T) { + query := mustSessionStateQuery(t, "select count(status) from stats group by status") + state := NewSessionState(query) + aggregate := NewAggregate("srv1", state) + + if err := aggregate.Aggregate("broken"); err == nil { + t.Fatalf("expected Aggregate() to reject malformed messages") + } +} + +func TestAggregateFlushMergesPendingLocalState(t *testing.T) { + query := mustSessionStateQuery(t, "select status,count(status) from stats group by status") + state := NewSessionState(query) + aggregate := NewAggregate("srv1", state) + countStorage := aggregateCountStorage(t, query) + + set := aggregate.group.GetSet("ERROR") + set.Samples = 3 + set.FValues[countStorage] = 3 + + if err := aggregate.Flush(); err != nil { + t.Fatalf("Flush() error = %v", err) + } + + result, numRows, err := state.Snapshot().GlobalGroup.Result(query, 10, nil) + if err != nil { + t.Fatalf("Result() error = %v", err) + } + if numRows != 1 { + t.Fatalf("numRows = %d, want 1", numRows) + } + if !strings.Contains(result, "3") { + t.Fatalf("expected flushed aggregate row, got %q", result) + } +} + +func aggregateCountStorage(t *testing.T, query *mapr.Query) string { + t.Helper() + + for _, selectCondition := range query.Select { + if selectCondition.Operation == mapr.Count { + return selectCondition.FieldStorage + } + } + t.Fatalf("query %q does not contain count() storage", query.RawQuery) + return "" +} diff --git a/internal/mapr/client/session_state.go b/internal/mapr/client/session_state.go new file mode 100644 index 0000000..1983644 --- /dev/null +++ b/internal/mapr/client/session_state.go @@ -0,0 +1,95 @@ +package client + +import ( + "fmt" + "sync" + + "github.com/mimecast/dtail/internal/mapr" +) + +// SessionSnapshot captures the current client-side mapreduce session state. +type SessionSnapshot struct { + Generation uint64 + Query *mapr.Query + GlobalGroup *mapr.GlobalGroupSet + LastResult string +} + +// SessionState keeps the mutable mapreduce query state shared by the client +// reporter and per-server handlers. +type SessionState struct { + mu sync.RWMutex + generation uint64 + query *mapr.Query + global *mapr.GlobalGroupSet + lastResult string + changedCh chan struct{} +} + +// NewSessionState returns a new shared mapreduce session state. +func NewSessionState(query *mapr.Query) *SessionState { + return &SessionState{ + query: query, + global: mapr.NewGlobalGroupSet(), + changedCh: make(chan struct{}, 1), + } +} + +// Snapshot returns a point-in-time copy of the shared mapreduce state. +func (s *SessionState) Snapshot() SessionSnapshot { + s.mu.RLock() + defer s.mu.RUnlock() + + return SessionSnapshot{ + Generation: s.generation, + Query: s.query, + GlobalGroup: s.global, + LastResult: s.lastResult, + } +} + +// Changes returns a channel that is signaled whenever a new generation is committed. +func (s *SessionState) Changes() <-chan struct{} { + return s.changedCh +} + +// CommitQuery resets the shared aggregation state for a newly accepted query generation. +func (s *SessionState) CommitQuery(rawQuery string, generation uint64) (*mapr.Query, error) { + query, err := mapr.NewQuery(rawQuery) + if err != nil { + return nil, fmt.Errorf("parse session query: %w", err) + } + + s.mu.Lock() + s.generation = generation + s.query = query + s.global = mapr.NewGlobalGroupSet() + s.lastResult = "" + s.mu.Unlock() + + s.notifyChange() + return query, nil +} + +// CommitRenderedResult stores the last rendered result for the active generation. +func (s *SessionState) CommitRenderedResult(generation uint64, result string) (changed bool, ok bool) { + s.mu.Lock() + defer s.mu.Unlock() + + if s.generation != generation { + return false, false + } + if s.lastResult == result { + return false, true + } + + s.lastResult = result + return true, true +} + +func (s *SessionState) notifyChange() { + select { + case s.changedCh <- struct{}{}: + default: + } +} diff --git a/internal/mapr/client/session_state_test.go b/internal/mapr/client/session_state_test.go new file mode 100644 index 0000000..f43ca70 --- /dev/null +++ b/internal/mapr/client/session_state_test.go @@ -0,0 +1,81 @@ +package client + +import ( + "testing" + + "github.com/mimecast/dtail/internal/mapr" +) + +func TestSessionStateCommitQueryResetsGenerationAndResults(t *testing.T) { + query := mustSessionStateQuery(t, "select count(status) from stats group by status") + state := NewSessionState(query) + + initial := state.Snapshot() + group := mapr.NewGroupSet() + set := group.GetSet("ERROR") + set.Samples = 1 + set.FValues[query.Select[0].FieldStorage] = 1 + if err := initial.GlobalGroup.Merge(query, group); err != nil { + t.Fatalf("Merge() error = %v", err) + } + if changed, ok := state.CommitRenderedResult(initial.Generation, "old-result"); !ok || !changed { + t.Fatalf("CommitRenderedResult() = changed:%v ok:%v, want changed and ok", changed, ok) + } + + rawQuery := "select count(status) from warnings group by status" + updatedQuery, err := state.CommitQuery(rawQuery, 3) + if err != nil { + t.Fatalf("CommitQuery() error = %v", err) + } + if updatedQuery == nil || updatedQuery.RawQuery != rawQuery { + t.Fatalf("unexpected updated query: %#v", updatedQuery) + } + + select { + case <-state.Changes(): + default: + t.Fatalf("expected change notification after CommitQuery") + } + + updated := state.Snapshot() + if updated.Generation != 3 { + t.Fatalf("generation = %d, want 3", updated.Generation) + } + if updated.Query == nil || updated.Query.RawQuery != rawQuery { + t.Fatalf("unexpected query after commit: %#v", updated.Query) + } + if !updated.GlobalGroup.IsEmpty() { + t.Fatalf("expected committed global group to be reset") + } + if updated.LastResult != "" { + t.Fatalf("last result = %q, want empty", updated.LastResult) + } +} + +func TestSessionStateCommitQueryRejectsInvalidQuery(t *testing.T) { + query := mustSessionStateQuery(t, "select count(status) from stats group by status") + state := NewSessionState(query) + before := state.Snapshot() + + if _, err := state.CommitQuery("select from", 5); err == nil { + t.Fatalf("expected CommitQuery() to reject invalid query") + } + + after := state.Snapshot() + if after.Generation != before.Generation { + t.Fatalf("generation changed on invalid query: got %d want %d", after.Generation, before.Generation) + } + if after.Query == nil || after.Query.RawQuery != before.Query.RawQuery { + t.Fatalf("query changed on invalid query: before=%#v after=%#v", before.Query, after.Query) + } +} + +func mustSessionStateQuery(t *testing.T, queryStr string) *mapr.Query { + t.Helper() + + query, err := mapr.NewQuery(queryStr) + if err != nil { + t.Fatalf("NewQuery(%q) error = %v", queryStr, err) + } + return query +} diff --git a/internal/mapr/funcs/function.go b/internal/mapr/funcs/function.go index 418d86f..2f21d5a 100644 --- a/internal/mapr/funcs/function.go +++ b/internal/mapr/funcs/function.go @@ -20,20 +20,10 @@ type Function struct { type FunctionStack []Function // NewFunctionStack parses the input string, e.g. foo(bar("arg")) and returns -// a corresponding function stack. +// a corresponding function stack. It returns an error for malformed inputs +// such as unbalanced parentheses (e.g. "foo(", "foo(bar)baz"). func NewFunctionStack(in string) (FunctionStack, string, error) { var fs FunctionStack - getCallback := func(name string) (CallbackFunc, error) { - var cb CallbackFunc - switch name { - case "md5sum": - return Md5Sum, nil - case "maskdigits": - return MaskDigits, nil - default: - return cb, fmt.Errorf("unknown function '%s'", name) - } - } aux := in for strings.HasSuffix(aux, ")") { @@ -43,16 +33,64 @@ func NewFunctionStack(in string) (FunctionStack, string, error) { } name := aux[0:index] - call, err := getCallback(name) + call, err := lookupCallback(name) if err != nil { return fs, "", err } fs = append(fs, Function{name, call}) + // Strip the outer function name and its enclosing parens, leaving + // only the argument expression for the next iteration. aux = aux[index+1 : len(aux)-1] } + + // Validate that no unbalanced parentheses remain in the argument string. + // Inputs like "foo(bar)baz" leave "bar)baz" after stripping, and inputs + // ending with "(" (no closing ")") are accepted as plain field literals + // without this check — both produce silently wrong behavior. + if err := validateParenBalance(aux, in); err != nil { + return fs, "", err + } + return fs, aux, nil } +// lookupCallback maps a function name to its CallbackFunc implementation. +// It returns an error for unrecognised names so callers get a clear message. +func lookupCallback(name string) (CallbackFunc, error) { + switch name { + case "md5sum": + return Md5Sum, nil + case "maskdigits": + return MaskDigits, nil + default: + var zero CallbackFunc + return zero, fmt.Errorf("unknown function '%s'", name) + } +} + +// validateParenBalance checks that the remaining argument string contains no +// unbalanced parentheses. A negative depth means a stray ')' was found; a +// non-zero depth after the loop means an unclosed '(' was found. The original +// full expression is included in the error message for context. +func validateParenBalance(aux, original string) error { + depth := 0 + for _, r := range aux { + switch r { + case '(': + depth++ + case ')': + depth-- + } + if depth < 0 { + return fmt.Errorf("malformed function expression %q: unexpected ')' in argument", original) + } + } + if depth != 0 { + return fmt.Errorf("malformed function expression %q: unclosed '(' in argument", original) + } + return nil +} + // Call the function stack. func (fs FunctionStack) Call(str string) string { for i := len(fs) - 1; i >= 0; i-- { diff --git a/internal/mapr/funcs/function_test.go b/internal/mapr/funcs/function_test.go index 8b5d8b7..8227817 100644 --- a/internal/mapr/funcs/function_test.go +++ b/internal/mapr/funcs/function_test.go @@ -2,51 +2,91 @@ package funcs import "testing" -func TestFunction(t *testing.T) { - input := "md5sum($line)" - fs, arg, err := NewFunctionStack(input) - if err != nil { - t.Errorf("error parsing function input '%s': %s (%v)\n", - input, err.Error(), fs) - } - if arg != "$line" { - t.Errorf("error parsing function input '%s': expected argument '$line' but "+ - "got '%s' (%v)\n", input, arg, fs) - } - t.Log(input, fs, arg) +func TestFunctionStackValid(t *testing.T) { + t.Parallel() - result := fs.Call(input) - if result != "b38699013d79e50d9d122433753959c1" { - t.Errorf("error executing function stack '%s': expected result "+ - "'b38699013d79e50d9d122433753959c1' but got '%s' (%v)\n", input, result, fs) + type want struct { + arg string + // result of calling the returned function stack on the original input + callResult string } - input = "maskdigits(md5sum(maskdigits($line)))" - fs, arg, err = NewFunctionStack(input) - if err != nil { - t.Errorf("error parsing function input '%s': %s (%v)\n", input, err.Error(), fs) - } - if arg != "$line" { - t.Errorf("error parsing function input '%s': expected argument '$line' but "+ - "got '%s' (%v)\n", input, arg, fs) + cases := []struct { + input string + want want + }{ + { + input: "md5sum($line)", + want: want{arg: "$line", callResult: "b38699013d79e50d9d122433753959c1"}, + }, + { + input: "maskdigits(md5sum(maskdigits($line)))", + want: want{arg: "$line", callResult: ".fac.bbe..bb.........d...a.c..b."}, + }, + { + // An argument containing nested parens that are balanced is valid. + input: "md5sum($foo)", + want: want{arg: "$foo"}, + }, + { + // Plain field with no function wrapper is a degenerate stack (empty). + input: "$line", + want: want{arg: "$line"}, + }, } - t.Log(input, fs, arg) - result = fs.Call(input) - if result != ".fac.bbe..bb.........d...a.c..b." { - t.Errorf("error executing function stack '%s': expected result "+ - "'.fac.bbe..bb.........d...a.c..b.' but got '%s' (%v)\n", input, result, fs) + for _, tc := range cases { + tc := tc + t.Run(tc.input, func(t *testing.T) { + t.Parallel() + fs, arg, err := NewFunctionStack(tc.input) + if err != nil { + t.Fatalf("unexpected error for input %q: %v (stack %v)", tc.input, err, fs) + } + if arg != tc.want.arg { + t.Errorf("arg: got %q, want %q", arg, tc.want.arg) + } + if tc.want.callResult != "" { + got := fs.Call(tc.input) + if got != tc.want.callResult { + t.Errorf("Call(%q) = %q, want %q", tc.input, got, tc.want.callResult) + } + } + }) } +} + +// TestFunctionStackMalformed verifies that NewFunctionStack rejects expressions +// that are structurally invalid. Before the fix, several of these were silently +// accepted and produced wrong results. +func TestFunctionStackMalformed(t *testing.T) { + t.Parallel() - input = "md5sum$line)" - if fs, _, err := NewFunctionStack(input); err == nil { - t.Errorf("Expected error parsing function input '%s' (%v) but got no error\n", - input, fs) + cases := []string{ + // Missing opening paren — no function call syntax at all. + "md5sum$line)", + // Known outer function but inner call is missing its closing paren. + "md5sum(makedigits$line))", + // Stray ')' inside the argument after stripping: "bar)baz" remains. + // Before the fix this was silently accepted and produced wrong output. + "md5sum(bar)baz)", + // Input ends with '(' — no closing ')' so the loop never strips, + // but the argument string itself contains an unclosed '('. + // Before the fix this was accepted as a plain field literal. + "foo(", + // Empty outer call — the name portion is empty (index == 0) which + // is caught by the existing index <= 0 guard. + "()", } - input = "md5sum(makedigits$line))" - if fs, _, err := NewFunctionStack(input); err == nil { - t.Errorf("Expected error parsing function input '%s' (%v) but got no error\n", - input, fs) + for _, input := range cases { + input := input + t.Run(input, func(t *testing.T) { + t.Parallel() + fs, _, err := NewFunctionStack(input) + if err == nil { + t.Errorf("expected error for malformed input %q but got none (stack %v)", input, fs) + } + }) } } diff --git a/internal/mapr/globalgroupset.go b/internal/mapr/globalgroupset.go index 2b12898..cbee303 100644 --- a/internal/mapr/globalgroupset.go +++ b/internal/mapr/globalgroupset.go @@ -33,12 +33,13 @@ func (g *GlobalGroupSet) Merge(query *Query, group *GroupSet) error { } // MergeNoblock merges (non-blocking) a group set into the global group set. +// The semaphore is released via defer so it is always returned even if +// g.merge panics, mirroring the guarantee provided by the blocking Merge. func (g *GlobalGroupSet) MergeNoblock(query *Query, group *GroupSet) (bool, error) { select { case g.semaphore <- struct{}{}: - err := g.merge(query, group) - <-g.semaphore - return true, err + defer func() { <-g.semaphore }() + return true, g.merge(query, group) default: return false, nil } @@ -86,8 +87,8 @@ func (g *GlobalGroupSet) WriteResult(query *Query, finalResult bool) error { } // Result returns the result of the mapreduce aggregation as a string. -func (g *GlobalGroupSet) Result(query *Query, rowsLimit int) (string, int, error) { +func (g *GlobalGroupSet) Result(query *Query, rowsLimit int, renderer ResultRenderer) (string, int, error) { g.semaphore <- struct{}{} defer func() { <-g.semaphore }() - return g.GroupSet.Result(query, rowsLimit) + return g.GroupSet.Result(query, rowsLimit, renderer) } diff --git a/internal/mapr/globalgroupset_test.go b/internal/mapr/globalgroupset_test.go new file mode 100644 index 0000000..3e2a9e5 --- /dev/null +++ b/internal/mapr/globalgroupset_test.go @@ -0,0 +1,95 @@ +package mapr + +import ( + "testing" + "time" +) + +// TestMergeNoblockSemaphoreReleasedOnPanic verifies that MergeNoblock releases +// the semaphore even when g.merge panics (e.g. due to a nil GroupSet). +// Without the fix (using defer), the semaphore would be leaked and subsequent +// calls like NumSets would deadlock forever. +func TestMergeNoblockSemaphoreReleasedOnPanic(t *testing.T) { + g := NewGlobalGroupSet() + + // Calling MergeNoblock with a nil *GroupSet causes a nil-pointer dereference + // inside g.merge when it iterates over group.sets. We catch the panic in a + // goroutine and verify that the GlobalGroupSet is still usable afterwards. + done := make(chan struct{}) + go func() { + defer func() { + // Recover the expected panic so the goroutine exits cleanly. + if r := recover(); r == nil { + t.Errorf("expected a panic from MergeNoblock with nil GroupSet, got none") + } + close(done) + }() + // This must panic internally; with the bug the semaphore is never released. + //nolint:staticcheck // intentional nil dereference to exercise the panic path + g.MergeNoblock(nil, nil) //nolint:errcheck + }() + + // Wait for the goroutine to finish (panic recovered). + select { + case <-done: + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for MergeNoblock panic to be recovered") + } + + // After the panic the semaphore must have been released by the deferred + // release in MergeNoblock. If the bug is present NumSets acquires the same + // 1-slot semaphore and blocks forever, causing the test to time out. + result := make(chan int, 1) + go func() { + result <- g.NumSets() + }() + + select { + case n := <-result: + if n != 0 { + t.Errorf("expected 0 sets in empty GlobalGroupSet, got %d", n) + } + case <-time.After(5 * time.Second): + t.Fatal("NumSets deadlocked: semaphore was not released after MergeNoblock panic (bug reproduced)") + } +} + +// TestMergeNoblockNormalOperation verifies the non-panic happy path still works +// correctly: a successful merge returns (true, nil) and NumSets reflects the +// merged data. +func TestMergeNoblockNormalOperation(t *testing.T) { + g := NewGlobalGroupSet() + group := NewGroupSet() + + // Populate the group set with one entry so there is something to merge. + set := NewAggregateSet() + set.FValues["count"] = 1 + group.sets["key1"] = set + + // A minimal query is enough; the merge loop only needs query.Select which + // can be empty for this structural test (no select conditions to iterate). + query := &Query{} + + merged, err := g.MergeNoblock(query, group) + if err != nil { + t.Errorf("unexpected error from MergeNoblock: %v", err) + } + if !merged { + t.Error("expected MergeNoblock to return merged=true when semaphore is free") + } + + // After merging, NumSets must return 1 and must not deadlock. + result := make(chan int, 1) + go func() { + result <- g.NumSets() + }() + + select { + case n := <-result: + if n != 1 { + t.Errorf("expected 1 set after merge, got %d", n) + } + case <-time.After(5 * time.Second): + t.Fatal("NumSets deadlocked after normal MergeNoblock (unexpected)") + } +} diff --git a/internal/mapr/groupset.go b/internal/mapr/groupset.go index 9d7661a..f544b35 100644 --- a/internal/mapr/groupset.go +++ b/internal/mapr/groupset.go @@ -24,6 +24,11 @@ type result struct { orderBy float64 } +type resultStats struct { + percentageTotals map[string]float64 + percentileValues map[string][]float64 +} + // NewGroupSet returns a new empty group set. func NewGroupSet() *GroupSet { g := GroupSet{} @@ -51,14 +56,53 @@ func (g *GroupSet) GetSet(groupKey string) *AggregateSet { return set } -// Serialize the group set (e.g. to send it over the wire). -func (g *GroupSet) Serialize(ctx context.Context, ch chan<- string) { +// Serialize the group set (e.g. to send it over the wire). If the context is +// cancelled mid-iteration, the remaining unsent aggregate sets are returned +// so callers can retry them (for example, by re-merging them into the live +// aggregation state). The returned map is nil when every entry was sent. +func (g *GroupSet) Serialize(ctx context.Context, ch chan<- string) map[string]*AggregateSet { + var remaining map[string]*AggregateSet + aborted := false for groupKey, set := range g.sets { - set.Serialize(ctx, groupKey, ch) + if aborted { + if remaining == nil { + remaining = make(map[string]*AggregateSet, len(g.sets)) + } + remaining[groupKey] = set + continue + } + if !set.Serialize(ctx, groupKey, ch) { + aborted = true + if remaining == nil { + remaining = make(map[string]*AggregateSet, len(g.sets)) + } + remaining[groupKey] = set + } } + return remaining +} + +// ResetWith replaces the underlying sets map. A nil argument is equivalent +// to InitSet. This is the supported way for callers in the same package to +// restore unsent data returned from Serialize without reaching into the +// unexported sets field. +func (g *GroupSet) ResetWith(sets map[string]*AggregateSet) { + if sets == nil { + g.InitSet() + return + } + g.sets = sets } // Return a sorted result slice of the query from the group set. +// +// Rows are built in lexicographic groupKey order first. This guarantees a +// stable, deterministic base ordering before any OrderBy sort is applied. +// Without the pre-sort, Go's intentionally randomised map iteration would +// make output order non-deterministic when OrderBy is empty, and would +// produce non-deterministic tie-breaks when multiple rows share the same +// OrderBy value (SortStable preserves incoming order, so random map order +// propagated directly into tied rows). func (g *GroupSet) result(query *Query, gathercolumnWidths bool) ([]result, []int, error) { var err error var rows []result @@ -67,12 +111,19 @@ func (g *GroupSet) result(query *Query, gathercolumnWidths bool) ([]result, []in // not a CSV file). columnWidths := make([]int, len(query.Select)) var valueStrLen int + stats := g.makeResultStats(query) - for groupKey, set := range g.sets { + // Collect and sort group keys lexicographically so that the row slice is + // built in a deterministic order. SortStable in resultOrderBy then + // preserves this order for tied OrderBy values. + keys := sortedGroupKeys(g.sets) + + for _, groupKey := range keys { + set := g.sets[groupKey] result := result{groupKey: groupKey} for i, sc := range query.Select { - if valueStrLen, err = g.resultSelect(query, &sc, set, &result); err != nil { + if valueStrLen, err = g.resultSelect(query, &sc, set, &result, &stats); err != nil { return rows, columnWidths, err } @@ -95,8 +146,21 @@ func (g *GroupSet) result(query *Query, gathercolumnWidths bool) ([]result, []in return rows, columnWidths, nil } +// sortedGroupKeys returns the keys of the given sets map sorted +// lexicographically. This helper centralises the deterministic key extraction +// used by result() and makeResultStats() to guarantee consistent iteration +// order regardless of Go's runtime map randomisation. +func sortedGroupKeys(sets map[string]*AggregateSet) []string { + keys := make([]string, 0, len(sets)) + for k := range sets { + keys = append(keys, k) + } + sort.Strings(keys) + return keys +} + func (*GroupSet) resultSelect(query *Query, sc *selectCondition, set *AggregateSet, - result *result) (int, error) { + result *result, stats *resultStats) (int, error) { var valueStr string var value float64 @@ -118,7 +182,26 @@ func (*GroupSet) resultSelect(query *Query, sc *selectCondition, set *AggregateS valueStr = set.SValues[sc.FieldStorage] value, _ = strconv.ParseFloat(valueStr, 64) case Avg: - value = set.FValues[sc.FieldStorage] / float64(set.Samples) + // Guard against division by zero when an empty aggregate set (Samples==0) + // is received from the server. Without this guard, 0/0 yields NaN, which + // propagates as the string "NaN" into CSV/terminal output. + if set.Samples == 0 { + value = 0 + } else { + value = set.FValues[sc.FieldStorage] / float64(set.Samples) + } + valueStr = fmt.Sprintf("%f", value) + case Percentage: + value = set.FValues[sc.FieldStorage] + total := stats.percentageTotals[sc.FieldStorage] + if total == 0 { + value = 0 + } else { + value = (value / total) * 100 + } + valueStr = fmt.Sprintf("%f", value) + case Percentile: + value = percentileRank(set.FValues[sc.FieldStorage], stats.percentileValues[sc.FieldStorage]) valueStr = fmt.Sprintf("%f", value) default: return 0, fmt.Errorf("Unknown aggregation method '%v'", sc.Operation) @@ -132,6 +215,42 @@ func (*GroupSet) resultSelect(query *Query, sc *selectCondition, set *AggregateS return len(valueStr), nil } +func (g *GroupSet) makeResultStats(query *Query) resultStats { + stats := resultStats{ + percentageTotals: make(map[string]float64), + percentileValues: make(map[string][]float64), + } + + for _, set := range g.sets { + for _, sc := range query.Select { + value := set.FValues[sc.FieldStorage] + switch sc.Operation { + case Percentage: + stats.percentageTotals[sc.FieldStorage] += value + case Percentile: + stats.percentileValues[sc.FieldStorage] = append(stats.percentileValues[sc.FieldStorage], value) + } + } + } + + for storage := range stats.percentileValues { + sort.Float64s(stats.percentileValues[storage]) + } + + return stats +} + +func percentileRank(value float64, sortedValues []float64) float64 { + if len(sortedValues) == 0 { + return 0 + } + + upperBound := sort.Search(len(sortedValues), func(i int) bool { + return sortedValues[i] > value + }) + return (float64(upperBound) / float64(len(sortedValues))) * 100 +} + func (*GroupSet) resultOrderBy(query *Query, rows []result) { if query.OrderBy == "" { return diff --git a/internal/mapr/groupset_avg_nan_test.go b/internal/mapr/groupset_avg_nan_test.go new file mode 100644 index 0000000..c3d48a4 --- /dev/null +++ b/internal/mapr/groupset_avg_nan_test.go @@ -0,0 +1,79 @@ +package mapr + +import ( + "math" + "strconv" + "strings" + "testing" +) + +// TestGroupSetAvgZeroSamplesDoesNotProduceNaN is a negative test that reproduces +// the bug where an empty aggregate set (Samples==0) causes 0/0 = NaN in the Avg +// case of resultSelect. This happens when the server creates a group-set entry +// via GetSet before confirming that any select fields matched, then serialises and +// sends the empty set to the client. The client-side resultSelect must guard the +// Avg division so that Samples==0 yields 0 instead of NaN. +func TestGroupSetAvgZeroSamplesDoesNotProduceNaN(t *testing.T) { + t.Parallel() + + query, err := NewQuery("select avg(latency) from stats group by host") + if err != nil { + t.Fatalf("Unable to parse query: %v", err) + } + + groupSet := NewGroupSet() + + // Simulate what the server does when no log line fields match the select + // clause: GetSet creates the entry, but Samples stays 0 and FValues is + // never populated. This is the bug trigger — previously 0/0 = NaN. + _ = groupSet.GetSet("host-a") + + rows, _, err := groupSet.result(query, false) + if err != nil { + t.Fatalf("result() returned unexpected error: %v", err) + } + if len(rows) != 1 { + t.Fatalf("Expected 1 result row (even for empty set), got %d", len(rows)) + } + + // Before the fix each floating-point value in the row was the string "NaN". + for _, row := range rows { + for _, v := range row.values { + trimmed := strings.TrimSpace(v) + f, parseErr := strconv.ParseFloat(trimmed, 64) + if parseErr != nil { + // Non-numeric values (e.g. integer count or last-string fields) + // are fine; only floating-point results can be NaN. + continue + } + if math.IsNaN(f) { + t.Errorf("avg on empty set produced NaN in output %q; expected 0", v) + } + } + } +} + +// TestGroupSetAvgZeroSamplesResultOutputContainsNoNaN verifies that the +// higher-level Result method (which drives terminal output) also never emits +// "NaN" strings when aggregate sets have zero samples. +func TestGroupSetAvgZeroSamplesResultOutputContainsNoNaN(t *testing.T) { + t.Parallel() + + query, err := NewQuery("select avg(latency) from stats group by host") + if err != nil { + t.Fatalf("Unable to parse query: %v", err) + } + + groupSet := NewGroupSet() + // Empty set — Samples==0, no FValues populated. + _ = groupSet.GetSet("host-a") + + output, _, err := groupSet.Result(query, 100, nil) + if err != nil { + t.Fatalf("Result() returned unexpected error: %v", err) + } + + if strings.Contains(output, "NaN") { + t.Errorf("Result output must not contain 'NaN', got:\n%s", output) + } +} diff --git a/internal/mapr/groupset_ordering_test.go b/internal/mapr/groupset_ordering_test.go new file mode 100644 index 0000000..18b845f --- /dev/null +++ b/internal/mapr/groupset_ordering_test.go @@ -0,0 +1,125 @@ +package mapr + +import ( + "reflect" + "testing" +) + +// TestGroupSetResultOrderIsDeterministicWithoutOrderBy is a negative test that +// reproduces the non-determinism bug in result(): when OrderBy is unset, the +// output row order depended on Go's map iteration order, which is intentionally +// randomised per runtime invocation. Two consecutive calls to result() on the +// same GroupSet could return rows in different orders. +// +// The fix collects group keys, sorts them lexicographically before building +// rows, and only then applies SortStable for the OrderBy pass. Ties on OrderBy +// (or no OrderBy) therefore resolve to lexicographic groupKey order rather than +// to random map iteration order. +func TestGroupSetResultOrderIsDeterministicWithoutOrderBy(t *testing.T) { + t.Parallel() + + // Query with no ORDER BY clause — the bug case where map iteration order + // was the sole determinant of row order. + query, err := NewQuery("select count(line) from logs group by host") + if err != nil { + t.Fatalf("Unable to parse query: %v", err) + } + + groupSet := NewGroupSet() + + // Insert keys in reverse lexicographic order to ensure the expected sorted + // order cannot coincide with insertion order. + for _, host := range []string{"host-z", "host-m", "host-a", "host-b"} { + set := groupSet.GetSet(host) + if err := set.Aggregate("count(line)", Count, "1", false); err != nil { + t.Fatalf("Aggregate failed for %s: %v", host, err) + } + } + + // Run result() many times. Before the fix a handful of iterations was + // enough to observe a different ordering; with the fix every call must + // return exactly the same lexicographically sorted sequence of groupKeys. + var firstKeys []string + const iterations = 50 + for i := range iterations { + rows, _, err := groupSet.result(query, false) + if err != nil { + t.Fatalf("result() iteration %d returned error: %v", i, err) + } + if len(rows) != 4 { + t.Fatalf("Expected 4 rows, got %d on iteration %d", len(rows), i) + } + + keys := make([]string, len(rows)) + for j, r := range rows { + keys[j] = r.groupKey + } + + if i == 0 { + firstKeys = keys + // Verify the order is lexicographic (the contract of the fix). + expected := []string{"host-a", "host-b", "host-m", "host-z"} + if !reflect.DeepEqual(keys, expected) { + t.Fatalf("First result not in lexicographic order: got %v, want %v", keys, expected) + } + continue + } + + // Every subsequent call must return the identical key sequence. + if !reflect.DeepEqual(keys, firstKeys) { + t.Fatalf("Non-deterministic ordering detected on iteration %d: got %v, want %v", i, keys, firstKeys) + } + } +} + +// TestGroupSetResultOrderIsDeterministicWithOrderByTies verifies that when +// multiple rows share the same OrderBy value (a tie), the tie-break falls back +// to lexicographic groupKey order rather than to random map iteration order. +// SortStable preserves the relative order of equal elements, so the pre-sort of +// keys guarantees a deterministic tie-break. +func TestGroupSetResultOrderIsDeterministicWithOrderByTies(t *testing.T) { + t.Parallel() + + // ORDER BY count(line) — all rows will have the same count (1), creating a + // full tie that must resolve to lexicographic groupKey order. + query, err := NewQuery("select count(line) from logs group by host order by count(line)") + if err != nil { + t.Fatalf("Unable to parse query: %v", err) + } + + groupSet := NewGroupSet() + + // All hosts receive the same count value to force a tie. + for _, host := range []string{"host-z", "host-m", "host-a", "host-b"} { + set := groupSet.GetSet(host) + if err := set.Aggregate("count(line)", Count, "1", false); err != nil { + t.Fatalf("Aggregate failed for %s: %v", host, err) + } + } + + var firstKeys []string + const iterations = 50 + for i := range iterations { + rows, _, err := groupSet.result(query, false) + if err != nil { + t.Fatalf("result() iteration %d returned error: %v", i, err) + } + if len(rows) != 4 { + t.Fatalf("Expected 4 rows, got %d on iteration %d", len(rows), i) + } + + keys := make([]string, len(rows)) + for j, r := range rows { + keys[j] = r.groupKey + } + + if i == 0 { + firstKeys = keys + continue + } + + if !reflect.DeepEqual(keys, firstKeys) { + t.Fatalf("Non-deterministic tie-break detected on iteration %d: got %v, want %v", i, keys, firstKeys) + } + } +} diff --git a/internal/mapr/groupset_percentage_test.go b/internal/mapr/groupset_percentage_test.go new file mode 100644 index 0000000..1273859 --- /dev/null +++ b/internal/mapr/groupset_percentage_test.go @@ -0,0 +1,143 @@ +package mapr + +import ( + "strconv" + "testing" +) + +func TestGroupSetResultPercentageAndPercentile(t *testing.T) { + query, err := NewQuery("select percentage(value),percentile(value) from stats group by host order by percentage(value)") + if err != nil { + t.Fatalf("Unable to parse query: %v", err) + } + + groupSet := NewGroupSet() + + setA := groupSet.GetSet("host-a") + if err := setA.Aggregate("percentage(value)", Percentage, "10", false); err != nil { + t.Fatalf("Unable to aggregate percentage for host-a: %v", err) + } + if err := setA.Aggregate("percentile(value)", Percentile, "10", false); err != nil { + t.Fatalf("Unable to aggregate percentile for host-a: %v", err) + } + + setB := groupSet.GetSet("host-b") + if err := setB.Aggregate("percentage(value)", Percentage, "30", false); err != nil { + t.Fatalf("Unable to aggregate percentage for host-b: %v", err) + } + if err := setB.Aggregate("percentile(value)", Percentile, "30", false); err != nil { + t.Fatalf("Unable to aggregate percentile for host-b: %v", err) + } + + setC := groupSet.GetSet("host-c") + if err := setC.Aggregate("percentage(value)", Percentage, "20", false); err != nil { + t.Fatalf("Unable to aggregate percentage for host-c: %v", err) + } + if err := setC.Aggregate("percentile(value)", Percentile, "20", false); err != nil { + t.Fatalf("Unable to aggregate percentile for host-c: %v", err) + } + + rows, _, err := groupSet.result(query, false) + if err != nil { + t.Fatalf("Unable to build result rows: %v", err) + } + if len(rows) != 3 { + t.Fatalf("Expected 3 result rows, got %d", len(rows)) + } + + if rows[0].groupKey != "host-b" { + t.Fatalf("Expected rows to be ordered by percentage descending, first row=%s", rows[0].groupKey) + } + + valuesByGroup := map[string][]float64{} + for _, row := range rows { + parsedValues := make([]float64, 0, len(row.values)) + for _, value := range row.values { + parsedValue, err := strconv.ParseFloat(value, 64) + if err != nil { + t.Fatalf("Unable to parse result value %q: %v", value, err) + } + parsedValues = append(parsedValues, parsedValue) + } + valuesByGroup[row.groupKey] = parsedValues + } + + assertAlmostEqual(t, valuesByGroup["host-a"][0], 16.6666666667, 0.0001, "host-a percentage") + assertAlmostEqual(t, valuesByGroup["host-a"][1], 33.3333333333, 0.0001, "host-a percentile") + assertAlmostEqual(t, valuesByGroup["host-b"][0], 50.0, 0.0001, "host-b percentage") + assertAlmostEqual(t, valuesByGroup["host-b"][1], 100.0, 0.0001, "host-b percentile") + assertAlmostEqual(t, valuesByGroup["host-c"][0], 33.3333333333, 0.0001, "host-c percentage") + assertAlmostEqual(t, valuesByGroup["host-c"][1], 66.6666666667, 0.0001, "host-c percentile") +} + +func TestGroupSetPercentageReturnsZeroWhenTotalIsZero(t *testing.T) { + query, err := NewQuery("select percentage(value) from stats group by host") + if err != nil { + t.Fatalf("Unable to parse query: %v", err) + } + + groupSet := NewGroupSet() + for _, host := range []string{"host-a", "host-b"} { + set := groupSet.GetSet(host) + if err := set.Aggregate("percentage(value)", Percentage, "0", false); err != nil { + t.Fatalf("Unable to aggregate percentage for %s: %v", host, err) + } + } + + rows, _, err := groupSet.result(query, false) + if err != nil { + t.Fatalf("Unable to build result rows: %v", err) + } + if len(rows) != 2 { + t.Fatalf("Expected 2 result rows, got %d", len(rows)) + } + for _, row := range rows { + if len(row.values) != 1 { + t.Fatalf("Expected one result value, got %d for %s", len(row.values), row.groupKey) + } + value, err := strconv.ParseFloat(row.values[0], 64) + if err != nil { + t.Fatalf("Unable to parse percentage result %q: %v", row.values[0], err) + } + assertAlmostEqual(t, value, 0.0, 0.0001, row.groupKey+" percentage") + } +} + +func TestPercentileRank(t *testing.T) { + sortedValues := []float64{10, 20, 30} + + tests := []struct { + name string + value float64 + expected float64 + }{ + {name: "below minimum", value: 5, expected: 0}, + {name: "first bucket", value: 10, expected: 33.3333333333}, + {name: "middle bucket", value: 20, expected: 66.6666666667}, + {name: "maximum", value: 30, expected: 100}, + {name: "above maximum", value: 40, expected: 100}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assertAlmostEqual(t, percentileRank(tt.value, sortedValues), tt.expected, 0.0001, tt.name) + }) + } + + assertAlmostEqual(t, percentileRank(10, []float64{10, 10, 30}), 66.6666666667, 0.0001, "duplicate percentile rank") + if got := percentileRank(10, nil); got != 0 { + t.Fatalf("Expected empty percentile input to return 0, got %f", got) + } +} + +func assertAlmostEqual(t *testing.T, got, expected, tolerance float64, label string) { + t.Helper() + + diff := got - expected + if diff < 0 { + diff = -diff + } + if diff > tolerance { + t.Fatalf("Unexpected %s: got=%f expected=%f tolerance=%f", label, got, expected, tolerance) + } +} diff --git a/internal/mapr/groupsetresult.go b/internal/mapr/groupsetresult.go index 47bdab8..c87c22f 100644 --- a/internal/mapr/groupsetresult.go +++ b/internal/mapr/groupsetresult.go @@ -6,15 +6,13 @@ import ( "os" "strings" - "github.com/mimecast/dtail/internal/color" - "github.com/mimecast/dtail/internal/config" "github.com/mimecast/dtail/internal/io/dlog" "github.com/mimecast/dtail/internal/io/pool" "github.com/mimecast/dtail/internal/protocol" ) // Result returns a nicely formated result of the query from the group set. -func (g *GroupSet) Result(query *Query, rowsLimit int) (string, int, error) { +func (g *GroupSet) Result(query *Query, rowsLimit int, renderer ResultRenderer) (string, int, error) { rows, columnWidths, err := g.result(query, true) if err != nil { return "", 0, err @@ -27,97 +25,69 @@ func (g *GroupSet) Result(query *Query, rowsLimit int) (string, int, error) { sb := pool.BuilderBuffer.Get().(*strings.Builder) defer pool.RecycleBuilderBuffer(sb) - g.resultWriteFormattedHeader(query, sb, lastColumn, rowsLimit, columnWidths) - g.resultWriteFormattedHeaderRowSeparator(query, sb, lastColumn, columnWidths) - g.resultWriteFormattedData(query, sb, lastColumn, rowsLimit, columnWidths, rows) + if renderer == nil { + renderer = PlainResultRenderer() + } + + g.resultWriteFormattedHeader(query, renderer, sb, lastColumn, rowsLimit, columnWidths) + g.resultWriteFormattedHeaderRowSeparator(query, renderer, sb, lastColumn, columnWidths) + g.resultWriteFormattedData(query, renderer, sb, lastColumn, rowsLimit, columnWidths, rows) return sb.String(), len(rows), nil } // Write a nicely formatted header for the result data. -func (g *GroupSet) resultWriteFormattedHeader(query *Query, sb *strings.Builder, +func (g *GroupSet) resultWriteFormattedHeader(query *Query, renderer ResultRenderer, sb *strings.Builder, lastColumn, rowsLimit int, columnWidths []int) { for i, sc := range query.Select { format := fmt.Sprintf(" %%%ds ", columnWidths[i]) str := fmt.Sprintf(format, sc.FieldStorage) - g.resultWriteFormattedHeaderEntry(query, sb, sc, str) + g.resultWriteFormattedHeaderEntry(query, renderer, sb, sc, str) if i == lastColumn { continue } - g.resultWriteFormattedHeaderEntrySeparator(query, sb) + g.resultWriteFormattedHeaderEntrySeparator(renderer, sb) } sb.WriteString("\n") } -func (g *GroupSet) resultWriteFormattedHeaderEntry(query *Query, sb *strings.Builder, +func (g *GroupSet) resultWriteFormattedHeaderEntry(query *Query, renderer ResultRenderer, sb *strings.Builder, sc selectCondition, str string) { - if config.Client.TermColorsEnable { - attrs := []color.Attribute{config.Client.TermColors.MaprTable.HeaderAttr} - if sc.FieldStorage == query.OrderBy { - attrs = append(attrs, config.Client.TermColors.MaprTable.HeaderSortKeyAttr) - } - for _, groupBy := range query.GroupBy { - if sc.FieldStorage == groupBy { - attrs = append(attrs, config.Client.TermColors.MaprTable.HeaderGroupKeyAttr) - break - } + isGroupKey := false + for _, groupBy := range query.GroupBy { + if sc.FieldStorage == groupBy { + isGroupKey = true + break } - color.PaintWithAttrs(sb, str, - config.Client.TermColors.MaprTable.HeaderFg, - config.Client.TermColors.MaprTable.HeaderBg, - attrs) - - } else { - sb.WriteString(str) } + renderer.WriteHeaderEntry(sb, str, sc.FieldStorage == query.OrderBy, isGroupKey) } -func (g *GroupSet) resultWriteFormattedHeaderEntrySeparator(query *Query, sb *strings.Builder) { - if config.Client.TermColorsEnable { - color.PaintWithAttr(sb, protocol.FieldDelimiter, - config.Client.TermColors.MaprTable.HeaderDelimiterFg, - config.Client.TermColors.MaprTable.HeaderDelimiterBg, - config.Client.TermColors.MaprTable.HeaderDelimiterAttr) - } else { - sb.WriteString(protocol.FieldDelimiter) - } +func (g *GroupSet) resultWriteFormattedHeaderEntrySeparator(renderer ResultRenderer, sb *strings.Builder) { + renderer.WriteHeaderDelimiter(sb, protocol.FieldDelimiter) } // This writes a nicely formatted line separating the header and the data. -func (g *GroupSet) resultWriteFormattedHeaderRowSeparator(query *Query, sb *strings.Builder, +func (g *GroupSet) resultWriteFormattedHeaderRowSeparator(query *Query, renderer ResultRenderer, sb *strings.Builder, lastColumn int, columnWidths []int) { for i := 0; i < len(query.Select); i++ { str := fmt.Sprintf("-%s-", strings.Repeat("-", columnWidths[i])) - if config.Client.TermColorsEnable { - color.PaintWithAttr(sb, str, - config.Client.TermColors.MaprTable.HeaderDelimiterFg, - config.Client.TermColors.MaprTable.HeaderDelimiterBg, - config.Client.TermColors.MaprTable.HeaderDelimiterAttr) - } else { - sb.WriteString(str) - } + renderer.WriteHeaderDelimiter(sb, str) if i == lastColumn { continue } - if config.Client.TermColorsEnable { - color.PaintWithAttr(sb, protocol.FieldDelimiter, - config.Client.TermColors.MaprTable.HeaderDelimiterFg, - config.Client.TermColors.MaprTable.HeaderDelimiterBg, - config.Client.TermColors.MaprTable.HeaderDelimiterAttr) - } else { - sb.WriteString(protocol.FieldDelimiter) - } + renderer.WriteHeaderDelimiter(sb, protocol.FieldDelimiter) } sb.WriteString("\n") } // Write the result data nicely formatted. -func (g *GroupSet) resultWriteFormattedData(query *Query, sb *strings.Builder, +func (g *GroupSet) resultWriteFormattedData(query *Query, renderer ResultRenderer, sb *strings.Builder, lastColumn, rowsLimit int, columnWidths []int, rows []result) { for i, r := range rows { @@ -125,37 +95,22 @@ func (g *GroupSet) resultWriteFormattedData(query *Query, sb *strings.Builder, break } for j, value := range r.values { - g.resultWriteFormattedDataEntry(query, sb, columnWidths, j, value) + g.resultWriteFormattedDataEntry(renderer, sb, columnWidths, j, value) if j == lastColumn { continue } - // Now, write the data entry separator. - if config.Client.TermColorsEnable { - color.PaintWithAttr(sb, protocol.FieldDelimiter, - config.Client.TermColors.MaprTable.DelimiterFg, - config.Client.TermColors.MaprTable.DelimiterBg, - config.Client.TermColors.MaprTable.DelimiterAttr) - } else { - sb.WriteString(protocol.FieldDelimiter) - } + renderer.WriteDataDelimiter(sb, protocol.FieldDelimiter) } sb.WriteString("\n") } } -func (g *GroupSet) resultWriteFormattedDataEntry(query *Query, sb *strings.Builder, +func (g *GroupSet) resultWriteFormattedDataEntry(renderer ResultRenderer, sb *strings.Builder, columnWidths []int, j int, value string) { format := fmt.Sprintf(" %%%ds ", columnWidths[j]) str := fmt.Sprintf(format, value) - if config.Client.TermColorsEnable { - color.PaintWithAttr(sb, str, - config.Client.TermColors.MaprTable.DataFg, - config.Client.TermColors.MaprTable.DataBg, - config.Client.TermColors.MaprTable.DataAttr) - } else { - sb.WriteString(str) - } + renderer.WriteDataEntry(sb, str) } func (*GroupSet) writeQueryFile(query *Query) error { @@ -248,12 +203,17 @@ func (g *GroupSet) resultWriteUnformatted(query *Query, rows []result, fd *os.Fi } } - if !query.Outfile.AppendMode && finalResult { + // Always rename .tmp to .csv after writing (not just on final result) + // This ensures the .csv file is updated at every interval + if !query.Outfile.AppendMode { tmpOutfile := fmt.Sprintf("%s.tmp", query.Outfile.FilePath) + dlog.Common.Debug("Renaming outfile", tmpOutfile, "to", query.Outfile.FilePath) if err := os.Rename(tmpOutfile, query.Outfile.FilePath); err != nil { + dlog.Common.Error("Failed to rename outfile", tmpOutfile, "error", err) os.Remove(tmpOutfile) return err } + dlog.Common.Info("Successfully renamed outfile to", query.Outfile.FilePath) } return nil diff --git a/internal/mapr/groupsetresult_renderer_test.go b/internal/mapr/groupsetresult_renderer_test.go new file mode 100644 index 0000000..53f45d5 --- /dev/null +++ b/internal/mapr/groupsetresult_renderer_test.go @@ -0,0 +1,113 @@ +package mapr + +import ( + "strings" + "testing" +) + +func TestGroupSetResultUsesProvidedRenderer(t *testing.T) { + query, err := NewQuery("select host,count(value) from stats group by host order by count(value)") + if err != nil { + t.Fatalf("Unable to parse query: %v", err) + } + + groupSet := NewGroupSet() + set := groupSet.GetSet("host-a") + if err := set.Aggregate("host", Last, "host-a", false); err != nil { + t.Fatalf("Unable to aggregate host field: %v", err) + } + if err := set.Aggregate("count(value)", Count, "", false); err != nil { + t.Fatalf("Unable to aggregate count field: %v", err) + } + + renderer := &recordingRenderer{} + result, numRows, err := groupSet.Result(query, 10, renderer) + if err != nil { + t.Fatalf("Unable to render result: %v", err) + } + if numRows != 1 { + t.Fatalf("Expected one row, got %d", numRows) + } + if len(renderer.headerCalls) != 2 { + t.Fatalf("Expected two header calls, got %d", len(renderer.headerCalls)) + } + if renderer.headerCalls[0].isSortKey || !renderer.headerCalls[0].isGroupKey { + t.Fatalf("Unexpected flags for group key header: %+v", renderer.headerCalls[0]) + } + if !renderer.headerCalls[1].isSortKey || renderer.headerCalls[1].isGroupKey { + t.Fatalf("Unexpected flags for sort key header: %+v", renderer.headerCalls[1]) + } + if len(renderer.headerDelimiters) == 0 { + t.Fatal("Expected header delimiters to be rendered") + } + if len(renderer.dataDelimiters) == 0 { + t.Fatal("Expected data delimiters to be rendered") + } + if !strings.Contains(result, "host-a") || !strings.Contains(result, "1") { + t.Fatalf("Expected rendered output to contain row data, got %q", result) + } +} + +func TestGroupSetResultFallsBackToPlainRenderer(t *testing.T) { + query, err := NewQuery("select count(value) from stats") + if err != nil { + t.Fatalf("Unable to parse query: %v", err) + } + + groupSet := NewGroupSet() + set := groupSet.GetSet("") + if err := set.Aggregate("count(value)", Count, "", false); err != nil { + t.Fatalf("Unable to aggregate count field: %v", err) + } + + result, numRows, err := groupSet.Result(query, 10, nil) + if err != nil { + t.Fatalf("Unable to render result with nil renderer: %v", err) + } + if numRows != 1 { + t.Fatalf("Expected one row, got %d", numRows) + } + if !strings.Contains(result, "count(value)") || !strings.Contains(result, "1") { + t.Fatalf("Expected plain rendered output, got %q", result) + } + if strings.Contains(result, "\x1b[") { + t.Fatalf("Expected plain output without ANSI escapes, got %q", result) + } +} + +type recordingRenderer struct { + headerCalls []headerCall + headerDelimiters []string + dataEntries []string + dataDelimiters []string +} + +type headerCall struct { + text string + isSortKey bool + isGroupKey bool +} + +func (r *recordingRenderer) WriteHeaderEntry(sb *strings.Builder, text string, isSortKey, isGroupKey bool) { + r.headerCalls = append(r.headerCalls, headerCall{ + text: text, + isSortKey: isSortKey, + isGroupKey: isGroupKey, + }) + sb.WriteString(text) +} + +func (r *recordingRenderer) WriteHeaderDelimiter(sb *strings.Builder, text string) { + r.headerDelimiters = append(r.headerDelimiters, text) + sb.WriteString(text) +} + +func (r *recordingRenderer) WriteDataEntry(sb *strings.Builder, text string) { + r.dataEntries = append(r.dataEntries, text) + sb.WriteString(text) +} + +func (r *recordingRenderer) WriteDataDelimiter(sb *strings.Builder, text string) { + r.dataDelimiters = append(r.dataDelimiters, text) + sb.WriteString(text) +} diff --git a/internal/mapr/logformat/csv.go b/internal/mapr/logformat/csv.go index ea85ca9..d82b238 100644 --- a/internal/mapr/logformat/csv.go +++ b/internal/mapr/logformat/csv.go @@ -2,52 +2,103 @@ package logformat import ( "fmt" - "strings" + "sync" "github.com/mimecast/dtail/internal/protocol" ) +// csvParser parses CSV log lines. The first line encountered for a given +// sourceID is treated as the column header and stored so that subsequent +// lines from the same source can be mapped to named fields. State is kept +// per sourceID because a single parser instance is shared across every +// file/stream processed within a mapreduce session; without this, the +// header row of every file after the first one would silently be mapped +// as a data row, corrupting aggregates. type csvParser struct { defaultParser - header []string - hasHeader bool + mu sync.RWMutex + headers map[string][]string } +var _ Parser = (*csvParser)(nil) + func newCSVParser(hostname, timeZoneName string, timeZoneOffset int) (*c