diff --git a/alembic/versions/6d1bae3dc723_link_codes.py b/alembic/versions/6d1bae3dc723_link_codes.py new file mode 100644 index 0000000..eb5aae4 --- /dev/null +++ b/alembic/versions/6d1bae3dc723_link_codes.py @@ -0,0 +1,44 @@ +"""+link_codes + +Revision ID: 6d1bae3dc723 +Revises: 551c0ad261cd +Create Date: 2026-08-18 20:19:35.961227 + +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision: str = '6d1bae3dc723' +down_revision: Union[str, Sequence[str], None] = '551c0ad261cd' +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.create_table('link_codes', + sa.Column('code', sa.TEXT(), nullable=False), + sa.Column('user_id', sa.INTEGER(), nullable=False), + sa.Column('status', sa.Enum('ACTIVE', 'USED', 'EXPIRED', name='linkcodestatus'), nullable=False), + sa.Column('created_at', sa.DateTime(timezone=True), nullable=False), + sa.Column('expires_at', sa.DateTime(timezone=True), nullable=False), + sa.Column('used_at', sa.DateTime(timezone=True), nullable=True), + sa.ForeignKeyConstraint(['user_id'], ['users.id'], ), + sa.PrimaryKeyConstraint('code'), + sa.UniqueConstraint('code') + ) + op.create_unique_constraint(None, 'service_signatures', ['kid']) + # ### end Alembic commands ### + + +def downgrade() -> None: + """Downgrade schema.""" + # ### commands auto generated by Alembic - please adjust! ### + op.drop_constraint(None, 'service_signatures', type_='unique') + op.drop_table('link_codes') + # ### end Alembic commands ### diff --git a/config.py b/config.py index 98505f1..114fffa 100644 --- a/config.py +++ b/config.py @@ -7,19 +7,15 @@ from pydantic_settings import BaseSettings, SettingsConfigDict class Settings(BaseSettings): model_config = SettingsConfigDict(env_file=".env") + ### Internal settings ### postgres_user: str = Field() postgres_password: str = Field() postgres_host: str = Field() postgres_port: str = Field() postgres_db: str = Field() - private_key_fp: str = Field() - public_key_fp: str = Field() - - access_token_ttl: int = Field(description="Access token TTL (minutes)") - - pally_shop_id: str = Field() - pally_token: str = Field() + ### Telegram ### + bot_username: str = Field() ### Remnawave ### remnawave_base_url: str = Field() @@ -27,12 +23,23 @@ class Settings(BaseSettings): remnawave_token: str = Field() remnawave_default_squads_raw: str = Field(alias="REMNAWAVE_DEFAULT_SQUADS_UUIDS") - ### + ### Finance ### minimal_deposit: int = Field() referal_bonus: int = Field(30) + #### Pally #### + pally_shop_id: str = Field() + pally_token: str = Field() + ### Security related ### + private_key_fp: str = Field() + public_key_fp: str = Field() + + access_token_ttl: int = Field(description="Access token TTL (minutes)") + link_code_ttl: int = Field(8, description="Link Code TTL (minutes)") + min_password_length: int = Field(8) + link_code_length: int = Field(8) password_security_threshold: int = Field(2) @computed_field @@ -61,5 +68,10 @@ class Settings(BaseSettings): with open(self.public_key_fp, "rb") as f: return f.read() + @computed_field + @property + def bot_url(self) -> str: + return "https://t.me/" + self.bot_username.lstrip("@") + cfg = Settings() # type: ignore diff --git a/core/deps.py b/core/deps.py index e3239c7..daefab1 100644 --- a/core/deps.py +++ b/core/deps.py @@ -3,9 +3,11 @@ from sqlalchemy.ext.asyncio import AsyncSession from config import cfg from core.auth import jwt +from core.secrets import get_kid_from_token from db.session import get_db from external.pally import PallyClient -from schemas.dto import AuthContext +from repositories.service_signatures import get_active_signature_by_kid +from schemas.dto import AuthContext, ServiceIdentity from services.subscriptions import sync_user_subscription @@ -25,5 +27,21 @@ async def get_auth_context( return ctx +async def get_service_identity( + request: Request, session: AsyncSession = Depends(get_db) +) -> ServiceIdentity | None: + auth = request.headers.get("Authorization") + + if auth.startswith("Bearer"): + token = auth.removeprefix("Bearer ").strip() + kid = get_kid_from_token(token) + if not kid: + raise HTTPException(401, detail="No kid provided.") + signature = await get_active_signature_by_kid(session, kid) + if not signature: + raise HTTPException(401, detail="Invalid signature") + return ServiceIdentity(service=signature.kid) + + def get_pally_client() -> PallyClient: return PallyClient(api_token=cfg.pally_token, payer_pays_commission=True) diff --git a/db/models/__init__.py b/db/models/__init__.py index cb3fbbb..98b1087 100644 --- a/db/models/__init__.py +++ b/db/models/__init__.py @@ -1,5 +1,6 @@ from .addons import Addon from .invoice import Invoice +from .link_codes import LinkCode from .orders import Order, OrderAddon from .pricing import PricingConfig from .service_signatures import ServiceSignature @@ -13,6 +14,7 @@ __all__ = [ "Addon", "BalanceTransaction", "Invoice", + "LinkCode", "Order", "OrderAddon", "PricingConfig", diff --git a/db/models/link_codes.py b/db/models/link_codes.py new file mode 100644 index 0000000..f12ca56 --- /dev/null +++ b/db/models/link_codes.py @@ -0,0 +1,28 @@ +from datetime import UTC, datetime + +from sqlalchemy import TEXT, DateTime, Enum, ForeignKey +from sqlalchemy.orm import Mapped, mapped_column + +from db.base import Base +from schemas.enums import LinkCodeStatus + + +class LinkCode(Base): + __tablename__ = "link_codes" + + code: Mapped[str] = mapped_column(TEXT, nullable=False, unique=True, primary_key=True) + user_id: Mapped[int] = mapped_column(ForeignKey("users.id"), nullable=False) + status: Mapped[LinkCodeStatus] = mapped_column( + Enum(LinkCodeStatus, name="linkcodestatus"), nullable=False, default=LinkCodeStatus.ACTIVE + ) + + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), + nullable=False, + default=lambda: datetime.now(UTC), + ) + expires_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), + nullable=False, + ) + used_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=True) diff --git a/repositories/link_codes.py b/repositories/link_codes.py new file mode 100644 index 0000000..e59301e --- /dev/null +++ b/repositories/link_codes.py @@ -0,0 +1,34 @@ +from datetime import datetime + +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from db.models.link_codes import LinkCode +from schemas.enums import LinkCodeStatus + + +async def create_link_code( + session: AsyncSession, + *, + code: str, + user_id: int, + status: LinkCodeStatus, + expires_at: datetime, +) -> LinkCode: + link_code = LinkCode( + code=code, + user_id=user_id, + status=status, + expires_at=expires_at, + ) + + session.add(link_code) + await session.commit() + return link_code + + +async def get_link_code_by_code(session: AsyncSession, code: str) -> LinkCode | None: + stmt = select(LinkCode).where(LinkCode.code == code) + r = await session.execute(stmt) + + return r.scalar_one_or_none() diff --git a/repositories/users.py b/repositories/users.py index 28c271b..24cf919 100644 --- a/repositories/users.py +++ b/repositories/users.py @@ -69,3 +69,9 @@ class UserRepository: user.balance += amount await self.session.commit() return user + + async def update_telegram_id(self, user: User, telegram_id: int) -> User: + user.telegram_id = telegram_id + await self.session.commit() + + return user diff --git a/routes/__init__.py b/routes/__init__.py index df33297..9da82aa 100644 --- a/routes/__init__.py +++ b/routes/__init__.py @@ -2,6 +2,7 @@ from fastapi import APIRouter from .auth import router as auth_router from .health import router as health_router +from .link_codes import router as link_code_router from .orders import router as orders_router from .payments import payment_routers from .plans import router as plans_router @@ -13,5 +14,6 @@ routers: list[APIRouter] = [ orders_router, users_routers, health_router, + link_code_router, *payment_routers, ] diff --git a/routes/link_codes.py b/routes/link_codes.py new file mode 100644 index 0000000..21ea6ab --- /dev/null +++ b/routes/link_codes.py @@ -0,0 +1,57 @@ +import secrets +from datetime import UTC, datetime, timedelta + +from fastapi import APIRouter, Depends, HTTPException +from sqlalchemy.ext.asyncio import AsyncSession + +from config import cfg +from core.deps import get_auth_context, get_service_identity +from db.session import get_db +from repositories.link_codes import create_link_code, get_link_code_by_code +from repositories.users import UserRepository +from schemas.dto import AuthContext +from schemas.enums import LinkCodeStatus +from schemas.link_codes import LinkCodeConsume, LinkCodeResponse +from schemas.user import UserInfo + +router = APIRouter(prefix="/link-codes") + + +@router.post("", response_model=LinkCodeResponse, status_code=201) +async def gen_link_code( + ctx: AuthContext = Depends(get_auth_context), session: AsyncSession = Depends(get_db) +): + code = secrets.token_urlsafe(cfg.link_code_length) + exp = datetime.now(UTC) + timedelta(minutes=cfg.link_code_ttl) + link_code = await create_link_code( + session, code=code, user_id=ctx.user.id, status=LinkCodeStatus.ACTIVE, expires_at=exp + ) + + return LinkCodeResponse(code=link_code.code, expires_at=link_code.expires_at) + + +@router.post("/consume", response_model=UserInfo) +async def consume_link_code( + payload: LinkCodeConsume, + ctx: AuthContext = Depends(get_service_identity), + session: AsyncSession = Depends(get_db), +): + link_code = await get_link_code_by_code(session, payload.code) + + if not link_code: + raise HTTPException(404, detail="Code not found") + + users_repo = UserRepository(session) + user = await users_repo.get_user_by_id(link_code.user_id) + + if not user: + raise HTTPException(404, detail="User not found") + + user = await users_repo.update_telegram_id(user, payload.telegram_id) + + return UserInfo( + username=user.username, + telegram_id=user.telegram_id, + referal_code=user.referal_code, + bonus_balance=user.balance, + ) diff --git a/schemas/dto.py b/schemas/dto.py index d62e095..6982638 100644 --- a/schemas/dto.py +++ b/schemas/dto.py @@ -16,3 +16,12 @@ class AuthContext: user: User auth_method: Literal["jwt", "service"] service: str | None = None + + +@dataclass +class ServiceIdentity: + """ + Proves that request is coming from verified source, no user context provided. + """ + + service: str diff --git a/schemas/enums.py b/schemas/enums.py index 57751ec..bc48fb0 100644 --- a/schemas/enums.py +++ b/schemas/enums.py @@ -9,3 +9,9 @@ class SubscriptionStatus(StrEnum): class ServiceSignatureStatus(StrEnum): ACTIVE = "active" INACTIVE = "inactive" + + +class LinkCodeStatus(StrEnum): + ACTIVE = "active" + USED = "used" + EXPIRED = "expired" diff --git a/schemas/link_codes.py b/schemas/link_codes.py new file mode 100644 index 0000000..6f57d3a --- /dev/null +++ b/schemas/link_codes.py @@ -0,0 +1,20 @@ +from datetime import datetime + +from pydantic import BaseModel, Field, computed_field + +from config import cfg + + +class LinkCodeResponse(BaseModel): + code: str = Field() + expires_at: datetime = Field() + + @computed_field + @property + def deep_link(self) -> str: + return cfg.bot_url + "?start=" + self.code + + +class LinkCodeConsume(BaseModel): + code: str = Field() + telegram_id: int = Field()