Skip to content

Commit 0c670c5

Browse files
committed
Fix docker scheduler state aggregation
1 parent 416a0a7 commit 0c670c5

2 files changed

Lines changed: 39 additions & 23 deletions

File tree

nemo_run/run/torchx_backend/schedulers/docker.py

Lines changed: 5 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -189,12 +189,11 @@ def describe(self, app_id: str) -> Optional[DescribeAppResponse]:
189189
states.append(state)
190190

191191
state = AppState.UNKNOWN
192-
if any(is_terminal(state) for state in states):
193-
if any(state == AppState.SUCCEEDED for state in states):
194-
state = AppState.SUCCEEDED
195-
else:
196-
state = AppState.FAILED
197-
elif len(states) > 0:
192+
if any(state == AppState.FAILED for state in states):
193+
state = AppState.FAILED
194+
elif len(states) == len(req.containers) and all(state == AppState.SUCCEEDED for state in states):
195+
state = AppState.SUCCEEDED
196+
elif any(not is_terminal(state) for state in states):
198197
state = next(state for state in states if not is_terminal(state))
199198

200199
return DescribeAppResponse(

test/run/torchx_backend/schedulers/test_docker.py

Lines changed: 34 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -187,37 +187,54 @@ def test_describe_failed(docker_scheduler, docker_executor):
187187
assert len(response.roles) == 1
188188

189189

190-
@pytest.mark.xfail
191-
def test_describe_failure_not_detected(docker_scheduler, docker_executor):
190+
@pytest.mark.parametrize(
191+
("container_states", "expected_state"),
192+
[
193+
([AppState.SUCCEEDED, AppState.FAILED], AppState.FAILED),
194+
([AppState.SUCCEEDED, AppState.RUNNING], AppState.RUNNING),
195+
([AppState.SUCCEEDED, AppState.SUCCEEDED], AppState.SUCCEEDED),
196+
],
197+
)
198+
def test_describe_aggregates_container_states(
199+
docker_scheduler, docker_executor, container_states, expected_state
200+
):
192201
with (
193202
mock.patch.object(DockerJobRequest, "load") as mock_load,
194203
mock.patch.object(DockerContainer, "get_container") as mock_get_container,
195204
mock.patch.object(PersistentDockerScheduler, "_get_app_state") as mock_get_app_state,
205+
mock.patch.object(
206+
PersistentDockerScheduler, "_docker_client", new_callable=mock.PropertyMock
207+
) as mock_docker_client,
196208
):
197-
container = DockerContainer(
198-
name="test_role",
199-
command=["test"],
200-
executor=docker_executor,
201-
extra_env={},
202-
)
209+
mock_docker_client.return_value = mock.Mock()
210+
containers = [
211+
DockerContainer(
212+
name="test_role",
213+
command=["test"],
214+
executor=docker_executor,
215+
extra_env={},
216+
),
217+
DockerContainer(
218+
name="test_role_2",
219+
command=["test"],
220+
executor=docker_executor,
221+
extra_env={},
222+
),
223+
]
203224
req = DockerJobRequest(
204225
id="test_session___test_role___test_container_id",
205226
executor=docker_executor,
206-
containers=[container],
227+
containers=containers,
207228
)
208229
mock_load.return_value = req
209-
mock_get_container.return_value = container
210-
mock_get_app_state.return_value = None
211-
status_file = os.path.join(req.executor.job_dir, f"status_{req.containers[0].name}.out")
212-
213-
with open(status_file, "w") as f:
214-
f.write(json.dumps({"exit_code": 1}))
230+
mock_get_container.side_effect = containers
231+
mock_get_app_state.side_effect = container_states
215232

216233
response = docker_scheduler.describe(req.id)
217234
assert response is not None
218235
assert response.app_id == req.id
219-
assert "SUCCEEDED" in str(response.state)
220-
assert len(response.roles) == 1
236+
assert response.state == expected_state
237+
assert len(response.roles) == 2
221238

222239

223240
def test_save_and_get_job_dirs():

0 commit comments

Comments
 (0)