diff --git a/scripts/aws.py b/scripts/aws.py index 0025e44e..0d10701c 100755 --- a/scripts/aws.py +++ b/scripts/aws.py @@ -10,16 +10,9 @@ logger = logging.getLogger(__name__) class AWSConnection: - def __init__(self, config): + def __init__(self, cluster_name): self.available = False - self.config = config - - if 'cluster_name' in config: - self.cluster_name = config.get('cluster_name') - elif 'etcd' in config and isinstance(config['etcd'], dict): - self.cluster_name = config['etcd'].get('scope', 'unknown') - else: - self.cluster_name = 'unknown' + self.cluster_name = cluster_name if cluster_name is not None else 'unknown' try: # get the instance id r = requests.get('http://169.254.169.254/latest/dynamic/instance-identity/document', timeout=0.1) @@ -73,12 +66,7 @@ class AWSConnection: if __name__ == '__main__': - if len(sys.argv) != 4: - print ("Usage: {0} action role name".format(sys.argv[0])) - sys.exit(1) - action, role, name = sys.argv[1:] - if action in ('on_start', 'on_stop', 'on_role_change'): - aws = AWSConnection({'cluster_name': name}) - aws.on_role_change(role) - sys.exit(0) - sys.exit(2) + if len(sys.argv) != 4 and sys.argv[1] in ('on_start', 'on_stop', 'on_role_change'): + AWSConnection(cluster_name=sys.argv[3]).on_role_change(sys.argv[2]) + else: + sys.exit("Usage: {0} action role name".format(sys.argv[0])) diff --git a/tests/test_aws.py b/tests/test_aws.py index 4331badb..84d495fe 100644 --- a/tests/test_aws.py +++ b/tests/test_aws.py @@ -4,7 +4,6 @@ import boto.ec2 from collections import namedtuple from scripts.aws import AWSConnection from requests.exceptions import RequestException -import yaml class MockEc2Connection: @@ -42,8 +41,8 @@ class TestAWSConnection(unittest.TestCase): def set_error(self): self.error = True - def set_ok(self): - self.error = False + def set_json_error(self): + self.json_error = True def boto_ec2_connect_to_region(self, region): return MockEc2Connection(self.error) @@ -53,7 +52,7 @@ class TestAWSConnection(unittest.TestCase): raise RequestException("foo") result = namedtuple('Request', 'ok content') result.ok = True - if url.split('/')[-1] == 'document': + if url.split('/')[-1] == 'document' and not self.json_error: result = {"instanceId": "012345", "region": "eu-west-1"} else: result = 'foo' @@ -61,50 +60,10 @@ class TestAWSConnection(unittest.TestCase): def setUp(self): self.error = False + self.json_error = False requests.get = self.requests_get boto.ec2.connect_to_region = self.boto_ec2_connect_to_region - self.config_string = """ -scope: &scope test -ttl: &ttl 30 -loop_wait: &loop_wait 10 -restapi: - listen: 0.0.0.0:8008 - connect_address: 127.0.0.1:5432 -etcd: - scope: *scope - ttl: *ttl - host: 127.0.0.1:8080 -postgresql: - scope: *scope - name: postgresql_foo - listen: 0.0.0.0:5432 - connect_address: 127.0.0.1:5432 - data_dir: /home/postgres/pgdata/data - replication: - username: standby - password: standby - network: 0.0.0.0/0 - superuser: - password: zalando - admin: - username: admin - password: admin - callbacks: - on_start: patroni/scripts/aws.py - on_stop: patroni/scripts/aws.py - on_restart: patroni/scripts/aws.py - on_role_change: patroni/scripts/aws.py - parameters: - archive_mode: "on" - wal_level: hot_standby - max_wal_senders: 5 - wal_keep_segments: 8 - archive_timeout: 1800s - max_replication_slots: 5 - hot_standby: "on" - ssl: "on" -""" - self.conn = AWSConnection(yaml.load(self.config_string)) + self.conn = AWSConnection('test') def test_aws_available(self): self.assertTrue(self.conn.aws_available()) @@ -116,11 +75,16 @@ postgresql: def test_non_aws(self): self.set_error() - conn = AWSConnection(yaml.load(self.config_string)) + conn = AWSConnection('test') self.assertFalse(conn.aws_available()) self.assertFalse(conn._tag_ebs('master')) self.assertFalse(conn._tag_ec2('master')) + def test_aws_bizare_response(self): + self.set_json_error() + conn = AWSConnection('test') + self.assertFalse(conn.aws_available()) + def test_aws_tag_ebs_error(self): self.set_error() self.assertFalse(self.conn._tag_ebs("master"))