diff options
| author | SangTran-127 <tranquangsang12.7@gmail.com> | 2026-08-09 18:32:44 +0700 |
|---|---|---|
| committer | SangTran-127 <tranquangsang12.7@gmail.com> | 2026-08-09 18:32:44 +0700 |
| commit | 4ca27aec303f54dc9ee670c67c2fefa80dd004a1 (patch) | |
| tree | 5628714f47e374bfb98c1cea99970b19892e8b3e /runtime | |
| parent | a892cb5b1d121177881fa56a79abbafef5880a80 (diff) | |
| download | goflink-4ca27aec303f54dc9ee670c67c2fefa80dd004a1.tar.gz goflink-4ca27aec303f54dc9ee670c67c2fefa80dd004a1.zip | |
feat: add KeyByOperator and wordCounter for stateful processing
Diffstat (limited to 'runtime')
| -rw-r--r-- | runtime/task.go | 14 | ||||
| -rw-r--r-- | runtime/task_test.go | 116 |
2 files changed, 118 insertions, 12 deletions
diff --git a/runtime/task.go b/runtime/task.go index 3c887bb..b3d1caa 100644 --- a/runtime/task.go +++ b/runtime/task.go @@ -5,19 +5,19 @@ import ( "fmt" "time" - "goflink/core" + "goflink/core/operator" "goflink/transport" ) // RunTask runs a stream task in the background. The returned channel yields the // first error that stopped it (context.Canceled when ctx ends early) and is // closed when the task is done, so callers can join on it. -func RunTask[IN, OUT any](ctx context.Context, input transport.Receiver[IN], output transport.Emitter[OUT], operator core.Operator[IN, OUT]) <-chan error { +func RunTask[IN, OUT any](ctx context.Context, input transport.Receiver[IN], output transport.Emitter[OUT], op operator.Operator[IN, OUT]) <-chan error { done := make(chan error, 1) go func() { defer close(done) - if err := runTask(ctx, input, output, operator); err != nil { + if err := runTask(ctx, input, output, op); err != nil { done <- err } }() @@ -25,7 +25,7 @@ func RunTask[IN, OUT any](ctx context.Context, input transport.Receiver[IN], out return done } -func runTask[IN, OUT any](ctx context.Context, input transport.Receiver[IN], output transport.Emitter[OUT], operator core.Operator[IN, OUT]) (err error) { +func runTask[IN, OUT any](ctx context.Context, input transport.Receiver[IN], output transport.Emitter[OUT], op operator.Operator[IN, OUT]) (err error) { // ponytail: assumes one writer per output channel. Fan-in needs a refcount // or a dedicated closer, otherwise the second Close panics. defer func() { @@ -34,7 +34,7 @@ func runTask[IN, OUT any](ctx context.Context, input transport.Receiver[IN], out } }() - if err := operator.Open(ctx); err != nil { + if err := op.Open(ctx); err != nil { return fmt.Errorf("open operator: %w", err) } defer func() { @@ -44,7 +44,7 @@ func runTask[IN, OUT any](ctx context.Context, input transport.Receiver[IN], out closeCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) defer cancel() - if cerr := operator.Close(closeCtx); cerr != nil && err == nil { + if cerr := op.Close(closeCtx); cerr != nil && err == nil { err = fmt.Errorf("close operator: %w", cerr) } }() @@ -55,7 +55,7 @@ func runTask[IN, OUT any](ctx context.Context, input transport.Receiver[IN], out break } - if err := operator.Process(ctx, record, output); err != nil { + if err := op.Process(ctx, record, output); err != nil { return fmt.Errorf("process record: %w", err) } } diff --git a/runtime/task_test.go b/runtime/task_test.go index f27a5eb..ae72c14 100644 --- a/runtime/task_test.go +++ b/runtime/task_test.go @@ -6,6 +6,8 @@ import ( "testing" "goflink/core" + "goflink/core/operator" + "goflink/state" "goflink/transport" ) @@ -21,12 +23,12 @@ func TestPipeline(t *testing.T) { mapped := RunTask(ctx, transport.NewLocalReceiveChannel(source), transport.NewLocalEmitChannel(mid), - core.NewMapOperator(func(in int) (int, error) { return in * 10, nil }), + operator.NewMapOperator(func(in int) (int, error) { return in * 10, nil }), ) filtered := RunTask(ctx, transport.NewLocalReceiveChannel(mid), transport.NewLocalEmitChannel(sink), - core.NewFilterOperator(func(in int) bool { return in > 20 }), + operator.NewFilterOperator(func(in int) bool { return in > 20 }), ) go func() { @@ -58,6 +60,110 @@ func TestPipeline(t *testing.T) { } } +// wordCounter is the canonical stateful user func: read the old count for the +// current key, +1, store it back. +type wordCounter struct { + count core.ValueState[any] +} + +func (w *wordCounter) Open(ctx context.Context) error { + be, ok := core.ExtractStateBackend(ctx) + if !ok { + return errors.New("no state backend in context") + } + s, err := be.GetValueState(ctx, "word-count") + if err != nil { + return err + } + w.count = s + return nil +} + +func (w *wordCounter) Map(ctx context.Context, in string) (int, error) { + v, err := w.count.Value(ctx) + if err != nil { + return 0, err + } + n, _ := v.(int) // unseen key -> nil -> 0 + n++ + _, err = w.count.Update(ctx, n) + return n, err +} + +func (w *wordCounter) Close(ctx context.Context) error { return nil } + +// keyBy -> statefulMap: state must be per key and survive across records. +func TestStatefulMapCountsPerKey(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + ctx = core.InjectStateBackend(ctx, state.NewMemoryStateBackend(core.DefaultMaxParallelism)) + + source := make(chan core.StreamRecord[string]) + mid := make(chan core.StreamRecord[string]) + sink := make(chan core.StreamRecord[int]) + + keyed := RunTask(ctx, + transport.NewLocalReceiveChannel(source), + transport.NewLocalEmitChannel(mid), + operator.NewKeyByOperator(func(in string) string { return in }), + ) + counted := RunTask(ctx, + transport.NewLocalReceiveChannel(mid), + transport.NewLocalEmitChannel(sink), + operator.NewStatefulMapOperator[string, int](&wordCounter{}), + ) + + go func() { + defer close(source) + for _, name := range []string{"Sang", "Tran", "Sang", "Sang"} { + source <- core.NewDataRecord("", name, 0) + } + }() + + var got []int + for record := range sink { + got = append(got, record.Value) + } + + want := []int{1, 1, 2, 3} + if len(got) != len(want) { + t.Fatalf("got %v, want %v", got, want) + } + for i := range want { + if got[i] != want[i] { + t.Fatalf("got %v, want %v", got, want) + } + } + + for _, errCh := range []<-chan error{keyed, counted} { + if err := <-errCh; err != nil { + t.Fatalf("task failed: %v", err) + } + } +} + +// An unkeyed record must fail loudly instead of silently sharing one state slot. +func TestStatefulMapRejectsUnkeyed(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + ctx = core.InjectStateBackend(ctx, state.NewMemoryStateBackend(core.DefaultMaxParallelism)) + + source := make(chan core.StreamRecord[string], 1) + sink := make(chan core.StreamRecord[int], 1) + + done := RunTask(ctx, + transport.NewLocalReceiveChannel(source), + transport.NewLocalEmitChannel(sink), + operator.NewStatefulMapOperator[string, int](&wordCounter{}), + ) + + source <- core.NewDataRecord("", "Sang", 0) + + if err := <-done; err == nil { + t.Fatal("unkeyed record was accepted, want an error") + } +} + // A cancelled ctx must unblock a task parked in Receive. func TestCancelStopsTask(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) @@ -68,7 +174,7 @@ func TestCancelStopsTask(t *testing.T) { done := RunTask(ctx, transport.NewLocalReceiveChannel(source), transport.NewLocalEmitChannel(sink), - core.NewMapOperator(func(in int) (int, error) { return in, nil }), + operator.NewMapOperator(func(in int) (int, error) { return in, nil }), ) cancel() @@ -90,7 +196,7 @@ func TestProcessErrorPropagates(t *testing.T) { done := RunTask(ctx, transport.NewLocalReceiveChannel(source), transport.NewLocalEmitChannel(sink), - core.NewMapOperator(func(in int) (int, error) { return 0, boom }), + operator.NewMapOperator(func(in int) (int, error) { return 0, boom }), ) source <- core.NewDataRecord("k", 1, 0) @@ -107,7 +213,7 @@ type closeProbe struct { func (o *closeProbe) Open(ctx context.Context) error { return nil } -func (o *closeProbe) Process(ctx context.Context, r core.StreamRecord[int], out core.Emitter[int]) error { +func (o *closeProbe) Process(ctx context.Context, r core.StreamRecord[int], out operator.Emitter[int]) error { return out.Emit(ctx, r) } |