proposal-system/lambdas/pdf-extract/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

494 lines
17 KiB
Python

"""Proposal System - PDF Extract Lambda.
Parses vendor proposal PDFs and extracts structured line item data.
Falls back to Claude multimodal for scanned/image-based PDFs.
"""
import base64
import json
import logging
import os
import re
import tempfile
import time
import boto3
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
import httpx
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")
INTERNAL_API_KEY_SECRET_ARN = os.environ.get("INTERNAL_API_KEY_SECRET_ARN", "")
s3 = boto3.client("s3")
bedrock_runtime = boto3.client("bedrock-runtime")
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-M3 — maximum PDF file size (50 MB)
_MAX_PDF_SIZE_BYTES = 50 * 1024 * 1024
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 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 = []
# 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.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:
# 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_pdf(proposal_id: str, s3_key: str, vendor_proposal_id: str):
update_processing_status(vendor_proposal_id, "Processing")
pdf_path = None
try:
pdf_path = download_pdf(s3_key)
extracted = extract_with_pdfplumber(pdf_path)
if not extracted["lineItems"] and extracted["rawText"].strip():
extracted = extract_with_claude_multimodal(pdf_path)
if not extracted["lineItems"] and not extracted["rawText"].strip():
extracted = extract_with_claude_multimodal(pdf_path)
save_extraction(vendor_proposal_id, extracted)
except Exception as e:
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error processing PDF: %s", e)
update_processing_status(vendor_proposal_id, "Failed")
finally:
if pdf_path:
try:
os.unlink(pdf_path)
except Exception:
pass
def download_pdf(s3_key: str) -> str:
# Fix: LAM-M3 — check file size before downloading to avoid processing
# excessively large PDFs that could exhaust Lambda memory/tmp storage.
head = s3.head_object(Bucket=UPLOADS_BUCKET, Key=s3_key)
file_size = head.get("ContentLength", 0)
if file_size > _MAX_PDF_SIZE_BYTES:
logger.warning(
"Fix: LAM-M3 — PDF too large: key=%s size=%d bytes (max=%d)",
s3_key,
file_size,
_MAX_PDF_SIZE_BYTES,
)
raise ValueError(
f"PDF file size ({file_size} bytes) exceeds maximum "
f"allowed size ({_MAX_PDF_SIZE_BYTES} bytes) for key: {s3_key}"
)
tmp = tempfile.NamedTemporaryFile(delete=False, suffix=".pdf")
s3.download_file(UPLOADS_BUCKET, s3_key, tmp.name)
tmp.close()
return tmp.name
def extract_with_pdfplumber(pdf_path: str) -> dict:
result = {
"vendorName": "",
"lineItems": [],
"rawText": "",
"totalVendorCost": 0.0,
}
try:
with pdfplumber.open(pdf_path) as pdf:
all_text = ""
all_tables = []
for page in pdf.pages:
text = page.extract_text() or ""
all_text += text + "\n"
tables = page.extract_tables()
for table in tables:
all_tables.append(table)
result["rawText"] = all_text.strip()
if all_tables:
result["lineItems"] = parse_tables(all_tables)
result["totalVendorCost"] = sum(
li.get("total", 0) for li in result["lineItems"]
)
if not result["vendorName"] and all_text:
lines = all_text.split("\n")
for line in lines[:5]:
stripped = line.strip()
if stripped and len(stripped) > 3 and not stripped[0].isdigit():
result["vendorName"] = stripped
break
except Exception as e:
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("pdfplumber extraction failed: %s", e)
return result
def parse_tables(tables: list) -> list[dict]:
line_items = []
for table in tables:
if not table or len(table) < 2:
continue
header = [str(cell).lower().strip() if cell else "" for cell in table[0]]
desc_col = find_column(
header, ["description", "item", "service", "work", "scope"]
)
qty_col = find_column(header, ["qty", "quantity", "count"])
price_col = find_column(
header, ["unit price", "rate", "price/unit", "unit cost"]
)
total_col = find_column(
header, ["total", "amount", "ext", "extended", "line total"]
)
if desc_col is None:
continue
for row in table[1:]:
if not row or len(row) <= desc_col:
continue
description = str(row[desc_col]).strip() if row[desc_col] else ""
if not description or description.lower() in (
"",
"total",
"subtotal",
"grand total",
):
continue
quantity = (
parse_number(row[qty_col])
if qty_col is not None and qty_col < len(row)
else None
)
unit_price = (
parse_number(row[price_col])
if price_col is not None and price_col < len(row)
else None
)
total = (
parse_number(row[total_col])
if total_col is not None and total_col < len(row)
else None
)
if total is None and quantity and unit_price:
total = quantity * unit_price
if description and (total or unit_price):
line_items.append(
{
"description": description,
"quantity": quantity,
"unitPrice": unit_price,
"total": total,
}
)
return line_items
def find_column(header: list[str], keywords: list[str]) -> int | None:
for i, col in enumerate(header):
for kw in keywords:
if kw in col:
return i
return None
def parse_number(value) -> float | None:
if value is None:
return None
try:
cleaned = str(value).replace("$", "").replace(",", "").strip()
if not cleaned or cleaned == "-":
return None
return float(cleaned)
except (ValueError, TypeError):
return None
def extract_with_claude_multimodal(pdf_path: str) -> dict:
try:
with open(pdf_path, "rb") as f:
pdf_bytes = f.read()
pdf_b64 = base64.standard_b64encode(pdf_bytes).decode("utf-8")
response = bedrock_runtime.invoke_model(
modelId=MODEL_ID,
contentType="application/json",
accept="application/json",
body=json.dumps(
{
"anthropic_version": "bedrock-2023-05-31",
"max_tokens": 4096,
"messages": [
{
"role": "user",
"content": [
{
"type": "document",
"source": {
"type": "base64",
"media_type": "application/pdf",
"data": pdf_b64,
},
},
{
"type": "text",
"text": """Extract all line items from this vendor proposal PDF.
Return a JSON object with these fields:
- vendorName: the vendor/company name
- lineItems: array of objects with: description, quantity (number or null), unitPrice (number or null), total (number or null)
- totalVendorCost: the grand total amount
Respond ONLY with the JSON object, no additional text.""",
},
],
}
],
"temperature": 0.1,
}
),
)
response_body = json.loads(response["body"].read())
content = response_body["content"][0]["text"]
content = content.strip()
if content.startswith("```"):
content = content.split("\n", 1)[1]
content = content.rsplit("```", 1)[0]
parsed = json.loads(content)
return {
"vendorName": parsed.get("vendorName", ""),
"lineItems": parsed.get("lineItems", []),
"rawText": "",
"totalVendorCost": float(parsed.get("totalVendorCost", 0)),
}
except Exception as e:
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Claude multimodal extraction failed: %s", e)
return {
"vendorName": "",
"lineItems": [],
"rawText": "",
"totalVendorCost": 0.0,
}
def save_extraction(vendor_proposal_id: str, extracted: dict):
extracted_data = {
"lineItems": extracted["lineItems"],
"vendorName": extracted["vendorName"],
}
try:
resp = _retry_request(
"PUT",
f"{API_BASE_URL}/api/vendor-proposals/{vendor_proposal_id}",
json={
"vendorName": extracted["vendorName"],
"extractedData": json.dumps(extracted_data),
"totalVendorCost": extracted["totalVendorCost"],
"processingStatus": "Complete",
},
headers=_api_headers(),
)
if resp.status_code not in (200, 204):
logger.error(
"Failed to save extraction: %s %s", resp.status_code, resp.text
)
except Exception as 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):
if not vendor_proposal_id:
return
try:
_retry_request(
"PUT",
f"{API_BASE_URL}/api/vendor-proposals/{vendor_proposal_id}/status",
json={"processingStatus": status},
headers=_api_headers(),
)
except Exception as e:
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error updating status: %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}")