fix!: removed .commit() from repository level
This commit is contained in:
@@ -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])
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
15
core/deps.py
15
core/deps.py
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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.")
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user