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

279 lines
11 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 opentelemetry.propagate import extract
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 TelemetryRuntime, add_trace_context, init_telemetry
log = structlog.get_logger()
MAX_CONNECT_ATTEMPTS = 3
def origin_links(traceparent: str | None) -> list[trace.Link]:
if not traceparent:
return []
context = extract({"traceparent": traceparent})
span_context = trace.get_current_span(context).get_span_context()
return [trace.Link(span_context)] if span_context.is_valid else []
def configure_logging(level: str, telemetry: TelemetryRuntime | None = None) -> None:
logging.basicConfig(level=level, format="%(message)s")
if telemetry and telemetry.logging_handler not in logging.getLogger().handlers:
telemetry.logging_handler.addFilter(
lambda record: not record.name.startswith("opentelemetry")
)
logging.getLogger().addHandler(telemetry.logging_handler)
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(),
],
logger_factory=structlog.stdlib.LoggerFactory(),
wrapper_class=structlog.stdlib.BoundLogger,
cache_logger_on_first_use=True,
)
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, telemetry)
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,
links=origin_links(message.traceparent),
)
with (
trace.use_span(process_span, end_on_exit=True),
structlog.contextvars.bound_contextvars(
sms_message_id=str(message.id)
),
):
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()