diff --git a/README.md b/README.md index 3e5de6e..9e4013c 100644 --- a/README.md +++ b/README.md @@ -85,6 +85,11 @@ aws ssm put-parameter \ --name "/banking-app/dev/SesSenderEmail" \ --value "sender@yourdomain.co.uk" \ --type "String" + +aws ssm put-parameter \ + --name "/banking-app/dev/SesNoReplyEmail" \ + --value "no-reply@yourdomain.co.uk" \ + --type "String" aws ssm put-parameter \ --name "/banking-app/dev/SesReplyEmail" \ diff --git a/dev-requirements.txt b/dev-requirements.txt index 0888e83..9f60e9a 100644 --- a/dev-requirements.txt +++ b/dev-requirements.txt @@ -5,8 +5,10 @@ black==25.1.0 ruff==0.11.9 # Running -aws_lambda_powertools==3.12.0 +aws_lambda_powertools==3.17.0 boto3==1.38.13 +xhtml2pdf==0.2.17 +Jinja2==3.1.6 # Tests pytest==8.3.5 diff --git a/layers/python/__init__.py b/functions/accounts/get_account_transactions/__init__.py similarity index 100% rename from layers/python/__init__.py rename to functions/accounts/get_account_transactions/__init__.py diff --git a/functions/accounts/get_account_transactions/get_account_transactions/__init__.py b/functions/accounts/get_account_transactions/get_account_transactions/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/functions/accounts/get_account_transactions/get_account_transactions/app.py b/functions/accounts/get_account_transactions/get_account_transactions/app.py new file mode 100644 index 0000000..31c87b9 --- /dev/null +++ b/functions/accounts/get_account_transactions/get_account_transactions/app.py @@ -0,0 +1,99 @@ +import os +import json + +from aws_lambda_powertools import Logger +from aws_lambda_powertools.event_handler import ( + APIGatewayRestResolver, + CORSConfig, +) +from aws_lambda_powertools.event_handler.exceptions import ( + InternalServerError, + BadRequestError, +) +from aws_lambda_powertools.utilities.typing import LambdaContext + +from dynamodb import get_dynamodb_resource +from .exceptions import ValidationError +from .transaction_helpers import query_transactions + +TRANSACTIONS_TABLE_NAME = os.environ.get("TRANSACTIONS_TABLE_NAME") +ENVIRONMENT_NAME = os.environ.get("ENVIRONMENT_NAME", "dev") +POWERTOOLS_LOG_LEVEL = os.environ.get("POWERTOOLS_LOG_LEVEL", "INFO").upper() +DYNAMODB_ENDPOINT = os.environ.get("DYNAMODB_ENDPOINT") +AWS_REGION = os.environ.get("AWS_REGION", "eu-west-2") + +logger = Logger(service="GetAccountTransactions", level=POWERTOOLS_LOG_LEVEL) + +app = APIGatewayRestResolver( + cors=CORSConfig(allow_headers=["Content-Type", "Authorization"]) +) + +dynamodb = get_dynamodb_resource(DYNAMODB_ENDPOINT, AWS_REGION, logger) +if TRANSACTIONS_TABLE_NAME: + table = dynamodb.Table(TRANSACTIONS_TABLE_NAME) + logger.debug(f"Initialized DynamoDB table: {TRANSACTIONS_TABLE_NAME}") +else: + logger.critical("FATAL: TRANSACTIONS_TABLE_NAME environment variable not set!") + table = None + + +@app.get("/accounts//transactions") +def get_account_transactions(account_id: str): + try: + period = app.current_event.get_query_string_value("period", default_value=None) + start = app.current_event.get_query_string_value("start", default_value=None) + end = app.current_event.get_query_string_value("end", default_value=None) + + result = query_transactions( + table=table, + account_id=account_id, + logger=logger, + period=period, + start=start, + end=end, + ) + return result + + except ValidationError as ve: + logger.warning(f"Validation error: {ve}") + raise BadRequestError(str(ve)) + except Exception as e: + logger.error(f"Error fetching transactions: {e}", exc_info=True) + raise InternalServerError("Internal server error") + + +@logger.inject_lambda_context +def lambda_handler(event, context: LambdaContext): + logger.append_keys(request_id=context.aws_request_id) + logger.info(f"Processing request in {ENVIRONMENT_NAME}") + + if not table: + logger.error("DynamoDB table resource is not initialized") + raise InternalServerError("Server configuration error") + + # Detect Step Functions or API Gateway + if "httpMethod" in event or "requestContext" in event: + return app.resolve(event, context) + else: + account_id = event.get("accountId") + if not account_id: + return { + "statusCode": 400, + "body": json.dumps({"error": "Missing accountId"}), + } + + try: + result = query_transactions( + table=table, account_id=account_id, logger=logger + ) + + response = { + **event, + "transactions": result.get("transactions", result), + } + + return response + + except Exception as e: + logger.error(f"Error fetching transactions: {e}", exc_info=True) + return {"statusCode": 500, "body": json.dumps({"error": str(e)})} diff --git a/functions/accounts/get_account_transactions/get_account_transactions/date_helpers.py b/functions/accounts/get_account_transactions/get_account_transactions/date_helpers.py new file mode 100644 index 0000000..4ce5905 --- /dev/null +++ b/functions/accounts/get_account_transactions/get_account_transactions/date_helpers.py @@ -0,0 +1,68 @@ +from datetime import datetime, timedelta, timezone +import calendar + +from .exceptions import ValidationError + + +def get_date_range(period: str = None, start: str = None, end: str = None): + # --- Validation rules --- + if period and (start or end): + raise ValidationError("Cannot combine 'period' with 'start'/'end'") + + if (start and not end) or (end and not start): + raise ValidationError("Both 'start' and 'end' must be provided together") + + # --- Custom range --- + if start and end: + try: + start_dt = datetime.strptime(start, "%Y-%m-%d").replace(tzinfo=timezone.utc) + end_dt = datetime.strptime(end, "%Y-%m-%d").replace( + hour=23, minute=59, second=59, tzinfo=timezone.utc + ) + except ValueError: + raise ValidationError("Invalid date format, must be YYYY-MM-DD") + + if end_dt < start_dt: + raise ValidationError("'end' date must be after 'start' date") + + statement_period = ( + f"{start_dt.strftime('%Y-%m-%d')}_to_{end_dt.strftime('%Y-%m-%d')}" + ) + + # --- Period (month) --- + elif period: + try: + year, month = map(int, period.split("-")) + start_dt = datetime(year, month, 1, tzinfo=timezone.utc) + last_day_num = calendar.monthrange(year, month)[1] + end_dt = datetime( + year, month, last_day_num, 23, 59, 59, tzinfo=timezone.utc + ) + except Exception: + raise ValidationError("Invalid period format, must be YYYY-MM") + + statement_period = start_dt.strftime("%Y-%m") + + # --- Default: last month --- + else: + today = datetime.now(timezone.utc) + first_day_this_month = datetime(today.year, today.month, 1, tzinfo=timezone.utc) + last_day_last_month = first_day_this_month - timedelta(days=1) + start_dt = datetime( + last_day_last_month.year, last_day_last_month.month, 1, tzinfo=timezone.utc + ) + end_dt = datetime( + last_day_last_month.year, + last_day_last_month.month, + last_day_last_month.day, + 23, + 59, + 59, + tzinfo=timezone.utc, + ) + statement_period = start_dt.strftime("%Y-%m") + + start_iso = start_dt.strftime("%Y-%m-%dT%H:%M:%SZ") + end_iso = end_dt.strftime("%Y-%m-%dT%H:%M:%SZ") + + return statement_period, start_iso, end_iso diff --git a/functions/accounts/get_account_transactions/get_account_transactions/exceptions.py b/functions/accounts/get_account_transactions/get_account_transactions/exceptions.py new file mode 100644 index 0000000..15e676c --- /dev/null +++ b/functions/accounts/get_account_transactions/get_account_transactions/exceptions.py @@ -0,0 +1,2 @@ +class ValidationError(Exception): + pass diff --git a/functions/accounts/get_account_transactions/get_account_transactions/transaction_helpers.py b/functions/accounts/get_account_transactions/get_account_transactions/transaction_helpers.py new file mode 100644 index 0000000..ee5f873 --- /dev/null +++ b/functions/accounts/get_account_transactions/get_account_transactions/transaction_helpers.py @@ -0,0 +1,35 @@ +from aws_lambda_powertools import Logger +from boto3.dynamodb.conditions import Key + +from . import date_helpers + + +def query_transactions( + table, + account_id: str, + logger: Logger, + period: str = None, + start: str = None, + end: str = None, + descending=False, +): + statement_period, start_iso, end_iso = date_helpers.get_date_range( + period, start, end + ) + + logger.info( + f"Querying transactions for account {account_id} " + f"from {start_iso} to {end_iso} (period {statement_period})" + ) + + response = table.query( + IndexName="AccountDateIndex", + KeyConditionExpression=Key("accountId").eq(account_id) + & Key("createdAt").between(start_iso, end_iso), + ScanIndexForward=not descending, + ) + + return { + "statementPeriod": statement_period, + "transactions": response.get("Items", []), + } diff --git a/functions/accounts/get_account_transactions/requirements.txt b/functions/accounts/get_account_transactions/requirements.txt new file mode 100644 index 0000000..216efe3 --- /dev/null +++ b/functions/accounts/get_account_transactions/requirements.txt @@ -0,0 +1,2 @@ +aws_lambda_powertools==3.17.0 +boto3==1.38.13 \ No newline at end of file diff --git a/functions/accounts/get_accounts/requirements.txt b/functions/accounts/get_accounts/requirements.txt index 8637f8f..216efe3 100644 --- a/functions/accounts/get_accounts/requirements.txt +++ b/functions/accounts/get_accounts/requirements.txt @@ -1,2 +1,2 @@ -aws_lambda_powertools==3.12.0 +aws_lambda_powertools==3.17.0 boto3==1.38.13 \ No newline at end of file diff --git a/functions/auth/requirements.txt b/functions/auth/requirements.txt index 8637f8f..216efe3 100644 --- a/functions/auth/requirements.txt +++ b/functions/auth/requirements.txt @@ -1,2 +1,2 @@ -aws_lambda_powertools==3.12.0 +aws_lambda_powertools==3.17.0 boto3==1.38.13 \ No newline at end of file diff --git a/functions/cognito/post_sign_up/requirements.txt b/functions/cognito/post_sign_up/requirements.txt index 8637f8f..216efe3 100644 --- a/functions/cognito/post_sign_up/requirements.txt +++ b/functions/cognito/post_sign_up/requirements.txt @@ -1,2 +1,2 @@ -aws_lambda_powertools==3.12.0 +aws_lambda_powertools==3.17.0 boto3==1.38.13 \ No newline at end of file diff --git a/functions/monthly_reports/__init__.py b/functions/monthly_reports/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/functions/monthly_reports/accounts/__init__.py b/functions/monthly_reports/accounts/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/functions/monthly_reports/accounts/create_report/__init__.py b/functions/monthly_reports/accounts/create_report/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/functions/monthly_reports/accounts/create_report/create_report/__init__.py b/functions/monthly_reports/accounts/create_report/create_report/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/functions/monthly_reports/accounts/create_report/create_report/app.py b/functions/monthly_reports/accounts/create_report/create_report/app.py new file mode 100644 index 0000000..e8f53ab --- /dev/null +++ b/functions/monthly_reports/accounts/create_report/create_report/app.py @@ -0,0 +1,78 @@ +import os + +from aws_lambda_powertools import Logger +from aws_lambda_powertools.utilities.typing import LambdaContext +from botocore.exceptions import ClientError + +from .exceptions import ReportGenerationError, ReportTemplateError, ReportUploadError +from s3 import get_s3_client +from .generate_pdf import generate_transactions_pdf + +REPORTS_BUCKET = os.environ.get("REPORTS_BUCKET") +POWERTOOLS_LOG_LEVEL = os.environ.get("POWERTOOLS_LOG_LEVEL") +AWS_REGION = os.environ.get("AWS_REGION") + +logger = Logger(service="CreateAccountsReport", level=POWERTOOLS_LOG_LEVEL) + +s3 = get_s3_client(AWS_REGION, logger) + + +def lambda_handler(event, _context: LambdaContext): + logger.info(f"Received event: {event}") + + try: + required = [ + "accountId", + "userId", + "statementPeriod", + "transactions", + "accountBalance", + ] + missing = [k for k in required if k not in event] + + if missing: + logger.error(f"Missing required fields: {missing}") + raise ReportGenerationError(f"Invalid event: missing {missing}") + + # Generate PDF + pdf_bytes = generate_transactions_pdf(event=event, logger=logger) + + logger.info("PDF generated successfully") + + # Store in S3 + s3_key = f"{event['accountId']}/{event['statementPeriod']}.pdf" + try: + s3.put_object( + Bucket=REPORTS_BUCKET, + Key=s3_key, + Body=pdf_bytes, + ContentType="application/pdf", + ) + except ClientError as e: + logger.exception("Failed to upload report to S3") + raise ReportUploadError(f"S3 upload failed: {str(e)}") from e + + logger.info("Generating PDF uploaded to S3") + + try: + presigned_url = s3.generate_presigned_url( + "get_object", + Params={"Bucket": REPORTS_BUCKET, "Key": s3_key}, + ExpiresIn=3600, + ) + except ClientError as e: + logger.exception("Failed to generate presigned URL") + raise ReportUploadError(f"Presigned URL generation failed: {str(e)}") from e + + logger.info("Presigned URL generated successfully") + + return { + "reportUrl": presigned_url, + "accountId": event["accountId"], + "userId": event["userId"], + "statementPeriod": event["statementPeriod"], + } + + except (ReportGenerationError, ReportTemplateError, ReportUploadError): + logger.exception("Report generation failed") + raise diff --git a/functions/monthly_reports/accounts/create_report/create_report/exceptions.py b/functions/monthly_reports/accounts/create_report/create_report/exceptions.py new file mode 100644 index 0000000..d00df24 --- /dev/null +++ b/functions/monthly_reports/accounts/create_report/create_report/exceptions.py @@ -0,0 +1,10 @@ +class ReportGenerationError(Exception): + """Raised when PDF generation fails.""" + + +class ReportTemplateError(Exception): + """Raised when the Jinja2 template is missing or invalid.""" + + +class ReportUploadError(Exception): + """Raised when uploading to S3 fails.""" diff --git a/functions/monthly_reports/accounts/create_report/create_report/generate_pdf.py b/functions/monthly_reports/accounts/create_report/create_report/generate_pdf.py new file mode 100644 index 0000000..c5b8223 --- /dev/null +++ b/functions/monthly_reports/accounts/create_report/create_report/generate_pdf.py @@ -0,0 +1,44 @@ +import io +import os +from datetime import datetime, timezone + +from aws_lambda_powertools import Logger +from jinja2 import Environment, FileSystemLoader, TemplateNotFound, select_autoescape +from xhtml2pdf import pisa + +from .exceptions import ReportGenerationError, ReportTemplateError + + +def generate_transactions_pdf(event: dict, logger: Logger) -> bytes: + current_dir = os.path.dirname(os.path.abspath(__file__)) + env = Environment( + loader=FileSystemLoader(current_dir), + autoescape=select_autoescape(["html", "xml"]), + ) + + try: + template = env.get_template("template.html") + except TemplateNotFound as e: + logger.error("Template 'template.html' not found") + raise ReportTemplateError("Missing template: template.html") from e + + html_out = template.render( + accountId=event["accountId"], + statementPeriod=event["statementPeriod"], + transactions=event["transactions"], + accountBalance=event["accountBalance"], + generationDate=datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M:%S UTC"), + ) + + pdf_buffer = io.BytesIO() + pisa_status = pisa.CreatePDF(io.StringIO(html_out), dest=pdf_buffer) + + if pisa_status.err: + logger.error("xhtml2pdf failed to generate PDF") + raise ReportGenerationError("Error generating PDF") + + pdf_buffer.seek(0) + pdf_bytes = pdf_buffer.getvalue() + + logger.debug("PDF generated (%d bytes).", len(pdf_bytes)) + return pdf_bytes diff --git a/functions/monthly_reports/accounts/create_report/create_report/template.html b/functions/monthly_reports/accounts/create_report/create_report/template.html new file mode 100644 index 0000000..7141492 --- /dev/null +++ b/functions/monthly_reports/accounts/create_report/create_report/template.html @@ -0,0 +1,243 @@ + + + + + + Account Statement + + + +
+
+

Account Statement for {{ accountId }}

+
+ + +
+ Period: + {{ statementPeriod }} +
+ +
+
+ + + + + + + + + + + + + {% for txn in transactions %} + + + + + + + + + {% endfor %} + + + + + +
Transaction IDStatusDescriptionDateTypeAmount
+ {{ txn.id[:8] }}... + + {{ txn.status }} + {{ txn.description }}{{ txn.createdAt[:10] }} + {{ txn.type }} + £{{ "%.2f"|format(txn.amount) }}
Total Balance£{{ "%.2f"|format(accountBalance) }}
+
+ +
+

Statement generated on {{ generationDate[:10] }}

+
+
+
+ + \ No newline at end of file diff --git a/functions/monthly_reports/accounts/create_report/requirements.txt b/functions/monthly_reports/accounts/create_report/requirements.txt new file mode 100644 index 0000000..353f449 --- /dev/null +++ b/functions/monthly_reports/accounts/create_report/requirements.txt @@ -0,0 +1,2 @@ +Jinja2==3.1.6 +xhtml2pdf==0.2.17 \ No newline at end of file diff --git a/functions/monthly_reports/accounts/notify_client/__init__.py b/functions/monthly_reports/accounts/notify_client/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/functions/monthly_reports/accounts/notify_client/notify_client/__init__.py b/functions/monthly_reports/accounts/notify_client/notify_client/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/functions/monthly_reports/accounts/notify_client/notify_client/app.py b/functions/monthly_reports/accounts/notify_client/notify_client/app.py new file mode 100644 index 0000000..840a557 --- /dev/null +++ b/functions/monthly_reports/accounts/notify_client/notify_client/app.py @@ -0,0 +1,137 @@ +import json +import os + +from aws_lambda_powertools import Logger +from aws_lambda_powertools.event_handler import ( + APIGatewayRestResolver, + CORSConfig, +) +from aws_lambda_powertools.event_handler.exceptions import ( + InternalServerError, + UnauthorizedError, + BadRequestError, +) +from aws_lambda_powertools.utilities.typing import LambdaContext + +from authentication.authenticate_request import authenticate_request +from checks import check_user_owns_account +from dynamodb import get_dynamodb_resource +from s3 import get_s3_client +from .date_helpers import period_is_in_future +from .processing import process_report + +SES_NO_REPLY_EMAIL = os.environ.get("SES_NO_REPLY_EMAIL") +REPORTS_BUCKET = os.environ.get("REPORTS_BUCKET") +AWS_REGION = os.environ.get("AWS_REGION") +POWERTOOLS_LOG_LEVEL = os.environ.get("POWERTOOLS_LOG_LEVEL") +COGNITO_USER_POOL_ID = os.environ.get("COGNITO_USER_POOL_ID") +COGNITO_CLIENT_ID = os.environ.get("COGNITO_CLIENT_ID") +DYNAMODB_ENDPOINT = os.environ.get("DYNAMODB_ENDPOINT") +ACCOUNTS_TABLE_NAME = os.environ.get("ACCOUNTS_TABLE_NAME") + +logger = Logger(service="MonthlyAccountReportsNotifyClient", level=POWERTOOLS_LOG_LEVEL) + +s3 = get_s3_client(AWS_REGION, logger) + +dynamodb = get_dynamodb_resource(DYNAMODB_ENDPOINT, AWS_REGION, logger) +if ACCOUNTS_TABLE_NAME: + table = dynamodb.Table(ACCOUNTS_TABLE_NAME) + logger.debug(f"Initialized DynamoDB table: {ACCOUNTS_TABLE_NAME}") +else: + logger.critical("FATAL: ACCOUNTS_TABLE_NAME environment variable not set!") + table = None + +# SES limit: 10 MB total, ~7 MB usable for attachments +MAX_ATTACHMENT_SIZE = 7 * 1024 * 1024 # 7 MB + +app = APIGatewayRestResolver( + cors=CORSConfig(allow_headers=["Content-Type", "Authorization"]) +) + + +@app.get("/accounts//reports/") +def get_account_report(account_id: str, statement_period: str): + try: + event = app.current_event + raw_headers = event.get("headers") or {} + headers = {k.lower(): v for k, v in raw_headers.items()} + + user_id = authenticate_request( + event, + headers, + COGNITO_USER_POOL_ID, + COGNITO_CLIENT_ID, + AWS_REGION.lower(), + logger, + ) + + if not user_id: + raise UnauthorizedError("Unauthorized") + + user_owns_account = check_user_owns_account( + account_id=account_id, user_id=user_id, table=table + ) + + if not user_owns_account: + raise UnauthorizedError("Unauthorized") + + if period_is_in_future(statement_period): + raise BadRequestError("Statement period is in the future") + + result = process_report( + account_id=account_id, + user_id=user_id, + statement_period=statement_period, + cognito_user_pool_id=COGNITO_USER_POOL_ID, + aws_region=AWS_REGION, + reports_bucket=REPORTS_BUCKET, + ses_no_reply_email=SES_NO_REPLY_EMAIL, + max_attachment_size=MAX_ATTACHMENT_SIZE, + logger=logger, + s3_client=s3, + ) + return result + + except UnauthorizedError: + raise + except Exception as e: + logger.error(f"Error processing report: {e}", exc_info=True) + raise InternalServerError("Internal server error") + + +@logger.inject_lambda_context +def lambda_handler(event, context: LambdaContext): + logger.append_keys(request_id=context.aws_request_id) + + if "httpMethod" in event or "requestContext" in event: + return app.resolve(event, context) + else: + account_id = event.get("accountId") + user_id = event.get("userId") + statement_period = event.get("statementPeriod") + + if not account_id or not user_id or not statement_period: + return { + "statusCode": 400, + "body": json.dumps( + {"error": "Missing accountId, userId, or statementPeriod"} + ), + } + + try: + result = process_report( + account_id=account_id, + user_id=user_id, + statement_period=statement_period, + cognito_user_pool_id=COGNITO_USER_POOL_ID, + aws_region=AWS_REGION, + reports_bucket=REPORTS_BUCKET, + ses_no_reply_email=SES_NO_REPLY_EMAIL, + max_attachment_size=MAX_ATTACHMENT_SIZE, + logger=logger, + s3_client=s3, + ) + return result + except Exception as e: + logger.error(f"Error processing report: {e}", exc_info=True) + return {"statusCode": 500, "body": json.dumps({"error": str(e)})} diff --git a/functions/monthly_reports/accounts/notify_client/notify_client/date_helpers.py b/functions/monthly_reports/accounts/notify_client/notify_client/date_helpers.py new file mode 100644 index 0000000..8bb10de --- /dev/null +++ b/functions/monthly_reports/accounts/notify_client/notify_client/date_helpers.py @@ -0,0 +1,19 @@ +import datetime + + +def period_is_in_future(statement_period: str) -> bool: + try: + requested_date = datetime.datetime.strptime(statement_period, "%Y-%m") + except ValueError: + raise ValueError("Invalid statement_period format. Use 'YYYY-MM'.") + + requested_month = datetime.datetime( + requested_date.year, requested_date.month, 1, tzinfo=datetime.timezone.utc + ) + + today = datetime.datetime.now(datetime.timezone.utc) + current_month = datetime.datetime( + today.year, today.month, 1, tzinfo=datetime.timezone.utc + ) + + return requested_month >= current_month diff --git a/functions/monthly_reports/accounts/notify_client/notify_client/processing.py b/functions/monthly_reports/accounts/notify_client/notify_client/processing.py new file mode 100644 index 0000000..a915c02 --- /dev/null +++ b/functions/monthly_reports/accounts/notify_client/notify_client/processing.py @@ -0,0 +1,76 @@ +from aws_lambda_powertools import Logger +from botocore.exceptions import ClientError + +from authentication.user_details import get_user_attributes +from functions.monthly_reports.accounts.notify_client.notify_client.send_report import ( + send_report_as_attachment, + send_report_as_link, +) + + +def process_report( + account_id: str, + user_id: str, + statement_period: str, + cognito_user_pool_id: str, + aws_region: str, + reports_bucket: str, + max_attachment_size: int, + ses_no_reply_email: str, + logger: Logger, + s3_client, +): + s3_key = f"{account_id}/{statement_period}.pdf" + subject = f"Your Account Statement for {statement_period}" + + try: + user_attributes = get_user_attributes( + aws_region=aws_region, + logger=logger, + username=user_id, + user_pool_id=cognito_user_pool_id, + ) + + recipient = user_attributes.get("email") + user_name = user_attributes.get("name", "Customer") + + if not recipient: + raise ValueError(f"User {user_id} has no email attribute in Cognito") + + # Get object metadata first (to check size without downloading the full file) + head = s3_client.head_object(Bucket=reports_bucket, Key=s3_key) + file_size = head["ContentLength"] + + if file_size <= max_attachment_size: + logger.info("PDF is small enough, sending as attachment") + return send_report_as_attachment( + recipient=recipient, + user_name=user_name, + subject=subject, + s3_key=s3_key, + aws_region=aws_region, + reports_bucket=reports_bucket, + ses_no_reply_email=ses_no_reply_email, + logger=logger, + s3_client=s3_client, + ) + else: + logger.info("PDF too large, sending presigned URL") + return send_report_as_link( + recipient=recipient, + user_name=user_name, + subject=subject, + s3_key=s3_key, + aws_region=aws_region, + reports_bucket=reports_bucket, + ses_no_reply_email=ses_no_reply_email, + logger=logger, + s3_client=s3_client, + ) + + except ClientError: + logger.exception("Failed to fetch report from S3") + raise + except Exception: + logger.exception("Exception processing email") + raise diff --git a/functions/monthly_reports/accounts/notify_client/notify_client/send_report.py b/functions/monthly_reports/accounts/notify_client/notify_client/send_report.py new file mode 100644 index 0000000..5de980d --- /dev/null +++ b/functions/monthly_reports/accounts/notify_client/notify_client/send_report.py @@ -0,0 +1,79 @@ +from aws_lambda_powertools import Logger + +from ses import send_user_email, send_user_email_with_attachment + + +def send_report_as_attachment( + recipient: str, + user_name: str, + subject: str, + s3_key: str, + aws_region: str, + reports_bucket: str, + ses_no_reply_email: str, + logger: Logger, + s3_client, +): + # Download PDF from S3 + pdf_obj = s3_client.get_object(Bucket=reports_bucket, Key=s3_key) + pdf_bytes = pdf_obj["Body"].read() + + body_text = f"Hello {user_name},\n\nPlease find your account statement attached.\n\nKind Regards." + + response = send_user_email_with_attachment( + aws_region=aws_region, + logger=logger, + sender_email=ses_no_reply_email, + to_addresses=[recipient], + subject_data=subject, + body_text=body_text, + attachment_bytes=pdf_bytes, + attachment_filename="statement.pdf", + ) + + return { + "status": "success" if response else "failed", + "messageId": response.get("MessageId") if response else None, + "mode": "attachment", + } + + +def send_report_as_link( + recipient: str, + user_name: str, + subject: str, + s3_key: str, + aws_region: str, + reports_bucket: str, + ses_no_reply_email: str, + logger: Logger, + s3_client, +): + presigned_url = s3_client.generate_presigned_url( + "get_object", + Params={"Bucket": reports_bucket, "Key": s3_key}, + ExpiresIn=3600, # 1 hour + ) + + body_text = ( + f"Hello {user_name},\n\n" + f"Your account statement is ready.\n\n" + f"Download it here (valid for 1 hour):\n{presigned_url}\n\n" + f"If you need a new link please request one through the API.\n\n" + f"Kind Regards." + ) + + response = send_user_email( + aws_region=aws_region, + logger=logger, + sender_email=ses_no_reply_email, + to_addresses=[recipient], + subject_data=subject, + text_body_data=body_text, + ) + + return { + "status": "success" if response else "failed", + "messageId": response.get("MessageId") if response else None, + "mode": "link", + } diff --git a/functions/monthly_reports/accounts/notify_client/requirements.txt b/functions/monthly_reports/accounts/notify_client/requirements.txt new file mode 100644 index 0000000..e69de29 diff --git a/functions/monthly_reports/accounts/process_pending_reports/__init__.py b/functions/monthly_reports/accounts/process_pending_reports/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/functions/monthly_reports/accounts/process_pending_reports/process_pending_reports/__init__.py b/functions/monthly_reports/accounts/process_pending_reports/process_pending_reports/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/functions/monthly_reports/accounts/process_pending_reports/process_pending_reports/app.py b/functions/monthly_reports/accounts/process_pending_reports/process_pending_reports/app.py new file mode 100644 index 0000000..60572c8 --- /dev/null +++ b/functions/monthly_reports/accounts/process_pending_reports/process_pending_reports/app.py @@ -0,0 +1,183 @@ +import json +import os + +from aws_lambda_powertools import Logger +from aws_lambda_powertools.utilities.typing import LambdaContext + +from dynamodb import get_dynamodb_resource +from monthly_reports.metrics import initialize_metrics, merge_metrics +from monthly_reports.processing import ( + process_accounts_scan_continuation, + process_batch_continuation, +) +from monthly_reports.responses import create_response +from monthly_reports.sqs import send_bad_account_to_dlq + +from sqs import get_sqs_client +from sfn import get_sfn_client + +ENVIRONMENT_NAME = os.environ.get("ENVIRONMENT_NAME", "dev") +POWERTOOLS_LOG_LEVEL = os.environ.get("POWERTOOLS_LOG_LEVEL", "INFO").upper() +ACCOUNTS_TABLE_NAME = os.environ.get("ACCOUNTS_TABLE_NAME") +STATE_MACHINE_ARN = os.environ.get("STATE_MACHINE_ARN") +DYNAMODB_ENDPOINT = os.environ.get("DYNAMODB_ENDPOINT") +SQS_ENDPOINT = os.environ.get("SQS_ENDPOINT") +CONTINUATION_QUEUE_URL = os.environ.get("CONTINUATION_QUEUE_URL") +DLQ_URL = os.environ.get("DLQ_URL") +AWS_REGION = os.environ.get("AWS_REGION", "eu-west-2") + +PAGE_SIZE = 50 +BATCH_SIZE = 10 +SAFETY_BUFFER = 30 + +logger = Logger(service="MonthlyAccountReportsContinuation", level=POWERTOOLS_LOG_LEVEL) + +dynamodb = get_dynamodb_resource(DYNAMODB_ENDPOINT, AWS_REGION, logger) +sqs_client = get_sqs_client(SQS_ENDPOINT, AWS_REGION, logger) +sfn_client = get_sfn_client(AWS_REGION, logger) + +if ACCOUNTS_TABLE_NAME: + accounts_table = dynamodb.Table(ACCOUNTS_TABLE_NAME) + logger.debug(f"Initialized DynamoDB table: {ACCOUNTS_TABLE_NAME}") +else: + logger.critical("FATAL: ACCOUNTS_TABLE_NAME environment variable not set!") + accounts_table = None + + +@logger.inject_lambda_context +def lambda_handler(event, context: LambdaContext): + logger.info("Processing SQS continuation messages") + + metrics = initialize_metrics() + + try: + # Process each SQS record + for record in event.get("Records", []): + try: + message_body = json.loads(record["body"]) + except json.JSONDecodeError as e: + logger.error(f"Failed to parse message body as JSON: {e}") + if DLQ_URL and AWS_REGION: + try: + error_data = { + "lambda_function": "process-pending-reports", + "error_type": "json_parse_error", + "raw_message": record.get("body"), + "message_id": record.get("messageId"), + } + send_bad_account_to_dlq( + error_data, + "unknown", + f"JSON parse error: {str(e)}", + SQS_ENDPOINT, + DLQ_URL, + AWS_REGION, + logger, + ) + except Exception as dlq_error: + logger.error(f"Failed to send parse error to DLQ: {dlq_error}") + continue + + message_attributes = record.get("messageAttributes", {}) + + continuation_type = message_attributes.get("continuation_type", {}).get( + "stringValue" + ) + + if continuation_type == "accounts_scan": + logger.info("Processing accounts scan continuation") + scan_metrics = process_accounts_scan_continuation( + message_body["scan_params"], + message_body["statement_period"], + context, + logger, + accounts_table, + sfn_client, + STATE_MACHINE_ARN, + SQS_ENDPOINT, + CONTINUATION_QUEUE_URL, + AWS_REGION, + PAGE_SIZE, + BATCH_SIZE, + SAFETY_BUFFER, + DLQ_URL, + ) + merge_metrics(metrics, scan_metrics) + + elif continuation_type == "batch_continuation": + logger.info("Processing batch continuation") + batch_metrics = process_batch_continuation( + message_body["scan_params"], + message_body["statement_period"], + message_body["remaining_accounts"], + message_body.get("last_evaluated_key"), + context, + logger, + accounts_table, + sfn_client, + STATE_MACHINE_ARN, + SQS_ENDPOINT, + CONTINUATION_QUEUE_URL, + AWS_REGION, + PAGE_SIZE, + BATCH_SIZE, + SAFETY_BUFFER, + DLQ_URL, + ) + merge_metrics(metrics, batch_metrics) + + else: + logger.warning(f"Unknown continuation type: {continuation_type}") + if DLQ_URL and AWS_REGION: + try: + error_data = { + "lambda_function": "process-pending-reports", + "error_type": "unknown_continuation_type", + "continuation_type": continuation_type, + "message_body": message_body, + "message_id": record.get("messageId"), + } + statement_period = message_body.get( + "statement_period", "unknown" + ) + send_bad_account_to_dlq( + error_data, + statement_period, + f"Unknown continuation type: {continuation_type}", + SQS_ENDPOINT, + DLQ_URL, + AWS_REGION, + logger, + ) + except Exception as dlq_error: + logger.error( + f"Failed to send unknown continuation type to DLQ: {dlq_error}" + ) + + except Exception as e: + logger.error( + f"Critical error during continuation processing: {e}", exc_info=True + ) + if DLQ_URL and AWS_REGION: + try: + + error_account = { + "lambda_function": "process-pending-reports", + "error_type": "critical_lambda_error", + "error_details": str(e), + "event": event, + } + send_bad_account_to_dlq( + error_account, + "unknown", + f"Critical lambda error: {str(e)}", + SQS_ENDPOINT, + DLQ_URL, + AWS_REGION, + logger, + ) + except Exception as dlq_error: + logger.error(f"Failed to send critical error to DLQ: {dlq_error}") + raise + + return create_response(metrics, "COMPLETED", logger) diff --git a/functions/monthly_reports/accounts/process_pending_reports/requirements.txt b/functions/monthly_reports/accounts/process_pending_reports/requirements.txt new file mode 100644 index 0000000..e69de29 diff --git a/functions/monthly_reports/accounts/trigger/__init__.py b/functions/monthly_reports/accounts/trigger/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/functions/monthly_reports/accounts/trigger/trigger/__init__.py b/functions/monthly_reports/accounts/trigger/trigger/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/functions/monthly_reports/accounts/trigger/trigger/app.py b/functions/monthly_reports/accounts/trigger/trigger/app.py new file mode 100644 index 0000000..06ee463 --- /dev/null +++ b/functions/monthly_reports/accounts/trigger/trigger/app.py @@ -0,0 +1,144 @@ +import os + +from aws_lambda_powertools import Logger +from aws_lambda_powertools.utilities.typing import LambdaContext + +from dynamodb import get_dynamodb_resource, get_paginated_table_data +from monthly_reports.helpers import get_statement_period +from monthly_reports.metrics import initialize_metrics, merge_metrics +from monthly_reports.processing import process_accounts_page +from monthly_reports.responses import create_response +from monthly_reports.sqs import send_continuation_message +from monthly_reports.sqs import send_bad_account_to_dlq + +from sqs import get_sqs_client +from sfn import get_sfn_client + +ENVIRONMENT_NAME = os.environ.get("ENVIRONMENT_NAME", "dev") +POWERTOOLS_LOG_LEVEL = os.environ.get("POWERTOOLS_LOG_LEVEL", "INFO").upper() +ACCOUNTS_TABLE_NAME = os.environ.get("ACCOUNTS_TABLE_NAME") +STATE_MACHINE_ARN = os.environ.get("STATE_MACHINE_ARN") +DYNAMODB_ENDPOINT = os.environ.get("DYNAMODB_ENDPOINT") +SQS_ENDPOINT = os.environ.get("SQS_ENDPOINT") +CONTINUATION_QUEUE_URL = os.environ.get("CONTINUATION_QUEUE_URL") +DLQ_URL = os.environ.get("DLQ_URL") +AWS_REGION = os.environ.get("AWS_REGION", "eu-west-2") + +PAGE_SIZE = 50 +BATCH_SIZE = 10 +SAFETY_BUFFER = 30 + +logger = Logger(service="MonthlyAccountReportsTrigger", level=POWERTOOLS_LOG_LEVEL) + +dynamodb = get_dynamodb_resource(DYNAMODB_ENDPOINT, AWS_REGION, logger) +sqs_client = get_sqs_client(SQS_ENDPOINT, AWS_REGION, logger) +sfn_client = get_sfn_client(AWS_REGION, logger) + +if ACCOUNTS_TABLE_NAME: + accounts_table = dynamodb.Table(ACCOUNTS_TABLE_NAME) + logger.debug(f"Initialized DynamoDB table: {ACCOUNTS_TABLE_NAME}") +else: + logger.critical("FATAL: ACCOUNTS_TABLE_NAME environment variable not set!") + accounts_table = None + + +@logger.inject_lambda_context +def lambda_handler(_event, context: LambdaContext): + logger.info("Starting monthly account reports processing from EventBridge trigger") + + statement_period = get_statement_period() + logger.info(f"Starting new processing for period: {statement_period}") + + metrics = initialize_metrics() + + if not CONTINUATION_QUEUE_URL: + logger.critical("CONTINUATION_QUEUE_URL is not set. Cannot test SQS sending.") + return create_response(metrics, "ERROR_NO_CONTINUATION_QUEUE", logger) + + scan_params = { + "ProjectionExpression": "accountId, userId, balance", + } + + try: + while True: + remaining_time = context.get_remaining_time_in_millis() / 1000.0 + if remaining_time < SAFETY_BUFFER: + logger.warning( + f"Approaching Lambda timeout. Processed {metrics['pages_processed']} pages." + ) + send_continuation_message( + scan_params, + statement_period, + None, + None, + "accounts_scan", + SQS_ENDPOINT, + CONTINUATION_QUEUE_URL, + AWS_REGION, + logger, + ) + return create_response(metrics, "TIMEOUT_CONTINUATION", logger) + + accounts_page, last_evaluated_key = get_paginated_table_data( + scan_params=scan_params, + index_name=None, + table=accounts_table, + logger=logger, + page_size=PAGE_SIZE, + ) + + metrics["pages_processed"] += 1 + + if not accounts_page: + logger.info("No more accounts to process") + break + + page_metrics = process_accounts_page( + accounts_page, + statement_period, + context, + logger, + sfn_client, + STATE_MACHINE_ARN, + scan_params, + last_evaluated_key, + SQS_ENDPOINT, + CONTINUATION_QUEUE_URL, + AWS_REGION, + BATCH_SIZE, + SAFETY_BUFFER, + DLQ_URL, + ) + + merge_metrics(metrics, page_metrics) + + if last_evaluated_key: + scan_params["ExclusiveStartKey"] = last_evaluated_key + logger.debug("More pages available, continuing...") + else: + logger.info("All pages processed successfully") + break + + except Exception as e: + logger.error(f"Critical error during processing: {e}", exc_info=True) + if DLQ_URL and AWS_REGION: + try: + error_account = { + "lambda_function": "monthly-reports-trigger", + "error_type": "critical_lambda_error", + "error_details": str(e), + } + send_bad_account_to_dlq( + error_account, + statement_period, + f"Critical lambda error: {str(e)}", + SQS_ENDPOINT, + DLQ_URL, + AWS_REGION, + logger, + ) + except Exception as dlq_error: + logger.error(f"Failed to send critical error to DLQ: {dlq_error}") + raise + + return create_response(metrics, "COMPLETED", logger) diff --git a/functions/transactions/get_transactions/requirements.txt b/functions/transactions/get_transactions/requirements.txt index 8637f8f..216efe3 100644 --- a/functions/transactions/get_transactions/requirements.txt +++ b/functions/transactions/get_transactions/requirements.txt @@ -1,2 +1,2 @@ -aws_lambda_powertools==3.12.0 +aws_lambda_powertools==3.17.0 boto3==1.38.13 \ No newline at end of file diff --git a/functions/transactions/process_transactions/process_transactions/app.py b/functions/transactions/process_transactions/process_transactions/app.py index b4083ba..23275e7 100644 --- a/functions/transactions/process_transactions/process_transactions/app.py +++ b/functions/transactions/process_transactions/process_transactions/app.py @@ -5,8 +5,9 @@ from aws_lambda_powertools.utilities.typing import LambdaContext from dynamodb import get_dynamodb_resource -from sqs import send_dynamodb_record_to_dlq +from sqs import send_message_to_sqs from .exceptions import BusinessLogicError, TransactionSystemError +from .sqs import format_sqs_message, get_message_attributes from .transaction_helpers import process_single_transaction, update_transaction_status ENVIRONMENT_NAME = os.environ.get("ENVIRONMENT_NAME", "dev") @@ -121,23 +122,36 @@ def lambda_handler(event, _context: LambdaContext): logger.error( f"Failed to update transaction status to FAILED: {update_error}" ) - if not send_dynamodb_record_to_dlq( - record=record, - sqs_url=SQS_ENDPOINT, - dlq_url=TRANSACTION_PROCESSING_DLQ_URL, + if not send_message_to_sqs( + message=format_sqs_message( + record, + f"Failed to update status after business logic error: {e}", + ), + message_attributes=get_message_attributes( + error_type="StatusUpdateError", + environment_name=ENVIRONMENT_NAME, + idempotency_key=idempotency_key, + ), + sqs_endpoint=SQS_ENDPOINT, + sqs_url=TRANSACTION_PROCESSING_DLQ_URL, aws_region=AWS_REGION, - error_message=f"Failed to update status after business logic error: {e}", logger=logger, ): critical_failures += 1 else: logger.error(f"No idempotency key found for business logic error: {e}") - if not send_dynamodb_record_to_dlq( - record=record, - sqs_url=SQS_ENDPOINT, - dlq_url=TRANSACTION_PROCESSING_DLQ_URL, + if not send_message_to_sqs( + message=format_sqs_message( + record, f"Business logic error without idempotency key: {e}" + ), + message_attributes=get_message_attributes( + error_type="BusinessLogicError", + environment_name=ENVIRONMENT_NAME, + idempotency_key=idempotency_key, + ), + sqs_endpoint=SQS_ENDPOINT, + sqs_url=TRANSACTION_PROCESSING_DLQ_URL, aws_region=AWS_REGION, - error_message=f"Business logic error without idempotency key: {e}", logger=logger, ): critical_failures += 1 @@ -146,12 +160,15 @@ def lambda_handler(event, _context: LambdaContext): system_failures += 1 logger.error(f"System error for record {sequence_number}: {e}") - if not send_dynamodb_record_to_dlq( - record=record, - sqs_url=SQS_ENDPOINT, - dlq_url=TRANSACTION_PROCESSING_DLQ_URL, + if not send_message_to_sqs( + message=format_sqs_message(record, str(e)), + message_attributes=get_message_attributes( + error_type="TransactionSystemError", + environment_name=ENVIRONMENT_NAME, + ), + sqs_endpoint=SQS_ENDPOINT, + sqs_url=TRANSACTION_PROCESSING_DLQ_URL, aws_region=AWS_REGION, - error_message=str(e), logger=logger, ): critical_failures += 1 @@ -165,12 +182,14 @@ def lambda_handler(event, _context: LambdaContext): f"Unknown error for record {sequence_number}: {e}", exc_info=True ) - if not send_dynamodb_record_to_dlq( - record=record, - sqs_url=SQS_ENDPOINT, - dlq_url=TRANSACTION_PROCESSING_DLQ_URL, + if not send_message_to_sqs( + message=format_sqs_message(record, f"Unknown error: {str(e)}"), + message_attributes=get_message_attributes( + error_type="UnknownError", environment_name=ENVIRONMENT_NAME + ), + sqs_endpoint=SQS_ENDPOINT, + sqs_url=TRANSACTION_PROCESSING_DLQ_URL, aws_region=AWS_REGION, - error_message=f"Unknown error: {str(e)}", logger=logger, ): critical_failures += 1 diff --git a/functions/transactions/process_transactions/process_transactions/sqs.py b/functions/transactions/process_transactions/process_transactions/sqs.py new file mode 100644 index 0000000..deb842f --- /dev/null +++ b/functions/transactions/process_transactions/process_transactions/sqs.py @@ -0,0 +1,67 @@ +import datetime + + +def format_sqs_message(record: dict, error_message: str = ""): + if not isinstance(record, dict): + raise ValueError("Record must be a dictionary") + return { + "originalRecord": record, + "errorMessage": error_message, + "timestamp": record.get("dynamodb", {}).get("ApproximateCreationDateTime"), + "sequenceNumber": record.get("dynamodb", {}).get("SequenceNumber"), + } + + +def get_message_attributes( + error_type: str, environment_name: str, idempotency_key: str = None +) -> dict: + base_attributes = { + "Source": {"StringValue": "ProcessTransactions", "DataType": "String"}, + "Environment": {"StringValue": environment_name, "DataType": "String"}, + "Timestamp": { + "StringValue": datetime.datetime.now(datetime.UTC).isoformat(), + "DataType": "String", + }, + } + + if error_type == "BusinessLogicError": + base_attributes.update( + { + "ErrorType": { + "StringValue": "BusinessLogicError", + "DataType": "String", + }, + "ErrorCategory": {"StringValue": "RECOVERABLE", "DataType": "String"}, + "HasIdempotencyKey": { + "StringValue": str(bool(idempotency_key)), + "DataType": "String", + }, + } + ) + elif error_type == "TransactionSystemError": + base_attributes.update( + { + "ErrorType": { + "StringValue": "TransactionSystemError", + "DataType": "String", + }, + "ErrorCategory": { + "StringValue": "SYSTEM_FAILURE", + "DataType": "String", + }, + "RequiresRetry": {"StringValue": "true", "DataType": "String"}, + } + ) + else: + base_attributes.update( + { + "ErrorType": {"StringValue": "UnknownError", "DataType": "String"}, + "ErrorCategory": { + "StringValue": "SYSTEM_FAILURE", + "DataType": "String", + }, + "RequiresRetry": {"StringValue": "true", "DataType": "String"}, + } + ) + + return base_attributes diff --git a/functions/transactions/request_transaction/requirements.txt b/functions/transactions/request_transaction/requirements.txt index 8637f8f..216efe3 100644 --- a/functions/transactions/request_transaction/requirements.txt +++ b/functions/transactions/request_transaction/requirements.txt @@ -1,2 +1,2 @@ -aws_lambda_powertools==3.12.0 +aws_lambda_powertools==3.17.0 boto3==1.38.13 \ No newline at end of file diff --git a/layers/python/authentication/authentication/user_details.py b/layers/python/authentication/authentication/user_details.py new file mode 100644 index 0000000..bb051d1 --- /dev/null +++ b/layers/python/authentication/authentication/user_details.py @@ -0,0 +1,20 @@ +import boto3 +from aws_lambda_powertools import Logger + + +def get_user_attributes( + aws_region: str, logger: Logger, username: str, user_pool_id: str +) -> dict: + try: + cognito = boto3.client("cognito-idp", region_name=aws_region) + + response = cognito.admin_get_user( + UserPoolId=user_pool_id, + Username=username, + ) + attrs = {attr["Name"]: attr["Value"] for attr in response["UserAttributes"]} + logger.info(f"Fetched attributes for user: {username}.") + return attrs + except Exception as e: + logger.exception(f"Failed to fetch user {username} from Cognito") + raise e diff --git a/layers/python/authentication/requirements.txt b/layers/python/authentication/requirements.txt index fe84f6b..b9710d0 100644 --- a/layers/python/authentication/requirements.txt +++ b/layers/python/authentication/requirements.txt @@ -1,4 +1,4 @@ -aws_lambda_powertools==3.12.0 +aws_lambda_powertools==3.17.0 PyJWT==2.10.1 requests==2.32.3 cryptography==44.0.3 \ No newline at end of file diff --git a/layers/python/helpers/dynamodb.py b/layers/python/helpers/dynamodb.py index aec9f6b..504e916 100644 --- a/layers/python/helpers/dynamodb.py +++ b/layers/python/helpers/dynamodb.py @@ -1,5 +1,6 @@ import boto3 from aws_lambda_powertools import Logger +from botocore.exceptions import ClientError def get_dynamodb_resource(dynamodb_endpoint: str, aws_region: str, logger: Logger): @@ -26,3 +27,29 @@ def get_dynamodb_resource(dynamodb_endpoint: str, aws_region: str, logger: Logge except Exception: logger.error("Failed to initialize DynamoDB resource", exc_info=True) raise + + +def get_paginated_table_data( + scan_params, index_name, table, logger: Logger, page_size: int = 10 +): + if scan_params is None: + scan_params = {} + + scan_params = scan_params.copy() + scan_params["Limit"] = page_size + + if index_name: + scan_params["IndexName"] = index_name + + try: + response = table.scan(**scan_params) + items = response.get("Items", []) + last_evaluated_key = response.get("LastEvaluatedKey") + + logger.info(f"Fetched {len(items)} items from DynamoDB") + + return items, last_evaluated_key + + except ClientError as exception: + logger.error(f"Error during DynamoDB scan: {exception}", exc_info=True) + raise exception diff --git a/layers/python/helpers/requirements.txt b/layers/python/helpers/requirements.txt index 8637f8f..216efe3 100644 --- a/layers/python/helpers/requirements.txt +++ b/layers/python/helpers/requirements.txt @@ -1,2 +1,2 @@ -aws_lambda_powertools==3.12.0 +aws_lambda_powertools==3.17.0 boto3==1.38.13 \ No newline at end of file diff --git a/layers/python/helpers/s3.py b/layers/python/helpers/s3.py new file mode 100644 index 0000000..619d396 --- /dev/null +++ b/layers/python/helpers/s3.py @@ -0,0 +1,13 @@ +from logging import Logger + +import boto3 + + +def get_s3_client(aws_region: str, logger: Logger): + try: + client = boto3.client("s3", region_name=aws_region) + logger.info("Initialized S3 client with default endpoint") + return client + except Exception: + logger.error("Failed to initialize S3 client", exc_info=True) + raise diff --git a/layers/python/helpers/ses.py b/layers/python/helpers/ses.py index 522d526..3d3798c 100644 --- a/layers/python/helpers/ses.py +++ b/layers/python/helpers/ses.py @@ -3,14 +3,14 @@ import boto3 from aws_lambda_powertools import Logger +from email.mime.multipart import MIMEMultipart +from email.mime.application import MIMEApplication +from email.mime.text import MIMEText def get_ses_client(aws_region: str, logger: Logger): """ Initialise and return an AWS SES client for the specified region. - - Raises: - Exception: If the SES client cannot be initialised. """ try: logger.info("Initialized SES client with default endpoint") @@ -26,11 +26,11 @@ def send_user_email( sender_email: str, to_addresses: List[str], subject_data: str, - subject_charset: str, + subject_charset: str = "UTF-8", text_body_data: Optional[str] = None, - text_body_charset: Optional[str] = None, + text_body_charset: str = "UTF-8", html_body_data: Optional[str] = None, - html_body_charset: Optional[str] = None, + html_body_charset: str = "UTF-8", cc_addresses: Optional[List[str]] = None, bcc_addresses: Optional[List[str]] = None, reply_to_addresses: Optional[List[str]] = None, @@ -38,45 +38,19 @@ def send_user_email( tags: Optional[List[Dict[str, str]]] = None, ): """ - Send an email using AWS SES with configurable sender, recipients, subject, and body content. - - At least one of `text_body_data` or `html_body_data` must be provided. Supports optional CC, BCC, reply-to addresses, return path, and message tags. Returns `True` if the email is sent successfully, or `False` if sending fails or required body content is missing. - - Parameters: - sender_email (str): The email address of the sender. - to_addresses (List[str]): List of recipient email addresses. - subject_data (str): The subject line of the email. - subject_charset (str): Character set for the subject line. - text_body_data (Optional[str]): Plain text content of the email body. - text_body_charset (Optional[str]): Character set for the plain text body. - html_body_data (Optional[str]): HTML content of the email body. - html_body_charset (Optional[str]): Character set for the HTML body. - cc_addresses (Optional[List[str]]): List of CC recipient email addresses. - bcc_addresses (Optional[List[str]]): List of BCC recipient email addresses. - reply_to_addresses (Optional[List[str]]): List of reply-to email addresses. - return_path (Optional[str]): Email address for bounce and complaint notifications. - tags (Optional[List[Dict[str, str]]]): List of tags to apply to the email. - - Returns: - bool: True if the email was sent successfully, False otherwise. + Send a simple email (text and/or HTML) using AWS SES. """ ses_client = get_ses_client(aws_region=aws_region, logger=logger) message_body = {} if text_body_data: - message_body["Text"] = { - "Data": text_body_data, - "Charset": text_body_charset if text_body_charset else "UTF-8", - } + message_body["Text"] = {"Data": text_body_data, "Charset": text_body_charset} if html_body_data: - message_body["Html"] = { - "Data": html_body_data, - "Charset": html_body_charset if html_body_charset else "UTF-8", - } + message_body["Html"] = {"Data": html_body_data, "Charset": html_body_charset} if not message_body: logger.error("Email must contain at least a text or HTML body.") - return False + raise Exception("Email must contain at least a text or HTML body.") destination = {"ToAddresses": to_addresses} if cc_addresses: @@ -84,25 +58,89 @@ def send_user_email( if bcc_addresses: destination["BccAddresses"] = bcc_addresses + request_params = { + "Source": sender_email, + "Destination": destination, + "Message": { + "Subject": {"Data": subject_data, "Charset": subject_charset}, + "Body": message_body, + }, + } + + if reply_to_addresses: + request_params["ReplyToAddresses"] = reply_to_addresses + if return_path: + request_params["ReturnPath"] = return_path + if tags: + request_params["Tags"] = tags + try: - ses_client.send_email( + response = ses_client.send_email(**request_params) + + logger.info( + f"Successfully sent email to {json.dumps(to_addresses)}, " + f"MessageId={response['MessageId']}" + ) + return response + + except Exception as e: + logger.error(f"Failed to send email: {e}", exc_info=True) + raise e + + +def send_user_email_with_attachment( + aws_region: str, + logger: Logger, + sender_email: str, + to_addresses: List[str], + subject_data: str, + body_text: str, + attachment_bytes: bytes, + attachment_filename: str, + cc_addresses: Optional[List[str]] = None, + bcc_addresses: Optional[List[str]] = None, +): + """ + Send an email with a single attachment using AWS SES (send_raw_email). + """ + ses_client = get_ses_client(aws_region=aws_region, logger=logger) + + msg = MIMEMultipart() + msg["Subject"] = subject_data + msg["From"] = sender_email + msg["To"] = ", ".join(to_addresses) + if cc_addresses: + msg["Cc"] = ", ".join(cc_addresses) + if bcc_addresses: + msg["Bcc"] = ", ".join(bcc_addresses) + + # Attach plain text body + msg.attach(MIMEText(body_text, "plain")) + + # Attach the file + part = MIMEApplication(attachment_bytes) + part.add_header("Content-Disposition", "attachment", filename=attachment_filename) + msg.attach(part) + + try: + destinations = list(to_addresses) + if cc_addresses: + destinations.extend(cc_addresses) + if bcc_addresses: + destinations.extend(bcc_addresses) + + response = ses_client.send_raw_email( Source=sender_email, - Destination=destination, - Message={ - "Subject": { - "Data": subject_data, - "Charset": subject_charset if subject_charset else "UTF-8", - }, - "Body": message_body, - }, - ReplyToAddresses=reply_to_addresses, - ReturnPath=return_path, - Tags=tags, + Destinations=destinations, + RawMessage={"Data": msg.as_string()}, ) - logger.info(f"Successfully sent email to users: {json.dumps(to_addresses)}") - return True + logger.info( + f"Successfully sent email with attachment to {json.dumps(to_addresses)}, " + f"MessageId={response['MessageId']}" + ) + return response except Exception as e: - logger.error(f"Failed to send email: {e}") - return False + logger.error(f"Failed to send email with attachment: {e}", exc_info=True) + raise e diff --git a/layers/python/helpers/sfn.py b/layers/python/helpers/sfn.py new file mode 100644 index 0000000..fcc1468 --- /dev/null +++ b/layers/python/helpers/sfn.py @@ -0,0 +1,12 @@ +import boto3 +from aws_lambda_powertools import Logger + + +def get_sfn_client(aws_region: str, logger: Logger): + try: + client = boto3.client("stepfunctions", region_name=aws_region) + logger.info("Initialized SFN client with default endpoint") + return client + except Exception: + logger.error("Failed to initialize SFN client", exc_info=True) + raise diff --git a/layers/python/helpers/sqs.py b/layers/python/helpers/sqs.py index 5eee409..84ec3b5 100644 --- a/layers/python/helpers/sqs.py +++ b/layers/python/helpers/sqs.py @@ -28,31 +28,20 @@ def get_sqs_client(sqs_endpoint: str, aws_region: str, logger: Logger): raise -def send_dynamodb_record_to_dlq( - record: dict, +def send_message_to_sqs( + message: dict, + message_attributes: dict, sqs_endpoint: str, - dlq_url: str, + sqs_url: str, aws_region: str, - error_message: str, logger: Logger, ): - """ - Send a DynamoDB stream record to an SQS Dead Letter Queue (DLQ). - - If the DLQ URL is not provided, logs an error and returns False. Constructs a message containing the original record, an error message, and relevant metadata, then sends it to the specified DLQ. Returns True if the message is sent successfully, otherwise logs the failure and returns False. - - Parameters: - record (dict): The DynamoDB stream record to send. - sqs_endpoint (str): Optional custom SQS endpoint URL. - dlq_url (str): The URL of the SQS Dead Letter Queue. - aws_region (str): AWS region for the SQS client. - error_message (str): Description of the error that triggered the DLQ send. + if not sqs_url: + logger.error("SQS URL not configured, cannot send message to DLQ") + return False - Returns: - bool: True if the message was sent successfully, False otherwise. - """ - if not dlq_url: - logger.error("DLQ URL not configured, cannot send message to DLQ") + if not message: + logger.error("Message is required to send to SQS") return False sqs_client = get_sqs_client( @@ -60,26 +49,15 @@ def send_dynamodb_record_to_dlq( ) try: - dlq_message = { - "originalRecord": record, - "errorMessage": error_message, - "timestamp": record.get("dynamodb", {}).get("ApproximateCreationDateTime"), - "sequenceNumber": record.get("dynamodb", {}).get("SequenceNumber"), - } - sqs_client.send_message( - QueueUrl=dlq_url, - MessageBody=json.dumps(dlq_message), - MessageAttributes={ - "ErrorType": {"StringValue": "SystemError", "DataType": "String"} - }, + QueueUrl=sqs_url, + MessageBody=json.dumps(message), + MessageAttributes=message_attributes, ) - logger.info( - f"Successfully sent record to DLQ: {record.get('dynamodb', {}).get('SequenceNumber')}" - ) + logger.info("Successfully sent message to SQS queue.") return True except Exception as e: - logger.error(f"Failed to send message to DLQ: {e}") + logger.error(f"Failed to send message to SQS: {e}") return False diff --git a/layers/python/monthly_reports/monthly_reports/__init__.py b/layers/python/monthly_reports/monthly_reports/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/layers/python/monthly_reports/monthly_reports/helpers.py b/layers/python/monthly_reports/monthly_reports/helpers.py new file mode 100644 index 0000000..0d4fad3 --- /dev/null +++ b/layers/python/monthly_reports/monthly_reports/helpers.py @@ -0,0 +1,11 @@ +import datetime + + +def get_statement_period(): + """Get the statement period for the previous month""" + today = datetime.datetime.now(datetime.UTC) + first_day_of_current_month = today.replace( + day=1, hour=0, minute=0, second=0, microsecond=0 + ) + last_day_of_previous_month = first_day_of_current_month - datetime.timedelta(days=1) + return last_day_of_previous_month.strftime("%Y-%m") diff --git a/layers/python/monthly_reports/monthly_reports/metrics.py b/layers/python/monthly_reports/monthly_reports/metrics.py new file mode 100644 index 0000000..3a94b93 --- /dev/null +++ b/layers/python/monthly_reports/monthly_reports/metrics.py @@ -0,0 +1,17 @@ +def initialize_metrics(): + """Initialize metrics dictionary""" + return { + "processed_count": 0, + "failed_starts_count": 0, + "skipped_count": 0, + "already_exists_count": 0, + "batches_processed": 0, + "pages_processed": 0, + } + + +def merge_metrics(target_metrics, source_metrics): + """Merge metrics from source into target""" + for key, value in source_metrics.items(): + if key in target_metrics: + target_metrics[key] += value diff --git a/layers/python/monthly_reports/monthly_reports/processing.py b/layers/python/monthly_reports/monthly_reports/processing.py new file mode 100644 index 0000000..a4c553c --- /dev/null +++ b/layers/python/monthly_reports/monthly_reports/processing.py @@ -0,0 +1,374 @@ +from dynamodb import get_paginated_table_data + +from .metrics import merge_metrics, initialize_metrics +from .sfn import start_sfn_execution_with_retry +from .sqs import send_continuation_message, send_bad_account_to_dlq + + +def chunk_accounts(accounts, chunk_size=10): + for i in range(0, len(accounts), chunk_size): + yield accounts[i : i + chunk_size] + + +def process_account_batch( + accounts_batch, + statement_period, + sfn_client, + logger, + state_machine_arn, + sqs_endpoint=None, + dlq_url=None, + aws_region=None, +): + processed_count = 0 + skipped_count = 0 + already_exists_count = 0 + failed_starts_count = 0 + + for account in accounts_batch: + account_id = account.get("accountId") + user_id = account.get("userId") + + if not all([account_id, user_id]): + error_reason = ( + f"Missing required fields - accountId: {bool(account_id)}, " + f"userId: {bool(user_id)}" + ) + logger.warning(f"Skipping account with missing data: {account}") + + if dlq_url and aws_region: + send_bad_account_to_dlq( + account, + statement_period, + error_reason, + sqs_endpoint, + dlq_url, + aws_region, + logger, + ) + + skipped_count += 1 + continue + + sf_input = { + "accountId": account_id, + "userId": user_id, + "accountBalance": float(account.get("balance", 0)), + "statementPeriod": statement_period, + } + + base_name = f"Stmt-{statement_period}-{account_id}" + execution_name = base_name[:80] + + try: + result = start_sfn_execution_with_retry( + sfn_client, state_machine_arn, execution_name, sf_input, logger + ) + + if result == "processed": + processed_count += 1 + elif result == "already_exists": + already_exists_count += 1 + else: + failed_starts_count += 1 + if dlq_url and aws_region: + send_bad_account_to_dlq( + account, + statement_period, + f"Step Function execution failed: {result}", + sqs_endpoint, + dlq_url, + aws_region, + logger, + ) + + except Exception as e: + logger.error(f"Failed to start SF execution for account {account_id}: {e}") + failed_starts_count += 1 + if dlq_url and aws_region: + send_bad_account_to_dlq( + account, + statement_period, + f"Step Function execution exception: {str(e)}", + sqs_endpoint, + dlq_url, + aws_region, + logger, + ) + + return { + "processed": processed_count, + "already_exists": already_exists_count, + "failed_starts": failed_starts_count, + "skipped": skipped_count, + } + + +def process_accounts_page( + accounts_page, + statement_period, + context, + logger, + sfn_client, + state_machine_arn, + scan_params, + last_evaluated_key, + sqs_endpoint, + continuation_queue_url, + aws_region, + batch_size=10, + safety_buffer=30, + dlq_url=None, +): + metrics = initialize_metrics() + + account_batches = list(chunk_accounts(accounts_page, chunk_size=batch_size)) + logger.info( + f"Processing {len(accounts_page)} accounts in {len(account_batches)} batches" + ) + + batch_metrics = process_account_batches( + account_batches, + statement_period, + context, + logger, + sfn_client, + state_machine_arn, + scan_params, + last_evaluated_key, + sqs_endpoint, + continuation_queue_url, + aws_region, + safety_buffer, + dlq_url, + ) + merge_metrics(metrics, batch_metrics) + + return metrics + + +def process_account_batches( + account_batches, + statement_period, + context, + logger, + sfn_client, + state_machine_arn, + scan_params, + last_evaluated_key, + sqs_endpoint, + continuation_queue_url, + aws_region, + safety_buffer=30, + dlq_url=None, +): + """Process multiple batches of accounts""" + metrics = initialize_metrics() + + for i, batch in enumerate(account_batches): + remaining_time = context.get_remaining_time_in_millis() / 1000.0 + + if remaining_time < safety_buffer: + logger.warning("Timeout approaching during batch processing") + + remaining_batches = account_batches[i:] + remaining_accounts = [] + for remaining_batch in remaining_batches: + remaining_accounts.extend(remaining_batch) + + send_continuation_message( + scan_params, + statement_period, + remaining_accounts, + last_evaluated_key, + "batch_continuation", + sqs_endpoint, + continuation_queue_url, + aws_region, + logger, + ) + break + + try: + logger.info( + f"Processing batch {i + 1}/{len(account_batches)} with {len(batch)} accounts" + ) + + batch_result = process_account_batch( + batch, + statement_period, + sfn_client, + logger, + state_machine_arn, + sqs_endpoint, + dlq_url, + aws_region, + ) + + for key, value in batch_result.items(): + metrics_key = f"{key}_count" + if metrics_key in metrics: + metrics[metrics_key] += value + + metrics["batches_processed"] += 1 + + except Exception as e: + logger.error(f"Error processing batch {i + 1}: {e}") + if dlq_url and aws_region: + for account in batch: + send_bad_account_to_dlq( + account, + statement_period, + f"Batch processing exception: {str(e)}", + sqs_endpoint, + dlq_url, + aws_region, + logger, + ) + metrics["failed_starts_count"] += len(batch) + + return metrics + + +def process_accounts_scan_continuation( + scan_params, + statement_period, + context, + logger, + accounts_table, + sfn_client, + state_machine_arn, + sqs_endpoint, + continuation_queue_url, + aws_region, + page_size=50, + batch_size=10, + safety_buffer=30, + dlq_url=None, +): + metrics = initialize_metrics() + + logger.info(f"Continuing scan for period: {statement_period}") + + while True: + remaining_time = context.get_remaining_time_in_millis() / 1000.0 + if remaining_time < safety_buffer: + logger.warning("Approaching timeout, sending continuation message") + send_continuation_message( + scan_params, + statement_period, + None, + scan_params.get("ExclusiveStartKey"), + "accounts_scan", + sqs_endpoint, + continuation_queue_url, + aws_region, + logger, + ) + break + + accounts_page, last_evaluated_key = get_paginated_table_data( + scan_params=scan_params, + index_name=None, + table=accounts_table, + logger=logger, + page_size=page_size, + ) + + metrics["pages_processed"] += 1 + + if not accounts_page: + logger.info("No more accounts to process") + break + + page_metrics = process_accounts_page( + accounts_page, + statement_period, + context, + logger, + sfn_client, + state_machine_arn, + scan_params, + last_evaluated_key, + sqs_endpoint, + continuation_queue_url, + aws_region, + batch_size, + safety_buffer, + dlq_url, + ) + merge_metrics(metrics, page_metrics) + + if last_evaluated_key: + scan_params["ExclusiveStartKey"] = last_evaluated_key + else: + logger.info("All pages processed successfully") + break + + return metrics + + +def process_batch_continuation( + scan_params, + statement_period, + remaining_accounts, + last_evaluated_key, + context, + logger, + accounts_table, + sfn_client, + state_machine_arn, + sqs_endpoint, + continuation_queue_url, + aws_region, + page_size=50, + batch_size=10, + safety_buffer=30, + dlq_url=None, +): + metrics = initialize_metrics() + + logger.info(f"Processing {len(remaining_accounts)} remaining accounts") + + if remaining_accounts: + remaining_batches = list( + chunk_accounts(remaining_accounts, chunk_size=batch_size) + ) + batch_metrics = process_account_batches( + remaining_batches, + statement_period, + context, + logger, + sfn_client, + state_machine_arn, + scan_params, + last_evaluated_key, + sqs_endpoint, + continuation_queue_url, + aws_region, + safety_buffer, + dlq_url, + ) + merge_metrics(metrics, batch_metrics) + + if last_evaluated_key: + scan_params["ExclusiveStartKey"] = last_evaluated_key + scan_metrics = process_accounts_scan_continuation( + scan_params, + statement_period, + context, + logger, + accounts_table, + sfn_client, + state_machine_arn, + sqs_endpoint, + continuation_queue_url, + aws_region, + page_size, + batch_size, + safety_buffer, + dlq_url, + ) + merge_metrics(metrics, scan_metrics) + + return metrics diff --git a/layers/python/monthly_reports/monthly_reports/responses.py b/layers/python/monthly_reports/monthly_reports/responses.py new file mode 100644 index 0000000..64abfc7 --- /dev/null +++ b/layers/python/monthly_reports/monthly_reports/responses.py @@ -0,0 +1,35 @@ +from aws_lambda_powertools import Logger + + +def create_response(metrics, status, logger: Logger): + total_accounts_processed = ( + metrics["processed_count"] + + metrics["failed_starts_count"] + + metrics["skipped_count"] + + metrics["already_exists_count"] + ) + + logger.info( + f"Processing finished with status: {status}. " + f"Processed {total_accounts_processed} accounts. " + f"Metrics: {metrics}" + ) + + status_code_map = { + "ERROR_NO_CONTINUATION_QUEUE": 500, + "CRITICAL_ERROR": 500, + "TIMEOUT_CONTINUATION": 202, + "COMPLETED": 200, + } + + status_code = status_code_map.get(status, 500) + + return { + "statusCode": status_code, + "body": { + "message": f"Monthly Account reports processing {status.lower()}", + "status": status, + "totalAccountsProcessed": total_accounts_processed, + **metrics, + }, + } diff --git a/layers/python/monthly_reports/monthly_reports/sfn.py b/layers/python/monthly_reports/monthly_reports/sfn.py new file mode 100644 index 0000000..9703d63 --- /dev/null +++ b/layers/python/monthly_reports/monthly_reports/sfn.py @@ -0,0 +1,47 @@ +import json +import random +import time + +from botocore.exceptions import ClientError + + +def start_sfn_execution_with_retry( + sfn_client, state_machine_arn, execution_name, sf_input, logger, max_retries=3 +): + for attempt in range(max_retries): + try: + sfn_client.start_execution( + stateMachineArn=state_machine_arn, + name=execution_name, + input=json.dumps(sf_input), + ) + return "processed" + except ClientError as e: + error_code = e.response["Error"]["Code"] + + if error_code == "ExecutionAlreadyExistsException": + logger.info(f"SF execution {execution_name} already exists. Skipping.") + return "already_exists" + + if error_code in [ + "ThrottlingException", + "ServiceUnavailable", + "InternalFailure", + ]: + if attempt < max_retries - 1: + wait_time = (2**attempt) + random.uniform(0, 1) + logger.warning( + f"Retrying SF execution {execution_name} after {wait_time:.2f}s (attempt {attempt + 1}/{max_retries})" + ) + time.sleep(wait_time) + continue + else: + logger.error( + f"Max retries exceeded for SF execution {execution_name}: {e}" + ) + else: + logger.error( + f"Non-retryable error for SF execution {execution_name}: {e}" + ) + + raise e diff --git a/layers/python/monthly_reports/monthly_reports/sqs.py b/layers/python/monthly_reports/monthly_reports/sqs.py new file mode 100644 index 0000000..e337715 --- /dev/null +++ b/layers/python/monthly_reports/monthly_reports/sqs.py @@ -0,0 +1,97 @@ +import datetime +from typing import Optional, Dict, Any, List + +from aws_lambda_powertools import Logger + +from sqs import send_message_to_sqs + + +def send_continuation_message( + scan_params: Dict[str, Any], + statement_period: str, + remaining_accounts: Optional[List[Dict[str, Any]]], + last_evaluated_key: Optional[Dict[str, Any]], + continuation_type: str, + sqs_endpoint: str, + continuation_queue_url: str, + aws_region: str, + logger: Logger, +): + """Send continuation message to SQS""" + if not continuation_queue_url: + logger.error("Cannot send continuation message: CONTINUATION_QUEUE_URL not set") + return + + message_body: Dict[str, Any] = { + "scan_params": scan_params, + "statement_period": statement_period, + } + + if remaining_accounts: + message_body["remaining_accounts"] = remaining_accounts + if last_evaluated_key: + message_body["last_evaluated_key"] = last_evaluated_key + + message_attributes = { + "continuation_type": { + "DataType": "String", + "StringValue": continuation_type, + } + } + + send_message_to_sqs( + message=message_body, + message_attributes=message_attributes, + sqs_endpoint=sqs_endpoint, + sqs_url=continuation_queue_url, + aws_region=aws_region, + logger=logger, + ) + + +def send_bad_account_to_dlq( + account: Dict[str, Any], + statement_period: str, + error_reason: str, + sqs_endpoint: str, + dlq_url: str, + aws_region: str, + logger: Logger, +): + """Send bad account to DLQ""" + if not dlq_url: + logger.warning("Cannot send bad account to DLQ: DLQ_URL not set") + return + + message_body = { + "account": account, + "statement_period": statement_period, + "error_reason": error_reason, + "timestamp": datetime.datetime.now(datetime.UTC).isoformat(), + } + + message_attributes = { + "error_type": { + "DataType": "String", + "StringValue": "bad_account", + }, + "error_reason": { + "DataType": "String", + "StringValue": error_reason, + }, + } + + try: + send_message_to_sqs( + message=message_body, + message_attributes=message_attributes, + sqs_endpoint=sqs_endpoint, + sqs_url=dlq_url, + aws_region=aws_region, + logger=logger, + ) + logger.info( + f"Sent bad account to DLQ: {account.get('accountId', 'unknown')} - {error_reason}" + ) + except Exception as e: + logger.error(f"Failed to send bad account to DLQ: {e}") diff --git a/layers/python/monthly_reports/requirements.txt b/layers/python/monthly_reports/requirements.txt new file mode 100644 index 0000000..9add2d8 --- /dev/null +++ b/layers/python/monthly_reports/requirements.txt @@ -0,0 +1 @@ +aws_lambda_powertools==3.17.0 \ No newline at end of file diff --git a/pytest.ini b/pytest.ini index 3693aeb..75f8211 100644 --- a/pytest.ini +++ b/pytest.ini @@ -2,4 +2,4 @@ filterwarnings = ignore::DeprecationWarning:botocore.* testpaths = tests -pythonpath = layers/python/helpers layers/python/accounts layers/python/authentication \ No newline at end of file +pythonpath = layers/python/helpers layers/python/accounts layers/python/authentication layers/python/monthly_reports \ No newline at end of file diff --git a/template.yml b/template.yml index 5445c8e..ad215e4 100644 --- a/template.yml +++ b/template.yml @@ -38,6 +38,7 @@ Globals: ENVIRONMENT_NAME: !Ref Environment SES_ENABLED: !Ref SesEnabled SES_SENDER_EMAIL: !Sub '{{resolve:ssm:/banking-app/${Environment}/SesSenderEmail}}' + SES_NO_REPLY_EMAIL: !Sub '{{resolve:ssm:/banking-app/${Environment}/SesNoReplyEmail}}' SES_REPLY_EMAIL: !Sub '{{resolve:ssm:/banking-app/${Environment}/SesReplyEmail}}' SES_BOUNCE_EMAIL: !Sub '{{resolve:ssm:/banking-app/${Environment}/SESBounceEmail}}' @@ -187,7 +188,9 @@ Resources: Variables: ACCOUNTS_TABLE_NAME: !Ref AccountsTable TRANSACTIONS_TABLE_NAME: !Ref TransactionsTable + TRANSACTION_PROCESSING_DLQ_URL: !Ref TransactionProcessingDLQ DYNAMODB_ENDPOINT: '' + SQS_ENDPOINT: '' COGNITO_USER_POOL_ID: !Ref BankingUserPool Layers: - !Ref PythonHelpersLambdaLayer @@ -268,6 +271,34 @@ Resources: - DynamoDBCrudPolicy: TableName: !Ref AccountsTable + GetAccountTransactionsFunction: + Type: AWS::Serverless::Function + Properties: + FunctionName: !Sub ${AWS::StackName}-get-account-transactions + Description: !Sub + - Stack ${AWS::StackName} Function ${ResourceName} for Stage ${Environment} + - ResourceName: GetAccountTransactionsFunction + CodeUri: functions/accounts/get_account_transactions + Handler: get_account_transactions.app.lambda_handler + Environment: + Variables: + TRANSACTIONS_TABLE_NAME: !Ref TransactionsTable + DYNAMODB_ENDPOINT: '' + Layers: + - !Ref PythonHelpersLambdaLayer + Events: + ApiGetAccountTransactionsEvent: + Type: Api + Properties: + Path: /accounts/{account_id}/transactions + Method: GET + RestApiId: !Ref BankingApiGateway + Auth: + Authorizer: CognitoAuthorizer + Policies: + - DynamoDBReadPolicy: + TableName: !Ref TransactionsTable + AuthFunction: Type: AWS::Serverless::Function Properties: @@ -337,6 +368,176 @@ Resources: - ses:SendRawEmail Resource: '*' + MonthlyAccountReportsTriggerFunction: + Type: AWS::Serverless::Function + Properties: + FunctionName: !Sub ${AWS::StackName}-monthly-account-reports-trigger-function + Description: !Sub + - Stack ${AWS::StackName} Function ${ResourceName} for Stage ${Environment} + - ResourceName: MonthlyAccountReportsTriggerFunction + CodeUri: functions/monthly_reports/accounts/trigger/ + Handler: trigger.app.lambda_handler + Timeout: 300 + Environment: + Variables: + ACCOUNTS_TABLE_NAME: !Ref AccountsTable + STATE_MACHINE_ARN: !Ref StateMachine + CONTINUATION_QUEUE_URL: !Ref MonthlyAccountReportsContinuationQueue + DLQ_URL: !Ref MonthlyAccountReportsContinuationDLQ + DYNAMODB_ENDPOINT: '' + SQS_ENDPOINT: '' + Events: + MonthlySchedule: + Type: Schedule + Properties: + Name: !Sub ${AWS::StackName}-monthly-accounts-reports-trigger + Schedule: cron(0 0 1 * ? *) + Description: !Sub Monthly accounts reports trigger for ${Environment} + Layers: + - !Ref PythonHelpersLambdaLayer + - !Ref PythonMonthlyReportsLambdaLayer + ReservedConcurrentExecutions: 1 + Policies: + - DynamoDBCrudPolicy: + TableName: !Ref AccountsTable + - StepFunctionsExecutionPolicy: + StateMachineName: !GetAtt StateMachine.Name + - Statement: + Effect: Allow + Action: + - sqs:SendMessage + - sqs:GetQueueAttributes + - sqs:GetQueueUrl + Resource: !GetAtt MonthlyAccountReportsContinuationQueue.Arn + - Statement: + Effect: Allow + Action: + - sqs:SendMessage + - sqs:GetQueueAttributes + - sqs:GetQueueUrl + Resource: !GetAtt MonthlyAccountReportsContinuationDLQ.Arn + + MonthlyAccountReportsProcessPendingReportsFunction: + Type: AWS::Serverless::Function + Properties: + FunctionName: !Sub ${AWS::StackName}-monthly-accounts-reports-continuation + Description: !Sub + - Stack ${AWS::StackName} Function ${ResourceName} for Stage ${Environment} + - ResourceName: MonthlyAccountReportsProcessPendingReportsFunction + CodeUri: functions/monthly_reports/accounts/process_pending_reports/ + Handler: process_pending_reports.app.lambda_handler + Timeout: 150 + Environment: + Variables: + ACCOUNTS_TABLE_NAME: !Ref AccountsTable + STATE_MACHINE_ARN: !Ref StateMachine + CONTINUATION_QUEUE_URL: !Ref MonthlyAccountReportsContinuationQueue + DLQ_URL: !Ref MonthlyAccountReportsContinuationDLQ + DYNAMODB_ENDPOINT: '' + SQS_ENDPOINT: '' + Events: + SQSEvent: + Type: SQS + Properties: + Queue: !GetAtt MonthlyAccountReportsContinuationQueue.Arn + BatchSize: 1 + Enabled: true + Layers: + - !Ref PythonHelpersLambdaLayer + - !Ref PythonMonthlyReportsLambdaLayer + Policies: + - DynamoDBCrudPolicy: + TableName: !Ref AccountsTable + - StepFunctionsExecutionPolicy: + StateMachineName: !GetAtt StateMachine.Name + - Statement: + Effect: Allow + Action: + - sqs:ReceiveMessage + - sqs:DeleteMessage + - sqs:GetQueueAttributes + - sqs:GetQueueUrl + Resource: !GetAtt MonthlyAccountReportsContinuationQueue.Arn + - Statement: + Effect: Allow + Action: + - sqs:SendMessage + - sqs:GetQueueAttributes + - sqs:GetQueueUrl + Resource: !GetAtt MonthlyAccountReportsContinuationDLQ.Arn + + CreateAccountsReportFunction: + Type: AWS::Serverless::Function + Properties: + FunctionName: !Sub ${AWS::StackName}-create-accounts-report + Description: !Sub + - Stack ${AWS::StackName} Function ${ResourceName} for Stage ${Environment} + - ResourceName: CreateAccountsReportFunction + CodeUri: functions/monthly_reports/accounts/create_report + Handler: create_report.app.lambda_handler + Environment: + Variables: + REPORTS_BUCKET: !Ref MonthlyReportsBucket + Layers: + - !Ref PythonHelpersLambdaLayer + Policies: + - S3WritePolicy: + BucketName: !Ref MonthlyReportsBucket + - Statement: + Effect: Allow + Action: + - s3:GetObject + Resource: !Sub ${MonthlyReportsBucket.Arn}/* + + AccountsReportsNotifyClientFunction: + Type: AWS::Serverless::Function + Properties: + FunctionName: !Sub ${AWS::StackName}-account-reports-notify-client + Description: !Sub + - Stack ${AWS::StackName} Function ${ResourceName} for Stage ${Environment} + - ResourceName: AccountsReportsNotifyClientFunction + CodeUri: functions/monthly_reports/accounts/notify_client + Handler: notify_client.app.lambda_handler + Environment: + Variables: + REPORTS_BUCKET: !Ref MonthlyReportsBucket + COGNITO_USER_POOL_ID: !Ref BankingUserPool + COGNITO_CLIENT_ID: !Ref BankingUserPoolClient + ACCOUNTS_TABLE_NAME: !Ref AccountsTable + DYNAMODB_ENDPOINT: '' + Events: + ApiRequestNewMonthlyAccountReportEvent: + Type: Api + Properties: + Path: /accounts/{account_id}/reports/{statement_period} + Method: GET + RestApiId: !Ref BankingApiGateway + Auth: + Authorizer: CognitoAuthorizer + Layers: + - !Ref PythonHelpersLambdaLayer + - !Ref PythonAuthenticationLambdaLayer + - !Ref PythonAccountsLambdaLayer + Policies: + - Statement: + Effect: Allow + Action: + - s3:GetObject + Resource: !Sub ${MonthlyReportsBucket.Arn}/* + - Statement: + Effect: Allow + Action: + - cognito-idp:AdminGetUser + Resource: !GetAtt BankingUserPool.Arn + - DynamoDBReadPolicy: + TableName: !Ref AccountsTable + - Statement: + - Effect: Allow + Action: + - ses:SendEmail + - ses:SendRawEmail + Resource: '*' + # --- Lambda Layers --- PythonHelpersLambdaLayer: Type: AWS::Serverless::LayerVersion @@ -371,6 +572,17 @@ Resources: BuildMethod: python3.12 BuildArchitecture: x86_64 + PythonMonthlyReportsLambdaLayer: + Type: AWS::Serverless::LayerVersion + Properties: + LayerName: !Sub ${AWS::StackName}-banking-python-monthly-reports-layer + ContentUri: layers/python/monthly_reports + CompatibleRuntimes: + - python3.12 + Metadata: + BuildMethod: python3.12 + BuildArchitecture: x86_64 + # --- Log Groups --- RequestTransactionFunctionLogGroup: Type: AWS::Logs::LogGroup @@ -412,6 +624,55 @@ Resources: LogGroupName: !Sub /aws/lambda/${CognitoPostSignUpFunction} RetentionInDays: 7 + GetAccountsFunctionLogGroup: + Type: AWS::Logs::LogGroup + DeletionPolicy: Delete + UpdateReplacePolicy: Delete + Properties: + LogGroupName: !Sub /aws/lambda/${GetAccountsFunction} + RetentionInDays: 7 + + MonthlyAccountReportsTriggerFunctionLogGroup: + Type: AWS::Logs::LogGroup + DeletionPolicy: Delete + UpdateReplacePolicy: Delete + Properties: + LogGroupName: !Sub /aws/lambda/${MonthlyAccountReportsTriggerFunction} + RetentionInDays: 7 + + MonthlyAccountReportsProcessPendingReportsFunctionLogGroup: + Type: AWS::Logs::LogGroup + DeletionPolicy: Delete + UpdateReplacePolicy: Delete + Properties: + LogGroupName: !Sub /aws/lambda/${MonthlyAccountReportsProcessPendingReportsFunction} + RetentionInDays: 7 + + GetAccountTransactionsFunctionLogGroup: + Type: AWS::Logs::LogGroup + DeletionPolicy: Delete + UpdateReplacePolicy: Delete + Properties: + LogGroupName: !Sub /aws/lambda/${GetAccountTransactionsFunction} + RetentionInDays: 7 + + CreateAccountsReportFunctionLogGroup: + Type: AWS::Logs::LogGroup + DeletionPolicy: Delete + UpdateReplacePolicy: Delete + Properties: + LogGroupName: !Sub /aws/lambda/${CreateAccountsReportFunction} + RetentionInDays: 7 + + AccountsReportsNotifyClientLogGroup: + Type: AWS::Logs::LogGroup + DeletionPolicy: Delete + UpdateReplacePolicy: Delete + Properties: + LogGroupName: !Sub /aws/lambda/${AccountsReportsNotifyClientFunction} + RetentionInDays: 7 + + # --- DynamoDB Table for Transactions --- TransactionsTable: Type: AWS::DynamoDB::Table @@ -427,6 +688,10 @@ Resources: AttributeType: S - AttributeName: userId AttributeType: S + - AttributeName: accountId + AttributeType: S + - AttributeName: createdAt + AttributeType: S BillingMode: PAY_PER_REQUEST PointInTimeRecoverySpecification: PointInTimeRecoveryEnabled: true @@ -445,6 +710,14 @@ Resources: KeyType: HASH Projection: ProjectionType: ALL + - IndexName: AccountDateIndex + KeySchema: + - AttributeName: accountId + KeyType: HASH + - AttributeName: createdAt + KeyType: RANGE + Projection: + ProjectionType: ALL TimeToLiveSpecification: AttributeName: ttlTimestamp Enabled: true @@ -481,6 +754,27 @@ Resources: Enabled: true # --- SQS Queues --- + MonthlyAccountReportsContinuationQueue: + Type: AWS::SQS::Queue + Properties: + QueueName: !Sub ${AWS::StackName}-monthly-account-reports-continuation-queue + VisibilityTimeout: 200 + MessageRetentionPeriod: 1209600 # 14 days + ReceiveMessageWaitTimeSeconds: 20 + RedrivePolicy: + deadLetterTargetArn: !GetAtt MonthlyAccountReportsContinuationDLQ.Arn + maxReceiveCount: 3 + + MonthlyAccountReportsContinuationDLQ: + Type: AWS::SQS::Queue + Properties: + QueueName: !Sub ${AWS::StackName}-monthly-account-reports-dlq + + MonthlyAccountReportsCreationDLQ: + Type: AWS::SQS::Queue + Properties: + QueueName: !Sub ${AWS::StackName}-monthly-account-reports-creation-dlq + TransactionProcessingDLQ: Type: AWS::SQS::Queue Properties: @@ -495,6 +789,28 @@ Resources: Principal: cognito-idp.amazonaws.com SourceArn: !GetAtt BankingUserPool.Arn + # S3 Buckets + MonthlyReportsBucket: + Type: AWS::S3::Bucket + Properties: + BucketName: !Sub ${AWS::StackName}-monthly-reports-bucket + VersioningConfiguration: + Status: Enabled + PublicAccessBlockConfiguration: + BlockPublicAcls: true + BlockPublicPolicy: true + IgnorePublicAcls: true + RestrictPublicBuckets: true + LifecycleConfiguration: + Rules: + - Id: TransitionToIA + Status: Enabled + Transitions: + - TransitionInDays: 90 + StorageClass: STANDARD_IA + - TransitionInDays: 365 + StorageClass: GLACIER + # Domain Configuration ApiCertificate: Type: AWS::CertificateManager::Certificate @@ -504,3 +820,71 @@ Resources: DomainValidationOptions: - DomainName: !Sub '{{resolve:ssm:/banking-app/${Environment}/DomainName}}' HostedZoneId: !Sub '{{resolve:ssm:/banking-app/${Environment}/Route53HostedZoneId}}' + + # State machines + StateMachine: + Type: AWS::Serverless::StateMachine + Properties: + Name: !Sub ${AWS::StackName}-monthly-reports-creation-state-machine + Type: STANDARD + Definition: + StartAt: GetAccountTransactions + States: + GetAccountTransactions: + Type: Task + Resource: !GetAtt GetAccountTransactionsFunction.Arn + Next: CreateAccountsReport + Retry: + - ErrorEquals: [ "States.ALL" ] + IntervalSeconds: 5 + MaxAttempts: 3 + BackoffRate: 2.0 + Catch: + - ErrorEquals: [ "States.ALL" ] + ResultPath: $.error + Next: SendToDLQ + + CreateAccountsReport: + Type: Task + Resource: !GetAtt CreateAccountsReportFunction.Arn + Next: NotifyClient + Retry: + - ErrorEquals: [ "States.ALL" ] + IntervalSeconds: 5 + MaxAttempts: 3 + BackoffRate: 2.0 + Catch: + - ErrorEquals: [ "States.ALL" ] + ResultPath: $.error + Next: SendToDLQ + + NotifyClient: + Type: Task + Resource: !GetAtt AccountsReportsNotifyClientFunction.Arn + End: true + Retry: + - ErrorEquals: [ "States.ALL" ] + IntervalSeconds: 10 + MaxAttempts: 2 + BackoffRate: 2.0 + Catch: + - ErrorEquals: [ "States.ALL" ] + ResultPath: $.error + Next: SendToDLQ + + SendToDLQ: + Type: Task + Resource: arn:aws:states:::sqs:sendMessage + Parameters: + QueueUrl: !Ref MonthlyAccountReportsCreationDLQ + MessageBody.$: $ + End: true + Policies: + - LambdaInvokePolicy: + FunctionName: !Ref GetAccountTransactionsFunction + - LambdaInvokePolicy: + FunctionName: !Ref CreateAccountsReportFunction + - LambdaInvokePolicy: + FunctionName: !Ref AccountsReportsNotifyClientFunction + - SQSSendMessagePolicy: + QueueName: !GetAtt MonthlyAccountReportsCreationDLQ.QueueName \ No newline at end of file diff --git a/tests/conftest.py b/tests/conftest.py index f55ea3f..6dd7b5d 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -254,6 +254,39 @@ def mock_sqs_client(): yield client +@pytest.fixture +def mock_sfn_client(): + with mock_aws(): + client = boto3.client("stepfunctions", region_name=AWS_REGION) + + yield client + + +@pytest.fixture +def magic_mock_ses_client(): + mock_client = MagicMock() + return mock_client + + +@pytest.fixture +def mock_s3_client(): + """Mock S3 client for testing.""" + mock_client = MagicMock() + return mock_client + + +@pytest.fixture +def mock_cognito_client(): + """Mock Cognito client for testing.""" + mock_client = MagicMock() + return mock_client + + +@pytest.fixture(scope="function") +def magic_mock_sfn_client(): + return MagicMock() + + @pytest.fixture def mock_context(): """ diff --git a/tests/functions/accounts/get_account_transactions/__init__.py b/tests/functions/accounts/get_account_transactions/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/functions/accounts/get_account_transactions/conftest.py b/tests/functions/accounts/get_account_transactions/conftest.py new file mode 100644 index 0000000..5de061e --- /dev/null +++ b/tests/functions/accounts/get_account_transactions/conftest.py @@ -0,0 +1,70 @@ +import uuid +from importlib import reload +from unittest.mock import patch + +import pytest + +from functions.accounts.get_account_transactions.get_account_transactions import app + +VALID_UUID = str(uuid.uuid4()) + + +@pytest.fixture +def valid_get_transactions_event(): + """ + Return a dictionary representing a valid HTTP GET event for retrieving account transactions. + + The returned event includes the HTTP method, endpoint path, authorisation header, and a unique request ID in the request context. + """ + account_id = VALID_UUID + return { + "httpMethod": "GET", + "path": f"/accounts/{account_id}/transactions", + "pathParameters": {"account_id": account_id}, + "headers": { + "Authorization": "Bearer valid-token", + }, + "requestContext": { + "requestId": str(uuid.uuid4()), + }, + } + + +@pytest.fixture +def step_functions_event(): + """ + Return a dictionary representing a Step Functions event for retrieving account transactions. + + The returned event includes the accountId field that Step Functions would pass. + """ + return { + "accountId": VALID_UUID, + "someOtherData": "value", + } + + +@pytest.fixture(scope="function") +def get_account_transactions_app_with_mocked_tables( + monkeypatch, + dynamo_resource, + mock_transactions_dynamo_table, +): + """ + Pytest fixture that configures the get_account_transactions app with mocked DynamoDB tables and environment variables for isolated testing. + + Yields: + The app instance with its transactions table and environment fully mocked for test execution. + """ + transactions_table_name = mock_transactions_dynamo_table + + monkeypatch.setenv("TRANSACTIONS_TABLE_NAME", transactions_table_name) + monkeypatch.setenv("ENVIRONMENT_NAME", "test") + monkeypatch.setenv("POWERTOOLS_LOG_LEVEL", "INFO") + monkeypatch.setenv("AWS_REGION", "eu-west-2") + + with patch("boto3.resource", return_value=dynamo_resource): + reload(app) + + app.table = dynamo_resource.Table(transactions_table_name) + + yield app diff --git a/tests/functions/accounts/get_account_transactions/test_app.py b/tests/functions/accounts/get_account_transactions/test_app.py new file mode 100644 index 0000000..8929384 --- /dev/null +++ b/tests/functions/accounts/get_account_transactions/test_app.py @@ -0,0 +1,251 @@ +import json +import uuid +from unittest.mock import patch + +import pytest +from aws_lambda_powertools.event_handler.exceptions import ( + InternalServerError, +) +from botocore.exceptions import ClientError + +from functions.accounts.get_account_transactions.get_account_transactions.app import ( + lambda_handler, +) +from functions.accounts.get_account_transactions.get_account_transactions.exceptions import ( + ValidationError, +) + + +class TestGetAccountTransactionsAPI: + + def test_get_account_transactions_success( + self, valid_get_transactions_event, mock_context + ): + account_id = valid_get_transactions_event["pathParameters"]["account_id"] + + with patch( + "functions.accounts.get_account_transactions.get_account_transactions.app.table" + ) as mock_table: + mock_table.query.return_value = { + "Items": [ + { + "id": str(uuid.uuid4()), + "accountId": account_id, + "amount": "100.50", + "type": "DEPOSIT", + "description": "Test transaction", + "status": "COMPLETED", + "createdAt": "2023-01-01T12:00:00Z", + } + ] + } + + response = lambda_handler(valid_get_transactions_event, mock_context) + response_body = json.loads(response["body"]) + + assert response["statusCode"] == 200 + assert "transactions" in response_body + assert len(response_body["transactions"]) == 1 + assert response_body["transactions"][0]["accountId"] == account_id + + def test_get_account_transactions_with_period_param( + self, valid_get_transactions_event, mock_context + ): + account_id = valid_get_transactions_event["pathParameters"]["account_id"] + valid_get_transactions_event["queryStringParameters"] = {"period": "2023-01"} + + with patch( + "functions.accounts.get_account_transactions.get_account_transactions.app.table" + ) as mock_table: + mock_table.query.return_value = { + "Items": [ + { + "id": str(uuid.uuid4()), + "accountId": account_id, + "amount": "100.50", + "type": "DEPOSIT", + "description": "Test transaction", + "status": "COMPLETED", + "createdAt": "2023-01-15T12:00:00Z", + } + ] + } + + response = lambda_handler(valid_get_transactions_event, mock_context) + response_body = json.loads(response["body"]) + + assert response["statusCode"] == 200 + assert "transactions" in response_body + assert len(response_body["transactions"]) == 1 + + def test_get_account_transactions_with_date_range_params( + self, valid_get_transactions_event, mock_context + ): + account_id = valid_get_transactions_event["pathParameters"]["account_id"] + valid_get_transactions_event["queryStringParameters"] = { + "start": "2023-01-01", + "end": "2023-01-31", + } + + with patch( + "functions.accounts.get_account_transactions.get_account_transactions.app.table" + ) as mock_table: + mock_table.query.return_value = { + "Items": [ + { + "id": str(uuid.uuid4()), + "accountId": account_id, + "amount": "100.50", + "type": "DEPOSIT", + "description": "Test transaction", + "status": "COMPLETED", + "createdAt": "2023-01-15T12:00:00Z", + } + ] + } + + response = lambda_handler(valid_get_transactions_event, mock_context) + response_body = json.loads(response["body"]) + + assert response["statusCode"] == 200 + assert "transactions" in response_body + assert len(response_body["transactions"]) == 1 + + def test_get_account_transactions_validation_error( + self, valid_get_transactions_event, mock_context + ): + with patch( + "functions.accounts.get_account_transactions.get_account_transactions.app.table" + ) as mock_table: + mock_table.query.side_effect = ValidationError("Invalid date range") + + response = lambda_handler(valid_get_transactions_event, mock_context) + + assert response["statusCode"] == 400 + response_body = json.loads(response["body"]) + assert "Invalid date range" in response_body["message"] + + def test_get_account_transactions_client_error( + self, valid_get_transactions_event, mock_context + ): + """Test handling of DynamoDB client errors""" + error_response = { + "Error": {"Code": "InternalServerError", "Message": "Internal server error"} + } + with patch( + "functions.accounts.get_account_transactions.get_account_transactions.app.table" + ) as mock_table: + mock_table.query.side_effect = ClientError(error_response, "Query") + + response = lambda_handler(valid_get_transactions_event, mock_context) + + assert response["statusCode"] == 500 + response_body = json.loads(response["body"]) + assert "Internal server error" in response_body["message"] + + def test_get_account_transactions_general_exception( + self, valid_get_transactions_event, mock_context + ): + with patch( + "functions.accounts.get_account_transactions.get_account_transactions.app.table" + ) as mock_table: + mock_table.query.side_effect = Exception("Unexpected error") + + response = lambda_handler(valid_get_transactions_event, mock_context) + + assert response["statusCode"] == 500 + response_body = json.loads(response["body"]) + assert "Internal server error" in response_body["message"] + + +class TestGetAccountTransactionsStepFunctions: + + def test_step_functions_request_success(self, step_functions_event, mock_context): + account_id = step_functions_event["accountId"] + + with patch( + "functions.accounts.get_account_transactions.get_account_transactions.app.table" + ) as mock_table: + mock_table.query.return_value = { + "Items": [ + { + "id": str(uuid.uuid4()), + "accountId": account_id, + "amount": "100.50", + "type": "DEPOSIT", + "description": "Test transaction", + "status": "COMPLETED", + "createdAt": "2023-01-01T12:00:00Z", + } + ] + } + + response = lambda_handler(step_functions_event, mock_context) + + assert "transactions" in response + assert len(response["transactions"]) == 1 + assert response["transactions"][0]["accountId"] == account_id + assert response["accountId"] == account_id + + def test_step_functions_request_missing_account_id(self, mock_context): + event = {"someOtherField": "value"} + + with patch( + "functions.accounts.get_account_transactions.get_account_transactions.app.table" + ) as mock_table: + mock_table.query.return_value = {"Items": []} + + response = lambda_handler(event, mock_context) + + assert response["statusCode"] == 400 + response_body = json.loads(response["body"]) + assert "Missing accountId" in response_body["error"] + + def test_step_functions_request_exception(self, step_functions_event, mock_context): + with patch( + "functions.accounts.get_account_transactions.get_account_transactions.app.table" + ) as mock_table: + mock_table.query.side_effect = Exception("Database error") + + response = lambda_handler(step_functions_event, mock_context) + + assert response["statusCode"] == 500 + response_body = json.loads(response["body"]) + assert "Database error" in response_body["error"] + + +class TestLambdaHandlerConfiguration: + + def test_lambda_handler_missing_table_configuration(self, mock_context): + with patch( + "functions.accounts.get_account_transactions.get_account_transactions.app.table", + None, + ): + event = {"httpMethod": "GET", "path": "/accounts/123/transactions"} + + with pytest.raises(InternalServerError) as exc_info: + lambda_handler(event, mock_context) + + assert "Server configuration error" in str(exc_info.value) + + def test_transactions_table_not_initialized( + self, mock_context, step_functions_event + ): + with patch( + "functions.accounts.get_account_transactions.get_account_transactions.app.table", + None, + ): + with pytest.raises(InternalServerError) as exc_info: + lambda_handler(step_functions_event, mock_context) + + assert str(exc_info.value) == "Server configuration error" + + def test_table_initialization_with_environment_variables( + self, get_account_transactions_app_with_mocked_tables + ): + assert get_account_transactions_app_with_mocked_tables.table is not None + + assert ( + get_account_transactions_app_with_mocked_tables.TRANSACTIONS_TABLE_NAME + == "test-transactions-table" + ) diff --git a/tests/functions/accounts/get_account_transactions/test_date_helpers.py b/tests/functions/accounts/get_account_transactions/test_date_helpers.py new file mode 100644 index 0000000..81aece3 --- /dev/null +++ b/tests/functions/accounts/get_account_transactions/test_date_helpers.py @@ -0,0 +1,82 @@ +import pytest +from functions.accounts.get_account_transactions.get_account_transactions.date_helpers import ( + get_date_range, +) +from functions.accounts.get_account_transactions.get_account_transactions.exceptions import ( + ValidationError, +) +from datetime import datetime, timezone + + +class TestGetDateRange: + def test_period_and_start(self): + with pytest.raises( + ValidationError, match="Cannot combine 'period' with 'start'/'end'" + ): + get_date_range(period="june", start="june") + + def test_end_and_no_start(self): + with pytest.raises( + ValidationError, match="Both 'start' and 'end' must be provided together" + ): + get_date_range(end="june") + + def test_invalid_custom_date_format(self): + with pytest.raises( + ValidationError, match="Invalid date format, must be YYYY-MM-DD" + ): + get_date_range(start="2024/01/01", end="2024-01-10") + + def test_end_before_start(self): + with pytest.raises( + ValidationError, match="'end' date must be after 'start' date" + ): + get_date_range(start="2024-06-10", end="2024-06-09") + + def test_valid_custom_range_outputs(self): + period, start_iso, end_iso = get_date_range( + start="2024-06-10", end="2024-06-12" + ) + assert period == "2024-06-10_to_2024-06-12" + assert start_iso == "2024-06-10T00:00:00Z" + assert end_iso == "2024-06-12T23:59:59Z" + + @pytest.mark.parametrize( + "bad_period", + ["june", "2024-13", "2024/06", "202406", "2024--06", "2024-0a"], + ) + def test_invalid_period_format(self, bad_period): + with pytest.raises( + ValidationError, match="Invalid period format, must be YYYY-MM" + ): + get_date_range(period=bad_period) + + def test_valid_period_leap_february(self): + period, start_iso, end_iso = get_date_range(period="2024-02") + assert period == "2024-02" + assert start_iso == "2024-02-01T00:00:00Z" + assert end_iso == "2024-02-29T23:59:59Z" # leap year + + def test_valid_period_april(self): + period, start_iso, end_iso = get_date_range(period="2023-04") + assert period == "2023-04" + assert start_iso == "2023-04-01T00:00:00Z" + assert end_iso == "2023-04-30T23:59:59Z" + + def test_base_case(self, monkeypatch): + fake_now = datetime(2023, 3, 5, tzinfo=timezone.utc) + + class FixedDateTime(datetime): + @classmethod + def now(cls, tz=None): + return fake_now + + monkeypatch.setattr( + "functions.accounts.get_account_transactions.get_account_transactions.date_helpers.datetime", + FixedDateTime, + ) + + period, start_iso, end_iso = get_date_range() + assert period == "2023-02" + assert start_iso == "2023-02-01T00:00:00Z" + assert end_iso == "2023-02-28T23:59:59Z" diff --git a/tests/functions/accounts/get_account_transactions/test_transaction_helpers.py b/tests/functions/accounts/get_account_transactions/test_transaction_helpers.py new file mode 100644 index 0000000..34fd506 --- /dev/null +++ b/tests/functions/accounts/get_account_transactions/test_transaction_helpers.py @@ -0,0 +1,82 @@ +from unittest.mock import MagicMock, patch + +from functions.accounts.get_account_transactions.get_account_transactions.transaction_helpers import ( + query_transactions, +) + + +class TestQueryTransactions: + @patch( + "functions.accounts.get_account_transactions.get_account_transactions.date_helpers.get_date_range" + ) + def test_query_transactions_success(self, mock_get_date_range): + mock_get_date_range.return_value = ( + "2024-06", + "2024-06-01T00:00:00Z", + "2024-06-30T23:59:59Z", + ) + + mock_logger = MagicMock() + mock_table = MagicMock() + mock_table.query.return_value = {"Items": [{"id": "txn1"}, {"id": "txn2"}]} + + result = query_transactions( + table=mock_table, + account_id="acc123", + logger=mock_logger, + period="2024-06", + ) + + mock_get_date_range.assert_called_once_with("2024-06", None, None) + mock_logger.info.assert_called_once() + mock_table.query.assert_called_once() + assert result["statementPeriod"] == "2024-06" + assert result["transactions"] == [{"id": "txn1"}, {"id": "txn2"}] + + @patch( + "functions.accounts.get_account_transactions.get_account_transactions.date_helpers.get_date_range" + ) + def test_query_transactions_descending(self, mock_get_date_range): + mock_get_date_range.return_value = ( + "2024-07", + "2024-07-01T00:00:00Z", + "2024-07-31T23:59:59Z", + ) + + mock_logger = MagicMock() + mock_table = MagicMock() + mock_table.query.return_value = {"Items": []} + + result = query_transactions( + table=mock_table, + account_id="acc456", + logger=mock_logger, + period="2024-07", + descending=True, + ) + + assert mock_table.query.call_args[1]["ScanIndexForward"] is False + assert result["transactions"] == [] + + @patch( + "functions.accounts.get_account_transactions.get_account_transactions.date_helpers.get_date_range" + ) + def test_query_transactions_no_items(self, mock_get_date_range): + mock_get_date_range.return_value = ( + "2024-08", + "2024-08-01T00:00:00Z", + "2024-08-31T23:59:59Z", + ) + + mock_logger = MagicMock() + mock_table = MagicMock() + mock_table.query.return_value = {} # no Items key + + result = query_transactions( + table=mock_table, + account_id="acc789", + logger=mock_logger, + period="2024-08", + ) + + assert result["transactions"] == [] diff --git a/tests/functions/monthly_reports/__init__.py b/tests/functions/monthly_reports/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/functions/monthly_reports/accounts/__init__.py b/tests/functions/monthly_reports/accounts/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/functions/monthly_reports/accounts/create_report/__init__.py b/tests/functions/monthly_reports/accounts/create_report/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/functions/monthly_reports/accounts/create_report/conftest.py b/tests/functions/monthly_reports/accounts/create_report/conftest.py new file mode 100644 index 0000000..180dc9a --- /dev/null +++ b/tests/functions/monthly_reports/accounts/create_report/conftest.py @@ -0,0 +1,72 @@ +from importlib import reload +from unittest.mock import MagicMock +import pytest +from functions.monthly_reports.accounts.create_report.create_report import app + + +@pytest.fixture(scope="function") +def create_report_app_with_mocks(monkeypatch, mock_s3_client): + """Fixture that sets up the create_report app with mocked dependencies.""" + + # Set environment variables + monkeypatch.setenv("REPORTS_BUCKET", "test-reports-bucket") + monkeypatch.setenv("POWERTOOLS_LOG_LEVEL", "INFO") + monkeypatch.setenv("AWS_REGION", "eu-west-2") + + # Mock the S3 client methods + mock_s3_client.put_object.return_value = { + "ResponseMetadata": {"HTTPStatusCode": 200} + } + mock_s3_client.generate_presigned_url.return_value = "https://test-reports-bucket.s3.eu-west-2.amazonaws.com/test-account-123/2024-01.pdf?AWSAccessKeyId=test&Signature=test&Expires=1234567890" + + # Load the app module + reload(app) + + # Replace the S3 client in the app module directly + app.s3 = mock_s3_client + + yield app + + +@pytest.fixture +def mock_s3_client(): + """Mock S3 client for testing.""" + mock_client = MagicMock() + return mock_client + + +@pytest.fixture +def sample_event(): + """Sample event data for testing.""" + return { + "accountId": "test-account-123", + "userId": "test-user-456", + "statementPeriod": "2024-01", + "transactions": [ + { + "id": "txn-1", + "amount": 100.00, + "description": "Test transaction 1", + "date": "2024-01-15", + }, + { + "id": "txn-2", + "amount": -50.00, + "description": "Test transaction 2", + "date": "2024-01-20", + }, + ], + "accountBalance": 1500.00, + } + + +@pytest.fixture +def mock_pdf_bytes(): + """Mock PDF bytes for testing.""" + return b"%PDF-1.4\n%Test PDF content\n%%EOF" + + +@pytest.fixture +def mock_presigned_url(): + """Mock presigned URL for testing.""" + return "https://test-reports-bucket.s3.eu-west-2.amazonaws.com/test-account-123/2024-01.pdf?AWSAccessKeyId=test&Signature=test&Expires=1234567890" diff --git a/tests/functions/monthly_reports/accounts/create_report/test_app.py b/tests/functions/monthly_reports/accounts/create_report/test_app.py new file mode 100644 index 0000000..d175d14 --- /dev/null +++ b/tests/functions/monthly_reports/accounts/create_report/test_app.py @@ -0,0 +1,320 @@ +import pytest +from unittest.mock import patch +from botocore.exceptions import ClientError +from functions.monthly_reports.accounts.create_report.create_report.exceptions import ( + ReportGenerationError, + ReportTemplateError, + ReportUploadError, +) + + +class TestCreateReportLambdaHandler: + """Test cases for the create_report Lambda handler.""" + + def test_successful_report_creation( + self, + create_report_app_with_mocks, + sample_event, + mock_pdf_bytes, + mock_presigned_url, + mock_context, + ): + """Test successful report creation and upload.""" + app = create_report_app_with_mocks + + # Mock the PDF generation + with patch( + "functions.monthly_reports.accounts.create_report.create_report.app.generate_transactions_pdf" + ) as mock_generate_pdf: + mock_generate_pdf.return_value = mock_pdf_bytes + + # Mock S3 operations + app.s3.put_object.return_value = { + "ResponseMetadata": {"HTTPStatusCode": 200} + } + app.s3.generate_presigned_url.return_value = mock_presigned_url + + # Call the handler + result = app.lambda_handler(sample_event, mock_context) + + # Verify PDF generation was called + mock_generate_pdf.assert_called_once_with( + event=sample_event, logger=app.logger + ) + + # Verify S3 upload was called with correct parameters + app.s3.put_object.assert_called_once_with( + Bucket="test-reports-bucket", + Key=f"{sample_event['accountId']}/{sample_event['statementPeriod']}.pdf", + Body=mock_pdf_bytes, + ContentType="application/pdf", + ) + + # Verify presigned URL generation was called + app.s3.generate_presigned_url.assert_called_once_with( + "get_object", + Params={ + "Bucket": "test-reports-bucket", + "Key": f"{sample_event['accountId']}/{sample_event['statementPeriod']}.pdf", + }, + ExpiresIn=3600, + ) + + # Verify the response + expected_response = { + "reportUrl": mock_presigned_url, + "accountId": sample_event["accountId"], + "userId": sample_event["userId"], + "statementPeriod": sample_event["statementPeriod"], + } + assert result == expected_response + + def test_pdf_generation_error( + self, create_report_app_with_mocks, sample_event, mock_context + ): + """Test handling of PDF generation errors.""" + app = create_report_app_with_mocks + + # Mock PDF generation to raise an error + with patch( + "functions.monthly_reports.accounts.create_report.create_report.app.generate_transactions_pdf" + ) as mock_generate_pdf: + mock_generate_pdf.side_effect = ReportGenerationError( + "PDF generation failed" + ) + + # Call the handler and expect the error to be re-raised + with pytest.raises(ReportGenerationError, match="PDF generation failed"): + app.lambda_handler(sample_event, mock_context) + + # Verify S3 operations were not called + app.s3.put_object.assert_not_called() + app.s3.generate_presigned_url.assert_not_called() + + def test_template_error( + self, create_report_app_with_mocks, sample_event, mock_context + ): + """Test handling of template errors.""" + app = create_report_app_with_mocks + + # Mock PDF generation to raise a template error + with patch( + "functions.monthly_reports.accounts.create_report.create_report.app.generate_transactions_pdf" + ) as mock_generate_pdf: + mock_generate_pdf.side_effect = ReportTemplateError("Template not found") + + # Call the handler and expect the error to be re-raised + with pytest.raises(ReportTemplateError, match="Template not found"): + app.lambda_handler(sample_event, mock_context) + + # Verify S3 operations were not called + app.s3.put_object.assert_not_called() + app.s3.generate_presigned_url.assert_not_called() + + def test_s3_upload_error( + self, create_report_app_with_mocks, sample_event, mock_pdf_bytes, mock_context + ): + """Test handling of S3 upload errors.""" + app = create_report_app_with_mocks + + # Mock the PDF generation + with patch( + "functions.monthly_reports.accounts.create_report.create_report.app.generate_transactions_pdf" + ) as mock_generate_pdf: + mock_generate_pdf.return_value = mock_pdf_bytes + + # Mock S3 upload to raise an error + error_response = { + "Error": {"Code": "NoSuchBucket", "Message": "Bucket does not exist"} + } + app.s3.put_object.side_effect = ClientError(error_response, "PutObject") + + # Call the handler and expect a ReportUploadError + with pytest.raises(ReportUploadError, match="S3 upload failed"): + app.lambda_handler(sample_event, mock_context) + + # Verify presigned URL generation was not called + app.s3.generate_presigned_url.assert_not_called() + + def test_presigned_url_generation_error( + self, create_report_app_with_mocks, sample_event, mock_pdf_bytes, mock_context + ): + """Test handling of presigned URL generation errors.""" + app = create_report_app_with_mocks + + # Mock the PDF generation + with patch( + "functions.monthly_reports.accounts.create_report.create_report.app.generate_transactions_pdf" + ) as mock_generate_pdf: + mock_generate_pdf.return_value = mock_pdf_bytes + + # Mock S3 upload to succeed + app.s3.put_object.return_value = { + "ResponseMetadata": {"HTTPStatusCode": 200} + } + + # Mock presigned URL generation to raise an error + error_response = { + "Error": {"Code": "AccessDenied", "Message": "Access denied"} + } + app.s3.generate_presigned_url.side_effect = ClientError( + error_response, "GeneratePresignedUrl" + ) + + # Call the handler and expect a ReportUploadError + with pytest.raises( + ReportUploadError, match="Presigned URL generation failed" + ): + app.lambda_handler(sample_event, mock_context) + + def test_missing_required_event_fields( + self, create_report_app_with_mocks, mock_context + ): + """Test handling of missing required event fields.""" + app = create_report_app_with_mocks + + # Create event with missing fields + incomplete_event = { + "accountId": "test-account-123" + # Missing userId, statementPeriod, transactions, accountBalance + } + + with pytest.raises(ReportGenerationError): + app.lambda_handler(incomplete_event, mock_context) + + def test_empty_transactions_list( + self, create_report_app_with_mocks, mock_presigned_url, mock_context + ): + """Test handling of empty transactions list.""" + app = create_report_app_with_mocks + + # Create event with empty transactions + event_with_empty_transactions = { + "accountId": "test-account-123", + "userId": "test-user-456", + "statementPeriod": "2024-01", + "transactions": [], + "accountBalance": 1500.00, + } + + mock_pdf_bytes = b"%PDF-1.4\n%Empty transactions PDF\n%%EOF" + + # Mock the PDF generation + with patch( + "functions.monthly_reports.accounts.create_report.create_report.app.generate_transactions_pdf" + ) as mock_generate_pdf: + mock_generate_pdf.return_value = mock_pdf_bytes + + # Mock S3 operations + app.s3.put_object.return_value = { + "ResponseMetadata": {"HTTPStatusCode": 200} + } + app.s3.generate_presigned_url.return_value = mock_presigned_url + + # Call the handler + result = app.lambda_handler(event_with_empty_transactions, mock_context) + + # Verify the response is correct + expected_response = { + "reportUrl": mock_presigned_url, + "accountId": event_with_empty_transactions["accountId"], + "userId": event_with_empty_transactions["userId"], + "statementPeriod": event_with_empty_transactions["statementPeriod"], + } + assert result == expected_response + + def test_logger_integration( + self, + create_report_app_with_mocks, + sample_event, + mock_pdf_bytes, + mock_presigned_url, + mock_context, + ): + """Test that logging is properly integrated.""" + app = create_report_app_with_mocks + + # Mock the PDF generation + with patch( + "functions.monthly_reports.accounts.create_report.create_report.app.generate_transactions_pdf" + ) as mock_generate_pdf: + mock_generate_pdf.return_value = mock_pdf_bytes + + # Mock S3 operations + app.s3.put_object.return_value = { + "ResponseMetadata": {"HTTPStatusCode": 200} + } + app.s3.generate_presigned_url.return_value = mock_presigned_url + + # Call the handler + result = app.lambda_handler(sample_event, mock_context) + + # Verify that the logger was used (we can't easily verify the exact calls due to the way powertools works) + # But we can verify the function completed successfully, which means logging worked + assert result is not None + assert "reportUrl" in result + + def test_s3_key_format( + self, + create_report_app_with_mocks, + sample_event, + mock_pdf_bytes, + mock_presigned_url, + mock_context, + ): + """Test that S3 key is formatted correctly.""" + app = create_report_app_with_mocks + + # Mock the PDF generation + with patch( + "functions.monthly_reports.accounts.create_report.create_report.app.generate_transactions_pdf" + ) as mock_generate_pdf: + mock_generate_pdf.return_value = mock_pdf_bytes + + # Mock S3 operations + app.s3.put_object.return_value = { + "ResponseMetadata": {"HTTPStatusCode": 200} + } + app.s3.generate_presigned_url.return_value = mock_presigned_url + + # Call the handler + app.lambda_handler(sample_event, mock_context) + + # Verify S3 key format + expected_key = ( + f"{sample_event['accountId']}/{sample_event['statementPeriod']}.pdf" + ) + app.s3.put_object.assert_called_once() + call_args = app.s3.put_object.call_args + assert call_args[1]["Key"] == expected_key + + def test_presigned_url_expiration( + self, + create_report_app_with_mocks, + sample_event, + mock_pdf_bytes, + mock_presigned_url, + mock_context, + ): + """Test that presigned URL is generated with correct expiration.""" + app = create_report_app_with_mocks + + # Mock the PDF generation + with patch( + "functions.monthly_reports.accounts.create_report.create_report.app.generate_transactions_pdf" + ) as mock_generate_pdf: + mock_generate_pdf.return_value = mock_pdf_bytes + + # Mock S3 operations + app.s3.put_object.return_value = { + "ResponseMetadata": {"HTTPStatusCode": 200} + } + app.s3.generate_presigned_url.return_value = mock_presigned_url + + # Call the handler + app.lambda_handler(sample_event, mock_context) + + # Verify presigned URL expiration + app.s3.generate_presigned_url.assert_called_once() + call_args = app.s3.generate_presigned_url.call_args + assert call_args[1]["ExpiresIn"] == 3600 diff --git a/tests/functions/monthly_reports/accounts/create_report/test_exceptions.py b/tests/functions/monthly_reports/accounts/create_report/test_exceptions.py new file mode 100644 index 0000000..a0ab1ac --- /dev/null +++ b/tests/functions/monthly_reports/accounts/create_report/test_exceptions.py @@ -0,0 +1,170 @@ +import pytest +from functions.monthly_reports.accounts.create_report.create_report.exceptions import ( + ReportGenerationError, + ReportTemplateError, + ReportUploadError, +) + + +class TestReportExceptions: + """Test cases for the custom exceptions used in the create_report module.""" + + def test_report_generation_error(self): + """Test ReportGenerationError exception.""" + error_message = "PDF generation failed" + error = ReportGenerationError(error_message) + + assert str(error) == error_message + assert isinstance(error, Exception) + assert isinstance(error, ReportGenerationError) + + def test_report_template_error(self): + """Test ReportTemplateError exception.""" + error_message = "Template not found" + error = ReportTemplateError(error_message) + + assert str(error) == error_message + assert isinstance(error, Exception) + assert isinstance(error, ReportTemplateError) + + def test_report_upload_error(self): + """Test ReportUploadError exception.""" + error_message = "S3 upload failed" + error = ReportUploadError(error_message) + + assert str(error) == error_message + assert isinstance(error, Exception) + assert isinstance(error, ReportUploadError) + + def test_exception_inheritance(self): + """Test that all custom exceptions inherit from Exception.""" + exceptions = [ReportGenerationError, ReportTemplateError, ReportUploadError] + + for exception_class in exceptions: + assert issubclass(exception_class, Exception) + + def test_exception_uniqueness(self): + """Test that all exceptions are unique classes.""" + exceptions = [ReportGenerationError, ReportTemplateError, ReportUploadError] + + for i, exception_class in enumerate(exceptions): + for j, other_exception_class in enumerate(exceptions): + if i != j: + assert exception_class != other_exception_class + + def test_exception_with_empty_message(self): + """Test exceptions with empty messages.""" + empty_message = "" + + generation_error = ReportGenerationError(empty_message) + template_error = ReportTemplateError(empty_message) + upload_error = ReportUploadError(empty_message) + + assert str(generation_error) == empty_message + assert str(template_error) == empty_message + assert str(upload_error) == empty_message + + def test_exception_with_none_message(self): + """Test exceptions with None messages.""" + none_message = None + + generation_error = ReportGenerationError(none_message) + template_error = ReportTemplateError(none_message) + upload_error = ReportUploadError(none_message) + + assert str(generation_error) == "None" + assert str(template_error) == "None" + assert str(upload_error) == "None" + + def test_exception_with_long_message(self): + """Test exceptions with long messages.""" + long_message = "A" * 1000 + + generation_error = ReportGenerationError(long_message) + template_error = ReportTemplateError(long_message) + upload_error = ReportUploadError(long_message) + + assert str(generation_error) == long_message + assert str(template_error) == long_message + assert str(upload_error) == long_message + + def test_exception_with_special_characters(self): + """Test exceptions with special characters in messages.""" + special_message = "Error with special chars: !@#$%^&*()_+-=[]{}|;':\",./<>?" + + generation_error = ReportGenerationError(special_message) + template_error = ReportTemplateError(special_message) + upload_error = ReportUploadError(special_message) + + assert str(generation_error) == special_message + assert str(template_error) == special_message + assert str(upload_error) == special_message + + def test_exception_with_unicode_characters(self): + """Test exceptions with unicode characters in messages.""" + unicode_message = "Error with unicode: 你好世界 🌍" + + generation_error = ReportGenerationError(unicode_message) + template_error = ReportTemplateError(unicode_message) + upload_error = ReportUploadError(unicode_message) + + assert str(generation_error) == unicode_message + assert str(template_error) == unicode_message + assert str(upload_error) == unicode_message + + def test_exception_raising_and_catching(self): + """Test that exceptions can be raised and caught properly.""" + error_message = "Test error message" + + # Test ReportGenerationError + with pytest.raises(ReportGenerationError) as exc_info: + raise ReportGenerationError(error_message) + assert str(exc_info.value) == error_message + + # Test ReportTemplateError + with pytest.raises(ReportTemplateError) as exc_info: + raise ReportTemplateError(error_message) + assert str(exc_info.value) == error_message + + # Test ReportUploadError + with pytest.raises(ReportUploadError) as exc_info: + raise ReportUploadError(error_message) + assert str(exc_info.value) == error_message + + def test_exception_in_except_block(self): + """Test that exceptions work properly in except blocks.""" + + def function_that_raises_generation_error(): + raise ReportGenerationError("Generation failed") + + def function_that_raises_template_error(): + raise ReportTemplateError("Template failed") + + def function_that_raises_upload_error(): + raise ReportUploadError("Upload failed") + + # Test catching and re-raising + with pytest.raises(ReportGenerationError) as exc_info: + try: + function_that_raises_generation_error() + except ReportGenerationError as e: + assert str(e) == "Generation failed" + # Re-raise to test it works + raise + assert str(exc_info.value) == "Generation failed" + + with pytest.raises(ReportTemplateError) as exc_info: + try: + function_that_raises_template_error() + except ReportTemplateError as e: + assert str(e) == "Template failed" + raise + assert str(exc_info.value) == "Template failed" + + with pytest.raises(ReportUploadError) as exc_info: + try: + function_that_raises_upload_error() + except ReportUploadError as e: + assert str(e) == "Upload failed" + raise + assert str(exc_info.value) == "Upload failed" diff --git a/tests/functions/monthly_reports/accounts/create_report/test_generate_pdf.py b/tests/functions/monthly_reports/accounts/create_report/test_generate_pdf.py new file mode 100644 index 0000000..aabd860 --- /dev/null +++ b/tests/functions/monthly_reports/accounts/create_report/test_generate_pdf.py @@ -0,0 +1,336 @@ +import pytest +from unittest.mock import patch, MagicMock +from functions.monthly_reports.accounts.create_report.create_report.generate_pdf import ( + generate_transactions_pdf, +) +from functions.monthly_reports.accounts.create_report.create_report.exceptions import ( + ReportGenerationError, + ReportTemplateError, +) + + +class TestGenerateTransactionsPDF: + """Test cases for the generate_transactions_pdf function.""" + + @pytest.fixture + def sample_event(self): + """Sample event data for testing.""" + return { + "accountId": "test-account-123", + "statementPeriod": "2024-01", + "transactions": [ + { + "id": "txn-1", + "amount": 100.00, + "description": "Test transaction 1", + "date": "2024-01-15", + }, + { + "id": "txn-2", + "amount": -50.00, + "description": "Test transaction 2", + "date": "2024-01-20", + }, + ], + "accountBalance": 1500.00, + } + + @pytest.fixture + def mock_logger(self): + """Mock logger for testing.""" + return MagicMock() + + @pytest.fixture + def mock_template_content(self): + """Mock HTML template content.""" + return """ + + Account Statement + +

Account Statement

+

Account ID: {{ accountId }}

+

Statement Period: {{ statementPeriod }}

+

Account Balance: {{ accountBalance }}

+

Generated: {{ generationDate }}

+ + + {% for transaction in transactions %} + + + + + + + {% endfor %} +
IDAmountDescriptionDate
{{ transaction.id }}{{ transaction.amount }}{{ transaction.description }}{{ transaction.date }}
+ + + """ + + def test_successful_pdf_generation( + self, sample_event, mock_logger, mock_template_content + ): + """Test successful PDF generation.""" + with patch("os.path.dirname") as mock_dirname: + mock_dirname.return_value = "/mock/path" + + # Mock the Jinja2 Environment and template + with patch( + "functions.monthly_reports.accounts.create_report.create_report.generate_pdf.Environment" + ) as mock_env: + mock_template = MagicMock() + mock_template.render.return_value = "Test PDF" + mock_env_instance = MagicMock() + mock_env_instance.get_template.return_value = mock_template + mock_env.return_value = mock_env_instance + + # Mock xhtml2pdf + with patch("xhtml2pdf.pisa.CreatePDF") as mock_pisa: + mock_pisa.return_value.err = False + + # Mock tempfile + with patch("tempfile.NamedTemporaryFile") as mock_tempfile: + mock_tempfile_instance = MagicMock() + mock_tempfile_instance.name = "/tmp/test.pdf" + mock_tempfile.return_value.__enter__.return_value = ( + mock_tempfile_instance + ) + mock_tempfile.return_value.__exit__.return_value = None + + # Call the function + result = generate_transactions_pdf(sample_event, mock_logger) + + # Verify template was rendered with correct data + mock_template.render.assert_called_once() + call_args = mock_template.render.call_args[1] + assert call_args["accountId"] == sample_event["accountId"] + assert ( + call_args["statementPeriod"] + == sample_event["statementPeriod"] + ) + assert call_args["transactions"] == sample_event["transactions"] + assert ( + call_args["accountBalance"] + == sample_event["accountBalance"] + ) + assert "generationDate" in call_args + + # Verify PDF generation was called + mock_pisa.assert_called_once() + + # Verify result is bytes + assert isinstance(result, bytes) + + def test_template_not_found_error(self, sample_event, mock_logger): + """Test handling of template not found error.""" + with patch("os.path.dirname") as mock_dirname: + mock_dirname.return_value = "/mock/path" + + # Mock Jinja2 to raise TemplateNotFound + with patch( + "functions.monthly_reports.accounts.create_report.create_report.generate_pdf.Environment" + ) as mock_env: + from jinja2 import TemplateNotFound + + mock_env_instance = MagicMock() + mock_env_instance.get_template.side_effect = TemplateNotFound( + "template.html", "template.html" + ) + mock_env.return_value = mock_env_instance + + # Call the function and expect ReportTemplateError + with pytest.raises( + ReportTemplateError, match="Missing template: template.html" + ): + generate_transactions_pdf(sample_event, mock_logger) + + # Verify error was logged + mock_logger.error.assert_called_with( + "Template 'template.html' not found" + ) + + def test_pdf_generation_error( + self, sample_event, mock_logger, mock_template_content + ): + """Test handling of PDF generation error.""" + with patch("os.path.dirname") as mock_dirname: + mock_dirname.return_value = "/mock/path" + + # Mock the Jinja2 Environment and template + with patch( + "functions.monthly_reports.accounts.create_report.create_report.generate_pdf.Environment" + ) as mock_env: + mock_template = MagicMock() + mock_template.render.return_value = "Test PDF" + mock_env_instance = MagicMock() + mock_env_instance.get_template.return_value = mock_template + mock_env.return_value = mock_env_instance + + # Mock xhtml2pdf to return an error + with patch("xhtml2pdf.pisa.CreatePDF") as mock_pisa: + mock_pisa.return_value.err = True + + # Call the function and expect ReportGenerationError + with pytest.raises( + ReportGenerationError, match="Error generating PDF" + ): + generate_transactions_pdf(sample_event, mock_logger) + + # Verify error was logged + mock_logger.error.assert_called_with( + "xhtml2pdf failed to generate PDF" + ) + + def test_empty_transactions(self, mock_logger, mock_template_content): + """Test PDF generation with empty transactions list.""" + event_with_empty_transactions = { + "accountId": "test-account-123", + "statementPeriod": "2024-01", + "transactions": [], + "accountBalance": 1500.00, + } + + with patch("os.path.dirname") as mock_dirname: + mock_dirname.return_value = "/mock/path" + + # Mock the Jinja2 Environment and template + with patch( + "functions.monthly_reports.accounts.create_report.create_report.generate_pdf.Environment" + ) as mock_env: + mock_template = MagicMock() + mock_template.render.return_value = ( + "Empty PDF" + ) + mock_env_instance = MagicMock() + mock_env_instance.get_template.return_value = mock_template + mock_env.return_value = mock_env_instance + + # Mock xhtml2pdf + with patch("xhtml2pdf.pisa.CreatePDF") as mock_pisa: + mock_pisa.return_value.err = False + + # Mock tempfile + with patch("tempfile.NamedTemporaryFile") as mock_tempfile: + mock_tempfile_instance = MagicMock() + mock_tempfile_instance.name = "/tmp/test.pdf" + mock_tempfile.return_value.__enter__.return_value = ( + mock_tempfile_instance + ) + mock_tempfile.return_value.__exit__.return_value = None + + # Call the function + result = generate_transactions_pdf( + event_with_empty_transactions, mock_logger + ) + + # Verify template was rendered with empty transactions + mock_template.render.assert_called_once() + call_args = mock_template.render.call_args[1] + assert call_args["transactions"] == [] + + # Verify result is bytes + assert isinstance(result, bytes) + + def test_large_transactions_list(self, mock_logger, mock_template_content): + """Test PDF generation with a large number of transactions.""" + large_transactions = [ + { + "id": f"txn-{i}", + "amount": 100.00 + i, + "description": f"Transaction {i}", + "date": "2024-01-15", + } + for i in range(100) + ] + + event_with_large_transactions = { + "accountId": "test-account-123", + "statementPeriod": "2024-01", + "transactions": large_transactions, + "accountBalance": 1500.00, + } + + with patch("os.path.dirname") as mock_dirname: + mock_dirname.return_value = "/mock/path" + + # Mock the Jinja2 Environment and template + with patch( + "functions.monthly_reports.accounts.create_report.create_report.generate_pdf.Environment" + ) as mock_env: + mock_template = MagicMock() + mock_template.render.return_value = ( + "Large PDF" + ) + mock_env_instance = MagicMock() + mock_env_instance.get_template.return_value = mock_template + mock_env.return_value = mock_env_instance + + # Mock xhtml2pdf + with patch("xhtml2pdf.pisa.CreatePDF") as mock_pisa: + mock_pisa.return_value.err = False + + # Mock tempfile + with patch("tempfile.NamedTemporaryFile") as mock_tempfile: + mock_tempfile_instance = MagicMock() + mock_tempfile_instance.name = "/tmp/test.pdf" + mock_tempfile.return_value.__enter__.return_value = ( + mock_tempfile_instance + ) + mock_tempfile.return_value.__exit__.return_value = None + + # Call the function + result = generate_transactions_pdf( + event_with_large_transactions, mock_logger + ) + + # Verify template was rendered with all transactions + mock_template.render.assert_called_once() + call_args = mock_template.render.call_args[1] + assert len(call_args["transactions"]) == 100 + + # Verify result is bytes + assert isinstance(result, bytes) + + def test_generation_date_format( + self, sample_event, mock_logger, mock_template_content + ): + """Test that generation date is properly formatted.""" + with patch("os.path.dirname") as mock_dirname: + mock_dirname.return_value = "/mock/path" + + # Mock the Jinja2 Environment and template + with patch( + "functions.monthly_reports.accounts.create_report.create_report.generate_pdf.Environment" + ) as mock_env: + mock_template = MagicMock() + mock_template.render.return_value = "Test PDF" + mock_env_instance = MagicMock() + mock_env_instance.get_template.return_value = mock_template + mock_env.return_value = mock_env_instance + + # Mock xhtml2pdf + with patch("xhtml2pdf.pisa.CreatePDF") as mock_pisa: + mock_pisa.return_value.err = False + + # Mock tempfile + with patch("tempfile.NamedTemporaryFile") as mock_tempfile: + mock_tempfile_instance = MagicMock() + mock_tempfile_instance.name = "/tmp/test.pdf" + mock_tempfile.return_value.__enter__.return_value = ( + mock_tempfile_instance + ) + mock_tempfile.return_value.__exit__.return_value = None + + # Call the function + generate_transactions_pdf(sample_event, mock_logger) + + # Verify generation date format + mock_template.render.assert_called_once() + call_args = mock_template.render.call_args[1] + generation_date = call_args["generationDate"] + + # Check that it matches the expected format: YYYY-MM-DD HH:MM:SS UTC + import re + + pattern = r"^\d{4}-\d{2}-\d{2} \d{2}:\d{2}:\d{2} UTC$" + assert re.match(pattern, generation_date) is not None diff --git a/tests/functions/monthly_reports/accounts/notify_client/__init__.py b/tests/functions/monthly_reports/accounts/notify_client/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/functions/monthly_reports/accounts/notify_client/conftest.py b/tests/functions/monthly_reports/accounts/notify_client/conftest.py new file mode 100644 index 0000000..15214d1 --- /dev/null +++ b/tests/functions/monthly_reports/accounts/notify_client/conftest.py @@ -0,0 +1,124 @@ +from importlib import reload +from unittest.mock import MagicMock +import pytest +from functions.monthly_reports.accounts.notify_client.notify_client import app + + +@pytest.fixture(scope="function") +def notify_client_app_with_mocks( + monkeypatch, mock_s3_client, magic_mock_ses_client, mock_cognito_client +): + + monkeypatch.setenv("SES_NO_REPLY_EMAIL", "noreply@testbank.com") + monkeypatch.setenv("REPORTS_BUCKET", "test-reports-bucket") + monkeypatch.setenv("AWS_REGION", "eu-west-2") + monkeypatch.setenv("POWERTOOLS_LOG_LEVEL", "INFO") + monkeypatch.setenv("COGNITO_USER_POOL_ID", "eu-west-2_testpool123") + monkeypatch.setenv("COGNITO_CLIENT_ID", "test-client-id-123") + monkeypatch.setenv("DYNAMODB_ENDPOINT", "") + monkeypatch.setenv("ACCOUNTS_TABLE_NAME", "test-accounts-table") + + mock_s3_client.head_object.return_value = {"ContentLength": 1024 * 1024} # 1MB + mock_s3_client.get_object.return_value = { + "Body": MagicMock(read=lambda: b"%PDF-1.4\n%Test PDF content\n%%EOF") + } + mock_s3_client.generate_presigned_url.return_value = "https://test-reports-bucket.s3.eu-west-2.amazonaws.com/test-account-123/2024-01.pdf?AWSAccessKeyId=test&Signature=test&Expires=1234567890" + + magic_mock_ses_client.send_email.return_value = {"MessageId": "test-message-id-123"} + magic_mock_ses_client.send_raw_email.return_value = { + "MessageId": "test-message-id-456" + } + + mock_cognito_client.admin_get_user.return_value = { + "UserAttributes": [ + {"Name": "email", "Value": "test@example.com"}, + {"Name": "name", "Value": "John Doe"}, + ] + } + + reload(app) + + app.s3 = mock_s3_client + + yield app + + +@pytest.fixture +def sample_event(): + """Sample event data for testing.""" + return { + "accountId": "test-account-123", + "userId": "test-user-456", + "statementPeriod": "2024-01", + } + + +@pytest.fixture +def mock_context(): + """Mock Lambda context for testing.""" + context = MagicMock() + context.function_name = "notify-client" + context.function_version = "$LATEST" + context.invoked_function_arn = ( + "arn:aws:lambda:eu-west-2:123456789012:function:notify-client" + ) + context.memory_limit_in_mb = 128 + context.remaining_time_in_millis = lambda: 30000 + context.aws_request_id = "test-request-id-123" + return context + + +@pytest.fixture +def mock_user_attributes(): + """Mock user attributes from Cognito.""" + return {"email": "test@example.com", "name": "John Doe", "sub": "test-user-456"} + + +@pytest.fixture +def mock_pdf_bytes(): + """Mock PDF bytes for testing.""" + return b"%PDF-1.4\n%Test PDF content\n%%EOF" + + +@pytest.fixture +def mock_presigned_url(): + """Mock presigned URL for testing.""" + return "https://test-reports-bucket.s3.eu-west-2.amazonaws.com/test-account-123/2024-01.pdf?AWSAccessKeyId=test&Signature=test&Expires=1234567890" + + +@pytest.fixture +def mock_dynamodb_table(): + """Mock DynamoDB table for testing.""" + mock_table = MagicMock() + mock_table.get_item.return_value = { + "Item": { + "accountId": "test-account-123", + "userId": "test-user-456", + "balance": 1000.0, + } + } + return mock_table + + +@pytest.fixture +def api_gateway_event(): + """Mock API Gateway event for testing.""" + return { + "httpMethod": "GET", + "path": "/accounts/test-account-123/reports/2024-01", + "headers": { + "Authorization": "Bearer test-jwt-token", + "Content-Type": "application/json", + }, + "requestContext": { + "requestId": "test-request-id-123", + "http": { + "method": "GET", + "path": "/accounts/test-account-123/reports/2024-01", + }, + }, + "pathParameters": { + "account_id": "test-account-123", + "statement_period": "2024-01", + }, + } diff --git a/tests/functions/monthly_reports/accounts/notify_client/test_app.py b/tests/functions/monthly_reports/accounts/notify_client/test_app.py new file mode 100644 index 0000000..5078eda --- /dev/null +++ b/tests/functions/monthly_reports/accounts/notify_client/test_app.py @@ -0,0 +1,522 @@ +import json +from unittest.mock import patch +from botocore.exceptions import ClientError + + +class TestNotifyClientLambdaHandler: + """Test cases for the notify_client Lambda handler.""" + + def test_successful_notification_with_attachment( + self, + notify_client_app_with_mocks, + sample_event, + mock_context, + mock_user_attributes, + mock_pdf_bytes, + ): + """Test successful notification with PDF attachment.""" + app = notify_client_app_with_mocks + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.processing.get_user_attributes" + ) as mock_get_user: + mock_get_user.return_value = mock_user_attributes + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.send_report.send_user_email_with_attachment" + ) as mock_send_email: + mock_send_email.return_value = {"MessageId": "test-message-id-123"} + + result = app.lambda_handler(sample_event, mock_context) + + mock_get_user.assert_called_once_with( + aws_region="eu-west-2", + logger=app.logger, + username=sample_event["userId"], + user_pool_id="eu-west-2_testpool123", + ) + + app.s3.head_object.assert_called_once_with( + Bucket="test-reports-bucket", + Key=f"{sample_event['accountId']}/{sample_event['statementPeriod']}.pdf", + ) + + mock_send_email.assert_called_once_with( + aws_region="eu-west-2", + logger=app.logger, + sender_email="noreply@testbank.com", + to_addresses=[mock_user_attributes["email"]], + subject_data=f"Your Account Statement for {sample_event['statementPeriod']}", + body_text=f"Hello {mock_user_attributes['name']},\n\nPlease find your account statement attached.\n\nKind Regards.", + attachment_bytes=mock_pdf_bytes, + attachment_filename="statement.pdf", + ) + + expected_response = { + "status": "success", + "messageId": "test-message-id-123", + "mode": "attachment", + } + assert result == expected_response + + def test_successful_notification_with_link( + self, + notify_client_app_with_mocks, + sample_event, + mock_context, + mock_user_attributes, + mock_presigned_url, + ): + app = notify_client_app_with_mocks + + app.s3.head_object.return_value = {"ContentLength": 8 * 1024 * 1024} # 8MB + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.processing.get_user_attributes" + ) as mock_get_user: + mock_get_user.return_value = mock_user_attributes + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.send_report.send_user_email" + ) as mock_send_email: + mock_send_email.return_value = {"MessageId": "test-message-id-456"} + + result = app.lambda_handler(sample_event, mock_context) + + mock_get_user.assert_called_once_with( + aws_region="eu-west-2", + logger=app.logger, + username=sample_event["userId"], + user_pool_id="eu-west-2_testpool123", + ) + + app.s3.head_object.assert_called_once_with( + Bucket="test-reports-bucket", + Key=f"{sample_event['accountId']}/{sample_event['statementPeriod']}.pdf", + ) + + app.s3.generate_presigned_url.assert_called_once_with( + "get_object", + Params={ + "Bucket": "test-reports-bucket", + "Key": f"{sample_event['accountId']}/{sample_event['statementPeriod']}.pdf", + }, + ExpiresIn=3600, + ) + + mock_send_email.assert_called_once_with( + aws_region="eu-west-2", + logger=app.logger, + sender_email="noreply@testbank.com", + to_addresses=[mock_user_attributes["email"]], + subject_data=f"Your Account Statement for {sample_event['statementPeriod']}", + text_body_data=( + f"Hello {mock_user_attributes['name']},\n\n" + f"Your account statement is ready.\n\n" + f"Download it here (valid for 1 hour):\n{mock_presigned_url}\n\n" + f"If you need a new link please request one through the API.\n\n" + f"Kind Regards." + ), + ) + + expected_response = { + "status": "success", + "messageId": "test-message-id-456", + "mode": "link", + } + assert result == expected_response + + def test_user_without_email_attribute( + self, notify_client_app_with_mocks, sample_event, mock_context + ): + app = notify_client_app_with_mocks + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.processing.get_user_attributes" + ) as mock_get_user: + mock_get_user.return_value = {"name": "John Doe"} + + result = app.lambda_handler(sample_event, mock_context) + + assert result["statusCode"] == 500 + assert ( + "User test-user-456 has no email attribute in Cognito" in result["body"] + ) + + app.s3.head_object.assert_not_called() + app.s3.get_object.assert_not_called() + + def test_s3_client_error( + self, + notify_client_app_with_mocks, + sample_event, + mock_context, + mock_user_attributes, + ): + app = notify_client_app_with_mocks + + app.s3.head_object.side_effect = ClientError( + error_response={ + "Error": { + "Code": "NoSuchKey", + "Message": "The specified key does not exist.", + } + }, + operation_name="HeadObject", + ) + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.processing.get_user_attributes" + ) as mock_get_user: + mock_get_user.return_value = mock_user_attributes + + result = app.lambda_handler(sample_event, mock_context) + + assert result["statusCode"] == 500 + assert "NoSuchKey" in result["body"] + + def test_email_sending_failure( + self, + notify_client_app_with_mocks, + sample_event, + mock_context, + mock_user_attributes, + ): + app = notify_client_app_with_mocks + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.processing.get_user_attributes" + ) as mock_get_user: + mock_get_user.return_value = mock_user_attributes + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.send_report.send_user_email_with_attachment" + ) as mock_send_email: + mock_send_email.return_value = None + + result = app.lambda_handler(sample_event, mock_context) + + expected_response = { + "status": "failed", + "messageId": None, + "mode": "attachment", + } + assert result == expected_response + + def test_email_sending_exception( + self, + notify_client_app_with_mocks, + sample_event, + mock_context, + mock_user_attributes, + ): + app = notify_client_app_with_mocks + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.processing.get_user_attributes" + ) as mock_get_user: + mock_get_user.return_value = mock_user_attributes + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.send_report.send_user_email_with_attachment" + ) as mock_send_email: + mock_send_email.side_effect = Exception("SES service unavailable") + + result = app.lambda_handler(sample_event, mock_context) + + assert result["statusCode"] == 500 + assert "SES service unavailable" in result["body"] + + def test_user_attributes_retrieval_failure( + self, notify_client_app_with_mocks, sample_event, mock_context + ): + app = notify_client_app_with_mocks + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.processing.get_user_attributes" + ) as mock_get_user: + mock_get_user.side_effect = Exception("Cognito service unavailable") + + result = app.lambda_handler(sample_event, mock_context) + + assert result["statusCode"] == 500 + assert "Cognito service unavailable" in result["body"] + + def test_user_without_name_attribute( + self, notify_client_app_with_mocks, sample_event, mock_context, mock_pdf_bytes + ): + app = notify_client_app_with_mocks + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.processing.get_user_attributes" + ) as mock_get_user: + mock_get_user.return_value = {"email": "test@example.com"} + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.send_report.send_user_email_with_attachment" + ) as mock_send_email: + mock_send_email.return_value = {"MessageId": "test-message-id-123"} + + result = app.lambda_handler(sample_event, mock_context) + + mock_send_email.assert_called_once_with( + aws_region="eu-west-2", + logger=app.logger, + sender_email="noreply@testbank.com", + to_addresses=["test@example.com"], + subject_data=f"Your Account Statement for {sample_event['statementPeriod']}", + body_text="Hello Customer,\n\nPlease find your account statement attached.\n\nKind Regards.", + attachment_bytes=mock_pdf_bytes, + attachment_filename="statement.pdf", + ) + + expected_response = { + "status": "success", + "messageId": "test-message-id-123", + "mode": "attachment", + } + assert result == expected_response + + def test_exact_file_size_threshold( + self, + notify_client_app_with_mocks, + sample_event, + mock_context, + mock_user_attributes, + ): + app = notify_client_app_with_mocks + + app.s3.head_object.return_value = { + "ContentLength": 7 * 1024 * 1024 + } # Exactly 7MB + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.processing.get_user_attributes" + ) as mock_get_user: + mock_get_user.return_value = mock_user_attributes + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.send_report.send_user_email_with_attachment" + ) as mock_send_email: + mock_send_email.return_value = {"MessageId": "test-message-id-123"} + + result = app.lambda_handler(sample_event, mock_context) + + mock_send_email.assert_called_once() + assert result["mode"] == "attachment" + + def test_missing_required_fields_direct_invocation( + self, notify_client_app_with_mocks, mock_context + ): + """Test lambda handler with missing required fields for direct invocation.""" + app = notify_client_app_with_mocks + + # Test missing accountId + event_missing_account = { + "userId": "test-user-456", + "statementPeriod": "2024-01", + } + result = app.lambda_handler(event_missing_account, mock_context) + + assert result["statusCode"] == 400 + assert "Missing accountId, userId, or statementPeriod" in result["body"] + + # Test missing userId + event_missing_user = { + "accountId": "test-account-123", + "statementPeriod": "2024-01", + } + result = app.lambda_handler(event_missing_user, mock_context) + + assert result["statusCode"] == 400 + assert "Missing accountId, userId, or statementPeriod" in result["body"] + + # Test missing statementPeriod + event_missing_period = { + "accountId": "test-account-123", + "userId": "test-user-456", + } + result = app.lambda_handler(event_missing_period, mock_context) + + assert result["statusCode"] == 400 + assert "Missing accountId, userId, or statementPeriod" in result["body"] + + def test_direct_invocation_exception_handling( + self, notify_client_app_with_mocks, sample_event, mock_context + ): + """Test lambda handler exception handling for direct invocation.""" + app = notify_client_app_with_mocks + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.app.process_report" + ) as mock_process_report: + mock_process_report.side_effect = Exception("Test exception") + + result = app.lambda_handler(sample_event, mock_context) + + assert result["statusCode"] == 500 + assert "Test exception" in result["body"] + + +class TestNotifyClientAPIGateway: + + def test_successful_api_gateway_request( + self, + notify_client_app_with_mocks, + api_gateway_event, + mock_context, + mock_user_attributes, + mock_pdf_bytes, + ): + app = notify_client_app_with_mocks + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.app.authenticate_request" + ) as mock_auth: + mock_auth.return_value = "test-user-456" + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.app.check_user_owns_account" + ) as mock_check_ownership: + mock_check_ownership.return_value = True + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.processing.get_user_attributes" + ) as mock_get_user: + mock_get_user.return_value = mock_user_attributes + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.send_report.send_user_email_with_attachment" + ) as mock_send_email: + mock_send_email.return_value = { + "MessageId": "test-message-id-123" + } + + result = app.lambda_handler(api_gateway_event, mock_context) + + assert "statusCode" in result + assert result["statusCode"] == 200 + assert "body" in result + + response_body = json.loads(result["body"]) + + assert response_body["status"] == "success" + assert response_body["messageId"] == "test-message-id-123" + assert response_body["mode"] == "attachment" + + def test_api_gateway_no_user_id( + self, + notify_client_app_with_mocks, + api_gateway_event, + mock_context, + ): + """Test API Gateway request with authorization failure.""" + app = notify_client_app_with_mocks + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.app.authenticate_request" + ) as mock_auth: + mock_auth.return_value = "" + + result = app.lambda_handler(api_gateway_event, mock_context) + + assert "statusCode" in result + assert result["statusCode"] == 401 + assert "body" in result + + response_body = json.loads(result["body"]) + assert "Unauthorized" in response_body.get("message", "") + + def test_api_gateway_authorization_failure( + self, + notify_client_app_with_mocks, + api_gateway_event, + mock_context, + ): + """Test API Gateway request with authorization failure.""" + app = notify_client_app_with_mocks + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.app.authenticate_request" + ) as mock_auth: + mock_auth.return_value = "test-user-456" + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.app.check_user_owns_account" + ) as mock_check_ownership: + mock_check_ownership.return_value = False + + result = app.lambda_handler(api_gateway_event, mock_context) + + assert "statusCode" in result + assert result["statusCode"] == 401 + assert "body" in result + + response_body = json.loads(result["body"]) + assert "Unauthorized" in response_body.get("message", "") + + def test_api_gateway_internal_server_error( + self, + notify_client_app_with_mocks, + api_gateway_event, + mock_context, + ): + """Test API Gateway request with internal server error.""" + app = notify_client_app_with_mocks + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.app.authenticate_request" + ) as mock_auth: + mock_auth.return_value = "test-user-456" + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.app.check_user_owns_account" + ) as mock_check_ownership: + mock_check_ownership.return_value = True + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.app.process_report" + ) as mock_process_report: + mock_process_report.side_effect = Exception("Internal error") + + result = app.lambda_handler(api_gateway_event, mock_context) + + assert "statusCode" in result + assert result["statusCode"] == 500 + assert "body" in result + + response_body = json.loads(result["body"]) + assert "Internal server error" in response_body.get("message", "") + + def test_api_gateway_period_in_future( + self, + notify_client_app_with_mocks, + api_gateway_event, + mock_context, + ): + """Test API Gateway request with statement period in the future.""" + app = notify_client_app_with_mocks + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.app.authenticate_request" + ) as mock_auth: + mock_auth.return_value = "test-user-456" + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.app.check_user_owns_account" + ) as mock_check_ownership: + mock_check_ownership.return_value = True + + with patch( + "functions.monthly_reports.accounts.notify_client.notify_client.app.period_is_in_future" + ) as mock_period_check: + mock_period_check.return_value = True + + result = app.lambda_handler(api_gateway_event, mock_context) + + assert "statusCode" in result + assert result["statusCode"] == 500 + assert "body" in result + + response_body = json.loads(result["body"]) + assert "Internal server error" in response_body.get("message", "") diff --git a/tests/functions/monthly_reports/accounts/notify_client/test_date_helpers.py b/tests/functions/monthly_reports/accounts/notify_client/test_date_helpers.py new file mode 100644 index 0000000..4d50380 --- /dev/null +++ b/tests/functions/monthly_reports/accounts/notify_client/test_date_helpers.py @@ -0,0 +1,46 @@ +import datetime + +import pytest +from dateutil.relativedelta import relativedelta + +from functions.monthly_reports.accounts.notify_client.notify_client.date_helpers import ( + period_is_in_future, +) + + +class TestPeriodIsInFuture: + def test_period_is_current_month(self): + today = datetime.datetime.now(datetime.UTC).strftime("%Y-%m") + result = period_is_in_future(today) + assert result is True + + def test_period_is_future_month(self): + today = datetime.datetime.now(datetime.timezone.utc) + one_year_later = today + relativedelta(years=1) + result = period_is_in_future(one_year_later.strftime("%Y-%m")) + assert result is True + + def test_period_is_in_the_past(self): + today = datetime.datetime.now(datetime.timezone.utc) + one_year_before = today - relativedelta(years=1) + result = period_is_in_future(one_year_before.strftime("%Y-%m")) + assert result is False + + def test_invalid_format(self): + with pytest.raises( + ValueError, match="Invalid statement_period format. Use 'YYYY-MM'." + ): + today = datetime.datetime.now(datetime.UTC) + period_is_in_future(today.strftime("%Y/%m")) + + def test_invalid_month(self): + with pytest.raises( + ValueError, match="Invalid statement_period format. Use 'YYYY-MM'." + ): + period_is_in_future("2025-13") + + def test_empty_string(self): + with pytest.raises( + ValueError, match="Invalid statement_period format. Use 'YYYY-MM'." + ): + period_is_in_future("") diff --git a/tests/functions/monthly_reports/accounts/process_pending_reports/__init__.py b/tests/functions/monthly_reports/accounts/process_pending_reports/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/functions/monthly_reports/accounts/process_pending_reports/conftest.py b/tests/functions/monthly_reports/accounts/process_pending_reports/conftest.py new file mode 100644 index 0000000..3942871 --- /dev/null +++ b/tests/functions/monthly_reports/accounts/process_pending_reports/conftest.py @@ -0,0 +1,35 @@ +from importlib import reload +from functions.monthly_reports.accounts.process_pending_reports.process_pending_reports import ( + app, +) +from unittest.mock import patch + +import pytest + + +@pytest.fixture(scope="function") +def monthly_reports_continuation_app_with_mocks( + monkeypatch, dynamo_resource, mock_accounts_dynamo_table +): + accounts_table_name = mock_accounts_dynamo_table + + monkeypatch.setenv("ACCOUNTS_TABLE_NAME", accounts_table_name) + monkeypatch.setenv( + "CONTINUATION_QUEUE_URL", + "https://sqs.eu-west-2.amazonaws.com/123456789012/continuation-queue", + ) + monkeypatch.setenv("STATE_MACHINE_ARN", "mock_arn") + monkeypatch.setenv("ENVIRONMENT_NAME", "test") + monkeypatch.setenv("POWERTOOLS_LOG_LEVEL", "INFO") + monkeypatch.setenv("AWS_REGION", "eu-west-2") + monkeypatch.setenv( + "DLQ_URL", "https://sqs.eu-west-2.amazonaws.com/123456789012/dlq" + ) + monkeypatch.setenv("SQS_ENDPOINT", "https://sqs.eu-west-2.amazonaws.com") + + with patch("boto3.resource", return_value=dynamo_resource): + reload(app) + + app.accounts_table = dynamo_resource.Table(accounts_table_name) + + yield app diff --git a/tests/functions/monthly_reports/accounts/process_pending_reports/test_app.py b/tests/functions/monthly_reports/accounts/process_pending_reports/test_app.py new file mode 100644 index 0000000..45bf623 --- /dev/null +++ b/tests/functions/monthly_reports/accounts/process_pending_reports/test_app.py @@ -0,0 +1,703 @@ +import json +from importlib import reload +from unittest.mock import MagicMock, patch + +import pytest + +from functions.monthly_reports.accounts.process_pending_reports.process_pending_reports import ( + app, +) +from functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app import ( + lambda_handler, +) + + +class TestLambdaHandler: + def test_accounts_scan_continuation_success( + self, monthly_reports_continuation_app_with_mocks + ): + mock_event = { + "Records": [ + { + "body": json.dumps( + { + "scan_params": { + "ProjectionExpression": "accountId, userId" + }, + "statement_period": "2024-01", + } + ), + "messageAttributes": { + "continuation_type": {"stringValue": "accounts_scan"} + }, + } + ] + } + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 300000 + + with patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.process_accounts_scan_continuation" + ) as mock_process_scan: + mock_process_scan.return_value = { + "processed_count": 5, + "failed_starts_count": 0, + "skipped_count": 0, + "already_exists_count": 0, + "batches_processed": 1, + "pages_processed": 1, + "status": "COMPLETED", + } + + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 200 + body = ( + response["body"] + if isinstance(response["body"], dict) + else json.loads(response["body"]) + ) + assert body["message"] == "Monthly Account reports processing completed" + assert body["processed_count"] == 5 + assert body["batches_processed"] == 1 + + mock_process_scan.assert_called_once_with( + {"ProjectionExpression": "accountId, userId"}, + "2024-01", + mock_context, + app.logger, + app.accounts_table, + app.sfn_client, + app.STATE_MACHINE_ARN, + app.SQS_ENDPOINT, + app.CONTINUATION_QUEUE_URL, + app.AWS_REGION, + app.PAGE_SIZE, + app.BATCH_SIZE, + app.SAFETY_BUFFER, + app.DLQ_URL, + ) + + def test_batch_continuation_success( + self, monthly_reports_continuation_app_with_mocks + ): + mock_event = { + "Records": [ + { + "body": json.dumps( + { + "scan_params": { + "ProjectionExpression": "accountId, userId" + }, + "statement_period": "2024-01", + "remaining_accounts": [ + {"accountId": "acc1", "userId": "user1"}, + {"accountId": "acc2", "userId": "user2"}, + ], + "last_evaluated_key": {"accountId": "acc_last"}, + } + ), + "messageAttributes": { + "continuation_type": {"stringValue": "batch_continuation"} + }, + } + ] + } + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 300000 + + with patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.process_batch_continuation" + ) as mock_process_batch: + mock_process_batch.return_value = { + "processed_count": 2, + "failed_starts_count": 0, + "skipped_count": 0, + "already_exists_count": 0, + "batches_processed": 1, + "status": "COMPLETED", + } + + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 200 + body = ( + response["body"] + if isinstance(response["body"], dict) + else json.loads(response["body"]) + ) + assert body["message"] == "Monthly Account reports processing completed" + assert body["processed_count"] == 2 + + mock_process_batch.assert_called_once_with( + {"ProjectionExpression": "accountId, userId"}, + "2024-01", + [ + {"accountId": "acc1", "userId": "user1"}, + {"accountId": "acc2", "userId": "user2"}, + ], + {"accountId": "acc_last"}, + mock_context, + app.logger, + app.accounts_table, + app.sfn_client, + app.STATE_MACHINE_ARN, + app.SQS_ENDPOINT, + app.CONTINUATION_QUEUE_URL, + app.AWS_REGION, + app.PAGE_SIZE, + app.BATCH_SIZE, + app.SAFETY_BUFFER, + app.DLQ_URL, + ) + + def test_batch_continuation_without_last_evaluated_key( + self, monthly_reports_continuation_app_with_mocks + ): + mock_event = { + "Records": [ + { + "body": json.dumps( + { + "scan_params": { + "ProjectionExpression": "accountId, userId" + }, + "statement_period": "2024-01", + "remaining_accounts": [ + {"accountId": "acc1", "userId": "user1"} + ], + } + ), + "messageAttributes": { + "continuation_type": {"stringValue": "batch_continuation"} + }, + } + ] + } + mock_context = MagicMock() + + with patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.process_batch_continuation" + ) as mock_process_batch: + mock_process_batch.return_value = { + "processed_count": 1, + "failed_starts_count": 0, + "skipped_count": 0, + "already_exists_count": 0, + "batches_processed": 1, + "status": "COMPLETED", + } + + lambda_handler(mock_event, mock_context) + + mock_process_batch.assert_called_once() + call_args = mock_process_batch.call_args[0] + assert call_args[3] is None + + def test_unknown_continuation_type( + self, monthly_reports_continuation_app_with_mocks + ): + mock_event = { + "Records": [ + { + "body": json.dumps( + { + "scan_params": { + "ProjectionExpression": "accountId, userId" + }, + "statement_period": "2024-01", + } + ), + "messageAttributes": { + "continuation_type": {"stringValue": "unknown_type"} + }, + } + ] + } + mock_context = MagicMock() + + with patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.logger" + ) as mock_logger: + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 200 + mock_logger.warning.assert_called_once_with( + "Unknown continuation type: unknown_type" + ) + + def test_missing_continuation_type( + self, monthly_reports_continuation_app_with_mocks + ): + mock_event = { + "Records": [ + { + "body": json.dumps( + { + "scan_params": { + "ProjectionExpression": "accountId, userId" + }, + "statement_period": "2024-01", + } + ), + "messageAttributes": {}, + } + ] + } + mock_context = MagicMock() + + with patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.logger" + ) as mock_logger: + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 200 + mock_logger.warning.assert_called_once_with( + "Unknown continuation type: None" + ) + + def test_multiple_sqs_records(self, monthly_reports_continuation_app_with_mocks): + mock_event = { + "Records": [ + { + "body": json.dumps( + { + "scan_params": { + "ProjectionExpression": "accountId, userId" + }, + "statement_period": "2024-01", + } + ), + "messageAttributes": { + "continuation_type": {"stringValue": "accounts_scan"} + }, + }, + { + "body": json.dumps( + { + "scan_params": { + "ProjectionExpression": "accountId, userId" + }, + "statement_period": "2024-01", + "remaining_accounts": [ + {"accountId": "acc1", "userId": "user1"} + ], + } + ), + "messageAttributes": { + "continuation_type": {"stringValue": "batch_continuation"} + }, + }, + ] + } + mock_context = MagicMock() + + with patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.process_accounts_scan_continuation" + ) as mock_process_scan, patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.process_batch_continuation" + ) as mock_process_batch: + mock_process_scan.return_value = { + "processed_count": 3, + "failed_starts_count": 0, + "skipped_count": 0, + "already_exists_count": 0, + "batches_processed": 1, + "pages_processed": 1, + "status": "COMPLETED", + } + mock_process_batch.return_value = { + "processed_count": 1, + "failed_starts_count": 0, + "skipped_count": 0, + "already_exists_count": 0, + "batches_processed": 1, + "status": "COMPLETED", + } + + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 200 + body = ( + response["body"] + if isinstance(response["body"], dict) + else json.loads(response["body"]) + ) + assert body["processed_count"] == 4 + assert body["batches_processed"] == 2 + + mock_process_scan.assert_called_once() + mock_process_batch.assert_called_once() + + def test_empty_records_list(self, monthly_reports_continuation_app_with_mocks): + mock_event = {"Records": []} + mock_context = MagicMock() + + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 200 + body = ( + response["body"] + if isinstance(response["body"], dict) + else json.loads(response["body"]) + ) + assert body["processed_count"] == 0 + assert body["batches_processed"] == 0 + + def test_missing_records_key(self, monthly_reports_continuation_app_with_mocks): + mock_event = {} + mock_context = MagicMock() + + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 200 + body = ( + response["body"] + if isinstance(response["body"], dict) + else json.loads(response["body"]) + ) + assert body["processed_count"] == 0 + + def test_critical_error_during_processing_raises_exception( + self, monthly_reports_continuation_app_with_mocks + ): + mock_event = { + "Records": [ + { + "body": json.dumps( + { + "scan_params": { + "ProjectionExpression": "accountId, userId" + }, + "statement_period": "2024-01", + } + ), + "messageAttributes": { + "continuation_type": {"stringValue": "accounts_scan"} + }, + } + ] + } + mock_context = MagicMock() + + with patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.process_accounts_scan_continuation" + ) as mock_process_scan: + mock_process_scan.side_effect = Exception("Processing failed") + + with pytest.raises(Exception, match="Processing failed"): + lambda_handler(mock_event, mock_context) + + def test_invalid_json_in_message_body_raises_exception( + self, monthly_reports_continuation_app_with_mocks + ): + mock_event = { + "Records": [ + { + "body": "invalid json", + "messageAttributes": { + "continuation_type": {"stringValue": "accounts_scan"} + }, + } + ] + } + mock_context = MagicMock() + + with patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.logger" + ) as mock_logger, patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.send_bad_account_to_dlq" + ) as mock_send_dlq: + lambda_handler(mock_event, mock_context) + + # Verify the JSON parsing error was logged + mock_logger.error.assert_any_call( + "Failed to parse message body as JSON: Expecting value: line 1 column 1 (char 0)" + ) + + # Verify DLQ function was called + mock_send_dlq.assert_called_once() + + # Check the DLQ call arguments + call_args = mock_send_dlq.call_args + error_data = call_args[0][0] + assert error_data["error_type"] == "json_parse_error" + assert error_data["raw_message"] == "invalid json" + + def test_missing_accounts_table_name(self, monkeypatch): + monkeypatch.delenv("ACCOUNTS_TABLE_NAME", raising=False) + monkeypatch.setenv( + "CONTINUATION_QUEUE_URL", + "https://sqs.eu-west-2.amazonaws.com/123456789012/continuation-queue", + ) + reload(app) + + mock_event = { + "Records": [ + { + "body": json.dumps( + { + "scan_params": { + "ProjectionExpression": "accountId, userId" + }, + "statement_period": "2024-01", + } + ), + "messageAttributes": { + "continuation_type": {"stringValue": "accounts_scan"} + }, + } + ] + } + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 300000 + + with pytest.raises(Exception): + lambda_handler(mock_event, mock_context) + + def test_environment_variables_initialization(self, monkeypatch): + monkeypatch.setenv("ENVIRONMENT_NAME", "production") + monkeypatch.setenv("POWERTOOLS_LOG_LEVEL", "DEBUG") + monkeypatch.setenv("AWS_REGION", "us-east-1") + reload(app) + + assert app.ENVIRONMENT_NAME == "production" + assert app.POWERTOOLS_LOG_LEVEL == "DEBUG" + assert app.AWS_REGION == "us-east-1" + + def test_constants_are_set(self): + assert app.PAGE_SIZE == 50 + assert app.BATCH_SIZE == 10 + assert app.SAFETY_BUFFER == 30 + + def test_lambda_handler_logger_injection( + self, monthly_reports_continuation_app_with_mocks + ): + mock_event = {"Records": []} + mock_context = MagicMock() + + with patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.logger" + ) as mock_logger: + response = lambda_handler(mock_event, mock_context) + + mock_logger.info.assert_any_call("Processing SQS continuation messages") + assert response["statusCode"] == 200 + + def test_invalid_json_with_dlq_send_success( + self, monthly_reports_continuation_app_with_mocks + ): + mock_event = { + "Records": [ + { + "body": "invalid json", + "messageId": "test-message-id", + "messageAttributes": { + "continuation_type": {"stringValue": "accounts_scan"} + }, + } + ] + } + mock_context = MagicMock() + + with patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.logger" + ) as mock_logger, patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.send_bad_account_to_dlq" + ) as mock_send_dlq: + + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 200 + mock_logger.error.assert_called_once() + mock_send_dlq.assert_called_once() + + def test_invalid_json_with_dlq_send_failure( + self, monthly_reports_continuation_app_with_mocks + ): + mock_event = { + "Records": [ + { + "body": "invalid json", + "messageId": "test-message-id", + "messageAttributes": { + "continuation_type": {"stringValue": "accounts_scan"} + }, + } + ] + } + mock_context = MagicMock() + + with patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.logger" + ) as mock_logger, patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.send_bad_account_to_dlq" + ) as mock_send_dlq: + mock_send_dlq.side_effect = Exception("DLQ send failed") + + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 200 + mock_logger.error.assert_any_call( + "Failed to parse message body as JSON: Expecting value: line 1 column 1 (char 0)" + ) + mock_logger.error.assert_any_call( + "Failed to send parse error to DLQ: DLQ send failed" + ) + + def test_unknown_continuation_type_with_dlq_send_success( + self, monthly_reports_continuation_app_with_mocks + ): + mock_event = { + "Records": [ + { + "body": json.dumps( + { + "scan_params": { + "ProjectionExpression": "accountId, userId" + }, + "statement_period": "2024-01", + } + ), + "messageId": "test-message-id", + "messageAttributes": { + "continuation_type": {"stringValue": "unknown_type"} + }, + } + ] + } + mock_context = MagicMock() + + with patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.logger" + ) as mock_logger, patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.send_bad_account_to_dlq" + ) as mock_send_dlq: + + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 200 + mock_logger.warning.assert_called_once_with( + "Unknown continuation type: unknown_type" + ) + mock_send_dlq.assert_called_once() + + def test_unknown_continuation_type_with_dlq_send_failure( + self, monthly_reports_continuation_app_with_mocks + ): + mock_event = { + "Records": [ + { + "body": json.dumps( + { + "scan_params": { + "ProjectionExpression": "accountId, userId" + }, + "statement_period": "2024-01", + } + ), + "messageId": "test-message-id", + "messageAttributes": { + "continuation_type": {"stringValue": "unknown_type"} + }, + } + ] + } + mock_context = MagicMock() + + with patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.logger" + ) as mock_logger, patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.send_bad_account_to_dlq" + ) as mock_send_dlq: + mock_send_dlq.side_effect = Exception("DLQ send failed") + + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 200 + mock_logger.warning.assert_called_once_with( + "Unknown continuation type: unknown_type" + ) + mock_logger.error.assert_called_once_with( + "Failed to send unknown continuation type to DLQ: DLQ send failed" + ) + + def test_critical_error_with_dlq_send_success( + self, monthly_reports_continuation_app_with_mocks + ): + mock_event = { + "Records": [ + { + "body": json.dumps( + { + "scan_params": { + "ProjectionExpression": "accountId, userId" + }, + "statement_period": "2024-01", + } + ), + "messageAttributes": { + "continuation_type": {"stringValue": "accounts_scan"} + }, + } + ] + } + mock_context = MagicMock() + + with patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.process_accounts_scan_continuation" + ) as mock_process_scan, patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.logger" + ) as mock_logger, patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.send_bad_account_to_dlq" + ) as mock_send_dlq: + mock_process_scan.side_effect = Exception("Processing failed") + + with pytest.raises(Exception, match="Processing failed"): + lambda_handler(mock_event, mock_context) + + mock_logger.error.assert_called_once() + mock_send_dlq.assert_called_once() + + def test_critical_error_with_dlq_send_failure( + self, monthly_reports_continuation_app_with_mocks + ): + mock_event = { + "Records": [ + { + "body": json.dumps( + { + "scan_params": { + "ProjectionExpression": "accountId, userId" + }, + "statement_period": "2024-01", + } + ), + "messageAttributes": { + "continuation_type": {"stringValue": "accounts_scan"} + }, + } + ] + } + mock_context = MagicMock() + + with patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.process_accounts_scan_continuation" + ) as mock_process_scan, patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.logger" + ) as mock_logger, patch( + "functions.monthly_reports.accounts.process_pending_reports.process_pending_reports.app.send_bad_account_to_dlq" + ) as mock_send_dlq: + mock_process_scan.side_effect = Exception("Processing failed") + mock_send_dlq.side_effect = Exception("DLQ send failed") + + with pytest.raises(Exception, match="Processing failed"): + lambda_handler(mock_event, mock_context) + + mock_logger.error.assert_any_call( + "Critical error during continuation processing: Processing failed", + exc_info=True, + ) + mock_logger.error.assert_any_call( + "Failed to send critical error to DLQ: DLQ send failed" + ) diff --git a/tests/functions/monthly_reports/accounts/trigger_tests/__init__.py b/tests/functions/monthly_reports/accounts/trigger_tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/functions/monthly_reports/accounts/trigger_tests/conftest.py b/tests/functions/monthly_reports/accounts/trigger_tests/conftest.py new file mode 100644 index 0000000..b2f1c0a --- /dev/null +++ b/tests/functions/monthly_reports/accounts/trigger_tests/conftest.py @@ -0,0 +1,33 @@ +from importlib import reload +from functions.monthly_reports.accounts.trigger.trigger import app +from unittest.mock import patch + +import pytest + + +@pytest.fixture(scope="function") +def monthly_accounts_reports_app_with_mocks( + monkeypatch, dynamo_resource, mock_accounts_dynamo_table +): + accounts_table_name = mock_accounts_dynamo_table + + monkeypatch.setenv("ACCOUNTS_TABLE_NAME", accounts_table_name) + monkeypatch.setenv( + "CONTINUATION_QUEUE_URL", + "https://sqs.eu-west-2.amazonaws.com/123456789012/continuation-queue", + ) + monkeypatch.setenv("STATE_MACHINE_ARN", "mock_arn") + monkeypatch.setenv("ENVIRONMENT_NAME", "test") + monkeypatch.setenv("POWERTOOLS_LOG_LEVEL", "INFO") + monkeypatch.setenv("AWS_REGION", "eu-west-2") + monkeypatch.setenv( + "DLQ_URL", "https://sqs.eu-west-2.amazonaws.com/123456789012/dlq" + ) + monkeypatch.setenv("SQS_ENDPOINT", "https://sqs.eu-west-2.amazonaws.com") + + with patch("boto3.resource", return_value=dynamo_resource): + reload(app) + + app.accounts_table = dynamo_resource.Table(accounts_table_name) + + yield app diff --git a/tests/functions/monthly_reports/accounts/trigger_tests/test_app.py b/tests/functions/monthly_reports/accounts/trigger_tests/test_app.py new file mode 100644 index 0000000..bff0bdc --- /dev/null +++ b/tests/functions/monthly_reports/accounts/trigger_tests/test_app.py @@ -0,0 +1,313 @@ +import json +from importlib import reload +from unittest.mock import MagicMock, patch + +import pytest + +from functions.monthly_reports.accounts.trigger.trigger import app +from functions.monthly_reports.accounts.trigger.trigger.app import ( + lambda_handler, +) + + +class TestLambdaHandler: + def test_success(self, monthly_accounts_reports_app_with_mocks): + mock_event = {} + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 300000 + + with patch( + "functions.monthly_reports.accounts.trigger.trigger.app.get_paginated_table_data" + ) as mock_get_data, patch( + "functions.monthly_reports.accounts.trigger.trigger.app.process_accounts_page" + ) as mock_process_page: + mock_get_data.return_value = [ + {"accountId": "acc1", "userId": "user1"}, + {"accountId": "acc2", "userId": "user2"}, + {"accountId": "acc3", "userId": "user3"}, + ], {} + + mock_process_page.return_value = { + "processed_count": 3, + "failed_starts_count": 0, + "skipped_count": 0, + "already_exists_count": 0, + "batches_processed": 1, + "status": "COMPLETED", + } + + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 200 + # The response body is already a dict, not a JSON string + body = ( + response["body"] + if isinstance(response["body"], dict) + else json.loads(response["body"]) + ) + assert body["message"] == "Monthly Account reports processing completed" + assert body["processed_count"] == 3 + assert body["batches_processed"] == 1 + + def test_timeout_warning( + self, monthly_accounts_reports_app_with_mocks, mock_logger + ): + mock_event = {} + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 10 + + with patch( + "functions.monthly_reports.accounts.trigger.trigger.app.logger" + ) as mock_logger_instance: + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 202 + body = ( + response["body"] + if isinstance(response["body"], dict) + else json.loads(response["body"]) + ) + assert ( + body["message"] + == "Monthly Account reports processing timeout_continuation" + ) + assert body["processed_count"] == 0 + assert body["pages_processed"] == 0 + + mock_logger_instance.warning.assert_called_once() + warning_call_args = mock_logger_instance.warning.call_args[0][0] + assert "Approaching Lambda timeout" in warning_call_args + + def test_critical_error_during_data_retrieval_raises_exception( + self, monthly_accounts_reports_app_with_mocks + ): + mock_event = {} + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 300000 + + with patch( + "functions.monthly_reports.accounts.trigger.trigger.app.get_paginated_table_data" + ) as mock_get_data: + mock_get_data.side_effect = Exception("Database connection failed") + + with pytest.raises(Exception, match="Database connection failed"): + lambda_handler(mock_event, mock_context) + + def test_page_processing_with_timeout_continuation( + self, monthly_accounts_reports_app_with_mocks + ): + mock_event = {} + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 10 + + with patch( + "functions.monthly_reports.accounts.trigger.trigger.app.get_paginated_table_data" + ) as mock_get_data, patch( + "functions.monthly_reports.accounts.trigger.trigger.app.process_accounts_page" + ) as mock_process_page, patch( + "functions.monthly_reports.accounts.trigger.trigger.app.send_continuation_message" + ) as mock_send_continuation: + mock_get_data.return_value = [ + {"accountId": f"acc{i}", "userId": f"user{i}"} for i in range(1, 11) + ], {} + + mock_process_page.return_value = { + "processed_count": 5, + "failed_starts_count": 0, + "skipped_count": 0, + "already_exists_count": 0, + "batches_processed": 1, + "status": "TIMEOUT_CONTINUATION", + "remaining_accounts": [ + {"accountId": f"acc{i}", "userId": f"user{i}"} for i in range(6, 11) + ], + } + + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 202 + body = ( + response["body"] + if isinstance(response["body"], dict) + else json.loads(response["body"]) + ) + assert body["status"] == "TIMEOUT_CONTINUATION" + + mock_send_continuation.assert_called_once() + + def test_no_accounts_pages_left(self, monthly_accounts_reports_app_with_mocks): + mock_event = {} + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 300000 + + with patch( + "functions.monthly_reports.accounts.trigger.trigger.app.get_paginated_table_data" + ) as mock_get_data, patch( + "functions.monthly_reports.accounts.trigger.trigger.app.process_accounts_page" + ) as mock_process_page, patch( + "functions.monthly_reports.accounts.trigger.trigger.app.logger" + ) as mock_logger: + mock_get_data.side_effect = [ + ( + [ + {"accountId": f"acc{i}", "userId": f"user{i}"} + for i in range(1, 11) + ], + {"LastEvaluatedKey": "some_key"}, + ), + ([], {}), + ] + + mock_process_page.return_value = { + "processed_count": 10, + "failed_starts_count": 0, + "skipped_count": 0, + "already_exists_count": 0, + "batches_processed": 1, + "status": "COMPLETED", + } + + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 200 + body = ( + response["body"] + if isinstance(response["body"], dict) + else json.loads(response["body"]) + ) + assert body["processed_count"] == 10 + assert body["batches_processed"] == 1 + + mock_logger.info.assert_any_call("No more accounts to process") + + def test_multiple_pages_processing(self, monthly_accounts_reports_app_with_mocks): + mock_event = {} + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 300000 + + with patch( + "functions.monthly_reports.accounts.trigger.trigger.app.get_paginated_table_data" + ) as mock_get_data, patch( + "functions.monthly_reports.accounts.trigger.trigger.app.process_accounts_page" + ) as mock_process_page: + # Simulate multiple pages + mock_get_data.side_effect = [ + ( + [ + {"accountId": f"acc{i}", "userId": f"user{i}"} + for i in range(1, 6) + ], + {"LastEvaluatedKey": "key1"}, + ), + ( + [ + {"accountId": f"acc{i}", "userId": f"user{i}"} + for i in range(6, 11) + ], + {}, + ), + ] + + mock_process_page.side_effect = [ + { + "processed_count": 5, + "failed_starts_count": 0, + "skipped_count": 0, + "already_exists_count": 0, + "batches_processed": 1, + "status": "COMPLETED", + }, + { + "processed_count": 5, + "failed_starts_count": 0, + "skipped_count": 0, + "already_exists_count": 0, + "batches_processed": 1, + "status": "COMPLETED", + }, + ] + + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 200 + body = ( + response["body"] + if isinstance(response["body"], dict) + else json.loads(response["body"]) + ) + assert body["processed_count"] == 10 + assert body["batches_processed"] == 2 + assert body["pages_processed"] == 2 + + def test_missing_queue_url( + self, monthly_accounts_reports_app_with_mocks, monkeypatch + ): + monkeypatch.delenv("CONTINUATION_QUEUE_URL", raising=False) + reload(app) + + mock_event = {} + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 300000 + + response = lambda_handler(mock_event, mock_context) + + assert response["statusCode"] == 500 + body = ( + response["body"] + if isinstance(response["body"], dict) + else json.loads(response["body"]) + ) + assert ( + body["message"] + == "Monthly Account reports processing error_no_continuation_queue" + ) + + def test_critical_error_with_dlq_send_success( + self, monthly_accounts_reports_app_with_mocks + ): + mock_event = {} + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 300000 + + with patch( + "functions.monthly_reports.accounts.trigger.trigger.app.get_paginated_table_data" + ) as mock_get_data, patch( + "functions.monthly_reports.accounts.trigger.trigger.app.logger" + ) as mock_logger, patch( + "functions.monthly_reports.accounts.trigger.trigger.app.send_bad_account_to_dlq" + ) as mock_send_dlq: + mock_get_data.side_effect = Exception("Database connection failed") + + with pytest.raises(Exception, match="Database connection failed"): + lambda_handler(mock_event, mock_context) + + mock_logger.error.assert_called_once() + mock_send_dlq.assert_called_once() + + def test_critical_error_with_dlq_send_failure( + self, monthly_accounts_reports_app_with_mocks + ): + mock_event = {} + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 300000 + + with patch( + "functions.monthly_reports.accounts.trigger.trigger.app.get_paginated_table_data" + ) as mock_get_data, patch( + "functions.monthly_reports.accounts.trigger.trigger.app.logger" + ) as mock_logger, patch( + "functions.monthly_reports.accounts.trigger.trigger.app.send_bad_account_to_dlq" + ) as mock_send_dlq: + mock_get_data.side_effect = Exception("Database connection failed") + mock_send_dlq.side_effect = Exception("DLQ send failed") + + with pytest.raises(Exception, match="Database connection failed"): + lambda_handler(mock_event, mock_context) + + mock_logger.error.assert_any_call( + "Critical error during processing: Database connection failed", + exc_info=True, + ) + mock_logger.error.assert_any_call( + "Failed to send critical error to DLQ: DLQ send failed" + ) diff --git a/tests/functions/transactions/process_transactions/test_app.py b/tests/functions/transactions/process_transactions/test_app.py index 58be16c..609d3d5 100644 --- a/tests/functions/transactions/process_transactions/test_app.py +++ b/tests/functions/transactions/process_transactions/test_app.py @@ -211,7 +211,7 @@ def test_business_logic_error_with_idempotency_key( "functions.transactions.process_transactions.process_transactions.app.process_single_transaction" ) @patch( - "functions.transactions.process_transactions.process_transactions.app.send_dynamodb_record_to_dlq" + "functions.transactions.process_transactions.process_transactions.app.send_message_to_sqs" ) def test_business_logic_error_without_idempotency_key( self, @@ -268,7 +268,7 @@ def test_business_logic_error_without_idempotency_key( "functions.transactions.process_transactions.process_transactions.app.process_single_transaction" ) @patch( - "functions.transactions.process_transactions.process_transactions.app.send_dynamodb_record_to_dlq", + "functions.transactions.process_transactions.process_transactions.app.send_message_to_sqs", return_value=False, ) def test_error_without_idempotency_key_and_dlq_fails( @@ -315,7 +315,7 @@ def test_error_without_idempotency_key_and_dlq_fails( "functions.transactions.process_transactions.process_transactions.app.update_transaction_status" ) @patch( - "functions.transactions.process_transactions.process_transactions.app.send_dynamodb_record_to_dlq" + "functions.transactions.process_transactions.process_transactions.app.send_message_to_sqs" ) def test_business_logic_error_and_update_status_fails( self, @@ -327,9 +327,9 @@ def test_business_logic_error_and_update_status_fails( environment_variables, ): """ - Test that when a business logic error occurs and updating transaction status fails, the record is sent to the DLQ and the handler reports a business logic failure. + Test that when a business logic error occurs and updating transaction status fails, the record is sent to the DLQ and the handler monthly_reports a business logic failure. - Simulates `process_single_transaction` raising a `BusinessLogicError`, `update_transaction_status` raising an exception, and `send_dynamodb_record_to_dlq` succeeding. Verifies the handler returns a 200 response with one business logic failure and that the DLQ function is called once. + Simulates `process_single_transaction` raising a `BusinessLogicError`, `update_transaction_status` raising an exception, and `send_message_to_sqs` succeeding. Verifies the handler returns a 200 response with one business logic failure and that the DLQ function is called once. """ mock_process_single_transaction.side_effect = BusinessLogicError( "Test business logic error" @@ -363,7 +363,7 @@ def test_business_logic_error_and_update_status_fails( "functions.transactions.process_transactions.process_transactions.app.update_transaction_status" ) @patch( - "functions.transactions.process_transactions.process_transactions.app.send_dynamodb_record_to_dlq" + "functions.transactions.process_transactions.process_transactions.app.send_message_to_sqs" ) def test_business_logic_error_and_dlq_fails( self, @@ -405,7 +405,7 @@ def test_business_logic_error_and_dlq_fails( "functions.transactions.process_transactions.process_transactions.app.process_single_transaction" ) @patch( - "functions.transactions.process_transactions.process_transactions.app.send_dynamodb_record_to_dlq" + "functions.transactions.process_transactions.process_transactions.app.send_message_to_sqs" ) def test_transaction_system_error( self, @@ -446,7 +446,7 @@ def test_transaction_system_error( "functions.transactions.process_transactions.process_transactions.app.process_single_transaction" ) @patch( - "functions.transactions.process_transactions.process_transactions.app.send_dynamodb_record_to_dlq" + "functions.transactions.process_transactions.process_transactions.app.send_message_to_sqs" ) def test_transaction_system_error_and_dlq_fails( self, @@ -486,7 +486,7 @@ def test_transaction_system_error_and_dlq_fails( "functions.transactions.process_transactions.process_transactions.app.process_single_transaction" ) @patch( - "functions.transactions.process_transactions.process_transactions.app.send_dynamodb_record_to_dlq" + "functions.transactions.process_transactions.process_transactions.app.send_message_to_sqs" ) def test_lambda_handler_generic_exception( self, @@ -527,7 +527,7 @@ def test_lambda_handler_generic_exception( "functions.transactions.process_transactions.process_transactions.app.process_single_transaction" ) @patch( - "functions.transactions.process_transactions.process_transactions.app.send_dynamodb_record_to_dlq" + "functions.transactions.process_transactions.process_transactions.app.send_message_to_sqs" ) def test_generic_exception_and_dlq_fails( self, @@ -568,7 +568,7 @@ def test_generic_exception_and_dlq_fails( "functions.transactions.process_transactions.process_transactions.app.update_transaction_status" ) @patch( - "functions.transactions.process_transactions.process_transactions.app.send_dynamodb_record_to_dlq" + "functions.transactions.process_transactions.process_transactions.app.send_message_to_sqs" ) def test_lambda_handler_success_and_failure( self, diff --git a/tests/functions/transactions/process_transactions/test_sqs.py b/tests/functions/transactions/process_transactions/test_sqs.py new file mode 100644 index 0000000..ddee475 --- /dev/null +++ b/tests/functions/transactions/process_transactions/test_sqs.py @@ -0,0 +1,16 @@ +import pytest + +from functions.transactions.process_transactions.process_transactions.sqs import ( + format_sqs_message, +) + + +class TestSqsHelpers: + + def test_format_sqs_message_incorrect_type(self): + + with pytest.raises(ValueError) as exception_info: + format_sqs_message("", "") + + assert exception_info.type is ValueError + assert exception_info.value.args[0] == "Record must be a dictionary" diff --git a/tests/layers/authentication/test_user_details.py b/tests/layers/authentication/test_user_details.py new file mode 100644 index 0000000..8528fca --- /dev/null +++ b/tests/layers/authentication/test_user_details.py @@ -0,0 +1,127 @@ +import uuid +from unittest.mock import patch + +import pytest + +from authentication.user_details import get_user_attributes +from tests.layers.authentication.conftest import TEST_AWS_REGION, TEST_USER_POOL_ID +from botocore.exceptions import ClientError + + +class TestUserDetails: + + def test_success(self, mock_cognito_client, mock_logger): + username = "test_user" + expected_attributes = { + "sub": str(uuid.uuid4()), + "email": "test@example.com", + "name": "Test User", + "email_verified": "true", + } + + mock_response = { + "UserAttributes": [ + {"Name": "sub", "Value": expected_attributes["sub"]}, + {"Name": "email", "Value": expected_attributes["email"]}, + {"Name": "name", "Value": expected_attributes["name"]}, + { + "Name": "email_verified", + "Value": expected_attributes["email_verified"], + }, + ] + } + + mock_cognito_client.admin_get_user.return_value = mock_response + + with patch("authentication.user_details.boto3.client") as mock_boto3_client: + mock_boto3_client.return_value = mock_cognito_client + + result = get_user_attributes( + aws_region=TEST_AWS_REGION, + logger=mock_logger, + username=username, + user_pool_id=TEST_USER_POOL_ID, + ) + + assert result == expected_attributes + mock_cognito_client.admin_get_user.assert_called_once_with( + UserPoolId=TEST_USER_POOL_ID, Username=username + ) + mock_logger.info.assert_called_once_with( + f"Fetched attributes for user: {username}." + ) + mock_boto3_client.assert_called_once_with( + "cognito-idp", region_name=TEST_AWS_REGION + ) + + def test_cognito_exception(self, mock_cognito_client, mock_logger): + username = "test_user" + expected_exception = Exception("Cognito service error") + + mock_cognito_client.admin_get_user.side_effect = expected_exception + + with patch("authentication.user_details.boto3.client") as mock_boto3_client: + mock_boto3_client.return_value = mock_cognito_client + + with pytest.raises(Exception) as exception_info: + get_user_attributes( + aws_region=TEST_AWS_REGION, + logger=mock_logger, + username=username, + user_pool_id=TEST_USER_POOL_ID, + ) + + assert exception_info.value == expected_exception + mock_logger.exception.assert_called_once_with( + f"Failed to fetch user {username} from Cognito" + ) + mock_cognito_client.admin_get_user.assert_called_once_with( + UserPoolId=TEST_USER_POOL_ID, Username=username + ) + + def test_empty_user_attributes(self, mock_cognito_client, mock_logger): + username = "test_user" + mock_response = {"UserAttributes": []} + + mock_cognito_client.admin_get_user.return_value = mock_response + + with patch("authentication.user_details.boto3.client") as mock_boto3_client: + mock_boto3_client.return_value = mock_cognito_client + + result = get_user_attributes( + aws_region=TEST_AWS_REGION, + logger=mock_logger, + username=username, + user_pool_id=TEST_USER_POOL_ID, + ) + + assert result == {} + mock_logger.info.assert_called_once_with( + f"Fetched attributes for user: {username}." + ) + + def test_specific_cognito_exceptions(self, mock_cognito_client, mock_logger): + username = "test_user" + + error_response = { + "Error": {"Code": "UserNotFoundException", "Message": "User does not exist"} + } + user_not_found_exception = ClientError(error_response, "AdminGetUser") + + mock_cognito_client.admin_get_user.side_effect = user_not_found_exception + + with patch("authentication.user_details.boto3.client") as mock_boto3_client: + mock_boto3_client.return_value = mock_cognito_client + + with pytest.raises(ClientError) as exception_info: + get_user_attributes( + aws_region=TEST_AWS_REGION, + logger=mock_logger, + username=username, + user_pool_id=TEST_USER_POOL_ID, + ) + + assert exception_info.value == user_not_found_exception + mock_logger.exception.assert_called_once_with( + f"Failed to fetch user {username} from Cognito" + ) diff --git a/tests/layers/helpers/test_dynamodb.py b/tests/layers/helpers/test_dynamodb.py index 8b86715..6660de9 100644 --- a/tests/layers/helpers/test_dynamodb.py +++ b/tests/layers/helpers/test_dynamodb.py @@ -1,8 +1,10 @@ +import uuid from unittest.mock import patch, MagicMock import pytest +from botocore.exceptions import ClientError -from dynamodb import get_dynamodb_resource +from dynamodb import get_dynamodb_resource, get_paginated_table_data class TestGetDynamoDBResource: @@ -66,3 +68,72 @@ def test_get_dynamodb_resource_error_handling(self): mock_logger.error.assert_called_once_with( "Failed to initialize DynamoDB resource", exc_info=True ) + + +class TestGetPaginatedTableData: + + def test_success(self, magic_mock_accounts_table, mock_logger): + item_id = str(uuid.uuid4()) + magic_mock_accounts_table.scan.return_value = {"Items": [{"id": item_id}]} + + result = get_paginated_table_data( + None, None, magic_mock_accounts_table, mock_logger + ) + + assert result[0] == [{"id": item_id}] + + def test_success_with_scan_params(self, magic_mock_accounts_table, mock_logger): + item_id = str(uuid.uuid4()) + magic_mock_accounts_table.scan.return_value = {"Items": [{"id": item_id}]} + + result = get_paginated_table_data( + { + "ProjectionExpression": "accountId, userId", + }, + None, + magic_mock_accounts_table, + mock_logger, + ) + + assert result[0] == [{"id": item_id}] + assert magic_mock_accounts_table.scan.call_args[1] == { + "ProjectionExpression": "accountId, userId", + "Limit": 10, + } + + def test_success_with_index(self, magic_mock_accounts_table, mock_logger): + item_id = str(uuid.uuid4()) + magic_mock_accounts_table.scan.return_value = {"Items": [{"id": item_id}]} + + result = get_paginated_table_data( + None, "id", magic_mock_accounts_table, mock_logger + ) + + assert result[0] == [{"id": item_id}] + assert magic_mock_accounts_table.scan.call_args[1] == { + "IndexName": "id", + "Limit": 10, + } + + def test_error(self, magic_mock_accounts_table, mock_logger): + magic_mock_accounts_table.scan.side_effect = ClientError( + operation_name="scan", + error_response={ + "Error": { + "Code": "ResourceNotFoundException", + "Message": "Requested resource not found.", + } + }, + ) + + with pytest.raises(Exception) as exception_info: + get_paginated_table_data( + { + "ProjectionExpression": "accountId, userId", + }, + None, + magic_mock_accounts_table, + mock_logger, + ) + + assert "Requested resource not found." in str(exception_info.value) diff --git a/tests/layers/helpers/test_s3.py b/tests/layers/helpers/test_s3.py new file mode 100644 index 0000000..d4b442d --- /dev/null +++ b/tests/layers/helpers/test_s3.py @@ -0,0 +1,39 @@ +from unittest.mock import patch, MagicMock + +import pytest + +from s3 import get_s3_client + + +class TestGetS3Client: + + def test_get_s3_client_success(self): + mock_logger = MagicMock() + region = "eu-west-2" + + with patch("boto3.client") as mock_boto3_client: + mock_client = MagicMock() + mock_boto3_client.return_value = mock_client + + result = get_s3_client(region, mock_logger) + + mock_boto3_client.assert_called_once_with("s3", region_name=region) + assert result == mock_client + mock_logger.info.assert_called_once_with( + "Initialized S3 client with default endpoint" + ) + + def test_get_s3_client_exception(self): + mock_logger = MagicMock() + region = "eu-west-2" + + with patch("boto3.client") as mock_boto3_client: + mock_boto3_client.side_effect = Exception("Connection error") + + with pytest.raises(Exception) as exc_info: + get_s3_client(region, mock_logger) + + assert "Connection error" in str(exc_info.value) + mock_logger.error.assert_called_once_with( + "Failed to initialize S3 client", exc_info=True + ) diff --git a/tests/layers/helpers/test_ses.py b/tests/layers/helpers/test_ses.py index eda7ba4..10b5828 100644 --- a/tests/layers/helpers/test_ses.py +++ b/tests/layers/helpers/test_ses.py @@ -3,7 +3,7 @@ import pytest -from ses import get_ses_client, send_user_email +from ses import get_ses_client, send_user_email, send_user_email_with_attachment class TestGetSesClient: @@ -77,6 +77,10 @@ def test_send_user_email_success(self, mock_ses_client): Verifies that the SES client's send_email method is called with the correct arguments, the logger records a success message, and the function returns True. """ + # Prepare a mock SES response + mock_response = {"MessageId": "test-message-id-123"} + self.mock_ses_client.send_email.return_value = mock_response + result = send_user_email( aws_region=self.aws_region, logger=self.mock_logger, @@ -113,46 +117,127 @@ def test_send_user_email_success(self, mock_ses_client): ) self.mock_logger.info.assert_called_once_with( - f"Successfully sent email to users: {json.dumps(self.to_addresses)}" + f"Successfully sent email to {json.dumps(self.to_addresses)}, MessageId={mock_response['MessageId']}" ) - assert result is True + assert result == mock_response def test_send_user_email_exception(self, mock_ses_client): mock_exception = Exception("Simulated SES send error") self.mock_ses_client.send_email.side_effect = mock_exception - result = send_user_email( - aws_region=self.aws_region, - logger=self.mock_logger, - sender_email=self.sender_email, - to_addresses=self.to_addresses, - subject_data=self.subject_data, - subject_charset=self.subject_charset, - text_body_data=self.text_body_data, - ) + with pytest.raises(Exception) as exc_info: + send_user_email( + aws_region=self.aws_region, + logger=self.mock_logger, + sender_email=self.sender_email, + to_addresses=self.to_addresses, + subject_data=self.subject_data, + subject_charset=self.subject_charset, + text_body_data=self.text_body_data, + ) + assert "Simulated SES send error" in str(exc_info.value) self.mock_ses_client.send_email.assert_called_once() self.mock_logger.error.assert_called_once_with( - f"Failed to send email: {mock_exception}" + f"Failed to send email: {mock_exception}", exc_info=True ) - assert result is False def test_send_user_email_no_body(self, mock_ses_client): """ Test that send_user_email returns False and logs an error when neither text nor HTML body is provided. """ - result = send_user_email( + with pytest.raises(Exception) as exc_info: + send_user_email( + aws_region=self.aws_region, + logger=self.mock_logger, + sender_email=self.sender_email, + to_addresses=self.to_addresses, + subject_data=self.subject_data, + subject_charset=self.subject_charset, + ) + + self.mock_ses_client.send_email.assert_not_called() + self.mock_logger.error.assert_called_once_with( + "Email must contain at least a text or HTML body." + ) + + assert str(exc_info.value) == "Email must contain at least a text or HTML body." + + +class TestSendEmailWithAttachment: + @pytest.fixture(autouse=True) + def setup(self, mock_get_ses_client): + self.mock_logger = MagicMock() + self.aws_region = "eu-west-2" + self.sender_email = "sender@example.com" + self.to_addresses = ["recipient@example.com"] + self.cc_addresses = ["cc@example.com"] + self.bcc_addresses = ["bcc@example.com"] + self.subject_data = "Monthly Report" + self.body_text = "Please find the report attached." + self.attachment_bytes = b"dummy-bytes" + self.attachment_filename = "report.pdf" + self.mock_get_client, self.mock_ses_client = mock_get_ses_client + + def test_send_user_email_with_attachment_success(self): + mock_response = {"MessageId": "raw-123"} + self.mock_ses_client.send_raw_email.return_value = mock_response + + result = send_user_email_with_attachment( aws_region=self.aws_region, logger=self.mock_logger, sender_email=self.sender_email, to_addresses=self.to_addresses, subject_data=self.subject_data, - subject_charset=self.subject_charset, + body_text=self.body_text, + attachment_bytes=self.attachment_bytes, + attachment_filename=self.attachment_filename, + cc_addresses=self.cc_addresses, + bcc_addresses=self.bcc_addresses, ) - self.mock_ses_client.send_email.assert_not_called() + assert self.mock_ses_client.send_raw_email.call_count == 1 + kwargs = self.mock_ses_client.send_raw_email.call_args.kwargs + + assert kwargs["Source"] == self.sender_email + assert set(kwargs["Destinations"]) == set( + self.to_addresses + self.cc_addresses + self.bcc_addresses + ) + + assert isinstance(kwargs["RawMessage"]["Data"], str) + raw_data = kwargs["RawMessage"]["Data"] + assert self.subject_data in raw_data + assert self.body_text in raw_data + assert f'filename="{self.attachment_filename}"' in raw_data + + assert self.subject_data in raw_data + assert self.body_text in raw_data + assert f'filename="{self.attachment_filename}"' in raw_data + + self.mock_logger.info.assert_called_once_with( + f"Successfully sent email with attachment to {json.dumps(self.to_addresses)}, " + f"MessageId={mock_response['MessageId']}" + ) + assert result == mock_response + + def test_send_user_email_with_attachment_exception(self): + err = Exception("attachment send fail") + self.mock_ses_client.send_raw_email.side_effect = err + + with pytest.raises(Exception) as exc: + send_user_email_with_attachment( + aws_region=self.aws_region, + logger=self.mock_logger, + sender_email=self.sender_email, + to_addresses=self.to_addresses, + subject_data=self.subject_data, + body_text=self.body_text, + attachment_bytes=self.attachment_bytes, + attachment_filename=self.attachment_filename, + ) + + assert "attachment send fail" in str(exc.value) self.mock_logger.error.assert_called_once_with( - "Email must contain at least a text or HTML body." + f"Failed to send email with attachment: {err}", exc_info=True ) - assert result is False diff --git a/tests/layers/helpers/test_sfn.py b/tests/layers/helpers/test_sfn.py new file mode 100644 index 0000000..c1cc442 --- /dev/null +++ b/tests/layers/helpers/test_sfn.py @@ -0,0 +1,41 @@ +from unittest.mock import patch, MagicMock + +import pytest + +from sfn import get_sfn_client + + +class TestGetSfnClient: + + def test_get_sfn_client_success(self): + mock_logger = MagicMock() + region = "eu-west-2" + + with patch("boto3.client") as mock_boto3_client: + mock_client = MagicMock() + mock_boto3_client.return_value = mock_client + + result = get_sfn_client(region, mock_logger) + + mock_boto3_client.assert_called_once_with( + "stepfunctions", region_name=region + ) + assert result == mock_client + mock_logger.info.assert_called_once_with( + "Initialized SFN client with default endpoint" + ) + + def test_get_sfn_client_exception(self): + mock_logger = MagicMock() + region = "eu-west-2" + + with patch("boto3.client") as mock_boto3_client: + mock_boto3_client.side_effect = Exception("Connection error") + + with pytest.raises(Exception) as exc_info: + get_sfn_client(region, mock_logger) + + assert "Connection error" in str(exc_info.value) + mock_logger.error.assert_called_once_with( + "Failed to initialize SFN client", exc_info=True + ) diff --git a/tests/layers/helpers/test_sqs.py b/tests/layers/helpers/test_sqs.py index f58e630..f2254e6 100644 --- a/tests/layers/helpers/test_sqs.py +++ b/tests/layers/helpers/test_sqs.py @@ -2,7 +2,7 @@ import pytest -from sqs import get_sqs_client, send_dynamodb_record_to_dlq +from sqs import get_sqs_client, send_message_to_sqs class TestGetSqsClient: @@ -57,29 +57,34 @@ def test_get_sqs_client_error_handling(self): ) -class TestSendDynamoDbRecordToDLQ: - def test_no_dlq_url(self, mock_sqs_client): - """ - Test that send_dynamodb_record_to_dlq returns False when the DLQ URL is empty. - """ +class TestSendDynamoDbRecordToSQS: + def test_no_sqs_url(self, mock_sqs_client): mock_logger = MagicMock() - result = send_dynamodb_record_to_dlq( - record={}, + result = send_message_to_sqs( + message={}, + message_attributes={}, sqs_endpoint="", - dlq_url="", + sqs_url="", aws_region="", - error_message="", logger=mock_logger, ) assert result is False - def test_send_message_success(self): - """ - Tests that a DynamoDB record is successfully sent to the DLQ and logs the operation. + def test_no_sqs_message(self, mock_sqs_client): + mock_logger = MagicMock() + result = send_message_to_sqs( + message={}, + message_attributes={}, + sqs_endpoint="", + sqs_url="http://localhost:4566/queue/test-queue", + aws_region="", + logger=mock_logger, + ) - Verifies that the SQS client is initialised with the correct parameters, the message is sent, and a success log entry is created. - """ + assert result is False + + def test_send_message_success(self): mock_logger = MagicMock() mock_sqs_client = MagicMock() @@ -89,20 +94,28 @@ def test_send_message_success(self): "SequenceNumber": "123456789012345678901", } } + error_message = "Test error message" + + message = { + "originalRecord": record, + "errorMessage": error_message, + "timestamp": record.get("dynamodb", {}).get("ApproximateCreationDateTime"), + "sequenceNumber": record.get("dynamodb", {}).get("SequenceNumber"), + } + sqs_endpoint = "http://localhost:4566" - dlq_url = "http://localhost:4566/queue/dlq" + sqs_url = "http://localhost:4566/queue/dlq" aws_region = "eu-west-2" - error_message = "Test error message" with patch( "sqs.get_sqs_client", return_value=mock_sqs_client ) as mock_get_client: - result = send_dynamodb_record_to_dlq( - record=record, + result = send_message_to_sqs( + message=message, + message_attributes={}, sqs_endpoint=sqs_endpoint, - dlq_url=dlq_url, + sqs_url=sqs_url, aws_region=aws_region, - error_message=error_message, logger=mock_logger, ) @@ -113,7 +126,7 @@ def test_send_message_success(self): mock_sqs_client.send_message.assert_called_once() mock_logger.info.assert_called_once_with( - f"Successfully sent record to DLQ: {record.get('dynamodb', {}).get('SequenceNumber')}" + "Successfully sent message to SQS queue." ) def test_send_message_failure(self): @@ -129,23 +142,31 @@ def test_send_message_failure(self): "SequenceNumber": "123456789012345678901", } } + error_message = "Test error message" + + message = { + "originalRecord": record, + "errorMessage": error_message, + "timestamp": record.get("dynamodb", {}).get("ApproximateCreationDateTime"), + "sequenceNumber": record.get("dynamodb", {}).get("SequenceNumber"), + } + sqs_endpoint = "http://localhost:4566" - dlq_url = "http://localhost:4566/queue/dlq" + sqs_url = "http://localhost:4566/queue/dlq" aws_region = "eu-west-2" - error_message = "Test error message" mock_sqs_client.send_message.side_effect = Exception("Connection error") with patch("sqs.get_sqs_client", return_value=mock_sqs_client): - result = send_dynamodb_record_to_dlq( - record=record, + result = send_message_to_sqs( + message=message, + message_attributes={}, sqs_endpoint=sqs_endpoint, - dlq_url=dlq_url, + sqs_url=sqs_url, aws_region=aws_region, - error_message=error_message, logger=mock_logger, ) assert result is False mock_logger.error.assert_called_once_with( - "Failed to send message to DLQ: Connection error" + "Failed to send message to SQS: Connection error" ) diff --git a/tests/layers/monthly_reports/__init__.py b/tests/layers/monthly_reports/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/layers/monthly_reports/test_helpers.py b/tests/layers/monthly_reports/test_helpers.py new file mode 100644 index 0000000..6c0127b --- /dev/null +++ b/tests/layers/monthly_reports/test_helpers.py @@ -0,0 +1,76 @@ +import datetime +from unittest.mock import patch + +import pytest + +from monthly_reports.helpers import get_statement_period + + +@pytest.mark.parametrize( + "current_date, expected_period, description", + [ + ( + datetime.datetime(2024, 1, 15, 14, 30, 45, 123456), + "2023-12", + "January mid-month", + ), + (datetime.datetime(2024, 1, 1, 0, 0, 0, 0), "2023-12", "January 1st midnight"), + ( + datetime.datetime(2024, 1, 31, 23, 59, 59, 999999), + "2023-12", + "January 31st last second", + ), + ( + datetime.datetime(2024, 2, 28, 23, 59, 59, 999999), + "2024-01", + "February 28th non-leap year", + ), + ( + datetime.datetime(2024, 2, 29, 12, 0, 0, 0), + "2024-01", + "February 29th leap year", + ), + ( + datetime.datetime(2023, 2, 28, 12, 0, 0, 0), + "2023-01", + "February 28th leap year", + ), + (datetime.datetime(2024, 3, 1, 0, 0, 0, 0), "2024-02", "March 1st leap year"), + ( + datetime.datetime(2023, 3, 1, 12, 0, 0, 0), + "2023-02", + "March 1st non-leap year", + ), + (datetime.datetime(2024, 6, 1, 0, 0, 0, 0), "2024-05", "June 1st"), + (datetime.datetime(2024, 7, 31, 23, 59, 59, 999999), "2024-06", "July 31st"), + ( + datetime.datetime(2024, 8, 15, 12, 30, 45, 123456), + "2024-07", + "August mid-month", + ), + ( + datetime.datetime(2024, 12, 31, 23, 59, 59, 999999), + "2024-11", + "December 31st", + ), + (datetime.datetime(2024, 5, 1, 0, 0, 0, 0), "2024-04", "First day midnight"), + ( + datetime.datetime(2024, 5, 31, 23, 59, 59, 999999), + "2024-04", + "Last day last second", + ), + ], +) +@patch("monthly_reports.helpers.datetime") +def test_get_statement_period_parametrized( + mock_datetime, current_date, expected_period, description +): + mock_datetime.datetime.now.return_value = current_date + mock_datetime.UTC = datetime.UTC + mock_datetime.timedelta = datetime.timedelta + + result = get_statement_period() + + assert ( + result == expected_period + ), f"Failed for {description}: expected {expected_period}, got {result}" diff --git a/tests/layers/monthly_reports/test_metrics.py b/tests/layers/monthly_reports/test_metrics.py new file mode 100644 index 0000000..02bac5f --- /dev/null +++ b/tests/layers/monthly_reports/test_metrics.py @@ -0,0 +1,128 @@ +import pytest + +from monthly_reports.metrics import initialize_metrics, merge_metrics + + +class TestInitialiseMetrics: + + def test_initialise_metrics(self): + metrics = initialize_metrics() + + assert metrics["processed_count"] == 0 + assert metrics["failed_starts_count"] == 0 + assert metrics["skipped_count"] == 0 + assert metrics["already_exists_count"] == 0 + assert metrics["batches_processed"] == 0 + assert metrics["pages_processed"] == 0 + + +class TestMergeMetrics: + @pytest.mark.parametrize( + "test_case,target,source,expected,should_modify_target", + [ + ( + "basic_merge", + {"processed_count": 5, "failed_starts_count": 2, "skipped_count": 1}, + {"processed_count": 3, "failed_starts_count": 1, "skipped_count": 2}, + {"processed_count": 8, "failed_starts_count": 3, "skipped_count": 3}, + True, + ), + ( + "merge_zeros", + {"processed_count": 5, "failed_starts_count": 2}, + {"processed_count": 0, "failed_starts_count": 0}, + {"processed_count": 5, "failed_starts_count": 2}, + False, + ), + ( + "empty_target", + {"processed_count": 0, "pages_processed": 0}, + {"processed_count": 10, "pages_processed": 20}, + {"processed_count": 10, "pages_processed": 20}, + True, + ), + ( + "negative_values", + {"processed_count": 10, "failed_starts_count": 5}, + {"processed_count": -3, "failed_starts_count": -2}, + {"processed_count": 7, "failed_starts_count": 3}, + True, + ), + ( + "partial_keys", + {"processed_count": 5, "failed_starts_count": 2, "pages_processed": 0}, + {"processed_count": 3, "pages_processed": 10}, + {"processed_count": 8, "failed_starts_count": 2, "pages_processed": 10}, + True, + ), + ( + "unknown_keys", + {"processed_count": 5}, + {"processed_count": 3, "unknown_key": 10, "another_unknown": 20}, + {"processed_count": 8}, + True, + ), + ( + "large_numbers", + {"processed_count": 1000000, "batches_processed": 500}, + {"processed_count": 2000000, "batches_processed": 300}, + {"processed_count": 3000000, "batches_processed": 800}, + True, + ), + ( + "single_key", + {"processed_count": 100, "failed_starts_count": 50}, + {"processed_count": 25}, + {"processed_count": 125, "failed_starts_count": 50}, + True, + ), + ( + "full_metrics", + { + "processed_count": 100, + "failed_starts_count": 10, + "skipped_count": 5, + "already_exists_count": 15, + "batches_processed": 2, + "pages_processed": 50, + }, + { + "processed_count": 50, + "failed_starts_count": 5, + "skipped_count": 3, + "already_exists_count": 8, + "batches_processed": 1, + "pages_processed": 25, + }, + { + "processed_count": 150, + "failed_starts_count": 15, + "skipped_count": 8, + "already_exists_count": 23, + "batches_processed": 3, + "pages_processed": 75, + }, + True, + ), + ], + ) + def test_metrics_functionality( + self, test_case, target, source, expected, should_modify_target + ): + original_target = target.copy() + original_id = id(target) + + result = merge_metrics(target, source) + + assert result is None + assert target == expected + assert id(target) == original_id + + for key in source: + if key not in original_target: + assert key not in target + + if should_modify_target: + assert target != original_target + else: + assert target == original_target diff --git a/tests/layers/monthly_reports/test_processing.py b/tests/layers/monthly_reports/test_processing.py new file mode 100644 index 0000000..ecbfbad --- /dev/null +++ b/tests/layers/monthly_reports/test_processing.py @@ -0,0 +1,654 @@ +import uuid +from unittest.mock import patch, MagicMock + +from monthly_reports.processing import ( + process_account_batch, + chunk_accounts, + process_batch_continuation, + process_accounts_scan_continuation, + process_account_batches, + process_accounts_page, +) + + +class TestProcessAccountBatch: + + def test_success(self, magic_mock_sfn_client, mock_logger): + accounts_batch = [ + {"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())}, + {"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())}, + ] + + result = process_account_batch( + accounts_batch, "2024-1", magic_mock_sfn_client, mock_logger, "" + ) + + assert result["processed"] == 2 + assert result["skipped"] == 0 + + def test_invalid_account_mix(self, magic_mock_sfn_client, mock_logger): + accounts_batch = [ + {}, + {"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())}, + ] + + result = process_account_batch( + accounts_batch, "2024-1", magic_mock_sfn_client, mock_logger, "" + ) + + assert result["skipped"] == 1 + assert result["processed"] == 1 + + def test_all_accounts_invalid(self, magic_mock_sfn_client, mock_logger): + accounts_batch = [ + {}, + ] + + process_account_batch( + accounts_batch, "2024-1", magic_mock_sfn_client, mock_logger, "" + ) + + @patch("monthly_reports.processing.send_bad_account_to_dlq") + def test_invalid_account_with_dlq_parameters( + self, mock_send_dlq, magic_mock_sfn_client, mock_logger + ): + accounts_batch = [ + {"accountId": "", "userId": ""}, + ] + + result = process_account_batch( + accounts_batch, + "2024-1", + magic_mock_sfn_client, + mock_logger, + "", + sqs_endpoint="https://sqs.amazonaws.com", + dlq_url="https://sqs.amazonaws.com/queue/dlq", + aws_region="us-east-1", + ) + + assert result["skipped"] == 1 + mock_send_dlq.assert_called_once() + + def test_execution_already_exists(self, magic_mock_sfn_client, mock_logger): + accounts_batch = [ + {"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())}, + ] + + with patch( + "monthly_reports.processing.start_sfn_execution_with_retry" + ) as mock_start_sfn_execution_with_retry: + mock_start_sfn_execution_with_retry.return_value = "already_exists" + + result = process_account_batch( + accounts_batch, "2024-1", magic_mock_sfn_client, mock_logger, "" + ) + + assert result["already_exists"] == 1 + assert result["skipped"] == 0 + + def test_failed_executions(self, magic_mock_sfn_client, mock_logger): + accounts_batch = [ + {"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())}, + ] + + with patch( + "monthly_reports.processing.start_sfn_execution_with_retry" + ) as mock_start_sfn_execution_with_retry: + mock_start_sfn_execution_with_retry.return_value = "failed" + + result = process_account_batch( + accounts_batch, "2024-1", magic_mock_sfn_client, mock_logger, "" + ) + + assert result["failed_starts"] == 1 + assert result["skipped"] == 0 + + def test_exception_raised(self, magic_mock_sfn_client, mock_logger): + accounts_batch = [ + {"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())}, + ] + + with patch( + "monthly_reports.processing.start_sfn_execution_with_retry" + ) as mock_start_sfn_execution_with_retry: + mock_start_sfn_execution_with_retry.side_effect = Exception( + "Test exception" + ) + + result = process_account_batch( + accounts_batch, "2024-1", magic_mock_sfn_client, mock_logger, "" + ) + + assert result["failed_starts"] == 1 + assert result["skipped"] == 0 + + @patch("monthly_reports.processing.send_bad_account_to_dlq") + def test_failed_executions_with_dlq( + self, mock_send_dlq, magic_mock_sfn_client, mock_logger + ): + accounts_batch = [ + {"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())}, + ] + + with patch( + "monthly_reports.processing.start_sfn_execution_with_retry" + ) as mock_start_sfn_execution_with_retry: + mock_start_sfn_execution_with_retry.return_value = "failed" + + result = process_account_batch( + accounts_batch, + "2024-1", + magic_mock_sfn_client, + mock_logger, + "", + sqs_endpoint="https://sqs.amazonaws.com", + dlq_url="https://sqs.amazonaws.com/queue/dlq", + aws_region="us-east-1", + ) + + assert result["failed_starts"] == 1 + assert result["skipped"] == 0 + mock_send_dlq.assert_called_once() + + @patch("monthly_reports.processing.send_bad_account_to_dlq") + def test_exception_raised_with_dlq( + self, mock_send_dlq, magic_mock_sfn_client, mock_logger + ): + accounts_batch = [ + {"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())}, + ] + + with patch( + "monthly_reports.processing.start_sfn_execution_with_retry" + ) as mock_start_sfn_execution_with_retry: + mock_start_sfn_execution_with_retry.side_effect = Exception( + "Test exception" + ) + + result = process_account_batch( + accounts_batch, + "2024-1", + magic_mock_sfn_client, + mock_logger, + "", + sqs_endpoint="https://sqs.amazonaws.com", + dlq_url="https://sqs.amazonaws.com/queue/dlq", + aws_region="us-east-1", + ) + + assert result["failed_starts"] == 1 + assert result["skipped"] == 0 + mock_send_dlq.assert_called_once() + + +class TestChunkAccounts: + + def test_chunk_accounts_basic(self): + accounts = list(range(20)) + chunks = list(chunk_accounts(accounts, chunk_size=10)) + + assert len(chunks) == 2 + assert chunks[0] == list(range(10)) + assert chunks[1] == list(range(10, 20)) + + def test_chunk_accounts_smaller_than_chunk_size(self): + accounts = [1, 2, 3] + chunks = list(chunk_accounts(accounts, chunk_size=10)) + + assert len(chunks) == 1 + assert chunks[0] == [1, 2, 3] + + def test_chunk_accounts_empty_list(self): + accounts = [] + chunks = list(chunk_accounts(accounts, chunk_size=10)) + + assert len(chunks) == 0 + + def test_chunk_accounts_custom_chunk_size(self): + accounts = list(range(7)) + chunks = list(chunk_accounts(accounts, chunk_size=3)) + + assert len(chunks) == 3 + assert chunks[0] == [0, 1, 2] + assert chunks[1] == [3, 4, 5] + assert chunks[2] == [6] + + +class TestProcessAccountsPage: + + @patch("monthly_reports.processing.process_account_batches") + @patch("monthly_reports.processing.initialize_metrics") + def test_process_accounts_page_success( + self, mock_initialize_metrics, mock_process_batches, mock_logger + ): + mock_initialize_metrics.return_value = {"processed_count": 0} + mock_process_batches.return_value = {"processed_count": 5} + + accounts_page = [ + {"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())} + for _ in range(15) + ] + + with patch("monthly_reports.processing.merge_metrics") as mock_merge: + result = process_accounts_page( + accounts_page=accounts_page, + statement_period="2024-1", + context=MagicMock(), + logger=mock_logger, + sfn_client=MagicMock(), + state_machine_arn="arn:aws:states:us-east-1:123456789012:stateMachine:test", + scan_params={}, + last_evaluated_key=None, + sqs_endpoint="https://sqs.us-east-1.amazonaws.com", + continuation_queue_url="https://sqs.us-east-1.amazonaws.com/123456789012/test-queue", + aws_region="us-east-1", + batch_size=10, + ) + + mock_process_batches.assert_called_once() + mock_merge.assert_called_once() + assert result == {"processed_count": 0} + + +class TestProcessAccountBatches: + + def test_process_account_batches_success(self, mock_logger): + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 60000 # 60 seconds + + account_batches = [ + [{"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())}], + [{"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())}], + ] + + with patch( + "monthly_reports.processing.process_account_batch" + ) as mock_process_batch: + mock_process_batch.return_value = {"processed": 1, "skipped": 0} + + with patch( + "monthly_reports.processing.initialize_metrics" + ) as mock_init_metrics: + mock_init_metrics.return_value = { + "processed_count": 0, + "skipped_count": 0, + "batches_processed": 0, + "failed_starts_count": 0, + } + + result = process_account_batches( + account_batches=account_batches, + statement_period="2024-1", + context=mock_context, + logger=mock_logger, + sfn_client=MagicMock(), + state_machine_arn="arn:aws:states:us-east-1:123456789012:stateMachine:test", + scan_params={}, + last_evaluated_key=None, + sqs_endpoint="https://sqs.us-east-1.amazonaws.com", + continuation_queue_url="https://sqs.us-east-1.amazonaws.com/123456789012/test-queue", + aws_region="us-east-1", + ) + + assert result["processed_count"] == 2 + assert result["batches_processed"] == 2 + assert mock_process_batch.call_count == 2 + + def test_process_account_batches_timeout_approaching(self, mock_logger): + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 20000 # 20 seconds + + account_batches = [ + [{"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())}], + [{"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())}], + ] + + with patch( + "monthly_reports.processing.send_continuation_message" + ) as mock_send_continuation: + with patch( + "monthly_reports.processing.initialize_metrics" + ) as mock_init_metrics: + mock_init_metrics.return_value = { + "processed_count": 0, + "skipped_count": 0, + "batches_processed": 0, + "failed_starts_count": 0, + } + + result = process_account_batches( + account_batches=account_batches, + statement_period="2024-1", + context=mock_context, + logger=mock_logger, + sfn_client=MagicMock(), + state_machine_arn="arn:aws:states:us-east-1:123456789012:stateMachine:test", + scan_params={}, + last_evaluated_key={"id": "test"}, + sqs_endpoint="https://sqs.us-east-1.amazonaws.com", + continuation_queue_url="https://sqs.us-east-1.amazonaws.com/123456789012/test-queue", + aws_region="us-east-1", + safety_buffer=30, + ) + + mock_send_continuation.assert_called_once() + assert result["batches_processed"] == 0 + + def test_process_account_batches_exception_handling(self, mock_logger): + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 60000 + + account_batches = [ + [{"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())}], + ] + + with patch( + "monthly_reports.processing.process_account_batch" + ) as mock_process_batch: + mock_process_batch.side_effect = Exception("Test exception") + + with patch( + "monthly_reports.processing.initialize_metrics" + ) as mock_init_metrics: + mock_init_metrics.return_value = { + "processed_count": 0, + "skipped_count": 0, + "batches_processed": 0, + "failed_starts_count": 0, + } + + with patch( + "monthly_reports.processing.send_bad_account_to_dlq" + ) as mock_send_to_dlq: + result = process_account_batches( + account_batches=account_batches, + statement_period="2024-1", + context=mock_context, + logger=mock_logger, + sfn_client=MagicMock(), + state_machine_arn="arn:aws:states:us-east-1:123456789012:stateMachine:test", + scan_params={}, + last_evaluated_key=None, + sqs_endpoint="https://sqs.us-east-1.amazonaws.com", + continuation_queue_url="https://sqs.us-east-1.amazonaws.com/123456789012/test-queue", + aws_region="us-east-1", + dlq_url="https://sqs.us-east-1.amazonaws.com/123456789012/dlq-queue", + ) + + assert result["failed_starts_count"] == 1 + assert result["batches_processed"] == 0 + + mock_send_to_dlq.assert_called_once() + call_args = mock_send_to_dlq.call_args[0] + assert call_args[0] == account_batches[0][0] + assert call_args[1] == "2024-1" + assert "Batch processing exception: Test exception" in call_args[2] + + +class TestProcessAccountsScanContinuation: + + @patch("monthly_reports.processing.merge_metrics") + @patch("monthly_reports.processing.initialize_metrics") + @patch("monthly_reports.processing.process_accounts_page") + @patch("monthly_reports.processing.get_paginated_table_data") + def test_process_accounts_scan_continuation_success( + self, + mock_get_paginated_data, + mock_process_page, + mock_initialize_metrics, + mock_merge_metrics, + mock_logger, + ): + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 60000 + + mock_initialize_metrics.return_value = {"pages_processed": 0} + mock_get_paginated_data.side_effect = [ + ([{"accountId": "123", "userId": "456"}], {"id": "next"}), + ([], None), + ] + mock_process_page.return_value = {"processed_count": 1} + + result = process_accounts_scan_continuation( + scan_params={}, + statement_period="2024-1", + context=mock_context, + logger=mock_logger, + accounts_table=MagicMock(), + sfn_client=MagicMock(), + state_machine_arn="arn:aws:states:us-east-1:123456789012:stateMachine:test", + sqs_endpoint="https://sqs.us-east-1.amazonaws.com", + continuation_queue_url="https://sqs.us-east-1.amazonaws.com/123456789012/test-queue", + aws_region="us-east-1", + ) + + assert mock_get_paginated_data.call_count == 2 + assert mock_process_page.call_count == 1 + assert result["pages_processed"] == 2 + mock_merge_metrics.assert_called_once() + + @patch("monthly_reports.processing.merge_metrics") + @patch("monthly_reports.processing.initialize_metrics") + @patch("monthly_reports.processing.process_accounts_page") + @patch("monthly_reports.processing.get_paginated_table_data") + def test_process_accounts_scan_continuation_no_last_evaluated_key( + self, + mock_get_paginated_data, + mock_process_page, + mock_initialize_metrics, + mock_merge_metrics, + mock_logger, + ): + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 60000 + + mock_initialize_metrics.return_value = {"pages_processed": 0} + mock_get_paginated_data.side_effect = [ + ([{"accountId": "123", "userId": "456"}], None) + ] + mock_process_page.return_value = {"processed_count": 1} + + result = process_accounts_scan_continuation( + scan_params={}, + statement_period="2024-1", + context=mock_context, + logger=mock_logger, + accounts_table=MagicMock(), + sfn_client=MagicMock(), + state_machine_arn="arn:aws:states:us-east-1:123456789012:stateMachine:test", + sqs_endpoint="https://sqs.us-east-1.amazonaws.com", + continuation_queue_url="https://sqs.us-east-1.amazonaws.com/123456789012/test-queue", + aws_region="us-east-1", + ) + + assert mock_get_paginated_data.call_count == 1 + assert mock_process_page.call_count == 1 + assert result["pages_processed"] == 1 + + mock_merge_metrics.assert_called_once() + + @patch("monthly_reports.processing.send_continuation_message") + @patch("monthly_reports.processing.initialize_metrics") + def test_process_accounts_scan_continuation_timeout( + self, mock_initialize_metrics, mock_send_continuation, mock_logger + ): + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 20000 # 20 seconds + + mock_initialize_metrics.return_value = {"pages_processed": 0} + + result = process_accounts_scan_continuation( + scan_params={}, + statement_period="2024-1", + context=mock_context, + logger=mock_logger, + accounts_table=MagicMock(), + sfn_client=MagicMock(), + state_machine_arn="arn:aws:states:us-east-1:123456789012:stateMachine:test", + sqs_endpoint="https://sqs.us-east-1.amazonaws.com", + continuation_queue_url="https://sqs.us-east-1.amazonaws.com/123456789012/test-queue", + aws_region="us-east-1", + safety_buffer=30, + ) + + mock_send_continuation.assert_called_once() + assert result["pages_processed"] == 0 + + @patch("monthly_reports.processing.get_paginated_table_data") + @patch("monthly_reports.processing.initialize_metrics") + def test_process_accounts_scan_continuation_no_accounts( + self, mock_initialize_metrics, mock_get_paginated_data, mock_logger + ): + mock_context = MagicMock() + mock_context.get_remaining_time_in_millis.return_value = 60000 + + mock_initialize_metrics.return_value = {"pages_processed": 0} + mock_get_paginated_data.return_value = ([], None) + + result = process_accounts_scan_continuation( + scan_params={}, + statement_period="2024-1", + context=mock_context, + logger=mock_logger, + accounts_table=MagicMock(), + sfn_client=MagicMock(), + state_machine_arn="arn:aws:states:us-east-1:123456789012:stateMachine:test", + sqs_endpoint="https://sqs.us-east-1.amazonaws.com", + continuation_queue_url="https://sqs.us-east-1.amazonaws.com/123456789012/test-queue", + aws_region="us-east-1", + ) + + assert result["pages_processed"] == 1 + + +class TestProcessBatchContinuation: + + @patch("monthly_reports.processing.process_accounts_scan_continuation") + @patch("monthly_reports.processing.process_account_batches") + @patch("monthly_reports.processing.initialize_metrics") + @patch("monthly_reports.processing.merge_metrics") + def test_process_batch_continuation_with_remaining_accounts_and_key( + self, + mock_merge_metrics, + mock_initialize_metrics, + mock_process_batches, + mock_process_scan, + mock_logger, + ): + mock_initialize_metrics.return_value = {"processed_count": 0} + mock_process_batches.return_value = {"processed_count": 2} + mock_process_scan.return_value = {"processed_count": 5} + + remaining_accounts = [ + {"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())}, + {"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())}, + ] + + process_batch_continuation( + scan_params={}, + statement_period="2024-1", + remaining_accounts=remaining_accounts, + last_evaluated_key={"id": "test"}, + context=MagicMock(), + logger=mock_logger, + accounts_table=MagicMock(), + sfn_client=MagicMock(), + state_machine_arn="arn:aws:states:us-east-1:123456789012:stateMachine:test", + sqs_endpoint="https://sqs.us-east-1.amazonaws.com", + continuation_queue_url="https://sqs.us-east-1.amazonaws.com/123456789012/test-queue", + aws_region="us-east-1", + ) + + mock_process_batches.assert_called_once() + mock_process_scan.assert_called_once() + assert mock_merge_metrics.call_count == 2 + + @patch("monthly_reports.processing.process_accounts_scan_continuation") + @patch("monthly_reports.processing.initialize_metrics") + @patch("monthly_reports.processing.merge_metrics") + def test_process_batch_continuation_no_remaining_accounts_with_key( + self, + mock_merge_metrics, + mock_initialize_metrics, + mock_process_scan, + mock_logger, + ): + mock_initialize_metrics.return_value = {"processed_count": 0} + mock_process_scan.return_value = {"processed_count": 5} + + process_batch_continuation( + scan_params={}, + statement_period="2024-1", + remaining_accounts=[], + last_evaluated_key={"id": "test"}, + context=MagicMock(), + logger=mock_logger, + accounts_table=MagicMock(), + sfn_client=MagicMock(), + state_machine_arn="arn:aws:states:us-east-1:123456789012:stateMachine:test", + sqs_endpoint="https://sqs.us-east-1.amazonaws.com", + continuation_queue_url="https://sqs.us-east-1.amazonaws.com/123456789012/test-queue", + aws_region="us-east-1", + ) + + mock_process_scan.assert_called_once() + assert mock_merge_metrics.call_count == 1 + + @patch("monthly_reports.processing.process_account_batches") + @patch("monthly_reports.processing.initialize_metrics") + @patch("monthly_reports.processing.merge_metrics") + def test_process_batch_continuation_remaining_accounts_no_key( + self, + mock_merge_metrics, + mock_initialize_metrics, + mock_process_batches, + mock_logger, + ): + mock_initialize_metrics.return_value = {"processed_count": 0} + mock_process_batches.return_value = {"processed_count": 2} + + remaining_accounts = [ + {"accountId": str(uuid.uuid4()), "userId": str(uuid.uuid4())}, + ] + + process_batch_continuation( + scan_params={}, + statement_period="2024-1", + remaining_accounts=remaining_accounts, + last_evaluated_key=None, + context=MagicMock(), + logger=mock_logger, + accounts_table=MagicMock(), + sfn_client=MagicMock(), + state_machine_arn="arn:aws:states:us-east-1:123456789012:stateMachine:test", + sqs_endpoint="https://sqs.us-east-1.amazonaws.com", + continuation_queue_url="https://sqs.us-east-1.amazonaws.com/123456789012/test-queue", + aws_region="us-east-1", + ) + + mock_process_batches.assert_called_once() + assert mock_merge_metrics.call_count == 1 + + @patch("monthly_reports.processing.initialize_metrics") + def test_process_batch_continuation_no_remaining_accounts_no_key( + self, mock_initialize_metrics, mock_logger + ): + mock_initialize_metrics.return_value = {"processed_count": 0} + + result = process_batch_continuation( + scan_params={}, + statement_period="2024-1", + remaining_accounts=[], + last_evaluated_key=None, + context=MagicMock(), + logger=mock_logger, + accounts_table=MagicMock(), + sfn_client=MagicMock(), + state_machine_arn="arn:aws:states:us-east-1:123456789012:stateMachine:test", + sqs_endpoint="https://sqs.us-east-1.amazonaws.com", + continuation_queue_url="https://sqs.us-east-1.amazonaws.com/123456789012/test-queue", + aws_region="us-east-1", + ) + + assert result == {"processed_count": 0} diff --git a/tests/layers/monthly_reports/test_responses.py b/tests/layers/monthly_reports/test_responses.py new file mode 100644 index 0000000..2807876 --- /dev/null +++ b/tests/layers/monthly_reports/test_responses.py @@ -0,0 +1,28 @@ +import pytest + +from monthly_reports.responses import create_response + + +class TestResponses: + + @pytest.mark.parametrize( + "status,expected_code", + [ + ("COMPLETED", 200), + ("TIMEOUT_CONTINUATION", 202), + ("ERROR_NO_CONTINUATION_QUEUE", 500), + ("CRITICAL_ERROR", 500), + ("UNKNOWN_STATUS", 500), + ], + ) + def test_create_response_status_codes(self, mock_logger, status, expected_code): + metrics = { + "processed_count": 1, + "failed_starts_count": 0, + "skipped_count": 0, + "already_exists_count": 0, + } + result = create_response(metrics, status, mock_logger) + + assert result + assert result.get("statusCode") == expected_code diff --git a/tests/layers/monthly_reports/test_sfn.py b/tests/layers/monthly_reports/test_sfn.py new file mode 100644 index 0000000..9c53804 --- /dev/null +++ b/tests/layers/monthly_reports/test_sfn.py @@ -0,0 +1,145 @@ +import uuid +from unittest.mock import patch + +import pytest +from botocore.exceptions import ClientError + +from monthly_reports.sfn import start_sfn_execution_with_retry + + +class TestStartExecution: + + def test_success(self, mock_logger, magic_mock_sfn_client): + + input_id = str(uuid.uuid4()) + + result = start_sfn_execution_with_retry( + magic_mock_sfn_client, + "test-state-machine-arn", + "test-input", + {"id": input_id}, + mock_logger, + ) + + assert result == "processed" + assert magic_mock_sfn_client.start_execution.call_count == 1 + + def test_execution_already_exists(self, mock_logger, magic_mock_sfn_client): + + magic_mock_sfn_client.start_execution.side_effect = ClientError( + {"Error": {"Code": "ExecutionAlreadyExistsException"}}, "StartExecution" + ) + + input_id = str(uuid.uuid4()) + + result = start_sfn_execution_with_retry( + magic_mock_sfn_client, + "test-state-machine-arn", + "test-input", + {"id": input_id}, + mock_logger, + ) + + assert result == "already_exists" + + @pytest.mark.parametrize( + "error_code", ["ThrottlingException", "ServiceUnavailable", "InternalFailure"] + ) + def test_retryable_error_success_on_retry( + self, mock_logger, magic_mock_sfn_client, error_code + ): + with patch("time.sleep"): + + magic_mock_sfn_client.start_execution.side_effect = [ + ClientError( + error_response={"Error": {"Code": error_code}}, + operation_name="StartExecution", + ), + None, + ] + + input_id = str(uuid.uuid4()) + + result = start_sfn_execution_with_retry( + magic_mock_sfn_client, + "test-state-machine-arn", + "test-execution", + {"id": input_id}, + mock_logger, + ) + + assert result == "processed" + assert magic_mock_sfn_client.start_execution.call_count == 2 + + @pytest.mark.parametrize( + "error_code", ["ThrottlingException", "ServiceUnavailable", "InternalFailure"] + ) + def test_retryable_error_max_retries_exceeded( + self, mock_logger, magic_mock_sfn_client, error_code, mocker + ): + with patch("time.sleep") as mock_sleep: + + client_error = ClientError( + error_response={"Error": {"Code": error_code}}, + operation_name="StartExecution", + ) + magic_mock_sfn_client.start_execution.side_effect = client_error + + input_id = str(uuid.uuid4()) + + with pytest.raises(ClientError): + start_sfn_execution_with_retry( + magic_mock_sfn_client, + "test-state-machine-arn", + "test-execution", + {"id": input_id}, + mock_logger, + max_retries=3, + ) + + assert magic_mock_sfn_client.start_execution.call_count == 3 + assert mock_sleep.call_count == 2 + + def test_non_retryable_error(self, mock_logger, magic_mock_sfn_client): + client_error = ClientError( + error_response={"Error": {"Code": "InvalidParameterValue"}}, + operation_name="StartExecution", + ) + magic_mock_sfn_client.start_execution.side_effect = client_error + + input_id = str(uuid.uuid4()) + + with pytest.raises(ClientError): + start_sfn_execution_with_retry( + magic_mock_sfn_client, + "test-state-machine-arn", + "test-execution", + {"id": input_id}, + mock_logger, + ) + + assert magic_mock_sfn_client.start_execution.call_count == 1 + mock_logger.error.assert_called_once() + + def test_custom_max_retries(self, mock_logger, magic_mock_sfn_client, mocker): + with patch("time.sleep") as mock_sleep: + client_error = ClientError( + error_response={"Error": {"Code": "ThrottlingException"}}, + operation_name="StartExecution", + ) + magic_mock_sfn_client.start_execution.side_effect = client_error + + input_id = str(uuid.uuid4()) + + with pytest.raises(ClientError): + start_sfn_execution_with_retry( + magic_mock_sfn_client, + "test-state-machine-arn", + "test-execution", + {"id": input_id}, + mock_logger, + max_retries=2, + ) + + assert magic_mock_sfn_client.start_execution.call_count == 2 + assert mock_sleep.call_count == 1 diff --git a/tests/layers/monthly_reports/test_sqs.py b/tests/layers/monthly_reports/test_sqs.py new file mode 100644 index 0000000..4e00dbf --- /dev/null +++ b/tests/layers/monthly_reports/test_sqs.py @@ -0,0 +1,92 @@ +from unittest.mock import patch + +from monthly_reports.sqs import send_continuation_message, send_bad_account_to_dlq + + +class TestSqsHelpers: + + def test_no_continuation_queue_url(self, mock_logger): + + result = send_continuation_message({}, "", [], {}, "", "", "", "", mock_logger) + + assert result is None + assert mock_logger.error.call_count == 1 + assert ( + mock_logger.error.call_args[0][0] + == "Cannot send continuation message: CONTINUATION_QUEUE_URL not set" + ) + + @patch("monthly_reports.sqs.send_message_to_sqs") + def test_send_message_success(self, mock_send_sqs, mock_logger): + """Test successful message sending with all data types""" + scan_params = {"TableName": "accounts"} + accounts = [{"accountId": "acc1", "userId": "user1"}] + last_key = {"accountId": "acc123"} + + send_continuation_message( + scan_params=scan_params, + statement_period="2024-01", + remaining_accounts=accounts, + last_evaluated_key=last_key, + continuation_type="batch_continuation", + sqs_endpoint="https://sqs.us-east-1.amazonaws.com", + continuation_queue_url="https://queue-url", + aws_region="us-east-1", + logger=mock_logger, + ) + + mock_send_sqs.assert_called_once() + + call_args = mock_send_sqs.call_args + actual_message = call_args[1]["message"] + expected_message = { + "scan_params": scan_params, + "statement_period": "2024-01", + "remaining_accounts": accounts, + "last_evaluated_key": last_key, + } + assert actual_message == expected_message + + expected_attributes = { + "continuation_type": { + "DataType": "String", + "StringValue": "batch_continuation", + } + } + assert call_args[1]["message_attributes"] == expected_attributes + + def test_send_bad_account_to_dlq_no_dlq_url(self, mock_logger): + """Test warning when DLQ URL is not set""" + result = send_bad_account_to_dlq( + account={"accountId": "acc1"}, + statement_period="2024-01", + error_reason="Test error", + sqs_endpoint="https://sqs.us-east-1.amazonaws.com", + dlq_url="", # Empty DLQ URL + aws_region="us-east-1", + logger=mock_logger, + ) + + assert result is None + mock_logger.warning.assert_called_once_with( + "Cannot send bad account to DLQ: DLQ_URL not set" + ) + + @patch("monthly_reports.sqs.send_message_to_sqs") + def test_send_bad_account_to_dlq_exception(self, mock_send_sqs, mock_logger): + """Test exception handling when sending to DLQ fails""" + mock_send_sqs.side_effect = Exception("SQS send failed") + + send_bad_account_to_dlq( + account={"accountId": "acc1"}, + statement_period="2024-01", + error_reason="Test error", + sqs_endpoint="https://sqs.us-east-1.amazonaws.com", + dlq_url="https://queue-url", + aws_region="us-east-1", + logger=mock_logger, + ) + + mock_logger.error.assert_called_once_with( + "Failed to send bad account to DLQ: SQS send failed" + ) diff --git a/tests/requirements.txt b/tests/requirements.txt index 2591b1f..8fae10e 100644 --- a/tests/requirements.txt +++ b/tests/requirements.txt @@ -1,4 +1,4 @@ -aws_lambda_powertools==3.12.0 +aws_lambda_powertools==3.17.0 boto3==1.38.13 pytest==8.3.5 moto==5.1.6