From a060a8ad8e0ddcc611744a2dd70b77ba3f29e193 Mon Sep 17 00:00:00 2001 From: Your Name Date: Tue, 29 Sep 2026 11:44:00 -0700 Subject: [PATCH 1/2] fix(audit): skip LogAuditEvent when Auditor is disabled Honor IsEnabled() in middleware so custom Auditor mocks are not invoked when disabled, matching the documented NewAuditMiddleware contract. Co-authored-by: Cursor --- pkg/audit/middleware.go | 6 +++-- pkg/audit/middleware_test.go | 45 ++++++++++++++++++++++++++++++++++++ 2 files changed, 49 insertions(+), 2 deletions(-) create mode 100644 pkg/audit/middleware_test.go diff --git a/pkg/audit/middleware.go b/pkg/audit/middleware.go index 25b3a07..bb04879 100644 --- a/pkg/audit/middleware.go +++ b/pkg/audit/middleware.go @@ -41,9 +41,11 @@ func (m *AuditMiddleware) Client() Auditor { } // LogAuditEvent sends an audit event to the audit service API -// This function is used to log audit events using the unified audit log structure +// This function is used to log audit events using the unified audit log structure. +// Skips when the client is nil or IsEnabled() returns false so custom Auditor +// implementations that do not no-op themselves are not invoked when disabled. func (m *AuditMiddleware) LogAuditEvent(ctx context.Context, auditRequest *AuditLogRequest) { - if m.client == nil { + if m.client == nil || !m.client.IsEnabled() { return } m.client.LogEvent(ctx, auditRequest) diff --git a/pkg/audit/middleware_test.go b/pkg/audit/middleware_test.go new file mode 100644 index 0000000..27b7ac0 --- /dev/null +++ b/pkg/audit/middleware_test.go @@ -0,0 +1,45 @@ +package audit + +import ( + "context" + "crypto" + "testing" +) + +// disabledAuditor is an Auditor whose IsEnabled returns false. +// LogEvent fails the test if called, asserting the middleware skips it. +type disabledAuditor struct { + t *testing.T +} + +func (d *disabledAuditor) LogEvent(ctx context.Context, event *AuditLogRequest) bool { + d.t.Error("LogEvent must not be called when IsEnabled returns false") + return false +} + +func (d *disabledAuditor) SignEvent(ctx context.Context, event *AuditLogRequest) error { + return nil +} + +func (d *disabledAuditor) SignMessageBytes(ctx context.Context, message []byte) (string, error) { + return "", nil +} + +func (d *disabledAuditor) LogSignedEvent(ctx context.Context, event *AuditLogRequest) {} + +func (d *disabledAuditor) VerifyIntegrity(event *AuditLogRequest, publicKey crypto.PublicKey) (bool, error) { + return false, nil +} + +func (d *disabledAuditor) IsEnabled() bool { return false } + +func (d *disabledAuditor) Close(ctx context.Context) error { return nil } + +func TestLogAuditEvent_SkipsWhenDisabled(t *testing.T) { + m := &AuditMiddleware{client: &disabledAuditor{t: t}} + m.LogAuditEvent(context.Background(), &AuditLogRequest{ + Action: "TEST", + ActorID: "actor-1", + ActorType: "USER", + }) +} From 85558a30c82138a251257a0a69406ba03cd522fa Mon Sep 17 00:00:00 2001 From: Your Name Date: Tue, 29 Sep 2026 11:53:15 -0700 Subject: [PATCH 2/2] test(audit): expand middleware coverage for IsEnabled contract Cover nil/disabled/enabled LogAuditEvent paths, global init once-semantics, and uninitialized warn behavior so custom Auditor mocks are protected. Co-authored-by: Cursor --- pkg/audit/middleware_test.go | 371 +++++++++++++++++++++++++++++++++-- 1 file changed, 352 insertions(+), 19 deletions(-) diff --git a/pkg/audit/middleware_test.go b/pkg/audit/middleware_test.go index 27b7ac0..2726f5d 100644 --- a/pkg/audit/middleware_test.go +++ b/pkg/audit/middleware_test.go @@ -1,45 +1,378 @@ package audit import ( + "bytes" "context" "crypto" + "fmt" + "log/slog" + "sync" "testing" ) -// disabledAuditor is an Auditor whose IsEnabled returns false. -// LogEvent fails the test if called, asserting the middleware skips it. -type disabledAuditor struct { - t *testing.T +// mockAuditor is a controllable Auditor for middleware tests. +// It records LogEvent calls and can be configured enabled/disabled. +type mockAuditor struct { + mu sync.Mutex + enabled bool + logEventCalls int + receivedEvents []*AuditLogRequest + receivedCtxs []context.Context + failOnLogEvent bool // if true, LogEvent fails the test when called + t *testing.T } -func (d *disabledAuditor) LogEvent(ctx context.Context, event *AuditLogRequest) bool { - d.t.Error("LogEvent must not be called when IsEnabled returns false") - return false +func (m *mockAuditor) LogEvent(ctx context.Context, event *AuditLogRequest) bool { + m.mu.Lock() + defer m.mu.Unlock() + if m.failOnLogEvent { + m.t.Error("LogEvent must not be called") + return false + } + m.logEventCalls++ + m.receivedCtxs = append(m.receivedCtxs, ctx) + if event != nil { + cp := *event + m.receivedEvents = append(m.receivedEvents, &cp) + } else { + m.receivedEvents = append(m.receivedEvents, nil) + } + return true } -func (d *disabledAuditor) SignEvent(ctx context.Context, event *AuditLogRequest) error { +func (m *mockAuditor) SignEvent(ctx context.Context, event *AuditLogRequest) error { return nil } -func (d *disabledAuditor) SignMessageBytes(ctx context.Context, message []byte) (string, error) { +func (m *mockAuditor) SignMessageBytes(ctx context.Context, message []byte) (string, error) { return "", nil } -func (d *disabledAuditor) LogSignedEvent(ctx context.Context, event *AuditLogRequest) {} +func (m *mockAuditor) LogSignedEvent(ctx context.Context, event *AuditLogRequest) {} -func (d *disabledAuditor) VerifyIntegrity(event *AuditLogRequest, publicKey crypto.PublicKey) (bool, error) { +func (m *mockAuditor) VerifyIntegrity(event *AuditLogRequest, publicKey crypto.PublicKey) (bool, error) { return false, nil } -func (d *disabledAuditor) IsEnabled() bool { return false } +func (m *mockAuditor) IsEnabled() bool { + m.mu.Lock() + defer m.mu.Unlock() + return m.enabled +} + +func (m *mockAuditor) Close(ctx context.Context) error { return nil } + +func (m *mockAuditor) callCount() int { + m.mu.Lock() + defer m.mu.Unlock() + return m.logEventCalls +} + +func (m *mockAuditor) events() []*AuditLogRequest { + m.mu.Lock() + defer m.mu.Unlock() + out := make([]*AuditLogRequest, len(m.receivedEvents)) + copy(out, m.receivedEvents) + return out +} + +func sampleRequest() *AuditLogRequest { + trace := "trace-123" + target := "target-1" + return &AuditLogRequest{ + TraceID: &trace, + Timestamp: "2023-01-01T00:00:00Z", + EventType: "USER_MANAGEMENT", + Action: "CREATE", + Status: StatusSuccess, + ActorType: "USER", + ActorID: "actor-1", + TargetType: "RESOURCE", + TargetID: &target, + Message: []byte(`{"k":"v"}`), + Metadata: map[string]interface{}{"source": "test"}, + } +} + +func TestAuditMiddleware_LogAuditEvent(t *testing.T) { + t.Run("skips when client is nil", func(t *testing.T) { + m := &AuditMiddleware{client: nil} + // Must not panic. + m.LogAuditEvent(context.Background(), sampleRequest()) + }) + + t.Run("skips when IsEnabled returns false", func(t *testing.T) { + client := &mockAuditor{t: t, enabled: false, failOnLogEvent: true} + m := &AuditMiddleware{client: client} + m.LogAuditEvent(context.Background(), sampleRequest()) + if got := client.callCount(); got != 0 { + t.Fatalf("LogEvent call count = %d, want 0", got) + } + if len(client.events()) != 0 { + t.Fatalf("receivedEvents len = %d, want 0", len(client.events())) + } + }) + + t.Run("calls LogEvent when enabled", func(t *testing.T) { + client := &mockAuditor{enabled: true} + m := &AuditMiddleware{client: client} + req := sampleRequest() + ctx := context.WithValue(context.Background(), struct{}{}, "marker") + + m.LogAuditEvent(ctx, req) + + if got := client.callCount(); got != 1 { + t.Fatalf("LogEvent call count = %d, want 1", got) + } + events := client.events() + if len(events) != 1 { + t.Fatalf("receivedEvents len = %d, want 1", len(events)) + } + got := events[0] + if got == nil { + t.Fatal("received event is nil") + } + if got.Action != req.Action || got.ActorID != req.ActorID || got.ActorType != req.ActorType { + t.Errorf("event mismatch: got action=%q actor=%q/%q, want action=%q actor=%q/%q", + got.Action, got.ActorType, got.ActorID, req.Action, req.ActorType, req.ActorID) + } + if got.Status != req.Status || got.EventType != req.EventType { + t.Errorf("event classification mismatch: got type=%q status=%q, want type=%q status=%q", + got.EventType, got.Status, req.EventType, req.Status) + } + if string(got.Message) != string(req.Message) { + t.Errorf("message = %q, want %q", got.Message, req.Message) + } + client.mu.Lock() + receivedCtx := client.receivedCtxs[0] + client.mu.Unlock() + if receivedCtx != ctx { + t.Error("LogEvent was not passed the same context") + } + }) + + t.Run("forwards nil request when enabled", func(t *testing.T) { + client := &mockAuditor{enabled: true} + m := &AuditMiddleware{client: client} + m.LogAuditEvent(context.Background(), nil) + if got := client.callCount(); got != 1 { + t.Fatalf("LogEvent call count = %d, want 1", got) + } + events := client.events() + if len(events) != 1 || events[0] != nil { + t.Fatalf("expected a single nil event, got %#v", events) + } + }) + + t.Run("does not call LogEvent after toggling to disabled", func(t *testing.T) { + client := &mockAuditor{enabled: true} + m := &AuditMiddleware{client: client} + + m.LogAuditEvent(context.Background(), sampleRequest()) + if got := client.callCount(); got != 1 { + t.Fatalf("after enabled call: count = %d, want 1", got) + } + + client.mu.Lock() + client.enabled = false + client.failOnLogEvent = true + client.t = t + client.mu.Unlock() + + m.LogAuditEvent(context.Background(), sampleRequest()) + if got := client.callCount(); got != 1 { + t.Fatalf("after disable: count = %d, want 1 (no additional calls)", got) + } + }) + + t.Run("multiple events when enabled", func(t *testing.T) { + client := &mockAuditor{enabled: true} + m := &AuditMiddleware{client: client} + + for i := 0; i < 3; i++ { + req := sampleRequest() + req.ActorID = fmt.Sprintf("actor-%d", i+1) + m.LogAuditEvent(context.Background(), req) + } + events := client.events() + if got := client.callCount(); got != 3 { + t.Fatalf("LogEvent call count = %d, want 3", got) + } + if len(events) != 3 { + t.Fatalf("receivedEvents len = %d, want 3", len(events)) + } + for i, ev := range events { + want := fmt.Sprintf("actor-%d", i+1) + if ev.ActorID != want { + t.Errorf("event[%d].ActorID = %q, want %q", i, ev.ActorID, want) + } + } + }) +} + +func TestNewAuditMiddleware(t *testing.T) { + ResetGlobalAuditMiddleware() + t.Cleanup(ResetGlobalAuditMiddleware) + + t.Run("returns middleware with client", func(t *testing.T) { + ResetGlobalAuditMiddleware() + client := &mockAuditor{enabled: true} + m := NewAuditMiddleware(client) + if m == nil { + t.Fatal("NewAuditMiddleware returned nil") + } + if m.Client() != client { + t.Error("Client() did not return the provided auditor") + } + }) + + t.Run("sets global instance on first call only", func(t *testing.T) { + ResetGlobalAuditMiddleware() + first := &mockAuditor{enabled: true} + second := &mockAuditor{enabled: false, failOnLogEvent: true, t: t} -func (d *disabledAuditor) Close(ctx context.Context) error { return nil } + m1 := NewAuditMiddleware(first) + m2 := NewAuditMiddleware(second) -func TestLogAuditEvent_SkipsWhenDisabled(t *testing.T) { - m := &AuditMiddleware{client: &disabledAuditor{t: t}} - m.LogAuditEvent(context.Background(), &AuditLogRequest{ - Action: "TEST", - ActorID: "actor-1", - ActorType: "USER", + if m1.Client() != first { + t.Error("first middleware should hold first client") + } + if m2.Client() != second { + t.Error("second middleware instance should hold second client") + } + + global := GetGlobalAuditMiddleware() + if global == nil { + t.Fatal("global middleware is nil") + } + if global.Client() != first { + t.Error("global middleware should remain the first-initialized instance") + } + + // Global path should still invoke the first (enabled) client. + LogAuditEvent(context.Background(), sampleRequest()) + if got := first.callCount(); got != 1 { + t.Fatalf("first client LogEvent count = %d, want 1 via global", got) + } + if got := second.callCount(); got != 0 { + t.Fatalf("second client LogEvent count = %d, want 0", got) + } + }) + + t.Run("nil client middleware skips logging", func(t *testing.T) { + ResetGlobalAuditMiddleware() + m := NewAuditMiddleware(nil) + m.LogAuditEvent(context.Background(), sampleRequest()) + LogAuditEvent(context.Background(), sampleRequest()) // via global; also nil client + }) +} + +func TestInitializeGlobalAudit(t *testing.T) { + ResetGlobalAuditMiddleware() + t.Cleanup(ResetGlobalAuditMiddleware) + + t.Run("disabled client skips via global LogAuditEvent", func(t *testing.T) { + ResetGlobalAuditMiddleware() + client := &mockAuditor{enabled: false, failOnLogEvent: true, t: t} + InitializeGlobalAudit(client) + + LogAuditEvent(context.Background(), sampleRequest()) + + if got := client.callCount(); got != 0 { + t.Fatalf("LogEvent call count = %d, want 0", got) + } + global := GetGlobalAuditMiddleware() + if global == nil || global.Client() != client { + t.Fatal("global middleware not initialized with disabled client") + } + }) + + t.Run("enabled client receives events via global LogAuditEvent", func(t *testing.T) { + ResetGlobalAuditMiddleware() + client := &mockAuditor{enabled: true} + InitializeGlobalAudit(client) + + req := sampleRequest() + LogAuditEvent(context.Background(), req) + + if got := client.callCount(); got != 1 { + t.Fatalf("LogEvent call count = %d, want 1", got) + } + events := client.events() + if len(events) != 1 || events[0].ActorID != req.ActorID { + t.Fatalf("unexpected events: %#v", events) + } }) + + t.Run("subsequent InitializeGlobalAudit is ignored", func(t *testing.T) { + ResetGlobalAuditMiddleware() + first := &mockAuditor{enabled: true} + second := &mockAuditor{enabled: true} + InitializeGlobalAudit(first) + InitializeGlobalAudit(second) + + LogAuditEvent(context.Background(), sampleRequest()) + if first.callCount() != 1 { + t.Fatalf("first client count = %d, want 1", first.callCount()) + } + if second.callCount() != 0 { + t.Fatalf("second client count = %d, want 0 (initialize ignored)", second.callCount()) + } + }) +} + +func TestLogAuditEvent_GlobalUninitialized(t *testing.T) { + ResetGlobalAuditMiddleware() + t.Cleanup(ResetGlobalAuditMiddleware) + + var buf bytes.Buffer + prev := slog.Default() + slog.SetDefault(slog.New(slog.NewTextHandler(&buf, &slog.HandlerOptions{Level: slog.LevelWarn}))) + t.Cleanup(func() { slog.SetDefault(prev) }) + + // Must not panic; should warn that global middleware is missing. + LogAuditEvent(context.Background(), sampleRequest()) + + if !bytes.Contains(buf.Bytes(), []byte("Global AuditMiddleware is not initialized")) { + t.Fatalf("expected warn log about uninitialized middleware, got: %s", buf.String()) + } +} + +func TestGetGlobalAuditMiddleware(t *testing.T) { + ResetGlobalAuditMiddleware() + t.Cleanup(ResetGlobalAuditMiddleware) + + if GetGlobalAuditMiddleware() != nil { + t.Fatal("expected nil global before init") + } + + client := &mockAuditor{enabled: true} + InitializeGlobalAudit(client) + got := GetGlobalAuditMiddleware() + if got == nil || got.Client() != client { + t.Fatal("GetGlobalAuditMiddleware did not return initialized instance") + } +} + +func TestResetGlobalAuditMiddleware(t *testing.T) { + ResetGlobalAuditMiddleware() + t.Cleanup(ResetGlobalAuditMiddleware) + + InitializeGlobalAudit(&mockAuditor{enabled: true}) + if GetGlobalAuditMiddleware() == nil { + t.Fatal("expected global after init") + } + + ResetGlobalAuditMiddleware() + if GetGlobalAuditMiddleware() != nil { + t.Fatal("expected nil global after reset") + } + + // After reset, a new InitializeGlobalAudit should take effect. + client := &mockAuditor{enabled: true} + InitializeGlobalAudit(client) + LogAuditEvent(context.Background(), sampleRequest()) + if client.callCount() != 1 { + t.Fatalf("after reset+reinit: count = %d, want 1", client.callCount()) + } }