jackhammer

Utilities for Go
Log | Files | Refs | README | LICENSE

para_test.go (4145B)


      1 package para
      2 
      3 import (
      4 	"context"
      5 	"fmt"
      6 	"slices"
      7 	"strconv"
      8 	"sync"
      9 	"sync/atomic"
     10 	"testing"
     11 )
     12 
     13 func TestFanout(t *testing.T) {
     14 	t.Run("happy", func(t *testing.T) {
     15 		n := 10
     16 
     17 		input := []int{}
     18 		for ii := range n {
     19 			input = append(input, ii)
     20 		}
     21 
     22 		var calls atomic.Int32
     23 
     24 		Fanout(slices.Values(input), func(v int) {
     25 			calls.Add(1)
     26 		})
     27 
     28 		if got := calls.Load(); got != int32(n) {
     29 			t.Fatalf("calls mismatch: want=%v, got=%v", n, got)
     30 		}
     31 	})
     32 }
     33 
     34 func TestBatch(t *testing.T) {
     35 	t.Run("happy", func(t *testing.T) {
     36 		ctx, cancel := context.WithCancel(context.Background())
     37 		defer cancel()
     38 
     39 		input := []string{}
     40 		want := []int64{}
     41 
     42 		for ii := range 100 {
     43 			input = append(input, strconv.Itoa(ii))
     44 			want = append(want, int64(ii))
     45 		}
     46 
     47 		got := make([]int64, len(input))
     48 
     49 		err := Batch(ctx, slices.Values(input), func(ctx context.Context, ii int, v string) error {
     50 			n, err := strconv.ParseInt(v, 10, 64)
     51 			got[ii] = n
     52 			return err
     53 		})
     54 
     55 		if err != nil {
     56 			t.Fatalf("unexpected error: %v", err)
     57 		}
     58 
     59 		if !slices.Equal(got, want) {
     60 			t.Fatalf("want=%v, got=%v", want, got)
     61 		}
     62 	})
     63 
     64 	t.Run("sad", func(t *testing.T) {
     65 		ctx, cancel := context.WithCancel(context.Background())
     66 		defer cancel()
     67 
     68 		input := []string{}
     69 
     70 		n := 5
     71 
     72 		for ii := range n {
     73 			input = append(input, strconv.Itoa(ii))
     74 		}
     75 
     76 		trigger := make(chan any, 1)
     77 		ready := sync.WaitGroup{}
     78 		ready.Add(n)
     79 
     80 		go func() {
     81 			ready.Wait()
     82 			// trigger one closure to return an error
     83 			trigger <- nil
     84 		}()
     85 
     86 		var nextID atomic.Int32
     87 		var canceled atomic.Int32
     88 		var errored atomic.Int32
     89 
     90 		err := Batch(ctx, slices.Values(input), func(ctx context.Context, ii int, v string) error {
     91 			id := nextID.Add(1)
     92 
     93 			ready.Done()
     94 
     95 			select {
     96 			case <-ctx.Done():
     97 				// most should get canceled due to the one that failed
     98 				// should be n-1
     99 				canceled.Add(1)
    100 				t.Logf("%d: canceled", id)
    101 				return ctx.Err()
    102 			case <-trigger:
    103 				// should be one
    104 				errored.Add(1)
    105 				t.Logf("%d: errored", id)
    106 				return fmt.Errorf("an error")
    107 			}
    108 		})
    109 
    110 		if err == nil {
    111 			t.Fatalf("expected error")
    112 		}
    113 
    114 		if got := canceled.Load(); got != int32(n-1) {
    115 			t.Errorf("expected %d cancels, got %v", int32(n-1), got)
    116 		}
    117 
    118 		if got := errored.Load(); got != 1 {
    119 			t.Errorf("expected 1 error, got %v", got)
    120 		}
    121 	})
    122 }
    123 
    124 func TestMapErr(t *testing.T) {
    125 	t.Run("happy", func(t *testing.T) {
    126 		ctx, cancel := context.WithCancel(context.Background())
    127 		defer cancel()
    128 
    129 		input := []string{}
    130 		want := []int64{}
    131 
    132 		for ii := range 100 {
    133 			input = append(input, strconv.Itoa(ii))
    134 			want = append(want, int64(ii))
    135 		}
    136 
    137 		got, err := MapErr(ctx, input, func(ctx context.Context, ii int, v string) (int64, error) {
    138 			return strconv.ParseInt(v, 10, 64)
    139 		})
    140 
    141 		if err != nil {
    142 			t.Fatalf("unexpected error: %v", err)
    143 		}
    144 
    145 		if !slices.Equal(got, want) {
    146 			t.Fatalf("want=%v, got=%v", want, got)
    147 		}
    148 	})
    149 
    150 	t.Run("sad", func(t *testing.T) {
    151 		ctx, cancel := context.WithCancel(context.Background())
    152 		defer cancel()
    153 
    154 		input := []string{}
    155 
    156 		n := 5
    157 
    158 		for ii := range n {
    159 			input = append(input, strconv.Itoa(ii))
    160 		}
    161 
    162 		trigger := make(chan any, 1)
    163 		ready := sync.WaitGroup{}
    164 		ready.Add(n)
    165 
    166 		go func() {
    167 			ready.Wait()
    168 			// trigger one closure to return an error
    169 			trigger <- nil
    170 		}()
    171 
    172 		var nextID atomic.Int32
    173 		var canceled atomic.Int32
    174 		var errored atomic.Int32
    175 
    176 		got, err := MapErr(ctx, input, func(ctx context.Context, ii int, v string) (int64, error) {
    177 			id := nextID.Add(1)
    178 
    179 			ready.Done()
    180 
    181 			select {
    182 			case <-ctx.Done():
    183 				// most should get canceled due to the one that failed
    184 				// should be n-1
    185 				canceled.Add(1)
    186 				t.Logf("%d: canceled", id)
    187 				return 0, ctx.Err()
    188 			case <-trigger:
    189 				// should be one
    190 				errored.Add(1)
    191 				t.Logf("%d: errored", id)
    192 				return 0, fmt.Errorf("an error")
    193 			}
    194 		})
    195 
    196 		if err == nil {
    197 			t.Fatalf("expected error")
    198 		}
    199 
    200 		if len(got) != 0 {
    201 			t.Fatalf("expected result to be empty, got %v", got)
    202 		}
    203 
    204 		if got := canceled.Load(); got != int32(n-1) {
    205 			t.Errorf("expected %d cancels, got %v", int32(n-1), got)
    206 		}
    207 
    208 		if got := errored.Load(); got != 1 {
    209 			t.Errorf("expected 1 error, got %v", got)
    210 		}
    211 	})
    212 }