"""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}")