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 sqlalchemy import BIGINT, INTEGER, REAL, TEXT, VARCHAR, ForeignKey
@@ -9,6 +11,11 @@ if TYPE_CHECKING:
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):
__tablename__ = "users"
@@ -19,6 +26,9 @@ class User(Base):
hashed_password: Mapped[str] = mapped_column(VARCHAR(255), 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_code: Mapped[str] = mapped_column(
TEXT, nullable=False, unique=True, default=generate_ref_code
)
balance: Mapped[float] = mapped_column(REAL, nullable=False, default=0)

View File

@@ -24,17 +24,24 @@ class UserRepository:
res = await self.session.execute(stmt)
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(
self,
*,
username: str | None = None,
hashed_password: str | None = None,
telegram_id: int | None = None,
referal_id: int | None = None,
) -> User:
obj = User(
username=username,
hashed_password=hashed_password,
telegram_id=telegram_id,
referal_id=referal_id,
)
self.session.add(obj)
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")
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(
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,
)
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],
)
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)
password: str | None = Field(None)
referal_code: str | None = Field(None)
provider: ProvidersType

View File

@@ -13,3 +13,4 @@ class UserInfo(BaseModel):
username: str | None = Field(None)
telegram_id: str | 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 schemas.login import UserLogin
from schemas.providers import ProvidersType
from schemas.user import UserInfo
from schemas.user import SubscriptionData, UserInfo
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)
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(
access_token=key_pair.access_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,
)