feat: advanced RW sync system
This commit is contained in:
46
alembic/versions/426b372286c3_add_rw_sync_outbox.py
Normal file
46
alembic/versions/426b372286c3_add_rw_sync_outbox.py
Normal file
@@ -0,0 +1,46 @@
|
||||
"""add rw sync outbox
|
||||
|
||||
Revision ID: 426b372286c3
|
||||
Revises: af1c3d7e4b20
|
||||
Create Date: 2026-08-20 21:31:20.949847
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = '426b372286c3'
|
||||
down_revision: Union[str, Sequence[str], None] = 'af1c3d7e4b20'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Upgrade schema."""
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.create_table('rw_sync_outbox',
|
||||
sa.Column('id', sa.INTEGER(), autoincrement=True, nullable=False),
|
||||
sa.Column('subscription_id', sa.INTEGER(), nullable=False),
|
||||
sa.Column('revision', sa.INTEGER(), nullable=False),
|
||||
sa.Column('attempts', sa.INTEGER(), nullable=False),
|
||||
sa.Column('next_attempt_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('locked_until', sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(['subscription_id'], ['subscriptions.id'], ),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
sa.UniqueConstraint('subscription_id')
|
||||
)
|
||||
op.create_unique_constraint(None, 'service_notifications', ['id'])
|
||||
# ### end Alembic commands ###
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Downgrade schema."""
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.drop_constraint(None, 'service_notifications', type_='unique')
|
||||
op.drop_table('rw_sync_outbox')
|
||||
# ### end Alembic commands ###
|
||||
@@ -38,6 +38,8 @@ class Settings(BaseSettings):
|
||||
access_token_ttl: int = Field(description="Access token TTL (minutes)")
|
||||
link_code_ttl: int = Field(8, description="Link Code TTL (minutes)")
|
||||
notification_scan_interval: int = Field(5, ge=1)
|
||||
rw_sync_interval_minutes: int = Field(1, ge=1)
|
||||
rw_sync_reconcile_interval_minutes: int = Field(60, ge=1)
|
||||
|
||||
min_password_length: int = Field(8)
|
||||
link_code_length: int = Field(8)
|
||||
|
||||
@@ -31,6 +31,8 @@ async def get_service_identity(
|
||||
request: Request, session: AsyncSession = Depends(get_db)
|
||||
) -> ServiceIdentity | None:
|
||||
auth = request.headers.get("Authorization")
|
||||
if not auth:
|
||||
raise HTTPException(401)
|
||||
|
||||
if auth.startswith("Bearer"):
|
||||
token = auth.removeprefix("Bearer ").strip()
|
||||
|
||||
@@ -3,6 +3,7 @@ from .invoice import Invoice
|
||||
from .link_codes import LinkCode
|
||||
from .orders import Order, OrderAddon
|
||||
from .pricing import PricingConfig
|
||||
from .rw_sync_outbox import RWSyncOutbox
|
||||
from .service_notifications import ServiceNotification
|
||||
from .service_signatures import ServiceSignature
|
||||
from .sessions import Session
|
||||
@@ -19,6 +20,7 @@ __all__ = [
|
||||
"Order",
|
||||
"OrderAddon",
|
||||
"PricingConfig",
|
||||
"RWSyncOutbox",
|
||||
"ServiceNotification",
|
||||
"ServiceSignature",
|
||||
"Session",
|
||||
|
||||
34
db/models/rw_sync_outbox.py
Normal file
34
db/models/rw_sync_outbox.py
Normal file
@@ -0,0 +1,34 @@
|
||||
from datetime import UTC, datetime
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sqlalchemy import INTEGER, DateTime, ForeignKey
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from db.base import Base
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from db.models import Subscription
|
||||
|
||||
|
||||
class RWSyncOutbox(Base):
|
||||
__tablename__ = "rw_sync_outbox"
|
||||
|
||||
id: Mapped[int] = mapped_column(INTEGER, primary_key=True, autoincrement=True)
|
||||
subscription_id: Mapped[int] = mapped_column(
|
||||
ForeignKey("subscriptions.id"), nullable=False, unique=True
|
||||
)
|
||||
revision: Mapped[int] = mapped_column(INTEGER, nullable=False, default=1)
|
||||
attempts: Mapped[int] = mapped_column(INTEGER, nullable=False, default=0)
|
||||
next_attempt_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
locked_until: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC)
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
nullable=False,
|
||||
default=lambda: datetime.now(UTC),
|
||||
onupdate=lambda: datetime.now(UTC),
|
||||
)
|
||||
|
||||
subscription: Mapped["Subscription"] = relationship("Subscription", lazy="selectin")
|
||||
10
external/rw.py
vendored
10
external/rw.py
vendored
@@ -412,7 +412,7 @@ async def delete_hwid(sdk: RemnawaveSDK, user_uuid: str, hwid: str) -> bool:
|
||||
|
||||
|
||||
async def update_expire_at(sdk: RemnawaveSDK, user_uuid: str, expire_at: datetime):
|
||||
dto = UpdateUserRequestDto(uuid=user_uuid, expire_at=expire_at) # type: ignore
|
||||
dto = UpdateUserRequestDto(uuid=UUID(user_uuid), expire_at=expire_at)
|
||||
try:
|
||||
await sdk.users.update_user(body=dto)
|
||||
return True
|
||||
@@ -463,6 +463,11 @@ async def sync_subscription_by_telegram_id(
|
||||
return False
|
||||
|
||||
rw_user = await get_rw_user(sdk, telegram_id=telegram_id, username=username)
|
||||
if expires_at <= datetime.now(UTC):
|
||||
if rw_user is None:
|
||||
return True
|
||||
return await disable_user(sdk, rw_user.uuid)
|
||||
|
||||
if rw_user is None:
|
||||
rw_username = _build_rw_username(username=username)
|
||||
rw_user = await create_user(
|
||||
@@ -483,7 +488,8 @@ async def sync_subscription_by_telegram_id(
|
||||
|
||||
expire_synced = await update_expire_at(sdk, rw_user.uuid, expires_at)
|
||||
devices_synced = await set_hwid_limit(sdk, rw_user.uuid, devices)
|
||||
return expire_synced and devices_synced
|
||||
enabled = await enable_user(sdk, rw_user.uuid)
|
||||
return expire_synced and devices_synced and enabled
|
||||
|
||||
|
||||
def build_subscription_link(short_uuid: str):
|
||||
|
||||
15
main.py
15
main.py
@@ -5,17 +5,24 @@ from fastapi import FastAPI
|
||||
|
||||
from routes import routers
|
||||
from services.notifications import run_subscription_notifications
|
||||
from services.rw_sync import run_rw_sync_reconciler, run_rw_sync_worker
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(_: FastAPI):
|
||||
task = asyncio.create_task(run_subscription_notifications())
|
||||
tasks = [
|
||||
asyncio.create_task(run_subscription_notifications()),
|
||||
asyncio.create_task(run_rw_sync_worker()),
|
||||
asyncio.create_task(run_rw_sync_reconciler()),
|
||||
]
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
task.cancel()
|
||||
with suppress(asyncio.CancelledError):
|
||||
await task
|
||||
for task in tasks:
|
||||
task.cancel()
|
||||
for task in tasks:
|
||||
with suppress(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
|
||||
app = FastAPI(debug=True, lifespan=lifespan)
|
||||
|
||||
@@ -1,14 +1,19 @@
|
||||
from fastapi import APIRouter, Depends
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from core.deps import get_service_identity
|
||||
from db.session import get_db
|
||||
from repositories.service_notifications import (
|
||||
ack_notification,
|
||||
get_pending_notifications,
|
||||
mark_notification_as_dispatched,
|
||||
)
|
||||
from schemas.dto import ServiceIdentity
|
||||
from schemas.notifications import NotificationResponse, UserNotificationData
|
||||
from schemas.notifications import (
|
||||
NotificationAcknowledgeRequest,
|
||||
NotificationResponse,
|
||||
UserNotificationData,
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/renewal")
|
||||
|
||||
@@ -25,6 +30,7 @@ async def get_pending(
|
||||
for notification in notifications:
|
||||
response_users.append(
|
||||
UserNotificationData(
|
||||
notification_id=notification.id,
|
||||
username=notification.subscription.user.username,
|
||||
telegram_id=notification.subscription.user.telegram_id,
|
||||
expires_at=notification.sub_expires_at,
|
||||
@@ -35,3 +41,19 @@ async def get_pending(
|
||||
await mark_notification_as_dispatched(session, notification.id)
|
||||
|
||||
return NotificationResponse(users=response_users, issued_by=ctx.service)
|
||||
|
||||
|
||||
@router.post("/ack")
|
||||
async def acknowledge(
|
||||
req: NotificationAcknowledgeRequest,
|
||||
ctx: ServiceIdentity = Depends(get_service_identity),
|
||||
session: AsyncSession = Depends(get_db),
|
||||
):
|
||||
try:
|
||||
r = await ack_notification(session, req.notification_id)
|
||||
except Exception as e:
|
||||
raise HTTPException(500, detail=str(e)) from None
|
||||
|
||||
if r:
|
||||
return "OK"
|
||||
raise HTTPException(500, detail="No such notification found.")
|
||||
|
||||
@@ -6,6 +6,7 @@ from schemas.enums import NotificationType
|
||||
|
||||
|
||||
class UserNotificationData(BaseModel):
|
||||
notification_id: int = Field()
|
||||
username: str = Field()
|
||||
telegram_id: int | None = Field()
|
||||
expires_at: datetime = Field()
|
||||
@@ -17,3 +18,7 @@ class NotificationResponse(BaseModel):
|
||||
issued_by: str = Field()
|
||||
|
||||
users: list[UserNotificationData] = Field()
|
||||
|
||||
|
||||
class NotificationAcknowledgeRequest(BaseModel):
|
||||
notification_id: int = Field()
|
||||
|
||||
@@ -7,7 +7,6 @@ from datetime import UTC, datetime
|
||||
from config import cfg
|
||||
from db.models.orders import OrderStatus
|
||||
from db.models.transactions import BalanceTransaction, BalanceTxType
|
||||
from external.rw import sync_subscription_by_telegram_id
|
||||
from repositories import AddonsRepository
|
||||
from repositories.invoices import InvoiceRepository
|
||||
from repositories.orders import OrderRepository
|
||||
@@ -15,6 +14,7 @@ from repositories.pricing import PricingRepository
|
||||
from repositories.users import UserRepository
|
||||
from schemas.invoices import InvoiceStatus
|
||||
from services.plans import get_pricing_model
|
||||
from services.rw_sync import enqueue_rw_sync
|
||||
from services.subscriptions import (
|
||||
apply_order_now,
|
||||
deduct_order_balance,
|
||||
@@ -129,12 +129,7 @@ async def process_subscription_purchase( # noqa: PLR0911, PLR0912
|
||||
applied_subscription = await apply_order_now(
|
||||
session, user=user, order=order, pricing=pricing, now=now
|
||||
)
|
||||
await sync_subscription_by_telegram_id(
|
||||
expires_at=applied_subscription.expires_at,
|
||||
devices=applied_subscription.devices,
|
||||
telegram_id=user.telegram_id,
|
||||
username=user.username,
|
||||
)
|
||||
await enqueue_rw_sync(session, applied_subscription.id)
|
||||
else:
|
||||
if subscription is None:
|
||||
logger.critical(
|
||||
|
||||
124
services/rw_sync.py
Normal file
124
services/rw_sync.py
Normal file
@@ -0,0 +1,124 @@
|
||||
import asyncio
|
||||
import logging
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sqlalchemy import delete, select, update
|
||||
from sqlalchemy.dialects.postgresql import insert
|
||||
|
||||
from config import cfg
|
||||
from db.models import RWSyncOutbox, Subscription
|
||||
from db.session import async_session
|
||||
from external.rw import sync_subscription_by_telegram_id
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
LOCK_DURATION = timedelta(minutes=5)
|
||||
MAX_RETRY_DELAY = timedelta(hours=1)
|
||||
|
||||
|
||||
async def enqueue_rw_sync(session: "AsyncSession", subscription_id: int) -> None:
|
||||
now = datetime.now(UTC)
|
||||
stmt = insert(RWSyncOutbox).values(
|
||||
subscription_id=subscription_id,
|
||||
revision=1,
|
||||
next_attempt_at=now,
|
||||
)
|
||||
await session.execute(
|
||||
stmt.on_conflict_do_update(
|
||||
index_elements=[RWSyncOutbox.subscription_id],
|
||||
set_={
|
||||
"revision": RWSyncOutbox.revision + 1,
|
||||
"next_attempt_at": now,
|
||||
"locked_until": None,
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
async def enqueue_all_rw_syncs(session: "AsyncSession") -> int:
|
||||
subscription_ids = await session.scalars(select(Subscription.id))
|
||||
count = 0
|
||||
for subscription_id in subscription_ids:
|
||||
await enqueue_rw_sync(session, subscription_id)
|
||||
count += 1
|
||||
return count
|
||||
|
||||
|
||||
async def process_rw_syncs(session: "AsyncSession", batch_size: int = 50) -> int:
|
||||
now = datetime.now(UTC)
|
||||
jobs = await session.scalars(
|
||||
select(RWSyncOutbox)
|
||||
.where(
|
||||
RWSyncOutbox.next_attempt_at <= now,
|
||||
(RWSyncOutbox.locked_until.is_(None)) | (RWSyncOutbox.locked_until <= now),
|
||||
)
|
||||
.order_by(RWSyncOutbox.next_attempt_at)
|
||||
.with_for_update(skip_locked=True)
|
||||
.limit(batch_size)
|
||||
)
|
||||
jobs = list(jobs)
|
||||
for job in jobs:
|
||||
job.locked_until = now + LOCK_DURATION
|
||||
await session.commit()
|
||||
|
||||
for job in jobs:
|
||||
subscription = job.subscription
|
||||
user = subscription.user
|
||||
synced = await sync_subscription_by_telegram_id(
|
||||
expires_at=subscription.expires_at,
|
||||
devices=subscription.devices,
|
||||
telegram_id=user.telegram_id,
|
||||
username=user.username,
|
||||
)
|
||||
if synced:
|
||||
await session.execute(
|
||||
delete(RWSyncOutbox).where(
|
||||
RWSyncOutbox.id == job.id,
|
||||
RWSyncOutbox.revision == job.revision,
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
continue
|
||||
|
||||
delay = min(timedelta(minutes=2**job.attempts), MAX_RETRY_DELAY)
|
||||
await session.execute(
|
||||
update(RWSyncOutbox)
|
||||
.where(RWSyncOutbox.id == job.id, RWSyncOutbox.revision == job.revision)
|
||||
.values(
|
||||
attempts=RWSyncOutbox.attempts + 1,
|
||||
next_attempt_at=datetime.now(UTC) + delay,
|
||||
locked_until=None,
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
logger.warning("RW sync failed for subscription_id=%s", subscription.id)
|
||||
|
||||
return len(jobs)
|
||||
|
||||
|
||||
async def run_rw_sync_worker() -> None:
|
||||
interval = cfg.rw_sync_interval_minutes * 60
|
||||
while True:
|
||||
try:
|
||||
async with async_session() as session:
|
||||
await process_rw_syncs(session)
|
||||
except Exception:
|
||||
logger.exception("RW sync worker iteration failed")
|
||||
await asyncio.sleep(interval)
|
||||
|
||||
|
||||
async def run_rw_sync_reconciler() -> None:
|
||||
interval = cfg.rw_sync_reconcile_interval_minutes * 60
|
||||
while True:
|
||||
try:
|
||||
async with async_session() as session:
|
||||
count = await enqueue_all_rw_syncs(session)
|
||||
await session.commit()
|
||||
logger.info("Queued %s subscriptions for RW reconciliation", count)
|
||||
except Exception:
|
||||
logger.exception("RW reconciliation iteration failed")
|
||||
await asyncio.sleep(interval)
|
||||
@@ -6,9 +6,9 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from db.models import Subscription, SubscriptionAddon, User
|
||||
from db.models.orders import Order, OrderStatus
|
||||
from db.models.transactions import BalanceTransaction, BalanceTxType
|
||||
from external.rw import sync_subscription_by_telegram_id
|
||||
from schemas.enums import SubscriptionStatus
|
||||
from schemas.plans import PricingPlans
|
||||
from services.rw_sync import enqueue_rw_sync
|
||||
|
||||
|
||||
def calculate_plan_monthly_price(
|
||||
@@ -194,9 +194,4 @@ async def sync_user_subscription(
|
||||
subscription.status = SubscriptionStatus.EXPIRED
|
||||
|
||||
if applied_due_orders:
|
||||
await sync_subscription_by_telegram_id(
|
||||
expires_at=subscription.expires_at,
|
||||
devices=subscription.devices,
|
||||
telegram_id=user.telegram_id,
|
||||
username=user.username,
|
||||
)
|
||||
await enqueue_rw_sync(session, subscription.id)
|
||||
|
||||
Reference in New Issue
Block a user