From ae855d6f3d2dee9ef81da8bff27137522eff7b58 Mon Sep 17 00:00:00 2001 From: Mike Ellis Date: Tue, 12 Nov 2024 14:23:34 -0500 Subject: [PATCH 001/802] Adding DMS serverless operators --- airflow/providers/amazon/aws/hooks/dms.py | 160 +++++ airflow/providers/amazon/aws/operators/dms.py | 561 ++++++++++++++++- airflow/providers/amazon/aws/triggers/dms.py | 225 +++++++ airflow/providers/amazon/aws/waiters/dms.json | 88 +++ airflow/providers/amazon/provider.yaml | 4 + tests/providers/amazon/aws/hooks/test_dms.py | 297 +++++++++ .../amazon/aws/operators/test_dms.py | 582 +++++++++++++++++- .../providers/amazon/aws/triggers/test_dms.py | 187 ++++++ .../providers/amazon/aws/waiters/test_dms.py | 139 +++++ .../amazon/aws/example_dms_serverless.py | 456 ++++++++++++++ 10 files changed, 2696 insertions(+), 3 deletions(-) create mode 100644 airflow/providers/amazon/aws/triggers/dms.py create mode 100644 airflow/providers/amazon/aws/waiters/dms.json create mode 100644 tests/providers/amazon/aws/triggers/test_dms.py create mode 100644 tests/providers/amazon/aws/waiters/test_dms.py create mode 100644 tests/system/providers/amazon/aws/example_dms_serverless.py diff --git a/airflow/providers/amazon/aws/hooks/dms.py b/airflow/providers/amazon/aws/hooks/dms.py index f4bc29cbe9be6..ae2a9f253a5c6 100644 --- a/airflow/providers/amazon/aws/hooks/dms.py +++ b/airflow/providers/amazon/aws/hooks/dms.py @@ -18,7 +18,12 @@ from __future__ import annotations import json +from datetime import datetime from enum import Enum +from typing import Any + +from botocore.exceptions import ClientError +from dateutil import parser from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseHook @@ -219,3 +224,158 @@ def wait_for_task_status(self, replication_task_arn: str, status: DmsTaskWaiterS ], WithoutSettings=True, ) + + def describe_replication_configs(self, filters: list[dict] | None = None, **kwargs) -> list[dict]: + """ + Return list of serverless replication configs. + + .. seealso:: + - :external+boto3:py:meth:`DatabaseMigrationService.Client.describe_replication_configs` + + :param filters: List of filter objects + :return: List of replication tasks + """ + filters = filters if filters is not None else [] + + try: + resp = self.conn.describe_replication_configs(Filters=filters, **kwargs) + return resp.get("ReplicationConfigs", []) + except Exception as ex: + self.log.error("Error while describing replication configs: %s", str(ex)) + return [] + + def create_replication_config( + self, + replication_config_id: str, + source_endpoint_arn: str, + target_endpoint_arn: str, + compute_config: dict[str, Any], + replication_type: str, + table_mappings: str, + additional_config_kwargs: dict[str, Any] | None = None, + **kwargs, + ): + """ + Create an AWS DMS Serverless configuration that can be used to start an DMS Serverless replication. + + .. seealso:: + - :external+boto3:py:meth:`DatabaseMigrationService.Client.create_replication_config` + + :param replicationConfigId: Unique identifier used to create a ReplicationConfigArn. + :param sourceEndpointArn: ARN of the source endpoint + :param targetEndpointArn: ARN of the target endpoint + :param computeConfig: Parameters for provisioning an DMS Serverless replication. + :param replicationType: type of DMS Serverless replication + :param tableMappings: JSON table mappings + :param tags: Key-value tag pairs + :param resourceId: Unique value or name that you set for a given resource that can be used to construct an Amazon Resource Name (ARN) for that resource. + :param supplementalSettings: JSON settings for specifying supplemental data + :param replicationSettings: JSON settings for DMS Serverless replications + + :return: ReplicationConfigArn + + """ + if additional_config_kwargs is None: + additional_config_kwargs = {} + try: + resp = self.conn.create_replication_config( + ReplicationConfigIdentifier=replication_config_id, + SourceEndpointArn=source_endpoint_arn, + TargetEndpointArn=target_endpoint_arn, + ComputeConfig=compute_config, + ReplicationType=replication_type, + TableMappings=table_mappings, + **additional_config_kwargs, + ) + arn = resp.get("ReplicationConfig", {}).get("ReplicationConfigArn") + self.log.info("Successfully created replication config: %s", arn) + return arn + + except ClientError as err: + err_str = f"Error: {err.get('Error','').get('Code','')}: {err.get('Error','').get('Message','')}" + self.log.error("Error while creating replication config: %s", err_str) + raise err + + def describe_replications(self, filters: list[dict[str, Any]] | None = None, **kwargs) -> list[dict]: + """ + Return list of serverless replications. + + .. seealso:: + - :external+boto3:py:meth:`DatabaseMigrationService.Client.describe_replications` + + :param filters: List of filter objects + :return: List of replications + """ + filters = filters if filters is not None else [] + try: + resp = self.conn.describe_replications(Filters=filters, **kwargs) + return resp.get("Replications", []) + except Exception: + return [] + + def delete_replication_config( + self, replication_config_arn: str, delay: int = 60, max_attempts: int = 120 + ): + """ + Delete an AWS DMS Serverless configuration. + + .. seealso:: + - :external+boto3:py:meth:`DatabaseMigrationService.Client.delete_replication_config` + + :param replication_config_arn: ReplicationConfigArn + """ + try: + self.log.info("Deleting replication config: %s", replication_config_arn) + + self.conn.delete_replication_config(ReplicationConfigArn=replication_config_arn) + + except ClientError as err: + err_str = ( + f"Error: {err.get('Error', '').get('Code', '')}: {err.get('Error', '').get('Message', '')}" + ) + self.log.error("Error while deleting replication config: %s", err_str) + raise err + + def start_replication( + self, + replication_config_arn: str, + start_replication_type: str, + cdc_start_time: datetime | str | None = None, + cdc_start_pos: str | None = None, + cdc_stop_pos: str | None = None, + ): + additional_args: dict[str, Any] = {} + + if cdc_start_time: + additional_args["CdcStartTime"] = ( + cdc_start_time if isinstance(cdc_start_time, datetime) else parser.parse(cdc_start_time) + ) + if cdc_start_pos: + additional_args["CdcStartPosition"] = cdc_start_pos + if cdc_stop_pos: + additional_args["CdcStopPosition"] = cdc_stop_pos + + try: + resp = self.conn.start_replication( + ReplicationConfigArn=replication_config_arn, + StartReplicationType=start_replication_type, + **additional_args, + ) + + return resp + except Exception as ex: + self.log.error("Error while starting replication: %s", str(ex)) + raise ex + + def stop_replication(self, replication_config_arn: str): + resp = self.conn.stop_replication(ReplicationConfigArn=replication_config_arn) + return resp + + def get_provision_status(self, replication_config_arn: str) -> str: + """Get the provisioning status for a serverless replication.""" + result = self.describe_replications( + filters=[{"Name": "replication-config-arn", "Values": [replication_config_arn]}] + ) + + provision_status = result[0].get("ProvisionData", {}).get("ProvisionState", "") + return provision_status diff --git a/airflow/providers/amazon/aws/operators/dms.py b/airflow/providers/amazon/aws/operators/dms.py index c564f802185a3..9fb33173884f3 100644 --- a/airflow/providers/amazon/aws/operators/dms.py +++ b/airflow/providers/amazon/aws/operators/dms.py @@ -17,11 +17,22 @@ # under the License. from __future__ import annotations -from typing import TYPE_CHECKING, Sequence +from datetime import datetime +from typing import TYPE_CHECKING, Any, Sequence +from airflow.configuration import conf +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.dms import DmsHook from airflow.providers.amazon.aws.operators.base_aws import AwsBaseOperator +from airflow.providers.amazon.aws.triggers.dms import ( + DmsReplicationCompleteTrigger, + DmsReplicationConfigDeletedTrigger, + DmsReplicationDeprovisionedTrigger, + DmsReplicationStoppedTrigger, + DmsReplicationTerminalStatusTrigger, +) from airflow.providers.amazon.aws.utils.mixins import aws_template_fields +from airflow.utils.context import Context if TYPE_CHECKING: from airflow.utils.context import Context @@ -277,3 +288,551 @@ def execute(self, context: Context): """Stop AWS DMS replication task from Airflow.""" self.hook.stop_replication_task(replication_task_arn=self.replication_task_arn) self.log.info("DMS replication task(%s) is stopping.", self.replication_task_arn) + + +class DmsDescribeReplicationConfigsOperator(AwsBaseOperator[DmsHook]): + """ + Describes AWS DMS Serverless replication configurations. + + .. seealso:: + For more information on how to use this operator, take a look at the guide: + :ref:`howto/operator:DmsDescribeReplicationConfigsOperator` + + :param describe_config_filter: Filters block for filtering results. + :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 + """ + + aws_hook_class = DmsHook + template_fields: Sequence[str] = aws_template_fields("filter") + template_fields_renderers: dict[str, Any] = {"filter": "json"} + + def __init__( + self, + *, + filter: list[dict] | None = None, + aws_conn_id: str | None = "aws_default", + **kwargs, + ): + super().__init__(aws_conn_id=aws_conn_id, **kwargs) + self.filter = filter + + def execute(self, context: Context) -> list: + """ + Describe AWS DMS replication configurations. + + :return: List of replication configurations + """ + return self.hook.describe_replication_configs(filters=self.filter) + + +class DmsCreateReplicationConfigOperator(AwsBaseOperator[DmsHook]): + """ + + Creates an AWS DMS Serverless replication configuration. + + .. seealso:: + For more information on how to use this operator, take a look at the guide: + :ref:`howto/operator:DmsCreateReplicationConfigOperator` + + :param replication_config_id: Unique identifier used to create a ReplicationConfigArn. + :param source_endpoint_arn: ARN of the source endpoint + :param target_endpoint_arn: ARN of the target endpoint + :param compute_config: Parameters for provisioning an DMS Serverless replication. + :param replication_type: type of DMS Serverless replication + :param table_mappings: JSON table mappings + :param tags: Key-value tag pairs + :param additional_config_kwargs: Additional configuration parameters for DMS Serverless replication. Passed directly to the API + :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 + """ + + aws_hook_class = DmsHook + template_fields: Sequence[str] = aws_template_fields( + "replication_config_id", + "source_endpoint_arn", + "target_endpoint_arn", + "compute_config", + "replication_type", + "table_mappings", + ) + + template_fields_renderers: dict[str, Any] = {"compute_config": "json", "tableMappings": "json"} + + def __init__( + self, + *, + replication_config_id: str, + source_endpoint_arn: str, + target_endpoint_arn: str, + compute_config: dict[str, Any], + replication_type: str, + table_mappings: str, + additional_config_kwargs: dict | None = None, + aws_conn_id: str | None = "aws_default", + **kwargs, + ): + super().__init__( + aws_conn_id=aws_conn_id, + **kwargs, + ) + + self.replication_config_id = replication_config_id + self.source_endpoint_arn = source_endpoint_arn + self.target_endpoint_arn = target_endpoint_arn + self.compute_config = compute_config + self.replication_type = replication_type + self.table_mappings = table_mappings + self.additional_config_kwargs = additional_config_kwargs or {} + + def execute(self, context: Context) -> str: + resp = self.hook.create_replication_config( + replication_config_id=self.replication_config_id, + source_endpoint_arn=self.source_endpoint_arn, + target_endpoint_arn=self.target_endpoint_arn, + compute_config=self.compute_config, + replication_type=self.replication_type, + table_mappings=self.table_mappings, + additional_config_kwargs=self.additional_config_kwargs, + ) + + self.log.info("DMS replication config(%s) has been created.", self.replication_config_id) + return resp + + +class DmsDeleteReplicationConfigOperator(AwsBaseOperator[DmsHook]): + """ + + Deletes an AWS DMS Serverless replication configuration. + + .. seealso:: + For more information on how to use this operator, take a look at the guide: + :ref:`howto/operator:DmsDeleteReplicationConfigOperator` + + :param replication_config_arn: ARN of the replication config + :param wait_for_completion: If True, waits for the replication config to be deleted before returning. + If False, the operator will return immediately after the request is made. + :param deferrable: Run the operator in deferrable mode. + :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 + """ + + aws_hook_class = DmsHook + template_fields: Sequence[str] = aws_template_fields("replication_config_arn") + + VALID_STATES = ["failed", "stopped", "created"] + DELETING_STATES = ["deleting"] + TERMINAL_PROVISION_STATES = ["deprovisioned", ""] + + def __init__( + self, + *, + replication_config_arn: str, + wait_for_completion: bool = True, + deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), + waiter_delay: int = 5, + waiter_max_attempts: int = 60, + aws_conn_id: str | None = "aws_default", + **kwargs, + ): + super().__init__( + aws_conn_id=aws_conn_id, + **kwargs, + ) + + self.replication_config_arn = replication_config_arn + self.wait_for_completion = wait_for_completion + self.deferrable = deferrable + self.waiter_delay = waiter_delay + self.waiter_max_attempts = waiter_max_attempts + + def execute(self, context: Context) -> None: + results = self.hook.describe_replications( + filters=[{"Name": "replication-config-arn", "Values": [self.replication_config_arn]}] + ) + + if len(results) > 0: + current_state = results[0].get("Status", "") + self.log.info( + "Current state of replication config(%s) is %s.", self.replication_config_arn, current_state + ) + # replication must be deprovisioned before deleting + provision_status = self.hook.get_provision_status( + replication_config_arn=self.replication_config_arn + ) + + if ( + current_state.lower() in self.VALID_STATES + and provision_status in self.TERMINAL_PROVISION_STATES + ): + self.log.info("DMS replication config(%s) is in valid state.", self.replication_config_arn) + + self.hook.delete_replication_config( + replication_config_arn=self.replication_config_arn, + delay=self.waiter_delay, + max_attempts=self.waiter_max_attempts, + ) + self.handle_delete_wait() + + else: + # Must be in a terminal state to delete + self.log.info( + "DMS replication config(%s) cannot be deleted until replication is in terminal state and deprovisioned. Waiting for terminal state.", + self.replication_config_arn, + ) + if self.deferrable: + if current_state.lower() not in self.VALID_STATES: + self.log.info("Deferring until terminal status reached.") + self.defer( + trigger=DmsReplicationTerminalStatusTrigger( + replication_config_arn=self.replication_config_arn, + waiter_delay=self.waiter_delay, + waiter_max_attempts=self.waiter_max_attempts, + aws_conn_id=self.aws_conn_id, + ), + method_name="retry_execution", + ) + if provision_status not in self.TERMINAL_PROVISION_STATES: # not deprovisioned: + self.log.info("Deferring until deprovisioning completes.") + self.defer( + trigger=DmsReplicationDeprovisionedTrigger( + replication_config_arn=self.replication_config_arn, + waiter_delay=self.waiter_delay, + waiter_max_attempts=self.waiter_max_attempts, + aws_conn_id=self.aws_conn_id, + ), + method_name="retry_execution", + ) + + else: + self.hook.get_waiter("replication_terminal_status").wait( + Filters=[{"Name": "replication-config-arn", "Values": [self.replication_config_arn]}], + WaiterConfig={"Delay": self.waiter_delay, "MaxAttempts": self.waiter_max_attempts}, + ) + self.hook.get_waiter("replication_deprovisioned").wait( + Filters=[{"Name": "replication-config-arn", "Values": [self.replication_config_arn]}], + WaiterConfig={"Delay": self.waiter_delay, "MaxAttempts": self.waiter_max_attempts}, + ) + self.hook.delete_replication_config(self.replication_config_arn) + self.handle_delete_wait() + + else: + self.log.info("DMS replication config(%s) does not exist.", self.replication_config_arn) + + def handle_delete_wait(self): + if self.deferrable: + self.log.info("Deferring until replication config is deleted.") + self.defer( + trigger=DmsReplicationConfigDeletedTrigger( + replication_config_arn=self.replication_config_arn, + waiter_delay=self.waiter_delay, + waiter_max_attempts=self.waiter_max_attempts, + aws_conn_id=self.aws_conn_id, + ), + method_name="execute_complete", + ) + + if self.wait_for_completion: + self.hook.get_waiter("replication_config_deleted").wait( + Filters=[{"Name": "replication-config-arn", "Values": [self.replication_config_arn]}], + WaiterConfig={"Delay": self.waiter_delay, "MaxAttempts": self.waiter_max_attempts}, + ) + self.log.info("DMS replication config(%s) deleted.", self.replication_config_arn) + + def execute_complete(self, context, event=None): + self.replication_config_arn = event.get("replication_config_arn") + self.log.info("DMS replication config(%s) deleted.", self.replication_config_arn) + + def retry_execution(self, context, event=None): + self.replication_config_arn = event.get("replication_config_arn") + self.log.info("Retrying replication config(%s) deletion.", self.replication_config_arn) + self.execute(context) + + +class DmsDescribeReplicationsOperator(AwsBaseOperator[DmsHook]): + """ + Describes AWS DMS Serverless replications. + + .. seealso:: + For more information on how to use this operator, take a look at the guide: + :ref:`howto/operator:DmsDescribeReplicationsOperator` + + :param filter: Filters block for filtering results. + + :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 + """ + + aws_hook_class = DmsHook + template_fields: Sequence[str] = aws_template_fields("filter") + template_fields_renderers: dict[str, Any] = {"filter": "json"} + + def __init__( + self, + *, + filter: list[dict[str, Any]] | None = None, + aws_conn_id: str | None = "aws_default", + **kwargs, + ): + super().__init__( + aws_conn_id=aws_conn_id, + **kwargs, + ) + + self.filter = filter + + def execute(self, context: Context) -> list[dict[str, Any]]: + """ + Describe AWS DMS replications. + + :return: Replications + """ + return self.hook.describe_replications(self.filter) + + +class DmsStartReplicationOperator(AwsBaseOperator[DmsHook]): + """ + Starts an AWS DMS Serverless replication. + + .. seealso:: + For more information on how to use this operator, take a look at the guide: + :ref:`howto/operator:DmsStartReplicationOperator` + + :param replication_config_arn: ARN of the replication config + :param replication_start_type: Type of replication. + :param cdc_start_time: Start time of CDC + :param cdc_start_pos: Indicates when to start CDC. + :param cdc_stop_pos: Indicates when to stop CDC. + :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 + """ + + RUNNING_STATES = ["running"] + STARTABLE_STATES = ["stopped", "failed", "created"] + TERMINAL_STATES = ["failed", "stopped", "created"] + TERMINAL_PROVISION_STATES = ["deprovisioned", ""] + + aws_hook_class = DmsHook + template_fields: Sequence[str] = aws_template_fields( + "replication_config_arn", "replication_start_type", "cdc_start_time", "cdc_start_pos", "cdc_stop_pos" + ) + + def __init__( + self, + *, + replication_config_arn: str, + replication_start_type: str, + cdc_start_time: datetime | str | None = None, + cdc_start_pos: str | None = None, + cdc_stop_pos: str | None = None, + wait_for_completion: bool = True, + waiter_delay: int = 30, + waiter_max_attempts: int = 60, + deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), + aws_conn_id: str | None = "aws_default", + **kwargs, + ): + super().__init__( + aws_conn_id=aws_conn_id, + **kwargs, + ) + + self.replication_config_arn = replication_config_arn + self.replication_start_type = replication_start_type + self.cdc_start_time = cdc_start_time + self.cdc_start_pos = cdc_start_pos + self.cdc_stop_pos = cdc_stop_pos + self.deferrable = deferrable + self.waiter_delay = waiter_delay + self.waiter_max_attempts = waiter_max_attempts + self.wait_for_completion = wait_for_completion + + if self.cdc_start_time and self.cdc_start_pos: + raise AirflowException("Only one of cdc_start_time or cdc_start_pos should be provided.") + + def execute(self, context: Context): + result = self.hook.describe_replications( + filters=[{"Name": "replication-config-arn", "Values": [self.replication_config_arn]}] + ) + + try: + current_status = result[0].get("Status", "") + except Exception as ex: + self.log.error("Error while getting replication status: %s. Unable to start replication", str(ex)) + raise ex + + provision_status = self.hook.get_provision_status(replication_config_arn=self.replication_config_arn) + + if provision_status == "deprovisioning": + # wait for deprovisioning to complete before start/restart + self.log.info( + "Replication is deprovisioning. Must wait for deprovisioning before running replication" + ) + if self.deferrable: + self.log.info("Deferring until deprovisioning completes.") + self.defer( + trigger=DmsReplicationDeprovisionedTrigger( + replication_config_arn=self.replication_config_arn, + waiter_delay=self.waiter_delay, + waiter_max_attempts=self.waiter_max_attempts, + aws_conn_id=self.aws_conn_id, + ), + method_name="retry_execution", + ) + else: + self.hook.get_waiter("replication_deprovisioned").wait( + Filters=[{"Name": "replication-config-arn", "Values": [self.replication_config_arn]}], + WaiterConfig={"Delay": self.waiter_delay, "MaxAttempts": self.waiter_max_attempts}, + ) + provision_status = self.hook.get_provision_status( + replication_config_arn=self.replication_config_arn + ) + self.log.info("Replication deprovisioning complete. Provision status: %s", provision_status) + + if ( + current_status.lower() in self.STARTABLE_STATES + and provision_status in self.TERMINAL_PROVISION_STATES + ): + resp = self.hook.start_replication( + replication_config_arn=self.replication_config_arn, + start_replication_type=self.replication_start_type, + cdc_start_time=self.cdc_start_time, + cdc_start_pos=self.cdc_start_pos, + cdc_stop_pos=self.cdc_stop_pos, + ) + + current_status = resp.get("Replication", {}).get("Status", "Unknown") + self.log.info( + "Replication(%s) started with status %s.", + self.replication_config_arn, + current_status, + ) + + if self.deferrable: + self.log.info("Deferring until %s replication completes.", self.replication_config_arn) + self.defer( + trigger=DmsReplicationCompleteTrigger( + replication_config_arn=self.replication_config_arn, + waiter_delay=self.waiter_delay, + waiter_max_attempts=self.waiter_max_attempts, + aws_conn_id=self.aws_conn_id, + ), + method_name="execute_complete", + ) + + if self.wait_for_completion: + self.log.info("Waiting for %s replication to complete.", self.replication_config_arn) + + self.hook.get_waiter("replication_complete").wait( + Filters=[{"Name": "replication-config-arn", "Values": [self.replication_config_arn]}], + WaiterConfig={"Delay": self.waiter_delay, "MaxAttempts": self.waiter_max_attempts}, + ) + self.log.info("Replication(%s) has completed.", self.replication_config_arn) + + else: + self.log.info("Replication(%s) is not in startable state.", self.replication_config_arn) + self.log.info("Status: %s Provision status: %s", current_status, provision_status) + + def execute_complete(self, context, event=None): + self.replication_config_arn = event.get("replication_config_arn") + self.log.info("Replication(%s) has completed.", self.replication_config_arn) + + def retry_execution(self, context, event=None): + self.replication_config_arn = event.get("replication_config_arn") + self.log.info("Retrying replication %s.", self.replication_config_arn) + self.execute(context) + + +class DmsStopReplicationOperator(AwsBaseOperator[DmsHook]): + """ + Stops an AWS DMS Serverless replication. + + .. seealso:: + For more information on how to use this operator, take a look at the guide: + :ref:`howto/operator:DmsStopReplicationOperator` + + :param replication_config_arn: ARN of the replication config + :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 + """ + + STOPPED_STATES = ["stopped"] + NON_STOPPABLE_STATES = ["stopped"] + + aws_hook_class = DmsHook + template_fields: Sequence[str] = aws_template_fields("replication_config_arn") + + def __init__( + self, + *, + replication_config_arn: str, + wait_for_completion: bool = True, + waiter_delay: int = 30, + waiter_max_attempts: int = 60, + deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), + aws_conn_id: str | None = "aws_default", + **kwargs, + ): + super().__init__( + aws_conn_id=aws_conn_id, + **kwargs, + ) + + self.replication_config_arn = replication_config_arn + self.wait_for_completion = wait_for_completion + self.waiter_delay = waiter_delay + self.waiter_max_attempts = waiter_max_attempts + self.deferrable = deferrable + + def execute(self, context: Context) -> None: + results = self.hook.describe_replications( + filters=[{"Name": "replication-config-arn", "Values": [self.replication_config_arn]}] + ) + + current_state = results[0].get("Status", "") + self.log.info( + "Current state of replication config(%s) is %s.", self.replication_config_arn, current_state + ) + + if current_state.lower() in self.STOPPED_STATES: + self.log.info("DMS replication config(%s) is already stopped.", self.replication_config_arn) + else: + resp = self.hook.stop_replication(self.replication_config_arn) + status = resp.get("Replication", {}).get("Status", "Unknown") + self.log.info( + "Stopping DMS replication config(%s). Current status: %s", self.replication_config_arn, status + ) + + if self.deferrable: + self.log.info("Deferring until %s replication stops.", self.replication_config_arn) + self.defer( + trigger=DmsReplicationStoppedTrigger( + replication_config_arn=self.replication_config_arn, + waiter_delay=self.waiter_delay, + waiter_max_attempts=self.waiter_max_attempts, + aws_conn_id=self.aws_conn_id, + ), + method_name="execute_complete", + ) + if self.wait_for_completion: + self.log.info("Waiting for %s replication to stop.", self.replication_config_arn) + self.hook.get_waiter("replication_stopped").wait( + Filters=[{"Name": "replication-config-arn", "Values": [self.replication_config_arn]}], + WaiterConfig={"Delay": self.waiter_delay, "MaxAttempts": self.waiter_max_attempts}, + ) + + def execute_complete(self, context, event=None): + self.replication_config_arn = event.get("replication_config_arn") + self.log.info("Replication(%s) has stopped.", self.replication_config_arn) diff --git a/airflow/providers/amazon/aws/triggers/dms.py b/airflow/providers/amazon/aws/triggers/dms.py new file mode 100644 index 0000000000000..3e0e563f2f4f7 --- /dev/null +++ b/airflow/providers/amazon/aws/triggers/dms.py @@ -0,0 +1,225 @@ +# 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 + +from typing import TYPE_CHECKING + +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 + +if TYPE_CHECKING: + from airflow.providers.amazon.aws.hooks.base_aws import AwsGenericHook + + +class DmsReplicationTerminalStatusTrigger(AwsBaseWaiterTrigger): + """ + Trigger when an AWS DMS Serverless replication is in a terminal state. + + :param replication_config_arn: The ARN of the replication config. + :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. + """ + + def __init__( + self, + replication_config_arn: str, + waiter_delay: int = 30, + waiter_max_attempts: int = 60, + aws_conn_id: str | None = "aws_default", + ) -> None: + super().__init__( + serialized_fields={"replication_config_arn": replication_config_arn}, + waiter_name="replication_terminal_status", + waiter_delay=waiter_delay, + waiter_args={"Filters": [{"Name": "replication-config-arn", "Values": [replication_config_arn]}]}, + waiter_max_attempts=waiter_max_attempts, + failure_message="Replication failed to reach terminal status.", + status_message="Status replication is", + status_queries=["Replications[0].Status"], + return_key="replication_config_arn", + return_value=replication_config_arn, + aws_conn_id=aws_conn_id, + ) + + def hook(self) -> AwsGenericHook: + return DmsHook( + self.aws_conn_id, + verify=self.verify, + config=self.botocore_config, + ) + + +class DmsReplicationConfigDeletedTrigger(AwsBaseWaiterTrigger): + """ + Trigger when an AWS DMS Serverless replication config is deleted. + + :param replication_config_arn: The ARN of the replication config. + :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. + """ + + def __init__( + self, + replication_config_arn: str, + waiter_delay: int = 30, + waiter_max_attempts: int = 60, + aws_conn_id: str | None = "aws_default", + ) -> None: + super().__init__( + serialized_fields={"replication_config_arn": replication_config_arn}, + # serialized_fields={ + # "Filters": [{"Name": "replication-config-arn", "Values": [replication_config_arn]}] + # }, + waiter_name="replication_config_deleted", + waiter_delay=waiter_delay, + waiter_args={"Filters": [{"Name": "replication-config-arn", "Values": [replication_config_arn]}]}, + waiter_max_attempts=waiter_max_attempts, + failure_message="Replication config failed to be deleted.", + status_message="Status replication config is", + status_queries=["ReplicationConfigs[0].Status"], + return_key="replication_config_arn", + return_value=replication_config_arn, + aws_conn_id=aws_conn_id, + ) + + def hook(self) -> AwsGenericHook: + return DmsHook( + self.aws_conn_id, + verify=self.verify, + config=self.botocore_config, + ) + + +class DmsReplicationCompleteTrigger(AwsBaseWaiterTrigger): + """ + Trigger when an AWS DMS Serverless replication completes. + + :param replication_config_arn: The ARN of the replication config. + :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. + """ + + def __init__( + self, + replication_config_arn: str, + waiter_delay: int = 30, + waiter_max_attempts: int = 60, + aws_conn_id: str | None = "aws_default", + ) -> None: + super().__init__( + serialized_fields={"replication_config_arn": replication_config_arn}, + # "Filters": [{"Name": "replication-config-arn", "Values": [replication_config_arn]}]}, + waiter_name="replication_complete", + waiter_delay=waiter_delay, + waiter_args={"Filters": [{"Name": "replication-config-arn", "Values": [replication_config_arn]}]}, + waiter_max_attempts=waiter_max_attempts, + failure_message="Replication failed to reach terminal status.", + status_message="Status replication is", + status_queries=["Replications[0].Status"], + return_key="replication_config_arn", + return_value=replication_config_arn, + aws_conn_id=aws_conn_id, + ) + + def hook(self) -> AwsGenericHook: + return DmsHook( + self.aws_conn_id, + verify=self.verify, + config=self.botocore_config, + ) + + +class DmsReplicationStoppedTrigger(AwsBaseWaiterTrigger): + """ + Trigger when an AWS DMS Serverless replication is stopped. + + :param replication_config_arn: The ARN of the replication config. + :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. + """ + + def __init__( + self, + replication_config_arn: str, + waiter_delay: int = 30, + waiter_max_attempts: int = 60, + aws_conn_id: str | None = "aws_default", + ) -> None: + super().__init__( + serialized_fields={"replication_config_arn": replication_config_arn}, + waiter_name="replication_stopped", + waiter_delay=waiter_delay, + waiter_args={"Filters": [{"Name": "replication-config-arn", "Values": [replication_config_arn]}]}, + waiter_max_attempts=waiter_max_attempts, + failure_message="Replication failed to stop.", + status_message="Status replication is", + status_queries=["Replications[0].Status"], + return_key="replication_config_arn", + return_value=replication_config_arn, + aws_conn_id=aws_conn_id, + ) + + def hook(self) -> AwsGenericHook: + return DmsHook( + self.aws_conn_id, + verify=self.verify, + config=self.botocore_config, + ) + + +class DmsReplicationDeprovisionedTrigger(AwsBaseWaiterTrigger): + """ + Trigger when an AWS DMS Serverless replication is deprovisioned. + + :param replication_config_arn: The ARN of the replication config. + :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. + """ + + def __init__( + self, + replication_config_arn: str, + waiter_delay: int = 30, + waiter_max_attempts: int = 60, + aws_conn_id: str | None = "aws_default", + ) -> None: + super().__init__( + serialized_fields={"replication_config_arn": replication_config_arn}, + waiter_name="replication_deprovisioned", + waiter_delay=waiter_delay, + waiter_args={"Filters": [{"Name": "replication-config-arn", "Values": [replication_config_arn]}]}, + waiter_max_attempts=waiter_max_attempts, + failure_message="Replication failed to deprovision.", + status_message="Status replication is", + status_queries=["Replications[0].ProvisionData.ProvisionState"], + return_key="replication_config_arn", + return_value=replication_config_arn, + aws_conn_id=aws_conn_id, + ) + + def hook(self) -> AwsGenericHook: + return DmsHook( + self.aws_conn_id, + verify=self.verify, + config=self.botocore_config, + ) diff --git a/airflow/providers/amazon/aws/waiters/dms.json b/airflow/providers/amazon/aws/waiters/dms.json new file mode 100644 index 0000000000000..4fab7825ae1a4 --- /dev/null +++ b/airflow/providers/amazon/aws/waiters/dms.json @@ -0,0 +1,88 @@ +{ + "version": 2, + "waiters": { + "replication_deprovisioned": { + "operation": "DescribeReplications", + "delay": 30, + "maxAttempts": 60, + "acceptors": [ + { + "matcher": "path", + "argument": "Replications[0].ProvisionData.ProvisionState", + "expected": "deprovisioned", + "state": "success" + } + ] + }, + "replication_terminal_status": { + "operation": "DescribeReplications", + "delay": 30, + "maxAttempts": 60, + "acceptors": [ + { + "matcher": "path", + "argument": "Replications[0].ProvisionData.ProvisionState", + "expected": "deprovisioning", + "state": "retry" + }, + { + "matcher": "path", + "argument": "Replications[0].Status", + "expected": "failed", + "state": "success" + }, + { + "matcher": "path", + "argument": "Replications[0].Status", + "expected": "stopped", + "state": "success" + } + ] + }, + "replication_complete": { + "operation": "DescribeReplications", + "delay": 30, + "maxAttempts": 60, + "acceptors": [ + { + "matcher": "path", + "argument": "Replications[0].Status", + "expected": "failed", + "state": "failure" + }, + { + "matcher": "path", + "argument": "Replications[0].Status", + "expected": "stopped", + "state": "success" + } + ] + }, + "replication_config_deleted": { + "operation": "DescribeReplicationConfigs", + "delay": 5, + "maxAttempts": 60, + "acceptors": [ + { + "matcher": "error", + "expected": "ResourceNotFoundFault", + "state": "success", + "argument": "Error.Code" + } + ] + }, + "replication_stopped": { + "operation": "DescribeReplications", + "delay": 5, + "maxAttempts": 60, + "acceptors": [ + { + "matcher": "path", + "argument": "Replications[0].Status", + "expected": "stopped", + "state": "success" + } + ] + } + } +} diff --git a/airflow/providers/amazon/provider.yaml b/airflow/providers/amazon/provider.yaml index 83d66de69a85b..e98a7b53a7253 100644 --- a/airflow/providers/amazon/provider.yaml +++ b/airflow/providers/amazon/provider.yaml @@ -769,6 +769,10 @@ triggers: - integration-name: Amazon Neptune python-modules: - airflow.providers.amazon.aws.triggers.neptune + - integration-name: AWS Database Migration Service + python-modules: + - airflow.providers.amazon.aws.triggers.dms + transfers: - source-integration-name: Amazon DynamoDB diff --git a/tests/providers/amazon/aws/hooks/test_dms.py b/tests/providers/amazon/aws/hooks/test_dms.py index 9d66df55c25eb..ef7095d54a1cf 100644 --- a/tests/providers/amazon/aws/hooks/test_dms.py +++ b/tests/providers/amazon/aws/hooks/test_dms.py @@ -17,10 +17,12 @@ from __future__ import annotations import json +from datetime import datetime from typing import Any from unittest import mock import pytest +from dateutil import parser from airflow.providers.amazon.aws.hooks.dms import DmsHook, DmsTaskWaiterStatus @@ -66,6 +68,187 @@ 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": [ + { + "ReplicationConfigIdentifier": "1111", + "ReplicationConfigArn": "arn:aws:my-arn", + "SourceEndpointArn": "source-endpoint-arn", + "TargetEndpointArn": "target-endpoint-arn", + "ReplicationType": "cdc", + "ComputeConfig": { + "AvailabilityZone": "az1", + "MaxCapacityUnits": 10, + "MinCapacityUnits": 20, + "MultiAZ": True, + "PreferredMaintenanceWindow": "string", + "ReplicationSubnetGroupId": "string", + "VpcSecurityGroupIds": [ + "string", + ], + }, + "ReplicationSettings": "string", + "SupplementalSettings": "string", + "TableMappings": "string", + "ReplicationConfigCreateTime": datetime(2015, 1, 1), + "ReplicationConfigUpdateTime": datetime(2015, 1, 1), + }, + { + "ReplicationConfigIdentifier": "2222", + "ReplicationConfigArn": "arn:aws:my-arn", + "SourceEndpointArn": "source-endpoint-arn", + "TargetEndpointArn": "target-endpoint-arn", + "ReplicationType": "full-load-and-cdc", + "ComputeConfig": { + "AvailabilityZone": "string", + "DnsNameServers": "string", + "KmsKeyId": "string", + "MaxCapacityUnits": 1, + "MinCapacityUnits": 30, + "MultiAZ": False, + "PreferredMaintenanceWindow": "string", + "ReplicationSubnetGroupId": "string", + "VpcSecurityGroupIds": [ + "string", + ], + }, + "ReplicationSettings": "string", + "SupplementalSettings": "string", + "TableMappings": "string", + "ReplicationConfigCreateTime": datetime(2015, 1, 1), + "ReplicationConfigUpdateTime": datetime(2015, 2, 1), + }, + ], +} + +MOCK_REPLICATION_CONFIG: dict[str, Any] = { + "ReplicationConfigIdentifier": "test-config", + "SourceEndpointArn": "arn:aws:dms:us-east-1:123456789012:endpoint:RZZK4EZW5UANC7Y3P4E776WHBE", + "TargetEndpointArn": "arn:aws:dms:us-east-1:123456789012:endpoint:GVBUJQXJZASXWHTWCLN2WNT57E", + "ComputeConfig": { + "MaxCapacityUnits": 2, + "MinCapacityUnits": 4, + }, + "ReplicationType": "full-load", + "TableMappings": json.dumps( + { + "TableMappings": [ + { + "Type": "Selection", + "RuleId": 123, + "RuleName": "test-rule", + "SourceSchema": "/", + "SourceTable": "/", + } + ] + } + ), + "ReplicationSettings": "string", + "SupplementalSettings": "string", + "ResourceIdentifier": "string", +} + +MOCK_REPLICATION_CONFIG_RESP: dict[str, Any] = { + "ReplicationConfig": { + "ReplicationConfigIdentifier": "test-config", + "ReplicationConfigArn": "arn:aws:dms:us-east-1:123456789012:replication-config/test-config", + "SourceEndpointArn": "arn:aws:dms:us-east-1:123456789012:endpoint:RZZK4EZW5UANC7Y3P4E776WHBE", + "TargetEndpointArn": "arn:aws:dms:us-east-1:123456789012:endpoint:GVBUJQXJZASXWHTWCLN2WNT57E", + "ReplicationType": "full-load", + } +} + +MOCK_DESCRIBE_REPLICATIONS_RESP = { + "Marker": "string", + "Replications": [ + { + "ReplicationConfigIdentifier": "test-config", + "ReplicationConfigArn": "arn:aws:dms:us-east-1:123456789012:replication-config/test-config", + "SourceEndpointArn": "arn:aws:dms:us-east-1:123456789012:endpoint:RZZK4EZW5UANC7Y3P4E776WHBE", + "TargetEndpointArn": "arn:aws:dms:us-east-1:123456789012:endpoint:GVBUJQXJZASXWHTWCLN2WNT57E", + "ReplicationType": "full-load", + "Status": "CREATED", + "ProvisionData": { + "ProvisionState": "string", + "ProvisionedCapacityUnits": 123, + "DateProvisioned": datetime(2015, 1, 1), + "IsNewProvisioningAvailable": False, + "DateNewProvisioningDataAvailable": datetime(2015, 1, 1), + "ReasonForNewProvisioningData": "string", + }, + "StopReason": "string", + "FailureMessages": [ + "string", + ], + "ReplicationStats": { + "FullLoadProgressPercent": 123, + "ElapsedTimeMillis": 123, + "TablesLoaded": 123, + "TablesLoading": 123, + "TablesQueued": 123, + "TablesErrored": 123, + "FreshStartDate": datetime(2015, 1, 1), + "StartDate": datetime(2015, 1, 1), + "StopDate": datetime(2015, 1, 1), + "FullLoadStartDate": datetime(2015, 1, 1), + "FullLoadFinishDate": datetime(2015, 1, 1), + }, + "StartReplicationType": "string", + "CdcStartTime": datetime(2015, 1, 1), + "CdcStartPosition": "string", + "CdcStopPosition": "string", + "RecoveryCheckpoint": "string", + "ReplicationCreateTime": datetime(2015, 1, 1), + "ReplicationUpdateTime": datetime(2015, 1, 1), + "ReplicationLastStopTime": datetime(2015, 1, 1), + "ReplicationDeprovisionTime": datetime(2015, 1, 1), + }, + { + "ReplicationConfigIdentifier": "test-config-2", + "ReplicationConfigArn": "arn:aws:dms:us-east-1:123456789012:replication-config/test-config-2", + "SourceEndpointArn": "arn:aws:dms:us-east-1:123456789012:endpoint:RZZK4EZW5UANC7Y3P4E776WHBE", + "TargetEndpointArn": "arn:aws:dms:us-east-1:123456789012:endpoint:GVBUJQXJZASXWHTWCLN2WNT57E", + "ReplicationType": "cdc-only", + "Status": "CREATED", + "ProvisionData": { + "ProvisionState": "string", + "ProvisionedCapacityUnits": 123, + "DateProvisioned": datetime(2015, 1, 1), + "IsNewProvisioningAvailable": False, + "DateNewProvisioningDataAvailable": datetime(2015, 1, 1), + "ReasonForNewProvisioningData": "string", + }, + "StopReason": "string", + "FailureMessages": [ + "string", + ], + "ReplicationStats": { + "FullLoadProgressPercent": 123, + "ElapsedTimeMillis": 123, + "TablesLoaded": 123, + "TablesLoading": 123, + "TablesQueued": 123, + "TablesErrored": 123, + "FreshStartDate": datetime(2015, 1, 1), + "StartDate": datetime(2015, 1, 1), + "StopDate": datetime(2015, 1, 1), + "FullLoadStartDate": datetime(2015, 1, 1), + "FullLoadFinishDate": datetime(2015, 1, 1), + }, + "StartReplicationType": "string", + "CdcStartTime": datetime(2015, 1, 1), + "CdcStartPosition": "string", + "CdcStopPosition": "string", + "RecoveryCheckpoint": "string", + "ReplicationCreateTime": datetime(2015, 1, 1), + "ReplicationUpdateTime": datetime(2015, 1, 1), + "ReplicationLastStopTime": datetime(2015, 1, 1), + "ReplicationDeprovisionTime": datetime(2015, 1, 1), + }, + ], +} + class TestDmsHook: def setup_method(self): @@ -213,3 +396,117 @@ def test_wait_for_task_status(self, mock_conn): } mock_conn.return_value.get_waiter.assert_called_with("replication_task_deleted") mock_conn.return_value.get_waiter.return_value.wait.assert_called_with(**expected_waiter_call_params) + + @mock.patch.object(DmsHook, "conn") + def test_describe_config_no_filter(self, mock_conn): + mock_conn.describe_replication_configs.return_value = MOCK_CONFIG_RESPONSE + resp = self.dms.describe_replication_configs() + + assert len(resp) == 2 + + @mock.patch.object(DmsHook, "conn") + def test_describe_config_filter(self, mock_conn): + filter = [{"Name": "replication-type", "Values": ["cdc"]}] + self.dms.describe_replication_configs(filters=filter) + mock_conn.describe_replication_configs.assert_called_with(Filters=filter) + + @mock.patch.object(DmsHook, "conn") + def test_create_repl_config_kwargs(self, mock_conn): + self.dms.create_replication_config( + replication_config_id=MOCK_REPLICATION_CONFIG["ReplicationConfigIdentifier"], + source_endpoint_arn=MOCK_REPLICATION_CONFIG["SourceEndpointArn"], + target_endpoint_arn=MOCK_REPLICATION_CONFIG["TargetEndpointArn"], + compute_config=MOCK_REPLICATION_CONFIG["ComputeConfig"], + replication_type=MOCK_REPLICATION_CONFIG["ReplicationType"], + table_mappings=MOCK_REPLICATION_CONFIG["TableMappings"], + additional_config_kwargs={ + "ReplicationSettings": MOCK_REPLICATION_CONFIG["ReplicationSettings"], + "SupplementalSettings": MOCK_REPLICATION_CONFIG["SupplementalSettings"], + }, + ) + + mock_conn.create_replication_config.assert_called_with( + ReplicationConfigIdentifier=MOCK_REPLICATION_CONFIG["ReplicationConfigIdentifier"], + SourceEndpointArn=MOCK_REPLICATION_CONFIG["SourceEndpointArn"], + TargetEndpointArn=MOCK_REPLICATION_CONFIG["TargetEndpointArn"], + ComputeConfig=MOCK_REPLICATION_CONFIG["ComputeConfig"], + ReplicationType=MOCK_REPLICATION_CONFIG["ReplicationType"], + TableMappings=MOCK_REPLICATION_CONFIG["TableMappings"], + ReplicationSettings=MOCK_REPLICATION_CONFIG["ReplicationSettings"], + SupplementalSettings=MOCK_REPLICATION_CONFIG["SupplementalSettings"], + ) + + self.dms.create_replication_config( + replication_config_id=MOCK_REPLICATION_CONFIG["ReplicationConfigIdentifier"], + source_endpoint_arn=MOCK_REPLICATION_CONFIG["SourceEndpointArn"], + target_endpoint_arn=MOCK_REPLICATION_CONFIG["TargetEndpointArn"], + compute_config=MOCK_REPLICATION_CONFIG["ComputeConfig"], + replication_type=MOCK_REPLICATION_CONFIG["ReplicationType"], + table_mappings=MOCK_REPLICATION_CONFIG["TableMappings"], + ) + mock_conn.create_replication_config.assert_called_with( + ReplicationConfigIdentifier=MOCK_REPLICATION_CONFIG["ReplicationConfigIdentifier"], + SourceEndpointArn=MOCK_REPLICATION_CONFIG["SourceEndpointArn"], + TargetEndpointArn=MOCK_REPLICATION_CONFIG["TargetEndpointArn"], + ComputeConfig=MOCK_REPLICATION_CONFIG["ComputeConfig"], + ReplicationType=MOCK_REPLICATION_CONFIG["ReplicationType"], + TableMappings=MOCK_REPLICATION_CONFIG["TableMappings"], + ) + + @mock.patch.object(DmsHook, "conn") + def test_create_repl_config(self, mock_conn): + mock_conn.create_replication_config.return_value = MOCK_REPLICATION_CONFIG_RESP + + resp = self.dms.create_replication_config( + replication_config_id=MOCK_REPLICATION_CONFIG["ReplicationConfigIdentifier"], + source_endpoint_arn=MOCK_REPLICATION_CONFIG["SourceEndpointArn"], + target_endpoint_arn=MOCK_REPLICATION_CONFIG["TargetEndpointArn"], + compute_config=MOCK_REPLICATION_CONFIG["ComputeConfig"], + replication_type=MOCK_REPLICATION_CONFIG["ReplicationType"], + table_mappings=MOCK_REPLICATION_CONFIG["TableMappings"], + ) + + assert resp == MOCK_REPLICATION_CONFIG_RESP["ReplicationConfig"]["ReplicationConfigArn"] + + @mock.patch.object(DmsHook, "conn") + def test_describe_replications(self, mock_conn): + mock_conn.describe_replication_tasks.return_value = MOCK_DESCRIBE_REPLICATIONS_RESP + resp = self.dms.describe_replication_tasks() + assert len(resp) == 2 + + @mock.patch.object(DmsHook, "conn") + def test_describe_replications_filter(self, mock_conn): + filter = [ + { + "Name": "replication-task-id", + "Values": MOCK_DESCRIBE_REPLICATIONS_RESP["Replications"][0]["ReplicationConfigArn"], + } + ] + self.dms.describe_replication_tasks(filters=filter) + mock_conn.describe_replication_tasks.assert_called_with(filters=filter) + + @mock.patch.object(DmsHook, "conn") + def test_start_replication_args(self, mock_conn): + self.dms.start_replication( + replication_config_arn=MOCK_TASK_ARN, + start_replication_type="cdc", + ) + mock_conn.start_replication.assert_called_with( + ReplicationConfigArn=MOCK_TASK_ARN, + StartReplicationType="cdc", + ) + + @mock.patch.object(DmsHook, "conn") + def test_start_replication_kwargs(self, mock_conn): + self.dms.start_replication( + replication_config_arn=MOCK_TASK_ARN, + start_replication_type="cdc", + cdc_start_time="2022-01-01T00:00:00Z", + cdc_start_pos=None, + cdc_stop_pos=None, + ) + mock_conn.start_replication.assert_called_with( + ReplicationConfigArn=MOCK_TASK_ARN, + StartReplicationType="cdc", + CdcStartTime=parser.parse("2022-01-01T00:00:00Z"), + ) diff --git a/tests/providers/amazon/aws/operators/test_dms.py b/tests/providers/amazon/aws/operators/test_dms.py index 2528edaef9e0a..acfe660293bd6 100644 --- a/tests/providers/amazon/aws/operators/test_dms.py +++ b/tests/providers/amazon/aws/operators/test_dms.py @@ -17,21 +17,33 @@ from __future__ import annotations import json +from typing import Any from unittest import mock import pendulum import pytest -from airflow import DAG -from airflow.models import DagRun, TaskInstance +from airflow.exceptions import AirflowException, TaskDeferred +from airflow.models import DAG, DagRun, TaskInstance +from airflow.models.variable import Variable from airflow.providers.amazon.aws.hooks.dms import DmsHook from airflow.providers.amazon.aws.operators.dms import ( + DmsCreateReplicationConfigOperator, DmsCreateTaskOperator, + DmsDeleteReplicationConfigOperator, DmsDeleteTaskOperator, + DmsDescribeReplicationConfigsOperator, + DmsDescribeReplicationsOperator, DmsDescribeTasksOperator, + DmsStartReplicationOperator, DmsStartTaskOperator, + DmsStopReplicationOperator, DmsStopTaskOperator, ) +from airflow.providers.amazon.aws.triggers.dms import ( + DmsReplicationDeprovisionedTrigger, + DmsReplicationTerminalStatusTrigger, +) from airflow.utils import timezone from airflow.utils.types import DagRunType from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields @@ -440,3 +452,569 @@ def test_template_fields(self): ) validate_template_fields(op) + + +class TestDmsDescribeReplicationConfigsOperator: + filter = [{"Name": "replication-type", "Values": ["cdc"]}] + + def test_init(self): + op = DmsDescribeReplicationConfigsOperator(task_id="test_task") + assert op.filter is None + + @pytest.mark.db_test + @mock.patch.object(DmsHook, "conn") + def test_template_fields_native(self, mock_conn, session): + execution_date = timezone.datetime(2020, 1, 1) + Variable.set("test_filter", self.filter, session=session) + + dag = DAG( + "test_dms", + schedule=None, + start_date=execution_date, + render_template_as_native_obj=True, + ) + op = DmsDescribeReplicationConfigsOperator( + task_id="test_task", filter="{{ var.value.test_filter }}", dag=dag + ) + + dag_run = DagRun( + dag_id=dag.dag_id, + execution_date=execution_date, + run_id="test", + run_type=DagRunType.MANUAL, + ) + ti = TaskInstance(task=op) + ti.dag_run = dag_run + session.add(ti) + session.commit() + context = ti.get_template_context(session) + ti.render_templates(context) + + assert op.filter == self.filter + + +class TestDmsCreateReplicationConfigOperator: + TASK_DATA = { + "ReplicationConfigIdentifier": "test-config", + "SourceEndpointArn": "arn:aws:dms:us-east-1:123456789012:endpoint:RZZK4EZW5UANC7Y3P4E776WHBE", + "TargetEndpointArn": "arn:aws:dms:us-east-1:123456789012:endpoint:GVBUJQXJZASXWHTWCLN2WNT57E", + "ComputeConfig": { + "MaxCapacityUnits": 2, + "MinCapacityUnits": 4, + }, + "ReplicationType": "full-load", + "TableMappings": json.dumps( + { + "TableMappings": [ + { + "Type": "Selection", + "RuleId": 123, + "RuleName": "test-rule", + "SourceSchema": "/", + "SourceTable": "/", + } + ] + } + ), + "ReplicationSettings": "string", + "SupplementalSettings": "string", + "ResourceIdentifier": "string", + } + + MOCK_REPLICATION_CONFIG_RESP: dict[str, Any] = { + "ReplicationConfig": { + "ReplicationConfigIdentifier": "test-config", + "ReplicationConfigArn": "arn:aws:dms:us-east-1:123456789012:replication-config/test-config", + "SourceEndpointArn": "arn:aws:dms:us-east-1:123456789012:endpoint:RZZK4EZW5UANC7Y3P4E776WHBE", + "TargetEndpointArn": "arn:aws:dms:us-east-1:123456789012:endpoint:GVBUJQXJZASXWHTWCLN2WNT57E", + "ReplicationType": "full-load", + } + } + + def test_init(self): + DmsCreateReplicationConfigOperator( + task_id="create_replication_config", + replication_config_id=self.TASK_DATA["ReplicationConfigIdentifier"], + source_endpoint_arn=self.TASK_DATA["SourceEndpointArn"], + target_endpoint_arn=self.TASK_DATA["TargetEndpointArn"], + replication_type=self.TASK_DATA["ReplicationType"], + table_mappings=self.TASK_DATA["TableMappings"], + compute_config=self.TASK_DATA["ComputeConfig"], + ) + + @mock.patch.object(DmsHook, "conn") + def test_operator(self, mock_hook): + mock_hook.create_replication_config.return_value = self.MOCK_REPLICATION_CONFIG_RESP + op = DmsCreateReplicationConfigOperator( + task_id="create_replication_config", + replication_config_id=self.TASK_DATA["ReplicationConfigIdentifier"], + source_endpoint_arn=self.TASK_DATA["SourceEndpointArn"], + target_endpoint_arn=self.TASK_DATA["TargetEndpointArn"], + replication_type=self.TASK_DATA["ReplicationType"], + table_mappings=self.TASK_DATA["TableMappings"], + compute_config=self.TASK_DATA["ComputeConfig"], + ) + resp = op.execute(None) + assert resp == self.MOCK_REPLICATION_CONFIG_RESP["ReplicationConfig"]["ReplicationConfigArn"] + + +class TestDmsDeleteReplicationConfigOperator: + TASK_DATA = { + "ReplicationConfigIdentifier": "test-config", + "ReplicationConfigArn": "arn:xxxxxx", + "SourceEndpointArn": "arn:aws:dms:us-east-1:123456789012:endpoint:RZZK4EZW5UANC7Y3P4E776WHBE", + "TargetEndpointArn": "arn:aws:dms:us-east-1:123456789012:endpoint:GVBUJQXJZASXWHTWCLN2WNT57E", + "ComputeConfig": { + "MaxCapacityUnits": 2, + "MinCapacityUnits": 4, + }, + "ReplicationType": "full-load", + "TableMappings": json.dumps( + { + "TableMappings": [ + { + "Type": "Selection", + "RuleId": 123, + "RuleName": "test-rule", + "SourceSchema": "/", + "SourceTable": "/", + } + ] + } + ), + "ReplicationSettings": "string", + "SupplementalSettings": "string", + "ResourceIdentifier": "string", + } + + def get_replication_status(self, status: str, deprovisioned: str = "deprovisioned"): + return [ + { + "Status": status, + "ReplicationArn": "XXXXXXXXXXXXXXXXXXXXXXXXX", + "ReplicationIdentifier": "test-config", + "SourceEndpointArn": "XXXXXXXXXXXXXXXXXXXXXXXXX", + "TargetEndpointArn": "XXXXXXXXXXXXXXXXXXXXXXXXX", + "ProvisionData": {"ProvisionState": deprovisioned, "ProvisionedCapacityUnits": 2}, + } + ] + + @mock.patch.object(DmsHook, "conn") + @mock.patch.object(DmsHook, "describe_replications") + @mock.patch.object(DmsDeleteReplicationConfigOperator, "handle_delete_wait") + @mock.patch.object(DmsHook, "get_waiter") + def test_happy_path(self, mock_waiter, mock_handle, mock_describe_replications, mock_conn): + # testing all good statuses and no waiting + mock_describe_replications.return_value = self.get_replication_status( + status="stopped", deprovisioned="deprovisioned" + ) + op = DmsDeleteReplicationConfigOperator( + task_id="delete_replication_config", + replication_config_arn=self.TASK_DATA["ReplicationConfigArn"], + deferrable=False, + wait_for_completion=False, + ) + op.execute({}) + + mock_conn.delete_replication_config.assert_called_once() + mock_waiter.assert_not_called() + mock_handle.assert_called_once() + + @mock.patch.object(DmsHook, "conn") + @mock.patch.object(DmsHook, "describe_replications") + def test_defer_not_ready(self, mock_describe, mock_conn): + mock_describe.return_value = self.get_replication_status("running") + + op = DmsDeleteReplicationConfigOperator( + task_id="delete_replication_config", + replication_config_arn=self.TASK_DATA["ReplicationConfigArn"], + deferrable=True, + ) + + with pytest.raises(TaskDeferred) as defer: + op.execute({}) + + assert isinstance(defer.value.trigger, DmsReplicationTerminalStatusTrigger) + + @mock.patch.object(DmsHook, "conn") + @mock.patch.object(DmsHook, "describe_replications") + @mock.patch.object(DmsHook, "get_waiter") + def test_wait_for_completion(self, mock_waiter, mock_describe_replications, mock_conn): + mock_describe_replications.return_value = self.get_replication_status( + status="failed", deprovisioned="deprovisioned" + ) + + op = DmsDeleteReplicationConfigOperator( + task_id="delete_replication_config", + replication_config_arn=self.TASK_DATA["ReplicationConfigArn"], + deferrable=False, + wait_for_completion=True, + ) + op.execute({}) + + mock_waiter.assert_called_with("replication_config_deleted") + mock_waiter.assert_called_once() + + @mock.patch.object(DmsHook, "conn") + @mock.patch.object(DmsHook, "describe_replications") + @mock.patch.object(DmsHook, "get_waiter") + def test_wait_for_completion_not_ready(self, mock_waiter, mock_describe_replications, mock_conn): + mock_describe_replications.return_value = self.get_replication_status( + status="failed", deprovisioned="xxx" + ) + + op = DmsDeleteReplicationConfigOperator( + task_id="delete_replication_config", + replication_config_arn=self.TASK_DATA["ReplicationConfigArn"], + deferrable=False, + wait_for_completion=True, + ) + op.execute({}) + + mock_waiter.assert_has_calls( + [ + mock.call("replication_deprovisioned"), + mock.call().wait( + Filters=[{"Name": "replication-config-arn", "Values": ["arn:xxxxxx"]}], + WaiterConfig={"Delay": 5, "MaxAttempts": 60}, + ), + mock.call("replication_config_deleted"), + mock.call().wait( + Filters=[{"Name": "replication-config-arn", "Values": ["arn:xxxxxx"]}], + WaiterConfig={"Delay": 5, "MaxAttempts": 60}, + ), + ] + ) + + @mock.patch.object(DmsHook, "conn") + @mock.patch.object(DmsHook, "describe_replications") + @mock.patch.object(DmsDeleteReplicationConfigOperator, "handle_delete_wait") + @mock.patch.object(DmsHook, "get_waiter") + def test_not_ready_state(self, mock_waiter, mock_handle, mock_describe, mock_conn): + mock_describe.return_value = self.get_replication_status("running") + + op = DmsDeleteReplicationConfigOperator( + task_id="delete_replication_config", + replication_config_arn=self.TASK_DATA["ReplicationConfigArn"], + deferrable=False, + wait_for_completion=False, + ) + op.execute({}) + + mock_waiter.assert_has_calls( + [ + mock.call("replication_terminal_status"), + mock.call().wait( + Filters=[{"Name": "replication-config-arn", "Values": ["arn:xxxxxx"]}], + WaiterConfig={"Delay": 5, "MaxAttempts": 60}, + ), + mock.call("replication_deprovisioned"), + mock.call().wait( + Filters=[{"Name": "replication-config-arn", "Values": ["arn:xxxxxx"]}], + WaiterConfig={"Delay": 5, "MaxAttempts": 60}, + ), + ] + ) + mock_handle.assert_called_once() + mock_conn.delete_replication_config.assert_called_once() + + @mock.patch.object(DmsHook, "conn") + @mock.patch.object(DmsHook, "describe_replications") + @mock.patch.object(DmsDeleteReplicationConfigOperator, "handle_delete_wait") + @mock.patch.object(DmsHook, "get_waiter") + def test_not_deprovisioned(self, mock_waiter, mock_handle, mock_describe, mock_conn): + mock_describe.return_value = self.get_replication_status("stopped", "deprovisioning") + op = DmsDeleteReplicationConfigOperator( + task_id="delete_replication_config", + replication_config_arn=self.TASK_DATA["ReplicationConfigArn"], + deferrable=False, + wait_for_completion=False, + ) + op.execute({}) + + mock_waiter.assert_has_calls( + [ + mock.call("replication_terminal_status"), + mock.call().wait( + Filters=[{"Name": "replication-config-arn", "Values": ["arn:xxxxxx"]}], + WaiterConfig={"Delay": 5, "MaxAttempts": 60}, + ), + mock.call("replication_deprovisioned"), + mock.call().wait( + Filters=[{"Name": "replication-config-arn", "Values": ["arn:xxxxxx"]}], + WaiterConfig={"Delay": 5, "MaxAttempts": 60}, + ), + ] + ) + mock_handle.assert_called_once() + + @mock.patch.object(DmsHook, "conn") + @mock.patch.object(DmsHook, "describe_replications") + @mock.patch.object(DmsHook, "get_waiter") + def test_config_not_found(self, mock_waiter, mock_describe, mock_conn): + mock_describe.return_value = [] + + op = DmsDeleteReplicationConfigOperator( + task_id="delete_replication_config", + replication_config_arn=self.TASK_DATA["ReplicationConfigArn"], + deferrable=False, + wait_for_completion=False, + ) + op.execute({}) + mock_waiter.assert_not_called() + mock_conn.delete_replication_config.assert_not_called() + + @mock.patch.object(DmsHook, "conn") + @mock.patch.object(DmsHook, "describe_replications") + @mock.patch.object(DmsDeleteReplicationConfigOperator, "defer") + @mock.patch.object(DmsHook, "get_waiter") + def test_handle_delete(self, mock_waiter, mock_defer, mock_describe, mock_conn): + mock_describe.return_value = self.get_replication_status( + status="stopped", deprovisioned="deprovisioned" + ) + + op = DmsDeleteReplicationConfigOperator( + task_id="delete_replication_config", + replication_config_arn=self.TASK_DATA["ReplicationConfigArn"], + deferrable=False, + wait_for_completion=True, + ) + op.execute({}) + mock_waiter.assert_called_with("replication_config_deleted") + + mock_waiter.reset_mock() + + op = DmsDeleteReplicationConfigOperator( + task_id="delete_replication_config", + replication_config_arn=self.TASK_DATA["ReplicationConfigArn"], + deferrable=False, + wait_for_completion=False, + ) + + op.execute({}) + mock_waiter.assert_not_called() + mock_defer.assert_not_called() + + @mock.patch.object(DmsHook, "conn") + @mock.patch.object(DmsHook, "describe_replications") + @mock.patch.object(DmsHook, "get_waiter") + def test_defer_not_deprovisioned(self, mock_waiter, mock_describe, mock_conn): + # not deprovisioned + mock_describe.return_value = self.get_replication_status("stopped", "deprovisioning") + op = DmsDeleteReplicationConfigOperator( + task_id="delete_replication_config", + replication_config_arn=self.TASK_DATA["ReplicationConfigArn"], + deferrable=True, + wait_for_completion=False, + ) + + with pytest.raises(TaskDeferred) as defer: + op.execute({}) + + assert isinstance(defer.value.trigger, DmsReplicationDeprovisionedTrigger) + + # not in terminal status + mock_describe.return_value = self.get_replication_status("running", "deprovisioning") + op = DmsDeleteReplicationConfigOperator( + task_id="delete_replication_config", + replication_config_arn=self.TASK_DATA["ReplicationConfigArn"], + deferrable=True, + wait_for_completion=False, + ) + + with pytest.raises(TaskDeferred) as defer: + op.execute({}) + + assert isinstance(defer.value.trigger, DmsReplicationTerminalStatusTrigger) + + +class TestDmsDescribeReplicationsOperator: + FILTER = [{"Name": "replication-type", "Values": ["cdc"]}] + + @mock.patch.object(DmsHook, "conn") + def test_filter(self, mock_conn): + mock_conn.describe_replications.return_value = [] + + op = DmsDescribeReplicationsOperator( + task_id="test_task", + filter=self.FILTER, + ) + + res = op.execute({}) + + mock_conn.describe_replications.assert_called_once_with(Filters=self.FILTER) + assert isinstance(res, list) + + @mock.patch.object(DmsHook, "conn") + def test_filter_none(self, mock_conn): + mock_conn.describe_replications.return_value = [] + + op = DmsDescribeReplicationsOperator( + task_id="test_task", + ) + + res = op.execute({}) + + mock_conn.describe_replications.assert_called_once_with(Filters=[]) + assert isinstance(res, list) + + +class TestDmsStartReplicationOperator: + def mock_describe_replication_response(self, status: str): + return [ + { + "ReplicationConfigIdentifier": "string", + "ReplicationConfigArn": "string", + "SourceEndpointArn": "string", + "TargetEndpointArn": "string", + "ReplicationType": "full-load", + "Status": status, + } + ] + + def mock_replication_response(self, status: str): + return { + "Replication": { + "ReplicationConfigIdentifier": "xxxx", + "ReplicationConfigArn": "xxxx", + "Status": status, + } + } + + def test_arg_validation(self): + with pytest.raises(AirflowException): + DmsStartReplicationOperator( + task_id="start_replication", + replication_config_arn="XXXXXXXXXXXXXXX", + replication_start_type="cdc", + cdc_start_pos=1, + cdc_start_time="2024-01-01 00:00:00", + ) + DmsStartReplicationOperator( + task_id="start_replication", + replication_config_arn="XXXXXXXXXXXXXXX", + replication_start_type="cdc", + cdc_start_pos=1, + ) + + DmsStartReplicationOperator( + task_id="start_replication", + replication_config_arn="XXXXXXXXXXXXXXX", + replication_start_type="cdc", + cdc_start_time="2024-01-01 00:00:00", + ) + + @mock.patch.object(DmsHook, "describe_replications") + @mock.patch.object(DmsHook, "start_replication") + def test_already_running(self, mock_replication, mock_describe): + mock_describe.return_value = self.mock_describe_replication_response("test") + + op = DmsStartReplicationOperator( + task_id="start_replication", + replication_config_arn="XXXXXXXXXXXXXXX", + replication_start_type="cdc", + cdc_start_pos=1, + wait_for_completion=False, + deferrable=False, + ) + + op.execute({}) + assert mock_replication.call_count == 0 + + mock_describe.return_value = self.mock_describe_replication_response("failed") + op.execute({}) + mock_replication.return_value = self.mock_replication_response("running") + assert mock_replication.call_count == 1 + + @mock.patch.object(DmsHook, "conn") + @mock.patch.object(DmsHook, "get_waiter") + @mock.patch.object(DmsHook, "describe_replications") + def test_wait_for_completion(self, mock_describe, mock_waiter, mock_conn): + mock_describe.return_value = self.mock_describe_replication_response("stopped") + op = DmsStartReplicationOperator( + task_id="start_replication", + replication_config_arn="XXXXXXXXXXXXXXX", + replication_start_type="cdc", + cdc_start_pos=1, + wait_for_completion=True, + deferrable=False, + ) + + op.execute({}) + mock_waiter.assert_called_with("replication_complete") + mock_waiter.assert_called_once() + + @mock.patch.object(DmsHook, "conn") + @mock.patch.object(DmsHook, "describe_replications") + def test_execute(self, mock_describe, mock_conn): + mock_describe.return_value = self.mock_describe_replication_response("stopped") + + op = DmsStartReplicationOperator( + task_id="start_replication", + replication_config_arn="XXXXXXXXXXXXXXX", + replication_start_type="cdc", + cdc_start_pos=1, + wait_for_completion=False, + deferrable=False, + ) + op.execute({}) + assert mock_conn.start_replication.call_count == 1 + + +class TestDmsStopReplicationOperator: + def mock_describe_replication_response(self, status: str): + return [ + { + "ReplicationConfigIdentifier": "string", + "ReplicationConfigArn": "string", + "SourceEndpointArn": "string", + "TargetEndpointArn": "string", + "ReplicationType": "full-load", + "Status": status, + } + ] + + @mock.patch.object(DmsHook, "describe_replications") + @mock.patch.object(DmsHook, "conn") + def test_already_stopped(self, mock_conn, mock_describe_replications): + mock_describe_replications.return_value = self.mock_describe_replication_response("stopped") + + op = DmsStopReplicationOperator( + task_id="stop_replication", + replication_config_arn="XXXXXXXXXXXXXXX", + wait_for_completion=False, + deferrable=False, + ) + op.execute({}) + + assert mock_conn.stop_replication.call_count == 0 + + @mock.patch.object(DmsHook, "stop_replication") + @mock.patch.object(DmsHook, "describe_replications") + def test_execute(self, mock_describe_replications, mock_stop): + mock_describe_replications.return_value = self.mock_describe_replication_response("started") + + op = DmsStopReplicationOperator( + task_id="stop_replication", + replication_config_arn="XXXXXXXXXXXXXXX", + wait_for_completion=False, + deferrable=False, + ) + op.execute({}) + assert mock_stop.call_count == 1 + + @mock.patch.object(DmsHook, "conn") + @mock.patch.object(DmsHook, "describe_replications") + @mock.patch.object(DmsHook, "get_waiter") + def test_wait_for_completion(self, mock_get_waiter, mock_describe_replications, mock_conn): + mock_describe_replications.return_value = self.mock_describe_replication_response("started") + op = DmsStopReplicationOperator( + task_id="stop_replication", + replication_config_arn="XXXXXXXXXXXXXXX", + wait_for_completion=True, + deferrable=False, + ) + + op.execute({}) + mock_get_waiter.assert_called_with("replication_stopped") + mock_get_waiter.assert_called_once() diff --git a/tests/providers/amazon/aws/triggers/test_dms.py b/tests/providers/amazon/aws/triggers/test_dms.py new file mode 100644 index 0000000000000..4835b3a2336ff --- /dev/null +++ b/tests/providers/amazon/aws/triggers/test_dms.py @@ -0,0 +1,187 @@ +# 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 + +from unittest import mock +from unittest.mock import AsyncMock + +import pytest + +from airflow.providers.amazon.aws.hooks.dms import DmsHook +from airflow.providers.amazon.aws.triggers.dms import ( + DmsReplicationCompleteTrigger, + DmsReplicationConfigDeletedTrigger, + DmsReplicationDeprovisionedTrigger, + DmsReplicationStoppedTrigger, + DmsReplicationTerminalStatusTrigger, +) +from airflow.triggers.base import TriggerEvent +from tests.providers.amazon.aws.utils.test_waiter import assert_expected_waiter_type + +BASE_TRIGGER_CLASSPATH = "airflow.providers.amazon.aws.triggers.dms." + + +class TestBaseDmsTrigger: + EXPECTED_WAITER_NAME: str | None = None + + def test_setup(self): + if self.__class__.__name__ != "TestBaseDmsTrigger": + assert isinstance(self.EXPECTED_WAITER_NAME, str) + + +class TestDmsReplicationCompleteTrigger(TestBaseDmsTrigger): + EXPECTED_WAITER_NAME = "replication_complete" + REPLICATION_CONFIG_ARN = "arn:aws:dms:region:account:config" + + def test_serialization(self): + trigger = DmsReplicationCompleteTrigger(replication_config_arn=self.REPLICATION_CONFIG_ARN) + + classpath, kwargs = trigger.serialize() + assert classpath == BASE_TRIGGER_CLASSPATH + "DmsReplicationCompleteTrigger" + + assert kwargs.get("replication_config_arn") == self.REPLICATION_CONFIG_ARN + + @pytest.mark.asyncio + @mock.patch.object(DmsHook, "get_waiter") + @mock.patch.object(DmsHook, "get_conn") + async def test_complete(self, mock_async_conn, mock_get_waiter): + mock_async_conn.__aenter__.return_value = mock.MagicMock() + mock_get_waiter().wait = AsyncMock() + trigger = DmsReplicationCompleteTrigger(replication_config_arn=self.REPLICATION_CONFIG_ARN) + generator = trigger.run() + response = await generator.asend(None) + assert response == TriggerEvent( + {"status": "success", "replication_config_arn": self.REPLICATION_CONFIG_ARN} + ) + mock_get_waiter().wait.assert_called_once() + + +class TestDmsReplicationTerminalStatusTrigger(TestBaseDmsTrigger): + EXPECTED_WAITER_NAME = "replication_terminal_status" + REPLICATION_CONFIG_ARN = "arn:aws:dms:region:account:config" + + def test_serialization(self): + trigger = DmsReplicationTerminalStatusTrigger(replication_config_arn=self.REPLICATION_CONFIG_ARN) + + classpath, kwargs = trigger.serialize() + assert classpath == BASE_TRIGGER_CLASSPATH + "DmsReplicationTerminalStatusTrigger" + + assert kwargs.get("replication_config_arn") == self.REPLICATION_CONFIG_ARN + + @pytest.mark.asyncio + @mock.patch.object(DmsHook, "get_waiter") + @mock.patch.object(DmsHook, "get_conn") + async def test_complete(self, mock_async_conn, mock_get_waiter): + mock_async_conn.__aenter__.return_value = mock.MagicMock() + mock_get_waiter().wait = AsyncMock() + trigger = DmsReplicationTerminalStatusTrigger(replication_config_arn=self.REPLICATION_CONFIG_ARN) + generator = trigger.run() + response = await generator.asend(None) + assert response == TriggerEvent( + {"status": "success", "replication_config_arn": self.REPLICATION_CONFIG_ARN} + ) + assert_expected_waiter_type(mock_get_waiter, self.EXPECTED_WAITER_NAME) + + mock_get_waiter().wait.assert_called_once() + + +class TestDmsReplicationConfigDeletedTrigger(TestBaseDmsTrigger): + EXPECTED_WAITER_NAME = "replication_config_deleted" + REPLICATION_CONFIG_ARN = "arn:aws:dms:region:account:config" + + def test_serialization(self): + trigger = DmsReplicationConfigDeletedTrigger(replication_config_arn=self.REPLICATION_CONFIG_ARN) + + classpath, kwargs = trigger.serialize() + assert classpath == BASE_TRIGGER_CLASSPATH + "DmsReplicationConfigDeletedTrigger" + + assert kwargs.get("replication_config_arn") == self.REPLICATION_CONFIG_ARN + + @pytest.mark.asyncio + @mock.patch.object(DmsHook, "get_waiter") + @mock.patch.object(DmsHook, "get_conn") + async def test_complete(self, mock_async_conn, mock_get_waiter): + mock_async_conn.__aenter__.return_value = mock.MagicMock() + mock_get_waiter().wait = AsyncMock() + trigger = DmsReplicationConfigDeletedTrigger(replication_config_arn=self.REPLICATION_CONFIG_ARN) + generator = trigger.run() + response = await generator.asend(None) + assert response == TriggerEvent( + {"status": "success", "replication_config_arn": self.REPLICATION_CONFIG_ARN} + ) + assert_expected_waiter_type(mock_get_waiter, self.EXPECTED_WAITER_NAME) + + mock_get_waiter().wait.assert_called_once() + + +class TestDmsReplicationStoppedTrigger(TestBaseDmsTrigger): + EXPECTED_WAITER_NAME = "replication_stopped" + REPLICATION_CONFIG_ARN = "arn:aws:dms:region:account:config" + + def test_serialization(self): + trigger = DmsReplicationStoppedTrigger(replication_config_arn=self.REPLICATION_CONFIG_ARN) + + classpath, kwargs = trigger.serialize() + assert classpath == BASE_TRIGGER_CLASSPATH + "DmsReplicationStoppedTrigger" + + """ assert kwargs.get("Filters") == [ + {"Name": "replication-config-arn", "Values": ["arn:aws:dms:region:account:config"]} + ] """ + assert kwargs.get("replication_config_arn") == self.REPLICATION_CONFIG_ARN + + @pytest.mark.asyncio + @mock.patch.object(DmsHook, "get_waiter") + @mock.patch.object(DmsHook, "get_conn") + async def test_complete(self, mock_async_conn, mock_get_waiter): + mock_async_conn.__aenter__.return_value = mock.MagicMock() + mock_get_waiter().wait = AsyncMock() + trigger = DmsReplicationStoppedTrigger(replication_config_arn=self.REPLICATION_CONFIG_ARN) + generator = trigger.run() + response = await generator.asend(None) + assert response == TriggerEvent( + {"status": "success", "replication_config_arn": self.REPLICATION_CONFIG_ARN} + ) + assert_expected_waiter_type(mock_get_waiter, self.EXPECTED_WAITER_NAME) + mock_get_waiter().wait.assert_called_once() + + +class TestDmsReplicationDeprovisionedTrigger(TestBaseDmsTrigger): + EXPECTED_WAITER_NAME = "replication_deprovisioned" + REPLICATION_CONFIG_ARN = "arn:aws:dms:region:account:config" + + def test_serialization(self): + trigger = DmsReplicationDeprovisionedTrigger(replication_config_arn=self.REPLICATION_CONFIG_ARN) + + classpath, kwargs = trigger.serialize() + assert classpath == BASE_TRIGGER_CLASSPATH + "DmsReplicationDeprovisionedTrigger" + + assert kwargs.get("replication_config_arn") == self.REPLICATION_CONFIG_ARN + + @pytest.mark.asyncio + @mock.patch.object(DmsHook, "get_waiter") + @mock.patch.object(DmsHook, "get_conn") + async def test_complete(self, mock_async_conn, mock_get_waiter): + mock_async_conn.__aenter__.return_value = mock.MagicMock() + mock_get_waiter().wait = AsyncMock() + trigger = DmsReplicationDeprovisionedTrigger(replication_config_arn=self.REPLICATION_CONFIG_ARN) + generator = trigger.run() + response = await generator.asend(None) + assert response == TriggerEvent( + {"status": "success", "replication_config_arn": self.REPLICATION_CONFIG_ARN} + ) + assert_expected_waiter_type(mock_get_waiter, self.EXPECTED_WAITER_NAME) + mock_get_waiter().wait.assert_called_once() diff --git a/tests/providers/amazon/aws/waiters/test_dms.py b/tests/providers/amazon/aws/waiters/test_dms.py new file mode 100644 index 0000000000000..86d36ee9c3e0c --- /dev/null +++ b/tests/providers/amazon/aws/waiters/test_dms.py @@ -0,0 +1,139 @@ +# 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 + +from unittest import mock + +import boto3 +import pytest + +from airflow.providers.amazon.aws.hooks.dms import DmsHook + + +class TestCustomDmsWaiters: + """Test waiters from ``amazon/aws/waiters/dms.json```.""" + + @pytest.fixture(autouse=True) + def setup_test_cases(self, monkeypatch): + self.client = boto3.client("dms", region_name="us-east-1") + monkeypatch.setattr(DmsHook, "conn", self.client) + + def test_service_waiters(self): + hook_waiters = DmsHook(aws_conn_id=None).list_waiters() + assert "replication_terminal_status" in hook_waiters + assert "replication_config_deleted" in hook_waiters + assert "replication_stopped" in hook_waiters + assert "replication_complete" in hook_waiters + + @pytest.fixture + def mock_describe_replication(self): + with mock.patch.object(self.client, "describe_replications") as m: + yield m + + @pytest.fixture + def mock_describe_replication_configs(self): + with mock.patch.object(self.client, "describe_replication_configs") as m: + yield m + + def test_wait_for_replication_terminal_status(self, mock_describe_replication): + mock_describe_replication.return_value = { + "Replications": [ + { + "ReplicationConfigArn": "XXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXX", + "Status": "failed", + } + ] + } + + hook = DmsHook(aws_conn_id=None) + waiter = hook.get_waiter("replication_terminal_status") + waiter.wait( + Filters=[ + { + "Name": "replication-instance-arn", + "Values": ["XXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXX"], + } + ] + ) + + mock_describe_replication.assert_called_once_with( + Filters=[ + { + "Name": "replication-instance-arn", + "Values": ["XXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXX"], + } + ] + ) + + def test_wait_for_replication_config_delete(self, mock_describe_replication_configs): + mock_describe_replication_configs.side_effect = [ + {"Replications": [{"ReplicationArn": "MyArn"}]}, + {"Error": {"Code": "ResourceNotFoundFault"}}, + ] + + hook = DmsHook(aws_conn_id=None) + waiter = hook.get_waiter("replication_config_deleted") + waiter.wait( + Filters=[ + { + "Name": "replication-config-arn", + "Values": ["XXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXX"], + } + ], + WaiterConfig={"Delay": 0.01, "MaxAttempts": 3}, + ) + + mock_describe_replication_configs.assert_called_with( + Filters=[ + { + "Name": "replication-config-arn", + "Values": ["XXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXX"], + } + ] + ) + + def test_wait_for_replication_stopped(self, mock_describe_replication): + mock_describe_replication.return_value = { + "Replications": [ + { + "ReplicationConfigArn": "XXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXX", + "Status": "stopped", + } + ] + } + + hook = DmsHook(aws_conn_id=None) + waiter = hook.get_waiter("replication_stopped") + waiter.wait( + Filters=[ + { + "Name": "replication_config_arn", + "Values": ["XXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXX"], + } + ], + WaiterConfig={"Delay": 0.01, "MaxAttempts": 3}, + ) + + mock_describe_replication.assert_called_once_with( + Filters=[ + { + "Name": "replication_config_arn", + "Values": ["XXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXX"], + } + ] + ) diff --git a/tests/system/providers/amazon/aws/example_dms_serverless.py b/tests/system/providers/amazon/aws/example_dms_serverless.py new file mode 100644 index 0000000000000..01796bf6a6126 --- /dev/null +++ b/tests/system/providers/amazon/aws/example_dms_serverless.py @@ -0,0 +1,456 @@ +# +# 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. +""" +Note: DMS requires you to configure specific IAM roles/permissions. For more information, see +https://docs.aws.amazon.com/dms/latest/userguide/security-iam.html#CHAP_Security.APIRole +""" + +from __future__ import annotations + +import json +from datetime import datetime + +import boto3 +from sqlalchemy import Column, MetaData, String, Table, create_engine + +from airflow.decorators import task +from airflow.models.baseoperator import chain +from airflow.models.dag import DAG +from airflow.providers.amazon.aws.operators.dms import ( + DmsCreateReplicationConfigOperator, + DmsDeleteReplicationConfigOperator, + DmsDescribeReplicationConfigsOperator, + DmsDescribeReplicationsOperator, + DmsStartReplicationOperator, +) +from airflow.providers.amazon.aws.operators.rds import ( + RdsCreateDbInstanceOperator, + RdsDeleteDbInstanceOperator, +) +from airflow.providers.amazon.aws.operators.s3 import S3CreateBucketOperator, S3DeleteBucketOperator +from airflow.utils.trigger_rule import TriggerRule +from tests.system.providers.amazon.aws.utils import ENV_ID_KEY, SystemTestContextBuilder +from tests.system.providers.amazon.aws.utils.ec2 import get_default_vpc_id + +""" +This example demonstrates how to use the DMS operators to create a serverless replication task to replicate data +from a PostgreSQL database to Amazon S3. + +The IAM role used for the replication must have the permissions defined in the [Amazon S3 target](https://docs.aws.amazon.com/dms/latest/userguide/CHAP_Target.S3.html#CHAP_Target.S3.Prerequisites) +documentation. +""" + +DAG_ID = "example_dms_serverless" +ROLE_ARN_KEY = "ROLE_ARN" + +sys_test_context_task = SystemTestContextBuilder().add_variable(ROLE_ARN_KEY).build() + +# Config values for setting up the "Source" database. +CA_CERT_ID = "rds-ca-rsa2048-g1" +RDS_ENGINE = "postgres" +RDS_PROTOCOL = "postgresql" +RDS_USERNAME = "username" +# NEVER store your production password in plaintext in a DAG like this. +# Use Airflow Secrets or a secret manager for this in production. +RDS_PASSWORD = "rds_password" +TABLE_HEADERS = ["apache_project", "release_year"] +SAMPLE_DATA = [ + ("Airflow", "2015"), + ("OpenOffice", "2012"), + ("Subversion", "2000"), + ("NiFi", "2006"), +] +SG_IP_PERMISSION = { + "FromPort": 5432, + "IpProtocol": "All", + "IpRanges": [{"CidrIp": "0.0.0.0/0"}], +} + + +def _get_rds_instance_endpoint(instance_name: str): + print("Retrieving RDS instance endpoint.") + rds_client = boto3.client("rds") + + response = rds_client.describe_db_instances(DBInstanceIdentifier=instance_name) + rds_instance_endpoint = response["DBInstances"][0]["Endpoint"] + return rds_instance_endpoint + + +@task +def create_security_group(security_group_name: str, vpc_id: str): + client = boto3.client("ec2") + security_group = client.create_security_group( + GroupName=security_group_name, + Description="Created for DMS system test", + VpcId=vpc_id, + ) + client.get_waiter("security_group_exists").wait( + GroupIds=[security_group["GroupId"]], + ) + client.authorize_security_group_ingress( + GroupId=security_group["GroupId"], + IpPermissions=[SG_IP_PERMISSION], + ) + + return security_group["GroupId"] + + +@task +def create_sample_table(instance_name: str, db_name: str, table_name: str): + print("Creating sample table.") + + 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}/{db_name}" + engine = create_engine(rds_url) + + table = Table( + table_name, + MetaData(engine), + Column(TABLE_HEADERS[0], String, primary_key=True), + Column(TABLE_HEADERS[1], String), + ) + + with engine.connect() as connection: + # Create the Table. + table.create() + load_data = table.insert().values(SAMPLE_DATA) + connection.execute(load_data) + + # Read the data back to verify everything is working. + connection.execute(table.select()) + + +@task(trigger_rule=TriggerRule.ALL_SUCCESS) +def create_vpc_endpoints(vpc_id: str): + print("Creating VPC endpoints in vpc: %s", vpc_id) + client = boto3.client("ec2") + session = boto3.session.Session() + region = session.region_name + route_tbls = client.describe_route_tables(Filters=[{"Name": "vpc-id", "Values": [vpc_id]}]) + endpoints = client.create_vpc_endpoint( + VpcId=vpc_id, + ServiceName=f"com.amazonaws.{region}.s3", + VpcEndpointType="Gateway", + RouteTableIds=[tbl["RouteTableId"] for tbl in route_tbls["RouteTables"]], + ) + + return endpoints.get("VpcEndpoint", {}).get("VpcEndpointId") + + +@task(trigger_rule=TriggerRule.ALL_DONE) +def delete_vpc_endpoints(endpoint_ids: list[str]): + if len(endpoint_ids) == 0: + print("No VPC endpoints to delete.") + return + + print("Deleting VPC endpoints.") + client = boto3.client("ec2") + + client.delete_vpc_endpoints(VpcEndpointIds=endpoint_ids, DryRun=False) + + print("Deleted endpoints: %s", endpoint_ids) + + +@task(multiple_outputs=True) +def create_dms_assets( + db_name: str, + 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") + rds_instance_endpoint = _get_rds_instance_endpoint(instance_name) + + print("Creating DMS source endpoint.") + source_endpoint_arn = dms_client.create_endpoint( + EndpointIdentifier=source_endpoint_identifier, + EndpointType="source", + EngineName=RDS_ENGINE, + Username=RDS_USERNAME, + Password=RDS_PASSWORD, + ServerName=rds_instance_endpoint["Address"], + Port=rds_instance_endpoint["Port"], + DatabaseName=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), + }, + )["Endpoint"]["EndpointArn"] + + return { + "source_endpoint_arn": source_endpoint_arn, + "target_endpoint_arn": target_endpoint_arn, + } + + +@task(trigger_rule=TriggerRule.ALL_DONE) +def delete_dms_assets( + source_endpoint_arn: str, + target_endpoint_arn: str, + source_endpoint_identifier: str, + target_endpoint_identifier: str, +): + dms_client = boto3.client("dms") + + print("Deleting DMS assets.") + + print(source_endpoint_arn) + print(target_endpoint_arn) + + try: + dms_client.delete_endpoint(EndpointArn=source_endpoint_arn) + dms_client.delete_endpoint(EndpointArn=target_endpoint_arn) + except Exception as ex: + print("Exception while cleaning up endpoints:%s", ex) + + print("Awaiting DMS assets tear-down.") + + dms_client.get_waiter("endpoint_deleted").wait( + Filters=[ + { + "Name": "endpoint-id", + "Values": [source_endpoint_identifier, target_endpoint_identifier], + } + ] + ) + + +@task(trigger_rule=TriggerRule.ALL_DONE) +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) + + +# setup +# source: aurora serverless +# dest: S3 +# S3 + +with DAG( + dag_id=DAG_ID, + schedule="@once", + start_date=datetime(2021, 1, 1), + tags=["example"], + catchup=False, +) 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_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" + replication_id = f"{env_id}-replication-id" + + 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) + + create_db_instance = RdsCreateDbInstanceOperator( + task_id="create_db_instance", + db_instance_identifier=rds_instance_name, + db_instance_class="db.t3.micro", + engine=RDS_ENGINE, + rds_kwargs={ + "DBName": rds_db_name, + "AllocatedStorage": 20, + "MasterUsername": RDS_USERNAME, + "MasterUserPassword": RDS_PASSWORD, + "PubliclyAccessible": True, + "VpcSecurityGroupIds": [ + create_sg, + ], + }, + ) + + # 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": [ + { + "rule-type": "selection", + "rule-id": "1", + "rule-name": "1", + "object-locator": { + "schema-name": "public", + "table-name": rds_table_name, + }, + "rule-action": "include", + } + ] + } + + create_assets = create_dms_assets( + db_name=rds_db_name, + instance_name=rds_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, + ) + + create_replication_config = DmsCreateReplicationConfigOperator( + task_id="create_replication_config", + replication_config_id=replication_id, + source_endpoint_arn=create_assets["source_endpoint_arn"], + target_endpoint_arn=create_assets["target_endpoint_arn"], + compute_config={ + "MaxCapacityUnits": 4, + "MinCapacityUnits": 1, + "MultiAZ": False, + "ReplicationSubnetGroupId": "default", + }, + replication_type="full-load", + table_mappings=json.dumps(table_mappings), + trigger_rule=TriggerRule.ALL_SUCCESS, + ) + + describe_replication_configs = DmsDescribeReplicationConfigsOperator( + task_id="describe_replication_configs", + trigger_rule=TriggerRule.ALL_SUCCESS, + ) + + describe_replications = DmsDescribeReplicationsOperator( + task_id="describe_replications", + trigger_rule=TriggerRule.ALL_SUCCESS, + ) + + replicate = DmsStartReplicationOperator( + task_id="replicate", + replication_config_arn="{{ task_instance.xcom_pull(task_ids='create_replication_config', key='return_value') }}", + replication_start_type="start-replication", + wait_for_completion=True, + waiter_delay=60, + waiter_max_attempts=200, + trigger_rule=TriggerRule.ALL_SUCCESS, + deferrable=False, + ) + + delete_replication_config = DmsDeleteReplicationConfigOperator( + task_id="delete_replication_config", + wait_for_completion=True, + waiter_delay=60, + waiter_max_attempts=200, + deferrable=False, + replication_config_arn="{{ task_instance.xcom_pull(task_ids='create_replication_config', key='return_value') }}", + trigger_rule=TriggerRule.ALL_DONE, + ) + + delete_assets = delete_dms_assets( + source_endpoint_arn=create_assets["source_endpoint_arn"], + target_endpoint_arn=create_assets["target_endpoint_arn"], + source_endpoint_identifier=source_endpoint_identifier, + target_endpoint_identifier=target_endpoint_identifier, + ) + + delete_db_instance = RdsDeleteDbInstanceOperator( + task_id="delete_db_instance", + db_instance_identifier=rds_instance_name, + rds_kwargs={ + "SkipFinalSnapshot": True, + }, + 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, + ) + + chain( + # TEST SETUP + create_s3_bucket, + get_vpc_id, + create_sg, + create_db_instance, + create_sample_table(rds_instance_name, rds_db_name, rds_table_name), + create_vpc_endpoints( + vpc_id="{{ task_instance.xcom_pull(task_ids='get_default_vpc_id',key='return_value')}}" + ), + create_assets, + # TEST BODY + create_replication_config, + describe_replication_configs, + replicate, + describe_replications, + delete_replication_config, + # TEST TEARDOWN + delete_vpc_endpoints( + endpoint_ids=[ + "{{ task_instance.xcom_pull(task_ids='create_vpc_endpoints', key='return_value') }}" + ] + ), + delete_assets, + delete_db_instance, + delete_security_group(create_sg, security_group_name), + delete_s3_bucket, + ) + + from tests.system.utils.watcher import watcher + + # This test needs watcher in order to properly mark success/failure + # when "tearDown" task with trigger rule is part of the DAG + list(dag.tasks) >> watcher() + +from tests.system.utils import get_test_run # noqa: E402 + +# Needed to run the example DAG with pytest (see: tests/system/README.md#run_via_pytest) +test_run = get_test_run(dag) From d88ff816d0dc0895ab8cd678adc05b7fe6e08d52 Mon Sep 17 00:00:00 2001 From: Mike Ellis Date: Wed, 13 Nov 2024 08:45:31 -0500 Subject: [PATCH 002/802] doc updates --- airflow/providers/amazon/aws/triggers/dms.py | 6 +- .../operators/dms.rst | 63 +++++++++++++++++++ .../amazon/aws/example_dms_serverless.py | 10 +++ 3 files changed, 74 insertions(+), 5 deletions(-) diff --git a/airflow/providers/amazon/aws/triggers/dms.py b/airflow/providers/amazon/aws/triggers/dms.py index 3e0e563f2f4f7..fa18cd009fad2 100644 --- a/airflow/providers/amazon/aws/triggers/dms.py +++ b/airflow/providers/amazon/aws/triggers/dms.py @@ -84,9 +84,6 @@ def __init__( ) -> None: super().__init__( serialized_fields={"replication_config_arn": replication_config_arn}, - # serialized_fields={ - # "Filters": [{"Name": "replication-config-arn", "Values": [replication_config_arn]}] - # }, waiter_name="replication_config_deleted", waiter_delay=waiter_delay, waiter_args={"Filters": [{"Name": "replication-config-arn", "Values": [replication_config_arn]}]}, @@ -126,7 +123,6 @@ def __init__( ) -> None: super().__init__( serialized_fields={"replication_config_arn": replication_config_arn}, - # "Filters": [{"Name": "replication-config-arn", "Values": [replication_config_arn]}]}, waiter_name="replication_complete", waiter_delay=waiter_delay, waiter_args={"Filters": [{"Name": "replication-config-arn", "Values": [replication_config_arn]}]}, @@ -188,7 +184,7 @@ def hook(self) -> AwsGenericHook: class DmsReplicationDeprovisionedTrigger(AwsBaseWaiterTrigger): """ - Trigger when an AWS DMS Serverless replication is deprovisioned. + Trigger when an AWS DMS Serverless replication is de-provisioned. :param replication_config_arn: The ARN of the replication config. :param waiter_delay: The amount of time in seconds to wait between attempts. diff --git a/docs/apache-airflow-providers-amazon/operators/dms.rst b/docs/apache-airflow-providers-amazon/operators/dms.rst index 2c30e3ca6ec88..56e7c85ce5077 100644 --- a/docs/apache-airflow-providers-amazon/operators/dms.rst +++ b/docs/apache-airflow-providers-amazon/operators/dms.rst @@ -114,6 +114,69 @@ To delete a replication task you can use :start-after: [START howto_operator_dms_delete_task] :end-before: [END howto_operator_dms_delete_task] + +Create a serverless replication config +====================================== + +To create a serverless replication config use +:class:`~airflow.providers.amazon.aws.operators.dms.DmsCreateReplicationConfigOperator`. + +.. exampleinclude:: /../../tests/system/providers/amazon/aws/example_dms_serverless.py + :language: python + :dedent: 4 + :start-after: [START howto_operator_dms_create_replication_config] + :end-before: [END howto_operator_dms_create_replication_config] + +Describe a serverless replication config +======================================== + +To describe a serverless replication config use +:class:`~airflow.providers.amazon.aws.operators.dms.DmsDescribeReplicationConfigsOperator`. + +.. exampleinclude:: /../../tests/system/providers/amazon/aws/example_dms_serverless.py + :language: python + :dedent: 4 + :start-after: [START howto_operator_dms_describe_replication_config] + :end-before: [END howto_operator_dms_describe_replication_config] + +Start a serverless replication +============================== + +To start a serverless replication use +:class:`~airflow.providers.amazon.aws.operators.dms.DmsStartReplicationOperator`. + +.. exampleinclude:: /../../tests/system/providers/amazon/aws/example_dms_serverless.py + :language: python + :dedent: 4 + :start-after: [START howto_operator_dms_serverless_start_replication] + :end-before: [END howto_operator_dms_serverless_start_replication] + +Get the status of a serverless replication +========================================== + +To get the status of a serverless replication use +:class:`~airflow.providers.amazon.aws.operators.dms.DmsDescribeReplicationsOperator`. + +.. exampleinclude:: /../../tests/system/providers/amazon/aws/example_dms_serverless.py + :language: python + :dedent: 4 + :start-after: [START howto_operator_dms_serverless_describe_replication] + :end-before: [END howto_operator_dms_serverless_describe_replication] + +Delete a serverless replication configuration +============================================= + +To delete a serverless replication config use +:class:`~airflow.providers.amazon.aws.operators.dms.DmsDescribeReplicationsOperator`. + +.. exampleinclude:: /../../tests/system/providers/amazon/aws/example_dms_serverless.py + :language: python + :dedent: 4 + :start-after: [START howto_operator_dms_serverless_delete_replication_config] + :end-before: [END howto_operator_dms_serverless_delete_replication_config] + + + Sensors ------- diff --git a/tests/system/providers/amazon/aws/example_dms_serverless.py b/tests/system/providers/amazon/aws/example_dms_serverless.py index 01796bf6a6126..9404d8aaffd38 100644 --- a/tests/system/providers/amazon/aws/example_dms_serverless.py +++ b/tests/system/providers/amazon/aws/example_dms_serverless.py @@ -345,6 +345,7 @@ def delete_security_group(security_group_id: str, security_group_name: str): table_definition=table_definition, ) + # [START howto_operator_dms_create_replication_config] create_replication_config = DmsCreateReplicationConfigOperator( task_id="create_replication_config", replication_config_id=replication_id, @@ -360,17 +361,23 @@ def delete_security_group(security_group_id: str, security_group_name: str): table_mappings=json.dumps(table_mappings), trigger_rule=TriggerRule.ALL_SUCCESS, ) + # [END howto_operator_dms_create_replication_config] + # [START howto_operator_dms_describe_replication_config] describe_replication_configs = DmsDescribeReplicationConfigsOperator( task_id="describe_replication_configs", trigger_rule=TriggerRule.ALL_SUCCESS, ) + # [END howto_operator_dms_describe_replication_config] + # [START howto_operator_dms_serverless_describe_replication] describe_replications = DmsDescribeReplicationsOperator( task_id="describe_replications", trigger_rule=TriggerRule.ALL_SUCCESS, ) + # [END howto_operator_dms_serverless_describe_replication] + # [START howto_operator_dms_serverless_start_replication] replicate = DmsStartReplicationOperator( task_id="replicate", replication_config_arn="{{ task_instance.xcom_pull(task_ids='create_replication_config', key='return_value') }}", @@ -381,7 +388,9 @@ def delete_security_group(security_group_id: str, security_group_name: str): trigger_rule=TriggerRule.ALL_SUCCESS, deferrable=False, ) + # [END howto_operator_dms_serverless_start_replication] + # [START howto_operator_dms_serverless_delete_replication_config] delete_replication_config = DmsDeleteReplicationConfigOperator( task_id="delete_replication_config", wait_for_completion=True, @@ -391,6 +400,7 @@ def delete_security_group(security_group_id: str, security_group_name: str): replication_config_arn="{{ task_instance.xcom_pull(task_ids='create_replication_config', key='return_value') }}", trigger_rule=TriggerRule.ALL_DONE, ) + # [END howto_operator_dms_serverless_delete_replication_config] delete_assets = delete_dms_assets( source_endpoint_arn=create_assets["source_endpoint_arn"], From 60f5009bfca9b127691738a66db12db2401aeee1 Mon Sep 17 00:00:00 2001 From: Elad Kalif <45845474+eladkal@users.noreply.github.com> Date: Tue, 24 Sep 2024 17:07:35 +0300 Subject: [PATCH 003/802] Update providers metadata 2024-09-24 (#42445) --- generated/provider_metadata.json | 376 +++++++++++++++++++------------ 1 file changed, 237 insertions(+), 139 deletions(-) diff --git a/generated/provider_metadata.json b/generated/provider_metadata.json index 4ca06608c7cab..a73e3da9f6fce 100644 --- a/generated/provider_metadata.json +++ b/generated/provider_metadata.json @@ -85,8 +85,12 @@ "date_released": "2024-05-30T06:38:15Z" }, "3.9.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "4.0.0": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "alibaba": { @@ -179,8 +183,12 @@ "date_released": "2024-05-30T06:38:15Z" }, "2.9.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "2.9.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "amazon": { @@ -433,8 +441,12 @@ "date_released": "2024-08-06T20:34:43Z" }, "8.28.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "8.29.0": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "apache.beam": { @@ -567,7 +579,7 @@ "date_released": "2024-08-06T20:34:43Z" }, "5.8.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -649,7 +661,7 @@ "date_released": "2024-05-30T06:38:15Z" }, "3.6.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:59Z" } }, @@ -751,7 +763,7 @@ "date_released": "2024-08-06T20:34:44Z" }, "2.8.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -877,7 +889,7 @@ "date_released": "2024-08-06T20:34:43Z" }, "3.11.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" } }, @@ -927,8 +939,12 @@ "date_released": "2024-06-27T07:50:54Z" }, "1.5.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "1.5.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "apache.hdfs": { @@ -1033,8 +1049,12 @@ "date_released": "2024-06-27T07:50:54Z" }, "4.5.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:59Z" + }, + "4.5.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "apache.hive": { @@ -1215,7 +1235,7 @@ "date_released": "2024-06-27T07:50:54Z" }, "8.2.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -1225,7 +1245,7 @@ "date_released": "2024-05-17T16:07:16Z" }, "1.1.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" } }, @@ -1275,8 +1295,12 @@ "date_released": "2024-08-06T20:34:44Z" }, "1.5.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "1.5.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "apache.kafka": { @@ -1321,7 +1345,7 @@ "date_released": "2024-06-27T07:50:54Z" }, "1.6.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -1395,7 +1419,7 @@ "date_released": "2024-06-27T07:50:54Z" }, "3.7.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -1505,8 +1529,12 @@ "date_released": "2024-05-30T06:38:15Z" }, "3.9.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "3.9.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "apache.pig": { @@ -1575,7 +1603,7 @@ "date_released": "2024-05-30T06:38:15Z" }, "4.5.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" } }, @@ -1677,7 +1705,7 @@ "date_released": "2024-08-06T20:34:44Z" }, "4.5.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -1815,8 +1843,12 @@ "date_released": "2024-07-25T14:17:37Z" }, "4.10.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" + }, + "4.11.0": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "apprise": { @@ -1861,7 +1893,7 @@ "date_released": "2024-08-06T20:34:43Z" }, "1.4.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" } }, @@ -1915,7 +1947,7 @@ "date_released": "2024-05-30T06:38:16Z" }, "2.6.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -1985,7 +2017,7 @@ "date_released": "2024-05-30T06:38:15Z" }, "2.6.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" } }, @@ -2043,7 +2075,7 @@ "date_released": "2024-05-30T06:38:15Z" }, "2.7.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:59Z" } }, @@ -2169,8 +2201,12 @@ "date_released": "2024-08-22T10:37:58Z" }, "3.8.1": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-28T10:31:24Z" + }, + "3.8.2": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "cloudant": { @@ -2243,8 +2279,12 @@ "date_released": "2024-06-27T07:50:54Z" }, "3.6.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "4.0.0": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "cncf.kubernetes": { @@ -2497,8 +2537,12 @@ "date_released": "2024-08-22T10:37:58Z" }, "8.4.1": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-28T10:31:24Z" + }, + "8.4.2": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "cohere": { @@ -2531,7 +2575,7 @@ "date_released": "2024-05-30T06:38:14Z" }, "1.3.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -2545,7 +2589,7 @@ "date_released": "2024-08-06T20:34:43Z" }, "1.2.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" } }, @@ -2581,6 +2625,10 @@ "1.4.0": { "associated_airflow_version": "2.10.0", "date_released": "2024-08-06T20:34:43Z" + }, + "1.4.1": { + "associated_airflow_version": "2.10.0", + "date_released": "2024-09-24T13:49:56Z" } }, "common.sql": { @@ -2709,8 +2757,12 @@ "date_released": "2024-08-06T20:34:43Z" }, "1.16.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" + }, + "1.17.0": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "databricks": { @@ -2875,8 +2927,12 @@ "date_released": "2024-08-06T20:34:43Z" }, "6.9.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "6.10.0": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "datadog": { @@ -2953,8 +3009,12 @@ "date_released": "2024-05-30T06:38:15Z" }, "3.7.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" + }, + "3.7.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "dbt.cloud": { @@ -3067,8 +3127,12 @@ "date_released": "2024-06-27T07:50:54Z" }, "3.10.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "3.10.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "dingding": { @@ -3137,7 +3201,7 @@ "date_released": "2024-05-30T06:38:16Z" }, "3.6.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -3219,7 +3283,7 @@ "date_released": "2024-05-30T06:38:15Z" }, "3.8.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -3397,8 +3461,12 @@ "date_released": "2024-08-06T20:34:43Z" }, "3.13.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" + }, + "3.14.0": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "elasticsearch": { @@ -3559,8 +3627,12 @@ "date_released": "2024-08-06T20:34:43Z" }, "5.5.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" + }, + "5.5.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "exasol": { @@ -3693,7 +3765,7 @@ "date_released": "2024-08-06T20:34:43Z" }, "4.6.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" } }, @@ -3739,8 +3811,12 @@ "date_released": "2024-07-31T14:18:50Z" }, "1.3.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "1.4.0": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "facebook": { @@ -3829,7 +3905,7 @@ "date_released": "2024-06-27T07:50:54Z" }, "3.6.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -3943,8 +4019,12 @@ "date_released": "2024-08-06T20:34:44Z" }, "3.11.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "3.11.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "github": { @@ -4017,7 +4097,7 @@ "date_released": "2024-06-27T07:50:54Z" }, "2.7.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:59Z" } }, @@ -4259,8 +4339,12 @@ "date_released": "2024-08-06T20:34:43Z" }, "10.22.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "10.23.0": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "grpc": { @@ -4341,7 +4425,7 @@ "date_released": "2024-06-27T07:50:54Z" }, "3.6.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -4459,7 +4543,7 @@ "date_released": "2024-05-30T06:38:15Z" }, "3.8.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" } }, @@ -4593,8 +4677,12 @@ "date_released": "2024-06-27T07:50:54Z" }, "4.13.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "4.13.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "imap": { @@ -4687,7 +4775,7 @@ "date_released": "2024-05-30T06:38:16Z" }, "3.7.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -4761,8 +4849,12 @@ "date_released": "2024-07-12T12:38:31Z" }, "2.7.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "2.7.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "jdbc": { @@ -4863,8 +4955,12 @@ "date_released": "2024-08-06T20:34:44Z" }, "4.5.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" + }, + "4.5.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "jenkins": { @@ -4965,8 +5061,12 @@ "date_released": "2024-05-30T06:38:15Z" }, "3.7.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" + }, + "3.7.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "microsoft.azure": { @@ -5195,8 +5295,12 @@ "date_released": "2024-08-06T20:34:43Z" }, "10.4.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "10.5.0": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "microsoft.mssql": { @@ -5305,8 +5409,12 @@ "date_released": "2024-08-06T20:34:43Z" }, "3.9.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "3.9.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "microsoft.psrp": { @@ -5387,7 +5495,7 @@ "date_released": "2024-05-30T06:38:15Z" }, "2.8.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" } }, @@ -5473,7 +5581,7 @@ "date_released": "2024-05-30T06:38:15Z" }, "3.6.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -5571,8 +5679,12 @@ "date_released": "2024-06-27T07:50:54Z" }, "4.2.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" + }, + "4.2.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "mysql": { @@ -5725,8 +5837,12 @@ "date_released": "2024-08-06T20:34:43Z" }, "5.7.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "5.7.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "neo4j": { @@ -5815,7 +5931,7 @@ "date_released": "2024-05-30T06:38:15Z" }, "3.7.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" } }, @@ -5921,8 +6037,12 @@ "date_released": "2024-08-06T20:34:43Z" }, "4.7.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "4.7.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "openai": { @@ -5951,8 +6071,12 @@ "date_released": "2024-06-27T07:50:54Z" }, "1.3.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "1.4.0": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "openfaas": { @@ -6017,7 +6141,7 @@ "date_released": "2024-05-30T06:38:15Z" }, "3.6.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" } }, @@ -6095,8 +6219,12 @@ "date_released": "2024-08-06T20:34:43Z" }, "1.11.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-28T10:31:24Z" + }, + "1.12.0": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "opensearch": { @@ -6129,7 +6257,7 @@ "date_released": "2024-06-27T07:50:54Z" }, "1.4.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -6215,7 +6343,7 @@ "date_released": "2024-05-30T06:38:16Z" }, "5.7.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -6345,7 +6473,7 @@ "date_released": "2024-07-12T12:38:31Z" }, "3.11.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:59Z" } }, @@ -6439,7 +6567,7 @@ "date_released": "2024-06-27T07:50:54Z" }, "3.8.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" } }, @@ -6537,8 +6665,12 @@ "date_released": "2024-06-27T07:50:54Z" }, "3.8.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "3.8.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "pgvector": { @@ -6563,7 +6695,7 @@ "date_released": "2024-08-06T20:34:43Z" }, "1.3.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -6593,7 +6725,7 @@ "date_released": "2024-06-27T07:50:54Z" }, "2.1.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" } }, @@ -6743,8 +6875,12 @@ "date_released": "2024-08-06T20:34:43Z" }, "5.12.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:59Z" + }, + "5.13.0": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "presto": { @@ -6881,7 +7017,7 @@ "date_released": "2024-06-27T07:50:54Z" }, "5.6.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -6903,7 +7039,7 @@ "date_released": "2024-08-06T20:34:43Z" }, "1.2.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -6993,7 +7129,7 @@ "date_released": "2024-05-30T06:38:16Z" }, "3.8.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -7119,7 +7255,7 @@ "date_released": "2024-06-27T07:50:54Z" }, "5.8.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -7201,7 +7337,7 @@ "date_released": "2024-05-30T06:38:15Z" }, "4.8.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -7267,7 +7403,7 @@ "date_released": "2024-05-30T06:38:16Z" }, "3.6.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -7341,7 +7477,7 @@ "date_released": "2024-05-30T06:38:14Z" }, "3.6.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -7499,8 +7635,12 @@ "date_released": "2024-08-06T20:34:43Z" }, "4.11.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "4.11.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "singularity": { @@ -7573,7 +7713,7 @@ "date_released": "2024-05-30T06:38:16Z" }, "3.6.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -7711,7 +7851,7 @@ "date_released": "2024-08-06T20:34:43Z" }, "8.9.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -7773,7 +7913,7 @@ "date_released": "2024-05-30T06:38:15Z" }, "1.8.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -7975,8 +8115,12 @@ "date_released": "2024-08-06T20:34:43Z" }, "5.7.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" + }, + "5.7.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "sqlite": { @@ -8089,7 +8233,7 @@ "date_released": "2024-08-06T20:34:43Z" }, "3.9.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -8235,7 +8379,7 @@ "date_released": "2024-08-22T10:37:57Z" }, "3.13.1": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-28T10:31:24Z" } }, @@ -8341,58 +8485,12 @@ "date_released": "2024-06-27T07:50:54Z" }, "4.6.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" - } - }, - "tabular": { - "1.0.0": { - "associated_airflow_version": "2.3.4", - "date_released": "2022-07-13T20:26:57Z" }, - "1.0.1": { - "associated_airflow_version": "2.4.0", - "date_released": "2022-07-17T09:00:32Z" - }, - "1.1.0": { - "associated_airflow_version": "2.5.0", - "date_released": "2022-11-18T10:44:03Z" - }, - "1.2.0": { - "associated_airflow_version": "2.6.2", - "date_released": "2023-05-23T14:20:25Z" - }, - "1.2.1": { - "associated_airflow_version": "2.6.3", - "date_released": "2023-06-23T15:38:46Z" - }, - "1.3.0": { - "associated_airflow_version": "2.7.3", - "date_released": "2023-10-17T07:49:17Z" - }, - "1.4.0": { - "associated_airflow_version": "2.8.0", - "date_released": "2023-12-12T07:17:13Z" - }, - "1.4.1": { - "associated_airflow_version": "2.8.1", - "date_released": "2023-12-27T23:07:27Z" - }, - "1.5.0": { - "associated_airflow_version": "2.8.1", - "date_released": "2024-05-06T08:35:20Z" - }, - "1.5.1": { - "associated_airflow_version": "2.9.2", - "date_released": "2024-05-17T16:07:16Z" - }, - "1.6.0": { - "associated_airflow_version": "2.9.2", - "date_released": "2024-08-22T10:37:58Z" - }, - "1.6.1": { - "associated_airflow_version": "2.9.2", - "date_released": "2024-08-28T10:31:24Z" + "4.6.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-09-24T13:49:56Z" } }, "telegram": { @@ -8481,7 +8579,7 @@ "date_released": "2024-06-27T07:50:54Z" }, "4.6.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -8515,7 +8613,7 @@ "date_released": "2024-08-06T20:34:43Z" }, "2.6.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -8661,7 +8759,7 @@ "date_released": "2024-06-27T07:50:54Z" }, "5.8.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -8767,7 +8865,7 @@ "date_released": "2024-06-27T07:50:54Z" }, "3.9.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -8817,7 +8915,7 @@ "date_released": "2024-07-15T11:42:08Z" }, "2.1.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:59Z" } }, @@ -8915,7 +9013,7 @@ "date_released": "2024-06-27T07:50:54Z" }, "3.12.0": { - "associated_airflow_version": "2.9.3", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:58Z" } }, @@ -8933,7 +9031,7 @@ "date_released": "2024-08-06T20:34:43Z" }, "1.3.0": { - "associated_airflow_version": "2.10.0", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:59Z" } }, @@ -9015,7 +9113,7 @@ "date_released": "2024-05-30T06:38:15Z" }, "4.8.0": { - "associated_airflow_version": "2.9.2", + "associated_airflow_version": "2.10.1", "date_released": "2024-08-22T10:37:57Z" } } From d534cc01d8059d7a86e2316ea07c777bfdcf69ae Mon Sep 17 00:00:00 2001 From: Wei Lee Date: Tue, 24 Sep 2024 08:04:52 -0700 Subject: [PATCH 004/802] fix(providers/amazon): handle ClientError raised after key is missing during table.get_item (#42408) --- .../providers/amazon/aws/sensors/dynamodb.py | 30 ++++++++++++++----- .../amazon/aws/sensors/test_dynamodb.py | 17 +++++++++++ 2 files changed, 39 insertions(+), 8 deletions(-) diff --git a/airflow/providers/amazon/aws/sensors/dynamodb.py b/airflow/providers/amazon/aws/sensors/dynamodb.py index dbb7f973041e6..ead8c123a621a 100644 --- a/airflow/providers/amazon/aws/sensors/dynamodb.py +++ b/airflow/providers/amazon/aws/sensors/dynamodb.py @@ -18,6 +18,8 @@ from typing import TYPE_CHECKING, Any, Iterable, Sequence +from botocore.exceptions import ClientError + from airflow.providers.amazon.aws.hooks.dynamodb import DynamoDBHook from airflow.providers.amazon.aws.sensors.base_aws import AwsBaseSensor from airflow.providers.amazon.aws.utils.mixins import aws_template_fields @@ -102,14 +104,26 @@ def poke(self, context: Context) -> bool: table = self.hook.conn.Table(self.table_name) self.log.info("Table: %s", table) self.log.info("Key: %s", key) - response = table.get_item(Key=key) + try: - item_attribute_value = response["Item"][self.attribute_name] - self.log.info("Response: %s", response) - self.log.info("Want: %s = %s", self.attribute_name, self.attribute_value) - self.log.info("Got: {response['Item'][self.attribute_name]} = %s", item_attribute_value) - return item_attribute_value in ( - [self.attribute_value] if isinstance(self.attribute_value, str) else self.attribute_value + response = table.get_item(Key=key) + except ClientError as err: + self.log.error( + "Couldn't get %s from table %s.\nError Code: %s\nError Message: %s", + key, + self.table_name, + err.response["Error"]["Code"], + err.response["Error"]["Message"], ) - except KeyError: return False + else: + try: + item_attribute_value = response["Item"][self.attribute_name] + self.log.info("Response: %s", response) + self.log.info("Want: %s = %s", self.attribute_name, self.attribute_value) + self.log.info("Got: {response['Item'][self.attribute_name]} = %s", item_attribute_value) + return item_attribute_value in ( + [self.attribute_value] if isinstance(self.attribute_value, str) else self.attribute_value + ) + except KeyError: + return False diff --git a/tests/providers/amazon/aws/sensors/test_dynamodb.py b/tests/providers/amazon/aws/sensors/test_dynamodb.py index d8b31b48c5e3c..93ca01d26275d 100644 --- a/tests/providers/amazon/aws/sensors/test_dynamodb.py +++ b/tests/providers/amazon/aws/sensors/test_dynamodb.py @@ -104,6 +104,23 @@ def test_sensor_with_pk_and_sk(self): assert self.sensor_pk_sk.poke(None) + @mock_aws + def test_sensor_with_client_error(self): + hook = DynamoDBHook(table_name=self.table_name, table_keys=[self.pk_name]) + + hook.conn.create_table( + TableName=self.table_name, + KeySchema=[{"AttributeName": self.pk_name, "KeyType": "HASH"}], + AttributeDefinitions=[{"AttributeName": self.pk_name, "AttributeType": "S"}], + ProvisionedThroughput={"ReadCapacityUnits": 10, "WriteCapacityUnits": 10}, + ) + + items = [{self.pk_name: self.pk_value, self.attribute_name: self.attribute_value}] + hook.write_batch_data(items) + + self.sensor_pk.partition_key_name = "no such key" + assert self.sensor_pk.poke(None) is False + class TestDynamoDBMultipleValuesSensor: def setup_method(self): From 10278bdf4d49ee0f9aa6472843bf24dbea2a60d6 Mon Sep 17 00:00:00 2001 From: Vincent <97131062+vincbeck@users.noreply.github.com> Date: Tue, 24 Sep 2024 08:13:51 -0700 Subject: [PATCH 005/802] Simple auth manager documentation (#42390) --- .../auth-manager/index.rst | 3 - .../core-concepts/auth-manager.rst | 179 ------------------ .../core-concepts/auth-manager/index.rst | 2 +- docs/apache-airflow/core-concepts/index.rst | 1 - 4 files changed, 1 insertion(+), 184 deletions(-) delete mode 100644 docs/apache-airflow/core-concepts/auth-manager.rst diff --git a/docs/apache-airflow-providers-amazon/auth-manager/index.rst b/docs/apache-airflow-providers-amazon/auth-manager/index.rst index 7d9b226037cf3..c01fc5403e288 100644 --- a/docs/apache-airflow-providers-amazon/auth-manager/index.rst +++ b/docs/apache-airflow-providers-amazon/auth-manager/index.rst @@ -22,9 +22,6 @@ AWS auth manager .. warning:: The AWS auth manager is alpha/experimental at the moment and may be subject to change without warning. -Before reading this, you should be familiar with the concept of auth manager. -See :doc:`apache-airflow:core-concepts/auth-manager`. - The AWS auth manager is an auth manager powered by AWS. It uses two services: * `AWS IAM Identity Center `_ for authentication purposes diff --git a/docs/apache-airflow/core-concepts/auth-manager.rst b/docs/apache-airflow/core-concepts/auth-manager.rst deleted file mode 100644 index 521264fd78ba7..0000000000000 --- a/docs/apache-airflow/core-concepts/auth-manager.rst +++ /dev/null @@ -1,179 +0,0 @@ - .. 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. - -Auth manager -============ - -Auth (for authentication/authorization) manager is the component in Airflow to handle user authentication and user authorization. They have a common -API and are "pluggable", meaning you can swap auth managers based on your installation needs. - -.. image:: ../img/diagram_auth_manager_airflow_architecture.png - -Airflow can only have one auth manager configured at a time; this is set by the ``auth_manager`` option in the -``[core]`` section of :doc:`the configuration file `. - -.. note:: - For more information on Airflow's configuration, see :doc:`/howto/set-config`. - -If you want to check which auth manager is currently set, you can use the -``airflow config get-value core auth_manager`` command: - -.. code-block:: bash - - $ airflow config get-value core auth_manager - airflow.providers.fab.auth_manager.fab_auth_manager.FabAuthManager - - -Why pluggable auth managers? ----------------------------- - -Airflow is used by a lot of different users with a lot of different configurations. Some Airflow environment might be -used by only one user and some might be used by thousand of users. An Airflow environment with only one (or very few) -users does not need the same user management as an environment used by thousand of them. - -This is why the whole user management (user authentication and user authorization) is packaged in one component -called auth manager. So that it is easy to plug-and-play an auth manager that suits your specific needs. - -By default, Airflow comes with the :doc:`apache-airflow-providers-fab:auth-manager/index`. - -.. note:: - Switching to a different auth manager is a heavy operation and should be considered as such. It will - impact users of the environment. The sign-in and sign-off experience will very likely change and disturb them if - they are not advised. Plus, all current users and permissions will have to be copied over from the previous auth - manager to the next. - -Writing your own auth manager ------------------------------ - -All Airflow auth managers implement a common interface so that they are pluggable and any auth manager has access -to all abilities and integrations within Airflow. This interface is used across Airflow to perform all user -authentication and user authorization related operation. - -The public interface is :class:`~airflow.auth.managers.base_auth_manager.BaseAuthManager`. -You can look through the code for the most detailed and up to date interface, but some important highlights are -outlined below. - -.. note:: - For more information about Airflow's public interface see :doc:`/public-airflow-interface`. - -Some reasons you may want to write a custom auth manager include: - -* An auth manager does not exist which fits your specific use case, such as a specific tool or service for user management. -* You'd like to use an auth manager that leverages an identity provider from your preferred cloud provider. -* You have a private user management tool that is only available to you or your organization. - - -Authentication related BaseAuthManager methods -^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ - -* ``is_logged_in``: Return whether the user is signed-in. -* ``get_user``: Return the signed-in user. -* ``get_url_login``: Return the URL the user is redirected to for signing in. -* ``get_url_logout``: Return the URL the user is redirected to for signing out. - -Authorization related BaseAuthManager methods -^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ - -Most of authorization methods in :class:`~airflow.auth.managers.base_auth_manager.BaseAuthManager` look the same. -Let's go over the different parameters used by most of these methods. - -* ``method``: Use HTTP method naming to determine the type of action being done on a specific resource. - - * ``GET``: Can the user read the resource? - * ``POST``: Can the user create a resource? - * ``PUT``: Can the user modify the resource? - * ``DELETE``: Can the user delete the resource? - * ``MENU``: Can the user see the resource in the menu? - -* ``details``: Optional details about the resource being accessed. -* ``user``: The user trying to access the resource. - -These authorization methods are: - -* ``is_authorized_configuration``: Return whether the user is authorized to access Airflow configuration. Some details about the configuration can be provided (e.g. the config section). -* ``is_authorized_connection``: Return whether the user is authorized to access Airflow connections. Some details about the connection can be provided (e.g. the connection ID). -* ``is_authorized_dag``: Return whether the user is authorized to access a DAG. Some details about the DAG can be provided (e.g. the DAG ID). - Also, ``is_authorized_dag`` is called for any entity related to DAGs (e.g. task instances, dag runs, ...). This information is passed in ``access_entity``. - Example: ``auth_manager.is_authorized_dag(method="GET", access_entity=DagAccessEntity.Run, details=DagDetails(id="dag-1"))`` asks - whether the user has permission to read the Dag runs of the dag "dag-1". -* ``is_authorized_dataset``: Return whether the user is authorized to access Airflow datasets. Some details about the dataset can be provided (e.g. the dataset uri). -* ``is_authorized_pool``: Return whether the user is authorized to access Airflow pools. Some details about the pool can be provided (e.g. the pool name). -* ``is_authorized_variable``: Return whether the user is authorized to access Airflow variables. Some details about the variable can be provided (e.g. the variable key). -* ``is_authorized_view``: Return whether the user is authorized to access a specific view in Airflow. The view is specified through ``access_view`` (e.g. ``AccessView.CLUSTER_ACTIVITY``). -* ``is_authorized_custom_view``: Return whether the user is authorized to access a specific view not defined in Airflow. This view can be provided by the auth manager itself or a plugin defined by the user. - -Optional methods recommended to override for optimization -^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ - -The following methods aren't required to override to have a functional Airflow auth manager. However, it is recommended to override these to make your auth manager faster (and potentially less costly): - -* ``batch_is_authorized_dag``: Batch version of ``is_authorized_dag``. If not overridden, it will call ``is_authorized_dag`` for every single item. -* ``batch_is_authorized_connection``: Batch version of ``is_authorized_connection``. If not overridden, it will call ``is_authorized_connection`` for every single item. -* ``batch_is_authorized_pool``: Batch version of ``is_authorized_pool``. If not overridden, it will call ``is_authorized_pool`` for every single item. -* ``batch_is_authorized_variable``: Batch version of ``is_authorized_variable``. If not overridden, it will call ``is_authorized_variable`` for every single item. -* ``get_permitted_dag_ids``: Return the list of DAG IDs the user has access to. If not overridden, it will call ``is_authorized_dag`` for every single DAG available in the environment. -* ``filter_permitted_menu_items``: Return the menu items the user has access to. If not overridden, it will call ``has_access`` in :class:`~airflow.www.security_manager.AirflowSecurityManagerV2` for every single menu item. - -CLI -^^^ - -Auth managers may vend CLI commands which will be included in the ``airflow`` command line tool by implementing the ``get_cli_commands`` method. The commands can be used to setup required resources. Commands are only vended for the currently configured auth manager. A pseudo-code example of implementing CLI command vending from an auth manager can be seen below: - -.. code-block:: python - - @staticmethod - def get_cli_commands() -> list[CLICommand]: - sub_commands = [ - ActionCommand( - name="command_name", - help="Description of what this specific command does", - func=lazy_load_command("path.to.python.function.for.command"), - args=(), - ), - ] - - return [ - GroupCommand( - name="my_cool_auth_manager", - help="Description of what this group of commands do", - subcommands=sub_commands, - ), - ] - -.. note:: - Currently there are no strict rules in place for the Airflow command namespace. It is up to developers to use names for their CLI commands that are sufficiently unique so as to not cause conflicts with other Airflow components. - -.. note:: - When creating a new auth manager, or updating any existing auth manager, be sure to not import or execute any expensive operations/code at the module level. Auth manager classes are imported in several places and if they are slow to import this will negatively impact the performance of your Airflow environment, especially for CLI commands. - -Rest API -^^^^^^^^ - -Auth managers may vend Rest API endpoints which will be included in the :doc:`/stable-rest-api-ref` by implementing the ``get_api_endpoints`` method. The endpoints can be used to manage resources such as users, groups, roles (if any) handled by your auth manager. Endpoints are only vended for the currently configured auth manager. - -Next Steps -^^^^^^^^^^ - -Once you have created a new auth manager class implementing the :class:`~airflow.auth.managers.base_auth_manager.BaseAuthManager` interface, you can configure Airflow to use it by setting the ``core.auth_manager`` configuration value to the module path of your auth manager: - -.. code-block:: ini - - [core] - auth_manager = my_company.auth_managers.MyCustomAuthManager - -.. note:: - For more information on Airflow's configuration, see :doc:`/howto/set-config` and for more information on managing Python modules in Airflow see :doc:`/administration-and-deployment/modules_management`. diff --git a/docs/apache-airflow/core-concepts/auth-manager/index.rst b/docs/apache-airflow/core-concepts/auth-manager/index.rst index cf64a8e96000a..b61b44ae39ec4 100644 --- a/docs/apache-airflow/core-concepts/auth-manager/index.rst +++ b/docs/apache-airflow/core-concepts/auth-manager/index.rst @@ -49,7 +49,7 @@ Provided by Airflow: simple -Provided by providers +Provided by providers: * :doc:`apache-airflow-providers-fab:auth-manager/index` * :doc:`apache-airflow-providers-amazon:auth-manager/index` diff --git a/docs/apache-airflow/core-concepts/index.rst b/docs/apache-airflow/core-concepts/index.rst index 739111d4be090..0fded6ba495f4 100644 --- a/docs/apache-airflow/core-concepts/index.rst +++ b/docs/apache-airflow/core-concepts/index.rst @@ -40,7 +40,6 @@ Here you can find detailed documentation about each one of the core concepts of sensors taskflow executor/index - auth-manager auth-manager/index objectstorage From 91529d32bbc472624ff34cd1fe1262ad9f77b02e Mon Sep 17 00:00:00 2001 From: Pierre Jeambrun Date: Tue, 24 Sep 2024 23:15:14 +0800 Subject: [PATCH 006/802] Fix UI pre commit hook (#42435) --- .pre-commit-config.yaml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 4ca7f304d1e6d..942b34ca2e6d5 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1174,7 +1174,7 @@ repos: description: TS types generation / ESLint / Prettier new UI files language: node types_or: [javascript, ts, tsx, yaml, css, json] - files: ^airflow/ui/|^airflow/api_connexion/openapi/v1\.yaml$ + files: ^airflow/ui/|^airflow/api_fastapi/openapi/v1-generated\.yaml$ entry: ./scripts/ci/pre_commit/lint_ui.py additional_dependencies: ['pnpm@9.7.1'] pass_filenames: false From db46c9257ec6c6396d8c02263fb8db40fefd0a29 Mon Sep 17 00:00:00 2001 From: Vincent <97131062+vincbeck@users.noreply.github.com> Date: Tue, 24 Sep 2024 08:55:38 -0700 Subject: [PATCH 007/802] Fix logout in AWS auth manager (#42447) --- airflow/providers/amazon/aws/auth_manager/views/auth.py | 2 +- tests/providers/amazon/aws/auth_manager/views/test_auth.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/airflow/providers/amazon/aws/auth_manager/views/auth.py b/airflow/providers/amazon/aws/auth_manager/views/auth.py index 7ea602d0dd45d..e08c2a7a6e100 100644 --- a/airflow/providers/amazon/aws/auth_manager/views/auth.py +++ b/airflow/providers/amazon/aws/auth_manager/views/auth.py @@ -61,7 +61,7 @@ def login(self): saml_auth = self._init_saml_auth() return redirect(saml_auth.login()) - @expose("/logout") + @expose("/logout", methods=("GET", "POST")) def logout(self): """Start logout process.""" session.clear() diff --git a/tests/providers/amazon/aws/auth_manager/views/test_auth.py b/tests/providers/amazon/aws/auth_manager/views/test_auth.py index 435dd8d2c32fe..05d2fb84b51cf 100644 --- a/tests/providers/amazon/aws/auth_manager/views/test_auth.py +++ b/tests/providers/amazon/aws/auth_manager/views/test_auth.py @@ -69,7 +69,7 @@ def aws_app(): ) as mock_is_policy_store_schema_up_to_date: mock_is_policy_store_schema_up_to_date.return_value = True mock_parser.parse_remote.return_value = SAML_METADATA_PARSED - return application.create_app(testing=True) + return application.create_app(testing=True, config={"WTF_CSRF_ENABLED": False}) @pytest.mark.db_test @@ -82,7 +82,7 @@ def test_login_redirect_to_identity_center(self, aws_app): def test_logout_redirect_to_identity_center(self, aws_app): with aws_app.test_client() as client: - response = client.get("/logout") + response = client.post("/logout") assert response.status_code == 302 assert response.location.startswith("https://portal.sso.us-east-1.amazonaws.com/saml/logout/") From c53b9c6f0455818b6c73ec9c8d639e699caea57d Mon Sep 17 00:00:00 2001 From: Daniel Standish <15932138+dstandish@users.noreply.github.com> Date: Tue, 24 Sep 2024 10:30:05 -0700 Subject: [PATCH 008/802] Split next_dagruns_to_examine function into two (#42386) The behavior is different enough to merit two different functions. In fact I noticed that we actually are using a bad index hint for the QUEUED case. And this becomes more apparent with introduction of backfill handling into scheduler, which is forthcoming. --- airflow/jobs/scheduler_job_runner.py | 10 +--- airflow/models/dagrun.py | 84 +++++++++++++++++++--------- tests/jobs/test_scheduler_job.py | 4 +- tests/models/test_dagrun.py | 8 ++- 4 files changed, 71 insertions(+), 35 deletions(-) diff --git a/airflow/jobs/scheduler_job_runner.py b/airflow/jobs/scheduler_job_runner.py index eb0abbc296e6a..a49c2361ec423 100644 --- a/airflow/jobs/scheduler_job_runner.py +++ b/airflow/jobs/scheduler_job_runner.py @@ -1218,9 +1218,10 @@ def _do_scheduling(self, session: Session) -> int: self._start_queued_dagruns(session) guard.commit() - dag_runs = self._get_next_dagruns_to_examine(DagRunState.RUNNING, session) + # Bulk fetch the currently active dag runs for the dags we are # examining, rather than making one query per DagRun + dag_runs = DagRun.get_running_dag_runs_to_examine(session=session) callback_tuples = self._schedule_all_dag_runs(guard, dag_runs, session) @@ -1274,11 +1275,6 @@ def _do_scheduling(self, session: Session) -> int: return num_queued_tis - @retry_db_transaction - def _get_next_dagruns_to_examine(self, state: DagRunState, session: Session) -> Query: - """Get Next DagRuns to Examine with retries.""" - return DagRun.next_dagruns_to_examine(state, session) - @retry_db_transaction def _create_dagruns_for_dags(self, guard: CommitProhibitorGuard, session: Session) -> None: """Find Dag Models needing DagRuns and Create Dag Runs with retries in case of OperationalError.""" @@ -1512,7 +1508,7 @@ def _should_update_dag_next_dagruns( def _start_queued_dagruns(self, session: Session) -> None: """Find DagRuns in queued state and decide moving them to running state.""" # added all() to save runtime, otherwise query is executed more than once - dag_runs: Collection[DagRun] = self._get_next_dagruns_to_examine(DagRunState.QUEUED, session).all() + dag_runs: Collection[DagRun] = DagRun.get_queued_dag_runs_to_set_running(session).all() active_runs_of_dags = Counter( DagRun.active_runs_of_dags((dr.dag_id for dr in dag_runs), only_running=True, session=session), diff --git a/airflow/models/dagrun.py b/airflow/models/dagrun.py index c932958861f7a..3ef1c18f152a4 100644 --- a/airflow/models/dagrun.py +++ b/airflow/models/dagrun.py @@ -67,6 +67,7 @@ from airflow.utils.dates import datetime_to_nano from airflow.utils.helpers import chunks, is_container, prune_dict from airflow.utils.log.logging_mixin import LoggingMixin +from airflow.utils.retries import retry_db_transaction from airflow.utils.session import NEW_SESSION, provide_session from airflow.utils.sqlalchemy import UtcDateTime, nulls_first, tuple_in_condition, with_row_locks from airflow.utils.state import DagRunState, State, TaskInstanceState @@ -388,12 +389,8 @@ def active_runs_of_dags( return dict(iter(session.execute(query))) @classmethod - def next_dagruns_to_examine( - cls, - state: DagRunState, - session: Session, - max_number: int | None = None, - ) -> Query: + @retry_db_transaction + def get_running_dag_runs_to_examine(cls, session: Session) -> Query: """ Return the next DagRuns that the scheduler should attempt to schedule. @@ -401,42 +398,79 @@ def next_dagruns_to_examine( query, you should ensure that any scheduling decisions are made in a single transaction -- as soon as the transaction is committed it will be unlocked. + :meta private: """ from airflow.models.dag import DagModel - if max_number is None: - max_number = cls.DEFAULT_DAGRUNS_TO_EXAMINE - - # TODO: Bake this query, it is run _A lot_ query = ( select(cls) .with_hint(cls, "USE INDEX (idx_dag_run_running_dags)", dialect_name="mysql") - .where(cls.state == state, cls.run_type != DagRunType.BACKFILL_JOB) + .where(cls.state == DagRunState.RUNNING, cls.run_type != DagRunType.BACKFILL_JOB) .join(DagModel, DagModel.dag_id == cls.dag_id) .where(DagModel.is_paused == false(), DagModel.is_active == true()) + .order_by( + nulls_first(cls.last_scheduling_decision, session=session), + cls.execution_date, + ) ) - if state == DagRunState.QUEUED: - # For dag runs in the queued state, we check if they have reached the max_active_runs limit - # and if so we drop them - running_drs = ( - select(DagRun.dag_id, func.count(DagRun.state).label("num_running")) - .where(DagRun.state == DagRunState.RUNNING) - .group_by(DagRun.dag_id) - .subquery() + + if not settings.ALLOW_FUTURE_EXEC_DATES: + query = query.where(DagRun.execution_date <= func.now()) + + return session.scalars( + with_row_locks( + query.limit(cls.DEFAULT_DAGRUNS_TO_EXAMINE), + of=cls, + session=session, + skip_locked=True, ) - query = query.outerjoin(running_drs, running_drs.c.dag_id == DagRun.dag_id).where( - func.coalesce(running_drs.c.num_running, 0) < DagModel.max_active_runs + ) + + @classmethod + @retry_db_transaction + def get_queued_dag_runs_to_set_running(cls, session: Session) -> Query: + """ + Return the next queued DagRuns that the scheduler should attempt to schedule. + + This will return zero or more DagRun rows that are row-level-locked with a "SELECT ... FOR UPDATE" + query, you should ensure that any scheduling decisions are made in a single transaction -- as soon as + the transaction is committed it will be unlocked. + + :meta private: + """ + from airflow.models.dag import DagModel + + # For dag runs in the queued state, we check if they have reached the max_active_runs limit + # and if so we drop them + running_drs = ( + select(DagRun.dag_id, func.count(DagRun.state).label("num_running")) + .where(DagRun.state == DagRunState.RUNNING) + .group_by(DagRun.dag_id) + .subquery() + ) + query = ( + select(cls) + .where(cls.state == DagRunState.QUEUED, cls.run_type != DagRunType.BACKFILL_JOB) + .join(DagModel, DagModel.dag_id == cls.dag_id) + .where(DagModel.is_paused == false(), DagModel.is_active == true()) + .outerjoin(running_drs, running_drs.c.dag_id == DagRun.dag_id) + .where(func.coalesce(running_drs.c.num_running, 0) < DagModel.max_active_runs) + .order_by( + nulls_first(cls.last_scheduling_decision, session=session), + cls.execution_date, ) - query = query.order_by( - nulls_first(cls.last_scheduling_decision, session=session), - cls.execution_date, ) if not settings.ALLOW_FUTURE_EXEC_DATES: query = query.where(DagRun.execution_date <= func.now()) return session.scalars( - with_row_locks(query.limit(max_number), of=cls, session=session, skip_locked=True) + with_row_locks( + query.limit(cls.DEFAULT_DAGRUNS_TO_EXAMINE), + of=cls, + session=session, + skip_locked=True, + ) ) @classmethod diff --git a/tests/jobs/test_scheduler_job.py b/tests/jobs/test_scheduler_job.py index 9113a2dee1bd1..4067f4fa17902 100644 --- a/tests/jobs/test_scheduler_job.py +++ b/tests/jobs/test_scheduler_job.py @@ -6036,7 +6036,9 @@ def test_execute_queries_count_with_harvested_dags(self, expected_query_count, d self.job_runner.processor_agent = mock_agent with assert_queries_count(expected_query_count, margin=15): - with mock.patch.object(DagRun, "next_dagruns_to_examine") as mock_dagruns: + with mock.patch.object( + DagRun, DagRun.get_running_dag_runs_to_examine.__name__ + ) as mock_dagruns: query = MagicMock() query.all.return_value = dagruns mock_dagruns.return_value = query diff --git a/tests/models/test_dagrun.py b/tests/models/test_dagrun.py index f72f3b2b794fc..d2f70ce69314b 100644 --- a/tests/models/test_dagrun.py +++ b/tests/models/test_dagrun.py @@ -931,14 +931,18 @@ def test_next_dagruns_to_examine_only_unpaused(self, session, state): **triggered_by_kwargs, ) - runs = DagRun.next_dagruns_to_examine(state, session).all() + if state == DagRunState.RUNNING: + func = DagRun.get_running_dag_runs_to_examine + else: + func = DagRun.get_queued_dag_runs_to_set_running + runs = func(session).all() assert runs == [dr] orm_dag.is_paused = True session.flush() - runs = DagRun.next_dagruns_to_examine(state, session).all() + runs = func(session).all() assert runs == [] @mock.patch.object(Stats, "timing") From 57bea991643a296a7135fd8944ee68e7add73fa8 Mon Sep 17 00:00:00 2001 From: Boris Morel <2323800+borismo@users.noreply.github.com> Date: Wed, 25 Sep 2024 02:22:54 +0800 Subject: [PATCH 009/802] Support session reuse in `RedshiftDataOperator` (#42218) --- airflow/providers/amazon/CHANGELOG.rst | 27 ++++ .../amazon/aws/hooks/redshift_data.py | 66 +++++++-- .../amazon/aws/operators/redshift_data.py | 21 ++- .../amazon/aws/transfers/redshift_to_s3.py | 25 ++-- .../amazon/aws/transfers/s3_to_redshift.py | 12 +- .../providers/amazon/aws/utils/openlineage.py | 4 +- .../operators/redshift/redshift_data.rst | 12 ++ .../amazon/aws/hooks/test_redshift_data.py | 127 +++++++++++++++++- .../aws/operators/test_redshift_data.py | 86 ++++++++++-- .../amazon/aws/utils/test_openlineage.py | 4 +- .../providers/amazon/aws/example_redshift.py | 41 +++++- .../aws/example_redshift_s3_transfers.py | 109 +++++++++++++-- 12 files changed, 468 insertions(+), 66 deletions(-) diff --git a/airflow/providers/amazon/CHANGELOG.rst b/airflow/providers/amazon/CHANGELOG.rst index 126da03ad630f..7596ad3886c7a 100644 --- a/airflow/providers/amazon/CHANGELOG.rst +++ b/airflow/providers/amazon/CHANGELOG.rst @@ -26,6 +26,33 @@ Changelog --------- +Main +...... + +Breaking changes +~~~~~~~~~~~~~~~~ + +.. warning:: + In order to support session reuse in RedshiftData operators, the following breaking changes were introduced: + + The ``database`` argument is now optional and as a result was moved after the ``sql`` argument which is a positional + one. Update your DAGs accordingly if they rely on argument order. Applies to: + * ``RedshiftDataHook``'s ``execute_query`` method + * ``RedshiftDataOperator`` + + ``RedshiftDataHook``'s ``execute_query`` method now returns a ``QueryExecutionOutput`` object instead of just the + statement ID as a string. + + ``RedshiftDataHook``'s ``parse_statement_resposne`` method was renamed to ``parse_statement_response``. + + ``S3ToRedshiftOperator``'s ``schema`` argument is now optional and was moved after the ``s3_key`` positional argument. + Update your DAGs accordingly if they rely on argument order. + +Features +~~~~~~~~ + +* ``Support session reuse in RedshiftDataOperator, RedshiftToS3Operator and S3ToRedshiftOperator (#42218)`` + 8.29.0 ...... diff --git a/airflow/providers/amazon/aws/hooks/redshift_data.py b/airflow/providers/amazon/aws/hooks/redshift_data.py index 3c1f84b1f694c..b2f46c0ef6049 100644 --- a/airflow/providers/amazon/aws/hooks/redshift_data.py +++ b/airflow/providers/amazon/aws/hooks/redshift_data.py @@ -18,8 +18,12 @@ from __future__ import annotations import time +from dataclasses import dataclass from pprint import pformat from typing import TYPE_CHECKING, Any, Iterable +from uuid import UUID + +from pendulum import duration from airflow.providers.amazon.aws.hooks.base_aws import AwsGenericHook from airflow.providers.amazon.aws.utils import trim_none_values @@ -35,6 +39,14 @@ RUNNING_STATES = {"PICKED", "STARTED", "SUBMITTED"} +@dataclass +class QueryExecutionOutput: + """Describes the output of a query execution.""" + + statement_id: str + session_id: str | None + + class RedshiftDataQueryFailedError(ValueError): """Raise an error that redshift data query failed.""" @@ -65,8 +77,8 @@ def __init__(self, *args, **kwargs) -> None: def execute_query( self, - database: str, sql: str | list[str], + database: str | None = None, cluster_identifier: str | None = None, db_user: str | None = None, parameters: Iterable | None = None, @@ -76,23 +88,28 @@ def execute_query( wait_for_completion: bool = True, poll_interval: int = 10, workgroup_name: str | None = None, - ) -> str: + session_id: str | None = None, + session_keep_alive_seconds: int | None = None, + ) -> QueryExecutionOutput: """ Execute a statement against Amazon Redshift. - :param database: the name of the database :param sql: the SQL statement or list of SQL statement to run + :param database: the name of the database :param cluster_identifier: unique identifier of a cluster :param db_user: the database username :param parameters: the parameters for the SQL statement :param secret_arn: the name or ARN of the secret that enables db access :param statement_name: the name of the SQL statement - :param with_event: indicates whether to send an event to EventBridge - :param wait_for_completion: indicates whether to wait for a result, if True wait, if False don't wait + :param with_event: whether to send an event to EventBridge + :param wait_for_completion: whether to wait for a result :param poll_interval: how often in seconds to check the query status :param workgroup_name: name of the Redshift Serverless workgroup. Mutually exclusive with `cluster_identifier`. Specify this parameter to query Redshift Serverless. More info https://docs.aws.amazon.com/redshift/latest/mgmt/working-with-serverless.html + :param session_id: the session identifier of the query + :param session_keep_alive_seconds: duration in seconds to keep the session alive after the query + finishes. The maximum time a session can keep alive is 24 hours :returns statement_id: str, the UUID of the statement """ @@ -105,7 +122,28 @@ def execute_query( "SecretArn": secret_arn, "StatementName": statement_name, "WorkgroupName": workgroup_name, + "SessionId": session_id, + "SessionKeepAliveSeconds": session_keep_alive_seconds, } + + if sum(x is not None for x in (cluster_identifier, workgroup_name, session_id)) != 1: + raise ValueError( + "Exactly one of cluster_identifier, workgroup_name, or session_id must be provided" + ) + + if session_id is not None: + msg = "session_id must be a valid UUID4" + try: + if UUID(session_id).version != 4: + raise ValueError(msg) + except ValueError: + raise ValueError(msg) + + if session_keep_alive_seconds is not None and ( + session_keep_alive_seconds < 0 or duration(seconds=session_keep_alive_seconds).hours > 24 + ): + raise ValueError("Session keep alive duration must be between 0 and 86400 seconds.") + if isinstance(sql, list): kwargs["Sqls"] = sql resp = self.conn.batch_execute_statement(**trim_none_values(kwargs)) @@ -115,13 +153,10 @@ def execute_query( statement_id = resp["Id"] - if bool(cluster_identifier) is bool(workgroup_name): - raise ValueError("Either 'cluster_identifier' or 'workgroup_name' must be specified.") - if wait_for_completion: self.wait_for_results(statement_id, poll_interval=poll_interval) - return statement_id + return QueryExecutionOutput(statement_id=statement_id, session_id=resp.get("SessionId")) def wait_for_results(self, statement_id: str, poll_interval: int) -> str: while True: @@ -135,9 +170,9 @@ def wait_for_results(self, statement_id: str, poll_interval: int) -> str: def check_query_is_finished(self, statement_id: str) -> bool: """Check whether query finished, raise exception is failed.""" resp = self.conn.describe_statement(Id=statement_id) - return self.parse_statement_resposne(resp) + return self.parse_statement_response(resp) - def parse_statement_resposne(self, resp: DescribeStatementResponseTypeDef) -> bool: + def parse_statement_response(self, resp: DescribeStatementResponseTypeDef) -> bool: """Parse the response of describe_statement.""" status = resp["Status"] if status == FINISHED_STATE: @@ -179,8 +214,10 @@ def get_table_primary_key( :param table: Name of the target table :param database: the name of the database :param schema: Name of the target schema, public by default - :param sql: the SQL statement or list of SQL statement to run :param cluster_identifier: unique identifier of a cluster + :param workgroup_name: name of the Redshift Serverless workgroup. Mutually exclusive with + `cluster_identifier`. Specify this parameter to query Redshift Serverless. More info + https://docs.aws.amazon.com/redshift/latest/mgmt/working-with-serverless.html :param db_user: the database username :param secret_arn: the name or ARN of the secret that enables db access :param statement_name: the name of the SQL statement @@ -212,7 +249,8 @@ def get_table_primary_key( with_event=with_event, wait_for_completion=wait_for_completion, poll_interval=poll_interval, - ) + ).statement_id + pk_columns = [] token = "" while True: @@ -251,4 +289,4 @@ async def check_query_is_finished_async(self, statement_id: str) -> bool: """ async with self.async_conn as client: resp = await client.describe_statement(Id=statement_id) - return self.parse_statement_resposne(resp) + return self.parse_statement_response(resp) diff --git a/airflow/providers/amazon/aws/operators/redshift_data.py b/airflow/providers/amazon/aws/operators/redshift_data.py index 45fee2a919483..3d00c6d22edf7 100644 --- a/airflow/providers/amazon/aws/operators/redshift_data.py +++ b/airflow/providers/amazon/aws/operators/redshift_data.py @@ -56,13 +56,16 @@ class RedshiftDataOperator(AwsBaseOperator[RedshiftDataHook]): :param workgroup_name: name of the Redshift Serverless workgroup. Mutually exclusive with `cluster_identifier`. Specify this parameter to query Redshift Serverless. More info https://docs.aws.amazon.com/redshift/latest/mgmt/working-with-serverless.html + :param session_id: the session identifier of the query + :param session_keep_alive_seconds: duration in seconds to keep the session alive after the query + finishes. The maximum time a session can keep alive is 24 hours :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: + :param verify: Whether 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 @@ -77,6 +80,7 @@ class RedshiftDataOperator(AwsBaseOperator[RedshiftDataHook]): "parameters", "statement_name", "workgroup_name", + "session_id", ) template_ext = (".sql",) template_fields_renderers = {"sql": "sql"} @@ -84,8 +88,8 @@ class RedshiftDataOperator(AwsBaseOperator[RedshiftDataHook]): def __init__( self, - database: str, sql: str | list, + database: str | None = None, cluster_identifier: str | None = None, db_user: str | None = None, parameters: list | None = None, @@ -97,6 +101,8 @@ def __init__( return_sql_result: bool = False, workgroup_name: str | None = None, deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), + session_id: str | None = None, + session_keep_alive_seconds: int | None = None, **kwargs, ) -> None: super().__init__(**kwargs) @@ -120,6 +126,8 @@ def __init__( self.return_sql_result = return_sql_result self.statement_id: str | None = None self.deferrable = deferrable + self.session_id = session_id + self.session_keep_alive_seconds = session_keep_alive_seconds def execute(self, context: Context) -> GetStatementResultResponseTypeDef | str: """Execute a statement against Amazon Redshift.""" @@ -130,7 +138,7 @@ def execute(self, context: Context) -> GetStatementResultResponseTypeDef | str: if self.deferrable: wait_for_completion = False - self.statement_id = self.hook.execute_query( + query_execution_output = self.hook.execute_query( database=self.database, sql=self.sql, cluster_identifier=self.cluster_identifier, @@ -142,8 +150,15 @@ def execute(self, context: Context) -> GetStatementResultResponseTypeDef | str: with_event=self.with_event, wait_for_completion=wait_for_completion, poll_interval=self.poll_interval, + session_id=self.session_id, + session_keep_alive_seconds=self.session_keep_alive_seconds, ) + self.statement_id = query_execution_output.statement_id + + if query_execution_output.session_id: + self.xcom_push(context, key="session_id", value=query_execution_output.session_id) + if self.deferrable and self.wait_for_completion: is_finished = self.hook.check_query_is_finished(self.statement_id) if not is_finished: diff --git a/airflow/providers/amazon/aws/transfers/redshift_to_s3.py b/airflow/providers/amazon/aws/transfers/redshift_to_s3.py index ef3cebdae9838..8538b1dfc313c 100644 --- a/airflow/providers/amazon/aws/transfers/redshift_to_s3.py +++ b/airflow/providers/amazon/aws/transfers/redshift_to_s3.py @@ -45,7 +45,8 @@ class RedshiftToS3Operator(BaseOperator): :param s3_key: reference to a specific S3 key. If ``table_as_file_name`` is set to False, this param must include the desired file name :param schema: reference to a specific schema in redshift database, - used when ``table`` param provided and ``select_query`` param not provided + used when ``table`` param provided and ``select_query`` param not provided. + Do not provide when unloading a temporary table :param table: reference to a specific table in redshift database, used when ``schema`` param provided and ``select_query`` param not provided :param select_query: custom select query to fetch data from redshift database, @@ -55,8 +56,8 @@ class RedshiftToS3Operator(BaseOperator): If the AWS connection contains 'aws_iam_role' in ``extras`` the operator will use AWS STS credentials with a token https://docs.aws.amazon.com/redshift/latest/dg/copy-parameters-authorization.html#copy-credentials - :param verify: Whether or not to verify SSL certificates for S3 connection. - By default SSL certificates are verified. + :param verify: Whether to verify SSL certificates for S3 connection. + By default, SSL certificates are verified. You can provide the following values: - ``False``: do not validate SSL certificates. SSL will still be used @@ -67,7 +68,7 @@ class RedshiftToS3Operator(BaseOperator): CA cert bundle than the one used by botocore. :param unload_options: reference to a list of UNLOAD options :param autocommit: If set to True it will automatically commit the UNLOAD statement. - Otherwise it will be committed right before the redshift connection gets closed. + Otherwise, it will be committed right before the redshift connection gets closed. :param include_header: If set to True the s3 file contains the header columns. :param parameters: (optional) the parameters to render the SQL query with. :param table_as_file_name: If set to True, the s3 file will be named as the table. @@ -141,9 +142,15 @@ def _build_unload_query( @property def default_select_query(self) -> str | None: - if self.schema and self.table: - return f"SELECT * FROM {self.schema}.{self.table}" - return None + if not self.table: + return None + + if self.schema: + table = f"{self.schema}.{self.table}" + else: + # Relevant when unloading a temporary table + table = self.table + return f"SELECT * FROM {table}" def execute(self, context: Context) -> None: if self.table and self.table_as_file_name: @@ -152,9 +159,7 @@ def execute(self, context: Context) -> None: self.select_query = self.select_query or self.default_select_query if self.select_query is None: - raise ValueError( - "Please provide both `schema` and `table` params or `select_query` to fetch the data." - ) + raise ValueError("Please specify either a table or `select_query` to fetch the data.") if self.include_header and "HEADER" not in [uo.upper().strip() for uo in self.unload_options]: self.unload_options = [*self.unload_options, "HEADER"] diff --git a/airflow/providers/amazon/aws/transfers/s3_to_redshift.py b/airflow/providers/amazon/aws/transfers/s3_to_redshift.py index 127ee07a60bbd..792119bfebb55 100644 --- a/airflow/providers/amazon/aws/transfers/s3_to_redshift.py +++ b/airflow/providers/amazon/aws/transfers/s3_to_redshift.py @@ -28,7 +28,6 @@ if TYPE_CHECKING: from airflow.utils.context import Context - AVAILABLE_METHODS = ["APPEND", "REPLACE", "UPSERT"] @@ -40,17 +39,18 @@ class S3ToRedshiftOperator(BaseOperator): For more information on how to use this operator, take a look at the guide: :ref:`howto/operator:S3ToRedshiftOperator` - :param schema: reference to a specific schema in redshift database :param table: reference to a specific table in redshift database :param s3_bucket: reference to a specific S3 bucket :param s3_key: key prefix that selects single or multiple objects from S3 + :param schema: reference to a specific schema in redshift database. + Do not provide when copying into a temporary table :param redshift_conn_id: reference to a specific redshift database OR a redshift data-api connection :param aws_conn_id: reference to a specific S3 connection If the AWS connection contains 'aws_iam_role' in ``extras`` the operator will use AWS STS credentials with a token https://docs.aws.amazon.com/redshift/latest/dg/copy-parameters-authorization.html#copy-credentials - :param verify: Whether or not to verify SSL certificates for S3 connection. - By default SSL certificates are verified. + :param verify: Whether to verify SSL certificates for S3 connection. + By default, SSL certificates are verified. You can provide the following values: - ``False``: do not validate SSL certificates. SSL will still be used @@ -87,10 +87,10 @@ class S3ToRedshiftOperator(BaseOperator): def __init__( self, *, - schema: str, table: str, s3_bucket: str, s3_key: str, + schema: str | None = None, redshift_conn_id: str = "redshift_default", aws_conn_id: str | None = "aws_default", verify: bool | str | None = None, @@ -160,7 +160,7 @@ def execute(self, context: Context) -> None: credentials_block = build_credentials_block(credentials) copy_options = "\n\t\t\t".join(self.copy_options) - destination = f"{self.schema}.{self.table}" + destination = f"{self.schema}.{self.table}" if self.schema else self.table copy_destination = f"#{self.table}" if self.method == "UPSERT" else destination copy_statement = self._build_copy_query( diff --git a/airflow/providers/amazon/aws/utils/openlineage.py b/airflow/providers/amazon/aws/utils/openlineage.py index db472a3e46c5f..be5703e2f6e80 100644 --- a/airflow/providers/amazon/aws/utils/openlineage.py +++ b/airflow/providers/amazon/aws/utils/openlineage.py @@ -86,7 +86,9 @@ def get_facets_from_redshift_table( ] ) else: - statement_id = redshift_hook.execute_query(sql=sql, poll_interval=1, **redshift_data_api_kwargs) + statement_id = redshift_hook.execute_query( + sql=sql, poll_interval=1, **redshift_data_api_kwargs + ).statement_id response = redshift_hook.conn.get_statement_result(Id=statement_id) table_schema = SchemaDatasetFacet( diff --git a/docs/apache-airflow-providers-amazon/operators/redshift/redshift_data.rst b/docs/apache-airflow-providers-amazon/operators/redshift/redshift_data.rst index 0b314d34f3193..2638e1732cd6c 100644 --- a/docs/apache-airflow-providers-amazon/operators/redshift/redshift_data.rst +++ b/docs/apache-airflow-providers-amazon/operators/redshift/redshift_data.rst @@ -54,6 +54,18 @@ the necessity of a Postgres connection. :start-after: [START howto_operator_redshift_data] :end-before: [END howto_operator_redshift_data] +Reuse a session when executing multiple statements +================================================== + +Specify the ``session_keep_alive_seconds`` parameter on an upstream task. In a downstream task, get the session ID from +the XCom and pass it to the ``session_id`` parameter. This is useful when you work with temporary tables. + +.. exampleinclude:: /../../tests/system/providers/amazon/aws/example_redshift.py + :language: python + :dedent: 4 + :start-after: [START howto_operator_redshift_data_session_reuse] + :end-before: [END howto_operator_redshift_data_session_reuse] + Reference --------- diff --git a/tests/providers/amazon/aws/hooks/test_redshift_data.py b/tests/providers/amazon/aws/hooks/test_redshift_data.py index a0952e5ba7259..d548086449812 100644 --- a/tests/providers/amazon/aws/hooks/test_redshift_data.py +++ b/tests/providers/amazon/aws/hooks/test_redshift_data.py @@ -19,6 +19,7 @@ import logging from unittest import mock +from uuid import uuid4 import pytest @@ -63,15 +64,18 @@ def test_execute_without_waiting(self, mock_conn): mock_conn.describe_statement.assert_not_called() @pytest.mark.parametrize( - "cluster_identifier, workgroup_name", + "cluster_identifier, workgroup_name, session_id", [ - (None, None), - ("some_cluster", "some_workgroup"), + (None, None, None), + ("some_cluster", "some_workgroup", None), + (None, "some_workgroup", None), + ("some_cluster", None, None), + (None, None, "some_session_id"), ], ) @mock.patch("airflow.providers.amazon.aws.hooks.redshift_data.RedshiftDataHook.conn") - def test_execute_requires_either_cluster_identifier_or_workgroup_name( - self, mock_conn, cluster_identifier, workgroup_name + def test_execute_requires_one_of_cluster_identifier_or_workgroup_name_or_session_id( + self, mock_conn, cluster_identifier, workgroup_name, session_id ): mock_conn.execute_statement.return_value = {"Id": STATEMENT_ID} cluster_identifier = "cluster_identifier" @@ -84,6 +88,51 @@ def test_execute_requires_either_cluster_identifier_or_workgroup_name( workgroup_name=workgroup_name, sql=SQL, wait_for_completion=False, + session_id=session_id, + ) + + @pytest.mark.parametrize( + "cluster_identifier, workgroup_name, session_id", + [ + (None, None, None), + ("some_cluster", "some_workgroup", None), + (None, "some_workgroup", None), + ("some_cluster", None, None), + (None, None, "some_session_id"), + ], + ) + @mock.patch("airflow.providers.amazon.aws.hooks.redshift_data.RedshiftDataHook.conn") + def test_execute_session_keep_alive_seconds_valid( + self, mock_conn, cluster_identifier, workgroup_name, session_id + ): + mock_conn.execute_statement.return_value = {"Id": STATEMENT_ID} + cluster_identifier = "cluster_identifier" + workgroup_name = "workgroup_name" + hook = RedshiftDataHook() + with pytest.raises(ValueError): + hook.execute_query( + database=DATABASE, + cluster_identifier=cluster_identifier, + workgroup_name=workgroup_name, + sql=SQL, + wait_for_completion=False, + session_id=session_id, + ) + + @mock.patch("airflow.providers.amazon.aws.hooks.redshift_data.RedshiftDataHook.conn") + def test_execute_session_id_valid(self, mock_conn): + mock_conn.execute_statement.return_value = {"Id": STATEMENT_ID} + cluster_identifier = "cluster_identifier" + workgroup_name = "workgroup_name" + hook = RedshiftDataHook() + with pytest.raises(ValueError): + hook.execute_query( + database=DATABASE, + cluster_identifier=cluster_identifier, + workgroup_name=workgroup_name, + sql=SQL, + wait_for_completion=False, + session_id="not_a_uuid", ) @mock.patch("airflow.providers.amazon.aws.hooks.redshift_data.RedshiftDataHook.conn") @@ -156,6 +205,74 @@ def test_execute_with_all_parameters_workgroup_name(self, mock_conn): Id=STATEMENT_ID, ) + @mock.patch("airflow.providers.amazon.aws.hooks.redshift_data.RedshiftDataHook.conn") + def test_execute_with_new_session(self, mock_conn): + cluster_identifier = "cluster_identifier" + db_user = "db_user" + secret_arn = "secret_arn" + statement_name = "statement_name" + parameters = [{"name": "id", "value": "1"}] + mock_conn.execute_statement.return_value = {"Id": STATEMENT_ID, "SessionId": "session_id"} + mock_conn.describe_statement.return_value = {"Status": "FINISHED"} + + hook = RedshiftDataHook() + output = hook.execute_query( + sql=SQL, + database=DATABASE, + cluster_identifier=cluster_identifier, + db_user=db_user, + secret_arn=secret_arn, + statement_name=statement_name, + parameters=parameters, + session_keep_alive_seconds=123, + ) + assert output.statement_id == STATEMENT_ID + assert output.session_id == "session_id" + + mock_conn.execute_statement.assert_called_once_with( + Database=DATABASE, + Sql=SQL, + ClusterIdentifier=cluster_identifier, + DbUser=db_user, + SecretArn=secret_arn, + StatementName=statement_name, + Parameters=parameters, + WithEvent=False, + SessionKeepAliveSeconds=123, + ) + mock_conn.describe_statement.assert_called_once_with( + Id=STATEMENT_ID, + ) + + @mock.patch("airflow.providers.amazon.aws.hooks.redshift_data.RedshiftDataHook.conn") + def test_execute_reuse_session(self, mock_conn): + statement_name = "statement_name" + parameters = [{"name": "id", "value": "1"}] + mock_conn.execute_statement.return_value = {"Id": STATEMENT_ID, "SessionId": "session_id"} + mock_conn.describe_statement.return_value = {"Status": "FINISHED"} + hook = RedshiftDataHook() + session_id = str(uuid4()) + output = hook.execute_query( + database=None, + sql=SQL, + statement_name=statement_name, + parameters=parameters, + session_id=session_id, + ) + assert output.statement_id == STATEMENT_ID + assert output.session_id == "session_id" + + mock_conn.execute_statement.assert_called_once_with( + Sql=SQL, + StatementName=statement_name, + Parameters=parameters, + WithEvent=False, + SessionId=session_id, + ) + mock_conn.describe_statement.assert_called_once_with( + Id=STATEMENT_ID, + ) + @mock.patch("airflow.providers.amazon.aws.hooks.redshift_data.RedshiftDataHook.conn") def test_batch_execute(self, mock_conn): mock_conn.execute_statement.return_value = {"Id": STATEMENT_ID} diff --git a/tests/providers/amazon/aws/operators/test_redshift_data.py b/tests/providers/amazon/aws/operators/test_redshift_data.py index abfa2b038b98b..c22d776a94b44 100644 --- a/tests/providers/amazon/aws/operators/test_redshift_data.py +++ b/tests/providers/amazon/aws/operators/test_redshift_data.py @@ -22,6 +22,7 @@ import pytest from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning, TaskDeferred +from airflow.providers.amazon.aws.hooks.redshift_data import QueryExecutionOutput from airflow.providers.amazon.aws.operators.redshift_data import RedshiftDataOperator from airflow.providers.amazon.aws.triggers.redshift_data import RedshiftDataTrigger from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields @@ -31,6 +32,7 @@ SQL = "sql" DATABASE = "database" STATEMENT_ID = "statement_id" +SESSION_ID = "session_id" @pytest.fixture @@ -98,6 +100,8 @@ def test_execute(self, mock_exec_query): poll_interval = 5 wait_for_completion = True + mock_exec_query.return_value = QueryExecutionOutput(statement_id=STATEMENT_ID, session_id=None) + operator = RedshiftDataOperator( aws_conn_id=CONN_ID, task_id=TASK_ID, @@ -111,7 +115,8 @@ def test_execute(self, mock_exec_query): wait_for_completion=True, poll_interval=poll_interval, ) - operator.execute(None) + mock_ti = mock.MagicMock(name="MockedTaskInstance") + operator.execute({"ti": mock_ti}) mock_exec_query.assert_called_once_with( sql=SQL, database=DATABASE, @@ -124,8 +129,12 @@ def test_execute(self, mock_exec_query): with_event=False, wait_for_completion=wait_for_completion, poll_interval=poll_interval, + session_id=None, + session_keep_alive_seconds=None, ) + mock_ti.xcom_push.assert_not_called() + @mock.patch("airflow.providers.amazon.aws.hooks.redshift_data.RedshiftDataHook.execute_query") def test_execute_with_workgroup_name(self, mock_exec_query): cluster_identifier = None @@ -150,7 +159,54 @@ def test_execute_with_workgroup_name(self, mock_exec_query): wait_for_completion=True, poll_interval=poll_interval, ) - operator.execute(None) + mock_ti = mock.MagicMock(name="MockedTaskInstance") + operator.execute({"ti": mock_ti}) + mock_exec_query.assert_called_once_with( + sql=SQL, + database=DATABASE, + cluster_identifier=cluster_identifier, + workgroup_name=workgroup_name, + db_user=db_user, + secret_arn=secret_arn, + statement_name=statement_name, + parameters=parameters, + with_event=False, + wait_for_completion=wait_for_completion, + poll_interval=poll_interval, + session_id=None, + session_keep_alive_seconds=None, + ) + + @mock.patch("airflow.providers.amazon.aws.hooks.redshift_data.RedshiftDataHook.execute_query") + def test_execute_new_session(self, mock_exec_query): + cluster_identifier = "cluster_identifier" + workgroup_name = None + db_user = "db_user" + secret_arn = "secret_arn" + statement_name = "statement_name" + parameters = [{"name": "id", "value": "1"}] + poll_interval = 5 + wait_for_completion = True + + mock_exec_query.return_value = QueryExecutionOutput(statement_id=STATEMENT_ID, session_id=SESSION_ID) + + operator = RedshiftDataOperator( + aws_conn_id=CONN_ID, + task_id=TASK_ID, + sql=SQL, + database=DATABASE, + cluster_identifier=cluster_identifier, + db_user=db_user, + secret_arn=secret_arn, + statement_name=statement_name, + parameters=parameters, + wait_for_completion=True, + poll_interval=poll_interval, + session_keep_alive_seconds=123, + ) + + mock_ti = mock.MagicMock(name="MockedTaskInstance") + operator.execute({"ti": mock_ti}) mock_exec_query.assert_called_once_with( sql=SQL, database=DATABASE, @@ -163,7 +219,11 @@ def test_execute_with_workgroup_name(self, mock_exec_query): with_event=False, wait_for_completion=wait_for_completion, poll_interval=poll_interval, + session_id=None, + session_keep_alive_seconds=123, ) + assert mock_ti.xcom_push.call_args.kwargs["key"] == "session_id" + assert mock_ti.xcom_push.call_args.kwargs["value"] == SESSION_ID @mock.patch("airflow.providers.amazon.aws.hooks.redshift_data.RedshiftDataHook.conn") def test_on_kill_without_query(self, mock_conn): @@ -180,7 +240,7 @@ def test_on_kill_without_query(self, mock_conn): @mock.patch("airflow.providers.amazon.aws.hooks.redshift_data.RedshiftDataHook.conn") def test_on_kill_with_query(self, mock_conn): - mock_conn.execute_statement.return_value = {"Id": STATEMENT_ID} + mock_conn.execute_statement.return_value = {"Id": STATEMENT_ID, "SessionId": SESSION_ID} operator = RedshiftDataOperator( aws_conn_id=CONN_ID, task_id=TASK_ID, @@ -189,7 +249,8 @@ def test_on_kill_with_query(self, mock_conn): database=DATABASE, wait_for_completion=False, ) - operator.execute(None) + mock_ti = mock.MagicMock(name="MockedTaskInstance") + operator.execute({"ti": mock_ti}) operator.on_kill() mock_conn.cancel_statement.assert_called_once_with( Id=STATEMENT_ID, @@ -198,7 +259,7 @@ def test_on_kill_with_query(self, mock_conn): @mock.patch("airflow.providers.amazon.aws.hooks.redshift_data.RedshiftDataHook.conn") def test_return_sql_result(self, mock_conn): expected_result = {"Result": True} - mock_conn.execute_statement.return_value = {"Id": STATEMENT_ID} + mock_conn.execute_statement.return_value = {"Id": STATEMENT_ID, "SessionId": SESSION_ID} mock_conn.describe_statement.return_value = {"Status": "FINISHED"} mock_conn.get_statement_result.return_value = expected_result cluster_identifier = "cluster_identifier" @@ -216,7 +277,8 @@ def test_return_sql_result(self, mock_conn): aws_conn_id=CONN_ID, return_sql_result=True, ) - actual_result = operator.execute(None) + mock_ti = mock.MagicMock(name="MockedTaskInstance") + actual_result = operator.execute({"ti": mock_ti}) assert actual_result == expected_result mock_conn.execute_statement.assert_called_once_with( Database=DATABASE, @@ -260,7 +322,9 @@ def test_execute_finished_before_defer(self, mock_exec_query, check_query_is_fin poll_interval=poll_interval, deferrable=True, ) - operator.execute(None) + + mock_ti = mock.MagicMock(name="MockedTaskInstance") + operator.execute({"ti": mock_ti}) assert not mock_defer.called mock_exec_query.assert_called_once_with( @@ -275,6 +339,8 @@ def test_execute_finished_before_defer(self, mock_exec_query, check_query_is_fin with_event=False, wait_for_completion=False, poll_interval=poll_interval, + session_id=None, + session_keep_alive_seconds=None, ) @mock.patch( @@ -283,8 +349,9 @@ def test_execute_finished_before_defer(self, mock_exec_query, check_query_is_fin ) @mock.patch("airflow.providers.amazon.aws.hooks.redshift_data.RedshiftDataHook.execute_query") def test_execute_defer(self, mock_exec_query, check_query_is_finished, deferrable_operator): + mock_ti = mock.MagicMock(name="MockedTaskInstance") with pytest.raises(TaskDeferred) as exc: - deferrable_operator.execute(None) + deferrable_operator.execute({"ti": mock_ti}) assert isinstance(exc.value.trigger, RedshiftDataTrigger) @@ -346,7 +413,8 @@ def test_no_wait_for_completion(self, mock_exec_query, mock_check_query_is_finis poll_interval=poll_interval, deferrable=deferrable, ) - operator.execute(None) + mock_ti = mock.MagicMock(name="MockedTaskInstance") + operator.execute({"ti": mock_ti}) assert not mock_check_query_is_finished.called assert not mock_defer.called diff --git a/tests/providers/amazon/aws/utils/test_openlineage.py b/tests/providers/amazon/aws/utils/test_openlineage.py index b3e820b58185e..195db068d3092 100644 --- a/tests/providers/amazon/aws/utils/test_openlineage.py +++ b/tests/providers/amazon/aws/utils/test_openlineage.py @@ -21,7 +21,7 @@ import pytest -from airflow.providers.amazon.aws.hooks.redshift_data import RedshiftDataHook +from airflow.providers.amazon.aws.hooks.redshift_data import QueryExecutionOutput, RedshiftDataHook from airflow.providers.amazon.aws.hooks.redshift_sql import RedshiftSQLHook from airflow.providers.amazon.aws.utils.openlineage import ( get_facets_from_redshift_table, @@ -58,7 +58,7 @@ def test_get_facets_from_redshift_table_sql_hook(mock_get_records): @mock.patch("airflow.providers.amazon.aws.hooks.redshift_data.RedshiftDataHook.execute_query") @mock.patch("airflow.providers.amazon.aws.hooks.redshift_data.RedshiftDataHook.conn") def test_get_facets_from_redshift_table_data_hook(mock_connection, mock_execute_query): - mock_execute_query.return_value = "statement_id" + mock_execute_query.return_value = QueryExecutionOutput(statement_id="statement_id", session_id=None) mock_connection.get_statement_result.return_value = { "Records": [ [ diff --git a/tests/system/providers/amazon/aws/example_redshift.py b/tests/system/providers/amazon/aws/example_redshift.py index cc92076dcba0a..67b822d41ef55 100644 --- a/tests/system/providers/amazon/aws/example_redshift.py +++ b/tests/system/providers/amazon/aws/example_redshift.py @@ -50,7 +50,6 @@ DB_NAME = "dev" POLL_INTERVAL = 10 - with DAG( dag_id=DAG_ID, start_date=datetime(2021, 1, 1), @@ -175,6 +174,37 @@ wait_for_completion=True, ) + # [START howto_operator_redshift_data_session_reuse] + create_tmp_table_data_api = RedshiftDataOperator( + task_id="create_tmp_table_data_api", + cluster_identifier=redshift_cluster_identifier, + database=DB_NAME, + db_user=DB_LOGIN, + sql=""" + CREATE TEMPORARY TABLE tmp_people ( + id INTEGER, + first_name VARCHAR(100), + age INTEGER + ); + """, + poll_interval=POLL_INTERVAL, + wait_for_completion=True, + session_keep_alive_seconds=600, + ) + + insert_data_reuse_session = RedshiftDataOperator( + task_id="insert_data_reuse_session", + sql=""" + INSERT INTO tmp_people VALUES ( 1, 'Bob', 30); + INSERT INTO tmp_people VALUES ( 2, 'Alice', 35); + INSERT INTO tmp_people VALUES ( 3, 'Charlie', 40); + """, + poll_interval=POLL_INTERVAL, + wait_for_completion=True, + session_id="{{ task_instance.xcom_pull(task_ids='create_tmp_table_data_api', key='session_id') }}", + ) + # [END howto_operator_redshift_data_session_reuse] + # [START howto_operator_redshift_delete_cluster] delete_cluster = RedshiftDeleteClusterOperator( task_id="delete_cluster", @@ -209,13 +239,20 @@ delete_cluster, ) + # Test session reuse in parallel + chain( + wait_cluster_available_after_resume, + create_tmp_table_data_api, + insert_data_reuse_session, + delete_cluster_snapshot, + ) + from tests.system.utils.watcher import watcher # This test needs watcher in order to properly mark success/failure # when "tearDown" task with trigger rule is part of the DAG list(dag.tasks) >> watcher() - from tests.system.utils import get_test_run # noqa: E402 # Needed to run the example DAG with pytest (see: tests/system/README.md#run_via_pytest) diff --git a/tests/system/providers/amazon/aws/example_redshift_s3_transfers.py b/tests/system/providers/amazon/aws/example_redshift_s3_transfers.py index 0691046190507..9fb989ec53697 100644 --- a/tests/system/providers/amazon/aws/example_redshift_s3_transfers.py +++ b/tests/system/providers/amazon/aws/example_redshift_s3_transfers.py @@ -53,22 +53,34 @@ S3_KEY = "s3_output_" S3_KEY_2 = "s3_key_2" +S3_KEY_3 = "s3_output_tmp_table_" S3_KEY_PREFIX = "s3_k" REDSHIFT_TABLE = "test_table" +REDSHIFT_TMP_TABLE = "tmp_table" -SQL_CREATE_TABLE = f""" - CREATE TABLE IF NOT EXISTS {REDSHIFT_TABLE} ( - fruit_id INTEGER, - name VARCHAR NOT NULL, - color VARCHAR NOT NULL - ); -""" +DATA = "0, 'Airflow', 'testing'" -SQL_INSERT_DATA = f"INSERT INTO {REDSHIFT_TABLE} VALUES ( 1, 'Banana', 'Yellow');" -SQL_DROP_TABLE = f"DROP TABLE IF EXISTS {REDSHIFT_TABLE};" +def _drop_table(table_name: str) -> str: + return f"DROP TABLE IF EXISTS {table_name};" -DATA = "0, 'Airflow', 'testing'" + +def _create_table(table_name: str, is_temp: bool = False) -> str: + temp_keyword = "TEMPORARY" if is_temp else "" + return ( + _drop_table(table_name) + + f""" + CREATE {temp_keyword} TABLE {table_name} ( + fruit_id INTEGER, + name VARCHAR NOT NULL, + color VARCHAR NOT NULL + ); + """ + ) + + +def _insert_data(table_name: str) -> str: + return f"INSERT INTO {table_name} VALUES ( 1, 'Banana', 'Yellow');" with DAG( @@ -124,7 +136,7 @@ cluster_identifier=redshift_cluster_identifier, database=DB_NAME, db_user=DB_LOGIN, - sql=SQL_CREATE_TABLE, + sql=_create_table(REDSHIFT_TABLE), wait_for_completion=True, ) @@ -133,7 +145,7 @@ cluster_identifier=redshift_cluster_identifier, database=DB_NAME, db_user=DB_LOGIN, - sql=SQL_INSERT_DATA, + sql=_insert_data(REDSHIFT_TABLE), wait_for_completion=True, ) @@ -159,6 +171,33 @@ bucket_key=f"{S3_KEY}/{REDSHIFT_TABLE}_0000_part_00", ) + create_tmp_table = RedshiftDataOperator( + task_id="create_tmp_table", + cluster_identifier=redshift_cluster_identifier, + database=DB_NAME, + db_user=DB_LOGIN, + sql=_create_table(REDSHIFT_TMP_TABLE, is_temp=True) + _insert_data(REDSHIFT_TMP_TABLE), + wait_for_completion=True, + session_keep_alive_seconds=600, + ) + + transfer_redshift_to_s3_reuse_session = RedshiftToS3Operator( + task_id="transfer_redshift_to_s3_reuse_session", + redshift_data_api_kwargs={ + "wait_for_completion": True, + "session_id": "{{ task_instance.xcom_pull(task_ids='create_tmp_table', key='session_id') }}", + }, + s3_bucket=bucket_name, + s3_key=S3_KEY_3, + table=REDSHIFT_TMP_TABLE, + ) + + check_if_tmp_table_key_exists = S3KeySensor( + task_id="check_if_tmp_table_key_exists", + bucket_name=bucket_name, + bucket_key=f"{S3_KEY_3}/{REDSHIFT_TMP_TABLE}_0000_part_00", + ) + # [START howto_transfer_s3_to_redshift] transfer_s3_to_redshift = S3ToRedshiftOperator( task_id="transfer_s3_to_redshift", @@ -176,6 +215,28 @@ ) # [END howto_transfer_s3_to_redshift] + create_dest_tmp_table = RedshiftDataOperator( + task_id="create_dest_tmp_table", + cluster_identifier=redshift_cluster_identifier, + database=DB_NAME, + db_user=DB_LOGIN, + sql=_create_table(REDSHIFT_TMP_TABLE, is_temp=True), + wait_for_completion=True, + session_keep_alive_seconds=600, + ) + + transfer_s3_to_redshift_tmp_table = S3ToRedshiftOperator( + task_id="transfer_s3_to_redshift_tmp_table", + redshift_data_api_kwargs={ + "session_id": "{{ task_instance.xcom_pull(task_ids='create_dest_tmp_table', key='session_id') }}", + "wait_for_completion": True, + }, + s3_bucket=bucket_name, + s3_key=S3_KEY_2, + table=REDSHIFT_TMP_TABLE, + copy_options=["csv"], + ) + # [START howto_transfer_s3_to_redshift_multiple_keys] transfer_s3_to_redshift_multiple = S3ToRedshiftOperator( task_id="transfer_s3_to_redshift_multiple", @@ -198,7 +259,7 @@ cluster_identifier=redshift_cluster_identifier, database=DB_NAME, db_user=DB_LOGIN, - sql=SQL_DROP_TABLE, + sql=_drop_table(REDSHIFT_TABLE), wait_for_completion=True, trigger_rule=TriggerRule.ALL_DONE, ) @@ -235,13 +296,33 @@ delete_bucket, ) + chain( + # TEST SETUP + wait_cluster_available, + create_tmp_table, + # TEST BODY + transfer_redshift_to_s3_reuse_session, + check_if_tmp_table_key_exists, + # TEST TEARDOWN + delete_cluster, + ) + + chain( + # TEST SETUP + wait_cluster_available, + create_dest_tmp_table, + # TEST BODY + transfer_s3_to_redshift_tmp_table, + # TEST TEARDOWN + delete_cluster, + ) + from tests.system.utils.watcher import watcher # This test needs watcher in order to properly mark success/failure # when "tearDown" task with trigger rule is part of the DAG list(dag.tasks) >> watcher() - from tests.system.utils import get_test_run # noqa: E402 # Needed to run the example DAG with pytest (see: tests/system/README.md#run_via_pytest) From 7af8e46974b52b701852719e0964c4d6e9679375 Mon Sep 17 00:00:00 2001 From: Jens Scheffler <95105677+jscheffl@users.noreply.github.com> Date: Tue, 24 Sep 2024 21:59:35 +0200 Subject: [PATCH 010/802] Fix broken main: generated JS types (#42451) --- airflow/ui/openapi-gen/queries/common.ts | 16 ++++++++++++++- airflow/ui/openapi-gen/queries/prefetch.ts | 10 ++++++++++ airflow/ui/openapi-gen/queries/queries.ts | 20 ++++++++++++++++++- airflow/ui/openapi-gen/queries/suspense.ts | 20 ++++++++++++++++++- .../ui/openapi-gen/requests/services.gen.ts | 4 ++++ airflow/ui/openapi-gen/requests/types.gen.ts | 2 ++ 6 files changed, 69 insertions(+), 3 deletions(-) diff --git a/airflow/ui/openapi-gen/queries/common.ts b/airflow/ui/openapi-gen/queries/common.ts index d942d51c91e9c..143ec83c55627 100644 --- a/airflow/ui/openapi-gen/queries/common.ts +++ b/airflow/ui/openapi-gen/queries/common.ts @@ -35,19 +35,23 @@ export const useDagServiceGetDagsPublicDagsGetKey = "DagServiceGetDagsPublicDagsGet"; export const UseDagServiceGetDagsPublicDagsGetKeyFn = ( { + dagDisplayNamePattern, dagIdPattern, limit, offset, onlyActive, orderBy, + owners, paused, tags, }: { + dagDisplayNamePattern?: string; dagIdPattern?: string; limit?: number; offset?: number; onlyActive?: boolean; orderBy?: string; + owners?: string[]; paused?: boolean; tags?: string[]; } = {}, @@ -55,6 +59,16 @@ export const UseDagServiceGetDagsPublicDagsGetKeyFn = ( ) => [ useDagServiceGetDagsPublicDagsGetKey, ...(queryKey ?? [ - { dagIdPattern, limit, offset, onlyActive, orderBy, paused, tags }, + { + dagDisplayNamePattern, + dagIdPattern, + limit, + offset, + onlyActive, + orderBy, + owners, + paused, + tags, + }, ]), ]; diff --git a/airflow/ui/openapi-gen/queries/prefetch.ts b/airflow/ui/openapi-gen/queries/prefetch.ts index 44b8f373534f5..f8e1bf616d143 100644 --- a/airflow/ui/openapi-gen/queries/prefetch.ts +++ b/airflow/ui/openapi-gen/queries/prefetch.ts @@ -35,7 +35,9 @@ export const prefetchUseDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGet = * @param data.limit * @param data.offset * @param data.tags + * @param data.owners * @param data.dagIdPattern + * @param data.dagDisplayNamePattern * @param data.onlyActive * @param data.paused * @param data.orderBy @@ -45,40 +47,48 @@ export const prefetchUseDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGet = export const prefetchUseDagServiceGetDagsPublicDagsGet = ( queryClient: QueryClient, { + dagDisplayNamePattern, dagIdPattern, limit, offset, onlyActive, orderBy, + owners, paused, tags, }: { + dagDisplayNamePattern?: string; dagIdPattern?: string; limit?: number; offset?: number; onlyActive?: boolean; orderBy?: string; + owners?: string[]; paused?: boolean; tags?: string[]; } = {}, ) => queryClient.prefetchQuery({ queryKey: Common.UseDagServiceGetDagsPublicDagsGetKeyFn({ + dagDisplayNamePattern, dagIdPattern, limit, offset, onlyActive, orderBy, + owners, paused, tags, }), queryFn: () => DagService.getDagsPublicDagsGet({ + dagDisplayNamePattern, dagIdPattern, limit, offset, onlyActive, orderBy, + owners, paused, tags, }), diff --git a/airflow/ui/openapi-gen/queries/queries.ts b/airflow/ui/openapi-gen/queries/queries.ts index 55653b622fa07..9dce528f2a503 100644 --- a/airflow/ui/openapi-gen/queries/queries.ts +++ b/airflow/ui/openapi-gen/queries/queries.ts @@ -43,7 +43,9 @@ export const useDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGet = < * @param data.limit * @param data.offset * @param data.tags + * @param data.owners * @param data.dagIdPattern + * @param data.dagDisplayNamePattern * @param data.onlyActive * @param data.paused * @param data.orderBy @@ -56,19 +58,23 @@ export const useDagServiceGetDagsPublicDagsGet = < TQueryKey extends Array = unknown[], >( { + dagDisplayNamePattern, dagIdPattern, limit, offset, onlyActive, orderBy, + owners, paused, tags, }: { + dagDisplayNamePattern?: string; dagIdPattern?: string; limit?: number; offset?: number; onlyActive?: boolean; orderBy?: string; + owners?: string[]; paused?: boolean; tags?: string[]; } = {}, @@ -77,16 +83,28 @@ export const useDagServiceGetDagsPublicDagsGet = < ) => useQuery({ queryKey: Common.UseDagServiceGetDagsPublicDagsGetKeyFn( - { dagIdPattern, limit, offset, onlyActive, orderBy, paused, tags }, + { + dagDisplayNamePattern, + dagIdPattern, + limit, + offset, + onlyActive, + orderBy, + owners, + paused, + tags, + }, queryKey, ), queryFn: () => DagService.getDagsPublicDagsGet({ + dagDisplayNamePattern, dagIdPattern, limit, offset, onlyActive, orderBy, + owners, paused, tags, }) as TData, diff --git a/airflow/ui/openapi-gen/queries/suspense.ts b/airflow/ui/openapi-gen/queries/suspense.ts index 1e4fb671a1130..bcc95a53e18ff 100644 --- a/airflow/ui/openapi-gen/queries/suspense.ts +++ b/airflow/ui/openapi-gen/queries/suspense.ts @@ -44,7 +44,9 @@ export const useDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGetSuspense = * @param data.limit * @param data.offset * @param data.tags + * @param data.owners * @param data.dagIdPattern + * @param data.dagDisplayNamePattern * @param data.onlyActive * @param data.paused * @param data.orderBy @@ -57,19 +59,23 @@ export const useDagServiceGetDagsPublicDagsGetSuspense = < TQueryKey extends Array = unknown[], >( { + dagDisplayNamePattern, dagIdPattern, limit, offset, onlyActive, orderBy, + owners, paused, tags, }: { + dagDisplayNamePattern?: string; dagIdPattern?: string; limit?: number; offset?: number; onlyActive?: boolean; orderBy?: string; + owners?: string[]; paused?: boolean; tags?: string[]; } = {}, @@ -78,16 +84,28 @@ export const useDagServiceGetDagsPublicDagsGetSuspense = < ) => useSuspenseQuery({ queryKey: Common.UseDagServiceGetDagsPublicDagsGetKeyFn( - { dagIdPattern, limit, offset, onlyActive, orderBy, paused, tags }, + { + dagDisplayNamePattern, + dagIdPattern, + limit, + offset, + onlyActive, + orderBy, + owners, + paused, + tags, + }, queryKey, ), queryFn: () => DagService.getDagsPublicDagsGet({ + dagDisplayNamePattern, dagIdPattern, limit, offset, onlyActive, orderBy, + owners, paused, tags, }) as TData, diff --git a/airflow/ui/openapi-gen/requests/services.gen.ts b/airflow/ui/openapi-gen/requests/services.gen.ts index cf28e39ab109f..e0786e9137156 100644 --- a/airflow/ui/openapi-gen/requests/services.gen.ts +++ b/airflow/ui/openapi-gen/requests/services.gen.ts @@ -41,7 +41,9 @@ export class DagService { * @param data.limit * @param data.offset * @param data.tags + * @param data.owners * @param data.dagIdPattern + * @param data.dagDisplayNamePattern * @param data.onlyActive * @param data.paused * @param data.orderBy @@ -58,7 +60,9 @@ export class DagService { limit: data.limit, offset: data.offset, tags: data.tags, + owners: data.owners, dag_id_pattern: data.dagIdPattern, + dag_display_name_pattern: data.dagDisplayNamePattern, only_active: data.onlyActive, paused: data.paused, order_by: data.orderBy, diff --git a/airflow/ui/openapi-gen/requests/types.gen.ts b/airflow/ui/openapi-gen/requests/types.gen.ts index cb98e2b769ad2..917dca6626c08 100644 --- a/airflow/ui/openapi-gen/requests/types.gen.ts +++ b/airflow/ui/openapi-gen/requests/types.gen.ts @@ -70,11 +70,13 @@ export type NextRunDatasetsUiNextRunDatasetsDagIdGetResponse = { }; export type GetDagsPublicDagsGetData = { + dagDisplayNamePattern?: string | null; dagIdPattern?: string | null; limit?: number; offset?: number; onlyActive?: boolean; orderBy?: string; + owners?: Array; paused?: boolean | null; tags?: Array; }; From bfee9bbaa9c99bea25d952832ad0a31fcdc3be4c Mon Sep 17 00:00:00 2001 From: Jens Scheffler <95105677+jscheffl@users.noreply.github.com> Date: Tue, 24 Sep 2024 22:49:36 +0200 Subject: [PATCH 011/802] AIP-69: Add CLI to Edge Provider (#42050) * Add CLI to Edge Provider * Review feedback --- airflow/providers/edge/cli/__init__.py | 16 + airflow/providers/edge/cli/edge_command.py | 313 ++++++++++++++++++ tests/providers/edge/cli/__init__.py | 17 + tests/providers/edge/cli/test_edge_command.py | 259 +++++++++++++++ .../providers/edge/models/test_edge_worker.py | 29 ++ 5 files changed, 634 insertions(+) create mode 100644 airflow/providers/edge/cli/__init__.py create mode 100644 airflow/providers/edge/cli/edge_command.py create mode 100644 tests/providers/edge/cli/__init__.py create mode 100644 tests/providers/edge/cli/test_edge_command.py diff --git a/airflow/providers/edge/cli/__init__.py b/airflow/providers/edge/cli/__init__.py new file mode 100644 index 0000000000000..13a83393a9124 --- /dev/null +++ b/airflow/providers/edge/cli/__init__.py @@ -0,0 +1,16 @@ +# 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. diff --git a/airflow/providers/edge/cli/edge_command.py b/airflow/providers/edge/cli/edge_command.py new file mode 100644 index 0000000000000..09998ffe80281 --- /dev/null +++ b/airflow/providers/edge/cli/edge_command.py @@ -0,0 +1,313 @@ +# 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 logging +import os +import platform +import signal +import sys +from dataclasses import dataclass +from datetime import datetime +from pathlib import Path +from subprocess import Popen +from time import sleep + +import psutil +from lockfile.pidlockfile import read_pid_from_pidfile, remove_existing_pidfile, write_pid_to_pidfile + +from airflow import __version__ as airflow_version, settings +from airflow.api_internal.internal_api_call import InternalApiConfig +from airflow.cli.cli_config import ARG_PID, ARG_VERBOSE, ActionCommand, Arg +from airflow.configuration import conf +from airflow.exceptions import AirflowException +from airflow.providers.edge import __version__ as edge_provider_version +from airflow.providers.edge.models.edge_job import EdgeJob +from airflow.providers.edge.models.edge_logs import EdgeLogs +from airflow.providers.edge.models.edge_worker import EdgeWorker, EdgeWorkerState +from airflow.utils import cli as cli_utils +from airflow.utils.platform import IS_WINDOWS +from airflow.utils.providers_configuration_loader import providers_configuration_loaded +from airflow.utils.state import TaskInstanceState + +logger = logging.getLogger(__name__) +EDGE_WORKER_PROCESS_NAME = "edge-worker" +EDGE_WORKER_HEADER = "\n".join( + [ + r" ____ __ _ __ __", + r" / __/__/ /__ ____ | | /| / /__ ____/ /_____ ____", + r" / _// _ / _ `/ -_) | |/ |/ / _ \/ __/ '_/ -_) __/", + r"/___/\_,_/\_, /\__/ |__/|__/\___/_/ /_/\_\\__/_/", + r" /___/", + r"", + ] +) + + +@providers_configuration_loaded +def force_use_internal_api_on_edge_worker(): + """ + Ensure that the environment is configured for the internal API without needing to declare it outside. + + This is only required for an Edge worker and must to be done before the Click CLI wrapper is initiated. + That is because the CLI wrapper will attempt to establish a DB connection, which will fail before the + function call can take effect. In an Edge worker, we need to "patch" the environment before starting. + """ + if "airflow" in sys.argv[0] and sys.argv[1:3] == ["edge", "worker"]: + api_url = conf.get("edge", "api_url") + if not api_url: + raise SystemExit("Error: API URL is not configured, please correct configuration.") + logger.info("Starting worker with API endpoint %s", api_url) + # export Edge API to be used for internal API + os.environ["AIRFLOW_ENABLE_AIP_44"] = "True" + os.environ["AIRFLOW__CORE__INTERNAL_API_URL"] = api_url + InternalApiConfig.set_use_internal_api("edge-worker") + # Disable mini-scheduler post task execution and leave next task schedule to core scheduler + os.environ["AIRFLOW__SCHEDULER__SCHEDULE_AFTER_TASK_EXECUTION"] = "False" + + +force_use_internal_api_on_edge_worker() + + +def _hostname() -> str: + if IS_WINDOWS: + return platform.uname().node + else: + return os.uname()[1] + + +def _get_sysinfo() -> dict: + """Produce the sysinfo from worker to post to central site.""" + return { + "airflow_version": airflow_version, + "edge_provider_version": edge_provider_version, + } + + +def _pid_file_path(pid_file: str | None) -> str: + return cli_utils.setup_locations(process=EDGE_WORKER_PROCESS_NAME, pid=pid_file)[0] + + +@dataclass +class _Job: + """Holds all information for a task/job to be executed as bundle.""" + + edge_job: EdgeJob + process: Popen + logfile: Path + logsize: int + """Last size of log file, point of last chunk push.""" + + +class _EdgeWorkerCli: + """Runner instance which executes the Edge Worker.""" + + jobs: list[_Job] = [] + """List of jobs that the worker is running currently.""" + last_hb: datetime | None = None + """Timestamp of last heart beat sent to server.""" + drain: bool = False + """Flag if job processing should be completed and no new jobs fetched for a graceful stop/shutdown.""" + + def __init__( + self, + pid_file_path: Path, + hostname: str, + queues: list[str] | None, + concurrency: int, + job_poll_interval: int, + heartbeat_interval: int, + ): + self.pid_file_path = pid_file_path + self.job_poll_interval = job_poll_interval + self.hb_interval = heartbeat_interval + self.hostname = hostname + self.queues = queues + self.concurrency = concurrency + + @staticmethod + def signal_handler(sig, frame): + logger.info("Request to show down Edge Worker received, waiting for jobs to complete.") + _EdgeWorkerCli.drain = True + + def start(self): + """Start the execution in a loop until terminated.""" + try: + self.last_hb = EdgeWorker.register_worker( + self.hostname, EdgeWorkerState.STARTING, self.queues, _get_sysinfo() + ).last_update + except AirflowException as e: + if "404:NOT FOUND" in str(e): + raise SystemExit("Error: API endpoint is not ready, please set [edge] api_enabled=True.") + raise SystemExit(str(e)) + write_pid_to_pidfile(self.pid_file_path) + signal.signal(signal.SIGINT, _EdgeWorkerCli.signal_handler) + try: + while not _EdgeWorkerCli.drain or self.jobs: + self.loop() + + logger.info("Quitting worker, signal being offline.") + EdgeWorker.set_state(self.hostname, EdgeWorkerState.OFFLINE, 0, _get_sysinfo()) + finally: + remove_existing_pidfile(self.pid_file_path) + + def loop(self): + """Run a loop of scheduling and monitoring tasks.""" + new_job = False + if not _EdgeWorkerCli.drain and len(self.jobs) < self.concurrency: + new_job = self.fetch_job() + self.check_running_jobs() + + if _EdgeWorkerCli.drain or datetime.now().timestamp() - self.last_hb.timestamp() > self.hb_interval: + self.heartbeat() + self.last_hb = datetime.now() + + if not new_job: + self.interruptible_sleep() + + def fetch_job(self) -> bool: + """Fetch and start a new job from central site.""" + logger.debug("Attempting to fetch a new job...") + edge_job = EdgeJob.reserve_task(self.hostname, self.queues) + if edge_job: + logger.info("Received job: %s", edge_job) + env = os.environ.copy() + env["AIRFLOW__CORE__DATABASE_ACCESS_ISOLATION"] = "True" + env["AIRFLOW__CORE__INTERNAL_API_URL"] = conf.get("edge", "api_url") + env["_AIRFLOW__SKIP_DATABASE_EXECUTOR_COMPATIBILITY_CHECK"] = "1" + process = Popen(edge_job.command, close_fds=True, env=env) + logfile = EdgeLogs.logfile_path(edge_job.key) + self.jobs.append(_Job(edge_job, process, logfile, 0)) + EdgeJob.set_state(edge_job.key, TaskInstanceState.RUNNING) + return True + + logger.info("No new job to process%s", f", {len(self.jobs)} still running" if self.jobs else "") + return False + + def check_running_jobs(self) -> None: + """Check which of the running tasks/jobs are completed and report back.""" + for i in range(len(self.jobs) - 1, -1, -1): + job = self.jobs[i] + job.process.poll() + if job.process.returncode is not None: + self.jobs.remove(job) + if job.process.returncode == 0: + logger.info("Job completed: %s", job.edge_job) + EdgeJob.set_state(job.edge_job.key, TaskInstanceState.SUCCESS) + else: + logger.error("Job failed: %s", job.edge_job) + EdgeJob.set_state(job.edge_job.key, TaskInstanceState.FAILED) + if job.logfile.exists() and job.logfile.stat().st_size > job.logsize: + with job.logfile.open("r") as logfile: + logfile.seek(job.logsize, os.SEEK_SET) + logdata = logfile.read() + EdgeLogs.push_logs( + task=job.edge_job.key, + log_chunk_time=datetime.now(), + log_chunk_data=logdata, + ) + job.logsize += len(logdata) + + def heartbeat(self) -> None: + """Report liveness state of worker to central site with stats.""" + state = ( + (EdgeWorkerState.TERMINATING if _EdgeWorkerCli.drain else EdgeWorkerState.RUNNING) + if self.jobs + else EdgeWorkerState.IDLE + ) + sysinfo = _get_sysinfo() + EdgeWorker.set_state(self.hostname, state, len(self.jobs), sysinfo) + + def interruptible_sleep(self): + """Sleeps but stops sleeping if drain is made.""" + drain_before_sleep = _EdgeWorkerCli.drain + for _ in range(0, self.job_poll_interval * 10): + sleep(0.1) + if drain_before_sleep != _EdgeWorkerCli.drain: + return + + +@cli_utils.action_cli(check_db=False) +@providers_configuration_loaded +def worker(args): + """Start Airflow Edge Worker.""" + print(settings.HEADER) + print(EDGE_WORKER_HEADER) + + edge_worker = _EdgeWorkerCli( + pid_file_path=_pid_file_path(args.pid), + hostname=args.edge_hostname or _hostname(), + queues=args.queues.split(",") if args.queues else None, + concurrency=args.concurrency, + job_poll_interval=conf.getint("edge", "job_poll_interval"), + heartbeat_interval=conf.getint("edge", "heartbeat_interval"), + ) + edge_worker.start() + + +@cli_utils.action_cli(check_db=False) +@providers_configuration_loaded +def stop(args): + """Stop a running Airflow Edge Worker.""" + pid = read_pid_from_pidfile(_pid_file_path(args.pid)) + # Send SIGINT + if pid: + logger.warning("Sending SIGINT to worker pid %i.", pid) + worker_process = psutil.Process(pid) + worker_process.send_signal(signal.SIGINT) + else: + logger.warning("Could not find PID of worker.") + + +ARG_CONCURRENCY = Arg( + ("-c", "--concurrency"), + type=int, + help="The number of worker processes", + default=conf.getint("edge", "worker_concurrency", fallback=8), +) +ARG_QUEUES = Arg( + ("-q", "--queues"), + help="Comma delimited list of queues to serve, serve all queues if not provided.", +) +ARG_EDGE_HOSTNAME = Arg( + ("-H", "--edge-hostname"), + help="Set the hostname of worker if you have multiple workers on a single machine", +) +EDGE_COMMANDS: list[ActionCommand] = [ + ActionCommand( + name=worker.__name__, + help=worker.__doc__, + func=worker, + args=( + ARG_CONCURRENCY, + ARG_QUEUES, + ARG_EDGE_HOSTNAME, + ARG_PID, + ARG_VERBOSE, + ), + ), + ActionCommand( + name=stop.__name__, + help=stop.__doc__, + func=stop, + args=( + ARG_PID, + ARG_VERBOSE, + ), + ), +] diff --git a/tests/providers/edge/cli/__init__.py b/tests/providers/edge/cli/__init__.py new file mode 100644 index 0000000000000..217e5db960782 --- /dev/null +++ b/tests/providers/edge/cli/__init__.py @@ -0,0 +1,17 @@ +# +# 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. diff --git a/tests/providers/edge/cli/test_edge_command.py b/tests/providers/edge/cli/test_edge_command.py new file mode 100644 index 0000000000000..398c221db02f9 --- /dev/null +++ b/tests/providers/edge/cli/test_edge_command.py @@ -0,0 +1,259 @@ +# 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 + +from datetime import datetime +from pathlib import Path +from subprocess import Popen +from unittest.mock import patch + +import pytest +import time_machine + +from airflow.exceptions import AirflowException +from airflow.providers.edge.cli.edge_command import ( + _EdgeWorkerCli, + _get_sysinfo, + _Job, +) +from airflow.providers.edge.models.edge_job import EdgeJob +from airflow.providers.edge.models.edge_worker import EdgeWorker, EdgeWorkerState +from airflow.utils.state import TaskInstanceState +from tests.test_utils.config import conf_vars + +pytest.importorskip("pydantic", minversion="2.0.0") + +# Ignore the following error for mocking +# mypy: disable-error-code="attr-defined" + + +def test_get_sysinfo(): + sysinfo = _get_sysinfo() + assert "airflow_version" in sysinfo + assert "edge_provider_version" in sysinfo + + +class TestEdgeWorkerCli: + @pytest.fixture + def dummy_joblist(self, tmp_path: Path) -> list[_Job]: + logfile = tmp_path / "file.log" + logfile.touch() + + class MockPopen(Popen): + generated_returncode = None + + def __init__(self): + pass + + def poll(self): + pass + + @property + def returncode(self): + return self.generated_returncode + + return [ + _Job( + edge_job=EdgeJob( + dag_id="test", + task_id="test1", + run_id="test", + map_index=-1, + try_number=1, + state=TaskInstanceState.RUNNING, + queue="test", + command=["test", "command"], + queued_dttm=datetime.now(), + edge_worker=None, + last_update=None, + ), + process=MockPopen(), + logfile=logfile, + logsize=0, + ), + ] + + @pytest.fixture + def worker_with_job(self, tmp_path: Path, dummy_joblist: list[_Job]) -> _EdgeWorkerCli: + test_worker = _EdgeWorkerCli(tmp_path / "dummy.pid", "dummy", None, 8, 5, 5) + test_worker.jobs = dummy_joblist + return test_worker + + @pytest.mark.parametrize( + "reserve_result, fetch_result, expected_calls", + [ + pytest.param(None, False, (0, 0), id="no_job"), + pytest.param( + EdgeJob( + dag_id="test", + task_id="test", + run_id="test", + map_index=-1, + try_number=1, + state=TaskInstanceState.QUEUED, + queue="test", + command=["test", "command"], + queued_dttm=datetime.now(), + edge_worker=None, + last_update=None, + ), + True, + (1, 1), + id="new_job", + ), + ], + ) + @patch("airflow.providers.edge.models.edge_job.EdgeJob.reserve_task") + @patch("airflow.providers.edge.models.edge_logs.EdgeLogs.logfile_path") + @patch("airflow.providers.edge.models.edge_job.EdgeJob.set_state") + @patch("subprocess.Popen") + def test_fetch_job( + self, + mock_popen, + mock_set_state, + mock_logfile_path, + mock_reserve_task, + reserve_result, + fetch_result, + expected_calls, + worker_with_job: _EdgeWorkerCli, + ): + logfile_path_call_count, set_state_call_count = expected_calls + mock_reserve_task.side_effect = [reserve_result] + mock_popen.side_effect = ["dummy"] + with conf_vars({("edge", "api_url"): "https://mock.server"}): + got_job = worker_with_job.fetch_job() + mock_reserve_task.assert_called_once() + assert got_job == fetch_result + assert mock_logfile_path.call_count == logfile_path_call_count + assert mock_set_state.call_count == set_state_call_count + + def test_check_running_jobs_running(self, worker_with_job: _EdgeWorkerCli): + worker_with_job.jobs[0].process.generated_returncode = None + with conf_vars({("edge", "api_url"): "https://mock.server"}): + worker_with_job.check_running_jobs() + assert len(worker_with_job.jobs) == 1 + + @patch("airflow.providers.edge.models.edge_job.EdgeJob.set_state") + def test_check_running_jobs_success(self, mock_set_state, worker_with_job: _EdgeWorkerCli): + job = worker_with_job.jobs[0] + job.process.generated_returncode = 0 + with conf_vars({("edge", "api_url"): "https://mock.server"}): + worker_with_job.check_running_jobs() + assert len(worker_with_job.jobs) == 0 + mock_set_state.assert_called_once_with(job.edge_job.key, TaskInstanceState.SUCCESS) + + @patch("airflow.providers.edge.models.edge_job.EdgeJob.set_state") + def test_check_running_jobs_failed(self, mock_set_state, worker_with_job: _EdgeWorkerCli): + job = worker_with_job.jobs[0] + job.process.generated_returncode = 42 + with conf_vars({("edge", "api_url"): "https://mock.server"}): + worker_with_job.check_running_jobs() + assert len(worker_with_job.jobs) == 0 + mock_set_state.assert_called_once_with(job.edge_job.key, TaskInstanceState.FAILED) + + @time_machine.travel(datetime.now(), tick=False) + @patch("airflow.providers.edge.models.edge_logs.EdgeLogs.push_logs") + def test_check_running_jobs_log_push(self, mock_push_logs, worker_with_job: _EdgeWorkerCli): + job = worker_with_job.jobs[0] + job.process.generated_returncode = None + job.logfile.write_text("some log content") + with conf_vars({("edge", "api_url"): "https://mock.server"}): + worker_with_job.check_running_jobs() + assert len(worker_with_job.jobs) == 1 + mock_push_logs.assert_called_once_with( + task=job.edge_job.key, log_chunk_time=datetime.now(), log_chunk_data="some log content" + ) + + @time_machine.travel(datetime.now(), tick=False) + @patch("airflow.providers.edge.models.edge_logs.EdgeLogs.push_logs") + def test_check_running_jobs_log_push_increment(self, mock_push_logs, worker_with_job: _EdgeWorkerCli): + job = worker_with_job.jobs[0] + job.process.generated_returncode = None + job.logfile.write_text("hello ") + job.logsize = job.logfile.stat().st_size + job.logfile.write_text("hello world") + with conf_vars({("edge", "api_url"): "https://mock.server"}): + worker_with_job.check_running_jobs() + assert len(worker_with_job.jobs) == 1 + mock_push_logs.assert_called_once_with( + task=job.edge_job.key, log_chunk_time=datetime.now(), log_chunk_data="world" + ) + + @pytest.mark.parametrize( + "drain, jobs, expected_state", + [ + pytest.param(False, True, EdgeWorkerState.RUNNING, id="running_jobs"), + pytest.param(True, True, EdgeWorkerState.TERMINATING, id="shutting_down"), + pytest.param(False, False, EdgeWorkerState.IDLE, id="idle"), + ], + ) + @patch("airflow.providers.edge.models.edge_worker.EdgeWorker.set_state") + def test_heartbeat(self, mock_set_state, drain, jobs, expected_state, worker_with_job: _EdgeWorkerCli): + if not jobs: + worker_with_job.jobs = [] + _EdgeWorkerCli.drain = drain + with conf_vars({("edge", "api_url"): "https://mock.server"}): + worker_with_job.heartbeat() + assert mock_set_state.call_args.args[1] == expected_state + + @patch("airflow.providers.edge.models.edge_worker.EdgeWorker.register_worker") + def test_start_missing_apiserver(self, mock_register_worker, worker_with_job: _EdgeWorkerCli): + mock_register_worker.side_effect = AirflowException( + "Something with 404:NOT FOUND means API is not active" + ) + with pytest.raises(SystemExit, match=r"API endpoint is not ready"): + worker_with_job.start() + + @patch("airflow.providers.edge.models.edge_worker.EdgeWorker.register_worker") + def test_start_server_error(self, mock_register_worker, worker_with_job: _EdgeWorkerCli): + mock_register_worker.side_effect = AirflowException("Something other error not FourhundretFour") + with pytest.raises(SystemExit, match=r"Something other"): + worker_with_job.start() + + @patch("airflow.providers.edge.models.edge_worker.EdgeWorker.register_worker") + @patch("airflow.providers.edge.cli.edge_command._EdgeWorkerCli.loop") + @patch("airflow.providers.edge.models.edge_worker.EdgeWorker.set_state") + def test_start_and_run_one( + self, mock_set_state, mock_loop, mock_register_worker, worker_with_job: _EdgeWorkerCli + ): + mock_register_worker.side_effect = [ + EdgeWorker( + worker_name="test", + state=EdgeWorkerState.STARTING, + queues=None, + first_online=datetime.now(), + last_update=datetime.now(), + jobs_active=0, + jobs_taken=0, + jobs_success=0, + jobs_failed=0, + sysinfo="", + ) + ] + + def stop_running(): + _EdgeWorkerCli.drain = True + worker_with_job.jobs = [] + + mock_loop.side_effect = stop_running + + worker_with_job.start() + + mock_register_worker.assert_called_once() + mock_loop.assert_called_once() + mock_set_state.assert_called_once() diff --git a/tests/providers/edge/models/test_edge_worker.py b/tests/providers/edge/models/test_edge_worker.py index 9eca293bafe3f..f0e0ac9dfa056 100644 --- a/tests/providers/edge/models/test_edge_worker.py +++ b/tests/providers/edge/models/test_edge_worker.py @@ -20,11 +20,14 @@ import pytest +from airflow.providers.edge.cli.edge_command import _get_sysinfo from airflow.providers.edge.models.edge_worker import ( EdgeWorker, EdgeWorkerModel, + EdgeWorkerState, EdgeWorkerVersionException, ) +from airflow.utils import timezone if TYPE_CHECKING: from sqlalchemy.orm import Session @@ -63,3 +66,29 @@ def test_assert_version(self): EdgeWorker.assert_version( {"airflow_version": airflow_version, "edge_provider_version": edge_provider_version} ) + + def test_register_worker(self, session: Session): + EdgeWorker.register_worker( + "test_worker", EdgeWorkerState.STARTING, queues=None, sysinfo=_get_sysinfo() + ) + + worker: list[EdgeWorkerModel] = session.query(EdgeWorkerModel).all() + assert len(worker) == 1 + assert worker[0].worker_name == "test_worker" + + def test_set_state(self, session: Session): + rwm = EdgeWorkerModel( + worker_name="test2_worker", + state=EdgeWorkerState.IDLE, + queues=["default"], + first_online=timezone.utcnow(), + ) + session.add(rwm) + session.commit() + + EdgeWorker.set_state("test2_worker", EdgeWorkerState.RUNNING, 1, _get_sysinfo()) + + worker: list[EdgeWorkerModel] = session.query(EdgeWorkerModel).all() + assert len(worker) == 1 + assert worker[0].worker_name == "test2_worker" + assert worker[0].state == EdgeWorkerState.RUNNING From 39843c704e07374b23253a4b14f5c5267b7c2354 Mon Sep 17 00:00:00 2001 From: "D. Ferruzzi" Date: Tue, 24 Sep 2024 15:07:40 -0700 Subject: [PATCH 012/802] Add STOPPED to the failure cases for Sagemaker Training Jobs (#42423) --- airflow/providers/amazon/aws/hooks/sagemaker.py | 3 ++- airflow/providers/amazon/aws/sensors/sagemaker.py | 2 +- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/airflow/providers/amazon/aws/hooks/sagemaker.py b/airflow/providers/amazon/aws/hooks/sagemaker.py index af131697a5e8d..2c0f4fb25edc5 100644 --- a/airflow/providers/amazon/aws/hooks/sagemaker.py +++ b/airflow/providers/amazon/aws/hooks/sagemaker.py @@ -155,6 +155,7 @@ class SageMakerHook(AwsBaseHook): endpoint_non_terminal_states = {"Creating", "Updating", "SystemUpdating", "RollingBack", "Deleting"} pipeline_non_terminal_states = {"Executing", "Stopping"} failed_states = {"Failed"} + training_failed_states = {*failed_states, "Stopped"} def __init__(self, *args, **kwargs): super().__init__(client_type="sagemaker", *args, **kwargs) @@ -309,7 +310,7 @@ def create_training_job( self.check_training_status_with_log( config["TrainingJobName"], self.non_terminal_states, - self.failed_states, + self.training_failed_states, wait_for_completion, check_interval, max_ingestion_time, diff --git a/airflow/providers/amazon/aws/sensors/sagemaker.py b/airflow/providers/amazon/aws/sensors/sagemaker.py index b01e24cd5b815..af07c504aa29d 100644 --- a/airflow/providers/amazon/aws/sensors/sagemaker.py +++ b/airflow/providers/amazon/aws/sensors/sagemaker.py @@ -238,7 +238,7 @@ def non_terminal_states(self): return SageMakerHook.non_terminal_states def failed_states(self): - return SageMakerHook.failed_states + return SageMakerHook.training_failed_states def get_sagemaker_response(self): if self.print_log: From 7020d501e7e9251519e805fda7a105109621ee78 Mon Sep 17 00:00:00 2001 From: Tzu-ping Chung Date: Tue, 24 Sep 2024 18:50:43 -0700 Subject: [PATCH 013/802] Refactor _register_dataset_changes (#42343) --- airflow/dag_processing/collection.py | 37 ++++----- airflow/datasets/manager.py | 77 +++++++++++++++---- airflow/listeners/spec/dataset.py | 9 ++- airflow/models/dataset.py | 6 ++ airflow/models/taskinstance.py | 51 ++++++------ .../listeners.rst | 1 + newsfragments/42343.feature.rst | 1 + newsfragments/42343.significant.rst | 7 ++ tests/datasets/test_manager.py | 7 +- 9 files changed, 128 insertions(+), 68 deletions(-) create mode 100644 newsfragments/42343.feature.rst create mode 100644 newsfragments/42343.significant.rst diff --git a/airflow/dag_processing/collection.py b/airflow/dag_processing/collection.py index 5d54d17b87072..bcac479d875a3 100644 --- a/airflow/dag_processing/collection.py +++ b/airflow/dag_processing/collection.py @@ -299,21 +299,15 @@ def add_datasets(self, *, session: Session) -> dict[str, DatasetModel]: dm.uri: dm for dm in session.scalars(select(DatasetModel).where(DatasetModel.uri.in_(self.datasets))) } - - def _resolve_dataset_addition() -> Iterator[DatasetModel]: - for uri, dataset in self.datasets.items(): - try: - dm = orm_datasets[uri] - except KeyError: - dm = orm_datasets[uri] = DatasetModel.from_public(dataset) - yield dm - else: - # The orphaned flag was bulk-set to True before parsing, so we - # don't need to handle rows in the db without a public entry. - dm.is_orphaned = expression.false() - dm.extra = dataset.extra - - dataset_manager.create_datasets(list(_resolve_dataset_addition()), session=session) + for model in orm_datasets.values(): + model.is_orphaned = expression.false() + orm_datasets.update( + (model.uri, model) + for model in dataset_manager.create_datasets( + [dataset for uri, dataset in self.datasets.items() if uri not in orm_datasets], + session=session, + ) + ) return orm_datasets def add_dataset_aliases(self, *, session: Session) -> dict[str, DatasetAliasModel]: @@ -326,12 +320,13 @@ def add_dataset_aliases(self, *, session: Session) -> dict[str, DatasetAliasMode select(DatasetAliasModel).where(DatasetAliasModel.name.in_(self.dataset_aliases)) ) } - for name, alias in self.dataset_aliases.items(): - try: - da = orm_aliases[name] - except KeyError: - da = orm_aliases[name] = DatasetAliasModel.from_public(alias) - session.add(da) + orm_aliases.update( + (model.name, model) + for model in dataset_manager.create_dataset_aliases( + [alias for name, alias in self.dataset_aliases.items() if name not in orm_aliases], + session=session, + ) + ) return orm_aliases def add_dag_dataset_references( diff --git a/airflow/datasets/manager.py b/airflow/datasets/manager.py index 19f6913fffbeb..c5ebb2e6d7eff 100644 --- a/airflow/datasets/manager.py +++ b/airflow/datasets/manager.py @@ -17,7 +17,7 @@ # under the License. from __future__ import annotations -from collections.abc import Iterable +from collections.abc import Collection, Iterable from typing import TYPE_CHECKING from sqlalchemy import exc, select @@ -25,7 +25,6 @@ from airflow.api_internal.internal_api_call import internal_api_call from airflow.configuration import conf -from airflow.datasets import Dataset from airflow.listeners.listener import get_listener_manager from airflow.models.dagbag import DagPriorityParsingRequest from airflow.models.dataset import ( @@ -43,6 +42,7 @@ if TYPE_CHECKING: from sqlalchemy.orm.session import Session + from airflow.datasets import Dataset, DatasetAlias from airflow.models.dag import DagModel from airflow.models.taskinstance import TaskInstance @@ -58,12 +58,51 @@ class DatasetManager(LoggingMixin): def __init__(self, **kwargs): super().__init__(**kwargs) - def create_datasets(self, dataset_models: list[DatasetModel], session: Session) -> None: + def create_datasets(self, datasets: list[Dataset], *, session: Session) -> list[DatasetModel]: """Create new datasets.""" - for dataset_model in dataset_models: - session.add(dataset_model) - for dataset_model in dataset_models: - self.notify_dataset_created(dataset=Dataset(uri=dataset_model.uri, extra=dataset_model.extra)) + + def _add_one(dataset: Dataset) -> DatasetModel: + model = DatasetModel.from_public(dataset) + session.add(model) + self.notify_dataset_created(dataset=dataset) + return model + + return [_add_one(d) for d in datasets] + + def create_dataset_aliases( + self, + dataset_aliases: list[DatasetAlias], + *, + session: Session, + ) -> list[DatasetAliasModel]: + """Create new dataset aliases.""" + + def _add_one(dataset_alias: DatasetAlias) -> DatasetAliasModel: + model = DatasetAliasModel.from_public(dataset_alias) + session.add(model) + self.notify_dataset_alias_created(dataset_alias=dataset_alias) + return model + + return [_add_one(a) for a in dataset_aliases] + + @classmethod + def _add_dataset_alias_association( + cls, + alias_names: Collection[str], + dataset: DatasetModel, + *, + session: Session, + ) -> None: + already_related = {m.name for m in dataset.aliases} + existing_aliases = { + m.name: m + for m in session.scalars(select(DatasetAliasModel).where(DatasetAliasModel.name.in_(alias_names))) + } + dataset.aliases.extend( + existing_aliases.get(name, DatasetAliasModel(name=name)) + for name in alias_names + if name not in already_related + ) @classmethod @internal_api_call @@ -74,8 +113,9 @@ def register_dataset_change( task_instance: TaskInstance | None = None, dataset: Dataset, extra=None, - session: Session = NEW_SESSION, + aliases: Collection[DatasetAlias] = (), source_alias_names: Iterable[str] | None = None, + session: Session = NEW_SESSION, **kwargs, ) -> DatasetEvent | None: """ @@ -88,24 +128,27 @@ def register_dataset_change( dataset_model = session.scalar( select(DatasetModel) .where(DatasetModel.uri == dataset.uri) - .options(joinedload(DatasetModel.consuming_dags).joinedload(DagScheduleDatasetReference.dag)) + .options( + joinedload(DatasetModel.aliases), + joinedload(DatasetModel.consuming_dags).joinedload(DagScheduleDatasetReference.dag), + ) ) if not dataset_model: cls.logger().warning("DatasetModel %s not found", dataset) return None + cls._add_dataset_alias_association({alias.name for alias in aliases}, dataset_model, session=session) + event_kwargs = { "dataset_id": dataset_model.id, "extra": extra, } if task_instance: event_kwargs.update( - { - "source_task_id": task_instance.task_id, - "source_dag_id": task_instance.dag_id, - "source_run_id": task_instance.run_id, - "source_map_index": task_instance.map_index, - } + source_task_id=task_instance.task_id, + source_dag_id=task_instance.dag_id, + source_run_id=task_instance.run_id, + source_map_index=task_instance.map_index, ) dataset_event = DatasetEvent(**event_kwargs) @@ -155,6 +198,10 @@ def notify_dataset_created(self, dataset: Dataset): """Run applicable notification actions when a dataset is created.""" get_listener_manager().hook.on_dataset_created(dataset=dataset) + def notify_dataset_alias_created(self, dataset_alias: DatasetAlias): + """Run applicable notification actions when a dataset alias is created.""" + get_listener_manager().hook.on_dataset_alias_created(dataset_alias=dataset_alias) + @classmethod def notify_dataset_changed(cls, dataset: Dataset): """Run applicable notification actions when a dataset is changed.""" diff --git a/airflow/listeners/spec/dataset.py b/airflow/listeners/spec/dataset.py index 214ddad3ffb13..eee1a10dd7d89 100644 --- a/airflow/listeners/spec/dataset.py +++ b/airflow/listeners/spec/dataset.py @@ -22,7 +22,7 @@ from pluggy import HookspecMarker if TYPE_CHECKING: - from airflow.datasets import Dataset + from airflow.datasets import Dataset, DatasetAlias hookspec = HookspecMarker("airflow") @@ -34,6 +34,13 @@ def on_dataset_created( """Execute when a new dataset is created.""" +@hookspec +def on_dataset_alias_created( + dataset_alias: DatasetAlias, +): + """Execute when a new dataset alias is created.""" + + @hookspec def on_dataset_changed( dataset: Dataset, diff --git a/airflow/models/dataset.py b/airflow/models/dataset.py index 5033da48a3059..489d6b68a6f15 100644 --- a/airflow/models/dataset.py +++ b/airflow/models/dataset.py @@ -138,6 +138,9 @@ def __eq__(self, other): else: return NotImplemented + def to_public(self) -> DatasetAlias: + return DatasetAlias(name=self.name) + class DatasetModel(Base): """ @@ -200,6 +203,9 @@ def __hash__(self): def __repr__(self): return f"{self.__class__.__name__}(uri={self.uri!r}, extra={self.extra!r})" + def to_public(self) -> Dataset: + return Dataset(uri=self.uri, extra=self.extra) + class DagScheduleDatasetAliasReference(Base): """References from a DAG to a dataset alias of which it is a consumer.""" diff --git a/airflow/models/taskinstance.py b/airflow/models/taskinstance.py index 954e5ed4d0c80..d3300207abfdc 100644 --- a/airflow/models/taskinstance.py +++ b/airflow/models/taskinstance.py @@ -31,7 +31,7 @@ from contextlib import nullcontext from datetime import timedelta from enum import Enum -from typing import TYPE_CHECKING, Any, Callable, Collection, Generator, Iterable, Mapping, Tuple +from typing import TYPE_CHECKING, Any, Callable, Collection, Dict, Generator, Iterable, Mapping, Tuple from urllib.parse import quote import dill @@ -89,7 +89,7 @@ from airflow.listeners.listener import get_listener_manager from airflow.models.base import Base, StringID, TaskInstanceDependencies, _sentinel from airflow.models.dagbag import DagBag -from airflow.models.dataset import DatasetAliasModel, DatasetModel +from airflow.models.dataset import DatasetModel from airflow.models.log import Log from airflow.models.param import process_params from airflow.models.renderedtifields import get_serialized_template_fields @@ -2893,7 +2893,7 @@ def _register_dataset_changes(self, *, events: OutletEventAccessors, session: Se # One task only triggers one dataset event for each dataset with the same extra. # This tuple[dataset uri, extra] to sets alias names mapping is used to find whether # there're datasets with same uri but different extra that we need to emit more than one dataset events. - dataset_tuple_to_alias_names_mapping: dict[tuple[str, frozenset], set[str]] = defaultdict(set) + dataset_alias_names: dict[tuple[str, frozenset], set[str]] = defaultdict(set) for obj in self.task.outlets or []: self.log.debug("outlet obj %s", obj) # Lineage can have other types of objects besides datasets @@ -2908,33 +2908,27 @@ def _register_dataset_changes(self, *, events: OutletEventAccessors, session: Se for dataset_alias_event in events[obj].dataset_alias_events: dataset_alias_name = dataset_alias_event["source_alias_name"] dataset_uri = dataset_alias_event["dest_dataset_uri"] - extra = dataset_alias_event["extra"] - frozen_extra = frozenset(extra.items()) + frozen_extra = frozenset(dataset_alias_event["extra"].items()) + dataset_alias_names[(dataset_uri, frozen_extra)].add(dataset_alias_name) - dataset_tuple_to_alias_names_mapping[(dataset_uri, frozen_extra)].add(dataset_alias_name) + class _DatasetModelCache(Dict[str, DatasetModel]): + log = self.log - dataset_objs_cache: dict[str, DatasetModel] = {} - for (uri, extra_items), alias_names in dataset_tuple_to_alias_names_mapping.items(): - if uri not in dataset_objs_cache: - dataset_obj = session.scalar(select(DatasetModel).where(DatasetModel.uri == uri).limit(1)) - dataset_objs_cache[uri] = dataset_obj - else: - dataset_obj = dataset_objs_cache[uri] - - if not dataset_obj: - dataset_obj = DatasetModel(uri=uri) - dataset_manager.create_datasets(dataset_models=[dataset_obj], session=session) - self.log.warning("Created a new %r as it did not exist.", dataset_obj) + def __missing__(self, key: str) -> DatasetModel: + (dataset_obj,) = dataset_manager.create_datasets([Dataset(uri=key)], session=session) session.flush() - dataset_objs_cache[uri] = dataset_obj - - for alias in alias_names: - alias_obj = session.scalar( - select(DatasetAliasModel).where(DatasetAliasModel.name == alias).limit(1) - ) - dataset_obj.aliases.append(alias_obj) + self.log.warning("Created a new %r as it did not exist.", dataset_obj) + self[key] = dataset_obj + return dataset_obj - extra = {k: v for k, v in extra_items} + dataset_objs_cache = _DatasetModelCache( + (dataset_obj.uri, dataset_obj) + for dataset_obj in session.scalars( + select(DatasetModel).where(DatasetModel.uri.in_(uri for uri, _ in dataset_alias_names)) + ) + ) + for (uri, extra_items), alias_names in dataset_alias_names.items(): + dataset_obj = dataset_objs_cache[uri] self.log.info( 'Creating event for %r through aliases "%s"', dataset_obj, @@ -2942,8 +2936,9 @@ def _register_dataset_changes(self, *, events: OutletEventAccessors, session: Se ) dataset_manager.register_dataset_change( task_instance=self, - dataset=dataset_obj, - extra=extra, + dataset=dataset_obj.to_public(), + aliases=[DatasetAlias(name) for name in alias_names], + extra=dict(extra_items), session=session, source_alias_names=alias_names, ) diff --git a/docs/apache-airflow/administration-and-deployment/listeners.rst b/docs/apache-airflow/administration-and-deployment/listeners.rst index 34909e225aaa9..4926b12ed6c6d 100644 --- a/docs/apache-airflow/administration-and-deployment/listeners.rst +++ b/docs/apache-airflow/administration-and-deployment/listeners.rst @@ -95,6 +95,7 @@ Dataset Events -------------- - ``on_dataset_created`` +- ``on_dataset_alias_created`` - ``on_dataset_changed`` Dataset events occur when Dataset management operations are run. diff --git a/newsfragments/42343.feature.rst b/newsfragments/42343.feature.rst new file mode 100644 index 0000000000000..8a7cdf335a06e --- /dev/null +++ b/newsfragments/42343.feature.rst @@ -0,0 +1 @@ +New function ``create_dataset_aliases`` added to DatasetManager for DatasetAlias creation. diff --git a/newsfragments/42343.significant.rst b/newsfragments/42343.significant.rst new file mode 100644 index 0000000000000..d9e1ba6b1229b --- /dev/null +++ b/newsfragments/42343.significant.rst @@ -0,0 +1,7 @@ +``DatasetManager.create_datasets`` now takes ``Dataset`` objects + +This function previously accepts a list of ``DatasetModel`` objects. it now +receives ``Dataset`` objects instead. A list of ``DatasetModel`` objects are +created inside, and returned by the function. + +Also, the ``session`` argument is now keyword-only. diff --git a/tests/datasets/test_manager.py b/tests/datasets/test_manager.py index 1e7b4fda40cee..d3013aef60c29 100644 --- a/tests/datasets/test_manager.py +++ b/tests/datasets/test_manager.py @@ -169,10 +169,11 @@ def test_create_datasets_notifies_dataset_listener(self, session): dataset_listener.clear() get_listener_manager().add_listener(dataset_listener) - dsm = DatasetModel(uri="test_dataset_uri_3") + ds = Dataset(uri="test_dataset_uri_3") - dsem.create_datasets([dsm], session) + dsms = dsem.create_datasets([ds], session=session) # Ensure the listener was notified assert len(dataset_listener.created) == 1 - assert dataset_listener.created[0].uri == dsm.uri + assert len(dsms) == 1 + assert dataset_listener.created[0].uri == ds.uri == dsms[0].uri From b74f7264b9dbf0bbe4c8e46075e907d8a28bf4bc Mon Sep 17 00:00:00 2001 From: pgvishnuram <81585115+pgvishnuram@users.noreply.github.com> Date: Wed, 25 Sep 2024 10:32:45 +0530 Subject: [PATCH 014/802] add env support for migratedatabase job (#42345) * add env support for migratedatabase job * add test case for env config * update schema json for migrateDatabaseJob * fix ci failures * fix pre-commit ci for json schema --- chart/templates/jobs/migrate-database-job.yaml | 3 +++ chart/values.schema.json | 10 ++++++++++ chart/values.yaml | 1 + helm_tests/airflow_aux/test_migrate_database_job.py | 12 ++++++++++++ 4 files changed, 26 insertions(+) diff --git a/chart/templates/jobs/migrate-database-job.yaml b/chart/templates/jobs/migrate-database-job.yaml index d7747970b88d9..297253e871335 100644 --- a/chart/templates/jobs/migrate-database-job.yaml +++ b/chart/templates/jobs/migrate-database-job.yaml @@ -117,6 +117,9 @@ spec: - name: PYTHONUNBUFFERED value: "1" {{- include "standard_airflow_environment" . | indent 10 }} + {{- if .Values.migrateDatabaseJob.env }} + {{- tpl (toYaml .Values.migrateDatabaseJob.env) $ | nindent 12 }} + {{- end }} resources: {{- toYaml .Values.migrateDatabaseJob.resources | nindent 12 }} volumeMounts: {{- include "airflow_config_mount" . | nindent 12 }} diff --git a/chart/values.schema.json b/chart/values.schema.json index 22679d764a654..948f09f3b9a4d 100644 --- a/chart/values.schema.json +++ b/chart/values.schema.json @@ -4651,6 +4651,16 @@ "null" ], "default": 300 + }, + "env": { + "description": "Add additional env vars to migrate database job.", + "items": { + "$ref": "#/definitions/io.k8s.api.core.v1.EnvVar" + }, + "type": "array", + "default": [], + "x-kubernetes-patch-merge-key": "name", + "x-kubernetes-patch-strategy": "merge" } } }, diff --git a/chart/values.yaml b/chart/values.yaml index ff8d726415936..7bfa733a905b4 100644 --- a/chart/values.yaml +++ b/chart/values.yaml @@ -1236,6 +1236,7 @@ migrateDatabaseJob: # Disable this if you are using ArgoCD for example useHelmHooks: true applyCustomEnv: true + env: [] # rpcServer support is experimental / dev purpose only and will later be renamed _rpcServer: diff --git a/helm_tests/airflow_aux/test_migrate_database_job.py b/helm_tests/airflow_aux/test_migrate_database_job.py index 56ac1d1cd50ab..426a35edc424e 100644 --- a/helm_tests/airflow_aux/test_migrate_database_job.py +++ b/helm_tests/airflow_aux/test_migrate_database_job.py @@ -455,3 +455,15 @@ def test_overridden_automount_service_account_token(self): show_only=["templates/jobs/migrate-database-job-serviceaccount.yaml"], ) assert jmespath.search("automountServiceAccountToken", docs[0]) is False + + def test_should_add_component_specific_env(self): + env = {"name": "test_env_key", "value": "test_env_value"} + docs = render_chart( + values={ + "migrateDatabaseJob": { + "env": [env], + }, + }, + show_only=["templates/jobs/migrate-database-job.yaml"], + ) + assert env in jmespath.search("spec.template.spec.containers[0].env", docs[0]) From bc8f79830e098f026a14bf90b7ab3582dbbf1f66 Mon Sep 17 00:00:00 2001 From: Andor Markus <51825189+andormarkus@users.noreply.github.com> Date: Wed, 25 Sep 2024 12:39:06 +0200 Subject: [PATCH 015/802] fix: Fixing Helm chart flower ingress service reference (#41179) * fix: Fixing Helm chart flower ingress service reference * fix: Fixing Helm chart flower ingress service reference * feat: Add helm unit test for the ingress backend service name * feat: Add helm unit test for the ingress backend service name * fix: Run linter for new files --------- Co-authored-by: Andor Markus (AllCloud) --- chart/templates/flower/flower-ingress.yaml | 5 +++-- helm_tests/webserver/test_ingress_flower.py | 25 +++++++++++++++++++++ helm_tests/webserver/test_ingress_web.py | 24 ++++++++++++++++++++ 3 files changed, 52 insertions(+), 2 deletions(-) diff --git a/chart/templates/flower/flower-ingress.yaml b/chart/templates/flower/flower-ingress.yaml index 7c798ad9fbcfb..1b24d82588069 100644 --- a/chart/templates/flower/flower-ingress.yaml +++ b/chart/templates/flower/flower-ingress.yaml @@ -22,10 +22,11 @@ ################################# {{- if .Values.flower.enabled }} {{- if and (or .Values.ingress.flower.enabled .Values.ingress.enabled) (or (eq .Values.executor "CeleryExecutor") (eq .Values.executor "CeleryKubernetesExecutor")) }} +{{- $fullname := (include "airflow.fullname" .) }} apiVersion: networking.k8s.io/v1 kind: Ingress metadata: - name: {{ include "airflow.fullname" . }}-flower-ingress + name: {{ $fullname }}-flower-ingress labels: tier: airflow component: flower-ingress @@ -72,7 +73,7 @@ spec: paths: - backend: service: - name: {{ $.Release.Name }}-flower + name: {{ $fullname }}-flower port: name: flower-ui {{- if $.Values.ingress.flower.path }} diff --git a/helm_tests/webserver/test_ingress_flower.py b/helm_tests/webserver/test_ingress_flower.py index e3d9ff171d16d..107bf5b270f9c 100644 --- a/helm_tests/webserver/test_ingress_flower.py +++ b/helm_tests/webserver/test_ingress_flower.py @@ -220,3 +220,28 @@ def test_can_ingress_hosts_be_templated(self): "cc.example.com", "dd.example.com", ] == jmespath.search("spec.rules[*].host", docs[0]) + + def test_backend_service_name(self): + docs = render_chart( + values={"ingress": {"enabled": True}, "flower": {"enabled": True}}, + show_only=["templates/flower/flower-ingress.yaml"], + ) + + assert "release-name-flower" == jmespath.search( + "spec.rules[0].http.paths[0].backend.service.name", docs[0] + ) + + def test_backend_service_name_with_fullname_override(self): + docs = render_chart( + values={ + "fullnameOverride": "test-basic", + "useStandardNaming": True, + "ingress": {"enabled": True}, + "flower": {"enabled": True}, + }, + show_only=["templates/flower/flower-ingress.yaml"], + ) + + assert "test-basic-flower" == jmespath.search( + "spec.rules[0].http.paths[0].backend.service.name", docs[0] + ) diff --git a/helm_tests/webserver/test_ingress_web.py b/helm_tests/webserver/test_ingress_web.py index 798da6c719594..38c258c93b9c4 100644 --- a/helm_tests/webserver/test_ingress_web.py +++ b/helm_tests/webserver/test_ingress_web.py @@ -200,3 +200,27 @@ def test_can_ingress_hosts_be_templated(self): "cc.example.com", "dd.example.com", ] == jmespath.search("spec.rules[*].host", docs[0]) + + def test_backend_service_name(self): + docs = render_chart( + values={"ingress": {"web": {"enabled": True}}}, + show_only=["templates/webserver/webserver-ingress.yaml"], + ) + + assert "release-name-webserver" == jmespath.search( + "spec.rules[0].http.paths[0].backend.service.name", docs[0] + ) + + def test_backend_service_name_with_fullname_override(self): + docs = render_chart( + values={ + "fullnameOverride": "test-basic", + "useStandardNaming": True, + "ingress": {"web": {"enabled": True}}, + }, + show_only=["templates/webserver/webserver-ingress.yaml"], + ) + + assert "test-basic-webserver" == jmespath.search( + "spec.rules[0].http.paths[0].backend.service.name", docs[0] + ) From 2a07514c108157199aa12a06ebbbd67531ebbae6 Mon Sep 17 00:00:00 2001 From: Tzu-ping Chung Date: Wed, 25 Sep 2024 04:47:37 -0700 Subject: [PATCH 016/802] Flush less in dataset manager (#42458) --- .../endpoints/dataset_endpoint.py | 1 + airflow/datasets/manager.py | 30 +++++++++---------- airflow/models/taskinstance.py | 28 ++++++++--------- tests/datasets/test_manager.py | 11 ++++--- tests/jobs/test_scheduler_job.py | 2 ++ tests/models/test_taskinstance.py | 4 ++- 6 files changed, 40 insertions(+), 36 deletions(-) diff --git a/airflow/api_connexion/endpoints/dataset_endpoint.py b/airflow/api_connexion/endpoints/dataset_endpoint.py index bfdb8d0a5e7ee..1a1578266838c 100644 --- a/airflow/api_connexion/endpoints/dataset_endpoint.py +++ b/airflow/api_connexion/endpoints/dataset_endpoint.py @@ -352,5 +352,6 @@ def create_dataset_event(session: Session = NEW_SESSION) -> APIResponse: ) if not dataset_event: raise NotFound(title="Dataset not found", detail=f"Dataset with uri: '{uri}' not found") + session.flush() # So we can dump the timestamp. event = dataset_event_schema.dump(dataset_event) return event diff --git a/airflow/datasets/manager.py b/airflow/datasets/manager.py index c5ebb2e6d7eff..6322414bb8499 100644 --- a/airflow/datasets/manager.py +++ b/airflow/datasets/manager.py @@ -37,7 +37,6 @@ ) from airflow.stats import Stats from airflow.utils.log.logging_mixin import LoggingMixin -from airflow.utils.session import NEW_SESSION, provide_session if TYPE_CHECKING: from sqlalchemy.orm.session import Session @@ -55,22 +54,21 @@ class DatasetManager(LoggingMixin): Airflow deployments can use plugins that broadcast dataset events to each other. """ - def __init__(self, **kwargs): - super().__init__(**kwargs) - - def create_datasets(self, datasets: list[Dataset], *, session: Session) -> list[DatasetModel]: + @classmethod + def create_datasets(cls, datasets: list[Dataset], *, session: Session) -> list[DatasetModel]: """Create new datasets.""" def _add_one(dataset: Dataset) -> DatasetModel: model = DatasetModel.from_public(dataset) session.add(model) - self.notify_dataset_created(dataset=dataset) + cls.notify_dataset_created(dataset=dataset) return model return [_add_one(d) for d in datasets] + @classmethod def create_dataset_aliases( - self, + cls, dataset_aliases: list[DatasetAlias], *, session: Session, @@ -80,7 +78,7 @@ def create_dataset_aliases( def _add_one(dataset_alias: DatasetAlias) -> DatasetAliasModel: model = DatasetAliasModel.from_public(dataset_alias) session.add(model) - self.notify_dataset_alias_created(dataset_alias=dataset_alias) + cls.notify_dataset_alias_created(dataset_alias=dataset_alias) return model return [_add_one(a) for a in dataset_aliases] @@ -106,7 +104,6 @@ def _add_dataset_alias_association( @classmethod @internal_api_call - @provide_session def register_dataset_change( cls, *, @@ -115,7 +112,7 @@ def register_dataset_change( extra=None, aliases: Collection[DatasetAlias] = (), source_alias_names: Iterable[str] | None = None, - session: Session = NEW_SESSION, + session: Session, **kwargs, ) -> DatasetEvent | None: """ @@ -153,6 +150,7 @@ def register_dataset_change( dataset_event = DatasetEvent(**event_kwargs) session.add(dataset_event) + session.flush() # Ensure the event is written earlier than DDRQ entries below. dags_to_queue_from_dataset = { ref.dag for ref in dataset_model.consuming_dags if ref.dag.is_active and not ref.dag.is_paused @@ -183,7 +181,6 @@ def register_dataset_change( if dags_to_reparse: file_locs = {dag.fileloc for dag in dags_to_reparse} cls._send_dag_priority_parsing_request(file_locs, session) - session.flush() cls.notify_dataset_changed(dataset=dataset) @@ -191,19 +188,20 @@ def register_dataset_change( dags_to_queue = dags_to_queue_from_dataset | dags_to_queue_from_dataset_alias cls._queue_dagruns(dataset_id=dataset_model.id, dags_to_queue=dags_to_queue, session=session) - session.flush() return dataset_event - def notify_dataset_created(self, dataset: Dataset): + @staticmethod + def notify_dataset_created(dataset: Dataset): """Run applicable notification actions when a dataset is created.""" get_listener_manager().hook.on_dataset_created(dataset=dataset) - def notify_dataset_alias_created(self, dataset_alias: DatasetAlias): + @staticmethod + def notify_dataset_alias_created(dataset_alias: DatasetAlias): """Run applicable notification actions when a dataset alias is created.""" get_listener_manager().hook.on_dataset_alias_created(dataset_alias=dataset_alias) - @classmethod - def notify_dataset_changed(cls, dataset: Dataset): + @staticmethod + def notify_dataset_changed(dataset: Dataset): """Run applicable notification actions when a dataset is changed.""" get_listener_manager().hook.on_dataset_changed(dataset=dataset) diff --git a/airflow/models/taskinstance.py b/airflow/models/taskinstance.py index d3300207abfdc..c17acdd2b7212 100644 --- a/airflow/models/taskinstance.py +++ b/airflow/models/taskinstance.py @@ -31,7 +31,7 @@ from contextlib import nullcontext from datetime import timedelta from enum import Enum -from typing import TYPE_CHECKING, Any, Callable, Collection, Dict, Generator, Iterable, Mapping, Tuple +from typing import TYPE_CHECKING, Any, Callable, Collection, Generator, Iterable, Mapping, Tuple from urllib.parse import quote import dill @@ -2911,24 +2911,22 @@ def _register_dataset_changes(self, *, events: OutletEventAccessors, session: Se frozen_extra = frozenset(dataset_alias_event["extra"].items()) dataset_alias_names[(dataset_uri, frozen_extra)].add(dataset_alias_name) - class _DatasetModelCache(Dict[str, DatasetModel]): - log = self.log - - def __missing__(self, key: str) -> DatasetModel: - (dataset_obj,) = dataset_manager.create_datasets([Dataset(uri=key)], session=session) - session.flush() - self.log.warning("Created a new %r as it did not exist.", dataset_obj) - self[key] = dataset_obj - return dataset_obj - - dataset_objs_cache = _DatasetModelCache( - (dataset_obj.uri, dataset_obj) + dataset_models: dict[str, DatasetModel] = { + dataset_obj.uri: dataset_obj for dataset_obj in session.scalars( select(DatasetModel).where(DatasetModel.uri.in_(uri for uri, _ in dataset_alias_names)) ) - ) + } + if missing_datasets := [Dataset(uri=u) for u, _ in dataset_alias_names if u not in dataset_models]: + dataset_models.update( + (dataset_obj.uri, dataset_obj) + for dataset_obj in dataset_manager.create_datasets(missing_datasets, session=session) + ) + self.log.warning("Created new datasets for alias reference: %s", missing_datasets) + session.flush() # Needed because we need the id for fk. + for (uri, extra_items), alias_names in dataset_alias_names.items(): - dataset_obj = dataset_objs_cache[uri] + dataset_obj = dataset_models[uri] self.log.info( 'Creating event for %r through aliases "%s"', dataset_obj, diff --git a/tests/datasets/test_manager.py b/tests/datasets/test_manager.py index d3013aef60c29..9b8b0c180d48e 100644 --- a/tests/datasets/test_manager.py +++ b/tests/datasets/test_manager.py @@ -119,9 +119,10 @@ def test_register_dataset_change(self, session, dag_maker, mock_task_instance): session.add(dsm) dsm.consuming_dags = [DagScheduleDatasetReference(dag_id=dag.dag_id) for dag in (dag1, dag2)] session.execute(delete(DatasetDagRunQueue)) - session.commit() + session.flush() dsem.register_dataset_change(task_instance=mock_task_instance, dataset=ds, session=session) + session.flush() # Ensure we've created a dataset assert session.query(DatasetEvent).filter_by(dataset_id=dsm.id).count() == 1 @@ -134,9 +135,10 @@ def test_register_dataset_change_no_downstreams(self, session, mock_task_instanc dsm = DatasetModel(uri="never_consumed") session.add(dsm) session.execute(delete(DatasetDagRunQueue)) - session.commit() + session.flush() dsem.register_dataset_change(task_instance=mock_task_instance, dataset=ds, session=session) + session.flush() # Ensure we've created a dataset assert session.query(DatasetEvent).filter_by(dataset_id=dsm.id).count() == 1 @@ -150,14 +152,15 @@ def test_register_dataset_change_notifies_dataset_listener(self, session, mock_t ds = Dataset(uri="test_dataset_uri_2") dag1 = DagModel(dag_id="dag3") - session.add_all([dag1]) + session.add(dag1) dsm = DatasetModel(uri="test_dataset_uri_2") session.add(dsm) dsm.consuming_dags = [DagScheduleDatasetReference(dag_id=dag1.dag_id)] - session.commit() + session.flush() dsem.register_dataset_change(task_instance=mock_task_instance, dataset=ds, session=session) + session.flush() # Ensure the listener was notified assert len(dataset_listener.changed) == 1 diff --git a/tests/jobs/test_scheduler_job.py b/tests/jobs/test_scheduler_job.py index 4067f4fa17902..2292f0130e323 100644 --- a/tests/jobs/test_scheduler_job.py +++ b/tests/jobs/test_scheduler_job.py @@ -4300,6 +4300,7 @@ def test_no_create_dag_runs_when_dag_disabled(self, session, dag_maker, disable, dataset=ds, session=session, ) + session.flush() assert session.scalars(dse_q).one().source_run_id == dr1.run_id assert session.scalars(ddrq_q).one_or_none() is None @@ -4313,6 +4314,7 @@ def test_no_create_dag_runs_when_dag_disabled(self, session, dag_maker, disable, dataset=ds, session=session, ) + session.flush() assert [e.source_run_id for e in session.scalars(dse_q)] == [dr1.run_id, dr2.run_id] assert session.scalars(ddrq_q).one().target_dag_id == "consumer" diff --git a/tests/models/test_taskinstance.py b/tests/models/test_taskinstance.py index 773c68915cefa..d2922db267805 100644 --- a/tests/models/test_taskinstance.py +++ b/tests/models/test_taskinstance.py @@ -2325,7 +2325,9 @@ def test_outlet_datasets(self, create_task_instance): ddrq_timestamps = ( session.query(DatasetDagRunQueue.created_at).filter_by(dataset_id=event.dataset.id).all() ) - assert all([event.timestamp < ddrq_timestamp for (ddrq_timestamp,) in ddrq_timestamps]) + assert all( + event.timestamp < ddrq_timestamp for (ddrq_timestamp,) in ddrq_timestamps + ), f"Some items in {[str(t) for t in ddrq_timestamps]} are earlier than {event.timestamp}" @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode def test_outlet_datasets_failed(self, create_task_instance): From 54b96b89efddcb8364d71d3e8f4204e2da169601 Mon Sep 17 00:00:00 2001 From: Niko Oliveira Date: Wed, 25 Sep 2024 06:58:22 -0700 Subject: [PATCH 017/802] Refactor AWS Auth manager user output (#42454) AWS auth manager has incredible tooling to setup the required resources, however one piece needs to be done manually. This PR updates the docs and user output to make it more clear what needs to happen next. Removing the stacktrace (which usually indicates a critical failure in a piece of code) and replacing with a more clearly marked output message. Also update the docs to more clearly indicate that the script will most likely need user intervention. --- .../amazon/aws/auth_manager/cli/idc_commands.py | 10 +++++++--- .../auth-manager/setup/identity-center.rst | 6 +----- 2 files changed, 8 insertions(+), 8 deletions(-) diff --git a/airflow/providers/amazon/aws/auth_manager/cli/idc_commands.py b/airflow/providers/amazon/aws/auth_manager/cli/idc_commands.py index 388948765ace6..c4901351b2cff 100644 --- a/airflow/providers/amazon/aws/auth_manager/cli/idc_commands.py +++ b/airflow/providers/amazon/aws/auth_manager/cli/idc_commands.py @@ -19,6 +19,7 @@ from __future__ import annotations import logging +import sys from typing import TYPE_CHECKING import boto3 @@ -139,10 +140,13 @@ def _create_application(client: BaseClient, instance_arn: str | None, args) -> s # Remove this part when it is supported if "is not supported for this action" in e.response["Error"]["Message"]: print( - "Creation of SAML applications is only supported in AWS console today. " - "Please create the application through the console." + "*************************************************************************\n" + "* ACTION REQUIRED *\n" + "* Creation of SAML applications is only supported in AWS console today. *\n" + "* Please create the application through the console. *\n" + "*************************************************************************\n" ) - raise + sys.exit(1) print(f"Application created: '{response['ApplicationArn']}'") diff --git a/docs/apache-airflow-providers-amazon/auth-manager/setup/identity-center.rst b/docs/apache-airflow-providers-amazon/auth-manager/setup/identity-center.rst index a134dfe0ddf7c..acf3727bf9c7f 100644 --- a/docs/apache-airflow-providers-amazon/auth-manager/setup/identity-center.rst +++ b/docs/apache-airflow-providers-amazon/auth-manager/setup/identity-center.rst @@ -48,11 +48,7 @@ To create the resources, please run the following command: airflow aws-auth-manager init-identity-center -The CLI command should exit successfully with the message: :: - - AWS IAM Identity Center resources created successfully. - -If the CLI command exited with an error, please look carefully at the CLI command output to understand which resource(s) +The CLI command will ask you to create any resources manually if they cannot be automatically created. Please look carefully at the CLI command output to understand which resource(s) have or have not been created successfully. The resource(s) which have not been successfully created need to be :ref:`created manually `. From 229b97692bbebb7e67889f01a6d5866aa7bf348a Mon Sep 17 00:00:00 2001 From: David Blain Date: Wed, 25 Sep 2024 16:01:36 +0200 Subject: [PATCH 018/802] (bugfix): Paginated results in MSGraphAsyncOperator (#42414) * fix: Make sure that when paginated results are returned that when the last MSGraph calls occurs we get the all pages instead of only the last one * refactor: Refactored paginate function and added context parameter so that XCom's get fetched each time * refactor: Added unit tests for paginate function * refactor: Reformatted code like recommended by static checks * refactor: Added missing white line MSGraphAsyncOperator test * refactor: Changed keyword arguments to regular ones for calling pagination function --------- Co-authored-by: David Blain --- .../microsoft/azure/operators/msgraph.py | 56 ++++++++++--------- .../microsoft/azure/operators/test_msgraph.py | 36 +++++++++++- 2 files changed, 66 insertions(+), 26 deletions(-) diff --git a/airflow/providers/microsoft/azure/operators/msgraph.py b/airflow/providers/microsoft/azure/operators/msgraph.py index 74409f3600a1e..b3d14b14a57ec 100644 --- a/airflow/providers/microsoft/azure/operators/msgraph.py +++ b/airflow/providers/microsoft/azure/operators/msgraph.py @@ -100,7 +100,7 @@ def __init__( timeout: float | None = None, proxies: dict | None = None, api_version: APIVersion | str | None = None, - pagination_function: Callable[[MSGraphAsyncOperator, dict], tuple[str, dict]] | None = None, + pagination_function: Callable[[MSGraphAsyncOperator, dict, Context], tuple[str, dict]] | None = None, result_processor: Callable[[Context, Any], Any] = lambda context, result: result, serializer: type[ResponseSerializer] = ResponseSerializer, **kwargs: Any, @@ -122,7 +122,6 @@ def __init__( self.pagination_function = pagination_function or self.paginate self.result_processor = result_processor self.serializer: ResponseSerializer = serializer() - self.results: list[Any] | None = None def execute(self, context: Context) -> None: self.defer( @@ -166,6 +165,8 @@ def execute_complete( self.log.debug("response: %s", response) + results = self.pull_xcom(context=context) + if response: response = self.serializer.deserialize(response) @@ -178,39 +179,46 @@ def execute_complete( event["response"] = result try: - self.trigger_next_link(response=response, method_name=self.execute_complete.__name__) + self.trigger_next_link( + response=response, method_name=self.execute_complete.__name__, context=context + ) except TaskDeferred as exception: - self.results = self.pull_xcom(context=context) self.append_result( + results=results, result=result, append_result_as_list_if_absent=True, ) - self.push_xcom(context=context, value=self.results) + self.push_xcom(context=context, value=results) raise exception - self.append_result(result=result) + if not results: + return result - return self.results + self.append_result(results=results, result=result) + return results return None + @classmethod def append_result( - self, + cls, + results: list[Any], result: Any, append_result_as_list_if_absent: bool = False, - ): - if isinstance(self.results, list): + ) -> list[Any]: + if isinstance(results, list): if isinstance(result, list): - self.results.extend(result) + results.extend(result) else: - self.results.append(result) + results.append(result) else: if append_result_as_list_if_absent: if isinstance(result, list): - self.results = result + return result else: - self.results = [result] + return [result] else: - self.results = result + return result + return results def pull_xcom(self, context: Context) -> list: map_index = context["ti"].map_index @@ -251,27 +259,25 @@ def push_xcom(self, context: Context, value) -> None: self.xcom_push(context=context, key=self.key, value=value) @staticmethod - def paginate(operator: MSGraphAsyncOperator, response: dict) -> tuple[Any, dict[str, Any] | None]: + def paginate( + operator: MSGraphAsyncOperator, response: dict, context: Context + ) -> tuple[Any, dict[str, Any] | None]: odata_count = response.get("@odata.count") if odata_count and operator.query_parameters: query_parameters = deepcopy(operator.query_parameters) top = query_parameters.get("$top") - odata_count = response.get("@odata.count") if top and odata_count: - if len(response.get("value", [])) == top: - skip = ( - sum(map(lambda result: len(result["value"]), operator.results)) + top - if operator.results - else top - ) + if len(response.get("value", [])) == top and context: + results = operator.pull_xcom(context=context) + skip = sum(map(lambda result: len(result["value"]), results)) + top if results else top query_parameters["$skip"] = skip return operator.url, query_parameters return response.get("@odata.nextLink"), operator.query_parameters - def trigger_next_link(self, response, method_name="execute_complete") -> None: + def trigger_next_link(self, response, method_name: str, context: Context) -> None: if isinstance(response, dict): - url, query_parameters = self.pagination_function(self, response) + url, query_parameters = self.pagination_function(self, response, context) self.log.debug("url: %s", url) self.log.debug("query_parameters: %s", query_parameters) diff --git a/tests/providers/microsoft/azure/operators/test_msgraph.py b/tests/providers/microsoft/azure/operators/test_msgraph.py index b7520d731544c..754b653ccdaf0 100644 --- a/tests/providers/microsoft/azure/operators/test_msgraph.py +++ b/tests/providers/microsoft/azure/operators/test_msgraph.py @@ -26,7 +26,13 @@ from airflow.providers.microsoft.azure.operators.msgraph import MSGraphAsyncOperator from airflow.triggers.base import TriggerEvent from tests.providers.microsoft.azure.base import Base -from tests.providers.microsoft.conftest import load_file, load_json, mock_json_response, mock_response +from tests.providers.microsoft.conftest import ( + load_file, + load_json, + mock_context, + mock_json_response, + mock_response, +) class TestMSGraphAsyncOperator(Base): @@ -127,3 +133,31 @@ def test_template_fields(self): for template_field in MSGraphAsyncOperator.template_fields: getattr(operator, template_field) + + def test_paginate_without_query_parameters(self): + operator = MSGraphAsyncOperator( + task_id="user_license_details", + conn_id="msgraph_api", + url="users", + ) + context = mock_context(task=operator) + response = load_json("resources", "users.json") + next_link, query_parameters = MSGraphAsyncOperator.paginate(operator, response, context) + + assert next_link == response["@odata.nextLink"] + assert query_parameters is None + + def test_paginate_with_context_query_parameters(self): + operator = MSGraphAsyncOperator( + task_id="user_license_details", + conn_id="msgraph_api", + url="users", + query_parameters={"$top": 12}, + ) + context = mock_context(task=operator) + response = load_json("resources", "users.json") + response["@odata.count"] = 100 + url, query_parameters = MSGraphAsyncOperator.paginate(operator, response, context) + + assert url == "users" + assert query_parameters == {"$skip": 12, "$top": 12} From af43a4efb6daae5da42c0a244b66f07176fc1176 Mon Sep 17 00:00:00 2001 From: Gopal Dirisala <39794726+dirrao@users.noreply.github.com> Date: Wed, 25 Sep 2024 21:42:29 +0530 Subject: [PATCH 019/802] uv version bump to 0.4.7 (#42274) --- Dockerfile | 2 +- Dockerfile.ci | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/Dockerfile b/Dockerfile index 5cd1caec434ee..68f1ed166f12a 100644 --- a/Dockerfile +++ b/Dockerfile @@ -50,7 +50,7 @@ ARG AIRFLOW_VERSION="2.10.1" ARG PYTHON_BASE_IMAGE="python:3.8-slim-bookworm" ARG AIRFLOW_PIP_VERSION=24.2 -ARG AIRFLOW_UV_VERSION=0.4.1 +ARG AIRFLOW_UV_VERSION=0.4.7 ARG AIRFLOW_USE_UV="false" ARG UV_HTTP_TIMEOUT="300" ARG AIRFLOW_IMAGE_REPOSITORY="https://github.com/apache/airflow" diff --git a/Dockerfile.ci b/Dockerfile.ci index 9d9de62dd1a4b..ad944d151adcb 100644 --- a/Dockerfile.ci +++ b/Dockerfile.ci @@ -1262,7 +1262,7 @@ ARG DEFAULT_CONSTRAINTS_BRANCH="constraints-main" ARG AIRFLOW_CI_BUILD_EPOCH="10" ARG AIRFLOW_PRE_CACHED_PIP_PACKAGES="true" ARG AIRFLOW_PIP_VERSION=24.2 -ARG AIRFLOW_UV_VERSION=0.4.1 +ARG AIRFLOW_UV_VERSION=0.4.7 ARG AIRFLOW_USE_UV="true" # Setup PIP # By default PIP install run without cache to make image smaller @@ -1286,7 +1286,7 @@ ARG AIRFLOW_VERSION="" ARG ADDITIONAL_PIP_INSTALL_FLAGS="" ARG AIRFLOW_PIP_VERSION=24.2 -ARG AIRFLOW_UV_VERSION=0.4.1 +ARG AIRFLOW_UV_VERSION=0.4.7 ARG AIRFLOW_USE_UV="true" ENV AIRFLOW_REPO=${AIRFLOW_REPO}\ From 71bc5282b84fad27c1d16808f897b53e64cc2c9d Mon Sep 17 00:00:00 2001 From: "D. Ferruzzi" Date: Wed, 25 Sep 2024 10:14:25 -0700 Subject: [PATCH 020/802] Purge existing SLA implementation (#42285) SLA will be reimplemented in either 3.0 or 3.1 --- .../endpoints/rpc_api_endpoint.py | 1 - airflow/callbacks/callback_requests.py | 20 - airflow/config_templates/config.yml | 7 - airflow/dag_processing/manager.py | 47 +-- airflow/dag_processing/processor.py | 196 +-------- airflow/example_dags/example_sla_dag.py | 66 --- airflow/jobs/scheduler_job_runner.py | 28 +- airflow/models/baseoperator.py | 18 +- airflow/models/dag.py | 18 +- airflow/models/mappedoperator.py | 15 +- airflow/serialization/enums.py | 1 - airflow/serialization/serialized_objects.py | 8 +- airflow/settings.py | 3 - .../chime_notifier_howto_guide.rst | 4 - .../notifications/sns.rst | 5 - .../notifications/sqs.rst | 5 - .../pagerduty_notifier_howto_guide.rst | 4 - .../slack_notifier_howto_guide.rst | 4 - .../slackwebhook_notifier_howto_guide.rst | 4 - .../smtp_notifier_howto_guide.rst | 4 - .../logging-monitoring/callbacks.rst | 1 - .../logging-monitoring/metrics.rst | 4 - docs/apache-airflow/core-concepts/tasks.rst | 73 +--- docs/conf.py | 1 - newsfragments/42285.significant.rst | 1 + tests/callbacks/test_callback_requests.py | 9 - tests/dag_processing/test_job_runner.py | 34 +- tests/dag_processing/test_processor.py | 394 +----------------- tests/jobs/test_scheduler_job.py | 78 +--- tests/models/test_baseoperator.py | 45 -- tests/serialization/test_dag_serialization.py | 5 - 31 files changed, 39 insertions(+), 1064 deletions(-) delete mode 100644 airflow/example_dags/example_sla_dag.py create mode 100644 newsfragments/42285.significant.rst diff --git a/airflow/api_internal/endpoints/rpc_api_endpoint.py b/airflow/api_internal/endpoints/rpc_api_endpoint.py index e4a5069b29bcc..8716d9c9cc49d 100644 --- a/airflow/api_internal/endpoints/rpc_api_endpoint.py +++ b/airflow/api_internal/endpoints/rpc_api_endpoint.py @@ -101,7 +101,6 @@ def initialize_method_map() -> dict[str, Callable]: DagFileProcessor._execute_task_callbacks, DagFileProcessor.execute_callbacks, DagFileProcessor.execute_callbacks_without_dag, - DagFileProcessor.manage_slas, DagFileProcessor.save_dag_to_db, DagFileProcessor.update_import_errors, DagFileProcessor._validate_task_pools_and_update_dag_warnings, diff --git a/airflow/callbacks/callback_requests.py b/airflow/callbacks/callback_requests.py index 7158c45d44d91..07ad648e9630f 100644 --- a/airflow/callbacks/callback_requests.py +++ b/airflow/callbacks/callback_requests.py @@ -137,23 +137,3 @@ def __init__( self.dag_id = dag_id self.run_id = run_id self.is_failure_callback = is_failure_callback - - -class SlaCallbackRequest(CallbackRequest): - """ - A class with information about the SLA callback to be executed. - - :param full_filepath: File Path to use to run the callback - :param dag_id: DAG ID - :param processor_subdir: Directory used by Dag Processor when parsed the dag. - """ - - def __init__( - self, - full_filepath: str, - dag_id: str, - processor_subdir: str | None, - msg: str | None = None, - ): - super().__init__(full_filepath, processor_subdir=processor_subdir, msg=msg) - self.dag_id = dag_id diff --git a/airflow/config_templates/config.yml b/airflow/config_templates/config.yml index 3bef18058dfbb..c9abee3c85065 100644 --- a/airflow/config_templates/config.yml +++ b/airflow/config_templates/config.yml @@ -395,13 +395,6 @@ core: type: integer example: ~ default: "30" - check_slas: - description: | - On each dagrun check against defined SLAs - version_added: 1.10.8 - type: string - example: ~ - default: "True" xcom_backend: description: | Path to custom XCom class that will be used to store and resolve operators results diff --git a/airflow/dag_processing/manager.py b/airflow/dag_processing/manager.py index 6df8060f3a311..05fb72daee602 100644 --- a/airflow/dag_processing/manager.py +++ b/airflow/dag_processing/manager.py @@ -42,7 +42,7 @@ import airflow.models from airflow.api_internal.internal_api_call import internal_api_call -from airflow.callbacks.callback_requests import CallbackRequest, SlaCallbackRequest +from airflow.callbacks.callback_requests import CallbackRequest from airflow.configuration import conf from airflow.dag_processing.processor import DagFileProcessorProcess from airflow.models.dag import DagModel @@ -752,40 +752,17 @@ def _fetch_callbacks_with_retries( return callback_queue def _add_callback_to_queue(self, request: CallbackRequest): - # requests are sent by dag processors. SLAs exist per-dag, but can be generated once per SLA-enabled - # task in the dag. If treated like other callbacks, SLAs can cause feedback where a SLA arrives, - # goes to the front of the queue, gets processed, triggers more SLAs from the same DAG, which go to - # the front of the queue, and we never get round to picking stuff off the back of the queue - if isinstance(request, SlaCallbackRequest): - if request in self._callback_to_execute[request.full_filepath]: - self.log.debug("Skipping already queued SlaCallbackRequest") - return - - # not already queued, queue the callback - # do NOT add the file of this SLA to self._file_path_queue. SLAs can arrive so rapidly that - # they keep adding to the file queue and never letting it drain. This in turn prevents us from - # ever rescanning the dags folder for changes to existing dags. We simply store the callback, and - # periodically, when self._file_path_queue is drained, we rescan and re-queue all DAG files. - # The SLAs will be picked up then. It means a delay in reacting to the SLAs (as controlled by the - # min_file_process_interval config) but stops SLAs from DoS'ing the queue. - self.log.debug("Queuing SlaCallbackRequest for %s", request.dag_id) - self._callback_to_execute[request.full_filepath].append(request) - Stats.incr("dag_processing.sla_callback_count") - - # Other callbacks have a higher priority over DAG Run scheduling, so those callbacks gazump, even if - # already in the file path queue - else: - self.log.debug("Queuing %s CallbackRequest: %s", type(request).__name__, request) - self._callback_to_execute[request.full_filepath].append(request) - if request.full_filepath in self._file_path_queue: - # Remove file paths matching request.full_filepath from self._file_path_queue - # Since we are already going to use that filepath to run callback, - # there is no need to have same file path again in the queue - self._file_path_queue = deque( - file_path for file_path in self._file_path_queue if file_path != request.full_filepath - ) - self._add_paths_to_queue([request.full_filepath], True) - Stats.incr("dag_processing.other_callback_count") + self.log.debug("Queuing %s CallbackRequest: %s", type(request).__name__, request) + self._callback_to_execute[request.full_filepath].append(request) + if request.full_filepath in self._file_path_queue: + # Remove file paths matching request.full_filepath from self._file_path_queue + # Since we are already going to use that filepath to run callback, + # there is no need to have same file path again in the queue + self._file_path_queue = deque( + file_path for file_path in self._file_path_queue if file_path != request.full_filepath + ) + self._add_paths_to_queue([request.full_filepath], True) + Stats.incr("dag_processing.other_callback_count") def _refresh_requested_filelocs(self) -> None: """Refresh filepaths from dag dir as requested by users via APIs.""" diff --git a/airflow/dag_processing/processor.py b/airflow/dag_processing/processor.py index 0b19d8f2db76c..f030cb75019e5 100644 --- a/airflow/dag_processing/processor.py +++ b/airflow/dag_processing/processor.py @@ -25,33 +25,28 @@ import zipfile from contextlib import contextmanager, redirect_stderr, redirect_stdout, suppress from dataclasses import dataclass -from datetime import timedelta -from typing import TYPE_CHECKING, Generator, Iterable, Iterator +from typing import TYPE_CHECKING, Generator, Iterable from setproctitle import setproctitle -from sqlalchemy import delete, event, func, or_, select +from sqlalchemy import delete, event from airflow import settings -from airflow.api_internal.internal_api_call import InternalApiConfig, internal_api_call +from airflow.api_internal.internal_api_call import internal_api_call from airflow.callbacks.callback_requests import ( DagCallbackRequest, - SlaCallbackRequest, TaskCallbackRequest, ) from airflow.configuration import conf -from airflow.exceptions import AirflowException, TaskNotFound +from airflow.exceptions import AirflowException from airflow.listeners.listener import get_listener_manager -from airflow.models import SlaMiss from airflow.models.dag import DAG, DagModel from airflow.models.dagbag import DagBag -from airflow.models.dagrun import DagRun as DR from airflow.models.dagwarning import DagWarning, DagWarningType from airflow.models.errors import ParseImportError from airflow.models.serialized_dag import SerializedDagModel -from airflow.models.taskinstance import TaskInstance, TaskInstance as TI, _run_finished_callback +from airflow.models.taskinstance import TaskInstance, _run_finished_callback from airflow.stats import Stats from airflow.utils import timezone -from airflow.utils.email import get_email_address_list, send_email from airflow.utils.file import iter_airflow_imports, might_contain_dag from airflow.utils.log.logging_mixin import LoggingMixin, StreamLogWriter, set_context from airflow.utils.mixins import MultiprocessingStartMethodMixin @@ -440,180 +435,6 @@ def __init__(self, dag_ids: list[str] | None, dag_directory: str, log: logging.L self.dag_warnings: set[tuple[str, str]] = set() self._last_num_of_db_queries = 0 - @classmethod - @internal_api_call - @provide_session - def manage_slas(cls, dag_folder, dag_id: str, session: Session = NEW_SESSION) -> None: - """ - Find all tasks that have SLAs defined, and send alert emails when needed. - - New SLA misses are also recorded in the database. - - We are assuming that the scheduler runs often, so we only check for - tasks that should have succeeded in the past hour. - """ - dagbag = DagFileProcessor._get_dagbag(dag_folder) - dag = dagbag.get_dag(dag_id) - cls.logger().info("Running SLA Checks for %s", dag.dag_id) - if not any(isinstance(ti.sla, timedelta) for ti in dag.tasks): - cls.logger().info("Skipping SLA check for %s because no tasks in DAG have SLAs", dag) - return - qry = ( - select(TI.task_id, func.max(DR.execution_date).label("max_ti")) - .join(TI.dag_run) - .where(TI.dag_id == dag.dag_id) - .where(or_(TI.state == TaskInstanceState.SUCCESS, TI.state == TaskInstanceState.SKIPPED)) - .where(TI.task_id.in_(dag.task_ids)) - .group_by(TI.task_id) - .subquery("sq") - ) - # get recorded SlaMiss - recorded_slas_query = set( - session.execute( - select(SlaMiss.dag_id, SlaMiss.task_id, SlaMiss.execution_date).where( - SlaMiss.dag_id == dag.dag_id, SlaMiss.task_id.in_(dag.task_ids) - ) - ) - ) - max_tis: Iterator[TI] = session.scalars( - select(TI) - .join(TI.dag_run) - .where(TI.dag_id == dag.dag_id, TI.task_id == qry.c.task_id, DR.execution_date == qry.c.max_ti) - ) - - ts = timezone.utcnow() - - for ti in max_tis: - task = dag.get_task(ti.task_id) - if not task.sla: - continue - - if not isinstance(task.sla, timedelta): - raise TypeError( - f"SLA is expected to be timedelta object, got " - f"{type(task.sla)} in {task.dag_id}:{task.task_id}" - ) - - sla_misses = [] - next_info = dag.next_dagrun_info(dag.get_run_data_interval(ti.dag_run), restricted=False) - while next_info and next_info.logical_date < ts: - next_info = dag.next_dagrun_info(next_info.data_interval, restricted=False) - - if next_info is None: - break - if (ti.dag_id, ti.task_id, next_info.logical_date) in recorded_slas_query: - continue - if next_info.logical_date + task.sla < ts: - sla_miss = SlaMiss( - task_id=ti.task_id, - dag_id=ti.dag_id, - execution_date=next_info.logical_date, - timestamp=ts, - ) - sla_misses.append(sla_miss) - Stats.incr("sla_missed", tags={"dag_id": ti.dag_id, "task_id": ti.task_id}) - if sla_misses: - session.add_all(sla_misses) - session.commit() - slas: list[SlaMiss] = session.scalars( - select(SlaMiss).where(~SlaMiss.notification_sent, SlaMiss.dag_id == dag.dag_id) - ).all() - if slas: - sla_dates: list[datetime] = [sla.execution_date for sla in slas] - fetched_tis: list[TI] = session.scalars( - select(TI).where( - TI.dag_id == dag.dag_id, - TI.execution_date.in_(sla_dates), - TI.state != TaskInstanceState.SUCCESS, - ) - ).all() - blocking_tis: list[TI] = [] - for ti in fetched_tis: - if ti.task_id in dag.task_ids: - ti.task = dag.get_task(ti.task_id) - blocking_tis.append(ti) - else: - session.delete(ti) - session.commit() - - task_list = "\n".join(sla.task_id + " on " + sla.execution_date.isoformat() for sla in slas) - blocking_task_list = "\n".join( - ti.task_id + " on " + ti.execution_date.isoformat() for ti in blocking_tis - ) - # Track whether email or any alert notification sent - # We consider email or the alert callback as notifications - email_sent = False - notification_sent = False - if dag.sla_miss_callback: - # Execute the alert callback - callbacks = ( - dag.sla_miss_callback - if isinstance(dag.sla_miss_callback, list) - else [dag.sla_miss_callback] - ) - for callback in callbacks: - cls.logger().info("Calling SLA miss callback %s", callback) - try: - callback(dag, task_list, blocking_task_list, slas, blocking_tis) - notification_sent = True - except Exception: - Stats.incr( - "sla_callback_notification_failure", - tags={ - "dag_id": dag.dag_id, - "func_name": callback.__name__, - }, - ) - cls.logger().exception( - "Could not call sla_miss_callback(%s) for DAG %s", - callback.__name__, - dag.dag_id, - ) - email_content = f"""\ - Here's a list of tasks that missed their SLAs: -
{task_list}\n
- Blocking tasks: -
{blocking_task_list}
- Airflow Webserver URL: {conf.get(section='webserver', key='base_url')} - """ - - tasks_missed_sla = [] - for sla in slas: - try: - task = dag.get_task(sla.task_id) - except TaskNotFound: - # task already deleted from DAG, skip it - cls.logger().warning( - "Task %s doesn't exist in DAG anymore, skipping SLA miss notification.", sla.task_id - ) - else: - tasks_missed_sla.append(task) - - emails: set[str] = set() - for task in tasks_missed_sla: - if task.email: - if isinstance(task.email, str): - emails.update(get_email_address_list(task.email)) - elif isinstance(task.email, (list, tuple)): - emails.update(task.email) - if emails: - try: - send_email(emails, f"[airflow] SLA miss on DAG={dag.dag_id}", email_content) - email_sent = True - notification_sent = True - except Exception: - Stats.incr("sla_email_notification_failure", tags={"dag_id": dag.dag_id}) - cls.logger().exception( - "Could not send SLA Miss email notification for DAG %s", dag.dag_id - ) - # If we sent any notification, update the sla_miss table - if notification_sent: - for sla in slas: - sla.email_sent = email_sent - sla.notification_sent = True - session.merge(sla) - session.commit() - @staticmethod @internal_api_call @provide_session @@ -748,13 +569,6 @@ def execute_callbacks( try: if isinstance(request, TaskCallbackRequest): cls._execute_task_callbacks(dagbag, request, unit_test_mode, session=session) - elif isinstance(request, SlaCallbackRequest): - if InternalApiConfig.get_use_internal_api(): - cls.logger().warning( - "SlaCallbacks are not supported when the Internal API is enabled" - ) - else: - DagFileProcessor.manage_slas(dagbag.dag_folder, request.dag_id, session=session) elif isinstance(request, DagCallbackRequest): cls._execute_dag_callbacks(dagbag, request, session=session) except Exception: diff --git a/airflow/example_dags/example_sla_dag.py b/airflow/example_dags/example_sla_dag.py deleted file mode 100644 index aca1277e88799..0000000000000 --- a/airflow/example_dags/example_sla_dag.py +++ /dev/null @@ -1,66 +0,0 @@ -# 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. -"""Example DAG demonstrating SLA use in Tasks""" - -from __future__ import annotations - -import datetime -import time - -import pendulum - -from airflow.decorators import dag, task - - -# [START howto_task_sla] -def sla_callback(dag, task_list, blocking_task_list, slas, blocking_tis): - print( - "The callback arguments are: ", - { - "dag": dag, - "task_list": task_list, - "blocking_task_list": blocking_task_list, - "slas": slas, - "blocking_tis": blocking_tis, - }, - ) - - -@dag( - schedule="*/2 * * * *", - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - catchup=False, - sla_miss_callback=sla_callback, - default_args={"email": "email@example.com"}, -) -def example_sla_dag(): - @task(sla=datetime.timedelta(seconds=10)) - def sleep_20(): - """Sleep for 20 seconds""" - time.sleep(20) - - @task - def sleep_30(): - """Sleep for 30 seconds""" - time.sleep(30) - - sleep_20() >> sleep_30() - - -example_dag = example_sla_dag() - -# [END howto_task_sla] diff --git a/airflow/jobs/scheduler_job_runner.py b/airflow/jobs/scheduler_job_runner.py index a49c2361ec423..9438edd4d9187 100644 --- a/airflow/jobs/scheduler_job_runner.py +++ b/airflow/jobs/scheduler_job_runner.py @@ -36,7 +36,7 @@ from sqlalchemy.sql import expression from airflow import settings -from airflow.callbacks.callback_requests import DagCallbackRequest, SlaCallbackRequest, TaskCallbackRequest +from airflow.callbacks.callback_requests import DagCallbackRequest, TaskCallbackRequest from airflow.callbacks.pipe_callback_sink import PipeCallbackSink from airflow.configuration import conf from airflow.exceptions import UnknownExecutorException @@ -1724,37 +1724,11 @@ def _verify_integrity_if_dag_changed(self, dag_run: DagRun, session: Session) -> return True def _send_dag_callbacks_to_processor(self, dag: DAG, callback: DagCallbackRequest | None = None) -> None: - self._send_sla_callbacks_to_processor(dag) if callback: self.job.executor.send_callback(callback) else: self.log.debug("callback is empty") - def _send_sla_callbacks_to_processor(self, dag: DAG) -> None: - """Send SLA Callbacks to DagFileProcessor if tasks have SLAs set and check_slas=True.""" - if not settings.CHECK_SLAS: - return - - if not any(isinstance(task.sla, timedelta) for task in dag.tasks): - self.log.debug("Skipping SLA check for %s because no tasks in DAG have SLAs", dag) - return - - if not dag.timetable.periodic: - self.log.debug("Skipping SLA check for %s because DAG is not scheduled", dag) - return - - dag_model = DagModel.get_dagmodel(dag.dag_id) - if not dag_model: - self.log.error("Couldn't find DAG %s in database!", dag.dag_id) - return - - request = SlaCallbackRequest( - full_filepath=dag.fileloc, - dag_id=dag.dag_id, - processor_subdir=dag_model.processor_subdir, - ) - self.job.executor.send_callback(request) - @provide_session def _fail_tasks_stuck_in_queued(self, session: Session = NEW_SESSION) -> None: """ diff --git a/airflow/models/baseoperator.py b/airflow/models/baseoperator.py index 8f95d1eee7302..20656586ba01e 100644 --- a/airflow/models/baseoperator.py +++ b/airflow/models/baseoperator.py @@ -677,17 +677,7 @@ class derived from this one results in the creation of a task object, way to limit concurrency for certain tasks :param pool_slots: the number of pool slots this task should use (>= 1) Values less than 1 are not allowed. - :param sla: time by which the job is expected to succeed. Note that - this represents the ``timedelta`` after the period is closed. For - example if you set an SLA of 1 hour, the scheduler would send an email - soon after 1:00AM on the ``2016-01-02`` if the ``2016-01-01`` instance - has not succeeded yet. - The scheduler pays special attention for jobs with an SLA and - sends alert - emails for SLA misses. SLA misses are also recorded in the database - for future reference. All tasks that share the same SLA time - get bundled in a single email, sent soon after that time. SLA - notification are sent once and only once for each task instance. + :param sla: DEPRECATED - The SLA feature is removed in Airflow 3.0, to be replaced with a new implementation in 3.1 :param execution_timeout: max time allowed for the execution of this task instance, if it goes beyond it will raise and fail. :param on_failure_callback: a function or list of functions to be called when a task instance @@ -975,7 +965,11 @@ def __init__( if self.pool_slots < 1: dag_str = f" in dag {dag.dag_id}" if dag else "" raise ValueError(f"pool slots for {self.task_id}{dag_str} cannot be less than 1") - self.sla = sla + + if sla: + self.log.warning( + "The SLA feature is removed in Airflow 3.0, to be replaced with a new implementation in 3.1" + ) if not TriggerRule.is_valid(trigger_rule): raise AirflowException( diff --git a/airflow/models/dag.py b/airflow/models/dag.py index 215dae298f106..91f8aec7302cb 100644 --- a/airflow/models/dag.py +++ b/airflow/models/dag.py @@ -42,7 +42,6 @@ Container, Iterable, Iterator, - List, MutableSet, Pattern, Sequence, @@ -146,7 +145,6 @@ from airflow.decorators import TaskDecoratorCollection from airflow.models.dagbag import DagBag from airflow.models.operator import Operator - from airflow.models.slamiss import SlaMiss from airflow.serialization.pydantic.dag import DagModelPydantic from airflow.serialization.pydantic.dag_run import DagRunPydantic from airflow.typing_compat import Literal @@ -169,8 +167,6 @@ Collection[Union["Dataset", "DatasetAlias"]], ] -SLAMissCallback = Callable[["DAG", str, str, List["SlaMiss"], List[TaskInstance]], None] - class InconsistentDataInterval(AirflowException): """ @@ -430,10 +426,7 @@ class DAG(LoggingMixin): beyond this the scheduler will disable the DAG :param dagrun_timeout: Specify the duration a DagRun should be allowed to run before it times out or fails. Task instances that are running when a DagRun is timed out will be marked as skipped. - :param sla_miss_callback: specify a function or list of functions to call when reporting SLA - timeouts. See :ref:`sla_miss_callback` for - more information about the function signature and parameters that are - passed to the callback. + :param sla_miss_callback: DEPRECATED - The SLA feature is removed in Airflow 3.0, to be replaced with a new implementation in 3.1 :param default_view: Specify DAG default view (grid, graph, duration, gantt, landing_times), default grid :param orientation: Specify DAG orientation in graph view (LR, TB, RL, BT), default LR @@ -519,7 +512,7 @@ def __init__( "core", "max_consecutive_failed_dag_runs_per_dag" ), dagrun_timeout: timedelta | None = None, - sla_miss_callback: None | SLAMissCallback | list[SLAMissCallback] = None, + sla_miss_callback: Any = None, default_view: str = airflow_conf.get_mandatory_value("webserver", "dag_default_view").lower(), orientation: str = airflow_conf.get_mandatory_value("webserver", "dag_orientation"), catchup: bool = airflow_conf.getboolean("scheduler", "catchup_by_default"), @@ -639,7 +632,10 @@ def __init__( f"requires max_active_runs <= {self.timetable.active_runs_limit}" ) self.dagrun_timeout = dagrun_timeout - self.sla_miss_callback = sla_miss_callback + if sla_miss_callback: + log.warning( + "The SLA feature is removed in Airflow 3.0, to be replaced with a new implementation in 3.1" + ) if default_view in DEFAULT_VIEW_PRESETS: self._default_view: str = default_view else: @@ -3297,7 +3293,7 @@ def dag( "core", "max_consecutive_failed_dag_runs_per_dag" ), dagrun_timeout: timedelta | None = None, - sla_miss_callback: None | SLAMissCallback | list[SLAMissCallback] = None, + sla_miss_callback: Any = None, default_view: str = airflow_conf.get_mandatory_value("webserver", "dag_default_view").lower(), orientation: str = airflow_conf.get_mandatory_value("webserver", "dag_orientation"), catchup: bool = airflow_conf.getboolean("scheduler", "catchup_by_default"), diff --git a/airflow/models/mappedoperator.py b/airflow/models/mappedoperator.py index 2cb7d993fc9f9..8a9e790ea7fc6 100644 --- a/airflow/models/mappedoperator.py +++ b/airflow/models/mappedoperator.py @@ -26,7 +26,7 @@ import attr import methodtools -from airflow.exceptions import AirflowException, UnmappableOperator +from airflow.exceptions import UnmappableOperator from airflow.models.abstractoperator import ( DEFAULT_EXECUTOR, DEFAULT_IGNORE_FIRST_DEPENDS_ON_PAST, @@ -328,11 +328,6 @@ def __attrs_post_init__(self): for k, v in self.partial_kwargs.items(): if k in self.template_fields: XComArg.apply_upstream_relationship(self, v) - if self.partial_kwargs.get("sla") is not None: - raise AirflowException( - f"SLAs are unsupported with mapped tasks. Please set `sla=None` for task " - f"{self.task_id!r}." - ) @methodtools.lru_cache(maxsize=None) @classmethod @@ -547,14 +542,6 @@ def weight_rule(self) -> PriorityWeightStrategy: # type: ignore[override] def weight_rule(self, value: str | PriorityWeightStrategy) -> None: self.partial_kwargs["weight_rule"] = validate_and_load_priority_weight_strategy(value) - @property - def sla(self) -> datetime.timedelta | None: - return self.partial_kwargs.get("sla") - - @sla.setter - def sla(self, value: datetime.timedelta | None) -> None: - self.partial_kwargs["sla"] = value - @property def max_active_tis_per_dag(self) -> int | None: return self.partial_kwargs.get("max_active_tis_per_dag") diff --git a/airflow/serialization/enums.py b/airflow/serialization/enums.py index f216ce7316103..49a3de3d774c4 100644 --- a/airflow/serialization/enums.py +++ b/airflow/serialization/enums.py @@ -71,6 +71,5 @@ class DagAttributeTypes(str, Enum): ARG_NOT_SET = "arg_not_set" TASK_CALLBACK_REQUEST = "task_callback_request" DAG_CALLBACK_REQUEST = "dag_callback_request" - SLA_CALLBACK_REQUEST = "sla_callback_request" TASK_INSTANCE_KEY = "task_instance_key" TRIGGER = "trigger" diff --git a/airflow/serialization/serialized_objects.py b/airflow/serialization/serialized_objects.py index 12310685ec692..c9c1f11835277 100644 --- a/airflow/serialization/serialized_objects.py +++ b/airflow/serialization/serialized_objects.py @@ -34,7 +34,7 @@ from pendulum.tz.timezone import FixedTimezone, Timezone from airflow import macros -from airflow.callbacks.callback_requests import DagCallbackRequest, SlaCallbackRequest, TaskCallbackRequest +from airflow.callbacks.callback_requests import DagCallbackRequest, TaskCallbackRequest from airflow.compat.functools import cache from airflow.datasets import ( BaseDataset, @@ -758,8 +758,6 @@ def serialize( return cls._encode(var.to_json(), type_=DAT.TASK_CALLBACK_REQUEST) elif isinstance(var, DagCallbackRequest): return cls._encode(var.to_json(), type_=DAT.DAG_CALLBACK_REQUEST) - elif isinstance(var, SlaCallbackRequest): - return cls._encode(var.to_json(), type_=DAT.SLA_CALLBACK_REQUEST) elif var.__class__ == Context: d = {} for k, v in var._context.items(): @@ -890,8 +888,6 @@ def deserialize(cls, encoded_var: Any, use_pydantic_models=False) -> Any: return TaskCallbackRequest.from_json(var) elif type_ == DAT.DAG_CALLBACK_REQUEST: return DagCallbackRequest.from_json(var) - elif type_ == DAT.SLA_CALLBACK_REQUEST: - return SlaCallbackRequest.from_json(var) elif type_ == DAT.TASK_INSTANCE_KEY: return TaskInstanceKey(**var) elif use_pydantic_models and _ENABLE_AIP_44: @@ -1289,7 +1285,7 @@ def populate_operator(cls, op: Operator, encoded_op: dict[str, Any]) -> None: continue elif k == "downstream_task_ids": v = set(v) - elif k in {"retry_delay", "execution_timeout", "sla", "max_retry_delay"}: + elif k in {"retry_delay", "execution_timeout", "max_retry_delay"}: v = cls._deserialize_timedelta(v) elif k in encoded_op["template_fields"]: pass diff --git a/airflow/settings.py b/airflow/settings.py index a242ce4da7694..7a805f64a29c7 100644 --- a/airflow/settings.py +++ b/airflow/settings.py @@ -781,9 +781,6 @@ def is_usage_data_collection_enabled() -> bool: ALLOW_FUTURE_EXEC_DATES = conf.getboolean("scheduler", "allow_trigger_in_future", fallback=False) -# Whether or not to check each dagrun against defined SLAs -CHECK_SLAS = conf.getboolean("core", "check_slas", fallback=True) - USE_JOB_SCHEDULE = conf.getboolean("scheduler", "use_job_schedule", fallback=True) # By default Airflow plugins are lazily-loaded (only loaded when required). Set it to False, diff --git a/docs/apache-airflow-providers-amazon/notifications/chime_notifier_howto_guide.rst b/docs/apache-airflow-providers-amazon/notifications/chime_notifier_howto_guide.rst index c10b8cbae4142..a52540fe78282 100644 --- a/docs/apache-airflow-providers-amazon/notifications/chime_notifier_howto_guide.rst +++ b/docs/apache-airflow-providers-amazon/notifications/chime_notifier_howto_guide.rst @@ -23,10 +23,6 @@ Introduction Chime notifier (:class:`airflow.providers.amazon.aws.notifications.chime.ChimeNotifier`) allows users to send messages to a Chime chat room setup via a webhook using the various ``on_*_callbacks`` at both the DAG level and Task level -You can also use a notifier with ``sla_miss_callback``. - -.. note:: - When notifiers are used with `sla_miss_callback` the context will contain only values passed to the callback, refer :ref:`sla_miss_callback`. Example Code: ------------- diff --git a/docs/apache-airflow-providers-amazon/notifications/sns.rst b/docs/apache-airflow-providers-amazon/notifications/sns.rst index 337e82cf62eb4..bbaad4f814712 100644 --- a/docs/apache-airflow-providers-amazon/notifications/sns.rst +++ b/docs/apache-airflow-providers-amazon/notifications/sns.rst @@ -25,11 +25,6 @@ Introduction `Amazon SNS `__ notifier :class:`~airflow.providers.amazon.aws.notifications.sns.SnsNotifier` allows users to push messages to a SNS Topic using the various ``on_*_callbacks`` at both the DAG level and Task level. -You can also use a notifier with ``sla_miss_callback``. - -.. note:: - When notifiers are used with ``sla_miss_callback`` the context will contain only values passed to the callback, - refer :ref:`sla_miss_callback`. Example Code: ------------- diff --git a/docs/apache-airflow-providers-amazon/notifications/sqs.rst b/docs/apache-airflow-providers-amazon/notifications/sqs.rst index 4a2232b006a03..6951caa9fdd67 100644 --- a/docs/apache-airflow-providers-amazon/notifications/sqs.rst +++ b/docs/apache-airflow-providers-amazon/notifications/sqs.rst @@ -25,11 +25,6 @@ Introduction `Amazon SQS `__ notifier :class:`~airflow.providers.amazon.aws.notifications.sqs.SqsNotifier` allows users to push messages to an Amazon SQS Queue using the various ``on_*_callbacks`` at both the DAG level and Task level. -You can also use a notifier with ``sla_miss_callback``. - -.. note:: - When notifiers are used with ``sla_miss_callback`` the context will contain only values passed to the callback, - refer :ref:`sla_miss_callback`. Example Code: ------------- diff --git a/docs/apache-airflow-providers-pagerduty/notifications/pagerduty_notifier_howto_guide.rst b/docs/apache-airflow-providers-pagerduty/notifications/pagerduty_notifier_howto_guide.rst index d93d5a2fc5757..d16f9b2b9e48a 100644 --- a/docs/apache-airflow-providers-pagerduty/notifications/pagerduty_notifier_howto_guide.rst +++ b/docs/apache-airflow-providers-pagerduty/notifications/pagerduty_notifier_howto_guide.rst @@ -23,10 +23,6 @@ Introduction The Pagerduty notifier (:class:`airflow.providers.pagerduty.notifications.pagerduty.PagerdutyNotifier`) allows users to send messages to Pagerduty using the various ``on_*_callbacks`` at both the DAG level and Task level. -You can also use a notifier with ``sla_miss_callback``. - -.. note:: - When notifiers are used with `sla_miss_callback` the context will contain only values passed to the callback, refer :ref:`sla_miss_callback`. Example Code: ------------- diff --git a/docs/apache-airflow-providers-slack/notifications/slack_notifier_howto_guide.rst b/docs/apache-airflow-providers-slack/notifications/slack_notifier_howto_guide.rst index d967779cee9c5..a4f891f8a57bb 100644 --- a/docs/apache-airflow-providers-slack/notifications/slack_notifier_howto_guide.rst +++ b/docs/apache-airflow-providers-slack/notifications/slack_notifier_howto_guide.rst @@ -23,10 +23,6 @@ Introduction Slack notifier (:class:`airflow.providers.slack.notifications.slack.SlackNotifier`) allows users to send messages to a slack channel using the various ``on_*_callbacks`` at both the DAG level and Task level -You can also use a notifier with ``sla_miss_callback``. - -.. note:: - When notifiers are used with `sla_miss_callback` the context will contain only values passed to the callback, refer :ref:`sla_miss_callback`. Example Code: ------------- diff --git a/docs/apache-airflow-providers-slack/notifications/slackwebhook_notifier_howto_guide.rst b/docs/apache-airflow-providers-slack/notifications/slackwebhook_notifier_howto_guide.rst index bb9e85c67466f..66ced818a7d18 100644 --- a/docs/apache-airflow-providers-slack/notifications/slackwebhook_notifier_howto_guide.rst +++ b/docs/apache-airflow-providers-slack/notifications/slackwebhook_notifier_howto_guide.rst @@ -24,10 +24,6 @@ Slack Incoming Webhook notifier (:class:`airflow.providers.slack.notifications.s allows users to send messages to a slack channel through `Incoming Webhook `__ using the various ``on_*_callbacks`` at both the DAG level and Task level -You can also use a notifier with ``sla_miss_callback``. - -.. note:: - When notifiers are used with `sla_miss_callback` the context will contain only values passed to the callback, refer :ref:`sla_miss_callback`. Example Code: ------------- diff --git a/docs/apache-airflow-providers-smtp/notifications/smtp_notifier_howto_guide.rst b/docs/apache-airflow-providers-smtp/notifications/smtp_notifier_howto_guide.rst index c7183c5e56874..4cb1bf310e03d 100644 --- a/docs/apache-airflow-providers-smtp/notifications/smtp_notifier_howto_guide.rst +++ b/docs/apache-airflow-providers-smtp/notifications/smtp_notifier_howto_guide.rst @@ -23,10 +23,6 @@ Introduction The SMTP notifier (:class:`airflow.providers.smtp.notifications.smtp.SmtpNotifier`) allows users to send messages to SMTP servers using the various ``on_*_callbacks`` at both the DAG level and Task level. -You can also use a notifier with ``sla_miss_callback``. - -.. note:: - When notifiers are used with `sla_miss_callback` the context will contain only values passed to the callback, refer :ref:`sla_miss_callback`. Example Code: ------------- diff --git a/docs/apache-airflow/administration-and-deployment/logging-monitoring/callbacks.rst b/docs/apache-airflow/administration-and-deployment/logging-monitoring/callbacks.rst index a70a876ba347e..b54071373cf09 100644 --- a/docs/apache-airflow/administration-and-deployment/logging-monitoring/callbacks.rst +++ b/docs/apache-airflow/administration-and-deployment/logging-monitoring/callbacks.rst @@ -46,7 +46,6 @@ Name Description =========================================== ================================================================ ``on_success_callback`` Invoked when the task :ref:`succeeds ` ``on_failure_callback`` Invoked when the task :ref:`fails ` -``sla_miss_callback`` Invoked when a task misses its defined :ref:`SLA ` ``on_retry_callback`` Invoked when the task is :ref:`up for retry ` ``on_execute_callback`` Invoked right before the task begins executing. ``on_skipped_callback`` Invoked when the task is :ref:`running ` and AirflowSkipException raised. diff --git a/docs/apache-airflow/administration-and-deployment/logging-monitoring/metrics.rst b/docs/apache-airflow/administration-and-deployment/logging-monitoring/metrics.rst index c8522bee3ba10..61985cecea9b0 100644 --- a/docs/apache-airflow/administration-and-deployment/logging-monitoring/metrics.rst +++ b/docs/apache-airflow/administration-and-deployment/logging-monitoring/metrics.rst @@ -164,7 +164,6 @@ Name Descripti Metric with file_path and action tagging. ``dag_processing.processor_timeouts`` Number of file processors that have been killed due to taking too long. Metric with file_path tagging. -``dag_processing.sla_callback_count`` Number of SLA callbacks received ``dag_processing.other_callback_count`` Number of non-SLA callbacks received ``dag_processing.file_path_queue_update_count`` Number of times we've scanned the filesystem and queued all existing dags ``dag_file_processor_timeouts`` (DEPRECATED) same behavior as ``dag_processing.processor_timeouts`` @@ -176,9 +175,6 @@ Name Descripti ``scheduler.critical_section_busy`` Count of times a scheduler process tried to get a lock on the critical section (needed to send tasks to the executor) and found it locked by another process. -``sla_missed`` Number of SLA misses. Metric with dag_id and task_id tagging. -``sla_callback_notification_failure`` Number of failed SLA miss callback notification attempts. Metric with dag_id and func_name tagging. -``sla_email_notification_failure`` Number of failed SLA miss email notification attempts. Metric with dag_id tagging. ``ti.start..`` Number of started task in a given dag. Similar to _start but for task ``ti.start`` Number of started task in a given dag. Similar to _start but for task. Metric with dag_id and task_id tagging. diff --git a/docs/apache-airflow/core-concepts/tasks.rst b/docs/apache-airflow/core-concepts/tasks.rst index 0e05f55bcf5c8..ad03283ef772d 100644 --- a/docs/apache-airflow/core-concepts/tasks.rst +++ b/docs/apache-airflow/core-concepts/tasks.rst @@ -149,82 +149,11 @@ is periodically executed and rescheduled until it succeeds. mode="reschedule", ) -If you merely want to be notified if a task runs over but still let it run to completion, you want :ref:`concepts:slas` instead. - - -.. _concepts:slas: SLAs ---- -An SLA, or a Service Level Agreement, is an expectation for the maximum time a Task should be completed relative to the Dag Run start time. If a task takes longer than this to run, it is then visible in the "SLA Misses" part of the user interface, as well as going out in an email of all tasks that missed their SLA. - -Tasks over their SLA are not cancelled, though - they are allowed to run to completion. If you want to cancel a task after a certain runtime is reached, you want :ref:`concepts:timeouts` instead. - -To set an SLA for a task, pass a ``datetime.timedelta`` object to the Task/Operator's ``sla`` parameter. You can also supply an ``sla_miss_callback`` that will be called when the SLA is missed if you want to run your own logic. - -If you want to disable SLA checking entirely, you can set ``check_slas = False`` in Airflow's ``[core]`` configuration. - -To read more about configuring the emails, see :doc:`/howto/email-config`. - -.. note:: - - Manually-triggered tasks and tasks in event-driven DAGs will not be checked for an SLA miss. For more information on DAG ``schedule`` values see :doc:`DAG Run `. - -.. _concepts:sla_miss_callback: - -sla_miss_callback -~~~~~~~~~~~~~~~~~ - -You can also supply an ``sla_miss_callback`` that will be called when the SLA is missed if you want to run your own logic. -The function signature of an ``sla_miss_callback`` requires 5 parameters. - -#. ``dag`` - - * Parent :ref:`DAG ` Object for the :doc:`DAGRun ` in which tasks missed their - :ref:`SLA `. - -#. ``task_list`` - - * String list (new-line separated, \\n) of all tasks that missed their :ref:`SLA ` - since the last time that the ``sla_miss_callback`` ran. - -#. ``blocking_task_list`` - - * Any task in the :doc:`DAGRun(s)` (with the same ``execution_date`` as a task that missed - :ref:`SLA `) that is not in a **SUCCESS** state at the time that the ``sla_miss_callback`` - runs. i.e. 'running', 'failed'. These tasks are described as tasks that are blocking itself or another - task from completing before its SLA window is complete. - -#. ``slas`` - - * List of :py:mod:`SlaMiss` objects associated with the tasks in the - ``task_list`` parameter. - -#. ``blocking_tis`` - - * List of the :ref:`TaskInstance ` objects that are associated with the tasks - in the ``blocking_task_list`` parameter. - -Examples of ``sla_miss_callback`` function signature: - -.. code-block:: python - - def my_sla_miss_callback(dag, task_list, blocking_task_list, slas, blocking_tis): - ... - -.. code-block:: python - - def my_sla_miss_callback(*args): - ... - -Example DAG: - -.. exampleinclude:: /../../airflow/example_dags/example_sla_dag.py - :language: python - :start-after: [START howto_task_sla] - :end-before: [END howto_task_sla] - +The SLA feature from Airflow 2 has been removed in 3.0 and will be replaced with a new implementation in Airflow 3.1 Special Exceptions ------------------ diff --git a/docs/conf.py b/docs/conf.py index c87871e7ede6d..4d01e402195a5 100644 --- a/docs/conf.py +++ b/docs/conf.py @@ -755,7 +755,6 @@ def _get_params(root_schema: dict, prefix: str = "", default_section: str = "") "*/node_modules/*", "*/migrations/*", "*/contrib/*", - "**/example_sla_dag.py", "**/example_taskflow_api_docker_virtualenv.py", "**/example_dag_decorator.py", ] diff --git a/newsfragments/42285.significant.rst b/newsfragments/42285.significant.rst new file mode 100644 index 0000000000000..8f8cfa0dee298 --- /dev/null +++ b/newsfragments/42285.significant.rst @@ -0,0 +1 @@ +The SLA feature is removed in Airflow 3.0, to be replaced with Airflow Alerts in 3.1 diff --git a/tests/callbacks/test_callback_requests.py b/tests/callbacks/test_callback_requests.py index 6d900c8bd3571..5992ee6fbbf70 100644 --- a/tests/callbacks/test_callback_requests.py +++ b/tests/callbacks/test_callback_requests.py @@ -23,7 +23,6 @@ from airflow.callbacks.callback_requests import ( CallbackRequest, DagCallbackRequest, - SlaCallbackRequest, TaskCallbackRequest, ) from airflow.models.dag import DAG @@ -55,14 +54,6 @@ class TestCallbackRequest: ), DagCallbackRequest, ), - ( - SlaCallbackRequest( - full_filepath="filepath", - dag_id="fake_dag", - processor_subdir="/test_dir", - ), - SlaCallbackRequest, - ), ], ) def test_from_json(self, input, request_class): diff --git a/tests/dag_processing/test_job_runner.py b/tests/dag_processing/test_job_runner.py index 8112b7222a697..1d3fefdf12d5f 100644 --- a/tests/dag_processing/test_job_runner.py +++ b/tests/dag_processing/test_job_runner.py @@ -39,7 +39,7 @@ import time_machine from sqlalchemy import func -from airflow.callbacks.callback_requests import CallbackRequest, DagCallbackRequest, SlaCallbackRequest +from airflow.callbacks.callback_requests import CallbackRequest, DagCallbackRequest from airflow.config_templates.airflow_local_settings import DEFAULT_LOGGING_CONFIG from airflow.configuration import conf from airflow.dag_processing.manager import ( @@ -1179,16 +1179,10 @@ def test_fetch_callbacks_from_database(self, tmp_path): processor_subdir=os.fspath(tmp_path), run_id="456", ) - callback3 = SlaCallbackRequest( - dag_id="test_start_date_scheduling", - full_filepath=str(dag_filepath), - processor_subdir=os.fspath(tmp_path), - ) with create_session() as session: session.add(DbCallbackRequest(callback=callback1, priority_weight=11)) session.add(DbCallbackRequest(callback=callback2, priority_weight=10)) - session.add(DbCallbackRequest(callback=callback3, priority_weight=9)) child_pipe, parent_pipe = multiprocessing.Pipe() manager = DagProcessorJobRunner( @@ -1371,16 +1365,6 @@ def test_callback_queue(self, tmp_path): processor_subdir=tmp_path, msg=None, ) - dag1_sla1 = SlaCallbackRequest( - full_filepath="/green_eggs/ham/file1.py", - dag_id="dag1", - processor_subdir=tmp_path, - ) - dag1_sla2 = SlaCallbackRequest( - full_filepath="/green_eggs/ham/file1.py", - dag_id="dag1", - processor_subdir=tmp_path, - ) dag2_req1 = DagCallbackRequest( full_filepath="/green_eggs/ham/file2.py", @@ -1391,15 +1375,8 @@ def test_callback_queue(self, tmp_path): msg=None, ) - dag3_sla1 = SlaCallbackRequest( - full_filepath="/green_eggs/ham/file3.py", - dag_id="dag3", - processor_subdir=tmp_path, - ) - # when manager.processor._add_callback_to_queue(dag1_req1) - manager.processor._add_callback_to_queue(dag1_sla1) manager.processor._add_callback_to_queue(dag2_req1) # then - requests should be in manager's queue, with dag2 ahead of dag1 (because it was added last) @@ -1408,18 +1385,10 @@ def test_callback_queue(self, tmp_path): dag1_req1.full_filepath, dag2_req1.full_filepath, } - assert manager.processor._callback_to_execute[dag1_req1.full_filepath] == [dag1_req1, dag1_sla1] assert manager.processor._callback_to_execute[dag2_req1.full_filepath] == [dag2_req1] - # when - manager.processor._add_callback_to_queue(dag1_sla2) - manager.processor._add_callback_to_queue(dag3_sla1) - - # then - since sla2 == sla1, should not have brought dag1 to the fore, and an SLA on dag3 doesn't # update the queue, although the callback is registered assert manager.processor._file_path_queue == deque([dag2_req1.full_filepath, dag1_req1.full_filepath]) - assert manager.processor._callback_to_execute[dag1_req1.full_filepath] == [dag1_req1, dag1_sla1] - assert manager.processor._callback_to_execute[dag3_sla1.full_filepath] == [dag3_sla1] # when manager.processor._add_callback_to_queue(dag1_req2) @@ -1428,7 +1397,6 @@ def test_callback_queue(self, tmp_path): assert manager.processor._file_path_queue == deque([dag1_req1.full_filepath, dag2_req1.full_filepath]) assert manager.processor._callback_to_execute[dag1_req1.full_filepath] == [ dag1_req1, - dag1_sla1, dag1_req2, ] diff --git a/tests/dag_processing/test_processor.py b/tests/dag_processing/test_processor.py index 2b250ae8c55ed..d7b2b2116653e 100644 --- a/tests/dag_processing/test_processor.py +++ b/tests/dag_processing/test_processor.py @@ -32,10 +32,9 @@ from airflow.configuration import TEST_DAGS_FOLDER, conf from airflow.dag_processing.manager import DagFileProcessorAgent from airflow.dag_processing.processor import DagFileProcessor, DagFileProcessorProcess -from airflow.models import DagBag, DagModel, SlaMiss, TaskInstance +from airflow.models import DagBag, DagModel, TaskInstance from airflow.models.serialized_dag import SerializedDagModel from airflow.models.taskinstance import SimpleTaskInstance -from airflow.operators.empty import EmptyOperator from airflow.utils import timezone from airflow.utils.session import create_session from airflow.utils.state import State @@ -50,7 +49,6 @@ clear_db_pools, clear_db_runs, clear_db_serialized_dags, - clear_db_sla_miss, ) from tests.test_utils.mock_executor import MockExecutor @@ -89,7 +87,6 @@ def clean_db(): clear_db_runs() clear_db_pools() clear_db_dags() - clear_db_sla_miss() clear_db_import_errors() clear_db_jobs() clear_db_serialized_dags() @@ -116,395 +113,6 @@ def _process_file(self, file_path, dag_directory, session): dag_file_processor.process_file(file_path, [], False) - @pytest.mark.skip_if_database_isolation_mode - @mock.patch("airflow.dag_processing.processor.DagFileProcessor._get_dagbag") - def test_dag_file_processor_sla_miss_callback(self, mock_get_dagbag, create_dummy_dag, get_test_dag): - """ - Test that the dag file processor calls the sla miss callback - """ - session = settings.Session() - sla_callback = MagicMock() - - # Create dag with a start of 1 day ago, but a sla of 0, so we'll already have a sla_miss on the books. - test_start_date = timezone.utcnow() - datetime.timedelta(days=1) - test_run_id = DagRunType.SCHEDULED.generate_run_id(test_start_date) - dag, task = create_dummy_dag( - dag_id="test_sla_miss", - task_id="dummy", - sla_miss_callback=sla_callback, - default_args={"start_date": test_start_date, "sla": datetime.timedelta()}, - ) - - session.merge( - TaskInstance( - task=task, - run_id=test_run_id, - state=State.SUCCESS, - ) - ) - session.merge(SlaMiss(task_id="dummy", dag_id="test_sla_miss", execution_date=test_start_date)) - - mock_dagbag = mock.Mock() - mock_dagbag.get_dag.return_value = dag - mock_get_dagbag.return_value = mock_dagbag - session.commit() - - DagFileProcessor.manage_slas(dag_folder=dag.fileloc, dag_id="test_sla_miss", session=session) - - assert sla_callback.called - - @pytest.mark.skip_if_database_isolation_mode - @mock.patch("airflow.dag_processing.processor.DagFileProcessor._get_dagbag") - def test_dag_file_processor_sla_miss_callback_invalid_sla(self, mock_get_dagbag, create_dummy_dag): - """ - Test that the dag file processor does not call the sla miss callback when - given an invalid sla - """ - session = settings.Session() - - sla_callback = MagicMock() - - # Create dag with a start of 1 day ago, but an sla of 0 - # so we'll already have an sla_miss on the books. - # Pass anything besides a timedelta object to the sla argument. - test_start_date = timezone.utcnow() - datetime.timedelta(days=1) - test_run_id = DagRunType.SCHEDULED.generate_run_id(test_start_date) - dag, task = create_dummy_dag( - dag_id="test_sla_miss", - task_id="dummy", - sla_miss_callback=sla_callback, - default_args={"start_date": test_start_date, "sla": None}, - ) - - session.merge(TaskInstance(task=task, run_id=test_run_id, state=State.SUCCESS)) - session.merge(SlaMiss(task_id="dummy", dag_id="test_sla_miss", execution_date=test_start_date)) - - mock_dagbag = mock.Mock() - mock_dagbag.get_dag.return_value = dag - mock_get_dagbag.return_value = mock_dagbag - - DagFileProcessor.manage_slas(dag_folder=dag.fileloc, dag_id="test_sla_miss", session=session) - sla_callback.assert_not_called() - - @pytest.mark.skip_if_database_isolation_mode - @mock.patch("airflow.dag_processing.processor.DagFileProcessor._get_dagbag") - def test_dag_file_processor_sla_miss_callback_sent_notification(self, mock_get_dagbag, create_dummy_dag): - """ - Test that the dag file processor does not call the sla_miss_callback when a - notification has already been sent - """ - session = settings.Session() - - # Mock the callback function so we can verify that it was not called - sla_callback = MagicMock() - - # Create dag with a start of 2 days ago, but an sla of 1 day - # ago so we'll already have an sla_miss on the books - test_start_date = timezone.utcnow() - datetime.timedelta(days=2) - test_run_id = DagRunType.SCHEDULED.generate_run_id(test_start_date) - dag, task = create_dummy_dag( - dag_id="test_sla_miss", - task_id="dummy", - sla_miss_callback=sla_callback, - default_args={"start_date": test_start_date, "sla": datetime.timedelta(days=1)}, - ) - - # Create a TaskInstance for two days ago - session.merge(TaskInstance(task=task, run_id=test_run_id, state=State.SUCCESS)) - - # Create an SlaMiss where notification was sent, but email was not - session.merge( - SlaMiss( - task_id="dummy", - dag_id="test_sla_miss", - execution_date=test_start_date, - email_sent=False, - notification_sent=True, - ) - ) - - mock_dagbag = mock.Mock() - mock_dagbag.get_dag.return_value = dag - mock_get_dagbag.return_value = mock_dagbag - - # Now call manage_slas and see if the sla_miss callback gets called - DagFileProcessor.manage_slas(dag_folder=dag.fileloc, dag_id="test_sla_miss", session=session) - - sla_callback.assert_not_called() - - @pytest.mark.skip_if_database_isolation_mode - @mock.patch("airflow.dag_processing.processor.Stats.incr") - @mock.patch("airflow.dag_processing.processor.DagFileProcessor._get_dagbag") - def test_dag_file_processor_sla_miss_doesnot_raise_integrity_error( - self, mock_get_dagbag, mock_stats_incr, dag_maker - ): - """ - Test that the dag file processor does not try to insert already existing item into the database - """ - session = settings.Session() - - # Create dag with a start of 2 days ago, but an sla of 1 day - # ago so we'll already have an sla_miss on the books - test_start_date = timezone.utcnow() - datetime.timedelta(days=2) - with dag_maker( - dag_id="test_sla_miss", - default_args={"start_date": test_start_date, "sla": datetime.timedelta(days=1)}, - ) as dag: - task = EmptyOperator(task_id="dummy") - - dr = dag_maker.create_dagrun(execution_date=test_start_date, state=State.SUCCESS) - - # Create a TaskInstance for two days ago - ti = TaskInstance(task=task, run_id=dr.run_id, state=State.SUCCESS) - session.merge(ti) - session.flush() - - mock_dagbag = mock.Mock() - mock_dagbag.get_dag.return_value = dag - mock_get_dagbag.return_value = mock_dagbag - - DagFileProcessor.manage_slas(dag_folder=dag.fileloc, dag_id="test_sla_miss", session=session) - sla_miss_count = ( - session.query(SlaMiss) - .filter( - SlaMiss.dag_id == dag.dag_id, - SlaMiss.task_id == task.task_id, - ) - .count() - ) - assert sla_miss_count == 1 - mock_stats_incr.assert_called_with("sla_missed", tags={"dag_id": "test_sla_miss", "task_id": "dummy"}) - # Now call manage_slas and see that it runs without errors - # because of existing SlaMiss above. - # Since this is run often, it's possible that it runs before another - # ti is successful thereby trying to insert a duplicate record. - DagFileProcessor.manage_slas(dag_folder=dag.fileloc, dag_id="test_sla_miss", session=session) - - @pytest.mark.skip_if_database_isolation_mode - @mock.patch("airflow.dag_processing.processor.Stats.incr") - @mock.patch("airflow.dag_processing.processor.DagFileProcessor._get_dagbag") - def test_dag_file_processor_sla_miss_continue_checking_the_task_instances_after_recording_missing_sla( - self, mock_get_dagbag, mock_stats_incr, dag_maker - ): - """ - Test that the dag file processor continue checking subsequent task instances - even if the preceding task instance misses the sla ahead - """ - session = settings.Session() - - # Create a dag with a start of 3 days ago and sla of 1 day, - # so we have 2 missing slas - now = timezone.utcnow() - test_start_date = now - datetime.timedelta(days=3) - # test_run_id = DagRunType.SCHEDULED.generate_run_id(test_start_date) - with dag_maker( - dag_id="test_sla_miss", - default_args={"start_date": test_start_date, "sla": datetime.timedelta(days=1)}, - ) as dag: - task = EmptyOperator(task_id="dummy") - - dr = dag_maker.create_dagrun(execution_date=test_start_date, state=State.SUCCESS) - - session.merge(TaskInstance(task=task, run_id=dr.run_id, state="success")) - session.merge( - SlaMiss(task_id=task.task_id, dag_id=dag.dag_id, execution_date=now - datetime.timedelta(days=2)) - ) - session.flush() - - mock_dagbag = mock.Mock() - mock_dagbag.get_dag.return_value = dag - mock_get_dagbag.return_value = mock_dagbag - - DagFileProcessor.manage_slas(dag_folder=dag.fileloc, dag_id="test_sla_miss", session=session) - sla_miss_count = ( - session.query(SlaMiss) - .filter( - SlaMiss.dag_id == dag.dag_id, - SlaMiss.task_id == task.task_id, - ) - .count() - ) - assert sla_miss_count == 2 - mock_stats_incr.assert_called_with("sla_missed", tags={"dag_id": "test_sla_miss", "task_id": "dummy"}) - - @pytest.mark.skip_if_database_isolation_mode - @patch.object(DagFileProcessor, "logger") - @mock.patch("airflow.dag_processing.processor.Stats.incr") - @mock.patch("airflow.dag_processing.processor.DagFileProcessor._get_dagbag") - def test_dag_file_processor_sla_miss_callback_exception( - self, - mock_get_dagbag, - mock_stats_incr, - mock_get_log, - create_dummy_dag, - ): - """ - Test that the dag file processor gracefully logs an exception if there is a problem - calling the sla_miss_callback - """ - session = settings.Session() - - sla_callback = MagicMock( - __name__="function_name", side_effect=RuntimeError("Could not call function") - ) - - test_start_date = timezone.utcnow() - datetime.timedelta(days=1) - test_run_id = DagRunType.SCHEDULED.generate_run_id(test_start_date) - - for i, callback in enumerate([[sla_callback], sla_callback]): - dag, task = create_dummy_dag( - dag_id=f"test_sla_miss_{i}", - task_id="dummy", - sla_miss_callback=callback, - default_args={"start_date": test_start_date, "sla": datetime.timedelta(hours=1)}, - ) - mock_stats_incr.reset_mock() - - session.merge(TaskInstance(task=task, run_id=test_run_id, state=State.SUCCESS)) - - # Create an SlaMiss where notification was sent, but email was not - session.merge( - SlaMiss(task_id="dummy", dag_id=f"test_sla_miss_{i}", execution_date=test_start_date) - ) - - # Now call manage_slas and see if the sla_miss callback gets called - mock_log = mock.Mock() - mock_get_log.return_value = mock_log - mock_dagbag = mock.Mock() - mock_dagbag.get_dag.return_value = dag - mock_get_dagbag.return_value = mock_dagbag - - DagFileProcessor.manage_slas(dag_folder=dag.fileloc, dag_id="test_sla_miss", session=session) - assert sla_callback.called - mock_log.exception.assert_called_once_with( - "Could not call sla_miss_callback(%s) for DAG %s", - sla_callback.__name__, - f"test_sla_miss_{i}", - ) - mock_stats_incr.assert_called_once_with( - "sla_callback_notification_failure", - tags={"dag_id": f"test_sla_miss_{i}", "func_name": sla_callback.__name__}, - ) - - @pytest.mark.skip_if_database_isolation_mode - @mock.patch("airflow.dag_processing.processor.send_email") - @mock.patch("airflow.dag_processing.processor.DagFileProcessor._get_dagbag") - def test_dag_file_processor_only_collect_emails_from_sla_missed_tasks( - self, mock_get_dagbag, mock_send_email, create_dummy_dag - ): - session = settings.Session() - - test_start_date = timezone.utcnow() - datetime.timedelta(days=1) - test_run_id = DagRunType.SCHEDULED.generate_run_id(test_start_date) - email1 = "test1@test.com" - dag, task = create_dummy_dag( - dag_id="test_sla_miss", - task_id="sla_missed", - email=email1, - default_args={"start_date": test_start_date, "sla": datetime.timedelta(hours=1)}, - ) - session.merge(TaskInstance(task=task, run_id=test_run_id, state=State.SUCCESS)) - - email2 = "test2@test.com" - EmptyOperator(task_id="sla_not_missed", dag=dag, owner="airflow", email=email2) - - session.merge(SlaMiss(task_id="sla_missed", dag_id="test_sla_miss", execution_date=test_start_date)) - - mock_dagbag = mock.Mock() - mock_dagbag.get_dag.return_value = dag - mock_get_dagbag.return_value = mock_dagbag - - DagFileProcessor.manage_slas(dag_folder=dag.fileloc, dag_id="test_sla_miss", session=session) - - assert len(mock_send_email.call_args_list) == 1 - - send_email_to = mock_send_email.call_args_list[0][0][0] - assert email1 in send_email_to - assert email2 not in send_email_to - - @pytest.mark.skip_if_database_isolation_mode - @patch.object(DagFileProcessor, "logger") - @mock.patch("airflow.dag_processing.processor.Stats.incr") - @mock.patch("airflow.utils.email.send_email") - @mock.patch("airflow.dag_processing.processor.DagFileProcessor._get_dagbag") - def test_dag_file_processor_sla_miss_email_exception( - self, - mock_get_dagbag, - mock_send_email, - mock_stats_incr, - mock_get_log, - create_dummy_dag, - ): - """ - Test that the dag file processor gracefully logs an exception if there is a problem - sending an email - """ - session = settings.Session() - dag_id = "test_sla_miss" - task_id = "test_ti" - email = "test@test.com" - - # Mock the callback function so we can verify that it was not called - mock_send_email.side_effect = RuntimeError("Could not send an email") - - test_start_date = timezone.utcnow() - datetime.timedelta(days=1) - test_run_id = DagRunType.SCHEDULED.generate_run_id(test_start_date) - dag, task = create_dummy_dag( - dag_id=dag_id, - task_id=task_id, - email=email, - default_args={"start_date": test_start_date, "sla": datetime.timedelta(hours=1)}, - ) - mock_stats_incr.reset_mock() - - session.merge(TaskInstance(task=task, run_id=test_run_id, state=State.SUCCESS)) - - # Create an SlaMiss where notification was sent, but email was not - session.merge(SlaMiss(task_id=task_id, dag_id=dag_id, execution_date=test_start_date)) - - mock_log = mock.Mock() - mock_get_log.return_value = mock_log - mock_dagbag = mock.Mock() - mock_dagbag.get_dag.return_value = dag - mock_get_dagbag.return_value = mock_dagbag - - DagFileProcessor.manage_slas(dag_folder=dag.fileloc, dag_id=dag_id, session=session) - mock_log.exception.assert_called_once_with( - "Could not send SLA Miss email notification for DAG %s", dag_id - ) - mock_stats_incr.assert_called_once_with("sla_email_notification_failure", tags={"dag_id": dag_id}) - - @pytest.mark.skip_if_database_isolation_mode - @mock.patch("airflow.dag_processing.processor.DagFileProcessor._get_dagbag") - def test_dag_file_processor_sla_miss_deleted_task(self, mock_get_dagbag, create_dummy_dag): - """ - Test that the dag file processor will not crash when trying to send - sla miss notification for a deleted task - """ - session = settings.Session() - - test_start_date = timezone.utcnow() - datetime.timedelta(days=1) - test_run_id = DagRunType.SCHEDULED.generate_run_id(test_start_date) - dag, task = create_dummy_dag( - dag_id="test_sla_miss", - task_id="dummy", - email="test@test.com", - default_args={"start_date": test_start_date, "sla": datetime.timedelta(hours=1)}, - ) - - session.merge(TaskInstance(task=task, run_id=test_run_id, state=State.SUCCESS)) - - # Create an SlaMiss where notification was sent, but email was not - session.merge( - SlaMiss(task_id="dummy_deleted", dag_id="test_sla_miss", execution_date=test_start_date) - ) - - mock_dagbag = mock.Mock() - mock_dagbag.get_dag.return_value = dag - mock_get_dagbag.return_value = mock_dagbag - - DagFileProcessor.manage_slas(dag_folder=dag.fileloc, dag_id="test_sla_miss", session=session) - @pytest.mark.skip_if_database_isolation_mode # Test is broken in db isolation mode @patch.object(TaskInstance, "handle_failure") def test_execute_on_failure_callbacks(self, mock_ti_handle_failure): diff --git a/tests/jobs/test_scheduler_job.py b/tests/jobs/test_scheduler_job.py index 2292f0130e323..52e9dbdeb1a04 100644 --- a/tests/jobs/test_scheduler_job.py +++ b/tests/jobs/test_scheduler_job.py @@ -36,7 +36,7 @@ import airflow.example_dags from airflow import settings -from airflow.callbacks.callback_requests import DagCallbackRequest, SlaCallbackRequest, TaskCallbackRequest +from airflow.callbacks.callback_requests import DagCallbackRequest, TaskCallbackRequest from airflow.callbacks.database_callback_sink import DatabaseCallbackSink from airflow.callbacks.pipe_callback_sink import PipeCallbackSink from airflow.dag_processing.manager import DagFileProcessorAgent @@ -3987,82 +3987,6 @@ def test_adopt_or_reset_orphaned_tasks_only_fails_scheduler_jobs(self, caplog): assert old_task_job.state == State.RUNNING assert "Marked 1 SchedulerJob instances as failed" in caplog.messages - def test_send_sla_callbacks_to_processor_sla_disabled(self, dag_maker): - """Test SLA Callbacks are not sent when check_slas is False""" - dag_id = "test_send_sla_callbacks_to_processor_sla_disabled" - with dag_maker(dag_id=dag_id, schedule="@daily") as dag: - EmptyOperator(task_id="task1") - - with patch.object(settings, "CHECK_SLAS", False): - scheduler_job = Job() - self.job_runner = SchedulerJobRunner(job=scheduler_job, subdir=os.devnull) - scheduler_job.executor = MockExecutor() - self.job_runner._send_sla_callbacks_to_processor(dag) - scheduler_job.executor.callback_sink.send.assert_not_called() - - def test_send_sla_callbacks_to_processor_sla_no_task_slas(self, dag_maker): - """Test SLA Callbacks are not sent when no task SLAs are defined""" - dag_id = "test_send_sla_callbacks_to_processor_sla_no_task_slas" - with dag_maker(dag_id=dag_id, schedule="@daily") as dag: - EmptyOperator(task_id="task1") - - with patch.object(settings, "CHECK_SLAS", True): - scheduler_job = Job() - self.job_runner = SchedulerJobRunner(job=scheduler_job, subdir=os.devnull) - scheduler_job.executor = MockExecutor() - self.job_runner._send_sla_callbacks_to_processor(dag) - scheduler_job.executor.callback_sink.send.assert_not_called() - - @pytest.mark.parametrize( - "schedule", - [ - "@daily", - "0 10 * * *", - timedelta(hours=2), - ], - ) - def test_send_sla_callbacks_to_processor_sla_with_task_slas(self, schedule, dag_maker): - """Test SLA Callbacks are sent to the DAG Processor when SLAs are defined on tasks""" - dag_id = "test_send_sla_callbacks_to_processor_sla_with_task_slas" - with dag_maker( - dag_id=dag_id, - schedule=schedule, - processor_subdir=TEST_DAG_FOLDER, - ) as dag: - EmptyOperator(task_id="task1", sla=timedelta(seconds=60)) - - with patch.object(settings, "CHECK_SLAS", True): - scheduler_job = Job() - self.job_runner = SchedulerJobRunner(job=scheduler_job, subdir=os.devnull) - scheduler_job.executor = MockExecutor() - self.job_runner._send_sla_callbacks_to_processor(dag) - expected_callback = SlaCallbackRequest( - full_filepath=dag.fileloc, - dag_id=dag.dag_id, - processor_subdir=TEST_DAG_FOLDER, - ) - scheduler_job.executor.callback_sink.send.assert_called_once_with(expected_callback) - - @pytest.mark.parametrize( - "schedule", - [ - None, - [Dataset("foo")], - ], - ) - def test_send_sla_callbacks_to_processor_sla_dag_not_scheduled(self, schedule, dag_maker): - """Test SLA Callbacks are not sent when DAG isn't scheduled""" - dag_id = "test_send_sla_callbacks_to_processor_sla_no_task_slas" - with dag_maker(dag_id=dag_id, schedule=schedule) as dag: - EmptyOperator(task_id="task1", sla=timedelta(seconds=5)) - - with patch.object(settings, "CHECK_SLAS", True): - scheduler_job = Job() - self.job_runner = SchedulerJobRunner(job=scheduler_job, subdir=os.devnull) - scheduler_job.executor = MockExecutor() - self.job_runner._send_sla_callbacks_to_processor(dag) - scheduler_job.executor.callback_sink.send.assert_not_called() - @pytest.mark.parametrize( "schedule, number_running, excepted", [ diff --git a/tests/models/test_baseoperator.py b/tests/models/test_baseoperator.py index 2aa5b76b22c03..3c5b7634d5a99 100644 --- a/tests/models/test_baseoperator.py +++ b/tests/models/test_baseoperator.py @@ -304,51 +304,6 @@ def test_render_template_with_native_envs(self, content, context, expected_outpu result = task.render_template(content, context) assert result == expected_output - def test_mapped_dag_slas_disabled_classic(self): - class MyOp(BaseOperator): - def __init__(self, x, **kwargs): - self.x = x - super().__init__(**kwargs) - - def execute(self, context): - print(self.x) - - with DAG( - dag_id="test-dag", - schedule=None, - start_date=DEFAULT_DATE, - default_args={"sla": timedelta(minutes=30)}, - ) as dag: - - @dag.task - def get_values(): - return [0, 1, 2] - - task1 = get_values() - with pytest.raises(AirflowException, match="SLAs are unsupported with mapped tasks"): - MyOp.partial(task_id="hi").expand(x=task1) - - def test_mapped_dag_slas_disabled_taskflow(self): - with DAG( - dag_id="test-dag", - schedule=None, - start_date=DEFAULT_DATE, - default_args={"sla": timedelta(minutes=30)}, - ) as dag: - - @dag.task - def get_values(): - return [0, 1, 2] - - task1 = get_values() - - @dag.task - def print_val(x): - print(x) - - with pytest.raises(AirflowException, match="SLAs are unsupported with mapped tasks"): - print_val.expand(x=task1) - @pytest.mark.db_test def test_render_template_fields(self): """Verify if operator attributes are correctly templated.""" diff --git a/tests/serialization/test_dag_serialization.py b/tests/serialization/test_dag_serialization.py index 7dfe57054c60f..758c7f496ed93 100644 --- a/tests/serialization/test_dag_serialization.py +++ b/tests/serialization/test_dag_serialization.py @@ -143,7 +143,6 @@ def detect_task_dependencies(task: Operator) -> DagDependency | None: # type: i "retries": 1, "retry_delay": {"__type": "timedelta", "__var": 300.0}, "max_retry_delay": {"__type": "timedelta", "__var": 600.0}, - "sla": {"__type": "timedelta", "__var": 100.0}, }, }, "start_date": 1564617600.0, @@ -179,7 +178,6 @@ def detect_task_dependencies(task: Operator) -> DagDependency | None: # type: i "retries": 1, "retry_delay": 300.0, "max_retry_delay": 600.0, - "sla": 100.0, "downstream_task_ids": [], "_is_empty": False, "ui_color": "#f0ede4", @@ -218,7 +216,6 @@ def detect_task_dependencies(task: Operator) -> DagDependency | None: # type: i "retries": 1, "retry_delay": 300.0, "max_retry_delay": 600.0, - "sla": 100.0, "downstream_task_ids": [], "_is_empty": False, "_operator_extra_links": [{"tests.test_utils.mock_operators.CustomOpLink": {}}], @@ -290,7 +287,6 @@ def make_simple_dag(): "retry_delay": timedelta(minutes=5), "max_retry_delay": timedelta(minutes=10), "depends_on_past": False, - "sla": timedelta(seconds=100), }, start_date=datetime(2019, 8, 1), is_paused_upon_creation=False, @@ -1299,7 +1295,6 @@ def test_no_new_fields_added_to_base_operator(self): "retry_delay": timedelta(0, 300), "retry_exponential_backoff": False, "run_as_user": None, - "sla": None, "task_id": "10", "trigger_rule": "all_success", "wait_for_downstream": False, From 51a0351d6f1c8052b0cdf1d3223c417d9568ac26 Mon Sep 17 00:00:00 2001 From: smsm1-ito Date: Wed, 25 Sep 2024 18:24:06 +0100 Subject: [PATCH 021/802] #42442 Make the AWS logging faster by reducing the amount of sleep (#42449) Longer running tasks with lots of logs were slow to write all of the logs to AWS. This reduces the sleep delay to prevent the issue. --- airflow/providers/amazon/aws/utils/task_log_fetcher.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/airflow/providers/amazon/aws/utils/task_log_fetcher.py b/airflow/providers/amazon/aws/utils/task_log_fetcher.py index 83c42f685792a..5a344b507e8ce 100644 --- a/airflow/providers/amazon/aws/utils/task_log_fetcher.py +++ b/airflow/providers/amazon/aws/utils/task_log_fetcher.py @@ -70,7 +70,7 @@ def run(self) -> None: # timestamp) # When a slight delay is added before logging the event, that solves the issue # See https://github.com/apache/airflow/issues/40875 - time.sleep(0.1) + time.sleep(0.001) self.logger.info(self.event_to_str(log_event)) prev_timestamp_event = current_timestamp_event From 69877ae295e6e99b579b449271a0ff531dcb8edd Mon Sep 17 00:00:00 2001 From: "D. Ferruzzi" Date: Wed, 25 Sep 2024 10:52:22 -0700 Subject: [PATCH 022/802] bugfix: create_vector_index task gets marked successful even when it fails (#42472) --- .../amazon/aws/example_bedrock_retrieve_and_generate.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/tests/system/providers/amazon/aws/example_bedrock_retrieve_and_generate.py b/tests/system/providers/amazon/aws/example_bedrock_retrieve_and_generate.py index 224ddc21d8e87..a1d1211da4c46 100644 --- a/tests/system/providers/amazon/aws/example_bedrock_retrieve_and_generate.py +++ b/tests/system/providers/amazon/aws/example_bedrock_retrieve_and_generate.py @@ -261,7 +261,10 @@ def create_vector_index(index_name: str, collection_id: str, region: str): ) log.debug(e) retries -= 1 - sleep(2) + if retries: + sleep(2) + else: + raise @task From 6ec110671dc8a477748cea59d053fa5ddaded55d Mon Sep 17 00:00:00 2001 From: Dmitry Pustoshilov Date: Wed, 25 Sep 2024 21:11:19 +0300 Subject: [PATCH 023/802] docs: fix Executor alias syntax (#42471) wrong example due to https://github.com/apache/airflow/blob/193defd2898772e7e989cbee85815d49e9f0d8f0/airflow/executors/executor_loader.py#L116 --- docs/apache-airflow/core-concepts/executor/index.rst | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/apache-airflow/core-concepts/executor/index.rst b/docs/apache-airflow/core-concepts/executor/index.rst index 1bb11f2335ab6..c7ad952f21d62 100644 --- a/docs/apache-airflow/core-concepts/executor/index.rst +++ b/docs/apache-airflow/core-concepts/executor/index.rst @@ -154,7 +154,7 @@ To make it easier to specify executors on tasks and DAGs, executor configuration .. code-block:: ini [core] - executor = 'LocalExecutor,my.custom.module.ExecutorClass:ShortName' + executor = 'LocalExecutor,ShortName:my.custom.module.ExecutorClass' .. note:: If a DAG specifies a task to use an executor that is not configured, the DAG will fail to parse and a warning dialog will be shown in the Airflow UI. Please ensure that all executors you wish to use are specified in Airflow configuration on *any* host/container that is running an Airflow component (scheduler, workers, etc). From 742a92c430e66f1b8afe36d9a9f8f343cb703897 Mon Sep 17 00:00:00 2001 From: Jens Scheffler <95105677+jscheffl@users.noreply.github.com> Date: Wed, 25 Sep 2024 20:12:03 +0200 Subject: [PATCH 024/802] Bugfix task execution from runner in Windows (#42426) --- airflow/jobs/local_task_job_runner.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/airflow/jobs/local_task_job_runner.py b/airflow/jobs/local_task_job_runner.py index 95a471f239a66..cdc3c1b624694 100644 --- a/airflow/jobs/local_task_job_runner.py +++ b/airflow/jobs/local_task_job_runner.py @@ -261,8 +261,8 @@ def handle_task_exit(self, return_code: int) -> None: _set_task_deferred_context_var() else: message = f"Task exited with return code {return_code}" - if return_code == -signal.SIGKILL: - message += "For more information, see https://airflow.apache.org/docs/apache-airflow/stable/troubleshooting.html#LocalTaskJob-killed" + if not IS_WINDOWS and return_code == -signal.SIGKILL: + message += ". For more information, see https://airflow.apache.org/docs/apache-airflow/stable/troubleshooting.html#LocalTaskJob-killed" self.log.info(message) if not (self.task_instance.test_mode or is_deferral): From 8b4d2bfd0a3180af227da1d49be7b5401bdd3492 Mon Sep 17 00:00:00 2001 From: Howard Yoo <32691630+howardyoo@users.noreply.github.com> Date: Wed, 25 Sep 2024 13:12:52 -0500 Subject: [PATCH 025/802] Fix the span link of task instance to point to the correct span in the scheduler_job_loop (#42430) --- airflow/traces/otel_tracer.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/airflow/traces/otel_tracer.py b/airflow/traces/otel_tracer.py index 1f87e458d3c1b..c6d493db1427a 100644 --- a/airflow/traces/otel_tracer.py +++ b/airflow/traces/otel_tracer.py @@ -199,7 +199,12 @@ def start_span_from_taskinstance( _links.append( Link( - context=trace.get_current_span().get_span_context(), + context=SpanContext( + trace_id=trace.get_current_span().get_span_context().trace_id, + span_id=span_id, + is_remote=True, + trace_flags=TraceFlags(0x01), + ), attributes={"meta.annotation_type": "link", "from": "parenttrace"}, ) ) From 8641a341fde63a8fd0d381221e3b44e736ee40ca Mon Sep 17 00:00:00 2001 From: max <42827971+moiseenkov@users.noreply.github.com> Date: Wed, 25 Sep 2024 18:20:58 +0000 Subject: [PATCH 026/802] Remove deprecated CloudSQL HA functionality from the system test (#42461) --- .../operators/cloud/cloud_sql.rst | 19 --------- .../cloud/cloud_sql/example_cloud_sql.py | 40 ------------------- 2 files changed, 59 deletions(-) diff --git a/docs/apache-airflow-providers-google/operators/cloud/cloud_sql.rst b/docs/apache-airflow-providers-google/operators/cloud/cloud_sql.rst index ec334c0895513..42a32712867d9 100644 --- a/docs/apache-airflow-providers-google/operators/cloud/cloud_sql.rst +++ b/docs/apache-airflow-providers-google/operators/cloud/cloud_sql.rst @@ -180,15 +180,6 @@ it will be retrieved from the Google Cloud connection used. Both variants are sh :start-after: [START howto_operator_cloudsql_delete] :end-before: [END howto_operator_cloudsql_delete] -Note: If the instance has read or failover replicas you need to delete them before you delete the primary instance. -Replicas are deleted the same way as primary instances: - -.. exampleinclude:: /../../tests/system/providers/google/cloud/cloud_sql/example_cloud_sql.py - :language: python - :dedent: 4 - :start-after: [START howto_operator_cloudsql_replicas_delete] - :end-before: [END howto_operator_cloudsql_replicas_delete] - Templating """""""""" @@ -393,16 +384,6 @@ Example body defining the instance with failover replica: :start-after: [START howto_operator_cloudsql_create_body] :end-before: [END howto_operator_cloudsql_create_body] -Example body defining read replica for the instance above: - -.. exampleinclude:: /../../tests/system/providers/google/cloud/cloud_sql/example_cloud_sql.py - :language: python - :start-after: [START howto_operator_cloudsql_create_replica] - :end-before: [END howto_operator_cloudsql_create_replica] - -Note: Failover replicas are created together with the instance in a single task. -Read replicas need to be created in separate tasks. - Using the operator """""""""""""""""" diff --git a/tests/system/providers/google/cloud/cloud_sql/example_cloud_sql.py b/tests/system/providers/google/cloud/cloud_sql/example_cloud_sql.py index d4b0d39e2e2d8..41d374f71488c 100644 --- a/tests/system/providers/google/cloud/cloud_sql/example_cloud_sql.py +++ b/tests/system/providers/google/cloud/cloud_sql/example_cloud_sql.py @@ -62,8 +62,6 @@ FILE_URI = f"gs://{BUCKET_NAME}/{FILE_NAME}" FILE_URI_DEFERRABLE = f"gs://{BUCKET_NAME}/{FILE_NAME_DEFERRABLE}" -FAILOVER_REPLICA_NAME = f"{INSTANCE_NAME}-failover-replica" -READ_REPLICA_NAME = f"{INSTANCE_NAME}-read-replica" CLONED_INSTANCE_NAME = f"{INSTANCE_NAME}-clone" # Bodies below represent Cloud SQL instance resources: @@ -86,30 +84,15 @@ "locationPreference": {"zone": "europe-west4-a"}, "maintenanceWindow": {"hour": 5, "day": 7, "updateTrack": "canary"}, "pricingPlan": "PER_USE", - "replicationType": "ASYNCHRONOUS", "storageAutoResize": True, "storageAutoResizeLimit": 0, "userLabels": {"my-key": "my-value"}, }, - "failoverReplica": {"name": FAILOVER_REPLICA_NAME}, "databaseVersion": "MYSQL_5_7", "region": "europe-west4", } # [END howto_operator_cloudsql_create_body] -# [START howto_operator_cloudsql_create_replica] -read_replica_body = { - "name": READ_REPLICA_NAME, - "settings": { - "tier": "db-n1-standard-1", - }, - "databaseVersion": "MYSQL_5_7", - "region": "europe-west4", - "masterInstanceName": INSTANCE_NAME, -} -# [END howto_operator_cloudsql_create_replica] - - # [START howto_operator_cloudsql_patch_body] patch_body = { "name": INSTANCE_NAME, @@ -169,12 +152,6 @@ ) # [END howto_operator_cloudsql_create] - sql_instance_read_replica_create = CloudSQLCreateInstanceOperator( - body=read_replica_body, - instance=READ_REPLICA_NAME, - task_id="sql_instance_read_replica_create", - ) - # ############################################## # # ### MODIFYING INSTANCE AND ITS DATABASE ###### # # ############################################## # @@ -277,20 +254,6 @@ # ### INSTANCES TEAR DOWN ###################### # # ############################################## # - # [START howto_operator_cloudsql_replicas_delete] - sql_instance_failover_replica_delete_task = CloudSQLDeleteInstanceOperator( - instance=FAILOVER_REPLICA_NAME, - task_id="sql_instance_failover_replica_delete_task", - trigger_rule=TriggerRule.ALL_DONE, - ) - - sql_instance_read_replica_delete_task = CloudSQLDeleteInstanceOperator( - instance=READ_REPLICA_NAME, - task_id="sql_instance_read_replica_delete_task", - trigger_rule=TriggerRule.ALL_DONE, - ) - # [END howto_operator_cloudsql_replicas_delete] - sql_instance_clone_delete_task = CloudSQLDeleteInstanceOperator( instance=CLONED_INSTANCE_NAME, task_id="sql_instance_clone_delete_task", @@ -312,7 +275,6 @@ create_bucket # TEST BODY >> sql_instance_create_task - >> sql_instance_read_replica_create >> sql_instance_patch_task >> sql_db_create_task >> sql_db_patch_task @@ -323,8 +285,6 @@ >> sql_import_task >> sql_instance_clone >> sql_db_delete_task - >> sql_instance_failover_replica_delete_task - >> sql_instance_read_replica_delete_task >> sql_instance_clone_delete_task >> sql_instance_delete_task # TEST TEARDOWN From 00c2f3b09e388a7f515e45e8a33446603adb2786 Mon Sep 17 00:00:00 2001 From: Niko Oliveira Date: Wed, 25 Sep 2024 12:54:19 -0700 Subject: [PATCH 027/802] Small fix to AWS AVP cli init script (#42479) --- .../providers/amazon/aws/auth_manager/cli/avp_commands.py | 7 ++----- 1 file changed, 2 insertions(+), 5 deletions(-) diff --git a/airflow/providers/amazon/aws/auth_manager/cli/avp_commands.py b/airflow/providers/amazon/aws/auth_manager/cli/avp_commands.py index 62452b722977c..fcd9bddaceded 100644 --- a/airflow/providers/amazon/aws/auth_manager/cli/avp_commands.py +++ b/airflow/providers/amazon/aws/auth_manager/cli/avp_commands.py @@ -18,7 +18,6 @@ from __future__ import annotations -import json import logging from pathlib import Path from typing import TYPE_CHECKING @@ -64,10 +63,8 @@ def init_avp(args): _set_schema(client, policy_store_id, args) if not args.dry_run: - print( - "Please set configs below in Airflow configuration under AIRFLOW__AWS_AUTH_MANAGER__." - ) - print(json.dumps({"avp_policy_store_id": policy_store_id}, indent=4)) + print("Please set configs below in Airflow configuration:\n") + print(f"AIRFLOW__AWS_AUTH_MANAGER__AVP_POLICY_STORE_ID={policy_store_id}\n") @cli_utils.action_cli From 6b2d67b9bbca0370d9d971f6f48353f12cad8867 Mon Sep 17 00:00:00 2001 From: Jens Scheffler <95105677+jscheffl@users.noreply.github.com> Date: Wed, 25 Sep 2024 23:56:58 +0200 Subject: [PATCH 028/802] Do not attempt to provide not stringified objects to UI via xcom if pickling is active (#42388) * Do not attempt to provide not stringified objects to UI via xcom if pickling is active * Add pytest --- .../api_connexion/endpoints/xcom_endpoint.py | 2 +- airflow/api_connexion/openapi/v1.yaml | 2 ++ .../endpoints/test_xcom_endpoint.py | 30 +++++++++++++++++++ 3 files changed, 33 insertions(+), 1 deletion(-) diff --git a/airflow/api_connexion/endpoints/xcom_endpoint.py b/airflow/api_connexion/endpoints/xcom_endpoint.py index 59fa9f5acaaa5..5ba0ffa71594d 100644 --- a/airflow/api_connexion/endpoints/xcom_endpoint.py +++ b/airflow/api_connexion/endpoints/xcom_endpoint.py @@ -125,7 +125,7 @@ def get_xcom_entry( stub.value = XCom.deserialize_value(stub) item = stub - if stringify: + if stringify or conf.getboolean("core", "enable_xcom_pickling"): return xcom_schema_string.dump(item) return xcom_schema_native.dump(item) diff --git a/airflow/api_connexion/openapi/v1.yaml b/airflow/api_connexion/openapi/v1.yaml index 07cb7fcb747a6..0c4b0414775f1 100644 --- a/airflow/api_connexion/openapi/v1.yaml +++ b/airflow/api_connexion/openapi/v1.yaml @@ -2040,6 +2040,8 @@ paths: If set to true (default) the Any value will be returned as string, e.g. a Python representation of a dict. If set to false it will return the raw data as dict, list, string or whatever was stored. + This parameter is not meaningful when using XCom pickling, then it is always returned as string. + *New in version 2.10.0* responses: "200": diff --git a/tests/api_connexion/endpoints/test_xcom_endpoint.py b/tests/api_connexion/endpoints/test_xcom_endpoint.py index 9f2d652500694..7a51714c5b299 100644 --- a/tests/api_connexion/endpoints/test_xcom_endpoint.py +++ b/tests/api_connexion/endpoints/test_xcom_endpoint.py @@ -174,6 +174,36 @@ def test_should_respond_200_native(self): "value": {"key": "value"}, } + @conf_vars({("core", "enable_xcom_pickling"): "True"}) + def test_should_respond_200_native_for_pickled(self): + dag_id = "test-dag-id" + task_id = "test-task-id" + execution_date = "2005-04-02T00:00:00+00:00" + xcom_key = "test-xcom-key" + execution_date_parsed = parse_execution_date(execution_date) + run_id = DagRun.generate_run_id(DagRunType.MANUAL, execution_date_parsed) + value_non_serializable_key = {("201009_NB502104_0421_AHJY23BGXG (SEQ_WF: 138898)", None): 82359} + self._create_xcom_entry( + dag_id, run_id, execution_date_parsed, task_id, xcom_key, {"key": value_non_serializable_key} + ) + response = self.client.get( + f"/api/v1/dags/{dag_id}/dagRuns/{run_id}/taskInstances/{task_id}/xcomEntries/{xcom_key}", + environ_overrides={"REMOTE_USER": "test"}, + ) + assert 200 == response.status_code + + current_data = response.json + current_data["timestamp"] = "TIMESTAMP" + assert current_data == { + "dag_id": dag_id, + "execution_date": execution_date, + "key": xcom_key, + "task_id": task_id, + "map_index": -1, + "timestamp": "TIMESTAMP", + "value": f"{{'key': {str(value_non_serializable_key)}}}", + } + def test_should_raise_404_for_non_existent_xcom(self): dag_id = "test-dag-id" task_id = "test-task-id" From 200bff424f110a97375bd25ff040dac40fa26c5e Mon Sep 17 00:00:00 2001 From: Niko Oliveira Date: Wed, 25 Sep 2024 15:09:06 -0700 Subject: [PATCH 029/802] Remove identity center auth manager cli (#42481) * Remove identity center auth manager cli The CLI command to setup identity center could only setup part of the required resources, since adding an application must be done from the console. As of November 15, 2023 it is now required to have an AWS Organization setup to create the required type of Identity Center Instance. The script would have to be change majorly to achieve this but it is also something that should be done with great care and intention since creating an organization in your AWS account has implications. If we automate it, many users won't know it's being created. Instead have users run through the wizard provided in the AWS console. * Missing test change --- .../amazon/aws/auth_manager/cli/definition.py | 6 - .../aws/auth_manager/cli/idc_commands.py | 153 ------------------ .../auth-manager/setup/identity-center.rst | 46 ++---- .../aws/auth_manager/cli/test_definition.py | 2 +- .../aws/auth_manager/cli/test_idc_commands.py | 140 ---------------- 5 files changed, 10 insertions(+), 337 deletions(-) delete mode 100644 airflow/providers/amazon/aws/auth_manager/cli/idc_commands.py delete mode 100644 tests/providers/amazon/aws/auth_manager/cli/test_idc_commands.py diff --git a/airflow/providers/amazon/aws/auth_manager/cli/definition.py b/airflow/providers/amazon/aws/auth_manager/cli/definition.py index bb1236d5c4c94..b5f37136f51d9 100644 --- a/airflow/providers/amazon/aws/auth_manager/cli/definition.py +++ b/airflow/providers/amazon/aws/auth_manager/cli/definition.py @@ -55,12 +55,6 @@ ################ AWS_AUTH_MANAGER_COMMANDS = ( - ActionCommand( - name="init-identity-center", - help="Initialize AWS IAM identity Center resources to be used by AWS manager", - func=lazy_load_command("airflow.providers.amazon.aws.auth_manager.cli.idc_commands.init_idc"), - args=(ARG_INSTANCE_NAME, ARG_APPLICATION_NAME, ARG_DRY_RUN, ARG_VERBOSE), - ), ActionCommand( name="init-avp", help="Initialize Amazon Verified resources to be used by AWS manager", diff --git a/airflow/providers/amazon/aws/auth_manager/cli/idc_commands.py b/airflow/providers/amazon/aws/auth_manager/cli/idc_commands.py deleted file mode 100644 index c4901351b2cff..0000000000000 --- a/airflow/providers/amazon/aws/auth_manager/cli/idc_commands.py +++ /dev/null @@ -1,153 +0,0 @@ -# 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. -"""User sub-commands.""" - -from __future__ import annotations - -import logging -import sys -from typing import TYPE_CHECKING - -import boto3 -from botocore.exceptions import ClientError - -from airflow.configuration import conf -from airflow.exceptions import AirflowOptionalProviderFeatureException -from airflow.providers.amazon.aws.auth_manager.constants import CONF_REGION_NAME_KEY, CONF_SECTION_NAME -from airflow.utils import cli as cli_utils - -try: - from airflow.utils.providers_configuration_loader import providers_configuration_loaded -except ImportError: - raise AirflowOptionalProviderFeatureException( - "Failed to import avp_commands. This feature is only available in Airflow " - "version >= 2.8.0 where Auth Managers are introduced." - ) - -if TYPE_CHECKING: - from botocore.client import BaseClient - -log = logging.getLogger(__name__) - - -@cli_utils.action_cli -@providers_configuration_loaded -def init_idc(args): - """Initialize AWS IAM Identity Center resources.""" - client = _get_client() - - # Create the instance if needed - instance_arn = _create_instance(client, args) - - # Create the application if needed - _create_application(client, instance_arn, args) - - if not args.dry_run: - print("AWS IAM Identity Center resources created successfully.") - - -def _get_client(): - """Return AWS IAM Identity Center client.""" - region_name = conf.get(CONF_SECTION_NAME, CONF_REGION_NAME_KEY) - return boto3.client("sso-admin", region_name=region_name) - - -def _create_instance(client: BaseClient, args) -> str | None: - """Create if needed AWS IAM Identity Center instance.""" - instances = client.list_instances() - - if args.verbose: - log.debug("Instances found: %s", instances) - - if len(instances["Instances"]) > 0: - print( - f"There is already an instance configured in AWS IAM Identity Center: '{instances['Instances'][0]['InstanceArn']}'. " - "No need to create a new one." - ) - return instances["Instances"][0]["InstanceArn"] - else: - print("No instance configured in AWS IAM Identity Center, creating one.") - if args.dry_run: - print("Dry run, not creating the instance.") - return None - - response = client.create_instance(Name=args.instance_name) - if args.verbose: - log.debug("Response from create_instance: %s", response) - - print(f"Instance created: '{response['InstanceArn']}'") - - return response["InstanceArn"] - - -def _create_application(client: BaseClient, instance_arn: str | None, args) -> str | None: - """Create if needed AWS IAM identity Center application.""" - paginator = client.get_paginator("list_applications") - pages = paginator.paginate(InstanceArn=instance_arn or "") - applications = [application for page in pages for application in page["Applications"]] - existing_applications = [ - application for application in applications if application["Name"] == args.application_name - ] - - if args.verbose: - log.debug("Applications found: %s", applications) - log.debug("Existing applications found: %s", existing_applications) - - if len(existing_applications) > 0: - print( - f"There is already an application named '{args.application_name}' in AWS IAM Identity Center: '{existing_applications[0]['ApplicationArn']}'. " - "Using this application." - ) - return existing_applications[0]["ApplicationArn"] - else: - print(f"No application named {args.application_name} found, creating one.") - if args.dry_run: - print("Dry run, not creating the application.") - return None - - try: - response = client.create_application( - ApplicationProviderArn="arn:aws:sso::aws:applicationProvider/custom-saml", - Description="Application automatically created through the Airflow CLI. This application is used to access Airflow environment.", - InstanceArn=instance_arn, - Name=args.application_name, - PortalOptions={ - "SignInOptions": { - "Origin": "IDENTITY_CENTER", - }, - "Visibility": "ENABLED", - }, - Status="ENABLED", - ) - if args.verbose: - log.debug("Response from create_application: %s", response) - except ClientError as e: - # This is needed because as of today, the create_application in AWS Identity Center does not support SAML application - # Remove this part when it is supported - if "is not supported for this action" in e.response["Error"]["Message"]: - print( - "*************************************************************************\n" - "* ACTION REQUIRED *\n" - "* Creation of SAML applications is only supported in AWS console today. *\n" - "* Please create the application through the console. *\n" - "*************************************************************************\n" - ) - sys.exit(1) - - print(f"Application created: '{response['ApplicationArn']}'") - - return response["ApplicationArn"] diff --git a/docs/apache-airflow-providers-amazon/auth-manager/setup/identity-center.rst b/docs/apache-airflow-providers-amazon/auth-manager/setup/identity-center.rst index acf3727bf9c7f..ff2dc6295eb83 100644 --- a/docs/apache-airflow-providers-amazon/auth-manager/setup/identity-center.rst +++ b/docs/apache-airflow-providers-amazon/auth-manager/setup/identity-center.rst @@ -27,51 +27,23 @@ Create resources ================ The AWS auth manager needs two resources in AWS IAM Identity Center: an instance and an application. -You can create them either through the provided CLI command or manually. +You can must create them manually. -Create resources with CLI -------------------------- - -.. note:: - The CLI command is not compatible with AWS accounts that are managed through AWS organizations. - If your AWS account is managed through an AWS organization, please follow the - :ref:`manual configuration `. - -.. note:: - To create all necessary resources for the AWS Auth Manager, you can utilize the CLI command provided as part of the - AWS auth manager. Before executing the command, ensure the AWS auth manager is configured as the auth manager - for the Airflow instance. See :doc:`/auth-manager/setup/config`. - -To create the resources, please run the following command: - -.. code-block:: bash - - airflow aws-auth-manager init-identity-center - -The CLI command will ask you to create any resources manually if they cannot be automatically created. Please look carefully at the CLI command output to understand which resource(s) -have or have not been created successfully. The resource(s) which have not been successfully created need to be -:ref:`created manually `. - -If the error message below is raised, please create the AWS IAM Identity Center application through the console -following :ref:`these instructions `: :: - - Creation of SAML applications is only supported in AWS console today. Please create the application through the console. - -.. _identity_center_manual_configuration: +Create the instance +------------------- -Create resources manually -------------------------- +The AWS auth manager leverages SAML 2.0 as the underlying technology powering authentication against AWS Identity Center. -Create the instance -~~~~~~~~~~~~~~~~~~~ +There are several instance types, but only Organization level instances can use SAML 2.0 applications. See more details +about instances types `here `_. -Please follow `AWS documentation `_ -to create the AWS IAM Identity Center instance. +Please follow `AWS documentation `_ +to create the AWS IAM Identity Center instance at the organization level. .. _identity_center_manual_configuration_application: Create the application -~~~~~~~~~~~~~~~~~~~~~~ +---------------------- Please follow the instructions below to create the AWS IAM Identity Center application. diff --git a/tests/providers/amazon/aws/auth_manager/cli/test_definition.py b/tests/providers/amazon/aws/auth_manager/cli/test_definition.py index 5866aa594f8ec..079df886f6039 100644 --- a/tests/providers/amazon/aws/auth_manager/cli/test_definition.py +++ b/tests/providers/amazon/aws/auth_manager/cli/test_definition.py @@ -21,4 +21,4 @@ class TestAwsCliDefinition: def test_aws_auth_manager_cli_commands(self): - assert len(AWS_AUTH_MANAGER_COMMANDS) == 3 + assert len(AWS_AUTH_MANAGER_COMMANDS) == 2 diff --git a/tests/providers/amazon/aws/auth_manager/cli/test_idc_commands.py b/tests/providers/amazon/aws/auth_manager/cli/test_idc_commands.py deleted file mode 100644 index 394704474f1bb..0000000000000 --- a/tests/providers/amazon/aws/auth_manager/cli/test_idc_commands.py +++ /dev/null @@ -1,140 +0,0 @@ -# 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 importlib -from unittest.mock import Mock, patch - -import pytest - -from airflow.cli import cli_parser -from airflow.providers.amazon.aws.auth_manager.cli.idc_commands import init_idc -from tests.test_utils.compat import AIRFLOW_V_2_8_PLUS -from tests.test_utils.config import conf_vars - -mock_boto3 = Mock() - -pytestmark = [ - pytest.mark.skipif(not AIRFLOW_V_2_8_PLUS, reason="Test requires Airflow 2.8+"), - pytest.mark.skip_if_database_isolation_mode, -] - - -@pytest.mark.db_test -class TestIdcCommands: - def setup_method(self): - mock_boto3.reset_mock() - - @classmethod - def setup_class(cls): - with conf_vars( - { - ( - "core", - "auth_manager", - ): "airflow.providers.amazon.aws.auth_manager.aws_auth_manager.AwsAuthManager" - } - ): - importlib.reload(cli_parser) - cls.arg_parser = cli_parser.get_parser() - - @pytest.mark.parametrize( - "dry_run, verbose", - [ - (False, False), - (True, True), - ], - ) - @patch("airflow.providers.amazon.aws.auth_manager.cli.idc_commands._get_client") - def test_init_idc_with_no_existing_resources(self, mock_get_client, dry_run, verbose): - mock_get_client.return_value = mock_boto3 - - instance_name = "test-instance" - instance_arn = "test-instance-arn" - application_name = "test-application" - application_arn = "test-application-arn" - - paginator = Mock() - paginator.paginate.return_value = [] - - mock_boto3.list_instances.return_value = {"Instances": []} - mock_boto3.create_instance.return_value = {"InstanceArn": instance_arn} - mock_boto3.get_paginator.return_value = paginator - mock_boto3.create_application.return_value = {"ApplicationArn": application_arn} - - with conf_vars({("database", "check_migrations"): "False"}): - params = [ - "aws-auth-manager", - "init-identity-center", - "--instance-name", - instance_name, - "--application-name", - application_name, - ] - if dry_run: - params.append("--dry-run") - if verbose: - params.append("--verbose") - init_idc(self.arg_parser.parse_args(params)) - - mock_boto3.list_instances.assert_called_once_with() - if not dry_run: - mock_boto3.create_instance.assert_called_once_with(Name=instance_name) - mock_boto3.create_application.assert_called_once() - - @pytest.mark.parametrize( - "dry_run, verbose", - [ - (False, False), - (True, True), - ], - ) - @patch("airflow.providers.amazon.aws.auth_manager.cli.idc_commands._get_client") - def test_init_idc_with_existing_resources(self, mock_get_client, dry_run, verbose): - mock_get_client.return_value = mock_boto3 - - instance_name = "test-instance" - instance_arn = "test-instance-arn" - application_name = "test-application" - application_arn = "test-application-arn" - - paginator = Mock() - paginator.paginate.return_value = [ - {"Applications": [{"Name": application_name, "ApplicationArn": application_arn}]} - ] - - mock_boto3.list_instances.return_value = {"Instances": [{"InstanceArn": instance_arn}]} - mock_boto3.get_paginator.return_value = paginator - - with conf_vars({("database", "check_migrations"): "False"}): - params = [ - "aws-auth-manager", - "init-identity-center", - "--instance-name", - instance_name, - "--application-name", - application_name, - ] - if dry_run: - params.append("--dry-run") - if verbose: - params.append("--verbose") - init_idc(self.arg_parser.parse_args(params)) - - mock_boto3.list_instances.assert_called_once_with() - mock_boto3.create_instance.assert_not_called() - mock_boto3.create_application.assert_not_called() From e89d398cd1e1412a5b77c948812150657e1afffd Mon Sep 17 00:00:00 2001 From: Daniel Standish <15932138+dstandish@users.noreply.github.com> Date: Wed, 25 Sep 2024 18:31:55 -0700 Subject: [PATCH 030/802] Remove jhtimmins as code owner for security (#42477) --- .github/CODEOWNERS | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS index 75ceb48d78504..a0a7b82331be6 100644 --- a/.github/CODEOWNERS +++ b/.github/CODEOWNERS @@ -33,9 +33,9 @@ /airflow/ui/ @bbovenzi @pierrejeambrun @ryanahamilton @jscheffl # Security/Permissions -/airflow/api_connexion/security.py @jhtimmins -/airflow/security/permissions.py @jhtimmins -/airflow/www/security.py @jhtimmins +/airflow/api_connexion/security.py @vincbeck +/airflow/security/permissions.py @vincbeck +/airflow/www/security.py @vincbeck # Calendar/Timetables /airflow/timetables/ @uranusjr From 76db850ef2a6e9cf6fd57f5c4d83473fd11cfee8 Mon Sep 17 00:00:00 2001 From: Daniel Standish <15932138+dstandish@users.noreply.github.com> Date: Wed, 25 Sep 2024 18:32:25 -0700 Subject: [PATCH 031/802] Simplify expression for get_permitted_dag_ids query (#42484) --- airflow/providers/fab/auth_manager/fab_auth_manager.py | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/airflow/providers/fab/auth_manager/fab_auth_manager.py b/airflow/providers/fab/auth_manager/fab_auth_manager.py index 2de8db2f56414..3d0f102650935 100644 --- a/airflow/providers/fab/auth_manager/fab_auth_manager.py +++ b/airflow/providers/fab/auth_manager/fab_auth_manager.py @@ -342,10 +342,7 @@ def get_permitted_dag_ids( resources.add(resource[len(permissions.RESOURCE_DAG_PREFIX) :]) else: resources.add(resource) - return { - dag.dag_id - for dag in session.execute(select(DagModel.dag_id).where(DagModel.dag_id.in_(resources))) - } + return set(session.scalars(select(DagModel.dag_id).where(DagModel.dag_id.in_(resources)))) @cached_property def security_manager(self) -> FabAirflowSecurityManagerOverride: From 5c65a25232cddd55ed08a6a02df0417fe9838e8d Mon Sep 17 00:00:00 2001 From: Daniel Standish <15932138+dstandish@users.noreply.github.com> Date: Wed, 25 Sep 2024 18:33:16 -0700 Subject: [PATCH 032/802] Split up the return statement in _is_authorized_callback for clarity (#42473) Co-authored-by: Vincent <97131062+vincbeck@users.noreply.github.com> --- airflow/api_connexion/security.py | 13 ++++++------- 1 file changed, 6 insertions(+), 7 deletions(-) diff --git a/airflow/api_connexion/security.py b/airflow/api_connexion/security.py index 7b0a026e095d0..7da83a76168bb 100644 --- a/airflow/api_connexion/security.py +++ b/airflow/api_connexion/security.py @@ -126,13 +126,12 @@ def callback(): if dag_id or access or access_entity: return access - # No DAG id is provided, the user is not authorized to access all DAGs and authorization is done - # on DAG level - # If method is "GET", return whether the user has read access to any DAGs - # If method is "PUT", return whether the user has edit access to any DAGs - return (method == "GET" and any(get_auth_manager().get_permitted_dag_ids(methods=["GET"]))) or ( - method == "PUT" and any(get_auth_manager().get_permitted_dag_ids(methods=["PUT"])) - ) + # dag_id is not provided, and the user is not authorized to access *all* DAGs + # so we check that the user can access at least *one* dag + # but we leave it to the endpoint function to properly restrict access beyond that + if method not in ("GET", "PUT"): + return False + return any(get_auth_manager().get_permitted_dag_ids(methods=[method])) return callback From 2d535c2da23408dec6ac4fa682155c7633614ede Mon Sep 17 00:00:00 2001 From: Daniel Standish <15932138+dstandish@users.noreply.github.com> Date: Thu, 26 Sep 2024 01:33:05 -0700 Subject: [PATCH 033/802] Fix typo in fixture set_auth_role_public (#42488) auto -> auth --- .../test_role_and_permission_endpoint.py | 36 +++++++++---------- tests/providers/fab/auth_manager/conftest.py | 2 +- 2 files changed, 19 insertions(+), 19 deletions(-) diff --git a/tests/providers/fab/auth_manager/api_endpoints/test_role_and_permission_endpoint.py b/tests/providers/fab/auth_manager/api_endpoints/test_role_and_permission_endpoint.py index 77e3107a0b51f..30cfaeb227903 100644 --- a/tests/providers/fab/auth_manager/api_endpoints/test_role_and_permission_endpoint.py +++ b/tests/providers/fab/auth_manager/api_endpoints/test_role_and_permission_endpoint.py @@ -108,11 +108,11 @@ def test_should_raise_403_forbidden(self): assert response.status_code == 403 @pytest.mark.parametrize( - "set_auto_role_public, expected_status_code", + "set_auth_role_public, expected_status_code", (("Public", 403), ("Admin", 200)), - indirect=["set_auto_role_public"], + indirect=["set_auth_role_public"], ) - def test_with_auth_role_public_set(self, set_auto_role_public, expected_status_code): + def test_with_auth_role_public_set(self, set_auth_role_public, expected_status_code): response = self.client.get("/auth/fab/v1/roles/Admin") assert response.status_code == expected_status_code, response.json @@ -146,11 +146,11 @@ def test_should_raise_403_forbidden(self): assert response.status_code == 403 @pytest.mark.parametrize( - "set_auto_role_public, expected_status_code", + "set_auth_role_public, expected_status_code", (("Public", 403), ("Admin", 200)), - indirect=["set_auto_role_public"], + indirect=["set_auth_role_public"], ) - def test_with_auth_role_public_set(self, set_auto_role_public, expected_status_code): + def test_with_auth_role_public_set(self, set_auth_role_public, expected_status_code): response = self.client.get("/auth/fab/v1/roles") assert response.status_code == expected_status_code, response.json @@ -208,11 +208,11 @@ def test_should_raise_403_forbidden(self): assert response.status_code == 403 @pytest.mark.parametrize( - "set_auto_role_public, expected_status_code", + "set_auth_role_public, expected_status_code", (("Public", 403), ("Admin", 200)), - indirect=["set_auto_role_public"], + indirect=["set_auth_role_public"], ) - def test_with_auth_role_public_set(self, set_auto_role_public, expected_status_code): + def test_with_auth_role_public_set(self, set_auth_role_public, expected_status_code): response = self.client.get("/auth/fab/v1/permissions") assert response.status_code == expected_status_code, response.json @@ -346,11 +346,11 @@ def test_should_raise_403_forbidden(self): assert response.status_code == 403 @pytest.mark.parametrize( - "set_auto_role_public, expected_status_code", + "set_auth_role_public, expected_status_code", (("Public", 403), ("Admin", 200)), - indirect=["set_auto_role_public"], + indirect=["set_auth_role_public"], ) - def test_with_auth_role_public_set(self, set_auto_role_public, expected_status_code): + def test_with_auth_role_public_set(self, set_auth_role_public, expected_status_code): payload = { "name": "Test2", "actions": [{"resource": {"name": "Connections"}, "action": {"name": "can_create"}}], @@ -393,11 +393,11 @@ def test_should_raise_403_forbidden(self): assert response.status_code == 403 @pytest.mark.parametrize( - "set_auto_role_public, expected_status_code", + "set_auth_role_public, expected_status_code", (("Public", 403), ("Admin", 204)), - indirect=["set_auto_role_public"], + indirect=["set_auth_role_public"], ) - def test_with_auth_role_public_set(self, set_auto_role_public, expected_status_code): + def test_with_auth_role_public_set(self, set_auth_role_public, expected_status_code): role = create_role(self.app, "mytestrole") response = self.client.delete(f"/auth/fab/v1/roles/{role.name}") assert response.status_code == expected_status_code, response.location @@ -579,11 +579,11 @@ def test_should_raise_403_forbidden(self): assert response.status_code == 403 @pytest.mark.parametrize( - "set_auto_role_public, expected_status_code", + "set_auth_role_public, expected_status_code", (("Public", 403), ("Admin", 200)), - indirect=["set_auto_role_public"], + indirect=["set_auth_role_public"], ) - def test_with_auth_role_public_set(self, set_auto_role_public, expected_status_code): + def test_with_auth_role_public_set(self, set_auth_role_public, expected_status_code): role = create_role(self.app, "mytestrole") response = self.client.patch( f"/auth/fab/v1/roles/{role.name}", diff --git a/tests/providers/fab/auth_manager/conftest.py b/tests/providers/fab/auth_manager/conftest.py index da18f9d6c06be..22c29dd229fa1 100644 --- a/tests/providers/fab/auth_manager/conftest.py +++ b/tests/providers/fab/auth_manager/conftest.py @@ -50,7 +50,7 @@ def factory(): @pytest.fixture -def set_auto_role_public(request): +def set_auth_role_public(request): app = request.getfixturevalue("minimal_app_for_auth_api") auto_role_public = app.config["AUTH_ROLE_PUBLIC"] app.config["AUTH_ROLE_PUBLIC"] = request.param From b151f8e14a12cf6f58322c5f8e6917b0dca45165 Mon Sep 17 00:00:00 2001 From: Pierre Jeambrun Date: Thu, 26 Sep 2024 18:16:12 +0800 Subject: [PATCH 034/802] Migrate patch dag to FastAPI API (#42469) --- .../api_connexion/endpoints/dag_endpoint.py | 1 + airflow/api_fastapi/openapi/v1-generated.yaml | 59 ++++++++++++++++++- airflow/api_fastapi/serializers/dags.py | 10 +++- airflow/api_fastapi/views/public/dags.py | 34 +++++++++-- airflow/ui/openapi-gen/queries/common.ts | 3 + airflow/ui/openapi-gen/queries/queries.ts | 55 ++++++++++++++++- .../ui/openapi-gen/requests/schemas.gen.ts | 19 +++++- .../ui/openapi-gen/requests/services.gen.ts | 32 ++++++++++ airflow/ui/openapi-gen/requests/types.gen.ts | 34 ++++++++++- airflow/ui/src/pages/DagsList.tsx | 4 +- tests/api_fastapi/views/public/test_dags.py | 20 +++++++ 11 files changed, 254 insertions(+), 17 deletions(-) diff --git a/airflow/api_connexion/endpoints/dag_endpoint.py b/airflow/api_connexion/endpoints/dag_endpoint.py index 08d36f9978c33..6fca5ae7c93d5 100644 --- a/airflow/api_connexion/endpoints/dag_endpoint.py +++ b/airflow/api_connexion/endpoints/dag_endpoint.py @@ -141,6 +141,7 @@ def get_dags( raise BadRequest("DAGCollectionSchema error", detail=str(e)) +@mark_fastapi_migration_done @security.requires_access_dag("PUT") @action_logging @provide_session diff --git a/airflow/api_fastapi/openapi/v1-generated.yaml b/airflow/api_fastapi/openapi/v1-generated.yaml index b0037b372bd4e..6d77056d0574d 100644 --- a/airflow/api_fastapi/openapi/v1-generated.yaml +++ b/airflow/api_fastapi/openapi/v1-generated.yaml @@ -123,13 +123,56 @@ paths: application/json: schema: $ref: '#/components/schemas/HTTPValidationError' + /public/dags/{dag_id}: + patch: + tags: + - DAG + summary: Patch Dag + description: Update the specific DAG. + operationId: patch_dag_public_dags__dag_id__patch + parameters: + - name: dag_id + in: path + required: true + schema: + type: string + title: Dag Id + - name: update_mask + in: query + required: false + schema: + anyOf: + - type: array + items: + type: string + - type: 'null' + title: Update Mask + requestBody: + required: true + content: + application/json: + schema: + $ref: '#/components/schemas/DAGPatchBody' + responses: + '200': + description: Successful Response + content: + application/json: + schema: + $ref: '#/components/schemas/DAGResponse' + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' components: schemas: DAGCollectionResponse: properties: dags: items: - $ref: '#/components/schemas/DAGModelResponse' + $ref: '#/components/schemas/DAGResponse' type: array title: Dags total_entries: @@ -141,7 +184,17 @@ components: - total_entries title: DAGCollectionResponse description: DAG Collection serializer for responses. - DAGModelResponse: + DAGPatchBody: + properties: + is_paused: + type: boolean + title: Is Paused + type: object + required: + - is_paused + title: DAGPatchBody + description: Dag Serializer for updatable body. + DAGResponse: properties: dag_id: type: string @@ -292,7 +345,7 @@ components: - next_dagrun_create_after - owners - file_token - title: DAGModelResponse + title: DAGResponse description: DAG serializer for responses. DagTagPydantic: properties: diff --git a/airflow/api_fastapi/serializers/dags.py b/airflow/api_fastapi/serializers/dags.py index 264f549e298a5..59b47bdef9e98 100644 --- a/airflow/api_fastapi/serializers/dags.py +++ b/airflow/api_fastapi/serializers/dags.py @@ -31,7 +31,7 @@ from airflow.serialization.pydantic.dag import DagTagPydantic -class DAGModelResponse(BaseModel): +class DAGResponse(BaseModel): """DAG serializer for responses.""" dag_id: str @@ -82,8 +82,14 @@ def file_token(self) -> str: return serializer.dumps(self.fileloc) +class DAGPatchBody(BaseModel): + """Dag Serializer for updatable body.""" + + is_paused: bool + + class DAGCollectionResponse(BaseModel): """DAG Collection serializer for responses.""" - dags: list[DAGModelResponse] + dags: list[DAGResponse] total_entries: int diff --git a/airflow/api_fastapi/views/public/dags.py b/airflow/api_fastapi/views/public/dags.py index a1957d30739ec..433e5ef862447 100644 --- a/airflow/api_fastapi/views/public/dags.py +++ b/airflow/api_fastapi/views/public/dags.py @@ -17,7 +17,7 @@ from __future__ import annotations -from fastapi import APIRouter, Depends, HTTPException +from fastapi import APIRouter, Depends, HTTPException, Query from sqlalchemy import select from sqlalchemy.orm import Session from typing_extensions import Annotated @@ -34,7 +34,7 @@ QueryTagsFilter, SortParam, ) -from airflow.api_fastapi.serializers.dags import DAGCollectionResponse, DAGModelResponse +from airflow.api_fastapi.serializers.dags import DAGCollectionResponse, DAGPatchBody, DAGResponse from airflow.models import DagModel from airflow.utils.db import get_query_count @@ -43,7 +43,6 @@ @dags_router.get("/dags") async def get_dags( - *, limit: QueryLimit, offset: QueryOffset, tags: QueryTagsFilter, @@ -74,8 +73,35 @@ async def get_dags( try: return DAGCollectionResponse( - dags=[DAGModelResponse.model_validate(dag, from_attributes=True) for dag in dags], + dags=[DAGResponse.model_validate(dag, from_attributes=True) for dag in dags], total_entries=total_entries, ) except ValueError as e: raise HTTPException(400, f"DAGCollectionSchema error: {str(e)}") + + +@dags_router.patch("/dags/{dag_id}") +async def patch_dag( + dag_id: str, + patch_body: DAGPatchBody, + session: Annotated[Session, Depends(get_session)], + update_mask: list[str] | None = Query(None), +) -> DAGResponse: + """Update the specific DAG.""" + dag = session.get(DagModel, dag_id) + + if dag is None: + raise HTTPException(404, f"Dag with id: {dag_id} was not found") + + if update_mask: + if update_mask != ["is_paused"]: + raise HTTPException(400, "Only `is_paused` field can be updated through the REST API") + + else: + update_mask = ["is_paused"] + + for attr_name in update_mask: + attr_value = getattr(patch_body, attr_name) + setattr(dag, attr_name, attr_value) + + return DAGResponse.model_validate(dag, from_attributes=True) diff --git a/airflow/ui/openapi-gen/queries/common.ts b/airflow/ui/openapi-gen/queries/common.ts index 143ec83c55627..2818b48a33e1a 100644 --- a/airflow/ui/openapi-gen/queries/common.ts +++ b/airflow/ui/openapi-gen/queries/common.ts @@ -72,3 +72,6 @@ export const UseDagServiceGetDagsPublicDagsGetKeyFn = ( }, ]), ]; +export type DagServicePatchDagPublicDagsDagIdPatchMutationResult = Awaited< + ReturnType +>; diff --git a/airflow/ui/openapi-gen/queries/queries.ts b/airflow/ui/openapi-gen/queries/queries.ts index 9dce528f2a503..2a0c6b6821978 100644 --- a/airflow/ui/openapi-gen/queries/queries.ts +++ b/airflow/ui/openapi-gen/queries/queries.ts @@ -1,7 +1,13 @@ // generated with @7nohe/openapi-react-query-codegen@1.6.0 -import { useQuery, UseQueryOptions } from "@tanstack/react-query"; +import { + useMutation, + UseMutationOptions, + useQuery, + UseQueryOptions, +} from "@tanstack/react-query"; import { DagService, DatasetService } from "../requests/services.gen"; +import { DAGPatchBody } from "../requests/types.gen"; import * as Common from "./common"; /** @@ -110,3 +116,50 @@ export const useDagServiceGetDagsPublicDagsGet = < }) as TData, ...options, }); +/** + * Patch Dag + * Update the specific DAG. + * @param data The data for the request. + * @param data.dagId + * @param data.requestBody + * @param data.updateMask + * @returns DAGResponse Successful Response + * @throws ApiError + */ +export const useDagServicePatchDagPublicDagsDagIdPatch = < + TData = Common.DagServicePatchDagPublicDagsDagIdPatchMutationResult, + TError = unknown, + TContext = unknown, +>( + options?: Omit< + UseMutationOptions< + TData, + TError, + { + dagId: string; + requestBody: DAGPatchBody; + updateMask?: string[]; + }, + TContext + >, + "mutationFn" + >, +) => + useMutation< + TData, + TError, + { + dagId: string; + requestBody: DAGPatchBody; + updateMask?: string[]; + }, + TContext + >({ + mutationFn: ({ dagId, requestBody, updateMask }) => + DagService.patchDagPublicDagsDagIdPatch({ + dagId, + requestBody, + updateMask, + }) as unknown as Promise, + ...options, + }); diff --git a/airflow/ui/openapi-gen/requests/schemas.gen.ts b/airflow/ui/openapi-gen/requests/schemas.gen.ts index 64cddab30b1c6..83d3670507f78 100644 --- a/airflow/ui/openapi-gen/requests/schemas.gen.ts +++ b/airflow/ui/openapi-gen/requests/schemas.gen.ts @@ -4,7 +4,7 @@ export const $DAGCollectionResponse = { properties: { dags: { items: { - $ref: "#/components/schemas/DAGModelResponse", + $ref: "#/components/schemas/DAGResponse", }, type: "array", title: "Dags", @@ -20,7 +20,20 @@ export const $DAGCollectionResponse = { description: "DAG Collection serializer for responses.", } as const; -export const $DAGModelResponse = { +export const $DAGPatchBody = { + properties: { + is_paused: { + type: "boolean", + title: "Is Paused", + }, + }, + type: "object", + required: ["is_paused"], + title: "DAGPatchBody", + description: "Dag Serializer for updatable body.", +} as const; + +export const $DAGResponse = { properties: { dag_id: { type: "string", @@ -271,7 +284,7 @@ export const $DAGModelResponse = { "owners", "file_token", ], - title: "DAGModelResponse", + title: "DAGResponse", description: "DAG serializer for responses.", } as const; diff --git a/airflow/ui/openapi-gen/requests/services.gen.ts b/airflow/ui/openapi-gen/requests/services.gen.ts index e0786e9137156..a4c36d5990c78 100644 --- a/airflow/ui/openapi-gen/requests/services.gen.ts +++ b/airflow/ui/openapi-gen/requests/services.gen.ts @@ -7,6 +7,8 @@ import type { NextRunDatasetsUiNextRunDatasetsDagIdGetResponse, GetDagsPublicDagsGetData, GetDagsPublicDagsGetResponse, + PatchDagPublicDagsDagIdPatchData, + PatchDagPublicDagsDagIdPatchResponse, } from "./types.gen"; export class DatasetService { @@ -72,4 +74,34 @@ export class DagService { }, }); } + + /** + * Patch Dag + * Update the specific DAG. + * @param data The data for the request. + * @param data.dagId + * @param data.requestBody + * @param data.updateMask + * @returns DAGResponse Successful Response + * @throws ApiError + */ + public static patchDagPublicDagsDagIdPatch( + data: PatchDagPublicDagsDagIdPatchData, + ): CancelablePromise { + return __request(OpenAPI, { + method: "PATCH", + url: "/public/dags/{dag_id}", + path: { + dag_id: data.dagId, + }, + query: { + update_mask: data.updateMask, + }, + body: data.requestBody, + mediaType: "application/json", + errors: { + 422: "Validation Error", + }, + }); + } } diff --git a/airflow/ui/openapi-gen/requests/types.gen.ts b/airflow/ui/openapi-gen/requests/types.gen.ts index 917dca6626c08..2f6bc263d4289 100644 --- a/airflow/ui/openapi-gen/requests/types.gen.ts +++ b/airflow/ui/openapi-gen/requests/types.gen.ts @@ -4,14 +4,21 @@ * DAG Collection serializer for responses. */ export type DAGCollectionResponse = { - dags: Array; + dags: Array; total_entries: number; }; +/** + * Dag Serializer for updatable body. + */ +export type DAGPatchBody = { + is_paused: boolean; +}; + /** * DAG serializer for responses. */ -export type DAGModelResponse = { +export type DAGResponse = { dag_id: string; dag_display_name: string; is_paused: boolean; @@ -83,6 +90,14 @@ export type GetDagsPublicDagsGetData = { export type GetDagsPublicDagsGetResponse = DAGCollectionResponse; +export type PatchDagPublicDagsDagIdPatchData = { + dagId: string; + requestBody: DAGPatchBody; + updateMask?: Array | null; +}; + +export type PatchDagPublicDagsDagIdPatchResponse = DAGResponse; + export type $OpenApiTs = { "/ui/next_run_datasets/{dag_id}": { get: { @@ -116,4 +131,19 @@ export type $OpenApiTs = { }; }; }; + "/public/dags/{dag_id}": { + patch: { + req: PatchDagPublicDagsDagIdPatchData; + res: { + /** + * Successful Response + */ + 200: DAGResponse; + /** + * Validation Error + */ + 422: HTTPValidationError; + }; + }; + }; }; diff --git a/airflow/ui/src/pages/DagsList.tsx b/airflow/ui/src/pages/DagsList.tsx index e93f281c50137..fe764f117e45d 100644 --- a/airflow/ui/src/pages/DagsList.tsx +++ b/airflow/ui/src/pages/DagsList.tsx @@ -31,7 +31,7 @@ import { type ChangeEventHandler, useCallback } from "react"; import { useSearchParams } from "react-router-dom"; import { useDagServiceGetDagsPublicDagsGet } from "openapi/queries"; -import type { DAGModelResponse } from "openapi/requests/types.gen"; +import type { DAGResponse } from "openapi/requests/types.gen"; import { DataTable } from "../components/DataTable"; import { useTableURLState } from "../components/DataTable/useTableUrlState"; @@ -39,7 +39,7 @@ import { QuickFilterButton } from "../components/QuickFilterButton"; import { SearchBar } from "../components/SearchBar"; import { pluralize } from "../utils/pluralize"; -const columns: Array> = [ +const columns: Array> = [ { accessorKey: "dag_id", cell: ({ row }) => row.original.dag_display_name, diff --git a/tests/api_fastapi/views/public/test_dags.py b/tests/api_fastapi/views/public/test_dags.py index dfba5437a8af3..b508a1448352d 100644 --- a/tests/api_fastapi/views/public/test_dags.py +++ b/tests/api_fastapi/views/public/test_dags.py @@ -115,3 +115,23 @@ def test_get_dags(test_client, query_params, expected_total_entries, expected_id assert body["total_entries"] == expected_total_entries assert [dag["dag_id"] for dag in body["dags"]] == expected_ids + + +@pytest.mark.parametrize( + "query_params, dag_id, body, expected_status_code, expected_is_paused", + [ + ({}, "fake_dag_id", {"is_paused": True}, 404, None), + ({"update_mask": ["field_1", "is_paused"]}, DAG1_ID, {"is_paused": True}, 400, None), + ({}, DAG1_ID, {"is_paused": True}, 200, True), + ({}, DAG1_ID, {"is_paused": False}, 200, False), + ({"update_mask": ["is_paused"]}, DAG1_ID, {"is_paused": True}, 200, True), + ({"update_mask": ["is_paused"]}, DAG1_ID, {"is_paused": False}, 200, False), + ], +) +def test_patch_dag(test_client, query_params, dag_id, body, expected_status_code, expected_is_paused): + response = test_client.patch(f"/public/dags/{dag_id}", json=body, params=query_params) + + assert response.status_code == expected_status_code + if expected_status_code == 200: + body = response.json() + assert body["is_paused"] == expected_is_paused From 693180304f8114169f89152eb6f284d6c33bf70e Mon Sep 17 00:00:00 2001 From: Danny Liu Date: Thu, 26 Sep 2024 04:28:08 -0700 Subject: [PATCH 035/802] fix: ensure DAG trigger form submits with updated parameters upon keyboard submit (#42487) --- airflow/www/templates/airflow/trigger.html | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/airflow/www/templates/airflow/trigger.html b/airflow/www/templates/airflow/trigger.html index 7cdcd337beddf..71d09e79076f1 100644 --- a/airflow/www/templates/airflow/trigger.html +++ b/airflow/www/templates/airflow/trigger.html @@ -163,7 +163,7 @@

{{ dag.description[0:150] + '…' if dag.description and dag.description|length > 150 else dag.description|default('', true) }}

{{ dag_docs(doc_md, False) }} -
+ From 9d21a7c9bb8ed4242a649c3cc2c80b433380e504 Mon Sep 17 00:00:00 2001 From: Wei Lee Date: Thu, 26 Sep 2024 21:39:00 +0900 Subject: [PATCH 036/802] fix(providers/common/sql): add dummy connection setter for backward compatibility (#42490) the introduction of connection property breaks apache-airflow-providers-mysql<5.7.1, apache-airflow-providers-elasticsearch<5.5.1 and apache-airflow-providers-postgres<5.13.0 --- airflow/providers/common/sql/hooks/sql.py | 11 +++++++++++ airflow/providers/common/sql/hooks/sql.pyi | 2 ++ tests/providers/mysql/hooks/test_mysql.py | 15 +++++++++++++++ 3 files changed, 28 insertions(+) diff --git a/airflow/providers/common/sql/hooks/sql.py b/airflow/providers/common/sql/hooks/sql.py index f2d11b21d7642..c0a3ed5f9a51c 100644 --- a/airflow/providers/common/sql/hooks/sql.py +++ b/airflow/providers/common/sql/hooks/sql.py @@ -210,6 +210,17 @@ def connection(self) -> Connection: self._connection = self.get_connection(self.get_conn_id()) return self._connection + @connection.setter + def connection(self, value: Any) -> None: + # This setter is for backward compatibility and should not be used. + # Since the introduction of connection property, the providers listed below + # breaks due to assigning value to self.connection + # + # apache-airflow-providers-mysql<5.7.1 + # apache-airflow-providers-elasticsearch<5.5.1 + # apache-airflow-providers-postgres<5.13.0 + pass + @property def connection_extra(self) -> dict: return self.connection.extra_dejson diff --git a/airflow/providers/common/sql/hooks/sql.pyi b/airflow/providers/common/sql/hooks/sql.pyi index 21081f06d36cf..e54b033991412 100644 --- a/airflow/providers/common/sql/hooks/sql.pyi +++ b/airflow/providers/common/sql/hooks/sql.pyi @@ -70,6 +70,8 @@ class DbApiHook(BaseHook): def placeholder(self): ... @property def connection(self) -> Connection: ... + @connection.setter + def connection(self, value: Any) -> None: ... @property def connection_extra(self) -> dict: ... @cached_property diff --git a/tests/providers/mysql/hooks/test_mysql.py b/tests/providers/mysql/hooks/test_mysql.py index cb6005ca8cf0c..48fc62fe2c220 100644 --- a/tests/providers/mysql/hooks/test_mysql.py +++ b/tests/providers/mysql/hooks/test_mysql.py @@ -67,6 +67,21 @@ def test_get_conn(self, mock_connect): assert kwargs["host"] == "host" assert kwargs["db"] == "schema" + @mock.patch("MySQLdb.connect") + def test_dummy_connection_setter(self, mock_connect): + self.db_hook.get_conn() + + self.db_hook.connection = "Won't affect anything" + assert self.db_hook.connection != "Won't affect anything" + + assert mock_connect.call_count == 1 + args, kwargs = mock_connect.call_args + assert args == () + assert kwargs["user"] == "login" + assert kwargs["passwd"] == "password" + assert kwargs["host"] == "host" + assert kwargs["db"] == "schema" + @mock.patch("MySQLdb.connect") def test_get_uri(self, mock_connect): self.connection.extra = json.dumps({"charset": "utf-8"}) From a3abbbd8bf046418e500505044a9df3030dbcc82 Mon Sep 17 00:00:00 2001 From: Daniel Standish <15932138+dstandish@users.noreply.github.com> Date: Thu, 26 Sep 2024 06:39:46 -0700 Subject: [PATCH 037/802] Fix broken main re generated api typescript comment (#42500) --- airflow/www/static/js/types/api-generated.ts | 2 ++ 1 file changed, 2 insertions(+) diff --git a/airflow/www/static/js/types/api-generated.ts b/airflow/www/static/js/types/api-generated.ts index 09def0ac66b6a..60fd384df00a7 100644 --- a/airflow/www/static/js/types/api-generated.ts +++ b/airflow/www/static/js/types/api-generated.ts @@ -4563,6 +4563,8 @@ export interface operations { * If set to true (default) the Any value will be returned as string, e.g. a Python representation * of a dict. If set to false it will return the raw data as dict, list, string or whatever was stored. * + * This parameter is not meaningful when using XCom pickling, then it is always returned as string. + * * *New in version 2.10.0* */ stringify?: boolean; From 0df4c7e3b6c05551c7bb7f53e6d6c33a11c8e53e Mon Sep 17 00:00:00 2001 From: ellisms <114107920+ellisms@users.noreply.github.com> Date: Thu, 26 Sep 2024 10:37:09 -0400 Subject: [PATCH 038/802] `S3DeleteObjects` Operator: Handle dates passed as strings (#42464) --- airflow/providers/amazon/aws/operators/s3.py | 14 ++++++- .../providers/amazon/aws/operators/test_s3.py | 42 +++++++++++++++++++ 2 files changed, 54 insertions(+), 2 deletions(-) diff --git a/airflow/providers/amazon/aws/operators/s3.py b/airflow/providers/amazon/aws/operators/s3.py index 669a6ad25aff3..998c7a81065dc 100644 --- a/airflow/providers/amazon/aws/operators/s3.py +++ b/airflow/providers/amazon/aws/operators/s3.py @@ -24,6 +24,9 @@ from tempfile import NamedTemporaryFile from typing import TYPE_CHECKING, Sequence +import pytz +from dateutil import parser + from airflow.exceptions import AirflowException from airflow.models import BaseOperator from airflow.providers.amazon.aws.hooks.s3 import S3Hook @@ -498,8 +501,8 @@ def __init__( bucket: str, keys: str | list | None = None, prefix: str | None = None, - from_datetime: datetime | None = None, - to_datetime: datetime | None = None, + from_datetime: datetime | str | None = None, + to_datetime: datetime | str | None = None, aws_conn_id: str | None = "aws_default", verify: str | bool | None = None, **kwargs, @@ -530,6 +533,13 @@ def execute(self, context: Context): if isinstance(self.keys, (list, str)) and not self.keys: return + # handle case where dates are strings, specifically when sent as template fields and macros. + if isinstance(self.to_datetime, str): + self.to_datetime = parser.parse(self.to_datetime).replace(tzinfo=pytz.UTC) + + if isinstance(self.from_datetime, str): + self.from_datetime = parser.parse(self.from_datetime).replace(tzinfo=pytz.UTC) + s3_hook = S3Hook(aws_conn_id=self.aws_conn_id, verify=self.verify) keys = self.keys or s3_hook.list_keys( diff --git a/tests/providers/amazon/aws/operators/test_s3.py b/tests/providers/amazon/aws/operators/test_s3.py index 267b678c8dd89..937baefde59ea 100644 --- a/tests/providers/amazon/aws/operators/test_s3.py +++ b/tests/providers/amazon/aws/operators/test_s3.py @@ -21,6 +21,7 @@ import os import shutil import sys +from datetime import timedelta from io import BytesIO from tempfile import mkdtemp from unittest import mock @@ -29,7 +30,10 @@ import pytest from moto import mock_aws +from airflow import DAG from airflow.exceptions import AirflowException +from airflow.models.dagrun import DagRun +from airflow.models.taskinstance import TaskInstance from airflow.providers.amazon.aws.hooks.s3 import S3Hook from airflow.providers.amazon.aws.operators.s3 import ( S3CopyObjectOperator, @@ -52,6 +56,7 @@ ) from airflow.providers.openlineage.extractors import OperatorLineage from airflow.utils.timezone import datetime, utcnow +from airflow.utils.types import DagRunType from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields BUCKET_NAME = os.environ.get("BUCKET_NAME", "test-airflow-bucket") @@ -623,6 +628,43 @@ def test_s3_delete_multiple_objects(self): # There should be no object found in the bucket created earlier assert "Contents" not in conn.list_objects(Bucket=bucket, Prefix=key_pattern) + @pytest.mark.db_test + def test_dates_from_template(self, session): + """Specifically test for dates passed from templating that could be strings""" + bucket = "testbucket" + key_pattern = "path/data" + n_keys = 3 + keys = [key_pattern + str(i) for i in range(n_keys)] + + conn = boto3.client("s3") + conn.create_bucket(Bucket=bucket) + for k in keys: + conn.upload_fileobj(Bucket=bucket, Key=k, Fileobj=BytesIO(b"input")) + + execution_date = utcnow() + dag = DAG("test_dag", start_date=datetime(2020, 1, 1), schedule=timedelta(days=1)) + # use macros.ds_add since it returns a string, not a date + op = S3DeleteObjectsOperator( + task_id="XXXXXXXXXXXXXXXXXXXXXXX", + bucket=bucket, + from_datetime="{{ macros.ds_add(ds, -1) }}", + to_datetime="{{ macros.ds_add(ds, 1) }}", + dag=dag, + ) + + dag_run = DagRun( + dag_id=dag.dag_id, execution_date=execution_date, run_id="test", run_type=DagRunType.MANUAL + ) + ti = TaskInstance(task=op) + ti.dag_run = dag_run + session.add(ti) + session.commit() + context = ti.get_template_context(session) + + ti.render_templates(context) + op.execute(None) + assert "Contents" not in conn.list_objects(Bucket=bucket) + def test_s3_delete_from_to_datetime(self): bucket = "testbucket" key_pattern = "path/data" From 8088bb0f2f0d6fa8124cc306124df4413c159331 Mon Sep 17 00:00:00 2001 From: Niko Oliveira Date: Thu, 26 Sep 2024 07:37:30 -0700 Subject: [PATCH 039/802] Move ECS executor to stable (#42483) The AWS ECS executor has been released for just under one year. We have made changes as users have found issues but it has remained stable for quite a while now. This PR proposes to remove the experimental warning for the executor and move it to respecting semver (so no more breaking changes in minor releases). --- .../apache-airflow-providers-amazon/executors/ecs-executor.rst | 3 --- docs/apache-airflow-providers-amazon/executors/index.rst | 2 +- 2 files changed, 1 insertion(+), 4 deletions(-) diff --git a/docs/apache-airflow-providers-amazon/executors/ecs-executor.rst b/docs/apache-airflow-providers-amazon/executors/ecs-executor.rst index d4289e629a039..b00704dc864f2 100644 --- a/docs/apache-airflow-providers-amazon/executors/ecs-executor.rst +++ b/docs/apache-airflow-providers-amazon/executors/ecs-executor.rst @@ -16,9 +16,6 @@ under the License. -.. warning:: - The ECS Executor is alpha/experimental at the moment and may be subject to change without warning. - .. |executorName| replace:: ECS .. |dockerfileLink| replace:: `here `__ .. |configKwargs| replace:: SUBMIT_JOB_KWARGS diff --git a/docs/apache-airflow-providers-amazon/executors/index.rst b/docs/apache-airflow-providers-amazon/executors/index.rst index e100cd845d0e5..117cd1facccaa 100644 --- a/docs/apache-airflow-providers-amazon/executors/index.rst +++ b/docs/apache-airflow-providers-amazon/executors/index.rst @@ -24,5 +24,5 @@ Amazon Executors .. toctree:: :maxdepth: 1 - ECS Executor (experimental) + ECS Executor Batch Executor (experimental) From 4ab708b7b8259088e81ff5bd5b8ef063919a5b20 Mon Sep 17 00:00:00 2001 From: Ash Berlin-Taylor Date: Thu, 26 Sep 2024 15:48:23 +0100 Subject: [PATCH 040/802] Add in backport of astunparse where needed in static checks (#42503) --- hatch_build.py | 1 + scripts/ci/pre_commit/check_deferrable_default.py | 8 +++++++- 2 files changed, 8 insertions(+), 1 deletion(-) diff --git a/hatch_build.py b/hatch_build.py index efd3ccd560e56..e4cd1e6030cca 100644 --- a/hatch_build.py +++ b/hatch_build.py @@ -239,6 +239,7 @@ "blinker>=1.7.0", ], "devel-static-checks": [ + "astunparse>=1.6.3; python_version < '3.9'", "black>=23.12.0", "pre-commit>=3.5.0", "ruff==0.5.5", diff --git a/scripts/ci/pre_commit/check_deferrable_default.py b/scripts/ci/pre_commit/check_deferrable_default.py index bfde61f231643..a007739083ab4 100755 --- a/scripts/ci/pre_commit/check_deferrable_default.py +++ b/scripts/ci/pre_commit/check_deferrable_default.py @@ -25,6 +25,12 @@ import sys from typing import Iterator +if hasattr(ast, "unparse"): + # Py 3.9+ + unparse = ast.unparse +else: + from astunparse import unparse # type: ignore[no-redef] + import libcst as cst from libcst.codemod import CodemodContext from libcst.codemod.visitors import AddImportsVisitor @@ -78,7 +84,7 @@ def leave_Param(self, original_node: cst.Param, updated_node: cst.Param) -> cst. def _is_valid_deferrable_default(default: ast.AST) -> bool: """Check whether default is 'conf.getboolean("operators", "default_deferrable", fallback=False)'""" - return ast.unparse(default) == "conf.getboolean('operators', 'default_deferrable', fallback=False)" + return unparse(default) == "conf.getboolean('operators', 'default_deferrable', fallback=False)" def iter_check_deferrable_default_errors(module_filename: str) -> Iterator[str]: From 977b5e0e1ed17fb41f21b529a2e40a73c38476f4 Mon Sep 17 00:00:00 2001 From: Vincent <97131062+vincbeck@users.noreply.github.com> Date: Thu, 26 Sep 2024 08:40:23 -0700 Subject: [PATCH 041/802] Remove deprecated stuff from Amazon provider package (#42450) --- airflow/providers/amazon/CHANGELOG.rst | 92 +++++++ airflow/providers/amazon/aws/hooks/athena.py | 20 +- .../providers/amazon/aws/hooks/base_aws.py | 166 +----------- airflow/providers/amazon/aws/hooks/logs.py | 21 +- .../providers/amazon/aws/hooks/quicksight.py | 18 +- .../amazon/aws/hooks/redshift_cluster.py | 126 +-------- airflow/providers/amazon/aws/hooks/s3.py | 12 +- .../providers/amazon/aws/hooks/sagemaker.py | 49 +--- .../providers/amazon/aws/operators/appflow.py | 11 +- .../providers/amazon/aws/operators/batch.py | 30 +-- .../amazon/aws/operators/datasync.py | 9 +- airflow/providers/amazon/aws/operators/ecs.py | 26 +- airflow/providers/amazon/aws/operators/eks.py | 53 +--- airflow/providers/amazon/aws/operators/emr.py | 248 ++---------------- .../amazon/aws/operators/glue_databrew.py | 11 +- airflow/providers/amazon/aws/operators/rds.py | 20 +- .../amazon/aws/operators/sagemaker.py | 42 +-- .../amazon/aws/secrets/secrets_manager.py | 41 +-- airflow/providers/amazon/aws/sensors/batch.py | 9 +- airflow/providers/amazon/aws/sensors/dms.py | 9 +- airflow/providers/amazon/aws/sensors/emr.py | 7 - .../aws/sensors/glue_catalog_partition.py | 9 +- .../amazon/aws/sensors/glue_crawler.py | 9 +- .../amazon/aws/sensors/quicksight.py | 30 +-- .../amazon/aws/sensors/redshift_cluster.py | 9 +- airflow/providers/amazon/aws/sensors/s3.py | 9 +- .../providers/amazon/aws/sensors/sagemaker.py | 9 +- airflow/providers/amazon/aws/sensors/sqs.py | 9 +- .../amazon/aws/sensors/step_function.py | 9 +- .../providers/amazon/aws/transfers/base.py | 15 +- .../amazon/aws/transfers/gcs_to_s3.py | 38 +-- .../providers/amazon/aws/triggers/batch.py | 169 +----------- airflow/providers/amazon/aws/triggers/eks.py | 21 +- airflow/providers/amazon/aws/triggers/emr.py | 32 --- .../amazon/aws/triggers/glue_crawler.py | 11 - .../amazon/aws/triggers/glue_databrew.py | 21 -- airflow/providers/amazon/aws/triggers/rds.py | 79 ------ .../amazon/aws/triggers/redshift_cluster.py | 69 +---- .../amazon/aws/triggers/sagemaker.py | 95 +------ .../amazon/aws/utils/connection_wrapper.py | 168 +----------- airflow/providers/amazon/aws/utils/mixins.py | 21 -- .../connections/aws.rst | 16 -- .../amazon/aws/deferrable/__init__.py | 16 -- .../amazon/aws/deferrable/hooks/__init__.py | 16 -- .../aws/deferrable/hooks/test_base_aws.py | 101 ------- .../deferrable/hooks/test_redshift_cluster.py | 121 --------- .../amazon/aws/hooks/test_base_aws.py | 27 +- .../amazon/aws/hooks/test_quicksight.py | 12 +- .../amazon/aws/hooks/test_sagemaker.py | 63 +---- .../amazon/aws/operators/test_appflow.py | 9 - .../amazon/aws/operators/test_base_aws.py | 77 ------ .../amazon/aws/operators/test_batch.py | 81 +----- .../amazon/aws/operators/test_ecs.py | 140 +--------- .../amazon/aws/operators/test_eks.py | 73 ++---- .../aws/operators/test_emr_serverless.py | 137 +++------- .../aws/operators/test_glue_databrew.py | 20 -- .../aws/operators/test_redshift_data.py | 5 +- .../aws/secrets/test_secrets_manager.py | 114 -------- .../amazon/aws/sensors/test_base_aws.py | 78 ------ .../amazon/aws/sensors/test_quicksight.py | 15 +- .../amazon/aws/transfers/test_base.py | 13 - .../aws/transfers/test_dynamodb_to_s3.py | 37 --- .../amazon/aws/transfers/test_gcs_to_s3.py | 225 ++++++---------- .../providers/amazon/aws/triggers/test_emr.py | 83 ------ .../amazon/aws/triggers/test_glue_crawler.py | 12 +- .../amazon/aws/triggers/test_glue_databrew.py | 22 -- .../aws/triggers/test_redshift_cluster.py | 14 +- .../amazon/aws/triggers/test_serialization.py | 30 +-- .../aws/utils/test_connection_wrapper.py | 169 +----------- .../providers/amazon/aws/example_batch.py | 2 +- 70 files changed, 363 insertions(+), 3217 deletions(-) delete mode 100644 tests/providers/amazon/aws/deferrable/__init__.py delete mode 100644 tests/providers/amazon/aws/deferrable/hooks/__init__.py delete mode 100644 tests/providers/amazon/aws/deferrable/hooks/test_base_aws.py delete mode 100644 tests/providers/amazon/aws/deferrable/hooks/test_redshift_cluster.py diff --git a/airflow/providers/amazon/CHANGELOG.rst b/airflow/providers/amazon/CHANGELOG.rst index 7596ad3886c7a..f4837c54fc366 100644 --- a/airflow/providers/amazon/CHANGELOG.rst +++ b/airflow/providers/amazon/CHANGELOG.rst @@ -32,11 +32,103 @@ Main Breaking changes ~~~~~~~~~~~~~~~~ +.. warning:: + All deprecated classes, parameters and features have been removed from the Amazon provider package. + The following breaking changes were introduced: + + * Hooks + + * Removed ``sleep_time`` parameter from ``AthenaHook``. Use ``poll_query_status`` instead + * Removed ``BaseAsyncSessionFactory`` + * Removed ``AwsBaseAsyncHook`` + * Removed ``start_from_head`` parameter from ``AwsLogsHook.get_log_events`` method + * Removed ``sts_hook`` property from ``QuickSightHook`` + * Removed ``RedshiftAsyncHook`` + * Removed S3 connection type. Please use ``aws`` as ``conn_type`` instead, and specify ``bucket_name`` in ``service_config.s3`` within ``extras`` + * Removed ``wait_for_completion``, ``check_interval`` and ``verbose`` parameters from ``SageMakerHook.start_pipeline`` method + * Removed ``wait_for_completion``, ``check_interval`` and ``verbose`` parameters from ``SageMakerHook.stop_pipeline`` method + + * Operators + + * Removed ``source`` parameter from ``AppflowRunOperator`` + * Removed ``overrides`` parameter from ``BatchOperator``. Use ``container_overrides`` instead + * Removed ``status_retries`` parameter from ``BatchCreateComputeEnvironmentOperator`` + * Removed ``get_hook`` method from ``DataSyncOperator``. Use ``hook`` property instead + * Removed ``wait_for_completion``, ``waiter_delay`` and ``waiter_max_attempts`` parameters from ``EcsDeregisterTaskDefinitionOperator``. Please use ``waiter_max_attempts`` and ``waiter_delay`` instead + * Removed ``wait_for_completion``, ``waiter_delay`` and ``waiter_max_attempts`` parameters from ``EcsRegisterTaskDefinitionOperator``. Please use ``waiter_max_attempts`` and ``waiter_delay`` instead + * Removed ``eks_hook`` property from ``EksCreateClusterOperator``. Use ``hook`` property instead + * Removed ``pod_context``, ``pod_username`` and ``is_delete_operator_pod`` parameters from ``EksPodOperator`` + * Removed ``waiter_countdown`` and ``waiter_check_interval_seconds`` parameters from ``EmrStartNotebookExecutionOperator``. Please use ``waiter_max_attempts`` and ``waiter_delay`` instead + * Removed ``waiter_countdown`` and ``waiter_check_interval_seconds`` parameters from ``EmrStopNotebookExecutionOperator``. Please use ``waiter_max_attempts`` and ``waiter_delay`` instead + * Removed ``max_tries`` parameter from ``EmrContainerOperator``. Use ``max_polling_attempts`` instead + * Removed ``waiter_countdown`` and ``waiter_check_interval_seconds`` parameters from ``EmrCreateJobFlowOperator``. Please use ``waiter_max_attempts`` and ``waiter_delay`` instead + * Removed ``waiter_countdown`` and ``waiter_check_interval_seconds`` parameters from ``EmrServerlessCreateApplicationOperator``. Please use ``waiter_max_attempts`` and ``waiter_delay`` instead + * Removed ``waiter_countdown`` and ``waiter_check_interval_seconds`` parameters from ``EmrServerlessStartJobOperator``. Please use ``waiter_max_attempts`` and ``waiter_delay`` instead + * Removed ``waiter_countdown`` and ``waiter_check_interval_seconds`` parameters from ``EmrServerlessStopApplicationOperator``. Please use ``waiter_max_attempts`` and ``waiter_delay`` instead + * Removed ``waiter_countdown`` and ``waiter_check_interval_seconds`` parameters from ``EmrServerlessDeleteApplicationOperator``. Please use ``waiter_max_attempts`` and ``waiter_delay`` instead + * Removed ``delay`` parameter from ``GlueDataBrewStartJobOperator``. Use ``waiter_delay`` instead + * Removed ``hook_params`` parameter from ``RdsBaseOperator`` + * Removed ``increment`` as possible value from ``action_if_job_exists`` parameter from ``SageMakerProcessingOperator`` + * Removed ``increment`` as possible value from ``action_if_job_exists`` parameter from ``SageMakerTransformOperator`` + * Removed ``increment`` as possible value from ``action_if_job_exists`` parameter from ``SageMakerTrainingOperator`` + + * Secrets + + * Removed from ``full_url_mode`` and ``are_secret_values_urlencoded`` as possible key in ``kwargs`` from ``SecretsManagerBackend`` + + * Sensors + + * Removed ``get_hook`` method from ``BatchSensor``. Use ``hook`` property instead + * Removed ``get_hook`` method from ``DmsTaskBaseSensor``. Use ``hook`` property instead + * Removed ``get_hook`` method from ``EmrBaseSensor``. Use ``hook`` property instead + * Removed ``get_hook`` method from ``GlueCatalogPartitionSensor``. Use ``hook`` property instead + * Removed ``get_hook`` method from ``GlueCrawlerSensor``. Use ``hook`` property instead + * Removed ``quicksight_hook`` property from ``QuickSightSensor``. Use ``QuickSightSensor.hook`` instead + * Removed ``sts_hook`` property from ``QuickSightSensor`` + * Removed ``get_hook`` method from ``RedshiftClusterSensor``. Use ``hook`` property instead + * Removed ``get_hook`` method from ``S3KeySensor``. Use ``hook`` property instead + * Removed ``get_hook`` method from ``SageMakerBaseSensor``. Use ``hook`` property instead + * Removed ``get_hook`` method from ``SqsSensor``. Use ``hook`` property instead + * Removed ``get_hook`` method from ``StepFunctionExecutionSensor``. Use ``hook`` property instead + + * Transfers + + * Removed ``aws_conn_id`` parameter from ``AwsToAwsBaseOperator``. Use ``source_aws_conn_id`` instead + * Removed ``bucket`` and ``delimiter`` parameters from ``GCSToS3Operator``. Use ``gcs_bucket`` instead of ``bucket`` + + * Triggers + + * Removed ``BatchOperatorTrigger``. Use ``BatchJobTrigger`` instead + * Removed ``BatchSensorTrigger``. Use ``BatchJobTrigger`` instead + * Removed ``region`` parameter from ``EksCreateFargateProfileTrigger``. Use ``region_name`` instead + * Removed ``region`` parameter from ``EksDeleteFargateProfileTrigger``. Use ``region_name`` instead + * Removed ``poll_interval`` and ``max_attempts`` parameters from ``EmrCreateJobFlowTrigger``. Use ``waiter_delay`` and ``waiter_max_attempts`` instead + * Removed ``poll_interval`` and ``max_attempts`` parameters from ``EmrTerminateJobFlowTrigger``. Use ``waiter_delay`` and ``waiter_max_attempts`` instead + * Removed ``poll_interval`` parameter from ``EmrContainerTrigger``. Use ``waiter_delay`` instead + * Removed ``poll_interval`` parameter from ``GlueCrawlerCompleteTrigger``. Use ``waiter_delay`` instead + * Removed ``delay`` and ``max_attempts`` parameters from ``GlueDataBrewJobCompleteTrigger``. Use ``waiter_delay`` and ``waiter_max_attempts`` instead + * Removed ``RdsDbInstanceTrigger``. Use the other RDS triggers such as ``RdsDbDeletedTrigger``, ``RdsDbStoppedTrigger`` or ``RdsDbAvailableTrigger`` + * Removed ``poll_interval`` and ``max_attempts`` parameters from ``RedshiftCreateClusterTrigger``. Use ``waiter_delay`` and ``waiter_max_attempts`` instead + * Removed ``poll_interval`` and ``max_attempts`` parameters from ``RedshiftPauseClusterTrigger``. Use ``waiter_delay`` and ``waiter_max_attempts`` instead + * Removed ``poll_interval`` and ``max_attempts`` parameters from ``RedshiftCreateClusterSnapshotTrigger``. Use ``waiter_delay`` and ``waiter_max_attempts`` instead + * Removed ``poll_interval`` and ``max_attempts`` parameters from ``RedshiftResumeClusterTrigger``. Use ``waiter_delay`` and ``waiter_max_attempts`` instead + * Removed ``poll_interval`` and ``max_attempts`` parameters from ``RedshiftDeleteClusterTrigger``. Use ``waiter_delay`` and ``waiter_max_attempts`` instead + * Removed ``SageMakerTrainingPrintLogTrigger``. Use ``SageMakerTrigger`` instead + + * Utils + + * Removed ``test_endpoint_url`` as possible key in ``extra_config`` from ``AwsConnectionWrapper``. Please set ``endpoint_url`` in ``service_config.sts`` within ``extras`` + * Removed ``s3`` as possible value in ``conn_type`` from ``AwsConnectionWrapper``. Please update your connection to have ``conn_type='aws'`` + * Removed ``session_kwargs`` as key in connection extra config. Please specify arguments passed to boto3 session directly + * Removed ``host`` from AWS connection, please set it in ``extra['endpoint_url']`` instead + * Removed ``region`` parameter from ``AwsHookParams``. Use ``region_name`` instead + .. warning:: In order to support session reuse in RedshiftData operators, the following breaking changes were introduced: The ``database`` argument is now optional and as a result was moved after the ``sql`` argument which is a positional one. Update your DAGs accordingly if they rely on argument order. Applies to: + * ``RedshiftDataHook``'s ``execute_query`` method * ``RedshiftDataOperator`` diff --git a/airflow/providers/amazon/aws/hooks/athena.py b/airflow/providers/amazon/aws/hooks/athena.py index 79360c48b927e..4969f339dba51 100644 --- a/airflow/providers/amazon/aws/hooks/athena.py +++ b/airflow/providers/amazon/aws/hooks/athena.py @@ -25,10 +25,9 @@ from __future__ import annotations -import warnings from typing import TYPE_CHECKING, Any, Collection -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseHook from airflow.providers.amazon.aws.utils.waiter_with_logging import wait @@ -56,7 +55,6 @@ class AthenaHook(AwsBaseHook): Provide thick wrapper around :external+boto3:py:class:`boto3.client("athena") `. - :param sleep_time: obsolete, please use the parameter of `poll_query_status` method instead :param log_query: Whether to log athena query and other execution params when it's executed. Defaults to *True*. @@ -82,20 +80,8 @@ class AthenaHook(AwsBaseHook): "CANCELLED", ) - def __init__( - self, *args: Any, sleep_time: int | None = None, log_query: bool = True, **kwargs: Any - ) -> None: + def __init__(self, *args: Any, log_query: bool = True, **kwargs: Any) -> None: super().__init__(client_type="athena", *args, **kwargs) # type: ignore - if sleep_time is not None: - self.sleep_time = sleep_time - warnings.warn( - "The `sleep_time` parameter of the Athena hook is deprecated, " - "please pass this parameter to the poll_query_status method instead.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - else: - self.sleep_time = 30 # previous default value self.log_query = log_query self.__query_results: dict[str, Any] = {} @@ -291,7 +277,7 @@ def poll_query_status( try: wait( waiter=self.get_waiter("query_complete"), - waiter_delay=self.sleep_time if sleep_time is None else sleep_time, + waiter_delay=30 if sleep_time is None else sleep_time, waiter_max_attempts=max_polling_attempts or 120, args={"QueryExecutionId": query_execution_id}, failure_message=f"Error while waiting for query {query_execution_id} to complete", diff --git a/airflow/providers/amazon/aws/hooks/base_aws.py b/airflow/providers/amazon/aws/hooks/base_aws.py index 0d07bb16494f2..2c919a4ff5183 100644 --- a/airflow/providers/amazon/aws/hooks/base_aws.py +++ b/airflow/providers/amazon/aws/hooks/base_aws.py @@ -44,14 +44,12 @@ from botocore.config import Config from botocore.waiter import Waiter, WaiterModel from dateutil.tz import tzlocal -from deprecated import deprecated from slugify import slugify from airflow.configuration import conf from airflow.exceptions import ( AirflowException, AirflowNotFoundException, - AirflowProviderDeprecationWarning, ) from airflow.hooks.base import BaseHook from airflow.providers.amazon.aws.utils.connection_wrapper import AwsConnectionWrapper @@ -65,10 +63,11 @@ BaseAwsConnection = TypeVar("BaseAwsConnection", bound=Union[boto3.client, boto3.resource]) if TYPE_CHECKING: + from aiobotocore.session import AioSession from botocore.client import ClientMeta from botocore.credentials import ReadOnlyCredentials - from airflow.models.connection import Connection # Avoid circular imports. + from airflow.models.connection import Connection _loader = botocore.loaders.Loader() """ @@ -172,9 +171,7 @@ def get_async_session(self): session.register_component("data_loader", _loader) return session - def create_session( - self, deferrable: bool = False - ) -> boto3.session.Session | aiobotocore.session.AioSession: + def create_session(self, deferrable: bool = False) -> boto3.session.Session | AioSession: """Create boto3 or aiobotocore Session from connection config.""" if not self.conn: self.log.info( @@ -216,7 +213,7 @@ def _create_basic_session(self, session_kwargs: dict[str, Any]) -> boto3.session def _create_session_with_assume_role( self, session_kwargs: dict[str, Any], deferrable: bool = False - ) -> boto3.session.Session | aiobotocore.session.AioSession: + ) -> boto3.session.Session | AioSession: if self.conn.assume_role_method == "assume_role_with_web_identity": # Deferred credentials have no initial credentials credential_fetcher = self._get_web_identity_credential_fetcher() @@ -1029,158 +1026,3 @@ def resolve_session_factory() -> type[BaseSessionFactory]: SessionFactory = resolve_session_factory() - - -def _parse_s3_config(config_file_name: str, config_format: str | None = "boto", profile: str | None = None): - """For compatibility with airflow.contrib.hooks.aws_hook.""" - from airflow.providers.amazon.aws.utils.connection_wrapper import _parse_s3_config - - return _parse_s3_config( - config_file_name=config_file_name, - config_format=config_format, - profile=profile, - ) - - -try: - import aiobotocore.credentials - from aiobotocore.session import AioSession, get_session -except ImportError: - pass - - -@deprecated( - reason=( - "`airflow.providers.amazon.aws.hook.base_aws.BaseAsyncSessionFactory` " - "has been deprecated and will be removed in future" - ), - category=AirflowProviderDeprecationWarning, -) -class BaseAsyncSessionFactory(BaseSessionFactory): - """ - Base AWS Session Factory class to handle aiobotocore session creation. - - It currently, handles ENV, AWS secret key and STS client method ``assume_role`` - provided in Airflow connection - """ - - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - - async def get_role_credentials(self) -> dict: - """Get the role_arn, method credentials from connection and get the role credentials.""" - async with self._basic_session.create_client("sts", region_name=self.region_name) as client: - response = await client.assume_role( - RoleArn=self.role_arn, - RoleSessionName=self._strip_invalid_session_name_characters(f"Airflow_{self.conn.conn_id}"), - **self.conn.assume_role_kwargs, - ) - return response["Credentials"] - - async def _get_refresh_credentials(self) -> dict[str, Any]: - self.log.debug("Refreshing credentials") - assume_role_method = self.conn.assume_role_method - if assume_role_method != "assume_role": - raise NotImplementedError(f"assume_role_method={assume_role_method} not expected") - - credentials = await self.get_role_credentials() - - expiry_time = credentials["Expiration"].isoformat() - self.log.debug("New credentials expiry_time: %s", expiry_time) - credentials = { - "access_key": credentials.get("AccessKeyId"), - "secret_key": credentials.get("SecretAccessKey"), - "token": credentials.get("SessionToken"), - "expiry_time": expiry_time, - } - return credentials - - def _get_session_with_assume_role(self) -> AioSession: - assume_role_method = self.conn.assume_role_method - if assume_role_method != "assume_role": - raise NotImplementedError(f"assume_role_method={assume_role_method} not expected") - - credentials = aiobotocore.credentials.AioRefreshableCredentials.create_from_metadata( - metadata=self._get_refresh_credentials(), - refresh_using=self._get_refresh_credentials, - method="sts-assume-role", - ) - - session = aiobotocore.session.get_session() - session._credentials = credentials - return session - - @cached_property - def _basic_session(self) -> AioSession: - """Cached property with basic aiobotocore.session.AioSession.""" - session_kwargs = self.conn.session_kwargs - aws_access_key_id = session_kwargs.get("aws_access_key_id") - aws_secret_access_key = session_kwargs.get("aws_secret_access_key") - aws_session_token = session_kwargs.get("aws_session_token") - region_name = session_kwargs.get("region_name") - profile_name = session_kwargs.get("profile_name") - - aio_session = get_session() - if profile_name is not None: - aio_session.set_config_variable("profile", profile_name) - if aws_access_key_id or aws_secret_access_key or aws_session_token: - aio_session.set_credentials( - access_key=aws_access_key_id, - secret_key=aws_secret_access_key, - token=aws_session_token, - ) - if region_name is not None: - aio_session.set_config_variable("region", region_name) - return aio_session - - def create_session(self, deferrable: bool = False) -> AioSession: - """Create aiobotocore Session from connection and config.""" - if not self._conn: - self.log.info("No connection ID provided. Fallback on boto3 credential strategy") - return get_session() - elif not self.role_arn: - return self._basic_session - return self._get_session_with_assume_role() - - -@deprecated( - reason=( - "`airflow.providers.amazon.aws.hook.base_aws.AwsBaseAsyncHook` " - "has been deprecated and will be removed in future" - ), - category=AirflowProviderDeprecationWarning, -) -class AwsBaseAsyncHook(AwsBaseHook): - """ - Interacts with AWS using aiobotocore asynchronously. - - :param aws_conn_id: The Airflow connection used for AWS credentials. - If this is None or empty then the default botocore behaviour is used. If - running Airflow in a distributed manner and aws_conn_id is None or - empty, then default botocore configuration would be used (and must be - maintained on each worker node). - :param verify: Whether to verify SSL certificates. - :param region_name: AWS region_name. If not specified then the default boto3 behaviour is used. - :param client_type: boto3.client client_type. Eg 's3', 'emr' etc - :param resource_type: boto3.resource resource_type. Eg 'dynamodb' etc - :param config: Configuration for botocore client. - """ - - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - - def get_async_session(self) -> AioSession: - """Get the underlying aiobotocore.session.AioSession(...).""" - return BaseAsyncSessionFactory( - conn=self.conn_config, region_name=self.region_name, config=self.config - ).create_session() - - async def get_client_async(self): - """Get the underlying aiobotocore client using aiobotocore session.""" - return self.get_async_session().create_client( - self.client_type, - region_name=self.region_name, - verify=self.verify, - endpoint_url=self.conn_config.endpoint_url, - config=self.config, - ) diff --git a/airflow/providers/amazon/aws/hooks/logs.py b/airflow/providers/amazon/aws/hooks/logs.py index 82b0bdd59dbe1..e3f34e33bd919 100644 --- a/airflow/providers/amazon/aws/hooks/logs.py +++ b/airflow/providers/amazon/aws/hooks/logs.py @@ -18,12 +18,10 @@ from __future__ import annotations import asyncio -import warnings from typing import Any, AsyncGenerator, Generator from botocore.exceptions import ClientError -from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseHook from airflow.utils.helpers import prune_dict @@ -80,8 +78,6 @@ def get_log_events( :param start_time: The timestamp value in ms to start reading the logs from (default: 0). :param skip: The number of log entries to skip at the start (default: 0). This is for when there are multiple entries at the same timestamp. - :param start_from_head: Deprecated. Do not use with False, logs would be retrieved out of order. - If possible, retrieve logs in one query, or implement pagination yourself. :param continuation_token: a token indicating where to read logs from. Will be updated as this method reads new logs, to be reused in subsequent calls. :param end_time: The timestamp value in ms to stop reading the logs from (default: None). @@ -91,21 +87,6 @@ def get_log_events( | 'message' (str): The log event data. | 'ingestionTime' (int): The time in milliseconds the event was ingested. """ - if start_from_head is not None: - message = ( - "start_from_head is deprecated, please remove this parameter." - if start_from_head - else "Do not use this method with start_from_head=False, logs will be returned out of order. " - "If possible, retrieve logs in one query, or implement pagination yourself." - ) - warnings.warn( - message, - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - else: - start_from_head = True - if continuation_token is None: continuation_token = AwsLogsHook.ContinuationToken() @@ -123,7 +104,7 @@ def get_log_events( "logStreamName": log_stream_name, "startTime": start_time, "endTime": end_time, - "startFromHead": start_from_head, + "startFromHead": True, **token_arg, } ) diff --git a/airflow/providers/amazon/aws/hooks/quicksight.py b/airflow/providers/amazon/aws/hooks/quicksight.py index 3a3a683597abd..f5d98ec32a08d 100644 --- a/airflow/providers/amazon/aws/hooks/quicksight.py +++ b/airflow/providers/amazon/aws/hooks/quicksight.py @@ -18,12 +18,10 @@ from __future__ import annotations import time -from functools import cached_property from botocore.exceptions import ClientError -from deprecated import deprecated -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseHook @@ -170,17 +168,3 @@ def wait_for_state( self.log.info("QuickSight Ingestion completed") return status - - @cached_property - @deprecated( - reason=( - "`QuickSightHook.sts_hook` property is deprecated and will be removed in the future. " - "This property used for obtain AWS Account ID, " - "please consider to use `QuickSightHook.account_id` instead" - ), - category=AirflowProviderDeprecationWarning, - ) - def sts_hook(self): - from airflow.providers.amazon.aws.hooks.sts import StsHook - - return StsHook(aws_conn_id=self.aws_conn_id) diff --git a/airflow/providers/amazon/aws/hooks/redshift_cluster.py b/airflow/providers/amazon/aws/hooks/redshift_cluster.py index 7e6dd01cf3948..ee365e0ab4d13 100644 --- a/airflow/providers/amazon/aws/hooks/redshift_cluster.py +++ b/airflow/providers/amazon/aws/hooks/redshift_cluster.py @@ -16,14 +16,9 @@ # under the License. from __future__ import annotations -import asyncio from typing import Any, Sequence -import botocore.exceptions -from deprecated import deprecated - -from airflow.exceptions import AirflowProviderDeprecationWarning -from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseAsyncHook, AwsBaseHook +from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseHook class RedshiftHook(AwsBaseHook): @@ -89,8 +84,6 @@ def cluster_status(self, cluster_identifier: str) -> str: - :external+boto3:py:meth:`Redshift.Client.describe_clusters` :param cluster_identifier: unique identifier of a cluster - :param skip_final_cluster_snapshot: determines cluster snapshot creation - :param final_cluster_snapshot_identifier: Optional[str] """ try: response = self.get_conn().describe_clusters(ClusterIdentifier=cluster_identifier)["Clusters"] @@ -98,6 +91,11 @@ def cluster_status(self, cluster_identifier: str) -> str: except self.get_conn().exceptions.ClusterNotFoundFault: return "cluster_not_found" + async def cluster_status_async(self, cluster_identifier: str) -> str: + async with self.async_conn as client: + response = await client.describe_clusters(ClusterIdentifier=cluster_identifier)["Clusters"] + return response[0]["ClusterStatus"] if response else None + def delete_cluster( self, cluster_identifier: str, @@ -201,115 +199,3 @@ def get_cluster_snapshot_status(self, snapshot_identifier: str): return snapshot_status except self.get_conn().exceptions.ClusterSnapshotNotFoundFault: return None - - -@deprecated( - reason=( - "`airflow.providers.amazon.aws.hook.base_aws.RedshiftAsyncHook` " - "has been deprecated and will be removed in future" - ), - category=AirflowProviderDeprecationWarning, -) -class RedshiftAsyncHook(AwsBaseAsyncHook): - """Interact with AWS Redshift using aiobotocore library.""" - - def __init__(self, *args, **kwargs): - kwargs["client_type"] = "redshift" - super().__init__(*args, **kwargs) - - async def cluster_status(self, cluster_identifier: str, delete_operation: bool = False) -> dict[str, Any]: - """ - Get the cluster status. - - :param cluster_identifier: unique identifier of a cluster - :param delete_operation: whether the method has been called as part of delete cluster operation - """ - async with await self.get_client_async() as client: - try: - response = await client.describe_clusters(ClusterIdentifier=cluster_identifier) - cluster_state = ( - response["Clusters"][0]["ClusterStatus"] if response and response["Clusters"] else None - ) - return {"status": "success", "cluster_state": cluster_state} - except botocore.exceptions.ClientError as error: - if delete_operation and error.response.get("Error", {}).get("Code", "") == "ClusterNotFound": - return {"status": "success", "cluster_state": "cluster_not_found"} - return {"status": "error", "message": str(error)} - - async def pause_cluster(self, cluster_identifier: str, poll_interval: float = 5.0) -> dict[str, Any]: - """ - Pause the cluster. - - :param cluster_identifier: unique identifier of a cluster - :param poll_interval: polling period in seconds to check for the status - """ - try: - async with await self.get_client_async() as client: - response = await client.pause_cluster(ClusterIdentifier=cluster_identifier) - status = response["Cluster"]["ClusterStatus"] if response and response["Cluster"] else None - if status == "pausing": - flag = asyncio.Event() - while True: - expected_response = await asyncio.create_task( - self.get_cluster_status(cluster_identifier, "paused", flag) - ) - await asyncio.sleep(poll_interval) - if flag.is_set(): - return expected_response - return {"status": "error", "cluster_state": status} - except botocore.exceptions.ClientError as error: - return {"status": "error", "message": str(error)} - - async def resume_cluster( - self, - cluster_identifier: str, - polling_period_seconds: float = 5.0, - ) -> dict[str, Any]: - """ - Resume the cluster. - - :param cluster_identifier: unique identifier of a cluster - :param polling_period_seconds: polling period in seconds to check for the status - """ - async with await self.get_client_async() as client: - try: - response = await client.resume_cluster(ClusterIdentifier=cluster_identifier) - status = response["Cluster"]["ClusterStatus"] if response and response["Cluster"] else None - if status == "resuming": - flag = asyncio.Event() - while True: - expected_response = await asyncio.create_task( - self.get_cluster_status(cluster_identifier, "available", flag) - ) - await asyncio.sleep(polling_period_seconds) - if flag.is_set(): - return expected_response - return {"status": "error", "cluster_state": status} - except botocore.exceptions.ClientError as error: - return {"status": "error", "message": str(error)} - - async def get_cluster_status( - self, - cluster_identifier: str, - expected_state: str, - flag: asyncio.Event, - delete_operation: bool = False, - ) -> dict[str, Any]: - """ - Check for expected Redshift cluster state. - - :param cluster_identifier: unique identifier of a cluster - :param expected_state: expected_state example("available", "pausing", "paused"") - :param flag: asyncio even flag set true if success and if any error - :param delete_operation: whether the method has been called as part of delete cluster operation - """ - try: - response = await self.cluster_status(cluster_identifier, delete_operation=delete_operation) - if ("cluster_state" in response and response["cluster_state"] == expected_state) or response[ - "status" - ] == "error": - flag.set() - return response - except botocore.exceptions.ClientError as error: - flag.set() - return {"status": "error", "message": str(error)} diff --git a/airflow/providers/amazon/aws/hooks/s3.py b/airflow/providers/amazon/aws/hooks/s3.py index 76aed19782a8f..b609259f846ba 100644 --- a/airflow/providers/amazon/aws/hooks/s3.py +++ b/airflow/providers/amazon/aws/hooks/s3.py @@ -28,7 +28,6 @@ import re import shutil import time -import warnings from contextlib import suppress from copy import deepcopy from datetime import datetime @@ -55,7 +54,7 @@ from boto3.s3.transfer import S3Transfer, TransferConfig from botocore.exceptions import ClientError -from airflow.exceptions import AirflowException, AirflowNotFoundException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException, AirflowNotFoundException from airflow.providers.amazon.aws.exceptions import S3HookUriParseFailure from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseHook from airflow.providers.amazon.aws.utils.tags import format_tags @@ -119,15 +118,6 @@ def wrapper(*args, **kwargs) -> Callable: if "bucket_name" in self.service_config: bound_args.arguments["bucket_name"] = self.service_config["bucket_name"] - elif self.conn_config and self.conn_config.schema: - warnings.warn( - "s3 conn_type, and the associated schema field, is deprecated. " - "Please use aws conn_type instead, and specify `bucket_name` " - "in `service_config.s3` within `extras`.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - bound_args.arguments["bucket_name"] = self.conn_config.schema return func(*bound_args.args, **bound_args.kwargs) diff --git a/airflow/providers/amazon/aws/hooks/sagemaker.py b/airflow/providers/amazon/aws/hooks/sagemaker.py index 2c0f4fb25edc5..e16ab11b0c95c 100644 --- a/airflow/providers/amazon/aws/hooks/sagemaker.py +++ b/airflow/providers/amazon/aws/hooks/sagemaker.py @@ -22,7 +22,6 @@ import tarfile import tempfile import time -import warnings from collections import Counter, namedtuple from datetime import datetime from functools import partial @@ -31,7 +30,7 @@ from asgiref.sync import sync_to_async from botocore.exceptions import ClientError -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseHook from airflow.providers.amazon.aws.hooks.logs import AwsLogsHook from airflow.providers.amazon.aws.hooks.s3 import S3Hook @@ -1098,9 +1097,6 @@ def start_pipeline( pipeline_name: str, display_name: str = "airflow-triggered-execution", pipeline_params: dict | None = None, - wait_for_completion: bool = False, - check_interval: int | None = None, - verbose: bool = True, ) -> str: """ Start a new execution for a SageMaker pipeline. @@ -1115,16 +1111,6 @@ def start_pipeline( :return: the ARN of the pipeline execution launched. """ - if wait_for_completion or check_interval is not None: - warnings.warn( - "parameter `wait_for_completion` and `check_interval` are deprecated, " - "remove them and call check_status yourself if you want to wait for completion", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - if check_interval is None: - check_interval = 30 - formatted_params = format_tags(pipeline_params, key_label="Name") try: @@ -1137,23 +1123,11 @@ def start_pipeline( self.log.error("Failed to start pipeline execution, error: %s", ce) raise - arn = res["PipelineExecutionArn"] - if wait_for_completion: - self.check_status( - arn, - "PipelineExecutionStatus", - lambda p: self.describe_pipeline_exec(p, verbose), - check_interval, - non_terminal_states=self.pipeline_non_terminal_states, - ) - return arn + return res["PipelineExecutionArn"] def stop_pipeline( self, pipeline_exec_arn: str, - wait_for_completion: bool = False, - check_interval: int | None = None, - verbose: bool = True, fail_if_not_running: bool = False, ) -> str: """ @@ -1172,16 +1146,6 @@ def stop_pipeline( :return: Status of the pipeline execution after the operation. One of 'Executing'|'Stopping'|'Stopped'|'Failed'|'Succeeded'. """ - if wait_for_completion or check_interval is not None: - warnings.warn( - "parameter `wait_for_completion` and `check_interval` are deprecated, " - "remove them and call check_status yourself if you want to wait for completion", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - if check_interval is None: - check_interval = 10 - for retries in reversed(range(5)): try: self.conn.stop_pipeline_execution(PipelineExecutionArn=pipeline_exec_arn) @@ -1213,15 +1177,6 @@ def stop_pipeline( res = self.describe_pipeline_exec(pipeline_exec_arn) - if wait_for_completion and res["PipelineExecutionStatus"] in self.pipeline_non_terminal_states: - res = self.check_status( - pipeline_exec_arn, - "PipelineExecutionStatus", - lambda p: self.describe_pipeline_exec(p, verbose), - check_interval, - non_terminal_states=self.pipeline_non_terminal_states, - ) - return res["PipelineExecutionStatus"] def create_model_package_group(self, package_group_name: str, package_group_desc: str = "") -> bool: diff --git a/airflow/providers/amazon/aws/operators/appflow.py b/airflow/providers/amazon/aws/operators/appflow.py index e338aa8071360..4aced5bea0c19 100644 --- a/airflow/providers/amazon/aws/operators/appflow.py +++ b/airflow/providers/amazon/aws/operators/appflow.py @@ -17,11 +17,10 @@ from __future__ import annotations import time -import warnings from datetime import datetime, timedelta from typing import TYPE_CHECKING, cast -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.operators.python import ShortCircuitOperator from airflow.providers.amazon.aws.hooks.appflow import AppflowHook from airflow.providers.amazon.aws.operators.base_aws import AwsBaseOperator @@ -140,7 +139,6 @@ class AppflowRunOperator(AppflowBaseOperator): For more information on how to use this operator, take a look at the guide: :ref:`howto/operator:AppflowRunOperator` - :param source: Obsolete, unnecessary for this operator :param flow_name: The flow name :param poll_interval: how often in seconds to check the query status :param aws_conn_id: The Airflow connection used for AWS credentials. @@ -155,17 +153,10 @@ class AppflowRunOperator(AppflowBaseOperator): def __init__( self, flow_name: str, - source: str | None = None, poll_interval: int = 20, wait_for_completion: bool = True, **kwargs, ) -> None: - if source is not None: - warnings.warn( - "The `source` parameter is unused when simply running a flow, please remove it.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) super().__init__( flow_name=flow_name, flow_update=False, diff --git a/airflow/providers/amazon/aws/operators/batch.py b/airflow/providers/amazon/aws/operators/batch.py index ca4ba8bfad8c8..f61292154799f 100644 --- a/airflow/providers/amazon/aws/operators/batch.py +++ b/airflow/providers/amazon/aws/operators/batch.py @@ -26,13 +26,12 @@ from __future__ import annotations -import warnings from datetime import timedelta from functools import cached_property from typing import TYPE_CHECKING, Any, Sequence from airflow.configuration import conf -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.models import BaseOperator from airflow.models.mappedoperator import MappedOperator from airflow.providers.amazon.aws.hooks.batch_client import BatchClientHook @@ -48,7 +47,6 @@ ) from airflow.providers.amazon.aws.utils import trim_none_values, validate_execute_complete_event from airflow.providers.amazon.aws.utils.task_log_fetcher import AwsTaskLogFetcher -from airflow.utils.types import NOTSET if TYPE_CHECKING: from airflow.utils.context import Context @@ -65,7 +63,6 @@ class BatchOperator(BaseOperator): :param job_name: the name for the job that will run on AWS Batch (templated) :param job_definition: the job definition name on AWS Batch :param job_queue: the queue name on AWS Batch - :param overrides: DEPRECATED, use container_overrides instead with the same value. :param container_overrides: the `containerOverrides` parameter for boto3 (templated) :param ecs_properties_override: the `ecsPropertiesOverride` parameter for boto3 (templated) :param eks_properties_override: the `eksPropertiesOverride` parameter for boto3 (templated) @@ -165,7 +162,6 @@ def __init__( job_name: str, job_definition: str, job_queue: str, - overrides: dict | None = None, # deprecated container_overrides: dict | None = None, array_properties: dict | None = None, ecs_properties_override: dict | None = None, @@ -196,21 +192,6 @@ def __init__( self.job_queue = job_queue self.container_overrides = container_overrides - # handle `overrides` deprecation in favor of `container_overrides` - if overrides: - if container_overrides: - # disallow setting both old and new params - raise AirflowException( - "'container_overrides' replaces the 'overrides' parameter. " - "You cannot specify both. Please remove assignation to the deprecated 'overrides'." - ) - self.container_overrides = overrides - warnings.warn( - "Parameter `overrides` is deprecated, Please use `container_overrides` instead.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - self.ecs_properties_override = ecs_properties_override self.eks_properties_override = eks_properties_override self.node_overrides = node_overrides @@ -501,17 +482,8 @@ def __init__( aws_conn_id: str | None = None, region_name: str | None = None, deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), - status_retries=NOTSET, **kwargs, ): - if status_retries is not NOTSET: - warnings.warn( - "The `status_retries` parameter is unused and should be removed. " - "It'll be deleted in a future version.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - super().__init__(**kwargs) self.compute_environment_name = compute_environment_name diff --git a/airflow/providers/amazon/aws/operators/datasync.py b/airflow/providers/amazon/aws/operators/datasync.py index 36d69d5079f6c..0c199b59fb7a1 100644 --- a/airflow/providers/amazon/aws/operators/datasync.py +++ b/airflow/providers/amazon/aws/operators/datasync.py @@ -22,9 +22,7 @@ import random from typing import TYPE_CHECKING, Any, Sequence -from deprecated.classic import deprecated - -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning, AirflowTaskTimeout +from airflow.exceptions import AirflowException, AirflowTaskTimeout from airflow.providers.amazon.aws.hooks.datasync import DataSyncHook from airflow.providers.amazon.aws.operators.base_aws import AwsBaseOperator from airflow.providers.amazon.aws.utils.mixins import aws_template_fields @@ -199,11 +197,6 @@ def __init__( def _hook_parameters(self) -> dict[str, Any]: return {**super()._hook_parameters, "wait_interval_seconds": self.wait_interval_seconds} - @deprecated(reason="use `hook` property instead.", category=AirflowProviderDeprecationWarning) - def get_hook(self) -> DataSyncHook: - """Create and return DataSyncHook.""" - return self.hook - def execute(self, context: Context): # If task_arn was not specified then try to # find 0, 1 or many candidate DataSync Tasks to run diff --git a/airflow/providers/amazon/aws/operators/ecs.py b/airflow/providers/amazon/aws/operators/ecs.py index 433fd88cd636e..6f2906f5ad6e2 100644 --- a/airflow/providers/amazon/aws/operators/ecs.py +++ b/airflow/providers/amazon/aws/operators/ecs.py @@ -18,13 +18,12 @@ from __future__ import annotations import re -import warnings from datetime import timedelta from functools import cached_property from typing import TYPE_CHECKING, Any, Sequence from airflow.configuration import conf -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.exceptions import EcsOperatorError, EcsTaskFailToStart from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseHook from airflow.providers.amazon.aws.hooks.ecs import EcsClusterStates, EcsHook, should_retry_eni @@ -40,7 +39,6 @@ from airflow.providers.amazon.aws.utils.mixins import aws_template_fields from airflow.providers.amazon.aws.utils.task_log_fetcher import AwsTaskLogFetcher from airflow.utils.helpers import prune_dict -from airflow.utils.types import NOTSET if TYPE_CHECKING: import boto3 @@ -258,19 +256,8 @@ def __init__( self, *, task_definition: str, - wait_for_completion=NOTSET, - waiter_delay=NOTSET, - waiter_max_attempts=NOTSET, **kwargs, ): - if any(arg is not NOTSET for arg in [wait_for_completion, waiter_delay, waiter_max_attempts]): - warnings.warn( - "'wait_for_completion' and waiter related params have no effect and are deprecated, " - "please remove them.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - super().__init__(**kwargs) self.task_definition = task_definition @@ -311,19 +298,8 @@ def __init__( family: str, container_definitions: list[dict], register_task_kwargs: dict | None = None, - wait_for_completion=NOTSET, - waiter_delay=NOTSET, - waiter_max_attempts=NOTSET, **kwargs, ): - if any(arg is not NOTSET for arg in [wait_for_completion, waiter_delay, waiter_max_attempts]): - warnings.warn( - "'wait_for_completion' and waiter related params have no effect and are deprecated, " - "please remove them.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - super().__init__(**kwargs) self.family = family self.container_definitions = container_definitions diff --git a/airflow/providers/amazon/aws/operators/eks.py b/airflow/providers/amazon/aws/operators/eks.py index a62fe59d67345..fa82cdcc72dfd 100644 --- a/airflow/providers/amazon/aws/operators/eks.py +++ b/airflow/providers/amazon/aws/operators/eks.py @@ -19,17 +19,15 @@ from __future__ import annotations import logging -import warnings from ast import literal_eval from datetime import timedelta from functools import cached_property from typing import TYPE_CHECKING, Any, List, Sequence, cast from botocore.exceptions import ClientError, WaiterError -from deprecated import deprecated from airflow.configuration import conf -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.models import BaseOperator from airflow.providers.amazon.aws.hooks.eks import EksHook from airflow.providers.amazon.aws.triggers.eks import ( @@ -267,17 +265,6 @@ def __init__( def hook(self) -> EksHook: return EksHook(aws_conn_id=self.aws_conn_id, region_name=self.region) - @property - @deprecated( - reason=( - "`eks_hook` property is deprecated and will be removed in the future. " - "Please use `hook` property instead." - ), - category=AirflowProviderDeprecationWarning, - ) - def eks_hook(self): - return self.hook - def execute(self, context: Context): if self.compute: if self.compute not in SUPPORTED_COMPUTE_VALUES: @@ -397,7 +384,7 @@ def deferrable_create_cluster_next(self, context: Context, event: dict[str, Any] waiter_delay=self.waiter_delay, waiter_max_attempts=self.waiter_max_attempts, aws_conn_id=self.aws_conn_id, - region=self.region, + region_name=self.region, ), method_name="execute_complete", timeout=timedelta(seconds=self.waiter_max_attempts * self.waiter_delay), @@ -656,7 +643,7 @@ def execute(self, context: Context): aws_conn_id=self.aws_conn_id, waiter_delay=self.waiter_delay, waiter_max_attempts=self.waiter_max_attempts, - region=self.region, + region_name=self.region, ), method_name="execute_complete", # timeout is set to ensure that if a trigger dies, the timeout does not restart @@ -968,7 +955,7 @@ def execute(self, context: Context): aws_conn_id=self.aws_conn_id, waiter_delay=self.waiter_delay, waiter_max_attempts=self.waiter_max_attempts, - region=self.region, + region_name=self.region, ), method_name="execute_complete", # timeout is set to ensure that if a trigger dies, the timeout does not restart @@ -1016,11 +1003,6 @@ class EksPodOperator(KubernetesPodOperator): If "delete_pod", the pod will be deleted regardless its state; if "delete_succeeded_pod", only succeeded pod will be deleted. You can set to "keep_pod" to keep the pod. Current default is `keep_pod`, but this will be changed in the next major release of this provider. - :param is_delete_operator_pod: What to do when the pod reaches its final - state, or the execution is interrupted. If True, delete the - pod; if False, leave the pod. Current default is False, but this will be - changed in the next major release of this provider. - Deprecated - use `on_finish_action` instead. """ @@ -1043,37 +1025,16 @@ def __init__( # file is stored locally in the worker and not in the cluster. in_cluster: bool = False, namespace: str = DEFAULT_NAMESPACE_NAME, - pod_context: str | None = None, pod_name: str | None = None, - pod_username: str | None = None, aws_conn_id: str | None = DEFAULT_CONN_ID, region: str | None = None, on_finish_action: str | None = None, - is_delete_operator_pod: bool | None = None, **kwargs, ) -> None: - if is_delete_operator_pod is not None: - warnings.warn( - "`is_delete_operator_pod` parameter is deprecated, please use `on_finish_action`", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - kwargs["on_finish_action"] = ( - OnFinishAction.DELETE_POD if is_delete_operator_pod else OnFinishAction.KEEP_POD - ) + if on_finish_action is not None: + kwargs["on_finish_action"] = OnFinishAction(on_finish_action) else: - if on_finish_action is not None: - kwargs["on_finish_action"] = OnFinishAction(on_finish_action) - else: - warnings.warn( - f"You have not set parameter `on_finish_action` in class {self.__class__.__name__}. " - "Currently the default for this parameter is `keep_pod` but in a future release" - " the default will be changed to `delete_pod`. To ensure pods are not deleted in" - " the future you will need to set `on_finish_action=keep_pod` explicitly.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - kwargs["on_finish_action"] = OnFinishAction.KEEP_POD + kwargs["on_finish_action"] = OnFinishAction.DELETE_POD self.cluster_name = cluster_name self.in_cluster = in_cluster diff --git a/airflow/providers/amazon/aws/operators/emr.py b/airflow/providers/amazon/aws/operators/emr.py index fb2f5de47849f..6aba14b350910 100644 --- a/airflow/providers/amazon/aws/operators/emr.py +++ b/airflow/providers/amazon/aws/operators/emr.py @@ -18,14 +18,13 @@ from __future__ import annotations import ast -import warnings from datetime import timedelta from functools import cached_property from typing import TYPE_CHECKING, Any, Sequence from uuid import uuid4 from airflow.configuration import conf -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.models import BaseOperator from airflow.providers.amazon.aws.hooks.emr import EmrContainerHook, EmrHook, EmrServerlessHook from airflow.providers.amazon.aws.links.emr import ( @@ -227,11 +226,6 @@ class EmrStartNotebookExecutionOperator(BaseOperator): :param tags: Optional list of key value pair to associate with the notebook execution. :param waiter_max_attempts: Maximum number of tries before failing. :param waiter_delay: Number of seconds between polling the state of the notebook. - - :param waiter_countdown: Total amount of time the operator will wait for the notebook to stop. - Defaults to 25 * 60 seconds. (Deprecated. Please use waiter_max_attempts.) - :param waiter_check_interval_seconds: Number of seconds between polling the state of the notebook. - Defaults to 60 seconds. (Deprecated. Please use waiter_delay.) """ template_fields: Sequence[str] = ( @@ -261,35 +255,10 @@ def __init__( tags: list | None = None, wait_for_completion: bool = False, aws_conn_id: str | None = "aws_default", - # TODO: waiter_max_attempts and waiter_delay should default to None when the other two are deprecated. waiter_max_attempts: int | None = None, waiter_delay: int | None = None, - waiter_countdown: int | None = None, - waiter_check_interval_seconds: int | None = None, **kwargs: Any, ): - if waiter_check_interval_seconds: - warnings.warn( - "The parameter `waiter_check_interval_seconds` has been deprecated to " - "standardize naming conventions. Please `use waiter_delay instead`. In the " - "future this will default to None and defer to the waiter's default value.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - else: - waiter_check_interval_seconds = 60 - if waiter_countdown: - warnings.warn( - "The parameter waiter_countdown has been deprecated to standardize " - "naming conventions. Please use waiter_max_attempts instead. In the " - "future this will default to None and defer to the waiter's default value.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - # waiter_countdown defaults to never timing out, which is not supported - # by boto waiters, so we will set it here to "a very long time" for now. - waiter_max_attempts = (waiter_countdown or 999) // waiter_check_interval_seconds - super().__init__(**kwargs) self.editor_id = editor_id self.relative_path = relative_path @@ -302,7 +271,7 @@ def __init__( self.cluster_id = cluster_id self.aws_conn_id = aws_conn_id self.waiter_max_attempts = waiter_max_attempts or 25 - self.waiter_delay = waiter_delay or waiter_check_interval_seconds or 60 + self.waiter_delay = waiter_delay or 60 self.master_instance_security_group_id = master_instance_security_group_id def execute(self, context: Context): @@ -371,11 +340,6 @@ class EmrStopNotebookExecutionOperator(BaseOperator): maintained on each worker node). :param waiter_max_attempts: Maximum number of tries before failing. :param waiter_delay: Number of seconds between polling the state of the notebook. - - :param waiter_countdown: Total amount of time the operator will wait for the notebook to stop. - Defaults to 25 * 60 seconds. (Deprecated. Please use waiter_max_attempts.) - :param waiter_check_interval_seconds: Number of seconds between polling the state of the notebook. - Defaults to 60 seconds. (Deprecated. Please use waiter_delay.) """ template_fields: Sequence[str] = ( @@ -389,41 +353,16 @@ def __init__( notebook_execution_id: str, wait_for_completion: bool = False, aws_conn_id: str | None = "aws_default", - # TODO: waiter_max_attempts and waiter_delay should default to None when the other two are deprecated. waiter_max_attempts: int | None = None, waiter_delay: int | None = None, - waiter_countdown: int | None = None, - waiter_check_interval_seconds: int | None = None, **kwargs: Any, ): - if waiter_check_interval_seconds: - warnings.warn( - "The parameter `waiter_check_interval_seconds` has been deprecated to " - "standardize naming conventions. Please `use waiter_delay instead`. In the " - "future this will default to None and defer to the waiter's default value.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - else: - waiter_check_interval_seconds = 60 - if waiter_countdown: - warnings.warn( - "The parameter waiter_countdown has been deprecated to standardize " - "naming conventions. Please use waiter_max_attempts instead. In the " - "future this will default to None and defer to the waiter's default value.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - # waiter_countdown defaults to never timing out, which is not supported - # by boto waiters, so we will set it here to "a very long time" for now. - waiter_max_attempts = (waiter_countdown or 999) // waiter_check_interval_seconds - super().__init__(**kwargs) self.notebook_execution_id = notebook_execution_id self.wait_for_completion = wait_for_completion self.aws_conn_id = aws_conn_id self.waiter_max_attempts = waiter_max_attempts or 25 - self.waiter_delay = waiter_delay or waiter_check_interval_seconds or 60 + self.waiter_delay = waiter_delay or 60 def execute(self, context: Context) -> None: emr_hook = EmrHook(aws_conn_id=self.aws_conn_id) @@ -518,7 +457,6 @@ class EmrContainerOperator(BaseOperator): :param aws_conn_id: The Airflow connection used for AWS credentials. :param wait_for_completion: Whether or not to wait in the operator for the job to complete. :param poll_interval: Time (in seconds) to wait between two consecutive calls to check query status on EMR - :param max_tries: Deprecated - use max_polling_attempts instead. :param max_polling_attempts: Maximum number of times to wait for the job run to finish. Defaults to None, which will poll until the job is *not* in a pending, submitted, or running state. :param job_retry_max_attempts: Maximum number of times to retry when the EMR job fails. @@ -551,7 +489,6 @@ def __init__( aws_conn_id: str | None = "aws_default", wait_for_completion: bool = True, poll_interval: int = 30, - max_tries: int | None = None, tags: dict | None = None, max_polling_attempts: int | None = None, job_retry_max_attempts: int | None = None, @@ -575,18 +512,6 @@ def __init__( self.job_id: str | None = None self.deferrable = deferrable - if max_tries: - warnings.warn( - f"Parameter `{self.__class__.__name__}.max_tries` is deprecated and will be removed " - "in a future release. Please use method `max_polling_attempts` instead.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - if max_polling_attempts and max_polling_attempts != max_tries: - raise ValueError("max_polling_attempts must be the same value as max_tries") - else: - self.max_polling_attempts = max_tries - @cached_property def hook(self) -> EmrContainerHook: """Create and return an EmrContainerHook.""" @@ -715,11 +640,6 @@ class EmrCreateJobFlowOperator(BaseOperator): completion (True) :param waiter_max_attempts: Maximum number of tries before failing. :param waiter_delay: Number of seconds between polling the state of the notebook. - - :param waiter_countdown: Max. seconds to wait for jobflow completion (only in combination with - wait_for_completion=True, None = no limit) (Deprecated. Please use waiter_max_attempts.) - :param waiter_check_interval_seconds: Number of seconds between polling the jobflow state. Defaults to 60 - seconds. (Deprecated. Please use waiter_delay.) :param deferrable: If True, the operator will wait asynchronously for the crawl to complete. This implies waiting for completion. This mode requires aiobotocore module to be installed. (default: False) @@ -748,33 +668,9 @@ def __init__( wait_for_completion: bool = False, waiter_max_attempts: int | None = None, waiter_delay: int | None = None, - waiter_countdown: int | None = None, - waiter_check_interval_seconds: int | None = None, deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), **kwargs: Any, ): - if waiter_check_interval_seconds: - warnings.warn( - "The parameter `waiter_check_interval_seconds` has been deprecated to " - "standardize naming conventions. Please `use waiter_delay instead`. In the " - "future this will default to None and defer to the waiter's default value.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - else: - waiter_check_interval_seconds = 60 - if waiter_countdown: - warnings.warn( - "The parameter waiter_countdown has been deprecated to standardize " - "naming conventions. Please use waiter_max_attempts instead. In the " - "future this will default to None and defer to the waiter's default value.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - # waiter_countdown defaults to never timing out, which is not supported - # by boto waiters, so we will set it here to "a very long time" for now. - waiter_max_attempts = (waiter_countdown or 999) // waiter_check_interval_seconds - super().__init__(**kwargs) self.aws_conn_id = aws_conn_id self.emr_conn_id = emr_conn_id @@ -782,7 +678,7 @@ def __init__( self.region_name = region_name self.wait_for_completion = wait_for_completion self.waiter_max_attempts = waiter_max_attempts or 60 - self.waiter_delay = waiter_delay or waiter_check_interval_seconds or 60 + self.waiter_delay = waiter_delay or 60 self.deferrable = deferrable @cached_property @@ -1054,10 +950,6 @@ class EmrServerlessCreateApplicationOperator(BaseOperator): 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 waiter_countdown: (deprecated) Total amount of time, in seconds, the operator will wait for - the application to start. Defaults to 25 minutes. - :param waiter_check_interval_seconds: (deprecated) Number of seconds between polling the state - of the application. Defaults to 60 seconds. :waiter_max_attempts: Number of times the waiter should poll the application to check the state. If not set, the waiter will use its default value. :param waiter_delay: Number of seconds between polling the state of the application. @@ -1074,38 +966,14 @@ def __init__( config: dict | None = None, wait_for_completion: bool = True, aws_conn_id: str | None = "aws_default", - waiter_countdown: int | ArgNotSet = NOTSET, - waiter_check_interval_seconds: int | ArgNotSet = NOTSET, waiter_max_attempts: int | ArgNotSet = NOTSET, waiter_delay: int | ArgNotSet = NOTSET, deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), **kwargs, ): - if waiter_check_interval_seconds is NOTSET: - waiter_delay = 60 if waiter_delay is NOTSET else waiter_delay - else: - waiter_delay = waiter_check_interval_seconds if waiter_delay is NOTSET else waiter_delay - warnings.warn( - "The parameter waiter_check_interval_seconds has been deprecated to standardize " - "naming conventions. Please use waiter_delay instead. In the " - "future this will default to None and defer to the waiter's default value.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - if waiter_countdown is NOTSET: - waiter_max_attempts = 25 if waiter_max_attempts is NOTSET else waiter_max_attempts - else: - if waiter_max_attempts is NOTSET: - # ignoring mypy because it doesn't like ArgNotSet as an operand, but neither variables - # are of type ArgNotSet at this point. - waiter_max_attempts = waiter_countdown // waiter_delay # type: ignore[operator] - warnings.warn( - "The parameter waiter_countdown has been deprecated to standardize " - "naming conventions. Please use waiter_max_attempts instead. In the " - "future this will default to None and defer to the waiter's default value.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) + waiter_delay = 60 if waiter_delay is NOTSET else waiter_delay + waiter_max_attempts = 25 if waiter_max_attempts is NOTSET else waiter_max_attempts + self.aws_conn_id = aws_conn_id self.release_label = release_label self.job_type = job_type @@ -1228,10 +1096,6 @@ class EmrServerlessStartJobOperator(BaseOperator): empty, then default boto3 configuration would be used (and must be maintained on each worker node). :param name: Name for the EMR Serverless job. If not provided, a default name will be assigned. - :param waiter_countdown: (deprecated) Total amount of time, in seconds, the operator will wait for - the job finish. Defaults to 25 minutes. - :param waiter_check_interval_seconds: (deprecated) Number of seconds between polling the state of the job. - Defaults to 60 seconds. :waiter_max_attempts: Number of times the waiter should poll the application to check the state. If not set, the waiter will use its default value. :param waiter_delay: Number of seconds between polling the state of the job run. @@ -1276,39 +1140,15 @@ def __init__( wait_for_completion: bool = True, aws_conn_id: str | None = "aws_default", name: str | None = None, - waiter_countdown: int | ArgNotSet = NOTSET, - waiter_check_interval_seconds: int | ArgNotSet = NOTSET, waiter_max_attempts: int | ArgNotSet = NOTSET, waiter_delay: int | ArgNotSet = NOTSET, deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), enable_application_ui_links: bool = False, **kwargs, ): - if waiter_check_interval_seconds is NOTSET: - waiter_delay = 60 if waiter_delay is NOTSET else waiter_delay - else: - waiter_delay = waiter_check_interval_seconds if waiter_delay is NOTSET else waiter_delay - warnings.warn( - "The parameter waiter_check_interval_seconds has been deprecated to standardize " - "naming conventions. Please use waiter_delay instead. In the " - "future this will default to None and defer to the waiter's default value.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - if waiter_countdown is NOTSET: - waiter_max_attempts = 25 if waiter_max_attempts is NOTSET else waiter_max_attempts - else: - if waiter_max_attempts is NOTSET: - # ignoring mypy because it doesn't like ArgNotSet as an operand, but neither variables - # are of type ArgNotSet at this point. - waiter_max_attempts = waiter_countdown // waiter_delay # type: ignore[operator] - warnings.warn( - "The parameter waiter_countdown has been deprecated to standardize " - "naming conventions. Please use waiter_max_attempts instead. In the " - "future this will default to None and defer to the waiter's default value.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) + waiter_delay = 60 if waiter_delay is NOTSET else waiter_delay + waiter_max_attempts = 25 if waiter_max_attempts is NOTSET else waiter_max_attempts + self.aws_conn_id = aws_conn_id self.application_id = application_id self.execution_role_arn = execution_role_arn @@ -1566,10 +1406,6 @@ class EmrServerlessStopApplicationOperator(BaseOperator): 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 waiter_countdown: (deprecated) Total amount of time, in seconds, the operator will wait for - the application be stopped. Defaults to 5 minutes. - :param waiter_check_interval_seconds: (deprecated) Number of seconds between polling the state of the - application. Defaults to 60 seconds. :param force_stop: If set to True, any job for that app that is not in a terminal state will be cancelled. Otherwise, trying to stop an app with running jobs will return an error. If you want to wait for the jobs to finish gracefully, use @@ -1590,39 +1426,15 @@ def __init__( application_id: str, wait_for_completion: bool = True, aws_conn_id: str | None = "aws_default", - waiter_countdown: int | ArgNotSet = NOTSET, - waiter_check_interval_seconds: int | ArgNotSet = NOTSET, waiter_max_attempts: int | ArgNotSet = NOTSET, waiter_delay: int | ArgNotSet = NOTSET, force_stop: bool = False, deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), **kwargs, ): - if waiter_check_interval_seconds is NOTSET: - waiter_delay = 60 if waiter_delay is NOTSET else waiter_delay - else: - waiter_delay = waiter_check_interval_seconds if waiter_delay is NOTSET else waiter_delay - warnings.warn( - "The parameter waiter_check_interval_seconds has been deprecated to standardize " - "naming conventions. Please use waiter_delay instead. In the " - "future this will default to None and defer to the waiter's default value.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - if waiter_countdown is NOTSET: - waiter_max_attempts = 25 if waiter_max_attempts is NOTSET else waiter_max_attempts - else: - if waiter_max_attempts is NOTSET: - # ignoring mypy because it doesn't like ArgNotSet as an operand, but neither variables - # are of type ArgNotSet at this point. - waiter_max_attempts = waiter_countdown // waiter_delay # type: ignore[operator] - warnings.warn( - "The parameter waiter_countdown has been deprecated to standardize " - "naming conventions. Please use waiter_max_attempts instead. In the " - "future this will default to None and defer to the waiter's default value.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) + waiter_delay = 60 if waiter_delay is NOTSET else waiter_delay + waiter_max_attempts = 25 if waiter_max_attempts is NOTSET else waiter_max_attempts + self.aws_conn_id = aws_conn_id self.application_id = application_id self.wait_for_completion = False if deferrable else wait_for_completion @@ -1734,10 +1546,6 @@ class EmrServerlessDeleteApplicationOperator(EmrServerlessStopApplicationOperato 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 waiter_countdown: (deprecated) Total amount of time, in seconds, the operator will wait for each - step of first,the application to be stopped, and then deleted. Defaults to 25 minutes. - :param waiter_check_interval_seconds: (deprecated) Number of seconds between polling the state - of the application. Defaults to 60 seconds. :waiter_max_attempts: Number of times the waiter should poll the application to check the state. Defaults to 25. :param waiter_delay: Number of seconds between polling the state of the application. @@ -1758,39 +1566,15 @@ def __init__( application_id: str, wait_for_completion: bool = True, aws_conn_id: str | None = "aws_default", - waiter_countdown: int | ArgNotSet = NOTSET, - waiter_check_interval_seconds: int | ArgNotSet = NOTSET, waiter_max_attempts: int | ArgNotSet = NOTSET, waiter_delay: int | ArgNotSet = NOTSET, force_stop: bool = False, deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), **kwargs, ): - if waiter_check_interval_seconds is NOTSET: - waiter_delay = 60 if waiter_delay is NOTSET else waiter_delay - else: - waiter_delay = waiter_check_interval_seconds if waiter_delay is NOTSET else waiter_delay - warnings.warn( - "The parameter waiter_check_interval_seconds has been deprecated to standardize " - "naming conventions. Please use waiter_delay instead. In the " - "future this will default to None and defer to the waiter's default value.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - if waiter_countdown is NOTSET: - waiter_max_attempts = 25 if waiter_max_attempts is NOTSET else waiter_max_attempts - else: - if waiter_max_attempts is NOTSET: - # ignoring mypy because it doesn't like ArgNotSet as an operand, but neither variables - # are of type ArgNotSet at this point. - waiter_max_attempts = waiter_countdown // waiter_delay # type: ignore[operator] - warnings.warn( - "The parameter waiter_countdown has been deprecated to standardize " - "naming conventions. Please use waiter_max_attempts instead. In the " - "future this will default to None and defer to the waiter's default value.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) + waiter_delay = 60 if waiter_delay is NOTSET else waiter_delay + waiter_max_attempts = 25 if waiter_max_attempts is NOTSET else waiter_max_attempts + self.wait_for_delete_completion = wait_for_completion # super stops the app super().__init__( diff --git a/airflow/providers/amazon/aws/operators/glue_databrew.py b/airflow/providers/amazon/aws/operators/glue_databrew.py index 2ea3257131525..f0b361ec176a9 100644 --- a/airflow/providers/amazon/aws/operators/glue_databrew.py +++ b/airflow/providers/amazon/aws/operators/glue_databrew.py @@ -17,11 +17,10 @@ # under the License. from __future__ import annotations -import warnings from typing import TYPE_CHECKING, Any, Sequence from airflow.configuration import conf -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.glue_databrew import GlueDataBrewHook from airflow.providers.amazon.aws.operators.base_aws import AwsBaseOperator from airflow.providers.amazon.aws.triggers.glue_databrew import GlueDataBrewJobCompleteTrigger @@ -49,7 +48,6 @@ class GlueDataBrewStartJobOperator(AwsBaseOperator[GlueDataBrewHook]): :param deferrable: If True, the operator will wait asynchronously for the job to complete. This implies waiting for completion. This mode requires aiobotocore module to be installed. (default: False) - :param delay: Time in seconds to wait between status checks. (Deprecated). :param waiter_delay: Time in seconds to wait between status checks. Default is 30. :param waiter_max_attempts: Maximum number of attempts to check for job completion. (default: 60) :return: dictionary with key run_id and value of the resulting job's run_id. @@ -92,13 +90,6 @@ def __init__( self.waiter_delay = waiter_delay self.waiter_max_attempts = waiter_max_attempts self.deferrable = deferrable - if delay is not None: - warnings.warn( - "please use `waiter_delay` instead of delay.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - self.waiter_delay = delay def execute(self, context: Context): job = self.hook.conn.start_job_run(Name=self.job_name) diff --git a/airflow/providers/amazon/aws/operators/rds.py b/airflow/providers/amazon/aws/operators/rds.py index f37c698d8796a..a0ab45d946846 100644 --- a/airflow/providers/amazon/aws/operators/rds.py +++ b/airflow/providers/amazon/aws/operators/rds.py @@ -18,13 +18,12 @@ from __future__ import annotations import json -import warnings from datetime import timedelta from functools import cached_property from typing import TYPE_CHECKING, Any, Sequence from airflow.configuration import conf -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.models import BaseOperator from airflow.providers.amazon.aws.hooks.rds import RdsHook from airflow.providers.amazon.aws.triggers.rds import ( @@ -55,30 +54,17 @@ def __init__( *args, aws_conn_id: str | None = "aws_conn_id", region_name: str | None = None, - hook_params: dict | None = None, **kwargs, ): - if hook_params is not None: - warnings.warn( - "The parameter hook_params is deprecated and will be removed. " - "Note that it is also incompatible with deferrable mode. " - "You can use the region_name parameter to specify the region. " - "If you were using hook_params for other purposes, please get in touch either on " - "airflow slack, or by opening a github issue on the project. " - "You can mention https://github.com/apache/airflow/pull/32352", - AirflowProviderDeprecationWarning, - stacklevel=3, # 2 is in the operator's init, 3 is in the user code creating the operator - ) - self.hook_params = hook_params or {} self.aws_conn_id = aws_conn_id - self.region_name = region_name or self.hook_params.pop("region_name", None) + self.region_name = region_name super().__init__(*args, **kwargs) self._await_interval = 60 # seconds @cached_property def hook(self) -> RdsHook: - return RdsHook(aws_conn_id=self.aws_conn_id, region_name=self.region_name, **self.hook_params) + return RdsHook(aws_conn_id=self.aws_conn_id, region_name=self.region_name) def execute(self, context: Context) -> str: """Different implementations for snapshots, tasks and events.""" diff --git a/airflow/providers/amazon/aws/operators/sagemaker.py b/airflow/providers/amazon/aws/operators/sagemaker.py index 4da238ff2c0e1..57a9194526227 100644 --- a/airflow/providers/amazon/aws/operators/sagemaker.py +++ b/airflow/providers/amazon/aws/operators/sagemaker.py @@ -19,14 +19,13 @@ import datetime import json import time -import warnings from functools import cached_property from typing import TYPE_CHECKING, Any, Callable, Sequence from botocore.exceptions import ClientError from airflow.configuration import conf -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.models import BaseOperator from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseHook from airflow.providers.amazon.aws.hooks.sagemaker import ( @@ -239,7 +238,7 @@ class SageMakerProcessingOperator(SageMakerBaseOperator): doesn't finish within max_ingestion_time seconds. If you set this parameter to None, the operation does not timeout. :param action_if_job_exists: Behaviour if the job name already exists. Possible options are "timestamp" - (default), "increment" (deprecated) and "fail". + (default) and "fail". :param deferrable: Run operator in the deferrable mode. This is only effective if wait_for_completion is set to True. :return Dict: Returns The ARN of the processing job created in Amazon SageMaker. @@ -260,18 +259,11 @@ def __init__( **kwargs, ): super().__init__(config=config, aws_conn_id=aws_conn_id, **kwargs) - if action_if_job_exists not in ("increment", "fail", "timestamp"): + if action_if_job_exists not in ("fail", "timestamp"): raise AirflowException( - f"Argument action_if_job_exists accepts only 'timestamp', 'increment' and 'fail'. \ + f"Argument action_if_job_exists accepts only 'timestamp' and 'fail'. \ Provided value: '{action_if_job_exists}'." ) - if action_if_job_exists == "increment": - warnings.warn( - "Action 'increment' on job name conflict has been deprecated for performance reasons." - "The alternative to 'fail' is now 'timestamp'.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) self.action_if_job_exists = action_if_job_exists self.wait_for_completion = wait_for_completion self.print_log = print_log @@ -657,7 +649,7 @@ class SageMakerTransformOperator(SageMakerBaseOperator): :param check_if_job_exists: If set to true, then the operator will check whether a transform job already exists for the name in the config. :param action_if_job_exists: Behaviour if the job name already exists. Possible options are "timestamp" - (default), "increment" (deprecated) and "fail". + (default) and "fail". This is only relevant if check_if_job_exists is True. :return Dict: Returns The ARN of the model created in Amazon SageMaker. """ @@ -684,18 +676,11 @@ def __init__( self.max_attempts = max_attempts or 60 self.max_ingestion_time = max_ingestion_time self.check_if_job_exists = check_if_job_exists - if action_if_job_exists in ("increment", "fail", "timestamp"): - if action_if_job_exists == "increment": - warnings.warn( - "Action 'increment' on job name conflict has been deprecated for performance reasons." - "The alternative to 'fail' is now 'timestamp'.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) + if action_if_job_exists in ("fail", "timestamp"): self.action_if_job_exists = action_if_job_exists else: raise AirflowException( - f"Argument action_if_job_exists accepts only 'timestamp', 'increment' and 'fail'. \ + f"Argument action_if_job_exists accepts only 'timestamp' and 'fail'. \ Provided value: '{action_if_job_exists}'." ) self.check_if_model_exists = check_if_model_exists @@ -1064,7 +1049,7 @@ class SageMakerTrainingOperator(SageMakerBaseOperator): :param check_if_job_exists: If set to true, then the operator will check whether a training job already exists for the name in the config. :param action_if_job_exists: Behaviour if the job name already exists. Possible options are "timestamp" - (default), "increment" (deprecated) and "fail". + (default) and "fail". This is only relevant if check_if_job_exists is True. :param deferrable: Run operator in the deferrable mode. This is only effective if wait_for_completion is set to True. @@ -1093,18 +1078,11 @@ def __init__( self.max_attempts = max_attempts or 60 self.max_ingestion_time = max_ingestion_time self.check_if_job_exists = check_if_job_exists - if action_if_job_exists in {"timestamp", "increment", "fail"}: - if action_if_job_exists == "increment": - warnings.warn( - "Action 'increment' on job name conflict has been deprecated for performance reasons." - "The alternative to 'fail' is now 'timestamp'.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) + if action_if_job_exists in {"timestamp", "fail"}: self.action_if_job_exists = action_if_job_exists else: raise AirflowException( - f"Argument action_if_job_exists accepts only 'timestamp', 'increment' and 'fail'. \ + f"Argument action_if_job_exists accepts only 'timestamp' and 'fail'. \ Provided value: '{action_if_job_exists}'." ) self.deferrable = deferrable diff --git a/airflow/providers/amazon/aws/secrets/secrets_manager.py b/airflow/providers/amazon/aws/secrets/secrets_manager.py index e4916b8e1bcd1..4d675771b23df 100644 --- a/airflow/providers/amazon/aws/secrets/secrets_manager.py +++ b/airflow/providers/amazon/aws/secrets/secrets_manager.py @@ -21,12 +21,9 @@ import json import re -import warnings from functools import cached_property from typing import Any -from urllib.parse import unquote -from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.providers.amazon.aws.utils import trim_none_values from airflow.secrets import BaseSecretsBackend from airflow.utils.log.logging_mixin import LoggingMixin @@ -145,28 +142,7 @@ def __init__( self.variables_lookup_pattern = variables_lookup_pattern self.config_lookup_pattern = config_lookup_pattern self.sep = sep - - if kwargs.pop("full_url_mode", None) is not None: - warnings.warn( - "The `full_url_mode` kwarg is deprecated. Going forward, the `SecretsManagerBackend`" - " will support both URL-encoded and JSON-encoded secrets at the same time. The encoding" - " of the secret will be determined automatically.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - - if kwargs.get("are_secret_values_urlencoded") is not None: - warnings.warn( - "The `secret_values_are_urlencoded` is deprecated. This kwarg only exists to assist in" - " migrating away from URL-encoding secret values for JSON secrets." - " To remove this warning, make sure your JSON secrets are *NOT* URL-encoded, and then" - " remove this kwarg from backend_kwargs.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - self.are_secret_values_urlencoded = kwargs.pop("are_secret_values_urlencoded", None) - else: - self.are_secret_values_urlencoded = False + self.are_secret_values_urlencoded = False self.extra_conn_words = extra_conn_words or {} @@ -222,19 +198,6 @@ def _standardize_secret_keys(self, secret: dict[str, Any]) -> dict[str, Any]: return conn_d - def _remove_escaping_in_secret_dict(self, secret: dict[str, Any]) -> dict[str, Any]: - """Un-escape secret values that are URL-encoded.""" - for k, v in secret.copy().items(): - if k == "extra" and isinstance(v, dict): - # The old behavior was that extras were _not_ urlencoded inside the secret. - # So we should just allow the extra dict to remain as-is. - continue - - elif v is not None: - secret[k] = unquote(v) - - return secret - def get_conn_value(self, conn_id: str) -> str | None: """ Get serialized representation of Connection. @@ -259,8 +222,6 @@ def get_conn_value(self, conn_id: str) -> str | None: secret_dict = json.loads(secret) standardized_secret_dict = self._standardize_secret_keys(secret_dict) - if self.are_secret_values_urlencoded: - standardized_secret_dict = self._remove_escaping_in_secret_dict(standardized_secret_dict) standardized_secret = json.dumps(standardized_secret_dict) return standardized_secret else: diff --git a/airflow/providers/amazon/aws/sensors/batch.py b/airflow/providers/amazon/aws/sensors/batch.py index 6ba1da17eb35b..3d0c7799bf519 100644 --- a/airflow/providers/amazon/aws/sensors/batch.py +++ b/airflow/providers/amazon/aws/sensors/batch.py @@ -20,10 +20,8 @@ from functools import cached_property from typing import TYPE_CHECKING, Any, Sequence -from deprecated import deprecated - from airflow.configuration import conf -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.batch_client import BatchClientHook from airflow.providers.amazon.aws.triggers.batch import BatchJobTrigger from airflow.sensors.base import BaseSensorOperator @@ -120,11 +118,6 @@ def execute_complete(self, context: Context, event: dict[str, Any]) -> None: job_id = event["job_id"] self.log.info("Batch Job %s complete", job_id) - @deprecated(reason="use `hook` property instead.", category=AirflowProviderDeprecationWarning) - def get_hook(self) -> BatchClientHook: - """Create and return a BatchClientHook.""" - return self.hook - @cached_property def hook(self) -> BatchClientHook: return BatchClientHook( diff --git a/airflow/providers/amazon/aws/sensors/dms.py b/airflow/providers/amazon/aws/sensors/dms.py index 2ea52ea0b5c35..11867cb538d20 100644 --- a/airflow/providers/amazon/aws/sensors/dms.py +++ b/airflow/providers/amazon/aws/sensors/dms.py @@ -19,9 +19,7 @@ from typing import TYPE_CHECKING, Iterable, Sequence -from deprecated import deprecated - -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.dms import DmsHook from airflow.providers.amazon.aws.sensors.base_aws import AwsBaseSensor from airflow.providers.amazon.aws.utils.mixins import aws_template_fields @@ -68,11 +66,6 @@ def __init__( self.target_statuses: Iterable[str] = target_statuses or [] self.termination_statuses: Iterable[str] = termination_statuses or [] - @deprecated(reason="use `hook` property instead.", category=AirflowProviderDeprecationWarning) - def get_hook(self) -> DmsHook: - """Get DmsHook.""" - return self.hook - def poke(self, context: Context): if not (status := self.hook.get_task_status(self.replication_task_arn)): raise AirflowException( diff --git a/airflow/providers/amazon/aws/sensors/emr.py b/airflow/providers/amazon/aws/sensors/emr.py index e79642d35c693..50c9a836d96a1 100644 --- a/airflow/providers/amazon/aws/sensors/emr.py +++ b/airflow/providers/amazon/aws/sensors/emr.py @@ -21,12 +21,9 @@ from functools import cached_property from typing import TYPE_CHECKING, Any, Iterable, Sequence -from deprecated import deprecated - from airflow.configuration import conf from airflow.exceptions import ( AirflowException, - AirflowProviderDeprecationWarning, ) from airflow.providers.amazon.aws.hooks.emr import EmrContainerHook, EmrHook, EmrServerlessHook from airflow.providers.amazon.aws.links.emr import EmrClusterLink, EmrLogsLink, get_log_uri @@ -68,10 +65,6 @@ def __init__(self, *, aws_conn_id: str | None = "aws_default", **kwargs): self.target_states: Iterable[str] = [] # will be set in subclasses self.failed_states: Iterable[str] = [] # will be set in subclasses - @deprecated(reason="use `hook` property instead.", category=AirflowProviderDeprecationWarning) - def get_hook(self) -> EmrHook: - return self.hook - @cached_property def hook(self) -> EmrHook: return EmrHook(aws_conn_id=self.aws_conn_id) diff --git a/airflow/providers/amazon/aws/sensors/glue_catalog_partition.py b/airflow/providers/amazon/aws/sensors/glue_catalog_partition.py index af125e2dda6a8..0171644dfa31c 100644 --- a/airflow/providers/amazon/aws/sensors/glue_catalog_partition.py +++ b/airflow/providers/amazon/aws/sensors/glue_catalog_partition.py @@ -20,10 +20,8 @@ from datetime import timedelta from typing import TYPE_CHECKING, Any, Sequence -from deprecated import deprecated - from airflow.configuration import conf -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.glue_catalog import GlueCatalogHook from airflow.providers.amazon.aws.sensors.base_aws import AwsBaseSensor from airflow.providers.amazon.aws.triggers.glue import GlueCatalogPartitionTrigger @@ -129,8 +127,3 @@ def execute_complete(self, context: Context, event: dict | None = None) -> None: if event["status"] != "success": raise AirflowException(f"Trigger error: event is {event}") self.log.info("Partition exists in the Glue Catalog") - - @deprecated(reason="use `hook` property instead.", category=AirflowProviderDeprecationWarning) - def get_hook(self) -> GlueCatalogHook: - """Get the GlueCatalogHook.""" - return self.hook diff --git a/airflow/providers/amazon/aws/sensors/glue_crawler.py b/airflow/providers/amazon/aws/sensors/glue_crawler.py index 2d4396c010c39..c1aea4af58a54 100644 --- a/airflow/providers/amazon/aws/sensors/glue_crawler.py +++ b/airflow/providers/amazon/aws/sensors/glue_crawler.py @@ -19,9 +19,7 @@ from typing import TYPE_CHECKING, Sequence -from deprecated import deprecated - -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.glue_crawler import GlueCrawlerHook from airflow.providers.amazon.aws.sensors.base_aws import AwsBaseSensor from airflow.providers.amazon.aws.utils.mixins import aws_template_fields @@ -78,8 +76,3 @@ def poke(self, context: Context): raise AirflowException(f"Status: {crawler_status}") else: return False - - @deprecated(reason="use `hook` property instead.", category=AirflowProviderDeprecationWarning) - def get_hook(self) -> GlueCrawlerHook: - """Return a new or pre-existing GlueCrawlerHook.""" - return self.hook diff --git a/airflow/providers/amazon/aws/sensors/quicksight.py b/airflow/providers/amazon/aws/sensors/quicksight.py index 848c0dc7048fc..b0a9477706c57 100644 --- a/airflow/providers/amazon/aws/sensors/quicksight.py +++ b/airflow/providers/amazon/aws/sensors/quicksight.py @@ -17,12 +17,9 @@ # under the License. from __future__ import annotations -from functools import cached_property from typing import TYPE_CHECKING, Sequence -from deprecated import deprecated - -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.quicksight import QuickSightHook from airflow.providers.amazon.aws.sensors.base_aws import AwsBaseSensor @@ -76,28 +73,3 @@ def poke(self, context: Context) -> bool: error = self.hook.get_error_info(None, self.data_set_id, self.ingestion_id) raise AirflowException(f"The QuickSight Ingestion failed. Error info: {error}") return quicksight_ingestion_state == self.success_status - - @cached_property - @deprecated( - reason=( - "`QuickSightSensor.quicksight_hook` property is deprecated, " - "please use `QuickSightSensor.hook` property instead." - ), - category=AirflowProviderDeprecationWarning, - ) - def quicksight_hook(self): - return self.hook - - @cached_property - @deprecated( - reason=( - "`QuickSightSensor.sts_hook` property is deprecated and will be removed in the future. " - "This property used for obtain AWS Account ID, " - "please consider to use `QuickSightSensor.hook.account_id` instead" - ), - category=AirflowProviderDeprecationWarning, - ) - def sts_hook(self): - from airflow.providers.amazon.aws.hooks.sts import StsHook - - return StsHook(aws_conn_id=self.aws_conn_id) diff --git a/airflow/providers/amazon/aws/sensors/redshift_cluster.py b/airflow/providers/amazon/aws/sensors/redshift_cluster.py index 243c71e61fe78..11ff83123022a 100644 --- a/airflow/providers/amazon/aws/sensors/redshift_cluster.py +++ b/airflow/providers/amazon/aws/sensors/redshift_cluster.py @@ -20,10 +20,8 @@ from functools import cached_property from typing import TYPE_CHECKING, Any, Sequence -from deprecated import deprecated - from airflow.configuration import conf -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.redshift_cluster import RedshiftHook from airflow.providers.amazon.aws.triggers.redshift_cluster import RedshiftClusterTrigger from airflow.providers.amazon.aws.utils import validate_execute_complete_event @@ -98,11 +96,6 @@ def execute_complete(self, context: Context, event: dict[str, Any] | None = None self.log.info("%s completed successfully.", self.task_id) self.log.info("Cluster Identifier %s is in %s state", self.cluster_identifier, self.target_status) - @deprecated(reason="use `hook` property instead.", category=AirflowProviderDeprecationWarning) - def get_hook(self) -> RedshiftHook: - """Create and return a RedshiftHook.""" - return self.hook - @cached_property def hook(self) -> RedshiftHook: return RedshiftHook(aws_conn_id=self.aws_conn_id) diff --git a/airflow/providers/amazon/aws/sensors/s3.py b/airflow/providers/amazon/aws/sensors/s3.py index 2f32fff3d30ac..fd0d70a648466 100644 --- a/airflow/providers/amazon/aws/sensors/s3.py +++ b/airflow/providers/amazon/aws/sensors/s3.py @@ -25,15 +25,13 @@ from functools import cached_property from typing import TYPE_CHECKING, Any, Callable, Sequence, cast -from deprecated import deprecated - from airflow.configuration import conf from airflow.providers.amazon.aws.utils import validate_execute_complete_event if TYPE_CHECKING: from airflow.utils.context import Context -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.s3 import S3Hook from airflow.providers.amazon.aws.triggers.s3 import S3KeysUnchangedTrigger, S3KeyTrigger from airflow.sensors.base import BaseSensorOperator, poke_mode_only @@ -221,11 +219,6 @@ def execute_complete(self, context: Context, event: dict[str, Any]) -> None: elif event["status"] == "error": raise AirflowException(event["message"]) - @deprecated(reason="use `hook` property instead.", category=AirflowProviderDeprecationWarning) - def get_hook(self) -> S3Hook: - """Create and return an S3Hook.""" - return self.hook - @cached_property def hook(self) -> S3Hook: return S3Hook(aws_conn_id=self.aws_conn_id, verify=self.verify) diff --git a/airflow/providers/amazon/aws/sensors/sagemaker.py b/airflow/providers/amazon/aws/sensors/sagemaker.py index af07c504aa29d..e77628cf8d596 100644 --- a/airflow/providers/amazon/aws/sensors/sagemaker.py +++ b/airflow/providers/amazon/aws/sensors/sagemaker.py @@ -20,9 +20,7 @@ from functools import cached_property from typing import TYPE_CHECKING, Sequence -from deprecated import deprecated - -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.sagemaker import LogState, SageMakerHook from airflow.sensors.base import BaseSensorOperator @@ -45,11 +43,6 @@ def __init__(self, *, aws_conn_id: str | None = "aws_default", resource_type: st self.aws_conn_id = aws_conn_id self.resource_type = resource_type # only used for logs, to say what kind of resource we are sensing - @deprecated(reason="use `hook` property instead.", category=AirflowProviderDeprecationWarning) - def get_hook(self) -> SageMakerHook: - """Get SageMakerHook.""" - return self.hook - @cached_property def hook(self) -> SageMakerHook: return SageMakerHook(aws_conn_id=self.aws_conn_id) diff --git a/airflow/providers/amazon/aws/sensors/sqs.py b/airflow/providers/amazon/aws/sensors/sqs.py index d04a8cf820b01..99991c83ceb58 100644 --- a/airflow/providers/amazon/aws/sensors/sqs.py +++ b/airflow/providers/amazon/aws/sensors/sqs.py @@ -22,10 +22,8 @@ from datetime import timedelta from typing import TYPE_CHECKING, Any, Collection, Sequence -from deprecated import deprecated - from airflow.configuration import conf -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.sqs import SqsHook from airflow.providers.amazon.aws.sensors.base_aws import AwsBaseSensor from airflow.providers.amazon.aws.triggers.sqs import SqsSensorTrigger @@ -223,8 +221,3 @@ def poke(self, context: Context): return True else: return False - - @deprecated(reason="use `hook` property instead.", category=AirflowProviderDeprecationWarning) - def get_hook(self) -> SqsHook: - """Create and return an SqsHook.""" - return self.hook diff --git a/airflow/providers/amazon/aws/sensors/step_function.py b/airflow/providers/amazon/aws/sensors/step_function.py index 8af3bb6fe9c67..d4e58e2177c18 100644 --- a/airflow/providers/amazon/aws/sensors/step_function.py +++ b/airflow/providers/amazon/aws/sensors/step_function.py @@ -19,9 +19,7 @@ import json from typing import TYPE_CHECKING, Sequence -from deprecated import deprecated - -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.step_function import StepFunctionHook from airflow.providers.amazon.aws.sensors.base_aws import AwsBaseSensor from airflow.providers.amazon.aws.utils.mixins import aws_template_fields @@ -84,8 +82,3 @@ def poke(self, context: Context): self.log.info("Doing xcom_push of output") self.xcom_push(context, "output", output) return True - - @deprecated(reason="use `hook` property instead.", category=AirflowProviderDeprecationWarning) - def get_hook(self) -> StepFunctionHook: - """Create and return a StepFunctionHook.""" - return self.hook diff --git a/airflow/providers/amazon/aws/transfers/base.py b/airflow/providers/amazon/aws/transfers/base.py index 09a72bcf3875f..50458bb1631dd 100644 --- a/airflow/providers/amazon/aws/transfers/base.py +++ b/airflow/providers/amazon/aws/transfers/base.py @@ -19,18 +19,12 @@ from __future__ import annotations -import warnings from typing import Sequence -from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.models import BaseOperator from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseHook from airflow.utils.types import NOTSET, ArgNotSet -_DEPRECATION_MSG = ( - "The aws_conn_id parameter has been deprecated. Use the source_aws_conn_id parameter instead." -) - class AwsToAwsBaseOperator(BaseOperator): """ @@ -43,8 +37,6 @@ class AwsToAwsBaseOperator(BaseOperator): would be used (and must be maintained on each worker node). :param dest_aws_conn_id: The Airflow connection used for AWS credentials to access S3. If this is not set then the source_aws_conn_id connection is used. - :param aws_conn_id: The Airflow connection used for AWS credentials (deprecated; use source_aws_conn_id). - """ template_fields: Sequence[str] = ( @@ -57,17 +49,12 @@ def __init__( *, source_aws_conn_id: str | None = AwsBaseHook.default_conn_name, dest_aws_conn_id: str | None | ArgNotSet = NOTSET, - aws_conn_id: str | None | ArgNotSet = NOTSET, **kwargs, ) -> None: super().__init__(**kwargs) self.source_aws_conn_id = source_aws_conn_id self.dest_aws_conn_id = dest_aws_conn_id - if not isinstance(aws_conn_id, ArgNotSet): - warnings.warn(_DEPRECATION_MSG, AirflowProviderDeprecationWarning, stacklevel=3) - self.source_aws_conn_id = aws_conn_id - else: - self.source_aws_conn_id = source_aws_conn_id + self.source_aws_conn_id = source_aws_conn_id if isinstance(dest_aws_conn_id, ArgNotSet): self.dest_aws_conn_id = self.source_aws_conn_id else: diff --git a/airflow/providers/amazon/aws/transfers/gcs_to_s3.py b/airflow/providers/amazon/aws/transfers/gcs_to_s3.py index c0a3762b5a6c8..edfbeaa2a239d 100644 --- a/airflow/providers/amazon/aws/transfers/gcs_to_s3.py +++ b/airflow/providers/amazon/aws/transfers/gcs_to_s3.py @@ -20,12 +20,11 @@ from __future__ import annotations import os -import warnings from typing import TYPE_CHECKING, Sequence from packaging.version import Version -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.models import BaseOperator from airflow.providers.amazon.aws.hooks.s3 import S3Hook from airflow.providers.google.cloud.hooks.gcs import GCSHook @@ -43,12 +42,8 @@ class GCSToS3Operator(BaseOperator): :ref:`howto/operator:GCSToS3Operator` :param gcs_bucket: The Google Cloud Storage bucket to find the objects. (templated) - :param bucket: (Deprecated) Use ``gcs_bucket`` instead. :param prefix: Prefix string which filters objects whose name begin with this prefix. (templated) - :param delimiter: (Deprecated) The delimiter by which you want to filter the objects. (templated) - For e.g to lists the CSV files from in a directory in GCS you would use - delimiter='.csv'. :param gcp_conn_id: (Optional) The connection ID used to connect to Google Cloud. :param dest_aws_conn_id: The destination S3 connection :param dest_s3_key: The base S3 key to be used to store the files. (templated) @@ -91,7 +86,6 @@ class GCSToS3Operator(BaseOperator): template_fields: Sequence[str] = ( "gcs_bucket", "prefix", - "delimiter", "dest_s3_key", "google_impersonation_chain", "gcp_user_project", @@ -101,10 +95,8 @@ class GCSToS3Operator(BaseOperator): def __init__( self, *, - gcs_bucket: str | None = None, - bucket: str | None = None, + gcs_bucket: str, prefix: str | None = None, - delimiter: str | None = None, gcp_conn_id: str = "google_cloud_default", dest_aws_conn_id: str | None = "aws_default", dest_s3_key: str, @@ -119,17 +111,7 @@ def __init__( **kwargs, ) -> None: super().__init__(**kwargs) - if bucket: - warnings.warn( - "The ``bucket`` parameter is deprecated and will be removed in a future version. " - "Please use ``gcs_bucket`` instead.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - self.gcs_bucket = gcs_bucket or bucket - if not (bucket or gcs_bucket): - raise ValueError("You must pass either ``bucket`` or ``gcs_bucket``.") - + self.gcs_bucket = gcs_bucket self.prefix = prefix self.gcp_conn_id = gcp_conn_id self.dest_aws_conn_id = dest_aws_conn_id @@ -149,18 +131,10 @@ def __init__( self.__is_match_glob_supported = False except ImportError: # __version__ was added in 10.1.0, so this means it's < 10.3.0 self.__is_match_glob_supported = False - if self.__is_match_glob_supported: - if delimiter: - warnings.warn( - "Usage of 'delimiter' is deprecated, please use 'match_glob' instead", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - elif match_glob: + if not self.__is_match_glob_supported and match_glob: raise AirflowException( "The 'match_glob' parameter requires 'apache-airflow-providers-google>=10.3.0'." ) - self.delimiter = delimiter self.match_glob = match_glob self.gcp_user_project = gcp_user_project @@ -172,16 +146,14 @@ def execute(self, context: Context) -> list[str]: ) self.log.info( - "Getting list of the files. Bucket: %s; Delimiter: %s; Prefix: %s", + "Getting list of the files. Bucket: %s; Prefix: %s", self.gcs_bucket, - self.delimiter, self.prefix, ) list_kwargs = { "bucket_name": self.gcs_bucket, "prefix": self.prefix, - "delimiter": self.delimiter, "user_project": self.gcp_user_project, } if self.__is_match_glob_supported: diff --git a/airflow/providers/amazon/aws/triggers/batch.py b/airflow/providers/amazon/aws/triggers/batch.py index e5e14705eebf0..9d97ead007b83 100644 --- a/airflow/providers/amazon/aws/triggers/batch.py +++ b/airflow/providers/amazon/aws/triggers/batch.py @@ -16,182 +16,15 @@ # under the License. from __future__ import annotations -import asyncio -import itertools -from functools import cached_property -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING -from botocore.exceptions import WaiterError -from deprecated import deprecated - -from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.providers.amazon.aws.hooks.batch_client import BatchClientHook from airflow.providers.amazon.aws.triggers.base import AwsBaseWaiterTrigger -from airflow.triggers.base import BaseTrigger, TriggerEvent if TYPE_CHECKING: from airflow.providers.amazon.aws.hooks.base_aws import AwsGenericHook -@deprecated(reason="use BatchJobTrigger instead", category=AirflowProviderDeprecationWarning) -class BatchOperatorTrigger(BaseTrigger): - """ - Asynchronously poll the boto3 API and wait for the Batch job to be in the `SUCCEEDED` state. - - :param job_id: A unique identifier for the cluster. - :param max_retries: The maximum number of attempts to be made. - :param aws_conn_id: The Airflow connection used for AWS credentials. - :param region_name: region name to use in AWS Hook - :param poll_interval: The amount of time in seconds to wait between attempts. - """ - - def __init__( - self, - job_id: str | None = None, - max_retries: int = 10, - aws_conn_id: str | None = "aws_default", - region_name: str | None = None, - poll_interval: int = 30, - ): - super().__init__() - self.job_id = job_id - self.max_retries = max_retries - self.aws_conn_id = aws_conn_id - self.region_name = region_name - self.poll_interval = poll_interval - - def serialize(self) -> tuple[str, dict[str, Any]]: - """Serialize BatchOperatorTrigger arguments and classpath.""" - return ( - "airflow.providers.amazon.aws.triggers.batch.BatchOperatorTrigger", - { - "job_id": self.job_id, - "max_retries": self.max_retries, - "aws_conn_id": self.aws_conn_id, - "region_name": self.region_name, - "poll_interval": self.poll_interval, - }, - ) - - @cached_property - def hook(self) -> BatchClientHook: - return BatchClientHook(aws_conn_id=self.aws_conn_id, region_name=self.region_name) - - async def run(self): - async with self.hook.async_conn as client: - waiter = self.hook.get_waiter("batch_job_complete", deferrable=True, client=client) - for attempt in range(1, 1 + self.max_retries): - try: - await waiter.wait( - jobs=[self.job_id], - WaiterConfig={ - "Delay": self.poll_interval, - "MaxAttempts": 1, - }, - ) - except WaiterError as error: - if "terminal failure" in str(error): - yield TriggerEvent( - {"status": "failure", "message": f"Delete Cluster Failed: {error}"} - ) - break - self.log.info( - "Job status is %s. Retrying attempt %s/%s", - error.last_response["jobs"][0]["status"], - attempt, - self.max_retries, - ) - await asyncio.sleep(int(self.poll_interval)) - else: - yield TriggerEvent({"status": "success", "job_id": self.job_id}) - break - else: - yield TriggerEvent({"status": "failure", "message": "Job Failed - max attempts reached."}) - - -@deprecated(reason="use BatchJobTrigger instead", category=AirflowProviderDeprecationWarning) -class BatchSensorTrigger(BaseTrigger): - """ - Checks for the status of a submitted job_id to AWS Batch until it reaches a failure or a success state. - - BatchSensorTrigger is fired as deferred class with params to poll the job state in Triggerer. - - :param job_id: the job ID, to poll for job completion or not - :param region_name: AWS region name to use - Override the region_name in connection (if provided) - :param aws_conn_id: connection id of AWS credentials / region name. If None, - credential boto3 strategy will be used - :param poke_interval: polling period in seconds to check for the status of the job - """ - - def __init__( - self, - job_id: str, - region_name: str | None, - aws_conn_id: str | None = "aws_default", - poke_interval: float = 5, - ): - super().__init__() - self.job_id = job_id - self.aws_conn_id = aws_conn_id - self.region_name = region_name - self.poke_interval = poke_interval - - def serialize(self) -> tuple[str, dict[str, Any]]: - """Serialize BatchSensorTrigger arguments and classpath.""" - return ( - "airflow.providers.amazon.aws.triggers.batch.BatchSensorTrigger", - { - "job_id": self.job_id, - "aws_conn_id": self.aws_conn_id, - "region_name": self.region_name, - "poke_interval": self.poke_interval, - }, - ) - - @cached_property - def hook(self) -> BatchClientHook: - return BatchClientHook(aws_conn_id=self.aws_conn_id, region_name=self.region_name) - - async def run(self): - """ - Make async connection using aiobotocore library to AWS Batch, periodically poll for the job status. - - The status that indicates job completion are: 'SUCCEEDED'|'FAILED'. - """ - async with self.hook.async_conn as client: - waiter = self.hook.get_waiter("batch_job_complete", deferrable=True, client=client) - for attempt in itertools.count(1): - try: - await waiter.wait( - jobs=[self.job_id], - WaiterConfig={ - "Delay": int(self.poke_interval), - "MaxAttempts": 1, - }, - ) - except WaiterError as error: - if "error" in str(error): - yield TriggerEvent({"status": "failure", "message": f"Job Failed: {error}"}) - break - self.log.info( - "Job response is %s. Retrying attempt %s", - error.last_response["Error"]["Message"], - attempt, - ) - await asyncio.sleep(int(self.poke_interval)) - else: - break - - yield TriggerEvent( - { - "status": "success", - "job_id": self.job_id, - "message": f"Job {self.job_id} Succeeded", - } - ) - - class BatchJobTrigger(AwsBaseWaiterTrigger): """ Checks for the status of a submitted job_id to AWS Batch until it reaches a failure or a success state. diff --git a/airflow/providers/amazon/aws/triggers/eks.py b/airflow/providers/amazon/aws/triggers/eks.py index 02f17853b791c..f9b3625644d19 100644 --- a/airflow/providers/amazon/aws/triggers/eks.py +++ b/airflow/providers/amazon/aws/triggers/eks.py @@ -16,12 +16,11 @@ # under the License. from __future__ import annotations -import warnings from typing import TYPE_CHECKING, Any from botocore.exceptions import ClientError -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.eks import EksHook from airflow.providers.amazon.aws.triggers.base import AwsBaseWaiterTrigger from airflow.providers.amazon.aws.utils.waiter_with_logging import async_wait @@ -235,17 +234,8 @@ def __init__( waiter_delay: int, waiter_max_attempts: int, aws_conn_id: str | None, - region: str | None = None, region_name: str | None = None, ): - if region is not None: - warnings.warn( - "please use region_name param instead of region", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - region_name = region - super().__init__( serialized_fields={"cluster_name": cluster_name, "fargate_profile_name": fargate_profile_name}, waiter_name="fargate_profile_active", @@ -282,17 +272,8 @@ def __init__( waiter_delay: int, waiter_max_attempts: int, aws_conn_id: str | None, - region: str | None = None, region_name: str | None = None, ): - if region is not None: - warnings.warn( - "please use region_name param instead of region", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - region_name = region - super().__init__( serialized_fields={"cluster_name": cluster_name, "fargate_profile_name": fargate_profile_name}, waiter_name="fargate_profile_deleted", diff --git a/airflow/providers/amazon/aws/triggers/emr.py b/airflow/providers/amazon/aws/triggers/emr.py index 9abfe120d24c6..f1921826d76f1 100644 --- a/airflow/providers/amazon/aws/triggers/emr.py +++ b/airflow/providers/amazon/aws/triggers/emr.py @@ -17,10 +17,8 @@ from __future__ import annotations import sys -import warnings from typing import TYPE_CHECKING -from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.providers.amazon.aws.hooks.emr import EmrContainerHook, EmrHook, EmrServerlessHook from airflow.providers.amazon.aws.triggers.base import AwsBaseWaiterTrigger @@ -81,21 +79,10 @@ class EmrCreateJobFlowTrigger(AwsBaseWaiterTrigger): def __init__( self, job_flow_id: str, - poll_interval: int | None = None, # deprecated - max_attempts: int | None = None, # deprecated aws_conn_id: str | None = None, waiter_delay: int = 30, waiter_max_attempts: int = 60, ): - if poll_interval is not None or max_attempts is not None: - warnings.warn( - "please use waiter_delay instead of poll_interval " - "and waiter_max_attempts instead of max_attempts", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - waiter_delay = poll_interval or waiter_delay - waiter_max_attempts = max_attempts or waiter_max_attempts super().__init__( serialized_fields={"job_flow_id": job_flow_id}, waiter_name="job_flow_waiting", @@ -131,21 +118,10 @@ class EmrTerminateJobFlowTrigger(AwsBaseWaiterTrigger): def __init__( self, job_flow_id: str, - poll_interval: int | None = None, # deprecated - max_attempts: int | None = None, # deprecated aws_conn_id: str | None = None, waiter_delay: int = 30, waiter_max_attempts: int = 60, ): - if poll_interval is not None or max_attempts is not None: - warnings.warn( - "please use waiter_delay instead of poll_interval " - "and waiter_max_attempts instead of max_attempts", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - waiter_delay = poll_interval or waiter_delay - waiter_max_attempts = max_attempts or waiter_max_attempts super().__init__( serialized_fields={"job_flow_id": job_flow_id}, waiter_name="job_flow_terminated", @@ -183,17 +159,9 @@ def __init__( virtual_cluster_id: str, job_id: str, aws_conn_id: str | None = "aws_default", - poll_interval: int | None = None, # deprecated waiter_delay: int = 30, waiter_max_attempts: int = sys.maxsize, ): - if poll_interval is not None: - warnings.warn( - "please use waiter_delay instead of poll_interval.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - waiter_delay = poll_interval or waiter_delay super().__init__( serialized_fields={"virtual_cluster_id": virtual_cluster_id, "job_id": job_id}, waiter_name="container_job_complete", diff --git a/airflow/providers/amazon/aws/triggers/glue_crawler.py b/airflow/providers/amazon/aws/triggers/glue_crawler.py index 6bd52bef6f554..2116740a8faa1 100644 --- a/airflow/providers/amazon/aws/triggers/glue_crawler.py +++ b/airflow/providers/amazon/aws/triggers/glue_crawler.py @@ -16,10 +16,8 @@ # under the License. from __future__ import annotations -import warnings from typing import TYPE_CHECKING -from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.providers.amazon.aws.hooks.glue_crawler import GlueCrawlerHook from airflow.providers.amazon.aws.triggers.base import AwsBaseWaiterTrigger @@ -32,26 +30,17 @@ class GlueCrawlerCompleteTrigger(AwsBaseWaiterTrigger): Watches for a glue crawl, triggers when it finishes. :param crawler_name: name of the crawler to watch - :param poll_interval: The amount of time in seconds to wait between attempts. :param aws_conn_id: The Airflow connection used for AWS credentials. """ def __init__( self, crawler_name: str, - poll_interval: int | None = None, aws_conn_id: str | None = "aws_default", waiter_delay: int = 5, waiter_max_attempts: int = 1500, **kwargs, ): - if poll_interval is not None: - warnings.warn( - "please use waiter_delay instead of poll_interval.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - waiter_delay = poll_interval or waiter_delay super().__init__( serialized_fields={"crawler_name": crawler_name}, waiter_name="crawler_ready", diff --git a/airflow/providers/amazon/aws/triggers/glue_databrew.py b/airflow/providers/amazon/aws/triggers/glue_databrew.py index a57dc8a0d0359..e0bb992753ea4 100644 --- a/airflow/providers/amazon/aws/triggers/glue_databrew.py +++ b/airflow/providers/amazon/aws/triggers/glue_databrew.py @@ -17,9 +17,6 @@ from __future__ import annotations -import warnings - -from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.providers.amazon.aws.hooks.glue_databrew import GlueDataBrewHook from airflow.providers.amazon.aws.triggers.base import AwsBaseWaiterTrigger @@ -30,9 +27,7 @@ class GlueDataBrewJobCompleteTrigger(AwsBaseWaiterTrigger): :param job_name: Glue DataBrew job name :param run_id: the ID of the specific run to watch for that job - :param delay: Number of seconds to wait between two checks.(Deprecated). :param waiter_delay: Number of seconds to wait between two checks. Default is 30 seconds. - :param max_attempts: Maximum number of attempts to wait for the job to complete.(Deprecated). :param waiter_max_attempts: Maximum number of attempts to wait for the job to complete. Default is 60 attempts. :param aws_conn_id: The Airflow connection used for AWS credentials. """ @@ -41,27 +36,11 @@ def __init__( self, job_name: str, run_id: str, - delay: int | None = None, - max_attempts: int | None = None, waiter_delay: int = 30, waiter_max_attempts: int = 60, aws_conn_id: str | None = "aws_default", **kwargs, ): - if delay is not None: - warnings.warn( - "please use `waiter_delay` instead of delay.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - waiter_delay = delay or waiter_delay - if max_attempts is not None: - warnings.warn( - "please use `waiter_max_attempts` instead of max_attempts.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - waiter_max_attempts = max_attempts or waiter_max_attempts super().__init__( serialized_fields={"job_name": job_name, "run_id": run_id}, waiter_name="job_complete", diff --git a/airflow/providers/amazon/aws/triggers/rds.py b/airflow/providers/amazon/aws/triggers/rds.py index 7aab22155a414..dd60037c69787 100644 --- a/airflow/providers/amazon/aws/triggers/rds.py +++ b/airflow/providers/amazon/aws/triggers/rds.py @@ -16,95 +16,16 @@ # under the License. from __future__ import annotations -from functools import cached_property from typing import TYPE_CHECKING, Any -from deprecated import deprecated - -from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.providers.amazon.aws.hooks.rds import RdsHook from airflow.providers.amazon.aws.triggers.base import AwsBaseWaiterTrigger from airflow.providers.amazon.aws.utils.rds import RdsDbType -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 -@deprecated( - reason=( - "This trigger is deprecated, please use the other RDS triggers " - "such as RdsDbDeletedTrigger, RdsDbStoppedTrigger or RdsDbAvailableTrigger" - ), - category=AirflowProviderDeprecationWarning, -) -class RdsDbInstanceTrigger(BaseTrigger): - """ - Deprecated Trigger for RDS operations. Do not use. - - :param waiter_name: Name of the waiter to use, for instance 'db_instance_available' - or 'db_instance_deleted'. - :param db_instance_identifier: The DB instance identifier for the DB instance to be polled. - :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 where the DB is located, if different from the default one. - :param response: The response from the RdsHook, to be passed back to the operator. - """ - - def __init__( - self, - waiter_name: str, - db_instance_identifier: str, - waiter_delay: int, - waiter_max_attempts: int, - aws_conn_id: str | None, - region_name: str | None, - response: dict[str, Any], - ): - self.db_instance_identifier = db_instance_identifier - 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.waiter_name = waiter_name - self.response = response - - def serialize(self) -> tuple[str, dict[str, Any]]: - return ( - # dynamically generate the fully qualified name of the class - self.__class__.__module__ + "." + self.__class__.__qualname__, - { - "db_instance_identifier": self.db_instance_identifier, - "waiter_delay": str(self.waiter_delay), - "waiter_max_attempts": str(self.waiter_max_attempts), - "aws_conn_id": self.aws_conn_id, - "region_name": self.region_name, - "waiter_name": self.waiter_name, - "response": self.response, - }, - ) - - @cached_property - def hook(self) -> RdsHook: - return RdsHook(aws_conn_id=self.aws_conn_id, region_name=self.region_name) - - async def run(self): - async with self.hook.async_conn as client: - waiter = client.get_waiter(self.waiter_name) - await async_wait( - waiter=waiter, - waiter_delay=int(self.waiter_delay), - waiter_max_attempts=int(self.waiter_max_attempts), - args={"DBInstanceIdentifier": self.db_instance_identifier}, - failure_message="Error checking DB Instance status", - status_message="DB instance status is", - status_args=["DBInstances[0].DBInstanceStatus"], - ) - yield TriggerEvent({"status": "success", "response": self.response}) - - _waiter_arg = { RdsDbType.INSTANCE.value: "DBInstanceIdentifier", RdsDbType.CLUSTER.value: "DBClusterIdentifier", diff --git a/airflow/providers/amazon/aws/triggers/redshift_cluster.py b/airflow/providers/amazon/aws/triggers/redshift_cluster.py index 37ca6b6d9fb12..0fd55e322e5aa 100644 --- a/airflow/providers/amazon/aws/triggers/redshift_cluster.py +++ b/airflow/providers/amazon/aws/triggers/redshift_cluster.py @@ -17,11 +17,9 @@ from __future__ import annotations import asyncio -import warnings from typing import TYPE_CHECKING, Any, AsyncIterator -from airflow.exceptions import AirflowProviderDeprecationWarning -from airflow.providers.amazon.aws.hooks.redshift_cluster import RedshiftAsyncHook, RedshiftHook +from airflow.providers.amazon.aws.hooks.redshift_cluster import RedshiftHook from airflow.providers.amazon.aws.triggers.base import AwsBaseWaiterTrigger from airflow.triggers.base import BaseTrigger, TriggerEvent @@ -45,21 +43,10 @@ class RedshiftCreateClusterTrigger(AwsBaseWaiterTrigger): def __init__( self, cluster_identifier: str, - poll_interval: int | None = None, - max_attempt: int | None = None, aws_conn_id: str | None = "aws_default", waiter_delay: int = 15, waiter_max_attempts: int = 999999, ): - if poll_interval is not None or max_attempt is not None: - warnings.warn( - "please use waiter_delay instead of poll_interval " - "and waiter_max_attempts instead of max_attempt.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - waiter_delay = poll_interval or waiter_delay - waiter_max_attempts = max_attempt or waiter_max_attempts super().__init__( serialized_fields={"cluster_identifier": cluster_identifier}, waiter_name="cluster_available", @@ -93,21 +80,10 @@ class RedshiftPauseClusterTrigger(AwsBaseWaiterTrigger): def __init__( self, cluster_identifier: str, - poll_interval: int | None = None, - max_attempts: int | None = None, aws_conn_id: str | None = "aws_default", waiter_delay: int = 15, waiter_max_attempts: int = 999999, ): - if poll_interval is not None or max_attempts is not None: - warnings.warn( - "please use waiter_delay instead of poll_interval " - "and waiter_max_attempts instead of max_attempt.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - waiter_delay = poll_interval or waiter_delay - waiter_max_attempts = max_attempts or waiter_max_attempts super().__init__( serialized_fields={"cluster_identifier": cluster_identifier}, waiter_name="cluster_paused", @@ -141,21 +117,10 @@ class RedshiftCreateClusterSnapshotTrigger(AwsBaseWaiterTrigger): def __init__( self, cluster_identifier: str, - poll_interval: int | None = None, - max_attempts: int | None = None, aws_conn_id: str | None = "aws_default", waiter_delay: int = 15, waiter_max_attempts: int = 999999, ): - if poll_interval is not None or max_attempts is not None: - warnings.warn( - "please use waiter_delay instead of poll_interval " - "and waiter_max_attempts instead of max_attempt.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - waiter_delay = poll_interval or waiter_delay - waiter_max_attempts = max_attempts or waiter_max_attempts super().__init__( serialized_fields={"cluster_identifier": cluster_identifier}, waiter_name="snapshot_available", @@ -189,21 +154,10 @@ class RedshiftResumeClusterTrigger(AwsBaseWaiterTrigger): def __init__( self, cluster_identifier: str, - poll_interval: int | None = None, - max_attempts: int | None = None, aws_conn_id: str | None = "aws_default", waiter_delay: int = 15, waiter_max_attempts: int = 999999, ): - if poll_interval is not None or max_attempts is not None: - warnings.warn( - "please use waiter_delay instead of poll_interval " - "and waiter_max_attempts instead of max_attempt.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - waiter_delay = poll_interval or waiter_delay - waiter_max_attempts = max_attempts or waiter_max_attempts super().__init__( serialized_fields={"cluster_identifier": cluster_identifier}, waiter_name="cluster_resumed", @@ -234,21 +188,10 @@ class RedshiftDeleteClusterTrigger(AwsBaseWaiterTrigger): def __init__( self, cluster_identifier: str, - poll_interval: int | None = None, - max_attempts: int | None = None, aws_conn_id: str | None = "aws_default", waiter_delay: int = 30, waiter_max_attempts: int = 30, ): - if poll_interval is not None or max_attempts is not None: - warnings.warn( - "please use waiter_delay instead of poll_interval " - "and waiter_max_attempts instead of max_attempt.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - waiter_delay = poll_interval or waiter_delay - waiter_max_attempts = max_attempts or waiter_max_attempts super().__init__( serialized_fields={"cluster_identifier": cluster_identifier}, waiter_name="cluster_deleted", @@ -304,13 +247,11 @@ def serialize(self) -> tuple[str, dict[str, Any]]: async def run(self) -> AsyncIterator[TriggerEvent]: """Run async until the cluster status matches the target status.""" try: - hook = RedshiftAsyncHook(aws_conn_id=self.aws_conn_id) + hook = RedshiftHook(aws_conn_id=self.aws_conn_id) while True: - res = await hook.cluster_status(self.cluster_identifier) - if (res["status"] == "success" and res["cluster_state"] == self.target_status) or res[ - "status" - ] == "error": - yield TriggerEvent(res) + status = await hook.cluster_status_async(self.cluster_identifier) + if status == self.target_status: + yield TriggerEvent({"status": "success", "message": "target state met"}) return await asyncio.sleep(self.poke_interval) except Exception as e: diff --git a/airflow/providers/amazon/aws/triggers/sagemaker.py b/airflow/providers/amazon/aws/triggers/sagemaker.py index e8ee8cd9c2c9b..5ac6ea1abbacf 100644 --- a/airflow/providers/amazon/aws/triggers/sagemaker.py +++ b/airflow/providers/amazon/aws/triggers/sagemaker.py @@ -18,17 +18,15 @@ from __future__ import annotations import asyncio -import time from collections import Counter from enum import IntEnum from functools import cached_property from typing import Any, AsyncIterator from botocore.exceptions import WaiterError -from deprecated import deprecated -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning -from airflow.providers.amazon.aws.hooks.sagemaker import LogState, SageMakerHook +from airflow.exceptions import AirflowException +from airflow.providers.amazon.aws.hooks.sagemaker import SageMakerHook from airflow.providers.amazon.aws.utils.waiter_with_logging import async_wait from airflow.triggers.base import BaseTrigger, TriggerEvent @@ -198,92 +196,3 @@ async def run(self) -> AsyncIterator[TriggerEvent]: await asyncio.sleep(int(self.waiter_delay)) raise AirflowException("Waiter error: max attempts reached") - - -@deprecated( - reason=( - "`airflow.providers.amazon.aws.triggers.sagemaker.SageMakerTrainingPrintLogTrigger` " - "has been deprecated and will be removed in future. Please use ``SageMakerTrigger`` instead." - ), - category=AirflowProviderDeprecationWarning, -) -class SageMakerTrainingPrintLogTrigger(BaseTrigger): - """ - SageMakerTrainingPrintLogTrigger is fired as deferred class with params to run the task in triggerer. - - :param job_name: name of the job to check status - :param poke_interval: polling period in seconds to check for the status - :param aws_conn_id: AWS connection ID for sagemaker - """ - - def __init__( - self, - job_name: str, - poke_interval: float, - aws_conn_id: str | None = "aws_default", - ): - super().__init__() - self.job_name = job_name - self.poke_interval = poke_interval - self.aws_conn_id = aws_conn_id - - def serialize(self) -> tuple[str, dict[str, Any]]: - """Serialize SageMakerTrainingPrintLogTrigger arguments and classpath.""" - return ( - "airflow.providers.amazon.aws.triggers.sagemaker.SageMakerTrainingPrintLogTrigger", - { - "poke_interval": self.poke_interval, - "aws_conn_id": self.aws_conn_id, - "job_name": self.job_name, - }, - ) - - @cached_property - def hook(self) -> SageMakerHook: - return SageMakerHook(aws_conn_id=self.aws_conn_id) - - async def run(self) -> AsyncIterator[TriggerEvent]: - """Make async connection to sagemaker async hook and gets job status for a job submitted by the operator.""" - stream_names: list[str] = [] # The list of log streams - positions: dict[str, Any] = {} # The current position in each stream, map of stream name -> position - - last_description = await self.hook.describe_training_job_async(self.job_name) - instance_count = last_description["ResourceConfig"]["InstanceCount"] - status = last_description["TrainingJobStatus"] - job_already_completed = status not in self.hook.non_terminal_states - state = LogState.COMPLETE if job_already_completed else LogState.TAILING - last_describe_job_call = time.time() - try: - while True: - ( - state, - last_description, - last_describe_job_call, - ) = await self.hook.describe_training_job_with_log_async( - self.job_name, - positions, - stream_names, - instance_count, - state, - last_description, - last_describe_job_call, - ) - status = last_description["TrainingJobStatus"] - if status in self.hook.non_terminal_states: - await asyncio.sleep(self.poke_interval) - elif status in self.hook.failed_states: - reason = last_description.get("FailureReason", "(No reason provided)") - error_message = f"SageMaker job failed because {reason}" - yield TriggerEvent({"status": "error", "message": error_message}) - return - else: - billable_seconds = SageMakerHook.count_billable_seconds( - training_start_time=last_description["TrainingStartTime"], - training_end_time=last_description["TrainingEndTime"], - instance_count=instance_count, - ) - self.log.info("Billable seconds: %d", billable_seconds) - yield TriggerEvent({"status": "success", "message": last_description}) - return - except Exception as e: - yield TriggerEvent({"status": "error", "message": str(e)}) diff --git a/airflow/providers/amazon/aws/utils/connection_wrapper.py b/airflow/providers/amazon/aws/utils/connection_wrapper.py index eaeb320a5a1a6..de5a30120fdf0 100644 --- a/airflow/providers/amazon/aws/utils/connection_wrapper.py +++ b/airflow/providers/amazon/aws/utils/connection_wrapper.py @@ -25,12 +25,10 @@ from botocore import UNSIGNED from botocore.config import Config -from deprecated import deprecated -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.utils import trim_none_values from airflow.utils.log.logging_mixin import LoggingMixin -from airflow.utils.log.secrets_masker import mask_secret from airflow.utils.types import NOTSET, ArgNotSet if TYPE_CHECKING: @@ -157,15 +155,6 @@ def get_service_endpoint_url( "Can't resolve STS endpoint when both " "`sts_connection` and `sts_test_connection` set to True." ) - elif sts_test_connection: - if "test_endpoint_url" in self.extra_config: - warnings.warn( - "extra['test_endpoint_url'] is deprecated and will be removed in a future release." - " Please set `endpoint_url` in `service_config.sts` within `extras`.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - global_endpoint_url = self.extra_config["test_endpoint_url"] return service_config.get("endpoint_url", global_endpoint_url) @@ -208,15 +197,7 @@ def __post_init__(self, conn: Connection | AwsConnectionWrapper | _ConnectionMet self.schema = conn.schema or None self.extra_config = deepcopy(conn.extra_dejson) - if self.conn_type.lower() == "s3": - warnings.warn( - f"{self.conn_repr} has connection type 's3', " - "which has been replaced by connection type 'aws'. " - "Please update your connection to have `conn_type='aws'`.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - elif self.conn_type != "aws": + if self.conn_type != "aws": warnings.warn( f"{self.conn_repr} expected connection type 'aws', got {self.conn_type!r}. " "This connection might not work correctly. " @@ -228,15 +209,6 @@ def __post_init__(self, conn: Connection | AwsConnectionWrapper | _ConnectionMet extra = deepcopy(conn.extra_dejson) self.service_config = extra.get("service_config", {}) - session_kwargs = extra.get("session_kwargs", {}) - if session_kwargs: - warnings.warn( - "'session_kwargs' in extra config is deprecated and will be removed in a future releases. " - f"Please specify arguments passed to boto3 Session directly in {self.conn_repr} extra.", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - # Retrieve initial connection credentials init_credentials = self._get_credentials(**extra) self.aws_access_key_id, self.aws_secret_access_key, self.aws_session_token = init_credentials @@ -245,13 +217,6 @@ def __post_init__(self, conn: Connection | AwsConnectionWrapper | _ConnectionMet if "region_name" in extra: self.region_name = extra["region_name"] self.log.debug("Retrieving region_name=%s from %s extra.", self.region_name, self.conn_repr) - elif "region_name" in session_kwargs: - self.region_name = session_kwargs["region_name"] - self.log.debug( - "Retrieving region_name=%s from %s extra['session_kwargs'].", - self.region_name, - self.conn_repr, - ) if self.verify is None and "verify" in extra: self.verify = extra["verify"] @@ -260,13 +225,6 @@ def __post_init__(self, conn: Connection | AwsConnectionWrapper | _ConnectionMet if "profile_name" in extra: self.profile_name = extra["profile_name"] self.log.debug("Retrieving profile_name=%s from %s extra.", self.profile_name, self.conn_repr) - elif "profile_name" in session_kwargs: - self.profile_name = session_kwargs["profile_name"] - self.log.debug( - "Retrieving profile_name=%s from %s extra['session_kwargs'].", - self.profile_name, - self.conn_repr, - ) # Warn the user that an invalid parameter is being used which actually not related to 'profile_name'. # ToDo: Remove this check entirely as soon as drop support credentials from s3_config_file @@ -287,24 +245,7 @@ def __post_init__(self, conn: Connection | AwsConnectionWrapper | _ConnectionMet config_kwargs["signature_version"] = UNSIGNED self.botocore_config = Config(**config_kwargs) - if conn.host: - warnings.warn( - f"Host {conn.host} specified in the connection is not used." - " Please, set it on extra['endpoint_url'] instead", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - - self.endpoint_url = extra.get("host") - if self.endpoint_url: - warnings.warn( - "extra['host'] is deprecated and will be removed in a future release." - " Please set extra['endpoint_url'] instead", - AirflowProviderDeprecationWarning, - stacklevel=2, - ) - else: - self.endpoint_url = extra.get("endpoint_url") + self.endpoint_url = extra.get("endpoint_url") # Retrieve Assume Role Configuration assume_role_configs = self._get_assume_role_configs(**extra) @@ -359,10 +300,6 @@ def _get_credentials( aws_access_key_id: str | None = None, aws_secret_access_key: str | None = None, aws_session_token: str | None = None, - # Deprecated Values - s3_config_file: str | None = None, - s3_config_format: str | None = None, - profile: str | None = None, session_kwargs: dict[str, Any] | None = None, **kwargs, ) -> tuple[str | None, str | None, str | None]: @@ -393,13 +330,6 @@ def _get_credentials( aws_access_key_id = session_aws_access_key_id aws_secret_access_key = session_aws_secret_access_key self.log.info("%s credentials retrieved from extra['session_kwargs'].", self.conn_repr) - elif s3_config_file: - aws_access_key_id, aws_secret_access_key = _parse_s3_config( - s3_config_file, - s3_config_format, - profile, - ) - self.log.info("%s credentials retrieved from extra['s3_config_file']", self.conn_repr) if aws_session_token: self.log.info( @@ -422,31 +352,12 @@ def _get_assume_role_configs( role_arn: str | None = None, assume_role_method: str = "assume_role", assume_role_kwargs: dict[str, Any] | None = None, - # Deprecated Values - aws_account_id: str | None = None, - aws_iam_role: str | None = None, - external_id: str | None = None, **kwargs, ) -> tuple[str | None, str | None, dict[Any, str]]: """Get assume role configs from Connection extra.""" if role_arn: self.log.debug("Retrieving role_arn=%r from %s extra.", role_arn, self.conn_repr) - elif aws_account_id and aws_iam_role: - warnings.warn( - "Constructing 'role_arn' from extra['aws_account_id'] and extra['aws_iam_role'] is deprecated" - f" and will be removed in a future releases." - f" Please set 'role_arn' in {self.conn_repr} extra.", - AirflowProviderDeprecationWarning, - stacklevel=3, - ) - role_arn = f"arn:aws:iam::{aws_account_id}:role/{aws_iam_role}" - self.log.debug( - "Constructions role_arn=%r from %s extra['aws_account_id'] and extra['aws_iam_role'].", - role_arn, - self.conn_repr, - ) - - if not role_arn: + else: # There is no reason obtain `assume_role_method` and `assume_role_kwargs` if `role_arn` not set. return None, None, {} @@ -460,76 +371,5 @@ def _get_assume_role_configs( self.log.debug("Retrieve assume_role_method=%r from %s.", assume_role_method, self.conn_repr) assume_role_kwargs = assume_role_kwargs or {} - if "ExternalId" not in assume_role_kwargs and external_id: - warnings.warn( - "'external_id' in extra config is deprecated and will be removed in a future releases. " - f"Please set 'ExternalId' in 'assume_role_kwargs' in {self.conn_repr} extra.", - AirflowProviderDeprecationWarning, - stacklevel=3, - ) - assume_role_kwargs["ExternalId"] = external_id return role_arn, assume_role_method, assume_role_kwargs - - -@deprecated( - reason=( - "Use local credentials file is never documented and well tested. " - "Obtain credentials by this way deprecated and will be removed in a future releases." - ), - category=AirflowProviderDeprecationWarning, -) -def _parse_s3_config( - config_file_name: str, config_format: str | None = "boto", profile: str | None = None -) -> tuple[str | None, str | None]: - """ - Parse a config file for S3 credentials. - - Can currently parse boto, s3cmd.conf and AWS SDK config formats. - - :param config_file_name: path to the config file - :param config_format: config type. One of "boto", "s3cmd" or "aws". - Defaults to "boto" - :param profile: profile name in AWS type config file - """ - import configparser - - config = configparser.ConfigParser() - try: - if config.read(config_file_name): # pragma: no cover - sections = config.sections() - else: - raise AirflowException(f"Couldn't read {config_file_name}") - except Exception as e: - raise AirflowException("Exception when parsing %s: %s", config_file_name, e.__class__.__name__) - # Setting option names depending on file format - if config_format is None: - config_format = "boto" - conf_format = config_format.lower() - if conf_format == "boto": # pragma: no cover - if profile is not None and "profile " + profile in sections: - cred_section = "profile " + profile - else: - cred_section = "Credentials" - elif conf_format == "aws" and profile is not None: - cred_section = profile - else: - cred_section = "default" - # Option names - if conf_format in ("boto", "aws"): # pragma: no cover - key_id_option = "aws_access_key_id" - secret_key_option = "aws_secret_access_key" - else: - key_id_option = "access_key" - secret_key_option = "secret_key" - # Actual Parsing - if cred_section not in sections: - raise AirflowException("This config file format is not recognized") - else: - try: - access_key = config.get(cred_section, key_id_option) - secret_key = config.get(cred_section, secret_key_option) - mask_secret(secret_key) - except Exception: - raise AirflowException("Option Error in parsing s3 config file") - return access_key, secret_key diff --git a/airflow/providers/amazon/aws/utils/mixins.py b/airflow/providers/amazon/aws/utils/mixins.py index b9c7e30c23dad..9dbbde914874c 100644 --- a/airflow/providers/amazon/aws/utils/mixins.py +++ b/airflow/providers/amazon/aws/utils/mixins.py @@ -27,19 +27,15 @@ from __future__ import annotations -import warnings from functools import cached_property from typing import Any, Generic, NamedTuple, TypeVar -from deprecated import deprecated from typing_extensions import final from airflow.compat.functools import cache -from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.providers.amazon.aws.hooks.base_aws import AwsGenericHook AwsHookType = TypeVar("AwsHookType", bound=AwsGenericHook) -REGION_MSG = "`region` is deprecated and will be removed in the future. Please use `region_name` instead." class AwsHookParams(NamedTuple): @@ -90,13 +86,6 @@ def __init__( self.botocore_config = params.botocore_config self.foo = foo """ - if region := additional_params.pop("region", None): - warnings.warn(REGION_MSG, AirflowProviderDeprecationWarning, stacklevel=3) - if region_name and region_name != region: - raise ValueError( - f"Conflicting `region_name` provided, region_name={region_name!r}, region={region!r}." - ) - region_name = region return cls(aws_conn_id, region_name, verify, botocore_config) @@ -160,16 +149,6 @@ def hook(self) -> AwsHookType: """ return self.aws_hook_class(**self._hook_parameters) - @property - @final - @deprecated( - reason="`region` is deprecated and will be removed in the future. Please use `region_name` instead.", - category=AirflowProviderDeprecationWarning, - ) - def region(self) -> str | None: - """Alias for ``region_name``, used for compatibility (deprecated).""" - return self.region_name - @cache def aws_template_fields(*template_fields: str) -> tuple[str, ...]: diff --git a/docs/apache-airflow-providers-amazon/connections/aws.rst b/docs/apache-airflow-providers-amazon/connections/aws.rst index 3354da06350e6..a5850f0a8305b 100644 --- a/docs/apache-airflow-providers-amazon/connections/aws.rst +++ b/docs/apache-airflow-providers-amazon/connections/aws.rst @@ -136,22 +136,6 @@ Extra (optional) * ``service_config``: json used to specify configuration/parameters per AWS service / Amazon provider hook, for more details please refer to :ref:`howto/connection:aws:per-service-configuration`. -.. warning:: The extra parameters below are deprecated and will be removed in a future version of this provider. - - * ``aws_account_id``: Used to construct ``role_arn`` if it was not specified. - * ``aws_iam_role``: Used to construct ``role_arn`` if it was not specified. - * ``external_id``: A unique identifier that might be required when you assume a role in another account. - Used if ``ExternalId`` in ``assume_role_kwargs`` was not specified. - * ``s3_config_file``: Path to local credentials file. - * ``s3_config_format``: ``s3_config_file`` format, one of - `aws `_, - `boto `_ or - `s3cmd `_ if not specified then **boto** is used. - * ``profile``: If you are getting your credentials from the ``s3_config_file`` - you can specify the profile with this parameter. - * ``host``: Used as connection's URL. Use ``endpoint_url`` instead. - * ``session_kwargs``: Additional **kwargs** passed to :external:py:class:`boto3.session.Session`. - If you are configuring the connection via a URI, ensure that all components of the URI are URL-encoded. Examples diff --git a/tests/providers/amazon/aws/deferrable/__init__.py b/tests/providers/amazon/aws/deferrable/__init__.py deleted file mode 100644 index 13a83393a9124..0000000000000 --- a/tests/providers/amazon/aws/deferrable/__init__.py +++ /dev/null @@ -1,16 +0,0 @@ -# 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. diff --git a/tests/providers/amazon/aws/deferrable/hooks/__init__.py b/tests/providers/amazon/aws/deferrable/hooks/__init__.py deleted file mode 100644 index 13a83393a9124..0000000000000 --- a/tests/providers/amazon/aws/deferrable/hooks/__init__.py +++ /dev/null @@ -1,16 +0,0 @@ -# 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. diff --git a/tests/providers/amazon/aws/deferrable/hooks/test_base_aws.py b/tests/providers/amazon/aws/deferrable/hooks/test_base_aws.py deleted file mode 100644 index 1e9526ef89839..0000000000000 --- a/tests/providers/amazon/aws/deferrable/hooks/test_base_aws.py +++ /dev/null @@ -1,101 +0,0 @@ -# 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 contextlib -from unittest import mock -from unittest.mock import ANY - -import pytest - -from airflow.models.connection import Connection -from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseAsyncHook - -pytest.importorskip("aiobotocore") - -with contextlib.suppress(ImportError): - from aiobotocore.credentials import AioCredentials - - -class TestAwsBaseAsyncHook: - @staticmethod - def compare_aio_cred(first, second): - if type(first) is not type(second): - return False - if first.access_key != second.access_key: - return False - if first.secret_key != second.secret_key: - return False - if first.method != second.method: - return False - if first.token != second.token: - return False - return True - - class Matcher: - def __init__(self, compare, obj): - self.compare = compare - self.obj = obj - - def __eq__(self, other): - return self.compare(self.obj, other) - - @pytest.mark.asyncio - @mock.patch("airflow.hooks.base.BaseHook.get_connection") - @mock.patch("aiobotocore.session.AioClientCreator.create_client") - async def test_get_client_async(self, mock_client, mock_get_connection): - """Check the connection credential passed while creating client""" - mock_get_connection.return_value = Connection( - conn_id="aws_default1", - extra="""{ - "aws_secret_access_key": "mock_aws_access_key", - "aws_access_key_id": "mock_aws_access_key_id", - "region_name": "us-east-2" - }""", - ) - - hook = AwsBaseAsyncHook(client_type="S3", aws_conn_id="aws_default1") - cred = await hook.get_async_session().get_credentials() - - # check credential have same values as intended - assert cred.__dict__ == { - "access_key": "mock_aws_access_key_id", - "method": "explicit", - "secret_key": "mock_aws_access_key", - "token": None, - } - - aio_cred = AioCredentials( - access_key="mock_aws_access_key_id", - method="explicit", - secret_key="mock_aws_access_key", - ) - - await (await hook.get_client_async())._coro - # Test the aiobotocore client created with right param - mock_client.assert_called_once_with( - service_name="S3", - region_name="us-east-2", - is_secure=True, - endpoint_url=None, - verify=None, - credentials=self.Matcher(self.compare_aio_cred, aio_cred), - scoped_config={}, - client_config=ANY, - api_version=None, - auth_token=None, - ) diff --git a/tests/providers/amazon/aws/deferrable/hooks/test_redshift_cluster.py b/tests/providers/amazon/aws/deferrable/hooks/test_redshift_cluster.py deleted file mode 100644 index 2693524231681..0000000000000 --- a/tests/providers/amazon/aws/deferrable/hooks/test_redshift_cluster.py +++ /dev/null @@ -1,121 +0,0 @@ -# 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 -from unittest import mock - -import pytest -from botocore.exceptions import ClientError - -from airflow.providers.amazon.aws.hooks.redshift_cluster import RedshiftAsyncHook - -pytest.importorskip("aiobotocore") - - -class TestRedshiftAsyncHook: - @pytest.mark.asyncio - @mock.patch("aiobotocore.client.AioBaseClient._make_api_call") - async def test_cluster_status(self, mock_make_api_call): - """Test that describe_clusters get called with correct param""" - hook = RedshiftAsyncHook(aws_conn_id="aws_default", client_type="redshift", resource_type="redshift") - await hook.cluster_status(cluster_identifier="redshift_cluster_1") - mock_make_api_call.assert_called_once_with( - "DescribeClusters", {"ClusterIdentifier": "redshift_cluster_1"} - ) - - @pytest.mark.asyncio - @mock.patch("aiobotocore.client.AioBaseClient._make_api_call") - async def test_pause_cluster(self, mock_make_api_call): - """Test that pause_cluster get called with correct param""" - hook = RedshiftAsyncHook(aws_conn_id="aws_default", client_type="redshift", resource_type="redshift") - await hook.pause_cluster(cluster_identifier="redshift_cluster_1") - mock_make_api_call.assert_called_once_with( - "PauseCluster", {"ClusterIdentifier": "redshift_cluster_1"} - ) - - @pytest.mark.asyncio - @mock.patch("airflow.providers.amazon.aws.hooks.redshift_cluster.RedshiftAsyncHook.get_client_async") - @mock.patch("airflow.providers.amazon.aws.hooks.redshift_cluster.RedshiftAsyncHook.cluster_status") - async def test_get_cluster_status(self, cluster_status, mock_client): - """Test get_cluster_status async function with success response""" - flag = asyncio.Event() - cluster_status.return_value = {"status": "success", "cluster_state": "available"} - hook = RedshiftAsyncHook(aws_conn_id="aws_default") - result = await hook.get_cluster_status("redshift_cluster_1", "available", flag) - assert result == {"status": "success", "cluster_state": "available"} - - @pytest.mark.asyncio - @mock.patch("airflow.providers.amazon.aws.hooks.redshift_cluster.RedshiftAsyncHook.cluster_status") - async def test_get_cluster_status_exception(self, cluster_status): - """Test get_cluster_status async function with exception response""" - flag = asyncio.Event() - cluster_status.side_effect = ClientError( - { - "Error": { - "Code": "SomeServiceException", - "Message": "Details/context around the exception or error", - }, - }, - operation_name="redshift", - ) - hook = RedshiftAsyncHook(aws_conn_id="aws_default") - result = await hook.get_cluster_status("test-identifier", "available", flag) - assert result == { - "status": "error", - "message": "An error occurred (SomeServiceException) when calling the " - "redshift operation: Details/context around the exception or error", - } - - @pytest.mark.asyncio - @mock.patch("aiobotocore.client.AioBaseClient._make_api_call") - async def test_resume_cluster(self, mock_make_api_call): - """Test Resume cluster async hook function by mocking return value of resume_cluster""" - - hook = RedshiftAsyncHook() - await hook.resume_cluster(cluster_identifier="redshift_cluster_1") - mock_make_api_call.assert_called_once_with( - "ResumeCluster", {"ClusterIdentifier": "redshift_cluster_1"} - ) - - @pytest.mark.asyncio - @mock.patch("airflow.providers.amazon.aws.hooks.redshift_cluster.RedshiftAsyncHook.get_client_async") - async def test_resume_cluster_exception(self, mock_client): - """Test Resume cluster async hook function with exception by mocking return value of resume_cluster""" - mock_client.return_value.__aenter__.return_value.resume_cluster.side_effect = ClientError( - { - "Error": { - "Code": "SomeServiceException", - "Message": "Details/context around the exception or error", - }, - "ResponseMetadata": { - "RequestId": "1234567890ABCDEF", - "HostId": "host ID data will appear here as a hash", - "HTTPStatusCode": 500, - "HTTPHeaders": {"header metadata key/values will appear here"}, - "RetryAttempts": 0, - }, - }, - operation_name="redshift", - ) - hook = RedshiftAsyncHook(aws_conn_id="test_aws_connection_id") - result = await hook.resume_cluster(cluster_identifier="test") - assert result == { - "status": "error", - "message": "An error occurred (SomeServiceException) when calling the " - "redshift operation: Details/context around the exception or error", - } diff --git a/tests/providers/amazon/aws/hooks/test_base_aws.py b/tests/providers/amazon/aws/hooks/test_base_aws.py index 984de7e223f85..0957e6a928aae 100644 --- a/tests/providers/amazon/aws/hooks/test_base_aws.py +++ b/tests/providers/amazon/aws/hooks/test_base_aws.py @@ -39,7 +39,7 @@ from moto import mock_aws from moto.core import DEFAULT_ACCOUNT_ID -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.models.connection import Connection from airflow.providers.amazon.aws.executors.ecs.ecs_executor import AwsEcsExecutor from airflow.providers.amazon.aws.hooks.base_aws import ( @@ -917,26 +917,15 @@ def mock_error(): assert hook.client_type == "ec2" @pytest.mark.parametrize( - "sts_service_endpoint_url, test_endpoint_url, result_url", + "sts_service_endpoint_url, result_url", [ - pytest.param(None, None, None, id="not-set"), - pytest.param( - "https://sts.service:1234", None, "https://sts.service:1234", id="sts-service-endpoint" - ), - pytest.param( - None, "http://deprecated.test", "http://deprecated.test", id="deprecated-test-parameter" - ), - pytest.param( - "https://sts.service:1234", - "http://deprecated.test", - "https://sts.service:1234", - id="mixin-resolve", - ), + pytest.param(None, None, id="not-set"), + pytest.param("https://sts.service:1234", "https://sts.service:1234", id="sts-service-endpoint"), ], ) @mock.patch("boto3.session.Session") def test_hook_connection_endpoint_url_valid( - self, mock_boto3_session, sts_service_endpoint_url, test_endpoint_url, result_url, monkeypatch + self, mock_boto3_session, sts_service_endpoint_url, result_url, monkeypatch ): """Test if test_endpoint_url is valid in test connection""" @@ -945,12 +934,6 @@ def test_hook_connection_endpoint_url_valid( warn_context = nullcontext() fake_extra = {"endpoint_url": "https://test.conn:777/should/ignore/global/endpoint/url"} - if test_endpoint_url: - fake_extra["test_endpoint_url"] = test_endpoint_url - # If `test_endpoint_url` set than we raise warning message - warn_context = pytest.warns( - AirflowProviderDeprecationWarning, match=r"extra\['test_endpoint_url'\] is deprecated" - ) if sts_service_endpoint_url: fake_extra["service_config"] = {"sts": {"endpoint_url": sts_service_endpoint_url}} diff --git a/tests/providers/amazon/aws/hooks/test_quicksight.py b/tests/providers/amazon/aws/hooks/test_quicksight.py index 6a7795843b57d..687fad00fdc95 100644 --- a/tests/providers/amazon/aws/hooks/test_quicksight.py +++ b/tests/providers/amazon/aws/hooks/test_quicksight.py @@ -22,7 +22,7 @@ import pytest from botocore.exceptions import ClientError -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.quicksight import QuickSightHook DEFAULT_AWS_ACCOUNT_ID = "123456789012" @@ -255,13 +255,3 @@ def test_create_ingestion_exception(self, mocked_account_id, mocked_client, capl ingestion_type="INCREMENTAL_REFRESH", ) assert "create_ingestion API, error: Fake Error" in caplog.text - - def test_deprecated_properties(self): - hook = QuickSightHook(aws_conn_id=None, region_name="us-east-1") - with mock.patch("airflow.providers.amazon.aws.hooks.sts.StsHook") as mocked_class, pytest.warns( - AirflowProviderDeprecationWarning, match="consider to use `.*account_id` instead" - ): - mocked_sts_hook = mock.MagicMock(name="FakeStsHook") - mocked_class.return_value = mocked_sts_hook - assert hook.sts_hook is mocked_sts_hook - mocked_class.assert_called_once_with(aws_conn_id=None) diff --git a/tests/providers/amazon/aws/hooks/test_sagemaker.py b/tests/providers/amazon/aws/hooks/test_sagemaker.py index 7ff7e687a19aa..4915144c75fa6 100644 --- a/tests/providers/amazon/aws/hooks/test_sagemaker.py +++ b/tests/providers/amazon/aws/hooks/test_sagemaker.py @@ -27,7 +27,7 @@ from dateutil.tz import tzlocal from moto import mock_aws -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.logs import AwsLogsHook from airflow.providers.amazon.aws.hooks.s3 import S3Hook from airflow.providers.amazon.aws.hooks.sagemaker import ( @@ -728,23 +728,6 @@ def test_start_pipeline_returns_arn(self, mock_conn): # Value contains the value associated with the key in Name assert transformed_param["Value"] == params_dict[transformed_param["Name"]] - @patch("airflow.providers.amazon.aws.hooks.sagemaker.SageMakerHook.conn", new_callable=mock.PropertyMock) - def test_start_pipeline_waits_for_completion(self, mock_conn): - mock_conn().describe_pipeline_execution.side_effect = [ - {"PipelineExecutionStatus": "Executing"}, - {"PipelineExecutionStatus": "Executing"}, - {"PipelineExecutionStatus": "Succeeded"}, - ] - - hook = SageMakerHook(aws_conn_id="aws_default") - with pytest.warns( - AirflowProviderDeprecationWarning, - match="parameter `wait_for_completion` and `check_interval` are deprecated, remove them and call check_status yourself if you want to wait for completion", - ): - hook.start_pipeline(pipeline_name="test_name", wait_for_completion=True, check_interval=0) - - assert mock_conn().describe_pipeline_execution.call_count == 3 - @patch("airflow.providers.amazon.aws.hooks.sagemaker.SageMakerHook.conn", new_callable=mock.PropertyMock) def test_stop_pipeline_returns_status(self, mock_conn): mock_conn().describe_pipeline_execution.return_value = {"PipelineExecutionStatus": "Stopping"} @@ -755,50 +738,6 @@ def test_stop_pipeline_returns_status(self, mock_conn): assert pipeline_status == "Stopping" mock_conn().stop_pipeline_execution.assert_called_once_with(PipelineExecutionArn="test") - @patch("airflow.providers.amazon.aws.hooks.sagemaker.SageMakerHook.conn", new_callable=mock.PropertyMock) - def test_stop_pipeline_waits_for_completion(self, mock_conn): - mock_conn().describe_pipeline_execution.side_effect = [ - {"PipelineExecutionStatus": "Stopping"}, - {"PipelineExecutionStatus": "Stopping"}, - {"PipelineExecutionStatus": "Stopped"}, - ] - - hook = SageMakerHook(aws_conn_id="aws_default") - with pytest.warns( - AirflowProviderDeprecationWarning, - match="parameter `wait_for_completion` and `check_interval` are deprecated, remove them and call check_status yourself if you want to wait for completion", - ): - pipeline_status = hook.stop_pipeline( - pipeline_exec_arn="test", wait_for_completion=True, check_interval=0 - ) - - assert pipeline_status == "Stopped" - assert mock_conn().describe_pipeline_execution.call_count == 3 - - @patch("airflow.providers.amazon.aws.hooks.sagemaker.SageMakerHook.conn", new_callable=mock.PropertyMock) - def test_stop_pipeline_waits_for_completion_even_when_already_stopped(self, mock_conn): - mock_conn().stop_pipeline_execution.side_effect = ClientError( - error_response={ - "Error": {"Message": "Only pipelines with 'Executing' status can be stopped", "Code": "0"} - }, - operation_name="empty", - ) - mock_conn().describe_pipeline_execution.side_effect = [ - {"PipelineExecutionStatus": "Stopping"}, - {"PipelineExecutionStatus": "Stopped"}, - ] - - hook = SageMakerHook(aws_conn_id="aws_default") - with pytest.warns( - AirflowProviderDeprecationWarning, - match="parameter `wait_for_completion` and `check_interval` are deprecated, remove them and call check_status yourself if you want to wait for completion", - ): - pipeline_status = hook.stop_pipeline( - pipeline_exec_arn="test", wait_for_completion=True, check_interval=0 - ) - - assert pipeline_status == "Stopped" - @patch("airflow.providers.amazon.aws.hooks.sagemaker.SageMakerHook.conn", new_callable=mock.PropertyMock) def test_stop_pipeline_raises_when_already_stopped_if_specified(self, mock_conn): error = ClientError( diff --git a/tests/providers/amazon/aws/operators/test_appflow.py b/tests/providers/amazon/aws/operators/test_appflow.py index 58f79f6ea39df..6b9d3972eb69d 100644 --- a/tests/providers/amazon/aws/operators/test_appflow.py +++ b/tests/providers/amazon/aws/operators/test_appflow.py @@ -256,12 +256,3 @@ def test_base_aws_op_attributes(op_class, op_base_args): assert hook._region_name == "eu-west-1" assert hook._verify is False assert hook._config.read_timeout == 42 - - # Compatibility check: previously Appflow Operators use `region` instead of `region_name` - warning_message = "`region` is deprecated and will be removed in the future" - with pytest.warns(DeprecationWarning, match=warning_message): - op = op_class(**op_base_args, region="us-west-1") - assert op.region_name == "us-west-1" - - with pytest.warns(DeprecationWarning, match=warning_message): - assert op.region == "us-west-1" diff --git a/tests/providers/amazon/aws/operators/test_base_aws.py b/tests/providers/amazon/aws/operators/test_base_aws.py index 510ed03de243c..10d594363e66d 100644 --- a/tests/providers/amazon/aws/operators/test_base_aws.py +++ b/tests/providers/amazon/aws/operators/test_base_aws.py @@ -20,7 +20,6 @@ import pytest -from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.hooks.base import BaseHook from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseHook from airflow.providers.amazon.aws.operators.base_aws import AwsBaseOperator @@ -121,40 +120,6 @@ def test_execute(self, op_kwargs, dag_maker): tis = {ti.task_id: ti for ti in dagrun.task_instances} tis["fake-task-id"].run() - @pytest.mark.parametrize( - "region, region_name", - [ - pytest.param("eu-west-1", None, id="region-only"), - pytest.param("us-east-1", "us-east-1", id="non-ambiguous-params"), - ], - ) - def test_deprecated_region_name(self, region, region_name): - warning_match = r"`region` is deprecated and will be removed" - with pytest.warns(AirflowProviderDeprecationWarning, match=warning_match): - op = FakeS3Operator( - task_id="fake-task-id", - aws_conn_id=TEST_CONN, - region=region, - region_name=region_name, - ) - assert op.region_name == region - - with pytest.warns(AirflowProviderDeprecationWarning, match=warning_match): - assert op.region == region - - def test_conflicting_region_name(self): - error_match = r"Conflicting `region_name` provided, region_name='us-west-1', region='eu-west-1'" - with pytest.raises(ValueError, match=error_match), pytest.warns( - AirflowProviderDeprecationWarning, - match="`region` is deprecated and will be removed in the future. Please use `region_name` instead.", - ): - FakeS3Operator( - task_id="fake-task-id", - aws_conn_id=TEST_CONN, - region="eu-west-1", - region_name="us-west-1", - ) - def test_no_aws_hook_class_attr(self): class NoAwsHookClassOperator(AwsBaseOperator): ... @@ -179,45 +144,3 @@ class SoWrongOperator(AwsBaseOperator): error_match = r"Class attribute 'SoWrongOperator.aws_hook_class' is not a subclass of AwsGenericHook" with pytest.raises(AttributeError, match=error_match): SoWrongOperator(task_id="fake-task-id") - - @pytest.mark.skip_if_database_isolation_mode - @pytest.mark.parametrize( - "region, region_name, expected_region_name", - [ - pytest.param("ca-west-1", None, "ca-west-1", id="region-only"), - pytest.param("us-west-1", "us-west-1", "us-west-1", id="non-ambiguous-params"), - ], - ) - @pytest.mark.db_test - def test_region_in_partial_operator(self, region, region_name, expected_region_name, dag_maker): - with dag_maker("test_region_in_partial_operator", serialized=True): - FakeS3Operator.partial( - task_id="fake-task-id", - region=region, - region_name=region_name, - ).expand(value=[1, 2, 3]) - - dr = dag_maker.create_dagrun(execution_date=timezone.utcnow()) - warning_match = r"`region` is deprecated and will be removed" - for ti in dr.task_instances: - with pytest.warns(AirflowProviderDeprecationWarning, match=warning_match): - ti.run() - assert ti.task.region_name == expected_region_name - - @pytest.mark.skip_if_database_isolation_mode - @pytest.mark.db_test - def test_ambiguous_region_in_partial_operator(self, dag_maker): - with dag_maker("test_ambiguous_region_in_partial_operator", serialized=True): - FakeS3Operator.partial( - task_id="fake-task-id", - region="eu-west-1", - region_name="us-east-1", - ).expand(value=[1, 2, 3]) - - dr = dag_maker.create_dagrun(execution_date=timezone.utcnow()) - warning_match = r"`region` is deprecated and will be removed" - for ti in dr.task_instances: - with pytest.warns(AirflowProviderDeprecationWarning, match=warning_match), pytest.raises( - ValueError, match="Conflicting `region_name` provided" - ): - ti.run() diff --git a/tests/providers/amazon/aws/operators/test_batch.py b/tests/providers/amazon/aws/operators/test_batch.py index 7d9f27a6f4c5a..1389099e44452 100644 --- a/tests/providers/amazon/aws/operators/test_batch.py +++ b/tests/providers/amazon/aws/operators/test_batch.py @@ -23,7 +23,7 @@ import botocore.client import pytest -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning, TaskDeferred +from airflow.exceptions import AirflowException, TaskDeferred from airflow.providers.amazon.aws.hooks.batch_client import BatchClientHook from airflow.providers.amazon.aws.operators.batch import BatchCreateComputeEnvironmentOperator, BatchOperator @@ -32,7 +32,6 @@ BatchCreateComputeEnvironmentTrigger, BatchJobTrigger, ) -from airflow.utils.task_instance_session import set_current_task_instance_session AWS_REGION = "eu-west-1" AWS_ACCESS_KEY_ID = "airflow_dummy_key" @@ -397,7 +396,7 @@ def test_kill_job(self): self.client_mock.terminate_job.assert_called_once_with(jobId=JOB_ID, reason="Task killed by the user") @pytest.mark.parametrize( - "override", ["overrides", "node_overrides", "ecs_properties_override", "eks_properties_override"] + "override", ["node_overrides", "ecs_properties_override", "eks_properties_override"] ) @patch( "airflow.providers.amazon.aws.hooks.batch_client.BatchClientHook.client", @@ -409,32 +408,16 @@ def test_override_not_sent_if_not_set(self, client_mock, override): in the API call (which would create a validation error from boto) """ override_arg = {override: {"a": "a"}} - if override == "overrides": - with pytest.warns( - AirflowProviderDeprecationWarning, - match="Parameter `overrides` is deprecated, Please use `container_overrides` instead.", - ): - batch = BatchOperator( - task_id="task", - job_name=JOB_NAME, - job_queue="queue", - job_definition="hello-world", - **override_arg, - # setting those to bypass code that is not relevant here - do_xcom_push=False, - wait_for_completion=False, - ) - else: - batch = BatchOperator( - task_id="task", - job_name=JOB_NAME, - job_queue="queue", - job_definition="hello-world", - **override_arg, - # setting those to bypass code that is not relevant here - do_xcom_push=False, - wait_for_completion=False, - ) + batch = BatchOperator( + task_id="task", + job_name=JOB_NAME, + job_queue="queue", + job_definition="hello-world", + **override_arg, + # setting those to bypass code that is not relevant here + do_xcom_push=False, + wait_for_completion=False, + ) batch.execute(None) @@ -457,16 +440,6 @@ def test_override_not_sent_if_not_set(self, client_mock, override): client_mock().submit_job.assert_called_once_with(**expected_args) - def test_deprecated_override_param(self): - with pytest.warns(AirflowProviderDeprecationWarning): - _ = BatchOperator( - task_id="task", - job_name=JOB_NAME, - job_queue="queue", - job_definition="hello-world", - overrides={"a": "b"}, # <- the deprecated field - ) - def test_cant_set_old_and_new_override_param(self): with pytest.raises(AirflowException): _ = BatchOperator( @@ -607,36 +580,6 @@ def test_execute(self, mock_conn): tags=tags, ) - def test_deprecation(self): - with pytest.warns(AirflowProviderDeprecationWarning, match=self.warn_message): - BatchCreateComputeEnvironmentOperator( - task_id="id", - compute_environment_name="environment_name", - environment_type="environment_type", - state="environment_state", - compute_resources={}, - status_retries="Huh?", - ) - - @pytest.mark.db_test - def test_partial_deprecation(self, dag_maker, session): - with dag_maker(dag_id="test_partial_deprecation_waiters_params_reg_ecs", session=session): - BatchCreateComputeEnvironmentOperator.partial( - task_id="id", - compute_environment_name="environment_name", - environment_type="environment_type", - state="environment_state", - status_retries="Huh?", - ).expand(compute_resources=[{}, {}]) - - dr = dag_maker.create_dagrun() - tis = dr.get_task_instances(session=session) - with set_current_task_instance_session(session=session): - for ti in tis: - with pytest.warns(AirflowProviderDeprecationWarning, match=self.warn_message): - ti.render_templates() - assert not hasattr(ti.task, "status_retries") - @mock.patch.object(BatchClientHook, "client") def test_defer(self, client_mock): client_mock.create_compute_environment.return_value = {"computeEnvironmentArn": "my_arn"} diff --git a/tests/providers/amazon/aws/operators/test_ecs.py b/tests/providers/amazon/aws/operators/test_ecs.py index 5c2ba16c4cff2..d7b6c0d4e8716 100644 --- a/tests/providers/amazon/aws/operators/test_ecs.py +++ b/tests/providers/amazon/aws/operators/test_ecs.py @@ -24,7 +24,7 @@ import boto3 import pytest -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning, TaskDeferred +from airflow.exceptions import AirflowException, TaskDeferred from airflow.providers.amazon.aws.exceptions import EcsOperatorError, EcsTaskFailToStart from airflow.providers.amazon.aws.hooks.ecs import EcsClusterStates, EcsHook from airflow.providers.amazon.aws.operators.ecs import ( @@ -37,7 +37,6 @@ ) from airflow.providers.amazon.aws.triggers.ecs import TaskDoneTrigger from airflow.providers.amazon.aws.utils.task_log_fetcher import AwsTaskLogFetcher -from airflow.utils.task_instance_session import set_current_task_instance_session from airflow.utils.types import NOTSET from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields @@ -684,54 +683,6 @@ def test_execute_complete(self, client_mock): # task gets described to assert its success client_mock().describe_tasks.assert_called_once_with(cluster="test_cluster", tasks=["my_arn"]) - @pytest.mark.db_test - @pytest.mark.parametrize( - "region, region_name, expected_region_name", - [ - pytest.param("ca-west-1", None, "ca-west-1", id="region-only"), - pytest.param("us-west-1", "us-west-1", "us-west-1", id="non-ambiguous-params"), - ], - ) - def test_partial_deprecated_region(self, region, region_name, expected_region_name, dag_maker, session): - with dag_maker(dag_id="test_partial_deprecated_region_ecs", session=session): - EcsRunTaskOperator.partial( - task_id="fake-task-id", - region=region, - region_name=region_name, - cluster="foo", - task_definition="bar", - ).expand(overrides=[{}, {}, {}]) - - dr = dag_maker.create_dagrun() - tis = dr.get_task_instances(session=session) - with set_current_task_instance_session(session=session): - warning_match = r"`region` is deprecated and will be removed" - for ti in tis: - with pytest.warns(AirflowProviderDeprecationWarning, match=warning_match): - ti.render_templates() - assert ti.task.region_name == expected_region_name - - @pytest.mark.db_test - def test_partial_ambiguous_region(self, dag_maker, session): - with dag_maker("test_partial_ambiguous_region_ecs", session=session): - EcsRunTaskOperator.partial( - task_id="fake-task-id", - region="eu-west-1", - region_name="us-west-1", - cluster="foo", - task_definition="bar", - ).expand(overrides=[{}, {}, {}]) - - dr = dag_maker.create_dagrun(session=session) - tis = dr.get_task_instances(session=session) - with set_current_task_instance_session(session=session): - warning_match = r"`region` is deprecated and will be removed" - for ti in tis: - with pytest.warns(AirflowProviderDeprecationWarning, match=warning_match), pytest.raises( - ValueError, match="Conflicting `region_name` provided" - ): - ti.render_templates() - class TestEcsCreateClusterOperator(EcsBaseTestCase): @pytest.mark.parametrize("waiter_delay, waiter_max_attempts", WAITERS_TEST_CASES) @@ -883,8 +834,6 @@ def test_template_fields(self): class TestEcsDeregisterTaskDefinitionOperator(EcsBaseTestCase): - warn_message = "'wait_for_completion' and waiter related params have no effect" - def test_execute_immediate_delete(self): """Test if task definition deleted during initial request.""" op = EcsDeregisterTaskDefinitionOperator(task_id="task", task_definition=TASK_DEFINITION_NAME) @@ -896,47 +845,6 @@ def test_execute_immediate_delete(self): mock_client_method.assert_called_once_with(taskDefinition=TASK_DEFINITION_NAME) assert result == "foo-bar" - def test_deprecation(self): - with pytest.warns(AirflowProviderDeprecationWarning, match=self.warn_message): - EcsDeregisterTaskDefinitionOperator(task_id="id", task_definition="def", wait_for_completion=True) - - @pytest.mark.db_test - @pytest.mark.parametrize( - "wait_for_completion, waiter_delay, waiter_max_attempts", - [ - pytest.param(True, 10, 42, id="all-params"), - pytest.param(False, None, None, id="wait-for-completion-only"), - pytest.param(None, 10, None, id="waiter-delay-only"), - pytest.param(None, None, 42, id="waiter-max-attempts-delay-only"), - ], - ) - def test_partial_deprecation_waiters_params( - self, wait_for_completion, waiter_delay, waiter_max_attempts, dag_maker, session - ): - op_kwargs = {} - if wait_for_completion is not None: - op_kwargs["wait_for_completion"] = wait_for_completion - if waiter_delay is not None: - op_kwargs["waiter_delay"] = waiter_delay - if waiter_max_attempts is not None: - op_kwargs["waiter_max_attempts"] = waiter_max_attempts - - with dag_maker(dag_id="test_partial_deprecation_waiters_params_dereg_ecs", session=session): - EcsDeregisterTaskDefinitionOperator.partial( - task_id="fake-task-id", - **op_kwargs, - ).expand(task_definition=["foo", "bar"]) - - dr = dag_maker.create_dagrun() - tis = dr.get_task_instances(session=session) - with set_current_task_instance_session(session=session): - for ti in tis: - with pytest.warns(AirflowProviderDeprecationWarning, match=self.warn_message): - ti.render_templates() - assert not hasattr(ti.task, "wait_for_completion") - assert not hasattr(ti.task, "waiter_delay") - assert not hasattr(ti.task, "waiter_max_attempts") - def test_template_fields(self): op = EcsDeregisterTaskDefinitionOperator(task_id="task", task_definition=TASK_DEFINITION_NAME) @@ -944,8 +852,6 @@ def test_template_fields(self): class TestEcsRegisterTaskDefinitionOperator(EcsBaseTestCase): - warn_message = "'wait_for_completion' and waiter related params have no effect" - def test_execute_immediate_create(self): """Test if task definition created during initial request.""" mock_ti = mock.MagicMock(name="MockedTaskInstance") @@ -976,50 +882,6 @@ def test_execute_immediate_create(self): mock_ti.xcom_push.assert_called_once_with(key="task_definition_arn", value="foo-bar") assert result == "foo-bar" - def test_deprecation(self): - with pytest.warns(AirflowProviderDeprecationWarning, match=self.warn_message): - EcsRegisterTaskDefinitionOperator( - task_id="id", wait_for_completion=True, **TASK_DEFINITION_CONFIG - ) - - @pytest.mark.db_test - @pytest.mark.parametrize( - "wait_for_completion, waiter_delay, waiter_max_attempts", - [ - pytest.param(True, 10, 42, id="all-params"), - pytest.param(False, None, None, id="wait-for-completion-only"), - pytest.param(None, 10, None, id="waiter-delay-only"), - pytest.param(None, None, 42, id="waiter-max-attempts-delay-only"), - ], - ) - def test_partial_deprecation_waiters_params( - self, wait_for_completion, waiter_delay, waiter_max_attempts, dag_maker, session - ): - op_kwargs = {} - if wait_for_completion is not None: - op_kwargs["wait_for_completion"] = wait_for_completion - if waiter_delay is not None: - op_kwargs["waiter_delay"] = waiter_delay - if waiter_max_attempts is not None: - op_kwargs["waiter_max_attempts"] = waiter_max_attempts - - with dag_maker(dag_id="test_partial_deprecation_waiters_params_reg_ecs", session=session): - EcsRegisterTaskDefinitionOperator.partial( - task_id="fake-task-id", - family="family_name", - **op_kwargs, - ).expand(container_definitions=[{}, {}]) - - dr = dag_maker.create_dagrun() - tis = dr.get_task_instances(session=session) - with set_current_task_instance_session(session=session): - for ti in tis: - with pytest.warns(AirflowProviderDeprecationWarning, match=self.warn_message): - ti.render_templates() - assert not hasattr(ti.task, "wait_for_completion") - assert not hasattr(ti.task, "waiter_delay") - assert not hasattr(ti.task, "waiter_max_attempts") - def test_template_fields(self): op = EcsRegisterTaskDefinitionOperator(task_id="task", **TASK_DEFINITION_CONFIG) diff --git a/tests/providers/amazon/aws/operators/test_eks.py b/tests/providers/amazon/aws/operators/test_eks.py index 399c8e40823ae..2daa484626879 100644 --- a/tests/providers/amazon/aws/operators/test_eks.py +++ b/tests/providers/amazon/aws/operators/test_eks.py @@ -23,7 +23,7 @@ import pytest from botocore.waiter import Waiter -from airflow.exceptions import AirflowProviderDeprecationWarning, TaskDeferred +from airflow.exceptions import TaskDeferred from airflow.providers.amazon.aws.hooks.eks import ClusterStates, EksHook from airflow.providers.amazon.aws.operators.eks import ( EksCreateClusterOperator, @@ -717,84 +717,41 @@ def test_existing_nodegroup( assert mock_generate_config_file.return_value.__enter__.return_value == op.config_file @pytest.mark.parametrize( - "compatible_kpo, kwargs, expected_attributes, warning, warning_message", + "compatible_kpo, kwargs, expected_attributes", [ ( True, {"on_finish_action": "delete_succeeded_pod"}, {"on_finish_action": OnFinishAction.DELETE_SUCCEEDED_POD}, - False, - None, - ), - ( - # test that priority for deprecated param - True, - {"on_finish_action": "keep_pod", "is_delete_operator_pod": True}, - {"on_finish_action": OnFinishAction.DELETE_POD, "is_delete_operator_pod": True}, - True, - "`is_delete_operator_pod` parameter is deprecated, please use `on_finish_action`", ), ( # test default True, {}, - {"on_finish_action": OnFinishAction.KEEP_POD, "is_delete_operator_pod": False}, - True, - "You have not set parameter `on_finish_action` in class EksPodOperator. Currently the default for this parameter is `keep_pod` but in a future release the default will be changed to `delete_pod`. To ensure pods are not deleted in the future you will need to set `on_finish_action=keep_pod` explicitly.", - ), - ( - False, - {"is_delete_operator_pod": True}, - {"is_delete_operator_pod": True}, - True, - "`is_delete_operator_pod` parameter is deprecated, please use `on_finish_action`", - ), - ( - False, - {"is_delete_operator_pod": False}, - {"is_delete_operator_pod": False}, - True, - "`is_delete_operator_pod` parameter is deprecated, please use `on_finish_action`", + {"on_finish_action": OnFinishAction.DELETE_POD}, ), ( # test default False, {}, - {"is_delete_operator_pod": False}, - True, - "You have not set parameter `on_finish_action` in class EksPodOperator. Currently the default for this parameter is `keep_pod` but in a future release the default will be changed to `delete_pod`. To ensure pods are not deleted in the future you will need to set `on_finish_action=keep_pod` explicitly.", + {}, ), ], ) - def test_on_finish_action_handler( - self, compatible_kpo, kwargs, expected_attributes, warning, warning_message - ): + def test_on_finish_action_handler(self, compatible_kpo, kwargs, expected_attributes): kpo_init_args_mock = mock.MagicMock(**{"parameters": ["on_finish_action"] if compatible_kpo else []}) with mock.patch("inspect.signature", return_value=kpo_init_args_mock): - if warning: - with pytest.warns(AirflowProviderDeprecationWarning, match=warning_message): - op = EksPodOperator( - task_id="run_pod", - pod_name="run_pod", - cluster_name=CLUSTER_NAME, - image="amazon/aws-cli:latest", - cmds=["sh", "-c", "ls"], - labels={"demo": "hello_world"}, - get_logs=True, - **kwargs, - ) - else: - op = EksPodOperator( - task_id="run_pod", - pod_name="run_pod", - cluster_name=CLUSTER_NAME, - image="amazon/aws-cli:latest", - cmds=["sh", "-c", "ls"], - labels={"demo": "hello_world"}, - get_logs=True, - **kwargs, - ) + op = EksPodOperator( + task_id="run_pod", + pod_name="run_pod", + cluster_name=CLUSTER_NAME, + image="amazon/aws-cli:latest", + cmds=["sh", "-c", "ls"], + labels={"demo": "hello_world"}, + get_logs=True, + **kwargs, + ) for expected_attr in expected_attributes: assert op.__getattribute__(expected_attr) == expected_attributes[expected_attr] diff --git a/tests/providers/amazon/aws/operators/test_emr_serverless.py b/tests/providers/amazon/aws/operators/test_emr_serverless.py index e7a43cf079f0b..c84d1032bcf71 100644 --- a/tests/providers/amazon/aws/operators/test_emr_serverless.py +++ b/tests/providers/amazon/aws/operators/test_emr_serverless.py @@ -23,7 +23,7 @@ import pytest from botocore.exceptions import WaiterError -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning, TaskDeferred +from airflow.exceptions import AirflowException, TaskDeferred from airflow.providers.amazon.aws.hooks.emr import EmrServerlessHook from airflow.providers.amazon.aws.operators.emr import ( EmrServerlessCreateApplicationOperator, @@ -327,51 +327,27 @@ def test_application_in_failure_state(self, mock_conn, mock_get_waiter): ) @pytest.mark.parametrize( - "waiter_delay, waiter_max_attempts, waiter_countdown, waiter_check_interval_seconds, expected, warning", + "waiter_delay, waiter_max_attempts, expected", [ - (NOTSET, NOTSET, NOTSET, NOTSET, [60, 25], False), - (30, 10, NOTSET, NOTSET, [30, 10], False), - (NOTSET, NOTSET, 30 * 15, 15, [15, 30], True), - (10, 20, 30, 40, [10, 20], True), + (NOTSET, NOTSET, [60, 25]), + (30, 10, [30, 10]), ], ) def test_create_application_waiter_params( self, waiter_delay, waiter_max_attempts, - waiter_countdown, - waiter_check_interval_seconds, expected, - warning, ): - if warning: - with pytest.warns( - AirflowProviderDeprecationWarning, - match="The parameter waiter_.* has been deprecated to standardize naming conventions. Please use waiter_.* instead. .*In the future this will default to None and defer to the waiter's default value.", - ): - operator = EmrServerlessCreateApplicationOperator( - task_id=task_id, - release_label=release_label, - job_type=job_type, - client_request_token=client_request_token, - config=config, - waiter_delay=waiter_delay, - waiter_max_attempts=waiter_max_attempts, - waiter_countdown=waiter_countdown, - waiter_check_interval_seconds=waiter_check_interval_seconds, - ) - else: - operator = EmrServerlessCreateApplicationOperator( - task_id=task_id, - release_label=release_label, - job_type=job_type, - client_request_token=client_request_token, - config=config, - waiter_delay=waiter_delay, - waiter_max_attempts=waiter_max_attempts, - waiter_countdown=waiter_countdown, - waiter_check_interval_seconds=waiter_check_interval_seconds, - ) + operator = EmrServerlessCreateApplicationOperator( + task_id=task_id, + release_label=release_label, + job_type=job_type, + client_request_token=client_request_token, + config=config, + waiter_delay=waiter_delay, + waiter_max_attempts=waiter_max_attempts, + ) assert operator.wait_for_completion is True assert operator.waiter_delay == expected[0] assert operator.waiter_max_attempts == expected[1] @@ -788,51 +764,27 @@ def test_cancel_job_run(self, mock_conn): ) @pytest.mark.parametrize( - "waiter_delay, waiter_max_attempts, waiter_countdown, waiter_check_interval_seconds, expected, warning", + "waiter_delay, waiter_max_attempts, expected", [ - (NOTSET, NOTSET, NOTSET, NOTSET, [60, 25], False), - (30, 10, NOTSET, NOTSET, [30, 10], False), - (NOTSET, NOTSET, 30 * 15, 15, [15, 30], True), - (10, 20, 30, 40, [10, 20], True), + (NOTSET, NOTSET, [60, 25]), + (30, 10, [30, 10]), ], ) def test_start_job_waiter_params( self, waiter_delay, waiter_max_attempts, - waiter_countdown, - waiter_check_interval_seconds, expected, - warning, ): - if warning: - with pytest.warns( - AirflowProviderDeprecationWarning, - match="The parameter waiter_.* has been deprecated to standardize naming conventions. Please use waiter_.* instead. .*In the future this will default to None and defer to the waiter's default value.", - ): - operator = EmrServerlessStartJobOperator( - task_id=task_id, - application_id=application_id, - execution_role_arn=execution_role_arn, - job_driver=job_driver, - configuration_overrides=configuration_overrides, - waiter_delay=waiter_delay, - waiter_max_attempts=waiter_max_attempts, - waiter_countdown=waiter_countdown, - waiter_check_interval_seconds=waiter_check_interval_seconds, - ) - else: - operator = EmrServerlessStartJobOperator( - task_id=task_id, - application_id=application_id, - execution_role_arn=execution_role_arn, - job_driver=job_driver, - configuration_overrides=configuration_overrides, - waiter_delay=waiter_delay, - waiter_max_attempts=waiter_max_attempts, - waiter_countdown=waiter_countdown, - waiter_check_interval_seconds=waiter_check_interval_seconds, - ) + operator = EmrServerlessStartJobOperator( + task_id=task_id, + application_id=application_id, + execution_role_arn=execution_role_arn, + job_driver=job_driver, + configuration_overrides=configuration_overrides, + waiter_delay=waiter_delay, + waiter_max_attempts=waiter_max_attempts, + ) assert operator.wait_for_completion is True assert operator.waiter_delay == expected[0] assert operator.waiter_max_attempts == expected[1] @@ -1260,45 +1212,24 @@ def test_delete_application_failed_deletion(self, mock_conn, mock_get_waiter): mock_conn.delete_application.assert_called_once_with(applicationId=application_id_delete_operator) @pytest.mark.parametrize( - "waiter_delay, waiter_max_attempts, waiter_countdown, waiter_check_interval_seconds, expected, warning", + "waiter_delay, waiter_max_attempts, expected", [ - (NOTSET, NOTSET, NOTSET, NOTSET, [60, 25], False), - (30, 10, NOTSET, NOTSET, [30, 10], False), - (NOTSET, NOTSET, 30 * 15, 15, [15, 30], True), - (10, 20, 30, 40, [10, 20], True), + (NOTSET, NOTSET, [60, 25]), + (30, 10, [30, 10]), ], ) def test_delete_application_waiter_params( self, waiter_delay, waiter_max_attempts, - waiter_countdown, - waiter_check_interval_seconds, expected, - warning, ): - if warning: - with pytest.warns( - AirflowProviderDeprecationWarning, - match="The parameter waiter_.* has been deprecated to standardize naming conventions. Please use waiter_.* instead. .*In the future this will default to None and defer to the waiter's default value.", - ): - operator = EmrServerlessDeleteApplicationOperator( - task_id=task_id, - application_id=application_id, - waiter_delay=waiter_delay, - waiter_max_attempts=waiter_max_attempts, - waiter_countdown=waiter_countdown, - waiter_check_interval_seconds=waiter_check_interval_seconds, - ) - else: - operator = EmrServerlessDeleteApplicationOperator( - task_id=task_id, - application_id=application_id, - waiter_delay=waiter_delay, - waiter_max_attempts=waiter_max_attempts, - waiter_countdown=waiter_countdown, - waiter_check_interval_seconds=waiter_check_interval_seconds, - ) + operator = EmrServerlessDeleteApplicationOperator( + task_id=task_id, + application_id=application_id, + waiter_delay=waiter_delay, + waiter_max_attempts=waiter_max_attempts, + ) assert operator.wait_for_completion is True assert operator.waiter_delay == expected[0] assert operator.waiter_max_attempts == expected[1] diff --git a/tests/providers/amazon/aws/operators/test_glue_databrew.py b/tests/providers/amazon/aws/operators/test_glue_databrew.py index a18c6ddd4a41a..698b206acfb1c 100644 --- a/tests/providers/amazon/aws/operators/test_glue_databrew.py +++ b/tests/providers/amazon/aws/operators/test_glue_databrew.py @@ -23,7 +23,6 @@ import pytest from moto import mock_aws -from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.providers.amazon.aws.hooks.glue_databrew import GlueDataBrewHook from airflow.providers.amazon.aws.operators.glue_databrew import GlueDataBrewStartJobOperator from tests.providers.amazon.aws.utils.test_template_fields import validate_template_fields @@ -84,25 +83,6 @@ def test_start_job_no_wait(self, mock_hook_get_waiter, mock_conn): operator.execute(None) mock_hook_get_waiter.assert_not_called() - @mock.patch.object(GlueDataBrewHook, "conn") - @mock.patch.object(GlueDataBrewHook, "get_waiter") - def test_start_job_with_deprecation_parameters(self, mock_hook_get_waiter, mock_conn): - TEST_RUN_ID = "12345" - - with pytest.warns(AirflowProviderDeprecationWarning): - operator = GlueDataBrewStartJobOperator( - task_id="task_test", - job_name=JOB_NAME, - wait_for_completion=False, - aws_conn_id="aws_default", - delay=15, - ) - - mock_conn.start_job_run(mock.MagicMock(), return_value=TEST_RUN_ID) - assert operator.waiter_delay == 15 - operator.execute(None) - mock_hook_get_waiter.assert_not_called() - def test_template_fields(self): operator = GlueDataBrewStartJobOperator(task_id="fake_task_id", job_name=JOB_NAME) validate_template_fields(operator) diff --git a/tests/providers/amazon/aws/operators/test_redshift_data.py b/tests/providers/amazon/aws/operators/test_redshift_data.py index c22d776a94b44..c4972e4c42e7e 100644 --- a/tests/providers/amazon/aws/operators/test_redshift_data.py +++ b/tests/providers/amazon/aws/operators/test_redshift_data.py @@ -21,7 +21,7 @@ import pytest -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning, TaskDeferred +from airflow.exceptions import AirflowException, TaskDeferred from airflow.providers.amazon.aws.hooks.redshift_data import QueryExecutionOutput from airflow.providers.amazon.aws.operators.redshift_data import RedshiftDataOperator from airflow.providers.amazon.aws.triggers.redshift_data import RedshiftDataTrigger @@ -72,9 +72,6 @@ def test_init(self): verify="/spam/egg.pem", botocore_config={"read_timeout": 42}, ) - with pytest.warns(AirflowProviderDeprecationWarning): - # Check deprecated region argument - assert op.region == "eu-central-1" assert op.hook.client_type == "redshift-data" assert op.hook.resource_type is None assert op.hook.aws_conn_id == "fake-conn-id" diff --git a/tests/providers/amazon/aws/secrets/test_secrets_manager.py b/tests/providers/amazon/aws/secrets/test_secrets_manager.py index dcf9c6f138d20..fa824e3d7d474 100644 --- a/tests/providers/amazon/aws/secrets/test_secrets_manager.py +++ b/tests/providers/amazon/aws/secrets/test_secrets_manager.py @@ -16,13 +16,11 @@ # under the License. from __future__ import annotations -import json from unittest import mock import pytest from moto import mock_aws -from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.providers.amazon.aws.secrets.secrets_manager import SecretsManagerBackend @@ -47,118 +45,6 @@ def test_get_conn_value_full_url_mode(self): returned_uri = secrets_manager_backend.get_conn_value(conn_id="test_postgres") assert "postgresql://airflow:airflow@host:5432/airflow" == returned_uri - @pytest.mark.parametrize( - "are_secret_values_urlencoded, login, host", - [ - (True, "is url encoded", "not%20idempotent"), - (False, "is%20url%20encoded", "not%2520idempotent"), - ], - ) - @mock_aws - def test_get_connection_broken_field_mode_url_encoding(self, are_secret_values_urlencoded, login, host): - secret_id = "airflow/connections/test_postgres" - create_param = { - "Name": secret_id, - "SecretString": json.dumps( - { - "conn_type": "postgresql", - "login": "is%20url%20encoded", - "password": "not url encoded", - "host": "not%2520idempotent", - "extra": json.dumps({"foo": "bar"}), - } - ), - } - - with pytest.warns( - AirflowProviderDeprecationWarning, - match=r"The `secret_values_are_urlencoded` is deprecated. This kwarg only exists to assist in migrating away from URL-encoding secret values for JSON secrets. To remove this warning, make sure your JSON secrets are \*NOT\* URL-encoded, and then remove this kwarg from backend_kwargs.", - ): - secrets_manager_backend = SecretsManagerBackend( - are_secret_values_urlencoded=are_secret_values_urlencoded - ) - secrets_manager_backend.client.create_secret(**create_param) - - conn = secrets_manager_backend.get_connection(conn_id="test_postgres") - - assert conn.login == login - assert conn.password == "not url encoded" - assert conn.host == host - assert conn.conn_id == "test_postgres" - assert conn.extra_dejson["foo"] == "bar" - - @mock_aws - def test_get_connection_broken_field_mode_extra_allows_nested_json(self): - secret_id = "airflow/connections/test_postgres" - create_param = { - "Name": secret_id, - "SecretString": json.dumps( - { - "conn_type": "postgresql", - "user": "airflow", - "password": "airflow", - "host": "airflow", - "extra": {"foo": "bar"}, - } - ), - } - - with pytest.warns( - AirflowProviderDeprecationWarning, - match="The `full_url_mode` kwarg is deprecated. Going forward, the `SecretsManagerBackend` will support both URL-encoded and JSON-encoded secrets at the same time. The encoding of the secret will be determined automatically.", - ): - secrets_manager_backend = SecretsManagerBackend(full_url_mode=False) - secrets_manager_backend.client.create_secret(**create_param) - - conn = secrets_manager_backend.get_connection(conn_id="test_postgres") - assert conn.extra_dejson["foo"] == "bar" - - @mock_aws - def test_get_conn_value_broken_field_mode(self): - secret_id = "airflow/connections/test_postgres" - create_param = { - "Name": secret_id, - "SecretString": ( - '{"user": "airflow", "pass": "airflow", "host": "host", ' - '"port": 5432, "schema": "airflow", "engine": "postgresql"}' - ), - } - - with pytest.warns( - AirflowProviderDeprecationWarning, - match="The `full_url_mode` kwarg is deprecated. Going forward, the `SecretsManagerBackend` will support both URL-encoded and JSON-encoded secrets at the same time. The encoding of the secret will be determined automatically.", - ): - secrets_manager_backend = SecretsManagerBackend(full_url_mode=False) - secrets_manager_backend.client.create_secret(**create_param) - - conn = secrets_manager_backend.get_connection(conn_id="test_postgres") - returned_uri = conn.get_uri() - assert "postgres://airflow:airflow@host:5432/airflow" == returned_uri - - @mock_aws - def test_get_conn_value_broken_field_mode_extra_words_added(self): - secret_id = "airflow/connections/test_postgres" - create_param = { - "Name": secret_id, - "SecretString": ( - '{"usuario": "airflow", "pass": "airflow", "host": "host", ' - '"port": 5432, "schema": "airflow", "engine": "postgresql"}' - ), - } - - with pytest.warns( - AirflowProviderDeprecationWarning, - match="The `full_url_mode` kwarg is deprecated. Going forward, the `SecretsManagerBackend` will support both URL-encoded and JSON-encoded secrets at the same time. The encoding of the secret will be determined automatically.", - ): - secrets_manager_backend = SecretsManagerBackend( - full_url_mode=False, extra_conn_words={"user": ["usuario"]} - ) - secrets_manager_backend.client.create_secret(**create_param) - - conn = secrets_manager_backend.get_connection(conn_id="test_postgres") - returned_uri = conn.get_uri() - assert "postgres://airflow:airflow@host:5432/airflow" == returned_uri - @mock_aws def test_get_conn_value_non_existent_key(self): """ diff --git a/tests/providers/amazon/aws/sensors/test_base_aws.py b/tests/providers/amazon/aws/sensors/test_base_aws.py index a10960585f38a..f099f5fcfe04f 100644 --- a/tests/providers/amazon/aws/sensors/test_base_aws.py +++ b/tests/providers/amazon/aws/sensors/test_base_aws.py @@ -20,7 +20,6 @@ import pytest -from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.hooks.base import BaseHook from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseHook from airflow.providers.amazon.aws.sensors.base_aws import AwsBaseSensor @@ -123,41 +122,6 @@ def test_execute(self, dag_maker, op_kwargs): tis = {ti.task_id: ti for ti in dagrun.task_instances} tis["fake-task-id"].run() - @pytest.mark.skip_if_database_isolation_mode - @pytest.mark.parametrize( - "region, region_name", - [ - pytest.param("eu-west-1", None, id="region-only"), - pytest.param("us-east-1", "us-east-1", id="non-ambiguous-params"), - ], - ) - def test_deprecated_region_name(self, region, region_name): - warning_match = r"`region` is deprecated and will be removed" - with pytest.warns(AirflowProviderDeprecationWarning, match=warning_match): - op = FakeDynamoDBSensor( - task_id="fake-task-id", - aws_conn_id=TEST_CONN, - region=region, - region_name=region_name, - ) - assert op.region_name == region - - with pytest.warns(AirflowProviderDeprecationWarning, match=warning_match): - assert op.region == region - - def test_conflicting_region_name(self): - error_match = r"Conflicting `region_name` provided, region_name='us-west-1', region='eu-west-1'" - with pytest.raises(ValueError, match=error_match), pytest.warns( - AirflowProviderDeprecationWarning, - match="`region` is deprecated and will be removed in the future. Please use `region_name` instead.", - ): - FakeDynamoDBSensor( - task_id="fake-task-id", - aws_conn_id=TEST_CONN, - region="eu-west-1", - region_name="us-west-1", - ) - def test_no_aws_hook_class_attr(self): class NoAwsHookClassSensor(AwsBaseSensor): ... @@ -180,45 +144,3 @@ class SoWrongSensor(AwsBaseSensor): error_match = r"Class attribute 'SoWrongSensor.aws_hook_class' is not a subclass of AwsGenericHook" with pytest.raises(AttributeError, match=error_match): SoWrongSensor(task_id="fake-task-id") - - @pytest.mark.skip_if_database_isolation_mode - @pytest.mark.parametrize( - "region, region_name, expected_region_name", - [ - pytest.param("ca-west-1", None, "ca-west-1", id="region-only"), - pytest.param("us-west-1", "us-west-1", "us-west-1", id="non-ambiguous-params"), - ], - ) - @pytest.mark.db_test - def test_region_in_partial_sensor(self, region, region_name, expected_region_name, dag_maker): - with dag_maker("test_region_in_partial_sensor", serialized=True): - FakeDynamoDBSensor.partial( - task_id="fake-task-id", - region=region, - region_name=region_name, - ).expand(value=[1, 2, 3]) - - dr = dag_maker.create_dagrun(execution_date=timezone.utcnow()) - warning_match = r"`region` is deprecated and will be removed" - for ti in dr.task_instances: - with pytest.warns(AirflowProviderDeprecationWarning, match=warning_match): - ti.run() - assert ti.task.region_name == expected_region_name - - @pytest.mark.skip_if_database_isolation_mode - @pytest.mark.db_test - def test_ambiguous_region_in_partial_sensor(self, dag_maker): - with dag_maker("test_ambiguous_region_in_partial_sensor", serialized=True): - FakeDynamoDBSensor.partial( - task_id="fake-task-id", - region="eu-west-1", - region_name="us-east-1", - ).expand(value=[1, 2, 3]) - - dr = dag_maker.create_dagrun(execution_date=timezone.utcnow()) - warning_match = r"`region` is deprecated and will be removed" - for ti in dr.task_instances: - with pytest.warns(AirflowProviderDeprecationWarning, match=warning_match), pytest.raises( - ValueError, match="Conflicting `region_name` provided" - ): - ti.run() diff --git a/tests/providers/amazon/aws/sensors/test_quicksight.py b/tests/providers/amazon/aws/sensors/test_quicksight.py index 46890a69cbfff..3e61d31239274 100644 --- a/tests/providers/amazon/aws/sensors/test_quicksight.py +++ b/tests/providers/amazon/aws/sensors/test_quicksight.py @@ -21,7 +21,7 @@ import pytest -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.providers.amazon.aws.hooks.quicksight import QuickSightHook from airflow.providers.amazon.aws.sensors.quicksight import QuickSightSensor @@ -95,16 +95,3 @@ def test_poke_terminated_status(self, status, mocked_get_status, mocked_get_erro QuickSightSensor(**self.default_op_kwargs).poke({}) mocked_get_status.assert_called_once_with(None, DATA_SET_ID, INGESTION_ID) mocked_get_error_info.assert_called_once_with(None, DATA_SET_ID, INGESTION_ID) - - def test_deprecated_properties(self): - sensor = QuickSightSensor(**self.default_op_kwargs) - with pytest.warns(AirflowProviderDeprecationWarning, match="please use `.*hook` property instead"): - assert sensor.quicksight_hook is sensor.hook - - with mock.patch("airflow.providers.amazon.aws.hooks.sts.StsHook") as mocked_class, pytest.warns( - AirflowProviderDeprecationWarning, match=r"consider to use `.*hook\.account_id` instead" - ): - mocked_sts_hook = mock.MagicMock(name="FakeStsHook") - mocked_class.return_value = mocked_sts_hook - assert sensor.sts_hook is mocked_sts_hook - mocked_class.assert_called_once_with(aws_conn_id=None) diff --git a/tests/providers/amazon/aws/transfers/test_base.py b/tests/providers/amazon/aws/transfers/test_base.py index b5144f4a7f64c..a60fdeba06244 100644 --- a/tests/providers/amazon/aws/transfers/test_base.py +++ b/tests/providers/amazon/aws/transfers/test_base.py @@ -20,7 +20,6 @@ import pytest from airflow import DAG -from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.models import DagRun, TaskInstance from airflow.providers.amazon.aws.transfers.base import AwsToAwsBaseOperator from airflow.utils import timezone @@ -54,15 +53,3 @@ def test_render_template(self, session, clean_dags_and_dagruns): ti.render_templates() assert "2020-01-01" == getattr(operator, "source_aws_conn_id") assert "2020-01-01" == getattr(operator, "dest_aws_conn_id") - - def test_deprecation(self): - with pytest.warns( - AirflowProviderDeprecationWarning, - match="The aws_conn_id parameter has been deprecated." - " Use the source_aws_conn_id parameter instead.", - ): - AwsToAwsBaseOperator( - task_id="transfer", - dag=self.dag, - aws_conn_id="my_conn", - ) diff --git a/tests/providers/amazon/aws/transfers/test_dynamodb_to_s3.py b/tests/providers/amazon/aws/transfers/test_dynamodb_to_s3.py index d608464b900c8..bcd0a82632d8b 100644 --- a/tests/providers/amazon/aws/transfers/test_dynamodb_to_s3.py +++ b/tests/providers/amazon/aws/transfers/test_dynamodb_to_s3.py @@ -25,9 +25,7 @@ import pytest from airflow import DAG -from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.models import DagRun, TaskInstance -from airflow.providers.amazon.aws.transfers.base import _DEPRECATION_MSG from airflow.providers.amazon.aws.transfers.dynamodb_to_s3 import ( DynamoDBToS3Operator, JSONEncoder, @@ -150,41 +148,6 @@ def test_dynamodb_to_s3_default_connection(self, mock_aws_dynamodb_hook, mock_s3 mock_s3_hook.assert_called_with(aws_conn_id=aws_conn_id) mock_aws_dynamodb_hook.assert_called_with(aws_conn_id=aws_conn_id) - @patch("airflow.providers.amazon.aws.transfers.dynamodb_to_s3.S3Hook") - @patch("airflow.providers.amazon.aws.transfers.dynamodb_to_s3.DynamoDBHook") - def test_dynamodb_to_s3_with_aws_conn_id(self, mock_aws_dynamodb_hook, mock_s3_hook): - responses = [ - { - "Items": [{"a": 1}, {"b": 2}], - "LastEvaluatedKey": "123", - }, - { - "Items": [{"c": 3}], - }, - ] - table = MagicMock() - table.return_value.scan.side_effect = responses - mock_aws_dynamodb_hook.return_value.get_conn.return_value.Table = table - - s3_client = MagicMock() - s3_client.return_value.upload_file = self.mock_upload_file - mock_s3_hook.return_value.get_conn = s3_client - - aws_conn_id = "test-conn-id" - with pytest.warns(AirflowProviderDeprecationWarning, match=_DEPRECATION_MSG): - dynamodb_to_s3_operator = DynamoDBToS3Operator( - task_id="dynamodb_to_s3", - dynamodb_table_name="airflow_rocks", - s3_bucket_name="airflow-bucket", - file_size=4000, - aws_conn_id=aws_conn_id, - ) - - dynamodb_to_s3_operator.execute(context={}) - - mock_s3_hook.assert_called_with(aws_conn_id=aws_conn_id) - mock_aws_dynamodb_hook.assert_called_with(aws_conn_id=aws_conn_id) - @patch("airflow.providers.amazon.aws.transfers.dynamodb_to_s3.S3Hook") @patch("airflow.providers.amazon.aws.transfers.dynamodb_to_s3.DynamoDBHook") def test_dynamodb_to_s3_with_different_aws_conn_id(self, mock_aws_dynamodb_hook, mock_s3_hook): diff --git a/tests/providers/amazon/aws/transfers/test_gcs_to_s3.py b/tests/providers/amazon/aws/transfers/test_gcs_to_s3.py index 7f7802ca23931..4eb4dfa260554 100644 --- a/tests/providers/amazon/aws/transfers/test_gcs_to_s3.py +++ b/tests/providers/amazon/aws/transfers/test_gcs_to_s3.py @@ -70,7 +70,6 @@ def test_execute__match_glob(self, mock_hook): operator.execute(None) mock_hook.return_value.list.assert_called_once_with( bucket_name=GCS_BUCKET, - delimiter=None, match_glob=f"**/*{DELIMITER}", prefix=PREFIX, user_project=None, @@ -83,16 +82,14 @@ def test_execute_incremental(self, mock_hook): gcs_provide_file = mock_hook.return_value.provide_file gcs_provide_file.return_value.__enter__.return_value.name = f.name - with pytest.deprecated_call(match=deprecated_call_match): - operator = GCSToS3Operator( - task_id=TASK_ID, - gcs_bucket=GCS_BUCKET, - prefix=PREFIX, - delimiter=DELIMITER, - dest_aws_conn_id="aws_default", - dest_s3_key=S3_BUCKET, - replace=False, - ) + operator = GCSToS3Operator( + task_id=TASK_ID, + gcs_bucket=GCS_BUCKET, + prefix=PREFIX, + dest_aws_conn_id="aws_default", + dest_s3_key=S3_BUCKET, + replace=False, + ) hook, bucket = _create_test_bucket() bucket.put_object(Key=MOCK_FILES[0], Body=b"testing") @@ -112,16 +109,14 @@ def test_execute_without_replace(self, mock_hook): gcs_provide_file = mock_hook.return_value.provide_file gcs_provide_file.return_value.__enter__.return_value.name = f.name - with pytest.deprecated_call(match=deprecated_call_match): - operator = GCSToS3Operator( - task_id=TASK_ID, - gcs_bucket=GCS_BUCKET, - prefix=PREFIX, - delimiter=DELIMITER, - dest_aws_conn_id="aws_default", - dest_s3_key=S3_BUCKET, - replace=False, - ) + operator = GCSToS3Operator( + task_id=TASK_ID, + gcs_bucket=GCS_BUCKET, + prefix=PREFIX, + dest_aws_conn_id="aws_default", + dest_s3_key=S3_BUCKET, + replace=False, + ) hook, bucket = _create_test_bucket() for mock_file in MOCK_FILES: bucket.put_object(Key=mock_file, Body=b"testing") @@ -150,16 +145,14 @@ def test_execute_without_replace_with_folder_structure(self, mock_hook, dest_s3_ gcs_provide_file = mock_hook.return_value.provide_file gcs_provide_file.return_value.__enter__.return_value.name = f.name - with pytest.deprecated_call(match=deprecated_call_match): - operator = GCSToS3Operator( - task_id=TASK_ID, - gcs_bucket=GCS_BUCKET, - prefix=PREFIX, - delimiter=DELIMITER, - dest_aws_conn_id="aws_default", - dest_s3_key=dest_s3_url, - replace=False, - ) + operator = GCSToS3Operator( + task_id=TASK_ID, + gcs_bucket=GCS_BUCKET, + prefix=PREFIX, + dest_aws_conn_id="aws_default", + dest_s3_key=dest_s3_url, + replace=False, + ) # we expect nothing to be uploaded # and all the MOCK_FILES to be present at the S3 bucket @@ -178,57 +171,21 @@ def test_execute(self, mock_hook): gcs_provide_file = mock_hook.return_value.provide_file gcs_provide_file.return_value.__enter__.return_value.name = f.name - with pytest.deprecated_call(match=deprecated_call_match): - operator = GCSToS3Operator( - task_id=TASK_ID, - gcs_bucket=GCS_BUCKET, - prefix=PREFIX, - delimiter=DELIMITER, - dest_aws_conn_id="aws_default", - dest_s3_key=S3_BUCKET, - replace=False, - ) - hook, _ = _create_test_bucket() - - # we expect all MOCK_FILES to be uploaded - # and all MOCK_FILES to be present at the S3 bucket - uploaded_files = operator.execute(None) - assert sorted(MOCK_FILES) == sorted(uploaded_files) - assert sorted(MOCK_FILES) == sorted(hook.list_keys("bucket", delimiter="/")) - - @mock.patch("airflow.providers.amazon.aws.transfers.gcs_to_s3.GCSHook") - def test_execute_gcs_bucket_rename_compatibility(self, mock_hook): - """ - Tests the same conditions as `test_execute` using the deprecated `bucket` parameter instead of - `gcs_bucket`. This test can be removed when the `bucket` parameter is removed. - """ - mock_hook.return_value.list.return_value = MOCK_FILES - with NamedTemporaryFile() as f: - gcs_provide_file = mock_hook.return_value.provide_file - gcs_provide_file.return_value.__enter__.return_value.name = f.name - bucket_param_deprecated_message = ( - "The ``bucket`` parameter is deprecated and will be removed in a future version. " - "Please use ``gcs_bucket`` instead." + operator = GCSToS3Operator( + task_id=TASK_ID, + gcs_bucket=GCS_BUCKET, + prefix=PREFIX, + dest_aws_conn_id="aws_default", + dest_s3_key=S3_BUCKET, + replace=False, ) - with pytest.deprecated_call(match=bucket_param_deprecated_message): - operator = GCSToS3Operator( - task_id=TASK_ID, - bucket=GCS_BUCKET, - prefix=PREFIX, - match_glob=DELIMITER, - dest_aws_conn_id="aws_default", - dest_s3_key=S3_BUCKET, - replace=False, - ) hook, _ = _create_test_bucket() + # we expect all MOCK_FILES to be uploaded # and all MOCK_FILES to be present at the S3 bucket uploaded_files = operator.execute(None) assert sorted(MOCK_FILES) == sorted(uploaded_files) assert sorted(MOCK_FILES) == sorted(hook.list_keys("bucket", delimiter="/")) - with pytest.raises(ValueError) as excinfo: - GCSToS3Operator(task_id=TASK_ID, dest_s3_key=S3_BUCKET) - assert str(excinfo.value) == "You must pass either ``bucket`` or ``gcs_bucket``." @mock.patch("airflow.providers.amazon.aws.transfers.gcs_to_s3.GCSHook") def test_execute_with_replace(self, mock_hook): @@ -237,16 +194,14 @@ def test_execute_with_replace(self, mock_hook): gcs_provide_file = mock_hook.return_value.provide_file gcs_provide_file.return_value.__enter__.return_value.name = f.name - with pytest.deprecated_call(match=deprecated_call_match): - operator = GCSToS3Operator( - task_id=TASK_ID, - gcs_bucket=GCS_BUCKET, - prefix=PREFIX, - delimiter=DELIMITER, - dest_aws_conn_id="aws_default", - dest_s3_key=S3_BUCKET, - replace=True, - ) + operator = GCSToS3Operator( + task_id=TASK_ID, + gcs_bucket=GCS_BUCKET, + prefix=PREFIX, + dest_aws_conn_id="aws_default", + dest_s3_key=S3_BUCKET, + replace=True, + ) hook, bucket = _create_test_bucket() for mock_file in MOCK_FILES: bucket.put_object(Key=mock_file, Body=b"testing") @@ -264,16 +219,14 @@ def test_execute_incremental_with_replace(self, mock_hook): gcs_provide_file = mock_hook.return_value.provide_file gcs_provide_file.return_value.__enter__.return_value.name = f.name - with pytest.deprecated_call(match=deprecated_call_match): - operator = GCSToS3Operator( - task_id=TASK_ID, - gcs_bucket=GCS_BUCKET, - prefix=PREFIX, - delimiter=DELIMITER, - dest_aws_conn_id="aws_default", - dest_s3_key=S3_BUCKET, - replace=True, - ) + operator = GCSToS3Operator( + task_id=TASK_ID, + gcs_bucket=GCS_BUCKET, + prefix=PREFIX, + dest_aws_conn_id="aws_default", + dest_s3_key=S3_BUCKET, + replace=True, + ) hook, bucket = _create_test_bucket() for mock_file in MOCK_FILES[:2]: bucket.put_object(Key=mock_file, Body=b"testing") @@ -292,16 +245,14 @@ def test_execute_should_handle_with_default_dest_s3_extra_args(self, s3_mock_hoo s3_mock_hook.return_value = mock.Mock() s3_mock_hook.parse_s3_url.return_value = mock.Mock() - with pytest.deprecated_call(match=deprecated_call_match): - operator = GCSToS3Operator( - task_id=TASK_ID, - gcs_bucket=GCS_BUCKET, - prefix=PREFIX, - delimiter=DELIMITER, - dest_aws_conn_id="aws_default", - dest_s3_key=S3_BUCKET, - replace=True, - ) + operator = GCSToS3Operator( + task_id=TASK_ID, + gcs_bucket=GCS_BUCKET, + prefix=PREFIX, + dest_aws_conn_id="aws_default", + dest_s3_key=S3_BUCKET, + replace=True, + ) operator.execute(None) s3_mock_hook.assert_called_once_with(aws_conn_id="aws_default", extra_args={}, verify=None) @@ -315,19 +266,17 @@ def test_execute_should_pass_dest_s3_extra_args_to_s3_hook(self, s3_mock_hook, m s3_mock_hook.return_value = mock.Mock() s3_mock_hook.parse_s3_url.return_value = mock.Mock() - with pytest.deprecated_call(match=deprecated_call_match): - operator = GCSToS3Operator( - task_id=TASK_ID, - gcs_bucket=GCS_BUCKET, - prefix=PREFIX, - delimiter=DELIMITER, - dest_aws_conn_id="aws_default", - dest_s3_key=S3_BUCKET, - replace=True, - dest_s3_extra_args={ - "ContentLanguage": "value", - }, - ) + operator = GCSToS3Operator( + task_id=TASK_ID, + gcs_bucket=GCS_BUCKET, + prefix=PREFIX, + dest_aws_conn_id="aws_default", + dest_s3_key=S3_BUCKET, + replace=True, + dest_s3_extra_args={ + "ContentLanguage": "value", + }, + ) operator.execute(None) s3_mock_hook.assert_called_once_with( aws_conn_id="aws_default", extra_args={"ContentLanguage": "value"}, verify=None @@ -341,17 +290,15 @@ def test_execute_with_s3_acl_policy(self, mock_load_file, mock_gcs_hook): gcs_provide_file = mock_gcs_hook.return_value.provide_file gcs_provide_file.return_value.__enter__.return_value.name = f.name - with pytest.deprecated_call(match=deprecated_call_match): - operator = GCSToS3Operator( - task_id=TASK_ID, - gcs_bucket=GCS_BUCKET, - prefix=PREFIX, - delimiter=DELIMITER, - dest_aws_conn_id="aws_default", - dest_s3_key=S3_BUCKET, - replace=False, - s3_acl_policy=S3_ACL_POLICY, - ) + operator = GCSToS3Operator( + task_id=TASK_ID, + gcs_bucket=GCS_BUCKET, + prefix=PREFIX, + dest_aws_conn_id="aws_default", + dest_s3_key=S3_BUCKET, + replace=False, + s3_acl_policy=S3_ACL_POLICY, + ) _create_test_bucket() operator.execute(None) @@ -366,17 +313,15 @@ def test_execute_without_keep_director_structure(self, mock_hook): gcs_provide_file = mock_hook.return_value.provide_file gcs_provide_file.return_value.__enter__.return_value.name = f.name - with pytest.deprecated_call(match=deprecated_call_match): - operator = GCSToS3Operator( - task_id=TASK_ID, - gcs_bucket=GCS_BUCKET, - prefix=PREFIX, - delimiter=DELIMITER, - dest_aws_conn_id="aws_default", - dest_s3_key=S3_BUCKET, - replace=False, - keep_directory_structure=False, - ) + operator = GCSToS3Operator( + task_id=TASK_ID, + gcs_bucket=GCS_BUCKET, + prefix=PREFIX, + dest_aws_conn_id="aws_default", + dest_s3_key=S3_BUCKET, + replace=False, + keep_directory_structure=False, + ) hook, _ = _create_test_bucket() # we expect all except first file in MOCK_FILES to be uploaded diff --git a/tests/providers/amazon/aws/triggers/test_emr.py b/tests/providers/amazon/aws/triggers/test_emr.py index 3469ee4c13a7d..eb7f1851155ac 100644 --- a/tests/providers/amazon/aws/triggers/test_emr.py +++ b/tests/providers/amazon/aws/triggers/test_emr.py @@ -58,34 +58,6 @@ def test_serialization(self): class TestEmrCreateJobFlowTrigger: - def test_init_with_deprecated_params(self): - import warnings - - with warnings.catch_warnings(record=True) as catch_warns: - warnings.simplefilter("always") - - job_flow_id = "test_job_flow_id" - poll_interval = 10 - max_attempts = 5 - aws_conn_id = "aws_default" - waiter_delay = 30 - waiter_max_attempts = 60 - - trigger = EmrCreateJobFlowTrigger( - job_flow_id=job_flow_id, - poll_interval=poll_interval, - max_attempts=max_attempts, - aws_conn_id=aws_conn_id, - waiter_delay=waiter_delay, - waiter_max_attempts=waiter_max_attempts, - ) - - assert trigger.waiter_delay == poll_interval - assert len(catch_warns) == 1 - assert issubclass(catch_warns[-1].category, DeprecationWarning) - assert "please use waiter_delay instead of poll_interval" in str(catch_warns[-1].message) - assert "and waiter_max_attempts instead of max_attempts" in str(catch_warns[-1].message) - def test_serialization(self): job_flow_id = "test_job_flow_id" waiter_delay = 30 @@ -109,34 +81,6 @@ def test_serialization(self): class TestEmrTerminateJobFlowTrigger: - def test_init_with_deprecated_params(self): - import warnings - - with warnings.catch_warnings(record=True) as catch_warns: - warnings.simplefilter("always") - - job_flow_id = "test_job_flow_id" - poll_interval = 10 - max_attempts = 5 - aws_conn_id = "aws_default" - waiter_delay = 30 - waiter_max_attempts = 60 - - trigger = EmrTerminateJobFlowTrigger( - job_flow_id=job_flow_id, - poll_interval=poll_interval, - max_attempts=max_attempts, - aws_conn_id=aws_conn_id, - waiter_delay=waiter_delay, - waiter_max_attempts=waiter_max_attempts, - ) - - assert trigger.waiter_delay == poll_interval # Assert deprecated parameter is correctly used - assert len(catch_warns) == 1 - assert issubclass(catch_warns[-1].category, DeprecationWarning) - assert "please use waiter_delay instead of poll_interval" in str(catch_warns[-1].message) - assert "and waiter_max_attempts instead of max_attempts" in str(catch_warns[-1].message) - def test_serialization(self): job_flow_id = "test_job_flow_id" waiter_delay = 30 @@ -160,33 +104,6 @@ def test_serialization(self): class TestEmrContainerTrigger: - def test_init_with_deprecated_params(self): - import warnings - - with warnings.catch_warnings(record=True) as catch_warns: - warnings.simplefilter("always") - - virtual_cluster_id = "test_virtual_cluster_id" - job_id = "test_job_id" - aws_conn_id = "aws_default" - poll_interval = 10 - waiter_delay = 30 - waiter_max_attempts = 600 - - trigger = EmrContainerTrigger( - virtual_cluster_id=virtual_cluster_id, - job_id=job_id, - aws_conn_id=aws_conn_id, - poll_interval=poll_interval, - waiter_delay=waiter_delay, - waiter_max_attempts=waiter_max_attempts, - ) - - assert trigger.waiter_delay == poll_interval # Assert deprecated parameter is correctly used - assert len(catch_warns) == 1 - assert issubclass(catch_warns[-1].category, DeprecationWarning) - assert "please use waiter_delay instead of poll_interval" in str(catch_warns[-1].message) - def test_serialization(self): virtual_cluster_id = "test_virtual_cluster_id" job_id = "test_job_id" diff --git a/tests/providers/amazon/aws/triggers/test_glue_crawler.py b/tests/providers/amazon/aws/triggers/test_glue_crawler.py index dd46b1882a79d..8975aa1aff505 100644 --- a/tests/providers/amazon/aws/triggers/test_glue_crawler.py +++ b/tests/providers/amazon/aws/triggers/test_glue_crawler.py @@ -17,7 +17,7 @@ from __future__ import annotations from unittest import mock -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock import pytest @@ -28,23 +28,17 @@ class TestGlueCrawlerCompleteTrigger: - @patch("airflow.providers.amazon.aws.triggers.glue_crawler.warnings.warn") - def test_serialization(self, mock_warn): + def test_serialization(self): crawler_name = "test_crawler" poll_interval = 10 aws_conn_id = "aws_default" trigger = GlueCrawlerCompleteTrigger( crawler_name=crawler_name, - poll_interval=poll_interval, + waiter_delay=poll_interval, aws_conn_id=aws_conn_id, ) - assert mock_warn.call_count == 1 - args, kwargs = mock_warn.call_args - assert args[0] == "please use waiter_delay instead of poll_interval." - assert kwargs == {"stacklevel": 2} - classpath, kwargs = trigger.serialize() assert classpath == "airflow.providers.amazon.aws.triggers.glue_crawler.GlueCrawlerCompleteTrigger" assert kwargs == { diff --git a/tests/providers/amazon/aws/triggers/test_glue_databrew.py b/tests/providers/amazon/aws/triggers/test_glue_databrew.py index c39892c247f4e..09137a0c7d6ac 100644 --- a/tests/providers/amazon/aws/triggers/test_glue_databrew.py +++ b/tests/providers/amazon/aws/triggers/test_glue_databrew.py @@ -18,7 +18,6 @@ import pytest -from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.providers.amazon.aws.triggers.glue_databrew import GlueDataBrewJobCompleteTrigger TEST_JOB_NAME = "test_job_name" @@ -45,24 +44,3 @@ def test_serialize(self, trigger): assert class_path == class_path2 assert args == args2 - - def test_serialize_with_deprecated_parameters(self, trigger): - with pytest.warns(AirflowProviderDeprecationWarning): - class_path, args = GlueDataBrewJobCompleteTrigger( - aws_conn_id="aws_default", - job_name=TEST_JOB_NAME, - run_id=TEST_JOB_RUN_ID, - delay=1, - max_attempts=1, - ).serialize() - - class_name = class_path.split(".")[-1] - clazz = globals()[class_name] - instance = clazz(**args) - - class_path2, args2 = instance.serialize() - - assert class_path == class_path2 - assert args == args2 - assert args.get("waiter_delay") == 1 - assert args.get("waiter_max_attempts") == 1 diff --git a/tests/providers/amazon/aws/triggers/test_redshift_cluster.py b/tests/providers/amazon/aws/triggers/test_redshift_cluster.py index 5d5cc2c4241ab..75055fe5ad6fe 100644 --- a/tests/providers/amazon/aws/triggers/test_redshift_cluster.py +++ b/tests/providers/amazon/aws/triggers/test_redshift_cluster.py @@ -52,12 +52,12 @@ def test_redshift_cluster_sensor_trigger_serialization(self): } @pytest.mark.asyncio - @mock.patch("airflow.providers.amazon.aws.hooks.redshift_cluster.RedshiftAsyncHook.cluster_status") + @mock.patch("airflow.providers.amazon.aws.hooks.redshift_cluster.RedshiftHook.cluster_status_async") async def test_redshift_cluster_sensor_trigger_success(self, mock_cluster_status): """ Test RedshiftClusterTrigger with the success status """ - expected_result = {"status": "success", "cluster_state": "available"} + expected_result = "available" mock_cluster_status.return_value = expected_result trigger = RedshiftClusterTrigger( @@ -69,16 +69,14 @@ async def test_redshift_cluster_sensor_trigger_success(self, mock_cluster_status generator = trigger.run() actual = await generator.asend(None) - assert TriggerEvent(expected_result) == actual + assert TriggerEvent({"status": "success", "message": "target state met"}) == actual @pytest.mark.asyncio @pytest.mark.parametrize( "expected_result", - [ - ({"status": "success", "cluster_state": "Resuming"}), - ], + ["Resuming"], ) - @mock.patch("airflow.providers.amazon.aws.hooks.redshift_cluster.RedshiftAsyncHook.cluster_status") + @mock.patch("airflow.providers.amazon.aws.hooks.redshift_cluster.RedshiftHook.cluster_status_async") async def test_redshift_cluster_sensor_trigger_resuming_status( self, mock_cluster_status, expected_result ): @@ -100,7 +98,7 @@ async def test_redshift_cluster_sensor_trigger_resuming_status( asyncio.get_event_loop().stop() @pytest.mark.asyncio - @mock.patch("airflow.providers.amazon.aws.hooks.redshift_cluster.RedshiftAsyncHook.cluster_status") + @mock.patch("airflow.providers.amazon.aws.hooks.redshift_cluster.RedshiftHook.cluster_status_async") async def test_redshift_cluster_sensor_trigger_exception(self, mock_cluster_status): """Test RedshiftClusterTrigger with exception""" mock_cluster_status.side_effect = Exception("Test exception") diff --git a/tests/providers/amazon/aws/triggers/test_serialization.py b/tests/providers/amazon/aws/triggers/test_serialization.py index 39000f3ab6fbb..d5b6a0a74b4cc 100644 --- a/tests/providers/amazon/aws/triggers/test_serialization.py +++ b/tests/providers/amazon/aws/triggers/test_serialization.py @@ -204,20 +204,20 @@ class TestTriggersSerialization: EmrCreateJobFlowTrigger( job_flow_id=TEST_JOB_FLOW_ID, aws_conn_id=AWS_CONN_ID, - poll_interval=WAITER_DELAY, - max_attempts=MAX_ATTEMPTS, + waiter_delay=WAITER_DELAY, + waiter_max_attempts=MAX_ATTEMPTS, ), EmrTerminateJobFlowTrigger( job_flow_id=TEST_JOB_FLOW_ID, aws_conn_id=AWS_CONN_ID, - poll_interval=WAITER_DELAY, - max_attempts=MAX_ATTEMPTS, + waiter_delay=WAITER_DELAY, + waiter_max_attempts=MAX_ATTEMPTS, ), EmrContainerTrigger( virtual_cluster_id=VIRTUAL_CLUSTER_ID, job_id=JOB_ID, aws_conn_id=AWS_CONN_ID, - poll_interval=WAITER_DELAY, + waiter_delay=WAITER_DELAY, ), EmrStepSensorTrigger( job_flow_id=TEST_JOB_FLOW_ID, @@ -278,32 +278,32 @@ class TestTriggersSerialization: ), RedshiftCreateClusterTrigger( cluster_identifier=TEST_CLUSTER_IDENTIFIER, - poll_interval=WAITER_DELAY, - max_attempt=MAX_ATTEMPTS, + waiter_delay=WAITER_DELAY, + waiter_max_attempts=MAX_ATTEMPTS, aws_conn_id=AWS_CONN_ID, ), RedshiftPauseClusterTrigger( cluster_identifier=TEST_CLUSTER_IDENTIFIER, - poll_interval=WAITER_DELAY, - max_attempts=MAX_ATTEMPTS, + waiter_delay=WAITER_DELAY, + waiter_max_attempts=MAX_ATTEMPTS, aws_conn_id=AWS_CONN_ID, ), RedshiftCreateClusterSnapshotTrigger( cluster_identifier=TEST_CLUSTER_IDENTIFIER, - poll_interval=WAITER_DELAY, - max_attempts=MAX_ATTEMPTS, + waiter_delay=WAITER_DELAY, + waiter_max_attempts=MAX_ATTEMPTS, aws_conn_id=AWS_CONN_ID, ), RedshiftResumeClusterTrigger( cluster_identifier=TEST_CLUSTER_IDENTIFIER, - poll_interval=WAITER_DELAY, - max_attempts=MAX_ATTEMPTS, + waiter_delay=WAITER_DELAY, + waiter_max_attempts=MAX_ATTEMPTS, aws_conn_id=AWS_CONN_ID, ), RedshiftDeleteClusterTrigger( cluster_identifier=TEST_CLUSTER_IDENTIFIER, - poll_interval=WAITER_DELAY, - max_attempts=MAX_ATTEMPTS, + waiter_delay=WAITER_DELAY, + waiter_max_attempts=MAX_ATTEMPTS, aws_conn_id=AWS_CONN_ID, ), RdsDbAvailableTrigger( diff --git a/tests/providers/amazon/aws/utils/test_connection_wrapper.py b/tests/providers/amazon/aws/utils/test_connection_wrapper.py index f08b7477a010f..95aff74d3b5f7 100644 --- a/tests/providers/amazon/aws/utils/test_connection_wrapper.py +++ b/tests/providers/amazon/aws/utils/test_connection_wrapper.py @@ -25,7 +25,7 @@ from botocore import UNSIGNED from botocore.config import Config -from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning +from airflow.exceptions import AirflowException from airflow.models import Connection from airflow.providers.amazon.aws.utils.connection_wrapper import AwsConnectionWrapper, _ConnectionMetadata @@ -122,16 +122,6 @@ def test_unexpected_aws_connection_type(self, conn_type): wrap_conn = AwsConnectionWrapper(conn=mock_connection_factory(conn_type=conn_type)) assert wrap_conn.conn_type == conn_type - @pytest.mark.parametrize("conn_type", ["s3", "S3"]) - def test_deprecated_s3_connection_type(self, conn_type): - warning_message = ( - r".* has connection type 's3', which has been replaced by connection type 'aws'\. " - r"Please update your connection to have `conn_type='aws'`." - ) - with pytest.warns(AirflowProviderDeprecationWarning, match=warning_message): - wrap_conn = AwsConnectionWrapper(conn=mock_connection_factory(conn_type=conn_type)) - assert wrap_conn.conn_type == conn_type - @pytest.mark.parametrize("aws_session_token", [None, "mock-aws-session-token"]) @pytest.mark.parametrize("aws_secret_access_key", ["mock-aws-secret-access-key"]) @pytest.mark.parametrize("aws_access_key_id", ["mock-aws-access-key-id"]) @@ -164,53 +154,6 @@ def test_get_credentials_from_extra(self, aws_access_key_id, aws_secret_access_k assert wrap_conn.aws_secret_access_key == aws_secret_access_key assert wrap_conn.aws_session_token == aws_session_token - @pytest.mark.parametrize("aws_access_key_id", ["mock-aws-access-key-id"]) - @pytest.mark.parametrize("aws_secret_access_key", ["mock-aws-secret-access-key"]) - @pytest.mark.parametrize("aws_session_token", [None, "mock-aws-session-token"]) - def test_get_credentials_from_session_kwargs( - self, aws_access_key_id, aws_secret_access_key, aws_session_token - ): - mock_conn_extra = { - "session_kwargs": { - "aws_access_key_id": aws_access_key_id, - "aws_secret_access_key": aws_secret_access_key, - "aws_session_token": aws_session_token, - }, - } - mock_conn = mock_connection_factory(login=None, password=None, extra=mock_conn_extra) - - with pytest.warns( - AirflowProviderDeprecationWarning, match=r"'session_kwargs' in extra config is deprecated" - ): - wrap_conn = AwsConnectionWrapper(conn=mock_conn) - assert wrap_conn.aws_access_key_id == aws_access_key_id - assert wrap_conn.aws_secret_access_key == aws_secret_access_key - assert wrap_conn.aws_session_token == aws_session_token - - # This function never tested and mark as deprecated. Only test expected output - @mock.patch("airflow.providers.amazon.aws.utils.connection_wrapper._parse_s3_config") - @pytest.mark.parametrize("aws_session_token", [None, "mock-aws-session-token"]) - @pytest.mark.parametrize("aws_secret_access_key", ["mock-aws-secret-access-key"]) - @pytest.mark.parametrize("aws_access_key_id", ["mock-aws-access-key-id"]) - def test_get_credentials_from_s3_config( - self, mock_parse_s3_config, aws_access_key_id, aws_secret_access_key, aws_session_token - ): - mock_parse_s3_config.return_value = (aws_access_key_id, aws_secret_access_key) - mock_conn_extra = { - "s3_config_format": "aws", - "profile": "test", - "s3_config_file": "aws-credentials", - } - if aws_session_token: - mock_conn_extra["aws_session_token"] = aws_session_token - mock_conn = mock_connection_factory(login=None, password=None, extra=mock_conn_extra) - - wrap_conn = AwsConnectionWrapper(conn=mock_conn) - mock_parse_s3_config.assert_called_once_with("aws-credentials", "aws", "test") - assert wrap_conn.aws_access_key_id == aws_access_key_id - assert wrap_conn.aws_secret_access_key == aws_secret_access_key - assert wrap_conn.aws_session_token == aws_session_token - @pytest.mark.parametrize("aws_access_key_id", [None, "mock-aws-access-key-id"]) @pytest.mark.parametrize("aws_secret_access_key", [None, "mock-aws-secret-access-key"]) @pytest.mark.parametrize("aws_session_token", [None, "mock-aws-session-token"]) @@ -249,42 +192,6 @@ def test_get_session_kwargs_from_wrapper( assert wrap_conn.session_kwargs == expected assert wrap_conn.session_kwargs != session_kwargs - @pytest.mark.parametrize("aws_access_key_id", [None, "mock-aws-access-key-id"]) - @pytest.mark.parametrize("aws_secret_access_key", [None, "mock-aws-secret-access-key"]) - @pytest.mark.parametrize("aws_session_token", [None, "mock-aws-session-token"]) - @pytest.mark.parametrize("profile_name", [None, "mock-profile"]) - @pytest.mark.parametrize("region_name", [None, "mock-region-name"]) - def test_get_session_kwargs_deprecation( - self, aws_access_key_id, aws_secret_access_key, aws_session_token, profile_name, region_name - ): - mock_conn_extra_session_kwargs = { - "aws_access_key_id": aws_access_key_id, - "aws_secret_access_key": aws_secret_access_key, - "aws_session_token": aws_session_token, - "profile_name": profile_name, - "region_name": region_name, - } - mock_conn = mock_connection_factory(extra={"session_kwargs": mock_conn_extra_session_kwargs}) - expected = {} - if aws_access_key_id and aws_secret_access_key: - expected["aws_access_key_id"] = aws_access_key_id - expected["aws_secret_access_key"] = aws_secret_access_key - if aws_session_token: - expected["aws_session_token"] = aws_session_token - if profile_name: - expected["profile_name"] = profile_name - if region_name: - expected["region_name"] = region_name - - warning_message = ( - r"'session_kwargs' in extra config is deprecated and will be removed in a future releases. " - r"Please specify arguments passed to boto3 Session directly in .* extra." - ) - with pytest.warns(AirflowProviderDeprecationWarning, match=warning_message): - wrap_conn = AwsConnectionWrapper(conn=mock_conn) - session_kwargs = wrap_conn.session_kwargs - assert session_kwargs == expected - @pytest.mark.parametrize( "region_name,conn_region_name", [ @@ -340,30 +247,6 @@ def test_get_botocore_config(self, mock_botocore_config, botocore_config, botoco botocore_config_kwargs["signature_version"] = UNSIGNED assert mock.call(**botocore_config_kwargs) in mock_botocore_config.mock_calls - @pytest.mark.parametrize( - "extra, expected", - [ - ({"host": "https://host.aws"}, "https://host.aws"), - ({"endpoint_url": "https://endpoint.aws"}, "https://endpoint.aws"), - ({"host": "https://host.aws", "endpoint_url": "https://endpoint.aws"}, "https://host.aws"), - ], - ids=["'host' is used", "'endpoint_url' is used", "'host' preferred over 'endpoint_url'"], - ) - def test_get_endpoint_url_from_extra(self, extra, expected): - mock_conn = mock_connection_factory(extra=extra) - expected_deprecation_message = ( - r"extra\['host'\] is deprecated and will be removed in a future release." - r" Please set extra\['endpoint_url'\] instead" - ) - - if extra.get("host"): - with pytest.warns(AirflowProviderDeprecationWarning, match=expected_deprecation_message): - wrap_conn = AwsConnectionWrapper(conn=mock_conn) - else: - wrap_conn = AwsConnectionWrapper(conn=mock_conn) - - assert wrap_conn.endpoint_url == expected - @pytest.mark.parametrize("aws_account_id, aws_iam_role", [(None, None), ("111111111111", "another-role")]) def test_get_role_arn(self, aws_account_id, aws_iam_role): mock_conn = mock_connection_factory( @@ -376,24 +259,6 @@ def test_get_role_arn(self, aws_account_id, aws_iam_role): wrap_conn = AwsConnectionWrapper(conn=mock_conn) assert wrap_conn.role_arn == MOCK_ROLE_ARN - @pytest.mark.parametrize( - "aws_account_id, aws_iam_role, expected", - [ - ("222222222222", "mock-role", "arn:aws:iam::222222222222:role/mock-role"), - ("333333333333", "role-path/mock-role", "arn:aws:iam::333333333333:role/role-path/mock-role"), - ], - ) - def test_constructing_role_arn(self, aws_account_id, aws_iam_role, expected): - mock_conn = mock_connection_factory( - extra={ - "aws_account_id": aws_account_id, - "aws_iam_role": aws_iam_role, - } - ) - with pytest.warns(AirflowProviderDeprecationWarning, match="Please set 'role_arn' in .* extra"): - wrap_conn = AwsConnectionWrapper(conn=mock_conn) - assert wrap_conn.role_arn == expected - def test_empty_role_arn(self): wrap_conn = AwsConnectionWrapper(conn=mock_connection_factory()) assert wrap_conn.role_arn is None @@ -453,17 +318,6 @@ def test_get_assume_role_kwargs_external_id_in_kwargs(self, external_id_in_extra assert wrap_conn.assume_role_kwargs["ExternalId"] == mock_external_id_in_kwargs assert wrap_conn.assume_role_kwargs["ExternalId"] != external_id_in_extra - def test_get_assume_role_kwargs_external_id_in_extra(self): - mock_external_id_in_extra = "mock-external-id-in-extra" - mock_conn_extra = {"role_arn": MOCK_ROLE_ARN, "external_id": mock_external_id_in_extra} - mock_conn = mock_connection_factory(extra=mock_conn_extra) - - warning_message = "Please set 'ExternalId' in 'assume_role_kwargs' in .* extra." - with pytest.warns(AirflowProviderDeprecationWarning, match=warning_message): - wrap_conn = AwsConnectionWrapper(conn=mock_conn) - assert "ExternalId" in wrap_conn.assume_role_kwargs - assert wrap_conn.assume_role_kwargs["ExternalId"] == mock_external_id_in_extra - @pytest.mark.parametrize( "orig_wrapper", [ @@ -509,17 +363,6 @@ def test_wrap_wrapper(self, orig_wrapper, region_name, botocore_config): assert wrap_conn.region_name == (region_name or orig_wrapper.region_name) assert wrap_conn.botocore_config == (botocore_config or orig_wrapper.botocore_config) - def test_connection_host_raises_deprecation(self): - mock_conn = mock_connection_factory(host="https://aws.com") - expected_deprecation_message = ( - f"Host {mock_conn.host} specified in the connection is not used." - " Please, set it on extra['endpoint_url'] instead" - ) - with pytest.warns(AirflowProviderDeprecationWarning) as record: - AwsConnectionWrapper(conn=mock_conn) - - assert str(record[0].message) == expected_deprecation_message - @pytest.mark.parametrize("conn_id", [None, "mock-conn-id"]) @pytest.mark.parametrize("profile_name", [None, "mock-profile"]) @pytest.mark.parametrize("role_arn", [None, MOCK_ROLE_ARN]) @@ -595,16 +438,6 @@ def test_get_service_endpoint_url_sts( assert wrap_conn.get_service_endpoint_url("sts", sts_connection_assume=True) == expected_endpoint_url assert wrap_conn.get_service_endpoint_url("sts", sts_test_connection=True) == expected_endpoint_url - def test_get_service_endpoint_url_sts_deprecated_test_connection(self): - fake_conn = mock_connection_factory( - conn_id="foo-bar", - extra={"endpoint_url": "https://spam.egg", "test_endpoint_url": "https://foo.bar"}, - ) - wrap_conn = AwsConnectionWrapper(conn=fake_conn) - warning_message = r"extra\['test_endpoint_url'\] is deprecated" - with pytest.warns(AirflowProviderDeprecationWarning, match=warning_message): - assert wrap_conn.get_service_endpoint_url("sts", sts_test_connection=True) == "https://foo.bar" - def test_get_service_endpoint_url_sts_unsupported(self): wrap_conn = AwsConnectionWrapper(conn=mock_connection_factory()) with pytest.raises(AirflowException, match=r"Can't resolve STS endpoint when both"): diff --git a/tests/system/providers/amazon/aws/example_batch.py b/tests/system/providers/amazon/aws/example_batch.py index d47fc683e5e84..0b79bb5a82b29 100644 --- a/tests/system/providers/amazon/aws/example_batch.py +++ b/tests/system/providers/amazon/aws/example_batch.py @@ -207,7 +207,7 @@ def delete_job_queue(job_queue_name): job_name=batch_job_name, job_queue=batch_job_queue_name, job_definition=batch_job_definition_name, - overrides=JOB_OVERRIDES, + ecs_properties_override=JOB_OVERRIDES, ) # [END howto_operator_batch] From 2c928ce10e1557e502b9b8767b096b55a4c42757 Mon Sep 17 00:00:00 2001 From: Daniel Standish <15932138+dstandish@users.noreply.github.com> Date: Thu, 26 Sep 2024 08:53:40 -0700 Subject: [PATCH 042/802] Clarify logic in callback func in is authorized callback (#42475) I think this makes it a little clearer what the logic is doing. --- airflow/api_connexion/security.py | 31 +++++++++++++++++++------------ 1 file changed, 19 insertions(+), 12 deletions(-) diff --git a/airflow/api_connexion/security.py b/airflow/api_connexion/security.py index 7da83a76168bb..445ded913e56a 100644 --- a/airflow/api_connexion/security.py +++ b/airflow/api_connexion/security.py @@ -113,18 +113,25 @@ def requires_access_dag( method: ResourceMethod, access_entity: DagAccessEntity | None = None ) -> Callable[[T], T]: def _is_authorized_callback(dag_id: str): - def callback(): - access = get_auth_manager().is_authorized_dag( - method=method, - access_entity=access_entity, - details=DagDetails(id=dag_id), - ) - - # ``access`` means here: - # - if a DAG id is provided (``dag_id`` not None): is the user authorized to access this DAG - # - if no DAG id is provided: is the user authorized to access all DAGs - if dag_id or access or access_entity: - return access + def callback() -> bool | DagAccessEntity: + if dag_id: + # a DAG id is provided; is the user authorized to access this DAG? + return get_auth_manager().is_authorized_dag( + method=method, + access_entity=access_entity, + details=DagDetails(id=dag_id), + ) + else: + # here we know dag_id is not provided. + # check is the user authorized to access all DAGs? + if get_auth_manager().is_authorized_dag( + method=method, + access_entity=access_entity, + ): + return True + elif access_entity: + # no dag_id provided, and user does not have access to all dags + return False # dag_id is not provided, and the user is not authorized to access *all* DAGs # so we check that the user can access at least *one* dag From 45b9d28d1ffb70394ff525f27a1649403e1a6203 Mon Sep 17 00:00:00 2001 From: Daniel Standish <15932138+dstandish@users.noreply.github.com> Date: Thu, 26 Sep 2024 09:26:13 -0700 Subject: [PATCH 043/802] Add basic endpoints for managing backfill entities (#42455) More logic will be added for `create` and `cancel`. We'll need to create dag runs and fail them accordingly. But I'll add that logic separately to make it easier to scrutinize it more closely. Will also follow up with some changes to the security implementation. --------- Co-authored-by: Jed Cunningham <66968678+jedcunningham@users.noreply.github.com> --- .../endpoints/backfill_endpoint.py | 181 +++++++ airflow/api_connexion/openapi/v1.yaml | 291 ++++++++++++ airflow/models/backfill.py | 7 - airflow/www/static/js/types/api-generated.ts | 240 ++++++++++ .../endpoints/test_backfill_endpoint.py | 440 ++++++++++++++++++ tests/test_utils/db.py | 7 + 6 files changed, 1159 insertions(+), 7 deletions(-) create mode 100644 airflow/api_connexion/endpoints/backfill_endpoint.py create mode 100644 tests/api_connexion/endpoints/test_backfill_endpoint.py diff --git a/airflow/api_connexion/endpoints/backfill_endpoint.py b/airflow/api_connexion/endpoints/backfill_endpoint.py new file mode 100644 index 0000000000000..f974be4d75d82 --- /dev/null +++ b/airflow/api_connexion/endpoints/backfill_endpoint.py @@ -0,0 +1,181 @@ +# 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 logging +from functools import wraps +from typing import TYPE_CHECKING + +import pendulum +from sqlalchemy import select + +from airflow.api_connexion import security +from airflow.api_connexion.exceptions import Conflict, NotFound +from airflow.api_connexion.schemas.backfill_schema import ( + BackfillCollection, + backfill_collection_schema, + backfill_schema, +) +from airflow.models.backfill import Backfill +from airflow.models.serialized_dag import SerializedDagModel +from airflow.utils import timezone +from airflow.utils.session import NEW_SESSION, provide_session +from airflow.www.decorators import action_logging + +if TYPE_CHECKING: + from sqlalchemy.orm import Session + + from airflow.api_connexion.types import APIResponse + +log = logging.getLogger(__name__) + +RESOURCE_EVENT_PREFIX = "dag" + + +def backfill_to_dag(func): + """ + Enrich the request with dag_id. + + :meta private: + """ + + @wraps(func) + def wrapper(*, backfill_id, session, **kwargs): + backfill = session.get(Backfill, backfill_id) + if not backfill: + raise NotFound("Backfill not found") + return func(dag_id=backfill.dag_id, backfill_id=backfill_id, session=session, **kwargs) + + return wrapper + + +@provide_session +def _create_backfill( + *, + dag_id: str, + from_date: str, + to_date: str, + max_active_runs: int, + reverse: bool, + dag_run_conf: dict | None, + session: Session = NEW_SESSION, +) -> Backfill: + serdag = session.get(SerializedDagModel, dag_id) + if not serdag: + raise NotFound(f"Could not find dag {dag_id}") + + br = Backfill( + dag_id=dag_id, + from_date=pendulum.parse(from_date), + to_date=pendulum.parse(to_date), + max_active_runs=max_active_runs, + dag_run_conf=dag_run_conf, + ) + session.add(br) + session.commit() + return br + + +@security.requires_access_dag("GET") +@action_logging +@provide_session +def list_backfills(dag_id, session): + backfills = session.scalars(select(Backfill).where(Backfill.dag_id == dag_id)).all() + obj = BackfillCollection( + backfills=backfills, + total_entries=len(backfills), + ) + return backfill_collection_schema.dump(obj) + + +@provide_session +@backfill_to_dag +@security.requires_access_dag("PUT") +@action_logging +def pause_backfill(*, backfill_id, session, **kwargs): + br = session.get(Backfill, backfill_id) + if br.completed_at: + raise Conflict("Backfill is already completed.") + if br.is_paused is False: + br.is_paused = True + session.commit() + return backfill_schema.dump(br) + + +@provide_session +@backfill_to_dag +@security.requires_access_dag("PUT") +@action_logging +def unpause_backfill(*, backfill_id, session, **kwargs): + br = session.get(Backfill, backfill_id) + if br.completed_at: + raise Conflict("Backfill is already completed.") + if br.is_paused: + br.is_paused = False + session.commit() + return backfill_schema.dump(br) + + +@provide_session +@backfill_to_dag +@security.requires_access_dag("PUT") +@action_logging +def cancel_backfill(*, backfill_id, session, **kwargs): + br: Backfill = session.get(Backfill, backfill_id) + if br.completed_at is not None: + raise Conflict("Backfill is already completed.") + + br.completed_at = timezone.utcnow() + + # first, pause + if not br.is_paused: + br.is_paused = True + session.commit() + return backfill_schema.dump(br) + + +@provide_session +@backfill_to_dag +@security.requires_access_dag("GET") +@action_logging +def get_backfill(*, backfill_id: int, session: Session = NEW_SESSION, **kwargs): + backfill = session.get(Backfill, backfill_id) + if backfill: + return backfill_schema.dump(backfill) + raise NotFound("Backfill not found") + + +@security.requires_access_dag("PUT") +@action_logging +def create_backfill( + dag_id: str, + from_date: str, + to_date: str, + max_active_runs: int = 10, + reverse: bool = False, + dag_run_conf: dict | None = None, +) -> APIResponse: + backfill_obj = _create_backfill( + dag_id=dag_id, + from_date=from_date, + to_date=to_date, + max_active_runs=max_active_runs, + reverse=reverse, + dag_run_conf=dag_run_conf, + ) + return backfill_schema.dump(backfill_obj) diff --git a/airflow/api_connexion/openapi/v1.yaml b/airflow/api_connexion/openapi/v1.yaml index 0c4b0414775f1..15ad6fd8a4f63 100644 --- a/airflow/api_connexion/openapi/v1.yaml +++ b/airflow/api_connexion/openapi/v1.yaml @@ -245,6 +245,200 @@ servers: description: Apache Airflow Stable API. paths: + # Database entities + /backfills: + get: + summary: List backfills + x-openapi-router-controller: airflow.api_connexion.endpoints.backfill_endpoint + operationId: list_backfills + tags: [Backfill] + parameters: + - name: dag_id + in: query + schema: + type: string + required: true + description: | + List backfills for this dag. + responses: + "200": + description: Success. + content: + application/json: + schema: + $ref: "#/components/schemas/BackfillCollection" + "401": + $ref: "#/components/responses/Unauthenticated" + "403": + $ref: "#/components/responses/PermissionDenied" + + post: + summary: Create a backfill job. + x-openapi-router-controller: airflow.api_connexion.endpoints.backfill_endpoint + operationId: create_backfill + tags: [Backfill] + parameters: + - name: dag_id + in: query + schema: + type: string + required: true + description: | + Create dag runs for this dag. + + - name: from_date + in: query + schema: + type: string + format: date-time + required: true + description: | + Create dag runs with logical dates from this date onward, including this date. + + - name: to_date + in: query + schema: + type: string + format: date-time + required: true + description: | + Create dag runs for logical dates up to but not including this date. + + - name: max_active_runs + in: query + schema: + type: integer + required: false + description: | + Maximum number of active DAG runs for the the backfill. + + - name: reverse + in: query + schema: + type: boolean + required: false + description: | + If true, run the dag runs in descending order of logical date. + + - name: config + in: query + schema: + # todo: AIP-78 make this object + type: string + required: false + description: | + If true, run the dag runs in descending order of logical date. + responses: + "200": + description: Success. + content: + application/json: + schema: + $ref: "#/components/schemas/Backfill" + "400": + $ref: "#/components/responses/BadRequest" + "401": + $ref: "#/components/responses/Unauthenticated" + "403": + $ref: "#/components/responses/PermissionDenied" + + /backfills/{backfill_id}: + parameters: + - $ref: "#/components/parameters/BackfillIdPath" + get: + summary: Get a backfill + x-openapi-router-controller: airflow.api_connexion.endpoints.backfill_endpoint + operationId: get_backfill + tags: [Backfill] + responses: + "200": + description: Success. + content: + application/json: + schema: + $ref: "#/components/schemas/Backfill" + "401": + $ref: "#/components/responses/Unauthenticated" + "403": + $ref: "#/components/responses/PermissionDenied" + "404": + $ref: "#/components/responses/NotFound" + + /backfills/{backfill_id}/pause: + parameters: + - $ref: "#/components/parameters/BackfillIdPath" + post: + summary: Pause a backfill + x-openapi-router-controller: airflow.api_connexion.endpoints.backfill_endpoint + operationId: pause_backfill + tags: [Backfill] + responses: + "200": + description: Success. + content: + application/json: + schema: + $ref: "#/components/schemas/Backfill" + "401": + $ref: "#/components/responses/Unauthenticated" + "403": + $ref: "#/components/responses/PermissionDenied" + "404": + $ref: "#/components/responses/NotFound" + "409": + $ref: "#/components/responses/Conflict" + + /backfills/{backfill_id}/unpause: + parameters: + - $ref: "#/components/parameters/BackfillIdPath" + post: + summary: Pause a backfill + x-openapi-router-controller: airflow.api_connexion.endpoints.backfill_endpoint + operationId: unpause_backfill + tags: [Backfill] + responses: + "200": + description: Success. + content: + application/json: + schema: + $ref: "#/components/schemas/Backfill" + "401": + $ref: "#/components/responses/Unauthenticated" + "403": + $ref: "#/components/responses/PermissionDenied" + "404": + $ref: "#/components/responses/NotFound" + "409": + $ref: "#/components/responses/Conflict" + + /backfills/{backfill_id}/cancel: + parameters: + - $ref: "#/components/parameters/BackfillIdPath" + post: + summary: Cancel a backfill + description: | + When a backfill is cancelled, all queued dag runs will be marked as failed. + Running dag runs will be allowed to continue. + x-openapi-router-controller: airflow.api_connexion.endpoints.backfill_endpoint + operationId: cancel_backfill + tags: [Backfill] + responses: + "200": + description: Success. + content: + application/json: + schema: + $ref: "#/components/schemas/Backfill" + "401": + $ref: "#/components/responses/Unauthenticated" + "403": + $ref: "#/components/responses/PermissionDenied" + "404": + $ref: "#/components/responses/NotFound" + "409": + $ref: "#/components/responses/Conflict" + # Database entities /connections: get: @@ -2704,6 +2898,66 @@ components: $ref: "#/components/schemas/UserCollectionItem" - $ref: "#/components/schemas/CollectionInfo" + Backfill: + description: > + Backfill entity object. + + Represents one backfill run / request. + type: object + properties: + id: + type: integer + description: id + dag_id: + type: string + description: The dag_id for the backfill. + from_date: + type: string + nullable: true + description: From date of the backfill (inclusive). + to_date: + type: string + nullable: true + description: To date of the backfill (exclusive). + dag_run_conf: + type: string + nullable: true + description: Dag run conf to be forwarded to the dag runs. + is_paused: + type: boolean + nullable: true + description: is_paused + max_active_runs: + type: integer + nullable: true + description: max_active_runs + created_at: + type: string + nullable: true + description: created_at + completed_at: + type: string + nullable: true + description: completed_at + updated_at: + type: string + nullable: true + description: updated_at + + + BackfillCollection: + type: object + description: | + Collection of backfill entities. + allOf: + - type: object + properties: + backfills: + type: array + items: + $ref: "#/components/schemas/Backfill" + - $ref: "#/components/schemas/CollectionInfo" + ConnectionCollectionItem: description: > Connection collection item. @@ -5125,6 +5379,36 @@ components: # Reusable path, query, header and cookie parameters parameters: + + BackfillIdPath: + in: path + name: backfill_id + schema: + type: integer + required: true + description: | + The integer id identifying the backfill entity. + + FromDate: + in: query + name: from_date + schema: + type: string + format: date-time + required: false + description: | + From date. + + ToDate: + in: query + name: to_date + schema: + type: string + format: date-time + required: false + description: | + To date. + # Pagination parameters PageOffset: in: query @@ -5691,6 +5975,13 @@ components: schema: $ref: "#/components/schemas/Error" # 409 + "Conflict": + description: There is some kind of conflict with the request. + content: + application/json: + schema: + $ref: "#/components/schemas/Error" + # 409 "AlreadyExists": description: An existing resource conflicts with the request. content: diff --git a/airflow/models/backfill.py b/airflow/models/backfill.py index cefe16d863f52..8ff2541353688 100644 --- a/airflow/models/backfill.py +++ b/airflow/models/backfill.py @@ -41,7 +41,6 @@ class Backfill(Base): Controls whether new dag runs will be created for this backfill. Does not pause existing dag runs. - todo: AIP-78 Add test """ max_active_runs = Column(Integer, default=10, nullable=False) created_at = Column(UtcDateTime, default=timezone.utcnow, nullable=False) @@ -49,12 +48,6 @@ class Backfill(Base): updated_at = Column(UtcDateTime, default=timezone.utcnow, onupdate=timezone.utcnow, nullable=False) -# todo: AIP-78 implement clear_failed_tasks?` -# todo: AIP-78 implement clear_dag_run? - -# todo: (AIP-78) should backfill be supported for things with no schedule, or statically partitioned assets? - - class BackfillDagRun(Base): """Mapping table between backfill run and dag run.""" diff --git a/airflow/www/static/js/types/api-generated.ts b/airflow/www/static/js/types/api-generated.ts index 60fd384df00a7..3616be30a1fac 100644 --- a/airflow/www/static/js/types/api-generated.ts +++ b/airflow/www/static/js/types/api-generated.ts @@ -6,6 +6,50 @@ import type { CamelCasedPropertiesDeep } from "type-fest"; */ export interface paths { + "/backfills": { + get: operations["list_backfills"]; + post: operations["create_backfill"]; + }; + "/backfills/{backfill_id}": { + get: operations["get_backfill"]; + parameters: { + path: { + /** The integer id identifying the backfill entity. */ + backfill_id: components["parameters"]["BackfillIdPath"]; + }; + }; + }; + "/backfills/{backfill_id}/pause": { + post: operations["pause_backfill"]; + parameters: { + path: { + /** The integer id identifying the backfill entity. */ + backfill_id: components["parameters"]["BackfillIdPath"]; + }; + }; + }; + "/backfills/{backfill_id}/unpause": { + post: operations["unpause_backfill"]; + parameters: { + path: { + /** The integer id identifying the backfill entity. */ + backfill_id: components["parameters"]["BackfillIdPath"]; + }; + }; + }; + "/backfills/{backfill_id}/cancel": { + /** + * When a backfill is cancelled, all queued dag runs will be marked as failed. + * Running dag runs will be allowed to continue. + */ + post: operations["cancel_backfill"]; + parameters: { + path: { + /** The integer id identifying the backfill entity. */ + backfill_id: components["parameters"]["BackfillIdPath"]; + }; + }; + }; "/connections": { get: operations["get_connections"]; post: operations["post_connection"]; @@ -861,6 +905,36 @@ export interface components { UserCollection: { users?: components["schemas"]["UserCollectionItem"][]; } & components["schemas"]["CollectionInfo"]; + /** + * @description Backfill entity object. + * Represents one backfill run / request. + */ + Backfill: { + /** @description id */ + id?: number; + /** @description The dag_id for the backfill. */ + dag_id?: string; + /** @description From date of the backfill (inclusive). */ + from_date?: string | null; + /** @description To date of the backfill (exclusive). */ + to_date?: string | null; + /** @description Dag run conf to be forwarded to the dag runs. */ + dag_run_conf?: string | null; + /** @description is_paused */ + is_paused?: boolean | null; + /** @description max_active_runs */ + max_active_runs?: number | null; + /** @description created_at */ + created_at?: string | null; + /** @description completed_at */ + completed_at?: string | null; + /** @description updated_at */ + updated_at?: string | null; + }; + /** @description Collection of backfill entities. */ + BackfillCollection: { + backfills?: components["schemas"]["Backfill"][]; + } & components["schemas"]["CollectionInfo"]; /** * @description Connection collection item. * The password and extra fields are only available when retrieving a single object due to the sensitivity of this data. @@ -2405,6 +2479,12 @@ export interface components { "application/json": components["schemas"]["Error"]; }; }; + /** There is some kind of conflict with the request. */ + Conflict: { + content: { + "application/json": components["schemas"]["Error"]; + }; + }; /** An existing resource conflicts with the request. */ AlreadyExists: { content: { @@ -2419,6 +2499,12 @@ export interface components { }; }; parameters: { + /** @description The integer id identifying the backfill entity. */ + BackfillIdPath: number; + /** @description From date. */ + FromDate: string; + /** @description To date. */ + ToDate: string; /** @description The number of items to skip before starting to collect the result set. */ PageOffset: number; /** @description The numbers of items to return. */ @@ -2621,6 +2707,136 @@ export interface components { } export interface operations { + list_backfills: { + parameters: { + query: { + /** List backfills for this dag. */ + dag_id: string; + }; + }; + responses: { + /** Success. */ + 200: { + content: { + "application/json": components["schemas"]["BackfillCollection"]; + }; + }; + 401: components["responses"]["Unauthenticated"]; + 403: components["responses"]["PermissionDenied"]; + }; + }; + create_backfill: { + parameters: { + query: { + /** Create dag runs for this dag. */ + dag_id: string; + /** Create dag runs with logical dates from this date onward, including this date. */ + from_date: string; + /** Create dag runs for logical dates up to but not including this date. */ + to_date: string; + /** Maximum number of active DAG runs for the the backfill. */ + max_active_runs?: number; + /** If true, run the dag runs in descending order of logical date. */ + reverse?: boolean; + /** If true, run the dag runs in descending order of logical date. */ + config?: string; + }; + }; + responses: { + /** Success. */ + 200: { + content: { + "application/json": components["schemas"]["Backfill"]; + }; + }; + 400: components["responses"]["BadRequest"]; + 401: components["responses"]["Unauthenticated"]; + 403: components["responses"]["PermissionDenied"]; + }; + }; + get_backfill: { + parameters: { + path: { + /** The integer id identifying the backfill entity. */ + backfill_id: components["parameters"]["BackfillIdPath"]; + }; + }; + responses: { + /** Success. */ + 200: { + content: { + "application/json": components["schemas"]["Backfill"]; + }; + }; + 401: components["responses"]["Unauthenticated"]; + 403: components["responses"]["PermissionDenied"]; + 404: components["responses"]["NotFound"]; + }; + }; + pause_backfill: { + parameters: { + path: { + /** The integer id identifying the backfill entity. */ + backfill_id: components["parameters"]["BackfillIdPath"]; + }; + }; + responses: { + /** Success. */ + 200: { + content: { + "application/json": components["schemas"]["Backfill"]; + }; + }; + 401: components["responses"]["Unauthenticated"]; + 403: components["responses"]["PermissionDenied"]; + 404: components["responses"]["NotFound"]; + 409: components["responses"]["Conflict"]; + }; + }; + unpause_backfill: { + parameters: { + path: { + /** The integer id identifying the backfill entity. */ + backfill_id: components["parameters"]["BackfillIdPath"]; + }; + }; + responses: { + /** Success. */ + 200: { + content: { + "application/json": components["schemas"]["Backfill"]; + }; + }; + 401: components["responses"]["Unauthenticated"]; + 403: components["responses"]["PermissionDenied"]; + 404: components["responses"]["NotFound"]; + 409: components["responses"]["Conflict"]; + }; + }; + /** + * When a backfill is cancelled, all queued dag runs will be marked as failed. + * Running dag runs will be allowed to continue. + */ + cancel_backfill: { + parameters: { + path: { + /** The integer id identifying the backfill entity. */ + backfill_id: components["parameters"]["BackfillIdPath"]; + }; + }; + responses: { + /** Success. */ + 200: { + content: { + "application/json": components["schemas"]["Backfill"]; + }; + }; + 401: components["responses"]["Unauthenticated"]; + 403: components["responses"]["PermissionDenied"]; + 404: components["responses"]["NotFound"]; + 409: components["responses"]["Conflict"]; + }; + }; get_connections: { parameters: { query: { @@ -5048,6 +5264,12 @@ export type User = CamelCasedPropertiesDeep; export type UserCollection = CamelCasedPropertiesDeep< components["schemas"]["UserCollection"] >; +export type Backfill = CamelCasedPropertiesDeep< + components["schemas"]["Backfill"] +>; +export type BackfillCollection = CamelCasedPropertiesDeep< + components["schemas"]["BackfillCollection"] +>; export type ConnectionCollectionItem = CamelCasedPropertiesDeep< components["schemas"]["ConnectionCollectionItem"] >; @@ -5305,6 +5527,24 @@ export type HealthStatus = CamelCasedPropertiesDeep< export type Operations = operations; /* Types for operation variables */ +export type ListBackfillsVariables = CamelCasedPropertiesDeep< + operations["list_backfills"]["parameters"]["query"] +>; +export type CreateBackfillVariables = CamelCasedPropertiesDeep< + operations["create_backfill"]["parameters"]["query"] +>; +export type GetBackfillVariables = CamelCasedPropertiesDeep< + operations["get_backfill"]["parameters"]["path"] +>; +export type PauseBackfillVariables = CamelCasedPropertiesDeep< + operations["pause_backfill"]["parameters"]["path"] +>; +export type UnpauseBackfillVariables = CamelCasedPropertiesDeep< + operations["unpause_backfill"]["parameters"]["path"] +>; +export type CancelBackfillVariables = CamelCasedPropertiesDeep< + operations["cancel_backfill"]["parameters"]["path"] +>; export type GetConnectionsVariables = CamelCasedPropertiesDeep< operations["get_connections"]["parameters"]["query"] >; diff --git a/tests/api_connexion/endpoints/test_backfill_endpoint.py b/tests/api_connexion/endpoints/test_backfill_endpoint.py new file mode 100644 index 0000000000000..51a4faf40055c --- /dev/null +++ b/tests/api_connexion/endpoints/test_backfill_endpoint.py @@ -0,0 +1,440 @@ +# 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 os +from datetime import datetime +from unittest import mock +from urllib.parse import urlencode + +import pendulum +import pytest + +from airflow.models import DagBag, DagModel +from airflow.models.backfill import Backfill +from airflow.models.dag import DAG +from airflow.models.serialized_dag import SerializedDagModel +from airflow.operators.empty import EmptyOperator +from airflow.security import permissions +from airflow.utils import timezone +from airflow.utils.session import provide_session +from tests.test_utils.api_connexion_utils import create_user, delete_user +from tests.test_utils.db import clear_db_backfills, clear_db_dags, clear_db_runs, clear_db_serialized_dags + +pytestmark = [pytest.mark.db_test] + + +DAG_ID = "test_dag" +TASK_ID = "op1" +DAG2_ID = "test_dag2" +DAG3_ID = "test_dag3" +UTC_JSON_REPR = "UTC" if pendulum.__version__.startswith("3") else "Timezone('UTC')" + + +@pytest.fixture(scope="module") +def configured_app(minimal_app_for_api): + app = minimal_app_for_api + + create_user( + app, # type: ignore + username="test", + role_name="Test", + permissions=[ + (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), + (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG), + (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_DAG), + ], + ) + create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_granular_permissions", role_name="TestGranularDag") # type: ignore + app.appbuilder.sm.sync_perm_for_dag( # type: ignore + "TEST_DAG_1", + access_control={ + "TestGranularDag": { + permissions.RESOURCE_DAG: {permissions.ACTION_CAN_EDIT, permissions.ACTION_CAN_READ} + }, + }, + ) + + with DAG( + DAG_ID, + schedule=None, + start_date=datetime(2020, 6, 15), + doc_md="details", + params={"foo": 1}, + tags=["example"], + ) as dag: + EmptyOperator(task_id=TASK_ID) + + with DAG(DAG2_ID, schedule=None, start_date=datetime(2020, 6, 15)) as dag2: # no doc_md + EmptyOperator(task_id=TASK_ID) + + with DAG(DAG3_ID, schedule=None) as dag3: # DAG start_date set to None + EmptyOperator(task_id=TASK_ID, start_date=datetime(2019, 6, 12)) + + dag_bag = DagBag(os.devnull, include_examples=False) + dag_bag.dags = {dag.dag_id: dag, dag2.dag_id: dag2, dag3.dag_id: dag3} + + app.dag_bag = dag_bag + + yield app + + delete_user(app, username="test") # type: ignore + delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test_granular_permissions") # type: ignore + + +class TestBackfillEndpoint: + @staticmethod + def clean_db(): + clear_db_backfills() + clear_db_runs() + clear_db_dags() + clear_db_serialized_dags() + + @pytest.fixture(autouse=True) + def setup_attrs(self, configured_app) -> None: + self.clean_db() + self.app = configured_app + self.client = self.app.test_client() # type:ignore + self.dag_id = DAG_ID + self.dag2_id = DAG2_ID + self.dag3_id = DAG3_ID + + def teardown_method(self) -> None: + self.clean_db() + + @provide_session + def _create_dag_models(self, *, count=1, dag_id_prefix="TEST_DAG", is_paused=False, session=None): + dags = [] + for num in range(1, count + 1): + dag_model = DagModel( + dag_id=f"{dag_id_prefix}_{num}", + fileloc=f"/tmp/dag_{num}.py", + is_active=True, + timetable_summary="0 0 * * *", + is_paused=is_paused, + ) + session.add(dag_model) + dags.append(dag_model) + return dags + + @provide_session + def _create_deactivated_dag(self, session=None): + dag_model = DagModel( + dag_id="TEST_DAG_DELETED_1", + fileloc="/tmp/dag_del_1.py", + schedule_interval="2 2 * * *", + is_active=False, + ) + session.add(dag_model) + + +class TestListBackfills(TestBackfillEndpoint): + def test_should_respond_200(self, session): + (dag,) = self._create_dag_models() + from_date = timezone.utcnow() + to_date = timezone.utcnow() + b = Backfill(dag_id=dag.dag_id, from_date=from_date, to_date=to_date) + session.add(b) + session.commit() + response = self.client.get( + f"/api/v1/backfills?dag_id={dag.dag_id}", + environ_overrides={"REMOTE_USER": "test"}, + ) + assert response.status_code == 200 + assert response.json == { + "backfills": [ + { + "completed_at": mock.ANY, + "created_at": mock.ANY, + "dag_id": "TEST_DAG_1", + "dag_run_conf": None, + "from_date": from_date.isoformat(), + "id": b.id, + "is_paused": False, + "max_active_runs": 10, + "to_date": to_date.isoformat(), + "updated_at": mock.ANY, + } + ], + "total_entries": 1, + } + + @pytest.mark.parametrize( + "user, expected", + [ + ("test_granular_permissions", 200), + ("test_no_permissions", 403), + ("test", 200), + (None, 401), + ], + ) + def test_should_respond_200_with_granular_dag_access(self, user, expected, session): + (dag,) = self._create_dag_models() + from_date = timezone.utcnow() + to_date = timezone.utcnow() + b = Backfill( + dag_id=dag.dag_id, + from_date=from_date, + to_date=to_date, + ) + + session.add(b) + session.commit() + kwargs = {} + if user: + kwargs.update(environ_overrides={"REMOTE_USER": user}) + response = self.client.get("/api/v1/backfills?dag_id=TEST_DAG_1", **kwargs) + assert response.status_code == expected + + +class TestGetBackfill(TestBackfillEndpoint): + def test_should_respond_200(self, session): + (dag,) = self._create_dag_models() + from_date = timezone.utcnow() + to_date = timezone.utcnow() + backfill = Backfill(dag_id=dag.dag_id, from_date=from_date, to_date=to_date) + session.add(backfill) + session.commit() + response = self.client.get( + f"/api/v1/backfills/{backfill.id}", + environ_overrides={"REMOTE_USER": "test"}, + ) + assert response.status_code == 200 + assert response.json == { + "completed_at": mock.ANY, + "created_at": mock.ANY, + "dag_id": "TEST_DAG_1", + "dag_run_conf": None, + "from_date": from_date.isoformat(), + "id": backfill.id, + "is_paused": False, + "max_active_runs": 10, + "to_date": to_date.isoformat(), + "updated_at": mock.ANY, + } + + def test_no_exist(self, session): + response = self.client.get( + f"/api/v1/backfills/{23198409834208}", + environ_overrides={"REMOTE_USER": "test"}, + ) + assert response.status_code == 404 + assert response.json.get("title") == "Backfill not found" + + @pytest.mark.parametrize( + "user, expected", + [ + ("test_granular_permissions", 200), + ("test_no_permissions", 403), + ("test", 200), + (None, 401), + ], + ) + def test_should_respond_200_with_granular_dag_access(self, user, expected, session): + (dag,) = self._create_dag_models() + from_date = timezone.utcnow() + to_date = timezone.utcnow() + backfill = Backfill( + dag_id=dag.dag_id, + from_date=from_date, + to_date=to_date, + ) + session.add(backfill) + session.commit() + kwargs = {} + if user: + kwargs.update(environ_overrides={"REMOTE_USER": user}) + response = self.client.get(f"/api/v1/backfills/{backfill.id}", **kwargs) + assert response.status_code == expected + + +class TestCreateBackfill(TestBackfillEndpoint): + @pytest.mark.parametrize( + "user, expected", + [ + ("test_granular_permissions", 200), + ("test_no_permissions", 403), + ("test", 200), + (None, 401), + ], + ) + def test_create_backfill(self, user, expected, session, dag_maker): + with dag_maker(session=session, dag_id="TEST_DAG_1", schedule="0 * * * *") as dag: + EmptyOperator(task_id="mytask") + session.add(SerializedDagModel(dag)) + session.commit() + session.query(DagModel).all() + from_date = pendulum.parse("2024-01-01") + from_date_iso = from_date.isoformat() + to_date = pendulum.parse("2024-02-01") + to_date_iso = to_date.isoformat() + max_active_runs = 5 + query = urlencode( + query={ + "dag_id": dag.dag_id, + "from_date": f"{from_date_iso}", + "to_date": f"{to_date_iso}", + "max_active_runs": max_active_runs, + "reverse": False, + } + ) + kwargs = {} + if user: + kwargs.update(environ_overrides={"REMOTE_USER": user}) + + response = self.client.post( + f"/api/v1/backfills?{query}", + **kwargs, + ) + assert response.status_code == expected + if expected < 300: + assert response.json == { + "completed_at": mock.ANY, + "created_at": mock.ANY, + "dag_id": "TEST_DAG_1", + "dag_run_conf": None, + "from_date": from_date_iso, + "id": mock.ANY, + "is_paused": False, + "max_active_runs": 5, + "to_date": to_date_iso, + "updated_at": mock.ANY, + } + + +class TestPauseBackfill(TestBackfillEndpoint): + def test_should_respond_200(self, session): + (dag,) = self._create_dag_models() + from_date = timezone.utcnow() + to_date = timezone.utcnow() + backfill = Backfill(dag_id=dag.dag_id, from_date=from_date, to_date=to_date) + session.add(backfill) + session.commit() + response = self.client.post( + f"/api/v1/backfills/{backfill.id}/pause", + environ_overrides={"REMOTE_USER": "test"}, + ) + assert response.status_code == 200 + assert response.json == { + "completed_at": mock.ANY, + "created_at": mock.ANY, + "dag_id": "TEST_DAG_1", + "dag_run_conf": None, + "from_date": from_date.isoformat(), + "id": backfill.id, + "is_paused": True, + "max_active_runs": 10, + "to_date": to_date.isoformat(), + "updated_at": mock.ANY, + } + + @pytest.mark.parametrize( + "user, expected", + [ + ("test_granular_permissions", 200), + ("test_no_permissions", 403), + ("test", 200), + (None, 401), + ], + ) + def test_should_respond_200_with_granular_dag_access(self, user, expected, session): + (dag,) = self._create_dag_models() + from_date = timezone.utcnow() + to_date = timezone.utcnow() + backfill = Backfill( + dag_id=dag.dag_id, + from_date=from_date, + to_date=to_date, + ) + session.add(backfill) + session.commit() + kwargs = {} + if user: + kwargs.update(environ_overrides={"REMOTE_USER": user}) + response = self.client.post(f"/api/v1/backfills/{backfill.id}/pause", **kwargs) + assert response.status_code == expected + + +class TestCancelBackfill(TestBackfillEndpoint): + def test_should_respond_200(self, session): + (dag,) = self._create_dag_models() + from_date = timezone.utcnow() + to_date = timezone.utcnow() + backfill = Backfill(dag_id=dag.dag_id, from_date=from_date, to_date=to_date) + session.add(backfill) + session.commit() + response = self.client.post( + f"/api/v1/backfills/{backfill.id}/cancel", + environ_overrides={"REMOTE_USER": "test"}, + ) + assert response.status_code == 200 + assert response.json == { + "completed_at": mock.ANY, + "created_at": mock.ANY, + "dag_id": "TEST_DAG_1", + "dag_run_conf": None, + "from_date": from_date.isoformat(), + "id": backfill.id, + "is_paused": True, + "max_active_runs": 10, + "to_date": to_date.isoformat(), + "updated_at": mock.ANY, + } + assert pendulum.parse(response.json["completed_at"]) + # now it is marked as completed + assert pendulum.parse(response.json["completed_at"]) + + # get conflict when canceling already-canceled backfill + response = self.client.post( + f"/api/v1/backfills/{backfill.id}/cancel", environ_overrides={"REMOTE_USER": "test"} + ) + assert response.status_code == 409 + + @pytest.mark.parametrize( + "user, expected", + [ + ("test_granular_permissions", 200), + ("test_no_permissions", 403), + ("test", 200), + (None, 401), + ], + ) + def test_should_respond_200_with_granular_dag_access(self, user, expected, session): + (dag,) = self._create_dag_models() + from_date = timezone.utcnow() + to_date = timezone.utcnow() + backfill = Backfill( + dag_id=dag.dag_id, + from_date=from_date, + to_date=to_date, + ) + session.add(backfill) + session.commit() + kwargs = {} + if user: + kwargs.update(environ_overrides={"REMOTE_USER": user}) + response = self.client.post(f"/api/v1/backfills/{backfill.id}/cancel", **kwargs) + assert response.status_code == expected + if response.status_code < 300: + # now it is marked as completed + assert pendulum.parse(response.json["completed_at"]) + + # get conflict when canceling already-canceled backfill + response = self.client.post(f"/api/v1/backfills/{backfill.id}/cancel", **kwargs) + assert response.status_code == 409 diff --git a/tests/test_utils/db.py b/tests/test_utils/db.py index ceb6bc94b8dce..77875bb03ec51 100644 --- a/tests/test_utils/db.py +++ b/tests/test_utils/db.py @@ -35,6 +35,7 @@ Variable, XCom, ) +from airflow.models.backfill import Backfill, BackfillDagRun from airflow.models.dag import DagOwnerAttributes from airflow.models.dagcode import DagCode from airflow.models.dagwarning import DagWarning @@ -66,6 +67,12 @@ def clear_db_runs(): pass +def clear_db_backfills(): + with create_session() as session: + session.query(BackfillDagRun).delete() + session.query(Backfill).delete() + + def clear_db_datasets(): with create_session() as session: session.query(DatasetEvent).delete() From 01844582432dbd3626ea85bd361470ded226fbcb Mon Sep 17 00:00:00 2001 From: Kacper Muda Date: Thu, 26 Sep 2024 13:34:41 -0400 Subject: [PATCH 044/802] fix: OL dag start event not being emitted (#42448) Signed-off-by: Kacper Muda --- airflow/providers/openlineage/plugins/listener.py | 1 - 1 file changed, 1 deletion(-) diff --git a/airflow/providers/openlineage/plugins/listener.py b/airflow/providers/openlineage/plugins/listener.py index fbe50e4a728a5..9c568aa374196 100644 --- a/airflow/providers/openlineage/plugins/listener.py +++ b/airflow/providers/openlineage/plugins/listener.py @@ -439,7 +439,6 @@ def on_dag_run_running(self, dag_run: DagRun, msg: str) -> None: self.submit_callable( self.adapter.dag_started, dag_id=dag_run.dag_id, - run_id=dag_run.run_id, logical_date=dag_run.logical_date, start_date=dag_run.start_date, nominal_start_time=data_interval_start, From 7e2cf994015d6f9d795ec2728e518d531d499c99 Mon Sep 17 00:00:00 2001 From: Jens Scheffler <95105677+jscheffl@users.noreply.github.com> Date: Thu, 26 Sep 2024 23:44:20 +0200 Subject: [PATCH 045/802] Fix DB utils for Airflow Backwards compatability tests - import not existing (#42522) --- tests/test_utils/db.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/test_utils/db.py b/tests/test_utils/db.py index 77875bb03ec51..bd56ed9175cc4 100644 --- a/tests/test_utils/db.py +++ b/tests/test_utils/db.py @@ -35,7 +35,6 @@ Variable, XCom, ) -from airflow.models.backfill import Backfill, BackfillDagRun from airflow.models.dag import DagOwnerAttributes from airflow.models.dagcode import DagCode from airflow.models.dagwarning import DagWarning @@ -68,6 +67,8 @@ def clear_db_runs(): def clear_db_backfills(): + from airflow.models.backfill import Backfill, BackfillDagRun + with create_session() as session: session.query(BackfillDagRun).delete() session.query(Backfill).delete() From 6f0303ea1e73889e8abca9b0e999d26c7cc25b5c Mon Sep 17 00:00:00 2001 From: Elad Kalif <45845474+eladkal@users.noreply.github.com> Date: Fri, 27 Sep 2024 08:40:54 +0700 Subject: [PATCH 046/802] Prepare docs for Sep 2nd adhoc wave of providers (#42519) --- airflow/providers/common/sql/CHANGELOG.rst | 9 +++++++++ airflow/providers/common/sql/__init__.py | 2 +- airflow/providers/common/sql/provider.yaml | 3 ++- airflow/providers/openlineage/CHANGELOG.rst | 9 +++++++++ airflow/providers/openlineage/__init__.py | 2 +- airflow/providers/openlineage/provider.yaml | 3 ++- .../commits.rst | 15 ++++++++++++++- .../apache-airflow-providers-common-sql/index.rst | 6 +++--- .../commits.rst | 15 ++++++++++++++- .../index.rst | 6 +++--- 10 files changed, 58 insertions(+), 12 deletions(-) diff --git a/airflow/providers/common/sql/CHANGELOG.rst b/airflow/providers/common/sql/CHANGELOG.rst index ff4a7a74d5c02..531353f8c800f 100644 --- a/airflow/providers/common/sql/CHANGELOG.rst +++ b/airflow/providers/common/sql/CHANGELOG.rst @@ -25,6 +25,15 @@ Changelog --------- +1.17.1 +...... + +Bug Fixes +~~~~~~~~~ + +* ``fix(providers/common/sql): add dummy connection setter for backward compatibility (#42490)`` +* ``Changed type hinting for handler function (#42275)`` + 1.17.0 ...... diff --git a/airflow/providers/common/sql/__init__.py b/airflow/providers/common/sql/__init__.py index e1c93c3efb9f1..6ef37aa0ed669 100644 --- a/airflow/providers/common/sql/__init__.py +++ b/airflow/providers/common/sql/__init__.py @@ -29,7 +29,7 @@ __all__ = ["__version__"] -__version__ = "1.17.0" +__version__ = "1.17.1" if packaging.version.parse(packaging.version.parse(airflow_version).base_version) < packaging.version.parse( "2.8.0" diff --git a/airflow/providers/common/sql/provider.yaml b/airflow/providers/common/sql/provider.yaml index f600acc1fa44b..ec487aca3f001 100644 --- a/airflow/providers/common/sql/provider.yaml +++ b/airflow/providers/common/sql/provider.yaml @@ -22,9 +22,10 @@ description: | `Common SQL Provider `__ state: ready -source-date-epoch: 1723970051 +source-date-epoch: 1727372263 # note that those versions are maintained by release manager - do not update them manually versions: + - 1.17.1 - 1.17.0 - 1.16.0 - 1.15.0 diff --git a/airflow/providers/openlineage/CHANGELOG.rst b/airflow/providers/openlineage/CHANGELOG.rst index 318d0d92b271b..0e35dab6deaa2 100644 --- a/airflow/providers/openlineage/CHANGELOG.rst +++ b/airflow/providers/openlineage/CHANGELOG.rst @@ -26,6 +26,15 @@ Changelog --------- +1.12.1 +...... + +Bug Fixes +~~~~~~~~~ + +* ``fix: OpenLineage dag start event not being emitted (#42448)`` +* ``fix: typo in error stack trace formatting for clearer output (#42017)`` + 1.12.0 ...... diff --git a/airflow/providers/openlineage/__init__.py b/airflow/providers/openlineage/__init__.py index 6c3c88bb926f4..664e5530ebf97 100644 --- a/airflow/providers/openlineage/__init__.py +++ b/airflow/providers/openlineage/__init__.py @@ -29,7 +29,7 @@ __all__ = ["__version__"] -__version__ = "1.12.0" +__version__ = "1.12.1" if packaging.version.parse(packaging.version.parse(airflow_version).base_version) < packaging.version.parse( "2.8.0" diff --git a/airflow/providers/openlineage/provider.yaml b/airflow/providers/openlineage/provider.yaml index af13b1954b71e..b249ff46c8591 100644 --- a/airflow/providers/openlineage/provider.yaml +++ b/airflow/providers/openlineage/provider.yaml @@ -22,9 +22,10 @@ description: | `OpenLineage `__ state: ready -source-date-epoch: 1726861079 +source-date-epoch: 1727372276 # note that those versions are maintained by release manager - do not update them manually versions: + - 1.12.1 - 1.12.0 - 1.11.0 - 1.10.0 diff --git a/docs/apache-airflow-providers-common-sql/commits.rst b/docs/apache-airflow-providers-common-sql/commits.rst index 95e835b0bf80a..f719dd7b39811 100644 --- a/docs/apache-airflow-providers-common-sql/commits.rst +++ b/docs/apache-airflow-providers-common-sql/commits.rst @@ -35,14 +35,27 @@ For high-level changelog, see :doc:`package information including changelog `_ 2024-09-26 ``fix(providers/common/sql): add dummy connection setter for backward compatibility (#42490)`` +`47c71108a8 `_ 2024-09-22 ``Changed type hinting for handler function (#42275)`` +================================================================================================= =========== ============================================================================================== + 1.17.0 ...... -Latest change: 2024-09-05 +Latest change: 2024-09-21 ================================================================================================= =========== ================================================================================= Commit Committed Subject ================================================================================================= =========== ================================================================================= +`7628d47d04 `_ 2024-09-21 ``Prepare docs for Sep 1st wave of providers (#42387)`` `17c30b4f21 `_ 2024-09-05 ``feat: log client db messages for provider postgres (#40171)`` `2e813eb87d `_ 2024-09-04 ``Generalize caching of connection in DbApiHook to improve performance (#40751)`` `1613e9ec1c `_ 2024-08-25 ``remove soft_fail (#41710)`` diff --git a/docs/apache-airflow-providers-common-sql/index.rst b/docs/apache-airflow-providers-common-sql/index.rst index e707a38902632..573603b2d7d04 100644 --- a/docs/apache-airflow-providers-common-sql/index.rst +++ b/docs/apache-airflow-providers-common-sql/index.rst @@ -77,7 +77,7 @@ apache-airflow-providers-common-sql package `Common SQL Provider `__ -Release: 1.17.0 +Release: 1.17.1 Provider package ---------------- @@ -130,5 +130,5 @@ Downloading official packages You can download officially released packages and verify their checksums and signatures from the `Official Apache Download site `_ -* `The apache-airflow-providers-common-sql 1.17.0 sdist package `_ (`asc `__, `sha512 `__) -* `The apache-airflow-providers-common-sql 1.17.0 wheel package `_ (`asc `__, `sha512 `__) +* `The apache-airflow-providers-common-sql 1.17.1 sdist package `_ (`asc `__, `sha512 `__) +* `The apache-airflow-providers-common-sql 1.17.1 wheel package `_ (`asc `__, `sha512 `__) diff --git a/docs/apache-airflow-providers-openlineage/commits.rst b/docs/apache-airflow-providers-openlineage/commits.rst index 299dc020aa36a..d2e20868233c3 100644 --- a/docs/apache-airflow-providers-openlineage/commits.rst +++ b/docs/apache-airflow-providers-openlineage/commits.rst @@ -35,14 +35,27 @@ For high-level changelog, see :doc:`package information including changelog `_ 2024-09-26 ``fix: OL dag start event not being emitted (#42448)`` +`ffff0e8b33 `_ 2024-09-23 ``Fix typo in error stack trace formatting for clearer output (#42017)`` +================================================================================================= =========== ======================================================================== + 1.12.0 ...... -Latest change: 2024-09-10 +Latest change: 2024-09-21 ================================================================================================= =========== ======================================================================================================================================================= Commit Committed Subject ================================================================================================= =========== ======================================================================================================================================================= +`7628d47d04 `_ 2024-09-21 ``Prepare docs for Sep 1st wave of providers (#42387)`` `e05c0358af `_ 2024-09-10 ``chore: bump OL provider dependencies versions (#42059)`` `aa23bfdbc7 `_ 2024-09-02 ``feat: notify about potential serialization failures when sending DagRun, don't serialize unnecessary params, guard listener for exceptions (#41690)`` `8640f3e397 `_ 2024-09-02 ``move to dag_run.logical_date from execution date in OpenLineage provider (#41889)`` diff --git a/docs/apache-airflow-providers-openlineage/index.rst b/docs/apache-airflow-providers-openlineage/index.rst index 623bb4580f3f6..817bddd6c0b74 100644 --- a/docs/apache-airflow-providers-openlineage/index.rst +++ b/docs/apache-airflow-providers-openlineage/index.rst @@ -73,7 +73,7 @@ apache-airflow-providers-openlineage package `OpenLineage `__ -Release: 1.12.0 +Release: 1.12.1 Provider package ---------------- @@ -129,5 +129,5 @@ Downloading official packages You can download officially released packages and verify their checksums and signatures from the `Official Apache Download site `_ -* `The apache-airflow-providers-openlineage 1.12.0 sdist package `_ (`asc `__, `sha512 `__) -* `The apache-airflow-providers-openlineage 1.12.0 wheel package `_ (`asc `__, `sha512 `__) +* `The apache-airflow-providers-openlineage 1.12.1 sdist package `_ (`asc `__, `sha512 `__) +* `The apache-airflow-providers-openlineage 1.12.1 wheel package `_ (`asc `__, `sha512 `__) From 016b43198e13df3eb12664d22095c700de5ea114 Mon Sep 17 00:00:00 2001 From: Pierre Jeambrun Date: Fri, 27 Sep 2024 15:19:18 +0800 Subject: [PATCH 047/802] Add more sort and filters to get dags endpoint (#42462) --- airflow/api_fastapi/db.py | 11 +++ airflow/api_fastapi/openapi/v1-generated.yaml | 24 ++++++ airflow/api_fastapi/parameters.py | 84 ++++++++++++------- airflow/api_fastapi/views/public/dags.py | 32 ++++++- airflow/ui/openapi-gen/queries/common.ts | 4 + airflow/ui/openapi-gen/queries/prefetch.ts | 6 ++ airflow/ui/openapi-gen/queries/queries.ts | 7 +- airflow/ui/openapi-gen/queries/suspense.ts | 6 ++ .../ui/openapi-gen/requests/schemas.gen.ts | 11 +++ .../ui/openapi-gen/requests/services.gen.ts | 2 + airflow/ui/openapi-gen/requests/types.gen.ts | 10 +++ tests/api_fastapi/views/public/test_dags.py | 60 ++++++++++--- 12 files changed, 209 insertions(+), 48 deletions(-) diff --git a/airflow/api_fastapi/db.py b/airflow/api_fastapi/db.py index 51faee25ed5a0..c3ed01a0aefec 100644 --- a/airflow/api_fastapi/db.py +++ b/airflow/api_fastapi/db.py @@ -19,6 +19,9 @@ from typing import TYPE_CHECKING +from sqlalchemy import func, select + +from airflow.models.dagrun import DagRun from airflow.utils.session import create_session if TYPE_CHECKING: @@ -52,3 +55,11 @@ def apply_filters_to_select(base_select: Select, filters: list[BaseParam]) -> Se select = filter.to_orm(select) return select + + +latest_dag_run_per_dag_id_cte = ( + select(DagRun.dag_id, func.max(DagRun.start_date).label("start_date")) + .where() + .group_by(DagRun.dag_id) + .cte() +) diff --git a/airflow/api_fastapi/openapi/v1-generated.yaml b/airflow/api_fastapi/openapi/v1-generated.yaml index 6d77056d0574d..f488825449c3a 100644 --- a/airflow/api_fastapi/openapi/v1-generated.yaml +++ b/airflow/api_fastapi/openapi/v1-generated.yaml @@ -103,6 +103,14 @@ paths: - type: boolean - type: 'null' title: Paused + - name: last_dag_run_state + in: query + required: false + schema: + anyOf: + - $ref: '#/components/schemas/DagRunState' + - type: 'null' + title: Last Dag Run State - name: order_by in: query required: false @@ -347,6 +355,22 @@ components: - file_token title: DAGResponse description: DAG serializer for responses. + DagRunState: + type: string + enum: + - queued + - running + - success + - failed + title: DagRunState + description: 'All possible states that a DagRun can be in. + + + These are "shared" with TaskInstanceState in some parts of the code, + + so please ensure that their values always match the ones with the + + same name in TaskInstanceState.' DagTagPydantic: properties: name: diff --git a/airflow/api_fastapi/parameters.py b/airflow/api_fastapi/parameters.py index 589403cc4e960..09eea5f6e055b 100644 --- a/airflow/api_fastapi/parameters.py +++ b/airflow/api_fastapi/parameters.py @@ -18,13 +18,15 @@ from __future__ import annotations from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Generic, List, TypeVar, Union +from typing import TYPE_CHECKING, Any, Generic, List, TypeVar from fastapi import Depends, HTTPException, Query from sqlalchemy import case, or_ from typing_extensions import Annotated, Self from airflow.models.dag import DagModel, DagTag +from airflow.models.dagrun import DagRun +from airflow.utils.state import DagRunState if TYPE_CHECKING: from sqlalchemy.sql import ColumnElement, Select @@ -43,14 +45,14 @@ def __init__(self) -> None: def to_orm(self, select: Select) -> Select: pass - @abstractmethod - def __call__(self, *args: Any, **kwarg: Any) -> BaseParam: - pass - - def set_value(self, value: T) -> Self: + def set_value(self, value: T | None) -> Self: self.value = value return self + @abstractmethod + def depends(self, *args: Any, **kwargs: Any) -> Self: + pass + class _LimitFilter(BaseParam[int]): """Filter on the limit.""" @@ -61,7 +63,7 @@ def to_orm(self, select: Select) -> Select: return select.limit(self.value) - def __call__(self, limit: int = 100) -> _LimitFilter: + def depends(self, limit: int = 100) -> _LimitFilter: return self.set_value(limit) @@ -73,11 +75,11 @@ def to_orm(self, select: Select) -> Select: return select return select.offset(self.value) - def __call__(self, offset: int = 0) -> _OffsetFilter: + def depends(self, offset: int = 0) -> _OffsetFilter: return self.set_value(offset) -class _PausedFilter(BaseParam[Union[bool, None]]): +class _PausedFilter(BaseParam[bool]): """Filter on is_paused.""" def to_orm(self, select: Select) -> Select: @@ -85,7 +87,7 @@ def to_orm(self, select: Select) -> Select: return select return select.where(DagModel.is_paused == self.value) - def __call__(self, paused: bool | None = Query(default=None)) -> _PausedFilter: + def depends(self, paused: bool | None = None) -> _PausedFilter: return self.set_value(paused) @@ -97,11 +99,11 @@ def to_orm(self, select: Select) -> Select: return select.where(DagModel.is_active == self.value) return select - def __call__(self, only_active: bool = Query(default=True)) -> _OnlyActiveFilter: + def depends(self, only_active: bool = True) -> _OnlyActiveFilter: return self.set_value(only_active) -class _SearchParam(BaseParam[Union[str, None]]): +class _SearchParam(BaseParam[str]): """Search on attribute.""" def __init__(self, attribute: ColumnElement) -> None: @@ -120,7 +122,7 @@ class _DagIdPatternSearch(_SearchParam): def __init__(self) -> None: super().__init__(DagModel.dag_id) - def __call__(self, dag_id_pattern: str | None = Query(default=None)) -> _DagIdPatternSearch: + def depends(self, dag_id_pattern: str | None = None) -> _DagIdPatternSearch: return self.set_value(dag_id_pattern) @@ -130,15 +132,18 @@ class _DagDisplayNamePatternSearch(_SearchParam): def __init__(self) -> None: super().__init__(DagModel.dag_display_name) - def __call__( - self, dag_display_name_pattern: str | None = Query(default=None) - ) -> _DagDisplayNamePatternSearch: + def depends(self, dag_display_name_pattern: str | None = None) -> _DagDisplayNamePatternSearch: return self.set_value(dag_display_name_pattern) -class SortParam(BaseParam[Union[str]]): +class SortParam(BaseParam[str]): """Order result by the attribute.""" + attr_mapping = { + "last_run_state": DagRun.state, + "last_run_start_date": DagRun.start_date, + } + def __init__(self, allowed_attrs: list[str]) -> None: super().__init__() self.allowed_attrs = allowed_attrs @@ -155,17 +160,17 @@ def to_orm(self, select: Select) -> Select: f"the attribute does not exist on the model", ) - column = getattr(DagModel, lstriped_orderby) + column = self.attr_mapping.get(lstriped_orderby, None) or getattr(DagModel, lstriped_orderby) # MySQL does not support `nullslast`, and True/False ordering depends on the - # database implementation + # database implementation. nullscheck = case((column.isnot(None), 0), else_=1) if self.value[0] == "-": - return select.order_by(nullscheck, column.desc(), DagModel.dag_id) + return select.order_by(nullscheck, column.desc(), DagModel.dag_id.desc()) else: - return select.order_by(nullscheck, column.asc(), DagModel.dag_id) + return select.order_by(nullscheck, column.asc(), DagModel.dag_id.asc()) - def __call__(self, order_by: str = Query(default="dag_id")) -> SortParam: + def depends(self, order_by: str = "dag_id") -> SortParam: return self.set_value(order_by) @@ -179,7 +184,7 @@ def to_orm(self, select: Select) -> Select: conditions = [DagModel.tags.any(DagTag.name == tag) for tag in self.value] return select.where(or_(*conditions)) - def __call__(self, tags: list[str] = Query(default_factory=list)) -> _TagsFilter: + def depends(self, tags: list[str] = Query(default_factory=list)) -> _TagsFilter: return self.set_value(tags) @@ -193,17 +198,32 @@ def to_orm(self, select: Select) -> Select: conditions = [DagModel.owners.ilike(f"%{owner}%") for owner in self.value] return select.where(or_(*conditions)) - def __call__(self, owners: list[str] = Query(default_factory=list)) -> _OwnersFilter: + def depends(self, owners: list[str] = Query(default_factory=list)) -> _OwnersFilter: return self.set_value(owners) -QueryLimit = Annotated[_LimitFilter, Depends(_LimitFilter())] -QueryOffset = Annotated[_OffsetFilter, Depends(_OffsetFilter())] -QueryPausedFilter = Annotated[_PausedFilter, Depends(_PausedFilter())] -QueryOnlyActiveFilter = Annotated[_OnlyActiveFilter, Depends(_OnlyActiveFilter())] -QueryDagIdPatternSearch = Annotated[_DagIdPatternSearch, Depends(_DagIdPatternSearch())] +class _LastDagRunStateFilter(BaseParam[DagRunState]): + """Filter on the state of the latest DagRun.""" + + def to_orm(self, select: Select) -> Select: + if self.value is None: + return select + return select.where(DagRun.state == self.value) + + def depends(self, last_dag_run_state: DagRunState | None = None) -> _LastDagRunStateFilter: + return self.set_value(last_dag_run_state) + + +# DAG +QueryLimit = Annotated[_LimitFilter, Depends(_LimitFilter().depends)] +QueryOffset = Annotated[_OffsetFilter, Depends(_OffsetFilter().depends)] +QueryPausedFilter = Annotated[_PausedFilter, Depends(_PausedFilter().depends)] +QueryOnlyActiveFilter = Annotated[_OnlyActiveFilter, Depends(_OnlyActiveFilter().depends)] +QueryDagIdPatternSearch = Annotated[_DagIdPatternSearch, Depends(_DagIdPatternSearch().depends)] QueryDagDisplayNamePatternSearch = Annotated[ - _DagDisplayNamePatternSearch, Depends(_DagDisplayNamePatternSearch()) + _DagDisplayNamePatternSearch, Depends(_DagDisplayNamePatternSearch().depends) ] -QueryTagsFilter = Annotated[_TagsFilter, Depends(_TagsFilter())] -QueryOwnersFilter = Annotated[_OwnersFilter, Depends(_OwnersFilter())] +QueryTagsFilter = Annotated[_TagsFilter, Depends(_TagsFilter().depends)] +QueryOwnersFilter = Annotated[_OwnersFilter, Depends(_OwnersFilter().depends)] +# DagRun +QueryLastDagRunStateFilter = Annotated[_LastDagRunStateFilter, Depends(_LastDagRunStateFilter().depends)] diff --git a/airflow/api_fastapi/views/public/dags.py b/airflow/api_fastapi/views/public/dags.py index 433e5ef862447..07ab968adc975 100644 --- a/airflow/api_fastapi/views/public/dags.py +++ b/airflow/api_fastapi/views/public/dags.py @@ -22,10 +22,11 @@ from sqlalchemy.orm import Session from typing_extensions import Annotated -from airflow.api_fastapi.db import apply_filters_to_select, get_session +from airflow.api_fastapi.db import apply_filters_to_select, get_session, latest_dag_run_per_dag_id_cte from airflow.api_fastapi.parameters import ( QueryDagDisplayNamePatternSearch, QueryDagIdPatternSearch, + QueryLastDagRunStateFilter, QueryLimit, QueryOffset, QueryOnlyActiveFilter, @@ -36,6 +37,7 @@ ) from airflow.api_fastapi.serializers.dags import DAGCollectionResponse, DAGPatchBody, DAGResponse from airflow.models import DagModel +from airflow.models.dagrun import DagRun from airflow.utils.db import get_query_count dags_router = APIRouter(tags=["DAG"]) @@ -51,14 +53,36 @@ async def get_dags( dag_display_name_pattern: QueryDagDisplayNamePatternSearch, only_active: QueryOnlyActiveFilter, paused: QueryPausedFilter, - order_by: Annotated[SortParam, Depends(SortParam(["dag_id", "dag_display_name", "next_dagrun"]))], + last_dag_run_state: QueryLastDagRunStateFilter, + order_by: Annotated[ + SortParam, + Depends( + SortParam( + ["dag_id", "dag_display_name", "next_dagrun", "last_run_state", "last_run_start_date"] + ).depends + ), + ], session: Annotated[Session, Depends(get_session)], ) -> DAGCollectionResponse: """Get all DAGs.""" - dags_query = select(DagModel) + dags_query = ( + select(DagModel) + .join( + latest_dag_run_per_dag_id_cte, + DagModel.dag_id == latest_dag_run_per_dag_id_cte.c.dag_id, + isouter=True, + ) + .join( + DagRun, + DagRun.start_date == latest_dag_run_per_dag_id_cte.c.start_date + and DagRun.dag_id == latest_dag_run_per_dag_id_cte.c.dag_id, + isouter=True, + ) + ) dags_query = apply_filters_to_select( - dags_query, [only_active, paused, dag_id_pattern, dag_display_name_pattern, tags, owners] + dags_query, + [only_active, paused, dag_id_pattern, dag_display_name_pattern, tags, owners, last_dag_run_state], ) # TODO: Re-enable when permissions are handled. diff --git a/airflow/ui/openapi-gen/queries/common.ts b/airflow/ui/openapi-gen/queries/common.ts index 2818b48a33e1a..b8021fed9be3c 100644 --- a/airflow/ui/openapi-gen/queries/common.ts +++ b/airflow/ui/openapi-gen/queries/common.ts @@ -2,6 +2,7 @@ import { UseQueryResult } from "@tanstack/react-query"; import { DagService, DatasetService } from "../requests/services.gen"; +import { DagRunState } from "../requests/types.gen"; export type DatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGetDefaultResponse = Awaited< @@ -37,6 +38,7 @@ export const UseDagServiceGetDagsPublicDagsGetKeyFn = ( { dagDisplayNamePattern, dagIdPattern, + lastDagRunState, limit, offset, onlyActive, @@ -47,6 +49,7 @@ export const UseDagServiceGetDagsPublicDagsGetKeyFn = ( }: { dagDisplayNamePattern?: string; dagIdPattern?: string; + lastDagRunState?: DagRunState; limit?: number; offset?: number; onlyActive?: boolean; @@ -62,6 +65,7 @@ export const UseDagServiceGetDagsPublicDagsGetKeyFn = ( { dagDisplayNamePattern, dagIdPattern, + lastDagRunState, limit, offset, onlyActive, diff --git a/airflow/ui/openapi-gen/queries/prefetch.ts b/airflow/ui/openapi-gen/queries/prefetch.ts index f8e1bf616d143..6dd99f96b8425 100644 --- a/airflow/ui/openapi-gen/queries/prefetch.ts +++ b/airflow/ui/openapi-gen/queries/prefetch.ts @@ -2,6 +2,7 @@ import { type QueryClient } from "@tanstack/react-query"; import { DagService, DatasetService } from "../requests/services.gen"; +import { DagRunState } from "../requests/types.gen"; import * as Common from "./common"; /** @@ -40,6 +41,7 @@ export const prefetchUseDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGet = * @param data.dagDisplayNamePattern * @param data.onlyActive * @param data.paused + * @param data.lastDagRunState * @param data.orderBy * @returns DAGCollectionResponse Successful Response * @throws ApiError @@ -49,6 +51,7 @@ export const prefetchUseDagServiceGetDagsPublicDagsGet = ( { dagDisplayNamePattern, dagIdPattern, + lastDagRunState, limit, offset, onlyActive, @@ -59,6 +62,7 @@ export const prefetchUseDagServiceGetDagsPublicDagsGet = ( }: { dagDisplayNamePattern?: string; dagIdPattern?: string; + lastDagRunState?: DagRunState; limit?: number; offset?: number; onlyActive?: boolean; @@ -72,6 +76,7 @@ export const prefetchUseDagServiceGetDagsPublicDagsGet = ( queryKey: Common.UseDagServiceGetDagsPublicDagsGetKeyFn({ dagDisplayNamePattern, dagIdPattern, + lastDagRunState, limit, offset, onlyActive, @@ -84,6 +89,7 @@ export const prefetchUseDagServiceGetDagsPublicDagsGet = ( DagService.getDagsPublicDagsGet({ dagDisplayNamePattern, dagIdPattern, + lastDagRunState, limit, offset, onlyActive, diff --git a/airflow/ui/openapi-gen/queries/queries.ts b/airflow/ui/openapi-gen/queries/queries.ts index 2a0c6b6821978..b771fccfeb947 100644 --- a/airflow/ui/openapi-gen/queries/queries.ts +++ b/airflow/ui/openapi-gen/queries/queries.ts @@ -7,7 +7,7 @@ import { } from "@tanstack/react-query"; import { DagService, DatasetService } from "../requests/services.gen"; -import { DAGPatchBody } from "../requests/types.gen"; +import { DAGPatchBody, DagRunState } from "../requests/types.gen"; import * as Common from "./common"; /** @@ -54,6 +54,7 @@ export const useDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGet = < * @param data.dagDisplayNamePattern * @param data.onlyActive * @param data.paused + * @param data.lastDagRunState * @param data.orderBy * @returns DAGCollectionResponse Successful Response * @throws ApiError @@ -66,6 +67,7 @@ export const useDagServiceGetDagsPublicDagsGet = < { dagDisplayNamePattern, dagIdPattern, + lastDagRunState, limit, offset, onlyActive, @@ -76,6 +78,7 @@ export const useDagServiceGetDagsPublicDagsGet = < }: { dagDisplayNamePattern?: string; dagIdPattern?: string; + lastDagRunState?: DagRunState; limit?: number; offset?: number; onlyActive?: boolean; @@ -92,6 +95,7 @@ export const useDagServiceGetDagsPublicDagsGet = < { dagDisplayNamePattern, dagIdPattern, + lastDagRunState, limit, offset, onlyActive, @@ -106,6 +110,7 @@ export const useDagServiceGetDagsPublicDagsGet = < DagService.getDagsPublicDagsGet({ dagDisplayNamePattern, dagIdPattern, + lastDagRunState, limit, offset, onlyActive, diff --git a/airflow/ui/openapi-gen/queries/suspense.ts b/airflow/ui/openapi-gen/queries/suspense.ts index bcc95a53e18ff..7743ce92d2855 100644 --- a/airflow/ui/openapi-gen/queries/suspense.ts +++ b/airflow/ui/openapi-gen/queries/suspense.ts @@ -2,6 +2,7 @@ import { UseQueryOptions, useSuspenseQuery } from "@tanstack/react-query"; import { DagService, DatasetService } from "../requests/services.gen"; +import { DagRunState } from "../requests/types.gen"; import * as Common from "./common"; /** @@ -49,6 +50,7 @@ export const useDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGetSuspense = * @param data.dagDisplayNamePattern * @param data.onlyActive * @param data.paused + * @param data.lastDagRunState * @param data.orderBy * @returns DAGCollectionResponse Successful Response * @throws ApiError @@ -61,6 +63,7 @@ export const useDagServiceGetDagsPublicDagsGetSuspense = < { dagDisplayNamePattern, dagIdPattern, + lastDagRunState, limit, offset, onlyActive, @@ -71,6 +74,7 @@ export const useDagServiceGetDagsPublicDagsGetSuspense = < }: { dagDisplayNamePattern?: string; dagIdPattern?: string; + lastDagRunState?: DagRunState; limit?: number; offset?: number; onlyActive?: boolean; @@ -87,6 +91,7 @@ export const useDagServiceGetDagsPublicDagsGetSuspense = < { dagDisplayNamePattern, dagIdPattern, + lastDagRunState, limit, offset, onlyActive, @@ -101,6 +106,7 @@ export const useDagServiceGetDagsPublicDagsGetSuspense = < DagService.getDagsPublicDagsGet({ dagDisplayNamePattern, dagIdPattern, + lastDagRunState, limit, offset, onlyActive, diff --git a/airflow/ui/openapi-gen/requests/schemas.gen.ts b/airflow/ui/openapi-gen/requests/schemas.gen.ts index 83d3670507f78..d9ce0528c396c 100644 --- a/airflow/ui/openapi-gen/requests/schemas.gen.ts +++ b/airflow/ui/openapi-gen/requests/schemas.gen.ts @@ -288,6 +288,17 @@ export const $DAGResponse = { description: "DAG serializer for responses.", } as const; +export const $DagRunState = { + type: "string", + enum: ["queued", "running", "success", "failed"], + title: "DagRunState", + description: `All possible states that a DagRun can be in. + +These are "shared" with TaskInstanceState in some parts of the code, +so please ensure that their values always match the ones with the +same name in TaskInstanceState.`, +} as const; + export const $DagTagPydantic = { properties: { name: { diff --git a/airflow/ui/openapi-gen/requests/services.gen.ts b/airflow/ui/openapi-gen/requests/services.gen.ts index a4c36d5990c78..9c261b3039000 100644 --- a/airflow/ui/openapi-gen/requests/services.gen.ts +++ b/airflow/ui/openapi-gen/requests/services.gen.ts @@ -48,6 +48,7 @@ export class DagService { * @param data.dagDisplayNamePattern * @param data.onlyActive * @param data.paused + * @param data.lastDagRunState * @param data.orderBy * @returns DAGCollectionResponse Successful Response * @throws ApiError @@ -67,6 +68,7 @@ export class DagService { dag_display_name_pattern: data.dagDisplayNamePattern, only_active: data.onlyActive, paused: data.paused, + last_dag_run_state: data.lastDagRunState, order_by: data.orderBy, }, errors: { diff --git a/airflow/ui/openapi-gen/requests/types.gen.ts b/airflow/ui/openapi-gen/requests/types.gen.ts index 2f6bc263d4289..803bcd84270c7 100644 --- a/airflow/ui/openapi-gen/requests/types.gen.ts +++ b/airflow/ui/openapi-gen/requests/types.gen.ts @@ -50,6 +50,15 @@ export type DAGResponse = { readonly file_token: string; }; +/** + * All possible states that a DagRun can be in. + * + * These are "shared" with TaskInstanceState in some parts of the code, + * so please ensure that their values always match the ones with the + * same name in TaskInstanceState. + */ +export type DagRunState = "queued" | "running" | "success" | "failed"; + /** * Serializable representation of the DagTag ORM SqlAlchemyModel used by internal API. */ @@ -79,6 +88,7 @@ export type NextRunDatasetsUiNextRunDatasetsDagIdGetResponse = { export type GetDagsPublicDagsGetData = { dagDisplayNamePattern?: string | null; dagIdPattern?: string | null; + lastDagRunState?: DagRunState | null; limit?: number; offset?: number; onlyActive?: boolean; diff --git a/tests/api_fastapi/views/public/test_dags.py b/tests/api_fastapi/views/public/test_dags.py index b508a1448352d..6e400f11cc0d2 100644 --- a/tests/api_fastapi/views/public/test_dags.py +++ b/tests/api_fastapi/views/public/test_dags.py @@ -20,9 +20,12 @@ import pytest -from airflow.models.dag import DAG, DagModel +from airflow.models.dag import DagModel +from airflow.models.dagrun import DagRun from airflow.operators.empty import EmptyOperator from airflow.utils.session import provide_session +from airflow.utils.state import DagRunState +from airflow.utils.types import DagRunType from tests.test_utils.db import clear_db_dags, clear_db_runs, clear_db_serialized_dags pytestmark = pytest.mark.db_test @@ -46,41 +49,62 @@ def _create_deactivated_paused_dag(session=None): owners="test_owner,another_test_owner", next_dagrun=datetime(2021, 1, 1, 12, 0, 0, tzinfo=timezone.utc), ) + + dagrun_failed = DagRun( + dag_id=DAG3_ID, + run_id="run1", + execution_date=datetime(2018, 1, 1, 12, 0, 0, tzinfo=timezone.utc), + start_date=datetime(2018, 1, 1, 12, 0, 0, tzinfo=timezone.utc), + run_type=DagRunType.SCHEDULED, + state=DagRunState.FAILED, + ) + + dagrun_success = DagRun( + dag_id=DAG3_ID, + run_id="run2", + execution_date=datetime(2019, 1, 1, 12, 0, 0, tzinfo=timezone.utc), + start_date=datetime(2019, 1, 1, 12, 0, 0, tzinfo=timezone.utc), + run_type=DagRunType.MANUAL, + state=DagRunState.SUCCESS, + ) + session.add(dag_model) + session.add(dagrun_failed) + session.add(dagrun_success) @pytest.fixture(autouse=True) -def setup() -> None: +def setup(dag_maker) -> None: clear_db_runs() clear_db_dags() clear_db_serialized_dags() - with DAG( + with dag_maker( DAG1_ID, dag_display_name=DAG1_DISPLAY_NAME, schedule=None, - start_date=datetime(2020, 6, 15), + start_date=datetime(2018, 6, 15, 0, 0, tzinfo=timezone.utc), doc_md="details", params={"foo": 1}, tags=["example"], - ) as dag1: + ): EmptyOperator(task_id=TASK_ID) - with DAG( + dag_maker.create_dagrun(state=DagRunState.FAILED) + + with dag_maker( DAG2_ID, dag_display_name=DAG2_DISPLAY_NAME, schedule=None, start_date=datetime( - 2020, + 2021, 6, 15, ), - ) as dag2: + ): EmptyOperator(task_id=TASK_ID) - dag1.sync_to_db() - dag2.sync_to_db() - + dag_maker.dagbag.sync_to_db() _create_deactivated_paused_dag() @@ -97,11 +121,25 @@ def setup() -> None: ({"paused": False}, 2, ["test_dag1", "test_dag2"]), ({"owners": ["airflow"]}, 2, ["test_dag1", "test_dag2"]), ({"owners": ["test_owner"], "only_active": False}, 1, ["test_dag3"]), + ({"last_dag_run_state": "success", "only_active": False}, 1, ["test_dag3"]), + ({"last_dag_run_state": "failed", "only_active": False}, 1, ["test_dag1"]), # # Sort ({"order_by": "-dag_id"}, 2, ["test_dag2", "test_dag1"]), ({"order_by": "-dag_display_name"}, 2, ["test_dag2", "test_dag1"]), ({"order_by": "dag_display_name"}, 2, ["test_dag1", "test_dag2"]), ({"order_by": "next_dagrun", "only_active": False}, 3, ["test_dag3", "test_dag1", "test_dag2"]), + ({"order_by": "last_run_state", "only_active": False}, 3, ["test_dag1", "test_dag3", "test_dag2"]), + ({"order_by": "-last_run_state", "only_active": False}, 3, ["test_dag3", "test_dag1", "test_dag2"]), + ( + {"order_by": "last_run_start_date", "only_active": False}, + 3, + ["test_dag1", "test_dag3", "test_dag2"], + ), + ( + {"order_by": "-last_run_start_date", "only_active": False}, + 3, + ["test_dag3", "test_dag1", "test_dag2"], + ), # Search ({"dag_id_pattern": "1"}, 1, ["test_dag1"]), ({"dag_display_name_pattern": "display2"}, 1, ["test_dag2"]), From 2d76042c800c076ff480d70ed9e0dbce9940b0a6 Mon Sep 17 00:00:00 2001 From: Jason <46563896+jsjasonseba@users.noreply.github.com> Date: Fri, 27 Sep 2024 15:11:08 +0700 Subject: [PATCH 048/802] Refactor ``bucket.get_blob`` calls in ``GCSHook`` to handle validation for non-existent objects. (#42474) * Fix implicit exception in GCSToLocalFilesystemOperator when reading non-existing files. * change method to internal method * Fix unit test expectations * Refactor bucket.get_blob calls in GCSHook to handle validation for non-existent objects --- airflow/providers/google/cloud/hooks/gcs.py | 43 +++++++------ newsfragments/42439.bugfix.rst | 1 + .../providers/google/cloud/hooks/test_gcs.py | 60 +++++++++++++------ 3 files changed, 68 insertions(+), 36 deletions(-) create mode 100644 newsfragments/42439.bugfix.rst diff --git a/airflow/providers/google/cloud/hooks/gcs.py b/airflow/providers/google/cloud/hooks/gcs.py index 83118e7079ca5..fb48fcd190609 100644 --- a/airflow/providers/google/cloud/hooks/gcs.py +++ b/airflow/providers/google/cloud/hooks/gcs.py @@ -59,6 +59,7 @@ from aiohttp import ClientSession from google.api_core.retry import Retry + from google.cloud.storage.blob import Blob RT = TypeVar("RT") @@ -597,11 +598,7 @@ def get_blob_update_time(self, bucket_name: str, object_name: str): :param object_name: The name of the blob to get updated time from the Google cloud storage bucket. """ - client = self.get_conn() - bucket = client.bucket(bucket_name) - blob = bucket.get_blob(blob_name=object_name) - if blob is None: - raise ValueError(f"Object ({object_name}) not found in Bucket ({bucket_name})") + blob = self._get_blob(bucket_name, object_name) return blob.updated def is_updated_after(self, bucket_name: str, object_name: str, ts: datetime) -> bool: @@ -957,19 +954,35 @@ def list_by_timespan( break return ids - def get_size(self, bucket_name: str, object_name: str) -> int: + def _get_blob(self, bucket_name: str, object_name: str) -> Blob: """ - Get the size of a file in Google Cloud Storage. + Get a blob object in Google Cloud Storage. :param bucket_name: The Google Cloud Storage bucket where the blob_name is. :param object_name: The name of the object to check in the Google cloud storage bucket_name. """ - self.log.info("Checking the file size of object: %s in bucket_name: %s", object_name, bucket_name) client = self.get_conn() bucket = client.bucket(bucket_name) blob = bucket.get_blob(blob_name=object_name) + + if blob is None: + raise ValueError(f"Object ({object_name}) not found in Bucket ({bucket_name})") + + return blob + + def get_size(self, bucket_name: str, object_name: str) -> int: + """ + Get the size of a file in Google Cloud Storage. + + :param bucket_name: The Google Cloud Storage bucket where the blob_name is. + :param object_name: The name of the object to check in the Google + cloud storage bucket_name. + + """ + self.log.info("Checking the file size of object: %s in bucket_name: %s", object_name, bucket_name) + blob = self._get_blob(bucket_name, object_name) blob_size = blob.size self.log.info("The file size of %s is %s bytes.", object_name, blob_size) return blob_size @@ -987,9 +1000,7 @@ def get_crc32c(self, bucket_name: str, object_name: str): object_name, bucket_name, ) - client = self.get_conn() - bucket = client.bucket(bucket_name) - blob = bucket.get_blob(blob_name=object_name) + blob = self._get_blob(bucket_name, object_name) blob_crc32c = blob.crc32c self.log.info("The crc32c checksum of %s is %s", object_name, blob_crc32c) return blob_crc32c @@ -1003,9 +1014,7 @@ def get_md5hash(self, bucket_name: str, object_name: str) -> str: storage bucket_name. """ self.log.info("Retrieving the MD5 hash of object: %s in bucket: %s", object_name, bucket_name) - client = self.get_conn() - bucket = client.bucket(bucket_name) - blob = bucket.get_blob(blob_name=object_name) + blob = self._get_blob(bucket_name, object_name) blob_md5hash = blob.md5_hash self.log.info("The md5Hash of %s is %s", object_name, blob_md5hash) return blob_md5hash @@ -1019,11 +1028,7 @@ def get_metadata(self, bucket_name: str, object_name: str) -> dict | None: :return: The metadata associated with the object """ self.log.info("Retrieving the metadata dict of object (%s) in bucket (%s)", object_name, bucket_name) - client = self.get_conn() - bucket = client.bucket(bucket_name) - blob = bucket.get_blob(blob_name=object_name) - if blob is None: - raise ValueError("Object (%s) not found in bucket (%s)", object_name, bucket_name) + blob = self._get_blob(bucket_name, object_name) blob_metadata = blob.metadata if blob_metadata: self.log.info("Retrieved metadata of object (%s) with %s fields", object_name, len(blob_metadata)) diff --git a/newsfragments/42439.bugfix.rst b/newsfragments/42439.bugfix.rst new file mode 100644 index 0000000000000..4ecd73eb636dd --- /dev/null +++ b/newsfragments/42439.bugfix.rst @@ -0,0 +1 @@ +Refactor ``bucket.get_blob`` calls in ``GCSHook`` to handle validation for non-existent objects. diff --git a/tests/providers/google/cloud/hooks/test_gcs.py b/tests/providers/google/cloud/hooks/test_gcs.py index 5dee8cbf8cdce..5d2735834a958 100644 --- a/tests/providers/google/cloud/hooks/test_gcs.py +++ b/tests/providers/google/cloud/hooks/test_gcs.py @@ -523,57 +523,57 @@ def test_delete_nonexisting_bucket(self, mock_service, caplog): mock_service.return_value.bucket.return_value.delete.assert_called_once() assert "Bucket test bucket not exist" in caplog.text - @mock.patch(GCS_STRING.format("GCSHook.get_conn")) + @mock.patch(GCS_STRING.format("GCSHook._get_blob")) def test_object_get_size(self, mock_service): test_bucket = "test_bucket" test_object = "test_object" returned_file_size = 1200 - bucket_method = mock_service.return_value.bucket - get_blob_method = bucket_method.return_value.get_blob - get_blob_method.return_value.size = returned_file_size + mock_blob = MagicMock() + mock_blob.size = returned_file_size + mock_service.return_value = mock_blob response = self.gcs_hook.get_size(bucket_name=test_bucket, object_name=test_object) assert response == returned_file_size - @mock.patch(GCS_STRING.format("GCSHook.get_conn")) + @mock.patch(GCS_STRING.format("GCSHook._get_blob")) def test_object_get_crc32c(self, mock_service): test_bucket = "test_bucket" test_object = "test_object" returned_file_crc32c = "xgdNfQ==" - bucket_method = mock_service.return_value.bucket - get_blob_method = bucket_method.return_value.get_blob - get_blob_method.return_value.crc32c = returned_file_crc32c + mock_blob = MagicMock() + mock_blob.crc32c = returned_file_crc32c + mock_service.return_value = mock_blob response = self.gcs_hook.get_crc32c(bucket_name=test_bucket, object_name=test_object) assert response == returned_file_crc32c - @mock.patch(GCS_STRING.format("GCSHook.get_conn")) + @mock.patch(GCS_STRING.format("GCSHook._get_blob")) def test_object_get_md5hash(self, mock_service): test_bucket = "test_bucket" test_object = "test_object" returned_file_md5hash = "leYUJBUWrRtks1UeUFONJQ==" - bucket_method = mock_service.return_value.bucket - get_blob_method = bucket_method.return_value.get_blob - get_blob_method.return_value.md5_hash = returned_file_md5hash + mock_blob = MagicMock() + mock_blob.md5_hash = returned_file_md5hash + mock_service.return_value = mock_blob response = self.gcs_hook.get_md5hash(bucket_name=test_bucket, object_name=test_object) assert response == returned_file_md5hash - @mock.patch(GCS_STRING.format("GCSHook.get_conn")) + @mock.patch(GCS_STRING.format("GCSHook._get_blob")) def test_object_get_metadata(self, mock_service): test_bucket = "test_bucket" test_object = "test_object" returned_file_metadata = {"test_metadata_key": "test_metadata_val"} - bucket_method = mock_service.return_value.bucket - get_blob_method = bucket_method.return_value.get_blob - get_blob_method.return_value.metadata = returned_file_metadata + mock_blob = MagicMock() + mock_blob.metadata = returned_file_metadata + mock_service.return_value = mock_blob response = self.gcs_hook.get_metadata(bucket_name=test_bucket, object_name=test_object) @@ -588,7 +588,7 @@ def test_nonexisting_object_get_metadata(self, mock_service): get_blob_method = bucket_method.return_value.get_blob get_blob_method.return_value = None - with pytest.raises(ValueError, match=r"Object \((.*?)\) not found in bucket \((.*?)\)"): + with pytest.raises(ValueError, match=r"Object \((.*?)\) not found in Bucket \((.*?)\)"): self.gcs_hook.get_metadata(bucket_name=test_bucket, object_name=test_object) @mock.patch("google.cloud.storage.Bucket") @@ -1509,6 +1509,32 @@ def test_should_not_overwrite_when_overwrite_is_disabled( mock_rewrite.assert_not_called() mock_copy.assert_not_called() + @mock.patch(GCS_STRING.format("GCSHook.get_conn")) + def test_object_get_blob(self, mock_service): + test_bucket = "test_bucket" + test_object = "test_object" + mock_blob = mock.MagicMock() + + bucket_method = mock_service.return_value.bucket + get_blob_method = bucket_method.return_value.get_blob + get_blob_method.return_value = mock_blob + + response = self.gcs_hook._get_blob(bucket_name=test_bucket, object_name=test_object) + + assert response == mock_blob + + @mock.patch(GCS_STRING.format("GCSHook.get_conn")) + def test_nonexisting_object_get_blob(self, mock_service): + test_bucket = "test_bucket" + test_object = "test_object" + + bucket_method = mock_service.return_value.bucket + get_blob_method = bucket_method.return_value.get_blob + get_blob_method.return_value = None + + with pytest.raises(ValueError, match=r"Object \((.*?)\) not found in Bucket \((.*?)\)"): + self.gcs_hook._get_blob(bucket_name=test_bucket, object_name=test_object) + def _create_blob( self, name: str, From 4c44959943fec00663776bf9d86ea25280f06ead Mon Sep 17 00:00:00 2001 From: olegkachur-e Date: Fri, 27 Sep 2024 10:20:21 +0200 Subject: [PATCH 049/802] Update docs regarding the AutoMLText deprecation (#42415) - Suggest ways to replace old functionality. - Move all related info to vertex_ai doc. Co-authored-by: Oleg Kachur --- .../operators/cloud/automl.rst | 5 ----- .../operators/cloud/vertex_ai.rst | 10 ++++++---- 2 files changed, 6 insertions(+), 9 deletions(-) diff --git a/docs/apache-airflow-providers-google/operators/cloud/automl.rst b/docs/apache-airflow-providers-google/operators/cloud/automl.rst index fdedb46ea8bf2..f8f35aa993fa0 100644 --- a/docs/apache-airflow-providers-google/operators/cloud/automl.rst +++ b/docs/apache-airflow-providers-google/operators/cloud/automl.rst @@ -109,11 +109,6 @@ available on the Vertex AI platform. Please use :class:`~airflow.providers.google.cloud.operators.vertex_ai.auto_ml.CreateAutoMLImageTrainingJobOperator` or :class:`~airflow.providers.google.cloud.operators.vertex_ai.auto_ml.CreateAutoMLVideoTrainingJobOperator`. -The Vertex AutoMLText API for model training is deprecated on September 15, 2024 and the other part will be deprecated -on June 15, 2025. -Please consider using fine tuning with Gemini model - -https://cloud.google.com/vertex-ai/generative-ai/docs/models/gemini-tuning. - You can find example on how to use VertexAI operators for AutoML Vision classification here: .. exampleinclude:: /../../tests/system/providers/google/cloud/automl/example_automl_vision_classification.py diff --git a/docs/apache-airflow-providers-google/operators/cloud/vertex_ai.rst b/docs/apache-airflow-providers-google/operators/cloud/vertex_ai.rst index 9a03ed2924107..8fb76cd80fdee 100644 --- a/docs/apache-airflow-providers-google/operators/cloud/vertex_ai.rst +++ b/docs/apache-airflow-providers-google/operators/cloud/vertex_ai.rst @@ -260,11 +260,13 @@ put dataset id to ``dataset_id`` parameter in operator. How to run AutoML Text Training Job :class:`~airflow.providers.google.cloud.operators.vertex_ai.auto_ml.CreateAutoMLTextTrainingJobOperator` -Operator is deprecated, please use -:class:`~airflow.providers.google.cloud.operators.vertex_ai.generative_model.SupervisedFineTuningTrainOperator` over -the Gemini model. -More info: https://cloud.google.com/vertex-ai/generative-ai/docs/models/gemini-tuning#tuning-gemini +Operator is deprecated, the non-training (existing legacy models) AutoMLText API will be deprecated +on June 15, 2025. +There are 2 options for text classification, extraction, and sentiment analysis tasks replacement: +- Prompts with pre-trained Gemini model, using :class:`~airflow.providers.google.cloud.operators.vertex_ai.generative_model.TextGenerationModelPredictOperator`. +- Fine tuning over Gemini model, For more tailored results, using :class:`~airflow.providers.google.cloud.operators.vertex_ai.generative_model.SupervisedFineTuningTrainOperator`. +Please visit the https://cloud.google.com/vertex-ai/generative-ai/docs/models/gemini-tuning for more details. How to run AutoML Video Training Job :class:`~airflow.providers.google.cloud.operators.vertex_ai.auto_ml.CreateAutoMLVideoTrainingJobOperator` From 139ae9d4afc3074bdd2971160c7d46126311ace8 Mon Sep 17 00:00:00 2001 From: paolomoriello <61800102+paolo-moriello@users.noreply.github.com> Date: Fri, 27 Sep 2024 10:20:57 +0200 Subject: [PATCH 050/802] KubernetesPodOperator never stops if credentials are refreshed (#42361) * Never stop retrying 401s in k8s pod operator * Try reading pod after refreshing credentials * Never stop retrying until credentials are refreshed * Linting --------- Co-authored-by: pmoriello --- .../cncf/kubernetes/operators/pod.py | 12 ++++++-- .../cncf/kubernetes/operators/test_pod.py | 29 +++++++++++++++++-- 2 files changed, 35 insertions(+), 6 deletions(-) diff --git a/airflow/providers/cncf/kubernetes/operators/pod.py b/airflow/providers/cncf/kubernetes/operators/pod.py index 709583d78e2ed..5b9e57ec01743 100644 --- a/airflow/providers/cncf/kubernetes/operators/pod.py +++ b/airflow/providers/cncf/kubernetes/operators/pod.py @@ -80,7 +80,6 @@ PodNotFoundException, PodOperatorHookProtocol, PodPhase, - check_exception_is_kubernetes_api_unauthorized, container_is_succeeded, get_container_termination_message, ) @@ -113,6 +112,10 @@ class PodReattachFailure(AirflowException): """When we expect to be able to find a pod but cannot.""" +class PodCredentialsExpiredFailure(AirflowException): + """When pod fails to refresh credentials.""" + + class KubernetesPodOperator(BaseOperator): """ Execute a task in a Kubernetes Pod. @@ -652,9 +655,8 @@ def execute_sync(self, context: Context): return result @tenacity.retry( - stop=tenacity.stop_after_attempt(3), wait=tenacity.wait_exponential(max=15), - retry=tenacity.retry_if_exception(lambda exc: check_exception_is_kubernetes_api_unauthorized(exc)), + retry=tenacity.retry_if_exception_type(PodCredentialsExpiredFailure), reraise=True, ) def await_pod_completion(self, pod: k8s.V1Pod): @@ -675,6 +677,10 @@ def await_pod_completion(self, pod: k8s.V1Pod): "Failed to check container status due to permission error. Refreshing credentials and retrying." ) self._refresh_cached_properties() + self.pod_manager.read_pod( + pod=pod + ) # attempt using refreshed credentials, raises if still invalid + raise PodCredentialsExpiredFailure("Kubernetes credentials expired, retrying after refresh.") raise exc def _refresh_cached_properties(self): diff --git a/tests/providers/cncf/kubernetes/operators/test_pod.py b/tests/providers/cncf/kubernetes/operators/test_pod.py index 1ca0d9851b18c..be8279bcabd9a 100644 --- a/tests/providers/cncf/kubernetes/operators/test_pod.py +++ b/tests/providers/cncf/kubernetes/operators/test_pod.py @@ -1629,8 +1629,9 @@ def test_execute_async_callbacks(self): @pytest.mark.parametrize("get_logs", [True, False]) @patch(f"{POD_MANAGER_CLASS}.fetch_requested_container_logs") @patch(f"{POD_MANAGER_CLASS}.await_container_completion") + @patch(f"{POD_MANAGER_CLASS}.read_pod") def test_await_container_completion_refreshes_properties_on_exception( - self, mock_await_container_completion, fetch_requested_container_logs, get_logs + self, mock_read_pod, mock_await_container_completion, fetch_requested_container_logs, get_logs ): k = KubernetesPodOperator(task_id="task", get_logs=get_logs) pod = self.run_pod(k) @@ -1655,6 +1656,28 @@ def test_await_container_completion_refreshes_properties_on_exception( mock_await_container_completion.assert_has_calls( [mock.call(pod=pod, container_name=k.base_container_name)] * 3 ) + mock_read_pod.assert_called() + assert client != k.client + assert hook != k.hook + assert pod_manager != k.pod_manager + + @patch(f"{POD_MANAGER_CLASS}.await_container_completion") + @patch(f"{POD_MANAGER_CLASS}.read_pod") + def test_await_container_completion_raises_unauthorized_if_credentials_still_invalid_after_refresh( + self, mock_read_pod, mock_await_container_completion + ): + k = KubernetesPodOperator(task_id="task", get_logs=False) + pod = self.run_pod(k) + client, hook, pod_manager = k.client, k.hook, k.pod_manager + + mock_await_container_completion.side_effect = [ApiException(status=401)] + mock_read_pod.side_effect = [ApiException(status=401)] + + with pytest.raises(ApiException): + k.await_pod_completion(pod) + + mock_read_pod.assert_called() + # assert cache was refreshed assert client != k.client assert hook != k.hook assert pod_manager != k.pod_manager @@ -1663,7 +1686,7 @@ def test_await_container_completion_refreshes_properties_on_exception( "side_effect, exception_type, expect_exc", [ ([ApiException(401), mock.DEFAULT], ApiException, True), # works after one 401 - ([ApiException(401)] * 10, ApiException, False), # exc after 3 retries on 401 + ([ApiException(401)] * 3 + [mock.DEFAULT], ApiException, True), # works after 3 retries ([ApiException(402)], ApiException, False), # exc on non-401 ([ApiException(500)], ApiException, False), # exc on non-401 ([Exception], Exception, False), # exc on different exception @@ -1684,7 +1707,7 @@ def test_await_container_completion_retries_on_specific_exception( else: with pytest.raises(exception_type): k.await_pod_completion(pod) - expected_call_count = min(len(side_effect), 3) # retry max 3 times + expected_call_count = len(side_effect) mock_await_container_completion.assert_has_calls( [mock.call(pod=pod, container_name=k.base_container_name)] * expected_call_count ) From 8de1dca01c01a6ba71c9004326fdf0f6f55c804a Mon Sep 17 00:00:00 2001 From: Jens Scheffler <95105677+jscheffl@users.noreply.github.com> Date: Fri, 27 Sep 2024 10:21:29 +0200 Subject: [PATCH 051/802] Handle ENTER key correctly in trigger form and allow manual JSON (#42525) --- airflow/www/static/js/trigger.js | 30 +++++++++++++++++++++- airflow/www/templates/airflow/trigger.html | 2 +- 2 files changed, 30 insertions(+), 2 deletions(-) diff --git a/airflow/www/static/js/trigger.js b/airflow/www/static/js/trigger.js index 2ded629240146..7a3444f460f41 100644 --- a/airflow/www/static/js/trigger.js +++ b/airflow/www/static/js/trigger.js @@ -96,6 +96,33 @@ function updateJSONconf() { jsonForm.setValue(JSON.stringify(params, null, 4)); } +/** + * If the user hits ENTER key inside an input, ensure JSON data is updated. + */ +function handleEnter() { + updateJSONconf(); + // somehow following is needed to enforce form is submitted correctly from CodeMirror + document.getElementById("json").value = jsonForm.getValue(); +} + +/** + * Track user changes in input fields, ensure JSON is updated when user presses enter + * See https://github.com/apache/airflow/issues/42157 + */ +function enterInputField() { + const form = document.getElementById("trigger_form"); + form.addEventListener("submit", handleEnter); +} + +/** + * Stop tracking user changes in input fields + */ +function leaveInputField() { + const form = document.getElementById("trigger_form"); + form.removeEventListener("submit", handleEnter); + updateJSONconf(); +} + /** * Initialize the form during load of the web page */ @@ -148,7 +175,8 @@ function initForm() { } else if (elements[i].type === "checkbox") { elements[i].addEventListener("change", updateJSONconf); } else { - elements[i].addEventListener("blur", updateJSONconf); + elements[i].addEventListener("focus", enterInputField); + elements[i].addEventListener("blur", leaveInputField); } } } diff --git a/airflow/www/templates/airflow/trigger.html b/airflow/www/templates/airflow/trigger.html index 71d09e79076f1..7cdcd337beddf 100644 --- a/airflow/www/templates/airflow/trigger.html +++ b/airflow/www/templates/airflow/trigger.html @@ -163,7 +163,7 @@

{{ dag.description[0:150] + '…' if dag.description and dag.description|length > 150 else dag.description|default('', true) }}

{{ dag_docs(doc_md, False) }} - + From c37e23dc58dd45b8fe36d2a393eb7934d2a300f6 Mon Sep 17 00:00:00 2001 From: olegkachur-e Date: Fri, 27 Sep 2024 10:30:01 +0200 Subject: [PATCH 052/802] Deprecate AutoMLBatchPredictOperator and refactor AutoMl system tests (#42260) - Remove example_automl_model.py as obsolete. - Move example hooks to example_automl_translation.py - Update documentation - Deprecate AutoMLPredictBatchOperator completely, as it cannot be used with translation model, previous usage has been already deprecated. Co-authored-by: Oleg Kachur --- .../google/cloud/operators/automl.py | 8 +- .../operators/cloud/automl.rst | 9 +- tests/always/test_project_structure.py | 1 + .../google/cloud/links/test_translate.py | 22 +- .../google/cloud/operators/test_automl.py | 70 +++--- .../cloud/automl/example_automl_model.py | 225 ------------------ .../automl/example_automl_translation.py | 27 ++- 7 files changed, 84 insertions(+), 278 deletions(-) delete mode 100644 tests/system/providers/google/cloud/automl/example_automl_model.py diff --git a/airflow/providers/google/cloud/operators/automl.py b/airflow/providers/google/cloud/operators/automl.py index 7cbc610444d07..8b3a86ba5250c 100644 --- a/airflow/providers/google/cloud/operators/automl.py +++ b/airflow/providers/google/cloud/operators/automl.py @@ -34,7 +34,7 @@ TableSpec, ) -from airflow.exceptions import AirflowException +from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning from airflow.providers.google.cloud.hooks.automl import CloudAutoMLHook from airflow.providers.google.cloud.hooks.vertex_ai.prediction_service import PredictionServiceHook from airflow.providers.google.cloud.links.translate import ( @@ -45,6 +45,7 @@ TranslationLegacyModelTrainLink, ) from airflow.providers.google.cloud.operators.cloud_base import GoogleCloudBaseOperator +from airflow.providers.google.common.deprecated import deprecated from airflow.providers.google.common.hooks.base_google import PROVIDE_PROJECT_ID if TYPE_CHECKING: @@ -338,6 +339,11 @@ def execute(self, context: Context): return PredictResponse.to_dict(result) +@deprecated( + planned_removal_date="January 01, 2025", + use_instead="airflow.providers.google.cloud.operators.vertex_ai.batch_prediction_job", + category=AirflowProviderDeprecationWarning, +) class AutoMLBatchPredictOperator(GoogleCloudBaseOperator): """ Perform a batch prediction on Google Cloud AutoML. diff --git a/docs/apache-airflow-providers-google/operators/cloud/automl.rst b/docs/apache-airflow-providers-google/operators/cloud/automl.rst index f8f35aa993fa0..8a92f49ac34f1 100644 --- a/docs/apache-airflow-providers-google/operators/cloud/automl.rst +++ b/docs/apache-airflow-providers-google/operators/cloud/automl.rst @@ -131,7 +131,7 @@ datasets. To create and import data to the dataset please use and :class:`~airflow.providers.google.cloud.operators.vertex_ai.dataset.ImportDataOperator` -.. exampleinclude:: /../../tests/system/providers/google/cloud/automl/example_automl_model.py +.. exampleinclude:: /../../tests/system/providers/google/cloud/automl/example_automl_translation.py :language: python :dedent: 4 :start-after: [START howto_operator_automl_create_model] @@ -190,17 +190,12 @@ To obtain predictions from Google Cloud AutoML model you can use :class:`~airflow.providers.google.cloud.operators.automl.AutoMLBatchPredictOperator`. In the first case the model must be deployed. -.. exampleinclude:: /../../tests/system/providers/google/cloud/automl/example_automl_model.py +.. exampleinclude:: /../../tests/system/providers/google/cloud/automl/example_automl_translation.py :language: python :dedent: 4 :start-after: [START howto_operator_prediction] :end-before: [END howto_operator_prediction] -.. exampleinclude:: /../../tests/system/providers/google/cloud/automl/example_automl_model.py - :language: python - :dedent: 4 - :start-after: [START howto_operator_batch_prediction] - :end-before: [END howto_operator_batch_prediction] Th :class:`~airflow.providers.google.cloud.operators.automl.AutoMLBatchPredictOperator` deprecated for tables, video intelligence, vision and natural language is deprecated and will be removed after 31.03.2024. Please use diff --git a/tests/always/test_project_structure.py b/tests/always/test_project_structure.py index b6ef9a89f9669..c387f6173ca2f 100644 --- a/tests/always/test_project_structure.py +++ b/tests/always/test_project_structure.py @@ -430,6 +430,7 @@ class TestGoogleProviderProjectStructure(ExampleCoverageTest, AssetsCoverageTest "airflow.providers.google.cloud.operators.vertex_ai.generative_model.GenerateTextEmbeddingsOperator", "airflow.providers.google.cloud.operators.vertex_ai.generative_model.PromptMultimodalModelOperator", "airflow.providers.google.cloud.operators.vertex_ai.generative_model.PromptMultimodalModelWithMediaOperator", + "airflow.providers.google.cloud.operators.automl.AutoMLBatchPredictOperator", } ASSETS_NOT_REQUIRED = { diff --git a/tests/providers/google/cloud/links/test_translate.py b/tests/providers/google/cloud/links/test_translate.py index 82547907a83f6..bb2f3f0811837 100644 --- a/tests/providers/google/cloud/links/test_translate.py +++ b/tests/providers/google/cloud/links/test_translate.py @@ -24,6 +24,7 @@ from google.cloud.automl_v1beta1 import Model +from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.providers.google.cloud.links.translate import ( TRANSLATION_BASE_LINK, TranslationDatasetListLink, @@ -146,16 +147,17 @@ def test_get_link(self, create_task_instance_of_operator, session): f"predict;modelId={MODEL}?project={GCP_PROJECT_ID}" ) link = TranslationLegacyModelPredictLink() - ti = create_task_instance_of_operator( - AutoMLBatchPredictOperator, - dag_id="test_legacy_model_predict_link_dag", - task_id="test_legacy_model_predict_link_task", - model_id=MODEL, - project_id=GCP_PROJECT_ID, - location=GCP_LOCATION, - input_config="input_config", - output_config="input_config", - ) + with pytest.warns(AirflowProviderDeprecationWarning): + ti = create_task_instance_of_operator( + AutoMLBatchPredictOperator, + dag_id="test_legacy_model_predict_link_dag", + task_id="test_legacy_model_predict_link_task", + model_id=MODEL, + project_id=GCP_PROJECT_ID, + location=GCP_LOCATION, + input_config="input_config", + output_config="input_config", + ) ti.task.model = Model(dataset_id=DATASET, display_name=MODEL) session.add(ti) session.commit() diff --git a/tests/providers/google/cloud/operators/test_automl.py b/tests/providers/google/cloud/operators/test_automl.py index fbe17537535dc..00cd4d396d832 100644 --- a/tests/providers/google/cloud/operators/test_automl.py +++ b/tests/providers/google/cloud/operators/test_automl.py @@ -28,7 +28,7 @@ from google.api_core.gapic_v1.method import DEFAULT from google.cloud.automl_v1beta1 import BatchPredictResult, Dataset, Model, PredictResponse -from airflow.exceptions import AirflowException +from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning from airflow.providers.google.cloud.hooks.automl import CloudAutoMLHook from airflow.providers.google.cloud.hooks.vertex_ai.prediction_service import PredictionServiceHook from airflow.providers.google.cloud.operators.automl import ( @@ -148,15 +148,16 @@ def test_execute(self, mock_hook, mock_link_persist): mock_hook.return_value.extract_object_id = extract_object_id mock_hook.return_value.wait_for_operation.return_value = BatchPredictResult() mock_context = {"ti": mock.MagicMock()} - op = AutoMLBatchPredictOperator( - model_id=MODEL_ID, - location=GCP_LOCATION, - project_id=GCP_PROJECT_ID, - input_config=INPUT_CONFIG, - output_config=OUTPUT_CONFIG, - task_id=TASK_ID, - prediction_params={}, - ) + with pytest.warns(AirflowProviderDeprecationWarning): + op = AutoMLBatchPredictOperator( + model_id=MODEL_ID, + location=GCP_LOCATION, + project_id=GCP_PROJECT_ID, + input_config=INPUT_CONFIG, + output_config=OUTPUT_CONFIG, + task_id=TASK_ID, + prediction_params={}, + ) op.execute(context=mock_context) mock_hook.return_value.batch_predict.assert_called_once_with( input_config=INPUT_CONFIG, @@ -182,16 +183,16 @@ def test_execute_deprecated(self, mock_hook): del returned_model.translation_model_metadata mock_hook.return_value.get_model.return_value = returned_model mock_hook.return_value.extract_object_id = extract_object_id - - op = AutoMLBatchPredictOperator( - model_id=MODEL_ID, - location=GCP_LOCATION, - project_id=GCP_PROJECT_ID, - input_config=INPUT_CONFIG, - output_config=OUTPUT_CONFIG, - task_id=TASK_ID, - prediction_params={}, - ) + with pytest.warns(AirflowProviderDeprecationWarning): + op = AutoMLBatchPredictOperator( + model_id=MODEL_ID, + location=GCP_LOCATION, + project_id=GCP_PROJECT_ID, + input_config=INPUT_CONFIG, + output_config=OUTPUT_CONFIG, + task_id=TASK_ID, + prediction_params={}, + ) expected_exception_str = ( "AutoMLBatchPredictOperator for text, image, and video prediction has been " "deprecated and no longer available" @@ -210,20 +211,21 @@ def test_execute_deprecated(self, mock_hook): @pytest.mark.db_test def test_templating(self, create_task_instance_of_operator, session): - ti = create_task_instance_of_operator( - AutoMLBatchPredictOperator, - # Templated fields - model_id="{{ 'model' }}", - input_config="{{ 'input-config' }}", - output_config="{{ 'output-config' }}", - location="{{ 'location' }}", - project_id="{{ 'project-id' }}", - impersonation_chain="{{ 'impersonation-chain' }}", - # Other parameters - dag_id="test_template_body_templating_dag", - task_id="test_template_body_templating_task", - execution_date=timezone.datetime(2024, 2, 1, tzinfo=timezone.utc), - ) + with pytest.warns(AirflowProviderDeprecationWarning): + ti = create_task_instance_of_operator( + AutoMLBatchPredictOperator, + # Templated fields + model_id="{{ 'model' }}", + input_config="{{ 'input-config' }}", + output_config="{{ 'output-config' }}", + location="{{ 'location' }}", + project_id="{{ 'project-id' }}", + impersonation_chain="{{ 'impersonation-chain' }}", + # Other parameters + dag_id="test_template_body_templating_dag", + task_id="test_template_body_templating_task", + execution_date=timezone.datetime(2024, 2, 1, tzinfo=timezone.utc), + ) session.add(ti) session.commit() ti.render_templates() diff --git a/tests/system/providers/google/cloud/automl/example_automl_model.py b/tests/system/providers/google/cloud/automl/example_automl_model.py deleted file mode 100644 index 1603595600023..0000000000000 --- a/tests/system/providers/google/cloud/automl/example_automl_model.py +++ /dev/null @@ -1,225 +0,0 @@ -# -# 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. - -"""Example Airflow DAG for Google AutoML service testing model operations.""" - -from __future__ import annotations - -import os -from datetime import datetime - -from google.protobuf.struct_pb2 import Value - -from airflow.models.dag import DAG -from airflow.providers.google.cloud.operators.automl import ( - AutoMLBatchPredictOperator, - AutoMLCreateDatasetOperator, - AutoMLDeleteDatasetOperator, - AutoMLDeleteModelOperator, - AutoMLGetModelOperator, - AutoMLImportDataOperator, - AutoMLPredictOperator, - AutoMLTrainModelOperator, -) -from airflow.providers.google.cloud.operators.gcs import ( - GCSCreateBucketOperator, - GCSDeleteBucketOperator, - GCSSynchronizeBucketsOperator, -) -from airflow.utils.trigger_rule import TriggerRule - -ENV_ID = os.environ.get("SYSTEM_TESTS_ENV_ID", "default") -DAG_ID = "automl_model" -GCP_PROJECT_ID = os.environ.get("SYSTEM_TESTS_GCP_PROJECT", "default") - -GCP_AUTOML_LOCATION = "us-central1" - -DATA_SAMPLE_GCS_BUCKET_NAME = f"bucket_{DAG_ID}_{ENV_ID}".replace("_", "-") -RESOURCE_DATA_BUCKET = "airflow-system-tests-resources" - -DATASET_NAME = f"ds_{DAG_ID}_{ENV_ID}".replace("-", "_") -DATASET = { - "display_name": DATASET_NAME, - "tables_dataset_metadata": {"target_column_spec_id": ""}, -} -AUTOML_DATASET_BUCKET = f"gs://{DATA_SAMPLE_GCS_BUCKET_NAME}/automl/bank-marketing-split.csv" -IMPORT_INPUT_CONFIG = {"gcs_source": {"input_uris": [AUTOML_DATASET_BUCKET]}} -IMPORT_OUTPUT_CONFIG = { - "gcs_destination": {"output_uri_prefix": f"gs://{DATA_SAMPLE_GCS_BUCKET_NAME}/automl"} -} - -# change the name here -MODEL_NAME = f"md_{DAG_ID}_{ENV_ID}".replace("-", "_") -MODEL = { - "display_name": MODEL_NAME, - "tables_model_metadata": {"train_budget_milli_node_hours": 1000}, -} - -PREDICT_VALUES = [ - Value(string_value="TRAINING"), - Value(string_value="51"), - Value(string_value="blue-collar"), - Value(string_value="married"), - Value(string_value="primary"), - Value(string_value="no"), - Value(string_value="620"), - Value(string_value="yes"), - Value(string_value="yes"), - Value(string_value="cellular"), - Value(string_value="29"), - Value(string_value="jul"), - Value(string_value="88"), - Value(string_value="10"), - Value(string_value="-1"), - Value(string_value="0"), - Value(string_value="unknown"), -] - - -with DAG( - dag_id=DAG_ID, - schedule="@once", - start_date=datetime(2021, 1, 1), - catchup=False, - tags=["example", "automl", "model"], -) as dag: - create_bucket = GCSCreateBucketOperator( - task_id="create_bucket", - bucket_name=DATA_SAMPLE_GCS_BUCKET_NAME, - storage_class="REGIONAL", - location=GCP_AUTOML_LOCATION, - ) - - move_dataset_file = GCSSynchronizeBucketsOperator( - task_id="move_data_to_bucket", - source_bucket=RESOURCE_DATA_BUCKET, - source_object="automl/datasets/model", - destination_bucket=DATA_SAMPLE_GCS_BUCKET_NAME, - destination_object="automl", - recursive=True, - ) - - create_dataset = AutoMLCreateDatasetOperator( - task_id="create_dataset", - dataset=DATASET, - location=GCP_AUTOML_LOCATION, - project_id=GCP_PROJECT_ID, - ) - - dataset_id = create_dataset.output["dataset_id"] - MODEL["dataset_id"] = dataset_id - import_dataset = AutoMLImportDataOperator( - task_id="import_dataset", - dataset_id=dataset_id, - location=GCP_AUTOML_LOCATION, - input_config=IMPORT_INPUT_CONFIG, - ) - - # [START howto_operator_automl_create_model] - create_model = AutoMLTrainModelOperator( - task_id="create_model", - model=MODEL, - location=GCP_AUTOML_LOCATION, - project_id=GCP_PROJECT_ID, - ) - model_id = create_model.output["model_id"] - # [END howto_operator_automl_create_model] - - # [START howto_operator_get_model] - get_model = AutoMLGetModelOperator( - task_id="get_model", - model_id=model_id, - location=GCP_AUTOML_LOCATION, - project_id=GCP_PROJECT_ID, - ) - # [END howto_operator_get_model] - - # [START howto_operator_prediction] - predict_task = AutoMLPredictOperator( - task_id="predict_task", - model_id=model_id, - payload={ - "row": { - "values": PREDICT_VALUES, - } - }, - location=GCP_AUTOML_LOCATION, - project_id=GCP_PROJECT_ID, - ) - # [END howto_operator_prediction] - - # [START howto_operator_batch_prediction] - batch_predict_task = AutoMLBatchPredictOperator( - task_id="batch_predict_task", - model_id=model_id, - input_config=IMPORT_INPUT_CONFIG, - output_config=IMPORT_OUTPUT_CONFIG, - location=GCP_AUTOML_LOCATION, - project_id=GCP_PROJECT_ID, - ) - # [END howto_operator_batch_prediction] - - # [START howto_operator_automl_delete_model] - delete_model = AutoMLDeleteModelOperator( - task_id="delete_model", - model_id=model_id, - location=GCP_AUTOML_LOCATION, - project_id=GCP_PROJECT_ID, - ) - # [END howto_operator_automl_delete_model] - - delete_dataset = AutoMLDeleteDatasetOperator( - task_id="delete_dataset", - dataset_id=dataset_id, - location=GCP_AUTOML_LOCATION, - project_id=GCP_PROJECT_ID, - trigger_rule=TriggerRule.ALL_DONE, - ) - - delete_bucket = GCSDeleteBucketOperator( - task_id="delete_bucket", - bucket_name=DATA_SAMPLE_GCS_BUCKET_NAME, - trigger_rule=TriggerRule.ALL_DONE, - ) - - ( - # TEST SETUP - [create_bucket >> move_dataset_file, create_dataset] - >> import_dataset - # TEST BODY - >> create_model - >> get_model - >> predict_task - >> batch_predict_task - # TEST TEARDOWN - >> delete_model - >> delete_dataset - >> delete_bucket - ) - - from tests.system.utils.watcher import watcher - - # This test needs watcher in order to properly mark success/failure - # when "tearDown" task with trigger rule is part of the DAG - list(dag.tasks) >> watcher() - - -from tests.system.utils import get_test_run # noqa: E402 - -# Needed to run the example DAG with pytest (see: tests/system/README.md#run_via_pytest) -test_run = get_test_run(dag) diff --git a/tests/system/providers/google/cloud/automl/example_automl_translation.py b/tests/system/providers/google/cloud/automl/example_automl_translation.py index 41acdf764b8ec..cda70693fb57c 100644 --- a/tests/system/providers/google/cloud/automl/example_automl_translation.py +++ b/tests/system/providers/google/cloud/automl/example_automl_translation.py @@ -34,7 +34,9 @@ AutoMLCreateDatasetOperator, AutoMLDeleteDatasetOperator, AutoMLDeleteModelOperator, + AutoMLGetModelOperator, AutoMLImportDataOperator, + AutoMLPredictOperator, AutoMLTrainModelOperator, ) from airflow.providers.google.cloud.operators.gcs import GCSCreateBucketOperator, GCSDeleteBucketOperator @@ -126,10 +128,31 @@ def upload_csv_file_to_gcs(): ) MODEL["dataset_id"] = dataset_id - + # [START howto_operator_automl_create_model] create_model = AutoMLTrainModelOperator(task_id="create_model", model=MODEL, location=GCP_AUTOML_LOCATION) + # [END howto_operator_automl_create_model] model_id = cast(str, XComArg(create_model, key="model_id")) + # [START howto_operator_get_model] + get_model = AutoMLGetModelOperator( + task_id="get_model", + model_id=model_id, + location=GCP_AUTOML_LOCATION, + project_id=GCP_PROJECT_ID, + ) + # [END howto_operator_get_model] + + # [START howto_operator_prediction] + TRANSLATION_STR = "A Dog walks down the street" + predict_task = AutoMLPredictOperator( + task_id="predict_task", + model_id=model_id, + payload={"text_snippet": {"content": TRANSLATION_STR}}, + location=GCP_AUTOML_LOCATION, + project_id=GCP_PROJECT_ID, + ) + # [END howto_operator_prediction] + delete_model = AutoMLDeleteModelOperator( task_id="delete_model", model_id=model_id, @@ -157,6 +180,8 @@ def upload_csv_file_to_gcs(): >> create_dataset >> import_dataset >> create_model + >> get_model + >> predict_task # TEST TEARDOWN >> delete_dataset >> delete_model From 3f01388b74cca77ec495efedb31317a4e24e8aaa Mon Sep 17 00:00:00 2001 From: olegkachur-e Date: Fri, 27 Sep 2024 10:43:31 +0200 Subject: [PATCH 053/802] Fix gcp text to speech uri fetch (#42309) - Fix acces to the uri attribute, if it's provided via the RecognitionAudio model. Co-authored-by: Oleg Kachur --- .../google/cloud/operators/speech_to_text.py | 17 +++--- .../cloud/operators/translate_speech.py | 15 ++--- .../cloud/operators/test_speech_to_text.py | 28 ++++++++- .../cloud/operators/test_translate_speech.py | 57 ++++++++++++++++--- 4 files changed, 90 insertions(+), 27 deletions(-) diff --git a/airflow/providers/google/cloud/operators/speech_to_text.py b/airflow/providers/google/cloud/operators/speech_to_text.py index de26a8ba8216f..f8c3e4703f9e0 100644 --- a/airflow/providers/google/cloud/operators/speech_to_text.py +++ b/airflow/providers/google/cloud/operators/speech_to_text.py @@ -113,15 +113,14 @@ def execute(self, context: Context): gcp_conn_id=self.gcp_conn_id, impersonation_chain=self.impersonation_chain, ) - - FileDetailsLink.persist( - context=context, - task_instance=self, - # Slice from: "gs://{BUCKET_NAME}/{FILE_NAME}" to: "{BUCKET_NAME}/{FILE_NAME}" - uri=self.audio["uri"][5:], - project_id=self.project_id or hook.project_id, - ) - + if self.audio.uri: + FileDetailsLink.persist( + context=context, + task_instance=self, + # Slice from: "gs://{BUCKET_NAME}/{FILE_NAME}" to: "{BUCKET_NAME}/{FILE_NAME}" + uri=self.audio.uri[5:], + project_id=self.project_id or hook.project_id, + ) response = hook.recognize_speech( config=self.config, audio=self.audio, retry=self.retry, timeout=self.timeout ) diff --git a/airflow/providers/google/cloud/operators/translate_speech.py b/airflow/providers/google/cloud/operators/translate_speech.py index fb3bdccb1abee..b0b540c31e086 100644 --- a/airflow/providers/google/cloud/operators/translate_speech.py +++ b/airflow/providers/google/cloud/operators/translate_speech.py @@ -169,7 +169,14 @@ def execute(self, context: Context) -> dict: raise AirflowException( f"Wrong response '{recognize_dict}' returned - it should contain {key} field" ) - + if self.audio.uri: + FileDetailsLink.persist( + context=context, + task_instance=self, + # Slice from: "gs://{BUCKET_NAME}/{FILE_NAME}" to: "{BUCKET_NAME}/{FILE_NAME}" + uri=self.audio.uri[5:], + project_id=self.project_id or translate_hook.project_id, + ) try: translation = translate_hook.translate( values=transcript, @@ -179,12 +186,6 @@ def execute(self, context: Context) -> dict: model=self.model, ) self.log.info("Translated output: %s", translation) - FileDetailsLink.persist( - context=context, - task_instance=self, - uri=self.audio["uri"][5:], - project_id=self.project_id or translate_hook.project_id, - ) return translation except ValueError as e: self.log.error("An error has been thrown from translate speech method:") diff --git a/tests/providers/google/cloud/operators/test_speech_to_text.py b/tests/providers/google/cloud/operators/test_speech_to_text.py index 51dd6dd8db7c0..1d7fa9ca37fea 100644 --- a/tests/providers/google/cloud/operators/test_speech_to_text.py +++ b/tests/providers/google/cloud/operators/test_speech_to_text.py @@ -21,7 +21,7 @@ import pytest from google.api_core.gapic_v1.method import DEFAULT -from google.cloud.speech_v1 import RecognizeResponse +from google.cloud.speech_v1 import RecognitionAudio, RecognitionConfig, RecognizeResponse from airflow.exceptions import AirflowException from airflow.providers.google.cloud.operators.speech_to_text import CloudSpeechToTextRecognizeSpeechOperator @@ -29,8 +29,8 @@ PROJECT_ID = "project-id" GCP_CONN_ID = "gcp-conn-id" IMPERSONATION_CHAIN = ["ACCOUNT_1", "ACCOUNT_2", "ACCOUNT_3"] -CONFIG = {"encoding": "LINEAR16"} -AUDIO = {"uri": "gs://bucket/object"} +CONFIG = RecognitionConfig({"encoding": "LINEAR16"}) +AUDIO = RecognitionAudio({"uri": "gs://bucket/object"}) class TestCloudSpeechToTextRecognizeSpeechOperator: @@ -80,3 +80,25 @@ def test_missing_audio(self, mock_hook): err = ctx.value assert "audio" in str(err) mock_hook.assert_not_called() + + @patch("airflow.providers.google.cloud.operators.speech_to_text.FileDetailsLink.persist") + @patch("airflow.providers.google.cloud.operators.speech_to_text.CloudSpeechToTextHook") + def test_no_audio_uri(self, mock_hook, mock_file_link): + mock_hook.return_value.recognize_speech.return_value = RecognizeResponse() + AUDIO_NO_URI = RecognitionAudio({"content": b"set content data instead of uri"}) + + op = CloudSpeechToTextRecognizeSpeechOperator( + project_id=PROJECT_ID, + gcp_conn_id=GCP_CONN_ID, + config=CONFIG, + audio=AUDIO_NO_URI, + task_id="id", + impersonation_chain=IMPERSONATION_CHAIN, + ) + op.execute(context=MagicMock()) + + mock_hook.return_value.recognize_speech.assert_called_once_with( + config=CONFIG, audio=AUDIO_NO_URI, retry=DEFAULT, timeout=None + ) + assert op.audio.uri == "" + mock_file_link.assert_not_called() diff --git a/tests/providers/google/cloud/operators/test_translate_speech.py b/tests/providers/google/cloud/operators/test_translate_speech.py index 6dd000504cef5..8e6beb79b9702 100644 --- a/tests/providers/google/cloud/operators/test_translate_speech.py +++ b/tests/providers/google/cloud/operators/test_translate_speech.py @@ -21,6 +21,8 @@ import pytest from google.cloud.speech_v1 import ( + RecognitionAudio, + RecognitionConfig, RecognizeResponse, SpeechRecognitionAlternative, SpeechRecognitionResult, @@ -54,8 +56,8 @@ def test_minimal_green_path(self, mock_translate_hook, mock_speech_hook): ] op = CloudTranslateSpeechOperator( - audio={"uri": "gs://bucket/object"}, - config={"encoding": "LINEAR16"}, + audio=RecognitionAudio({"uri": "gs://bucket/object"}), + config=RecognitionConfig({"encoding": "LINEAR16"}), target_language="pl", format_="text", source_language=None, @@ -77,8 +79,8 @@ def test_minimal_green_path(self, mock_translate_hook, mock_speech_hook): ) mock_speech_hook.return_value.recognize_speech.assert_called_once_with( - audio={"uri": "gs://bucket/object"}, - config={"encoding": "LINEAR16"}, + audio=RecognitionAudio({"uri": "gs://bucket/object"}), + config=RecognitionConfig({"encoding": "LINEAR16"}), ) mock_translate_hook.return_value.translate.assert_called_once_with( @@ -104,8 +106,8 @@ def test_bad_recognition_response(self, mock_translate_hook, mock_speech_hook): results=[SpeechRecognitionResult()] ) op = CloudTranslateSpeechOperator( - audio={"uri": "gs://bucket/object"}, - config={"encoding": "LINEAR16"}, + audio=RecognitionAudio({"uri": "gs://bucket/object"}), + config=RecognitionConfig({"encoding": "LINEAR16"}), target_language="pl", format_="text", source_language=None, @@ -128,8 +130,47 @@ def test_bad_recognition_response(self, mock_translate_hook, mock_speech_hook): ) mock_speech_hook.return_value.recognize_speech.assert_called_once_with( - audio={"uri": "gs://bucket/object"}, - config={"encoding": "LINEAR16"}, + audio=RecognitionAudio({"uri": "gs://bucket/object"}), + config=RecognitionConfig({"encoding": "LINEAR16"}), ) mock_translate_hook.return_value.translate.assert_not_called() + + @mock.patch("airflow.providers.google.cloud.operators.translate_speech.FileDetailsLink.persist") + @mock.patch("airflow.providers.google.cloud.operators.translate_speech.CloudSpeechToTextHook") + @mock.patch("airflow.providers.google.cloud.operators.translate_speech.CloudTranslateHook") + def test_no_audio_uri(self, mock_translate_hook, mock_speech_hook, file_link_mock): + mock_speech_hook.return_value.recognize_speech.return_value = RecognizeResponse( + results=[ + SpeechRecognitionResult( + alternatives=[SpeechRecognitionAlternative(transcript="test speech recognition result")] + ) + ] + ) + mock_translate_hook.return_value.translate.return_value = [ + { + "translatedText": "sprawdzić wynik rozpoznawania mowy", + "detectedSourceLanguage": "en", + "model": "base", + "input": "test speech recognition result", + } + ] + op = CloudTranslateSpeechOperator( + audio=RecognitionAudio({"content": b"set content data instead of uri"}), + config=RecognitionConfig({"encoding": "LINEAR16"}), + target_language="pl", + format_="text", + source_language=None, + model="base", + gcp_conn_id=GCP_CONN_ID, + task_id="id", + impersonation_chain=IMPERSONATION_CHAIN, + ) + op.execute(context=mock.MagicMock()) + + mock_speech_hook.return_value.recognize_speech.assert_called_once_with( + audio=RecognitionAudio({"content": b"set content data instead of uri"}), + config=RecognitionConfig({"encoding": "LINEAR16"}), + ) + assert op.audio.uri == "" + file_link_mock.assert_not_called() From 409b23dd3ec1936002b85642b10409e32feba14f Mon Sep 17 00:00:00 2001 From: Pierre Jeambrun Date: Fri, 27 Sep 2024 16:52:29 +0800 Subject: [PATCH 054/802] Update code owners (#42536) --- .github/CODEOWNERS | 1 + 1 file changed, 1 insertion(+) diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS index a0a7b82331be6..4de511fdfafc2 100644 --- a/.github/CODEOWNERS +++ b/.github/CODEOWNERS @@ -25,6 +25,7 @@ # API /airflow/api/ @ephraimbuddy @pierrejeambrun /airflow/api_connexion/ @ephraimbuddy @pierrejeambrun +/airflow/api_fastapi/ @ephraimbuddy @pierrejeambrun # WWW /airflow/www/ @ryanahamilton @ashb @bbovenzi @pierrejeambrun @jscheffl From 5bf6b1678176cefc0d958b5b816a833deb08926d Mon Sep 17 00:00:00 2001 From: Vincent <97131062+vincbeck@users.noreply.github.com> Date: Fri, 27 Sep 2024 05:43:55 -0400 Subject: [PATCH 055/802] Fix AWS system test `example_batch` (#42518) --- tests/system/providers/amazon/aws/example_batch.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/system/providers/amazon/aws/example_batch.py b/tests/system/providers/amazon/aws/example_batch.py index 0b79bb5a82b29..b33078407a296 100644 --- a/tests/system/providers/amazon/aws/example_batch.py +++ b/tests/system/providers/amazon/aws/example_batch.py @@ -207,7 +207,7 @@ def delete_job_queue(job_queue_name): job_name=batch_job_name, job_queue=batch_job_queue_name, job_definition=batch_job_definition_name, - ecs_properties_override=JOB_OVERRIDES, + container_overrides=JOB_OVERRIDES, ) # [END howto_operator_batch] From 56378b77e27f3909b278f4e78b93333f12d49547 Mon Sep 17 00:00:00 2001 From: max <42827971+moiseenkov@users.noreply.github.com> Date: Fri, 27 Sep 2024 10:57:14 +0000 Subject: [PATCH 056/802] Undo partition exclusion from the table name when splitting a full BigQuery table name (#42541) --- airflow/providers/google/cloud/hooks/bigquery.py | 4 ---- tests/providers/google/cloud/hooks/test_bigquery.py | 7 ++----- 2 files changed, 2 insertions(+), 9 deletions(-) diff --git a/airflow/providers/google/cloud/hooks/bigquery.py b/airflow/providers/google/cloud/hooks/bigquery.py index b1aed15c458ca..5f8881f45854d 100644 --- a/airflow/providers/google/cloud/hooks/bigquery.py +++ b/airflow/providers/google/cloud/hooks/bigquery.py @@ -2418,10 +2418,6 @@ def var_print(var_name): f"{var_print(var_name)}Expect format of (., " f"got {table_input}" ) - - # Exclude partition from the table name - table_id = table_id.split("$")[0] - if project_id is None: if var_name is not None: self.log.info( diff --git a/tests/providers/google/cloud/hooks/test_bigquery.py b/tests/providers/google/cloud/hooks/test_bigquery.py index 81db43c0f5310..02f442cfc7caa 100644 --- a/tests/providers/google/cloud/hooks/test_bigquery.py +++ b/tests/providers/google/cloud/hooks/test_bigquery.py @@ -1034,7 +1034,6 @@ def test_split_tablename_internal_need_default_project(self): with pytest.raises(ValueError, match="INTERNAL: No default project is specified"): self.hook.split_tablename("dataset.table", None) - @pytest.mark.parametrize("partition", ["$partition", ""]) @pytest.mark.parametrize( "project_expected, dataset_expected, table_expected, table_input", [ @@ -1045,11 +1044,9 @@ def test_split_tablename_internal_need_default_project(self): ("alt1:alt", "dataset", "table", "alt1:alt:dataset.table"), ], ) - def test_split_tablename( - self, project_expected, dataset_expected, table_expected, table_input, partition - ): + def test_split_tablename(self, project_expected, dataset_expected, table_expected, table_input): default_project_id = "project" - project, dataset, table = self.hook.split_tablename(table_input + partition, default_project_id) + project, dataset, table = self.hook.split_tablename(table_input, default_project_id) assert project_expected == project assert dataset_expected == dataset assert table_expected == table From 487a6451ecc09659339ed90bd6c125a2c010fe2d Mon Sep 17 00:00:00 2001 From: rom sharon <33751805+romsharon98@users.noreply.github.com> Date: Fri, 27 Sep 2024 15:07:49 +0300 Subject: [PATCH 057/802] Add slack notification for canary build failures (#42394) * add slack notifier --- .github/workflows/ci.yml | 29 +++++++++++++++++++++++++++++ 1 file changed, 29 insertions(+) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 14fe0bbe4baa9..8625aee73d9e1 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -39,6 +39,7 @@ env: GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} GITHUB_USERNAME: ${{ github.actor }} IMAGE_TAG: "${{ github.event.pull_request.head.sha || github.sha }}" + SLACK_BOT_TOKEN: ${{ secrets.SLACK_BOT_TOKEN }} VERBOSE: "true" concurrency: @@ -669,3 +670,31 @@ jobs: include-success-outputs: ${{ needs.build-info.outputs.include-success-outputs }} docker-cache: ${{ needs.build-info.outputs.docker-cache }} canary-run: ${{ needs.build-info.outputs.canary-run }} + + notify-slack-failure: + name: "Notify Slack on Failure" + if: github.event_name == 'schedule' && failure() + runs-on: ["ubuntu-22.04"] + steps: + - name: Notify Slack + id: slack + uses: slackapi/slack-github-action@v1.27.0 + with: + channel-id: 'zzz_webhook_test' + # yamllint disable rule:line-length + payload: | + { + "text": "🚨🕒 Scheduled CI Failure Alert 🕒🚨\n\n*Details:* ", + "blocks": [ + { + "type": "section", + "text": { + "type": "mrkdwn", + "text": "🚨🕒 Scheduled CI Failure Alert 🕒🚨\n\n*Details:* " + } + } + ] + } + # yamllint enable rule:line-length + env: + SLACK_BOT_TOKEN: ${{ env.SLACK_BOT_TOKEN }} From 034fd4dd0f2dca32c6fdbf5856a12f20be600f5c Mon Sep 17 00:00:00 2001 From: Pierre Jeambrun Date: Fri, 27 Sep 2024 20:27:27 +0800 Subject: [PATCH 058/802] AIP-84 Add HTTPException openapi documentation (#42508) * Add HTTPException openapi documentation * Update following code review --- airflow/api_fastapi/openapi/__init__.py | 16 ++++++++ airflow/api_fastapi/openapi/exceptions.py | 41 +++++++++++++++++++ airflow/api_fastapi/openapi/v1-generated.yaml | 36 ++++++++++++++++ airflow/api_fastapi/views/public/dags.py | 14 +++---- airflow/api_fastapi/views/ui/datasets.py | 2 - .../ui/openapi-gen/requests/schemas.gen.ts | 20 +++++++++ .../ui/openapi-gen/requests/services.gen.ts | 4 ++ airflow/ui/openapi-gen/requests/types.gen.ts | 27 ++++++++++++ 8 files changed, 150 insertions(+), 10 deletions(-) create mode 100644 airflow/api_fastapi/openapi/__init__.py create mode 100644 airflow/api_fastapi/openapi/exceptions.py diff --git a/airflow/api_fastapi/openapi/__init__.py b/airflow/api_fastapi/openapi/__init__.py new file mode 100644 index 0000000000000..13a83393a9124 --- /dev/null +++ b/airflow/api_fastapi/openapi/__init__.py @@ -0,0 +1,16 @@ +# 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. diff --git a/airflow/api_fastapi/openapi/exceptions.py b/airflow/api_fastapi/openapi/exceptions.py new file mode 100644 index 0000000000000..b3eaf204cc063 --- /dev/null +++ b/airflow/api_fastapi/openapi/exceptions.py @@ -0,0 +1,41 @@ +# 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 + +from pydantic import BaseModel + + +class HTTPExceptionResponse(BaseModel): + """HTTPException Model used for error response.""" + + detail: str | dict + + +def create_openapi_http_exception_doc(responses_status_code: list[int]) -> dict: + """ + Will create additional response example for errors raised by the endpoint. + + There is no easy way to introspect the code and automatically see what HTTPException are actually + raised by the endpoint implementation. This piece of documentation needs to be kept + in sync with the endpoint code manually. + + Validation error i.e 422 are natively added to the openapi documentation by FastAPI. + """ + responses_status_code = sorted(responses_status_code) + + return {status_code: {"model": HTTPExceptionResponse} for status_code in responses_status_code} diff --git a/airflow/api_fastapi/openapi/v1-generated.yaml b/airflow/api_fastapi/openapi/v1-generated.yaml index f488825449c3a..64e475aeb6baa 100644 --- a/airflow/api_fastapi/openapi/v1-generated.yaml +++ b/airflow/api_fastapi/openapi/v1-generated.yaml @@ -168,6 +168,30 @@ paths: application/json: schema: $ref: '#/components/schemas/DAGResponse' + '400': + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPExceptionResponse' + description: Bad Request + '401': + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPExceptionResponse' + description: Unauthorized + '403': + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPExceptionResponse' + description: Forbidden + '404': + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPExceptionResponse' + description: Not Found '422': description: Validation Error content: @@ -386,6 +410,18 @@ components: title: DagTagPydantic description: Serializable representation of the DagTag ORM SqlAlchemyModel used by internal API. + HTTPExceptionResponse: + properties: + detail: + anyOf: + - type: string + - type: object + title: Detail + type: object + required: + - detail + title: HTTPExceptionResponse + description: HTTPException Model used for error response. HTTPValidationError: properties: detail: diff --git a/airflow/api_fastapi/views/public/dags.py b/airflow/api_fastapi/views/public/dags.py index 07ab968adc975..a9fe87eef0953 100644 --- a/airflow/api_fastapi/views/public/dags.py +++ b/airflow/api_fastapi/views/public/dags.py @@ -23,6 +23,7 @@ from typing_extensions import Annotated from airflow.api_fastapi.db import apply_filters_to_select, get_session, latest_dag_run_per_dag_id_cte +from airflow.api_fastapi.openapi.exceptions import create_openapi_http_exception_doc from airflow.api_fastapi.parameters import ( QueryDagDisplayNamePatternSearch, QueryDagIdPatternSearch, @@ -95,16 +96,13 @@ async def get_dags( dags = session.scalars(dags_query).all() - try: - return DAGCollectionResponse( - dags=[DAGResponse.model_validate(dag, from_attributes=True) for dag in dags], - total_entries=total_entries, - ) - except ValueError as e: - raise HTTPException(400, f"DAGCollectionSchema error: {str(e)}") + return DAGCollectionResponse( + dags=[DAGResponse.model_validate(dag, from_attributes=True) for dag in dags], + total_entries=total_entries, + ) -@dags_router.patch("/dags/{dag_id}") +@dags_router.patch("/dags/{dag_id}", responses=create_openapi_http_exception_doc([400, 401, 403, 404])) async def patch_dag( dag_id: str, patch_body: DAGPatchBody, diff --git a/airflow/api_fastapi/views/ui/datasets.py b/airflow/api_fastapi/views/ui/datasets.py index 484385031a23d..f5dd2cacb126d 100644 --- a/airflow/api_fastapi/views/ui/datasets.py +++ b/airflow/api_fastapi/views/ui/datasets.py @@ -29,8 +29,6 @@ datasets_router = APIRouter(tags=["Dataset"]) -# Ultimately we want async routes, with async sqlalchemy session / context manager. -# Additional effort to make airflow utility code async, not handled for now and most likely part of the AIP-70 @datasets_router.get("/next_run_datasets/{dag_id}", include_in_schema=False) async def next_run_datasets( dag_id: str, diff --git a/airflow/ui/openapi-gen/requests/schemas.gen.ts b/airflow/ui/openapi-gen/requests/schemas.gen.ts index d9ce0528c396c..e8c9b5d70cdf4 100644 --- a/airflow/ui/openapi-gen/requests/schemas.gen.ts +++ b/airflow/ui/openapi-gen/requests/schemas.gen.ts @@ -317,6 +317,26 @@ export const $DagTagPydantic = { "Serializable representation of the DagTag ORM SqlAlchemyModel used by internal API.", } as const; +export const $HTTPExceptionResponse = { + properties: { + detail: { + anyOf: [ + { + type: "string", + }, + { + type: "object", + }, + ], + title: "Detail", + }, + }, + type: "object", + required: ["detail"], + title: "HTTPExceptionResponse", + description: "HTTPException Model used for error response.", +} as const; + export const $HTTPValidationError = { properties: { detail: { diff --git a/airflow/ui/openapi-gen/requests/services.gen.ts b/airflow/ui/openapi-gen/requests/services.gen.ts index 9c261b3039000..37a4d11873acf 100644 --- a/airflow/ui/openapi-gen/requests/services.gen.ts +++ b/airflow/ui/openapi-gen/requests/services.gen.ts @@ -102,6 +102,10 @@ export class DagService { body: data.requestBody, mediaType: "application/json", errors: { + 400: "Bad Request", + 401: "Unauthorized", + 403: "Forbidden", + 404: "Not Found", 422: "Validation Error", }, }); diff --git a/airflow/ui/openapi-gen/requests/types.gen.ts b/airflow/ui/openapi-gen/requests/types.gen.ts index 803bcd84270c7..16977004e79d6 100644 --- a/airflow/ui/openapi-gen/requests/types.gen.ts +++ b/airflow/ui/openapi-gen/requests/types.gen.ts @@ -67,6 +67,17 @@ export type DagTagPydantic = { dag_id: string; }; +/** + * HTTPException Model used for error response. + */ +export type HTTPExceptionResponse = { + detail: + | string + | { + [key: string]: unknown; + }; +}; + export type HTTPValidationError = { detail?: Array; }; @@ -149,6 +160,22 @@ export type $OpenApiTs = { * Successful Response */ 200: DAGResponse; + /** + * Bad Request + */ + 400: HTTPExceptionResponse; + /** + * Unauthorized + */ + 401: HTTPExceptionResponse; + /** + * Forbidden + */ + 403: HTTPExceptionResponse; + /** + * Not Found + */ + 404: HTTPExceptionResponse; /** * Validation Error */ From c28ce62a94944e66ca434906e33648cba77e1dae Mon Sep 17 00:00:00 2001 From: GPK Date: Fri, 27 Sep 2024 13:55:24 +0100 Subject: [PATCH 059/802] Fix SparkKubernetesOperator spark name. (#42427) * use name parameter from spark yaml config or from operator argument parameter * update tests and name usage condition check * adding test, to check spark name starts with task_id * use set_name function in create_job * remove lower --- .../kubernetes/operators/spark_kubernetes.py | 15 +- ...ication_test_with_no_name_from_config.json | 57 +++++++ ...ication_test_with_no_name_from_config.yaml | 55 +++++++ .../operators/test_spark_kubernetes.py | 143 ++++++++++++++++++ 4 files changed, 265 insertions(+), 5 deletions(-) create mode 100644 tests/providers/cncf/kubernetes/data_files/spark/application_test_with_no_name_from_config.json create mode 100644 tests/providers/cncf/kubernetes/data_files/spark/application_test_with_no_name_from_config.yaml diff --git a/airflow/providers/cncf/kubernetes/operators/spark_kubernetes.py b/airflow/providers/cncf/kubernetes/operators/spark_kubernetes.py index 39fadae90e5bd..9bcf46d0d4f57 100644 --- a/airflow/providers/cncf/kubernetes/operators/spark_kubernetes.py +++ b/airflow/providers/cncf/kubernetes/operators/spark_kubernetes.py @@ -17,7 +17,6 @@ # under the License. from __future__ import annotations -import re from functools import cached_property from pathlib import Path from typing import TYPE_CHECKING, Any @@ -83,7 +82,7 @@ def __init__( image: str | None = None, code_path: str | None = None, namespace: str = "default", - name: str = "default", + name: str | None = None, application_file: str | None = None, template_spec=None, get_logs: bool = True, @@ -103,7 +102,6 @@ def __init__( self.code_path = code_path self.application_file = application_file self.template_spec = template_spec - self.name = self.create_job_name() self.kubernetes_conn_id = kubernetes_conn_id self.startup_timeout_seconds = startup_timeout_seconds self.reattach_on_restart = reattach_on_restart @@ -161,8 +159,13 @@ def manage_template_specs(self): return template_body def create_job_name(self): - initial_name = add_unique_suffix(name=self.task_id, max_len=MAX_LABEL_LEN) - return re.sub(r"[^a-z0-9-]+", "-", initial_name.lower()) + name = ( + self.name or self.template_body.get("spark", {}).get("metadata", {}).get("name") or self.task_id + ) + + updated_name = add_unique_suffix(name=name, max_len=MAX_LABEL_LEN) + + return self._set_name(updated_name) @staticmethod def _get_pod_identifying_label_string(labels) -> str: @@ -282,6 +285,8 @@ def custom_obj_api(self) -> CustomObjectsApi: return CustomObjectsApi() def execute(self, context: Context): + self.name = self.create_job_name() + self.log.info("Creating sparkApplication.") self.launcher = CustomObjectLauncher( name=self.name, diff --git a/tests/providers/cncf/kubernetes/data_files/spark/application_test_with_no_name_from_config.json b/tests/providers/cncf/kubernetes/data_files/spark/application_test_with_no_name_from_config.json new file mode 100644 index 0000000000000..1504c40fbd1e9 --- /dev/null +++ b/tests/providers/cncf/kubernetes/data_files/spark/application_test_with_no_name_from_config.json @@ -0,0 +1,57 @@ +{ + "apiVersion":"sparkoperator.k8s.io/v1beta2", + "kind":"SparkApplication", + "metadata":{ + "namespace":"default" + }, + "spec":{ + "type":"Scala", + "mode":"cluster", + "image":"gcr.io/spark-operator/spark:v2.4.5", + "imagePullPolicy":"Always", + "mainClass":"org.apache.spark.examples.SparkPi", + "mainApplicationFile":"local:///opt/spark/examples/jars/spark-examples_2.11-2.4.5.jar", + "sparkVersion":"2.4.5", + "restartPolicy":{ + "type":"Never" + }, + "volumes":[ + { + "name":"test-volume", + "hostPath":{ + "path":"/tmp", + "type":"Directory" + } + } + ], + "driver":{ + "cores":1, + "coreLimit":"1200m", + "memory":"512m", + "labels":{ + "version":"2.4.5" + }, + "serviceAccount":"spark", + "volumeMounts":[ + { + "name":"test-volume", + "mountPath":"/tmp" + } + ] + }, + "executor":{ + "cores":1, + "instances":1, + "memory":"512m", + "labels":{ + "version":"2.4.5" + }, + "volumeMounts":[ + { + "name":"test-volume", + "mountPath":"/tmp" + } + ] + } + } +} diff --git a/tests/providers/cncf/kubernetes/data_files/spark/application_test_with_no_name_from_config.yaml b/tests/providers/cncf/kubernetes/data_files/spark/application_test_with_no_name_from_config.yaml new file mode 100644 index 0000000000000..91723980954ee --- /dev/null +++ b/tests/providers/cncf/kubernetes/data_files/spark/application_test_with_no_name_from_config.yaml @@ -0,0 +1,55 @@ +# 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. +--- +apiVersion: "sparkoperator.k8s.io/v1beta2" +kind: SparkApplication +metadata: + namespace: default +spec: + type: Scala + mode: cluster + image: "gcr.io/spark-operator/spark:v2.4.5" + imagePullPolicy: Always + mainClass: org.apache.spark.examples.SparkPi + mainApplicationFile: "local:///opt/spark/examples/jars/spark-examples_2.11-2.4.5.jar" + sparkVersion: "2.4.5" + restartPolicy: + type: Never + volumes: + - name: "test-volume" + hostPath: + path: "/tmp" + type: Directory + driver: + cores: 1 + coreLimit: "1200m" + memory: "512m" + labels: + version: 2.4.5 + serviceAccount: spark + volumeMounts: + - name: "test-volume" + mountPath: "/tmp" + executor: + cores: 1 + instances: 1 + memory: "512m" + labels: + version: 2.4.5 + volumeMounts: + - name: "test-volume" + mountPath: "/tmp" diff --git a/tests/providers/cncf/kubernetes/operators/test_spark_kubernetes.py b/tests/providers/cncf/kubernetes/operators/test_spark_kubernetes.py index bc8404aa85607..9c8c40de6558d 100644 --- a/tests/providers/cncf/kubernetes/operators/test_spark_kubernetes.py +++ b/tests/providers/cncf/kubernetes/operators/test_spark_kubernetes.py @@ -273,6 +273,149 @@ def test_create_application_from_yaml_json( version="v1beta2", ) + def test_create_application_from_yaml_json_and_use_name_from_metadata( + self, + mock_create_namespaced_crd, + mock_get_namespaced_custom_object_status, + mock_cleanup, + mock_create_job_name, + mock_get_kube_client, + mock_create_pod, + mock_await_pod_start, + mock_await_pod_completion, + mock_fetch_requested_container_logs, + data_file, + ): + op = SparkKubernetesOperator( + application_file=data_file("spark/application_test.yaml").as_posix(), + kubernetes_conn_id="kubernetes_default_kube_config", + task_id="create_app_and_use_name_from_metadata", + ) + context = create_context(op) + op.execute(context) + TEST_APPLICATION_DICT["metadata"]["name"] = op.name + mock_create_namespaced_crd.assert_called_with( + body=TEST_APPLICATION_DICT, + group="sparkoperator.k8s.io", + namespace="default", + plural="sparkapplications", + version="v1beta2", + ) + assert op.name.startswith("default_yaml") + + op = SparkKubernetesOperator( + application_file=data_file("spark/application_test.json").as_posix(), + kubernetes_conn_id="kubernetes_default_kube_config", + task_id="create_app_and_use_name_from_metadata", + ) + context = create_context(op) + op.execute(context) + TEST_APPLICATION_DICT["metadata"]["name"] = op.name + mock_create_namespaced_crd.assert_called_with( + body=TEST_APPLICATION_DICT, + group="sparkoperator.k8s.io", + namespace="default", + plural="sparkapplications", + version="v1beta2", + ) + assert op.name.startswith("default_json") + + def test_create_application_from_yaml_json_and_use_name_from_operator_args( + self, + mock_create_namespaced_crd, + mock_get_namespaced_custom_object_status, + mock_cleanup, + mock_create_job_name, + mock_get_kube_client, + mock_create_pod, + mock_await_pod_start, + mock_await_pod_completion, + mock_fetch_requested_container_logs, + data_file, + ): + op = SparkKubernetesOperator( + application_file=data_file("spark/application_test.yaml").as_posix(), + kubernetes_conn_id="kubernetes_default_kube_config", + task_id="default_yaml", + name="test-spark", + ) + context = create_context(op) + op.execute(context) + TEST_APPLICATION_DICT["metadata"]["name"] = op.name + mock_create_namespaced_crd.assert_called_with( + body=TEST_APPLICATION_DICT, + group="sparkoperator.k8s.io", + namespace="default", + plural="sparkapplications", + version="v1beta2", + ) + assert op.name.startswith("test-spark") + + op = SparkKubernetesOperator( + application_file=data_file("spark/application_test.json").as_posix(), + kubernetes_conn_id="kubernetes_default_kube_config", + task_id="default_json", + name="test-spark", + ) + context = create_context(op) + op.execute(context) + TEST_APPLICATION_DICT["metadata"]["name"] = op.name + mock_create_namespaced_crd.assert_called_with( + body=TEST_APPLICATION_DICT, + group="sparkoperator.k8s.io", + namespace="default", + plural="sparkapplications", + version="v1beta2", + ) + assert op.name.startswith("test-spark") + + def test_create_application_from_yaml_json_and_use_name_task_id( + self, + mock_create_namespaced_crd, + mock_get_namespaced_custom_object_status, + mock_cleanup, + mock_create_job_name, + mock_get_kube_client, + mock_create_pod, + mock_await_pod_start, + mock_await_pod_completion, + mock_fetch_requested_container_logs, + data_file, + ): + op = SparkKubernetesOperator( + application_file=data_file("spark/application_test_with_no_name_from_config.yaml").as_posix(), + kubernetes_conn_id="kubernetes_default_kube_config", + task_id="create_app_and_use_name_from_task_id", + ) + context = create_context(op) + op.execute(context) + TEST_APPLICATION_DICT["metadata"]["name"] = op.name + mock_create_namespaced_crd.assert_called_with( + body=TEST_APPLICATION_DICT, + group="sparkoperator.k8s.io", + namespace="default", + plural="sparkapplications", + version="v1beta2", + ) + assert op.name.startswith("create_app_and_use_name_from_task_id") + + op = SparkKubernetesOperator( + application_file=data_file("spark/application_test_with_no_name_from_config.json").as_posix(), + kubernetes_conn_id="kubernetes_default_kube_config", + task_id="create_app_and_use_name_from_task_id", + ) + context = create_context(op) + op.execute(context) + TEST_APPLICATION_DICT["metadata"]["name"] = op.name + mock_create_namespaced_crd.assert_called_with( + body=TEST_APPLICATION_DICT, + group="sparkoperator.k8s.io", + namespace="default", + plural="sparkapplications", + version="v1beta2", + ) + assert op.name.startswith("create_app_and_use_name_from_task_id") + def test_new_template_from_yaml( self, mock_create_namespaced_crd, From 924978423ec5ef5463c9a1d1a4ab1e577e6018ca Mon Sep 17 00:00:00 2001 From: dan-js <50807588+dan-js@users.noreply.github.com> Date: Fri, 27 Sep 2024 13:56:48 +0100 Subject: [PATCH 060/802] Fix incorrect operator name in FileTransferOperator example (#42543) --- docs/apache-airflow-providers-common-io/operators.rst | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/apache-airflow-providers-common-io/operators.rst b/docs/apache-airflow-providers-common-io/operators.rst index 7170c18980238..12b4a1c207ff0 100644 --- a/docs/apache-airflow-providers-common-io/operators.rst +++ b/docs/apache-airflow-providers-common-io/operators.rst @@ -38,7 +38,7 @@ location to another. Parameters of the operator are: If the ``src`` and the ``dst`` are both on the same object storage, copy will be performed in the object storage. Otherwise the data will be streamed from the source to the destination. -The example below shows how to instantiate the SQLExecuteQueryOperator task. +The example below shows how to instantiate the FileTransferOperator task. .. exampleinclude:: /../../tests/system/providers/common/io/example_file_transfer_local_to_s3.py :language: python From 68cfe0d531674d507619ad746ef44dc639aeac44 Mon Sep 17 00:00:00 2001 From: GPK Date: Fri, 27 Sep 2024 14:24:42 +0100 Subject: [PATCH 061/802] Pre commit script to validate template fields (#42284) --- .pre-commit-config.yaml | 7 + contributing-docs/08_static_code_checks.rst | 2 + .../doc/images/output_static-checks.svg | 12 +- .../doc/images/output_static-checks.txt | 2 +- .../src/airflow_breeze/pre_commit_ids.py | 1 + .../pre_commit/check_provider_yaml_files.py | 15 +- .../ci/pre_commit/check_template_fields.py | 40 ++++ .../ci/pre_commit/common_precommit_utils.py | 17 ++ scripts/ci/pre_commit/migration_reference.py | 14 +- scripts/ci/pre_commit/update_er_diagram.py | 13 +- .../ci/pre_commit/update_fastapi_api_spec.py | 13 +- .../in_container/run_template_fields_check.py | 180 ++++++++++++++++++ 12 files changed, 279 insertions(+), 37 deletions(-) create mode 100755 scripts/ci/pre_commit/check_template_fields.py create mode 100644 scripts/in_container/run_template_fields_check.py diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 942b34ca2e6d5..2263086335bc2 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1343,6 +1343,13 @@ repos: files: ^airflow/providers/.*/provider\.yaml$ additional_dependencies: ['rich>=12.4.4'] require_serial: true + - id: check-template-fields-valid + name: Check templated fields mapped in operators/sensors + language: python + entry: ./scripts/ci/pre_commit/check_template_fields.py + files: ^airflow/.*/sensors/.*\.py$|^airflow/.*/operators/.*\.py$ + additional_dependencies: [ 'rich>=12.4.4' ] + require_serial: true - id: update-migration-references name: Update migration ref doc language: python diff --git a/contributing-docs/08_static_code_checks.rst b/contributing-docs/08_static_code_checks.rst index 0a3dcacd9e070..d50b9db3e607f 100644 --- a/contributing-docs/08_static_code_checks.rst +++ b/contributing-docs/08_static_code_checks.rst @@ -236,6 +236,8 @@ require Breeze Docker image to be built locally. +-----------------------------------------------------------+--------------------------------------------------------+---------+ | check-template-context-variable-in-sync | Sync template context variable refs | | +-----------------------------------------------------------+--------------------------------------------------------+---------+ +| check-template-fields-valid | Check templated fields mapped in operators/sensors | * | ++-----------------------------------------------------------+--------------------------------------------------------+---------+ | check-tests-in-the-right-folders | Check if tests are in the right folders | | +-----------------------------------------------------------+--------------------------------------------------------+---------+ | check-tests-unittest-testcase | Unit tests do not inherit from unittest.TestCase | | diff --git a/dev/breeze/doc/images/output_static-checks.svg b/dev/breeze/doc/images/output_static-checks.svg index ed52a596def64..36b88513a56ba 100644 --- a/dev/breeze/doc/images/output_static-checks.svg +++ b/dev/breeze/doc/images/output_static-checks.svg @@ -356,12 +356,12 @@ │check-safe-filter-usage-in-html | check-sql-dependency-common-data-structure |   │ │check-start-date-not-used-in-defaults | check-system-tests-present |             │ │check-system-tests-tocs | check-taskinstance-tis-attrs |                         │ -│check-template-context-variable-in-sync | check-tests-in-the-right-folders |     │ -│check-tests-unittest-testcase | check-urlparse-usage-in-code |                   │ -│check-usage-of-re2-over-re | check-xml | codespell | compile-ui-assets |         │ -│compile-ui-assets-dev | compile-www-assets | compile-www-assets-dev |            │ -│create-missing-init-py-files-tests | debug-statements | detect-private-key |     │ -│doctoc | end-of-file-fixer | fix-encoding-pragma | flynt |                       │ +│check-template-context-variable-in-sync | check-template-fields-valid |          │ +│check-tests-in-the-right-folders | check-tests-unittest-testcase |               │ +│check-urlparse-usage-in-code | check-usage-of-re2-over-re | check-xml | codespell│ +│| compile-ui-assets | compile-ui-assets-dev | compile-www-assets |               │ +│compile-www-assets-dev | create-missing-init-py-files-tests | debug-statements | │ +│detect-private-key | doctoc | end-of-file-fixer | fix-encoding-pragma | flynt |  │ │generate-airflow-diagrams | generate-openapi-spec | generate-pypi-readme |       │ │identity | insert-license | kubeconform | lint-chart-schema | lint-css |         │ │lint-dockerfile | lint-helm-chart | lint-json-schema | lint-markdown |           │ diff --git a/dev/breeze/doc/images/output_static-checks.txt b/dev/breeze/doc/images/output_static-checks.txt index 3a3837fbb15bb..9e3ae46130640 100644 --- a/dev/breeze/doc/images/output_static-checks.txt +++ b/dev/breeze/doc/images/output_static-checks.txt @@ -1 +1 @@ -5c6ba60b1865538bce04fc940cd240c6 +e33cdf5f43d8c63290e44e92dc19d2c4 diff --git a/dev/breeze/src/airflow_breeze/pre_commit_ids.py b/dev/breeze/src/airflow_breeze/pre_commit_ids.py index 9a48df5e3f69c..457379f5b90ba 100644 --- a/dev/breeze/src/airflow_breeze/pre_commit_ids.py +++ b/dev/breeze/src/airflow_breeze/pre_commit_ids.py @@ -83,6 +83,7 @@ "check-system-tests-tocs", "check-taskinstance-tis-attrs", "check-template-context-variable-in-sync", + "check-template-fields-valid", "check-tests-in-the-right-folders", "check-tests-unittest-testcase", "check-urlparse-usage-in-code", diff --git a/scripts/ci/pre_commit/check_provider_yaml_files.py b/scripts/ci/pre_commit/check_provider_yaml_files.py index fcbe2512910a3..f848e38afa0b2 100755 --- a/scripts/ci/pre_commit/check_provider_yaml_files.py +++ b/scripts/ci/pre_commit/check_provider_yaml_files.py @@ -17,12 +17,15 @@ # under the License. from __future__ import annotations -import os import sys from pathlib import Path sys.path.insert(0, str(Path(__file__).parent.resolve())) -from common_precommit_utils import console, initialize_breeze_precommit, run_command_via_breeze_shell +from common_precommit_utils import ( + initialize_breeze_precommit, + run_command_via_breeze_shell, + validate_cmd_result, +) initialize_breeze_precommit(__name__, __file__) @@ -33,10 +36,4 @@ warn_image_upgrade_needed=True, extra_env={"PYTHONWARNINGS": "default"}, ) -if cmd_result.returncode != 0 and os.environ.get("CI") != "true": - console.print( - "\n[yellow]If you see strange stacktraces above, especially about missing imports " - "run this command:[/]\n" - ) - console.print("[magenta]breeze ci-image build --python 3.8 --upgrade-to-newer-dependencies[/]\n") -sys.exit(cmd_result.returncode) +validate_cmd_result(cmd_result, include_ci_env_check=True) diff --git a/scripts/ci/pre_commit/check_template_fields.py b/scripts/ci/pre_commit/check_template_fields.py new file mode 100755 index 0000000000000..da0b60fbd978f --- /dev/null +++ b/scripts/ci/pre_commit/check_template_fields.py @@ -0,0 +1,40 @@ +#!/usr/bin/env python +# 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 sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).parent.resolve())) +from common_precommit_utils import ( + initialize_breeze_precommit, + run_command_via_breeze_shell, + validate_cmd_result, +) + +initialize_breeze_precommit(__name__, __file__) +py_files_to_test = sys.argv[1:] + +cmd_result = run_command_via_breeze_shell( + ["python3", "/opt/airflow/scripts/in_container/run_template_fields_check.py", *py_files_to_test], + backend="sqlite", + warn_image_upgrade_needed=True, + extra_env={"PYTHONWARNINGS": "default"}, +) + +validate_cmd_result(cmd_result, include_ci_env_check=True) diff --git a/scripts/ci/pre_commit/common_precommit_utils.py b/scripts/ci/pre_commit/common_precommit_utils.py index 41bc3a5eeaf93..4f62c50cabeaa 100644 --- a/scripts/ci/pre_commit/common_precommit_utils.py +++ b/scripts/ci/pre_commit/common_precommit_utils.py @@ -211,3 +211,20 @@ def check_list_sorted(the_list: list[str], message: str, errors: list[str]) -> b console.print() errors.append(f"ERROR in {message}. The elements are not sorted/unique.") return False + + +def validate_cmd_result(cmd_result, include_ci_env_check=False): + if include_ci_env_check: + if cmd_result.returncode != 0 and os.environ.get("CI") != "true": + console.print( + "\n[yellow]If you see strange stacktraces above, especially about missing imports " + "run this command:[/]\n" + ) + console.print("[magenta]breeze ci-image build --python 3.8 --upgrade-to-newer-dependencies[/]\n") + + elif cmd_result.returncode != 0: + console.print( + "[warning]\nIf you see strange stacktraces above, " + "run `breeze ci-image build --python 3.8` and try again." + ) + sys.exit(cmd_result.returncode) diff --git a/scripts/ci/pre_commit/migration_reference.py b/scripts/ci/pre_commit/migration_reference.py index 34d3a94c6a90d..505bea5ca91af 100755 --- a/scripts/ci/pre_commit/migration_reference.py +++ b/scripts/ci/pre_commit/migration_reference.py @@ -21,7 +21,11 @@ from pathlib import Path sys.path.insert(0, str(Path(__file__).parent.resolve())) -from common_precommit_utils import console, initialize_breeze_precommit, run_command_via_breeze_shell +from common_precommit_utils import ( + initialize_breeze_precommit, + run_command_via_breeze_shell, + validate_cmd_result, +) initialize_breeze_precommit(__name__, __file__) @@ -29,9 +33,5 @@ ["python3", "/opt/airflow/scripts/in_container/run_migration_reference.py"], backend="sqlite", ) -if cmd_result.returncode != 0: - console.print( - "[warning]\nIf you see strange stacktraces above, " - "run `breeze ci-image build --python 3.8` and try again." - ) -sys.exit(cmd_result.returncode) + +validate_cmd_result(cmd_result) diff --git a/scripts/ci/pre_commit/update_er_diagram.py b/scripts/ci/pre_commit/update_er_diagram.py index e660b47c6e6ae..c4f3cb797cf21 100755 --- a/scripts/ci/pre_commit/update_er_diagram.py +++ b/scripts/ci/pre_commit/update_er_diagram.py @@ -21,7 +21,11 @@ from pathlib import Path sys.path.insert(0, str(Path(__file__).parent.resolve())) -from common_precommit_utils import console, initialize_breeze_precommit, run_command_via_breeze_shell +from common_precommit_utils import ( + initialize_breeze_precommit, + run_command_via_breeze_shell, + validate_cmd_result, +) initialize_breeze_precommit(__name__, __file__) @@ -36,9 +40,4 @@ }, ) -if cmd_result.returncode != 0: - console.print( - "[warning]\nIf you see strange stacktraces above, " - "run `breeze ci-image build --python 3.8` and try again." - ) - sys.exit(cmd_result.returncode) +validate_cmd_result(cmd_result) diff --git a/scripts/ci/pre_commit/update_fastapi_api_spec.py b/scripts/ci/pre_commit/update_fastapi_api_spec.py index 15ccaa5ac209e..3d7731c7ef2e2 100755 --- a/scripts/ci/pre_commit/update_fastapi_api_spec.py +++ b/scripts/ci/pre_commit/update_fastapi_api_spec.py @@ -21,7 +21,11 @@ from pathlib import Path sys.path.insert(0, str(Path(__file__).parent.resolve())) -from common_precommit_utils import console, initialize_breeze_precommit, run_command_via_breeze_shell +from common_precommit_utils import ( + initialize_breeze_precommit, + run_command_via_breeze_shell, + validate_cmd_result, +) initialize_breeze_precommit(__name__, __file__) @@ -31,9 +35,4 @@ skip_environment_initialization=False, ) -if cmd_result.returncode != 0: - console.print( - "[warning]\nIf you see strange stacktraces above, " - "run `breeze ci-image build --python 3.8` and try again." - ) -sys.exit(cmd_result.returncode) +validate_cmd_result(cmd_result) diff --git a/scripts/in_container/run_template_fields_check.py b/scripts/in_container/run_template_fields_check.py new file mode 100644 index 0000000000000..202dce35c5745 --- /dev/null +++ b/scripts/in_container/run_template_fields_check.py @@ -0,0 +1,180 @@ +# 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 ast +import importlib.util +import inspect +import itertools +import pathlib +import sys +import warnings + +import yaml +from rich.console import Console + +try: + from yaml import CSafeLoader as SafeLoader +except ImportError: + from yaml import SafeLoader # type: ignore + +console = Console(width=400, color_system="standard") +ROOT_DIR = pathlib.Path(__file__).resolve().parents[2] + +provider_files_pattern = pathlib.Path(ROOT_DIR, "airflow", "providers").rglob("provider.yaml") +errors: list[str] = [] + +OPERATORS: list[str] = ["sensors", "operators"] +CLASS_IDENTIFIERS: list[str] = ["sensor", "operator"] + +TEMPLATE_TYPES: list[str] = ["template_fields"] + + +class InstanceFieldExtractor(ast.NodeVisitor): + def __init__(self): + self.current_class = None + self.instance_fields = [] + + def visit_FunctionDef(self, node: ast.FunctionDef) -> ast.FunctionDef: + if node.name == "__init__": + self.generic_visit(node) + return node + + def visit_Assign(self, node: ast.Assign) -> ast.Assign: + fields = [] + for target in node.targets: + if isinstance(target, ast.Attribute): + fields.append(target.attr) + if fields: + self.instance_fields.extend(fields) + return node + + def visit_AnnAssign(self, node: ast.AnnAssign) -> ast.AnnAssign: + if isinstance(node.target, ast.Attribute): + self.instance_fields.append(node.target.attr) + return node + + +def get_template_fields_and_class_instance_fields(cls): + """ + 1.This method retrieves the operator class and obtains all its parent classes using the method resolution order (MRO). + 2. It then gathers the templated fields declared in both the operator class and its parent classes. + 3. Finally, it retrieves the instance fields of the operator class, specifically the self.fields attributes. + """ + all_template_fields = [] + class_instance_fields = [] + + all_classes = cls.__mro__ + for current_class in all_classes: + if current_class.__init__ is not object.__init__: + cls_attr = current_class.__dict__ + for template_type in TEMPLATE_TYPES: + fields = cls_attr.get(template_type) + if fields: + all_template_fields.extend(fields) + + tree = ast.parse(inspect.getsource(current_class)) + visitor = InstanceFieldExtractor() + visitor.visit(tree) + if visitor.instance_fields: + class_instance_fields.extend(visitor.instance_fields) + return all_template_fields, class_instance_fields + + +def load_yaml_data() -> dict: + """ + It loads all the provider YAML files and retrieves the module referenced within each YAML file. + """ + package_paths = sorted(str(path) for path in provider_files_pattern) + result = {} + for provider_yaml_path in package_paths: + with open(provider_yaml_path) as yaml_file: + provider = yaml.load(yaml_file, SafeLoader) + rel_path = pathlib.Path(provider_yaml_path).relative_to(ROOT_DIR).as_posix() + result[rel_path] = provider + return result + + +def get_providers_modules() -> list[str]: + modules_container = [] + result = load_yaml_data() + + for (_, provider_data), resource_type in itertools.product(result.items(), OPERATORS): + if provider_data.get(resource_type): + for data in provider_data.get(resource_type): + modules_container.extend(data.get("python-modules")) + + return modules_container + + +def is_class_eligible(name: str) -> bool: + for op in CLASS_IDENTIFIERS: + if name.lower().endswith(op): + return True + return False + + +def get_eligible_classes(all_classes): + """ + Filter the results to include only classes that end with `Sensor` or `Operator`. + + """ + + eligible_classes = [(name, cls) for name, cls in all_classes if is_class_eligible(name)] + return eligible_classes + + +def iter_check_template_fields(module: str): + """ + 1. This method imports the providers module and retrieves all the classes defined within it. + 2. It then filters and selects classes related to operators or sensors by checking if the class name ends with "Operator" or "Sensor." + 3. For each operator class, it validates the template fields by inspecting the class instance fields. + """ + with warnings.catch_warnings(record=True): + imported_module = importlib.import_module(module) + classes = inspect.getmembers(imported_module, inspect.isclass) + op_classes = get_eligible_classes(classes) + + for op_class_name, cls in op_classes: + if cls.__module__ == module: + templated_fields, class_instance_fields = get_template_fields_and_class_instance_fields(cls) + + for field in templated_fields: + if field not in class_instance_fields: + errors.append(f"{module}: {op_class_name}: {field}") + + +if __name__ == "__main__": + provider_modules = get_providers_modules() + + if len(sys.argv) > 1: + py_files = sorted(sys.argv[1:]) + modules_to_validate = [ + module_name + for pyfile in py_files + if (module_name := pyfile.rstrip(".py").replace("/", ".")) in provider_modules + ] + else: + modules_to_validate = provider_modules + + [iter_check_template_fields(module) for module in modules_to_validate] + if errors: + console.print("[red]Found Invalid template fields:") + for error in errors: + console.print(f"[red]Error:[/] {error}") + + sys.exit(len(errors)) From 76f072727bb8a6acf7beb606a5d7634adda7ca3b Mon Sep 17 00:00:00 2001 From: Daniel Standish <15932138+dstandish@users.noreply.github.com> Date: Fri, 27 Sep 2024 13:25:19 -0700 Subject: [PATCH 062/802] Remove DagRun.is_backfill attribute (#42548) This attribute is only used in one place and is not very useful. --- airflow/models/dagrun.py | 6 +----- newsfragments/42548.significant.rst | 1 + tests/jobs/test_scheduler_job.py | 5 ++--- 3 files changed, 4 insertions(+), 8 deletions(-) create mode 100644 newsfragments/42548.significant.rst diff --git a/airflow/models/dagrun.py b/airflow/models/dagrun.py index 3ef1c18f152a4..5d53e51763dff 100644 --- a/airflow/models/dagrun.py +++ b/airflow/models/dagrun.py @@ -1267,7 +1267,7 @@ def verify_integrity(self, *, session: Session = NEW_SESSION) -> None: def task_filter(task: Operator) -> bool: return task.task_id not in task_ids and ( - self.is_backfill + self.run_type == DagRunType.BACKFILL_JOB or (task.start_date is None or task.start_date <= self.execution_date) and (task.end_date is None or self.execution_date <= task.end_date) ) @@ -1538,10 +1538,6 @@ def _revise_map_indexes_if_mapped(self, task: Operator, *, session: Session) -> session.flush() yield ti - @property - def is_backfill(self) -> bool: - return self.run_type == DagRunType.BACKFILL_JOB - @classmethod @provide_session def get_latest_runs(cls, session: Session = NEW_SESSION) -> list[DagRun]: diff --git a/newsfragments/42548.significant.rst b/newsfragments/42548.significant.rst new file mode 100644 index 0000000000000..28d6795eebcc6 --- /dev/null +++ b/newsfragments/42548.significant.rst @@ -0,0 +1 @@ +Remove is_backfill attribute from DagRun object diff --git a/tests/jobs/test_scheduler_job.py b/tests/jobs/test_scheduler_job.py index 52e9dbdeb1a04..78a911153dab2 100644 --- a/tests/jobs/test_scheduler_job.py +++ b/tests/jobs/test_scheduler_job.py @@ -601,8 +601,7 @@ def test_execute_task_instances_backfill_tasks_wont_execute(self, dag_maker): ti1.state = State.SCHEDULED session.merge(ti1) session.flush() - - assert dr1.is_backfill + assert dr1.run_type == DagRunType.BACKFILL_JOB self.job_runner._critical_section_enqueue_task_instances(session) session.flush() @@ -3851,7 +3850,7 @@ def test_adopt_or_reset_orphaned_tasks_backfill_dag(self, dag_maker): session.merge(dr1) session.flush() - assert dr1.is_backfill + assert dr1.run_type == DagRunType.BACKFILL_JOB assert 0 == self.job_runner.adopt_or_reset_orphaned_tasks(session=session) session.rollback() From a5455d46243defa780f6084151341bfa7b53ed6d Mon Sep 17 00:00:00 2001 From: Jens Scheffler <95105677+jscheffl@users.noreply.github.com> Date: Sat, 28 Sep 2024 09:34:19 +0200 Subject: [PATCH 063/802] Attempt to correct dependency for Slack notification for canary build (#42551) --- .github/workflows/ci.yml | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 8625aee73d9e1..8828a30ce3ecd 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -673,6 +673,16 @@ jobs: notify-slack-failure: name: "Notify Slack on Failure" + needs: + - basic-tests + - additional-ci-image-checks + - providers + - tests-helm + - tests-special + - tests-with-lowest-direct-resolution + - additional-prod-image-tests + - tests-kubernetes + - finalize-tests if: github.event_name == 'schedule' && failure() runs-on: ["ubuntu-22.04"] steps: From 7ece767cf7dc2b3c8c104e40b7c1d48c774efbe8 Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Sat, 28 Sep 2024 15:41:38 +0100 Subject: [PATCH 064/802] Airflow 2.10.2 has been released (#42405) --- .github/ISSUE_TEMPLATE/airflow_bug_report.yml | 2 +- Dockerfile | 2 +- README.md | 12 +++--- RELEASE_NOTES.rst | 40 ++++++++++++++++++- airflow/reproducible_build.yaml | 4 +- .../installation/supported-versions.rst | 2 +- generated/PYPI_README.md | 10 ++--- scripts/ci/pre_commit/supported_versions.py | 2 +- 8 files changed, 55 insertions(+), 19 deletions(-) diff --git a/.github/ISSUE_TEMPLATE/airflow_bug_report.yml b/.github/ISSUE_TEMPLATE/airflow_bug_report.yml index 853b102ef07f8..f835c879f8380 100644 --- a/.github/ISSUE_TEMPLATE/airflow_bug_report.yml +++ b/.github/ISSUE_TEMPLATE/airflow_bug_report.yml @@ -25,7 +25,7 @@ body: the latest release or main to see if the issue is fixed before reporting it. multiple: false options: - - "2.10.1" + - "2.10.2" - "main (development)" - "Other Airflow 2 version (please specify below)" validations: diff --git a/Dockerfile b/Dockerfile index 68f1ed166f12a..3053e0779540d 100644 --- a/Dockerfile +++ b/Dockerfile @@ -45,7 +45,7 @@ ARG AIRFLOW_UID="50000" ARG AIRFLOW_USER_HOME_DIR=/home/airflow # latest released version here -ARG AIRFLOW_VERSION="2.10.1" +ARG AIRFLOW_VERSION="2.10.2" ARG PYTHON_BASE_IMAGE="python:3.8-slim-bookworm" diff --git a/README.md b/README.md index 91ddf5e927245..3169ac5144844 100644 --- a/README.md +++ b/README.md @@ -97,7 +97,7 @@ Airflow is not a streaming solution, but it is often used to process real-time d Apache Airflow is tested with: -| | Main version (dev) | Stable version (2.10.1) | +| | Main version (dev) | Stable version (2.10.2) | |------------|----------------------------|----------------------------| | Python | 3.8, 3.9, 3.10, 3.11, 3.12 | 3.8, 3.9, 3.10, 3.11, 3.12 | | Platform | AMD64/ARM64(\*) | AMD64/ARM64(\*) | @@ -177,15 +177,15 @@ them to the appropriate format and workflow that your tool requires. ```bash -pip install 'apache-airflow==2.10.1' \ - --constraint "https://raw.githubusercontent.com/apache/airflow/constraints-2.10.1/constraints-3.8.txt" +pip install 'apache-airflow==2.10.2' \ + --constraint "https://raw.githubusercontent.com/apache/airflow/constraints-2.10.2/constraints-3.8.txt" ``` 2. Installing with extras (i.e., postgres, google) ```bash -pip install 'apache-airflow[postgres,google]==2.10.1' \ - --constraint "https://raw.githubusercontent.com/apache/airflow/constraints-2.10.1/constraints-3.8.txt" +pip install 'apache-airflow[postgres,google]==2.10.2' \ + --constraint "https://raw.githubusercontent.com/apache/airflow/constraints-2.10.2/constraints-3.8.txt" ``` For information on installing provider packages, check @@ -290,7 +290,7 @@ Apache Airflow version life cycle: | Version | Current Patch/Minor | State | First Release | Limited Support | EOL/Terminated | |-----------|-----------------------|-----------|-----------------|-------------------|------------------| -| 2 | 2.10.1 | Supported | Dec 17, 2020 | TBD | TBD | +| 2 | 2.10.2 | Supported | Dec 17, 2020 | TBD | TBD | | 1.10 | 1.10.15 | EOL | Aug 27, 2018 | Dec 17, 2020 | June 17, 2021 | | 1.9 | 1.9.0 | EOL | Jan 03, 2018 | Aug 27, 2018 | Aug 27, 2018 | | 1.8 | 1.8.2 | EOL | Mar 19, 2017 | Jan 03, 2018 | Jan 03, 2018 | diff --git a/RELEASE_NOTES.rst b/RELEASE_NOTES.rst index d42074b1146a7..6c84e45d8aca0 100644 --- a/RELEASE_NOTES.rst +++ b/RELEASE_NOTES.rst @@ -21,6 +21,43 @@ .. towncrier release notes start +Airflow 2.10.2 (2024-09-18) +--------------------------- + +Significant Changes +^^^^^^^^^^^^^^^^^^^ + +No significant changes. + +Bug Fixes +""""""""" +- Revert "Fix: DAGs are not marked as stale if the dags folder change" (#42220, #42217) +- Add missing open telemetry span and correct scheduled slots documentation (#41985) +- Fix require_confirmation_dag_change (#42063) (#42211) +- Only treat null/undefined as falsy when rendering XComEntry (#42199) (#42213) +- Add extra and ``renderedTemplates`` as keys to skip ``camelCasing`` (#42206) (#42208) +- Do not ``camelcase`` xcom entries (#42182) (#42187) +- Fix task_instance and dag_run links from list views (#42138) (#42143) +- Support multi-line input for Params of type string in trigger UI form (#40414) (#42139) +- Fix details tab log url detection (#42104) (#42114) +- Add new type of exception to catch timeout (#42064) (#42078) +- Rewrite how DAG to dataset / dataset alias are stored (#41987) (#42055) +- Allow dataset alias to add more than one dataset events (#42189) (#42247) + +Miscellaneous +""""""""""""" +- Limit universal-pathlib below ``0.2.4`` as it breaks our integration (#42101) +- Auto-fix default deferrable with ``LibCST`` (#42089) +- Deprecate ``--tree`` flag for ``tasks list`` cli command (#41965) + +Doc Only Changes +"""""""""""""""" +- Update ``security_model.rst`` to clear unauthenticated endpoints exceptions (#42085) +- Add note about dataclasses and attrs to XComs page (#42056) +- Improve docs on markdown docs in DAGs (#42013) +- Add warning that listeners can be dangerous (#41968) + + Airflow 2.10.1 (2024-09-05) --------------------------- @@ -38,7 +75,7 @@ Bug Fixes - Fix compatibility with FAB provider versions <1.3.0 (#41809) - Don't Fail LocalTaskJob on heartbeat (#41810) - Remove deprecation warning for cgitb in Plugins Manager (#41793) -- Fix log for notifier(instance) without __name__ (#41699) +- Fix log for notifier(instance) without ``__name__`` (#41699) - Splitting syspath preparation into stages (#41694) - Adding url sanitization for extra links (#41680) - Fix InletEventsAccessors type stub (#41607) @@ -64,7 +101,6 @@ Doc Only Changes - Add an example for auth with ``keycloak`` (#41791) - Airflow 2.10.0 (2024-08-15) --------------------------- diff --git a/airflow/reproducible_build.yaml b/airflow/reproducible_build.yaml index 31e63fbce742b..1bf308b87a705 100644 --- a/airflow/reproducible_build.yaml +++ b/airflow/reproducible_build.yaml @@ -1,2 +1,2 @@ -release-notes-hash: aa948d55b0b6062659dbcd0293d73838 -source-date-epoch: 1725624671 +release-notes-hash: 828fa8d5e93e215963c0a3e52e7f1e3d +source-date-epoch: 1727075869 diff --git a/docs/apache-airflow/installation/supported-versions.rst b/docs/apache-airflow/installation/supported-versions.rst index 0a7694abbda3d..d82500728ce3b 100644 --- a/docs/apache-airflow/installation/supported-versions.rst +++ b/docs/apache-airflow/installation/supported-versions.rst @@ -29,7 +29,7 @@ Apache Airflow® version life cycle: ========= ===================== ========= =============== ================= ================ Version Current Patch/Minor State First Release Limited Support EOL/Terminated ========= ===================== ========= =============== ================= ================ -2 2.10.1 Supported Dec 17, 2020 TBD TBD +2 2.10.2 Supported Dec 17, 2020 TBD TBD 1.10 1.10.15 EOL Aug 27, 2018 Dec 17, 2020 June 17, 2021 1.9 1.9.0 EOL Jan 03, 2018 Aug 27, 2018 Aug 27, 2018 1.8 1.8.2 EOL Mar 19, 2017 Jan 03, 2018 Jan 03, 2018 diff --git a/generated/PYPI_README.md b/generated/PYPI_README.md index 2b80e73a45f5e..50802f301b753 100644 --- a/generated/PYPI_README.md +++ b/generated/PYPI_README.md @@ -54,7 +54,7 @@ Use Airflow to author workflows as directed acyclic graphs (DAGs) of tasks. The Apache Airflow is tested with: -| | Main version (dev) | Stable version (2.10.1) | +| | Main version (dev) | Stable version (2.10.2) | |------------|----------------------------|----------------------------| | Python | 3.8, 3.9, 3.10, 3.11, 3.12 | 3.8, 3.9, 3.10, 3.11, 3.12 | | Platform | AMD64/ARM64(\*) | AMD64/ARM64(\*) | @@ -130,15 +130,15 @@ them to the appropriate format and workflow that your tool requires. ```bash -pip install 'apache-airflow==2.10.1' \ - --constraint "https://raw.githubusercontent.com/apache/airflow/constraints-2.10.1/constraints-3.8.txt" +pip install 'apache-airflow==2.10.2' \ + --constraint "https://raw.githubusercontent.com/apache/airflow/constraints-2.10.2/constraints-3.8.txt" ``` 2. Installing with extras (i.e., postgres, google) ```bash -pip install 'apache-airflow[postgres,google]==2.10.1' \ - --constraint "https://raw.githubusercontent.com/apache/airflow/constraints-2.10.1/constraints-3.8.txt" +pip install 'apache-airflow[postgres,google]==2.10.2' \ + --constraint "https://raw.githubusercontent.com/apache/airflow/constraints-2.10.2/constraints-3.8.txt" ``` For information on installing provider packages, check diff --git a/scripts/ci/pre_commit/supported_versions.py b/scripts/ci/pre_commit/supported_versions.py index b392eaf6d4e01..ab8204ab03baa 100755 --- a/scripts/ci/pre_commit/supported_versions.py +++ b/scripts/ci/pre_commit/supported_versions.py @@ -27,7 +27,7 @@ HEADERS = ("Version", "Current Patch/Minor", "State", "First Release", "Limited Support", "EOL/Terminated") SUPPORTED_VERSIONS = ( - ("2", "2.10.1", "Supported", "Dec 17, 2020", "TBD", "TBD"), + ("2", "2.10.2", "Supported", "Dec 17, 2020", "TBD", "TBD"), ("1.10", "1.10.15", "EOL", "Aug 27, 2018", "Dec 17, 2020", "June 17, 2021"), ("1.9", "1.9.0", "EOL", "Jan 03, 2018", "Aug 27, 2018", "Aug 27, 2018"), ("1.8", "1.8.2", "EOL", "Mar 19, 2017", "Jan 03, 2018", "Jan 03, 2018"), From 4753dc4a257ca06944f01e98e14eb6f9b9407245 Mon Sep 17 00:00:00 2001 From: GPK Date: Sun, 29 Sep 2024 12:40:37 +0100 Subject: [PATCH 065/802] remove callable functions parameter from kafka operator template_fields (#42555) --- airflow/providers/apache/kafka/operators/consume.py | 1 - airflow/providers/apache/kafka/operators/produce.py | 1 - 2 files changed, 2 deletions(-) diff --git a/airflow/providers/apache/kafka/operators/consume.py b/airflow/providers/apache/kafka/operators/consume.py index 91d2f4f052daf..377b58a46df5e 100644 --- a/airflow/providers/apache/kafka/operators/consume.py +++ b/airflow/providers/apache/kafka/operators/consume.py @@ -68,7 +68,6 @@ class ConsumeFromTopicOperator(BaseOperator): ui_color = BLUE template_fields = ( "topics", - "apply_function", "apply_function_args", "apply_function_kwargs", "kafka_config_id", diff --git a/airflow/providers/apache/kafka/operators/produce.py b/airflow/providers/apache/kafka/operators/produce.py index 04090811b9ec7..e0623128a1f7d 100644 --- a/airflow/providers/apache/kafka/operators/produce.py +++ b/airflow/providers/apache/kafka/operators/produce.py @@ -67,7 +67,6 @@ class ProduceToTopicOperator(BaseOperator): template_fields = ( "topic", - "producer_function", "producer_function_args", "producer_function_kwargs", "kafka_config_id", From e13173c530a6f15e1ceaee58ed026f7fa8306005 Mon Sep 17 00:00:00 2001 From: Danny Liu Date: Sun, 29 Sep 2024 04:42:01 -0700 Subject: [PATCH 066/802] fix PyDocStyle checks (#42557) --- tests/www/views/test_views_task_norun.py | 2 +- tests/www/views/test_views_tasks.py | 4 ++-- tests/www/views/test_views_trigger_dag.py | 2 +- tests/www/views/test_views_variable.py | 2 +- 4 files changed, 5 insertions(+), 5 deletions(-) diff --git a/tests/www/views/test_views_task_norun.py b/tests/www/views/test_views_task_norun.py index a0709c4303d99..2a39b2a60134e 100644 --- a/tests/www/views/test_views_task_norun.py +++ b/tests/www/views/test_views_task_norun.py @@ -32,7 +32,7 @@ @pytest.fixture(scope="module", autouse=True) -def reset_dagruns(): +def _reset_dagruns(): """Clean up stray garbage from other tests.""" clear_db_runs() diff --git a/tests/www/views/test_views_tasks.py b/tests/www/views/test_views_tasks.py index 0b52c1f9aef3c..f5cc011fb6f0e 100644 --- a/tests/www/views/test_views_tasks.py +++ b/tests/www/views/test_views_tasks.py @@ -64,13 +64,13 @@ @pytest.fixture(scope="module", autouse=True) -def reset_dagruns(): +def _reset_dagruns(): """Clean up stray garbage from other tests.""" clear_db_runs() @pytest.fixture(autouse=True) -def init_dagruns(app): +def _init_dagruns(app): with time_machine.travel(DEFAULT_DATE, tick=False): triggered_by_kwargs = {"triggered_by": DagRunTriggeredByType.TEST} if AIRFLOW_V_3_0_PLUS else {} app.dag_bag.get_dag("example_bash_operator").create_dagrun( diff --git a/tests/www/views/test_views_trigger_dag.py b/tests/www/views/test_views_trigger_dag.py index 01b6713600af3..0c9384a195f5e 100644 --- a/tests/www/views/test_views_trigger_dag.py +++ b/tests/www/views/test_views_trigger_dag.py @@ -40,7 +40,7 @@ @pytest.fixture(autouse=True) -def initialize_one_dag(): +def _initialize_one_dag(): with create_session() as session: DagBag().get_dag("example_bash_operator").sync_to_db(session=session) yield diff --git a/tests/www/views/test_views_variable.py b/tests/www/views/test_views_variable.py index fcdad2bdb0bdd..a91a12ddc470b 100644 --- a/tests/www/views/test_views_variable.py +++ b/tests/www/views/test_views_variable.py @@ -43,7 +43,7 @@ @pytest.fixture(autouse=True) -def clear_variables(): +def _clear_variables(): with create_session() as session: session.query(Variable).delete() From b45259d124b465bf2ae7e0e4d23bb42c1a9175e8 Mon Sep 17 00:00:00 2001 From: Kunal Bhattacharya Date: Sun, 29 Sep 2024 21:55:20 +0530 Subject: [PATCH 067/802] Documentation change to highlight difference in usage between params and parameters attributes in SQLExecuteQueryOperator for Postgres (#42564) * Documentation change to call out the difference in usage between params and parameters attributes * Static checks fix --- .../operators/postgres_operator_howto_guide.rst | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/docs/apache-airflow-providers-postgres/operators/postgres_operator_howto_guide.rst b/docs/apache-airflow-providers-postgres/operators/postgres_operator_howto_guide.rst index f9dafe34196b1..09402178aa057 100644 --- a/docs/apache-airflow-providers-postgres/operators/postgres_operator_howto_guide.rst +++ b/docs/apache-airflow-providers-postgres/operators/postgres_operator_howto_guide.rst @@ -123,10 +123,11 @@ Passing Parameters into SQLExecuteQueryOperator for Postgres SQLExecuteQueryOperator provides ``parameters`` attribute which makes it possible to dynamically inject values into your SQL requests during runtime. The BaseOperator class has the ``params`` attribute which is available to the SQLExecuteQueryOperator -by virtue of inheritance. Both ``parameters`` and ``params`` make it possible to dynamically pass in parameters in many -interesting ways. +by virtue of inheritance. While both ``parameters`` and ``params`` make it possible to dynamically pass in parameters in many +interesting ways, their usage is slightly different as demonstrated in the examples below. -To find the owner of the pet called 'Lester': +To find the birth dates of all pets between two dates, when we use the SQL statements directly in our code, we will use the +``parameters`` attribute: .. code-block:: python @@ -137,16 +138,15 @@ To find the owner of the pet called 'Lester': parameters={"begin_date": "2020-01-01", "end_date": "2020-12-31"}, ) -Now lets refactor our ``get_birth_date`` task. Instead of dumping SQL statements directly into our code, let's tidy things up -by creating a sql file. +Now lets refactor our ``get_birth_date`` task. Now, instead of dumping SQL statements directly into our code, let's tidy things up +by creating a sql file. And this time we will use the ``params`` attribute which we get for free from the parent ``BaseOperator`` +class. :: -- dags/sql/birth_date.sql SELECT * FROM pet WHERE birth_date BETWEEN SYMMETRIC {{ params.begin_date }} AND {{ params.end_date }}; -And this time we will use the ``params`` attribute which we get for free from the parent ``BaseOperator`` -class. .. code-block:: python From a81cedc1532ae5299d2949f3b5292dd91341db72 Mon Sep 17 00:00:00 2001 From: Jens Scheffler <95105677+jscheffl@users.noreply.github.com> Date: Sun, 29 Sep 2024 23:02:07 +0200 Subject: [PATCH 068/802] Bugfix/42575 workaround pin azure kusto data (#42576) * Workaround, pin azure-kusto-data to not 4.6.0 --- airflow/providers/microsoft/azure/provider.yaml | 3 ++- generated/provider_dependencies.json | 2 +- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/airflow/providers/microsoft/azure/provider.yaml b/airflow/providers/microsoft/azure/provider.yaml index 9110a3046b5d1..45fe28eecffc7 100644 --- a/airflow/providers/microsoft/azure/provider.yaml +++ b/airflow/providers/microsoft/azure/provider.yaml @@ -101,7 +101,8 @@ dependencies: - azure-synapse-artifacts>=0.17.0 - adal>=1.2.7 - azure-storage-file-datalake>=12.9.1 - - azure-kusto-data>=4.1.0 + # azure-kusto-data 4.6.0 breaks main - see https://github.com/apache/airflow/issues/42575 + - azure-kusto-data>=4.1.0,!=4.6.0 - azure-mgmt-datafactory>=2.0.0 - azure-mgmt-containerregistry>=8.0.0 - azure-mgmt-containerinstance>=10.1.0 diff --git a/generated/provider_dependencies.json b/generated/provider_dependencies.json index 074c5dd41e93b..b9bc363b15e33 100644 --- a/generated/provider_dependencies.json +++ b/generated/provider_dependencies.json @@ -804,7 +804,7 @@ "azure-datalake-store>=0.0.45", "azure-identity>=1.3.1", "azure-keyvault-secrets>=4.1.0", - "azure-kusto-data>=4.1.0", + "azure-kusto-data>=4.1.0,!=4.6.0", "azure-mgmt-containerinstance>=10.1.0", "azure-mgmt-containerregistry>=8.0.0", "azure-mgmt-cosmosdb>=3.0.0", From 3645539c4191d7fd313ee438f79dfcde4d00ed42 Mon Sep 17 00:00:00 2001 From: Topher Anderson <48180628+topherinternational@users.noreply.github.com> Date: Sun, 29 Sep 2024 19:21:11 -0500 Subject: [PATCH 069/802] Bump uv to 0.4.17 (#42574) --- Dockerfile | 2 +- Dockerfile.ci | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/Dockerfile b/Dockerfile index 3053e0779540d..cfb894ac87d22 100644 --- a/Dockerfile +++ b/Dockerfile @@ -50,7 +50,7 @@ ARG AIRFLOW_VERSION="2.10.2" ARG PYTHON_BASE_IMAGE="python:3.8-slim-bookworm" ARG AIRFLOW_PIP_VERSION=24.2 -ARG AIRFLOW_UV_VERSION=0.4.7 +ARG AIRFLOW_UV_VERSION=0.4.17 ARG AIRFLOW_USE_UV="false" ARG UV_HTTP_TIMEOUT="300" ARG AIRFLOW_IMAGE_REPOSITORY="https://github.com/apache/airflow" diff --git a/Dockerfile.ci b/Dockerfile.ci index ad944d151adcb..f7b7bb4172025 100644 --- a/Dockerfile.ci +++ b/Dockerfile.ci @@ -1262,7 +1262,7 @@ ARG DEFAULT_CONSTRAINTS_BRANCH="constraints-main" ARG AIRFLOW_CI_BUILD_EPOCH="10" ARG AIRFLOW_PRE_CACHED_PIP_PACKAGES="true" ARG AIRFLOW_PIP_VERSION=24.2 -ARG AIRFLOW_UV_VERSION=0.4.7 +ARG AIRFLOW_UV_VERSION=0.4.17 ARG AIRFLOW_USE_UV="true" # Setup PIP # By default PIP install run without cache to make image smaller From 161786c46d7724d230004e27bd934eaedf71fe19 Mon Sep 17 00:00:00 2001 From: Tzu-ping Chung Date: Sun, 29 Sep 2024 21:43:44 -0700 Subject: [PATCH 070/802] Change default .airflowignore syntax to glob (#42436) Co-authored-by: Shahar Epstein <60007259+shahar1@users.noreply.github.com> --- airflow/config_templates/config.yml | 2 +- airflow/utils/file.py | 6 ++-- .../modules_management.rst | 9 +----- docs/apache-airflow/core-concepts/dags.rst | 31 +++++++------------ .../howto/dynamic-dag-generation.rst | 2 +- newsfragments/42436.significant.rst | 7 +++++ tests/dags/.airflowignore | 5 ++- tests/dags/subdir1/.airflowignore | 2 +- tests/plugins/test_plugin_ignore.py | 2 +- 9 files changed, 29 insertions(+), 37 deletions(-) create mode 100644 newsfragments/42436.significant.rst diff --git a/airflow/config_templates/config.yml b/airflow/config_templates/config.yml index c9abee3c85065..7317fce60e4e6 100644 --- a/airflow/config_templates/config.yml +++ b/airflow/config_templates/config.yml @@ -310,7 +310,7 @@ core: version_added: 2.3.0 type: string example: ~ - default: "regexp" + default: "glob" default_task_retries: description: | The number of retries each task is going to have by default. Can be overridden at dag or task level. diff --git a/airflow/utils/file.py b/airflow/utils/file.py index 7081113d5bd46..2e39eb7dd7b52 100644 --- a/airflow/utils/file.py +++ b/airflow/utils/file.py @@ -221,7 +221,7 @@ def _find_path_from_directory( def find_path_from_directory( base_dir_path: str | os.PathLike[str], ignore_file_name: str, - ignore_file_syntax: str = conf.get_mandatory_value("core", "DAG_IGNORE_FILE_SYNTAX", fallback="regexp"), + ignore_file_syntax: str = conf.get_mandatory_value("core", "DAG_IGNORE_FILE_SYNTAX", fallback="glob"), ) -> Generator[str, None, None]: """ Recursively search the base path for a list of file paths that should not be ignored. @@ -232,9 +232,9 @@ def find_path_from_directory( :return: a generator of file paths. """ - if ignore_file_syntax == "glob": + if ignore_file_syntax == "glob" or not ignore_file_syntax: return _find_path_from_directory(base_dir_path, ignore_file_name, _GlobIgnoreRule) - elif ignore_file_syntax == "regexp" or not ignore_file_syntax: + elif ignore_file_syntax == "regexp": return _find_path_from_directory(base_dir_path, ignore_file_name, _RegexpIgnoreRule) else: raise ValueError(f"Unsupported ignore_file_syntax: {ignore_file_syntax}") diff --git a/docs/apache-airflow/administration-and-deployment/modules_management.rst b/docs/apache-airflow/administration-and-deployment/modules_management.rst index dc6be49b1d43d..25adb5f333c91 100644 --- a/docs/apache-airflow/administration-and-deployment/modules_management.rst +++ b/docs/apache-airflow/administration-and-deployment/modules_management.rst @@ -125,14 +125,7 @@ for the paths that should be ignored. You do not need to have that file in any o In the example above the DAGs are only in ``my_custom_dags`` folder, the ``common_package`` should not be scanned by scheduler when searching for DAGS, so we should ignore ``common_package`` folder. You also want to ignore the ``base_dag.py`` if you keep a base DAG there that ``my_dag1.py`` and ``my_dag2.py`` derives -from. Your ``.airflowignore`` should look then like this: - -.. code-block:: none - - my_company/common_package/.* - my_company/my_custom_dags/base_dag\.py - -If ``DAG_IGNORE_FILE_SYNTAX`` is set to ``glob``, the equivalent ``.airflowignore`` file would be: +from. Your ``.airflowignore`` should look then like this (using the default ``glob`` syntax): .. code-block:: none diff --git a/docs/apache-airflow/core-concepts/dags.rst b/docs/apache-airflow/core-concepts/dags.rst index fbef745e46d34..f9dc7d64c72e0 100644 --- a/docs/apache-airflow/core-concepts/dags.rst +++ b/docs/apache-airflow/core-concepts/dags.rst @@ -712,19 +712,9 @@ configuration parameter (*added in Airflow 2.3*): ``regexp`` and ``glob``. .. note:: - The default ``DAG_IGNORE_FILE_SYNTAX`` is ``regexp`` to ensure backwards compatibility. + The default ``DAG_IGNORE_FILE_SYNTAX`` is ``glob`` in Airflow 3 or later (in previous versions it was ``regexp``). -For the ``regexp`` pattern syntax (the default), each line in ``.airflowignore`` -specifies a regular expression pattern, and directories or files whose names (not DAG id) -match any of the patterns would be ignored (under the hood, ``Pattern.search()`` is used -to match the pattern). Use the ``#`` character to indicate a comment; all characters -on lines starting with ``#`` will be ignored. - -As with most regexp matching in Airflow, the regexp engine is ``re2``, which explicitly -doesn't support many advanced features, please check its -`documentation `_ for more information. - -With the ``glob`` syntax, the patterns work just like those in a ``.gitignore`` file: +With the ``glob`` syntax (the default), the patterns work just like those in a ``.gitignore`` file: * The ``*`` character will match any number of characters, except ``/`` * The ``?`` character will match any single character, except ``/`` @@ -738,15 +728,18 @@ With the ``glob`` syntax, the patterns work just like those in a ``.gitignore`` is relative to the directory level of the particular .airflowignore file itself. Otherwise the pattern may also match at any level below the .airflowignore level. -The ``.airflowignore`` file should be put in your ``DAG_FOLDER``. For example, you can prepare -a ``.airflowignore`` file using the ``regexp`` syntax with content - -.. code-block:: +For the ``regexp`` pattern syntax, each line in ``.airflowignore`` +specifies a regular expression pattern, and directories or files whose names (not DAG id) +match any of the patterns would be ignored (under the hood, ``Pattern.search()`` is used +to match the pattern). Use the ``#`` character to indicate a comment; all characters +on lines starting with ``#`` will be ignored. - project_a - tenant_[\d] +As with most regexp matching in Airflow, the regexp engine is ``re2``, which explicitly +doesn't support many advanced features, please check its +`documentation `_ for more information. -Or, equivalently, in the ``glob`` syntax +The ``.airflowignore`` file should be put in your ``DAG_FOLDER``. For example, you can prepare +a ``.airflowignore`` file with the ``glob`` syntax .. code-block:: diff --git a/docs/apache-airflow/howto/dynamic-dag-generation.rst b/docs/apache-airflow/howto/dynamic-dag-generation.rst index 5d542a29320b7..9aa988f28bdb1 100644 --- a/docs/apache-airflow/howto/dynamic-dag-generation.rst +++ b/docs/apache-airflow/howto/dynamic-dag-generation.rst @@ -91,7 +91,7 @@ Then you can import and use the ``ALL_TASKS`` constant in all your DAGs like tha ... Don't forget that in this case you need to add empty ``__init__.py`` file in the ``my_company_utils`` folder -and you should add the ``my_company_utils/.*`` line to ``.airflowignore`` file (if using the regexp ignore +and you should add the ``my_company_utils/*`` line to ``.airflowignore`` file (using the default glob syntax), so that the whole folder is ignored by the scheduler when it looks for DAGs. diff --git a/newsfragments/42436.significant.rst b/newsfragments/42436.significant.rst new file mode 100644 index 0000000000000..d9dbcfc4c9f5d --- /dev/null +++ b/newsfragments/42436.significant.rst @@ -0,0 +1,7 @@ +Default ``.airflowignore`` syntax changed to ``glob`` + +The default value to the configuration ``[core] dag_ignore_file_syntax`` has +been changed to ``glob``, which better matches the ignore file behavior of many +popular tools. + +To revert to the previous behavior, set the configuration to ``regexp``. diff --git a/tests/dags/.airflowignore b/tests/dags/.airflowignore index 313b04ef81cd4..7daaf22e65efc 100644 --- a/tests/dags/.airflowignore +++ b/tests/dags/.airflowignore @@ -1,3 +1,2 @@ -.*_invalid.* # Skip invalid files -subdir3 # Skip the nested subdir3 directory -# *badrule # This rule is an invalid regex. It would be warned about and skipped. +*_invalid_* # Skip invalid files +subdir3 # Skip the nested subdir3 directory diff --git a/tests/dags/subdir1/.airflowignore b/tests/dags/subdir1/.airflowignore index 8b69a752e69fb..0bfa43be300a1 100644 --- a/tests/dags/subdir1/.airflowignore +++ b/tests/dags/subdir1/.airflowignore @@ -1 +1 @@ -.*_ignore_this.py # Ignore files ending with "_ignore_this.py" +*_ignore_this.py # Ignore files ending with "_ignore_this.py" diff --git a/tests/plugins/test_plugin_ignore.py b/tests/plugins/test_plugin_ignore.py index d995fabd080f8..92951304d2b9f 100644 --- a/tests/plugins/test_plugin_ignore.py +++ b/tests/plugins/test_plugin_ignore.py @@ -77,7 +77,7 @@ def test_find_not_should_ignore_path_regexp(self, tmp_path): "test_load_sub1.py", } ignore_list_file = ".airflowignore" - for file_path in find_path_from_directory(plugin_folder_path, ignore_list_file): + for file_path in find_path_from_directory(plugin_folder_path, ignore_list_file, "regexp"): file_path = Path(file_path) if file_path.is_file() and file_path.suffix == ".py": detected_files.add(file_path.name) From 9a3832011a1c784dd393532db2a1bad6164b6917 Mon Sep 17 00:00:00 2001 From: Wei Lee Date: Mon, 30 Sep 2024 14:30:24 +0900 Subject: [PATCH 071/802] Rename dataset related python variable names to asset (#41348) --- RELEASE_NOTES.rst | 4 +- airflow/__init__.py | 6 +- .../endpoints/dag_run_endpoint.py | 16 +- .../endpoints/dataset_endpoint.py | 190 +++--- .../{dataset_schema.py => asset_schema.py} | 88 +-- airflow/api_connexion/security.py | 8 +- airflow/api_fastapi/openapi/v1-generated.yaml | 8 +- airflow/api_fastapi/views/ui/__init__.py | 4 +- .../views/ui/{datasets.py => assets.py} | 34 +- .../endpoints/rpc_api_endpoint.py | 8 +- airflow/{datasets => assets}/__init__.py | 206 +++--- airflow/{datasets => assets}/manager.py | 195 +++--- airflow/{datasets => assets}/metadata.py | 10 +- airflow/auth/managers/base_auth_manager.py | 10 +- .../auth/managers/models/resource_details.py | 4 +- .../managers/simple/simple_auth_manager.py | 6 +- airflow/config_templates/config.yml | 18 +- airflow/dag_processing/collection.py | 147 +++-- airflow/decorators/base.py | 10 +- airflow/example_dags/example_asset_alias.py | 101 +++ .../example_asset_alias_with_no_taskflow.py | 108 ++++ airflow/example_dags/example_assets.py | 192 ++++++ airflow/example_dags/example_dataset_alias.py | 101 --- .../example_dataset_alias_with_no_taskflow.py | 108 ---- airflow/example_dags/example_datasets.py | 192 ------ .../example_dags/example_inlet_event_extra.py | 22 +- .../example_outlet_event_extra.py | 28 +- airflow/io/path.py | 12 +- airflow/jobs/scheduler_job_runner.py | 90 ++- airflow/lineage/__init__.py | 2 +- airflow/lineage/hook.py | 126 ++-- airflow/listeners/listener.py | 4 +- .../listeners/spec/{dataset.py => asset.py} | 18 +- airflow/models/__init__.py | 2 +- airflow/models/{dataset.py => asset.py} | 104 ++-- airflow/models/dag.py | 91 +-- airflow/models/taskinstance.py | 74 +-- airflow/operators/python.py | 6 +- airflow/provider.yaml.schema.json | 30 +- .../aws/{datasets => assets}/__init__.py | 0 .../amazon/aws/{datasets => assets}/s3.py | 16 +- .../amazon/aws/auth_manager/avp/entities.py | 2 +- .../aws/auth_manager/aws_auth_manager.py | 16 +- airflow/providers/amazon/aws/hooks/s3.py | 39 +- .../utils/asset_compat_lineage_collector.py | 106 ++++ airflow/providers/amazon/provider.yaml | 14 +- .../common/compat/assets/__init__.py | 77 +++ .../providers/common/compat/lineage/hook.py | 73 ++- .../openlineage/utils}/__init__.py | 0 .../common/compat/openlineage/utils/utils.py | 43 ++ .../compat/security}/__init__.py | 0 .../common/compat/security/permissions.py | 30 + .../datasets => common/io/assets}/__init__.py | 0 .../common/io/{datasets => assets}/file.py | 15 +- airflow/providers/common/io/provider.yaml | 14 +- .../fab/auth_manager/fab_auth_manager.py | 16 +- .../auth_manager/security_manager/override.py | 16 +- airflow/providers/fab/provider.yaml | 1 + airflow/providers/google/provider.yaml | 8 + .../datasets => mysql/assets}/__init__.py | 0 .../mysql/{datasets => assets}/mysql.py | 0 airflow/providers/mysql/provider.yaml | 8 +- .../openlineage/extractors/manager.py | 27 +- .../utils/asset_compat_lineage_collector.py | 108 ++++ airflow/providers/openlineage/utils/utils.py | 38 +- .../providers/postgres/assets}/__init__.py | 0 .../postgres/{datasets => assets}/postgres.py | 0 airflow/providers/postgres/provider.yaml | 8 +- .../providers/trino/assets}/__init__.py | 0 .../trino/{datasets => assets}/trino.py | 0 airflow/providers/trino/provider.yaml | 8 +- airflow/providers_manager.py | 48 +- airflow/reproducible_build.yaml | 4 +- airflow/security/permissions.py | 2 +- airflow/serialization/dag_dependency.py | 2 +- airflow/serialization/enums.py | 13 +- .../pydantic/{dataset.py => asset.py} | 22 +- airflow/serialization/pydantic/dag_run.py | 4 +- .../serialization/pydantic/taskinstance.py | 4 +- airflow/serialization/serialized_objects.py | 111 ++-- airflow/timetables/{datasets.py => assets.py} | 36 +- airflow/timetables/base.py | 24 +- airflow/timetables/simple.py | 40 +- airflow/ui/openapi-gen/queries/common.ts | 18 +- airflow/ui/openapi-gen/queries/prefetch.ts | 36 +- airflow/ui/openapi-gen/queries/queries.ts | 21 +- airflow/ui/openapi-gen/queries/suspense.ts | 52 +- .../ui/openapi-gen/requests/services.gen.ts | 14 +- airflow/ui/openapi-gen/requests/types.gen.ts | 6 +- airflow/utils/context.py | 102 +-- airflow/utils/context.pyi | 28 +- airflow/utils/operator_helpers.py | 2 +- airflow/www/auth.py | 6 +- airflow/www/security_manager.py | 4 +- airflow/www/static/css/graph.css | 8 +- .../www/static/js/dag/details/graph/Node.tsx | 2 +- .../www/static/js/dag/details/graph/index.tsx | 6 +- .../www/static/js/dag/details/graph/utils.ts | 2 +- airflow/www/static/js/datasets/Graph/Node.tsx | 4 +- .../www/static/js/datasets/Graph/index.tsx | 2 +- airflow/www/static/js/datasets/SearchBar.tsx | 2 +- airflow/www/static/js/types/index.ts | 4 +- airflow/www/templates/airflow/dag.html | 16 +- .../templates/airflow/dag_dependencies.html | 4 +- airflow/www/templates/airflow/dags.html | 22 +- airflow/www/views.py | 102 ++- dev/breeze/tests/test_packages.py | 3 + .../tests/test_pytest_args_for_test_types.py | 2 +- dev/breeze/tests/test_selective_checks.py | 26 +- .../auth-manager/manage/index.rst | 12 +- .../auth-manager/access-control.rst | 10 +- ...{dataset-schemes.rst => asset-schemes.rst} | 8 +- .../howto/create-custom-providers.rst | 4 +- .../administration-and-deployment/lineage.rst | 12 +- .../listeners.rst | 8 +- .../logging-monitoring/metrics.rst | 6 +- .../authoring-and-scheduling/assets.rst | 532 ++++++++++++++++ .../authoring-and-scheduling/datasets.rst | 532 ---------------- .../authoring-and-scheduling/index.rst | 2 +- .../authoring-and-scheduling/timetable.rst | 18 +- docs/apache-airflow/core-concepts/dag-run.rst | 4 +- .../apache-airflow/core-concepts/taskflow.rst | 14 +- ...uled-dags.png => asset-scheduled-dags.png} | Bin .../img/{datasets.png => assets.png} | Bin docs/apache-airflow/templates-ref.rst | 12 +- .../apache-airflow/tutorial/objectstorage.rst | 2 +- docs/apache-airflow/ui.rst | 4 +- docs/exts/operators_and_hooks_ref.py | 6 +- ...st.jinja2 => asset-uri-schemes.rst.jinja2} | 2 +- docs/spelling_wordlist.txt | 2 + generated/provider_dependencies.json | 10 +- newsfragments/41348.significant.rst | 240 +++++++ .../check_tests_in_right_folders.py | 2 +- scripts/cov/core_coverage.py | 2 +- scripts/cov/other_coverage.py | 4 +- tests/always/test_project_structure.py | 2 + .../endpoints/test_dag_run_endpoint.py | 22 +- .../endpoints/test_dag_source_endpoint.py | 2 +- .../endpoints/test_dataset_endpoint.py | 170 ++--- .../api_connexion/schemas/test_dag_schema.py | 12 +- .../schemas/test_dataset_schema.py | 90 ++- .../ui/{test_datasets.py => test_assets.py} | 4 +- .../common/io/datasets => assets}/__init__.py | 0 tests/{datasets => assets}/test_manager.py | 104 ++-- tests/assets/tests_asset.py | 586 +++++++++++++++++ .../simple/test_simple_auth_manager.py | 6 +- tests/auth/managers/test_base_auth_manager.py | 6 +- tests/conftest.py | 4 +- .../dags/{test_datasets.py => test_assets.py} | 6 +- tests/dags/test_only_empty_tasks.py | 4 +- tests/datasets/test_dataset.py | 588 ------------------ tests/decorators/test_python.py | 10 +- tests/io/test_path.py | 32 +- tests/io/test_wrapper.py | 14 +- tests/jobs/test_scheduler_job.py | 186 +++--- tests/lineage/test_hook.py | 140 ++--- ...{dataset_listener.py => asset_listener.py} | 14 +- ...set_listener.py => test_asset_listener.py} | 26 +- .../models/{test_dataset.py => test_asset.py} | 12 +- tests/models/test_dag.py | 326 +++++----- tests/models/test_dagrun.py | 2 +- tests/models/test_serialized_dag.py | 14 +- tests/models/test_taskinstance.py | 453 +++++++------- tests/operators/test_python.py | 4 +- .../aws/assets}/__init__.py | 0 .../aws/{datasets => assets}/test_s3.py | 22 +- .../aws/auth_manager/test_aws_auth_manager.py | 53 +- tests/providers/amazon/aws/hooks/test_s3.py | 50 +- .../compat/openlineage/utils}/__init__.py | 0 .../compat/openlineage/utils/test_utils.py | 23 + .../compat/security}/__init__.py | 0 .../compat/security/test_permissions.py | 23 + tests/providers/common/io/assets/__init__.py | 16 + .../io/{datasets => assets}/test_file.py | 16 +- .../fab/auth_manager/test_fab_auth_manager.py | 12 +- .../fab/auth_manager/test_security.py | 14 +- tests/providers/mysql/assets/__init__.py | 16 + .../mysql/{datasets => assets}/test_mysql.py | 2 +- .../openlineage/extractors/test_manager.py | 35 +- tests/providers/postgres/assets/__init__.py | 16 + .../{datasets => assets}/test_postgres.py | 2 +- tests/providers/trino/assets/__init__.py | 16 + .../trino/{datasets => assets}/test_trino.py | 2 +- tests/serialization/test_dag_serialization.py | 86 +-- tests/serialization/test_pydantic_models.py | 52 +- tests/serialization/test_serde.py | 6 +- .../serialization/test_serialized_objects.py | 48 +- .../microsoft/azure/example_msfabric.py | 4 +- .../providers/papermill/input_notebook.ipynb | 2 + tests/test_utils/compat.py | 33 + tests/test_utils/db.py | 35 +- ..._timetable.py => test_assets_timetable.py} | 134 ++-- tests/utils/test_context.py | 40 +- tests/utils/test_db_cleanup.py | 6 +- tests/utils/test_json.py | 10 +- tests/www/test_auth.py | 2 +- tests/www/views/test_views_acl.py | 8 +- tests/www/views/test_views_dataset.py | 197 +++--- tests/www/views/test_views_grid.py | 50 +- 199 files changed, 5061 insertions(+), 4127 deletions(-) rename airflow/api_connexion/schemas/{dataset_schema.py => asset_schema.py} (65%) rename airflow/api_fastapi/views/ui/{datasets.py => assets.py} (66%) rename airflow/{datasets => assets}/__init__.py (60%) rename airflow/{datasets => assets}/manager.py (53%) rename airflow/{datasets => assets}/metadata.py (80%) create mode 100644 airflow/example_dags/example_asset_alias.py create mode 100644 airflow/example_dags/example_asset_alias_with_no_taskflow.py create mode 100644 airflow/example_dags/example_assets.py delete mode 100644 airflow/example_dags/example_dataset_alias.py delete mode 100644 airflow/example_dags/example_dataset_alias_with_no_taskflow.py delete mode 100644 airflow/example_dags/example_datasets.py rename airflow/listeners/spec/{dataset.py => asset.py} (76%) rename airflow/models/{dataset.py => asset.py} (83%) rename airflow/providers/amazon/aws/{datasets => assets}/__init__.py (100%) rename airflow/providers/amazon/aws/{datasets => assets}/s3.py (73%) create mode 100644 airflow/providers/amazon/aws/utils/asset_compat_lineage_collector.py create mode 100644 airflow/providers/common/compat/assets/__init__.py rename airflow/providers/common/{io/datasets => compat/openlineage/utils}/__init__.py (100%) create mode 100644 airflow/providers/common/compat/openlineage/utils/utils.py rename airflow/providers/{mysql/datasets => common/compat/security}/__init__.py (100%) create mode 100644 airflow/providers/common/compat/security/permissions.py rename airflow/providers/{postgres/datasets => common/io/assets}/__init__.py (100%) rename airflow/providers/common/io/{datasets => assets}/file.py (76%) rename airflow/providers/{trino/datasets => mysql/assets}/__init__.py (100%) rename airflow/providers/mysql/{datasets => assets}/mysql.py (100%) create mode 100644 airflow/providers/openlineage/utils/asset_compat_lineage_collector.py rename {tests/datasets => airflow/providers/postgres/assets}/__init__.py (100%) rename airflow/providers/postgres/{datasets => assets}/postgres.py (100%) rename {tests/providers/amazon/aws/datasets => airflow/providers/trino/assets}/__init__.py (100%) rename airflow/providers/trino/{datasets => assets}/trino.py (100%) rename airflow/serialization/pydantic/{dataset.py => asset.py} (68%) rename airflow/timetables/{datasets.py => assets.py} (71%) rename docs/apache-airflow-providers/core-extensions/{dataset-schemes.rst => asset-schemes.rst} (82%) create mode 100644 docs/apache-airflow/authoring-and-scheduling/assets.rst delete mode 100644 docs/apache-airflow/authoring-and-scheduling/datasets.rst rename docs/apache-airflow/img/{dataset-scheduled-dags.png => asset-scheduled-dags.png} (100%) rename docs/apache-airflow/img/{datasets.png => assets.png} (100%) rename docs/exts/templates/{dataset-uri-schemes.rst.jinja2 => asset-uri-schemes.rst.jinja2} (95%) create mode 100644 newsfragments/41348.significant.rst rename tests/api_fastapi/views/ui/{test_datasets.py => test_assets.py} (92%) rename tests/{providers/common/io/datasets => assets}/__init__.py (100%) rename tests/{datasets => assets}/test_manager.py (52%) create mode 100644 tests/assets/tests_asset.py rename tests/dags/{test_datasets.py => test_assets.py} (91%) delete mode 100644 tests/datasets/test_dataset.py rename tests/listeners/{dataset_listener.py => asset_listener.py} (80%) rename tests/listeners/{test_dataset_listener.py => test_asset_listener.py} (72%) rename tests/models/{test_dataset.py => test_asset.py} (73%) rename tests/providers/{mysql/datasets => amazon/aws/assets}/__init__.py (100%) rename tests/providers/amazon/aws/{datasets => assets}/test_s3.py (75%) rename tests/providers/{postgres/datasets => common/compat/openlineage/utils}/__init__.py (100%) create mode 100644 tests/providers/common/compat/openlineage/utils/test_utils.py rename tests/providers/{trino/datasets => common/compat/security}/__init__.py (100%) create mode 100644 tests/providers/common/compat/security/test_permissions.py create mode 100644 tests/providers/common/io/assets/__init__.py rename tests/providers/common/io/{datasets => assets}/test_file.py (83%) create mode 100644 tests/providers/mysql/assets/__init__.py rename tests/providers/mysql/{datasets => assets}/test_mysql.py (97%) create mode 100644 tests/providers/postgres/assets/__init__.py rename tests/providers/postgres/{datasets => assets}/test_postgres.py (97%) create mode 100644 tests/providers/trino/assets/__init__.py rename tests/providers/trino/{datasets => assets}/test_trino.py (97%) rename tests/timetables/{test_datasets_timetable.py => test_assets_timetable.py} (57%) diff --git a/RELEASE_NOTES.rst b/RELEASE_NOTES.rst index 6c84e45d8aca0..69e666461efae 100644 --- a/RELEASE_NOTES.rst +++ b/RELEASE_NOTES.rst @@ -642,7 +642,7 @@ Dataset URIs are now validated on input (#37005) Datasets must use a URI that conform to rules laid down in AIP-60, and the value will be automatically normalized when the DAG file is parsed. See -`documentation on Datasets `_ for +`documentation on Datasets `_ for a more detailed description on the rules. You may need to change your Dataset identifiers if they look like a URI, but are @@ -3264,7 +3264,7 @@ If you have the producer and consumer in different files you do not need to use Datasets represent the abstract concept of a dataset, and (for now) do not have any direct read or write capability - in this release we are adding the foundational feature that we will build upon. -For more info on Datasets please see :doc:`/authoring-and-scheduling/datasets`. +For more info on Datasets please see `Datasets documentation `_. Expanded dynamic task mapping support """"""""""""""""""""""""""""""""""""" diff --git a/airflow/__init__.py b/airflow/__init__.py index 8930f190130a7..18f4cc3e3c28c 100644 --- a/airflow/__init__.py +++ b/airflow/__init__.py @@ -55,7 +55,7 @@ __all__ = [ "__version__", "DAG", - "Dataset", + "Asset", "XComArg", ] @@ -76,7 +76,7 @@ # Things to lazy import in form {local_name: ('target_module', 'target_name', 'deprecated')} __lazy_imports: dict[str, tuple[str, str, bool]] = { "DAG": (".models.dag", "DAG", False), - "Dataset": (".datasets", "Dataset", False), + "Asset": (".assets", "Asset", False), "XComArg": (".models.xcom_arg", "XComArg", False), "version": (".version", "", False), # Deprecated lazy imports @@ -86,8 +86,8 @@ # These objects are imported by PEP-562, however, static analyzers and IDE's # have no idea about typing of these objects. # Add it under TYPE_CHECKING block should help with it. + from airflow.models.asset import Asset from airflow.models.dag import DAG - from airflow.models.dataset import Dataset from airflow.models.xcom_arg import XComArg diff --git a/airflow/api_connexion/endpoints/dag_run_endpoint.py b/airflow/api_connexion/endpoints/dag_run_endpoint.py index 02847f0a00e92..02d4663837f4e 100644 --- a/airflow/api_connexion/endpoints/dag_run_endpoint.py +++ b/airflow/api_connexion/endpoints/dag_run_endpoint.py @@ -39,6 +39,10 @@ format_datetime, format_parameters, ) +from airflow.api_connexion.schemas.asset_schema import ( + AssetEventCollection, + asset_event_collection_schema, +) from airflow.api_connexion.schemas.dag_run_schema import ( DAGRunCollection, DAGRunCollectionSchema, @@ -50,10 +54,6 @@ set_dagrun_note_form_schema, set_dagrun_state_form_schema, ) -from airflow.api_connexion.schemas.dataset_schema import ( - DatasetEventCollection, - dataset_event_collection_schema, -) from airflow.api_connexion.schemas.task_instance_schema import ( TaskInstanceReferenceCollection, task_instance_reference_collection_schema, @@ -112,12 +112,12 @@ def get_dag_run( @security.requires_access_dag("GET", DagAccessEntity.RUN) -@security.requires_access_dataset("GET") +@security.requires_access_asset("GET") @provide_session def get_upstream_dataset_events( *, dag_id: str, dag_run_id: str, session: Session = NEW_SESSION ) -> APIResponse: - """If dag run is dataset-triggered, return the dataset events that triggered it.""" + """If dag run is dataset-triggered, return the asset events that triggered it.""" dag_run: DagRun | None = session.scalar( select(DagRun).where( DagRun.dag_id == dag_id, @@ -130,8 +130,8 @@ def get_upstream_dataset_events( detail=f"DAGRun with DAG ID: '{dag_id}' and DagRun ID: '{dag_run_id}' not found", ) events = dag_run.consumed_dataset_events - return dataset_event_collection_schema.dump( - DatasetEventCollection(dataset_events=events, total_entries=len(events)) + return asset_event_collection_schema.dump( + AssetEventCollection(dataset_events=events, total_entries=len(events)) ) diff --git a/airflow/api_connexion/endpoints/dataset_endpoint.py b/airflow/api_connexion/endpoints/dataset_endpoint.py index 1a1578266838c..95c3bead3da52 100644 --- a/airflow/api_connexion/endpoints/dataset_endpoint.py +++ b/airflow/api_connexion/endpoints/dataset_endpoint.py @@ -28,24 +28,24 @@ from airflow.api_connexion.endpoints.request_dict import get_json_request_dict from airflow.api_connexion.exceptions import BadRequest, NotFound from airflow.api_connexion.parameters import apply_sorting, check_limit, format_datetime, format_parameters -from airflow.api_connexion.schemas.dataset_schema import ( - DagScheduleDatasetReference, - DatasetCollection, - DatasetEventCollection, +from airflow.api_connexion.schemas.asset_schema import ( + AssetCollection, + AssetEventCollection, + DagScheduleAssetReference, QueuedEvent, QueuedEventCollection, - TaskOutletDatasetReference, - create_dataset_event_schema, - dataset_collection_schema, - dataset_event_collection_schema, - dataset_event_schema, - dataset_schema, + TaskOutletAssetReference, + asset_collection_schema, + asset_event_collection_schema, + asset_event_schema, + asset_schema, + create_asset_event_schema, queued_event_collection_schema, queued_event_schema, ) -from airflow.datasets import Dataset -from airflow.datasets.manager import dataset_manager -from airflow.models.dataset import DatasetDagRunQueue, DatasetEvent, DatasetModel +from airflow.assets import Asset +from airflow.assets.manager import asset_manager +from airflow.models.asset import AssetDagRunQueue, AssetEvent, AssetModel from airflow.utils import timezone from airflow.utils.db import get_query_count from airflow.utils.session import NEW_SESSION, provide_session @@ -60,24 +60,24 @@ RESOURCE_EVENT_PREFIX = "dataset" -@security.requires_access_dataset("GET") +@security.requires_access_asset("GET") @provide_session def get_dataset(*, uri: str, session: Session = NEW_SESSION) -> APIResponse: - """Get a Dataset.""" - dataset = session.scalar( - select(DatasetModel) - .where(DatasetModel.uri == uri) - .options(joinedload(DatasetModel.consuming_dags), joinedload(DatasetModel.producing_tasks)) + """Get an asset .""" + asset = session.scalar( + select(AssetModel) + .where(AssetModel.uri == uri) + .options(joinedload(AssetModel.consuming_dags), joinedload(AssetModel.producing_tasks)) ) - if not dataset: + if not asset: raise NotFound( - "Dataset not found", - detail=f"The Dataset with uri: `{uri}` was not found", + "Asset not found", + detail=f"The Asset with uri: `{uri}` was not found", ) - return dataset_schema.dump(dataset) + return asset_schema.dump(asset) -@security.requires_access_dataset("GET") +@security.requires_access_asset("GET") @format_parameters({"limit": check_limit}) @provide_session def get_datasets( @@ -89,30 +89,30 @@ def get_datasets( order_by: str = "id", session: Session = NEW_SESSION, ) -> APIResponse: - """Get datasets.""" + """Get assets.""" allowed_attrs = ["id", "uri", "created_at", "updated_at"] - total_entries = session.scalars(select(func.count(DatasetModel.id))).one() - query = select(DatasetModel) + total_entries = session.scalars(select(func.count(AssetModel.id))).one() + query = select(AssetModel) if dag_ids: dags_list = dag_ids.split(",") query = query.filter( - (DatasetModel.consuming_dags.any(DagScheduleDatasetReference.dag_id.in_(dags_list))) - | (DatasetModel.producing_tasks.any(TaskOutletDatasetReference.dag_id.in_(dags_list))) + (AssetModel.consuming_dags.any(DagScheduleAssetReference.dag_id.in_(dags_list))) + | (AssetModel.producing_tasks.any(TaskOutletAssetReference.dag_id.in_(dags_list))) ) if uri_pattern: - query = query.where(DatasetModel.uri.ilike(f"%{uri_pattern}%")) + query = query.where(AssetModel.uri.ilike(f"%{uri_pattern}%")) query = apply_sorting(query, order_by, {}, allowed_attrs) - datasets = session.scalars( - query.options(subqueryload(DatasetModel.consuming_dags), subqueryload(DatasetModel.producing_tasks)) + assets = session.scalars( + query.options(subqueryload(AssetModel.consuming_dags), subqueryload(AssetModel.producing_tasks)) .offset(offset) .limit(limit) ).all() - return dataset_collection_schema.dump(DatasetCollection(datasets=datasets, total_entries=total_entries)) + return asset_collection_schema.dump(AssetCollection(datasets=assets, total_entries=total_entries)) -@security.requires_access_dataset("GET") +@security.requires_access_asset("GET") @provide_session @format_parameters({"limit": check_limit}) def get_dataset_events( @@ -127,29 +127,29 @@ def get_dataset_events( source_map_index: int | None = None, session: Session = NEW_SESSION, ) -> APIResponse: - """Get dataset events.""" + """Get asset events.""" allowed_attrs = ["source_dag_id", "source_task_id", "source_run_id", "source_map_index", "timestamp"] - query = select(DatasetEvent) + query = select(AssetEvent) if dataset_id: - query = query.where(DatasetEvent.dataset_id == dataset_id) + query = query.where(AssetEvent.dataset_id == dataset_id) if source_dag_id: - query = query.where(DatasetEvent.source_dag_id == source_dag_id) + query = query.where(AssetEvent.source_dag_id == source_dag_id) if source_task_id: - query = query.where(DatasetEvent.source_task_id == source_task_id) + query = query.where(AssetEvent.source_task_id == source_task_id) if source_run_id: - query = query.where(DatasetEvent.source_run_id == source_run_id) + query = query.where(AssetEvent.source_run_id == source_run_id) if source_map_index: - query = query.where(DatasetEvent.source_map_index == source_map_index) + query = query.where(AssetEvent.source_map_index == source_map_index) - query = query.options(subqueryload(DatasetEvent.created_dagruns)) + query = query.options(subqueryload(AssetEvent.created_dagruns)) total_entries = get_query_count(query, session=session) query = apply_sorting(query, order_by, {}, allowed_attrs) events = session.scalars(query.offset(offset).limit(limit)).all() - return dataset_event_collection_schema.dump( - DatasetEventCollection(dataset_events=events, total_entries=total_entries) + return asset_event_collection_schema.dump( + AssetEventCollection(dataset_events=events, total_entries=total_entries) ) @@ -161,79 +161,77 @@ def _generate_queued_event_where_clause( before: str | None = None, permitted_dag_ids: set[str] | None = None, ) -> list: - """Get DatasetDagRunQueue where clause.""" + """Get AssetDagRunQueue where clause.""" where_clause = [] if dag_id is not None: - where_clause.append(DatasetDagRunQueue.target_dag_id == dag_id) + where_clause.append(AssetDagRunQueue.target_dag_id == dag_id) if dataset_id is not None: - where_clause.append(DatasetDagRunQueue.dataset_id == dataset_id) + where_clause.append(AssetDagRunQueue.dataset_id == dataset_id) if uri is not None: where_clause.append( - DatasetDagRunQueue.dataset_id.in_( - select(DatasetModel.id).where(DatasetModel.uri == uri), + AssetDagRunQueue.dataset_id.in_( + select(AssetModel.id).where(AssetModel.uri == uri), ), ) if before is not None: - where_clause.append(DatasetDagRunQueue.created_at < format_datetime(before)) + where_clause.append(AssetDagRunQueue.created_at < format_datetime(before)) if permitted_dag_ids is not None: - where_clause.append(DatasetDagRunQueue.target_dag_id.in_(permitted_dag_ids)) + where_clause.append(AssetDagRunQueue.target_dag_id.in_(permitted_dag_ids)) return where_clause -@security.requires_access_dataset("GET") +@security.requires_access_asset("GET") @security.requires_access_dag("GET") @provide_session def get_dag_dataset_queued_event( *, dag_id: str, uri: str, before: str | None = None, session: Session = NEW_SESSION ) -> APIResponse: - """Get a queued Dataset event for a DAG.""" + """Get a queued asset event for a DAG.""" where_clause = _generate_queued_event_where_clause(dag_id=dag_id, uri=uri, before=before) - ddrq = session.scalar( - select(DatasetDagRunQueue) - .join(DatasetModel, DatasetDagRunQueue.dataset_id == DatasetModel.id) + adrq = session.scalar( + select(AssetDagRunQueue) + .join(AssetModel, AssetDagRunQueue.dataset_id == AssetModel.id) .where(*where_clause) ) - if ddrq is None: + if adrq is None: raise NotFound( "Queue event not found", - detail=f"Queue event with dag_id: `{dag_id}` and dataset uri: `{uri}` was not found", + detail=f"Queue event with dag_id: `{dag_id}` and asset uri: `{uri}` was not found", ) - queued_event = {"created_at": ddrq.created_at, "dag_id": dag_id, "uri": uri} + queued_event = {"created_at": adrq.created_at, "dag_id": dag_id, "uri": uri} return queued_event_schema.dump(queued_event) -@security.requires_access_dataset("DELETE") +@security.requires_access_asset("DELETE") @security.requires_access_dag("GET") @provide_session @action_logging def delete_dag_dataset_queued_event( *, dag_id: str, uri: str, before: str | None = None, session: Session = NEW_SESSION ) -> APIResponse: - """Delete a queued Dataset event for a DAG.""" + """Delete a queued asset event for a DAG.""" where_clause = _generate_queued_event_where_clause(dag_id=dag_id, uri=uri, before=before) - delete_stmt = ( - delete(DatasetDagRunQueue).where(*where_clause).execution_options(synchronize_session="fetch") - ) + delete_stmt = delete(AssetDagRunQueue).where(*where_clause).execution_options(synchronize_session="fetch") result = session.execute(delete_stmt) if result.rowcount > 0: return NoContent, HTTPStatus.NO_CONTENT raise NotFound( "Queue event not found", - detail=f"Queue event with dag_id: `{dag_id}` and dataset uri: `{uri}` was not found", + detail=f"Queue event with dag_id: `{dag_id}` and asset uri: `{uri}` was not found", ) -@security.requires_access_dataset("GET") +@security.requires_access_asset("GET") @security.requires_access_dag("GET") @provide_session def get_dag_dataset_queued_events( *, dag_id: str, before: str | None = None, session: Session = NEW_SESSION ) -> APIResponse: - """Get queued Dataset events for a DAG.""" + """Get queued asset events for a DAG.""" where_clause = _generate_queued_event_where_clause(dag_id=dag_id, before=before) query = ( - select(DatasetDagRunQueue, DatasetModel.uri) - .join(DatasetModel, DatasetDagRunQueue.dataset_id == DatasetModel.id) + select(AssetDagRunQueue, AssetModel.uri) + .join(AssetModel, AssetDagRunQueue.dataset_id == AssetModel.id) .where(*where_clause) ) result = session.execute(query).all() @@ -244,23 +242,23 @@ def get_dag_dataset_queued_events( detail=f"Queue event with dag_id: `{dag_id}` was not found", ) queued_events = [ - QueuedEvent(created_at=ddrq.created_at, dag_id=ddrq.target_dag_id, uri=uri) for ddrq, uri in result + QueuedEvent(created_at=adrq.created_at, dag_id=adrq.target_dag_id, uri=uri) for adrq, uri in result ] return queued_event_collection_schema.dump( QueuedEventCollection(queued_events=queued_events, total_entries=total_entries) ) -@security.requires_access_dataset("DELETE") +@security.requires_access_asset("DELETE") @security.requires_access_dag("GET") @action_logging @provide_session def delete_dag_dataset_queued_events( *, dag_id: str, before: str | None = None, session: Session = NEW_SESSION ) -> APIResponse: - """Delete queued Dataset events for a DAG.""" + """Delete queued asset events for a DAG.""" where_clause = _generate_queued_event_where_clause(dag_id=dag_id, before=before) - delete_stmt = delete(DatasetDagRunQueue).where(*where_clause) + delete_stmt = delete(AssetDagRunQueue).where(*where_clause) result = session.execute(delete_stmt) if result.rowcount > 0: return NoContent, HTTPStatus.NO_CONTENT @@ -271,87 +269,85 @@ def delete_dag_dataset_queued_events( ) -@security.requires_access_dataset("GET") +@security.requires_access_asset("GET") @provide_session def get_dataset_queued_events( *, uri: str, before: str | None = None, session: Session = NEW_SESSION ) -> APIResponse: - """Get queued Dataset events for a Dataset.""" + """Get queued asset events for an asset.""" permitted_dag_ids = get_auth_manager().get_permitted_dag_ids(methods=["GET"]) where_clause = _generate_queued_event_where_clause( uri=uri, before=before, permitted_dag_ids=permitted_dag_ids ) query = ( - select(DatasetDagRunQueue, DatasetModel.uri) - .join(DatasetModel, DatasetDagRunQueue.dataset_id == DatasetModel.id) + select(AssetDagRunQueue, AssetModel.uri) + .join(AssetModel, AssetDagRunQueue.dataset_id == AssetModel.id) .where(*where_clause) ) total_entries = get_query_count(query, session=session) result = session.execute(query).all() if total_entries > 0: queued_events = [ - QueuedEvent(created_at=ddrq.created_at, dag_id=ddrq.target_dag_id, uri=uri) - for ddrq, uri in result + QueuedEvent(created_at=adrq.created_at, dag_id=adrq.target_dag_id, uri=uri) + for adrq, uri in result ] return queued_event_collection_schema.dump( QueuedEventCollection(queued_events=queued_events, total_entries=total_entries) ) raise NotFound( "Queue event not found", - detail=f"Queue event with dataset uri: `{uri}` was not found", + detail=f"Queue event with asset uri: `{uri}` was not found", ) -@security.requires_access_dataset("DELETE") +@security.requires_access_asset("DELETE") @action_logging @provide_session def delete_dataset_queued_events( *, uri: str, before: str | None = None, session: Session = NEW_SESSION ) -> APIResponse: - """Delete queued Dataset events for a Dataset.""" + """Delete queued asset events for an asset.""" permitted_dag_ids = get_auth_manager().get_permitted_dag_ids(methods=["GET"]) where_clause = _generate_queued_event_where_clause( uri=uri, before=before, permitted_dag_ids=permitted_dag_ids ) - delete_stmt = ( - delete(DatasetDagRunQueue).where(*where_clause).execution_options(synchronize_session="fetch") - ) + delete_stmt = delete(AssetDagRunQueue).where(*where_clause).execution_options(synchronize_session="fetch") result = session.execute(delete_stmt) if result.rowcount > 0: return NoContent, HTTPStatus.NO_CONTENT raise NotFound( "Queue event not found", - detail=f"Queue event with dataset uri: `{uri}` was not found", + detail=f"Queue event with asset uri: `{uri}` was not found", ) -@security.requires_access_dataset("POST") +@security.requires_access_asset("POST") @provide_session @action_logging def create_dataset_event(session: Session = NEW_SESSION) -> APIResponse: - """Create dataset event.""" + """Create asset event.""" body = get_json_request_dict() try: - json_body = create_dataset_event_schema.load(body) + json_body = create_asset_event_schema.load(body) except ValidationError as err: raise BadRequest(detail=str(err)) uri = json_body["dataset_uri"] - dataset = session.scalar(select(DatasetModel).where(DatasetModel.uri == uri).limit(1)) - if not dataset: - raise NotFound(title="Dataset not found", detail=f"Dataset with uri: '{uri}' not found") + asset = session.scalar(select(AssetModel).where(AssetModel.uri == uri).limit(1)) + if not asset: + raise NotFound(title="Asset not found", detail=f"Asset with uri: '{uri}' not found") timestamp = timezone.utcnow() extra = json_body.get("extra", {}) extra["from_rest_api"] = True - dataset_event = dataset_manager.register_dataset_change( - dataset=Dataset(uri), + asset_event = asset_manager.register_asset_change( + asset=Asset(uri), timestamp=timestamp, extra=extra, session=session, ) - if not dataset_event: - raise NotFound(title="Dataset not found", detail=f"Dataset with uri: '{uri}' not found") + if not asset_event: + raise NotFound(title="Asset not found", detail=f"Asset with uri: '{uri}' not found") session.flush() # So we can dump the timestamp. - event = dataset_event_schema.dump(dataset_event) + event = asset_event_schema.dump(asset_event) return event diff --git a/airflow/api_connexion/schemas/dataset_schema.py b/airflow/api_connexion/schemas/asset_schema.py similarity index 65% rename from airflow/api_connexion/schemas/dataset_schema.py rename to airflow/api_connexion/schemas/asset_schema.py index b8aaf2f8fa30e..791941f42016d 100644 --- a/airflow/api_connexion/schemas/dataset_schema.py +++ b/airflow/api_connexion/schemas/asset_schema.py @@ -23,23 +23,23 @@ from marshmallow_sqlalchemy import SQLAlchemySchema, auto_field from airflow.api_connexion.schemas.common_schema import JsonObjectField -from airflow.models.dagrun import DagRun -from airflow.models.dataset import ( - DagScheduleDatasetReference, - DatasetAliasModel, - DatasetEvent, - DatasetModel, - TaskOutletDatasetReference, +from airflow.models.asset import ( + AssetAliasModel, + AssetEvent, + AssetModel, + DagScheduleAssetReference, + TaskOutletAssetReference, ) +from airflow.models.dagrun import DagRun -class TaskOutletDatasetReferenceSchema(SQLAlchemySchema): - """TaskOutletDatasetReference DB schema.""" +class TaskOutletAssetReferenceSchema(SQLAlchemySchema): + """TaskOutletAssetReference DB schema.""" class Meta: """Meta.""" - model = TaskOutletDatasetReference + model = TaskOutletAssetReference dag_id = auto_field() task_id = auto_field() @@ -47,65 +47,65 @@ class Meta: updated_at = auto_field() -class DagScheduleDatasetReferenceSchema(SQLAlchemySchema): - """DagScheduleDatasetReference DB schema.""" +class DagScheduleAssetReferenceSchema(SQLAlchemySchema): + """DagScheduleAssetReference DB schema.""" class Meta: """Meta.""" - model = DagScheduleDatasetReference + model = DagScheduleAssetReference dag_id = auto_field() created_at = auto_field() updated_at = auto_field() -class DatasetAliasSchema(SQLAlchemySchema): - """DatasetAlias DB schema.""" +class AssetAliasSchema(SQLAlchemySchema): + """AssetAlias DB schema.""" class Meta: """Meta.""" - model = DatasetAliasModel + model = AssetAliasModel id = auto_field() name = auto_field() -class DatasetSchema(SQLAlchemySchema): - """Dataset DB schema.""" +class AssetSchema(SQLAlchemySchema): + """Asset DB schema.""" class Meta: """Meta.""" - model = DatasetModel + model = AssetModel id = auto_field() uri = auto_field() extra = JsonObjectField() created_at = auto_field() updated_at = auto_field() - producing_tasks = fields.List(fields.Nested(TaskOutletDatasetReferenceSchema)) - consuming_dags = fields.List(fields.Nested(DagScheduleDatasetReferenceSchema)) - aliases = fields.List(fields.Nested(DatasetAliasSchema)) + producing_tasks = fields.List(fields.Nested(TaskOutletAssetReferenceSchema)) + consuming_dags = fields.List(fields.Nested(DagScheduleAssetReferenceSchema)) + aliases = fields.List(fields.Nested(AssetAliasSchema)) -class DatasetCollection(NamedTuple): - """List of Datasets with meta.""" +class AssetCollection(NamedTuple): + """List of Assets with meta.""" - datasets: list[DatasetModel] + datasets: list[AssetModel] total_entries: int -class DatasetCollectionSchema(Schema): - """Dataset Collection Schema.""" +class AssetCollectionSchema(Schema): + """Asset Collection Schema.""" - datasets = fields.List(fields.Nested(DatasetSchema)) + datasets = fields.List(fields.Nested(AssetSchema)) total_entries = fields.Int() -dataset_schema = DatasetSchema() -dataset_collection_schema = DatasetCollectionSchema() +asset_schema = AssetSchema() +asset_collection_schema = AssetCollectionSchema() class BasicDAGRunSchema(SQLAlchemySchema): @@ -127,13 +127,13 @@ class Meta: data_interval_end = auto_field(dump_only=True) -class DatasetEventSchema(SQLAlchemySchema): - """Dataset Event DB schema.""" +class AssetEventSchema(SQLAlchemySchema): + """Asset Event DB schema.""" class Meta: """Meta.""" - model = DatasetEvent + model = AssetEvent id = auto_field() dataset_id = auto_field() @@ -147,30 +147,30 @@ class Meta: timestamp = auto_field() -class DatasetEventCollection(NamedTuple): - """List of Dataset events with meta.""" +class AssetEventCollection(NamedTuple): + """List of Asset events with meta.""" - dataset_events: list[DatasetEvent] + dataset_events: list[AssetEvent] total_entries: int -class DatasetEventCollectionSchema(Schema): - """Dataset Event Collection Schema.""" +class AssetEventCollectionSchema(Schema): + """Asset Event Collection Schema.""" - dataset_events = fields.List(fields.Nested(DatasetEventSchema)) + dataset_events = fields.List(fields.Nested(AssetEventSchema)) total_entries = fields.Int() -class CreateDatasetEventSchema(Schema): - """Create Dataset Event Schema.""" +class CreateAssetEventSchema(Schema): + """Create Asset Event Schema.""" dataset_uri = fields.String() extra = JsonObjectField() -dataset_event_schema = DatasetEventSchema() -dataset_event_collection_schema = DatasetEventCollectionSchema() -create_dataset_event_schema = CreateDatasetEventSchema() +asset_event_schema = AssetEventSchema() +asset_event_collection_schema = AssetEventCollectionSchema() +create_asset_event_schema = CreateAssetEventSchema() class QueuedEvent(NamedTuple): diff --git a/airflow/api_connexion/security.py b/airflow/api_connexion/security.py index 445ded913e56a..1098de3a1f474 100644 --- a/airflow/api_connexion/security.py +++ b/airflow/api_connexion/security.py @@ -24,11 +24,11 @@ from airflow.api_connexion.exceptions import PermissionDenied, Unauthenticated from airflow.auth.managers.models.resource_details import ( AccessView, + AssetDetails, ConfigurationDetails, ConnectionDetails, DagAccessEntity, DagDetails, - DatasetDetails, PoolDetails, VariableDetails, ) @@ -158,14 +158,14 @@ def decorated(*args, **kwargs): return requires_access_decorator -def requires_access_dataset(method: ResourceMethod) -> Callable[[T], T]: +def requires_access_asset(method: ResourceMethod) -> Callable[[T], T]: def requires_access_decorator(func: T): @wraps(func) def decorated(*args, **kwargs): uri: str | None = kwargs.get("uri") return _requires_access( - is_authorized_callback=lambda: get_auth_manager().is_authorized_dataset( - method=method, details=DatasetDetails(uri=uri) + is_authorized_callback=lambda: get_auth_manager().is_authorized_asset( + method=method, details=AssetDetails(uri=uri) ), func=func, args=args, diff --git a/airflow/api_fastapi/openapi/v1-generated.yaml b/airflow/api_fastapi/openapi/v1-generated.yaml index 64e475aeb6baa..c130f3162c6e6 100644 --- a/airflow/api_fastapi/openapi/v1-generated.yaml +++ b/airflow/api_fastapi/openapi/v1-generated.yaml @@ -10,9 +10,9 @@ paths: /ui/next_run_datasets/{dag_id}: get: tags: - - Dataset - summary: Next Run Datasets - operationId: next_run_datasets_ui_next_run_datasets__dag_id__get + - Asset + summary: Next Run Assets + operationId: next_run_assets_ui_next_run_datasets__dag_id__get parameters: - name: dag_id in: path @@ -27,7 +27,7 @@ paths: application/json: schema: type: object - title: Response Next Run Datasets Ui Next Run Datasets Dag Id Get + title: Response Next Run Assets Ui Next Run Datasets Dag Id Get '422': description: Validation Error content: diff --git a/airflow/api_fastapi/views/ui/__init__.py b/airflow/api_fastapi/views/ui/__init__.py index 2d95e040403a7..edba930c3d1d1 100644 --- a/airflow/api_fastapi/views/ui/__init__.py +++ b/airflow/api_fastapi/views/ui/__init__.py @@ -18,8 +18,8 @@ from fastapi import APIRouter -from airflow.api_fastapi.views.ui.datasets import datasets_router +from airflow.api_fastapi.views.ui.assets import assets_router ui_router = APIRouter(prefix="/ui") -ui_router.include_router(datasets_router) +ui_router.include_router(assets_router) diff --git a/airflow/api_fastapi/views/ui/datasets.py b/airflow/api_fastapi/views/ui/assets.py similarity index 66% rename from airflow/api_fastapi/views/ui/datasets.py rename to airflow/api_fastapi/views/ui/assets.py index f5dd2cacb126d..458d531facf6a 100644 --- a/airflow/api_fastapi/views/ui/datasets.py +++ b/airflow/api_fastapi/views/ui/assets.py @@ -24,13 +24,13 @@ from airflow.api_fastapi.db import get_session from airflow.models import DagModel -from airflow.models.dataset import DagScheduleDatasetReference, DatasetDagRunQueue, DatasetEvent, DatasetModel +from airflow.models.asset import AssetDagRunQueue, AssetEvent, AssetModel, DagScheduleAssetReference -datasets_router = APIRouter(tags=["Dataset"]) +assets_router = APIRouter(tags=["Asset"]) -@datasets_router.get("/next_run_datasets/{dag_id}", include_in_schema=False) -async def next_run_datasets( +@assets_router.get("/next_run_datasets/{dag_id}", include_in_schema=False) +async def next_run_assets( dag_id: str, request: Request, session: Annotated[Session, Depends(get_session)], @@ -51,34 +51,34 @@ async def next_run_datasets( dict(info._mapping) for info in session.execute( select( - DatasetModel.id, - DatasetModel.uri, - func.max(DatasetEvent.timestamp).label("lastUpdate"), + AssetModel.id, + AssetModel.uri, + func.max(AssetEvent.timestamp).label("lastUpdate"), ) - .join(DagScheduleDatasetReference, DagScheduleDatasetReference.dataset_id == DatasetModel.id) + .join(DagScheduleAssetReference, DagScheduleAssetReference.dataset_id == AssetModel.id) .join( - DatasetDagRunQueue, + AssetDagRunQueue, and_( - DatasetDagRunQueue.dataset_id == DatasetModel.id, - DatasetDagRunQueue.target_dag_id == DagScheduleDatasetReference.dag_id, + AssetDagRunQueue.dataset_id == AssetModel.id, + AssetDagRunQueue.target_dag_id == DagScheduleAssetReference.dag_id, ), isouter=True, ) .join( - DatasetEvent, + AssetEvent, and_( - DatasetEvent.dataset_id == DatasetModel.id, + AssetEvent.dataset_id == AssetModel.id, ( - DatasetEvent.timestamp >= latest_run.execution_date + AssetEvent.timestamp >= latest_run.execution_date if latest_run and latest_run.execution_date else True ), ), isouter=True, ) - .where(DagScheduleDatasetReference.dag_id == dag_id, ~DatasetModel.is_orphaned) - .group_by(DatasetModel.id, DatasetModel.uri) - .order_by(DatasetModel.uri) + .where(DagScheduleAssetReference.dag_id == dag_id, ~AssetModel.is_orphaned) + .group_by(AssetModel.id, AssetModel.uri) + .order_by(AssetModel.uri) ) ] data = {"dataset_expression": dag_model.dataset_expression, "events": events} diff --git a/airflow/api_internal/endpoints/rpc_api_endpoint.py b/airflow/api_internal/endpoints/rpc_api_endpoint.py index 8716d9c9cc49d..1cdd2536e1354 100644 --- a/airflow/api_internal/endpoints/rpc_api_endpoint.py +++ b/airflow/api_internal/endpoints/rpc_api_endpoint.py @@ -54,11 +54,11 @@ @functools.lru_cache def initialize_method_map() -> dict[str, Callable]: from airflow.api.common.trigger_dag import trigger_dag + from airflow.assets import expand_alias_to_assets + from airflow.assets.manager import AssetManager from airflow.cli.commands.task_command import _get_ti_db_access from airflow.dag_processing.manager import DagFileProcessorManager from airflow.dag_processing.processor import DagFileProcessor - from airflow.datasets import expand_alias_to_datasets - from airflow.datasets.manager import DatasetManager from airflow.models import Trigger, Variable, XCom from airflow.models.dag import DAG, DagModel from airflow.models.dagrun import DagRun @@ -109,8 +109,8 @@ def initialize_method_map() -> dict[str, Callable]: DagFileProcessorManager.clear_nonexistent_import_errors, DagFileProcessorManager.deactivate_stale_dags, DagWarning.purge_inactive_dag_warnings, - expand_alias_to_datasets, - DatasetManager.register_dataset_change, + expand_alias_to_assets, + AssetManager.register_asset_change, FileTaskHandler._render_filename_db_access, Job._add_to_db, Job._fetch_from_db, diff --git a/airflow/datasets/__init__.py b/airflow/assets/__init__.py similarity index 60% rename from airflow/datasets/__init__.py rename to airflow/assets/__init__.py index 6f7ae99ff7417..9727e408edc2e 100644 --- a/airflow/datasets/__init__.py +++ b/airflow/assets/__init__.py @@ -38,7 +38,7 @@ from airflow.configuration import conf -__all__ = ["Dataset", "DatasetAll", "DatasetAny"] +__all__ = ["Asset", "AssetAll", "AssetAny"] def normalize_noop(parts: SplitResult) -> SplitResult: @@ -55,7 +55,7 @@ def _get_uri_normalizer(scheme: str) -> Callable[[SplitResult], SplitResult] | N return normalize_noop from airflow.providers_manager import ProvidersManager - return ProvidersManager().dataset_uri_handlers.get(scheme) + return ProvidersManager().asset_uri_handlers.get(scheme) def _get_normalized_scheme(uri: str) -> str: @@ -65,17 +65,17 @@ def _get_normalized_scheme(uri: str) -> str: def _sanitize_uri(uri: str) -> str: """ - Sanitize a dataset URI. + Sanitize an asset URI. This checks for URI validity, and normalizes the URI if needed. A fully normalized URI is returned. """ if not uri: - raise ValueError("Dataset URI cannot be empty") + raise ValueError("Asset URI cannot be empty") if uri.isspace(): - raise ValueError("Dataset URI cannot be just whitespace") + raise ValueError("Asset URI cannot be just whitespace") if not uri.isascii(): - raise ValueError("Dataset URI must only consist of ASCII characters") + raise ValueError("Asset URI must only consist of ASCII characters") parsed = urllib.parse.urlsplit(uri) if not parsed.scheme and not parsed.netloc: # Does not look like a URI. return uri @@ -84,12 +84,12 @@ def _sanitize_uri(uri: str) -> str: if normalized_scheme.startswith("x-"): return uri if normalized_scheme == "airflow": - raise ValueError("Dataset scheme 'airflow' is reserved") + raise ValueError("Asset scheme 'airflow' is reserved") _, auth_exists, normalized_netloc = parsed.netloc.rpartition("@") if auth_exists: # TODO: Collect this into a DagWarning. warnings.warn( - "A dataset URI should not contain auth info (e.g. username or " + "An Asset URI should not contain auth info (e.g. username or " "password). It has been automatically dropped.", UserWarning, stacklevel=3, @@ -109,10 +109,10 @@ def _sanitize_uri(uri: str) -> str: try: parsed = normalizer(parsed) except ValueError as exception: - if conf.getboolean("core", "strict_dataset_uri_validation", fallback=False): + if conf.getboolean("core", "strict_asset_uri_validation", fallback=False): raise warnings.warn( - f"The dataset URI {uri} is not AIP-60 compliant: {exception}. " + f"The Asset URI {uri} is not AIP-60 compliant: {exception}. " f"In Airflow 3, this will raise an exception.", UserWarning, stacklevel=3, @@ -120,46 +120,44 @@ def _sanitize_uri(uri: str) -> str: return urllib.parse.urlunsplit(parsed) -def extract_event_key(value: str | Dataset | DatasetAlias) -> str: +def extract_event_key(value: str | Asset | AssetAlias) -> str: """ Extract the key of an inlet or an outlet event. If the input value is a string, it is treated as a URI and sanitized. If the - input is a :class:`Dataset`, the URI it contains is considered sanitized and - returned directly. If the input is a :class:`DatasetAlias`, the name it contains + input is a :class:`Asset`, the URI it contains is considered sanitized and + returned directly. If the input is a :class:`AssetAlias`, the name it contains will be returned directly. :meta private: """ - if isinstance(value, DatasetAlias): + if isinstance(value, AssetAlias): return value.name - if isinstance(value, Dataset): + if isinstance(value, Asset): return value.uri return _sanitize_uri(str(value)) @internal_api_call @provide_session -def expand_alias_to_datasets( - alias: str | DatasetAlias, *, session: Session = NEW_SESSION -) -> list[BaseDataset]: - """Expand dataset alias to resolved datasets.""" - from airflow.models.dataset import DatasetAliasModel +def expand_alias_to_assets(alias: str | AssetAlias, *, session: Session = NEW_SESSION) -> list[BaseAsset]: + """Expand asset alias to resolved assets.""" + from airflow.models.asset import AssetAliasModel - alias_name = alias.name if isinstance(alias, DatasetAlias) else alias + alias_name = alias.name if isinstance(alias, AssetAlias) else alias - dataset_alias_obj = session.scalar( - select(DatasetAliasModel).where(DatasetAliasModel.name == alias_name).limit(1) + asset_alias_obj = session.scalar( + select(AssetAliasModel).where(AssetAliasModel.name == alias_name).limit(1) ) - if dataset_alias_obj: - return [Dataset(uri=dataset.uri, extra=dataset.extra) for dataset in dataset_alias_obj.datasets] + if asset_alias_obj: + return [Asset(uri=asset.uri, extra=asset.extra) for asset in asset_alias_obj.datasets] return [] -class BaseDataset: +class BaseAsset: """ - Protocol for all dataset triggers to use in ``DAG(schedule=...)``. + Protocol for all asset triggers to use in ``DAG(schedule=...)``. :meta private: """ @@ -167,19 +165,19 @@ class BaseDataset: def __bool__(self) -> bool: return True - def __or__(self, other: BaseDataset) -> BaseDataset: - if not isinstance(other, BaseDataset): + def __or__(self, other: BaseAsset) -> BaseAsset: + if not isinstance(other, BaseAsset): return NotImplemented - return DatasetAny(self, other) + return AssetAny(self, other) - def __and__(self, other: BaseDataset) -> BaseDataset: - if not isinstance(other, BaseDataset): + def __and__(self, other: BaseAsset) -> BaseAsset: + if not isinstance(other, BaseAsset): return NotImplemented - return DatasetAll(self, other) + return AssetAll(self, other) def as_expression(self) -> Any: """ - Serialize the dataset into its scheduling expression. + Serialize the asset into its scheduling expression. The return value is stored in DagModel for display purposes. It must be JSON-compatible. @@ -191,15 +189,15 @@ def as_expression(self) -> Any: def evaluate(self, statuses: dict[str, bool]) -> bool: raise NotImplementedError - def iter_datasets(self) -> Iterator[tuple[str, Dataset]]: + def iter_assets(self) -> Iterator[tuple[str, Asset]]: raise NotImplementedError - def iter_dataset_aliases(self) -> Iterator[tuple[str, DatasetAlias]]: + def iter_asset_aliases(self) -> Iterator[tuple[str, AssetAlias]]: raise NotImplementedError def iter_dag_dependencies(self, *, source: str, target: str) -> Iterator[DagDependency]: """ - Iterate a base dataset as dag dependency. + Iterate a base asset as dag dependency. :meta private: """ @@ -207,36 +205,36 @@ def iter_dag_dependencies(self, *, source: str, target: str) -> Iterator[DagDepe @attr.define(unsafe_hash=False) -class DatasetAlias(BaseDataset): - """A represeation of dataset alias which is used to create dataset during the runtime.""" +class AssetAlias(BaseAsset): + """A represeation of asset alias which is used to create asset during the runtime.""" name: str - def iter_datasets(self) -> Iterator[tuple[str, Dataset]]: + def iter_assets(self) -> Iterator[tuple[str, Asset]]: return iter(()) - def iter_dataset_aliases(self) -> Iterator[tuple[str, DatasetAlias]]: + def iter_asset_aliases(self) -> Iterator[tuple[str, AssetAlias]]: yield self.name, self def iter_dag_dependencies(self, *, source: str, target: str) -> Iterator[DagDependency]: """ - Iterate a dataset alias as dag dependency. + Iterate an asset alias as dag dependency. :meta private: """ yield DagDependency( - source=source or "dataset-alias", - target=target or "dataset-alias", - dependency_type="dataset-alias", + source=source or "asset-alias", + target=target or "asset-alias", + dependency_type="asset-alias", dependency_id=self.name, ) -class DatasetAliasEvent(TypedDict): - """A represeation of dataset event to be triggered by a dataset alias.""" +class AssetAliasEvent(TypedDict): + """A represeation of asset event to be triggered by an asset alias.""" source_alias_name: str - dest_dataset_uri: str + dest_asset_uri: str extra: dict[str, Any] @@ -244,7 +242,7 @@ def _set_extra_default(extra: dict | None) -> dict: """ Automatically convert None to an empty dict. - This allows the caller site to continue doing ``Dataset(uri, extra=None)``, + This allows the caller site to continue doing ``Asset(uri, extra=None)``, but still allow the ``extra`` attribute to always be a dict. """ if extra is None: @@ -253,7 +251,7 @@ def _set_extra_default(extra: dict | None) -> dict: @attr.define(unsafe_hash=False) -class Dataset(os.PathLike, BaseDataset): +class Asset(os.PathLike, BaseAsset): """A representation of data dependencies between workflows.""" uri: str = attr.field( @@ -291,16 +289,16 @@ def normalized_uri(self) -> str | None: def as_expression(self) -> Any: """ - Serialize the dataset into its scheduling expression. + Serialize the asset into its scheduling expression. :meta private: """ return self.uri - def iter_datasets(self) -> Iterator[tuple[str, Dataset]]: + def iter_assets(self) -> Iterator[tuple[str, Asset]]: yield self.uri, self - def iter_dataset_aliases(self) -> Iterator[tuple[str, DatasetAlias]]: + def iter_asset_aliases(self) -> Iterator[tuple[str, AssetAlias]]: return iter(()) def evaluate(self, statuses: dict[str, bool]) -> bool: @@ -308,51 +306,51 @@ def evaluate(self, statuses: dict[str, bool]) -> bool: def iter_dag_dependencies(self, *, source: str, target: str) -> Iterator[DagDependency]: """ - Iterate a dataset as dag dependency. + Iterate an asset as dag dependency. :meta private: """ yield DagDependency( - source=source or "dataset", - target=target or "dataset", - dependency_type="dataset", + source=source or "asset", + target=target or "asset", + dependency_type="asset", dependency_id=self.uri, ) -class _DatasetBooleanCondition(BaseDataset): - """Base class for dataset boolean logic.""" +class _AssetBooleanCondition(BaseAsset): + """Base class for asset boolean logic.""" agg_func: Callable[[Iterable], bool] - def __init__(self, *objects: BaseDataset) -> None: - if not all(isinstance(o, BaseDataset) for o in objects): - raise TypeError("expect dataset expressions in condition") + def __init__(self, *objects: BaseAsset) -> None: + if not all(isinstance(o, BaseAsset) for o in objects): + raise TypeError("expect asset expressions in condition") self.objects = [ - _DatasetAliasCondition(obj.name) if isinstance(obj, DatasetAlias) else obj for obj in objects + _AssetAliasCondition(obj.name) if isinstance(obj, AssetAlias) else obj for obj in objects ] def evaluate(self, statuses: dict[str, bool]) -> bool: return self.agg_func(x.evaluate(statuses=statuses) for x in self.objects) - def iter_datasets(self) -> Iterator[tuple[str, Dataset]]: + def iter_assets(self) -> Iterator[tuple[str, Asset]]: seen = set() # We want to keep the first instance. for o in self.objects: - for k, v in o.iter_datasets(): + for k, v in o.iter_assets(): if k in seen: continue yield k, v seen.add(k) - def iter_dataset_aliases(self) -> Iterator[tuple[str, DatasetAlias]]: - """Filter dataest aliases in the condition.""" + def iter_asset_aliases(self) -> Iterator[tuple[str, AssetAlias]]: + """Filter asset aliases in the condition.""" for o in self.objects: - yield from o.iter_dataset_aliases() + yield from o.iter_asset_aliases() def iter_dag_dependencies(self, *, source: str, target: str) -> Iterator[DagDependency]: """ - Iterate dataset, dataset aliases and their resolved datasets as dag dependency. + Iterate asset, asset aliases and their resolved assets as dag dependency. :meta private: """ @@ -360,104 +358,104 @@ def iter_dag_dependencies(self, *, source: str, target: str) -> Iterator[DagDepe yield from obj.iter_dag_dependencies(source=source, target=target) -class DatasetAny(_DatasetBooleanCondition): - """Use to combine datasets schedule references in an "and" relationship.""" +class AssetAny(_AssetBooleanCondition): + """Use to combine assets schedule references in an "and" relationship.""" agg_func = any - def __or__(self, other: BaseDataset) -> BaseDataset: - if not isinstance(other, BaseDataset): + def __or__(self, other: BaseAsset) -> BaseAsset: + if not isinstance(other, BaseAsset): return NotImplemented # Optimization: X | (Y | Z) is equivalent to X | Y | Z. - return DatasetAny(*self.objects, other) + return AssetAny(*self.objects, other) def __repr__(self) -> str: - return f"DatasetAny({', '.join(map(str, self.objects))})" + return f"AssetAny({', '.join(map(str, self.objects))})" def as_expression(self) -> dict[str, Any]: """ - Serialize the dataset into its scheduling expression. + Serialize the asset into its scheduling expression. :meta private: """ return {"any": [o.as_expression() for o in self.objects]} -class _DatasetAliasCondition(DatasetAny): +class _AssetAliasCondition(AssetAny): """ - Use to expand DataAlias as DatasetAny of its resolved Datasets. + Use to expand AssetAlias as AssetAny of its resolved Assets. :meta private: """ def __init__(self, name: str) -> None: self.name = name - self.objects = expand_alias_to_datasets(name) + self.objects = expand_alias_to_assets(name) def __repr__(self) -> str: - return f"_DatasetAliasCondition({', '.join(map(str, self.objects))})" + return f"_AssetAliasCondition({', '.join(map(str, self.objects))})" def as_expression(self) -> Any: """ - Serialize the dataset into its scheduling expression. + Serialize the asset alias into its scheduling expression. :meta private: """ return {"alias": self.name} - def iter_dataset_aliases(self) -> Iterator[tuple[str, DatasetAlias]]: - yield self.name, DatasetAlias(self.name) + def iter_asset_aliases(self) -> Iterator[tuple[str, AssetAlias]]: + yield self.name, AssetAlias(self.name) def iter_dag_dependencies(self, *, source: str = "", target: str = "") -> Iterator[DagDependency]: """ - Iterate a dataset alias and its resolved datasets as dag dependency. + Iterate an asset alias and its resolved assets as dag dependency. :meta private: """ if self.objects: for obj in self.objects: - dataset = cast(Dataset, obj) - uri = dataset.uri - # dataset + asset = cast(Asset, obj) + uri = asset.uri + # asset yield DagDependency( - source=f"dataset-alias:{self.name}" if source else "dataset", - target="dataset" if source else f"dataset-alias:{self.name}", - dependency_type="dataset", + source=f"asset-alias:{self.name}" if source else "asset", + target="asset" if source else f"asset-alias:{self.name}", + dependency_type="asset", dependency_id=uri, ) - # dataset alias + # asset alias yield DagDependency( - source=source or f"dataset:{uri}", - target=target or f"dataset:{uri}", - dependency_type="dataset-alias", + source=source or f"asset:{uri}", + target=target or f"asset:{uri}", + dependency_type="asset-alias", dependency_id=self.name, ) else: yield DagDependency( - source=source or "dataset-alias", - target=target or "dataset-alias", - dependency_type="dataset-alias", + source=source or "asset-alias", + target=target or "asset-alias", + dependency_type="asset-alias", dependency_id=self.name, ) -class DatasetAll(_DatasetBooleanCondition): - """Use to combine datasets schedule references in an "or" relationship.""" +class AssetAll(_AssetBooleanCondition): + """Use to combine assets schedule references in an "or" relationship.""" agg_func = all - def __and__(self, other: BaseDataset) -> BaseDataset: - if not isinstance(other, BaseDataset): + def __and__(self, other: BaseAsset) -> BaseAsset: + if not isinstance(other, BaseAsset): return NotImplemented # Optimization: X & (Y & Z) is equivalent to X & Y & Z. - return DatasetAll(*self.objects, other) + return AssetAll(*self.objects, other) def __repr__(self) -> str: - return f"DatasetAll({', '.join(map(str, self.objects))})" + return f"AssetAll({', '.join(map(str, self.objects))})" def as_expression(self) -> Any: """ - Serialize the dataset into its scheduling expression. + Serialize the assets into its scheduling expression. :meta private: """ diff --git a/airflow/datasets/manager.py b/airflow/assets/manager.py similarity index 53% rename from airflow/datasets/manager.py rename to airflow/assets/manager.py index 6322414bb8499..d68a0efc87d12 100644 --- a/airflow/datasets/manager.py +++ b/airflow/assets/manager.py @@ -24,120 +24,121 @@ from sqlalchemy.orm import joinedload from airflow.api_internal.internal_api_call import internal_api_call +from airflow.assets import Asset from airflow.configuration import conf from airflow.listeners.listener import get_listener_manager -from airflow.models.dagbag import DagPriorityParsingRequest -from airflow.models.dataset import ( - DagScheduleDatasetAliasReference, - DagScheduleDatasetReference, - DatasetAliasModel, - DatasetDagRunQueue, - DatasetEvent, - DatasetModel, +from airflow.models.asset import ( + AssetAliasModel, + AssetDagRunQueue, + AssetEvent, + AssetModel, + DagScheduleAssetAliasReference, + DagScheduleAssetReference, ) +from airflow.models.dagbag import DagPriorityParsingRequest from airflow.stats import Stats from airflow.utils.log.logging_mixin import LoggingMixin if TYPE_CHECKING: from sqlalchemy.orm.session import Session - from airflow.datasets import Dataset, DatasetAlias + from airflow.assets import Asset, AssetAlias from airflow.models.dag import DagModel from airflow.models.taskinstance import TaskInstance -class DatasetManager(LoggingMixin): +class AssetManager(LoggingMixin): """ - A pluggable class that manages operations for datasets. + A pluggable class that manages operations for assets. - The intent is to have one place to handle all Dataset-related operations, so different - Airflow deployments can use plugins that broadcast dataset events to each other. + The intent is to have one place to handle all Asset-related operations, so different + Airflow deployments can use plugins that broadcast Asset events to each other. """ @classmethod - def create_datasets(cls, datasets: list[Dataset], *, session: Session) -> list[DatasetModel]: - """Create new datasets.""" + def create_assets(cls, assets: list[Asset], *, session: Session) -> list[AssetModel]: + """Create new assets.""" - def _add_one(dataset: Dataset) -> DatasetModel: - model = DatasetModel.from_public(dataset) + def _add_one(asset: Asset) -> AssetModel: + model = AssetModel.from_public(asset) session.add(model) - cls.notify_dataset_created(dataset=dataset) + cls.notify_asset_created(asset=asset) return model - return [_add_one(d) for d in datasets] + return [_add_one(a) for a in assets] @classmethod - def create_dataset_aliases( + def create_asset_aliases( cls, - dataset_aliases: list[DatasetAlias], + asset_aliases: list[AssetAlias], *, session: Session, - ) -> list[DatasetAliasModel]: - """Create new dataset aliases.""" + ) -> list[AssetAliasModel]: + """Create new asset aliases.""" - def _add_one(dataset_alias: DatasetAlias) -> DatasetAliasModel: - model = DatasetAliasModel.from_public(dataset_alias) + def _add_one(asset_alias: AssetAlias) -> AssetAliasModel: + model = AssetAliasModel.from_public(asset_alias) session.add(model) - cls.notify_dataset_alias_created(dataset_alias=dataset_alias) + cls.notify_asset_alias_created(asset_assets=asset_alias) return model - return [_add_one(a) for a in dataset_aliases] + return [_add_one(a) for a in asset_aliases] @classmethod - def _add_dataset_alias_association( + def _add_asset_alias_association( cls, alias_names: Collection[str], - dataset: DatasetModel, + asset: AssetModel, *, session: Session, ) -> None: - already_related = {m.name for m in dataset.aliases} + already_related = {m.name for m in asset.aliases} existing_aliases = { m.name: m - for m in session.scalars(select(DatasetAliasModel).where(DatasetAliasModel.name.in_(alias_names))) + for m in session.scalars(select(AssetAliasModel).where(AssetAliasModel.name.in_(alias_names))) } - dataset.aliases.extend( - existing_aliases.get(name, DatasetAliasModel(name=name)) + asset.aliases.extend( + existing_aliases.get(name, AssetAliasModel(name=name)) for name in alias_names if name not in already_related ) @classmethod @internal_api_call - def register_dataset_change( + def register_asset_change( cls, *, task_instance: TaskInstance | None = None, - dataset: Dataset, + asset: Asset, extra=None, - aliases: Collection[DatasetAlias] = (), + aliases: Collection[AssetAlias] = (), source_alias_names: Iterable[str] | None = None, session: Session, **kwargs, - ) -> DatasetEvent | None: + ) -> AssetEvent | None: """ - Register dataset related changes. + Register asset related changes. - For local datasets, look them up, record the dataset event, queue dagruns, and broadcast - the dataset event + For local assets, look them up, record the asset event, queue dagruns, and broadcast + the asset event """ # todo: add test so that all usages of internal_api_call are added to rpc endpoint - dataset_model = session.scalar( - select(DatasetModel) - .where(DatasetModel.uri == dataset.uri) + asset_model = session.scalar( + select(AssetModel) + .where(AssetModel.uri == asset.uri) .options( - joinedload(DatasetModel.aliases), - joinedload(DatasetModel.consuming_dags).joinedload(DagScheduleDatasetReference.dag), + joinedload(AssetModel.aliases), + joinedload(AssetModel.consuming_dags).joinedload(DagScheduleAssetReference.dag), ) ) - if not dataset_model: - cls.logger().warning("DatasetModel %s not found", dataset) + if not asset_model: + cls.logger().warning("AssetModel %s not found", asset) return None - cls._add_dataset_alias_association({alias.name for alias in aliases}, dataset_model, session=session) + cls._add_asset_alias_association({alias.name for alias in aliases}, asset_model, session=session) event_kwargs = { - "dataset_id": dataset_model.id, + "dataset_id": asset_model.id, "extra": extra, } if task_instance: @@ -148,67 +149,65 @@ def register_dataset_change( source_map_index=task_instance.map_index, ) - dataset_event = DatasetEvent(**event_kwargs) - session.add(dataset_event) + asset_event = AssetEvent(**event_kwargs) + session.add(asset_event) session.flush() # Ensure the event is written earlier than DDRQ entries below. - dags_to_queue_from_dataset = { - ref.dag for ref in dataset_model.consuming_dags if ref.dag.is_active and not ref.dag.is_paused + dags_to_queue_from_asset = { + ref.dag for ref in asset_model.consuming_dags if ref.dag.is_active and not ref.dag.is_paused } - dags_to_queue_from_dataset_alias = set() + dags_to_queue_from_asset_alias = set() if source_alias_names: - dataset_alias_models = session.scalars( - select(DatasetAliasModel) - .where(DatasetAliasModel.name.in_(source_alias_names)) + asset_alias_models = session.scalars( + select(AssetAliasModel) + .where(AssetAliasModel.name.in_(source_alias_names)) .options( - joinedload(DatasetAliasModel.consuming_dags).joinedload( - DagScheduleDatasetAliasReference.dag - ) + joinedload(AssetAliasModel.consuming_dags).joinedload(DagScheduleAssetAliasReference.dag) ) ).unique() - for dsa in dataset_alias_models: - dsa.dataset_events.append(dataset_event) - session.add(dsa) + for asset_alias_model in asset_alias_models: + asset_alias_model.dataset_events.append(asset_event) + session.add(asset_alias_model) - dags_to_queue_from_dataset_alias |= { + dags_to_queue_from_asset_alias |= { alias_ref.dag - for alias_ref in dsa.consuming_dags + for alias_ref in asset_alias_model.consuming_dags if alias_ref.dag.is_active and not alias_ref.dag.is_paused } - dags_to_reparse = dags_to_queue_from_dataset_alias - dags_to_queue_from_dataset + dags_to_reparse = dags_to_queue_from_asset_alias - dags_to_queue_from_asset if dags_to_reparse: file_locs = {dag.fileloc for dag in dags_to_reparse} cls._send_dag_priority_parsing_request(file_locs, session) - cls.notify_dataset_changed(dataset=dataset) + cls.notify_asset_changed(asset=asset) - Stats.incr("dataset.updates") + Stats.incr("asset.updates") - dags_to_queue = dags_to_queue_from_dataset | dags_to_queue_from_dataset_alias - cls._queue_dagruns(dataset_id=dataset_model.id, dags_to_queue=dags_to_queue, session=session) - return dataset_event + dags_to_queue = dags_to_queue_from_asset | dags_to_queue_from_asset_alias + cls._queue_dagruns(asset_id=asset_model.id, dags_to_queue=dags_to_queue, session=session) + return asset_event @staticmethod - def notify_dataset_created(dataset: Dataset): - """Run applicable notification actions when a dataset is created.""" - get_listener_manager().hook.on_dataset_created(dataset=dataset) + def notify_asset_created(asset: Asset): + """Run applicable notification actions when an asset is created.""" + get_listener_manager().hook.on_asset_created(asset=asset) @staticmethod - def notify_dataset_alias_created(dataset_alias: DatasetAlias): - """Run applicable notification actions when a dataset alias is created.""" - get_listener_manager().hook.on_dataset_alias_created(dataset_alias=dataset_alias) + def notify_asset_alias_created(asset_assets: AssetAlias): + """Run applicable notification actions when an asset alias is created.""" + get_listener_manager().hook.on_asset_alias_created(asset_alias=asset_assets) @staticmethod - def notify_dataset_changed(dataset: Dataset): - """Run applicable notification actions when a dataset is changed.""" - get_listener_manager().hook.on_dataset_changed(dataset=dataset) + def notify_asset_changed(asset: Asset): + """Run applicable notification actions when an asset is changed.""" + get_listener_manager().hook.on_asset_changed(asset=asset) @classmethod - def _queue_dagruns(cls, dataset_id: int, dags_to_queue: set[DagModel], session: Session) -> None: + def _queue_dagruns(cls, asset_id: int, dags_to_queue: set[DagModel], session: Session) -> None: # Possible race condition: if multiple dags or multiple (usually - # mapped) tasks update the same dataset, this can fail with a unique + # mapped) tasks update the same asset, this can fail with a unique # constraint violation. # # If we support it, use ON CONFLICT to do nothing, otherwise @@ -219,15 +218,13 @@ def _queue_dagruns(cls, dataset_id: int, dags_to_queue: set[DagModel], session: return if session.bind.dialect.name == "postgresql": - return cls._postgres_queue_dagruns(dataset_id, dags_to_queue, session) - return cls._slow_path_queue_dagruns(dataset_id, dags_to_queue, session) + return cls._postgres_queue_dagruns(asset_id, dags_to_queue, session) + return cls._slow_path_queue_dagruns(asset_id, dags_to_queue, session) @classmethod - def _slow_path_queue_dagruns( - cls, dataset_id: int, dags_to_queue: set[DagModel], session: Session - ) -> None: + def _slow_path_queue_dagruns(cls, asset_id: int, dags_to_queue: set[DagModel], session: Session) -> None: def _queue_dagrun_if_needed(dag: DagModel) -> str | None: - item = DatasetDagRunQueue(target_dag_id=dag.dag_id, dataset_id=dataset_id) + item = AssetDagRunQueue(target_dag_id=dag.dag_id, dataset_id=asset_id) # Don't error whole transaction when a single RunQueue item conflicts. # https://docs.sqlalchemy.org/en/14/orm/session_transaction.html#using-savepoint try: @@ -242,11 +239,11 @@ def _queue_dagrun_if_needed(dag: DagModel) -> str | None: cls.logger().debug("consuming dag ids %s", queued_dag_ids) @classmethod - def _postgres_queue_dagruns(cls, dataset_id: int, dags_to_queue: set[DagModel], session: Session) -> None: + def _postgres_queue_dagruns(cls, asset_id: int, dags_to_queue: set[DagModel], session: Session) -> None: from sqlalchemy.dialects.postgresql import insert values = [{"target_dag_id": dag.dag_id} for dag in dags_to_queue] - stmt = insert(DatasetDagRunQueue).values(dataset_id=dataset_id).on_conflict_do_nothing() + stmt = insert(AssetDagRunQueue).values(dataset_id=asset_id).on_conflict_do_nothing() session.execute(stmt, values) @classmethod @@ -279,19 +276,19 @@ def _postgres_send_dag_priority_parsing_request(cls, file_locs: Iterable[str], s session.execute(stmt, {"fileloc": fileloc for fileloc in file_locs}) -def resolve_dataset_manager() -> DatasetManager: - """Retrieve the dataset manager.""" - _dataset_manager_class = conf.getimport( +def resolve_asset_manager() -> AssetManager: + """Retrieve the asset manager.""" + _asset_manager_class = conf.getimport( section="core", - key="dataset_manager_class", - fallback="airflow.datasets.manager.DatasetManager", + key="asset_manager_class", + fallback="airflow.assets.manager.AssetManager", ) - _dataset_manager_kwargs = conf.getjson( + _asset_manager_kwargs = conf.getjson( section="core", - key="dataset_manager_kwargs", + key="asset_manager_kwargs", fallback={}, ) - return _dataset_manager_class(**_dataset_manager_kwargs) + return _asset_manager_class(**_asset_manager_kwargs) -dataset_manager = resolve_dataset_manager() +asset_manager = resolve_asset_manager() diff --git a/airflow/datasets/metadata.py b/airflow/assets/metadata.py similarity index 80% rename from airflow/datasets/metadata.py rename to airflow/assets/metadata.py index 43dff9287365c..4fd2902afc8bf 100644 --- a/airflow/datasets/metadata.py +++ b/airflow/assets/metadata.py @@ -21,26 +21,26 @@ import attrs -from airflow.datasets import DatasetAlias, extract_event_key +from airflow.assets import AssetAlias, extract_event_key if TYPE_CHECKING: - from airflow.datasets import Dataset + from airflow.assets import Asset @attrs.define(init=False) class Metadata: - """Metadata to attach to a DatasetEvent.""" + """Metadata to attach to a AssetEvent.""" uri: str extra: dict[str, Any] alias_name: str | None = None def __init__( - self, target: str | Dataset, extra: dict[str, Any], alias: DatasetAlias | str | None = None + self, target: str | Asset, extra: dict[str, Any], alias: AssetAlias | str | None = None ) -> None: self.uri = extract_event_key(target) self.extra = extra - if isinstance(alias, DatasetAlias): + if isinstance(alias, AssetAlias): self.alias_name = alias.name else: self.alias_name = alias diff --git a/airflow/auth/managers/base_auth_manager.py b/airflow/auth/managers/base_auth_manager.py index 6be53da0807e0..69b5969c827c6 100644 --- a/airflow/auth/managers/base_auth_manager.py +++ b/airflow/auth/managers/base_auth_manager.py @@ -46,10 +46,10 @@ ) from airflow.auth.managers.models.resource_details import ( AccessView, + AssetDetails, ConfigurationDetails, ConnectionDetails, DagAccessEntity, - DatasetDetails, PoolDetails, VariableDetails, ) @@ -178,18 +178,18 @@ def is_authorized_dag( """ @abstractmethod - def is_authorized_dataset( + def is_authorized_asset( self, *, method: ResourceMethod, - details: DatasetDetails | None = None, + details: AssetDetails | None = None, user: BaseUser | None = None, ) -> bool: """ - Return whether the user is authorized to perform a given action on a dataset. + Return whether the user is authorized to perform a given action on an asset. :param method: the method to perform - :param details: optional details about the dataset + :param details: optional details about the asset :param user: the user to perform the action on. If not provided (or None), it uses the current user """ diff --git a/airflow/auth/managers/models/resource_details.py b/airflow/auth/managers/models/resource_details.py index fcbee5a2ad299..6dec2236bf233 100644 --- a/airflow/auth/managers/models/resource_details.py +++ b/airflow/auth/managers/models/resource_details.py @@ -43,8 +43,8 @@ class DagDetails: @dataclass -class DatasetDetails: - """Represents the details of a dataset.""" +class AssetDetails: + """Represents the details of an asset.""" uri: str | None = None diff --git a/airflow/auth/managers/simple/simple_auth_manager.py b/airflow/auth/managers/simple/simple_auth_manager.py index a683aa5472cef..451068733667c 100644 --- a/airflow/auth/managers/simple/simple_auth_manager.py +++ b/airflow/auth/managers/simple/simple_auth_manager.py @@ -36,11 +36,11 @@ from airflow.auth.managers.models.base_user import BaseUser from airflow.auth.managers.models.resource_details import ( AccessView, + AssetDetails, ConfigurationDetails, ConnectionDetails, DagAccessEntity, DagDetails, - DatasetDetails, PoolDetails, VariableDetails, ) @@ -163,8 +163,8 @@ def is_authorized_dag( allow_role=SimpleAuthManagerRole.USER, ) - def is_authorized_dataset( - self, *, method: ResourceMethod, details: DatasetDetails | None = None, user: BaseUser | None = None + def is_authorized_asset( + self, *, method: ResourceMethod, details: AssetDetails | None = None, user: BaseUser | None = None ) -> bool: return self._is_authorized( method=method, diff --git a/airflow/config_templates/config.yml b/airflow/config_templates/config.yml index 7317fce60e4e6..a6d40a48c039e 100644 --- a/airflow/config_templates/config.yml +++ b/airflow/config_templates/config.yml @@ -470,22 +470,22 @@ core: type: string default: "0o077" example: ~ - dataset_manager_class: - description: Class to use as dataset manager. - version_added: 2.4.0 + asset_manager_class: + description: Class to use as asset manager. + version_added: 3.0.0 type: string default: ~ - example: 'airflow.datasets.manager.DatasetManager' - dataset_manager_kwargs: - description: Kwargs to supply to dataset manager. - version_added: 2.4.0 + example: 'airflow.datasets.manager.AssetManager' + asset_manager_kwargs: + description: Kwargs to supply to asset manager. + version_added: 3.0.0 type: string sensitive: true default: ~ example: '{"some_param": "some_value"}' - strict_dataset_uri_validation: + strict_asset_uri_validation: description: | - Dataset URI validation should raise an exception if it is not compliant with AIP-60. + Asset URI validation should raise an exception if it is not compliant with AIP-60. By default this configuration is false, meaning that Airflow 2.x only warns the user. In Airflow 3, this configuration will be enabled by default. default: "False" diff --git a/airflow/dag_processing/collection.py b/airflow/dag_processing/collection.py index bcac479d875a3..c8ce5dc873afa 100644 --- a/airflow/dag_processing/collection.py +++ b/airflow/dag_processing/collection.py @@ -35,17 +35,17 @@ from sqlalchemy.orm import joinedload, load_only from sqlalchemy.sql import expression -from airflow.datasets import Dataset, DatasetAlias -from airflow.datasets.manager import dataset_manager +from airflow.assets import Asset, AssetAlias +from airflow.assets.manager import asset_manager +from airflow.models.asset import ( + AssetAliasModel, + AssetModel, + DagScheduleAssetAliasReference, + DagScheduleAssetReference, + TaskOutletAssetReference, +) from airflow.models.dag import DAG, DagModel, DagOwnerAttributes, DagTag from airflow.models.dagrun import DagRun -from airflow.models.dataset import ( - DagScheduleDatasetAliasReference, - DagScheduleDatasetReference, - DatasetAliasModel, - DatasetModel, - TaskOutletDatasetReference, -) from airflow.utils.sqlalchemy import with_row_locks from airflow.utils.timezone import utcnow from airflow.utils.types import DagRunType @@ -209,7 +209,7 @@ def update_dags( ) dm.timetable_summary = dag.timetable.summary dm.timetable_description = dag.timetable.description - dm.dataset_expression = dag.timetable.dataset_condition.as_expression() + dm.dataset_expression = dag.timetable.asset_condition.as_expression() dm.processor_subdir = processor_subdir last_automated_run: DagRun | None = run_info.latest_runs.get(dag.dag_id) @@ -222,7 +222,7 @@ def update_dags( else: dm.calculate_dagrun_date_fields(dag, last_automated_data_interval) - if not dag.timetable.dataset_condition: + if not dag.timetable.asset_condition: dm.schedule_dataset_references = [] dm.schedule_dataset_alias_references = [] # FIXME: STORE NEW REFERENCES. @@ -237,44 +237,44 @@ def update_dags( dm.dag_owner_links = [] -def _find_all_datasets(dags: Iterable[DAG]) -> Iterator[Dataset]: +def _find_all_assets(dags: Iterable[DAG]) -> Iterator[Asset]: for dag in dags: - for _, dataset in dag.timetable.dataset_condition.iter_datasets(): - yield dataset + for _, asset in dag.timetable.asset_condition.iter_assets(): + yield asset for task in dag.task_dict.values(): for obj in itertools.chain(task.inlets, task.outlets): - if isinstance(obj, Dataset): + if isinstance(obj, Asset): yield obj -def _find_all_dataset_aliases(dags: Iterable[DAG]) -> Iterator[DatasetAlias]: +def _find_all_asset_aliases(dags: Iterable[DAG]) -> Iterator[AssetAlias]: for dag in dags: - for _, alias in dag.timetable.dataset_condition.iter_dataset_aliases(): + for _, alias in dag.timetable.asset_condition.iter_asset_aliases(): yield alias for task in dag.task_dict.values(): for obj in itertools.chain(task.inlets, task.outlets): - if isinstance(obj, DatasetAlias): + if isinstance(obj, AssetAlias): yield obj -class DatasetModelOperation(NamedTuple): - """Collect dataset/alias objects from DAGs and perform database operations for them.""" +class AssetModelOperation(NamedTuple): + """Collect asset/alias objects from DAGs and perform database operations for them.""" - schedule_dataset_references: dict[str, list[Dataset]] - schedule_dataset_alias_references: dict[str, list[DatasetAlias]] - outlet_references: dict[str, list[tuple[str, Dataset]]] - datasets: dict[str, Dataset] - dataset_aliases: dict[str, DatasetAlias] + schedule_asset_references: dict[str, list[Asset]] + schedule_asset_alias_references: dict[str, list[AssetAlias]] + outlet_references: dict[str, list[tuple[str, Asset]]] + assets: dict[str, Asset] + asset_aliases: dict[str, AssetAlias] @classmethod def collect(cls, dags: dict[str, DAG]) -> Self: coll = cls( - schedule_dataset_references={ - dag_id: [dataset for _, dataset in dag.timetable.dataset_condition.iter_datasets()] + schedule_asset_references={ + dag_id: [asset for _, asset in dag.timetable.asset_condition.iter_assets()] for dag_id, dag in dags.items() }, - schedule_dataset_alias_references={ - dag_id: [alias for _, alias in dag.timetable.dataset_condition.iter_dataset_aliases()] + schedule_asset_alias_references={ + dag_id: [alias for _, alias in dag.timetable.asset_condition.iter_asset_aliases()] for dag_id, dag in dags.items() }, outlet_references={ @@ -282,90 +282,89 @@ def collect(cls, dags: dict[str, DAG]) -> Self: (task_id, outlet) for task_id, task in dag.task_dict.items() for outlet in task.outlets - if isinstance(outlet, Dataset) + if isinstance(outlet, Asset) ] for dag_id, dag in dags.items() }, - datasets={dataset.uri: dataset for dataset in _find_all_datasets(dags.values())}, - dataset_aliases={alias.name: alias for alias in _find_all_dataset_aliases(dags.values())}, + assets={asset.uri: asset for asset in _find_all_assets(dags.values())}, + asset_aliases={alias.name: alias for alias in _find_all_asset_aliases(dags.values())}, ) return coll - def add_datasets(self, *, session: Session) -> dict[str, DatasetModel]: - # Optimization: skip all database calls if no datasets were collected. - if not self.datasets: + def add_assets(self, *, session: Session) -> dict[str, AssetModel]: + # Optimization: skip all database calls if no assets were collected. + if not self.assets: return {} - orm_datasets: dict[str, DatasetModel] = { - dm.uri: dm - for dm in session.scalars(select(DatasetModel).where(DatasetModel.uri.in_(self.datasets))) + orm_assets: dict[str, AssetModel] = { + am.uri: am for am in session.scalars(select(AssetModel).where(AssetModel.uri.in_(self.assets))) } - for model in orm_datasets.values(): + for model in orm_assets.values(): model.is_orphaned = expression.false() - orm_datasets.update( + orm_assets.update( (model.uri, model) - for model in dataset_manager.create_datasets( - [dataset for uri, dataset in self.datasets.items() if uri not in orm_datasets], + for model in asset_manager.create_assets( + [asset for uri, asset in self.assets.items() if uri not in orm_assets], session=session, ) ) - return orm_datasets + return orm_assets - def add_dataset_aliases(self, *, session: Session) -> dict[str, DatasetAliasModel]: - # Optimization: skip all database calls if no dataset aliases were collected. - if not self.dataset_aliases: + def add_asset_aliases(self, *, session: Session) -> dict[str, AssetAliasModel]: + # Optimization: skip all database calls if no asset aliases were collected. + if not self.asset_aliases: return {} - orm_aliases: dict[str, DatasetAliasModel] = { + orm_aliases: dict[str, AssetAliasModel] = { da.name: da for da in session.scalars( - select(DatasetAliasModel).where(DatasetAliasModel.name.in_(self.dataset_aliases)) + select(AssetAliasModel).where(AssetAliasModel.name.in_(self.asset_aliases)) ) } orm_aliases.update( (model.name, model) - for model in dataset_manager.create_dataset_aliases( - [alias for name, alias in self.dataset_aliases.items() if name not in orm_aliases], + for model in asset_manager.create_asset_aliases( + [alias for name, alias in self.asset_aliases.items() if name not in orm_aliases], session=session, ) ) return orm_aliases - def add_dag_dataset_references( + def add_dag_asset_references( self, dags: dict[str, DagModel], - datasets: dict[str, DatasetModel], + assets: dict[str, AssetModel], *, session: Session, ) -> None: - # Optimization: No datasets means there are no references to update. - if not datasets: + # Optimization: No assets means there are no references to update. + if not assets: return - for dag_id, references in self.schedule_dataset_references.items(): + for dag_id, references in self.schedule_asset_references.items(): # Optimization: no references at all; this is faster than repeated delete(). if not references: dags[dag_id].schedule_dataset_references = [] continue - referenced_dataset_ids = {dataset.id for dataset in (datasets[r.uri] for r in references)} + referenced_asset_ids = {asset.id for asset in (assets[r.uri] for r in references)} orm_refs = {r.dataset_id: r for r in dags[dag_id].schedule_dataset_references} - for dataset_id, ref in orm_refs.items(): - if dataset_id not in referenced_dataset_ids: + for asset_id, ref in orm_refs.items(): + if asset_id not in referenced_asset_ids: session.delete(ref) session.bulk_save_objects( - DagScheduleDatasetReference(dataset_id=dataset_id, dag_id=dag_id) - for dataset_id in referenced_dataset_ids - if dataset_id not in orm_refs + DagScheduleAssetReference(dataset_id=asset_id, dag_id=dag_id) + for asset_id in referenced_asset_ids + if asset_id not in orm_refs ) - def add_dag_dataset_alias_references( + def add_dag_asset_alias_references( self, dags: dict[str, DagModel], - aliases: dict[str, DatasetAliasModel], + aliases: dict[str, AssetAliasModel], *, session: Session, ) -> None: # Optimization: No aliases means there are no references to update. if not aliases: return - for dag_id, references in self.schedule_dataset_alias_references.items(): + for dag_id, references in self.schedule_asset_alias_references.items(): # Optimization: no references at all; this is faster than repeated delete(). if not references: dags[dag_id].schedule_dataset_alias_references = [] @@ -376,20 +375,20 @@ def add_dag_dataset_alias_references( if alias_id not in referenced_alias_ids: session.delete(ref) session.bulk_save_objects( - DagScheduleDatasetAliasReference(alias_id=alias_id, dag_id=dag_id) + DagScheduleAssetAliasReference(alias_id=alias_id, dag_id=dag_id) for alias_id in referenced_alias_ids if alias_id not in orm_refs ) - def add_task_dataset_references( + def add_task_asset_references( self, dags: dict[str, DagModel], - datasets: dict[str, DatasetModel], + assets: dict[str, AssetModel], *, session: Session, ) -> None: - # Optimization: No datasets means there are no references to update. - if not datasets: + # Optimization: No assets means there are no references to update. + if not assets: return for dag_id, references in self.outlet_references.items(): # Optimization: no references at all; this is faster than repeated delete(). @@ -397,15 +396,15 @@ def add_task_dataset_references( dags[dag_id].task_outlet_dataset_references = [] continue referenced_outlets = { - (task_id, dataset.id) - for task_id, dataset in ((task_id, datasets[d.uri]) for task_id, d in references) + (task_id, asset.id) + for task_id, asset in ((task_id, assets[d.uri]) for task_id, d in references) } orm_refs = {(r.task_id, r.dataset_id): r for r in dags[dag_id].task_outlet_dataset_references} for key, ref in orm_refs.items(): if key not in referenced_outlets: session.delete(ref) session.bulk_save_objects( - TaskOutletDatasetReference(dataset_id=dataset_id, dag_id=dag_id, task_id=task_id) - for task_id, dataset_id in referenced_outlets - if (task_id, dataset_id) not in orm_refs + TaskOutletAssetReference(dataset_id=asset_id, dag_id=dag_id, task_id=task_id) + for task_id, asset_id in referenced_outlets + if (task_id, asset_id) not in orm_refs ) diff --git a/airflow/decorators/base.py b/airflow/decorators/base.py index 611d363961c51..1ef2c12c702f2 100644 --- a/airflow/decorators/base.py +++ b/airflow/decorators/base.py @@ -40,7 +40,7 @@ import re2 import typing_extensions -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.models.abstractoperator import DEFAULT_RETRIES, DEFAULT_RETRY_DELAY from airflow.models.baseoperator import ( BaseOperator, @@ -261,7 +261,7 @@ def execute(self, context: Context): # todo make this more generic (move to prepare_lineage) so it deals with non taskflow operators # as well for arg in itertools.chain(self.op_args, self.op_kwargs.values()): - if isinstance(arg, Dataset): + if isinstance(arg, Asset): self.inlets.append(arg) return_value = super().execute(context) return self._handle_output(return_value=return_value, context=context, xcom_push=self.xcom_push) @@ -270,17 +270,17 @@ def _handle_output(self, return_value: Any, context: Context, xcom_push: Callabl """ Handle logic for whether a decorator needs to push a single return value or multiple return values. - It sets outlets if any datasets are found in the returned value(s) + It sets outlets if any assets are found in the returned value(s) :param return_value: :param context: :param xcom_push: """ - if isinstance(return_value, Dataset): + if isinstance(return_value, Asset): self.outlets.append(return_value) if isinstance(return_value, list): for item in return_value: - if isinstance(item, Dataset): + if isinstance(item, Asset): self.outlets.append(item) return return_value diff --git a/airflow/example_dags/example_asset_alias.py b/airflow/example_dags/example_asset_alias.py new file mode 100644 index 0000000000000..4970b1eda2660 --- /dev/null +++ b/airflow/example_dags/example_asset_alias.py @@ -0,0 +1,101 @@ +# 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. +""" +Example DAG for demonstrating the behavior of the AssetAlias feature in Airflow, including conditional and +asset expression-based scheduling. + +Notes on usage: + +Turn on all the DAGs. + +Before running any DAG, the schedule of the "asset_alias_example_alias_consumer" DAG will show as "Unresolved AssetAlias". +This is expected because the asset alias has not been resolved into any asset yet. + +Once the "asset_s3_bucket_producer" DAG is triggered, the "asset_s3_bucket_consumer" DAG should be triggered upon completion. +This is because the asset alias "example-alias" is used to add an asset event to the asset "s3://bucket/my-task" +during the "produce_asset_events_through_asset_alias" task. +As the DAG "asset-alias-consumer" relies on asset alias "example-alias" which was previously unresolved, +the DAG "asset-alias-consumer" (along with all the DAGs in the same file) will be re-parsed and +thus update its schedule to the asset "s3://bucket/my-task" and will also be triggered. +""" + +from __future__ import annotations + +import pendulum + +from airflow import DAG +from airflow.assets import Asset, AssetAlias +from airflow.decorators import task + +with DAG( + dag_id="asset_s3_bucket_producer", + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=None, + catchup=False, + tags=["producer", "asset"], +): + + @task(outlets=[Asset("s3://bucket/my-task")]) + def produce_asset_events(): + pass + + produce_asset_events() + +with DAG( + dag_id="asset_alias_example_alias_producer", + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=None, + catchup=False, + tags=["producer", "asset-alias"], +): + + @task(outlets=[AssetAlias("example-alias")]) + def produce_asset_events_through_asset_alias(*, outlet_events=None): + bucket_name = "bucket" + object_path = "my-task" + outlet_events["example-alias"].add(Asset(f"s3://{bucket_name}/{object_path}")) + + produce_asset_events_through_asset_alias() + +with DAG( + dag_id="asset_s3_bucket_consumer", + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=[Asset("s3://bucket/my-task")], + catchup=False, + tags=["consumer", "asset"], +): + + @task + def consume_asset_event(): + pass + + consume_asset_event() + +with DAG( + dag_id="asset_alias_example_alias_consumer", + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=[AssetAlias("example-alias")], + catchup=False, + tags=["consumer", "asset-alias"], +): + + @task(inlets=[AssetAlias("example-alias")]) + def consume_asset_event_from_asset_alias(*, inlet_events=None): + for event in inlet_events[AssetAlias("example-alias")]: + print(event) + + consume_asset_event_from_asset_alias() diff --git a/airflow/example_dags/example_asset_alias_with_no_taskflow.py b/airflow/example_dags/example_asset_alias_with_no_taskflow.py new file mode 100644 index 0000000000000..3293f7e45bb94 --- /dev/null +++ b/airflow/example_dags/example_asset_alias_with_no_taskflow.py @@ -0,0 +1,108 @@ +# 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. +""" +Example DAG for demonstrating the behavior of the AssetAlias feature in Airflow, including conditional and +asset expression-based scheduling. + +Notes on usage: + +Turn on all the DAGs. + +Before running any DAG, the schedule of the "asset_alias_example_alias_consumer_with_no_taskflow" DAG will show as "unresolved AssetAlias". +This is expected because the asset alias has not been resolved into any asset yet. + +Once the "asset_s3_bucket_producer_with_no_taskflow" DAG is triggered, the "asset_s3_bucket_consumer_with_no_taskflow" DAG should be triggered upon completion. +This is because the asset alias "example-alias-no-taskflow" is used to add an asset event to the asset "s3://bucket/my-task-with-no-taskflow" +during the "produce_asset_events_through_asset_alias_with_no_taskflow" task. Also, the schedule of the "asset_alias_example_alias_consumer_with_no_taskflow" DAG should change to "Asset" as +the asset alias "example-alias-no-taskflow" is now resolved to the asset "s3://bucket/my-task-with-no-taskflow" and this DAG should also be triggered. +""" + +from __future__ import annotations + +import pendulum + +from airflow import DAG +from airflow.assets import Asset, AssetAlias +from airflow.operators.python import PythonOperator + +with DAG( + dag_id="asset_s3_bucket_producer_with_no_taskflow", + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=None, + catchup=False, + tags=["producer", "asset"], +): + + def produce_asset_events(): + pass + + PythonOperator( + task_id="produce_asset_events", + outlets=[Asset("s3://bucket/my-task-with-no-taskflow")], + python_callable=produce_asset_events, + ) + + +with DAG( + dag_id="asset_alias_example_alias_producer_with_no_taskflow", + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=None, + catchup=False, + tags=["producer", "asset-alias"], +): + + def produce_asset_events_through_asset_alias_with_no_taskflow(*, outlet_events=None): + bucket_name = "bucket" + object_path = "my-task" + outlet_events["example-alias-no-taskflow"].add(Asset(f"s3://{bucket_name}/{object_path}")) + + PythonOperator( + task_id="produce_asset_events_through_asset_alias_with_no_taskflow", + outlets=[AssetAlias("example-alias-no-taskflow")], + python_callable=produce_asset_events_through_asset_alias_with_no_taskflow, + ) + +with DAG( + dag_id="asset_s3_bucket_consumer_with_no_taskflow", + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=[Asset("s3://bucket/my-task-with-no-taskflow")], + catchup=False, + tags=["consumer", "asset"], +): + + def consume_asset_event(): + pass + + PythonOperator(task_id="consume_asset_event", python_callable=consume_asset_event) + +with DAG( + dag_id="asset_alias_example_alias_consumer_with_no_taskflow", + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=[AssetAlias("example-alias-no-taskflow")], + catchup=False, + tags=["consumer", "asset-alias"], +): + + def consume_asset_event_from_asset_alias(*, inlet_events=None): + for event in inlet_events[AssetAlias("example-alias-no-taskflow")]: + print(event) + + PythonOperator( + task_id="consume_asset_event_from_asset_alias", + python_callable=consume_asset_event_from_asset_alias, + inlets=[AssetAlias("example-alias-no-taskflow")], + ) diff --git a/airflow/example_dags/example_assets.py b/airflow/example_dags/example_assets.py new file mode 100644 index 0000000000000..66369794ed999 --- /dev/null +++ b/airflow/example_dags/example_assets.py @@ -0,0 +1,192 @@ +# 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. +""" +Example DAG for demonstrating the behavior of the Assets feature in Airflow, including conditional and +asset expression-based scheduling. + +Notes on usage: + +Turn on all the DAGs. + +asset_produces_1 is scheduled to run daily. Once it completes, it triggers several DAGs due to its asset +being updated. asset_consumes_1 is triggered immediately, as it depends solely on the asset produced by +asset_produces_1. consume_1_or_2_with_asset_expressions will also be triggered, as its condition of +either asset_produces_1 or asset_produces_2 being updated is satisfied with asset_produces_1. + +asset_consumes_1_and_2 will not be triggered after asset_produces_1 runs because it requires the asset +from asset_produces_2, which has no schedule and must be manually triggered. + +After manually triggering asset_produces_2, several DAGs will be affected. asset_consumes_1_and_2 should +run because both its asset dependencies are now met. consume_1_and_2_with_asset_expressions will be +triggered, as it requires both asset_produces_1 and asset_produces_2 assets to be updated. +consume_1_or_2_with_asset_expressions will be triggered again, since it's conditionally set to run when +either asset is updated. + +consume_1_or_both_2_and_3_with_asset_expressions demonstrates complex asset dependency logic. +This DAG triggers if asset_produces_1 is updated or if both asset_produces_2 and dag3_asset +are updated. This example highlights the capability to combine updates from multiple assets with logical +expressions for advanced scheduling. + +conditional_asset_and_time_based_timetable illustrates the integration of time-based scheduling with +asset dependencies. This DAG is configured to execute either when both asset_produces_1 and +asset_produces_2 assets have been updated or according to a specific cron schedule, showcasing +Airflow's versatility in handling mixed triggers for asset and time-based scheduling. + +The DAGs asset_consumes_1_never_scheduled and asset_consumes_unknown_never_scheduled will not run +automatically as they depend on assets that do not get updated or are not produced by any scheduled tasks. +""" + +from __future__ import annotations + +import pendulum + +from airflow.assets import Asset +from airflow.models.dag import DAG +from airflow.operators.bash import BashOperator +from airflow.timetables.assets import AssetOrTimeSchedule +from airflow.timetables.trigger import CronTriggerTimetable + +# [START asset_def] +dag1_asset = Asset("s3://dag1/output_1.txt", extra={"hi": "bye"}) +# [END asset_def] +dag2_asset = Asset("s3://dag2/output_1.txt", extra={"hi": "bye"}) +dag3_asset = Asset("s3://dag3/output_3.txt", extra={"hi": "bye"}) + +with DAG( + dag_id="asset_produces_1", + catchup=False, + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule="@daily", + tags=["produces", "asset-scheduled"], +) as dag1: + # [START task_outlet] + BashOperator(outlets=[dag1_asset], task_id="producing_task_1", bash_command="sleep 5") + # [END task_outlet] + +with DAG( + dag_id="asset_produces_2", + catchup=False, + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=None, + tags=["produces", "asset-scheduled"], +) as dag2: + BashOperator(outlets=[dag2_asset], task_id="producing_task_2", bash_command="sleep 5") + +# [START dag_dep] +with DAG( + dag_id="asset_consumes_1", + catchup=False, + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=[dag1_asset], + tags=["consumes", "asset-scheduled"], +) as dag3: + # [END dag_dep] + BashOperator( + outlets=[Asset("s3://consuming_1_task/asset_other.txt")], + task_id="consuming_1", + bash_command="sleep 5", + ) + +with DAG( + dag_id="asset_consumes_1_and_2", + catchup=False, + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=[dag1_asset, dag2_asset], + tags=["consumes", "asset-scheduled"], +) as dag4: + BashOperator( + outlets=[Asset("s3://consuming_2_task/asset_other_unknown.txt")], + task_id="consuming_2", + bash_command="sleep 5", + ) + +with DAG( + dag_id="asset_consumes_1_never_scheduled", + catchup=False, + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=[ + dag1_asset, + Asset("s3://unrelated/this-asset-doesnt-get-triggered"), + ], + tags=["consumes", "asset-scheduled"], +) as dag5: + BashOperator( + outlets=[Asset("s3://consuming_2_task/asset_other_unknown.txt")], + task_id="consuming_3", + bash_command="sleep 5", + ) + +with DAG( + dag_id="asset_consumes_unknown_never_scheduled", + catchup=False, + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=[ + Asset("s3://unrelated/asset3.txt"), + Asset("s3://unrelated/asset_other_unknown.txt"), + ], + tags=["asset-scheduled"], +) as dag6: + BashOperator( + task_id="unrelated_task", + outlets=[Asset("s3://unrelated_task/asset_other_unknown.txt")], + bash_command="sleep 5", + ) + +with DAG( + dag_id="consume_1_and_2_with_asset_expressions", + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=(dag1_asset & dag2_asset), +) as dag5: + BashOperator( + outlets=[Asset("s3://consuming_2_task/asset_other_unknown.txt")], + task_id="consume_1_and_2_with_asset_expressions", + bash_command="sleep 5", + ) +with DAG( + dag_id="consume_1_or_2_with_asset_expressions", + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=(dag1_asset | dag2_asset), +) as dag6: + BashOperator( + outlets=[Asset("s3://consuming_2_task/asset_other_unknown.txt")], + task_id="consume_1_or_2_with_asset_expressions", + bash_command="sleep 5", + ) +with DAG( + dag_id="consume_1_or_both_2_and_3_with_asset_expressions", + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=(dag1_asset | (dag2_asset & dag3_asset)), +) as dag7: + BashOperator( + outlets=[Asset("s3://consuming_2_task/asset_other_unknown.txt")], + task_id="consume_1_or_both_2_and_3_with_asset_expressions", + bash_command="sleep 5", + ) +with DAG( + dag_id="conditional_asset_and_time_based_timetable", + catchup=False, + start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), + schedule=AssetOrTimeSchedule( + timetable=CronTriggerTimetable("0 1 * * 3", timezone="UTC"), assets=(dag1_asset & dag2_asset) + ), + tags=["asset-time-based-timetable"], +) as dag8: + BashOperator( + outlets=[Asset("s3://asset_time_based/asset_other_unknown.txt")], + task_id="conditional_asset_and_time_based_timetable", + bash_command="sleep 5", + ) diff --git a/airflow/example_dags/example_dataset_alias.py b/airflow/example_dags/example_dataset_alias.py deleted file mode 100644 index c50a89e34fb8c..0000000000000 --- a/airflow/example_dags/example_dataset_alias.py +++ /dev/null @@ -1,101 +0,0 @@ -# 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. -""" -Example DAG for demonstrating the behavior of the DatasetAlias feature in Airflow, including conditional and -dataset expression-based scheduling. - -Notes on usage: - -Turn on all the DAGs. - -Before running any DAG, the schedule of the "dataset_alias_example_alias_consumer" DAG will show as "Unresolved DatasetAlias". -This is expected because the dataset alias has not been resolved into any dataset yet. - -Once the "dataset_s3_bucket_producer" DAG is triggered, the "dataset_s3_bucket_consumer" DAG should be triggered upon completion. -This is because the dataset alias "example-alias" is used to add a dataset event to the dataset "s3://bucket/my-task" -during the "produce_dataset_events_through_dataset_alias" task. -As the DAG "dataset-alias-consumer" relies on dataset alias "example-alias" which was previously unresolved, -the DAG "dataset-alias-consumer" (along with all the DAGs in the same file) will be re-parsed and -thus update its schedule to the dataset "s3://bucket/my-task" and will also be triggered. -""" - -from __future__ import annotations - -import pendulum - -from airflow import DAG -from airflow.datasets import Dataset, DatasetAlias -from airflow.decorators import task - -with DAG( - dag_id="dataset_s3_bucket_producer", - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=None, - catchup=False, - tags=["producer", "dataset"], -): - - @task(outlets=[Dataset("s3://bucket/my-task")]) - def produce_dataset_events(): - pass - - produce_dataset_events() - -with DAG( - dag_id="dataset_alias_example_alias_producer", - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=None, - catchup=False, - tags=["producer", "dataset-alias"], -): - - @task(outlets=[DatasetAlias("example-alias")]) - def produce_dataset_events_through_dataset_alias(*, outlet_events=None): - bucket_name = "bucket" - object_path = "my-task" - outlet_events["example-alias"].add(Dataset(f"s3://{bucket_name}/{object_path}")) - - produce_dataset_events_through_dataset_alias() - -with DAG( - dag_id="dataset_s3_bucket_consumer", - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=[Dataset("s3://bucket/my-task")], - catchup=False, - tags=["consumer", "dataset"], -): - - @task - def consume_dataset_event(): - pass - - consume_dataset_event() - -with DAG( - dag_id="dataset_alias_example_alias_consumer", - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=[DatasetAlias("example-alias")], - catchup=False, - tags=["consumer", "dataset-alias"], -): - - @task(inlets=[DatasetAlias("example-alias")]) - def consume_dataset_event_from_dataset_alias(*, inlet_events=None): - for event in inlet_events[DatasetAlias("example-alias")]: - print(event) - - consume_dataset_event_from_dataset_alias() diff --git a/airflow/example_dags/example_dataset_alias_with_no_taskflow.py b/airflow/example_dags/example_dataset_alias_with_no_taskflow.py deleted file mode 100644 index 7d7227af39f50..0000000000000 --- a/airflow/example_dags/example_dataset_alias_with_no_taskflow.py +++ /dev/null @@ -1,108 +0,0 @@ -# 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. -""" -Example DAG for demonstrating the behavior of the DatasetAlias feature in Airflow, including conditional and -dataset expression-based scheduling. - -Notes on usage: - -Turn on all the DAGs. - -Before running any DAG, the schedule of the "dataset_alias_example_alias_consumer_with_no_taskflow" DAG will show as "unresolved DatasetAlias". -This is expected because the dataset alias has not been resolved into any dataset yet. - -Once the "dataset_s3_bucket_producer_with_no_taskflow" DAG is triggered, the "dataset_s3_bucket_consumer_with_no_taskflow" DAG should be triggered upon completion. -This is because the dataset alias "example-alias-no-taskflow" is used to add a dataset event to the dataset "s3://bucket/my-task-with-no-taskflow" -during the "produce_dataset_events_through_dataset_alias_with_no_taskflow" task. Also, the schedule of the "dataset_alias_example_alias_consumer_with_no_taskflow" DAG should change to "Dataset" as -the dataset alias "example-alias-no-taskflow" is now resolved to the dataset "s3://bucket/my-task-with-no-taskflow" and this DAG should also be triggered. -""" - -from __future__ import annotations - -import pendulum - -from airflow import DAG -from airflow.datasets import Dataset, DatasetAlias -from airflow.operators.python import PythonOperator - -with DAG( - dag_id="dataset_s3_bucket_producer_with_no_taskflow", - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=None, - catchup=False, - tags=["producer", "dataset"], -): - - def produce_dataset_events(): - pass - - PythonOperator( - task_id="produce_dataset_events", - outlets=[Dataset("s3://bucket/my-task-with-no-taskflow")], - python_callable=produce_dataset_events, - ) - - -with DAG( - dag_id="dataset_alias_example_alias_producer_with_no_taskflow", - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=None, - catchup=False, - tags=["producer", "dataset-alias"], -): - - def produce_dataset_events_through_dataset_alias_with_no_taskflow(*, outlet_events=None): - bucket_name = "bucket" - object_path = "my-task" - outlet_events["example-alias-no-taskflow"].add(Dataset(f"s3://{bucket_name}/{object_path}")) - - PythonOperator( - task_id="produce_dataset_events_through_dataset_alias_with_no_taskflow", - outlets=[DatasetAlias("example-alias-no-taskflow")], - python_callable=produce_dataset_events_through_dataset_alias_with_no_taskflow, - ) - -with DAG( - dag_id="dataset_s3_bucket_consumer_with_no_taskflow", - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=[Dataset("s3://bucket/my-task-with-no-taskflow")], - catchup=False, - tags=["consumer", "dataset"], -): - - def consume_dataset_event(): - pass - - PythonOperator(task_id="consume_dataset_event", python_callable=consume_dataset_event) - -with DAG( - dag_id="dataset_alias_example_alias_consumer_with_no_taskflow", - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=[DatasetAlias("example-alias-no-taskflow")], - catchup=False, - tags=["consumer", "dataset-alias"], -): - - def consume_dataset_event_from_dataset_alias(*, inlet_events=None): - for event in inlet_events[DatasetAlias("example-alias-no-taskflow")]: - print(event) - - PythonOperator( - task_id="consume_dataset_event_from_dataset_alias", - python_callable=consume_dataset_event_from_dataset_alias, - inlets=[DatasetAlias("example-alias-no-taskflow")], - ) diff --git a/airflow/example_dags/example_datasets.py b/airflow/example_dags/example_datasets.py deleted file mode 100644 index 54f15d8a2d802..0000000000000 --- a/airflow/example_dags/example_datasets.py +++ /dev/null @@ -1,192 +0,0 @@ -# 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. -""" -Example DAG for demonstrating the behavior of the Datasets feature in Airflow, including conditional and -dataset expression-based scheduling. - -Notes on usage: - -Turn on all the DAGs. - -dataset_produces_1 is scheduled to run daily. Once it completes, it triggers several DAGs due to its dataset -being updated. dataset_consumes_1 is triggered immediately, as it depends solely on the dataset produced by -dataset_produces_1. consume_1_or_2_with_dataset_expressions will also be triggered, as its condition of -either dataset_produces_1 or dataset_produces_2 being updated is satisfied with dataset_produces_1. - -dataset_consumes_1_and_2 will not be triggered after dataset_produces_1 runs because it requires the dataset -from dataset_produces_2, which has no schedule and must be manually triggered. - -After manually triggering dataset_produces_2, several DAGs will be affected. dataset_consumes_1_and_2 should -run because both its dataset dependencies are now met. consume_1_and_2_with_dataset_expressions will be -triggered, as it requires both dataset_produces_1 and dataset_produces_2 datasets to be updated. -consume_1_or_2_with_dataset_expressions will be triggered again, since it's conditionally set to run when -either dataset is updated. - -consume_1_or_both_2_and_3_with_dataset_expressions demonstrates complex dataset dependency logic. -This DAG triggers if dataset_produces_1 is updated or if both dataset_produces_2 and dag3_dataset -are updated. This example highlights the capability to combine updates from multiple datasets with logical -expressions for advanced scheduling. - -conditional_dataset_and_time_based_timetable illustrates the integration of time-based scheduling with -dataset dependencies. This DAG is configured to execute either when both dataset_produces_1 and -dataset_produces_2 datasets have been updated or according to a specific cron schedule, showcasing -Airflow's versatility in handling mixed triggers for dataset and time-based scheduling. - -The DAGs dataset_consumes_1_never_scheduled and dataset_consumes_unknown_never_scheduled will not run -automatically as they depend on datasets that do not get updated or are not produced by any scheduled tasks. -""" - -from __future__ import annotations - -import pendulum - -from airflow.datasets import Dataset -from airflow.models.dag import DAG -from airflow.operators.bash import BashOperator -from airflow.timetables.datasets import DatasetOrTimeSchedule -from airflow.timetables.trigger import CronTriggerTimetable - -# [START dataset_def] -dag1_dataset = Dataset("s3://dag1/output_1.txt", extra={"hi": "bye"}) -# [END dataset_def] -dag2_dataset = Dataset("s3://dag2/output_1.txt", extra={"hi": "bye"}) -dag3_dataset = Dataset("s3://dag3/output_3.txt", extra={"hi": "bye"}) - -with DAG( - dag_id="dataset_produces_1", - catchup=False, - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule="@daily", - tags=["produces", "dataset-scheduled"], -) as dag1: - # [START task_outlet] - BashOperator(outlets=[dag1_dataset], task_id="producing_task_1", bash_command="sleep 5") - # [END task_outlet] - -with DAG( - dag_id="dataset_produces_2", - catchup=False, - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=None, - tags=["produces", "dataset-scheduled"], -) as dag2: - BashOperator(outlets=[dag2_dataset], task_id="producing_task_2", bash_command="sleep 5") - -# [START dag_dep] -with DAG( - dag_id="dataset_consumes_1", - catchup=False, - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=[dag1_dataset], - tags=["consumes", "dataset-scheduled"], -) as dag3: - # [END dag_dep] - BashOperator( - outlets=[Dataset("s3://consuming_1_task/dataset_other.txt")], - task_id="consuming_1", - bash_command="sleep 5", - ) - -with DAG( - dag_id="dataset_consumes_1_and_2", - catchup=False, - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=[dag1_dataset, dag2_dataset], - tags=["consumes", "dataset-scheduled"], -) as dag4: - BashOperator( - outlets=[Dataset("s3://consuming_2_task/dataset_other_unknown.txt")], - task_id="consuming_2", - bash_command="sleep 5", - ) - -with DAG( - dag_id="dataset_consumes_1_never_scheduled", - catchup=False, - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=[ - dag1_dataset, - Dataset("s3://unrelated/this-dataset-doesnt-get-triggered"), - ], - tags=["consumes", "dataset-scheduled"], -) as dag5: - BashOperator( - outlets=[Dataset("s3://consuming_2_task/dataset_other_unknown.txt")], - task_id="consuming_3", - bash_command="sleep 5", - ) - -with DAG( - dag_id="dataset_consumes_unknown_never_scheduled", - catchup=False, - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=[ - Dataset("s3://unrelated/dataset3.txt"), - Dataset("s3://unrelated/dataset_other_unknown.txt"), - ], - tags=["dataset-scheduled"], -) as dag6: - BashOperator( - task_id="unrelated_task", - outlets=[Dataset("s3://unrelated_task/dataset_other_unknown.txt")], - bash_command="sleep 5", - ) - -with DAG( - dag_id="consume_1_and_2_with_dataset_expressions", - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=(dag1_dataset & dag2_dataset), -) as dag5: - BashOperator( - outlets=[Dataset("s3://consuming_2_task/dataset_other_unknown.txt")], - task_id="consume_1_and_2_with_dataset_expressions", - bash_command="sleep 5", - ) -with DAG( - dag_id="consume_1_or_2_with_dataset_expressions", - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=(dag1_dataset | dag2_dataset), -) as dag6: - BashOperator( - outlets=[Dataset("s3://consuming_2_task/dataset_other_unknown.txt")], - task_id="consume_1_or_2_with_dataset_expressions", - bash_command="sleep 5", - ) -with DAG( - dag_id="consume_1_or_both_2_and_3_with_dataset_expressions", - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=(dag1_dataset | (dag2_dataset & dag3_dataset)), -) as dag7: - BashOperator( - outlets=[Dataset("s3://consuming_2_task/dataset_other_unknown.txt")], - task_id="consume_1_or_both_2_and_3_with_dataset_expressions", - bash_command="sleep 5", - ) -with DAG( - dag_id="conditional_dataset_and_time_based_timetable", - catchup=False, - start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), - schedule=DatasetOrTimeSchedule( - timetable=CronTriggerTimetable("0 1 * * 3", timezone="UTC"), datasets=(dag1_dataset & dag2_dataset) - ), - tags=["dataset-time-based-timetable"], -) as dag8: - BashOperator( - outlets=[Dataset("s3://dataset_time_based/dataset_other_unknown.txt")], - task_id="conditional_dataset_and_time_based_timetable", - bash_command="sleep 5", - ) diff --git a/airflow/example_dags/example_inlet_event_extra.py b/airflow/example_dags/example_inlet_event_extra.py index 4b7567fc2f87e..974534c295b79 100644 --- a/airflow/example_dags/example_inlet_event_extra.py +++ b/airflow/example_dags/example_inlet_event_extra.py @@ -16,7 +16,7 @@ # under the License. """ -Example DAG to demonstrate reading dataset events annotated with extra information. +Example DAG to demonstrate reading asset events annotated with extra information. Also see examples in ``example_outlet_event_extra.py``. """ @@ -25,37 +25,37 @@ import datetime -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.decorators import task from airflow.models.dag import DAG from airflow.operators.bash import BashOperator -ds = Dataset("s3://output/1.txt") +asset = Asset("s3://output/1.txt") with DAG( - dag_id="read_dataset_event", + dag_id="read_asset_event", catchup=False, start_date=datetime.datetime.min, schedule="@daily", tags=["consumes"], ): - @task(inlets=[ds]) - def read_dataset_event(*, inlet_events=None): - for event in inlet_events[ds][:-2]: + @task(inlets=[asset]) + def read_asset_event(*, inlet_events=None): + for event in inlet_events[asset][:-2]: print(event.extra["hi"]) - read_dataset_event() + read_asset_event() with DAG( - dag_id="read_dataset_event_from_classic", + dag_id="read_asset_event_from_classic", catchup=False, start_date=datetime.datetime.min, schedule="@daily", tags=["consumes"], ): BashOperator( - task_id="read_dataset_event_from_classic", - inlets=[ds], + task_id="read_asset_event_from_classic", + inlets=[asset], bash_command="echo '{{ inlet_events['s3://output/1.txt'][-1].extra | tojson }}'", ) diff --git a/airflow/example_dags/example_outlet_event_extra.py b/airflow/example_dags/example_outlet_event_extra.py index 5f7d986e90fdf..893090460b538 100644 --- a/airflow/example_dags/example_outlet_event_extra.py +++ b/airflow/example_dags/example_outlet_event_extra.py @@ -16,7 +16,7 @@ # under the License. """ -Example DAG to demonstrate annotating a dataset event with extra information. +Example DAG to demonstrate annotating an asset event with extra information. Also see examples in ``example_inlet_event_extra.py``. """ @@ -25,16 +25,16 @@ import datetime -from airflow.datasets import Dataset -from airflow.datasets.metadata import Metadata +from airflow.assets import Asset +from airflow.assets.metadata import Metadata from airflow.decorators import task from airflow.models.dag import DAG from airflow.operators.bash import BashOperator -ds = Dataset("s3://output/1.txt") +ds = Asset("s3://output/1.txt") with DAG( - dag_id="dataset_with_extra_by_yield", + dag_id="asset_with_extra_by_yield", catchup=False, start_date=datetime.datetime.min, schedule="@daily", @@ -42,13 +42,13 @@ ): @task(outlets=[ds]) - def dataset_with_extra_by_yield(): + def asset_with_extra_by_yield(): yield Metadata(ds, {"hi": "bye"}) - dataset_with_extra_by_yield() + asset_with_extra_by_yield() with DAG( - dag_id="dataset_with_extra_by_context", + dag_id="asset_with_extra_by_context", catchup=False, start_date=datetime.datetime.min, schedule="@daily", @@ -56,25 +56,25 @@ def dataset_with_extra_by_yield(): ): @task(outlets=[ds]) - def dataset_with_extra_by_context(*, outlet_events=None): + def asset_with_extra_by_context(*, outlet_events=None): outlet_events[ds].extra = {"hi": "bye"} - dataset_with_extra_by_context() + asset_with_extra_by_context() with DAG( - dag_id="dataset_with_extra_from_classic_operator", + dag_id="asset_with_extra_from_classic_operator", catchup=False, start_date=datetime.datetime.min, schedule="@daily", tags=["produces"], ): - def _dataset_with_extra_from_classic_operator_post_execute(context, result): + def _asset_with_extra_from_classic_operator_post_execute(context, result): context["outlet_events"][ds].extra = {"hi": "bye"} BashOperator( - task_id="dataset_with_extra_from_classic_operator", + task_id="asset_with_extra_from_classic_operator", outlets=[ds], bash_command=":", - post_execute=_dataset_with_extra_from_classic_operator_post_execute, + post_execute=_asset_with_extra_from_classic_operator_post_execute, ) diff --git a/airflow/io/path.py b/airflow/io/path.py index 6deafae004959..3526050d12883 100644 --- a/airflow/io/path.py +++ b/airflow/io/path.py @@ -56,9 +56,9 @@ def __getattr__(self, name): def wrapper(*args, **kwargs): self.log.debug("Calling method: %s", name) if name == "read": - get_hook_lineage_collector().add_input_dataset(context=self._path, uri=str(self._path)) + get_hook_lineage_collector().add_input_asset(context=self._path, uri=str(self._path)) elif name == "write": - get_hook_lineage_collector().add_output_dataset(context=self._path, uri=str(self._path)) + get_hook_lineage_collector().add_output_asset(context=self._path, uri=str(self._path)) result = attr(*args, **kwargs) return result @@ -316,8 +316,8 @@ def copy(self, dst: str | ObjectStoragePath, recursive: bool = False, **kwargs) if self.samestore(dst) or self.protocol == "file" or dst.protocol == "file": # only emit this in "optimized" variants - else lineage will be captured by file writes/reads - get_hook_lineage_collector().add_input_dataset(context=self, uri=str(self)) - get_hook_lineage_collector().add_output_dataset(context=dst, uri=str(dst)) + get_hook_lineage_collector().add_input_asset(context=self, uri=str(self)) + get_hook_lineage_collector().add_output_asset(context=dst, uri=str(dst)) # same -> same if self.samestore(dst): @@ -381,8 +381,8 @@ def move(self, path: str | ObjectStoragePath, recursive: bool = False, **kwargs) path = ObjectStoragePath(path) if self.samestore(path): - get_hook_lineage_collector().add_input_dataset(context=self, uri=str(self)) - get_hook_lineage_collector().add_output_dataset(context=path, uri=str(path)) + get_hook_lineage_collector().add_input_asset(context=self, uri=str(self)) + get_hook_lineage_collector().add_output_asset(context=path, uri=str(path)) return self.fs.move(self.path, path.path, recursive=recursive, **kwargs) # non-local copy diff --git a/airflow/jobs/scheduler_job_runner.py b/airflow/jobs/scheduler_job_runner.py index 9438edd4d9187..242154820df9e 100644 --- a/airflow/jobs/scheduler_job_runner.py +++ b/airflow/jobs/scheduler_job_runner.py @@ -44,21 +44,21 @@ from airflow.jobs.base_job_runner import BaseJobRunner from airflow.jobs.job import Job, perform_heartbeat from airflow.models import Log +from airflow.models.asset import ( + AssetDagRunQueue, + AssetEvent, + AssetModel, + DagScheduleAssetReference, + TaskOutletAssetReference, +) from airflow.models.dag import DAG, DagModel from airflow.models.dagbag import DagBag from airflow.models.dagrun import DagRun -from airflow.models.dataset import ( - DagScheduleDatasetReference, - DatasetDagRunQueue, - DatasetEvent, - DatasetModel, - TaskOutletDatasetReference, -) from airflow.models.serialized_dag import SerializedDagModel from airflow.models.taskinstance import SimpleTaskInstance, TaskInstance from airflow.stats import Stats from airflow.ti_deps.dependencies_states import EXECUTION_STATES -from airflow.timetables.simple import DatasetTriggeredTimetable +from airflow.timetables.simple import AssetTriggeredTimetable from airflow.traces import utils as trace_utils from airflow.traces.tracer import Trace, add_span from airflow.utils import timezone @@ -1086,7 +1086,7 @@ def _run_scheduler_loop(self) -> None: timers.call_regular_interval( conf.getfloat("scheduler", "parsing_cleanup_interval"), - self._orphan_unreferenced_datasets, + self._orphan_unreferenced_assets, ) if self._standalone_dag_processor: @@ -1286,9 +1286,7 @@ def _create_dagruns_for_dags(self, guard: CommitProhibitorGuard, session: Sessio non_dataset_dags = all_dags_needing_dag_runs.difference(dataset_triggered_dags) self._create_dag_runs(non_dataset_dags, session) if dataset_triggered_dags: - self._create_dag_runs_dataset_triggered( - dataset_triggered_dags, dataset_triggered_dag_info, session - ) + self._create_dag_runs_asset_triggered(dataset_triggered_dags, dataset_triggered_dag_info, session) # commit the session - Release the write lock on DagModel table. guard.commit() @@ -1367,13 +1365,13 @@ def _create_dag_runs(self, dag_models: Collection[DagModel], session: Session) - # TODO[HA]: Should we do a session.flush() so we don't have to keep lots of state/object in # memory for larger dags? or expunge_all() - def _create_dag_runs_dataset_triggered( + def _create_dag_runs_asset_triggered( self, dag_models: Collection[DagModel], dataset_triggered_dag_info: dict[str, tuple[datetime, datetime]], session: Session, ) -> None: - """For DAGs that are triggered by datasets, create dag runs.""" + """For DAGs that are triggered by assets, create dag runs.""" # Bulk Fetch DagRuns with dag_id and execution_date same # as DagModel.dag_id and DagModel.next_dagrun # This list is used to verify if the DagRun already exist so that we don't attempt to create @@ -1396,9 +1394,9 @@ def _create_dag_runs_dataset_triggered( self.log.error("DAG '%s' not found in serialized_dag table", dag_model.dag_id) continue - if not isinstance(dag.timetable, DatasetTriggeredTimetable): + if not isinstance(dag.timetable, AssetTriggeredTimetable): self.log.error( - "DAG '%s' was dataset-scheduled, but didn't have a DatasetTriggeredTimetable!", + "DAG '%s' was asset-scheduled, but didn't have a AssetTriggeredTimetable!", dag_model.dag_id, ) continue @@ -1425,29 +1423,29 @@ def _create_dag_runs_dataset_triggered( .order_by(DagRun.execution_date.desc()) .limit(1) ) - dataset_event_filters = [ - DagScheduleDatasetReference.dag_id == dag.dag_id, - DatasetEvent.timestamp <= exec_date, + asset_event_filters = [ + DagScheduleAssetReference.dag_id == dag.dag_id, + AssetEvent.timestamp <= exec_date, ] if previous_dag_run: - dataset_event_filters.append(DatasetEvent.timestamp > previous_dag_run.execution_date) + asset_event_filters.append(AssetEvent.timestamp > previous_dag_run.execution_date) - dataset_events = session.scalars( - select(DatasetEvent) + asset_events = session.scalars( + select(AssetEvent) .join( - DagScheduleDatasetReference, - DatasetEvent.dataset_id == DagScheduleDatasetReference.dataset_id, + DagScheduleAssetReference, + AssetEvent.dataset_id == DagScheduleAssetReference.dataset_id, ) - .where(*dataset_event_filters) + .where(*asset_event_filters) ).all() - data_interval = dag.timetable.data_interval_for_events(exec_date, dataset_events) + data_interval = dag.timetable.data_interval_for_events(exec_date, asset_events) run_id = dag.timetable.generate_run_id( run_type=DagRunType.DATASET_TRIGGERED, logical_date=exec_date, data_interval=data_interval, session=session, - events=dataset_events, + events=asset_events, ) dag_run = dag.create_dagrun( @@ -1462,10 +1460,10 @@ def _create_dag_runs_dataset_triggered( creating_job_id=self.job.id, triggered_by=DagRunTriggeredByType.DATASET, ) - Stats.incr("dataset.triggered_dagruns") - dag_run.consumed_dataset_events.extend(dataset_events) + Stats.incr("asset.triggered_dagruns") + dag_run.consumed_dataset_events.extend(asset_events) session.execute( - delete(DatasetDagRunQueue).where(DatasetDagRunQueue.target_dag_id == dag_run.dag_id) + delete(AssetDagRunQueue).where(AssetDagRunQueue.target_dag_id == dag_run.dag_id) ) def _should_update_dag_next_dagruns( @@ -2014,40 +2012,40 @@ def _cleanup_stale_dags(self, session: Session = NEW_SESSION) -> None: SerializedDagModel.remove_dag(dag_id=dag.dag_id, session=session) session.flush() - def _set_orphaned(self, dataset: DatasetModel) -> int: - self.log.info("Orphaning unreferenced dataset '%s'", dataset.uri) - dataset.is_orphaned = expression.true() + def _set_orphaned(self, asset: AssetModel) -> int: + self.log.info("Orphaning unreferenced asset '%s'", asset.uri) + asset.is_orphaned = expression.true() return 1 @provide_session - def _orphan_unreferenced_datasets(self, session: Session = NEW_SESSION) -> None: + def _orphan_unreferenced_assets(self, session: Session = NEW_SESSION) -> None: """ - Detect orphaned datasets and set is_orphaned flag to True. + Detect orphaned assets and set is_orphaned flag to True. - An orphaned dataset is no longer referenced in any DAG schedule parameters or task outlets. + An orphaned asset is no longer referenced in any DAG schedule parameters or task outlets. """ - orphaned_dataset_query = session.scalars( - select(DatasetModel) + orphaned_asset_query = session.scalars( + select(AssetModel) .join( - DagScheduleDatasetReference, + DagScheduleAssetReference, isouter=True, ) .join( - TaskOutletDatasetReference, + TaskOutletAssetReference, isouter=True, ) - .group_by(DatasetModel.id) - .where(~DatasetModel.is_orphaned) + .group_by(AssetModel.id) + .where(~AssetModel.is_orphaned) .having( and_( - func.count(DagScheduleDatasetReference.dag_id) == 0, - func.count(TaskOutletDatasetReference.dag_id) == 0, + func.count(DagScheduleAssetReference.dag_id) == 0, + func.count(TaskOutletAssetReference.dag_id) == 0, ) ) ) - updated_count = sum(self._set_orphaned(dataset) for dataset in orphaned_dataset_query) - Stats.gauge("dataset.orphaned", updated_count) + updated_count = sum(self._set_orphaned(asset) for asset in orphaned_asset_query) + Stats.gauge("asset.orphaned", updated_count) def _executor_to_tis(self, tis: list[TaskInstance]) -> dict[BaseExecutor, list[TaskInstance]]: """Organize TIs into lists per their respective executor.""" diff --git a/airflow/lineage/__init__.py b/airflow/lineage/__init__.py index 332a04e7250bf..4385f3fbaf586 100644 --- a/airflow/lineage/__init__.py +++ b/airflow/lineage/__init__.py @@ -104,7 +104,7 @@ def prepare_lineage(func: T) -> T: * "auto" -> picks up any outlets from direct upstream tasks that have outlets defined, as such that if A -> B -> C and B does not have outlets but A does, these are provided as inlets. * "list of task_ids" -> picks up outlets from the upstream task_ids - * "list of datasets" -> manually defined list of data + * "list of datasets" -> manually defined list of dataset """ diff --git a/airflow/lineage/hook.py b/airflow/lineage/hook.py index 4ff35e4d9ce82..fd321bcab49cf 100644 --- a/airflow/lineage/hook.py +++ b/airflow/lineage/hook.py @@ -24,7 +24,7 @@ import attr -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.providers_manager import ProvidersManager from airflow.utils.log.logging_mixin import LoggingMixin @@ -39,15 +39,15 @@ @attr.define -class DatasetLineageInfo: +class AssetLineageInfo: """ - Holds lineage information for a single dataset. + Holds lineage information for a single asset. - This class represents the lineage information for a single dataset, including the dataset itself, + This class represents the lineage information for a single asset, including the asset itself, the count of how many times it has been encountered, and the context in which it was encountered. """ - dataset: Dataset + asset: Asset count: int context: LineageContext @@ -58,133 +58,129 @@ class HookLineage: Holds lineage collected by HookLineageCollector. This class represents the lineage information collected by the `HookLineageCollector`. It stores - the input and output datasets, each with an associated count indicating how many times the dataset + the input and output assets, each with an associated count indicating how many times the asset has been encountered during the hook execution. """ - inputs: list[DatasetLineageInfo] = attr.ib(factory=list) - outputs: list[DatasetLineageInfo] = attr.ib(factory=list) + inputs: list[AssetLineageInfo] = attr.ib(factory=list) + outputs: list[AssetLineageInfo] = attr.ib(factory=list) class HookLineageCollector(LoggingMixin): """ HookLineageCollector is a base class for collecting hook lineage information. - It is used to collect the input and output datasets of a hook execution. + It is used to collect the input and output assets of a hook execution. """ def __init__(self, **kwargs): super().__init__(**kwargs) - # Dictionary to store input datasets, counted by unique key (dataset URI, MD5 hash of extra + # Dictionary to store input assets, counted by unique key (asset URI, MD5 hash of extra # dictionary, and LineageContext's unique identifier) - self._inputs: dict[str, tuple[Dataset, LineageContext]] = {} - self._outputs: dict[str, tuple[Dataset, LineageContext]] = {} + self._inputs: dict[str, tuple[Asset, LineageContext]] = {} + self._outputs: dict[str, tuple[Asset, LineageContext]] = {} self._input_counts: dict[str, int] = defaultdict(int) self._output_counts: dict[str, int] = defaultdict(int) - def _generate_key(self, dataset: Dataset, context: LineageContext) -> str: + def _generate_key(self, asset: Asset, context: LineageContext) -> str: """ - Generate a unique key for the given dataset and context. + Generate a unique key for the given asset and context. - This method creates a unique key by combining the dataset URI, the MD5 hash of the dataset's extra + This method creates a unique key by combining the asset URI, the MD5 hash of the asset's extra dictionary, and the LineageContext's unique identifier. This ensures that the generated key is - unique for each combination of dataset and context. + unique for each combination of asset and context. """ - extra_str = json.dumps(dataset.extra, sort_keys=True) + extra_str = json.dumps(asset.extra, sort_keys=True) extra_hash = hashlib.md5(extra_str.encode()).hexdigest() - return f"{dataset.uri}_{extra_hash}_{id(context)}" + return f"{asset.uri}_{extra_hash}_{id(context)}" - def create_dataset( - self, scheme: str | None, uri: str | None, dataset_kwargs: dict | None, dataset_extra: dict | None - ) -> Dataset | None: + def create_asset( + self, scheme: str | None, uri: str | None, asset_kwargs: dict | None, asset_extra: dict | None + ) -> Asset | None: """ - Create a Dataset instance using the provided parameters. + Create an asset instance using the provided parameters. - This method attempts to create a Dataset instance using the given parameters. - It first checks if a URI is provided and falls back to using the default dataset factory + This method attempts to create an asset instance using the given parameters. + It first checks if a URI is provided and falls back to using the default asset factory with the given URI if no other information is available. - If a scheme is provided but no URI, it attempts to find a dataset factory that matches + If a scheme is provided but no URI, it attempts to find an asset factory that matches the given scheme. If no such factory is found, it logs an error message and returns None. - If dataset_kwargs is provided, it is used to pass additional parameters to the Dataset - factory. The dataset_extra parameter is also passed to the factory as an ``extra`` parameter. + If asset_kwargs is provided, it is used to pass additional parameters to the asset + factory. The asset_extra parameter is also passed to the factory as an ``extra`` parameter. """ if uri: # Fallback to default factory using the provided URI - return Dataset(uri=uri, extra=dataset_extra) + return Asset(uri=uri, extra=asset_extra) if not scheme: self.log.debug( - "Missing required parameter: either 'uri' or 'scheme' must be provided to create a Dataset." + "Missing required parameter: either 'uri' or 'scheme' must be provided to create an asset." ) return None - dataset_factory = ProvidersManager().dataset_factories.get(scheme) - if not dataset_factory: - self.log.debug("Unsupported scheme: %s. Please provide a valid URI to create a Dataset.", scheme) + asset_factory = ProvidersManager().asset_factories.get(scheme) + if not asset_factory: + self.log.debug("Unsupported scheme: %s. Please provide a valid URI to create an asset.", scheme) return None - dataset_kwargs = dataset_kwargs or {} + asset_kwargs = asset_kwargs or {} try: - return dataset_factory(**dataset_kwargs, extra=dataset_extra) + return asset_factory(**asset_kwargs, extra=asset_extra) except Exception as e: - self.log.debug("Failed to create dataset. Skipping. Error: %s", e) + self.log.debug("Failed to create asset. Skipping. Error: %s", e) return None - def add_input_dataset( + def add_input_asset( self, context: LineageContext, scheme: str | None = None, uri: str | None = None, - dataset_kwargs: dict | None = None, - dataset_extra: dict | None = None, + asset_kwargs: dict | None = None, + asset_extra: dict | None = None, ): - """Add the input dataset and its corresponding hook execution context to the collector.""" - dataset = self.create_dataset( - scheme=scheme, uri=uri, dataset_kwargs=dataset_kwargs, dataset_extra=dataset_extra - ) - if dataset: - key = self._generate_key(dataset, context) + """Add the input asset and its corresponding hook execution context to the collector.""" + asset = self.create_asset(scheme=scheme, uri=uri, asset_kwargs=asset_kwargs, asset_extra=asset_extra) + if asset: + key = self._generate_key(asset, context) if key not in self._inputs: - self._inputs[key] = (dataset, context) + self._inputs[key] = (asset, context) self._input_counts[key] += 1 - def add_output_dataset( + def add_output_asset( self, context: LineageContext, scheme: str | None = None, uri: str | None = None, - dataset_kwargs: dict | None = None, - dataset_extra: dict | None = None, + asset_kwargs: dict | None = None, + asset_extra: dict | None = None, ): - """Add the output dataset and its corresponding hook execution context to the collector.""" - dataset = self.create_dataset( - scheme=scheme, uri=uri, dataset_kwargs=dataset_kwargs, dataset_extra=dataset_extra - ) - if dataset: - key = self._generate_key(dataset, context) + """Add the output asset and its corresponding hook execution context to the collector.""" + asset = self.create_asset(scheme=scheme, uri=uri, asset_kwargs=asset_kwargs, asset_extra=asset_extra) + if asset: + key = self._generate_key(asset, context) if key not in self._outputs: - self._outputs[key] = (dataset, context) + self._outputs[key] = (asset, context) self._output_counts[key] += 1 @property - def collected_datasets(self) -> HookLineage: + def collected_assets(self) -> HookLineage: """Get the collected hook lineage information.""" return HookLineage( [ - DatasetLineageInfo(dataset=dataset, count=self._input_counts[key], context=context) - for key, (dataset, context) in self._inputs.items() + AssetLineageInfo(asset=asset, count=self._input_counts[key], context=context) + for key, (asset, context) in self._inputs.items() ], [ - DatasetLineageInfo(dataset=dataset, count=self._output_counts[key], context=context) - for key, (dataset, context) in self._outputs.items() + AssetLineageInfo(asset=asset, count=self._output_counts[key], context=context) + for key, (asset, context) in self._outputs.items() ], ) @property def has_collected(self) -> bool: - """Check if any datasets have been collected.""" + """Check if any assets have been collected.""" return len(self._inputs) != 0 or len(self._outputs) != 0 @@ -195,14 +191,14 @@ class NoOpCollector(HookLineageCollector): It is used when you want to disable lineage collection. """ - def add_input_dataset(self, *_, **__): + def add_input_asset(self, *_, **__): pass - def add_output_dataset(self, *_, **__): + def add_output_asset(self, *_, **__): pass @property - def collected_datasets( + def collected_assets( self, ) -> HookLineage: self.log.warning( @@ -219,7 +215,7 @@ def __init__(self, **kwargs): def retrieve_hook_lineage(self) -> HookLineage: """Retrieve hook lineage from HookLineageCollector.""" - hook_lineage = self.lineage_collector.collected_datasets + hook_lineage = self.lineage_collector.collected_assets return hook_lineage diff --git a/airflow/listeners/listener.py b/airflow/listeners/listener.py index 57d0360487bbf..5e8fba55d4395 100644 --- a/airflow/listeners/listener.py +++ b/airflow/listeners/listener.py @@ -46,13 +46,13 @@ class ListenerManager: """Manage listener registration and provides hook property for calling them.""" def __init__(self): - from airflow.listeners.spec import dagrun, dataset, importerrors, lifecycle, taskinstance + from airflow.listeners.spec import asset, dagrun, importerrors, lifecycle, taskinstance self.pm = pluggy.PluginManager("airflow") self.pm.add_hookcall_monitoring(_before_hookcall, _after_hookcall) self.pm.add_hookspecs(lifecycle) self.pm.add_hookspecs(dagrun) - self.pm.add_hookspecs(dataset) + self.pm.add_hookspecs(asset) self.pm.add_hookspecs(taskinstance) self.pm.add_hookspecs(importerrors) diff --git a/airflow/listeners/spec/dataset.py b/airflow/listeners/spec/asset.py similarity index 76% rename from airflow/listeners/spec/dataset.py rename to airflow/listeners/spec/asset.py index eee1a10dd7d89..78b14c8b10aeb 100644 --- a/airflow/listeners/spec/dataset.py +++ b/airflow/listeners/spec/asset.py @@ -22,27 +22,21 @@ from pluggy import HookspecMarker if TYPE_CHECKING: - from airflow.datasets import Dataset, DatasetAlias + from airflow.assets import Asset, AssetAlias hookspec = HookspecMarker("airflow") @hookspec -def on_dataset_created( - dataset: Dataset, -): - """Execute when a new dataset is created.""" +def on_asset_created(asset: Asset): + """Execute when a new asset is created.""" @hookspec -def on_dataset_alias_created( - dataset_alias: DatasetAlias, -): +def on_asset_alias_created(dataset_alias: AssetAlias): """Execute when a new dataset alias is created.""" @hookspec -def on_dataset_changed( - dataset: Dataset, -): - """Execute when dataset change is registered.""" +def on_asset_changed(asset: Asset): + """Execute when asset change is registered.""" diff --git a/airflow/models/__init__.py b/airflow/models/__init__.py index 7bf23e1bbb7d1..375761bc20f52 100644 --- a/airflow/models/__init__.py +++ b/airflow/models/__init__.py @@ -58,9 +58,9 @@ def import_all_models(): for name in __lazy_imports: __getattr__(name) + import airflow.models.asset import airflow.models.backfill import airflow.models.dagwarning - import airflow.models.dataset import airflow.models.errors import airflow.models.serialized_dag import airflow.models.taskinstancehistory diff --git a/airflow/models/dataset.py b/airflow/models/asset.py similarity index 83% rename from airflow/models/dataset.py rename to airflow/models/asset.py index 489d6b68a6f15..b99aa86f2c889 100644 --- a/airflow/models/dataset.py +++ b/airflow/models/asset.py @@ -34,7 +34,7 @@ ) from sqlalchemy.orm import relationship -from airflow.datasets import Dataset, DatasetAlias +from airflow.assets import Asset, AssetAlias from airflow.models.base import Base, StringID from airflow.settings import json from airflow.utils import timezone @@ -83,11 +83,11 @@ ) -class DatasetAliasModel(Base): +class AssetAliasModel(Base): """ - A table to store dataset alias. + A table to store asset alias. - :param uri: a string that uniquely identifies the dataset alias + :param uri: a string that uniquely identifies the asset alias """ id = Column(Integer, primary_key=True, autoincrement=True) @@ -111,19 +111,19 @@ class DatasetAliasModel(Base): ) datasets = relationship( - "DatasetModel", + "AssetModel", secondary=alias_association_table, backref="aliases", ) dataset_events = relationship( - "DatasetEvent", + "AssetEvent", secondary=dataset_alias_dataset_event_assocation_table, back_populates="source_aliases", ) - consuming_dags = relationship("DagScheduleDatasetAliasReference", back_populates="dataset_alias") + consuming_dags = relationship("DagScheduleAssetAliasReference", back_populates="dataset_alias") @classmethod - def from_public(cls, obj: DatasetAlias) -> DatasetAliasModel: + def from_public(cls, obj: AssetAlias) -> AssetAliasModel: return cls(name=obj.name) def __repr__(self): @@ -133,20 +133,20 @@ def __hash__(self): return hash(self.name) def __eq__(self, other): - if isinstance(other, (self.__class__, DatasetAlias)): + if isinstance(other, (self.__class__, AssetAlias)): return self.name == other.name else: return NotImplemented - def to_public(self) -> DatasetAlias: - return DatasetAlias(name=self.name) + def to_public(self) -> AssetAlias: + return AssetAlias(name=self.name) -class DatasetModel(Base): +class AssetModel(Base): """ - A table to store datasets. + A table to store assets. - :param uri: a string that uniquely identifies the dataset + :param uri: a string that uniquely identifies the asset :param extra: JSON field for arbitrary extra info """ @@ -168,8 +168,8 @@ class DatasetModel(Base): updated_at = Column(UtcDateTime, default=timezone.utcnow, onupdate=timezone.utcnow, nullable=False) is_orphaned = Column(Boolean, default=False, nullable=False, server_default="0") - consuming_dags = relationship("DagScheduleDatasetReference", back_populates="dataset") - producing_tasks = relationship("TaskOutletDatasetReference", back_populates="dataset") + consuming_dags = relationship("DagScheduleAssetReference", back_populates="dataset") + producing_tasks = relationship("TaskOutletAssetReference", back_populates="dataset") __tablename__ = "dataset" __table_args__ = ( @@ -178,7 +178,7 @@ class DatasetModel(Base): ) @classmethod - def from_public(cls, obj: Dataset) -> DatasetModel: + def from_public(cls, obj: Asset) -> AssetModel: return cls(uri=obj.uri, extra=obj.extra) def __init__(self, uri: str, **kwargs): @@ -192,7 +192,7 @@ def __init__(self, uri: str, **kwargs): super().__init__(uri=uri, **kwargs) def __eq__(self, other): - if isinstance(other, (self.__class__, Dataset)): + if isinstance(other, (self.__class__, Asset)): return self.uri == other.uri else: return NotImplemented @@ -203,19 +203,19 @@ def __hash__(self): def __repr__(self): return f"{self.__class__.__name__}(uri={self.uri!r}, extra={self.extra!r})" - def to_public(self) -> Dataset: - return Dataset(uri=self.uri, extra=self.extra) + def to_public(self) -> Asset: + return Asset(uri=self.uri, extra=self.extra) -class DagScheduleDatasetAliasReference(Base): - """References from a DAG to a dataset alias of which it is a consumer.""" +class DagScheduleAssetAliasReference(Base): + """References from a DAG to an asset alias of which it is a consumer.""" alias_id = Column(Integer, primary_key=True, nullable=False) dag_id = Column(StringID(), primary_key=True, nullable=False) created_at = Column(UtcDateTime, default=timezone.utcnow, nullable=False) updated_at = Column(UtcDateTime, default=timezone.utcnow, onupdate=timezone.utcnow, nullable=False) - dataset_alias = relationship("DatasetAliasModel", back_populates="consuming_dags") + dataset_alias = relationship("AssetAliasModel", back_populates="consuming_dags") dag = relationship("DagModel", back_populates="schedule_dataset_alias_references") __tablename__ = "dag_schedule_dataset_alias_reference" @@ -251,22 +251,22 @@ def __repr__(self): return f"{self.__class__.__name__}({', '.join(args)})" -class DagScheduleDatasetReference(Base): - """References from a DAG to a dataset of which it is a consumer.""" +class DagScheduleAssetReference(Base): + """References from a DAG to an asset of which it is a consumer.""" dataset_id = Column(Integer, primary_key=True, nullable=False) dag_id = Column(StringID(), primary_key=True, nullable=False) created_at = Column(UtcDateTime, default=timezone.utcnow, nullable=False) updated_at = Column(UtcDateTime, default=timezone.utcnow, onupdate=timezone.utcnow, nullable=False) - dataset = relationship("DatasetModel", back_populates="consuming_dags") + dataset = relationship("AssetModel", back_populates="consuming_dags") dag = relationship("DagModel", back_populates="schedule_dataset_references") queue_records = relationship( - "DatasetDagRunQueue", + "AssetDagRunQueue", primaryjoin="""and_( - DagScheduleDatasetReference.dataset_id == foreign(DatasetDagRunQueue.dataset_id), - DagScheduleDatasetReference.dag_id == foreign(DatasetDagRunQueue.target_dag_id), + DagScheduleAssetReference.dataset_id == foreign(AssetDagRunQueue.dataset_id), + DagScheduleAssetReference.dag_id == foreign(AssetDagRunQueue.target_dag_id), )""", cascade="all, delete, delete-orphan", ) @@ -305,8 +305,8 @@ def __repr__(self): return f"{self.__class__.__name__}({', '.join(args)})" -class TaskOutletDatasetReference(Base): - """References from a task to a dataset that it updates / produces.""" +class TaskOutletAssetReference(Base): + """References from a task to an asset that it updates / produces.""" dataset_id = Column(Integer, primary_key=True, nullable=False) dag_id = Column(StringID(), primary_key=True, nullable=False) @@ -314,7 +314,7 @@ class TaskOutletDatasetReference(Base): created_at = Column(UtcDateTime, default=timezone.utcnow, nullable=False) updated_at = Column(UtcDateTime, default=timezone.utcnow, onupdate=timezone.utcnow, nullable=False) - dataset = relationship("DatasetModel", back_populates="producing_tasks") + dataset = relationship("AssetModel", back_populates="producing_tasks") __tablename__ = "task_outlet_dataset_reference" __table_args__ = ( @@ -354,13 +354,13 @@ def __repr__(self): return f"{self.__class__.__name__}({', '.join(args)})" -class DatasetDagRunQueue(Base): - """Model for storing dataset events that need processing.""" +class AssetDagRunQueue(Base): + """Model for storing asset events that need processing.""" dataset_id = Column(Integer, primary_key=True, nullable=False) target_dag_id = Column(StringID(), primary_key=True, nullable=False) created_at = Column(UtcDateTime, default=timezone.utcnow, nullable=False) - dataset = relationship("DatasetModel", viewonly=True) + dataset = relationship("AssetModel", viewonly=True) __tablename__ = "dataset_dag_run_queue" __table_args__ = ( PrimaryKeyConstraint(dataset_id, target_dag_id, name="datasetdagrunqueue_pkey"), @@ -405,19 +405,19 @@ def __repr__(self): ) -class DatasetEvent(Base): +class AssetEvent(Base): """ - A table to store datasets events. + A table to store assets events. - :param dataset_id: reference to DatasetModel record + :param dataset_id: reference to AssetModel record :param extra: JSON field for arbitrary extra info - :param source_task_id: the task_id of the TI which updated the dataset - :param source_dag_id: the dag_id of the TI which updated the dataset - :param source_run_id: the run_id of the TI which updated the dataset - :param source_map_index: the map_index of the TI which updated the dataset + :param source_task_id: the task_id of the TI which updated the asset + :param source_dag_id: the dag_id of the TI which updated the asset + :param source_run_id: the run_id of the TI which updated the asset + :param source_map_index: the map_index of the TI which updated the asset :param timestamp: the time the event was logged - We use relationships instead of foreign keys so that dataset events are not deleted even + We use relationships instead of foreign keys so that asset events are not deleted even if the foreign key object is. """ @@ -443,7 +443,7 @@ class DatasetEvent(Base): ) source_aliases = relationship( - "DatasetAliasModel", + "AssetAliasModel", secondary=dataset_alias_dataset_event_assocation_table, back_populates="dataset_events", ) @@ -451,10 +451,10 @@ class DatasetEvent(Base): source_task_instance = relationship( "TaskInstance", primaryjoin="""and_( - DatasetEvent.source_dag_id == foreign(TaskInstance.dag_id), - DatasetEvent.source_run_id == foreign(TaskInstance.run_id), - DatasetEvent.source_task_id == foreign(TaskInstance.task_id), - DatasetEvent.source_map_index == foreign(TaskInstance.map_index), + AssetEvent.source_dag_id == foreign(TaskInstance.dag_id), + AssetEvent.source_run_id == foreign(TaskInstance.run_id), + AssetEvent.source_task_id == foreign(TaskInstance.task_id), + AssetEvent.source_map_index == foreign(TaskInstance.map_index), )""", viewonly=True, lazy="select", @@ -463,16 +463,16 @@ class DatasetEvent(Base): source_dag_run = relationship( "DagRun", primaryjoin="""and_( - DatasetEvent.source_dag_id == foreign(DagRun.dag_id), - DatasetEvent.source_run_id == foreign(DagRun.run_id), + AssetEvent.source_dag_id == foreign(DagRun.dag_id), + AssetEvent.source_run_id == foreign(DagRun.run_id), )""", viewonly=True, lazy="select", uselist=False, ) dataset = relationship( - DatasetModel, - primaryjoin="DatasetEvent.dataset_id == foreign(DatasetModel.id)", + AssetModel, + primaryjoin="AssetEvent.dataset_id == foreign(AssetModel.id)", viewonly=True, lazy="select", uselist=False, diff --git a/airflow/models/dag.py b/airflow/models/dag.py index 91f8aec7302cb..0632819952ae4 100644 --- a/airflow/models/dag.py +++ b/airflow/models/dag.py @@ -81,8 +81,8 @@ import airflow.templates from airflow import settings, utils from airflow.api_internal.internal_api_call import internal_api_call +from airflow.assets import Asset, AssetAlias, AssetAll, BaseAsset from airflow.configuration import conf as airflow_conf, secrets_backend_list -from airflow.datasets import BaseDataset, Dataset, DatasetAlias, DatasetAll from airflow.exceptions import ( AirflowException, DuplicateTaskIdFound, @@ -96,12 +96,15 @@ from airflow.executors.executor_loader import ExecutorLoader from airflow.jobs.job import run_job from airflow.models.abstractoperator import AbstractOperator, TaskStateChangeCallback +from airflow.models.asset import ( + AssetDagRunQueue, + AssetModel, +) from airflow.models.base import Base, StringID from airflow.models.baseoperator import BaseOperator from airflow.models.dagcode import DagCode from airflow.models.dagpickle import DagPickle from airflow.models.dagrun import RUN_ID_REGEX, DagRun -from airflow.models.dataset import DatasetDagRunQueue from airflow.models.param import DagParam, ParamsDict from airflow.models.taskinstance import ( Context, @@ -118,8 +121,8 @@ from airflow.timetables.base import DagRunInfo, DataInterval, TimeRestriction, Timetable from airflow.timetables.interval import CronDataIntervalTimetable, DeltaDataIntervalTimetable from airflow.timetables.simple import ( + AssetTriggeredTimetable, ContinuousTimetable, - DatasetTriggeredTimetable, NullTimetable, OnceTimetable, ) @@ -163,8 +166,8 @@ ScheduleArg = Union[ ScheduleInterval, Timetable, - BaseDataset, - Collection[Union["Dataset", "DatasetAlias"]], + BaseAsset, + Collection[Union["Asset", "AssetAlias"]], ] @@ -240,16 +243,16 @@ def get_last_dagrun(dag_id, session, include_externally_triggered=False): return session.scalar(query.limit(1)) -def get_dataset_triggered_next_run_info( +def get_asset_triggered_next_run_info( dag_ids: list[str], *, session: Session ) -> dict[str, dict[str, int | str]]: """ Get next run info for a list of dag_ids. - Given a list of dag_ids, get string representing how close any that are dataset triggered are - their next run, e.g. "1 of 2 datasets updated". + Given a list of dag_ids, get string representing how close any that are asset triggered are + their next run, e.g. "1 of 2 assets updated". """ - from airflow.models.dataset import DagScheduleDatasetReference, DatasetDagRunQueue as DDRQ, DatasetModel + from airflow.models.asset import AssetDagRunQueue as ADRQ, DagScheduleAssetReference return { x.dag_id: { @@ -259,24 +262,24 @@ def get_dataset_triggered_next_run_info( } for x in session.execute( select( - DagScheduleDatasetReference.dag_id, + DagScheduleAssetReference.dag_id, # This is a dirty hack to workaround group by requiring an aggregate, - # since grouping by dataset is not what we want to do here...but it works - case((func.count() == 1, func.max(DatasetModel.uri)), else_="").label("uri"), + # since grouping by asset is not what we want to do here...but it works + case((func.count() == 1, func.max(AssetModel.uri)), else_="").label("uri"), func.count().label("total"), - func.sum(case((DDRQ.target_dag_id.is_not(None), 1), else_=0)).label("ready"), + func.sum(case((ADRQ.target_dag_id.is_not(None), 1), else_=0)).label("ready"), ) .join( - DDRQ, + ADRQ, and_( - DDRQ.dataset_id == DagScheduleDatasetReference.dataset_id, - DDRQ.target_dag_id == DagScheduleDatasetReference.dag_id, + ADRQ.dataset_id == DagScheduleAssetReference.dataset_id, + ADRQ.target_dag_id == DagScheduleAssetReference.dag_id, ), isouter=True, ) - .join(DatasetModel, DatasetModel.id == DagScheduleDatasetReference.dataset_id) - .group_by(DagScheduleDatasetReference.dag_id) - .where(DagScheduleDatasetReference.dag_id.in_(dag_ids)) + .join(AssetModel, AssetModel.id == DagScheduleAssetReference.dataset_id) + .group_by(DagScheduleAssetReference.dag_id) + .where(DagScheduleAssetReference.dag_id.in_(dag_ids)) ).all() } @@ -386,7 +389,7 @@ class DAG(LoggingMixin): :param description: The description for the DAG to e.g. be shown on the webserver :param schedule: If provided, this defines the rules according to which DAG runs are scheduled. Possible values include a cron expression string, - timedelta object, Timetable, or list of Dataset objects. + timedelta object, Timetable, or list of Asset objects. See also :doc:`/howto/timetable`. :param start_date: The timestamp from which the scheduler will attempt to backfill. If this is not provided, backfilling must be done @@ -595,12 +598,12 @@ def __init__( if isinstance(schedule, Timetable): self.timetable = schedule - elif isinstance(schedule, BaseDataset): - self.timetable = DatasetTriggeredTimetable(schedule) + elif isinstance(schedule, BaseAsset): + self.timetable = AssetTriggeredTimetable(schedule) elif isinstance(schedule, Collection) and not isinstance(schedule, str): - if not all(isinstance(x, (Dataset, DatasetAlias)) for x in schedule): - raise ValueError("All elements in 'schedule' should be datasets or dataset aliases") - self.timetable = DatasetTriggeredTimetable(DatasetAll(*schedule)) + if not all(isinstance(x, (Asset, AssetAlias)) for x in schedule): + raise ValueError("All elements in 'schedule' should be assets or asset aliases") + self.timetable = AssetTriggeredTimetable(AssetAll(*schedule)) else: self.timetable = create_timetable(schedule, self.timezone) @@ -873,7 +876,7 @@ def infer_automated_data_interval(self, logical_date: datetime) -> DataInterval: :meta private: """ timetable_type = type(self.timetable) - if issubclass(timetable_type, (NullTimetable, OnceTimetable, DatasetTriggeredTimetable)): + if issubclass(timetable_type, (NullTimetable, OnceTimetable, AssetTriggeredTimetable)): return DataInterval.exact(timezone.coerce_datetime(logical_date)) start = timezone.coerce_datetime(logical_date) if issubclass(timetable_type, CronDataIntervalTimetable): @@ -2649,7 +2652,7 @@ def bulk_write_to_db( if not dags: return - from airflow.dag_processing.collection import DagModelOperation, DatasetModelOperation + from airflow.dag_processing.collection import AssetModelOperation, DagModelOperation log.info("Sync %s DAGs", len(dags)) dag_op = DagModelOperation({dag.dag_id: dag for dag in dags}) @@ -2658,15 +2661,15 @@ def bulk_write_to_db( dag_op.update_dags(orm_dags, processor_subdir=processor_subdir, session=session) DagCode.bulk_sync_to_db((dag.fileloc for dag in dags), session=session) - dataset_op = DatasetModelOperation.collect(dag_op.dags) + asset_op = AssetModelOperation.collect(dag_op.dags) - orm_datasets = dataset_op.add_datasets(session=session) - orm_dataset_aliases = dataset_op.add_dataset_aliases(session=session) + orm_assets = asset_op.add_assets(session=session) + orm_asset_aliases = asset_op.add_asset_aliases(session=session) session.flush() # This populates id so we can create fks in later calls. - dataset_op.add_dag_dataset_references(orm_dags, orm_datasets, session=session) - dataset_op.add_dag_dataset_alias_references(orm_dags, orm_dataset_aliases, session=session) - dataset_op.add_task_dataset_references(orm_dags, orm_datasets, session=session) + asset_op.add_dag_asset_references(orm_dags, orm_assets, session=session) + asset_op.add_dag_asset_alias_references(orm_dags, orm_asset_aliases, session=session) + asset_op.add_task_asset_references(orm_dags, orm_assets, session=session) session.flush() @provide_session @@ -2963,18 +2966,18 @@ class DagModel(Base): __table_args__ = (Index("idx_next_dagrun_create_after", next_dagrun_create_after, unique=False),) schedule_dataset_references = relationship( - "DagScheduleDatasetReference", + "DagScheduleAssetReference", back_populates="dag", cascade="all, delete, delete-orphan", ) schedule_dataset_alias_references = relationship( - "DagScheduleDatasetAliasReference", + "DagScheduleAssetAliasReference", back_populates="dag", cascade="all, delete, delete-orphan", ) schedule_datasets = association_proxy("schedule_dataset_references", "dataset") task_outlet_dataset_references = relationship( - "TaskOutletDatasetReference", + "TaskOutletAssetReference", cascade="all, delete, delete-orphan", ) NUM_DAGS_PER_DAGRUN_QUERY = airflow_conf.getint( @@ -3155,7 +3158,7 @@ def dags_needing_dagruns(cls, session: Session) -> tuple[Query, dict[str, tuple[ """ from airflow.models.serialized_dag import SerializedDagModel - def dag_ready(dag_id: str, cond: BaseDataset, statuses: dict) -> bool | None: + def dag_ready(dag_id: str, cond: BaseAsset, statuses: dict) -> bool | None: # if dag was serialized before 2.9 and we *just* upgraded, # we may be dealing with old version. In that case, # just wait for the dag to be reserialized. @@ -3165,8 +3168,8 @@ def dag_ready(dag_id: str, cond: BaseDataset, statuses: dict) -> bool | None: log.warning("dag '%s' has old serialization; skipping DAG run creation.", dag_id) return None - # this loads all the DDRQ records.... may need to limit num dags - all_records = session.scalars(select(DatasetDagRunQueue)).all() + # this loads all the ADRQ records.... may need to limit num dags + all_records = session.scalars(select(AssetDagRunQueue)).all() by_dag = defaultdict(list) for r in all_records: by_dag[r.target_dag_id].append(r) @@ -3181,7 +3184,7 @@ def dag_ready(dag_id: str, cond: BaseDataset, statuses: dict) -> bool | None: dag_id = ser_dag.dag_id statuses = dag_statuses[dag_id] - if not dag_ready(dag_id, cond=ser_dag.dag.timetable.dataset_condition, statuses=statuses): + if not dag_ready(dag_id, cond=ser_dag.dag.timetable.asset_condition, statuses=statuses): del by_dag[dag_id] del dag_statuses[dag_id] del dag_statuses @@ -3265,13 +3268,13 @@ def calculate_dagrun_date_fields( ) @provide_session - def get_dataset_triggered_next_run_info(self, *, session=NEW_SESSION) -> dict[str, int | str] | None: + def get_asset_triggered_next_run_info(self, *, session=NEW_SESSION) -> dict[str, int | str] | None: if self.dataset_expression is None: return None - # When a dataset alias does not resolve into datasets, get_dataset_triggered_next_run_info returns - # an empty dict as there's no dataset info to get. This method should thus return None. - return get_dataset_triggered_next_run_info([self.dag_id], session=session).get(self.dag_id, None) + # When an asset alias does not resolve into assets, get_asset_triggered_next_run_info returns + # an empty dict as there's no asset info to get. This method should thus return None. + return get_asset_triggered_next_run_info([self.dag_id], session=session).get(self.dag_id, None) # NOTE: Please keep the list of arguments in sync with DAG.__init__. diff --git a/airflow/models/taskinstance.py b/airflow/models/taskinstance.py index c17acdd2b7212..b19e65486307d 100644 --- a/airflow/models/taskinstance.py +++ b/airflow/models/taskinstance.py @@ -67,10 +67,10 @@ from airflow import settings from airflow.api_internal.internal_api_call import InternalApiConfig, internal_api_call +from airflow.assets import Asset, AssetAlias +from airflow.assets.manager import asset_manager from airflow.compat.functools import cache from airflow.configuration import conf -from airflow.datasets import Dataset, DatasetAlias -from airflow.datasets.manager import dataset_manager from airflow.exceptions import ( AirflowException, AirflowFailException, @@ -87,9 +87,9 @@ XComForMappingNotPushed, ) from airflow.listeners.listener import get_listener_manager +from airflow.models.asset import AssetEvent, AssetModel from airflow.models.base import Base, StringID, TaskInstanceDependencies, _sentinel from airflow.models.dagbag import DagBag -from airflow.models.dataset import DatasetModel from airflow.models.log import Log from airflow.models.param import process_params from airflow.models.renderedtifields import get_serialized_template_fields @@ -154,13 +154,13 @@ from sqlalchemy.sql.expression import ColumnOperators from airflow.models.abstractoperator import TaskStateChangeCallback + from airflow.models.asset import AssetEvent from airflow.models.baseoperator import BaseOperator from airflow.models.dag import DAG, DagModel from airflow.models.dagrun import DagRun - from airflow.models.dataset import DatasetEvent from airflow.models.operator import Operator + from airflow.serialization.pydantic.asset import AssetEventPydantic from airflow.serialization.pydantic.dag import DagModelPydantic - from airflow.serialization.pydantic.dataset import DatasetEventPydantic from airflow.serialization.pydantic.taskinstance import TaskInstancePydantic from airflow.timetables.base import DataInterval from airflow.typing_compat import Literal, TypeGuard @@ -366,7 +366,7 @@ def _run_raw_task( if not test_mode: _add_log(event=ti.state, task_instance=ti, session=session) if ti.state == TaskInstanceState.SUCCESS: - ti._register_dataset_changes(events=context["outlet_events"], session=session) + ti._register_asset_changes(events=context["outlet_events"], session=session) TaskInstance.save_to_db(ti=ti, session=session) if ti.state == TaskInstanceState.SUCCESS: @@ -1077,7 +1077,7 @@ def get_prev_ds_nodash() -> str | None: return None return prev_ds.replace("-", "") - def get_triggering_events() -> dict[str, list[DatasetEvent | DatasetEventPydantic]]: + def get_triggering_events() -> dict[str, list[AssetEvent | AssetEventPydantic]]: if TYPE_CHECKING: assert session is not None @@ -1087,9 +1087,9 @@ def get_triggering_events() -> dict[str, list[DatasetEvent | DatasetEventPydanti nonlocal dag_run if dag_run not in session: dag_run = session.merge(dag_run, load=False) - dataset_events = dag_run.consumed_dataset_events - triggering_events: dict[str, list[DatasetEvent | DatasetEventPydantic]] = defaultdict(list) - for event in dataset_events: + asset_events = dag_run.consumed_dataset_events + triggering_events: dict[str, list[AssetEvent | AssetEventPydantic]] = defaultdict(list) + for event in asset_events: if event.dataset: triggering_events[event.dataset.uri].append(event) @@ -1144,7 +1144,7 @@ def get_triggering_events() -> dict[str, list[DatasetEvent | DatasetEventPydanti "ti": task_instance, "tomorrow_ds": get_tomorrow_ds(), "tomorrow_ds_nodash": get_tomorrow_ds_nodash(), - "triggering_dataset_events": lazy_object_proxy.Proxy(get_triggering_events), + "triggering_asset_events": lazy_object_proxy.Proxy(get_triggering_events), "ts": ts, "ts_nodash": ts_nodash, "ts_nodash_with_tz": ts_nodash_with_tz, @@ -2886,56 +2886,56 @@ def _run_raw_task( session=session, ) - def _register_dataset_changes(self, *, events: OutletEventAccessors, session: Session) -> None: + def _register_asset_changes(self, *, events: OutletEventAccessors, session: Session) -> None: if TYPE_CHECKING: assert self.task - # One task only triggers one dataset event for each dataset with the same extra. - # This tuple[dataset uri, extra] to sets alias names mapping is used to find whether - # there're datasets with same uri but different extra that we need to emit more than one dataset events. - dataset_alias_names: dict[tuple[str, frozenset], set[str]] = defaultdict(set) + # One task only triggers one asset event for each asset with the same extra. + # This tuple[asset uri, extra] to sets alias names mapping is used to find whether + # there're assets with same uri but different extra that we need to emit more than one asset events. + asset_alias_names: dict[tuple[str, frozenset], set[str]] = defaultdict(set) for obj in self.task.outlets or []: self.log.debug("outlet obj %s", obj) - # Lineage can have other types of objects besides datasets - if isinstance(obj, Dataset): - dataset_manager.register_dataset_change( + # Lineage can have other types of objects besides assets + if isinstance(obj, Asset): + asset_manager.register_asset_change( task_instance=self, - dataset=obj, + asset=obj, extra=events[obj].extra, session=session, ) - elif isinstance(obj, DatasetAlias): - for dataset_alias_event in events[obj].dataset_alias_events: - dataset_alias_name = dataset_alias_event["source_alias_name"] - dataset_uri = dataset_alias_event["dest_dataset_uri"] - frozen_extra = frozenset(dataset_alias_event["extra"].items()) - dataset_alias_names[(dataset_uri, frozen_extra)].add(dataset_alias_name) - - dataset_models: dict[str, DatasetModel] = { + elif isinstance(obj, AssetAlias): + for asset_alias_event in events[obj].asset_alias_events: + asset_alias_name = asset_alias_event["source_alias_name"] + asset_uri = asset_alias_event["dest_asset_uri"] + frozen_extra = frozenset(asset_alias_event["extra"].items()) + asset_alias_names[(asset_uri, frozen_extra)].add(asset_alias_name) + + dataset_models: dict[str, AssetModel] = { dataset_obj.uri: dataset_obj for dataset_obj in session.scalars( - select(DatasetModel).where(DatasetModel.uri.in_(uri for uri, _ in dataset_alias_names)) + select(AssetModel).where(AssetModel.uri.in_(uri for uri, _ in asset_alias_names)) ) } - if missing_datasets := [Dataset(uri=u) for u, _ in dataset_alias_names if u not in dataset_models]: + if missing_datasets := [Asset(uri=u) for u, _ in asset_alias_names if u not in dataset_models]: dataset_models.update( (dataset_obj.uri, dataset_obj) - for dataset_obj in dataset_manager.create_datasets(missing_datasets, session=session) + for dataset_obj in asset_manager.create_assets(missing_datasets, session=session) ) self.log.warning("Created new datasets for alias reference: %s", missing_datasets) session.flush() # Needed because we need the id for fk. - for (uri, extra_items), alias_names in dataset_alias_names.items(): - dataset_obj = dataset_models[uri] + for (uri, extra_items), alias_names in asset_alias_names.items(): + asset_obj = dataset_models[uri] self.log.info( 'Creating event for %r through aliases "%s"', - dataset_obj, + asset_obj, ", ".join(alias_names), ) - dataset_manager.register_dataset_change( + asset_manager.register_asset_change( task_instance=self, - dataset=dataset_obj.to_public(), - aliases=[DatasetAlias(name) for name in alias_names], + asset=asset_obj, + aliases=[AssetAlias(name) for name in alias_names], extra=dict(extra_items), session=session, source_alias_names=alias_names, diff --git a/airflow/operators/python.py b/airflow/operators/python.py index 25c0b8ca68632..a4788caedf438 100644 --- a/airflow/operators/python.py +++ b/airflow/operators/python.py @@ -234,7 +234,7 @@ def __init__( def execute(self, context: Context) -> Any: context_merge(context, self.op_kwargs, templates_dict=self.templates_dict) self.op_kwargs = self.determine_kwargs(context) - self._dataset_events = context_get_outlet_events(context) + self._asset_events = context_get_outlet_events(context) return_value = self.execute_callable() if self.show_return_value_in_logs: @@ -253,7 +253,7 @@ def execute_callable(self) -> Any: :return: the return value of the call. """ - runner = ExecutionCallableRunner(self.python_callable, self._dataset_events, logger=self.log) + runner = ExecutionCallableRunner(self.python_callable, self._asset_events, logger=self.log) return runner.run(*self.op_args, **self.op_kwargs) @@ -424,7 +424,7 @@ class _BasePythonVirtualenvOperator(PythonOperator, metaclass=ABCMeta): "dag_run", "task", "params", - "triggering_dataset_events", + "triggering_asset_events", } def __init__( diff --git a/airflow/provider.yaml.schema.json b/airflow/provider.yaml.schema.json index 8f11833ee1c8f..35e266c310ac3 100644 --- a/airflow/provider.yaml.schema.json +++ b/airflow/provider.yaml.schema.json @@ -196,9 +196,37 @@ "type": "string" } }, + "asset-uris": { + "type": "array", + "description": "Asset URI formats", + "items": { + "type": "object", + "properties": { + "schemes": { + "type": "array", + "description": "List of supported URI schemes", + "items": { + "type": "string" + } + }, + "handler": { + "type": ["string", "null"], + "description": "Normalization function for specified URI schemes. Import path to a callable taking and returning a SplitResult. 'null' specifies a no-op." + }, + "factory": { + "type": ["string", "null"], + "description": "Dataset factory for specified URI. Creates AIP-60 compliant Dataset." + }, + "to_openlineage_converter": { + "type": ["string", "null"], + "description": "OpenLineage converter function for specified URI schemes. Import path to a callable accepting a Dataset and LineageContext and returning OpenLineage dataset." + } + } + } + }, "dataset-uris": { "type": "array", - "description": "Dataset URI formats", + "description": "Dataset URI formats (will be removed in Airflow 3.0)", "items": { "type": "object", "properties": { diff --git a/airflow/providers/amazon/aws/datasets/__init__.py b/airflow/providers/amazon/aws/assets/__init__.py similarity index 100% rename from airflow/providers/amazon/aws/datasets/__init__.py rename to airflow/providers/amazon/aws/assets/__init__.py diff --git a/airflow/providers/amazon/aws/datasets/s3.py b/airflow/providers/amazon/aws/assets/s3.py similarity index 73% rename from airflow/providers/amazon/aws/datasets/s3.py rename to airflow/providers/amazon/aws/assets/s3.py index c42ec2bb1cc03..378e7e977ab56 100644 --- a/airflow/providers/amazon/aws/datasets/s3.py +++ b/airflow/providers/amazon/aws/assets/s3.py @@ -18,17 +18,21 @@ from typing import TYPE_CHECKING -from airflow.datasets import Dataset from airflow.providers.amazon.aws.hooks.s3 import S3Hook +try: + from airflow.assets import Asset +except ModuleNotFoundError: + from airflow.datasets import Dataset as Asset # type: ignore[no-redef] + if TYPE_CHECKING: from urllib.parse import SplitResult from airflow.providers.common.compat.openlineage.facet import Dataset as OpenLineageDataset -def create_dataset(*, bucket: str, key: str, extra=None) -> Dataset: - return Dataset(uri=f"s3://{bucket}/{key}", extra=extra) +def create_asset(*, bucket: str, key: str, extra=None) -> Asset: + return Asset(uri=f"s3://{bucket}/{key}", extra=extra) def sanitize_uri(uri: SplitResult) -> SplitResult: @@ -37,9 +41,9 @@ def sanitize_uri(uri: SplitResult) -> SplitResult: return uri -def convert_dataset_to_openlineage(dataset: Dataset, lineage_context) -> OpenLineageDataset: - """Translate Dataset with valid AIP-60 uri to OpenLineage with assistance from the hook.""" +def convert_asset_to_openlineage(asset: Asset, lineage_context) -> OpenLineageDataset: + """Translate Asset with valid AIP-60 uri to OpenLineage with assistance from the hook.""" from airflow.providers.common.compat.openlineage.facet import Dataset as OpenLineageDataset - bucket, key = S3Hook.parse_s3_url(dataset.uri) + bucket, key = S3Hook.parse_s3_url(asset.uri) return OpenLineageDataset(namespace=f"s3://{bucket}", name=key if key else "/") diff --git a/airflow/providers/amazon/aws/auth_manager/avp/entities.py b/airflow/providers/amazon/aws/auth_manager/avp/entities.py index 8c2e8855b877d..4db9aed340208 100644 --- a/airflow/providers/amazon/aws/auth_manager/avp/entities.py +++ b/airflow/providers/amazon/aws/auth_manager/avp/entities.py @@ -33,11 +33,11 @@ class AvpEntities(Enum): USER = "User" # Resource types + ASSET = "Asset" CONFIGURATION = "Configuration" CONNECTION = "Connection" CUSTOM = "Custom" DAG = "Dag" - DATASET = "Dataset" MENU = "Menu" POOL = "Pool" VARIABLE = "Variable" diff --git a/airflow/providers/amazon/aws/auth_manager/aws_auth_manager.py b/airflow/providers/amazon/aws/auth_manager/aws_auth_manager.py index c8693da3382e0..face67c38fb57 100644 --- a/airflow/providers/amazon/aws/auth_manager/aws_auth_manager.py +++ b/airflow/providers/amazon/aws/auth_manager/aws_auth_manager.py @@ -63,10 +63,7 @@ IsAuthorizedPoolRequest, IsAuthorizedVariableRequest, ) - from airflow.auth.managers.models.resource_details import ( - ConfigurationDetails, - DatasetDetails, - ) + from airflow.auth.managers.models.resource_details import AssetDetails, ConfigurationDetails from airflow.providers.amazon.aws.auth_manager.user import AwsAuthManagerUser from airflow.www.extensions.init_appbuilder import AirflowAppBuilder @@ -161,15 +158,12 @@ def is_authorized_dag( context=context, ) - def is_authorized_dataset( - self, *, method: ResourceMethod, details: DatasetDetails | None = None, user: BaseUser | None = None + def is_authorized_asset( + self, *, method: ResourceMethod, details: AssetDetails | None = None, user: BaseUser | None = None ) -> bool: - dataset_uri = details.uri if details else None + asset_uri = details.uri if details else None return self.avp_facade.is_authorized( - method=method, - entity_type=AvpEntities.DATASET, - user=user or self.get_user(), - entity_id=dataset_uri, + method=method, entity_type=AvpEntities.ASSET, user=user or self.get_user(), entity_id=asset_uri ) def is_authorized_pool( diff --git a/airflow/providers/amazon/aws/hooks/s3.py b/airflow/providers/amazon/aws/hooks/s3.py index b609259f846ba..6efb5953b8eee 100644 --- a/airflow/providers/amazon/aws/hooks/s3.py +++ b/airflow/providers/amazon/aws/hooks/s3.py @@ -40,8 +40,6 @@ from urllib.parse import urlsplit from uuid import uuid4 -from airflow.providers.common.compat.lineage.hook import get_hook_lineage_collector - if TYPE_CHECKING: from mypy_boto3_s3.service_resource import Bucket as S3Bucket, Object as S3ResourceObject @@ -50,6 +48,8 @@ with suppress(ImportError): from aiobotocore.client import AioBaseClient +from importlib.util import find_spec + from asgiref.sync import sync_to_async from boto3.s3.transfer import S3Transfer, TransferConfig from botocore.exceptions import ClientError @@ -60,6 +60,13 @@ from airflow.providers.amazon.aws.utils.tags import format_tags from airflow.utils.helpers import chunks +if find_spec("airflow.assets"): + from airflow.lineage.hook import get_hook_lineage_collector +else: + # TODO: import from common.compat directly after common.compat providers with + # asset_compat_lineage_collector released + from airflow.providers.amazon.aws.utils.asset_compat_lineage_collector import get_hook_lineage_collector + logger = logging.getLogger(__name__) @@ -1103,11 +1110,11 @@ def load_file( client = self.get_conn() client.upload_file(filename, bucket_name, key, ExtraArgs=extra_args, Config=self.transfer_config) - get_hook_lineage_collector().add_input_dataset( - context=self, scheme="file", dataset_kwargs={"path": filename} + get_hook_lineage_collector().add_input_asset( + context=self, scheme="file", asset_kwargs={"path": filename} ) - get_hook_lineage_collector().add_output_dataset( - context=self, scheme="s3", dataset_kwargs={"bucket": bucket_name, "key": key} + get_hook_lineage_collector().add_output_asset( + context=self, scheme="s3", asset_kwargs={"bucket": bucket_name, "key": key} ) @unify_bucket_name_and_key @@ -1250,8 +1257,8 @@ def _upload_file_obj( Config=self.transfer_config, ) # No input because file_obj can be anything - handle in calling function if possible - get_hook_lineage_collector().add_output_dataset( - context=self, scheme="s3", dataset_kwargs={"bucket": bucket_name, "key": key} + get_hook_lineage_collector().add_output_asset( + context=self, scheme="s3", asset_kwargs={"bucket": bucket_name, "key": key} ) def copy_object( @@ -1308,11 +1315,11 @@ def copy_object( response = self.get_conn().copy_object( Bucket=dest_bucket_name, Key=dest_bucket_key, CopySource=copy_source, **kwargs ) - get_hook_lineage_collector().add_input_dataset( - context=self, scheme="s3", dataset_kwargs={"bucket": source_bucket_name, "key": source_bucket_key} + get_hook_lineage_collector().add_input_asset( + context=self, scheme="s3", asset_kwargs={"bucket": source_bucket_name, "key": source_bucket_key} ) - get_hook_lineage_collector().add_output_dataset( - context=self, scheme="s3", dataset_kwargs={"bucket": dest_bucket_name, "key": dest_bucket_key} + get_hook_lineage_collector().add_output_asset( + context=self, scheme="s3", asset_kwargs={"bucket": dest_bucket_name, "key": dest_bucket_key} ) return response @@ -1433,10 +1440,10 @@ def download_file( file_path.parent.mkdir(exist_ok=True, parents=True) - get_hook_lineage_collector().add_output_dataset( + get_hook_lineage_collector().add_output_asset( context=self, scheme="file", - dataset_kwargs={"path": file_path if file_path.is_absolute() else file_path.absolute()}, + asset_kwargs={"path": file_path if file_path.is_absolute() else file_path.absolute()}, ) file = open(file_path, "wb") else: @@ -1448,8 +1455,8 @@ def download_file( ExtraArgs=self.extra_args, Config=self.transfer_config, ) - get_hook_lineage_collector().add_input_dataset( - context=self, scheme="s3", dataset_kwargs={"bucket": bucket_name, "key": key} + get_hook_lineage_collector().add_input_asset( + context=self, scheme="s3", asset_kwargs={"bucket": bucket_name, "key": key} ) return file.name diff --git a/airflow/providers/amazon/aws/utils/asset_compat_lineage_collector.py b/airflow/providers/amazon/aws/utils/asset_compat_lineage_collector.py new file mode 100644 index 0000000000000..50fbc3d0996aa --- /dev/null +++ b/airflow/providers/amazon/aws/utils/asset_compat_lineage_collector.py @@ -0,0 +1,106 @@ +# 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 + +from importlib.util import find_spec + + +def _get_asset_compat_hook_lineage_collector(): + from airflow.lineage.hook import get_hook_lineage_collector + + collector = get_hook_lineage_collector() + + if all( + getattr(collector, asset_method_name, None) + for asset_method_name in ("add_input_asset", "add_output_asset", "collected_assets") + ): + return collector + + # dataset is renamed as asset in Airflow 3.0 + + from functools import wraps + + from airflow.lineage.hook import DatasetLineageInfo, HookLineage + + DatasetLineageInfo.asset = DatasetLineageInfo.dataset + + def rename_dataset_kwargs_as_assets_kwargs(function): + @wraps(function) + def wrapper(*args, **kwargs): + if "asset_kwargs" in kwargs: + kwargs["dataset_kwargs"] = kwargs.pop("asset_kwargs") + + if "asset_extra" in kwargs: + kwargs["dataset_extra"] = kwargs.pop("asset_extra") + + return function(*args, **kwargs) + + return wrapper + + collector.create_asset = rename_dataset_kwargs_as_assets_kwargs(collector.create_dataset) + collector.add_input_asset = rename_dataset_kwargs_as_assets_kwargs(collector.add_input_dataset) + collector.add_output_asset = rename_dataset_kwargs_as_assets_kwargs(collector.add_output_dataset) + + def collected_assets_compat(collector) -> HookLineage: + """Get the collected hook lineage information.""" + lineage = collector.collected_datasets + return HookLineage( + [ + DatasetLineageInfo(dataset=item.dataset, count=item.count, context=item.context) + for item in lineage.inputs + ], + [ + DatasetLineageInfo(dataset=item.dataset, count=item.count, context=item.context) + for item in lineage.outputs + ], + ) + + setattr( + collector.__class__, + "collected_assets", + property(lambda collector: collected_assets_compat(collector)), + ) + + return collector + + +def get_hook_lineage_collector(): + # HookLineageCollector added in 2.10 + try: + if find_spec("airflow.assets"): + # Dataset has been renamed as Asset in 3.0 + from airflow.lineage.hook import get_hook_lineage_collector + + return get_hook_lineage_collector() + + return _get_asset_compat_hook_lineage_collector() + except ImportError: + + class NoOpCollector: + """ + NoOpCollector is a hook lineage collector that does nothing. + + It is used when you want to disable lineage collection. + """ + + def add_input_asset(self, *_, **__): + pass + + def add_output_asset(self, *_, **__): + pass + + return NoOpCollector() diff --git a/airflow/providers/amazon/provider.yaml b/airflow/providers/amazon/provider.yaml index e98a7b53a7253..3b4fe2aec4118 100644 --- a/airflow/providers/amazon/provider.yaml +++ b/airflow/providers/amazon/provider.yaml @@ -563,11 +563,19 @@ sensors: python-modules: - airflow.providers.amazon.aws.sensors.quicksight +asset-uris: + - schemes: [s3] + handler: airflow.providers.amazon.aws.assets.s3.sanitize_uri + to_openlineage_converter: airflow.providers.amazon.aws.assets.s3.convert_asset_to_openlineage + factory: airflow.providers.amazon.aws.assets.s3.create_asset + +# dataset has been renamed to asset in Airflow 3.0 +# This is kept for backward compatibility. dataset-uris: - schemes: [s3] - handler: airflow.providers.amazon.aws.datasets.s3.sanitize_uri - to_openlineage_converter: airflow.providers.amazon.aws.datasets.s3.convert_dataset_to_openlineage - factory: airflow.providers.amazon.aws.datasets.s3.create_dataset + handler: airflow.providers.amazon.aws.assets.s3.sanitize_uri + to_openlineage_converter: airflow.providers.amazon.aws.assets.s3.convert_asset_to_openlineage + factory: airflow.providers.amazon.aws.assets.s3.create_asset filesystems: - airflow.providers.amazon.aws.fs.s3 diff --git a/airflow/providers/common/compat/assets/__init__.py b/airflow/providers/common/compat/assets/__init__.py new file mode 100644 index 0000000000000..460204a4e417f --- /dev/null +++ b/airflow/providers/common/compat/assets/__init__.py @@ -0,0 +1,77 @@ +# 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 + +from typing import TYPE_CHECKING + +from airflow import __version__ as AIRFLOW_VERSION + +if TYPE_CHECKING: + from airflow.assets import ( + Asset, + AssetAlias, + AssetAliasEvent, + AssetAll, + AssetAny, + expand_alias_to_assets, + ) + from airflow.auth.managers.models.resource_details import AssetDetails +else: + try: + from airflow.assets import ( + Asset, + AssetAlias, + AssetAliasEvent, + AssetAll, + AssetAny, + expand_alias_to_assets, + ) + from airflow.auth.managers.models.resource_details import AssetDetails + except ModuleNotFoundError: + from packaging.version import Version + + _IS_AIRFLOW_2_10_OR_HIGHER = Version(Version(AIRFLOW_VERSION).base_version) >= Version("2.10.0") + _IS_AIRFLOW_2_9_OR_HIGHER = Version(Version(AIRFLOW_VERSION).base_version) >= Version("2.9.0") + + # dataset is renamed to asset since Airflow 3.0 + from airflow.auth.managers.models.resource_details import DatasetDetails as AssetDetails + from airflow.datasets import Dataset as Asset + + if _IS_AIRFLOW_2_9_OR_HIGHER: + from airflow.datasets import ( + DatasetAll as AssetAll, + DatasetAny as AssetAny, + ) + + if _IS_AIRFLOW_2_10_OR_HIGHER: + from airflow.datasets import ( + DatasetAlias as AssetAlias, + DatasetAliasEvent as AssetAliasEvent, + expand_alias_to_datasets as expand_alias_to_assets, + ) + + +__all__ = [ + "Asset", + "AssetAlias", + "AssetAliasEvent", + "AssetAll", + "AssetAny", + "AssetDetails", + "expand_alias_to_assets", +] diff --git a/airflow/providers/common/compat/lineage/hook.py b/airflow/providers/common/compat/lineage/hook.py index dbdbc5bf86f4d..50fbc3d0996aa 100644 --- a/airflow/providers/common/compat/lineage/hook.py +++ b/airflow/providers/common/compat/lineage/hook.py @@ -16,13 +16,78 @@ # under the License. from __future__ import annotations +from importlib.util import find_spec + + +def _get_asset_compat_hook_lineage_collector(): + from airflow.lineage.hook import get_hook_lineage_collector + + collector = get_hook_lineage_collector() + + if all( + getattr(collector, asset_method_name, None) + for asset_method_name in ("add_input_asset", "add_output_asset", "collected_assets") + ): + return collector + + # dataset is renamed as asset in Airflow 3.0 + + from functools import wraps + + from airflow.lineage.hook import DatasetLineageInfo, HookLineage + + DatasetLineageInfo.asset = DatasetLineageInfo.dataset + + def rename_dataset_kwargs_as_assets_kwargs(function): + @wraps(function) + def wrapper(*args, **kwargs): + if "asset_kwargs" in kwargs: + kwargs["dataset_kwargs"] = kwargs.pop("asset_kwargs") + + if "asset_extra" in kwargs: + kwargs["dataset_extra"] = kwargs.pop("asset_extra") + + return function(*args, **kwargs) + + return wrapper + + collector.create_asset = rename_dataset_kwargs_as_assets_kwargs(collector.create_dataset) + collector.add_input_asset = rename_dataset_kwargs_as_assets_kwargs(collector.add_input_dataset) + collector.add_output_asset = rename_dataset_kwargs_as_assets_kwargs(collector.add_output_dataset) + + def collected_assets_compat(collector) -> HookLineage: + """Get the collected hook lineage information.""" + lineage = collector.collected_datasets + return HookLineage( + [ + DatasetLineageInfo(dataset=item.dataset, count=item.count, context=item.context) + for item in lineage.inputs + ], + [ + DatasetLineageInfo(dataset=item.dataset, count=item.count, context=item.context) + for item in lineage.outputs + ], + ) + + setattr( + collector.__class__, + "collected_assets", + property(lambda collector: collected_assets_compat(collector)), + ) + + return collector + def get_hook_lineage_collector(): # HookLineageCollector added in 2.10 try: - from airflow.lineage.hook import get_hook_lineage_collector + if find_spec("airflow.assets"): + # Dataset has been renamed as Asset in 3.0 + from airflow.lineage.hook import get_hook_lineage_collector + + return get_hook_lineage_collector() - return get_hook_lineage_collector() + return _get_asset_compat_hook_lineage_collector() except ImportError: class NoOpCollector: @@ -32,10 +97,10 @@ class NoOpCollector: It is used when you want to disable lineage collection. """ - def add_input_dataset(self, *_, **__): + def add_input_asset(self, *_, **__): pass - def add_output_dataset(self, *_, **__): + def add_output_asset(self, *_, **__): pass return NoOpCollector() diff --git a/airflow/providers/common/io/datasets/__init__.py b/airflow/providers/common/compat/openlineage/utils/__init__.py similarity index 100% rename from airflow/providers/common/io/datasets/__init__.py rename to airflow/providers/common/compat/openlineage/utils/__init__.py diff --git a/airflow/providers/common/compat/openlineage/utils/utils.py b/airflow/providers/common/compat/openlineage/utils/utils.py new file mode 100644 index 0000000000000..5492c76d55dcc --- /dev/null +++ b/airflow/providers/common/compat/openlineage/utils/utils.py @@ -0,0 +1,43 @@ +# 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 + +from functools import wraps +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from airflow.providers.openlineage.utils.utils import translate_airflow_asset +else: + try: + from airflow.providers.openlineage.utils.utils import translate_airflow_asset + except ImportError: + from airflow.providers.openlineage.utils.utils import translate_airflow_dataset + + def rename_asset_as_dataset(function): + @wraps(function) + def wrapper(*args, **kwargs): + if "asset" in kwargs: + kwargs["dataset"] = kwargs.pop("asset") + return function(*args, **kwargs) + + return wrapper + + translate_airflow_asset = rename_asset_as_dataset(translate_airflow_dataset) + + +__all__ = ["translate_airflow_asset"] diff --git a/airflow/providers/mysql/datasets/__init__.py b/airflow/providers/common/compat/security/__init__.py similarity index 100% rename from airflow/providers/mysql/datasets/__init__.py rename to airflow/providers/common/compat/security/__init__.py diff --git a/airflow/providers/common/compat/security/permissions.py b/airflow/providers/common/compat/security/permissions.py new file mode 100644 index 0000000000000..d5c351bdad31e --- /dev/null +++ b/airflow/providers/common/compat/security/permissions.py @@ -0,0 +1,30 @@ +# 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 + +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from airflow.security.permissions import RESOURCE_ASSET +else: + try: + from airflow.security.permissions import RESOURCE_ASSET + except ImportError: + from airflow.security.permissions import RESOURCE_DATASET as RESOURCE_ASSET + + +__all__ = ["RESOURCE_ASSET"] diff --git a/airflow/providers/postgres/datasets/__init__.py b/airflow/providers/common/io/assets/__init__.py similarity index 100% rename from airflow/providers/postgres/datasets/__init__.py rename to airflow/providers/common/io/assets/__init__.py diff --git a/airflow/providers/common/io/datasets/file.py b/airflow/providers/common/io/assets/file.py similarity index 76% rename from airflow/providers/common/io/datasets/file.py rename to airflow/providers/common/io/assets/file.py index 35d3b227e5223..fadc4cbe1bdc8 100644 --- a/airflow/providers/common/io/datasets/file.py +++ b/airflow/providers/common/io/assets/file.py @@ -19,7 +19,10 @@ import urllib.parse from typing import TYPE_CHECKING -from airflow.datasets import Dataset +try: + from airflow.assets import Asset +except ModuleNotFoundError: + from airflow.datasets import Dataset as Asset # type: ignore[no-redef] if TYPE_CHECKING: from urllib.parse import SplitResult @@ -27,9 +30,9 @@ from airflow.providers.common.compat.openlineage.facet import Dataset as OpenLineageDataset -def create_dataset(*, path: str, extra=None) -> Dataset: +def create_asset(*, path: str, extra=None) -> Asset: # We assume that we get absolute path starting with / - return Dataset(uri=f"file://{path}", extra=extra) + return Asset(uri=f"file://{path}", extra=extra) def sanitize_uri(uri: SplitResult) -> SplitResult: @@ -38,13 +41,13 @@ def sanitize_uri(uri: SplitResult) -> SplitResult: return uri -def convert_dataset_to_openlineage(dataset: Dataset, lineage_context) -> OpenLineageDataset: +def convert_asset_to_openlineage(asset: Asset, lineage_context) -> OpenLineageDataset: """ - Translate Dataset with valid AIP-60 uri to OpenLineage with assistance from the context. + Translate Asset with valid AIP-60 uri to OpenLineage with assistance from the context. Windows paths are not standardized and can produce unexpected behaviour. """ from airflow.providers.common.compat.openlineage.facet import Dataset as OpenLineageDataset - parsed = urllib.parse.urlsplit(dataset.uri) + parsed = urllib.parse.urlsplit(asset.uri) return OpenLineageDataset(namespace=f"file://{parsed.netloc}", name=parsed.path) diff --git a/airflow/providers/common/io/provider.yaml b/airflow/providers/common/io/provider.yaml index 870605be33b05..6743cfff86c40 100644 --- a/airflow/providers/common/io/provider.yaml +++ b/airflow/providers/common/io/provider.yaml @@ -53,11 +53,19 @@ operators: xcom: - airflow.providers.common.io.xcom.backend +asset-uris: + - schemes: [file] + handler: airflow.providers.common.io.assets.file.sanitize_uri + to_openlineage_converter: airflow.providers.common.io.assets.file.convert_asset_to_openlineage + factory: airflow.providers.common.io.assets.file.create_asset + +# dataset has been renamed to asset in Airflow 3.0 +# This is kept for backward compatibility. dataset-uris: - schemes: [file] - handler: airflow.providers.common.io.datasets.file.sanitize_uri - to_openlineage_converter: airflow.providers.common.io.datasets.file.convert_dataset_to_openlineage - factory: airflow.providers.common.io.datasets.file.create_dataset + handler: airflow.providers.common.io.assets.file.sanitize_uri + to_openlineage_converter: airflow.providers.common.io.assets.file.convert_asset_to_openlineage + factory: airflow.providers.common.io.assets.file.create_asset config: common.io: diff --git a/airflow/providers/fab/auth_manager/fab_auth_manager.py b/airflow/providers/fab/auth_manager/fab_auth_manager.py index 3d0f102650935..425f2d6d2124f 100644 --- a/airflow/providers/fab/auth_manager/fab_auth_manager.py +++ b/airflow/providers/fab/auth_manager/fab_auth_manager.py @@ -37,7 +37,6 @@ ConnectionDetails, DagAccessEntity, DagDetails, - DatasetDetails, PoolDetails, VariableDetails, ) @@ -67,7 +66,6 @@ RESOURCE_DAG_DEPENDENCIES, RESOURCE_DAG_RUN, RESOURCE_DAG_WARNING, - RESOURCE_DATASET, RESOURCE_DOCS, RESOURCE_IMPORT_ERROR, RESOURCE_JOB, @@ -94,7 +92,15 @@ from airflow.cli.cli_config import ( CLICommand, ) + from airflow.providers.common.compat.assets import AssetDetails from airflow.providers.fab.auth_manager.security_manager.override import FabAirflowSecurityManagerOverride + from airflow.security.permissions import RESOURCE_ASSET +else: + try: + from airflow.security.permissions import RESOURCE_ASSET + except ImportError: + from airflow.security.permissions import RESOURCE_DATASET as RESOURCE_ASSET + _MAP_DAG_ACCESS_ENTITY_TO_FAB_RESOURCE_TYPE: dict[DagAccessEntity, tuple[str, ...]] = { DagAccessEntity.AUDIT_LOG: (RESOURCE_AUDIT_LOG,), @@ -263,10 +269,10 @@ def is_authorized_dag( for resource_type in resource_types ) - def is_authorized_dataset( - self, *, method: ResourceMethod, details: DatasetDetails | None = None, user: BaseUser | None = None + def is_authorized_asset( + self, *, method: ResourceMethod, details: AssetDetails | None = None, user: BaseUser | None = None ) -> bool: - return self._is_authorized(method=method, resource_type=RESOURCE_DATASET, user=user) + return self._is_authorized(method=method, resource_type=RESOURCE_ASSET, user=user) def is_authorized_pool( self, *, method: ResourceMethod, details: PoolDetails | None = None, user: BaseUser | None = None diff --git a/airflow/providers/fab/auth_manager/security_manager/override.py b/airflow/providers/fab/auth_manager/security_manager/override.py index 0f4b79b4f1f68..640d406f81942 100644 --- a/airflow/providers/fab/auth_manager/security_manager/override.py +++ b/airflow/providers/fab/auth_manager/security_manager/override.py @@ -115,6 +115,12 @@ if TYPE_CHECKING: from airflow.auth.managers.base_auth_manager import ResourceMethod + from airflow.security.permissions import RESOURCE_ASSET +else: + try: + from airflow.security.permissions import RESOURCE_ASSET + except ImportError: + from airflow.security.permissions import RESOURCE_DATASET as RESOURCE_ASSET log = logging.getLogger(__name__) @@ -234,7 +240,7 @@ class FabAirflowSecurityManagerOverride(AirflowSecurityManagerV2): (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_DEPENDENCIES), (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_CODE), (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DATASET), + (permissions.ACTION_CAN_READ, RESOURCE_ASSET), (permissions.ACTION_CAN_READ, permissions.RESOURCE_CLUSTER_ACTIVITY), (permissions.ACTION_CAN_READ, permissions.RESOURCE_POOL), (permissions.ACTION_CAN_READ, permissions.RESOURCE_IMPORT_ERROR), @@ -253,7 +259,7 @@ class FabAirflowSecurityManagerOverride(AirflowSecurityManagerV2): (permissions.ACTION_CAN_ACCESS_MENU, permissions.RESOURCE_DAG), (permissions.ACTION_CAN_ACCESS_MENU, permissions.RESOURCE_DAG_DEPENDENCIES), (permissions.ACTION_CAN_ACCESS_MENU, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_ACCESS_MENU, permissions.RESOURCE_DATASET), + (permissions.ACTION_CAN_ACCESS_MENU, RESOURCE_ASSET), (permissions.ACTION_CAN_ACCESS_MENU, permissions.RESOURCE_CLUSTER_ACTIVITY), (permissions.ACTION_CAN_ACCESS_MENU, permissions.RESOURCE_DOCS), (permissions.ACTION_CAN_ACCESS_MENU, permissions.RESOURCE_DOCS_MENU), @@ -273,7 +279,7 @@ class FabAirflowSecurityManagerOverride(AirflowSecurityManagerV2): (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_DAG_RUN), (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG_RUN), (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_DATASET), + (permissions.ACTION_CAN_CREATE, RESOURCE_ASSET), ] # [END security_user_perms] @@ -302,8 +308,8 @@ class FabAirflowSecurityManagerOverride(AirflowSecurityManagerV2): (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_VARIABLE), (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_VARIABLE), (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_XCOM), - (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_DATASET), - (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_DATASET), + (permissions.ACTION_CAN_DELETE, RESOURCE_ASSET), + (permissions.ACTION_CAN_CREATE, RESOURCE_ASSET), ] # [END security_op_perms] diff --git a/airflow/providers/fab/provider.yaml b/airflow/providers/fab/provider.yaml index 63be11c264938..1d5cc820f1f24 100644 --- a/airflow/providers/fab/provider.yaml +++ b/airflow/providers/fab/provider.yaml @@ -47,6 +47,7 @@ versions: dependencies: - apache-airflow>=2.9.0 + - apache-airflow-providers-common-compat>=1.2.0 - flask>=2.2,<2.3 # We are tightly coupled with FAB version as we vendored-in part of FAB code related to security manager # This is done as part of preparation to removing FAB as dependency, but we are not ready for it yet diff --git a/airflow/providers/google/provider.yaml b/airflow/providers/google/provider.yaml index 612fd8e29bac3..a64b2ce17a76e 100644 --- a/airflow/providers/google/provider.yaml +++ b/airflow/providers/google/provider.yaml @@ -768,6 +768,14 @@ sensors: filesystems: - airflow.providers.google.cloud.fs.gcs +asset-uris: + - schemes: [gcp] + handler: null + - schemes: [bigquery] + handler: airflow.providers.google.datasets.bigquery.sanitize_uri + +# dataset has been renamed to asset in Airflow 3.0 +# This is kept for backward compatibility. dataset-uris: - schemes: [gcp] handler: null diff --git a/airflow/providers/trino/datasets/__init__.py b/airflow/providers/mysql/assets/__init__.py similarity index 100% rename from airflow/providers/trino/datasets/__init__.py rename to airflow/providers/mysql/assets/__init__.py diff --git a/airflow/providers/mysql/datasets/mysql.py b/airflow/providers/mysql/assets/mysql.py similarity index 100% rename from airflow/providers/mysql/datasets/mysql.py rename to airflow/providers/mysql/assets/mysql.py diff --git a/airflow/providers/mysql/provider.yaml b/airflow/providers/mysql/provider.yaml index 28ba986b64ccd..f0f77f28d0e94 100644 --- a/airflow/providers/mysql/provider.yaml +++ b/airflow/providers/mysql/provider.yaml @@ -113,6 +113,12 @@ connection-types: - hook-class-name: airflow.providers.mysql.hooks.mysql.MySqlHook connection-type: mysql +asset-uris: + - schemes: [mysql, mariadb] + handler: airflow.providers.mysql.assets.mysql.sanitize_uri + +# dataset has been renamed to asset in Airflow 3.0 +# This is kept for backward compatibility. dataset-uris: - schemes: [mysql, mariadb] - handler: airflow.providers.mysql.datasets.mysql.sanitize_uri + handler: airflow.providers.mysql.assets.mysql.sanitize_uri diff --git a/airflow/providers/openlineage/extractors/manager.py b/airflow/providers/openlineage/extractors/manager.py index 74be9e01f4b7b..c72c989d8936e 100644 --- a/airflow/providers/openlineage/extractors/manager.py +++ b/airflow/providers/openlineage/extractors/manager.py @@ -18,6 +18,7 @@ from typing import TYPE_CHECKING, Iterator +from airflow.providers.common.compat.openlineage.utils.utils import translate_airflow_asset from airflow.providers.openlineage import conf from airflow.providers.openlineage.extractors import BaseExtractor, OperatorLineage from airflow.providers.openlineage.extractors.base import DefaultExtractor @@ -25,7 +26,6 @@ from airflow.providers.openlineage.extractors.python import PythonExtractor from airflow.providers.openlineage.utils.utils import ( get_unknown_source_attribute_run_facet, - translate_airflow_dataset, try_import_from_string, ) from airflow.utils.log.logging_mixin import LoggingMixin @@ -178,7 +178,16 @@ def extract_inlets_and_outlets( def get_hook_lineage(self) -> tuple[list[Dataset], list[Dataset]] | None: try: - from airflow.lineage.hook import get_hook_lineage_collector + from importlib.util import find_spec + + if find_spec("airflow.assets"): + from airflow.lineage.hook import get_hook_lineage_collector + else: + # TODO: import from common.compat directly after common.compat providers with + # asset_compat_lineage_collector released + from airflow.providers.openlineage.utils.asset_compat_lineage_collector import ( + get_hook_lineage_collector, + ) except ImportError: return None @@ -187,16 +196,14 @@ def get_hook_lineage(self) -> tuple[list[Dataset], list[Dataset]] | None: return ( [ - dataset - for dataset_info in get_hook_lineage_collector().collected_datasets.inputs - if (dataset := translate_airflow_dataset(dataset_info.dataset, dataset_info.context)) - is not None + asset + for asset_info in get_hook_lineage_collector().collected_assets.inputs + if (asset := translate_airflow_asset(asset_info.asset, asset_info.context)) is not None ], [ - dataset - for dataset_info in get_hook_lineage_collector().collected_datasets.outputs - if (dataset := translate_airflow_dataset(dataset_info.dataset, dataset_info.context)) - is not None + asset + for asset_info in get_hook_lineage_collector().collected_assets.outputs + if (asset := translate_airflow_asset(asset_info.asset, asset_info.context)) is not None ], ) diff --git a/airflow/providers/openlineage/utils/asset_compat_lineage_collector.py b/airflow/providers/openlineage/utils/asset_compat_lineage_collector.py new file mode 100644 index 0000000000000..8a4d2b61914ff --- /dev/null +++ b/airflow/providers/openlineage/utils/asset_compat_lineage_collector.py @@ -0,0 +1,108 @@ +# 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 + +from importlib.util import find_spec + +# TODO: replace this module with common.compat provider once common.compat 1.3.0 released + + +def _get_asset_compat_hook_lineage_collector(): + from airflow.lineage.hook import get_hook_lineage_collector + + collector = get_hook_lineage_collector() + + if all( + getattr(collector, asset_method_name, None) + for asset_method_name in ("add_input_asset", "add_output_asset", "collected_assets") + ): + return collector + + # dataset is renamed as asset in Airflow 3.0 + + from functools import wraps + + from airflow.lineage.hook import DatasetLineageInfo, HookLineage + + DatasetLineageInfo.asset = DatasetLineageInfo.dataset + + def rename_dataset_kwargs_as_assets_kwargs(function): + @wraps(function) + def wrapper(*args, **kwargs): + if "asset_kwargs" in kwargs: + kwargs["dataset_kwargs"] = kwargs.pop("asset_kwargs") + + if "asset_extra" in kwargs: + kwargs["dataset_extra"] = kwargs.pop("asset_extra") + + return function(*args, **kwargs) + + return wrapper + + collector.create_asset = rename_dataset_kwargs_as_assets_kwargs(collector.create_dataset) + collector.add_input_asset = rename_dataset_kwargs_as_assets_kwargs(collector.add_input_dataset) + collector.add_output_asset = rename_dataset_kwargs_as_assets_kwargs(collector.add_output_dataset) + + def collected_assets_compat(collector) -> HookLineage: + """Get the collected hook lineage information.""" + lineage = collector.collected_datasets + return HookLineage( + [ + DatasetLineageInfo(dataset=item.dataset, count=item.count, context=item.context) + for item in lineage.inputs + ], + [ + DatasetLineageInfo(dataset=item.dataset, count=item.count, context=item.context) + for item in lineage.outputs + ], + ) + + setattr( + collector.__class__, + "collected_assets", + property(lambda collector: collected_assets_compat(collector)), + ) + + return collector + + +def get_hook_lineage_collector(): + # HookLineageCollector added in 2.10 + try: + if find_spec("airflow.assets"): + # Dataset has been renamed as Asset in 3.0 + from airflow.lineage.hook import get_hook_lineage_collector + + return get_hook_lineage_collector() + + return _get_asset_compat_hook_lineage_collector() + except ImportError: + + class NoOpCollector: + """ + NoOpCollector is a hook lineage collector that does nothing. + + It is used when you want to disable lineage collection. + """ + + def add_input_asset(self, *_, **__): + pass + + def add_output_asset(self, *_, **__): + pass + + return NoOpCollector() diff --git a/airflow/providers/openlineage/utils/utils.py b/airflow/providers/openlineage/utils/utils.py index f283c09e8759c..ca57755d692ea 100644 --- a/airflow/providers/openlineage/utils/utils.py +++ b/airflow/providers/openlineage/utils/utils.py @@ -31,7 +31,6 @@ from packaging.version import Version from airflow import __version__ as AIRFLOW_VERSION -from airflow.datasets import Dataset from airflow.exceptions import AirflowProviderDeprecationWarning # TODO: move this maybe to Airflow's logic? from airflow.models import DAG, BaseOperator, DagRun, MappedOperator from airflow.providers.openlineage import conf @@ -54,6 +53,11 @@ from airflow.utils.log.secrets_masker import Redactable, Redacted, SecretsMasker, should_hide_value_for_key from airflow.utils.module_loading import import_string +try: + from airflow.assets import Asset +except ModuleNotFoundError: + from airflow.datasets import Dataset as Asset # type: ignore[no-redef] + if TYPE_CHECKING: from openlineage.client.event_v2 import Dataset as OpenLineageDataset from openlineage.client.facet_v2 import RunFacet @@ -283,8 +287,8 @@ class TaskInstanceInfo(InfoJsonEncodable): } -class DatasetInfo(InfoJsonEncodable): - """Defines encoding Airflow Dataset object to JSON.""" +class AssetInfo(InfoJsonEncodable): + """Defines encoding Airflow Asset object to JSON.""" includes = ["uri", "extra"] @@ -335,8 +339,8 @@ class TaskInfo(InfoJsonEncodable): if hasattr(task, "task_group") and getattr(task.task_group, "_group_id", None) else None ), - "inlets": lambda task: [DatasetInfo(i) for i in task.inlets if isinstance(i, Dataset)], - "outlets": lambda task: [DatasetInfo(o) for o in task.outlets if isinstance(o, Dataset)], + "inlets": lambda task: [AssetInfo(i) for i in task.inlets if isinstance(i, Asset)], + "outlets": lambda task: [AssetInfo(o) for o in task.outlets if isinstance(o, Asset)], } @@ -641,19 +645,29 @@ def should_use_external_connection(hook) -> bool: return True -def translate_airflow_dataset(dataset: Dataset, lineage_context) -> OpenLineageDataset | None: +def translate_airflow_asset(asset: Asset, lineage_context) -> OpenLineageDataset | None: """ - Convert a Dataset with an AIP-60 compliant URI to an OpenLineageDataset. + Convert a Asset with an AIP-60 compliant URI to an OpenLineageDataset. - This function returns None if no URI normalizer is defined, no dataset converter is found or + This function returns None if no URI normalizer is defined, no asset converter is found or some core Airflow changes are missing and ImportError is raised. """ try: - from airflow.datasets import _get_normalized_scheme + from airflow.assets import _get_normalized_scheme + except ModuleNotFoundError: + try: + from airflow.datasets import _get_normalized_scheme # type: ignore[no-redef] + except ImportError: + return None + + try: from airflow.providers_manager import ProvidersManager - ol_converters = ProvidersManager().dataset_to_openlineage_converters - normalized_uri = dataset.normalized_uri + ol_converters = getattr(ProvidersManager(), "asset_to_openlineage_converters", None) + if not ol_converters: + ol_converters = ProvidersManager().dataset_to_openlineage_converters # type: ignore[attr-defined] + + normalized_uri = asset.normalized_uri except (ImportError, AttributeError): return None @@ -666,4 +680,4 @@ def translate_airflow_dataset(dataset: Dataset, lineage_context) -> OpenLineageD if (airflow_to_ol_converter := ol_converters.get(normalized_scheme)) is None: return None - return airflow_to_ol_converter(Dataset(uri=normalized_uri, extra=dataset.extra), lineage_context) + return airflow_to_ol_converter(Asset(uri=normalized_uri, extra=asset.extra), lineage_context) diff --git a/tests/datasets/__init__.py b/airflow/providers/postgres/assets/__init__.py similarity index 100% rename from tests/datasets/__init__.py rename to airflow/providers/postgres/assets/__init__.py diff --git a/airflow/providers/postgres/datasets/postgres.py b/airflow/providers/postgres/assets/postgres.py similarity index 100% rename from airflow/providers/postgres/datasets/postgres.py rename to airflow/providers/postgres/assets/postgres.py diff --git a/airflow/providers/postgres/provider.yaml b/airflow/providers/postgres/provider.yaml index 7ce95986fd65a..edbbaeb1da2c1 100644 --- a/airflow/providers/postgres/provider.yaml +++ b/airflow/providers/postgres/provider.yaml @@ -96,6 +96,12 @@ connection-types: - hook-class-name: airflow.providers.postgres.hooks.postgres.PostgresHook connection-type: postgres +asset-uris: + - schemes: [postgres, postgresql] + handler: airflow.providers.postgres.assets.postgres.sanitize_uri + +# dataset has been renamed to asset in Airflow 3.0 +# This is kept for backward compatibility. dataset-uris: - schemes: [postgres, postgresql] - handler: airflow.providers.postgres.datasets.postgres.sanitize_uri + handler: airflow.providers.postgres.assets.postgres.sanitize_uri diff --git a/tests/providers/amazon/aws/datasets/__init__.py b/airflow/providers/trino/assets/__init__.py similarity index 100% rename from tests/providers/amazon/aws/datasets/__init__.py rename to airflow/providers/trino/assets/__init__.py diff --git a/airflow/providers/trino/datasets/trino.py b/airflow/providers/trino/assets/trino.py similarity index 100% rename from airflow/providers/trino/datasets/trino.py rename to airflow/providers/trino/assets/trino.py diff --git a/airflow/providers/trino/provider.yaml b/airflow/providers/trino/provider.yaml index d4000baaa063f..424be2cca67d9 100644 --- a/airflow/providers/trino/provider.yaml +++ b/airflow/providers/trino/provider.yaml @@ -86,9 +86,15 @@ operators: python-modules: - airflow.providers.trino.operators.trino +asset-uris: + - schemes: [trino] + handler: airflow.providers.trino.assets.trino.sanitize_uri + +# dataset has been renamed to asset in Airflow 3.0 +# This is kept for backward compatibility. dataset-uris: - schemes: [trino] - handler: airflow.providers.trino.datasets.trino.sanitize_uri + handler: airflow.providers.trino.assets.trino.sanitize_uri hooks: - integration-name: Trino diff --git a/airflow/providers_manager.py b/airflow/providers_manager.py index dd3e841fa1662..2c673063cb23e 100644 --- a/airflow/providers_manager.py +++ b/airflow/providers_manager.py @@ -91,7 +91,7 @@ def ensure_prefix(field): if TYPE_CHECKING: from urllib.parse import SplitResult - from airflow.datasets import Dataset + from airflow.assets import Asset from airflow.decorators.base import TaskDecorator from airflow.hooks.base import BaseHook from airflow.typing_compat import Literal @@ -426,9 +426,9 @@ def __init__(self): # Keeps dict of hooks keyed by connection type self._hooks_dict: dict[str, HookInfo] = {} self._fs_set: set[str] = set() - self._dataset_uri_handlers: dict[str, Callable[[SplitResult], SplitResult]] = {} - self._dataset_factories: dict[str, Callable[..., Dataset]] = {} - self._dataset_to_openlineage_converters: dict[str, Callable] = {} + self._asset_uri_handlers: dict[str, Callable[[SplitResult], SplitResult]] = {} + self._asset_factories: dict[str, Callable[..., Asset]] = {} + self._asset_to_openlineage_converters: dict[str, Callable] = {} self._taskflow_decorators: dict[str, Callable] = LazyDictWithCache() # type: ignore[assignment] # keeps mapping between connection_types and hook class, package they come from self._hook_provider_dict: dict[str, HookClassProvider] = {} @@ -525,11 +525,11 @@ def initialize_providers_filesystems(self): self.initialize_providers_list() self._discover_filesystems() - @provider_info_cache("dataset_uris") - def initialize_providers_dataset_uri_resources(self): - """Lazy initialization of provider dataset URI handlers, factories, converters etc.""" + @provider_info_cache("asset_uris") + def initialize_providers_asset_uri_resources(self): + """Lazy initialization of provider asset URI handlers, factories, converters etc.""" self.initialize_providers_list() - self._discover_dataset_uri_resources() + self._discover_asset_uri_resources() @provider_info_cache("hook_lineage_writers") @provider_info_cache("taskflow_decorators") @@ -882,9 +882,9 @@ def _discover_filesystems(self) -> None: self._fs_set.add(fs_module_name) self._fs_set = set(sorted(self._fs_set)) - def _discover_dataset_uri_resources(self) -> None: - """Discovers and registers dataset URI handlers, factories, and converters for all providers.""" - from airflow.datasets import normalize_noop + def _discover_asset_uri_resources(self) -> None: + """Discovers and registers asset URI handlers, factories, and converters for all providers.""" + from airflow.assets import normalize_noop def _safe_register_resource( provider_package_name: str, @@ -908,24 +908,24 @@ def _safe_register_resource( resource_registry.update((scheme, resource) for scheme in schemes_list) for provider_name, provider in self._provider_dict.items(): - for uri_info in provider.data.get("dataset-uris", []): + for uri_info in provider.data.get("asset-uris", []): if "schemes" not in uri_info or "handler" not in uri_info: continue # Both schemas and handler must be explicitly set, handler can be set to null common_args = {"schemes_list": uri_info["schemes"], "provider_package_name": provider_name} _safe_register_resource( resource_path=uri_info["handler"], - resource_registry=self._dataset_uri_handlers, + resource_registry=self._asset_uri_handlers, default_resource=normalize_noop, **common_args, ) _safe_register_resource( resource_path=uri_info.get("factory"), - resource_registry=self._dataset_factories, + resource_registry=self._asset_factories, **common_args, ) _safe_register_resource( resource_path=uri_info.get("to_openlineage_converter"), - resource_registry=self._dataset_to_openlineage_converters, + resource_registry=self._asset_to_openlineage_converters, **common_args, ) @@ -1325,21 +1325,21 @@ def filesystem_module_names(self) -> list[str]: return sorted(self._fs_set) @property - def dataset_factories(self) -> dict[str, Callable[..., Dataset]]: - self.initialize_providers_dataset_uri_resources() - return self._dataset_factories + def asset_factories(self) -> dict[str, Callable[..., Asset]]: + self.initialize_providers_asset_uri_resources() + return self._asset_factories @property - def dataset_uri_handlers(self) -> dict[str, Callable[[SplitResult], SplitResult]]: - self.initialize_providers_dataset_uri_resources() - return self._dataset_uri_handlers + def asset_uri_handlers(self) -> dict[str, Callable[[SplitResult], SplitResult]]: + self.initialize_providers_asset_uri_resources() + return self._asset_uri_handlers @property - def dataset_to_openlineage_converters( + def asset_to_openlineage_converters( self, ) -> dict[str, Callable]: - self.initialize_providers_dataset_uri_resources() - return self._dataset_to_openlineage_converters + self.initialize_providers_asset_uri_resources() + return self._asset_to_openlineage_converters @property def provider_configs(self) -> list[tuple[str, dict[str, Any]]]: diff --git a/airflow/reproducible_build.yaml b/airflow/reproducible_build.yaml index 1bf308b87a705..8a35282492059 100644 --- a/airflow/reproducible_build.yaml +++ b/airflow/reproducible_build.yaml @@ -1,2 +1,2 @@ -release-notes-hash: 828fa8d5e93e215963c0a3e52e7f1e3d -source-date-epoch: 1727075869 +release-notes-hash: cc9c5c2ea1cade5d714aa4832587e13a +source-date-epoch: 1727595745 diff --git a/airflow/security/permissions.py b/airflow/security/permissions.py index 45b56c342b44e..acd245865a4ad 100644 --- a/airflow/security/permissions.py +++ b/airflow/security/permissions.py @@ -33,7 +33,7 @@ RESOURCE_DAG_RUN_PREFIX = "DAG Run:" RESOURCE_DAG_WARNING = "DAG Warnings" RESOURCE_CLUSTER_ACTIVITY = "Cluster Activity" -RESOURCE_DATASET = "Datasets" +RESOURCE_ASSET = "Assets" RESOURCE_DOCS = "Documentation" RESOURCE_DOCS_MENU = "Docs" RESOURCE_IMPORT_ERROR = "ImportError" diff --git a/airflow/serialization/dag_dependency.py b/airflow/serialization/dag_dependency.py index bff1b39ebe04b..bede95ba9235b 100644 --- a/airflow/serialization/dag_dependency.py +++ b/airflow/serialization/dag_dependency.py @@ -36,7 +36,7 @@ class DagDependency: def node_id(self): """Node ID for graph rendering.""" val = f"{self.dependency_type}" - if self.dependency_type not in ("dataset", "dataset-alias"): + if self.dependency_type not in ("asset", "asset-alias"): val += f":{self.source}:{self.target}" if self.dependency_id: val += f":{self.dependency_id}" diff --git a/airflow/serialization/enums.py b/airflow/serialization/enums.py index 49a3de3d774c4..dd63366b8a958 100644 --- a/airflow/serialization/enums.py +++ b/airflow/serialization/enums.py @@ -37,8 +37,8 @@ class DagAttributeTypes(str, Enum): """Enum of supported attribute types of DAG.""" DAG = "dag" - DATASET_EVENT_ACCESSORS = "dataset_event_accessors" - DATASET_EVENT_ACCESSOR = "dataset_event_accessor" + ASSET_EVENT_ACCESSORS = "asset_event_accessors" + ASSET_EVENT_ACCESSOR = "asset_event_accessor" OP = "operator" DATETIME = "datetime" TIMEDELTA = "timedelta" @@ -55,16 +55,15 @@ class DagAttributeTypes(str, Enum): EDGE_INFO = "edgeinfo" PARAM = "param" XCOM_REF = "xcomref" - DATASET = "dataset" - DATASET_ALIAS = "dataset_alias" - DATASET_ANY = "dataset_any" - DATASET_ALL = "dataset_all" + ASSET = "asset" + ASSET_ALIAS = "asset_alias" + ASSET_ANY = "asset_any" + ASSET_ALL = "asset_all" SIMPLE_TASK_INSTANCE = "simple_task_instance" BASE_JOB = "Job" TASK_INSTANCE = "task_instance" DAG_RUN = "dag_run" DAG_MODEL = "dag_model" - DATA_SET = "data_set" LOG_TEMPLATE = "log_template" CONNECTION = "connection" TASK_CONTEXT = "task_context" diff --git a/airflow/serialization/pydantic/dataset.py b/airflow/serialization/pydantic/asset.py similarity index 68% rename from airflow/serialization/pydantic/dataset.py rename to airflow/serialization/pydantic/asset.py index 0c233a3fd67c6..29806d3bdf911 100644 --- a/airflow/serialization/pydantic/dataset.py +++ b/airflow/serialization/pydantic/asset.py @@ -20,8 +20,8 @@ from pydantic import BaseModel as BaseModelPydantic, ConfigDict -class DagScheduleDatasetReferencePydantic(BaseModelPydantic): - """Serializable version of the DagScheduleDatasetReference ORM SqlAlchemyModel used by internal API.""" +class DagScheduleAssetReferencePydantic(BaseModelPydantic): + """Serializable version of the DagScheduleAssetReference ORM SqlAlchemyModel used by internal API.""" dataset_id: int dag_id: str @@ -31,8 +31,8 @@ class DagScheduleDatasetReferencePydantic(BaseModelPydantic): model_config = ConfigDict(from_attributes=True) -class TaskOutletDatasetReferencePydantic(BaseModelPydantic): - """Serializable version of the TaskOutletDatasetReference ORM SqlAlchemyModel used by internal API.""" +class TaskOutletAssetReferencePydantic(BaseModelPydantic): + """Serializable version of the TaskOutletAssetReference ORM SqlAlchemyModel used by internal API.""" dataset_id: int dag_id: str @@ -43,8 +43,8 @@ class TaskOutletDatasetReferencePydantic(BaseModelPydantic): model_config = ConfigDict(from_attributes=True) -class DatasetPydantic(BaseModelPydantic): - """Serializable representation of the Dataset ORM SqlAlchemyModel used by internal API.""" +class AssetPydantic(BaseModelPydantic): + """Serializable representation of the Asset ORM SqlAlchemyModel used by internal API.""" id: int uri: str @@ -53,14 +53,14 @@ class DatasetPydantic(BaseModelPydantic): updated_at: datetime is_orphaned: bool - consuming_dags: List[DagScheduleDatasetReferencePydantic] - producing_tasks: List[TaskOutletDatasetReferencePydantic] + consuming_dags: List[DagScheduleAssetReferencePydantic] + producing_tasks: List[TaskOutletAssetReferencePydantic] model_config = ConfigDict(from_attributes=True) -class DatasetEventPydantic(BaseModelPydantic): - """Serializable representation of the DatasetEvent ORM SqlAlchemyModel used by internal API.""" +class AssetEventPydantic(BaseModelPydantic): + """Serializable representation of the AssetEvent ORM SqlAlchemyModel used by internal API.""" id: int dataset_id: Optional[int] @@ -70,6 +70,6 @@ class DatasetEventPydantic(BaseModelPydantic): source_run_id: Optional[str] source_map_index: Optional[int] timestamp: datetime - dataset: Optional[DatasetPydantic] + dataset: Optional[AssetPydantic] model_config = ConfigDict(from_attributes=True, arbitrary_types_allowed=True) diff --git a/airflow/serialization/pydantic/dag_run.py b/airflow/serialization/pydantic/dag_run.py index a3a53c6d941f4..86857452e8310 100644 --- a/airflow/serialization/pydantic/dag_run.py +++ b/airflow/serialization/pydantic/dag_run.py @@ -22,8 +22,8 @@ from pydantic import BaseModel as BaseModelPydantic, ConfigDict from airflow.models.dagrun import DagRun +from airflow.serialization.pydantic.asset import AssetEventPydantic from airflow.serialization.pydantic.dag import PydanticDag -from airflow.serialization.pydantic.dataset import DatasetEventPydantic from airflow.utils.types import DagRunTriggeredByType if TYPE_CHECKING: @@ -55,7 +55,7 @@ class DagRunPydantic(BaseModelPydantic): dag_hash: Optional[str] updated_at: Optional[datetime] dag: Optional[PydanticDag] - consumed_dataset_events: List[DatasetEventPydantic] # noqa: UP006 + consumed_dataset_events: List[AssetEventPydantic] # noqa: UP006 log_template_id: Optional[int] triggered_by: Optional[DagRunTriggeredByType] diff --git a/airflow/serialization/pydantic/taskinstance.py b/airflow/serialization/pydantic/taskinstance.py index 549b03680df83..caf44bea4c673 100644 --- a/airflow/serialization/pydantic/taskinstance.py +++ b/airflow/serialization/pydantic/taskinstance.py @@ -509,8 +509,8 @@ def command_as_list( cfg_path=cfg_path, ) - def _register_dataset_changes(self, *, events, session: Session | None = None) -> None: - TaskInstance._register_dataset_changes(self=self, events=events, session=session) # type: ignore[arg-type] + def _register_asset_changes(self, *, events, session: Session | None = None) -> None: + TaskInstance._register_asset_changes(self=self, events=events, session=session) # type: ignore[arg-type] def defer_task(self, exception: TaskDeferred, session: Session | None = None): """Defer task.""" diff --git a/airflow/serialization/serialized_objects.py b/airflow/serialization/serialized_objects.py index c9c1f11835277..a4801b767acc5 100644 --- a/airflow/serialization/serialized_objects.py +++ b/airflow/serialization/serialized_objects.py @@ -34,16 +34,16 @@ from pendulum.tz.timezone import FixedTimezone, Timezone from airflow import macros +from airflow.assets import ( + Asset, + AssetAlias, + AssetAll, + AssetAny, + BaseAsset, + _AssetAliasCondition, +) from airflow.callbacks.callback_requests import DagCallbackRequest, TaskCallbackRequest from airflow.compat.functools import cache -from airflow.datasets import ( - BaseDataset, - Dataset, - DatasetAlias, - DatasetAll, - DatasetAny, - _DatasetAliasCondition, -) from airflow.exceptions import AirflowException, SerializationError, TaskDeferred from airflow.jobs.job import Job from airflow.models import Trigger @@ -63,9 +63,9 @@ from airflow.serialization.enums import DagAttributeTypes as DAT, Encoding from airflow.serialization.helpers import serialize_template_field from airflow.serialization.json_schema import load_dag_schema +from airflow.serialization.pydantic.asset import AssetPydantic from airflow.serialization.pydantic.dag import DagModelPydantic from airflow.serialization.pydantic.dag_run import DagRunPydantic -from airflow.serialization.pydantic.dataset import DatasetPydantic from airflow.serialization.pydantic.job import JobPydantic from airflow.serialization.pydantic.taskinstance import TaskInstancePydantic from airflow.serialization.pydantic.tasklog import LogTemplatePydantic @@ -246,38 +246,38 @@ def __str__(self) -> str: ) -def encode_dataset_condition(var: BaseDataset) -> dict[str, Any]: +def encode_asset_condition(var: BaseAsset) -> dict[str, Any]: """ - Encode a dataset condition. + Encode an asset condition. :meta private: """ - if isinstance(var, Dataset): - return {"__type": DAT.DATASET, "uri": var.uri, "extra": var.extra} - if isinstance(var, DatasetAlias): - return {"__type": DAT.DATASET_ALIAS, "name": var.name} - if isinstance(var, DatasetAll): - return {"__type": DAT.DATASET_ALL, "objects": [encode_dataset_condition(x) for x in var.objects]} - if isinstance(var, DatasetAny): - return {"__type": DAT.DATASET_ANY, "objects": [encode_dataset_condition(x) for x in var.objects]} + if isinstance(var, Asset): + return {"__type": DAT.ASSET, "uri": var.uri, "extra": var.extra} + if isinstance(var, AssetAlias): + return {"__type": DAT.ASSET_ALIAS, "name": var.name} + if isinstance(var, AssetAll): + return {"__type": DAT.ASSET_ALL, "objects": [encode_asset_condition(x) for x in var.objects]} + if isinstance(var, AssetAny): + return {"__type": DAT.ASSET_ANY, "objects": [encode_asset_condition(x) for x in var.objects]} raise ValueError(f"serialization not implemented for {type(var).__name__!r}") -def decode_dataset_condition(var: dict[str, Any]) -> BaseDataset: +def decode_asset_condition(var: dict[str, Any]) -> BaseAsset: """ Decode a previously serialized dataset condition. :meta private: """ dat = var["__type"] - if dat == DAT.DATASET: - return Dataset(var["uri"], extra=var["extra"]) - if dat == DAT.DATASET_ALL: - return DatasetAll(*(decode_dataset_condition(x) for x in var["objects"])) - if dat == DAT.DATASET_ANY: - return DatasetAny(*(decode_dataset_condition(x) for x in var["objects"])) - if dat == DAT.DATASET_ALIAS: - return DatasetAlias(name=var["name"]) + if dat == DAT.ASSET: + return Asset(var["uri"], extra=var["extra"]) + if dat == DAT.ASSET_ALL: + return AssetAll(*(decode_asset_condition(x) for x in var["objects"])) + if dat == DAT.ASSET_ANY: + return AssetAny(*(decode_asset_condition(x) for x in var["objects"])) + if dat == DAT.ASSET_ALIAS: + return AssetAlias(name=var["name"]) raise ValueError(f"deserialization not implemented for DAT {dat!r}") @@ -285,23 +285,18 @@ def encode_outlet_event_accessor(var: OutletEventAccessor) -> dict[str, Any]: raw_key = var.raw_key return { "extra": var.extra, - "dataset_alias_events": var.dataset_alias_events, + "asset_alias_events": var.asset_alias_events, "raw_key": BaseSerialization.serialize(raw_key), } def decode_outlet_event_accessor(var: dict[str, Any]) -> OutletEventAccessor: - # This is added for compatibility. The attribute used to be dataset_alias_event and - # is now dataset_alias_events. - if dataset_alias_event := var.get("dataset_alias_event", None): - dataset_alias_events = [dataset_alias_event] - else: - dataset_alias_events = var.get("dataset_alias_events", []) + asset_alias_events = var.get("asset_alias_events", []) outlet_event_accessor = OutletEventAccessor( extra=var["extra"], raw_key=BaseSerialization.deserialize(var["raw_key"]), - dataset_alias_events=dataset_alias_events, + asset_alias_events=asset_alias_events, ) return outlet_event_accessor @@ -482,7 +477,7 @@ def deref(self, dag: DAG) -> ExpandInput: DagRun: DagRunPydantic, DagModel: DagModelPydantic, LogTemplate: LogTemplatePydantic, - Dataset: DatasetPydantic, + Asset: AssetPydantic, Trigger: TriggerPydantic, } _type_to_class: dict[DAT | str, list] = { @@ -491,7 +486,7 @@ def deref(self, dag: DAG) -> ExpandInput: DAT.DAG_RUN: [DagRunPydantic, DagRun], DAT.DAG_MODEL: [DagModelPydantic, DagModel], DAT.LOG_TEMPLATE: [LogTemplatePydantic, LogTemplate], - DAT.DATA_SET: [DatasetPydantic, Dataset], + DAT.ASSET: [AssetPydantic, Asset], DAT.TRIGGER: [TriggerPydantic, Trigger], } _class_to_type = {cls_: type_ for type_, classes in _type_to_class.items() for cls_ in classes} @@ -661,12 +656,12 @@ def serialize( elif isinstance(var, OutletEventAccessors): return cls._encode( cls.serialize(var._dict, strict=strict, use_pydantic_models=use_pydantic_models), # type: ignore[attr-defined] - type_=DAT.DATASET_EVENT_ACCESSORS, + type_=DAT.ASSET_EVENT_ACCESSORS, ) elif isinstance(var, OutletEventAccessor): return cls._encode( encode_outlet_event_accessor(var), - type_=DAT.DATASET_EVENT_ACCESSOR, + type_=DAT.ASSET_EVENT_ACCESSOR, ) elif isinstance(var, DAG): return cls._encode(SerializedDAG.serialize_dag(var), type_=DAT.DAG) @@ -744,8 +739,8 @@ def serialize( return cls._encode(serialize_xcom_arg(var), type_=DAT.XCOM_REF) elif isinstance(var, LazySelectSequence): return cls.serialize(list(var)) - elif isinstance(var, BaseDataset): - serialized_dataset = encode_dataset_condition(var) + elif isinstance(var, BaseAsset): + serialized_dataset = encode_asset_condition(var) return cls._encode(serialized_dataset, type_=serialized_dataset.pop("__type")) elif isinstance(var, SimpleTaskInstance): return cls._encode( @@ -826,11 +821,11 @@ def deserialize(cls, encoded_var: Any, use_pydantic_models=False) -> Any: return Context(**d) elif type_ == DAT.DICT: return {k: cls.deserialize(v, use_pydantic_models) for k, v in var.items()} - elif type_ == DAT.DATASET_EVENT_ACCESSORS: + elif type_ == DAT.ASSET_EVENT_ACCESSORS: d = OutletEventAccessors() # type: ignore[assignment] d._dict = cls.deserialize(var) # type: ignore[attr-defined] return d - elif type_ == DAT.DATASET_EVENT_ACCESSOR: + elif type_ == DAT.ASSET_EVENT_ACCESSOR: return decode_outlet_event_accessor(var) elif type_ == DAT.DAG: return SerializedDAG.deserialize_dag(var) @@ -872,14 +867,14 @@ def deserialize(cls, encoded_var: Any, use_pydantic_models=False) -> Any: return cls._deserialize_param(var) elif type_ == DAT.XCOM_REF: return _XComRef(var) # Delay deserializing XComArg objects until we have the entire DAG. - elif type_ == DAT.DATASET: - return Dataset(**var) - elif type_ == DAT.DATASET_ALIAS: - return DatasetAlias(**var) - elif type_ == DAT.DATASET_ANY: - return DatasetAny(*(decode_dataset_condition(x) for x in var["objects"])) - elif type_ == DAT.DATASET_ALL: - return DatasetAll(*(decode_dataset_condition(x) for x in var["objects"])) + elif type_ == DAT.ASSET: + return Asset(**var) + elif type_ == DAT.ASSET_ALIAS: + return AssetAlias(**var) + elif type_ == DAT.ASSET_ANY: + return AssetAny(*(decode_asset_condition(x) for x in var["objects"])) + elif type_ == DAT.ASSET_ALL: + return AssetAll(*(decode_asset_condition(x) for x in var["objects"])) elif type_ == DAT.SIMPLE_TASK_INSTANCE: return SimpleTaskInstance(**cls.deserialize(var)) elif type_ == DAT.CONNECTION: @@ -1041,17 +1036,17 @@ def detect_task_dependencies(task: Operator) -> list[DagDependency]: ) ) for obj in task.outlets or []: - if isinstance(obj, Dataset): + if isinstance(obj, Asset): deps.append( DagDependency( source=task.dag_id, - target="dataset", - dependency_type="dataset", + target="asset", + dependency_type="asset", dependency_id=obj.uri, ) ) - elif isinstance(obj, DatasetAlias): - cond = _DatasetAliasCondition(obj.name) + elif isinstance(obj, AssetAlias): + cond = _AssetAliasCondition(obj.name) deps.extend(cond.iter_dag_dependencies(source=task.dag_id, target="")) return deps @@ -1062,7 +1057,7 @@ def detect_dag_dependencies(dag: DAG | None) -> Iterable[DagDependency]: if not dag: return - yield from dag.timetable.dataset_condition.iter_dag_dependencies(source="", target=dag.dag_id) + yield from dag.timetable.asset_condition.iter_dag_dependencies(source="", target=dag.dag_id) class SerializedBaseOperator(BaseOperator, BaseSerialization): diff --git a/airflow/timetables/datasets.py b/airflow/timetables/assets.py similarity index 71% rename from airflow/timetables/datasets.py rename to airflow/timetables/assets.py index 05db0d66cc2df..b158555590ad5 100644 --- a/airflow/timetables/datasets.py +++ b/airflow/timetables/assets.py @@ -19,9 +19,9 @@ import typing -from airflow.datasets import BaseDataset, DatasetAll +from airflow.assets import AssetAll, BaseAsset from airflow.exceptions import AirflowTimetableInvalid -from airflow.timetables.simple import DatasetTriggeredTimetable as DatasetTriggeredSchedule +from airflow.timetables.simple import AssetTriggeredTimetable from airflow.utils.types import DagRunType if typing.TYPE_CHECKING: @@ -29,56 +29,56 @@ import pendulum - from airflow.datasets import Dataset + from airflow.assets import Asset from airflow.timetables.base import DagRunInfo, DataInterval, TimeRestriction, Timetable -class DatasetOrTimeSchedule(DatasetTriggeredSchedule): +class AssetOrTimeSchedule(AssetTriggeredTimetable): """Combine time-based scheduling with event-based scheduling.""" def __init__( self, *, timetable: Timetable, - datasets: Collection[Dataset] | BaseDataset, + assets: Collection[Asset] | BaseAsset, ) -> None: self.timetable = timetable - if isinstance(datasets, BaseDataset): - self.dataset_condition = datasets + if isinstance(assets, BaseAsset): + self.asset_condition = assets else: - self.dataset_condition = DatasetAll(*datasets) + self.asset_condition = AssetAll(*assets) - self.description = f"Triggered by datasets or {timetable.description}" + self.description = f"Triggered by assets or {timetable.description}" self.periodic = timetable.periodic self.can_be_scheduled = timetable.can_be_scheduled self.active_runs_limit = timetable.active_runs_limit @classmethod def deserialize(cls, data: dict[str, typing.Any]) -> Timetable: - from airflow.serialization.serialized_objects import decode_dataset_condition, decode_timetable + from airflow.serialization.serialized_objects import decode_asset_condition, decode_timetable return cls( - datasets=decode_dataset_condition(data["dataset_condition"]), + assets=decode_asset_condition(data["asset_condition"]), timetable=decode_timetable(data["timetable"]), ) def serialize(self) -> dict[str, typing.Any]: - from airflow.serialization.serialized_objects import encode_dataset_condition, encode_timetable + from airflow.serialization.serialized_objects import encode_asset_condition, encode_timetable return { - "dataset_condition": encode_dataset_condition(self.dataset_condition), + "asset_condition": encode_asset_condition(self.asset_condition), "timetable": encode_timetable(self.timetable), } def validate(self) -> None: - if isinstance(self.timetable, DatasetTriggeredSchedule): - raise AirflowTimetableInvalid("cannot nest dataset timetables") - if not isinstance(self.dataset_condition, BaseDataset): - raise AirflowTimetableInvalid("all elements in 'datasets' must be datasets") + if isinstance(self.timetable, AssetTriggeredTimetable): + raise AirflowTimetableInvalid("cannot nest asset timetables") + if not isinstance(self.asset_condition, BaseAsset): + raise AirflowTimetableInvalid("all elements in 'assets' must be assets") @property def summary(self) -> str: - return f"Dataset or {self.timetable.summary}" + return f"Asset or {self.timetable.summary}" def infer_manual_data_interval(self, *, run_after: pendulum.DateTime) -> DataInterval: return self.timetable.infer_manual_data_interval(run_after=run_after) diff --git a/airflow/timetables/base.py b/airflow/timetables/base.py index 5d97591856b5a..64a2612026517 100644 --- a/airflow/timetables/base.py +++ b/airflow/timetables/base.py @@ -18,20 +18,20 @@ from typing import TYPE_CHECKING, Any, Iterator, NamedTuple, Sequence -from airflow.datasets import BaseDataset +from airflow.assets import BaseAsset from airflow.typing_compat import Protocol, runtime_checkable if TYPE_CHECKING: from pendulum import DateTime - from airflow.datasets import Dataset, DatasetAlias + from airflow.assets import Asset, AssetAlias from airflow.serialization.dag_dependency import DagDependency from airflow.utils.types import DagRunType -class _NullDataset(BaseDataset): +class _NullAsset(BaseAsset): """ - Sentinel type that represents "no datasets". + Sentinel type that represents "no assets". This is only implemented to make typing easier in timetables, and not expected to be used anywhere else. @@ -42,10 +42,10 @@ class _NullDataset(BaseDataset): def __bool__(self) -> bool: return False - def __or__(self, other: BaseDataset) -> BaseDataset: + def __or__(self, other: BaseAsset) -> BaseAsset: return NotImplemented - def __and__(self, other: BaseDataset) -> BaseDataset: + def __and__(self, other: BaseAsset) -> BaseAsset: return NotImplemented def as_expression(self) -> Any: @@ -54,10 +54,10 @@ def as_expression(self) -> Any: def evaluate(self, statuses: dict[str, bool]) -> bool: return False - def iter_datasets(self) -> Iterator[tuple[str, Dataset]]: + def iter_assets(self) -> Iterator[tuple[str, Asset]]: return iter(()) - def iter_dataset_aliases(self) -> Iterator[tuple[str, DatasetAlias]]: + def iter_asset_aliases(self) -> Iterator[tuple[str, AssetAlias]]: return iter(()) def iter_dag_dependencies(self, source, target) -> Iterator[DagDependency]: @@ -189,11 +189,11 @@ class Timetable(Protocol): as for :class:`~airflow.timetable.simple.ContinuousTimetable`. """ - dataset_condition: BaseDataset = _NullDataset() - """The dataset condition that triggers a DAG using this timetable. + asset_condition: BaseAsset = _NullAsset() + """The asset condition that triggers a DAG using this timetable. - If this is not *None*, this should be a dataset, or a combination of, that - controls the DAG's dataset triggers. + If this is not *None*, this should be an asset, or a combination of, that + controls the DAG's asset triggers. """ @classmethod diff --git a/airflow/timetables/simple.py b/airflow/timetables/simple.py index ad166a641378a..5a931b40dd11d 100644 --- a/airflow/timetables/simple.py +++ b/airflow/timetables/simple.py @@ -18,7 +18,7 @@ from typing import TYPE_CHECKING, Any, Collection, Sequence -from airflow.datasets import DatasetAlias, _DatasetAliasCondition +from airflow.assets import AssetAlias, _AssetAliasCondition from airflow.timetables.base import DagRunInfo, DataInterval, Timetable from airflow.utils import timezone @@ -26,8 +26,8 @@ from pendulum import DateTime from sqlalchemy import Session - from airflow.datasets import BaseDataset - from airflow.models.dataset import DatasetEvent + from airflow.assets import BaseAsset + from airflow.models.asset import AssetEvent from airflow.timetables.base import TimeRestriction from airflow.utils.types import DagRunType @@ -152,44 +152,44 @@ def next_dagrun_info( return DagRunInfo.interval(start, end) -class DatasetTriggeredTimetable(_TrivialTimetable): +class AssetTriggeredTimetable(_TrivialTimetable): """ Timetable that never schedules anything. - This should not be directly used anywhere, but only set if a DAG is triggered by datasets. + This should not be directly used anywhere, but only set if a DAG is triggered by assets. :meta private: """ - UNRESOLVED_ALIAS_SUMMARY = "Unresolved DatasetAlias" + UNRESOLVED_ALIAS_SUMMARY = "Unresolved AssetAlias" - description: str = "Triggered by datasets" + description: str = "Triggered by assets" - def __init__(self, datasets: BaseDataset) -> None: + def __init__(self, assets: BaseAsset) -> None: super().__init__() - self.dataset_condition = datasets - if isinstance(self.dataset_condition, DatasetAlias): - self.dataset_condition = _DatasetAliasCondition(self.dataset_condition.name) + self.asset_condition = assets + if isinstance(self.asset_condition, AssetAlias): + self.asset_condition = _AssetAliasCondition(self.asset_condition.name) - if not next(self.dataset_condition.iter_datasets(), False): - self._summary = DatasetTriggeredTimetable.UNRESOLVED_ALIAS_SUMMARY + if not next(self.asset_condition.iter_assets(), False): + self._summary = AssetTriggeredTimetable.UNRESOLVED_ALIAS_SUMMARY else: - self._summary = "Dataset" + self._summary = "Asset" @classmethod def deserialize(cls, data: dict[str, Any]) -> Timetable: - from airflow.serialization.serialized_objects import decode_dataset_condition + from airflow.serialization.serialized_objects import decode_asset_condition - return cls(decode_dataset_condition(data["dataset_condition"])) + return cls(decode_asset_condition(data["asset_condition"])) @property def summary(self) -> str: return self._summary def serialize(self) -> dict[str, Any]: - from airflow.serialization.serialized_objects import encode_dataset_condition + from airflow.serialization.serialized_objects import encode_asset_condition - return {"dataset_condition": encode_dataset_condition(self.dataset_condition)} + return {"asset_condition": encode_asset_condition(self.asset_condition)} def generate_run_id( self, @@ -198,7 +198,7 @@ def generate_run_id( logical_date: DateTime, data_interval: DataInterval | None, session: Session | None = None, - events: Collection[DatasetEvent] | None = None, + events: Collection[AssetEvent] | None = None, **extra, ) -> str: from airflow.models.dagrun import DagRun @@ -208,7 +208,7 @@ def generate_run_id( def data_interval_for_events( self, logical_date: DateTime, - events: Collection[DatasetEvent], + events: Collection[AssetEvent], ) -> DataInterval: if not events: return DataInterval(logical_date, logical_date) diff --git a/airflow/ui/openapi-gen/queries/common.ts b/airflow/ui/openapi-gen/queries/common.ts index b8021fed9be3c..46694939ed74e 100644 --- a/airflow/ui/openapi-gen/queries/common.ts +++ b/airflow/ui/openapi-gen/queries/common.ts @@ -1,20 +1,20 @@ // generated with @7nohe/openapi-react-query-codegen@1.6.0 import { UseQueryResult } from "@tanstack/react-query"; -import { DagService, DatasetService } from "../requests/services.gen"; +import { AssetService, DagService } from "../requests/services.gen"; import { DagRunState } from "../requests/types.gen"; -export type DatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGetDefaultResponse = +export type AssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetDefaultResponse = Awaited< - ReturnType + ReturnType >; -export type DatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGetQueryResult< - TData = DatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGetDefaultResponse, +export type AssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetQueryResult< + TData = AssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetDefaultResponse, TError = unknown, > = UseQueryResult; -export const useDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGetKey = - "DatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGet"; -export const UseDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGetKeyFn = ( +export const useAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetKey = + "AssetServiceNextRunAssetsUiNextRunDatasetsDagIdGet"; +export const UseAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetKeyFn = ( { dagId, }: { @@ -22,7 +22,7 @@ export const UseDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGetKeyFn = ( }, queryKey?: Array, ) => [ - useDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGetKey, + useAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetKey, ...(queryKey ?? [{ dagId }]), ]; export type DagServiceGetDagsPublicDagsGetDefaultResponse = Awaited< diff --git a/airflow/ui/openapi-gen/queries/prefetch.ts b/airflow/ui/openapi-gen/queries/prefetch.ts index 6dd99f96b8425..7de7282a9bd01 100644 --- a/airflow/ui/openapi-gen/queries/prefetch.ts +++ b/airflow/ui/openapi-gen/queries/prefetch.ts @@ -1,34 +1,32 @@ // generated with @7nohe/openapi-react-query-codegen@1.6.0 import { type QueryClient } from "@tanstack/react-query"; -import { DagService, DatasetService } from "../requests/services.gen"; +import { AssetService, DagService } from "../requests/services.gen"; import { DagRunState } from "../requests/types.gen"; import * as Common from "./common"; /** - * Next Run Datasets + * Next Run Assets * @param data The data for the request. * @param data.dagId * @returns unknown Successful Response * @throws ApiError */ -export const prefetchUseDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGet = - ( - queryClient: QueryClient, - { - dagId, - }: { - dagId: string; - }, - ) => - queryClient.prefetchQuery({ - queryKey: - Common.UseDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGetKeyFn({ - dagId, - }), - queryFn: () => - DatasetService.nextRunDatasetsUiNextRunDatasetsDagIdGet({ dagId }), - }); +export const prefetchUseAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGet = ( + queryClient: QueryClient, + { + dagId, + }: { + dagId: string; + }, +) => + queryClient.prefetchQuery({ + queryKey: Common.UseAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetKeyFn( + { dagId }, + ), + queryFn: () => + AssetService.nextRunAssetsUiNextRunDatasetsDagIdGet({ dagId }), + }); /** * Get Dags * Get all DAGs. diff --git a/airflow/ui/openapi-gen/queries/queries.ts b/airflow/ui/openapi-gen/queries/queries.ts index b771fccfeb947..7cbaac5b2c77d 100644 --- a/airflow/ui/openapi-gen/queries/queries.ts +++ b/airflow/ui/openapi-gen/queries/queries.ts @@ -6,19 +6,19 @@ import { UseQueryOptions, } from "@tanstack/react-query"; -import { DagService, DatasetService } from "../requests/services.gen"; +import { AssetService, DagService } from "../requests/services.gen"; import { DAGPatchBody, DagRunState } from "../requests/types.gen"; import * as Common from "./common"; /** - * Next Run Datasets + * Next Run Assets * @param data The data for the request. * @param data.dagId * @returns unknown Successful Response * @throws ApiError */ -export const useDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGet = < - TData = Common.DatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGetDefaultResponse, +export const useAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGet = < + TData = Common.AssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetDefaultResponse, TError = unknown, TQueryKey extends Array = unknown[], >( @@ -31,15 +31,12 @@ export const useDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGet = < options?: Omit, "queryKey" | "queryFn">, ) => useQuery({ - queryKey: - Common.UseDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGetKeyFn( - { dagId }, - queryKey, - ), + queryKey: Common.UseAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetKeyFn( + { dagId }, + queryKey, + ), queryFn: () => - DatasetService.nextRunDatasetsUiNextRunDatasetsDagIdGet({ - dagId, - }) as TData, + AssetService.nextRunAssetsUiNextRunDatasetsDagIdGet({ dagId }) as TData, ...options, }); /** diff --git a/airflow/ui/openapi-gen/queries/suspense.ts b/airflow/ui/openapi-gen/queries/suspense.ts index 7743ce92d2855..18dba7acb4b5b 100644 --- a/airflow/ui/openapi-gen/queries/suspense.ts +++ b/airflow/ui/openapi-gen/queries/suspense.ts @@ -1,43 +1,39 @@ // generated with @7nohe/openapi-react-query-codegen@1.6.0 import { UseQueryOptions, useSuspenseQuery } from "@tanstack/react-query"; -import { DagService, DatasetService } from "../requests/services.gen"; +import { AssetService, DagService } from "../requests/services.gen"; import { DagRunState } from "../requests/types.gen"; import * as Common from "./common"; /** - * Next Run Datasets + * Next Run Assets * @param data The data for the request. * @param data.dagId * @returns unknown Successful Response * @throws ApiError */ -export const useDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGetSuspense = - < - TData = Common.DatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGetDefaultResponse, - TError = unknown, - TQueryKey extends Array = unknown[], - >( - { - dagId, - }: { - dagId: string; - }, - queryKey?: TQueryKey, - options?: Omit, "queryKey" | "queryFn">, - ) => - useSuspenseQuery({ - queryKey: - Common.UseDatasetServiceNextRunDatasetsUiNextRunDatasetsDagIdGetKeyFn( - { dagId }, - queryKey, - ), - queryFn: () => - DatasetService.nextRunDatasetsUiNextRunDatasetsDagIdGet({ - dagId, - }) as TData, - ...options, - }); +export const useAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetSuspense = < + TData = Common.AssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetDefaultResponse, + TError = unknown, + TQueryKey extends Array = unknown[], +>( + { + dagId, + }: { + dagId: string; + }, + queryKey?: TQueryKey, + options?: Omit, "queryKey" | "queryFn">, +) => + useSuspenseQuery({ + queryKey: Common.UseAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetKeyFn( + { dagId }, + queryKey, + ), + queryFn: () => + AssetService.nextRunAssetsUiNextRunDatasetsDagIdGet({ dagId }) as TData, + ...options, + }); /** * Get Dags * Get all DAGs. diff --git a/airflow/ui/openapi-gen/requests/services.gen.ts b/airflow/ui/openapi-gen/requests/services.gen.ts index 37a4d11873acf..5aa5876d112ad 100644 --- a/airflow/ui/openapi-gen/requests/services.gen.ts +++ b/airflow/ui/openapi-gen/requests/services.gen.ts @@ -3,25 +3,25 @@ import type { CancelablePromise } from "./core/CancelablePromise"; import { OpenAPI } from "./core/OpenAPI"; import { request as __request } from "./core/request"; import type { - NextRunDatasetsUiNextRunDatasetsDagIdGetData, - NextRunDatasetsUiNextRunDatasetsDagIdGetResponse, + NextRunAssetsUiNextRunDatasetsDagIdGetData, + NextRunAssetsUiNextRunDatasetsDagIdGetResponse, GetDagsPublicDagsGetData, GetDagsPublicDagsGetResponse, PatchDagPublicDagsDagIdPatchData, PatchDagPublicDagsDagIdPatchResponse, } from "./types.gen"; -export class DatasetService { +export class AssetService { /** - * Next Run Datasets + * Next Run Assets * @param data The data for the request. * @param data.dagId * @returns unknown Successful Response * @throws ApiError */ - public static nextRunDatasetsUiNextRunDatasetsDagIdGet( - data: NextRunDatasetsUiNextRunDatasetsDagIdGetData, - ): CancelablePromise { + public static nextRunAssetsUiNextRunDatasetsDagIdGet( + data: NextRunAssetsUiNextRunDatasetsDagIdGetData, + ): CancelablePromise { return __request(OpenAPI, { method: "GET", url: "/ui/next_run_datasets/{dag_id}", diff --git a/airflow/ui/openapi-gen/requests/types.gen.ts b/airflow/ui/openapi-gen/requests/types.gen.ts index 16977004e79d6..bc455f63b6449 100644 --- a/airflow/ui/openapi-gen/requests/types.gen.ts +++ b/airflow/ui/openapi-gen/requests/types.gen.ts @@ -88,11 +88,11 @@ export type ValidationError = { type: string; }; -export type NextRunDatasetsUiNextRunDatasetsDagIdGetData = { +export type NextRunAssetsUiNextRunDatasetsDagIdGetData = { dagId: string; }; -export type NextRunDatasetsUiNextRunDatasetsDagIdGetResponse = { +export type NextRunAssetsUiNextRunDatasetsDagIdGetResponse = { [key: string]: unknown; }; @@ -122,7 +122,7 @@ export type PatchDagPublicDagsDagIdPatchResponse = DAGResponse; export type $OpenApiTs = { "/ui/next_run_datasets/{dag_id}": { get: { - req: NextRunDatasetsUiNextRunDatasetsDagIdGetData; + req: NextRunAssetsUiNextRunDatasetsDagIdGetData; res: { /** * Successful Response diff --git a/airflow/utils/context.py b/airflow/utils/context.py index a72885401f7b2..e5d30b1e2d7d2 100644 --- a/airflow/utils/context.py +++ b/airflow/utils/context.py @@ -40,14 +40,14 @@ import lazy_object_proxy from sqlalchemy import select -from airflow.datasets import ( - Dataset, - DatasetAlias, - DatasetAliasEvent, +from airflow.assets import ( + Asset, + AssetAlias, + AssetAliasEvent, extract_event_key, ) from airflow.exceptions import RemovedInAirflow3Warning -from airflow.models.dataset import DatasetAliasModel, DatasetEvent, DatasetModel +from airflow.models.asset import AssetAliasModel, AssetEvent, AssetModel from airflow.utils.db import LazySelectSequence from airflow.utils.types import NOTSET @@ -102,7 +102,7 @@ "ti", "tomorrow_ds", "tomorrow_ds_nodash", - "triggering_dataset_events", + "triggering_asset_events", "ts", "ts_nodash", "ts_nodash_with_tz", @@ -165,40 +165,40 @@ def get(self, key: str, default_conn: Any = None) -> Any: @attrs.define() class OutletEventAccessor: """ - Wrapper to access an outlet dataset event in template. + Wrapper to access an outlet asset event in template. :meta private: """ - raw_key: str | Dataset | DatasetAlias + raw_key: str | Asset | AssetAlias extra: dict[str, Any] = attrs.Factory(dict) - dataset_alias_events: list[DatasetAliasEvent] = attrs.field(factory=list) - - def add(self, dataset: Dataset | str, extra: dict[str, Any] | None = None) -> None: - """Add a DatasetEvent to an existing Dataset.""" - if isinstance(dataset, str): - dataset_uri = dataset - elif isinstance(dataset, Dataset): - dataset_uri = dataset.uri + asset_alias_events: list[AssetAliasEvent] = attrs.field(factory=list) + + def add(self, asset: Asset | str, extra: dict[str, Any] | None = None) -> None: + """Add an AssetEvent to an existing Asset.""" + if isinstance(asset, str): + asset_uri = asset + elif isinstance(asset, Asset): + asset_uri = asset.uri else: return if isinstance(self.raw_key, str): - dataset_alias_name = self.raw_key - elif isinstance(self.raw_key, DatasetAlias): - dataset_alias_name = self.raw_key.name + asset_alias_name = self.raw_key + elif isinstance(self.raw_key, AssetAlias): + asset_alias_name = self.raw_key.name else: return - event = DatasetAliasEvent( - source_alias_name=dataset_alias_name, dest_dataset_uri=dataset_uri, extra=extra or {} + event = AssetAliasEvent( + source_alias_name=asset_alias_name, dest_asset_uri=asset_uri, extra=extra or {} ) - self.dataset_alias_events.append(event) + self.asset_alias_events.append(event) class OutletEventAccessors(Mapping[str, OutletEventAccessor]): """ - Lazy mapping of outlet dataset event accessors. + Lazy mapping of outlet asset event accessors. :meta private: """ @@ -215,53 +215,53 @@ def __iter__(self) -> Iterator[str]: def __len__(self) -> int: return len(self._dict) - def __getitem__(self, key: str | Dataset | DatasetAlias) -> OutletEventAccessor: + def __getitem__(self, key: str | Asset | AssetAlias) -> OutletEventAccessor: event_key = extract_event_key(key) if event_key not in self._dict: self._dict[event_key] = OutletEventAccessor(extra={}, raw_key=key) return self._dict[event_key] -class LazyDatasetEventSelectSequence(LazySelectSequence[DatasetEvent]): +class LazyAssetEventSelectSequence(LazySelectSequence[AssetEvent]): """ - List-like interface to lazily access DatasetEvent rows. + List-like interface to lazily access AssetEvent rows. :meta private: """ @staticmethod def _rebuild_select(stmt: TextClause) -> Select: - return select(DatasetEvent).from_statement(stmt) + return select(AssetEvent).from_statement(stmt) @staticmethod - def _process_row(row: Row) -> DatasetEvent: + def _process_row(row: Row) -> AssetEvent: return row[0] @attrs.define(init=False) -class InletEventsAccessors(Mapping[str, LazyDatasetEventSelectSequence]): +class InletEventsAccessors(Mapping[str, LazyAssetEventSelectSequence]): """ - Lazy mapping for inlet dataset events accessors. + Lazy mapping for inlet asset events accessors. :meta private: """ _inlets: list[Any] - _datasets: dict[str, Dataset] - _dataset_aliases: dict[str, DatasetAlias] + _assets: dict[str, Asset] + _asset_aliases: dict[str, AssetAlias] _session: Session def __init__(self, inlets: list, *, session: Session) -> None: self._inlets = inlets self._session = session - self._datasets = {} - self._dataset_aliases = {} + self._assets = {} + self._asset_aliases = {} for inlet in inlets: - if isinstance(inlet, Dataset): - self._datasets[inlet.uri] = inlet - elif isinstance(inlet, DatasetAlias): - self._dataset_aliases[inlet.name] = inlet + if isinstance(inlet, Asset): + self._assets[inlet.uri] = inlet + elif isinstance(inlet, AssetAlias): + self._asset_aliases[inlet.name] = inlet def __iter__(self) -> Iterator[str]: return iter(self._inlets) @@ -269,28 +269,28 @@ def __iter__(self) -> Iterator[str]: def __len__(self) -> int: return len(self._inlets) - def __getitem__(self, key: int | str | Dataset | DatasetAlias) -> LazyDatasetEventSelectSequence: + def __getitem__(self, key: int | str | Asset | AssetAlias) -> LazyAssetEventSelectSequence: if isinstance(key, int): # Support index access; it's easier for trivial cases. obj = self._inlets[key] - if not isinstance(obj, (Dataset, DatasetAlias)): + if not isinstance(obj, (Asset, AssetAlias)): raise IndexError(key) else: obj = key - if isinstance(obj, DatasetAlias): - dataset_alias = self._dataset_aliases[obj.name] - join_clause = DatasetEvent.source_aliases - where_clause = DatasetAliasModel.name == dataset_alias.name - elif isinstance(obj, (Dataset, str)): - dataset = self._datasets[extract_event_key(obj)] - join_clause = DatasetEvent.dataset - where_clause = DatasetModel.uri == dataset.uri + if isinstance(obj, AssetAlias): + asset_alias = self._asset_aliases[obj.name] + join_clause = AssetEvent.source_aliases + where_clause = AssetAliasModel.name == asset_alias.name + elif isinstance(obj, (Asset, str)): + asset = self._assets[extract_event_key(obj)] + join_clause = AssetEvent.dataset + where_clause = AssetModel.uri == asset.uri else: raise ValueError(key) - return LazyDatasetEventSelectSequence.from_select( - select(DatasetEvent).join(join_clause).where(where_clause), - order_by=[DatasetEvent.timestamp], + return LazyAssetEventSelectSequence.from_select( + select(AssetEvent).join(join_clause).where(where_clause), + order_by=[AssetEvent.timestamp], session=self._session, ) diff --git a/airflow/utils/context.pyi b/airflow/utils/context.pyi index 658aac5839ec5..4dc4659548ac0 100644 --- a/airflow/utils/context.pyi +++ b/airflow/utils/context.pyi @@ -31,16 +31,16 @@ from typing import Any, Collection, Container, Iterable, Iterator, Mapping, Sequ from pendulum import DateTime from sqlalchemy.orm import Session +from airflow.assets import Asset, AssetAlias, AssetAliasEvent from airflow.configuration import AirflowConfigParser -from airflow.datasets import Dataset, DatasetAlias, DatasetAliasEvent +from airflow.models.asset import AssetEvent from airflow.models.baseoperator import BaseOperator from airflow.models.dag import DAG from airflow.models.dagrun import DagRun -from airflow.models.dataset import DatasetEvent from airflow.models.param import ParamsDict from airflow.models.taskinstance import TaskInstance +from airflow.serialization.pydantic.asset import AssetEventPydantic from airflow.serialization.pydantic.dag_run import DagRunPydantic -from airflow.serialization.pydantic.dataset import DatasetEventPydantic from airflow.serialization.pydantic.taskinstance import TaskInstancePydantic from airflow.typing_compat import TypedDict @@ -62,31 +62,31 @@ class OutletEventAccessor: self, *, extra: dict[str, Any], - raw_key: str | Dataset | DatasetAlias, - dataset_alias_events: list[DatasetAliasEvent], + raw_key: str | Asset | AssetAlias, + asset_alias_events: list[AssetAliasEvent], ) -> None: ... - def add(self, dataset: Dataset | str, extra: dict[str, Any] | None = None) -> None: ... + def add(self, asset: Asset | str, extra: dict[str, Any] | None = None) -> None: ... extra: dict[str, Any] - raw_key: str | Dataset | DatasetAlias - dataset_alias_events: list[DatasetAliasEvent] + raw_key: str | Asset | AssetAlias + asset_alias_events: list[AssetAliasEvent] class OutletEventAccessors(Mapping[str, OutletEventAccessor]): def __iter__(self) -> Iterator[str]: ... def __len__(self) -> int: ... - def __getitem__(self, key: str | Dataset | DatasetAlias) -> OutletEventAccessor: ... + def __getitem__(self, key: str | Asset | AssetAlias) -> OutletEventAccessor: ... -class InletEventsAccessor(Sequence[DatasetEvent]): +class InletEventsAccessor(Sequence[AssetEvent]): @overload - def __getitem__(self, key: int) -> DatasetEvent: ... + def __getitem__(self, key: int) -> AssetEvent: ... @overload - def __getitem__(self, key: slice) -> Sequence[DatasetEvent]: ... + def __getitem__(self, key: slice) -> Sequence[AssetEvent]: ... def __len__(self) -> int: ... class InletEventsAccessors(Mapping[str, InletEventsAccessor]): def __init__(self, inlets: list, *, session: Session) -> None: ... def __iter__(self) -> Iterator[str]: ... def __len__(self) -> int: ... - def __getitem__(self, key: int | str | Dataset | DatasetAlias) -> InletEventsAccessor: ... + def __getitem__(self, key: int | str | Asset | AssetAlias) -> InletEventsAccessor: ... # NOTE: Please keep this in sync with the following: # * KNOWN_CONTEXT_KEYS in airflow/utils/context.py @@ -132,7 +132,7 @@ class Context(TypedDict, total=False): ti: TaskInstance | TaskInstancePydantic tomorrow_ds: str tomorrow_ds_nodash: str - triggering_dataset_events: Mapping[str, Collection[DatasetEvent | DatasetEventPydantic]] + triggering_asset_events: Mapping[str, Collection[AssetEvent | AssetEventPydantic]] ts: str ts_nodash: str ts_nodash_with_tz: str diff --git a/airflow/utils/operator_helpers.py b/airflow/utils/operator_helpers.py index 108a84a9eabb9..6e5e03d8d163b 100644 --- a/airflow/utils/operator_helpers.py +++ b/airflow/utils/operator_helpers.py @@ -251,7 +251,7 @@ def __init__( def run(self, *args, **kwargs) -> Any: import inspect - from airflow.datasets.metadata import Metadata + from airflow.assets.metadata import Metadata from airflow.utils.types import NOTSET if not inspect.isgeneratorfunction(self.func): diff --git a/airflow/www/auth.py b/airflow/www/auth.py index 47a06f52e94bc..74f31d135c1aa 100644 --- a/airflow/www/auth.py +++ b/airflow/www/auth.py @@ -262,9 +262,9 @@ def decorated(*args, **kwargs): return has_access_decorator -def has_access_dataset(method: ResourceMethod) -> Callable[[T], T]: - """Check current user's permissions against required permissions for datasets.""" - return _has_access_no_details(lambda: get_auth_manager().is_authorized_dataset(method=method)) +def has_access_asset(method: ResourceMethod) -> Callable[[T], T]: + """Check current user's permissions against required permissions for assets.""" + return _has_access_no_details(lambda: get_auth_manager().is_authorized_asset(method=method)) def has_access_pool(method: ResourceMethod) -> Callable[[T], T]: diff --git a/airflow/www/security_manager.py b/airflow/www/security_manager.py index 926148f7eba86..77fd653b5f416 100644 --- a/airflow/www/security_manager.py +++ b/airflow/www/security_manager.py @@ -40,6 +40,7 @@ from airflow.models import Connection, DagRun, Pool, TaskInstance, Variable from airflow.security.permissions import ( RESOURCE_ADMIN_MENU, + RESOURCE_ASSET, RESOURCE_AUDIT_LOG, RESOURCE_BROWSE_MENU, RESOURCE_CLUSTER_ACTIVITY, @@ -49,7 +50,6 @@ RESOURCE_DAG_CODE, RESOURCE_DAG_DEPENDENCIES, RESOURCE_DAG_RUN, - RESOURCE_DATASET, RESOURCE_DOCS, RESOURCE_DOCS_MENU, RESOURCE_JOB, @@ -253,7 +253,7 @@ def _is_authorized_dag(entity_=None, details_func_=None): details=ConnectionDetails(conn_id=get_connection_id(resource_pk)), user=user, ), - RESOURCE_DATASET: lambda action, resource_pk, user: auth_manager.is_authorized_dataset( + RESOURCE_ASSET: lambda action, resource_pk, user: auth_manager.is_authorized_asset( method=methods[action], user=user, ), diff --git a/airflow/www/static/css/graph.css b/airflow/www/static/css/graph.css index f175a7e025d78..16dc5186af14f 100644 --- a/airflow/www/static/css/graph.css +++ b/airflow/www/static/css/graph.css @@ -161,12 +161,12 @@ g.node text { background-color: #e6f1f2; } -.legend-item.dataset { +.legend-item.asset { float: left; background-color: #fcecd4; } -.legend-item.dataset-alias { +.legend-item.asset-alias { float: left; background-color: #e8cfe4; } @@ -183,10 +183,10 @@ g.node.sensor rect { fill: #e6f1f2; } -g.node.dataset rect { +g.node.asset rect { fill: #fcecd4; } -g.node.dataset-alias rect { +g.node.asset-alias rect { fill: #e8cfe4; } diff --git a/airflow/www/static/js/dag/details/graph/Node.tsx b/airflow/www/static/js/dag/details/graph/Node.tsx index a4e9dee4c8074..daedfb8524e0b 100644 --- a/airflow/www/static/js/dag/details/graph/Node.tsx +++ b/airflow/www/static/js/dag/details/graph/Node.tsx @@ -94,7 +94,7 @@ const Node = (props: NodeProps) => { ); } - if (data.class === "dataset") return ; + if (data.class === "asset") return ; return ; }; diff --git a/airflow/www/static/js/dag/details/graph/index.tsx b/airflow/www/static/js/dag/details/graph/index.tsx index 51dd20d88b105..edafc99fe9c34 100644 --- a/airflow/www/static/js/dag/details/graph/index.tsx +++ b/airflow/www/static/js/dag/details/graph/index.tsx @@ -105,7 +105,7 @@ const getUpstreamDatasets = ( nodes.push({ id: d, value: { - class: "dataset", + class: "asset", label: d, }, }); @@ -202,7 +202,7 @@ const Graph = ({ openGroupIds, onToggleGroups, hoveredTaskState }: Props) => { datasetNodes.push({ id: dataset.uri, value: { - class: "dataset", + class: "asset", label: dataset.uri, }, }); @@ -221,7 +221,7 @@ const Graph = ({ openGroupIds, onToggleGroups, hoveredTaskState }: Props) => { datasetNodes.push({ id: de.datasetUri, value: { - class: "dataset", + class: "asset", label: de.datasetUri, }, }); diff --git a/airflow/www/static/js/dag/details/graph/utils.ts b/airflow/www/static/js/dag/details/graph/utils.ts index 93c8b9c253016..2fb7351525e71 100644 --- a/airflow/www/static/js/dag/details/graph/utils.ts +++ b/airflow/www/static/js/dag/details/graph/utils.ts @@ -92,7 +92,7 @@ export const flattenNodes = ({ onToggleGroups(newGroupIds); }, datasetEvent: - node.value.class === "dataset" + node.value.class === "asset" ? datasetEvents?.find((de) => de.datasetUri === node.value.label) : undefined, ...node.value, diff --git a/airflow/www/static/js/datasets/Graph/Node.tsx b/airflow/www/static/js/datasets/Graph/Node.tsx index baef11aa83663..dfb0cf8ed4deb 100644 --- a/airflow/www/static/js/datasets/Graph/Node.tsx +++ b/airflow/www/static/js/datasets/Graph/Node.tsx @@ -70,10 +70,10 @@ const BaseNode = ({ justifyContent="space-between" alignItems="center" > - {type === "dataset" && } + {type === "asset" && } {type === "sensor" && } {type === "trigger" && } - {type === "dataset-alias" && } + {type === "asset-alias" && } {label} )} diff --git a/airflow/www/static/js/datasets/Graph/index.tsx b/airflow/www/static/js/datasets/Graph/index.tsx index 9157a8a1a2617..e960c48ff63a2 100644 --- a/airflow/www/static/js/datasets/Graph/index.tsx +++ b/airflow/www/static/js/datasets/Graph/index.tsx @@ -87,7 +87,7 @@ const Graph = ({ selectedNodeId, onSelect }: Props) => { height: c.height, onSelect: () => { if (onSelect) { - if (c.value.class === "dataset") onSelect({ uri: c.value.label }); + if (c.value.class === "asset") onSelect({ uri: c.value.label }); else if (c.value.class === "dag") onSelect({ dagId: c.value.label }); } diff --git a/airflow/www/static/js/datasets/SearchBar.tsx b/airflow/www/static/js/datasets/SearchBar.tsx index fc47215389a57..33476b8419fdb 100644 --- a/airflow/www/static/js/datasets/SearchBar.tsx +++ b/airflow/www/static/js/datasets/SearchBar.tsx @@ -46,7 +46,7 @@ const SearchBar = ({ (datasetDependencies?.nodes || []).forEach((node) => { if (node.value.class === "dag") dagOptions.push({ value: node.id, label: node.value.label }); - if (node.value.class === "dataset") + if (node.value.class === "asset") datasetOptions.push({ value: node.id, label: node.value.label }); }); diff --git a/airflow/www/static/js/types/index.ts b/airflow/www/static/js/types/index.ts index b390568c8cf3b..1ce07bb350795 100644 --- a/airflow/www/static/js/types/index.ts +++ b/airflow/www/static/js/types/index.ts @@ -135,12 +135,12 @@ interface DepNode { id?: string; class: | "dag" - | "dataset" + | "asset" | "trigger" | "sensor" | "or-gate" | "and-gate" - | "dataset-alias"; + | "asset-alias"; label: string; rx?: number; ry?: number; diff --git a/airflow/www/templates/airflow/dag.html b/airflow/www/templates/airflow/dag.html index b0c00bcd5c88f..0d3a2cf1770ca 100644 --- a/airflow/www/templates/airflow/dag.html +++ b/airflow/www/templates/airflow/dag.html @@ -149,29 +149,29 @@

- {% if ds_info.total == 1 -%} - On {{ ds_info.uri }} + {% if asset_info.total == 1 -%} + On {{ asset_info.uri }} {%- else -%} - {{ ds_info.ready }} of {{ ds_info.total }} datasets updated + {{ asset_info.ready }} of {{ asset_info.total }} datasets updated {%- endif %}

diff --git a/airflow/www/templates/airflow/dag_dependencies.html b/airflow/www/templates/airflow/dag_dependencies.html index 542a5b7b47b8f..393b8d0de694d 100644 --- a/airflow/www/templates/airflow/dag_dependencies.html +++ b/airflow/www/templates/airflow/dag_dependencies.html @@ -43,8 +43,8 @@

dag trigger sensor - dataset - dataset alias + asset + asset alias
Last refresh:
diff --git a/airflow/www/templates/airflow/dags.html b/airflow/www/templates/airflow/dags.html index ca374c665aea0..c629936df7c00 100644 --- a/airflow/www/templates/airflow/dags.html +++ b/airflow/www/templates/airflow/dags.html @@ -304,31 +304,31 @@

{{ page_title }}

info -
+ {table.getHeaderGroups().map((headerGroup) => ( {headerGroup.headers.map( diff --git a/airflow/ui/src/components/TogglePause.tsx b/airflow/ui/src/components/TogglePause.tsx new file mode 100644 index 0000000000000..50362187c8ad7 --- /dev/null +++ b/airflow/ui/src/components/TogglePause.tsx @@ -0,0 +1,56 @@ +/*! + * 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. + */ +import { Switch } from "@chakra-ui/react"; +import { useQueryClient } from "@tanstack/react-query"; +import { useCallback } from "react"; + +import { + useDagServiceGetDagsKey, + useDagServicePatchDag, +} from "openapi/queries"; + +type Props = { + readonly dagId: string; + readonly isPaused: boolean; +}; + +export const TogglePause = ({ dagId, isPaused }: Props) => { + const queryClient = useQueryClient(); + + const onSuccess = async () => { + await queryClient.invalidateQueries({ + queryKey: [useDagServiceGetDagsKey], + }); + }; + + const { mutate } = useDagServicePatchDag({ + onSuccess, + }); + + const onChange = useCallback(() => { + mutate({ + dagId, + requestBody: { + is_paused: !isPaused, + }, + }); + }, [dagId, isPaused, mutate]); + + return ; +}; diff --git a/airflow/ui/src/pages/DagsList/DagsFilters.tsx b/airflow/ui/src/pages/DagsList/DagsFilters.tsx new file mode 100644 index 0000000000000..cb2be8322e500 --- /dev/null +++ b/airflow/ui/src/pages/DagsList/DagsFilters.tsx @@ -0,0 +1,86 @@ +/*! + * 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. + */ +import { HStack, Select, Text, Box } from "@chakra-ui/react"; +import { Select as ReactSelect } from "chakra-react-select"; +import { useCallback } from "react"; +import { useSearchParams } from "react-router-dom"; + +import { useTableURLState } from "src/components/DataTable/useTableUrlState"; +import { QuickFilterButton } from "src/components/QuickFilterButton"; + +const PAUSED_PARAM = "paused"; + +export const DagsFilters = () => { + const [searchParams, setSearchParams] = useSearchParams(); + + const showPaused = searchParams.get(PAUSED_PARAM); + + const { setTableURLState, tableURLState } = useTableURLState(); + const { pagination, sorting } = tableURLState; + + const handlePausedChange: React.ChangeEventHandler = + useCallback( + ({ target: { value } }) => { + if (value === "All") { + searchParams.delete(PAUSED_PARAM); + } else { + searchParams.set(PAUSED_PARAM, value); + } + setSearchParams(searchParams); + setTableURLState({ + pagination: { ...pagination, pageIndex: 0 }, + sorting, + }); + }, + [pagination, searchParams, setSearchParams, setTableURLState, sorting], + ); + + return ( + + + + + State: + + + All + Failed + Running + Successful + + + + + Active: + + + + + + + ); +}; diff --git a/airflow/ui/src/pages/DagsList.tsx b/airflow/ui/src/pages/DagsList/DagsList.tsx similarity index 72% rename from airflow/ui/src/pages/DagsList.tsx rename to airflow/ui/src/pages/DagsList/DagsList.tsx index ab480d2cbabdb..d58e3eaa2038c 100644 --- a/airflow/ui/src/pages/DagsList.tsx +++ b/airflow/ui/src/pages/DagsList/DagsList.tsx @@ -18,7 +18,6 @@ */ import { Badge, - Checkbox, Heading, HStack, Select, @@ -26,30 +25,36 @@ import { VStack, } from "@chakra-ui/react"; import type { ColumnDef } from "@tanstack/react-table"; -import { Select as ReactSelect } from "chakra-react-select"; import { type ChangeEventHandler, useCallback } from "react"; import { useSearchParams } from "react-router-dom"; import { useDagServiceGetDags } from "openapi/queries"; import type { DAGResponse } from "openapi/requests/types.gen"; +import { DataTable } from "src/components/DataTable"; +import { useTableURLState } from "src/components/DataTable/useTableUrlState"; +import { SearchBar } from "src/components/SearchBar"; +import { TogglePause } from "src/components/TogglePause"; +import { pluralize } from "src/utils/pluralize"; -import { DataTable } from "../components/DataTable"; -import { useTableURLState } from "../components/DataTable/useTableUrlState"; -import { QuickFilterButton } from "../components/QuickFilterButton"; -import { SearchBar } from "../components/SearchBar"; -import { pluralize } from "../utils/pluralize"; +import { DagsFilters } from "./DagsFilters"; const columns: Array> = [ + { + accessorKey: "is_paused", + cell: ({ row }) => ( + + ), + enableSorting: false, + header: "", + }, { accessorKey: "dag_id", cell: ({ row }) => row.original.dag_display_name, header: "DAG", }, - { - accessorKey: "is_paused", - enableSorting: false, - header: () => "Is Paused", - }, { accessorKey: "timetable_description", cell: (info) => @@ -82,9 +87,9 @@ const PAUSED_PARAM = "paused"; // eslint-disable-next-line complexity export const DagsList = ({ cardView = false }) => { - const [searchParams, setSearchParams] = useSearchParams(); + const [searchParams] = useSearchParams(); - const showPaused = searchParams.get(PAUSED_PARAM) === "true"; + const showPaused = searchParams.get(PAUSED_PARAM); const { setTableURLState, tableURLState } = useTableURLState(); const { pagination, sorting } = tableURLState; @@ -98,22 +103,9 @@ export const DagsList = ({ cardView = false }) => { offset: pagination.pageIndex * pagination.pageSize, onlyActive: true, orderBy, - paused: showPaused, + paused: showPaused === null ? undefined : showPaused === "true", }); - const handlePausedChange = useCallback(() => { - searchParams[showPaused ? "delete" : "set"](PAUSED_PARAM, "true"); - setSearchParams(searchParams); - setTableURLState({ pagination: { ...pagination, pageIndex: 0 }, sorting }); - }, [ - pagination, - searchParams, - setSearchParams, - setTableURLState, - showPaused, - sorting, - ]); - const handleSortChange = useCallback>( ({ currentTarget: { value } }) => { setTableURLState({ @@ -136,20 +128,7 @@ export const DagsList = ({ cardView = false }) => { buttonProps={{ isDisabled: true }} inputProps={{ isDisabled: true }} /> - - - - All - Failed - Running - Successful - - - Show Paused DAGs - - - - + {pluralize("DAG", data?.total_entries)} diff --git a/airflow/ui/src/pages/DagsList/index.tsx b/airflow/ui/src/pages/DagsList/index.tsx new file mode 100644 index 0000000000000..df59a682abb67 --- /dev/null +++ b/airflow/ui/src/pages/DagsList/index.tsx @@ -0,0 +1,20 @@ +/*! + * 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. + */ + +export { DagsList } from "./DagsList"; From a28665a8b05532978fe4fa0e4669e9781b4f96f7 Mon Sep 17 00:00:00 2001 From: jonhspyro <121674572+jonhspyro@users.noreply.github.com> Date: Wed, 2 Oct 2024 11:36:27 +0100 Subject: [PATCH 107/802] Correctly select task in DAG Graph View when clicking on its name (#38782) * Fix in DAG Graph View, clicking Task on it's name doesn't select the task. (#37932) * Updated TaskName onClick * Fixed missing onToggleCollapse * Added missing changes * Updated: rebase * fixed providers error message * undo fab changes * Update user_command.py --------- Co-authored-by: Brent Bovenzi --- .../static/js/dag/details/graph/DagNode.test.tsx | 14 +++++++++++++- .../www/static/js/dag/details/graph/DagNode.tsx | 8 +++++--- 2 files changed, 18 insertions(+), 4 deletions(-) diff --git a/airflow/www/static/js/dag/details/graph/DagNode.test.tsx b/airflow/www/static/js/dag/details/graph/DagNode.test.tsx index 34ddac7506c71..7c6dea7584a0b 100644 --- a/airflow/www/static/js/dag/details/graph/DagNode.test.tsx +++ b/airflow/www/static/js/dag/details/graph/DagNode.test.tsx @@ -20,7 +20,7 @@ /* global describe, test, expect */ import React from "react"; -import { render } from "@testing-library/react"; +import { fireEvent, render } from "@testing-library/react"; import { Wrapper } from "src/utils/testUtils"; @@ -124,4 +124,16 @@ describe("Test Graph Node", () => { expect(getByTestId("node")).toHaveStyle("opacity: 0.3"); }); + + test("Clicks on taskName", async () => { + const { getByText } = render(, { + wrapper: Wrapper, + }); + + const taskName = getByText("task_id"); + + fireEvent.click(taskName); + + expect(taskName).toBeInTheDocument(); + }); }); diff --git a/airflow/www/static/js/dag/details/graph/DagNode.tsx b/airflow/www/static/js/dag/details/graph/DagNode.tsx index c2f9b01296c35..4ac1be8ef4ade 100644 --- a/airflow/www/static/js/dag/details/graph/DagNode.tsx +++ b/airflow/www/static/js/dag/details/graph/DagNode.tsx @@ -42,10 +42,10 @@ const DagNode = ({ task, isSelected, latestDagRunId, - onToggleCollapse, isOpen, isActive, setupTeardownType, + onToggleCollapse, labelStyle, style, isZoomedOut, @@ -139,8 +139,10 @@ const DagNode = ({ isOpen={isOpen} isGroup={!!childCount} onClick={(e) => { - e.stopPropagation(); - onToggleCollapse(); + if (childCount) { + e.stopPropagation(); + onToggleCollapse(); + } }} setupTeardownType={setupTeardownType} isZoomedOut={isZoomedOut} From e54a2820db70a4f248e8e3f2329aeee43eec2eb2 Mon Sep 17 00:00:00 2001 From: GPK Date: Wed, 2 Oct 2024 12:48:12 +0100 Subject: [PATCH 108/802] =?UTF-8?q?Revert=20"Move=20FSHook/PackageIndexHoo?= =?UTF-8?q?k/SubprocessHook=20to=20standard=20provider=20(#42=E2=80=A6"=20?= =?UTF-8?q?(#42659)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit This reverts commit 61d1dbbc7feb9728da125dc00ad05314758036eb. --- .../standard => }/hooks/filesystem.py | 0 .../standard => }/hooks/package_index.py | 0 .../standard => }/hooks/subprocess.py | 4 +- airflow/operators/bash.py | 2 +- airflow/providers/standard/hooks/__init__.py | 16 -------- airflow/providers/standard/provider.yaml | 7 ---- airflow/providers_manager.py | 4 +- airflow/sensors/filesystem.py | 2 +- .../logging-monitoring/errors.rst | 2 +- .../operators-and-hooks-ref.rst | 4 +- .../standard => }/hooks/test_package_index.py | 6 +-- .../standard => }/hooks/test_subprocess.py | 6 +-- tests/providers/standard/hooks/__init__.py | 16 -------- .../standard/hooks/test_filesystem.py | 39 ------------------- tests/sensors/test_filesystem.py | 2 +- 15 files changed, 16 insertions(+), 94 deletions(-) rename airflow/{providers/standard => }/hooks/filesystem.py (100%) rename airflow/{providers/standard => }/hooks/package_index.py (100%) rename airflow/{providers/standard => }/hooks/subprocess.py (96%) delete mode 100644 airflow/providers/standard/hooks/__init__.py rename tests/{providers/standard => }/hooks/test_package_index.py (93%) rename tests/{providers/standard => }/hooks/test_subprocess.py (95%) delete mode 100644 tests/providers/standard/hooks/__init__.py delete mode 100644 tests/providers/standard/hooks/test_filesystem.py diff --git a/airflow/providers/standard/hooks/filesystem.py b/airflow/hooks/filesystem.py similarity index 100% rename from airflow/providers/standard/hooks/filesystem.py rename to airflow/hooks/filesystem.py diff --git a/airflow/providers/standard/hooks/package_index.py b/airflow/hooks/package_index.py similarity index 100% rename from airflow/providers/standard/hooks/package_index.py rename to airflow/hooks/package_index.py diff --git a/airflow/providers/standard/hooks/subprocess.py b/airflow/hooks/subprocess.py similarity index 96% rename from airflow/providers/standard/hooks/subprocess.py rename to airflow/hooks/subprocess.py index 9e578a7d8034b..bc20b5c20b4c5 100644 --- a/airflow/providers/standard/hooks/subprocess.py +++ b/airflow/hooks/subprocess.py @@ -52,8 +52,8 @@ def run_command( :param env: Optional dict containing environment variables to be made available to the shell environment in which ``command`` will be executed. If omitted, ``os.environ`` will be used. Note, that in case you have Sentry configured, original variables from the environment - will also be passed to the subprocess with ``SUBPROCESS_`` prefix. See: - https://airflow.apache.org/docs/apache-airflow/stable/administration-and-deployment/logging-monitoring/errors.html for details. + will also be passed to the subprocess with ``SUBPROCESS_`` prefix. See + :doc:`/administration-and-deployment/logging-monitoring/errors` for details. :param output_encoding: encoding to use for decoding stdout :param cwd: Working directory to run the command in. If None (default), the command is run in a temporary directory. diff --git a/airflow/operators/bash.py b/airflow/operators/bash.py index bf4a943df6e08..2ec0341a0d1e2 100644 --- a/airflow/operators/bash.py +++ b/airflow/operators/bash.py @@ -24,8 +24,8 @@ from typing import TYPE_CHECKING, Any, Callable, Container, Sequence, cast from airflow.exceptions import AirflowException, AirflowSkipException +from airflow.hooks.subprocess import SubprocessHook from airflow.models.baseoperator import BaseOperator -from airflow.providers.standard.hooks.subprocess import SubprocessHook from airflow.utils.operator_helpers import context_to_airflow_vars from airflow.utils.types import ArgNotSet diff --git a/airflow/providers/standard/hooks/__init__.py b/airflow/providers/standard/hooks/__init__.py deleted file mode 100644 index 13a83393a9124..0000000000000 --- a/airflow/providers/standard/hooks/__init__.py +++ /dev/null @@ -1,16 +0,0 @@ -# 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. diff --git a/airflow/providers/standard/provider.yaml b/airflow/providers/standard/provider.yaml index 068fde1fe3761..83d8acf0a68b3 100644 --- a/airflow/providers/standard/provider.yaml +++ b/airflow/providers/standard/provider.yaml @@ -50,10 +50,3 @@ sensors: - airflow.providers.standard.sensors.time_delta - airflow.providers.standard.sensors.time - airflow.providers.standard.sensors.weekday - -hooks: - - integration-name: Standard - python-modules: - - airflow.providers.standard.hooks.filesystem - - airflow.providers.standard.hooks.package_index - - airflow.providers.standard.hooks.subprocess diff --git a/airflow/providers_manager.py b/airflow/providers_manager.py index e276c465ef689..2c673063cb23e 100644 --- a/airflow/providers_manager.py +++ b/airflow/providers_manager.py @@ -36,8 +36,8 @@ from packaging.utils import canonicalize_name from airflow.exceptions import AirflowOptionalProviderFeatureException -from airflow.providers.standard.hooks.filesystem import FSHook -from airflow.providers.standard.hooks.package_index import PackageIndexHook +from airflow.hooks.filesystem import FSHook +from airflow.hooks.package_index import PackageIndexHook from airflow.typing_compat import ParamSpec from airflow.utils import yaml from airflow.utils.entry_points import entry_points_with_dist diff --git a/airflow/sensors/filesystem.py b/airflow/sensors/filesystem.py index 4496f5d6abfa4..5d32ab07ad4e7 100644 --- a/airflow/sensors/filesystem.py +++ b/airflow/sensors/filesystem.py @@ -25,7 +25,7 @@ from airflow.configuration import conf from airflow.exceptions import AirflowException -from airflow.providers.standard.hooks.filesystem import FSHook +from airflow.hooks.filesystem import FSHook from airflow.sensors.base import BaseSensorOperator from airflow.triggers.base import StartTriggerArgs from airflow.triggers.file import FileTrigger diff --git a/docs/apache-airflow/administration-and-deployment/logging-monitoring/errors.rst b/docs/apache-airflow/administration-and-deployment/logging-monitoring/errors.rst index 0ad3fa8c5127a..cb09843422321 100644 --- a/docs/apache-airflow/administration-and-deployment/logging-monitoring/errors.rst +++ b/docs/apache-airflow/administration-and-deployment/logging-monitoring/errors.rst @@ -96,7 +96,7 @@ Impact of Sentry on Environment variables passed to Subprocess Hook When Sentry is enabled, by default it changes the standard library to pass all environment variables to subprocesses opened by Airflow. This changes the default behaviour of -:class:`airflow.providers.standard.hooks.subprocess.SubprocessHook` - always all environment variables are passed to the +:class:`airflow.hooks.subprocess.SubprocessHook` - always all environment variables are passed to the subprocess executed with specific set of environment variables. In this case not only the specified environment variables are passed but also all existing environment variables are passed with ``SUBPROCESS_`` prefix added. This happens also for all other subprocesses. diff --git a/docs/apache-airflow/operators-and-hooks-ref.rst b/docs/apache-airflow/operators-and-hooks-ref.rst index d4ac6bda74c34..16b74305a958b 100644 --- a/docs/apache-airflow/operators-and-hooks-ref.rst +++ b/docs/apache-airflow/operators-and-hooks-ref.rst @@ -106,8 +106,8 @@ For details see: :doc:`apache-airflow-providers:operators-and-hooks-ref/index`. * - Hooks - Guides - * - :mod:`airflow.providers.standard.hooks.filesystem` + * - :mod:`airflow.hooks.filesystem` - - * - :mod:`airflow.providers.standard.hooks.subprocess` + * - :mod:`airflow.hooks.subprocess` - diff --git a/tests/providers/standard/hooks/test_package_index.py b/tests/hooks/test_package_index.py similarity index 93% rename from tests/providers/standard/hooks/test_package_index.py rename to tests/hooks/test_package_index.py index 6a90db0715d81..9da429c5a09cf 100644 --- a/tests/providers/standard/hooks/test_package_index.py +++ b/tests/hooks/test_package_index.py @@ -21,8 +21,8 @@ import pytest +from airflow.hooks.package_index import PackageIndexHook from airflow.models.connection import Connection -from airflow.providers.standard.hooks.package_index import PackageIndexHook class MockConnection(Connection): @@ -73,7 +73,7 @@ def mock_get_connection(monkeypatch: pytest.MonkeyPatch, request: pytest.Fixture password: str | None = testdata.get("password", None) expected_result: str | None = testdata.get("expected_result", None) monkeypatch.setattr( - "airflow.providers.standard.hooks.package_index.PackageIndexHook.get_connection", + "airflow.hooks.package_index.PackageIndexHook.get_connection", lambda *_: MockConnection(host, login, password), ) return expected_result @@ -104,7 +104,7 @@ class MockProc: return MockProc() - monkeypatch.setattr("airflow.providers.standard.hooks.package_index.subprocess.run", mock_run) + monkeypatch.setattr("airflow.hooks.package_index.subprocess.run", mock_run) hook_instance = PackageIndexHook() if mock_get_connection: diff --git a/tests/providers/standard/hooks/test_subprocess.py b/tests/hooks/test_subprocess.py similarity index 95% rename from tests/providers/standard/hooks/test_subprocess.py rename to tests/hooks/test_subprocess.py index 2b2e9473359e5..0f625be816887 100644 --- a/tests/providers/standard/hooks/test_subprocess.py +++ b/tests/hooks/test_subprocess.py @@ -26,7 +26,7 @@ import pytest -from airflow.providers.standard.hooks.subprocess import SubprocessHook +from airflow.hooks.subprocess import SubprocessHook OS_ENV_KEY = "SUBPROCESS_ENV_TEST" OS_ENV_VAL = "this-is-from-os-environ" @@ -81,11 +81,11 @@ def test_return_value(self, val, expected): @mock.patch.dict("os.environ", clear=True) @mock.patch( - "airflow.providers.standard.hooks.subprocess.TemporaryDirectory", + "airflow.hooks.subprocess.TemporaryDirectory", return_value=MagicMock(__enter__=MagicMock(return_value="/tmp/airflowtmpcatcat")), ) @mock.patch( - "airflow.providers.standard.hooks.subprocess.Popen", + "airflow.hooks.subprocess.Popen", return_value=MagicMock(stdout=MagicMock(readline=MagicMock(side_effect=StopIteration), returncode=0)), ) def test_should_exec_subprocess(self, mock_popen, mock_temporary_directory): diff --git a/tests/providers/standard/hooks/__init__.py b/tests/providers/standard/hooks/__init__.py deleted file mode 100644 index 13a83393a9124..0000000000000 --- a/tests/providers/standard/hooks/__init__.py +++ /dev/null @@ -1,16 +0,0 @@ -# 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. diff --git a/tests/providers/standard/hooks/test_filesystem.py b/tests/providers/standard/hooks/test_filesystem.py deleted file mode 100644 index bbcd22dc94219..0000000000000 --- a/tests/providers/standard/hooks/test_filesystem.py +++ /dev/null @@ -1,39 +0,0 @@ -# -# 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 pytest - -from airflow.providers.standard.hooks.filesystem import FSHook - -pytestmark = pytest.mark.db_test - - -class TestFSHook: - def test_get_ui_field_behaviour(self): - fs_hook = FSHook() - assert fs_hook.get_ui_field_behaviour() == { - "hidden_fields": ["host", "schema", "port", "login", "password", "extra"], - "relabeling": {}, - "placeholders": {}, - } - - def test_get_path(self): - fs_hook = FSHook(fs_conn_id="fs_default") - - assert fs_hook.get_path() == "/" diff --git a/tests/sensors/test_filesystem.py b/tests/sensors/test_filesystem.py index 641f2f218f2db..1fb123cfe7248 100644 --- a/tests/sensors/test_filesystem.py +++ b/tests/sensors/test_filesystem.py @@ -40,7 +40,7 @@ @pytest.mark.skip_if_database_isolation_mode # Test is broken in db isolation mode class TestFileSensor: def setup_method(self): - from airflow.providers.standard.hooks.filesystem import FSHook + from airflow.hooks.filesystem import FSHook hook = FSHook() args = {"owner": "airflow", "start_date": DEFAULT_DATE} From d8f44b8445acabdb714b3994a0492ea974e7edd3 Mon Sep 17 00:00:00 2001 From: Daniel Standish <15932138+dstandish@users.noreply.github.com> Date: Wed, 2 Oct 2024 05:59:14 -0700 Subject: [PATCH 109/802] Add backfill cancellation logic (#42530) --- .../endpoints/backfill_endpoint.py | 40 ++++++++--------- airflow/models/backfill.py | 45 +++++++++++++++++-- tests/models/test_backfill.py | 45 ++++++++++++++++++- 3 files changed, 104 insertions(+), 26 deletions(-) diff --git a/airflow/api_connexion/endpoints/backfill_endpoint.py b/airflow/api_connexion/endpoints/backfill_endpoint.py index baafdeea4f992..a0e728c5bc464 100644 --- a/airflow/api_connexion/endpoints/backfill_endpoint.py +++ b/airflow/api_connexion/endpoints/backfill_endpoint.py @@ -32,8 +32,12 @@ backfill_collection_schema, backfill_schema, ) -from airflow.models.backfill import AlreadyRunningBackfill, Backfill, _create_backfill -from airflow.utils import timezone +from airflow.models.backfill import ( + AlreadyRunningBackfill, + Backfill, + _cancel_backfill, + _create_backfill, +) from airflow.utils.session import NEW_SESSION, provide_session from airflow.www.decorators import action_logging @@ -104,24 +108,6 @@ def unpause_backfill(*, backfill_id, session, **kwargs): return backfill_schema.dump(br) -@provide_session -@backfill_to_dag -@security.requires_access_dag("PUT") -@action_logging -def cancel_backfill(*, backfill_id, session, **kwargs): - br: Backfill = session.get(Backfill, backfill_id) - if br.completed_at is not None: - raise Conflict("Backfill is already completed.") - - br.completed_at = timezone.utcnow() - - # first, pause - if not br.is_paused: - br.is_paused = True - session.commit() - return backfill_schema.dump(br) - - @provide_session @backfill_to_dag @security.requires_access_dag("GET") @@ -155,3 +141,17 @@ def create_backfill( return backfill_schema.dump(backfill_obj) except AlreadyRunningBackfill: raise Conflict(f"There is already a running backfill for dag {dag_id}") + + +@provide_session +@backfill_to_dag +@security.requires_access_dag("PUT") +@action_logging +def cancel_backfill( + *, + backfill_id, + session: Session = NEW_SESSION, # used by backfill_to_dag decorator + **kwargs, +): + br = _cancel_backfill(backfill_id=backfill_id) + return backfill_schema.dump(br) diff --git a/airflow/models/backfill.py b/airflow/models/backfill.py index 6d3a8ee4fa922..db10c804aac0d 100644 --- a/airflow/models/backfill.py +++ b/airflow/models/backfill.py @@ -26,12 +26,13 @@ import logging from typing import TYPE_CHECKING -from sqlalchemy import Boolean, Column, ForeignKeyConstraint, Integer, UniqueConstraint, func, select +from sqlalchemy import Boolean, Column, ForeignKeyConstraint, Integer, UniqueConstraint, func, select, update from sqlalchemy.orm import relationship from sqlalchemy_jsonfield import JSONField -from airflow.api_connexion.exceptions import NotFound +from airflow.api_connexion.exceptions import Conflict, NotFound from airflow.exceptions import AirflowException +from airflow.models import DagRun from airflow.models.base import Base, StringID from airflow.models.serialized_dag import SerializedDagModel from airflow.settings import json @@ -48,7 +49,11 @@ class AlreadyRunningBackfill(AirflowException): - """Raised when attempting to create backfill and one already active.""" + """ + Raised when attempting to create backfill and one already active. + + :meta private: + """ class Backfill(Base): @@ -172,7 +177,11 @@ def _create_backfill( session=session, ) except Exception: - dag.log.exception("something failed") + dag.log.exception( + "Error while attempting to create a dag run dag_id='%s' logical_date='%s'", + dag.dag_id, + info.logical_date, + ) session.rollback() session.add( BackfillDagRun( @@ -183,3 +192,31 @@ def _create_backfill( ) session.commit() return br + + +def _cancel_backfill(backfill_id) -> Backfill: + with create_session() as session: + b: Backfill = session.get(Backfill, backfill_id) + if b.completed_at is not None: + raise Conflict("Backfill is already completed.") + + b.completed_at = timezone.utcnow() + + # first, pause + if not b.is_paused: + b.is_paused = True + + session.commit() + + # now, let's mark all queued dag runs as failed + query = ( + update(DagRun) + .where( + DagRun.id.in_(select(BackfillDagRun.dag_run_id).where(BackfillDagRun.backfill_id == b.id)), + DagRun.state == DagRunState.QUEUED, + ) + .values(state=DagRunState.FAILED) + .execution_options(synchronize_session=False) + ) + session.execute(query) + return b diff --git a/tests/models/test_backfill.py b/tests/models/test_backfill.py index 9a845f86803e0..c45625db335de 100644 --- a/tests/models/test_backfill.py +++ b/tests/models/test_backfill.py @@ -24,7 +24,13 @@ from sqlalchemy import select from airflow.models import DagRun -from airflow.models.backfill import AlreadyRunningBackfill, Backfill, BackfillDagRun, _create_backfill +from airflow.models.backfill import ( + AlreadyRunningBackfill, + Backfill, + BackfillDagRun, + _cancel_backfill, + _create_backfill, +) from airflow.operators.python import PythonOperator from airflow.utils.state import DagRunState from tests.test_utils.db import clear_db_backfills, clear_db_dags, clear_db_runs, clear_db_serialized_dags @@ -71,7 +77,7 @@ def test_reverse_and_depends_on_past_fails(dep_on_past, dag_maker, session): @pytest.mark.parametrize("reverse", [True, False]) -def test_simple(reverse, dag_maker, session): +def test_create_backfill_simple(reverse, dag_maker, session): """ Verify simple case behavior. @@ -150,3 +156,38 @@ def test_active_dag_run(dag_maker, session): reverse=False, dag_run_conf={"this": "param"}, ) + + +def test_cancel_backfill(dag_maker, session): + """ + Queued runs should be marked *failed*. + Every other dag run should be left alone. + """ + with dag_maker(schedule="@daily") as dag: + PythonOperator(task_id="hi", python_callable=print) + b = _create_backfill( + dag_id=dag.dag_id, + from_date=pendulum.parse("2021-01-01"), + to_date=pendulum.parse("2021-01-05"), + max_active_runs=2, + reverse=False, + dag_run_conf={}, + ) + query = ( + select(DagRun) + .join(BackfillDagRun.dag_run) + .where(BackfillDagRun.backfill_id == b.id) + .order_by(BackfillDagRun.sort_ordinal) + ) + dag_runs = session.scalars(query).all() + dates = [str(x.logical_date.date()) for x in dag_runs] + expected_dates = ["2021-01-01", "2021-01-02", "2021-01-03", "2021-01-04", "2021-01-05"] + assert dates == expected_dates + assert all(x.state == DagRunState.QUEUED for x in dag_runs) + dag_runs[0].state = "running" + session.commit() + _cancel_backfill(backfill_id=b.id) + session.expunge_all() + dag_runs = session.scalars(query).all() + states = [x.state for x in dag_runs] + assert states == ["running", "failed", "failed", "failed", "failed"] From a3287bddd35bfe52807c1a4afc17ba61b8c5e24d Mon Sep 17 00:00:00 2001 From: Vincent <97131062+vincbeck@users.noreply.github.com> Date: Wed, 2 Oct 2024 09:48:32 -0400 Subject: [PATCH 110/802] Use FAB auth manager in `test_google_openid` (#42622) --- .../google/common/auth_backend/test_google_openid.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/tests/providers/google/common/auth_backend/test_google_openid.py b/tests/providers/google/common/auth_backend/test_google_openid.py index dab0eae07a23d..260ae0d6fb5e1 100644 --- a/tests/providers/google/common/auth_backend/test_google_openid.py +++ b/tests/providers/google/common/auth_backend/test_google_openid.py @@ -22,6 +22,7 @@ from google.auth.exceptions import GoogleAuthError from airflow.www.app import create_app +from tests.test_utils.compat import AIRFLOW_V_2_9_PLUS from tests.test_utils.config import conf_vars from tests.test_utils.db import clear_db_pools from tests.test_utils.decorators import dont_initialize_flask_app_submodules @@ -41,7 +42,13 @@ def google_openid_app(): ) def factory(): with conf_vars( - {("api", "auth_backends"): "airflow.providers.google.common.auth_backend.google_openid"} + { + ("api", "auth_backends"): "airflow.providers.google.common.auth_backend.google_openid", + ( + "core", + "auth_manager", + ): "airflow.providers.fab.auth_manager.fab_auth_manager.FabAuthManager", + } ): _app = create_app(testing=True, config={"WTF_CSRF_ENABLED": False}) # type:ignore _app.config["AUTH_ROLE_PUBLIC"] = None @@ -67,6 +74,7 @@ def admin_user(google_openid_app): return role_admin +@pytest.mark.skipif(not AIRFLOW_V_2_9_PLUS, reason="The tests should be skipped for Airflow < 2.9") @pytest.mark.skip_if_database_isolation_mode @pytest.mark.db_test class TestGoogleOpenID: From 5900bf8f2a8bf1b3d9365efd5c2a1c8c6a4eaef7 Mon Sep 17 00:00:00 2001 From: Jed Cunningham <66968678+jedcunningham@users.noreply.github.com> Date: Wed, 2 Oct 2024 08:38:37 -0600 Subject: [PATCH 111/802] Remove "project" from log path in callback docs (#42666) Airflow doesn't have the concept of a "project", unless DAG authors add that layer themselves. --- .../logging-monitoring/callbacks.rst | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/apache-airflow/administration-and-deployment/logging-monitoring/callbacks.rst b/docs/apache-airflow/administration-and-deployment/logging-monitoring/callbacks.rst index b54071373cf09..4f74626ab29ba 100644 --- a/docs/apache-airflow/administration-and-deployment/logging-monitoring/callbacks.rst +++ b/docs/apache-airflow/administration-and-deployment/logging-monitoring/callbacks.rst @@ -34,7 +34,7 @@ For example, you may wish to alert when certain tasks have failed, or have the l Callback functions are executed after tasks are completed. Errors in callback functions will show up in scheduler logs rather than task logs. By default, scheduler logs do not show up in the UI and instead can be found in - ``$AIRFLOW_HOME/logs/scheduler/latest/PROJECT/DAG_FILE.py.log`` + ``$AIRFLOW_HOME/logs/scheduler/latest/DAG_FILE.py.log`` Callback Types -------------- From 9e357cca93cafe9b0a5532e710aaee9aa0a3e060 Mon Sep 17 00:00:00 2001 From: Shahar Epstein <60007259+shahar1@users.noreply.github.com> Date: Wed, 2 Oct 2024 19:07:31 +0300 Subject: [PATCH 112/802] Fix invalid path in lineage.rst (#42655) --- docs/apache-airflow/administration-and-deployment/lineage.rst | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/apache-airflow/administration-and-deployment/lineage.rst b/docs/apache-airflow/administration-and-deployment/lineage.rst index d2ef63d869755..b274809175c03 100644 --- a/docs/apache-airflow/administration-and-deployment/lineage.rst +++ b/docs/apache-airflow/administration-and-deployment/lineage.rst @@ -101,7 +101,7 @@ The collector then uses this data to construct AIP-60 compliant Assets, a standa .. code-block:: python - from airflow.lineage.hook_lineage import get_hook_lineage_collector + from airflow.lineage.hook.lineage import get_hook_lineage_collector class CustomHook(BaseHook): From 642532004487be78172bb7c3e0cc3ae25afafd2d Mon Sep 17 00:00:00 2001 From: Jed Cunningham <66968678+jedcunningham@users.noreply.github.com> Date: Wed, 2 Oct 2024 10:36:30 -0600 Subject: [PATCH 113/802] Add support for PostgreSQL 17 in Breeze (#42644) * Add support for PostgreSQL 17 in Breeze * Fix tests --------- Co-authored-by: Tzu-ping Chung --- README.md | 2 +- dev/breeze/doc/images/output-commands.svg | 42 ++--- dev/breeze/doc/images/output_setup_config.svg | 2 +- dev/breeze/doc/images/output_setup_config.txt | 2 +- dev/breeze/doc/images/output_shell.svg | 140 +++++++-------- dev/breeze/doc/images/output_shell.txt | 2 +- .../doc/images/output_start-airflow.svg | 2 +- .../doc/images/output_start-airflow.txt | 2 +- .../doc/images/output_testing_db-tests.svg | 158 ++++++++--------- .../doc/images/output_testing_db-tests.txt | 2 +- .../output_testing_integration-tests.svg | 50 +++--- .../output_testing_integration-tests.txt | 2 +- .../doc/images/output_testing_tests.svg | 162 +++++++++--------- .../doc/images/output_testing_tests.txt | 2 +- .../src/airflow_breeze/global_constants.py | 4 +- dev/breeze/tests/test_selective_checks.py | 4 +- generated/PYPI_README.md | 2 +- 17 files changed, 296 insertions(+), 284 deletions(-) diff --git a/README.md b/README.md index 3169ac5144844..3cd6416e93405 100644 --- a/README.md +++ b/README.md @@ -102,7 +102,7 @@ Apache Airflow is tested with: | Python | 3.8, 3.9, 3.10, 3.11, 3.12 | 3.8, 3.9, 3.10, 3.11, 3.12 | | Platform | AMD64/ARM64(\*) | AMD64/ARM64(\*) | | Kubernetes | 1.28, 1.29, 1.30, 1.31 | 1.27, 1.28, 1.29, 1.30 | -| PostgreSQL | 12, 13, 14, 15, 16 | 12, 13, 14, 15, 16 | +| PostgreSQL | 12, 13, 14, 15, 16, 17 | 12, 13, 14, 15, 16 | | MySQL | 8.0, 8.4, Innovation | 8.0, 8.4, Innovation | | SQLite | 3.15.0+ | 3.15.0+ | diff --git a/dev/breeze/doc/images/output-commands.svg b/dev/breeze/doc/images/output-commands.svg index 1556dfef6f5a7..78c753526e449 100644 --- a/dev/breeze/doc/images/output-commands.svg +++ b/dev/breeze/doc/images/output-commands.svg @@ -301,53 +301,53 @@ Usage:breeze[OPTIONS] COMMAND [ARGS]... ╭─ Execution mode ─────────────────────────────────────────────────────────────────────────────────────────────────────╮ -│--python-pPython major/minor version used in Airflow image for images.│ +│--python-pPython major/minor version used in Airflow image for images.│ │(>3.8< | 3.9 | 3.10 | 3.11 | 3.12)                          │ │[default: 3.8]                                              │ -│--integrationIntegration(s) to enable when running (can be more than one).                       │ +│--integrationIntegration(s) to enable when running (can be more than one).                       │ │(all | all-testable | cassandra | celery | drill | kafka | kerberos | mongo | mssql │ │| openlineage | otel | pinot | qdrant | redis | statsd | trino | ydb)               │ -│--standalone-dag-processorRun standalone dag processor for start-airflow.│ -│--database-isolationRun airflow in database isolation mode.│ +│--standalone-dag-processorRun standalone dag processor for start-airflow.│ +│--database-isolationRun airflow in database isolation mode.│ ╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ ╭─ Docker Compose selection and cleanup ───────────────────────────────────────────────────────────────────────────────╮ -│--project-nameName of the docker-compose project to bring down. The `docker-compose` is for legacy breeze       │ -│project name and you can use `breeze down --project-name docker-compose` to stop all containers   │ +│--project-nameName of the docker-compose project to bring down. The `docker-compose` is for legacy breeze       │ +│project name and you can use `breeze down --project-name docker-compose` to stop all containers   │ │belonging to it.                                                                                  │ │(breeze | pre-commit | docker-compose)                                                            │ │[default: breeze]                                                                                 │ -│--docker-hostOptional - docker host to use when running docker commands. When set, the `--builder` option is   │ +│--docker-hostOptional - docker host to use when running docker commands. When set, the `--builder` option is   │ │ignored when building images.                                                                     │ │(TEXT)                                                                                            │ ╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ ╭─ Database ───────────────────────────────────────────────────────────────────────────────────────────────────────────╮ -│--backend-bDatabase backend to use. If 'none' is chosen, Breeze will start with an invalid database    │ +│--backend-bDatabase backend to use. If 'none' is chosen, Breeze will start with an invalid database    │ │configuration, meaning there will be no database available, and any attempts to connect to  │ │the Airflow database will fail.                                                             │ │(>sqlite< | mysql | postgres | none)                                                        │ │[default: sqlite]                                                                           │ -│--postgres-version-PVersion of Postgres used.(>12< | 13 | 14 | 15 | 16)[default: 12]│ -│--mysql-version-MVersion of MySQL used.(>8.0< | 8.4)[default: 8.0]│ -│--db-reset-dReset DB when entering the container.│ +│--postgres-version-PVersion of Postgres used.(>12< | 13 | 14 | 15 | 16 | 17)[default: 12]│ +│--mysql-version-MVersion of MySQL used.(>8.0< | 8.4)[default: 8.0]│ +│--db-reset-dReset DB when entering the container.│ ╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ ╭─ Build CI image (before entering shell) ─────────────────────────────────────────────────────────────────────────────╮ -│--github-repository-gGitHub repository used to pull, push run images.(TEXT)[default: apache/airflow]│ -│--builderBuildx builder used to perform `docker buildx build` commands.(TEXT)│ +│--github-repository-gGitHub repository used to pull, push run images.(TEXT)[default: apache/airflow]│ +│--builderBuildx builder used to perform `docker buildx build` commands.(TEXT)│ │[default: autodetect]                                         │ -│--use-uv/--no-use-uvUse uv instead of pip as packaging tool to build the image.[default: use-uv]│ -│--uv-http-timeoutTimeout for requests that UV makes (only used in case of UV builds).(INTEGER RANGE)│ +│--use-uv/--no-use-uvUse uv instead of pip as packaging tool to build the image.[default: use-uv]│ +│--uv-http-timeoutTimeout for requests that UV makes (only used in case of UV builds).(INTEGER RANGE)│ │[default: 300; x>=1]                                                │ ╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ ╭─ Other options ──────────────────────────────────────────────────────────────────────────────────────────────────────╮ -│--forward-credentials-fForward local credentials to container when running.│ -│--max-timeMaximum time that the command should take - if it takes longer, the command will fail.│ +│--forward-credentials-fForward local credentials to container when running.│ +│--max-timeMaximum time that the command should take - if it takes longer, the command will fail.│ │(INTEGER RANGE)                                                                       │ ╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ ╭─ Common options ─────────────────────────────────────────────────────────────────────────────────────────────────────╮ -│--answer-aForce answer to questions.(y | n | q | yes | no | quit)│ -│--dry-run-DIf dry-run is set, commands are only printed, not executed.│ -│--verbose-vPrint verbose information about performed steps.│ -│--help-hShow this message and exit.│ +│--answer-aForce answer to questions.(y | n | q | yes | no | quit)│ +│--dry-run-DIf dry-run is set, commands are only printed, not executed.│ +│--verbose-vPrint verbose information about performed steps.│ +│--help-hShow this message and exit.│ ╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ ╭─ Developer commands ─────────────────────────────────────────────────────────────────────────────────────────────────╮ │start-airflow          Enter breeze environment and starts all Airflow components in the tmux session. Compile    │ diff --git a/dev/breeze/doc/images/output_setup_config.svg b/dev/breeze/doc/images/output_setup_config.svg index 2fb7fab65273c..69780cb5426af 100644 --- a/dev/breeze/doc/images/output_setup_config.svg +++ b/dev/breeze/doc/images/output_setup_config.svg @@ -137,7 +137,7 @@ │attempts to connect to the Airflow database will fail.                         │ │(>sqlite< | mysql | postgres | none)                                           │ │[default: sqlite]                                                              │ -│--postgres-version-PVersion of Postgres used.(>12< | 13 | 14 | 15 | 16)[default: 12]│ +│--postgres-version-PVersion of Postgres used.(>12< | 13 | 14 | 15 | 16 | 17)[default: 12]│ │--mysql-version-MVersion of MySQL used.(>8.0< | 8.4)[default: 8.0]│ │--cheatsheet/--no-cheatsheet-C/-cEnable/disable cheatsheet.│ │--asciiart/--no-asciiart-A/-aEnable/disable ASCIIart.│ diff --git a/dev/breeze/doc/images/output_setup_config.txt b/dev/breeze/doc/images/output_setup_config.txt index f47fa38e42c7b..97d022c37b5e1 100644 --- a/dev/breeze/doc/images/output_setup_config.txt +++ b/dev/breeze/doc/images/output_setup_config.txt @@ -1 +1 @@ -422c8c524b557fcf5924da4c8590935d +783acef079cbdd31cd7880618c20fae5 diff --git a/dev/breeze/doc/images/output_shell.svg b/dev/breeze/doc/images/output_shell.svg index bf8fbc3ee81f6..1e86993b2c466 100644 --- a/dev/breeze/doc/images/output_shell.svg +++ b/dev/breeze/doc/images/output_shell.svg @@ -573,7 +573,7 @@ │the Airflow database will fail.                                                             │ │(>sqlite< | mysql | postgres | none)                                                        │ │[default: sqlite]                                                                           │ -│--postgres-version-PVersion of Postgres used.(>12< | 13 | 14 | 15 | 16)[default: 12]│ +│--postgres-version-PVersion of Postgres used.(>12< | 13 | 14 | 15 | 16 | 17)[default: 12]│ │--mysql-version-MVersion of MySQL used.(>8.0< | 8.4)[default: 8.0]│ │--db-reset-dReset DB when entering the container.│ ╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ @@ -620,75 +620,75 @@ │--airflow-skip-constraintsDo not use constraints when installing airflow.│ │--clean-airflow-installationClean the airflow installation before installing version│ │specified by --use-airflow-version.                     │ -│--force-lowest-dependenciesRun tests for the lowest direct dependencies of Airflow │ -│or selected provider if `Provider[PROVIDER_ID]` is used │ -│as test type.                                           │ -│--install-airflow-with-constraints/--no-install-airflow…Install airflow in a separate step, with constraints    │ -│determined from package or airflow version.             │ -│[default: install-airflow-with-constraints]             │ -│--install-selected-providersComma-separated list of providers selected to be        │ -│installed (implies --use-packages-from-dist).           │ -│(TEXT)                                                  │ -│--package-formatFormat of packages that should be installed from dist.│ -│(wheel | sdist)                                       │ -│[default: wheel]                                      │ -│--providers-constraints-locationLocation of providers constraints to use (remote URL or │ -│local context file).                                    │ -│(TEXT)                                                  │ -│--providers-constraints-modeMode of constraints for Providers for CI image building.│ -│(constraints-source-providers | constraints |           │ -│constraints-no-providers)                               │ -│[default: constraints-source-providers]                 │ -│--providers-constraints-referenceConstraint reference to use for providers installation  │ -│(used in calculated constraints URL). Can be 'default'  │ -│in which case the default constraints-reference is used.│ -│(TEXT)                                                  │ -│--providers-skip-constraintsDo not use constraints when installing providers.│ -│--test-typeType of test to run. With Providers, you can specify    │ -│tests of which providers should be run:                 │ -│`Providers[airbyte,http]` or excluded from the full test│ -│suite: `Providers[-amazon,google]`                      │ -│(All | Default | API | Always | BranchExternalPython |  │ -│BranchPythonVenv | CLI | Core | ExternalPython |        │ -│Operators | Other | PlainAsserts | Providers |          │ -│PythonVenv | Serialization | WWW | All-Postgres |       │ -│All-MySQL | All-Quarantined)                            │ -│[default: Default]                                      │ -│--use-airflow-versionUse (reinstall at entry) Airflow version from PyPI. It  │ -│can also be version (to install from PyPI), `none`,     │ -│`wheel`, or `sdist` to install from `dist` folder, or   │ -│VCS URL to install from                                 │ -│(https://pip.pypa.io/en/stable/topics/vcs-support/).    │ -│Implies --mount-sources `remove`.                       │ -│(none | wheel | sdist | <airflow_version>)              │ -│--use-packages-from-distInstall all found packages (--package-format determines │ -│type) from 'dist' folder when entering breeze.          │ -╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ -╭─ Upgrading/downgrading/removing selected packages ───────────────────────────────────────────────────────────────────╮ -│--upgrade-botoRemove aiobotocore and upgrade botocore and boto to the latest version.│ -│--downgrade-sqlalchemyDowngrade SQLAlchemy to minimum supported version.│ -│--downgrade-pendulumDowngrade Pendulum to minimum supported version.│ -╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ -╭─ DB test flags ──────────────────────────────────────────────────────────────────────────────────────────────────────╮ -│--run-db-tests-onlyOnly runs tests that require a database│ -│--skip-db-testsSkip tests that require a database│ -╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ -╭─ Other options ──────────────────────────────────────────────────────────────────────────────────────────────────────╮ -│--forward-credentials-fForward local credentials to container when running.│ -│--max-timeMaximum time that the command should take - if it takes longer, the command will fail.│ -│(INTEGER RANGE)                                                                       │ -│--verbose-commandsShow details of commands executed.│ -│--keep-env-variablesDo not clear environment variables that might have side effect while running tests│ -│--no-db-cleanupDo not clear the database before each test module│ -╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ -╭─ Common options ─────────────────────────────────────────────────────────────────────────────────────────────────────╮ -│--answer-aForce answer to questions.(y | n | q | yes | no | quit)│ -│--dry-run-DIf dry-run is set, commands are only printed, not executed.│ -│--excluded-providersJSON-string of dictionary containing excluded providers per python version ({'3.12':      │ -│['provider']})                                                                            │ -│(TEXT)                                                                                    │ -│--verbose-vPrint verbose information about performed steps.│ -│--help-hShow this message and exit.│ +│--excluded-providersJSON-string of dictionary containing excluded providers │ +│per python version ({'3.12': ['provider']})             │ +│(TEXT)                                                  │ +│--force-lowest-dependenciesRun tests for the lowest direct dependencies of Airflow │ +│or selected provider if `Provider[PROVIDER_ID]` is used │ +│as test type.                                           │ +│--install-airflow-with-constraints/--no-install-airflow…Install airflow in a separate step, with constraints    │ +│determined from package or airflow version.             │ +│[default: install-airflow-with-constraints]             │ +│--install-selected-providersComma-separated list of providers selected to be        │ +│installed (implies --use-packages-from-dist).           │ +│(TEXT)                                                  │ +│--package-formatFormat of packages that should be installed from dist.│ +│(wheel | sdist)                                       │ +│[default: wheel]                                      │ +│--providers-constraints-locationLocation of providers constraints to use (remote URL or │ +│local context file).                                    │ +│(TEXT)                                                  │ +│--providers-constraints-modeMode of constraints for Providers for CI image building.│ +│(constraints-source-providers | constraints |           │ +│constraints-no-providers)                               │ +│[default: constraints-source-providers]                 │ +│--providers-constraints-referenceConstraint reference to use for providers installation  │ +│(used in calculated constraints URL). Can be 'default'  │ +│in which case the default constraints-reference is used.│ +│(TEXT)                                                  │ +│--providers-skip-constraintsDo not use constraints when installing providers.│ +│--test-typeType of test to run. With Providers, you can specify    │ +│tests of which providers should be run:                 │ +│`Providers[airbyte,http]` or excluded from the full test│ +│suite: `Providers[-amazon,google]`                      │ +│(All | Default | API | Always | BranchExternalPython |  │ +│BranchPythonVenv | CLI | Core | ExternalPython |        │ +│Operators | Other | PlainAsserts | Providers |          │ +│PythonVenv | Serialization | WWW | All-Postgres |       │ +│All-MySQL | All-Quarantined)                            │ +│[default: Default]                                      │ +│--use-airflow-versionUse (reinstall at entry) Airflow version from PyPI. It  │ +│can also be version (to install from PyPI), `none`,     │ +│`wheel`, or `sdist` to install from `dist` folder, or   │ +│VCS URL to install from                                 │ +│(https://pip.pypa.io/en/stable/topics/vcs-support/).    │ +│Implies --mount-sources `remove`.                       │ +│(none | wheel | sdist | <airflow_version>)              │ +│--use-packages-from-distInstall all found packages (--package-format determines │ +│type) from 'dist' folder when entering breeze.          │ +╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ +╭─ Upgrading/downgrading/removing selected packages ───────────────────────────────────────────────────────────────────╮ +│--upgrade-botoRemove aiobotocore and upgrade botocore and boto to the latest version.│ +│--downgrade-sqlalchemyDowngrade SQLAlchemy to minimum supported version.│ +│--downgrade-pendulumDowngrade Pendulum to minimum supported version.│ +╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ +╭─ DB test flags ──────────────────────────────────────────────────────────────────────────────────────────────────────╮ +│--run-db-tests-onlyOnly runs tests that require a database│ +│--skip-db-testsSkip tests that require a database│ +╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ +╭─ Other options ──────────────────────────────────────────────────────────────────────────────────────────────────────╮ +│--forward-credentials-fForward local credentials to container when running.│ +│--max-timeMaximum time that the command should take - if it takes longer, the command will fail.│ +│(INTEGER RANGE)                                                                       │ +│--verbose-commandsShow details of commands executed.│ +│--keep-env-variablesDo not clear environment variables that might have side effect while running tests│ +│--no-db-cleanupDo not clear the database before each test module│ +╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ +╭─ Common options ─────────────────────────────────────────────────────────────────────────────────────────────────────╮ +│--answer-aForce answer to questions.(y | n | q | yes | no | quit)│ +│--dry-run-DIf dry-run is set, commands are only printed, not executed.│ +│--verbose-vPrint verbose information about performed steps.│ +│--help-hShow this message and exit.│ ╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ diff --git a/dev/breeze/doc/images/output_shell.txt b/dev/breeze/doc/images/output_shell.txt index 71be1f4ed5fa6..8529a31200925 100644 --- a/dev/breeze/doc/images/output_shell.txt +++ b/dev/breeze/doc/images/output_shell.txt @@ -1 +1 @@ -4d7e652e8a79290f5ca783e94662ada1 +12f9e4a84051e05a5e0b9ec4fe3c8632 diff --git a/dev/breeze/doc/images/output_start-airflow.svg b/dev/breeze/doc/images/output_start-airflow.svg index 377745370a5b4..55bac0da8cd10 100644 --- a/dev/breeze/doc/images/output_start-airflow.svg +++ b/dev/breeze/doc/images/output_start-airflow.svg @@ -432,7 +432,7 @@ │the Airflow database will fail.                                                             │ │(>sqlite< | mysql | postgres | none)                                                        │ │[default: sqlite]                                                                           │ -│--postgres-version-PVersion of Postgres used.(>12< | 13 | 14 | 15 | 16)[default: 12]│ +│--postgres-version-PVersion of Postgres used.(>12< | 13 | 14 | 15 | 16 | 17)[default: 12]│ │--mysql-version-MVersion of MySQL used.(>8.0< | 8.4)[default: 8.0]│ │--db-reset-dReset DB when entering the container.│ ╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ diff --git a/dev/breeze/doc/images/output_start-airflow.txt b/dev/breeze/doc/images/output_start-airflow.txt index 6399738d71e5d..3aaa148b8ca83 100644 --- a/dev/breeze/doc/images/output_start-airflow.txt +++ b/dev/breeze/doc/images/output_start-airflow.txt @@ -1 +1 @@ -74f2c1895c08408a8caa90eaf96f98cf +f6365a250b86242436df9236025d447f diff --git a/dev/breeze/doc/images/output_testing_db-tests.svg b/dev/breeze/doc/images/output_testing_db-tests.svg index 916e8c6005c29..d9f6d92eef100 100644 --- a/dev/breeze/doc/images/output_testing_db-tests.svg +++ b/dev/breeze/doc/images/output_testing_db-tests.svg @@ -1,4 +1,4 @@ - + │--python-pPython major/minor version used in Airflow image for images.│ │(>3.8< | 3.9 | 3.10 | 3.11 | 3.12)                          │ │[default: 3.8]                                              │ -│--postgres-version-PVersion of Postgres used.(>12< | 13 | 14 | 15 | 16)[default: 12]│ -│--mysql-version-MVersion of MySQL used.(>8.0< | 8.4)[default: 8.0]│ -│--forward-credentials-fForward local credentials to container when running.│ -│--force-sa-warnings/--no-force-sa-warningsEnable `sqlalchemy.exc.MovedIn20Warning` during the tests runs.│ -│[default: force-sa-warnings]                                   │ -╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ -╭─ Options for parallel test commands ─────────────────────────────────────────────────────────────────────────────────╮ -│--parallelismMaximum number of processes to use while running the operation in parallel.│ -│(INTEGER RANGE)                                                            │ -│[default: 4; 1<=x<=8]                                                      │ -│--skip-cleanupSkip cleanup of temporary files created during parallel run.│ -│--debug-resourcesWhether to show resource information while running in parallel.│ -│--include-success-outputsWhether to include outputs of successful parallel runs (skipped by default).│ -╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ -╭─ Upgrading/downgrading/removing selected packages ───────────────────────────────────────────────────────────────────╮ -│--upgrade-botoRemove aiobotocore and upgrade botocore and boto to the latest version.│ -│--downgrade-sqlalchemyDowngrade SQLAlchemy to minimum supported version.│ -│--downgrade-pendulumDowngrade Pendulum to minimum supported version.│ -│--remove-arm-packagesRemoves arm packages from the image to test if ARM collection works│ -╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ -╭─ Advanced flag for tests command ────────────────────────────────────────────────────────────────────────────────────╮ -│--airflow-constraints-referenceConstraint reference to use for airflow installation   │ -│(used in calculated constraints URL).                  │ -│(TEXT)                                                 │ -│--clean-airflow-installationClean the airflow installation before installing       │ -│version specified by --use-airflow-version.            │ -│--excluded-providersJSON-string of dictionary containing excluded providers│ -│per python version ({'3.12': ['provider']})            │ -│(TEXT)                                                 │ -│--force-lowest-dependenciesRun tests for the lowest direct dependencies of Airflow│ -│or selected provider if `Provider[PROVIDER_ID]` is used│ -│as test type.                                          │ -│--github-repository-gGitHub repository used to pull, push run images.(TEXT)│ -│[default: apache/airflow]                       │ -│--image-tagTag of the image which is used to run the image        │ -│(implies --mount-sources=skip).                        │ -│(TEXT)                                                 │ -│[default: latest]                                      │ -│--install-airflow-with-constraints/--no-install-airflo…Install airflow in a separate step, with constraints   │ -│determined from package or airflow version.            │ -│[default: no-install-airflow-with-constraints]         │ -│--package-formatFormat of packages.(wheel | sdist | both)│ -│[default: wheel]   │ -│--providers-constraints-locationLocation of providers constraints to use (remote URL or│ -│local context file).                                   │ -│(TEXT)                                                 │ -│--providers-skip-constraintsDo not use constraints when installing providers.│ -│--use-airflow-versionUse (reinstall at entry) Airflow version from PyPI. It │ -│can also be version (to install from PyPI), `none`,    │ -│`wheel`, or `sdist` to install from `dist` folder, or  │ -│VCS URL to install from                                │ -│(https://pip.pypa.io/en/stable/topics/vcs-support/).   │ -│Implies --mount-sources `remove`.                      │ -│(none | wheel | sdist | <airflow_version>)             │ -│--use-packages-from-distInstall all found packages (--package-format determines│ -│type) from 'dist' folder when entering breeze.         │ -│--mount-sourcesChoose scope of local sources that should be mounted,  │ -│skipped, or removed (default = selected).              │ -│(selected | all | skip | remove | tests |              │ -│providers-and-tests)                                   │ -│[default: selected]                                    │ -│--skip-docker-compose-downSkips running docker-compose down after tests│ -│--skip-providersSpace-separated list of provider ids to skip when      │ -│running tests                                          │ -│(TEXT)                                                 │ -│--keep-env-variablesDo not clear environment variables that might have side│ -│effect while running tests                             │ -│--no-db-cleanupDo not clear the database before each test module│ -╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ -╭─ Common options ─────────────────────────────────────────────────────────────────────────────────────────────────────╮ -│--dry-run-DIf dry-run is set, commands are only printed, not executed.│ -│--verbose-vPrint verbose information about performed steps.│ -│--help-hShow this message and exit.│ -╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ +│--postgres-version-PVersion of Postgres used.(>12< | 13 | 14 | 15 | 16 | 17)│ +│[default: 12]            │ +│--mysql-version-MVersion of MySQL used.(>8.0< | 8.4)[default: 8.0]│ +│--forward-credentials-fForward local credentials to container when running.│ +│--force-sa-warnings/--no-force-sa-warningsEnable `sqlalchemy.exc.MovedIn20Warning` during the tests runs.│ +│[default: force-sa-warnings]                                   │ +╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ +╭─ Options for parallel test commands ─────────────────────────────────────────────────────────────────────────────────╮ +│--parallelismMaximum number of processes to use while running the operation in parallel.│ +│(INTEGER RANGE)                                                            │ +│[default: 4; 1<=x<=8]                                                      │ +│--skip-cleanupSkip cleanup of temporary files created during parallel run.│ +│--debug-resourcesWhether to show resource information while running in parallel.│ +│--include-success-outputsWhether to include outputs of successful parallel runs (skipped by default).│ +╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ +╭─ Upgrading/downgrading/removing selected packages ───────────────────────────────────────────────────────────────────╮ +│--upgrade-botoRemove aiobotocore and upgrade botocore and boto to the latest version.│ +│--downgrade-sqlalchemyDowngrade SQLAlchemy to minimum supported version.│ +│--downgrade-pendulumDowngrade Pendulum to minimum supported version.│ +│--remove-arm-packagesRemoves arm packages from the image to test if ARM collection works│ +╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ +╭─ Advanced flag for tests command ────────────────────────────────────────────────────────────────────────────────────╮ +│--airflow-constraints-referenceConstraint reference to use for airflow installation   │ +│(used in calculated constraints URL).                  │ +│(TEXT)                                                 │ +│--clean-airflow-installationClean the airflow installation before installing       │ +│version specified by --use-airflow-version.            │ +│--excluded-providersJSON-string of dictionary containing excluded providers│ +│per python version ({'3.12': ['provider']})            │ +│(TEXT)                                                 │ +│--force-lowest-dependenciesRun tests for the lowest direct dependencies of Airflow│ +│or selected provider if `Provider[PROVIDER_ID]` is used│ +│as test type.                                          │ +│--github-repository-gGitHub repository used to pull, push run images.(TEXT)│ +│[default: apache/airflow]                       │ +│--image-tagTag of the image which is used to run the image        │ +│(implies --mount-sources=skip).                        │ +│(TEXT)                                                 │ +│[default: latest]                                      │ +│--install-airflow-with-constraints/--no-install-airflo…Install airflow in a separate step, with constraints   │ +│determined from package or airflow version.            │ +│[default: no-install-airflow-with-constraints]         │ +│--package-formatFormat of packages.(wheel | sdist | both)│ +│[default: wheel]   │ +│--providers-constraints-locationLocation of providers constraints to use (remote URL or│ +│local context file).                                   │ +│(TEXT)                                                 │ +│--providers-skip-constraintsDo not use constraints when installing providers.│ +│--use-airflow-versionUse (reinstall at entry) Airflow version from PyPI. It │ +│can also be version (to install from PyPI), `none`,    │ +│`wheel`, or `sdist` to install from `dist` folder, or  │ +│VCS URL to install from                                │ +│(https://pip.pypa.io/en/stable/topics/vcs-support/).   │ +│Implies --mount-sources `remove`.                      │ +│(none | wheel | sdist | <airflow_version>)             │ +│--use-packages-from-distInstall all found packages (--package-format determines│ +│type) from 'dist' folder when entering breeze.         │ +│--mount-sourcesChoose scope of local sources that should be mounted,  │ +│skipped, or removed (default = selected).              │ +│(selected | all | skip | remove | tests |              │ +│providers-and-tests)                                   │ +│[default: selected]                                    │ +│--skip-docker-compose-downSkips running docker-compose down after tests│ +│--skip-providersSpace-separated list of provider ids to skip when      │ +│running tests                                          │ +│(TEXT)                                                 │ +│--keep-env-variablesDo not clear environment variables that might have side│ +│effect while running tests                             │ +│--no-db-cleanupDo not clear the database before each test module│ +╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ +╭─ Common options ─────────────────────────────────────────────────────────────────────────────────────────────────────╮ +│--dry-run-DIf dry-run is set, commands are only printed, not executed.│ +│--verbose-vPrint verbose information about performed steps.│ +│--help-hShow this message and exit.│ +╰──────────────────────────────────────────────────────────────────────────────────────────────────────────────────────╯ diff --git a/dev/breeze/doc/images/output_testing_db-tests.txt b/dev/breeze/doc/images/output_testing_db-tests.txt index 41be88099eb91..683eef1d0509b 100644 --- a/dev/breeze/doc/images/output_testing_db-tests.txt +++ b/dev/breeze/doc/images/output_testing_db-tests.txt @@ -1 +1 @@ -97cc799ed5c1244b2aeb680f3215021b +5a6490989a911c538b427ca2806d0e3e diff --git a/dev/breeze/doc/images/output_testing_integration-tests.svg b/dev/breeze/doc/images/output_testing_integration-tests.svg index ac4ce00627101..4f86fcef38b29 100644 --- a/dev/breeze/doc/images/output_testing_integration-tests.svg +++ b/dev/breeze/doc/images/output_testing_integration-tests.svg @@ -1,4 +1,4 @@ - +

Edge Worker Hosts

+ {% if hosts|length == 0 %} +

No Edge Workers connected or known currently.

+ {% else %} + +
- {% if dag.dag_id in dataset_triggered_next_run_info %} - {%- with ds_info = dataset_triggered_next_run_info[dag.dag_id] -%} + + {% if dag.dag_id in asset_triggered_next_run_info %} + {%- with asset_info = asset_triggered_next_run_info[dag.dag_id] -%}
- {% if ds_info.total == 1 -%} - On {{ ds_info.uri[0:40] + '…' if ds_info.uri and ds_info.uri|length > 40 else ds_info.uri|default('', true) }} + {% if asset_info.total == 1 -%} + On {{ asset_info.uri[0:40] + '…' if asset_info.uri and asset_info.uri|length > 40 else asset_info.uri|default('', true) }} {%- else -%} - {{ ds_info.ready }} of {{ ds_info.total }} datasets updated + {{ asset_info.ready }} of {{ asset_info.total }} datasets updated {%- endif %}
diff --git a/airflow/www/views.py b/airflow/www/views.py index 0ef37f71d336c..b3300b517e757 100644 --- a/airflow/www/views.py +++ b/airflow/www/views.py @@ -87,10 +87,10 @@ set_dag_run_state_to_success, set_state, ) +from airflow.assets import Asset, AssetAlias from airflow.auth.managers.models.resource_details import AccessView, DagAccessEntity, DagDetails from airflow.compat.functools import cache from airflow.configuration import AIRFLOW_CONFIG, conf -from airflow.datasets import Dataset, DatasetAlias from airflow.exceptions import ( AirflowConfigException, AirflowException, @@ -104,9 +104,9 @@ from airflow.jobs.scheduler_job_runner import SchedulerJobRunner from airflow.jobs.triggerer_job_runner import TriggererJobRunner from airflow.models import Connection, DagModel, DagTag, Log, SlaMiss, Trigger, XCom -from airflow.models.dag import get_dataset_triggered_next_run_info +from airflow.models.asset import AssetDagRunQueue, AssetEvent, AssetModel, DagScheduleAssetReference +from airflow.models.dag import get_asset_triggered_next_run_info from airflow.models.dagrun import RUN_ID_REGEX, DagRun, DagRunType -from airflow.models.dataset import DagScheduleDatasetReference, DatasetDagRunQueue, DatasetEvent, DatasetModel from airflow.models.errors import ParseImportError from airflow.models.serialized_dag import SerializedDagModel from airflow.models.taskinstance import TaskInstance, TaskInstanceNote @@ -469,9 +469,7 @@ def set_overall_state(record): "label": item.label, "extra_links": item.extra_links, "is_mapped": item_is_mapped, - "has_outlet_datasets": any( - isinstance(i, (Dataset, DatasetAlias)) for i in (item.outlets or []) - ), + "has_outlet_datasets": any(isinstance(i, (Asset, AssetAlias)) for i in (item.outlets or [])), "operator": item.operator_name, "trigger_rule": item.trigger_rule, **setup_teardown_type, @@ -1005,13 +1003,13 @@ def index(self): .all() ) - dataset_triggered_dag_ids = {dag.dag_id for dag in dags if dag.dataset_expression is not None} - if dataset_triggered_dag_ids: - dataset_triggered_next_run_info = get_dataset_triggered_next_run_info( - dataset_triggered_dag_ids, session=session + asset_triggered_dag_ids = {dag.dag_id for dag in dags if dag.dataset_expression is not None} + if asset_triggered_dag_ids: + asset_triggered_next_run_info = get_asset_triggered_next_run_info( + asset_triggered_dag_ids, session=session ) else: - dataset_triggered_next_run_info = {} + asset_triggered_next_run_info = {} file_tokens = {} for dag in dags: @@ -1168,15 +1166,15 @@ def _iter_parsed_moved_data_table_names(): sorting_key=arg_sorting_key, sorting_direction=arg_sorting_direction, auto_refresh_interval=conf.getint("webserver", "auto_refresh_interval"), - dataset_triggered_next_run_info=dataset_triggered_next_run_info, + asset_triggered_next_run_info=asset_triggered_next_run_info, scarf_url=scarf_url, file_tokens=file_tokens, ) @expose("/datasets") - @auth.has_access_dataset("GET") + @auth.has_access_asset("GET") def datasets(self): - """Datasets view.""" + """Assets view.""" state_color_mapping = State.state_color.copy() state_color_mapping["null"] = state_color_mapping.pop(None) return self.render_template( @@ -1222,11 +1220,11 @@ def next_run_datasets_summary(self, session: Session = NEW_SESSION): .where(DagModel.dataset_expression.is_not(None)) ).all() - dataset_triggered_next_run_info = get_dataset_triggered_next_run_info( + asset_triggered_next_run_info = get_asset_triggered_next_run_info( dataset_triggered_dag_ids, session=session ) - return flask.json.jsonify(dataset_triggered_next_run_info) + return flask.json.jsonify(asset_triggered_next_run_info) @expose("/dag_stats", methods=["POST"]) @auth.has_access_dag("GET", DagAccessEntity.RUN) @@ -3407,7 +3405,7 @@ def historical_metrics_data(self): @expose("/object/next_run_datasets/") @auth.has_access_dag("GET", DagAccessEntity.RUN) - @auth.has_access_dataset("GET") + @auth.has_access_asset("GET") @mark_fastapi_migration_done def next_run_datasets(self, dag_id): """Return datasets necessary, and their status, for the next dag run.""" @@ -3424,36 +3422,34 @@ def next_run_datasets(self, dag_id): dict(info._mapping) for info in session.execute( select( - DatasetModel.id, - DatasetModel.uri, - func.max(DatasetEvent.timestamp).label("lastUpdate"), - ) - .join( - DagScheduleDatasetReference, DagScheduleDatasetReference.dataset_id == DatasetModel.id + AssetModel.id, + AssetModel.uri, + func.max(AssetEvent.timestamp).label("lastUpdate"), ) + .join(DagScheduleAssetReference, DagScheduleAssetReference.dataset_id == AssetModel.id) .join( - DatasetDagRunQueue, + AssetDagRunQueue, and_( - DatasetDagRunQueue.dataset_id == DatasetModel.id, - DatasetDagRunQueue.target_dag_id == DagScheduleDatasetReference.dag_id, + AssetDagRunQueue.dataset_id == AssetModel.id, + AssetDagRunQueue.target_dag_id == DagScheduleAssetReference.dag_id, ), isouter=True, ) .join( - DatasetEvent, + AssetEvent, and_( - DatasetEvent.dataset_id == DatasetModel.id, + AssetEvent.dataset_id == AssetModel.id, ( - DatasetEvent.timestamp >= latest_run.execution_date + AssetEvent.timestamp >= latest_run.execution_date if latest_run and latest_run.execution_date else True ), ), isouter=True, ) - .where(DagScheduleDatasetReference.dag_id == dag_id, ~DatasetModel.is_orphaned) - .group_by(DatasetModel.id, DatasetModel.uri) - .order_by(DatasetModel.uri) + .where(DagScheduleAssetReference.dag_id == dag_id, ~AssetModel.is_orphaned) + .group_by(AssetModel.id, AssetModel.uri) + .order_by(AssetModel.uri) ) ] data = {"dataset_expression": dag_model.dataset_expression, "events": events} @@ -3473,7 +3469,7 @@ def dataset_dependencies(self): dag_node_id = f"dag:{dag}" if dag_node_id not in nodes_dict: for dep in dependencies: - if dep.dependency_type in ("dag", "dataset", "dataset-alias"): + if dep.dependency_type in ("dag", "asset", "asset-alias"): # add node nodes_dict[dag_node_id] = node_dict(dag_node_id, dag, "dag") if dep.node_id not in nodes_dict: @@ -3509,7 +3505,7 @@ def dataset_dependencies(self): ) @expose("/object/datasets_summary") - @auth.has_access_dataset("GET") + @auth.has_access_asset("GET") def datasets_summary(self): """ Get a summary of datasets. @@ -3543,54 +3539,54 @@ def datasets_summary(self): with create_session() as session: if lstripped_orderby == "uri": if order_by.startswith("-"): - order_by = (DatasetModel.uri.desc(),) + order_by = (AssetModel.uri.desc(),) else: - order_by = (DatasetModel.uri.asc(),) + order_by = (AssetModel.uri.asc(),) elif lstripped_orderby == "last_dataset_update": if order_by.startswith("-"): order_by = ( - func.max(DatasetEvent.timestamp).desc(), - DatasetModel.uri.asc(), + func.max(AssetEvent.timestamp).desc(), + AssetModel.uri.asc(), ) if session.bind.dialect.name == "postgresql": order_by = (order_by[0].nulls_last(), *order_by[1:]) else: order_by = ( - func.max(DatasetEvent.timestamp).asc(), - DatasetModel.uri.desc(), + func.max(AssetEvent.timestamp).asc(), + AssetModel.uri.desc(), ) if session.bind.dialect.name == "postgresql": order_by = (order_by[0].nulls_first(), *order_by[1:]) - count_query = select(func.count(DatasetModel.id)) + count_query = select(func.count(AssetModel.id)) has_event_filters = bool(updated_before or updated_after) query = ( select( - DatasetModel.id, - DatasetModel.uri, - func.max(DatasetEvent.timestamp).label("last_dataset_update"), - func.sum(case((DatasetEvent.id.is_not(None), 1), else_=0)).label("total_updates"), + AssetModel.id, + AssetModel.uri, + func.max(AssetEvent.timestamp).label("last_dataset_update"), + func.sum(case((AssetEvent.id.is_not(None), 1), else_=0)).label("total_updates"), ) - .join(DatasetEvent, DatasetEvent.dataset_id == DatasetModel.id, isouter=not has_event_filters) + .join(AssetEvent, AssetEvent.dataset_id == AssetModel.id, isouter=not has_event_filters) .group_by( - DatasetModel.id, - DatasetModel.uri, + AssetModel.id, + AssetModel.uri, ) .order_by(*order_by) ) if has_event_filters: - count_query = count_query.join(DatasetEvent, DatasetEvent.dataset_id == DatasetModel.id) + count_query = count_query.join(AssetEvent, AssetEvent.dataset_id == AssetModel.id) - filters = [~DatasetModel.is_orphaned] + filters = [~AssetModel.is_orphaned] if uri_pattern: - filters.append(DatasetModel.uri.ilike(f"%{uri_pattern}%")) + filters.append(AssetModel.uri.ilike(f"%{uri_pattern}%")) if updated_after: - filters.append(DatasetEvent.timestamp >= updated_after) + filters.append(AssetEvent.timestamp >= updated_after) if updated_before: - filters.append(DatasetEvent.timestamp <= updated_before) + filters.append(AssetEvent.timestamp <= updated_before) query = query.where(*filters).offset(offset).limit(limit) count_query = count_query.where(*filters) diff --git a/dev/breeze/tests/test_packages.py b/dev/breeze/tests/test_packages.py index 9556ae695be8e..39ee245cf36bd 100644 --- a/dev/breeze/tests/test_packages.py +++ b/dev/breeze/tests/test_packages.py @@ -165,6 +165,7 @@ def test_get_documentation_package_path(): "fab", "", """ + "apache-airflow-providers-common-compat>=1.2.0", "apache-airflow>=2.9.0", "flask-appbuilder==4.5.0", "flask-login>=0.6.2", @@ -178,6 +179,7 @@ def test_get_documentation_package_path(): "fab", "dev0", """ + "apache-airflow-providers-common-compat>=1.2.0.dev0", "apache-airflow>=2.9.0.dev0", "flask-appbuilder==4.5.0", "flask-login>=0.6.2", @@ -191,6 +193,7 @@ def test_get_documentation_package_path(): "fab", "beta0", """ + "apache-airflow-providers-common-compat>=1.2.0b0", "apache-airflow>=2.9.0b0", "flask-appbuilder==4.5.0", "flask-login>=0.6.2", diff --git a/dev/breeze/tests/test_pytest_args_for_test_types.py b/dev/breeze/tests/test_pytest_args_for_test_types.py index 36a4b157794d2..7ecbbf4b5bf3c 100644 --- a/dev/breeze/tests/test_pytest_args_for_test_types.py +++ b/dev/breeze/tests/test_pytest_args_for_test_types.py @@ -151,13 +151,13 @@ ( "Other", [ + "tests/assets", "tests/auth", "tests/callbacks", "tests/charts", "tests/cluster_policies", "tests/config_templates", "tests/dag_processing", - "tests/datasets", "tests/decorators", "tests/hooks", "tests/io", diff --git a/dev/breeze/tests/test_selective_checks.py b/dev/breeze/tests/test_selective_checks.py index 6161c44f6ebf3..4483cae573359 100644 --- a/dev/breeze/tests/test_selective_checks.py +++ b/dev/breeze/tests/test_selective_checks.py @@ -136,7 +136,7 @@ def assert_outputs_are_printed(expected_outputs: dict[str, str], stderr: str): pytest.param( ("airflow/api/file.py",), { - "affected-providers-list-as-string": "fab", + "affected-providers-list-as-string": "common.compat fab", "all-python-versions": "['3.8']", "all-python-versions-list-as-string": "3.8", "python-versions": "['3.8']", @@ -150,13 +150,13 @@ def assert_outputs_are_printed(expected_outputs: dict[str, str], stderr: str): "skip-pre-commits": "check-provider-yaml-valid,identity,lint-helm-chart,mypy-airflow,mypy-dev," "mypy-docs,mypy-providers,ts-compile-format-lint-ui,ts-compile-format-lint-www", "upgrade-to-newer-dependencies": "false", - "parallel-test-types-list-as-string": "API Always Providers[fab]", - "providers-test-types-list-as-string": "Providers[fab]", - "separate-test-types-list-as-string": "API Always Providers[fab]", + "parallel-test-types-list-as-string": "API Always Providers[common.compat,fab]", + "providers-test-types-list-as-string": "Providers[common.compat,fab]", + "separate-test-types-list-as-string": "API Always Providers[common.compat] Providers[fab]", "needs-mypy": "true", "mypy-folders": "['airflow']", }, - id="Only API tests and DOCS and FAB provider should run", + id="Only API tests and DOCS and common.compat, FAB providers should run", ) ), ( @@ -324,7 +324,7 @@ def assert_outputs_are_printed(expected_outputs: dict[str, str], stderr: str): "tests/providers/postgres/file.py", ), { - "affected-providers-list-as-string": "amazon common.sql fab google openlineage " + "affected-providers-list-as-string": "amazon common.compat common.sql fab google openlineage " "pgvector postgres", "all-python-versions": "['3.8']", "all-python-versions-list-as-string": "3.8", @@ -340,10 +340,10 @@ def assert_outputs_are_printed(expected_outputs: dict[str, str], stderr: str): "ts-compile-format-lint-ui,ts-compile-format-lint-www", "upgrade-to-newer-dependencies": "false", "parallel-test-types-list-as-string": "API Always Providers[amazon] " - "Providers[common.sql,fab,openlineage,pgvector,postgres] Providers[google]", + "Providers[common.compat,common.sql,fab,openlineage,pgvector,postgres] Providers[google]", "providers-test-types-list-as-string": "Providers[amazon] " - "Providers[common.sql,fab,openlineage,pgvector,postgres] Providers[google]", - "separate-test-types-list-as-string": "API Always Providers[amazon] Providers[common.sql] " + "Providers[common.compat,common.sql,fab,openlineage,pgvector,postgres] Providers[google]", + "separate-test-types-list-as-string": "API Always Providers[amazon] Providers[common.compat] Providers[common.sql] " "Providers[fab] Providers[google] Providers[openlineage] Providers[pgvector] " "Providers[postgres]", "needs-mypy": "true", @@ -1390,7 +1390,7 @@ def test_expected_output_pull_request_v2_7( "airflow/api/file.py", ), { - "affected-providers-list-as-string": "fab", + "affected-providers-list-as-string": "common.compat fab", "all-python-versions": "['3.8']", "all-python-versions-list-as-string": "3.8", "ci-image-build": "true", @@ -1398,17 +1398,17 @@ def test_expected_output_pull_request_v2_7( "needs-helm-tests": "false", "run-tests": "true", "docs-build": "true", - "docs-list-as-string": "apache-airflow fab", + "docs-list-as-string": "apache-airflow common.compat fab", "skip-pre-commits": "check-provider-yaml-valid,identity,lint-helm-chart,mypy-airflow,mypy-dev,mypy-docs,mypy-providers," "ts-compile-format-lint-ui,ts-compile-format-lint-www", "run-kubernetes-tests": "false", "upgrade-to-newer-dependencies": "false", "skip-provider-tests": "false", - "parallel-test-types-list-as-string": "API Always CLI Operators Providers[fab] WWW", + "parallel-test-types-list-as-string": "API Always CLI Operators Providers[common.compat,fab] WWW", "needs-mypy": "true", "mypy-folders": "['airflow']", }, - id="No providers tests except fab should run if only CLI/API/Operators/WWW file changed", + id="No providers tests except common.compat fab should run if only CLI/API/Operators/WWW file changed", ), pytest.param( ("airflow/models/test.py",), diff --git a/docs/apache-airflow-providers-amazon/auth-manager/manage/index.rst b/docs/apache-airflow-providers-amazon/auth-manager/manage/index.rst index 0a540b8d32b4a..359f3cfff040c 100644 --- a/docs/apache-airflow-providers-amazon/auth-manager/manage/index.rst +++ b/docs/apache-airflow-providers-amazon/auth-manager/manage/index.rst @@ -164,7 +164,7 @@ This is equivalent to the :doc:`Viewer role in Flask AppBuilder ` for details on how dataset URIs work. +See :doc:`documentation on assets ` for details on how asset URIs work. -.. airflow-dataset-schemes:: +.. airflow-asset-schemes:: :tags: None :header-separator: " diff --git a/docs/apache-airflow-providers/howto/create-custom-providers.rst b/docs/apache-airflow-providers/howto/create-custom-providers.rst index 70588d4532b4d..d95719e38bc6c 100644 --- a/docs/apache-airflow-providers/howto/create-custom-providers.rst +++ b/docs/apache-airflow-providers/howto/create-custom-providers.rst @@ -96,9 +96,9 @@ Exposing customized functionality to the Airflow's core: * ``filesystems`` - this field should contain the list of all the filesystem module names. See :doc:`apache-airflow:core-concepts/objectstorage` for description of the filesystems. -* ``dataset-uris`` - this field should contain the list of the URI schemes together with +* ``asset-uris`` - this field should contain the list of the URI schemes together with class names implementing normalization functions. - See :doc:`apache-airflow:authoring-and-scheduling/datasets` for description of the dataset URIs. + See :doc:`apache-airflow:authoring-and-scheduling/assets` for description of the asset URIs. .. note:: Deprecated values diff --git a/docs/apache-airflow/administration-and-deployment/lineage.rst b/docs/apache-airflow/administration-and-deployment/lineage.rst index de20dd8f1d802..d2ef63d869755 100644 --- a/docs/apache-airflow/administration-and-deployment/lineage.rst +++ b/docs/apache-airflow/administration-and-deployment/lineage.rst @@ -96,8 +96,8 @@ Airflow provides a powerful feature for tracking data lineage not only between t This functionality helps you understand how data flows throughout your Airflow pipelines. A global instance of ``HookLineageCollector`` serves as the central hub for collecting lineage information. -Hooks can send details about datasets they interact with to this collector. -The collector then uses this data to construct AIP-60 compliant Datasets, a standard format for describing datasets. +Hooks can send details about assets they interact with to this collector. +The collector then uses this data to construct AIP-60 compliant Assets, a standard format for describing assets. .. code-block:: python @@ -108,8 +108,8 @@ The collector then uses this data to construct AIP-60 compliant Datasets, a stan def run(self): # run actual code collector = get_hook_lineage_collector() - collector.add_input_dataset(self, dataset_kwargs={"scheme": "file", "path": "/tmp/in"}) - collector.add_output_dataset(self, dataset_kwargs={"scheme": "file", "path": "/tmp/out"}) + collector.add_input_asset(self, asset_kwargs={"scheme": "file", "path": "/tmp/in"}) + collector.add_output_asset(self, asset_kwargs={"scheme": "file", "path": "/tmp/out"}) Lineage data collected by the ``HookLineageCollector`` can be accessed using an instance of ``HookLineageReader``, which is registered in an Airflow plugin. @@ -122,7 +122,7 @@ which is registered in an Airflow plugin. class CustomHookLineageReader(HookLineageReader): def get_inputs(self): - return self.lineage_collector.collected_datasets.inputs + return self.lineage_collector.collected_assets.inputs class HookLineageCollectionPlugin(AirflowPlugin): @@ -130,7 +130,7 @@ which is registered in an Airflow plugin. hook_lineage_readers = [CustomHookLineageReader] If no ``HookLineageReader`` is registered within Airflow, a default ``NoOpCollector`` is used instead. -This collector does not create AIP-60 compliant datasets or collect lineage information. +This collector does not create AIP-60 compliant assets or collect lineage information. Lineage Backend diff --git a/docs/apache-airflow/administration-and-deployment/listeners.rst b/docs/apache-airflow/administration-and-deployment/listeners.rst index 4926b12ed6c6d..1fca915a6f1df 100644 --- a/docs/apache-airflow/administration-and-deployment/listeners.rst +++ b/docs/apache-airflow/administration-and-deployment/listeners.rst @@ -91,14 +91,14 @@ You can use these events to react to ``LocalTaskJob`` state changes. :end-before: [END howto_listen_ti_failure_task] -Dataset Events +Asset Events -------------- -- ``on_dataset_created`` +- ``on_asset_created`` - ``on_dataset_alias_created`` -- ``on_dataset_changed`` +- ``on_asset_changed`` -Dataset events occur when Dataset management operations are run. +Asset events occur when Asset management operations are run. Dag Import Error Events diff --git a/docs/apache-airflow/administration-and-deployment/logging-monitoring/metrics.rst b/docs/apache-airflow/administration-and-deployment/logging-monitoring/metrics.rst index 61985cecea9b0..ac44d1acba9c0 100644 --- a/docs/apache-airflow/administration-and-deployment/logging-monitoring/metrics.rst +++ b/docs/apache-airflow/administration-and-deployment/logging-monitoring/metrics.rst @@ -201,10 +201,10 @@ Name Descripti fully asynchronous) ``triggers.failed`` Number of triggers that errored before they could fire an event ``triggers.succeeded`` Number of triggers that have fired at least one event -``dataset.updates`` Number of updated datasets -``dataset.orphaned`` Number of datasets marked as orphans because they are no longer referenced in DAG +``asset.updates`` Number of updated assets +``asset.orphaned`` Number of assets marked as orphans because they are no longer referenced in DAG schedule parameters or task outlets -``dataset.triggered_dagruns`` Number of DAG runs triggered by a dataset update +``asset.triggered_dagruns`` Number of DAG runs triggered by a asset update ====================================================================== ================================================================ Gauges diff --git a/docs/apache-airflow/authoring-and-scheduling/assets.rst b/docs/apache-airflow/authoring-and-scheduling/assets.rst new file mode 100644 index 0000000000000..d37143367fabe --- /dev/null +++ b/docs/apache-airflow/authoring-and-scheduling/assets.rst @@ -0,0 +1,532 @@ + .. 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. + +Data-aware scheduling +===================== + +.. versionadded:: 2.4 + +Quickstart +---------- + +In addition to scheduling DAGs based on time, you can also schedule DAGs to run based on when a task updates a asset. + +.. code-block:: python + + from airflow.assets import asset + + with DAG(...): + MyOperator( + # this task updates example.csv + outlets=[asset("s3://asset-bucket/example.csv")], + ..., + ) + + + with DAG( + # this DAG should be run when example.csv is updated (by dag1) + schedule=[asset("s3://asset-bucket/example.csv")], + ..., + ): + ... + + +.. image:: /img/asset-scheduled-dags.png + + +What is a "asset"? +-------------------- + +An Airflow asset is a logical grouping of data. Upstream producer tasks can update assets, and asset updates contribute to scheduling downstream consumer DAGs. + +`Uniform Resource Identifier (URI) `_ define assets: + +.. code-block:: python + + from airflow.assets import asset + + example_asset = asset("s3://asset-bucket/example.csv") + +Airflow makes no assumptions about the content or location of the data represented by the URI, and treats the URI like a string. This means that Airflow treats any regular expressions, like ``input_\d+.csv``, or file glob patterns, such as ``input_2022*.csv``, as an attempt to create multiple assets from one declaration, and they will not work. + +You must create assets with a valid URI. Airflow core and providers define various URI schemes that you can use, such as ``file`` (core), ``postgres`` (by the Postgres provider), and ``s3`` (by the Amazon provider). Third-party providers and plugins might also provide their own schemes. These pre-defined schemes have individual semantics that are expected to be followed. + +What is valid URI? +------------------ + +Technically, the URI must conform to the valid character set in RFC 3986, which is basically ASCII alphanumeric characters, plus ``%``, ``-``, ``_``, ``.``, and ``~``. To identify a resource that cannot be represented by URI-safe characters, encode the resource name with `percent-encoding `_. + +The URI is also case sensitive, so ``s3://example/asset`` and ``s3://Example/asset`` are considered different. Note that the *host* part of the URI is also case sensitive, which differs from RFC 3986. + +Do not use the ``airflow`` scheme, which is is reserved for Airflow's internals. + +Airflow always prefers using lower cases in schemes, and case sensitivity is needed in the host part of the URI to correctly distinguish between resources. + +.. code-block:: python + + # invalid assets: + reserved = asset("airflow://example_asset") + not_ascii = asset("èxample_datašet") + +If you want to define assets with a scheme that doesn't include additional semantic constraints, use a scheme with the prefix ``x-``. Airflow skips any semantic validation on URIs with these schemes. + +.. code-block:: python + + # valid asset, treated as a plain string + my_ds = asset("x-my-thing://foobarbaz") + +The identifier does not have to be absolute; it can be a scheme-less, relative URI, or even just a simple path or string: + +.. code-block:: python + + # valid assets: + schemeless = asset("//example/asset") + csv_file = asset("example_asset") + +Non-absolute identifiers are considered plain strings that do not carry any semantic meanings to Airflow. + +Extra information on asset +---------------------------- + +If needed, you can include an extra dictionary in a asset: + +.. code-block:: python + + example_asset = asset( + "s3://asset/example.csv", + extra={"team": "trainees"}, + ) + +This can be used to supply custom description to the asset, such as who has ownership to the target file, or what the file is for. The extra information does not affect a asset's identity. This means a DAG will be triggered by a asset with an identical URI, even if the extra dict is different: + +.. code-block:: python + + with DAG( + dag_id="consumer", + schedule=[asset("s3://asset/example.csv", extra={"different": "extras"})], + ): + ... + + with DAG(dag_id="producer", ...): + MyOperator( + # triggers "consumer" with the given extra! + outlets=[asset("s3://asset/example.csv", extra={"team": "trainees"})], + ..., + ) + +.. note:: **Security Note:** asset URI and extra fields are not encrypted, they are stored in cleartext in Airflow's metadata database. Do NOT store any sensitive values, especially credentials, in either asset URIs or extra key values! + +How to use assets in your DAGs +-------------------------------- + +You can use assets to specify data dependencies in your DAGs. The following example shows how after the ``producer`` task in the ``producer`` DAG successfully completes, Airflow schedules the ``consumer`` DAG. Airflow marks a asset as ``updated`` only if the task completes successfully. If the task fails or if it is skipped, no update occurs, and Airflow doesn't schedule the ``consumer`` DAG. + +.. code-block:: python + + example_asset = asset("s3://asset/example.csv") + + with DAG(dag_id="producer", ...): + BashOperator(task_id="producer", outlets=[example_asset], ...) + + with DAG(dag_id="consumer", schedule=[example_asset], ...): + ... + + +You can find a listing of the relationships between assets and DAGs in the +:ref:`assets View` + +Multiple assets +----------------- + +Because the ``schedule`` parameter is a list, DAGs can require multiple assets. Airflow schedules a DAG after **all** assets the DAG consumes have been updated at least once since the last time the DAG ran: + +.. code-block:: python + + with DAG( + dag_id="multiple_assets_example", + schedule=[ + example_asset_1, + example_asset_2, + example_asset_3, + ], + ..., + ): + ... + + +If one asset is updated multiple times before all consumed assets update, the downstream DAG still only runs once, as shown in this illustration: + +.. :: + ASCII art representation of this diagram + + example_asset_1 x----x---x---x----------------------x- + example_asset_2 -------x---x-------x------x----x------ + example_asset_3 ---------------x-----x------x--------- + DAG runs created * * + +.. graphviz:: + + graph asset_event_timeline { + graph [layout=neato] + { + node [margin=0 fontcolor=blue width=0.1 shape=point label=""] + e1 [pos="1,2.5!"] + e2 [pos="2,2.5!"] + e3 [pos="2.5,2!"] + e4 [pos="4,2.5!"] + e5 [pos="5,2!"] + e6 [pos="6,2.5!"] + e7 [pos="7,1.5!"] + r7 [pos="7,1!" shape=star width=0.25 height=0.25 fixedsize=shape] + e8 [pos="8,2!"] + e9 [pos="9,1.5!"] + e10 [pos="10,2!"] + e11 [pos="11,1.5!"] + e12 [pos="12,2!"] + e13 [pos="13,2.5!"] + r13 [pos="13,1!" shape=star width=0.25 height=0.25 fixedsize=shape] + } + { + node [shape=none label="" width=0] + end_ds1 [pos="14,2.5!"] + end_ds2 [pos="14,2!"] + end_ds3 [pos="14,1.5!"] + } + + { + node [shape=none margin=0.25 fontname="roboto,sans-serif"] + example_asset_1 [ pos="-0.5,2.5!"] + example_asset_2 [ pos="-0.5,2!"] + example_asset_3 [ pos="-0.5,1.5!"] + dag_runs [label="DagRuns created" pos="-0.5,1!"] + } + + edge [color=lightgrey] + + example_asset_1 -- e1 -- e2 -- e4 -- e6 -- e13 -- end_ds1 + example_asset_2 -- e3 -- e5 -- e8 -- e10 -- e12 -- end_ds2 + example_asset_3 -- e7 -- e9 -- e11 -- end_ds3 + + } + +Attaching extra information to an emitting asset event +-------------------------------------------------------- + +.. versionadded:: 2.10.0 + +A task with a asset outlet can optionally attach extra information before it emits a asset event. This is different +from `Extra information on asset`_. Extra information on a asset statically describes the entity pointed to by the asset URI; extra information on the *asset event* instead should be used to annotate the triggering data change, such as how many rows in the database are changed by the update, or the date range covered by it. + +The easiest way to attach extra information to the asset event is by ``yield``-ing a ``Metadata`` object from a task: + +.. code-block:: python + + from airflow.assets import asset + from airflow.assets.metadata import Metadata + + example_s3_asset = asset("s3://asset/example.csv") + + + @task(outlets=[example_s3_asset]) + def write_to_s3(): + df = ... # Get a Pandas DataFrame to write. + # Write df to asset... + yield Metadata(example_s3_asset, {"row_count": len(df)}) + +Airflow automatically collects all yielded metadata, and populates asset events with extra information for corresponding metadata objects. + +This can also be done in classic operators. The best way is to subclass the operator and override ``execute``. Alternatively, extras can also be added in a task's ``pre_execute`` or ``post_execute`` hook. If you choose to use hooks, however, remember that they are not rerun when a task is retried, and may cause the extra information to not match actual data in certain scenarios. + +Another way to achieve the same is by accessing ``outlet_events`` in a task's execution context directly: + +.. code-block:: python + + @task(outlets=[example_s3_asset]) + def write_to_s3(*, outlet_events): + outlet_events[example_s3_asset].extra = {"row_count": len(df)} + +There's minimal magic here---Airflow simply writes the yielded values to the exact same accessor. This also works in classic operators, including ``execute``, ``pre_execute``, and ``post_execute``. + +.. _fetching_information_from_previously_emitted_asset_events: + +Fetching information from previously emitted asset events +----------------------------------------------------------- + +.. versionadded:: 2.10.0 + +Events of a asset defined in a task's ``outlets``, as described in the previous section, can be read by a task that declares the same asset in its ``inlets``. A asset event entry contains ``extra`` (see previous section for details), ``timestamp`` indicating when the event was emitted from a task, and ``source_task_instance`` linking the event back to its source. + +Inlet asset events can be read with the ``inlet_events`` accessor in the execution context. Continuing from the ``write_to_s3`` task in the previous section: + +.. code-block:: python + + @task(inlets=[example_s3_asset]) + def post_process_s3_file(*, inlet_events): + events = inlet_events[example_s3_asset] + last_row_count = events[-1].extra["row_count"] + +Each value in the ``inlet_events`` mapping is a sequence-like object that orders past events of a given asset by ``timestamp``, earliest to latest. It supports most of Python's list interface, so you can use ``[-1]`` to access the last event, ``[-2:]`` for the last two, etc. The accessor is lazy and only hits the database when you access items inside it. + + +Fetching information from a triggering asset event +---------------------------------------------------- + +A triggered DAG can fetch information from the asset that triggered it using the ``triggering_asset_events`` template or parameter. See more at :ref:`templates-ref`. + +Example: + +.. code-block:: python + + example_snowflake_asset = asset("snowflake://my_db/my_schema/my_table") + + with DAG(dag_id="load_snowflake_data", schedule="@hourly", ...): + SQLExecuteQueryOperator( + task_id="load", conn_id="snowflake_default", outlets=[example_snowflake_asset], ... + ) + + with DAG(dag_id="query_snowflake_data", schedule=[example_snowflake_asset], ...): + SQLExecuteQueryOperator( + task_id="query", + conn_id="snowflake_default", + sql=""" + SELECT * + FROM my_db.my_schema.my_table + WHERE "updated_at" >= '{{ (triggering_asset_events.values() | first | first).source_dag_run.data_interval_start }}' + AND "updated_at" < '{{ (triggering_asset_events.values() | first | first).source_dag_run.data_interval_end }}'; + """, + ) + + @task + def print_triggering_asset_events(triggering_asset_events=None): + for asset, asset_list in triggering_asset_events.items(): + print(asset, asset_list) + print(asset_list[0].source_dag_run.dag_id) + + print_triggering_asset_events() + +Note that this example is using `(.values() | first | first) `_ to fetch the first of one asset given to the DAG, and the first of one AssetEvent for that asset. An implementation can be quite complex if you have multiple assets, potentially with multiple AssetEvents. + + +Manipulating queued asset events through REST API +--------------------------------------------------- + +.. versionadded:: 2.9 + +In this example, the DAG ``waiting_for_asset_1_and_2`` will be triggered when tasks update both assets "asset-1" and "asset-2". Once "asset-1" is updated, Airflow creates a record. This ensures that Airflow knows to trigger the DAG when "asset-2" is updated. We call such records queued asset events. + +.. code-block:: python + + with DAG( + dag_id="waiting_for_asset_1_and_2", + schedule=[asset("asset-1"), asset("asset-2")], + ..., + ): + ... + + +``queuedEvent`` API endpoints are introduced to manipulate such records. + +* Get a queued asset event for a DAG: ``/assets/queuedEvent/{uri}`` +* Get queued asset events for a DAG: ``/dags/{dag_id}/assets/queuedEvent`` +* Delete a queued asset event for a DAG: ``/assets/queuedEvent/{uri}`` +* Delete queued asset events for a DAG: ``/dags/{dag_id}/assets/queuedEvent`` +* Get queued asset events for a asset: ``/dags/{dag_id}/assets/queuedEvent/{uri}`` +* Delete queued asset events for a asset: ``DELETE /dags/{dag_id}/assets/queuedEvent/{uri}`` + + For how to use REST API and the parameters needed for these endpoints, please refer to :doc:`Airflow API `. + +Advanced asset scheduling with conditional expressions +-------------------------------------------------------- + +Apache Airflow includes advanced scheduling capabilities that use conditional expressions with assets. This feature allows you to define complex dependencies for DAG executions based on asset updates, using logical operators for more control on workflow triggers. + +Logical operators for assets +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +Airflow supports two logical operators for combining asset conditions: + +- **AND (``&``)**: Specifies that the DAG should be triggered only after all of the specified assets have been updated. +- **OR (``|``)**: Specifies that the DAG should be triggered when any of the specified assets is updated. + +These operators enable you to configure your Airflow workflows to use more complex asset update conditions, making them more dynamic and flexible. + +Example Use +------------- + +**Scheduling based on multiple asset updates** + +To schedule a DAG to run only when two specific assets have both been updated, use the AND operator (``&``): + +.. code-block:: python + + dag1_asset = asset("s3://dag1/output_1.txt") + dag2_asset = asset("s3://dag2/output_1.txt") + + with DAG( + # Consume asset 1 and 2 with asset expressions + schedule=(dag1_asset & dag2_asset), + ..., + ): + ... + +**Scheduling based on any asset update** + +To trigger a DAG execution when either one of two assets is updated, apply the OR operator (``|``): + +.. code-block:: python + + with DAG( + # Consume asset 1 or 2 with asset expressions + schedule=(dag1_asset | dag2_asset), + ..., + ): + ... + +**Complex Conditional Logic** + +For scenarios requiring more intricate conditions, such as triggering a DAG when one asset is updated or when both of two other assets are updated, combine the OR and AND operators: + +.. code-block:: python + + dag3_asset = asset("s3://dag3/output_3.txt") + + with DAG( + # Consume asset 1 or both 2 and 3 with asset expressions + schedule=(dag1_asset | (dag2_asset & dag3_asset)), + ..., + ): + ... + + +Dynamic data events emitting and asset creation through AssetAlias +----------------------------------------------------------------------- +An asset alias can be used to emit asset events of assets with association to the aliases. Downstreams can depend on resolved asset. This feature allows you to define complex dependencies for DAG executions based on asset updates. + +How to use AssetAlias +~~~~~~~~~~~~~~~~~~~~~~~ + +``AssetAlias`` has one single argument ``name`` that uniquely identifies the asset. The task must first declare the alias as an outlet, and use ``outlet_events`` or yield ``Metadata`` to add events to it. + +The following example creates a asset event against the S3 URI ``f"s3://bucket/my-task"`` with optional extra information ``extra``. If the asset does not exist, Airflow will dynamically create it and log a warning message. + +**Emit a asset event during task execution through outlet_events** + +.. code-block:: python + + from airflow.assets import AssetAlias + + + @task(outlets=[AssetAlias("my-task-outputs")]) + def my_task_with_outlet_events(*, outlet_events): + outlet_events["my-task-outputs"].add(asset("s3://bucket/my-task"), extra={"k": "v"}) + + +**Emit a asset event during task execution through yielding Metadata** + +.. code-block:: python + + from airflow.assets.metadata import Metadata + + + @task(outlets=[AssetAlias("my-task-outputs")]) + def my_task_with_metadata(): + s3_asset = asset("s3://bucket/my-task") + yield Metadata(s3_asset, extra={"k": "v"}, alias="my-task-outputs") + +Only one asset event is emitted for an added asset, even if it is added to the alias multiple times, or added to multiple aliases. However, if different ``extra`` values are passed, it can emit multiple asset events. In the following example, two asset events will be emitted. + +.. code-block:: python + + from airflow.assets import AssetAlias + + + @task( + outlets=[ + AssetAlias("my-task-outputs-1"), + AssetAlias("my-task-outputs-2"), + AssetAlias("my-task-outputs-3"), + ] + ) + def my_task_with_outlet_events(*, outlet_events): + outlet_events["my-task-outputs-1"].add(asset("s3://bucket/my-task"), extra={"k": "v"}) + # This line won't emit an additional asset event as the asset and extra are the same as the previous line. + outlet_events["my-task-outputs-2"].add(asset("s3://bucket/my-task"), extra={"k": "v"}) + # This line will emit an additional asset event as the extra is different. + outlet_events["my-task-outputs-3"].add(asset("s3://bucket/my-task"), extra={"k2": "v2"}) + +Scheduling based on asset aliases +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +Since asset events added to an alias are just simple asset events, a downstream DAG depending on the actual asset can read asset events of it normally, without considering the associated aliases. A downstream DAG can also depend on an asset alias. The authoring syntax is referencing the ``AssetAlias`` by name, and the associated asset events are picked up for scheduling. Note that a DAG can be triggered by a task with ``outlets=AssetAlias("xxx")`` if and only if the alias is resolved into ``asset("s3://bucket/my-task")``. The DAG runs whenever a task with outlet ``AssetAlias("out")`` gets associated with at least one asset at runtime, regardless of the asset's identity. The downstream DAG is not triggered if no assets are associated to the alias for a particular given task run. This also means we can do conditional asset-triggering. + +The asset alias is resolved to the assets during DAG parsing. Thus, if the "min_file_process_interval" configuration is set to a high value, there is a possibility that the asset alias may not be resolved. To resolve this issue, you can trigger DAG parsing. + +.. code-block:: python + + with DAG(dag_id="asset-producer"): + + @task(outlets=[asset("example-alias")]) + def produce_asset_events(): + pass + + + with DAG(dag_id="asset-alias-producer"): + + @task(outlets=[AssetAlias("example-alias")]) + def produce_asset_events(*, outlet_events): + outlet_events["example-alias"].add(asset("s3://bucket/my-task")) + + + with DAG(dag_id="asset-consumer", schedule=asset("s3://bucket/my-task")): + ... + + with DAG(dag_id="asset-alias-consumer", schedule=AssetAlias("example-alias")): + ... + + +In the example provided, once the DAG ``asset-alias-producer`` is executed, the asset alias ``AssetAlias("example-alias")`` will be resolved to ``asset("s3://bucket/my-task")``. However, the DAG ``asset-alias-consumer`` will have to wait for the next DAG re-parsing to update its schedule. To address this, Airflow will re-parse the DAGs relying on the asset alias ``AssetAlias("example-alias")`` when it's resolved into assets that these DAGs did not previously depend on. As a result, both the "asset-consumer" and "asset-alias-consumer" DAGs will be triggered after the execution of DAG ``asset-alias-producer``. + + +Fetching information from previously emitted asset events through resolved asset aliases +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +As mentioned in :ref:`Fetching information from previously emitted asset events`, inlet asset events can be read with the ``inlet_events`` accessor in the execution context, and you can also use asset aliases to access the asset events triggered by them. + +.. code-block:: python + + with DAG(dag_id="asset-alias-producer"): + + @task(outlets=[AssetAlias("example-alias")]) + def produce_asset_events(*, outlet_events): + outlet_events["example-alias"].add(asset("s3://bucket/my-task"), extra={"row_count": 1}) + + + with DAG(dag_id="asset-alias-consumer", schedule=None): + + @task(inlets=[AssetAlias("example-alias")]) + def consume_asset_alias_events(*, inlet_events): + events = inlet_events[AssetAlias("example-alias")] + last_row_count = events[-1].extra["row_count"] + + +Combining asset and time-based schedules +------------------------------------------ + +AssetTimetable Integration +~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +You can schedule DAGs based on both asset events and time-based schedules using ``AssetOrTimeSchedule``. This allows you to create workflows when a DAG needs both to be triggered by data updates and run periodically according to a fixed timetable. + +For more detailed information on ``AssetOrTimeSchedule``, refer to the corresponding section in :ref:`AssetOrTimeSchedule `. diff --git a/docs/apache-airflow/authoring-and-scheduling/datasets.rst b/docs/apache-airflow/authoring-and-scheduling/datasets.rst deleted file mode 100644 index a69c09bc13b0f..0000000000000 --- a/docs/apache-airflow/authoring-and-scheduling/datasets.rst +++ /dev/null @@ -1,532 +0,0 @@ - .. 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. - -Data-aware scheduling -===================== - -.. versionadded:: 2.4 - -Quickstart ----------- - -In addition to scheduling DAGs based on time, you can also schedule DAGs to run based on when a task updates a dataset. - -.. code-block:: python - - from airflow.datasets import Dataset - - with DAG(...): - MyOperator( - # this task updates example.csv - outlets=[Dataset("s3://dataset-bucket/example.csv")], - ..., - ) - - - with DAG( - # this DAG should be run when example.csv is updated (by dag1) - schedule=[Dataset("s3://dataset-bucket/example.csv")], - ..., - ): - ... - - -.. image:: /img/dataset-scheduled-dags.png - - -What is a "dataset"? --------------------- - -An Airflow dataset is a logical grouping of data. Upstream producer tasks can update datasets, and dataset updates contribute to scheduling downstream consumer DAGs. - -`Uniform Resource Identifier (URI) `_ define datasets: - -.. code-block:: python - - from airflow.datasets import Dataset - - example_dataset = Dataset("s3://dataset-bucket/example.csv") - -Airflow makes no assumptions about the content or location of the data represented by the URI, and treats the URI like a string. This means that Airflow treats any regular expressions, like ``input_\d+.csv``, or file glob patterns, such as ``input_2022*.csv``, as an attempt to create multiple datasets from one declaration, and they will not work. - -You must create datasets with a valid URI. Airflow core and providers define various URI schemes that you can use, such as ``file`` (core), ``postgres`` (by the Postgres provider), and ``s3`` (by the Amazon provider). Third-party providers and plugins might also provide their own schemes. These pre-defined schemes have individual semantics that are expected to be followed. - -What is valid URI? ------------------- - -Technically, the URI must conform to the valid character set in RFC 3986, which is basically ASCII alphanumeric characters, plus ``%``, ``-``, ``_``, ``.``, and ``~``. To identify a resource that cannot be represented by URI-safe characters, encode the resource name with `percent-encoding `_. - -The URI is also case sensitive, so ``s3://example/dataset`` and ``s3://Example/Dataset`` are considered different. Note that the *host* part of the URI is also case sensitive, which differs from RFC 3986. - -Do not use the ``airflow`` scheme, which is is reserved for Airflow's internals. - -Airflow always prefers using lower cases in schemes, and case sensitivity is needed in the host part of the URI to correctly distinguish between resources. - -.. code-block:: python - - # invalid datasets: - reserved = Dataset("airflow://example_dataset") - not_ascii = Dataset("èxample_datašet") - -If you want to define datasets with a scheme that doesn't include additional semantic constraints, use a scheme with the prefix ``x-``. Airflow skips any semantic validation on URIs with these schemes. - -.. code-block:: python - - # valid dataset, treated as a plain string - my_ds = Dataset("x-my-thing://foobarbaz") - -The identifier does not have to be absolute; it can be a scheme-less, relative URI, or even just a simple path or string: - -.. code-block:: python - - # valid datasets: - schemeless = Dataset("//example/dataset") - csv_file = Dataset("example_dataset") - -Non-absolute identifiers are considered plain strings that do not carry any semantic meanings to Airflow. - -Extra information on dataset ----------------------------- - -If needed, you can include an extra dictionary in a dataset: - -.. code-block:: python - - example_dataset = Dataset( - "s3://dataset/example.csv", - extra={"team": "trainees"}, - ) - -This can be used to supply custom description to the dataset, such as who has ownership to the target file, or what the file is for. The extra information does not affect a dataset's identity. This means a DAG will be triggered by a dataset with an identical URI, even if the extra dict is different: - -.. code-block:: python - - with DAG( - dag_id="consumer", - schedule=[Dataset("s3://dataset/example.csv", extra={"different": "extras"})], - ): - ... - - with DAG(dag_id="producer", ...): - MyOperator( - # triggers "consumer" with the given extra! - outlets=[Dataset("s3://dataset/example.csv", extra={"team": "trainees"})], - ..., - ) - -.. note:: **Security Note:** Dataset URI and extra fields are not encrypted, they are stored in cleartext in Airflow's metadata database. Do NOT store any sensitive values, especially credentials, in either dataset URIs or extra key values! - -How to use datasets in your DAGs --------------------------------- - -You can use datasets to specify data dependencies in your DAGs. The following example shows how after the ``producer`` task in the ``producer`` DAG successfully completes, Airflow schedules the ``consumer`` DAG. Airflow marks a dataset as ``updated`` only if the task completes successfully. If the task fails or if it is skipped, no update occurs, and Airflow doesn't schedule the ``consumer`` DAG. - -.. code-block:: python - - example_dataset = Dataset("s3://dataset/example.csv") - - with DAG(dag_id="producer", ...): - BashOperator(task_id="producer", outlets=[example_dataset], ...) - - with DAG(dag_id="consumer", schedule=[example_dataset], ...): - ... - - -You can find a listing of the relationships between datasets and DAGs in the -:ref:`Datasets View` - -Multiple Datasets ------------------ - -Because the ``schedule`` parameter is a list, DAGs can require multiple datasets. Airflow schedules a DAG after **all** datasets the DAG consumes have been updated at least once since the last time the DAG ran: - -.. code-block:: python - - with DAG( - dag_id="multiple_datasets_example", - schedule=[ - example_dataset_1, - example_dataset_2, - example_dataset_3, - ], - ..., - ): - ... - - -If one dataset is updated multiple times before all consumed datasets update, the downstream DAG still only runs once, as shown in this illustration: - -.. :: - ASCII art representation of this diagram - - example_dataset_1 x----x---x---x----------------------x- - example_dataset_2 -------x---x-------x------x----x------ - example_dataset_3 ---------------x-----x------x--------- - DAG runs created * * - -.. graphviz:: - - graph dataset_event_timeline { - graph [layout=neato] - { - node [margin=0 fontcolor=blue width=0.1 shape=point label=""] - e1 [pos="1,2.5!"] - e2 [pos="2,2.5!"] - e3 [pos="2.5,2!"] - e4 [pos="4,2.5!"] - e5 [pos="5,2!"] - e6 [pos="6,2.5!"] - e7 [pos="7,1.5!"] - r7 [pos="7,1!" shape=star width=0.25 height=0.25 fixedsize=shape] - e8 [pos="8,2!"] - e9 [pos="9,1.5!"] - e10 [pos="10,2!"] - e11 [pos="11,1.5!"] - e12 [pos="12,2!"] - e13 [pos="13,2.5!"] - r13 [pos="13,1!" shape=star width=0.25 height=0.25 fixedsize=shape] - } - { - node [shape=none label="" width=0] - end_ds1 [pos="14,2.5!"] - end_ds2 [pos="14,2!"] - end_ds3 [pos="14,1.5!"] - } - - { - node [shape=none margin=0.25 fontname="roboto,sans-serif"] - example_dataset_1 [ pos="-0.5,2.5!"] - example_dataset_2 [ pos="-0.5,2!"] - example_dataset_3 [ pos="-0.5,1.5!"] - dag_runs [label="DagRuns created" pos="-0.5,1!"] - } - - edge [color=lightgrey] - - example_dataset_1 -- e1 -- e2 -- e4 -- e6 -- e13 -- end_ds1 - example_dataset_2 -- e3 -- e5 -- e8 -- e10 -- e12 -- end_ds2 - example_dataset_3 -- e7 -- e9 -- e11 -- end_ds3 - - } - -Attaching extra information to an emitting dataset event --------------------------------------------------------- - -.. versionadded:: 2.10.0 - -A task with a dataset outlet can optionally attach extra information before it emits a dataset event. This is different -from `Extra information on dataset`_. Extra information on a dataset statically describes the entity pointed to by the dataset URI; extra information on the *dataset event* instead should be used to annotate the triggering data change, such as how many rows in the database are changed by the update, or the date range covered by it. - -The easiest way to attach extra information to the dataset event is by ``yield``-ing a ``Metadata`` object from a task: - -.. code-block:: python - - from airflow.datasets import Dataset - from airflow.datasets.metadata import Metadata - - example_s3_dataset = Dataset("s3://dataset/example.csv") - - - @task(outlets=[example_s3_dataset]) - def write_to_s3(): - df = ... # Get a Pandas DataFrame to write. - # Write df to dataset... - yield Metadata(example_s3_dataset, {"row_count": len(df)}) - -Airflow automatically collects all yielded metadata, and populates dataset events with extra information for corresponding metadata objects. - -This can also be done in classic operators. The best way is to subclass the operator and override ``execute``. Alternatively, extras can also be added in a task's ``pre_execute`` or ``post_execute`` hook. If you choose to use hooks, however, remember that they are not rerun when a task is retried, and may cause the extra information to not match actual data in certain scenarios. - -Another way to achieve the same is by accessing ``outlet_events`` in a task's execution context directly: - -.. code-block:: python - - @task(outlets=[example_s3_dataset]) - def write_to_s3(*, outlet_events): - outlet_events[example_s3_dataset].extra = {"row_count": len(df)} - -There's minimal magic here---Airflow simply writes the yielded values to the exact same accessor. This also works in classic operators, including ``execute``, ``pre_execute``, and ``post_execute``. - -.. _fetching_information_from_previously_emitted_dataset_events: - -Fetching information from previously emitted dataset events ------------------------------------------------------------ - -.. versionadded:: 2.10.0 - -Events of a dataset defined in a task's ``outlets``, as described in the previous section, can be read by a task that declares the same dataset in its ``inlets``. A dataset event entry contains ``extra`` (see previous section for details), ``timestamp`` indicating when the event was emitted from a task, and ``source_task_instance`` linking the event back to its source. - -Inlet dataset events can be read with the ``inlet_events`` accessor in the execution context. Continuing from the ``write_to_s3`` task in the previous section: - -.. code-block:: python - - @task(inlets=[example_s3_dataset]) - def post_process_s3_file(*, inlet_events): - events = inlet_events[example_s3_dataset] - last_row_count = events[-1].extra["row_count"] - -Each value in the ``inlet_events`` mapping is a sequence-like object that orders past events of a given dataset by ``timestamp``, earliest to latest. It supports most of Python's list interface, so you can use ``[-1]`` to access the last event, ``[-2:]`` for the last two, etc. The accessor is lazy and only hits the database when you access items inside it. - - -Fetching information from a triggering dataset event ----------------------------------------------------- - -A triggered DAG can fetch information from the dataset that triggered it using the ``triggering_dataset_events`` template or parameter. See more at :ref:`templates-ref`. - -Example: - -.. code-block:: python - - example_snowflake_dataset = Dataset("snowflake://my_db/my_schema/my_table") - - with DAG(dag_id="load_snowflake_data", schedule="@hourly", ...): - SQLExecuteQueryOperator( - task_id="load", conn_id="snowflake_default", outlets=[example_snowflake_dataset], ... - ) - - with DAG(dag_id="query_snowflake_data", schedule=[example_snowflake_dataset], ...): - SQLExecuteQueryOperator( - task_id="query", - conn_id="snowflake_default", - sql=""" - SELECT * - FROM my_db.my_schema.my_table - WHERE "updated_at" >= '{{ (triggering_dataset_events.values() | first | first).source_dag_run.data_interval_start }}' - AND "updated_at" < '{{ (triggering_dataset_events.values() | first | first).source_dag_run.data_interval_end }}'; - """, - ) - - @task - def print_triggering_dataset_events(triggering_dataset_events=None): - for dataset, dataset_list in triggering_dataset_events.items(): - print(dataset, dataset_list) - print(dataset_list[0].source_dag_run.dag_id) - - print_triggering_dataset_events() - -Note that this example is using `(.values() | first | first) `_ to fetch the first of one dataset given to the DAG, and the first of one DatasetEvent for that dataset. An implementation can be quite complex if you have multiple datasets, potentially with multiple DatasetEvents. - - -Manipulating queued dataset events through REST API ---------------------------------------------------- - -.. versionadded:: 2.9 - -In this example, the DAG ``waiting_for_dataset_1_and_2`` will be triggered when tasks update both datasets "dataset-1" and "dataset-2". Once "dataset-1" is updated, Airflow creates a record. This ensures that Airflow knows to trigger the DAG when "dataset-2" is updated. We call such records queued dataset events. - -.. code-block:: python - - with DAG( - dag_id="waiting_for_dataset_1_and_2", - schedule=[Dataset("dataset-1"), Dataset("dataset-2")], - ..., - ): - ... - - -``queuedEvent`` API endpoints are introduced to manipulate such records. - -* Get a queued Dataset event for a DAG: ``/datasets/queuedEvent/{uri}`` -* Get queued Dataset events for a DAG: ``/dags/{dag_id}/datasets/queuedEvent`` -* Delete a queued Dataset event for a DAG: ``/datasets/queuedEvent/{uri}`` -* Delete queued Dataset events for a DAG: ``/dags/{dag_id}/datasets/queuedEvent`` -* Get queued Dataset events for a Dataset: ``/dags/{dag_id}/datasets/queuedEvent/{uri}`` -* Delete queued Dataset events for a Dataset: ``DELETE /dags/{dag_id}/datasets/queuedEvent/{uri}`` - - For how to use REST API and the parameters needed for these endpoints, please refer to :doc:`Airflow API `. - -Advanced dataset scheduling with conditional expressions --------------------------------------------------------- - -Apache Airflow includes advanced scheduling capabilities that use conditional expressions with datasets. This feature allows you to define complex dependencies for DAG executions based on dataset updates, using logical operators for more control on workflow triggers. - -Logical operators for datasets -~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ - -Airflow supports two logical operators for combining dataset conditions: - -- **AND (``&``)**: Specifies that the DAG should be triggered only after all of the specified datasets have been updated. -- **OR (``|``)**: Specifies that the DAG should be triggered when any of the specified datasets is updated. - -These operators enable you to configure your Airflow workflows to use more complex dataset update conditions, making them more dynamic and flexible. - -Example Use -------------- - -**Scheduling based on multiple dataset updates** - -To schedule a DAG to run only when two specific datasets have both been updated, use the AND operator (``&``): - -.. code-block:: python - - dag1_dataset = Dataset("s3://dag1/output_1.txt") - dag2_dataset = Dataset("s3://dag2/output_1.txt") - - with DAG( - # Consume dataset 1 and 2 with dataset expressions - schedule=(dag1_dataset & dag2_dataset), - ..., - ): - ... - -**Scheduling based on any dataset update** - -To trigger a DAG execution when either one of two datasets is updated, apply the OR operator (``|``): - -.. code-block:: python - - with DAG( - # Consume dataset 1 or 2 with dataset expressions - schedule=(dag1_dataset | dag2_dataset), - ..., - ): - ... - -**Complex Conditional Logic** - -For scenarios requiring more intricate conditions, such as triggering a DAG when one dataset is updated or when both of two other datasets are updated, combine the OR and AND operators: - -.. code-block:: python - - dag3_dataset = Dataset("s3://dag3/output_3.txt") - - with DAG( - # Consume dataset 1 or both 2 and 3 with dataset expressions - schedule=(dag1_dataset | (dag2_dataset & dag3_dataset)), - ..., - ): - ... - - -Dynamic data events emitting and dataset creation through DatasetAlias ------------------------------------------------------------------------ -A dataset alias can be used to emit dataset events of datasets with association to the aliases. Downstreams can depend on resolved dataset. This feature allows you to define complex dependencies for DAG executions based on dataset updates. - -How to use DatasetAlias -~~~~~~~~~~~~~~~~~~~~~~~ - -``DatasetAlias`` has one single argument ``name`` that uniquely identifies the dataset. The task must first declare the alias as an outlet, and use ``outlet_events`` or yield ``Metadata`` to add events to it. - -The following example creates a dataset event against the S3 URI ``f"s3://bucket/my-task"`` with optional extra information ``extra``. If the dataset does not exist, Airflow will dynamically create it and log a warning message. - -**Emit a dataset event during task execution through outlet_events** - -.. code-block:: python - - from airflow.datasets import DatasetAlias - - - @task(outlets=[DatasetAlias("my-task-outputs")]) - def my_task_with_outlet_events(*, outlet_events): - outlet_events["my-task-outputs"].add(Dataset("s3://bucket/my-task"), extra={"k": "v"}) - - -**Emit a dataset event during task execution through yielding Metadata** - -.. code-block:: python - - from airflow.datasets.metadata import Metadata - - - @task(outlets=[DatasetAlias("my-task-outputs")]) - def my_task_with_metadata(): - s3_dataset = Dataset("s3://bucket/my-task") - yield Metadata(s3_dataset, extra={"k": "v"}, alias="my-task-outputs") - -Only one dataset event is emitted for an added dataset, even if it is added to the alias multiple times, or added to multiple aliases. However, if different ``extra`` values are passed, it can emit multiple dataset events. In the following example, two dataset events will be emitted. - -.. code-block:: python - - from airflow.datasets import DatasetAlias - - - @task( - outlets=[ - DatasetAlias("my-task-outputs-1"), - DatasetAlias("my-task-outputs-2"), - DatasetAlias("my-task-outputs-3"), - ] - ) - def my_task_with_outlet_events(*, outlet_events): - outlet_events["my-task-outputs-1"].add(Dataset("s3://bucket/my-task"), extra={"k": "v"}) - # This line won't emit an additional dataset event as the dataset and extra are the same as the previous line. - outlet_events["my-task-outputs-2"].add(Dataset("s3://bucket/my-task"), extra={"k": "v"}) - # This line will emit an additional dataset event as the extra is different. - outlet_events["my-task-outputs-3"].add(Dataset("s3://bucket/my-task"), extra={"k2": "v2"}) - -Scheduling based on dataset aliases -~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ -Since dataset events added to an alias are just simple dataset events, a downstream DAG depending on the actual dataset can read dataset events of it normally, without considering the associated aliases. A downstream DAG can also depend on a dataset alias. The authoring syntax is referencing the ``DatasetAlias`` by name, and the associated dataset events are picked up for scheduling. Note that a DAG can be triggered by a task with ``outlets=DatasetAlias("xxx")`` if and only if the alias is resolved into ``Dataset("s3://bucket/my-task")``. The DAG runs whenever a task with outlet ``DatasetAlias("out")`` gets associated with at least one dataset at runtime, regardless of the dataset's identity. The downstream DAG is not triggered if no datasets are associated to the alias for a particular given task run. This also means we can do conditional dataset-triggering. - -The dataset alias is resolved to the datasets during DAG parsing. Thus, if the "min_file_process_interval" configuration is set to a high value, there is a possibility that the dataset alias may not be resolved. To resolve this issue, you can trigger DAG parsing. - -.. code-block:: python - - with DAG(dag_id="dataset-producer"): - - @task(outlets=[Dataset("example-alias")]) - def produce_dataset_events(): - pass - - - with DAG(dag_id="dataset-alias-producer"): - - @task(outlets=[DatasetAlias("example-alias")]) - def produce_dataset_events(*, outlet_events): - outlet_events["example-alias"].add(Dataset("s3://bucket/my-task")) - - - with DAG(dag_id="dataset-consumer", schedule=Dataset("s3://bucket/my-task")): - ... - - with DAG(dag_id="dataset-alias-consumer", schedule=DatasetAlias("example-alias")): - ... - - -In the example provided, once the DAG ``dataset-alias-producer`` is executed, the dataset alias ``DatasetAlias("example-alias")`` will be resolved to ``Dataset("s3://bucket/my-task")``. However, the DAG ``dataset-alias-consumer`` will have to wait for the next DAG re-parsing to update its schedule. To address this, Airflow will re-parse the DAGs relying on the dataset alias ``DatasetAlias("example-alias")`` when it's resolved into datasets that these DAGs did not previously depend on. As a result, both the "dataset-consumer" and "dataset-alias-consumer" DAGs will be triggered after the execution of DAG ``dataset-alias-producer``. - - -Fetching information from previously emitted dataset events through resolved dataset aliases -~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ - -As mentioned in :ref:`Fetching information from previously emitted dataset events`, inlet dataset events can be read with the ``inlet_events`` accessor in the execution context, and you can also use dataset aliases to access the dataset events triggered by them. - -.. code-block:: python - - with DAG(dag_id="dataset-alias-producer"): - - @task(outlets=[DatasetAlias("example-alias")]) - def produce_dataset_events(*, outlet_events): - outlet_events["example-alias"].add(Dataset("s3://bucket/my-task"), extra={"row_count": 1}) - - - with DAG(dag_id="dataset-alias-consumer", schedule=None): - - @task(inlets=[DatasetAlias("example-alias")]) - def consume_dataset_alias_events(*, inlet_events): - events = inlet_events[DatasetAlias("example-alias")] - last_row_count = events[-1].extra["row_count"] - - -Combining dataset and time-based schedules ------------------------------------------- - -DatasetTimetable Integration -~~~~~~~~~~~~~~~~~~~~~~~~~~~~ -You can schedule DAGs based on both dataset events and time-based schedules using ``DatasetOrTimeSchedule``. This allows you to create workflows when a DAG needs both to be triggered by data updates and run periodically according to a fixed timetable. - -For more detailed information on ``DatasetOrTimeSchedule``, refer to the corresponding section in :ref:`DatasetOrTimeSchedule `. diff --git a/docs/apache-airflow/authoring-and-scheduling/index.rst b/docs/apache-airflow/authoring-and-scheduling/index.rst index 1a042918fc1ba..5ec94d6ca7301 100644 --- a/docs/apache-airflow/authoring-and-scheduling/index.rst +++ b/docs/apache-airflow/authoring-and-scheduling/index.rst @@ -41,5 +41,5 @@ It's recommended that you first review the pages in :doc:`core concepts dict: + def retrieve(src: Asset) -> dict: resp = requests.get(url=src.uri) data = resp.json() return data["data"] @@ -137,14 +137,14 @@ a ``Dataset``, which is ``@attr.define`` decorated, together with TaskFlow. return ret @task() - def load(fahrenheit: dict[int, float]) -> Dataset: + def load(fahrenheit: dict[int, float]) -> Asset: filename = "/tmp/fahrenheit.json" s = json.dumps(fahrenheit) f = open(filename, "w") f.write(s) f.close() - return Dataset(f"file:///{filename}") + return Asset(f"file:///{filename}") data = retrieve(SRC) fahrenheit = to_fahrenheit(data) diff --git a/docs/apache-airflow/img/dataset-scheduled-dags.png b/docs/apache-airflow/img/asset-scheduled-dags.png similarity index 100% rename from docs/apache-airflow/img/dataset-scheduled-dags.png rename to docs/apache-airflow/img/asset-scheduled-dags.png diff --git a/docs/apache-airflow/img/datasets.png b/docs/apache-airflow/img/assets.png similarity index 100% rename from docs/apache-airflow/img/datasets.png rename to docs/apache-airflow/img/assets.png diff --git a/docs/apache-airflow/templates-ref.rst b/docs/apache-airflow/templates-ref.rst index 05d4b10accca5..5524c82f8cc95 100644 --- a/docs/apache-airflow/templates-ref.rst +++ b/docs/apache-airflow/templates-ref.rst @@ -62,10 +62,10 @@ Variable Type Description ``{{ prev_end_date_success }}`` `pendulum.DateTime`_ End date from prior successful :class:`~airflow.models.dagrun.DagRun` (if available). | ``None`` ``{{ inlets }}`` list List of inlets declared on the task. -``{{ inlet_events }}`` dict[str, ...] Access past events of inlet datasets. See :doc:`Datasets `. Added in version 2.10. +``{{ inlet_events }}`` dict[str, ...] Access past events of inlet assets. See :doc:`Assets `. Added in version 2.10. ``{{ outlets }}`` list List of outlets declared on the task. -``{{ outlet_events }}`` dict[str, ...] | Accessors to attach information to dataset events that will be emitted by the current task. - | See :doc:`Datasets `. Added in version 2.10. +``{{ outlet_events }}`` dict[str, ...] | Accessors to attach information to asset events that will be emitted by the current task. + | See :doc:`Assets `. Added in version 2.10. ``{{ dag }}`` DAG The currently running :class:`~airflow.models.dag.DAG`. You can read more about DAGs in :doc:`DAGs `. ``{{ task }}`` BaseOperator | The currently running :class:`~airflow.models.baseoperator.BaseOperator`. You can read more about Tasks in :doc:`core-concepts/operators` ``{{ macros }}`` | A reference to the macros package. See Macros_ below. @@ -88,9 +88,9 @@ Variable Type Description ``{{ expanded_ti_count }}`` int | ``None`` | Number of task instances that a mapped task was expanded into. If | the current task is not mapped, this should be ``None``. | Added in version 2.5. -``{{ triggering_dataset_events }}`` dict[str, | If in a Dataset Scheduled DAG, a map of Dataset URI to a list of triggering :class:`~airflow.models.dataset.DatasetEvent` - list[DatasetEvent]] | (there may be more than one, if there are multiple Datasets with different frequencies). - | Read more here :doc:`Datasets `. +``{{ triggering_asset_events }}`` dict[str, | If in a Asset Scheduled DAG, a map of Asset URI to a list of triggering :class:`~airflow.models.asset.AssetEvent` + list[AssetEvent]] | (there may be more than one, if there are multiple Assets with different frequencies). + | Read more here :doc:`Assets `. | Added in version 2.4. =========================================== ===================== =================================================================== diff --git a/docs/apache-airflow/tutorial/objectstorage.rst b/docs/apache-airflow/tutorial/objectstorage.rst index 943e8031a7e58..39e42ddd76627 100644 --- a/docs/apache-airflow/tutorial/objectstorage.rst +++ b/docs/apache-airflow/tutorial/objectstorage.rst @@ -65,7 +65,7 @@ The connection ID can alternatively be passed in with a keyword argument: ObjectStoragePath("s3://airflow-tutorial-data/", conn_id="aws_default") -This is useful when reusing a URL defined for another purpose (e.g. Dataset), +This is useful when reusing a URL defined for another purpose (e.g. Asset), which generally does not contain a username part. The explicit keyword argument takes precedence over the URL's username value if both are specified. diff --git a/docs/apache-airflow/ui.rst b/docs/apache-airflow/ui.rst index 4238cb6ba427c..05d71c9a96176 100644 --- a/docs/apache-airflow/ui.rst +++ b/docs/apache-airflow/ui.rst @@ -61,7 +61,7 @@ Native Airflow dashboard page into the UI to collect several useful metrics for ------------ -.. _ui:datasets-view: +.. _ui:assets-view: Datasets View ............. @@ -72,7 +72,7 @@ Clicking on any dataset in either the list or the graph will highlight it and it ------------ -.. image:: img/datasets.png +.. image:: img/assets.png ------------ diff --git a/docs/exts/operators_and_hooks_ref.py b/docs/exts/operators_and_hooks_ref.py index 43f954ebb0c37..fe6cd5d3300d2 100644 --- a/docs/exts/operators_and_hooks_ref.py +++ b/docs/exts/operators_and_hooks_ref.py @@ -519,8 +519,8 @@ class DatasetSchemeDirective(BaseJinjaReferenceDirective): def render_content(self, *, tags: set[str] | None, header_separator: str = DEFAULT_HEADER_SEPARATOR): return _common_render_list_content( header_separator=header_separator, - resource_type="dataset-uris", - template="dataset-uri-schemes.rst.jinja2", + resource_type="asset-uris", + template="asset-uri-schemes.rst.jinja2", ) @@ -538,7 +538,7 @@ def setup(app): app.add_directive("airflow-executors", ExecutorsDirective) app.add_directive("airflow-deferrable-operators", DeferrableOperatorDirective) app.add_directive("airflow-deprecations", DeprecationsDirective) - app.add_directive("airflow-dataset-schemes", DatasetSchemeDirective) + app.add_directive("airflow-asset-schemes", DatasetSchemeDirective) return {"parallel_read_safe": True, "parallel_write_safe": True} diff --git a/docs/exts/templates/dataset-uri-schemes.rst.jinja2 b/docs/exts/templates/asset-uri-schemes.rst.jinja2 similarity index 95% rename from docs/exts/templates/dataset-uri-schemes.rst.jinja2 rename to docs/exts/templates/asset-uri-schemes.rst.jinja2 index aa247507becc4..14cdbcf9aab15 100644 --- a/docs/exts/templates/dataset-uri-schemes.rst.jinja2 +++ b/docs/exts/templates/asset-uri-schemes.rst.jinja2 @@ -27,7 +27,7 @@ Core {{ provider['name'] }} {{ header_separator * (provider['name']|length) }} -{% for uri_entry in provider['dataset-uris'] -%} +{% for uri_entry in provider['asset-uris'] -%} - {% for scheme in uri_entry['schemes'] %}``{{ scheme }}``{% if not loop.last %}, {% endif %}{% endfor %} {% endfor %} diff --git a/docs/spelling_wordlist.txt b/docs/spelling_wordlist.txt index b834ccc9b2005..eb6e612e0992b 100644 --- a/docs/spelling_wordlist.txt +++ b/docs/spelling_wordlist.txt @@ -88,6 +88,8 @@ asctime asend asia assertEqualIgnoreMultipleSpaces +AssetEvent +AssetEvents assigment ast astroid diff --git a/generated/provider_dependencies.json b/generated/provider_dependencies.json index b9bc363b15e33..2a81933c6c584 100644 --- a/generated/provider_dependencies.json +++ b/generated/provider_dependencies.json @@ -397,7 +397,9 @@ ], "devel-deps": [], "plugins": [], - "cross-providers-deps": [], + "cross-providers-deps": [ + "openlineage" + ], "excluded-python-versions": [], "state": "ready" }, @@ -563,6 +565,7 @@ }, "fab": { "deps": [ + "apache-airflow-providers-common-compat>=1.2.0", "apache-airflow>=2.9.0", "flask-appbuilder==4.5.0", "flask-login>=0.6.2", @@ -574,7 +577,9 @@ "kerberos>=1.3.0" ], "plugins": [], - "cross-providers-deps": [], + "cross-providers-deps": [ + "common.compat" + ], "excluded-python-versions": [], "state": "ready" }, @@ -967,6 +972,7 @@ } ], "cross-providers-deps": [ + "common.compat", "common.sql" ], "excluded-python-versions": [], diff --git a/newsfragments/41348.significant.rst b/newsfragments/41348.significant.rst new file mode 100644 index 0000000000000..8b5cc54dd40dc --- /dev/null +++ b/newsfragments/41348.significant.rst @@ -0,0 +1,240 @@ +**Breaking Change** + +* Rename module ``airflow.api_connexion.schemas.dataset_schema`` as ``airflow.api_connexion.schemas.asset_schema`` + + * Rename variable ``create_dataset_event_schema`` as ``create_asset_event_schema`` + * Rename variable ``dataset_collection_schema`` as ``asset_collection_schema`` + * Rename variable ``dataset_event_collection_schema`` as ``asset_event_collection_schema`` + * Rename variable ``dataset_event_schema`` as ``asset_event_schema`` + * Rename variable ``dataset_schema`` as ``asset_schema`` + * Rename class ``TaskOutletDatasetReferenceSchema`` as ``TaskOutletAssetReferenceSchema`` + * Rename class ``DagScheduleDatasetReferenceSchema`` as ``DagScheduleAssetReferenceSchema`` + * Rename class ``DatasetAliasSchema`` as ``AssetAliasSchema`` + * Rename class ``DatasetSchema`` as ``AssetSchema`` + * Rename class ``DatasetCollection`` as ``AssetCollection`` + * Rename class ``DatasetEventSchema`` as ``AssetEventSchema`` + * Rename class ``DatasetEventCollection`` as ``AssetEventCollection`` + * Rename class ``DatasetEventCollectionSchema`` as ``AssetEventCollectionSchema`` + * Rename class ``CreateDatasetEventSchema`` as ``CreateAssetEventSchema`` + +* Rename module ``airflow.datasets`` as ``airflow.assets`` + + * Rename class ``DatasetAlias`` as ``AssetAlias`` + * Rename class ``DatasetAll`` as ``AssetAll`` + * Rename class ``DatasetAny`` as ``AssetAny`` + * Rename function ``expand_alias_to_datasets`` as ``expand_alias_to_assets`` + * Rename class ``DatasetAliasEvent`` as ``AssetAliasEvent`` + + * Rename method ``dest_dataset_uri`` as ``dest_asset_uri`` + + * Rename class ``BaseDataset`` as ``BaseAsset`` + + * Rename method ``iter_datasets`` as ``iter_assets`` + * Rename method ``iter_dataset_aliases`` as ``iter_asset_aliases`` + + * Rename class ``Dataset`` as ``Asset`` + + * Rename method ``iter_datasets`` as ``iter_assets`` + * Rename method ``iter_dataset_aliases`` as ``iter_asset_aliases`` + + * Rename class ``_DatasetBooleanCondition`` as ``_AssetBooleanCondition`` + + * Rename method ``iter_datasets`` as ``iter_assets`` + * Rename method ``iter_dataset_aliases`` as ``iter_asset_aliases`` + +* Rename module ``airflow.datasets.manager`` as ``airflow.assets.manager`` + + * Rename variable ``dataset_manager`` as ``asset_manager`` + * Rename function ``resolve_dataset_manager`` as ``resolve_asset_manager`` + * Rename class ``DatasetManager`` as ``AssetManager`` + + * Rename method ``register_dataset_change`` as ``register_asset_change`` + * Rename method ``create_datasets`` as ``create_assets`` + * Rename method ``register_dataset_change`` as ``notify_asset_created`` + * Rename method ``notify_dataset_changed`` as ``notify_asset_changed`` + * Renme method ``notify_dataset_alias_created`` as ``notify_asset_alias_created`` + +* Rename module ``airflow.models.dataset`` as ``airflow.models.asset`` + + * Rename class ``DatasetDagRunQueue`` as ``AssetDagRunQueue`` + * Rename class ``DatasetEvent`` as ``AssetEvent`` + * Rename class ``DatasetModel`` as ``AssetModel`` + * Rename class ``DatasetAliasModel`` as ``AssetAliasModel`` + * Rename class ``DagScheduleDatasetReference`` as ``DagScheduleAssetReference`` + * Rename class ``TaskOutletDatasetReference`` as ``TaskOutletAssetReference`` + * Rename class ``DagScheduleDatasetAliasReference`` as ``DagScheduleAssetAliasReference`` + +* Rename module ``airflow.api_ui.views.datasets`` as ``airflow.api_ui.views.assets`` + + * Rename variable ``dataset_router`` as ``asset_rounter`` + +* Rename module ``airflow.listeners.spec.dataset`` as ``airflow.listeners.spec.asset`` + + * Rename function ``on_dataset_created`` as ``on_asset_created`` + * Rename function ``on_dataset_changed`` as ``on_asset_changed`` + +* Rename module ``airflow.timetables.datasets`` as ``airflow.timetables.assets`` + + * Rename class ``DatasetOrTimeSchedule`` as ``AssetOrTimeSchedule`` + +* Rename module ``airflow.serialization.pydantic.dataset`` as ``airflow.serialization.pydantic.asset`` + + * Rename class ``DagScheduleDatasetReferencePydantic`` as ``DagScheduleAssetReferencePydantic`` + * Rename class ``TaskOutletDatasetReferencePydantic`` as ``TaskOutletAssetReferencePydantic`` + * Rename class ``DatasetPydantic`` as ``AssetPydantic`` + * Rename class ``DatasetEventPydantic`` as ``AssetEventPydantic`` + +* Rename module ``airflow.datasets.metadata`` as ``airflow.assets.metadata`` + +* In module ``airflow.jobs.scheduler_job_runner`` + + * and its class ``SchedulerJobRunner`` + + * Rename method ``_create_dag_runs_dataset_triggered`` as ``_create_dag_runs_asset_triggered`` + * Rename method ``_orphan_unreferenced_datasets`` as ``_orphan_unreferenced_datasets`` + +* In module ``airflow.api_connexion.security`` + + * Rename decorator ``requires_access_dataset`` as ``requires_access_asset`` + +* In module ``airflow.auth.managers.models.resource_details`` + + * Rename class ``DatasetDetails`` as ``AssetDetails`` + +* In module ``airflow.auth.managers.base_auth_manager`` + + * Rename function ``is_authorized_dataset`` as ``is_authorized_asset`` + +* In module ``airflow.timetables.simple`` + + * Rename class ``DatasetTriggeredTimetable`` as ``AssetTriggeredTimetable`` + +* In module ``airflow.lineage.hook`` + + * Rename class ``DatasetLineageInfo`` as ``AssetLineageInfo`` + + * Rename attribute ``dataset`` as ``asset`` + + * In its class ``HookLineageCollector`` + + * Rename method ``create_dataset`` as ``create_asset`` + * Rename method ``add_input_dataset`` as ``add_input_asset`` + * Rename method ``add_output_dataset`` as ``add_output_asset`` + * Rename method ``collected_datasets`` as ``collected_assets`` + +* In module ``airflow.models.dag`` + + * Rename function ``get_dataset_triggered_next_run_info`` as ``get_asset_triggered_next_run_info`` + + * In its class ``DagModel`` + + * Rename method ``get_dataset_triggered_next_run_info`` as ``get_asset_triggered_next_run_info`` + +* In module ``airflow.models.taskinstance`` + + * and its class ``TaskInstance`` + + * Rename method ``_register_dataset_changes`` as ``_register_asset_changes`` + +* In module ``airflow.providers_manager`` + + * and its class ``ProvidersManager`` + + * Rename method ``initialize_providers_dataset_uri_resources`` as ``initialize_providers_asset_uri_resources`` + * Rename attribute ``_discover_dataset_uri_resources`` as ``_discover_asset_uri_resources`` + * Rename property ``dataset_factories`` as ``asset_factories`` + * Rename property ``dataset_uri_handlers`` as ``asset_uri_handlers`` + * Rename property ``dataset_to_openlineage_converters`` as ``asset_to_openlineage_converters`` + +* In module ``airflow.security.permissions`` + + * Rename constant ``RESOURCE_DATASET`` as ``RESOURCE_ASSET`` + +* In module ``airflow.serialization.enums`` + + * and its class DagAttributeTypes + + * Rename attribute ``DATASET_EVENT_ACCESSORS`` as ``ASSET_EVENT_ACCESSORS`` + * Rename attribute ``DATASET_EVENT_ACCESSOR`` as ``ASSET_EVENT_ACCESSOR`` + * Rename attribute ``DATASET`` as ``ASSET`` + * Rename attribute ``DATASET_ALIAS`` as ``ASSET_ALIAS`` + * Rename attribute ``DATASET_ANY`` as ``ASSET_ANY`` + * Rename attribute ``DATASET_ALL`` as ``ASSET_ALL`` + +* In module ``airflow.serialization.pydantic.taskinstance`` + + * and its class ``TaskInstancePydantic`` + + * Rename method ``_register_dataset_changes`` as ``_register_dataset_changes`` + +* In module ``airflow.serialization.serialized_objects`` + + * Rename function ``encode_dataset_condition`` as ``encode_asset_condition`` + * Rename function ``decode_dataset_condition`` as ``decode_asset_condition`` + +* In module ``airflow.timetables.base`` + + * Rename class ```_NullDataset``` as ```_NullAsset``` + + * Rename method ``iter_datasets`` as ``iter_assets`` + * Rename method ``iter_dataset_aliases`` as ``iter_assets_aliases`` + +* In module ``airflow.utils.context`` + + * Rename class ``LazyDatasetEventSelectSequence`` as ``LazyAssetEventSelectSequence`` + +* In module ``airflow.www.auth`` + + * Rename function ``has_access_dataset`` as ``has_access_asset`` + +* Rename configuration ``core.strict_dataset_uri_validation`` as ``core.strict_asset_uri_validation``, ``core.dataset_manager_class`` as ``core.asset_manager_class`` and ``core.dataset_manager_class`` as ``core.asset_manager_class`` +* Rename example dags ``example_dataset_alias.py``, ``example_dataset_alias_with_no_taskflow.py``, ``example_datasets.py`` as ``example_asset_alias.py``, ``example_asset_alias_with_no_taskflow.py``, ``example_assets.py`` +* Rename DagDependency name ``dataset-alias``, ``dataset`` as ``asset-alias``, ``asset`` +* Rename context key ``triggering_dataset_events`` as ``triggering_asset_events`` +* Rename resource key ``dataset-uris`` as ``asset-uris`` for providers amazon, common.io, mysql, fab, postgres, trino + +* In provider ``airflow.providers.amazon.aws`` + + * Rename package ``datasets`` as ``assets`` + + * In its module ``s3`` + + * Rename method ``create_dataset`` as ``create_asset`` + * Rename method ``convert_dataset_to_openlineage`` as ``convert_asset_to_openlineage`` + + * and its module ``auth_manager.avp.entities`` + + * Rename attribute ``AvpEntities.DATASET`` as ``AvpEntities.ASSET`` + + * and its module ``auth_manager.auth_manager.aws_auth_manager`` + + * Rename function ``is_authorized_dataset`` as ``is_authorized_asset`` + +* In provider ``airflow.providers.common.io`` + + * Rename package ``datasets`` as ``assets`` + + * in its module ``file`` + + * Rename method ``create_dataset`` as ``create_asset`` + * Rename method ``convert_dataset_to_openlineage`` as ``convert_asset_to_openlineage`` + +* In provider ``airflow.providers.fab`` + + * in its module ``auth_manager.fab_auth_manager`` + + * Rename function ``is_authorized_dataset`` as ``is_authorized_asset`` + +* In provider ``airflow.providers.openlineage`` + + * in its module ``utils.utils`` + + * Rename class ``DatasetInfo`` as ``AssetInfo`` + * Rename function ``translate_airflow_dataset`` as ``translate_airflow_asset`` + +* Rename package ``airflow.providers.postgres.datasets`` as ``airflow.providers.postgres.assets`` +* Rename package ``airflow.providers.mysql.datasets`` as ``airflow.providers.mysql.assets`` +* Rename package ``airflow.providers.trino.datasets`` as ``airflow.providers.trino.assets`` +* Add module ``airflow.providers.common.compat.assets`` +* Add module ``airflow.providers.common.compat.openlineage.utils.utils`` +* Add module ``airflow.providers.common.compat.security.permissions`` diff --git a/scripts/ci/pre_commit/check_tests_in_right_folders.py b/scripts/ci/pre_commit/check_tests_in_right_folders.py index 8260b6ad0d578..11d44efd407a7 100755 --- a/scripts/ci/pre_commit/check_tests_in_right_folders.py +++ b/scripts/ci/pre_commit/check_tests_in_right_folders.py @@ -34,6 +34,7 @@ "api_connexion", "api_internal", "api_fastapi", + "assets", "auth", "callbacks", "charts", @@ -45,7 +46,6 @@ "dags", "dags_corrupted", "dags_with_system_exit", - "datasets", "decorators", "executors", "hooks", diff --git a/scripts/cov/core_coverage.py b/scripts/cov/core_coverage.py index 0facd4bb1c5d7..2d8ac091c6e0d 100644 --- a/scripts/cov/core_coverage.py +++ b/scripts/cov/core_coverage.py @@ -47,6 +47,7 @@ "airflow/jobs/triggerer_job_runner.py", # models "airflow/models/abstractoperator.py", + "airflow/models/asset.py", "airflow/models/base.py", "airflow/models/baseoperator.py", "airflow/models/connection.py", @@ -57,7 +58,6 @@ "airflow/models/dagpickle.py", "airflow/models/dagrun.py", "airflow/models/dagwarning.py", - "airflow/models/dataset.py", "airflow/models/expandinput.py", "airflow/models/log.py", "airflow/models/mappedoperator.py", diff --git a/scripts/cov/other_coverage.py b/scripts/cov/other_coverage.py index 6543d2fc780e0..dae7733ec5c15 100644 --- a/scripts/cov/other_coverage.py +++ b/scripts/cov/other_coverage.py @@ -37,7 +37,7 @@ "airflow/callbacks", "airflow/config_templates", "airflow/dag_processing", - "airflow/datasets", + "airflow/assets", "airflow/decorators", "airflow/hooks", "airflow/io", @@ -79,7 +79,7 @@ "tests/cluster_policies", "tests/config_templates", "tests/dag_processing", - "tests/datasets", + "tests/assets", "tests/decorators", "tests/hooks", "tests/io", diff --git a/tests/always/test_project_structure.py b/tests/always/test_project_structure.py index c387f6173ca2f..b27729a68a261 100644 --- a/tests/always/test_project_structure.py +++ b/tests/always/test_project_structure.py @@ -75,6 +75,7 @@ def test_providers_modules_should_have_tests(self): "tests/providers/amazon/aws/triggers/test_step_function.py", "tests/providers/amazon/aws/utils/test_rds.py", "tests/providers/amazon/aws/utils/test_sagemaker.py", + "tests/providers/amazon/aws/utils/test_asset_compat_lineage_collector.py", "tests/providers/amazon/aws/waiters/test_base_waiter.py", "tests/providers/apache/cassandra/hooks/test_cassandra.py", "tests/providers/apache/drill/operators/test_drill.py", @@ -150,6 +151,7 @@ def test_providers_modules_should_have_tests(self): "tests/providers/google/test_go_module_utils.py", "tests/providers/microsoft/azure/operators/test_adls.py", "tests/providers/microsoft/azure/transfers/test_azure_blob_to_gcs.py", + "tests/providers/openlineage/utils/test_asset_compat_lineage_collector.py", "tests/providers/slack/notifications/test_slack_notifier.py", "tests/providers/snowflake/triggers/test_snowflake_trigger.py", "tests/providers/yandex/hooks/test_yandexcloud_dataproc.py", diff --git a/tests/api_connexion/endpoints/test_dag_run_endpoint.py b/tests/api_connexion/endpoints/test_dag_run_endpoint.py index deb5fe0af2daa..f3921da7b9c29 100644 --- a/tests/api_connexion/endpoints/test_dag_run_endpoint.py +++ b/tests/api_connexion/endpoints/test_dag_run_endpoint.py @@ -24,10 +24,10 @@ import time_machine from airflow.api_connexion.exceptions import EXCEPTIONS_LINK_MAP -from airflow.datasets import Dataset +from airflow.assets import Asset +from airflow.models.asset import AssetEvent, AssetModel from airflow.models.dag import DAG, DagModel from airflow.models.dagrun import DagRun -from airflow.models.dataset import DatasetEvent, DatasetModel from airflow.models.param import Param from airflow.operators.empty import EmptyOperator from airflow.security import permissions @@ -57,7 +57,7 @@ def configured_app(minimal_app_for_api): role_name="Test", permissions=[ (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DATASET), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_ASSET), (permissions.ACTION_CAN_READ, permissions.RESOURCE_CLUSTER_ACTIVITY), (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG), (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_DAG_RUN), @@ -72,7 +72,7 @@ def configured_app(minimal_app_for_api): role_name="TestNoDagRunCreatePermission", permissions=[ (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DATASET), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_ASSET), (permissions.ACTION_CAN_READ, permissions.RESOURCE_CLUSTER_ACTIVITY), (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG), (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), @@ -1911,16 +1911,16 @@ def test_should_respond_404(self): @pytest.mark.need_serialized_dag class TestGetDagRunDatasetTriggerEvents(TestDagRunEndpoint): def test_should_respond_200(self, dag_maker, session): - dataset1 = Dataset(uri="ds1") + asset1 = Asset(uri="ds1") with dag_maker(dag_id="source_dag", start_date=timezone.utcnow(), session=session): - EmptyOperator(task_id="task", outlets=[dataset1]) + EmptyOperator(task_id="task", outlets=[asset1]) dr = dag_maker.create_dagrun() ti = dr.task_instances[0] - ds1_id = session.query(DatasetModel.id).filter_by(uri=dataset1.uri).scalar() - event = DatasetEvent( - dataset_id=ds1_id, + asset1_id = session.query(AssetModel.id).filter_by(uri=asset1.uri).scalar() + event = AssetEvent( + dataset_id=asset1_id, source_task_id=ti.task_id, source_dag_id=ti.dag_id, source_run_id=ti.run_id, @@ -1945,8 +1945,8 @@ def test_should_respond_200(self, dag_maker, session): "dataset_events": [ { "timestamp": event.timestamp.isoformat(), - "dataset_id": ds1_id, - "dataset_uri": dataset1.uri, + "dataset_id": asset1_id, + "dataset_uri": asset1.uri, "extra": {}, "id": event.id, "source_dag_id": ti.dag_id, diff --git a/tests/api_connexion/endpoints/test_dag_source_endpoint.py b/tests/api_connexion/endpoints/test_dag_source_endpoint.py index 1e5389d377440..a8d1224e034c3 100644 --- a/tests/api_connexion/endpoints/test_dag_source_endpoint.py +++ b/tests/api_connexion/endpoints/test_dag_source_endpoint.py @@ -37,7 +37,7 @@ EXAMPLE_DAG_ID = "example_bash_operator" TEST_DAG_ID = "latest_only" NOT_READABLE_DAG_ID = "latest_only_with_trigger" -TEST_MULTIPLE_DAGS_ID = "dataset_produces_1" +TEST_MULTIPLE_DAGS_ID = "asset_produces_1" @pytest.fixture(scope="module") diff --git a/tests/api_connexion/endpoints/test_dataset_endpoint.py b/tests/api_connexion/endpoints/test_dataset_endpoint.py index 25f8012039109..5caec0ac2a131 100644 --- a/tests/api_connexion/endpoints/test_dataset_endpoint.py +++ b/tests/api_connexion/endpoints/test_dataset_endpoint.py @@ -25,14 +25,14 @@ from airflow.api_connexion.exceptions import EXCEPTIONS_LINK_MAP from airflow.models import DagModel -from airflow.models.dagrun import DagRun -from airflow.models.dataset import ( - DagScheduleDatasetReference, - DatasetDagRunQueue, - DatasetEvent, - DatasetModel, - TaskOutletDatasetReference, +from airflow.models.asset import ( + AssetDagRunQueue, + AssetEvent, + AssetModel, + DagScheduleAssetReference, + TaskOutletAssetReference, ) +from airflow.models.dagrun import DagRun from airflow.security import permissions from airflow.utils import timezone from airflow.utils.session import provide_session @@ -40,7 +40,7 @@ from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_user from tests.test_utils.asserts import assert_queries_count from tests.test_utils.config import conf_vars -from tests.test_utils.db import clear_db_datasets, clear_db_runs +from tests.test_utils.db import clear_db_assets, clear_db_runs from tests.test_utils.www import _check_last_log pytestmark = [pytest.mark.db_test, pytest.mark.skip_if_database_isolation_mode] @@ -54,8 +54,8 @@ def configured_app(minimal_app_for_api): username="test", role_name="Test", permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DATASET), - (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_DATASET), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_ASSET), + (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_ASSET), ], ) create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore @@ -65,8 +65,8 @@ def configured_app(minimal_app_for_api): role_name="TestQueuedEvent", permissions=[ (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DATASET), - (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_DATASET), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_ASSET), + (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_ASSET), ], ) @@ -84,30 +84,30 @@ class TestDatasetEndpoint: def setup_attrs(self, configured_app) -> None: self.app = configured_app self.client = self.app.test_client() - clear_db_datasets() + clear_db_assets() clear_db_runs() def teardown_method(self) -> None: - clear_db_datasets() + clear_db_assets() clear_db_runs() def _create_dataset(self, session): - dataset_model = DatasetModel( + asset_model = AssetModel( id=1, uri="s3://bucket/key", extra={"foo": "bar"}, created_at=timezone.parse(self.default_time), updated_at=timezone.parse(self.default_time), ) - session.add(dataset_model) + session.add(asset_model) session.commit() - return dataset_model + return asset_model class TestGetDatasetEndpoint(TestDatasetEndpoint): def test_should_respond_200(self, session): self._create_dataset(session) - assert session.query(DatasetModel).count() == 1 + assert session.query(AssetModel).count() == 1 with assert_queries_count(6): response = self.client.get( @@ -133,9 +133,9 @@ def test_should_respond_404(self): ) assert response.status_code == 404 assert { - "detail": "The Dataset with uri: `s3://bucket/key` was not found", + "detail": "The Asset with uri: `s3://bucket/key` was not found", "status": 404, - "title": "Dataset not found", + "title": "Asset not found", "type": EXCEPTIONS_LINK_MAP[404], } == response.json @@ -147,8 +147,8 @@ def test_should_raises_401_unauthenticated(self, session): class TestGetDatasets(TestDatasetEndpoint): def test_should_respond_200(self, session): - datasets = [ - DatasetModel( + assets = [ + AssetModel( id=i, uri=f"s3://bucket/key/{i}", extra={"foo": "bar"}, @@ -157,9 +157,9 @@ def test_should_respond_200(self, session): ) for i in [1, 2] ] - session.add_all(datasets) + session.add_all(assets) session.commit() - assert session.query(DatasetModel).count() == 2 + assert session.query(AssetModel).count() == 2 with assert_queries_count(10): response = self.client.get("/api/v1/datasets", environ_overrides={"REMOTE_USER": "test"}) @@ -193,8 +193,8 @@ def test_should_respond_200(self, session): } def test_order_by_raises_400_for_invalid_attr(self, session): - datasets = [ - DatasetModel( + assets = [ + AssetModel( uri=f"s3://bucket/key/{i}", extra={"foo": "bar"}, created_at=timezone.parse(self.default_time), @@ -202,9 +202,9 @@ def test_order_by_raises_400_for_invalid_attr(self, session): ) for i in [1, 2] ] - session.add_all(datasets) + session.add_all(assets) session.commit() - assert session.query(DatasetModel).count() == 2 + assert session.query(AssetModel).count() == 2 response = self.client.get( "/api/v1/datasets?order_by=fake", environ_overrides={"REMOTE_USER": "test"} @@ -215,8 +215,8 @@ def test_order_by_raises_400_for_invalid_attr(self, session): assert response.json["detail"] == msg def test_should_raises_401_unauthenticated(self, session): - datasets = [ - DatasetModel( + assets = [ + AssetModel( uri=f"s3://bucket/key/{i}", extra={"foo": "bar"}, created_at=timezone.parse(self.default_time), @@ -224,9 +224,9 @@ def test_should_raises_401_unauthenticated(self, session): ) for i in [1, 2] ] - session.add_all(datasets) + session.add_all(assets) session.commit() - assert session.query(DatasetModel).count() == 2 + assert session.query(AssetModel).count() == 2 response = self.client.get("/api/v1/datasets") @@ -254,11 +254,11 @@ def test_should_raises_401_unauthenticated(self, session): ) @provide_session def test_filter_datasets_by_uri_pattern_works(self, url, expected_datasets, session): - dataset1 = DatasetModel("s3://folder/key") - dataset2 = DatasetModel("gcp://bucket/key") - dataset3 = DatasetModel("somescheme://dataset/key") - dataset4 = DatasetModel("wasb://some_dataset_bucket_/key") - session.add_all([dataset1, dataset2, dataset3, dataset4]) + asset1 = AssetModel("s3://folder/key") + asset2 = AssetModel("gcp://bucket/key") + asset3 = AssetModel("somescheme://dataset/key") + asset4 = AssetModel("wasb://some_dataset_bucket_/key") + session.add_all([asset1, asset2, asset3, asset4]) session.commit() response = self.client.get(url, environ_overrides={"REMOTE_USER": "test"}) assert response.status_code == 200 @@ -273,12 +273,12 @@ def test_filter_datasets_by_dag_ids_works(self, dag_ids, expected_num, session): dag1 = DagModel(dag_id="dag1") dag2 = DagModel(dag_id="dag2") dag3 = DagModel(dag_id="dag3") - dataset1 = DatasetModel("s3://folder/key") - dataset2 = DatasetModel("gcp://bucket/key") - dataset3 = DatasetModel("somescheme://dataset/key") - dag_ref1 = DagScheduleDatasetReference(dag_id="dag1", dataset=dataset1) - dag_ref2 = DagScheduleDatasetReference(dag_id="dag2", dataset=dataset2) - task_ref1 = TaskOutletDatasetReference(dag_id="dag3", task_id="task1", dataset=dataset3) + dataset1 = AssetModel("s3://folder/key") + dataset2 = AssetModel("gcp://bucket/key") + dataset3 = AssetModel("somescheme://dataset/key") + dag_ref1 = DagScheduleAssetReference(dag_id="dag1", dataset=dataset1) + dag_ref2 = DagScheduleAssetReference(dag_id="dag2", dataset=dataset2) + task_ref1 = TaskOutletAssetReference(dag_id="dag3", task_id="task1", dataset=dataset3) session.add_all([dataset1, dataset2, dataset3, dag1, dag2, dag3, dag_ref1, dag_ref2, task_ref1]) session.commit() response = self.client.get( @@ -300,13 +300,13 @@ def test_filter_datasets_by_dag_ids_and_uri_pattern_works( dag1 = DagModel(dag_id="dag1") dag2 = DagModel(dag_id="dag2") dag3 = DagModel(dag_id="dag3") - dataset1 = DatasetModel("s3://folder/key") - dataset2 = DatasetModel("gcp://bucket/key") - dataset3 = DatasetModel("somescheme://dataset/key") - dag_ref1 = DagScheduleDatasetReference(dag_id="dag1", dataset=dataset1) - dag_ref2 = DagScheduleDatasetReference(dag_id="dag2", dataset=dataset2) - task_ref1 = TaskOutletDatasetReference(dag_id="dag3", task_id="task1", dataset=dataset3) - session.add_all([dataset1, dataset2, dataset3, dag1, dag2, dag3, dag_ref1, dag_ref2, task_ref1]) + asset1 = AssetModel("s3://folder/key") + asset2 = AssetModel("gcp://bucket/key") + asset3 = AssetModel("somescheme://dataset/key") + dag_ref1 = DagScheduleAssetReference(dag_id="dag1", dataset=asset1) + dag_ref2 = DagScheduleAssetReference(dag_id="dag2", dataset=asset2) + task_ref1 = TaskOutletAssetReference(dag_id="dag3", task_id="task1", dataset=asset3) + session.add_all([asset1, asset2, asset3, dag1, dag2, dag3, dag_ref1, dag_ref2, task_ref1]) session.commit() response = self.client.get( f"/api/v1/datasets?dag_ids={dag_ids}&uri_pattern={uri_pattern}", @@ -333,8 +333,8 @@ class TestGetDatasetsEndpointPagination(TestDatasetEndpoint): ) @provide_session def test_limit_and_offset(self, url, expected_dataset_uris, session): - datasets = [ - DatasetModel( + assets = [ + AssetModel( uri=f"s3://bucket/key/{i}", extra={"foo": "bar"}, created_at=timezone.parse(self.default_time), @@ -342,7 +342,7 @@ def test_limit_and_offset(self, url, expected_dataset_uris, session): ) for i in range(1, 110) ] - session.add_all(datasets) + session.add_all(assets) session.commit() response = self.client.get(url, environ_overrides={"REMOTE_USER": "test"}) @@ -352,8 +352,8 @@ def test_limit_and_offset(self, url, expected_dataset_uris, session): assert dataset_uris == expected_dataset_uris def test_should_respect_page_size_limit_default(self, session): - datasets = [ - DatasetModel( + assets = [ + AssetModel( uri=f"s3://bucket/key/{i}", extra={"foo": "bar"}, created_at=timezone.parse(self.default_time), @@ -361,7 +361,7 @@ def test_should_respect_page_size_limit_default(self, session): ) for i in range(1, 110) ] - session.add_all(datasets) + session.add_all(assets) session.commit() response = self.client.get("/api/v1/datasets", environ_overrides={"REMOTE_USER": "test"}) @@ -371,8 +371,8 @@ def test_should_respect_page_size_limit_default(self, session): @conf_vars({("api", "maximum_page_limit"): "150"}) def test_should_return_conf_max_if_req_max_above_conf(self, session): - datasets = [ - DatasetModel( + assets = [ + AssetModel( uri=f"s3://bucket/key/{i}", extra={"foo": "bar"}, created_at=timezone.parse(self.default_time), @@ -380,7 +380,7 @@ def test_should_return_conf_max_if_req_max_above_conf(self, session): ) for i in range(1, 200) ] - session.add_all(datasets) + session.add_all(assets) session.commit() response = self.client.get("/api/v1/datasets?limit=180", environ_overrides={"REMOTE_USER": "test"}) @@ -402,10 +402,10 @@ def test_should_respond_200(self, session): "created_dagruns": [], } - events = [DatasetEvent(id=i, timestamp=timezone.parse(self.default_time), **common) for i in [1, 2]] + events = [AssetEvent(id=i, timestamp=timezone.parse(self.default_time), **common) for i in [1, 2]] session.add_all(events) session.commit() - assert session.query(DatasetEvent).count() == 2 + assert session.query(AssetEvent).count() == 2 response = self.client.get("/api/v1/datasets/events", environ_overrides={"REMOTE_USER": "test"}) @@ -441,8 +441,8 @@ def test_should_respond_200(self, session): ) @provide_session def test_filtering(self, attr, value, session): - datasets = [ - DatasetModel( + assets = [ + AssetModel( id=i, uri=f"s3://bucket/key/{i}", extra={"foo": "bar"}, @@ -451,10 +451,10 @@ def test_filtering(self, attr, value, session): ) for i in [1, 2, 3] ] - session.add_all(datasets) + session.add_all(assets) session.commit() events = [ - DatasetEvent( + AssetEvent( id=i, dataset_id=i, source_dag_id=f"dag{i}", @@ -467,7 +467,7 @@ def test_filtering(self, attr, value, session): ] session.add_all(events) session.commit() - assert session.query(DatasetEvent).count() == 3 + assert session.query(AssetEvent).count() == 3 response = self.client.get( f"/api/v1/datasets/events?{attr}={value}", environ_overrides={"REMOTE_USER": "test"} @@ -480,7 +480,7 @@ def test_filtering(self, attr, value, session): { "id": 2, "dataset_id": 2, - "dataset_uri": datasets[1].uri, + "dataset_uri": assets[1].uri, "extra": {}, "source_dag_id": "dag2", "source_task_id": "task2", @@ -496,7 +496,7 @@ def test_filtering(self, attr, value, session): def test_order_by_raises_400_for_invalid_attr(self, session): self._create_dataset(session) events = [ - DatasetEvent( + AssetEvent( dataset_id=1, extra="{'foo': 'bar'}", source_dag_id="foo", @@ -509,7 +509,7 @@ def test_order_by_raises_400_for_invalid_attr(self, session): ] session.add_all(events) session.commit() - assert session.query(DatasetEvent).count() == 2 + assert session.query(AssetEvent).count() == 2 response = self.client.get( "/api/v1/datasets/events?order_by=fake", environ_overrides={"REMOTE_USER": "test"} @@ -525,7 +525,7 @@ def test_should_raises_401_unauthenticated(self, session): def test_includes_created_dagrun(self, session): self._create_dataset(session) - event = DatasetEvent( + event = AssetEvent( id=1, dataset_id=1, timestamp=timezone.parse(self.default_time), @@ -685,7 +685,7 @@ class TestGetDatasetEventsEndpointPagination(TestDatasetEndpoint): def test_limit_and_offset(self, url, expected_event_runids, session): self._create_dataset(session) events = [ - DatasetEvent( + AssetEvent( dataset_id=1, source_dag_id="foo", source_task_id="bar", @@ -707,7 +707,7 @@ def test_limit_and_offset(self, url, expected_event_runids, session): def test_should_respect_page_size_limit_default(self, session): self._create_dataset(session) events = [ - DatasetEvent( + AssetEvent( dataset_id=1, source_dag_id="foo", source_task_id="bar", @@ -729,7 +729,7 @@ def test_should_respect_page_size_limit_default(self, session): def test_should_return_conf_max_if_req_max_above_conf(self, session): self._create_dataset(session) events = [ - DatasetEvent( + AssetEvent( dataset_id=1, source_dag_id="foo", source_task_id="bar", @@ -761,10 +761,10 @@ def time_freezer(self) -> Generator: freezer.stop() def _create_dataset_dag_run_queues(self, dag_id, dataset_id, session): - ddrq = DatasetDagRunQueue(target_dag_id=dag_id, dataset_id=dataset_id) - session.add(ddrq) + adrq = AssetDagRunQueue(target_dag_id=dag_id, dataset_id=dataset_id) + session.add(adrq) session.commit() - return ddrq + return adrq class TestGetDagDatasetQueuedEvent(TestQueuedEventEndpoint): @@ -799,7 +799,7 @@ def test_should_respond_404(self): assert response.status_code == 404 assert { - "detail": "Queue event with dag_id: `not_exists` and dataset uri: `not_exists` was not found", + "detail": "Queue event with dag_id: `not_exists` and asset uri: `not_exists` was not found", "status": 404, "title": "Queue event not found", "type": EXCEPTIONS_LINK_MAP[404], @@ -832,10 +832,10 @@ def test_delete_should_respond_204(self, session, create_dummy_dag): dataset_uri = "s3://bucket/key" dataset_id = self._create_dataset(session).id - ddrq = DatasetDagRunQueue(target_dag_id=dag_id, dataset_id=dataset_id) - session.add(ddrq) + adrq = AssetDagRunQueue(target_dag_id=dag_id, dataset_id=dataset_id) + session.add(adrq) session.commit() - conn = session.query(DatasetDagRunQueue).all() + conn = session.query(AssetDagRunQueue).all() assert len(conn) == 1 response = self.client.delete( @@ -844,7 +844,7 @@ def test_delete_should_respond_204(self, session, create_dummy_dag): ) assert response.status_code == 204 - conn = session.query(DatasetDagRunQueue).all() + conn = session.query(AssetDagRunQueue).all() assert len(conn) == 0 _check_last_log( session, dag_id=dag_id, event="api.delete_dag_dataset_queued_event", execution_date=None @@ -861,7 +861,7 @@ def test_should_respond_404(self): assert response.status_code == 404 assert { - "detail": "Queue event with dag_id: `not_exists` and dataset uri: `not_exists` was not found", + "detail": "Queue event with dag_id: `not_exists` and asset uri: `not_exists` was not found", "status": 404, "title": "Queue event not found", "type": EXCEPTIONS_LINK_MAP[404], @@ -1013,7 +1013,7 @@ def test_should_respond_404(self): assert response.status_code == 404 assert { - "detail": "Queue event with dataset uri: `not_exists` was not found", + "detail": "Queue event with asset uri: `not_exists` was not found", "status": 404, "title": "Queue event not found", "type": EXCEPTIONS_LINK_MAP[404], @@ -1051,7 +1051,7 @@ def test_delete_should_respond_204(self, session, create_dummy_dag): ) assert response.status_code == 204 - conn = session.query(DatasetDagRunQueue).all() + conn = session.query(AssetDagRunQueue).all() assert len(conn) == 0 _check_last_log(session, dag_id=None, event="api.delete_dataset_queued_events", execution_date=None) @@ -1065,7 +1065,7 @@ def test_should_respond_404(self): assert response.status_code == 404 assert { - "detail": "Queue event with dataset uri: `not_exists` was not found", + "detail": "Queue event with asset uri: `not_exists` was not found", "status": 404, "title": "Queue event not found", "type": EXCEPTIONS_LINK_MAP[404], diff --git a/tests/api_connexion/schemas/test_dag_schema.py b/tests/api_connexion/schemas/test_dag_schema.py index a4a86bc05cc9b..1a7e345421c62 100644 --- a/tests/api_connexion/schemas/test_dag_schema.py +++ b/tests/api_connexion/schemas/test_dag_schema.py @@ -27,7 +27,7 @@ DAGDetailSchema, DAGSchema, ) -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.models import DagModel, DagTag from airflow.models.dag import DAG @@ -210,9 +210,9 @@ def test_serialize_test_dag_detail_schema(url_safe_serializer): @pytest.mark.skip_if_database_isolation_mode @pytest.mark.db_test -def test_serialize_test_dag_with_dataset_schedule_detail_schema(url_safe_serializer): - dataset1 = Dataset(uri="s3://bucket/obj1") - dataset2 = Dataset(uri="s3://bucket/obj2") +def test_serialize_test_dag_with_asset_schedule_detail_schema(url_safe_serializer): + asset1 = Asset(uri="s3://bucket/obj1") + asset2 = Asset(uri="s3://bucket/obj2") dag = DAG( dag_id="test_dag", start_date=datetime(2020, 6, 19), @@ -220,7 +220,7 @@ def test_serialize_test_dag_with_dataset_schedule_detail_schema(url_safe_seriali orientation="LR", default_view="duration", params={"foo": 1}, - schedule=dataset1 & dataset2, + schedule=asset1 & asset2, tags=["example1", "example2"], ) schema = DAGDetailSchema() @@ -255,7 +255,7 @@ def test_serialize_test_dag_with_dataset_schedule_detail_schema(url_safe_seriali key=lambda val: val["name"], ), "template_searchpath": None, - "timetable_summary": "Dataset", + "timetable_summary": "Asset", "timezone": UTC_JSON_REPR, "max_active_runs": 16, "max_consecutive_failed_dag_runs": 0, diff --git a/tests/api_connexion/schemas/test_dataset_schema.py b/tests/api_connexion/schemas/test_dataset_schema.py index c07eed2236a67..a9a5ce9e9673b 100644 --- a/tests/api_connexion/schemas/test_dataset_schema.py +++ b/tests/api_connexion/schemas/test_dataset_schema.py @@ -19,26 +19,26 @@ import pytest import time_machine -from airflow.api_connexion.schemas.dataset_schema import ( - DatasetCollection, - DatasetEventCollection, - dataset_collection_schema, - dataset_event_collection_schema, - dataset_event_schema, - dataset_schema, +from airflow.api_connexion.schemas.asset_schema import ( + AssetCollection, + AssetEventCollection, + asset_collection_schema, + asset_event_collection_schema, + asset_event_schema, + asset_schema, ) -from airflow.datasets import Dataset -from airflow.models.dataset import DatasetAliasModel, DatasetEvent, DatasetModel +from airflow.assets import Asset +from airflow.models.asset import AssetAliasModel, AssetEvent, AssetModel from airflow.operators.empty import EmptyOperator -from tests.test_utils.db import clear_db_dags, clear_db_datasets +from tests.test_utils.db import clear_db_assets, clear_db_dags pytestmark = [pytest.mark.db_test, pytest.mark.skip_if_database_isolation_mode] -class TestDatasetSchemaBase: +class TestAssetSchemaBase: def setup_method(self) -> None: clear_db_dags() - clear_db_datasets() + clear_db_assets() self.timestamp = "2022-06-10T12:02:44+00:00" self.freezer = time_machine.travel(self.timestamp, tick=False) self.freezer.start() @@ -46,12 +46,12 @@ def setup_method(self) -> None: def teardown_method(self) -> None: self.freezer.stop() clear_db_dags() - clear_db_datasets() + clear_db_assets() -class TestDatasetSchema(TestDatasetSchemaBase): +class TestAssetSchema(TestAssetSchemaBase): def test_serialize(self, dag_maker, session): - dataset = Dataset( + dataset = Asset( uri="s3://bucket/key", extra={"foo": "bar"}, ) @@ -62,9 +62,9 @@ def test_serialize(self, dag_maker, session): ): EmptyOperator(task_id="task2") - dataset_model = session.query(DatasetModel).filter_by(uri=dataset.uri).one() + asset_model = session.query(AssetModel).filter_by(uri=dataset.uri).one() - serialized_data = dataset_schema.dump(dataset_model) + serialized_data = asset_schema.dump(asset_model) serialized_data["id"] = 1 assert serialized_data == { "id": 1, @@ -91,24 +91,22 @@ def test_serialize(self, dag_maker, session): } -class TestDatasetCollectionSchema(TestDatasetSchemaBase): +class TestAssetCollectionSchema(TestAssetSchemaBase): def test_serialize(self, session): - datasets = [ - DatasetModel( + assets = [ + AssetModel( uri=f"s3://bucket/key/{i+1}", extra={"foo": "bar"}, ) for i in range(2) ] - dataset_aliases = [DatasetAliasModel(name=f"alias_{i}") for i in range(2)] - for dataset_alias in dataset_aliases: - dataset_alias.datasets.append(datasets[0]) - session.add_all(datasets) - session.add_all(dataset_aliases) + asset_aliases = [AssetAliasModel(name=f"alias_{i}") for i in range(2)] + for asset_alias in asset_aliases: + asset_alias.datasets.append(assets[0]) + session.add_all(assets) + session.add_all(asset_aliases) session.flush() - serialized_data = dataset_collection_schema.dump( - DatasetCollection(datasets=datasets, total_entries=2) - ) + serialized_data = asset_collection_schema.dump(AssetCollection(datasets=assets, total_entries=2)) serialized_data["datasets"][0]["id"] = 1 serialized_data["datasets"][1]["id"] = 2 serialized_data["datasets"][0]["aliases"][0]["id"] = 1 @@ -143,14 +141,14 @@ def test_serialize(self, session): } -class TestDatasetEventSchema(TestDatasetSchemaBase): +class TestAssetEventSchema(TestAssetSchemaBase): def test_serialize(self, session): - d = DatasetModel("s3://abc") - session.add(d) + assetssetsset = AssetModel("s3://abc") + session.add(assetssetsset) session.commit() - event = DatasetEvent( + event = AssetEvent( id=1, - dataset_id=d.id, + dataset_id=assetssetsset.id, extra={"foo": "bar"}, source_dag_id="foo", source_task_id="bar", @@ -159,10 +157,10 @@ def test_serialize(self, session): ) session.add(event) session.flush() - serialized_data = dataset_event_schema.dump(event) + serialized_data = asset_event_schema.dump(event) assert serialized_data == { "id": 1, - "dataset_id": d.id, + "dataset_id": assetssetsset.id, "dataset_uri": "s3://abc", "extra": {"foo": "bar"}, "source_dag_id": "foo", @@ -174,14 +172,14 @@ def test_serialize(self, session): } -class TestDatasetEventCreateSchema(TestDatasetSchemaBase): +class TestDatasetEventCreateSchema(TestAssetSchemaBase): def test_serialize(self, session): - d = DatasetModel("s3://abc") - session.add(d) + asset = AssetModel("s3://abc") + session.add(asset) session.commit() - event = DatasetEvent( + event = AssetEvent( id=1, - dataset_id=d.id, + dataset_id=asset.id, extra={"foo": "bar"}, source_dag_id=None, source_task_id=None, @@ -190,10 +188,10 @@ def test_serialize(self, session): ) session.add(event) session.flush() - serialized_data = dataset_event_schema.dump(event) + serialized_data = asset_event_schema.dump(event) assert serialized_data == { "id": 1, - "dataset_id": d.id, + "dataset_id": asset.id, "dataset_uri": "s3://abc", "extra": {"foo": "bar"}, "source_dag_id": None, @@ -205,7 +203,7 @@ def test_serialize(self, session): } -class TestDatasetEventCollectionSchema(TestDatasetSchemaBase): +class TestAssetEventCollectionSchema(TestAssetSchemaBase): def test_serialize(self, session): common = { "dataset_id": 10, @@ -217,11 +215,11 @@ def test_serialize(self, session): "created_dagruns": [], } - events = [DatasetEvent(id=i, **common) for i in [1, 2]] + events = [AssetEvent(id=i, **common) for i in [1, 2]] session.add_all(events) session.flush() - serialized_data = dataset_event_collection_schema.dump( - DatasetEventCollection(dataset_events=events, total_entries=2) + serialized_data = asset_event_collection_schema.dump( + AssetEventCollection(dataset_events=events, total_entries=2) ) assert serialized_data == { "dataset_events": [ diff --git a/tests/api_fastapi/views/ui/test_datasets.py b/tests/api_fastapi/views/ui/test_assets.py similarity index 92% rename from tests/api_fastapi/views/ui/test_datasets.py rename to tests/api_fastapi/views/ui/test_assets.py index 12b22e4bbb9ef..7aff14249b4db 100644 --- a/tests/api_fastapi/views/ui/test_datasets.py +++ b/tests/api_fastapi/views/ui/test_assets.py @@ -18,7 +18,7 @@ import pytest -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.operators.empty import EmptyOperator from tests.conftest import initial_db_init @@ -35,7 +35,7 @@ def cleanup(): def test_next_run_datasets(test_client, dag_maker): - with dag_maker(dag_id="upstream", schedule=[Dataset(uri="s3://bucket/key/1")], serialized=True): + with dag_maker(dag_id="upstream", schedule=[Asset(uri="s3://bucket/key/1")], serialized=True): EmptyOperator(task_id="task1") dag_maker.create_dagrun() diff --git a/tests/providers/common/io/datasets/__init__.py b/tests/assets/__init__.py similarity index 100% rename from tests/providers/common/io/datasets/__init__.py rename to tests/assets/__init__.py diff --git a/tests/datasets/test_manager.py b/tests/assets/test_manager.py similarity index 52% rename from tests/datasets/test_manager.py rename to tests/assets/test_manager.py index 9b8b0c180d48e..0539fdace52ba 100644 --- a/tests/datasets/test_manager.py +++ b/tests/assets/test_manager.py @@ -24,13 +24,13 @@ import pytest from sqlalchemy import delete -from airflow.datasets import Dataset -from airflow.datasets.manager import DatasetManager +from airflow.assets import Asset +from airflow.assets.manager import AssetManager from airflow.listeners.listener import get_listener_manager +from airflow.models.asset import AssetDagRunQueue, AssetEvent, AssetModel, DagScheduleAssetReference from airflow.models.dag import DagModel -from airflow.models.dataset import DagScheduleDatasetReference, DatasetDagRunQueue, DatasetEvent, DatasetModel from airflow.serialization.pydantic.taskinstance import TaskInstancePydantic -from tests.listeners import dataset_listener +from tests.listeners import asset_listener pytestmark = pytest.mark.db_test @@ -90,93 +90,93 @@ def create_mock_dag(): yield mock_dag -class TestDatasetManager: - def test_register_dataset_change_dataset_doesnt_exist(self, mock_task_instance): - dsem = DatasetManager() +class TestAssetManager: + def test_register_asset_change_asset_doesnt_exist(self, mock_task_instance): + dsem = AssetManager() - dataset = Dataset(uri="dataset_doesnt_exist") + asset = Asset(uri="asset_doesnt_exist") mock_session = mock.Mock() # Gotta mock up the query results mock_session.scalar.return_value = None - dsem.register_dataset_change(task_instance=mock_task_instance, dataset=dataset, session=mock_session) + dsem.register_asset_change(task_instance=mock_task_instance, asset=asset, session=mock_session) - # Ensure that we have ignored the dataset and _not_ created a DatasetEvent or - # DatasetDagRunQueue rows + # Ensure that we have ignored the asset and _not_ created a AssetEvent or + # AssetDagRunQueue rows mock_session.add.assert_not_called() mock_session.merge.assert_not_called() - def test_register_dataset_change(self, session, dag_maker, mock_task_instance): - dsem = DatasetManager() + def test_register_asset_change(self, session, dag_maker, mock_task_instance): + dsem = AssetManager() - ds = Dataset(uri="test_dataset_uri") + ds = Asset(uri="test_asset_uri") dag1 = DagModel(dag_id="dag1", is_active=True) dag2 = DagModel(dag_id="dag2", is_active=True) session.add_all([dag1, dag2]) - dsm = DatasetModel(uri="test_dataset_uri") - session.add(dsm) - dsm.consuming_dags = [DagScheduleDatasetReference(dag_id=dag.dag_id) for dag in (dag1, dag2)] - session.execute(delete(DatasetDagRunQueue)) + asm = AssetModel(uri="test_asset_uri") + session.add(asm) + asm.consuming_dags = [DagScheduleAssetReference(dag_id=dag.dag_id) for dag in (dag1, dag2)] + session.execute(delete(AssetDagRunQueue)) session.flush() - dsem.register_dataset_change(task_instance=mock_task_instance, dataset=ds, session=session) + dsem.register_asset_change(task_instance=mock_task_instance, asset=ds, session=session) session.flush() - # Ensure we've created a dataset - assert session.query(DatasetEvent).filter_by(dataset_id=dsm.id).count() == 1 - assert session.query(DatasetDagRunQueue).count() == 2 + # Ensure we've created an asset + assert session.query(AssetEvent).filter_by(dataset_id=asm.id).count() == 1 + assert session.query(AssetDagRunQueue).count() == 2 - def test_register_dataset_change_no_downstreams(self, session, mock_task_instance): - dsem = DatasetManager() + def test_register_asset_change_no_downstreams(self, session, mock_task_instance): + dsem = AssetManager() - ds = Dataset(uri="never_consumed") - dsm = DatasetModel(uri="never_consumed") - session.add(dsm) - session.execute(delete(DatasetDagRunQueue)) + ds = Asset(uri="never_consumed") + asm = AssetModel(uri="never_consumed") + session.add(asm) + session.execute(delete(AssetDagRunQueue)) session.flush() - dsem.register_dataset_change(task_instance=mock_task_instance, dataset=ds, session=session) + dsem.register_asset_change(task_instance=mock_task_instance, asset=ds, session=session) session.flush() - # Ensure we've created a dataset - assert session.query(DatasetEvent).filter_by(dataset_id=dsm.id).count() == 1 - assert session.query(DatasetDagRunQueue).count() == 0 + # Ensure we've created an asset + assert session.query(AssetEvent).filter_by(dataset_id=asm.id).count() == 1 + assert session.query(AssetDagRunQueue).count() == 0 @pytest.mark.skip_if_database_isolation_mode - def test_register_dataset_change_notifies_dataset_listener(self, session, mock_task_instance): - dsem = DatasetManager() - dataset_listener.clear() - get_listener_manager().add_listener(dataset_listener) + def test_register_asset_change_notifies_asset_listener(self, session, mock_task_instance): + dsem = AssetManager() + asset_listener.clear() + get_listener_manager().add_listener(asset_listener) - ds = Dataset(uri="test_dataset_uri_2") + ds = Asset(uri="test_asset_uri_2") dag1 = DagModel(dag_id="dag3") session.add(dag1) - dsm = DatasetModel(uri="test_dataset_uri_2") - session.add(dsm) - dsm.consuming_dags = [DagScheduleDatasetReference(dag_id=dag1.dag_id)] + asm = AssetModel(uri="test_asset_uri_2") + session.add(asm) + asm.consuming_dags = [DagScheduleAssetReference(dag_id=dag1.dag_id)] session.flush() - dsem.register_dataset_change(task_instance=mock_task_instance, dataset=ds, session=session) + dsem.register_asset_change(task_instance=mock_task_instance, asset=ds, session=session) session.flush() # Ensure the listener was notified - assert len(dataset_listener.changed) == 1 - assert dataset_listener.changed[0].uri == ds.uri + assert len(asset_listener.changed) == 1 + assert asset_listener.changed[0].uri == ds.uri @pytest.mark.skip_if_database_isolation_mode - def test_create_datasets_notifies_dataset_listener(self, session): - dsem = DatasetManager() - dataset_listener.clear() - get_listener_manager().add_listener(dataset_listener) + def test_create_assets_notifies_asset_listener(self, session): + asset_manager = AssetManager() + asset_listener.clear() + get_listener_manager().add_listener(asset_listener) - ds = Dataset(uri="test_dataset_uri_3") + asset = Asset(uri="test_asset_uri_3") - dsms = dsem.create_datasets([ds], session=session) + asms = asset_manager.create_assets([asset], session=session) # Ensure the listener was notified - assert len(dataset_listener.created) == 1 - assert len(dsms) == 1 - assert dataset_listener.created[0].uri == ds.uri == dsms[0].uri + assert len(asset_listener.created) == 1 + assert len(asms) == 1 + assert asset_listener.created[0].uri == asset.uri == asms[0].uri diff --git a/tests/assets/tests_asset.py b/tests/assets/tests_asset.py new file mode 100644 index 0000000000000..da6ef8ee79e39 --- /dev/null +++ b/tests/assets/tests_asset.py @@ -0,0 +1,586 @@ +# 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 os +from collections import defaultdict +from typing import Callable +from unittest.mock import patch + +import pytest +from sqlalchemy.sql import select + +from airflow.assets import ( + Asset, + AssetAlias, + AssetAll, + AssetAny, + BaseAsset, + _AssetAliasCondition, + _get_normalized_scheme, + _sanitize_uri, +) +from airflow.models.asset import AssetAliasModel, AssetDagRunQueue, AssetModel +from airflow.models.serialized_dag import SerializedDagModel +from airflow.operators.empty import EmptyOperator +from airflow.serialization.serialized_objects import BaseSerialization, SerializedDAG +from tests.test_utils.config import conf_vars + + +@pytest.fixture +def clear_assets(): + from tests.test_utils.db import clear_db_assets + + clear_db_assets() + yield + clear_db_assets() + + +@pytest.mark.parametrize( + ["uri"], + [ + pytest.param("", id="empty"), + pytest.param("\n\t", id="whitespace"), + pytest.param("a" * 3001, id="too_long"), + pytest.param("airflow://xcom/dag/task", id="reserved_scheme"), + pytest.param("😊", id="non-ascii"), + ], +) +def test_invalid_uris(uri): + with pytest.raises(ValueError): + Asset(uri=uri) + + +@pytest.mark.parametrize( + "uri, normalized", + [ + pytest.param("foobar", "foobar", id="scheme-less"), + pytest.param("foo:bar", "foo:bar", id="scheme-less-colon"), + pytest.param("foo/bar", "foo/bar", id="scheme-less-slash"), + pytest.param("s3://bucket/key/path", "s3://bucket/key/path", id="normal"), + pytest.param("file:///123/456/", "file:///123/456", id="trailing-slash"), + ], +) +def test_uri_with_scheme(uri: str, normalized: str) -> None: + asset = Asset(uri) + EmptyOperator(task_id="task1", outlets=[asset]) + assert asset.uri == normalized + assert os.fspath(asset) == normalized + + +def test_uri_with_auth() -> None: + with pytest.warns(UserWarning) as record: + asset = Asset("ftp://user@localhost/foo.txt") + assert len(record) == 1 + assert str(record[0].message) == ( + "An Asset URI should not contain auth info (e.g. username or " + "password). It has been automatically dropped." + ) + EmptyOperator(task_id="task1", outlets=[asset]) + assert asset.uri == "ftp://localhost/foo.txt" + assert os.fspath(asset) == "ftp://localhost/foo.txt" + + +def test_uri_without_scheme(): + asset = Asset(uri="example_asset") + EmptyOperator(task_id="task1", outlets=[asset]) + + +def test_fspath(): + uri = "s3://example/asset" + asset = Asset(uri=uri) + assert os.fspath(asset) == uri + + +def test_equal_when_same_uri(): + uri = "s3://example/asset" + asset1 = Asset(uri=uri) + asset2 = Asset(uri=uri) + assert asset1 == asset2 + + +def test_not_equal_when_different_uri(): + asset1 = Asset(uri="s3://example/asset") + asset2 = Asset(uri="s3://other/asset") + assert asset1 != asset2 + + +def test_asset_logic_operations(): + result_or = asset1 | asset2 + assert isinstance(result_or, AssetAny) + result_and = asset1 & asset2 + assert isinstance(result_and, AssetAll) + + +def test_asset_iter_assets(): + assert list(asset1.iter_assets()) == [("s3://bucket1/data1", asset1)] + + +@pytest.mark.db_test +def test_asset_iter_asset_aliases(): + base_asset = AssetAll( + AssetAlias("example-alias-1"), + Asset("1"), + AssetAny( + Asset("2"), + AssetAlias("example-alias-2"), + Asset("3"), + AssetAll(AssetAlias("example-alias-3"), Asset("4"), AssetAlias("example-alias-4")), + ), + AssetAll(AssetAlias("example-alias-5"), Asset("5")), + ) + assert list(base_asset.iter_asset_aliases()) == [ + (f"example-alias-{i}", AssetAlias(f"example-alias-{i}")) for i in range(1, 6) + ] + + +def test_asset_evaluate(): + assert asset1.evaluate({"s3://bucket1/data1": True}) is True + assert asset1.evaluate({"s3://bucket1/data1": False}) is False + + +def test_asset_any_operations(): + result_or = (asset1 | asset2) | asset3 + assert isinstance(result_or, AssetAny) + assert len(result_or.objects) == 3 + result_and = (asset1 | asset2) & asset3 + assert isinstance(result_and, AssetAll) + + +def test_asset_all_operations(): + result_or = (asset1 & asset2) | asset3 + assert isinstance(result_or, AssetAny) + result_and = (asset1 & asset2) & asset3 + assert isinstance(result_and, AssetAll) + + +def test_assset_boolean_condition_evaluate_iter(): + """ + Tests _AssetBooleanCondition's evaluate and iter_assets methods through AssetAny and AssetAll. + Ensures AssetAny evaluate returns True with any true condition, AssetAll evaluate returns False if + any condition is false, and both classes correctly iterate over assets without duplication. + """ + any_condition = AssetAny(asset1, asset2) + all_condition = AssetAll(asset1, asset2) + assert any_condition.evaluate({"s3://bucket1/data1": False, "s3://bucket2/data2": True}) is True + assert all_condition.evaluate({"s3://bucket1/data1": True, "s3://bucket2/data2": False}) is False + + # Testing iter_assets indirectly through the subclasses + assets_any = dict(any_condition.iter_assets()) + assets_all = dict(all_condition.iter_assets()) + assert assets_any == {"s3://bucket1/data1": asset1, "s3://bucket2/data2": asset2} + assert assets_all == {"s3://bucket1/data1": asset1, "s3://bucket2/data2": asset2} + + +@pytest.mark.parametrize( + "inputs, scenario, expected", + [ + # Scenarios for AssetAny + ((True, True, True), "any", True), + ((True, True, False), "any", True), + ((True, False, True), "any", True), + ((True, False, False), "any", True), + ((False, False, True), "any", True), + ((False, True, False), "any", True), + ((False, True, True), "any", True), + ((False, False, False), "any", False), + # Scenarios for AssetAll + ((True, True, True), "all", True), + ((True, True, False), "all", False), + ((True, False, True), "all", False), + ((True, False, False), "all", False), + ((False, False, True), "all", False), + ((False, True, False), "all", False), + ((False, True, True), "all", False), + ((False, False, False), "all", False), + ], +) +def test_asset_logical_conditions_evaluation_and_serialization(inputs, scenario, expected): + class_ = AssetAny if scenario == "any" else AssetAll + assets = [Asset(uri=f"s3://abc/{i}") for i in range(123, 126)] + condition = class_(*assets) + + statuses = {asset.uri: status for asset, status in zip(assets, inputs)} + assert ( + condition.evaluate(statuses) == expected + ), f"Condition evaluation failed for inputs {inputs} and scenario '{scenario}'" + + # Serialize and deserialize the condition to test persistence + serialized = BaseSerialization.serialize(condition) + deserialized = BaseSerialization.deserialize(serialized) + assert deserialized.evaluate(statuses) == expected, "Serialization round-trip failed" + + +@pytest.mark.parametrize( + "status_values, expected_evaluation", + [ + ((False, True, True), False), # AssetAll requires all conditions to be True, but d1 is False + ((True, True, True), True), # All conditions are True + ((True, False, True), True), # d1 is True, and AssetAny condition (d2 or d3 being True) is met + ((True, False, False), False), # d1 is True, but neither d2 nor d3 meet the AssetAny condition + ], +) +def test_nested_asset_conditions_with_serialization(status_values, expected_evaluation): + # Define assets + d1 = Asset(uri="s3://abc/123") + d2 = Asset(uri="s3://abc/124") + d3 = Asset(uri="s3://abc/125") + + # Create a nested condition: AssetAll with d1 and AssetAny with d2 and d3 + nested_condition = AssetAll(d1, AssetAny(d2, d3)) + + statuses = { + d1.uri: status_values[0], + d2.uri: status_values[1], + d3.uri: status_values[2], + } + + assert nested_condition.evaluate(statuses) == expected_evaluation, "Initial evaluation mismatch" + + serialized_condition = BaseSerialization.serialize(nested_condition) + deserialized_condition = BaseSerialization.deserialize(serialized_condition) + + assert ( + deserialized_condition.evaluate(statuses) == expected_evaluation + ), "Post-serialization evaluation mismatch" + + +@pytest.fixture +def create_test_assets(session): + """Fixture to create test assets and corresponding models.""" + assets = [Asset(uri=f"hello{i}") for i in range(1, 3)] + for asset in assets: + session.add(AssetModel(uri=asset.uri)) + session.commit() + return assets + + +@pytest.mark.db_test +@pytest.mark.usefixtures("clear_assets") +def test_asset_trigger_setup_and_serialization(session, dag_maker, create_test_assets): + assets = create_test_assets + + # Create DAG with asset triggers + with dag_maker(schedule=AssetAny(*assets)) as dag: + EmptyOperator(task_id="hello") + + # Verify assets are set up correctly + assert isinstance(dag.timetable.asset_condition, AssetAny), "DAG assets should be an instance of AssetAny" + + # Round-trip the DAG through serialization + deserialized_dag = SerializedDAG.deserialize_dag(SerializedDAG.serialize_dag(dag)) + + # Verify serialization and deserialization integrity + assert isinstance( + deserialized_dag.timetable.asset_condition, AssetAny + ), "Deserialized assets should maintain type AssetAny" + assert ( + deserialized_dag.timetable.asset_condition.objects == dag.timetable.asset_condition.objects + ), "Deserialized assets should match original" + + +@pytest.mark.db_test +@pytest.mark.usefixtures("clear_assets") +def test_asset_dag_run_queue_processing(session, clear_assets, dag_maker, create_test_assets): + assets = create_test_assets + asset_models = session.query(AssetModel).all() + + with dag_maker(schedule=AssetAny(*assets)) as dag: + EmptyOperator(task_id="hello") + + # Add AssetDagRunQueue entries to simulate asset event processing + for am in asset_models: + session.add(AssetDagRunQueue(dataset_id=am.id, target_dag_id=dag.dag_id)) + session.commit() + + # Fetch and evaluate asset triggers for all DAGs affected by asset events + records = session.scalars(select(AssetDagRunQueue)).all() + dag_statuses = defaultdict(lambda: defaultdict(bool)) + for record in records: + dag_statuses[record.target_dag_id][record.dataset.uri] = True + + serialized_dags = session.execute( + select(SerializedDagModel).where(SerializedDagModel.dag_id.in_(dag_statuses.keys())) + ).fetchall() + + for (serialized_dag,) in serialized_dags: + dag = SerializedDAG.deserialize(serialized_dag.data) + for asset_uri, status in dag_statuses[dag.dag_id].items(): + cond = dag.timetable.asset_condition + assert cond.evaluate({asset_uri: status}), "DAG trigger evaluation failed" + + +@pytest.mark.db_test +@pytest.mark.usefixtures("clear_assets") +def test_dag_with_complex_asset_condition(session, dag_maker): + # Create Asset instances + d1 = Asset(uri="hello1") + d2 = Asset(uri="hello2") + + # Create and add AssetModel instances to the session + am1 = AssetModel(uri=d1.uri) + am2 = AssetModel(uri=d2.uri) + session.add_all([am1, am2]) + session.commit() + + # Setup a DAG with complex asset triggers (AssetAny with AssetAll) + with dag_maker(schedule=AssetAny(d1, AssetAll(d2, d1))) as dag: + EmptyOperator(task_id="hello") + + assert isinstance( + dag.timetable.asset_condition, AssetAny + ), "DAG's asset trigger should be an instance of AssetAny" + assert any( + isinstance(trigger, AssetAll) for trigger in dag.timetable.asset_condition.objects + ), "DAG's asset trigger should include AssetAll" + + serialized_triggers = SerializedDAG.serialize(dag.timetable.asset_condition) + + deserialized_triggers = SerializedDAG.deserialize(serialized_triggers) + + assert isinstance( + deserialized_triggers, AssetAny + ), "Deserialized triggers should be an instance of AssetAny" + assert any( + isinstance(trigger, AssetAll) for trigger in deserialized_triggers.objects + ), "Deserialized triggers should include AssetAll" + + serialized_timetable_dict = SerializedDAG.to_dict(dag)["dag"]["timetable"]["__var"] + assert ( + "asset_condition" in serialized_timetable_dict + ), "Serialized timetable should contain 'asset_condition'" + assert isinstance( + serialized_timetable_dict["asset_condition"], dict + ), "Serialized 'asset_condition' should be a dict" + + +def assets_equal(a1: BaseAsset, a2: BaseAsset) -> bool: + if type(a1) is not type(a2): + return False + + if isinstance(a1, Asset) and isinstance(a2, Asset): + return a1.uri == a2.uri + + elif isinstance(a1, (AssetAny, AssetAll)) and isinstance(a2, (AssetAny, AssetAll)): + if len(a1.objects) != len(a2.objects): + return False + + # Compare each pair of objects + for obj1, obj2 in zip(a1.objects, a2.objects): + # If obj1 or obj2 is a Asset, AssetAny, or AssetAll instance, + # recursively call assets_equal + if not assets_equal(obj1, obj2): + return False + return True + + return False + + +asset1 = Asset(uri="s3://bucket1/data1") +asset2 = Asset(uri="s3://bucket2/data2") +asset3 = Asset(uri="s3://bucket3/data3") +asset4 = Asset(uri="s3://bucket4/data4") +asset5 = Asset(uri="s3://bucket5/data5") + +test_cases = [ + (lambda: asset1, asset1), + (lambda: asset1 & asset2, AssetAll(asset1, asset2)), + (lambda: asset1 | asset2, AssetAny(asset1, asset2)), + (lambda: asset1 | (asset2 & asset3), AssetAny(asset1, AssetAll(asset2, asset3))), + (lambda: asset1 | asset2 & asset3, AssetAny(asset1, AssetAll(asset2, asset3))), + ( + lambda: ((asset1 & asset2) | asset3) & (asset4 | asset5), + AssetAll(AssetAny(AssetAll(asset1, asset2), asset3), AssetAny(asset4, asset5)), + ), + (lambda: asset1 & asset2 | asset3, AssetAny(AssetAll(asset1, asset2), asset3)), + ( + lambda: (asset1 | asset2) & (asset3 | asset4), + AssetAll(AssetAny(asset1, asset2), AssetAny(asset3, asset4)), + ), + ( + lambda: (asset1 & asset2) | (asset3 & (asset4 | asset5)), + AssetAny(AssetAll(asset1, asset2), AssetAll(asset3, AssetAny(asset4, asset5))), + ), + ( + lambda: (asset1 & asset2) & (asset3 & asset4), + AssetAll(asset1, asset2, AssetAll(asset3, asset4)), + ), + (lambda: asset1 | asset2 | asset3, AssetAny(asset1, asset2, asset3)), + (lambda: asset1 & asset2 & asset3, AssetAll(asset1, asset2, asset3)), + ( + lambda: ((asset1 & asset2) | asset3) & (asset4 | asset5), + AssetAll(AssetAny(AssetAll(asset1, asset2), asset3), AssetAny(asset4, asset5)), + ), +] + + +@pytest.mark.parametrize("expression, expected", test_cases) +def test_evaluate_assets_expression(expression, expected): + expr = expression() + assert assets_equal(expr, expected) + + +@pytest.mark.parametrize( + "expression, error", + [ + pytest.param( + lambda: asset1 & 1, # type: ignore[operator] + "unsupported operand type(s) for &: 'Asset' and 'int'", + id="&", + ), + pytest.param( + lambda: asset1 | 1, # type: ignore[operator] + "unsupported operand type(s) for |: 'Asset' and 'int'", + id="|", + ), + pytest.param( + lambda: AssetAll(1, asset1), # type: ignore[arg-type] + "expect asset expressions in condition", + id="AssetAll", + ), + pytest.param( + lambda: AssetAny(1, asset1), # type: ignore[arg-type] + "expect asset expressions in condition", + id="AssetAny", + ), + ], +) +def test_assets_expression_error(expression: Callable[[], None], error: str) -> None: + with pytest.raises(TypeError) as info: + expression() + assert str(info.value) == error + + +def test_get_normalized_scheme(): + assert _get_normalized_scheme("http://example.com") == "http" + assert _get_normalized_scheme("HTTPS://example.com") == "https" + assert _get_normalized_scheme("ftp://example.com") == "ftp" + assert _get_normalized_scheme("file://") == "file" + + assert _get_normalized_scheme("example.com") == "" + assert _get_normalized_scheme("") == "" + assert _get_normalized_scheme(" ") == "" + + +def _mock_get_uri_normalizer_raising_error(normalized_scheme): + def normalizer(uri): + raise ValueError("Incorrect URI format") + + return normalizer + + +def _mock_get_uri_normalizer_noop(normalized_scheme): + def normalizer(uri): + return uri + + return normalizer + + +@patch("airflow.assets._get_uri_normalizer", _mock_get_uri_normalizer_raising_error) +@patch("airflow.assets.warnings.warn") +def test_sanitize_uri_raises_warning(mock_warn): + _sanitize_uri("postgres://localhost:5432/database.schema.table") + msg = mock_warn.call_args.args[0] + assert "The Asset URI postgres://localhost:5432/database.schema.table is not AIP-60 compliant" in msg + assert "In Airflow 3, this will raise an exception." in msg + + +@patch("airflow.assets._get_uri_normalizer", _mock_get_uri_normalizer_raising_error) +@conf_vars({("core", "strict_asset_uri_validation"): "True"}) +def test_sanitize_uri_raises_exception(): + with pytest.raises(ValueError) as e_info: + _sanitize_uri("postgres://localhost:5432/database.schema.table") + assert isinstance(e_info.value, ValueError) + assert str(e_info.value) == "Incorrect URI format" + + +@patch("airflow.assets._get_uri_normalizer", lambda x: None) +def test_normalize_uri_no_normalizer_found(): + asset = Asset(uri="any_uri_without_normalizer_defined") + assert asset.normalized_uri is None + + +@patch("airflow.assets._get_uri_normalizer", _mock_get_uri_normalizer_raising_error) +def test_normalize_uri_invalid_uri(): + asset = Asset(uri="any_uri_not_aip60_compliant") + assert asset.normalized_uri is None + + +@patch("airflow.assets._get_uri_normalizer", _mock_get_uri_normalizer_noop) +@patch("airflow.assets._get_normalized_scheme", lambda x: "valid_scheme") +def test_normalize_uri_valid_uri(): + asset = Asset(uri="valid_aip60_uri") + assert asset.normalized_uri == "valid_aip60_uri" + + +@pytest.mark.skip_if_database_isolation_mode +@pytest.mark.db_test +@pytest.mark.usefixtures("clear_assets") +class Test_AssetAliasCondition: + @pytest.fixture + def asset_1(self, session): + """Example asset links to asset alias resolved_asset_alias_2.""" + asset_uri = "test_uri" + asset_1 = AssetModel(id=1, uri=asset_uri) + + session.add(asset_1) + session.commit() + + return asset_1 + + @pytest.fixture + def asset_alias_1(self, session): + """Example asset alias links to no assets.""" + alias_name = "test_name" + asset_alias_model = AssetAliasModel(name=alias_name) + + session.add(asset_alias_model) + session.commit() + + return asset_alias_model + + @pytest.fixture + def resolved_asset_alias_2(self, session, asset_1): + """Example asset alias links to asset asset_alias_1.""" + asset_name = "test_name_2" + asset_alias_2 = AssetAliasModel(name=asset_name) + asset_alias_2.datasets.append(asset_1) + + session.add(asset_alias_2) + session.commit() + + return asset_alias_2 + + def test_init(self, asset_alias_1, asset_1, resolved_asset_alias_2): + cond = _AssetAliasCondition(name=asset_alias_1.name) + assert cond.objects == [] + + cond = _AssetAliasCondition(name=resolved_asset_alias_2.name) + assert cond.objects == [Asset(uri=asset_1.uri)] + + def test_as_expression(self, asset_alias_1, resolved_asset_alias_2): + for assset_alias in (asset_alias_1, resolved_asset_alias_2): + cond = _AssetAliasCondition(assset_alias.name) + assert cond.as_expression() == {"alias": assset_alias.name} + + def test_evalute(self, asset_alias_1, resolved_asset_alias_2, asset_1): + cond = _AssetAliasCondition(asset_alias_1.name) + assert cond.evaluate({asset_1.uri: True}) is False + + cond = _AssetAliasCondition(resolved_asset_alias_2.name) + assert cond.evaluate({asset_1.uri: True}) is True diff --git a/tests/auth/managers/simple/test_simple_auth_manager.py b/tests/auth/managers/simple/test_simple_auth_manager.py index a11c79063d042..d4bd4e4fbfed2 100644 --- a/tests/auth/managers/simple/test_simple_auth_manager.py +++ b/tests/auth/managers/simple/test_simple_auth_manager.py @@ -140,7 +140,7 @@ def test_get_user_return_none_when_not_logged_in(self, mock_is_logged_in, auth_m "is_authorized_configuration", "is_authorized_connection", "is_authorized_dag", - "is_authorized_dataset", + "is_authorized_asset", "is_authorized_pool", "is_authorized_variable", ], @@ -206,7 +206,7 @@ def test_is_authorized_view_methods( [ "is_authorized_configuration", "is_authorized_connection", - "is_authorized_dataset", + "is_authorized_asset", "is_authorized_pool", "is_authorized_variable", ], @@ -258,7 +258,7 @@ def test_is_authorized_methods_user_role_required( @patch.object(SimpleAuthManager, "is_logged_in") @pytest.mark.parametrize( "api", - ["is_authorized_dag", "is_authorized_dataset", "is_authorized_pool"], + ["is_authorized_dag", "is_authorized_asset", "is_authorized_pool"], ) @pytest.mark.parametrize( "role, method, result", diff --git a/tests/auth/managers/test_base_auth_manager.py b/tests/auth/managers/test_base_auth_manager.py index cd9652fb465d0..82efe20048b71 100644 --- a/tests/auth/managers/test_base_auth_manager.py +++ b/tests/auth/managers/test_base_auth_manager.py @@ -35,9 +35,9 @@ from airflow.auth.managers.models.base_user import BaseUser from airflow.auth.managers.models.resource_details import ( AccessView, + AssetDetails, ConfigurationDetails, DagAccessEntity, - DatasetDetails, ) @@ -73,8 +73,8 @@ def is_authorized_dag( ) -> bool: raise NotImplementedError() - def is_authorized_dataset( - self, *, method: ResourceMethod, details: DatasetDetails | None = None, user: BaseUser | None = None + def is_authorized_asset( + self, *, method: ResourceMethod, details: AssetDetails | None = None, user: BaseUser | None = None ) -> bool: raise NotImplementedError() diff --git a/tests/conftest.py b/tests/conftest.py index 7e3affd2a1b27..60d009416fe8e 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -973,10 +973,10 @@ def __call__( def cleanup(self): from airflow.models import DagModel, DagRun, TaskInstance, XCom - from airflow.models.dataset import DatasetEvent from airflow.models.serialized_dag import SerializedDagModel from airflow.models.taskmap import TaskMap from airflow.utils.retries import run_with_db_retries + from tests.test_utils.compat import AssetEvent for attempt in run_with_db_retries(logger=self.log): with attempt: @@ -1004,7 +1004,7 @@ def cleanup(self): self.session.query(TaskMap).filter(TaskMap.dag_id.in_(dag_ids)).delete( synchronize_session=False, ) - self.session.query(DatasetEvent).filter(DatasetEvent.source_dag_id.in_(dag_ids)).delete( + self.session.query(AssetEvent).filter(AssetEvent.source_dag_id.in_(dag_ids)).delete( synchronize_session=False, ) self.session.commit() diff --git a/tests/dags/test_datasets.py b/tests/dags/test_assets.py similarity index 91% rename from tests/dags/test_datasets.py rename to tests/dags/test_assets.py index 4bdef9f6978cb..a4ecd6aad4a6a 100644 --- a/tests/dags/test_datasets.py +++ b/tests/dags/test_assets.py @@ -19,14 +19,14 @@ from datetime import datetime -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.exceptions import AirflowFailException, AirflowSkipException from airflow.models.dag import DAG from airflow.operators.bash import BashOperator from airflow.operators.python import PythonOperator -skip_task_dag_dataset = Dataset("s3://dag_with_skip_task/output_1.txt", extra={"hi": "bye"}) -fail_task_dag_dataset = Dataset("s3://dag_with_fail_task/output_1.txt", extra={"hi": "bye"}) +skip_task_dag_dataset = Asset("s3://dag_with_skip_task/output_1.txt", extra={"hi": "bye"}) +fail_task_dag_dataset = Asset("s3://dag_with_fail_task/output_1.txt", extra={"hi": "bye"}) def raise_skip_exc(): diff --git a/tests/dags/test_only_empty_tasks.py b/tests/dags/test_only_empty_tasks.py index 68f1dc5e897ae..2cea9c3c6b173 100644 --- a/tests/dags/test_only_empty_tasks.py +++ b/tests/dags/test_only_empty_tasks.py @@ -20,7 +20,7 @@ from datetime import datetime from typing import Sequence -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.models.dag import DAG from airflow.operators.empty import EmptyOperator @@ -56,4 +56,4 @@ def __init__(self, body, *args, **kwargs): EmptyOperator(task_id="test_task_on_success", on_success_callback=lambda *args, **kwargs: None) - EmptyOperator(task_id="test_task_outlets", outlets=[Dataset("hello")]) + EmptyOperator(task_id="test_task_outlets", outlets=[Asset("hello")]) diff --git a/tests/datasets/test_dataset.py b/tests/datasets/test_dataset.py deleted file mode 100644 index 8221a5aea8aa3..0000000000000 --- a/tests/datasets/test_dataset.py +++ /dev/null @@ -1,588 +0,0 @@ -# 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 os -from collections import defaultdict -from typing import Callable -from unittest.mock import patch - -import pytest -from sqlalchemy.sql import select - -from airflow.datasets import ( - BaseDataset, - Dataset, - DatasetAlias, - DatasetAll, - DatasetAny, - _DatasetAliasCondition, - _get_normalized_scheme, - _sanitize_uri, -) -from airflow.models.dataset import DatasetAliasModel, DatasetDagRunQueue, DatasetModel -from airflow.models.serialized_dag import SerializedDagModel -from airflow.operators.empty import EmptyOperator -from airflow.serialization.serialized_objects import BaseSerialization, SerializedDAG -from tests.test_utils.config import conf_vars - - -@pytest.fixture -def clear_datasets(): - from tests.test_utils.db import clear_db_datasets - - clear_db_datasets() - yield - clear_db_datasets() - - -@pytest.mark.parametrize( - ["uri"], - [ - pytest.param("", id="empty"), - pytest.param("\n\t", id="whitespace"), - pytest.param("a" * 3001, id="too_long"), - pytest.param("airflow://xcom/dag/task", id="reserved_scheme"), - pytest.param("😊", id="non-ascii"), - ], -) -def test_invalid_uris(uri): - with pytest.raises(ValueError): - Dataset(uri=uri) - - -@pytest.mark.parametrize( - "uri, normalized", - [ - pytest.param("foobar", "foobar", id="scheme-less"), - pytest.param("foo:bar", "foo:bar", id="scheme-less-colon"), - pytest.param("foo/bar", "foo/bar", id="scheme-less-slash"), - pytest.param("s3://bucket/key/path", "s3://bucket/key/path", id="normal"), - pytest.param("file:///123/456/", "file:///123/456", id="trailing-slash"), - ], -) -def test_uri_with_scheme(uri: str, normalized: str) -> None: - dataset = Dataset(uri) - EmptyOperator(task_id="task1", outlets=[dataset]) - assert dataset.uri == normalized - assert os.fspath(dataset) == normalized - - -def test_uri_with_auth() -> None: - with pytest.warns(UserWarning) as record: - dataset = Dataset("ftp://user@localhost/foo.txt") - assert len(record) == 1 - assert str(record[0].message) == ( - "A dataset URI should not contain auth info (e.g. username or " - "password). It has been automatically dropped." - ) - EmptyOperator(task_id="task1", outlets=[dataset]) - assert dataset.uri == "ftp://localhost/foo.txt" - assert os.fspath(dataset) == "ftp://localhost/foo.txt" - - -def test_uri_without_scheme(): - dataset = Dataset(uri="example_dataset") - EmptyOperator(task_id="task1", outlets=[dataset]) - - -def test_fspath(): - uri = "s3://example/dataset" - dataset = Dataset(uri=uri) - assert os.fspath(dataset) == uri - - -def test_equal_when_same_uri(): - uri = "s3://example/dataset" - dataset1 = Dataset(uri=uri) - dataset2 = Dataset(uri=uri) - assert dataset1 == dataset2 - - -def test_not_equal_when_different_uri(): - dataset1 = Dataset(uri="s3://example/dataset") - dataset2 = Dataset(uri="s3://other/dataset") - assert dataset1 != dataset2 - - -def test_dataset_logic_operations(): - result_or = dataset1 | dataset2 - assert isinstance(result_or, DatasetAny) - result_and = dataset1 & dataset2 - assert isinstance(result_and, DatasetAll) - - -def test_dataset_iter_datasets(): - assert list(dataset1.iter_datasets()) == [("s3://bucket1/data1", dataset1)] - - -@pytest.mark.db_test -def test_dataset_iter_dataset_aliases(): - base_dataset = DatasetAll( - DatasetAlias("example-alias-1"), - Dataset("1"), - DatasetAny( - Dataset("2"), - DatasetAlias("example-alias-2"), - Dataset("3"), - DatasetAll(DatasetAlias("example-alias-3"), Dataset("4"), DatasetAlias("example-alias-4")), - ), - DatasetAll(DatasetAlias("example-alias-5"), Dataset("5")), - ) - assert list(base_dataset.iter_dataset_aliases()) == [ - (f"example-alias-{i}", DatasetAlias(f"example-alias-{i}")) for i in range(1, 6) - ] - - -def test_dataset_evaluate(): - assert dataset1.evaluate({"s3://bucket1/data1": True}) is True - assert dataset1.evaluate({"s3://bucket1/data1": False}) is False - - -def test_dataset_any_operations(): - result_or = (dataset1 | dataset2) | dataset3 - assert isinstance(result_or, DatasetAny) - assert len(result_or.objects) == 3 - result_and = (dataset1 | dataset2) & dataset3 - assert isinstance(result_and, DatasetAll) - - -def test_dataset_all_operations(): - result_or = (dataset1 & dataset2) | dataset3 - assert isinstance(result_or, DatasetAny) - result_and = (dataset1 & dataset2) & dataset3 - assert isinstance(result_and, DatasetAll) - - -def test_datasetbooleancondition_evaluate_iter(): - """ - Tests _DatasetBooleanCondition's evaluate and iter_datasets methods through DatasetAny and DatasetAll. - Ensures DatasetAny evaluate returns True with any true condition, DatasetAll evaluate returns False if - any condition is false, and both classes correctly iterate over datasets without duplication. - """ - any_condition = DatasetAny(dataset1, dataset2) - all_condition = DatasetAll(dataset1, dataset2) - assert any_condition.evaluate({"s3://bucket1/data1": False, "s3://bucket2/data2": True}) is True - assert all_condition.evaluate({"s3://bucket1/data1": True, "s3://bucket2/data2": False}) is False - - # Testing iter_datasets indirectly through the subclasses - datasets_any = dict(any_condition.iter_datasets()) - datasets_all = dict(all_condition.iter_datasets()) - assert datasets_any == {"s3://bucket1/data1": dataset1, "s3://bucket2/data2": dataset2} - assert datasets_all == {"s3://bucket1/data1": dataset1, "s3://bucket2/data2": dataset2} - - -@pytest.mark.parametrize( - "inputs, scenario, expected", - [ - # Scenarios for DatasetAny - ((True, True, True), "any", True), - ((True, True, False), "any", True), - ((True, False, True), "any", True), - ((True, False, False), "any", True), - ((False, False, True), "any", True), - ((False, True, False), "any", True), - ((False, True, True), "any", True), - ((False, False, False), "any", False), - # Scenarios for DatasetAll - ((True, True, True), "all", True), - ((True, True, False), "all", False), - ((True, False, True), "all", False), - ((True, False, False), "all", False), - ((False, False, True), "all", False), - ((False, True, False), "all", False), - ((False, True, True), "all", False), - ((False, False, False), "all", False), - ], -) -def test_dataset_logical_conditions_evaluation_and_serialization(inputs, scenario, expected): - class_ = DatasetAny if scenario == "any" else DatasetAll - datasets = [Dataset(uri=f"s3://abc/{i}") for i in range(123, 126)] - condition = class_(*datasets) - - statuses = {dataset.uri: status for dataset, status in zip(datasets, inputs)} - assert ( - condition.evaluate(statuses) == expected - ), f"Condition evaluation failed for inputs {inputs} and scenario '{scenario}'" - - # Serialize and deserialize the condition to test persistence - serialized = BaseSerialization.serialize(condition) - deserialized = BaseSerialization.deserialize(serialized) - assert deserialized.evaluate(statuses) == expected, "Serialization round-trip failed" - - -@pytest.mark.parametrize( - "status_values, expected_evaluation", - [ - ((False, True, True), False), # DatasetAll requires all conditions to be True, but d1 is False - ((True, True, True), True), # All conditions are True - ((True, False, True), True), # d1 is True, and DatasetAny condition (d2 or d3 being True) is met - ((True, False, False), False), # d1 is True, but neither d2 nor d3 meet the DatasetAny condition - ], -) -def test_nested_dataset_conditions_with_serialization(status_values, expected_evaluation): - # Define datasets - d1 = Dataset(uri="s3://abc/123") - d2 = Dataset(uri="s3://abc/124") - d3 = Dataset(uri="s3://abc/125") - - # Create a nested condition: DatasetAll with d1 and DatasetAny with d2 and d3 - nested_condition = DatasetAll(d1, DatasetAny(d2, d3)) - - statuses = { - d1.uri: status_values[0], - d2.uri: status_values[1], - d3.uri: status_values[2], - } - - assert nested_condition.evaluate(statuses) == expected_evaluation, "Initial evaluation mismatch" - - serialized_condition = BaseSerialization.serialize(nested_condition) - deserialized_condition = BaseSerialization.deserialize(serialized_condition) - - assert ( - deserialized_condition.evaluate(statuses) == expected_evaluation - ), "Post-serialization evaluation mismatch" - - -@pytest.fixture -def create_test_datasets(session): - """Fixture to create test datasets and corresponding models.""" - datasets = [Dataset(uri=f"hello{i}") for i in range(1, 3)] - for dataset in datasets: - session.add(DatasetModel(uri=dataset.uri)) - session.commit() - return datasets - - -@pytest.mark.db_test -@pytest.mark.usefixtures("clear_datasets") -def test_dataset_trigger_setup_and_serialization(session, dag_maker, create_test_datasets): - datasets = create_test_datasets - - # Create DAG with dataset triggers - with dag_maker(schedule=DatasetAny(*datasets)) as dag: - EmptyOperator(task_id="hello") - - # Verify datasets are set up correctly - assert isinstance( - dag.timetable.dataset_condition, DatasetAny - ), "DAG datasets should be an instance of DatasetAny" - - # Round-trip the DAG through serialization - deserialized_dag = SerializedDAG.deserialize_dag(SerializedDAG.serialize_dag(dag)) - - # Verify serialization and deserialization integrity - assert isinstance( - deserialized_dag.timetable.dataset_condition, DatasetAny - ), "Deserialized datasets should maintain type DatasetAny" - assert ( - deserialized_dag.timetable.dataset_condition.objects == dag.timetable.dataset_condition.objects - ), "Deserialized datasets should match original" - - -@pytest.mark.db_test -@pytest.mark.usefixtures("clear_datasets") -def test_dataset_dag_run_queue_processing(session, clear_datasets, dag_maker, create_test_datasets): - datasets = create_test_datasets - dataset_models = session.query(DatasetModel).all() - - with dag_maker(schedule=DatasetAny(*datasets)) as dag: - EmptyOperator(task_id="hello") - - # Add DatasetDagRunQueue entries to simulate dataset event processing - for dm in dataset_models: - session.add(DatasetDagRunQueue(dataset_id=dm.id, target_dag_id=dag.dag_id)) - session.commit() - - # Fetch and evaluate dataset triggers for all DAGs affected by dataset events - records = session.scalars(select(DatasetDagRunQueue)).all() - dag_statuses = defaultdict(lambda: defaultdict(bool)) - for record in records: - dag_statuses[record.target_dag_id][record.dataset.uri] = True - - serialized_dags = session.execute( - select(SerializedDagModel).where(SerializedDagModel.dag_id.in_(dag_statuses.keys())) - ).fetchall() - - for (serialized_dag,) in serialized_dags: - dag = SerializedDAG.deserialize(serialized_dag.data) - for dataset_uri, status in dag_statuses[dag.dag_id].items(): - cond = dag.timetable.dataset_condition - assert cond.evaluate({dataset_uri: status}), "DAG trigger evaluation failed" - - -@pytest.mark.db_test -@pytest.mark.usefixtures("clear_datasets") -def test_dag_with_complex_dataset_condition(session, dag_maker): - # Create Dataset instances - d1 = Dataset(uri="hello1") - d2 = Dataset(uri="hello2") - - # Create and add DatasetModel instances to the session - dm1 = DatasetModel(uri=d1.uri) - dm2 = DatasetModel(uri=d2.uri) - session.add_all([dm1, dm2]) - session.commit() - - # Setup a DAG with complex dataset triggers (DatasetAny with DatasetAll) - with dag_maker(schedule=DatasetAny(d1, DatasetAll(d2, d1))) as dag: - EmptyOperator(task_id="hello") - - assert isinstance( - dag.timetable.dataset_condition, DatasetAny - ), "DAG's dataset trigger should be an instance of DatasetAny" - assert any( - isinstance(trigger, DatasetAll) for trigger in dag.timetable.dataset_condition.objects - ), "DAG's dataset trigger should include DatasetAll" - - serialized_triggers = SerializedDAG.serialize(dag.timetable.dataset_condition) - - deserialized_triggers = SerializedDAG.deserialize(serialized_triggers) - - assert isinstance( - deserialized_triggers, DatasetAny - ), "Deserialized triggers should be an instance of DatasetAny" - assert any( - isinstance(trigger, DatasetAll) for trigger in deserialized_triggers.objects - ), "Deserialized triggers should include DatasetAll" - - serialized_timetable_dict = SerializedDAG.to_dict(dag)["dag"]["timetable"]["__var"] - assert ( - "dataset_condition" in serialized_timetable_dict - ), "Serialized timetable should contain 'dataset_condition'" - assert isinstance( - serialized_timetable_dict["dataset_condition"], dict - ), "Serialized 'dataset_condition' should be a dict" - - -def datasets_equal(d1: BaseDataset, d2: BaseDataset) -> bool: - if type(d1) is not type(d2): - return False - - if isinstance(d1, Dataset) and isinstance(d2, Dataset): - return d1.uri == d2.uri - - elif isinstance(d1, (DatasetAny, DatasetAll)) and isinstance(d2, (DatasetAny, DatasetAll)): - if len(d1.objects) != len(d2.objects): - return False - - # Compare each pair of objects - for obj1, obj2 in zip(d1.objects, d2.objects): - # If obj1 or obj2 is a Dataset, DatasetAny, or DatasetAll instance, - # recursively call datasets_equal - if not datasets_equal(obj1, obj2): - return False - return True - - return False - - -dataset1 = Dataset(uri="s3://bucket1/data1") -dataset2 = Dataset(uri="s3://bucket2/data2") -dataset3 = Dataset(uri="s3://bucket3/data3") -dataset4 = Dataset(uri="s3://bucket4/data4") -dataset5 = Dataset(uri="s3://bucket5/data5") - -test_cases = [ - (lambda: dataset1, dataset1), - (lambda: dataset1 & dataset2, DatasetAll(dataset1, dataset2)), - (lambda: dataset1 | dataset2, DatasetAny(dataset1, dataset2)), - (lambda: dataset1 | (dataset2 & dataset3), DatasetAny(dataset1, DatasetAll(dataset2, dataset3))), - (lambda: dataset1 | dataset2 & dataset3, DatasetAny(dataset1, DatasetAll(dataset2, dataset3))), - ( - lambda: ((dataset1 & dataset2) | dataset3) & (dataset4 | dataset5), - DatasetAll(DatasetAny(DatasetAll(dataset1, dataset2), dataset3), DatasetAny(dataset4, dataset5)), - ), - (lambda: dataset1 & dataset2 | dataset3, DatasetAny(DatasetAll(dataset1, dataset2), dataset3)), - ( - lambda: (dataset1 | dataset2) & (dataset3 | dataset4), - DatasetAll(DatasetAny(dataset1, dataset2), DatasetAny(dataset3, dataset4)), - ), - ( - lambda: (dataset1 & dataset2) | (dataset3 & (dataset4 | dataset5)), - DatasetAny(DatasetAll(dataset1, dataset2), DatasetAll(dataset3, DatasetAny(dataset4, dataset5))), - ), - ( - lambda: (dataset1 & dataset2) & (dataset3 & dataset4), - DatasetAll(dataset1, dataset2, DatasetAll(dataset3, dataset4)), - ), - (lambda: dataset1 | dataset2 | dataset3, DatasetAny(dataset1, dataset2, dataset3)), - (lambda: dataset1 & dataset2 & dataset3, DatasetAll(dataset1, dataset2, dataset3)), - ( - lambda: ((dataset1 & dataset2) | dataset3) & (dataset4 | dataset5), - DatasetAll(DatasetAny(DatasetAll(dataset1, dataset2), dataset3), DatasetAny(dataset4, dataset5)), - ), -] - - -@pytest.mark.parametrize("expression, expected", test_cases) -def test_evaluate_datasets_expression(expression, expected): - expr = expression() - assert datasets_equal(expr, expected) - - -@pytest.mark.parametrize( - "expression, error", - [ - pytest.param( - lambda: dataset1 & 1, # type: ignore[operator] - "unsupported operand type(s) for &: 'Dataset' and 'int'", - id="&", - ), - pytest.param( - lambda: dataset1 | 1, # type: ignore[operator] - "unsupported operand type(s) for |: 'Dataset' and 'int'", - id="|", - ), - pytest.param( - lambda: DatasetAll(1, dataset1), # type: ignore[arg-type] - "expect dataset expressions in condition", - id="DatasetAll", - ), - pytest.param( - lambda: DatasetAny(1, dataset1), # type: ignore[arg-type] - "expect dataset expressions in condition", - id="DatasetAny", - ), - ], -) -def test_datasets_expression_error(expression: Callable[[], None], error: str) -> None: - with pytest.raises(TypeError) as info: - expression() - assert str(info.value) == error - - -def test_get_normalized_scheme(): - assert _get_normalized_scheme("http://example.com") == "http" - assert _get_normalized_scheme("HTTPS://example.com") == "https" - assert _get_normalized_scheme("ftp://example.com") == "ftp" - assert _get_normalized_scheme("file://") == "file" - - assert _get_normalized_scheme("example.com") == "" - assert _get_normalized_scheme("") == "" - assert _get_normalized_scheme(" ") == "" - - -def _mock_get_uri_normalizer_raising_error(normalized_scheme): - def normalizer(uri): - raise ValueError("Incorrect URI format") - - return normalizer - - -def _mock_get_uri_normalizer_noop(normalized_scheme): - def normalizer(uri): - return uri - - return normalizer - - -@patch("airflow.datasets._get_uri_normalizer", _mock_get_uri_normalizer_raising_error) -@patch("airflow.datasets.warnings.warn") -def test_sanitize_uri_raises_warning(mock_warn): - _sanitize_uri("postgres://localhost:5432/database.schema.table") - msg = mock_warn.call_args.args[0] - assert "The dataset URI postgres://localhost:5432/database.schema.table is not AIP-60 compliant" in msg - assert "In Airflow 3, this will raise an exception." in msg - - -@patch("airflow.datasets._get_uri_normalizer", _mock_get_uri_normalizer_raising_error) -@conf_vars({("core", "strict_dataset_uri_validation"): "True"}) -def test_sanitize_uri_raises_exception(): - with pytest.raises(ValueError) as e_info: - _sanitize_uri("postgres://localhost:5432/database.schema.table") - assert isinstance(e_info.value, ValueError) - assert str(e_info.value) == "Incorrect URI format" - - -@patch("airflow.datasets._get_uri_normalizer", lambda x: None) -def test_normalize_uri_no_normalizer_found(): - dataset = Dataset(uri="any_uri_without_normalizer_defined") - assert dataset.normalized_uri is None - - -@patch("airflow.datasets._get_uri_normalizer", _mock_get_uri_normalizer_raising_error) -def test_normalize_uri_invalid_uri(): - dataset = Dataset(uri="any_uri_not_aip60_compliant") - assert dataset.normalized_uri is None - - -@patch("airflow.datasets._get_uri_normalizer", _mock_get_uri_normalizer_noop) -@patch("airflow.datasets._get_normalized_scheme", lambda x: "valid_scheme") -def test_normalize_uri_valid_uri(): - dataset = Dataset(uri="valid_aip60_uri") - assert dataset.normalized_uri == "valid_aip60_uri" - - -@pytest.mark.skip_if_database_isolation_mode -@pytest.mark.db_test -@pytest.mark.usefixtures("clear_datasets") -class Test_DatasetAliasCondition: - @pytest.fixture - def ds_1(self, session): - """Example dataset links to dataset alias resolved_dsa_2.""" - ds_uri = "test_uri" - ds_1 = DatasetModel(id=1, uri=ds_uri) - - session.add(ds_1) - session.commit() - - return ds_1 - - @pytest.fixture - def dsa_1(self, session): - """Example dataset alias links to no datasets.""" - dsa_name = "test_name" - dsa_1 = DatasetAliasModel(name=dsa_name) - - session.add(dsa_1) - session.commit() - - return dsa_1 - - @pytest.fixture - def resolved_dsa_2(self, session, ds_1): - """Example dataset alias links to no dataset dsa_1.""" - dsa_name = "test_name_2" - dsa_2 = DatasetAliasModel(name=dsa_name) - dsa_2.datasets.append(ds_1) - - session.add(dsa_2) - session.commit() - - return dsa_2 - - def test_init(self, dsa_1, ds_1, resolved_dsa_2): - cond = _DatasetAliasCondition(name=dsa_1.name) - assert cond.objects == [] - - cond = _DatasetAliasCondition(name=resolved_dsa_2.name) - assert cond.objects == [Dataset(uri=ds_1.uri)] - - def test_as_expression(self, dsa_1, resolved_dsa_2): - for dsa in (dsa_1, resolved_dsa_2): - cond = _DatasetAliasCondition(dsa.name) - assert cond.as_expression() == {"alias": dsa.name} - - def test_evalute(self, dsa_1, resolved_dsa_2, ds_1): - cond = _DatasetAliasCondition(dsa_1.name) - assert cond.evaluate({ds_1.uri: True}) is False - - cond = _DatasetAliasCondition(resolved_dsa_2.name) - assert cond.evaluate({ds_1.uri: True}) is True diff --git a/tests/decorators/test_python.py b/tests/decorators/test_python.py index 96473518cc982..adbf96a0f41ba 100644 --- a/tests/decorators/test_python.py +++ b/tests/decorators/test_python.py @@ -983,8 +983,8 @@ def other(x): ... @pytest.mark.skip_if_database_isolation_mode # Test is broken in db isolation mode -def test_task_decorator_dataset(dag_maker, session): - from airflow.datasets import Dataset +def test_task_decorator_asset(dag_maker, session): + from airflow.assets import Asset result = None uri = "s3://bucket/name" @@ -992,11 +992,11 @@ def test_task_decorator_dataset(dag_maker, session): with dag_maker(session=session) as dag: @dag.task() - def up1() -> Dataset: - return Dataset(uri) + def up1() -> Asset: + return Asset(uri) @dag.task() - def up2(src: Dataset) -> str: + def up2(src: Asset) -> str: return src.uri @dag.task() diff --git a/tests/io/test_path.py b/tests/io/test_path.py index 195c2423b1822..0e504b586b3d2 100644 --- a/tests/io/test_path.py +++ b/tests/io/test_path.py @@ -29,7 +29,7 @@ from fsspec.implementations.memory import MemoryFileSystem from fsspec.registry import _registry as _fsspec_registry, register_implementation -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.io import _register_filesystems, get_fs from airflow.io.path import ObjectStoragePath from airflow.io.store import _STORE_CACHE, ObjectStore, attach @@ -280,12 +280,12 @@ def test_move_local(self, hook_lineage_collector): _to.unlink() - collected_datasets = hook_lineage_collector.collected_datasets + collected_assets = hook_lineage_collector.collected_assets - assert len(collected_datasets.inputs) == 1 - assert len(collected_datasets.outputs) == 1 - assert collected_datasets.inputs[0].dataset == Dataset(uri=_from_path) - assert collected_datasets.outputs[0].dataset == Dataset(uri=_to_path) + assert len(collected_assets.inputs) == 1 + assert len(collected_assets.outputs) == 1 + assert collected_assets.inputs[0].asset == Asset(uri=_from_path) + assert collected_assets.outputs[0].asset == Asset(uri=_to_path) def test_move_remote(self, hook_lineage_collector): attach("fakefs", fs=FakeRemoteFileSystem()) @@ -303,12 +303,12 @@ def test_move_remote(self, hook_lineage_collector): _to.unlink() - collected_datasets = hook_lineage_collector.collected_datasets + collected_assets = hook_lineage_collector.collected_assets - assert len(collected_datasets.inputs) == 1 - assert len(collected_datasets.outputs) == 1 - assert collected_datasets.inputs[0].dataset == Dataset(uri=str(_from)) - assert collected_datasets.outputs[0].dataset == Dataset(uri=str(_to)) + assert len(collected_assets.inputs) == 1 + assert len(collected_assets.outputs) == 1 + assert collected_assets.inputs[0].asset == Asset(uri=str(_from)) + assert collected_assets.outputs[0].asset == Asset(uri=str(_to)) def test_copy_remote_remote(self, hook_lineage_collector): attach("ffs", fs=FakeRemoteFileSystem(skip_instance_cache=True)) @@ -338,11 +338,11 @@ def test_copy_remote_remote(self, hook_lineage_collector): _from.rmdir(recursive=True) _to.rmdir(recursive=True) - assert len(hook_lineage_collector.collected_datasets.inputs) == 1 - assert hook_lineage_collector.collected_datasets.inputs[0].dataset == Dataset(uri=str(_from_file)) + assert len(hook_lineage_collector.collected_assets.inputs) == 1 + assert hook_lineage_collector.collected_assets.inputs[0].asset == Asset(uri=str(_from_file)) # Empty file - shutil.copyfileobj does nothing - assert len(hook_lineage_collector.collected_datasets.outputs) == 0 + assert len(hook_lineage_collector.collected_assets.outputs) == 0 def test_serde_objectstoragepath(self): path = "file:///bucket/key/part1/part2" @@ -402,12 +402,12 @@ def test_backwards_compat(self): # Reset the cache to avoid side effects _register_filesystems.cache_clear() - def test_dataset(self): + def test_asset(self): attach("s3", fs=FakeRemoteFileSystem()) p = "s3" f = "/tmp/foo" - i = Dataset(uri=f"{p}://{f}", extra={"foo": "bar"}) + i = Asset(uri=f"{p}://{f}", extra={"foo": "bar"}) o = ObjectStoragePath(i) assert o.protocol == p assert o.path == f diff --git a/tests/io/test_wrapper.py b/tests/io/test_wrapper.py index e00c5ab22bf64..641eda84d1a4f 100644 --- a/tests/io/test_wrapper.py +++ b/tests/io/test_wrapper.py @@ -19,13 +19,13 @@ import uuid from unittest.mock import patch -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.io.path import ObjectStoragePath @patch("airflow.providers_manager.ProvidersManager") def test_wrapper_catches_reads_writes(providers_manager, hook_lineage_collector): - providers_manager.return_value._dataset_factories = lambda x: Dataset(uri=x) + providers_manager.return_value._asset_factories = lambda x: Asset(uri=x) uri = f"file:///tmp/{str(uuid.uuid4())}" path = ObjectStoragePath(uri) file = path.open("w") @@ -33,7 +33,7 @@ def test_wrapper_catches_reads_writes(providers_manager, hook_lineage_collector) file.close() assert len(hook_lineage_collector._outputs) == 1 - assert next(iter(hook_lineage_collector._outputs.values()))[0] == Dataset(uri=uri) + assert next(iter(hook_lineage_collector._outputs.values()))[0] == Asset(uri=uri) file = path.open("r") file.read() @@ -42,23 +42,23 @@ def test_wrapper_catches_reads_writes(providers_manager, hook_lineage_collector) path.unlink(missing_ok=True) assert len(hook_lineage_collector._inputs) == 1 - assert next(iter(hook_lineage_collector._inputs.values()))[0] == Dataset(uri=uri) + assert next(iter(hook_lineage_collector._inputs.values()))[0] == Asset(uri=uri) @patch("airflow.providers_manager.ProvidersManager") def test_wrapper_works_with_contextmanager(providers_manager, hook_lineage_collector): - providers_manager.return_value._dataset_factories = lambda x: Dataset(uri=x) + providers_manager.return_value._asset_factories = lambda x: Asset(uri=x) uri = f"file:///tmp/{str(uuid.uuid4())}" path = ObjectStoragePath(uri) with path.open("w") as file: file.write("asdf") assert len(hook_lineage_collector._outputs) == 1 - assert next(iter(hook_lineage_collector._outputs.values()))[0] == Dataset(uri=uri) + assert next(iter(hook_lineage_collector._outputs.values()))[0] == Asset(uri=uri) with path.open("r") as file: file.read() path.unlink(missing_ok=True) assert len(hook_lineage_collector._inputs) == 1 - assert next(iter(hook_lineage_collector._inputs.values()))[0] == Dataset(uri=uri) + assert next(iter(hook_lineage_collector._inputs.values()))[0] == Asset(uri=uri) diff --git a/tests/jobs/test_scheduler_job.py b/tests/jobs/test_scheduler_job.py index 78a911153dab2..32662d7d873db 100644 --- a/tests/jobs/test_scheduler_job.py +++ b/tests/jobs/test_scheduler_job.py @@ -36,12 +36,12 @@ import airflow.example_dags from airflow import settings +from airflow.assets import Asset +from airflow.assets.manager import AssetManager from airflow.callbacks.callback_requests import DagCallbackRequest, TaskCallbackRequest from airflow.callbacks.database_callback_sink import DatabaseCallbackSink from airflow.callbacks.pipe_callback_sink import PipeCallbackSink from airflow.dag_processing.manager import DagFileProcessorAgent -from airflow.datasets import Dataset -from airflow.datasets.manager import DatasetManager from airflow.exceptions import AirflowException, RemovedInAirflow3Warning from airflow.executors.base_executor import BaseExecutor from airflow.executors.executor_constants import MOCK_EXECUTOR @@ -50,10 +50,10 @@ from airflow.jobs.job import Job, run_job from airflow.jobs.local_task_job_runner import LocalTaskJobRunner from airflow.jobs.scheduler_job_runner import SchedulerJobRunner +from airflow.models.asset import AssetDagRunQueue, AssetEvent, AssetModel from airflow.models.dag import DAG, DagModel from airflow.models.dagbag import DagBag from airflow.models.dagrun import DagRun -from airflow.models.dataset import DatasetDagRunQueue, DatasetEvent, DatasetModel from airflow.models.db_callback_request import DbCallbackRequest from airflow.models.pool import Pool from airflow.models.serialized_dag import SerializedDagModel @@ -74,8 +74,8 @@ from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS from tests.test_utils.config import conf_vars, env_vars from tests.test_utils.db import ( + clear_db_assets, clear_db_dags, - clear_db_datasets, clear_db_import_errors, clear_db_jobs, clear_db_pools, @@ -141,7 +141,7 @@ def clean_db(): clear_db_sla_miss() clear_db_import_errors() clear_db_jobs() - clear_db_datasets() + clear_db_assets() # DO NOT try to run clear_db_serialized_dags() here - this will break the tests # The tests expect DAGs to be fully loaded here via setUpClass method below @@ -4094,7 +4094,7 @@ def test_create_dag_runs(self, dag_maker): assert dag.get_last_dagrun().creating_job_id == scheduler_job.id @pytest.mark.need_serialized_dag - def test_create_dag_runs_datasets(self, session, dag_maker): + def test_create_dag_runs_assets(self, session, dag_maker): """ Test various invariants of _create_dag_runs. @@ -4103,21 +4103,21 @@ def test_create_dag_runs_datasets(self, session, dag_maker): - That dag_model has next_dagrun """ - dataset1 = Dataset(uri="ds1") - dataset2 = Dataset(uri="ds2") + asset1 = Asset(uri="ds1") + asset2 = Asset(uri="ds2") - with dag_maker(dag_id="datasets-1", start_date=timezone.utcnow(), session=session): - BashOperator(task_id="task", bash_command="echo 1", outlets=[dataset1]) + with dag_maker(dag_id="assets-1", start_date=timezone.utcnow(), session=session): + BashOperator(task_id="task", bash_command="echo 1", outlets=[asset1]) dr = dag_maker.create_dagrun( run_id="run1", execution_date=(DEFAULT_DATE + timedelta(days=100)), data_interval=(DEFAULT_DATE + timedelta(days=10), DEFAULT_DATE + timedelta(days=11)), ) - ds1_id = session.query(DatasetModel.id).filter_by(uri=dataset1.uri).scalar() + asset1_id = session.query(AssetModel.id).filter_by(uri=asset1.uri).scalar() - event1 = DatasetEvent( - dataset_id=ds1_id, + event1 = AssetEvent( + dataset_id=asset1_id, source_task_id="task", source_dag_id=dr.dag_id, source_run_id=dr.run_id, @@ -4132,8 +4132,8 @@ def test_create_dag_runs_datasets(self, session, dag_maker): data_interval=(DEFAULT_DATE + timedelta(days=5), DEFAULT_DATE + timedelta(days=6)), ) - event2 = DatasetEvent( - dataset_id=ds1_id, + event2 = AssetEvent( + dataset_id=asset1_id, source_task_id="task", source_dag_id=dr.dag_id, source_run_id=dr.run_id, @@ -4141,18 +4141,18 @@ def test_create_dag_runs_datasets(self, session, dag_maker): ) session.add(event2) - with dag_maker(dag_id="datasets-consumer-multiple", schedule=[dataset1, dataset2]): + with dag_maker(dag_id="assets-consumer-multiple", schedule=[asset1, asset2]): pass dag2 = dag_maker.dag - with dag_maker(dag_id="datasets-consumer-single", schedule=[dataset1]): + with dag_maker(dag_id="assets-consumer-single", schedule=[asset1]): pass dag3 = dag_maker.dag session = dag_maker.session session.add_all( [ - DatasetDagRunQueue(dataset_id=ds1_id, target_dag_id=dag2.dag_id), - DatasetDagRunQueue(dataset_id=ds1_id, target_dag_id=dag3.dag_id), + AssetDagRunQueue(dataset_id=asset1_id, target_dag_id=dag2.dag_id), + AssetDagRunQueue(dataset_id=asset1_id, target_dag_id=dag3.dag_id), ] ) session.flush() @@ -4169,24 +4169,24 @@ def dict_from_obj(obj): """Get dict of column attrs from SqlAlchemy object.""" return {k.key: obj.__dict__.get(k) for k in obj.__mapper__.column_attrs} - # dag3 should be triggered since it only depends on dataset1, and it's been queued + # dag3 should be triggered since it only depends on asset1, and it's been queued created_run = session.query(DagRun).filter(DagRun.dag_id == dag3.dag_id).one() assert created_run.state == State.QUEUED assert created_run.start_date is None - # we don't have __eq__ defined on DatasetEvent because... given the fact that in the future - # we may register events from other systems, dataset_id + timestamp might not be enough PK + # we don't have __eq__ defined on AssetEvent because... given the fact that in the future + # we may register events from other systems, asset_id + timestamp might not be enough PK assert list(map(dict_from_obj, created_run.consumed_dataset_events)) == list( map(dict_from_obj, [event1, event2]) ) assert created_run.data_interval_start == DEFAULT_DATE + timedelta(days=5) assert created_run.data_interval_end == DEFAULT_DATE + timedelta(days=11) - # dag2 DDRQ record should still be there since the dag run was *not* triggered - assert session.query(DatasetDagRunQueue).filter_by(target_dag_id=dag2.dag_id).one() is not None - # dag2 should not be triggered since it depends on both dataset 1 and 2 + # dag2 ADRQ record should still be there since the dag run was *not* triggered + assert session.query(AssetDagRunQueue).filter_by(target_dag_id=dag2.dag_id).one() is not None + # dag2 should not be triggered since it depends on both asset 1 and 2 assert session.query(DagRun).filter(DagRun.dag_id == dag2.dag_id).one_or_none() is None - # dag3 DDRQ record should be deleted since the dag run was triggered - assert session.query(DatasetDagRunQueue).filter_by(target_dag_id=dag3.dag_id).one_or_none() is None + # dag3 ADRQ record should be deleted since the dag run was triggered + assert session.query(AssetDagRunQueue).filter_by(target_dag_id=dag3.dag_id).one_or_none() is None assert dag3.get_last_dagrun().creating_job_id == scheduler_job.id @@ -4199,47 +4199,47 @@ def dict_from_obj(obj): ], ) def test_no_create_dag_runs_when_dag_disabled(self, session, dag_maker, disable, enable): - ds = Dataset("ds") + ds = Asset("ds") with dag_maker(dag_id="consumer", schedule=[ds], session=session): pass with dag_maker(dag_id="producer", schedule="@daily", session=session): BashOperator(task_id="task", bash_command="echo 1", outlets=ds) - dsm = DatasetManager() + asset_manger = AssetManager() - ds_id = session.scalars(select(DatasetModel.id).filter_by(uri=ds.uri)).one() + asset_id = session.scalars(select(AssetModel.id).filter_by(uri=ds.uri)).one() - dse_q = select(DatasetEvent).where(DatasetEvent.dataset_id == ds_id).order_by(DatasetEvent.timestamp) - ddrq_q = select(DatasetDagRunQueue).where( - DatasetDagRunQueue.dataset_id == ds_id, DatasetDagRunQueue.target_dag_id == "consumer" + ase_q = select(AssetEvent).where(AssetEvent.dataset_id == asset_id).order_by(AssetEvent.timestamp) + adrq_q = select(AssetDagRunQueue).where( + AssetDagRunQueue.dataset_id == asset_id, AssetDagRunQueue.target_dag_id == "consumer" ) # Simulate the consumer DAG being disabled. session.execute(update(DagModel).where(DagModel.dag_id == "consumer").values(**disable)) - # A DDRQ is not scheduled although an event is emitted. + # An ADRQ is not scheduled although an event is emitted. dr1: DagRun = dag_maker.create_dagrun(run_type=DagRunType.SCHEDULED) - dsm.register_dataset_change( + asset_manger.register_asset_change( task_instance=dr1.get_task_instance("task", session=session), - dataset=ds, + asset=ds, session=session, ) session.flush() - assert session.scalars(dse_q).one().source_run_id == dr1.run_id - assert session.scalars(ddrq_q).one_or_none() is None + assert session.scalars(ase_q).one().source_run_id == dr1.run_id + assert session.scalars(adrq_q).one_or_none() is None # Simulate the consumer DAG being enabled. session.execute(update(DagModel).where(DagModel.dag_id == "consumer").values(**enable)) - # A DDRQ should be scheduled for the new event, but not the previous one. + # An ADRQ should be scheduled for the new event, but not the previous one. dr2: DagRun = dag_maker.create_dagrun_after(dr1, run_type=DagRunType.SCHEDULED) - dsm.register_dataset_change( + asset_manger.register_asset_change( task_instance=dr2.get_task_instance("task", session=session), - dataset=ds, + asset=ds, session=session, ) session.flush() - assert [e.source_run_id for e in session.scalars(dse_q)] == [dr1.run_id, dr2.run_id] - assert session.scalars(ddrq_q).one().target_dag_id == "consumer" + assert [e.source_run_id for e in session.scalars(ase_q)] == [dr1.run_id, dr2.run_id] + assert session.scalars(adrq_q).one().target_dag_id == "consumer" @time_machine.travel(DEFAULT_DATE + datetime.timedelta(days=1, seconds=9), tick=False) @mock.patch("airflow.jobs.scheduler_job_runner.Stats.timing") @@ -5728,87 +5728,85 @@ def test_update_dagrun_state_for_paused_dag_not_for_backfill(self, dag_maker, se (backfill_run,) = DagRun.find(dag_id=dag.dag_id, run_type=DagRunType.BACKFILL_JOB, session=session) assert backfill_run.state == State.RUNNING - def test_dataset_orphaning(self, dag_maker, session): - dataset1 = Dataset(uri="ds1") - dataset2 = Dataset(uri="ds2") - dataset3 = Dataset(uri="ds3") - dataset4 = Dataset(uri="ds4") + def test_asset_orphaning(self, dag_maker, session): + asset1 = Asset(uri="ds1") + asset2 = Asset(uri="ds2") + asset3 = Asset(uri="ds3") + asset4 = Asset(uri="ds4") - with dag_maker(dag_id="datasets-1", schedule=[dataset1, dataset2], session=session): - BashOperator(task_id="task", bash_command="echo 1", outlets=[dataset3, dataset4]) + with dag_maker(dag_id="assets-1", schedule=[asset1, asset2], session=session): + BashOperator(task_id="task", bash_command="echo 1", outlets=[asset3, asset4]) - non_orphaned_dataset_count = session.query(DatasetModel).filter(~DatasetModel.is_orphaned).count() - assert non_orphaned_dataset_count == 4 - orphaned_dataset_count = session.query(DatasetModel).filter(DatasetModel.is_orphaned).count() - assert orphaned_dataset_count == 0 + non_orphaned_asset_count = session.query(AssetModel).filter(~AssetModel.is_orphaned).count() + assert non_orphaned_asset_count == 4 + orphaned_asset_count = session.query(AssetModel).filter(AssetModel.is_orphaned).count() + assert orphaned_asset_count == 0 - # now remove 2 dataset references - with dag_maker(dag_id="datasets-1", schedule=[dataset1], session=session): - BashOperator(task_id="task", bash_command="echo 1", outlets=[dataset3]) + # now remove 2 asset references + with dag_maker(dag_id="assets-1", schedule=[asset1], session=session): + BashOperator(task_id="task", bash_command="echo 1", outlets=[asset3]) scheduler_job = Job() self.job_runner = SchedulerJobRunner(job=scheduler_job, subdir=os.devnull) - self.job_runner._orphan_unreferenced_datasets(session=session) + self.job_runner._orphan_unreferenced_assets(session=session) session.flush() # and find the orphans - non_orphaned_datasets = [ - dataset.uri - for dataset in session.query(DatasetModel.uri) - .filter(~DatasetModel.is_orphaned) - .order_by(DatasetModel.uri) + non_orphaned_assets = [ + asset.uri + for asset in session.query(AssetModel.uri) + .filter(~AssetModel.is_orphaned) + .order_by(AssetModel.uri) ] - assert non_orphaned_datasets == ["ds1", "ds3"] - orphaned_datasets = [ - dataset.uri - for dataset in session.query(DatasetModel.uri) - .filter(DatasetModel.is_orphaned) - .order_by(DatasetModel.uri) + assert non_orphaned_assets == ["ds1", "ds3"] + orphaned_assets = [ + asset.uri + for asset in session.query(AssetModel.uri).filter(AssetModel.is_orphaned).order_by(AssetModel.uri) ] - assert orphaned_datasets == ["ds2", "ds4"] + assert orphaned_assets == ["ds2", "ds4"] - def test_dataset_orphaning_ignore_orphaned_datasets(self, dag_maker, session): - dataset1 = Dataset(uri="ds1") + def test_asset_orphaning_ignore_orphaned_assets(self, dag_maker, session): + asset1 = Asset(uri="ds1") - with dag_maker(dag_id="datasets-1", schedule=[dataset1], session=session): + with dag_maker(dag_id="assets-1", schedule=[asset1], session=session): BashOperator(task_id="task", bash_command="echo 1") - non_orphaned_dataset_count = session.query(DatasetModel).filter(~DatasetModel.is_orphaned).count() - assert non_orphaned_dataset_count == 1 - orphaned_dataset_count = session.query(DatasetModel).filter(DatasetModel.is_orphaned).count() - assert orphaned_dataset_count == 0 + non_orphaned_asset_count = session.query(AssetModel).filter(~AssetModel.is_orphaned).count() + assert non_orphaned_asset_count == 1 + orphaned_asset_count = session.query(AssetModel).filter(AssetModel.is_orphaned).count() + assert orphaned_asset_count == 0 - # now remove dataset1 reference - with dag_maker(dag_id="datasets-1", schedule=None, session=session): + # now remove asset1 reference + with dag_maker(dag_id="assets-1", schedule=None, session=session): BashOperator(task_id="task", bash_command="echo 1") scheduler_job = Job() self.job_runner = SchedulerJobRunner(job=scheduler_job, subdir=os.devnull) - self.job_runner._orphan_unreferenced_datasets(session=session) + self.job_runner._orphan_unreferenced_assets(session=session) session.flush() - orphaned_datasets_before_rerun = ( - session.query(DatasetModel.updated_at, DatasetModel.uri) - .filter(DatasetModel.is_orphaned) - .order_by(DatasetModel.uri) + orphaned_assets_before_rerun = ( + session.query(AssetModel.updated_at, AssetModel.uri) + .filter(AssetModel.is_orphaned) + .order_by(AssetModel.uri) ) - assert [dataset.uri for dataset in orphaned_datasets_before_rerun] == ["ds1"] - updated_at_timestamps = [dataset.updated_at for dataset in orphaned_datasets_before_rerun] + assert [asset.uri for asset in orphaned_assets_before_rerun] == ["ds1"] + updated_at_timestamps = [asset.updated_at for asset in orphaned_assets_before_rerun] - # when rerunning we should ignore the already orphaned datasets and thus the updated_at timestamp + # when rerunning we should ignore the already orphaned assets and thus the updated_at timestamp # should remain the same - self.job_runner._orphan_unreferenced_datasets(session=session) + self.job_runner._orphan_unreferenced_assets(session=session) session.flush() - orphaned_datasets_after_rerun = ( - session.query(DatasetModel.updated_at, DatasetModel.uri) - .filter(DatasetModel.is_orphaned) - .order_by(DatasetModel.uri) + orphaned_assets_after_rerun = ( + session.query(AssetModel.updated_at, AssetModel.uri) + .filter(AssetModel.is_orphaned) + .order_by(AssetModel.uri) ) - assert [dataset.uri for dataset in orphaned_datasets_after_rerun] == ["ds1"] - assert updated_at_timestamps == [dataset.updated_at for dataset in orphaned_datasets_after_rerun] + assert [asset.uri for asset in orphaned_assets_after_rerun] == ["ds1"] + assert updated_at_timestamps == [asset.updated_at for asset in orphaned_assets_after_rerun] def test_misconfigured_dags_doesnt_crash_scheduler(self, session, dag_maker, caplog): """Test that if dagrun creation throws an exception, the scheduler doesn't crash""" diff --git a/tests/lineage/test_hook.py b/tests/lineage/test_hook.py index 67059f91b4da1..c076b19aecedd 100644 --- a/tests/lineage/test_hook.py +++ b/tests/lineage/test_hook.py @@ -22,11 +22,11 @@ import pytest from airflow import plugins_manager -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.hooks.base import BaseHook from airflow.lineage import hook from airflow.lineage.hook import ( - DatasetLineageInfo, + AssetLineageInfo, HookLineage, HookLineageCollector, HookLineageReader, @@ -40,156 +40,124 @@ class TestHookLineageCollector: def setup_method(self): self.collector = HookLineageCollector() - def test_are_datasets_collected(self): + def test_are_assets_collected(self): assert self.collector is not None - assert self.collector.collected_datasets == HookLineage() + assert self.collector.collected_assets == HookLineage() input_hook = BaseHook() output_hook = BaseHook() - self.collector.add_input_dataset(input_hook, uri="s3://in_bucket/file") - self.collector.add_output_dataset( - output_hook, uri="postgres://example.com:5432/database/default/table" - ) - assert self.collector.collected_datasets == HookLineage( - [DatasetLineageInfo(dataset=Dataset("s3://in_bucket/file"), count=1, context=input_hook)], + self.collector.add_input_asset(input_hook, uri="s3://in_bucket/file") + self.collector.add_output_asset(output_hook, uri="postgres://example.com:5432/database/default/table") + assert self.collector.collected_assets == HookLineage( + [AssetLineageInfo(asset=Asset("s3://in_bucket/file"), count=1, context=input_hook)], [ - DatasetLineageInfo( - dataset=Dataset("postgres://example.com:5432/database/default/table"), + AssetLineageInfo( + asset=Asset("postgres://example.com:5432/database/default/table"), count=1, context=output_hook, ) ], ) - @patch("airflow.lineage.hook.Dataset") - def test_add_input_dataset(self, mock_dataset): - dataset = MagicMock(spec=Dataset, extra={}) - mock_dataset.return_value = dataset + @patch("airflow.lineage.hook.Asset") + def test_add_input_asset(self, mock_asset): + asset = MagicMock(spec=Asset, extra={}) + mock_asset.return_value = asset hook = MagicMock() - self.collector.add_input_dataset(hook, uri="test_uri") + self.collector.add_input_asset(hook, uri="test_uri") - assert next(iter(self.collector._inputs.values())) == (dataset, hook) - mock_dataset.assert_called_once_with(uri="test_uri", extra=None) + assert next(iter(self.collector._inputs.values())) == (asset, hook) + mock_asset.assert_called_once_with(uri="test_uri", extra=None) - def test_grouping_datasets(self): + def test_grouping_assets(self): hook_1 = MagicMock() hook_2 = MagicMock() uri = "test://uri/" - self.collector.add_input_dataset(context=hook_1, uri=uri) - self.collector.add_input_dataset(context=hook_2, uri=uri) - self.collector.add_input_dataset(context=hook_1, uri=uri, dataset_extra={"key": "value"}) + self.collector.add_input_asset(context=hook_1, uri=uri) + self.collector.add_input_asset(context=hook_2, uri=uri) + self.collector.add_input_asset(context=hook_1, uri=uri, asset_extra={"key": "value"}) - collected_inputs = self.collector.collected_datasets.inputs + collected_inputs = self.collector.collected_assets.inputs assert len(collected_inputs) == 3 - assert collected_inputs[0].dataset.uri == "test://uri/" - assert collected_inputs[0].dataset == collected_inputs[1].dataset + assert collected_inputs[0].asset.uri == "test://uri/" + assert collected_inputs[0].asset == collected_inputs[1].asset assert collected_inputs[0].count == 1 assert collected_inputs[0].context == collected_inputs[2].context == hook_1 assert collected_inputs[1].count == 1 assert collected_inputs[1].context == hook_2 assert collected_inputs[2].count == 1 - assert collected_inputs[2].dataset.extra == {"key": "value"} + assert collected_inputs[2].asset.extra == {"key": "value"} @patch("airflow.lineage.hook.ProvidersManager") - def test_create_dataset(self, mock_providers_manager): - def create_dataset(arg1, arg2="default", extra=None): - return Dataset(uri=f"myscheme://{arg1}/{arg2}", extra=extra) - - test_scheme = "myscheme" - mock_providers_manager.return_value.dataset_factories = {test_scheme: create_dataset} - - test_uri = "urischeme://value_a/value_b" - test_kwargs = {"arg1": "value_1"} - test_kwargs_uri = "myscheme://value_1/default" - test_extra = {"key": "value"} - - # test uri arg - should take precedence over the keyword args + scheme - assert self.collector.create_dataset( - scheme=test_scheme, uri=test_uri, dataset_kwargs=test_kwargs, dataset_extra=None - ) == Dataset(test_uri) - assert self.collector.create_dataset( - scheme=test_scheme, uri=test_uri, dataset_kwargs=test_kwargs, dataset_extra={} - ) == Dataset(test_uri) - assert self.collector.create_dataset( - scheme=test_scheme, uri=test_uri, dataset_kwargs=test_kwargs, dataset_extra=test_extra - ) == Dataset(test_uri, extra=test_extra) - - # test keyword args - assert self.collector.create_dataset( - scheme=test_scheme, uri=None, dataset_kwargs=test_kwargs, dataset_extra=None - ) == Dataset(test_kwargs_uri) - assert self.collector.create_dataset( - scheme=test_scheme, uri=None, dataset_kwargs=test_kwargs, dataset_extra={} - ) == Dataset(test_kwargs_uri) - assert self.collector.create_dataset( - scheme=test_scheme, + def test_create_asset(self, mock_providers_manager): + def create_asset(arg1, arg2="default", extra=None): + return Asset(uri=f"myscheme://{arg1}/{arg2}", extra=extra or {}) + + mock_providers_manager.return_value.asset_factories = {"myscheme": create_asset} + assert self.collector.create_asset( + scheme="myscheme", uri=None, asset_kwargs={"arg1": "value_1"}, asset_extra=None + ) == Asset("myscheme://value_1/default") + assert self.collector.create_asset( + scheme="myscheme", uri=None, - dataset_kwargs={**test_kwargs, "arg2": "value_2"}, - dataset_extra=test_extra, - ) == Dataset("myscheme://value_1/value_2", extra=test_extra) - - # missing both uri and scheme - assert ( - self.collector.create_dataset( - scheme=None, uri=None, dataset_kwargs=test_kwargs, dataset_extra=None - ) - is None - ) + asset_kwargs={"arg1": "value_1", "arg2": "value_2"}, + asset_extra={"key": "value"}, + ) == Asset("myscheme://value_1/value_2", extra={"key": "value"}) @patch("airflow.lineage.hook.ProvidersManager") - def test_create_dataset_no_factory(self, mock_providers_manager): + def test_create_asset_no_factory(self, mock_providers_manager): test_scheme = "myscheme" - mock_providers_manager.return_value.dataset_factories = {} + mock_providers_manager.return_value.asset_factories = {} test_kwargs = {"arg1": "value_1"} assert ( - self.collector.create_dataset( - scheme=test_scheme, uri=None, dataset_kwargs=test_kwargs, dataset_extra=None + self.collector.create_asset( + scheme=test_scheme, uri=None, asset_kwargs=test_kwargs, asset_extra=None ) is None ) @patch("airflow.lineage.hook.ProvidersManager") - def test_create_dataset_factory_exception(self, mock_providers_manager): - def create_dataset(extra=None, **kwargs): + def test_create_asset_factory_exception(self, mock_providers_manager): + def create_asset(extra=None, **kwargs): raise RuntimeError("Factory error") test_scheme = "myscheme" - mock_providers_manager.return_value.dataset_factories = {test_scheme: create_dataset} + mock_providers_manager.return_value.asset_factories = {test_scheme: create_asset} test_kwargs = {"arg1": "value_1"} assert ( - self.collector.create_dataset( - scheme=test_scheme, uri=None, dataset_kwargs=test_kwargs, dataset_extra=None + self.collector.create_asset( + scheme=test_scheme, uri=None, asset_kwargs=test_kwargs, asset_extra=None ) is None ) - def test_collected_datasets(self): + def test_collected_assets(self): context_input = MagicMock() context_output = MagicMock() - self.collector.add_input_dataset(context_input, uri="test://input") - self.collector.add_output_dataset(context_output, uri="test://output") + self.collector.add_input_asset(context_input, uri="test://input") + self.collector.add_output_asset(context_output, uri="test://output") - hook_lineage = self.collector.collected_datasets + hook_lineage = self.collector.collected_assets assert len(hook_lineage.inputs) == 1 - assert hook_lineage.inputs[0].dataset.uri == "test://input/" + assert hook_lineage.inputs[0].asset.uri == "test://input/" assert hook_lineage.inputs[0].context == context_input assert len(hook_lineage.outputs) == 1 - assert hook_lineage.outputs[0].dataset.uri == "test://output/" + assert hook_lineage.outputs[0].asset.uri == "test://output/" def test_has_collected(self): collector = HookLineageCollector() assert not collector.has_collected - collector._inputs = {"unique_key": (MagicMock(spec=Dataset), MagicMock())} + collector._inputs = {"unique_key": (MagicMock(spec=Asset), MagicMock())} assert collector.has_collected diff --git a/tests/listeners/dataset_listener.py b/tests/listeners/asset_listener.py similarity index 80% rename from tests/listeners/dataset_listener.py rename to tests/listeners/asset_listener.py index 0e4b768c696f1..e7adf580363b8 100644 --- a/tests/listeners/dataset_listener.py +++ b/tests/listeners/asset_listener.py @@ -23,21 +23,21 @@ from airflow.listeners import hookimpl if typing.TYPE_CHECKING: - from airflow.datasets import Dataset + from airflow.assets import Asset -changed: list[Dataset] = [] -created: list[Dataset] = [] +changed: list[Asset] = [] +created: list[Asset] = [] @hookimpl -def on_dataset_changed(dataset): - changed.append(copy.deepcopy(dataset)) +def on_asset_changed(asset): + changed.append(copy.deepcopy(asset)) @hookimpl -def on_dataset_created(dataset): - created.append(copy.deepcopy(dataset)) +def on_asset_created(asset): + created.append(copy.deepcopy(asset)) def clear(): diff --git a/tests/listeners/test_dataset_listener.py b/tests/listeners/test_asset_listener.py similarity index 72% rename from tests/listeners/test_dataset_listener.py rename to tests/listeners/test_asset_listener.py index b0ac6223e79ea..bb93acd8a0fff 100644 --- a/tests/listeners/test_dataset_listener.py +++ b/tests/listeners/test_asset_listener.py @@ -18,33 +18,33 @@ import pytest -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.listeners.listener import get_listener_manager -from airflow.models.dataset import DatasetModel +from airflow.models.asset import AssetModel from airflow.operators.empty import EmptyOperator from airflow.utils.session import provide_session -from tests.listeners import dataset_listener +from tests.listeners import asset_listener @pytest.fixture(autouse=True) def clean_listener_manager(): lm = get_listener_manager() lm.clear() - lm.add_listener(dataset_listener) + lm.add_listener(asset_listener) yield lm = get_listener_manager() lm.clear() - dataset_listener.clear() + asset_listener.clear() @pytest.mark.skip_if_database_isolation_mode # Test is broken in db isolation mode @pytest.mark.db_test @provide_session -def test_dataset_listener_on_dataset_changed_gets_calls(create_task_instance_of_operator, session): - dataset_uri = "test_dataset_uri" - ds = Dataset(uri=dataset_uri) - ds_model = DatasetModel(uri=dataset_uri) - session.add(ds_model) +def test_asset_listener_on_asset_changed_gets_calls(create_task_instance_of_operator, session): + asset_uri = "test_asset_uri" + asset = Asset(uri=asset_uri) + asset_model = AssetModel(uri=asset_uri) + session.add(asset_model) session.flush() @@ -53,9 +53,9 @@ def test_dataset_listener_on_dataset_changed_gets_calls(create_task_instance_of_ dag_id="producing_dag", task_id="test_task", session=session, - outlets=[ds], + outlets=[asset], ) ti.run() - assert len(dataset_listener.changed) == 1 - assert dataset_listener.changed[0].uri == dataset_uri + assert len(asset_listener.changed) == 1 + assert asset_listener.changed[0].uri == asset_uri diff --git a/tests/models/test_dataset.py b/tests/models/test_asset.py similarity index 73% rename from tests/models/test_dataset.py rename to tests/models/test_asset.py index f562e2347b008..5b35a0c89529e 100644 --- a/tests/models/test_dataset.py +++ b/tests/models/test_asset.py @@ -17,13 +17,13 @@ from __future__ import annotations -from airflow.datasets import DatasetAlias -from airflow.models.dataset import DatasetAliasModel +from airflow.assets import AssetAlias +from airflow.models.asset import AssetAliasModel -class TestDatasetAliasModel: +class TestAssetAliasModel: def test_from_public(self): - dataset_alias = DatasetAlias(name="test_alias") - dataset_alias_model = DatasetAliasModel.from_public(dataset_alias) + asset_alias = AssetAlias(name="test_alias") + asset_alias_model = AssetAliasModel.from_public(asset_alias) - assert dataset_alias_model.name == "test_alias" + assert asset_alias_model.name == "test_alias" diff --git a/tests/models/test_dag.py b/tests/models/test_dag.py index df4a892768816..ab67c3778c262 100644 --- a/tests/models/test_dag.py +++ b/tests/models/test_dag.py @@ -38,8 +38,8 @@ from sqlalchemy import inspect, select from airflow import settings +from airflow.assets import Asset, AssetAlias, AssetAll, AssetAny from airflow.configuration import conf -from airflow.datasets import Dataset, DatasetAlias, DatasetAll, DatasetAny from airflow.decorators import setup, task as task_decorator, teardown from airflow.exceptions import ( AirflowException, @@ -51,6 +51,13 @@ from airflow.executors import executor_loader from airflow.executors.local_executor import LocalExecutor from airflow.executors.sequential_executor import SequentialExecutor +from airflow.models.asset import ( + AssetAliasModel, + AssetDagRunQueue, + AssetEvent, + AssetModel, + TaskOutletAssetReference, +) from airflow.models.baseoperator import BaseOperator from airflow.models.dag import ( DAG, @@ -60,16 +67,9 @@ DagTag, ExecutorLoader, dag as dag_decorator, - get_dataset_triggered_next_run_info, + get_asset_triggered_next_run_info, ) from airflow.models.dagrun import DagRun -from airflow.models.dataset import ( - DatasetAliasModel, - DatasetDagRunQueue, - DatasetEvent, - DatasetModel, - TaskOutletDatasetReference, -) from airflow.models.param import DagParam, Param, ParamsDict from airflow.models.serialized_dag import SerializedDagModel from airflow.models.taskfail import TaskFail @@ -81,8 +81,8 @@ from airflow.templates import NativeEnvironment, SandboxedEnvironment from airflow.timetables.base import DagRunInfo, DataInterval, TimeRestriction, Timetable from airflow.timetables.simple import ( + AssetTriggeredTimetable, ContinuousTimetable, - DatasetTriggeredTimetable, NullTimetable, OnceTimetable, ) @@ -105,7 +105,7 @@ from tests.test_utils.asserts import assert_queries_count from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS from tests.test_utils.config import conf_vars -from tests.test_utils.db import clear_db_dags, clear_db_datasets, clear_db_runs, clear_db_serialized_dags +from tests.test_utils.db import clear_db_assets, clear_db_dags, clear_db_runs, clear_db_serialized_dags from tests.test_utils.mapping import expand_mapped_task from tests.test_utils.mock_plugins import mock_plugin_manager from tests.test_utils.timetables import cron_timetable, delta_timetable @@ -133,24 +133,24 @@ def clear_dags(): @pytest.fixture -def clear_datasets(): - clear_db_datasets() +def clear_assets(): + clear_db_assets() yield - clear_db_datasets() + clear_db_assets() class TestDag: def setup_method(self) -> None: clear_db_runs() clear_db_dags() - clear_db_datasets() + clear_db_assets() self.patcher_dag_code = mock.patch("airflow.models.dag.DagCode.bulk_sync_to_db") self.patcher_dag_code.start() def teardown_method(self) -> None: clear_db_runs() clear_db_dags() - clear_db_datasets() + clear_db_assets() self.patcher_dag_code.stop() @staticmethod @@ -1004,47 +1004,47 @@ def test_bulk_write_to_db_has_import_error(self): assert not model.has_import_errors session.close() - def test_bulk_write_to_db_datasets(self): + def test_bulk_write_to_db_assets(self): """ - Ensure that datasets referenced in a dag are correctly loaded into the database. + Ensure that assets referenced in a dag are correctly loaded into the database. """ - dag_id1 = "test_dataset_dag1" - dag_id2 = "test_dataset_dag2" - task_id = "test_dataset_task" - uri1 = "s3://dataset/1" - d1 = Dataset(uri1, extra={"not": "used"}) - d2 = Dataset("s3://dataset/2") - d3 = Dataset("s3://dataset/3") + dag_id1 = "test_asset_dag1" + dag_id2 = "test_asset_dag2" + task_id = "test_asset_task" + uri1 = "s3://asset/1" + d1 = Asset(uri1, extra={"not": "used"}) + d2 = Asset("s3://asset/2") + d3 = Asset("s3://asset/3") dag1 = DAG(dag_id=dag_id1, start_date=DEFAULT_DATE, schedule=[d1]) EmptyOperator(task_id=task_id, dag=dag1, outlets=[d2, d3]) dag2 = DAG(dag_id=dag_id2, start_date=DEFAULT_DATE, schedule=None) - EmptyOperator(task_id=task_id, dag=dag2, outlets=[Dataset(uri1, extra={"should": "be used"})]) + EmptyOperator(task_id=task_id, dag=dag2, outlets=[Asset(uri1, extra={"should": "be used"})]) session = settings.Session() dag1.clear() DAG.bulk_write_to_db([dag1, dag2], session=session) session.commit() - stored_datasets = {x.uri: x for x in session.query(DatasetModel).all()} - d1_orm = stored_datasets[d1.uri] - d2_orm = stored_datasets[d2.uri] - d3_orm = stored_datasets[d3.uri] - assert stored_datasets[uri1].extra == {"should": "be used"} - assert [x.dag_id for x in d1_orm.consuming_dags] == [dag_id1] - assert [(x.task_id, x.dag_id) for x in d1_orm.producing_tasks] == [(task_id, dag_id2)] + stored_assets = {x.uri: x for x in session.query(AssetModel).all()} + asset1_orm = stored_assets[d1.uri] + asset2_orm = stored_assets[d2.uri] + asset3_orm = stored_assets[d3.uri] + assert stored_assets[uri1].extra == {"should": "be used"} + assert [x.dag_id for x in asset1_orm.consuming_dags] == [dag_id1] + assert [(x.task_id, x.dag_id) for x in asset1_orm.producing_tasks] == [(task_id, dag_id2)] assert set( session.query( - TaskOutletDatasetReference.task_id, - TaskOutletDatasetReference.dag_id, - TaskOutletDatasetReference.dataset_id, + TaskOutletAssetReference.task_id, + TaskOutletAssetReference.dag_id, + TaskOutletAssetReference.dataset_id, ) - .filter(TaskOutletDatasetReference.dag_id.in_((dag_id1, dag_id2))) + .filter(TaskOutletAssetReference.dag_id.in_((dag_id1, dag_id2))) .all() ) == { - (task_id, dag_id1, d2_orm.id), - (task_id, dag_id1, d3_orm.id), - (task_id, dag_id2, d1_orm.id), + (task_id, dag_id1, asset2_orm.id), + (task_id, dag_id1, asset3_orm.id), + (task_id, dag_id2, asset1_orm.id), } - # now that we have verified that a new dag has its dataset references recorded properly, + # now that we have verified that a new dag has its asset references recorded properly, # we need to verify that *changes* are recorded properly. # so if any references are *removed*, they should also be deleted from the DB # so let's remove some references and see what happens @@ -1055,96 +1055,96 @@ def test_bulk_write_to_db_datasets(self): DAG.bulk_write_to_db([dag1, dag2], session=session) session.commit() session.expunge_all() - stored_datasets = {x.uri: x for x in session.query(DatasetModel).all()} - d1_orm = stored_datasets[d1.uri] - d2_orm = stored_datasets[d2.uri] - assert [x.dag_id for x in d1_orm.consuming_dags] == [] + stored_assets = {x.uri: x for x in session.query(AssetModel).all()} + asset1_orm = stored_assets[d1.uri] + asset2_orm = stored_assets[d2.uri] + assert [x.dag_id for x in asset1_orm.consuming_dags] == [] assert set( session.query( - TaskOutletDatasetReference.task_id, - TaskOutletDatasetReference.dag_id, - TaskOutletDatasetReference.dataset_id, + TaskOutletAssetReference.task_id, + TaskOutletAssetReference.dag_id, + TaskOutletAssetReference.dataset_id, ) - .filter(TaskOutletDatasetReference.dag_id.in_((dag_id1, dag_id2))) + .filter(TaskOutletAssetReference.dag_id.in_((dag_id1, dag_id2))) .all() - ) == {(task_id, dag_id1, d2_orm.id)} + ) == {(task_id, dag_id1, asset2_orm.id)} - def test_bulk_write_to_db_unorphan_datasets(self): + def test_bulk_write_to_db_unorphan_assets(self): """ - Datasets can lose their last reference and be orphaned, but then if a reference to them reappears, we - need to un-orphan those datasets + Assets can lose their last reference and be orphaned, but then if a reference to them reappears, we + need to un-orphan those assets """ with create_session() as session: - # Create four datasets - two that have references and two that are unreferenced and marked as + # Create four assets - two that have references and two that are unreferenced and marked as # orphans - dataset1 = Dataset(uri="ds1") - dataset2 = Dataset(uri="ds2") - session.add(DatasetModel(uri=dataset2.uri, is_orphaned=True)) - dataset3 = Dataset(uri="ds3") - dataset4 = Dataset(uri="ds4") - session.add(DatasetModel(uri=dataset4.uri, is_orphaned=True)) + asset1 = Asset(uri="ds1") + asset2 = Asset(uri="ds2") + session.add(AssetModel(uri=asset2.uri, is_orphaned=True)) + asset3 = Asset(uri="ds3") + asset4 = Asset(uri="ds4") + session.add(AssetModel(uri=asset4.uri, is_orphaned=True)) session.flush() - dag1 = DAG(dag_id="datasets-1", start_date=DEFAULT_DATE, schedule=[dataset1]) - BashOperator(dag=dag1, task_id="task", bash_command="echo 1", outlets=[dataset3]) + dag1 = DAG(dag_id="assets-1", start_date=DEFAULT_DATE, schedule=[asset1]) + BashOperator(dag=dag1, task_id="task", bash_command="echo 1", outlets=[asset3]) DAG.bulk_write_to_db([dag1], session=session) # Double check - non_orphaned_datasets = [ - dataset.uri - for dataset in session.query(DatasetModel.uri) - .filter(~DatasetModel.is_orphaned) - .order_by(DatasetModel.uri) + non_orphaned_assets = [ + asset.uri + for asset in session.query(AssetModel.uri) + .filter(~AssetModel.is_orphaned) + .order_by(AssetModel.uri) ] - assert non_orphaned_datasets == ["ds1", "ds3"] - orphaned_datasets = [ - dataset.uri - for dataset in session.query(DatasetModel.uri) - .filter(DatasetModel.is_orphaned) - .order_by(DatasetModel.uri) + assert non_orphaned_assets == ["ds1", "ds3"] + orphaned_assets = [ + asset.uri + for asset in session.query(AssetModel.uri) + .filter(AssetModel.is_orphaned) + .order_by(AssetModel.uri) ] - assert orphaned_datasets == ["ds2", "ds4"] + assert orphaned_assets == ["ds2", "ds4"] - # Now add references to the two unreferenced datasets - dag1 = DAG(dag_id="datasets-1", start_date=DEFAULT_DATE, schedule=[dataset1, dataset2]) - BashOperator(dag=dag1, task_id="task", bash_command="echo 1", outlets=[dataset3, dataset4]) + # Now add references to the two unreferenced assets + dag1 = DAG(dag_id="assets-1", start_date=DEFAULT_DATE, schedule=[asset1, asset2]) + BashOperator(dag=dag1, task_id="task", bash_command="echo 1", outlets=[asset3, asset4]) DAG.bulk_write_to_db([dag1], session=session) # and count the orphans and non-orphans - non_orphaned_dataset_count = session.query(DatasetModel).filter(~DatasetModel.is_orphaned).count() - assert non_orphaned_dataset_count == 4 - orphaned_dataset_count = session.query(DatasetModel).filter(DatasetModel.is_orphaned).count() - assert orphaned_dataset_count == 0 - - def test_bulk_write_to_db_dataset_aliases(self): - """ - Ensure that dataset aliases referenced in a dag are correctly loaded into the database. - """ - dag_id1 = "test_dataset_alias_dag1" - dag_id2 = "test_dataset_alias_dag2" - task_id = "test_dataset_task" - da1 = DatasetAlias(name="da1") - da2 = DatasetAlias(name="da2") - da2_2 = DatasetAlias(name="da2") - da3 = DatasetAlias(name="da3") + non_orphaned_asset_count = session.query(AssetModel).filter(~AssetModel.is_orphaned).count() + assert non_orphaned_asset_count == 4 + orphaned_asset_count = session.query(AssetModel).filter(AssetModel.is_orphaned).count() + assert orphaned_asset_count == 0 + + def test_bulk_write_to_db_asset_aliases(self): + """ + Ensure that asset aliases referenced in a dag are correctly loaded into the database. + """ + dag_id1 = "test_asset_alias_dag1" + dag_id2 = "test_asset_alias_dag2" + task_id = "test_asset_task" + asset_alias_1 = AssetAlias(name="asset_alias_1") + asset_alias_2 = AssetAlias(name="asset_alias_2") + asset_alias_2_2 = AssetAlias(name="asset_alias_2") + asset_alias_3 = AssetAlias(name="asset_alias_3") dag1 = DAG(dag_id=dag_id1, start_date=DEFAULT_DATE, schedule=None) - EmptyOperator(task_id=task_id, dag=dag1, outlets=[da1, da2, da3]) + EmptyOperator(task_id=task_id, dag=dag1, outlets=[asset_alias_1, asset_alias_2, asset_alias_3]) dag2 = DAG(dag_id=dag_id2, start_date=DEFAULT_DATE, schedule=None) - EmptyOperator(task_id=task_id, dag=dag2, outlets=[da2_2, da3]) + EmptyOperator(task_id=task_id, dag=dag2, outlets=[asset_alias_2_2, asset_alias_3]) session = settings.Session() DAG.bulk_write_to_db([dag1, dag2], session=session) session.commit() - stored_dataset_aliases = {x.name: x for x in session.query(DatasetAliasModel).all()} - da1_orm = stored_dataset_aliases[da1.name] - da2_orm = stored_dataset_aliases[da2.name] - da3_orm = stored_dataset_aliases[da3.name] - assert da1_orm.name == "da1" - assert da2_orm.name == "da2" - assert da3_orm.name == "da3" - assert len(stored_dataset_aliases) == 3 + stored_asset_alias_models = {x.name: x for x in session.query(AssetAliasModel).all()} + asset_alias_1_orm = stored_asset_alias_models[asset_alias_1.name] + asset_alias_2_orm = stored_asset_alias_models[asset_alias_2.name] + asset_alias_3_orm = stored_asset_alias_models[asset_alias_3.name] + assert asset_alias_1_orm.name == "asset_alias_1" + assert asset_alias_2_orm.name == "asset_alias_2" + assert asset_alias_3_orm.name == "asset_alias_3" + assert len(stored_asset_alias_models) == 3 def test_sync_to_db(self): dag = DAG("dag", start_date=DEFAULT_DATE, schedule=None) @@ -1664,10 +1664,10 @@ def test_timetable_and_description_from_schedule_arg( assert dag.timetable == expected_timetable assert dag.timetable.description == interval_description - def test_timetable_and_description_from_dataset(self): - dag = DAG("test_schedule_arg", schedule=[Dataset(uri="hello")], start_date=TEST_DATE) - assert dag.timetable == DatasetTriggeredTimetable(Dataset(uri="hello")) - assert dag.timetable.description == "Triggered by datasets" + def test_timetable_and_description_from_asset(self): + dag = DAG("test_schedule_interval_arg", schedule=[Asset(uri="hello")], start_date=TEST_DATE) + assert dag.timetable == AssetTriggeredTimetable(Asset(uri="hello")) + assert dag.timetable.description == "Triggered by assets" @pytest.mark.parametrize( "timetable, expected_description", @@ -2400,7 +2400,7 @@ def test_continuous_schedule_linmits_max_active_runs(self): class TestDagModel: def _clean(self): clear_db_dags() - clear_db_datasets() + clear_db_assets() clear_db_runs() def setup_method(self): @@ -2432,13 +2432,13 @@ def test_dags_needing_dagruns_not_too_early(self): session.rollback() session.close() - def test_dags_needing_dagruns_datasets(self, dag_maker, session): - dataset = Dataset(uri="hello") + def test_dags_needing_dagruns_assets(self, dag_maker, session): + asset = Asset(uri="hello") with dag_maker( session=session, dag_id="my_dag", max_active_runs=1, - schedule=[dataset], + schedule=[asset], start_date=pendulum.now().add(days=-2), ) as dag: EmptyOperator(task_id="dummy") @@ -2450,8 +2450,8 @@ def test_dags_needing_dagruns_datasets(self, dag_maker, session): # add queue records so we'll need a run dag_model = session.query(DagModel).filter(DagModel.dag_id == dag.dag_id).one() - dataset_model: DatasetModel = dag_model.schedule_datasets[0] - session.add(DatasetDagRunQueue(dataset_id=dataset_model.id, target_dag_id=dag_model.dag_id)) + asset_model: AssetModel = dag_model.schedule_datasets[0] + session.add(AssetDagRunQueue(dataset_id=asset_model.id, target_dag_id=dag_model.dag_id)) session.flush() query, _ = DagModel.dags_needing_dagruns(session) dag_models = query.all() @@ -2474,19 +2474,19 @@ def test_dags_needing_dagruns_datasets(self, dag_maker, session): dag_models = query.all() assert dag_models == [dag_model] - def test_dags_needing_dagruns_dataset_aliases(self, dag_maker, session): - # link dataset_alias hello_alias to dataset hello - dataset_model = DatasetModel(uri="hello") - dataset_alias_model = DatasetAliasModel(name="hello_alias") - dataset_alias_model.datasets.append(dataset_model) - session.add_all([dataset_model, dataset_alias_model]) + def test_dags_needing_dagruns_asset_aliases(self, dag_maker, session): + # link asset_alias hello_alias to asset hello + asset_model = AssetModel(uri="hello") + asset_alias_model = AssetAliasModel(name="hello_alias") + asset_alias_model.datasets.append(asset_model) + session.add_all([asset_model, asset_alias_model]) session.commit() with dag_maker( session=session, dag_id="my_dag", max_active_runs=1, - schedule=[DatasetAlias(name="hello_alias")], + schedule=[AssetAlias(name="hello_alias")], start_date=pendulum.now().add(days=-2), ): EmptyOperator(task_id="dummy") @@ -2498,8 +2498,8 @@ def test_dags_needing_dagruns_dataset_aliases(self, dag_maker, session): # add queue records so we'll need a run dag_model = dag_maker.dag_model - dataset_model: DatasetModel = dag_model.schedule_datasets[0] - session.add(DatasetDagRunQueue(dataset_id=dataset_model.id, target_dag_id=dag_model.dag_id)) + asset_model: AssetModel = dag_model.schedule_datasets[0] + session.add(AssetDagRunQueue(dataset_id=asset_model.id, target_dag_id=dag_model.dag_id)) session.flush() query, _ = DagModel.dags_needing_dagruns(session) dag_models = query.all() @@ -2663,20 +2663,20 @@ def test__processor_dags_folder(self, session): assert sdm.dag._processor_dags_folder == settings.DAGS_FOLDER @pytest.mark.need_serialized_dag - def test_dags_needing_dagruns_dataset_triggered_dag_info_queued_times(self, session, dag_maker): - dataset1 = Dataset(uri="ds1") - dataset2 = Dataset(uri="ds2") + def test_dags_needing_dagruns_asset_triggered_dag_info_queued_times(self, session, dag_maker): + asset1 = Asset(uri="ds1") + asset2 = Asset(uri="ds2") - for dag_id, dataset in [("datasets-1", dataset1), ("datasets-2", dataset2)]: + for dag_id, asset in [("assets-1", asset1), ("assets-2", asset2)]: with dag_maker(dag_id=dag_id, start_date=timezone.utcnow(), session=session): - EmptyOperator(task_id="task", outlets=[dataset]) + EmptyOperator(task_id="task", outlets=[asset]) dr = dag_maker.create_dagrun() - ds_id = session.query(DatasetModel.id).filter_by(uri=dataset.uri).scalar() + asset_id = session.query(AssetModel.id).filter_by(uri=asset.uri).scalar() session.add( - DatasetEvent( - dataset_id=ds_id, + AssetEvent( + dataset_id=asset_id, source_task_id="task", source_dag_id=dr.dag_id, source_run_id=dr.run_id, @@ -2684,18 +2684,20 @@ def test_dags_needing_dagruns_dataset_triggered_dag_info_queued_times(self, sess ) ) - ds1_id = session.query(DatasetModel.id).filter_by(uri=dataset1.uri).scalar() - ds2_id = session.query(DatasetModel.id).filter_by(uri=dataset2.uri).scalar() + asset1_id = session.query(AssetModel.id).filter_by(uri=asset1.uri).scalar() + asset2_id = session.query(AssetModel.id).filter_by(uri=asset2.uri).scalar() - with dag_maker(dag_id="datasets-consumer-multiple", schedule=[dataset1, dataset2]) as dag: + with dag_maker(dag_id="assets-consumer-multiple", schedule=[asset1, asset2]) as dag: pass session.flush() session.add_all( [ - DatasetDagRunQueue(dataset_id=ds1_id, target_dag_id=dag.dag_id, created_at=DEFAULT_DATE), - DatasetDagRunQueue( - dataset_id=ds2_id, target_dag_id=dag.dag_id, created_at=DEFAULT_DATE + timedelta(hours=1) + AssetDagRunQueue(dataset_id=asset1_id, target_dag_id=dag.dag_id, created_at=DEFAULT_DATE), + AssetDagRunQueue( + dataset_id=asset2_id, + target_dag_id=dag.dag_id, + created_at=DEFAULT_DATE + timedelta(hours=1), ), ] ) @@ -2708,16 +2710,16 @@ def test_dags_needing_dagruns_dataset_triggered_dag_info_queued_times(self, sess assert first_queued_time == DEFAULT_DATE assert last_queued_time == DEFAULT_DATE + timedelta(hours=1) - def test_dataset_expression(self, session: Session) -> None: + def test_asset_expression(self, session: Session) -> None: dag = DAG( - dag_id="test_dag_dataset_expression", - schedule=DatasetAny( - Dataset("s3://dag1/output_1.txt", {"hi": "bye"}), - DatasetAll( - Dataset("s3://dag2/output_1.txt", {"hi": "bye"}), - Dataset("s3://dag3/output_3.txt", {"hi": "bye"}), + dag_id="test_dag_asset_expression", + schedule=AssetAny( + Asset("s3://dag1/output_1.txt", {"hi": "bye"}), + AssetAll( + Asset("s3://dag2/output_1.txt", {"hi": "bye"}), + Asset("s3://dag3/output_3.txt", {"hi": "bye"}), ), - DatasetAlias(name="test_name"), + AssetAlias(name="test_name"), ), start_date=datetime.datetime.min, ) @@ -3424,43 +3426,43 @@ def test__tags_mutable(): @pytest.mark.need_serialized_dag -def test_get_dataset_triggered_next_run_info(dag_maker, clear_datasets): - dataset1 = Dataset(uri="ds1") - dataset2 = Dataset(uri="ds2") - dataset3 = Dataset(uri="ds3") - with dag_maker(dag_id="datasets-1", schedule=[dataset2]): +def test_get_asset_triggered_next_run_info(dag_maker, clear_assets): + asset1 = Asset(uri="ds1") + asset2 = Asset(uri="ds2") + asset3 = Asset(uri="ds3") + with dag_maker(dag_id="assets-1", schedule=[asset2]): pass dag1 = dag_maker.dag - with dag_maker(dag_id="datasets-2", schedule=[dataset1, dataset2]): + with dag_maker(dag_id="assets-2", schedule=[asset1, asset2]): pass dag2 = dag_maker.dag - with dag_maker(dag_id="datasets-3", schedule=[dataset1, dataset2, dataset3]): + with dag_maker(dag_id="assets-3", schedule=[asset1, asset2, asset3]): pass dag3 = dag_maker.dag session = dag_maker.session - ds1_id = session.query(DatasetModel.id).filter_by(uri=dataset1.uri).scalar() + asset1_id = session.query(AssetModel.id).filter_by(uri=asset1.uri).scalar() session.bulk_save_objects( [ - DatasetDagRunQueue(dataset_id=ds1_id, target_dag_id=dag2.dag_id), - DatasetDagRunQueue(dataset_id=ds1_id, target_dag_id=dag3.dag_id), + AssetDagRunQueue(dataset_id=asset1_id, target_dag_id=dag2.dag_id), + AssetDagRunQueue(dataset_id=asset1_id, target_dag_id=dag3.dag_id), ] ) session.flush() - datasets = session.query(DatasetModel.uri).order_by(DatasetModel.id).all() + assets = session.query(AssetModel.uri).order_by(AssetModel.id).all() - info = get_dataset_triggered_next_run_info([dag1.dag_id], session=session) + info = get_asset_triggered_next_run_info([dag1.dag_id], session=session) assert info[dag1.dag_id] == { "ready": 0, "total": 1, - "uri": datasets[0].uri, + "uri": assets[0].uri, } # This time, check both dag2 and dag3 at the same time (tests filtering) - info = get_dataset_triggered_next_run_info([dag2.dag_id, dag3.dag_id], session=session) + info = get_asset_triggered_next_run_info([dag2.dag_id, dag3.dag_id], session=session) assert info[dag2.dag_id] == { "ready": 1, "total": 2, @@ -3474,19 +3476,19 @@ def test_get_dataset_triggered_next_run_info(dag_maker, clear_datasets): @pytest.mark.need_serialized_dag -def test_get_dataset_triggered_next_run_info_with_unresolved_dataset_alias(dag_maker, clear_datasets): - dataset_alias1 = DatasetAlias(name="alias") +def test_get_dataset_triggered_next_run_info_with_unresolved_dataset_alias(dag_maker, clear_assets): + dataset_alias1 = AssetAlias(name="alias") with dag_maker(dag_id="dag-1", schedule=[dataset_alias1]): pass dag1 = dag_maker.dag session = dag_maker.session session.flush() - info = get_dataset_triggered_next_run_info([dag1.dag_id], session=session) + info = get_asset_triggered_next_run_info([dag1.dag_id], session=session) assert info == {} dag1_model = DagModel.get_dagmodel(dag1.dag_id) - assert dag1_model.get_dataset_triggered_next_run_info(session=session) is None + assert dag1_model.get_asset_triggered_next_run_info(session=session) is None def test_dag_uses_timetable_for_run_id(session): diff --git a/tests/models/test_dagrun.py b/tests/models/test_dagrun.py index d2f70ce69314b..c7dacaeb291e4 100644 --- a/tests/models/test_dagrun.py +++ b/tests/models/test_dagrun.py @@ -85,7 +85,7 @@ def _clean_db(): db.clear_db_pools() db.clear_db_dags() db.clear_db_variables() - db.clear_db_datasets() + db.clear_db_assets() db.clear_db_xcom() db.clear_db_task_fail() diff --git a/tests/models/test_serialized_dag.py b/tests/models/test_serialized_dag.py index 9f83280f8eb38..b8fddc655dae5 100644 --- a/tests/models/test_serialized_dag.py +++ b/tests/models/test_serialized_dag.py @@ -25,7 +25,7 @@ import pytest import airflow.example_dags as example_dags_module -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.models.dag import DAG from airflow.models.dagbag import DagBag from airflow.models.dagcode import DagCode @@ -237,16 +237,16 @@ def test_order_of_deps_is_consistent(self): dag_id="example", start_date=pendulum.datetime(2021, 1, 1, tz="UTC"), schedule=[ - Dataset("1"), - Dataset("2"), - Dataset("3"), - Dataset("4"), - Dataset("5"), + Asset("1"), + Asset("2"), + Asset("3"), + Asset("4"), + Asset("5"), ], ) as dag6: BashOperator( task_id="any", - outlets=[Dataset("0*"), Dataset("6*")], + outlets=[Asset("0*"), Asset("6*")], bash_command="sleep 5", ) deps_order = [x["dependency_id"] for x in SerializedDAG.serialize_dag(dag6)["dag_dependencies"]] diff --git a/tests/models/test_taskinstance.py b/tests/models/test_taskinstance.py index d2922db267805..8c334366f0488 100644 --- a/tests/models/test_taskinstance.py +++ b/tests/models/test_taskinstance.py @@ -39,7 +39,7 @@ from sqlalchemy import select from airflow import settings -from airflow.datasets import DatasetAlias +from airflow.assets import AssetAlias from airflow.decorators import task, task_group from airflow.example_dags.plugins.workday import AfterWorkdayTimetable from airflow.exceptions import ( @@ -53,11 +53,11 @@ UnmappableXComTypePushed, XComForMappingNotPushed, ) +from airflow.models.asset import AssetAliasModel, AssetDagRunQueue, AssetEvent, AssetModel from airflow.models.connection import Connection from airflow.models.dag import DAG from airflow.models.dagbag import DagBag from airflow.models.dagrun import DagRun -from airflow.models.dataset import DatasetAliasModel, DatasetDagRunQueue, DatasetEvent, DatasetModel from airflow.models.expandinput import EXPAND_INPUT_EMPTY, NotFullyPopulated from airflow.models.param import process_params from airflow.models.pool import Pool @@ -158,7 +158,7 @@ def clean_db(): db.clear_db_task_fail() db.clear_rendered_ti_fields() db.clear_db_task_reschedule() - db.clear_db_datasets() + db.clear_db_assets() db.clear_db_xcom() def setup_method(self): @@ -2269,16 +2269,16 @@ def test_success_callback_no_race_condition(self, create_task_instance): assert ti.state == State.SUCCESS @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode - def test_outlet_datasets(self, create_task_instance): + def test_outlet_assets(self, create_task_instance): """ - Verify that when we have an outlet dataset on a task, and the task - completes successfully, a DatasetDagRunQueue is logged. + Verify that when we have an outlet asset on a task, and the task + completes successfully, a AssetDagRunQueue is logged. """ - from airflow.example_dags import example_datasets - from airflow.example_dags.example_datasets import dag1 + from airflow.example_dags import example_assets + from airflow.example_dags.example_assets import dag1 session = settings.Session() - dagbag = DagBag(dag_folder=example_datasets.__file__) + dagbag = DagBag(dag_folder=example_assets.__file__) dagbag.collect_dags(only_if_updated=False, safe_mode=False) dagbag.sync_to_db(session=session) run_id = str(uuid4()) @@ -2293,54 +2293,54 @@ def test_outlet_datasets(self, create_task_instance): ti.refresh_from_db() assert ti.state == TaskInstanceState.SUCCESS - # check that no other dataset events recorded + # check that no other asset events recorded event = ( - session.query(DatasetEvent) - .join(DatasetEvent.dataset) - .filter(DatasetEvent.source_task_instance == ti) + session.query(AssetEvent) + .join(AssetEvent.dataset) + .filter(AssetEvent.source_task_instance == ti) .one() ) assert event assert event.dataset - # check that one queue record created for each dag that depends on dataset 1 - assert session.query(DatasetDagRunQueue.target_dag_id).filter_by( - dataset_id=event.dataset.id - ).order_by(DatasetDagRunQueue.target_dag_id).all() == [ - ("conditional_dataset_and_time_based_timetable",), - ("consume_1_and_2_with_dataset_expressions",), - ("consume_1_or_2_with_dataset_expressions",), - ("consume_1_or_both_2_and_3_with_dataset_expressions",), - ("dataset_consumes_1",), - ("dataset_consumes_1_and_2",), - ("dataset_consumes_1_never_scheduled",), + # check that one queue record created for each dag that depends on asset 1 + assert session.query(AssetDagRunQueue.target_dag_id).filter_by(dataset_id=event.dataset.id).order_by( + AssetDagRunQueue.target_dag_id + ).all() == [ + ("asset_consumes_1",), + ("asset_consumes_1_and_2",), + ("asset_consumes_1_never_scheduled",), + ("conditional_asset_and_time_based_timetable",), + ("consume_1_and_2_with_asset_expressions",), + ("consume_1_or_2_with_asset_expressions",), + ("consume_1_or_both_2_and_3_with_asset_expressions",), ] - # check that one event record created for dataset1 and this TI - assert session.query(DatasetModel.uri).join(DatasetEvent.dataset).filter( - DatasetEvent.source_task_instance == ti + # check that one event record created for asset1 and this TI + assert session.query(AssetModel.uri).join(AssetEvent.dataset).filter( + AssetEvent.source_task_instance == ti ).one() == ("s3://dag1/output_1.txt",) - # check that the dataset event has an earlier timestamp than the DDRQ's - ddrq_timestamps = ( - session.query(DatasetDagRunQueue.created_at).filter_by(dataset_id=event.dataset.id).all() + # check that the asset event has an earlier timestamp than the ADRQ's + adrq_timestamps = ( + session.query(AssetDagRunQueue.created_at).filter_by(dataset_id=event.dataset.id).all() ) assert all( - event.timestamp < ddrq_timestamp for (ddrq_timestamp,) in ddrq_timestamps - ), f"Some items in {[str(t) for t in ddrq_timestamps]} are earlier than {event.timestamp}" + event.timestamp < adrq_timestamp for (adrq_timestamp,) in adrq_timestamps + ), f"Some items in {[str(t) for t in adrq_timestamps]} are earlier than {event.timestamp}" @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode - def test_outlet_datasets_failed(self, create_task_instance): + def test_outlet_assets_failed(self, create_task_instance): """ - Verify that when we have an outlet dataset on a task, and the task - failed, a DatasetDagRunQueue is not logged, and a DatasetEvent is + Verify that when we have an outlet asset on a task, and the task + failed, a AssetDagRunQueue is not logged, and an AssetEvent is not generated """ - from tests.dags import test_datasets - from tests.dags.test_datasets import dag_with_fail_task + from tests.dags import test_assets + from tests.dags.test_assets import dag_with_fail_task session = settings.Session() - dagbag = DagBag(dag_folder=test_datasets.__file__) + dagbag = DagBag(dag_folder=test_assets.__file__) dagbag.collect_dags(only_if_updated=False, safe_mode=False) dagbag.sync_to_db(session=session) run_id = str(uuid4()) @@ -2356,10 +2356,10 @@ def test_outlet_datasets_failed(self, create_task_instance): assert ti.state == TaskInstanceState.FAILED # check that no dagruns were queued - assert session.query(DatasetDagRunQueue).count() == 0 + assert session.query(AssetDagRunQueue).count() == 0 - # check that no dataset events were generated - assert session.query(DatasetEvent).count() == 0 + # check that no asset events were generated + assert session.query(AssetEvent).count() == 0 @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode def test_mapped_current_state(self, dag_maker): @@ -2386,17 +2386,17 @@ def raise_an_exception(placeholder: int): assert task_instance.current_state() == TaskInstanceState.SUCCESS @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode - def test_outlet_datasets_skipped(self): + def test_outlet_assets_skipped(self): """ - Verify that when we have an outlet dataset on a task, and the task - is skipped, a DatasetDagRunQueue is not logged, and a DatasetEvent is + Verify that when we have an outlet asset on a task, and the task + is skipped, a AssetDagRunQueue is not logged, and an AssetEvent is not generated """ - from tests.dags import test_datasets - from tests.dags.test_datasets import dag_with_skip_task + from tests.dags import test_assets + from tests.dags.test_assets import dag_with_skip_task session = settings.Session() - dagbag = DagBag(dag_folder=test_datasets.__file__) + dagbag = DagBag(dag_folder=test_assets.__file__) dagbag.collect_dags(only_if_updated=False, safe_mode=False) dagbag.sync_to_db(session=session) run_id = str(uuid4()) @@ -2411,30 +2411,30 @@ def test_outlet_datasets_skipped(self): assert ti.state == TaskInstanceState.SKIPPED # check that no dagruns were queued - assert session.query(DatasetDagRunQueue).count() == 0 + assert session.query(AssetDagRunQueue).count() == 0 - # check that no dataset events were generated - assert session.query(DatasetEvent).count() == 0 + # check that no asset events were generated + assert session.query(AssetEvent).count() == 0 @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode - def test_outlet_dataset_extra(self, dag_maker, session): - from airflow.datasets import Dataset + def test_outlet_asset_extra(self, dag_maker, session): + from airflow.assets import Asset with dag_maker(schedule=None, session=session) as dag: - @task(outlets=Dataset("test_outlet_dataset_extra_1")) + @task(outlets=Asset("test_outlet_asset_extra_1")) def write1(*, outlet_events): - outlet_events["test_outlet_dataset_extra_1"].extra = {"foo": "bar"} + outlet_events["test_outlet_asset_extra_1"].extra = {"foo": "bar"} write1() def _write2_post_execute(context, _): - context["outlet_events"]["test_outlet_dataset_extra_2"].extra = {"x": 1} + context["outlet_events"]["test_outlet_asset_extra_2"].extra = {"x": 1} BashOperator( task_id="write2", bash_command=":", - outlets=Dataset("test_outlet_dataset_extra_2"), + outlets=Asset("test_outlet_asset_extra_2"), post_execute=_write2_post_execute, ) @@ -2443,30 +2443,30 @@ def _write2_post_execute(context, _): ti.refresh_from_task(dag.get_task(ti.task_id)) ti.run(session=session) - events = dict(iter(session.execute(select(DatasetEvent.source_task_id, DatasetEvent)))) + events = dict(iter(session.execute(select(AssetEvent.source_task_id, AssetEvent)))) assert set(events) == {"write1", "write2"} assert events["write1"].source_dag_id == dr.dag_id assert events["write1"].source_run_id == dr.run_id assert events["write1"].source_task_id == "write1" - assert events["write1"].dataset.uri == "test_outlet_dataset_extra_1" + assert events["write1"].dataset.uri == "test_outlet_asset_extra_1" assert events["write1"].extra == {"foo": "bar"} assert events["write2"].source_dag_id == dr.dag_id assert events["write2"].source_run_id == dr.run_id assert events["write2"].source_task_id == "write2" - assert events["write2"].dataset.uri == "test_outlet_dataset_extra_2" + assert events["write2"].dataset.uri == "test_outlet_asset_extra_2" assert events["write2"].extra == {"x": 1} @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode - def test_outlet_dataset_extra_ignore_different(self, dag_maker, session): - from airflow.datasets import Dataset + def test_outlet_asset_extra_ignore_different(self, dag_maker, session): + from airflow.assets import Asset with dag_maker(schedule=None, session=session): - @task(outlets=Dataset("test_outlet_dataset_extra")) + @task(outlets=Asset("test_outlet_asset_extra")) def write(*, outlet_events): - outlet_events["test_outlet_dataset_extra"].extra = {"one": 1} + outlet_events["test_outlet_asset_extra"].extra = {"one": 1} outlet_events["different_uri"].extra = {"foo": "bar"} # Will be silently dropped. write() @@ -2474,34 +2474,34 @@ def write(*, outlet_events): dr: DagRun = dag_maker.create_dagrun() dr.get_task_instance("write").run(session=session) - event = session.scalars(select(DatasetEvent)).one() + event = session.scalars(select(AssetEvent)).one() assert event.source_dag_id == dr.dag_id assert event.source_run_id == dr.run_id assert event.source_task_id == "write" assert event.extra == {"one": 1} @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode - def test_outlet_dataset_extra_yield(self, dag_maker, session): - from airflow.datasets import Dataset - from airflow.datasets.metadata import Metadata + def test_outlet_asset_extra_yield(self, dag_maker, session): + from airflow.assets import Asset + from airflow.assets.metadata import Metadata with dag_maker(schedule=None, session=session) as dag: - @task(outlets=Dataset("test_outlet_dataset_extra_1")) + @task(outlets=Asset("test_outlet_asset_extra_1")) def write1(): result = "write_1 result" - yield Metadata("test_outlet_dataset_extra_1", {"foo": "bar"}) + yield Metadata("test_outlet_asset_extra_1", {"foo": "bar"}) return result write1() def _write2_post_execute(context, result): - yield Metadata("test_outlet_dataset_extra_2", {"x": 1}) + yield Metadata("test_outlet_asset_extra_2", {"x": 1}) BashOperator( task_id="write2", bash_command=":", - outlets=Dataset("test_outlet_dataset_extra_2"), + outlets=Asset("test_outlet_asset_extra_2"), post_execute=_write2_post_execute, ) @@ -2515,37 +2515,37 @@ def _write2_post_execute(context, result): ).one() assert xcom.value == "write_1 result" - events = dict(iter(session.execute(select(DatasetEvent.source_task_id, DatasetEvent)))) + events = dict(iter(session.execute(select(AssetEvent.source_task_id, AssetEvent)))) assert set(events) == {"write1", "write2"} assert events["write1"].source_dag_id == dr.dag_id assert events["write1"].source_run_id == dr.run_id assert events["write1"].source_task_id == "write1" - assert events["write1"].dataset.uri == "test_outlet_dataset_extra_1" + assert events["write1"].dataset.uri == "test_outlet_asset_extra_1" assert events["write1"].extra == {"foo": "bar"} assert events["write2"].source_dag_id == dr.dag_id assert events["write2"].source_run_id == dr.run_id assert events["write2"].source_task_id == "write2" - assert events["write2"].dataset.uri == "test_outlet_dataset_extra_2" + assert events["write2"].dataset.uri == "test_outlet_asset_extra_2" assert events["write2"].extra == {"x": 1} @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode - def test_outlet_dataset_alias(self, dag_maker, session): - from airflow.datasets import Dataset, DatasetAlias + def test_outlet_asset_alias(self, dag_maker, session): + from airflow.assets import Asset, AssetAlias - ds_uri = "test_outlet_dataset_alias_test_case_ds" - dsa_name_1 = "test_outlet_dataset_alias_test_case_dsa_1" + asset_uri = "test_outlet_asset_alias_test_case_ds" + alias_name_1 = "test_outlet_asset_alias_test_case_asset_alias_1" - ds1 = DatasetModel(id=1, uri=ds_uri) + ds1 = AssetModel(id=1, uri=asset_uri) session.add(ds1) session.commit() with dag_maker(dag_id="producer_dag", schedule=None, session=session) as dag: - @task(outlets=DatasetAlias(dsa_name_1)) + @task(outlets=AssetAlias(alias_name_1)) def producer(*, outlet_events): - outlet_events[dsa_name_1].add(Dataset(ds_uri)) + outlet_events[alias_name_1].add(Asset(asset_uri)) producer() @@ -2556,7 +2556,7 @@ def producer(*, outlet_events): ti.run(session=session) producer_events = session.execute( - select(DatasetEvent).where(DatasetEvent.source_task_id == "producer") + select(AssetEvent).where(AssetEvent.source_task_id == "producer") ).fetchall() assert len(producer_events) == 1 @@ -2566,39 +2566,45 @@ def producer(*, outlet_events): assert producer_event.source_dag_id == "producer_dag" assert producer_event.source_run_id == "test" assert producer_event.source_map_index == -1 - assert producer_event.dataset.uri == ds_uri + assert producer_event.dataset.uri == asset_uri assert len(producer_event.source_aliases) == 1 assert producer_event.extra == {} - assert producer_event.source_aliases[0].name == dsa_name_1 + assert producer_event.source_aliases[0].name == alias_name_1 - ds_obj = session.scalar(select(DatasetModel).where(DatasetModel.uri == ds_uri)) - assert len(ds_obj.aliases) == 1 - assert ds_obj.aliases[0].name == dsa_name_1 + asset_obj = session.scalar(select(AssetModel).where(AssetModel.uri == asset_uri)) + assert len(asset_obj.aliases) == 1 + assert asset_obj.aliases[0].name == alias_name_1 - dsa_obj = session.scalar(select(DatasetAliasModel).where(DatasetAliasModel.name == dsa_name_1)) - assert len(dsa_obj.datasets) == 1 - assert dsa_obj.datasets[0].uri == ds_uri + asset_alias_obj = session.scalar(select(AssetAliasModel).where(AssetAliasModel.name == alias_name_1)) + assert len(asset_alias_obj.datasets) == 1 + assert asset_alias_obj.datasets[0].uri == asset_uri @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode - def test_outlet_multiple_dataset_alias(self, dag_maker, session): - from airflow.datasets import Dataset, DatasetAlias + def test_outlet_multiple_asset_alias(self, dag_maker, session): + from airflow.assets import Asset, AssetAlias - ds_uri = "test_outlet_mdsa_ds" - dsa_name_1 = "test_outlet_mdsa_dsa_1" - dsa_name_2 = "test_outlet_mdsa_dsa_2" - dsa_name_3 = "test_outlet_mdsa_dsa_3" + asset_uri = "test_outlet_maa_ds" + asset_alias_name_1 = "test_outlet_maa_asset_alias_1" + asset_alias_name_2 = "test_outlet_maa_asset_alias_2" + asset_alias_name_3 = "test_outlet_maa_asset_alias_3" - ds1 = DatasetModel(id=1, uri=ds_uri) + ds1 = AssetModel(id=1, uri=asset_uri) session.add(ds1) session.commit() with dag_maker(dag_id="producer_dag", schedule=None, session=session) as dag: - @task(outlets=[DatasetAlias(dsa_name_1), DatasetAlias(dsa_name_2), DatasetAlias(dsa_name_3)]) + @task( + outlets=[ + AssetAlias(asset_alias_name_1), + AssetAlias(asset_alias_name_2), + AssetAlias(asset_alias_name_3), + ] + ) def producer(*, outlet_events): - outlet_events[dsa_name_1].add(Dataset(ds_uri)) - outlet_events[dsa_name_2].add(Dataset(ds_uri)) - outlet_events[dsa_name_3].add(Dataset(ds_uri), extra={"k": "v"}) + outlet_events[asset_alias_name_1].add(Asset(asset_uri)) + outlet_events[asset_alias_name_2].add(Asset(asset_uri)) + outlet_events[asset_alias_name_3].add(Asset(asset_uri), extra={"k": "v"}) producer() @@ -2609,7 +2615,7 @@ def producer(*, outlet_events): ti.run(session=session) producer_events = session.execute( - select(DatasetEvent).where(DatasetEvent.source_task_id == "producer") + select(AssetEvent).where(AssetEvent.source_task_id == "producer") ).fetchall() assert len(producer_events) == 2 @@ -2619,44 +2625,51 @@ def producer(*, outlet_events): assert producer_event.source_dag_id == "producer_dag" assert producer_event.source_run_id == "test" assert producer_event.source_map_index == -1 - assert producer_event.dataset.uri == ds_uri + assert producer_event.dataset.uri == asset_uri if not producer_event.extra: assert producer_event.extra == {} assert len(producer_event.source_aliases) == 2 - assert {alias.name for alias in producer_event.source_aliases} == {dsa_name_1, dsa_name_2} + assert {alias.name for alias in producer_event.source_aliases} == { + asset_alias_name_1, + asset_alias_name_2, + } else: assert producer_event.extra == {"k": "v"} assert len(producer_event.source_aliases) == 1 - assert producer_event.source_aliases[0].name == dsa_name_3 - - ds_obj = session.scalar(select(DatasetModel).where(DatasetModel.uri == ds_uri)) - assert len(ds_obj.aliases) == 3 - assert {alias.name for alias in ds_obj.aliases} == {dsa_name_1, dsa_name_2, dsa_name_3} + assert producer_event.source_aliases[0].name == asset_alias_name_3 + + asset_obj = session.scalar(select(AssetModel).where(AssetModel.uri == asset_uri)) + assert len(asset_obj.aliases) == 3 + assert {alias.name for alias in asset_obj.aliases} == { + asset_alias_name_1, + asset_alias_name_2, + asset_alias_name_3, + } - dsa_objs = session.scalars(select(DatasetAliasModel)).all() - assert len(dsa_objs) == 3 - for dsa_obj in dsa_objs: - assert len(dsa_obj.datasets) == 1 - assert dsa_obj.datasets[0].uri == ds_uri + asset_alias_objs = session.scalars(select(AssetAliasModel)).all() + assert len(asset_alias_objs) == 3 + for asset_alias_obj in asset_alias_objs: + assert len(asset_alias_obj.datasets) == 1 + assert asset_alias_obj.datasets[0].uri == asset_uri @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode - def test_outlet_dataset_alias_through_metadata(self, dag_maker, session): - from airflow.datasets import DatasetAlias - from airflow.datasets.metadata import Metadata + def test_outlet_asset_alias_through_metadata(self, dag_maker, session): + from airflow.assets import AssetAlias + from airflow.assets.metadata import Metadata - ds_uri = "test_outlet_dataset_alias_through_metadata_ds" - dsa_name = "test_outlet_dataset_alias_through_metadata_dsa" + asset_uri = "test_outlet_asset_alias_through_metadata_ds" + asset_alias_name = "test_outlet_asset_alias_through_metadata_asset_alias" - ds1 = DatasetModel(id=1, uri="test_outlet_dataset_alias_through_metadata_ds") + ds1 = AssetModel(id=1, uri="test_outlet_asset_alias_through_metadata_ds") session.add(ds1) session.commit() with dag_maker(dag_id="producer_dag", schedule=None, session=session) as dag: - @task(outlets=DatasetAlias(dsa_name)) + @task(outlets=AssetAlias(asset_alias_name)) def producer(*, outlet_events): - yield Metadata(ds_uri, extra={"key": "value"}, alias=dsa_name) + yield Metadata(asset_uri, extra={"key": "value"}, alias=asset_alias_name) producer() @@ -2666,37 +2679,37 @@ def producer(*, outlet_events): ti.refresh_from_task(dag.get_task(ti.task_id)) ti.run(session=session) - producer_event = session.scalar(select(DatasetEvent).where(DatasetEvent.source_task_id == "producer")) + producer_event = session.scalar(select(AssetEvent).where(AssetEvent.source_task_id == "producer")) assert producer_event.source_task_id == "producer" assert producer_event.source_dag_id == "producer_dag" assert producer_event.source_run_id == "test" assert producer_event.source_map_index == -1 - assert producer_event.dataset.uri == ds_uri + assert producer_event.dataset.uri == asset_uri assert producer_event.extra == {"key": "value"} assert len(producer_event.source_aliases) == 1 - assert producer_event.source_aliases[0].name == dsa_name + assert producer_event.source_aliases[0].name == asset_alias_name - ds_obj = session.scalar(select(DatasetModel).where(DatasetModel.uri == ds_uri)) - assert len(ds_obj.aliases) == 1 - assert ds_obj.aliases[0].name == dsa_name + asset_obj = session.scalar(select(AssetModel).where(AssetModel.uri == asset_uri)) + assert len(asset_obj.aliases) == 1 + assert asset_obj.aliases[0].name == asset_alias_name - dsa_obj = session.scalar(select(DatasetAliasModel)) - assert len(dsa_obj.datasets) == 1 - assert dsa_obj.datasets[0].uri == ds_uri + asset_alias_obj = session.scalar(select(AssetAliasModel)) + assert len(asset_alias_obj.datasets) == 1 + assert asset_alias_obj.datasets[0].uri == asset_uri @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode - def test_outlet_dataset_alias_dataset_not_exists(self, dag_maker, session): - from airflow.datasets import Dataset, DatasetAlias + def test_outlet_asset_alias_asset_not_exists(self, dag_maker, session): + from airflow.assets import Asset, AssetAlias - dsa_name = "test_outlet_dataset_alias_dataset_not_exists_dsa" - ds_uri = "did_not_exists" + asset_alias_name = "test_outlet_asset_alias_asset_not_exists_asset_alias" + asset_uri = "did_not_exists" with dag_maker(dag_id="producer_dag", schedule=None, session=session) as dag: - @task(outlets=DatasetAlias(dsa_name)) + @task(outlets=AssetAlias(asset_alias_name)) def producer(*, outlet_events): - outlet_events[dsa_name].add(Dataset(ds_uri), extra={"key": "value"}) + outlet_events[asset_alias_name].add(Asset(asset_uri), extra={"key": "value"}) producer() @@ -2706,51 +2719,51 @@ def producer(*, outlet_events): ti.refresh_from_task(dag.get_task(ti.task_id)) ti.run(session=session) - producer_event = session.scalar(select(DatasetEvent).where(DatasetEvent.source_task_id == "producer")) + producer_event = session.scalar(select(AssetEvent).where(AssetEvent.source_task_id == "producer")) assert producer_event.source_task_id == "producer" assert producer_event.source_dag_id == "producer_dag" assert producer_event.source_run_id == "test" assert producer_event.source_map_index == -1 - assert producer_event.dataset.uri == ds_uri + assert producer_event.dataset.uri == asset_uri assert producer_event.extra == {"key": "value"} assert len(producer_event.source_aliases) == 1 - assert producer_event.source_aliases[0].name == dsa_name + assert producer_event.source_aliases[0].name == asset_alias_name - ds_obj = session.scalar(select(DatasetModel).where(DatasetModel.uri == ds_uri)) - assert len(ds_obj.aliases) == 1 - assert ds_obj.aliases[0].name == dsa_name + asset_obj = session.scalar(select(AssetModel).where(AssetModel.uri == asset_uri)) + assert len(asset_obj.aliases) == 1 + assert asset_obj.aliases[0].name == asset_alias_name - dsa_obj = session.scalar(select(DatasetAliasModel)) - assert len(dsa_obj.datasets) == 1 - assert dsa_obj.datasets[0].uri == ds_uri + asset_alias_obj = session.scalar(select(AssetAliasModel)) + assert len(asset_alias_obj.datasets) == 1 + assert asset_alias_obj.datasets[0].uri == asset_uri @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode - def test_inlet_dataset_extra(self, dag_maker, session): - from airflow.datasets import Dataset + def test_inlet_asset_extra(self, dag_maker, session): + from airflow.assets import Asset read_task_evaluated = False with dag_maker(schedule=None, session=session): - @task(outlets=Dataset("test_inlet_dataset_extra")) + @task(outlets=Asset("test_inlet_asset_extra")) def write(*, ti, outlet_events): - outlet_events["test_inlet_dataset_extra"].extra = {"from": ti.task_id} + outlet_events["test_inlet_asset_extra"].extra = {"from": ti.task_id} - @task(inlets=Dataset("test_inlet_dataset_extra")) + @task(inlets=Asset("test_inlet_asset_extra")) def read(*, inlet_events): - second_event = inlet_events["test_inlet_dataset_extra"][1] - assert second_event.uri == "test_inlet_dataset_extra" + second_event = inlet_events["test_inlet_asset_extra"][1] + assert second_event.uri == "test_inlet_asset_extra" assert second_event.extra == {"from": "write2"} - last_event = inlet_events["test_inlet_dataset_extra"][-1] - assert last_event.uri == "test_inlet_dataset_extra" + last_event = inlet_events["test_inlet_asset_extra"][-1] + assert last_event.uri == "test_inlet_asset_extra" assert last_event.extra == {"from": "write3"} with pytest.raises(KeyError): inlet_events["does_not_exist"] with pytest.raises(IndexError): - inlet_events["test_inlet_dataset_extra"][5] + inlet_events["test_inlet_asset_extra"][5] # TODO: Support slices. @@ -2780,42 +2793,42 @@ def read(*, inlet_events): assert read_task_evaluated @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode - def test_inlet_dataset_alias_extra(self, dag_maker, session): - ds_uri = "test_inlet_dataset_extra_ds" - dsa_name = "test_inlet_dataset_extra_dsa" - - ds_model = DatasetModel(id=1, uri=ds_uri) - dsa_model = DatasetAliasModel(name=dsa_name) - dsa_model.datasets.append(ds_model) - session.add_all([ds_model, dsa_model]) + def test_inlet_asset_alias_extra(self, dag_maker, session): + asset_uri = "test_inlet_asset_extra_ds" + asset_alias_name = "test_inlet_asset_extra_asset_alias" + + asset_model = AssetModel(id=1, uri=asset_uri) + asset_alias_model = AssetAliasModel(name=asset_alias_name) + asset_alias_model.datasets.append(asset_model) + session.add_all([asset_model, asset_alias_model]) session.commit() - from airflow.datasets import Dataset, DatasetAlias + from airflow.assets import Asset, AssetAlias read_task_evaluated = False with dag_maker(schedule=None, session=session): - @task(outlets=DatasetAlias(dsa_name)) + @task(outlets=AssetAlias(asset_alias_name)) def write(*, ti, outlet_events): - outlet_events[dsa_name].add(Dataset(ds_uri), extra={"from": ti.task_id}) + outlet_events[asset_alias_name].add(Asset(asset_uri), extra={"from": ti.task_id}) - @task(inlets=DatasetAlias(dsa_name)) + @task(inlets=AssetAlias(asset_alias_name)) def read(*, inlet_events): - second_event = inlet_events[DatasetAlias(dsa_name)][1] - assert second_event.uri == ds_uri + second_event = inlet_events[AssetAlias(asset_alias_name)][1] + assert second_event.uri == asset_uri assert second_event.extra == {"from": "write2"} - last_event = inlet_events[DatasetAlias(dsa_name)][-1] - assert last_event.uri == ds_uri + last_event = inlet_events[AssetAlias(asset_alias_name)][-1] + assert last_event.uri == asset_uri assert last_event.extra == {"from": "write3"} with pytest.raises(KeyError): inlet_events["does_not_exist"] with pytest.raises(KeyError): - inlet_events[DatasetAlias("does_not_exist")] + inlet_events[AssetAlias("does_not_exist")] with pytest.raises(IndexError): - inlet_events[DatasetAlias(dsa_name)][5] + inlet_events[AssetAlias(asset_alias_name)][5] nonlocal read_task_evaluated read_task_evaluated = True @@ -2842,21 +2855,21 @@ def read(*, inlet_events): assert not dr.task_instance_scheduling_decisions(session=session).schedulable_tis assert read_task_evaluated - def test_inlet_unresolved_dataset_alias(self, dag_maker, session): - dsa_name = "test_inlet_dataset_extra_dsa" + def test_inlet_unresolved_asset_alias(self, dag_maker, session): + asset_alias_name = "test_inlet_asset_extra_asset_alias" - dsa_model = DatasetAliasModel(name=dsa_name) - session.add(dsa_model) + asset_alias_model = AssetAliasModel(name=asset_alias_name) + session.add(asset_alias_model) session.commit() - from airflow.datasets import DatasetAlias + from airflow.assets import AssetAlias with dag_maker(schedule=None, session=session): - @task(inlets=DatasetAlias(dsa_name)) + @task(inlets=AssetAlias(asset_alias_name)) def read(*, inlet_events): with pytest.raises(IndexError): - inlet_events[DatasetAlias(dsa_name)][0] + inlet_events[AssetAlias(asset_alias_name)][0] read() @@ -2879,16 +2892,16 @@ def read(*, inlet_events): (lambda x: x[-5:5], []), ], ) - def test_inlet_dataset_extra_slice(self, dag_maker, session, slicer, expected): - from airflow.datasets import Dataset + def test_inlet_asset_extra_slice(self, dag_maker, session, slicer, expected): + from airflow.assets import Asset - ds_uri = "test_inlet_dataset_extra_slice" + asset_uri = "test_inlet_asset_extra_slice" with dag_maker(dag_id="write", schedule="@daily", params={"i": -1}, session=session): - @task(outlets=Dataset(ds_uri)) + @task(outlets=Asset(asset_uri)) def write(*, params, outlet_events): - outlet_events[ds_uri].extra = {"from": params["i"]} + outlet_events[asset_uri].extra = {"from": params["i"]} write() @@ -2905,10 +2918,10 @@ def write(*, params, outlet_events): with dag_maker(dag_id="read", schedule=None, session=session): - @task(inlets=Dataset(ds_uri)) + @task(inlets=Asset(asset_uri)) def read(*, inlet_events): nonlocal result - result = [e.extra for e in slicer(inlet_events[ds_uri])] + result = [e.extra for e in slicer(inlet_events[asset_uri])] read() @@ -2933,23 +2946,23 @@ def read(*, inlet_events): (lambda x: x[-5:5], []), ], ) - def test_inlet_dataset_alias_extra_slice(self, dag_maker, session, slicer, expected): - ds_uri = "test_inlet_dataset_alias_extra_slice_ds" - dsa_name = "test_inlet_dataset_alias_extra_slice_dsa" - - ds_model = DatasetModel(id=1, uri=ds_uri) - dsa_model = DatasetAliasModel(name=dsa_name) - dsa_model.datasets.append(ds_model) - session.add_all([ds_model, dsa_model]) + def test_inlet_asset_alias_extra_slice(self, dag_maker, session, slicer, expected): + asset_uri = "test_inlet_asset_alias_extra_slice_ds" + asset_alias_name = "test_inlet_asset_alias_extra_slice_asset_alias" + + asset_model = AssetModel(id=1, uri=asset_uri) + asset_alias_model = AssetAliasModel(name=asset_alias_name) + asset_alias_model.datasets.append(asset_model) + session.add_all([asset_model, asset_alias_model]) session.commit() - from airflow.datasets import Dataset + from airflow.assets import Asset with dag_maker(dag_id="write", schedule="@daily", params={"i": -1}, session=session): - @task(outlets=DatasetAlias(dsa_name)) + @task(outlets=AssetAlias(asset_alias_name)) def write(*, params, outlet_events): - outlet_events[dsa_name].add(Dataset(ds_uri), {"from": params["i"]}) + outlet_events[asset_alias_name].add(Asset(asset_uri), {"from": params["i"]}) write() @@ -2966,10 +2979,10 @@ def write(*, params, outlet_events): with dag_maker(dag_id="read", schedule=None, session=session): - @task(inlets=DatasetAlias(dsa_name)) + @task(inlets=AssetAlias(asset_alias_name)) def read(*, inlet_events): nonlocal result - result = [e.extra for e in slicer(inlet_events[DatasetAlias(dsa_name)])] + result = [e.extra for e in slicer(inlet_events[AssetAlias(asset_alias_name)])] read() @@ -2983,16 +2996,16 @@ def read(*, inlet_events): assert result == expected @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode - def test_changing_of_dataset_when_ddrq_is_already_populated(self, dag_maker): + def test_changing_of_asset_when_adrq_is_already_populated(self, dag_maker): """ - Test that when a task that produces dataset has ran, that changing the consumer - dag dataset will not cause primary key blank-out + Test that when a task that produces asset has ran, that changing the consumer + dag asset will not cause primary key blank-out """ - from airflow.datasets import Dataset + from airflow.assets import Asset with dag_maker(schedule=None, serialized=True) as dag1: - @task(outlets=Dataset("test/1")) + @task(outlets=Asset("test/1")) def test_task1(): print(1) @@ -3001,7 +3014,7 @@ def test_task1(): dr1 = dag_maker.create_dagrun() test_task1 = dag1.get_task("test_task1") - with dag_maker(dag_id="testdag", schedule=[Dataset("test/1")], serialized=True): + with dag_maker(dag_id="testdag", schedule=[Asset("test/1")], serialized=True): @task def test_task2(): @@ -3011,8 +3024,8 @@ def test_task2(): ti = dr1.get_task_instance(task_id="test_task1") ti.run() - # Change the dataset. - with dag_maker(dag_id="testdag", schedule=[Dataset("test2/1")], serialized=True): + # Change the asset. + with dag_maker(dag_id="testdag", schedule=[Asset("test2/1")], serialized=True): @task def test_task2(): @@ -3153,19 +3166,19 @@ def test_get_previous_start_date_none(self, dag_maker): assert ti_1.start_date is None @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode - def test_context_triggering_dataset_events_none(self, session, create_task_instance): + def test_context_triggering_asset_events_none(self, session, create_task_instance): ti = create_task_instance() template_context = ti.get_template_context() assert ti in session session.expunge_all() - assert template_context["triggering_dataset_events"] == {} + assert template_context["triggering_asset_events"] == {} @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode - def test_context_triggering_dataset_events(self, create_dummy_dag, session): - ds1 = DatasetModel(id=1, uri="one") - ds2 = DatasetModel(id=2, uri="two") + def test_context_triggering_asset_events(self, create_dummy_dag, session): + ds1 = AssetModel(id=1, uri="one") + ds2 = AssetModel(id=2, uri="two") session.add_all([ds1, ds2]) session.commit() @@ -3173,7 +3186,7 @@ def test_context_triggering_dataset_events(self, create_dummy_dag, session): triggered_by_kwargs = {"triggered_by": DagRunTriggeredByType.TEST} if AIRFLOW_V_3_0_PLUS else {} # it's easier to fake a manual run here dag, task1 = create_dummy_dag( - dag_id="test_triggering_dataset_events", + dag_id="test_triggering_asset_events", schedule=None, start_date=DEFAULT_DATE, task_id="test_context", @@ -3189,9 +3202,9 @@ def test_context_triggering_dataset_events(self, create_dummy_dag, session): data_interval=(execution_date, execution_date), **triggered_by_kwargs, ) - ds1_event = DatasetEvent(dataset_id=1) - ds2_event_1 = DatasetEvent(dataset_id=2) - ds2_event_2 = DatasetEvent(dataset_id=2) + ds1_event = AssetEvent(dataset_id=1) + ds2_event_1 = AssetEvent(dataset_id=2) + ds2_event_2 = AssetEvent(dataset_id=2) dr.consumed_dataset_events.append(ds1_event) dr.consumed_dataset_events.append(ds2_event_1) dr.consumed_dataset_events.append(ds2_event_2) @@ -3207,7 +3220,7 @@ def test_context_triggering_dataset_events(self, create_dummy_dag, session): template_context = ti.get_template_context() - assert template_context["triggering_dataset_events"] == { + assert template_context["triggering_asset_events"] == { "one": [ds1_event], "two": [ds2_event_1, ds2_event_2], } @@ -4187,7 +4200,7 @@ def _clean(): db.clear_db_dags() db.clear_db_sla_miss() db.clear_db_import_errors() - db.clear_db_datasets() + db.clear_db_assets() def setup_method(self) -> None: self._clean() diff --git a/tests/operators/test_python.py b/tests/operators/test_python.py index 28c66511a3d38..c98dc7018a4a1 100644 --- a/tests/operators/test_python.py +++ b/tests/operators/test_python.py @@ -894,8 +894,8 @@ def test_virtualenv_serializable_context_fields(self, create_task_instance): "ti", "var", # Accessor for Variable; var->json and var->value. "conn", # Accessor for Connection. - "inlet_events", # Accessor for inlet DatasetEvent. - "outlet_events", # Accessor for outlet DatasetEvent. + "inlet_events", # Accessor for inlet AssetEvent. + "outlet_events", # Accessor for outlet AssetEvent. ] ti = create_task_instance(dag_id=self.dag_id, task_id=self.task_id, schedule=None) diff --git a/tests/providers/mysql/datasets/__init__.py b/tests/providers/amazon/aws/assets/__init__.py similarity index 100% rename from tests/providers/mysql/datasets/__init__.py rename to tests/providers/amazon/aws/assets/__init__.py diff --git a/tests/providers/amazon/aws/datasets/test_s3.py b/tests/providers/amazon/aws/assets/test_s3.py similarity index 75% rename from tests/providers/amazon/aws/datasets/test_s3.py rename to tests/providers/amazon/aws/assets/test_s3.py index 893d6acf677bc..e918c9fdffa2f 100644 --- a/tests/providers/amazon/aws/datasets/test_s3.py +++ b/tests/providers/amazon/aws/assets/test_s3.py @@ -20,10 +20,10 @@ import pytest -from airflow.datasets import Dataset -from airflow.providers.amazon.aws.datasets.s3 import ( - convert_dataset_to_openlineage, - create_dataset, +from airflow.providers.amazon.aws.assets.s3 import ( + Asset, + convert_asset_to_openlineage, + create_asset, sanitize_uri, ) from airflow.providers.amazon.aws.hooks.s3 import S3Hook @@ -50,9 +50,9 @@ def test_sanitize_uri_no_path(): assert result.path == "" -def test_create_dataset(): - assert create_dataset(bucket="test-bucket", key="test-path") == Dataset(uri="s3://test-bucket/test-path") - assert create_dataset(bucket="test-bucket", key="test-dir/test-path") == Dataset( +def test_create_asset(): + assert create_asset(bucket="test-bucket", key="test-path") == Asset(uri="s3://test-bucket/test-path") + assert create_asset(bucket="test-bucket", key="test-dir/test-path") == Asset( uri="s3://test-bucket/test-dir/test-path" ) @@ -65,15 +65,15 @@ def test_sanitize_uri_trailing_slash(): assert result.path == "/" -def test_convert_dataset_to_openlineage_valid(): +def test_convert_asset_to_openlineage_valid(): uri = "s3://bucket/dir/file.txt" - ol_dataset = convert_dataset_to_openlineage(dataset=Dataset(uri=uri), lineage_context=S3Hook()) + ol_dataset = convert_asset_to_openlineage(asset=Asset(uri=uri), lineage_context=S3Hook()) assert ol_dataset.namespace == "s3://bucket" assert ol_dataset.name == "dir/file.txt" @pytest.mark.parametrize("uri", ("s3://bucket", "s3://bucket/")) -def test_convert_dataset_to_openlineage_no_path(uri): - ol_dataset = convert_dataset_to_openlineage(dataset=Dataset(uri=uri), lineage_context=S3Hook()) +def test_convert_asset_to_openlineage_no_path(uri): + ol_dataset = convert_asset_to_openlineage(asset=Asset(uri=uri), lineage_context=S3Hook()) assert ol_dataset.namespace == "s3://bucket" assert ol_dataset.name == "/" diff --git a/tests/providers/amazon/aws/auth_manager/test_aws_auth_manager.py b/tests/providers/amazon/aws/auth_manager/test_aws_auth_manager.py index f54a2a3e5fb1f..d827ba3ff0e6d 100644 --- a/tests/providers/amazon/aws/auth_manager/test_aws_auth_manager.py +++ b/tests/providers/amazon/aws/auth_manager/test_aws_auth_manager.py @@ -23,10 +23,24 @@ from flask import Flask, session from flask_appbuilder.menu import MenuItem +from airflow.providers.amazon.aws.auth_manager.avp.entities import AvpEntities +from airflow.providers.amazon.aws.auth_manager.avp.facade import AwsAuthManagerAmazonVerifiedPermissionsFacade +from airflow.providers.amazon.aws.auth_manager.aws_auth_manager import AwsAuthManager from airflow.providers.amazon.aws.auth_manager.security_manager.aws_security_manager_override import ( AwsSecurityManagerOverride, ) +from airflow.providers.amazon.aws.auth_manager.user import AwsAuthManagerUser +from airflow.security.permissions import ( + RESOURCE_AUDIT_LOG, + RESOURCE_CLUSTER_ACTIVITY, + RESOURCE_CONNECTION, + RESOURCE_VARIABLE, +) +from airflow.www import app as application +from airflow.www.extensions.init_appbuilder import init_appbuilder from tests.test_utils.compat import AIRFLOW_V_2_8_PLUS, AIRFLOW_V_2_9_PLUS +from tests.test_utils.config import conf_vars +from tests.test_utils.www import check_content_in_response try: from airflow.auth.managers.models.resource_details import ( @@ -35,7 +49,6 @@ ConnectionDetails, DagAccessEntity, DagDetails, - DatasetDetails, PoolDetails, VariableDetails, ) @@ -47,24 +60,18 @@ ) else: raise -from airflow.providers.amazon.aws.auth_manager.avp.entities import AvpEntities -from airflow.providers.amazon.aws.auth_manager.avp.facade import AwsAuthManagerAmazonVerifiedPermissionsFacade -from airflow.providers.amazon.aws.auth_manager.aws_auth_manager import AwsAuthManager -from airflow.providers.amazon.aws.auth_manager.user import AwsAuthManagerUser -from airflow.security.permissions import ( - RESOURCE_AUDIT_LOG, - RESOURCE_CLUSTER_ACTIVITY, - RESOURCE_CONNECTION, - RESOURCE_DATASET, - RESOURCE_VARIABLE, -) -from airflow.www import app as application -from airflow.www.extensions.init_appbuilder import init_appbuilder -from tests.test_utils.config import conf_vars -from tests.test_utils.www import check_content_in_response if TYPE_CHECKING: from airflow.auth.managers.base_auth_manager import ResourceMethod + from airflow.auth.managers.models.resource_details import AssetDetails + from airflow.security.permissions import RESOURCE_ASSET +else: + try: + from airflow.auth.managers.models.resource_details import AssetDetails + from airflow.security.permissions import RESOURCE_ASSET + except ImportError: + from airflow.auth.managers.models.resource_details import DatasetDetails as AssetDetails + from airflow.security.permissions import RESOURCE_DATASET as RESOURCE_ASSET pytestmark = [ pytest.mark.skipif(not AIRFLOW_V_2_9_PLUS, reason="Test requires Airflow 2.9+"), @@ -324,12 +331,12 @@ def test_is_authorized_dag( "details, user, expected_user, expected_entity_id", [ (None, None, ANY, None), - (DatasetDetails(uri="uri"), mock, mock, "uri"), + (AssetDetails(uri="uri"), mock, mock, "uri"), ], ) @patch.object(AwsAuthManager, "avp_facade") @patch.object(AwsAuthManager, "get_user") - def test_is_authorized_dataset( + def test_is_authorized_asset( self, mock_get_user, mock_avp_facade, @@ -343,12 +350,12 @@ def test_is_authorized_dataset( mock_avp_facade.is_authorized = is_authorized method: ResourceMethod = "GET" - result = auth_manager.is_authorized_dataset(method=method, details=details, user=user) + result = auth_manager.is_authorized_asset(method=method, details=details, user=user) if not user: mock_get_user.assert_called_once() is_authorized.assert_called_once_with( - method=method, entity_type=AvpEntities.DATASET, user=expected_user, entity_id=expected_entity_id + method=method, entity_type=AvpEntities.ASSET, user=expected_user, entity_id=expected_entity_id ) assert result @@ -611,7 +618,7 @@ def test_filter_permitted_menu_items(self, mock_get_user, auth_manager, test_use "request": { "principal": {"entityType": "Airflow::User", "entityId": "test_user_id"}, "action": {"actionType": "Airflow::Action", "actionId": "Menu.MENU"}, - "resource": {"entityType": "Airflow::Menu", "entityId": "Datasets"}, + "resource": {"entityType": "Airflow::Menu", "entityId": RESOURCE_ASSET}, }, "decision": "DENY", }, @@ -649,7 +656,7 @@ def test_filter_permitted_menu_items(self, mock_get_user, auth_manager, test_use result = auth_manager.filter_permitted_menu_items( [ MenuItem("Category1", childs=[MenuItem(RESOURCE_CONNECTION), MenuItem(RESOURCE_VARIABLE)]), - MenuItem("Category2", childs=[MenuItem(RESOURCE_DATASET)]), + MenuItem("Category2", childs=[MenuItem(RESOURCE_ASSET)]), MenuItem(RESOURCE_CLUSTER_ACTIVITY), MenuItem(RESOURCE_AUDIT_LOG), MenuItem("CustomPage"), @@ -679,7 +686,7 @@ def test_filter_permitted_menu_items(self, mock_get_user, auth_manager, test_use { "method": "MENU", "entity_type": AvpEntities.MENU, - "entity_id": "Datasets", + "entity_id": RESOURCE_ASSET, }, {"method": "MENU", "entity_type": AvpEntities.MENU, "entity_id": "Cluster Activity"}, {"method": "MENU", "entity_type": AvpEntities.MENU, "entity_id": "Audit Logs"}, diff --git a/tests/providers/amazon/aws/hooks/test_s3.py b/tests/providers/amazon/aws/hooks/test_s3.py index 97696c64b6e7a..43c4b94445b6f 100644 --- a/tests/providers/amazon/aws/hooks/test_s3.py +++ b/tests/providers/amazon/aws/hooks/test_s3.py @@ -31,9 +31,9 @@ from botocore.exceptions import ClientError from moto import mock_aws -from airflow.datasets import Dataset from airflow.exceptions import AirflowException from airflow.models import Connection +from airflow.providers.amazon.aws.assets.s3 import Asset from airflow.providers.amazon.aws.exceptions import S3HookUriParseFailure from airflow.providers.amazon.aws.hooks.s3 import ( NO_ACL, @@ -58,6 +58,21 @@ def s3_bucket(mocked_s3_res): return bucket +if AIRFLOW_V_2_10_PLUS: + + @pytest.fixture + def hook_lineage_collector(): + from airflow.lineage import hook + from airflow.providers.amazon.aws.hooks.s3 import get_hook_lineage_collector + + hook._hook_lineage_collector = None + hook._hook_lineage_collector = hook.HookLineageCollector() + + yield get_hook_lineage_collector() + + hook._hook_lineage_collector = None + + class TestAwsS3Hook: @mock_aws def test_get_conn(self): @@ -429,9 +444,10 @@ def test_load_string(self, s3_bucket): @pytest.mark.skipif(not AIRFLOW_V_2_10_PLUS, reason="Hook lineage works in Airflow >= 2.10.0") def test_load_string_exposes_lineage(self, s3_bucket, hook_lineage_collector): hook = S3Hook() + hook.load_string("Contént", "my_key", s3_bucket) - assert len(hook_lineage_collector.collected_datasets.outputs) == 1 - assert hook_lineage_collector.collected_datasets.outputs[0].dataset == Dataset( + assert len(hook_lineage_collector.collected_assets.outputs) == 1 + assert hook_lineage_collector.collected_assets.outputs[0].asset == Asset( uri=f"s3://{s3_bucket}/my_key" ) @@ -1023,8 +1039,8 @@ def test_load_file_exposes_lineage(self, s3_bucket, tmp_path, hook_lineage_colle path = tmp_path / "testfile" path.write_text("Content") hook.load_file(path, "my_key", s3_bucket) - assert len(hook_lineage_collector.collected_datasets.outputs) == 1 - assert hook_lineage_collector.collected_datasets.outputs[0].dataset == Dataset( + assert len(hook_lineage_collector.collected_assets.outputs) == 1 + assert hook_lineage_collector.collected_assets.outputs[0].asset == Asset( uri=f"s3://{s3_bucket}/my_key" ) @@ -1095,13 +1111,13 @@ def test_copy_object_ol_instrumentation(self, s3_bucket, hook_lineage_collector) "get_conn", ): mock_hook.copy_object("my_key", "my_key3", s3_bucket, s3_bucket) - assert len(hook_lineage_collector.collected_datasets.inputs) == 1 - assert hook_lineage_collector.collected_datasets.inputs[0].dataset == Dataset( + assert len(hook_lineage_collector.collected_assets.inputs) == 1 + assert hook_lineage_collector.collected_assets.inputs[0].asset == Asset( uri=f"s3://{s3_bucket}/my_key" ) - assert len(hook_lineage_collector.collected_datasets.outputs) == 1 - assert hook_lineage_collector.collected_datasets.outputs[0].dataset == Dataset( + assert len(hook_lineage_collector.collected_assets.outputs) == 1 + assert hook_lineage_collector.collected_assets.outputs[0].asset == Asset( uri=f"s3://{s3_bucket}/my_key3" ) @@ -1233,8 +1249,8 @@ def test_download_file_exposes_lineage(self, mock_temp_file, tmp_path, hook_line s3_hook.download_file(key=key, bucket_name=bucket) - assert len(hook_lineage_collector.collected_datasets.inputs) == 1 - assert hook_lineage_collector.collected_datasets.inputs[0].dataset == Dataset( + assert len(hook_lineage_collector.collected_assets.inputs) == 1 + assert hook_lineage_collector.collected_assets.inputs[0].asset == Asset( uri="s3://test_bucket/test_key" ) @@ -1285,14 +1301,14 @@ def test_download_file_with_preserve_name_exposes_lineage( use_autogenerated_subdir=False, ) - assert len(hook_lineage_collector.collected_datasets.inputs) == 1 - assert hook_lineage_collector.collected_datasets.inputs[0].dataset == Dataset( - uri="s3://test_bucket/test_key/test.log" + assert len(hook_lineage_collector.collected_assets.inputs) == 1 + assert hook_lineage_collector.collected_assets.inputs[0].asset == Asset( + uri="s3://test_bucket/test_key/test.log", extra={} ) - assert len(hook_lineage_collector.collected_datasets.outputs) == 1 - assert hook_lineage_collector.collected_datasets.outputs[0].dataset == Dataset( - uri=f"file://{local_path}/test.log", + assert len(hook_lineage_collector.collected_assets.outputs) == 1 + assert hook_lineage_collector.collected_assets.outputs[0].asset == Asset( + uri=f"file://{local_path}/test.log", extra={} ) @mock.patch("airflow.providers.amazon.aws.hooks.s3.open") diff --git a/tests/providers/postgres/datasets/__init__.py b/tests/providers/common/compat/openlineage/utils/__init__.py similarity index 100% rename from tests/providers/postgres/datasets/__init__.py rename to tests/providers/common/compat/openlineage/utils/__init__.py diff --git a/tests/providers/common/compat/openlineage/utils/test_utils.py b/tests/providers/common/compat/openlineage/utils/test_utils.py new file mode 100644 index 0000000000000..72af469a1becd --- /dev/null +++ b/tests/providers/common/compat/openlineage/utils/test_utils.py @@ -0,0 +1,23 @@ +# 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 + + +def test_import(): + from airflow.providers.common.compat.openlineage.utils.utils import translate_airflow_asset + + assert translate_airflow_asset is not None diff --git a/tests/providers/trino/datasets/__init__.py b/tests/providers/common/compat/security/__init__.py similarity index 100% rename from tests/providers/trino/datasets/__init__.py rename to tests/providers/common/compat/security/__init__.py diff --git a/tests/providers/common/compat/security/test_permissions.py b/tests/providers/common/compat/security/test_permissions.py new file mode 100644 index 0000000000000..40a13832f9e25 --- /dev/null +++ b/tests/providers/common/compat/security/test_permissions.py @@ -0,0 +1,23 @@ +# 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 + + +def test_import(): + from airflow.providers.common.compat.security.permissions import RESOURCE_ASSET + + assert RESOURCE_ASSET is not None diff --git a/tests/providers/common/io/assets/__init__.py b/tests/providers/common/io/assets/__init__.py new file mode 100644 index 0000000000000..13a83393a9124 --- /dev/null +++ b/tests/providers/common/io/assets/__init__.py @@ -0,0 +1,16 @@ +# 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. diff --git a/tests/providers/common/io/datasets/test_file.py b/tests/providers/common/io/assets/test_file.py similarity index 83% rename from tests/providers/common/io/datasets/test_file.py rename to tests/providers/common/io/assets/test_file.py index d8d53247a6796..21357f933fde8 100644 --- a/tests/providers/common/io/datasets/test_file.py +++ b/tests/providers/common/io/assets/test_file.py @@ -20,11 +20,11 @@ import pytest -from airflow.datasets import Dataset from airflow.providers.common.compat.openlineage.facet import Dataset as OpenLineageDataset -from airflow.providers.common.io.datasets.file import ( - convert_dataset_to_openlineage, - create_dataset, +from airflow.providers.common.io.assets.file import ( + Asset, + convert_asset_to_openlineage, + create_asset, sanitize_uri, ) @@ -47,8 +47,8 @@ def test_sanitize_uri_invalid(uri): sanitize_uri(urlsplit(uri)) -def test_file_dataset(): - assert create_dataset(path="/asdf/fdsa") == Dataset(uri="file:///asdf/fdsa") +def test_file_asset(): + assert create_asset(path="/asdf/fdsa") == Asset(uri="file:///asdf/fdsa") @pytest.mark.parametrize( @@ -62,6 +62,6 @@ def test_file_dataset(): ("file:///C://dir/file", OpenLineageDataset(namespace="file://", name="/C://dir/file")), ), ) -def test_convert_dataset_to_openlineage(uri, ol_dataset): - result = convert_dataset_to_openlineage(Dataset(uri=uri), None) +def test_convert_asset_to_openlineage(uri, ol_dataset): + result = convert_asset_to_openlineage(Asset(uri=uri), None) assert result == ol_dataset diff --git a/tests/providers/fab/auth_manager/test_fab_auth_manager.py b/tests/providers/fab/auth_manager/test_fab_auth_manager.py index b755afcc70d03..d727b6090822f 100644 --- a/tests/providers/fab/auth_manager/test_fab_auth_manager.py +++ b/tests/providers/fab/auth_manager/test_fab_auth_manager.py @@ -48,7 +48,6 @@ RESOURCE_CONNECTION, RESOURCE_DAG, RESOURCE_DAG_RUN, - RESOURCE_DATASET, RESOURCE_DOCS, RESOURCE_JOB, RESOURCE_PLUGIN, @@ -62,11 +61,20 @@ if TYPE_CHECKING: from airflow.auth.managers.base_auth_manager import ResourceMethod + from airflow.security.permissions import RESOURCE_ASSET +else: + try: + from airflow.security.permissions import RESOURCE_ASSET + except ImportError: + from airflow.security.permissions import ( + RESOURCE_DATASET as RESOURCE_ASSET, + ) + IS_AUTHORIZED_METHODS_SIMPLE = { "is_authorized_configuration": RESOURCE_CONFIG, "is_authorized_connection": RESOURCE_CONNECTION, - "is_authorized_dataset": RESOURCE_DATASET, + "is_authorized_asset": RESOURCE_ASSET, "is_authorized_variable": RESOURCE_VARIABLE, } diff --git a/tests/providers/fab/auth_manager/test_security.py b/tests/providers/fab/auth_manager/test_security.py index b6aca2d4513a5..156b5cf626271 100644 --- a/tests/providers/fab/auth_manager/test_security.py +++ b/tests/providers/fab/auth_manager/test_security.py @@ -22,6 +22,7 @@ import json import logging import os +from typing import TYPE_CHECKING from unittest import mock from unittest.mock import patch @@ -60,6 +61,15 @@ from tests.test_utils.mock_security_manager import MockSecurityManager from tests.test_utils.permissions import _resource_name +if TYPE_CHECKING: + from airflow.security.permissions import RESOURCE_ASSET +else: + try: + from airflow.security.permissions import RESOURCE_ASSET + except ImportError: + from airflow.security.permissions import RESOURCE_DATASET as RESOURCE_ASSET + + pytestmark = pytest.mark.db_test READ_WRITE = {permissions.RESOURCE_DAG: {permissions.ACTION_CAN_READ, permissions.ACTION_CAN_EDIT}} @@ -435,7 +445,7 @@ def test_get_user_roles_for_anonymous_user(app, security_manager): (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_DEPENDENCIES), (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_CODE), (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DATASET), + (permissions.ACTION_CAN_READ, RESOURCE_ASSET), (permissions.ACTION_CAN_READ, permissions.RESOURCE_CLUSTER_ACTIVITY), (permissions.ACTION_CAN_READ, permissions.RESOURCE_IMPORT_ERROR), (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_WARNING), @@ -454,7 +464,7 @@ def test_get_user_roles_for_anonymous_user(app, security_manager): (permissions.ACTION_CAN_ACCESS_MENU, permissions.RESOURCE_DAG), (permissions.ACTION_CAN_ACCESS_MENU, permissions.RESOURCE_DAG_DEPENDENCIES), (permissions.ACTION_CAN_ACCESS_MENU, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_ACCESS_MENU, permissions.RESOURCE_DATASET), + (permissions.ACTION_CAN_ACCESS_MENU, RESOURCE_ASSET), (permissions.ACTION_CAN_ACCESS_MENU, permissions.RESOURCE_CLUSTER_ACTIVITY), (permissions.ACTION_CAN_ACCESS_MENU, permissions.RESOURCE_JOB), (permissions.ACTION_CAN_ACCESS_MENU, permissions.RESOURCE_SLA_MISS), diff --git a/tests/providers/mysql/assets/__init__.py b/tests/providers/mysql/assets/__init__.py new file mode 100644 index 0000000000000..13a83393a9124 --- /dev/null +++ b/tests/providers/mysql/assets/__init__.py @@ -0,0 +1,16 @@ +# 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. diff --git a/tests/providers/mysql/datasets/test_mysql.py b/tests/providers/mysql/assets/test_mysql.py similarity index 97% rename from tests/providers/mysql/datasets/test_mysql.py rename to tests/providers/mysql/assets/test_mysql.py index 5f31d72991f27..28e44558f31d6 100644 --- a/tests/providers/mysql/datasets/test_mysql.py +++ b/tests/providers/mysql/assets/test_mysql.py @@ -21,7 +21,7 @@ import pytest -from airflow.providers.mysql.datasets.mysql import sanitize_uri +from airflow.providers.mysql.assets.mysql import sanitize_uri @pytest.mark.parametrize( diff --git a/tests/providers/openlineage/extractors/test_manager.py b/tests/providers/openlineage/extractors/test_manager.py index 479347179bd17..601a456604843 100644 --- a/tests/providers/openlineage/extractors/test_manager.py +++ b/tests/providers/openlineage/extractors/test_manager.py @@ -25,7 +25,6 @@ from openlineage.client.event_v2 import Dataset as OpenLineageDataset from openlineage.client.facet_v2 import documentation_dataset, ownership_dataset, schema_dataset -from airflow.datasets import Dataset from airflow.io.path import ObjectStoragePath from airflow.lineage.entities import Column, File, Table, User from airflow.models.baseoperator import BaseOperator @@ -33,12 +32,36 @@ from airflow.operators.python import PythonOperator from airflow.providers.openlineage.extractors import OperatorLineage from airflow.providers.openlineage.extractors.manager import ExtractorManager +from airflow.providers.openlineage.utils.utils import Asset from airflow.utils.state import State from tests.test_utils.compat import AIRFLOW_V_2_10_PLUS if TYPE_CHECKING: from airflow.utils.context import Context +if AIRFLOW_V_2_10_PLUS: + + @pytest.fixture + def hook_lineage_collector(): + from importlib.util import find_spec + + from airflow.lineage import hook + + if find_spec("airflow.assets"): + # Dataset has been renamed as Asset in 3.0 + from airflow.lineage.hook import get_hook_lineage_collector + else: + from airflow.providers.openlineage.utils.asset_compat_lineage_collector import ( + get_hook_lineage_collector, + ) + + hook._hook_lineage_collector = None + hook._hook_lineage_collector = hook.HookLineageCollector() + + yield get_hook_lineage_collector() + + hook._hook_lineage_collector = None + @pytest.mark.parametrize( ("uri", "dataset"), @@ -213,8 +236,8 @@ def test_extractor_manager_uses_hook_level_lineage(hook_lineage_collector): del task.get_openlineage_facets_on_complete ti = MagicMock() - hook_lineage_collector.add_input_dataset(None, uri="s3://bucket/input_key") - hook_lineage_collector.add_output_dataset(None, uri="s3://bucket/output_key") + hook_lineage_collector.add_input_asset(None, uri="s3://bucket/input_key") + hook_lineage_collector.add_output_asset(None, uri="s3://bucket/output_key") extractor_manager = ExtractorManager() metadata = extractor_manager.extract_metadata(dagrun=dagrun, task=task, complete=True, task_instance=ti) @@ -236,7 +259,7 @@ def get_openlineage_facets_on_start(self): dagrun = MagicMock() task = FakeSupportedOperator(task_id="test_task_extractor") ti = MagicMock() - hook_lineage_collector.add_input_dataset(None, uri="s3://bucket/input_key") + hook_lineage_collector.add_input_asset(None, uri="s3://bucket/input_key") extractor_manager = ExtractorManager() metadata = extractor_manager.extract_metadata(dagrun=dagrun, task=task, complete=True, task_instance=ti) @@ -269,7 +292,7 @@ def use_read(): ti.run() - datasets = hook_lineage_collector.collected_datasets + datasets = hook_lineage_collector.collected_assets assert len(datasets.outputs) == 1 - assert datasets.outputs[0].dataset == Dataset(uri=path) + assert datasets.outputs[0].asset == Asset(uri=path) diff --git a/tests/providers/postgres/assets/__init__.py b/tests/providers/postgres/assets/__init__.py new file mode 100644 index 0000000000000..13a83393a9124 --- /dev/null +++ b/tests/providers/postgres/assets/__init__.py @@ -0,0 +1,16 @@ +# 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. diff --git a/tests/providers/postgres/datasets/test_postgres.py b/tests/providers/postgres/assets/test_postgres.py similarity index 97% rename from tests/providers/postgres/datasets/test_postgres.py rename to tests/providers/postgres/assets/test_postgres.py index 40d6bf11d235d..82c64759a290a 100644 --- a/tests/providers/postgres/datasets/test_postgres.py +++ b/tests/providers/postgres/assets/test_postgres.py @@ -21,7 +21,7 @@ import pytest -from airflow.providers.postgres.datasets.postgres import sanitize_uri +from airflow.providers.postgres.assets.postgres import sanitize_uri @pytest.mark.parametrize( diff --git a/tests/providers/trino/assets/__init__.py b/tests/providers/trino/assets/__init__.py new file mode 100644 index 0000000000000..13a83393a9124 --- /dev/null +++ b/tests/providers/trino/assets/__init__.py @@ -0,0 +1,16 @@ +# 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. diff --git a/tests/providers/trino/datasets/test_trino.py b/tests/providers/trino/assets/test_trino.py similarity index 97% rename from tests/providers/trino/datasets/test_trino.py rename to tests/providers/trino/assets/test_trino.py index 12cacd4eb0cf2..4ebea16d9fc6b 100644 --- a/tests/providers/trino/datasets/test_trino.py +++ b/tests/providers/trino/assets/test_trino.py @@ -21,7 +21,7 @@ import pytest -from airflow.providers.trino.datasets.trino import sanitize_uri +from airflow.providers.trino.assets.trino import sanitize_uri @pytest.mark.parametrize( diff --git a/tests/serialization/test_dag_serialization.py b/tests/serialization/test_dag_serialization.py index 758c7f496ed93..6910514776afe 100644 --- a/tests/serialization/test_dag_serialization.py +++ b/tests/serialization/test_dag_serialization.py @@ -42,7 +42,7 @@ from kubernetes.client import models as k8s import airflow -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.decorators import teardown from airflow.decorators.base import DecoratedOperator from airflow.exceptions import ( @@ -1656,16 +1656,16 @@ class DerivedSensor(ExternalTaskSensor): ] @pytest.mark.db_test - def test_dag_deps_datasets_with_duplicate_dataset(self): + def test_dag_deps_assets_with_duplicate_asset(self): """ - Check that dag_dependencies node is populated correctly for a DAG with duplicate datasets. + Check that dag_dependencies node is populated correctly for a DAG with duplicate assets. """ from airflow.sensors.external_task import ExternalTaskSensor - d1 = Dataset("d1") - d2 = Dataset("d2") - d3 = Dataset("d3") - d4 = Dataset("d4") + d1 = Asset("d1") + d2 = Asset("d2") + d3 = Asset("d3") + d4 = Asset("d4") execution_date = datetime(2020, 1, 1) with DAG(dag_id="test", start_date=execution_date, schedule=[d1, d1, d1, d1, d1]) as dag: ExternalTaskSensor( @@ -1673,13 +1673,13 @@ def test_dag_deps_datasets_with_duplicate_dataset(self): external_dag_id="external_dag_id", mode="reschedule", ) - BashOperator(task_id="dataset_writer", bash_command="echo hello", outlets=[d2, d2, d2, d3]) + BashOperator(task_id="asset_writer", bash_command="echo hello", outlets=[d2, d2, d2, d3]) @dag.task(outlets=[d4]) - def other_dataset_writer(x): + def other_asset_writer(x): pass - other_dataset_writer.expand(x=[1, 2]) + other_asset_writer.expand(x=[1, 2]) dag = SerializedDAG.to_dict(dag) actual = sorted(dag["dag"]["dag_dependencies"], key=lambda x: tuple(x.values())) @@ -1687,8 +1687,8 @@ def other_dataset_writer(x): [ { "source": "test", - "target": "dataset", - "dependency_type": "dataset", + "target": "asset", + "dependency_type": "asset", "dependency_id": "d4", }, { @@ -1699,44 +1699,44 @@ def other_dataset_writer(x): }, { "source": "test", - "target": "dataset", - "dependency_type": "dataset", + "target": "asset", + "dependency_type": "asset", "dependency_id": "d3", }, { "source": "test", - "target": "dataset", - "dependency_type": "dataset", + "target": "asset", + "dependency_type": "asset", "dependency_id": "d2", }, { - "source": "dataset", + "source": "asset", "target": "test", - "dependency_type": "dataset", + "dependency_type": "asset", "dependency_id": "d1", }, { "dependency_id": "d1", - "dependency_type": "dataset", - "source": "dataset", + "dependency_type": "asset", + "source": "asset", "target": "test", }, { "dependency_id": "d1", - "dependency_type": "dataset", - "source": "dataset", + "dependency_type": "asset", + "source": "asset", "target": "test", }, { "dependency_id": "d1", - "dependency_type": "dataset", - "source": "dataset", + "dependency_type": "asset", + "source": "asset", "target": "test", }, { "dependency_id": "d1", - "dependency_type": "dataset", - "source": "dataset", + "dependency_type": "asset", + "source": "asset", "target": "test", }, ], @@ -1745,16 +1745,16 @@ def other_dataset_writer(x): assert actual == expected @pytest.mark.db_test - def test_dag_deps_datasets(self): + def test_dag_deps_assets(self): """ - Check that dag_dependencies node is populated correctly for a DAG with datasets. + Check that dag_dependencies node is populated correctly for a DAG with assets. """ from airflow.sensors.external_task import ExternalTaskSensor - d1 = Dataset("d1") - d2 = Dataset("d2") - d3 = Dataset("d3") - d4 = Dataset("d4") + d1 = Asset("d1") + d2 = Asset("d2") + d3 = Asset("d3") + d4 = Asset("d4") execution_date = datetime(2020, 1, 1) with DAG(dag_id="test", start_date=execution_date, schedule=[d1]) as dag: ExternalTaskSensor( @@ -1762,13 +1762,13 @@ def test_dag_deps_datasets(self): external_dag_id="external_dag_id", mode="reschedule", ) - BashOperator(task_id="dataset_writer", bash_command="echo hello", outlets=[d2, d3]) + BashOperator(task_id="asset_writer", bash_command="echo hello", outlets=[d2, d3]) @dag.task(outlets=[d4]) - def other_dataset_writer(x): + def other_asset_writer(x): pass - other_dataset_writer.expand(x=[1, 2]) + other_asset_writer.expand(x=[1, 2]) dag = SerializedDAG.to_dict(dag) actual = sorted(dag["dag"]["dag_dependencies"], key=lambda x: tuple(x.values())) @@ -1776,8 +1776,8 @@ def other_dataset_writer(x): [ { "source": "test", - "target": "dataset", - "dependency_type": "dataset", + "target": "asset", + "dependency_type": "asset", "dependency_id": "d4", }, { @@ -1788,20 +1788,20 @@ def other_dataset_writer(x): }, { "source": "test", - "target": "dataset", - "dependency_type": "dataset", + "target": "asset", + "dependency_type": "asset", "dependency_id": "d3", }, { "source": "test", - "target": "dataset", - "dependency_type": "dataset", + "target": "asset", + "dependency_type": "asset", "dependency_id": "d2", }, { - "source": "dataset", + "source": "asset", "target": "test", - "dependency_type": "dataset", + "dependency_type": "asset", "dependency_id": "d1", }, ], diff --git a/tests/serialization/test_pydantic_models.py b/tests/serialization/test_pydantic_models.py index d29cba23d86e7..55c61ea220263 100644 --- a/tests/serialization/test_pydantic_models.py +++ b/tests/serialization/test_pydantic_models.py @@ -27,16 +27,16 @@ from airflow.jobs.job import Job from airflow.jobs.local_task_job_runner import LocalTaskJobRunner from airflow.models import MappedOperator -from airflow.models.dag import DAG, DagModel, create_timetable -from airflow.models.dataset import ( - DagScheduleDatasetReference, - DatasetEvent, - DatasetModel, - TaskOutletDatasetReference, +from airflow.models.asset import ( + AssetEvent, + AssetModel, + DagScheduleAssetReference, + TaskOutletAssetReference, ) +from airflow.models.dag import DAG, DagModel, create_timetable +from airflow.serialization.pydantic.asset import AssetEventPydantic from airflow.serialization.pydantic.dag import DagModelPydantic from airflow.serialization.pydantic.dag_run import DagRunPydantic -from airflow.serialization.pydantic.dataset import DatasetEventPydantic from airflow.serialization.pydantic.job import JobPydantic from airflow.serialization.pydantic.taskinstance import TaskInstancePydantic from airflow.serialization.serialized_objects import BaseSerialization @@ -222,15 +222,15 @@ def test_serializing_pydantic_local_task_job(session, create_task_instance): @pytest.mark.skip_if_database_isolation_mode @pytest.mark.skipif(not _ENABLE_AIP_44, reason="AIP-44 is disabled") def test_serializing_pydantic_dataset_event(session, create_task_instance, create_dummy_dag): - ds1 = DatasetModel(id=1, uri="one", extra={"foo": "bar"}) - ds2 = DatasetModel(id=2, uri="two") + ds1 = AssetModel(id=1, uri="one", extra={"foo": "bar"}) + ds2 = AssetModel(id=2, uri="two") session.add_all([ds1, ds2]) session.commit() # it's easier to fake a manual run here dag, task1 = create_dummy_dag( - dag_id="test_triggering_dataset_events", + dag_id="test_triggering_asset_events", schedule=None, start_date=DEFAULT_DATE, task_id="test_context", @@ -250,29 +250,29 @@ def test_serializing_pydantic_dataset_event(session, create_task_instance, creat data_interval=(execution_date, execution_date), **triggered_by_kwargs, ) - ds1_event = DatasetEvent(dataset_id=1) - ds2_event_1 = DatasetEvent(dataset_id=2) - ds2_event_2 = DatasetEvent(dataset_id=2) - - dag_ds_ref = DagScheduleDatasetReference(dag_id=dag.dag_id) - session.add(dag_ds_ref) - dag_ds_ref.dataset = ds1 - task_ds_ref = TaskOutletDatasetReference(task_id=task1.task_id, dag_id=dag.dag_id) + asset1_event = AssetEvent(dataset_id=1) + asset2_event_1 = AssetEvent(dataset_id=2) + asset2_event_2 = AssetEvent(dataset_id=2) + + dag_asset_ref = DagScheduleAssetReference(dag_id=dag.dag_id) + session.add(dag_asset_ref) + dag_asset_ref.dataset = ds1 + task_ds_ref = TaskOutletAssetReference(task_id=task1.task_id, dag_id=dag.dag_id) session.add(task_ds_ref) task_ds_ref.dataset = ds1 - dr.consumed_dataset_events.append(ds1_event) - dr.consumed_dataset_events.append(ds2_event_1) - dr.consumed_dataset_events.append(ds2_event_2) + dr.consumed_dataset_events.append(asset1_event) + dr.consumed_dataset_events.append(asset2_event_1) + dr.consumed_dataset_events.append(asset2_event_2) session.commit() TracebackSessionForTests.set_allow_db_access(session, False) - print(ds2_event_2.dataset.consuming_dags) - pydantic_dse1 = DatasetEventPydantic.model_validate(ds1_event) + print(asset2_event_2.dataset.consuming_dags) + pydantic_dse1 = AssetEventPydantic.model_validate(asset1_event) json_string1 = pydantic_dse1.model_dump_json() print(json_string1) - pydantic_dse2 = DatasetEventPydantic.model_validate(ds2_event_1) + pydantic_dse2 = AssetEventPydantic.model_validate(asset2_event_1) json_string2 = pydantic_dse2.model_dump_json() print(json_string2) @@ -280,13 +280,13 @@ def test_serializing_pydantic_dataset_event(session, create_task_instance, creat json_string_dr = pydantic_dag_run.model_dump_json() print(json_string_dr) - deserialized_model1 = DatasetEventPydantic.model_validate_json(json_string1) + deserialized_model1 = AssetEventPydantic.model_validate_json(json_string1) assert deserialized_model1.dataset.id == 1 assert deserialized_model1.dataset.uri == "one" assert len(deserialized_model1.dataset.consuming_dags) == 1 assert len(deserialized_model1.dataset.producing_tasks) == 1 - deserialized_model2 = DatasetEventPydantic.model_validate_json(json_string2) + deserialized_model2 = AssetEventPydantic.model_validate_json(json_string2) assert deserialized_model2.dataset.id == 2 assert deserialized_model2.dataset.uri == "two" assert len(deserialized_model2.dataset.consuming_dags) == 0 diff --git a/tests/serialization/test_serde.py b/tests/serialization/test_serde.py index cc50e772248a7..a36013d20cfa7 100644 --- a/tests/serialization/test_serde.py +++ b/tests/serialization/test_serde.py @@ -28,7 +28,7 @@ import pytest from pydantic import BaseModel -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.serialization.serde import ( CLASSNAME, DATA, @@ -336,7 +336,7 @@ def test_backwards_compat(self): """ uri = "s3://does/not/exist" data = { - "__type": "airflow.datasets.Dataset", + "__type": "airflow.assets.Asset", "__source": None, "__var": { "__var": { @@ -364,7 +364,7 @@ def test_backwards_compat_wrapped(self): assert e["extra"] == {"hi": "bye"} def test_encode_dataset(self): - dataset = Dataset("mytest://dataset") + dataset = Asset("mytest://dataset") obj = deserialize(serialize(dataset)) assert dataset.uri == obj.uri diff --git a/tests/serialization/test_serialized_objects.py b/tests/serialization/test_serialized_objects.py index 104062ff6c318..0bc8a67ef879a 100644 --- a/tests/serialization/test_serialized_objects.py +++ b/tests/serialization/test_serialized_objects.py @@ -31,7 +31,7 @@ from pendulum.tz.timezone import Timezone from pydantic import BaseModel -from airflow.datasets import Dataset, DatasetAlias, DatasetAliasEvent +from airflow.assets import Asset, AssetAlias, AssetAliasEvent from airflow.exceptions import ( AirflowException, AirflowFailException, @@ -40,10 +40,10 @@ TaskDeferred, ) from airflow.jobs.job import Job +from airflow.models.asset import AssetEvent from airflow.models.connection import Connection from airflow.models.dag import DAG, DagModel, DagTag from airflow.models.dagrun import DagRun -from airflow.models.dataset import DatasetEvent from airflow.models.param import Param from airflow.models.taskinstance import SimpleTaskInstance, TaskInstance from airflow.models.tasklog import LogTemplate @@ -51,9 +51,9 @@ from airflow.operators.empty import EmptyOperator from airflow.operators.python import PythonOperator from airflow.serialization.enums import DagAttributeTypes as DAT, Encoding +from airflow.serialization.pydantic.asset import AssetEventPydantic, AssetPydantic from airflow.serialization.pydantic.dag import DagModelPydantic, DagTagPydantic from airflow.serialization.pydantic.dag_run import DagRunPydantic -from airflow.serialization.pydantic.dataset import DatasetEventPydantic, DatasetPydantic from airflow.serialization.pydantic.job import JobPydantic from airflow.serialization.pydantic.taskinstance import TaskInstancePydantic from airflow.serialization.pydantic.tasklog import LogTemplatePydantic @@ -163,7 +163,7 @@ def equal_exception(a: AirflowException, b: AirflowException) -> bool: def equal_outlet_event_accessor(a: OutletEventAccessor, b: OutletEventAccessor) -> bool: - return a.raw_key == b.raw_key and a.extra == b.extra and a.dataset_alias_events == b.dataset_alias_events + return a.raw_key == b.raw_key and a.extra == b.extra and a.asset_alias_events == b.asset_alias_events class MockLazySelectSequence(LazySelectSequence): @@ -232,7 +232,7 @@ def __len__(self) -> int: None, ), (MockLazySelectSequence(), None, lambda a, b: len(a) == len(b) and isinstance(b, list)), - (Dataset(uri="test"), DAT.DATASET, equals), + (Asset(uri="test"), DAT.ASSET, equals), (SimpleTaskInstance.from_ti(ti=TI), DAT.SIMPLE_TASK_INSTANCE, equals), ( Connection(conn_id="TEST_ID", uri="mysql://"), @@ -240,24 +240,24 @@ def __len__(self) -> int: lambda a, b: a.get_uri() == b.get_uri(), ), ( - OutletEventAccessor(raw_key=Dataset(uri="test"), extra={"key": "value"}, dataset_alias_events=[]), - DAT.DATASET_EVENT_ACCESSOR, + OutletEventAccessor(raw_key=Asset(uri="test"), extra={"key": "value"}, asset_alias_events=[]), + DAT.ASSET_EVENT_ACCESSOR, equal_outlet_event_accessor, ), ( OutletEventAccessor( - raw_key=DatasetAlias(name="test_alias"), + raw_key=AssetAlias(name="test_alias"), extra={"key": "value"}, - dataset_alias_events=[ - DatasetAliasEvent(source_alias_name="test_alias", dest_dataset_uri="test_uri", extra={}) + asset_alias_events=[ + AssetAliasEvent(source_alias_name="test_alias", dest_asset_uri="test_uri", extra={}) ], ), - DAT.DATASET_EVENT_ACCESSOR, + DAT.ASSET_EVENT_ACCESSOR, equal_outlet_event_accessor, ), ( - OutletEventAccessor(raw_key="test", extra={"key": "value"}, dataset_alias_events=[]), - DAT.DATASET_EVENT_ACCESSOR, + OutletEventAccessor(raw_key="test", extra={"key": "value"}, asset_alias_events=[]), + DAT.ASSET_EVENT_ACCESSOR, equal_outlet_event_accessor, ), ( @@ -326,8 +326,8 @@ def test_backcompat_deserialize_connection(conn_uri): id=1, filename="test_file", elasticsearch_id="test_id", created_at=datetime.now() ), DagTagPydantic: DagTag(), - DatasetPydantic: Dataset("uri", {}), - DatasetEventPydantic: DatasetEvent(), + AssetPydantic: Asset("uri", {}), + AssetEventPydantic: AssetEvent(), } @@ -354,14 +354,14 @@ def test_backcompat_deserialize_connection(conn_uri): lambda a, b: equal_time(a.execution_date, b.execution_date) and equal_time(a.start_date, b.start_date), ), - # DataSet is already serialized by non-Pydantic serialization. Is DatasetPydantic needed then? + # Asset is already serialized by non-Pydantic serialization. Is AssetPydantic needed then? # ( - # Dataset( + # Asset( # uri="foo://bar", # extra={"foo": "bar"}, # ), - # DatasetPydantic, - # DAT.DATA_SET, + # AssetPydantic, + # DAT.ASSET, # lambda a, b: a.uri == b.uri and a.extra == b.extra, # ), ( @@ -429,12 +429,12 @@ def test_all_pydantic_models_round_trip(): continue classes.add(obj) exclusion_list = { - "DatasetPydantic", + "AssetPydantic", "DagTagPydantic", - "DagScheduleDatasetReferencePydantic", - "TaskOutletDatasetReferencePydantic", + "DagScheduleAssetReferencePydantic", + "TaskOutletAssetReferencePydantic", "DagOwnerAttributesPydantic", - "DatasetEventPydantic", + "AssetEventPydantic", "TriggerPydantic", } for c in sorted(classes, key=str): @@ -490,7 +490,7 @@ def test_serialized_mapped_operator_unmap(dag_maker): assert serialized_unmapped_task.dag is serialized_dag -def test_ser_of_dataset_event_accessor(): +def test_ser_of_asset_event_accessor(): # todo: (Airflow 3.0) we should force reserialization on upgrade d = OutletEventAccessors() d["hi"].extra = "blah1" # todo: this should maybe be forbidden? i.e. can extra be any json or just dict? diff --git a/tests/system/providers/microsoft/azure/example_msfabric.py b/tests/system/providers/microsoft/azure/example_msfabric.py index 7d62a49e0bc31..5f8b0657c4019 100644 --- a/tests/system/providers/microsoft/azure/example_msfabric.py +++ b/tests/system/providers/microsoft/azure/example_msfabric.py @@ -19,7 +19,7 @@ from datetime import datetime from airflow import models -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.providers.microsoft.azure.operators.msgraph import MSGraphAsyncOperator DAG_ID = "example_msfabric" @@ -44,7 +44,7 @@ query_parameters={"jobType": "Pipeline"}, dag=dag, outlets=[ - Dataset( + Asset( "workspaces/e90b2873-4812-4dfb-9246-593638165644/items/65448530-e5ec-4aeb-a97e-7cebf5d67c18/jobs/instances?jobType=Pipeline" ) ], diff --git a/tests/system/providers/papermill/input_notebook.ipynb b/tests/system/providers/papermill/input_notebook.ipynb index 6c1d53a5a780c..e985432160a8b 100644 --- a/tests/system/providers/papermill/input_notebook.ipynb +++ b/tests/system/providers/papermill/input_notebook.ipynb @@ -36,6 +36,8 @@ "metadata": {}, "outputs": [], "source": [ + "from __future__ import annotations\n", + "\n", "import scrapbook as sb" ] }, diff --git a/tests/test_utils/compat.py b/tests/test_utils/compat.py index 5daf429cf641f..ca1d7e9c77dfa 100644 --- a/tests/test_utils/compat.py +++ b/tests/test_utils/compat.py @@ -55,6 +55,39 @@ from airflow.models.baseoperator import BaseOperatorLink +if TYPE_CHECKING: + from airflow.models.asset import ( + AssetAliasModel, + AssetDagRunQueue, + AssetEvent, + AssetModel, + DagScheduleAssetReference, + TaskOutletAssetReference, + ) +else: + try: + from airflow.models.asset import ( + AssetAliasModel, + AssetDagRunQueue, + AssetEvent, + AssetModel, + DagScheduleAssetReference, + TaskOutletAssetReference, + ) + except ModuleNotFoundError: + # dataset is renamed to asset since Airflow 3.0 + from airflow.models.dataset import ( + DagScheduleDatasetReference as DagScheduleAssetReference, + DatasetDagRunQueue as AssetDagRunQueue, + DatasetEvent as AssetEvent, + DatasetModel as AssetModel, + TaskOutletDatasetReference as TaskOutletAssetReference, + ) + + if AIRFLOW_V_2_10_PLUS: + from airflow.models.dataset import DatasetAliasModel as AssetAliasModel + + def deserialize_operator(serialized_operator: dict[str, Any]) -> Operator: if AIRFLOW_V_2_10_PLUS: # In airflow 2.10+ we can deserialize operator using regular deserialize method. diff --git a/tests/test_utils/db.py b/tests/test_utils/db.py index bd56ed9175cc4..a5dd94e2d009d 100644 --- a/tests/test_utils/db.py +++ b/tests/test_utils/db.py @@ -38,18 +38,19 @@ from airflow.models.dag import DagOwnerAttributes from airflow.models.dagcode import DagCode from airflow.models.dagwarning import DagWarning -from airflow.models.dataset import ( - DagScheduleDatasetReference, - DatasetDagRunQueue, - DatasetEvent, - DatasetModel, - TaskOutletDatasetReference, -) from airflow.models.serialized_dag import SerializedDagModel from airflow.security.permissions import RESOURCE_DAG_PREFIX from airflow.utils.db import add_default_pool_if_not_exists, create_default_connections, reflect_tables from airflow.utils.session import create_session -from tests.test_utils.compat import AIRFLOW_V_2_10_PLUS, ParseImportError +from tests.test_utils.compat import ( + AIRFLOW_V_2_10_PLUS, + AssetDagRunQueue, + AssetEvent, + AssetModel, + DagScheduleAssetReference, + ParseImportError, + TaskOutletAssetReference, +) def clear_db_runs(): @@ -74,17 +75,17 @@ def clear_db_backfills(): session.query(Backfill).delete() -def clear_db_datasets(): +def clear_db_assets(): with create_session() as session: - session.query(DatasetEvent).delete() - session.query(DatasetModel).delete() - session.query(DatasetDagRunQueue).delete() - session.query(DagScheduleDatasetReference).delete() - session.query(TaskOutletDatasetReference).delete() + session.query(AssetEvent).delete() + session.query(AssetModel).delete() + session.query(AssetDagRunQueue).delete() + session.query(DagScheduleAssetReference).delete() + session.query(TaskOutletAssetReference).delete() if AIRFLOW_V_2_10_PLUS: - from airflow.models.dataset import DatasetAliasModel + from tests.test_utils.compat import AssetAliasModel - session.query(DatasetAliasModel).delete() + session.query(AssetAliasModel).delete() def clear_db_dags(): @@ -231,7 +232,7 @@ def clear_dag_specific_permissions(): def clear_all(): clear_db_runs() - clear_db_datasets() + clear_db_assets() clear_db_dags() clear_db_serialized_dags() clear_db_sla_miss() diff --git a/tests/timetables/test_datasets_timetable.py b/tests/timetables/test_assets_timetable.py similarity index 57% rename from tests/timetables/test_datasets_timetable.py rename to tests/timetables/test_assets_timetable.py index b456b9bf5dc9c..f2105891c7298 100644 --- a/tests/timetables/test_datasets_timetable.py +++ b/tests/timetables/test_assets_timetable.py @@ -23,11 +23,11 @@ import pytest from pendulum import DateTime -from airflow.datasets import Dataset, DatasetAlias -from airflow.models.dataset import DatasetAliasModel, DatasetEvent, DatasetModel +from airflow.assets import Asset, AssetAlias +from airflow.models.asset import AssetAliasModel, AssetEvent, AssetModel +from airflow.timetables.assets import AssetOrTimeSchedule from airflow.timetables.base import DagRunInfo, DataInterval, TimeRestriction, Timetable -from airflow.timetables.datasets import DatasetOrTimeSchedule -from airflow.timetables.simple import DatasetTriggeredTimetable +from airflow.timetables.simple import AssetTriggeredTimetable from airflow.utils.types import DagRunType if TYPE_CHECKING: @@ -103,45 +103,45 @@ def test_timetable() -> MockTimetable: @pytest.fixture -def test_datasets() -> list[Dataset]: - """Pytest fixture for creating a list of Dataset objects.""" - return [Dataset("test_dataset")] +def test_assets() -> list[Asset]: + """Pytest fixture for creating a list of Asset objects.""" + return [Asset("test_asset")] @pytest.fixture -def dataset_timetable(test_timetable: MockTimetable, test_datasets: list[Dataset]) -> DatasetOrTimeSchedule: +def asset_timetable(test_timetable: MockTimetable, test_assets: list[Asset]) -> AssetOrTimeSchedule: """ - Pytest fixture for creating a DatasetTimetable object. + Pytest fixture for creating a AssetOrTimeSchedule object. :param test_timetable: The test timetable instance. - :param test_datasets: A list of Dataset instances. + :param test_assets: A list of Asset instances. """ - return DatasetOrTimeSchedule(timetable=test_timetable, datasets=test_datasets) + return AssetOrTimeSchedule(timetable=test_timetable, assets=test_assets) -def test_serialization(dataset_timetable: DatasetOrTimeSchedule, monkeypatch: Any) -> None: +def test_serialization(asset_timetable: AssetOrTimeSchedule, monkeypatch: Any) -> None: """ - Tests the serialization method of DatasetTimetable. + Tests the serialization method of AssetOrTimeSchedule. - :param dataset_timetable: The DatasetTimetable instance to test. + :param asset_timetable: The AssetOrTimeSchedule instance to test. :param monkeypatch: The monkeypatch fixture from pytest. """ monkeypatch.setattr( "airflow.serialization.serialized_objects.encode_timetable", lambda x: "mock_serialized_timetable" ) - serialized = dataset_timetable.serialize() + serialized = asset_timetable.serialize() assert serialized == { "timetable": "mock_serialized_timetable", - "dataset_condition": { - "__type": "dataset_all", - "objects": [{"__type": "dataset", "uri": "test_dataset", "extra": {}}], + "asset_condition": { + "__type": "asset_all", + "objects": [{"__type": "asset", "uri": "test_asset", "extra": {}}], }, } def test_deserialization(monkeypatch: Any) -> None: """ - Tests the deserialization method of DatasetTimetable. + Tests the deserialization method of AssetOrTimeSchedule. :param monkeypatch: The monkeypatch fixture from pytest. """ @@ -150,55 +150,55 @@ def test_deserialization(monkeypatch: Any) -> None: ) mock_serialized_data = { "timetable": "mock_serialized_timetable", - "dataset_condition": { - "__type": "dataset_all", - "objects": [{"__type": "dataset", "uri": "test_dataset", "extra": None}], + "asset_condition": { + "__type": "asset_all", + "objects": [{"__type": "asset", "uri": "test_asset", "extra": None}], }, } - deserialized = DatasetOrTimeSchedule.deserialize(mock_serialized_data) - assert isinstance(deserialized, DatasetOrTimeSchedule) + deserialized = AssetOrTimeSchedule.deserialize(mock_serialized_data) + assert isinstance(deserialized, AssetOrTimeSchedule) -def test_infer_manual_data_interval(dataset_timetable: DatasetOrTimeSchedule) -> None: +def test_infer_manual_data_interval(asset_timetable: AssetOrTimeSchedule) -> None: """ - Tests the infer_manual_data_interval method of DatasetTimetable. + Tests the infer_manual_data_interval method of AssetOrTimeSchedule. - :param dataset_timetable: The DatasetTimetable instance to test. + :param asset_timetable: The AssetOrTimeSchedule instance to test. """ run_after = DateTime.now() - result = dataset_timetable.infer_manual_data_interval(run_after=run_after) + result = asset_timetable.infer_manual_data_interval(run_after=run_after) assert isinstance(result, DataInterval) -def test_next_dagrun_info(dataset_timetable: DatasetOrTimeSchedule) -> None: +def test_next_dagrun_info(asset_timetable: AssetOrTimeSchedule) -> None: """ - Tests the next_dagrun_info method of DatasetTimetable. + Tests the next_dagrun_info method of AssetOrTimeSchedule. - :param dataset_timetable: The DatasetTimetable instance to test. + :param asset_timetable: The AssetOrTimeSchedule instance to test. """ last_interval = DataInterval.exact(DateTime.now()) restriction = TimeRestriction(earliest=DateTime.now(), latest=None, catchup=True) - result = dataset_timetable.next_dagrun_info( + result = asset_timetable.next_dagrun_info( last_automated_data_interval=last_interval, restriction=restriction ) assert result is None or isinstance(result, DagRunInfo) -def test_generate_run_id(dataset_timetable: DatasetOrTimeSchedule) -> None: +def test_generate_run_id(asset_timetable: AssetOrTimeSchedule) -> None: """ - Tests the generate_run_id method of DatasetTimetable. + Tests the generate_run_id method of AssetOrTimeSchedule. - :param dataset_timetable: The DatasetTimetable instance to test. + :param asset_timetable: The AssetOrTimeSchedule instance to test. """ - run_id = dataset_timetable.generate_run_id( + run_id = asset_timetable.generate_run_id( run_type=DagRunType.MANUAL, extra_args="test", logical_date=DateTime.now(), data_interval=None ) assert isinstance(run_id, str) @pytest.fixture -def dataset_events(mocker) -> list[DatasetEvent]: - """Pytest fixture for creating mock DatasetEvent objects.""" +def asset_events(mocker) -> list[AssetEvent]: + """Pytest fixture for creating mock AssetEvent objects.""" now = DateTime.now() earlier = now.subtract(days=1) later = now.add(days=1) @@ -212,9 +212,9 @@ def dataset_events(mocker) -> list[DatasetEvent]: mock_dag_run_later.data_interval_start = now mock_dag_run_later.data_interval_end = later - # Create DatasetEvent objects with mock source_dag_run - event_earlier = DatasetEvent(timestamp=earlier, dataset_id=1) - event_later = DatasetEvent(timestamp=later, dataset_id=1) + # Create AssetEvent objects with mock source_dag_run + event_earlier = AssetEvent(timestamp=earlier, dataset_id=1) + event_later = AssetEvent(timestamp=later, dataset_id=1) # Use mocker to set the source_dag_run attribute to avoid SQLAlchemy's instrumentation mocker.patch.object(event_earlier, "source_dag_run", new=mock_dag_run_earlier) @@ -224,54 +224,50 @@ def dataset_events(mocker) -> list[DatasetEvent]: def test_data_interval_for_events( - dataset_timetable: DatasetOrTimeSchedule, dataset_events: list[DatasetEvent] + asset_timetable: AssetOrTimeSchedule, asset_events: list[AssetEvent] ) -> None: """ - Tests the data_interval_for_events method of DatasetTimetable. + Tests the data_interval_for_events method of AssetOrTimeSchedule. - :param dataset_timetable: The DatasetTimetable instance to test. - :param dataset_events: A list of mock DatasetEvent instances. + :param asset_timetable: The AssetOrTimeSchedule instance to test. + :param asset_events: A list of mock AssetEvent instances. """ - data_interval = dataset_timetable.data_interval_for_events( - logical_date=DateTime.now(), events=dataset_events - ) + data_interval = asset_timetable.data_interval_for_events(logical_date=DateTime.now(), events=asset_events) assert data_interval.start == min( - event.timestamp for event in dataset_events + event.timestamp for event in asset_events ), "Data interval start does not match" assert data_interval.end == max( - event.timestamp for event in dataset_events + event.timestamp for event in asset_events ), "Data interval end does not match" -def test_run_ordering_inheritance(dataset_timetable: DatasetOrTimeSchedule) -> None: +def test_run_ordering_inheritance(asset_timetable: AssetOrTimeSchedule) -> None: """ - Tests that DatasetOrTimeSchedule inherits run_ordering from its parent class correctly. + Tests that AssetOrTimeSchedule inherits run_ordering from its parent class correctly. - :param dataset_timetable: The DatasetTimetable instance to test. + :param asset_timetable: The AssetOrTimeSchedule instance to test. """ assert hasattr( - dataset_timetable, "run_ordering" - ), "DatasetOrTimeSchedule should have 'run_ordering' attribute" - parent_run_ordering = getattr(DatasetTriggeredTimetable, "run_ordering", None) - assert ( - dataset_timetable.run_ordering == parent_run_ordering - ), "run_ordering does not match the parent class" + asset_timetable, "run_ordering" + ), "AssetOrTimeSchedule should have 'run_ordering' attribute" + parent_run_ordering = getattr(AssetTriggeredTimetable, "run_ordering", None) + assert asset_timetable.run_ordering == parent_run_ordering, "run_ordering does not match the parent class" @pytest.mark.db_test def test_summary(session: Session) -> None: - dataset_model = DatasetModel(uri="test_dataset") - dataset_alias_model = DatasetAliasModel(name="test_dataset_alias") - session.add_all([dataset_model, dataset_alias_model]) + asset_model = AssetModel(uri="test_asset") + asset_alias_model = AssetAliasModel(name="test_asset_alias") + session.add_all([asset_model, asset_alias_model]) session.commit() - dataset_alias = DatasetAlias("test_dataset_alias") - table = DatasetTriggeredTimetable(dataset_alias) - assert table.summary == "Unresolved DatasetAlias" + asset_alias = AssetAlias("test_asset_alias") + table = AssetTriggeredTimetable(asset_alias) + assert table.summary == "Unresolved AssetAlias" - dataset_alias_model.datasets.append(dataset_model) - session.add(dataset_alias_model) + asset_alias_model.datasets.append(asset_model) + session.add(asset_alias_model) session.commit() - table = DatasetTriggeredTimetable(dataset_alias) - assert table.summary == "Dataset" + table = AssetTriggeredTimetable(asset_alias) + assert table.summary == "Asset" diff --git a/tests/utils/test_context.py b/tests/utils/test_context.py index 0f4f80f36504c..5d2f7543b6299 100644 --- a/tests/utils/test_context.py +++ b/tests/utils/test_context.py @@ -20,55 +20,55 @@ import pytest -from airflow.datasets import Dataset, DatasetAlias, DatasetAliasEvent -from airflow.models.dataset import DatasetAliasModel, DatasetModel +from airflow.assets import Asset, AssetAlias, AssetAliasEvent +from airflow.models.asset import AssetAliasModel, AssetModel from airflow.utils.context import OutletEventAccessor, OutletEventAccessors class TestOutletEventAccessor: @pytest.mark.parametrize( - "raw_key, dataset_alias_events", + "raw_key, asset_alias_events", ( ( - DatasetAlias("test_alias"), - [DatasetAliasEvent(source_alias_name="test_alias", dest_dataset_uri="test_uri", extra={})], + AssetAlias("test_alias"), + [AssetAliasEvent(source_alias_name="test_alias", dest_asset_uri="test_uri", extra={})], ), - (Dataset("test_uri"), []), + (Asset("test_uri"), []), ), ) - def test_add(self, raw_key, dataset_alias_events): + def test_add(self, raw_key, asset_alias_events): outlet_event_accessor = OutletEventAccessor(raw_key=raw_key, extra={}) - outlet_event_accessor.add(Dataset("test_uri")) - assert outlet_event_accessor.dataset_alias_events == dataset_alias_events + outlet_event_accessor.add(Asset("test_uri")) + assert outlet_event_accessor.asset_alias_events == asset_alias_events @pytest.mark.db_test @pytest.mark.parametrize( - "raw_key, dataset_alias_events", + "raw_key, asset_alias_events", ( ( - DatasetAlias("test_alias"), - [DatasetAliasEvent(source_alias_name="test_alias", dest_dataset_uri="test_uri", extra={})], + AssetAlias("test_alias"), + [AssetAliasEvent(source_alias_name="test_alias", dest_asset_uri="test_uri", extra={})], ), ( "test_alias", - [DatasetAliasEvent(source_alias_name="test_alias", dest_dataset_uri="test_uri", extra={})], + [AssetAliasEvent(source_alias_name="test_alias", dest_asset_uri="test_uri", extra={})], ), - (Dataset("test_uri"), []), + (Asset("test_uri"), []), ), ) - def test_add_with_db(self, raw_key, dataset_alias_events, session): - dsm = DatasetModel(uri="test_uri") - dsam = DatasetAliasModel(name="test_alias") - session.add_all([dsm, dsam]) + def test_add_with_db(self, raw_key, asset_alias_events, session): + asm = AssetModel(uri="test_uri") + aam = AssetAliasModel(name="test_alias") + session.add_all([asm, aam]) session.flush() outlet_event_accessor = OutletEventAccessor(raw_key=raw_key, extra={"not": ""}) outlet_event_accessor.add("test_uri", extra={}) - assert outlet_event_accessor.dataset_alias_events == dataset_alias_events + assert outlet_event_accessor.asset_alias_events == asset_alias_events class TestOutletEventAccessors: - @pytest.mark.parametrize("key", ("test", Dataset("test"), DatasetAlias("test_alias"))) + @pytest.mark.parametrize("key", ("test", Asset("test"), AssetAlias("test_alias"))) def test____get_item___dict_key_not_exists(self, key): outlet_event_accessors = OutletEventAccessors() assert len(outlet_event_accessors) == 0 diff --git a/tests/utils/test_db_cleanup.py b/tests/utils/test_db_cleanup.py index 06e99523faa2f..0a8cd9c90c962 100644 --- a/tests/utils/test_db_cleanup.py +++ b/tests/utils/test_db_cleanup.py @@ -48,7 +48,7 @@ run_cleanup, ) from airflow.utils.session import create_session -from tests.test_utils.db import clear_db_dags, clear_db_datasets, clear_db_runs, drop_tables_with_prefix +from tests.test_utils.db import clear_db_assets, clear_db_dags, clear_db_runs, drop_tables_with_prefix pytestmark = [pytest.mark.db_test, pytest.mark.skip_if_database_isolation_mode] @@ -57,11 +57,11 @@ def clean_database(): """Fixture that cleans the database before and after every test.""" clear_db_runs() - clear_db_datasets() + clear_db_assets() clear_db_dags() yield # Test runs here clear_db_dags() - clear_db_datasets() + clear_db_assets() clear_db_runs() diff --git a/tests/utils/test_json.py b/tests/utils/test_json.py index 5e7e6eb1e5c56..5a58b5d790329 100644 --- a/tests/utils/test_json.py +++ b/tests/utils/test_json.py @@ -26,7 +26,7 @@ import pendulum import pytest -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.utils import json as utils_json @@ -85,11 +85,11 @@ def test_encode_raises(self): cls=utils_json.XComEncoder, ) - def test_encode_xcom_dataset(self): - dataset = Dataset("mytest://dataset") - s = json.dumps(dataset, cls=utils_json.XComEncoder) + def test_encode_xcom_asset(self): + asset = Asset("mytest://asset") + s = json.dumps(asset, cls=utils_json.XComEncoder) obj = json.loads(s, cls=utils_json.XComDecoder) - assert dataset.uri == obj.uri + assert asset.uri == obj.uri @pytest.mark.parametrize( "data", diff --git a/tests/www/test_auth.py b/tests/www/test_auth.py index 613812968d27f..40e731b060972 100644 --- a/tests/www/test_auth.py +++ b/tests/www/test_auth.py @@ -34,7 +34,7 @@ "decorator_name, is_authorized_method_name", [ ("has_access_configuration", "is_authorized_configuration"), - ("has_access_dataset", "is_authorized_dataset"), + ("has_access_asset", "is_authorized_asset"), ("has_access_view", "is_authorized_view"), ], ) diff --git a/tests/www/views/test_views_acl.py b/tests/www/views/test_views_acl.py index 053e0f339fb3f..139644f67a6da 100644 --- a/tests/www/views/test_views_acl.py +++ b/tests/www/views/test_views_acl.py @@ -264,22 +264,22 @@ def test_dag_autocomplete_success(client_all_dags): {"name": "airflow", "type": "owner", "dag_display_name": None}, { "dag_display_name": None, - "name": "dataset_alias_example_alias_consumer_with_no_taskflow", + "name": "asset_alias_example_alias_consumer_with_no_taskflow", "type": "dag", }, { "dag_display_name": None, - "name": "dataset_alias_example_alias_producer_with_no_taskflow", + "name": "asset_alias_example_alias_producer_with_no_taskflow", "type": "dag", }, { "dag_display_name": None, - "name": "dataset_s3_bucket_consumer_with_no_taskflow", + "name": "asset_s3_bucket_consumer_with_no_taskflow", "type": "dag", }, { "dag_display_name": None, - "name": "dataset_s3_bucket_producer_with_no_taskflow", + "name": "asset_s3_bucket_producer_with_no_taskflow", "type": "dag", }, { diff --git a/tests/www/views/test_views_dataset.py b/tests/www/views/test_views_dataset.py index 797ed40ba009e..3d3351bb6493a 100644 --- a/tests/www/views/test_views_dataset.py +++ b/tests/www/views/test_views_dataset.py @@ -21,11 +21,11 @@ import pytest from dateutil.tz import UTC -from airflow.datasets import Dataset -from airflow.models.dataset import DatasetEvent, DatasetModel +from airflow.assets import Asset +from airflow.models.asset import AssetEvent, AssetModel from airflow.operators.empty import EmptyOperator from tests.test_utils.asserts import assert_queries_count -from tests.test_utils.db import clear_db_datasets +from tests.test_utils.db import clear_db_assets pytestmark = pytest.mark.db_test @@ -33,23 +33,23 @@ class TestDatasetEndpoint: @pytest.fixture(autouse=True) def cleanup(self): - clear_db_datasets() + clear_db_assets() yield - clear_db_datasets() + clear_db_assets() class TestGetDatasets(TestDatasetEndpoint): def test_should_respond_200(self, admin_client, session): - datasets = [ - DatasetModel( + assets = [ + AssetModel( id=i, uri=f"s3://bucket/key/{i}", ) for i in [1, 2] ] - session.add_all(datasets) + session.add_all(assets) session.commit() - assert session.query(DatasetModel).count() == 2 + assert session.query(AssetModel).count() == 2 with assert_queries_count(10): response = admin_client.get("/object/datasets_summary") @@ -75,15 +75,15 @@ def test_should_respond_200(self, admin_client, session): } def test_order_by_raises_400_for_invalid_attr(self, admin_client, session): - datasets = [ - DatasetModel( + assets = [ + AssetModel( uri=f"s3://bucket/key/{i}", ) for i in [1, 2] ] - session.add_all(datasets) + session.add_all(assets) session.commit() - assert session.query(DatasetModel).count() == 2 + assert session.query(AssetModel).count() == 2 response = admin_client.get("/object/datasets_summary?order_by=fake") @@ -92,15 +92,10 @@ def test_order_by_raises_400_for_invalid_attr(self, admin_client, session): assert response.json["detail"] == msg def test_order_by_raises_400_for_invalid_datetimes(self, admin_client, session): - datasets = [ - DatasetModel( - uri=f"s3://bucket/key/{i}", - ) - for i in [1, 2] - ] - session.add_all(datasets) + assets = [AssetModel(uri=f"s3://bucket/key/{i}") for i in [1, 2]] + session.add_all(assets) session.commit() - assert session.query(DatasetModel).count() == 2 + assert session.query(AssetModel).count() == 2 response = admin_client.get("/object/datasets_summary?updated_before=null") @@ -115,25 +110,25 @@ def test_order_by_raises_400_for_invalid_datetimes(self, admin_client, session): def test_filter_by_datetimes(self, admin_client, session): today = pendulum.today("UTC") - datasets = [ - DatasetModel( + assets = [ + AssetModel( id=i, uri=f"s3://bucket/key/{i}", ) for i in range(1, 4) ] - session.add_all(datasets) - # Update datasets, one per day, starting with datasets[0], ending with datasets[2] - dataset_events = [ - DatasetEvent( - dataset_id=datasets[i].id, - timestamp=today.add(days=-len(datasets) + i + 1), + session.add_all(assets) + # Update assets, one per day, starting with assets[0], ending with assets[2] + asset_events = [ + AssetEvent( + dataset_id=assets[i].id, + timestamp=today.add(days=-len(assets) + i + 1), ) - for i in range(len(datasets)) + for i in range(len(assets)) ] - session.add_all(dataset_events) + session.add_all(asset_events) session.commit() - assert session.query(DatasetModel).count() == len(datasets) + assert session.query(AssetModel).count() == len(assets) cutoff = today.add(days=-1).add(minutes=-5).to_iso8601_string() response = admin_client.get(f"/object/datasets_summary?updated_after={cutoff}") @@ -150,7 +145,7 @@ def test_filter_by_datetimes(self, admin_client, session): assert [json_dict["id"] for json_dict in response.json["datasets"]] == [1, 2] @pytest.mark.parametrize( - "order_by, ordered_dataset_ids", + "order_by, ordered_asset_ids", [ ("uri", [1, 2, 3, 4]), ("-uri", [4, 3, 2, 1]), @@ -158,50 +153,50 @@ def test_filter_by_datetimes(self, admin_client, session): ("-last_dataset_update", [2, 3, 1, 4]), ], ) - def test_order_by(self, admin_client, session, order_by, ordered_dataset_ids): - datasets = [ - DatasetModel( + def test_order_by(self, admin_client, session, order_by, ordered_asset_ids): + assets = [ + AssetModel( id=i, uri=f"s3://bucket/key/{i}", ) - for i in range(1, len(ordered_dataset_ids) + 1) + for i in range(1, len(ordered_asset_ids) + 1) ] - session.add_all(datasets) - dataset_events = [ - DatasetEvent( - dataset_id=datasets[2].id, + session.add_all(assets) + asset_events = [ + AssetEvent( + dataset_id=assets[2].id, timestamp=pendulum.today("UTC").add(days=-3), ), - DatasetEvent( - dataset_id=datasets[1].id, + AssetEvent( + dataset_id=assets[1].id, timestamp=pendulum.today("UTC").add(days=-2), ), - DatasetEvent( - dataset_id=datasets[1].id, + AssetEvent( + dataset_id=assets[1].id, timestamp=pendulum.today("UTC").add(days=-1), ), ] - session.add_all(dataset_events) + session.add_all(asset_events) session.commit() - assert session.query(DatasetModel).count() == len(ordered_dataset_ids) + assert session.query(AssetModel).count() == len(ordered_asset_ids) response = admin_client.get(f"/object/datasets_summary?order_by={order_by}") assert response.status_code == 200 - assert ordered_dataset_ids == [json_dict["id"] for json_dict in response.json["datasets"]] - assert response.json["total_entries"] == len(ordered_dataset_ids) + assert ordered_asset_ids == [json_dict["id"] for json_dict in response.json["datasets"]] + assert response.json["total_entries"] == len(ordered_asset_ids) def test_search_uri_pattern(self, admin_client, session): - datasets = [ - DatasetModel( + assets = [ + AssetModel( id=i, uri=f"s3://bucket/key_{i}", ) for i in [1, 2] ] - session.add_all(datasets) + session.add_all(assets) session.commit() - assert session.query(DatasetModel).count() == 2 + assert session.query(AssetModel).count() == 2 uri_pattern = "key_2" response = admin_client.get(f"/object/datasets_summary?uri_pattern={uri_pattern}") @@ -246,92 +241,92 @@ def test_search_uri_pattern(self, admin_client, session): @pytest.mark.need_serialized_dag def test_correct_counts_update(self, admin_client, session, dag_maker, app, monkeypatch): with monkeypatch.context() as m: - datasets = [Dataset(uri=f"s3://bucket/key/{i}") for i in [1, 2, 3, 4, 5]] + assets = [Asset(uri=f"s3://bucket/key/{i}") for i in [1, 2, 3, 4, 5]] - # DAG that produces dataset #1 + # DAG that produces asset #1 with dag_maker(dag_id="upstream", schedule=None, serialized=True, session=session): - EmptyOperator(task_id="task1", outlets=[datasets[0]]) + EmptyOperator(task_id="task1", outlets=[assets[0]]) - # DAG that is consumes only datasets #1 and #2 - with dag_maker(dag_id="downstream", schedule=datasets[:2], serialized=True, session=session): + # DAG that is consumes only assets #1 and #2 + with dag_maker(dag_id="downstream", schedule=assets[:2], serialized=True, session=session): EmptyOperator(task_id="task1") - # We create multiple dataset-producing and dataset-consuming DAGs because the query requires + # We create multiple asset-producing and asset-consuming DAGs because the query requires # COUNT(DISTINCT ...) for total_updates, or else it returns a multiple of the correct number due - # to the outer joins with DagScheduleDatasetReference and TaskOutletDatasetReference - # Two independent DAGs that produce dataset #3 + # to the outer joins with DagScheduleAssetReference and TaskOutletAssetReference + # Two independent DAGs that produce asset #3 with dag_maker(dag_id="independent_producer_1", serialized=True, session=session): - EmptyOperator(task_id="task1", outlets=[datasets[2]]) + EmptyOperator(task_id="task1", outlets=[assets[2]]) with dag_maker(dag_id="independent_producer_2", serialized=True, session=session): - EmptyOperator(task_id="task1", outlets=[datasets[2]]) - # Two independent DAGs that consume dataset #4 + EmptyOperator(task_id="task1", outlets=[assets[2]]) + # Two independent DAGs that consume asset #4 with dag_maker( dag_id="independent_consumer_1", - schedule=[datasets[3]], + schedule=[assets[3]], serialized=True, session=session, ): EmptyOperator(task_id="task1") with dag_maker( dag_id="independent_consumer_2", - schedule=[datasets[3]], + schedule=[assets[3]], serialized=True, session=session, ): EmptyOperator(task_id="task1") - # Independent DAG that is produces and consumes the same dataset, #5 + # Independent DAG that is produces and consumes the same asset, #5 with dag_maker( dag_id="independent_producer_self_consumer", - schedule=[datasets[4]], + schedule=[assets[4]], serialized=True, session=session, ): - EmptyOperator(task_id="task1", outlets=[datasets[4]]) + EmptyOperator(task_id="task1", outlets=[assets[4]]) m.setattr(app, "dag_bag", dag_maker.dagbag) - ds1_id = session.query(DatasetModel.id).filter_by(uri=datasets[0].uri).scalar() - ds2_id = session.query(DatasetModel.id).filter_by(uri=datasets[1].uri).scalar() - ds3_id = session.query(DatasetModel.id).filter_by(uri=datasets[2].uri).scalar() - ds4_id = session.query(DatasetModel.id).filter_by(uri=datasets[3].uri).scalar() - ds5_id = session.query(DatasetModel.id).filter_by(uri=datasets[4].uri).scalar() + asset1_id = session.query(AssetModel.id).filter_by(uri=assets[0].uri).scalar() + asset2_id = session.query(AssetModel.id).filter_by(uri=assets[1].uri).scalar() + asset3_id = session.query(AssetModel.id).filter_by(uri=assets[2].uri).scalar() + asset4_id = session.query(AssetModel.id).filter_by(uri=assets[3].uri).scalar() + asset5_id = session.query(AssetModel.id).filter_by(uri=assets[4].uri).scalar() - # dataset 1 events + # asset 1 events session.add_all( [ - DatasetEvent( - dataset_id=ds1_id, + AssetEvent( + dataset_id=asset1_id, timestamp=pendulum.DateTime(2022, 8, 1, i, tzinfo=UTC), ) for i in range(3) ] ) - # dataset 3 events + # asset 3 events session.add_all( [ - DatasetEvent( - dataset_id=ds3_id, + AssetEvent( + dataset_id=asset3_id, timestamp=pendulum.DateTime(2022, 8, 1, i, tzinfo=UTC), ) for i in range(3) ] ) - # dataset 4 events + # asset 4 events session.add_all( [ - DatasetEvent( - dataset_id=ds4_id, + AssetEvent( + dataset_id=asset4_id, timestamp=pendulum.DateTime(2022, 8, 1, i, tzinfo=UTC), ) for i in range(4) ] ) - # dataset 5 events + # asset 5 events session.add_all( [ - DatasetEvent( - dataset_id=ds5_id, + AssetEvent( + dataset_id=asset5_id, timestamp=pendulum.DateTime(2022, 8, 1, i, tzinfo=UTC), ) for i in range(5) @@ -346,31 +341,31 @@ def test_correct_counts_update(self, admin_client, session, dag_maker, app, monk assert response_data == { "datasets": [ { - "id": ds1_id, + "id": asset1_id, "uri": "s3://bucket/key/1", "last_dataset_update": "2022-08-01T02:00:00+00:00", "total_updates": 3, }, { - "id": ds2_id, + "id": asset2_id, "uri": "s3://bucket/key/2", "last_dataset_update": None, "total_updates": 0, }, { - "id": ds3_id, + "id": asset3_id, "uri": "s3://bucket/key/3", "last_dataset_update": "2022-08-01T02:00:00+00:00", "total_updates": 3, }, { - "id": ds4_id, + "id": asset4_id, "uri": "s3://bucket/key/4", "last_dataset_update": "2022-08-01T03:00:00+00:00", "total_updates": 4, }, { - "id": ds5_id, + "id": asset5_id, "uri": "s3://bucket/key/5", "last_dataset_update": "2022-08-01T04:00:00+00:00", "total_updates": 5, @@ -395,14 +390,14 @@ class TestGetDatasetsEndpointPagination(TestDatasetEndpoint): ], ) def test_limit_and_offset(self, admin_client, session, url, expected_dataset_uris): - datasets = [ - DatasetModel( + assets = [ + AssetModel( uri=f"s3://bucket/key/{i}", extra={"foo": "bar"}, ) for i in range(1, 10) ] - session.add_all(datasets) + session.add_all(assets) session.commit() response = admin_client.get(url) @@ -412,14 +407,14 @@ def test_limit_and_offset(self, admin_client, session, url, expected_dataset_uri assert dataset_uris == expected_dataset_uris def test_should_respect_page_size_limit_default(self, admin_client, session): - datasets = [ - DatasetModel( + assets = [ + AssetModel( uri=f"s3://bucket/key/{i}", extra={"foo": "bar"}, ) for i in range(1, 60) ] - session.add_all(datasets) + session.add_all(assets) session.commit() response = admin_client.get("/object/datasets_summary") @@ -428,14 +423,14 @@ def test_should_respect_page_size_limit_default(self, admin_client, session): assert len(response.json["datasets"]) == 25 def test_should_return_max_if_req_above(self, admin_client, session): - datasets = [ - DatasetModel( + assets = [ + AssetModel( uri=f"s3://bucket/key/{i}", extra={"foo": "bar"}, ) for i in range(1, 60) ] - session.add_all(datasets) + session.add_all(assets) session.commit() response = admin_client.get("/object/datasets_summary?limit=180") @@ -446,7 +441,7 @@ def test_should_return_max_if_req_above(self, admin_client, session): class TestGetDatasetNextRunSummary(TestDatasetEndpoint): def test_next_run_dataset_summary(self, dag_maker, admin_client): - with dag_maker(dag_id="upstream", schedule=[Dataset(uri="s3://bucket/key/1")], serialized=True): + with dag_maker(dag_id="upstream", schedule=[Asset(uri="s3://bucket/key/1")], serialized=True): EmptyOperator(task_id="task1") response = admin_client.post("/next_run_datasets_summary", data={"dag_ids": ["upstream"]}) diff --git a/tests/www/views/test_views_grid.py b/tests/www/views/test_views_grid.py index 8726b67e1dad3..b4dd6f6082e57 100644 --- a/tests/www/views/test_views_grid.py +++ b/tests/www/views/test_views_grid.py @@ -24,11 +24,11 @@ import pytest from dateutil.tz import UTC -from airflow.datasets import Dataset +from airflow.assets import Asset from airflow.decorators import task_group from airflow.lineage.entities import File from airflow.models import DagBag -from airflow.models.dataset import DatasetDagRunQueue, DatasetEvent, DatasetModel +from airflow.models.asset import AssetDagRunQueue, AssetEvent, AssetModel from airflow.operators.empty import EmptyOperator from airflow.utils import timezone from airflow.utils.state import DagRunState, TaskInstanceState @@ -36,7 +36,7 @@ from airflow.utils.types import DagRunType from airflow.www.views import dag_to_grid from tests.test_utils.asserts import assert_queries_count -from tests.test_utils.db import clear_db_datasets, clear_db_runs +from tests.test_utils.db import clear_db_assets, clear_db_runs from tests.test_utils.mock_operators import MockOperator pytestmark = pytest.mark.db_test @@ -56,10 +56,10 @@ def examples_dag_bag(): @pytest.fixture(autouse=True) def clean(): clear_db_runs() - clear_db_datasets() + clear_db_assets() yield clear_db_runs() - clear_db_datasets() + clear_db_assets() @pytest.fixture @@ -419,7 +419,7 @@ def test_query_count(dag_with_runs, session): dag_to_grid(run1.dag, (run1, run2), session) -def test_has_outlet_dataset_flag(admin_client, dag_maker, session, app, monkeypatch): +def test_has_outlet_asset_flag(admin_client, dag_maker, session, app, monkeypatch): with monkeypatch.context() as m: # Remove global operator links for this test m.setattr("airflow.plugins_manager.global_operator_extra_links", []) @@ -430,8 +430,8 @@ def test_has_outlet_dataset_flag(admin_client, dag_maker, session, app, monkeypa lineagefile = File("/tmp/does_not_exist") EmptyOperator(task_id="task1") EmptyOperator(task_id="task2", outlets=[lineagefile]) - EmptyOperator(task_id="task3", outlets=[Dataset("foo"), lineagefile]) - EmptyOperator(task_id="task4", outlets=[Dataset("foo")]) + EmptyOperator(task_id="task3", outlets=[Asset("foo"), lineagefile]) + EmptyOperator(task_id="task4", outlets=[Asset("foo")]) m.setattr(app, "dag_bag", dag_maker.dagbag) resp = admin_client.get(f"/object/grid_data?dag_id={DAG_ID}", follow_redirects=True) @@ -470,37 +470,37 @@ def _expected_task_details(task_id, has_outlet_datasets): @pytest.mark.need_serialized_dag def test_next_run_datasets(admin_client, dag_maker, session, app, monkeypatch): with monkeypatch.context() as m: - datasets = [Dataset(uri=f"s3://bucket/key/{i}") for i in [1, 2]] + assets = [Asset(uri=f"s3://bucket/key/{i}") for i in [1, 2]] - with dag_maker(dag_id=DAG_ID, schedule=datasets, serialized=True, session=session): + with dag_maker(dag_id=DAG_ID, schedule=assets, serialized=True, session=session): EmptyOperator(task_id="task1") m.setattr(app, "dag_bag", dag_maker.dagbag) - ds1_id = session.query(DatasetModel.id).filter_by(uri=datasets[0].uri).scalar() - ds2_id = session.query(DatasetModel.id).filter_by(uri=datasets[1].uri).scalar() - ddrq = DatasetDagRunQueue( - target_dag_id=DAG_ID, dataset_id=ds1_id, created_at=pendulum.DateTime(2022, 8, 2, tzinfo=UTC) + asset1_id = session.query(AssetModel.id).filter_by(uri=assets[0].uri).scalar() + asset2_id = session.query(AssetModel.id).filter_by(uri=assets[1].uri).scalar() + adrq = AssetDagRunQueue( + target_dag_id=DAG_ID, dataset_id=asset1_id, created_at=pendulum.DateTime(2022, 8, 2, tzinfo=UTC) ) - session.add(ddrq) - dataset_events = [ - DatasetEvent( - dataset_id=ds1_id, + session.add(adrq) + asset_events = [ + AssetEvent( + dataset_id=asset1_id, extra={}, timestamp=pendulum.DateTime(2022, 8, 1, 1, tzinfo=UTC), ), - DatasetEvent( - dataset_id=ds1_id, + AssetEvent( + dataset_id=asset1_id, extra={}, timestamp=pendulum.DateTime(2022, 8, 2, 1, tzinfo=UTC), ), - DatasetEvent( - dataset_id=ds1_id, + AssetEvent( + dataset_id=asset1_id, extra={}, timestamp=pendulum.DateTime(2022, 8, 2, 2, tzinfo=UTC), ), ] - session.add_all(dataset_events) + session.add_all(asset_events) session.commit() resp = admin_client.get(f"/object/next_run_datasets/{DAG_ID}", follow_redirects=True) @@ -509,8 +509,8 @@ def test_next_run_datasets(admin_client, dag_maker, session, app, monkeypatch): assert resp.json == { "dataset_expression": {"all": ["s3://bucket/key/1", "s3://bucket/key/2"]}, "events": [ - {"id": ds1_id, "uri": "s3://bucket/key/1", "lastUpdate": "2022-08-02T02:00:00+00:00"}, - {"id": ds2_id, "uri": "s3://bucket/key/2", "lastUpdate": None}, + {"id": asset1_id, "uri": "s3://bucket/key/1", "lastUpdate": "2022-08-02T02:00:00+00:00"}, + {"id": asset2_id, "uri": "s3://bucket/key/2", "lastUpdate": None}, ], } From 96a5f68f00a95751447c962db4fa8c4087b05dda Mon Sep 17 00:00:00 2001 From: Tzu-ping Chung Date: Sun, 29 Sep 2024 23:27:41 -0700 Subject: [PATCH 072/802] Add 'name' and 'group' to DatasetModel (#42407) The unique index is also modified to include 'name', so now an asset is considered unique if *either* the name or URI is different. This makes no difference for the moment---the name is simply populated from URI. We'll add a public interface to set the name in a later PR. This PR strictly only touches the model so it does not conflict with too many things, and can be merged quickly. The unique index on DatasetAliasModel is also renamed since we were using a wrong naming convention on both models. Since the index namespace is shared in the entire database, the index name should include additional components. The idx_name_unique is still usable, but we should a better citizen and name this the right way(tm). --- airflow/assets/__init__.py | 2 +- ...4_3_0_0_add_name_field_to_dataset_model.py | 94 + airflow/models/asset.py | 43 +- airflow/utils/db.py | 2 +- docs/apache-airflow/img/airflow_erd.sha256 | 2 +- docs/apache-airflow/img/airflow_erd.svg | 2266 +++++++++-------- docs/apache-airflow/migrations-ref.rst | 4 +- 7 files changed, 1272 insertions(+), 1141 deletions(-) create mode 100644 airflow/migrations/versions/0034_3_0_0_add_name_field_to_dataset_model.py diff --git a/airflow/assets/__init__.py b/airflow/assets/__init__.py index 9727e408edc2e..deb9aa593ded5 100644 --- a/airflow/assets/__init__.py +++ b/airflow/assets/__init__.py @@ -256,7 +256,7 @@ class Asset(os.PathLike, BaseAsset): uri: str = attr.field( converter=_sanitize_uri, - validator=[attr.validators.min_len(1), attr.validators.max_len(3000)], + validator=[attr.validators.min_len(1), attr.validators.max_len(1500)], ) extra: dict[str, Any] = attr.field(factory=dict, converter=_set_extra_default) diff --git a/airflow/migrations/versions/0034_3_0_0_add_name_field_to_dataset_model.py b/airflow/migrations/versions/0034_3_0_0_add_name_field_to_dataset_model.py new file mode 100644 index 0000000000000..5c8aec69e9be9 --- /dev/null +++ b/airflow/migrations/versions/0034_3_0_0_add_name_field_to_dataset_model.py @@ -0,0 +1,94 @@ +# +# 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. + +""" +Add name and group fields to DatasetModel. + +The unique index on DatasetModel is also modified to include name. Existing rows +have their name copied from URI. + +While not strictly related to other changes, the index name on DatasetAliasModel +is also renamed. Index names are scoped to the entire database. Airflow generally +includes the table's name to manually scope the index, but ``idx_uri_unique`` +(on DatasetModel) and ``idx_name_unique`` (on DatasetAliasModel) do not do this. +The one on DatasetModel is already renamed in this PR (to include name), so we +also rename the one on DatasetAliasModel here for consistency. + +Revision ID: 0d9e73a75ee4 +Revises: 16cbcb1c8c36 +Create Date: 2024-08-13 09:45:32.213222 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op +from sqlalchemy.orm import Session + +# revision identifiers, used by Alembic. +revision = "0d9e73a75ee4" +down_revision = "16cbcb1c8c36" +branch_labels = None +depends_on = None +airflow_version = "3.0.0" + +_STRING_COLUMN_TYPE = sa.String(length=1500).with_variant( + sa.String(length=1500, collation="latin1_general_cs"), + dialect_name="mysql", +) + + +def upgrade(): + # Fix index name on DatasetAlias. + with op.batch_alter_table("dataset_alias", schema=None) as batch_op: + batch_op.drop_index("idx_name_unique") + batch_op.create_index("idx_dataset_alias_name_unique", ["name"], unique=True) + # Add 'name' column. Set it to nullable for now. + with op.batch_alter_table("dataset", schema=None) as batch_op: + batch_op.add_column(sa.Column("name", _STRING_COLUMN_TYPE)) + batch_op.add_column(sa.Column("group", _STRING_COLUMN_TYPE, default=str, nullable=False)) + # Fill name from uri column. + Session(bind=op.get_bind()).execute(sa.text("update dataset set name=uri")) + # Set the name column non-nullable. + # Now with values in there, we can create the new unique constraint and index. + # Due to MySQL restrictions, we are also reducing the length on uri. + with op.batch_alter_table("dataset", schema=None) as batch_op: + batch_op.alter_column("name", existing_type=_STRING_COLUMN_TYPE, nullable=False) + batch_op.alter_column("uri", type_=_STRING_COLUMN_TYPE, nullable=False) + batch_op.drop_index("idx_uri_unique") + batch_op.create_index("idx_dataset_name_uri_unique", ["name", "uri"], unique=True) + + +def downgrade(): + with op.batch_alter_table("dataset", schema=None) as batch_op: + batch_op.drop_index("idx_dataset_name_uri_unique") + batch_op.create_index("idx_uri_unique", ["uri"], unique=True) + with op.batch_alter_table("dataset", schema=None) as batch_op: + batch_op.drop_column("group") + batch_op.drop_column("name") + batch_op.alter_column( + "uri", + type_=sa.String(length=3000).with_variant( + sa.String(length=3000, collation="latin1_general_cs"), + dialect_name="mysql", + ), + nullable=False, + ) + with op.batch_alter_table("dataset_alias", schema=None) as batch_op: + batch_op.drop_index("idx_dataset_alias_name_unique") + batch_op.create_index("idx_name_unique", ["name"], unique=True) diff --git a/airflow/models/asset.py b/airflow/models/asset.py index b99aa86f2c889..fb56bc4bf1ecf 100644 --- a/airflow/models/asset.py +++ b/airflow/models/asset.py @@ -106,7 +106,7 @@ class AssetAliasModel(Base): __tablename__ = "dataset_alias" __table_args__ = ( - Index("idx_name_unique", name, unique=True), + Index("idx_dataset_alias_name_unique", name, unique=True), {"sqlite_autoincrement": True}, # ensures PK values not reused ) @@ -151,10 +151,22 @@ class AssetModel(Base): """ id = Column(Integer, primary_key=True, autoincrement=True) + name = Column( + String(length=1500).with_variant( + String( + length=1500, + # latin1 allows for more indexed length in mysql + # and this field should only be ascii chars + collation="latin1_general_cs", + ), + "mysql", + ), + nullable=False, + ) uri = Column( - String(length=3000).with_variant( + String(length=1500).with_variant( String( - length=3000, + length=1500, # latin1 allows for more indexed length in mysql # and this field should only be ascii chars collation="latin1_general_cs", @@ -163,7 +175,21 @@ class AssetModel(Base): ), nullable=False, ) + group = Column( + String(length=1500).with_variant( + String( + length=1500, + # latin1 allows for more indexed length in mysql + # and this field should only be ascii chars + collation="latin1_general_cs", + ), + "mysql", + ), + default=str, + nullable=False, + ) extra = Column(sqlalchemy_jsonfield.JSONField(json=json), nullable=False, default={}) + created_at = Column(UtcDateTime, default=timezone.utcnow, nullable=False) updated_at = Column(UtcDateTime, default=timezone.utcnow, onupdate=timezone.utcnow, nullable=False) is_orphaned = Column(Boolean, default=False, nullable=False, server_default="0") @@ -173,7 +199,7 @@ class AssetModel(Base): __tablename__ = "dataset" __table_args__ = ( - Index("idx_uri_unique", uri, unique=True), + Index("idx_dataset_name_uri_unique", name, uri, unique=True), {"sqlite_autoincrement": True}, # ensures PK values not reused ) @@ -189,16 +215,15 @@ def __init__(self, uri: str, **kwargs): parsed = urlsplit(uri) if parsed.scheme and parsed.scheme.lower() == "airflow": raise ValueError("Scheme `airflow` is reserved.") - super().__init__(uri=uri, **kwargs) + super().__init__(name=uri, uri=uri, **kwargs) def __eq__(self, other): if isinstance(other, (self.__class__, Asset)): - return self.uri == other.uri - else: - return NotImplemented + return self.name == other.name and self.uri == other.uri + return NotImplemented def __hash__(self): - return hash(self.uri) + return hash((self.name, self.uri)) def __repr__(self): return f"{self.__class__.__name__}(uri={self.uri!r}, extra={self.extra!r})" diff --git a/airflow/utils/db.py b/airflow/utils/db.py index 512195c3aa963..8a254d4fef4d8 100644 --- a/airflow/utils/db.py +++ b/airflow/utils/db.py @@ -96,7 +96,7 @@ class MappedClassProtocol(Protocol): "2.9.0": "1949afb29106", "2.9.2": "686269002441", "2.10.0": "22ed7efa9da2", - "3.0.0": "16cbcb1c8c36", + "3.0.0": "0d9e73a75ee4", } diff --git a/docs/apache-airflow/img/airflow_erd.sha256 b/docs/apache-airflow/img/airflow_erd.sha256 index 237c598ec1dc8..e4a952da1b9fd 100644 --- a/docs/apache-airflow/img/airflow_erd.sha256 +++ b/docs/apache-airflow/img/airflow_erd.sha256 @@ -1 +1 @@ -f4379048d3f13f35aaba824c00450c17ad4deea9af82b5498d755a12f8a85a37 \ No newline at end of file +c33e9a583a5b29eb748ebd50e117643e11bcb2a9b61ec017efd690621e22769b \ No newline at end of file diff --git a/docs/apache-airflow/img/airflow_erd.svg b/docs/apache-airflow/img/airflow_erd.svg index 65f94c58ad24a..76fbd8f841f25 100644 --- a/docs/apache-airflow/img/airflow_erd.svg +++ b/docs/apache-airflow/img/airflow_erd.svg @@ -4,11 +4,11 @@ - - + + %3 - + log @@ -527,244 +527,254 @@ dataset_alias_dataset - -dataset_alias_dataset - -alias_id - - [INTEGER] - NOT NULL - -dataset_id - - [INTEGER] - NOT NULL + +dataset_alias_dataset + +alias_id + + [INTEGER] + NOT NULL + +dataset_id + + [INTEGER] + NOT NULL dataset_alias--dataset_alias_dataset - -0..N -1 + +0..N +1 dataset_alias--dataset_alias_dataset - -0..N -1 + +0..N +1 dataset_alias_dataset_event - -dataset_alias_dataset_event - -alias_id - - [INTEGER] - NOT NULL - -event_id - - [INTEGER] - NOT NULL + +dataset_alias_dataset_event + +alias_id + + [INTEGER] + NOT NULL + +event_id + + [INTEGER] + NOT NULL dataset_alias--dataset_alias_dataset_event - -0..N -1 + +0..N +1 dataset_alias--dataset_alias_dataset_event - -0..N -1 + +0..N +1 dag_schedule_dataset_alias_reference - -dag_schedule_dataset_alias_reference - -alias_id - - [INTEGER] - NOT NULL - -dag_id - - [VARCHAR(250)] - NOT NULL - -created_at - - [TIMESTAMP] - NOT NULL - -updated_at - - [TIMESTAMP] - NOT NULL + +dag_schedule_dataset_alias_reference + +alias_id + + [INTEGER] + NOT NULL + +dag_id + + [VARCHAR(250)] + NOT NULL + +created_at + + [TIMESTAMP] + NOT NULL + +updated_at + + [TIMESTAMP] + NOT NULL dataset_alias--dag_schedule_dataset_alias_reference - -0..N -1 + +0..N +1 dataset - -dataset - -id - - [INTEGER] - NOT NULL - -created_at - - [TIMESTAMP] - NOT NULL - -extra - - [JSON] - NOT NULL - -is_orphaned - - [BOOLEAN] - NOT NULL - -updated_at - - [TIMESTAMP] - NOT NULL - -uri - - [VARCHAR(3000)] - NOT NULL + +dataset + +id + + [INTEGER] + NOT NULL + +created_at + + [TIMESTAMP] + NOT NULL + +extra + + [JSON] + NOT NULL + +group + + [VARCHAR(1500)] + NOT NULL + +is_orphaned + + [BOOLEAN] + NOT NULL + +name + + [VARCHAR(1500)] + NOT NULL + +updated_at + + [TIMESTAMP] + NOT NULL + +uri + + [VARCHAR(1500)] + NOT NULL dataset--dataset_alias_dataset - -0..N -1 + +0..N +1 dataset--dataset_alias_dataset - -0..N -1 + +0..N +1 dag_schedule_dataset_reference - -dag_schedule_dataset_reference - -dag_id - - [VARCHAR(250)] - NOT NULL - -dataset_id - - [INTEGER] - NOT NULL - -created_at - - [TIMESTAMP] - NOT NULL - -updated_at - - [TIMESTAMP] - NOT NULL + +dag_schedule_dataset_reference + +dag_id + + [VARCHAR(250)] + NOT NULL + +dataset_id + + [INTEGER] + NOT NULL + +created_at + + [TIMESTAMP] + NOT NULL + +updated_at + + [TIMESTAMP] + NOT NULL dataset--dag_schedule_dataset_reference - -0..N -1 + +0..N +1 task_outlet_dataset_reference - -task_outlet_dataset_reference - -dag_id - - [VARCHAR(250)] - NOT NULL - -dataset_id - - [INTEGER] - NOT NULL - -task_id - - [VARCHAR(250)] - NOT NULL - -created_at - - [TIMESTAMP] - NOT NULL - -updated_at - - [TIMESTAMP] - NOT NULL + +task_outlet_dataset_reference + +dag_id + + [VARCHAR(250)] + NOT NULL + +dataset_id + + [INTEGER] + NOT NULL + +task_id + + [VARCHAR(250)] + NOT NULL + +created_at + + [TIMESTAMP] + NOT NULL + +updated_at + + [TIMESTAMP] + NOT NULL dataset--task_outlet_dataset_reference - -0..N -1 + +0..N +1 dataset_dag_run_queue - -dataset_dag_run_queue - -dataset_id - - [INTEGER] - NOT NULL - -target_dag_id - - [VARCHAR(250)] - NOT NULL - -created_at - - [TIMESTAMP] - NOT NULL + +dataset_dag_run_queue + +dataset_id + + [INTEGER] + NOT NULL + +target_dag_id + + [VARCHAR(250)] + NOT NULL + +created_at + + [TIMESTAMP] + NOT NULL dataset--dataset_dag_run_queue - -0..N -1 + +0..N +1 @@ -811,39 +821,39 @@ dataset_event--dataset_alias_dataset_event - -0..N -1 + +0..N +1 dataset_event--dataset_alias_dataset_event - -0..N -1 + +0..N +1 dagrun_dataset_event - -dagrun_dataset_event - -dag_run_id - - [INTEGER] - NOT NULL - -event_id - - [INTEGER] - NOT NULL + +dagrun_dataset_event + +dag_run_id + + [INTEGER] + NOT NULL + +event_id + + [INTEGER] + NOT NULL dataset_event--dagrun_dataset_event - -0..N -1 + +0..N +1 @@ -962,114 +972,114 @@ dag--dag_schedule_dataset_alias_reference - -0..N -1 + +0..N +1 dag--dag_schedule_dataset_reference - -0..N -1 + +0..N +1 dag--task_outlet_dataset_reference - -0..N -1 + +0..N +1 dag--dataset_dag_run_queue - -0..N -1 + +0..N +1 dag_tag - -dag_tag - -dag_id - - [VARCHAR(250)] - NOT NULL - -name - - [VARCHAR(100)] - NOT NULL + +dag_tag + +dag_id + + [VARCHAR(250)] + NOT NULL + +name + + [VARCHAR(100)] + NOT NULL dag--dag_tag - -0..N -1 + +0..N +1 dag_owner_attributes - -dag_owner_attributes - -dag_id - - [VARCHAR(250)] - NOT NULL - -owner - - [VARCHAR(500)] - NOT NULL - -link - - [VARCHAR(500)] - NOT NULL + +dag_owner_attributes + +dag_id + + [VARCHAR(250)] + NOT NULL + +owner + + [VARCHAR(500)] + NOT NULL + +link + + [VARCHAR(500)] + NOT NULL dag--dag_owner_attributes - -0..N -1 + +0..N +1 dag_warning - -dag_warning - -dag_id - - [VARCHAR(250)] - NOT NULL - -warning_type - - [VARCHAR(50)] - NOT NULL - -message - - [TEXT] - NOT NULL - -timestamp - - [TIMESTAMP] - NOT NULL + +dag_warning + +dag_id + + [VARCHAR(250)] + NOT NULL + +warning_type + + [VARCHAR(50)] + NOT NULL + +message + + [TEXT] + NOT NULL + +timestamp + + [TIMESTAMP] + NOT NULL dag--dag_warning - -0..N -1 + +0..N +1 @@ -1199,813 +1209,813 @@ dag_run--dagrun_dataset_event - -0..N -1 + +0..N +1 task_instance - -task_instance - -dag_id - - [VARCHAR(250)] - NOT NULL - -map_index - - [INTEGER] - NOT NULL - -run_id - - [VARCHAR(250)] - NOT NULL - -task_id - - [VARCHAR(250)] - NOT NULL - -custom_operator_name - - [VARCHAR(1000)] - -duration - - [DOUBLE_PRECISION] - -end_date - - [TIMESTAMP] - -executor - - [VARCHAR(1000)] - -executor_config - - [BYTEA] - -external_executor_id - - [VARCHAR(250)] - -hostname - - [VARCHAR(1000)] - -job_id - - [INTEGER] - -max_tries - - [INTEGER] - -next_kwargs - - [JSON] - -next_method - - [VARCHAR(1000)] - -operator - - [VARCHAR(1000)] - -pid - - [INTEGER] - -pool - - [VARCHAR(256)] - NOT NULL - -pool_slots - - [INTEGER] - NOT NULL - -priority_weight - - [INTEGER] - -queue - - [VARCHAR(256)] - -queued_by_job_id - - [INTEGER] - -queued_dttm - - [TIMESTAMP] - -rendered_map_index - - [VARCHAR(250)] - -start_date - - [TIMESTAMP] - -state - - [VARCHAR(20)] - -task_display_name - - [VARCHAR(2000)] - -trigger_id - - [INTEGER] - -trigger_timeout - - [TIMESTAMP] - -try_number - - [INTEGER] - -unixname - - [VARCHAR(1000)] - -updated_at - - [TIMESTAMP] + +task_instance + +dag_id + + [VARCHAR(250)] + NOT NULL + +map_index + + [INTEGER] + NOT NULL + +run_id + + [VARCHAR(250)] + NOT NULL + +task_id + + [VARCHAR(250)] + NOT NULL + +custom_operator_name + + [VARCHAR(1000)] + +duration + + [DOUBLE_PRECISION] + +end_date + + [TIMESTAMP] + +executor + + [VARCHAR(1000)] + +executor_config + + [BYTEA] + +external_executor_id + + [VARCHAR(250)] + +hostname + + [VARCHAR(1000)] + +job_id + + [INTEGER] + +max_tries + + [INTEGER] + +next_kwargs + + [JSON] + +next_method + + [VARCHAR(1000)] + +operator + + [VARCHAR(1000)] + +pid + + [INTEGER] + +pool + + [VARCHAR(256)] + NOT NULL + +pool_slots + + [INTEGER] + NOT NULL + +priority_weight + + [INTEGER] + +queue + + [VARCHAR(256)] + +queued_by_job_id + + [INTEGER] + +queued_dttm + + [TIMESTAMP] + +rendered_map_index + + [VARCHAR(250)] + +start_date + + [TIMESTAMP] + +state + + [VARCHAR(20)] + +task_display_name + + [VARCHAR(2000)] + +trigger_id + + [INTEGER] + +trigger_timeout + + [TIMESTAMP] + +try_number + + [INTEGER] + +unixname + + [VARCHAR(1000)] + +updated_at + + [TIMESTAMP] dag_run--task_instance - -0..N -1 + +0..N +1 dag_run--task_instance - -0..N -1 + +0..N +1 dag_run_note - -dag_run_note - -dag_run_id - - [INTEGER] - NOT NULL - -content - - [VARCHAR(1000)] - -created_at - - [TIMESTAMP] - NOT NULL - -updated_at - - [TIMESTAMP] - NOT NULL - -user_id - - [INTEGER] + +dag_run_note + +dag_run_id + + [INTEGER] + NOT NULL + +content + + [VARCHAR(1000)] + +created_at + + [TIMESTAMP] + NOT NULL + +updated_at + + [TIMESTAMP] + NOT NULL + +user_id + + [INTEGER] dag_run--dag_run_note - -1 -1 + +1 +1 task_reschedule - -task_reschedule - -id - - [INTEGER] - NOT NULL - -dag_id - - [VARCHAR(250)] - NOT NULL - -duration - - [INTEGER] - NOT NULL - -end_date - - [TIMESTAMP] - NOT NULL - -map_index - - [INTEGER] - NOT NULL - -reschedule_date - - [TIMESTAMP] - NOT NULL - -run_id - - [VARCHAR(250)] - NOT NULL - -start_date - - [TIMESTAMP] - NOT NULL - -task_id - - [VARCHAR(250)] - NOT NULL - -try_number - - [INTEGER] - NOT NULL + +task_reschedule + +id + + [INTEGER] + NOT NULL + +dag_id + + [VARCHAR(250)] + NOT NULL + +duration + + [INTEGER] + NOT NULL + +end_date + + [TIMESTAMP] + NOT NULL + +map_index + + [INTEGER] + NOT NULL + +reschedule_date + + [TIMESTAMP] + NOT NULL + +run_id + + [VARCHAR(250)] + NOT NULL + +start_date + + [TIMESTAMP] + NOT NULL + +task_id + + [VARCHAR(250)] + NOT NULL + +try_number + + [INTEGER] + NOT NULL dag_run--task_reschedule - -0..N -1 + +0..N +1 dag_run--task_reschedule - -0..N -1 + +0..N +1 task_instance--task_reschedule - -0..N -1 + +0..N +1 task_instance--task_reschedule - -0..N -1 + +0..N +1 task_instance--task_reschedule - -0..N -1 + +0..N +1 task_instance--task_reschedule - -0..N -1 + +0..N +1 rendered_task_instance_fields - -rendered_task_instance_fields - -dag_id - - [VARCHAR(250)] - NOT NULL - -map_index - - [INTEGER] - NOT NULL - -run_id - - [VARCHAR(250)] - NOT NULL - -task_id - - [VARCHAR(250)] - NOT NULL - -k8s_pod_yaml - - [JSON] - -rendered_fields - - [JSON] - NOT NULL + +rendered_task_instance_fields + +dag_id + + [VARCHAR(250)] + NOT NULL + +map_index + + [INTEGER] + NOT NULL + +run_id + + [VARCHAR(250)] + NOT NULL + +task_id + + [VARCHAR(250)] + NOT NULL + +k8s_pod_yaml + + [JSON] + +rendered_fields + + [JSON] + NOT NULL task_instance--rendered_task_instance_fields - -0..N -1 + +0..N +1 task_instance--rendered_task_instance_fields - -0..N -1 + +0..N +1 task_instance--rendered_task_instance_fields - -0..N -1 + +0..N +1 task_instance--rendered_task_instance_fields - -0..N -1 + +0..N +1 task_fail - -task_fail - -id - - [INTEGER] - NOT NULL - -dag_id - - [VARCHAR(250)] - NOT NULL - -duration - - [INTEGER] - -end_date - - [TIMESTAMP] - -map_index - - [INTEGER] - NOT NULL - -run_id - - [VARCHAR(250)] - NOT NULL - -start_date - - [TIMESTAMP] - -task_id - - [VARCHAR(250)] - NOT NULL + +task_fail + +id + + [INTEGER] + NOT NULL + +dag_id + + [VARCHAR(250)] + NOT NULL + +duration + + [INTEGER] + +end_date + + [TIMESTAMP] + +map_index + + [INTEGER] + NOT NULL + +run_id + + [VARCHAR(250)] + NOT NULL + +start_date + + [TIMESTAMP] + +task_id + + [VARCHAR(250)] + NOT NULL task_instance--task_fail - -0..N -1 + +0..N +1 task_instance--task_fail - -0..N -1 + +0..N +1 task_instance--task_fail - -0..N -1 + +0..N +1 task_instance--task_fail - -0..N -1 + +0..N +1 task_map - -task_map - -dag_id - - [VARCHAR(250)] - NOT NULL - -map_index - - [INTEGER] - NOT NULL - -run_id - - [VARCHAR(250)] - NOT NULL - -task_id - - [VARCHAR(250)] - NOT NULL - -keys - - [JSON] - -length - - [INTEGER] - NOT NULL + +task_map + +dag_id + + [VARCHAR(250)] + NOT NULL + +map_index + + [INTEGER] + NOT NULL + +run_id + + [VARCHAR(250)] + NOT NULL + +task_id + + [VARCHAR(250)] + NOT NULL + +keys + + [JSON] + +length + + [INTEGER] + NOT NULL task_instance--task_map - -0..N -1 + +0..N +1 task_instance--task_map - -0..N -1 + +0..N +1 task_instance--task_map - -0..N -1 + +0..N +1 task_instance--task_map - -0..N -1 + +0..N +1 xcom - -xcom - -dag_run_id - - [INTEGER] - NOT NULL - -key - - [VARCHAR(512)] - NOT NULL - -map_index - - [INTEGER] - NOT NULL - -task_id - - [VARCHAR(250)] - NOT NULL - -dag_id - - [VARCHAR(250)] - NOT NULL - -run_id - - [VARCHAR(250)] - NOT NULL - -timestamp - - [TIMESTAMP] - NOT NULL - -value - - [BYTEA] + +xcom + +dag_run_id + + [INTEGER] + NOT NULL + +key + + [VARCHAR(512)] + NOT NULL + +map_index + + [INTEGER] + NOT NULL + +task_id + + [VARCHAR(250)] + NOT NULL + +dag_id + + [VARCHAR(250)] + NOT NULL + +run_id + + [VARCHAR(250)] + NOT NULL + +timestamp + + [TIMESTAMP] + NOT NULL + +value + + [BYTEA] task_instance--xcom - -0..N -1 + +0..N +1 task_instance--xcom - -0..N -1 + +0..N +1 task_instance--xcom - -0..N -1 + +0..N +1 task_instance--xcom - -0..N -1 + +0..N +1 task_instance_note - -task_instance_note - -dag_id - - [VARCHAR(250)] - NOT NULL - -map_index - - [INTEGER] - NOT NULL - -run_id - - [VARCHAR(250)] - NOT NULL - -task_id - - [VARCHAR(250)] - NOT NULL - -content - - [VARCHAR(1000)] - -created_at - - [TIMESTAMP] - NOT NULL - -updated_at - - [TIMESTAMP] - NOT NULL - -user_id - - [INTEGER] + +task_instance_note + +dag_id + + [VARCHAR(250)] + NOT NULL + +map_index + + [INTEGER] + NOT NULL + +run_id + + [VARCHAR(250)] + NOT NULL + +task_id + + [VARCHAR(250)] + NOT NULL + +content + + [VARCHAR(1000)] + +created_at + + [TIMESTAMP] + NOT NULL + +updated_at + + [TIMESTAMP] + NOT NULL + +user_id + + [INTEGER] task_instance--task_instance_note - -0..N -1 + +0..N +1 task_instance--task_instance_note - -0..N -1 + +0..N +1 task_instance--task_instance_note - -0..N -1 + +0..N +1 task_instance--task_instance_note - -0..N -1 + +0..N +1 task_instance_history - -task_instance_history - -id - - [INTEGER] - NOT NULL - -custom_operator_name - - [VARCHAR(1000)] - -dag_id - - [VARCHAR(250)] - NOT NULL - -duration - - [DOUBLE_PRECISION] - -end_date - - [TIMESTAMP] - -executor - - [VARCHAR(1000)] - -executor_config - - [BYTEA] - -external_executor_id - - [VARCHAR(250)] - -hostname - - [VARCHAR(1000)] - -job_id - - [INTEGER] - -map_index - - [INTEGER] - NOT NULL - -max_tries - - [INTEGER] - -next_kwargs - - [JSON] - -next_method - - [VARCHAR(1000)] - -operator - - [VARCHAR(1000)] - -pid - - [INTEGER] - -pool - - [VARCHAR(256)] - NOT NULL - -pool_slots - - [INTEGER] - NOT NULL - -priority_weight - - [INTEGER] - -queue - - [VARCHAR(256)] - -queued_by_job_id - - [INTEGER] - -queued_dttm - - [TIMESTAMP] - -rendered_map_index - - [VARCHAR(250)] - -run_id - - [VARCHAR(250)] - NOT NULL - -start_date - - [TIMESTAMP] - -state - - [VARCHAR(20)] - -task_display_name - - [VARCHAR(2000)] - -task_id - - [VARCHAR(250)] - NOT NULL - -trigger_id - - [INTEGER] - -trigger_timeout - - [TIMESTAMP] - -try_number - - [INTEGER] - NOT NULL - -unixname - - [VARCHAR(1000)] - -updated_at - - [TIMESTAMP] + +task_instance_history + +id + + [INTEGER] + NOT NULL + +custom_operator_name + + [VARCHAR(1000)] + +dag_id + + [VARCHAR(250)] + NOT NULL + +duration + + [DOUBLE_PRECISION] + +end_date + + [TIMESTAMP] + +executor + + [VARCHAR(1000)] + +executor_config + + [BYTEA] + +external_executor_id + + [VARCHAR(250)] + +hostname + + [VARCHAR(1000)] + +job_id + + [INTEGER] + +map_index + + [INTEGER] + NOT NULL + +max_tries + + [INTEGER] + +next_kwargs + + [JSON] + +next_method + + [VARCHAR(1000)] + +operator + + [VARCHAR(1000)] + +pid + + [INTEGER] + +pool + + [VARCHAR(256)] + NOT NULL + +pool_slots + + [INTEGER] + NOT NULL + +priority_weight + + [INTEGER] + +queue + + [VARCHAR(256)] + +queued_by_job_id + + [INTEGER] + +queued_dttm + + [TIMESTAMP] + +rendered_map_index + + [VARCHAR(250)] + +run_id + + [VARCHAR(250)] + NOT NULL + +start_date + + [TIMESTAMP] + +state + + [VARCHAR(20)] + +task_display_name + + [VARCHAR(2000)] + +task_id + + [VARCHAR(250)] + NOT NULL + +trigger_id + + [INTEGER] + +trigger_timeout + + [TIMESTAMP] + +try_number + + [INTEGER] + NOT NULL + +unixname + + [VARCHAR(1000)] + +updated_at + + [TIMESTAMP] task_instance--task_instance_history - -0..N -1 + +0..N +1 task_instance--task_instance_history - -0..N -1 + +0..N +1 task_instance--task_instance_history - -0..N -1 + +0..N +1 task_instance--task_instance_history - -0..N -1 + +0..N +1 @@ -2040,325 +2050,325 @@ trigger--task_instance - -0..N -{0,1} + +0..N +{0,1} session - -session - -id - - [INTEGER] - NOT NULL - -data - - [BYTEA] - -expiry - - [TIMESTAMP] - -session_id - - [VARCHAR(255)] + +session + +id + + [INTEGER] + NOT NULL + +data + + [BYTEA] + +expiry + + [TIMESTAMP] + +session_id + + [VARCHAR(255)] alembic_version - -alembic_version - -version_num - - [VARCHAR(32)] - NOT NULL + +alembic_version + +version_num + + [VARCHAR(32)] + NOT NULL ab_user - -ab_user + +ab_user + +id + + [INTEGER] + NOT NULL -id - - [INTEGER] - NOT NULL +active + + [BOOLEAN] -active - - [BOOLEAN] +changed_by_fk + + [INTEGER] -changed_by_fk - - [INTEGER] +changed_on + + [TIMESTAMP] -changed_on - - [TIMESTAMP] +created_by_fk + + [INTEGER] -created_by_fk - - [INTEGER] +created_on + + [TIMESTAMP] -created_on - - [TIMESTAMP] +email + + [VARCHAR(512)] + NOT NULL -email - - [VARCHAR(512)] - NOT NULL +fail_login_count + + [INTEGER] -fail_login_count - - [INTEGER] +first_name + + [VARCHAR(256)] + NOT NULL -first_name - - [VARCHAR(256)] - NOT NULL +last_login + + [TIMESTAMP] -last_login - - [TIMESTAMP] +last_name + + [VARCHAR(256)] + NOT NULL -last_name - - [VARCHAR(256)] - NOT NULL +login_count + + [INTEGER] -login_count - - [INTEGER] +password + + [VARCHAR(256)] -password - - [VARCHAR(256)] - -username - - [VARCHAR(512)] - NOT NULL +username + + [VARCHAR(512)] + NOT NULL ab_user--ab_user - -0..N -{0,1} + +0..N +{0,1} ab_user--ab_user - -0..N -{0,1} + +0..N +{0,1} ab_user_role - -ab_user_role - -id - - [INTEGER] - NOT NULL - -role_id - - [INTEGER] - -user_id - - [INTEGER] + +ab_user_role + +id + + [INTEGER] + NOT NULL + +role_id + + [INTEGER] + +user_id + + [INTEGER] ab_user--ab_user_role - -0..N -{0,1} + +0..N +{0,1} ab_register_user - -ab_register_user + +ab_register_user + +id + + [INTEGER] + NOT NULL -id - - [INTEGER] - NOT NULL +email + + [VARCHAR(512)] + NOT NULL -email - - [VARCHAR(512)] - NOT NULL +first_name + + [VARCHAR(256)] + NOT NULL -first_name - - [VARCHAR(256)] - NOT NULL +last_name + + [VARCHAR(256)] + NOT NULL -last_name - - [VARCHAR(256)] - NOT NULL +password + + [VARCHAR(256)] -password - - [VARCHAR(256)] +registration_date + + [TIMESTAMP] -registration_date - - [TIMESTAMP] +registration_hash + + [VARCHAR(256)] -registration_hash - - [VARCHAR(256)] - -username - - [VARCHAR(512)] - NOT NULL +username + + [VARCHAR(512)] + NOT NULL ab_permission - -ab_permission + +ab_permission + +id + + [INTEGER] + NOT NULL -id - - [INTEGER] - NOT NULL - -name - - [VARCHAR(100)] - NOT NULL +name + + [VARCHAR(100)] + NOT NULL ab_permission_view - -ab_permission_view + +ab_permission_view + +id + + [INTEGER] + NOT NULL -id - - [INTEGER] - NOT NULL +permission_id + + [INTEGER] -permission_id - - [INTEGER] - -view_menu_id - - [INTEGER] +view_menu_id + + [INTEGER] ab_permission--ab_permission_view - -0..N -{0,1} + +0..N +{0,1} ab_permission_view_role - -ab_permission_view_role - -id - - [INTEGER] - NOT NULL - -permission_view_id - - [INTEGER] - -role_id - - [INTEGER] + +ab_permission_view_role + +id + + [INTEGER] + NOT NULL + +permission_view_id + + [INTEGER] + +role_id + + [INTEGER] ab_permission_view--ab_permission_view_role - -0..N -{0,1} + +0..N +{0,1} ab_view_menu - -ab_view_menu + +ab_view_menu + +id + + [INTEGER] + NOT NULL -id - - [INTEGER] - NOT NULL - -name - - [VARCHAR(250)] - NOT NULL +name + + [VARCHAR(250)] + NOT NULL ab_view_menu--ab_permission_view - -0..N -{0,1} + +0..N +{0,1} ab_role - -ab_role + +ab_role + +id + + [INTEGER] + NOT NULL -id - - [INTEGER] - NOT NULL - -name - - [VARCHAR(64)] - NOT NULL +name + + [VARCHAR(64)] + NOT NULL ab_role--ab_user_role - -0..N -{0,1} + +0..N +{0,1} ab_role--ab_permission_view_role - -0..N -{0,1} + +0..N +{0,1} alembic_version_fab - -alembic_version_fab - -version_num - - [VARCHAR(32)] - NOT NULL + +alembic_version_fab + +version_num + + [VARCHAR(32)] + NOT NULL diff --git a/docs/apache-airflow/migrations-ref.rst b/docs/apache-airflow/migrations-ref.rst index ded2b290b5a4e..a547d03d75be6 100644 --- a/docs/apache-airflow/migrations-ref.rst +++ b/docs/apache-airflow/migrations-ref.rst @@ -39,7 +39,9 @@ Here's the list of all the Database Migrations that are executed via when you ru +-------------------------+------------------+-------------------+--------------------------------------------------------------+ | Revision ID | Revises ID | Airflow Version | Description | +=========================+==================+===================+==============================================================+ -| ``16cbcb1c8c36`` (head) | ``522625f6d606`` | ``3.0.0`` | Remove redundant index. | +| ``0d9e73a75ee4`` (head) | ``16cbcb1c8c36`` | ``3.0.0`` | Add name and group fields to DatasetModel. | ++-------------------------+------------------+-------------------+--------------------------------------------------------------+ +| ``16cbcb1c8c36`` | ``522625f6d606`` | ``3.0.0`` | Remove redundant index. | +-------------------------+------------------+-------------------+--------------------------------------------------------------+ | ``522625f6d606`` | ``1cdc775ca98f`` | ``3.0.0`` | Add tables for backfill. | +-------------------------+------------------+-------------------+--------------------------------------------------------------+ From 72dc70aceed13c9c5f25c2ae9c4816aac754b467 Mon Sep 17 00:00:00 2001 From: Ephraim Anierobi Date: Mon, 30 Sep 2024 12:10:22 +0100 Subject: [PATCH 073/802] Ensure consistent Seriailized DAG hashing (#42517) * Ensure consistent Seriailized DAG hashing The serialized DAG dictionary is not ordered correctly when creating hashes, and that causes inconsistent hashes, leading to frequent update of the serialized DAG table. Changes: Implemented sorting for serialized DAG dictionaries and nested structures to ensure consistent and predictable serialization order for hashing. Using `sort_keys` in `json.dumps` is not enough to sort the nested structures in the serialized DAG. Added serialize and deserialize methods for DagParam and ArgNotSet to allow for more structured serialization. Updated serialize_template_field to handle objects that implement the serialize method. This was done because of DagParam and ArgNotSet in the template fields. Previously, it produced an object, but with this change, it now serialises to a consistent object. * Move hashing to a method * fixup! Move hashing to a method * Add test --- airflow/models/param.py | 25 ++++++++++++++++++++++ airflow/models/serialized_dag.py | 32 ++++++++++++++++++++++++++--- airflow/serialization/helpers.py | 7 +++++-- airflow/utils/types.py | 8 ++++++++ tests/models/test_serialized_dag.py | 19 +++++++++++++++++ 5 files changed, 86 insertions(+), 5 deletions(-) diff --git a/airflow/models/param.py b/airflow/models/param.py index a4bbce2b768c6..895cd2af8bb42 100644 --- a/airflow/models/param.py +++ b/airflow/models/param.py @@ -290,6 +290,7 @@ def __init__(self, current_dag: DAG, name: str, default: Any = NOTSET): current_dag.params[name] = default self._name = name self._default = default + self.current_dag = current_dag def iter_references(self) -> Iterable[tuple[Operator, str]]: return () @@ -304,6 +305,30 @@ def resolve(self, context: Context, *, include_xcom: bool = True) -> Any: return context["params"][self._name] raise AirflowException(f"No value could be resolved for parameter {self._name}") + def serialize(self) -> dict: + """Serialize the DagParam object into a dictionary.""" + return { + "dag_id": self.current_dag.dag_id, + "name": self._name, + "default": self._default, + } + + @classmethod + def deserialize(cls, data: dict, dags: dict) -> DagParam: + """ + Deserializes the dictionary back into a DagParam object. + + :param data: The serialized representation of the DagParam. + :param dags: A dictionary of available DAGs to look up the DAG. + """ + dag_id = data["dag_id"] + # Retrieve the current DAG from the provided DAGs dictionary + current_dag = dags.get(dag_id) + if not current_dag: + raise ValueError(f"DAG with id {dag_id} not found.") + + return cls(current_dag=current_dag, name=data["name"], default=data["default"]) + def process_params( dag: DAG, diff --git a/airflow/models/serialized_dag.py b/airflow/models/serialized_dag.py index dec843451a98a..32be31d721e34 100644 --- a/airflow/models/serialized_dag.py +++ b/airflow/models/serialized_dag.py @@ -22,7 +22,7 @@ import logging import zlib from datetime import timedelta -from typing import TYPE_CHECKING, Collection +from typing import TYPE_CHECKING, Any, Collection import sqlalchemy_jsonfield from sqlalchemy import BigInteger, Column, Index, LargeBinary, String, and_, exc, or_, select @@ -114,9 +114,10 @@ def __init__(self, dag: DAG, processor_subdir: str | None = None) -> None: self.processor_subdir = processor_subdir dag_data = SerializedDAG.to_dict(dag) - dag_data_json = json.dumps(dag_data, sort_keys=True).encode("utf-8") + self.dag_hash = SerializedDagModel.hash(dag_data) - self.dag_hash = md5(dag_data_json).hexdigest() + # partially ordered json data + dag_data_json = json.dumps(dag_data, sort_keys=True).encode("utf-8") if COMPRESS_SERIALIZED_DAGS: self._data = None @@ -132,6 +133,30 @@ def __init__(self, dag: DAG, processor_subdir: str | None = None) -> None: def __repr__(self) -> str: return f"" + @classmethod + def hash(cls, dag_data): + """Hash the data to get the dag_hash.""" + dag_data = cls._sort_serialized_dag_dict(dag_data) + data_json = json.dumps(dag_data, sort_keys=True).encode("utf-8") + return md5(data_json).hexdigest() + + @classmethod + def _sort_serialized_dag_dict(cls, serialized_dag: Any): + """Recursively sort json_dict and its nested dictionaries and lists.""" + if isinstance(serialized_dag, dict): + return {k: cls._sort_serialized_dag_dict(v) for k, v in sorted(serialized_dag.items())} + elif isinstance(serialized_dag, list): + if all(isinstance(i, dict) for i in serialized_dag): + if all("task_id" in i.get("__var", {}) for i in serialized_dag): + return sorted( + [cls._sort_serialized_dag_dict(i) for i in serialized_dag], + key=lambda x: x["__var"]["task_id"], + ) + elif all(isinstance(item, str) for item in serialized_dag): + return sorted(serialized_dag) + return [cls._sort_serialized_dag_dict(i) for i in serialized_dag] + return serialized_dag + @classmethod @provide_session def write_dag( @@ -149,6 +174,7 @@ def write_dag( :param dag: a DAG to be written into database :param min_update_interval: minimal interval in seconds to update serialized DAG + :param processor_subdir: The dag directory of the processor :param session: ORM Session :returns: Boolean indicating if the DAG was written to the DB diff --git a/airflow/serialization/helpers.py b/airflow/serialization/helpers.py index 7f97e7f2ff299..85bf3a1cc551c 100644 --- a/airflow/serialization/helpers.py +++ b/airflow/serialization/helpers.py @@ -44,14 +44,17 @@ def is_jsonable(x): max_length = conf.getint("core", "max_templated_field_length") if not is_jsonable(template_field): - serialized = str(template_field) + try: + serialized = template_field.serialize() + except AttributeError: + serialized = str(template_field) if len(serialized) > max_length: rendered = redact(serialized, name) return ( "Truncated. You can change this behaviour in [core]max_templated_field_length. " f"{rendered[:max_length - 79]!r}... " ) - return str(template_field) + return serialized else: if not template_field: return template_field diff --git a/airflow/utils/types.py b/airflow/utils/types.py index 86af13832755b..a19b2534b03fb 100644 --- a/airflow/utils/types.py +++ b/airflow/utils/types.py @@ -41,6 +41,14 @@ def is_arg_passed(arg: Union[ArgNotSet, None] = NOTSET) -> bool: is_arg_passed(None) # True. """ + @staticmethod + def serialize(): + return "NOTSET" + + @classmethod + def deserialize(cls): + return cls + NOTSET = ArgNotSet() """Sentinel value for argument default. See ``ArgNotSet``.""" diff --git a/tests/models/test_serialized_dag.py b/tests/models/test_serialized_dag.py index b8fddc655dae5..d9a77e55edaf5 100644 --- a/tests/models/test_serialized_dag.py +++ b/tests/models/test_serialized_dag.py @@ -23,6 +23,7 @@ import pendulum import pytest +from sqlalchemy import select import airflow.example_dags as example_dags_module from airflow.assets import Asset @@ -264,3 +265,21 @@ def test_order_of_deps_is_consistent(self): # dag hash should not change without change in structure (we're in a loop) assert this_dag_hash == first_dag_hash + + def test_example_dag_hashes_are_always_consistent(self, session): + """ + This test asserts that the hashes of the example dags are always consistent. + """ + + def get_hash_set(): + example_dags = self._write_example_dags() + ordered_example_dags = dict(sorted(example_dags.items())) + hashes = set() + for dag_id in ordered_example_dags.keys(): + smd = session.execute(select(SDM.dag_hash).where(SDM.dag_id == dag_id)).one() + hashes.add(smd.dag_hash) + return hashes + + first_hashes = get_hash_set() + # assert that the hashes are the same + assert first_hashes == get_hash_set() From bece698b88b1b4c396ae95f9049e1fbd3acbdcc3 Mon Sep 17 00:00:00 2001 From: codecae Date: Mon, 30 Sep 2024 10:01:35 -0400 Subject: [PATCH 074/802] reduce eyestrain in dark mode with reduced contrast and saturation (#42567) * reduce eyestrain in dark mode with reduced contrast and saturation * feat: readjusted saturation --------- Co-authored-by: Curtis Bangert --- airflow/www/static/css/bootstrap-theme.css | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/airflow/www/static/css/bootstrap-theme.css b/airflow/www/static/css/bootstrap-theme.css index 921d795b96b12..94e1ee5887715 100644 --- a/airflow/www/static/css/bootstrap-theme.css +++ b/airflow/www/static/css/bootstrap-theme.css @@ -37,7 +37,7 @@ html { -webkit-text-size-adjust: 100%; } html[data-color-scheme="dark"] { - filter: invert(100%) hue-rotate(180deg); + filter: invert(100%) hue-rotate(180deg) saturate(90%) contrast(85%); } /* Default icons to not display until the data-color-scheme has been set */ From b602770c5f6125250b5bfbb818de494b13886bc8 Mon Sep 17 00:00:00 2001 From: Pierre Jeambrun Date: Mon, 30 Sep 2024 23:04:25 +0800 Subject: [PATCH 075/802] AIP-84 Migrate patch dags to FastAPI API (#42545) * AIP-84 Migrate patch dags to FastAPI API * Fix CI --- .../api_connexion/endpoints/dag_endpoint.py | 1 + airflow/api_fastapi/db/__init__.py | 16 +++ airflow/api_fastapi/db/common.py | 83 ++++++++++++ airflow/api_fastapi/{db.py => db/dags.py} | 54 +++----- airflow/api_fastapi/openapi/v1-generated.yaml | 123 +++++++++++++++++- airflow/api_fastapi/parameters.py | 50 +++++-- airflow/api_fastapi/views/public/dags.py | 93 ++++++++----- airflow/api_fastapi/views/ui/assets.py | 2 +- airflow/ui/openapi-gen/queries/common.ts | 3 + airflow/ui/openapi-gen/queries/queries.ts | 88 ++++++++++++- .../ui/openapi-gen/requests/services.gen.ts | 50 ++++++- airflow/ui/openapi-gen/requests/types.gen.ts | 44 +++++++ tests/api_fastapi/views/public/test_dags.py | 94 ++++++++++--- 13 files changed, 596 insertions(+), 105 deletions(-) create mode 100644 airflow/api_fastapi/db/__init__.py create mode 100644 airflow/api_fastapi/db/common.py rename airflow/api_fastapi/{db.py => db/dags.py} (55%) diff --git a/airflow/api_connexion/endpoints/dag_endpoint.py b/airflow/api_connexion/endpoints/dag_endpoint.py index 6fca5ae7c93d5..5d10a97dedce6 100644 --- a/airflow/api_connexion/endpoints/dag_endpoint.py +++ b/airflow/api_connexion/endpoints/dag_endpoint.py @@ -165,6 +165,7 @@ def patch_dag(*, dag_id: str, update_mask: UpdateMask = None, session: Session = return dag_schema.dump(dag) +@mark_fastapi_migration_done @security.requires_access_dag("PUT") @format_parameters({"limit": check_limit}) @action_logging diff --git a/airflow/api_fastapi/db/__init__.py b/airflow/api_fastapi/db/__init__.py new file mode 100644 index 0000000000000..13a83393a9124 --- /dev/null +++ b/airflow/api_fastapi/db/__init__.py @@ -0,0 +1,16 @@ +# 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. diff --git a/airflow/api_fastapi/db/common.py b/airflow/api_fastapi/db/common.py new file mode 100644 index 0000000000000..f611eaa64f07d --- /dev/null +++ b/airflow/api_fastapi/db/common.py @@ -0,0 +1,83 @@ +# 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 + +from typing import TYPE_CHECKING, Sequence + +from airflow.utils.db import get_query_count +from airflow.utils.session import NEW_SESSION, create_session, provide_session + +if TYPE_CHECKING: + from sqlalchemy.orm import Session + from sqlalchemy.sql import Select + + from airflow.api_fastapi.parameters import BaseParam + + +async def get_session() -> Session: + """ + Dependency for providing a session. + + For non route function please use the :class:`airflow.utils.session.provide_session` decorator. + + Example usage: + + .. code:: python + + @router.get("/your_path") + def your_route(session: Annotated[Session, Depends(get_session)]): + pass + """ + with create_session() as session: + yield session + + +def apply_filters_to_select(base_select: Select, filters: Sequence[BaseParam | None]) -> Select: + base_select = base_select + for filter in filters: + if filter is None: + continue + base_select = filter.to_orm(base_select) + + return base_select + + +@provide_session +def paginated_select( + base_select: Select, + filters: Sequence[BaseParam], + order_by: BaseParam | None = None, + offset: BaseParam | None = None, + limit: BaseParam | None = None, + session: Session = NEW_SESSION, +) -> Select: + base_select = apply_filters_to_select( + base_select, + filters, + ) + + total_entries = get_query_count(base_select, session=session) + + # TODO: Re-enable when permissions are handled. Readable / writable entities, + # for instance: + # readable_dags = get_auth_manager().get_permitted_dag_ids(user=g.user) + # dags_select = dags_select.where(DagModel.dag_id.in_(readable_dags)) + + base_select = apply_filters_to_select(base_select, [order_by, offset, limit]) + + return base_select, total_entries diff --git a/airflow/api_fastapi/db.py b/airflow/api_fastapi/db/dags.py similarity index 55% rename from airflow/api_fastapi/db.py rename to airflow/api_fastapi/db/dags.py index c3ed01a0aefec..7cd7cc9cd955d 100644 --- a/airflow/api_fastapi/db.py +++ b/airflow/api_fastapi/db/dags.py @@ -17,45 +17,10 @@ from __future__ import annotations -from typing import TYPE_CHECKING - from sqlalchemy import func, select +from airflow.models.dag import DagModel from airflow.models.dagrun import DagRun -from airflow.utils.session import create_session - -if TYPE_CHECKING: - from sqlalchemy.orm import Session - from sqlalchemy.sql import Select - - from airflow.api_fastapi.parameters import BaseParam - - -async def get_session() -> Session: - """ - Dependency for providing a session. - - For non route function please use the :class:`airflow.utils.session.provide_session` decorator. - - Example usage: - - .. code:: python - - @router.get("/your_path") - def your_route(session: Annotated[Session, Depends(get_session)]): - pass - """ - with create_session() as session: - yield session - - -def apply_filters_to_select(base_select: Select, filters: list[BaseParam]) -> Select: - select = base_select - for filter in filters: - select = filter.to_orm(select) - - return select - latest_dag_run_per_dag_id_cte = ( select(DagRun.dag_id, func.max(DagRun.start_date).label("start_date")) @@ -63,3 +28,20 @@ def apply_filters_to_select(base_select: Select, filters: list[BaseParam]) -> Se .group_by(DagRun.dag_id) .cte() ) + + +dags_select_with_latest_dag_run = ( + select(DagModel) + .join( + latest_dag_run_per_dag_id_cte, + DagModel.dag_id == latest_dag_run_per_dag_id_cte.c.dag_id, + isouter=True, + ) + .join( + DagRun, + DagRun.start_date == latest_dag_run_per_dag_id_cte.c.start_date + and DagRun.dag_id == latest_dag_run_per_dag_id_cte.c.dag_id, + isouter=True, + ) + .order_by(DagModel.dag_id) +) diff --git a/airflow/api_fastapi/openapi/v1-generated.yaml b/airflow/api_fastapi/openapi/v1-generated.yaml index c130f3162c6e6..a38a1021890d6 100644 --- a/airflow/api_fastapi/openapi/v1-generated.yaml +++ b/airflow/api_fastapi/openapi/v1-generated.yaml @@ -131,12 +131,133 @@ paths: application/json: schema: $ref: '#/components/schemas/HTTPValidationError' + patch: + tags: + - DAG + summary: Patch Dags + description: Patch multiple DAGs. + operationId: patch_dags_public_dags_patch + parameters: + - name: update_mask + in: query + required: false + schema: + anyOf: + - type: array + items: + type: string + - type: 'null' + title: Update Mask + - name: limit + in: query + required: false + schema: + type: integer + default: 100 + title: Limit + - name: offset + in: query + required: false + schema: + type: integer + default: 0 + title: Offset + - name: tags + in: query + required: false + schema: + type: array + items: + type: string + title: Tags + - name: owners + in: query + required: false + schema: + type: array + items: + type: string + title: Owners + - name: dag_id_pattern + in: query + required: false + schema: + anyOf: + - type: string + - type: 'null' + title: Dag Id Pattern + - name: only_active + in: query + required: false + schema: + type: boolean + default: true + title: Only Active + - name: paused + in: query + required: false + schema: + anyOf: + - type: boolean + - type: 'null' + title: Paused + - name: last_dag_run_state + in: query + required: false + schema: + anyOf: + - $ref: '#/components/schemas/DagRunState' + - type: 'null' + title: Last Dag Run State + requestBody: + required: true + content: + application/json: + schema: + $ref: '#/components/schemas/DAGPatchBody' + responses: + '200': + description: Successful Response + content: + application/json: + schema: + $ref: '#/components/schemas/DAGCollectionResponse' + '400': + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPExceptionResponse' + description: Bad Request + '401': + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPExceptionResponse' + description: Unauthorized + '403': + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPExceptionResponse' + description: Forbidden + '404': + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPExceptionResponse' + description: Not Found + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' /public/dags/{dag_id}: patch: tags: - DAG summary: Patch Dag - description: Update the specific DAG. + description: Patch the specific DAG. operationId: patch_dag_public_dags__dag_id__patch parameters: - name: dag_id diff --git a/airflow/api_fastapi/parameters.py b/airflow/api_fastapi/parameters.py index 09eea5f6e055b..504014602f3b5 100644 --- a/airflow/api_fastapi/parameters.py +++ b/airflow/api_fastapi/parameters.py @@ -37,9 +37,10 @@ class BaseParam(Generic[T], ABC): """Base class for filters.""" - def __init__(self) -> None: + def __init__(self, skip_none: bool = True) -> None: self.value: T | None = None self.attribute: ColumnElement | None = None + self.skip_none = skip_none @abstractmethod def to_orm(self, select: Select) -> Select: @@ -58,7 +59,7 @@ class _LimitFilter(BaseParam[int]): """Filter on the limit.""" def to_orm(self, select: Select) -> Select: - if self.value is None: + if self.value is None and self.skip_none: return select return select.limit(self.value) @@ -71,7 +72,7 @@ class _OffsetFilter(BaseParam[int]): """Filter on offset.""" def to_orm(self, select: Select) -> Select: - if self.value is None: + if self.value is None and self.skip_none: return select return select.offset(self.value) @@ -83,7 +84,7 @@ class _PausedFilter(BaseParam[bool]): """Filter on is_paused.""" def to_orm(self, select: Select) -> Select: - if self.value is None: + if self.value is None and self.skip_none: return select return select.where(DagModel.is_paused == self.value) @@ -95,7 +96,7 @@ class _OnlyActiveFilter(BaseParam[bool]): """Filter on is_active.""" def to_orm(self, select: Select) -> Select: - if self.value: + if self.value and self.skip_none: return select.where(DagModel.is_active == self.value) return select @@ -106,33 +107,40 @@ def depends(self, only_active: bool = True) -> _OnlyActiveFilter: class _SearchParam(BaseParam[str]): """Search on attribute.""" - def __init__(self, attribute: ColumnElement) -> None: - super().__init__() + def __init__(self, attribute: ColumnElement, skip_none: bool = True) -> None: + super().__init__(skip_none) self.attribute: ColumnElement = attribute def to_orm(self, select: Select) -> Select: - if self.value is None: + if self.value is None and self.skip_none: return select return select.where(self.attribute.ilike(f"%{self.value}")) + def transform_aliases(self, value: str | None) -> str | None: + if value == "~": + value = "%" + return value + class _DagIdPatternSearch(_SearchParam): """Search on dag_id.""" - def __init__(self) -> None: - super().__init__(DagModel.dag_id) + def __init__(self, skip_none: bool = True) -> None: + super().__init__(DagModel.dag_id, skip_none) def depends(self, dag_id_pattern: str | None = None) -> _DagIdPatternSearch: + dag_id_pattern = super().transform_aliases(dag_id_pattern) return self.set_value(dag_id_pattern) class _DagDisplayNamePatternSearch(_SearchParam): """Search on dag_display_name.""" - def __init__(self) -> None: - super().__init__(DagModel.dag_display_name) + def __init__(self, skip_none: bool = True) -> None: + super().__init__(DagModel.dag_display_name, skip_none) def depends(self, dag_display_name_pattern: str | None = None) -> _DagDisplayNamePatternSearch: + dag_display_name_pattern = super().transform_aliases(dag_display_name_pattern) return self.set_value(dag_display_name_pattern) @@ -149,6 +157,9 @@ def __init__(self, allowed_attrs: list[str]) -> None: self.allowed_attrs = allowed_attrs def to_orm(self, select: Select) -> Select: + if self.skip_none is False: + raise ValueError(f"Cannot set 'skip_none' to False on a {type(self)}") + if self.value is None: return select @@ -165,6 +176,10 @@ def to_orm(self, select: Select) -> Select: # MySQL does not support `nullslast`, and True/False ordering depends on the # database implementation. nullscheck = case((column.isnot(None), 0), else_=1) + + # Reset default sorting + select = select.order_by(None) + if self.value[0] == "-": return select.order_by(nullscheck, column.desc(), DagModel.dag_id.desc()) else: @@ -178,6 +193,9 @@ class _TagsFilter(BaseParam[List[str]]): """Filter on tags.""" def to_orm(self, select: Select) -> Select: + if self.skip_none is False: + raise ValueError(f"Cannot set 'skip_none' to False on a {type(self)}") + if not self.value: return select @@ -192,6 +210,9 @@ class _OwnersFilter(BaseParam[List[str]]): """Filter on owners.""" def to_orm(self, select: Select) -> Select: + if self.skip_none is False: + raise ValueError(f"Cannot set 'skip_none' to False on a {type(self)}") + if not self.value: return select @@ -206,7 +227,7 @@ class _LastDagRunStateFilter(BaseParam[DagRunState]): """Filter on the state of the latest DagRun.""" def to_orm(self, select: Select) -> Select: - if self.value is None: + if self.value is None and self.skip_none: return select return select.where(DagRun.state == self.value) @@ -223,6 +244,9 @@ def depends(self, last_dag_run_state: DagRunState | None = None) -> _LastDagRunS QueryDagDisplayNamePatternSearch = Annotated[ _DagDisplayNamePatternSearch, Depends(_DagDisplayNamePatternSearch().depends) ] +QueryDagIdPatternSearchWithNone = Annotated[ + _DagIdPatternSearch, Depends(_DagIdPatternSearch(skip_none=False).depends) +] QueryTagsFilter = Annotated[_TagsFilter, Depends(_TagsFilter().depends)] QueryOwnersFilter = Annotated[_OwnersFilter, Depends(_OwnersFilter().depends)] # DagRun diff --git a/airflow/api_fastapi/views/public/dags.py b/airflow/api_fastapi/views/public/dags.py index a9fe87eef0953..a6c25d6568c1e 100644 --- a/airflow/api_fastapi/views/public/dags.py +++ b/airflow/api_fastapi/views/public/dags.py @@ -18,15 +18,20 @@ from __future__ import annotations from fastapi import APIRouter, Depends, HTTPException, Query -from sqlalchemy import select +from sqlalchemy import update from sqlalchemy.orm import Session from typing_extensions import Annotated -from airflow.api_fastapi.db import apply_filters_to_select, get_session, latest_dag_run_per_dag_id_cte +from airflow.api_fastapi.db.common import ( + get_session, + paginated_select, +) +from airflow.api_fastapi.db.dags import dags_select_with_latest_dag_run from airflow.api_fastapi.openapi.exceptions import create_openapi_http_exception_doc from airflow.api_fastapi.parameters import ( QueryDagDisplayNamePatternSearch, QueryDagIdPatternSearch, + QueryDagIdPatternSearchWithNone, QueryLastDagRunStateFilter, QueryLimit, QueryOffset, @@ -38,8 +43,6 @@ ) from airflow.api_fastapi.serializers.dags import DAGCollectionResponse, DAGPatchBody, DAGResponse from airflow.models import DagModel -from airflow.models.dagrun import DagRun -from airflow.utils.db import get_query_count dags_router = APIRouter(tags=["DAG"]) @@ -66,35 +69,16 @@ async def get_dags( session: Annotated[Session, Depends(get_session)], ) -> DAGCollectionResponse: """Get all DAGs.""" - dags_query = ( - select(DagModel) - .join( - latest_dag_run_per_dag_id_cte, - DagModel.dag_id == latest_dag_run_per_dag_id_cte.c.dag_id, - isouter=True, - ) - .join( - DagRun, - DagRun.start_date == latest_dag_run_per_dag_id_cte.c.start_date - and DagRun.dag_id == latest_dag_run_per_dag_id_cte.c.dag_id, - isouter=True, - ) - ) - - dags_query = apply_filters_to_select( - dags_query, + dags_select, total_entries = paginated_select( + dags_select_with_latest_dag_run, [only_active, paused, dag_id_pattern, dag_display_name_pattern, tags, owners, last_dag_run_state], + order_by, + offset, + limit, + session, ) - # TODO: Re-enable when permissions are handled. - # readable_dags = get_auth_manager().get_permitted_dag_ids(user=g.user) - # dags_query = dags_query.where(DagModel.dag_id.in_(readable_dags)) - - total_entries = get_query_count(dags_query, session=session) - - dags_query = apply_filters_to_select(dags_query, [order_by, offset, limit]) - - dags = session.scalars(dags_query).all() + dags = session.scalars(dags_select).all() return DAGCollectionResponse( dags=[DAGResponse.model_validate(dag, from_attributes=True) for dag in dags], @@ -109,7 +93,7 @@ async def patch_dag( session: Annotated[Session, Depends(get_session)], update_mask: list[str] | None = Query(None), ) -> DAGResponse: - """Update the specific DAG.""" + """Patch the specific DAG.""" dag = session.get(DagModel, dag_id) if dag is None: @@ -127,3 +111,50 @@ async def patch_dag( setattr(dag, attr_name, attr_value) return DAGResponse.model_validate(dag, from_attributes=True) + + +@dags_router.patch("/dags", responses=create_openapi_http_exception_doc([400, 401, 403, 404])) +async def patch_dags( + patch_body: DAGPatchBody, + limit: QueryLimit, + offset: QueryOffset, + tags: QueryTagsFilter, + owners: QueryOwnersFilter, + dag_id_pattern: QueryDagIdPatternSearchWithNone, + only_active: QueryOnlyActiveFilter, + paused: QueryPausedFilter, + last_dag_run_state: QueryLastDagRunStateFilter, + session: Annotated[Session, Depends(get_session)], + update_mask: list[str] | None = Query(None), +) -> DAGCollectionResponse: + """Patch multiple DAGs.""" + if update_mask: + if update_mask != ["is_paused"]: + raise HTTPException(400, "Only `is_paused` field can be updated through the REST API") + else: + update_mask = ["is_paused"] + + dags_select, total_entries = paginated_select( + dags_select_with_latest_dag_run, + [only_active, paused, dag_id_pattern, tags, owners, last_dag_run_state], + None, + offset, + limit, + session, + ) + + dags = session.scalars(dags_select).all() + + dags_to_update = {dag.dag_id for dag in dags} + + session.execute( + update(DagModel) + .where(DagModel.dag_id.in_(dags_to_update)) + .values(is_paused=patch_body.is_paused) + .execution_options(synchronize_session="fetch") + ) + + return DAGCollectionResponse( + dags=[DAGResponse.model_validate(dag, from_attributes=True) for dag in dags], + total_entries=total_entries, + ) diff --git a/airflow/api_fastapi/views/ui/assets.py b/airflow/api_fastapi/views/ui/assets.py index 458d531facf6a..739c7d64af439 100644 --- a/airflow/api_fastapi/views/ui/assets.py +++ b/airflow/api_fastapi/views/ui/assets.py @@ -22,7 +22,7 @@ from sqlalchemy.orm import Session from typing_extensions import Annotated -from airflow.api_fastapi.db import get_session +from airflow.api_fastapi.db.common import get_session from airflow.models import DagModel from airflow.models.asset import AssetDagRunQueue, AssetEvent, AssetModel, DagScheduleAssetReference diff --git a/airflow/ui/openapi-gen/queries/common.ts b/airflow/ui/openapi-gen/queries/common.ts index 46694939ed74e..b1508c86c0c4b 100644 --- a/airflow/ui/openapi-gen/queries/common.ts +++ b/airflow/ui/openapi-gen/queries/common.ts @@ -76,6 +76,9 @@ export const UseDagServiceGetDagsPublicDagsGetKeyFn = ( }, ]), ]; +export type DagServicePatchDagsPublicDagsPatchMutationResult = Awaited< + ReturnType +>; export type DagServicePatchDagPublicDagsDagIdPatchMutationResult = Awaited< ReturnType >; diff --git a/airflow/ui/openapi-gen/queries/queries.ts b/airflow/ui/openapi-gen/queries/queries.ts index 7cbaac5b2c77d..5eda2a3d0e4d2 100644 --- a/airflow/ui/openapi-gen/queries/queries.ts +++ b/airflow/ui/openapi-gen/queries/queries.ts @@ -118,9 +118,95 @@ export const useDagServiceGetDagsPublicDagsGet = < }) as TData, ...options, }); +/** + * Patch Dags + * Patch multiple DAGs. + * @param data The data for the request. + * @param data.requestBody + * @param data.updateMask + * @param data.limit + * @param data.offset + * @param data.tags + * @param data.owners + * @param data.dagIdPattern + * @param data.onlyActive + * @param data.paused + * @param data.lastDagRunState + * @returns DAGCollectionResponse Successful Response + * @throws ApiError + */ +export const useDagServicePatchDagsPublicDagsPatch = < + TData = Common.DagServicePatchDagsPublicDagsPatchMutationResult, + TError = unknown, + TContext = unknown, +>( + options?: Omit< + UseMutationOptions< + TData, + TError, + { + dagIdPattern?: string; + lastDagRunState?: DagRunState; + limit?: number; + offset?: number; + onlyActive?: boolean; + owners?: string[]; + paused?: boolean; + requestBody: DAGPatchBody; + tags?: string[]; + updateMask?: string[]; + }, + TContext + >, + "mutationFn" + >, +) => + useMutation< + TData, + TError, + { + dagIdPattern?: string; + lastDagRunState?: DagRunState; + limit?: number; + offset?: number; + onlyActive?: boolean; + owners?: string[]; + paused?: boolean; + requestBody: DAGPatchBody; + tags?: string[]; + updateMask?: string[]; + }, + TContext + >({ + mutationFn: ({ + dagIdPattern, + lastDagRunState, + limit, + offset, + onlyActive, + owners, + paused, + requestBody, + tags, + updateMask, + }) => + DagService.patchDagsPublicDagsPatch({ + dagIdPattern, + lastDagRunState, + limit, + offset, + onlyActive, + owners, + paused, + requestBody, + tags, + updateMask, + }) as unknown as Promise, + ...options, + }); /** * Patch Dag - * Update the specific DAG. + * Patch the specific DAG. * @param data The data for the request. * @param data.dagId * @param data.requestBody diff --git a/airflow/ui/openapi-gen/requests/services.gen.ts b/airflow/ui/openapi-gen/requests/services.gen.ts index 5aa5876d112ad..7fb6306afbc67 100644 --- a/airflow/ui/openapi-gen/requests/services.gen.ts +++ b/airflow/ui/openapi-gen/requests/services.gen.ts @@ -7,6 +7,8 @@ import type { NextRunAssetsUiNextRunDatasetsDagIdGetResponse, GetDagsPublicDagsGetData, GetDagsPublicDagsGetResponse, + PatchDagsPublicDagsPatchData, + PatchDagsPublicDagsPatchResponse, PatchDagPublicDagsDagIdPatchData, PatchDagPublicDagsDagIdPatchResponse, } from "./types.gen"; @@ -77,9 +79,55 @@ export class DagService { }); } + /** + * Patch Dags + * Patch multiple DAGs. + * @param data The data for the request. + * @param data.requestBody + * @param data.updateMask + * @param data.limit + * @param data.offset + * @param data.tags + * @param data.owners + * @param data.dagIdPattern + * @param data.onlyActive + * @param data.paused + * @param data.lastDagRunState + * @returns DAGCollectionResponse Successful Response + * @throws ApiError + */ + public static patchDagsPublicDagsPatch( + data: PatchDagsPublicDagsPatchData, + ): CancelablePromise { + return __request(OpenAPI, { + method: "PATCH", + url: "/public/dags", + query: { + update_mask: data.updateMask, + limit: data.limit, + offset: data.offset, + tags: data.tags, + owners: data.owners, + dag_id_pattern: data.dagIdPattern, + only_active: data.onlyActive, + paused: data.paused, + last_dag_run_state: data.lastDagRunState, + }, + body: data.requestBody, + mediaType: "application/json", + errors: { + 400: "Bad Request", + 401: "Unauthorized", + 403: "Forbidden", + 404: "Not Found", + 422: "Validation Error", + }, + }); + } + /** * Patch Dag - * Update the specific DAG. + * Patch the specific DAG. * @param data The data for the request. * @param data.dagId * @param data.requestBody diff --git a/airflow/ui/openapi-gen/requests/types.gen.ts b/airflow/ui/openapi-gen/requests/types.gen.ts index bc455f63b6449..0fe7134ba8c31 100644 --- a/airflow/ui/openapi-gen/requests/types.gen.ts +++ b/airflow/ui/openapi-gen/requests/types.gen.ts @@ -111,6 +111,21 @@ export type GetDagsPublicDagsGetData = { export type GetDagsPublicDagsGetResponse = DAGCollectionResponse; +export type PatchDagsPublicDagsPatchData = { + dagIdPattern?: string | null; + lastDagRunState?: DagRunState | null; + limit?: number; + offset?: number; + onlyActive?: boolean; + owners?: Array; + paused?: boolean | null; + requestBody: DAGPatchBody; + tags?: Array; + updateMask?: Array | null; +}; + +export type PatchDagsPublicDagsPatchResponse = DAGCollectionResponse; + export type PatchDagPublicDagsDagIdPatchData = { dagId: string; requestBody: DAGPatchBody; @@ -151,6 +166,35 @@ export type $OpenApiTs = { 422: HTTPValidationError; }; }; + patch: { + req: PatchDagsPublicDagsPatchData; + res: { + /** + * Successful Response + */ + 200: DAGCollectionResponse; + /** + * Bad Request + */ + 400: HTTPExceptionResponse; + /** + * Unauthorized + */ + 401: HTTPExceptionResponse; + /** + * Forbidden + */ + 403: HTTPExceptionResponse; + /** + * Not Found + */ + 404: HTTPExceptionResponse; + /** + * Validation Error + */ + 422: HTTPValidationError; + }; + }; }; "/public/dags/{dag_id}": { patch: { diff --git a/tests/api_fastapi/views/public/test_dags.py b/tests/api_fastapi/views/public/test_dags.py index 6e400f11cc0d2..7b68ebe512a25 100644 --- a/tests/api_fastapi/views/public/test_dags.py +++ b/tests/api_fastapi/views/public/test_dags.py @@ -112,37 +112,37 @@ def setup(dag_maker) -> None: "query_params, expected_total_entries, expected_ids", [ # Filters - ({}, 2, ["test_dag1", "test_dag2"]), - ({"limit": 1}, 2, ["test_dag1"]), - ({"offset": 1}, 2, ["test_dag2"]), - ({"tags": ["example"]}, 1, ["test_dag1"]), - ({"only_active": False}, 3, ["test_dag1", "test_dag2", "test_dag3"]), - ({"paused": True, "only_active": False}, 1, ["test_dag3"]), - ({"paused": False}, 2, ["test_dag1", "test_dag2"]), - ({"owners": ["airflow"]}, 2, ["test_dag1", "test_dag2"]), - ({"owners": ["test_owner"], "only_active": False}, 1, ["test_dag3"]), - ({"last_dag_run_state": "success", "only_active": False}, 1, ["test_dag3"]), - ({"last_dag_run_state": "failed", "only_active": False}, 1, ["test_dag1"]), + ({}, 2, [DAG1_ID, DAG2_ID]), + ({"limit": 1}, 2, [DAG1_ID]), + ({"offset": 1}, 2, [DAG2_ID]), + ({"tags": ["example"]}, 1, [DAG1_ID]), + ({"only_active": False}, 3, [DAG1_ID, DAG2_ID, DAG3_ID]), + ({"paused": True, "only_active": False}, 1, [DAG3_ID]), + ({"paused": False}, 2, [DAG1_ID, DAG2_ID]), + ({"owners": ["airflow"]}, 2, [DAG1_ID, DAG2_ID]), + ({"owners": ["test_owner"], "only_active": False}, 1, [DAG3_ID]), + ({"last_dag_run_state": "success", "only_active": False}, 1, [DAG3_ID]), + ({"last_dag_run_state": "failed", "only_active": False}, 1, [DAG1_ID]), # # Sort - ({"order_by": "-dag_id"}, 2, ["test_dag2", "test_dag1"]), - ({"order_by": "-dag_display_name"}, 2, ["test_dag2", "test_dag1"]), - ({"order_by": "dag_display_name"}, 2, ["test_dag1", "test_dag2"]), - ({"order_by": "next_dagrun", "only_active": False}, 3, ["test_dag3", "test_dag1", "test_dag2"]), - ({"order_by": "last_run_state", "only_active": False}, 3, ["test_dag1", "test_dag3", "test_dag2"]), - ({"order_by": "-last_run_state", "only_active": False}, 3, ["test_dag3", "test_dag1", "test_dag2"]), + ({"order_by": "-dag_id"}, 2, [DAG2_ID, DAG1_ID]), + ({"order_by": "-dag_display_name"}, 2, [DAG2_ID, DAG1_ID]), + ({"order_by": "dag_display_name"}, 2, [DAG1_ID, DAG2_ID]), + ({"order_by": "next_dagrun", "only_active": False}, 3, [DAG3_ID, DAG1_ID, DAG2_ID]), + ({"order_by": "last_run_state", "only_active": False}, 3, [DAG1_ID, DAG3_ID, DAG2_ID]), + ({"order_by": "-last_run_state", "only_active": False}, 3, [DAG3_ID, DAG1_ID, DAG2_ID]), ( {"order_by": "last_run_start_date", "only_active": False}, 3, - ["test_dag1", "test_dag3", "test_dag2"], + [DAG1_ID, DAG3_ID, DAG2_ID], ), ( {"order_by": "-last_run_start_date", "only_active": False}, 3, - ["test_dag3", "test_dag1", "test_dag2"], + [DAG3_ID, DAG1_ID, DAG2_ID], ), # Search - ({"dag_id_pattern": "1"}, 1, ["test_dag1"]), - ({"dag_display_name_pattern": "display2"}, 1, ["test_dag2"]), + ({"dag_id_pattern": "1"}, 1, [DAG1_ID]), + ({"dag_display_name_pattern": "display2"}, 1, [DAG2_ID]), ], ) def test_get_dags(test_client, query_params, expected_total_entries, expected_ids): @@ -173,3 +173,55 @@ def test_patch_dag(test_client, query_params, dag_id, body, expected_status_code if expected_status_code == 200: body = response.json() assert body["is_paused"] == expected_is_paused + + +@pytest.mark.parametrize( + "query_params, body, expected_status_code, expected_ids, expected_paused_ids", + [ + ({"update_mask": ["field_1", "is_paused"]}, {"is_paused": True}, 400, None, None), + ( + {"only_active": False}, + {"is_paused": True}, + 200, + [], + [], + ), # no-op because the dag_id_pattern is not provided + ( + {"only_active": False, "dag_id_pattern": "~"}, + {"is_paused": True}, + 200, + [DAG1_ID, DAG2_ID, DAG3_ID], + [DAG1_ID, DAG2_ID, DAG3_ID], + ), + ( + {"only_active": False, "dag_id_pattern": "~"}, + {"is_paused": False}, + 200, + [DAG1_ID, DAG2_ID, DAG3_ID], + [], + ), + ( + {"dag_id_pattern": "~"}, + {"is_paused": True}, + 200, + [DAG1_ID, DAG2_ID], + [DAG1_ID, DAG2_ID], + ), + ( + {"dag_id_pattern": "dag1"}, + {"is_paused": True}, + 200, + [DAG1_ID], + [DAG1_ID], + ), + ], +) +def test_patch_dags(test_client, query_params, body, expected_status_code, expected_ids, expected_paused_ids): + response = test_client.patch("/public/dags", json=body, params=query_params) + + assert response.status_code == expected_status_code + if expected_status_code == 200: + body = response.json() + assert [dag["dag_id"] for dag in body["dags"]] == expected_ids + paused_dag_ids = [dag["dag_id"] for dag in body["dags"] if dag["is_paused"]] + assert paused_dag_ids == expected_paused_ids From 68bb306af8ac8a2d890892f3ae179162146d2b38 Mon Sep 17 00:00:00 2001 From: Daniel Standish <15932138+dstandish@users.noreply.github.com> Date: Mon, 30 Sep 2024 08:58:36 -0700 Subject: [PATCH 076/802] KubernetesHook kube_config extra can take dict (#41413) Previously had to be json-encoded string which is less convenient when defining the conn in json. --------- Co-authored-by: Jed Cunningham <66968678+jedcunningham@users.noreply.github.com> --- airflow/providers/cncf/kubernetes/hooks/kubernetes.py | 2 ++ tests/providers/cncf/kubernetes/hooks/test_kubernetes.py | 2 ++ 2 files changed, 4 insertions(+) diff --git a/airflow/providers/cncf/kubernetes/hooks/kubernetes.py b/airflow/providers/cncf/kubernetes/hooks/kubernetes.py index 9f7e33696eb87..a810e8f9ed522 100644 --- a/airflow/providers/cncf/kubernetes/hooks/kubernetes.py +++ b/airflow/providers/cncf/kubernetes/hooks/kubernetes.py @@ -250,6 +250,8 @@ def get_conn(self) -> client.ApiClient: if kubeconfig is not None: with tempfile.NamedTemporaryFile() as temp_config: self.log.debug("loading kube_config from: connection kube_config") + if isinstance(kubeconfig, dict): + kubeconfig = json.dumps(kubeconfig) temp_config.write(kubeconfig.encode()) temp_config.flush() self._is_in_cluster = False diff --git a/tests/providers/cncf/kubernetes/hooks/test_kubernetes.py b/tests/providers/cncf/kubernetes/hooks/test_kubernetes.py index 348974eacdfa6..065768def24ea 100644 --- a/tests/providers/cncf/kubernetes/hooks/test_kubernetes.py +++ b/tests/providers/cncf/kubernetes/hooks/test_kubernetes.py @@ -79,6 +79,7 @@ def setup_class(cls) -> None: ("in_cluster", {"in_cluster": True}), ("in_cluster_empty", {"in_cluster": ""}), ("kube_config", {"kube_config": '{"test": "kube"}'}), + ("kube_config_dict", {"kube_config": {"test": "kube"}}), ("kube_config_path", {"kube_config_path": "path/to/file"}), ("kube_config_empty", {"kube_config": ""}), ("kube_config_path_empty", {"kube_config_path": ""}), @@ -285,6 +286,7 @@ def test_kube_config_path( ( (None, False), ("kube_config", True), + ("kube_config_dict", True), ("kube_config_empty", False), ), ) From e606ae4150305d20684a10b735854b7f037079e7 Mon Sep 17 00:00:00 2001 From: Daniel Standish <15932138+dstandish@users.noreply.github.com> Date: Mon, 30 Sep 2024 12:58:31 -0700 Subject: [PATCH 077/802] Speed up boring cyborg consistency pre-commit check (#42589) This is typically the slowest pre-commit besides mypy, and it runs every time. Previously it loaded all filenames into memory and ran glob filter on that. It seems faster to apply glob against the file system directly. This makes pre-commit much faster. Previously took around 4 seconds, now about a half a second. --- scripts/ci/pre_commit/boring_cyborg.py | 19 +++++++++---------- 1 file changed, 9 insertions(+), 10 deletions(-) diff --git a/scripts/ci/pre_commit/boring_cyborg.py b/scripts/ci/pre_commit/boring_cyborg.py index cf852b12bb6da..ec674485b5457 100755 --- a/scripts/ci/pre_commit/boring_cyborg.py +++ b/scripts/ci/pre_commit/boring_cyborg.py @@ -17,13 +17,11 @@ # under the License. from __future__ import annotations -import subprocess import sys from pathlib import Path import yaml from termcolor import colored -from wcmatch import glob if __name__ not in ("__main__", "__mp_main__"): raise SystemExit( @@ -33,9 +31,8 @@ CONFIG_KEY = "labelPRBasedOnFilePath" -current_files = subprocess.check_output(["git", "ls-files"]).decode().splitlines() -git_root = Path(subprocess.check_output(["git", "rev-parse", "--show-toplevel"]).decode().strip()) -cyborg_config_path = git_root / ".github" / "boring-cyborg.yml" +repo_root = Path(__file__).parent.parent.parent.parent +cyborg_config_path = repo_root / ".github" / "boring-cyborg.yml" cyborg_config = yaml.safe_load(cyborg_config_path.read_text()) if CONFIG_KEY not in cyborg_config: raise SystemExit(f"Missing section {CONFIG_KEY}") @@ -43,12 +40,14 @@ errors = [] for label, patterns in cyborg_config[CONFIG_KEY].items(): for pattern in patterns: - if glob.globfilter(current_files, pattern, flags=glob.G | glob.E): + try: + next(Path(repo_root).glob(pattern)) continue - yaml_path = f"{CONFIG_KEY}.{label}" - errors.append( - f"Unused pattern [{colored(pattern, 'cyan')}] in [{colored(yaml_path, 'cyan')}] section." - ) + except StopIteration: + yaml_path = f"{CONFIG_KEY}.{label}" + errors.append( + f"Unused pattern [{colored(pattern, 'cyan')}] in [{colored(yaml_path, 'cyan')}] section." + ) if errors: print(f"Found {colored(str(len(errors)), 'red')} problems:") From 174ea479adf40e55dad82edf6815782fda91e073 Mon Sep 17 00:00:00 2001 From: JISHAN GARGACHARYA <34843832+jishangarg@users.noreply.github.com> Date: Tue, 1 Oct 2024 12:02:59 +0530 Subject: [PATCH 078/802] Doc update - Airflow local settings no longer importable from dags folder (#42231) --------- Co-authored-by: Jishan Garg --- docs/apache-airflow/howto/set-config.rst | 2 ++ 1 file changed, 2 insertions(+) diff --git a/docs/apache-airflow/howto/set-config.rst b/docs/apache-airflow/howto/set-config.rst index 4f19159a810d4..2a03b2bbf5ee4 100644 --- a/docs/apache-airflow/howto/set-config.rst +++ b/docs/apache-airflow/howto/set-config.rst @@ -179,6 +179,8 @@ where you can configure such local settings - This is usually done in the ``airf You should create a ``airflow_local_settings.py`` file and put it in a directory in ``sys.path`` or in the ``$AIRFLOW_HOME/config`` folder. (Airflow adds ``$AIRFLOW_HOME/config`` to ``sys.path`` when Airflow is initialized) +Starting from Airflow 2.10.1, the $AIRFLOW_HOME/dags folder is no longer included in sys.path at initialization, so any local settings in that folder will not be imported. Ensure that airflow_local_settings.py is located in a path that is part of sys.path during initialization, like $AIRFLOW_HOME/config. +For more context about this change, see the `mailing list announcement `_. You can see the example of such local settings here: From 10f2ea9ce5887ef2ab5f769bd1300b6c75344f76 Mon Sep 17 00:00:00 2001 From: Dewen Kong Date: Tue, 1 Oct 2024 02:40:33 -0400 Subject: [PATCH 079/802] add flexibility for redis service (#41811) * add service type options for redis * additional value * update based on testing * fix syntax * update description * Update chart/templates/redis/redis-service.yaml Co-authored-by: rom sharon <33751805+romsharon98@users.noreply.github.com> * Update chart/templates/redis/redis-service.yaml Co-authored-by: rom sharon <33751805+romsharon98@users.noreply.github.com> --------- Co-authored-by: rom sharon <33751805+romsharon98@users.noreply.github.com> --- chart/templates/redis/redis-service.yaml | 10 ++++++ chart/values.schema.json | 33 +++++++++++++++++++ chart/values.yaml | 8 +++++ helm_tests/other/test_redis.py | 41 ++++++++++++++++++++++++ 4 files changed, 92 insertions(+) diff --git a/chart/templates/redis/redis-service.yaml b/chart/templates/redis/redis-service.yaml index 17d4c8d5e4836..ee010901ef84e 100644 --- a/chart/templates/redis/redis-service.yaml +++ b/chart/templates/redis/redis-service.yaml @@ -35,7 +35,14 @@ metadata: {{- toYaml . | nindent 4 }} {{- end }} spec: +{{- if eq .Values.redis.service.type "ClusterIP" }} type: ClusterIP + {{- if .Values.redis.service.clusterIP }} + clusterIP: {{ .Values.redis.service.clusterIP }} + {{- end }} +{{- else }} + type: {{ .Values.redis.service.type }} +{{- end }} selector: tier: airflow component: redis @@ -45,4 +52,7 @@ spec: protocol: TCP port: {{ .Values.ports.redisDB }} targetPort: {{ .Values.ports.redisDB }} + {{- if (and (eq .Values.redis.service.type "NodePort") (not (empty .Values.redis.service.nodePort))) }} + nodePort: {{ .Values.redis.service.nodePort }} + {{- end }} {{- end }} diff --git a/chart/values.schema.json b/chart/values.schema.json index 948f09f3b9a4d..d8b5de41c8eb8 100644 --- a/chart/values.schema.json +++ b/chart/values.schema.json @@ -7670,6 +7670,39 @@ "type": "integer", "default": 600 }, + "service": { + "description": "service configuration.", + "type": "object", + "additionalProperties": false, + "properties": { + "type": { + "description": "Service type.", + "enum": [ + "ClusterIP", + "NodePort", + "LoadBalancer" + ], + "type": "string", + "default": "ClusterIP" + }, + "clusterIP": { + "description": "If using `ClusterIP` service type, custom IP address can be specified.", + "type": [ + "string", + "null" + ], + "default": null + }, + "nodePort": { + "description": "If using `NodePort` service type, custom node port can be specified.", + "type": [ + "integer", + "null" + ], + "default": null + } + } + }, "persistence": { "description": "Persistence configuration.", "type": "object", diff --git a/chart/values.yaml b/chart/values.yaml index 7bfa733a905b4..0edb9f2bd7cd3 100644 --- a/chart/values.yaml +++ b/chart/values.yaml @@ -2378,6 +2378,14 @@ redis: # Annotations to add to worker kubernetes service account. annotations: {} + service: + # service type, default: ClusterIP + type: "ClusterIP" + # If using ClusterIP service type, custom IP address can be specified + clusterIP: + # If using NodePort service type, custom node port can be specified + nodePort: + persistence: # Enable persistent volumes enabled: true diff --git a/helm_tests/other/test_redis.py b/helm_tests/other/test_redis.py index a5a6f2099e4ab..8c44567420314 100644 --- a/helm_tests/other/test_redis.py +++ b/helm_tests/other/test_redis.py @@ -452,3 +452,44 @@ def test_overridden_automount_service_account_token(self): show_only=["templates/redis/redis-serviceaccount.yaml"], ) assert jmespath.search("automountServiceAccountToken", docs[0]) is False + + +class TestRedisService: + """Tests redis service.""" + + @pytest.mark.parametrize( + "redis_values, expected", + [ + ({"redis": {"service": {"type": "ClusterIP"}}}, "ClusterIP"), + ({"redis": {"service": {"type": "NodePort"}}}, "NodePort"), + ({"redis": {"service": {"type": "LoadBalancer"}}}, "LoadBalancer"), + ], + ) + def test_redis_service_type(self, redis_values, expected): + docs = render_chart( + values=redis_values, + show_only=["templates/redis/redis-service.yaml"], + ) + assert expected == jmespath.search("spec.type", docs[0]) + + def test_redis_service_nodeport(self): + docs = render_chart( + values={ + "redis": { + "service": {"type": "NodePort", "nodePort": 11111}, + }, + }, + show_only=["templates/redis/redis-service.yaml"], + ) + assert 11111 == jmespath.search("spec.ports[0].nodePort", docs[0]) + + def test_redis_service_clusterIP(self): + docs = render_chart( + values={ + "redis": { + "service": {"type": "ClusterIP", "clusterIP": "127.0.0.1"}, + }, + }, + show_only=["templates/redis/redis-service.yaml"], + ) + assert "127.0.0.1" == jmespath.search("spec.clusterIP", docs[0]) From 491191cce9b3fb259e58b58acb702d7eeb6eddb4 Mon Sep 17 00:00:00 2001 From: Howard Yoo <32691630+howardyoo@users.noreply.github.com> Date: Tue, 1 Oct 2024 02:04:06 -0500 Subject: [PATCH 080/802] Support of host.name in OTEL metrics and usage of OTEL_RESOURCE_ATTRIBUTES in metrics (#42428) * fixes: 42425, and 42424 * fixed static type check failure --- airflow/metrics/otel_logger.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/airflow/metrics/otel_logger.py b/airflow/metrics/otel_logger.py index 14080eb2d8313..6d7d6e8fffa1c 100644 --- a/airflow/metrics/otel_logger.py +++ b/airflow/metrics/otel_logger.py @@ -28,7 +28,7 @@ from opentelemetry.metrics import Observation from opentelemetry.sdk.metrics import MeterProvider from opentelemetry.sdk.metrics._internal.export import ConsoleMetricExporter, PeriodicExportingMetricReader -from opentelemetry.sdk.resources import SERVICE_NAME, Resource +from opentelemetry.sdk.resources import HOST_NAME, SERVICE_NAME, Resource from airflow.configuration import conf from airflow.exceptions import AirflowProviderDeprecationWarning @@ -40,6 +40,7 @@ get_validator, stat_name_otel_handler, ) +from airflow.utils.net import get_hostname if TYPE_CHECKING: from opentelemetry.metrics import Instrument @@ -410,7 +411,7 @@ def get_otel_logger(cls) -> SafeOtelLogger: debug = conf.getboolean("metrics", "otel_debugging_on") service_name = conf.get("metrics", "otel_service") - resource = Resource(attributes={SERVICE_NAME: service_name}) + resource = Resource.create(attributes={HOST_NAME: get_hostname(), SERVICE_NAME: service_name}) protocol = "https" if ssl_active else "http" endpoint = f"{protocol}://{host}:{port}/v1/metrics" From e42c157422c199aa2c0d6af8ca978eaee4a3e50a Mon Sep 17 00:00:00 2001 From: Pierre Jeambrun Date: Tue, 1 Oct 2024 15:42:37 +0800 Subject: [PATCH 081/802] Update fastapi operation ids (#42588) * Update operation id automatically * Cherry pick Brent change --------- Co-authored-by: Brent Bovenzi --- airflow/api_fastapi/openapi/v1-generated.yaml | 10 +- airflow/api_fastapi/views/public/__init__.py | 5 +- airflow/api_fastapi/views/public/dags.py | 5 +- airflow/api_fastapi/views/router.py | 93 +++++++++++++++++++ airflow/api_fastapi/views/ui/__init__.py | 5 +- airflow/api_fastapi/views/ui/assets.py | 5 +- airflow/ui/openapi-gen/queries/common.ts | 44 ++++----- airflow/ui/openapi-gen/queries/prefetch.ts | 15 ++- airflow/ui/openapi-gen/queries/queries.ts | 32 +++---- airflow/ui/openapi-gen/queries/suspense.ts | 20 ++-- .../ui/openapi-gen/requests/services.gen.ts | 40 ++++---- airflow/ui/openapi-gen/requests/types.gen.ts | 24 ++--- airflow/ui/package.json | 2 +- airflow/ui/src/App.test.tsx | 7 +- airflow/ui/src/pages/DagsList.tsx | 4 +- 15 files changed, 193 insertions(+), 118 deletions(-) create mode 100644 airflow/api_fastapi/views/router.py diff --git a/airflow/api_fastapi/openapi/v1-generated.yaml b/airflow/api_fastapi/openapi/v1-generated.yaml index a38a1021890d6..b08ef42c16df1 100644 --- a/airflow/api_fastapi/openapi/v1-generated.yaml +++ b/airflow/api_fastapi/openapi/v1-generated.yaml @@ -12,7 +12,7 @@ paths: tags: - Asset summary: Next Run Assets - operationId: next_run_assets_ui_next_run_datasets__dag_id__get + operationId: next_run_assets parameters: - name: dag_id in: path @@ -27,7 +27,7 @@ paths: application/json: schema: type: object - title: Response Next Run Assets Ui Next Run Datasets Dag Id Get + title: Response Next Run Assets '422': description: Validation Error content: @@ -40,7 +40,7 @@ paths: - DAG summary: Get Dags description: Get all DAGs. - operationId: get_dags_public_dags_get + operationId: get_dags parameters: - name: limit in: query @@ -136,7 +136,7 @@ paths: - DAG summary: Patch Dags description: Patch multiple DAGs. - operationId: patch_dags_public_dags_patch + operationId: patch_dags parameters: - name: update_mask in: query @@ -258,7 +258,7 @@ paths: - DAG summary: Patch Dag description: Patch the specific DAG. - operationId: patch_dag_public_dags__dag_id__patch + operationId: patch_dag parameters: - name: dag_id in: path diff --git a/airflow/api_fastapi/views/public/__init__.py b/airflow/api_fastapi/views/public/__init__.py index b6466536c3359..1c2511fc82ac2 100644 --- a/airflow/api_fastapi/views/public/__init__.py +++ b/airflow/api_fastapi/views/public/__init__.py @@ -17,11 +17,10 @@ from __future__ import annotations -from fastapi import APIRouter - from airflow.api_fastapi.views.public.dags import dags_router +from airflow.api_fastapi.views.router import AirflowRouter -public_router = APIRouter(prefix="/public") +public_router = AirflowRouter(prefix="/public") public_router.include_router(dags_router) diff --git a/airflow/api_fastapi/views/public/dags.py b/airflow/api_fastapi/views/public/dags.py index a6c25d6568c1e..3761d593d2fd0 100644 --- a/airflow/api_fastapi/views/public/dags.py +++ b/airflow/api_fastapi/views/public/dags.py @@ -17,7 +17,7 @@ from __future__ import annotations -from fastapi import APIRouter, Depends, HTTPException, Query +from fastapi import Depends, HTTPException, Query from sqlalchemy import update from sqlalchemy.orm import Session from typing_extensions import Annotated @@ -42,9 +42,10 @@ SortParam, ) from airflow.api_fastapi.serializers.dags import DAGCollectionResponse, DAGPatchBody, DAGResponse +from airflow.api_fastapi.views.router import AirflowRouter from airflow.models import DagModel -dags_router = APIRouter(tags=["DAG"]) +dags_router = AirflowRouter(tags=["DAG"]) @dags_router.get("/dags") diff --git a/airflow/api_fastapi/views/router.py b/airflow/api_fastapi/views/router.py new file mode 100644 index 0000000000000..5bf07e0fe834a --- /dev/null +++ b/airflow/api_fastapi/views/router.py @@ -0,0 +1,93 @@ +# 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 + +from enum import Enum +from typing import Any, Callable, Sequence + +from fastapi import APIRouter, params +from fastapi.datastructures import Default +from fastapi.routing import APIRoute +from fastapi.types import DecoratedCallable, IncEx +from fastapi.utils import generate_unique_id +from starlette.responses import JSONResponse, Response +from starlette.routing import BaseRoute + + +class AirflowRouter(APIRouter): + """Extends the FastAPI default router.""" + + def api_route( + self, + path: str, + *, + response_model: Any = Default(None), + status_code: int | None = None, + tags: list[str | Enum] | None = None, + dependencies: Sequence[params.Depends] | None = None, + summary: str | None = None, + description: str | None = None, + response_description: str = "Successful Response", + responses: dict[int | str, dict[str, Any]] | None = None, + deprecated: bool | None = None, + methods: list[str] | None = None, + operation_id: str | None = None, + response_model_include: IncEx | None = None, + response_model_exclude: IncEx | None = None, + response_model_by_alias: bool = True, + response_model_exclude_unset: bool = False, + response_model_exclude_defaults: bool = False, + response_model_exclude_none: bool = False, + include_in_schema: bool = True, + response_class: type[Response] = Default(JSONResponse), + name: str | None = None, + callbacks: list[BaseRoute] | None = None, + openapi_extra: dict[str, Any] | None = None, + generate_unique_id_function: Callable[[APIRoute], str] = Default(generate_unique_id), + ) -> Callable[[DecoratedCallable], DecoratedCallable]: + def decorator(func: DecoratedCallable) -> DecoratedCallable: + self.add_api_route( + path, + func, + response_model=response_model, + status_code=status_code, + tags=tags, + dependencies=dependencies, + summary=summary, + description=description, + response_description=response_description, + responses=responses, + deprecated=deprecated, + methods=methods, + operation_id=operation_id or func.__name__, + response_model_include=response_model_include, + response_model_exclude=response_model_exclude, + response_model_by_alias=response_model_by_alias, + response_model_exclude_unset=response_model_exclude_unset, + response_model_exclude_defaults=response_model_exclude_defaults, + response_model_exclude_none=response_model_exclude_none, + include_in_schema=include_in_schema, + response_class=response_class, + name=name, + callbacks=callbacks, + openapi_extra=openapi_extra, + generate_unique_id_function=generate_unique_id_function, + ) + return func + + return decorator diff --git a/airflow/api_fastapi/views/ui/__init__.py b/airflow/api_fastapi/views/ui/__init__.py index edba930c3d1d1..8495ac5e5e6a4 100644 --- a/airflow/api_fastapi/views/ui/__init__.py +++ b/airflow/api_fastapi/views/ui/__init__.py @@ -16,10 +16,9 @@ # under the License. from __future__ import annotations -from fastapi import APIRouter - +from airflow.api_fastapi.views.router import AirflowRouter from airflow.api_fastapi.views.ui.assets import assets_router -ui_router = APIRouter(prefix="/ui") +ui_router = AirflowRouter(prefix="/ui") ui_router.include_router(assets_router) diff --git a/airflow/api_fastapi/views/ui/assets.py b/airflow/api_fastapi/views/ui/assets.py index 739c7d64af439..01cc9fd1cfbff 100644 --- a/airflow/api_fastapi/views/ui/assets.py +++ b/airflow/api_fastapi/views/ui/assets.py @@ -17,16 +17,17 @@ from __future__ import annotations -from fastapi import APIRouter, Depends, HTTPException, Request +from fastapi import Depends, HTTPException, Request from sqlalchemy import and_, func, select from sqlalchemy.orm import Session from typing_extensions import Annotated from airflow.api_fastapi.db.common import get_session +from airflow.api_fastapi.views.router import AirflowRouter from airflow.models import DagModel from airflow.models.asset import AssetDagRunQueue, AssetEvent, AssetModel, DagScheduleAssetReference -assets_router = APIRouter(tags=["Asset"]) +assets_router = AirflowRouter(tags=["Asset"]) @assets_router.get("/next_run_datasets/{dag_id}", include_in_schema=False) diff --git a/airflow/ui/openapi-gen/queries/common.ts b/airflow/ui/openapi-gen/queries/common.ts index b1508c86c0c4b..96e49cc6d7673 100644 --- a/airflow/ui/openapi-gen/queries/common.ts +++ b/airflow/ui/openapi-gen/queries/common.ts @@ -4,37 +4,31 @@ import { UseQueryResult } from "@tanstack/react-query"; import { AssetService, DagService } from "../requests/services.gen"; import { DagRunState } from "../requests/types.gen"; -export type AssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetDefaultResponse = - Awaited< - ReturnType - >; -export type AssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetQueryResult< - TData = AssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetDefaultResponse, +export type AssetServiceNextRunAssetsDefaultResponse = Awaited< + ReturnType +>; +export type AssetServiceNextRunAssetsQueryResult< + TData = AssetServiceNextRunAssetsDefaultResponse, TError = unknown, > = UseQueryResult; -export const useAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetKey = - "AssetServiceNextRunAssetsUiNextRunDatasetsDagIdGet"; -export const UseAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetKeyFn = ( +export const useAssetServiceNextRunAssetsKey = "AssetServiceNextRunAssets"; +export const UseAssetServiceNextRunAssetsKeyFn = ( { dagId, }: { dagId: string; }, queryKey?: Array, -) => [ - useAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetKey, - ...(queryKey ?? [{ dagId }]), -]; -export type DagServiceGetDagsPublicDagsGetDefaultResponse = Awaited< - ReturnType +) => [useAssetServiceNextRunAssetsKey, ...(queryKey ?? [{ dagId }])]; +export type DagServiceGetDagsDefaultResponse = Awaited< + ReturnType >; -export type DagServiceGetDagsPublicDagsGetQueryResult< - TData = DagServiceGetDagsPublicDagsGetDefaultResponse, +export type DagServiceGetDagsQueryResult< + TData = DagServiceGetDagsDefaultResponse, TError = unknown, > = UseQueryResult; -export const useDagServiceGetDagsPublicDagsGetKey = - "DagServiceGetDagsPublicDagsGet"; -export const UseDagServiceGetDagsPublicDagsGetKeyFn = ( +export const useDagServiceGetDagsKey = "DagServiceGetDags"; +export const UseDagServiceGetDagsKeyFn = ( { dagDisplayNamePattern, dagIdPattern, @@ -60,7 +54,7 @@ export const UseDagServiceGetDagsPublicDagsGetKeyFn = ( } = {}, queryKey?: Array, ) => [ - useDagServiceGetDagsPublicDagsGetKey, + useDagServiceGetDagsKey, ...(queryKey ?? [ { dagDisplayNamePattern, @@ -76,9 +70,9 @@ export const UseDagServiceGetDagsPublicDagsGetKeyFn = ( }, ]), ]; -export type DagServicePatchDagsPublicDagsPatchMutationResult = Awaited< - ReturnType +export type DagServicePatchDagsMutationResult = Awaited< + ReturnType >; -export type DagServicePatchDagPublicDagsDagIdPatchMutationResult = Awaited< - ReturnType +export type DagServicePatchDagMutationResult = Awaited< + ReturnType >; diff --git a/airflow/ui/openapi-gen/queries/prefetch.ts b/airflow/ui/openapi-gen/queries/prefetch.ts index 7de7282a9bd01..95c2c7b737348 100644 --- a/airflow/ui/openapi-gen/queries/prefetch.ts +++ b/airflow/ui/openapi-gen/queries/prefetch.ts @@ -12,7 +12,7 @@ import * as Common from "./common"; * @returns unknown Successful Response * @throws ApiError */ -export const prefetchUseAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGet = ( +export const prefetchUseAssetServiceNextRunAssets = ( queryClient: QueryClient, { dagId, @@ -21,11 +21,8 @@ export const prefetchUseAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGet = ( }, ) => queryClient.prefetchQuery({ - queryKey: Common.UseAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetKeyFn( - { dagId }, - ), - queryFn: () => - AssetService.nextRunAssetsUiNextRunDatasetsDagIdGet({ dagId }), + queryKey: Common.UseAssetServiceNextRunAssetsKeyFn({ dagId }), + queryFn: () => AssetService.nextRunAssets({ dagId }), }); /** * Get Dags @@ -44,7 +41,7 @@ export const prefetchUseAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGet = ( * @returns DAGCollectionResponse Successful Response * @throws ApiError */ -export const prefetchUseDagServiceGetDagsPublicDagsGet = ( +export const prefetchUseDagServiceGetDags = ( queryClient: QueryClient, { dagDisplayNamePattern, @@ -71,7 +68,7 @@ export const prefetchUseDagServiceGetDagsPublicDagsGet = ( } = {}, ) => queryClient.prefetchQuery({ - queryKey: Common.UseDagServiceGetDagsPublicDagsGetKeyFn({ + queryKey: Common.UseDagServiceGetDagsKeyFn({ dagDisplayNamePattern, dagIdPattern, lastDagRunState, @@ -84,7 +81,7 @@ export const prefetchUseDagServiceGetDagsPublicDagsGet = ( tags, }), queryFn: () => - DagService.getDagsPublicDagsGet({ + DagService.getDags({ dagDisplayNamePattern, dagIdPattern, lastDagRunState, diff --git a/airflow/ui/openapi-gen/queries/queries.ts b/airflow/ui/openapi-gen/queries/queries.ts index 5eda2a3d0e4d2..985bf952e3eb3 100644 --- a/airflow/ui/openapi-gen/queries/queries.ts +++ b/airflow/ui/openapi-gen/queries/queries.ts @@ -17,8 +17,8 @@ import * as Common from "./common"; * @returns unknown Successful Response * @throws ApiError */ -export const useAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGet = < - TData = Common.AssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetDefaultResponse, +export const useAssetServiceNextRunAssets = < + TData = Common.AssetServiceNextRunAssetsDefaultResponse, TError = unknown, TQueryKey extends Array = unknown[], >( @@ -31,12 +31,8 @@ export const useAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGet = < options?: Omit, "queryKey" | "queryFn">, ) => useQuery({ - queryKey: Common.UseAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetKeyFn( - { dagId }, - queryKey, - ), - queryFn: () => - AssetService.nextRunAssetsUiNextRunDatasetsDagIdGet({ dagId }) as TData, + queryKey: Common.UseAssetServiceNextRunAssetsKeyFn({ dagId }, queryKey), + queryFn: () => AssetService.nextRunAssets({ dagId }) as TData, ...options, }); /** @@ -56,8 +52,8 @@ export const useAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGet = < * @returns DAGCollectionResponse Successful Response * @throws ApiError */ -export const useDagServiceGetDagsPublicDagsGet = < - TData = Common.DagServiceGetDagsPublicDagsGetDefaultResponse, +export const useDagServiceGetDags = < + TData = Common.DagServiceGetDagsDefaultResponse, TError = unknown, TQueryKey extends Array = unknown[], >( @@ -88,7 +84,7 @@ export const useDagServiceGetDagsPublicDagsGet = < options?: Omit, "queryKey" | "queryFn">, ) => useQuery({ - queryKey: Common.UseDagServiceGetDagsPublicDagsGetKeyFn( + queryKey: Common.UseDagServiceGetDagsKeyFn( { dagDisplayNamePattern, dagIdPattern, @@ -104,7 +100,7 @@ export const useDagServiceGetDagsPublicDagsGet = < queryKey, ), queryFn: () => - DagService.getDagsPublicDagsGet({ + DagService.getDags({ dagDisplayNamePattern, dagIdPattern, lastDagRunState, @@ -135,8 +131,8 @@ export const useDagServiceGetDagsPublicDagsGet = < * @returns DAGCollectionResponse Successful Response * @throws ApiError */ -export const useDagServicePatchDagsPublicDagsPatch = < - TData = Common.DagServicePatchDagsPublicDagsPatchMutationResult, +export const useDagServicePatchDags = < + TData = Common.DagServicePatchDagsMutationResult, TError = unknown, TContext = unknown, >( @@ -190,7 +186,7 @@ export const useDagServicePatchDagsPublicDagsPatch = < tags, updateMask, }) => - DagService.patchDagsPublicDagsPatch({ + DagService.patchDags({ dagIdPattern, lastDagRunState, limit, @@ -214,8 +210,8 @@ export const useDagServicePatchDagsPublicDagsPatch = < * @returns DAGResponse Successful Response * @throws ApiError */ -export const useDagServicePatchDagPublicDagsDagIdPatch = < - TData = Common.DagServicePatchDagPublicDagsDagIdPatchMutationResult, +export const useDagServicePatchDag = < + TData = Common.DagServicePatchDagMutationResult, TError = unknown, TContext = unknown, >( @@ -244,7 +240,7 @@ export const useDagServicePatchDagPublicDagsDagIdPatch = < TContext >({ mutationFn: ({ dagId, requestBody, updateMask }) => - DagService.patchDagPublicDagsDagIdPatch({ + DagService.patchDag({ dagId, requestBody, updateMask, diff --git a/airflow/ui/openapi-gen/queries/suspense.ts b/airflow/ui/openapi-gen/queries/suspense.ts index 18dba7acb4b5b..dc8b99dfb2188 100644 --- a/airflow/ui/openapi-gen/queries/suspense.ts +++ b/airflow/ui/openapi-gen/queries/suspense.ts @@ -12,8 +12,8 @@ import * as Common from "./common"; * @returns unknown Successful Response * @throws ApiError */ -export const useAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetSuspense = < - TData = Common.AssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetDefaultResponse, +export const useAssetServiceNextRunAssetsSuspense = < + TData = Common.AssetServiceNextRunAssetsDefaultResponse, TError = unknown, TQueryKey extends Array = unknown[], >( @@ -26,12 +26,8 @@ export const useAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetSuspense = < options?: Omit, "queryKey" | "queryFn">, ) => useSuspenseQuery({ - queryKey: Common.UseAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetKeyFn( - { dagId }, - queryKey, - ), - queryFn: () => - AssetService.nextRunAssetsUiNextRunDatasetsDagIdGet({ dagId }) as TData, + queryKey: Common.UseAssetServiceNextRunAssetsKeyFn({ dagId }, queryKey), + queryFn: () => AssetService.nextRunAssets({ dagId }) as TData, ...options, }); /** @@ -51,8 +47,8 @@ export const useAssetServiceNextRunAssetsUiNextRunDatasetsDagIdGetSuspense = < * @returns DAGCollectionResponse Successful Response * @throws ApiError */ -export const useDagServiceGetDagsPublicDagsGetSuspense = < - TData = Common.DagServiceGetDagsPublicDagsGetDefaultResponse, +export const useDagServiceGetDagsSuspense = < + TData = Common.DagServiceGetDagsDefaultResponse, TError = unknown, TQueryKey extends Array = unknown[], >( @@ -83,7 +79,7 @@ export const useDagServiceGetDagsPublicDagsGetSuspense = < options?: Omit, "queryKey" | "queryFn">, ) => useSuspenseQuery({ - queryKey: Common.UseDagServiceGetDagsPublicDagsGetKeyFn( + queryKey: Common.UseDagServiceGetDagsKeyFn( { dagDisplayNamePattern, dagIdPattern, @@ -99,7 +95,7 @@ export const useDagServiceGetDagsPublicDagsGetSuspense = < queryKey, ), queryFn: () => - DagService.getDagsPublicDagsGet({ + DagService.getDags({ dagDisplayNamePattern, dagIdPattern, lastDagRunState, diff --git a/airflow/ui/openapi-gen/requests/services.gen.ts b/airflow/ui/openapi-gen/requests/services.gen.ts index 7fb6306afbc67..be216bd534c61 100644 --- a/airflow/ui/openapi-gen/requests/services.gen.ts +++ b/airflow/ui/openapi-gen/requests/services.gen.ts @@ -3,14 +3,14 @@ import type { CancelablePromise } from "./core/CancelablePromise"; import { OpenAPI } from "./core/OpenAPI"; import { request as __request } from "./core/request"; import type { - NextRunAssetsUiNextRunDatasetsDagIdGetData, - NextRunAssetsUiNextRunDatasetsDagIdGetResponse, - GetDagsPublicDagsGetData, - GetDagsPublicDagsGetResponse, - PatchDagsPublicDagsPatchData, - PatchDagsPublicDagsPatchResponse, - PatchDagPublicDagsDagIdPatchData, - PatchDagPublicDagsDagIdPatchResponse, + NextRunAssetsData, + NextRunAssetsResponse, + GetDagsData, + GetDagsResponse, + PatchDagsData, + PatchDagsResponse, + PatchDagData, + PatchDagResponse, } from "./types.gen"; export class AssetService { @@ -21,9 +21,9 @@ export class AssetService { * @returns unknown Successful Response * @throws ApiError */ - public static nextRunAssetsUiNextRunDatasetsDagIdGet( - data: NextRunAssetsUiNextRunDatasetsDagIdGetData, - ): CancelablePromise { + public static nextRunAssets( + data: NextRunAssetsData, + ): CancelablePromise { return __request(OpenAPI, { method: "GET", url: "/ui/next_run_datasets/{dag_id}", @@ -55,9 +55,9 @@ export class DagService { * @returns DAGCollectionResponse Successful Response * @throws ApiError */ - public static getDagsPublicDagsGet( - data: GetDagsPublicDagsGetData = {}, - ): CancelablePromise { + public static getDags( + data: GetDagsData = {}, + ): CancelablePromise { return __request(OpenAPI, { method: "GET", url: "/public/dags", @@ -96,9 +96,9 @@ export class DagService { * @returns DAGCollectionResponse Successful Response * @throws ApiError */ - public static patchDagsPublicDagsPatch( - data: PatchDagsPublicDagsPatchData, - ): CancelablePromise { + public static patchDags( + data: PatchDagsData, + ): CancelablePromise { return __request(OpenAPI, { method: "PATCH", url: "/public/dags", @@ -135,9 +135,9 @@ export class DagService { * @returns DAGResponse Successful Response * @throws ApiError */ - public static patchDagPublicDagsDagIdPatch( - data: PatchDagPublicDagsDagIdPatchData, - ): CancelablePromise { + public static patchDag( + data: PatchDagData, + ): CancelablePromise { return __request(OpenAPI, { method: "PATCH", url: "/public/dags/{dag_id}", diff --git a/airflow/ui/openapi-gen/requests/types.gen.ts b/airflow/ui/openapi-gen/requests/types.gen.ts index 0fe7134ba8c31..e1db8310a1dc1 100644 --- a/airflow/ui/openapi-gen/requests/types.gen.ts +++ b/airflow/ui/openapi-gen/requests/types.gen.ts @@ -88,15 +88,15 @@ export type ValidationError = { type: string; }; -export type NextRunAssetsUiNextRunDatasetsDagIdGetData = { +export type NextRunAssetsData = { dagId: string; }; -export type NextRunAssetsUiNextRunDatasetsDagIdGetResponse = { +export type NextRunAssetsResponse = { [key: string]: unknown; }; -export type GetDagsPublicDagsGetData = { +export type GetDagsData = { dagDisplayNamePattern?: string | null; dagIdPattern?: string | null; lastDagRunState?: DagRunState | null; @@ -109,9 +109,9 @@ export type GetDagsPublicDagsGetData = { tags?: Array; }; -export type GetDagsPublicDagsGetResponse = DAGCollectionResponse; +export type GetDagsResponse = DAGCollectionResponse; -export type PatchDagsPublicDagsPatchData = { +export type PatchDagsData = { dagIdPattern?: string | null; lastDagRunState?: DagRunState | null; limit?: number; @@ -124,20 +124,20 @@ export type PatchDagsPublicDagsPatchData = { updateMask?: Array | null; }; -export type PatchDagsPublicDagsPatchResponse = DAGCollectionResponse; +export type PatchDagsResponse = DAGCollectionResponse; -export type PatchDagPublicDagsDagIdPatchData = { +export type PatchDagData = { dagId: string; requestBody: DAGPatchBody; updateMask?: Array | null; }; -export type PatchDagPublicDagsDagIdPatchResponse = DAGResponse; +export type PatchDagResponse = DAGResponse; export type $OpenApiTs = { "/ui/next_run_datasets/{dag_id}": { get: { - req: NextRunAssetsUiNextRunDatasetsDagIdGetData; + req: NextRunAssetsData; res: { /** * Successful Response @@ -154,7 +154,7 @@ export type $OpenApiTs = { }; "/public/dags": { get: { - req: GetDagsPublicDagsGetData; + req: GetDagsData; res: { /** * Successful Response @@ -167,7 +167,7 @@ export type $OpenApiTs = { }; }; patch: { - req: PatchDagsPublicDagsPatchData; + req: PatchDagsData; res: { /** * Successful Response @@ -198,7 +198,7 @@ export type $OpenApiTs = { }; "/public/dags/{dag_id}": { patch: { - req: PatchDagPublicDagsDagIdPatchData; + req: PatchDagData; res: { /** * Successful Response diff --git a/airflow/ui/package.json b/airflow/ui/package.json index c7d79f792a59e..1f77334074f03 100644 --- a/airflow/ui/package.json +++ b/airflow/ui/package.json @@ -11,7 +11,7 @@ "lint:fix": "eslint --fix && tsc --p tsconfig.app.json", "format": "pnpm prettier --write .", "preview": "vite preview", - "codegen": "openapi-rq -i \"../api_fastapi/openapi/v1-generated.yaml\" -c axios --format prettier -o openapi-gen", + "codegen": "openapi-rq -i \"../api_fastapi/openapi/v1-generated.yaml\" -c axios --format prettier -o openapi-gen --operationId", "test": "vitest run", "coverage": "vitest run --coverage" }, diff --git a/airflow/ui/src/App.test.tsx b/airflow/ui/src/App.test.tsx index d34cf016befdb..5efcf90f1a05d 100644 --- a/airflow/ui/src/App.test.tsx +++ b/airflow/ui/src/App.test.tsx @@ -105,10 +105,9 @@ beforeEach(() => { isLoading: false, } as QueryObserverSuccessResult; - vi.spyOn( - openapiQueriesModule, - "useDagServiceGetDagsPublicDagsGet", - ).mockImplementation(() => returnValue); + vi.spyOn(openapiQueriesModule, "useDagServiceGetDags").mockImplementation( + () => returnValue, + ); }); afterEach(() => { diff --git a/airflow/ui/src/pages/DagsList.tsx b/airflow/ui/src/pages/DagsList.tsx index fe764f117e45d..ab480d2cbabdb 100644 --- a/airflow/ui/src/pages/DagsList.tsx +++ b/airflow/ui/src/pages/DagsList.tsx @@ -30,7 +30,7 @@ import { Select as ReactSelect } from "chakra-react-select"; import { type ChangeEventHandler, useCallback } from "react"; import { useSearchParams } from "react-router-dom"; -import { useDagServiceGetDagsPublicDagsGet } from "openapi/queries"; +import { useDagServiceGetDags } from "openapi/queries"; import type { DAGResponse } from "openapi/requests/types.gen"; import { DataTable } from "../components/DataTable"; @@ -93,7 +93,7 @@ export const DagsList = ({ cardView = false }) => { const [sort] = sorting; const orderBy = sort ? `${sort.desc ? "-" : ""}${sort.id}` : undefined; - const { data, isLoading } = useDagServiceGetDagsPublicDagsGet({ + const { data, isLoading } = useDagServiceGetDags({ limit: pagination.pageSize, offset: pagination.pageIndex * pagination.pageSize, onlyActive: true, From f0f53c2fe9a8747019ad4b39b743bd9cc1543545 Mon Sep 17 00:00:00 2001 From: Jarek Potiuk Date: Tue, 1 Oct 2024 01:02:38 -0700 Subject: [PATCH 082/802] Limit build-images workflow to main and v2-10 branches (#42601) There is no need to run image builds for PRs to old branches. --- .github/workflows/build-images.yml | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/.github/workflows/build-images.yml b/.github/workflows/build-images.yml index 1256fd2f0da6e..abf966faede02 100644 --- a/.github/workflows/build-images.yml +++ b/.github/workflows/build-images.yml @@ -21,6 +21,10 @@ run-name: > Build images for ${{ github.event.pull_request.title }} ${{ github.event.pull_request._links.html.href }} on: # yamllint disable-line rule:truthy pull_request_target: + branches: + - main + - v2-10-stable + - v2-10-test permissions: # all other permissions are set to none contents: read From da4ee4fe4dac46342e80e5c19cc8da0b4428fbd8 Mon Sep 17 00:00:00 2001 From: Jakub Dardzinski Date: Tue, 1 Oct 2024 14:16:03 +0200 Subject: [PATCH 083/802] openlineage: add unit test for listener hooks on dag run state changes. (#42554) openlineage: cover task instance failure in unit tests. Signed-off-by: Jakub Dardzinski --- tests/dags/test_openlineage_execution.py | 12 +++++- .../openlineage/plugins/test_execution.py | 11 +++++ .../openlineage/plugins/test_listener.py | 43 +++++++++++++++++++ 3 files changed, 65 insertions(+), 1 deletion(-) diff --git a/tests/dags/test_openlineage_execution.py b/tests/dags/test_openlineage_execution.py index 475e43ef6ac2e..f8db91611e848 100644 --- a/tests/dags/test_openlineage_execution.py +++ b/tests/dags/test_openlineage_execution.py @@ -27,13 +27,16 @@ class OpenLineageExecutionOperator(BaseOperator): - def __init__(self, *, stall_amount=0, **kwargs) -> None: + def __init__(self, *, stall_amount=0, fail=False, **kwargs) -> None: super().__init__(**kwargs) self.stall_amount = stall_amount + self.fail = fail def execute(self, context): self.log.error("STALL AMOUNT %s", self.stall_amount) time.sleep(1) + if self.fail: + raise Exception("Failed") def get_openlineage_facets_on_start(self): return OperatorLineage(inputs=[Dataset(namespace="test", name="on-start")]) @@ -43,6 +46,11 @@ def get_openlineage_facets_on_complete(self, task_instance): time.sleep(self.stall_amount) return OperatorLineage(inputs=[Dataset(namespace="test", name="on-complete")]) + def get_openlineage_facets_on_failure(self, task_instance): + self.log.error("STALL AMOUNT %s", self.stall_amount) + time.sleep(self.stall_amount) + return OperatorLineage(inputs=[Dataset(namespace="test", name="on-failure")]) + with DAG( dag_id="test_openlineage_execution", @@ -57,3 +65,5 @@ def get_openlineage_facets_on_complete(self, task_instance): mid_stall = OpenLineageExecutionOperator(task_id="execute_mid_stall", stall_amount=15) long_stall = OpenLineageExecutionOperator(task_id="execute_long_stall", stall_amount=30) + + fail = OpenLineageExecutionOperator(task_id="execute_fail", fail=True) diff --git a/tests/providers/openlineage/plugins/test_execution.py b/tests/providers/openlineage/plugins/test_execution.py index 3adaaac582dd7..8c0bdd55a1f96 100644 --- a/tests/providers/openlineage/plugins/test_execution.py +++ b/tests/providers/openlineage/plugins/test_execution.py @@ -124,6 +124,17 @@ def test_not_stalled_task_emits_proper_lineage(self): assert has_value_in_events(events, ["inputs", "name"], "on-start") assert has_value_in_events(events, ["inputs", "name"], "on-complete") + @pytest.mark.db_test + @conf_vars({("openlineage", "transport"): f'{{"type": "file", "log_file_path": "{listener_path}"}}'}) + def test_not_stalled_failing_task_emits_proper_lineage(self): + task_name = "execute_fail" + run_id = "test_failure" + self.setup_job(task_name, run_id) + + events = get_sorted_events(tmp_dir) + assert has_value_in_events(events, ["inputs", "name"], "on-start") + assert has_value_in_events(events, ["inputs", "name"], "on-failure") + @conf_vars( { ("openlineage", "transport"): f'{{"type": "file", "log_file_path": "{listener_path}"}}', diff --git a/tests/providers/openlineage/plugins/test_listener.py b/tests/providers/openlineage/plugins/test_listener.py index 92467a58af8c5..57c0134f79d82 100644 --- a/tests/providers/openlineage/plugins/test_listener.py +++ b/tests/providers/openlineage/plugins/test_listener.py @@ -606,6 +606,49 @@ def test_listener_on_dag_run_state_changes_configure_process_pool_size(mock_exec mock_executor.return_value.submit.assert_called_once() +class MockExecutor: + def __init__(self, *args, **kwargs): + self.submitted = False + self.succeeded = False + self.result = None + + def submit(self, fn, /, *args, **kwargs): + self.submitted = True + try: + fn(*args, **kwargs) + self.succeeded = True + except Exception: + pass + return MagicMock() + + def shutdown(self, *args, **kwargs): + print("Shutting down") + + +@pytest.mark.parametrize( + ("method", "dag_run_state"), + [ + ("on_dag_run_running", DagRunState.RUNNING), + ("on_dag_run_success", DagRunState.SUCCESS), + ("on_dag_run_failed", DagRunState.FAILED), + ], +) +@patch("airflow.providers.openlineage.plugins.adapter.OpenLineageAdapter.emit") +def test_listener_on_dag_run_state_changes(mock_emit, method, dag_run_state, create_task_instance): + mock_executor = MockExecutor() + ti = create_task_instance(dag_id="dag", task_id="op") + # Change the state explicitly to set end_date following the logic in the method + ti.dag_run.set_state(dag_run_state) + with mock.patch( + "airflow.providers.openlineage.plugins.listener.ProcessPoolExecutor", return_value=mock_executor + ): + listener = OpenLineageListener() + getattr(listener, method)(ti.dag_run, None) + assert mock_executor.submitted is True + assert mock_executor.succeeded is True + mock_emit.assert_called_once() + + def test_listener_logs_failed_serialization(): listener = OpenLineageListener() callback_future = Future() From eda8089f2f0c664722823915c3b2172594b82b77 Mon Sep 17 00:00:00 2001 From: Vincent <97131062+vincbeck@users.noreply.github.com> Date: Tue, 1 Oct 2024 09:59:11 -0400 Subject: [PATCH 084/802] Update Rest API tests to no longer rely on FAB auth manager. Move tests specific to FAB permissions to FAB provider (#42523) --- .../managers/simple/simple_auth_manager.py | 7 +- airflow/auth/managers/simple/user.py | 6 +- .../0034_3_0_0_update_user_id_type.py | 52 +++ ..._3_0_0_add_name_field_to_dataset_model.py} | 4 +- airflow/models/dagrun.py | 2 +- airflow/models/taskinstance.py | 2 +- .../api/auth/backend/basic_auth.py | 4 +- .../api/auth/backend/kerberos_auth.py | 2 +- .../fab/auth_manager/models/anonymous_user.py | 2 +- docs/apache-airflow/img/airflow_erd.sha256 | 2 +- docs/apache-airflow/img/airflow_erd.svg | 4 +- docs/apache-airflow/migrations-ref.rst | 5 +- tests/api_connexion/conftest.py | 11 +- .../endpoints/test_backfill_endpoint.py | 31 +- .../endpoints/test_config_endpoint.py | 12 +- .../endpoints/test_connection_endpoint.py | 17 +- .../endpoints/test_dag_endpoint.py | 100 +--- .../endpoints/test_dag_parsing.py | 16 +- .../endpoints/test_dag_run_endpoint.py | 154 +------ .../endpoints/test_dag_source_endpoint.py | 71 +-- .../endpoints/test_dag_stats_endpoint.py | 15 +- .../endpoints/test_dag_warning_endpoint.py | 33 +- .../endpoints/test_dataset_endpoint.py | 234 +--------- .../endpoints/test_event_log_endpoint.py | 80 +--- .../endpoints/test_extra_link_endpoint.py | 20 +- .../endpoints/test_import_error_endpoint.py | 170 +------ .../endpoints/test_log_endpoint.py | 9 +- .../test_mapped_task_instance_endpoint.py | 25 +- .../endpoints/test_plugin_endpoint.py | 12 +- .../endpoints/test_pool_endpoint.py | 17 +- .../endpoints/test_provider_endpoint.py | 12 +- .../endpoints/test_task_endpoint.py | 16 +- .../endpoints/test_task_instance_endpoint.py | 219 +-------- .../endpoints/test_variable_endpoint.py | 37 +- .../endpoints/test_xcom_endpoint.py | 74 +-- tests/api_connexion/test_auth.py | 188 ++------ tests/api_connexion/test_security.py | 8 +- .../api_endpoints/api_connexion_utils.py | 116 +++++ .../remote_user_api_auth_backend.py | 81 ++++ .../auth_manager/api_endpoints/test_auth.py | 176 ++++++++ .../api_endpoints/test_backfill_endpoint.py | 264 +++++++++++ .../auth_manager/api_endpoints}/test_cors.py | 35 +- .../api_endpoints/test_dag_endpoint.py | 252 +++++++++++ .../api_endpoints/test_dag_run_endpoint.py | 273 +++++++++++ .../api_endpoints/test_dag_source_endpoint.py | 144 ++++++ .../test_dag_warning_endpoint.py | 84 ++++ .../api_endpoints/test_dataset_endpoint.py | 327 ++++++++++++++ .../api_endpoints/test_event_log_endpoint.py | 151 +++++++ .../test_import_error_endpoint.py | 221 +++++++++ .../test_role_and_permission_endpoint.py | 22 +- .../test_role_and_permission_schema.py | 22 +- .../test_task_instance_endpoint.py | 427 ++++++++++++++++++ .../api_endpoints/test_user_endpoint.py | 15 +- .../api_endpoints/test_user_schema.py | 3 +- .../api_endpoints/test_variable_endpoint.py | 88 ++++ .../api_endpoints/test_xcom_endpoint.py | 230 ++++++++++ tests/providers/fab/auth_manager/conftest.py | 17 +- .../fab/auth_manager/test_security.py | 2 +- .../auth_manager/views/test_permissions.py | 2 +- .../fab/auth_manager/views/test_roles_list.py | 2 +- .../fab/auth_manager/views/test_user.py | 2 +- .../fab/auth_manager/views/test_user_edit.py | 2 +- .../fab/auth_manager/views/test_user_stats.py | 2 +- tests/test_utils/api_connexion_utils.py | 64 +-- .../remote_user_api_auth_backend.py | 32 +- .../www/views/test_views_custom_user_views.py | 5 +- tests/www/views/test_views_dagrun.py | 6 +- tests/www/views/test_views_home.py | 2 +- tests/www/views/test_views_tasks.py | 6 +- tests/www/views/test_views_variable.py | 2 +- 70 files changed, 3208 insertions(+), 1542 deletions(-) create mode 100644 airflow/migrations/versions/0034_3_0_0_update_user_id_type.py rename airflow/migrations/versions/{0034_3_0_0_add_name_field_to_dataset_model.py => 0035_3_0_0_add_name_field_to_dataset_model.py} (98%) create mode 100644 tests/providers/fab/auth_manager/api_endpoints/api_connexion_utils.py create mode 100644 tests/providers/fab/auth_manager/api_endpoints/remote_user_api_auth_backend.py create mode 100644 tests/providers/fab/auth_manager/api_endpoints/test_auth.py create mode 100644 tests/providers/fab/auth_manager/api_endpoints/test_backfill_endpoint.py rename tests/{api_connexion => providers/fab/auth_manager/api_endpoints}/test_cors.py (81%) create mode 100644 tests/providers/fab/auth_manager/api_endpoints/test_dag_endpoint.py create mode 100644 tests/providers/fab/auth_manager/api_endpoints/test_dag_run_endpoint.py create mode 100644 tests/providers/fab/auth_manager/api_endpoints/test_dag_source_endpoint.py create mode 100644 tests/providers/fab/auth_manager/api_endpoints/test_dag_warning_endpoint.py create mode 100644 tests/providers/fab/auth_manager/api_endpoints/test_dataset_endpoint.py create mode 100644 tests/providers/fab/auth_manager/api_endpoints/test_event_log_endpoint.py create mode 100644 tests/providers/fab/auth_manager/api_endpoints/test_import_error_endpoint.py rename tests/{api_connexion/schemas => providers/fab/auth_manager/api_endpoints}/test_role_and_permission_schema.py (85%) create mode 100644 tests/providers/fab/auth_manager/api_endpoints/test_task_instance_endpoint.py create mode 100644 tests/providers/fab/auth_manager/api_endpoints/test_variable_endpoint.py create mode 100644 tests/providers/fab/auth_manager/api_endpoints/test_xcom_endpoint.py diff --git a/airflow/auth/managers/simple/simple_auth_manager.py b/airflow/auth/managers/simple/simple_auth_manager.py index 451068733667c..4a9639a998c46 100644 --- a/airflow/auth/managers/simple/simple_auth_manager.py +++ b/airflow/auth/managers/simple/simple_auth_manager.py @@ -221,7 +221,12 @@ def _is_authorized( user = self.get_user() if not user: return False - role_str = user.get_role().upper() + + user_role = user.get_role() + if not user_role: + return False + + role_str = user_role.upper() role = SimpleAuthManagerRole[role_str] if role == SimpleAuthManagerRole.ADMIN: return True diff --git a/airflow/auth/managers/simple/user.py b/airflow/auth/managers/simple/user.py index fa032f596ee44..f4591b0b1c751 100644 --- a/airflow/auth/managers/simple/user.py +++ b/airflow/auth/managers/simple/user.py @@ -24,10 +24,10 @@ class SimpleAuthManagerUser(BaseUser): User model for users managed by the simple auth manager. :param username: The username - :param role: The role associated to the user + :param role: The role associated to the user. If not provided, the user has no permission """ - def __init__(self, *, username: str, role: str) -> None: + def __init__(self, *, username: str, role: str | None) -> None: self.username = username self.role = role @@ -37,5 +37,5 @@ def get_id(self) -> str: def get_name(self) -> str: return self.username - def get_role(self): + def get_role(self) -> str | None: return self.role diff --git a/airflow/migrations/versions/0034_3_0_0_update_user_id_type.py b/airflow/migrations/versions/0034_3_0_0_update_user_id_type.py new file mode 100644 index 0000000000000..321a1e2bbafa8 --- /dev/null +++ b/airflow/migrations/versions/0034_3_0_0_update_user_id_type.py @@ -0,0 +1,52 @@ +# +# 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. + +""" +Update dag_run_note.user_id and task_instance_note.user_id columns to String. + +Revision ID: 44eabb1904b4 +Revises: 16cbcb1c8c36 +Create Date: 2024-09-27 09:57:29.830521 + +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op + +# revision identifiers, used by Alembic. +revision = "44eabb1904b4" +down_revision = "16cbcb1c8c36" +branch_labels = None +depends_on = None +airflow_version = "3.0.0" + + +def upgrade(): + with op.batch_alter_table("dag_run_note") as batch_op: + batch_op.alter_column("user_id", type_=sa.String(length=128)) + with op.batch_alter_table("task_instance_note") as batch_op: + batch_op.alter_column("user_id", type_=sa.String(length=128)) + + +def downgrade(): + with op.batch_alter_table("dag_run_note") as batch_op: + batch_op.alter_column("user_id", type_=sa.Integer(), postgresql_using="user_id::integer") + with op.batch_alter_table("task_instance_note") as batch_op: + batch_op.alter_column("user_id", type_=sa.Integer(), postgresql_using="user_id::integer") diff --git a/airflow/migrations/versions/0034_3_0_0_add_name_field_to_dataset_model.py b/airflow/migrations/versions/0035_3_0_0_add_name_field_to_dataset_model.py similarity index 98% rename from airflow/migrations/versions/0034_3_0_0_add_name_field_to_dataset_model.py rename to airflow/migrations/versions/0035_3_0_0_add_name_field_to_dataset_model.py index 5c8aec69e9be9..6016dd9658908 100644 --- a/airflow/migrations/versions/0034_3_0_0_add_name_field_to_dataset_model.py +++ b/airflow/migrations/versions/0035_3_0_0_add_name_field_to_dataset_model.py @@ -30,7 +30,7 @@ also rename the one on DatasetAliasModel here for consistency. Revision ID: 0d9e73a75ee4 -Revises: 16cbcb1c8c36 +Revises: 44eabb1904b4 Create Date: 2024-08-13 09:45:32.213222 """ @@ -42,7 +42,7 @@ # revision identifiers, used by Alembic. revision = "0d9e73a75ee4" -down_revision = "16cbcb1c8c36" +down_revision = "44eabb1904b4" branch_labels = None depends_on = None airflow_version = "3.0.0" diff --git a/airflow/models/dagrun.py b/airflow/models/dagrun.py index 5d53e51763dff..4928c7fcbd8f7 100644 --- a/airflow/models/dagrun.py +++ b/airflow/models/dagrun.py @@ -1687,7 +1687,7 @@ class DagRunNote(Base): __tablename__ = "dag_run_note" - user_id = Column(Integer, nullable=True) + user_id = Column(String(128), nullable=True) dag_run_id = Column(Integer, primary_key=True, nullable=False) content = Column(String(1000).with_variant(Text(1000), "mysql")) created_at = Column(UtcDateTime, default=timezone.utcnow, nullable=False) diff --git a/airflow/models/taskinstance.py b/airflow/models/taskinstance.py index b19e65486307d..333a4cad91cbe 100644 --- a/airflow/models/taskinstance.py +++ b/airflow/models/taskinstance.py @@ -4002,7 +4002,7 @@ class TaskInstanceNote(TaskInstanceDependencies): __tablename__ = "task_instance_note" - user_id = Column(Integer, nullable=True) + user_id = Column(String(128), nullable=True) task_id = Column(StringID(), primary_key=True, nullable=False) dag_id = Column(StringID(), primary_key=True, nullable=False) run_id = Column(StringID(), primary_key=True, nullable=False) diff --git a/airflow/providers/fab/auth_manager/api/auth/backend/basic_auth.py b/airflow/providers/fab/auth_manager/api/auth/backend/basic_auth.py index 3a0328fe9962c..7b50338733453 100644 --- a/airflow/providers/fab/auth_manager/api/auth/backend/basic_auth.py +++ b/airflow/providers/fab/auth_manager/api/auth/backend/basic_auth.py @@ -62,9 +62,7 @@ def requires_authentication(function: T): @wraps(function) def decorated(*args, **kwargs): - if auth_current_user() is not None or current_app.appbuilder.get_app.config.get( - "AUTH_ROLE_PUBLIC", None - ): + if auth_current_user() is not None or current_app.config.get("AUTH_ROLE_PUBLIC", None): return function(*args, **kwargs) else: return Response("Unauthorized", 401, {"WWW-Authenticate": "Basic"}) diff --git a/airflow/providers/fab/auth_manager/api/auth/backend/kerberos_auth.py b/airflow/providers/fab/auth_manager/api/auth/backend/kerberos_auth.py index d8d5a95ee676b..f2038b27597c1 100644 --- a/airflow/providers/fab/auth_manager/api/auth/backend/kerberos_auth.py +++ b/airflow/providers/fab/auth_manager/api/auth/backend/kerberos_auth.py @@ -124,7 +124,7 @@ def requires_authentication(function: T, find_user: Callable[[str], BaseUser] | @wraps(function) def decorated(*args, **kwargs): - if current_app.appbuilder.get_app.config.get("AUTH_ROLE_PUBLIC", None): + if current_app.config.get("AUTH_ROLE_PUBLIC", None): response = function(*args, **kwargs) return make_response(response) diff --git a/airflow/providers/fab/auth_manager/models/anonymous_user.py b/airflow/providers/fab/auth_manager/models/anonymous_user.py index 2f294fd9e5d0e..9afb2cdff635f 100644 --- a/airflow/providers/fab/auth_manager/models/anonymous_user.py +++ b/airflow/providers/fab/auth_manager/models/anonymous_user.py @@ -35,7 +35,7 @@ class AnonymousUser(AnonymousUserMixin, BaseUser): @property def roles(self): if not self._roles: - public_role = current_app.appbuilder.get_app.config.get("AUTH_ROLE_PUBLIC", None) + public_role = current_app.config.get("AUTH_ROLE_PUBLIC", None) self._roles = {current_app.appbuilder.sm.find_role(public_role)} if public_role else set() return self._roles diff --git a/docs/apache-airflow/img/airflow_erd.sha256 b/docs/apache-airflow/img/airflow_erd.sha256 index e4a952da1b9fd..bca068fde6749 100644 --- a/docs/apache-airflow/img/airflow_erd.sha256 +++ b/docs/apache-airflow/img/airflow_erd.sha256 @@ -1 +1 @@ -c33e9a583a5b29eb748ebd50e117643e11bcb2a9b61ec017efd690621e22769b \ No newline at end of file +64dfad12dfd49f033c4723c2f3bb3bac58dd956136fb24a87a2e5a6ae176ec1a \ No newline at end of file diff --git a/docs/apache-airflow/img/airflow_erd.svg b/docs/apache-airflow/img/airflow_erd.svg index 76fbd8f841f25..4eb6c2ee70917 100644 --- a/docs/apache-airflow/img/airflow_erd.svg +++ b/docs/apache-airflow/img/airflow_erd.svg @@ -1394,7 +1394,7 @@ user_id - [INTEGER] + [VARCHAR(100)] @@ -1813,7 +1813,7 @@ user_id - [INTEGER] + [VARCHAR(100)] diff --git a/docs/apache-airflow/migrations-ref.rst b/docs/apache-airflow/migrations-ref.rst index a547d03d75be6..e4fb2dfa332eb 100644 --- a/docs/apache-airflow/migrations-ref.rst +++ b/docs/apache-airflow/migrations-ref.rst @@ -39,7 +39,10 @@ Here's the list of all the Database Migrations that are executed via when you ru +-------------------------+------------------+-------------------+--------------------------------------------------------------+ | Revision ID | Revises ID | Airflow Version | Description | +=========================+==================+===================+==============================================================+ -| ``0d9e73a75ee4`` (head) | ``16cbcb1c8c36`` | ``3.0.0`` | Add name and group fields to DatasetModel. | +| ``0d9e73a75ee4`` (head) | ``44eabb1904b4`` | ``3.0.0`` | Add name and group fields to DatasetModel. | ++-------------------------+------------------+-------------------+--------------------------------------------------------------+ +| ``44eabb1904b4`` | ``16cbcb1c8c36`` | ``3.0.0`` | Update dag_run_note.user_id and task_instance_note.user_id | +| | | | columns to String. | +-------------------------+------------------+-------------------+--------------------------------------------------------------+ | ``16cbcb1c8c36`` | ``522625f6d606`` | ``3.0.0`` | Remove redundant index. | +-------------------------+------------------+-------------------+--------------------------------------------------------------+ diff --git a/tests/api_connexion/conftest.py b/tests/api_connexion/conftest.py index 38e7b58cb5981..6a23b2cf11d93 100644 --- a/tests/api_connexion/conftest.py +++ b/tests/api_connexion/conftest.py @@ -36,9 +36,16 @@ def minimal_app_for_api(): ] ) def factory(): - with conf_vars({("api", "auth_backends"): "tests.test_utils.remote_user_api_auth_backend"}): + with conf_vars( + { + ("api", "auth_backends"): "tests.test_utils.remote_user_api_auth_backend", + ( + "core", + "auth_manager", + ): "airflow.auth.managers.simple.simple_auth_manager.SimpleAuthManager", + } + ): _app = app.create_app(testing=True, config={"WTF_CSRF_ENABLED": False}) # type:ignore - _app.config["AUTH_ROLE_PUBLIC"] = None return _app return factory() diff --git a/tests/api_connexion/endpoints/test_backfill_endpoint.py b/tests/api_connexion/endpoints/test_backfill_endpoint.py index 51a4faf40055c..07b2a3fd56c2d 100644 --- a/tests/api_connexion/endpoints/test_backfill_endpoint.py +++ b/tests/api_connexion/endpoints/test_backfill_endpoint.py @@ -29,7 +29,6 @@ from airflow.models.dag import DAG from airflow.models.serialized_dag import SerializedDagModel from airflow.operators.empty import EmptyOperator -from airflow.security import permissions from airflow.utils import timezone from airflow.utils.session import provide_session from tests.test_utils.api_connexion_utils import create_user, delete_user @@ -50,25 +49,11 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type: ignore + app, username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_DAG), - ], - ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore - create_user(app, username="test_granular_permissions", role_name="TestGranularDag") # type: ignore - app.appbuilder.sm.sync_perm_for_dag( # type: ignore - "TEST_DAG_1", - access_control={ - "TestGranularDag": { - permissions.RESOURCE_DAG: {permissions.ACTION_CAN_EDIT, permissions.ACTION_CAN_READ} - }, - }, + role_name="admin", ) + create_user(app, username="test_no_permissions", role_name=None) with DAG( DAG_ID, @@ -93,9 +78,8 @@ def configured_app(minimal_app_for_api): yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore - delete_user(app, username="test_granular_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestBackfillEndpoint: @@ -178,7 +162,6 @@ def test_should_respond_200(self, session): @pytest.mark.parametrize( "user, expected", [ - ("test_granular_permissions", 200), ("test_no_permissions", 403), ("test", 200), (None, 401), @@ -240,7 +223,6 @@ def test_no_exist(self, session): @pytest.mark.parametrize( "user, expected", [ - ("test_granular_permissions", 200), ("test_no_permissions", 403), ("test", 200), (None, 401), @@ -268,7 +250,6 @@ class TestCreateBackfill(TestBackfillEndpoint): @pytest.mark.parametrize( "user, expected", [ - ("test_granular_permissions", 200), ("test_no_permissions", 403), ("test", 200), (None, 401), @@ -347,7 +328,6 @@ def test_should_respond_200(self, session): @pytest.mark.parametrize( "user, expected", [ - ("test_granular_permissions", 200), ("test_no_permissions", 403), ("test", 200), (None, 401), @@ -409,7 +389,6 @@ def test_should_respond_200(self, session): @pytest.mark.parametrize( "user, expected", [ - ("test_granular_permissions", 200), ("test_no_permissions", 403), ("test", 200), (None, 401), diff --git a/tests/api_connexion/endpoints/test_config_endpoint.py b/tests/api_connexion/endpoints/test_config_endpoint.py index 475753a4a902e..bd88c491c952b 100644 --- a/tests/api_connexion/endpoints/test_config_endpoint.py +++ b/tests/api_connexion/endpoints/test_config_endpoint.py @@ -21,7 +21,6 @@ import pytest -from airflow.security import permissions from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_user from tests.test_utils.config import conf_vars @@ -54,18 +53,17 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type:ignore + app, username="test", - role_name="Test", - permissions=[(permissions.ACTION_CAN_READ, permissions.RESOURCE_CONFIG)], # type: ignore + role_name="admin", ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name=None) with conf_vars({("webserver", "expose_config"): "True"}): yield minimal_app_for_api - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestGetConfig: diff --git a/tests/api_connexion/endpoints/test_connection_endpoint.py b/tests/api_connexion/endpoints/test_connection_endpoint.py index a19b046aa2747..a140046656e31 100644 --- a/tests/api_connexion/endpoints/test_connection_endpoint.py +++ b/tests/api_connexion/endpoints/test_connection_endpoint.py @@ -24,7 +24,6 @@ from airflow.api_connexion.exceptions import EXCEPTIONS_LINK_MAP from airflow.models import Connection from airflow.secrets.environment_variables import CONN_ENV_PREFIX -from airflow.security import permissions from airflow.utils.session import provide_session from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_user from tests.test_utils.config import conf_vars @@ -38,22 +37,16 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type: ignore + app, username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_CONNECTION), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_CONNECTION), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_CONNECTION), - (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_CONNECTION), - ], + role_name="admin", ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestConnectionEndpoint: diff --git a/tests/api_connexion/endpoints/test_dag_endpoint.py b/tests/api_connexion/endpoints/test_dag_endpoint.py index 9905b4e27ab2c..6d4ffc2d06d2c 100644 --- a/tests/api_connexion/endpoints/test_dag_endpoint.py +++ b/tests/api_connexion/endpoints/test_dag_endpoint.py @@ -28,7 +28,6 @@ from airflow.models.dag import DAG from airflow.models.serialized_dag import SerializedDagModel from airflow.operators.empty import EmptyOperator -from airflow.security import permissions from airflow.utils.session import provide_session from airflow.utils.state import TaskInstanceState from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_user @@ -56,33 +55,11 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type: ignore + app, username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_DAG), - ], - ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore - create_user(app, username="test_granular_permissions", role_name="TestGranularDag") # type: ignore - app.appbuilder.sm.sync_perm_for_dag( # type: ignore - "TEST_DAG_1", - access_control={ - "TestGranularDag": { - permissions.RESOURCE_DAG: {permissions.ACTION_CAN_EDIT, permissions.ACTION_CAN_READ} - }, - }, - ) - app.appbuilder.sm.sync_perm_for_dag( # type: ignore - "TEST_DAG_1", - access_control={ - "TestGranularDag": { - permissions.RESOURCE_DAG: {permissions.ACTION_CAN_EDIT, permissions.ACTION_CAN_READ} - }, - }, + role_name="admin", ) + create_user(app, username="test_no_permissions", role_name=None) with DAG( DAG_ID, @@ -107,9 +84,8 @@ def configured_app(minimal_app_for_api): yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore - delete_user(app, username="test_granular_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestDagEndpoint: @@ -258,13 +234,6 @@ def test_should_respond_200_with_schedule_none(self, session): "pickle_id": None, } == response.json - def test_should_respond_200_with_granular_dag_access(self): - self._create_dag_models(1) - response = self.client.get( - "/api/v1/dags/TEST_DAG_1", environ_overrides={"REMOTE_USER": "test_granular_permissions"} - ) - assert response.status_code == 200 - def test_should_respond_404(self): response = self.client.get("/api/v1/dags/INVALID_DAG", environ_overrides={"REMOTE_USER": "test"}) assert response.status_code == 404 @@ -282,13 +251,6 @@ def test_should_raise_403_forbidden(self): ) assert response.status_code == 403 - def test_should_respond_403_with_granular_access_for_different_dag(self): - self._create_dag_models(3) - response = self.client.get( - "/api/v1/dags/TEST_DAG_2", environ_overrides={"REMOTE_USER": "test_granular_permissions"} - ) - assert response.status_code == 403 - @pytest.mark.parametrize( "fields", [ @@ -961,15 +923,6 @@ def test_filter_dags_by_dag_id_works(self, url, expected_dag_ids): assert expected_dag_ids == dag_ids - def test_should_respond_200_with_granular_dag_access(self): - self._create_dag_models(3) - response = self.client.get( - "/api/v1/dags", environ_overrides={"REMOTE_USER": "test_granular_permissions"} - ) - assert response.status_code == 200 - assert len(response.json["dags"]) == 1 - assert response.json["dags"][0]["dag_id"] == "TEST_DAG_1" - @pytest.mark.parametrize( "url, expected_dag_ids", [ @@ -1252,18 +1205,6 @@ def test_should_respond_200_on_patch_is_paused(self, url_safe_serializer, sessio session, dag_id="TEST_DAG_1", event="api.patch_dag", execution_date=None, expected_extra=payload ) - def test_should_respond_200_on_patch_with_granular_dag_access(self, session): - self._create_dag_models(1) - response = self.client.patch( - "/api/v1/dags/TEST_DAG_1", - json={ - "is_paused": False, - }, - environ_overrides={"REMOTE_USER": "test_granular_permissions"}, - ) - assert response.status_code == 200 - _check_last_log(session, dag_id="TEST_DAG_1", event="api.patch_dag", execution_date=None) - def test_should_respond_400_on_invalid_request(self): patch_body = { "is_paused": True, @@ -1279,24 +1220,6 @@ def test_should_respond_400_on_invalid_request(self): "type": EXCEPTIONS_LINK_MAP[400], } - def test_validation_error_raises_400(self): - patch_body = { - "ispaused": True, - } - dag_model = self._create_dag_model() - response = self.client.patch( - f"/api/v1/dags/{dag_model.dag_id}", - json=patch_body, - environ_overrides={"REMOTE_USER": "test_granular_permissions"}, - ) - assert response.status_code == 400 - assert response.json == { - "detail": "{'ispaused': ['Unknown field.']}", - "status": 400, - "title": "Bad Request", - "type": EXCEPTIONS_LINK_MAP[400], - } - def test_non_existing_dag_raises_not_found(self): patch_body = { "is_paused": True, @@ -1820,19 +1743,6 @@ def test_filter_dags_by_dag_id_works(self, url, expected_dag_ids): assert expected_dag_ids == dag_ids - def test_should_respond_200_with_granular_dag_access(self): - self._create_dag_models(3) - response = self.client.patch( - "api/v1/dags?dag_id_pattern=~", - json={ - "is_paused": False, - }, - environ_overrides={"REMOTE_USER": "test_granular_permissions"}, - ) - assert response.status_code == 200 - assert len(response.json["dags"]) == 1 - assert response.json["dags"][0]["dag_id"] == "TEST_DAG_1" - @pytest.mark.parametrize( "url, expected_dag_ids", [ diff --git a/tests/api_connexion/endpoints/test_dag_parsing.py b/tests/api_connexion/endpoints/test_dag_parsing.py index 521d8d9e8cd99..ae42a565dd052 100644 --- a/tests/api_connexion/endpoints/test_dag_parsing.py +++ b/tests/api_connexion/endpoints/test_dag_parsing.py @@ -24,7 +24,6 @@ from airflow.models import DagBag from airflow.models.dagbag import DagPriorityParsingRequest -from airflow.security import permissions from tests.test_utils.api_connexion_utils import create_user, delete_user from tests.test_utils.db import clear_db_dag_parsing_requests @@ -45,21 +44,16 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type:ignore + app, username="test", - role_name="Test", - permissions=[(permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG)], # type: ignore + role_name="admin", ) - app.appbuilder.sm.sync_perm_for_dag( # type: ignore - TEST_DAG_ID, - access_control={"Test": [permissions.ACTION_CAN_EDIT]}, - ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestDagParsingRequest: diff --git a/tests/api_connexion/endpoints/test_dag_run_endpoint.py b/tests/api_connexion/endpoints/test_dag_run_endpoint.py index f3921da7b9c29..73c75b98a43b1 100644 --- a/tests/api_connexion/endpoints/test_dag_run_endpoint.py +++ b/tests/api_connexion/endpoints/test_dag_run_endpoint.py @@ -30,12 +30,11 @@ from airflow.models.dagrun import DagRun from airflow.models.param import Param from airflow.operators.empty import EmptyOperator -from airflow.security import permissions from airflow.utils import timezone from airflow.utils.session import create_session, provide_session from airflow.utils.state import DagRunState, State from airflow.utils.types import DagRunType -from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_roles, delete_user +from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_user from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS from tests.test_utils.config import conf_vars from tests.test_utils.db import clear_db_dags, clear_db_runs, clear_db_serialized_dags @@ -52,79 +51,16 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type: ignore + app, username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_ASSET), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_CLUSTER_ACTIVITY), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_DAG_RUN), - ], - ) - create_user( - app, # type: ignore - username="test_no_dag_run_create_permission", - role_name="TestNoDagRunCreatePermission", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_ASSET), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_CLUSTER_ACTIVITY), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_DAG_RUN), - ], - ) - create_user( - app, # type: ignore - username="test_dag_view_only", - role_name="TestViewDags", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_DAG_RUN), - ], - ) - create_user( - app, # type: ignore - username="test_view_dags", - role_name="TestViewDags", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_DAG_RUN), - ], + role_name="admin", ) - create_user( - app, # type: ignore - username="test_granular_permissions", - role_name="TestGranularDag", - permissions=[(permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN)], - ) - app.appbuilder.sm.sync_perm_for_dag( # type: ignore - "TEST_DAG_ID", - access_control={ - "TestGranularDag": {permissions.ACTION_CAN_EDIT, permissions.ACTION_CAN_READ}, - "TestNoDagRunCreatePermission": {permissions.RESOURCE_DAG_RUN: {permissions.ACTION_CAN_CREATE}}, - }, - ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_dag_view_only") # type: ignore - delete_user(app, username="test_view_dags") # type: ignore - delete_user(app, username="test_granular_permissions") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore - delete_user(app, username="test_no_dag_run_create_permission") # type: ignore - delete_roles(app) + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestDagRunEndpoint: @@ -499,16 +435,6 @@ def test_should_return_all_with_tilde_as_dag_id_and_all_dag_permissions(self): dag_run_ids = [dag_run["dag_id"] for dag_run in response.json["dag_runs"]] assert dag_run_ids == expected_dag_run_ids - def test_should_return_accessible_with_tilde_as_dag_id_and_dag_level_permissions(self): - self._create_test_dag_run(extra_dag=True) - expected_dag_run_ids = ["TEST_DAG_ID", "TEST_DAG_ID"] - response = self.client.get( - "api/v1/dags/~/dagRuns", environ_overrides={"REMOTE_USER": "test_granular_permissions"} - ) - assert response.status_code == 200 - dag_run_ids = [dag_run["dag_id"] for dag_run in response.json["dag_runs"]] - assert dag_run_ids == expected_dag_run_ids - def test_should_raises_401_unauthenticated(self): self._create_test_dag_run() @@ -907,57 +833,6 @@ def test_order_by_raises_for_invalid_attr(self): msg = "Ordering with 'dag_ru' is disallowed or the attribute does not exist on the model" assert response.json["detail"] == msg - def test_should_return_accessible_with_tilde_as_dag_id_and_dag_level_permissions(self): - self._create_test_dag_run(extra_dag=True) - expected_response_json_1 = { - "dag_id": "TEST_DAG_ID", - "dag_run_id": "TEST_DAG_RUN_ID_1", - "end_date": None, - "state": "running", - "execution_date": self.default_time, - "logical_date": self.default_time, - "external_trigger": True, - "start_date": self.default_time, - "conf": {}, - "data_interval_end": None, - "data_interval_start": None, - "last_scheduling_decision": None, - "run_type": "manual", - "note": None, - } - expected_response_json_1.update({"triggered_by": "test"} if AIRFLOW_V_3_0_PLUS else {}) - expected_response_json_2 = { - "dag_id": "TEST_DAG_ID", - "dag_run_id": "TEST_DAG_RUN_ID_2", - "end_date": None, - "state": "running", - "execution_date": self.default_time_2, - "logical_date": self.default_time_2, - "external_trigger": True, - "start_date": self.default_time, - "conf": {}, - "data_interval_end": None, - "data_interval_start": None, - "last_scheduling_decision": None, - "run_type": "manual", - "note": None, - } - expected_response_json_2.update({"triggered_by": "test"} if AIRFLOW_V_3_0_PLUS else {}) - - response = self.client.post( - "api/v1/dags/~/dagRuns/list", - json={"dag_ids": []}, - environ_overrides={"REMOTE_USER": "test_granular_permissions"}, - ) - assert response.status_code == 200 - assert response.json == { - "dag_runs": [ - expected_response_json_1, - expected_response_json_2, - ], - "total_entries": 2, - } - @pytest.mark.parametrize( "payload, error", [ @@ -1328,15 +1203,6 @@ def test_raises_validation_error_for_invalid_params(self): assert response.status_code == 400 assert "Invalid input for param" in response.json["detail"] - def test_dagrun_trigger_with_dag_level_permissions(self): - self._create_dag("TEST_DAG_ID") - response = self.client.post( - "api/v1/dags/TEST_DAG_ID/dagRuns", - json={"conf": {"validated_number": 1}}, - environ_overrides={"REMOTE_USER": "test_no_dag_run_create_permission"}, - ) - assert response.status_code == 200 - @mock.patch("airflow.api_connexion.endpoints.dag_run_endpoint.get_airflow_app") def test_dagrun_creation_exception_is_handled(self, mock_get_app, session): self._create_dag("TEST_DAG_ID") @@ -1627,11 +1493,7 @@ def test_should_raises_401_unauthenticated(self): assert_401(response) - @pytest.mark.parametrize( - "username", - ["test_dag_view_only", "test_view_dags", "test_granular_permissions", "test_no_permissions"], - ) - def test_should_raises_403_unauthorized(self, username): + def test_should_raises_403_unauthorized(self): self._create_dag("TEST_DAG_ID") response = self.client.post( "api/v1/dags/TEST_DAG_ID/dagRuns", @@ -1639,7 +1501,7 @@ def test_should_raises_403_unauthorized(self, username): "dag_run_id": "TEST_DAG_RUN_ID_1", "execution_date": self.default_time, }, - environ_overrides={"REMOTE_USER": username}, + environ_overrides={"REMOTE_USER": "test_no_permissions"}, ) assert response.status_code == 403 diff --git a/tests/api_connexion/endpoints/test_dag_source_endpoint.py b/tests/api_connexion/endpoints/test_dag_source_endpoint.py index a8d1224e034c3..f4df56ba629ae 100644 --- a/tests/api_connexion/endpoints/test_dag_source_endpoint.py +++ b/tests/api_connexion/endpoints/test_dag_source_endpoint.py @@ -23,7 +23,6 @@ import pytest from airflow.models import DagBag -from airflow.security import permissions from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_user from tests.test_utils.db import clear_db_dag_code, clear_db_dags, clear_db_serialized_dags @@ -44,29 +43,16 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type:ignore + app, username="test", - role_name="Test", - permissions=[(permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_CODE)], # type: ignore + role_name="admin", ) - app.appbuilder.sm.sync_perm_for_dag( # type: ignore - TEST_DAG_ID, - access_control={"Test": [permissions.ACTION_CAN_READ]}, - ) - app.appbuilder.sm.sync_perm_for_dag( # type: ignore - EXAMPLE_DAG_ID, - access_control={"Test": [permissions.ACTION_CAN_READ]}, - ) - app.appbuilder.sm.sync_perm_for_dag( # type: ignore - TEST_MULTIPLE_DAGS_ID, - access_control={"Test": [permissions.ACTION_CAN_READ]}, - ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestGetSource: @@ -123,18 +109,6 @@ def test_should_respond_200_json(self, url_safe_serializer): assert dag_docstring in response.json["content"] assert "application/json" == response.headers["Content-Type"] - def test_should_respond_406(self, url_safe_serializer): - dagbag = DagBag(dag_folder=EXAMPLE_DAG_FILE) - dagbag.sync_to_db() - test_dag: DAG = dagbag.dags[TEST_DAG_ID] - - url = f"/api/v1/dagSources/{url_safe_serializer.dumps(test_dag.fileloc)}" - response = self.client.get( - url, headers={"Accept": "image/webp"}, environ_overrides={"REMOTE_USER": "test"} - ) - - assert 406 == response.status_code - def test_should_respond_404(self): wrong_fileloc = "abcd1234" url = f"/api/v1/dagSources/{wrong_fileloc}" @@ -167,38 +141,3 @@ def test_should_raise_403_forbidden(self, url_safe_serializer): environ_overrides={"REMOTE_USER": "test_no_permissions"}, ) assert response.status_code == 403 - - def test_should_respond_403_not_readable(self, url_safe_serializer): - dagbag = DagBag(dag_folder=EXAMPLE_DAG_FILE) - dagbag.sync_to_db() - dag: DAG = dagbag.dags[NOT_READABLE_DAG_ID] - - response = self.client.get( - f"/api/v1/dagSources/{url_safe_serializer.dumps(dag.fileloc)}", - headers={"Accept": "text/plain"}, - environ_overrides={"REMOTE_USER": "test"}, - ) - read_dag = self.client.get( - f"/api/v1/dags/{NOT_READABLE_DAG_ID}", - environ_overrides={"REMOTE_USER": "test"}, - ) - assert response.status_code == 403 - assert read_dag.status_code == 403 - - def test_should_respond_403_some_dags_not_readable_in_the_file(self, url_safe_serializer): - dagbag = DagBag(dag_folder=EXAMPLE_DAG_FILE) - dagbag.sync_to_db() - dag: DAG = dagbag.dags[TEST_MULTIPLE_DAGS_ID] - - response = self.client.get( - f"/api/v1/dagSources/{url_safe_serializer.dumps(dag.fileloc)}", - headers={"Accept": "text/plain"}, - environ_overrides={"REMOTE_USER": "test"}, - ) - - read_dag = self.client.get( - f"/api/v1/dags/{TEST_MULTIPLE_DAGS_ID}", - environ_overrides={"REMOTE_USER": "test"}, - ) - assert response.status_code == 403 - assert read_dag.status_code == 200 diff --git a/tests/api_connexion/endpoints/test_dag_stats_endpoint.py b/tests/api_connexion/endpoints/test_dag_stats_endpoint.py index 36fc54d3a5b17..9ab5b49765931 100644 --- a/tests/api_connexion/endpoints/test_dag_stats_endpoint.py +++ b/tests/api_connexion/endpoints/test_dag_stats_endpoint.py @@ -22,7 +22,6 @@ from airflow.models.dag import DAG, DagModel from airflow.models.dagrun import DagRun -from airflow.security import permissions from airflow.utils import timezone from airflow.utils.session import create_session from airflow.utils.state import DagRunState @@ -38,21 +37,17 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type: ignore + app, username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), - ], + role_name="admin", ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestDagStatsEndpoint: diff --git a/tests/api_connexion/endpoints/test_dag_warning_endpoint.py b/tests/api_connexion/endpoints/test_dag_warning_endpoint.py index 3e7c805173b39..f156d8921c0e6 100644 --- a/tests/api_connexion/endpoints/test_dag_warning_endpoint.py +++ b/tests/api_connexion/endpoints/test_dag_warning_endpoint.py @@ -22,7 +22,6 @@ from airflow.models.dag import DagModel from airflow.models.dagwarning import DagWarning -from airflow.security import permissions from airflow.utils.session import create_session from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_user from tests.test_utils.db import clear_db_dag_warnings, clear_db_dags @@ -34,30 +33,16 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type:ignore + app, username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_WARNING), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - ], # type: ignore - ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore - create_user( - app, # type:ignore - username="test_with_dag2_read", - role_name="TestWithDag2Read", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_WARNING), - (permissions.ACTION_CAN_READ, f"{permissions.RESOURCE_DAG_PREFIX}dag2"), - ], # type: ignore + role_name="admin", ) + create_user(app, username="test_no_permissions", role_name=None) yield minimal_app_for_api - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore - delete_user(app, username="test_with_dag2_read") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestBaseDagWarning: @@ -162,11 +147,3 @@ def test_should_raise_403_forbidden(self): "/api/v1/dagWarnings", environ_overrides={"REMOTE_USER": "test_no_permissions"} ) assert response.status_code == 403 - - def test_should_raise_403_forbidden_when_user_has_no_dag_read_permission(self): - response = self.client.get( - "/api/v1/dagWarnings", - environ_overrides={"REMOTE_USER": "test_with_dag2_read"}, - query_string={"dag_id": "dag1"}, - ) - assert response.status_code == 403 diff --git a/tests/api_connexion/endpoints/test_dataset_endpoint.py b/tests/api_connexion/endpoints/test_dataset_endpoint.py index 5caec0ac2a131..76c164654c9d8 100644 --- a/tests/api_connexion/endpoints/test_dataset_endpoint.py +++ b/tests/api_connexion/endpoints/test_dataset_endpoint.py @@ -33,7 +33,6 @@ TaskOutletAssetReference, ) from airflow.models.dagrun import DagRun -from airflow.security import permissions from airflow.utils import timezone from airflow.utils.session import provide_session from airflow.utils.types import DagRunType @@ -50,31 +49,16 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type: ignore + app, username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_ASSET), - (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_ASSET), - ], - ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore - create_user( - app, # type: ignore - username="test_queued_event", - role_name="TestQueuedEvent", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_ASSET), - (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_ASSET), - ], + role_name="admin", ) + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test_queued_event") # type: ignore - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestDatasetEndpoint: @@ -768,43 +752,6 @@ def _create_dataset_dag_run_queues(self, dag_id, dataset_id, session): class TestGetDagDatasetQueuedEvent(TestQueuedEventEndpoint): - @pytest.mark.usefixtures("time_freezer") - def test_should_respond_200(self, session, create_dummy_dag): - dag, _ = create_dummy_dag() - dag_id = dag.dag_id - dataset_id = self._create_dataset(session).id - self._create_dataset_dag_run_queues(dag_id, dataset_id, session) - dataset_uri = "s3://bucket/key" - - response = self.client.get( - f"/api/v1/dags/{dag_id}/datasets/queuedEvent/{dataset_uri}", - environ_overrides={"REMOTE_USER": "test_queued_event"}, - ) - - assert response.status_code == 200 - assert response.json == { - "created_at": self.default_time, - "uri": "s3://bucket/key", - "dag_id": "dag", - } - - def test_should_respond_404(self): - dag_id = "not_exists" - dataset_uri = "not_exists" - - response = self.client.get( - f"/api/v1/dags/{dag_id}/datasets/queuedEvent/{dataset_uri}", - environ_overrides={"REMOTE_USER": "test_queued_event"}, - ) - - assert response.status_code == 404 - assert { - "detail": "Queue event with dag_id: `not_exists` and asset uri: `not_exists` was not found", - "status": 404, - "title": "Queue event not found", - "type": EXCEPTIONS_LINK_MAP[404], - } == response.json - def test_should_raises_401_unauthenticated(self, session): dag_id = "dummy" dataset_uri = "dummy" @@ -826,47 +773,6 @@ def test_should_raise_403_forbidden(self, session): class TestDeleteDagDatasetQueuedEvent(TestDatasetEndpoint): - def test_delete_should_respond_204(self, session, create_dummy_dag): - dag, _ = create_dummy_dag() - dag_id = dag.dag_id - dataset_uri = "s3://bucket/key" - dataset_id = self._create_dataset(session).id - - adrq = AssetDagRunQueue(target_dag_id=dag_id, dataset_id=dataset_id) - session.add(adrq) - session.commit() - conn = session.query(AssetDagRunQueue).all() - assert len(conn) == 1 - - response = self.client.delete( - f"/api/v1/dags/{dag_id}/datasets/queuedEvent/{dataset_uri}", - environ_overrides={"REMOTE_USER": "test_queued_event"}, - ) - - assert response.status_code == 204 - conn = session.query(AssetDagRunQueue).all() - assert len(conn) == 0 - _check_last_log( - session, dag_id=dag_id, event="api.delete_dag_dataset_queued_event", execution_date=None - ) - - def test_should_respond_404(self): - dag_id = "not_exists" - dataset_uri = "not_exists" - - response = self.client.delete( - f"/api/v1/dags/{dag_id}/datasets/queuedEvent/{dataset_uri}", - environ_overrides={"REMOTE_USER": "test_queued_event"}, - ) - - assert response.status_code == 404 - assert { - "detail": "Queue event with dag_id: `not_exists` and asset uri: `not_exists` was not found", - "status": 404, - "title": "Queue event not found", - "type": EXCEPTIONS_LINK_MAP[404], - } == response.json - def test_should_raises_401_unauthenticated(self, session): dag_id = "dummy" dataset_uri = "dummy" @@ -884,46 +790,6 @@ def test_should_raise_403_forbidden(self, session): class TestGetDagDatasetQueuedEvents(TestQueuedEventEndpoint): - @pytest.mark.usefixtures("time_freezer") - def test_should_respond_200(self, session, create_dummy_dag): - dag, _ = create_dummy_dag() - dag_id = dag.dag_id - dataset_id = self._create_dataset(session).id - self._create_dataset_dag_run_queues(dag_id, dataset_id, session) - - response = self.client.get( - f"/api/v1/dags/{dag_id}/datasets/queuedEvent", - environ_overrides={"REMOTE_USER": "test_queued_event"}, - ) - - assert response.status_code == 200 - assert response.json == { - "queued_events": [ - { - "created_at": self.default_time, - "uri": "s3://bucket/key", - "dag_id": "dag", - } - ], - "total_entries": 1, - } - - def test_should_respond_404(self): - dag_id = "not_exists" - - response = self.client.get( - f"/api/v1/dags/{dag_id}/datasets/queuedEvent", - environ_overrides={"REMOTE_USER": "test_queued_event"}, - ) - - assert response.status_code == 404 - assert { - "detail": "Queue event with dag_id: `not_exists` was not found", - "status": 404, - "title": "Queue event not found", - "type": EXCEPTIONS_LINK_MAP[404], - } == response.json - def test_should_raises_401_unauthenticated(self): dag_id = "dummy" @@ -943,22 +809,6 @@ def test_should_raise_403_forbidden(self): class TestDeleteDagDatasetQueuedEvents(TestDatasetEndpoint): - def test_should_respond_404(self): - dag_id = "not_exists" - - response = self.client.delete( - f"/api/v1/dags/{dag_id}/datasets/queuedEvent", - environ_overrides={"REMOTE_USER": "test_queued_event"}, - ) - - assert response.status_code == 404 - assert { - "detail": "Queue event with dag_id: `not_exists` was not found", - "status": 404, - "title": "Queue event not found", - "type": EXCEPTIONS_LINK_MAP[404], - } == response.json - def test_should_raises_401_unauthenticated(self): dag_id = "dummy" @@ -978,47 +828,6 @@ def test_should_raise_403_forbidden(self): class TestGetDatasetQueuedEvents(TestQueuedEventEndpoint): - @pytest.mark.usefixtures("time_freezer") - def test_should_respond_200(self, session, create_dummy_dag): - dag, _ = create_dummy_dag() - dag_id = dag.dag_id - dataset_id = self._create_dataset(session).id - self._create_dataset_dag_run_queues(dag_id, dataset_id, session) - dataset_uri = "s3://bucket/key" - - response = self.client.get( - f"/api/v1/datasets/queuedEvent/{dataset_uri}", - environ_overrides={"REMOTE_USER": "test_queued_event"}, - ) - - assert response.status_code == 200 - assert response.json == { - "queued_events": [ - { - "created_at": self.default_time, - "uri": "s3://bucket/key", - "dag_id": "dag", - } - ], - "total_entries": 1, - } - - def test_should_respond_404(self): - dataset_uri = "not_exists" - - response = self.client.get( - f"/api/v1/datasets/queuedEvent/{dataset_uri}", - environ_overrides={"REMOTE_USER": "test_queued_event"}, - ) - - assert response.status_code == 404 - assert { - "detail": "Queue event with asset uri: `not_exists` was not found", - "status": 404, - "title": "Queue event not found", - "type": EXCEPTIONS_LINK_MAP[404], - } == response.json - def test_should_raises_401_unauthenticated(self): dataset_uri = "not_exists" @@ -1038,39 +847,6 @@ def test_should_raise_403_forbidden(self): class TestDeleteDatasetQueuedEvents(TestQueuedEventEndpoint): - def test_delete_should_respond_204(self, session, create_dummy_dag): - dag, _ = create_dummy_dag() - dag_id = dag.dag_id - dataset_id = self._create_dataset(session).id - self._create_dataset_dag_run_queues(dag_id, dataset_id, session) - dataset_uri = "s3://bucket/key" - - response = self.client.delete( - f"/api/v1/datasets/queuedEvent/{dataset_uri}", - environ_overrides={"REMOTE_USER": "test_queued_event"}, - ) - - assert response.status_code == 204 - conn = session.query(AssetDagRunQueue).all() - assert len(conn) == 0 - _check_last_log(session, dag_id=None, event="api.delete_dataset_queued_events", execution_date=None) - - def test_should_respond_404(self): - dataset_uri = "not_exists" - - response = self.client.delete( - f"/api/v1/datasets/queuedEvent/{dataset_uri}", - environ_overrides={"REMOTE_USER": "test_queued_event"}, - ) - - assert response.status_code == 404 - assert { - "detail": "Queue event with asset uri: `not_exists` was not found", - "status": 404, - "title": "Queue event not found", - "type": EXCEPTIONS_LINK_MAP[404], - } == response.json - def test_should_raises_401_unauthenticated(self): dataset_uri = "not_exists" diff --git a/tests/api_connexion/endpoints/test_event_log_endpoint.py b/tests/api_connexion/endpoints/test_event_log_endpoint.py index 0fdef1a3af2b6..e5ca3d301765a 100644 --- a/tests/api_connexion/endpoints/test_event_log_endpoint.py +++ b/tests/api_connexion/endpoints/test_event_log_endpoint.py @@ -20,7 +20,6 @@ from airflow.api_connexion.exceptions import EXCEPTIONS_LINK_MAP from airflow.models import Log -from airflow.security import permissions from airflow.utils import timezone from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_user from tests.test_utils.config import conf_vars @@ -33,32 +32,16 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type:ignore + app, username="test", - role_name="Test", - permissions=[(permissions.ACTION_CAN_READ, permissions.RESOURCE_AUDIT_LOG)], # type: ignore + role_name="admin", ) - create_user( - app, # type:ignore - username="test_granular", - role_name="TestGranular", - permissions=[(permissions.ACTION_CAN_READ, permissions.RESOURCE_AUDIT_LOG)], # type: ignore - ) - app.appbuilder.sm.sync_perm_for_dag( # type: ignore - "TEST_DAG_ID_1", - access_control={"TestGranular": [permissions.ACTION_CAN_READ]}, - ) - app.appbuilder.sm.sync_perm_for_dag( # type: ignore - "TEST_DAG_ID_2", - access_control={"TestGranular": [permissions.ACTION_CAN_READ]}, - ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_granular") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") @pytest.fixture @@ -274,33 +257,6 @@ def test_should_raises_401_unauthenticated(self, log_model): assert_401(response) - def test_should_filter_eventlogs_by_allowed_attributes(self, create_log_model, session): - eventlog1 = create_log_model( - event="TEST_EVENT_1", - dag_id="TEST_DAG_ID_1", - task_id="TEST_TASK_ID_1", - owner="TEST_OWNER_1", - when=self.default_time, - ) - eventlog2 = create_log_model( - event="TEST_EVENT_2", - dag_id="TEST_DAG_ID_2", - task_id="TEST_TASK_ID_2", - owner="TEST_OWNER_2", - when=self.default_time_2, - ) - session.add_all([eventlog1, eventlog2]) - session.commit() - for attr in ["dag_id", "task_id", "owner", "event"]: - attr_value = f"TEST_{attr}_1".upper() - response = self.client.get( - f"/api/v1/eventLogs?{attr}={attr_value}", environ_overrides={"REMOTE_USER": "test_granular"} - ) - assert response.status_code == 200 - assert response.json["total_entries"] == 1 - assert len(response.json["event_logs"]) == 1 - assert response.json["event_logs"][0][attr] == attr_value - def test_should_filter_eventlogs_by_when(self, create_log_model, session): eventlog1 = create_log_model(event="TEST_EVENT_1", when=self.default_time) eventlog2 = create_log_model(event="TEST_EVENT_2", when=self.default_time_2) @@ -339,32 +295,6 @@ def test_should_filter_eventlogs_by_run_id(self, create_log_model, session): assert {eventlog["event"] for eventlog in response.json["event_logs"]} == expected_eventlogs assert all({eventlog["run_id"] == run_id for eventlog in response.json["event_logs"]}) - def test_should_filter_eventlogs_by_included_events(self, create_log_model): - for event in ["TEST_EVENT_1", "TEST_EVENT_2", "cli_scheduler"]: - create_log_model(event=event, when=self.default_time) - response = self.client.get( - "/api/v1/eventLogs?included_events=TEST_EVENT_1,TEST_EVENT_2", - environ_overrides={"REMOTE_USER": "test_granular"}, - ) - assert response.status_code == 200 - response_data = response.json - assert len(response_data["event_logs"]) == 2 - assert response_data["total_entries"] == 2 - assert {"TEST_EVENT_1", "TEST_EVENT_2"} == {x["event"] for x in response_data["event_logs"]} - - def test_should_filter_eventlogs_by_excluded_events(self, create_log_model): - for event in ["TEST_EVENT_1", "TEST_EVENT_2", "cli_scheduler"]: - create_log_model(event=event, when=self.default_time) - response = self.client.get( - "/api/v1/eventLogs?excluded_events=TEST_EVENT_1,TEST_EVENT_2", - environ_overrides={"REMOTE_USER": "test_granular"}, - ) - assert response.status_code == 200 - response_data = response.json - assert len(response_data["event_logs"]) == 1 - assert response_data["total_entries"] == 1 - assert {"cli_scheduler"} == {x["event"] for x in response_data["event_logs"]} - class TestGetEventLogPagination(TestEventLogEndpoint): @pytest.mark.parametrize( diff --git a/tests/api_connexion/endpoints/test_extra_link_endpoint.py b/tests/api_connexion/endpoints/test_extra_link_endpoint.py index 1e9226ede9847..2c3eacdc91dc0 100644 --- a/tests/api_connexion/endpoints/test_extra_link_endpoint.py +++ b/tests/api_connexion/endpoints/test_extra_link_endpoint.py @@ -26,7 +26,6 @@ from airflow.models.dagbag import DagBag from airflow.models.xcom import XCom from airflow.plugins_manager import AirflowPlugin -from airflow.security import permissions from airflow.timetables.base import DataInterval from airflow.utils import timezone from airflow.utils.state import DagRunState @@ -48,21 +47,16 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type: ignore + app, username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_TASK_INSTANCE), - ], + role_name="admin", ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestGetExtraLinks: @@ -78,8 +72,8 @@ def setup_attrs(self, configured_app, session) -> None: self.dag = self._create_dag() self.app.dag_bag = DagBag(os.devnull, include_examples=False) - self.app.dag_bag.dags = {self.dag.dag_id: self.dag} # type: ignore - self.app.dag_bag.sync_to_db() # type: ignore + self.app.dag_bag.dags = {self.dag.dag_id: self.dag} + self.app.dag_bag.sync_to_db() triggered_by_kwargs = {"triggered_by": DagRunTriggeredByType.TEST} if AIRFLOW_V_3_0_PLUS else {} self.dag.create_dagrun( diff --git a/tests/api_connexion/endpoints/test_import_error_endpoint.py b/tests/api_connexion/endpoints/test_import_error_endpoint.py index 635e159bb292c..af2b83ebb1eed 100644 --- a/tests/api_connexion/endpoints/test_import_error_endpoint.py +++ b/tests/api_connexion/endpoints/test_import_error_endpoint.py @@ -21,15 +21,12 @@ import pytest from airflow.api_connexion.exceptions import EXCEPTIONS_LINK_MAP -from airflow.models.dag import DagModel -from airflow.security import permissions from airflow.utils import timezone from airflow.utils.session import provide_session from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_user from tests.test_utils.compat import ParseImportError from tests.test_utils.config import conf_vars from tests.test_utils.db import clear_db_dags, clear_db_import_errors -from tests.test_utils.permissions import _resource_name pytestmark = [pytest.mark.db_test, pytest.mark.skip_if_database_isolation_mode] @@ -40,42 +37,16 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type:ignore + app, username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_IMPORT_ERROR), - ], # type: ignore - ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore - create_user( - app, # type:ignore - username="test_single_dag", - role_name="TestSingleDAG", - permissions=[(permissions.ACTION_CAN_READ, permissions.RESOURCE_IMPORT_ERROR)], # type: ignore - ) - # For some reason, DAG level permissions are not synced when in the above list of perms, - # so do it manually here: - app.appbuilder.sm.bulk_sync_roles( - [ - { - "role": "TestSingleDAG", - "perms": [ - ( - permissions.ACTION_CAN_READ, - _resource_name(TEST_DAG_IDS[0], permissions.RESOURCE_DAG), - ) - ], - } - ] + role_name="admin", ) + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore - delete_user(app, username="test_single_dag") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestBaseImportError: @@ -152,72 +123,6 @@ def test_should_raise_403_forbidden(self): ) assert response.status_code == 403 - def test_should_raise_403_forbidden_without_dag_read(self, session): - import_error = ParseImportError( - filename="Lorem_ipsum.py", - stacktrace="Lorem ipsum", - timestamp=timezone.parse(self.timestamp, timezone="UTC"), - ) - session.add(import_error) - session.commit() - - response = self.client.get( - f"/api/v1/importErrors/{import_error.id}", environ_overrides={"REMOTE_USER": "test_single_dag"} - ) - - assert response.status_code == 403 - - def test_should_return_200_with_single_dag_read(self, session): - dag_model = DagModel(dag_id=TEST_DAG_IDS[0], fileloc="Lorem_ipsum.py") - session.add(dag_model) - import_error = ParseImportError( - filename="Lorem_ipsum.py", - stacktrace="Lorem ipsum", - timestamp=timezone.parse(self.timestamp, timezone="UTC"), - ) - session.add(import_error) - session.commit() - - response = self.client.get( - f"/api/v1/importErrors/{import_error.id}", environ_overrides={"REMOTE_USER": "test_single_dag"} - ) - - assert response.status_code == 200 - response_data = response.json - response_data["import_error_id"] = 1 - assert { - "filename": "Lorem_ipsum.py", - "import_error_id": 1, - "stack_trace": "Lorem ipsum", - "timestamp": "2020-06-10T12:00:00+00:00", - } == response_data - - def test_should_return_200_redacted_with_single_dag_read_in_dagfile(self, session): - for dag_id in TEST_DAG_IDS: - dag_model = DagModel(dag_id=dag_id, fileloc="Lorem_ipsum.py") - session.add(dag_model) - import_error = ParseImportError( - filename="Lorem_ipsum.py", - stacktrace="Lorem ipsum", - timestamp=timezone.parse(self.timestamp, timezone="UTC"), - ) - session.add(import_error) - session.commit() - - response = self.client.get( - f"/api/v1/importErrors/{import_error.id}", environ_overrides={"REMOTE_USER": "test_single_dag"} - ) - - assert response.status_code == 200 - response_data = response.json - response_data["import_error_id"] = 1 - assert { - "filename": "Lorem_ipsum.py", - "import_error_id": 1, - "stack_trace": "REDACTED - you do not have read permission on all DAGs in the file", - "timestamp": "2020-06-10T12:00:00+00:00", - } == response_data - class TestGetImportErrorsEndpoint(TestBaseImportError): def test_get_import_errors(self, session): @@ -328,71 +233,6 @@ def test_should_raises_401_unauthenticated(self, session): assert_401(response) - def test_get_import_errors_single_dag(self, session): - for dag_id in TEST_DAG_IDS: - fake_filename = f"/tmp/{dag_id}.py" - dag_model = DagModel(dag_id=dag_id, fileloc=fake_filename) - session.add(dag_model) - importerror = ParseImportError( - filename=fake_filename, - stacktrace="Lorem ipsum", - timestamp=timezone.parse(self.timestamp, timezone="UTC"), - ) - session.add(importerror) - session.commit() - - response = self.client.get( - "/api/v1/importErrors", environ_overrides={"REMOTE_USER": "test_single_dag"} - ) - - assert response.status_code == 200 - response_data = response.json - self._normalize_import_errors(response_data["import_errors"]) - assert { - "import_errors": [ - { - "filename": "/tmp/test_dag.py", - "import_error_id": 1, - "stack_trace": "Lorem ipsum", - "timestamp": "2020-06-10T12:00:00+00:00", - }, - ], - "total_entries": 1, - } == response_data - - def test_get_import_errors_single_dag_in_dagfile(self, session): - for dag_id in TEST_DAG_IDS: - fake_filename = "/tmp/all_in_one.py" - dag_model = DagModel(dag_id=dag_id, fileloc=fake_filename) - session.add(dag_model) - - importerror = ParseImportError( - filename="/tmp/all_in_one.py", - stacktrace="Lorem ipsum", - timestamp=timezone.parse(self.timestamp, timezone="UTC"), - ) - session.add(importerror) - session.commit() - - response = self.client.get( - "/api/v1/importErrors", environ_overrides={"REMOTE_USER": "test_single_dag"} - ) - - assert response.status_code == 200 - response_data = response.json - self._normalize_import_errors(response_data["import_errors"]) - assert { - "import_errors": [ - { - "filename": "/tmp/all_in_one.py", - "import_error_id": 1, - "stack_trace": "REDACTED - you do not have read permission on all DAGs in the file", - "timestamp": "2020-06-10T12:00:00+00:00", - }, - ], - "total_entries": 1, - } == response_data - class TestGetImportErrorsEndpointPagination(TestBaseImportError): @pytest.mark.parametrize( diff --git a/tests/api_connexion/endpoints/test_log_endpoint.py b/tests/api_connexion/endpoints/test_log_endpoint.py index 420d2dd65f89c..2b112e3221843 100644 --- a/tests/api_connexion/endpoints/test_log_endpoint.py +++ b/tests/api_connexion/endpoints/test_log_endpoint.py @@ -30,7 +30,6 @@ from airflow.decorators import task from airflow.models.dag import DAG from airflow.operators.empty import EmptyOperator -from airflow.security import permissions from airflow.utils import timezone from airflow.utils.types import DagRunType from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_user @@ -46,13 +45,9 @@ def configured_app(minimal_app_for_api): create_user( app, username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_TASK_LOG), - ], + role_name="admin", ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") + create_user(app, username="test_no_permissions", role_name=None) yield app diff --git a/tests/api_connexion/endpoints/test_mapped_task_instance_endpoint.py b/tests/api_connexion/endpoints/test_mapped_task_instance_endpoint.py index 72cdccdee68df..fc53b8952f4aa 100644 --- a/tests/api_connexion/endpoints/test_mapped_task_instance_endpoint.py +++ b/tests/api_connexion/endpoints/test_mapped_task_instance_endpoint.py @@ -28,12 +28,11 @@ from airflow.models.baseoperator import BaseOperator from airflow.models.dagbag import DagBag from airflow.models.taskmap import TaskMap -from airflow.security import permissions from airflow.utils.platform import getuser from airflow.utils.session import provide_session from airflow.utils.state import State, TaskInstanceState from airflow.utils.timezone import datetime -from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_roles, delete_user +from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_user from tests.test_utils.db import clear_db_runs, clear_db_sla_miss, clear_rendered_ti_fields from tests.test_utils.mock_operators import MockOperator @@ -50,24 +49,16 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type: ignore + app, username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_TASK_INSTANCE), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_TASK_INSTANCE), - ], + role_name="admin", ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore - delete_roles(app) + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestMappedTaskInstanceEndpoint: @@ -133,8 +124,8 @@ def create_dag_runs_with_mapped_tasks(self, dag_maker, session, dags=None): session.add(ti) self.app.dag_bag = DagBag(os.devnull, include_examples=False) - self.app.dag_bag.dags = {dag_id: dag_maker.dag} # type: ignore - self.app.dag_bag.sync_to_db() # type: ignore + self.app.dag_bag.dags = {dag_id: dag_maker.dag} + self.app.dag_bag.sync_to_db() session.flush() mapped.expand_mapped_task(dr.run_id, session=session) diff --git a/tests/api_connexion/endpoints/test_plugin_endpoint.py b/tests/api_connexion/endpoints/test_plugin_endpoint.py index edf925cf0fa73..0cd630375a282 100644 --- a/tests/api_connexion/endpoints/test_plugin_endpoint.py +++ b/tests/api_connexion/endpoints/test_plugin_endpoint.py @@ -24,7 +24,6 @@ from airflow.hooks.base import BaseHook from airflow.plugins_manager import AirflowPlugin -from airflow.security import permissions from airflow.ti_deps.deps.base_ti_dep import BaseTIDep from airflow.timetables.base import Timetable from airflow.utils.module_loading import qualname @@ -105,17 +104,16 @@ class MockPlugin(AirflowPlugin): def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type: ignore + app, username="test", - role_name="Test", - permissions=[(permissions.ACTION_CAN_READ, permissions.RESOURCE_PLUGIN)], + role_name="admin", ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestPluginsEndpoint: diff --git a/tests/api_connexion/endpoints/test_pool_endpoint.py b/tests/api_connexion/endpoints/test_pool_endpoint.py index 87439a5811945..2cc095d077aa9 100644 --- a/tests/api_connexion/endpoints/test_pool_endpoint.py +++ b/tests/api_connexion/endpoints/test_pool_endpoint.py @@ -20,7 +20,6 @@ from airflow.api_connexion.exceptions import EXCEPTIONS_LINK_MAP from airflow.models.pool import Pool -from airflow.security import permissions from airflow.utils.session import provide_session from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_user from tests.test_utils.config import conf_vars @@ -35,22 +34,16 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type: ignore + app, username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_POOL), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_POOL), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_POOL), - (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_POOL), - ], + role_name="admin", ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestBasePoolEndpoints: diff --git a/tests/api_connexion/endpoints/test_provider_endpoint.py b/tests/api_connexion/endpoints/test_provider_endpoint.py index 16e5989cc56db..b4cf8f10a92ae 100644 --- a/tests/api_connexion/endpoints/test_provider_endpoint.py +++ b/tests/api_connexion/endpoints/test_provider_endpoint.py @@ -21,7 +21,6 @@ import pytest from airflow.providers_manager import ProviderInfo -from airflow.security import permissions from tests.test_utils.api_connexion_utils import create_user, delete_user pytestmark = [pytest.mark.db_test, pytest.mark.skip_if_database_isolation_mode] @@ -54,17 +53,16 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type: ignore + app, username="test", - role_name="Test", - permissions=[(permissions.ACTION_CAN_READ, permissions.RESOURCE_PROVIDER)], + role_name="admin", ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestBaseProviderEndpoint: diff --git a/tests/api_connexion/endpoints/test_task_endpoint.py b/tests/api_connexion/endpoints/test_task_endpoint.py index d0a4fb903c8b8..b2e068bd507fe 100644 --- a/tests/api_connexion/endpoints/test_task_endpoint.py +++ b/tests/api_connexion/endpoints/test_task_endpoint.py @@ -27,7 +27,6 @@ from airflow.models.expandinput import EXPAND_INPUT_EMPTY from airflow.models.serialized_dag import SerializedDagModel from airflow.operators.empty import EmptyOperator -from airflow.security import permissions from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_user from tests.test_utils.db import clear_db_dags, clear_db_runs, clear_db_serialized_dags @@ -38,21 +37,16 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type: ignore + app, username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_TASK_INSTANCE), - ], + role_name="admin", ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestTaskEndpoint: diff --git a/tests/api_connexion/endpoints/test_task_instance_endpoint.py b/tests/api_connexion/endpoints/test_task_instance_endpoint.py index 25ded6c814b72..b5b3163e988d0 100644 --- a/tests/api_connexion/endpoints/test_task_instance_endpoint.py +++ b/tests/api_connexion/endpoints/test_task_instance_endpoint.py @@ -25,19 +25,17 @@ from sqlalchemy import select from sqlalchemy.orm import contains_eager -from airflow.api_connexion.exceptions import EXCEPTIONS_LINK_MAP from airflow.jobs.job import Job from airflow.jobs.triggerer_job_runner import TriggererJobRunner from airflow.models import DagRun, SlaMiss, TaskInstance, Trigger from airflow.models.renderedtifields import RenderedTaskInstanceFields as RTIF from airflow.models.taskinstancehistory import TaskInstanceHistory -from airflow.security import permissions from airflow.utils.platform import getuser from airflow.utils.session import provide_session from airflow.utils.state import State from airflow.utils.timezone import datetime from airflow.utils.types import DagRunType -from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_roles, delete_user +from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_user from tests.test_utils.db import clear_db_runs, clear_db_sla_miss, clear_rendered_ti_fields from tests.test_utils.www import _check_last_log @@ -55,69 +53,16 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type: ignore + app, username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_TASK_INSTANCE), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_TASK_INSTANCE), - ], - ) - create_user( - app, # type: ignore - username="test_dag_read_only", - role_name="TestDagReadOnly", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_TASK_INSTANCE), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_TASK_INSTANCE), - ], - ) - create_user( - app, # type: ignore - username="test_task_read_only", - role_name="TestTaskReadOnly", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_TASK_INSTANCE), - ], - ) - create_user( - app, # type: ignore - username="test_read_only_one_dag", - role_name="TestReadOnlyOneDag", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_TASK_INSTANCE), - ], - ) - # For some reason, "DAG:example_python_operator" is not synced when in the above list of perms, - # so do it manually here: - app.appbuilder.sm.bulk_sync_roles( - [ - { - "role": "TestReadOnlyOneDag", - "perms": [(permissions.ACTION_CAN_READ, "DAG:example_python_operator")], - } - ] + role_name="admin", ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_dag_read_only") # type: ignore - delete_user(app, username="test_task_read_only") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore - delete_user(app, username="test_read_only_one_dag") # type: ignore - delete_roles(app) + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestTaskInstanceEndpoint: @@ -219,9 +164,8 @@ def setup_method(self): def teardown_method(self): clear_db_runs() - @pytest.mark.parametrize("username", ["test", "test_dag_read_only", "test_task_read_only"]) @provide_session - def test_should_respond_200(self, username, session): + def test_should_respond_200(self, session): self.create_task_instances(session) # Update ti and set operator to None to # test that operator field is nullable. @@ -232,7 +176,7 @@ def test_should_respond_200(self, username, session): session.commit() response = self.client.get( "/api/v1/dags/example_python_operator/dagRuns/TEST_DAG_RUN_ID/taskInstances/print_the_context", - environ_overrides={"REMOTE_USER": username}, + environ_overrides={"REMOTE_USER": "test"}, ) assert response.status_code == 200 assert response.json == { @@ -723,36 +667,11 @@ def test_should_respond_200(self, task_instances, update_extras, url, expected_t assert response.json["total_entries"] == expected_ti assert len(response.json["task_instances"]) == expected_ti - @pytest.mark.parametrize( - "task_instances, user, expected_ti", - [ - pytest.param( - { - "example_python_operator": 2, - "example_skip_dag": 1, - }, - "test_read_only_one_dag", - 2, - ), - pytest.param( - { - "example_python_operator": 1, - "example_skip_dag": 2, - }, - "test_read_only_one_dag", - 1, - ), - pytest.param( - { - "example_python_operator": 1, - "example_skip_dag": 2, - }, - "test", - 3, - ), - ], - ) - def test_return_TI_only_from_readable_dags(self, task_instances, user, expected_ti, session): + def test_return_TI_only_from_readable_dags(self, session): + task_instances = { + "example_python_operator": 1, + "example_skip_dag": 2, + } for dag_id in task_instances: self.create_task_instances( session, @@ -763,11 +682,11 @@ def test_return_TI_only_from_readable_dags(self, task_instances, user, expected_ dag_id=dag_id, ) response = self.client.get( - "/api/v1/dags/~/dagRuns/~/taskInstances", environ_overrides={"REMOTE_USER": user} + "/api/v1/dags/~/dagRuns/~/taskInstances", environ_overrides={"REMOTE_USER": "test"} ) assert response.status_code == 200 - assert response.json["total_entries"] == expected_ti - assert len(response.json["task_instances"]) == expected_ti + assert response.json["total_entries"] == 3 + assert len(response.json["task_instances"]) == 3 def test_should_respond_200_for_dag_id_filter(self, session): self.create_task_instances(session) @@ -898,44 +817,6 @@ class TestGetTaskInstancesBatch(TestTaskInstanceEndpoint): "test", id="test executor filter", ), - pytest.param( - [ - {"pool": "test_pool_1"}, - {"pool": "test_pool_2"}, - {"pool": "test_pool_3"}, - ], - True, - {"pool": ["test_pool_1", "test_pool_2"]}, - 2, - "test_dag_read_only", - id="test pool filter", - ), - pytest.param( - [ - {"state": State.RUNNING}, - {"state": State.QUEUED}, - {"state": State.SUCCESS}, - {"state": State.NONE}, - ], - False, - {"state": ["running", "queued", "none"]}, - 3, - "test_task_read_only", - id="test state filter", - ), - pytest.param( - [ - {"state": State.NONE}, - {"state": State.NONE}, - {"state": State.NONE}, - {"state": State.NONE}, - ], - False, - {}, - 4, - "test_task_read_only", - id="test dag with null states", - ), pytest.param( [ {"duration": 100}, @@ -948,36 +829,6 @@ class TestGetTaskInstancesBatch(TestTaskInstanceEndpoint): "test", id="test duration filter", ), - pytest.param( - [ - {"end_date": DEFAULT_DATETIME_1}, - {"end_date": DEFAULT_DATETIME_1 + dt.timedelta(days=1)}, - {"end_date": DEFAULT_DATETIME_1 + dt.timedelta(days=2)}, - ], - True, - { - "end_date_gte": DEFAULT_DATETIME_STR_1, - "end_date_lte": DEFAULT_DATETIME_STR_2, - }, - 2, - "test_task_read_only", - id="test end date filter", - ), - pytest.param( - [ - {"start_date": DEFAULT_DATETIME_1}, - {"start_date": DEFAULT_DATETIME_1 + dt.timedelta(days=1)}, - {"start_date": DEFAULT_DATETIME_1 + dt.timedelta(days=2)}, - ], - True, - { - "start_date_gte": DEFAULT_DATETIME_STR_1, - "start_date_lte": DEFAULT_DATETIME_STR_2, - }, - 2, - "test_dag_read_only", - id="test start date filter", - ), pytest.param( [ {"execution_date": DEFAULT_DATETIME_1}, @@ -1162,24 +1013,6 @@ def test_should_raise_403_forbidden(self): ) assert response.status_code == 403 - def test_returns_403_forbidden_when_user_has_access_to_only_some_dags(self, session): - self.create_task_instances(session=session) - self.create_task_instances(session=session, dag_id="example_skip_dag") - payload = {"dag_ids": ["example_python_operator", "example_skip_dag"]} - - response = self.client.post( - "/api/v1/dags/~/dagRuns/~/taskInstances/list", - environ_overrides={"REMOTE_USER": "test_read_only_one_dag"}, - json=payload, - ) - assert response.status_code == 403 - assert response.json == { - "detail": "User not allowed to access some of these DAGs: ['example_python_operator', 'example_skip_dag']", - "status": 403, - "title": "Forbidden", - "type": EXCEPTIONS_LINK_MAP[403], - } - def test_should_raise_400_for_no_json(self): response = self.client.post( "/api/v1/dags/~/dagRuns/~/taskInstances/list", @@ -1794,11 +1627,10 @@ def test_should_raises_401_unauthenticated(self): ) assert_401(response) - @pytest.mark.parametrize("username", ["test_no_permissions", "test_dag_read_only", "test_task_read_only"]) - def test_should_raise_403_forbidden(self, username: str): + def test_should_raise_403_forbidden(self): response = self.client.post( "/api/v1/dags/example_python_operator/clearTaskInstances", - environ_overrides={"REMOTE_USER": username}, + environ_overrides={"REMOTE_USER": "test_no_permissions"}, json={ "dry_run": False, "reset_dag_runs": True, @@ -2043,11 +1875,10 @@ def test_should_raises_401_unauthenticated(self): ) assert_401(response) - @pytest.mark.parametrize("username", ["test_no_permissions", "test_dag_read_only", "test_task_read_only"]) - def test_should_raise_403_forbidden(self, username): + def test_should_raise_403_forbidden(self): response = self.client.post( "/api/v1/dags/example_python_operator/updateTaskInstancesState", - environ_overrides={"REMOTE_USER": username}, + environ_overrides={"REMOTE_USER": "test_no_permissions"}, json={ "dry_run": True, "task_id": "print_the_context", @@ -2386,11 +2217,10 @@ def test_should_raises_401_unauthenticated(self): ) assert_401(response) - @pytest.mark.parametrize("username", ["test_no_permissions", "test_dag_read_only", "test_task_read_only"]) - def test_should_raise_403_forbidden(self, username): + def test_should_raise_403_forbidden(self): response = self.client.patch( self.ENDPOINT_URL, - environ_overrides={"REMOTE_USER": username}, + environ_overrides={"REMOTE_USER": "test_no_permissions"}, json={ "dry_run": True, "new_state": "failed", @@ -2748,14 +2578,13 @@ def setup_method(self): def teardown_method(self): clear_db_runs() - @pytest.mark.parametrize("username", ["test", "test_dag_read_only", "test_task_read_only"]) @provide_session - def test_should_respond_200(self, username, session): + def test_should_respond_200(self, session): self.create_task_instances(session, task_instances=[{"state": State.SUCCESS}], with_ti_history=True) response = self.client.get( "/api/v1/dags/example_python_operator/dagRuns/TEST_DAG_RUN_ID/taskInstances/print_the_context/tries/1", - environ_overrides={"REMOTE_USER": username}, + environ_overrides={"REMOTE_USER": "test"}, ) assert response.status_code == 200 assert response.json == { diff --git a/tests/api_connexion/endpoints/test_variable_endpoint.py b/tests/api_connexion/endpoints/test_variable_endpoint.py index 81405df08b045..aa5f7c99674f8 100644 --- a/tests/api_connexion/endpoints/test_variable_endpoint.py +++ b/tests/api_connexion/endpoints/test_variable_endpoint.py @@ -22,7 +22,6 @@ from airflow.api_connexion.exceptions import EXCEPTIONS_LINK_MAP from airflow.models import Variable -from airflow.security import permissions from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_user from tests.test_utils.config import conf_vars from tests.test_utils.db import clear_db_variables @@ -36,40 +35,16 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type: ignore + app, username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_VARIABLE), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_VARIABLE), - (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_VARIABLE), - (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_VARIABLE), - ], - ) - create_user( - app, # type: ignore - username="test_read_only", - role_name="TestReadOnly", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_VARIABLE), - ], - ) - create_user( - app, # type: ignore - username="test_delete_only", - role_name="TestDeleteOnly", - permissions=[ - (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_VARIABLE), - ], + role_name="admin", ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_read_only") # type: ignore - delete_user(app, username="test_delete_only") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestVariableEndpoint: @@ -131,8 +106,6 @@ class TestGetVariable(TestVariableEndpoint): "user, expected_status_code", [ ("test", 200), - ("test_read_only", 200), - ("test_delete_only", 403), ("test_no_permissions", 403), ], ) diff --git a/tests/api_connexion/endpoints/test_xcom_endpoint.py b/tests/api_connexion/endpoints/test_xcom_endpoint.py index 7a51714c5b299..809e537f9f88d 100644 --- a/tests/api_connexion/endpoints/test_xcom_endpoint.py +++ b/tests/api_connexion/endpoints/test_xcom_endpoint.py @@ -26,7 +26,6 @@ from airflow.models.taskinstance import TaskInstance from airflow.models.xcom import BaseXCom, XCom, resolve_xcom_backend from airflow.operators.empty import EmptyOperator -from airflow.security import permissions from airflow.utils.dates import parse_execution_date from airflow.utils.session import create_session from airflow.utils.timezone import utcnow @@ -52,32 +51,16 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type: ignore + app, username="test", - role_name="Test", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), - (permissions.ACTION_CAN_READ, permissions.RESOURCE_XCOM), - ], - ) - create_user( - app, # type: ignore - username="test_granular_permissions", - role_name="TestGranularDag", - permissions=[ - (permissions.ACTION_CAN_READ, permissions.RESOURCE_XCOM), - ], - ) - app.appbuilder.sm.sync_perm_for_dag( # type: ignore - "test-dag-id-1", - access_control={"TestGranularDag": [permissions.ACTION_CAN_EDIT, permissions.ACTION_CAN_READ]}, + role_name="admin", ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name=None) yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") def _compare_xcom_collections(collection1: dict, collection_2: dict): @@ -435,53 +418,6 @@ def test_should_respond_200_with_tilde_and_access_to_all_dags(self): }, ) - def test_should_respond_200_with_tilde_and_granular_dag_access(self): - dag_id_1 = "test-dag-id-1" - task_id_1 = "test-task-id-1" - execution_date = "2005-04-02T00:00:00+00:00" - execution_date_parsed = parse_execution_date(execution_date) - dag_run_id_1 = DagRun.generate_run_id(DagRunType.MANUAL, execution_date_parsed) - self._create_xcom_entries(dag_id_1, dag_run_id_1, execution_date_parsed, task_id_1) - - dag_id_2 = "test-dag-id-2" - task_id_2 = "test-task-id-2" - run_id_2 = DagRun.generate_run_id(DagRunType.MANUAL, execution_date_parsed) - self._create_xcom_entries(dag_id_2, run_id_2, execution_date_parsed, task_id_2) - self._create_invalid_xcom_entries(execution_date_parsed) - response = self.client.get( - "/api/v1/dags/~/dagRuns/~/taskInstances/~/xcomEntries", - environ_overrides={"REMOTE_USER": "test_granular_permissions"}, - ) - - assert 200 == response.status_code - response_data = response.json - for xcom_entry in response_data["xcom_entries"]: - xcom_entry["timestamp"] = "TIMESTAMP" - _compare_xcom_collections( - response_data, - { - "xcom_entries": [ - { - "dag_id": dag_id_1, - "execution_date": execution_date, - "key": "test-xcom-key-1", - "task_id": task_id_1, - "timestamp": "TIMESTAMP", - "map_index": -1, - }, - { - "dag_id": dag_id_1, - "execution_date": execution_date, - "key": "test-xcom-key-2", - "task_id": task_id_1, - "timestamp": "TIMESTAMP", - "map_index": -1, - }, - ], - "total_entries": 2, - }, - ) - def test_should_respond_200_with_map_index(self): dag_id = "test-dag-id" task_id = "test-task-id" diff --git a/tests/api_connexion/test_auth.py b/tests/api_connexion/test_auth.py index 7d1dcc088273c..54e5632ad84d1 100644 --- a/tests/api_connexion/test_auth.py +++ b/tests/api_connexion/test_auth.py @@ -16,15 +16,15 @@ # under the License. from __future__ import annotations -from base64 import b64encode +from unittest.mock import patch import pytest -from flask_login import current_user +from airflow.auth.managers.simple.simple_auth_manager import SimpleAuthManager +from airflow.auth.managers.simple.user import SimpleAuthManagerUser from tests.test_utils.api_connexion_utils import assert_401 from tests.test_utils.config import conf_vars from tests.test_utils.db import clear_db_pools -from tests.test_utils.www import client_with_login pytestmark = [pytest.mark.db_test, pytest.mark.skip_if_database_isolation_mode] @@ -34,101 +34,6 @@ class BaseTestAuth: def set_attrs(self, minimal_app_for_api): self.app = minimal_app_for_api - sm = self.app.appbuilder.sm - tester = sm.find_user(username="test") - if not tester: - role_admin = sm.find_role("Admin") - sm.add_user( - username="test", - first_name="test", - last_name="test", - email="test@fab.org", - role=role_admin, - password="test", - ) - - -class TestBasicAuth(BaseTestAuth): - @pytest.fixture(autouse=True, scope="class") - def with_basic_auth_backend(self, minimal_app_for_api): - from airflow.www.extensions.init_security import init_api_auth - - old_auth = getattr(minimal_app_for_api, "api_auth") - - try: - with conf_vars( - {("api", "auth_backends"): "airflow.providers.fab.auth_manager.api.auth.backend.basic_auth"} - ): - init_api_auth(minimal_app_for_api) - yield - finally: - setattr(minimal_app_for_api, "api_auth", old_auth) - - def test_success(self): - token = "Basic " + b64encode(b"test:test").decode() - clear_db_pools() - - with self.app.test_client() as test_client: - response = test_client.get("/api/v1/pools", headers={"Authorization": token}) - assert current_user.email == "test@fab.org" - - assert response.status_code == 200 - assert response.json == { - "pools": [ - { - "name": "default_pool", - "slots": 128, - "occupied_slots": 0, - "running_slots": 0, - "queued_slots": 0, - "scheduled_slots": 0, - "deferred_slots": 0, - "open_slots": 128, - "description": "Default pool", - "include_deferred": False, - }, - ], - "total_entries": 1, - } - - @pytest.mark.parametrize( - "token", - [ - "basic", - "basic ", - "bearer", - "test:test", - b64encode(b"test:test").decode(), - "bearer ", - "basic: ", - "basic 123", - ], - ) - def test_malformed_headers(self, token): - with self.app.test_client() as test_client: - response = test_client.get("/api/v1/pools", headers={"Authorization": token}) - assert response.status_code == 401 - assert response.headers["Content-Type"] == "application/problem+json" - assert response.headers["WWW-Authenticate"] == "Basic" - assert_401(response) - - @pytest.mark.parametrize( - "token", - [ - "basic " + b64encode(b"test").decode(), - "basic " + b64encode(b"test:").decode(), - "basic " + b64encode(b"test:123").decode(), - "basic " + b64encode(b"test test").decode(), - ], - ) - def test_invalid_auth_header(self, token): - with self.app.test_client() as test_client: - response = test_client.get("/api/v1/pools", headers={"Authorization": token}) - assert response.status_code == 401 - assert response.headers["Content-Type"] == "application/problem+json" - assert response.headers["WWW-Authenticate"] == "Basic" - assert_401(response) - class TestSessionAuth(BaseTestAuth): @pytest.fixture(autouse=True, scope="class") @@ -144,74 +49,37 @@ def with_session_backend(self, minimal_app_for_api): finally: setattr(minimal_app_for_api, "api_auth", old_auth) - def test_success(self): + @patch.object(SimpleAuthManager, "is_logged_in", return_value=True) + @patch.object( + SimpleAuthManager, "get_user", return_value=SimpleAuthManagerUser(username="test", role="admin") + ) + def test_success(self, *args): clear_db_pools() - admin_user = client_with_login(self.app, username="test", password="test") - response = admin_user.get("/api/v1/pools") - assert response.status_code == 200 - assert response.json == { - "pools": [ - { - "name": "default_pool", - "slots": 128, - "occupied_slots": 0, - "running_slots": 0, - "queued_slots": 0, - "scheduled_slots": 0, - "deferred_slots": 0, - "open_slots": 128, - "description": "Default pool", - "include_deferred": False, - }, - ], - "total_entries": 1, - } - - def test_failure(self): with self.app.test_client() as test_client: response = test_client.get("/api/v1/pools") - assert response.status_code == 401 - assert response.headers["Content-Type"] == "application/problem+json" - assert_401(response) - - -class TestSessionWithBasicAuthFallback(BaseTestAuth): - @pytest.fixture(autouse=True, scope="class") - def with_basic_auth_backend(self, minimal_app_for_api): - from airflow.www.extensions.init_security import init_api_auth - - old_auth = getattr(minimal_app_for_api, "api_auth") - - try: - with conf_vars( - { - ( - "api", - "auth_backends", - ): "airflow.api.auth.backend.session,airflow.providers.fab.auth_manager.api.auth.backend.basic_auth" - } - ): - init_api_auth(minimal_app_for_api) - yield - finally: - setattr(minimal_app_for_api, "api_auth", old_auth) - - def test_basic_auth_fallback(self): - token = "Basic " + b64encode(b"test:test").decode() - clear_db_pools() - - # request uses session - admin_user = client_with_login(self.app, username="test", password="test") - response = admin_user.get("/api/v1/pools") - assert response.status_code == 200 - - # request uses basic auth - with self.app.test_client() as test_client: - response = test_client.get("/api/v1/pools", headers={"Authorization": token}) assert response.status_code == 200 + assert response.json == { + "pools": [ + { + "name": "default_pool", + "slots": 128, + "occupied_slots": 0, + "running_slots": 0, + "queued_slots": 0, + "scheduled_slots": 0, + "deferred_slots": 0, + "open_slots": 128, + "description": "Default pool", + "include_deferred": False, + }, + ], + "total_entries": 1, + } - # request without session or basic auth header + def test_failure(self): with self.app.test_client() as test_client: response = test_client.get("/api/v1/pools") assert response.status_code == 401 + assert response.headers["Content-Type"] == "application/problem+json" + assert_401(response) diff --git a/tests/api_connexion/test_security.py b/tests/api_connexion/test_security.py index 13a5dd4e25af1..c6a112b1a1bb9 100644 --- a/tests/api_connexion/test_security.py +++ b/tests/api_connexion/test_security.py @@ -18,7 +18,6 @@ import pytest -from airflow.security import permissions from tests.test_utils.api_connexion_utils import create_user, delete_user pytestmark = [pytest.mark.db_test, pytest.mark.skip_if_database_isolation_mode] @@ -28,15 +27,14 @@ def configured_app(minimal_app_for_api): app = minimal_app_for_api create_user( - app, # type:ignore + app, username="test", - role_name="Test", - permissions=[(permissions.ACTION_CAN_READ, permissions.RESOURCE_CONFIG)], # type: ignore + role_name="admin", ) yield minimal_app_for_api - delete_user(app, username="test") # type: ignore + delete_user(app, username="test") class TestSession: diff --git a/tests/providers/fab/auth_manager/api_endpoints/api_connexion_utils.py b/tests/providers/fab/auth_manager/api_endpoints/api_connexion_utils.py new file mode 100644 index 0000000000000..61d923d5ff125 --- /dev/null +++ b/tests/providers/fab/auth_manager/api_endpoints/api_connexion_utils.py @@ -0,0 +1,116 @@ +# 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 + +from contextlib import contextmanager + +from tests.test_utils.compat import ignore_provider_compatibility_error + +with ignore_provider_compatibility_error("2.9.0+", __file__): + from airflow.providers.fab.auth_manager.security_manager.override import EXISTING_ROLES + + +@contextmanager +def create_test_client(app, user_name, role_name, permissions): + """ + Helper function to create a client with a temporary user which will be deleted once done + """ + client = app.test_client() + with create_user_scope(app, username=user_name, role_name=role_name, permissions=permissions) as _: + resp = client.post("/login/", data={"username": user_name, "password": user_name}) + assert resp.status_code == 302 + yield client + + +@contextmanager +def create_user_scope(app, username, **kwargs): + """ + Helper function designed to be used with pytest fixture mainly. + It will create a user and provide it for the fixture via YIELD (generator) + then will tidy up once test is complete + """ + test_user = create_user(app, username, **kwargs) + + try: + yield test_user + finally: + delete_user(app, username) + + +def create_user(app, username, role_name=None, email=None, permissions=None): + appbuilder = app.appbuilder + + # Removes user and role so each test has isolated test data. + delete_user(app, username) + role = None + if role_name: + delete_role(app, role_name) + role = create_role(app, role_name, permissions) + else: + role = [] + + return appbuilder.sm.add_user( + username=username, + first_name=username, + last_name=username, + email=email or f"{username}@example.org", + role=role, + password=username, + ) + + +def create_role(app, name, permissions=None): + appbuilder = app.appbuilder + role = appbuilder.sm.find_role(name) + if not role: + role = appbuilder.sm.add_role(name) + if not permissions: + permissions = [] + for permission in permissions: + perm_object = appbuilder.sm.get_permission(*permission) + appbuilder.sm.add_permission_to_role(role, perm_object) + return role + + +def set_user_single_role(app, user, role_name): + role = create_role(app, role_name) + if role not in user.roles: + user.roles = [role] + app.appbuilder.sm.update_user(user) + user._perms = None + + +def delete_role(app, name): + if name not in EXISTING_ROLES: + if app.appbuilder.sm.find_role(name): + app.appbuilder.sm.delete_role(name) + + +def delete_roles(app): + for role in app.appbuilder.sm.get_all_roles(): + delete_role(app, role.name) + + +def delete_user(app, username): + appbuilder = app.appbuilder + for user in appbuilder.sm.get_all_users(): + if user.username == username: + _ = [ + delete_role(app, role.name) for role in user.roles if role and role.name not in EXISTING_ROLES + ] + appbuilder.sm.del_register_user(user) + break diff --git a/tests/providers/fab/auth_manager/api_endpoints/remote_user_api_auth_backend.py b/tests/providers/fab/auth_manager/api_endpoints/remote_user_api_auth_backend.py new file mode 100644 index 0000000000000..b7714e5192e6a --- /dev/null +++ b/tests/providers/fab/auth_manager/api_endpoints/remote_user_api_auth_backend.py @@ -0,0 +1,81 @@ +# +# 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. +"""Default authentication backend - everything is allowed""" + +from __future__ import annotations + +import logging +from functools import wraps +from typing import TYPE_CHECKING, Callable, TypeVar, cast + +from flask import Response, request +from flask_login import login_user + +from airflow.utils.airflow_flask_app import get_airflow_app + +if TYPE_CHECKING: + from requests.auth import AuthBase + +log = logging.getLogger(__name__) + +CLIENT_AUTH: tuple[str, str] | AuthBase | None = None + + +def init_app(_): + """Initializes authentication backend""" + + +T = TypeVar("T", bound=Callable) + + +def _lookup_user(user_email_or_username: str): + security_manager = get_airflow_app().appbuilder.sm + user = security_manager.find_user(email=user_email_or_username) or security_manager.find_user( + username=user_email_or_username + ) + if not user: + return None + + if not user.is_active: + return None + + return user + + +def requires_authentication(function: T): + """Decorator for functions that require authentication""" + + @wraps(function) + def decorated(*args, **kwargs): + user_id = request.remote_user + if not user_id: + log.debug("Missing REMOTE_USER.") + return Response("Forbidden", 403) + + log.debug("Looking for user: %s", user_id) + + user = _lookup_user(user_id) + if not user: + return Response("Forbidden", 403) + + log.debug("Found user: %s", user) + + login_user(user, remember=False) + return function(*args, **kwargs) + + return cast(T, decorated) diff --git a/tests/providers/fab/auth_manager/api_endpoints/test_auth.py b/tests/providers/fab/auth_manager/api_endpoints/test_auth.py new file mode 100644 index 0000000000000..d3012e2f1b43e --- /dev/null +++ b/tests/providers/fab/auth_manager/api_endpoints/test_auth.py @@ -0,0 +1,176 @@ +# 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 + +from base64 import b64encode + +import pytest +from flask_login import current_user + +from tests.test_utils.api_connexion_utils import assert_401 +from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS +from tests.test_utils.config import conf_vars +from tests.test_utils.db import clear_db_pools +from tests.test_utils.www import client_with_login + +pytestmark = [ + pytest.mark.db_test, + pytest.mark.skip_if_database_isolation_mode, + pytest.mark.skipif(not AIRFLOW_V_3_0_PLUS, reason="Test requires Airflow 3.0+"), +] + + +class BaseTestAuth: + @pytest.fixture(autouse=True) + def set_attrs(self, minimal_app_for_auth_api): + self.app = minimal_app_for_auth_api + + sm = self.app.appbuilder.sm + tester = sm.find_user(username="test") + if not tester: + role_admin = sm.find_role("Admin") + sm.add_user( + username="test", + first_name="test", + last_name="test", + email="test@fab.org", + role=role_admin, + password="test", + ) + + +class TestBasicAuth(BaseTestAuth): + @pytest.fixture(autouse=True, scope="class") + def with_basic_auth_backend(self, minimal_app_for_auth_api): + from airflow.www.extensions.init_security import init_api_auth + + old_auth = getattr(minimal_app_for_auth_api, "api_auth") + + try: + with conf_vars( + {("api", "auth_backends"): "airflow.providers.fab.auth_manager.api.auth.backend.basic_auth"} + ): + init_api_auth(minimal_app_for_auth_api) + yield + finally: + setattr(minimal_app_for_auth_api, "api_auth", old_auth) + + def test_success(self): + token = "Basic " + b64encode(b"test:test").decode() + clear_db_pools() + + with self.app.test_client() as test_client: + response = test_client.get("/api/v1/pools", headers={"Authorization": token}) + assert current_user.email == "test@fab.org" + + assert response.status_code == 200 + assert response.json == { + "pools": [ + { + "name": "default_pool", + "slots": 128, + "occupied_slots": 0, + "running_slots": 0, + "queued_slots": 0, + "scheduled_slots": 0, + "deferred_slots": 0, + "open_slots": 128, + "description": "Default pool", + "include_deferred": False, + }, + ], + "total_entries": 1, + } + + @pytest.mark.parametrize( + "token", + [ + "basic", + "basic ", + "bearer", + "test:test", + b64encode(b"test:test").decode(), + "bearer ", + "basic: ", + "basic 123", + ], + ) + def test_malformed_headers(self, token): + with self.app.test_client() as test_client: + response = test_client.get("/api/v1/pools", headers={"Authorization": token}) + assert response.status_code == 401 + assert response.headers["Content-Type"] == "application/problem+json" + assert response.headers["WWW-Authenticate"] == "Basic" + assert_401(response) + + @pytest.mark.parametrize( + "token", + [ + "basic " + b64encode(b"test").decode(), + "basic " + b64encode(b"test:").decode(), + "basic " + b64encode(b"test:123").decode(), + "basic " + b64encode(b"test test").decode(), + ], + ) + def test_invalid_auth_header(self, token): + with self.app.test_client() as test_client: + response = test_client.get("/api/v1/pools", headers={"Authorization": token}) + assert response.status_code == 401 + assert response.headers["Content-Type"] == "application/problem+json" + assert response.headers["WWW-Authenticate"] == "Basic" + assert_401(response) + + +class TestSessionWithBasicAuthFallback(BaseTestAuth): + @pytest.fixture(autouse=True, scope="class") + def with_basic_auth_backend(self, minimal_app_for_auth_api): + from airflow.www.extensions.init_security import init_api_auth + + old_auth = getattr(minimal_app_for_auth_api, "api_auth") + + try: + with conf_vars( + { + ( + "api", + "auth_backends", + ): "airflow.api.auth.backend.session,airflow.providers.fab.auth_manager.api.auth.backend.basic_auth" + } + ): + init_api_auth(minimal_app_for_auth_api) + yield + finally: + setattr(minimal_app_for_auth_api, "api_auth", old_auth) + + def test_basic_auth_fallback(self): + token = "Basic " + b64encode(b"test:test").decode() + clear_db_pools() + + # request uses session + admin_user = client_with_login(self.app, username="test", password="test") + response = admin_user.get("/api/v1/pools") + assert response.status_code == 200 + + # request uses basic auth + with self.app.test_client() as test_client: + response = test_client.get("/api/v1/pools", headers={"Authorization": token}) + assert response.status_code == 200 + + # request without session or basic auth header + with self.app.test_client() as test_client: + response = test_client.get("/api/v1/pools") + assert response.status_code == 401 diff --git a/tests/providers/fab/auth_manager/api_endpoints/test_backfill_endpoint.py b/tests/providers/fab/auth_manager/api_endpoints/test_backfill_endpoint.py new file mode 100644 index 0000000000000..56f135d457e9c --- /dev/null +++ b/tests/providers/fab/auth_manager/api_endpoints/test_backfill_endpoint.py @@ -0,0 +1,264 @@ +# 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 os +from datetime import datetime +from unittest import mock +from urllib.parse import urlencode + +import pendulum +import pytest + +from airflow.models import DagBag, DagModel +from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS + +try: + from airflow.models.backfill import Backfill +except ImportError: + if AIRFLOW_V_3_0_PLUS: + raise + else: + pass +from airflow.models.dag import DAG +from airflow.models.serialized_dag import SerializedDagModel +from airflow.operators.empty import EmptyOperator +from airflow.security import permissions +from airflow.utils import timezone +from airflow.utils.session import provide_session +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_user, delete_user +from tests.test_utils.db import clear_db_backfills, clear_db_dags, clear_db_runs, clear_db_serialized_dags + +pytestmark = [ + pytest.mark.db_test, + pytest.mark.skipif(not AIRFLOW_V_3_0_PLUS, reason="Test requires Airflow 3.0+"), +] + + +DAG_ID = "test_dag" +TASK_ID = "op1" +DAG2_ID = "test_dag2" +DAG3_ID = "test_dag3" +UTC_JSON_REPR = "UTC" if pendulum.__version__.startswith("3") else "Timezone('UTC')" + + +@pytest.fixture(scope="module") +def configured_app(minimal_app_for_auth_api): + app = minimal_app_for_auth_api + + create_user(app, username="test_granular_permissions", role_name="TestGranularDag") + app.appbuilder.sm.sync_perm_for_dag( + "TEST_DAG_1", + access_control={ + "TestGranularDag": { + permissions.RESOURCE_DAG: {permissions.ACTION_CAN_EDIT, permissions.ACTION_CAN_READ} + }, + }, + ) + + with DAG( + DAG_ID, + schedule=None, + start_date=datetime(2020, 6, 15), + doc_md="details", + params={"foo": 1}, + tags=["example"], + ) as dag: + EmptyOperator(task_id=TASK_ID) + + with DAG(DAG2_ID, schedule=None, start_date=datetime(2020, 6, 15)) as dag2: # no doc_md + EmptyOperator(task_id=TASK_ID) + + with DAG(DAG3_ID, schedule=None) as dag3: # DAG start_date set to None + EmptyOperator(task_id=TASK_ID, start_date=datetime(2019, 6, 12)) + + dag_bag = DagBag(os.devnull, include_examples=False) + dag_bag.dags = {dag.dag_id: dag, dag2.dag_id: dag2, dag3.dag_id: dag3} + + app.dag_bag = dag_bag + + yield app + + delete_user(app, username="test_granular_permissions") + + +class TestBackfillEndpoint: + @staticmethod + def clean_db(): + clear_db_backfills() + clear_db_runs() + clear_db_dags() + clear_db_serialized_dags() + + @pytest.fixture(autouse=True) + def setup_attrs(self, configured_app) -> None: + self.clean_db() + self.app = configured_app + self.client = self.app.test_client() # type:ignore + self.dag_id = DAG_ID + self.dag2_id = DAG2_ID + self.dag3_id = DAG3_ID + + def teardown_method(self) -> None: + self.clean_db() + + @provide_session + def _create_dag_models(self, *, count=1, dag_id_prefix="TEST_DAG", is_paused=False, session=None): + dags = [] + for num in range(1, count + 1): + dag_model = DagModel( + dag_id=f"{dag_id_prefix}_{num}", + fileloc=f"/tmp/dag_{num}.py", + is_active=True, + timetable_summary="0 0 * * *", + is_paused=is_paused, + ) + session.add(dag_model) + dags.append(dag_model) + return dags + + @provide_session + def _create_deactivated_dag(self, session=None): + dag_model = DagModel( + dag_id="TEST_DAG_DELETED_1", + fileloc="/tmp/dag_del_1.py", + schedule_interval="2 2 * * *", + is_active=False, + ) + session.add(dag_model) + + +class TestListBackfills(TestBackfillEndpoint): + def test_should_respond_200_with_granular_dag_access(self, session): + (dag,) = self._create_dag_models() + from_date = timezone.utcnow() + to_date = timezone.utcnow() + b = Backfill( + dag_id=dag.dag_id, + from_date=from_date, + to_date=to_date, + ) + + session.add(b) + session.commit() + kwargs = {} + kwargs.update(environ_overrides={"REMOTE_USER": "test_granular_permissions"}) + response = self.client.get("/api/v1/backfills?dag_id=TEST_DAG_1", **kwargs) + assert response.status_code == 200 + + +class TestGetBackfill(TestBackfillEndpoint): + def test_should_respond_200_with_granular_dag_access(self, session): + (dag,) = self._create_dag_models() + from_date = timezone.utcnow() + to_date = timezone.utcnow() + backfill = Backfill( + dag_id=dag.dag_id, + from_date=from_date, + to_date=to_date, + ) + session.add(backfill) + session.commit() + kwargs = {} + kwargs.update(environ_overrides={"REMOTE_USER": "test_granular_permissions"}) + response = self.client.get(f"/api/v1/backfills/{backfill.id}", **kwargs) + assert response.status_code == 200 + + +class TestCreateBackfill(TestBackfillEndpoint): + def test_create_backfill(self, session, dag_maker): + with dag_maker(session=session, dag_id="TEST_DAG_1", schedule="0 * * * *") as dag: + EmptyOperator(task_id="mytask") + session.add(SerializedDagModel(dag)) + session.commit() + session.query(DagModel).all() + from_date = pendulum.parse("2024-01-01") + from_date_iso = from_date.isoformat() + to_date = pendulum.parse("2024-02-01") + to_date_iso = to_date.isoformat() + max_active_runs = 5 + query = urlencode( + query={ + "dag_id": dag.dag_id, + "from_date": f"{from_date_iso}", + "to_date": f"{to_date_iso}", + "max_active_runs": max_active_runs, + "reverse": False, + } + ) + kwargs = {} + kwargs.update(environ_overrides={"REMOTE_USER": "test_granular_permissions"}) + + response = self.client.post( + f"/api/v1/backfills?{query}", + **kwargs, + ) + assert response.status_code == 200 + assert response.json == { + "completed_at": mock.ANY, + "created_at": mock.ANY, + "dag_id": "TEST_DAG_1", + "dag_run_conf": None, + "from_date": from_date_iso, + "id": mock.ANY, + "is_paused": False, + "max_active_runs": 5, + "to_date": to_date_iso, + "updated_at": mock.ANY, + } + + +class TestPauseBackfill(TestBackfillEndpoint): + def test_should_respond_200_with_granular_dag_access(self, session): + (dag,) = self._create_dag_models() + from_date = timezone.utcnow() + to_date = timezone.utcnow() + backfill = Backfill( + dag_id=dag.dag_id, + from_date=from_date, + to_date=to_date, + ) + session.add(backfill) + session.commit() + kwargs = {} + kwargs.update(environ_overrides={"REMOTE_USER": "test_granular_permissions"}) + response = self.client.post(f"/api/v1/backfills/{backfill.id}/pause", **kwargs) + assert response.status_code == 200 + + +class TestCancelBackfill(TestBackfillEndpoint): + def test_should_respond_200_with_granular_dag_access(self, session): + (dag,) = self._create_dag_models() + from_date = timezone.utcnow() + to_date = timezone.utcnow() + backfill = Backfill( + dag_id=dag.dag_id, + from_date=from_date, + to_date=to_date, + ) + session.add(backfill) + session.commit() + kwargs = {} + kwargs.update(environ_overrides={"REMOTE_USER": "test_granular_permissions"}) + response = self.client.post(f"/api/v1/backfills/{backfill.id}/cancel", **kwargs) + assert response.status_code == 200 + # now it is marked as completed + assert pendulum.parse(response.json["completed_at"]) + + # get conflict when canceling already-canceled backfill + response = self.client.post(f"/api/v1/backfills/{backfill.id}/cancel", **kwargs) + assert response.status_code == 409 diff --git a/tests/api_connexion/test_cors.py b/tests/providers/fab/auth_manager/api_endpoints/test_cors.py similarity index 81% rename from tests/api_connexion/test_cors.py rename to tests/providers/fab/auth_manager/api_endpoints/test_cors.py index a2b7f0ebca743..b44eab8820ec6 100644 --- a/tests/api_connexion/test_cors.py +++ b/tests/providers/fab/auth_manager/api_endpoints/test_cors.py @@ -20,16 +20,21 @@ import pytest +from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS from tests.test_utils.config import conf_vars from tests.test_utils.db import clear_db_pools -pytestmark = [pytest.mark.db_test, pytest.mark.skip_if_database_isolation_mode] +pytestmark = [ + pytest.mark.db_test, + pytest.mark.skip_if_database_isolation_mode, + pytest.mark.skipif(not AIRFLOW_V_3_0_PLUS, reason="Test requires Airflow 3.0+"), +] class BaseTestAuth: @pytest.fixture(autouse=True) - def set_attrs(self, minimal_app_for_api): - self.app = minimal_app_for_api + def set_attrs(self, minimal_app_for_auth_api): + self.app = minimal_app_for_auth_api sm = self.app.appbuilder.sm tester = sm.find_user(username="test") @@ -47,19 +52,19 @@ def set_attrs(self, minimal_app_for_api): class TestEmptyCors(BaseTestAuth): @pytest.fixture(autouse=True, scope="class") - def with_basic_auth_backend(self, minimal_app_for_api): + def with_basic_auth_backend(self, minimal_app_for_auth_api): from airflow.www.extensions.init_security import init_api_auth - old_auth = getattr(minimal_app_for_api, "api_auth") + old_auth = getattr(minimal_app_for_auth_api, "api_auth") try: with conf_vars( {("api", "auth_backends"): "airflow.providers.fab.auth_manager.api.auth.backend.basic_auth"} ): - init_api_auth(minimal_app_for_api) + init_api_auth(minimal_app_for_auth_api) yield finally: - setattr(minimal_app_for_api, "api_auth", old_auth) + setattr(minimal_app_for_auth_api, "api_auth", old_auth) def test_empty_cors_headers(self): token = "Basic " + b64encode(b"test:test").decode() @@ -75,10 +80,10 @@ def test_empty_cors_headers(self): class TestCorsOrigin(BaseTestAuth): @pytest.fixture(autouse=True, scope="class") - def with_basic_auth_backend(self, minimal_app_for_api): + def with_basic_auth_backend(self, minimal_app_for_auth_api): from airflow.www.extensions.init_security import init_api_auth - old_auth = getattr(minimal_app_for_api, "api_auth") + old_auth = getattr(minimal_app_for_auth_api, "api_auth") try: with conf_vars( @@ -90,10 +95,10 @@ def with_basic_auth_backend(self, minimal_app_for_api): ("api", "access_control_allow_origins"): "http://apache.org http://example.com", } ): - init_api_auth(minimal_app_for_api) + init_api_auth(minimal_app_for_auth_api) yield finally: - setattr(minimal_app_for_api, "api_auth", old_auth) + setattr(minimal_app_for_auth_api, "api_auth", old_auth) def test_cors_origin_reflection(self): token = "Basic " + b64encode(b"test:test").decode() @@ -119,10 +124,10 @@ def test_cors_origin_reflection(self): class TestCorsWildcard(BaseTestAuth): @pytest.fixture(autouse=True, scope="class") - def with_basic_auth_backend(self, minimal_app_for_api): + def with_basic_auth_backend(self, minimal_app_for_auth_api): from airflow.www.extensions.init_security import init_api_auth - old_auth = getattr(minimal_app_for_api, "api_auth") + old_auth = getattr(minimal_app_for_auth_api, "api_auth") try: with conf_vars( @@ -134,10 +139,10 @@ def with_basic_auth_backend(self, minimal_app_for_api): ("api", "access_control_allow_origins"): "*", } ): - init_api_auth(minimal_app_for_api) + init_api_auth(minimal_app_for_auth_api) yield finally: - setattr(minimal_app_for_api, "api_auth", old_auth) + setattr(minimal_app_for_auth_api, "api_auth", old_auth) def test_cors_origin_reflection(self): token = "Basic " + b64encode(b"test:test").decode() diff --git a/tests/providers/fab/auth_manager/api_endpoints/test_dag_endpoint.py b/tests/providers/fab/auth_manager/api_endpoints/test_dag_endpoint.py new file mode 100644 index 0000000000000..b78ac58e442e0 --- /dev/null +++ b/tests/providers/fab/auth_manager/api_endpoints/test_dag_endpoint.py @@ -0,0 +1,252 @@ +# 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 os +from datetime import datetime + +import pendulum +import pytest + +from airflow.api_connexion.exceptions import EXCEPTIONS_LINK_MAP +from airflow.models import DagBag, DagModel +from airflow.models.dag import DAG +from airflow.operators.empty import EmptyOperator +from airflow.security import permissions +from airflow.utils.session import provide_session +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_user, delete_user +from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS +from tests.test_utils.db import clear_db_dags, clear_db_runs, clear_db_serialized_dags +from tests.test_utils.www import _check_last_log + +pytestmark = [ + pytest.mark.db_test, + pytest.mark.skip_if_database_isolation_mode, + pytest.mark.skipif(not AIRFLOW_V_3_0_PLUS, reason="Test requires Airflow 3.0+"), +] + + +@pytest.fixture +def current_file_token(url_safe_serializer) -> str: + return url_safe_serializer.dumps(__file__) + + +DAG_ID = "test_dag" +TASK_ID = "op1" +DAG2_ID = "test_dag2" +DAG3_ID = "test_dag3" +UTC_JSON_REPR = "UTC" if pendulum.__version__.startswith("3") else "Timezone('UTC')" + + +@pytest.fixture(scope="module") +def configured_app(minimal_app_for_auth_api): + app = minimal_app_for_auth_api + + create_user(app, username="test_granular_permissions", role_name="TestGranularDag") + app.appbuilder.sm.sync_perm_for_dag( + "TEST_DAG_1", + access_control={ + "TestGranularDag": { + permissions.RESOURCE_DAG: {permissions.ACTION_CAN_EDIT, permissions.ACTION_CAN_READ} + }, + }, + ) + app.appbuilder.sm.sync_perm_for_dag( + "TEST_DAG_1", + access_control={ + "TestGranularDag": { + permissions.RESOURCE_DAG: {permissions.ACTION_CAN_EDIT, permissions.ACTION_CAN_READ} + }, + }, + ) + + with DAG( + DAG_ID, + schedule=None, + start_date=datetime(2020, 6, 15), + doc_md="details", + params={"foo": 1}, + tags=["example"], + ) as dag: + EmptyOperator(task_id=TASK_ID) + + with DAG(DAG2_ID, schedule=None, start_date=datetime(2020, 6, 15)) as dag2: # no doc_md + EmptyOperator(task_id=TASK_ID) + + with DAG(DAG3_ID, schedule=None) as dag3: # DAG start_date set to None + EmptyOperator(task_id=TASK_ID, start_date=datetime(2019, 6, 12)) + + dag_bag = DagBag(os.devnull, include_examples=False) + dag_bag.dags = {dag.dag_id: dag, dag2.dag_id: dag2, dag3.dag_id: dag3} + + app.dag_bag = dag_bag + + yield app + + delete_user(app, username="test_granular_permissions") + + +class TestDagEndpoint: + @staticmethod + def clean_db(): + clear_db_runs() + clear_db_dags() + clear_db_serialized_dags() + + @pytest.fixture(autouse=True) + def setup_attrs(self, configured_app) -> None: + self.clean_db() + self.app = configured_app + self.client = self.app.test_client() # type:ignore + self.dag_id = DAG_ID + self.dag2_id = DAG2_ID + self.dag3_id = DAG3_ID + + def teardown_method(self) -> None: + self.clean_db() + + @provide_session + def _create_dag_models(self, count, dag_id_prefix="TEST_DAG", is_paused=False, session=None): + for num in range(1, count + 1): + dag_model = DagModel( + dag_id=f"{dag_id_prefix}_{num}", + fileloc=f"/tmp/dag_{num}.py", + timetable_summary="2 2 * * *", + is_active=True, + is_paused=is_paused, + ) + session.add(dag_model) + + @provide_session + def _create_dag_model_for_details_endpoint(self, dag_id, session=None): + dag_model = DagModel( + dag_id=dag_id, + fileloc="/tmp/dag.py", + timetable_summary="2 2 * * *", + is_active=True, + is_paused=False, + ) + session.add(dag_model) + + @provide_session + def _create_dag_model_for_details_endpoint_with_dataset_expression(self, dag_id, session=None): + dag_model = DagModel( + dag_id=dag_id, + fileloc="/tmp/dag.py", + timetable_summary="2 2 * * *", + is_active=True, + is_paused=False, + dataset_expression={ + "any": [ + "s3://dag1/output_1.txt", + {"all": ["s3://dag2/output_1.txt", "s3://dag3/output_3.txt"]}, + ] + }, + ) + session.add(dag_model) + + @provide_session + def _create_deactivated_dag(self, session=None): + dag_model = DagModel( + dag_id="TEST_DAG_DELETED_1", + fileloc="/tmp/dag_del_1.py", + timetable_summary="2 2 * * *", + is_active=False, + ) + session.add(dag_model) + + +class TestGetDag(TestDagEndpoint): + def test_should_respond_200_with_granular_dag_access(self): + self._create_dag_models(1) + response = self.client.get( + "/api/v1/dags/TEST_DAG_1", environ_overrides={"REMOTE_USER": "test_granular_permissions"} + ) + assert response.status_code == 200 + + def test_should_respond_403_with_granular_access_for_different_dag(self): + self._create_dag_models(3) + response = self.client.get( + "/api/v1/dags/TEST_DAG_2", environ_overrides={"REMOTE_USER": "test_granular_permissions"} + ) + assert response.status_code == 403 + + +class TestGetDags(TestDagEndpoint): + def test_should_respond_200_with_granular_dag_access(self): + self._create_dag_models(3) + response = self.client.get( + "/api/v1/dags", environ_overrides={"REMOTE_USER": "test_granular_permissions"} + ) + assert response.status_code == 200 + assert len(response.json["dags"]) == 1 + assert response.json["dags"][0]["dag_id"] == "TEST_DAG_1" + + +class TestPatchDag(TestDagEndpoint): + @provide_session + def _create_dag_model(self, session=None): + dag_model = DagModel( + dag_id="TEST_DAG_1", fileloc="/tmp/dag_1.py", timetable_summary="2 2 * * *", is_paused=True + ) + session.add(dag_model) + return dag_model + + def test_should_respond_200_on_patch_with_granular_dag_access(self, session): + self._create_dag_models(1) + response = self.client.patch( + "/api/v1/dags/TEST_DAG_1", + json={ + "is_paused": False, + }, + environ_overrides={"REMOTE_USER": "test_granular_permissions"}, + ) + assert response.status_code == 200 + _check_last_log(session, dag_id="TEST_DAG_1", event="api.patch_dag", execution_date=None) + + def test_validation_error_raises_400(self): + patch_body = { + "ispaused": True, + } + dag_model = self._create_dag_model() + response = self.client.patch( + f"/api/v1/dags/{dag_model.dag_id}", + json=patch_body, + environ_overrides={"REMOTE_USER": "test_granular_permissions"}, + ) + assert response.status_code == 400 + assert response.json == { + "detail": "{'ispaused': ['Unknown field.']}", + "status": 400, + "title": "Bad Request", + "type": EXCEPTIONS_LINK_MAP[400], + } + + +class TestPatchDags(TestDagEndpoint): + def test_should_respond_200_with_granular_dag_access(self): + self._create_dag_models(3) + response = self.client.patch( + "api/v1/dags?dag_id_pattern=~", + json={ + "is_paused": False, + }, + environ_overrides={"REMOTE_USER": "test_granular_permissions"}, + ) + assert response.status_code == 200 + assert len(response.json["dags"]) == 1 + assert response.json["dags"][0]["dag_id"] == "TEST_DAG_1" diff --git a/tests/providers/fab/auth_manager/api_endpoints/test_dag_run_endpoint.py b/tests/providers/fab/auth_manager/api_endpoints/test_dag_run_endpoint.py new file mode 100644 index 0000000000000..a58ea08ff31cf --- /dev/null +++ b/tests/providers/fab/auth_manager/api_endpoints/test_dag_run_endpoint.py @@ -0,0 +1,273 @@ +# 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 + +from datetime import timedelta + +import pytest + +from airflow.models.dag import DAG, DagModel +from airflow.models.dagrun import DagRun +from airflow.models.param import Param +from airflow.security import permissions +from airflow.utils import timezone +from airflow.utils.session import create_session +from airflow.utils.state import DagRunState +from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS + +try: + from airflow.utils.types import DagRunTriggeredByType, DagRunType +except ImportError: + if AIRFLOW_V_3_0_PLUS: + raise + else: + pass +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import ( + create_user, + delete_roles, + delete_user, +) +from tests.test_utils.db import clear_db_dags, clear_db_runs, clear_db_serialized_dags + +pytestmark = [ + pytest.mark.db_test, + pytest.mark.skip_if_database_isolation_mode, + pytest.mark.skipif(not AIRFLOW_V_3_0_PLUS, reason="Test requires Airflow 3.0+"), +] + + +@pytest.fixture(scope="module") +def configured_app(minimal_app_for_auth_api): + app = minimal_app_for_auth_api + + create_user( + app, + username="test_no_dag_run_create_permission", + role_name="TestNoDagRunCreatePermission", + permissions=[ + (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_ASSET), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_CLUSTER_ACTIVITY), + (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), + (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG_RUN), + (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_DAG_RUN), + ], + ) + create_user( + app, + username="test_dag_view_only", + role_name="TestViewDags", + permissions=[ + (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), + (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_DAG_RUN), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), + (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG_RUN), + (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_DAG_RUN), + ], + ) + create_user( + app, + username="test_view_dags", + role_name="TestViewDags", + permissions=[ + (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), + (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_DAG_RUN), + ], + ) + create_user( + app, + username="test_granular_permissions", + role_name="TestGranularDag", + permissions=[(permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN)], + ) + app.appbuilder.sm.sync_perm_for_dag( + "TEST_DAG_ID", + access_control={ + "TestGranularDag": {permissions.ACTION_CAN_EDIT, permissions.ACTION_CAN_READ}, + "TestNoDagRunCreatePermission": {permissions.RESOURCE_DAG_RUN: {permissions.ACTION_CAN_CREATE}}, + }, + ) + + yield app + + delete_user(app, username="test_dag_view_only") + delete_user(app, username="test_view_dags") + delete_user(app, username="test_granular_permissions") + delete_user(app, username="test_no_dag_run_create_permission") + delete_roles(app) + + +class TestDagRunEndpoint: + default_time = "2020-06-11T18:00:00+00:00" + default_time_2 = "2020-06-12T18:00:00+00:00" + default_time_3 = "2020-06-13T18:00:00+00:00" + + @pytest.fixture(autouse=True) + def setup_attrs(self, configured_app) -> None: + self.app = configured_app + self.client = self.app.test_client() # type:ignore + clear_db_runs() + clear_db_serialized_dags() + clear_db_dags() + + def teardown_method(self) -> None: + clear_db_runs() + clear_db_dags() + clear_db_serialized_dags() + + def _create_dag(self, dag_id): + dag_instance = DagModel(dag_id=dag_id) + dag_instance.is_active = True + with create_session() as session: + session.add(dag_instance) + dag = DAG(dag_id=dag_id, schedule=None, params={"validated_number": Param(1, minimum=1, maximum=10)}) + self.app.dag_bag.bag_dag(dag) + return dag_instance + + def _create_test_dag_run(self, state=DagRunState.RUNNING, extra_dag=False, commit=True, idx_start=1): + dag_runs = [] + dags = [] + triggered_by_kwargs = {"triggered_by": DagRunTriggeredByType.TEST} if AIRFLOW_V_3_0_PLUS else {} + + for i in range(idx_start, idx_start + 2): + if i == 1: + dags.append(DagModel(dag_id="TEST_DAG_ID", is_active=True)) + dagrun_model = DagRun( + dag_id="TEST_DAG_ID", + run_id=f"TEST_DAG_RUN_ID_{i}", + run_type=DagRunType.MANUAL, + execution_date=timezone.parse(self.default_time) + timedelta(days=i - 1), + start_date=timezone.parse(self.default_time), + external_trigger=True, + state=state, + **triggered_by_kwargs, + ) + dagrun_model.updated_at = timezone.parse(self.default_time) + dag_runs.append(dagrun_model) + + if extra_dag: + for i in range(idx_start + 2, idx_start + 4): + dags.append(DagModel(dag_id=f"TEST_DAG_ID_{i}")) + dag_runs.append( + DagRun( + dag_id=f"TEST_DAG_ID_{i}", + run_id=f"TEST_DAG_RUN_ID_{i}", + run_type=DagRunType.MANUAL, + execution_date=timezone.parse(self.default_time_2), + start_date=timezone.parse(self.default_time), + external_trigger=True, + state=state, + ) + ) + if commit: + with create_session() as session: + session.add_all(dag_runs) + session.add_all(dags) + return dag_runs + + +class TestGetDagRuns(TestDagRunEndpoint): + def test_should_return_accessible_with_tilde_as_dag_id_and_dag_level_permissions(self): + self._create_test_dag_run(extra_dag=True) + expected_dag_run_ids = ["TEST_DAG_ID", "TEST_DAG_ID"] + response = self.client.get( + "api/v1/dags/~/dagRuns", environ_overrides={"REMOTE_USER": "test_granular_permissions"} + ) + assert response.status_code == 200 + dag_run_ids = [dag_run["dag_id"] for dag_run in response.json["dag_runs"]] + assert dag_run_ids == expected_dag_run_ids + + +class TestGetDagRunBatch(TestDagRunEndpoint): + def test_should_return_accessible_with_tilde_as_dag_id_and_dag_level_permissions(self): + self._create_test_dag_run(extra_dag=True) + expected_response_json_1 = { + "dag_id": "TEST_DAG_ID", + "dag_run_id": "TEST_DAG_RUN_ID_1", + "end_date": None, + "state": "running", + "execution_date": self.default_time, + "logical_date": self.default_time, + "external_trigger": True, + "start_date": self.default_time, + "conf": {}, + "data_interval_end": None, + "data_interval_start": None, + "last_scheduling_decision": None, + "run_type": "manual", + "note": None, + } + expected_response_json_1.update({"triggered_by": "test"} if AIRFLOW_V_3_0_PLUS else {}) + expected_response_json_2 = { + "dag_id": "TEST_DAG_ID", + "dag_run_id": "TEST_DAG_RUN_ID_2", + "end_date": None, + "state": "running", + "execution_date": self.default_time_2, + "logical_date": self.default_time_2, + "external_trigger": True, + "start_date": self.default_time, + "conf": {}, + "data_interval_end": None, + "data_interval_start": None, + "last_scheduling_decision": None, + "run_type": "manual", + "note": None, + } + expected_response_json_2.update({"triggered_by": "test"} if AIRFLOW_V_3_0_PLUS else {}) + + response = self.client.post( + "api/v1/dags/~/dagRuns/list", + json={"dag_ids": []}, + environ_overrides={"REMOTE_USER": "test_granular_permissions"}, + ) + assert response.status_code == 200 + assert response.json == { + "dag_runs": [ + expected_response_json_1, + expected_response_json_2, + ], + "total_entries": 2, + } + + +class TestPostDagRun(TestDagRunEndpoint): + def test_dagrun_trigger_with_dag_level_permissions(self): + self._create_dag("TEST_DAG_ID") + response = self.client.post( + "api/v1/dags/TEST_DAG_ID/dagRuns", + json={"conf": {"validated_number": 1}}, + environ_overrides={"REMOTE_USER": "test_no_dag_run_create_permission"}, + ) + assert response.status_code == 200 + + @pytest.mark.parametrize( + "username", + ["test_dag_view_only", "test_view_dags", "test_granular_permissions"], + ) + def test_should_raises_403_unauthorized(self, username): + self._create_dag("TEST_DAG_ID") + response = self.client.post( + "api/v1/dags/TEST_DAG_ID/dagRuns", + json={ + "dag_run_id": "TEST_DAG_RUN_ID_1", + "execution_date": self.default_time, + }, + environ_overrides={"REMOTE_USER": username}, + ) + assert response.status_code == 403 diff --git a/tests/providers/fab/auth_manager/api_endpoints/test_dag_source_endpoint.py b/tests/providers/fab/auth_manager/api_endpoints/test_dag_source_endpoint.py new file mode 100644 index 0000000000000..f0d9b0da298c6 --- /dev/null +++ b/tests/providers/fab/auth_manager/api_endpoints/test_dag_source_endpoint.py @@ -0,0 +1,144 @@ +# 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 ast +import os +from typing import TYPE_CHECKING + +import pytest + +from airflow.models import DagBag +from airflow.security import permissions +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_user, delete_user +from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS +from tests.test_utils.db import clear_db_dag_code, clear_db_dags, clear_db_serialized_dags + +pytestmark = [ + pytest.mark.db_test, + pytest.mark.skip_if_database_isolation_mode, + pytest.mark.skipif(not AIRFLOW_V_3_0_PLUS, reason="Test requires Airflow 3.0+"), +] + +if TYPE_CHECKING: + from airflow.models.dag import DAG + +ROOT_DIR = os.path.abspath(os.path.join(os.path.dirname(__file__), os.pardir, os.pardir)) +EXAMPLE_DAG_FILE = os.path.join("airflow", "example_dags", "example_bash_operator.py") +EXAMPLE_DAG_ID = "example_bash_operator" +TEST_DAG_ID = "latest_only" +NOT_READABLE_DAG_ID = "latest_only_with_trigger" +TEST_MULTIPLE_DAGS_ID = "asset_produces_1" + + +@pytest.fixture(scope="module") +def configured_app(minimal_app_for_auth_api): + app = minimal_app_for_auth_api + create_user( + app, + username="test", + role_name="Test", + permissions=[(permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_CODE)], + ) + app.appbuilder.sm.sync_perm_for_dag( + TEST_DAG_ID, + access_control={"Test": [permissions.ACTION_CAN_READ]}, + ) + app.appbuilder.sm.sync_perm_for_dag( + EXAMPLE_DAG_ID, + access_control={"Test": [permissions.ACTION_CAN_READ]}, + ) + app.appbuilder.sm.sync_perm_for_dag( + TEST_MULTIPLE_DAGS_ID, + access_control={"Test": [permissions.ACTION_CAN_READ]}, + ) + + yield app + + delete_user(app, username="test") + + +class TestGetSource: + @pytest.fixture(autouse=True) + def setup_attrs(self, configured_app) -> None: + self.app = configured_app + self.client = self.app.test_client() # type:ignore + self.clear_db() + + def teardown_method(self) -> None: + self.clear_db() + + @staticmethod + def clear_db(): + clear_db_dags() + clear_db_serialized_dags() + clear_db_dag_code() + + @staticmethod + def _get_dag_file_docstring(fileloc: str) -> str | None: + with open(fileloc) as f: + file_contents = f.read() + module = ast.parse(file_contents) + docstring = ast.get_docstring(module) + return docstring + + def test_should_respond_406(self, url_safe_serializer): + dagbag = DagBag(dag_folder=EXAMPLE_DAG_FILE) + dagbag.sync_to_db() + test_dag: DAG = dagbag.dags[TEST_DAG_ID] + + url = f"/api/v1/dagSources/{url_safe_serializer.dumps(test_dag.fileloc)}" + response = self.client.get( + url, headers={"Accept": "image/webp"}, environ_overrides={"REMOTE_USER": "test"} + ) + + assert 406 == response.status_code + + def test_should_respond_403_not_readable(self, url_safe_serializer): + dagbag = DagBag(dag_folder=EXAMPLE_DAG_FILE) + dagbag.sync_to_db() + dag: DAG = dagbag.dags[NOT_READABLE_DAG_ID] + + response = self.client.get( + f"/api/v1/dagSources/{url_safe_serializer.dumps(dag.fileloc)}", + headers={"Accept": "text/plain"}, + environ_overrides={"REMOTE_USER": "test"}, + ) + read_dag = self.client.get( + f"/api/v1/dags/{NOT_READABLE_DAG_ID}", + environ_overrides={"REMOTE_USER": "test"}, + ) + assert response.status_code == 403 + assert read_dag.status_code == 403 + + def test_should_respond_403_some_dags_not_readable_in_the_file(self, url_safe_serializer): + dagbag = DagBag(dag_folder=EXAMPLE_DAG_FILE) + dagbag.sync_to_db() + dag: DAG = dagbag.dags[TEST_MULTIPLE_DAGS_ID] + + response = self.client.get( + f"/api/v1/dagSources/{url_safe_serializer.dumps(dag.fileloc)}", + headers={"Accept": "text/plain"}, + environ_overrides={"REMOTE_USER": "test"}, + ) + + read_dag = self.client.get( + f"/api/v1/dags/{TEST_MULTIPLE_DAGS_ID}", + environ_overrides={"REMOTE_USER": "test"}, + ) + assert response.status_code == 403 + assert read_dag.status_code == 200 diff --git a/tests/providers/fab/auth_manager/api_endpoints/test_dag_warning_endpoint.py b/tests/providers/fab/auth_manager/api_endpoints/test_dag_warning_endpoint.py new file mode 100644 index 0000000000000..adfde1cc5b3eb --- /dev/null +++ b/tests/providers/fab/auth_manager/api_endpoints/test_dag_warning_endpoint.py @@ -0,0 +1,84 @@ +# 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 pytest + +from airflow.models.dag import DagModel +from airflow.models.dagwarning import DagWarning +from airflow.security import permissions +from airflow.utils.session import create_session +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_user, delete_user +from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS +from tests.test_utils.db import clear_db_dag_warnings, clear_db_dags + +pytestmark = [ + pytest.mark.db_test, + pytest.mark.skip_if_database_isolation_mode, + pytest.mark.skipif(not AIRFLOW_V_3_0_PLUS, reason="Test requires Airflow 3.0+"), +] + + +@pytest.fixture(scope="module") +def configured_app(minimal_app_for_auth_api): + app = minimal_app_for_auth_api + create_user( + app, # type:ignore + username="test_with_dag2_read", + role_name="TestWithDag2Read", + permissions=[ + (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_WARNING), + (permissions.ACTION_CAN_READ, f"{permissions.RESOURCE_DAG_PREFIX}dag2"), + ], + ) + + yield app + + delete_user(app, username="test_with_dag2_read") + + +class TestBaseDagWarning: + timestamp = "2020-06-10T12:00" + + @pytest.fixture(autouse=True) + def setup_attrs(self, configured_app) -> None: + self.app = configured_app + self.client = self.app.test_client() # type:ignore + + def teardown_method(self) -> None: + clear_db_dag_warnings() + clear_db_dags() + + +class TestGetDagWarningEndpoint(TestBaseDagWarning): + def setup_class(self): + clear_db_dag_warnings() + clear_db_dags() + + def setup_method(self): + with create_session() as session: + session.add(DagModel(dag_id="dag1")) + session.add(DagWarning("dag1", "non-existent pool", "test message")) + session.commit() + + def test_should_raise_403_forbidden_when_user_has_no_dag_read_permission(self): + response = self.client.get( + "/api/v1/dagWarnings", + environ_overrides={"REMOTE_USER": "test_with_dag2_read"}, + query_string={"dag_id": "dag1"}, + ) + assert response.status_code == 403 diff --git a/tests/providers/fab/auth_manager/api_endpoints/test_dataset_endpoint.py b/tests/providers/fab/auth_manager/api_endpoints/test_dataset_endpoint.py new file mode 100644 index 0000000000000..4d302722223d8 --- /dev/null +++ b/tests/providers/fab/auth_manager/api_endpoints/test_dataset_endpoint.py @@ -0,0 +1,327 @@ +# 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 + +from typing import Generator + +import pytest +import time_machine + +from airflow.api_connexion.exceptions import EXCEPTIONS_LINK_MAP +from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS + +try: + from airflow.models.asset import AssetDagRunQueue, AssetModel +except ImportError: + if AIRFLOW_V_3_0_PLUS: + raise + else: + pass +from airflow.security import permissions +from airflow.utils import timezone +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_user, delete_user +from tests.test_utils.db import clear_db_assets, clear_db_runs +from tests.test_utils.www import _check_last_log + +pytestmark = [ + pytest.mark.db_test, + pytest.mark.skip_if_database_isolation_mode, + pytest.mark.skipif(not AIRFLOW_V_3_0_PLUS, reason="Test requires Airflow 3.0+"), +] + + +@pytest.fixture(scope="module") +def configured_app(minimal_app_for_auth_api): + app = minimal_app_for_auth_api + create_user( + app, + username="test_queued_event", + role_name="TestQueuedEvent", + permissions=[ + (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_ASSET), + (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_ASSET), + ], + ) + + yield app + + delete_user(app, username="test_queued_event") + + +class TestAssetEndpoint: + default_time = "2020-06-11T18:00:00+00:00" + + @pytest.fixture(autouse=True) + def setup_attrs(self, configured_app) -> None: + self.app = configured_app + self.client = self.app.test_client() + clear_db_assets() + clear_db_runs() + + def teardown_method(self) -> None: + clear_db_assets() + clear_db_runs() + + def _create_asset(self, session): + asset_model = AssetModel( + id=1, + uri="s3://bucket/key", + extra={"foo": "bar"}, + created_at=timezone.parse(self.default_time), + updated_at=timezone.parse(self.default_time), + ) + session.add(asset_model) + session.commit() + return asset_model + + +class TestQueuedEventEndpoint(TestAssetEndpoint): + @pytest.fixture + def time_freezer(self) -> Generator: + freezer = time_machine.travel(self.default_time, tick=False) + freezer.start() + + yield + + freezer.stop() + + def _create_asset_dag_run_queues(self, dag_id, dataset_id, session): + ddrq = AssetDagRunQueue(target_dag_id=dag_id, dataset_id=dataset_id) + session.add(ddrq) + session.commit() + return ddrq + + +class TestGetDagDatasetQueuedEvent(TestQueuedEventEndpoint): + @pytest.mark.usefixtures("time_freezer") + def test_should_respond_200(self, session, create_dummy_dag): + dag, _ = create_dummy_dag() + dag_id = dag.dag_id + dataset_id = self._create_asset(session).id + self._create_asset_dag_run_queues(dag_id, dataset_id, session) + dataset_uri = "s3://bucket/key" + + response = self.client.get( + f"/api/v1/dags/{dag_id}/datasets/queuedEvent/{dataset_uri}", + environ_overrides={"REMOTE_USER": "test_queued_event"}, + ) + + assert response.status_code == 200 + assert response.json == { + "created_at": self.default_time, + "uri": "s3://bucket/key", + "dag_id": "dag", + } + + def test_should_respond_404(self): + dag_id = "not_exists" + dataset_uri = "not_exists" + + response = self.client.get( + f"/api/v1/dags/{dag_id}/datasets/queuedEvent/{dataset_uri}", + environ_overrides={"REMOTE_USER": "test_queued_event"}, + ) + + assert response.status_code == 404 + assert { + "detail": "Queue event with dag_id: `not_exists` and asset uri: `not_exists` was not found", + "status": 404, + "title": "Queue event not found", + "type": EXCEPTIONS_LINK_MAP[404], + } == response.json + + +class TestDeleteDagDatasetQueuedEvent(TestAssetEndpoint): + def test_delete_should_respond_204(self, session, create_dummy_dag): + dag, _ = create_dummy_dag() + dag_id = dag.dag_id + dataset_uri = "s3://bucket/key" + dataset_id = self._create_asset(session).id + + ddrq = AssetDagRunQueue(target_dag_id=dag_id, dataset_id=dataset_id) + session.add(ddrq) + session.commit() + conn = session.query(AssetDagRunQueue).all() + assert len(conn) == 1 + + response = self.client.delete( + f"/api/v1/dags/{dag_id}/datasets/queuedEvent/{dataset_uri}", + environ_overrides={"REMOTE_USER": "test_queued_event"}, + ) + + assert response.status_code == 204 + conn = session.query(AssetDagRunQueue).all() + assert len(conn) == 0 + _check_last_log( + session, dag_id=dag_id, event="api.delete_dag_dataset_queued_event", execution_date=None + ) + + def test_should_respond_404(self): + dag_id = "not_exists" + dataset_uri = "not_exists" + + response = self.client.delete( + f"/api/v1/dags/{dag_id}/datasets/queuedEvent/{dataset_uri}", + environ_overrides={"REMOTE_USER": "test_queued_event"}, + ) + + assert response.status_code == 404 + assert { + "detail": "Queue event with dag_id: `not_exists` and asset uri: `not_exists` was not found", + "status": 404, + "title": "Queue event not found", + "type": EXCEPTIONS_LINK_MAP[404], + } == response.json + + +class TestGetDagDatasetQueuedEvents(TestQueuedEventEndpoint): + @pytest.mark.usefixtures("time_freezer") + def test_should_respond_200(self, session, create_dummy_dag): + dag, _ = create_dummy_dag() + dag_id = dag.dag_id + dataset_id = self._create_asset(session).id + self._create_asset_dag_run_queues(dag_id, dataset_id, session) + + response = self.client.get( + f"/api/v1/dags/{dag_id}/datasets/queuedEvent", + environ_overrides={"REMOTE_USER": "test_queued_event"}, + ) + + assert response.status_code == 200 + assert response.json == { + "queued_events": [ + { + "created_at": self.default_time, + "uri": "s3://bucket/key", + "dag_id": "dag", + } + ], + "total_entries": 1, + } + + def test_should_respond_404(self): + dag_id = "not_exists" + + response = self.client.get( + f"/api/v1/dags/{dag_id}/datasets/queuedEvent", + environ_overrides={"REMOTE_USER": "test_queued_event"}, + ) + + assert response.status_code == 404 + assert { + "detail": "Queue event with dag_id: `not_exists` was not found", + "status": 404, + "title": "Queue event not found", + "type": EXCEPTIONS_LINK_MAP[404], + } == response.json + + +class TestDeleteDagDatasetQueuedEvents(TestAssetEndpoint): + def test_should_respond_404(self): + dag_id = "not_exists" + + response = self.client.delete( + f"/api/v1/dags/{dag_id}/datasets/queuedEvent", + environ_overrides={"REMOTE_USER": "test_queued_event"}, + ) + + assert response.status_code == 404 + assert { + "detail": "Queue event with dag_id: `not_exists` was not found", + "status": 404, + "title": "Queue event not found", + "type": EXCEPTIONS_LINK_MAP[404], + } == response.json + + +class TestGetDatasetQueuedEvents(TestQueuedEventEndpoint): + @pytest.mark.usefixtures("time_freezer") + def test_should_respond_200(self, session, create_dummy_dag): + dag, _ = create_dummy_dag() + dag_id = dag.dag_id + dataset_id = self._create_asset(session).id + self._create_asset_dag_run_queues(dag_id, dataset_id, session) + dataset_uri = "s3://bucket/key" + + response = self.client.get( + f"/api/v1/datasets/queuedEvent/{dataset_uri}", + environ_overrides={"REMOTE_USER": "test_queued_event"}, + ) + + assert response.status_code == 200 + assert response.json == { + "queued_events": [ + { + "created_at": self.default_time, + "uri": "s3://bucket/key", + "dag_id": "dag", + } + ], + "total_entries": 1, + } + + def test_should_respond_404(self): + dataset_uri = "not_exists" + + response = self.client.get( + f"/api/v1/datasets/queuedEvent/{dataset_uri}", + environ_overrides={"REMOTE_USER": "test_queued_event"}, + ) + + assert response.status_code == 404 + assert { + "detail": "Queue event with asset uri: `not_exists` was not found", + "status": 404, + "title": "Queue event not found", + "type": EXCEPTIONS_LINK_MAP[404], + } == response.json + + +class TestDeleteDatasetQueuedEvents(TestQueuedEventEndpoint): + def test_delete_should_respond_204(self, session, create_dummy_dag): + dag, _ = create_dummy_dag() + dag_id = dag.dag_id + dataset_id = self._create_asset(session).id + self._create_asset_dag_run_queues(dag_id, dataset_id, session) + dataset_uri = "s3://bucket/key" + + response = self.client.delete( + f"/api/v1/datasets/queuedEvent/{dataset_uri}", + environ_overrides={"REMOTE_USER": "test_queued_event"}, + ) + + assert response.status_code == 204 + conn = session.query(AssetDagRunQueue).all() + assert len(conn) == 0 + _check_last_log(session, dag_id=None, event="api.delete_dataset_queued_events", execution_date=None) + + def test_should_respond_404(self): + dataset_uri = "not_exists" + + response = self.client.delete( + f"/api/v1/datasets/queuedEvent/{dataset_uri}", + environ_overrides={"REMOTE_USER": "test_queued_event"}, + ) + + assert response.status_code == 404 + assert { + "detail": "Queue event with asset uri: `not_exists` was not found", + "status": 404, + "title": "Queue event not found", + "type": EXCEPTIONS_LINK_MAP[404], + } == response.json diff --git a/tests/providers/fab/auth_manager/api_endpoints/test_event_log_endpoint.py b/tests/providers/fab/auth_manager/api_endpoints/test_event_log_endpoint.py new file mode 100644 index 0000000000000..acf3ca62684a1 --- /dev/null +++ b/tests/providers/fab/auth_manager/api_endpoints/test_event_log_endpoint.py @@ -0,0 +1,151 @@ +# 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 pytest + +from airflow.models import Log +from airflow.security import permissions +from airflow.utils import timezone +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_user, delete_user +from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS +from tests.test_utils.db import clear_db_logs + +pytestmark = [ + pytest.mark.db_test, + pytest.mark.skip_if_database_isolation_mode, + pytest.mark.skipif(not AIRFLOW_V_3_0_PLUS, reason="Test requires Airflow 3.0+"), +] + + +@pytest.fixture(scope="module") +def configured_app(minimal_app_for_auth_api): + app = minimal_app_for_auth_api + create_user( + app, + username="test_granular", + role_name="TestGranular", + permissions=[(permissions.ACTION_CAN_READ, permissions.RESOURCE_AUDIT_LOG)], + ) + app.appbuilder.sm.sync_perm_for_dag( + "TEST_DAG_ID_1", + access_control={"TestGranular": [permissions.ACTION_CAN_READ]}, + ) + app.appbuilder.sm.sync_perm_for_dag( + "TEST_DAG_ID_2", + access_control={"TestGranular": [permissions.ACTION_CAN_READ]}, + ) + + yield app + + delete_user(app, username="test_granular") + + +@pytest.fixture +def task_instance(session, create_task_instance, request): + return create_task_instance( + session=session, + dag_id="TEST_DAG_ID", + task_id="TEST_TASK_ID", + run_id="TEST_RUN_ID", + execution_date=request.instance.default_time, + ) + + +@pytest.fixture +def create_log_model(create_task_instance, task_instance, session, request): + def maker(event, when, **kwargs): + log_model = Log( + event=event, + task_instance=task_instance, + **kwargs, + ) + log_model.dttm = when + + session.add(log_model) + session.flush() + return log_model + + return maker + + +class TestEventLogEndpoint: + @pytest.fixture(autouse=True) + def setup_attrs(self, configured_app) -> None: + self.app = configured_app + self.client = self.app.test_client() # type:ignore + clear_db_logs() + self.default_time = timezone.parse("2020-06-10T20:00:00+00:00") + self.default_time_2 = timezone.parse("2020-06-11T07:00:00+00:00") + + def teardown_method(self) -> None: + clear_db_logs() + + +class TestGetEventLogs(TestEventLogEndpoint): + def test_should_filter_eventlogs_by_allowed_attributes(self, create_log_model, session): + eventlog1 = create_log_model( + event="TEST_EVENT_1", + dag_id="TEST_DAG_ID_1", + task_id="TEST_TASK_ID_1", + owner="TEST_OWNER_1", + when=self.default_time, + ) + eventlog2 = create_log_model( + event="TEST_EVENT_2", + dag_id="TEST_DAG_ID_2", + task_id="TEST_TASK_ID_2", + owner="TEST_OWNER_2", + when=self.default_time_2, + ) + session.add_all([eventlog1, eventlog2]) + session.commit() + for attr in ["dag_id", "task_id", "owner", "event"]: + attr_value = f"TEST_{attr}_1".upper() + response = self.client.get( + f"/api/v1/eventLogs?{attr}={attr_value}", environ_overrides={"REMOTE_USER": "test_granular"} + ) + assert response.status_code == 200 + assert response.json["total_entries"] == 1 + assert len(response.json["event_logs"]) == 1 + assert response.json["event_logs"][0][attr] == attr_value + + def test_should_filter_eventlogs_by_included_events(self, create_log_model): + for event in ["TEST_EVENT_1", "TEST_EVENT_2", "cli_scheduler"]: + create_log_model(event=event, when=self.default_time) + response = self.client.get( + "/api/v1/eventLogs?included_events=TEST_EVENT_1,TEST_EVENT_2", + environ_overrides={"REMOTE_USER": "test_granular"}, + ) + assert response.status_code == 200 + response_data = response.json + assert len(response_data["event_logs"]) == 2 + assert response_data["total_entries"] == 2 + assert {"TEST_EVENT_1", "TEST_EVENT_2"} == {x["event"] for x in response_data["event_logs"]} + + def test_should_filter_eventlogs_by_excluded_events(self, create_log_model): + for event in ["TEST_EVENT_1", "TEST_EVENT_2", "cli_scheduler"]: + create_log_model(event=event, when=self.default_time) + response = self.client.get( + "/api/v1/eventLogs?excluded_events=TEST_EVENT_1,TEST_EVENT_2", + environ_overrides={"REMOTE_USER": "test_granular"}, + ) + assert response.status_code == 200 + response_data = response.json + assert len(response_data["event_logs"]) == 1 + assert response_data["total_entries"] == 1 + assert {"cli_scheduler"} == {x["event"] for x in response_data["event_logs"]} diff --git a/tests/providers/fab/auth_manager/api_endpoints/test_import_error_endpoint.py b/tests/providers/fab/auth_manager/api_endpoints/test_import_error_endpoint.py new file mode 100644 index 0000000000000..a2fa1d028a3f2 --- /dev/null +++ b/tests/providers/fab/auth_manager/api_endpoints/test_import_error_endpoint.py @@ -0,0 +1,221 @@ +# 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 pytest + +from airflow.models.dag import DagModel +from airflow.security import permissions +from airflow.utils import timezone +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_user, delete_user +from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS, ParseImportError +from tests.test_utils.db import clear_db_dags, clear_db_import_errors +from tests.test_utils.permissions import _resource_name + +pytestmark = [ + pytest.mark.db_test, + pytest.mark.skip_if_database_isolation_mode, + pytest.mark.skipif(not AIRFLOW_V_3_0_PLUS, reason="Test requires Airflow 3.0+"), +] + +TEST_DAG_IDS = ["test_dag", "test_dag2"] + + +@pytest.fixture(scope="module") +def configured_app(minimal_app_for_auth_api): + app = minimal_app_for_auth_api + create_user( + app, + username="test_single_dag", + role_name="TestSingleDAG", + permissions=[(permissions.ACTION_CAN_READ, permissions.RESOURCE_IMPORT_ERROR)], + ) + # For some reason, DAG level permissions are not synced when in the above list of perms, + # so do it manually here: + app.appbuilder.sm.bulk_sync_roles( + [ + { + "role": "TestSingleDAG", + "perms": [ + ( + permissions.ACTION_CAN_READ, + _resource_name(TEST_DAG_IDS[0], permissions.RESOURCE_DAG), + ) + ], + } + ] + ) + + yield app + + delete_user(app, username="test_single_dag") + + +class TestBaseImportError: + timestamp = "2020-06-10T12:00" + + @pytest.fixture(autouse=True) + def setup_attrs(self, configured_app) -> None: + self.app = configured_app + self.client = self.app.test_client() # type:ignore + + clear_db_import_errors() + clear_db_dags() + + def teardown_method(self) -> None: + clear_db_import_errors() + clear_db_dags() + + @staticmethod + def _normalize_import_errors(import_errors): + for i, import_error in enumerate(import_errors, 1): + import_error["import_error_id"] = i + + +class TestGetImportErrorEndpoint(TestBaseImportError): + def test_should_raise_403_forbidden_without_dag_read(self, session): + import_error = ParseImportError( + filename="Lorem_ipsum.py", + stacktrace="Lorem ipsum", + timestamp=timezone.parse(self.timestamp, timezone="UTC"), + ) + session.add(import_error) + session.commit() + + response = self.client.get( + f"/api/v1/importErrors/{import_error.id}", environ_overrides={"REMOTE_USER": "test_single_dag"} + ) + + assert response.status_code == 403 + + def test_should_return_200_with_single_dag_read(self, session): + dag_model = DagModel(dag_id=TEST_DAG_IDS[0], fileloc="Lorem_ipsum.py") + session.add(dag_model) + import_error = ParseImportError( + filename="Lorem_ipsum.py", + stacktrace="Lorem ipsum", + timestamp=timezone.parse(self.timestamp, timezone="UTC"), + ) + session.add(import_error) + session.commit() + + response = self.client.get( + f"/api/v1/importErrors/{import_error.id}", environ_overrides={"REMOTE_USER": "test_single_dag"} + ) + + assert response.status_code == 200 + response_data = response.json + response_data["import_error_id"] = 1 + assert { + "filename": "Lorem_ipsum.py", + "import_error_id": 1, + "stack_trace": "Lorem ipsum", + "timestamp": "2020-06-10T12:00:00+00:00", + } == response_data + + def test_should_return_200_redacted_with_single_dag_read_in_dagfile(self, session): + for dag_id in TEST_DAG_IDS: + dag_model = DagModel(dag_id=dag_id, fileloc="Lorem_ipsum.py") + session.add(dag_model) + import_error = ParseImportError( + filename="Lorem_ipsum.py", + stacktrace="Lorem ipsum", + timestamp=timezone.parse(self.timestamp, timezone="UTC"), + ) + session.add(import_error) + session.commit() + + response = self.client.get( + f"/api/v1/importErrors/{import_error.id}", environ_overrides={"REMOTE_USER": "test_single_dag"} + ) + + assert response.status_code == 200 + response_data = response.json + response_data["import_error_id"] = 1 + assert { + "filename": "Lorem_ipsum.py", + "import_error_id": 1, + "stack_trace": "REDACTED - you do not have read permission on all DAGs in the file", + "timestamp": "2020-06-10T12:00:00+00:00", + } == response_data + + +class TestGetImportErrorsEndpoint(TestBaseImportError): + def test_get_import_errors_single_dag(self, session): + for dag_id in TEST_DAG_IDS: + fake_filename = f"/tmp/{dag_id}.py" + dag_model = DagModel(dag_id=dag_id, fileloc=fake_filename) + session.add(dag_model) + importerror = ParseImportError( + filename=fake_filename, + stacktrace="Lorem ipsum", + timestamp=timezone.parse(self.timestamp, timezone="UTC"), + ) + session.add(importerror) + session.commit() + + response = self.client.get( + "/api/v1/importErrors", environ_overrides={"REMOTE_USER": "test_single_dag"} + ) + + assert response.status_code == 200 + response_data = response.json + self._normalize_import_errors(response_data["import_errors"]) + assert { + "import_errors": [ + { + "filename": "/tmp/test_dag.py", + "import_error_id": 1, + "stack_trace": "Lorem ipsum", + "timestamp": "2020-06-10T12:00:00+00:00", + }, + ], + "total_entries": 1, + } == response_data + + def test_get_import_errors_single_dag_in_dagfile(self, session): + for dag_id in TEST_DAG_IDS: + fake_filename = "/tmp/all_in_one.py" + dag_model = DagModel(dag_id=dag_id, fileloc=fake_filename) + session.add(dag_model) + + importerror = ParseImportError( + filename="/tmp/all_in_one.py", + stacktrace="Lorem ipsum", + timestamp=timezone.parse(self.timestamp, timezone="UTC"), + ) + session.add(importerror) + session.commit() + + response = self.client.get( + "/api/v1/importErrors", environ_overrides={"REMOTE_USER": "test_single_dag"} + ) + + assert response.status_code == 200 + response_data = response.json + self._normalize_import_errors(response_data["import_errors"]) + assert { + "import_errors": [ + { + "filename": "/tmp/all_in_one.py", + "import_error_id": 1, + "stack_trace": "REDACTED - you do not have read permission on all DAGs in the file", + "timestamp": "2020-06-10T12:00:00+00:00", + }, + ], + "total_entries": 1, + } == response_data diff --git a/tests/providers/fab/auth_manager/api_endpoints/test_role_and_permission_endpoint.py b/tests/providers/fab/auth_manager/api_endpoints/test_role_and_permission_endpoint.py index 30cfaeb227903..413a49a9d86a1 100644 --- a/tests/providers/fab/auth_manager/api_endpoints/test_role_and_permission_endpoint.py +++ b/tests/providers/fab/auth_manager/api_endpoints/test_role_and_permission_endpoint.py @@ -19,6 +19,13 @@ import pytest from airflow.api_connexion.exceptions import EXCEPTIONS_LINK_MAP +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import ( + create_role, + create_user, + delete_role, + delete_user, +) +from tests.test_utils.api_connexion_utils import assert_401 from tests.test_utils.compat import ignore_provider_compatibility_error with ignore_provider_compatibility_error("2.9.0+", __file__): @@ -27,13 +34,6 @@ from airflow.security import permissions -from tests.test_utils.api_connexion_utils import ( - assert_401, - create_role, - create_user, - delete_role, - delete_user, -) pytestmark = pytest.mark.db_test @@ -42,7 +42,7 @@ def configured_app(minimal_app_for_auth_api): app = minimal_app_for_auth_api create_user( - app, # type: ignore + app, username="test", role_name="Test", permissions=[ @@ -53,11 +53,11 @@ def configured_app(minimal_app_for_auth_api): (permissions.ACTION_CAN_READ, permissions.RESOURCE_ACTION), ], ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name="TestNoPermissions") yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") class TestRoleEndpoint: diff --git a/tests/api_connexion/schemas/test_role_and_permission_schema.py b/tests/providers/fab/auth_manager/api_endpoints/test_role_and_permission_schema.py similarity index 85% rename from tests/api_connexion/schemas/test_role_and_permission_schema.py rename to tests/providers/fab/auth_manager/api_endpoints/test_role_and_permission_schema.py index f2967d519794c..4a2f0068e5e4a 100644 --- a/tests/api_connexion/schemas/test_role_and_permission_schema.py +++ b/tests/providers/fab/auth_manager/api_endpoints/test_role_and_permission_schema.py @@ -31,19 +31,19 @@ class TestRoleCollectionItemSchema: @pytest.fixture(scope="class") - def role(self, minimal_app_for_api): + def role(self, minimal_app_for_auth_api): yield create_role( - minimal_app_for_api, # type: ignore + minimal_app_for_auth_api, # type: ignore name="Test", permissions=[ (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_CONNECTION), ], ) - delete_role(minimal_app_for_api, "Test") + delete_role(minimal_app_for_auth_api, "Test") @pytest.fixture(autouse=True) - def _set_attrs(self, minimal_app_for_api, role): - self.app = minimal_app_for_api + def _set_attrs(self, minimal_app_for_auth_api, role): + self.app = minimal_app_for_auth_api self.role = role def test_serialize(self): @@ -67,26 +67,26 @@ def test_deserialize(self): class TestRoleCollectionSchema: @pytest.fixture(scope="class") - def role1(self, minimal_app_for_api): + def role1(self, minimal_app_for_auth_api): yield create_role( - minimal_app_for_api, # type: ignore + minimal_app_for_auth_api, # type: ignore name="Test1", permissions=[ (permissions.ACTION_CAN_CREATE, permissions.RESOURCE_CONNECTION), ], ) - delete_role(minimal_app_for_api, "Test1") + delete_role(minimal_app_for_auth_api, "Test1") @pytest.fixture(scope="class") - def role2(self, minimal_app_for_api): + def role2(self, minimal_app_for_auth_api): yield create_role( - minimal_app_for_api, # type: ignore + minimal_app_for_auth_api, # type: ignore name="Test2", permissions=[ (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG), ], ) - delete_role(minimal_app_for_api, "Test2") + delete_role(minimal_app_for_auth_api, "Test2") def test_serialize(self, role1, role2): instance = RoleCollection([role1, role2], total_entries=2) diff --git a/tests/providers/fab/auth_manager/api_endpoints/test_task_instance_endpoint.py b/tests/providers/fab/auth_manager/api_endpoints/test_task_instance_endpoint.py new file mode 100644 index 0000000000000..69b3c221eae93 --- /dev/null +++ b/tests/providers/fab/auth_manager/api_endpoints/test_task_instance_endpoint.py @@ -0,0 +1,427 @@ +# 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 datetime as dt +import urllib + +import pytest + +from airflow.api_connexion.exceptions import EXCEPTIONS_LINK_MAP +from airflow.models import DagRun, TaskInstance +from airflow.security import permissions +from airflow.utils.session import provide_session +from airflow.utils.state import State +from airflow.utils.timezone import datetime +from airflow.utils.types import DagRunType +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import ( + create_user, + delete_roles, + delete_user, +) +from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS +from tests.test_utils.db import clear_db_runs, clear_db_sla_miss, clear_rendered_ti_fields + +pytestmark = [ + pytest.mark.db_test, + pytest.mark.skip_if_database_isolation_mode, + pytest.mark.skipif(not AIRFLOW_V_3_0_PLUS, reason="Test requires Airflow 3.0+"), +] + +DEFAULT_DATETIME_1 = datetime(2020, 1, 1) +DEFAULT_DATETIME_STR_1 = "2020-01-01T00:00:00+00:00" +DEFAULT_DATETIME_STR_2 = "2020-01-02T00:00:00+00:00" + +QUOTED_DEFAULT_DATETIME_STR_1 = urllib.parse.quote(DEFAULT_DATETIME_STR_1) +QUOTED_DEFAULT_DATETIME_STR_2 = urllib.parse.quote(DEFAULT_DATETIME_STR_2) + + +@pytest.fixture(scope="module") +def configured_app(minimal_app_for_auth_api): + app = minimal_app_for_auth_api + create_user( + app, + username="test_dag_read_only", + role_name="TestDagReadOnly", + permissions=[ + (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_TASK_INSTANCE), + (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_TASK_INSTANCE), + ], + ) + create_user( + app, + username="test_task_read_only", + role_name="TestTaskReadOnly", + permissions=[ + (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG), + (permissions.ACTION_CAN_EDIT, permissions.RESOURCE_DAG), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_TASK_INSTANCE), + ], + ) + create_user( + app, + username="test_read_only_one_dag", + role_name="TestReadOnlyOneDag", + permissions=[ + (permissions.ACTION_CAN_READ, permissions.RESOURCE_DAG_RUN), + (permissions.ACTION_CAN_READ, permissions.RESOURCE_TASK_INSTANCE), + ], + ) + # For some reason, "DAG:example_python_operator" is not synced when in the above list of perms, + # so do it manually here: + app.appbuilder.sm.bulk_sync_roles( + [ + { + "role": "TestReadOnlyOneDag", + "perms": [(permissions.ACTION_CAN_READ, "DAG:example_python_operator")], + } + ] + ) + + yield app + + delete_user(app, username="test_dag_read_only") + delete_user(app, username="test_task_read_only") + delete_user(app, username="test_read_only_one_dag") + delete_roles(app) + + +class TestTaskInstanceEndpoint: + @pytest.fixture(autouse=True) + def setup_attrs(self, configured_app, dagbag) -> None: + self.default_time = DEFAULT_DATETIME_1 + self.ti_init = { + "execution_date": self.default_time, + "state": State.RUNNING, + } + self.ti_extras = { + "start_date": self.default_time + dt.timedelta(days=1), + "end_date": self.default_time + dt.timedelta(days=2), + "pid": 100, + "duration": 10000, + "pool": "default_pool", + "queue": "default_queue", + "job_id": 0, + } + self.app = configured_app + self.client = self.app.test_client() # type:ignore + clear_db_runs() + clear_db_sla_miss() + clear_rendered_ti_fields() + self.dagbag = dagbag + + def create_task_instances( + self, + session, + dag_id: str = "example_python_operator", + update_extras: bool = True, + task_instances=None, + dag_run_state=State.RUNNING, + with_ti_history=False, + ): + """Method to create task instances using kwargs and default arguments""" + + dag = self.dagbag.get_dag(dag_id) + tasks = dag.tasks + counter = len(tasks) + if task_instances is not None: + counter = min(len(task_instances), counter) + + run_id = "TEST_DAG_RUN_ID" + execution_date = self.ti_init.pop("execution_date", self.default_time) + dr = None + + tis = [] + for i in range(counter): + if task_instances is None: + pass + elif update_extras: + self.ti_extras.update(task_instances[i]) + else: + self.ti_init.update(task_instances[i]) + + if "execution_date" in self.ti_init: + run_id = f"TEST_DAG_RUN_ID_{i}" + execution_date = self.ti_init.pop("execution_date") + dr = None + + if not dr: + dr = DagRun( + run_id=run_id, + dag_id=dag_id, + execution_date=execution_date, + run_type=DagRunType.MANUAL, + state=dag_run_state, + ) + session.add(dr) + ti = TaskInstance(task=tasks[i], **self.ti_init) + session.add(ti) + ti.dag_run = dr + ti.note = "placeholder-note" + + for key, value in self.ti_extras.items(): + setattr(ti, key, value) + tis.append(ti) + + session.commit() + if with_ti_history: + for ti in tis: + ti.try_number = 1 + session.merge(ti) + session.commit() + dag.clear() + for ti in tis: + ti.try_number = 2 + ti.queue = "default_queue" + session.merge(ti) + session.commit() + return tis + + +class TestGetTaskInstance(TestTaskInstanceEndpoint): + def setup_method(self): + clear_db_runs() + + def teardown_method(self): + clear_db_runs() + + @pytest.mark.parametrize("username", ["test_dag_read_only", "test_task_read_only"]) + @provide_session + def test_should_respond_200(self, username, session): + self.create_task_instances(session) + # Update ti and set operator to None to + # test that operator field is nullable. + # This prevents issue when users upgrade to 2.0+ + # from 1.10.x + # https://github.com/apache/airflow/issues/14421 + session.query(TaskInstance).update({TaskInstance.operator: None}, synchronize_session="fetch") + session.commit() + response = self.client.get( + "/api/v1/dags/example_python_operator/dagRuns/TEST_DAG_RUN_ID/taskInstances/print_the_context", + environ_overrides={"REMOTE_USER": username}, + ) + assert response.status_code == 200 + + +class TestGetTaskInstances(TestTaskInstanceEndpoint): + @pytest.mark.parametrize( + "task_instances, user, expected_ti", + [ + pytest.param( + { + "example_python_operator": 2, + "example_skip_dag": 1, + }, + "test_read_only_one_dag", + 2, + ), + pytest.param( + { + "example_python_operator": 1, + "example_skip_dag": 2, + }, + "test_read_only_one_dag", + 1, + ), + ], + ) + def test_return_TI_only_from_readable_dags(self, task_instances, user, expected_ti, session): + for dag_id in task_instances: + self.create_task_instances( + session, + task_instances=[ + {"execution_date": DEFAULT_DATETIME_1 + dt.timedelta(days=i)} + for i in range(task_instances[dag_id]) + ], + dag_id=dag_id, + ) + response = self.client.get( + "/api/v1/dags/~/dagRuns/~/taskInstances", environ_overrides={"REMOTE_USER": user} + ) + assert response.status_code == 200 + assert response.json["total_entries"] == expected_ti + assert len(response.json["task_instances"]) == expected_ti + + +class TestGetTaskInstancesBatch(TestTaskInstanceEndpoint): + @pytest.mark.parametrize( + "task_instances, update_extras, payload, expected_ti_count, username", + [ + pytest.param( + [ + {"pool": "test_pool_1"}, + {"pool": "test_pool_2"}, + {"pool": "test_pool_3"}, + ], + True, + {"pool": ["test_pool_1", "test_pool_2"]}, + 2, + "test_dag_read_only", + id="test pool filter", + ), + pytest.param( + [ + {"state": State.RUNNING}, + {"state": State.QUEUED}, + {"state": State.SUCCESS}, + {"state": State.NONE}, + ], + False, + {"state": ["running", "queued", "none"]}, + 3, + "test_task_read_only", + id="test state filter", + ), + pytest.param( + [ + {"state": State.NONE}, + {"state": State.NONE}, + {"state": State.NONE}, + {"state": State.NONE}, + ], + False, + {}, + 4, + "test_task_read_only", + id="test dag with null states", + ), + pytest.param( + [ + {"end_date": DEFAULT_DATETIME_1}, + {"end_date": DEFAULT_DATETIME_1 + dt.timedelta(days=1)}, + {"end_date": DEFAULT_DATETIME_1 + dt.timedelta(days=2)}, + ], + True, + { + "end_date_gte": DEFAULT_DATETIME_STR_1, + "end_date_lte": DEFAULT_DATETIME_STR_2, + }, + 2, + "test_task_read_only", + id="test end date filter", + ), + pytest.param( + [ + {"start_date": DEFAULT_DATETIME_1}, + {"start_date": DEFAULT_DATETIME_1 + dt.timedelta(days=1)}, + {"start_date": DEFAULT_DATETIME_1 + dt.timedelta(days=2)}, + ], + True, + { + "start_date_gte": DEFAULT_DATETIME_STR_1, + "start_date_lte": DEFAULT_DATETIME_STR_2, + }, + 2, + "test_dag_read_only", + id="test start date filter", + ), + ], + ) + def test_should_respond_200( + self, task_instances, update_extras, payload, expected_ti_count, username, session + ): + self.create_task_instances( + session, + update_extras=update_extras, + task_instances=task_instances, + ) + response = self.client.post( + "/api/v1/dags/~/dagRuns/~/taskInstances/list", + environ_overrides={"REMOTE_USER": username}, + json=payload, + ) + assert response.status_code == 200, response.json + assert expected_ti_count == response.json["total_entries"] + assert expected_ti_count == len(response.json["task_instances"]) + + def test_returns_403_forbidden_when_user_has_access_to_only_some_dags(self, session): + self.create_task_instances(session=session) + self.create_task_instances(session=session, dag_id="example_skip_dag") + payload = {"dag_ids": ["example_python_operator", "example_skip_dag"]} + + response = self.client.post( + "/api/v1/dags/~/dagRuns/~/taskInstances/list", + environ_overrides={"REMOTE_USER": "test_read_only_one_dag"}, + json=payload, + ) + assert response.status_code == 403 + assert response.json == { + "detail": "User not allowed to access some of these DAGs: ['example_python_operator', 'example_skip_dag']", + "status": 403, + "title": "Forbidden", + "type": EXCEPTIONS_LINK_MAP[403], + } + + +class TestPostSetTaskInstanceState(TestTaskInstanceEndpoint): + @pytest.mark.parametrize("username", ["test_dag_read_only", "test_task_read_only"]) + def test_should_raise_403_forbidden(self, username): + response = self.client.post( + "/api/v1/dags/example_python_operator/updateTaskInstancesState", + environ_overrides={"REMOTE_USER": username}, + json={ + "dry_run": True, + "task_id": "print_the_context", + "execution_date": DEFAULT_DATETIME_1.isoformat(), + "include_upstream": True, + "include_downstream": True, + "include_future": True, + "include_past": True, + "new_state": "failed", + }, + ) + assert response.status_code == 403 + + +class TestPatchTaskInstance(TestTaskInstanceEndpoint): + ENDPOINT_URL = ( + "/api/v1/dags/example_python_operator/dagRuns/TEST_DAG_RUN_ID/taskInstances/print_the_context" + ) + + @pytest.mark.parametrize("username", ["test_dag_read_only", "test_task_read_only"]) + def test_should_raise_403_forbidden(self, username): + response = self.client.patch( + self.ENDPOINT_URL, + environ_overrides={"REMOTE_USER": username}, + json={ + "dry_run": True, + "new_state": "failed", + }, + ) + assert response.status_code == 403 + + +class TestGetTaskInstanceTry(TestTaskInstanceEndpoint): + def setup_method(self): + clear_db_runs() + + def teardown_method(self): + clear_db_runs() + + @pytest.mark.parametrize("username", ["test_dag_read_only", "test_task_read_only"]) + @provide_session + def test_should_respond_200(self, username, session): + self.create_task_instances(session, task_instances=[{"state": State.SUCCESS}], with_ti_history=True) + + response = self.client.get( + "/api/v1/dags/example_python_operator/dagRuns/TEST_DAG_RUN_ID/taskInstances/print_the_context/tries/1", + environ_overrides={"REMOTE_USER": username}, + ) + assert response.status_code == 200 diff --git a/tests/providers/fab/auth_manager/api_endpoints/test_user_endpoint.py b/tests/providers/fab/auth_manager/api_endpoints/test_user_endpoint.py index bc400c8a43fad..7f2c885bab52c 100644 --- a/tests/providers/fab/auth_manager/api_endpoints/test_user_endpoint.py +++ b/tests/providers/fab/auth_manager/api_endpoints/test_user_endpoint.py @@ -30,7 +30,12 @@ with ignore_provider_compatibility_error("2.9.0+", __file__): from airflow.providers.fab.auth_manager.models import User -from tests.test_utils.api_connexion_utils import assert_401, create_user, delete_role, delete_user +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import ( + create_user, + delete_role, + delete_user, +) +from tests.test_utils.api_connexion_utils import assert_401 from tests.test_utils.config import conf_vars pytestmark = pytest.mark.db_test @@ -43,7 +48,7 @@ def configured_app(minimal_app_for_auth_api): app = minimal_app_for_auth_api create_user( - app, # type: ignore + app, username="test", role_name="Test", permissions=[ @@ -53,12 +58,12 @@ def configured_app(minimal_app_for_auth_api): (permissions.ACTION_CAN_READ, permissions.RESOURCE_USER), ], ) - create_user(app, username="test_no_permissions", role_name="TestNoPermissions") # type: ignore + create_user(app, username="test_no_permissions", role_name="TestNoPermissions") yield app - delete_user(app, username="test") # type: ignore - delete_user(app, username="test_no_permissions") # type: ignore + delete_user(app, username="test") + delete_user(app, username="test_no_permissions") delete_role(app, name="TestNoPermissions") diff --git a/tests/providers/fab/auth_manager/api_endpoints/test_user_schema.py b/tests/providers/fab/auth_manager/api_endpoints/test_user_schema.py index 265407622e269..f3399de6a9775 100644 --- a/tests/providers/fab/auth_manager/api_endpoints/test_user_schema.py +++ b/tests/providers/fab/auth_manager/api_endpoints/test_user_schema.py @@ -18,6 +18,7 @@ import pytest +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_role, delete_role from tests.test_utils.compat import ignore_provider_compatibility_error with ignore_provider_compatibility_error("2.9.0+", __file__): @@ -30,8 +31,6 @@ DEFAULT_TIME = "2021-01-09T13:59:56.336000+00:00" -from tests.test_utils.api_connexion_utils import create_role, delete_role # noqa: E402 - pytestmark = pytest.mark.db_test diff --git a/tests/providers/fab/auth_manager/api_endpoints/test_variable_endpoint.py b/tests/providers/fab/auth_manager/api_endpoints/test_variable_endpoint.py new file mode 100644 index 0000000000000..a8e71e1a82466 --- /dev/null +++ b/tests/providers/fab/auth_manager/api_endpoints/test_variable_endpoint.py @@ -0,0 +1,88 @@ +# 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 pytest + +from airflow.models import Variable +from airflow.security import permissions +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_user, delete_user +from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS +from tests.test_utils.db import clear_db_variables + +pytestmark = [ + pytest.mark.db_test, + pytest.mark.skip_if_database_isolation_mode, + pytest.mark.skipif(not AIRFLOW_V_3_0_PLUS, reason="Test requires Airflow 3.0+"), +] + + +@pytest.fixture(scope="module") +def configured_app(minimal_app_for_auth_api): + app = minimal_app_for_auth_api + + create_user( + app, + username="test_read_only", + role_name="TestReadOnly", + permissions=[ + (permissions.ACTION_CAN_READ, permissions.RESOURCE_VARIABLE), + ], + ) + create_user( + app, + username="test_delete_only", + role_name="TestDeleteOnly", + permissions=[ + (permissions.ACTION_CAN_DELETE, permissions.RESOURCE_VARIABLE), + ], + ) + + yield app + + delete_user(app, username="test_read_only") + delete_user(app, username="test_delete_only") + + +class TestVariableEndpoint: + @pytest.fixture(autouse=True) + def setup_method(self, configured_app) -> None: + self.app = configured_app + self.client = self.app.test_client() # type:ignore + clear_db_variables() + + def teardown_method(self) -> None: + clear_db_variables() + + +class TestGetVariable(TestVariableEndpoint): + @pytest.mark.parametrize( + "user, expected_status_code", + [ + ("test_read_only", 200), + ("test_delete_only", 403), + ], + ) + def test_read_variable(self, user, expected_status_code): + expected_value = '{"foo": 1}' + Variable.set("TEST_VARIABLE_KEY", expected_value) + response = self.client.get( + "/api/v1/variables/TEST_VARIABLE_KEY", environ_overrides={"REMOTE_USER": user} + ) + assert response.status_code == expected_status_code + if expected_status_code == 200: + assert response.json == {"key": "TEST_VARIABLE_KEY", "value": expected_value, "description": None} diff --git a/tests/providers/fab/auth_manager/api_endpoints/test_xcom_endpoint.py b/tests/providers/fab/auth_manager/api_endpoints/test_xcom_endpoint.py new file mode 100644 index 0000000000000..01336f9957c6d --- /dev/null +++ b/tests/providers/fab/auth_manager/api_endpoints/test_xcom_endpoint.py @@ -0,0 +1,230 @@ +# 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 + +from datetime import timedelta + +import pytest + +from airflow.models.dag import DagModel +from airflow.models.dagrun import DagRun +from airflow.models.taskinstance import TaskInstance +from airflow.models.xcom import BaseXCom, XCom +from airflow.operators.empty import EmptyOperator +from airflow.security import permissions +from airflow.utils.dates import parse_execution_date +from airflow.utils.session import create_session +from airflow.utils.types import DagRunType +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_user, delete_user +from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS +from tests.test_utils.db import clear_db_dags, clear_db_runs, clear_db_xcom + +pytestmark = [ + pytest.mark.db_test, + pytest.mark.skip_if_database_isolation_mode, + pytest.mark.skipif(not AIRFLOW_V_3_0_PLUS, reason="Test requires Airflow 3.0+"), +] + + +class CustomXCom(BaseXCom): + @classmethod + def deserialize_value(cls, xcom: XCom): + return f"real deserialized {super().deserialize_value(xcom)}" + + def orm_deserialize_value(self): + return f"orm deserialized {super().orm_deserialize_value()}" + + +@pytest.fixture(scope="module") +def configured_app(minimal_app_for_auth_api): + app = minimal_app_for_auth_api + + create_user( + app, + username="test_granular_permissions", + role_name="TestGranularDag", + permissions=[ + (permissions.ACTION_CAN_READ, permissions.RESOURCE_XCOM), + ], + ) + app.appbuilder.sm.sync_perm_for_dag( + "test-dag-id-1", + access_control={"TestGranularDag": [permissions.ACTION_CAN_EDIT, permissions.ACTION_CAN_READ]}, + ) + + yield app + + delete_user(app, username="test_granular_permissions") + + +def _compare_xcom_collections(collection1: dict, collection_2: dict): + assert collection1.get("total_entries") == collection_2.get("total_entries") + + def sort_key(record): + return ( + record.get("dag_id"), + record.get("task_id"), + record.get("execution_date"), + record.get("map_index"), + record.get("key"), + ) + + assert sorted(collection1.get("xcom_entries", []), key=sort_key) == sorted( + collection_2.get("xcom_entries", []), key=sort_key + ) + + +class TestXComEndpoint: + @staticmethod + def clean_db(): + clear_db_dags() + clear_db_runs() + clear_db_xcom() + + @pytest.fixture(autouse=True) + def setup_attrs(self, configured_app) -> None: + """ + Setup For XCom endpoint TC + """ + self.app = configured_app + self.client = self.app.test_client() # type:ignore + # clear existing xcoms + self.clean_db() + + def teardown_method(self) -> None: + """ + Clear Hanging XComs + """ + self.clean_db() + + +class TestGetXComEntries(TestXComEndpoint): + def test_should_respond_200_with_tilde_and_granular_dag_access(self): + dag_id_1 = "test-dag-id-1" + task_id_1 = "test-task-id-1" + execution_date = "2005-04-02T00:00:00+00:00" + execution_date_parsed = parse_execution_date(execution_date) + dag_run_id_1 = DagRun.generate_run_id(DagRunType.MANUAL, execution_date_parsed) + self._create_xcom_entries(dag_id_1, dag_run_id_1, execution_date_parsed, task_id_1) + + dag_id_2 = "test-dag-id-2" + task_id_2 = "test-task-id-2" + run_id_2 = DagRun.generate_run_id(DagRunType.MANUAL, execution_date_parsed) + self._create_xcom_entries(dag_id_2, run_id_2, execution_date_parsed, task_id_2) + self._create_invalid_xcom_entries(execution_date_parsed) + response = self.client.get( + "/api/v1/dags/~/dagRuns/~/taskInstances/~/xcomEntries", + environ_overrides={"REMOTE_USER": "test_granular_permissions"}, + ) + + assert 200 == response.status_code + response_data = response.json + for xcom_entry in response_data["xcom_entries"]: + xcom_entry["timestamp"] = "TIMESTAMP" + _compare_xcom_collections( + response_data, + { + "xcom_entries": [ + { + "dag_id": dag_id_1, + "execution_date": execution_date, + "key": "test-xcom-key-1", + "task_id": task_id_1, + "timestamp": "TIMESTAMP", + "map_index": -1, + }, + { + "dag_id": dag_id_1, + "execution_date": execution_date, + "key": "test-xcom-key-2", + "task_id": task_id_1, + "timestamp": "TIMESTAMP", + "map_index": -1, + }, + ], + "total_entries": 2, + }, + ) + + def _create_xcom_entries(self, dag_id, run_id, execution_date, task_id, mapped_ti=False): + with create_session() as session: + dag = DagModel(dag_id=dag_id) + session.add(dag) + dagrun = DagRun( + dag_id=dag_id, + run_id=run_id, + execution_date=execution_date, + start_date=execution_date, + run_type=DagRunType.MANUAL, + ) + session.add(dagrun) + if mapped_ti: + for i in [0, 1]: + ti = TaskInstance(EmptyOperator(task_id=task_id), run_id=run_id, map_index=i) + ti.dag_id = dag_id + session.add(ti) + else: + ti = TaskInstance(EmptyOperator(task_id=task_id), run_id=run_id) + ti.dag_id = dag_id + session.add(ti) + + for i in [1, 2]: + if mapped_ti: + key = "test-xcom-key" + map_index = i - 1 + else: + key = f"test-xcom-key-{i}" + map_index = -1 + + XCom.set( + key=key, value="TEST", run_id=run_id, task_id=task_id, dag_id=dag_id, map_index=map_index + ) + + def _create_invalid_xcom_entries(self, execution_date): + """ + Invalid XCom entries to test join query + """ + with create_session() as session: + dag = DagModel(dag_id="invalid_dag") + session.add(dag) + dagrun = DagRun( + dag_id="invalid_dag", + run_id="invalid_run_id", + execution_date=execution_date + timedelta(days=1), + start_date=execution_date, + run_type=DagRunType.MANUAL, + ) + session.add(dagrun) + dagrun1 = DagRun( + dag_id="invalid_dag", + run_id="not_this_run_id", + execution_date=execution_date, + start_date=execution_date, + run_type=DagRunType.MANUAL, + ) + session.add(dagrun1) + ti = TaskInstance(EmptyOperator(task_id="invalid_task"), run_id="not_this_run_id") + ti.dag_id = "invalid_dag" + session.add(ti) + for i in [1, 2]: + XCom.set( + key=f"invalid-xcom-key-{i}", + value="TEST", + run_id="not_this_run_id", + task_id="invalid_task", + dag_id="invalid_dag", + ) diff --git a/tests/providers/fab/auth_manager/conftest.py b/tests/providers/fab/auth_manager/conftest.py index 22c29dd229fa1..a8fbe5fbdaaae 100644 --- a/tests/providers/fab/auth_manager/conftest.py +++ b/tests/providers/fab/auth_manager/conftest.py @@ -30,7 +30,10 @@ def minimal_app_for_auth_api(): "init_appbuilder", "init_api_auth", "init_api_auth_provider", + "init_api_connexion", "init_api_error_handlers", + "init_airflow_session_interface", + "init_appbuilder_views", ] ) def factory(): @@ -39,7 +42,11 @@ def factory(): ( "api", "auth_backends", - ): "tests.test_utils.remote_user_api_auth_backend,airflow.api.auth.backend.session" + ): "tests.providers.fab.auth_manager.api_endpoints.remote_user_api_auth_backend,airflow.api.auth.backend.session", + ( + "core", + "auth_manager", + ): "airflow.providers.fab.auth_manager.fab_auth_manager.FabAuthManager", } ): _app = app.create_app(testing=True, config={"WTF_CSRF_ENABLED": False}) # type:ignore @@ -58,3 +65,11 @@ def set_auth_role_public(request): yield app.config["AUTH_ROLE_PUBLIC"] = auto_role_public + + +@pytest.fixture(scope="module") +def dagbag(): + from airflow.models import DagBag + + DagBag(include_examples=True, read_dags_from_db=False).sync_to_db() + return DagBag(include_examples=True, read_dags_from_db=True) diff --git a/tests/providers/fab/auth_manager/test_security.py b/tests/providers/fab/auth_manager/test_security.py index 156b5cf626271..bebb52c256fc8 100644 --- a/tests/providers/fab/auth_manager/test_security.py +++ b/tests/providers/fab/auth_manager/test_security.py @@ -49,7 +49,7 @@ from airflow.www.auth import get_access_denied_message from airflow.www.extensions.init_auth_manager import get_auth_manager from airflow.www.utils import CustomSQLAInterface -from tests.test_utils.api_connexion_utils import ( +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import ( create_user, create_user_scope, delete_role, diff --git a/tests/providers/fab/auth_manager/views/test_permissions.py b/tests/providers/fab/auth_manager/views/test_permissions.py index 0b1073df287fa..f24d9b738343b 100644 --- a/tests/providers/fab/auth_manager/views/test_permissions.py +++ b/tests/providers/fab/auth_manager/views/test_permissions.py @@ -21,7 +21,7 @@ from airflow.security import permissions from airflow.www import app as application -from tests.test_utils.api_connexion_utils import create_user, delete_user +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_user, delete_user from tests.test_utils.compat import AIRFLOW_V_2_9_PLUS from tests.test_utils.www import client_with_login diff --git a/tests/providers/fab/auth_manager/views/test_roles_list.py b/tests/providers/fab/auth_manager/views/test_roles_list.py index 156f07df41209..8de63ad5ba88a 100644 --- a/tests/providers/fab/auth_manager/views/test_roles_list.py +++ b/tests/providers/fab/auth_manager/views/test_roles_list.py @@ -21,7 +21,7 @@ from airflow.security import permissions from airflow.www import app as application -from tests.test_utils.api_connexion_utils import create_user, delete_user +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_user, delete_user from tests.test_utils.compat import AIRFLOW_V_2_9_PLUS from tests.test_utils.www import client_with_login diff --git a/tests/providers/fab/auth_manager/views/test_user.py b/tests/providers/fab/auth_manager/views/test_user.py index 6660ab926d886..62b03a99e7c2c 100644 --- a/tests/providers/fab/auth_manager/views/test_user.py +++ b/tests/providers/fab/auth_manager/views/test_user.py @@ -21,7 +21,7 @@ from airflow.security import permissions from airflow.www import app as application -from tests.test_utils.api_connexion_utils import create_user, delete_user +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_user, delete_user from tests.test_utils.compat import AIRFLOW_V_2_9_PLUS from tests.test_utils.www import client_with_login diff --git a/tests/providers/fab/auth_manager/views/test_user_edit.py b/tests/providers/fab/auth_manager/views/test_user_edit.py index 65937b6f83d33..8099f67948183 100644 --- a/tests/providers/fab/auth_manager/views/test_user_edit.py +++ b/tests/providers/fab/auth_manager/views/test_user_edit.py @@ -21,7 +21,7 @@ from airflow.security import permissions from airflow.www import app as application -from tests.test_utils.api_connexion_utils import create_user, delete_user +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_user, delete_user from tests.test_utils.compat import AIRFLOW_V_2_9_PLUS from tests.test_utils.www import client_with_login diff --git a/tests/providers/fab/auth_manager/views/test_user_stats.py b/tests/providers/fab/auth_manager/views/test_user_stats.py index 8cb260fcf1ec4..ae09cf92252c6 100644 --- a/tests/providers/fab/auth_manager/views/test_user_stats.py +++ b/tests/providers/fab/auth_manager/views/test_user_stats.py @@ -21,7 +21,7 @@ from airflow.security import permissions from airflow.www import app as application -from tests.test_utils.api_connexion_utils import create_user, delete_user +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_user, delete_user from tests.test_utils.compat import AIRFLOW_V_2_9_PLUS from tests.test_utils.www import client_with_login diff --git a/tests/test_utils/api_connexion_utils.py b/tests/test_utils/api_connexion_utils.py index af746b2d55468..48869ee48078d 100644 --- a/tests/test_utils/api_connexion_utils.py +++ b/tests/test_utils/api_connexion_utils.py @@ -17,6 +17,7 @@ from __future__ import annotations from contextlib import contextmanager +from typing import TYPE_CHECKING from airflow.api_connexion.exceptions import EXCEPTIONS_LINK_MAP from tests.test_utils.compat import ignore_provider_compatibility_error @@ -24,6 +25,9 @@ with ignore_provider_compatibility_error("2.9.0+", __file__): from airflow.providers.fab.auth_manager.security_manager.override import EXISTING_ROLES +if TYPE_CHECKING: + from flask import Flask + @contextmanager def create_test_client(app, user_name, role_name, permissions): @@ -44,7 +48,11 @@ def create_user_scope(app, username, **kwargs): It will create a user and provide it for the fixture via YIELD (generator) then will tidy up once test is complete """ - test_user = create_user(app, username, **kwargs) + from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import ( + create_user as create_user_fab, + ) + + test_user = create_user_fab(app, username, **kwargs) try: yield test_user @@ -52,27 +60,20 @@ def create_user_scope(app, username, **kwargs): delete_user(app, username) -def create_user(app, username, role_name=None, email=None, permissions=None): - appbuilder = app.appbuilder - +def create_user(app: Flask, username: str, role_name: str | None): # Removes user and role so each test has isolated test data. delete_user(app, username) - role = None - if role_name: - delete_role(app, role_name) - role = create_role(app, role_name, permissions) - else: - role = [] - - return appbuilder.sm.add_user( - username=username, - first_name=username, - last_name=username, - email=email or f"{username}@example.org", - role=role, - password=username, + + users = app.config.get("SIMPLE_AUTH_MANAGER_USERS", []) + users.append( + { + "username": username, + "role": role_name, + } ) + app.config["SIMPLE_AUTH_MANAGER_USERS"] = users + def create_role(app, name, permissions=None): appbuilder = app.appbuilder @@ -87,14 +88,6 @@ def create_role(app, name, permissions=None): return role -def set_user_single_role(app, user, role_name): - role = create_role(app, role_name) - if role not in user.roles: - user.roles = [role] - app.appbuilder.sm.update_user(user) - user._perms = None - - def delete_role(app, name): if name not in EXISTING_ROLES: if app.appbuilder.sm.find_role(name): @@ -106,20 +99,11 @@ def delete_roles(app): delete_role(app, role.name) -def delete_user(app, username): - appbuilder = app.appbuilder - for user in appbuilder.sm.get_all_users(): - if user.username == username: - _ = [ - delete_role(app, role.name) for role in user.roles if role and role.name not in EXISTING_ROLES - ] - appbuilder.sm.del_register_user(user) - break - - -def delete_users(app): - for user in app.appbuilder.sm.get_all_users(): - delete_user(app, user.username) +def delete_user(app: Flask, username): + users = app.config.get("SIMPLE_AUTH_MANAGER_USERS", []) + users = [user for user in users if user["username"] != username] + + app.config["SIMPLE_AUTH_MANAGER_USERS"] = users def assert_401(response): diff --git a/tests/test_utils/remote_user_api_auth_backend.py b/tests/test_utils/remote_user_api_auth_backend.py index b7714e5192e6a..59df201e530e4 100644 --- a/tests/test_utils/remote_user_api_auth_backend.py +++ b/tests/test_utils/remote_user_api_auth_backend.py @@ -15,17 +15,15 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. -"""Default authentication backend - everything is allowed""" - from __future__ import annotations import logging from functools import wraps from typing import TYPE_CHECKING, Callable, TypeVar, cast -from flask import Response, request -from flask_login import login_user +from flask import Response, request, session +from airflow.auth.managers.simple.user import SimpleAuthManagerUser from airflow.utils.airflow_flask_app import get_airflow_app if TYPE_CHECKING: @@ -36,25 +34,15 @@ CLIENT_AUTH: tuple[str, str] | AuthBase | None = None -def init_app(_): - """Initializes authentication backend""" +def init_app(_): ... T = TypeVar("T", bound=Callable) -def _lookup_user(user_email_or_username: str): - security_manager = get_airflow_app().appbuilder.sm - user = security_manager.find_user(email=user_email_or_username) or security_manager.find_user( - username=user_email_or_username - ) - if not user: - return None - - if not user.is_active: - return None - - return user +def _lookup_user(username: str): + users = get_airflow_app().config.get("SIMPLE_AUTH_MANAGER_USERS", []) + return next((user for user in users if user["username"] == username), None) def requires_authentication(function: T): @@ -69,13 +57,13 @@ def decorated(*args, **kwargs): log.debug("Looking for user: %s", user_id) - user = _lookup_user(user_id) - if not user: + user_dict = _lookup_user(user_id) + if not user_dict: return Response("Forbidden", 403) - log.debug("Found user: %s", user) + log.debug("Found user: %s", user_dict) + session["user"] = SimpleAuthManagerUser(username=user_dict["username"], role=user_dict["role"]) - login_user(user, remember=False) return function(*args, **kwargs) return cast(T, decorated) diff --git a/tests/www/views/test_views_custom_user_views.py b/tests/www/views/test_views_custom_user_views.py index ae6d0132827c2..84947a8e5f36f 100644 --- a/tests/www/views/test_views_custom_user_views.py +++ b/tests/www/views/test_views_custom_user_views.py @@ -27,7 +27,10 @@ from airflow import settings from airflow.security import permissions from airflow.www import app as application -from tests.test_utils.api_connexion_utils import create_user, delete_role +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import ( + create_user as create_user, + delete_role, +) from tests.test_utils.www import check_content_in_response, check_content_not_in_response, client_with_login pytestmark = pytest.mark.db_test diff --git a/tests/www/views/test_views_dagrun.py b/tests/www/views/test_views_dagrun.py index 39c17d086f379..d95955246ac78 100644 --- a/tests/www/views/test_views_dagrun.py +++ b/tests/www/views/test_views_dagrun.py @@ -24,7 +24,11 @@ from airflow.utils import timezone from airflow.utils.session import create_session from airflow.www.views import DagRunModelView -from tests.test_utils.api_connexion_utils import create_user, delete_roles, delete_user +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import ( + create_user, + delete_roles, + delete_user, +) from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS from tests.test_utils.www import check_content_in_response, check_content_not_in_response, client_with_login from tests.www.views.test_views_tasks import _get_appbuilder_pk_string diff --git a/tests/www/views/test_views_home.py b/tests/www/views/test_views_home.py index 5393115041392..ddec0c0bcfed3 100644 --- a/tests/www/views/test_views_home.py +++ b/tests/www/views/test_views_home.py @@ -27,7 +27,7 @@ from airflow.utils.state import State from airflow.www.utils import UIAlert from airflow.www.views import FILTER_LASTRUN_COOKIE, FILTER_STATUS_COOKIE, FILTER_TAGS_COOKIE -from tests.test_utils.api_connexion_utils import create_user +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_user from tests.test_utils.db import clear_db_dags, clear_db_import_errors, clear_db_serialized_dags from tests.test_utils.permissions import _resource_name from tests.test_utils.www import check_content_in_response, check_content_not_in_response, client_with_login diff --git a/tests/www/views/test_views_tasks.py b/tests/www/views/test_views_tasks.py index f5cc011fb6f0e..7b65051724c27 100644 --- a/tests/www/views/test_views_tasks.py +++ b/tests/www/views/test_views_tasks.py @@ -44,7 +44,11 @@ from airflow.utils.state import DagRunState, State from airflow.utils.types import DagRunType from airflow.www.views import TaskInstanceModelView, _safe_parse_datetime -from tests.test_utils.api_connexion_utils import create_user, delete_roles, delete_user +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import ( + create_user, + delete_roles, + delete_user, +) from tests.test_utils.compat import AIRFLOW_V_3_0_PLUS from tests.test_utils.config import conf_vars from tests.test_utils.db import clear_db_runs, clear_db_xcom diff --git a/tests/www/views/test_views_variable.py b/tests/www/views/test_views_variable.py index a91a12ddc470b..b7fa8b37c52c8 100644 --- a/tests/www/views/test_views_variable.py +++ b/tests/www/views/test_views_variable.py @@ -25,7 +25,7 @@ from airflow.models import Variable from airflow.security import permissions from airflow.utils.session import create_session -from tests.test_utils.api_connexion_utils import create_user +from tests.providers.fab.auth_manager.api_endpoints.api_connexion_utils import create_user from tests.test_utils.www import ( _check_last_log, check_content_in_response, From e2130c972772b4e218a49d3201519681bd74e4d9 Mon Sep 17 00:00:00 2001 From: Brent Bovenzi Date: Tue, 1 Oct 2024 17:25:18 +0200 Subject: [PATCH 085/802] Add Docs button to Nav (#42586) * Add Docs button to new UI nav * Add Docs menu button to Nav * Use src alias * Address PR feedback, update documentation * Delete airflow/ui/.env.local --- .gitignore | 1 + airflow/ui/.env.example | 23 +++++++ airflow/ui/src/layouts/Nav/DocsButton.tsx | 67 +++++++++++++++++++ airflow/ui/src/layouts/{ => Nav}/Nav.tsx | 9 ++- .../ui/src/layouts/{ => Nav}/NavButton.tsx | 17 ++--- airflow/ui/src/layouts/Nav/index.tsx | 20 ++++++ airflow/ui/src/layouts/Nav/navButtonProps.ts | 30 +++++++++ airflow/ui/src/main.tsx | 2 +- airflow/ui/src/vite-env.d.ts | 9 +++ .../14_node_environment_setup.rst | 17 +++++ 10 files changed, 178 insertions(+), 17 deletions(-) create mode 100644 airflow/ui/.env.example create mode 100644 airflow/ui/src/layouts/Nav/DocsButton.tsx rename airflow/ui/src/layouts/{ => Nav}/Nav.tsx (94%) rename airflow/ui/src/layouts/{ => Nav}/NavButton.tsx (83%) create mode 100644 airflow/ui/src/layouts/Nav/index.tsx create mode 100644 airflow/ui/src/layouts/Nav/navButtonProps.ts diff --git a/.gitignore b/.gitignore index 257331cb4e90b..a9c055041d980 100644 --- a/.gitignore +++ b/.gitignore @@ -111,6 +111,7 @@ celerybeat-schedule # dotenv .env +.env.local .autoenv*.zsh # virtualenv diff --git a/airflow/ui/.env.example b/airflow/ui/.env.example new file mode 100644 index 0000000000000..9374d93de6bca --- /dev/null +++ b/airflow/ui/.env.example @@ -0,0 +1,23 @@ +# +# 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. +#/ + + +# This is an example. You should make your own `.env.local` file for development + +VITE_FASTAPI_URL="http://localhost:29091" diff --git a/airflow/ui/src/layouts/Nav/DocsButton.tsx b/airflow/ui/src/layouts/Nav/DocsButton.tsx new file mode 100644 index 0000000000000..07a4b93dfaede --- /dev/null +++ b/airflow/ui/src/layouts/Nav/DocsButton.tsx @@ -0,0 +1,67 @@ +/*! + * 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. + */ +import { + IconButton, + Link, + Menu, + MenuButton, + MenuItem, + MenuList, +} from "@chakra-ui/react"; +import { FiBookOpen } from "react-icons/fi"; + +import { navButtonProps } from "./navButtonProps"; + +const links = [ + { + href: "https://airflow.apache.org/docs/", + title: "Documentation", + }, + { + href: "https://github.com/apache/airflow", + title: "GitHub Repo", + }, + { + href: `${import.meta.env.VITE_FASTAPI_URL}/docs`, + title: "REST API Reference", + }, +]; + +export const DocsButton = () => ( + + } + {...navButtonProps} + /> + + {links.map((link) => ( + + {link.title} + + ))} + + +); diff --git a/airflow/ui/src/layouts/Nav.tsx b/airflow/ui/src/layouts/Nav/Nav.tsx similarity index 94% rename from airflow/ui/src/layouts/Nav.tsx rename to airflow/ui/src/layouts/Nav/Nav.tsx index 4900540cd96d1..55bfd4480e0f4 100644 --- a/airflow/ui/src/layouts/Nav.tsx +++ b/airflow/ui/src/layouts/Nav/Nav.tsx @@ -37,8 +37,10 @@ import { FiSun, } from "react-icons/fi"; -import { AirflowPin } from "../assets/AirflowPin"; -import { DagIcon } from "../assets/DagIcon"; +import { AirflowPin } from "src/assets/AirflowPin"; +import { DagIcon } from "src/assets/DagIcon"; + +import { DocsButton } from "./DocsButton"; import { NavButton } from "./NavButton"; export const Nav = () => { @@ -78,7 +80,7 @@ export const Nav = () => { } isDisabled - title="Datasets" + title="Assets" /> } @@ -103,6 +105,7 @@ export const Nav = () => { icon={} title="Return to legacy UI" /> + ( - diff --git a/airflow/ui/src/layouts/Nav/index.tsx b/airflow/ui/src/layouts/Nav/index.tsx new file mode 100644 index 0000000000000..403e140919b04 --- /dev/null +++ b/airflow/ui/src/layouts/Nav/index.tsx @@ -0,0 +1,20 @@ +/*! + * 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. + */ + +export { Nav } from "./Nav"; diff --git a/airflow/ui/src/layouts/Nav/navButtonProps.ts b/airflow/ui/src/layouts/Nav/navButtonProps.ts new file mode 100644 index 0000000000000..740348bc9676b --- /dev/null +++ b/airflow/ui/src/layouts/Nav/navButtonProps.ts @@ -0,0 +1,30 @@ +/*! + * 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. + */ +import type { ButtonProps } from "@chakra-ui/react"; + +export const navButtonProps: ButtonProps = { + alignItems: "center", + borderRadius: "none", + flexDir: "column", + height: 16, + transition: "0.2s background-color ease-in-out", + variant: "ghost", + whiteSpace: "wrap", + width: 24, +}; diff --git a/airflow/ui/src/main.tsx b/airflow/ui/src/main.tsx index ca5fbed04b6ce..7b762508ea7b3 100644 --- a/airflow/ui/src/main.tsx +++ b/airflow/ui/src/main.tsx @@ -43,7 +43,7 @@ const queryClient = new QueryClient({ }, }); -axios.defaults.baseURL = "http://localhost:29091"; +axios.defaults.baseURL = import.meta.env.VITE_FASTAPI_URL; // redirect to login page if the API responds with unauthorized or forbidden errors axios.interceptors.response.use( diff --git a/airflow/ui/src/vite-env.d.ts b/airflow/ui/src/vite-env.d.ts index a1fdcdd1e6fc5..193866687bff9 100644 --- a/airflow/ui/src/vite-env.d.ts +++ b/airflow/ui/src/vite-env.d.ts @@ -1,3 +1,4 @@ +/* eslint-disable @typescript-eslint/consistent-type-definitions */ /*! * Licensed to the Apache Software Foundation (ASF) under one * or more contributor license agreements. See the NOTICE file @@ -18,3 +19,11 @@ */ /// + +interface ImportMetaEnv { + readonly VITE_FASTAPI_URL: string; +} + +interface ImportMeta { + readonly env: ImportMetaEnv; +} diff --git a/contributing-docs/14_node_environment_setup.rst b/contributing-docs/14_node_environment_setup.rst index 8d98f0860fc8b..7b10f0b0d5ed5 100644 --- a/contributing-docs/14_node_environment_setup.rst +++ b/contributing-docs/14_node_environment_setup.rst @@ -84,6 +84,23 @@ Project Structure - ``/src/components`` shared components across the UI - ``/dist`` build files +Local Environment Variables +--------------------------- + +Copy the example environment + +.. code-block:: bash + + cp .env.example .env.local + +If you run into CORS issues, you may need to add some variables to your Breeze config, ``files/airflow-breeze-config/variables.env``: + +.. code-block:: bash + + export AIRFLOW__API__ACCESS_CONTROL_ALLOW_HEADERS="Origin, Access-Control-Request-Method" + export AIRFLOW__API__ACCESS_CONTROL_ALLOW_METHODS="*" + export AIRFLOW__API__ACCESS_CONTROL_ALLOW_ORIGINS="http://localhost:28080,http://localhost:8080" + DEPRECATED Airflow WWW From 09bb6d1342ec971a3bf4c751ffc8b146e5c21d96 Mon Sep 17 00:00:00 2001 From: Elad Kalif <45845474+eladkal@users.noreply.github.com> Date: Tue, 1 Oct 2024 22:31:17 +0700 Subject: [PATCH 086/802] Update providers metadata 2024-10-01 (#42611) --- generated/provider_metadata.json | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/generated/provider_metadata.json b/generated/provider_metadata.json index a73e3da9f6fce..56199e2f82c3d 100644 --- a/generated/provider_metadata.json +++ b/generated/provider_metadata.json @@ -2763,6 +2763,10 @@ "1.17.0": { "associated_airflow_version": "2.10.1", "date_released": "2024-09-24T13:49:56Z" + }, + "1.17.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-10-01T09:05:14Z" } }, "databricks": { @@ -6225,6 +6229,10 @@ "1.12.0": { "associated_airflow_version": "2.10.1", "date_released": "2024-09-24T13:49:56Z" + }, + "1.12.1": { + "associated_airflow_version": "2.10.1", + "date_released": "2024-10-01T09:05:14Z" } }, "opensearch": { From 084eabe2014cd1a62b49a933370514f96d1dd505 Mon Sep 17 00:00:00 2001 From: Julian Maicher Date: Tue, 1 Oct 2024 21:23:40 +0200 Subject: [PATCH 087/802] Prevent redirect loop on /home with tags/lastrun filters (#42607) (#42609) Closes #42607 --- airflow/www/views.py | 17 +++++++------ tests/www/views/test_views_home.py | 38 ++++++++++++++++++++++++++++++ 2 files changed, 48 insertions(+), 7 deletions(-) diff --git a/airflow/www/views.py b/airflow/www/views.py index b3300b517e757..a361b9bd50c16 100644 --- a/airflow/www/views.py +++ b/airflow/www/views.py @@ -814,19 +814,22 @@ def index(self): return redirect(url_for("Airflow.index")) filter_tags_cookie_val = flask_session.get(FILTER_TAGS_COOKIE) + filter_lastrun_cookie_val = flask_session.get(FILTER_LASTRUN_COOKIE) + + # update filter args in url from session values if needed + if (not arg_tags_filter and filter_tags_cookie_val) or ( + not arg_lastrun_filter and filter_lastrun_cookie_val + ): + tags = arg_tags_filter or (filter_tags_cookie_val and filter_tags_cookie_val.split(",")) + lastrun = arg_lastrun_filter or filter_lastrun_cookie_val + return redirect(url_for("Airflow.index", tags=tags, lastrun=lastrun)) + if arg_tags_filter: flask_session[FILTER_TAGS_COOKIE] = ",".join(arg_tags_filter) - elif filter_tags_cookie_val: - # If tags exist in cookie, but not URL, add them to the URL - return redirect(url_for("Airflow.index", tags=filter_tags_cookie_val.split(","))) - filter_lastrun_cookie_val = flask_session.get(FILTER_LASTRUN_COOKIE) if arg_lastrun_filter: arg_lastrun_filter = arg_lastrun_filter.strip().lower() flask_session[FILTER_LASTRUN_COOKIE] = arg_lastrun_filter - elif filter_lastrun_cookie_val: - # If tags exist in cookie, but not URL, add them to the URL - return redirect(url_for("Airflow.index", lastrun=filter_lastrun_cookie_val)) if arg_status_filter is None: filter_status_cookie_val = flask_session.get(FILTER_STATUS_COOKIE) diff --git a/tests/www/views/test_views_home.py b/tests/www/views/test_views_home.py index ddec0c0bcfed3..44dda24feecbc 100644 --- a/tests/www/views/test_views_home.py +++ b/tests/www/views/test_views_home.py @@ -466,3 +466,41 @@ def test_analytics_pixel(user_client, is_enabled, should_have_pixel): check_content_in_response("apacheairflow.gateway.scarf.sh", resp) else: check_content_not_in_response("apacheairflow.gateway.scarf.sh", resp) + + +@pytest.mark.parametrize( + "url, filter_tags_cookie_val, filter_lastrun_cookie_val, expected_filter_tags, expected_filter_lastrun", + [ + ("home", None, None, [], None), + # from url only + ("home?tags=example&tags=test", None, None, ["example", "test"], None), + ("home?lastrun=running", None, None, [], "running"), + ("home?tags=example&tags=test&lastrun=running", None, None, ["example", "test"], "running"), + # from cookie only + ("home", "example,test", None, ["example", "test"], None), + ("home", None, "running", [], "running"), + ("home", "example,test", "running", ["example", "test"], "running"), + # from url and cookie + ("home?tags=example", "example,test", None, ["example"], None), + ("home?lastrun=failed", None, "running", [], "failed"), + ("home?tags=example", None, "running", ["example"], "running"), + ("home?lastrun=running", "example,test", None, ["example", "test"], "running"), + ("home?tags=example&lastrun=running", "example,test", "failed", ["example"], "running"), + ], +) +def test_filter_cookie_eval( + working_dags, + admin_client, + url, + filter_tags_cookie_val, + filter_lastrun_cookie_val, + expected_filter_tags, + expected_filter_lastrun, +): + with admin_client.session_transaction() as flask_session: + flask_session[FILTER_TAGS_COOKIE] = filter_tags_cookie_val + flask_session[FILTER_LASTRUN_COOKIE] = filter_lastrun_cookie_val + + resp = admin_client.get(url, follow_redirects=True) + assert resp.request.args.getlist("tags") == expected_filter_tags + assert resp.request.args.get("lastrun") == expected_filter_lastrun From 0383b09c58f3bcff6fa982fb09e0e9084c8a1670 Mon Sep 17 00:00:00 2001 From: Vincent <97131062+vincbeck@users.noreply.github.com> Date: Tue, 1 Oct 2024 15:37:19 -0400 Subject: [PATCH 088/802] Remove `AIRFLOW_V_2_7_PLUS` constant (#42627) --- contributing-docs/testing/unit_tests.rst | 4 ++-- tests/test_utils/compat.py | 1 - 2 files changed, 2 insertions(+), 3 deletions(-) diff --git a/contributing-docs/testing/unit_tests.rst b/contributing-docs/testing/unit_tests.rst index 935a7b9b602b4..cc6513eaa2e5f 100644 --- a/contributing-docs/testing/unit_tests.rst +++ b/contributing-docs/testing/unit_tests.rst @@ -1184,10 +1184,10 @@ are not part of the public API. We deal with it in one of the following ways: .. code-block:: python - from tests.test_utils.compat import AIRFLOW_V_2_7_PLUS + from tests.test_utils.compat import AIRFLOW_V_2_8_PLUS - @pytest.mark.skipif(not AIRFLOW_V_2_7_PLUS, reason="The tests should be skipped for Airflow < 2.7") + @pytest.mark.skipif(not AIRFLOW_V_2_8_PLUS, reason="The tests should be skipped for Airflow < 2.8") def some_test_that_only_works_for_airflow_2_7_plus(): pass diff --git a/tests/test_utils/compat.py b/tests/test_utils/compat.py index ca1d7e9c77dfa..09f3653db82d8 100644 --- a/tests/test_utils/compat.py +++ b/tests/test_utils/compat.py @@ -42,7 +42,6 @@ from airflow import __version__ as airflow_version AIRFLOW_VERSION = Version(airflow_version) -AIRFLOW_V_2_7_PLUS = Version(AIRFLOW_VERSION.base_version) >= Version("2.7.0") AIRFLOW_V_2_8_PLUS = Version(AIRFLOW_VERSION.base_version) >= Version("2.8.0") AIRFLOW_V_2_9_PLUS = Version(AIRFLOW_VERSION.base_version) >= Version("2.9.0") AIRFLOW_V_2_10_PLUS = Version(AIRFLOW_VERSION.base_version) >= Version("2.10.0") From f898ec78d55d12d1a31b738c7fe5ada22c140441 Mon Sep 17 00:00:00 2001 From: GPK Date: Tue, 1 Oct 2024 21:00:29 +0100 Subject: [PATCH 089/802] Move FSHook/PackageIndexHook/SubprocessHook to standard provider (#42506) * move hooks to standard providers * fix document build and adding hooks to provider yaml file * adding fshook tests * marking as db test * doc reference update to subprocess hook --- airflow/operators/bash.py | 2 +- airflow/providers/standard/hooks/__init__.py | 16 ++++++++ .../standard}/hooks/filesystem.py | 0 .../standard}/hooks/package_index.py | 0 .../standard}/hooks/subprocess.py | 4 +- airflow/providers/standard/provider.yaml | 7 ++++ airflow/providers_manager.py | 4 +- airflow/sensors/filesystem.py | 2 +- .../logging-monitoring/errors.rst | 2 +- .../operators-and-hooks-ref.rst | 4 +- tests/providers/standard/hooks/__init__.py | 16 ++++++++ .../standard/hooks/test_filesystem.py | 39 +++++++++++++++++++ .../standard}/hooks/test_package_index.py | 6 +-- .../standard}/hooks/test_subprocess.py | 6 +-- tests/sensors/test_filesystem.py | 2 +- 15 files changed, 94 insertions(+), 16 deletions(-) create mode 100644 airflow/providers/standard/hooks/__init__.py rename airflow/{ => providers/standard}/hooks/filesystem.py (100%) rename airflow/{ => providers/standard}/hooks/package_index.py (100%) rename airflow/{ => providers/standard}/hooks/subprocess.py (96%) create mode 100644 tests/providers/standard/hooks/__init__.py create mode 100644 tests/providers/standard/hooks/test_filesystem.py rename tests/{ => providers/standard}/hooks/test_package_index.py (93%) rename tests/{ => providers/standard}/hooks/test_subprocess.py (95%) diff --git a/airflow/operators/bash.py b/airflow/operators/bash.py index 2ec0341a0d1e2..bf4a943df6e08 100644 --- a/airflow/operators/bash.py +++ b/airflow/operators/bash.py @@ -24,8 +24,8 @@ from typing import TYPE_CHECKING, Any, Callable, Container, Sequence, cast from airflow.exceptions import AirflowException, AirflowSkipException -from airflow.hooks.subprocess import SubprocessHook from airflow.models.baseoperator import BaseOperator +from airflow.providers.standard.hooks.subprocess import SubprocessHook from airflow.utils.operator_helpers import context_to_airflow_vars from airflow.utils.types import ArgNotSet diff --git a/airflow/providers/standard/hooks/__init__.py b/airflow/providers/standard/hooks/__init__.py new file mode 100644 index 0000000000000..13a83393a9124 --- /dev/null +++ b/airflow/providers/standard/hooks/__init__.py @@ -0,0 +1,16 @@ +# 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. diff --git a/airflow/hooks/filesystem.py b/airflow/providers/standard/hooks/filesystem.py similarity index 100% rename from airflow/hooks/filesystem.py rename to airflow/providers/standard/hooks/filesystem.py diff --git a/airflow/hooks/package_index.py b/airflow/providers/standard/hooks/package_index.py similarity index 100% rename from airflow/hooks/package_index.py rename to airflow/providers/standard/hooks/package_index.py diff --git a/airflow/hooks/subprocess.py b/airflow/providers/standard/hooks/subprocess.py similarity index 96% rename from airflow/hooks/subprocess.py rename to airflow/providers/standard/hooks/subprocess.py index bc20b5c20b4c5..9e578a7d8034b 100644 --- a/airflow/hooks/subprocess.py +++ b/airflow/providers/standard/hooks/subprocess.py @@ -52,8 +52,8 @@ def run_command( :param env: Optional dict containing environment variables to be made available to the shell environment in which ``command`` will be executed. If omitted, ``os.environ`` will be used. Note, that in case you have Sentry configured, original variables from the environment - will also be passed to the subprocess with ``SUBPROCESS_`` prefix. See - :doc:`/administration-and-deployment/logging-monitoring/errors` for details. + will also be passed to the subprocess with ``SUBPROCESS_`` prefix. See: + https://airflow.apache.org/docs/apache-airflow/stable/administration-and-deployment/logging-monitoring/errors.html for details. :param output_encoding: encoding to use for decoding stdout :param cwd: Working directory to run the command in. If None (default), the command is run in a temporary directory. diff --git a/airflow/providers/standard/provider.yaml b/airflow/providers/standard/provider.yaml index 83d8acf0a68b3..068fde1fe3761 100644 --- a/airflow/providers/standard/provider.yaml +++ b/airflow/providers/standard/provider.yaml @@ -50,3 +50,10 @@ sensors: - airflow.providers.standard.sensors.time_delta - airflow.providers.standard.sensors.time - airflow.providers.standard.sensors.weekday + +hooks: + - integration-name: Standard + python-modules: + - airflow.providers.standard.hooks.filesystem + - airflow.providers.standard.hooks.package_index + - airflow.providers.standard.hooks.subprocess diff --git a/airflow/providers_manager.py b/airflow/providers_manager.py index 2c673063cb23e..e276c465ef689 100644 --- a/airflow/providers_manager.py +++ b/airflow/providers_manager.py @@ -36,8 +36,8 @@ from packaging.utils import canonicalize_name from airflow.exceptions import AirflowOptionalProviderFeatureException -from airflow.hooks.filesystem import FSHook -from airflow.hooks.package_index import PackageIndexHook +from airflow.providers.standard.hooks.filesystem import FSHook +from airflow.providers.standard.hooks.package_index import PackageIndexHook from airflow.typing_compat import ParamSpec from airflow.utils import yaml from airflow.utils.entry_points import entry_points_with_dist diff --git a/airflow/sensors/filesystem.py b/airflow/sensors/filesystem.py index 5d32ab07ad4e7..4496f5d6abfa4 100644 --- a/airflow/sensors/filesystem.py +++ b/airflow/sensors/filesystem.py @@ -25,7 +25,7 @@ from airflow.configuration import conf from airflow.exceptions import AirflowException -from airflow.hooks.filesystem import FSHook +from airflow.providers.standard.hooks.filesystem import FSHook from airflow.sensors.base import BaseSensorOperator from airflow.triggers.base import StartTriggerArgs from airflow.triggers.file import FileTrigger diff --git a/docs/apache-airflow/administration-and-deployment/logging-monitoring/errors.rst b/docs/apache-airflow/administration-and-deployment/logging-monitoring/errors.rst index cb09843422321..0ad3fa8c5127a 100644 --- a/docs/apache-airflow/administration-and-deployment/logging-monitoring/errors.rst +++ b/docs/apache-airflow/administration-and-deployment/logging-monitoring/errors.rst @@ -96,7 +96,7 @@ Impact of Sentry on Environment variables passed to Subprocess Hook When Sentry is enabled, by default it changes the standard library to pass all environment variables to subprocesses opened by Airflow. This changes the default behaviour of -:class:`airflow.hooks.subprocess.SubprocessHook` - always all environment variables are passed to the +:class:`airflow.providers.standard.hooks.subprocess.SubprocessHook` - always all environment variables are passed to the subprocess executed with specific set of environment variables. In this case not only the specified environment variables are passed but also all existing environment variables are passed with ``SUBPROCESS_`` prefix added. This happens also for all other subprocesses. diff --git a/docs/apache-airflow/operators-and-hooks-ref.rst b/docs/apache-airflow/operators-and-hooks-ref.rst index 16b74305a958b..d4ac6bda74c34 100644 --- a/docs/apache-airflow/operators-and-hooks-ref.rst +++ b/docs/apache-airflow/operators-and-hooks-ref.rst @@ -106,8 +106,8 @@ For details see: :doc:`apache-airflow-providers:operators-and-hooks-ref/index`. * - Hooks - Guides - * - :mod:`airflow.hooks.filesystem` + * - :mod:`airflow.providers.standard.hooks.filesystem` - - * - :mod:`airflow.hooks.subprocess` + * - :mod:`airflow.providers.standard.hooks.subprocess` - diff --git a/tests/providers/standard/hooks/__init__.py b/tests/providers/standard/hooks/__init__.py new file mode 100644 index 0000000000000..13a83393a9124 --- /dev/null +++ b/tests/providers/standard/hooks/__init__.py @@ -0,0 +1,16 @@ +# 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. diff --git a/tests/providers/standard/hooks/test_filesystem.py b/tests/providers/standard/hooks/test_filesystem.py new file mode 100644 index 0000000000000..bbcd22dc94219 --- /dev/null +++ b/tests/providers/standard/hooks/test_filesystem.py @@ -0,0 +1,39 @@ +# +# 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 pytest + +from airflow.providers.standard.hooks.filesystem import FSHook + +pytestmark = pytest.mark.db_test + + +class TestFSHook: + def test_get_ui_field_behaviour(self): + fs_hook = FSHook() + assert fs_hook.get_ui_field_behaviour() == { + "hidden_fields": ["host", "schema", "port", "login", "password", "extra"], + "relabeling": {}, + "placeholders": {}, + } + + def test_get_path(self): + fs_hook = FSHook(fs_conn_id="fs_default") + + assert fs_hook.get_path() == "/" diff --git a/tests/hooks/test_package_index.py b/tests/providers/standard/hooks/test_package_index.py similarity index 93% rename from tests/hooks/test_package_index.py rename to tests/providers/standard/hooks/test_package_index.py index 9da429c5a09cf..6a90db0715d81 100644 --- a/tests/hooks/test_package_index.py +++ b/tests/providers/standard/hooks/test_package_index.py @@ -21,8 +21,8 @@ import pytest -from airflow.hooks.package_index import PackageIndexHook from airflow.models.connection import Connection +from airflow.providers.standard.hooks.package_index import PackageIndexHook class MockConnection(Connection): @@ -73,7 +73,7 @@ def mock_get_connection(monkeypatch: pytest.MonkeyPatch, request: pytest.Fixture password: str | None = testdata.get("password", None) expected_result: str | None = testdata.get("expected_result", None) monkeypatch.setattr( - "airflow.hooks.package_index.PackageIndexHook.get_connection", + "airflow.providers.standard.hooks.package_index.PackageIndexHook.get_connection", lambda *_: MockConnection(host, login, password), ) return expected_result @@ -104,7 +104,7 @@ class MockProc: return MockProc() - monkeypatch.setattr("airflow.hooks.package_index.subprocess.run", mock_run) + monkeypatch.setattr("airflow.providers.standard.hooks.package_index.subprocess.run", mock_run) hook_instance = PackageIndexHook() if mock_get_connection: diff --git a/tests/hooks/test_subprocess.py b/tests/providers/standard/hooks/test_subprocess.py similarity index 95% rename from tests/hooks/test_subprocess.py rename to tests/providers/standard/hooks/test_subprocess.py index 0f625be816887..2b2e9473359e5 100644 --- a/tests/hooks/test_subprocess.py +++ b/tests/providers/standard/hooks/test_subprocess.py @@ -26,7 +26,7 @@ import pytest -from airflow.hooks.subprocess import SubprocessHook +from airflow.providers.standard.hooks.subprocess import SubprocessHook OS_ENV_KEY = "SUBPROCESS_ENV_TEST" OS_ENV_VAL = "this-is-from-os-environ" @@ -81,11 +81,11 @@ def test_return_value(self, val, expected): @mock.patch.dict("os.environ", clear=True) @mock.patch( - "airflow.hooks.subprocess.TemporaryDirectory", + "airflow.providers.standard.hooks.subprocess.TemporaryDirectory", return_value=MagicMock(__enter__=MagicMock(return_value="/tmp/airflowtmpcatcat")), ) @mock.patch( - "airflow.hooks.subprocess.Popen", + "airflow.providers.standard.hooks.subprocess.Popen", return_value=MagicMock(stdout=MagicMock(readline=MagicMock(side_effect=StopIteration), returncode=0)), ) def test_should_exec_subprocess(self, mock_popen, mock_temporary_directory): diff --git a/tests/sensors/test_filesystem.py b/tests/sensors/test_filesystem.py index 1fb123cfe7248..641f2f218f2db 100644 --- a/tests/sensors/test_filesystem.py +++ b/tests/sensors/test_filesystem.py @@ -40,7 +40,7 @@ @pytest.mark.skip_if_database_isolation_mode # Test is broken in db isolation mode class TestFileSensor: def setup_method(self): - from airflow.hooks.filesystem import FSHook + from airflow.providers.standard.hooks.filesystem import FSHook hook = FSHook() args = {"owner": "airflow", "start_date": DEFAULT_DATE} From 6ddb8ebc036a254661b0c01bb7848a8ce039e5af Mon Sep 17 00:00:00 2001 From: rom sharon <33751805+romsharon98@users.noreply.github.com> Date: Tue, 1 Oct 2024 23:05:45 +0300 Subject: [PATCH 090/802] send notification to internal-ci-cd channel (#42630) --- .github/workflows/ci.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 8828a30ce3ecd..716323cb9acfd 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -690,7 +690,7 @@ jobs: id: slack uses: slackapi/slack-github-action@v1.27.0 with: - channel-id: 'zzz_webhook_test' + channel-id: 'internal-airflow-ci-cd' # yamllint disable rule:line-length payload: | { From 2f90f75b3fed1f3174dc57a4c7a82afe3d19435f Mon Sep 17 00:00:00 2001 From: Kalyan Date: Wed, 2 Oct 2024 03:29:54 +0530 Subject: [PATCH 091/802] Add heartbeat metric for DAG processor (#42398) --------- Signed-off-by: kalyanr --- airflow/jobs/dag_processor_job_runner.py | 16 ++++++++++------ chart/files/statsd-mappings.yml | 6 ++++++ .../logging-monitoring/metrics.rst | 1 + 3 files changed, 17 insertions(+), 6 deletions(-) diff --git a/airflow/jobs/dag_processor_job_runner.py b/airflow/jobs/dag_processor_job_runner.py index 76b2ab5925540..28128efba474b 100644 --- a/airflow/jobs/dag_processor_job_runner.py +++ b/airflow/jobs/dag_processor_job_runner.py @@ -17,18 +17,18 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING from airflow.jobs.base_job_runner import BaseJobRunner from airflow.jobs.job import Job, perform_heartbeat +from airflow.stats import Stats from airflow.utils.log.logging_mixin import LoggingMixin +from airflow.utils.session import NEW_SESSION, provide_session if TYPE_CHECKING: - from airflow.dag_processing.manager import DagFileProcessorManager - + from sqlalchemy.orm import Session -def empty_callback(_: Any) -> None: - pass + from airflow.dag_processing.manager import DagFileProcessorManager class DagProcessorJobRunner(BaseJobRunner, LoggingMixin): @@ -52,7 +52,7 @@ def __init__( self.processor = processor self.processor.heartbeat = lambda: perform_heartbeat( job=self.job, - heartbeat_callback=empty_callback, + heartbeat_callback=self.heartbeat_callback, only_if_necessary=True, ) @@ -67,3 +67,7 @@ def _execute(self) -> int | None: self.processor.terminate() self.processor.end() return None + + @provide_session + def heartbeat_callback(self, session: Session = NEW_SESSION) -> None: + Stats.incr("dag_processor_heartbeat", 1, 1) diff --git a/chart/files/statsd-mappings.yml b/chart/files/statsd-mappings.yml index 86d773fd20b7f..cef9593dd16d3 100644 --- a/chart/files/statsd-mappings.yml +++ b/chart/files/statsd-mappings.yml @@ -46,6 +46,12 @@ mappings: labels: type: counter + - match: airflow.dag_processor_heartbeat + match_type: regex + name: "airflow_dag_processor_heartbeat" + labels: + type: counter + - match: airflow.dag.*.*.duration name: "airflow_task_duration" labels: diff --git a/docs/apache-airflow/administration-and-deployment/logging-monitoring/metrics.rst b/docs/apache-airflow/administration-and-deployment/logging-monitoring/metrics.rst index ac44d1acba9c0..079aa5d397a41 100644 --- a/docs/apache-airflow/administration-and-deployment/logging-monitoring/metrics.rst +++ b/docs/apache-airflow/administration-and-deployment/logging-monitoring/metrics.rst @@ -159,6 +159,7 @@ Name Descripti ``previously_succeeded`` Number of previously succeeded task instances. Metric with dag_id and task_id tagging. ``zombies_killed`` Zombie tasks killed. Metric with dag_id and task_id tagging. ``scheduler_heartbeat`` Scheduler heartbeats +``dag_processor_heartbeat`` Standalone DAG processor heartbeats ``dag_processing.processes`` Relative number of currently running DAG parsing processes (ie this delta is negative when, since the last metric was sent, processes have completed). Metric with file_path and action tagging. From a8a19b5f86355dc67fcb810761eadca778057db8 Mon Sep 17 00:00:00 2001 From: Alexander Millin Date: Wed, 2 Oct 2024 03:44:04 +0300 Subject: [PATCH 092/802] Fix the order of tasks during serialization (#42219) SerializedDagModel().dag_hash may change if the order of tasks is not fixed --- airflow/serialization/serialized_objects.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/airflow/serialization/serialized_objects.py b/airflow/serialization/serialized_objects.py index a4801b767acc5..08944391b8166 100644 --- a/airflow/serialization/serialized_objects.py +++ b/airflow/serialization/serialized_objects.py @@ -1604,7 +1604,9 @@ def serialize_dag(cls, dag: DAG) -> dict: try: serialized_dag = cls.serialize_to_json(dag, cls._decorated_fields) serialized_dag["_processor_dags_folder"] = DAGS_FOLDER - serialized_dag["tasks"] = [cls.serialize(task) for _, task in dag.task_dict.items()] + serialized_dag["tasks"] = [ + cls.serialize(dag.task_dict[task_id]) for task_id in sorted(dag.task_dict) + ] dag_deps = [ dep From 93295d21b37ba40d081e5e84ae7aa22dff4a6510 Mon Sep 17 00:00:00 2001 From: Kyle Thatcher <33584092+Kytha@users.noreply.github.com> Date: Tue, 1 Oct 2024 20:46:16 -0400 Subject: [PATCH 093/802] Remove state sync during celery task processing (#41870) --- airflow/providers/celery/executors/celery_executor.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/airflow/providers/celery/executors/celery_executor.py b/airflow/providers/celery/executors/celery_executor.py index 93037bb31c136..807c77ab98782 100644 --- a/airflow/providers/celery/executors/celery_executor.py +++ b/airflow/providers/celery/executors/celery_executor.py @@ -300,9 +300,6 @@ def _process_tasks(self, task_tuples: list[TaskTuple]) -> None: # which point we don't need the ID anymore anyway self.event_buffer[key] = (TaskInstanceState.QUEUED, result.task_id) - # If the task runs _really quickly_ we may already have a result! - self.update_task_state(key, result.state, getattr(result, "info", None)) - def _send_tasks_to_celery(self, task_tuples_to_send: list[TaskInstanceInCelery]): from airflow.providers.celery.executors.celery_executor_utils import send_task_to_executor From b4cdc0fda1e8e3e239a03b004c8d21eb1568e6d1 Mon Sep 17 00:00:00 2001 From: phi-friday Date: Wed, 2 Oct 2024 10:12:30 +0900 Subject: [PATCH 094/802] fix: rm `skip_if` and `run_if` in python source (#41832) --- airflow/utils/decorators.py | 2 +- tests/utils/test_decorators.py | 128 +++++++++++++++++++++++++++++++++ 2 files changed, 129 insertions(+), 1 deletion(-) create mode 100644 tests/utils/test_decorators.py diff --git a/airflow/utils/decorators.py b/airflow/utils/decorators.py index 4cad5ab9e6073..e299999423e56 100644 --- a/airflow/utils/decorators.py +++ b/airflow/utils/decorators.py @@ -49,7 +49,7 @@ def _remove_task_decorator(py_source, decorator_name): after_decorator = after_decorator[1:] return before_decorator + after_decorator - decorators = ["@setup", "@teardown", task_decorator_name] + decorators = ["@setup", "@teardown", "@task.skip_if", "@task.run_if", task_decorator_name] for decorator in decorators: python_source = _remove_task_decorator(python_source, decorator) return python_source diff --git a/tests/utils/test_decorators.py b/tests/utils/test_decorators.py new file mode 100644 index 0000000000000..19d3ec31d0311 --- /dev/null +++ b/tests/utils/test_decorators.py @@ -0,0 +1,128 @@ +# 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 + +from typing import TYPE_CHECKING + +import pytest + +from airflow.decorators import task + +if TYPE_CHECKING: + from airflow.decorators.base import Task, TaskDecorator + +_CONDITION_DECORATORS = frozenset({"skip_if", "run_if"}) +_NO_SOURCE_DECORATORS = frozenset({"sensor"}) +DECORATORS = sorted( + set(x for x in dir(task) if not x.startswith("_")) - _CONDITION_DECORATORS - _NO_SOURCE_DECORATORS +) +DECORATORS_USING_SOURCE = ("external_python", "virtualenv", "branch_virtualenv", "branch_external_python") + + +@pytest.fixture +def decorator(request: pytest.FixtureRequest) -> TaskDecorator: + decorator_factory = getattr(task, request.param) + + kwargs = {} + if "external" in request.param: + kwargs["python"] = "python3" + return decorator_factory(**kwargs) + + +@pytest.mark.parametrize("decorator", DECORATORS_USING_SOURCE, indirect=["decorator"]) +def test_task_decorator_using_source(decorator: TaskDecorator): + @decorator + def f(): + return ["some_task"] + + assert parse_python_source(f, "decorator") == 'def f():\n return ["some_task"]\n' + + +@pytest.mark.parametrize("decorator", DECORATORS, indirect=["decorator"]) +def test_skip_if(decorator: TaskDecorator): + @task.skip_if(lambda context: True) + @decorator + def f(): + return "hello world" + + assert parse_python_source(f, "decorator") == 'def f():\n return "hello world"\n' + + +@pytest.mark.parametrize("decorator", DECORATORS, indirect=["decorator"]) +def test_run_if(decorator: TaskDecorator): + @task.run_if(lambda context: True) + @decorator + def f(): + return "hello world" + + assert parse_python_source(f, "decorator") == 'def f():\n return "hello world"\n' + + +def test_skip_if_and_run_if(): + @task.skip_if(lambda context: True) + @task.run_if(lambda context: True) + @task.virtualenv() + def f(): + return "hello world" + + assert parse_python_source(f) == 'def f():\n return "hello world"\n' + + +def test_run_if_and_skip_if(): + @task.run_if(lambda context: True) + @task.skip_if(lambda context: True) + @task.virtualenv() + def f(): + return "hello world" + + assert parse_python_source(f) == 'def f():\n return "hello world"\n' + + +def test_skip_if_allow_decorator(): + def non_task_decorator(func): + return func + + @task.skip_if(lambda context: True) + @task.virtualenv() + @non_task_decorator + def f(): + return "hello world" + + assert parse_python_source(f) == '@non_task_decorator\ndef f():\n return "hello world"\n' + + +def test_run_if_allow_decorator(): + def non_task_decorator(func): + return func + + @task.run_if(lambda context: True) + @task.virtualenv() + @non_task_decorator + def f(): + return "hello world" + + assert parse_python_source(f) == '@non_task_decorator\ndef f():\n return "hello world"\n' + + +def parse_python_source(task: Task, custom_operator_name: str | None = None) -> str: + operator = task().operator + if custom_operator_name: + custom_operator_name = ( + custom_operator_name if custom_operator_name.startswith("@") else f"@{custom_operator_name}" + ) + operator.__dict__["custom_operator_name"] = custom_operator_name + return operator.get_python_source() From fd37e9c1a1b0044dae92be12d4044682e9ef25c9 Mon Sep 17 00:00:00 2001 From: phi-friday Date: Wed, 2 Oct 2024 10:13:38 +0900 Subject: [PATCH 095/802] fix: task flow dynamic mapping with default_args (#41592) --- airflow/decorators/base.py | 25 ++++++++++++++++++------- tests/decorators/test_mapped.py | 24 ++++++++++++++++++++++++ 2 files changed, 42 insertions(+), 7 deletions(-) diff --git a/airflow/decorators/base.py b/airflow/decorators/base.py index 1ef2c12c702f2..e650c1920a870 100644 --- a/airflow/decorators/base.py +++ b/airflow/decorators/base.py @@ -431,18 +431,29 @@ def _expand(self, expand_input: ExpandInput, *, strict: bool) -> XComArg: dag = task_kwargs.pop("dag", None) or DagContext.get_current_dag() task_group = task_kwargs.pop("task_group", None) or TaskGroupContext.get_current_task_group(dag) - partial_kwargs, partial_params = get_merged_defaults( + default_args, partial_params = get_merged_defaults( dag=dag, task_group=task_group, task_params=task_kwargs.pop("params", None), task_default_args=task_kwargs.pop("default_args", None), ) - partial_kwargs.update( - task_kwargs, - is_setup=self.is_setup, - is_teardown=self.is_teardown, - on_failure_fail_dagrun=self.on_failure_fail_dagrun, - ) + partial_kwargs: dict[str, Any] = { + "is_setup": self.is_setup, + "is_teardown": self.is_teardown, + "on_failure_fail_dagrun": self.on_failure_fail_dagrun, + } + base_signature = inspect.signature(BaseOperator) + ignore = { + "default_args", # This is target we are working on now. + "kwargs", # A common name for a keyword argument. + "do_xcom_push", # In the same boat as `multiple_outputs` + "multiple_outputs", # We will use `self.multiple_outputs` instead. + "params", # Already handled above `partial_params`. + "task_concurrency", # Deprecated(replaced by `max_active_tis_per_dag`). + } + partial_keys = set(base_signature.parameters) - ignore + partial_kwargs.update({key: value for key, value in default_args.items() if key in partial_keys}) + partial_kwargs.update(task_kwargs) task_id = get_unique_task_id(partial_kwargs.pop("task_id"), dag, task_group) if task_group: diff --git a/tests/decorators/test_mapped.py b/tests/decorators/test_mapped.py index 3812367425f8b..2d3747b5f34ef 100644 --- a/tests/decorators/test_mapped.py +++ b/tests/decorators/test_mapped.py @@ -17,6 +17,9 @@ # under the License. from __future__ import annotations +import pytest + +from airflow.decorators import task from airflow.models.dag import DAG from airflow.utils.task_group import TaskGroup from tests.models import DEFAULT_DATE @@ -36,3 +39,24 @@ def f(z): dag.get_task("t1") == x1.operator dag.get_task("g.t2") == x2.operator + + +@pytest.mark.db_test +def test_mapped_task_with_arbitrary_default_args(dag_maker, session): + default_args = {"some": "value", "not": "in", "the": "task", "or": "dag"} + with dag_maker(session=session, default_args=default_args): + + @task.python(do_xcom_push=True) + def f(x: int, y: int) -> int: + return x + y + + f.partial(y=10).expand(x=[1, 2, 3]) + + dag_run = dag_maker.create_dagrun(session=session) + decision = dag_run.task_instance_scheduling_decisions(session=session) + xcoms = set() + for ti in decision.schedulable_tis: + ti.run(session=session) + xcoms.add(ti.xcom_pull(session=session, task_ids=ti.task_id, map_indexes=ti.map_index)) + + assert xcoms == {11, 12, 13} From 1fc6633067a38bb0c4353fdcc645d60c20e4a719 Mon Sep 17 00:00:00 2001 From: Usiel Riedl Date: Wed, 2 Oct 2024 09:41:11 +0800 Subject: [PATCH 096/802] Adds new `triggerer.capacity_left[.]` metric (#41323) After reducing the default capacity our deployment, it is rather close to the total capacity at certain times, hence it would be useful to be able to create monitoring and alerting based on the left capacity. The new metric will enable better alerting (not relying on hardcoded capacity values) and even auto-scaling if wished for. --- airflow/jobs/triggerer_job_runner.py | 6 ++++++ .../logging-monitoring/metrics.rst | 3 +++ 2 files changed, 9 insertions(+) diff --git a/airflow/jobs/triggerer_job_runner.py b/airflow/jobs/triggerer_job_runner.py index b41af29f376ba..defde4a16471c 100644 --- a/airflow/jobs/triggerer_job_runner.py +++ b/airflow/jobs/triggerer_job_runner.py @@ -430,9 +430,15 @@ def emit_metrics(self): Stats.gauge( "triggers.running", len(self.trigger_runner.triggers), tags={"hostname": self.job.hostname} ) + + capacity_left = self.capacity - len(self.trigger_runner.triggers) + Stats.gauge(f"triggerer.capacity_left.{self.job.hostname}", capacity_left) + Stats.gauge("triggerer.capacity_left", capacity_left, tags={"hostname": self.job.hostname}) + span = Trace.get_current_span() span.set_attribute("trigger host", self.job.hostname) span.set_attribute("triggers running", len(self.trigger_runner.triggers)) + span.set_attribute("capacity left", capacity_left) class TriggerDetails(TypedDict): diff --git a/docs/apache-airflow/administration-and-deployment/logging-monitoring/metrics.rst b/docs/apache-airflow/administration-and-deployment/logging-monitoring/metrics.rst index 079aa5d397a41..7ce9b9b765a9c 100644 --- a/docs/apache-airflow/administration-and-deployment/logging-monitoring/metrics.rst +++ b/docs/apache-airflow/administration-and-deployment/logging-monitoring/metrics.rst @@ -248,6 +248,9 @@ Name Description ``triggers.running.`` Number of triggers currently running for a triggerer (described by hostname) ``triggers.running`` Number of triggers currently running for a triggerer (described by hostname). Metric with hostname tagging. +``triggerer.capacity_left.`` Capacity left on a triggerer to run triggers (described by hostname) +``triggerer.capacity_left`` Capacity left on a triggerer to run triggers (described by hostname). + Metric with hostname tagging. ==================================================== ======================================================================== Timers From 832f152c9bf66eda6c5d21b477b53f8f25e0d55e Mon Sep 17 00:00:00 2001 From: Andor Markus <51825189+andormarkus@users.noreply.github.com> Date: Wed, 2 Oct 2024 03:51:48 +0200 Subject: [PATCH 097/802] [HELM] - Add guide how to PgBouncer with Kubernetes Secret (#42460) * feat: Add guide how to PgBouncer with Kubernetes Secret * feat: Add guide how to PgBouncer with Kubernetes Secret --------- Co-authored-by: Andor Markus (AllCloud) --- docs/helm-chart/production-guide.rst | 74 ++++++++++++++++++++++++++++ 1 file changed, 74 insertions(+) diff --git a/docs/helm-chart/production-guide.rst b/docs/helm-chart/production-guide.rst index ee1fc2308be43..020394c8583fc 100644 --- a/docs/helm-chart/production-guide.rst +++ b/docs/helm-chart/production-guide.rst @@ -91,10 +91,84 @@ If you are using PostgreSQL as your database, you will likely want to enable `Pg Airflow can open a lot of database connections due to its distributed nature and using a connection pooler can significantly reduce the number of open connections on the database. +Database credentials stored Values file +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +.. code-block:: yaml + + pgbouncer: + enabled: true + + +Database credentials stored Kubernetes Secret +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +The default connection string in this case will not work you need to modify accordingly + +.. code-block:: bash + + kubectl create secret generic mydatabase --from-literal=connection=postgresql://user:pass@pgbouncer_svc_name.deployment_namespace:6543/airflow-metadata + +Two additional Kubernetes Secret required to PgBouncer able to properly work in this configuration: + +``airflow-pgbouncer-stats`` + +.. code-block:: bash + + kubectl create secret generic airflow-pgbouncer-stats --from-literal=connection=postgresql://user:pass@127.0.0.1:6543/pgbouncer?sslmode=disable + +``airflow-pgbouncer-config`` + +.. code-block:: yaml + + apiVersion: v1 + kind: Secret + metadata: + name: airflow-pgbouncer-config + data: + pgbouncer.ini: dmFsdWUtMg0KDQo= + users.txt: dmFsdWUtMg0KDQo= + + +``pgbouncer.ini`` equal to the base64 encoded version of this text + +.. code-block:: text + + [databases] + airflow-metadata = host={external_database_host} dbname={external_database_dbname} port=5432 pool_size=10 + + [pgbouncer] + pool_mode = transaction + listen_port = 6543 + listen_addr = * + auth_type = scram-sha-256 + auth_file = /etc/pgbouncer/users.txt + stats_users = postgres + ignore_startup_parameters = extra_float_digits + max_client_conn = 100 + verbose = 0 + log_disconnections = 0 + log_connections = 0 + + server_tls_sslmode = prefer + server_tls_ciphers = normal + +``users.txt`` equal to the base64 encoded version of this text + +.. code-block:: text + + "{ external_database_host }" "{ external_database_pass }" + +The ``values.yaml`` should looks like this + .. code-block:: yaml pgbouncer: enabled: true + configSecretName: airflow-pgbouncer-config + metricsExporterSidecar: + statsSecretName: airflow-pgbouncer-stats + Depending on the size of your Airflow instance, you may want to adjust the following as well (defaults are shown): From 46b9151530f941b842f0558564c96740982c4b5c Mon Sep 17 00:00:00 2001 From: Daniel Standish <15932138+dstandish@users.noreply.github.com> Date: Tue, 1 Oct 2024 18:54:15 -0700 Subject: [PATCH 098/802] Add dag run creation logic for backfill (#42529) Add basic backfill creation logic. This will be refined, but we're trying to be incremental here. --- .../endpoints/backfill_endpoint.py | 54 ++----- airflow/models/backfill.py | 127 ++++++++++++++- airflow/utils/types.py | 1 + .../endpoints/test_backfill_endpoint.py | 30 ++-- tests/models/test_backfill.py | 152 ++++++++++++++++++ 5 files changed, 308 insertions(+), 56 deletions(-) create mode 100644 tests/models/test_backfill.py diff --git a/airflow/api_connexion/endpoints/backfill_endpoint.py b/airflow/api_connexion/endpoints/backfill_endpoint.py index f974be4d75d82..baafdeea4f992 100644 --- a/airflow/api_connexion/endpoints/backfill_endpoint.py +++ b/airflow/api_connexion/endpoints/backfill_endpoint.py @@ -19,9 +19,10 @@ import logging from functools import wraps -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, cast import pendulum +from pendulum import DateTime from sqlalchemy import select from airflow.api_connexion import security @@ -31,8 +32,7 @@ backfill_collection_schema, backfill_schema, ) -from airflow.models.backfill import Backfill -from airflow.models.serialized_dag import SerializedDagModel +from airflow.models.backfill import AlreadyRunningBackfill, Backfill, _create_backfill from airflow.utils import timezone from airflow.utils.session import NEW_SESSION, provide_session from airflow.www.decorators import action_logging @@ -64,33 +64,6 @@ def wrapper(*, backfill_id, session, **kwargs): return wrapper -@provide_session -def _create_backfill( - *, - dag_id: str, - from_date: str, - to_date: str, - max_active_runs: int, - reverse: bool, - dag_run_conf: dict | None, - session: Session = NEW_SESSION, -) -> Backfill: - serdag = session.get(SerializedDagModel, dag_id) - if not serdag: - raise NotFound(f"Could not find dag {dag_id}") - - br = Backfill( - dag_id=dag_id, - from_date=pendulum.parse(from_date), - to_date=pendulum.parse(to_date), - max_active_runs=max_active_runs, - dag_run_conf=dag_run_conf, - ) - session.add(br) - session.commit() - return br - - @security.requires_access_dag("GET") @action_logging @provide_session @@ -170,12 +143,15 @@ def create_backfill( reverse: bool = False, dag_run_conf: dict | None = None, ) -> APIResponse: - backfill_obj = _create_backfill( - dag_id=dag_id, - from_date=from_date, - to_date=to_date, - max_active_runs=max_active_runs, - reverse=reverse, - dag_run_conf=dag_run_conf, - ) - return backfill_schema.dump(backfill_obj) + try: + backfill_obj = _create_backfill( + dag_id=dag_id, + from_date=cast(DateTime, pendulum.parse(from_date)), + to_date=cast(DateTime, pendulum.parse(to_date)), + max_active_runs=max_active_runs, + reverse=reverse, + dag_run_conf=dag_run_conf, + ) + return backfill_schema.dump(backfill_obj) + except AlreadyRunningBackfill: + raise Conflict(f"There is already a running backfill for dag {dag_id}") diff --git a/airflow/models/backfill.py b/airflow/models/backfill.py index 8ff2541353688..6d3a8ee4fa922 100644 --- a/airflow/models/backfill.py +++ b/airflow/models/backfill.py @@ -15,15 +15,40 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. +""" +Internal classes for management of dag backfills. + +:meta private: +""" + from __future__ import annotations -from sqlalchemy import Boolean, Column, Integer, UniqueConstraint +import logging +from typing import TYPE_CHECKING + +from sqlalchemy import Boolean, Column, ForeignKeyConstraint, Integer, UniqueConstraint, func, select +from sqlalchemy.orm import relationship from sqlalchemy_jsonfield import JSONField +from airflow.api_connexion.exceptions import NotFound +from airflow.exceptions import AirflowException from airflow.models.base import Base, StringID +from airflow.models.serialized_dag import SerializedDagModel from airflow.settings import json from airflow.utils import timezone +from airflow.utils.session import create_session from airflow.utils.sqlalchemy import UtcDateTime +from airflow.utils.state import DagRunState +from airflow.utils.types import DagRunTriggeredByType, DagRunType + +if TYPE_CHECKING: + from pendulum import DateTime + +log = logging.getLogger(__name__) + + +class AlreadyRunningBackfill(AirflowException): + """Raised when attempting to create backfill and one already active.""" class Backfill(Base): @@ -47,6 +72,11 @@ class Backfill(Base): completed_at = Column(UtcDateTime, nullable=True) updated_at = Column(UtcDateTime, default=timezone.utcnow, onupdate=timezone.utcnow, nullable=False) + backfill_dag_run_associations = relationship("BackfillDagRun", back_populates="backfill") + + def __repr__(self): + return f"Backfill({self.dag_id=}, {self.from_date=}, {self.to_date=})" + class BackfillDagRun(Base): """Mapping table between backfill run and dag run.""" @@ -59,4 +89,97 @@ class BackfillDagRun(Base): ) # the run might already exist; we could store the reason we did not create sort_ordinal = Column(Integer, nullable=False) - __table_args__ = (UniqueConstraint("backfill_id", "dag_run_id", name="ix_bdr_backfill_id_dag_run_id"),) + backfill = relationship("Backfill", back_populates="backfill_dag_run_associations") + dag_run = relationship("DagRun") + + __table_args__ = ( + UniqueConstraint("backfill_id", "dag_run_id", name="ix_bdr_backfill_id_dag_run_id"), + ForeignKeyConstraint( + [backfill_id], + ["backfill.id"], + name="bdr_backfill_fkey", + ondelete="cascade", + ), + ForeignKeyConstraint( + [dag_run_id], + ["dag_run.id"], + name="bdr_dag_run_fkey", + ondelete="set null", + ), + ) + + +def _create_backfill( + *, + dag_id: str, + from_date: DateTime, + to_date: DateTime, + max_active_runs: int, + reverse: bool, + dag_run_conf: dict | None, +) -> Backfill | None: + with create_session() as session: + serdag = session.get(SerializedDagModel, dag_id) + if not serdag: + raise NotFound(f"Could not find dag {dag_id}") + + num_active = session.scalar( + select(func.count()).where(Backfill.dag_id == dag_id, Backfill.completed_at.is_(None)) + ) + if num_active > 0: + raise AlreadyRunningBackfill( + f"Another backfill is running for dag {dag_id}. " + f"There can be only one running backfill per dag." + ) + + br = Backfill( + dag_id=dag_id, + from_date=from_date, + to_date=to_date, + max_active_runs=max_active_runs, + dag_run_conf=dag_run_conf, + ) + session.add(br) + session.commit() + + dag = serdag.dag + depends_on_past = any(x.depends_on_past for x in dag.tasks) + if depends_on_past: + if reverse is True: + raise ValueError( + "Backfill cannot be run in reverse when the dag has tasks where depends_on_past=True" + ) + + backfill_sort_ordinal = 0 + dagrun_info_list = dag.iter_dagrun_infos_between(from_date, to_date) + if reverse: + dagrun_info_list = reversed([x for x in dag.iter_dagrun_infos_between(from_date, to_date)]) + for info in dagrun_info_list: + backfill_sort_ordinal += 1 + log.info("creating backfill dag run %s dag_id=%s backfill_id=%s, info=", dag.dag_id, br.id, info) + dr = None + try: + dr = dag.create_dagrun( + triggered_by=DagRunTriggeredByType.BACKFILL, + execution_date=info.logical_date, + data_interval=info.data_interval, + start_date=timezone.utcnow(), + state=DagRunState.QUEUED, + external_trigger=False, + conf=br.dag_run_conf, + run_type=DagRunType.BACKFILL_JOB, + creating_job_id=None, + session=session, + ) + except Exception: + dag.log.exception("something failed") + session.rollback() + session.add( + BackfillDagRun( + backfill_id=br.id, + dag_run_id=dr.id if dr else None, # this means we failed to create the dag run + sort_ordinal=backfill_sort_ordinal, + ) + ) + session.commit() + return br diff --git a/airflow/utils/types.py b/airflow/utils/types.py index a19b2534b03fb..80ee1d644d4d2 100644 --- a/airflow/utils/types.py +++ b/airflow/utils/types.py @@ -119,3 +119,4 @@ class DagRunTriggeredByType(enum.Enum): TEST = "test" # for dag.test() TIMETABLE = "timetable" # for timetable based triggering DATASET = "dataset" # for dataset_triggered run type + BACKFILL = "backfill" diff --git a/tests/api_connexion/endpoints/test_backfill_endpoint.py b/tests/api_connexion/endpoints/test_backfill_endpoint.py index 07b2a3fd56c2d..dd086339b73ac 100644 --- a/tests/api_connexion/endpoints/test_backfill_endpoint.py +++ b/tests/api_connexion/endpoints/test_backfill_endpoint.py @@ -27,14 +27,13 @@ from airflow.models import DagBag, DagModel from airflow.models.backfill import Backfill from airflow.models.dag import DAG -from airflow.models.serialized_dag import SerializedDagModel from airflow.operators.empty import EmptyOperator from airflow.utils import timezone from airflow.utils.session import provide_session from tests.test_utils.api_connexion_utils import create_user, delete_user from tests.test_utils.db import clear_db_backfills, clear_db_dags, clear_db_runs, clear_db_serialized_dags -pytestmark = [pytest.mark.db_test] +pytestmark = [pytest.mark.db_test, pytest.mark.need_serialized_dag] DAG_ID = "test_dag" @@ -44,6 +43,20 @@ UTC_JSON_REPR = "UTC" if pendulum.__version__.startswith("3") else "Timezone('UTC')" +def _clean_db(): + clear_db_backfills() + clear_db_runs() + clear_db_dags() + clear_db_serialized_dags() + + +@pytest.fixture(autouse=True) +def clean_db(): + _clean_db() + yield + _clean_db() + + @pytest.fixture(scope="module") def configured_app(minimal_app_for_api): app = minimal_app_for_api @@ -83,25 +96,14 @@ def configured_app(minimal_app_for_api): class TestBackfillEndpoint: - @staticmethod - def clean_db(): - clear_db_backfills() - clear_db_runs() - clear_db_dags() - clear_db_serialized_dags() - @pytest.fixture(autouse=True) def setup_attrs(self, configured_app) -> None: - self.clean_db() self.app = configured_app self.client = self.app.test_client() # type:ignore self.dag_id = DAG_ID self.dag2_id = DAG2_ID self.dag3_id = DAG3_ID - def teardown_method(self) -> None: - self.clean_db() - @provide_session def _create_dag_models(self, *, count=1, dag_id_prefix="TEST_DAG", is_paused=False, session=None): dags = [] @@ -258,8 +260,6 @@ class TestCreateBackfill(TestBackfillEndpoint): def test_create_backfill(self, user, expected, session, dag_maker): with dag_maker(session=session, dag_id="TEST_DAG_1", schedule="0 * * * *") as dag: EmptyOperator(task_id="mytask") - session.add(SerializedDagModel(dag)) - session.commit() session.query(DagModel).all() from_date = pendulum.parse("2024-01-01") from_date_iso = from_date.isoformat() diff --git a/tests/models/test_backfill.py b/tests/models/test_backfill.py new file mode 100644 index 0000000000000..9a845f86803e0 --- /dev/null +++ b/tests/models/test_backfill.py @@ -0,0 +1,152 @@ +# 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 + +from contextlib import nullcontext + +import pendulum +import pytest +from sqlalchemy import select + +from airflow.models import DagRun +from airflow.models.backfill import AlreadyRunningBackfill, Backfill, BackfillDagRun, _create_backfill +from airflow.operators.python import PythonOperator +from airflow.utils.state import DagRunState +from tests.test_utils.db import clear_db_backfills, clear_db_dags, clear_db_runs, clear_db_serialized_dags + +pytestmark = [pytest.mark.db_test, pytest.mark.need_serialized_dag] + + +def _clean_db(): + clear_db_backfills() + clear_db_runs() + clear_db_dags() + clear_db_serialized_dags() + + +@pytest.fixture(autouse=True) +def clean_db(): + _clean_db() + yield + _clean_db() + + +@pytest.mark.parametrize("dep_on_past", [True, False]) +def test_reverse_and_depends_on_past_fails(dep_on_past, dag_maker, session): + with dag_maker() as dag: + PythonOperator(task_id="hi", python_callable=print, depends_on_past=dep_on_past) + session.commit() + cm = nullcontext() + if dep_on_past: + cm = pytest.raises(ValueError, match="cannot be run in reverse") + b = None + with cm: + b = _create_backfill( + dag_id=dag.dag_id, + from_date=pendulum.parse("2021-01-01"), + to_date=pendulum.parse("2021-01-05"), + max_active_runs=2, + reverse=True, + dag_run_conf={}, + ) + if dep_on_past: + assert b is None + else: + assert b is not None + + +@pytest.mark.parametrize("reverse", [True, False]) +def test_simple(reverse, dag_maker, session): + """ + Verify simple case behavior. + + This test verifies that runs in the range are created according + to schedule intervals, and the sort ordinal is correct. Also verifies + that dag runs are created in the queued state. + """ + with dag_maker(schedule="@daily") as dag: + PythonOperator(task_id="hi", python_callable=print) + b = _create_backfill( + dag_id=dag.dag_id, + from_date=pendulum.parse("2021-01-01"), + to_date=pendulum.parse("2021-01-05"), + max_active_runs=2, + reverse=reverse, + dag_run_conf={}, + ) + query = ( + select(DagRun) + .join(BackfillDagRun.dag_run) + .where(BackfillDagRun.backfill_id == b.id) + .order_by(BackfillDagRun.sort_ordinal) + ) + dag_runs = session.scalars(query).all() + dates = [str(x.logical_date.date()) for x in dag_runs] + expected_dates = ["2021-01-01", "2021-01-02", "2021-01-03", "2021-01-04", "2021-01-05"] + if reverse: + expected_dates = list(reversed(expected_dates)) + assert dates == expected_dates + assert all(x.state == DagRunState.QUEUED for x in dag_runs) + + +def test_params_stored_correctly(dag_maker, session): + with dag_maker(schedule="@daily") as dag: + PythonOperator(task_id="hi", python_callable=print) + b = _create_backfill( + dag_id=dag.dag_id, + from_date=pendulum.parse("2021-01-01"), + to_date=pendulum.parse("2021-01-05"), + max_active_runs=263, + reverse=False, + dag_run_conf={"this": "param"}, + ) + session.expunge_all() + b_stored = session.get(Backfill, b.id) + assert all( + ( + b_stored.dag_id == b.dag_id, + b_stored.from_date == b.from_date, + b_stored.to_date == b.to_date, + b_stored.max_active_runs == b.max_active_runs, + b_stored.dag_run_conf == b.dag_run_conf, + ) + ) + + +def test_active_dag_run(dag_maker, session): + with dag_maker(schedule="@daily") as dag: + PythonOperator(task_id="hi", python_callable=print) + session.commit() + b1 = _create_backfill( + dag_id=dag.dag_id, + from_date=pendulum.parse("2021-01-01"), + to_date=pendulum.parse("2021-01-05"), + max_active_runs=10, + reverse=False, + dag_run_conf={"this": "param"}, + ) + assert b1 is not None + with pytest.raises(AlreadyRunningBackfill, match="Another backfill is running for dag"): + _create_backfill( + dag_id=dag.dag_id, + from_date=pendulum.parse("2021-02-01"), + to_date=pendulum.parse("2021-02-05"), + max_active_runs=10, + reverse=False, + dag_run_conf={"this": "param"}, + ) From b155f458a6f1b1db5e137ae4d4c42c0454807e2e Mon Sep 17 00:00:00 2001 From: Daniel Standish <15932138+dstandish@users.noreply.github.com> Date: Tue, 1 Oct 2024 19:23:44 -0700 Subject: [PATCH 099/802] All executors should inherit from BaseExecutor (#41904) --- .../executors/celery_kubernetes_executor.py | 34 +++++++++++++++---- .../executors/local_kubernetes_executor.py | 34 +++++++++++++++---- 2 files changed, 56 insertions(+), 12 deletions(-) diff --git a/airflow/providers/celery/executors/celery_kubernetes_executor.py b/airflow/providers/celery/executors/celery_kubernetes_executor.py index bc2ed7904f5a5..acd1afcba995a 100644 --- a/airflow/providers/celery/executors/celery_kubernetes_executor.py +++ b/airflow/providers/celery/executors/celery_kubernetes_executor.py @@ -21,6 +21,7 @@ from typing import TYPE_CHECKING, Sequence from airflow.configuration import conf +from airflow.executors.base_executor import BaseExecutor from airflow.providers.celery.executors.celery_executor import CeleryExecutor try: @@ -30,18 +31,21 @@ raise AirflowOptionalProviderFeatureException(e) -from airflow.utils.log.logging_mixin import LoggingMixin from airflow.utils.providers_configuration_loader import providers_configuration_loaded if TYPE_CHECKING: from airflow.callbacks.base_callback_sink import BaseCallbackSink from airflow.callbacks.callback_requests import CallbackRequest - from airflow.executors.base_executor import CommandType, EventBufferValueType, QueuedTaskInstanceType + from airflow.executors.base_executor import ( + CommandType, + EventBufferValueType, + QueuedTaskInstanceType, + ) from airflow.models.taskinstance import SimpleTaskInstance, TaskInstance from airflow.models.taskinstancekey import TaskInstanceKey -class CeleryKubernetesExecutor(LoggingMixin): +class CeleryKubernetesExecutor(BaseExecutor): """ CeleryKubernetesExecutor consists of CeleryExecutor and KubernetesExecutor. @@ -71,11 +75,21 @@ def kubernetes_queue(self) -> str: def __init__(self, celery_executor: CeleryExecutor, kubernetes_executor: KubernetesExecutor): super().__init__() - self._job_id: int | None = None + self._job_id: int | str | None = None self.celery_executor = celery_executor self.kubernetes_executor = kubernetes_executor self.kubernetes_executor.kubernetes_queue = self.kubernetes_queue + @property + def _task_event_logs(self): + self.celery_executor._task_event_logs += self.kubernetes_executor._task_event_logs + self.kubernetes_executor._task_event_logs.clear() + return self.celery_executor._task_event_logs + + @_task_event_logs.setter + def _task_event_logs(self, value): + """Not implemented for hybrid executors.""" + @property def queued_tasks(self) -> dict[TaskInstanceKey, QueuedTaskInstanceType]: """Return queued tasks from celery and kubernetes executor.""" @@ -84,13 +98,21 @@ def queued_tasks(self) -> dict[TaskInstanceKey, QueuedTaskInstanceType]: return queued_tasks + @queued_tasks.setter + def queued_tasks(self, value) -> None: + """Not implemented for hybrid executors.""" + @property def running(self) -> set[TaskInstanceKey]: """Return running tasks from celery and kubernetes executor.""" return self.celery_executor.running.union(self.kubernetes_executor.running) + @running.setter + def running(self, value) -> None: + """Not implemented for hybrid executors.""" + @property - def job_id(self) -> int | None: + def job_id(self) -> int | str | None: """ Inherited attribute from BaseExecutor. @@ -100,7 +122,7 @@ def job_id(self) -> int | None: return self._job_id @job_id.setter - def job_id(self, value: int | None) -> None: + def job_id(self, value: int | str | None) -> None: """Expose job ID for SchedulerJob.""" self._job_id = value self.kubernetes_executor.job_id = value diff --git a/airflow/providers/cncf/kubernetes/executors/local_kubernetes_executor.py b/airflow/providers/cncf/kubernetes/executors/local_kubernetes_executor.py index 75de1101c59ba..63755d3d11a1c 100644 --- a/airflow/providers/cncf/kubernetes/executors/local_kubernetes_executor.py +++ b/airflow/providers/cncf/kubernetes/executors/local_kubernetes_executor.py @@ -20,18 +20,22 @@ from typing import TYPE_CHECKING, Sequence from airflow.configuration import conf +from airflow.executors.base_executor import BaseExecutor from airflow.providers.cncf.kubernetes.executors.kubernetes_executor import KubernetesExecutor -from airflow.utils.log.logging_mixin import LoggingMixin if TYPE_CHECKING: from airflow.callbacks.base_callback_sink import BaseCallbackSink from airflow.callbacks.callback_requests import CallbackRequest - from airflow.executors.base_executor import CommandType, EventBufferValueType, QueuedTaskInstanceType + from airflow.executors.base_executor import ( + CommandType, + EventBufferValueType, + QueuedTaskInstanceType, + ) from airflow.executors.local_executor import LocalExecutor from airflow.models.taskinstance import SimpleTaskInstance, TaskInstance, TaskInstanceKey -class LocalKubernetesExecutor(LoggingMixin): +class LocalKubernetesExecutor(BaseExecutor): """ Chooses between LocalExecutor and KubernetesExecutor based on the queue defined on the task. @@ -57,11 +61,21 @@ class LocalKubernetesExecutor(LoggingMixin): def __init__(self, local_executor: LocalExecutor, kubernetes_executor: KubernetesExecutor): super().__init__() - self._job_id: str | None = None + self._job_id: int | str | None = None self.local_executor = local_executor self.kubernetes_executor = kubernetes_executor self.kubernetes_executor.kubernetes_queue = self.KUBERNETES_QUEUE + @property + def _task_event_logs(self): + self.local_executor._task_event_logs += self.kubernetes_executor._task_event_logs + self.kubernetes_executor._task_event_logs.clear() + return self.local_executor._task_event_logs + + @_task_event_logs.setter + def _task_event_logs(self, value): + """Not implemented for hybrid executors.""" + @property def queued_tasks(self) -> dict[TaskInstanceKey, QueuedTaskInstanceType]: """Return queued tasks from local and kubernetes executor.""" @@ -70,13 +84,21 @@ def queued_tasks(self) -> dict[TaskInstanceKey, QueuedTaskInstanceType]: return queued_tasks + @queued_tasks.setter + def queued_tasks(self, value) -> None: + """Not implemented for hybrid executors.""" + @property def running(self) -> set[TaskInstanceKey]: """Return running tasks from local and kubernetes executor.""" return self.local_executor.running.union(self.kubernetes_executor.running) + @running.setter + def running(self, value) -> None: + """Not implemented for hybrid executors.""" + @property - def job_id(self) -> str | None: + def job_id(self) -> int | str | None: """ Inherited attribute from BaseExecutor. @@ -86,7 +108,7 @@ def job_id(self) -> str | None: return self._job_id @job_id.setter - def job_id(self, value: str | None) -> None: + def job_id(self, value: int | str | None) -> None: """Expose job ID for SchedulerJob.""" self._job_id = value self.kubernetes_executor.job_id = value From 85b4cdf11834c29f054dff97415f39f9f71807fa Mon Sep 17 00:00:00 2001 From: Daniel Standish <15932138+dstandish@users.noreply.github.com> Date: Tue, 1 Oct 2024 21:35:14 -0700 Subject: [PATCH 100/802] Revert "Fix the order of tasks during serialization (#42219)" (#42646) This reverts commit adb9466bd7ce1c92e51f11a90d39fd557c99dc5b a.k.a. PR #42219. Was causing tests to fail. --- airflow/serialization/serialized_objects.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/airflow/serialization/serialized_objects.py b/airflow/serialization/serialized_objects.py index 08944391b8166..a4801b767acc5 100644 --- a/airflow/serialization/serialized_objects.py +++ b/airflow/serialization/serialized_objects.py @@ -1604,9 +1604,7 @@ def serialize_dag(cls, dag: DAG) -> dict: try: serialized_dag = cls.serialize_to_json(dag, cls._decorated_fields) serialized_dag["_processor_dags_folder"] = DAGS_FOLDER - serialized_dag["tasks"] = [ - cls.serialize(dag.task_dict[task_id]) for task_id in sorted(dag.task_dict) - ] + serialized_dag["tasks"] = [cls.serialize(task) for _, task in dag.task_dict.items()] dag_deps = [ dep From dad6f23487fd6ab7d7fcbaa0b00fd03664d9a8d9 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Wed, 2 Oct 2024 02:20:43 -0400 Subject: [PATCH 101/802] Check pool_slots on partial task import instead of execution (#39724) Co-authored-by: Ryan Hatter <25823361+RNHTTR@users.noreply.github.com> --- airflow/decorators/base.py | 6 ++++++ airflow/models/baseoperator.py | 5 +++++ tests/models/test_mappedoperator.py | 9 +++++++++ 3 files changed, 20 insertions(+) diff --git a/airflow/decorators/base.py b/airflow/decorators/base.py index e650c1920a870..bb9602d50c1cd 100644 --- a/airflow/decorators/base.py +++ b/airflow/decorators/base.py @@ -468,6 +468,12 @@ def _expand(self, expand_input: ExpandInput, *, strict: bool) -> XComArg: end_date = timezone.convert_to_utc(partial_kwargs.pop("end_date", None)) if partial_kwargs.get("pool") is None: partial_kwargs["pool"] = Pool.DEFAULT_POOL_NAME + if "pool_slots" in partial_kwargs: + if partial_kwargs["pool_slots"] < 1: + dag_str = "" + if dag: + dag_str = f" in dag {dag.dag_id}" + raise ValueError(f"pool slots for {task_id}{dag_str} cannot be less than 1") partial_kwargs["retries"] = parse_retries(partial_kwargs.get("retries", DEFAULT_RETRIES)) partial_kwargs["retry_delay"] = coerce_timedelta( partial_kwargs.get("retry_delay", DEFAULT_RETRY_DELAY), diff --git a/airflow/models/baseoperator.py b/airflow/models/baseoperator.py index 20656586ba01e..9e0c8e1e69b61 100644 --- a/airflow/models/baseoperator.py +++ b/airflow/models/baseoperator.py @@ -358,6 +358,11 @@ def partial( partial_kwargs["end_date"] = timezone.convert_to_utc(partial_kwargs["end_date"]) if partial_kwargs["pool"] is None: partial_kwargs["pool"] = Pool.DEFAULT_POOL_NAME + if partial_kwargs["pool_slots"] < 1: + dag_str = "" + if dag: + dag_str = f" in dag {dag.dag_id}" + raise ValueError(f"pool slots for {task_id}{dag_str} cannot be less than 1") partial_kwargs["retries"] = parse_retries(partial_kwargs["retries"]) partial_kwargs["retry_delay"] = coerce_timedelta(partial_kwargs["retry_delay"], key="retry_delay") if partial_kwargs["max_retry_delay"] is not None: diff --git a/tests/models/test_mappedoperator.py b/tests/models/test_mappedoperator.py index 2b0cd50165c45..0571e07e671f8 100644 --- a/tests/models/test_mappedoperator.py +++ b/tests/models/test_mappedoperator.py @@ -220,6 +220,15 @@ def test_partial_on_class_invalid_ctor_args() -> None: MockOperator.partial(task_id="a", foo="bar", bar=2) +def test_partial_on_invalid_pool_slots_raises() -> None: + """Test that when we pass an invalid value to pool_slots in partial(), + + i.e. if the value is not an integer, an error is raised at import time.""" + + with pytest.raises(TypeError, match="'<' not supported between instances of 'str' and 'int'"): + MockOperator.partial(task_id="pool_slots_test", pool="test", pool_slots="a").expand(arg1=[1, 2, 3]) + + @pytest.mark.skip_if_database_isolation_mode # Does not work in db isolation mode @pytest.mark.parametrize( ["num_existing_tis", "expected"], From ac592d87a9120c28454656dc126025dece769139 Mon Sep 17 00:00:00 2001 From: TakawaAkirayo <153728772+TakawaAkirayo@users.noreply.github.com> Date: Wed, 2 Oct 2024 14:51:08 +0800 Subject: [PATCH 102/802] Add retry logic in the scheduler for updating trigger timeouts in case of deadlocks. (#41429) * Add retry in update trigger timeout * add ut for these cases * use OperationalError in ut to describe deadlock scenarios * [MINOR] add newsfragment for this PR * [MINOR] refactor UT for mypy check --- airflow/jobs/scheduler_job_runner.py | 36 +++++++------ newsfragments/41429.improvement.rst | 1 + tests/jobs/test_scheduler_job.py | 78 +++++++++++++++++++++++++++- 3 files changed, 98 insertions(+), 17 deletions(-) create mode 100644 newsfragments/41429.improvement.rst diff --git a/airflow/jobs/scheduler_job_runner.py b/airflow/jobs/scheduler_job_runner.py index 242154820df9e..de6ce5019b9de 100644 --- a/airflow/jobs/scheduler_job_runner.py +++ b/airflow/jobs/scheduler_job_runner.py @@ -1884,23 +1884,27 @@ def adopt_or_reset_orphaned_tasks(self, session: Session = NEW_SESSION) -> int: return len(to_reset) @provide_session - def check_trigger_timeouts(self, session: Session = NEW_SESSION) -> None: + def check_trigger_timeouts( + self, max_retries: int = MAX_DB_RETRIES, session: Session = NEW_SESSION + ) -> None: """Mark any "deferred" task as failed if the trigger or execution timeout has passed.""" - num_timed_out_tasks = session.execute( - update(TI) - .where( - TI.state == TaskInstanceState.DEFERRED, - TI.trigger_timeout < timezone.utcnow(), - ) - .values( - state=TaskInstanceState.SCHEDULED, - next_method="__fail__", - next_kwargs={"error": "Trigger/execution timeout"}, - trigger_id=None, - ) - ).rowcount - if num_timed_out_tasks: - self.log.info("Timed out %i deferred tasks without fired triggers", num_timed_out_tasks) + for attempt in run_with_db_retries(max_retries, logger=self.log): + with attempt: + num_timed_out_tasks = session.execute( + update(TI) + .where( + TI.state == TaskInstanceState.DEFERRED, + TI.trigger_timeout < timezone.utcnow(), + ) + .values( + state=TaskInstanceState.SCHEDULED, + next_method="__fail__", + next_kwargs={"error": "Trigger/execution timeout"}, + trigger_id=None, + ) + ).rowcount + if num_timed_out_tasks: + self.log.info("Timed out %i deferred tasks without fired triggers", num_timed_out_tasks) # [START find_zombies] def _find_zombies(self) -> None: diff --git a/newsfragments/41429.improvement.rst b/newsfragments/41429.improvement.rst new file mode 100644 index 0000000000000..6d04d5dfe61af --- /dev/null +++ b/newsfragments/41429.improvement.rst @@ -0,0 +1 @@ +Add ``run_with_db_retries`` when the scheduler updates the deferred Task as failed to tolerate database deadlock issues. diff --git a/tests/jobs/test_scheduler_job.py b/tests/jobs/test_scheduler_job.py index 32662d7d873db..40a7220698407 100644 --- a/tests/jobs/test_scheduler_job.py +++ b/tests/jobs/test_scheduler_job.py @@ -148,7 +148,7 @@ def clean_db(): @pytest.fixture(autouse=True) def per_test(self) -> Generator: self.clean_db() - self.job_runner = None + self.job_runner: SchedulerJobRunner | None = None yield @@ -5192,6 +5192,82 @@ def test_timeout_triggers(self, dag_maker): assert ti1.next_method == "__fail__" assert ti2.state == State.DEFERRED + def test_retry_on_db_error_when_update_timeout_triggers(self, dag_maker): + """ + Tests that it will retry on DB error like deadlock when updating timeout triggers. + """ + from sqlalchemy.exc import OperationalError + + retry_times = 3 + + session = settings.Session() + # Create the test DAG and task + with dag_maker( + dag_id="test_retry_on_db_error_when_update_timeout_triggers", + start_date=DEFAULT_DATE, + schedule="@once", + max_active_runs=1, + session=session, + ): + EmptyOperator(task_id="dummy1") + + # Mock the db failure within retry times + might_fail_session = MagicMock(wraps=session) + + def check_if_trigger_timeout(max_retries: int): + def make_side_effect(): + call_count = 0 + + def side_effect(*args, **kwargs): + nonlocal call_count + if call_count < retry_times - 1: + call_count += 1 + raise OperationalError("any_statement", "any_params", "any_orig") + else: + return session.execute(*args, **kwargs) + + return side_effect + + might_fail_session.execute.side_effect = make_side_effect() + + try: + # Create a Task Instance for the task that is allegedly deferred + # but past its timeout, and one that is still good. + # We don't actually need a linked trigger here; the code doesn't check. + dr1 = dag_maker.create_dagrun() + dr2 = dag_maker.create_dagrun( + run_id="test2", execution_date=DEFAULT_DATE + datetime.timedelta(seconds=1) + ) + ti1 = dr1.get_task_instance("dummy1", session) + ti2 = dr2.get_task_instance("dummy1", session) + ti1.state = State.DEFERRED + ti1.trigger_timeout = timezone.utcnow() - datetime.timedelta(seconds=60) + ti2.state = State.DEFERRED + ti2.trigger_timeout = timezone.utcnow() + datetime.timedelta(seconds=60) + session.flush() + + # Boot up the scheduler and make it check timeouts + scheduler_job = Job() + self.job_runner = SchedulerJobRunner(job=scheduler_job, subdir=os.devnull) + + self.job_runner.check_trigger_timeouts(max_retries=max_retries, session=might_fail_session) + + # Make sure that TI1 is now scheduled to fail, and 2 wasn't touched + session.refresh(ti1) + session.refresh(ti2) + assert ti1.state == State.SCHEDULED + assert ti1.next_method == "__fail__" + assert ti2.state == State.DEFERRED + finally: + self.clean_db() + + # Positive case, will retry until success before reach max retry times + check_if_trigger_timeout(retry_times) + + # Negative case: no retries, execute only once. + with pytest.raises(OperationalError): + check_if_trigger_timeout(1) + def test_find_zombies_nothing(self): executor = MockExecutor(do_update=False) scheduler_job = Job(executor=executor) From 82a89470834d01092a133064e4593a00e3d27ba8 Mon Sep 17 00:00:00 2001 From: GPK Date: Wed, 2 Oct 2024 07:59:17 +0100 Subject: [PATCH 103/802] Fix consistent return response from PubSubPullSensor (#42080) * fix consistent return response pubsubsensor * removed messages_callback argument to pubsub trigger and using it in execute_complete * updated variable name * updates as per comments, added return types and refactored logic * update types, tests and use inherit exception --- .../providers/google/cloud/sensors/pubsub.py | 24 +++++++- .../providers/google/cloud/triggers/pubsub.py | 22 +++----- .../google/cloud/sensors/test_pubsub.py | 48 ++++++++++++++++ .../google/cloud/triggers/test_pubsub.py | 55 ++++++++++++++++++- 4 files changed, 129 insertions(+), 20 deletions(-) diff --git a/airflow/providers/google/cloud/sensors/pubsub.py b/airflow/providers/google/cloud/sensors/pubsub.py index cb224d42979b7..aa74411f072e5 100644 --- a/airflow/providers/google/cloud/sensors/pubsub.py +++ b/airflow/providers/google/cloud/sensors/pubsub.py @@ -22,6 +22,7 @@ from datetime import timedelta from typing import TYPE_CHECKING, Any, Callable, Sequence +from google.cloud import pubsub_v1 from google.cloud.pubsub_v1.types import ReceivedMessage from airflow.configuration import conf @@ -34,6 +35,10 @@ from airflow.utils.context import Context +class PubSubMessageTransformException(AirflowException): + """Raise when messages failed to convert pubsub received format.""" + + class PubSubPullSensor(BaseSensorOperator): """ Pulls messages from a PubSub subscription and passes them through XCom. @@ -170,7 +175,6 @@ def execute(self, context: Context) -> None: subscription=self.subscription, max_messages=self.max_messages, ack_messages=self.ack_messages, - messages_callback=self.messages_callback, poke_interval=self.poke_interval, gcp_conn_id=self.gcp_conn_id, impersonation_chain=self.impersonation_chain, @@ -178,14 +182,28 @@ def execute(self, context: Context) -> None: method_name="execute_complete", ) - def execute_complete(self, context: dict[str, Any], event: dict[str, str | list[str]]) -> str | list[str]: - """Return immediately and relies on trigger to throw a success event. Callback for the trigger.""" + def execute_complete(self, context: Context, event: dict[str, str | list[str]]) -> Any: + """If messages_callback is provided, execute it; otherwise, return immediately with trigger event message.""" if event["status"] == "success": self.log.info("Sensor pulls messages: %s", event["message"]) + if self.messages_callback: + received_messages = self._convert_to_received_messages(event["message"]) + _return_value = self.messages_callback(received_messages, context) + return _return_value + return event["message"] self.log.info("Sensor failed: %s", event["message"]) raise AirflowException(event["message"]) + def _convert_to_received_messages(self, messages: Any) -> list[ReceivedMessage]: + try: + received_messages = [pubsub_v1.types.ReceivedMessage(msg) for msg in messages] + return received_messages + except Exception as e: + raise PubSubMessageTransformException( + f"Error converting triggerer event message back to received message format: {e}" + ) + def _default_message_callback( self, pulled_messages: list[ReceivedMessage], diff --git a/airflow/providers/google/cloud/triggers/pubsub.py b/airflow/providers/google/cloud/triggers/pubsub.py index 535bfe2ba1c68..db3fe409e942b 100644 --- a/airflow/providers/google/cloud/triggers/pubsub.py +++ b/airflow/providers/google/cloud/triggers/pubsub.py @@ -19,16 +19,13 @@ from __future__ import annotations import asyncio -from typing import TYPE_CHECKING, Any, AsyncIterator, Callable, Sequence +from typing import Any, AsyncIterator, Sequence + +from google.cloud.pubsub_v1.types import ReceivedMessage from airflow.providers.google.cloud.hooks.pubsub import PubSubAsyncHook from airflow.triggers.base import BaseTrigger, TriggerEvent -if TYPE_CHECKING: - from google.cloud.pubsub_v1.types import ReceivedMessage - - from airflow.utils.context import Context - class PubsubPullTrigger(BaseTrigger): """ @@ -41,11 +38,6 @@ class PubsubPullTrigger(BaseTrigger): :param ack_messages: If True, each message will be acknowledged immediately rather than by any downstream tasks :param gcp_conn_id: Reference to google cloud connection id - :param messages_callback: (Optional) Callback to process received messages. - Its return value will be saved to XCom. - If you are pulling large messages, you probably want to provide a custom callback. - If not provided, the default implementation will convert `ReceivedMessage` objects - into JSON-serializable dicts using `google.protobuf.json_format.MessageToDict` function. :param poke_interval: polling period in seconds to check for the status :param impersonation_chain: Optional service account to impersonate using short-term credentials, or chained list of accounts required to get the access_token @@ -64,7 +56,6 @@ def __init__( max_messages: int, ack_messages: bool, gcp_conn_id: str, - messages_callback: Callable[[list[ReceivedMessage], Context], Any] | None = None, poke_interval: float = 10.0, impersonation_chain: str | Sequence[str] | None = None, ): @@ -73,7 +64,6 @@ def __init__( self.subscription = subscription self.max_messages = max_messages self.ack_messages = ack_messages - self.messages_callback = messages_callback self.poke_interval = poke_interval self.gcp_conn_id = gcp_conn_id self.impersonation_chain = impersonation_chain @@ -88,7 +78,6 @@ def serialize(self) -> tuple[str, dict[str, Any]]: "subscription": self.subscription, "max_messages": self.max_messages, "ack_messages": self.ack_messages, - "messages_callback": self.messages_callback, "poke_interval": self.poke_interval, "gcp_conn_id": self.gcp_conn_id, "impersonation_chain": self.impersonation_chain, @@ -106,7 +95,10 @@ async def run(self) -> AsyncIterator[TriggerEvent]: # type: ignore[override] ): if self.ack_messages: await self.message_acknowledgement(pulled_messages) - yield TriggerEvent({"status": "success", "message": pulled_messages}) + + messages_json = [ReceivedMessage.to_dict(m) for m in pulled_messages] + + yield TriggerEvent({"status": "success", "message": messages_json}) return self.log.info("Sleeping for %s seconds.", self.poke_interval) await asyncio.sleep(self.poke_interval) diff --git a/tests/providers/google/cloud/sensors/test_pubsub.py b/tests/providers/google/cloud/sensors/test_pubsub.py index a77167dda3037..5a3fb170b7482 100644 --- a/tests/providers/google/cloud/sensors/test_pubsub.py +++ b/tests/providers/google/cloud/sensors/test_pubsub.py @@ -21,6 +21,7 @@ from unittest import mock import pytest +from google.cloud import pubsub_v1 from google.cloud.pubsub_v1.types import ReceivedMessage from airflow.exceptions import AirflowException, TaskDeferred @@ -197,3 +198,50 @@ def test_pubsub_pull_sensor_async_execute_complete(self): with mock.patch.object(operator.log, "info") as mock_log_info: operator.execute_complete(context={}, event={"status": "success", "message": test_message}) mock_log_info.assert_called_with("Sensor pulls messages: %s", test_message) + + @mock.patch("airflow.providers.google.cloud.sensors.pubsub.PubSubHook") + def test_pubsub_pull_sensor_async_execute_complete_use_message_callback(self, mock_hook): + test_message = [ + { + "ack_id": "UAYWLF1GSFE3GQhoUQ5PXiM_NSAoRRIJB08CKF15MU0sQVhwaFENGXJ9YHxrUxsDV0ECel1RGQdoTm11H4GglfRLQ1RrWBIHB01Vel5TEwxoX11wBnm4vPO6v8vgfwk9OpX-8tltO6ywsP9GZiM9XhJLLD5-LzlFQV5AEkwkDERJUytDCypYEU4EISE-MD5FU0Q", + "message": { + "data": "aGkgZnJvbSBjbG91ZCBjb25zb2xlIQ==", + "message_id": "12165864188103151", + "publish_time": "2024-08-28T11:49:50.962Z", + "attributes": {}, + "ordering_key": "", + }, + "delivery_attempt": 0, + } + ] + + received_messages = [pubsub_v1.types.ReceivedMessage(msg) for msg in test_message] + + messages_callback_return_value = "custom_message_from_callback" + + def messages_callback( + pulled_messages: list[ReceivedMessage], + context: dict[str, Any], + ): + assert pulled_messages == received_messages + + assert isinstance(context, dict) + for key in context.keys(): + assert isinstance(key, str) + + return messages_callback_return_value + + operator = PubSubPullSensor( + task_id="test_task", + ack_messages=True, + project_id=TEST_PROJECT, + subscription=TEST_SUBSCRIPTION, + deferrable=True, + messages_callback=messages_callback, + ) + mock_hook.return_value.pull.return_value = received_messages + + with mock.patch.object(operator.log, "info") as mock_log_info: + resp = operator.execute_complete(context={}, event={"status": "success", "message": test_message}) + mock_log_info.assert_called_with("Sensor pulls messages: %s", test_message) + assert resp == messages_callback_return_value diff --git a/tests/providers/google/cloud/triggers/test_pubsub.py b/tests/providers/google/cloud/triggers/test_pubsub.py index d2294eb61414b..e1a4e178d2918 100644 --- a/tests/providers/google/cloud/triggers/test_pubsub.py +++ b/tests/providers/google/cloud/triggers/test_pubsub.py @@ -16,9 +16,13 @@ # under the License. from __future__ import annotations +from unittest import mock + import pytest +from google.cloud.pubsub_v1.types import ReceivedMessage from airflow.providers.google.cloud.triggers.pubsub import PubsubPullTrigger +from airflow.triggers.base import TriggerEvent TEST_POLL_INTERVAL = 10 TEST_GCP_CONN_ID = "google_cloud_default" @@ -34,13 +38,25 @@ def trigger(): subscription="subscription", max_messages=MAX_MESSAGES, ack_messages=ACK_MESSAGES, - messages_callback=None, poke_interval=TEST_POLL_INTERVAL, gcp_conn_id=TEST_GCP_CONN_ID, impersonation_chain=None, ) +async def generate_messages(count: int) -> list[ReceivedMessage]: + return [ + ReceivedMessage( + ack_id=f"{i}", + message={ + "data": f"Message {i}".encode(), + "attributes": {"type": "generated message"}, + }, + ) + for i in range(1, count + 1) + ] + + class TestPubsubPullTrigger: def test_async_pubsub_pull_trigger_serialization_should_execute_successfully(self, trigger): """ @@ -54,8 +70,43 @@ def test_async_pubsub_pull_trigger_serialization_should_execute_successfully(sel "subscription": "subscription", "max_messages": MAX_MESSAGES, "ack_messages": ACK_MESSAGES, - "messages_callback": None, "poke_interval": TEST_POLL_INTERVAL, "gcp_conn_id": TEST_GCP_CONN_ID, "impersonation_chain": None, } + + @pytest.mark.asyncio + @mock.patch("airflow.providers.google.cloud.hooks.pubsub.PubSubAsyncHook.pull") + async def test_async_pubsub_pull_trigger_return_event(self, mock_pull): + mock_pull.return_value = generate_messages(1) + trigger = PubsubPullTrigger( + project_id=PROJECT_ID, + subscription="subscription", + max_messages=MAX_MESSAGES, + ack_messages=False, + poke_interval=TEST_POLL_INTERVAL, + gcp_conn_id=TEST_GCP_CONN_ID, + impersonation_chain=None, + ) + + expected_event = TriggerEvent( + { + "status": "success", + "message": [ + { + "ack_id": "1", + "message": { + "data": "TWVzc2FnZSAx", + "attributes": {"type": "generated message"}, + "message_id": "", + "ordering_key": "", + }, + "delivery_attempt": 0, + } + ], + } + ) + + response = await trigger.run().asend(None) + + assert response == expected_event From 7d54440f3ef3150e339eb1338159f2b95b04108f Mon Sep 17 00:00:00 2001 From: Tzu-ping Chung Date: Wed, 2 Oct 2024 01:35:43 -0700 Subject: [PATCH 104/802] Fix type-ignore comment for typing changes (#42656) --- tests/jobs/test_scheduler_job.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/jobs/test_scheduler_job.py b/tests/jobs/test_scheduler_job.py index 40a7220698407..97d84da9c4d58 100644 --- a/tests/jobs/test_scheduler_job.py +++ b/tests/jobs/test_scheduler_job.py @@ -5552,7 +5552,7 @@ def spy(*args, **kwargs): def watch_set_state(dr: DagRun, state, **kwargs): if state in (DagRunState.SUCCESS, DagRunState.FAILED): # Stop the scheduler - self.job_runner.num_runs = 1 # type: ignore[attr-defined] + self.job_runner.num_runs = 1 # type: ignore[union-attr] orig_set_state(dr, state, **kwargs) # type: ignore[call-arg] def watch_heartbeat(*args, **kwargs): From 9bd026d213a2ab4e8b5d11541ea135ac4ad3dbfa Mon Sep 17 00:00:00 2001 From: Bugra Ozturk Date: Wed, 2 Oct 2024 10:47:56 +0200 Subject: [PATCH 105/802] AIP-84 Migrate delete a connection to FastAPI API (#42571) * Include connections router and migrate delete a connection endpoint to fastapi * Mark tests as db_test * Use only pyfixture session * make method async * setup method to setup_attrs * Convert APIRouter tags, make setup method unified * Use AirflowRouter over fastapi.APIRouter --- .../endpoints/connection_endpoint.py | 2 + airflow/api_fastapi/openapi/v1-generated.yaml | 41 ++++++++++++ airflow/api_fastapi/views/public/__init__.py | 2 + .../api_fastapi/views/public/connections.py | 47 ++++++++++++++ airflow/ui/openapi-gen/queries/common.ts | 9 ++- airflow/ui/openapi-gen/queries/queries.ts | 45 ++++++++++++- .../ui/openapi-gen/requests/services.gen.ts | 30 +++++++++ airflow/ui/openapi-gen/requests/types.gen.ts | 33 ++++++++++ .../views/public/test_connections.py | 63 +++++++++++++++++++ 9 files changed, 270 insertions(+), 2 deletions(-) create mode 100644 airflow/api_fastapi/views/public/connections.py create mode 100644 tests/api_fastapi/views/public/test_connections.py diff --git a/airflow/api_connexion/endpoints/connection_endpoint.py b/airflow/api_connexion/endpoints/connection_endpoint.py index c17a9280d78f8..b28c9dfcafa79 100644 --- a/airflow/api_connexion/endpoints/connection_endpoint.py +++ b/airflow/api_connexion/endpoints/connection_endpoint.py @@ -40,6 +40,7 @@ from airflow.secrets.environment_variables import CONN_ENV_PREFIX from airflow.security import permissions from airflow.utils import helpers +from airflow.utils.api_migration import mark_fastapi_migration_done from airflow.utils.log.action_logger import action_event_from_permission from airflow.utils.session import NEW_SESSION, provide_session from airflow.utils.strings import get_random_string @@ -53,6 +54,7 @@ RESOURCE_EVENT_PREFIX = "connection" +@mark_fastapi_migration_done @security.requires_access_connection("DELETE") @provide_session @action_logging( diff --git a/airflow/api_fastapi/openapi/v1-generated.yaml b/airflow/api_fastapi/openapi/v1-generated.yaml index b08ef42c16df1..a54e0e4ca57dd 100644 --- a/airflow/api_fastapi/openapi/v1-generated.yaml +++ b/airflow/api_fastapi/openapi/v1-generated.yaml @@ -319,6 +319,47 @@ paths: application/json: schema: $ref: '#/components/schemas/HTTPValidationError' + /public/connections/{connection_id}: + delete: + tags: + - Connection + summary: Delete Connection + description: Delete a connection entry. + operationId: delete_connection + parameters: + - name: connection_id + in: path + required: true + schema: + type: string + title: Connection Id + responses: + '204': + description: Successful Response + '401': + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPExceptionResponse' + description: Unauthorized + '403': + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPExceptionResponse' + description: Forbidden + '404': + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPExceptionResponse' + description: Not Found + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' components: schemas: DAGCollectionResponse: diff --git a/airflow/api_fastapi/views/public/__init__.py b/airflow/api_fastapi/views/public/__init__.py index 1c2511fc82ac2..9c0eefebb875e 100644 --- a/airflow/api_fastapi/views/public/__init__.py +++ b/airflow/api_fastapi/views/public/__init__.py @@ -17,6 +17,7 @@ from __future__ import annotations +from airflow.api_fastapi.views.public.connections import connections_router from airflow.api_fastapi.views.public.dags import dags_router from airflow.api_fastapi.views.router import AirflowRouter @@ -24,3 +25,4 @@ public_router.include_router(dags_router) +public_router.include_router(connections_router) diff --git a/airflow/api_fastapi/views/public/connections.py b/airflow/api_fastapi/views/public/connections.py new file mode 100644 index 0000000000000..d418e10026796 --- /dev/null +++ b/airflow/api_fastapi/views/public/connections.py @@ -0,0 +1,47 @@ +# 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 + +from fastapi import Depends, HTTPException +from sqlalchemy import select +from sqlalchemy.orm import Session +from typing_extensions import Annotated + +from airflow.api_fastapi.db.common import get_session +from airflow.api_fastapi.openapi.exceptions import create_openapi_http_exception_doc +from airflow.api_fastapi.views.router import AirflowRouter +from airflow.models import Connection + +connections_router = AirflowRouter(tags=["Connection"]) + + +@connections_router.delete( + "/connections/{connection_id}", + status_code=204, + responses=create_openapi_http_exception_doc([401, 403, 404]), +) +async def delete_connection( + connection_id: str, + session: Annotated[Session, Depends(get_session)], +): + """Delete a connection entry.""" + connection = session.scalar(select(Connection).filter_by(conn_id=connection_id)) + + if connection is None: + raise HTTPException(404, f"The Connection with connection_id: `{connection_id}` was not found") + + session.delete(connection) diff --git a/airflow/ui/openapi-gen/queries/common.ts b/airflow/ui/openapi-gen/queries/common.ts index 96e49cc6d7673..fcddded7dc121 100644 --- a/airflow/ui/openapi-gen/queries/common.ts +++ b/airflow/ui/openapi-gen/queries/common.ts @@ -1,7 +1,11 @@ // generated with @7nohe/openapi-react-query-codegen@1.6.0 import { UseQueryResult } from "@tanstack/react-query"; -import { AssetService, DagService } from "../requests/services.gen"; +import { + AssetService, + ConnectionService, + DagService, +} from "../requests/services.gen"; import { DagRunState } from "../requests/types.gen"; export type AssetServiceNextRunAssetsDefaultResponse = Awaited< @@ -76,3 +80,6 @@ export type DagServicePatchDagsMutationResult = Awaited< export type DagServicePatchDagMutationResult = Awaited< ReturnType >; +export type ConnectionServiceDeleteConnectionMutationResult = Awaited< + ReturnType +>; diff --git a/airflow/ui/openapi-gen/queries/queries.ts b/airflow/ui/openapi-gen/queries/queries.ts index 985bf952e3eb3..f83c151b91e23 100644 --- a/airflow/ui/openapi-gen/queries/queries.ts +++ b/airflow/ui/openapi-gen/queries/queries.ts @@ -6,7 +6,11 @@ import { UseQueryOptions, } from "@tanstack/react-query"; -import { AssetService, DagService } from "../requests/services.gen"; +import { + AssetService, + ConnectionService, + DagService, +} from "../requests/services.gen"; import { DAGPatchBody, DagRunState } from "../requests/types.gen"; import * as Common from "./common"; @@ -247,3 +251,42 @@ export const useDagServicePatchDag = < }) as unknown as Promise, ...options, }); +/** + * Delete Connection + * Delete a connection entry. + * @param data The data for the request. + * @param data.connectionId + * @returns void Successful Response + * @throws ApiError + */ +export const useConnectionServiceDeleteConnection = < + TData = Common.ConnectionServiceDeleteConnectionMutationResult, + TError = unknown, + TContext = unknown, +>( + options?: Omit< + UseMutationOptions< + TData, + TError, + { + connectionId: string; + }, + TContext + >, + "mutationFn" + >, +) => + useMutation< + TData, + TError, + { + connectionId: string; + }, + TContext + >({ + mutationFn: ({ connectionId }) => + ConnectionService.deleteConnection({ + connectionId, + }) as unknown as Promise, + ...options, + }); diff --git a/airflow/ui/openapi-gen/requests/services.gen.ts b/airflow/ui/openapi-gen/requests/services.gen.ts index be216bd534c61..24c960d2b7d5f 100644 --- a/airflow/ui/openapi-gen/requests/services.gen.ts +++ b/airflow/ui/openapi-gen/requests/services.gen.ts @@ -11,6 +11,8 @@ import type { PatchDagsResponse, PatchDagData, PatchDagResponse, + DeleteConnectionData, + DeleteConnectionResponse, } from "./types.gen"; export class AssetService { @@ -159,3 +161,31 @@ export class DagService { }); } } + +export class ConnectionService { + /** + * Delete Connection + * Delete a connection entry. + * @param data The data for the request. + * @param data.connectionId + * @returns void Successful Response + * @throws ApiError + */ + public static deleteConnection( + data: DeleteConnectionData, + ): CancelablePromise { + return __request(OpenAPI, { + method: "DELETE", + url: "/public/connections/{connection_id}", + path: { + connection_id: data.connectionId, + }, + errors: { + 401: "Unauthorized", + 403: "Forbidden", + 404: "Not Found", + 422: "Validation Error", + }, + }); + } +} diff --git a/airflow/ui/openapi-gen/requests/types.gen.ts b/airflow/ui/openapi-gen/requests/types.gen.ts index e1db8310a1dc1..b38d5c00a69f3 100644 --- a/airflow/ui/openapi-gen/requests/types.gen.ts +++ b/airflow/ui/openapi-gen/requests/types.gen.ts @@ -134,6 +134,12 @@ export type PatchDagData = { export type PatchDagResponse = DAGResponse; +export type DeleteConnectionData = { + connectionId: string; +}; + +export type DeleteConnectionResponse = void; + export type $OpenApiTs = { "/ui/next_run_datasets/{dag_id}": { get: { @@ -227,4 +233,31 @@ export type $OpenApiTs = { }; }; }; + "/public/connections/{connection_id}": { + delete: { + req: DeleteConnectionData; + res: { + /** + * Successful Response + */ + 204: void; + /** + * Unauthorized + */ + 401: HTTPExceptionResponse; + /** + * Forbidden + */ + 403: HTTPExceptionResponse; + /** + * Not Found + */ + 404: HTTPExceptionResponse; + /** + * Validation Error + */ + 422: HTTPValidationError; + }; + }; + }; }; diff --git a/tests/api_fastapi/views/public/test_connections.py b/tests/api_fastapi/views/public/test_connections.py new file mode 100644 index 0000000000000..cfdca1d67984d --- /dev/null +++ b/tests/api_fastapi/views/public/test_connections.py @@ -0,0 +1,63 @@ +# 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 pytest + +from airflow.models import Connection +from airflow.utils.session import provide_session +from tests.test_utils.db import clear_db_connections + +pytestmark = pytest.mark.db_test + +TEST_CONN_ID = "test_connection_id" +TEST_CONN_TYPE = "test_type" + + +@provide_session +def _create_connection(session) -> None: + connection_model = Connection(conn_id=TEST_CONN_ID, conn_type=TEST_CONN_TYPE) + session.add(connection_model) + + +class TestConnectionEndpoint: + @pytest.fixture(autouse=True) + def setup(self) -> None: + clear_db_connections(False) + + def teardown_method(self) -> None: + clear_db_connections() + + def create_connection(self): + _create_connection() + + +class TestDeleteConnection(TestConnectionEndpoint): + def test_delete_should_respond_204(self, test_client, session): + self.create_connection() + conns = session.query(Connection).all() + assert len(conns) == 1 + response = test_client.delete(f"/public/connections/{TEST_CONN_ID}") + assert response.status_code == 204 + connection = session.query(Connection).all() + assert len(connection) == 0 + + def test_delete_should_respond_404(self, test_client): + response = test_client.delete(f"/public/connections/{TEST_CONN_ID}") + assert response.status_code == 404 + body = response.json() + assert f"The Connection with connection_id: `{TEST_CONN_ID}` was not found" == body["detail"] From e08998a916ea550df7f38030d2c01598bd8bbdce Mon Sep 17 00:00:00 2001 From: Brent Bovenzi Date: Wed, 2 Oct 2024 11:08:53 +0200 Subject: [PATCH 106/802] Add is_paused toggle (#42621) * Add pause/unpause DAG toggle * wire up onSuccess handler * Refactor query names --- airflow/ui/package.json | 1 + airflow/ui/pnpm-lock.yaml | 9 ++ airflow/ui/rules/react.js | 3 +- .../ui/src/components/DataTable/DataTable.tsx | 2 +- airflow/ui/src/components/TogglePause.tsx | 56 ++++++++++++ airflow/ui/src/pages/DagsList/DagsFilters.tsx | 86 +++++++++++++++++++ .../ui/src/pages/{ => DagsList}/DagsList.tsx | 63 +++++--------- airflow/ui/src/pages/DagsList/index.tsx | 20 +++++ 8 files changed, 196 insertions(+), 44 deletions(-) create mode 100644 airflow/ui/src/components/TogglePause.tsx create mode 100644 airflow/ui/src/pages/DagsList/DagsFilters.tsx rename airflow/ui/src/pages/{ => DagsList}/DagsList.tsx (72%) create mode 100644 airflow/ui/src/pages/DagsList/index.tsx diff --git a/airflow/ui/package.json b/airflow/ui/package.json index 1f77334074f03..82c6370f9dcba 100644 --- a/airflow/ui/package.json +++ b/airflow/ui/package.json @@ -32,6 +32,7 @@ }, "devDependencies": { "@7nohe/openapi-react-query-codegen": "^1.6.0", + "@eslint/compat": "^1.1.1", "@eslint/js": "^9.10.0", "@stylistic/eslint-plugin": "^2.8.0", "@tanstack/eslint-plugin-query": "^5.52.0", diff --git a/airflow/ui/pnpm-lock.yaml b/airflow/ui/pnpm-lock.yaml index 0f9f256941f5e..515e7fea5279d 100644 --- a/airflow/ui/pnpm-lock.yaml +++ b/airflow/ui/pnpm-lock.yaml @@ -51,6 +51,9 @@ importers: '@7nohe/openapi-react-query-codegen': specifier: ^1.6.0 version: 1.6.0(commander@12.1.0)(glob@11.0.0)(magicast@0.3.5)(ts-morph@23.0.0)(typescript@5.5.4) + '@eslint/compat': + specifier: ^1.1.1 + version: 1.1.1 '@eslint/js': specifier: ^9.10.0 version: 9.10.0 @@ -920,6 +923,10 @@ packages: resolution: {integrity: sha512-G/M/tIiMrTAxEWRfLfQJMmGNX28IxBg4PBz8XqQhqUHLFI6TL2htpIB1iQCj144V5ee/JaKyT9/WZ0MGZWfA7A==} engines: {node: ^12.0.0 || ^14.0.0 || >=16.0.0} + '@eslint/compat@1.1.1': + resolution: {integrity: sha512-lpHyRyplhGPL5mGEh6M9O5nnKk0Gz4bFI+Zu6tKlPpDUN7XshWvH9C/px4UVm87IAANE0W81CEsNGbS1KlzXpA==} + engines: {node: ^18.18.0 || ^20.9.0 || >=21.1.0} + '@eslint/config-array@0.18.0': resolution: {integrity: sha512-fTxvnS1sRMu3+JjXwJG0j/i4RT9u4qJ+lqS/yCGap4lH4zZGzQ7tu+xZqQmcMZq5OBZDL4QRxQzRjkWcGt8IVw==} engines: {node: ^18.18.0 || ^20.9.0 || >=21.1.0} @@ -4368,6 +4375,8 @@ snapshots: '@eslint-community/regexpp@4.11.0': {} + '@eslint/compat@1.1.1': {} + '@eslint/config-array@0.18.0': dependencies: '@eslint/object-schema': 2.1.4 diff --git a/airflow/ui/rules/react.js b/airflow/ui/rules/react.js index 8b4b610078c53..4c8d8b8ba5f09 100644 --- a/airflow/ui/rules/react.js +++ b/airflow/ui/rules/react.js @@ -20,6 +20,7 @@ /** * @import { FlatConfig } from "@typescript-eslint/utils/ts-eslint"; */ +import { fixupPluginRules } from "@eslint/compat"; import jsxA11y from "eslint-plugin-jsx-a11y"; import react from "eslint-plugin-react"; import reactHooks from "eslint-plugin-react-hooks"; @@ -57,7 +58,7 @@ export const reactRefreshNamespace = "react-refresh"; export const reactRules = /** @type {const} @satisfies {FlatConfig.Config} */ ({ plugins: { [jsxA11yNamespace]: jsxA11y, - [reactHooksNamespace]: reactHooks, + [reactHooksNamespace]: fixupPluginRules(reactHooks), [reactNamespace]: react, [reactRefreshNamespace]: reactRefresh, }, diff --git a/airflow/ui/src/components/DataTable/DataTable.tsx b/airflow/ui/src/components/DataTable/DataTable.tsx index a4bf1255a4ba6..705d7883f07d2 100644 --- a/airflow/ui/src/components/DataTable/DataTable.tsx +++ b/airflow/ui/src/components/DataTable/DataTable.tsx @@ -115,7 +115,7 @@ export const DataTable = ({ return ( -
+ + + + + + + + + + + {% for host in hosts %} + + + + + + + + + + + {% endfor %} +
HostnameStateQueuesFirst OnlineLast Heart BeatActive JobsSystem Information
{{ host.worker_name }} + {%- if host.state == "offline" -%} + {{ host.state }} + {%- elif host.last_update.timestamp() <= five_min_ago.timestamp() -%} + Reported {{ host.state }} + but no heartbeat + {%- elif host.state == "starting" -%} + {{ host.state }} + {%- elif host.state == "running" -%} + {{ host.state }} + {%- elif host.state == "idle" -%} + {{ host.state }} + {%- elif host.state == "terminating" -%} + {{ host.state }} + {%- elif host.state == "unknown" -%} + {{ host.state }} + {%- else -%} + {{ host.state }} + {%- endif -%} + {% if host.queues %}{{ host.queues }}{% else %}(all){% endif %}{% if host.last_update %}{% endif %}{{ host.jobs_active }} +
    + {% for item in host.sysinfo_json %} +
  • {{ item }}: {{ host.sysinfo_json[item] }}
  • + {% endfor %} +
+
+ {% endif %} + {% endblock %} + diff --git a/airflow/providers/edge/plugins/templates/edge_worker_jobs.html b/airflow/providers/edge/plugins/templates/edge_worker_jobs.html new file mode 100644 index 0000000000000..a73e0f1d485f4 --- /dev/null +++ b/airflow/providers/edge/plugins/templates/edge_worker_jobs.html @@ -0,0 +1,63 @@ +{# + 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. + #} + + {% extends base_template %} + + {% block title %} + Edge Worker Jobs + {% endblock %} + + {% block content %} +

Edge Worker Jobs

+ {% if jobs|length == 0 %} +

No jobs running currently

+ {% else %} + + + + + + + + + + + + + + + + {% for job in jobs %} + + + + + + + + + + + + + {% endfor %} +
DAG IDTask IDRun IDMap IndexTry NumberStateQueueQueued DTTMEdge WorkerLast Update
{{ job.dag_id }}{{ job.task_id }}{{ job.run_id }}{% if job.map_index >= 0 %}{{ job.map_index }}{% else %}-{% endif %}{{ job.try_number }}{{ html_states[job.state] }}{{ job.queue }}{% if job.edge_worker %}{{ job.edge_worker }}{% endif %}{% if job.last_update %}{% endif %}
+ {% endif %} + {% endblock %} + diff --git a/airflow/providers/edge/provider.yaml b/airflow/providers/edge/provider.yaml index cb775ee7cc7e4..6525b7bb846ff 100644 --- a/airflow/providers/edge/provider.yaml +++ b/airflow/providers/edge/provider.yaml @@ -32,6 +32,10 @@ dependencies: - apache-airflow>=2.10.0 - pydantic>=2.3.0 +plugins: + - name: edge_executor + plugin-class: airflow.providers.edge.plugins.edge_executor_plugin.EdgeExecutorPlugin + config: edge: description: | diff --git a/generated/provider_dependencies.json b/generated/provider_dependencies.json index 10631afb9b292..59da56f180744 100644 --- a/generated/provider_dependencies.json +++ b/generated/provider_dependencies.json @@ -528,7 +528,12 @@ "pydantic>=2.3.0" ], "devel-deps": [], - "plugins": [], + "plugins": [ + { + "name": "edge_executor", + "plugin-class": "airflow.providers.edge.plugins.edge_executor_plugin.EdgeExecutorPlugin" + } + ], "cross-providers-deps": [], "excluded-python-versions": [], "state": "not-ready" diff --git a/tests/plugins/test_plugins_manager.py b/tests/plugins/test_plugins_manager.py index cb59afd36742a..7e4bedbfb8c1b 100644 --- a/tests/plugins/test_plugins_manager.py +++ b/tests/plugins/test_plugins_manager.py @@ -417,7 +417,7 @@ def test_does_not_double_import_entrypoint_provider_plugins(self): assert len(plugins_manager.plugins) == 0 plugins_manager.load_entrypoint_plugins() plugins_manager.load_providers_plugins() - assert len(plugins_manager.plugins) == 3 + assert len(plugins_manager.plugins) == 4 class TestPluginsDirectorySource: diff --git a/tests/providers/edge/api_endpoints/__init__.py b/tests/providers/edge/api_endpoints/__init__.py new file mode 100644 index 0000000000000..217e5db960782 --- /dev/null +++ b/tests/providers/edge/api_endpoints/__init__.py @@ -0,0 +1,17 @@ +# +# 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. diff --git a/tests/providers/edge/api_endpoints/test_health_endpoint.py b/tests/providers/edge/api_endpoints/test_health_endpoint.py new file mode 100644 index 0000000000000..1bfc9e5c0c5bf --- /dev/null +++ b/tests/providers/edge/api_endpoints/test_health_endpoint.py @@ -0,0 +1,23 @@ +# 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 + +from airflow.providers.edge.api_endpoints.health_endpoint import health + + +def test_health(): + assert health() == {} diff --git a/tests/providers/edge/api_endpoints/test_rpc_api_endpoint.py b/tests/providers/edge/api_endpoints/test_rpc_api_endpoint.py new file mode 100644 index 0000000000000..becf2f9397e31 --- /dev/null +++ b/tests/providers/edge/api_endpoints/test_rpc_api_endpoint.py @@ -0,0 +1,281 @@ +# 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 json +from typing import TYPE_CHECKING, Generator +from unittest import mock + +import pytest + +from airflow.api_connexion.exceptions import PermissionDenied +from airflow.configuration import conf +from airflow.models.baseoperator import BaseOperator +from airflow.models.connection import Connection +from airflow.models.dagrun import DagRun +from airflow.models.taskinstance import TaskInstance +from airflow.models.xcom import XCom +from airflow.operators.empty import EmptyOperator +from airflow.providers.edge.api_endpoints.rpc_api_endpoint import _initialize_method_map +from airflow.providers.edge.models.edge_job import EdgeJob +from airflow.providers.edge.models.edge_logs import EdgeLogs +from airflow.providers.edge.models.edge_worker import EdgeWorker +from airflow.serialization.pydantic.taskinstance import TaskInstancePydantic +from airflow.serialization.serialized_objects import BaseSerialization +from airflow.settings import _ENABLE_AIP_44 +from airflow.utils.jwt_signer import JWTSigner +from airflow.utils.state import State +from airflow.www import app +from tests.test_utils.decorators import dont_initialize_flask_app_submodules +from tests.test_utils.mock_plugins import mock_plugin_manager + +# Note: Sounds a bit strange to disable internal API tests in isolation mode but... +# As long as the test is modelled to run its own internal API endpoints, it is conflicting +# to the test setup with a dedicated internal API server. +pytestmark = [pytest.mark.db_test, pytest.mark.skip_if_database_isolation_mode] + + +def test_initialize_method_map(): + method_map = _initialize_method_map() + assert len(method_map) > 70 + for method in [ + # Test some basics + XCom.get_value, + XCom.get_one, + XCom.clear, + XCom.set, + DagRun.get_previous_dagrun, + DagRun.get_previous_scheduled_dagrun, + DagRun.get_task_instances, + DagRun.fetch_task_instance, + # Test some for Edge + EdgeJob.reserve_task, + EdgeJob.set_state, + EdgeLogs.push_logs, + EdgeWorker.register_worker, + EdgeWorker.set_state, + ]: + method_key = f"{method.__module__}.{method.__qualname__}" + assert method_key in method_map.keys() + + +if TYPE_CHECKING: + from flask import Flask + +TEST_METHOD_NAME = "test_method" +TEST_METHOD_WITH_LOG_NAME = "test_method_with_log" +TEST_API_ENDPOINT = "/edge_worker/v1/rpcapi" + +mock_test_method = mock.MagicMock() + +pytest.importorskip("pydantic", minversion="2.0.0") + + +def equals(a, b) -> bool: + return a == b + + +@pytest.mark.skipif(not _ENABLE_AIP_44, reason="AIP-44 is disabled") +class TestRpcApiEndpoint: + @pytest.fixture(scope="session") + def minimal_app_for_edge_api(self) -> Flask: + @dont_initialize_flask_app_submodules( + skip_all_except=[ + "init_api_auth", # This is needed for Airflow 2.10 compat tests + "init_appbuilder", + "init_plugins", + ] + ) + def factory() -> Flask: + import airflow.providers.edge.plugins.edge_executor_plugin as plugin_module + + class TestingEdgeExecutorPlugin(plugin_module.EdgeExecutorPlugin): + flask_blueprints = [plugin_module._get_api_endpoints(), plugin_module.template_bp] + + testing_edge_plugin = TestingEdgeExecutorPlugin() + assert len(testing_edge_plugin.flask_blueprints) > 0 + with mock_plugin_manager(plugins=[testing_edge_plugin]): + return app.create_app(testing=True, config={"WTF_CSRF_ENABLED": False}) # type:ignore + + return factory() + + @pytest.fixture + def setup_attrs(self, minimal_app_for_edge_api: Flask) -> Generator: + self.app = minimal_app_for_edge_api + self.client = self.app.test_client() # type:ignore + mock_test_method.reset_mock() + mock_test_method.side_effect = None + with mock.patch( + "airflow.providers.edge.api_endpoints.rpc_api_endpoint._initialize_method_map" + ) as mock_initialize_method_map: + mock_initialize_method_map.return_value = { + TEST_METHOD_NAME: mock_test_method, + } + yield mock_initialize_method_map + + @pytest.fixture + def signer(self) -> JWTSigner: + return JWTSigner( + secret_key=conf.get("core", "internal_api_secret_key"), + expiration_time_in_seconds=conf.getint("core", "internal_api_clock_grace", fallback=30), + audience="api", + ) + + @pytest.mark.parametrize( + "input_params, method_result, result_cmp_func, method_params", + [ + ({}, None, lambda got, _: got == b"", {}), + ({}, "test_me", equals, {}), + ( + BaseSerialization.serialize({"dag_id": 15, "task_id": "fake-task"}), + ("dag_id_15", "fake-task", 1), + equals, + {"dag_id": 15, "task_id": "fake-task"}, + ), + ( + {}, + TaskInstance(task=EmptyOperator(task_id="task"), run_id="run_id", state=State.RUNNING), + lambda a, b: a.model_dump() == TaskInstancePydantic.model_validate(b).model_dump() + and isinstance(a.task, BaseOperator), + {}, + ), + ( + {}, + Connection(conn_id="test_conn", conn_type="http", host="", password=""), + lambda a, b: a.get_uri() == b.get_uri() and a.conn_id == b.conn_id, + {}, + ), + ], + ) + def test_method( + self, input_params, method_result, result_cmp_func, method_params, setup_attrs, signer: JWTSigner + ): + mock_test_method.return_value = method_result + headers = { + "Content-Type": "application/json", + "Accept": "application/json", + "Authorization": signer.generate_signed_token({"method": TEST_METHOD_NAME}), + } + input_data = { + "jsonrpc": "2.0", + "method": TEST_METHOD_NAME, + "params": input_params, + } + response = self.client.post( + TEST_API_ENDPOINT, + headers=headers, + data=json.dumps(input_data), + ) + assert response.status_code == 200 + if method_result: + response_data = BaseSerialization.deserialize(json.loads(response.data), use_pydantic_models=True) + else: + response_data = response.data + + assert result_cmp_func(response_data, method_result) + + mock_test_method.assert_called_once_with(**method_params, session=mock.ANY) + + def test_method_with_exception(self, setup_attrs, signer: JWTSigner): + headers = { + "Content-Type": "application/json", + "Accept": "application/json", + "Authorization": signer.generate_signed_token({"method": TEST_METHOD_NAME}), + } + mock_test_method.side_effect = ValueError("Error!!!") + data = {"jsonrpc": "2.0", "method": TEST_METHOD_NAME, "params": {}} + + response = self.client.post(TEST_API_ENDPOINT, headers=headers, data=json.dumps(data)) + assert response.status_code == 500 + assert response.data, b"Error executing method: test_method." + mock_test_method.assert_called_once() + + def test_unknown_method(self, setup_attrs, signer: JWTSigner): + UNKNOWN_METHOD = "i-bet-it-does-not-exist" + headers = { + "Content-Type": "application/json", + "Accept": "application/json", + "Authorization": signer.generate_signed_token({"method": UNKNOWN_METHOD}), + } + data = {"jsonrpc": "2.0", "method": UNKNOWN_METHOD, "params": {}} + + response = self.client.post(TEST_API_ENDPOINT, headers=headers, data=json.dumps(data)) + assert response.status_code == 400 + assert response.data.startswith(b"Unrecognized method: i-bet-it-does-not-exist.") + mock_test_method.assert_not_called() + + def test_invalid_jsonrpc(self, setup_attrs, signer: JWTSigner): + headers = { + "Content-Type": "application/json", + "Accept": "application/json", + "Authorization": signer.generate_signed_token({"method": TEST_METHOD_NAME}), + } + data = {"jsonrpc": "1.0", "method": TEST_METHOD_NAME, "params": {}} + + response = self.client.post(TEST_API_ENDPOINT, headers=headers, data=json.dumps(data)) + assert response.status_code == 400 + assert response.data.startswith(b"Expected jsonrpc 2.0 request.") + mock_test_method.assert_not_called() + + def test_missing_token(self, setup_attrs): + mock_test_method.return_value = None + + input_data = { + "jsonrpc": "2.0", + "method": TEST_METHOD_NAME, + "params": {}, + } + with pytest.raises(PermissionDenied, match="Unable to authenticate API via token."): + self.client.post( + TEST_API_ENDPOINT, + headers={"Content-Type": "application/json", "Accept": "application/json"}, + data=json.dumps(input_data), + ) + + def test_invalid_token(self, setup_attrs, signer: JWTSigner): + headers = { + "Content-Type": "application/json", + "Accept": "application/json", + "Authorization": signer.generate_signed_token({"method": "WRONG_METHOD_NAME"}), + } + data = {"jsonrpc": "1.0", "method": TEST_METHOD_NAME, "params": {}} + + with pytest.raises( + PermissionDenied, match="Bad Signature. Please use only the tokens provided by the API." + ): + self.client.post(TEST_API_ENDPOINT, headers=headers, data=json.dumps(data)) + + def test_missing_accept(self, setup_attrs, signer: JWTSigner): + headers = { + "Content-Type": "application/json", + "Authorization": signer.generate_signed_token({"method": "WRONG_METHOD_NAME"}), + } + data = {"jsonrpc": "1.0", "method": TEST_METHOD_NAME, "params": {}} + + with pytest.raises(PermissionDenied, match="Expected Accept: application/json"): + self.client.post(TEST_API_ENDPOINT, headers=headers, data=json.dumps(data)) + + def test_wrong_accept(self, setup_attrs, signer: JWTSigner): + headers = { + "Content-Type": "application/json", + "Accept": "application/html", + "Authorization": signer.generate_signed_token({"method": "WRONG_METHOD_NAME"}), + } + data = {"jsonrpc": "1.0", "method": TEST_METHOD_NAME, "params": {}} + + with pytest.raises(PermissionDenied, match="Expected Accept: application/json"): + self.client.post(TEST_API_ENDPOINT, headers=headers, data=json.dumps(data)) diff --git a/tests/providers/edge/plugins/__init__.py b/tests/providers/edge/plugins/__init__.py new file mode 100644 index 0000000000000..217e5db960782 --- /dev/null +++ b/tests/providers/edge/plugins/__init__.py @@ -0,0 +1,17 @@ +# +# 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. diff --git a/tests/providers/edge/plugins/test_edge_executor_plugin.py b/tests/providers/edge/plugins/test_edge_executor_plugin.py new file mode 100644 index 0000000000000..e3422b17da3c8 --- /dev/null +++ b/tests/providers/edge/plugins/test_edge_executor_plugin.py @@ -0,0 +1,66 @@ +# 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 importlib + +import pytest + +from airflow.plugins_manager import AirflowPlugin +from airflow.providers.edge.plugins import edge_executor_plugin +from tests.test_utils.config import conf_vars + + +def test_plugin_inactive(): + with conf_vars({("edge", "api_enabled"): "false"}): + importlib.reload(edge_executor_plugin) + + from airflow.providers.edge.plugins.edge_executor_plugin import ( + EDGE_EXECUTOR_ACTIVE, + EdgeExecutorPlugin, + ) + + rep = EdgeExecutorPlugin() + assert not EDGE_EXECUTOR_ACTIVE + assert len(rep.flask_blueprints) == 0 + assert len(rep.appbuilder_views) == 0 + + +def test_plugin_active(): + with conf_vars({("edge", "api_enabled"): "true"}): + importlib.reload(edge_executor_plugin) + + from airflow.providers.edge.plugins.edge_executor_plugin import ( + EDGE_EXECUTOR_ACTIVE, + EdgeExecutorPlugin, + ) + + rep = EdgeExecutorPlugin() + assert EDGE_EXECUTOR_ACTIVE + assert len(rep.flask_blueprints) == 2 + assert len(rep.appbuilder_views) == 2 + + +@pytest.fixture +def plugin(): + from airflow.providers.edge.plugins.edge_executor_plugin import EdgeExecutorPlugin + + return EdgeExecutorPlugin() + + +def test_plugin_is_airflow_plugin(plugin): + assert isinstance(plugin, AirflowPlugin) From 7f98de89e3e9bcfbc01b52f6962c755b54dd3ab4 Mon Sep 17 00:00:00 2001 From: Jarek Potiuk Date: Thu, 3 Oct 2024 03:00:54 -0700 Subject: [PATCH 132/802] Update min version of Pydantic to 2.6.4 (#42694) Pydantic 2.6.4 fixes problem with AliasGenerator to throw error when generating schema - see an issue in Pydantic repository https://github.com/pydantic/pydantic/issues/8768 --- airflow/providers/edge/provider.yaml | 2 +- generated/provider_dependencies.json | 2 +- hatch_build.py | 2 +- newsfragments/41857.significant.rst | 2 +- 4 files changed, 4 insertions(+), 4 deletions(-) diff --git a/airflow/providers/edge/provider.yaml b/airflow/providers/edge/provider.yaml index 6525b7bb846ff..d6644271a02f3 100644 --- a/airflow/providers/edge/provider.yaml +++ b/airflow/providers/edge/provider.yaml @@ -30,7 +30,7 @@ versions: dependencies: - apache-airflow>=2.10.0 - - pydantic>=2.3.0 + - pydantic>=2.6.4 plugins: - name: edge_executor diff --git a/generated/provider_dependencies.json b/generated/provider_dependencies.json index 59da56f180744..7dc5e337292b8 100644 --- a/generated/provider_dependencies.json +++ b/generated/provider_dependencies.json @@ -525,7 +525,7 @@ "edge": { "deps": [ "apache-airflow>=2.10.0", - "pydantic>=2.3.0" + "pydantic>=2.6.4" ], "devel-deps": [], "plugins": [ diff --git a/hatch_build.py b/hatch_build.py index 6e3d77981e3d4..765e71ff98962 100644 --- a/hatch_build.py +++ b/hatch_build.py @@ -467,7 +467,7 @@ 'pendulum>=3.0.0,<4.0;python_version>="3.12"', "pluggy>=1.5.0", "psutil>=5.8.0", - "pydantic>=2.6.0", + "pydantic>=2.6.4", "pygments>=2.0.1", "pyjwt>=2.0.0", "python-daemon>=3.0.0", diff --git a/newsfragments/41857.significant.rst b/newsfragments/41857.significant.rst index df3c85853eee7..f0b06f2811b1f 100644 --- a/newsfragments/41857.significant.rst +++ b/newsfragments/41857.significant.rst @@ -1,3 +1,3 @@ **Breaking Change** -Airflow core now depends on ``pydantic>=2.3.0``. If you have Pydantic v1 installed, please upgrade. +Airflow core now depends on Pydantic v2. If you have Pydantic v1 installed, please upgrade. From e65916fa2d1adf81477e5bc93ab42696c10f2e06 Mon Sep 17 00:00:00 2001 From: Lorin Dawson <22798188+R7L208@users.noreply.github.com> Date: Thu, 3 Oct 2024 06:36:30 -0600 Subject: [PATCH 133/802] Add `on_kill` to Databricks Workflow Operator (#42115) * add on_kill override to databricks workflow operator * on_kill equivalent for DatabricksSqlOperator * add tests for create_timeout_thread * add note for on_kill in DatabricksCopyIntoOperator * chore: static checks * remove changes for databricks_sql.py for PR isolated to databricks_workflows.py --------- Co-authored-by: Lorin --- .../operators/databricks_workflow.py | 27 ++++++++++++++++++- .../operators/test_databricks_workflow.py | 22 +++++++++++++++ 2 files changed, 48 insertions(+), 1 deletion(-) diff --git a/airflow/providers/databricks/operators/databricks_workflow.py b/airflow/providers/databricks/operators/databricks_workflow.py index 15333dc69118b..6df8e2d025cea 100644 --- a/airflow/providers/databricks/operators/databricks_workflow.py +++ b/airflow/providers/databricks/operators/databricks_workflow.py @@ -52,7 +52,7 @@ class WorkflowRunMetadata: """ conn_id: str - job_id: str + job_id: int run_id: int @@ -116,6 +116,7 @@ def __init__( self.notebook_params = notebook_params or {} self.tasks_to_convert = tasks_to_convert or [] self.relevant_upstreams = [task_id] + self.workflow_run_metadata: WorkflowRunMetadata | None = None super().__init__(task_id=task_id, **kwargs) def _get_hook(self, caller: str) -> DatabricksHook: @@ -212,12 +213,36 @@ def execute(self, context: Context) -> Any: self._wait_for_job_to_start(run_id) + self.workflow_run_metadata = WorkflowRunMetadata( + self.databricks_conn_id, + job_id, + run_id, + ) + return { "conn_id": self.databricks_conn_id, "job_id": job_id, "run_id": run_id, } + def on_kill(self) -> None: + if self.workflow_run_metadata: + run_id = self.workflow_run_metadata.run_id + job_id = self.workflow_run_metadata.job_id + + self._hook.cancel_run(run_id) + self.log.info( + "Run: %(run_id)s of job_id: %(job_id)s was requested to be cancelled.", + {"run_id": run_id, "job_id": job_id}, + ) + else: + self.log.error( + """ + Error: Workflow Run metadata is not populated, so the run was not canceled. This could be due + to the workflow not being started or an error in the workflow creation process. + """ + ) + class DatabricksWorkflowTaskGroup(TaskGroup): """ diff --git a/tests/providers/databricks/operators/test_databricks_workflow.py b/tests/providers/databricks/operators/test_databricks_workflow.py index 4c3f54b800ae9..fbc429ed1d9a8 100644 --- a/tests/providers/databricks/operators/test_databricks_workflow.py +++ b/tests/providers/databricks/operators/test_databricks_workflow.py @@ -28,6 +28,7 @@ from airflow.providers.databricks.hooks.databricks import RunLifeCycleState from airflow.providers.databricks.operators.databricks_workflow import ( DatabricksWorkflowTaskGroup, + WorkflowRunMetadata, _CreateDatabricksWorkflowOperator, _flatten_node, ) @@ -59,6 +60,11 @@ def mock_task_group(): return mock_group +@pytest.fixture +def mock_workflow_run_metadata(): + return MagicMock(spec=WorkflowRunMetadata) + + def test_flatten_node(): """Test that _flatten_node returns a flat list of operators.""" task_group = MagicMock(spec=DatabricksWorkflowTaskGroup) @@ -231,3 +237,19 @@ def test_task_group_root_tasks_set_upstream_to_operator(mock_databricks_workflow create_operator_instance = mock_databricks_workflow_operator.return_value task1.set_upstream.assert_called_once_with(create_operator_instance) + + +def test_on_kill(mock_databricks_hook, context, mock_workflow_run_metadata): + """Test that _CreateDatabricksWorkflowOperator.execute runs the task group.""" + operator = _CreateDatabricksWorkflowOperator(task_id="test_task", databricks_conn_id="databricks_default") + operator.workflow_run_metadata = mock_workflow_run_metadata + + RUN_ID = 789 + + mock_workflow_run_metadata.conn_id = operator.databricks_conn_id + mock_workflow_run_metadata.job_id = "123" + mock_workflow_run_metadata.run_id = RUN_ID + + operator.on_kill() + + operator._hook.cancel_run.assert_called_once_with(RUN_ID) From 49e2753cb52015124acc689616dc76d3546ba8e7 Mon Sep 17 00:00:00 2001 From: Pierre Jeambrun Date: Fri, 4 Oct 2024 00:50:30 +0800 Subject: [PATCH 134/802] AIP-84 Serve new UI from FastAPI API (#42663) * Serve new UI from FastAPI API * Fix CI * Fix CI another try --- airflow/api_fastapi/app.py | 32 ++++++++++++++- airflow/ui/.env.example | 3 +- airflow/ui/src/App.tsx | 1 - airflow/ui/src/layouts/Nav/DocsButton.tsx | 2 +- airflow/ui/src/layouts/Nav/Nav.tsx | 2 +- airflow/ui/src/main.tsx | 4 +- airflow/ui/src/vite-env.d.ts | 2 +- airflow/ui/vite.config.ts | 4 +- airflow/www/app.py | 2 - airflow/www/extensions/init_react_ui.py | 40 ------------------- airflow/www/templates/airflow/main.html | 2 +- .../14_node_environment_setup.rst | 10 ----- .../run_update_fastapi_api_spec.py | 2 + 13 files changed, 41 insertions(+), 65 deletions(-) delete mode 100644 airflow/www/extensions/init_react_ui.py diff --git a/airflow/api_fastapi/app.py b/airflow/api_fastapi/app.py index 6f8bbcdf149b4..6b9df0ed8b7f8 100644 --- a/airflow/api_fastapi/app.py +++ b/airflow/api_fastapi/app.py @@ -16,9 +16,16 @@ # under the License. from __future__ import annotations -from fastapi import FastAPI +import os +from pathlib import Path + +from fastapi import FastAPI, Request from fastapi.middleware.cors import CORSMiddleware +from fastapi.responses import HTMLResponse +from fastapi.staticfiles import StaticFiles +from fastapi.templating import Jinja2Templates +from airflow.settings import AIRFLOW_PATH from airflow.www.extensions.init_dagbag import get_dag_bag app: FastAPI | None = None @@ -70,6 +77,29 @@ def init_views(app) -> None: app.include_router(ui_router) app.include_router(public_router) + dev_mode = os.environ.get("DEV_MODE", False) == "true" + + directory = Path(AIRFLOW_PATH) / ("airflow/ui/dev" if dev_mode else "airflow/ui/dist") + + # During python tests or when the backend is run without having the frontend build + # those directories might not exist. App should not fail initializing in those scenarios. + Path(directory).mkdir(exist_ok=True) + + templates = Jinja2Templates(directory=directory) + + app.mount( + "/static", + StaticFiles( + directory=directory, + html=True, + ), + name="webapp_static_folder", + ) + + @app.get("/webapp/{rest_of_path:path}", response_class=HTMLResponse, include_in_schema=False) + def webapp(request: Request, rest_of_path: str): + return templates.TemplateResponse("/index.html", {"request": request}, media_type="text/html") + def cached_app(config=None, testing=False) -> FastAPI: """Return cached instance of Airflow UI app.""" diff --git a/airflow/ui/.env.example b/airflow/ui/.env.example index 9374d93de6bca..3e3c1569f1238 100644 --- a/airflow/ui/.env.example +++ b/airflow/ui/.env.example @@ -19,5 +19,4 @@ # This is an example. You should make your own `.env.local` file for development - -VITE_FASTAPI_URL="http://localhost:29091" +VITE_LEGACY_API_URL="http://localhost:28080" diff --git a/airflow/ui/src/App.tsx b/airflow/ui/src/App.tsx index 0eb603b46a330..3c5e9d866f0c9 100644 --- a/airflow/ui/src/App.tsx +++ b/airflow/ui/src/App.tsx @@ -22,7 +22,6 @@ import { DagsList } from "src/pages/DagsList"; import { BaseLayout } from "./layouts/BaseLayout"; -// Note: When changing routes, make sure to update init_react_ui.py too export const App = () => ( } path="/"> diff --git a/airflow/ui/src/layouts/Nav/DocsButton.tsx b/airflow/ui/src/layouts/Nav/DocsButton.tsx index 07a4b93dfaede..e85d923b88f54 100644 --- a/airflow/ui/src/layouts/Nav/DocsButton.tsx +++ b/airflow/ui/src/layouts/Nav/DocsButton.tsx @@ -38,7 +38,7 @@ const links = [ title: "GitHub Repo", }, { - href: `${import.meta.env.VITE_FASTAPI_URL}/docs`, + href: `/docs`, title: "REST API Reference", }, ]; diff --git a/airflow/ui/src/layouts/Nav/Nav.tsx b/airflow/ui/src/layouts/Nav/Nav.tsx index 55bfd4480e0f4..9886b5eb75760 100644 --- a/airflow/ui/src/layouts/Nav/Nav.tsx +++ b/airflow/ui/src/layouts/Nav/Nav.tsx @@ -101,7 +101,7 @@ export const Nav = () => { } title="Return to legacy UI" /> diff --git a/airflow/ui/src/main.tsx b/airflow/ui/src/main.tsx index 7b762508ea7b3..daf4bcd024cd6 100644 --- a/airflow/ui/src/main.tsx +++ b/airflow/ui/src/main.tsx @@ -43,8 +43,6 @@ const queryClient = new QueryClient({ }, }); -axios.defaults.baseURL = import.meta.env.VITE_FASTAPI_URL; - // redirect to login page if the API responds with unauthorized or forbidden errors axios.interceptors.response.use( (response: AxiosResponse) => response, @@ -61,7 +59,7 @@ axios.interceptors.response.use( const root = createRoot(document.querySelector("#root") as HTMLDivElement); root.render( - + diff --git a/airflow/ui/src/vite-env.d.ts b/airflow/ui/src/vite-env.d.ts index 193866687bff9..8a62dd17206eb 100644 --- a/airflow/ui/src/vite-env.d.ts +++ b/airflow/ui/src/vite-env.d.ts @@ -21,7 +21,7 @@ /// interface ImportMetaEnv { - readonly VITE_FASTAPI_URL: string; + readonly VITE_LEGACY_API_URL: string; } interface ImportMeta { diff --git a/airflow/ui/vite.config.ts b/airflow/ui/vite.config.ts index 06ad450f377a1..7bc48d640418a 100644 --- a/airflow/ui/vite.config.ts +++ b/airflow/ui/vite.config.ts @@ -29,8 +29,8 @@ export default defineConfig({ name: "transform-url-src", transformIndexHtml: (html) => html - .replace(`src="/assets/`, `src="/ui/assets/`) - .replace(`href="/`, `href="/ui/`), + .replace(`src="/assets/`, `src="/static/assets/`) + .replace(`href="/`, `href="/webapp/`), }, ], resolve: { alias: { openapi: "/openapi-gen", src: "/src" } }, diff --git a/airflow/www/app.py b/airflow/www/app.py index f5e1191fb43fb..3409510b5a1a6 100644 --- a/airflow/www/app.py +++ b/airflow/www/app.py @@ -40,7 +40,6 @@ from airflow.www.extensions.init_dagbag import init_dagbag from airflow.www.extensions.init_jinja_globals import init_jinja_globals from airflow.www.extensions.init_manifest_files import configure_manifest_files -from airflow.www.extensions.init_react_ui import init_react_ui from airflow.www.extensions.init_robots import init_robots from airflow.www.extensions.init_security import ( init_api_auth, @@ -155,7 +154,6 @@ def create_app(config=None, testing=False): with flask_app.app_context(): init_appbuilder(flask_app) - init_react_ui(flask_app) init_appbuilder_views(flask_app) init_appbuilder_links(flask_app) init_plugins(flask_app) diff --git a/airflow/www/extensions/init_react_ui.py b/airflow/www/extensions/init_react_ui.py deleted file mode 100644 index 872a22c059476..0000000000000 --- a/airflow/www/extensions/init_react_ui.py +++ /dev/null @@ -1,40 +0,0 @@ -# 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 os - -from flask import Blueprint - - -def init_react_ui(app): - dev_mode = os.environ.get("DEV_MODE", False) == "true" - - bp = Blueprint( - "ui", - __name__, - # The dev mode index file points to the vite dev server instead of static build files - static_folder="../../ui/dev" if dev_mode else "../../ui/dist", - static_url_path="/ui", - ) - - @bp.route("/ui", defaults={"page": ""}) - @bp.route("/ui/") - def index(page): - return bp.send_static_file("index.html") - - app.register_blueprint(bp) diff --git a/airflow/www/templates/airflow/main.html b/airflow/www/templates/airflow/main.html index 69aa6faaaf0de..008418d7e5b78 100644 --- a/airflow/www/templates/airflow/main.html +++ b/airflow/www/templates/airflow/main.html @@ -99,7 +99,7 @@ {% if auth_manager.is_logged_in() %} {% call show_message(category='info', dismissible=true) %} We have a new UI for Airflow 3.0 - Check it out now! + Check it out now! {% endcall %} {% endif %} {% endblock %} diff --git a/contributing-docs/14_node_environment_setup.rst b/contributing-docs/14_node_environment_setup.rst index 7b10f0b0d5ed5..81ced88240ac5 100644 --- a/contributing-docs/14_node_environment_setup.rst +++ b/contributing-docs/14_node_environment_setup.rst @@ -93,16 +93,6 @@ Copy the example environment cp .env.example .env.local -If you run into CORS issues, you may need to add some variables to your Breeze config, ``files/airflow-breeze-config/variables.env``: - -.. code-block:: bash - - export AIRFLOW__API__ACCESS_CONTROL_ALLOW_HEADERS="Origin, Access-Control-Request-Method" - export AIRFLOW__API__ACCESS_CONTROL_ALLOW_METHODS="*" - export AIRFLOW__API__ACCESS_CONTROL_ALLOW_ORIGINS="http://localhost:28080,http://localhost:8080" - - - DEPRECATED Airflow WWW ---------------------- diff --git a/scripts/in_container/run_update_fastapi_api_spec.py b/scripts/in_container/run_update_fastapi_api_spec.py index 4d78bc4afd585..5d31b0bee3f0c 100644 --- a/scripts/in_container/run_update_fastapi_api_spec.py +++ b/scripts/in_container/run_update_fastapi_api_spec.py @@ -29,6 +29,8 @@ # The persisted openapi spec will list all endpoints (public and ui), this # is used for code generation. for route in app.routes: + if getattr(route, "name") == "webapp": + continue route.__setattr__("include_in_schema", True) with open(OPENAPI_SPEC_FILE, "w+") as f: From 18efe3cbf7a5015b1467677772f74bf93a75b98d Mon Sep 17 00:00:00 2001 From: Josix Date: Fri, 4 Oct 2024 02:36:21 +0900 Subject: [PATCH 135/802] fix(shell_params): prevent generating `,celery` in extra when there is no other extra items (#42709) --- dev/breeze/src/airflow_breeze/params/shell_params.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/dev/breeze/src/airflow_breeze/params/shell_params.py b/dev/breeze/src/airflow_breeze/params/shell_params.py index af74be27c919b..36fa44bb8fed5 100644 --- a/dev/breeze/src/airflow_breeze/params/shell_params.py +++ b/dev/breeze/src/airflow_breeze/params/shell_params.py @@ -332,7 +332,9 @@ def compose_file(self) -> str: get_console().print( "[warning]Adding `celery` extras as it is implicitly needed by celery executor" ) - self.airflow_extras = ",".join(current_extras.split(",") + ["celery"]) + self.airflow_extras = ( + ",".join(current_extras.split(",") + ["celery"]) if current_extras else "celery" + ) compose_file_list.append(DOCKER_COMPOSE_DIR / "base.yml") self.add_docker_in_docker(compose_file_list) From 6d731eb0adfe8ebc68004701bd852d292f4454f9 Mon Sep 17 00:00:00 2001 From: Vincent <97131062+vincbeck@users.noreply.github.com> Date: Thu, 3 Oct 2024 13:37:42 -0400 Subject: [PATCH 136/802] Mention in simple auth manager doc how to read/update passwords directly form file (#42710) * Mention in simple auth manager doc how to read/update passwords directly form file * Fix static checks --- docs/apache-airflow/core-concepts/auth-manager/simple.rst | 2 ++ 1 file changed, 2 insertions(+) diff --git a/docs/apache-airflow/core-concepts/auth-manager/simple.rst b/docs/apache-airflow/core-concepts/auth-manager/simple.rst index bef2e5032f0d7..f418ca15f2981 100644 --- a/docs/apache-airflow/core-concepts/auth-manager/simple.rst +++ b/docs/apache-airflow/core-concepts/auth-manager/simple.rst @@ -51,6 +51,8 @@ Each user needs two pieces of information: The password is auto-generated for each user and printed out in the webserver logs. When generated, these passwords are also saved in your environment, therefore they will not change if you stop or restart your environment. +The passwords are saved in the file ``generated/simple_auth_manager_passwords.json.generated``, you can read and update them directly in the file as well if desired. + .. _roles-permissions: Manage roles and permissions From 87bd9e20cc1c5d9b68a0a8622ff699c3c22f10c4 Mon Sep 17 00:00:00 2001 From: Maksim Date: Thu, 3 Oct 2024 12:49:46 -0700 Subject: [PATCH 137/802] Update tensorflow image uris for VertexAI system tests (#42707) --- .../cloud/vertex_ai/example_vertex_ai_custom_container.py | 3 +-- .../google/cloud/vertex_ai/example_vertex_ai_custom_job.py | 4 ++-- .../vertex_ai/example_vertex_ai_custom_job_python_package.py | 4 ++-- .../google/cloud/vertex_ai/example_vertex_ai_model_service.py | 4 ++-- 4 files changed, 7 insertions(+), 8 deletions(-) diff --git a/tests/system/providers/google/cloud/vertex_ai/example_vertex_ai_custom_container.py b/tests/system/providers/google/cloud/vertex_ai/example_vertex_ai_custom_container.py index b8d01f8d71493..dc09a8be90ed7 100644 --- a/tests/system/providers/google/cloud/vertex_ai/example_vertex_ai_custom_container.py +++ b/tests/system/providers/google/cloud/vertex_ai/example_vertex_ai_custom_container.py @@ -68,9 +68,8 @@ def TABULAR_DATASET(bucket_name): } -CONTAINER_URI = "gcr.io/cloud-aiplatform/training/tf-cpu.2-2:latest" CUSTOM_CONTAINER_URI = "us-central1-docker.pkg.dev/airflow-system-tests-resources/system-tests/housing" -MODEL_SERVING_CONTAINER_URI = "gcr.io/cloud-aiplatform/prediction/tf2-cpu.2-2:latest" +MODEL_SERVING_CONTAINER_URI = "us-docker.pkg.dev/vertex-ai/prediction/tf2-cpu.2-2:latest" REPLICA_COUNT = 1 MACHINE_TYPE = "n1-standard-4" ACCELERATOR_TYPE = "ACCELERATOR_TYPE_UNSPECIFIED" diff --git a/tests/system/providers/google/cloud/vertex_ai/example_vertex_ai_custom_job.py b/tests/system/providers/google/cloud/vertex_ai/example_vertex_ai_custom_job.py index b2856a28a23d0..8762feb85ba39 100644 --- a/tests/system/providers/google/cloud/vertex_ai/example_vertex_ai_custom_job.py +++ b/tests/system/providers/google/cloud/vertex_ai/example_vertex_ai_custom_job.py @@ -68,8 +68,8 @@ def TABULAR_DATASET(bucket_name): } -CONTAINER_URI = "gcr.io/cloud-aiplatform/training/tf-cpu.2-2:latest" -MODEL_SERVING_CONTAINER_URI = "gcr.io/cloud-aiplatform/prediction/tf2-cpu.2-2:latest" +CONTAINER_URI = "us-docker.pkg.dev/vertex-ai/training/tf-cpu.2-2:latest" +MODEL_SERVING_CONTAINER_URI = "us-docker.pkg.dev/vertex-ai/prediction/tf2-cpu.2-2:latest" REPLICA_COUNT = 1 # LOCAL_TRAINING_SCRIPT_PATH should be set for Airflow which is running on distributed system. diff --git a/tests/system/providers/google/cloud/vertex_ai/example_vertex_ai_custom_job_python_package.py b/tests/system/providers/google/cloud/vertex_ai/example_vertex_ai_custom_job_python_package.py index 33105d273f159..49a8d870bc394 100644 --- a/tests/system/providers/google/cloud/vertex_ai/example_vertex_ai_custom_job_python_package.py +++ b/tests/system/providers/google/cloud/vertex_ai/example_vertex_ai_custom_job_python_package.py @@ -68,8 +68,8 @@ def TABULAR_DATASET(bucket_name): } -CONTAINER_URI = "gcr.io/cloud-aiplatform/training/tf-cpu.2-2:latest" -MODEL_SERVING_CONTAINER_URI = "gcr.io/cloud-aiplatform/prediction/tf2-cpu.2-2:latest" +CONTAINER_URI = "us-docker.pkg.dev/vertex-ai/training/tf-cpu.2-2:latest" +MODEL_SERVING_CONTAINER_URI = "us-docker.pkg.dev/vertex-ai/prediction/tf2-cpu.2-2:latest" REPLICA_COUNT = 1 MACHINE_TYPE = "n1-standard-4" ACCELERATOR_TYPE = "ACCELERATOR_TYPE_UNSPECIFIED" diff --git a/tests/system/providers/google/cloud/vertex_ai/example_vertex_ai_model_service.py b/tests/system/providers/google/cloud/vertex_ai/example_vertex_ai_model_service.py index e6ad1e710e4c3..b06f8287798df 100644 --- a/tests/system/providers/google/cloud/vertex_ai/example_vertex_ai_model_service.py +++ b/tests/system/providers/google/cloud/vertex_ai/example_vertex_ai_model_service.py @@ -85,7 +85,7 @@ ), } -CONTAINER_URI = "gcr.io/cloud-aiplatform/training/tf-cpu.2-2:latest" +CONTAINER_URI = "us-docker.pkg.dev/vertex-ai/training/tf-cpu.2-2:latest" # LOCAL_TRAINING_SCRIPT_PATH should be set for Airflow which is running on distributed system. # For example in Composer the correct path is `gcs/data/california_housing_training_script.py`. @@ -99,7 +99,7 @@ }, "export_format_id": "custom-trained", } -MODEL_SERVING_CONTAINER_URI = "gcr.io/cloud-aiplatform/prediction/tf2-cpu.2-2:latest" +MODEL_SERVING_CONTAINER_URI = "us-docker.pkg.dev/vertex-ai/prediction/tf2-cpu.2-2:latest" MODEL_OBJ = { "display_name": f"model-{ENV_ID}", "artifact_uri": "{{ti.xcom_pull('custom_task')['artifactUri']}}", From 55ff735ccb90c9afc46ef0dabaf5f76df87043d5 Mon Sep 17 00:00:00 2001 From: GPK Date: Thu, 3 Oct 2024 21:36:09 +0100 Subject: [PATCH 138/802] fix PubSubAsyncHook in PubsubPullTrigger to use gcp_conn_id (#42671) --- .../providers/google/cloud/triggers/pubsub.py | 10 +++++++++- .../google/cloud/triggers/test_pubsub.py | 20 +++++++++++++++++++ 2 files changed, 29 insertions(+), 1 deletion(-) diff --git a/airflow/providers/google/cloud/triggers/pubsub.py b/airflow/providers/google/cloud/triggers/pubsub.py index db3fe409e942b..e98603006f725 100644 --- a/airflow/providers/google/cloud/triggers/pubsub.py +++ b/airflow/providers/google/cloud/triggers/pubsub.py @@ -19,6 +19,7 @@ from __future__ import annotations import asyncio +from functools import cached_property from typing import Any, AsyncIterator, Sequence from google.cloud.pubsub_v1.types import ReceivedMessage @@ -67,7 +68,6 @@ def __init__( self.poke_interval = poke_interval self.gcp_conn_id = gcp_conn_id self.impersonation_chain = impersonation_chain - self.hook = PubSubAsyncHook() def serialize(self) -> tuple[str, dict[str, Any]]: """Serialize PubsubPullTrigger arguments and classpath.""" @@ -113,3 +113,11 @@ async def message_acknowledgement(self, pulled_messages): messages=pulled_messages, ) self.log.info("Acknowledged ack_ids from subscription %s", self.subscription) + + @cached_property + def hook(self) -> PubSubAsyncHook: + return PubSubAsyncHook( + gcp_conn_id=self.gcp_conn_id, + impersonation_chain=self.impersonation_chain, + project_id=self.project_id, + ) diff --git a/tests/providers/google/cloud/triggers/test_pubsub.py b/tests/providers/google/cloud/triggers/test_pubsub.py index e1a4e178d2918..60acd2b7d4c2e 100644 --- a/tests/providers/google/cloud/triggers/test_pubsub.py +++ b/tests/providers/google/cloud/triggers/test_pubsub.py @@ -110,3 +110,23 @@ async def test_async_pubsub_pull_trigger_return_event(self, mock_pull): response = await trigger.run().asend(None) assert response == expected_event + + @mock.patch("airflow.providers.google.cloud.triggers.pubsub.PubSubAsyncHook") + def test_hook(self, mock_async_hook): + trigger = PubsubPullTrigger( + project_id=PROJECT_ID, + subscription="subscription", + max_messages=MAX_MESSAGES, + ack_messages=False, + poke_interval=TEST_POLL_INTERVAL, + gcp_conn_id=TEST_GCP_CONN_ID, + impersonation_chain=None, + ) + async_hook_actual = trigger.hook + + mock_async_hook.assert_called_once_with( + gcp_conn_id=trigger.gcp_conn_id, + impersonation_chain=trigger.impersonation_chain, + project_id=trigger.project_id, + ) + assert async_hook_actual == mock_async_hook.return_value From 286acfb163813f2fdab08d83f271ac3e6295ee54 Mon Sep 17 00:00:00 2001 From: Wei Lee Date: Fri, 4 Oct 2024 08:57:51 +0900 Subject: [PATCH 139/802] Rename dataset endpoints as asset endpoints (#42579) * feat(api_connexion): rename dataset_endpoint module as asset_endpoint * feat(api_connexion/openapi): rename tag Dataset as Asset * feat(api_connexion): rename create_dataset_event as create_asset_event * feat(api_connexion): rename schema CreateDatasetEvent as CreateAssetEvent * test(api_connexion): rename test_dataset_endpoint as test_asset_endpoint * feat(api_connexion): rename delete_dataset_queued_events as delete_asset_queued_events * feat(api_connexion): rename get_dataset_queued_events as get_asset_queued_events * feat(api_connexion): rename delete_dag_dataset_queued_events as delete_dag_asset_queued_events * feat(api_connexion): rename delete_dag_dataset_queued_event as delete_dag_asset_queued_event * feat(api_connexion): rename get_dag_dataset_queued_events as get_dag_asset_queued_events * feat(api_connexion): rename get_dag_dataset_queued_event as get_dag_asset_queued_event * refactor(api_connexion): remove unused dataset_id in _generate_queued_event_where_clause * feat(api_connexion): rename get_dataset_events as get_asset_events * feat(api_connexion): rename get_datasets as get_assets * feat(api_connexion): rename get_dataset as get_asset * feat(api_connexion/openapi): update api docs * feat(js): rename DatasetEvents as AssetEvents * feat(js): rename DatasetDetails as AssetDetails * feat(js): rename DatasetList as AssetList * feat(js/api): rename useUpstreamDatasetEvents as useUpstreamAssetEvents * feat(js/api): rename useDatasetsSummary as useAssetsSummary * feat(js/api): rename useDatasetDependencies as useAssetDependencies * feat(js/api): rename useDatasetEvents as useAssetEvents * feat(js/api): rename useDatasets as useAssets * feat(js/api): rename useDataset as useAsset * feat(api_connexion): rename get_upstream_dataset_events as get_upstream_asset_events * feat(api_connexion/openapi/v1): rename DatasetURI as AssetURI * feat(api_connexion/openapi/v1): rename DatasetCollection as AssetCollection * feat(api_connexion/openapi/v1): rename DagScheduleDatasetReference as DagScheduleAssetReference * feat(api_connexion/openapi/v1): rename TaskOutletDatasetReference as TaskOutletAssetReference * feat(js/api): rename DatasetEventCollection as AssetEventCollection * feat(api_connexion/openapi/v1): rename DatasetEvent as AssetEvent * feat(api_connexion/openapi/v1): rename Dataset as Asset * docs(api_connexion/openapi/v1): update dataset to asset in v1.yaml * feat(api_connexion): rename endpoint datasets as assets * test(api_connexion): rename dataset as asset * fix(api_connexion/openapi/v1): fix queued_events property name error * feat(api_fastapi): rename next_run_datasets as next_run_assets * test: resolve test conflict * docs(newsfragments): add newsfragments for dataset to asset endpoint rename * feat(js/api): rename datasetEvents as assetEvents * feat(js/api): rename variable datasetEvent as assetEvent * feat(js/api): rename dataset_api as asset_api --- ...{dataset_endpoint.py => asset_endpoint.py} | 41 +-- .../endpoints/dag_run_endpoint.py | 8 +- airflow/api_connexion/openapi/v1.yaml | 258 ++++++------- airflow/api_connexion/schemas/asset_schema.py | 10 +- airflow/api_fastapi/openapi/v1-generated.yaml | 2 +- airflow/api_fastapi/views/ui/assets.py | 2 +- .../ui/openapi-gen/requests/services.gen.ts | 2 +- airflow/ui/openapi-gen/requests/types.gen.ts | 2 +- airflow/www/static/js/api/index.ts | 28 +- .../js/api/{useDataset.ts => useAsset.ts} | 6 +- ...ependencies.ts => useAssetDependencies.ts} | 6 +- ...{useDatasetEvents.ts => useAssetEvents.ts} | 24 +- .../js/api/{useDatasets.ts => useAssets.ts} | 4 +- ...DatasetsSummary.ts => useAssetsSummary.ts} | 2 +- ...DatasetEvent.ts => useCreateAssetEvent.ts} | 19 +- ...setEvents.ts => useUpstreamAssetEvents.ts} | 22 +- .../static/js/components/DatasetEventCard.tsx | 33 +- .../js/components/SourceTaskInstance.tsx | 12 +- .../details/dagRun/DatasetTriggerEvents.tsx | 14 +- .../js/dag/details/graph/DatasetNode.tsx | 24 +- .../www/static/js/dag/details/graph/Node.tsx | 4 +- .../www/static/js/dag/details/graph/index.tsx | 40 +- .../www/static/js/dag/details/graph/utils.ts | 10 +- .../taskInstance/DatasetUpdateEvents.tsx | 14 +- airflow/www/static/js/datasetUtils.js | 8 +- .../{DatasetDetails.tsx => AssetDetails.tsx} | 12 +- .../{DatasetEvents.tsx => AssetEvents.tsx} | 20 +- ...tasetsList.test.tsx => AssetList.test.tsx} | 20 +- .../{DatasetsList.tsx => AssetsList.tsx} | 12 +- ...eDatasetEvent.tsx => CreateAssetEvent.tsx} | 12 +- .../www/static/js/datasets/Graph/index.tsx | 4 +- airflow/www/static/js/datasets/Main.tsx | 20 +- airflow/www/static/js/datasets/SearchBar.tsx | 2 +- airflow/www/static/js/types/api-generated.ts | 346 +++++++++--------- airflow/www/static/js/types/index.ts | 2 +- airflow/www/templates/airflow/dag.html | 4 +- airflow/www/templates/airflow/datasets.html | 6 +- airflow/www/templates/airflow/grid.html | 2 +- clients/python/README.md | 24 +- .../auth-manager/access-control.rst | 6 +- newsfragments/42579.significant.rst | 20 + ...set_endpoint.py => test_asset_endpoint.py} | 250 ++++++------- .../endpoints/test_dag_run_endpoint.py | 8 +- .../schemas/test_dataset_schema.py | 16 +- tests/api_fastapi/views/ui/test_assets.py | 4 +- 45 files changed, 696 insertions(+), 689 deletions(-) rename airflow/api_connexion/endpoints/{dataset_endpoint.py => asset_endpoint.py} (92%) rename airflow/www/static/js/api/{useDataset.ts => useAsset.ts} (86%) rename airflow/www/static/js/api/{useDatasetDependencies.ts => useAssetDependencies.ts} (94%) rename airflow/www/static/js/api/{useDatasetEvents.ts => useAssetEvents.ts} (80%) rename airflow/www/static/js/api/{useDatasets.ts => useAssets.ts} (90%) rename airflow/www/static/js/api/{useDatasetsSummary.ts => useAssetsSummary.ts} (98%) rename airflow/www/static/js/api/{useCreateDatasetEvent.ts => useCreateAssetEvent.ts} (77%) rename airflow/www/static/js/api/{useUpstreamDatasetEvents.ts => useUpstreamAssetEvents.ts} (67%) rename airflow/www/static/js/datasets/{DatasetDetails.tsx => AssetDetails.tsx} (91%) rename airflow/www/static/js/datasets/{DatasetEvents.tsx => AssetEvents.tsx} (87%) rename airflow/www/static/js/datasets/{DatasetsList.test.tsx => AssetList.test.tsx} (87%) rename airflow/www/static/js/datasets/{DatasetsList.tsx => AssetsList.tsx} (95%) rename airflow/www/static/js/datasets/{CreateDatasetEvent.tsx => CreateAssetEvent.tsx} (88%) create mode 100644 newsfragments/42579.significant.rst rename tests/api_connexion/endpoints/{test_dataset_endpoint.py => test_asset_endpoint.py} (74%) diff --git a/airflow/api_connexion/endpoints/dataset_endpoint.py b/airflow/api_connexion/endpoints/asset_endpoint.py similarity index 92% rename from airflow/api_connexion/endpoints/dataset_endpoint.py rename to airflow/api_connexion/endpoints/asset_endpoint.py index 95c3bead3da52..cbbe542ea7987 100644 --- a/airflow/api_connexion/endpoints/dataset_endpoint.py +++ b/airflow/api_connexion/endpoints/asset_endpoint.py @@ -57,13 +57,13 @@ from airflow.api_connexion.types import APIResponse -RESOURCE_EVENT_PREFIX = "dataset" +RESOURCE_EVENT_PREFIX = "asset" @security.requires_access_asset("GET") @provide_session -def get_dataset(*, uri: str, session: Session = NEW_SESSION) -> APIResponse: - """Get an asset .""" +def get_asset(*, uri: str, session: Session = NEW_SESSION) -> APIResponse: + """Get an asset.""" asset = session.scalar( select(AssetModel) .where(AssetModel.uri == uri) @@ -80,7 +80,7 @@ def get_dataset(*, uri: str, session: Session = NEW_SESSION) -> APIResponse: @security.requires_access_asset("GET") @format_parameters({"limit": check_limit}) @provide_session -def get_datasets( +def get_assets( *, limit: int, offset: int = 0, @@ -109,18 +109,18 @@ def get_datasets( .offset(offset) .limit(limit) ).all() - return asset_collection_schema.dump(AssetCollection(datasets=assets, total_entries=total_entries)) + return asset_collection_schema.dump(AssetCollection(assets=assets, total_entries=total_entries)) @security.requires_access_asset("GET") @provide_session @format_parameters({"limit": check_limit}) -def get_dataset_events( +def get_asset_events( *, limit: int, offset: int = 0, order_by: str = "timestamp", - dataset_id: int | None = None, + asset_id: int | None = None, source_dag_id: str | None = None, source_task_id: str | None = None, source_run_id: str | None = None, @@ -132,8 +132,8 @@ def get_dataset_events( query = select(AssetEvent) - if dataset_id: - query = query.where(AssetEvent.dataset_id == dataset_id) + if asset_id: + query = query.where(AssetEvent.dataset_id == asset_id) if source_dag_id: query = query.where(AssetEvent.source_dag_id == source_dag_id) if source_task_id: @@ -149,14 +149,13 @@ def get_dataset_events( query = apply_sorting(query, order_by, {}, allowed_attrs) events = session.scalars(query.offset(offset).limit(limit)).all() return asset_event_collection_schema.dump( - AssetEventCollection(dataset_events=events, total_entries=total_entries) + AssetEventCollection(asset_events=events, total_entries=total_entries) ) def _generate_queued_event_where_clause( *, dag_id: str | None = None, - dataset_id: int | None = None, uri: str | None = None, before: str | None = None, permitted_dag_ids: set[str] | None = None, @@ -165,8 +164,6 @@ def _generate_queued_event_where_clause( where_clause = [] if dag_id is not None: where_clause.append(AssetDagRunQueue.target_dag_id == dag_id) - if dataset_id is not None: - where_clause.append(AssetDagRunQueue.dataset_id == dataset_id) if uri is not None: where_clause.append( AssetDagRunQueue.dataset_id.in_( @@ -183,7 +180,7 @@ def _generate_queued_event_where_clause( @security.requires_access_asset("GET") @security.requires_access_dag("GET") @provide_session -def get_dag_dataset_queued_event( +def get_dag_asset_queued_event( *, dag_id: str, uri: str, before: str | None = None, session: Session = NEW_SESSION ) -> APIResponse: """Get a queued asset event for a DAG.""" @@ -206,7 +203,7 @@ def get_dag_dataset_queued_event( @security.requires_access_dag("GET") @provide_session @action_logging -def delete_dag_dataset_queued_event( +def delete_dag_asset_queued_event( *, dag_id: str, uri: str, before: str | None = None, session: Session = NEW_SESSION ) -> APIResponse: """Delete a queued asset event for a DAG.""" @@ -224,7 +221,7 @@ def delete_dag_dataset_queued_event( @security.requires_access_asset("GET") @security.requires_access_dag("GET") @provide_session -def get_dag_dataset_queued_events( +def get_dag_asset_queued_events( *, dag_id: str, before: str | None = None, session: Session = NEW_SESSION ) -> APIResponse: """Get queued asset events for a DAG.""" @@ -253,7 +250,7 @@ def get_dag_dataset_queued_events( @security.requires_access_dag("GET") @action_logging @provide_session -def delete_dag_dataset_queued_events( +def delete_dag_asset_queued_events( *, dag_id: str, before: str | None = None, session: Session = NEW_SESSION ) -> APIResponse: """Delete queued asset events for a DAG.""" @@ -271,7 +268,7 @@ def delete_dag_dataset_queued_events( @security.requires_access_asset("GET") @provide_session -def get_dataset_queued_events( +def get_asset_queued_events( *, uri: str, before: str | None = None, session: Session = NEW_SESSION ) -> APIResponse: """Get queued asset events for an asset.""" @@ -303,7 +300,7 @@ def get_dataset_queued_events( @security.requires_access_asset("DELETE") @action_logging @provide_session -def delete_dataset_queued_events( +def delete_asset_queued_events( *, uri: str, before: str | None = None, session: Session = NEW_SESSION ) -> APIResponse: """Delete queued asset events for an asset.""" @@ -325,7 +322,7 @@ def delete_dataset_queued_events( @security.requires_access_asset("POST") @provide_session @action_logging -def create_dataset_event(session: Session = NEW_SESSION) -> APIResponse: +def create_asset_event(session: Session = NEW_SESSION) -> APIResponse: """Create asset event.""" body = get_json_request_dict() try: @@ -333,7 +330,7 @@ def create_dataset_event(session: Session = NEW_SESSION) -> APIResponse: except ValidationError as err: raise BadRequest(detail=str(err)) - uri = json_body["dataset_uri"] + uri = json_body["asset_uri"] asset = session.scalar(select(AssetModel).where(AssetModel.uri == uri).limit(1)) if not asset: raise NotFound(title="Asset not found", detail=f"Asset with uri: '{uri}' not found") @@ -341,7 +338,7 @@ def create_dataset_event(session: Session = NEW_SESSION) -> APIResponse: extra = json_body.get("extra", {}) extra["from_rest_api"] = True asset_event = asset_manager.register_asset_change( - asset=Asset(uri), + asset=Asset(uri=uri), timestamp=timestamp, extra=extra, session=session, diff --git a/airflow/api_connexion/endpoints/dag_run_endpoint.py b/airflow/api_connexion/endpoints/dag_run_endpoint.py index 02d4663837f4e..44891c0ef2c84 100644 --- a/airflow/api_connexion/endpoints/dag_run_endpoint.py +++ b/airflow/api_connexion/endpoints/dag_run_endpoint.py @@ -114,10 +114,8 @@ def get_dag_run( @security.requires_access_dag("GET", DagAccessEntity.RUN) @security.requires_access_asset("GET") @provide_session -def get_upstream_dataset_events( - *, dag_id: str, dag_run_id: str, session: Session = NEW_SESSION -) -> APIResponse: - """If dag run is dataset-triggered, return the asset events that triggered it.""" +def get_upstream_asset_events(*, dag_id: str, dag_run_id: str, session: Session = NEW_SESSION) -> APIResponse: + """If dag run is asset-triggered, return the asset events that triggered it.""" dag_run: DagRun | None = session.scalar( select(DagRun).where( DagRun.dag_id == dag_id, @@ -131,7 +129,7 @@ def get_upstream_dataset_events( ) events = dag_run.consumed_dataset_events return asset_event_collection_schema.dump( - AssetEventCollection(dataset_events=events, total_entries=len(events)) + AssetEventCollection(asset_events=events, total_entries=len(events)) ) diff --git a/airflow/api_connexion/openapi/v1.yaml b/airflow/api_connexion/openapi/v1.yaml index 15ad6fd8a4f63..828a3af25e879 100644 --- a/airflow/api_connexion/openapi/v1.yaml +++ b/airflow/api_connexion/openapi/v1.yaml @@ -1181,26 +1181,26 @@ paths: "404": $ref: "#/components/responses/NotFound" - /dags/{dag_id}/dagRuns/{dag_run_id}/upstreamDatasetEvents: + /dags/{dag_id}/dagRuns/{dag_run_id}/upstreamAssetEvents: parameters: - $ref: "#/components/parameters/DAGID" - $ref: "#/components/parameters/DAGRunID" get: - summary: Get dataset events for a DAG run + summary: Get asset events for a DAG run description: | - Get datasets for a dag run. + Get asset for a dag run. *New in version 2.4.0* x-openapi-router-controller: airflow.api_connexion.endpoints.dag_run_endpoint - operationId: get_upstream_dataset_events - tags: [DAGRun, Dataset] + operationId: get_upstream_asset_events + tags: [DAGRun, Asset] responses: "200": description: Success. content: application/json: schema: - $ref: "#/components/schemas/DatasetEventCollection" + $ref: "#/components/schemas/AssetEventCollection" "401": $ref: "#/components/responses/Unauthenticated" "403": @@ -1245,22 +1245,22 @@ paths: "404": $ref: "#/components/responses/NotFound" - /dags/{dag_id}/datasets/queuedEvent/{uri}: + /dags/{dag_id}/assets/queuedEvent/{uri}: parameters: - $ref: "#/components/parameters/DAGID" - - $ref: "#/components/parameters/DatasetURI" + - $ref: "#/components/parameters/AssetURI" get: - summary: Get a queued Dataset event for a DAG + summary: Get a queued asset event for a DAG description: | - Get a queued Dataset event for a DAG. + Get a queued asset event for a DAG. *New in version 2.9.0* - x-openapi-router-controller: airflow.api_connexion.endpoints.dataset_endpoint - operationId: get_dag_dataset_queued_event + x-openapi-router-controller: airflow.api_connexion.endpoints.asset_endpoint + operationId: get_dag_asset_queued_event parameters: - $ref: "#/components/parameters/Before" - tags: [Dataset] + tags: [Asset] responses: "200": description: Success. @@ -1276,16 +1276,16 @@ paths: $ref: "#/components/responses/NotFound" delete: - summary: Delete a queued Dataset event for a DAG. + summary: Delete a queued Asset event for a DAG. description: | - Delete a queued Dataset event for a DAG. + Delete a queued Asset event for a DAG. *New in version 2.9.0* - x-openapi-router-controller: airflow.api_connexion.endpoints.dataset_endpoint - operationId: delete_dag_dataset_queued_event + x-openapi-router-controller: airflow.api_connexion.endpoints.asset_endpoint + operationId: delete_dag_asset_queued_event parameters: - $ref: "#/components/parameters/Before" - tags: [Dataset] + tags: [Asset] responses: "204": description: Success. @@ -1298,21 +1298,21 @@ paths: "404": $ref: "#/components/responses/NotFound" - /dags/{dag_id}/datasets/queuedEvent: + /dags/{dag_id}/assets/queuedEvent: parameters: - $ref: "#/components/parameters/DAGID" get: - summary: Get queued Dataset events for a DAG. + summary: Get queued Asset events for a DAG. description: | - Get queued Dataset events for a DAG. + Get queued Asset events for a DAG. *New in version 2.9.0* - x-openapi-router-controller: airflow.api_connexion.endpoints.dataset_endpoint - operationId: get_dag_dataset_queued_events + x-openapi-router-controller: airflow.api_connexion.endpoints.asset_endpoint + operationId: get_dag_asset_queued_events parameters: - $ref: "#/components/parameters/Before" - tags: [Dataset] + tags: [Asset] responses: "200": description: Success. @@ -1328,16 +1328,16 @@ paths: $ref: "#/components/responses/NotFound" delete: - summary: Delete queued Dataset events for a DAG. + summary: Delete queued Asset events for a DAG. description: | - Delete queued Dataset events for a DAG. + Delete queued Asset events for a DAG. *New in version 2.9.0* - x-openapi-router-controller: airflow.api_connexion.endpoints.dataset_endpoint - operationId: delete_dag_dataset_queued_events + x-openapi-router-controller: airflow.api_connexion.endpoints.asset_endpoint + operationId: delete_dag_asset_queued_events parameters: - $ref: "#/components/parameters/Before" - tags: [Dataset] + tags: [Asset] responses: "204": description: Success. @@ -1371,21 +1371,21 @@ paths: "404": $ref: "#/components/responses/NotFound" - /datasets/queuedEvent/{uri}: + /assets/queuedEvent/{uri}: parameters: - - $ref: "#/components/parameters/DatasetURI" + - $ref: "#/components/parameters/AssetURI" get: - summary: Get queued Dataset events for a Dataset. + summary: Get queued Asset events for an Asset. description: | - Get queued Dataset events for a Dataset + Get queued Asset events for an Asset *New in version 2.9.0* - x-openapi-router-controller: airflow.api_connexion.endpoints.dataset_endpoint - operationId: get_dataset_queued_events + x-openapi-router-controller: airflow.api_connexion.endpoints.asset_endpoint + operationId: get_asset_queued_events parameters: - $ref: "#/components/parameters/Before" - tags: [Dataset] + tags: [Asset] responses: "200": description: Success. @@ -1401,16 +1401,16 @@ paths: $ref: "#/components/responses/NotFound" delete: - summary: Delete queued Dataset events for a Dataset. + summary: Delete queued Asset events for an Asset. description: | - Delete queued Dataset events for a Dataset. + Delete queued Asset events for a Asset. *New in version 2.9.0* - x-openapi-router-controller: airflow.api_connexion.endpoints.dataset_endpoint - operationId: delete_dataset_queued_events + x-openapi-router-controller: airflow.api_connexion.endpoints.asset_endpoint + operationId: delete_asset_queued_events parameters: - $ref: "#/components/parameters/Before" - tags: [Dataset] + tags: [Asset] responses: "204": description: Success. @@ -2517,12 +2517,12 @@ paths: "403": $ref: "#/components/responses/PermissionDenied" - /datasets: + /assets: get: - summary: List datasets - x-openapi-router-controller: airflow.api_connexion.endpoints.dataset_endpoint - operationId: get_datasets - tags: [Dataset] + summary: List assets + x-openapi-router-controller: airflow.api_connexion.endpoints.asset_endpoint + operationId: get_assets + tags: [Asset] parameters: - $ref: "#/components/parameters/PageLimit" - $ref: "#/components/parameters/PageOffset" @@ -2533,14 +2533,14 @@ paths: type: string required: false description: | - If set, only return datasets with uris matching this pattern. + If set, only return assets with uris matching this pattern. - name: dag_ids in: query schema: type: string required: false description: | - One or more DAG IDs separated by commas to filter datasets by associated DAGs either consuming or producing. + One or more DAG IDs separated by commas to filter assets by associated DAGs either consuming or producing. *New in version 2.9.0* responses: @@ -2549,28 +2549,28 @@ paths: content: application/json: schema: - $ref: "#/components/schemas/DatasetCollection" + $ref: "#/components/schemas/AssetCollection" "401": $ref: "#/components/responses/Unauthenticated" "403": $ref: "#/components/responses/PermissionDenied" - /datasets/{uri}: + /assets/{uri}: parameters: - - $ref: "#/components/parameters/DatasetURI" + - $ref: "#/components/parameters/AssetURI" get: - summary: Get a dataset - description: Get a dataset by uri. - x-openapi-router-controller: airflow.api_connexion.endpoints.dataset_endpoint - operationId: get_dataset - tags: [Dataset] + summary: Get an asset + description: Get an asset by uri. + x-openapi-router-controller: airflow.api_connexion.endpoints.asset_endpoint + operationId: get_asset + tags: [Asset] responses: "200": description: Success. content: application/json: schema: - $ref: "#/components/schemas/Dataset" + $ref: "#/components/schemas/Asset" "401": $ref: "#/components/responses/Unauthenticated" "403": @@ -2578,18 +2578,18 @@ paths: "404": $ref: "#/components/responses/NotFound" - /datasets/events: + /assets/events: get: - summary: Get dataset events - description: Get dataset events - x-openapi-router-controller: airflow.api_connexion.endpoints.dataset_endpoint - operationId: get_dataset_events - tags: [Dataset] + summary: Get asset events + description: Get asset events + x-openapi-router-controller: airflow.api_connexion.endpoints.asset_endpoint + operationId: get_asset_events + tags: [Asset] parameters: - $ref: "#/components/parameters/PageLimit" - $ref: "#/components/parameters/PageOffset" - $ref: "#/components/parameters/OrderBy" - - $ref: "#/components/parameters/FilterDatasetID" + - $ref: "#/components/parameters/FilterAssetID" - $ref: "#/components/parameters/FilterSourceDAGID" - $ref: "#/components/parameters/FilterSourceTaskID" - $ref: "#/components/parameters/FilterSourceRunID" @@ -2600,7 +2600,7 @@ paths: content: application/json: schema: - $ref: "#/components/schemas/DatasetEventCollection" + $ref: "#/components/schemas/AssetEventCollection" "401": $ref: "#/components/responses/Unauthenticated" "403": @@ -2608,24 +2608,24 @@ paths: "404": $ref: "#/components/responses/NotFound" post: - summary: Create dataset event - description: Create dataset event - x-openapi-router-controller: airflow.api_connexion.endpoints.dataset_endpoint - operationId: create_dataset_event - tags: [Dataset] + summary: Create asset event + description: Create asset event + x-openapi-router-controller: airflow.api_connexion.endpoints.asset_endpoint + operationId: create_asset_event + tags: [Asset] requestBody: required: true content: application/json: schema: - $ref: '#/components/schemas/CreateDatasetEvent' + $ref: '#/components/schemas/CreateAssetEvent' responses: '200': description: Success. content: application/json: schema: - $ref: '#/components/schemas/DatasetEvent' + $ref: '#/components/schemas/AssetEvent' "400": $ref: "#/components/responses/BadRequest" '401': @@ -4133,7 +4133,7 @@ components: nullable: true dataset_expression: type: object - description: Nested dataset any/all conditions + description: Nested asset any/all conditions nullable: true doc_md: type: string @@ -4507,133 +4507,133 @@ components: $ref: "#/components/schemas/Resource" description: The permission resource - Dataset: + Asset: description: | - A dataset item. + An asset item. *New in version 2.4.0* type: object properties: id: type: integer - description: The dataset id + description: The asset id uri: type: string - description: The dataset uri + description: The asset uri nullable: false extra: type: object - description: The dataset extra + description: The asset extra nullable: true created_at: type: string - description: The dataset creation time + description: The asset creation time nullable: false updated_at: type: string - description: The dataset update time + description: The asset update time nullable: false consuming_dags: type: array items: - $ref: "#/components/schemas/DagScheduleDatasetReference" + $ref: "#/components/schemas/DagScheduleAssetReference" producing_tasks: type: array items: - $ref: "#/components/schemas/TaskOutletDatasetReference" + $ref: "#/components/schemas/TaskOutletAssetReference" - TaskOutletDatasetReference: + TaskOutletAssetReference: description: | - A datasets reference to an upstream task. + An asset reference to an upstream task. *New in version 2.4.0* type: object properties: dag_id: type: string - description: The DAG ID that updates the dataset. + description: The DAG ID that updates the asset. nullable: true task_id: type: string - description: The task ID that updates the dataset. + description: The task ID that updates the asset. nullable: true created_at: type: string - description: The dataset creation time + description: The asset creation time nullable: false updated_at: type: string - description: The dataset update time + description: The asset update time nullable: false - DagScheduleDatasetReference: + DagScheduleAssetReference: description: | - A datasets reference to a downstream DAG. + An asset reference to a downstream DAG. *New in version 2.4.0* type: object properties: dag_id: type: string - description: The DAG ID that depends on the dataset. + description: The DAG ID that depends on the asset. nullable: true created_at: type: string - description: The dataset reference creation time + description: The asset reference creation time nullable: false updated_at: type: string - description: The dataset reference update time + description: The asset reference update time nullable: false - DatasetCollection: + AssetCollection: description: | - A collection of datasets. + A collection of assets. *New in version 2.4.0* type: object allOf: - type: object properties: - datasets: + assets: type: array items: - $ref: "#/components/schemas/Dataset" + $ref: "#/components/schemas/Asset" - $ref: "#/components/schemas/CollectionInfo" - DatasetEvent: + AssetEvent: description: | - A dataset event. + An asset event. *New in version 2.4.0* type: object properties: dataset_id: type: integer - description: The dataset id + description: The asset id dataset_uri: type: string - description: The URI of the dataset + description: The URI of the asset nullable: false extra: type: object - description: The dataset event extra + description: The asset event extra nullable: true source_dag_id: type: string - description: The DAG ID that updated the dataset. + description: The DAG ID that updated the asset. nullable: true source_task_id: type: string - description: The task ID that updated the dataset. + description: The task ID that updated the asset. nullable: true source_run_id: type: string - description: The DAG run ID that updated the dataset. + description: The DAG run ID that updated the asset. nullable: true source_map_index: type: integer - description: The task map index that updated the dataset. + description: The task map index that updated the asset. nullable: true created_dagruns: type: array @@ -4641,21 +4641,21 @@ components: $ref: "#/components/schemas/BasicDAGRun" timestamp: type: string - description: The dataset event creation time + description: The asset event creation time nullable: false - CreateDatasetEvent: + CreateAssetEvent: type: object required: - - dataset_uri + - asset_uri properties: - dataset_uri: + asset_uri: type: string - description: The URI of the dataset + description: The URI of the asset nullable: false extra: type: object - description: The dataset event extra + description: The asset event extra nullable: true QueuedEvent: @@ -4663,7 +4663,7 @@ components: properties: uri: type: string - description: The datata uri. + description: The asset uri. dag_id: type: string description: The DAG ID. @@ -4674,14 +4674,14 @@ components: QueuedEventCollection: description: | - A collection of Dataset Dag Run Queues. + A collection of asset Dag Run Queues. *New in version 2.9.0* type: object allOf: - type: object properties: - datasets: + queued_events: type: array items: $ref: "#/components/schemas/QueuedEvent" @@ -4737,19 +4737,19 @@ components: state: $ref: "#/components/schemas/DagState" - DatasetEventCollection: + AssetEventCollection: description: | - A collection of dataset events. + A collection of asset events. *New in version 2.4.0* type: object allOf: - type: object properties: - dataset_events: + asset_events: type: array items: - $ref: "#/components/schemas/DatasetEvent" + $ref: "#/components/schemas/AssetEvent" - $ref: "#/components/schemas/CollectionInfo" # Configuration @@ -5545,14 +5545,14 @@ components: required: true description: The import error ID. - DatasetURI: + AssetURI: in: path name: uri schema: type: string format: path required: true - description: The encoded Dataset URI + description: The encoded Asset URI PoolName: in: path @@ -5733,40 +5733,40 @@ components: *New in version 2.2.0* - FilterDatasetID: + FilterAssetID: in: query - name: dataset_id + name: asset_id schema: type: integer - description: The Dataset ID that updated the dataset. + description: The Asset ID that updated the asset. FilterSourceDAGID: in: query name: source_dag_id schema: type: string - description: The DAG ID that updated the dataset. + description: The DAG ID that updated the asset. FilterSourceTaskID: in: query name: source_task_id schema: type: string - description: The task ID that updated the dataset. + description: The task ID that updated the asset. FilterSourceRunID: in: query name: source_run_id schema: type: string - description: The DAG run ID that updated the dataset. + description: The DAG run ID that updated the asset. FilterSourceMapIndex: in: query name: source_map_index schema: type: integer - description: The map index that updated the dataset. + description: The map index that updated the asset. FilterMapIndex: in: query @@ -6024,12 +6024,12 @@ components: security: [] tags: + - name: Asset - name: Config - name: Connection - name: DAG - name: DAGRun - name: DagWarning - - name: Dataset - name: EventLog - name: ImportError - name: Monitoring diff --git a/airflow/api_connexion/schemas/asset_schema.py b/airflow/api_connexion/schemas/asset_schema.py index 791941f42016d..662f73a50d8b9 100644 --- a/airflow/api_connexion/schemas/asset_schema.py +++ b/airflow/api_connexion/schemas/asset_schema.py @@ -93,14 +93,14 @@ class Meta: class AssetCollection(NamedTuple): """List of Assets with meta.""" - datasets: list[AssetModel] + assets: list[AssetModel] total_entries: int class AssetCollectionSchema(Schema): """Asset Collection Schema.""" - datasets = fields.List(fields.Nested(AssetSchema)) + assets = fields.List(fields.Nested(AssetSchema)) total_entries = fields.Int() @@ -150,21 +150,21 @@ class Meta: class AssetEventCollection(NamedTuple): """List of Asset events with meta.""" - dataset_events: list[AssetEvent] + asset_events: list[AssetEvent] total_entries: int class AssetEventCollectionSchema(Schema): """Asset Event Collection Schema.""" - dataset_events = fields.List(fields.Nested(AssetEventSchema)) + asset_events = fields.List(fields.Nested(AssetEventSchema)) total_entries = fields.Int() class CreateAssetEventSchema(Schema): """Create Asset Event Schema.""" - dataset_uri = fields.String() + asset_uri = fields.String() extra = JsonObjectField() diff --git a/airflow/api_fastapi/openapi/v1-generated.yaml b/airflow/api_fastapi/openapi/v1-generated.yaml index ce488a996af47..23c4ecf545d9f 100644 --- a/airflow/api_fastapi/openapi/v1-generated.yaml +++ b/airflow/api_fastapi/openapi/v1-generated.yaml @@ -7,7 +7,7 @@ info: Users should not rely on those but use the public ones instead. version: 0.1.0 paths: - /ui/next_run_datasets/{dag_id}: + /ui/next_run_assets/{dag_id}: get: tags: - Asset diff --git a/airflow/api_fastapi/views/ui/assets.py b/airflow/api_fastapi/views/ui/assets.py index 01cc9fd1cfbff..4a4ad1d0df9b4 100644 --- a/airflow/api_fastapi/views/ui/assets.py +++ b/airflow/api_fastapi/views/ui/assets.py @@ -30,7 +30,7 @@ assets_router = AirflowRouter(tags=["Asset"]) -@assets_router.get("/next_run_datasets/{dag_id}", include_in_schema=False) +@assets_router.get("/next_run_assets/{dag_id}", include_in_schema=False) async def next_run_assets( dag_id: str, request: Request, diff --git a/airflow/ui/openapi-gen/requests/services.gen.ts b/airflow/ui/openapi-gen/requests/services.gen.ts index 0aefb56d06e66..0e91fa416571e 100644 --- a/airflow/ui/openapi-gen/requests/services.gen.ts +++ b/airflow/ui/openapi-gen/requests/services.gen.ts @@ -30,7 +30,7 @@ export class AssetService { ): CancelablePromise { return __request(OpenAPI, { method: "GET", - url: "/ui/next_run_datasets/{dag_id}", + url: "/ui/next_run_assets/{dag_id}", path: { dag_id: data.dagId, }, diff --git a/airflow/ui/openapi-gen/requests/types.gen.ts b/airflow/ui/openapi-gen/requests/types.gen.ts index c37106abc8fcd..b87a172363584 100644 --- a/airflow/ui/openapi-gen/requests/types.gen.ts +++ b/airflow/ui/openapi-gen/requests/types.gen.ts @@ -203,7 +203,7 @@ export type DeleteConnectionData = { export type DeleteConnectionResponse = void; export type $OpenApiTs = { - "/ui/next_run_datasets/{dag_id}": { + "/ui/next_run_assets/{dag_id}": { get: { req: NextRunAssetsData; res: { diff --git a/airflow/www/static/js/api/index.ts b/airflow/www/static/js/api/index.ts index a4a45a08bfef8..c2a9885b2c7ea 100644 --- a/airflow/www/static/js/api/index.ts +++ b/airflow/www/static/js/api/index.ts @@ -32,14 +32,14 @@ import useMarkTaskDryRun from "./useMarkTaskDryRun"; import useGraphData from "./useGraphData"; import useGridData from "./useGridData"; import useMappedInstances from "./useMappedInstances"; -import useDatasets from "./useDatasets"; -import useDatasetsSummary from "./useDatasetsSummary"; -import useDataset from "./useDataset"; -import useDatasetDependencies from "./useDatasetDependencies"; -import useDatasetEvents from "./useDatasetEvents"; +import useAssets from "./useAssets"; +import useAssetsSummary from "./useAssetsSummary"; +import useAsset from "./useAsset"; +import useAssetDependencies from "./useAssetDependencies"; +import useAssetEvents from "./useAssetEvents"; import useSetDagRunNote from "./useSetDagRunNote"; import useSetTaskInstanceNote from "./useSetTaskInstanceNote"; -import useUpstreamDatasetEvents from "./useUpstreamDatasetEvents"; +import useUpstreamAssetEvents from "./useUpstreamAssetEvents"; import useTaskInstance from "./useTaskInstance"; import useTaskFailedDependency from "./useTaskFailedDependency"; import useDag from "./useDag"; @@ -53,7 +53,7 @@ import useHistoricalMetricsData from "./useHistoricalMetricsData"; import { useTaskXcomEntry, useTaskXcomCollection } from "./useTaskXcom"; import useEventLogs from "./useEventLogs"; import useCalendarData from "./useCalendarData"; -import useCreateDatasetEvent from "./useCreateDatasetEvent"; +import useCreateAssetEvent from "./useCreateAssetEvent"; import useRenderedK8s from "./useRenderedK8s"; import useTaskDetail from "./useTaskDetail"; import useTIHistory from "./useTIHistory"; @@ -85,11 +85,11 @@ export { useDagDetails, useDagRuns, useDags, - useDataset, - useDatasets, - useDatasetDependencies, - useDatasetEvents, - useDatasetsSummary, + useAsset, + useAssets, + useAssetDependencies, + useAssetEvents, + useAssetsSummary, useExtraLinks, useGraphData, useGridData, @@ -105,14 +105,14 @@ export { useSetDagRunNote, useSetTaskInstanceNote, useTaskInstance, - useUpstreamDatasetEvents, + useUpstreamAssetEvents, useHistoricalMetricsData, useTaskXcomEntry, useTaskXcomCollection, useTaskFailedDependency, useEventLogs, useCalendarData, - useCreateDatasetEvent, + useCreateAssetEvent, useRenderedK8s, useTaskDetail, useTIHistory, diff --git a/airflow/www/static/js/api/useDataset.ts b/airflow/www/static/js/api/useAsset.ts similarity index 86% rename from airflow/www/static/js/api/useDataset.ts rename to airflow/www/static/js/api/useAsset.ts index 4793464fac378..b490ca6e46565 100644 --- a/airflow/www/static/js/api/useDataset.ts +++ b/airflow/www/static/js/api/useAsset.ts @@ -27,12 +27,12 @@ interface Props { uri: string; } -export default function useDataset({ uri }: Props) { +export default function useAsset({ uri }: Props) { return useQuery(["dataset", uri], () => { - const datasetUrl = getMetaValue("dataset_api").replace( + const datasetUrl = getMetaValue("asset_api").replace( "__URI__", encodeURIComponent(uri) ); - return axios.get(datasetUrl); + return axios.get(datasetUrl); }); } diff --git a/airflow/www/static/js/api/useDatasetDependencies.ts b/airflow/www/static/js/api/useAssetDependencies.ts similarity index 94% rename from airflow/www/static/js/api/useDatasetDependencies.ts rename to airflow/www/static/js/api/useAssetDependencies.ts index d2ba627f64458..11e7219c53fc8 100644 --- a/airflow/www/static/js/api/useDatasetDependencies.ts +++ b/airflow/www/static/js/api/useAssetDependencies.ts @@ -82,15 +82,15 @@ const formatDependencies = async ({ edges, nodes }: DatasetDependencies) => { return graph as DatasetGraph; }; -export default function useDatasetDependencies() { +export default function useAssetDependencies() { return useQuery("datasetDependencies", async () => { const datasetDepsUrl = getMetaValue("dataset_dependencies_url"); return axios.get(datasetDepsUrl); }); } -export const useDatasetGraphs = () => { - const { data: datasetDependencies } = useDatasetDependencies(); +export const useAssetGraphs = () => { + const { data: datasetDependencies } = useAssetDependencies(); return useQuery(["datasetGraphs", datasetDependencies], () => { if (datasetDependencies) { return formatDependencies(datasetDependencies); diff --git a/airflow/www/static/js/api/useDatasetEvents.ts b/airflow/www/static/js/api/useAssetEvents.ts similarity index 80% rename from airflow/www/static/js/api/useDatasetEvents.ts rename to airflow/www/static/js/api/useAssetEvents.ts index 30e4670a87d3e..068bb471ef64e 100644 --- a/airflow/www/static/js/api/useDatasetEvents.ts +++ b/airflow/www/static/js/api/useAssetEvents.ts @@ -23,16 +23,16 @@ import { useQuery, UseQueryOptions } from "react-query"; import { getMetaValue } from "src/utils"; import URLSearchParamsWrapper from "src/utils/URLSearchParamWrapper"; import type { - DatasetEventCollection, - GetDatasetEventsVariables, + AssetEventCollection, + GetAssetEventsVariables, } from "src/types/api-generated"; -interface Props extends GetDatasetEventsVariables { - options?: UseQueryOptions; +interface Props extends GetAssetEventsVariables { + options?: UseQueryOptions; } -const useDatasetEvents = ({ - datasetId, +const useAssetEvents = ({ + assetId, sourceDagId, sourceRunId, sourceTaskId, @@ -42,10 +42,10 @@ const useDatasetEvents = ({ orderBy, options, }: Props) => { - const query = useQuery( + const query = useQuery( [ "datasets-events", - datasetId, + assetId, sourceDagId, sourceRunId, sourceTaskId, @@ -55,14 +55,14 @@ const useDatasetEvents = ({ orderBy, ], () => { - const datasetsUrl = getMetaValue("dataset_events_api"); + const datasetsUrl = getMetaValue("asset_events_api"); const params = new URLSearchParamsWrapper(); if (limit) params.set("limit", limit.toString()); if (offset) params.set("offset", offset.toString()); if (orderBy) params.set("order_by", orderBy); - if (datasetId) params.set("dataset_id", datasetId.toString()); + if (assetId) params.set("asset_id", assetId.toString()); if (sourceDagId) params.set("source_dag_id", sourceDagId); if (sourceRunId) params.set("source_run_id", sourceRunId); if (sourceTaskId) params.set("source_task_id", sourceTaskId); @@ -80,8 +80,8 @@ const useDatasetEvents = ({ ); return { ...query, - data: query.data ?? { datasetEvents: [], totalEntries: 0 }, + data: query.data ?? { assetEvents: [], totalEntries: 0 }, }; }; -export default useDatasetEvents; +export default useAssetEvents; diff --git a/airflow/www/static/js/api/useDatasets.ts b/airflow/www/static/js/api/useAssets.ts similarity index 90% rename from airflow/www/static/js/api/useDatasets.ts rename to airflow/www/static/js/api/useAssets.ts index db46415062c1a..3654c583c12ef 100644 --- a/airflow/www/static/js/api/useDatasets.ts +++ b/airflow/www/static/js/api/useAssets.ts @@ -28,7 +28,7 @@ interface Props { enabled?: boolean; } -export default function useDatasets({ dagIds, enabled = true }: Props) { +export default function useAssets({ dagIds, enabled = true }: Props) { return useQuery( ["datasets", dagIds], () => { @@ -36,7 +36,7 @@ export default function useDatasets({ dagIds, enabled = true }: Props) { const dagIdsParam = dagIds && dagIds.length ? { dag_ids: dagIds.join(",") } : {}; - return axios.get(datasetsUrl, { + return axios.get(datasetsUrl, { params: { ...dagIdsParam, }, diff --git a/airflow/www/static/js/api/useDatasetsSummary.ts b/airflow/www/static/js/api/useAssetsSummary.ts similarity index 98% rename from airflow/www/static/js/api/useDatasetsSummary.ts rename to airflow/www/static/js/api/useAssetsSummary.ts index 6f902946f6296..66b56ca9f6925 100644 --- a/airflow/www/static/js/api/useDatasetsSummary.ts +++ b/airflow/www/static/js/api/useAssetsSummary.ts @@ -42,7 +42,7 @@ interface Props { updatedAfter?: DateOption; } -export default function useDatasetsSummary({ +export default function useAssetsSummary({ limit, offset, order, diff --git a/airflow/www/static/js/api/useCreateDatasetEvent.ts b/airflow/www/static/js/api/useCreateAssetEvent.ts similarity index 77% rename from airflow/www/static/js/api/useCreateDatasetEvent.ts rename to airflow/www/static/js/api/useCreateAssetEvent.ts index f14b35ee375fe..7d2322c33ce9d 100644 --- a/airflow/www/static/js/api/useCreateDatasetEvent.ts +++ b/airflow/www/static/js/api/useCreateAssetEvent.ts @@ -29,22 +29,19 @@ interface Props { uri?: string; } -const createDatasetUrl = getMetaValue("create_dataset_event_api"); +const createAssetUrl = getMetaValue("create_asset_event_api"); -export default function useCreateDatasetEvent({ datasetId, uri }: Props) { +export default function useCreateAssetEvent({ datasetId, uri }: Props) { const queryClient = useQueryClient(); const errorToast = useErrorToast(); return useMutation( - ["createDatasetEvent", uri], - (extra?: API.DatasetEvent["extra"]) => - axios.post( - createDatasetUrl, - { - dataset_uri: uri, - extra: extra || {}, - } - ), + ["createAssetEvent", uri], + (extra?: API.AssetEvent["extra"]) => + axios.post(createAssetUrl, { + asset_uri: uri, + extra: extra || {}, + }), { onSuccess: () => { queryClient.invalidateQueries(["datasets-events", datasetId]); diff --git a/airflow/www/static/js/api/useUpstreamDatasetEvents.ts b/airflow/www/static/js/api/useUpstreamAssetEvents.ts similarity index 67% rename from airflow/www/static/js/api/useUpstreamDatasetEvents.ts rename to airflow/www/static/js/api/useUpstreamAssetEvents.ts index 32d1c7aeff2d8..437205501d6c5 100644 --- a/airflow/www/static/js/api/useUpstreamDatasetEvents.ts +++ b/airflow/www/static/js/api/useUpstreamAssetEvents.ts @@ -22,30 +22,30 @@ import { useQuery, UseQueryOptions } from "react-query"; import { getMetaValue } from "src/utils"; import type { - DatasetEventCollection, - GetUpstreamDatasetEventsVariables, + AssetEventCollection, + GetUpstreamAssetEventsVariables, } from "src/types/api-generated"; -interface Props extends GetUpstreamDatasetEventsVariables { - options?: UseQueryOptions; +interface Props extends GetUpstreamAssetEventsVariables { + options?: UseQueryOptions; } -const useUpstreamDatasetEvents = ({ dagId, dagRunId, options }: Props) => { +const useUpstreamAssetEvents = ({ dagId, dagRunId, options }: Props) => { const upstreamEventsUrl = ( - getMetaValue("upstream_dataset_events_api") || - `api/v1/dags/${dagId}/dagRuns/_DAG_RUN_ID_/upstreamDatasetEvents` + getMetaValue("upstream_asset_events_api") || + `api/v1/dags/${dagId}/dagRuns/_DAG_RUN_ID_/upstreamAssetEvents` ).replace("_DAG_RUN_ID_", encodeURIComponent(dagRunId)); - const query = useQuery( - ["upstreamDatasetEvents", dagRunId], + const query = useQuery( + ["upstreamAssetEvents", dagRunId], () => axios.get(upstreamEventsUrl), options ); return { ...query, - data: query.data ?? { datasetEvents: [], totalEntries: 0 }, + data: query.data ?? { assetEvents: [], totalEntries: 0 }, }; }; -export default useUpstreamDatasetEvents; +export default useUpstreamAssetEvents; diff --git a/airflow/www/static/js/components/DatasetEventCard.tsx b/airflow/www/static/js/components/DatasetEventCard.tsx index 2367c8efa9b4a..9dd1ee91e3731 100644 --- a/airflow/www/static/js/components/DatasetEventCard.tsx +++ b/airflow/www/static/js/components/DatasetEventCard.tsx @@ -21,7 +21,7 @@ import React from "react"; import { isEmpty } from "lodash"; import { TbApi } from "react-icons/tb"; -import type { DatasetEvent } from "src/types/api-generated"; +import type { AssetEvent } from "src/types/api-generated"; import { Box, Flex, @@ -43,7 +43,7 @@ import SourceTaskInstance from "./SourceTaskInstance"; import TriggeredDagRuns from "./TriggeredDagRuns"; type CardProps = { - datasetEvent: DatasetEvent; + assetEvent: AssetEvent; showSource?: boolean; showTriggeredDagRuns?: boolean; }; @@ -51,7 +51,7 @@ type CardProps = { const datasetsUrl = getMetaValue("datasets_url"); const DatasetEventCard = ({ - datasetEvent, + assetEvent, showSource = true, showTriggeredDagRuns = true, }: CardProps) => { @@ -60,14 +60,16 @@ const DatasetEventCard = ({ const selectedUri = decodeURIComponent(searchParams.get("uri") || ""); const containerRef = useContainerRef(); - const { from_rest_api: fromRestApi, ...extra } = - datasetEvent?.extra as Record; + const { from_rest_api: fromRestApi, ...extra } = assetEvent?.extra as Record< + string, + string + >; return ( - @@ -111,17 +112,17 @@ const DatasetEventCard = ({ )} - {!!datasetEvent.sourceTaskId && ( - + {!!assetEvent.sourceTaskId && ( + )} )} - {showTriggeredDagRuns && !!datasetEvent?.createdDagruns?.length && ( + {showTriggeredDagRuns && !!assetEvent?.createdDagruns?.length && ( <> Triggered Dag Runs: - + )} diff --git a/airflow/www/static/js/components/SourceTaskInstance.tsx b/airflow/www/static/js/components/SourceTaskInstance.tsx index 4c63198c5f40c..4343d3ce82443 100644 --- a/airflow/www/static/js/components/SourceTaskInstance.tsx +++ b/airflow/www/static/js/components/SourceTaskInstance.tsx @@ -22,7 +22,7 @@ import { Box, Link, Tooltip, Flex } from "@chakra-ui/react"; import { FiLink } from "react-icons/fi"; import { useTaskInstance } from "src/api"; -import type { DatasetEvent } from "src/types/api-generated"; +import type { AssetEvent } from "src/types/api-generated"; import { useContainerRef } from "src/context/containerRef"; import { SimpleStatus } from "src/dag/StatusBox"; import InstanceTooltip from "src/components/InstanceTooltip"; @@ -30,20 +30,16 @@ import type { TaskInstance } from "src/types"; import { getMetaValue } from "src/utils"; type SourceTIProps = { - datasetEvent: DatasetEvent; + assetEvent: AssetEvent; showLink?: boolean; }; const gridUrl = getMetaValue("grid_url"); const dagId = getMetaValue("dag_id") || "__DAG_ID__"; -const SourceTaskInstance = ({ - datasetEvent, - showLink = true, -}: SourceTIProps) => { +const SourceTaskInstance = ({ assetEvent, showLink = true }: SourceTIProps) => { const containerRef = useContainerRef(); - const { sourceDagId, sourceRunId, sourceTaskId, sourceMapIndex } = - datasetEvent; + const { sourceDagId, sourceRunId, sourceTaskId, sourceMapIndex } = assetEvent; const { data: taskInstance } = useTaskInstance({ dagId: sourceDagId || "", diff --git a/airflow/www/static/js/dag/details/dagRun/DatasetTriggerEvents.tsx b/airflow/www/static/js/dag/details/dagRun/DatasetTriggerEvents.tsx index 5fa585830b437..6deedb073e8d9 100644 --- a/airflow/www/static/js/dag/details/dagRun/DatasetTriggerEvents.tsx +++ b/airflow/www/static/js/dag/details/dagRun/DatasetTriggerEvents.tsx @@ -19,10 +19,10 @@ import React, { useMemo } from "react"; import { Box, Text } from "@chakra-ui/react"; -import { useUpstreamDatasetEvents } from "src/api"; +import { useUpstreamAssetEvents } from "src/api"; import type { DagRun as DagRunType } from "src/types"; import { CardDef, CardList } from "src/components/Table"; -import type { DatasetEvent } from "src/types/api-generated"; +import type { AssetEvent } from "src/types/api-generated"; import DatasetEventCard from "src/components/DatasetEventCard"; import { getMetaValue } from "src/utils"; @@ -32,17 +32,17 @@ interface Props { const dagId = getMetaValue("dag_id"); -const cardDef: CardDef = { +const cardDef: CardDef = { card: ({ row }) => ( - + ), }; const DatasetTriggerEvents = ({ runId }: Props) => { const { - data: { datasetEvents = [] }, + data: { assetEvents = [] }, isLoading, - } = useUpstreamDatasetEvents({ dagRunId: runId, dagId }); + } = useUpstreamAssetEvents({ dagRunId: runId, dagId }); const columns = useMemo( () => [ @@ -66,7 +66,7 @@ const DatasetTriggerEvents = ({ runId }: Props) => { [] ); - const data = useMemo(() => datasetEvents, [datasetEvents]); + const data = useMemo(() => assetEvents, [assetEvents]); return ( diff --git a/airflow/www/static/js/dag/details/graph/DatasetNode.tsx b/airflow/www/static/js/dag/details/graph/DatasetNode.tsx index bfd288f072dc3..d80f399b032a3 100644 --- a/airflow/www/static/js/dag/details/graph/DatasetNode.tsx +++ b/airflow/www/static/js/dag/details/graph/DatasetNode.tsx @@ -47,11 +47,11 @@ import type { CustomNodeProps } from "./Node"; const datasetsUrl = getMetaValue("datasets_url"); const DatasetNode = ({ - data: { label, height, width, latestDagRunId, isZoomedOut, datasetEvent }, + data: { label, height, width, latestDagRunId, isZoomedOut, assetEvent }, }: NodeProps) => { const containerRef = useContainerRef(); - const { from_rest_api: fromRestApi } = (datasetEvent?.extra || {}) as Record< + const { from_rest_api: fromRestApi } = (assetEvent?.extra || {}) as Record< string, string >; @@ -61,8 +61,8 @@ const DatasetNode = ({ Dataset - {!!datasetEvent && ( + {!!assetEvent && ( {/* @ts-ignore */} - {moment(datasetEvent.timestamp).fromNow()} + {moment(assetEvent.timestamp).fromNow()} )} @@ -120,23 +120,23 @@ const DatasetNode = ({ {label} - {!!datasetEvent && ( + {!!assetEvent && ( -