proposal-system/lambdas/pdf-extract/app.py
Adam Moussa 99e0c16505 Merge main into feature/fix-phase-2, resolve infra conflicts
Keep both Phase 1 (JWT authorizer, webClientId/mobileClientId props) and
Phase 2 (alarmTopic, CloudWatch alarms) changes in CDK stacks.
2026-05-20 19:09:35 -04:00

365 lines
11 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 tempfile
import time
import boto3
import httpx
import pdfplumber
logger = logging.getLogger(__name__)
logger.setLevel(os.environ.get("LOG_LEVEL", "INFO"))
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
def _get_api_key() -> str:
global _cached_api_key
if _cached_api_key is None:
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 = ""
return _cached_api_key
def handler(event, context):
batch_item_failures = []
for record in event.get("Records", []):
try:
body = json.loads(record["body"])
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
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)
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:
logger.error("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:
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:
logger.error("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:
logger.error("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:
logger.error("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:
logger.error("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
def _retry_request(
method: str, url: str, *, max_retries: int = 3, **kwargs
) -> httpx.Response:
kwargs.setdefault("timeout", 10)
for attempt in range(max_retries):
try:
resp = httpx.request(method, url, **kwargs)
if resp.status_code < 500:
return 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))
return resp # type: ignore[possibly-undefined]