diff --git a/sdks/go/pkg/beam/runners/prism/internal/handlecombine.go b/sdks/go/pkg/beam/runners/prism/internal/handlecombine.go index 6b336043b8c9..d65ef63cccc9 100644 --- a/sdks/go/pkg/beam/runners/prism/internal/handlecombine.go +++ b/sdks/go/pkg/beam/runners/prism/internal/handlecombine.go @@ -64,43 +64,52 @@ func (h *combine) PrepareTransform(tid string, t *pipepb.PTransform, comps *pipe combineInput := comps.GetPcollections()[onlyInput] ws := comps.GetWindowingStrategies()[combineInput.GetWindowingStrategyId()] - var hasElementCount func(tpb *pipepb.Trigger) bool + var hasTriggerType func(tpb *pipepb.Trigger, targetTriggerType reflect.Type) bool - hasElementCount = func(tpb *pipepb.Trigger) bool { - elCount := false + hasTriggerType = func(tpb *pipepb.Trigger, targetTriggerType reflect.Type) bool { + if tpb == nil { + return false + } switch at := tpb.GetTrigger().(type) { - case *pipepb.Trigger_ElementCount_: - return true case *pipepb.Trigger_AfterAll_: for _, st := range at.AfterAll.GetSubtriggers() { - elCount = elCount || hasElementCount(st) + if hasTriggerType(st, targetTriggerType) { + return true + } } - return elCount + return false case *pipepb.Trigger_AfterAny_: for _, st := range at.AfterAny.GetSubtriggers() { - elCount = elCount || hasElementCount(st) + if hasTriggerType(st, targetTriggerType) { + return true + } } - return elCount + return false case *pipepb.Trigger_AfterEach_: for _, st := range at.AfterEach.GetSubtriggers() { - elCount = elCount || hasElementCount(st) + if hasTriggerType(st, targetTriggerType) { + return true + } } - return elCount + return false case *pipepb.Trigger_AfterEndOfWindow_: - return hasElementCount(at.AfterEndOfWindow.GetEarlyFirings()) || - hasElementCount(at.AfterEndOfWindow.GetLateFirings()) + return hasTriggerType(at.AfterEndOfWindow.GetEarlyFirings(), targetTriggerType) || + hasTriggerType(at.AfterEndOfWindow.GetLateFirings(), targetTriggerType) case *pipepb.Trigger_OrFinally_: - return hasElementCount(at.OrFinally.GetMain()) || - hasElementCount(at.OrFinally.GetFinally()) + return hasTriggerType(at.OrFinally.GetMain(), targetTriggerType) || + hasTriggerType(at.OrFinally.GetFinally(), targetTriggerType) case *pipepb.Trigger_Repeat_: - return hasElementCount(at.Repeat.GetSubtrigger()) + return hasTriggerType(at.Repeat.GetSubtrigger(), targetTriggerType) default: - return false + return reflect.TypeOf(at) == targetTriggerType } } // If we aren't lifting, the "default impl" for combines should be sufficient. - if !h.config.EnableLifting || hasElementCount(ws.GetTrigger()) { + // Disable lifting if there is any TriggerElementCount or TriggerAlways. + if (!h.config.EnableLifting || + hasTriggerType(ws.GetTrigger(), reflect.TypeOf(&pipepb.Trigger_ElementCount_{})) || + hasTriggerType(ws.GetTrigger(), reflect.TypeOf(&pipepb.Trigger_Always_{}))) { return prepareResult{} // Strip the composite layer when lifting is disabled. } diff --git a/sdks/go/pkg/beam/runners/prism/internal/handlecombine_test.go b/sdks/go/pkg/beam/runners/prism/internal/handlecombine_test.go index 7b38daa295ef..26be37e77d17 100644 --- a/sdks/go/pkg/beam/runners/prism/internal/handlecombine_test.go +++ b/sdks/go/pkg/beam/runners/prism/internal/handlecombine_test.go @@ -25,10 +25,14 @@ import ( "google.golang.org/protobuf/testing/protocmp" ) -func TestHandleCombine(t *testing.T) { - undertest := "UnderTest" +func makeWindowingStrategy(trigger *pipepb.Trigger) *pipepb.WindowingStrategy { + return &pipepb.WindowingStrategy{ + Trigger: trigger, + } +} - combineTransform := &pipepb.PTransform{ +func makeCombineTransform(inputPCollectionID string) *pipepb.PTransform { + return &pipepb.PTransform{ UniqueName: "COMBINE", Spec: &pipepb.FunctionSpec{ Urn: urns.TransformCombinePerKey, @@ -41,7 +45,7 @@ func TestHandleCombine(t *testing.T) { }), }, Inputs: map[string]string{ - "input": "combineIn", + "input": inputPCollectionID, }, Outputs: map[string]string{ "input": "combineOut", @@ -51,6 +55,15 @@ func TestHandleCombine(t *testing.T) { "combine_values", }, } +} + +func TestHandleCombine(t *testing.T) { + undertest := "UnderTest" + + combineTransform := makeCombineTransform("combineIn") + combineTransformWithTriggerElementCount := makeCombineTransform("combineInWithTriggerElementCount") + combineTransformWithTriggerAlways := makeCombineTransform("combineInWithTriggerAlways") + combineValuesTransform := &pipepb.PTransform{ UniqueName: "combine_values", Subtransforms: []string{ @@ -64,6 +77,14 @@ func TestHandleCombine(t *testing.T) { "combineOut": { CoderId: "outputCoder", }, + "combineInWithTriggerElementCount": { + CoderId: "inputCoder", + WindowingStrategyId: "wsElementCount", + }, + "combineInWithTriggerAlways": { + CoderId: "inputCoder", + WindowingStrategyId: "wsAlways", + }, } baseCoderMap := map[string]*pipepb.Coder{ "int": { @@ -84,7 +105,20 @@ func TestHandleCombine(t *testing.T) { ComponentCoderIds: []string{"int", "string"}, }, } - + baseWindowingStrategyMap := map[string]*pipepb.WindowingStrategy{ + "wsElementCount": makeWindowingStrategy(&pipepb.Trigger{ + Trigger: &pipepb.Trigger_ElementCount_{ + ElementCount: &pipepb.Trigger_ElementCount{ + ElementCount: 10, + }, + }, + }), + "wsAlways": makeWindowingStrategy(&pipepb.Trigger{ + Trigger: &pipepb.Trigger_Always_{ + Always: &pipepb.Trigger_Always{}, + }, + }), + } tests := []struct { name string lifted bool @@ -188,6 +222,32 @@ func TestHandleCombine(t *testing.T) { }, }, }, + }, { + name: "noLift_triggerElementCount", + lifted: true, // Lifting is enabled, but should be disabled in the present of the trigger + comps: &pipepb.Components{ + Transforms: map[string]*pipepb.PTransform{ + undertest: combineTransformWithTriggerElementCount, + "combine_values": combineValuesTransform, + }, + Pcollections: basePCollectionMap, + Coders: baseCoderMap, + WindowingStrategies: baseWindowingStrategyMap, + }, + want: prepareResult{}, + }, { + name: "noLift_triggerAlways", + lifted: true, // Lifting is enabled, but should be disabled in the present of the trigger + comps: &pipepb.Components{ + Transforms: map[string]*pipepb.PTransform{ + undertest: combineTransformWithTriggerAlways, + "combine_values": combineValuesTransform, + }, + Pcollections: basePCollectionMap, + Coders: baseCoderMap, + WindowingStrategies: baseWindowingStrategyMap, + }, + want: prepareResult{}, }, } for _, test := range tests { diff --git a/sdks/go/pkg/beam/runners/prism/internal/unimplemented_test.go b/sdks/go/pkg/beam/runners/prism/internal/unimplemented_test.go index 185940eada14..7a742c22d0fb 100644 --- a/sdks/go/pkg/beam/runners/prism/internal/unimplemented_test.go +++ b/sdks/go/pkg/beam/runners/prism/internal/unimplemented_test.go @@ -49,7 +49,6 @@ func TestUnimplemented(t *testing.T) { // See https://github.com/apache/beam/issues/31153. {pipeline: primitives.TriggerElementCount}, {pipeline: primitives.TriggerOrFinally}, - {pipeline: primitives.TriggerAlways}, // Currently unimplemented triggers. // https://github.com/apache/beam/issues/31438 @@ -87,6 +86,7 @@ func TestImplemented(t *testing.T) { {pipeline: primitives.ParDoProcessElementBundleFinalizer}, {pipeline: primitives.TriggerNever}, + {pipeline: primitives.TriggerAlways}, {pipeline: primitives.Panes}, {pipeline: primitives.TriggerAfterAll}, {pipeline: primitives.TriggerAfterAny},