Fix security gaps and improve code quality across API and Lambdas (#23)

Security: add system identity claims to InternalApiKeyMiddleware so
Lambda-to-API calls resolve a proper user, inject ICurrentUserService
into GeneratedPdfsController to replace Guid.Empty, and consolidate
CurrentUserService into a single ResolveAsync lookup chain.

Quality: replace four COUNT queries in ProposalService.GetStatsAsync
with a single grouped query, convert all Lambda print() to structured
logging, and add retry helpers for Lambda-to-API HTTP calls.
This commit is contained in:
Adam Moussa 2026-05-17 13:48:14 -04:00 • committed by GitHub
parent 36b39c73df
commit f051f74fde
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 221 additions and 142 deletions

View file

@ -1,6 +1,7 @@
using Microsoft.AspNetCore.Authorization; using Microsoft.AspNetCore.Authorization;
using Microsoft.AspNetCore.Mvc; using Microsoft.AspNetCore.Mvc;
using Microsoft.EntityFrameworkCore; using Microsoft.EntityFrameworkCore;
using ProposalSystem.Application.Interfaces;
using ProposalSystem.Domain.Entities; using ProposalSystem.Domain.Entities;
using ProposalSystem.Infrastructure.Data; using ProposalSystem.Infrastructure.Data;
@ -12,10 +13,12 @@ namespace ProposalSystem.Api.Controllers;
public class GeneratedPdfsController : ControllerBase public class GeneratedPdfsController : ControllerBase
{ {
private readonly ProposalDbContext _db; private readonly ProposalDbContext _db;
private readonly ICurrentUserService _currentUser;
public GeneratedPdfsController(ProposalDbContext db) public GeneratedPdfsController(ProposalDbContext db, ICurrentUserService currentUser)
{ {
_db = db; _db = db;
_currentUser = currentUser;
} }
[HttpPost] [HttpPost]
@ -24,6 +27,8 @@ public class GeneratedPdfsController : ControllerBase
var proposal = await _db.Proposals.FindAsync(new object[] { request.ProposalId }, ct); var proposal = await _db.Proposals.FindAsync(new object[] { request.ProposalId }, ct);
if (proposal == null) return NotFound(); if (proposal == null) return NotFound();
await _currentUser.ResolveAsync();
var pdf = new GeneratedPdf var pdf = new GeneratedPdf
{ {
Id = Guid.NewGuid(), Id = Guid.NewGuid(),
@ -31,7 +36,7 @@ public class GeneratedPdfsController : ControllerBase
Revision = proposal.CurrentRevision, Revision = proposal.CurrentRevision,
S3Key = request.S3Key, S3Key = request.S3Key,
GeneratedAt = DateTime.UtcNow, GeneratedAt = DateTime.UtcNow,
GeneratedById = Guid.Empty, GeneratedById = _currentUser.UserId,
}; };
_db.GeneratedPdfs.Add(pdf); _db.GeneratedPdfs.Add(pdf);

View file

@ -28,6 +28,10 @@ public class InternalApiKeyMiddleware
var claims = new[] var claims = new[]
{ {
new Claim(ClaimTypes.NameIdentifier, "system"), new Claim(ClaimTypes.NameIdentifier, "system"),
new Claim("sub", "system-lambda-caller"),
new Claim(ClaimTypes.Email, "system@proposal-system.internal"),
new Claim("email", "system@proposal-system.internal"),
new Claim("name", "System"),
new Claim(ClaimTypes.Role, "admins"), new Claim(ClaimTypes.Role, "admins"),
new Claim("cognito:groups", "admins"), new Claim("cognito:groups", "admins"),
}; };

View file

@ -18,9 +18,9 @@ public class CurrentUserService : ICurrentUserService
_db = db; _db = db;
} }
public Guid UserId => GetUser().Id; public Guid UserId => GetOrThrow().Id;
public string Email => GetUser().Email; public string Email => GetOrThrow().Email;
public UserRole Role => GetUser().Role; public UserRole Role => GetOrThrow().Role;
public string? IpAddress => public string? IpAddress =>
_httpContext.HttpContext?.Connection.RemoteIpAddress?.ToString(); _httpContext.HttpContext?.Connection.RemoteIpAddress?.ToString();
@ -32,122 +32,64 @@ public class CurrentUserService : ICurrentUserService
var principal = _httpContext.HttpContext?.User var principal = _httpContext.HttpContext?.User
?? throw new UnauthorizedAccessException("No authenticated user"); ?? throw new UnauthorizedAccessException("No authenticated user");
var userId = principal.FindFirstValue(ClaimTypes.NameIdentifier);
var sub = principal.FindFirstValue("sub"); var sub = principal.FindFirstValue("sub");
var email = principal.FindFirstValue(ClaimTypes.Email)
?? principal.FindFirstValue("email")
?? "unknown@seahaven.com";
var name = principal.FindFirstValue("name")
?? email.Split('@')[0];
var groups = principal.FindAll("cognito:groups")
.Select(c => c.Value).ToList();
var role = groups.Contains("sysadmins") ? UserRole.SysAdmin
: groups.Contains("admins") ? UserRole.Admin
: UserRole.Dispatcher;
var userId = principal.FindFirstValue(ClaimTypes.NameIdentifier);
if (userId != null && Guid.TryParse(userId, out var parsedId)) if (userId != null && Guid.TryParse(userId, out var parsedId))
{
_cachedUser = await _db.Users.FirstOrDefaultAsync(u => u.Id == parsedId); _cachedUser = await _db.Users.FirstOrDefaultAsync(u => u.Id == parsedId);
}
if (_cachedUser == null && sub != null) if (_cachedUser == null && sub != null)
{
_cachedUser = await _db.Users.FirstOrDefaultAsync(u => u.CognitoSub == sub); _cachedUser = await _db.Users.FirstOrDefaultAsync(u => u.CognitoSub == sub);
}
if (_cachedUser == null) if (_cachedUser == null)
{
var email = principal.FindFirstValue(ClaimTypes.Email)
?? principal.FindFirstValue("email")
?? "unknown@seahaven.com";
_cachedUser = await _db.Users.FirstOrDefaultAsync(u => u.Email == email); _cachedUser = await _db.Users.FirstOrDefaultAsync(u => u.Email == email);
}
if (_cachedUser == null) if (_cachedUser != null)
{ {
var email = principal.FindFirstValue(ClaimTypes.Email) var changed = false;
?? principal.FindFirstValue("email") if (_cachedUser.Email != email) { _cachedUser.Email = email; changed = true; }
?? "unknown@seahaven.com"; if (_cachedUser.DisplayName != name) { _cachedUser.DisplayName = name; changed = true; }
if (changed)
var name = principal.FindFirstValue("name")
?? email.Split('@')[0];
var groups = principal.FindAll("cognito:groups")
.Select(c => c.Value).ToList();
var role = groups.Contains("sysadmins") ? UserRole.SysAdmin
: groups.Contains("admins") ? UserRole.Admin
: UserRole.Dispatcher;
_cachedUser = new User
{ {
Id = Guid.NewGuid(), _cachedUser.UpdatedAt = DateTime.UtcNow;
CognitoSub = sub ?? $"auto-{Guid.NewGuid():N}", await _db.SaveChangesAsync();
Email = email, }
DisplayName = name, return;
Role = role,
IsActive = true,
CreatedAt = DateTime.UtcNow,
UpdatedAt = DateTime.UtcNow,
};
_db.Users.Add(_cachedUser);
await _db.SaveChangesAsync();
} }
_cachedUser = new User
{
Id = Guid.NewGuid(),
CognitoSub = sub ?? $"auto-{Guid.NewGuid():N}",
Email = email,
DisplayName = name,
Role = role,
IsActive = true,
CreatedAt = DateTime.UtcNow,
UpdatedAt = DateTime.UtcNow,
};
_db.Users.Add(_cachedUser);
await _db.SaveChangesAsync();
} }
private User GetUser() private User GetOrThrow()
{ {
if (_cachedUser != null) return _cachedUser; if (_cachedUser != null) return _cachedUser;
var principal = _httpContext.HttpContext?.User ResolveAsync().GetAwaiter().GetResult();
?? throw new UnauthorizedAccessException("No authenticated user");
var userId = principal.FindFirstValue(ClaimTypes.NameIdentifier); return _cachedUser
var sub = principal.FindFirstValue("sub"); ?? throw new UnauthorizedAccessException("Could not resolve current user");
if (userId != null && Guid.TryParse(userId, out var parsedId))
{
_cachedUser = _db.Users.FirstOrDefault(u => u.Id == parsedId);
}
if (_cachedUser == null && sub != null)
{
_cachedUser = _db.Users.FirstOrDefault(u => u.CognitoSub == sub);
}
if (_cachedUser == null)
{
var email = principal.FindFirstValue(ClaimTypes.Email)
?? principal.FindFirstValue("email")
?? "unknown@seahaven.com";
_cachedUser = _db.Users.FirstOrDefault(u => u.Email == email);
}
if (_cachedUser == null)
{
var email = principal.FindFirstValue(ClaimTypes.Email)
?? principal.FindFirstValue("email")
?? "unknown@seahaven.com";
var name = principal.FindFirstValue("name")
?? email.Split('@')[0];
var groups = principal.FindAll("cognito:groups")
.Select(c => c.Value).ToList();
var role = groups.Contains("sysadmins") ? UserRole.SysAdmin
: groups.Contains("admins") ? UserRole.Admin
: UserRole.Dispatcher;
_cachedUser = new User
{
Id = Guid.NewGuid(),
CognitoSub = sub ?? $"auto-{Guid.NewGuid():N}",
Email = email,
DisplayName = name,
Role = role,
IsActive = true,
CreatedAt = DateTime.UtcNow,
UpdatedAt = DateTime.UtcNow,
};
_db.Users.Add(_cachedUser);
_db.SaveChanges();
}
return _cachedUser;
} }
} }

View file

@ -303,14 +303,22 @@ public class ProposalService : IProposalService
public async Task<ProposalStatsResponse> GetStatsAsync(CancellationToken ct = default) public async Task<ProposalStatsResponse> GetStatsAsync(CancellationToken ct = default)
{ {
var userId = _currentUser.UserId; var userId = _currentUser.UserId;
var baseQuery = _db.Proposals.Where(p => p.SubmittedById == userId);
var total = await baseQuery.CountAsync(ct); var counts = await _db.Proposals
var inReview = await baseQuery.CountAsync(p => p.Status == ProposalStatus.InReview, ct); .Where(p => p.SubmittedById == userId)
var approved = await baseQuery.CountAsync(p => p.Status == ProposalStatus.Approved, ct); .GroupBy(_ => 1)
var sent = await baseQuery.CountAsync(p => p.Status == ProposalStatus.Sent, ct); .Select(g => new
{
Total = g.Count(),
InReview = g.Count(p => p.Status == ProposalStatus.InReview),
Approved = g.Count(p => p.Status == ProposalStatus.Approved),
Sent = g.Count(p => p.Status == ProposalStatus.Sent),
})
.FirstOrDefaultAsync(ct);
return new ProposalStatsResponse(total, inReview, approved, sent); return counts == null
? new ProposalStatsResponse(0, 0, 0, 0)
: new ProposalStatsResponse(counts.Total, counts.InReview, counts.Approved, counts.Sent);
} }
private static ProposalResponse MapToResponse(Proposal p) => new( private static ProposalResponse MapToResponse(Proposal p) => new(

View file

@ -6,12 +6,17 @@ then triggers a KB sync.
""" """
import json import json
import logging
import os import os
import time
from datetime import datetime from datetime import datetime
import boto3 import boto3
import httpx import httpx
logger = logging.getLogger(__name__)
logger.setLevel(os.environ.get("LOG_LEVEL", "INFO"))
LIBRARY_BUCKET = os.environ.get("LIBRARY_BUCKET", "") LIBRARY_BUCKET = os.environ.get("LIBRARY_BUCKET", "")
KNOWLEDGE_BASE_ID = os.environ.get("KNOWLEDGE_BASE_ID", "") KNOWLEDGE_BASE_ID = os.environ.get("KNOWLEDGE_BASE_ID", "")
DATA_SOURCE_ID = os.environ.get("DATA_SOURCE_ID", "") DATA_SOURCE_ID = os.environ.get("DATA_SOURCE_ID", "")
@ -48,7 +53,7 @@ def handler(event, context):
def process_ingestion(proposal_id: str): def process_ingestion(proposal_id: str):
proposal = fetch_proposal(proposal_id) proposal = fetch_proposal(proposal_id)
if not proposal: if not proposal:
print(f"Proposal {proposal_id} not found") logger.warning("Proposal %s not found", proposal_id)
return return
line_items = fetch_line_items(proposal_id) line_items = fetch_line_items(proposal_id)
@ -71,7 +76,7 @@ def fetch_proposal(proposal_id: str) -> dict | None:
if resp.status_code == 200: if resp.status_code == 200:
return resp.json() return resp.json()
except Exception as e: except Exception as e:
print(f"Error fetching proposal: {e}") logger.error("Error fetching proposal: %s", e)
return None return None
@ -85,7 +90,7 @@ def fetch_line_items(proposal_id: str) -> list[dict]:
if resp.status_code == 200: if resp.status_code == 200:
return resp.json() return resp.json()
except Exception as e: except Exception as e:
print(f"Error fetching line items: {e}") logger.error("Error fetching line items: %s", e)
return [] return []
@ -144,7 +149,7 @@ def format_proposal_document(proposal: dict, line_items: list[dict]) -> str:
def upload_to_library(proposal: dict, document: str) -> str | None: def upload_to_library(proposal: dict, document: str) -> str | None:
if not LIBRARY_BUCKET: if not LIBRARY_BUCKET:
print("No library bucket configured") logger.warning("No library bucket configured")
return None return None
proposal_number = proposal["proposalNumber"] proposal_number = proposal["proposalNumber"]
@ -165,16 +170,16 @@ def upload_to_library(proposal: dict, document: str) -> str | None:
"date-submitted": proposal.get("submittedAt", ""), "date-submitted": proposal.get("submittedAt", ""),
}, },
) )
print(f"Uploaded {s3_key} to library bucket") logger.info("Uploaded %s to library bucket", s3_key)
return s3_key return s3_key
except Exception as e: except Exception as e:
print(f"Error uploading to library: {e}") logger.error("Error uploading to library: %s", e)
return None return None
def trigger_kb_sync(): def trigger_kb_sync():
if not KNOWLEDGE_BASE_ID or not DATA_SOURCE_ID: if not KNOWLEDGE_BASE_ID or not DATA_SOURCE_ID:
print("KB or data source ID not configured, skipping sync") logger.info("KB or data source ID not configured, skipping sync")
return return
try: try:
@ -183,9 +188,9 @@ def trigger_kb_sync():
dataSourceId=DATA_SOURCE_ID, dataSourceId=DATA_SOURCE_ID,
) )
job_id = response.get("ingestionJob", {}).get("ingestionJobId", "") job_id = response.get("ingestionJob", {}).get("ingestionJobId", "")
print(f"Started KB ingestion job: {job_id}") logger.info("Started KB ingestion job: %s", job_id)
except Exception as e: except Exception as e:
print(f"Error triggering KB sync: {e}") logger.error("Error triggering KB sync: %s", e)
def _api_headers() -> dict: def _api_headers() -> dict:
@ -194,3 +199,27 @@ def _api_headers() -> dict:
if api_key: if api_key:
headers["X-Internal-Api-Key"] = api_key headers["X-Internal-Api-Key"] = api_key
return headers return headers
def _api_request(method: str, url: str, retries: int = 3, **kwargs) -> httpx.Response:
kwargs.setdefault("headers", _api_headers())
kwargs.setdefault("timeout", 10)
for attempt in range(retries):
try:
resp = httpx.request(method, url, **kwargs)
if resp.status_code < 500:
return resp
logger.warning(
"API returned %s on attempt %d for %s",
resp.status_code,
attempt + 1,
url,
)
except httpx.TransportError as e:
logger.warning(
"Transport error on attempt %d for %s: %s", attempt + 1, url, e
)
if attempt == retries - 1:
raise
time.sleep(min(2**attempt, 4))
return resp # type: ignore[possibly-undefined]

View file

@ -4,15 +4,20 @@ Parses vendor proposal PDFs and extracts structured line item data.
Falls back to Claude multimodal for scanned/image-based PDFs. Falls back to Claude multimodal for scanned/image-based PDFs.
""" """
import base64
import json import json
import logging
import os import os
import tempfile import tempfile
import base64 import time
import boto3 import boto3
import httpx import httpx
import pdfplumber import pdfplumber
logger = logging.getLogger(__name__)
logger.setLevel(os.environ.get("LOG_LEVEL", "INFO"))
UPLOADS_BUCKET = os.environ.get("UPLOADS_BUCKET", "") UPLOADS_BUCKET = os.environ.get("UPLOADS_BUCKET", "")
API_BASE_URL = os.environ.get("API_BASE_URL", "") 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") MODEL_ID = os.environ.get("MODEL_ID", "us.anthropic.claude-sonnet-4-5-20250929-v1:0")
@ -45,7 +50,7 @@ def handler(event, context):
vendor_proposal_id = payload.get("vendorProposalId", "") vendor_proposal_id = payload.get("vendorProposalId", "")
if not s3_key: if not s3_key:
print(f"No s3Key in payload for proposal {proposal_id}") logger.warning("No s3Key in payload for proposal %s", proposal_id)
continue continue
process_pdf(proposal_id, s3_key, vendor_proposal_id) process_pdf(proposal_id, s3_key, vendor_proposal_id)
@ -69,7 +74,7 @@ def process_pdf(proposal_id: str, s3_key: str, vendor_proposal_id: str):
save_extraction(vendor_proposal_id, extracted) save_extraction(vendor_proposal_id, extracted)
except Exception as e: except Exception as e:
print(f"Error processing PDF: {e}") logger.error("Error processing PDF: %s", e)
update_processing_status(vendor_proposal_id, "Failed") update_processing_status(vendor_proposal_id, "Failed")
finally: finally:
if pdf_path: if pdf_path:
@ -124,7 +129,7 @@ def extract_with_pdfplumber(pdf_path: str) -> dict:
break break
except Exception as e: except Exception as e:
print(f"pdfplumber extraction failed: {e}") logger.error("pdfplumber extraction failed: %s", e)
return result return result
@ -279,7 +284,7 @@ Respond ONLY with the JSON object, no additional text.""",
} }
except Exception as e: except Exception as e:
print(f"Claude multimodal extraction failed: {e}") logger.error("Claude multimodal extraction failed: %s", e)
return { return {
"vendorName": "", "vendorName": "",
"lineItems": [], "lineItems": [],
@ -307,9 +312,11 @@ def save_extraction(vendor_proposal_id: str, extracted: dict):
timeout=10, timeout=10,
) )
if resp.status_code not in (200, 204): if resp.status_code not in (200, 204):
print(f"Failed to save extraction: {resp.status_code} {resp.text}") logger.error(
"Failed to save extraction: %s %s", resp.status_code, resp.text
)
except Exception as e: except Exception as e:
print(f"Error saving extraction: {e}") logger.error("Error saving extraction: %s", e)
def update_processing_status(vendor_proposal_id: str, status: str): def update_processing_status(vendor_proposal_id: str, status: str):
@ -323,7 +330,7 @@ def update_processing_status(vendor_proposal_id: str, status: str):
timeout=10, timeout=10,
) )
except Exception as e: except Exception as e:
print(f"Error updating status: {e}") logger.error("Error updating status: %s", e)
def _api_headers() -> dict: def _api_headers() -> dict:
@ -332,3 +339,27 @@ def _api_headers() -> dict:
if api_key: if api_key:
headers["X-Internal-Api-Key"] = api_key headers["X-Internal-Api-Key"] = api_key
return headers return headers
def _api_request(method: str, url: str, retries: int = 3, **kwargs) -> httpx.Response:
kwargs.setdefault("headers", _api_headers())
kwargs.setdefault("timeout", 10)
for attempt in range(retries):
try:
resp = httpx.request(method, url, **kwargs)
if resp.status_code < 500:
return resp
logger.warning(
"API returned %s on attempt %d for %s",
resp.status_code,
attempt + 1,
url,
)
except httpx.TransportError as e:
logger.warning(
"Transport error on attempt %d for %s: %s", attempt + 1, url, e
)
if attempt == retries - 1:
raise
time.sleep(min(2**attempt, 4))
return resp # type: ignore[possibly-undefined]

View file

@ -5,7 +5,9 @@ Triggered via SQS when an admin requests PDF generation.
""" """
import json import json
import logging
import os import os
import time
from datetime import datetime from datetime import datetime
from io import BytesIO from io import BytesIO
@ -24,6 +26,9 @@ from reportlab.platypus import (
TableStyle, TableStyle,
) )
logger = logging.getLogger(__name__)
logger.setLevel(os.environ.get("LOG_LEVEL", "INFO"))
GENERATED_BUCKET = os.environ.get("GENERATED_BUCKET", "") GENERATED_BUCKET = os.environ.get("GENERATED_BUCKET", "")
API_BASE_URL = os.environ.get("API_BASE_URL", "") API_BASE_URL = os.environ.get("API_BASE_URL", "")
INTERNAL_API_KEY_SECRET_ARN = os.environ.get("INTERNAL_API_KEY_SECRET_ARN", "") INTERNAL_API_KEY_SECRET_ARN = os.environ.get("INTERNAL_API_KEY_SECRET_ARN", "")
@ -72,7 +77,7 @@ def handler(event, context):
def generate_pdf(proposal_id: str): def generate_pdf(proposal_id: str):
proposal = fetch_proposal(proposal_id) proposal = fetch_proposal(proposal_id)
if not proposal: if not proposal:
print(f"Proposal {proposal_id} not found") logger.warning("Proposal %s not found", proposal_id)
return return
line_items = fetch_line_items(proposal_id) line_items = fetch_line_items(proposal_id)
@ -87,7 +92,7 @@ def generate_pdf(proposal_id: str):
register_pdf(proposal_id, s3_key) register_pdf(proposal_id, s3_key)
print(f"Generated PDF: {s3_key} ({len(pdf_bytes)} bytes)") logger.info("Generated PDF: %s (%d bytes)", s3_key, len(pdf_bytes))
def fetch_proposal(proposal_id: str) -> dict | None: def fetch_proposal(proposal_id: str) -> dict | None:
@ -100,7 +105,7 @@ def fetch_proposal(proposal_id: str) -> dict | None:
if resp.status_code == 200: if resp.status_code == 200:
return resp.json() return resp.json()
except Exception as e: except Exception as e:
print(f"Error fetching proposal: {e}") logger.error("Error fetching proposal: %s", e)
return None return None
@ -114,7 +119,7 @@ def fetch_line_items(proposal_id: str) -> list[dict]:
if resp.status_code == 200: if resp.status_code == 200:
return resp.json() return resp.json()
except Exception as e: except Exception as e:
print(f"Error fetching line items: {e}") logger.error("Error fetching line items: %s", e)
return [] return []
@ -514,7 +519,7 @@ def upload_pdf(s3_key: str, pdf_bytes: bytes):
ContentType="application/pdf", ContentType="application/pdf",
) )
except Exception as e: except Exception as e:
print(f"Error uploading PDF: {e}") logger.error("Error uploading PDF: %s", e)
raise raise
@ -527,9 +532,9 @@ def register_pdf(proposal_id: str, s3_key: str):
timeout=10, timeout=10,
) )
if resp.status_code not in (200, 201): if resp.status_code not in (200, 201):
print(f"Failed to register PDF: {resp.status_code} {resp.text}") logger.error("Failed to register PDF: %s %s", resp.status_code, resp.text)
except Exception as e: except Exception as e:
print(f"Error registering PDF: {e}") logger.error("Error registering PDF: %s", e)
def _api_headers() -> dict: def _api_headers() -> dict:
@ -538,3 +543,27 @@ def _api_headers() -> dict:
if api_key: if api_key:
headers["X-Internal-Api-Key"] = api_key headers["X-Internal-Api-Key"] = api_key
return headers return headers
def _api_request(method: str, url: str, retries: int = 3, **kwargs) -> httpx.Response:
kwargs.setdefault("headers", _api_headers())
kwargs.setdefault("timeout", 10)
for attempt in range(retries):
try:
resp = httpx.request(method, url, **kwargs)
if resp.status_code < 500:
return resp
logger.warning(
"API returned %s on attempt %d for %s",
resp.status_code,
attempt + 1,
url,
)
except httpx.TransportError as e:
logger.warning(
"Transport error on attempt %d for %s: %s", attempt + 1, url, e
)
if attempt == retries - 1:
raise
time.sleep(min(2**attempt, 4))
return resp # type: ignore[possibly-undefined]

View file

@ -5,11 +5,16 @@ to generate line item suggestions for new proposals.
""" """
import json import json
import logging
import os import os
import time
import boto3 import boto3
import httpx import httpx
logger = logging.getLogger(__name__)
logger.setLevel(os.environ.get("LOG_LEVEL", "INFO"))
KNOWLEDGE_BASE_ID = os.environ.get("KNOWLEDGE_BASE_ID", "") 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") 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", "") API_BASE_URL = os.environ.get("API_BASE_URL", "")
@ -46,7 +51,7 @@ def handler(event, context):
def process_suggestion(proposal_id: str, trigger: str): def process_suggestion(proposal_id: str, trigger: str):
proposal = fetch_proposal(proposal_id) proposal = fetch_proposal(proposal_id)
if not proposal: if not proposal:
print(f"Proposal {proposal_id} not found") logger.warning("Proposal %s not found", proposal_id)
return return
scope = proposal.get("refinedScope") or proposal.get("scopeOfWork", "") scope = proposal.get("refinedScope") or proposal.get("scopeOfWork", "")
@ -76,7 +81,7 @@ def fetch_line_items(proposal_id: str) -> list[dict]:
if resp.status_code == 200: if resp.status_code == 200:
return resp.json() return resp.json()
except Exception as e: except Exception as e:
print(f"Error fetching line items: {e}") logger.error("Error fetching line items: %s", e)
return [] return []
@ -90,13 +95,13 @@ def fetch_proposal(proposal_id: str) -> dict | None:
if resp.status_code == 200: if resp.status_code == 200:
return resp.json() return resp.json()
except Exception as e: except Exception as e:
print(f"Error fetching proposal: {e}") logger.error("Error fetching proposal: %s", e)
return None return None
def retrieve_similar(scope: str, category: str) -> list[dict]: def retrieve_similar(scope: str, category: str) -> list[dict]:
if not KNOWLEDGE_BASE_ID: if not KNOWLEDGE_BASE_ID:
print("No Knowledge Base configured, skipping retrieval") logger.info("No Knowledge Base configured, skipping retrieval")
return [] return []
try: try:
@ -142,7 +147,7 @@ def retrieve_similar(scope: str, category: str) -> list[dict]:
return results return results
except Exception as e: except Exception as e:
print(f"Error retrieving from KB: {e}") logger.error("Error retrieving from KB: %s", e)
return [] return []
@ -210,7 +215,7 @@ Respond ONLY with the JSON array, no additional text."""
return line_items if isinstance(line_items, list) else [] return line_items if isinstance(line_items, list) else []
except Exception as e: except Exception as e:
print(f"Error generating line items: {e}") logger.error("Error generating line items: %s", e)
return [] return []
@ -266,9 +271,11 @@ def post_line_items(proposal_id: str, items: list[dict], existing_items: list[di
timeout=15, timeout=15,
) )
if resp.status_code not in (200, 201): if resp.status_code not in (200, 201):
print(f"Failed to post line items: {resp.status_code} {resp.text}") logger.error(
"Failed to post line items: %s %s", resp.status_code, resp.text
)
except Exception as e: except Exception as e:
print(f"Error posting line items: {e}") logger.error("Error posting line items: %s", e)
def store_similar_references(proposal_id: str, similar_proposals: list[dict]): def store_similar_references(proposal_id: str, similar_proposals: list[dict]):
@ -292,7 +299,7 @@ def store_similar_references(proposal_id: str, similar_proposals: list[dict]):
timeout=10, timeout=10,
) )
except Exception as e: except Exception as e:
print(f"Error storing similar reference: {e}") logger.error("Error storing similar reference: %s", e)
def update_status_to_in_review(proposal_id: str): def update_status_to_in_review(proposal_id: str):
@ -304,7 +311,7 @@ def update_status_to_in_review(proposal_id: str):
timeout=10, timeout=10,
) )
except Exception as e: except Exception as e:
print(f"Error updating status: {e}") logger.error("Error updating status: %s", e)
def _api_headers() -> dict: def _api_headers() -> dict:
@ -313,3 +320,27 @@ def _api_headers() -> dict:
if api_key: if api_key:
headers["X-Internal-Api-Key"] = api_key headers["X-Internal-Api-Key"] = api_key
return headers return headers
def _api_request(method: str, url: str, retries: int = 3, **kwargs) -> httpx.Response:
kwargs.setdefault("headers", _api_headers())
kwargs.setdefault("timeout", 10)
for attempt in range(retries):
try:
resp = httpx.request(method, url, **kwargs)
if resp.status_code < 500:
return resp
logger.warning(
"API returned %s on attempt %d for %s",
resp.status_code,
attempt + 1,
url,
)
except httpx.TransportError as e:
logger.warning(
"Transport error on attempt %d for %s: %s", attempt + 1, url, e
)
if attempt == retries - 1:
raise
time.sleep(min(2**attempt, 4))
return resp # type: ignore[possibly-undefined]