diff --git a/alembic/versions/af1c3d7e4b20_remove_service_notification_expiry_unique.py b/alembic/versions/af1c3d7e4b20_remove_service_notification_expiry_unique.py new file mode 100644 index 0000000..1128c7d --- /dev/null +++ b/alembic/versions/af1c3d7e4b20_remove_service_notification_expiry_unique.py @@ -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"], + ) diff --git a/config.py b/config.py index 114fffa..93c18a4 100644 --- a/config.py +++ b/config.py @@ -37,6 +37,7 @@ 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) min_password_length: int = Field(8) link_code_length: int = Field(8) diff --git a/db/models/service_notifications.py b/db/models/service_notifications.py index 8dfd5d0..a3045c4 100644 --- a/db/models/service_notifications.py +++ b/db/models/service_notifications.py @@ -23,9 +23,7 @@ class ServiceNotification(Base): notify_type: Mapped[NotificationType] = mapped_column( Enum(NotificationType, name="notificationtype"), nullable=False ) - sub_expires_at: Mapped[datetime] = mapped_column( - DateTime(timezone=True), unique=True, nullable=False - ) + sub_expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) status: Mapped[NotificationStatus] = mapped_column( Enum(NotificationStatus, name="notificationstatus"), diff --git a/main.py b/main.py index 53e9612..b65fed4 100644 --- a/main.py +++ b/main.py @@ -1,8 +1,24 @@ +import asyncio +from contextlib import asynccontextmanager, suppress + from fastapi import FastAPI from routes import routers +from services.notifications import run_subscription_notifications -app = FastAPI(debug=True) + +@asynccontextmanager +async def lifespan(_: FastAPI): + task = asyncio.create_task(run_subscription_notifications()) + try: + yield + finally: + task.cancel() + with suppress(asyncio.CancelledError): + await task + + +app = FastAPI(debug=True, lifespan=lifespan) for r in routers: app.include_router(r) diff --git a/services/notifications.py b/services/notifications.py new file mode 100644 index 0000000..9ee3002 --- /dev/null +++ b/services/notifications.py @@ -0,0 +1,64 @@ +import asyncio +from datetime import UTC, datetime, timedelta +from typing import TYPE_CHECKING + +from sqlalchemy import select +from sqlalchemy.dialects.postgresql import insert + +from config import cfg +from db.models import ServiceNotification, Subscription +from db.session import async_session +from schemas.enums import NotificationType + +if TYPE_CHECKING: + from sqlalchemy.ext.asyncio import AsyncSession + + +async def collect_subscription_notifications( + session: "AsyncSession", now: datetime | None = None +) -> int: + 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 session.commit() + return created + + +async def run_subscription_notifications() -> None: + interval = cfg.notification_scan_interval * 60 + while True: + try: + async with async_session() as session: + await collect_subscription_notifications(session) + except Exception: + # A failed iteration must not stop subsequent notification checks. + pass + await asyncio.sleep(interval) diff --git a/tests/test_notifications.py b/tests/test_notifications.py new file mode 100644 index 0000000..b1b46b4 --- /dev/null +++ b/tests/test_notifications.py @@ -0,0 +1,46 @@ +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 FakeSession: + def __init__(self, subscriptions_by_period): + 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) + session = FakeSession( + [ + [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(session, now)) + + assert created == len(session.executed) + assert session.committed + assert [statement.compile().params["notify_type_m0"] for statement in session.executed] == [ + NotificationType.SEVEN_DAYS, + NotificationType.THREE_DAYS, + NotificationType.ONE_DAY, + NotificationType.EXPIRED, + ]