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

View file

@ -28,6 +28,10 @@ public class InternalApiKeyMiddleware
var claims = new[]
{
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("cognito:groups", "admins"),
};

View file

@ -18,9 +18,9 @@ public class CurrentUserService : ICurrentUserService
_db = db;
}
public Guid UserId => GetUser().Id;
public string Email => GetUser().Email;
public UserRole Role => GetUser().Role;
public Guid UserId => GetOrThrow().Id;
public string Email => GetOrThrow().Email;
public UserRole Role => GetOrThrow().Role;
public string? IpAddress =>
_httpContext.HttpContext?.Connection.RemoteIpAddress?.ToString();
@ -32,122 +32,64 @@ public class CurrentUserService : ICurrentUserService
var principal = _httpContext.HttpContext?.User
?? throw new UnauthorizedAccessException("No authenticated user");
var userId = principal.FindFirstValue(ClaimTypes.NameIdentifier);
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))
{
_cachedUser = await _db.Users.FirstOrDefaultAsync(u => u.Id == parsedId);
}
if (_cachedUser == null && sub != null)
{
_cachedUser = await _db.Users.FirstOrDefaultAsync(u => u.CognitoSub == sub);
}
if (_cachedUser == null)
{
var email = principal.FindFirstValue(ClaimTypes.Email)
?? principal.FindFirstValue("email")
?? "unknown@seahaven.com";
_cachedUser = await _db.Users.FirstOrDefaultAsync(u => u.Email == email);
}
if (_cachedUser == null)
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
var changed = false;
if (_cachedUser.Email != email) { _cachedUser.Email = email; changed = true; }
if (_cachedUser.DisplayName != name) { _cachedUser.DisplayName = name; changed = true; }
if (changed)
{
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();
_cachedUser.UpdatedAt = DateTime.UtcNow;
await _db.SaveChangesAsync();
}
return;
}
_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;
var principal = _httpContext.HttpContext?.User
?? throw new UnauthorizedAccessException("No authenticated user");
ResolveAsync().GetAwaiter().GetResult();
var userId = principal.FindFirstValue(ClaimTypes.NameIdentifier);
var sub = principal.FindFirstValue("sub");
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;
return _cachedUser
?? throw new UnauthorizedAccessException("Could not resolve current user");
}
}

View file

@ -303,14 +303,22 @@ public class ProposalService : IProposalService
public async Task<ProposalStatsResponse> GetStatsAsync(CancellationToken ct = default)
{
var userId = _currentUser.UserId;
var baseQuery = _db.Proposals.Where(p => p.SubmittedById == userId);
var total = await baseQuery.CountAsync(ct);
var inReview = await baseQuery.CountAsync(p => p.Status == ProposalStatus.InReview, ct);
var approved = await baseQuery.CountAsync(p => p.Status == ProposalStatus.Approved, ct);
var sent = await baseQuery.CountAsync(p => p.Status == ProposalStatus.Sent, ct);
var counts = await _db.Proposals
.Where(p => p.SubmittedById == userId)
.GroupBy(_ => 1)
.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(

View file

@ -6,12 +6,17 @@ then triggers a KB sync.
"""
import json
import logging
import os
import time
from datetime import datetime
import boto3
import httpx
logger = logging.getLogger(__name__)
logger.setLevel(os.environ.get("LOG_LEVEL", "INFO"))
LIBRARY_BUCKET = os.environ.get("LIBRARY_BUCKET", "")
KNOWLEDGE_BASE_ID = os.environ.get("KNOWLEDGE_BASE_ID", "")
DATA_SOURCE_ID = os.environ.get("DATA_SOURCE_ID", "")
@ -48,7 +53,7 @@ def handler(event, context):
def process_ingestion(proposal_id: str):
proposal = fetch_proposal(proposal_id)
if not proposal:
print(f"Proposal {proposal_id} not found")
logger.warning("Proposal %s not found", proposal_id)
return
line_items = fetch_line_items(proposal_id)
@ -71,7 +76,7 @@ def fetch_proposal(proposal_id: str) -> dict | None:
if resp.status_code == 200:
return resp.json()
except Exception as e:
print(f"Error fetching proposal: {e}")
logger.error("Error fetching proposal: %s", e)
return None
@ -85,7 +90,7 @@ def fetch_line_items(proposal_id: str) -> list[dict]:
if resp.status_code == 200:
return resp.json()
except Exception as e:
print(f"Error fetching line items: {e}")
logger.error("Error fetching line items: %s", e)
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:
if not LIBRARY_BUCKET:
print("No library bucket configured")
logger.warning("No library bucket configured")
return None
proposal_number = proposal["proposalNumber"]
@ -165,16 +170,16 @@ def upload_to_library(proposal: dict, document: str) -> str | None:
"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
except Exception as e:
print(f"Error uploading to library: {e}")
logger.error("Error uploading to library: %s", e)
return None
def trigger_kb_sync():
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
try:
@ -183,9 +188,9 @@ def trigger_kb_sync():
dataSourceId=DATA_SOURCE_ID,
)
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:
print(f"Error triggering KB sync: {e}")
logger.error("Error triggering KB sync: %s", e)
def _api_headers() -> dict:
@ -194,3 +199,27 @@ def _api_headers() -> dict:
if api_key:
headers["X-Internal-Api-Key"] = api_key
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.
"""
import base64
import json
import logging
import os
import tempfile
import base64
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")
@ -45,7 +50,7 @@ def handler(event, context):
vendor_proposal_id = payload.get("vendorProposalId", "")
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
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)
except Exception as e:
print(f"Error processing PDF: {e}")
logger.error("Error processing PDF: %s", e)
update_processing_status(vendor_proposal_id, "Failed")
finally:
if pdf_path:
@ -124,7 +129,7 @@ def extract_with_pdfplumber(pdf_path: str) -> dict:
break
except Exception as e:
print(f"pdfplumber extraction failed: {e}")
logger.error("pdfplumber extraction failed: %s", e)
return result
@ -279,7 +284,7 @@ Respond ONLY with the JSON object, no additional text.""",
}
except Exception as e:
print(f"Claude multimodal extraction failed: {e}")
logger.error("Claude multimodal extraction failed: %s", e)
return {
"vendorName": "",
"lineItems": [],
@ -307,9 +312,11 @@ def save_extraction(vendor_proposal_id: str, extracted: dict):
timeout=10,
)
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:
print(f"Error saving extraction: {e}")
logger.error("Error saving extraction: %s", e)
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,
)
except Exception as e:
print(f"Error updating status: {e}")
logger.error("Error updating status: %s", e)
def _api_headers() -> dict:
@ -332,3 +339,27 @@ def _api_headers() -> dict:
if api_key:
headers["X-Internal-Api-Key"] = api_key
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 logging
import os
import time
from datetime import datetime
from io import BytesIO
@ -24,6 +26,9 @@ from reportlab.platypus import (
TableStyle,
)
logger = logging.getLogger(__name__)
logger.setLevel(os.environ.get("LOG_LEVEL", "INFO"))
GENERATED_BUCKET = os.environ.get("GENERATED_BUCKET", "")
API_BASE_URL = os.environ.get("API_BASE_URL", "")
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):
proposal = fetch_proposal(proposal_id)
if not proposal:
print(f"Proposal {proposal_id} not found")
logger.warning("Proposal %s not found", proposal_id)
return
line_items = fetch_line_items(proposal_id)
@ -87,7 +92,7 @@ def generate_pdf(proposal_id: str):
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:
@ -100,7 +105,7 @@ def fetch_proposal(proposal_id: str) -> dict | None:
if resp.status_code == 200:
return resp.json()
except Exception as e:
print(f"Error fetching proposal: {e}")
logger.error("Error fetching proposal: %s", e)
return None
@ -114,7 +119,7 @@ def fetch_line_items(proposal_id: str) -> list[dict]:
if resp.status_code == 200:
return resp.json()
except Exception as e:
print(f"Error fetching line items: {e}")
logger.error("Error fetching line items: %s", e)
return []
@ -514,7 +519,7 @@ def upload_pdf(s3_key: str, pdf_bytes: bytes):
ContentType="application/pdf",
)
except Exception as e:
print(f"Error uploading PDF: {e}")
logger.error("Error uploading PDF: %s", e)
raise
@ -527,9 +532,9 @@ def register_pdf(proposal_id: str, s3_key: str):
timeout=10,
)
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:
print(f"Error registering PDF: {e}")
logger.error("Error registering PDF: %s", e)
def _api_headers() -> dict:
@ -538,3 +543,27 @@ def _api_headers() -> dict:
if api_key:
headers["X-Internal-Api-Key"] = api_key
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 logging
import os
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", "")
@ -46,7 +51,7 @@ def handler(event, context):
def process_suggestion(proposal_id: str, trigger: str):
proposal = fetch_proposal(proposal_id)
if not proposal:
print(f"Proposal {proposal_id} not found")
logger.warning("Proposal %s not found", proposal_id)
return
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:
return resp.json()
except Exception as e:
print(f"Error fetching line items: {e}")
logger.error("Error fetching line items: %s", e)
return []
@ -90,13 +95,13 @@ def fetch_proposal(proposal_id: str) -> dict | None:
if resp.status_code == 200:
return resp.json()
except Exception as e:
print(f"Error fetching proposal: {e}")
logger.error("Error fetching proposal: %s", e)
return None
def retrieve_similar(scope: str, category: str) -> list[dict]:
if not KNOWLEDGE_BASE_ID:
print("No Knowledge Base configured, skipping retrieval")
logger.info("No Knowledge Base configured, skipping retrieval")
return []
try:
@ -142,7 +147,7 @@ def retrieve_similar(scope: str, category: str) -> list[dict]:
return results
except Exception as e:
print(f"Error retrieving from KB: {e}")
logger.error("Error retrieving from KB: %s", e)
return []
@ -210,7 +215,7 @@ Respond ONLY with the JSON array, no additional text."""
return line_items if isinstance(line_items, list) else []
except Exception as e:
print(f"Error generating line items: {e}")
logger.error("Error generating line items: %s", e)
return []
@ -266,9 +271,11 @@ def post_line_items(proposal_id: str, items: list[dict], existing_items: list[di
timeout=15,
)
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:
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]):
@ -292,7 +299,7 @@ def store_similar_references(proposal_id: str, similar_proposals: list[dict]):
timeout=10,
)
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):
@ -304,7 +311,7 @@ def update_status_to_in_review(proposal_id: str):
timeout=10,
)
except Exception as e:
print(f"Error updating status: {e}")
logger.error("Error updating status: %s", e)
def _api_headers() -> dict:
@ -313,3 +320,27 @@ def _api_headers() -> dict:
if api_key:
headers["X-Internal-Api-Key"] = api_key
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]