diff --git a/action/action.go b/action/action.go index 65ffafc..fd41820 100644 --- a/action/action.go +++ b/action/action.go @@ -15,6 +15,12 @@ import ( var ErrTypeAssertion = errors.New("critical type assertion failure") +type execStateKey struct{} + +type execState struct { + fromCache bool +} + type BuiltAction[Req, Res any] struct { meta *Meta exec Fn[Req, Res] @@ -41,16 +47,22 @@ func (a *BuiltAction[Req, Res]) GetMeta() *Meta { cp.RequiredFeatures = slices.Clone(a.meta.RequiredFeatures) return &cp } + func (a *BuiltAction[Req, Res]) GetBindings() []Binding { return append([]Binding(nil), a.bindings...) } + func (a *BuiltAction[Req, Res]) GetAnyHooks() []AnyHook { return append([]AnyHook(nil), a.anyHooksSnapshot()...) } + func (a *BuiltAction[Req, Res]) Describe() *Meta { return a.GetMeta() } -func (a *BuiltAction[Req, Res]) History() *History[Req, Res] { return a.history } + +func (a *BuiltAction[Req, Res]) History() *History[Req, Res] { + return a.history +} func (a *BuiltAction[Req, Res]) anyHooksSnapshot() []AnyHook { set := a.anyHooks.Load() @@ -60,11 +72,34 @@ func (a *BuiltAction[Req, Res]) anyHooksSnapshot() []AnyHook { return set.hooks } +func hasOnExecutedHook[Req, Res any](hooks []Hook[Req, Res], anyHooks []AnyHook) bool { + for i := range hooks { + if hooks[i].OnExecuted != nil { + return true + } + } + for i := range anyHooks { + if anyHooks[i].OnExecuted != nil { + return true + } + } + return false +} + func (a *BuiltAction[Req, Res]) Do(ctx context.Context, req Req) (res Res, err error) { var anyHooksRan, typedHooksRan int - finalCtx := ctx + anyHooks := a.anyHooksSnapshot() + needExecState := hasOnExecutedHook(a.hooks, anyHooks) + + var state *execState + finalCtx := ctx + if needExecState { + state = &execState{} + finalCtx = context.WithValue(ctx, execStateKey{}, state) + } + defer func() { if r := recover(); r != nil { err = xerr.PanicRecovery(r) @@ -80,64 +115,68 @@ func (a *BuiltAction[Req, Res]) Do(ctx context.Context, req Req) (res Res, err e for i := typedHooksRan - 1; i >= 0; i-- { h := a.hooks[i] - if h.After != nil { - h := h - callHook(a.meta, "After", func() { - h.After(finalCtx, req, res, err, a.meta) - }) - } - if err != nil { - if errors.Is(err, context.Canceled) { - if h.OnCancel != nil { - h := h - callHook(a.meta, "OnCancel", func() { - h.OnCancel(finalCtx, req, a.meta) - }) - } - } else { - if h.OnError != nil { - h := h - callHook(a.meta, "OnError", func() { - h.OnError(finalCtx, req, err, a.meta) - }) - } + + switch { + case err != nil && errors.Is(err, context.Canceled): + if h.OnCancel != nil { + h := h + callHook(a.meta, "OnCancel", func() { + h.OnCancel(finalCtx, req, a.meta) + }) + } + case err != nil: + if h.OnError != nil { + h := h + callHook(a.meta, "OnError", func() { + h.OnError(finalCtx, req, err, a.meta) + }) } - } else if h.OnExecuted != nil { + case h.OnExecuted != nil && state != nil && !state.fromCache: h := h callHook(a.meta, "OnExecuted", func() { h.OnExecuted(finalCtx, req, res, nil, a.meta) }) } + + if h.After != nil { + h := h + callHook(a.meta, "After", func() { + h.After(finalCtx, req, res, err, a.meta) + }) + } } for i := anyHooksRan - 1; i >= 0; i-- { h := anyHooks[i] - if errors.Is(err, context.Canceled) { + + switch { + case errors.Is(err, context.Canceled): if h.OnCancel != nil { h := h callHook(a.meta, "OnCancel", func() { h.OnCancel(finalCtx, any(req), a.meta) }) } - continue + case err != nil: + if h.OnError != nil { + h := h + callHook(a.meta, "OnError", func() { + h.OnError(finalCtx, any(req), err, a.meta) + }) + } + case h.OnExecuted != nil && state != nil && !state.fromCache: + h := h + callHook(a.meta, "OnExecuted", func() { + h.OnExecuted(finalCtx, any(req), any(res), nil, a.meta) + }) } + if h.After != nil { h := h callHook(a.meta, "After", func() { h.After(finalCtx, any(req), any(res), err, a.meta) }) } - if err != nil && h.OnError != nil { - h := h - callHook(a.meta, "OnError", func() { - h.OnError(finalCtx, any(req), err, a.meta) - }) - } else if err == nil && h.OnExecuted != nil { - h := h - callHook(a.meta, "OnExecuted", func() { - h.OnExecuted(finalCtx, any(req), any(res), nil, a.meta) - }) - } } }() @@ -175,6 +214,10 @@ func (a *BuiltAction[Req, Res]) Do(ctx context.Context, req Req) (res Res, err e } func (a *BuiltAction[Req, Res]) OnCacheHit(ctx context.Context, req Req, res Res) { + if s, ok := ctx.Value(execStateKey{}).(*execState); ok && s != nil { + s.fromCache = true + } + for _, h := range a.hooks { if h.OnCacheHit != nil { h.OnCacheHit(ctx, req, res, a.meta) diff --git a/action/action_test.go b/action/action_test.go index 064b550..9cf1466 100644 --- a/action/action_test.go +++ b/action/action_test.go @@ -3,87 +3,140 @@ package action_test import ( "context" "errors" - "fmt" + "sync/atomic" "testing" "github.com/nexssp/kernel/action" + "github.com/nexssp/kernel/xerr" ) -func TestExecuteDecoded_PointerRequest(t *testing.T) { - type Req struct{ A int } - act := action.New("test", func(ctx context.Context, req *Req) (int, error) { - return req.A, nil - }).Build() - res, err := act.ExecuteDecoded(context.Background(), func(v any) error { - req, ok := v.(*Req) - if !ok { - return fmt.Errorf("expected *Req, got %T", v) +func TestBuiltAction_Do_LifecycleOrder(t *testing.T) { + t.Parallel() + + var steps []string + act := action.New("order.test", func(ctx context.Context, req string) (string, error) { + steps = append(steps, "handler") + return "result_" + req, nil + }). + HookBefore(func(ctx context.Context, req string, meta *action.Meta) (context.Context, error) { + steps = append(steps, "before") + return ctx, nil + }). + HookAfter(func(ctx context.Context, req, res string, err error, meta *action.Meta) { + steps = append(steps, "after") + }). + HookExecuted(func(ctx context.Context, req, res string, meta *action.Meta) { + steps = append(steps, "executed") + }). + Build() + + res, err := act.Do(context.Background(), "input") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if res != "result_input" { + t.Fatalf("expected 'result_input', got %q", res) + } + + expectedSteps := []string{"before", "handler", "executed", "after"} + if len(steps) != len(expectedSteps) { + t.Fatalf("expected steps %v, got %v", expectedSteps, steps) + } + for i, step := range steps { + if step != expectedSteps[i] { + t.Errorf("step %d: expected %q, got %q", i, expectedSteps[i], step) } - *req = Req{A: 42} - return nil - }) - if err != nil || res != 42 { - t.Fail() } } -func TestExecuteDecoded_ValueRequest(t *testing.T) { +func TestBuiltAction_Do_ErrorLifecycle(t *testing.T) { t.Parallel() - type Req struct{ A int } // Value type, not pointer - act := action.New("test.value", func(ctx context.Context, req Req) (int, error) { - return req.A, nil - }).Build() - res, err := act.ExecuteDecoded(context.Background(), func(v any) error { - // v is passed as a pointer to the value type by ExecuteDecoded - req, ok := v.(*Req) - if !ok { - return fmt.Errorf("expected *Req, got %T", v) - } - *req = Req{A: 99} - return nil - }) - if err != nil || res != 99 { - t.Fatalf("expected 99, got %v (err: %v)", res, err) + var errorHookCalled, executedHookCalled bool + act := action.New("error.test", func(ctx context.Context, req string) (string, error) { + return "", errors.New("business failure") + }). + HookError(func(ctx context.Context, req string, err error, meta *action.Meta) { + errorHookCalled = true + }). + HookExecuted(func(ctx context.Context, req, res string, meta *action.Meta) { + executedHookCalled = true + }). + Build() + + _, err := act.Do(context.Background(), "req") + if err == nil { + t.Fatal("expected error, got nil") + } + if !errorHookCalled { + t.Fatal("expected HookError to be called") + } + if executedHookCalled { + t.Fatal("HookExecuted must not be called when execution fails") } } -func TestRace_AllFailures(t *testing.T) { +func TestBuiltAction_Do_PanicRecovery(t *testing.T) { t.Parallel() - act := action.New("race.fail", func(ctx context.Context, req int) (int, error) { - return 0, errors.New("always fails") - }).Build() - res, err := action.Race(context.Background(), act, []int{1, 2}) + var panicHookRan atomic.Bool + act := action.New("panic.test", func(ctx context.Context, req string) (string, error) { + panic("fatal unexpected crash") + }). + AnyHook(action.AnyHook{ + OnPanic: func(ctx context.Context, req, recovered any, meta *action.Meta) { + panicHookRan.Store(true) + }, + }). + Build() + + res, err := act.Do(context.Background(), "req") + if res != "" { + t.Fatalf("expected empty result on panic, got %q", res) + } if err == nil { - t.Fatal("expected error, got nil") + t.Fatal("expected error from recovered panic, got nil") + } + + var appErr *xerr.AppError + if !errors.As(err, &appErr) || appErr.Kind != xerr.KindInternal { + t.Fatalf("expected xerr.KindInternal, got %v", err) } - if res != 0 { - t.Fatalf("expected zero value, got %d", res) + if !panicHookRan.Load() { + t.Fatal("expected AnyHook.OnPanic to execute") } } -func TestAction_ContextCancelTriggersOnCancel(t *testing.T) { +func TestBuiltAction_ExecuteDecoded(t *testing.T) { t.Parallel() - var canceled bool - - act := action.New("test.cancel", func(ctx context.Context, req int) (int, error) { - return 0, context.Canceled - }).Hook(action.Hook[int, int]{ - OnCancel: func(ctx context.Context, req int, m *action.Meta) { - canceled = true - }, - OnError: func(ctx context.Context, req int, err error, m *action.Meta) { - t.Fatal("OnError should not be called when context is canceled") - }, + + type RequestDto struct { + Name string `json:"name"` + } + type ResponseDto struct { + Greeting string `json:"greeting"` + } + + act := action.New("decode.test", func(ctx context.Context, req RequestDto) (ResponseDto, error) { + return ResponseDto{Greeting: "Hello, " + req.Name}, nil }).Build() - _, err := act.Do(context.Background(), 1) - if err == nil || !errors.Is(err, context.Canceled) { - t.Fatalf("expected context.Canceled error, got %v", err) + decodeFunc := func(target any) error { + req, ok := target.(*RequestDto) + if !ok { + return errors.New("invalid target type") + } + req.Name = "Tester" + return nil + } + + rawRes, err := act.ExecuteDecoded(context.Background(), decodeFunc) + if err != nil { + t.Fatalf("unexpected error: %v", err) } - if !canceled { - t.Fatal("expected OnCancel to be called") + res, ok := rawRes.(ResponseDto) + if !ok || res.Greeting != "Hello, Tester" { + t.Fatalf("unexpected result: %+v", rawRes) } } diff --git a/action/builder.go b/action/builder.go index 1bda099..de5ad1e 100644 --- a/action/builder.go +++ b/action/builder.go @@ -192,5 +192,8 @@ func (b *Builder[Req, Res]) Build() *BuiltAction[Req, Res] { func (b *Builder[Req, Res]) LogSlowWhen(d time.Duration) *Builder[Req, Res] { b.meta.LogSlowThreshold = d + if d > 0 { + return b.Use(SlowLogMiddleware[Req, Res](d, b.meta.Name)) + } return b } diff --git a/action/builder_observe.go b/action/builder_observe.go index 33ba87b..a292e76 100644 --- a/action/builder_observe.go +++ b/action/builder_observe.go @@ -104,3 +104,20 @@ func (b *Builder[Req, Res]) Instrument( } type instrumentStartKey struct{} + +func SlowLogMiddleware[Req, Res any](threshold time.Duration, actionName string) Middleware[Req, Res] { + return func(next Fn[Req, Res]) Fn[Req, Res] { + return func(ctx context.Context, req Req) (Res, error) { + start := time.Now() + res, err := next(ctx, req) + if dur := time.Since(start); dur > threshold { + slog.WarnContext(ctx, "action_slow_execution", + "action", actionName, + "duration", dur, + "threshold", threshold, + ) + } + return res, err + } + } +} diff --git a/action/cache.go b/action/cache.go index c6707c9..de7534f 100644 --- a/action/cache.go +++ b/action/cache.go @@ -17,72 +17,139 @@ type CacheLayer[V any] interface { type CacheConfig[Req, Res any] struct { KeyFunc func(Req) string - Layers []CacheLayer[Res] // L1 (In-Memory) -> L2 (NATS/Redis) + Layers []CacheLayer[Res] TTL time.Duration + Timeout time.Duration } -// CacheMiddleware implements the Read-Through / Write-Behind pattern +// CacheMiddleware implements the Read-Through / Write-Behind pattern with isolated singleflight execution. func CacheMiddleware[Req, Res any](cfg CacheConfig[Req, Res]) DispatcherMiddleware[Req, Res] { + if cfg.KeyFunc == nil { + panic("cache: KeyFunc must not be nil") + } + var sf singleflight.Group + return func(next Fn[Req, Res], hooks HookDispatcher[Req, Res]) Fn[Req, Res] { return func(ctx context.Context, req Req) (Res, error) { key := cfg.KeyFunc(req) if key == "" { return next(ctx, req) } - // Try layers L1 → L2 (order matters!) + + if err := ctx.Err(); err != nil { + var zero Res + return zero, err + } + + // 1. Try layers L1 -> LN in order for i, layer := range cfg.Layers { + if layer == nil { + continue + } + val, hit, err := layer.Get(ctx, key) if err != nil { - // Log infrastructure degradation but do not crash the request slog.WarnContext(ctx, "cache_layer_get_failed", "key", key, "layer_index", i, "error", err) continue } + if hit { hooks.OnCacheHit(ctx, req, val) - if i > 0 && len(cfg.Layers) > 0 { - if fillErr := cfg.Layers[0].Set(ctx, key, val, cfg.TTL); fillErr != nil { - slog.WarnContext(ctx, "cache_backfill_failed", - "key", key, "error", fillErr) + // Backfill all faster layers (0 to i-1) + if i > 0 { + for j := 0; j < i; j++ { + if cfg.Layers[j] == nil { + continue + } + if fillErr := cfg.Layers[j].Set(ctx, key, val, cfg.TTL); fillErr != nil { + slog.WarnContext(ctx, "cache_backfill_failed", + "key", key, "layer_index", j, "error", fillErr) + } } } return val, nil } } - hooks.OnCacheMiss(ctx, req) + // 2. Singleflight execution with a context owned by the flight, + // independent of any individual caller's cancellation. + baseCtx := context.WithoutCancel(ctx) + var executedByThisCaller bool + + ch := sf.DoChan(key, func() (result any, execErr error) { + executedByThisCaller = true + + flightTimeout := cfg.Timeout + if flightTimeout <= 0 { + flightTimeout = 2 * time.Minute + } + + execCtx, execCancel := context.WithTimeout(baseCtx, flightTimeout) + defer execCancel() - v, err, _ := sf.Do(key, func() (result any, execErr error) { defer func() { if r := recover(); r != nil { execErr = xerr.PanicRecovery(r) result = nil } }() - res, err := next(ctx, req) + + hooks.OnCacheMiss(execCtx, req) + + res, err := next(execCtx, req) if err == nil { - for _, layer := range cfg.Layers { - if writeErr := layer.Set(ctx, key, res, cfg.TTL); writeErr != nil { - slog.WarnContext(ctx, "cache_write_through_failed", - "key", key, "error", writeErr) + for idx, layer := range cfg.Layers { + if layer == nil { + continue + } + if writeErr := layer.Set(execCtx, key, res, cfg.TTL); writeErr != nil { + slog.WarnContext(execCtx, "cache_write_through_failed", + "key", key, "layer_index", idx, "error", writeErr) } } } return res, err }) - if err != nil { - var zero Res - return zero, err - } - res, ok := v.(Res) - if !ok { + // 3. Wait for result. + select { + case <-ctx.Done(): var zero Res - return zero, fmt.Errorf("cache: unexpected stored type %T", v) + return zero, ctx.Err() + + case result := <-ch: + // Required: even if the flight completed, the caller must still + // observe its own cancellation if its context was canceled. + if err := ctx.Err(); err != nil { + var zero Res + return zero, err + } + + // Only notify joining waiters. + if result.Shared && !executedByThisCaller { + hooks.OnCoalesced(ctx, req) + } + + if result.Err != nil { + var zero Res + return zero, result.Err + } + + if result.Val == nil { + var zero Res + return zero, nil + } + + res, ok := result.Val.(Res) + if !ok { + var zero Res + return zero, fmt.Errorf("cache: unexpected stored type %T", result.Val) + } + return res, nil } - return res, nil } } } diff --git a/action/cache_test.go b/action/cache_test.go index 307fbab..ec679ce 100644 --- a/action/cache_test.go +++ b/action/cache_test.go @@ -4,27 +4,34 @@ import ( "context" "errors" "fmt" + "runtime" "strings" "sync" + "sync/atomic" "testing" "time" "github.com/nexssp/kernel/action" + "github.com/nexssp/kernel/xerr" ) -// mockKVStore implements ports.KVStore[string, string] for testing. type mockKVStore struct { - mu sync.Mutex - store map[string]string - // simulate error on Set if needed + mu sync.Mutex + store map[string]string failSet bool + onGet func(key string) } -func newMockStore() *mockKVStore { return &mockKVStore{store: make(map[string]string)} } +func newMockStore() *mockKVStore { + return &mockKVStore{store: make(map[string]string)} +} func (m *mockKVStore) Get(_ context.Context, key string) (string, bool, error) { m.mu.Lock() defer m.mu.Unlock() + if m.onGet != nil { + m.onGet(key) + } v, ok := m.store[key] return v, ok, nil } @@ -49,23 +56,28 @@ func (m *mockKVStore) Delete(_ context.Context, key string) error { func TestCache_Hit(t *testing.T) { t.Parallel() store := newMockStore() - // Pre-populate the cache if err := store.Set(context.Background(), "key1", "cached", 0); err != nil { t.Fatalf("failed to pre-populate cache: %v", err) } var handlerCalled bool + var hitHookCalled bool + var executedHookCalled bool + act := action.New("cache.hit", func(ctx context.Context, req string) (string, error) { handlerCalled = true return "fresh", nil }). Cache(10*time.Minute, func(r string) string { return r }, store). HookCacheHit(func(ctx context.Context, req string, res string, meta *action.Meta) { - // verify we were notified + hitHookCalled = true if res != "cached" { t.Errorf("expected cached value 'cached', got %q", res) } }). + HookExecuted(func(ctx context.Context, req string, res string, meta *action.Meta) { + executedHookCalled = true + }). Build() res, err := act.Do(context.Background(), "key1") @@ -78,12 +90,20 @@ func TestCache_Hit(t *testing.T) { if handlerCalled { t.Fatal("handler should not have been called on cache hit") } + if !hitHookCalled { + t.Fatal("HookCacheHit must be called on cache hit") + } + if executedHookCalled { + t.Fatal("HookExecuted must be suppressed on cache hit") + } } func TestCache_Miss(t *testing.T) { t.Parallel() store := newMockStore() var handlerHit bool + var executedHookCalled bool + act := action.New("cache.miss", func(ctx context.Context, req string) (string, error) { handlerHit = true return "computed", nil @@ -94,6 +114,9 @@ func TestCache_Miss(t *testing.T) { t.Errorf("expected req 'missKey', got %q", req) } }). + HookExecuted(func(ctx context.Context, req string, res string, meta *action.Meta) { + executedHookCalled = true + }). Build() res, err := act.Do(context.Background(), "missKey") @@ -106,7 +129,10 @@ func TestCache_Miss(t *testing.T) { if !handlerHit { t.Fatal("handler should have been called on cache miss") } - // verify that the value was stored in the cache + if !executedHookCalled { + t.Fatal("HookExecuted must be called on cache miss (real execution)") + } + stored, ok, err := store.Get(context.Background(), "missKey") if err != nil { t.Fatalf("unexpected get error: %v", err) @@ -118,7 +144,6 @@ func TestCache_Miss(t *testing.T) { func TestCache_BackFill(t *testing.T) { t.Parallel() - // L1 and L2, L1 empty, L2 has value → should backfill L1 l1 := newMockStore() l2 := newMockStore() if err := l2.Set(context.Background(), "backfill", "from_l2", 0); err != nil { @@ -138,7 +163,7 @@ func TestCache_BackFill(t *testing.T) { if res != "from_l2" { t.Fatalf("expected 'from_l2', got %q", res) } - // L1 should now have the value + v, ok, err := l1.Get(context.Background(), "backfill") if err != nil { t.Fatalf("unexpected get error: %v", err) @@ -148,6 +173,184 @@ func TestCache_BackFill(t *testing.T) { } } +func TestCache_Singleflight_SingleMissNotification(t *testing.T) { + t.Parallel() + store := newMockStore() + + started := make(chan struct{}) + block := make(chan struct{}) + + var handlerCalls atomic.Int32 + var missCount atomic.Int32 + var coalesceCount atomic.Int32 + + act := action.New("cache.concurrent_miss", func(ctx context.Context, req string) (string, error) { + handlerCalls.Add(1) + close(started) + select { + case <-ctx.Done(): + return "", ctx.Err() + case <-block: + return "flight_value", nil + } + }). + Cache(10*time.Minute, func(r string) string { return r }, store). + HookCacheMiss(func(ctx context.Context, req string, meta *action.Meta) { + missCount.Add(1) + }). + Hook(action.Hook[string, string]{ + OnCoalesced: func(ctx context.Context, req string, meta *action.Meta) { + coalesceCount.Add(1) + }, + }). + Build() + + const concurrentCallers = 10 + var wg sync.WaitGroup + wg.Add(concurrentCallers) + + // 1. Uruchom pierwszego callera (lidera flightu) + go func() { + defer wg.Done() + res, err := act.Do(context.Background(), "shared_key") + if err != nil || res != "flight_value" { + t.Errorf("leader failed: res=%q, err=%v", res, err) + } + }() + + <-started // Czekamy aż lider zablokuje się wewnątrz handlera + + // 2. Uruchom 9 współbieżnych callerów, które uderzą w trwający flight + for i := 0; i < concurrentCallers-1; i++ { + go func() { + defer wg.Done() + res, err := act.Do(context.Background(), "shared_key") + if err != nil || res != "flight_value" { + t.Errorf("waiter failed: res=%q, err=%v", res, err) + } + }() + } + + // Dajemy wątkom wystartować i zablokować się na trwającym flighcie + time.Sleep(20 * time.Millisecond) + + // 3. Zwalniamy blokadę handlera – flight kończy się dla wszystkich + close(block) + wg.Wait() + + // Twarde niezmienniki biznesowe: + if got := handlerCalls.Load(); got != 1 { + t.Fatalf("CRITICAL: handler executed %d times, expected exactly 1", got) + } + if got := missCount.Load(); got != 1 { + t.Fatalf("CRITICAL: OnCacheMiss executed %d times, expected exactly 1", got) + } + if got := coalesceCount.Load(); got == 0 { + t.Fatalf("expected at least one coalesced caller, got %d", got) + } +} + +func TestCache_Singleflight_ContextCancellation(t *testing.T) { + t.Parallel() + store := newMockStore() + + started := make(chan struct{}) + block := make(chan struct{}) + caller2CheckingCache := make(chan struct{}) + var getCalls atomic.Int32 + + store.onGet = func(key string) { + if getCalls.Add(1) == 2 { + close(caller2CheckingCache) + } + } + + act := action.New("cache.concurrent_cancel", func(ctx context.Context, req string) (string, error) { + close(started) + select { + case <-ctx.Done(): + return "", ctx.Err() + case <-block: + return "computed_value", nil + } + }). + Cache(10*time.Minute, func(r string) string { return r }, store). + Build() + + ctx1, cancel1 := context.WithCancel(context.Background()) + ctx2 := context.Background() + + var res1, res2 string + var err1, err2 error + var wg sync.WaitGroup + wg.Add(2) + + go func() { + defer wg.Done() + res1, err1 = act.Do(ctx1, "shared_key") + }() + + <-started + + go func() { + defer wg.Done() + res2, err2 = act.Do(ctx2, "shared_key") + }() + + <-caller2CheckingCache + for i := 0; i < 5; i++ { + runtime.Gosched() + } + + cancel1() + + close(block) + wg.Wait() + + if !errors.Is(err1, context.Canceled) { + t.Fatalf("expected caller 1 to get context.Canceled, got: %v", err1) + } + if res1 != "" { + t.Fatalf("expected caller 1 to receive zero string, got: %q", res1) + } + + if err2 != nil { + t.Fatalf("caller 2 failed unexpectedly: %v", err2) + } + if res2 != "computed_value" { + t.Fatalf("caller 2 expected 'computed_value', got %q", res2) + } + + cached, ok, _ := store.Get(context.Background(), "shared_key") + if !ok || cached != "computed_value" { + t.Fatalf("expected value to be cached in store, got %q (ok=%v)", cached, ok) + } +} + +func TestCache_Singleflight_PanicRecovery(t *testing.T) { + t.Parallel() + store := newMockStore() + + act := action.New("cache.panic", func(ctx context.Context, req string) (string, error) { + panic("database crashed inside cache flight") + }). + Cache(10*time.Minute, func(r string) string { return r }, store). + Build() + + res, err := act.Do(context.Background(), "panic_key") + if res != "" { + t.Fatalf("expected empty string on panic, got %q", res) + } + if err == nil { + t.Fatal("expected error from panic, got nil") + } + + var appErr *xerr.AppError + if !errors.As(err, &appErr) || appErr.Kind != xerr.KindInternal { + t.Fatalf("expected KindInternal error, got %v", err) + } +} + func TestOnce(t *testing.T) { t.Parallel() var callCount int @@ -156,7 +359,6 @@ func TestOnce(t *testing.T) { return fmt.Sprintf("call-%d", callCount), nil }).Once().Build() - // Run multiple times concurrently var wg sync.WaitGroup results := make([]string, 10) for i := range 10 { @@ -176,7 +378,6 @@ func TestOnce(t *testing.T) { if callCount != 1 { t.Fatalf("expected handler to be called exactly once, got %d", callCount) } - // all results must be the same "call-1" for i, r := range results { if r != "call-1" { t.Fatalf("expected result 'call-1' at index %d, got %q", i, r) diff --git a/action/composer.go b/action/composer.go index ca2c8d1..2913cc9 100644 --- a/action/composer.go +++ b/action/composer.go @@ -161,10 +161,14 @@ func Chain[T any]( name string, builders ...*Builder[T, T], ) *Builder[T, T] { + acts := make([]*BuiltAction[T, T], len(builders)) + for i, b := range builders { + acts[i] = b.Build() + } + return New(name, func(ctx context.Context, req T) (T, error) { cur := req - for _, b := range builders { - act := b.Build() + for _, act := range acts { var err error cur, err = act.Do(ctx, cur) if err != nil { diff --git a/action/fenced_lock.go b/action/fenced_lock.go index 435fbf0..65c25af 100644 --- a/action/fenced_lock.go +++ b/action/fenced_lock.go @@ -3,6 +3,7 @@ package action import ( "context" "fmt" + "log/slog" "sync" "time" @@ -27,6 +28,18 @@ type FencedMutex interface { Release(ctx context.Context, lease LockLease) (released bool, err error) } +type leaseCtxKey struct{} + +// LeaseFromContext retrieves the active LockLease from the execution context. +func LeaseFromContext(ctx context.Context) (LockLease, bool) { + lease, ok := ctx.Value(leaseCtxKey{}).(LockLease) + return lease, ok +} + +const ( + minTTL = 300 * time.Millisecond +) + // ExclusiveFenced runs an action under an ownership-checked lease. The lease // is renewed while the action is running; loss of the lease cancels the action // context and returns an unavailable error rather than claiming success. @@ -34,18 +47,21 @@ func (b *Builder[Req, Res]) ExclusiveFenced(m FencedMutex, ttl time.Duration, ke return b.Use(func(next Fn[Req, Res]) Fn[Req, Res] { return func(ctx context.Context, req Req) (Res, error) { var zero Res + if m == nil { return zero, xerr.Internal("fenced mutex is required") } - if ttl <= 0 { - return zero, xerr.BadRequest("fenced lock TTL must be positive") + if ttl < minTTL { + return zero, xerr.BadRequest(fmt.Sprintf("fenced lock TTL must be at least %v", minTTL)) } key := keyFn(req) if key == "" { - return next(ctx, req) + return zero, xerr.BadRequest("fenced lock key cannot be empty") } + lockKey := b.meta.Name + ":lock:" + key + lease, acquired, err := m.Acquire(ctx, lockKey, ttl) if err != nil { return zero, xerr.Unavailable("failed to acquire fenced distributed lock", err) @@ -55,19 +71,41 @@ func (b *Builder[Req, Res]) ExclusiveFenced(m FencedMutex, ttl time.Duration, ke } execCtx, cancel := context.WithCancel(ctx) - defer cancel() + + // Inject the lease so downstream handlers can verify the fence token. + execCtx = context.WithValue(execCtx, leaseCtxKey{}, lease) + done := make(chan struct{}) lost := make(chan error, 1) var renewWG sync.WaitGroup renewWG.Add(1) + + reportLoss := func(cause error) { + select { + case lost <- cause: + default: + slog.Error("fenced_lock_additional_error", + "action", b.meta.Name, + "lock", lockKey, + "error", cause, + ) + } + cancel() + } + + interval := ttl / 3 + go func() { defer renewWG.Done() - interval := ttl / 3 - if interval < 100*time.Millisecond { - interval = 100 * time.Millisecond - } + defer func() { + if r := recover(); r != nil { + reportLoss(fmt.Errorf("fenced lock renew panic: %v", r)) + } + }() + ticker := time.NewTicker(interval) defer ticker.Stop() + for { select { case <-done: @@ -75,42 +113,88 @@ func (b *Builder[Req, Res]) ExclusiveFenced(m FencedMutex, ttl time.Duration, ke case <-execCtx.Done(): return case <-ticker.C: - renewed, renewErr := m.Renew(context.WithoutCancel(ctx), lease, ttl) - if renewErr != nil { - select { - case lost <- fmt.Errorf("renew fenced lock: %w", renewErr): - default: - } - cancel() + } + + renewCtx, renewCancel := context.WithTimeout(execCtx, interval) + renewed, renewErr := m.Renew(renewCtx, lease, ttl) + renewCancel() + + if renewErr != nil { + // If we're shutting down normally, this is not a lease loss. + select { + case <-done: return + default: + } + + if execCtx.Err() == nil { + reportLoss(fmt.Errorf("renew fenced lock: %w", renewErr)) } - if !renewed { - select { - case lost <- fmt.Errorf("fenced lock lease lost"): - default: - } - cancel() + return + } + + if !renewed { + select { + case <-done: return + default: } + if execCtx.Err() != nil { + return + } + reportLoss(fmt.Errorf("fenced lock lease lost")) + return } } }() - res, execErr := next(execCtx, req) - close(done) - renewWG.Wait() - - releaseCtx, releaseCancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) - _, releaseErr := m.Release(releaseCtx, lease) - releaseCancel() - if releaseErr != nil { - return zero, xerr.Unavailable("failed to release fenced distributed lock", releaseErr) + var releaseOnce sync.Once + cleanup := func() { + releaseOnce.Do(func() { + defer func() { + if r := recover(); r != nil { + slog.Error("fenced_lock_release_panic", + "action", b.meta.Name, + "lock", lockKey, + "panic", r, + ) + } + }() + close(done) + renewWG.Wait() + + releaseCtx, releaseCancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) + defer releaseCancel() + + if _, releaseErr := m.Release(releaseCtx, lease); releaseErr != nil { + slog.Error("fenced_lock_release_failed", + "action", b.meta.Name, + "lock", lockKey, + "owner", lease.Owner, + "error", releaseErr, + ) + } + }) } + + // Panic safety: cancel first to stop the renewer, then release. + defer func() { + cancel() + cleanup() + }() + + res, execErr := next(execCtx, req) + + // Normal path: stop the renewer before checking for lease loss. + cancel() + cleanup() + select { case lostErr := <-lost: return zero, xerr.Unavailable("fenced distributed lock lease lost", lostErr) default: } + return res, execErr } }) diff --git a/action/fenced_lock_test.go b/action/fenced_lock_test.go index 9132931..27aa1d3 100644 --- a/action/fenced_lock_test.go +++ b/action/fenced_lock_test.go @@ -2,28 +2,47 @@ package action import ( "context" + "errors" "sync" "testing" "time" ) type fakeFencedMutex struct { - mu sync.Mutex - lease LockLease - renewOK bool - released bool + mu sync.Mutex + lease LockLease + acquireOK bool + renewOK bool + released bool + err error +} + +func newFakeMutex() *fakeFencedMutex { + return &fakeFencedMutex{ + acquireOK: true, + renewOK: true, + } } func (m *fakeFencedMutex) Acquire(_ context.Context, key string, _ time.Duration) (LockLease, bool, error) { m.mu.Lock() defer m.mu.Unlock() - m.lease = LockLease{Key: key, Owner: "owner-1", Fence: 1} + if m.err != nil { + return LockLease{}, false, m.err + } + if !m.acquireOK { + return LockLease{}, false, nil + } + m.lease = LockLease{Key: key, Owner: "owner-1", Fence: 42} return m.lease, true, nil } func (m *fakeFencedMutex) Renew(_ context.Context, lease LockLease, _ time.Duration) (bool, error) { m.mu.Lock() defer m.mu.Unlock() + if m.err != nil { + return false, m.err + } return m.renewOK && lease == m.lease, nil } @@ -37,34 +56,144 @@ func (m *fakeFencedMutex) Release(_ context.Context, lease LockLease) (bool, err return true, nil } -func TestExclusiveFenced_ReleasesExactLeaseAfterSuccess(t *testing.T) { +func TestExclusiveFenced_SuccessAndContextPropagation(t *testing.T) { t.Parallel() - mutex := &fakeFencedMutex{renewOK: true} - built := New("invoice.settle", func(context.Context, string) (string, error) { + mutex := newFakeMutex() + + var capturedLease LockLease + built := New("invoice.settle", func(ctx context.Context, _ string) (string, error) { + var ok bool + capturedLease, ok = LeaseFromContext(ctx) + if !ok { + t.Fatal("lease not found in context") + } return "settled", nil - }).ExclusiveFenced(mutex, time.Second, func(req string) string { return req }).Build() + }).ExclusiveFenced(mutex, 300*time.Millisecond, func(req string) string { return req }).Build() - result, err := built.Do(context.Background(), "invoice-1") + result, err := built.Do(context.Background(), "inv-1") if err != nil || result != "settled" { - t.Fatalf("result = %q, %v", result, err) + t.Fatalf("result = %q, err = %v", result, err) } + + if capturedLease.Fence != 42 || capturedLease.Owner != "owner-1" { + t.Fatalf("unexpected lease: %+v", capturedLease) + } + + mutex.mu.Lock() + defer mutex.mu.Unlock() if !mutex.released { - t.Fatal("successful fenced action did not release its lease") + t.Fatal("successful action did not release its lease") } } func TestExclusiveFenced_CancelsWhenLeaseIsLost(t *testing.T) { t.Parallel() - mutex := &fakeFencedMutex{renewOK: false} + mutex := newFakeMutex() + mutex.renewOK = false + built := New("invoice.settle", func(ctx context.Context, _ string) (string, error) { <-ctx.Done() return "", ctx.Err() - }).ExclusiveFenced(mutex, 10*time.Millisecond, func(req string) string { return req }).Build() + }).ExclusiveFenced(mutex, 300*time.Millisecond, func(req string) string { return req }).Build() - if _, err := built.Do(context.Background(), "invoice-2"); err == nil { - t.Fatal("lost fenced lease unexpectedly returned success") + _, err := built.Do(context.Background(), "inv-2") + if err == nil { + t.Fatal("expected error on lost lease, got nil") } + + mutex.mu.Lock() + defer mutex.mu.Unlock() if !mutex.released { - t.Fatal("lost fenced action did not attempt release") + t.Fatal("lost lease action did not cleanup/release") + } +} + +func TestExclusiveFenced_LockContention(t *testing.T) { + t.Parallel() + mutex := newFakeMutex() + mutex.acquireOK = false + + built := New("invoice.settle", func(_ context.Context, _ string) (string, error) { + return "ok", nil + }).ExclusiveFenced(mutex, 300*time.Millisecond, func(req string) string { return req }).Build() + + _, err := built.Do(context.Background(), "inv-3") + if !errors.Is(err, ErrLocked) { + t.Fatalf("expected ErrLocked, got %v", err) + } +} + +func TestExclusiveFenced_ValidationErrors(t *testing.T) { + t.Parallel() + mutex := newFakeMutex() + + t.Run("empty key rejected", func(t *testing.T) { + built := New("invoice.settle", func(_ context.Context, _ string) (string, error) { + return "ok", nil + }).ExclusiveFenced(mutex, 300*time.Millisecond, func(_ string) string { return "" }).Build() + + if _, err := built.Do(context.Background(), ""); err == nil { + t.Fatal("expected error for empty key") + } + }) + + t.Run("short TTL rejected", func(t *testing.T) { + built := New("invoice.settle", func(_ context.Context, _ string) (string, error) { + return "ok", nil + }).ExclusiveFenced(mutex, 50*time.Millisecond, func(req string) string { return req }).Build() + + if _, err := built.Do(context.Background(), "inv-4"); err == nil { + t.Fatal("expected error for TTL < minTTL") + } + }) + + t.Run("nil mutex rejected", func(t *testing.T) { + built := New("invoice.settle", func(_ context.Context, _ string) (string, error) { + return "ok", nil + }).ExclusiveFenced(nil, 300*time.Millisecond, func(req string) string { return req }).Build() + + if _, err := built.Do(context.Background(), "inv-5"); err == nil { + t.Fatal("expected error for nil mutex") + } + }) +} + +func TestExclusiveFenced_PanicSafety(t *testing.T) { + t.Parallel() + mutex := newFakeMutex() + + built := New("invoice.settle", func(_ context.Context, _ string) (string, error) { + panic("action handler panic") + }).ExclusiveFenced(mutex, 300*time.Millisecond, func(req string) string { return req }).Build() + + _, err := built.Do(context.Background(), "inv-6") + if err == nil { + t.Fatal("expected error from recovered panic") + } + + mutex.mu.Lock() + defer mutex.mu.Unlock() + if !mutex.released { + t.Fatal("lock lease was not released when handler panicked") + } +} + +func TestLeaderOnlyFenced(t *testing.T) { + t.Parallel() + mutex := newFakeMutex() + + built := New("cluster.sync", func(ctx context.Context, _ string) (string, error) { + lease, ok := LeaseFromContext(ctx) + if !ok { + t.Fatal("lease not found in context") + } + if lease.Key != "cluster.sync:lock:global_leader" { + t.Fatalf("unexpected lock key: %s", lease.Key) + } + return "ok", nil + }).LeaderOnlyFenced(mutex, 300*time.Millisecond).Build() + + if _, err := built.Do(context.Background(), ""); err != nil { + t.Fatalf("unexpected error: %v", err) } } diff --git a/action/idempotency.go b/action/idempotency.go index 2b5741d..14bdb8a 100644 --- a/action/idempotency.go +++ b/action/idempotency.go @@ -2,7 +2,6 @@ package action import ( "context" - "sync" "time" ) @@ -44,7 +43,15 @@ func (c IdempotencyConfig) Header() string { return "Idempotency-Key" } -// ── Store ───────────────────────────────────────────────────────────────────── +// EffectiveLeaseTTL returns the configured lease or the documented default. +func (c IdempotencyConfig) EffectiveLeaseTTL() time.Duration { + if c.LeaseTTL > 0 { + return c.LeaseTTL + } + return DefaultIdempotencyLeaseTTL +} + +// ── Store Contracts ─────────────────────────────────────────────────────────── // IdempotencyEntry is the captured response for a completed idempotent request. type IdempotencyEntry struct { @@ -97,100 +104,3 @@ type IdempotencyCoordinator interface { Complete(ctx context.Context, key, token string, entry IdempotencyEntry, ttl time.Duration) error Release(ctx context.Context, key, token string) error } - -// EffectiveLeaseTTL returns the configured lease or the documented default. -func (c IdempotencyConfig) EffectiveLeaseTTL() time.Duration { - if c.LeaseTTL > 0 { - return c.LeaseTTL - } - return DefaultIdempotencyLeaseTTL -} - -// ── In-memory default ───────────────────────────────────────────────────────── - -type memEntry struct { - IdempotencyEntry - ttl time.Duration -} - -type MemoryIdempotencyStore struct { - mu sync.RWMutex - entries map[string]memEntry - defTTL time.Duration - stopCh chan struct{} - - closeOnce sync.Once -} - -// NewMemoryIdempotencyStore creates a store with background TTL eviction. -// defTTL 0 → 24 h. -func NewMemoryIdempotencyStore(defTTL time.Duration) *MemoryIdempotencyStore { - if defTTL == 0 { - defTTL = DefaultIdempotencyTTL - } - s := &MemoryIdempotencyStore{ - entries: make(map[string]memEntry), - defTTL: defTTL, - stopCh: make(chan struct{}), // new field - } - go s.evict() - return s -} - -// Close stops the eviction goroutine. Safe to call multiple times. -func (s *MemoryIdempotencyStore) Close() { - s.closeOnce.Do(func() { - close(s.stopCh) - }) -} - -func (s *MemoryIdempotencyStore) Get(_ context.Context, key string) (IdempotencyEntry, bool) { - s.mu.RLock() - e, ok := s.entries[key] - s.mu.RUnlock() - if !ok { - return IdempotencyEntry{}, false - } - ttl := e.ttl - if ttl == 0 { - ttl = s.defTTL - } - if time.Since(e.StoredAt) > ttl { - // Lazy eviction to prevent unbounded map growth between janitor ticks - s.mu.Lock() - delete(s.entries, key) - s.mu.Unlock() - return IdempotencyEntry{}, false - } - return e.IdempotencyEntry, true -} - -func (s *MemoryIdempotencyStore) Set(_ context.Context, key string, entry IdempotencyEntry, ttl time.Duration) { - s.mu.Lock() - s.entries[key] = memEntry{IdempotencyEntry: entry, ttl: ttl} - s.mu.Unlock() -} - -func (s *MemoryIdempotencyStore) evict() { - t := time.NewTicker(time.Hour) - defer t.Stop() - for { - select { - case <-s.stopCh: - return // clean exit - case <-t.C: - now := time.Now() - s.mu.Lock() - for k, e := range s.entries { - ttl := e.ttl - if ttl == 0 { - ttl = s.defTTL - } - if now.Sub(e.StoredAt) > ttl { - delete(s.entries, k) - } - } - s.mu.Unlock() - } - } -} diff --git a/action/idempotency_default.go b/action/idempotency_default.go new file mode 100644 index 0000000..778caa1 --- /dev/null +++ b/action/idempotency_default.go @@ -0,0 +1,232 @@ +package action + +import ( + "context" + "crypto/rand" + "encoding/hex" + "errors" + "fmt" + "sync" + "time" +) + +type memEntry struct { + IdempotencyEntry + ttl time.Duration + seq uint64 +} + +type memClaim struct { + token string + requestHash string + expiresAt time.Time +} + +var _ IdempotencyStore = (*MemoryIdempotencyStore)(nil) +var _ IdempotencyCoordinator = (*MemoryIdempotencyStore)(nil) + +// MemoryIdempotencyStore provides an in-memory implementation of IdempotencyCoordinator +// with background TTL eviction and atomic lease claims. +type MemoryIdempotencyStore struct { + mu sync.RWMutex + entries map[string]memEntry + claims map[string]memClaim + defTTL time.Duration + stopCh chan struct{} + closeOnce sync.Once + nextSeq uint64 +} + +// NewMemoryIdempotencyStore creates a store with background TTL eviction. +// defTTL 0 → 24 h. +func NewMemoryIdempotencyStore(defTTL time.Duration) *MemoryIdempotencyStore { + if defTTL == 0 { + defTTL = DefaultIdempotencyTTL + } + s := &MemoryIdempotencyStore{ + entries: make(map[string]memEntry), + claims: make(map[string]memClaim), + defTTL: defTTL, + stopCh: make(chan struct{}), + } + go s.evict() + return s +} + +// Close stops the eviction goroutine. Safe to call multiple times. +func (s *MemoryIdempotencyStore) Close() { + s.closeOnce.Do(func() { + close(s.stopCh) + }) +} + +func (s *MemoryIdempotencyStore) Get(_ context.Context, key string) (IdempotencyEntry, bool) { + s.mu.RLock() + e, ok := s.entries[key] + s.mu.RUnlock() + + if !ok { + return IdempotencyEntry{}, false + } + + ttl := e.ttl + if ttl == 0 { + ttl = s.defTTL + } + + if time.Since(e.StoredAt) > ttl { + s.mu.Lock() + // Double-checked eviction: only delete if this exact sequence entry is still in place. + if cur, ok := s.entries[key]; ok && cur.seq == e.seq { + delete(s.entries, key) + } + s.mu.Unlock() + return IdempotencyEntry{}, false + } + + return e.IdempotencyEntry, true +} + +func (s *MemoryIdempotencyStore) setLocked(key string, entry IdempotencyEntry, ttl time.Duration) { + if entry.StoredAt.IsZero() { + entry.StoredAt = time.Now() + } + + s.nextSeq++ + s.entries[key] = memEntry{ + IdempotencyEntry: entry, + ttl: ttl, + seq: s.nextSeq, + } +} + +func (s *MemoryIdempotencyStore) Set(_ context.Context, key string, entry IdempotencyEntry, ttl time.Duration) { + s.mu.Lock() + defer s.mu.Unlock() + s.setLocked(key, entry, ttl) +} + +func (s *MemoryIdempotencyStore) Claim(ctx context.Context, key, requestHash string, leaseTTL time.Duration) (IdempotencyClaim, error) { + if err := ctx.Err(); err != nil { + return IdempotencyClaim{}, err + } + if leaseTTL <= 0 { + leaseTTL = DefaultIdempotencyLeaseTTL + } + now := time.Now() + + s.mu.Lock() + defer s.mu.Unlock() + + // 1. Check existing completed entry + if entry, ok := s.entries[key]; ok { + ttl := entry.ttl + if ttl == 0 { + ttl = s.defTTL + } + if now.After(entry.StoredAt.Add(ttl)) { + delete(s.entries, key) + } else { + if requestHash != "" && entry.RequestHash != "" && requestHash != entry.RequestHash { + return IdempotencyClaim{State: IdempotencyClaimConflict}, nil + } + return IdempotencyClaim{State: IdempotencyClaimCompleted, Entry: entry.IdempotencyEntry}, nil + } + } + + // 2. Check active in-progress lease + if claim, ok := s.claims[key]; ok { + if now.Before(claim.expiresAt) { + if requestHash != "" && claim.requestHash != "" && requestHash != claim.requestHash { + return IdempotencyClaim{State: IdempotencyClaimConflict}, nil + } + return IdempotencyClaim{State: IdempotencyClaimInProgress}, nil + } + delete(s.claims, key) + } + + // 3. Issue new atomic lease claim + token, err := newClaimToken() + if err != nil { + return IdempotencyClaim{}, err + } + + s.claims[key] = memClaim{ + token: token, + requestHash: requestHash, + expiresAt: now.Add(leaseTTL), + } + + return IdempotencyClaim{State: IdempotencyClaimAcquired, Token: token}, nil +} + +func (s *MemoryIdempotencyStore) Complete(_ context.Context, key, token string, entry IdempotencyEntry, ttl time.Duration) error { + s.mu.Lock() + defer s.mu.Unlock() + + claim, ok := s.claims[key] + if !ok { + return errors.New("idempotency: no active claim found for key") + } + if claim.token != token { + return errors.New("idempotency: token does not match active claim") + } + + delete(s.claims, key) + s.setLocked(key, entry, ttl) + return nil +} + +func (s *MemoryIdempotencyStore) Release(_ context.Context, key, token string) error { + s.mu.Lock() + defer s.mu.Unlock() + + claim, ok := s.claims[key] + if !ok { + // Claim was already completed or released + return nil + } + if claim.token != token { + return errors.New("idempotency: token does not match active claim") + } + + delete(s.claims, key) + return nil +} + +func (s *MemoryIdempotencyStore) evict() { + t := time.NewTicker(time.Hour) + defer t.Stop() + for { + select { + case <-s.stopCh: + return + case <-t.C: + now := time.Now() + s.mu.Lock() + for k, e := range s.entries { + ttl := e.ttl + if ttl == 0 { + ttl = s.defTTL + } + if now.Sub(e.StoredAt) > ttl { + delete(s.entries, k) + } + } + for k, c := range s.claims { + if now.After(c.expiresAt) { + delete(s.claims, k) + } + } + s.mu.Unlock() + } + } +} + +func newClaimToken() (string, error) { + var b [16]byte + if _, err := rand.Read(b[:]); err != nil { + return "", fmt.Errorf("idempotency: failed generating claim token: %w", err) + } + return hex.EncodeToString(b[:]), nil +} diff --git a/action/idempotency_internal_test.go b/action/idempotency_internal_test.go index 80c0e15..ad16706 100644 --- a/action/idempotency_internal_test.go +++ b/action/idempotency_internal_test.go @@ -37,3 +37,40 @@ func TestMemoryIdempotencyStore_EagerEviction(t *testing.T) { t.Fatalf("expected exactly 1 live entry, got %d", count) } } + +func TestMemoryIdempotencyStore_PhantomDeletePrevention(t *testing.T) { + t.Parallel() + + store := NewMemoryIdempotencyStore(time.Hour) + defer store.Close() + + ctx := context.Background() + + // 1. Setup: Inject a stale entry + staleStoredAt := time.Now().Add(-2 * time.Hour) + store.mu.Lock() + store.nextSeq++ + store.entries["test-key"] = memEntry{ + IdempotencyEntry: IdempotencyEntry{Status: 200, Body: []byte("stale"), StoredAt: staleStoredAt}, + ttl: 10 * time.Millisecond, + seq: store.nextSeq, + } + staleSeq := store.nextSeq + store.mu.Unlock() + + // 2. Writer sets a fresh entry with a new sequence number + store.Set(ctx, "test-key", IdempotencyEntry{Status: 200, Body: []byte("fresh"), StoredAt: time.Now()}, time.Hour) + + // 3. Simulate stale eviction attempt with the old sequence + store.mu.Lock() + if cur, ok := store.entries["test-key"]; ok && cur.seq == staleSeq { + delete(store.entries, "test-key") // MUST NOT EXECUTE + } + store.mu.Unlock() + + // 4. Verify fresh entry remains untouched + entry, found := store.Get(ctx, "test-key") + if !found || string(entry.Body) != "fresh" { + t.Fatalf("phantom delete occurred: expected 'fresh' entry to persist, found=%v", found) + } +} diff --git a/action/zero_alloc_test.go b/action/zero_alloc_test.go new file mode 100644 index 0000000..9fc70e6 --- /dev/null +++ b/action/zero_alloc_test.go @@ -0,0 +1,57 @@ +package action_test + +import ( + "context" + "testing" + + "github.com/nexssp/kernel/action" +) + +func newZeroAllocAction() *action.BuiltAction[int, int] { + return action.New("zeroalloc.action", func(ctx context.Context, req int) (int, error) { + return req * 2, nil + }). + Tag("bench"). + HookBefore(func(ctx context.Context, _ int, _ *action.Meta) (context.Context, error) { + return ctx, nil + }). + Build() +} + +func newZeroAllocAnyHookAction() *action.BuiltAction[int, int] { + return action.New("zeroalloc.anyhook", func(ctx context.Context, req int) (int, error) { + return req * 2, nil + }). + AnyHook(action.AnyHook{}). + Build() +} + +func TestActionDispatchZeroAlloc(t *testing.T) { + // Nie używamy t.Parallel(), aby inne goroutines nie zaburzyły pomiaru alokacji. + act := newZeroAllocAction() + + allocations := testing.AllocsPerRun(1000, func() { + if _, err := act.Do(context.Background(), 42); err != nil { + t.Fatal(err) + } + }) + + if allocations != 0 { + t.Fatalf("expected 0 allocations, got %.1f", allocations) + } +} + +func TestActionDispatchZeroAlloc_AnyHookSnapshot(t *testing.T) { + // Nie używamy t.Parallel(), aby inne goroutines nie zaburzyły pomiaru alokacji. + act := newZeroAllocAnyHookAction() + + allocations := testing.AllocsPerRun(1000, func() { + if _, err := act.Do(context.Background(), 42); err != nil { + t.Fatal(err) + } + }) + + if allocations != 0 { + t.Fatalf("expected 0 allocations, got %.1f", allocations) + } +}