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

504 lines
18 KiB
Python

"""Proposal System - Suggestion Engine Lambda.
Queries Bedrock Knowledge Base for similar proposals and invokes Claude
to generate line item suggestions for new proposals.
"""
import json
import logging
import math
import os
import re
import time
import boto3
import httpx
logger = logging.getLogger(__name__)
logger.setLevel(os.environ.get("LOG_LEVEL", "INFO"))
KNOWLEDGE_BASE_ID = os.environ.get("KNOWLEDGE_BASE_ID", "")
MODEL_ID = os.environ.get("MODEL_ID", "us.anthropic.claude-sonnet-4-5-20250929-v1:0")
API_BASE_URL = os.environ.get("API_BASE_URL", "")
INTERNAL_API_KEY_SECRET_ARN = os.environ.get("INTERNAL_API_KEY_SECRET_ARN", "")
bedrock_agent = boto3.client("bedrock-agent-runtime")
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
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
# ---------------------------------------------------------------------------
# Fix: LAM-M2 — Prompt injection mitigation
# ---------------------------------------------------------------------------
# Blocklist patterns that commonly appear in prompt injection attempts.
# This is a simple defense-in-depth measure, not a full NLP solution.
_INJECTION_PATTERNS = [
re.compile(r"ignore\s+(all\s+)?previous\s+instructions", re.IGNORECASE),
re.compile(r"ignore\s+(all\s+)?above\s+instructions", re.IGNORECASE),
re.compile(r"disregard\s+(all\s+)?previous", re.IGNORECASE),
re.compile(r"override\s+(all\s+)?instructions", re.IGNORECASE),
re.compile(r"you\s+are\s+now\s+(a|an)\s+", re.IGNORECASE),
re.compile(r"new\s+instructions?\s*:", re.IGNORECASE),
re.compile(r"system\s*:", re.IGNORECASE),
re.compile(r"<\|?\s*(system|im_start|endoftext)\s*\|?>", re.IGNORECASE),
re.compile(r"\[INST\]", re.IGNORECASE),
re.compile(r"```\s*(system|instruction)", re.IGNORECASE),
]
# Delimiter sequences that could be used to break out of the user-text section
_INJECTION_DELIMITERS = ["```", "---\n", "===\n", "***\n"]
def sanitize_user_text(text: str) -> str:
"""Strip common prompt injection patterns from user-supplied text.
Fix: LAM-M2 — basic prompt injection mitigation before including user
text in Bedrock prompts. Applies a blocklist of known injection patterns
and neutralises delimiter sequences.
"""
if not text:
return text
sanitized = text
# Remove blocklisted patterns
for pattern in _INJECTION_PATTERNS:
sanitized = pattern.sub("[removed]", sanitized)
# Neutralise delimiter sequences by replacing them with a safe alternative
for delim in _INJECTION_DELIMITERS:
sanitized = sanitized.replace(delim, " ")
return sanitized
# ---------------------------------------------------------------------------
# Fix: LAM-M6 — Numeric validation for suggestion amounts
# ---------------------------------------------------------------------------
_MAX_LINE_ITEM_VALUE = 10_000_000 # $10M ceiling for any single value
def validate_line_item_numerics(items: list[dict]) -> list[dict]:
"""Validate and filter line items with unreasonable numeric values.
Fix: LAM-M6 — after getting suggestions from Bedrock, reject items with
negative values, NaN, or values exceeding the $10M ceiling.
"""
validated = []
for item in items:
rejected = False
for field in ("quantity", "unitPrice", "totalPrice"):
value = item.get(field)
if value is None:
continue
try:
num = float(value)
except (ValueError, TypeError):
logger.warning(
"Fix: LAM-M6 — rejected line item: %s is not a valid number (%r) "
"in item %r",
field, value, item.get("description", "unknown"),
)
rejected = True
break
if math.isnan(num) or math.isinf(num):
logger.warning(
"Fix: LAM-M6 — rejected line item: %s is NaN/Inf in item %r",
field, item.get("description", "unknown"),
)
rejected = True
break
if num < 0:
logger.warning(
"Fix: LAM-M6 — rejected line item: %s is negative (%.2f) in item %r",
field, num, item.get("description", "unknown"),
)
rejected = True
break
if num > _MAX_LINE_ITEM_VALUE:
logger.warning(
"Fix: LAM-M6 — rejected line item: %s exceeds max (%.2f > %d) in item %r",
field, num, _MAX_LINE_ITEM_VALUE, item.get("description", "unknown"),
)
rejected = True
break
if not rejected:
validated.append(item)
return validated
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"]
trigger = payload.get("trigger", "generate")
process_suggestion(proposal_id, trigger)
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_suggestion(proposal_id: str, trigger: str):
proposal = fetch_proposal(proposal_id)
if not proposal:
logger.warning("Proposal %s not found", proposal_id)
return
scope = proposal.get("refinedScope") or proposal.get("scopeOfWork", "")
category = proposal.get("serviceCategory", "")
priority = proposal.get("priority", "")
existing_items = fetch_line_items(proposal_id)
has_ai_items = any(li.get("source") == "AI" for li in existing_items)
if has_ai_items:
logger.info("AI items already exist for %s, skipping regeneration", proposal_id)
return
status = proposal.get("status", "")
if status not in ("InReview", "Revised"):
logger.info("Proposal %s is in status %s, skipping suggestions", proposal_id, status)
return
similar_proposals = retrieve_similar(sanitize_user_text(scope), category)
suggested_items = generate_line_items(scope, category, priority, similar_proposals)
if not suggested_items and not existing_items:
logger.warning(
"No suggestions generated and no existing items for %s, skipping status update",
proposal_id,
)
return
post_line_items(proposal_id, suggested_items, existing_items)
store_similar_references(proposal_id, similar_proposals)
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 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 retrieve_similar(scope: str, category: str) -> list[dict]:
if not KNOWLEDGE_BASE_ID:
logger.info("No Knowledge Base configured, skipping retrieval")
return []
try:
filter_config = (
{"equals": {"key": "service_category", "value": category}}
if category
else None
)
params = {
"knowledgeBaseId": KNOWLEDGE_BASE_ID,
"retrievalQuery": {"text": scope},
"retrievalConfiguration": {
"vectorSearchConfiguration": {
"numberOfResults": 10,
}
},
}
if filter_config:
params["retrievalConfiguration"]["vectorSearchConfiguration"]["filter"] = (
filter_config
)
response = bedrock_agent.retrieve(**params)
results = []
for result in response.get("retrievalResults", []):
content = result.get("content", {}).get("text", "")
score = result.get("score", 0.0)
metadata = result.get("metadata", {})
source_uri = result.get("location", {}).get("s3Location", {}).get("uri", "")
results.append(
{
"content": content,
"score": score,
"metadata": metadata,
"sourceUri": source_uri,
}
)
return results
except Exception as e:
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error retrieving from KB: %s", e)
return []
def generate_line_items(
scope: str,
category: str,
priority: str,
similar_proposals: list[dict],
) -> list[dict]:
# Fix: LAM-M2 — sanitize user-supplied text before including in prompt
safe_scope = sanitize_user_text(scope)
safe_category = sanitize_user_text(category)
safe_priority = sanitize_user_text(priority)
context_block = ""
if similar_proposals:
context_block = "Here are similar historical proposals and their line items for reference:\n\n"
for i, sp in enumerate(similar_proposals[:5], 1):
context_block += (
f"--- Similar Proposal {i} (relevance: {sp['score']:.2f}) ---\n"
)
context_block += sp["content"] + "\n\n"
prompt = f"""You are a construction/facilities proposal estimator for Sea Haven Industries.
Based on the scope of work and similar historical proposals, generate a detailed list of line items
with quantities, units, and estimated pricing.
Service Category: {safe_category}
Priority: {safe_priority}
Scope of Work:
{safe_scope}
{context_block}
Generate line items as a JSON array. Each item should have:
- description: clear description of the work/material
- quantity: numeric quantity
- unit: unit of measurement (e.g., "sq ft", "hours", "each", "linear ft")
- unitPrice: price per unit in dollars (or null if lump sum)
- totalPrice: total price for this line item in dollars
- pricingMode: "UnitPrice" if unit price provided, "TotalPrice" if lump sum
Respond ONLY with the JSON array, no additional text."""
try:
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": prompt}],
"temperature": 0.3,
}
),
)
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]
line_items = json.loads(content)
if not isinstance(line_items, list):
return []
# Fix: LAM-M6 — validate numeric fields before returning suggestions
return validate_line_item_numerics(line_items)
except Exception as e:
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error generating line items: %s", e)
return []
def post_line_items(proposal_id: str, items: list[dict], existing_items: list[dict]):
if not items and not existing_items:
return
line_items_payload = []
# Preserve non-AI items (Manual, Vendor, Historical)
preserved = [li for li in existing_items if li.get("source") != "AI"]
for i, li in enumerate(preserved):
line_items_payload.append(
{
"id": li.get("id"),
"description": li["description"],
"quantity": float(li.get("quantity", 1)),
"unit": li.get("unit", "each"),
"unitPrice": li.get("unitPrice"),
"totalPrice": float(li.get("totalPrice", 0)),
"pricingMode": li.get("pricingMode", "TotalPrice"),
"sortOrder": i + 1,
"source": li.get("source", "Manual"),
}
)
# Add new AI-generated items after preserved ones
offset = len(line_items_payload)
for i, item in enumerate(items):
pricing_mode = item.get("pricingMode", "TotalPrice")
if pricing_mode not in ("UnitPrice", "TotalPrice", "Both"):
pricing_mode = "UnitPrice" if item.get("unitPrice") else "TotalPrice"
line_items_payload.append(
{
"id": None,
"description": item["description"],
"quantity": float(item.get("quantity", 1)),
"unit": item.get("unit", "each"),
"unitPrice": item.get("unitPrice"),
"totalPrice": float(item.get("totalPrice", 0)),
"pricingMode": pricing_mode,
"sortOrder": offset + i + 1,
"source": "AI",
}
)
try:
resp = _retry_request(
"PUT",
f"{API_BASE_URL}/api/proposals/{proposal_id}/line-items",
json={"lineItems": line_items_payload},
headers=_api_headers(),
timeout=15,
)
if resp.status_code not in (200, 201):
logger.error(
"Failed to post line items: %s %s", resp.status_code, resp.text
)
except Exception as 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]):
if not similar_proposals:
return
for sp in similar_proposals[:5]:
source_uri = sp.get("sourceUri", "")
library_item_id = source_uri.split("/")[-1] if source_uri else ""
if not library_item_id:
continue
try:
_retry_request(
"POST",
f"{API_BASE_URL}/api/proposals/{proposal_id}/similar-references",
json={
"referencedLibraryItemId": library_item_id,
"similarityScore": sp["score"],
},
headers=_api_headers(),
)
except Exception as e:
# Fix: LAM-M5 — include stack trace in error logging
logger.exception("Error storing similar reference: %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}")