fix!: removed .commit() from repository level

This commit is contained in:
2026-08-24 11:07:26 +07:00
parent 8d1b753b99
commit a374501e4d
22 changed files with 177 additions and 180 deletions

View File

@@ -1,14 +1,11 @@
from sqlalchemy.ext.asyncio import AsyncSession
from db.models import User from db.models import User
from db.session import UnitOfWork
from repositories.users import UserRepository from repositories.users import UserRepository
from schemas.jwt import ServiceJWTPayload from schemas.jwt import ServiceJWTPayload
async def fetch_subject_from_service( async def fetch_subject_from_service(payload: ServiceJWTPayload, uow: UnitOfWork) -> User | None:
payload: ServiceJWTPayload, session: AsyncSession repo = UserRepository(uow)
) -> User | None:
repo = UserRepository(session)
if payload.acting_as.startswith("telegram:"): if payload.acting_as.startswith("telegram:"):
telegram_id = int(payload.acting_as.split("telegram:")[1]) telegram_id = int(payload.acting_as.split("telegram:")[1])

View File

@@ -2,18 +2,18 @@ from datetime import UTC, datetime
from fastapi import HTTPException from fastapi import HTTPException
from pydantic import ValidationError from pydantic import ValidationError
from sqlalchemy.ext.asyncio import AsyncSession
from core.auth.fetch_sub import fetch_subject_from_service from core.auth.fetch_sub import fetch_subject_from_service
from core.secrets import decode_jwt, decode_user_jwt, get_kid_from_token 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.service_signatures import get_active_signature_by_kid
from repositories.users import UserRepository from repositories.users import UserRepository
from schemas.dto import AuthContext from schemas.dto import AuthContext
from schemas.jwt import ServiceJWTPayload, UserJWTPayload from schemas.jwt import ServiceJWTPayload, UserJWTPayload
async def authorize_bot(kid: str, token: str, session: AsyncSession) -> AuthContext: async def authorize_bot(kid: str, token: str, uow: UnitOfWork) -> AuthContext:
signature = await get_active_signature_by_kid(session, kid) signature = await get_active_signature_by_kid(uow.session, kid)
if not signature: if not signature:
raise HTTPException(401, detail="Invalid service 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(): if payload.exp < datetime.now(UTC).timestamp():
raise HTTPException(status_code=401, detail="Access token expired") 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: if not subject:
raise HTTPException(status_code=401, detail="User not found") raise HTTPException(status_code=401, detail="User not found")
return AuthContext(subject, auth_method="service", service=kid) 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) kid = get_kid_from_token(token)
if kid: if kid:
return await authorize_bot(kid, token, session) return await authorize_bot(kid, token, uow)
content = decode_user_jwt(token) 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(): if payload.exp < datetime.now(UTC).timestamp():
raise HTTPException(status_code=401, detail="Access token expired") 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)) user = await repo.get_user_by_id(int(payload.sub))
if not user: if not user:

View File

@@ -1,10 +1,9 @@
from fastapi import Depends, HTTPException, Request from fastapi import Depends, HTTPException, Request
from sqlalchemy.ext.asyncio import AsyncSession
from config import cfg from config import cfg
from core.auth import jwt from core.auth import jwt
from core.secrets import get_kid_from_token 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 external.pally import PallyClient
from repositories.service_signatures import get_active_signature_by_kid from repositories.service_signatures import get_active_signature_by_kid
from schemas.dto import AuthContext, ServiceIdentity from schemas.dto import AuthContext, ServiceIdentity
@@ -12,7 +11,7 @@ from services.subscriptions import sync_user_subscription
async def get_auth_context( async def get_auth_context(
request: Request, session: AsyncSession = Depends(get_db) request: Request, uow: UnitOfWork = Depends(get_uow)
) -> AuthContext | None: ) -> AuthContext | None:
auth = request.headers.get("Authorization") auth = request.headers.get("Authorization")
@@ -21,14 +20,14 @@ async def get_auth_context(
if auth.startswith("Bearer"): if auth.startswith("Bearer"):
token = auth.removeprefix("Bearer ").strip() token = auth.removeprefix("Bearer ").strip()
ctx = await jwt.authorize(token, session) ctx = await jwt.authorize(token, uow)
await sync_user_subscription(session, user=ctx.user) await sync_user_subscription(uow.session, user=ctx.user)
await session.commit() await uow.commit()
return ctx return ctx
async def get_service_identity( async def get_service_identity(
request: Request, session: AsyncSession = Depends(get_db) request: Request, uow: UnitOfWork = Depends(get_uow)
) -> ServiceIdentity | None: ) -> ServiceIdentity | None:
auth = request.headers.get("Authorization") auth = request.headers.get("Authorization")
if not auth: if not auth:
@@ -39,7 +38,7 @@ async def get_service_identity(
kid = get_kid_from_token(token) kid = get_kid_from_token(token)
if not kid: if not kid:
raise HTTPException(401, detail="No kid provided.") 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: if not signature:
raise HTTPException(401, detail="Invalid signature") raise HTTPException(401, detail="Invalid signature")
return ServiceIdentity(service=signature.kid) return ServiceIdentity(service=signature.kid)

View File

@@ -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 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) 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 def get_db():
async with async_session() as session: async with async_session() as session:
yield session yield session
async def get_uow():
async with async_session() as session, UnitOfWork(session) as uow:
yield uow

View File

@@ -1,12 +1,12 @@
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models import Addon from db.models import Addon
from db.session import UnitOfWork
class AddonsRepository: class AddonsRepository:
def __init__(self, session: AsyncSession) -> None: def __init__(self, uow: UnitOfWork) -> None:
self.session = session self.session = uow.session
async def get_all(self) -> list[Addon]: async def get_all(self) -> list[Addon]:
stmt = select(Addon) stmt = select(Addon)

View File

@@ -1,13 +1,14 @@
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models.invoice import Invoice from db.models.invoice import Invoice
from db.session import UnitOfWork
from schemas.invoices import InvoiceStatus from schemas.invoices import InvoiceStatus
class InvoiceRepository: class InvoiceRepository:
def __init__(self, session: AsyncSession) -> None: def __init__(self, uow: UnitOfWork) -> None:
self.session = session self.uow = uow
self.session = uow.session
async def get_by_id(self, id: int) -> Invoice | None: async def get_by_id(self, id: int) -> Invoice | None:
stmt = select(Invoice).where(Invoice.id == id) stmt = select(Invoice).where(Invoice.id == id)
@@ -32,13 +33,9 @@ class InvoiceRepository:
) )
self.session.add(obj) self.session.add(obj)
await self.session.commit()
return obj return obj
async def update_status_by_id(self, invoice_id: int, status: InvoiceStatus) -> Invoice | None: async def update_status_by_id(self, invoice_id: int, status: InvoiceStatus) -> Invoice | None:
invoice = await self.get_by_id(invoice_id) invoice = await self.get_by_id(invoice_id)
invoice.status = status invoice.status = status
await self.session.commit()
return invoice return invoice

View File

@@ -1,14 +1,14 @@
from datetime import datetime from datetime import datetime
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models.link_codes import LinkCode from db.models.link_codes import LinkCode
from db.session import UnitOfWork
from schemas.enums import LinkCodeStatus from schemas.enums import LinkCodeStatus
async def create_link_code( async def create_link_code(
session: AsyncSession, uow: UnitOfWork,
*, *,
code: str, code: str,
user_id: int, user_id: int,
@@ -22,20 +22,17 @@ async def create_link_code(
expires_at=expires_at, expires_at=expires_at,
) )
session.add(link_code) uow.session.add(link_code)
await session.commit()
return 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) stmt = select(LinkCode).where(LinkCode.code == code)
r = await session.execute(stmt) r = await uow.session.execute(stmt)
return r.scalar_one_or_none() 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 code.status = LinkCodeStatus.USED
await session.commit()
return code return code

View File

@@ -1,12 +1,13 @@
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models.orders import Order, OrderAddon, OrderStatus from db.models.orders import Order, OrderAddon, OrderStatus
from db.session import UnitOfWork
class OrderRepository: class OrderRepository:
def __init__(self, session: AsyncSession) -> None: def __init__(self, uow: UnitOfWork) -> None:
self.session = session self.uow = uow
self.session = uow.session
async def create( async def create(
self, self,
@@ -37,7 +38,6 @@ class OrderRepository:
addon = OrderAddon(order_id=order.id, addon_id=addon_id) addon = OrderAddon(order_id=order.id, addon_id=addon_id)
self.session.add(addon) self.session.add(addon)
await self.session.commit()
await self.session.refresh(order, attribute_names=["addons"]) await self.session.refresh(order, attribute_names=["addons"])
return order return order

View File

@@ -1,12 +1,12 @@
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models.pricing import PricingConfig from db.models.pricing import PricingConfig
from db.session import UnitOfWork
class PricingRepository: class PricingRepository:
def __init__(self, session: AsyncSession) -> None: def __init__(self, uow: UnitOfWork) -> None:
self.session = session self.session = uow.session
async def get(self) -> PricingConfig | None: async def get(self) -> PricingConfig | None:
stmt = select(PricingConfig).where(PricingConfig.id == 1) stmt = select(PricingConfig).where(PricingConfig.id == 1)

View File

@@ -1,19 +1,19 @@
from sqlalchemy import func, or_, select, text from sqlalchemy import func, or_, select, text
from sqlalchemy.ext.asyncio import AsyncSession
from db.models.service_notifications import ServiceNotification from db.models.service_notifications import ServiceNotification
from db.session import UnitOfWork
from schemas.enums import NotificationStatus 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) 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() return r.scalar_one_or_none()
async def get_pending_notifications( async def get_pending_notifications(
session: AsyncSession, batch_size: int = 50 uow: UnitOfWork, batch_size: int = 50
) -> list[ServiceNotification]: ) -> list[ServiceNotification]:
stmt = ( stmt = (
select(ServiceNotification) select(ServiceNotification)
@@ -33,26 +33,24 @@ async def get_pending_notifications(
.with_for_update(skip_locked=True) .with_for_update(skip_locked=True)
.limit(batch_size) .limit(batch_size)
) )
r = await session.execute(stmt) r = await uow.session.execute(stmt)
return list(r.scalars().all()) return list(r.scalars().all())
async def ack_notification(session: AsyncSession, n_id: int) -> ServiceNotification | None: async def ack_notification(uow: UnitOfWork, n_id: int) -> ServiceNotification | None:
notification = await get_notification_by_id(session, n_id) notification = await get_notification_by_id(uow, n_id)
if not notification: if not notification:
return return
notification.sent_at = func.now() notification.sent_at = func.now()
notification.status = NotificationStatus.SENT notification.status = NotificationStatus.SENT
await session.commit()
return notification return notification
async def mark_notification_as_dispatched(session: AsyncSession, n_id: int): async def mark_notification_as_dispatched(uow: UnitOfWork, n_id: int):
notification = await get_notification_by_id(session, n_id) notification = await get_notification_by_id(uow, n_id)
if not notification: if not notification:
return return
@@ -60,5 +58,4 @@ async def mark_notification_as_dispatched(session: AsyncSession, n_id: int):
notification.status = NotificationStatus.DISPATCHED notification.status = NotificationStatus.DISPATCHED
notification.attempts += 1 notification.attempts += 1
await session.commit()
return notification return notification

View File

@@ -1,13 +1,14 @@
from sqlalchemy import func, select from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models import Session from db.models import Session
from db.session import UnitOfWork
from schemas.providers import ProvidersType from schemas.providers import ProvidersType
class SessionsRepository: class SessionsRepository:
def __init__(self, session: AsyncSession) -> None: def __init__(self, uow: UnitOfWork) -> None:
self.session = session self.uow = uow
self.session = uow.session
async def get_session_by_id(self, id: int) -> Session | None: async def get_session_by_id(self, id: int) -> Session | None:
stmt = select(Session).where(Session.id == id) 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: 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) obj = Session(user_id=user_id, refresh_token_hash=refresh_token_hash, source=iss)
self.session.add(obj) self.session.add(obj)
await self.session.commit()
return obj return obj
async def revoke(self, token_id: int): async def revoke(self, token_id: int):
session = await self.get_session_by_id(token_id) session = await self.get_session_by_id(token_id)
session.is_revoked = True session.is_revoked = True
session.revoked_at = func.now() session.revoked_at = func.now()
await self.session.commit()
return session return session

View File

@@ -1,13 +1,14 @@
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models import User from db.models import User
from db.models.transactions import BalanceTransaction, BalanceTxType from db.models.transactions import BalanceTransaction, BalanceTxType
from db.session import UnitOfWork
class UserRepository: class UserRepository:
def __init__(self, session: AsyncSession) -> None: def __init__(self, uow: UnitOfWork) -> None:
self.session = session self.uow = uow
self.session = uow.session
async def get_user_by_id(self, id: int) -> User | None: async def get_user_by_id(self, id: int) -> User | None:
stmt = select(User).where(User.id == id) stmt = select(User).where(User.id == id)
@@ -44,8 +45,6 @@ class UserRepository:
referal_id=referal_id, referal_id=referal_id,
) )
self.session.add(obj) self.session.add(obj)
await self.session.commit()
return obj return obj
async def increase_balance( async def increase_balance(
@@ -67,11 +66,8 @@ class UserRepository:
self.session.add(obj) self.session.add(obj)
user.balance += amount user.balance += amount
await self.session.commit()
return user return user
async def update_telegram_id(self, user: User, telegram_id: int) -> User: async def update_telegram_id(self, user: User, telegram_id: int) -> User:
user.telegram_id = telegram_id user.telegram_id = telegram_id
await self.session.commit()
return user return user

View File

@@ -1,6 +1,5 @@
from fastapi import APIRouter, Depends, HTTPException from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
from sqlalchemy.ext.asyncio import AsyncSession
from core.secrets import ( from core.secrets import (
estimate_password_strength, estimate_password_strength,
@@ -8,7 +7,7 @@ from core.secrets import (
hash_refresh_token, hash_refresh_token,
verify_password, verify_password,
) )
from db.session import get_db from db.session import UnitOfWork, get_uow
from repositories.sessions import SessionsRepository from repositories.sessions import SessionsRepository
from repositories.users import UserRepository from repositories.users import UserRepository
from schemas.login import UserLogin, UserLoginData, UserTokens from schemas.login import UserLogin, UserLoginData, UserTokens
@@ -22,8 +21,8 @@ router = APIRouter(prefix="/auth")
@router.post("/signup") @router.post("/signup")
async def signup(req: UserRegistration, session: AsyncSession = Depends(get_db)): async def signup(req: UserRegistration, uow: UnitOfWork = Depends(get_uow)):
users_repo = UserRepository(session) users_repo = UserRepository(uow)
if req.provider == "credentials": if req.provider == "credentials":
if not req.username or not req.password: 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( user = await users_repo.create(
username=req.username, hashed_password=password_hash, referal_id=referal_id username=req.username, hashed_password=password_hash, referal_id=referal_id
) )
await uow.commit()
return JSONResponse( return JSONResponse(
UserInfo( UserInfo(
username=user.username, username=user.username,
@@ -57,9 +57,9 @@ async def signup(req: UserRegistration, session: AsyncSession = Depends(get_db))
@router.post("/login", response_model=UserLogin) @router.post("/login", response_model=UserLogin)
async def login(req: UserLoginData, session: AsyncSession = Depends(get_db)): async def login(req: UserLoginData, uow: UnitOfWork = Depends(get_uow)):
users_repo = UserRepository(session) users_repo = UserRepository(uow)
sessions_repo = SessionsRepository(session) sessions_repo = SessionsRepository(uow)
if req.provider == "credentials": if req.provider == "credentials":
if not req.username or not req.password: if not req.username or not req.password:
raise HTTPException(status_code=400, detail="Username or password is not provided.") 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") raise HTTPException(status_code=401, detail="Invalid password")
data = await authorize_user(sessions_repo, user, req.provider) data = await authorize_user(sessions_repo, user, req.provider)
await uow.commit()
return data return data
if req.provider == "telegram": if req.provider == "telegram":
@@ -82,8 +83,8 @@ async def login(req: UserLoginData, session: AsyncSession = Depends(get_db)):
@router.post("/refresh", response_model=UserTokens) @router.post("/refresh", response_model=UserTokens)
async def refresh(refresh_token: str, iss: ProvidersType, session: AsyncSession = Depends(get_db)): async def refresh(refresh_token: str, iss: ProvidersType, uow: UnitOfWork = Depends(get_uow)):
sessions_repo = SessionsRepository(session) sessions_repo = SessionsRepository(uow)
token_hash = hash_refresh_token(refresh_token) token_hash = hash_refresh_token(refresh_token)
token_entry = await sessions_repo.get_session_by_hash(token_hash) 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.") raise HTTPException(status_code=401, detail="Refresh token is invalid.")
key_pair = await refresh_token_rotation(sessions_repo, token_entry, iss) key_pair = await refresh_token_rotation(sessions_repo, token_entry, iss)
await uow.commit()
return UserTokens( return UserTokens(
access_token=key_pair.access_token, access_token=key_pair.access_token,
refresh_token=key_pair.refresh_token, refresh_token=key_pair.refresh_token,

View File

@@ -1,8 +1,7 @@
from fastapi import APIRouter, Depends, HTTPException from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.ext.asyncio import AsyncSession
from core.deps import get_service_identity 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 ( from repositories.service_notifications import (
ack_notification, ack_notification,
get_pending_notifications, get_pending_notifications,
@@ -22,9 +21,9 @@ router = APIRouter(prefix="/renewal")
async def get_pending( async def get_pending(
limit: int = 50, limit: int = 50,
ctx: ServiceIdentity = Depends(get_service_identity), 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] = [] response_users: list[UserNotificationData] = []
for notification in notifications: 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) return NotificationResponse(users=response_users, issued_by=ctx.service)
@@ -47,13 +47,14 @@ async def get_pending(
async def acknowledge( async def acknowledge(
req: NotificationAcknowledgeRequest, req: NotificationAcknowledgeRequest,
ctx: ServiceIdentity = Depends(get_service_identity), ctx: ServiceIdentity = Depends(get_service_identity),
session: AsyncSession = Depends(get_db), uow: UnitOfWork = Depends(get_uow),
): ):
try: try:
r = await ack_notification(session, req.notification_id) r = await ack_notification(uow, req.notification_id)
except Exception as e: except Exception as e:
raise HTTPException(500, detail=str(e)) from None raise HTTPException(500, detail=str(e)) from None
if r: if r:
await uow.commit()
return "OK" return "OK"
raise HTTPException(500, detail="No such notification found.") raise HTTPException(500, detail="No such notification found.")

View File

@@ -2,11 +2,10 @@ import secrets
from datetime import UTC, datetime, timedelta from datetime import UTC, datetime, timedelta
from fastapi import APIRouter, Depends, HTTPException from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.ext.asyncio import AsyncSession
from config import cfg from config import cfg
from core.deps import get_auth_context, get_service_identity 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.link_codes import create_link_code, get_link_code_by_code, use_link_code
from repositories.users import UserRepository from repositories.users import UserRepository
from schemas.dto import AuthContext from schemas.dto import AuthContext
@@ -19,13 +18,14 @@ router = APIRouter(prefix="/link-codes")
@router.post("", response_model=LinkCodeResponse, status_code=201) @router.post("", response_model=LinkCodeResponse, status_code=201)
async def gen_link_code( 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) code = secrets.token_urlsafe(cfg.link_code_length)
exp = datetime.now(UTC) + timedelta(minutes=cfg.link_code_ttl) exp = datetime.now(UTC) + timedelta(minutes=cfg.link_code_ttl)
link_code = await create_link_code( 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) 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( async def consume_link_code(
payload: LinkCodeConsume, payload: LinkCodeConsume,
ctx: AuthContext = Depends(get_service_identity), 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: if not link_code:
raise HTTPException(404, detail="Code not found") raise HTTPException(404, detail="Code not found")
@@ -44,7 +44,7 @@ async def consume_link_code(
if link_code.status != LinkCodeStatus.ACTIVE: if link_code.status != LinkCodeStatus.ACTIVE:
raise HTTPException(404, detail="Code expired or is invalid.") 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) user = await users_repo.get_user_by_id(link_code.user_id)
if not user: if not user:
@@ -52,7 +52,8 @@ async def consume_link_code(
user = await users_repo.update_telegram_id(user, payload.telegram_id) 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( return UserInfo(
username=user.username, username=user.username,

View File

@@ -2,12 +2,11 @@ import math
from datetime import UTC, datetime from datetime import UTC, datetime
from fastapi import APIRouter, Depends, HTTPException from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.ext.asyncio import AsyncSession
from config import cfg from config import cfg
from core.deps import get_auth_context, get_pally_client from core.deps import get_auth_context, get_pally_client
from db.models.orders import OrderStatus 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 external.pally import PallyClient
from repositories import AddonsRepository, PricingRepository from repositories import AddonsRepository, PricingRepository
from repositories.invoices import InvoiceRepository from repositories.invoices import InvoiceRepository
@@ -31,13 +30,13 @@ router = APIRouter(prefix="/orders")
async def checkout( async def checkout(
order: OrderDetails, order: OrderDetails,
ctx: AuthContext = Depends(get_auth_context), ctx: AuthContext = Depends(get_auth_context),
session: AsyncSession = Depends(get_db), uow: UnitOfWork = Depends(get_uow),
pally: PallyClient = Depends(get_pally_client), pally: PallyClient = Depends(get_pally_client),
): ):
addons_repo = AddonsRepository(session) addons_repo = AddonsRepository(uow)
pricing_repo = PricingRepository(session) pricing_repo = PricingRepository(uow)
invoices_repo = InvoiceRepository(session) invoices_repo = InvoiceRepository(uow)
orders_repo = OrderRepository(session) orders_repo = OrderRepository(uow)
pricing = await get_pricing_model(addons_repo, pricing_repo) pricing = await get_pricing_model(addons_repo, pricing_repo)
price = math.ceil(await calculate_price(addons_repo=addons_repo, order=order, pricing=pricing)) 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") raise HTTPException(500, detail="failed to create invoice")
else: else:
await deduct_order_balance( await deduct_order_balance(
session, uow.session,
user=ctx.user, user=ctx.user,
order=order_entry, order=order_entry,
description=f"order {order_entry.id} paid from balance", description=f"order {order_entry.id} paid from balance",
@@ -85,7 +84,7 @@ async def checkout(
now=now, now=now,
): ):
await apply_order_now( await apply_order_now(
session, uow.session,
user=ctx.user, user=ctx.user,
order=order_entry, order=order_entry,
pricing=pricing, pricing=pricing,
@@ -95,9 +94,12 @@ async def checkout(
await queue_order_for_later( await queue_order_for_later(
order=order_entry, subscription=ctx.user.subscription, now=now order=order_entry, subscription=ctx.user.subscription, now=now
) )
await session.commit() await uow.commit()
payment_link = None payment_link = None
if amount_to_pay > 0:
await uow.commit()
return CheckoutResponse( return CheckoutResponse(
order_id=str(order_entry.id), order_id=str(order_entry.id),
total_amount=price, total_amount=price,

View File

@@ -3,10 +3,9 @@ import logging
from fastapi import Depends, Form, HTTPException from fastapi import Depends, Form, HTTPException
from fastapi.routing import APIRouter from fastapi.routing import APIRouter
from sqlalchemy.ext.asyncio import AsyncSession
from core.deps import get_db from db.models.transactions import BalanceTxType
from db.models.transactions import BalanceTransaction, BalanceTxType from db.session import UnitOfWork, get_uow
from external.pally import BillStatus from external.pally import BillStatus
from repositories.invoices import InvoiceRepository from repositories.invoices import InvoiceRepository
from repositories.users import UserRepository from repositories.users import UserRepository
@@ -39,7 +38,7 @@ async def pally_callback(
PayerComment: str | None = Form(None), PayerComment: str | None = Form(None),
ErrorCode: int | None = Form(None), ErrorCode: int | None = Form(None),
ErrorMessage: str | None = Form(None), ErrorMessage: str | None = Form(None),
session: AsyncSession = Depends(get_db), uow: UnitOfWork = Depends(get_uow),
): ):
logger.info( logger.info(
"Pally webhook received - InvId: %s, OutSum: %s, Commission: %s, TrsId: %s, Status: %s, " "Pally webhook received - InvId: %s, OutSum: %s, Commission: %s, TrsId: %s, Status: %s, "
@@ -88,7 +87,7 @@ async def pally_callback(
try: try:
await process_subscription_purchase( await process_subscription_purchase(
session, uow,
invoice_id=int(invoice_id_str), invoice_id=int(invoice_id_str),
trs_id=TrsId, trs_id=TrsId,
amount=amount, amount=amount,
@@ -102,9 +101,9 @@ async def pally_callback(
invoice_id_str, invoice_id_str,
TrsId, 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)) invoice = await invoice_repo.get_by_id(int(invoice_id_str))
if invoice is None: if invoice is None:
logger.critical( logger.critical(
@@ -117,7 +116,8 @@ async def pally_callback(
if invoice.status == InvoiceStatus.PAID: if invoice.status == InvoiceStatus.PAID:
return "OK" 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: if user is None:
logger.critical( logger.critical(
"Cannot credit fallback balance: user %s was not found " "for bill %s (TrsId: %s)", "Cannot credit fallback balance: user %s was not found " "for bill %s (TrsId: %s)",
@@ -127,22 +127,14 @@ async def pally_callback(
) )
raise raise
balance_before = user.balance await users_repo.increase_balance(
user.balance += amount user.id,
invoice.status = InvoiceStatus.PAID amount,
session.add( BalanceTxType.DEPOSIT,
BalanceTransaction( f"fallback payment credit for invoice {invoice.id} (TrsId: {TrsId})",
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 session.commit() await invoice_repo.update_status_by_id(invoice.id, InvoiceStatus.PAID)
await uow.commit()
logger.info( logger.info(
"Fallback payment credit processed: user_id=%s, amount=%s, " "invoice_id=%s, TrsId=%s", "Fallback payment credit processed: user_id=%s, amount=%s, " "invoice_id=%s, TrsId=%s",
user.id, user.id,

View File

@@ -1,7 +1,6 @@
from fastapi import APIRouter, Depends 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.addons import AddonsRepository
from repositories.pricing import PricingRepository from repositories.pricing import PricingRepository
from schemas.plans import PricingPlans from schemas.plans import PricingPlans
@@ -11,9 +10,9 @@ router = APIRouter(prefix="/plans")
@router.get("/", response_model=PricingPlans) @router.get("/", response_model=PricingPlans)
async def get_plans(session: AsyncSession = Depends(get_db)): async def get_plans(uow: UnitOfWork = Depends(get_uow)):
addons_repo = AddonsRepository(session) addons_repo = AddonsRepository(uow)
pricing_repo = PricingRepository(session) pricing_repo = PricingRepository(uow)
res = await get_pricing_model(addons_repo, pricing_repo) res = await get_pricing_model(addons_repo, pricing_repo)
return res return res

View File

@@ -1,22 +1,17 @@
import asyncio import asyncio
from datetime import UTC, datetime, timedelta from datetime import UTC, datetime, timedelta
from typing import TYPE_CHECKING
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.dialects.postgresql import insert from sqlalchemy.dialects.postgresql import insert
from config import cfg from config import cfg
from db.models import ServiceNotification, Subscription from db.models import ServiceNotification, Subscription
from db.session import async_session from db.session import UnitOfWork, async_session
from schemas.enums import NotificationType from schemas.enums import NotificationType
if TYPE_CHECKING:
from sqlalchemy.ext.asyncio import AsyncSession
async def collect_subscription_notifications(uow: UnitOfWork, now: datetime | None = None) -> int:
async def collect_subscription_notifications( session = uow.session
session: "AsyncSession", now: datetime | None = None
) -> int:
now = now or datetime.now(UTC) now = now or datetime.now(UTC)
periods = ( periods = (
(NotificationType.SEVEN_DAYS, now + timedelta(days=3), now + timedelta(days=7)), (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) result = await session.execute(stmt)
created += result.rowcount or 0 created += result.rowcount or 0
await session.commit() await uow.commit()
return created return created
@@ -56,8 +51,8 @@ async def run_subscription_notifications() -> None:
interval = cfg.notification_scan_interval * 60 interval = cfg.notification_scan_interval * 60
while True: while True:
try: try:
async with async_session() as session: async with async_session() as session, UnitOfWork(session) as uow:
await collect_subscription_notifications(session) await collect_subscription_notifications(uow)
except Exception: except Exception:
# A failed iteration must not stop subsequent notification checks. # A failed iteration must not stop subsequent notification checks.
pass pass

View File

@@ -6,7 +6,8 @@ from datetime import UTC, datetime
from config import cfg from config import cfg
from db.models.orders import OrderStatus 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 import AddonsRepository
from repositories.invoices import InvoiceRepository from repositories.invoices import InvoiceRepository
from repositories.orders import OrderRepository 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 async def process_subscription_purchase( # noqa: PLR0911, PLR0912
session, uow: UnitOfWork,
*, *,
invoice_id: int, invoice_id: int,
trs_id: str, trs_id: str,
amount: int, amount: int,
) -> None: ) -> None:
invoice_repo = InvoiceRepository(session) session = uow.session
orders_repo = OrderRepository(session) invoice_repo = InvoiceRepository(uow)
users_repo = UserRepository(session) orders_repo = OrderRepository(uow)
pricing_repo = PricingRepository(session) users_repo = UserRepository(uow)
pricing_repo = PricingRepository(uow)
invoice = await invoice_repo.get_by_id(invoice_id) invoice = await invoice_repo.get_by_id(invoice_id)
@@ -99,7 +101,7 @@ async def process_subscription_purchase( # noqa: PLR0911, PLR0912
subscription = user.subscription subscription = user.subscription
now = datetime.now(UTC) now = datetime.now(UTC)
pricing = await get_pricing_model(AddonsRepository(session), pricing_repo) pricing = await get_pricing_model(AddonsRepository(uow), pricing_repo)
logger.info( logger.info(
"Processing payment: bill_id=%s, order_id=%s, user_id=%s, amount=%s", "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: if referal_id is not None:
referal_amount = math.floor(amount * (cfg.referal_bonus / 100)) referal_amount = math.floor(amount * (cfg.referal_bonus / 100))
if referal_amount > 0: 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: if referal_user is not None:
session.add( await users_repo.increase_balance(
BalanceTransaction( referal_id,
user_id=referal_id, referal_amount,
amount=referal_amount, BalanceTxType.REFERRAL_BONUS,
tx_type=BalanceTxType.REFERRAL_BONUS, f"referral reward for user {invoice.creator_id} (TrsId: {trs_id})",
balance_before=referal_user.balance,
balance_after=referal_user.balance + referal_amount,
description=f"referral reward for user {invoice.creator_id} (TrsId: {trs_id})",
)
) )
referal_user.balance += referal_amount
logger.info( logger.info(
"Referral bonus processed: referrer_id=%s, amount=%s", "Referral bonus processed: referrer_id=%s, amount=%s",
@@ -169,6 +166,6 @@ async def process_subscription_purchase( # noqa: PLR0911, PLR0912
trs_id, trs_id,
) )
await session.commit() await uow.commit()
logger.info("Bill %s marked as PAID for TrsId %s", invoice_id, trs_id) logger.info("Bill %s marked as PAID for TrsId %s", invoice_id, trs_id)

View File

@@ -8,7 +8,7 @@ from sqlalchemy.dialects.postgresql import insert
from config import cfg from config import cfg
from db.models import RWSyncOutbox, Subscription 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 from external.rw import sync_subscription_by_telegram_id
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -48,7 +48,8 @@ async def enqueue_all_rw_syncs(session: "AsyncSession") -> int:
return count 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) now = datetime.now(UTC)
jobs = await session.scalars( jobs = await session.scalars(
select(RWSyncOutbox) select(RWSyncOutbox)
@@ -63,7 +64,7 @@ async def process_rw_syncs(session: "AsyncSession", batch_size: int = 50) -> int
jobs = list(jobs) jobs = list(jobs)
for job in jobs: for job in jobs:
job.locked_until = now + LOCK_DURATION job.locked_until = now + LOCK_DURATION
await session.commit() await uow.commit()
for job in jobs: for job in jobs:
subscription = job.subscription subscription = job.subscription
@@ -81,7 +82,7 @@ async def process_rw_syncs(session: "AsyncSession", batch_size: int = 50) -> int
RWSyncOutbox.revision == job.revision, RWSyncOutbox.revision == job.revision,
) )
) )
await session.commit() await uow.commit()
continue continue
delay = min(timedelta(minutes=2**job.attempts), MAX_RETRY_DELAY) 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, locked_until=None,
) )
) )
await session.commit() await uow.commit()
logger.warning("RW sync failed for subscription_id=%s", subscription.id) logger.warning("RW sync failed for subscription_id=%s", subscription.id)
return len(jobs) return len(jobs)
@@ -104,8 +105,8 @@ async def run_rw_sync_worker() -> None:
interval = cfg.rw_sync_interval_minutes * 60 interval = cfg.rw_sync_interval_minutes * 60
while True: while True:
try: try:
async with async_session() as session: async with async_session() as session, UnitOfWork(session) as uow:
await process_rw_syncs(session) await process_rw_syncs(uow)
except Exception: except Exception:
logger.exception("RW sync worker iteration failed") logger.exception("RW sync worker iteration failed")
await asyncio.sleep(interval) await asyncio.sleep(interval)
@@ -116,8 +117,9 @@ async def run_rw_sync_reconciler() -> None:
while True: while True:
try: try:
async with async_session() as session: async with async_session() as session:
count = await enqueue_all_rw_syncs(session) async with UnitOfWork(session) as uow:
await session.commit() count = await enqueue_all_rw_syncs(uow.session)
await uow.commit()
logger.info("Queued %s subscriptions for RW reconciliation", count) logger.info("Queued %s subscriptions for RW reconciliation", count)
except Exception: except Exception:
logger.exception("RW reconciliation iteration failed") logger.exception("RW reconciliation iteration failed")

View File

@@ -6,8 +6,9 @@ from schemas.enums import NotificationType
from services.notifications import collect_subscription_notifications from services.notifications import collect_subscription_notifications
class FakeSession: class FakeUnitOfWork:
def __init__(self, subscriptions_by_period): def __init__(self, subscriptions_by_period):
self.session = self
self.subscriptions_by_period = iter(subscriptions_by_period) self.subscriptions_by_period = iter(subscriptions_by_period)
self.executed = [] self.executed = []
self.committed = False self.committed = False
@@ -25,7 +26,7 @@ class FakeSession:
def test_collects_notification_for_each_expiry_period(): def test_collects_notification_for_each_expiry_period():
now = datetime(2026, 8, 20, tzinfo=UTC) now = datetime(2026, 8, 20, tzinfo=UTC)
session = FakeSession( uow = FakeUnitOfWork(
[ [
[SimpleNamespace(id=1, expires_at=now + timedelta(days=6))], [SimpleNamespace(id=1, expires_at=now + timedelta(days=6))],
[SimpleNamespace(id=2, expires_at=now + timedelta(days=2))], [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 created == len(uow.executed)
assert session.committed assert uow.committed
assert [statement.compile().params["notify_type_m0"] for statement in session.executed] == [ assert [statement.compile().params["notify_type_m0"] for statement in uow.executed] == [
NotificationType.SEVEN_DAYS, NotificationType.SEVEN_DAYS,
NotificationType.THREE_DAYS, NotificationType.THREE_DAYS,
NotificationType.ONE_DAY, NotificationType.ONE_DAY,