import unittest from unittest.mock import MagicMock, patch from get_number_of_ecs_tasks import ( get_clusters, get_services, get_list_of_tasks, get_max_count, extract_cluster_name_from_arn, extract_service_name_from_arn, ) class TestGetNumberOfECSTasks(unittest.TestCase): def setUp(self): self.client = MagicMock() def test_get_clusters(self): paginator = MagicMock() paginator.paginate.return_value = [ {"clusterArns": ["cluster1", "cluster2"]}, {"clusterArns": ["cluster3"]}, ] self.client.get_paginator.return_value = paginator clusters = get_clusters(self.client) self.assertEqual(clusters, ["cluster1", "cluster2", "cluster3"]) self.client.get_paginator.assert_called_once_with("list_clusters") paginator.paginate.assert_called_once_with() def test_get_services(self): paginator = MagicMock() paginator.paginate.return_value = [ {"serviceArns": ["service1", "service2"]}, {"serviceArns": ["service3"]}, ] self.client.get_paginator.return_value = paginator services = get_services(self.client, "cluster") self.assertEqual(services, ["service1", "service2", "service3"]) self.client.get_paginator.assert_called_once_with("list_services") paginator.paginate.assert_called_once_with( cluster="cluster", launchType="FARGATE" ) def test_get_list_of_tasks(self): paginator = MagicMock() paginator.paginate.return_value = [ {"taskArns": ["task1", "task2"]}, {"taskArns": ["task3"]}, ] self.client.get_paginator.return_value = paginator tasks = get_list_of_tasks(self.client, "cluster", "service") self.assertEqual(tasks, ["task1", "task2", "task3"]) self.client.get_paginator.assert_called_once_with("list_tasks") paginator.paginate.assert_called_once_with( cluster="cluster", serviceName="service", launchType="FARGATE", desiredStatus="RUNNING" ) def test_get_max_count(self): self.client.describe_scalable_targets.return_value = { "ScalableTargets": [ {"MaxCapacity": 10}, ] } max_count = get_max_count(self.client, "cluster", "service") self.assertEqual(max_count, 10) self.client.describe_scalable_targets.assert_called_once_with( ResourceIds=["service/cluster/service"], ScalableDimension="ecs:service:DesiredCount", ServiceNamespace="ecs", ) def test_get_max_count_no_scalable_targets(self): self.client.describe_scalable_targets.return_value = { "ScalableTargets": [], } max_count = get_max_count(self.client, "cluster", "service") self.assertEqual(max_count, 0) self.client.describe_scalable_targets.assert_called_once_with( ResourceIds=["service/cluster/service"], ScalableDimension="ecs:service:DesiredCount", ServiceNamespace="ecs", ) def test_extract_cluster_name_from_arn(self): cluster_name = extract_cluster_name_from_arn("arn:aws:ecs:us-east-1:123456789012:cluster/cluster-name") self.assertEqual(cluster_name, "cluster-name") def test_extract_service_name_from_arn(self): service_name = extract_service_name_from_arn("arn:aws:ecs:us-east-1:123456789012:service/cluster-name/service-name") self.assertEqual(service_name, "service-name") if __name__ == "__main__": unittest.main()