From c7b4530689dc559341f8be7e5869029cbe1a4263 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 13:07:47 +0200 Subject: [PATCH] fix(race): make `make test-race` pass across ./pkg/... Add a test-race target and run test-unit over ./pkg/... Production races: - logger: guard Logger/errorTracker with an RWMutex - security: copy UserContext for the async session-activity goroutine and track it with a WaitGroup Test fixes: - eventbroker, websocketspec: use atomics for state shared with workers - security: wait for async activity updates before touching sqlmock - mqttspec: build full HookContext, set SubscriptionID, pin the in-memory SQLite to one connection Update the cross-cutting audit (X1) with status and findings. --- Makefile | 16 ++- audit/pkg/_CROSS-CUTTING.audit.md | 23 +++++ pkg/eventbroker/broker_test.go | 16 +-- pkg/logger/logger.go | 82 +++++++++++----- pkg/mqttspec/handler_test.go | 158 ++++++++++++++++++------------ pkg/security/providers.go | 10 +- pkg/security/providers_test.go | 44 +++++---- pkg/websocketspec/hooks_test.go | 7 +- 8 files changed, 235 insertions(+), 121 deletions(-) diff --git a/Makefile b/Makefile index eefef3f..311b5c4 100644 --- a/Makefile +++ b/Makefile @@ -1,11 +1,18 @@ -.PHONY: test test-unit test-integration docker-up docker-down clean +.PHONY: test test-unit test-race test-integration docker-up docker-down clean GOLANGCI_LINT := $(shell go env GOPATH)/bin/golangci-lint # Run all unit tests test-unit: @echo "Running unit tests..." - @go test ./pkg/resolvespec ./pkg/restheadspec -v -cover + @go test ./pkg/... -v -cover + +# Run all unit tests under the race detector (kept separate from coverage: +# race builds are 2-10x slower). Only races on executed paths are reported, +# so this covers every package rather than a subset. +test-race: + @echo "Running unit tests with the race detector..." + @go test -race -count=1 ./pkg/... # Run all integration tests (requires PostgreSQL) test-integration: @@ -13,7 +20,7 @@ test-integration: @go test -tags=integration ./pkg/resolvespec ./pkg/restheadspec -v # Run all tests (unit + integration) -test: test-unit test-integration +test: test-unit test-race test-integration release-version: ## Create and push a release with specific version (use: make release-version VERSION=v1.2.3 or make release-version to auto-increment) @if [ -z "$(VERSION)" ]; then \ @@ -113,7 +120,8 @@ coverage-integration: help: @echo "Available targets:" - @echo " test-unit - Run unit tests" + @echo " test-unit - Run unit tests for all packages (./pkg/...)" + @echo " test-race - Run unit tests for all packages with -race" @echo " test-integration - Run integration tests (requires PostgreSQL)" @echo " test - Run all tests" @echo " docker-up - Start PostgreSQL container" diff --git a/audit/pkg/_CROSS-CUTTING.audit.md b/audit/pkg/_CROSS-CUTTING.audit.md index b63b491..911bffa 100644 --- a/audit/pkg/_CROSS-CUTTING.audit.md +++ b/audit/pkg/_CROSS-CUTTING.audit.md @@ -105,6 +105,29 @@ that the race detector only reports races that **actually execute**, so X1 and X have to be fixed together: a race detector pointed at packages with no tests finds nothing. +**Status (2026-09-30) — partially resolved.** `make test-race` now exists +(`go test -race -count=1 ./pkg/...`), `test-unit` covers `./pkg/...`, and `test` +depends on both. The CI workflow (`.github/workflows/tests.yml`) still has no +race job, so nothing enforces it yet. The first full run was not clean: + +| Package | Race | Kind | +|---|---|---| +| `pkg/logger` | `Logger` / `errorTracker` reassigned while other goroutines log (hit via `pkg/server` tests) | **production** — now guarded by an `RWMutex` (`getLogger`, `setLogger`, `getErrorTracker`); the exported `Logger` var is kept for compatibility | +| `pkg/security` | `DatabaseAuthenticator.Authenticate` passed `&userCtx` to the async session-activity goroutine while also returning it to the caller | **production** — the goroutine now gets a copy and is tracked by a `WaitGroup` so tests can wait for it | +| `pkg/security` tests | async activity update used sqlmock concurrently with the test adding expectations | test — tests wait via `authenticateSync` | +| `pkg/eventbroker`, `pkg/websocketspec` tests | handler/hook closures mutated a plain `bool`/`int` from worker goroutines | test — now `atomic` | +| `pkg/mqttspec` tests | not a race: hand-built `HookContext` lacked `TableName`/`Model`/`ModelPtr`, the unsubscribe test set `Data` instead of `SubscriptionID`, and `:memory:` SQLite gave each pooled connection its own empty database | test — fixed; these were failing without `-race` too | + +`pkg/cache`, `pkg/config`, `pkg/modelregistry`, `pkg/tracing` and +`pkg/errortracking` are listed above but did **not** trip the detector: their +racing paths are not exercised by the current tests, which is the point made in +the paragraph above about X1 and X2 needing to be fixed together. Adding +concurrent tests for those globals is still outstanding. + +Known limitation: `pkg/security` tests are not repeatable with `-count>1` (a +package-level capability cache carries over between runs), so the race target +keeps `-count=1`. + --- ### X2. High — `go test` runs against 2 of 23 packages diff --git a/pkg/eventbroker/broker_test.go b/pkg/eventbroker/broker_test.go index 9217ac8..c764380 100644 --- a/pkg/eventbroker/broker_test.go +++ b/pkg/eventbroker/broker_test.go @@ -171,7 +171,7 @@ func TestBrokerPublishAsync(t *testing.T) { // Publish multiple events for i := 0; i < 5; i++ { event := NewEvent(EventSourceSystem, "test.event") - event.InstanceID = "test-instance" + event.InstanceID = "test-instance" if err := broker.PublishAsync(context.Background(), event); err != nil { t.Fatalf("PublishAsync failed: %v", err) } @@ -346,7 +346,7 @@ func TestBrokerStats(t *testing.T) { // Publish events for i := 0; i < 3; i++ { event := NewEvent(EventSourceSystem, "test.event") - event.InstanceID = "test-instance" + event.InstanceID = "test-instance" broker.PublishSync(context.Background(), event) } @@ -413,7 +413,7 @@ func TestBrokerConcurrentPublish(t *testing.T) { go func() { defer wg.Done() event := NewEvent(EventSourceSystem, "test.event") - event.InstanceID = "test-instance" + event.InstanceID = "test-instance" broker.PublishAsync(context.Background(), event) }() } @@ -450,7 +450,7 @@ func TestBrokerGracefulShutdown(t *testing.T) { // Publish events for i := 0; i < 5; i++ { event := NewEvent(EventSourceSystem, "test.event") - event.InstanceID = "test-instance" + event.InstanceID = "test-instance" broker.PublishAsync(context.Background(), event) } @@ -502,21 +502,21 @@ func TestBrokerProcessingModes(t *testing.T) { broker.Start(context.Background()) defer broker.Stop(context.Background()) - called := false + var called atomic.Bool broker.Subscribe("test.*", EventHandlerFunc(func(ctx context.Context, event *Event) error { - called = true + called.Store(true) return nil })) event := NewEvent(EventSourceSystem, "test.event") - event.InstanceID = "test-instance" + event.InstanceID = "test-instance" broker.Publish(context.Background(), event) if tt.mode == ProcessingModeAsync { time.Sleep(50 * time.Millisecond) } - if !called { + if !called.Load() { t.Error("Expected handler to be called") } }) diff --git a/pkg/logger/logger.go b/pkg/logger/logger.go index b58082f..54f9a8a 100644 --- a/pkg/logger/logger.go +++ b/pkg/logger/logger.go @@ -6,15 +6,40 @@ import ( "log" "os" "runtime/debug" + "sync" "go.uber.org/zap" errortracking "github.com/bitechdev/ResolveSpec/pkg/errortracking" ) +// Logger is the active logger. It is kept exported for compatibility, but +// inside this package it must only be accessed through getLogger/setLogger. var Logger *zap.SugaredLogger var errorTracker errortracking.Provider +// stateMu guards Logger and errorTracker, which may be replaced while other +// goroutines are logging. +var stateMu sync.RWMutex + +func getLogger() *zap.SugaredLogger { + stateMu.RLock() + defer stateMu.RUnlock() + return Logger +} + +func setLogger(l *zap.SugaredLogger) { + stateMu.Lock() + defer stateMu.Unlock() + Logger = l +} + +func getErrorTracker() errortracking.Provider { + stateMu.RLock() + defer stateMu.RUnlock() + return errorTracker +} + func Init(dev bool) { if dev { @@ -49,28 +74,30 @@ func UpdateLogger(config *zap.Config) { return } - Logger = logger.Sugar() + setLogger(logger.Sugar()) Info("ResolveSpec Logger initialized") } // InitErrorTracking initializes the error tracking provider func InitErrorTracking(provider errortracking.Provider) { + stateMu.Lock() errorTracker = provider - if errorTracker != nil { + stateMu.Unlock() + if provider != nil { Info("Error tracking initialized") } } // GetErrorTracker returns the current error tracking provider func GetErrorTracker() errortracking.Provider { - return errorTracker + return getErrorTracker() } // CloseErrorTracking flushes and closes the error tracking provider func CloseErrorTracking() error { - if errorTracker != nil { - errorTracker.Flush(5) - return errorTracker.Close() + if tracker := getErrorTracker(); tracker != nil { + tracker.Flush(5) + return tracker.Close() } return nil } @@ -98,53 +125,59 @@ func extractContext(args ...interface{}) (ctx context.Context, filteredArgs []in } func Info(template string, args ...interface{}) { - if Logger == nil { + lg := getLogger() + if lg == nil { log.Printf(template, args...) return } - Logger.Infow(fmt.Sprintf(template, args...), "process_id", os.Getpid()) + 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 Logger == nil { + if lg == nil { log.Printf("%s", message) } else { - Logger.Warnw(message, "process_id", os.Getpid()) + lg.Warnw(message, "process_id", os.Getpid()) } // Send to error tracker - if errorTracker != nil { - errorTracker.CaptureMessage(ctx, message, errortracking.SeverityWarning, map[string]interface{}{ + 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 Logger == nil { + if lg == nil { log.Printf("%s", message) } else { - Logger.Errorw(message, "process_id", os.Getpid()) + lg.Errorw(message, "process_id", os.Getpid()) } // Send to error tracker - if errorTracker != nil { - errorTracker.CaptureMessage(ctx, message, errortracking.SeverityError, map[string]interface{}{ + if tracker != nil { + tracker.CaptureMessage(ctx, message, errortracking.SeverityError, map[string]interface{}{ "process_id": os.Getpid(), }) } } func Debug(template string, args ...interface{}) { - if Logger == nil { + lg := getLogger() + if lg == nil { log.Printf(template, args...) return } - Logger.Debugw(fmt.Sprintf(template, args...), "process_id", os.Getpid()) + lg.Debugw(fmt.Sprintf(template, args...), "process_id", os.Getpid()) } // CatchPanic - Handle panic @@ -155,8 +188,10 @@ func CatchPanicCallback(location string, cb func(err any), args ...interface{}) return func() { if err := recover(); err != nil { callstack := debug.Stack() + lg := getLogger() + tracker := getErrorTracker() - if Logger != nil { + if lg != nil { Error("Panic in %s : %v", location, err, ctx) // Pass context implicitly } else { fmt.Printf("%s:PANIC->%+v", location, err) @@ -164,8 +199,8 @@ func CatchPanicCallback(location string, cb func(err any), args ...interface{}) } // Send to error tracker - if errorTracker != nil { - errorTracker.CapturePanic(ctx, err, callstack, map[string]interface{}{ + if tracker != nil { + tracker.CapturePanic(ctx, err, callstack, map[string]interface{}{ "location": location, "process_id": os.Getpid(), }) @@ -195,13 +230,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() Error("Panic in %s: %v\nStack trace:\n%s", methodName, r, string(stack), ctx) // Pass context implicitly // Send to error tracker - if errorTracker != nil { - errorTracker.CapturePanic(ctx, r, stack, map[string]interface{}{ + if tracker != nil { + tracker.CapturePanic(ctx, r, stack, map[string]interface{}{ "method": methodName, "process_id": os.Getpid(), }) diff --git a/pkg/mqttspec/handler_test.go b/pkg/mqttspec/handler_test.go index 49966e6..407271a 100644 --- a/pkg/mqttspec/handler_test.go +++ b/pkg/mqttspec/handler_test.go @@ -41,6 +41,12 @@ func setupTestHandler(t *testing.T) (*Handler, *gorm.DB) { db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) require.NoError(t, err) + // Each connection to ":memory:" gets its own database; pin to one so + // concurrent requests all see the migrated schema. + sqlDB, err := db.DB() + require.NoError(t, err) + sqlDB.SetMaxOpenConns(1) + // Auto-migrate test model err = db.AutoMigrate(&TestUser{}) require.NoError(t, err) @@ -93,9 +99,9 @@ func TestHandler_HandleRead_Single(t *testing.T) { // Insert test data user := &TestUser{ - ID: 1, - Name: "John Doe", - Email: "john@example.com", + ID: 1, + Name: "John Doe", + Email: "john@example.com", Status: "active", } db.Create(user) @@ -115,13 +121,16 @@ func TestHandler_HandleRead_Single(t *testing.T) { // Create hook context hookCtx := &HookContext{ - Context: context.Background(), - Handler: nil, - Schema: "public", - Entity: "users", - ID: "1", - Options: msg.Options, - Metadata: map[string]interface{}{"mqtt_client": client}, + Context: context.Background(), + TableName: "users", + Model: &TestUser{}, + ModelPtr: &TestUser{}, + Handler: nil, + Schema: "public", + Entity: "users", + ID: "1", + Options: msg.Options, + Metadata: map[string]interface{}{"mqtt_client": client}, } // Handle read @@ -164,12 +173,15 @@ func TestHandler_HandleRead_Multiple(t *testing.T) { // Create hook context hookCtx := &HookContext{ - Context: context.Background(), - Handler: nil, - Schema: "public", - Entity: "users", - Options: msg.Options, - Metadata: map[string]interface{}{"mqtt_client": client}, + Context: context.Background(), + TableName: "users", + Model: &TestUser{}, + ModelPtr: &TestUser{}, + Handler: nil, + Schema: "public", + Entity: "users", + Options: msg.Options, + Metadata: map[string]interface{}{"mqtt_client": client}, } // Handle read @@ -208,13 +220,16 @@ func TestHandler_HandleCreate(t *testing.T) { // Create hook context hookCtx := &HookContext{ - Context: context.Background(), - Handler: nil, - Schema: "public", - Entity: "users", - Data: newUser, - Options: msg.Options, - Metadata: map[string]interface{}{"mqtt_client": client}, + Context: context.Background(), + TableName: "users", + Model: &TestUser{}, + ModelPtr: &TestUser{}, + Handler: nil, + Schema: "public", + Entity: "users", + Data: newUser, + Options: msg.Options, + Metadata: map[string]interface{}{"mqtt_client": client}, } // Handle create @@ -266,14 +281,17 @@ func TestHandler_HandleUpdate(t *testing.T) { // Create hook context hookCtx := &HookContext{ - Context: context.Background(), - Handler: nil, - Schema: "public", - Entity: "users", - ID: "1", - Data: updateData, - Options: msg.Options, - Metadata: map[string]interface{}{"mqtt_client": client}, + Context: context.Background(), + TableName: "users", + Model: &TestUser{}, + ModelPtr: &TestUser{}, + Handler: nil, + Schema: "public", + Entity: "users", + ID: "1", + Data: updateData, + Options: msg.Options, + Metadata: map[string]interface{}{"mqtt_client": client}, } // Handle update @@ -318,13 +336,16 @@ func TestHandler_HandleDelete(t *testing.T) { // Create hook context hookCtx := &HookContext{ - Context: context.Background(), - Handler: nil, - Schema: "public", - Entity: "users", - ID: "1", - Options: msg.Options, - Metadata: map[string]interface{}{"mqtt_client": client}, + Context: context.Background(), + TableName: "users", + Model: &TestUser{}, + ModelPtr: &TestUser{}, + Handler: nil, + Schema: "public", + Entity: "users", + ID: "1", + Options: msg.Options, + Metadata: map[string]interface{}{"mqtt_client": client}, } // Handle delete @@ -388,13 +409,13 @@ func TestHandler_HandleUnsubscribe(t *testing.T) { sub := handler.subscriptionManager.Subscribe("sub-1", client.ID, "public", "users", &common.RequestOptions{}) client.AddSubscription(sub) - // Create unsubscribe message with subscription ID in Data + // Create unsubscribe message with the subscription ID msg := &Message{ - ID: "msg-7", - Type: MessageTypeSubscription, - Operation: OperationUnsubscribe, - Data: map[string]interface{}{"subscription_id": "sub-1"}, - Options: &common.RequestOptions{}, + ID: "msg-7", + Type: MessageTypeSubscription, + Operation: OperationUnsubscribe, + SubscriptionID: "sub-1", + Options: &common.RequestOptions{}, } // Handle unsubscribe @@ -490,12 +511,15 @@ func TestHandler_Hooks_BeforeRead(t *testing.T) { // Create hook context hookCtx := &HookContext{ - Context: context.Background(), - Handler: nil, - Schema: "public", - Entity: "users", - Options: msg.Options, - Metadata: map[string]interface{}{"mqtt_client": client}, + Context: context.Background(), + TableName: "users", + Model: &TestUser{}, + ModelPtr: &TestUser{}, + Handler: nil, + Schema: "public", + Entity: "users", + Options: msg.Options, + Metadata: map[string]interface{}{"mqtt_client": client}, } // Handle read @@ -544,13 +568,16 @@ func TestHandler_Hooks_BeforeCreate(t *testing.T) { } hookCtx := &HookContext{ - Context: context.Background(), - Handler: nil, - Schema: "public", - Entity: "users", - Data: newUser, - Options: msg.Options, - Metadata: map[string]interface{}{"mqtt_client": client}, + Context: context.Background(), + TableName: "users", + Model: &TestUser{}, + ModelPtr: &TestUser{}, + Handler: nil, + Schema: "public", + Entity: "users", + Data: newUser, + Options: msg.Options, + Metadata: map[string]interface{}{"mqtt_client": client}, } // Handle create @@ -600,13 +627,16 @@ func TestHandler_ConcurrentRequests(t *testing.T) { } hookCtx := &HookContext{ - Context: context.Background(), - Handler: nil, - Schema: "public", - Entity: "users", - Data: newUser, - Options: msg.Options, - Metadata: map[string]interface{}{"mqtt_client": client}, + Context: context.Background(), + TableName: "users", + Model: &TestUser{}, + ModelPtr: &TestUser{}, + Handler: nil, + Schema: "public", + Entity: "users", + Data: newUser, + Options: msg.Options, + Metadata: map[string]interface{}{"mqtt_client": client}, } handler.handleCreate(client, msg, hookCtx) diff --git a/pkg/security/providers.go b/pkg/security/providers.go index 57266eb..f9796fb 100644 --- a/pkg/security/providers.go +++ b/pkg/security/providers.go @@ -81,6 +81,9 @@ type DatabaseAuthenticator struct { queryMode QueryMode capability *dbCapability + // activityWG tracks in-flight asynchronous session activity updates + activityWG sync.WaitGroup + // Cookie session support (optional, gated by enableCookieSession) enableCookieSession bool cookieOptions SessionCookieOptions @@ -444,7 +447,12 @@ func (a *DatabaseAuthenticator) Authenticate(r *http.Request) (*UserContext, err // Authentication succeeded with this token // Update last activity timestamp asynchronously - go a.updateSessionActivity(r.Context(), token, &userCtx) + activityCtx := userCtx + a.activityWG.Add(1) + go func(ctx context.Context, token string) { + defer a.activityWG.Done() + a.updateSessionActivity(ctx, token, &activityCtx) + }(r.Context(), token) return &userCtx, nil } diff --git a/pkg/security/providers_test.go b/pkg/security/providers_test.go index 685d097..a5af9bb 100644 --- a/pkg/security/providers_test.go +++ b/pkg/security/providers_test.go @@ -194,7 +194,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) { WithArgs("cached-token-123", "authenticate"). WillReturnRows(rows) - userCtx1, err := auth.Authenticate(req) + userCtx1, err := authenticateSync(auth, req) if err != nil { t.Fatalf("first authenticate failed: %v", err) } @@ -203,7 +203,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) { } // Second call - should use cache, no database call expected - userCtx2, err := auth.Authenticate(req) + userCtx2, err := authenticateSync(auth, req) if err != nil { t.Fatalf("second authenticate failed: %v", err) } @@ -229,7 +229,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) { WithArgs("expire-token-456", "authenticate"). WillReturnRows(rows1) - _, err := auth.Authenticate(req) + _, err := authenticateSync(auth, req) if err != nil { t.Fatalf("first authenticate failed: %v", err) } @@ -245,7 +245,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) { WithArgs("expire-token-456", "authenticate"). WillReturnRows(rows2) - _, err = auth.Authenticate(req) + _, err = authenticateSync(auth, req) if err != nil { t.Fatalf("second authenticate after expiration failed: %v", err) } @@ -267,7 +267,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) { WithArgs("logout-token-789", "authenticate"). WillReturnRows(rows1) - _, err := auth.Authenticate(req) + _, err := authenticateSync(auth, req) if err != nil { t.Fatalf("authenticate failed: %v", err) } @@ -296,7 +296,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) { WithArgs("logout-token-789", "authenticate"). WillReturnRows(rows2) - _, err = auth.Authenticate(req) + _, err = authenticateSync(auth, req) if err != nil { t.Fatalf("authenticate after logout failed: %v", err) } @@ -318,7 +318,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) { WithArgs("manual-clear-token", "authenticate"). WillReturnRows(rows) - _, err := auth.Authenticate(req) + _, err := authenticateSync(auth, req) if err != nil { t.Fatalf("authenticate failed: %v", err) } @@ -334,7 +334,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) { WithArgs("manual-clear-token", "authenticate"). WillReturnRows(rows2) - _, err = auth.Authenticate(req) + _, err = authenticateSync(auth, req) if err != nil { t.Fatalf("authenticate after cache clear failed: %v", err) } @@ -356,7 +356,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) { WithArgs("user-token-1", "authenticate"). WillReturnRows(rows1) - _, err := auth.Authenticate(req1) + _, err := authenticateSync(auth, req1) if err != nil { t.Fatalf("first authenticate failed: %v", err) } @@ -371,7 +371,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) { WithArgs("user-token-2", "authenticate"). WillReturnRows(rows2) - _, err = auth.Authenticate(req2) + _, err = authenticateSync(auth, req2) if err != nil { t.Fatalf("second authenticate failed: %v", err) } @@ -387,7 +387,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) { WithArgs("user-token-1", "authenticate"). WillReturnRows(rows3) - _, err = auth.Authenticate(req1) + _, err = authenticateSync(auth, req1) if err != nil { t.Fatalf("authenticate after user cache clear failed: %v", err) } @@ -496,7 +496,7 @@ func TestDatabaseAuthenticator(t *testing.T) { WithArgs("test-token-123", "authenticate"). WillReturnRows(rows) - userCtx, err := auth.Authenticate(req) + userCtx, err := authenticateSync(auth, req) if err != nil { t.Fatalf("expected no error, got %v", err) } @@ -528,7 +528,7 @@ func TestDatabaseAuthenticator(t *testing.T) { WithArgs("cookie-token-456", "cookie"). WillReturnRows(rows) - userCtx, err := cookieAuth.Authenticate(req) + userCtx, err := authenticateSync(cookieAuth, req) if err != nil { t.Fatalf("expected no error, got %v", err) } @@ -545,7 +545,7 @@ func TestDatabaseAuthenticator(t *testing.T) { t.Run("authenticate missing token", func(t *testing.T) { req := httptest.NewRequest("GET", "/test", nil) - _, err := auth.Authenticate(req) + _, err := authenticateSync(auth, req) if err == nil { t.Fatal("expected error when token is missing") } @@ -571,7 +571,7 @@ func TestDatabaseAuthenticator(t *testing.T) { WithArgs("valid-token-123", "authenticate"). WillReturnRows(rows2) - userCtx, err := auth.Authenticate(req) + userCtx, err := authenticateSync(auth, req) if err != nil { t.Fatalf("expected no error, got %v", err) } @@ -597,7 +597,7 @@ func TestDatabaseAuthenticator(t *testing.T) { WithArgs("968CA5AE-4F83-4D55-A3C6-51AE4410E03A", "authenticate"). WillReturnRows(rows) - userCtx, err := auth.Authenticate(req) + userCtx, err := authenticateSync(auth, req) if err != nil { t.Fatalf("expected no error, got %v", err) } @@ -631,7 +631,7 @@ func TestDatabaseAuthenticator(t *testing.T) { WithArgs("bad-token-2", "authenticate"). WillReturnRows(rows2) - _, err := auth.Authenticate(req) + _, err := authenticateSync(auth, req) if err == nil { t.Fatal("expected error when all tokens fail") } @@ -891,7 +891,7 @@ func TestDatabaseAuthenticatorReconnectsClosedDBPaths(t *testing.T) { WithArgs("reconnect-auth-token", "authenticate"). WillReturnRows(reconnectRows) - userCtx, err := auth.Authenticate(req) + userCtx, err := authenticateSync(auth, req) if err != nil { t.Fatalf("expected authenticate to reconnect, got %v", err) } @@ -1328,3 +1328,11 @@ func TestConfigRowSecurityProvider(t *testing.T) { } }) } + +// authenticateSync authenticates and waits for the asynchronous session +// activity update so sqlmock expectations are never touched concurrently. +func authenticateSync(auth *DatabaseAuthenticator, req *http.Request) (*UserContext, error) { + userCtx, err := auth.Authenticate(req) + auth.activityWG.Wait() + return userCtx, err +} diff --git a/pkg/websocketspec/hooks_test.go b/pkg/websocketspec/hooks_test.go index 01be934..0b2043d 100644 --- a/pkg/websocketspec/hooks_test.go +++ b/pkg/websocketspec/hooks_test.go @@ -3,6 +3,7 @@ package websocketspec import ( "context" "errors" + "sync/atomic" "testing" "github.com/bitechdev/ResolveSpec/pkg/common" @@ -519,9 +520,9 @@ func TestHookRegistry_ConcurrentExecution(t *testing.T) { // This test verifies that concurrent hook executions don't cause race conditions // Run with: go test -race - counter := 0 + var counter atomic.Int64 hook := func(ctx *HookContext) error { - counter++ + counter.Add(1) return nil } @@ -543,5 +544,5 @@ func TestHookRegistry_ConcurrentExecution(t *testing.T) { <-done } - assert.Equal(t, 10, counter) + assert.Equal(t, int64(10), counter.Load()) }