Files

309 lines
11 KiB
Python

from __future__ import annotations
import hashlib
import uuid
from dataclasses import dataclass
from datetime import UTC, datetime, timedelta
from typing import Any, cast
import structlog
from sqlalchemy import func, select, text
from sqlalchemy.dialects.postgresql import insert as pg_insert
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession
from app.db import (
Channel,
DeliveryStatus,
SendStatus,
SmsCallbackEvent,
SmsOutboundMessage,
SmsSetting,
SmsTemplate,
)
from app.domain import (
DomainError,
delivery_transition,
destination_hmac,
normalize_phone,
render_template,
request_fingerprint,
)
from app.schemas import CallbackItem, MessageResponse, SendRequest, SendResponse
log = structlog.get_logger()
@dataclass(frozen=True)
class RuntimeSettings:
default_sender_name: str
connect_timeout_ms: int
request_timeout_ms: int
callback_enabled: bool
poll_interval_ms: int
lease_seconds: int
SETTING_RULES: dict[str, tuple[type, int | None, int | None]] = {
"provider.idgtl.default_sender_name": (str, 1, 64),
"provider.idgtl.connect_timeout_ms": (int, 100, 30_000),
"provider.idgtl.request_timeout_ms": (int, 1_000, 120_000),
"provider.idgtl.callback_enabled": (bool, None, None),
"worker.poll_interval_ms": (int, 100, 60_000),
"worker.lease_seconds": (int, 10, 600),
}
async def load_runtime_settings(db: AsyncSession) -> RuntimeSettings:
rows = (
await db.execute(select(SmsSetting).where(SmsSetting.setting_key.in_(SETTING_RULES)))
).scalars()
values = {row.setting_key: row.setting_value for row in rows}
if values.keys() != SETTING_RULES.keys():
raise RuntimeError("required sms settings are missing")
for key, (expected_type, minimum, maximum) in SETTING_RULES.items():
value = values[key]
if type(value) is not expected_type: # bool is an int subclass
raise RuntimeError(f"invalid sms setting type: {key}")
if isinstance(value, (int, str)):
size = value if isinstance(value, int) else len(value)
if minimum is not None and size < minimum:
raise RuntimeError(f"sms setting below minimum: {key}")
if maximum is not None and size > maximum:
raise RuntimeError(f"sms setting above maximum: {key}")
sender = str(values["provider.idgtl.default_sender_name"])
if sender.startswith("__"):
raise RuntimeError("provider sender name is not configured")
return RuntimeSettings(
default_sender_name=sender,
connect_timeout_ms=cast(int, values["provider.idgtl.connect_timeout_ms"]),
request_timeout_ms=cast(int, values["provider.idgtl.request_timeout_ms"]),
callback_enabled=cast(bool, values["provider.idgtl.callback_enabled"]),
poll_interval_ms=cast(int, values["worker.poll_interval_ms"]),
lease_seconds=cast(int, values["worker.lease_seconds"]),
)
def fingerprint_payload(body: SendRequest, phone_e164: str) -> dict[str, Any]:
return {
"idempotency_key": body.idempotency_key,
"template_code": body.template_code,
"locale": body.locale,
"phone_e164": phone_e164,
"substitutions": body.substitutions,
"customer_ref": body.customer_ref,
"message_ttl_sec": body.message_ttl_sec,
}
def validate_otp_request(body: SendRequest) -> None:
code = body.substitutions.get("code")
ttl_min = body.substitutions.get("ttl_min")
if (
not isinstance(code, str)
or not code.isascii()
or not code.isdigit()
or not 4 <= len(code) <= 10
or body.message_ttl_sec % 60 != 0
or str(ttl_min) != str(body.message_ttl_sec // 60)
):
raise DomainError(
"sms_request_invalid", 422, "OTP substitutions and message TTL are inconsistent"
)
def send_response(message: SmsOutboundMessage) -> SendResponse:
return SendResponse(sms_message_id=message.id, ordered_at=message.requested_at)
async def existing_order(
db: AsyncSession, idempotency_key: str, fingerprint: str
) -> SmsOutboundMessage | None:
message = await db.scalar(
select(SmsOutboundMessage).where(
SmsOutboundMessage.requester_service == "keycloak",
SmsOutboundMessage.idempotency_key == idempotency_key,
)
)
if message and message.request_fingerprint != fingerprint:
raise DomainError("idempotency_key_reused", 409, "Idempotency key was reused")
return message
async def enforce_rate_limit(db: AsyncSession, phone_e164: str, destination_key: bytes) -> None:
digest = destination_hmac(phone_e164, destination_key)
lock_id = int.from_bytes(bytes.fromhex(digest[:16]), byteorder="big", signed=True)
await db.execute(text("SELECT pg_advisory_xact_lock(:key)"), {"key": lock_id})
since = datetime.now(UTC) - timedelta(minutes=10)
count = await db.scalar(
select(func.count(SmsOutboundMessage.id)).where(
SmsOutboundMessage.requester_service == "keycloak",
SmsOutboundMessage.phone_e164 == phone_e164,
SmsOutboundMessage.created_at >= since,
)
)
if (count or 0) >= 5:
raise DomainError(
"rate_limit_exceeded",
429,
"Rate limit exceeded",
{"retry_after": 600},
)
async def create_order(
db: AsyncSession,
body: SendRequest,
request_id: str | None,
traceparent: str | None,
destination_key: bytes,
) -> tuple[SendResponse, bool]:
validate_otp_request(body)
phone_e164, phone_digits, phone_masked = normalize_phone(body.phone_e164)
fingerprint = request_fingerprint(fingerprint_payload(body, phone_e164))
existing = await existing_order(db, body.idempotency_key, fingerprint)
if existing:
return send_response(existing), False
runtime = await load_runtime_settings(db)
template = await db.scalar(
select(SmsTemplate).where(
SmsTemplate.code == body.template_code,
SmsTemplate.channel == Channel.SMS,
SmsTemplate.locale == body.locale,
SmsTemplate.is_active.is_(True),
SmsTemplate.approved_at.is_not(None),
)
)
if not template:
raise DomainError("sms_service_unavailable", 503, "SMS service is unavailable")
sender = template.sender_name or runtime.default_sender_name
rendered = render_template(
template.body_template, template.placeholders, body.substitutions, template.max_parts
)
await enforce_rate_limit(db, phone_e164, destination_key)
now = datetime.now(UTC)
message = SmsOutboundMessage(
id=uuid.uuid4(),
requested_at=now,
updated_at=now,
requester_service="keycloak",
process="auth_otp",
channel="SMS",
provider="idgtl",
phone_e164=phone_e164,
phone_digits=phone_digits,
phone_masked=phone_masked,
template_id=template.id,
template_code=template.code,
body_rendered=rendered,
substitutions=body.substitutions,
send_status=SendStatus.PENDING,
delivery_status=DeliveryStatus.UNKNOWN,
customer_ref=body.customer_ref,
idempotency_key=body.idempotency_key,
request_fingerprint=fingerprint,
request_id=request_id,
traceparent=traceparent,
sender_name=sender,
message_ttl_sec=body.message_ttl_sec,
attempt_count=0,
next_attempt_at=now,
)
db.add(message)
try:
await db.commit()
except IntegrityError:
await db.rollback()
concurrent = await existing_order(db, body.idempotency_key, fingerprint)
if concurrent:
return send_response(concurrent), False
raise
return send_response(message), True
def message_response(message: SmsOutboundMessage) -> MessageResponse:
return MessageResponse(
sms_message_id=message.id,
ordered_at=message.requested_at,
updated_at=message.updated_at,
requester_service=message.requester_service,
process=message.process,
channel=message.channel,
provider=message.provider,
phone_masked=message.phone_masked,
template_code=message.template_code,
customer_ref=message.customer_ref,
send_status=message.send_status.value,
delivery_status=message.delivery_status.value,
provider_message_id=message.provider_message_id,
accepted_at=message.accepted_at,
sent_at=message.sent_at,
delivered_at=message.delivered_at,
attempt_count=message.attempt_count,
provider_error_code=message.provider_error_code,
)
async def read_message(db: AsyncSession, message_id: uuid.UUID) -> MessageResponse:
message = await db.scalar(
select(SmsOutboundMessage).where(
SmsOutboundMessage.id == message_id,
SmsOutboundMessage.requester_service == "keycloak",
)
)
if not message:
raise DomainError("not_found", 404, "Resource was not found")
return message_response(message)
async def apply_callback(db: AsyncSession, item: CallbackItem) -> bool:
if item.channel_type.upper() != "SMS":
log.warning("callback.rejected", reason="wrong_channel")
return False
message = await db.scalar(
select(SmsOutboundMessage)
.where(
SmsOutboundMessage.provider == "idgtl",
SmsOutboundMessage.provider_message_id == item.message_uuid,
)
.with_for_update()
)
if not message or item.external_message_id != str(message.id):
digest = hashlib.sha256(item.message_uuid.encode()).hexdigest()[:16]
log.warning(
"callback.rejected", reason="unknown_or_conflicting_message", message_hash=digest
)
return False
target = delivery_transition(message.delivery_status, item.status)
if target is None:
log.warning("callback.rejected", reason="unknown_status", sms_message_id=str(message.id))
return False
inserted = await db.scalar(
pg_insert(SmsCallbackEvent)
.values(
id=uuid.uuid4(),
message_uuid=item.message_uuid,
callback_event=item.callback_event.lower(),
status=item.status.lower(),
status_time=item.status_time,
)
.on_conflict_do_nothing(constraint="uq_callback_event")
.returning(SmsCallbackEvent.id)
)
if inserted is None:
return True
now = datetime.now(UTC)
message.delivery_status = target
message.callback_last_at = now
message.updated_at = now
message.provider_error_code = item.error_code
message.parts = item.parts if item.parts is not None else message.parts
message.price = item.price if item.price is not None else message.price
message.currency = item.currency if item.currency is not None else message.currency
if target == DeliveryStatus.SENT and message.sent_at is None:
message.sent_at = item.status_time
elif target == DeliveryStatus.DELIVERED and message.delivered_at is None:
message.delivered_at = item.status_time
return True