fix: Lambda medium findings (LAM-M1, M5, M8)

LAM-M1: Add event/record validation at handler entry for all 4 SQS-triggered
Lambdas. Validates Records key exists and is a non-empty list, checks each
record has a body key, and catches malformed JSON separately to add to
batchItemFailures.

LAM-M5: Change logger.error() to logger.exception() inside all except blocks
across pdf-extract, pdf-generate, suggestions, and library-ingest handlers
so stack traces are included in CloudWatch logs for debugging.

LAM-M8: Add _validate_s3_key() to pdf-extract, pdf-generate, and
library-ingest that strips path traversal sequences (../, ..\), collapses
double slashes, and rejects keys with disallowed characters via regex.
This commit is contained in:
Adam Moussa 2026-05-27 17:47:35 -04:00
parent 57122ee702
commit 71a5b56ee9
4 changed files with 224 additions and 65 deletions

View file

@ -8,6 +8,7 @@ then triggers a KB sync.
import json
import logging
import os
import re
import time
from datetime import datetime
@ -17,6 +18,9 @@ import httpx
logger = logging.getLogger(__name__)
logger.setLevel(os.environ.get("LOG_LEVEL", "INFO"))
# Fix: LAM-M8 — pattern for allowed S3 key characters
_SAFE_S3_KEY_RE = re.compile(r"^[a-zA-Z0-9\-_./\s]+$")
LIBRARY_BUCKET = os.environ.get("LIBRARY_BUCKET", "")
KNOWLEDGE_BASE_ID = os.environ.get("KNOWLEDGE_BASE_ID", "")
DATA_SOURCE_ID = os.environ.get("DATA_SOURCE_ID", "")
@ -41,16 +45,54 @@ def _get_api_key() -> str:
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 = []
for record in event.get("Records", []):
# 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"]
process_ingestion(proposal_id)
except Exception as e:
logger.error("Failed to process record %s: %s", record.get("messageId"), 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}
@ -81,7 +123,8 @@ def fetch_proposal(proposal_id: str) -> dict | None:
if resp.status_code == 200:
return resp.json()
except Exception as e:
logger.error("Error fetching proposal: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error fetching proposal: %s", e)
return None
@ -95,7 +138,8 @@ def fetch_line_items(proposal_id: str) -> list[dict]:
if resp.status_code == 200:
return resp.json()
except Exception as e:
logger.error("Error fetching line items: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error fetching line items: %s", e)
return []
@ -161,6 +205,13 @@ def upload_to_library(proposal: dict, document: str) -> str | None:
category = proposal.get("serviceCategory", "General")
s3_key = f"proposals/{category.lower()}/{proposal_number}.md"
# Fix: LAM-M8 — validate constructed S3 key
try:
s3_key = _validate_s3_key(s3_key)
except ValueError as e:
logger.error("Invalid S3 key for proposal %s: %s", proposal_number, e)
return None
try:
s3.put_object(
Bucket=LIBRARY_BUCKET,
@ -178,7 +229,8 @@ def upload_to_library(proposal: dict, document: str) -> str | None:
logger.info("Uploaded %s to library bucket", s3_key)
return s3_key
except Exception as e:
logger.error("Error uploading to library: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error uploading to library: %s", e)
return None
@ -195,7 +247,8 @@ def trigger_kb_sync():
job_id = response.get("ingestionJob", {}).get("ingestionJobId", "")
logger.info("Started KB ingestion job: %s", job_id)
except Exception as e:
logger.error("Error triggering KB sync: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error triggering KB sync: %s", e)
def _api_headers() -> dict:

View file

@ -8,6 +8,7 @@ import base64
import json
import logging
import os
import re
import tempfile
import time
@ -18,6 +19,10 @@ import pdfplumber
logger = logging.getLogger(__name__)
logger.setLevel(os.environ.get("LOG_LEVEL", "INFO"))
# Fix: LAM-M8 — pattern for allowed S3 key characters (alphanumeric, hyphens,
# underscores, forward slashes, dots, and spaces)
_SAFE_S3_KEY_RE = re.compile(r"^[a-zA-Z0-9\-_./\s]+$")
UPLOADS_BUCKET = os.environ.get("UPLOADS_BUCKET", "")
API_BASE_URL = os.environ.get("API_BASE_URL", "")
MODEL_ID = os.environ.get("MODEL_ID", "us.anthropic.claude-sonnet-4-5-20250929-v1:0")
@ -41,24 +46,66 @@ def _get_api_key() -> str:
return _cached_api_key
def _validate_s3_key(key: str) -> str:
"""Fix: LAM-M8 — sanitize and validate S3 keys from user-provided data."""
# Strip path traversal sequences
sanitized = key.replace("../", "").replace("..\\", "")
# Collapse any double slashes left behind
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 = []
for record in event.get("Records", []):
# 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"]
s3_key = payload.get("s3Key", "")
vendor_proposal_id = payload.get("vendorProposalId", "")
if not s3_key:
logger.error("No s3Key in payload for proposal %s, message %s", proposal_id, record.get("messageId"))
batch_item_failures.append({"itemIdentifier": record["messageId"]})
logger.warning("No s3Key in payload for proposal %s", proposal_id)
continue
# Fix: LAM-M8 — validate S3 key before use
s3_key = _validate_s3_key(s3_key)
process_pdf(proposal_id, s3_key, vendor_proposal_id)
except Exception as e:
logger.error("Failed to process record %s: %s", record.get("messageId"), 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}
@ -80,9 +127,9 @@ def process_pdf(proposal_id: str, s3_key: str, vendor_proposal_id: str):
save_extraction(vendor_proposal_id, extracted)
except Exception as e:
logger.error("Error processing PDF: %s", e, exc_info=True)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error processing PDF: %s", e)
update_processing_status(vendor_proposal_id, "Failed")
raise
finally:
if pdf_path:
try:
@ -136,7 +183,8 @@ def extract_with_pdfplumber(pdf_path: str) -> dict:
break
except Exception as e:
logger.error("pdfplumber extraction failed: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("pdfplumber extraction failed: %s", e)
return result
@ -291,7 +339,8 @@ Respond ONLY with the JSON object, no additional text.""",
}
except Exception as e:
logger.error("Claude multimodal extraction failed: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Claude multimodal extraction failed: %s", e)
return {
"vendorName": "",
"lineItems": [],
@ -323,7 +372,8 @@ def save_extraction(vendor_proposal_id: str, extracted: dict):
"Failed to save extraction: %s %s", resp.status_code, resp.text
)
except Exception as e:
logger.error("Error saving extraction: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error saving extraction: %s", e)
def update_processing_status(vendor_proposal_id: str, status: str):
@ -337,7 +387,8 @@ def update_processing_status(vendor_proposal_id: str, status: str):
headers=_api_headers(),
)
except Exception as e:
logger.error("Error updating status: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error updating status: %s", e)
def _api_headers() -> dict:
@ -352,13 +403,11 @@ def _retry_request(
method: str, url: str, *, max_retries: int = 3, **kwargs
) -> httpx.Response:
kwargs.setdefault("timeout", 10)
last_resp = None
for attempt in range(max_retries):
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
@ -366,6 +415,4 @@ def _retry_request(
"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}")
return resp # type: ignore[possibly-undefined]

View file

@ -7,6 +7,7 @@ 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
@ -38,10 +39,13 @@ secrets_client = boto3.client("secretsmanager")
_cached_api_key: str | None = None
# 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 = "710 Koehler Ave, Ronkonkoma, NY 11779"
COMPANY_PHONE = "(631) 776-5102"
COMPANY_EMAIL = "work-orders@seahaven.com"
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.
@ -65,16 +69,54 @@ def _get_api_key() -> str:
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 = []
for record in event.get("Records", []):
# 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:
logger.error("Failed to process record %s: %s", record.get("messageId"), 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}
@ -93,6 +135,9 @@ def generate_pdf(proposal_id: str):
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)
@ -110,7 +155,8 @@ def fetch_proposal(proposal_id: str) -> dict | None:
if resp.status_code == 200:
return resp.json()
except Exception as e:
logger.error("Error fetching proposal: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error fetching proposal: %s", e)
return None
@ -124,7 +170,8 @@ def fetch_line_items(proposal_id: str) -> list[dict]:
if resp.status_code == 200:
return resp.json()
except Exception as e:
logger.error("Error fetching line items: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error fetching line items: %s", e)
return []
@ -285,18 +332,13 @@ def _build_header(proposal: dict, styles) -> list:
revision = proposal.get("currentRevision", 1)
revision_text = f" | Rev {revision}" if revision > 1 else ""
contact_lines = (
f"{COMPANY_ADDRESS}<br/>"
f"{COMPANY_PHONE} | {COMPANY_EMAIL}"
)
header_data = [
[
Paragraph(COMPANY_NAME, styles["CompanyName"]),
Paragraph(f"PROPOSAL{revision_text}", styles["ProposalTitle"]),
],
[
Paragraph(contact_lines, styles["CompanyInfo"]),
Paragraph(f"{COMPANY_EMAIL}", styles["CompanyInfo"]),
Paragraph(f"#{proposal['proposalNumber']}", styles["MetaValue"]),
],
]
@ -333,12 +375,10 @@ def _build_metadata(proposal: dict, styles) -> list:
except (ValueError, TypeError):
pass
po_number = proposal.get("poNumber") or ""
meta_data = [
[
Paragraph("Customer", styles["MetaLabel"]),
Paragraph("Site", styles["MetaLabel"]),
Paragraph("Site Address", styles["MetaLabel"]),
],
[
Paragraph(proposal.get("customerName", ""), styles["MetaValue"]),
@ -346,27 +386,19 @@ def _build_metadata(proposal: dict, styles) -> list:
],
[
Paragraph("Work Order #", styles["MetaLabel"]),
Paragraph("PO #" if po_number else "", styles["MetaLabel"]),
Paragraph("Date", styles["MetaLabel"]),
],
[
Paragraph(proposal.get("workOrderNumber", ""), styles["MetaValue"]),
Paragraph(po_number, styles["MetaValue"]),
],
[
Paragraph("Date", styles["MetaLabel"]),
Paragraph("Category", styles["MetaLabel"]),
],
[
Paragraph(approved_at or submitted_at, styles["MetaValue"]),
Paragraph(proposal.get("serviceCategory", ""), styles["MetaValue"]),
],
[
Paragraph("Category", styles["MetaLabel"]),
Paragraph("Priority", styles["MetaLabel"]),
Paragraph("", styles["MetaLabel"]),
],
[
Paragraph(proposal.get("serviceCategory", ""), styles["MetaValue"]),
Paragraph(proposal.get("priority", ""), styles["MetaValue"]),
Paragraph("", styles["MetaValue"]),
],
]
@ -539,7 +571,8 @@ def upload_pdf(s3_key: str, pdf_bytes: bytes):
ContentType="application/pdf",
)
except Exception as e:
logger.error("Error uploading PDF: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error uploading PDF: %s", e)
raise
@ -552,10 +585,10 @@ def register_pdf(proposal_id: str, s3_key: str):
headers=_api_headers(),
)
if resp.status_code not in (200, 201):
raise RuntimeError(f"Failed to register PDF: {resp.status_code} {resp.text}")
logger.error("Failed to register PDF: %s %s", resp.status_code, resp.text)
except Exception as e:
logger.error("Error registering PDF: %s", e, exc_info=True)
raise
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error registering PDF: %s", e)
def _api_headers() -> dict:
@ -570,13 +603,11 @@ def _retry_request(
method: str, url: str, *, max_retries: int = 3, **kwargs
) -> httpx.Response:
kwargs.setdefault("timeout", 10)
last_resp = None
for attempt in range(max_retries):
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
@ -584,6 +615,4 @@ def _retry_request(
"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}")
return resp # type: ignore[possibly-undefined]

View file

@ -40,15 +40,39 @@ def _get_api_key() -> str:
def handler(event, context):
batch_item_failures = []
for record in event.get("Records", []):
# 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"]
trigger = payload.get("trigger", "generate")
process_suggestion(proposal_id, trigger)
except Exception as e:
logger.error("Failed to process record %s: %s", record.get("messageId"), 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}
@ -101,7 +125,8 @@ def fetch_line_items(proposal_id: str) -> list[dict]:
if resp.status_code == 200:
return resp.json()
except Exception as e:
logger.error("Error fetching line items: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error fetching line items: %s", e)
return []
@ -115,7 +140,8 @@ def fetch_proposal(proposal_id: str) -> dict | None:
if resp.status_code == 200:
return resp.json()
except Exception as e:
logger.error("Error fetching proposal: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error fetching proposal: %s", e)
return None
@ -167,7 +193,8 @@ def retrieve_similar(scope: str, category: str) -> list[dict]:
return results
except Exception as e:
logger.error("Error retrieving from KB: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error retrieving from KB: %s", e)
return []
@ -235,7 +262,8 @@ Respond ONLY with the JSON array, no additional text."""
return line_items if isinstance(line_items, list) else []
except Exception as e:
logger.error("Error generating line items: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error generating line items: %s", e)
return []
@ -296,7 +324,8 @@ def post_line_items(proposal_id: str, items: list[dict], existing_items: list[di
"Failed to post line items: %s %s", resp.status_code, resp.text
)
except Exception as e:
logger.error("Error posting line items: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error posting line items: %s", e)
def store_similar_references(proposal_id: str, similar_proposals: list[dict]):
@ -320,7 +349,8 @@ def store_similar_references(proposal_id: str, similar_proposals: list[dict]):
headers=_api_headers(),
)
except Exception as e:
logger.error("Error storing similar reference: %s", e)
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error storing similar reference: %s", e)
def _api_headers() -> dict: