Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
95 changes: 61 additions & 34 deletions sdks/go/pkg/beam/runners/prism/internal/engine/elementmanager.go
Original file line number Diff line number Diff line change
Expand Up @@ -184,6 +184,8 @@ type Config struct {
MaxBundleSize int
// Whether to use real-time clock as processing time
EnableRTC bool
// Whether to process the data in a streaming mode
StreamingMode bool
}

// ElementManager handles elements, watermarks, and related errata to determine
Expand Down Expand Up @@ -1296,6 +1298,43 @@ func (ss *stageState) AddPending(em *ElementManager, newPending []element) int {
return ss.kind.addPending(ss, em, newPending)
}

func (ss *stageState) injectTriggeredBundlesIfReady(em *ElementManager, window typex.Window, key string) int {
// Check on triggers for this key.
// We use an empty linkID as the key into state for aggregations.
count := 0
if ss.state == nil {
ss.state = make(map[LinkID]map[typex.Window]map[string]StateData)
}
lv, ok := ss.state[LinkID{}]
if !ok {
lv = make(map[typex.Window]map[string]StateData)
ss.state[LinkID{}] = lv
}
wv, ok := lv[window]
if !ok {
wv = make(map[string]StateData)
lv[window] = wv
}
state := wv[key]
endOfWindowReached := window.MaxTimestamp() < ss.input
ready := ss.strat.IsTriggerReady(triggerInput{
newElementCount: 1,
endOfWindowReached: endOfWindowReached,
}, &state)

if ready {
state.Pane = computeNextTriggeredPane(state.Pane, endOfWindowReached)
}
// Store the state as triggers may have changed it.
ss.state[LinkID{}][window][key] = state

// If we're ready, it's time to fire!
if ready {
count += ss.buildTriggeredBundle(em, key, window)
}
return count
}

// addPending for aggregate stages behaves likes stateful stages, but don't need to handle timers or a separate window
// expiration condition.
func (*aggregateStageKind) addPending(ss *stageState, em *ElementManager, newPending []element) int {
Expand All @@ -1315,6 +1354,13 @@ func (*aggregateStageKind) addPending(ss *stageState, em *ElementManager, newPen
if ss.pendingByKeys == nil {
ss.pendingByKeys = map[string]*dataAndTimers{}
}

type windowKey struct {
window typex.Window
key string
}
pendingWindowKeys := set[windowKey]{}

count := 0
for _, e := range newPending {
count++
Expand All @@ -1327,37 +1373,18 @@ func (*aggregateStageKind) addPending(ss *stageState, em *ElementManager, newPen
ss.pendingByKeys[string(e.keyBytes)] = dnt
}
heap.Push(&dnt.elements, e)
// Check on triggers for this key.
// We use an empty linkID as the key into state for aggregations.
if ss.state == nil {
ss.state = make(map[LinkID]map[typex.Window]map[string]StateData)
}
lv, ok := ss.state[LinkID{}]
if !ok {
lv = make(map[typex.Window]map[string]StateData)
ss.state[LinkID{}] = lv
}
wv, ok := lv[e.window]
if !ok {
wv = make(map[string]StateData)
lv[e.window] = wv
}
state := wv[string(e.keyBytes)]
endOfWindowReached := e.window.MaxTimestamp() < ss.input
ready := ss.strat.IsTriggerReady(triggerInput{
newElementCount: 1,
endOfWindowReached: endOfWindowReached,
}, &state)

if ready {
state.Pane = computeNextTriggeredPane(state.Pane, endOfWindowReached)
if em.config.StreamingMode {
// In streaming mode, we check trigger readiness on each element
count += ss.injectTriggeredBundlesIfReady(em, e.window, string(e.keyBytes))
} else {
// In batch mode, we store key + window pairs here and check trigger readiness for each of them later.
pendingWindowKeys.insert(windowKey{window: e.window, key: string(e.keyBytes)})
}
// Store the state as triggers may have changed it.
ss.state[LinkID{}][e.window][string(e.keyBytes)] = state

// If we're ready, it's time to fire!
if ready {
count += ss.buildTriggeredBundle(em, e.keyBytes, e.window)
}
if !em.config.StreamingMode {
for wk := range pendingWindowKeys {
count += ss.injectTriggeredBundlesIfReady(em, wk.window, wk.key)
}
}
return count
Expand Down Expand Up @@ -1493,9 +1520,9 @@ func (ss *stageState) savePanes(bundID string, panesInBundle []bundlePane) {
// buildTriggeredBundle must be called with the stage.mu lock held.
// When in discarding mode, returns 0.
// When in accumulating mode, returns the number of fired elements to maintain a correct pending count.
func (ss *stageState) buildTriggeredBundle(em *ElementManager, key []byte, win typex.Window) int {
func (ss *stageState) buildTriggeredBundle(em *ElementManager, key string, win typex.Window) int {
var toProcess []element
dnt := ss.pendingByKeys[string(key)]

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'll note that Go does some magic when passing a string cast []byte as a map key inline, which avoids allocating. That's why this method took in a []byte for the key, instead of eagerly converting it.

dnt := ss.pendingByKeys[key]
var notYet []element

rb := RunBundle{StageID: ss.ID, BundleID: "agg-" + em.nextBundID(), Watermark: ss.input}
Expand Down Expand Up @@ -1524,7 +1551,7 @@ func (ss *stageState) buildTriggeredBundle(em *ElementManager, key []byte, win t
}
dnt.elements = append(dnt.elements, notYet...)
if dnt.elements.Len() == 0 {
delete(ss.pendingByKeys, string(key))
delete(ss.pendingByKeys, key)
} else {
// Ensure the heap invariants are maintained.
heap.Init(&dnt.elements)
Expand All @@ -1537,15 +1564,15 @@ func (ss *stageState) buildTriggeredBundle(em *ElementManager, key []byte, win t
{
win: win,
key: string(key),
pane: ss.state[LinkID{}][win][string(key)].Pane,
pane: ss.state[LinkID{}][win][key].Pane,
},
}

ss.makeInProgressBundle(
func() string { return rb.BundleID },
toProcess,
ss.input,
singleSet(string(key)),

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

But this over here probably made the string(key) (when key is []byte) moot anyway.

singleSet(key),
nil,
panesInBundle,
)
Expand Down
13 changes: 13 additions & 0 deletions sdks/go/pkg/beam/runners/prism/internal/execute.go
Original file line number Diff line number Diff line change
Expand Up @@ -153,6 +153,7 @@ func executePipeline(ctx context.Context, wks map[string]*worker.W, j *jobservic

topo := prepro.preProcessGraph(comps, j)
ts := comps.GetTransforms()
pcols := comps.GetPcollections()

config := engine.Config{}
m := j.PipelineOptions().AsMap()
Expand All @@ -167,6 +168,18 @@ func executePipeline(ctx context.Context, wks map[string]*worker.W, j *jobservic
}
}

if streaming, ok := m["beam:option:streaming:v1"].(bool); ok {
config.StreamingMode = streaming
}

// Set StreamingMode to true if there is any unbounded PCollection.
for _, pcoll := range pcols {
if pcoll.GetIsBounded() == pipepb.IsBounded_UNBOUNDED {
config.StreamingMode = true
break
}
}

em := engine.NewElementManager(config)

// TODO move this loop and code into the preprocessor instead.
Expand Down
70 changes: 69 additions & 1 deletion sdks/python/apache_beam/runners/portability/prism_runner_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,10 +35,14 @@
import apache_beam as beam
from apache_beam.options.pipeline_options import DebugOptions
from apache_beam.options.pipeline_options import PortableOptions
from apache_beam.options.pipeline_options import StandardOptions
from apache_beam.options.pipeline_options import TypeOptions
from apache_beam.runners.portability import portable_runner_test
from apache_beam.runners.portability import prism_runner
from apache_beam.testing.util import assert_that
from apache_beam.testing.util import equal_to
from apache_beam.transforms import trigger
from apache_beam.transforms import window
from apache_beam.utils import shared

# Run as
Expand All @@ -64,6 +68,8 @@ def __init__(self, *args, **kwargs):
self.environment_type = None
self.environment_config = None
self.enable_commit = False
self.streaming = False
self.allow_unsafe_triggers = False

def setUp(self):
self.enable_commit = False
Expand Down Expand Up @@ -175,6 +181,9 @@ def create_options(self):
options.view_as(
PortableOptions).environment_options = self.environment_options

options.view_as(StandardOptions).streaming = self.streaming
options.view_as(
TypeOptions).allow_unsafe_triggers = self.allow_unsafe_triggers
return options

# Can't read host files from within docker, read a "local" file there.
Expand Down Expand Up @@ -225,7 +234,66 @@ def test_custom_window_type(self):
def test_metrics(self):
super().test_metrics(check_bounded_trie=False)

# Inherits all other tests.
def construct_timestamped(k, t):
return window.TimestampedValue((k, t), t)

def format_result(k, vs):
return ('%s-%s' % (k, len(list(vs))), set(vs))

def test_after_count_trigger_batch(self):
self.allow_unsafe_triggers = True
with self.create_pipeline() as p:
result = (
p
| beam.Create([1, 2, 3, 4, 5, 10, 11])
| beam.FlatMap(lambda t: [('A', t), ('B', t + 5)])
#A1, A2, A3, A4, A5, A10, A11, B6, B7, B8, B9, B10, B15, B16
| beam.MapTuple(PrismRunnerTest.construct_timestamped)
| beam.WindowInto(
window.FixedWindows(10),
trigger=trigger.AfterCount(3),
accumulation_mode=trigger.AccumulationMode.DISCARDING,
)
| beam.GroupByKey()
| beam.MapTuple(PrismRunnerTest.format_result))
assert_that(
result,
equal_to(
list([
('A-5', {1, 2, 3, 4, 5}),
('A-2', {10, 11}),
('B-4', {6, 7, 8, 9}),
('B-3', {10, 15, 16}),
])))

def test_after_count_trigger_streaming(self):
self.allow_unsafe_triggers = True
self.streaming = True
with self.create_pipeline() as p:
result = (
p
| beam.Create([1, 2, 3, 4, 5, 10, 11])
| beam.FlatMap(lambda t: [('A', t), ('B', t + 5)])
#A1, A2, A3, A4, A5, A10, A11, B6, B7, B8, B9, B10, B15, B16
| beam.MapTuple(PrismRunnerTest.construct_timestamped)
| beam.WindowInto(
window.FixedWindows(10),
trigger=trigger.AfterCount(3),
accumulation_mode=trigger.AccumulationMode.DISCARDING,
)
| beam.GroupByKey()
| beam.MapTuple(PrismRunnerTest.format_result))
assert_that(
result,
equal_to(
list([
('A-3', {1, 2, 3}),
('A-2', {4, 5}),
('A-2', {10, 11}),
('B-3', {6, 7, 8}),
('B-1', {9}),
('B-3', {10, 15, 16}),
])))


class PrismJobServerTest(unittest.TestCase):
Expand Down
Loading