feat: cron worker for renewal notifs
This commit is contained in:
@@ -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"],
|
||||||
|
)
|
||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
18
main.py
@@ -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
64
services/notifications.py
Normal 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)
|
||||||
46
tests/test_notifications.py
Normal file
46
tests/test_notifications.py
Normal 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,
|
||||||
|
]
|
||||||
Reference in New Issue
Block a user