Skip to content

Commit 692c2c2

Browse files
committed
Improve Argo workflow label tests
1 parent cf6814f commit 692c2c2

2 files changed

Lines changed: 105 additions & 238 deletions

File tree

Lines changed: 56 additions & 238 deletions
Original file line numberDiff line numberDiff line change
@@ -1,246 +1,64 @@
11
import pytest
2-
from unittest.mock import patch, MagicMock, PropertyMock
32

3+
from metaflow.plugins.argo.argo_workflows import ArgoWorkflows
44
from metaflow.plugins.kubernetes.kube_utils import KubernetesException
55

66

7-
class TestArgoWorkflowsLabels:
8-
"""Test _base_argo_labels() method with mocked configuration."""
7+
@pytest.fixture
8+
def argo_workflows():
9+
return ArgoWorkflows.__new__(ArgoWorkflows)
910

10-
def _call_base_argo_labels(self):
11-
"""
12-
Call _base_argo_labels without full ArgoWorkflows instantiation.
1311

14-
The methods only use module-level config, not instance state, so we can
15-
call them on a minimal object with the methods bound to it.
16-
"""
17-
from metaflow.plugins.argo.argo_workflows import ArgoWorkflows
18-
19-
# Create a minimal object with both methods bound
20-
class MinimalArgo:
21-
_base_kubernetes_labels = ArgoWorkflows._base_kubernetes_labels
22-
_custom_argo_labels = ArgoWorkflows._custom_argo_labels
23-
_base_argo_labels = ArgoWorkflows._base_argo_labels
24-
25-
obj = MinimalArgo()
26-
return obj._base_argo_labels()
27-
28-
def test_default_labels_when_env_not_set(self):
29-
"""Should return default labels when ARGO_WORKFLOWS_LABELS is empty."""
30-
with patch("metaflow.plugins.argo.argo_workflows.ARGO_WORKFLOWS_LABELS", ""):
31-
labels = self._call_base_argo_labels()
32-
33-
assert labels == {"app.kubernetes.io/part-of": "metaflow"}
34-
35-
def test_adds_custom_labels_from_env(self):
36-
"""Should add custom labels from ARGO_WORKFLOWS_LABELS."""
37-
with patch(
38-
"metaflow.plugins.argo.argo_workflows.ARGO_WORKFLOWS_LABELS",
12+
@pytest.mark.parametrize(
13+
("configured_labels", "expected"),
14+
[
15+
(
16+
"",
17+
{"app.kubernetes.io/part-of": "metaflow"},
18+
),
19+
(
3920
"team=ml,env=prod",
40-
):
41-
labels = self._call_base_argo_labels()
42-
43-
assert labels == {
44-
"app.kubernetes.io/part-of": "metaflow",
45-
"team": "ml",
46-
"env": "prod",
47-
}
48-
49-
def test_custom_labels_do_not_override_internal_labels(self):
50-
"""Custom labels must not override internal/base labels."""
51-
with patch(
52-
"metaflow.plugins.argo.argo_workflows.ARGO_WORKFLOWS_LABELS",
53-
"app.kubernetes.io/part-of=custom-app",
54-
):
55-
labels = self._call_base_argo_labels()
56-
57-
assert labels == {"app.kubernetes.io/part-of": "metaflow"}
58-
59-
def test_custom_labels_alongside_part_of_override_attempt(self):
60-
"""Unrelated custom labels pass through even when part-of override is attempted."""
61-
with patch(
62-
"metaflow.plugins.argo.argo_workflows.ARGO_WORKFLOWS_LABELS",
63-
"app.kubernetes.io/part-of=custom-app,team=ml",
64-
):
65-
labels = self._call_base_argo_labels()
66-
67-
assert labels == {
68-
"app.kubernetes.io/part-of": "metaflow",
69-
"team": "ml",
70-
}
71-
72-
def test_single_label(self):
73-
"""Should handle a single label correctly."""
74-
with patch(
75-
"metaflow.plugins.argo.argo_workflows.ARGO_WORKFLOWS_LABELS",
76-
"cost-center=12345",
77-
):
78-
labels = self._call_base_argo_labels()
79-
80-
assert labels == {
81-
"app.kubernetes.io/part-of": "metaflow",
82-
"cost-center": "12345",
83-
}
84-
85-
def test_invalid_label_value_raises_exception(self):
86-
"""Should raise exception for invalid label values."""
87-
with patch(
88-
"metaflow.plugins.argo.argo_workflows.ARGO_WORKFLOWS_LABELS",
89-
"team=invalid value with spaces",
90-
):
91-
with pytest.raises(KubernetesException):
92-
self._call_base_argo_labels()
93-
94-
def test_label_value_too_long_raises_exception(self):
95-
"""Should raise exception for label values exceeding 63 chars."""
96-
long_value = "a" * 64
97-
with patch(
98-
"metaflow.plugins.argo.argo_workflows.ARGO_WORKFLOWS_LABELS",
99-
f"team={long_value}",
100-
):
101-
with pytest.raises(KubernetesException):
102-
self._call_base_argo_labels()
103-
104-
def test_invalid_label_key_empty_name_raises_exception(self):
105-
"""Should raise exception for a key with an empty name segment (e.g. '=ml')."""
106-
with patch(
107-
"metaflow.plugins.argo.argo_workflows.ARGO_WORKFLOWS_LABELS",
108-
"=ml",
109-
):
110-
with pytest.raises(KubernetesException):
111-
self._call_base_argo_labels()
112-
113-
def test_invalid_label_key_with_spaces_raises_exception(self):
114-
"""Should raise exception for a key containing spaces."""
115-
with patch(
116-
"metaflow.plugins.argo.argo_workflows.ARGO_WORKFLOWS_LABELS",
117-
" team=ml",
118-
):
119-
with pytest.raises(KubernetesException):
120-
self._call_base_argo_labels()
121-
122-
def test_invalid_label_key_too_long_raises_exception(self):
123-
"""Should raise exception for a key name segment exceeding 63 chars."""
124-
long_key = "a" * 64
125-
with patch(
126-
"metaflow.plugins.argo.argo_workflows.ARGO_WORKFLOWS_LABELS",
127-
f"{long_key}=ml",
128-
):
129-
with pytest.raises(KubernetesException):
130-
self._call_base_argo_labels()
131-
132-
133-
class TestArgoWorkflowsTemplateLabels:
134-
"""Test that labels appear in compiled WorkflowTemplate and Sensor."""
135-
136-
@pytest.fixture
137-
def mock_argo_workflows(self):
138-
"""Create an ArgoWorkflows instance with mocked dependencies."""
139-
# Mock all the complex dependencies
140-
patches = [
141-
patch(
142-
"metaflow.plugins.argo.argo_workflows.ARGO_WORKFLOWS_LABELS",
143-
"team=ml-platform,env=production",
144-
),
145-
patch(
146-
"metaflow.plugins.argo.argo_workflows.KUBERNETES_NAMESPACE", "test-ns"
147-
),
148-
patch("metaflow.plugins.argo.argo_workflows.ARGO_EVENTS_EVENT", None),
149-
patch(
150-
"metaflow.plugins.argo.argo_workflows.ARGO_EVENTS_EVENT_SOURCE", None
151-
),
152-
patch(
153-
"metaflow.plugins.argo.argo_workflows.ARGO_EVENTS_SERVICE_ACCOUNT", None
154-
),
155-
]
156-
157-
for p in patches:
158-
p.start()
159-
160-
from metaflow.plugins.argo.argo_workflows import ArgoWorkflows
161-
162-
# Create mock graph with minimal structure
163-
mock_node = MagicMock()
164-
mock_node.name = "start"
165-
mock_node.type = "linear"
166-
mock_node.out_funcs = ["end"]
167-
mock_node.is_inside_foreach = False
168-
mock_node.parallel_foreach = False
169-
170-
mock_end_node = MagicMock()
171-
mock_end_node.name = "end"
172-
mock_end_node.type = "end"
173-
mock_end_node.out_funcs = []
174-
mock_end_node.is_inside_foreach = False
175-
mock_end_node.parallel_foreach = False
176-
177-
mock_graph = MagicMock()
178-
mock_graph.nodes = {"start": mock_node, "end": mock_end_node}
179-
mock_graph.__iter__ = lambda self: iter([mock_node, mock_end_node])
180-
181-
# Create mock flow
182-
mock_flow = MagicMock()
183-
mock_flow.name = "TestFlow"
184-
mock_flow._flow_decorators = {}
185-
type(mock_flow)._parameters = PropertyMock(return_value={})
186-
type(mock_flow)._configs = PropertyMock(return_value={})
187-
188-
# Create mock environment
189-
mock_environment = MagicMock()
190-
mock_environment.get_package_commands.return_value = []
191-
mock_environment.bootstrap_commands.return_value = []
192-
193-
# Create mock datastore
194-
mock_datastore = MagicMock()
195-
mock_datastore.TYPE = "s3"
196-
197-
# Create the instance - we'll patch the complex compilation methods
198-
with patch.object(ArgoWorkflows, "_compile_workflow_template"), patch.object(
199-
ArgoWorkflows, "_compile_sensor"
200-
), patch.object(
201-
ArgoWorkflows, "_process_parameters", return_value=[]
202-
), patch.object(
203-
ArgoWorkflows, "_process_config_parameters", return_value=[]
204-
), patch.object(
205-
ArgoWorkflows, "_process_triggers", return_value=([], {})
206-
), patch.object(
207-
ArgoWorkflows, "_get_schedule", return_value=(None, None)
208-
), patch.object(
209-
ArgoWorkflows, "_parse_conditional_branches"
210-
):
211-
argo = ArgoWorkflows(
212-
name="test-flow",
213-
graph=mock_graph,
214-
flow=mock_flow,
215-
code_package_metadata={},
216-
code_package_sha="abc123",
217-
code_package_url="s3://bucket/code.tar.gz",
218-
production_token="prod-token",
219-
metadata=MagicMock(),
220-
flow_datastore=mock_datastore,
221-
environment=mock_environment,
222-
event_logger=MagicMock(),
223-
monitor=MagicMock(),
224-
username="testuser",
225-
)
226-
227-
for p in patches:
228-
p.stop()
229-
230-
return argo
231-
232-
def test_workflow_labels_includes_custom_labels(self, mock_argo_workflows):
233-
"""Verify _workflow_labels (used at WorkflowTemplate/Workflow level) contains custom labels from env var."""
234-
labels = mock_argo_workflows._workflow_labels
235-
236-
assert labels["app.kubernetes.io/part-of"] == "metaflow"
237-
assert labels["team"] == "ml-platform"
238-
assert labels["env"] == "production"
239-
240-
def test_base_labels_excludes_custom_labels(self, mock_argo_workflows):
241-
"""Verify _base_labels (used for pods/JobSet/Sensor) does NOT contain custom labels from env var."""
242-
labels = mock_argo_workflows._base_labels
243-
244-
assert labels == {"app.kubernetes.io/part-of": "metaflow"}
245-
assert "team" not in labels
246-
assert "env" not in labels
21+
{
22+
"app.kubernetes.io/part-of": "metaflow",
23+
"team": "ml",
24+
"env": "prod",
25+
},
26+
),
27+
(
28+
"app.kubernetes.io/part-of=custom,team=ml",
29+
{
30+
"app.kubernetes.io/part-of": "metaflow",
31+
"team": "ml",
32+
},
33+
),
34+
],
35+
ids=["default", "custom-labels", "protected-label"],
36+
)
37+
def test_base_argo_labels(mocker, argo_workflows, configured_labels, expected):
38+
mocker.patch(
39+
"metaflow.plugins.argo.argo_workflows.ARGO_WORKFLOWS_LABELS",
40+
configured_labels,
41+
)
42+
43+
assert argo_workflows._base_argo_labels() == expected
44+
45+
46+
@pytest.mark.parametrize(
47+
"configured_labels",
48+
[
49+
"missing-value",
50+
"team=value with spaces",
51+
"team=%s" % ("a" * 64),
52+
],
53+
ids=["missing-equals", "invalid-value", "value-too-long"],
54+
)
55+
def test_base_argo_labels_rejects_invalid_configuration(
56+
mocker, argo_workflows, configured_labels
57+
):
58+
mocker.patch(
59+
"metaflow.plugins.argo.argo_workflows.ARGO_WORKFLOWS_LABELS",
60+
configured_labels,
61+
)
62+
63+
with pytest.raises(KubernetesException):
64+
argo_workflows._base_argo_labels()

‎test/ux/core/test_argo_compilation.py‎

Lines changed: 49 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -64,6 +64,55 @@ def test_argo_only_json_exposes_workflow_template(
6464
assert workflow_template["spec"]["templates"]
6565

6666

67+
def test_configured_labels_are_emitted_on_argo_workflows(
68+
exec_mode, decospecs, compute_env, tag, scheduler_config
69+
):
70+
if exec_mode != "deployer":
71+
pytest.skip("Argo compilation tests require deployer mode")
72+
if scheduler_config.scheduler_type != "argo-workflows":
73+
pytest.skip("Argo compilation tests require the argo-workflows scheduler")
74+
75+
from metaflow import Deployer
76+
77+
from .test_utils import _resolve_flow_path, prepare_runner_deployer_args
78+
79+
env = dict(compute_env)
80+
env["METAFLOW_ARGO_WORKFLOWS_LABELS"] = (
81+
"team=ml-platform,environment=test,app.kubernetes.io/name=custom"
82+
)
83+
84+
deployed_flow = (
85+
Deployer(
86+
flow_file=_resolve_flow_path("basic/helloworld.py"),
87+
show_output=False,
88+
**prepare_runner_deployer_args({"decospecs": decospecs, "env": env}),
89+
)
90+
.argo_workflows()
91+
.create(
92+
only_json=True,
93+
tags=tag + ["test_configured_argo_labels"],
94+
**(scheduler_config.deploy_args or {}),
95+
)
96+
)
97+
98+
workflow_template = deployed_flow.workflow_template
99+
template_labels = workflow_template["metadata"]["labels"]
100+
workflow_labels = workflow_template["spec"]["workflowMetadata"]["labels"]
101+
102+
assert template_labels["team"] == "ml-platform"
103+
assert template_labels["environment"] == "test"
104+
assert template_labels["app.kubernetes.io/name"] == "metaflow-flow"
105+
106+
assert workflow_labels["team"] == "ml-platform"
107+
assert workflow_labels["environment"] == "test"
108+
assert workflow_labels["app.kubernetes.io/name"] == "metaflow-run"
109+
110+
# Custom Argo labels are intentionally workflow-level only.
111+
pod_labels = workflow_template["spec"]["podMetadata"]["labels"]
112+
assert "team" not in pod_labels
113+
assert "environment" not in pod_labels
114+
115+
67116
def test_foreach_split_switch_join_task_names_are_deduplicated(
68117
exec_mode, decospecs, tag, scheduler_config
69118
):

0 commit comments

Comments
 (0)