feat: cron worker for renewal notifs

This commit is contained in:
2026-08-20 21:28:43 +07:00
parent 0a58a41930
commit b39cae8046
6 changed files with 161 additions and 4 deletions

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

@@ -37,6 +37,7 @@ class Settings(BaseSettings):
access_token_ttl: int = Field(description="Access token TTL (minutes)") access_token_ttl: int = Field(description="Access token TTL (minutes)")
link_code_ttl: int = Field(8, description="Link Code 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) min_password_length: int = Field(8)
link_code_length: int = Field(8) link_code_length: int = Field(8)

View File

@@ -23,9 +23,7 @@ class ServiceNotification(Base):
notify_type: Mapped[NotificationType] = mapped_column( notify_type: Mapped[NotificationType] = mapped_column(
Enum(NotificationType, name="notificationtype"), nullable=False Enum(NotificationType, name="notificationtype"), nullable=False
) )
sub_expires_at: Mapped[datetime] = mapped_column( sub_expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
DateTime(timezone=True), unique=True, nullable=False
)
status: Mapped[NotificationStatus] = mapped_column( status: Mapped[NotificationStatus] = mapped_column(
Enum(NotificationStatus, name="notificationstatus"), Enum(NotificationStatus, name="notificationstatus"),

18
main.py
View File

@@ -1,8 +1,24 @@
import asyncio
from contextlib import asynccontextmanager, suppress
from fastapi import FastAPI from fastapi import FastAPI
from routes import routers 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: for r in routers:
app.include_router(r) app.include_router(r)

64
services/notifications.py Normal file
View File

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

View File

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