amoghrajesh commented on code in PR #71135:
URL: https://github.com/apache/airflow/pull/71135#discussion_r3869582537


##########
providers/amazon/src/airflow/providers/amazon/aws/triggers/kinesis.py:
##########
@@ -0,0 +1,336 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#   http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied.  See the License for the
+# specific language governing permissions and limitations
+# under the License.
+from __future__ import annotations
+
+import asyncio
+import base64
+import hashlib
+import json
+from collections.abc import AsyncIterator
+from typing import TYPE_CHECKING, Any
+
+from airflow.providers.amazon.aws.hooks.kinesis import KinesisHook
+from airflow.providers.amazon.version_compat import AIRFLOW_V_3_0_PLUS
+
+if AIRFLOW_V_3_0_PLUS:
+    from airflow.triggers.base import BaseEventTrigger, TriggerEvent
+else:
+    from airflow.triggers.base import (  # type: ignore
+        BaseTrigger as BaseEventTrigger,
+        TriggerEvent,
+    )
+
+if TYPE_CHECKING:
+    from airflow.providers.amazon.aws.hooks.base_aws import BaseAwsConnection
+
+_CHECKPOINT_KEY_PREFIX = "kinesis_shard_sequence_numbers"
+_ITERATOR_TYPES_WITHOUT_EXTRA_ARGS = frozenset({"LATEST", "TRIM_HORIZON"})
+
+
+class KinesisTrigger(BaseEventTrigger):
+    """
+    Wait asynchronously for records on an Amazon Kinesis Data Stream.
+
+    The trigger is long-running and emits one event for each non-empty shard 
response. Record data is
+    base64-encoded in the event payload and must be decoded by the consumer. 
Delivery is best-effort:
+    a triggerer failure can cause records to be repeated or missed around the 
failure window.
+
+    When Airflow provides an asset state store for a single watched asset, the 
trigger checkpoints the
+    last sequence number read from each shard. The same asset and stream 
identity share one logical cursor;
+    do not configure multiple watchers that require independent progress for 
the same stream on one asset.
+
+    :param stream_name: Name of the Kinesis Data Stream to watch.
+    :param aws_conn_id: AWS connection id.
+    :param shard_iterator_type: Position used when a shard has no checkpoint. 
``LATEST`` only sees records
+        that arrive after the watcher starts; ``TRIM_HORIZON`` starts from the 
oldest retained record.
+    :param batch_size: Maximum records per ``GetRecords`` call and trigger 
event. Must be between 1 and
+        10,000. Record data is base64-encoded before it is stored in the 
metadata database, so use a
+        conservative value for large records.
+    :param waiter_delay: Seconds between complete polling sweeps. Must be less 
than the five-minute
+        shard iterator lifetime. Kinesis permits at most five ``GetRecords`` 
calls per second per shard.
+        When reading a backlog with ``TRIM_HORIZON``, draining ``N`` records 
from one shard takes roughly
+        ``ceil(N / batch_size) * waiter_delay`` seconds when calls return full 
batches; with the defaults,
+        10,000 records take about 1,000 seconds.
+    :param region_name: AWS region for the Kinesis client.
+    :param verify: Whether to verify SSL certificates, or the path to a CA 
bundle.
+    :param botocore_config: Botocore configuration passed to the Kinesis 
client.
+    """
+
+    def __init__(
+        self,
+        stream_name: str,
+        aws_conn_id: str | None = "aws_default",
+        shard_iterator_type: str = "LATEST",
+        batch_size: int = 100,
+        waiter_delay: int = 10,
+        region_name: str | None = None,
+        verify: bool | str | None = None,
+        botocore_config: dict | None = None,
+    ) -> None:
+        super().__init__()
+        if shard_iterator_type not in _ITERATOR_TYPES_WITHOUT_EXTRA_ARGS:
+            raise ValueError(
+                "shard_iterator_type must be one of "
+                f"{sorted(_ITERATOR_TYPES_WITHOUT_EXTRA_ARGS)}; got 
{shard_iterator_type!r}"
+            )
+        if not 1 <= batch_size <= 10_000:
+            raise ValueError("batch_size must be between 1 and 10000")
+        if not 0 < waiter_delay < 300:
+            raise ValueError("waiter_delay must be between 1 and 299 seconds")
+
+        self.stream_name = stream_name
+        self.aws_conn_id = aws_conn_id
+        self.shard_iterator_type = shard_iterator_type
+        self.batch_size = batch_size
+        self.waiter_delay = waiter_delay
+        self.region_name = region_name
+        self.verify = verify
+        self.botocore_config = botocore_config
+        self._checkpoint_warning_logged = False
+
+    def serialize(self) -> tuple[str, dict[str, Any]]:
+        return (
+            self.__class__.__module__ + "." + self.__class__.__qualname__,
+            {
+                "stream_name": self.stream_name,
+                "aws_conn_id": self.aws_conn_id,
+                "shard_iterator_type": self.shard_iterator_type,
+                "batch_size": self.batch_size,
+                "waiter_delay": self.waiter_delay,
+                "region_name": self.region_name,
+                "verify": self.verify,
+                "botocore_config": self.botocore_config,
+            },
+        )
+
+    @property
+    def hook(self) -> KinesisHook:
+        return KinesisHook(
+            aws_conn_id=self.aws_conn_id,
+            region_name=self.region_name,
+            verify=self.verify,
+            config=self.botocore_config,
+        )
+
+    def _build_checkpoint_key(self) -> str:
+        identity = json.dumps(
+            {
+                "stream_name": self.stream_name,
+                "aws_conn_id": self.aws_conn_id,
+                "region_name": self.region_name,
+            },
+            sort_keys=True,
+            separators=(",", ":"),
+        ).encode()
+        return 
f"{_CHECKPOINT_KEY_PREFIX}:{hashlib.sha256(identity).hexdigest()}"
+
+    def _log_checkpoint_warning_once(self, message: str) -> None:
+        if self._checkpoint_warning_logged:
+            return
+        self.log.warning(message)
+        self._checkpoint_warning_logged = True
+
+    def _load_checkpoint(self) -> dict[str, str]:
+        store = getattr(self, "asset_state_store", None)
+        if store is None:
+            self._log_checkpoint_warning_once(
+                "Kinesis checkpointing is unavailable; using an in-memory 
cursor"
+            )
+            return {}
+
+        try:
+            checkpoint = store.get(self._build_checkpoint_key(), default={}) 
or {}
+        except ValueError:
+            self._log_checkpoint_warning_once(
+                "Kinesis checkpointing requires a single watched asset; using 
an in-memory cursor"
+            )
+            return {}
+
+        if not isinstance(checkpoint, dict) or not all(
+            isinstance(shard_id, str) and isinstance(sequence_number, str)
+            for shard_id, sequence_number in checkpoint.items()
+        ):
+            self._log_checkpoint_warning_once(
+                "Kinesis checkpoint data is invalid; using the configured 
initial position"
+            )
+            return {}
+        return dict(checkpoint)
+
+    def _save_checkpoint(self, sequence_numbers: dict[str, str]) -> None:
+        store = getattr(self, "asset_state_store", None)
+        if store is None:
+            self._log_checkpoint_warning_once(
+                "Kinesis checkpointing is unavailable; using an in-memory 
cursor"
+            )
+            return
+
+        try:
+            store.set(self._build_checkpoint_key(), dict(sequence_numbers))
+        except ValueError:

Review Comment:
   It's a good pattern to use (but slight correction: its asset state store not 
task state store), he Triggerer itself injects it at: 
https://github.com/apache/airflow/blob/main/airflow-core/src/airflow/jobs/triggerer_job_runner.py#L1405-L1408,
 so a `BaseEventTrigger` attached to a watched asset is designed to read and 
write asset state. Watermarking a stream cursor is one of the use case.



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]

Reply via email to