From a4e1abc1dfa314f81c47e7c112cae530e93b7838 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 13:19:52 +0200 Subject: [PATCH] fix(logger): address audit findings Redact and rate-limit error tracker fan-out, cap panic stack capture, sanitise stdlib fallback output, strip contexts in Info/Debug, sync the replaced logger, add Sync, UpdateLoggerE and CatchPanicRethrow, cache the PID, and add tests. --- pkg/logger/logger.go | 236 +++++++++++++++++++++++++++++--------- pkg/logger/logger_test.go | 129 +++++++++++++++++++++ 2 files changed, 311 insertions(+), 54 deletions(-) create mode 100644 pkg/logger/logger_test.go diff --git a/pkg/logger/logger.go b/pkg/logger/logger.go index 54f9a8a..0f0f8bc 100644 --- a/pkg/logger/logger.go +++ b/pkg/logger/logger.go @@ -5,8 +5,12 @@ import ( "fmt" "log" "os" + "regexp" + "runtime" "runtime/debug" + "strings" "sync" + "time" "go.uber.org/zap" @@ -28,10 +32,12 @@ func getLogger() *zap.SugaredLogger { return Logger } -func setLogger(l *zap.SugaredLogger) { +func swapLogger(l *zap.SugaredLogger) *zap.SugaredLogger { stateMu.Lock() defer stateMu.Unlock() + old := Logger Logger = l + return old } func getErrorTracker() errortracking.Provider { @@ -40,6 +46,103 @@ func getErrorTracker() errortracking.Provider { return errorTracker } +// pid is constant for the life of the process; cache it off the hot path. +var pid = os.Getpid() + +// maxStackBytes caps the stack trace captured for a recovered panic. +const maxStackBytes = 16 << 10 + +func captureStack() []byte { + buf := make([]byte, maxStackBytes) + return buf[:runtime.Stack(buf, false)] +} + +// Patterns scrubbed from messages before they leave the process for the +// error tracker. The local log is left untouched. +var redactPatterns = []struct { + re *regexp.Regexp + repl string +}{ + {regexp.MustCompile(`([a-zA-Z][a-zA-Z0-9+.\-]*://)[^\s/@:]+:[^\s/@]+@`), "${1}[REDACTED]@"}, + {regexp.MustCompile(`(?i)\b(password|passwd|pwd|secret|token|api[_-]?key|access[_-]?key)(\s*[=:]\s*)('[^']*'|"[^"]*"|[^\s,;&]+)`), "${1}${2}[REDACTED]"}, + {regexp.MustCompile(`(?i)\b(bearer|basic)\s+[A-Za-z0-9._~+/=\-]+`), "${1} [REDACTED]"}, +} + +func redact(s string) string { + for _, p := range redactPatterns { + s = p.re.ReplaceAllString(s, p.repl) + } + return s +} + +// sanitizeForStdlog escapes CR/LF and other control characters so untrusted +// values cannot forge additional log lines on the stdlib fallback path. +func sanitizeForStdlog(s string) string { + if !strings.ContainsFunc(s, func(r rune) bool { return r < 0x20 || r == 0x7f }) { + return s + } + var b strings.Builder + b.Grow(len(s) + 8) + for _, r := range s { + switch { + case r == '\n': + b.WriteString(`\n`) + case r == '\r': + b.WriteString(`\r`) + case r == '\t': + b.WriteString(`\t`) + case r < 0x20 || r == 0x7f: + fmt.Fprintf(&b, `\x%02x`, r) + default: + b.WriteRune(r) + } + } + return b.String() +} + +// Error tracker fan-out limiting: a global token bucket plus per-template +// dedup, so attacker-triggerable errors cannot burn quota or flood the queue. +const ( + trackerBurst = 50 + trackerRefillPerSec = 20.0 + trackerDedupWindow = time.Second + trackerMaxKeys = 1024 +) + +var limiter = struct { + sync.Mutex + tokens float64 + last time.Time + seen map[string]time.Time +}{tokens: trackerBurst, seen: map[string]time.Time{}} + +func allowTracker(key string) bool { + now := time.Now() + limiter.Lock() + defer limiter.Unlock() + + if !limiter.last.IsZero() { + limiter.tokens += now.Sub(limiter.last).Seconds() * trackerRefillPerSec + if limiter.tokens > trackerBurst { + limiter.tokens = trackerBurst + } + } + limiter.last = now + + if t, ok := limiter.seen[key]; ok && now.Sub(t) < trackerDedupWindow { + return false + } + if limiter.tokens < 1 { + return false + } + if len(limiter.seen) >= trackerMaxKeys { + limiter.seen = map[string]time.Time{} + } + limiter.seen[key] = now + limiter.tokens-- + return true +} + func Init(dev bool) { if dev { @@ -61,7 +164,16 @@ func UpdateLoggerPath(path string, dev bool) { UpdateLogger(&defaultConfig) } +// UpdateLogger rebuilds the logger from config. On failure the previous logger +// stays in place; use UpdateLoggerE to get the error. func UpdateLogger(config *zap.Config) { + if err := UpdateLoggerE(config); err != nil { + log.Printf("logger: failed to build logger, keeping previous: %s", sanitizeForStdlog(err.Error())) + } +} + +// UpdateLoggerE is UpdateLogger but returns the build error. +func UpdateLoggerE(config *zap.Config) error { defaultConfig := zap.NewProductionConfig() defaultConfig.OutputPaths = []string{"resolvespec.log"} if config == nil { @@ -70,12 +182,23 @@ func UpdateLogger(config *zap.Config) { logger, err := config.Build() if err != nil { - log.Print(err) - return + return err } - setLogger(logger.Sugar()) + old := swapLogger(logger.Sugar()) + if old != nil { + _ = old.Sync() + } Info("ResolveSpec Logger initialized") + return nil +} + +// Sync flushes buffered log entries. Call it on shutdown. +func Sync() error { + if lg := getLogger(); lg != nil { + return lg.Sync() + } + return nil } // InitErrorTracking initializes the error tracking provider @@ -125,59 +248,51 @@ func extractContext(args ...interface{}) (ctx context.Context, filteredArgs []in } func Info(template string, args ...interface{}) { - lg := getLogger() - if lg == nil { - log.Printf(template, args...) + _, args = extractContext(args...) + message := fmt.Sprintf(template, args...) + if lg := getLogger(); lg != nil { + lg.Infow(message, "process_id", pid) return } - lg.Infow(fmt.Sprintf(template, args...), "process_id", os.Getpid()) -} - -func Warn(template string, args ...interface{}) { - lg := getLogger() - tracker := getErrorTracker() - ctx, remainingArgs := extractContext(args...) - message := fmt.Sprintf(template, remainingArgs...) - if lg == nil { - log.Printf("%s", message) - } else { - lg.Warnw(message, "process_id", os.Getpid()) - } - - // Send to error tracker - if tracker != nil { - tracker.CaptureMessage(ctx, message, errortracking.SeverityWarning, map[string]interface{}{ - "process_id": os.Getpid(), - }) - } -} - -func Error(template string, args ...interface{}) { - lg := getLogger() - tracker := getErrorTracker() - ctx, remainingArgs := extractContext(args...) - message := fmt.Sprintf(template, remainingArgs...) - if lg == nil { - log.Printf("%s", message) - } else { - lg.Errorw(message, "process_id", os.Getpid()) - } - - // Send to error tracker - if tracker != nil { - tracker.CaptureMessage(ctx, message, errortracking.SeverityError, map[string]interface{}{ - "process_id": os.Getpid(), - }) - } + log.Printf("%s", sanitizeForStdlog(message)) } func Debug(template string, args ...interface{}) { - lg := getLogger() - if lg == nil { - log.Printf(template, args...) + _, args = extractContext(args...) + message := fmt.Sprintf(template, args...) + if lg := getLogger(); lg != nil { + lg.Debugw(message, "process_id", pid) return } - lg.Debugw(fmt.Sprintf(template, args...), "process_id", os.Getpid()) + log.Printf("%s", sanitizeForStdlog(message)) +} + +func Warn(template string, args ...interface{}) { + logAndTrack(errortracking.SeverityWarning, template, args) +} + +func Error(template string, args ...interface{}) { + logAndTrack(errortracking.SeverityError, template, args) +} + +func logAndTrack(sev errortracking.Severity, template string, args []interface{}) { + ctx, remainingArgs := extractContext(args...) + message := fmt.Sprintf(template, remainingArgs...) + if lg := getLogger(); lg == nil { + log.Printf("%s", sanitizeForStdlog(message)) + } else if sev == errortracking.SeverityWarning { + lg.Warnw(message, "process_id", pid) + } else { + lg.Errorw(message, "process_id", pid) + } + + tracker := getErrorTracker() + if tracker == nil || !allowTracker(string(sev)+"|"+template) { + return + } + tracker.CaptureMessage(ctx, redact(message), sev, map[string]interface{}{ + "process_id": pid, + }) } // CatchPanic - Handle panic @@ -187,7 +302,7 @@ func CatchPanicCallback(location string, cb func(err any), args ...interface{}) ctx, _ := extractContext(args...) return func() { if err := recover(); err != nil { - callstack := debug.Stack() + callstack := captureStack() lg := getLogger() tracker := getErrorTracker() @@ -202,7 +317,7 @@ func CatchPanicCallback(location string, cb func(err any), args ...interface{}) if tracker != nil { tracker.CapturePanic(ctx, err, callstack, map[string]interface{}{ "location": location, - "process_id": os.Getpid(), + "process_id": pid, }) } @@ -220,6 +335,19 @@ func CatchPanic(location string, args ...interface{}) func() { return CatchPanicCallback(location, nil, args...) } +// CatchPanicRethrow returns a function to defer that logs and reports a panic +// and then re-panics. Use it inside enforcement/internal code where swallowing +// the panic would let the caller proceed as if the work had succeeded. The +// swallowing CatchPanic is for outermost request/goroutine boundaries only. +func CatchPanicRethrow(location string, args ...interface{}) func() { + return func() { + if r := recover(); r != nil { + _ = HandlePanic(location, r, args...) + panic(r) + } + } +} + // HandlePanic logs a panic and returns it as an error // This should be called with the result of recover() from a deferred function // Example usage: @@ -232,14 +360,14 @@ func CatchPanic(location string, args ...interface{}) func() { func HandlePanic(methodName string, r any, args ...interface{}) error { tracker := getErrorTracker() ctx, _ := extractContext(args...) - stack := debug.Stack() + stack := captureStack() Error("Panic in %s: %v\nStack trace:\n%s", methodName, r, string(stack), ctx) // Pass context implicitly // Send to error tracker if tracker != nil { tracker.CapturePanic(ctx, r, stack, map[string]interface{}{ "method": methodName, - "process_id": os.Getpid(), + "process_id": pid, }) } diff --git a/pkg/logger/logger_test.go b/pkg/logger/logger_test.go new file mode 100644 index 0000000..b37d252 --- /dev/null +++ b/pkg/logger/logger_test.go @@ -0,0 +1,129 @@ +package logger + +import ( + "bytes" + "context" + "log" + "strings" + "sync" + "testing" + + "go.uber.org/zap" + + errortracking "github.com/bitechdev/ResolveSpec/pkg/errortracking" +) + +type fakeTracker struct { + mu sync.Mutex + msgs []string +} + +func (f *fakeTracker) CaptureError(context.Context, error, errortracking.Severity, map[string]interface{}) { +} +func (f *fakeTracker) CaptureMessage(_ context.Context, m string, _ errortracking.Severity, _ map[string]interface{}) { + f.mu.Lock() + f.msgs = append(f.msgs, m) + f.mu.Unlock() +} +func (f *fakeTracker) CapturePanic(context.Context, interface{}, []byte, map[string]interface{}) {} +func (f *fakeTracker) Flush(int) bool { return true } +func (f *fakeTracker) Close() error { return nil } + +func TestStdlibFallbackSafe(t *testing.T) { + swapLogger(nil) + var buf bytes.Buffer + log.SetOutput(&buf) + defer log.SetOutput(nil) + + Info("100%s done\nFAKE line", "") + Info("%d%% %s", 5, context.Background()) + out := buf.String() + if strings.Count(out, "\n") != 2 { + t.Fatalf("log injection: %q", out) + } +} + +func TestContextStrippedFromInfo(t *testing.T) { + swapLogger(nil) + var buf bytes.Buffer + log.SetOutput(&buf) + defer log.SetOutput(nil) + Info("saved %s", "widget", context.WithValue(context.Background(), "k", "secret")) + if strings.Contains(buf.String(), "EXTRA") || strings.Contains(buf.String(), "secret") { + t.Fatalf("context leaked: %q", buf.String()) + } +} + +func TestRedact(t *testing.T) { + in := "connect postgres://admin:hunter2@db:5432/x failed password=abc123 Authorization: Bearer eyJ.abc-d" + out := redact(in) + for _, leak := range []string{"hunter2", "abc123", "eyJ.abc-d"} { + if strings.Contains(out, leak) { + t.Errorf("%q leaked in %q", leak, out) + } + } +} + +func TestTrackerRateLimitAndRedaction(t *testing.T) { + swapLogger(zap.NewNop().Sugar()) + ft := &fakeTracker{} + InitErrorTracking(ft) + defer InitErrorTracking(nil) + + for i := 0; i < 100; i++ { + Error("same template %d password=x", i) + } + if len(ft.msgs) != 1 { + t.Fatalf("dedup failed: %d events", len(ft.msgs)) + } + if strings.Contains(ft.msgs[0], "password=x") { + t.Fatalf("not redacted: %q", ft.msgs[0]) + } +} + +func TestConcurrentUpdateAndLog(t *testing.T) { + InitErrorTracking(&fakeTracker{}) + defer InitErrorTracking(nil) + var wg sync.WaitGroup + for i := 0; i < 4; i++ { + wg.Add(2) + go func() { + defer wg.Done() + for j := 0; j < 50; j++ { + cfg := zap.NewProductionConfig() + cfg.OutputPaths = []string{"stderr"} + _ = UpdateLoggerE(&cfg) + } + }() + go func() { + defer wg.Done() + for j := 0; j < 200; j++ { + Error("e %d", j) + _ = CloseErrorTracking() + } + }() + } + wg.Wait() + _ = Sync() +} + +func TestCatchPanicRethrow(t *testing.T) { + swapLogger(zap.NewNop().Sugar()) + defer func() { + if recover() == nil { + t.Fatal("expected re-panic") + } + }() + func() { + defer CatchPanicRethrow("x")() + panic("boom") + }() +} + +func TestCatchPanicSwallows(t *testing.T) { + swapLogger(zap.NewNop().Sugar()) + func() { + defer CatchPanic("x")() + panic("boom") + }() +}