|
4 | 4 | import shutil |
5 | 5 | from pathlib import Path |
6 | 6 | from unittest.case import TestCase |
| 7 | +from urllib.parse import urlparse |
7 | 8 |
|
8 | 9 | import boto3 |
9 | 10 | import botocore |
10 | 11 | import pytest |
11 | 12 | import requests |
| 13 | +from botocore.auth import SigV4Auth |
| 14 | +from botocore.awsrequest import AWSRequest |
12 | 15 | from samtranslator.translator.arn_generator import ArnGenerator |
13 | 16 | from samtranslator.yaml_helper import yaml_parse |
14 | 17 | from tenacity import ( |
@@ -538,6 +541,25 @@ def verify_get_request_response(self, url, expected_status_code, headers=None): |
538 | 541 | ) |
539 | 542 | return response |
540 | 543 |
|
| 544 | + @retry( |
| 545 | + stop=stop_after_attempt(6), |
| 546 | + wait=wait_exponential(multiplier=1, min=16, max=64) + wait_random(0, 1), |
| 547 | + retry=retry_if_exception_type(StatusCodeError), |
| 548 | + after=after_log(LOG, logging.WARNING), |
| 549 | + reraise=True, |
| 550 | + ) |
| 551 | + def verify_get_request_response_sigv4(self, url, expected_status_code, headers=None): |
| 552 | + """ |
| 553 | + Verify if a SigV4-signed get request to a certain url returns the expected status code. |
| 554 | + Use this for APIs with IAM authorization. |
| 555 | + """ |
| 556 | + response = self.do_get_request_with_sigv4(url, headers) |
| 557 | + if response.status_code != expected_status_code: |
| 558 | + raise StatusCodeError( |
| 559 | + f"SigV4 request to {url} failed with status: {response.status_code}, expected status: {expected_status_code}" |
| 560 | + ) |
| 561 | + return response |
| 562 | + |
541 | 563 | @retry( |
542 | 564 | stop=stop_after_attempt(6), |
543 | 565 | wait=wait_exponential(multiplier=1, min=16, max=64) + wait_random(0, 1), |
@@ -581,6 +603,22 @@ def verify_post_request(self, url: str, body_obj, expected_status_code: int, hea |
581 | 603 | ) |
582 | 604 | return response |
583 | 605 |
|
| 606 | + @retry( |
| 607 | + stop=stop_after_attempt(6), |
| 608 | + wait=wait_exponential(multiplier=1, min=16, max=64) + wait_random(0, 1), |
| 609 | + retry=retry_if_exception_type(StatusCodeError), |
| 610 | + after=after_log(LOG, logging.WARNING), |
| 611 | + reraise=True, |
| 612 | + ) |
| 613 | + def verify_post_request_sigv4(self, url: str, body_obj, expected_status_code: int, headers=None): |
| 614 | + """Return response to SigV4-signed POST request and verify matches expected status code.""" |
| 615 | + response = self.do_post_request_with_sigv4(url, body_obj, headers) |
| 616 | + if response.status_code != expected_status_code: |
| 617 | + raise StatusCodeError( |
| 618 | + f"SigV4 POST request to {url} failed with status: {response.status_code}, expected status: {expected_status_code}" |
| 619 | + ) |
| 620 | + return response |
| 621 | + |
584 | 622 | def get_default_test_template_parameters(self): |
585 | 623 | """ |
586 | 624 | get the default template parameters |
@@ -636,6 +674,29 @@ def do_get_request_with_logging(self, url, headers=None): |
636 | 674 | ) |
637 | 675 | return response |
638 | 676 |
|
| 677 | + def do_get_request_with_sigv4(self, url, headers=None): |
| 678 | + """ |
| 679 | + Perform a SigV4-signed GET request to an APIGW endpoint with IAM auth. |
| 680 | + """ |
| 681 | + parsed = urlparse(url) |
| 682 | + request_headers = {"host": parsed.hostname} |
| 683 | + if headers: |
| 684 | + request_headers.update(headers) |
| 685 | + |
| 686 | + aws_request = AWSRequest(method="GET", url=url, headers=request_headers) |
| 687 | + session = botocore.session.Session() |
| 688 | + credentials = session.get_credentials().get_frozen_credentials() |
| 689 | + SigV4Auth(credentials, "execute-api", self.my_region).add_auth(aws_request) |
| 690 | + |
| 691 | + response = requests.get(url, headers=dict(aws_request.headers)) |
| 692 | + amazon_headers = RequestUtils(response).get_amazon_headers() |
| 693 | + if self.internal: |
| 694 | + REQUEST_LOGGER.info( |
| 695 | + "SigV4 request made to " + url, |
| 696 | + extra={"test": self.testcase, "status": response.status_code, "headers": amazon_headers}, |
| 697 | + ) |
| 698 | + return response |
| 699 | + |
639 | 700 | def do_options_request_with_logging(self, url, headers=None): |
640 | 701 | """ |
641 | 702 | Perform a options request to an APIGW endpoint and log relevant info |
@@ -669,3 +730,25 @@ def do_post_request_with_logging(self, url: str, body_obj, requestHeaders=None): |
669 | 730 | extra={"test": self.testcase, "status": response.status_code, "headers": amazon_headers}, |
670 | 731 | ) |
671 | 732 | return response |
| 733 | + |
| 734 | + def do_post_request_with_sigv4(self, url: str, body_obj, headers=None): |
| 735 | + """Perform a SigV4-signed POST request to an APIGW endpoint with IAM auth.""" |
| 736 | + parsed = urlparse(url) |
| 737 | + body = json.dumps(body_obj) |
| 738 | + request_headers = {"host": parsed.hostname, "content-type": "application/json"} |
| 739 | + if headers: |
| 740 | + request_headers.update(headers) |
| 741 | + |
| 742 | + aws_request = AWSRequest(method="POST", url=url, headers=request_headers, data=body) |
| 743 | + session = botocore.session.Session() |
| 744 | + credentials = session.get_credentials().get_frozen_credentials() |
| 745 | + SigV4Auth(credentials, "execute-api", self.my_region).add_auth(aws_request) |
| 746 | + |
| 747 | + response = requests.post(url, data=body, headers=dict(aws_request.headers)) |
| 748 | + amazon_headers = RequestUtils(response).get_amazon_headers() |
| 749 | + if self.internal: |
| 750 | + REQUEST_LOGGER.info( |
| 751 | + "SigV4 POST request made to " + url, |
| 752 | + extra={"test": self.testcase, "status": response.status_code, "headers": amazon_headers}, |
| 753 | + ) |
| 754 | + return response |
0 commit comments