diff --git a/providers/weaviate/src/airflow/providers/weaviate/operators/weaviate.py b/providers/weaviate/src/airflow/providers/weaviate/operators/weaviate.py index de080c7232036..a52317393f078 100644 --- a/providers/weaviate/src/airflow/providers/weaviate/operators/weaviate.py +++ b/providers/weaviate/src/airflow/providers/weaviate/operators/weaviate.py @@ -59,7 +59,7 @@ def __init__( self, conn_id: str, collection_name: str, - input_data: list[dict[str, Any]] | pd.DataFrame | None = None, + input_data: list[dict[str, Any]] | pd.DataFrame, vector_col: str = "Vector", uuid_column: str = "id", tenant: str | None = None, @@ -75,15 +75,14 @@ def __init__( self.input_data = input_data self.hook_params = hook_params or {} - if self.input_data is None: - raise TypeError("input_data is required") - @cached_property def hook(self) -> WeaviateHook: """Return an instance of the WeaviateHook.""" return WeaviateHook(conn_id=self.conn_id, **self.hook_params) def execute(self, context: Context) -> None: + if self.input_data is None: + raise TypeError("input_data is required") self.log.debug("Input data: %s", self.input_data) self.hook.batch_data( collection_name=self.collection_name, diff --git a/providers/weaviate/tests/unit/weaviate/operators/test_weaviate.py b/providers/weaviate/tests/unit/weaviate/operators/test_weaviate.py index 0f09fb35d528b..dbb94265f81ae 100644 --- a/providers/weaviate/tests/unit/weaviate/operators/test_weaviate.py +++ b/providers/weaviate/tests/unit/weaviate/operators/test_weaviate.py @@ -81,6 +81,21 @@ def test_execute_passes_tenant_to_hook(self): tenant="tenant-a", ) + def test_missing_input_data_raises_at_execute_not_init(self): + """ + input_data is a template field, so the required-value check runs in execute() + (after rendering), not in __init__. Constructing with input_data=None must not raise. + """ + operator = WeaviateIngestOperator( + task_id="weaviate_task", + conn_id="weaviate_conn", + collection_name="my_collection", + input_data=None, + ) + + with pytest.raises(TypeError, match="input_data is required"): + operator.execute(context=None) + @pytest.mark.db_test def test_templates(self, create_task_instance_of_operator): dag_id = "TestWeaviateIngestOperator" diff --git a/scripts/ci/prek/validate_operators_init_exemptions.txt b/scripts/ci/prek/validate_operators_init_exemptions.txt index 71a13973e4c4a..b78123db0dc66 100644 --- a/scripts/ci/prek/validate_operators_init_exemptions.txt +++ b/scripts/ci/prek/validate_operators_init_exemptions.txt @@ -69,4 +69,3 @@ providers/ssh/src/airflow/providers/ssh/operators/ssh_remote_job.py::SSHRemoteJo providers/standard/src/airflow/providers/standard/operators/bash.py::BashOperator providers/standard/src/airflow/providers/standard/operators/trigger_dagrun.py::TriggerDagRunOperator providers/standard/src/airflow/providers/standard/sensors/date_time.py::DateTimeSensor -providers/weaviate/src/airflow/providers/weaviate/operators/weaviate.py::WeaviateIngestOperator