import argparse import asyncio import os import sys from enum import Enum import aiohttp from atlas_client import AtlasClient from tqdm.auto import tqdm BASE_DIR = os.path.dirname(os.path.abspath(__file__)) ENDPOINT = "/api/v1/email/send" with open(os.path.join(BASE_DIR, "email_body.html"), "r") as f: BODY_TEXT = f.read() TEMPLATE_ID = 1 TEMPLATE_CONTEXT = {"test_var": "test_var_value"} class SESTestEmails(Enum): success = "success@simulator.amazonses.com" bounce = "bounce@simulator.amazonses.com" otoo = "ooto@simulator.amazonses.com" complaint = "complaint@simulator.amazonses.com" suppression = "suppressionlist@simulator.amazonses.com" class Loadtest: def __init__( self, env, atlas_env, client_id, client_secret, audience, max_requests, qty_batches, qty_per_batch, use_template, **kwargs, ): self.env = env self.atlas_env = atlas_env self.client_id = client_id self.client_secret = client_secret self.audience = audience self.max_requests = max_requests self.qty_batches = qty_batches self.qty_per_batch = qty_per_batch self.use_template = use_template self.atlas = AtlasClient( self.atlas_env, self.client_id, self.client_secret, self.audience ) self.session = aiohttp.ClientSession(f"https://{self.env}") self.sem = asyncio.Semaphore(self.max_requests) self.progress = tqdm(total=self.qty_batches, desc="Loadtesting...") async def run(self): try: await self.send_batches() finally: await self.session.close() async def send_email(self, to, subject, body): try: # it's ok to block event loop here to refresh the token # it will only block when the token is expired token = self.atlas.get_token() except Exception as e: print(e) sys.exit(1) async with self.sem: if self.use_template: payload = [ { "subject": subject, "to": to, "template": { "id": TEMPLATE_ID, "context": TEMPLATE_CONTEXT, }, } ] else: payload = [{"subject": subject, "to": to, "body": body}] async with self.session.post( ENDPOINT, json=payload, headers={"Authorization": f"Bearer {token}"}, ) as resp: if resp.status != 200: print(f"Error: {await resp.text()}") self.progress.update() self.progress.refresh() async def send_batches(self): tasks = [] for n in range(self.qty_batches): tasks.append( asyncio.create_task( self.send_email( [SESTestEmails.success.value] * self.qty_per_batch, "test email", BODY_TEXT, ) ) ) await asyncio.wait(tasks) async def main(): parser = argparse.ArgumentParser( description="Load test for core-notifications" ) parser.add_argument( "--max_requests", type=int, dest="max_requests", default=100, help="Max requests at the same time (default 100)", ) parser.add_argument( "--qty_batches", type=int, dest="qty_batches", default=1, help="Total quantity of requests to core-notifications (default 1)", ) parser.add_argument( "--qty_per_batch", type=int, dest="qty_per_batch", default=1, help="Quantity of emails per request (default 1)", ) parser.add_argument( "--use_template", action="store_true", default=False, help="Use template for emails", ) parser.add_argument( "--env", type=str, dest="env", default="dev-notifications-api.atlas.stream", help="Env (default dev-notifications-api.atlas.stream)", ) parser.add_argument( "--atlas_env", type=str, dest="atlas_env", default="dev-um.atlas.stream", help="Atlas env for auth (default dev-um.atlas.stream)", ) parser.add_argument( "--client_id", type=str, dest="client_id", help="M2M client ID", required=True, ) parser.add_argument( "--client_secret", type=str, dest="client_secret", help="M2M client secret", required=True, ) parser.add_argument( "--audience", type=str, dest="audience", default="atlasum|core_notifications", help="M2M audience (default 'atlasum|core_notifications')", ) args = parser.parse_args() await Loadtest(**vars(args)).run() if __name__ == "__main__": asyncio.run(main())