Harden auth, pricing, and reliability in order handlers

Enforce Google auth when configured (reject missing tokens with 403),
return 503 on token verification outages, switch to Decimal with
ROUND_HALF_UP for financial precision, clamp discount bounds 0-100,
use email-based slugs, add 5-min cache TTL with time.monotonic(),
wrap Slack invocation in try/except, add reopen_at timestamp to
closed form status, add reminder dedup guards for dual EST/EDT crons,
escape Slack mrkdwn special characters, and handle empty employee names.
This commit is contained in:
Adam Moussa 2026-05-13 12:58:56 -04:00
parent e70d5fbf18
commit bcce15b9e6
2 changed files with 149 additions and 54 deletions

View file

@ -1,11 +1,18 @@
import json import json
import os import os
from datetime import datetime
from decimal import Decimal from decimal import Decimal
from zoneinfo import ZoneInfo
from shared.db import current_week, get_orders, get_roster, get_summary from shared.db import current_week, get_orders, get_roster, get_summary
from shared.slack import post_channel_message, send_dm from shared.slack import post_channel_message, send_dm
def _escape_mrkdwn(text: str) -> str:
"""Escape Slack mrkdwn special characters."""
return text.replace("&", "&amp;").replace("<", "&lt;").replace(">", "&gt;")
class DecimalEncoder(json.JSONEncoder): class DecimalEncoder(json.JSONEncoder):
def default(self, o): def default(self, o):
if isinstance(o, Decimal): if isinstance(o, Decimal):
@ -21,6 +28,10 @@ def lambda_handler(event, context):
elif event_type == "orders_aggregated": elif event_type == "orders_aggregated":
return handle_orders_aggregated(event) return handle_orders_aggregated(event)
elif event_type == "reminder": elif event_type == "reminder":
# Guard against duplicate triggers from dual EST/EDT schedules
now_et = datetime.now(ZoneInfo("America/New_York"))
if now_et.hour != 10 or now_et.weekday() != 3: # 10am Thursday
return {"status": "skipped", "reason": "outside reminder window"}
return handle_reminder(event) return handle_reminder(event)
elif event_type == "order_confirmed": elif event_type == "order_confirmed":
return handle_order_confirmed(event) return handle_order_confirmed(event)
@ -69,7 +80,7 @@ def handle_orders_aggregated(event):
meal_lines = [] meal_lines = []
for m in summary.get("meals", []): for m in summary.get("meals", []):
meal_lines.append(f"{m['meal']}: *{int(m['quantity'])}*") meal_lines.append(f"{_escape_mrkdwn(m['meal'])}: *{int(m['quantity'])}*")
meal_list = "\n".join(meal_lines) meal_list = "\n".join(meal_lines)
has_subsidy = employee_total < grand_total has_subsidy = employee_total < grand_total
@ -126,7 +137,7 @@ def handle_order_confirmed(event):
for item in items: for item in items:
qty = int(item.get("quantity", 0)) qty = int(item.get("quantity", 0))
price = float(item.get("price", 0)) price = float(item.get("price", 0))
item_lines.append(f"{item['name']} x{qty} — ${price * qty:.2f}") item_lines.append(f"{_escape_mrkdwn(item['name'])} x{qty} — your cost: ${price * qty:.2f}")
item_list = "\n".join(item_lines) item_list = "\n".join(item_lines)
send_dm( send_dm(
@ -142,9 +153,9 @@ def handle_order_confirmed(event):
"text": { "text": {
"type": "mrkdwn", "type": "mrkdwn",
"text": ( "text": (
f"Hey {name.split()[0]}! Your meal order has been submitted.\n\n" f"Hey {name.split()[0] if name.strip() else 'there'}! Your meal order has been submitted.\n\n"
f"{item_list}\n\n" f"{item_list}\n\n"
f"*Total: ${total:.2f}*" f"*Your total: ${total:.2f}* (payroll deduction)"
), ),
}, },
}, },
@ -155,6 +166,11 @@ def handle_order_confirmed(event):
def handle_reminder(event): def handle_reminder(event):
# Guard against duplicate triggers from dual EST/EDT schedules
now_et = datetime.now(ZoneInfo("America/New_York"))
if now_et.hour != 10 or now_et.weekday() != 3: # 10am Thursday
return {"status": "skipped", "reason": "outside reminder window"}
week = event.get("week", current_week()) week = event.get("week", current_week())
form_url = os.environ.get("FORM_URL", "") form_url = os.environ.get("FORM_URL", "")
@ -181,7 +197,7 @@ def handle_reminder(event):
"text": { "text": {
"type": "mrkdwn", "type": "mrkdwn",
"text": ( "text": (
f"Hey {emp['name'].split()[0]}! Meal orders close today at *11:59pm*.\n\n" f"Hey {emp['name'].split()[0] if emp['name'].strip() else 'there'}! Meal orders close today at *11:59pm*.\n\n"
f"*<{form_url}|Place your order>*" f"*<{form_url}|Place your order>*"
), ),
}, },

View file

@ -1,7 +1,12 @@
import json import json
import logging
import os import os
import sys
import time
import urllib.error
import urllib.request import urllib.request
from datetime import datetime from datetime import datetime, timedelta
from decimal import Decimal, ROUND_HALF_UP
from zoneinfo import ZoneInfo from zoneinfo import ZoneInfo
import boto3 import boto3
@ -9,10 +14,19 @@ import boto3
from shared.db import current_week, get_form_status, get_roster, get_settings, put_order from shared.db import current_week, get_form_status, get_roster, get_settings, put_order
from shared.secrets import get_parameter, get_secret from shared.secrets import get_parameter, get_secret
logger = logging.getLogger(__name__)
logger.setLevel(logging.INFO)
if not logger.handlers:
logger.addHandler(logging.StreamHandler(sys.stderr))
EASTERN = ZoneInfo("America/New_York") EASTERN = ZoneInfo("America/New_York")
CACHE_TTL_SECONDS = 300 # 5-minute TTL for cached config values
_api_key = None _api_key = None
_settings = None _settings = None
_settings_ts = 0.0
_google_client_id = None _google_client_id = None
_google_client_id_ts = 0.0
_lambda = boto3.client("lambda") _lambda = boto3.client("lambda")
@ -23,32 +37,48 @@ def _get_api_key() -> str:
return _api_key return _api_key
def _get_discount_settings() -> tuple[float, float]: def _get_discount_settings() -> tuple[Decimal, Decimal]:
global _settings global _settings, _settings_ts
if _settings is None: now = time.monotonic()
if _settings is None or (now - _settings_ts) > CACHE_TTL_SECONDS:
s = get_settings() s = get_settings()
_settings = ( _settings = (
float(s.get("bulk_discount_percent", 0)), Decimal(str(s.get("bulk_discount_percent", 0))),
float(s.get("company_subsidy_percent", 0)), Decimal(str(s.get("company_subsidy_percent", 0))),
) )
_settings_ts = now
return _settings return _settings
def _get_google_client_id() -> str: def _get_google_client_id() -> str:
global _google_client_id global _google_client_id, _google_client_id_ts
if _google_client_id is None: now = time.monotonic()
if _google_client_id is None or (now - _google_client_id_ts) > CACHE_TTL_SECONDS:
param = os.environ.get("GOOGLE_CLIENT_ID_PARAM", "") param = os.environ.get("GOOGLE_CLIENT_ID_PARAM", "")
if param: if param:
_google_client_id = get_parameter(param, decrypt=False) or "" _google_client_id = get_parameter(param, decrypt=False) or ""
else: else:
_google_client_id = "" _google_client_id = ""
_google_client_id_ts = now
return _google_client_id return _google_client_id
def _verify_google_token(token: str) -> dict | None: def _verify_google_token(token: str) -> tuple[dict | None, str]:
"""Verify a Google ID token via the tokeninfo endpoint.
Returns a tuple of (user_info, error_kind) where:
- ({"name": ..., "email": ...}, "ok") on success
- (None, "invalid") for bad/expired tokens or wrong audience/domain
- (None, "unavailable") when the Google verification service is unreachable
NOTE: The token is passed as a query parameter to Google's tokeninfo endpoint.
This is acceptable because ID tokens are short-lived (typically ~1 hour) and
this is Google's own documented verification method, but be aware that the
token will appear in Google's server access logs.
"""
client_id = _get_google_client_id() client_id = _get_google_client_id()
if not client_id: if not client_id:
return None return None, "invalid"
try: try:
req = urllib.request.Request( req = urllib.request.Request(
f"https://oauth2.googleapis.com/tokeninfo?id_token={token}" f"https://oauth2.googleapis.com/tokeninfo?id_token={token}"
@ -56,12 +86,18 @@ def _verify_google_token(token: str) -> dict | None:
with urllib.request.urlopen(req, timeout=5) as resp: with urllib.request.urlopen(req, timeout=5) as resp:
data = json.loads(resp.read()) data = json.loads(resp.read())
if data.get("aud") != client_id: if data.get("aud") != client_id:
return None logger.warning("Google token audience mismatch: got %s", data.get("aud"))
return None, "invalid"
if data.get("hd") != "seahavenind.com": if data.get("hd") != "seahavenind.com":
return None logger.warning("Google token domain mismatch: got %s", data.get("hd"))
return {"name": data.get("name", ""), "email": data.get("email", "")} return None, "invalid"
except Exception: return {"name": data.get("name", ""), "email": data.get("email", "")}, "ok"
return None except (urllib.error.URLError, TimeoutError, OSError) as exc:
logger.error("Google token verification service unavailable: %s", exc)
return None, "unavailable"
except Exception as exc:
logger.error("Google token verification failed (bad token data): %s", exc)
return None, "invalid"
def lambda_handler(event, context): def lambda_handler(event, context):
@ -83,7 +119,19 @@ def lambda_handler(event, context):
def handle_form_status(event): def handle_form_status(event):
week = event.get("pathParameters", {}).get("week", current_week()) week = event.get("pathParameters", {}).get("week", current_week())
status = get_form_status(week) status = get_form_status(week)
return response(200, {"week": week, "status": status}) result = {"week": week, "status": status}
if status == "closed":
# Compute next Monday 8:00 AM Eastern, accounting for DST
now_et = datetime.now(EASTERN)
days_until_monday = (7 - now_et.weekday()) % 7
if days_until_monday == 0 and now_et.hour >= 8:
# If today is Monday past 8am, next Monday is 7 days away
days_until_monday = 7
next_monday = (now_et + timedelta(days=days_until_monday)).replace(
hour=8, minute=0, second=0, microsecond=0
)
result["reopen_at"] = int(next_monday.timestamp())
return response(200, result)
def handle_roster(): def handle_roster():
@ -102,14 +150,33 @@ def handle_submit(event):
except json.JSONDecodeError: except json.JSONDecodeError:
return response(400, {"error": "Invalid JSON"}) return response(400, {"error": "Invalid JSON"})
# --- Authentication ---
# If Google auth is configured, require a valid google_id_token.
# Manual name/email fallback is only allowed when Google auth is NOT configured.
google_auth_enabled = bool(_get_google_client_id())
google_token = body.get("google_id_token") google_token = body.get("google_id_token")
if google_token:
user_info = _verify_google_token(google_token) if google_auth_enabled:
if not user_info: if not google_token:
return response(403, {"error": "Google authentication is required"})
user_info, verify_status = _verify_google_token(google_token)
if verify_status == "unavailable":
return response(503, {"error": "Authentication service temporarily unavailable"})
if user_info is None:
return response(403, {"error": "Invalid or unauthorized Google account"})
name = user_info["name"]
email = user_info["email"]
elif google_token:
# Google auth not configured but token provided — verify it anyway
user_info, verify_status = _verify_google_token(google_token)
if verify_status == "unavailable":
return response(503, {"error": "Authentication service temporarily unavailable"})
if user_info is None:
return response(403, {"error": "Invalid or unauthorized Google account"}) return response(403, {"error": "Invalid or unauthorized Google account"})
name = user_info["name"] name = user_info["name"]
email = user_info["email"] email = user_info["email"]
else: else:
# Manual fallback — only when Google auth is not configured
name = body.get("employee_name", "").strip() name = body.get("employee_name", "").strip()
email = body.get("employee_email", "").strip() email = body.get("employee_email", "").strip()
@ -131,51 +198,63 @@ def handle_submit(event):
filtered_items = [i for i in items if i.get("quantity", 0) > 0] filtered_items = [i for i in items if i.get("quantity", 0) > 0]
# --- Price calculation using Decimal for financial precision ---
# Rounding approach (two-step intermediate rounding):
# 1. bulk_price = retail * bulk_mult, rounded to 2 decimal places
# 2. emp_price = bulk_price * subsidy_mult, rounded to 2 decimal places
# The frontend should match this two-step rounding to avoid discrepancies.
TWO_PLACES = Decimal("0.01")
bulk_pct, subsidy_pct = _get_discount_settings() bulk_pct, subsidy_pct = _get_discount_settings()
bulk_mult = 1 - (bulk_pct / 100) bulk_pct = max(Decimal("0"), min(Decimal("100"), bulk_pct))
subsidy_mult = 1 - (subsidy_pct / 100) subsidy_pct = max(Decimal("0"), min(Decimal("100"), subsidy_pct))
bulk_mult = Decimal("1") - (bulk_pct / Decimal("100"))
subsidy_mult = Decimal("1") - (subsidy_pct / Decimal("100"))
for item in filtered_items: for item in filtered_items:
retail = item.get("retail_price", item.get("price", 0)) or 0 retail = Decimal(str(item.get("retail_price", item.get("price", 0)) or 0))
qty = item.get("quantity", 0) qty = Decimal(str(item.get("quantity", 0)))
bulk_price = round(retail * bulk_mult, 2) # Step 1: apply bulk discount and round
emp_price = round(bulk_price * subsidy_mult, 2) bulk_price = (retail * bulk_mult).quantize(TWO_PLACES, rounding=ROUND_HALF_UP)
item["retail_price"] = retail # Step 2: apply company subsidy and round
item["bulk_price"] = bulk_price emp_price = (bulk_price * subsidy_mult).quantize(TWO_PLACES, rounding=ROUND_HALF_UP)
item["price"] = emp_price subtotal = (emp_price * qty).quantize(TWO_PLACES, rounding=ROUND_HALF_UP)
item["subtotal"] = round(emp_price * qty, 2) # Convert back to float for JSON serialization
item["retail_price"] = float(retail)
item["bulk_price"] = float(bulk_price)
item["price"] = float(emp_price)
item["subtotal"] = float(subtotal)
total = sum(i["subtotal"] for i in filtered_items) total = float(sum(Decimal(str(i["subtotal"])) for i in filtered_items).quantize(TWO_PLACES, rounding=ROUND_HALF_UP))
slug = ( slug = email.split("@")[0].lower().replace(".", "-")
"".join(c if c.isalnum() or c in "- " else "" for c in name)
.strip()
.replace(" ", "-")
.lower()
)
order_data = { order_data = {
"employee_name": name, "employee_name": name,
"employee_email": email, "employee_email": email,
"submitted_at": datetime.now(EASTERN).isoformat(), "submitted_at": datetime.now(EASTERN).isoformat(),
"items": filtered_items, "items": filtered_items,
"total": round(total, 2), "total": total,
} }
put_order(week, slug, order_data) put_order(week, slug, order_data)
_lambda.invoke( # Slack notification is best-effort — order is already persisted above,
FunctionName=os.environ["SLACK_NOTIFIER_ARN"], # so we return success to the user even if this invocation fails.
InvocationType="Event", try:
Payload=json.dumps({ _lambda.invoke(
"event": "order_confirmed", FunctionName=os.environ["SLACK_NOTIFIER_ARN"],
"employee_name": name, InvocationType="Event",
"employee_email": email, Payload=json.dumps({
"items": filtered_items, "event": "order_confirmed",
"total": order_data["total"], "employee_name": name,
"week": week, "employee_email": email,
}, default=float), "items": filtered_items,
) "total": order_data["total"],
"week": week,
}, default=float),
)
except Exception as exc:
logger.error("Slack notifier invocation failed (order already saved): %s", exc)
return response( return response(
200, 200,