aboutsummaryrefslogtreecommitdiff
path: root/runtime
diff options
context:
space:
mode:
Diffstat (limited to 'runtime')
-rw-r--r--runtime/task.go14
-rw-r--r--runtime/task_test.go116
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)
}