import time import unittest from fargate_tools.cli.utils import poll_task_status from fargate_tools.cli.constants import VerifyMode class FakePaginator: def __init__(self, pages): self.pages = pages def paginate(self, **kwargs): for page in self.pages: yield page class FakeECSClient: def __init__(self, task_sequences): self._task_sequences = task_sequences self._describe_calls = 0 self._list_tasks_pages = [] def get_paginator(self, name): assert name == 'list_tasks' return FakePaginator(self._list_tasks_pages) def set_list_tasks_pages(self, pages): self._list_tasks_pages = pages def describe_tasks(self, cluster, tasks): self._describe_calls += 1 described = [] for arn in tasks: seq = self._task_sequences.get(arn, []) idx = min(self._describe_calls - 1, len(seq) - 1) if seq: described.append(seq[idx]) return {"tasks": described} class UtilsPollingTests(unittest.TestCase): def test_timeout_before_task_appears(self): client = FakeECSClient(task_sequences={}) client.set_list_tasks_pages([{"taskArns": []}]) original_sleep = time.sleep try: time.sleep = lambda _s: None result = poll_task_status( client=client, verify_mode=VerifyMode.TASK_RUNNING, update_timeout=1, cluster_name="cluster", service_name="service", task_definition_arn="arn:td:1", container_name="container", grace_period=1, ) finally: time.sleep = original_sleep self.assertFalse(result) def test_task_running_success(self): task_arn = "arn:task:123" client = FakeECSClient(task_sequences={ task_arn: [ {"taskArn": task_arn, "taskDefinitionArn": "arn:td:1", "lastStatus": "PENDING"}, {"taskArn": task_arn, "taskDefinitionArn": "arn:td:1", "lastStatus": "RUNNING"}, ] }) client.set_list_tasks_pages([{"taskArns": [task_arn]}]) result = poll_task_status( client=client, verify_mode=VerifyMode.TASK_RUNNING, update_timeout=5, cluster_name="cluster", service_name="service", task_definition_arn="arn:td:1", container_name="container", grace_period=1, ) self.assertTrue(result) def test_health_check_unhealthy_after_grace(self): task_arn = "arn:task:999" client = FakeECSClient(task_sequences={ task_arn: [ {"taskArn": task_arn, "taskDefinitionArn": "arn:td:2", "lastStatus": "RUNNING", "healthStatus": "UNKNOWN"}, {"taskArn": task_arn, "taskDefinitionArn": "arn:td:2", "lastStatus": "RUNNING", "healthStatus": "UNHEALTHY"}, ] }) client.set_list_tasks_pages([{"taskArns": [task_arn]}]) result = poll_task_status( client=client, verify_mode=VerifyMode.HEALTH_CHECK, update_timeout=3, cluster_name="cluster", service_name="service", task_definition_arn="arn:td:2", container_name="container", grace_period=0, ) self.assertFalse(result) def test_exit_code_failure(self): task_arn = "arn:task:55" client = FakeECSClient(task_sequences={ task_arn: [ {"taskArn": task_arn, "taskDefinitionArn": "arn:td:3", "lastStatus": "PENDING"}, {"taskArn": task_arn, "taskDefinitionArn": "arn:td:3", "lastStatus": "STOPPED", "stopCode": "EssentialContainerExited", "containers": [{"name": "container", "exitCode": 2}]}, ] }) client.set_list_tasks_pages([{"taskArns": [task_arn]}]) result = poll_task_status( client=client, verify_mode=VerifyMode.EXIT_CODE, update_timeout=5, cluster_name="cluster", service_name="service", task_definition_arn="arn:td:3", container_name="container", grace_period=1, ) self.assertFalse(result) if __name__ == "__main__": unittest.main()