Compare commits

..

11 Commits

49 changed files with 901 additions and 192 deletions

View 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 ###

View File

@@ -0,0 +1,38 @@
"""enforce addon.is_enabled->NOT NULL
Revision ID: ad537d63c440
Revises: 6d1bae3dc723
Create Date: 2026-08-19 21:14:39.928563
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = 'ad537d63c440'
down_revision: Union[str, Sequence[str], None] = '6d1bae3dc723'
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.alter_column('addons', 'is_enabled',
existing_type=sa.BOOLEAN(),
nullable=False)
op.create_unique_constraint(None, 'link_codes', ['code'])
# ### end Alembic commands ###
def downgrade() -> None:
"""Downgrade schema."""
# ### commands auto generated by Alembic - please adjust! ###
op.drop_constraint(None, 'link_codes', type_='unique')
op.alter_column('addons', 'is_enabled',
existing_type=sa.BOOLEAN(),
nullable=True)
# ### end Alembic commands ###

View File

@@ -0,0 +1,32 @@
"""Remove global service notification expiry uniqueness.
Revision ID: af1c3d7e4b20
Revises: ee6174eeef14
Create Date: 2026-08-20 00:00:00.000000
"""
from typing import Sequence, Union
from alembic import op
revision: str = "af1c3d7e4b20"
down_revision: Union[str, Sequence[str], None] = "ee6174eeef14"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
op.drop_constraint(
"service_notifications_sub_expires_at_key",
"service_notifications",
type_="unique",
)
def downgrade() -> None:
op.create_unique_constraint(
"service_notifications_sub_expires_at_key",
"service_notifications",
["sub_expires_at"],
)

View File

@@ -0,0 +1,47 @@
"""+service_notifications
Revision ID: ee6174eeef14
Revises: ad537d63c440
Create Date: 2026-08-20 12:01:51.073160
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = 'ee6174eeef14'
down_revision: Union[str, Sequence[str], None] = 'ad537d63c440'
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('service_notifications',
sa.Column('id', sa.INTEGER(), autoincrement=True, nullable=False),
sa.Column('subscription_id', sa.INTEGER(), nullable=False),
sa.Column('notify_type', sa.Enum('SEVEN_DAYS', 'THREE_DAYS', 'ONE_DAY', 'EXPIRED', name='notificationtype'), nullable=False),
sa.Column('sub_expires_at', sa.DateTime(timezone=True), nullable=False),
sa.Column('status', sa.Enum('PENDING', 'DISPATCHED', 'SENT', 'FAILED', name='notificationstatus'), nullable=False),
sa.Column('dispatched_at', sa.DateTime(timezone=True), nullable=True),
sa.Column('sent_at', sa.DateTime(timezone=True), nullable=True),
sa.Column('attempts', sa.INTEGER(), nullable=False),
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
sa.ForeignKeyConstraint(['subscription_id'], ['subscriptions.id'], ),
sa.PrimaryKeyConstraint('id'),
sa.UniqueConstraint('id'),
sa.UniqueConstraint('sub_expires_at'),
sa.UniqueConstraint('subscription_id', 'notify_type', 'sub_expires_at', name='uq_subscription_notification')
)
# ### end Alembic commands ###
def downgrade() -> None:
"""Downgrade schema."""
# ### commands auto generated by Alembic - please adjust! ###
op.drop_table('service_notifications')
# ### end Alembic commands ###

4
compose.local.yml Normal file
View File

@@ -0,0 +1,4 @@
services:
postgres:
ports: !override
- "5432:5432"

View File

@@ -5,7 +5,7 @@ from pydantic_settings import BaseSettings, SettingsConfigDict
class Settings(BaseSettings):
model_config = SettingsConfigDict(env_file=".env")
model_config = SettingsConfigDict(env_file=".env", extra="ignore")
### Internal settings ###
postgres_user: str = Field()
@@ -37,6 +37,9 @@ 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)

View File

@@ -1,14 +1,11 @@
from sqlalchemy.ext.asyncio import AsyncSession
from db.models import User
from db.session import UnitOfWork
from repositories.users import UserRepository
from schemas.jwt import ServiceJWTPayload
async def fetch_subject_from_service(
payload: ServiceJWTPayload, session: AsyncSession
) -> User | None:
repo = UserRepository(session)
async def fetch_subject_from_service(payload: ServiceJWTPayload, uow: UnitOfWork) -> User | None:
repo = UserRepository(uow)
if payload.acting_as.startswith("telegram:"):
telegram_id = int(payload.acting_as.split("telegram:")[1])

View File

@@ -2,18 +2,18 @@ from datetime import UTC, datetime
from fastapi import HTTPException
from pydantic import ValidationError
from sqlalchemy.ext.asyncio import AsyncSession
from core.auth.fetch_sub import fetch_subject_from_service
from core.secrets import decode_jwt, decode_user_jwt, get_kid_from_token
from db.session import UnitOfWork
from repositories.service_signatures import get_active_signature_by_kid
from repositories.users import UserRepository
from schemas.dto import AuthContext
from schemas.jwt import ServiceJWTPayload, UserJWTPayload
async def authorize_bot(kid: str, token: str, session: AsyncSession) -> AuthContext:
signature = await get_active_signature_by_kid(session, kid)
async def authorize_bot(kid: str, token: str, uow: UnitOfWork) -> AuthContext:
signature = await get_active_signature_by_kid(uow.session, kid)
if not signature:
raise HTTPException(401, detail="Invalid service signature")
@@ -28,18 +28,18 @@ async def authorize_bot(kid: str, token: str, session: AsyncSession) -> AuthCont
if payload.exp < datetime.now(UTC).timestamp():
raise HTTPException(status_code=401, detail="Access token expired")
subject = await fetch_subject_from_service(payload, session)
subject = await fetch_subject_from_service(payload, uow)
if not subject:
raise HTTPException(status_code=401, detail="User not found")
return AuthContext(subject, auth_method="service", service=kid)
async def authorize(token: str, session: AsyncSession, service: str | None = None) -> AuthContext:
async def authorize(token: str, uow: UnitOfWork, service: str | None = None) -> AuthContext:
kid = get_kid_from_token(token)
if kid:
return await authorize_bot(kid, token, session)
return await authorize_bot(kid, token, uow)
content = decode_user_jwt(token)
@@ -51,7 +51,7 @@ async def authorize(token: str, session: AsyncSession, service: str | None = Non
if payload.exp < datetime.now(UTC).timestamp():
raise HTTPException(status_code=401, detail="Access token expired")
repo = UserRepository(session)
repo = UserRepository(uow)
user = await repo.get_user_by_id(int(payload.sub))
if not user:

View File

@@ -1,10 +1,9 @@
from fastapi import Depends, HTTPException, Request
from sqlalchemy.ext.asyncio import AsyncSession
from config import cfg
from core.auth import jwt
from core.secrets import get_kid_from_token
from db.session import get_db
from db.session import UnitOfWork, get_uow
from external.pally import PallyClient
from repositories.service_signatures import get_active_signature_by_kid
from schemas.dto import AuthContext, ServiceIdentity
@@ -12,7 +11,7 @@ from services.subscriptions import sync_user_subscription
async def get_auth_context(
request: Request, session: AsyncSession = Depends(get_db)
request: Request, uow: UnitOfWork = Depends(get_uow)
) -> AuthContext | None:
auth = request.headers.get("Authorization")
@@ -21,23 +20,25 @@ async def get_auth_context(
if auth.startswith("Bearer"):
token = auth.removeprefix("Bearer ").strip()
ctx = await jwt.authorize(token, session)
await sync_user_subscription(session, user=ctx.user)
await session.commit()
ctx = await jwt.authorize(token, uow)
await sync_user_subscription(uow.session, user=ctx.user)
await uow.commit()
return ctx
async def get_service_identity(
request: Request, session: AsyncSession = Depends(get_db)
request: Request, uow: UnitOfWork = Depends(get_uow)
) -> ServiceIdentity | None:
auth = request.headers.get("Authorization")
if not auth:
raise HTTPException(401)
if auth.startswith("Bearer"):
token = auth.removeprefix("Bearer ").strip()
kid = get_kid_from_token(token)
if not kid:
raise HTTPException(401, detail="No kid provided.")
signature = await get_active_signature_by_kid(session, kid)
signature = await get_active_signature_by_kid(uow.session, kid)
if not signature:
raise HTTPException(401, detail="Invalid signature")
return ServiceIdentity(service=signature.kid)

View File

@@ -37,7 +37,11 @@ def generate_jwt(payload: dict[str, Any]) -> str:
def get_kid_from_token(token: str) -> str | None:
return jwt.get_unverified_header(token).get("kid")
try:
header = jwt.get_unverified_header(token)
return header.get("kid")
except jwt.exceptions.PyJWTError:
return
def decode_jwt(
@@ -45,7 +49,7 @@ def decode_jwt(
) -> dict[str, Any] | None:
try:
return jwt.decode(token, public_key, algo)
except jwt.ExpiredSignatureError:
except jwt.exceptions.PyJWTError:
return

View File

@@ -3,6 +3,8 @@ 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
from .subscription_addons import SubscriptionAddon
@@ -18,6 +20,8 @@ __all__ = [
"Order",
"OrderAddon",
"PricingConfig",
"RWSyncOutbox",
"ServiceNotification",
"ServiceSignature",
"Session",
"Subscription",

View File

@@ -11,4 +11,4 @@ class Addon(Base):
name: Mapped[str] = mapped_column(TEXT, nullable=False, unique=False)
price: Mapped[float] = mapped_column(FLOAT, nullable=False)
free_threshold: Mapped[int] = mapped_column(INTEGER, default=-1)
is_enabled: Mapped[bool] = mapped_column(BOOLEAN, default=False, nullable=True)
is_enabled: Mapped[bool] = mapped_column(BOOLEAN, default=False, nullable=False)

View 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")

View File

@@ -0,0 +1,52 @@
from datetime import UTC, datetime
from typing import TYPE_CHECKING
from sqlalchemy import INTEGER, DateTime, Enum, ForeignKey, UniqueConstraint
from sqlalchemy.orm import Mapped, mapped_column, relationship
from db.base import Base
from schemas.enums import NotificationStatus, NotificationType
if TYPE_CHECKING:
from db.models import Subscription
class ServiceNotification(Base):
__tablename__ = "service_notifications"
id: Mapped[int] = mapped_column(
INTEGER, unique=True, autoincrement=True, primary_key=True, nullable=False
)
subscription_id: Mapped[int] = mapped_column(
ForeignKey("subscriptions.id"), nullable=False, unique=False
)
notify_type: Mapped[NotificationType] = mapped_column(
Enum(NotificationType, name="notificationtype"), nullable=False
)
sub_expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
status: Mapped[NotificationStatus] = mapped_column(
Enum(NotificationStatus, name="notificationstatus"),
nullable=False,
default=NotificationStatus.PENDING,
)
dispatched_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=True)
sent_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=True)
attempts: Mapped[int] = mapped_column(INTEGER, nullable=False, default=0)
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True),
nullable=False,
default=lambda: datetime.now(UTC),
)
__table_args__ = (
UniqueConstraint(
"subscription_id",
"notify_type",
"sub_expires_at",
name="uq_subscription_notification",
),
)
subscription: Mapped["Subscription"] = relationship("Subscription", lazy="selectin")

View File

@@ -37,7 +37,4 @@ class User(Base):
"Subscription", back_populates="user", lazy="selectin"
)
sessions: Mapped[list["Session"]] = relationship(back_populates="user", lazy="selectin")
referal: Mapped["User | None"] = relationship(
"User",
remote_side=[id],
)
referal: Mapped["User | None"] = relationship("User", remote_side=[id], lazy="selectin")

View File

@@ -1,4 +1,4 @@
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from config import cfg
@@ -6,6 +6,29 @@ engine = create_async_engine(cfg.db_url, echo=True)
async_session = async_sessionmaker(bind=engine, expire_on_commit=False)
class UnitOfWork:
def __init__(self, session: AsyncSession) -> None:
self.session = session
async def commit(self) -> None:
await self.session.commit()
async def rollback(self) -> None:
await self.session.rollback()
async def __aenter__(self):
return self
async def __aexit__(self, exc_type, *_):
if exc_type:
await self.rollback()
async def get_db():
async with async_session() as session:
yield session
async def get_uow():
async with async_session() as session, UnitOfWork(session) as uow:
yield uow

View File

@@ -19,8 +19,6 @@ services:
interval: 10s
timeout: 5s
retries: 5
ports:
- "5432:5432"
volumes:
postgres_data:

20
external/rw.py vendored
View File

@@ -76,10 +76,7 @@ def _parse_user(user_dto) -> RWUserInfo:
)
def _build_rw_username(*, telegram_id: int | None, username: str | None) -> str | None:
if telegram_id is not None:
return f"tg_{telegram_id}"
def _build_rw_username(*, username: str | None) -> str | None:
if not username:
return None
@@ -415,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
@@ -439,7 +436,7 @@ async def get_rw_user(
rw_user = await get_user_by_telegram_id(sdk, telegram_id)
if rw_user is None:
rw_username = _build_rw_username(telegram_id=telegram_id, username=username)
rw_username = _build_rw_username(username=username)
if rw_username is None:
logger.warning(
"Cannot build username for telegram_id=%s username=%s",
@@ -450,6 +447,7 @@ async def get_rw_user(
rw_user = await get_user_by_username(sdk, rw_username)
return rw_user
return rw_user
async def sync_subscription_by_telegram_id(
@@ -465,8 +463,13 @@ 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(telegram_id=telegram_id, username=username)
rw_username = _build_rw_username(username=username)
rw_user = await create_user(
sdk=sdk,
username=rw_username,
@@ -485,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):

25
main.py
View File

@@ -1,8 +1,31 @@
import asyncio
from contextlib import asynccontextmanager, suppress
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
app = FastAPI(debug=True)
@asynccontextmanager
async def lifespan(_: FastAPI):
tasks = [
asyncio.create_task(run_subscription_notifications()),
asyncio.create_task(run_rw_sync_worker()),
asyncio.create_task(run_rw_sync_reconciler()),
]
try:
yield
finally:
for task in tasks:
task.cancel()
for task in tasks:
with suppress(asyncio.CancelledError):
await task
app = FastAPI(debug=True, lifespan=lifespan)
for r in routers:
app.include_router(r)

View File

@@ -1,12 +1,12 @@
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models import Addon
from db.session import UnitOfWork
class AddonsRepository:
def __init__(self, session: AsyncSession) -> None:
self.session = session
def __init__(self, uow: UnitOfWork) -> None:
self.session = uow.session
async def get_all(self) -> list[Addon]:
stmt = select(Addon)

View File

@@ -1,13 +1,14 @@
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models.invoice import Invoice
from db.session import UnitOfWork
from schemas.invoices import InvoiceStatus
class InvoiceRepository:
def __init__(self, session: AsyncSession) -> None:
self.session = session
def __init__(self, uow: UnitOfWork) -> None:
self.uow = uow
self.session = uow.session
async def get_by_id(self, id: int) -> Invoice | None:
stmt = select(Invoice).where(Invoice.id == id)
@@ -32,13 +33,9 @@ class InvoiceRepository:
)
self.session.add(obj)
await self.session.commit()
return obj
async def update_status_by_id(self, invoice_id: int, status: InvoiceStatus) -> Invoice | None:
invoice = await self.get_by_id(invoice_id)
invoice.status = status
await self.session.commit()
return invoice

View File

@@ -1,14 +1,14 @@
from datetime import datetime
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models.link_codes import LinkCode
from db.session import UnitOfWork
from schemas.enums import LinkCodeStatus
async def create_link_code(
session: AsyncSession,
uow: UnitOfWork,
*,
code: str,
user_id: int,
@@ -22,13 +22,17 @@ async def create_link_code(
expires_at=expires_at,
)
session.add(link_code)
await session.commit()
uow.session.add(link_code)
return link_code
async def get_link_code_by_code(session: AsyncSession, code: str) -> LinkCode | None:
async def get_link_code_by_code(uow: UnitOfWork, code: str) -> LinkCode | None:
stmt = select(LinkCode).where(LinkCode.code == code)
r = await session.execute(stmt)
r = await uow.session.execute(stmt)
return r.scalar_one_or_none()
async def use_link_code(uow: UnitOfWork, code: LinkCode) -> LinkCode:
code.status = LinkCodeStatus.USED
return code

View File

@@ -1,12 +1,13 @@
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models.orders import Order, OrderAddon, OrderStatus
from db.session import UnitOfWork
class OrderRepository:
def __init__(self, session: AsyncSession) -> None:
self.session = session
def __init__(self, uow: UnitOfWork) -> None:
self.uow = uow
self.session = uow.session
async def create(
self,
@@ -37,7 +38,7 @@ class OrderRepository:
addon = OrderAddon(order_id=order.id, addon_id=addon_id)
self.session.add(addon)
await self.session.commit()
await self.session.refresh(order, attribute_names=["addons"])
return order

View File

@@ -1,12 +1,12 @@
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models.pricing import PricingConfig
from db.session import UnitOfWork
class PricingRepository:
def __init__(self, session: AsyncSession) -> None:
self.session = session
def __init__(self, uow: UnitOfWork) -> None:
self.session = uow.session
async def get(self) -> PricingConfig | None:
stmt = select(PricingConfig).where(PricingConfig.id == 1)

View File

@@ -0,0 +1,61 @@
from sqlalchemy import func, or_, select, text
from db.models.service_notifications import ServiceNotification
from db.session import UnitOfWork
from schemas.enums import NotificationStatus
async def get_notification_by_id(uow: UnitOfWork, n_id: int) -> ServiceNotification | None:
stmt = select(ServiceNotification).where(ServiceNotification.id == n_id)
r = await uow.session.execute(stmt)
return r.scalar_one_or_none()
async def get_pending_notifications(
uow: UnitOfWork, batch_size: int = 50
) -> list[ServiceNotification]:
stmt = (
select(ServiceNotification)
.where(
or_(
ServiceNotification.status == "pending",
(
(ServiceNotification.status == "dispatched")
& (
ServiceNotification.dispatched_at
< func.now() - text("interval '10 minutes'")
)
),
)
)
.order_by(ServiceNotification.created_at)
.with_for_update(skip_locked=True)
.limit(batch_size)
)
r = await uow.session.execute(stmt)
return list(r.scalars().all())
async def ack_notification(uow: UnitOfWork, n_id: int) -> ServiceNotification | None:
notification = await get_notification_by_id(uow, n_id)
if not notification:
return
notification.sent_at = func.now()
notification.status = NotificationStatus.SENT
return notification
async def mark_notification_as_dispatched(uow: UnitOfWork, n_id: int):
notification = await get_notification_by_id(uow, n_id)
if not notification:
return
notification.dispatched_at = func.now()
notification.status = NotificationStatus.DISPATCHED
notification.attempts += 1
return notification

View File

@@ -2,10 +2,15 @@ from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models.service_signatures import ServiceSignature
from schemas.enums import ServiceSignatureStatus
async def get_active_signature_by_kid(session: AsyncSession, kid: str) -> ServiceSignature | None:
stmt = select(ServiceSignature).where(ServiceSignature.kid == kid)
stmt = (
select(ServiceSignature)
.where(ServiceSignature.kid == kid)
.where(ServiceSignature.status == ServiceSignatureStatus.ACTIVE)
)
r = await session.execute(stmt)
return r.scalar_one_or_none()

View File

@@ -1,13 +1,14 @@
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models import Session
from db.session import UnitOfWork
from schemas.providers import ProvidersType
class SessionsRepository:
def __init__(self, session: AsyncSession) -> None:
self.session = session
def __init__(self, uow: UnitOfWork) -> None:
self.uow = uow
self.session = uow.session
async def get_session_by_id(self, id: int) -> Session | None:
stmt = select(Session).where(Session.id == id)
@@ -27,12 +28,10 @@ class SessionsRepository:
async def create(self, user_id: int, refresh_token_hash: str, iss: ProvidersType) -> Session:
obj = Session(user_id=user_id, refresh_token_hash=refresh_token_hash, source=iss)
self.session.add(obj)
await self.session.commit()
return obj
async def revoke(self, token_id: int):
session = await self.get_session_by_id(token_id)
session.is_revoked = True
session.revoked_at = func.now()
await self.session.commit()
return session

View File

@@ -1,13 +1,14 @@
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models import User
from db.models.transactions import BalanceTransaction, BalanceTxType
from db.session import UnitOfWork
class UserRepository:
def __init__(self, session: AsyncSession) -> None:
self.session = session
def __init__(self, uow: UnitOfWork) -> None:
self.uow = uow
self.session = uow.session
async def get_user_by_id(self, id: int) -> User | None:
stmt = select(User).where(User.id == id)
@@ -44,8 +45,6 @@ class UserRepository:
referal_id=referal_id,
)
self.session.add(obj)
await self.session.commit()
return obj
async def increase_balance(
@@ -67,11 +66,8 @@ class UserRepository:
self.session.add(obj)
user.balance += amount
await self.session.commit()
return user
async def update_telegram_id(self, user: User, telegram_id: int) -> User:
user.telegram_id = telegram_id
await self.session.commit()
return user

View File

@@ -8,7 +8,10 @@ asyncpg>=0.31.0
alembic>=1.18.0
aiohttp>=3.14.0
python-multipart==0.0.32
remnawave>=2.6.1
# Private remnawave SDK: bare URL without creds (safe to commit).
# Local install: export REMNAWAVE_SDK_TOKEN=xxx && git config --global url."https://agony:${REMNAWAVE_SDK_TOKEN}@git.mdevs.lat/".insteadOf "https://git.mdevs.lat/" && pip install -r requirements.txt
# Without the token git clone fails with 401/403 (repo is private, login is hardcoded to 'agony').
remnawave @ git+https://git.mdevs.lat/agony/remnawave-sdk.git
zxcvbn>=4.5.0
httpx2>=2.11.0
pytest>=9.1.0

View File

@@ -2,9 +2,10 @@ from fastapi import APIRouter
from .auth import router as auth_router
from .health import router as health_router
from .internal import internal_router
from .link_codes import router as link_code_router
from .orders import router as orders_router
from .payments import payment_routers
from .payments import payment_router
from .plans import router as plans_router
from .users import router as users_routers
@@ -15,5 +16,6 @@ routers: list[APIRouter] = [
users_routers,
health_router,
link_code_router,
*payment_routers,
payment_router,
internal_router,
]

View File

@@ -1,6 +1,4 @@
from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import JSONResponse
from sqlalchemy.ext.asyncio import AsyncSession
from core.secrets import (
estimate_password_strength,
@@ -8,10 +6,10 @@ from core.secrets import (
hash_refresh_token,
verify_password,
)
from db.session import get_db
from db.session import UnitOfWork, get_uow
from repositories.sessions import SessionsRepository
from repositories.users import UserRepository
from schemas.login import UserLogin, UserLoginData, UserTokens
from schemas.login import AuthenticatedUser, UserLoginData, UserTokens
from schemas.providers import ProvidersType
from schemas.registration import UserRegistration
from schemas.user import UserInfo
@@ -21,9 +19,10 @@ from services.users import authorize_user
router = APIRouter(prefix="/auth")
@router.post("/signup")
async def signup(req: UserRegistration, session: AsyncSession = Depends(get_db)):
users_repo = UserRepository(session)
@router.post("/signup", response_model=AuthenticatedUser)
async def signup(req: UserRegistration, uow: UnitOfWork = Depends(get_uow)) -> AuthenticatedUser:
users_repo = UserRepository(uow)
sessions_repo = SessionsRepository(uow)
if req.provider == "credentials":
if not req.username or not req.password:
@@ -44,22 +43,28 @@ async def signup(req: UserRegistration, session: AsyncSession = Depends(get_db))
user = await users_repo.create(
username=req.username, hashed_password=password_hash, referal_id=referal_id
)
return JSONResponse(
UserInfo(
username=user.username,
telegram_id=user.telegram_id,
referal_code=user.referal_code,
bonus_balance=user.balance,
).model_dump(),
status_code=201,
await uow.commit()
info = UserInfo(
username=user.username,
telegram_id=user.telegram_id,
referal_code=user.referal_code,
bonus_balance=user.balance,
)
user_login = await authorize_user(sessions_repo, user, req.provider)
return AuthenticatedUser(
access_token=user_login.access_token,
refresh_token=user_login.refresh_token,
expires_at=user_login.expires_at,
user=info,
)
raise HTTPException(status_code=400, detail="Unsupported provider")
@router.post("/login", response_model=UserLogin)
async def login(req: UserLoginData, session: AsyncSession = Depends(get_db)):
users_repo = UserRepository(session)
sessions_repo = SessionsRepository(session)
@router.post("/login", response_model=AuthenticatedUser)
async def login(req: UserLoginData, uow: UnitOfWork = Depends(get_uow)):
users_repo = UserRepository(uow)
sessions_repo = SessionsRepository(uow)
if req.provider == "credentials":
if not req.username or not req.password:
raise HTTPException(status_code=400, detail="Username or password is not provided.")
@@ -72,6 +77,7 @@ async def login(req: UserLoginData, session: AsyncSession = Depends(get_db)):
raise HTTPException(status_code=401, detail="Invalid password")
data = await authorize_user(sessions_repo, user, req.provider)
await uow.commit()
return data
if req.provider == "telegram":
@@ -82,8 +88,8 @@ async def login(req: UserLoginData, session: AsyncSession = Depends(get_db)):
@router.post("/refresh", response_model=UserTokens)
async def refresh(refresh_token: str, iss: ProvidersType, session: AsyncSession = Depends(get_db)):
sessions_repo = SessionsRepository(session)
async def refresh(refresh_token: str, iss: ProvidersType, uow: UnitOfWork = Depends(get_uow)):
sessions_repo = SessionsRepository(uow)
token_hash = hash_refresh_token(refresh_token)
token_entry = await sessions_repo.get_session_by_hash(token_hash)
@@ -92,6 +98,7 @@ async def refresh(refresh_token: str, iss: ProvidersType, session: AsyncSession
raise HTTPException(status_code=401, detail="Refresh token is invalid.")
key_pair = await refresh_token_rotation(sessions_repo, token_entry, iss)
await uow.commit()
return UserTokens(
access_token=key_pair.access_token,
refresh_token=key_pair.refresh_token,

View File

@@ -0,0 +1,6 @@
from fastapi import APIRouter
from .renewal import router as renewal_router
internal_router = APIRouter(prefix="/internal")
internal_router.include_router(renewal_router)

View File

@@ -0,0 +1,60 @@
from fastapi import APIRouter, Depends, HTTPException
from core.deps import get_service_identity
from db.session import UnitOfWork, get_uow
from repositories.service_notifications import (
ack_notification,
get_pending_notifications,
mark_notification_as_dispatched,
)
from schemas.dto import ServiceIdentity
from schemas.notifications import (
NotificationAcknowledgeRequest,
NotificationResponse,
UserNotificationData,
)
router = APIRouter(prefix="/renewal")
@router.get("/pending", response_model=NotificationResponse)
async def get_pending(
limit: int = 50,
ctx: ServiceIdentity = Depends(get_service_identity),
uow: UnitOfWork = Depends(get_uow),
):
notifications = await get_pending_notifications(uow, limit)
response_users: list[UserNotificationData] = []
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,
notification_type=notification.notify_type,
)
)
await mark_notification_as_dispatched(uow, notification.id)
await uow.commit()
return NotificationResponse(users=response_users, issued_by=ctx.service)
@router.post("/ack")
async def acknowledge(
req: NotificationAcknowledgeRequest,
ctx: ServiceIdentity = Depends(get_service_identity),
uow: UnitOfWork = Depends(get_uow),
):
try:
r = await ack_notification(uow, req.notification_id)
except Exception as e:
raise HTTPException(500, detail=str(e)) from None
if r:
await uow.commit()
return "OK"
raise HTTPException(500, detail="No such notification found.")

View File

@@ -2,12 +2,11 @@ import secrets
from datetime import UTC, datetime, timedelta
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.ext.asyncio import AsyncSession
from config import cfg
from core.deps import get_auth_context, get_service_identity
from db.session import get_db
from repositories.link_codes import create_link_code, get_link_code_by_code
from db.session import UnitOfWork, get_uow
from repositories.link_codes import create_link_code, get_link_code_by_code, use_link_code
from repositories.users import UserRepository
from schemas.dto import AuthContext
from schemas.enums import LinkCodeStatus
@@ -19,13 +18,14 @@ router = APIRouter(prefix="/link-codes")
@router.post("", response_model=LinkCodeResponse, status_code=201)
async def gen_link_code(
ctx: AuthContext = Depends(get_auth_context), session: AsyncSession = Depends(get_db)
ctx: AuthContext = Depends(get_auth_context), uow: UnitOfWork = Depends(get_uow)
):
code = secrets.token_urlsafe(cfg.link_code_length)
exp = datetime.now(UTC) + timedelta(minutes=cfg.link_code_ttl)
link_code = await create_link_code(
session, code=code, user_id=ctx.user.id, status=LinkCodeStatus.ACTIVE, expires_at=exp
uow, code=code, user_id=ctx.user.id, status=LinkCodeStatus.ACTIVE, expires_at=exp
)
await uow.commit()
return LinkCodeResponse(code=link_code.code, expires_at=link_code.expires_at)
@@ -34,14 +34,17 @@ async def gen_link_code(
async def consume_link_code(
payload: LinkCodeConsume,
ctx: AuthContext = Depends(get_service_identity),
session: AsyncSession = Depends(get_db),
uow: UnitOfWork = Depends(get_uow),
):
link_code = await get_link_code_by_code(session, payload.code)
link_code = await get_link_code_by_code(uow, payload.code)
if not link_code:
raise HTTPException(404, detail="Code not found")
users_repo = UserRepository(session)
if link_code.status != LinkCodeStatus.ACTIVE:
raise HTTPException(404, detail="Code expired or is invalid.")
users_repo = UserRepository(uow)
user = await users_repo.get_user_by_id(link_code.user_id)
if not user:
@@ -49,6 +52,9 @@ async def consume_link_code(
user = await users_repo.update_telegram_id(user, payload.telegram_id)
await use_link_code(uow, code=link_code)
await uow.commit()
return UserInfo(
username=user.username,
telegram_id=user.telegram_id,

View File

@@ -2,12 +2,11 @@ import math
from datetime import UTC, datetime
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.ext.asyncio import AsyncSession
from config import cfg
from core.deps import get_auth_context, get_pally_client
from db.models.orders import OrderStatus
from db.session import get_db
from db.session import UnitOfWork, get_uow
from external.pally import PallyClient
from repositories import AddonsRepository, PricingRepository
from repositories.invoices import InvoiceRepository
@@ -31,13 +30,13 @@ router = APIRouter(prefix="/orders")
async def checkout(
order: OrderDetails,
ctx: AuthContext = Depends(get_auth_context),
session: AsyncSession = Depends(get_db),
uow: UnitOfWork = Depends(get_uow),
pally: PallyClient = Depends(get_pally_client),
):
addons_repo = AddonsRepository(session)
pricing_repo = PricingRepository(session)
invoices_repo = InvoiceRepository(session)
orders_repo = OrderRepository(session)
addons_repo = AddonsRepository(uow)
pricing_repo = PricingRepository(uow)
invoices_repo = InvoiceRepository(uow)
orders_repo = OrderRepository(uow)
pricing = await get_pricing_model(addons_repo, pricing_repo)
price = math.ceil(await calculate_price(addons_repo=addons_repo, order=order, pricing=pricing))
@@ -72,7 +71,7 @@ async def checkout(
raise HTTPException(500, detail="failed to create invoice")
else:
await deduct_order_balance(
session,
uow.session,
user=ctx.user,
order=order_entry,
description=f"order {order_entry.id} paid from balance",
@@ -85,7 +84,7 @@ async def checkout(
now=now,
):
await apply_order_now(
session,
uow.session,
user=ctx.user,
order=order_entry,
pricing=pricing,
@@ -95,9 +94,12 @@ async def checkout(
await queue_order_for_later(
order=order_entry, subscription=ctx.user.subscription, now=now
)
await session.commit()
await uow.commit()
payment_link = None
if amount_to_pay > 0:
await uow.commit()
return CheckoutResponse(
order_id=str(order_entry.id),
total_amount=price,

View File

@@ -1,3 +1,6 @@
from fastapi import APIRouter
from .pally import router as pally_router
payment_routers = [pally_router]
payment_router = APIRouter(prefix="/payments")
payment_router.include_router(pally_router)

View File

@@ -3,17 +3,16 @@ import logging
from fastapi import Depends, Form, HTTPException
from fastapi.routing import APIRouter
from sqlalchemy.ext.asyncio import AsyncSession
from core.deps import get_db
from db.models.transactions import BalanceTransaction, BalanceTxType
from db.models.transactions import BalanceTxType
from db.session import UnitOfWork, get_uow
from external.pally import BillStatus
from repositories.invoices import InvoiceRepository
from repositories.users import UserRepository
from schemas.invoices import InvoiceStatus
from services.payments import process_subscription_purchase, validate_pally_signature
router = APIRouter(prefix="/payments/pally")
router = APIRouter(prefix="/pally")
logger = logging.getLogger(__name__)
@@ -39,7 +38,7 @@ async def pally_callback(
PayerComment: str | None = Form(None),
ErrorCode: int | None = Form(None),
ErrorMessage: str | None = Form(None),
session: AsyncSession = Depends(get_db),
uow: UnitOfWork = Depends(get_uow),
):
logger.info(
"Pally webhook received - InvId: %s, OutSum: %s, Commission: %s, TrsId: %s, Status: %s, "
@@ -88,7 +87,7 @@ async def pally_callback(
try:
await process_subscription_purchase(
session,
uow,
invoice_id=int(invoice_id_str),
trs_id=TrsId,
amount=amount,
@@ -102,9 +101,9 @@ async def pally_callback(
invoice_id_str,
TrsId,
)
await session.rollback()
await uow.rollback()
invoice_repo = InvoiceRepository(session)
invoice_repo = InvoiceRepository(uow)
invoice = await invoice_repo.get_by_id(int(invoice_id_str))
if invoice is None:
logger.critical(
@@ -117,7 +116,8 @@ async def pally_callback(
if invoice.status == InvoiceStatus.PAID:
return "OK"
user = await UserRepository(session).get_user_by_id(invoice.creator_id)
users_repo = UserRepository(uow)
user = await users_repo.get_user_by_id(invoice.creator_id)
if user is None:
logger.critical(
"Cannot credit fallback balance: user %s was not found " "for bill %s (TrsId: %s)",
@@ -127,22 +127,14 @@ async def pally_callback(
)
raise
balance_before = user.balance
user.balance += amount
invoice.status = InvoiceStatus.PAID
session.add(
BalanceTransaction(
user_id=user.id,
amount=amount,
tx_type=BalanceTxType.DEPOSIT,
balance_before=balance_before,
balance_after=user.balance,
description=(
f"fallback payment credit for invoice {invoice.id} " f"(TrsId: {TrsId})"
),
)
await users_repo.increase_balance(
user.id,
amount,
BalanceTxType.DEPOSIT,
f"fallback payment credit for invoice {invoice.id} (TrsId: {TrsId})",
)
await session.commit()
await invoice_repo.update_status_by_id(invoice.id, InvoiceStatus.PAID)
await uow.commit()
logger.info(
"Fallback payment credit processed: user_id=%s, amount=%s, " "invoice_id=%s, TrsId=%s",
user.id,

View File

@@ -1,7 +1,6 @@
from fastapi import APIRouter, Depends
from sqlalchemy.ext.asyncio import AsyncSession
from db.session import get_db
from db.session import UnitOfWork, get_uow
from repositories.addons import AddonsRepository
from repositories.pricing import PricingRepository
from schemas.plans import PricingPlans
@@ -11,9 +10,9 @@ router = APIRouter(prefix="/plans")
@router.get("/", response_model=PricingPlans)
async def get_plans(session: AsyncSession = Depends(get_db)):
addons_repo = AddonsRepository(session)
pricing_repo = PricingRepository(session)
async def get_plans(uow: UnitOfWork = Depends(get_uow)):
addons_repo = AddonsRepository(uow)
pricing_repo = PricingRepository(uow)
res = await get_pricing_model(addons_repo, pricing_repo)
return res

View File

@@ -15,3 +15,17 @@ class LinkCodeStatus(StrEnum):
ACTIVE = "active"
USED = "used"
EXPIRED = "expired"
class NotificationType(StrEnum):
SEVEN_DAYS = "7d"
THREE_DAYS = "3d"
ONE_DAY = "1d"
EXPIRED = "expired"
class NotificationStatus(StrEnum):
PENDING = "pending"
DISPATCHED = "dispatched"
SENT = "sent"
FAILED = "failed"

View File

@@ -16,7 +16,7 @@ class UserLoginData(BaseModel):
telegram: TelegramData | None = None
class UserLogin(BaseModel):
class AuthenticatedUser(BaseModel):
access_token: str
refresh_token: str
expires_at: float

24
schemas/notifications.py Normal file
View File

@@ -0,0 +1,24 @@
from datetime import UTC, datetime
from pydantic import BaseModel, Field
from schemas.enums import NotificationType
class UserNotificationData(BaseModel):
notification_id: int = Field()
username: str = Field()
telegram_id: int | None = Field()
expires_at: datetime = Field()
notification_type: NotificationType = Field()
class NotificationResponse(BaseModel):
created_at: datetime = Field(default_factory=lambda: datetime.now(UTC))
issued_by: str = Field()
users: list[UserNotificationData] = Field()
class NotificationAcknowledgeRequest(BaseModel):
notification_id: int = Field()

59
services/notifications.py Normal file
View File

@@ -0,0 +1,59 @@
import asyncio
from datetime import UTC, datetime, timedelta
from sqlalchemy import select
from sqlalchemy.dialects.postgresql import insert
from config import cfg
from db.models import ServiceNotification, Subscription
from db.session import UnitOfWork, async_session
from schemas.enums import NotificationType
async def collect_subscription_notifications(uow: UnitOfWork, now: datetime | None = None) -> int:
session = uow.session
now = now or datetime.now(UTC)
periods = (
(NotificationType.SEVEN_DAYS, now + timedelta(days=3), now + timedelta(days=7)),
(NotificationType.THREE_DAYS, now + timedelta(days=1), now + timedelta(days=3)),
(NotificationType.ONE_DAY, now, now + timedelta(days=1)),
(NotificationType.EXPIRED, None, now),
)
created = 0
for notification_type, starts_at, ends_at in periods:
conditions = [Subscription.expires_at <= ends_at]
if starts_at is not None:
conditions.append(Subscription.expires_at > starts_at)
subscriptions = await session.scalars(select(Subscription).where(*conditions))
rows = [
{
"subscription_id": subscription.id,
"notify_type": notification_type,
"sub_expires_at": subscription.expires_at,
}
for subscription in subscriptions
]
if not rows:
continue
stmt = insert(ServiceNotification).values(rows)
stmt = stmt.on_conflict_do_nothing(constraint="uq_subscription_notification")
result = await session.execute(stmt)
created += result.rowcount or 0
await uow.commit()
return created
async def run_subscription_notifications() -> None:
interval = cfg.notification_scan_interval * 60
while True:
try:
async with async_session() as session, UnitOfWork(session) as uow:
await collect_subscription_notifications(uow)
except Exception:
# A failed iteration must not stop subsequent notification checks.
pass
await asyncio.sleep(interval)

View File

@@ -6,8 +6,8 @@ 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 db.models.transactions import BalanceTxType
from db.session import UnitOfWork
from repositories import AddonsRepository
from repositories.invoices import InvoiceRepository
from repositories.orders import OrderRepository
@@ -15,6 +15,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,
@@ -32,16 +33,17 @@ def validate_pally_signature(out_sum: str, inv_id: str, signature_value: str) ->
async def process_subscription_purchase( # noqa: PLR0911, PLR0912
session,
uow: UnitOfWork,
*,
invoice_id: int,
trs_id: str,
amount: int,
) -> None:
invoice_repo = InvoiceRepository(session)
orders_repo = OrderRepository(session)
users_repo = UserRepository(session)
pricing_repo = PricingRepository(session)
session = uow.session
invoice_repo = InvoiceRepository(uow)
orders_repo = OrderRepository(uow)
users_repo = UserRepository(uow)
pricing_repo = PricingRepository(uow)
invoice = await invoice_repo.get_by_id(invoice_id)
@@ -99,7 +101,7 @@ async def process_subscription_purchase( # noqa: PLR0911, PLR0912
subscription = user.subscription
now = datetime.now(UTC)
pricing = await get_pricing_model(AddonsRepository(session), pricing_repo)
pricing = await get_pricing_model(AddonsRepository(uow), pricing_repo)
logger.info(
"Processing payment: bill_id=%s, order_id=%s, user_id=%s, amount=%s",
@@ -129,12 +131,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(
@@ -147,19 +144,14 @@ async def process_subscription_purchase( # noqa: PLR0911, PLR0912
if referal_id is not None:
referal_amount = math.floor(amount * (cfg.referal_bonus / 100))
if referal_amount > 0:
referal_user = await session.get(type(user), referal_id)
referal_user = await users_repo.get_user_by_id(referal_id)
if referal_user is not None:
session.add(
BalanceTransaction(
user_id=referal_id,
amount=referal_amount,
tx_type=BalanceTxType.REFERRAL_BONUS,
balance_before=referal_user.balance,
balance_after=referal_user.balance + referal_amount,
description=f"referral reward for user {invoice.creator_id} (TrsId: {trs_id})",
)
await users_repo.increase_balance(
referal_id,
referal_amount,
BalanceTxType.REFERRAL_BONUS,
f"referral reward for user {invoice.creator_id} (TrsId: {trs_id})",
)
referal_user.balance += referal_amount
logger.info(
"Referral bonus processed: referrer_id=%s, amount=%s",
@@ -174,6 +166,6 @@ async def process_subscription_purchase( # noqa: PLR0911, PLR0912
trs_id,
)
await session.commit()
await uow.commit()
logger.info("Bill %s marked as PAID for TrsId %s", invoice_id, trs_id)

126
services/rw_sync.py Normal file
View File

@@ -0,0 +1,126 @@
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 UnitOfWork, 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(uow: UnitOfWork, batch_size: int = 50) -> int:
session = uow.session
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 uow.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 uow.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 uow.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, UnitOfWork(session) as uow:
await process_rw_syncs(uow)
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:
async with UnitOfWork(session) as uow:
count = await enqueue_all_rw_syncs(uow.session)
await uow.commit()
logger.info("Queued %s subscriptions for RW reconciliation", count)
except Exception:
logger.exception("RW reconciliation iteration failed")
await asyncio.sleep(interval)

View File

@@ -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)

View File

@@ -1,21 +1,21 @@
from core.secrets import generate_pair, hash_refresh_token
from db.models.users import User
from repositories.sessions import SessionsRepository
from schemas.login import UserLogin
from schemas.login import AuthenticatedUser
from schemas.providers import ProvidersType
from schemas.user import UserInfo
async def authorize_user(
sessions_repo: SessionsRepository, user: User, iss: ProvidersType
) -> UserLogin:
) -> AuthenticatedUser:
key_pair = generate_pair(user.id, iss)
refresh_token_hash = hash_refresh_token(key_pair.refresh_token)
await sessions_repo.create(user_id=user.id, refresh_token_hash=refresh_token_hash, iss=iss)
return UserLogin(
return AuthenticatedUser(
access_token=key_pair.access_token,
refresh_token=key_pair.refresh_token,
user=UserInfo(

View File

@@ -49,15 +49,17 @@ def test_signup_rejects_weak_password(client):
def test_signup_accepts_unknown_referral_code(client):
created_user = SimpleNamespace(
username="alice", telegram_id=None, referal_code="new-code", balance=0
id=1, username="alice", telegram_id=None, referal_code="new-code", balance=0
)
repository = SimpleNamespace(
get_user_by_username=AsyncMock(return_value=None),
get_user_by_ref_code=AsyncMock(return_value=None),
create=AsyncMock(return_value=created_user),
)
sessions_repository = SimpleNamespace(create=AsyncMock(return_value=None))
with (
patch("routes.auth.UserRepository", return_value=repository),
patch("routes.auth.SessionsRepository", return_value=sessions_repository),
patch("routes.auth.estimate_password_strength", return_value=True),
patch("routes.auth.hash_password", return_value="hashed"),
):
@@ -71,11 +73,11 @@ def test_signup_accepts_unknown_referral_code(client):
},
)
assert response.status_code == 201
assert response.status_code == 200
repository.create.assert_awaited_once_with(
username="alice", hashed_password="hashed", referal_id=None
)
assert response.json()["referal_code"] == "new-code"
assert response.json()["user"]["referal_code"] == "new-code"
def test_login_distinguishes_unknown_user_and_bad_password(client):

View File

@@ -93,7 +93,7 @@ def test_consume_link_code_rejects_unknown_code(client):
def test_consume_link_code_rejects_code_for_deleted_user(client):
client.app.dependency_overrides[link_codes.get_service_identity] = service_identity
link_code = SimpleNamespace(user_id=42)
link_code = SimpleNamespace(user_id=42, status=LinkCodeStatus.ACTIVE)
repository = SimpleNamespace(get_user_by_id=AsyncMock(return_value=None))
with (
patch("routes.link_codes.get_link_code_by_code", new=AsyncMock(return_value=link_code)),
@@ -110,7 +110,7 @@ def test_consume_link_code_rejects_code_for_deleted_user(client):
def test_consume_link_code_updates_telegram_id_and_returns_user(client):
client.app.dependency_overrides[link_codes.get_service_identity] = service_identity
link_code = SimpleNamespace(user_id=42)
link_code = SimpleNamespace(user_id=42, status=LinkCodeStatus.ACTIVE)
user = SimpleNamespace(
username="alice", telegram_id=12345, referal_code="ref-code", balance=100
)

View File

@@ -0,0 +1,47 @@
import asyncio
from datetime import UTC, datetime, timedelta
from types import SimpleNamespace
from schemas.enums import NotificationType
from services.notifications import collect_subscription_notifications
class FakeUnitOfWork:
def __init__(self, subscriptions_by_period):
self.session = self
self.subscriptions_by_period = iter(subscriptions_by_period)
self.executed = []
self.committed = False
async def scalars(self, _statement):
return next(self.subscriptions_by_period)
async def execute(self, statement):
self.executed.append(statement)
return SimpleNamespace(rowcount=1)
async def commit(self):
self.committed = True
def test_collects_notification_for_each_expiry_period():
now = datetime(2026, 8, 20, tzinfo=UTC)
uow = FakeUnitOfWork(
[
[SimpleNamespace(id=1, expires_at=now + timedelta(days=6))],
[SimpleNamespace(id=2, expires_at=now + timedelta(days=2))],
[SimpleNamespace(id=3, expires_at=now + timedelta(hours=12))],
[SimpleNamespace(id=4, expires_at=now - timedelta(hours=1))],
]
)
created = asyncio.run(collect_subscription_notifications(uow, now))
assert created == len(uow.executed)
assert uow.committed
assert [statement.compile().params["notify_type_m0"] for statement in uow.executed] == [
NotificationType.SEVEN_DAYS,
NotificationType.THREE_DAYS,
NotificationType.ONE_DAY,
NotificationType.EXPIRED,
]