proposal-system/lambdas/library-ingest/app.py
Adam Moussa 669e9c0e43 fix(lambdas): LAM-M2, M3, M6, M9 — prompt injection, PDF size check, numeric validation, API key TTL
LAM-M2: Add sanitize_user_text() to suggestions Lambda that strips common
prompt injection patterns (blocklist + delimiter neutralisation) before
including user-supplied text in Bedrock prompts.

LAM-M3: Add file size check in pdf-extract before downloading — rejects
PDFs over 50 MB with a logged warning and ValueError.

LAM-M6: Add validate_line_item_numerics() to suggestions Lambda that
rejects Bedrock-generated line items with negative values, NaN/Inf, or
amounts exceeding $10M ceiling.

LAM-M9: Replace indefinite API key cache with 5-minute TTL in all four
Lambdas (suggestions, pdf-extract, pdf-generate, library-ingest) so
rotated Secrets Manager values take effect promptly.
2026-05-27 18:18:44 -04:00

292 lines
9.6 KiB
Python

"""Proposal System - Library Ingest Lambda.
Processes approved/sent proposals into the Bedrock Knowledge Base library.
Formats proposal data as structured markdown and uploads to the library bucket,
then triggers a KB sync.
"""
import json
import logging
import os
import re
import time
from datetime import datetime
import boto3
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", "")
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")
bedrock_agent = boto3.client("bedrock-agent")
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
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"]
process_ingestion(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 process_ingestion(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)
document = format_proposal_document(proposal, line_items)
s3_key = upload_to_library(proposal, document)
if s3_key:
trigger_kb_sync()
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 format_proposal_document(proposal: dict, line_items: list[dict]) -> str:
total = sum(li.get("totalPrice", 0) for li in line_items)
submitted_at = proposal.get("submittedAt", "")
if submitted_at:
try:
dt = datetime.fromisoformat(submitted_at.replace("Z", "+00:00"))
submitted_at = dt.strftime("%Y-%m-%d")
except (ValueError, TypeError):
pass
lines = [
f"# Proposal: {proposal['proposalNumber']}",
"",
f"**Customer:** {proposal['customerName']}",
f"**Address:** {proposal.get('customerAddress', '')}",
f"**Service Category:** {proposal['serviceCategory']}",
f"**Priority:** {proposal['priority']}",
f"**Date:** {submitted_at}",
f"**Total Bid Amount:** ${total:,.2f}",
f"**Work Order:** {proposal.get('workOrderNumber', '')}",
"",
"## Scope of Work",
"",
proposal.get("refinedScope") or proposal.get("scopeOfWork", ""),
"",
"## Line Items",
"",
"| # | Description | Qty | Unit | Unit Price | Total |",
"|---|---|---|---|---|---|",
]
for i, li in enumerate(line_items, 1):
desc = li.get("description", "")
qty = li.get("quantity", "")
unit = li.get("unit", "")
unit_price = li.get("unitPrice")
total_price = li.get("totalPrice", 0)
up_str = f"${unit_price:,.2f}" if unit_price else "-"
tp_str = f"${total_price:,.2f}"
lines.append(f"| {i} | {desc} | {qty} | {unit} | {up_str} | {tp_str} |")
lines.extend(
[
"",
f"**Total: ${total:,.2f}**",
]
)
return "\n".join(lines)
def upload_to_library(proposal: dict, document: str) -> str | None:
if not LIBRARY_BUCKET:
logger.warning("No library bucket configured")
return None
proposal_number = proposal["proposalNumber"]
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,
Key=s3_key,
Body=document.encode("utf-8"),
ContentType="text/markdown",
Metadata={
"service-category": category,
"proposal-number": proposal_number,
"customer-name": proposal.get("customerName", ""),
"total-amount": str(proposal.get("totalBidAmount", 0)),
"date-submitted": proposal.get("submittedAt", ""),
},
)
logger.info("Uploaded %s to library bucket", s3_key)
return s3_key
except Exception as e:
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error uploading to library: %s", e)
return None
def trigger_kb_sync():
if not KNOWLEDGE_BASE_ID or not DATA_SOURCE_ID:
logger.info("KB or data source ID not configured, skipping sync")
return
try:
response = bedrock_agent.start_ingestion_job(
knowledgeBaseId=KNOWLEDGE_BASE_ID,
dataSourceId=DATA_SOURCE_ID,
)
job_id = response.get("ingestionJob", {}).get("ingestionJobId", "")
logger.info("Started KB ingestion job: %s", job_id)
except Exception as e:
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error triggering KB sync: %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
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
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}")