|
1 | 1 | import pytest |
2 | | -from unittest.mock import patch, MagicMock, PropertyMock |
3 | 2 |
|
| 3 | +from metaflow.plugins.argo.argo_workflows import ArgoWorkflows |
4 | 4 | from metaflow.plugins.kubernetes.kube_utils import KubernetesException |
5 | 5 |
|
6 | 6 |
|
7 | | -class TestArgoWorkflowsLabels: |
8 | | - """Test _base_argo_labels() method with mocked configuration.""" |
| 7 | +@pytest.fixture |
| 8 | +def argo_workflows(): |
| 9 | + return ArgoWorkflows.__new__(ArgoWorkflows) |
9 | 10 |
|
10 | | - def _call_base_argo_labels(self): |
11 | | - """ |
12 | | - Call _base_argo_labels without full ArgoWorkflows instantiation. |
13 | 11 |
|
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 | + ( |
39 | 20 | "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() |
0 commit comments