Skip to content
Draft
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 @@ -242,7 +242,7 @@ class KubernetesPodOperator(BaseOperator):
:param configmaps: (Optional) A list of names of config maps from which it collects ConfigMaps
to populate the environment variables with. The contents of the target
ConfigMap's Data field will represent the key-value pairs as environment variables.
Extends env_from.
Extends env_from. (templated)
:param skip_on_exit_code: If task exits with this exit code, leave the task
in ``skipped`` state (default: None). If set to ``None``, any non-zero
exit code will be treated as a failure.
Expand Down Expand Up @@ -313,6 +313,7 @@ class KubernetesPodOperator(BaseOperator):
"volume_mounts",
"cluster_context",
"env_from",
"configmaps",
"node_selector",
"kubernetes_conn_id",
"base_container_name",
Expand Down Expand Up @@ -421,20 +422,16 @@ def __init__(
self.startup_check_interval_seconds = startup_check_interval_seconds
# New parameter startup_timeout_seconds adds breaking change, to handle this as smooth as possible just reuse startup time
self.schedule_timeout_seconds = schedule_timeout_seconds or startup_timeout_seconds
env_vars = convert_env_vars(env_vars) if env_vars else []
self.env_vars = env_vars
self.env_vars = env_vars or []
pod_runtime_info_envs = (
[convert_pod_runtime_info_env(p) for p in pod_runtime_info_envs] if pod_runtime_info_envs else []
)
self.pod_runtime_info_envs = pod_runtime_info_envs
self.env_from = env_from or []
if configmaps:
self.env_from.extend([convert_configmap(c) for c in configmaps])
self.configmaps = configmaps or []
self.ports = [convert_port(p) for p in ports] if ports else []
volume_mounts = [convert_volume_mount(v) for v in volume_mounts] if volume_mounts else []
self.volume_mounts = volume_mounts
volumes = [convert_volume(volume) for volume in volumes] if volumes else []
self.volumes = volumes
self.volume_mounts = volume_mounts or []
self.volumes = volumes or []
self.secrets = secrets or []
self.in_cluster = in_cluster
self.cluster_context = cluster_context
Expand All @@ -449,7 +446,7 @@ def __init__(
self.base_container_name = base_container_name or self.BASE_CONTAINER_NAME
self.base_container_status_polling_interval = base_container_status_polling_interval
self.init_container_logs = init_container_logs
self.container_logs = container_logs or self.base_container_name
self._container_logs = container_logs
self.image_pull_policy = image_pull_policy
self.runtime_class_name = runtime_class_name
self.node_selector = node_selector or {}
Expand Down Expand Up @@ -512,6 +509,23 @@ def __init__(
self.container_name_log_prefix_enabled = container_name_log_prefix_enabled
self.log_formatter = log_formatter

@property
def container_logs(self) -> Iterable[str] | str | Literal[True]:
# Falls back lazily rather than in __init__: base_container_name is a template field, so
# the fallback has to read it once rendering has happened.
return self._container_logs or self.base_container_name

@container_logs.setter
def container_logs(self, value: Iterable[str] | str | Literal[True] | None) -> None:
self._container_logs = value

def render_template_fields(self, context: Context, jinja_env: jinja2.Environment | None = None) -> None:
# A str-str mapping has to become V1EnvVar objects before rendering: rendering a dict
# covers only its values, whereas an env var name is a template field of V1EnvVar.
if isinstance(self.env_vars, dict):
self.env_vars = convert_env_vars(self.env_vars)
super().render_template_fields(context, jinja_env)

@cached_property
def _incluster_namespace(self):
from pathlib import Path
Expand Down Expand Up @@ -1573,6 +1587,9 @@ def build_pod_request_obj(self, context: Context | None = None, *, dry_run: bool
self.env_vars = convert_env_vars_or_raise_error(self.env_vars) if self.env_vars else []
if self.pod_runtime_info_envs:
self.env_vars.extend(self.pod_runtime_info_envs)
env_from = [*self.env_from, *(convert_configmap(c) for c in self.configmaps)]
volume_mounts = [convert_volume_mount(v) for v in self.volume_mounts]
volumes = [convert_volume(volume) for volume in self.volumes]

if self.pod_template_file:
self.log.debug("Pod template file found, will parse for base pod")
Expand Down Expand Up @@ -1613,10 +1630,10 @@ def build_pod_request_obj(self, context: Context | None = None, *, dry_run: bool
ports=self.ports,
image_pull_policy=self.image_pull_policy,
resources=self.container_resources,
volume_mounts=self.volume_mounts,
volume_mounts=volume_mounts,
args=self.arguments,
env=self.env_vars,
env_from=self.env_from,
env_from=env_from,
security_context=self.container_security_context,
termination_message_policy=self.termination_message_policy,
)
Expand All @@ -1633,7 +1650,7 @@ def build_pod_request_obj(self, context: Context | None = None, *, dry_run: bool
scheduler_name=self.schedulername,
restart_policy="Never",
priority_class_name=self.priority_class_name,
volumes=self.volumes,
volumes=volumes,
active_deadline_seconds=self.active_deadline_seconds,
termination_grace_period_seconds=self.termination_grace_period,
),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -118,7 +118,6 @@ def __init__(

# fix mypy typing
self.base_container_name: str
self.container_logs: list[str]

if self.base_container_name != self.BASE_CONTAINER_NAME:
self.log.warning(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -270,7 +270,7 @@ def test_templates(self, create_task_instance_of_operator, session):
assert dag_id == rendered.arguments
assert dag_id == rendered.env_vars[0]
assert dag_id == rendered.annotations["dag-id"]
assert dag_id == rendered.env_from[0].config_map_ref.name
assert [dag_id] == rendered.configmaps
assert dag_id == rendered.volumes[0].name
assert dag_id == rendered.volumes[0].config_map.name

Expand Down Expand Up @@ -414,6 +414,54 @@ def test_envs_from_configmaps_backcompat(self):
pod = k.build_pod_request_obj(create_context(k))
assert pod.spec.containers[0].env_from == expected

def test_envs_from_templated_configmaps(self):
env_from = [k8s.V1EnvFromSource(config_map_ref=k8s.V1ConfigMapEnvSource(name="from-env-from"))]
k = KubernetesPodOperator(
task_id="task",
env_from=env_from,
configmaps="{{ maps }}",
dag=DAG(
dag_id="dag",
schedule=None,
start_date=pendulum.now(),
render_template_as_native_obj=True,
),
)
k.render_template_fields(context={"maps": ["from-configmaps"]})
pod = k.build_pod_request_obj(create_context(k))
assert pod.spec.containers[0].env_from == [
*env_from,
k8s.V1EnvFromSource(config_map_ref=k8s.V1ConfigMapEnvSource(name="from-configmaps")),
]

def test_templated_volumes_are_converted_after_rendering(self):
volume = k8s.V1Volume(name="vol", empty_dir=k8s.V1EmptyDirVolumeSource())
volume_mount = k8s.V1VolumeMount(name="vol", mount_path="/mnt")
k = KubernetesPodOperator(
task_id="task",
volumes="{{ vols }}",
volume_mounts="{{ mounts }}",
dag=DAG(
dag_id="dag",
schedule=None,
start_date=pendulum.now(),
render_template_as_native_obj=True,
),
)
k.render_template_fields(context={"vols": [volume], "mounts": [volume_mount]})
pod = k.build_pod_request_obj(create_context(k))
assert pod.spec.volumes == [volume]
assert pod.spec.containers[0].volume_mounts == [volume_mount]

def test_container_logs_falls_back_to_rendered_base_container_name(self):
k = KubernetesPodOperator(
task_id="task",
base_container_name="{{ container }}",
dag=DAG(dag_id="dag", schedule=None, start_date=pendulum.now()),
)
k.render_template_fields(context={"container": "rendered-base"})
assert k.container_logs == "rendered-base"

def test_envs_from_secrets(self):
secret_ref = "secret_name"
secrets = [Secret("env", None, secret_ref)]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -222,7 +222,14 @@ def test_spark_kubernetes_operator(mock_kubernetes_hook, data_file):
assert "hook" not in operator.__dict__ # Cached property has not been accessed as part of construction.


def test_init_spark_kubernetes_operator(data_file):
@pytest.mark.parametrize(
("container_logs", "expected_container_logs"),
[
pytest.param(None, "spark-kubernetes-driver", id="default"),
pytest.param(["sidecar"], ["spark-kubernetes-driver"], id="requested"),
],
)
def test_init_spark_kubernetes_operator(data_file, container_logs, expected_container_logs):
operator = SparkKubernetesOperator(
task_id="task_id",
application_file=data_file("spark/application_test.yaml").as_posix(),
Expand All @@ -231,10 +238,11 @@ def test_init_spark_kubernetes_operator(data_file):
cluster_context="cluster_context",
config_file="config_file",
base_container_name="base",
container_logs=container_logs,
get_logs=True,
)
assert operator.base_container_name == "spark-kubernetes-driver"
assert operator.container_logs == ["spark-kubernetes-driver"]
assert operator.container_logs == expected_container_logs


@patch("airflow.providers.cncf.kubernetes.operators.spark_kubernetes.KubernetesHook")
Expand Down
1 change: 0 additions & 1 deletion scripts/ci/prek/validate_operators_init_exemptions.txt
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,6 @@
providers/amazon/src/airflow/providers/amazon/aws/operators/neptune.py::NeptuneStartDbClusterOperator
providers/amazon/src/airflow/providers/amazon/aws/operators/neptune.py::NeptuneStopDbClusterOperator
providers/amazon/src/airflow/providers/amazon/aws/transfers/gcs_to_s3.py::GCSToS3Operator
providers/cncf/kubernetes/src/airflow/providers/cncf/kubernetes/operators/pod.py::KubernetesPodOperator
providers/google/src/airflow/providers/google/cloud/operators/cloud_batch.py::CloudBatchSubmitJobOperator
providers/google/src/airflow/providers/google/cloud/operators/cloud_build.py::CloudBuildCreateBuildOperator
providers/google/src/airflow/providers/google/cloud/operators/cloud_storage_transfer_service.py::CloudDataTransferServiceCreateJobOperator
Expand Down
Loading