From a374501e4dca6517457fc649909c049a70746127 Mon Sep 17 00:00:00 2001 From: hexdev Date: Mon, 24 Aug 2026 11:07:26 +0700 Subject: [PATCH] fix!: removed .commit() from repository level --- core/auth/fetch_sub.py | 9 +++---- core/auth/jwt.py | 14 +++++----- core/deps.py | 15 +++++------ db/session.py | 25 +++++++++++++++++- repositories/addons.py | 6 ++--- repositories/invoices.py | 11 +++----- repositories/link_codes.py | 15 +++++------ repositories/orders.py | 8 +++--- repositories/pricing.py | 6 ++--- repositories/service_notifications.py | 21 +++++++-------- repositories/sessions.py | 9 +++---- repositories/users.py | 12 +++------ routes/auth.py | 20 +++++++------- routes/internal/renewal.py | 15 ++++++----- routes/link_codes.py | 17 ++++++------ routes/orders.py | 22 +++++++++------- routes/payments/pally.py | 38 +++++++++++---------------- routes/plans.py | 9 +++---- services/notifications.py | 17 +++++------- services/payments.py | 35 +++++++++++------------- services/rw_sync.py | 20 +++++++------- tests/test_notifications.py | 13 ++++----- 22 files changed, 177 insertions(+), 180 deletions(-) diff --git a/core/auth/fetch_sub.py b/core/auth/fetch_sub.py index 687e4a9..6f96a50 100644 --- a/core/auth/fetch_sub.py +++ b/core/auth/fetch_sub.py @@ -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]) diff --git a/core/auth/jwt.py b/core/auth/jwt.py index b12bb8e..0f0d8a7 100644 --- a/core/auth/jwt.py +++ b/core/auth/jwt.py @@ -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: diff --git a/core/deps.py b/core/deps.py index 51ea800..83bae3e 100644 --- a/core/deps.py +++ b/core/deps.py @@ -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,14 +20,14 @@ 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: @@ -39,7 +38,7 @@ async def get_service_identity( 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) diff --git a/db/session.py b/db/session.py index cf36752..8986066 100644 --- a/db/session.py +++ b/db/session.py @@ -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 diff --git a/repositories/addons.py b/repositories/addons.py index 0ebaad2..5bb48ad 100644 --- a/repositories/addons.py +++ b/repositories/addons.py @@ -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) diff --git a/repositories/invoices.py b/repositories/invoices.py index 2873484..b2427aa 100644 --- a/repositories/invoices.py +++ b/repositories/invoices.py @@ -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 diff --git a/repositories/link_codes.py b/repositories/link_codes.py index 78a4eba..99cb59f 100644 --- a/repositories/link_codes.py +++ b/repositories/link_codes.py @@ -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 diff --git a/repositories/orders.py b/repositories/orders.py index 8f66eef..adc04d6 100644 --- a/repositories/orders.py +++ b/repositories/orders.py @@ -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 diff --git a/repositories/pricing.py b/repositories/pricing.py index f6a01d8..de08709 100644 --- a/repositories/pricing.py +++ b/repositories/pricing.py @@ -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) diff --git a/repositories/service_notifications.py b/repositories/service_notifications.py index 4c497c0..2541825 100644 --- a/repositories/service_notifications.py +++ b/repositories/service_notifications.py @@ -1,19 +1,19 @@ from sqlalchemy import func, or_, select, text -from sqlalchemy.ext.asyncio import AsyncSession from db.models.service_notifications import ServiceNotification +from db.session import UnitOfWork from schemas.enums import NotificationStatus -async def get_notification_by_id(session: AsyncSession, n_id: int) -> ServiceNotification | None: +async def get_notification_by_id(uow: UnitOfWork, n_id: int) -> ServiceNotification | None: stmt = select(ServiceNotification).where(ServiceNotification.id == n_id) - r = await session.execute(stmt) + r = await uow.session.execute(stmt) return r.scalar_one_or_none() async def get_pending_notifications( - session: AsyncSession, batch_size: int = 50 + uow: UnitOfWork, batch_size: int = 50 ) -> list[ServiceNotification]: stmt = ( select(ServiceNotification) @@ -33,26 +33,24 @@ async def get_pending_notifications( .with_for_update(skip_locked=True) .limit(batch_size) ) - r = await session.execute(stmt) + r = await uow.session.execute(stmt) return list(r.scalars().all()) -async def ack_notification(session: AsyncSession, n_id: int) -> ServiceNotification | None: - notification = await get_notification_by_id(session, n_id) +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 - await session.commit() - return notification -async def mark_notification_as_dispatched(session: AsyncSession, n_id: int): - notification = await get_notification_by_id(session, n_id) +async def mark_notification_as_dispatched(uow: UnitOfWork, n_id: int): + notification = await get_notification_by_id(uow, n_id) if not notification: return @@ -60,5 +58,4 @@ async def mark_notification_as_dispatched(session: AsyncSession, n_id: int): notification.status = NotificationStatus.DISPATCHED notification.attempts += 1 - await session.commit() return notification diff --git a/repositories/sessions.py b/repositories/sessions.py index f54182e..4859e64 100644 --- a/repositories/sessions.py +++ b/repositories/sessions.py @@ -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 diff --git a/repositories/users.py b/repositories/users.py index 24cf919..57e39e1 100644 --- a/repositories/users.py +++ b/repositories/users.py @@ -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 diff --git a/routes/auth.py b/routes/auth.py index 21f7c27..e3e65a6 100644 --- a/routes/auth.py +++ b/routes/auth.py @@ -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, diff --git a/routes/internal/renewal.py b/routes/internal/renewal.py index dda7c94..432204a 100644 --- a/routes/internal/renewal.py +++ b/routes/internal/renewal.py @@ -1,8 +1,7 @@ from fastapi import APIRouter, Depends, HTTPException -from sqlalchemy.ext.asyncio import AsyncSession from core.deps import get_service_identity -from db.session import get_db +from db.session import UnitOfWork, get_uow from repositories.service_notifications import ( ack_notification, get_pending_notifications, @@ -22,9 +21,9 @@ router = APIRouter(prefix="/renewal") async def get_pending( limit: int = 50, ctx: ServiceIdentity = Depends(get_service_identity), - session: AsyncSession = Depends(get_db), + uow: UnitOfWork = Depends(get_uow), ): - notifications = await get_pending_notifications(session, limit) + notifications = await get_pending_notifications(uow, limit) response_users: list[UserNotificationData] = [] for notification in notifications: @@ -38,8 +37,9 @@ async def get_pending( ) ) - await mark_notification_as_dispatched(session, notification.id) + await mark_notification_as_dispatched(uow, notification.id) + await uow.commit() return NotificationResponse(users=response_users, issued_by=ctx.service) @@ -47,13 +47,14 @@ async def get_pending( async def acknowledge( req: NotificationAcknowledgeRequest, ctx: ServiceIdentity = Depends(get_service_identity), - session: AsyncSession = Depends(get_db), + uow: UnitOfWork = Depends(get_uow), ): try: - r = await ack_notification(session, req.notification_id) + 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.") diff --git a/routes/link_codes.py b/routes/link_codes.py index 99f6fd9..50094d9 100644 --- a/routes/link_codes.py +++ b/routes/link_codes.py @@ -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, diff --git a/routes/orders.py b/routes/orders.py index b4b9edd..8bd3b89 100644 --- a/routes/orders.py +++ b/routes/orders.py @@ -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, diff --git a/routes/payments/pally.py b/routes/payments/pally.py index dfc47a0..a141396 100644 --- a/routes/payments/pally.py +++ b/routes/payments/pally.py @@ -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, diff --git a/routes/plans.py b/routes/plans.py index fbe7471..67d83b2 100644 --- a/routes/plans.py +++ b/routes/plans.py @@ -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 diff --git a/services/notifications.py b/services/notifications.py index 9ee3002..e8ca0d7 100644 --- a/services/notifications.py +++ b/services/notifications.py @@ -1,22 +1,17 @@ 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 db.session import UnitOfWork, 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: +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)), @@ -48,7 +43,7 @@ async def collect_subscription_notifications( result = await session.execute(stmt) created += result.rowcount or 0 - await session.commit() + await uow.commit() return created @@ -56,8 +51,8 @@ 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) + 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 diff --git a/services/payments.py b/services/payments.py index 7f7e571..a9512e1 100644 --- a/services/payments.py +++ b/services/payments.py @@ -6,7 +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 db.models.transactions import BalanceTxType +from db.session import UnitOfWork from repositories import AddonsRepository from repositories.invoices import InvoiceRepository from repositories.orders import OrderRepository @@ -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", @@ -142,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", @@ -169,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) diff --git a/services/rw_sync.py b/services/rw_sync.py index a2adeea..ecbbfbb 100644 --- a/services/rw_sync.py +++ b/services/rw_sync.py @@ -8,7 +8,7 @@ from sqlalchemy.dialects.postgresql import insert from config import cfg from db.models import RWSyncOutbox, Subscription -from db.session import async_session +from db.session import UnitOfWork, async_session from external.rw import sync_subscription_by_telegram_id if TYPE_CHECKING: @@ -48,7 +48,8 @@ async def enqueue_all_rw_syncs(session: "AsyncSession") -> int: return count -async def process_rw_syncs(session: "AsyncSession", batch_size: int = 50) -> int: +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) @@ -63,7 +64,7 @@ async def process_rw_syncs(session: "AsyncSession", batch_size: int = 50) -> int jobs = list(jobs) for job in jobs: job.locked_until = now + LOCK_DURATION - await session.commit() + await uow.commit() for job in jobs: subscription = job.subscription @@ -81,7 +82,7 @@ async def process_rw_syncs(session: "AsyncSession", batch_size: int = 50) -> int RWSyncOutbox.revision == job.revision, ) ) - await session.commit() + await uow.commit() continue delay = min(timedelta(minutes=2**job.attempts), MAX_RETRY_DELAY) @@ -94,7 +95,7 @@ async def process_rw_syncs(session: "AsyncSession", batch_size: int = 50) -> int locked_until=None, ) ) - await session.commit() + await uow.commit() logger.warning("RW sync failed for subscription_id=%s", subscription.id) return len(jobs) @@ -104,8 +105,8 @@ async def run_rw_sync_worker() -> None: interval = cfg.rw_sync_interval_minutes * 60 while True: try: - async with async_session() as session: - await process_rw_syncs(session) + 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) @@ -116,8 +117,9 @@ async def run_rw_sync_reconciler() -> None: while True: try: async with async_session() as session: - count = await enqueue_all_rw_syncs(session) - await session.commit() + 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") diff --git a/tests/test_notifications.py b/tests/test_notifications.py index b1b46b4..398cb5a 100644 --- a/tests/test_notifications.py +++ b/tests/test_notifications.py @@ -6,8 +6,9 @@ from schemas.enums import NotificationType from services.notifications import collect_subscription_notifications -class FakeSession: +class FakeUnitOfWork: def __init__(self, subscriptions_by_period): + self.session = self self.subscriptions_by_period = iter(subscriptions_by_period) self.executed = [] self.committed = False @@ -25,7 +26,7 @@ class FakeSession: def test_collects_notification_for_each_expiry_period(): now = datetime(2026, 8, 20, tzinfo=UTC) - session = FakeSession( + uow = FakeUnitOfWork( [ [SimpleNamespace(id=1, expires_at=now + timedelta(days=6))], [SimpleNamespace(id=2, expires_at=now + timedelta(days=2))], @@ -34,11 +35,11 @@ def test_collects_notification_for_each_expiry_period(): ] ) - created = asyncio.run(collect_subscription_notifications(session, now)) + created = asyncio.run(collect_subscription_notifications(uow, now)) - assert created == len(session.executed) - assert session.committed - assert [statement.compile().params["notify_type_m0"] for statement in session.executed] == [ + 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,