mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-05 13:01:58 +00:00
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:
@@ -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"
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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(),
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
|
|||||||
@@ -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())
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user