This is an automated email from the ASF dual-hosted git repository.

o-nikolas pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git


The following commit(s) were added to refs/heads/main by this push:
     new ffe077e8cd3 Add `DmsReloadTablesOperator` (#70569)
ffe077e8cd3 is described below

commit ffe077e8cd32d120c859e9460170911ca9567722
Author: Alejandro Morgante <[email protected]>
AuthorDate: Wed Jul 29 22:00:54 2026 -0300

    Add `DmsReloadTablesOperator` (#70569)
    
    Add DmsReloadTablesOperator to reload selected target tables for an active 
AWS DMS replication task.
    
    The operator wraps the AWS DMS ReloadTables API and supports:
    
        data-reload and validate-only reload options.
        Synchronous and deferrable completion waits.
        Configurable waiter delay and maximum attempts.
        Templated replication task, table, and reload-option arguments.
---
 providers/amazon/docs/operators/dms.rst            |  32 +++
 .../airflow/providers/amazon/aws/operators/dms.py  | 130 ++++++++++++
 .../airflow/providers/amazon/aws/triggers/dms.py   | 129 +++++++++++-
 .../airflow/providers/amazon/aws/waiters/dms.json  | 122 ++++++++++++
 .../amazon/tests/system/amazon/aws/example_dms.py  | 193 ++++++++++++------
 .../amazon/tests/unit/amazon/aws/hooks/test_dms.py |   1 -
 .../tests/unit/amazon/aws/operators/test_dms.py    | 219 +++++++++++++++++++++
 .../tests/unit/amazon/aws/triggers/test_dms.py     | 153 ++++++++++++++
 .../tests/unit/amazon/aws/waiters/test_dms.py      | 101 ++++++++++
 9 files changed, 1022 insertions(+), 58 deletions(-)

diff --git a/providers/amazon/docs/operators/dms.rst 
b/providers/amazon/docs/operators/dms.rst
index 77a576ff5ed..a4422d5ceb9 100644
--- a/providers/amazon/docs/operators/dms.rst
+++ b/providers/amazon/docs/operators/dms.rst
@@ -89,6 +89,38 @@ To start a replication task you can use
     :start-after: [START howto_operator_dms_start_task]
     :end-before: [END howto_operator_dms_start_task]
 
+.. _howto/operator:DmsReloadTablesOperator:
+
+Reload tables for a replication task
+====================================
+
+To reload selected target tables from their source data, use
+:class:`~airflow.providers.amazon.aws.operators.dms.DmsReloadTablesOperator`.
+The required parameters are ``replication_task_arn`` and ``tables_to_reload``. 
Each item in
+``tables_to_reload`` must contain ``SchemaName`` and ``TableName``. The 
optional ``reload_option``
+defaults to ``data-reload``, which reloads the data and runs validation again 
when it is enabled.
+Use ``validate-only`` to revalidate without reloading data; this option only 
applies when validation
+is enabled for the task.
+The replication task must be in the ``RUNNING`` state. By default, the 
operator waits until every
+requested operation completes. With ``data-reload``, it waits for 
``TableState`` to reach
+``Table completed``. With ``validate-only``, it waits for ``ValidationState`` 
to reach ``Validated``
+and fails if DMS reports mismatched or suspended records, a table error, or 
another terminal
+validation failure. The first waiter poll runs immediately after the request 
returns. Set
+``deferrable=True`` to release the worker slot while waiting, or set 
``wait_for_completion=False``
+to return the replication task ARN immediately. Use ``waiter_delay`` and 
``waiter_max_attempts`` to
+control subsequent polls for each table.
+AWS DMS accepts up to 10 unique tables per request and supports tasks using 
the ``full-load`` or
+``full-load-and-cdc`` migration type. DMS applies the task's 
``TargetTablePrepMode`` setting before
+reloading each table; when it is ``DO_NOTHING``, truncate the target table 
manually before reloading.
+See `Reloading tables during a task
+<https://docs.aws.amazon.com/dms/latest/userguide/CHAP_Tasks.ReloadTables.html>`__
 for details.
+
+.. exampleinclude:: /../../amazon/tests/system/amazon/aws/example_dms.py
+    :language: python
+    :dedent: 4
+    :start-after: [START howto_operator_dms_reload_tables]
+    :end-before: [END howto_operator_dms_reload_tables]
+
 .. _howto/operator:DmsDescribeTasksOperator:
 
 Get details of replication tasks
diff --git a/providers/amazon/src/airflow/providers/amazon/aws/operators/dms.py 
b/providers/amazon/src/airflow/providers/amazon/aws/operators/dms.py
index a6667108491..1cdd9884467 100644
--- a/providers/amazon/src/airflow/providers/amazon/aws/operators/dms.py
+++ b/providers/amazon/src/airflow/providers/amazon/aws/operators/dms.py
@@ -29,6 +29,7 @@ from airflow.providers.amazon.aws.triggers.dms import (
     DmsReplicationDeprovisionedTrigger,
     DmsReplicationStoppedTrigger,
     DmsReplicationTerminalStatusTrigger,
+    DmsTableReloadCompleteTrigger,
     DmsTaskModifyCompleteTrigger,
 )
 from airflow.providers.amazon.aws.utils import validate_execute_complete_event
@@ -419,6 +420,135 @@ class DmsStartTaskOperator(AwsBaseOperator[DmsHook]):
         self.log.info("DMS replication task(%s) is starting.", 
self.replication_task_arn)
 
 
+class DmsReloadTablesOperator(AwsBaseOperator[DmsHook]):
+    """
+    Reload target tables for a running AWS DMS replication task.
+
+    AWS DMS supports up to 10 unique tables per request. The replication task 
must be running
+    and use either the ``full-load`` or ``full-load-and-cdc`` migration type.
+
+    .. seealso::
+        For more information on how to use this operator, take a look at the 
guide:
+        :ref:`howto/operator:DmsReloadTablesOperator`
+
+    :param replication_task_arn: Replication task ARN. (required)
+    :param tables_to_reload: Tables to reload. Each item must contain 
``SchemaName`` and ``TableName``.
+        (required)
+    :param reload_option: Use ``data-reload`` to reload data and revalidate it 
when validation is
+        enabled. Use ``validate-only`` to revalidate without reloading data; 
this option only applies
+        when validation is enabled. Defaults to ``data-reload``.
+    :param wait_for_completion: If True, wait until every data reload or 
validation completes.
+        Defaults to True.
+    :param deferrable: Run the operator in deferrable mode when waiting for 
completion.
+        Defaults to the ``operators.default_deferrable`` configuration.
+    :param waiter_delay: Seconds between table statistics polls (default: 30).
+    :param waiter_max_attempts: Maximum table statistics poll attempts per 
table (default: 60).
+    :param aws_conn_id: The Airflow connection used for AWS credentials.
+        If this is ``None`` or empty then the default boto3 behaviour is used. 
If
+        running Airflow in a distributed manner and aws_conn_id is None or
+        empty, then default boto3 configuration would be used (and must be
+        maintained on each worker node).
+    :param region_name: AWS region_name. If not specified then the default 
boto3 behaviour is used.
+    :param verify: Whether or not to verify SSL certificates. See:
+        
https://boto3.amazonaws.com/v1/documentation/api/latest/reference/core/session.html
+    :param botocore_config: Configuration dictionary (key-values) for botocore 
client. See:
+        
https://botocore.amazonaws.com/v1/documentation/api/latest/reference/config.html
+    """
+
+    aws_hook_class = DmsHook
+    template_fields: Sequence[str] = aws_template_fields(
+        "replication_task_arn",
+        "tables_to_reload",
+        "reload_option",
+    )
+    template_fields_renderers: ClassVar[dict[str, str]] = {"tables_to_reload": 
"json"}
+
+    def __init__(
+        self,
+        *,
+        replication_task_arn: str,
+        tables_to_reload: list[dict[str, str]],
+        reload_option: str = "data-reload",
+        wait_for_completion: bool = True,
+        deferrable: bool = conf.getboolean("operators", "default_deferrable", 
fallback=False),
+        waiter_delay: int = 30,
+        waiter_max_attempts: int = 60,
+        aws_conn_id: str | None = "aws_default",
+        **kwargs,
+    ):
+        super().__init__(aws_conn_id=aws_conn_id, **kwargs)
+        self.replication_task_arn = replication_task_arn
+        self.tables_to_reload = tables_to_reload
+        self.reload_option = reload_option
+        self.wait_for_completion = wait_for_completion
+        self.deferrable = deferrable
+        self.waiter_delay = waiter_delay
+        self.waiter_max_attempts = waiter_max_attempts
+
+    def _wait_for_reload_completion(self) -> None:
+        waiter_name = (
+            "table_validation_complete" if self.reload_option == 
"validate-only" else "table_reload_complete"
+        )
+        for table in self.tables_to_reload:
+            self.hook.get_waiter(waiter_name).wait(
+                ReplicationTaskArn=self.replication_task_arn,
+                Filters=[
+                    {"Name": "schema-name", "Values": [table["SchemaName"]]},
+                    {"Name": "table-name", "Values": [table["TableName"]]},
+                ],
+                WaiterConfig={
+                    "Delay": self.waiter_delay,
+                    "MaxAttempts": self.waiter_max_attempts,
+                },
+            )
+
+    def execute(self, context: Context) -> str:
+        """Start reloading target tables for an AWS DMS replication task."""
+        self.log.info(
+            "Reloading %s table(s) for DMS replication task(%s).",
+            len(self.tables_to_reload),
+            self.replication_task_arn,
+        )
+        response = self.hook.conn.reload_tables(
+            ReplicationTaskArn=self.replication_task_arn,
+            TablesToReload=self.tables_to_reload,
+            ReloadOption=self.reload_option,
+        )
+        replication_task_arn = response["ReplicationTaskArn"]
+        self.log.info("DMS table reload started for replication task(%s).", 
replication_task_arn)
+
+        if self.wait_for_completion:
+            if self.deferrable:
+                self.defer(
+                    trigger=DmsTableReloadCompleteTrigger(
+                        replication_task_arn=self.replication_task_arn,
+                        tables_to_reload=self.tables_to_reload,
+                        reload_option=self.reload_option,
+                        waiter_delay=self.waiter_delay,
+                        waiter_max_attempts=self.waiter_max_attempts,
+                        aws_conn_id=self.aws_conn_id,
+                        region_name=self.region_name,
+                        verify=self.verify,
+                        botocore_config=self.botocore_config,
+                    ),
+                    method_name="execute_complete",
+                )
+            else:
+                self._wait_for_reload_completion()
+                self.log.info("DMS table reloads completed for replication 
task(%s).", replication_task_arn)
+
+        return replication_task_arn
+
+    def execute_complete(self, context: Context, event: dict | None = None) -> 
str:
+        """Resume after the table reload trigger completes."""
+        validated_event = validate_execute_complete_event(event)
+        if validated_event["status"] != "success":
+            raise RuntimeError(f"Error waiting for DMS table reloads to 
complete: {validated_event}")
+        replication_task_arn = validated_event["replication_task_arn"]
+        self.log.info("DMS table reloads completed for replication task(%s).", 
replication_task_arn)
+        return replication_task_arn
+
+
 class DmsStopTaskOperator(AwsBaseOperator[DmsHook]):
     """
     Stops AWS DMS replication task.
diff --git a/providers/amazon/src/airflow/providers/amazon/aws/triggers/dms.py 
b/providers/amazon/src/airflow/providers/amazon/aws/triggers/dms.py
index 8de08abcfeb..9f784de3c40 100644
--- a/providers/amazon/src/airflow/providers/amazon/aws/triggers/dms.py
+++ b/providers/amazon/src/airflow/providers/amazon/aws/triggers/dms.py
@@ -16,11 +16,15 @@
 # under the License.
 from __future__ import annotations
 
-from typing import TYPE_CHECKING
+from collections.abc import AsyncIterator
+from typing import TYPE_CHECKING, Any
 
+from airflow.exceptions import AirflowException
 from airflow.providers.amazon.aws.hooks.base_aws import AwsGenericHook
 from airflow.providers.amazon.aws.hooks.dms import DmsHook
 from airflow.providers.amazon.aws.triggers.base import AwsBaseWaiterTrigger
+from airflow.providers.amazon.aws.utils.waiter_with_logging import async_wait
+from airflow.triggers.base import BaseTrigger, TriggerEvent
 
 if TYPE_CHECKING:
     from airflow.providers.amazon.aws.hooks.base_aws import AwsGenericHook
@@ -267,3 +271,126 @@ class DmsTaskModifyCompleteTrigger(AwsBaseWaiterTrigger):
             verify=self.verify,
             config=self.botocore_config,
         )
+
+
+class DmsTableReloadCompleteTrigger(BaseTrigger):
+    """
+    Trigger when AWS DMS finishes reloading or validating a set of tables.
+
+    :param replication_task_arn: The ARN of the replication task.
+    :param tables_to_reload: Tables being reloaded, including schema and table 
names.
+    :param reload_option: The reload operation whose completion state should 
be monitored.
+    :param waiter_delay: The amount of time in seconds to wait between 
attempts.
+    :param waiter_max_attempts: The maximum number of attempts to be made.
+    :param aws_conn_id: The Airflow connection used for AWS credentials.
+    :param region_name: AWS region name.
+    :param verify: Whether or not to verify SSL certificates.
+    :param botocore_config: Configuration dictionary (key-values) for botocore 
client.
+    """
+
+    def __init__(
+        self,
+        *,
+        replication_task_arn: str,
+        tables_to_reload: list[dict[str, str]],
+        reload_option: str = "data-reload",
+        waiter_delay: int = 30,
+        waiter_max_attempts: int = 60,
+        aws_conn_id: str | None = "aws_default",
+        region_name: str | None = None,
+        verify: bool | str | None = None,
+        botocore_config: dict | None = None,
+    ) -> None:
+        super().__init__()
+        self.replication_task_arn = replication_task_arn
+        self.tables_to_reload = tables_to_reload
+        self.reload_option = reload_option
+        self.waiter_delay = waiter_delay
+        self.waiter_max_attempts = waiter_max_attempts
+        self.aws_conn_id = aws_conn_id
+        self.region_name = region_name
+        self.verify = verify
+        self.botocore_config = botocore_config
+
+    def serialize(self) -> tuple[str, dict[str, Any]]:
+        """Serialize trigger arguments and classpath."""
+        return (
+            f"{self.__class__.__module__}.{self.__class__.__name__}",
+            {
+                "replication_task_arn": self.replication_task_arn,
+                "tables_to_reload": self.tables_to_reload,
+                "reload_option": self.reload_option,
+                "waiter_delay": self.waiter_delay,
+                "waiter_max_attempts": self.waiter_max_attempts,
+                "aws_conn_id": self.aws_conn_id,
+                "region_name": self.region_name,
+                "verify": self.verify,
+                "botocore_config": self.botocore_config,
+            },
+        )
+
+    def _build_waiter_args(self, table: dict[str, str]) -> dict[str, Any]:
+        return {
+            "ReplicationTaskArn": self.replication_task_arn,
+            "Filters": [
+                {"Name": "schema-name", "Values": [table["SchemaName"]]},
+                {"Name": "table-name", "Values": [table["TableName"]]},
+            ],
+        }
+
+    def _get_waiter_config(self) -> tuple[str, str, str]:
+        if self.reload_option == "validate-only":
+            return (
+                "table_validation_complete",
+                "validation",
+                "TableStatistics[0].ValidationState",
+            )
+        return (
+            "table_reload_complete",
+            "reload",
+            "TableStatistics[0].TableState",
+        )
+
+    async def run(self) -> AsyncIterator[TriggerEvent]:
+        """Poll table statistics until all requested operations finish."""
+        hook = DmsHook(
+            aws_conn_id=self.aws_conn_id,
+            region_name=self.region_name,
+            verify=self.verify,
+            config=self.botocore_config,
+        )
+
+        try:
+            async with await hook.get_async_conn() as client:
+                waiter_name, operation_name, status_query = 
self._get_waiter_config()
+                for table in self.tables_to_reload:
+                    waiter = hook.get_waiter(
+                        waiter_name,
+                        deferrable=True,
+                        client=client,
+                    )
+                    table_name = f"{table['SchemaName']}.{table['TableName']}"
+                    await async_wait(
+                        waiter,
+                        self.waiter_delay,
+                        self.waiter_max_attempts,
+                        self._build_waiter_args(table),
+                        f"DMS table {operation_name} failed for {table_name}.",
+                        f"Status of DMS table {operation_name} {table_name} 
is",
+                        [status_query],
+                    )
+        except AirflowException as error:
+            yield TriggerEvent(
+                {
+                    "status": "error",
+                    "message": str(error),
+                    "replication_task_arn": self.replication_task_arn,
+                }
+            )
+        else:
+            yield TriggerEvent(
+                {
+                    "status": "success",
+                    "replication_task_arn": self.replication_task_arn,
+                }
+            )
diff --git a/providers/amazon/src/airflow/providers/amazon/aws/waiters/dms.json 
b/providers/amazon/src/airflow/providers/amazon/aws/waiters/dms.json
index d0123e4e8d9..72b342b90dd 100644
--- a/providers/amazon/src/airflow/providers/amazon/aws/waiters/dms.json
+++ b/providers/amazon/src/airflow/providers/amazon/aws/waiters/dms.json
@@ -174,6 +174,128 @@
                     "state": "failure"
                 }
             ]
+        },
+        "table_reload_complete": {
+            "operation": "DescribeTableStatistics",
+            "delay": 30,
+            "maxAttempts": 60,
+            "acceptors": [
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].TableState",
+                    "expected": "Table completed",
+                    "state": "success"
+                },
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].TableState",
+                    "expected": "Table is being reloaded",
+                    "state": "retry"
+                },
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].TableState",
+                    "expected": "Before load",
+                    "state": "retry"
+                },
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].TableState",
+                    "expected": "Full load",
+                    "state": "retry"
+                },
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].TableState",
+                    "expected": "Table cancelled",
+                    "state": "failure"
+                },
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].TableState",
+                    "expected": "Table does not exist",
+                    "state": "failure"
+                },
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].TableState",
+                    "expected": "Table error",
+                    "state": "failure"
+                }
+            ]
+        },
+        "table_validation_complete": {
+            "operation": "DescribeTableStatistics",
+            "delay": 30,
+            "maxAttempts": 60,
+            "acceptors": [
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].ValidationState",
+                    "expected": "Validated",
+                    "state": "success"
+                },
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].ValidationState",
+                    "expected": "Pending records",
+                    "state": "retry"
+                },
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].ValidationState",
+                    "expected": "Pending validation",
+                    "state": "retry"
+                },
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].ValidationState",
+                    "expected": "Preparing table",
+                    "state": "retry"
+                },
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].ValidationState",
+                    "expected": "Pending revalidation",
+                    "state": "retry"
+                },
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].ValidationState",
+                    "expected": "Mismatched records",
+                    "state": "failure"
+                },
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].ValidationState",
+                    "expected": "Suspended records",
+                    "state": "failure"
+                },
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].ValidationState",
+                    "expected": "No primary key",
+                    "state": "failure"
+                },
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].ValidationState",
+                    "expected": "Table error",
+                    "state": "failure"
+                },
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].ValidationState",
+                    "expected": "Error",
+                    "state": "failure"
+                },
+                {
+                    "matcher": "path",
+                    "argument": "TableStatistics[0].ValidationState",
+                    "expected": "Not enabled",
+                    "state": "failure"
+                }
+            ]
         }
     }
 }
diff --git a/providers/amazon/tests/system/amazon/aws/example_dms.py 
b/providers/amazon/tests/system/amazon/aws/example_dms.py
index edf8d9bef7c..8ecfc07b735 100644
--- a/providers/amazon/tests/system/amazon/aws/example_dms.py
+++ b/providers/amazon/tests/system/amazon/aws/example_dms.py
@@ -27,13 +27,16 @@ from datetime import datetime
 from typing import cast
 
 import boto3
+import pendulum
 from sqlalchemy import Column, MetaData, String, Table, create_engine
 
+from airflow.providers.amazon.aws.hooks.dms import DmsHook
 from airflow.providers.amazon.aws.operators.dms import (
     DmsCreateTaskOperator,
     DmsDeleteTaskOperator,
     DmsDescribeTasksOperator,
     DmsModifyTaskOperator,
+    DmsReloadTablesOperator,
     DmsStartTaskOperator,
     DmsStopTaskOperator,
 )
@@ -41,8 +44,8 @@ from airflow.providers.amazon.aws.operators.rds import (
     RdsCreateDbInstanceOperator,
     RdsDeleteDbInstanceOperator,
 )
-from airflow.providers.amazon.aws.operators.s3 import S3CreateBucketOperator, 
S3DeleteBucketOperator
 from airflow.providers.amazon.aws.sensors.dms import DmsTaskBaseSensor, 
DmsTaskCompletedSensor
+from airflow.providers.standard.sensors.date_time import DateTimeSensorAsync
 
 from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS
 
@@ -64,11 +67,10 @@ from system.amazon.aws.utils import ENV_ID_KEY, 
SystemTestContextBuilder
 from system.amazon.aws.utils.ec2 import get_default_vpc_id
 
 DAG_ID = "example_dms"
-ROLE_ARN_KEY = "ROLE_ARN"
 
-sys_test_context_task = 
SystemTestContextBuilder().add_variable(ROLE_ARN_KEY).build()
+sys_test_context_task = SystemTestContextBuilder().build()
 
-# Config values for setting up the "Source" database.
+# Config values for setting up the RDS databases.
 RDS_ENGINE = "postgres"
 RDS_PROTOCOL = "postgresql"
 RDS_USERNAME = "username"
@@ -117,6 +119,35 @@ def create_security_group(security_group_name: str, 
vpc_id: str):
     return security_group["GroupId"]
 
 
+@task(multiple_outputs=True)
+def create_db_parameter_group(parameter_group_name: str):
+    rds_client = boto3.client("rds")
+    engine = rds_client.describe_db_engine_versions(Engine=RDS_ENGINE, 
DefaultOnly=True)["DBEngineVersions"][
+        0
+    ]
+
+    rds_client.create_db_parameter_group(
+        DBParameterGroupName=parameter_group_name,
+        DBParameterGroupFamily=engine["DBParameterGroupFamily"],
+        Description="Created for DMS system test logical replication",
+    )
+    rds_client.modify_db_parameter_group(
+        DBParameterGroupName=parameter_group_name,
+        Parameters=[
+            {
+                "ParameterName": "rds.logical_replication",
+                "ParameterValue": "1",
+                "ApplyMethod": "pending-reboot",
+            }
+        ],
+    )
+    return {
+        "name": parameter_group_name,
+        "engine_version": engine["EngineVersion"],
+        "available_at": pendulum.now("UTC").add(minutes=5).isoformat(),
+    }
+
+
 @task
 def create_sample_table(instance_name: str, db_name: str, table_name: str):
     print("Creating sample table.")
@@ -134,7 +165,7 @@ def create_sample_table(instance_name: str, db_name: str, 
table_name: str):
         Column(TABLE_HEADERS[1], String),
     )
 
-    with engine.connect() as connection:
+    with engine.begin() as connection:
         # Create the Table.
         table.create(bind=connection)
         load_data = table.insert().values(SAMPLE_DATA)
@@ -144,16 +175,41 @@ def create_sample_table(instance_name: str, db_name: str, 
table_name: str):
         connection.execute(table.select())
 
 
+@task
+def create_target_database(instance_name: str, source_db_name: str, 
target_db_name: str):
+    print("Creating target database.")
+
+    rds_endpoint = _get_rds_instance_endpoint(instance_name)
+    hostname = rds_endpoint["Address"]
+    port = rds_endpoint["Port"]
+    rds_url = 
f"{RDS_PROTOCOL}://{RDS_USERNAME}:{RDS_PASSWORD}@{hostname}:{port}/{source_db_name}"
+    engine = create_engine(rds_url)
+    quoted_target_db_name = 
engine.dialect.identifier_preparer.quote(target_db_name)
+
+    with engine.connect().execution_options(isolation_level="AUTOCOMMIT") as 
connection:
+        connection.exec_driver_sql(f"CREATE DATABASE {quoted_target_db_name}")
+
+
+@task
+def await_table_load(replication_task_arn: str, schema_name: str, table_name: 
str):
+    DmsHook().get_waiter("table_reload_complete").wait(
+        ReplicationTaskArn=replication_task_arn,
+        Filters=[
+            {"Name": "schema-name", "Values": [schema_name]},
+            {"Name": "table-name", "Values": [table_name]},
+        ],
+        WaiterConfig={"Delay": 10, "MaxAttempts": 60},
+    )
+
+
 @task(multiple_outputs=True)
 def create_dms_assets(
-    db_name: str,
+    source_db_name: str,
+    target_db_name: str,
     instance_name: str,
     replication_instance_name: str,
-    bucket_name: str,
-    role_arn,
     source_endpoint_identifier: str,
     target_endpoint_identifier: str,
-    table_definition: dict,
 ):
     print("Creating DMS assets.")
     dms_client = boto3.client("dms")
@@ -162,7 +218,7 @@ def create_dms_assets(
     print("Creating replication instance.")
     instance_arn = dms_client.create_replication_instance(
         ReplicationInstanceIdentifier=replication_instance_name,
-        ReplicationInstanceClass="dms.t3.micro",
+        ReplicationInstanceClass="dms.t3.small",
     )["ReplicationInstance"]["ReplicationInstanceArn"]
 
     print("Creating DMS source endpoint.")
@@ -174,20 +230,21 @@ def create_dms_assets(
         Password=RDS_PASSWORD,
         ServerName=rds_instance_endpoint["Address"],
         Port=rds_instance_endpoint["Port"],
-        DatabaseName=db_name,
+        DatabaseName=source_db_name,
+        SslMode="require",
     )["Endpoint"]["EndpointArn"]
 
     print("Creating DMS target endpoint.")
     target_endpoint_arn = dms_client.create_endpoint(
         EndpointIdentifier=target_endpoint_identifier,
         EndpointType="target",
-        EngineName="s3",
-        S3Settings={
-            "BucketName": bucket_name,
-            "BucketFolder": "folder",
-            "ServiceAccessRoleArn": role_arn,
-            "ExternalTableDefinition": json.dumps(table_definition),
-        },
+        EngineName=RDS_ENGINE,
+        Username=RDS_USERNAME,
+        Password=RDS_PASSWORD,
+        ServerName=rds_instance_endpoint["Address"],
+        Port=rds_instance_endpoint["Port"],
+        DatabaseName=target_db_name,
+        SslMode="require",
     )["Endpoint"]["EndpointArn"]
 
     print("Awaiting replication instance provisioning.")
@@ -237,6 +294,15 @@ def delete_security_group(security_group_id: str, 
security_group_name: str):
     boto3.client("ec2").delete_security_group(GroupId=security_group_id, 
GroupName=security_group_name)
 
 
+@task(trigger_rule=TriggerRule.ALL_DONE)
+def delete_db_parameter_group(parameter_group_name: str):
+    rds_client = boto3.client("rds")
+    try:
+        
rds_client.delete_db_parameter_group(DBParameterGroupName=parameter_group_name)
+    except rds_client.exceptions.DBParameterGroupNotFoundFault:
+        print(f"DB parameter group {parameter_group_name} is already deleted.")
+
+
 with DAG(
     DAG_ID,
     schedule="@once",
@@ -245,37 +311,23 @@ with DAG(
 ) as dag:
     test_context = sys_test_context_task()
     env_id = test_context[ENV_ID_KEY]
-    role_arn = test_context[ROLE_ARN_KEY]
 
-    bucket_name = f"{env_id}-dms-bucket"
     rds_instance_name = f"{env_id}-instance"
-    rds_db_name = f"{env_id}_source_database"  # dashes are not allowed in db 
name
+    rds_source_db_name = f"{env_id}_source_database"  # dashes are not allowed 
in db name
+    rds_target_db_name = f"{env_id}_target_database"
     rds_table_name = f"{env_id}-table"
     dms_replication_instance_name = f"{env_id}-replication-instance"
     dms_replication_task_id = f"{env_id}-replication-task"
     source_endpoint_identifier = f"{env_id}-source-endpoint"
     target_endpoint_identifier = f"{env_id}-target-endpoint"
     security_group_name = f"{env_id}-dms-security-group"
+    db_parameter_group_name = f"{env_id}-dms-parameter-group"
+    db_parameter_group = create_db_parameter_group(db_parameter_group_name)
+    await_db_parameter_group = DateTimeSensorAsync(
+        task_id="await_db_parameter_group",
+        target_time=db_parameter_group["available_at"],
+    )
 
-    # Sample data.
-    table_definition = {
-        "TableCount": "1",
-        "Tables": [
-            {
-                "TableName": rds_table_name,
-                "TableColumns": [
-                    {
-                        "ColumnName": TABLE_HEADERS[0],
-                        "ColumnType": "STRING",
-                        "ColumnNullable": "false",
-                        "ColumnIsPk": "true",
-                    },
-                    {"ColumnName": TABLE_HEADERS[1], "ColumnType": "STRING", 
"ColumnLength": "4"},
-                ],
-                "TableColumnsTotal": "2",
-            }
-        ],
-    }
     table_mappings = {
         "rules": [
             {
@@ -291,8 +343,6 @@ with DAG(
         ]
     }
 
-    create_s3_bucket = S3CreateBucketOperator(task_id="create_s3_bucket", 
bucket_name=bucket_name)
-
     get_vpc_id = get_default_vpc_id()
 
     create_sg = create_security_group(security_group_name, get_vpc_id)
@@ -303,26 +353,32 @@ with DAG(
         db_instance_class="db.t3.micro",
         engine=RDS_ENGINE,
         rds_kwargs={
-            "DBName": rds_db_name,
+            "DBName": rds_source_db_name,
             "AllocatedStorage": 20,
             "MasterUsername": RDS_USERNAME,
             "MasterUserPassword": RDS_PASSWORD,
             "PubliclyAccessible": True,
+            "EngineVersion": db_parameter_group["engine_version"],
+            "DBParameterGroupName": db_parameter_group["name"],
             "VpcSecurityGroupIds": [
                 create_sg,
             ],
         },
     )
 
+    create_target_db = create_target_database(
+        instance_name=rds_instance_name,
+        source_db_name=rds_source_db_name,
+        target_db_name=rds_target_db_name,
+    )
+
     create_assets = create_dms_assets(
-        db_name=rds_db_name,
+        source_db_name=rds_source_db_name,
+        target_db_name=rds_target_db_name,
         instance_name=rds_instance_name,
         replication_instance_name=dms_replication_instance_name,
-        bucket_name=bucket_name,
-        role_arn=role_arn,
         source_endpoint_identifier=source_endpoint_identifier,
         target_endpoint_identifier=target_endpoint_identifier,
-        table_definition=table_definition,
     )
 
     # [START howto_operator_dms_create_task]
@@ -333,6 +389,10 @@ with DAG(
         target_endpoint_arn=create_assets["target_endpoint_arn"],
         replication_instance_arn=create_assets["replication_instance_arn"],
         table_mappings=table_mappings,
+        migration_type="full-load-and-cdc",
+        create_task_kwargs={
+            "ReplicationTaskSettings": json.dumps({"ValidationSettings": 
{"EnableValidation": True}})
+        },
     )
     # [END howto_operator_dms_create_task]
 
@@ -368,6 +428,27 @@ with DAG(
         poke_interval=10,
     )
 
+    await_initial_table_load = await_table_load(task_arn, "public", 
rds_table_name)
+
+    # [START howto_operator_dms_reload_tables]
+    reload_tables = DmsReloadTablesOperator(
+        task_id="reload_tables",
+        replication_task_arn=task_arn,
+        tables_to_reload=[{"SchemaName": "public", "TableName": 
rds_table_name}],
+        reload_option="data-reload",
+        wait_for_completion=True,
+        deferrable=True,
+    )
+    revalidate_tables = DmsReloadTablesOperator(
+        task_id="revalidate_tables",
+        replication_task_arn=task_arn,
+        tables_to_reload=[{"SchemaName": "public", "TableName": 
rds_table_name}],
+        reload_option="validate-only",
+        wait_for_completion=True,
+        deferrable=True,
+    )
+    # [END howto_operator_dms_reload_tables]
+
     # [START howto_operator_dms_stop_task]
     stop_task = DmsStopTaskOperator(
         task_id="stop_task",
@@ -428,27 +509,27 @@ with DAG(
         trigger_rule=TriggerRule.ALL_DONE,
     )
 
-    delete_s3_bucket = S3DeleteBucketOperator(
-        task_id="delete_s3_bucket",
-        bucket_name=bucket_name,
-        force_delete=True,
-        trigger_rule=TriggerRule.ALL_DONE,
-    )
+    delete_parameter_group = delete_db_parameter_group(db_parameter_group_name)
 
     chain(
         # TEST SETUP
         test_context,
-        create_s3_bucket,
         get_vpc_id,
         create_sg,
+        db_parameter_group,
+        await_db_parameter_group,
         create_db_instance,
-        create_sample_table(rds_instance_name, rds_db_name, rds_table_name),
+        create_target_db,
+        create_sample_table(rds_instance_name, rds_source_db_name, 
rds_table_name),
         create_assets,
         # TEST BODY
         create_task,
         start_task,
         describe_tasks,
         await_task_start,
+        await_initial_table_load,
+        reload_tables,
+        revalidate_tables,
         stop_task,
         await_task_stop,
         modify_task,
@@ -456,8 +537,8 @@ with DAG(
         delete_task,
         delete_assets,
         delete_db_instance,
+        delete_parameter_group,
         delete_security_group(create_sg, security_group_name),
-        delete_s3_bucket,
     )
 
     from tests_common.test_utils.watcher import watcher
diff --git a/providers/amazon/tests/unit/amazon/aws/hooks/test_dms.py 
b/providers/amazon/tests/unit/amazon/aws/hooks/test_dms.py
index baea76e5871..c1dc96bd9f4 100644
--- a/providers/amazon/tests/unit/amazon/aws/hooks/test_dms.py
+++ b/providers/amazon/tests/unit/amazon/aws/hooks/test_dms.py
@@ -67,7 +67,6 @@ MOCK_CREATE_RESPONSE: dict[str, Any] = {"ReplicationTask": 
MOCK_TASK_RESPONSE_DA
 MOCK_START_RESPONSE: dict[str, Any] = {"ReplicationTask": 
{**MOCK_TASK_RESPONSE_DATA, "Status": "starting"}}
 MOCK_STOP_RESPONSE: dict[str, Any] = {"ReplicationTask": 
{**MOCK_TASK_RESPONSE_DATA, "Status": "stopping"}}
 MOCK_DELETE_RESPONSE: dict[str, Any] = {"ReplicationTask": 
{**MOCK_TASK_RESPONSE_DATA, "Status": "deleting"}}
-
 MOCK_CONFIG_RESPONSE: dict[str, Any] = {
     "Marker": "xxxxx",
     "ReplicationConfigs": [
diff --git a/providers/amazon/tests/unit/amazon/aws/operators/test_dms.py 
b/providers/amazon/tests/unit/amazon/aws/operators/test_dms.py
index e70149a8c18..2c2a299353d 100644
--- a/providers/amazon/tests/unit/amazon/aws/operators/test_dms.py
+++ b/providers/amazon/tests/unit/amazon/aws/operators/test_dms.py
@@ -37,6 +37,7 @@ from airflow.providers.amazon.aws.operators.dms import (
     DmsDescribeReplicationsOperator,
     DmsDescribeTasksOperator,
     DmsModifyTaskOperator,
+    DmsReloadTablesOperator,
     DmsStartReplicationOperator,
     DmsStartTaskOperator,
     DmsStopReplicationOperator,
@@ -45,6 +46,7 @@ from airflow.providers.amazon.aws.operators.dms import (
 from airflow.providers.amazon.aws.triggers.dms import (
     DmsReplicationDeprovisionedTrigger,
     DmsReplicationTerminalStatusTrigger,
+    DmsTableReloadCompleteTrigger,
     DmsTaskModifyCompleteTrigger,
 )
 from airflow.providers.common.compat.sdk import AirflowException, TaskDeferred
@@ -720,6 +722,223 @@ class TestDmsStartTaskOperator:
         validate_template_fields(op)
 
 
+class TestDmsReloadTablesOperator:
+    TABLES_TO_RELOAD = [{"SchemaName": "public", "TableName": "test_table"}]
+
+    def test_init(self):
+        op = DmsReloadTablesOperator(
+            task_id="reload_tables",
+            replication_task_arn=TASK_ARN,
+            tables_to_reload=self.TABLES_TO_RELOAD,
+            aws_conn_id="fake-conn-id",
+            region_name="us-west-1",
+            verify=False,
+            botocore_config={"read_timeout": 42},
+        )
+
+        assert op.replication_task_arn == TASK_ARN
+        assert op.tables_to_reload == self.TABLES_TO_RELOAD
+        assert op.reload_option == "data-reload"
+        assert op.wait_for_completion is True
+        assert op.deferrable is False
+        assert op.waiter_delay == 30
+        assert op.waiter_max_attempts == 60
+        assert op.hook.client_type == "dms"
+        assert op.hook.resource_type is None
+        assert op.hook.aws_conn_id == "fake-conn-id"
+        assert op.hook._region_name == "us-west-1"
+        assert op.hook._verify is False
+        assert op.hook._config is not None
+        assert op.hook._config.read_timeout == 42
+
+    @pytest.mark.parametrize("reload_option", ["data-reload", "validate-only"])
+    @mock.patch.object(DmsHook, "conn", new_callable=mock.PropertyMock)
+    def test_execute(self, mock_conn, reload_option):
+        mock_client = mock.MagicMock(spec=["reload_tables"])
+        mock_client.reload_tables.return_value = {"ReplicationTaskArn": 
TASK_ARN}
+        mock_conn.return_value = mock_client
+        op = DmsReloadTablesOperator(
+            task_id="reload_tables",
+            replication_task_arn=TASK_ARN,
+            tables_to_reload=self.TABLES_TO_RELOAD,
+            reload_option=reload_option,
+            wait_for_completion=False,
+        )
+
+        result = op.execute(None)
+
+        mock_client.reload_tables.assert_called_once_with(
+            ReplicationTaskArn=TASK_ARN,
+            TablesToReload=self.TABLES_TO_RELOAD,
+            ReloadOption=reload_option,
+        )
+        assert result == TASK_ARN
+
+    @pytest.mark.parametrize("reload_option", ["data-reload", "validate-only"])
+    @mock.patch.object(DmsHook, "conn", new_callable=mock.PropertyMock)
+    def test_execute_defers_for_completion(self, mock_conn, reload_option):
+        mock_client = mock.MagicMock(spec=["reload_tables"])
+        mock_client.reload_tables.return_value = {"ReplicationTaskArn": 
TASK_ARN}
+        mock_conn.return_value = mock_client
+        op = DmsReloadTablesOperator(
+            task_id="reload_tables",
+            replication_task_arn=TASK_ARN,
+            tables_to_reload=self.TABLES_TO_RELOAD,
+            reload_option=reload_option,
+            wait_for_completion=True,
+            deferrable=True,
+            waiter_delay=5,
+            waiter_max_attempts=10,
+            aws_conn_id="test_conn",
+            region_name="us-east-2",
+            verify=False,
+            botocore_config={"read_timeout": 42},
+        )
+
+        with pytest.raises(TaskDeferred) as exc_info:
+            op.execute(None)
+
+        trigger = exc_info.value.trigger
+        assert isinstance(trigger, DmsTableReloadCompleteTrigger)
+        assert trigger.replication_task_arn == TASK_ARN
+        assert trigger.tables_to_reload == self.TABLES_TO_RELOAD
+        assert trigger.reload_option == reload_option
+        assert trigger.waiter_delay == 5
+        assert trigger.waiter_max_attempts == 10
+        assert trigger.aws_conn_id == "test_conn"
+        assert trigger.region_name == "us-east-2"
+        assert trigger.verify is False
+        assert trigger.botocore_config == {"read_timeout": 42}
+        assert exc_info.value.method_name == "execute_complete"
+        mock_client.reload_tables.assert_called_once()
+
+    @pytest.mark.parametrize(
+        ("reload_option", "waiter_name"),
+        [
+            pytest.param("data-reload", "table_reload_complete", 
id="data-reload"),
+            pytest.param("validate-only", "table_validation_complete", 
id="validate-only"),
+        ],
+    )
+    @mock.patch.object(DmsHook, "get_waiter", autospec=True)
+    @mock.patch.object(DmsHook, "conn", new_callable=mock.PropertyMock)
+    def test_execute_waits_for_completion(
+        self,
+        mock_conn,
+        mock_get_waiter,
+        reload_option,
+        waiter_name,
+    ):
+        mock_client = mock.MagicMock(spec=["reload_tables"])
+        mock_client.reload_tables.return_value = {"ReplicationTaskArn": 
TASK_ARN}
+        mock_conn.return_value = mock_client
+        tables_to_reload = [
+            {"SchemaName": "public", "TableName": "first_table"},
+            {"SchemaName": "archive", "TableName": "second_table"},
+        ]
+        mock_waiter = mock.MagicMock(spec=["wait"])
+        mock_get_waiter.return_value = mock_waiter
+        op = DmsReloadTablesOperator(
+            task_id="reload_tables",
+            replication_task_arn=TASK_ARN,
+            tables_to_reload=tables_to_reload,
+            reload_option=reload_option,
+            waiter_delay=5,
+            waiter_max_attempts=10,
+        )
+
+        result = op.execute(None)
+
+        assert result == TASK_ARN
+        mock_client.reload_tables.assert_called_once()
+        assert mock_get_waiter.call_args_list == [
+            mock.call(op.hook, waiter_name),
+            mock.call(op.hook, waiter_name),
+        ]
+        assert mock_waiter.wait.call_args_list == [
+            mock.call(
+                ReplicationTaskArn=TASK_ARN,
+                Filters=[
+                    {"Name": "schema-name", "Values": ["public"]},
+                    {"Name": "table-name", "Values": ["first_table"]},
+                ],
+                WaiterConfig={"Delay": 5, "MaxAttempts": 10},
+            ),
+            mock.call(
+                ReplicationTaskArn=TASK_ARN,
+                Filters=[
+                    {"Name": "schema-name", "Values": ["archive"]},
+                    {"Name": "table-name", "Values": ["second_table"]},
+                ],
+                WaiterConfig={"Delay": 5, "MaxAttempts": 10},
+            ),
+        ]
+
+    @mock.patch.object(DmsHook, "get_waiter", autospec=True)
+    @mock.patch.object(DmsHook, "conn", new_callable=mock.PropertyMock)
+    def test_execute_propagates_waiter_error(
+        self,
+        mock_conn,
+        mock_get_waiter,
+    ):
+        mock_client = mock.MagicMock(spec=["reload_tables"])
+        mock_client.reload_tables.return_value = {"ReplicationTaskArn": 
TASK_ARN}
+        mock_conn.return_value = mock_client
+        mock_get_waiter.return_value.wait.side_effect = WaiterError(
+            name="table_reload_complete",
+            reason="Max attempts exceeded",
+            last_response={},
+        )
+        op = DmsReloadTablesOperator(
+            task_id="reload_tables",
+            replication_task_arn=TASK_ARN,
+            tables_to_reload=self.TABLES_TO_RELOAD,
+            wait_for_completion=True,
+            waiter_max_attempts=1,
+        )
+
+        with pytest.raises(WaiterError, match="Max attempts exceeded"):
+            op.execute(None)
+
+        mock_client.reload_tables.assert_called_once()
+
+    def test_execute_complete(self):
+        op = DmsReloadTablesOperator(
+            task_id="reload_tables",
+            replication_task_arn=TASK_ARN,
+            tables_to_reload=self.TABLES_TO_RELOAD,
+        )
+
+        assert (
+            op.execute_complete(
+                None,
+                event={"status": "success", "replication_task_arn": TASK_ARN},
+            )
+            == TASK_ARN
+        )
+
+    def test_execute_complete_raises_for_error(self):
+        op = DmsReloadTablesOperator(
+            task_id="reload_tables",
+            replication_task_arn=TASK_ARN,
+            tables_to_reload=self.TABLES_TO_RELOAD,
+        )
+
+        with pytest.raises(RuntimeError, match="Error waiting for DMS table 
reloads"):
+            op.execute_complete(
+                None,
+                event={"status": "error", "message": "reload failed"},
+            )
+
+    def test_template_fields(self):
+        op = DmsReloadTablesOperator(
+            task_id="reload_tables",
+            replication_task_arn=TASK_ARN,
+            tables_to_reload=self.TABLES_TO_RELOAD,
+        )
+
+        validate_template_fields(op)
+
+
 class TestDmsStopTaskOperator:
     TASK_DATA = {
         "replication_task_id": "task_id",
diff --git a/providers/amazon/tests/unit/amazon/aws/triggers/test_dms.py 
b/providers/amazon/tests/unit/amazon/aws/triggers/test_dms.py
index 3805ab4f5b8..25a1ef4c812 100644
--- a/providers/amazon/tests/unit/amazon/aws/triggers/test_dms.py
+++ b/providers/amazon/tests/unit/amazon/aws/triggers/test_dms.py
@@ -29,6 +29,7 @@ from airflow.providers.amazon.aws.triggers.dms import (
     DmsReplicationDeprovisionedTrigger,
     DmsReplicationStoppedTrigger,
     DmsReplicationTerminalStatusTrigger,
+    DmsTableReloadCompleteTrigger,
     DmsTaskModifyCompleteTrigger,
 )
 from airflow.triggers.base import TriggerEvent
@@ -229,3 +230,155 @@ class TestDmsTaskModifyCompleteTrigger:
                 "replication_task_arn": self.TASK_ARN,
             }
         )
+
+
+class TestDmsTableReloadCompleteTrigger:
+    TASK_ARN = "arn:aws:dms:us-east-1:123456789012:task:EXAMPLE"
+    TABLES = [
+        {"SchemaName": "public", "TableName": "first_table"},
+        {"SchemaName": "archive", "TableName": "second_table"},
+    ]
+
+    def build_trigger(self, **overrides):
+        kwargs = {
+            "replication_task_arn": self.TASK_ARN,
+            "tables_to_reload": self.TABLES,
+            "reload_option": "data-reload",
+            "waiter_delay": 5,
+            "waiter_max_attempts": 10,
+            "aws_conn_id": "test_conn",
+            "region_name": "us-east-2",
+            "verify": False,
+            "botocore_config": {"read_timeout": 42},
+            **overrides,
+        }
+        return DmsTableReloadCompleteTrigger(**kwargs)
+
+    def test_serialization(self):
+        trigger = self.build_trigger()
+
+        classpath, kwargs = trigger.serialize()
+
+        assert classpath == BASE_TRIGGER_CLASSPATH + 
"DmsTableReloadCompleteTrigger"
+        assert kwargs == {
+            "replication_task_arn": self.TASK_ARN,
+            "tables_to_reload": self.TABLES,
+            "reload_option": "data-reload",
+            "waiter_delay": 5,
+            "waiter_max_attempts": 10,
+            "aws_conn_id": "test_conn",
+            "region_name": "us-east-2",
+            "verify": False,
+            "botocore_config": {"read_timeout": 42},
+        }
+
+    @pytest.mark.asyncio
+    @pytest.mark.parametrize(
+        ("reload_option", "waiter_name", "operation_name", "status_query"),
+        [
+            pytest.param(
+                "data-reload",
+                "table_reload_complete",
+                "reload",
+                "TableStatistics[0].TableState",
+                id="data-reload",
+            ),
+            pytest.param(
+                "validate-only",
+                "table_validation_complete",
+                "validation",
+                "TableStatistics[0].ValidationState",
+                id="validate-only",
+            ),
+        ],
+    )
+    @mock.patch(f"{DmsTableReloadCompleteTrigger.__module__}.async_wait", 
autospec=True)
+    @mock.patch.object(DmsHook, "get_waiter", autospec=True)
+    @mock.patch.object(DmsHook, "get_async_conn", autospec=True)
+    async def test_run_success(
+        self,
+        mock_get_async_conn,
+        mock_get_waiter,
+        mock_async_wait,
+        reload_option,
+        waiter_name,
+        operation_name,
+        status_query,
+    ):
+        mock_client = mock.MagicMock(spec=["describe_table_statistics"])
+        mock_waiter = mock.MagicMock(spec=["wait"])
+        mock_get_async_conn.return_value.__aenter__.return_value = mock_client
+        mock_get_waiter.return_value = mock_waiter
+
+        [response] = [event async for event in 
self.build_trigger(reload_option=reload_option).run()]
+
+        assert response == TriggerEvent({"status": "success", 
"replication_task_arn": self.TASK_ARN})
+        assert mock_get_waiter.call_args_list == [
+            mock.call(
+                mock.ANY,
+                waiter_name,
+                deferrable=True,
+                client=mock_client,
+            ),
+            mock.call(
+                mock.ANY,
+                waiter_name,
+                deferrable=True,
+                client=mock_client,
+            ),
+        ]
+        assert mock_async_wait.await_args_list == [
+            mock.call(
+                mock_waiter,
+                5,
+                10,
+                {
+                    "ReplicationTaskArn": self.TASK_ARN,
+                    "Filters": [
+                        {"Name": "schema-name", "Values": ["public"]},
+                        {"Name": "table-name", "Values": ["first_table"]},
+                    ],
+                },
+                f"DMS table {operation_name} failed for public.first_table.",
+                f"Status of DMS table {operation_name} public.first_table is",
+                [status_query],
+            ),
+            mock.call(
+                mock_waiter,
+                5,
+                10,
+                {
+                    "ReplicationTaskArn": self.TASK_ARN,
+                    "Filters": [
+                        {"Name": "schema-name", "Values": ["archive"]},
+                        {"Name": "table-name", "Values": ["second_table"]},
+                    ],
+                },
+                f"DMS table {operation_name} failed for archive.second_table.",
+                f"Status of DMS table {operation_name} archive.second_table 
is",
+                [status_query],
+            ),
+        ]
+
+    @pytest.mark.asyncio
+    @mock.patch(
+        f"{DmsTableReloadCompleteTrigger.__module__}.async_wait",
+        autospec=True,
+        side_effect=AirflowException("DMS table reload failed."),
+    )
+    @mock.patch.object(DmsHook, "get_waiter", autospec=True)
+    @mock.patch.object(DmsHook, "get_async_conn", autospec=True)
+    async def test_run_failure(self, mock_get_async_conn, mock_get_waiter, 
mock_async_wait):
+        mock_get_async_conn.return_value.__aenter__.return_value = 
mock.MagicMock(
+            spec=["describe_table_statistics"]
+        )
+        mock_get_waiter.return_value = mock.MagicMock(spec=["wait"])
+
+        [response] = [event async for event in self.build_trigger().run()]
+
+        assert response.payload == {
+            "status": "error",
+            "message": "DMS table reload failed.",
+            "replication_task_arn": self.TASK_ARN,
+        }
+        mock_async_wait.assert_awaited_once()
diff --git a/providers/amazon/tests/unit/amazon/aws/waiters/test_dms.py 
b/providers/amazon/tests/unit/amazon/aws/waiters/test_dms.py
index 378a38bc158..34b037f8f54 100644
--- a/providers/amazon/tests/unit/amazon/aws/waiters/test_dms.py
+++ b/providers/amazon/tests/unit/amazon/aws/waiters/test_dms.py
@@ -21,6 +21,7 @@ from unittest import mock
 
 import boto3
 import pytest
+from botocore.exceptions import WaiterError
 
 from airflow.providers.amazon.aws.hooks.dms import DmsHook
 
@@ -40,6 +41,8 @@ class TestCustomDmsWaiters:
         assert "replication_stopped" in hook_waiters
         assert "replication_complete" in hook_waiters
         assert "replication_task_modified" in hook_waiters
+        assert "table_reload_complete" in hook_waiters
+        assert "table_validation_complete" in hook_waiters
 
     @pytest.fixture
     def mock_describe_replication(self):
@@ -56,6 +59,11 @@ class TestCustomDmsWaiters:
         with mock.patch.object(self.client, "describe_replication_configs") as 
m:
             yield m
 
+    @pytest.fixture
+    def mock_describe_table_statistics(self):
+        with mock.patch.object(self.client, "describe_table_statistics") as m:
+            yield m
+
     def test_wait_for_replication_terminal_status(self, 
mock_describe_replication):
         mock_describe_replication.return_value = {
             "Replications": [
@@ -177,3 +185,96 @@ class TestCustomDmsWaiters:
             ],
             WithoutSettings=True,
         )
+
+    def test_wait_for_table_reload_complete(self, 
mock_describe_table_statistics):
+        mock_describe_table_statistics.side_effect = [
+            {"TableStatistics": [{"TableState": "Before load"}]},
+            {"TableStatistics": [{"TableState": "Table is being reloaded"}]},
+            {"TableStatistics": [{"TableState": "Full load"}]},
+            {"TableStatistics": [{"TableState": "Table completed"}]},
+        ]
+
+        hook = DmsHook(aws_conn_id=None)
+        waiter = hook.get_waiter("table_reload_complete")
+        waiter.wait(
+            ReplicationTaskArn="task-arn",
+            Filters=[
+                {"Name": "schema-name", "Values": ["dmsreload"]},
+                {"Name": "table-name", "Values": ["reload_test"]},
+            ],
+            WaiterConfig={"Delay": 0.01, "MaxAttempts": 4},
+        )
+
+        assert mock_describe_table_statistics.call_count == 4
+        mock_describe_table_statistics.assert_called_with(
+            ReplicationTaskArn="task-arn",
+            Filters=[
+                {"Name": "schema-name", "Values": ["dmsreload"]},
+                {"Name": "table-name", "Values": ["reload_test"]},
+            ],
+        )
+
+    @pytest.mark.parametrize("table_state", ["Table cancelled", "Table does 
not exist", "Table error"])
+    def test_wait_for_table_reload_failure(self, 
mock_describe_table_statistics, table_state):
+        mock_describe_table_statistics.return_value = {"TableStatistics": 
[{"TableState": table_state}]}
+
+        waiter = DmsHook(aws_conn_id=None).get_waiter("table_reload_complete")
+
+        with pytest.raises(WaiterError, match="terminal failure"):
+            waiter.wait(
+                ReplicationTaskArn="task-arn",
+                Filters=[
+                    {"Name": "schema-name", "Values": ["dmsreload"]},
+                    {"Name": "table-name", "Values": ["reload_test"]},
+                ],
+                WaiterConfig={"Delay": 0.01, "MaxAttempts": 1},
+            )
+
+    def test_wait_for_table_validation_complete(self, 
mock_describe_table_statistics):
+        mock_describe_table_statistics.side_effect = [
+            {"TableStatistics": [{"ValidationState": "Pending validation"}]},
+            {"TableStatistics": [{"ValidationState": "Preparing table"}]},
+            {"TableStatistics": [{"ValidationState": "Pending revalidation"}]},
+            {"TableStatistics": [{"ValidationState": "Pending records"}]},
+            {"TableStatistics": [{"ValidationState": "Validated"}]},
+        ]
+
+        waiter = 
DmsHook(aws_conn_id=None).get_waiter("table_validation_complete")
+        waiter.wait(
+            ReplicationTaskArn="task-arn",
+            Filters=[
+                {"Name": "schema-name", "Values": ["dmsreload"]},
+                {"Name": "table-name", "Values": ["reload_test"]},
+            ],
+            WaiterConfig={"Delay": 0.01, "MaxAttempts": 5},
+        )
+
+        assert mock_describe_table_statistics.call_count == 5
+
+    @pytest.mark.parametrize(
+        "validation_state",
+        [
+            "Mismatched records",
+            "Suspended records",
+            "No primary key",
+            "Table error",
+            "Error",
+            "Not enabled",
+        ],
+    )
+    def test_wait_for_table_validation_failure(self, 
mock_describe_table_statistics, validation_state):
+        mock_describe_table_statistics.return_value = {
+            "TableStatistics": [{"ValidationState": validation_state}]
+        }
+
+        waiter = 
DmsHook(aws_conn_id=None).get_waiter("table_validation_complete")
+
+        with pytest.raises(WaiterError, match="terminal failure"):
+            waiter.wait(
+                ReplicationTaskArn="task-arn",
+                Filters=[
+                    {"Name": "schema-name", "Values": ["dmsreload"]},
+                    {"Name": "table-name", "Values": ["reload_test"]},
+                ],
+                WaiterConfig={"Delay": 0.01, "MaxAttempts": 1},
+            )

Reply via email to