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.
This commit is contained in:
Hein
2026-09-30 13:07:47 +02:00
parent e1cf72834e
commit c7b4530689
8 changed files with 235 additions and 121 deletions
+12 -4
View File
@@ -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 GOLANGCI_LINT := $(shell go env GOPATH)/bin/golangci-lint
# Run all unit tests # Run all unit tests
test-unit: test-unit:
@echo "Running unit tests..." @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) # Run all integration tests (requires PostgreSQL)
test-integration: test-integration:
@@ -13,7 +20,7 @@ test-integration:
@go test -tags=integration ./pkg/resolvespec ./pkg/restheadspec -v @go test -tags=integration ./pkg/resolvespec ./pkg/restheadspec -v
# Run all tests (unit + integration) # 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) 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 \ @if [ -z "$(VERSION)" ]; then \
@@ -113,7 +120,8 @@ coverage-integration:
help: help:
@echo "Available targets:" @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-integration - Run integration tests (requires PostgreSQL)"
@echo " test - Run all tests" @echo " test - Run all tests"
@echo " docker-up - Start PostgreSQL container" @echo " docker-up - Start PostgreSQL container"
+23
View File
@@ -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 have to be fixed together: a race detector pointed at packages with no tests
finds nothing. 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 ### X2. High — `go test` runs against 2 of 23 packages
+8 -8
View File
@@ -171,7 +171,7 @@ func TestBrokerPublishAsync(t *testing.T) {
// Publish multiple events // Publish multiple events
for i := 0; i < 5; i++ { for i := 0; i < 5; i++ {
event := NewEvent(EventSourceSystem, "test.event") event := NewEvent(EventSourceSystem, "test.event")
event.InstanceID = "test-instance" event.InstanceID = "test-instance"
if err := broker.PublishAsync(context.Background(), event); err != nil { if err := broker.PublishAsync(context.Background(), event); err != nil {
t.Fatalf("PublishAsync failed: %v", err) t.Fatalf("PublishAsync failed: %v", err)
} }
@@ -346,7 +346,7 @@ func TestBrokerStats(t *testing.T) {
// Publish events // Publish events
for i := 0; i < 3; i++ { for i := 0; i < 3; i++ {
event := NewEvent(EventSourceSystem, "test.event") event := NewEvent(EventSourceSystem, "test.event")
event.InstanceID = "test-instance" event.InstanceID = "test-instance"
broker.PublishSync(context.Background(), event) broker.PublishSync(context.Background(), event)
} }
@@ -413,7 +413,7 @@ func TestBrokerConcurrentPublish(t *testing.T) {
go func() { go func() {
defer wg.Done() defer wg.Done()
event := NewEvent(EventSourceSystem, "test.event") event := NewEvent(EventSourceSystem, "test.event")
event.InstanceID = "test-instance" event.InstanceID = "test-instance"
broker.PublishAsync(context.Background(), event) broker.PublishAsync(context.Background(), event)
}() }()
} }
@@ -450,7 +450,7 @@ func TestBrokerGracefulShutdown(t *testing.T) {
// Publish events // Publish events
for i := 0; i < 5; i++ { for i := 0; i < 5; i++ {
event := NewEvent(EventSourceSystem, "test.event") event := NewEvent(EventSourceSystem, "test.event")
event.InstanceID = "test-instance" event.InstanceID = "test-instance"
broker.PublishAsync(context.Background(), event) broker.PublishAsync(context.Background(), event)
} }
@@ -502,21 +502,21 @@ func TestBrokerProcessingModes(t *testing.T) {
broker.Start(context.Background()) broker.Start(context.Background())
defer broker.Stop(context.Background()) defer broker.Stop(context.Background())
called := false var called atomic.Bool
broker.Subscribe("test.*", EventHandlerFunc(func(ctx context.Context, event *Event) error { broker.Subscribe("test.*", EventHandlerFunc(func(ctx context.Context, event *Event) error {
called = true called.Store(true)
return nil return nil
})) }))
event := NewEvent(EventSourceSystem, "test.event") event := NewEvent(EventSourceSystem, "test.event")
event.InstanceID = "test-instance" event.InstanceID = "test-instance"
broker.Publish(context.Background(), event) broker.Publish(context.Background(), event)
if tt.mode == ProcessingModeAsync { if tt.mode == ProcessingModeAsync {
time.Sleep(50 * time.Millisecond) time.Sleep(50 * time.Millisecond)
} }
if !called { if !called.Load() {
t.Error("Expected handler to be called") t.Error("Expected handler to be called")
} }
}) })
+59 -23
View File
@@ -6,15 +6,40 @@ import (
"log" "log"
"os" "os"
"runtime/debug" "runtime/debug"
"sync"
"go.uber.org/zap" "go.uber.org/zap"
errortracking "github.com/bitechdev/ResolveSpec/pkg/errortracking" 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 Logger *zap.SugaredLogger
var errorTracker errortracking.Provider 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) { func Init(dev bool) {
if dev { if dev {
@@ -49,28 +74,30 @@ func UpdateLogger(config *zap.Config) {
return return
} }
Logger = logger.Sugar() setLogger(logger.Sugar())
Info("ResolveSpec Logger initialized") Info("ResolveSpec Logger initialized")
} }
// InitErrorTracking initializes the error tracking provider // InitErrorTracking initializes the error tracking provider
func InitErrorTracking(provider errortracking.Provider) { func InitErrorTracking(provider errortracking.Provider) {
stateMu.Lock()
errorTracker = provider errorTracker = provider
if errorTracker != nil { stateMu.Unlock()
if provider != nil {
Info("Error tracking initialized") Info("Error tracking initialized")
} }
} }
// GetErrorTracker returns the current error tracking provider // GetErrorTracker returns the current error tracking provider
func GetErrorTracker() errortracking.Provider { func GetErrorTracker() errortracking.Provider {
return errorTracker return getErrorTracker()
} }
// CloseErrorTracking flushes and closes the error tracking provider // CloseErrorTracking flushes and closes the error tracking provider
func CloseErrorTracking() error { func CloseErrorTracking() error {
if errorTracker != nil { if tracker := getErrorTracker(); tracker != nil {
errorTracker.Flush(5) tracker.Flush(5)
return errorTracker.Close() return tracker.Close()
} }
return nil return nil
} }
@@ -98,53 +125,59 @@ func extractContext(args ...interface{}) (ctx context.Context, filteredArgs []in
} }
func Info(template string, args ...interface{}) { func Info(template string, args ...interface{}) {
if Logger == nil { lg := getLogger()
if lg == nil {
log.Printf(template, args...) log.Printf(template, args...)
return 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{}) { func Warn(template string, args ...interface{}) {
lg := getLogger()
tracker := getErrorTracker()
ctx, remainingArgs := extractContext(args...) ctx, remainingArgs := extractContext(args...)
message := fmt.Sprintf(template, remainingArgs...) message := fmt.Sprintf(template, remainingArgs...)
if Logger == nil { if lg == nil {
log.Printf("%s", message) log.Printf("%s", message)
} else { } else {
Logger.Warnw(message, "process_id", os.Getpid()) lg.Warnw(message, "process_id", os.Getpid())
} }
// Send to error tracker // Send to error tracker
if errorTracker != nil { if tracker != nil {
errorTracker.CaptureMessage(ctx, message, errortracking.SeverityWarning, map[string]interface{}{ tracker.CaptureMessage(ctx, message, errortracking.SeverityWarning, map[string]interface{}{
"process_id": os.Getpid(), "process_id": os.Getpid(),
}) })
} }
} }
func Error(template string, args ...interface{}) { func Error(template string, args ...interface{}) {
lg := getLogger()
tracker := getErrorTracker()
ctx, remainingArgs := extractContext(args...) ctx, remainingArgs := extractContext(args...)
message := fmt.Sprintf(template, remainingArgs...) message := fmt.Sprintf(template, remainingArgs...)
if Logger == nil { if lg == nil {
log.Printf("%s", message) log.Printf("%s", message)
} else { } else {
Logger.Errorw(message, "process_id", os.Getpid()) lg.Errorw(message, "process_id", os.Getpid())
} }
// Send to error tracker // Send to error tracker
if errorTracker != nil { if tracker != nil {
errorTracker.CaptureMessage(ctx, message, errortracking.SeverityError, map[string]interface{}{ tracker.CaptureMessage(ctx, message, errortracking.SeverityError, map[string]interface{}{
"process_id": os.Getpid(), "process_id": os.Getpid(),
}) })
} }
} }
func Debug(template string, args ...interface{}) { func Debug(template string, args ...interface{}) {
if Logger == nil { lg := getLogger()
if lg == nil {
log.Printf(template, args...) log.Printf(template, args...)
return return
} }
Logger.Debugw(fmt.Sprintf(template, args...), "process_id", os.Getpid()) lg.Debugw(fmt.Sprintf(template, args...), "process_id", os.Getpid())
} }
// CatchPanic - Handle panic // CatchPanic - Handle panic
@@ -155,8 +188,10 @@ func CatchPanicCallback(location string, cb func(err any), args ...interface{})
return func() { return func() {
if err := recover(); err != nil { if err := recover(); err != nil {
callstack := debug.Stack() callstack := debug.Stack()
lg := getLogger()
tracker := getErrorTracker()
if Logger != nil { if lg != nil {
Error("Panic in %s : %v", location, err, ctx) // Pass context implicitly Error("Panic in %s : %v", location, err, ctx) // Pass context implicitly
} else { } else {
fmt.Printf("%s:PANIC->%+v", location, err) 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 // Send to error tracker
if errorTracker != nil { if tracker != nil {
errorTracker.CapturePanic(ctx, err, callstack, map[string]interface{}{ tracker.CapturePanic(ctx, err, callstack, map[string]interface{}{
"location": location, "location": location,
"process_id": os.Getpid(), "process_id": os.Getpid(),
}) })
@@ -195,13 +230,14 @@ func CatchPanic(location string, args ...interface{}) func() {
// } // }
// }() // }()
func HandlePanic(methodName string, r any, args ...interface{}) error { func HandlePanic(methodName string, r any, args ...interface{}) error {
tracker := getErrorTracker()
ctx, _ := extractContext(args...) ctx, _ := extractContext(args...)
stack := debug.Stack() stack := debug.Stack()
Error("Panic in %s: %v\nStack trace:\n%s", methodName, r, string(stack), ctx) // Pass context implicitly Error("Panic in %s: %v\nStack trace:\n%s", methodName, r, string(stack), ctx) // Pass context implicitly
// Send to error tracker // Send to error tracker
if errorTracker != nil { if tracker != nil {
errorTracker.CapturePanic(ctx, r, stack, map[string]interface{}{ tracker.CapturePanic(ctx, r, stack, map[string]interface{}{
"method": methodName, "method": methodName,
"process_id": os.Getpid(), "process_id": os.Getpid(),
}) })
+94 -64
View File
@@ -41,6 +41,12 @@ func setupTestHandler(t *testing.T) (*Handler, *gorm.DB) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err) 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 // Auto-migrate test model
err = db.AutoMigrate(&TestUser{}) err = db.AutoMigrate(&TestUser{})
require.NoError(t, err) require.NoError(t, err)
@@ -93,9 +99,9 @@ func TestHandler_HandleRead_Single(t *testing.T) {
// Insert test data // Insert test data
user := &TestUser{ user := &TestUser{
ID: 1, ID: 1,
Name: "John Doe", Name: "John Doe",
Email: "john@example.com", Email: "john@example.com",
Status: "active", Status: "active",
} }
db.Create(user) db.Create(user)
@@ -115,13 +121,16 @@ func TestHandler_HandleRead_Single(t *testing.T) {
// Create hook context // Create hook context
hookCtx := &HookContext{ hookCtx := &HookContext{
Context: context.Background(), Context: context.Background(),
Handler: nil, TableName: "users",
Schema: "public", Model: &TestUser{},
Entity: "users", ModelPtr: &TestUser{},
ID: "1", Handler: nil,
Options: msg.Options, Schema: "public",
Metadata: map[string]interface{}{"mqtt_client": client}, Entity: "users",
ID: "1",
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
} }
// Handle read // Handle read
@@ -164,12 +173,15 @@ func TestHandler_HandleRead_Multiple(t *testing.T) {
// Create hook context // Create hook context
hookCtx := &HookContext{ hookCtx := &HookContext{
Context: context.Background(), Context: context.Background(),
Handler: nil, TableName: "users",
Schema: "public", Model: &TestUser{},
Entity: "users", ModelPtr: &TestUser{},
Options: msg.Options, Handler: nil,
Metadata: map[string]interface{}{"mqtt_client": client}, Schema: "public",
Entity: "users",
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
} }
// Handle read // Handle read
@@ -208,13 +220,16 @@ func TestHandler_HandleCreate(t *testing.T) {
// Create hook context // Create hook context
hookCtx := &HookContext{ hookCtx := &HookContext{
Context: context.Background(), Context: context.Background(),
Handler: nil, TableName: "users",
Schema: "public", Model: &TestUser{},
Entity: "users", ModelPtr: &TestUser{},
Data: newUser, Handler: nil,
Options: msg.Options, Schema: "public",
Metadata: map[string]interface{}{"mqtt_client": client}, Entity: "users",
Data: newUser,
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
} }
// Handle create // Handle create
@@ -266,14 +281,17 @@ func TestHandler_HandleUpdate(t *testing.T) {
// Create hook context // Create hook context
hookCtx := &HookContext{ hookCtx := &HookContext{
Context: context.Background(), Context: context.Background(),
Handler: nil, TableName: "users",
Schema: "public", Model: &TestUser{},
Entity: "users", ModelPtr: &TestUser{},
ID: "1", Handler: nil,
Data: updateData, Schema: "public",
Options: msg.Options, Entity: "users",
Metadata: map[string]interface{}{"mqtt_client": client}, ID: "1",
Data: updateData,
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
} }
// Handle update // Handle update
@@ -318,13 +336,16 @@ func TestHandler_HandleDelete(t *testing.T) {
// Create hook context // Create hook context
hookCtx := &HookContext{ hookCtx := &HookContext{
Context: context.Background(), Context: context.Background(),
Handler: nil, TableName: "users",
Schema: "public", Model: &TestUser{},
Entity: "users", ModelPtr: &TestUser{},
ID: "1", Handler: nil,
Options: msg.Options, Schema: "public",
Metadata: map[string]interface{}{"mqtt_client": client}, Entity: "users",
ID: "1",
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
} }
// Handle delete // Handle delete
@@ -388,13 +409,13 @@ func TestHandler_HandleUnsubscribe(t *testing.T) {
sub := handler.subscriptionManager.Subscribe("sub-1", client.ID, "public", "users", &common.RequestOptions{}) sub := handler.subscriptionManager.Subscribe("sub-1", client.ID, "public", "users", &common.RequestOptions{})
client.AddSubscription(sub) client.AddSubscription(sub)
// Create unsubscribe message with subscription ID in Data // Create unsubscribe message with the subscription ID
msg := &Message{ msg := &Message{
ID: "msg-7", ID: "msg-7",
Type: MessageTypeSubscription, Type: MessageTypeSubscription,
Operation: OperationUnsubscribe, Operation: OperationUnsubscribe,
Data: map[string]interface{}{"subscription_id": "sub-1"}, SubscriptionID: "sub-1",
Options: &common.RequestOptions{}, Options: &common.RequestOptions{},
} }
// Handle unsubscribe // Handle unsubscribe
@@ -490,12 +511,15 @@ func TestHandler_Hooks_BeforeRead(t *testing.T) {
// Create hook context // Create hook context
hookCtx := &HookContext{ hookCtx := &HookContext{
Context: context.Background(), Context: context.Background(),
Handler: nil, TableName: "users",
Schema: "public", Model: &TestUser{},
Entity: "users", ModelPtr: &TestUser{},
Options: msg.Options, Handler: nil,
Metadata: map[string]interface{}{"mqtt_client": client}, Schema: "public",
Entity: "users",
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
} }
// Handle read // Handle read
@@ -544,13 +568,16 @@ func TestHandler_Hooks_BeforeCreate(t *testing.T) {
} }
hookCtx := &HookContext{ hookCtx := &HookContext{
Context: context.Background(), Context: context.Background(),
Handler: nil, TableName: "users",
Schema: "public", Model: &TestUser{},
Entity: "users", ModelPtr: &TestUser{},
Data: newUser, Handler: nil,
Options: msg.Options, Schema: "public",
Metadata: map[string]interface{}{"mqtt_client": client}, Entity: "users",
Data: newUser,
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
} }
// Handle create // Handle create
@@ -600,13 +627,16 @@ func TestHandler_ConcurrentRequests(t *testing.T) {
} }
hookCtx := &HookContext{ hookCtx := &HookContext{
Context: context.Background(), Context: context.Background(),
Handler: nil, TableName: "users",
Schema: "public", Model: &TestUser{},
Entity: "users", ModelPtr: &TestUser{},
Data: newUser, Handler: nil,
Options: msg.Options, Schema: "public",
Metadata: map[string]interface{}{"mqtt_client": client}, Entity: "users",
Data: newUser,
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
} }
handler.handleCreate(client, msg, hookCtx) handler.handleCreate(client, msg, hookCtx)
+9 -1
View File
@@ -81,6 +81,9 @@ type DatabaseAuthenticator struct {
queryMode QueryMode queryMode QueryMode
capability *dbCapability capability *dbCapability
// activityWG tracks in-flight asynchronous session activity updates
activityWG sync.WaitGroup
// Cookie session support (optional, gated by enableCookieSession) // Cookie session support (optional, gated by enableCookieSession)
enableCookieSession bool enableCookieSession bool
cookieOptions SessionCookieOptions cookieOptions SessionCookieOptions
@@ -444,7 +447,12 @@ func (a *DatabaseAuthenticator) Authenticate(r *http.Request) (*UserContext, err
// Authentication succeeded with this token // Authentication succeeded with this token
// Update last activity timestamp asynchronously // 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 return &userCtx, nil
} }
+26 -18
View File
@@ -194,7 +194,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("cached-token-123", "authenticate"). WithArgs("cached-token-123", "authenticate").
WillReturnRows(rows) WillReturnRows(rows)
userCtx1, err := auth.Authenticate(req) userCtx1, err := authenticateSync(auth, req)
if err != nil { if err != nil {
t.Fatalf("first authenticate failed: %v", err) 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 // Second call - should use cache, no database call expected
userCtx2, err := auth.Authenticate(req) userCtx2, err := authenticateSync(auth, req)
if err != nil { if err != nil {
t.Fatalf("second authenticate failed: %v", err) t.Fatalf("second authenticate failed: %v", err)
} }
@@ -229,7 +229,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("expire-token-456", "authenticate"). WithArgs("expire-token-456", "authenticate").
WillReturnRows(rows1) WillReturnRows(rows1)
_, err := auth.Authenticate(req) _, err := authenticateSync(auth, req)
if err != nil { if err != nil {
t.Fatalf("first authenticate failed: %v", err) t.Fatalf("first authenticate failed: %v", err)
} }
@@ -245,7 +245,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("expire-token-456", "authenticate"). WithArgs("expire-token-456", "authenticate").
WillReturnRows(rows2) WillReturnRows(rows2)
_, err = auth.Authenticate(req) _, err = authenticateSync(auth, req)
if err != nil { if err != nil {
t.Fatalf("second authenticate after expiration failed: %v", err) t.Fatalf("second authenticate after expiration failed: %v", err)
} }
@@ -267,7 +267,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("logout-token-789", "authenticate"). WithArgs("logout-token-789", "authenticate").
WillReturnRows(rows1) WillReturnRows(rows1)
_, err := auth.Authenticate(req) _, err := authenticateSync(auth, req)
if err != nil { if err != nil {
t.Fatalf("authenticate failed: %v", err) t.Fatalf("authenticate failed: %v", err)
} }
@@ -296,7 +296,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("logout-token-789", "authenticate"). WithArgs("logout-token-789", "authenticate").
WillReturnRows(rows2) WillReturnRows(rows2)
_, err = auth.Authenticate(req) _, err = authenticateSync(auth, req)
if err != nil { if err != nil {
t.Fatalf("authenticate after logout failed: %v", err) t.Fatalf("authenticate after logout failed: %v", err)
} }
@@ -318,7 +318,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("manual-clear-token", "authenticate"). WithArgs("manual-clear-token", "authenticate").
WillReturnRows(rows) WillReturnRows(rows)
_, err := auth.Authenticate(req) _, err := authenticateSync(auth, req)
if err != nil { if err != nil {
t.Fatalf("authenticate failed: %v", err) t.Fatalf("authenticate failed: %v", err)
} }
@@ -334,7 +334,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("manual-clear-token", "authenticate"). WithArgs("manual-clear-token", "authenticate").
WillReturnRows(rows2) WillReturnRows(rows2)
_, err = auth.Authenticate(req) _, err = authenticateSync(auth, req)
if err != nil { if err != nil {
t.Fatalf("authenticate after cache clear failed: %v", err) t.Fatalf("authenticate after cache clear failed: %v", err)
} }
@@ -356,7 +356,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("user-token-1", "authenticate"). WithArgs("user-token-1", "authenticate").
WillReturnRows(rows1) WillReturnRows(rows1)
_, err := auth.Authenticate(req1) _, err := authenticateSync(auth, req1)
if err != nil { if err != nil {
t.Fatalf("first authenticate failed: %v", err) t.Fatalf("first authenticate failed: %v", err)
} }
@@ -371,7 +371,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("user-token-2", "authenticate"). WithArgs("user-token-2", "authenticate").
WillReturnRows(rows2) WillReturnRows(rows2)
_, err = auth.Authenticate(req2) _, err = authenticateSync(auth, req2)
if err != nil { if err != nil {
t.Fatalf("second authenticate failed: %v", err) t.Fatalf("second authenticate failed: %v", err)
} }
@@ -387,7 +387,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("user-token-1", "authenticate"). WithArgs("user-token-1", "authenticate").
WillReturnRows(rows3) WillReturnRows(rows3)
_, err = auth.Authenticate(req1) _, err = authenticateSync(auth, req1)
if err != nil { if err != nil {
t.Fatalf("authenticate after user cache clear failed: %v", err) t.Fatalf("authenticate after user cache clear failed: %v", err)
} }
@@ -496,7 +496,7 @@ func TestDatabaseAuthenticator(t *testing.T) {
WithArgs("test-token-123", "authenticate"). WithArgs("test-token-123", "authenticate").
WillReturnRows(rows) WillReturnRows(rows)
userCtx, err := auth.Authenticate(req) userCtx, err := authenticateSync(auth, req)
if err != nil { if err != nil {
t.Fatalf("expected no error, got %v", err) t.Fatalf("expected no error, got %v", err)
} }
@@ -528,7 +528,7 @@ func TestDatabaseAuthenticator(t *testing.T) {
WithArgs("cookie-token-456", "cookie"). WithArgs("cookie-token-456", "cookie").
WillReturnRows(rows) WillReturnRows(rows)
userCtx, err := cookieAuth.Authenticate(req) userCtx, err := authenticateSync(cookieAuth, req)
if err != nil { if err != nil {
t.Fatalf("expected no error, got %v", err) 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) { t.Run("authenticate missing token", func(t *testing.T) {
req := httptest.NewRequest("GET", "/test", nil) req := httptest.NewRequest("GET", "/test", nil)
_, err := auth.Authenticate(req) _, err := authenticateSync(auth, req)
if err == nil { if err == nil {
t.Fatal("expected error when token is missing") t.Fatal("expected error when token is missing")
} }
@@ -571,7 +571,7 @@ func TestDatabaseAuthenticator(t *testing.T) {
WithArgs("valid-token-123", "authenticate"). WithArgs("valid-token-123", "authenticate").
WillReturnRows(rows2) WillReturnRows(rows2)
userCtx, err := auth.Authenticate(req) userCtx, err := authenticateSync(auth, req)
if err != nil { if err != nil {
t.Fatalf("expected no error, got %v", err) t.Fatalf("expected no error, got %v", err)
} }
@@ -597,7 +597,7 @@ func TestDatabaseAuthenticator(t *testing.T) {
WithArgs("968CA5AE-4F83-4D55-A3C6-51AE4410E03A", "authenticate"). WithArgs("968CA5AE-4F83-4D55-A3C6-51AE4410E03A", "authenticate").
WillReturnRows(rows) WillReturnRows(rows)
userCtx, err := auth.Authenticate(req) userCtx, err := authenticateSync(auth, req)
if err != nil { if err != nil {
t.Fatalf("expected no error, got %v", err) t.Fatalf("expected no error, got %v", err)
} }
@@ -631,7 +631,7 @@ func TestDatabaseAuthenticator(t *testing.T) {
WithArgs("bad-token-2", "authenticate"). WithArgs("bad-token-2", "authenticate").
WillReturnRows(rows2) WillReturnRows(rows2)
_, err := auth.Authenticate(req) _, err := authenticateSync(auth, req)
if err == nil { if err == nil {
t.Fatal("expected error when all tokens fail") t.Fatal("expected error when all tokens fail")
} }
@@ -891,7 +891,7 @@ func TestDatabaseAuthenticatorReconnectsClosedDBPaths(t *testing.T) {
WithArgs("reconnect-auth-token", "authenticate"). WithArgs("reconnect-auth-token", "authenticate").
WillReturnRows(reconnectRows) WillReturnRows(reconnectRows)
userCtx, err := auth.Authenticate(req) userCtx, err := authenticateSync(auth, req)
if err != nil { if err != nil {
t.Fatalf("expected authenticate to reconnect, got %v", err) 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
}
+4 -3
View File
@@ -3,6 +3,7 @@ package websocketspec
import ( import (
"context" "context"
"errors" "errors"
"sync/atomic"
"testing" "testing"
"github.com/bitechdev/ResolveSpec/pkg/common" "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 // This test verifies that concurrent hook executions don't cause race conditions
// Run with: go test -race // Run with: go test -race
counter := 0 var counter atomic.Int64
hook := func(ctx *HookContext) error { hook := func(ctx *HookContext) error {
counter++ counter.Add(1)
return nil return nil
} }
@@ -543,5 +544,5 @@ func TestHookRegistry_ConcurrentExecution(t *testing.T) {
<-done <-done
} }
assert.Equal(t, 10, counter) assert.Equal(t, int64(10), counter.Load())
} }