saga
packageAPI reference for the saga
package.
Imports
(14)RetryPolicy
RetryPolicy defines retry behavior for saga steps.
type RetryPolicy struct
Fields
| Name | Type | Description |
|---|---|---|
| MaxAttempts | int | |
| Delay | time.Duration | |
| Multiplier | float64 |
WithRetry
WithRetry wraps a function with retry logic according to the given policy.
Parameters
Returns
func WithRetry(policy RetryPolicy, do func(context.Context) error) func(context.Context) error
{
return func(ctx context.Context) error {
return resiliency.Retry(ctx, func() error {
return do(ctx)
},
resiliency.WithAttempts(policy.MaxAttempts),
resiliency.WithDelay(policy.Delay, 24*time.Hour),
resiliency.WithFactor(policy.Multiplier),
)
}
}
Uses
StepStatus
StepStatus represents the status of a saga step.
type StepStatus string
SagaState
SagaState holds the persisted state of a saga.
type SagaState struct
Fields
| Name | Type | Description |
|---|---|---|
| ID | string | json:"id" |
| Steps | []StepState | json:"steps" |
| Status | StepStatus | json:"status" |
| Error | string | json:"error,omitempty" |
| IdempotencyKey | string | json:"idempotency_key,omitempty" |
| RetryCount | int | json:"retry_count,omitempty" |
| MaxRetries | int | json:"max_retries,omitempty" |
| CreatedAt | time.Time | json:"created_at" |
| UpdatedAt | time.Time | json:"updated_at" |
Uses
StepState
StepState holds the persisted state of a single saga step.
type StepState struct
Fields
| Name | Type | Description |
|---|---|---|
| Name | string | json:"name" |
| Status | StepStatus | json:"status" |
| StepIndex | int | json:"step_index" |
Uses
SagaStore
SagaStore is the interface for saga state persistence.
type SagaStore interface
MemoryStore
MemoryStore is an in-memory saga store for testing and single-process use.
type MemoryStore struct
Methods
Parameters
Returns
func (*MemoryStore) Save(state *SagaState) error
{
if state == nil {
return errors.New("saga memory store: state cannot be nil")
}
m.mu.Lock()
defer m.mu.Unlock()
if _, exists := m.sagas[state.ID]; !exists && len(m.sagas) >= m.capacity {
return errors.New("saga memory store: capacity reached")
}
cp := cloneSagaState(state)
cp.UpdatedAt = time.Now()
if cp.CreatedAt.IsZero() {
cp.CreatedAt = time.Now()
}
m.sagas[state.ID] = cp
return nil
}
Parameters
Returns
func (*MemoryStore) Load(id string) (*SagaState, error)
{
m.mu.RLock()
defer m.mu.RUnlock()
s, ok := m.sagas[id]
if !ok {
return nil, fmt.Errorf("%w: %q", ErrSagaNotFound, id)
}
return cloneSagaState(s), nil
}
Parameters
Returns
func (*MemoryStore) Delete(id string) error
{
m.mu.Lock()
defer m.mu.Unlock()
delete(m.sagas, id)
return nil
}
Returns
func (*MemoryStore) ListIncomplete() ([]*SagaState, error)
{
m.mu.RLock()
defer m.mu.RUnlock()
var result []*SagaState
for _, s := range m.sagas {
if s.Status == StatusPending || s.Status == StatusFailed {
result = append(result, cloneSagaState(s))
}
}
return result, nil
}
Fields
| Name | Type | Description |
|---|---|---|
| mu | sync.RWMutex | |
| sagas | map[string]*SagaState | |
| capacity | int |
NewMemoryStore
NewMemoryStore creates a new in-memory saga store.
Parameters
Returns
func NewMemoryStore(capacity ...int) *MemoryStore
{
limit := 10_000
if len(capacity) > 0 && capacity[0] > 0 {
limit = capacity[0]
}
return &MemoryStore{sagas: make(map[string]*SagaState), capacity: limit}
}
IdempotencyRecorder
IdempotencyRecorder tracks processed idempotency keys to prevent double execution.
type IdempotencyRecorder struct
Methods
Parameters
Returns
func (*IdempotencyRecorder) Record(key string, state *SagaState) error
{
if key == "" || state == nil {
return errors.New("saga idempotency recorder: key and state are required")
}
r.mu.Lock()
defer r.mu.Unlock()
if existing, ok := r.seen[key]; ok {
return &IdempotencyError{Key: key, ExistingState: cloneSagaState(existing)}
}
if len(r.seen) >= r.capacity {
return errors.New("saga idempotency recorder: capacity reached")
}
r.seen[key] = cloneSagaState(state)
return nil
}
Fields
| Name | Type | Description |
|---|---|---|
| mu | sync.RWMutex | |
| seen | map[string]*SagaState | |
| capacity | int |
NewIdempotencyRecorder
NewIdempotencyRecorder creates a new IdempotencyRecorder.
Parameters
Returns
func NewIdempotencyRecorder(capacity ...int) *IdempotencyRecorder
{
limit := 10_000
if len(capacity) > 0 && capacity[0] > 0 {
limit = capacity[0]
}
return &IdempotencyRecorder{seen: make(map[string]*SagaState), capacity: limit}
}
IdempotencyError
IdempotencyError is returned when an idempotency key has already been processed.
type IdempotencyError struct
Methods
Fields
| Name | Type | Description |
|---|---|---|
| Key | string | |
| ExistingState | *SagaState |
DeadLetterQueue
DeadLetterQueue holds failed saga states that exceeded max retries.
type DeadLetterQueue struct
Methods
Parameters
Returns
func (*DeadLetterQueue) Enqueue(state *SagaState, reason string) error
{
if state == nil {
return errors.New("saga dead letter queue: state cannot be nil")
}
dlq.mu.Lock()
defer dlq.mu.Unlock()
if len(dlq.items) >= dlq.capacity {
return errors.New("saga dead letter queue: capacity reached")
}
dlq.items = append(dlq.items, &DLQEntry{
State: cloneSagaState(state),
Reason: reason,
DeadSince: time.Now(),
})
return nil
}
Returns
func (*DeadLetterQueue) List() []*DLQEntry
{
dlq.mu.RLock()
defer dlq.mu.RUnlock()
result := make([]*DLQEntry, 0, len(dlq.items))
for _, entry := range dlq.items {
copyEntry := *entry
copyEntry.State = cloneSagaState(entry.State)
result = append(result, ©Entry)
}
return result
}
Returns
func (*DeadLetterQueue) Len() int
{
dlq.mu.RLock()
defer dlq.mu.RUnlock()
return len(dlq.items)
}
Parameters
func (*DeadLetterQueue) Remove(id string)
{
dlq.mu.Lock()
defer dlq.mu.Unlock()
for i, entry := range dlq.items {
if entry.State.ID == id {
dlq.items = append(dlq.items[:i], dlq.items[i+1:]...)
return
}
}
}
Fields
| Name | Type | Description |
|---|---|---|
| mu | sync.RWMutex | |
| items | []*DLQEntry | |
| capacity | int |
NewDeadLetterQueue
NewDeadLetterQueue creates a new DeadLetterQueue.
Parameters
Returns
func NewDeadLetterQueue(capacity ...int) *DeadLetterQueue
{
limit := 10_000
if len(capacity) > 0 && capacity[0] > 0 {
limit = capacity[0]
}
return &DeadLetterQueue{items: make([]*DLQEntry, 0), capacity: limit}
}
sagaRunLock
type sagaRunLock struct
Fields
| Name | Type | Description |
|---|---|---|
| mu | sync.Mutex | |
| refs | int |
sagaRunLocks
type sagaRunLocks struct
Methods
Parameters
Returns
func (*sagaRunLocks) acquire(id string) func()
{
l.mu.Lock()
if l.locks == nil {
l.locks = make(map[string]*sagaRunLock)
}
lock := l.locks[id]
if lock == nil {
lock = &sagaRunLock{}
l.locks[id] = lock
}
lock.refs++
l.mu.Unlock()
lock.mu.Lock()
return func() {
lock.mu.Unlock()
l.mu.Lock()
lock.refs--
if lock.refs == 0 {
delete(l.locks, id)
}
l.mu.Unlock()
}
}
Fields
| Name | Type | Description |
|---|---|---|
| mu | sync.Mutex | |
| locks | map[string]*sagaRunLock |
ProcessWithDLQ
ProcessWithDLQ retries dead sagas up to maxRetries, then moves to dead state.
Parameters
Returns
func ProcessWithDLQ(store SagaStore, dlq *DeadLetterQueue, maxRetries int) func(id string, runFunc func(ctx context.Context) error) func(ctx context.Context) error
{
var locks sagaRunLocks
return func(id string, runFunc func(ctx context.Context) error) func(ctx context.Context) error {
return func(ctx context.Context) error {
release := locks.acquire(id)
defer release()
state, err := store.Load(id)
if err != nil {
return err
}
if state.Status == StatusCompleted {
return nil
}
if state.Status == StatusDead {
return fmt.Errorf("saga %q is dead", id)
}
if state.RetryCount >= maxRetries {
state.Status = StatusDead
state.Error = errors.New("exceeded max retries").Error()
if err := store.Save(state); err != nil {
return fmt.Errorf("saga %q persist dead state: %w", id, err)
}
if err := dlq.Enqueue(state, "max retries exceeded"); err != nil {
return fmt.Errorf("saga %q enqueue dead state: %w", id, err)
}
return fmt.Errorf("saga %q dead: max retries exceeded", id)
}
state.RetryCount++
if err := callSagaFunction(ctx, runFunc); err != nil {
state.Status = StatusFailed
state.Error = err.Error()
if saveErr := store.Save(state); saveErr != nil {
return errors.Join(err, fmt.Errorf("saga %q persist failure: %w", id, saveErr))
}
return err
}
state.Status = StatusCompleted
if err := store.Save(state); err != nil {
return fmt.Errorf("saga %q persist completion: %w", id, err)
}
return nil
}
}
}
Uses
callSagaFunction
Parameters
Returns
func callSagaFunction(ctx context.Context, fn func(context.Context) error) (err error)
{
if fn == nil {
return errors.New("saga: run function cannot be nil")
}
defer func() {
if recovered := recover(); recovered != nil {
err = fmt.Errorf("saga: run function panic: %v", recovered)
}
}()
return fn(ctx)
}
FileStore
FileStore persists saga state to disk.
type FileStore struct
Methods
Parameters
Returns
func (*FileStore) Save(state *SagaState) error
{
s.mu.Lock()
defer s.mu.Unlock()
if state == nil {
return errors.New("saga store: state cannot be nil")
}
if err := validateStoreID(state.ID); err != nil {
return err
}
state.UpdatedAt = time.Now()
if state.CreatedAt.IsZero() {
state.CreatedAt = time.Now()
}
data, err := json.MarshalIndent(state, "", " ")
if err != nil {
return fmt.Errorf("saga store: marshal: %w", err)
}
if len(data) > maxSagaStateSize {
return errors.New("saga store: state exceeds size limit")
}
path := filepath.Join(s.dir, state.ID+".json")
tmp, err := os.CreateTemp(s.dir, ".saga-*.tmp")
if err != nil {
return fmt.Errorf("saga store: create temp: %w", err)
}
tmpName := tmp.Name()
defer os.Remove(tmpName)
if err := tmp.Chmod(0600); err != nil {
tmp.Close()
return fmt.Errorf("saga store: secure temp: %w", err)
}
if _, err := tmp.Write(data); err != nil {
tmp.Close()
return fmt.Errorf("saga store: write: %w", err)
}
if err := tmp.Close(); err != nil {
return fmt.Errorf("saga store: close: %w", err)
}
if err := os.Rename(tmpName, path); err != nil {
return fmt.Errorf("saga store: rename: %w", err)
}
return nil
}
Parameters
Returns
func (*FileStore) Load(id string) (*SagaState, error)
{
s.mu.RLock()
defer s.mu.RUnlock()
if err := validateStoreID(id); err != nil {
return nil, err
}
data, err := readStoreFile(s.dir, id+".json", maxSagaStateSize)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return nil, fmt.Errorf("%w: %q", ErrSagaNotFound, id)
}
return nil, fmt.Errorf("saga store: read %s: %w", id, err)
}
var state SagaState
if err := json.Unmarshal(data, &state); err != nil {
return nil, fmt.Errorf("saga store: unmarshal %s: %w", id, err)
}
return &state, nil
}
Parameters
Returns
func (*FileStore) Delete(id string) error
{
s.mu.Lock()
defer s.mu.Unlock()
if err := validateStoreID(id); err != nil {
return err
}
root, err := os.OpenRoot(s.dir)
if err != nil {
return err
}
defer root.Close()
name := id + ".json"
info, err := root.Lstat(name)
if err != nil {
return err
}
if info.Mode()&os.ModeSymlink != 0 {
return errors.New("saga store: refusing to delete symbolic link")
}
return root.Remove(name)
}
Returns
func (*FileStore) ListIncomplete() ([]*SagaState, error)
{
s.mu.RLock()
defer s.mu.RUnlock()
entries, err := os.ReadDir(s.dir)
if err != nil {
return nil, fmt.Errorf("saga store: readdir: %w", err)
}
var result []*SagaState
var readErrors []error
for _, entry := range entries {
if filepath.Ext(entry.Name()) != ".json" || entry.Type()&os.ModeSymlink != 0 {
continue
}
data, err := readStoreFile(s.dir, entry.Name(), maxSagaStateSize)
if err != nil {
readErrors = append(readErrors, fmt.Errorf("saga store: read %s: %w", entry.Name(), err))
continue
}
var state SagaState
if err := json.Unmarshal(data, &state); err != nil {
readErrors = append(readErrors, fmt.Errorf("saga store: unmarshal %s: %w", entry.Name(), err))
continue
}
if state.Status == StatusPending || state.Status == StatusFailed {
result = append(result, &state)
}
}
return result, errors.Join(readErrors...)
}
Fields
| Name | Type | Description |
|---|---|---|
| dir | string | |
| mu | sync.RWMutex |
NewStore
NewStore creates a FileStore that writes state JSON files to dir.
Parameters
Returns
func NewStore(dir string) (*FileStore, error)
{
if err := os.MkdirAll(dir, 0700); err != nil {
return nil, fmt.Errorf("saga store: cannot create dir %s: %w", dir, err)
}
if err := os.Chmod(dir, 0700); err != nil {
return nil, fmt.Errorf("saga store: cannot secure dir %s: %w", dir, err)
}
return &FileStore{dir: dir}, nil
}
validateStoreID
Parameters
Returns
func validateStoreID(id string) error
{
if id == "" || id == "." || id == ".." ||
strings.ContainsAny(id, `/\`) {
return fmt.Errorf("saga store: invalid ID %q", id)
}
for _, character := range id {
if character == 0 {
return fmt.Errorf("saga store: invalid ID %q", id)
}
}
return nil
}
readStoreFile
Parameters
Returns
func readStoreFile(dir, name string, limit int64) ([]byte, error)
{
root, err := os.OpenRoot(dir)
if err != nil {
return nil, err
}
defer root.Close()
info, err := root.Lstat(name)
if err != nil {
return nil, err
}
if info.Mode()&os.ModeSymlink != 0 {
return nil, errors.New("saga store: refusing to read symbolic link")
}
file, err := root.Open(name)
if err != nil {
return nil, err
}
defer file.Close()
data, err := io.ReadAll(io.LimitReader(file, limit+1))
if err != nil {
return nil, err
}
if int64(len(data)) > limit {
return nil, errors.New("saga store: state file exceeds size limit")
}
return data, nil
}
RecoverableWorkflow
RecoverableWorkflow is a saga that persists state for crash recovery.
type RecoverableWorkflow struct
Methods
Parameters
func (*RecoverableWorkflow) Add(name string, do, compensate func(ctx context.Context) error)
{
rw.stateMu.Lock()
idx := len(rw.state.Steps)
rw.state.Steps = append(rw.state.Steps, StepState{
Name: name,
Status: StatusPending,
StepIndex: idx,
})
rw.stateMu.Unlock()
rw.Workflow.Add(name, do, compensate)
rw.indexes = append(rw.indexes, []int{idx})
}
AddGroup appends a parallel group and tracks each step independently.
Parameters
func (*RecoverableWorkflow) AddGroup(group Group)
{
rw.stateMu.Lock()
indexes := make([]int, 0, len(group))
for _, step := range group {
index := len(rw.state.Steps)
rw.state.Steps = append(rw.state.Steps, StepState{
Name: step.Name,
Status: StatusPending,
StepIndex: index,
})
indexes = append(indexes, index)
}
rw.stateMu.Unlock()
rw.Workflow.AddGroup(group)
rw.indexes = append(rw.indexes, indexes)
}
Parameters
Returns
func (*RecoverableWorkflow) Run(ctx context.Context) error
{
if rw.store == nil {
return errors.New("saga: recoverable workflow requires a store")
}
if !rw.running.CompareAndSwap(false, true) {
return errors.New("saga: workflow is already running")
}
defer rw.running.Store(false)
if err := rw.restoreState(); err != nil {
return err
}
rw.stateMu.Lock()
alreadyCompleted := rw.state.Status == StatusCompleted
rw.stateMu.Unlock()
if alreadyCompleted {
return nil
}
completed := rw.persistedCompletedSteps()
var completedMu sync.Mutex
if err := rw.updateAndSave(func(state *SagaState) {
state.Status = StatusPending
state.Error = ""
for index := range state.Steps {
if state.Steps[index].Status != StatusCompleted {
state.Steps[index].Status = StatusPending
}
}
}); err != nil {
return fmt.Errorf("saga %q persist pending state: %w", rw.id, err)
}
for i, item := range rw.steps {
if rw.stepsCompleted(rw.indexes[i]) {
continue
}
if ctx.Err() != nil {
saveErr := rw.updateAndSave(func(state *SagaState) {
state.Status = StatusFailed
state.Error = ctx.Err().Error()
})
return errors.Join(rw.rollback(ctx, ctx.Err(), completed), saveErr)
}
var err error
switch v := item.(type) {
case Step:
err = rw.runStepTracking(ctx, v, rw.indexes[i][0], &completed, &completedMu)
case Group:
var pending Group
var pendingIndexes []int
for groupIndex, stateIndex := range rw.indexes[i] {
if rw.stepsCompleted([]int{stateIndex}) {
continue
}
pending = append(pending, v[groupIndex])
pendingIndexes = append(pendingIndexes, stateIndex)
}
err = rw.runGroupTracking(
ctx,
pending,
pendingIndexes,
&completed,
&completedMu,
)
}
if err != nil {
saveErr := rw.updateAndSave(func(state *SagaState) {
state.Status = StatusFailed
state.Error = err.Error()
})
return errors.Join(rw.rollback(ctx, err, completed), saveErr)
}
}
if err := rw.updateAndSave(func(state *SagaState) {
state.Status = StatusCompleted
state.Error = ""
}); err != nil {
return fmt.Errorf("saga %q persist completion: %w", rw.id, err)
}
return nil
}
Returns
func (*RecoverableWorkflow) restoreState() error
{
persisted, err := rw.store.Load(rw.id)
if errors.Is(err, ErrSagaNotFound) {
return nil
}
if err != nil {
return fmt.Errorf("saga %q load state: %w", rw.id, err)
}
if persisted == nil || persisted.ID != rw.id {
return fmt.Errorf("saga %q loaded invalid state", rw.id)
}
rw.stateMu.Lock()
defer rw.stateMu.Unlock()
if len(persisted.Steps) != len(rw.state.Steps) {
return fmt.Errorf(
"saga %q persisted step count %d does not match workflow step count %d",
rw.id,
len(persisted.Steps),
len(rw.state.Steps),
)
}
for index := range persisted.Steps {
if persisted.Steps[index].Name != rw.state.Steps[index].Name {
return fmt.Errorf(
"saga %q persisted step %d is %q, want %q",
rw.id,
index,
persisted.Steps[index].Name,
rw.state.Steps[index].Name,
)
}
}
rw.state = cloneSagaState(persisted)
return nil
}
Parameters
Returns
func (*RecoverableWorkflow) stepsCompleted(indexes []int) bool
{
rw.stateMu.Lock()
defer rw.stateMu.Unlock()
for _, index := range indexes {
if index >= len(rw.state.Steps) || rw.state.Steps[index].Status != StatusCompleted {
return false
}
}
return true
}
Returns
func (*RecoverableWorkflow) persistedCompletedSteps() []Step
{
rw.stateMu.Lock()
statuses := make([]StepStatus, len(rw.state.Steps))
for index := range rw.state.Steps {
statuses[index] = rw.state.Steps[index].Status
}
rw.stateMu.Unlock()
var completed []Step
for itemIndex, item := range rw.steps {
switch typed := item.(type) {
case Step:
if statuses[rw.indexes[itemIndex][0]] == StatusCompleted && typed.Compensate != nil {
completed = append(completed, typed)
}
case Group:
for groupIndex, stateIndex := range rw.indexes[itemIndex] {
step := typed[groupIndex]
if statuses[stateIndex] == StatusCompleted && step.Compensate != nil {
completed = append(completed, step)
}
}
}
}
return completed
}
Parameters
Returns
func (*RecoverableWorkflow) runStepTracking(ctx context.Context, step Step, stepIndex int, completed *[]Step, completedMu *sync.Mutex) error
{
if err := executeStep(ctx, step); err != nil {
return err
}
if step.Compensate != nil {
completedMu.Lock()
*completed = append(*completed, step)
completedMu.Unlock()
}
if err := rw.updateAndSave(func(state *SagaState) {
if stepIndex < len(state.Steps) {
state.Steps[stepIndex].Status = StatusCompleted
}
}); err != nil {
return fmt.Errorf("saga %q persist step %q: %w", rw.id, step.Name, err)
}
return nil
}
Parameters
Returns
func (*RecoverableWorkflow) runGroupTracking(ctx context.Context, group Group, indexes []int, completed *[]Step, completedMu *sync.Mutex) error
{
var wg sync.WaitGroup
errChan := make(chan error, len(group))
for index, step := range group {
wg.Add(1)
go func(s Step, stateIndex int) {
defer wg.Done()
if err := rw.runStepTracking(ctx, s, stateIndex, completed, completedMu); err != nil {
errChan <- err
}
}(step, indexes[index])
}
wg.Wait()
close(errChan)
if len(errChan) > 0 {
var errs []error
for e := range errChan {
errs = append(errs, e)
}
return fmt.Errorf("group failed: %w", joinErrors(errs))
}
return nil
}
Parameters
Returns
func (*RecoverableWorkflow) rollback(ctx context.Context, triggerErr error, completed []Step) error
{
rollbackCtx := context.WithoutCancel(ctx)
var errs []error
errs = append(errs, triggerErr)
for i := len(completed) - 1; i >= 0; i-- {
step := completed[i]
if err := rw.safeCompensate(rollbackCtx, step); err != nil {
errs = append(errs, fmt.Errorf("rollback failed for '%s': %w", step.Name, err))
continue
}
if err := rw.updateAndSave(func(state *SagaState) {
for index := range state.Steps {
if state.Steps[index].Name == step.Name {
state.Steps[index].Status = StatusCompensated
break
}
}
}); err != nil {
errs = append(errs, fmt.Errorf("persist compensation for %q: %w", step.Name, err))
}
}
return joinErrors(errs)
}
Parameters
Returns
func (*RecoverableWorkflow) updateAndSave(update func(*SagaState)) error
{
rw.stateMu.Lock()
update(rw.state)
state := cloneSagaState(rw.state)
rw.stateMu.Unlock()
return rw.store.Save(state)
}
Fields
| Name | Type | Description |
|---|---|---|
| id | string | |
| store | SagaStore | |
| state | *SagaState | |
| stateMu | sync.Mutex | |
| indexes | [][]int |
Uses
NewRecoverable
NewRecoverable creates a saga that persists state using the given store.
Parameters
Returns
func NewRecoverable(id string, store SagaStore) *RecoverableWorkflow
{
return &RecoverableWorkflow{
Workflow: New(),
id: id,
store: store,
state: &SagaState{
ID: id,
Status: StatusPending,
},
}
}
Uses
joinErrors
Parameters
Returns
func joinErrors(errs []error) error
{
return errors.Join(errs...)
}
failingSagaStore
type failingSagaStore struct
Methods
Parameters
Returns
func (*failingSagaStore) Save(*SagaState) error
{
return s.saveErr
}
Parameters
Returns
func (*failingSagaStore) Load(string) (*SagaState, error)
{
if s.state == nil {
return nil, ErrSagaNotFound
}
return cloneSagaState(s.state), nil
}
Parameters
Returns
func (*failingSagaStore) Delete(string) error
{
return nil
}
Returns
func (*failingSagaStore) ListIncomplete() ([]*SagaState, error)
{
return nil, nil
}
Fields
| Name | Type | Description |
|---|---|---|
| state | *SagaState | |
| saveErr | error |
TestStore_SaveAndLoad
Parameters
func TestStore_SaveAndLoad(t *testing.T)
{
dir := t.TempDir()
store, err := NewStore(dir)
if err != nil {
t.Fatalf("NewStore: %v", err)
}
state := &SagaState{
ID: "test-1",
Status: StatusPending,
Steps: []StepState{
{Name: "reserve", Status: StatusCompleted, StepIndex: 0},
{Name: "charge", Status: StatusPending, StepIndex: 1},
},
}
if err := store.Save(state); err != nil {
t.Fatalf("Save: %v", err)
}
loaded, err := store.Load("test-1")
if err != nil {
t.Fatalf("Load: %v", err)
}
if loaded.ID != "test-1" {
t.Errorf("ID = %q, want %q", loaded.ID, "test-1")
}
if len(loaded.Steps) != 2 {
t.Errorf("len(Steps) = %d, want 2", len(loaded.Steps))
}
if loaded.Steps[0].Status != StatusCompleted {
t.Errorf("step[0] status = %q, want %q", loaded.Steps[0].Status, StatusCompleted)
}
}
TestStore_Delete
Parameters
func TestStore_Delete(t *testing.T)
{
dir := t.TempDir()
store, _ := NewStore(dir)
state := &SagaState{ID: "del-me", Status: StatusPending}
store.Save(state)
if err := store.Delete("del-me"); err != nil {
t.Fatalf("Delete: %v", err)
}
_, err := store.Load("del-me")
if err == nil {
t.Error("expected error after delete")
}
}
TestStore_ListIncomplete
Parameters
func TestStore_ListIncomplete(t *testing.T)
{
dir := t.TempDir()
store, _ := NewStore(dir)
store.Save(&SagaState{ID: "s1", Status: StatusPending})
store.Save(&SagaState{ID: "s2", Status: StatusCompleted})
store.Save(&SagaState{ID: "s3", Status: StatusFailed})
list, err := store.ListIncomplete()
if err != nil {
t.Fatalf("ListIncomplete: %v", err)
}
if len(list) != 2 {
t.Errorf("got %d incomplete, want 2", len(list))
}
}
TestStoreListIncompleteReportsCorruptFiles
Parameters
func TestStoreListIncompleteReportsCorruptFiles(t *testing.T)
{
dir := t.TempDir()
store, err := NewStore(dir)
if err != nil {
t.Fatal(err)
}
if err := store.Save(&SagaState{ID: "valid", Status: StatusPending}); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(dir, "corrupt.json"), []byte("{"), 0600); err != nil {
t.Fatal(err)
}
states, err := store.ListIncomplete()
if err == nil {
t.Fatal("ListIncomplete() ignored a corrupt state file")
}
if len(states) != 1 || states[0].ID != "valid" {
t.Fatalf("ListIncomplete() states = %v, want valid state", states)
}
}
TestStoreRejectsTraversalAndSymlinks
Parameters
func TestStoreRejectsTraversalAndSymlinks(t *testing.T)
{
dir := t.TempDir()
store, err := NewStore(dir)
if err != nil {
t.Fatal(err)
}
if err := store.Save(&SagaState{ID: "../../escape"}); err == nil {
t.Fatal("Save() accepted path traversal")
}
if _, err := store.Load("../../escape"); err == nil {
t.Fatal("Load() accepted path traversal")
}
if err := store.Delete("../../escape"); err == nil {
t.Fatal("Delete() accepted path traversal")
}
outside := filepath.Join(t.TempDir(), "outside.json")
if err := os.WriteFile(outside, []byte(`{"id":"outside"}`), 0600); err != nil {
t.Fatal(err)
}
if err := os.Symlink(outside, filepath.Join(dir, "linked.json")); err != nil {
t.Fatal(err)
}
if _, err := store.Load("linked"); err == nil {
t.Fatal("Load() followed a symbolic link")
}
list, err := store.ListIncomplete()
if err != nil {
t.Fatal(err)
}
if len(list) != 0 {
t.Fatalf("ListIncomplete() returned symlink state: %v", list)
}
}
TestStoreUsesPrivateModesAndSizeLimit
Parameters
func TestStoreUsesPrivateModesAndSizeLimit(t *testing.T)
{
dir := t.TempDir()
store, err := NewStore(dir)
if err != nil {
t.Fatal(err)
}
if err := store.Save(&SagaState{ID: "private"}); err != nil {
t.Fatal(err)
}
info, err := os.Stat(filepath.Join(dir, "private.json"))
if err != nil {
t.Fatal(err)
}
if info.Mode().Perm() != 0600 {
t.Fatalf("file mode = %o, want 600", info.Mode().Perm())
}
if err := store.Save(&SagaState{
ID: "large",
Error: string(make([]byte, maxSagaStateSize+1)),
}); err == nil {
t.Fatal("Save() accepted an oversized state")
}
}
TestRecoverableWorkflow_AllStepsSucceed
Parameters
func TestRecoverableWorkflow_AllStepsSucceed(t *testing.T)
{
dir := t.TempDir()
store, _ := NewStore(dir)
var executed []string
rw := NewRecoverable("order-1", store)
rw.Add("reserve",
func(ctx context.Context) error { executed = append(executed, "reserve"); return nil },
func(ctx context.Context) error { executed = append(executed, "undo-reserve"); return nil },
)
rw.Add("charge",
func(ctx context.Context) error { executed = append(executed, "charge"); return nil },
func(ctx context.Context) error { executed = append(executed, "undo-charge"); return nil },
)
err := rw.Run(context.Background())
if err != nil {
t.Fatalf("Run: %v", err)
}
if len(executed) != 2 {
t.Errorf("executed %v, want 2 steps", executed)
}
state, _ := store.Load("order-1")
if state.Status != StatusCompleted {
t.Errorf("status = %q, want %q", state.Status, StatusCompleted)
}
}
TestRecoverableWorkflow_StepFails_Persists
Parameters
func TestRecoverableWorkflow_StepFails_Persists(t *testing.T)
{
dir := t.TempDir()
store, _ := NewStore(dir)
rw := NewRecoverable("order-2", store)
rw.Add("reserve",
func(ctx context.Context) error { return nil },
func(ctx context.Context) error { return nil },
)
rw.Add("charge",
func(ctx context.Context) error { return errors.New("card declined") },
func(ctx context.Context) error { return nil },
)
err := rw.Run(context.Background())
if err == nil {
t.Fatal("expected error")
}
state, _ := store.Load("order-2")
if state.Status != StatusFailed {
t.Errorf("status = %q, want %q", state.Status, StatusFailed)
}
if state.Steps[0].Status != StatusCompensated {
t.Errorf("step[0] = %q, want %q", state.Steps[0].Status, StatusCompensated)
}
}
TestRecoverableWorkflow_CrashRecovery
Parameters
func TestRecoverableWorkflow_CrashRecovery(t *testing.T)
{
dir := t.TempDir()
store, _ := NewStore(dir)
rw := NewRecoverable("order-3", store)
rw.Add("reserve",
func(ctx context.Context) error { return nil },
func(ctx context.Context) error { return nil },
)
rw.Add("charge",
func(ctx context.Context) error { return errors.New("crash") },
func(ctx context.Context) error { return nil },
)
rw.Run(context.Background())
incomplete, err := store.ListIncomplete()
if err != nil {
t.Fatalf("ListIncomplete: %v", err)
}
if len(incomplete) != 1 {
t.Fatalf("got %d incomplete sagas, want 1", len(incomplete))
}
if incomplete[0].ID != "order-3" {
t.Errorf("ID = %q, want %q", incomplete[0].ID, "order-3")
}
_ = filepath.Join(dir, "order-3.json")
data, err := os.ReadFile(filepath.Join(dir, "order-3.json"))
if err != nil {
t.Fatalf("read file: %v", err)
}
t.Logf("persisted state: %s", string(data))
}
TestRecoverableWorkflowSkipsPersistedCompletedSteps
Parameters
func TestRecoverableWorkflowSkipsPersistedCompletedSteps(t *testing.T)
{
store := NewMemoryStore()
if err := store.Save(&SagaState{
ID: "recovery",
Status: StatusPending,
Steps: []StepState{
{Name: "completed", Status: StatusCompleted, StepIndex: 0},
{Name: "pending", Status: StatusPending, StepIndex: 1},
},
}); err != nil {
t.Fatal(err)
}
var completedCalls atomic.Int32
var pendingCalls atomic.Int32
workflow := NewRecoverable("recovery", store)
workflow.Add("completed", func(context.Context) error {
completedCalls.Add(1)
return nil
}, nil)
workflow.Add("pending", func(context.Context) error {
pendingCalls.Add(1)
return nil
}, nil)
if err := workflow.Run(context.Background()); err != nil {
t.Fatal(err)
}
if completedCalls.Load() != 0 {
t.Fatal("Run() repeated a persisted completed step")
}
if pendingCalls.Load() != 1 {
t.Fatalf("pending step calls = %d, want 1", pendingCalls.Load())
}
state, err := store.Load("recovery")
if err != nil {
t.Fatal(err)
}
if state.Status != StatusCompleted {
t.Fatalf("recovered saga status = %s, want completed", state.Status)
}
}
TestRecoverableWorkflowCompensatesPersistedCompletedSteps
Parameters
func TestRecoverableWorkflowCompensatesPersistedCompletedSteps(t *testing.T)
{
store := NewMemoryStore()
if err := store.Save(&SagaState{
ID: "rollback-recovery",
Status: StatusPending,
Steps: []StepState{
{Name: "completed", Status: StatusCompleted, StepIndex: 0},
{Name: "failure", Status: StatusPending, StepIndex: 1},
},
}); err != nil {
t.Fatal(err)
}
var compensationCalls atomic.Int32
workflow := NewRecoverable("rollback-recovery", store)
workflow.Add("completed", func(context.Context) error {
t.Fatal("persisted completed step was executed again")
return nil
}, func(context.Context) error {
compensationCalls.Add(1)
return nil
})
workflow.Add("failure", func(context.Context) error {
return errors.New("failed")
}, nil)
if err := workflow.Run(context.Background()); err == nil {
t.Fatal("Run() succeeded")
}
if compensationCalls.Load() != 1 {
t.Fatalf("compensation calls = %d, want 1", compensationCalls.Load())
}
state, err := store.Load("rollback-recovery")
if err != nil {
t.Fatal(err)
}
if state.Steps[0].Status != StatusCompensated {
t.Fatalf("persisted step status = %s, want compensated", state.Steps[0].Status)
}
}
TestMemoryStore_CRUD
Parameters
func TestMemoryStore_CRUD(t *testing.T)
{
store := NewMemoryStore()
state := &SagaState{ID: "m1", Status: StatusPending}
if err := store.Save(state); err != nil {
t.Fatalf("Save: %v", err)
}
loaded, err := store.Load("m1")
if err != nil {
t.Fatalf("Load: %v", err)
}
if loaded.ID != "m1" {
t.Errorf("ID = %q, want m1", loaded.ID)
}
store.Save(&SagaState{ID: "m2", Status: StatusCompleted})
incomplete, _ := store.ListIncomplete()
if len(incomplete) != 1 {
t.Errorf("incomplete = %d, want 1", len(incomplete))
}
store.Delete("m1")
_, err = store.Load("m1")
if err == nil {
t.Error("expected error after delete")
}
}
TestIdempotencyRecorder
Parameters
func TestIdempotencyRecorder(t *testing.T)
{
rec := NewIdempotencyRecorder()
state1 := &SagaState{ID: "s1", Status: StatusCompleted}
if err := rec.Record("key-1", state1); err != nil {
t.Fatalf("Record: %v", err)
}
err := rec.Record("key-1", &SagaState{ID: "s2"})
if err == nil {
t.Fatal("expected idempotency error")
}
var idErr *IdempotencyError
if !errors.As(err, &idErr) {
t.Errorf("error type = %T, want *IdempotencyError", err)
}
got, ok := rec.Get("key-1")
if !ok {
t.Fatal("expected key found")
}
if got.ID != "s1" {
t.Errorf("got ID = %q, want s1", got.ID)
}
_, ok = rec.Get("missing")
if ok {
t.Error("expected missing key not found")
}
}
TestDeadLetterQueue
Parameters
func TestDeadLetterQueue(t *testing.T)
{
dlq := NewDeadLetterQueue()
state := &SagaState{ID: "dead-1", Status: StatusFailed, RetryCount: 3}
dlq.Enqueue(state, "max retries exceeded")
if dlq.Len() != 1 {
t.Errorf("Len = %d, want 1", dlq.Len())
}
entries := dlq.List()
if len(entries) != 1 {
t.Fatalf("List len = %d, want 1", len(entries))
}
if entries[0].State.ID != "dead-1" {
t.Errorf("entry ID = %q, want dead-1", entries[0].State.ID)
}
if entries[0].Reason != "max retries exceeded" {
t.Errorf("reason = %q, want max retries exceeded", entries[0].Reason)
}
dlq.Remove("dead-1")
if dlq.Len() != 0 {
t.Errorf("Len after remove = %d, want 0", dlq.Len())
}
}
TestProcessWithDLQ_ExceedsRetries
Parameters
func TestProcessWithDLQ_ExceedsRetries(t *testing.T)
{
store := NewMemoryStore()
dlq := NewDeadLetterQueue()
state := &SagaState{ID: "dlq-1", Status: StatusFailed, RetryCount: 3, MaxRetries: 3}
store.Save(state)
processor := ProcessWithDLQ(store, dlq, 3)
run := processor("dlq-1", func(ctx context.Context) error {
return errors.New("still failing")
})
err := run(context.Background())
if err == nil {
t.Fatal("expected error")
}
updated, _ := store.Load("dlq-1")
if updated.Status != StatusDead {
t.Errorf("status = %q, want %q", updated.Status, StatusDead)
}
if dlq.Len() != 1 {
t.Errorf("dlq len = %d, want 1", dlq.Len())
}
}
TestMemoryStoresCloneAndEnforceCapacity
Parameters
func TestMemoryStoresCloneAndEnforceCapacity(t *testing.T)
{
store := NewMemoryStore(1)
state := &SagaState{
ID: "one",
Steps: []StepState{{Name: "step", Status: StatusPending}},
}
if err := store.Save(state); err != nil {
t.Fatal(err)
}
state.Steps[0].Status = StatusCompleted
loaded, err := store.Load("one")
if err != nil {
t.Fatal(err)
}
if loaded.Steps[0].Status != StatusPending {
t.Fatal("MemoryStore retained the caller's slice")
}
loaded.Steps[0].Status = StatusCompleted
again, _ := store.Load("one")
if again.Steps[0].Status != StatusPending {
t.Fatal("MemoryStore exposed its internal slice")
}
if err := store.Save(&SagaState{ID: "two"}); err == nil {
t.Fatal("MemoryStore exceeded its capacity")
}
recorder := NewIdempotencyRecorder(1)
if err := recorder.Record("one", state); err != nil {
t.Fatal(err)
}
if err := recorder.Record("two", &SagaState{ID: "two"}); err == nil {
t.Fatal("IdempotencyRecorder exceeded its capacity")
}
queue := NewDeadLetterQueue(1)
if err := queue.Enqueue(state, "failed"); err != nil {
t.Fatal(err)
}
if err := queue.Enqueue(&SagaState{ID: "two"}, "failed"); err == nil {
t.Fatal("DeadLetterQueue exceeded its capacity")
}
}
TestProcessWithDLQSurfacesSaveFailure
Parameters
func TestProcessWithDLQSurfacesSaveFailure(t *testing.T)
{
saveErr := errors.New("save failed")
store := &failingSagaStore{
state: &SagaState{ID: "one", Status: StatusPending},
saveErr: saveErr,
}
run := ProcessWithDLQ(store, NewDeadLetterQueue(), 3)(
"one",
func(context.Context) error { return nil },
)
if err := run(context.Background()); !errors.Is(err, saveErr) {
t.Fatalf("ProcessWithDLQ() error = %v, want save error", err)
}
}
TestRecoverableWorkflowSurfacesSaveFailure
Parameters
func TestRecoverableWorkflowSurfacesSaveFailure(t *testing.T)
{
saveErr := errors.New("save failed")
workflow := NewRecoverable("one", &failingSagaStore{saveErr: saveErr})
if err := workflow.Run(context.Background()); !errors.Is(err, saveErr) {
t.Fatalf("Run() error = %v, want save error", err)
}
}
TestProcessWithDLQ_RetrySucceeds
Parameters
func TestProcessWithDLQ_RetrySucceeds(t *testing.T)
{
store := NewMemoryStore()
dlq := NewDeadLetterQueue()
state := &SagaState{ID: "retry-1", Status: StatusFailed, RetryCount: 1, MaxRetries: 3}
store.Save(state)
processor := ProcessWithDLQ(store, dlq, 3)
run := processor("retry-1", func(ctx context.Context) error {
return nil
})
if err := run(context.Background()); err != nil {
t.Fatalf("unexpected error: %v", err)
}
updated, _ := store.Load("retry-1")
if updated.Status != StatusCompleted {
t.Errorf("status = %q, want %q", updated.Status, StatusCompleted)
}
if updated.RetryCount != 2 {
t.Errorf("retry count = %d, want 2", updated.RetryCount)
}
if dlq.Len() != 0 {
t.Errorf("dlq len = %d, want 0", dlq.Len())
}
}
TestProcessWithDLQSerializesSameSaga
Parameters
func TestProcessWithDLQSerializesSameSaga(t *testing.T)
{
store := NewMemoryStore()
if err := store.Save(&SagaState{ID: "shared", Status: StatusFailed}); err != nil {
t.Fatal(err)
}
var active atomic.Int32
var maximum atomic.Int32
runFunc := func(context.Context) error {
current := active.Add(1)
for {
observed := maximum.Load()
if current <= observed || maximum.CompareAndSwap(observed, current) {
break
}
}
time.Sleep(10 * time.Millisecond)
active.Add(-1)
return nil
}
processor := ProcessWithDLQ(store, NewDeadLetterQueue(), 10)
run := processor("shared", runFunc)
var wait sync.WaitGroup
wait.Add(2)
for range 2 {
go func() {
defer wait.Done()
_ = run(context.Background())
}()
}
wait.Wait()
if maximum.Load() != 1 {
t.Fatalf("concurrent saga executions = %d, want 1", maximum.Load())
}
state, err := store.Load("shared")
if err != nil {
t.Fatal(err)
}
if state.RetryCount != 1 {
t.Fatalf("retry count = %d, want 1", state.RetryCount)
}
}
TestRecoverableWorkflowRecoversStepPanic
Parameters
func TestRecoverableWorkflowRecoversStepPanic(t *testing.T)
{
store := NewMemoryStore()
workflow := NewRecoverable("panic", store)
workflow.Add("panic", func(context.Context) error {
panic("step failed")
}, nil)
if err := workflow.Run(context.Background()); err == nil {
t.Fatal("Run() accepted a step panic")
}
state, err := store.Load("panic")
if err != nil {
t.Fatal(err)
}
if state.Status != StatusFailed {
t.Fatalf("state status = %s, want failed", state.Status)
}
}
TestSagaStore_Interface
Parameters
func TestSagaStore_Interface(t *testing.T)
{
var _ SagaStore = NewMemoryStore()
var _ SagaStore = &FileStore{}
dir := t.TempDir()
fs, _ := NewStore(dir)
var _ SagaStore = fs
}
Step
Step defines a saga step with a do and compensate action.
type Step struct
Fields
| Name | Type | Description |
|---|---|---|
| Name | string | |
| Do | func(ctx context.Context) error | |
| Compensate | func(ctx context.Context) error |
Group
Group is a collection of steps executed in parallel.
type Group []Step
Workflow
Workflow orchestrates a saga with rollback support.
type Workflow struct
Methods
Add appends a step to the workflow.
Parameters
func (*Workflow) Add(name string, do, compensate func(ctx context.Context) error)
{
w.steps = append(w.steps, Step{
Name: name,
Do: do,
Compensate: compensate,
})
}
AddGroup appends a parallel step group to the workflow.
Parameters
func (*Workflow) AddGroup(g Group)
{
w.steps = append(w.steps, g)
}
Run executes all steps in order, rolling back on failure.
Parameters
Returns
func (*Workflow) Run(ctx context.Context) error
{
if !w.running.CompareAndSwap(false, true) {
return errors.New("saga: workflow is already running")
}
defer w.running.Store(false)
var completed []Step
var completedMu sync.Mutex
for _, item := range w.steps {
if ctx.Err() != nil {
return w.rollback(ctx, ctx.Err(), completed)
}
var err error
switch v := item.(type) {
case Step:
err = w.runStep(ctx, v, &completed, &completedMu)
case Group:
err = w.runGroup(ctx, v, &completed, &completedMu)
}
if err != nil {
return w.rollback(ctx, err, completed)
}
}
return nil
}
Parameters
Returns
func (*Workflow) runStep(ctx context.Context, step Step, completed *[]Step, completedMu *sync.Mutex) error
{
if err := executeStep(ctx, step); err != nil {
return err
}
if step.Compensate != nil {
completedMu.Lock()
*completed = append(*completed, step)
completedMu.Unlock()
}
return nil
}
Parameters
Returns
func (*Workflow) runGroup(ctx context.Context, group Group, completed *[]Step, completedMu *sync.Mutex) error
{
var wg sync.WaitGroup
errChan := make(chan error, len(group))
for _, step := range group {
wg.Add(1)
go func(s Step) {
defer wg.Done()
if err := w.runStep(ctx, s, completed, completedMu); err != nil {
errChan <- err
}
}(step)
}
wg.Wait()
close(errChan)
if len(errChan) > 0 {
var errs []error
for e := range errChan {
errs = append(errs, e)
}
return errors.Join(errs...)
}
return nil
}
Parameters
Returns
func (*Workflow) rollback(ctx context.Context, triggerErr error, completed []Step) error
{
rollbackCtx := context.WithoutCancel(ctx)
var errs []error
errs = append(errs, triggerErr)
for i := len(completed) - 1; i >= 0; i-- {
step := completed[i]
if err := w.safeCompensate(rollbackCtx, step); err != nil {
errs = append(errs, fmt.Errorf("rollback failed for '%s': %w", step.Name, err))
}
}
return errors.Join(errs...)
}
Parameters
Returns
func (*Workflow) safeCompensate(ctx context.Context, step Step) (err error)
{
defer func() {
if r := recover(); r != nil {
err = fmt.Errorf("panic during compensation: %v", r)
}
}()
return step.Compensate(ctx)
}
Fields
| Name | Type | Description |
|---|---|---|
| steps | []any | |
| running | atomic.Bool |
New
New creates an empty Workflow.
Returns
func New() *Workflow
{
return &Workflow{}
}
executeStep
Parameters
Returns
func executeStep(ctx context.Context, step Step) (err error)
{
defer func() {
if r := recover(); r != nil {
err = fmt.Errorf("panic in step '%s': %v", step.Name, r)
}
}()
if err := step.Do(ctx); err != nil {
return fmt.Errorf("step '%s' failed: %w", step.Name, err)
}
return nil
}
Uses
TestWorkflow_Run_AllStepsSucceed
Parameters
func TestWorkflow_Run_AllStepsSucceed(t *testing.T)
{
var executed []string
wf := New()
wf.Add("step1",
func(ctx context.Context) error { executed = append(executed, "step1"); return nil },
func(ctx context.Context) error { executed = append(executed, "undo1"); return nil },
)
wf.Add("step2",
func(ctx context.Context) error { executed = append(executed, "step2"); return nil },
func(ctx context.Context) error { executed = append(executed, "undo2"); return nil },
)
err := wf.Run(context.Background())
if err != nil {
t.Fatalf("Run failed: %v", err)
}
if len(executed) != 2 {
t.Fatalf("expected 2 steps executed, got %v", executed)
}
if executed[0] != "step1" || executed[1] != "step2" {
t.Errorf("unexpected order: %v", executed)
}
}
TestWorkflow_Run_StepFails_Compensates
Parameters
func TestWorkflow_Run_StepFails_Compensates(t *testing.T)
{
var executed []string
wf := New()
wf.Add("step1",
func(ctx context.Context) error { executed = append(executed, "step1"); return nil },
func(ctx context.Context) error { executed = append(executed, "undo1"); return nil },
)
wf.Add("step2",
func(ctx context.Context) error { executed = append(executed, "step2"); return errors.New("fail") },
func(ctx context.Context) error { executed = append(executed, "undo2"); return nil },
)
wf.Add("step3",
func(ctx context.Context) error { executed = append(executed, "step3"); return nil },
func(ctx context.Context) error { executed = append(executed, "undo3"); return nil },
)
err := wf.Run(context.Background())
if err == nil {
t.Fatal("expected error")
}
expect := []string{"step1", "step2", "undo1"}
if len(executed) != len(expect) {
t.Fatalf("expected %v, got %v", expect, executed)
}
for i, v := range expect {
if executed[i] != v {
t.Errorf("executed[%d] = %q, want %q", i, executed[i], v)
}
}
}
TestWorkflow_Run_ContextCancelled
Parameters
func TestWorkflow_Run_ContextCancelled(t *testing.T)
{
var executed []string
ctx, cancel := context.WithCancel(context.Background())
wf := New()
wf.Add("step1",
func(ctx context.Context) error { executed = append(executed, "step1"); return nil },
func(ctx context.Context) error { executed = append(executed, "undo1"); return nil },
)
wf.Add("step2",
func(ctx context.Context) error {
cancel()
return ctx.Err()
},
func(ctx context.Context) error { return nil },
)
err := wf.Run(ctx)
if err == nil {
t.Fatal("expected error from cancelled context")
}
if len(executed) != 2 {
t.Fatalf("expected step1 and step2 executed, got %v", executed)
}
}
TestWorkflow_Run_PanicInStep
Parameters
func TestWorkflow_Run_PanicInStep(t *testing.T)
{
var executed []string
wf := New()
wf.Add("step1",
func(ctx context.Context) error { executed = append(executed, "step1"); return nil },
func(ctx context.Context) error { executed = append(executed, "undo1"); return nil },
)
wf.Add("panic-step",
func(ctx context.Context) error {
panic("something went wrong")
},
func(ctx context.Context) error { return nil },
)
err := wf.Run(context.Background())
if err == nil {
t.Fatal("expected error from panic")
}
expect := []string{"step1", "undo1"}
if len(executed) != len(expect) {
t.Fatalf("expected %v, got %v", expect, executed)
}
for i, v := range expect {
if executed[i] != v {
t.Errorf("executed[%d] = %q, want %q", i, executed[i], v)
}
}
}
TestWorkflow_Run_GroupParallel
Parameters
func TestWorkflow_Run_GroupParallel(t *testing.T)
{
var mu sync.Mutex
var executed []string
wf := New()
group := Group{
{Name: "g1", Do: func(ctx context.Context) error { mu.Lock(); executed = append(executed, "g1"); mu.Unlock(); return nil },
Compensate: func(ctx context.Context) error { return nil }},
{Name: "g2", Do: func(ctx context.Context) error { mu.Lock(); executed = append(executed, "g2"); mu.Unlock(); return nil },
Compensate: func(ctx context.Context) error { return nil }},
}
wf.AddGroup(group)
err := wf.Run(context.Background())
if err != nil {
t.Fatalf("Run failed: %v", err)
}
if len(executed) != 2 {
t.Errorf("expected 2 group steps, got %v", executed)
}
}
TestWorkflow_Run_CompensatePanicSafe
Parameters
func TestWorkflow_Run_CompensatePanicSafe(t *testing.T)
{
wf := New()
wf.Add("step1",
func(ctx context.Context) error { return nil },
func(ctx context.Context) error { panic("compensate panic") },
)
wf.Add("step2",
func(ctx context.Context) error { return errors.New("fail") },
nil,
)
err := wf.Run(context.Background())
if err == nil {
t.Fatal("expected error")
}
}
TestWorkflow_Run_NoCompensateOnSuccess
Parameters
func TestWorkflow_Run_NoCompensateOnSuccess(t *testing.T)
{
var compensated bool
wf := New()
wf.Add("step1",
func(ctx context.Context) error { return nil },
func(ctx context.Context) error { compensated = true; return nil },
)
err := wf.Run(context.Background())
if err != nil {
t.Fatalf("Run failed: %v", err)
}
if compensated {
t.Error("expected no compensation on success")
}
}
TestWorkflowDoesNotReuseCompensationsAcrossRuns
Parameters
func TestWorkflowDoesNotReuseCompensationsAcrossRuns(t *testing.T)
{
workflow := New()
var fail bool
var compensations int
workflow.Add(
"first",
func(context.Context) error { return nil },
func(context.Context) error {
compensations++
return nil
},
)
workflow.Add(
"second",
func(context.Context) error {
if fail {
return errors.New("failed")
}
return nil
},
nil,
)
if err := workflow.Run(context.Background()); err != nil {
t.Fatal(err)
}
fail = true
if err := workflow.Run(context.Background()); err == nil {
t.Fatal("second Run() succeeded")
}
if compensations != 1 {
t.Fatalf("compensations = %d, want 1 for the failing run", compensations)
}
}
TestWorkflowCompensationCanReenterRun
Parameters
func TestWorkflowCompensationCanReenterRun(t *testing.T)
{
workflow := New()
reentered := make(chan error, 1)
workflow.Add(
"first",
func(context.Context) error { return nil },
func(context.Context) error {
reentered <- workflow.Run(context.Background())
return nil
},
)
workflow.Add(
"failure",
func(context.Context) error { return errors.New("failed") },
nil,
)
done := make(chan error, 1)
go func() {
done <- workflow.Run(context.Background())
}()
select {
case err := <-done:
if err == nil {
t.Fatal("Run() succeeded")
}
case <-time.After(time.Second):
t.Fatal("compensation deadlocked while reentering the workflow")
}
if err := <-reentered; err == nil {
t.Fatal("reentrant Run() was accepted")
}
}
TestWorkflow_Compensate_WithRetry
Parameters
func TestWorkflow_Compensate_WithRetry(t *testing.T)
{
attempts := 0
do := WithRetry(RetryPolicy{MaxAttempts: 3, Delay: 10 * time.Millisecond, Multiplier: 1.0},
func(ctx context.Context) error {
attempts++
return nil
},
)
err := do(context.Background())
if err != nil {
t.Fatalf("WithRetry failed: %v", err)
}
if attempts != 1 {
t.Errorf("expected 1 attempt on success, got %d", attempts)
}
}
TestWorkflow_Compensate_WithRetryExhausted
Parameters
func TestWorkflow_Compensate_WithRetryExhausted(t *testing.T)
{
attempts := 0
do := WithRetry(RetryPolicy{MaxAttempts: 3, Delay: 10 * time.Millisecond, Multiplier: 1.0},
func(ctx context.Context) error {
attempts++
return errors.New("always fail")
},
)
err := do(context.Background())
if err == nil {
t.Fatal("expected error after retry exhausted")
}
if attempts != 3 {
t.Errorf("expected 3 attempts, got %d", attempts)
}
}