Проект разделен на два репозитория
This commit is contained in:
@@ -0,0 +1,308 @@
|
||||
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
|
||||
Reference in New Issue
Block a user