from typing import Any, Mapping from airflow.providers.snowflake.operators.snowflake import SnowflakeOperator __all__ = ["XCOMSnowflakeOperator"] class XCOMSnowflakeOperator(SnowflakeOperator): """ Extension for SnowflakeOperator to pass parameters from xcom """ def __init__(self, xcom_parameters: dict[str, dict[str, Any]] | None = None, **kwargs): super().__init__(**kwargs) self.xcom_parameters = xcom_parameters def _get_xcom_parameters(self, context: Any): parameters = {} if self.xcom_parameters is not None: for name, lookup in self.xcom_parameters.items(): parameters[name] = self.xcom_pull(context, **lookup) return parameters def _get_parameters(self, context: Any) -> Mapping | None: parameters = self.parameters if self.parameters is not None else {} xcom_parameters = self._get_xcom_parameters(context) return dict(parameters, **xcom_parameters) def pre_execute(self, context: Any): self.parameters = self._get_parameters(context)