Skip to content

Commit 3b00dc8

Browse files
authored
Fix SNS FilterPolicy configuration; add kinesis/ListStreams API Gateway integration (localstack#1760)
1 parent 5d492fc commit 3b00dc8

7 files changed

Lines changed: 389 additions & 330 deletions

File tree

localstack/services/apigateway/apigateway_listener.py

Lines changed: 21 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -156,34 +156,38 @@ def invoke_rest_api(api_id, stage, method, invocation_path, data, headers, path=
156156
return make_error('Unable to find integration for path %s' % path, 404)
157157

158158
uri = integration.get('uri')
159-
if method == 'POST' and integration['type'] == 'AWS':
160-
if uri.endswith('kinesis:action/PutRecords'):
159+
if integration['type'] == 'AWS':
160+
if 'kinesis:action/' in uri:
161+
if uri.endswith('kinesis:action/PutRecords'):
162+
target = kinesis_listener.ACTION_PUT_RECORDS
163+
if uri.endswith('kinesis:action/ListStreams'):
164+
target = kinesis_listener.ACTION_LIST_STREAMS
165+
161166
template = integration['requestTemplates'][APPLICATION_JSON]
162167
new_request = aws_stack.render_velocity_template(template, data)
163-
164168
# forward records to target kinesis stream
165169
headers = aws_stack.mock_aws_request_headers(service='kinesis')
166-
headers['X-Amz-Target'] = kinesis_listener.ACTION_PUT_RECORDS
170+
headers['X-Amz-Target'] = target
167171
result = common.make_http_request(url=TEST_KINESIS_URL,
168172
method='POST', data=new_request, headers=headers)
169173
return result
170174

171-
elif uri.startswith('arn:aws:apigateway:') and ':sqs:path' in uri:
172-
template = integration['requestTemplates'][APPLICATION_JSON]
173-
account_id, queue = uri.split('/')[-2:]
174-
region_name = uri.split(':')[3]
175+
if method == 'POST':
176+
if uri.startswith('arn:aws:apigateway:') and ':sqs:path' in uri:
177+
template = integration['requestTemplates'][APPLICATION_JSON]
178+
account_id, queue = uri.split('/')[-2:]
179+
region_name = uri.split(':')[3]
175180

176-
new_request = aws_stack.render_velocity_template(template, data) + '&QueueName=%s' % queue
177-
headers = aws_stack.mock_aws_request_headers(service='sqs', region_name=region_name)
181+
new_request = aws_stack.render_velocity_template(template, data) + '&QueueName=%s' % queue
182+
headers = aws_stack.mock_aws_request_headers(service='sqs', region_name=region_name)
178183

179-
url = urljoin(TEST_SQS_URL, '%s/%s' % (account_id, queue))
180-
result = common.make_http_request(url, method='POST', headers=headers, data=new_request)
181-
return result
184+
url = urljoin(TEST_SQS_URL, '%s/%s' % (account_id, queue))
185+
result = common.make_http_request(url, method='POST', headers=headers, data=new_request)
186+
return result
182187

183-
else:
184-
msg = 'API Gateway action uri "%s" not yet implemented' % uri
185-
LOGGER.warning(msg)
186-
return make_error(msg, 404)
188+
msg = 'API Gateway AWS integration action URI "%s", method "%s" not yet implemented' % (uri, method)
189+
LOGGER.warning(msg)
190+
return make_error(msg, 404)
187191

188192
elif integration['type'] == 'AWS_PROXY':
189193
if uri.startswith('arn:aws:apigateway:') and ':lambda:path' in uri:

localstack/services/es/es_starter.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@ def start_elasticsearch(port=None, delete_data=True, asynchronous=False, update_
3030
es_tmp_dir = '%s/infra/elasticsearch/tmp' % (ROOT_PATH)
3131
es_mods_dir = '%s/infra/elasticsearch/modules' % (ROOT_PATH)
3232
if config.DATA_DIR:
33+
delete_data = False
3334
es_data_dir = '%s/elasticsearch' % config.DATA_DIR
3435
# Elasticsearch 5.x cannot be bound to 0.0.0.0 in some Docker environments,
3536
# hence we use the default bind address 127.0.0.0 and put a proxy in front of it

localstack/services/kinesis/kinesis_listener.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111
ACTION_PREFIX = 'Kinesis_20131202'
1212
ACTION_PUT_RECORD = '%s.PutRecord' % ACTION_PREFIX
1313
ACTION_PUT_RECORDS = '%s.PutRecords' % ACTION_PREFIX
14+
ACTION_LIST_STREAMS = '%s.ListStreams' % ACTION_PREFIX
1415
ACTION_CREATE_STREAM = '%s.CreateStream' % ACTION_PREFIX
1516
ACTION_DELETE_STREAM = '%s.DeleteStream' % ACTION_PREFIX
1617
ACTION_UPDATE_SHARD_COUNT = '%s.UpdateShardCount' % ACTION_PREFIX

localstack/services/sns/sns_listener.py

Lines changed: 10 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -140,14 +140,16 @@ def return_response(self, method, path, data, headers, response):
140140
if req_action == 'Subscribe' and response.status_code < 400:
141141
response_data = xmltodict.parse(response.content)
142142
topic_arn = (req_data.get('TargetArn') or req_data.get('TopicArn'))[0]
143+
filter_policy = (req_data.get('FilterPolicy') or [None])[0]
143144
attributes = get_subscribe_attributes(req_data)
144145
sub_arn = response_data['SubscribeResponse']['SubscribeResult']['SubscriptionArn']
145146
do_subscribe(
146147
topic_arn,
147148
req_data['Endpoint'][0],
148149
req_data['Protocol'][0],
149150
sub_arn,
150-
attributes
151+
attributes,
152+
filter_policy
151153
)
152154
if req_action == 'CreateTopic' and response.status_code < 400:
153155
response_data = xmltodict.parse(response.content)
@@ -173,7 +175,7 @@ def publish_message(topic_arn, req_data, subscription_arn=None):
173175
for subscriber in SNS_SUBSCRIPTIONS.get(topic_arn, []):
174176
if subscription_arn not in [None, subscriber['SubscriptionArn']]:
175177
continue
176-
filter_policy = json.loads(subscriber.get('FilterPolicy', '{}'))
178+
filter_policy = json.loads(subscriber.get('FilterPolicy') or '{}')
177179
message_attributes = get_message_attributes(req_data)
178180
if not check_filter_policy(filter_policy, message_attributes):
179181
continue
@@ -231,7 +233,7 @@ def do_delete_topic(topic_arn):
231233
SNS_SUBSCRIPTIONS.pop(topic_arn, None)
232234

233235

234-
def do_subscribe(topic_arn, endpoint, protocol, subscription_arn, attributes):
236+
def do_subscribe(topic_arn, endpoint, protocol, subscription_arn, attributes, filter_policy=None):
235237
# An endpoint may only be subscribed to a topic once. Subsequent
236238
# subscribe calls do nothing (subscribe is idempotent).
237239
for existing_topic_subscription in SNS_SUBSCRIPTIONS.get(topic_arn, []):
@@ -244,6 +246,7 @@ def do_subscribe(topic_arn, endpoint, protocol, subscription_arn, attributes):
244246
'Endpoint': endpoint,
245247
'Protocol': protocol,
246248
'SubscriptionArn': subscription_arn,
249+
'FilterPolicy': filter_policy
247250
}
248251
subscription.update(attributes)
249252
SNS_SUBSCRIPTIONS[topic_arn].append(subscription)
@@ -379,7 +382,7 @@ def create_sqs_message_attributes(subscriber, attributes):
379382
if value['Type'] == 'Binary':
380383
attribute['BinaryValue'] = value['Value']
381384
else:
382-
attribute['StringValue'] = value['Value']
385+
attribute['StringValue'] = str(value['Value'])
383386
message_attributes[key] = attribute
384387

385388
return message_attributes
@@ -400,6 +403,9 @@ def get_message_attributes(req_data):
400403
elif binary_value is not None:
401404
attribute['Value'] = binary_value
402405

406+
if attribute['Type'] == 'Number':
407+
attribute['Value'] = float(attribute['Value'])
408+
403409
attributes[name] = attribute
404410
x += 1
405411
else:

tests/integration/test_api_gateway.py

Lines changed: 17 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -79,9 +79,15 @@ def test_api_gateway_kinesis_integration(self):
7979
stage_name=self.TEST_STAGE_NAME,
8080
path=self.API_PATH_DATA_INBOUND
8181
)
82-
result = requests.post(url, data=json.dumps(test_data))
82+
83+
# list Kinesis streams via API Gateway
84+
result = requests.get(url)
8385
result = json.loads(to_str(result.content))
86+
self.assertIn('StreamNames', result)
8487

88+
# post test data to Kinesis via API Gateway
89+
result = requests.post(url, data=json.dumps(test_data))
90+
result = json.loads(to_str(result.content))
8591
self.assertEqual(result['FailedRecordCount'], 0)
8692
self.assertEqual(len(result['Records']), len(test_data['records']))
8793

@@ -312,6 +318,16 @@ def connect_api_gateway_to_kinesis(self, gateway_name, kinesis_stream):
312318
'application/json': template
313319
}
314320
}]
321+
}, {
322+
'httpMethod': 'GET',
323+
'authorizationType': 'NONE',
324+
'integrations': [{
325+
'type': 'AWS',
326+
'uri': 'arn:aws:apigateway:%s:kinesis:action/ListStreams' % DEFAULT_REGION,
327+
'requestTemplates': {
328+
'application/json': '{}'
329+
}
330+
}]
315331
}]
316332
return aws_stack.create_api_gateway(
317333
name=gateway_name,

tests/integration/test_sns.py

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010

1111
TEST_TOPIC_NAME = 'TestTopic_snsTest'
1212
TEST_QUEUE_NAME = 'TestQueue_snsTest'
13+
TEST_QUEUE_NAME_2 = 'TestQueue_snsTest2'
1314

1415

1516
class SNSTest(unittest.TestCase):
@@ -19,6 +20,7 @@ def setUp(self):
1920
self.sns_client = aws_stack.connect_to_service('sns')
2021
self.topic_arn = self.sns_client.create_topic(Name=TEST_TOPIC_NAME)['TopicArn']
2122
self.queue_url = self.sqs_client.create_queue(QueueName=TEST_QUEUE_NAME)['QueueUrl']
23+
self.queue_url_2 = self.sqs_client.create_queue(QueueName=TEST_QUEUE_NAME_2)['QueueUrl']
2224

2325
def tearDown(self):
2426
self.sqs_client.delete_queue(QueueUrl=self.queue_url)
@@ -89,6 +91,36 @@ def test_attribute_raw_subscribe(self):
8991
msg_received = msgs['Messages'][0]
9092
self.assertEqual(message, msg_received['Body'])
9193

94+
def test_filter_policy(self):
95+
# connect SNS topic to an SQS queue
96+
queue_arn = aws_stack.sqs_queue_arn(TEST_QUEUE_NAME_2)
97+
filter_policy = {'attr1': [{'numeric': ['>', 0, '<=', 100]}]}
98+
self.sns_client.subscribe(
99+
TopicArn=self.topic_arn,
100+
Protocol='sqs',
101+
Endpoint=queue_arn,
102+
Attributes={
103+
'FilterPolicy': json.dumps(filter_policy)
104+
}
105+
)
106+
107+
# get number of messages
108+
num_msgs_0 = len(self.sqs_client.receive_message(QueueUrl=self.queue_url_2).get('Messages', []))
109+
110+
# publish message that satisfies the filter policy, assert that message is received
111+
message = u'This is a test message'
112+
self.sns_client.publish(TopicArn=self.topic_arn, Message=message,
113+
MessageAttributes={'attr1': {'DataType': 'Number', 'StringValue': '99'}})
114+
num_msgs_1 = len(self.sqs_client.receive_message(QueueUrl=self.queue_url_2, VisibilityTimeout=0)['Messages'])
115+
self.assertEqual(num_msgs_1, num_msgs_0 + 1)
116+
117+
# publish message that does not satisfy the filter policy, assert that message is not received
118+
message = u'This is a test message'
119+
self.sns_client.publish(TopicArn=self.topic_arn, Message=message,
120+
MessageAttributes={'attr1': {'DataType': 'Number', 'StringValue': '111'}})
121+
num_msgs_2 = len(self.sqs_client.receive_message(QueueUrl=self.queue_url_2, VisibilityTimeout=0)['Messages'])
122+
self.assertEqual(num_msgs_2, num_msgs_1)
123+
92124
def test_unknown_topic_publish(self):
93125
fake_arn = 'arn:aws:sns:us-east-1:123456789012:i_dont_exist'
94126
message = u'This is a test message'

0 commit comments

Comments
 (0)