diff --git a/sdks/python/apache_beam/io/watch.py b/sdks/python/apache_beam/io/watch.py index f2eadaf4ace8..40b7e451ce6e 100644 --- a/sdks/python/apache_beam/io/watch.py +++ b/sdks/python/apache_beam/io/watch.py @@ -34,6 +34,13 @@ not passed explicitly and converted to its deterministic form, so equal keys hash equally across workers and restarts. +By default, the Watch transform internally stores the hash of all items +seen. If the incremental items returned by the poll function guarantee +monotonic timestamp growth (new items on the next poll have timestamps +larger than the largest of the previous poll), consider setting +``timestamp_cursor=True`` for better performance, as it replaces the hash +dedup with an O(1) event-time cursor; see :class:`Watch`. + Example:: from apache_beam.io.watch import Watch, PollResult, after_total_of @@ -55,8 +62,10 @@ def poll(prefix) -> PollResult[str]: import collections import dataclasses +import enum import hashlib import inspect +import logging import time import typing from collections.abc import Iterable @@ -91,6 +100,8 @@ def poll(prefix) -> PollResult[str]: 'after_total_of', ] +_LOGGER = logging.getLogger(__name__) + _HASH_DIGEST_SIZE = 16 # 128-bit digest width. OutputT = TypeVar('OutputT') @@ -120,6 +131,7 @@ def is_complete(self) -> bool: @staticmethod def _normalize(outputs, timestamp) -> tuple[TimestampedValue, ...]: + # One default timestamp per call, so raw outputs share an event time. if timestamp is None: default_ts = Timestamp.now() else: @@ -137,7 +149,9 @@ def incomplete(outputs: Iterable, timestamp=None) -> 'PollResult': """Reports outputs and expects more; the transform infers the watermark. A raw (non-:class:`TimestampedValue`) output is stamped with ``timestamp`` - when given, else with the current processing time. + when given, else with the current processing time. The inferred watermark + is safe only for non-decreasing event-time enumerations; out-of-order + sources should call :meth:`with_watermark`. """ return PollResult(PollResult._normalize(outputs, timestamp), watermark=None) @@ -146,12 +160,15 @@ def complete(outputs: Iterable, timestamp=None) -> 'PollResult': """Reports the final outputs for an input, after which polling stops. A raw (non-:class:`TimestampedValue`) output is stamped with ``timestamp`` - when given, else with the current processing time. + when given, else with the current processing time. The watermark is + released to ``MAX_TIMESTAMP`` so downstream event-time windows close. """ return PollResult( PollResult._normalize(outputs, timestamp), watermark=MAX_TIMESTAMP) def with_watermark(self, watermark) -> 'PollResult': + """Sets an explicit watermark, a promise that no future output for this + input will have an event time below ``watermark``.""" return dataclasses.replace(self, watermark=Timestamp.of(watermark)) @@ -254,15 +271,17 @@ class _GrowthState: @dataclasses.dataclass(frozen=True) class _PollingGrowthState(_GrowthState): - """Keep-polling state: emitted-output hashes, watermark, termination state. + """Keep-polling state: dedup state, watermark, termination state. ``completed`` maps a 16-byte output-key hash to the event time it was first - seen. It is insertion-ordered and treated as immutable; a new mapping is - built for each residual. + seen; it is insertion-ordered and treated as immutable. In timestamp-cursor + mode ``completed`` is empty and ``cursor`` is the greatest emitted event + time. """ completed: 'collections.OrderedDict[bytes, Timestamp]' poll_watermark: Optional[Timestamp] termination_state: Any + cursor: Optional[Timestamp] = None @dataclasses.dataclass(frozen=True) @@ -305,13 +324,23 @@ def is_deterministic(self) -> bool: return self._tuple_coder.is_deterministic() +class _StateTag(enum.IntEnum): + """Envelope tag selecting the encoded restriction variant.""" + POLLING = 0 + NON_POLLING = 1 + CURSOR_POLLING = 2 + + class _GrowthStateCoder(Coder): """Encodes a :class:`_PollingGrowthState` or :class:`_NonPollingGrowthState`. A ``(tag, payload)`` envelope selects the variant; the payload is a variant-specific :class:`TupleCoder`. ``completed`` is encoded as an ordered - list of ``(hash, timestamp)`` pairs so insertion order survives a round trip. - This format is internal to the Python SDK. + list of ``(hash, timestamp)`` pairs so insertion order survives a round + trip. A cursor state encodes only its termination state and cursor; the + watermark is restored from the estimator state the runner persists. Hash + states keep the pre-cursor byte format. This format is internal to the + Python SDK. """ def __init__(self, output_coder: Coder, termination: TerminationCondition): nullable_ts = NullableCoder(TimestampCoder()) @@ -322,6 +351,10 @@ def __init__(self, output_coder: Coder, termination: TerminationCondition): nullable_ts, coders.ListCoder(TupleCoder([coders.BytesCoder(), TimestampCoder()])), ]) + self._cursor_polling_coder = TupleCoder([ + termination.state_coder(), + TimestampCoder(), + ]) self._non_polling_coder = TupleCoder([ nullable_ts, coders.ListCoder(_TimestampedValueCoder(output_coder)), @@ -329,25 +362,33 @@ def __init__(self, output_coder: Coder, termination: TerminationCondition): def encode(self, state: _GrowthState) -> bytes: if isinstance(state, _PollingGrowthState): - payload = self._polling_coder.encode(( - state.termination_state, - state.poll_watermark, - list(state.completed.items()))) - return self._envelope_coder.encode((0, payload)) + if state.cursor is None: + payload = self._polling_coder.encode(( + state.termination_state, + state.poll_watermark, + list(state.completed.items()))) + return self._envelope_coder.encode((_StateTag.POLLING, payload)) + payload = self._cursor_polling_coder.encode( + (state.termination_state, state.cursor)) + return self._envelope_coder.encode((_StateTag.CURSOR_POLLING, payload)) payload = self._non_polling_coder.encode( (state.pending.watermark, list(state.pending.outputs))) - return self._envelope_coder.encode((1, payload)) + return self._envelope_coder.encode((_StateTag.NON_POLLING, payload)) def decode(self, encoded: bytes) -> _GrowthState: tag, payload = self._envelope_coder.decode(encoded) - if tag == 0: + if tag == _StateTag.POLLING: termination_state, poll_watermark, items = self._polling_coder.decode( payload) return _PollingGrowthState( collections.OrderedDict(items), poll_watermark, termination_state) - if tag == 1: + if tag == _StateTag.NON_POLLING: watermark, outputs = self._non_polling_coder.decode(payload) return _NonPollingGrowthState(PollResult(tuple(outputs), watermark)) + if tag == _StateTag.CURSOR_POLLING: + termination_state, cursor = self._cursor_polling_coder.decode(payload) + return _PollingGrowthState( + collections.OrderedDict(), None, termination_state, cursor) raise ValueError('unknown Watch growth state tag: %r' % (tag, )) def is_deterministic(self) -> bool: @@ -400,6 +441,30 @@ def _never_seen_before( return dataclasses.replace(result, outputs=tuple(new_outputs)) +def _cursor_of(restriction: _PollingGrowthState) -> Optional[Timestamp]: + """The dedup cursor: the stored one, or for a restriction switched over + from hash dedup, the greatest event time its hash map recorded.""" + if restriction.cursor is not None: + return restriction.cursor + if restriction.completed: + return max(restriction.completed.values()) + return None + + +def _past_cursor( + restriction: _PollingGrowthState, result: PollResult) -> PollResult: + """Filters a poll result down to outputs strictly past the cursor, sorted + by timestamp so the earliest infers the watermark and the latest advances + the cursor.""" + cursor = _cursor_of(restriction) + new_outputs = [ + output for output in result.outputs + if cursor is None or output.timestamp > cursor + ] + new_outputs.sort(key=lambda output: output.timestamp) + return dataclasses.replace(result, outputs=tuple(new_outputs)) + + class _GrowthRestrictionTracker(iobase.RestrictionTracker): """Tracks one input's polling restriction over claimed poll rounds. @@ -413,10 +478,12 @@ def __init__( self, restriction: _GrowthState, key_fn: Callable[[Any], Any], - key_coder: Coder): + key_coder: Coder, + timestamp_cursor: bool = False): self._restriction = restriction self._key_fn = key_fn self._key_coder = key_coder + self._timestamp_cursor = timestamp_cursor self._claimed_result = None # type: Optional[PollResult] self._claimed_termination_state = None # type: Any self._claimed_hashes = None # type: Optional[collections.OrderedDict] @@ -438,19 +505,35 @@ def try_claim(self, position: tuple[PollResult, Any]) -> bool: if self._should_stop: return False result, termination_state = position - claimed_hashes = collections.OrderedDict() - for output in result.outputs: - claimed_hashes[self._hash(output.value)] = output.timestamp - if isinstance(self._restriction, _PollingGrowthState): - if any(key_hash in self._restriction.completed - for key_hash in claimed_hashes): - return False + claimed_hashes = None + if self._timestamp_cursor: + # Cursor mode validates by timestamps and never hashes. + if isinstance(self._restriction, _PollingGrowthState): + cursor = _cursor_of(self._restriction) + if cursor is not None and any(output.timestamp <= cursor + for output in result.outputs): + return False + else: + # Values may lack stable equality without a deterministic coder, so a + # replay is identified by its timestamps. + expected = sorted( + output.timestamp for output in self._restriction.pending.outputs) + if expected != sorted(output.timestamp for output in result.outputs): + return False else: - expected = set( - self._hash(output.value) - for output in self._restriction.pending.outputs) - if expected != set(claimed_hashes): - return False + claimed_hashes = collections.OrderedDict() + for output in result.outputs: + claimed_hashes[self._hash(output.value)] = output.timestamp + if isinstance(self._restriction, _PollingGrowthState): + if any(key_hash in self._restriction.completed + for key_hash in claimed_hashes): + return False + else: + expected = set( + self._hash(output.value) + for output in self._restriction.pending.outputs) + if expected != set(claimed_hashes): + return False self._should_stop = True self._claimed_result = result self._claimed_termination_state = termination_state @@ -470,14 +553,31 @@ def try_split(self, fraction_of_remainder): residual = _EMPTY_STATE else: # The primary becomes a replay of the claimed round; the residual - # resumes polling with the claimed keys marked completed. - merged = collections.OrderedDict(self._restriction.completed) - merged.update(self._claimed_hashes) + # resumes polling with the claimed round folded into the dedup state. + # A state holds hashes or a cursor, never both, so each mode drops the + # other mode's leftovers after a switch. + if self._timestamp_cursor: + completed = self._restriction.completed + if completed: + completed = collections.OrderedDict() + if self._claimed_result.outputs: + cursor = self._claimed_result.outputs[-1].timestamp + else: + cursor = _cursor_of(self._restriction) + elif self._claimed_hashes: + completed = collections.OrderedDict(self._restriction.completed) + completed.update(self._claimed_hashes) + cursor = None + else: + # An idle round reuses the parent map so empty polls stay O(1). + completed = self._restriction.completed + cursor = None residual = _PollingGrowthState( - merged, + completed, _max_watermark( self._restriction.poll_watermark, self._claimed_result.watermark), - self._claimed_termination_state) + self._claimed_termination_state, + cursor) self._restriction = _NonPollingGrowthState(self._claimed_result) self._should_stop = True return self._restriction, residual @@ -522,6 +622,7 @@ def __init__( output_coder: Coder, key_fn: Callable[[Any], Any], key_coder: Coder, + timestamp_cursor: bool = False, now_fn: Optional[Callable[[], float]] = None): self._poll_fn = poll_fn self._termination = termination @@ -529,8 +630,11 @@ def __init__( self._output_coder = output_coder self._key_fn = key_fn self._key_coder = key_coder + self._timestamp_cursor = timestamp_cursor self._now = now_fn or time.time self._restriction_coder = _GrowthStateCoder(output_coder, termination) + # Count of late emissions seen on this worker, for throttled warnings. + self._late_count = 0 def initial_restriction(self, element) -> _PollingGrowthState: now = Timestamp.of(self._now()) @@ -540,7 +644,8 @@ def initial_restriction(self, element) -> _PollingGrowthState: self._termination.for_new_input(now, element)) def create_tracker(self, restriction) -> _GrowthRestrictionTracker: - return _GrowthRestrictionTracker(restriction, self._key_fn, self._key_coder) + return _GrowthRestrictionTracker( + restriction, self._key_fn, self._key_coder, self._timestamp_cursor) def restriction_coder(self) -> Coder: return self._restriction_coder @@ -570,13 +675,21 @@ def process( for output in restriction.pending.outputs: yield TimestampedValue((element, output.value), output.timestamp) return + if (self._timestamp_cursor and restriction.cursor is not None and + restriction.cursor >= MAX_TIMESTAMP): + # Nothing can be past a cursor at MAX; claim an empty round and stop. + tracker.try_claim((PollResult(()), restriction.termination_state)) + return # Poll before claiming so a slow poll never holds the tracker lock, which # would block runner progress checks and checkpoints. result = self._poll_fn(element) # Read the clock after the poll so a slow poll counts against termination. now = Timestamp.of(self._now()) - new_results = _never_seen_before( - restriction, result, self._key_fn, self._key_coder) + if self._timestamp_cursor: + new_results = _past_cursor(restriction, result) + else: + new_results = _never_seen_before( + restriction, result, self._key_fn, self._key_coder) termination_state = restriction.termination_state if new_results.outputs: termination_state = self._termination.on_seen_new_output( @@ -585,7 +698,15 @@ def process( if not tracker.try_claim((new_results, termination_state)): # A checkpoint already stopped this invocation; emit nothing. return + # Emit before advancing the watermark so a round's own watermark cannot + # make its outputs late. Late outputs are warned about only once the + # watermark has advanced past the element-timestamp seed. + current_watermark = watermark_estimator.current_watermark() + warn_on_late = ( + current_watermark is not None and current_watermark > timestamp) for output in new_results.outputs: + if warn_on_late and output.timestamp < current_watermark: + self._warn_late(element, output.timestamp, current_watermark) yield TimestampedValue((element, output.value), output.timestamp) if new_results.watermark is not None: watermark = new_results.watermark @@ -594,6 +715,13 @@ def process( watermark = new_results.outputs[0].timestamp else: watermark = None + if self._timestamp_cursor: + new_cursor = ( + new_results.outputs[-1].timestamp + if new_results.outputs else restriction.cursor) + if new_cursor is not None and new_cursor >= MAX_TIMESTAMP: + # A cursor at MAX is terminal; polling on would only drop outputs. + return if self._termination.can_stop_polling(now, termination_state): return if watermark is not None and watermark >= MAX_TIMESTAMP: @@ -603,6 +731,20 @@ def process( _set_watermark_if_greater(watermark_estimator, watermark) tracker.defer_remainder(self._poll_interval) + def _warn_late(self, element, output_timestamp, watermark) -> None: + # Log at powers of two to keep an ongoing problem visible without spam. + self._late_count += 1 + if self._late_count & (self._late_count - 1) == 0: + _LOGGER.warning( + 'Watch emitted output for input %r at %s, behind the watermark %s; ' + 'downstream event-time windowing may drop it as late. Use ' + 'PollResult.with_watermark for out-of-order sources. ' + '(%d late emissions on this worker)', + element, + output_timestamp, + watermark, + self._late_count) + def _set_watermark_if_greater(watermark_estimator, new_watermark) -> None: # set_watermark raises on regression, so only ever advance the watermark. @@ -670,6 +812,14 @@ class Watch(PTransform): inferred like ``output_coder`` when omitted. It is converted with ``as_deterministic_coder`` so equal keys always hash equally; a coder with no deterministic form is rejected. + timestamp_cursor: dedup by event time instead of by key. Each round emits + only outputs strictly past the greatest event time already emitted, so + the per-input state is a single timestamp. Requires every new output to + carry an event time strictly greater than all previously emitted ones; + re-listed old outputs at or below the cursor are dropped as already + seen. For sources whose new outputs can arrive at or below the cursor, + keep the default hash dedup. Incompatible with ``output_key_fn`` and + ``output_key_coder``. now_fn: clock used for termination decisions; tests can inject one. """ def __init__( @@ -680,16 +830,23 @@ def __init__( output_coder: Optional[Coder] = None, output_key_fn: Optional[Callable[[Any], Any]] = None, output_key_coder: Optional[Coder] = None, + timestamp_cursor: bool = False, now_fn: Optional[Callable[[], float]] = None): super().__init__() if poll_interval is None: raise ValueError('Watch requires a poll_interval') + if timestamp_cursor and (output_key_fn is not None or + output_key_coder is not None): + raise ValueError( + 'timestamp_cursor dedups by event time, not by key; do not pass ' + 'output_key_fn or output_key_coder with timestamp_cursor=True.') self._poll_fn = poll_fn self._poll_interval = _as_duration(poll_interval) self._termination = termination or never() self._output_coder = output_coder self._output_key_fn = output_key_fn self._output_key_coder = output_key_coder + self._timestamp_cursor = timestamp_cursor self._now = now_fn def expand(self, pcoll): @@ -698,22 +855,28 @@ def expand(self, pcoll): output_coder = self._poll_fn.default_output_coder() if output_coder is None: output_coder = _coder_for_hint(_poll_output_type(self._poll_fn)) - if self._output_key_fn is None: - # The output is its own dedup key, so the key coder is the output coder. + if self._timestamp_cursor: + # Cursor dedup never hashes, so no deterministic key coder is needed. key_fn = _identity - key_coder = self._output_key_coder or output_coder + key_coder = output_coder else: - key_fn = self._output_key_fn - key_coder = self._output_key_coder or _coder_for_hint( - _return_type(self._output_key_fn)) - # Dedup hashes the encoded key, so equal keys must encode equally; use the - # coder's deterministic form and reject coders that have none. - key_coder = key_coder.as_deterministic_coder( - self.label, - 'Watch dedups by hashing the encoded output key, so the key coder ' - 'must be deterministic. %s has no deterministic form; pass a ' - 'deterministic output_key_coder (or output_coder).' % - type(key_coder).__name__) + if self._output_key_fn is None: + # The output is its own dedup key, so the key coder is the output + # coder. + key_fn = _identity + key_coder = self._output_key_coder or output_coder + else: + key_fn = self._output_key_fn + key_coder = self._output_key_coder or _coder_for_hint( + _return_type(self._output_key_fn)) + # Dedup hashes the encoded key, so equal keys must encode equally; use + # the coder's deterministic form and reject coders that have none. + key_coder = key_coder.as_deterministic_coder( + self.label, + 'Watch dedups by hashing the encoded output key, so the key coder ' + 'must be deterministic. %s has no deterministic form; pass a ' + 'deterministic output_key_coder (or output_coder).' % + type(key_coder).__name__) # Type the (input, output) pairs from the input type and the resolved # coder's type, so downstream transforms are typed and coder inference does # not fall back to pickling. @@ -730,6 +893,7 @@ def expand(self, pcoll): output_coder, key_fn, key_coder, + self._timestamp_cursor, self._now)).with_output_types(tuple[input_type, value_type]) diff --git a/sdks/python/apache_beam/io/watch_test.py b/sdks/python/apache_beam/io/watch_test.py index 472177ceaee6..a07f98bfa8d7 100644 --- a/sdks/python/apache_beam/io/watch_test.py +++ b/sdks/python/apache_beam/io/watch_test.py @@ -22,8 +22,14 @@ import unittest import apache_beam as beam +from apache_beam.coders.coders import BytesCoder from apache_beam.coders.coders import Coder +from apache_beam.coders.coders import ListCoder +from apache_beam.coders.coders import NullableCoder from apache_beam.coders.coders import StrUtf8Coder +from apache_beam.coders.coders import TimestampCoder +from apache_beam.coders.coders import TupleCoder +from apache_beam.coders.coders import VarIntCoder from apache_beam.io.watch import PollFn from apache_beam.io.watch import PollResult from apache_beam.io.watch import Watch @@ -31,6 +37,7 @@ from apache_beam.io.watch import _GrowthStateCoder from apache_beam.io.watch import _never_seen_before from apache_beam.io.watch import _NonPollingGrowthState +from apache_beam.io.watch import _past_cursor from apache_beam.io.watch import _PollingGrowthState from apache_beam.io.watch import _WatchGrowthDoFn from apache_beam.io.watch import after_total_of @@ -70,12 +77,45 @@ def _tracker(restriction): return _GrowthRestrictionTracker(restriction, _identity, StrUtf8Coder()) +def _cursor_tracker(restriction): + return _GrowthRestrictionTracker( + restriction, _identity, StrUtf8Coder(), timestamp_cursor=True) + + def _initial_polling(termination=None, now=Timestamp(0)): termination = termination or never() return _PollingGrowthState( collections.OrderedDict(), None, termination.for_new_input(now, 'input')) +class PollResultTest(unittest.TestCase): + def test_normalize_stamps_one_processing_time_when_timestamp_none(self): + before = Timestamp.now() + result = PollResult.incomplete(['a', 'b']) + after = Timestamp.now() + # Raw outputs share a single processing-time stamp (no per-output jitter). + stamps = {o.timestamp for o in result.outputs} + self.assertEqual(1, len(stamps)) + ts = stamps.pop() + self.assertTrue(before <= ts <= after) + + def test_normalize_preserves_timestamped_and_applies_explicit_default(self): + result = PollResult.incomplete([_ts('a', 1), 'b'], timestamp=7) + by_value = {o.value: o.timestamp for o in result.outputs} + self.assertEqual(Timestamp(1), by_value['a']) # TimestampedValue preserved + self.assertEqual(Timestamp(7), by_value['b']) # raw stamped with default + + def test_complete_releases_watermark_to_max(self): + self.assertEqual( + MAX_TIMESTAMP, PollResult.complete([_ts('a', 1)]).watermark) + self.assertTrue(PollResult.complete([]).is_complete) + + def test_with_watermark_overrides(self): + self.assertEqual( + Timestamp(0), + PollResult.incomplete([_ts('a', 9)]).with_watermark(0).watermark) + + class GrowthStateCoderTest(unittest.TestCase): def test_polling_round_trip_preserves_resume_state(self): termination = after_total_of(Duration(30)) @@ -91,6 +131,42 @@ def test_polling_round_trip_preserves_resume_state(self): self.assertEqual(list(completed.items()), list(decoded.completed.items())) self.assertEqual(Timestamp(5), decoded.poll_watermark) self.assertEqual(termination_state, decoded.termination_state) + self.assertIsNone(decoded.cursor) + + def test_polling_round_trip_preserves_cursor(self): + coder = _GrowthStateCoder(StrUtf8Coder(), never()) + state = _PollingGrowthState( + collections.OrderedDict(), + Timestamp(5), + never().for_new_input(Timestamp(0), 'input'), + Timestamp(42)) + decoded = coder.decode(coder.encode(state)) + self.assertEqual(Timestamp(42), decoded.cursor) + self.assertEqual(0, len(decoded.completed)) + self.assertIsNone(decoded.poll_watermark) # not part of the payload + + def test_cursorless_state_keeps_the_pre_cursor_byte_format(self): + # A polling state without a cursor must encode exactly as before the + # cursor existed, so in-flight hash-mode restrictions decode across an + # upgrade in either direction. + termination = never() + coder = _GrowthStateCoder(StrUtf8Coder(), termination) + completed = collections.OrderedDict([(b'a' * 16, Timestamp(1))]) + termination_state = termination.for_new_input(Timestamp(0), 'input') + state = _PollingGrowthState(completed, Timestamp(5), termination_state) + legacy_polling_coder = TupleCoder([ + termination.state_coder(), + NullableCoder(TimestampCoder()), + ListCoder(TupleCoder([BytesCoder(), TimestampCoder()])), + ]) + legacy_payload = legacy_polling_coder.encode( + (termination_state, Timestamp(5), list(completed.items()))) + legacy_encoded = TupleCoder([VarIntCoder(), BytesCoder()]).encode( + (0, legacy_payload)) + self.assertEqual(legacy_encoded, coder.encode(state)) + decoded = coder.decode(legacy_encoded) + self.assertEqual(list(completed.items()), list(decoded.completed.items())) + self.assertIsNone(decoded.cursor) def test_non_polling_round_trip_preserves_pending_outputs(self): coder = _GrowthStateCoder(StrUtf8Coder(), never()) @@ -215,6 +291,167 @@ def test_wrapper_chain_defers_merged_residual(self): self.assertIsInstance(residual, _PollingGrowthState) self.assertEqual(2, len(residual.completed)) + def test_idle_round_reuses_completed_map_object(self): + # A round that discovers nothing must reuse the parent dedup map rather + # than copying it O(N), so a steady-state empty poll stays cheap. + state = _initial_polling() + first = _new_results(state, PollResult.incomplete([_ts('a', 1)])) + tracker = _tracker(state) + self.assertTrue(tracker.try_claim((first, 0))) + _, residual1 = tracker.try_split(0) + resumed = _tracker(residual1) + empty = _new_results(residual1, PollResult.incomplete([])) + self.assertTrue(resumed.try_claim((empty, 0))) + _, residual2 = resumed.try_split(0) + self.assertIs(residual1.completed, residual2.completed) + + +class TimestampCursorTest(unittest.TestCase): + """Cursor-mode dedup: high-water-mark timestamp instead of a hash set.""" + def test_keeps_state_o1_and_tracks_high_water_mark(self): + state = _initial_polling() + result = PollResult.incomplete([_ts('a', 1), _ts('b', 2), _ts('c', 3)]) + new_results = _past_cursor(state, result) + self.assertEqual(['a', 'b', 'c'], [o.value for o in new_results.outputs]) + tracker = _cursor_tracker(state) + self.assertTrue(tracker.try_claim((new_results, 0))) + _, residual = tracker.try_split(0) + self.assertIsInstance(residual, _PollingGrowthState) + self.assertEqual(0, len(residual.completed)) # no hash set + self.assertEqual(Timestamp(3), residual.cursor) # high-water mark + + def test_emits_only_outputs_after_the_cursor(self): + # A later round emits only outputs strictly past the cursor; a re-listed + # output (== cursor) and an earlier output (< cursor) are both dropped. + state = _initial_polling() + tracker = _cursor_tracker(state) + first = _past_cursor(state, PollResult.incomplete([_ts('a', 10)])) + self.assertTrue(tracker.try_claim((first, 0))) + _, residual = tracker.try_split(0) + self.assertEqual(Timestamp(10), residual.cursor) + second = _past_cursor( + residual, + PollResult.incomplete([_ts('early', 5), _ts('a', 10), _ts('c', 20)])) + self.assertEqual(['c'], [o.value for o in second.outputs]) # only 20 > 10 + resumed = _cursor_tracker(residual) + self.assertTrue(resumed.try_claim((second, 0))) + _, residual = resumed.try_split(0) + self.assertEqual(Timestamp(20), residual.cursor) + + def test_relist_emits_each_output_exactly_once(self): + # A full re-list of a growing collection at strictly increasing event + # times emits each output once; the state never accumulates a hash set. + state = _initial_polling() + emitted = collections.Counter() + for round_index in range(10): + result = PollResult.incomplete( + [_ts('f%d' % i, i + 1) for i in range(round_index + 1)]) + new_results = _past_cursor(state, result) + tracker = _cursor_tracker(state) + self.assertTrue(tracker.try_claim((new_results, 0))) + for output in new_results.outputs: + emitted[output.value] += 1 + _, state = tracker.try_split(0) + self.assertEqual(0, len(state.completed)) # O(1) throughout + self.assertEqual([1] * 10, [emitted['f%d' % i] for i in range(10)]) + self.assertEqual(Timestamp(10), state.cursor) + + def test_round_below_high_water_mark_keeps_cursor_and_reuses_state(self): + # A round whose outputs are all at or below the cursor emits nothing and + # leaves the cursor unchanged; the (empty) completed map is reused as-is. + state = _initial_polling() + tracker = _cursor_tracker(state) + first = _past_cursor(state, PollResult.incomplete([_ts('a', 10)])) + self.assertTrue(tracker.try_claim((first, 0))) + _, residual1 = tracker.try_split(0) + stale = _past_cursor( + residual1, PollResult.incomplete([_ts('a', 10), _ts('old', 4)])) + self.assertEqual((), stale.outputs) + resumed = _cursor_tracker(residual1) + self.assertTrue(resumed.try_claim((stale, 0))) + _, residual2 = resumed.try_split(0) + self.assertEqual(Timestamp(10), residual2.cursor) # unchanged + self.assertIs(residual1.completed, residual2.completed) + + def test_claim_rejects_outputs_at_or_below_the_cursor(self): + # The tracker re-validates a claim, so a round that was not filtered + # against the cursor is rejected instead of emitting already-seen outputs. + state = _initial_polling() + tracker = _cursor_tracker(state) + first = _past_cursor(state, PollResult.incomplete([_ts('a', 10)])) + self.assertTrue(tracker.try_claim((first, 0))) + _, residual = tracker.try_split(0) + stale = PollResult.incomplete([_ts('a', 10)]) + self.assertFalse(_cursor_tracker(residual).try_claim((stale, 0))) + + def test_replay_validates_by_timestamps(self): + # Cursor mode never hashes, so a replay is validated by its timestamps. + pending = PollResult((_ts('a', 1), _ts('b', 2)), MAX_TIMESTAMP) + tracker = _cursor_tracker(_NonPollingGrowthState(pending)) + partial = PollResult((_ts('a', 1), ), None) + self.assertFalse(tracker.try_claim((partial, None))) + self.assertTrue(tracker.try_claim((pending, None))) + + def test_switching_hash_state_to_cursor_drops_the_hash_map(self): + # A restriction carried over from hash dedup still holds completed hashes; + # cursor mode ignores them, so the first cursor round must drop them and + # make the state O(1) rather than carry dead hashes forever. + legacy = _PollingGrowthState( + collections.OrderedDict([(b'a' * 16, Timestamp(1))]), + None, + never().for_new_input(Timestamp(0), 'input')) + result = _past_cursor(legacy, PollResult.incomplete([_ts('a', 100)])) + tracker = _cursor_tracker(legacy) + self.assertTrue(tracker.try_claim((result, 0))) + _, residual = tracker.try_split(0) + self.assertEqual(0, len(residual.completed)) + self.assertEqual(Timestamp(100), residual.cursor) + + def test_switching_hash_state_to_cursor_seeds_the_cursor(self): + # Outputs at or below the hash map's greatest recorded event time are + # already seen and must not re-emit after the switch. + legacy = _PollingGrowthState( + collections.OrderedDict([(b'a' * 16, Timestamp(5)), + (b'b' * 16, Timestamp(10))]), + None, + never().for_new_input(Timestamp(0), 'input')) + relist = PollResult.incomplete([_ts('a', 5), _ts('b', 10), _ts('c', 20)]) + new_results = _past_cursor(legacy, relist) + self.assertEqual(['c'], [o.value for o in new_results.outputs]) + tracker = _cursor_tracker(legacy) + self.assertTrue(tracker.try_claim((new_results, 0))) + _, residual = tracker.try_split(0) + self.assertEqual(0, len(residual.completed)) + self.assertEqual(Timestamp(20), residual.cursor) + + def test_hash_round_drops_a_stale_cursor(self): + # The reverse switch: a hash round drops the cursor, so a state never + # holds hashes and a cursor at the same time. + state = _PollingGrowthState( + collections.OrderedDict(), None, 0, cursor=Timestamp(10)) + tracker = _tracker(state) + result = _new_results(state, PollResult.incomplete([_ts('a', 20)])) + self.assertTrue(tracker.try_claim((result, 0))) + _, residual = tracker.try_split(0) + self.assertIsNone(residual.cursor) + self.assertEqual(1, len(residual.completed)) + + def test_cursor_state_encoding_size_is_independent_of_outputs(self): + coder = _GrowthStateCoder(StrUtf8Coder(), never()) + + def encoded_residual_after_claiming(count): + state = _initial_polling() + result = PollResult.incomplete( + [_ts('output%d' % i, i + 1) for i in range(count)]) + tracker = _cursor_tracker(state) + self.assertTrue(tracker.try_claim((_past_cursor(state, result), 0))) + _, residual = tracker.try_split(0) + return coder.encode(residual) + + self.assertEqual( + len(encoded_residual_after_claiming(1)), + len(encoded_residual_after_claiming(100))) + class TerminationConditionTest(unittest.TestCase): def test_never_does_not_stop(self): @@ -256,6 +493,20 @@ def _empty_poll(unused_element): return PollResult.incomplete([]) +def _out_of_order_poll(prefix): + # Round 1 emits late_after@10 (advances the watermark to 10); round 2 emits + # early@5, which is behind the watermark and therefore late. + _POLL_CALLS[prefix] += 1 + if _POLL_CALLS[prefix] == 1: + return PollResult.incomplete([_ts(prefix + 'late_after', 10)]) + return PollResult.complete([_ts(prefix + 'early', 5)]) + + +def _max_timestamp_poll(unused_element): + return PollResult.incomplete( + [_ts('a', 10), TimestampedValue('b', MAX_TIMESTAMP)]) + + def _keyed_poll(prefix): # 'a1' and 'a2' share the dedup key 'a', so only 'a1' is emitted. return PollResult.complete([_ts('a1', 1), _ts('a2', 2), _ts('b1', 3)]) @@ -286,14 +537,21 @@ def _windowed_group(kv, window=beam.DoFn.WindowParam): class WatchDoFnProcessTest(unittest.TestCase): def _process( - self, poll_fn, element, timestamp, restriction=None, watermark=None): + self, + poll_fn, + element, + timestamp, + restriction=None, + watermark=None, + timestamp_cursor=False): dofn = _WatchGrowthDoFn( poll_fn, never(), Duration(1), StrUtf8Coder(), _identity, - StrUtf8Coder()) + StrUtf8Coder(), + timestamp_cursor) if restriction is None: restriction = dofn.initial_restriction(element) threadsafe = ThreadsafeRestrictionTracker(dofn.create_tracker(restriction)) @@ -368,6 +626,107 @@ def test_terminal_round_after_deferring_leaves_no_residual(self): self.assertIsNone(threadsafe.deferred_status()) self.assertTrue(threadsafe.check_done()) + def test_cursor_at_max_timestamp_stops_polling(self): + # A cursor reaching MAX is terminal: nothing can be strictly past it, so + # the round stops instead of polling forever and dropping every output. + outputs, threadsafe, _ = self._process( + _max_timestamp_poll, 'k:', Timestamp(0), timestamp_cursor=True) + self.assertEqual([('k:', 'a'), ('k:', 'b')], + [value.value for value in outputs]) + self.assertIsNone(threadsafe.deferred_status()) + self.assertTrue(threadsafe.check_done()) + + def test_resumed_cursor_at_max_stops_without_polling(self): + # A restriction resumed with the cursor already at MAX (persisted by a + # checkpoint after a MAX-timestamped round) must stop without invoking the + # poll function at all. + polls = [] + + def poll(unused_element): + polls.append(1) + return PollResult.incomplete([]) + + resumed = _PollingGrowthState( + collections.OrderedDict(), + None, + never().for_new_input(Timestamp(0), 'input'), + MAX_TIMESTAMP) + outputs, threadsafe, _ = self._process( + poll, 'k:', Timestamp(0), restriction=resumed, timestamp_cursor=True) + self.assertEqual([], outputs) + self.assertEqual([], polls) # the poll function never ran + self.assertIsNone(threadsafe.deferred_status()) + self.assertTrue(threadsafe.check_done()) + + def test_out_of_order_new_output_emits_late_and_warns(self): + # Round 1 surfaces late_after@10 and parks the watermark there; round 2 + # surfaces a brand-new early@5. The output is emitted at its true (earlier) + # time, so it is late for downstream windowing, and Watch warns about it. + _POLL_CALLS.clear() + _, threadsafe, estimator = self._process( + _out_of_order_poll, 'k:', Timestamp(0)) + self.assertEqual(Timestamp(10), estimator.current_watermark()) + residual, _ = threadsafe.deferred_status() + with self.assertLogs('apache_beam.io.watch', level='WARNING') as logs: + outputs, _, _ = self._process( + _out_of_order_poll, + 'k:', + Timestamp(0), + restriction=residual, + watermark=estimator.current_watermark()) + self.assertEqual([('k:', 'k:early')], [value.value for value in outputs]) + self.assertEqual([Timestamp(5)], [value.timestamp for value in outputs]) + self.assertTrue( + any('behind the watermark' in line for line in logs.output), + 'expected a late-emission warning, got: %s' % logs.output) + + def test_first_round_early_output_does_not_warn(self): + # While the estimator holds the input element's timestamp seed, an output + # behind it must not trigger the out-of-order warning: the seed is not a + # poll-order signal. + def poll(unused_element): + return PollResult.incomplete([_ts('a', 5)]) + + with self.assertNoLogs('apache_beam.io.watch', level='WARNING'): + outputs, _, _ = self._process(poll, 'k:', Timestamp(10)) + self.assertEqual([Timestamp(5)], [value.timestamp for value in outputs]) + + def test_early_output_after_empty_poll_does_not_warn(self): + # An empty first poll defers with the watermark still at the element seed; + # the next round's first real output must not be treated as out-of-order + # either; the watermark has not advanced past the seed. + polls = [] + + def poll(unused_element): + polls.append(len(polls)) + if len(polls) == 1: + return PollResult.incomplete([]) + return PollResult.incomplete([_ts('a', 5)]) + + _, threadsafe, estimator = self._process(poll, 'k:', Timestamp(10)) + self.assertEqual(Timestamp(10), estimator.current_watermark()) + residual, _ = threadsafe.deferred_status() + with self.assertNoLogs('apache_beam.io.watch', level='WARNING'): + outputs, _, _ = self._process( + poll, + 'k:', + Timestamp(10), + restriction=residual, + watermark=estimator.current_watermark()) + self.assertEqual([Timestamp(5)], [value.timestamp for value in outputs]) + + def test_explicit_watermark_holds_below_output_time(self): + # An explicit watermark below the output's own event time is honored, so + # a later, earlier-timestamped output stays on time (the out-of-order-safe + # path). + def poll(unused_element): + return PollResult.incomplete([_ts('a', 10)]).with_watermark(0) + + _, threadsafe, estimator = self._process(poll, 'k:', Timestamp(0)) + self.assertEqual(Timestamp(0), estimator.current_watermark()) + residual, _ = threadsafe.deferred_status() + self.assertEqual(Timestamp(0), residual.poll_watermark) + class WatchEndToEndTest(unittest.TestCase): def _in_memory_pipeline(self): @@ -417,6 +776,30 @@ def test_multi_round_dedups_stops_and_is_per_input(self): self.assertEqual(3, _POLL_CALLS['x:']) self.assertEqual(3, _POLL_CALLS['y:']) + def test_timestamp_cursor_dedups_growing_source(self): + _POLL_CALLS.clear() + with self._in_memory_pipeline() as p: + output = ( + p | beam.Create(['x:', 'y:']) + | Watch( + _growing_poll, + poll_interval=Duration(0.05), + timestamp_cursor=True)) + # Each output is emitted exactly once via the high-water-mark cursor, + # with no hash set kept, across poll rounds and checkpoints. + assert_that( + output, + equal_to([('x:', 'x:0'), ('x:', 'x:1'), ('x:', 'x:2'), ('y:', 'y:0'), + ('y:', 'y:1'), ('y:', 'y:2')])) + + def test_timestamp_cursor_rejects_key_spec(self): + with self.assertRaises(ValueError): + Watch( + _complete_poll, + poll_interval=Duration(1), + output_key_fn=_first_char, + timestamp_cursor=True) + def test_output_key_dedups_across_pipeline(self): with self._in_memory_pipeline() as p: output = (