diff --git a/main.py b/main.py index 6a6e0ca..91cb4c8 100644 --- a/main.py +++ b/main.py @@ -15,6 +15,10 @@ from src.presentation.middlewares import MetricsMiddleware from src.infrastructure.tasks.cleanup_magic_tokens import cleanup_telegram_tokens from src.config import get_database_url +from src.infrastructure.tasks.remove_bans import remove_expired_bans +from src.infrastructure.db import database + +scheduler = AsyncIOScheduler() # Планировщик @asynccontextmanager async def lifespan(app: FastAPI): @@ -24,17 +28,31 @@ async def lifespan(app: FastAPI): DB_URL ) - SessionLocal = async_sessionmaker( - bind=engine, - class_=AsyncSession, + database.SessionLocal = async_sessionmaker( + bind=engine, + class_=AsyncSession, expire_on_commit=False ) - app.state.SessionLocal = SessionLocal + @app.on_event("startup") + async def startup(): + scheduler.add_job( + cleanup_telegram_tokens, + "interval", + minutes=5 + ) + scheduler.add_job( + remove_expired_bans, + "interval", + minutes=10 + ) + scheduler.start() + # При старте выполняется ко до yiled yield # выполняется код после yiled и остановка + scheduler.dispose() await engine.dispose() @@ -69,18 +87,6 @@ async def lifespan(app: FastAPI): allow_credentials=True ) -scheduler = AsyncIOScheduler() # Планировщик -@app.on_event("startup") -async def startup(): - scheduler.add_job( - cleanup_telegram_tokens, - "interval", - minutes=5 - ) - scheduler.start() - - - @app.get("/") @limiter.limit("5/minute") def welcome(request: Request): diff --git a/src/domain/models/ban_model.py b/src/domain/models/ban_model.py index 76c9467..9732977 100644 --- a/src/domain/models/ban_model.py +++ b/src/domain/models/ban_model.py @@ -1,7 +1,7 @@ from datetime import datetime, timezone from typing import TYPE_CHECKING -from sqlalchemy import CheckConstraint, DateTime, ForeignKey, Index, Integer, String, func +from sqlalchemy import CheckConstraint, DateTime, ForeignKey, Index, Integer, String from sqlalchemy.orm import Mapped, mapped_column, relationship, validates from src.infrastructure.db.database import Base @@ -56,7 +56,6 @@ class Ban(Base): unique=True, postgresql_where=( revoked_at.is_(None) - & (expires_at.is_(None) | (expires_at > func.now())) ), ), ) diff --git a/src/infrastructure/db/database.py b/src/infrastructure/db/database.py index a560c70..c4cffe0 100644 --- a/src/infrastructure/db/database.py +++ b/src/infrastructure/db/database.py @@ -1,15 +1,18 @@ from sqlalchemy.orm import declarative_base -from fastapi import Request +from sqlalchemy.ext.asyncio import ( + AsyncSession, + async_sessionmaker +) Base = declarative_base() -async def get_db(request: Request): - - async with request.app.state.SessionLocal() as session: - try: +SessionLocal: async_sessionmaker[AsyncSession] | None = None + +async def get_db(): + if SessionLocal is None: + raise RuntimeError("Database is not initialied") + async with SessionLocal() as session: yield session - finally: - await session.close() diff --git a/src/infrastructure/tasks/cleanup_magic_tokens.py b/src/infrastructure/tasks/cleanup_magic_tokens.py index 3386ea9..7a456c6 100644 --- a/src/infrastructure/tasks/cleanup_magic_tokens.py +++ b/src/infrastructure/tasks/cleanup_magic_tokens.py @@ -1,9 +1,9 @@ from sqlalchemy import delete, or_ from datetime import datetime, timezone -from src.infrastructure.db.database import SessionLocal from src.domain.models import MagicToken from src.logger import logger +from src.infrastructure.db.database import SessionLocal async def cleanup_telegram_tokens(): """Очистка истекших токенов""" diff --git a/src/infrastructure/tasks/remove_bans.py b/src/infrastructure/tasks/remove_bans.py new file mode 100644 index 0000000..8cca3e1 --- /dev/null +++ b/src/infrastructure/tasks/remove_bans.py @@ -0,0 +1,19 @@ +from sqlalchemy import update +from datetime import datetime, timezone + +from src.domain.models.ban_model import Ban +from src.infrastructure.db.database import SessionLocal + +async def remove_expired_bans(): + async with SessionLocal() as db: + await db.execute(update(Ban).where( + Ban.expires_at.is_not(None), + Ban.expires_at < datetime.now(timezone.utc), + Ban.revoked_at.is_(None) + ).values( + revoked_at = datetime.now(timezone.utc), + revoked_reason="Expiration of the Term" + ) + ) + + await db.commit() \ No newline at end of file