proposal-system/lambdas/pdf-generate/app.py
Adam Moussa 3d050bcf8e
Some checks are pending
Deploy / Deploy to AWS (push) Waiting to run
fix(lambdas): SigV4-sign internal API calls and bundle Lambda dependencies (#122)
The .NET API Lambda Function URL uses authType=AWS_IAM, but the four workload
Lambdas (suggestions, pdf-extract, pdf-generate, library-ingest) sent unsigned
requests with only X-Internal-Api-Key -> every internal call 403s. They also
used bare fromAsset() with no pip bundling -> ImportError at cold start. Both
made the SQS->Lambda->API pipeline non-functional when deployed (v1 pre-flight).

- Add _sign_request_headers (botocore SigV4Auth, service "lambda"); serialize the
  JSON body once and send via httpx content= so the signed payload hash matches
  the bytes sent; preserve X-Internal-Api-Key for the app-layer check. Sign per
  retry attempt to avoid SigV4 timestamp expiry on slow retries.
- Add CDK pip bundling (--platform manylinux2014_aarch64 --only-binary=:all:) to
  all four Lambdas so ARM64 wheels (reportlab, Pillow, pdfplumber) ship.
- Converge _retry_request across all four (fixes possibly-undefined return in
  pdf-extract/pdf-generate).
- Add SigV4 signing regression tests.

Verified: ruff clean, infra tsc clean, aarch64 wheels resolve for all four,
23 pytest pass. GPT-4.1 cross-family review: no BLOCK (FIX + NIT applied).
2026-06-12 17:13:08 -04:00

675 lines
20 KiB
Python

"""Proposal System - PDF Generate Lambda.
Generates professional branded proposal PDFs using reportlab Platypus.
Triggered via SQS when an admin requests PDF generation.
"""
import json
import logging
import os
import re
import time
from datetime import datetime
from io import BytesIO
import boto3
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
import httpx
from reportlab.lib import colors
from reportlab.lib.enums import TA_CENTER, TA_RIGHT
from reportlab.lib.pagesizes import letter
from reportlab.lib.styles import ParagraphStyle, getSampleStyleSheet
from reportlab.lib.units import inch
from reportlab.platypus import (
Paragraph,
SimpleDocTemplate,
Spacer,
Table,
TableStyle,
)
logger = logging.getLogger(__name__)
logger.setLevel(os.environ.get("LOG_LEVEL", "INFO"))
GENERATED_BUCKET = os.environ.get("GENERATED_BUCKET", "")
API_BASE_URL = os.environ.get("API_BASE_URL", "")
INTERNAL_API_KEY_SECRET_ARN = os.environ.get("INTERNAL_API_KEY_SECRET_ARN", "")
s3 = boto3.client("s3")
secrets_client = boto3.client("secretsmanager")
_cached_api_key: str | None = None
_cached_api_key_ts: float = 0.0
_API_KEY_TTL_SECONDS = 300 # Fix: LAM-M9 — re-fetch every 5 minutes
# Fix: LAM-M8 — pattern for allowed S3 key characters
_SAFE_S3_KEY_RE = re.compile(r"^[a-zA-Z0-9\-_./\s]+$")
COMPANY_NAME = "Sea Haven Industries"
COMPANY_ADDRESS = "Sea Haven Industries LLC"
COMPANY_PHONE = ""
COMPANY_EMAIL = "info@seahavenind.com"
TERMS_AND_CONDITIONS = """
1. This proposal is valid for 30 days from the date of issue.
2. Payment terms: Net 30 days from invoice date.
3. Any changes to the scope of work may result in additional charges.
4. Work will be scheduled upon acceptance of this proposal.
5. All materials and workmanship are guaranteed for one (1) year from completion.
6. Client is responsible for providing access to the work area.
7. This proposal does not include permits unless specifically noted in the line items.
""".strip()
def _get_api_key() -> str:
"""Fetch internal API key from Secrets Manager with a 5-minute TTL cache.
Fix: LAM-M9 — the key was previously cached indefinitely. Now the cached
value expires after _API_KEY_TTL_SECONDS so rotated secrets take effect.
"""
global _cached_api_key, _cached_api_key_ts
now = time.monotonic()
if (
_cached_api_key is not None
and (now - _cached_api_key_ts) < _API_KEY_TTL_SECONDS
):
return _cached_api_key
if INTERNAL_API_KEY_SECRET_ARN:
resp = secrets_client.get_secret_value(SecretId=INTERNAL_API_KEY_SECRET_ARN)
_cached_api_key = resp["SecretString"]
else:
_cached_api_key = ""
_cached_api_key_ts = now
return _cached_api_key
def _validate_s3_key(key: str) -> str:
"""Fix: LAM-M8 — sanitize and validate S3 keys before use."""
sanitized = key.replace("../", "").replace("..\\", "")
while "//" in sanitized:
sanitized = sanitized.replace("//", "/")
sanitized = sanitized.strip("/").strip()
if not sanitized:
raise ValueError("S3 key is empty after sanitization")
if not _SAFE_S3_KEY_RE.match(sanitized):
raise ValueError(f"S3 key contains disallowed characters: {sanitized!r}")
return sanitized
def handler(event, context):
batch_item_failures = []
# Fix: LAM-M1 — validate event structure before processing
records = event.get("Records")
if not isinstance(records, list) or not records:
logger.warning("Event has no Records or Records is not a list, returning early")
return {"batchItemFailures": []}
for record in records:
if "body" not in record:
logger.warning(
"Record missing 'body' key, skipping: %s",
record.get("messageId", "unknown"),
)
continue
try:
body = json.loads(record["body"])
except (json.JSONDecodeError, TypeError) as e:
# Fix: LAM-M1 — malformed JSON cannot be retried, add to failures
logger.error("Malformed JSON in record %s: %s", record.get("messageId"), e)
batch_item_failures.append({"itemIdentifier": record["messageId"]})
continue
try:
payload = body.get("payload", body)
proposal_id = payload["proposalId"]
generate_pdf(proposal_id)
except Exception as e:
# Fix: LAM-M5 — include stack trace in error logging
logger.exception(
"Failed to process record %s: %s", record.get("messageId"), e
)
batch_item_failures.append({"itemIdentifier": record["messageId"]})
return {"batchItemFailures": batch_item_failures}
def generate_pdf(proposal_id: str):
proposal = fetch_proposal(proposal_id)
if not proposal:
logger.warning("Proposal %s not found", proposal_id)
return
line_items = fetch_line_items(proposal_id)
pdf_bytes = build_pdf(proposal, line_items)
proposal_number = proposal["proposalNumber"]
revision = proposal.get("currentRevision", 1)
s3_key = f"{proposal_number}/rev-{revision}.pdf"
# Fix: LAM-M8 — validate constructed S3 key
s3_key = _validate_s3_key(s3_key)
upload_pdf(s3_key, pdf_bytes)
register_pdf(proposal_id, s3_key)
logger.info("Generated PDF: %s (%d bytes)", s3_key, len(pdf_bytes))
def fetch_proposal(proposal_id: str) -> dict | None:
try:
resp = _retry_request(
"GET",
f"{API_BASE_URL}/api/proposals/{proposal_id}",
headers=_api_headers(),
)
if resp.status_code == 200:
return resp.json()
except Exception as e:
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error fetching proposal: %s", e)
return None
def fetch_line_items(proposal_id: str) -> list[dict]:
try:
resp = _retry_request(
"GET",
f"{API_BASE_URL}/api/proposals/{proposal_id}/line-items",
headers=_api_headers(),
)
if resp.status_code == 200:
return resp.json()
except Exception as e:
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error fetching line items: %s", e)
return []
def build_pdf(proposal: dict, line_items: list[dict]) -> bytes:
buffer = BytesIO()
doc = SimpleDocTemplate(
buffer,
pagesize=letter,
leftMargin=0.75 * inch,
rightMargin=0.75 * inch,
topMargin=0.75 * inch,
bottomMargin=0.75 * inch,
)
styles = _get_styles()
elements = []
# Header
elements.extend(_build_header(proposal, styles))
elements.append(Spacer(1, 0.3 * inch))
# Proposal metadata
elements.extend(_build_metadata(proposal, styles))
elements.append(Spacer(1, 0.3 * inch))
# Scope of work
elements.extend(_build_scope(proposal, styles))
elements.append(Spacer(1, 0.3 * inch))
# Line items table
elements.extend(_build_line_items_table(line_items, styles))
elements.append(Spacer(1, 0.4 * inch))
# Terms and conditions
elements.extend(_build_terms(styles))
doc.build(elements, onFirstPage=_page_footer, onLaterPages=_page_footer)
return buffer.getvalue()
def _get_styles():
styles = getSampleStyleSheet()
styles.add(
ParagraphStyle(
"CompanyName",
parent=styles["Heading1"],
fontSize=18,
leading=22,
textColor=colors.HexColor("#1a237e"),
spaceAfter=2,
)
)
styles.add(
ParagraphStyle(
"CompanyInfo",
parent=styles["Normal"],
fontSize=9,
leading=12,
textColor=colors.HexColor("#555555"),
)
)
styles.add(
ParagraphStyle(
"ProposalTitle",
parent=styles["Heading2"],
fontSize=14,
leading=18,
textColor=colors.HexColor("#1a237e"),
spaceBefore=6,
spaceAfter=12,
)
)
styles.add(
ParagraphStyle(
"SectionHeader",
parent=styles["Heading3"],
fontSize=11,
leading=14,
textColor=colors.HexColor("#1a237e"),
spaceBefore=8,
spaceAfter=6,
borderWidth=0,
)
)
styles.add(
ParagraphStyle(
"MetaLabel",
parent=styles["Normal"],
fontSize=9,
leading=12,
textColor=colors.HexColor("#666666"),
)
)
styles.add(
ParagraphStyle(
"MetaValue",
parent=styles["Normal"],
fontSize=10,
leading=13,
fontName="Helvetica-Bold",
)
)
styles.add(
ParagraphStyle(
"ScopeText",
parent=styles["Normal"],
fontSize=10,
leading=14,
spaceBefore=4,
)
)
styles.add(
ParagraphStyle(
"TermsText",
parent=styles["Normal"],
fontSize=8,
leading=11,
textColor=colors.HexColor("#555555"),
)
)
styles.add(
ParagraphStyle(
"TotalLabel",
parent=styles["Normal"],
fontSize=11,
leading=14,
fontName="Helvetica-Bold",
alignment=TA_RIGHT,
)
)
styles.add(
ParagraphStyle(
"FooterText",
parent=styles["Normal"],
fontSize=8,
leading=10,
textColor=colors.HexColor("#888888"),
alignment=TA_CENTER,
)
)
return styles
def _build_header(proposal: dict, styles) -> list:
revision = proposal.get("currentRevision", 1)
revision_text = f" | Rev {revision}" if revision > 1 else ""
header_data = [
[
Paragraph(COMPANY_NAME, styles["CompanyName"]),
Paragraph(f"PROPOSAL{revision_text}", styles["ProposalTitle"]),
],
[
Paragraph(f"{COMPANY_EMAIL}", styles["CompanyInfo"]),
Paragraph(f"#{proposal['proposalNumber']}", styles["MetaValue"]),
],
]
header_table = Table(header_data, colWidths=[3.5 * inch, 3.5 * inch])
header_table.setStyle(
TableStyle(
[
("VALIGN", (0, 0), (-1, -1), "TOP"),
("ALIGN", (1, 0), (1, -1), "RIGHT"),
("LINEBELOW", (0, -1), (-1, -1), 1.5, colors.HexColor("#1a237e")),
("BOTTOMPADDING", (0, -1), (-1, -1), 8),
]
)
)
return [header_table]
def _build_metadata(proposal: dict, styles) -> list:
submitted_at = proposal.get("submittedAt", "")
if submitted_at:
try:
dt = datetime.fromisoformat(submitted_at.replace("Z", "+00:00"))
submitted_at = dt.strftime("%B %d, %Y")
except (ValueError, TypeError):
pass
approved_at = proposal.get("approvedAt", "")
if approved_at:
try:
dt = datetime.fromisoformat(approved_at.replace("Z", "+00:00"))
approved_at = dt.strftime("%B %d, %Y")
except (ValueError, TypeError):
pass
meta_data = [
[
Paragraph("Customer", styles["MetaLabel"]),
Paragraph("Site Address", styles["MetaLabel"]),
],
[
Paragraph(proposal.get("customerName", ""), styles["MetaValue"]),
Paragraph(proposal.get("customerAddress", ""), styles["MetaValue"]),
],
[
Paragraph("Work Order #", styles["MetaLabel"]),
Paragraph("Date", styles["MetaLabel"]),
],
[
Paragraph(proposal.get("workOrderNumber", ""), styles["MetaValue"]),
Paragraph(approved_at or submitted_at, styles["MetaValue"]),
],
[
Paragraph("Category", styles["MetaLabel"]),
Paragraph("Priority", styles["MetaLabel"]),
],
[
Paragraph(proposal.get("serviceCategory", ""), styles["MetaValue"]),
Paragraph(proposal.get("priority", ""), styles["MetaValue"]),
],
]
meta_table = Table(meta_data, colWidths=[3.5 * inch, 3.5 * inch])
meta_table.setStyle(
TableStyle(
[
("VALIGN", (0, 0), (-1, -1), "TOP"),
("TOPPADDING", (0, 0), (-1, -1), 2),
("BOTTOMPADDING", (0, 0), (-1, -1), 2),
]
)
)
return [meta_table]
def _build_scope(proposal: dict, styles) -> list:
scope = proposal.get("refinedScope") or proposal.get("scopeOfWork", "")
if not scope:
return []
return [
Paragraph("Scope of Work", styles["SectionHeader"]),
Paragraph(scope, styles["ScopeText"]),
]
def _build_line_items_table(line_items: list[dict], styles) -> list:
if not line_items:
return [Paragraph("No line items", styles["Normal"])]
elements = [Paragraph("Itemized Pricing", styles["SectionHeader"])]
header = ["#", "Description", "Qty", "Unit", "Unit Price", "Total"]
table_data = [header]
subtotal = 0.0
for i, li in enumerate(line_items, 1):
qty = li.get("quantity", "")
unit = li.get("unit", "")
unit_price = li.get("unitPrice")
total_price = li.get("totalPrice", 0)
pricing_mode = li.get("pricingMode", "TotalPrice")
subtotal += float(total_price or 0)
if pricing_mode == "TotalPrice":
up_str = "-"
elif unit_price is not None:
up_str = f"${float(unit_price):,.2f}"
else:
up_str = "-"
tp_str = f"${float(total_price):,.2f}" if total_price else "-"
row = [
str(i),
Paragraph(li.get("description", ""), styles["Normal"]),
str(qty) if qty else "",
unit,
up_str,
tp_str,
]
table_data.append(row)
col_widths = [
0.35 * inch,
3.15 * inch,
0.55 * inch,
0.7 * inch,
1.0 * inch,
1.0 * inch,
]
table = Table(table_data, colWidths=col_widths, repeatRows=1)
table.setStyle(
TableStyle(
[
# Header row
("BACKGROUND", (0, 0), (-1, 0), colors.HexColor("#1a237e")),
("TEXTCOLOR", (0, 0), (-1, 0), colors.white),
("FONTNAME", (0, 0), (-1, 0), "Helvetica-Bold"),
("FONTSIZE", (0, 0), (-1, 0), 9),
("BOTTOMPADDING", (0, 0), (-1, 0), 6),
("TOPPADDING", (0, 0), (-1, 0), 6),
# Data rows
("FONTSIZE", (0, 1), (-1, -1), 9),
("TOPPADDING", (0, 1), (-1, -1), 4),
("BOTTOMPADDING", (0, 1), (-1, -1), 4),
("VALIGN", (0, 0), (-1, -1), "MIDDLE"),
# Alignment
("ALIGN", (0, 0), (0, -1), "CENTER"),
("ALIGN", (2, 0), (2, -1), "CENTER"),
("ALIGN", (3, 0), (3, -1), "CENTER"),
("ALIGN", (4, 0), (4, -1), "RIGHT"),
("ALIGN", (5, 0), (5, -1), "RIGHT"),
# Grid
("LINEBELOW", (0, 0), (-1, 0), 1, colors.HexColor("#1a237e")),
("LINEBELOW", (0, 1), (-1, -2), 0.5, colors.HexColor("#e0e0e0")),
("LINEBELOW", (0, -1), (-1, -1), 1, colors.HexColor("#1a237e")),
# Alternating row colors
*[
("BACKGROUND", (0, i), (-1, i), colors.HexColor("#f5f5f5"))
for i in range(2, len(table_data), 2)
],
]
)
)
elements.append(table)
elements.append(Spacer(1, 0.15 * inch))
# Total row
total_data = [
["", "", "", "", "TOTAL:", f"${subtotal:,.2f}"],
]
total_table = Table(total_data, colWidths=col_widths)
total_table.setStyle(
TableStyle(
[
("FONTNAME", (0, 0), (-1, -1), "Helvetica-Bold"),
("FONTSIZE", (0, 0), (-1, -1), 11),
("ALIGN", (4, 0), (4, 0), "RIGHT"),
("ALIGN", (5, 0), (5, 0), "RIGHT"),
("TOPPADDING", (0, 0), (-1, -1), 4),
("LINEABOVE", (4, 0), (5, 0), 1.5, colors.HexColor("#1a237e")),
]
)
)
elements.append(total_table)
return elements
def _build_terms(styles) -> list:
elements = [
Spacer(1, 0.2 * inch),
Paragraph("Terms & Conditions", styles["SectionHeader"]),
]
for line in TERMS_AND_CONDITIONS.split("\n"):
elements.append(Paragraph(line, styles["TermsText"]))
return elements
def _page_footer(canvas, doc):
canvas.saveState()
page_num = canvas.getPageNumber()
footer_text = f"Page {page_num}"
canvas.setFont("Helvetica", 8)
canvas.setFillColor(colors.HexColor("#888888"))
canvas.drawCentredString(letter[0] / 2, 0.4 * inch, footer_text)
canvas.drawString(
0.75 * inch,
0.4 * inch,
f"{COMPANY_NAME} — Confidential",
)
canvas.restoreState()
def upload_pdf(s3_key: str, pdf_bytes: bytes):
try:
s3.put_object(
Bucket=GENERATED_BUCKET,
Key=s3_key,
Body=pdf_bytes,
ContentType="application/pdf",
)
except Exception as e:
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error uploading PDF: %s", e)
raise
def register_pdf(proposal_id: str, s3_key: str):
try:
resp = _retry_request(
"POST",
f"{API_BASE_URL}/api/generated-pdfs",
json={"proposalId": proposal_id, "s3Key": s3_key},
headers=_api_headers(),
)
if resp.status_code not in (200, 201):
logger.error("Failed to register PDF: %s %s", resp.status_code, resp.text)
except Exception as e:
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error registering PDF: %s", e)
def _api_headers() -> dict:
headers = {"Content-Type": "application/json"}
api_key = _get_api_key()
if api_key:
headers["X-Internal-Api-Key"] = api_key
return headers
# Fix (v1 PR1): SigV4-sign internal calls to the .NET API Function URL (authType=AWS_IAM).
# The X-Internal-Api-Key header is preserved for the application-layer check; SigV4
# satisfies the transport-layer IAM auth the Function URL enforces.
_SIGV4_SERVICE = "lambda"
_boto_session = boto3.Session()
def _sign_request_headers(method: str, url: str, body: bytes, headers: dict) -> dict:
"""Return headers with a SigV4 signature for the AWS_IAM Function URL call.
Falls back to the unsigned headers when no AWS credentials are resolvable, so a
non-IAM target still works in local development.
"""
creds = _boto_session.get_credentials()
if creds is None:
return headers
region = os.environ.get("AWS_REGION") or os.environ.get(
"AWS_DEFAULT_REGION", "us-east-1"
)
aws_request = AWSRequest(method=method, url=url, data=body, headers=dict(headers))
SigV4Auth(creds, _SIGV4_SERVICE, region).add_auth(aws_request)
return dict(aws_request.headers)
def _retry_request(
method: str, url: str, *, max_retries: int = 3, **kwargs
) -> httpx.Response:
kwargs.setdefault("timeout", 10)
base_headers = dict(kwargs.pop("headers", None) or {})
# Serialize the body once so the bytes we sign are exactly the bytes we send:
# SigV4 hashes the payload, so httpx must not re-serialize a json= kwarg.
if "json" in kwargs:
body = json.dumps(kwargs.pop("json")).encode("utf-8")
base_headers.setdefault("Content-Type", "application/json")
elif "content" in kwargs:
raw = kwargs.pop("content")
body = raw if isinstance(raw, bytes) else (raw or "").encode("utf-8")
else:
body = b""
kwargs["content"] = body
last_resp = None
for attempt in range(max_retries):
# Sign per attempt so a slow retry never sends an expired SigV4 timestamp.
kwargs["headers"] = _sign_request_headers(method, url, body, base_headers)
try:
resp = httpx.request(method, url, **kwargs)
if resp.status_code < 500:
return resp
last_resp = resp
except (httpx.ConnectError, httpx.ReadTimeout, httpx.WriteTimeout) as exc:
if attempt == max_retries - 1:
raise
logger.warning(
"Retryable error (attempt %d/%d): %s", attempt + 1, max_retries, exc
)
time.sleep(min(2**attempt, 4))
if last_resp is not None:
return last_resp
raise RuntimeError(f"All {max_retries} retries failed for {method} {url}")