feat: referal spec on registration

This commit is contained in:
2026-08-03 10:35:26 +07:00
parent 680c15414c
commit 801014c392
8 changed files with 88 additions and 5 deletions

View File

@@ -0,0 +1,34 @@
"""+user.referal_code
Revision ID: 135672cdf14a
Revises: 72e78a8a43cf
Create Date: 2026-08-03 10:26:22.290547
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = '135672cdf14a'
down_revision: Union[str, Sequence[str], None] = '72e78a8a43cf'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
"""Upgrade schema."""
# ### commands auto generated by Alembic - please adjust! ###
op.add_column('users', sa.Column('referal_code', sa.TEXT(), nullable=False))
op.create_unique_constraint(None, 'users', ['referal_code'])
# ### end Alembic commands ###
def downgrade() -> None:
"""Downgrade schema."""
# ### commands auto generated by Alembic - please adjust! ###
op.drop_constraint(None, 'users', type_='unique')
op.drop_column('users', 'referal_code')
# ### end Alembic commands ###

View File

@@ -1,3 +1,5 @@
import secrets
import string
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from sqlalchemy import BIGINT, INTEGER, REAL, TEXT, VARCHAR, ForeignKey from sqlalchemy import BIGINT, INTEGER, REAL, TEXT, VARCHAR, ForeignKey
@@ -9,6 +11,11 @@ if TYPE_CHECKING:
from db.models import Order, Session, Subscription from db.models import Order, Session, Subscription
def generate_ref_code(length: int = 8):
alphabet = string.ascii_letters + string.digits
return "".join(secrets.choice(alphabet) for _ in range(length))
class User(Base): class User(Base):
__tablename__ = "users" __tablename__ = "users"
@@ -19,6 +26,9 @@ class User(Base):
hashed_password: Mapped[str] = mapped_column(VARCHAR(255), nullable=True) hashed_password: Mapped[str] = mapped_column(VARCHAR(255), nullable=True)
telegram_id: Mapped[int] = mapped_column(BIGINT, unique=True, nullable=True) telegram_id: Mapped[int] = mapped_column(BIGINT, unique=True, nullable=True)
referal_id: Mapped[int] = mapped_column(ForeignKey("users.id"), nullable=True) referal_id: Mapped[int] = mapped_column(ForeignKey("users.id"), nullable=True)
referal_code: Mapped[str] = mapped_column(
TEXT, nullable=False, unique=True, default=generate_ref_code
)
balance: Mapped[float] = mapped_column(REAL, nullable=False, default=0) balance: Mapped[float] = mapped_column(REAL, nullable=False, default=0)

View File

@@ -24,17 +24,24 @@ class UserRepository:
res = await self.session.execute(stmt) res = await self.session.execute(stmt)
return res.scalar_one_or_none() return res.scalar_one_or_none()
async def get_user_by_ref_code(self, ref_code: str) -> User | None:
stmt = select(User).where(User.referal_code == ref_code)
res = await self.session.execute(stmt)
return res.scalar_one_or_none()
async def create( async def create(
self, self,
*, *,
username: str | None = None, username: str | None = None,
hashed_password: str | None = None, hashed_password: str | None = None,
telegram_id: int | None = None, telegram_id: int | None = None,
referal_id: int | None = None,
) -> User: ) -> User:
obj = User( obj = User(
username=username, username=username,
hashed_password=hashed_password, hashed_password=hashed_password,
telegram_id=telegram_id, telegram_id=telegram_id,
referal_id=referal_id,
) )
self.session.add(obj) self.session.add(obj)
await self.session.commit() await self.session.commit()

View File

@@ -28,9 +28,21 @@ async def signup(req: UserRegistration, session: AsyncSession = Depends(get_db))
raise HTTPException(status_code=409, detail="User already exists") raise HTTPException(status_code=409, detail="User already exists")
password_hash = hash_password(req.password) password_hash = hash_password(req.password)
user = await users_repo.create(username=req.username, hashed_password=password_hash) referal_id = None
if req.referal_code:
referal = await users_repo.get_user_by_ref_code(req.referal_code)
referal_id = referal.id if referal else None
user = await users_repo.create(
username=req.username, hashed_password=password_hash, referal_id=referal_id
)
return JSONResponse( return JSONResponse(
UserInfo(username=user.username, telegram_id=user.telegram_id).model_dump(), UserInfo(
username=user.username,
telegram_id=user.telegram_id,
subscription=None,
referal_code=user.referal_code,
).model_dump(),
status_code=201, status_code=201,
) )
raise HTTPException(status_code=400, detail="Unsupported provider") raise HTTPException(status_code=400, detail="Unsupported provider")

View File

@@ -15,5 +15,8 @@ async def get_me(ctx: AuthContext = Depends(get_auth_context)):
addon_ids=[a.id for a in ctx.user.subscription.addons], addon_ids=[a.id for a in ctx.user.subscription.addons],
) )
return UserInfo( return UserInfo(
username=ctx.user.username, telegram_id=ctx.user.telegram_id, subscription=sub_data username=ctx.user.username,
telegram_id=ctx.user.telegram_id,
subscription=sub_data,
referal_code=ctx.user.referal_code,
) )

View File

@@ -9,4 +9,6 @@ class UserRegistration(BaseModel):
username: str | None = Field(None) username: str | None = Field(None)
password: str | None = Field(None) password: str | None = Field(None)
referal_code: str | None = Field(None)
provider: ProvidersType provider: ProvidersType

View File

@@ -13,3 +13,4 @@ class UserInfo(BaseModel):
username: str | None = Field(None) username: str | None = Field(None)
telegram_id: str | None = Field(None) telegram_id: str | None = Field(None)
subscription: SubscriptionData | None = Field(None) subscription: SubscriptionData | None = Field(None)
referal_code: str = Field()

View File

@@ -3,7 +3,7 @@ from db.models.users import User
from repositories.sessions import SessionsRepository from repositories.sessions import SessionsRepository
from schemas.login import UserLogin from schemas.login import UserLogin
from schemas.providers import ProvidersType from schemas.providers import ProvidersType
from schemas.user import UserInfo from schemas.user import SubscriptionData, UserInfo
async def authorize_user( async def authorize_user(
@@ -15,9 +15,23 @@ async def authorize_user(
await sessions_repo.create(user_id=user.id, refresh_token_hash=refresh_token_hash, iss=iss) await sessions_repo.create(user_id=user.id, refresh_token_hash=refresh_token_hash, iss=iss)
sub = (
SubscriptionData(
devices=user.subscription.devices,
expires_at=user.subscription.expires_at,
addon_ids=[a.id for a in user.subscription.addons],
)
if user.subscription
else None
)
return UserLogin( return UserLogin(
access_token=key_pair.access_token, access_token=key_pair.access_token,
refresh_token=key_pair.refresh_token, refresh_token=key_pair.refresh_token,
user=UserInfo(username=user.username, telegram_id=user.telegram_id), user=UserInfo(
username=user.username,
telegram_id=user.telegram_id,
referal_code=user.referal_code,
subscription=sub,
),
expires_at=key_pair.expires_at, expires_at=key_pair.expires_at,
) )