diff --git a/alembic/versions/135672cdf14a_user_referal_code.py b/alembic/versions/135672cdf14a_user_referal_code.py new file mode 100644 index 0000000..c92208e --- /dev/null +++ b/alembic/versions/135672cdf14a_user_referal_code.py @@ -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 ### diff --git a/db/models/users.py b/db/models/users.py index 3aa68c5..d84af5d 100644 --- a/db/models/users.py +++ b/db/models/users.py @@ -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) diff --git a/repositories/users.py b/repositories/users.py index f63299b..28c271b 100644 --- a/repositories/users.py +++ b/repositories/users.py @@ -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() diff --git a/routes/auth.py b/routes/auth.py index 7419d70..5eb61b0 100644 --- a/routes/auth.py +++ b/routes/auth.py @@ -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") diff --git a/routes/users/me.py b/routes/users/me.py index 53000f3..7478f83 100644 --- a/routes/users/me.py +++ b/routes/users/me.py @@ -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, ) diff --git a/schemas/registration.py b/schemas/registration.py index 0a2a3eb..cb2fcc8 100644 --- a/schemas/registration.py +++ b/schemas/registration.py @@ -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 diff --git a/schemas/user.py b/schemas/user.py index 464980f..45edb04 100644 --- a/schemas/user.py +++ b/schemas/user.py @@ -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() diff --git a/services/users.py b/services/users.py index 69c8b3f..12cfc81 100644 --- a/services/users.py +++ b/services/users.py @@ -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, )