Files
han-app/VM1_app/codebase/backend/sms-service/app/worker.py
T

253 lines
9.8 KiB
Python

from __future__ import annotations
import asyncio
import logging
import random
import signal
import time
from datetime import UTC, datetime, timedelta
import httpx
import structlog
from opentelemetry import trace
from prometheus_client import start_http_server
from sqlalchemy import and_, func, or_, select, update
from app.db import Database, SendStatus, SmsOutboundMessage
from app.logging_security import redact_event
from app.metrics import (
JOURNAL_ROWS,
PENDING_AGE,
PROVIDER_LATENCY,
SEND_TOTAL,
SETTINGS_VALID,
UNCERTAIN_TOTAL,
)
from app.provider import IdgtlClient, IdgtlConfig
from app.service import RuntimeSettings, load_runtime_settings
from app.settings import Settings, get_settings
from app.telemetry import add_trace_context, init_telemetry
log = structlog.get_logger()
MAX_CONNECT_ATTEMPTS = 3
def configure_logging(level: str) -> None:
logging.basicConfig(level=level, format="%(message)s")
structlog.configure(
processors=[
structlog.contextvars.merge_contextvars,
add_trace_context,
redact_event,
structlog.processors.TimeStamper(fmt="iso", utc=True, key="timestamp"),
structlog.stdlib.add_log_level,
structlog.processors.JSONRenderer(),
]
)
async def reconcile_expired_leases(db: Database) -> int:
now = datetime.now(UTC)
async with db.sessions.begin() as session:
result = await session.execute(
update(SmsOutboundMessage)
.where(
SmsOutboundMessage.send_status == SendStatus.PENDING,
SmsOutboundMessage.attempt_count > 0,
SmsOutboundMessage.worker_locked_until < now,
)
.values(
send_status=SendStatus.UNCERTAIN,
worker_locked_until=None,
next_attempt_at=None,
updated_at=now,
provider_error_code="worker_lease_expired",
provider_error_message="provider_result_uncertain",
)
.returning(SmsOutboundMessage.id)
)
ids = list(result.scalars())
for message_id in ids:
SEND_TOTAL.labels("idgtl", SendStatus.UNCERTAIN.value).inc()
UNCERTAIN_TOTAL.labels("idgtl").inc()
log.error("worker.lease_expired", sms_message_id=str(message_id))
return len(ids)
async def lease_message(db: Database, runtime: RuntimeSettings) -> SmsOutboundMessage | None:
now = datetime.now(UTC)
eligible = or_(
and_(
SmsOutboundMessage.send_status == SendStatus.PENDING,
SmsOutboundMessage.attempt_count == 0,
),
and_(
SmsOutboundMessage.send_status == SendStatus.FAILED,
SmsOutboundMessage.attempt_count < MAX_CONNECT_ATTEMPTS,
),
)
async with db.sessions.begin() as session:
message = await session.scalar(
select(SmsOutboundMessage)
.where(
eligible,
SmsOutboundMessage.next_attempt_at <= now,
or_(
SmsOutboundMessage.worker_locked_until.is_(None),
SmsOutboundMessage.worker_locked_until < now,
),
)
.order_by(SmsOutboundMessage.next_attempt_at, SmsOutboundMessage.created_at)
.with_for_update(skip_locked=True)
.limit(1)
)
if message:
message.send_status = SendStatus.PENDING
message.attempt_count += 1
message.last_attempt_at = now
message.worker_locked_until = now + timedelta(seconds=runtime.lease_seconds)
message.updated_at = now
return message
async def save_result(db: Database, message_id, result, attempt_count: int) -> None:
now = datetime.now(UTC)
status = result.send_status
next_attempt = None
if result.retry_safe and attempt_count < MAX_CONNECT_ATTEMPTS:
next_attempt = now + timedelta(seconds=(2**attempt_count) + random.uniform(0, 1)) # noqa: S311
async with db.sessions.begin() as session:
values = {
"send_status": status,
"provider_http_status": result.http_status,
"provider_message_id": result.message_uuid,
"provider_external_id": result.external_id,
"provider_error_code": result.error_code,
"provider_error_message": result.error_message,
"worker_locked_until": None,
"next_attempt_at": next_attempt,
"updated_at": now,
}
if status == SendStatus.ACCEPTED:
values["accepted_at"] = now
await session.execute(
update(SmsOutboundMessage)
.where(
SmsOutboundMessage.id == message_id,
SmsOutboundMessage.send_status == SendStatus.PENDING,
SmsOutboundMessage.attempt_count == attempt_count,
)
.values(**values)
)
SEND_TOTAL.labels("idgtl", status.value).inc()
if status == SendStatus.UNCERTAIN:
UNCERTAIN_TOTAL.labels("idgtl").inc()
if result.contract_violation:
log.error("provider.contract_violation", sms_message_id=str(message_id))
def provider_config(settings: Settings, runtime: RuntimeSettings) -> IdgtlConfig:
if settings.idgtl_api_key is None:
raise RuntimeError("IDGTL_SMS_API_KEY is required by sms-worker")
return IdgtlConfig(
base_url=str(settings.idgtl_base_url),
api_key=settings.idgtl_api_key.get_secret_value(),
callback_url=str(settings.callback_public_url),
callback_username=settings.callback_username.get_secret_value(),
callback_password=settings.callback_password.get_secret_value(),
connect_timeout_ms=runtime.connect_timeout_ms,
request_timeout_ms=runtime.request_timeout_ms,
callback_enabled=runtime.callback_enabled,
)
async def update_queue_metrics(db: Database) -> None:
async with db.sessions() as session:
oldest = await session.scalar(
select(func.min(SmsOutboundMessage.created_at)).where(
SmsOutboundMessage.send_status == SendStatus.PENDING
)
)
count = await session.scalar(select(func.count(SmsOutboundMessage.id)))
age = max(0.0, (datetime.now(UTC) - oldest).total_seconds()) if oldest else 0.0
PENDING_AGE.set(age)
JOURNAL_ROWS.set(count or 0)
async def worker_loop(stop: asyncio.Event) -> None:
settings = get_settings()
telemetry = init_telemetry("sms-worker")
configure_logging(settings.log_level)
structlog.contextvars.bind_contextvars(**{"service.name": "sms-worker"})
metrics_server, metrics_thread = start_http_server(
settings.metrics_port,
addr="0.0.0.0", # noqa: S104 - internal Docker-network listener
)
tracer = trace.get_tracer("han.sms.worker")
db = Database(settings.database_url)
async with httpx.AsyncClient() as http:
try:
while not stop.is_set():
try:
await reconcile_expired_leases(db)
async with db.sessions() as session:
runtime = await load_runtime_settings(session)
SETTINGS_VALID.set(1)
claim_started_ns = time.time_ns()
message = await lease_message(db, runtime)
claim_finished_ns = time.time_ns()
if message is None:
await update_queue_metrics(db)
await asyncio.wait_for(stop.wait(), timeout=runtime.poll_interval_ms / 1000)
continue
process_span = tracer.start_span("sms.process", start_time=claim_started_ns)
with trace.use_span(process_span, end_on_exit=True):
claim_span = tracer.start_span("sms.claim", start_time=claim_started_ns)
claim_span.set_attribute("sms.claimed", True)
claim_span.end(end_time=claim_finished_ns)
span = trace.get_current_span()
span.set_attribute("messaging.operation.name", "send")
span.set_attribute("messaging.system", "idgtl")
span.set_attribute("sms.attempt", message.attempt_count)
client = IdgtlClient(http, provider_config(settings, runtime))
with tracer.start_as_current_span("sms.provider"):
started = time.monotonic()
result = await client.send(message)
PROVIDER_LATENCY.labels("idgtl").observe(time.monotonic() - started)
span.set_attribute("sms.outcome", result.send_status.value)
with tracer.start_as_current_span("sms.save_result"):
await save_result(db, message.id, result, message.attempt_count)
except TimeoutError:
continue
except Exception as exc:
SETTINGS_VALID.set(0)
log.error(
"worker.iteration_failed",
error_type=type(exc).__name__,
)
try:
await asyncio.wait_for(stop.wait(), timeout=5)
except TimeoutError:
pass
finally:
await db.close()
metrics_server.shutdown()
metrics_thread.join(timeout=5)
if telemetry:
telemetry.shutdown()
def run() -> None:
stop = asyncio.Event()
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
for name in (signal.SIGINT, signal.SIGTERM):
try:
loop.add_signal_handler(name, stop.set)
except NotImplementedError:
pass
try:
loop.run_until_complete(worker_loop(stop))
finally:
loop.close()