Move KYC polling to locked worker
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
27
src/main.py
27
src/main.py
@@ -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
47
src/worker.py
Normal 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())
|
||||
Reference in New Issue
Block a user