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