Move KYC polling to locked worker

This commit is contained in:
Codex
2026-07-02 13:13:11 +03:00
parent aa01e67a09
commit f97492ce63
9 changed files with 206 additions and 142 deletions

View File

@@ -1,7 +1,10 @@
from __future__ import annotations
import asyncio
from datetime import datetime,timezone,timedelta
from typing import Any
from sqlalchemy import text
from src.application.abstractions import IUnitOfWork
from src.application.contracts import IBeorgService,ILogger
from src.application.domain.dto import BeorgKycCreateResponse,BeorgKycResultResponse,KycPersonalData,KycSessionResponse
@@ -165,6 +168,7 @@ class GetKycSessionCommand:
class PollKycSessionsCommand:
_ADVISORY_LOCK_ID = 50291076
def __init__(
self,
@@ -173,30 +177,59 @@ class PollKycSessionsCommand:
logger: ILogger,
beorg_service: IBeorgService,
batch_size: int,
concurrency: int = 4,
) -> None:
self._unit_of_work = unit_of_work
self._logger = logger
self._beorg_service = beorg_service
self._batch_size = batch_size
self._concurrency = max(1,concurrency)
async def __call__(self) -> None:
session_factory = getattr(self._unit_of_work,'session_factory',None)
if session_factory is None:
await self._poll_batch()
return
async with session_factory() as lock_session:
locked = await lock_session.scalar(
text('select pg_try_advisory_lock(:lock_id)'),
{'lock_id': self._ADVISORY_LOCK_ID},
)
if not locked:
self._logger.debug('KYC polling skipped: another worker holds advisory lock')
return
try:
await self._poll_batch()
finally:
await lock_session.execute(
text('select pg_advisory_unlock(:lock_id)'),
{'lock_id': self._ADVISORY_LOCK_ID},
)
async def _poll_batch(self) -> None:
now = _utc_now()
async with self._unit_of_work as unit_of_work:
await unit_of_work.kyc_repository.expire_all_started_sessions(now=now)
sessions = await unit_of_work.kyc_repository.get_started_sessions(now=now,limit=self._batch_size)
for session in sessions:
if not session.user_id or not session.user_token:
continue
semaphore = asyncio.Semaphore(self._concurrency)
tasks = [self._poll_session_safe(session,semaphore) for session in sessions if session.user_id and session.user_token]
if tasks:
await asyncio.gather(*tasks)
async def _poll_session_safe(self,session: KycEntity,semaphore: asyncio.Semaphore) -> None:
async with semaphore:
try:
await self._poll_session(user_id=session.user_id,user_token=session.user_token)
await self._poll_session(user_id=str(session.user_id),user_token=str(session.user_token))
except ApplicationException as exc:
if exc.status_code in (400,403,422):
async with self._unit_of_work as unit_of_work:
await unit_of_work.kyc_repository.update_session_result(
user_id=session.user_id,
user_token=session.user_token,
user_id=str(session.user_id),
user_token=str(session.user_token),
status='failed',
done_state=True,
set_id=None,

View File

@@ -39,6 +39,7 @@ class Settings(BaseSettings):
DATABASE_ECHO: bool = False
KYC_POLL_SECONDS: int = 10
KYC_POLL_BATCH_SIZE: int = 20
KYC_POLL_CONCURRENCY: int = 4
KYC_SESSION_TTL_SECONDS: int = 900
EXCLUDED_PATHS: tuple[str,...] = ('/docs','/redoc','/openapi.json','/ping')
BEORG_TIMEOUT: int = 120

View File

@@ -1,52 +1,40 @@
from __future__ import annotations
import asyncio
import os
import time
from collections import defaultdict
from collections.abc import Awaitable, Callable
from datetime import timedelta
from typing import Iterable
from typing import Any
from fastapi import Request, Response
from prometheus_client import CollectorRegistry, Counter, Histogram, CONTENT_TYPE_LATEST, generate_latest, multiprocess
from sqlalchemy import func, select
from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint
from starlette.responses import Response as StarletteResponse
from src.infrastructure.database.context import async_session_maker
from src.infrastructure.database.models.kyc import KycModel
SERVICE_NAME = "kyc"
WINDOWS: tuple[tuple[str, timedelta | None], ...] = (
("all", None),
("30d", timedelta(days=30)),
("7d", timedelta(days=7)),
("24h", timedelta(days=1)),
("1h", timedelta(hours=1)),
)
BUSINESS_CACHE_TTL_SECONDS = 30
LATENCY_BUCKETS = (0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0)
WINDOWS = (("all", None), ("30d", timedelta(days=30)), ("7d", timedelta(days=7)), ("24h", timedelta(days=1)), ("1h", timedelta(hours=1)))
_http_requests: dict[tuple[str, str, str], int] = defaultdict(int)
_http_latency_sum: dict[tuple[str, str], float] = defaultdict(float)
_http_latency_count: dict[tuple[str, str], int] = defaultdict(int)
_http_latency_buckets: dict[tuple[str, str, float], int] = defaultdict(int)
def _labels(labels: dict[str, str]) -> str:
return ",".join(f'{key}="{value}"' for key, value in labels.items())
def _sample(name: str, value: int | float, labels: dict[str, str] | None = None) -> str:
return f"{name}{{{_labels(labels)}}} {value}" if labels else f"{name} {value}"
HTTP_REQUESTS = Counter("kyc_http_requests_total", "Total HTTP requests handled by kyc service.", ("service", "method", "route", "status_code"))
HTTP_LATENCY = Histogram("kyc_http_request_duration_seconds", "HTTP request duration in seconds for kyc service.", ("service", "method", "route"), buckets=LATENCY_BUCKETS)
def _route_template(request: Request) -> str:
route = request.scope.get("route")
return str(getattr(route, "path", None) or request.url.path)
path = getattr(route, "path", None)
return str(path) if path else "__unmatched__"
class PrometheusMetricsMiddleware(BaseHTTPMiddleware):
async def dispatch(self, request: Request, call_next: RequestResponseEndpoint) -> StarletteResponse:
if request.url.path == "/metrics":
return await call_next(request)
started = time.perf_counter()
status_code = "500"
try:
@@ -56,91 +44,52 @@ class PrometheusMetricsMiddleware(BaseHTTPMiddleware):
finally:
elapsed = time.perf_counter() - started
route = _route_template(request)
method = request.method
_http_requests[(method, route, status_code)] += 1
_http_latency_sum[(method, route)] += elapsed
_http_latency_count[(method, route)] += 1
for bucket in LATENCY_BUCKETS:
if elapsed <= bucket:
_http_latency_buckets[(method, route, bucket)] += 1
HTTP_REQUESTS.labels(SERVICE_NAME, request.method, route, status_code).inc()
HTTP_LATENCY.labels(SERVICE_NAME, request.method, route).observe(elapsed)
async def _count(statement) -> int:
async with async_session_maker() as session:
result = await session.execute(statement)
return int(result.scalar_one() or 0)
class CachedBusinessMetrics:
def __init__(self, ttl_seconds: int, collector: Callable[[], Awaitable[list[str]]]) -> None:
self._ttl_seconds = ttl_seconds
self._collector = collector
self._expires_at = 0.0
self._lines: list[str] = []
self._lock = asyncio.Lock()
async def get(self) -> list[str]:
now = time.monotonic()
if now < self._expires_at:
return self._lines
async with self._lock:
now = time.monotonic()
if now < self._expires_at:
return self._lines
self._lines = await self._collector()
self._expires_at = now + self._ttl_seconds
return self._lines
async def _rows(statement):
async with async_session_maker() as session:
result = await session.execute(statement)
return result.all()
def _labels(labels: dict[str, str]) -> str:
escaped = {key: value.replace("\\", "\\\\").replace('"', '\\"').replace("\n", "\\n") for key, value in labels.items()}
return ",".join(f'{key}="{value}"' for key, value in escaped.items())
def _http_metrics(prefix: str) -> Iterable[str]:
yield f"# TYPE {prefix}_http_requests_total counter"
for (method, route, status_code), value in sorted(_http_requests.items()):
yield _sample(
f"{prefix}_http_requests_total",
value,
{"service": SERVICE_NAME, "method": method, "route": route, "status_code": status_code},
)
yield f"# TYPE {prefix}_http_request_duration_seconds histogram"
for method, route in sorted(_http_latency_count):
for bucket in LATENCY_BUCKETS:
yield _sample(
f"{prefix}_http_request_duration_seconds_bucket",
_http_latency_buckets.get((method, route, bucket), 0),
{"service": SERVICE_NAME, "method": method, "route": route, "le": str(bucket)},
)
yield _sample(
f"{prefix}_http_request_duration_seconds_bucket",
_http_latency_count[(method, route)],
{"service": SERVICE_NAME, "method": method, "route": route, "le": "+Inf"},
)
yield _sample(
f"{prefix}_http_request_duration_seconds_sum",
_http_latency_sum[(method, route)],
{"service": SERVICE_NAME, "method": method, "route": route},
)
yield _sample(
f"{prefix}_http_request_duration_seconds_count",
_http_latency_count[(method, route)],
{"service": SERVICE_NAME, "method": method, "route": route},
)
from src.infrastructure.database.models.kyc import KycModel
def _sample(name: str, value: int | float, labels: dict[str, str] | None = None) -> str:
return f"{name}{{{_labels(labels)}}} {value}" if labels else f"{name} {value}"
async def _business_metrics() -> list[str]:
lines: list[str] = []
total = await _count(select(func.count()).select_from(KycModel))
lines.append(_sample("kyc_requests_total", total, {"service": SERVICE_NAME}))
for status, count in await _rows(select(KycModel.status, func.count()).group_by(KycModel.status)):
status_value = str(status)
lines.append(_sample("kyc_requests_by_status_total", int(count), {"service": SERVICE_NAME, "status": status_value}))
approved = await _count(select(func.count()).select_from(KycModel).where(KycModel.done_state.is_(True)))
failed = await _count(select(func.count()).select_from(KycModel).where(KycModel.done_state.is_(False)))
lines.append(_sample("kyc_approvals_total", approved, {"service": SERVICE_NAME}))
lines.append(_sample("kyc_fails_total", failed, {"service": SERVICE_NAME}))
for window_name, delta in WINDOWS:
conditions = []
completed_conditions = []
if delta is not None:
conditions.append(KycModel.created_at >= func.now() - delta)
completed_conditions.append(KycModel.completed_at >= func.now() - delta)
created = await _count(select(func.count()).select_from(KycModel).where(*conditions))
approvals = await _count(select(func.count()).select_from(KycModel).where(KycModel.done_state.is_(True), *completed_conditions))
fails = await _count(select(func.count()).select_from(KycModel).where(KycModel.done_state.is_(False), *completed_conditions))
lines.append(_sample("kyc_requests_created_total", created, {"service": SERVICE_NAME, "window": window_name}))
lines.append(_sample("kyc_approvals_window_total", approvals, {"service": SERVICE_NAME, "window": window_name}))
lines.append(_sample("kyc_fails_window_total", fails, {"service": SERVICE_NAME, "window": window_name}))
return lines
async def _scalar(session: Any, statement: Any) -> int:
result = await session.execute(statement)
return int(result.scalar_one() or 0)
async def metrics_response() -> Response:
lines = [
async def _rows(session: Any, statement: Any) -> list[Any]:
result = await session.execute(statement)
return list(result.all())
async def _collect_business_metrics() -> list[str]:
lines: list[str] = [
"# TYPE kyc_requests_total gauge",
"# TYPE kyc_requests_by_status_total gauge",
"# TYPE kyc_requests_created_total gauge",
@@ -149,10 +98,48 @@ async def metrics_response() -> Response:
"# TYPE kyc_approvals_window_total gauge",
"# TYPE kyc_fails_window_total gauge",
]
async with async_session_maker() as session:
total = await _scalar(session, select(func.count()).select_from(KycModel))
lines.append(_sample("kyc_requests_total", total, {"service": SERVICE_NAME}))
for status, count in await _rows(session, select(KycModel.status, func.count()).group_by(KycModel.status)):
lines.append(_sample("kyc_requests_by_status_total", int(count), {"service": SERVICE_NAME, "status": str(status)}))
approved = await _scalar(session, select(func.count()).select_from(KycModel).where(KycModel.done_state.is_(True)))
failed = await _scalar(session, select(func.count()).select_from(KycModel).where(KycModel.done_state.is_(False)))
lines.append(_sample("kyc_approvals_total", approved, {"service": SERVICE_NAME}))
lines.append(_sample("kyc_fails_total", failed, {"service": SERVICE_NAME}))
for window_name, delta in WINDOWS:
conditions = []
completed_conditions = []
if delta is not None:
conditions.append(KycModel.created_at >= func.now() - delta)
completed_conditions.append(KycModel.completed_at >= func.now() - delta)
created = await _scalar(session, select(func.count()).select_from(KycModel).where(*conditions))
approvals = await _scalar(session, select(func.count()).select_from(KycModel).where(KycModel.done_state.is_(True), *completed_conditions))
fails = await _scalar(session, select(func.count()).select_from(KycModel).where(KycModel.done_state.is_(False), *completed_conditions))
labels = {"service": SERVICE_NAME, "window": window_name}
lines.append(_sample("kyc_requests_created_total", created, labels))
lines.append(_sample("kyc_approvals_window_total", approvals, labels))
lines.append(_sample("kyc_fails_window_total", fails, labels))
lines.append(_sample("kyc_metrics_scrape_success", 1, {"service": SERVICE_NAME}))
return lines
_business_metrics_cache = CachedBusinessMetrics(BUSINESS_CACHE_TTL_SECONDS, _collect_business_metrics)
def _client_metrics() -> bytes:
if os.environ.get("PROMETHEUS_MULTIPROC_DIR"):
registry = CollectorRegistry()
multiprocess.MultiProcessCollector(registry)
return generate_latest(registry)
return generate_latest()
async def metrics_response() -> Response:
lines: list[str] = []
try:
lines.extend(await _business_metrics())
lines.append(_sample("kyc_metrics_scrape_success", 1, {"service": SERVICE_NAME}))
lines.extend(await _business_metrics_cache.get())
except Exception:
lines.append(_sample("kyc_metrics_scrape_success", 0, {"service": SERVICE_NAME}))
lines.extend(_http_metrics("kyc"))
return Response("\n".join(lines) + "\n", media_type="text/plain; version=0.0.4; charset=utf-8")
body = _client_metrics().decode("utf-8") + "\n".join(lines) + "\n"
return Response(body, media_type=CONTENT_TYPE_LATEST)

View File

@@ -2,20 +2,15 @@ from __future__ import annotations
from contextlib import asynccontextmanager
import secrets
from typing import AsyncGenerator
from apscheduler.schedulers.asyncio import AsyncIOScheduler
from fastapi import Depends,FastAPI
from fastapi.exceptions import RequestValidationError
from fastapi.middleware.cors import CORSMiddleware
from fastapi.openapi.docs import get_redoc_html, get_swagger_ui_html
from fastapi.responses import HTMLResponse
from fastapi.security import HTTPBasic, HTTPBasicCredentials
from src.application.commands import PollKycSessionsCommand
from src.application.domain.enums import LogFormat,LogLevel
from src.application.domain.exceptions import ApplicationException,UnauthorizedException
from src.infrastructure.beorg import BeorgService
from src.infrastructure.config.settings import get_settings
from src.infrastructure.database.context import async_session_maker
from src.infrastructure.database.unit_of_work import UnitOfWork
from src.infrastructure.vault import JwtKeyStore, start_jwt_keys_scheduler
from src.infrastructure.utils import generate_instance_id
from src.infrastructure.logger import logger
@@ -61,34 +56,12 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]:
await jwt_store.refresh()
jwt_scheduler = start_jwt_keys_scheduler(jwt_store, refresh_seconds=settings.JWT_KEYS_REFRESH_SECONDS)
kyc_poll_command = PollKycSessionsCommand(
unit_of_work=UnitOfWork(session_factory=async_session_maker,logger=logger),
logger=logger,
beorg_service=BeorgService(
project_id=settings.BEORG_PROJECT_ID,
machine_uid=settings.BEORG_MACHINE_UID,
token=settings.BEORG_TOKEN,
process_info=settings.BEORG_PROCESS_INFO,
timeout=settings.BEORG_TIMEOUT,
),
batch_size=settings.KYC_POLL_BATCH_SIZE,
)
kyc_scheduler = AsyncIOScheduler()
kyc_scheduler.add_job(
kyc_poll_command.__call__,
'interval',
seconds=settings.KYC_POLL_SECONDS,
max_instances=1,
)
kyc_scheduler.start()
app.state.jwt_key_store = jwt_store
app.state.jwt_keys_scheduler = jwt_scheduler
app.state.kyc_scheduler = kyc_scheduler
try:
yield
finally:
app.state.kyc_scheduler.shutdown(wait=False)
app.state.jwt_keys_scheduler.shutdown(wait=False)
logger.info(f'KYC service instance ended with id {instance_id}')

47
src/worker.py Normal file
View File

@@ -0,0 +1,47 @@
from __future__ import annotations
import asyncio
from src.application.commands import PollKycSessionsCommand
from src.application.domain.enums import LogFormat, LogLevel
from src.infrastructure.beorg import BeorgService
from src.infrastructure.config import settings
from src.infrastructure.database.context import async_session_maker
from src.infrastructure.database.unit_of_work import UnitOfWork
from src.infrastructure.logger import logger
from src.infrastructure.utils import generate_instance_id
def _build_poll_command() -> PollKycSessionsCommand:
return PollKycSessionsCommand(
unit_of_work=UnitOfWork(session_factory=async_session_maker,logger=logger),
logger=logger,
beorg_service=BeorgService(
project_id=settings.BEORG_PROJECT_ID,
machine_uid=settings.BEORG_MACHINE_UID,
token=settings.BEORG_TOKEN,
process_info=settings.BEORG_PROCESS_INFO,
timeout=settings.BEORG_TIMEOUT,
),
batch_size=settings.KYC_POLL_BATCH_SIZE,
concurrency=settings.KYC_POLL_CONCURRENCY,
)
async def main() -> None:
instance_id = generate_instance_id()
logger.set_format(LogFormat(settings.LOG_FORMAT.lower()))
logger.set_min_level(LogLevel[settings.LOG_LEVEL.upper()])
logger.set_instance_id(instance_id)
logger.info(f'KYC poller started with id {instance_id}')
command = _build_poll_command()
try:
while True:
await command()
await asyncio.sleep(settings.KYC_POLL_SECONDS)
finally:
logger.info(f'KYC poller stopped with id {instance_id}')
if __name__ == '__main__':
asyncio.run(main())