Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -63,11 +63,9 @@ def __init__(
**kwargs: Any,
):
super().__init__(poke_interval=poke_interval, **kwargs)
if not partition:
partition = "ds='{{ ds }}'"
self.metastore_conn_id = metastore_conn_id
self.table = table
self.partition = partition
self.partition = partition or "ds='{{ ds }}'"
self.schema = schema

def poke(self, context: Context) -> bool:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -56,9 +56,6 @@ def __init__(
super().__init__(poke_interval=poke_interval, **kwargs)

self.next_index_to_poke = 0
if isinstance(partition_names, str):
raise TypeError("partition_names must be an array of strings")

self.metastore_conn_id = metastore_conn_id
self.partition_names = partition_names
self.hook = hook
Expand Down Expand Up @@ -95,6 +92,8 @@ def poke_partition(self, partition: str) -> Any:
return self.hook.check_for_named_partition(schema, table, partition)

def poke(self, context: Context) -> bool:
if isinstance(self.partition_names, str):
raise TypeError("partition_names must be an array of strings")
Comment thread
shahar1 marked this conversation as resolved.
number_of_partitions = len(self.partition_names)
poke_index_start = self.next_index_to_poke
for i in range(number_of_partitions):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -114,6 +114,38 @@ def test_poke_non_existing(self):
self.database, self.table, f"{self.partition_by}={self.next_day}"
)

def test_native_templated_partition_names(self):
self.hook.metastore.__enter__().check_for_named_partition.return_value = True
partitions = [f"{self.database}.{self.table}/{self.partition_by}={DEFAULT_DATE_DS}"]
dag = DAG(
"named_hive_native",
schedule=None,
start_date=DEFAULT_DATE,
render_template_as_native_obj=True,
)
sensor = NamedHivePartitionSensor(
partition_names="{{ params.names }}",
task_id="test_native_templated_partition_names",
poke_interval=1,
hook=self.hook,
params={"names": partitions},
dag=dag,
)
sensor.render_template_fields({"params": {"names": partitions}})
assert sensor.partition_names == partitions
assert sensor.poke(None)

def test_poke_rejects_str_partition_names(self):
"""partition_names is validated at poke now, so a rendered str raises there, not in __init__."""
sensor = NamedHivePartitionSensor(
partition_names="not-a-list",
task_id="test_poke_rejects_str_partition_names",
poke_interval=1,
hook=self.hook,
)
with pytest.raises(TypeError, match="partition_names must be an array of strings"):
sensor.poke(None)


@pytest.mark.skipif(
"AIRFLOW_RUNALL_TESTS" not in os.environ, reason="Skipped because AIRFLOW_RUNALL_TESTS is not set"
Expand Down
2 changes: 0 additions & 2 deletions scripts/ci/prek/validate_operators_init_exemptions.txt
Original file line number Diff line number Diff line change
Expand Up @@ -21,8 +21,6 @@ providers/amazon/src/airflow/providers/amazon/aws/transfers/base.py::AwsToAwsBas
providers/amazon/src/airflow/providers/amazon/aws/transfers/gcs_to_s3.py::GCSToS3Operator
providers/amazon/src/airflow/providers/amazon/aws/transfers/s3_to_redshift.py::S3ToRedshiftOperator
providers/anthropic/src/airflow/providers/anthropic/operators/agent.py::AnthropicAgentSessionOperator
providers/apache/hive/src/airflow/providers/apache/hive/sensors/hive_partition.py::HivePartitionSensor
providers/apache/hive/src/airflow/providers/apache/hive/sensors/named_hive_partition.py::NamedHivePartitionSensor
providers/apache/kafka/src/airflow/providers/apache/kafka/operators/produce.py::ProduceToTopicOperator
providers/cncf/kubernetes/src/airflow/providers/cncf/kubernetes/operators/kueue.py::KubernetesInstallKueueOperator
providers/cncf/kubernetes/src/airflow/providers/cncf/kubernetes/operators/pod.py::KubernetesPodOperator
Expand Down