import logging from typing import Any, Dict from typing import Optional from airflow.providers.amazon.aws.hooks.dynamodb import DynamoDBHook from airflow.sensors.base import BaseSensorOperator logger = logging.getLogger(__name__) class MyDynamoDBHook(DynamoDBHook): """ Interact with AWS DynamoDB. Extends AwsDynamoDBHook with required functionality """ def get_item(self, key) -> Optional[Dict]: table = self.get_conn().Table(self.table_name) item_response = table.get_item(Key=key) logger.info(f'item_response {item_response}') return item_response.get('Item') class FeedIngestionStatusSensor(BaseSensorOperator): """ Check status of feed ingestion in DynamoDB. :param feed_name: feed_name to check the state of :type feed_name: str :param target_state: target state. Default to 'INGESTED' :type target_state: str :param max_retries: Number of times to poll for query state before returning the current state, defaults to None :type max_retries: int :param aws_conn_id: aws connection to use, defaults to 'aws_default' :type aws_conn_id: str :param sleep_time: Time in seconds to wait between two consecutive call to check query status on athena, defaults to 10 :type sleep_time: int """ TARGET_STATE_DEFAULT = 'INGESTED' template_fields = [] template_ext = () ui_color = '#66c3ff' def __init__( self, *args, feed_name: str, target_state: str = TARGET_STATE_DEFAULT, table_name: str, aws_conn_id: str = 'aws_default', **kwargs: Any, ) -> None: super().__init__(**kwargs) self.aws_conn_id = aws_conn_id self.feed_name = feed_name self.target_state = target_state self.table_name = table_name assert self.target_state def poke(self, context: dict) -> bool: logger.warning(context) key = { 'feed_name': self.feed_name, 'date': context['ds'], } logger.info(f'Lookup status table {self.table_name} for key: {key}') state_item = self.get_hook().get_item(key=key) logger.info(f'State item returned: {state_item}') if state_item and state_item.get('state') == self.target_state: return True return False def get_hook(self) -> MyDynamoDBHook: return MyDynamoDBHook( aws_conn_id=self.aws_conn_id, table_name=self.table_name, )