Compare commits

...

4 Commits

Author SHA1 Message Date
a374501e4d fix!: removed .commit() from repository level 2026-08-24 11:07:26 +07:00
8d1b753b99 feat: advanced RW sync system 2026-08-24 10:22:14 +07:00
b39cae8046 feat: cron worker for renewal notifs 2026-08-20 21:28:43 +07:00
0a58a41930 feat: /internal/renewal/pending 2026-08-20 21:01:22 +07:00
37 changed files with 796 additions and 153 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,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 ###

View File

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

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

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

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

10
external/rw.py vendored
View File

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

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,20 +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(session: AsyncSession, code: LinkCode) -> LinkCode:
async def use_link_code(uow: UnitOfWork, code: LinkCode) -> LinkCode:
code.status = LinkCodeStatus.USED
await session.commit()
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,6 @@ 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

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

@@ -2,6 +2,7 @@ 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_router
@@ -16,4 +17,5 @@ routers: list[APIRouter] = [
health_router,
link_code_router,
payment_router,
internal_router,
]

View File

@@ -1,6 +1,5 @@
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,7 +7,7 @@ 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
@@ -22,8 +21,8 @@ router = APIRouter(prefix="/auth")
@router.post("/signup")
async def signup(req: UserRegistration, session: AsyncSession = Depends(get_db)):
users_repo = UserRepository(session)
async def signup(req: UserRegistration, uow: UnitOfWork = Depends(get_uow)):
users_repo = UserRepository(uow)
if req.provider == "credentials":
if not req.username or not req.password:
@@ -44,6 +43,7 @@ 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
)
await uow.commit()
return JSONResponse(
UserInfo(
username=user.username,
@@ -57,9 +57,9 @@ async def signup(req: UserRegistration, session: AsyncSession = Depends(get_db))
@router.post("/login", response_model=UserLogin)
async def login(req: UserLoginData, session: AsyncSession = Depends(get_db)):
users_repo = UserRepository(session)
sessions_repo = SessionsRepository(session)
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 +72,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 +83,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 +93,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,11 +2,10 @@ 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 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
@@ -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,9 +34,9 @@ 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")
@@ -44,7 +44,7 @@ async def consume_link_code(
if link_code.status != LinkCodeStatus.ACTIVE:
raise HTTPException(404, detail="Code expired or is invalid.")
users_repo = UserRepository(session)
users_repo = UserRepository(uow)
user = await users_repo.get_user_by_id(link_code.user_id)
if not user:
@@ -52,7 +52,8 @@ async def consume_link_code(
user = await users_repo.update_telegram_id(user, payload.telegram_id)
await use_link_code(session, code=link_code)
await use_link_code(uow, code=link_code)
await uow.commit()
return UserInfo(
username=user.username,

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

@@ -3,10 +3,9 @@ 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
@@ -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"

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

@@ -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,
]