proposal-system/lambdas/suggestions/app.py

576 lines
20 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
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
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:
# Fix: CONC-L1 — the bulk PUT is version-guarded (ADR 0004): echo the
# proposal's current rowVersion as proposalVersion, and retry once with
# a fresh token if a concurrent edit wins the race (409/stale 422).
for attempt in range(2):
proposal = fetch_proposal(proposal_id)
token = (proposal or {}).get("rowVersion")
if not token:
logger.error(
"Cannot post line items: no rowVersion for proposal %s", proposal_id
)
return
resp = _retry_request(
"PUT",
f"{API_BASE_URL}/api/proposals/{proposal_id}/line-items",
json={"lineItems": line_items_payload, "proposalVersion": token},
headers=_api_headers(),
timeout=15,
)
if resp.status_code in (200, 201):
return
if resp.status_code == 409 and attempt == 0:
logger.warning(
"Concurrency conflict posting line items for %s; retrying with fresh token",
proposal_id,
)
continue
logger.error(
"Failed to post line items: %s %s", resp.status_code, resp.text
)
return
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
# 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}")