Compare commits
2 Commits
7ae3c98585
...
2e01b6502c
| Author | SHA1 | Date | |
|---|---|---|---|
| 2e01b6502c | |||
| 771d44d34c |
44
alembic/versions/6d1bae3dc723_link_codes.py
Normal file
44
alembic/versions/6d1bae3dc723_link_codes.py
Normal file
@@ -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 ###
|
||||||
28
config.py
28
config.py
@@ -7,19 +7,15 @@ from pydantic_settings import BaseSettings, SettingsConfigDict
|
|||||||
class Settings(BaseSettings):
|
class Settings(BaseSettings):
|
||||||
model_config = SettingsConfigDict(env_file=".env")
|
model_config = SettingsConfigDict(env_file=".env")
|
||||||
|
|
||||||
|
### Internal settings ###
|
||||||
postgres_user: str = Field()
|
postgres_user: str = Field()
|
||||||
postgres_password: str = Field()
|
postgres_password: str = Field()
|
||||||
postgres_host: str = Field()
|
postgres_host: str = Field()
|
||||||
postgres_port: str = Field()
|
postgres_port: str = Field()
|
||||||
postgres_db: str = Field()
|
postgres_db: str = Field()
|
||||||
|
|
||||||
private_key_fp: str = Field()
|
### Telegram ###
|
||||||
public_key_fp: str = Field()
|
bot_username: str = Field()
|
||||||
|
|
||||||
access_token_ttl: int = Field(description="Access token TTL (minutes)")
|
|
||||||
|
|
||||||
pally_shop_id: str = Field()
|
|
||||||
pally_token: str = Field()
|
|
||||||
|
|
||||||
### Remnawave ###
|
### Remnawave ###
|
||||||
remnawave_base_url: str = Field()
|
remnawave_base_url: str = Field()
|
||||||
@@ -27,12 +23,23 @@ class Settings(BaseSettings):
|
|||||||
remnawave_token: str = Field()
|
remnawave_token: str = Field()
|
||||||
remnawave_default_squads_raw: str = Field(alias="REMNAWAVE_DEFAULT_SQUADS_UUIDS")
|
remnawave_default_squads_raw: str = Field(alias="REMNAWAVE_DEFAULT_SQUADS_UUIDS")
|
||||||
|
|
||||||
###
|
### Finance ###
|
||||||
minimal_deposit: int = Field()
|
minimal_deposit: int = Field()
|
||||||
referal_bonus: int = Field(30)
|
referal_bonus: int = Field(30)
|
||||||
|
|
||||||
|
#### Pally ####
|
||||||
|
pally_shop_id: str = Field()
|
||||||
|
pally_token: str = Field()
|
||||||
|
|
||||||
### Security related ###
|
### 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)
|
min_password_length: int = Field(8)
|
||||||
|
link_code_length: int = Field(8)
|
||||||
password_security_threshold: int = Field(2)
|
password_security_threshold: int = Field(2)
|
||||||
|
|
||||||
@computed_field
|
@computed_field
|
||||||
@@ -61,5 +68,10 @@ class Settings(BaseSettings):
|
|||||||
with open(self.public_key_fp, "rb") as f:
|
with open(self.public_key_fp, "rb") as f:
|
||||||
return f.read()
|
return f.read()
|
||||||
|
|
||||||
|
@computed_field
|
||||||
|
@property
|
||||||
|
def bot_url(self) -> str:
|
||||||
|
return "https://t.me/" + self.bot_username.lstrip("@")
|
||||||
|
|
||||||
|
|
||||||
cfg = Settings() # type: ignore
|
cfg = Settings() # type: ignore
|
||||||
|
|||||||
20
core/deps.py
20
core/deps.py
@@ -3,9 +3,11 @@ 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 db.session import get_db
|
from db.session import get_db
|
||||||
from external.pally import PallyClient
|
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
|
from services.subscriptions import sync_user_subscription
|
||||||
|
|
||||||
|
|
||||||
@@ -25,5 +27,21 @@ async def get_auth_context(
|
|||||||
return ctx
|
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:
|
def get_pally_client() -> PallyClient:
|
||||||
return PallyClient(api_token=cfg.pally_token, payer_pays_commission=True)
|
return PallyClient(api_token=cfg.pally_token, payer_pays_commission=True)
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
from .addons import Addon
|
from .addons import Addon
|
||||||
from .invoice import Invoice
|
from .invoice import Invoice
|
||||||
|
from .link_codes import LinkCode
|
||||||
from .orders import Order, OrderAddon
|
from .orders import Order, OrderAddon
|
||||||
from .pricing import PricingConfig
|
from .pricing import PricingConfig
|
||||||
from .service_signatures import ServiceSignature
|
from .service_signatures import ServiceSignature
|
||||||
@@ -13,6 +14,7 @@ __all__ = [
|
|||||||
"Addon",
|
"Addon",
|
||||||
"BalanceTransaction",
|
"BalanceTransaction",
|
||||||
"Invoice",
|
"Invoice",
|
||||||
|
"LinkCode",
|
||||||
"Order",
|
"Order",
|
||||||
"OrderAddon",
|
"OrderAddon",
|
||||||
"PricingConfig",
|
"PricingConfig",
|
||||||
|
|||||||
28
db/models/link_codes.py
Normal file
28
db/models/link_codes.py
Normal file
@@ -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)
|
||||||
34
repositories/link_codes.py
Normal file
34
repositories/link_codes.py
Normal file
@@ -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()
|
||||||
@@ -69,3 +69,9 @@ class UserRepository:
|
|||||||
user.balance += amount
|
user.balance += amount
|
||||||
await self.session.commit()
|
await self.session.commit()
|
||||||
return user
|
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
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ from fastapi import APIRouter
|
|||||||
|
|
||||||
from .auth import router as auth_router
|
from .auth import router as auth_router
|
||||||
from .health import router as health_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 .orders import router as orders_router
|
||||||
from .payments import payment_routers
|
from .payments import payment_routers
|
||||||
from .plans import router as plans_router
|
from .plans import router as plans_router
|
||||||
@@ -13,5 +14,6 @@ routers: list[APIRouter] = [
|
|||||||
orders_router,
|
orders_router,
|
||||||
users_routers,
|
users_routers,
|
||||||
health_router,
|
health_router,
|
||||||
|
link_code_router,
|
||||||
*payment_routers,
|
*payment_routers,
|
||||||
]
|
]
|
||||||
|
|||||||
57
routes/link_codes.py
Normal file
57
routes/link_codes.py
Normal file
@@ -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,
|
||||||
|
)
|
||||||
@@ -16,3 +16,12 @@ class AuthContext:
|
|||||||
user: User
|
user: User
|
||||||
auth_method: Literal["jwt", "service"]
|
auth_method: Literal["jwt", "service"]
|
||||||
service: str | None = None
|
service: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ServiceIdentity:
|
||||||
|
"""
|
||||||
|
Proves that request is coming from verified source, no user context provided.
|
||||||
|
"""
|
||||||
|
|
||||||
|
service: str
|
||||||
|
|||||||
@@ -9,3 +9,9 @@ class SubscriptionStatus(StrEnum):
|
|||||||
class ServiceSignatureStatus(StrEnum):
|
class ServiceSignatureStatus(StrEnum):
|
||||||
ACTIVE = "active"
|
ACTIVE = "active"
|
||||||
INACTIVE = "inactive"
|
INACTIVE = "inactive"
|
||||||
|
|
||||||
|
|
||||||
|
class LinkCodeStatus(StrEnum):
|
||||||
|
ACTIVE = "active"
|
||||||
|
USED = "used"
|
||||||
|
EXPIRED = "expired"
|
||||||
|
|||||||
20
schemas/link_codes.py
Normal file
20
schemas/link_codes.py
Normal file
@@ -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()
|
||||||
137
tests/test_link_code_routes.py
Normal file
137
tests/test_link_code_routes.py
Normal file
@@ -0,0 +1,137 @@
|
|||||||
|
# ruff: noqa: PLR2004
|
||||||
|
|
||||||
|
from datetime import UTC, datetime, timedelta
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import AsyncMock, patch
|
||||||
|
|
||||||
|
from config import cfg
|
||||||
|
from routes import link_codes
|
||||||
|
from schemas.enums import LinkCodeStatus
|
||||||
|
|
||||||
|
|
||||||
|
def auth_context(user_id=42):
|
||||||
|
return SimpleNamespace(user=SimpleNamespace(id=user_id))
|
||||||
|
|
||||||
|
|
||||||
|
def service_identity():
|
||||||
|
return SimpleNamespace(service="telegram-bot")
|
||||||
|
|
||||||
|
|
||||||
|
def test_generate_link_code_requires_authorization(client):
|
||||||
|
response = client.post("/link-codes")
|
||||||
|
|
||||||
|
assert response.status_code == 403
|
||||||
|
assert response.json()["detail"] == "No authorization provided."
|
||||||
|
|
||||||
|
|
||||||
|
def test_generate_link_code_creates_active_code_with_expiry_and_deep_link(client):
|
||||||
|
client.app.dependency_overrides[link_codes.get_auth_context] = auth_context
|
||||||
|
created_codes = []
|
||||||
|
|
||||||
|
async def create_code(session, *, code, user_id, status, expires_at):
|
||||||
|
created_codes.append(
|
||||||
|
SimpleNamespace(
|
||||||
|
session=session,
|
||||||
|
code=code,
|
||||||
|
user_id=user_id,
|
||||||
|
status=status,
|
||||||
|
expires_at=expires_at,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return created_codes[-1]
|
||||||
|
|
||||||
|
before = datetime.now(UTC)
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"routes.link_codes.secrets.token_urlsafe", return_value="one-time-code"
|
||||||
|
) as token_urlsafe,
|
||||||
|
patch("routes.link_codes.create_link_code", side_effect=create_code),
|
||||||
|
):
|
||||||
|
response = client.post("/link-codes")
|
||||||
|
after = datetime.now(UTC)
|
||||||
|
|
||||||
|
assert response.status_code == 201
|
||||||
|
assert response.json() == {
|
||||||
|
"code": "one-time-code",
|
||||||
|
"expires_at": created_codes[0].expires_at.isoformat().replace("+00:00", "Z"),
|
||||||
|
"deep_link": f"{cfg.bot_url}?start=one-time-code",
|
||||||
|
}
|
||||||
|
assert created_codes[0].user_id == 42
|
||||||
|
assert created_codes[0].status is LinkCodeStatus.ACTIVE
|
||||||
|
token_urlsafe.assert_called_once_with(cfg.link_code_length)
|
||||||
|
assert before + timedelta(minutes=cfg.link_code_ttl) <= created_codes[0].expires_at
|
||||||
|
assert created_codes[0].expires_at <= after + timedelta(minutes=cfg.link_code_ttl)
|
||||||
|
|
||||||
|
|
||||||
|
def test_consume_link_code_requires_code_and_telegram_id(client):
|
||||||
|
client.app.dependency_overrides[link_codes.get_service_identity] = service_identity
|
||||||
|
|
||||||
|
missing_code = client.post("/link-codes/consume", json={"telegram_id": 12345})
|
||||||
|
missing_telegram_id = client.post("/link-codes/consume", json={"code": "link-code"})
|
||||||
|
invalid_telegram_id = client.post(
|
||||||
|
"/link-codes/consume", json={"code": "link-code", "telegram_id": "not-an-id"}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert missing_code.status_code == 422
|
||||||
|
assert missing_telegram_id.status_code == 422
|
||||||
|
assert invalid_telegram_id.status_code == 422
|
||||||
|
|
||||||
|
|
||||||
|
def test_consume_link_code_rejects_unknown_code(client):
|
||||||
|
client.app.dependency_overrides[link_codes.get_service_identity] = service_identity
|
||||||
|
lookup = AsyncMock(return_value=None)
|
||||||
|
with patch("routes.link_codes.get_link_code_by_code", lookup):
|
||||||
|
response = client.post(
|
||||||
|
"/link-codes/consume", json={"code": "missing", "telegram_id": 12345}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 404
|
||||||
|
assert response.json()["detail"] == "Code not found"
|
||||||
|
lookup.assert_awaited_once()
|
||||||
|
assert lookup.await_args.args[1] == "missing"
|
||||||
|
|
||||||
|
|
||||||
|
def test_consume_link_code_rejects_code_for_deleted_user(client):
|
||||||
|
client.app.dependency_overrides[link_codes.get_service_identity] = service_identity
|
||||||
|
link_code = SimpleNamespace(user_id=42)
|
||||||
|
repository = SimpleNamespace(get_user_by_id=AsyncMock(return_value=None))
|
||||||
|
with (
|
||||||
|
patch("routes.link_codes.get_link_code_by_code", new=AsyncMock(return_value=link_code)),
|
||||||
|
patch("routes.link_codes.UserRepository", return_value=repository),
|
||||||
|
):
|
||||||
|
response = client.post(
|
||||||
|
"/link-codes/consume", json={"code": "orphaned", "telegram_id": 12345}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 404
|
||||||
|
assert response.json()["detail"] == "User not found"
|
||||||
|
repository.get_user_by_id.assert_awaited_once_with(42)
|
||||||
|
|
||||||
|
|
||||||
|
def test_consume_link_code_updates_telegram_id_and_returns_user(client):
|
||||||
|
client.app.dependency_overrides[link_codes.get_service_identity] = service_identity
|
||||||
|
link_code = SimpleNamespace(user_id=42)
|
||||||
|
user = SimpleNamespace(
|
||||||
|
username="alice", telegram_id=12345, referal_code="ref-code", balance=100
|
||||||
|
)
|
||||||
|
repository = SimpleNamespace(
|
||||||
|
get_user_by_id=AsyncMock(return_value=user),
|
||||||
|
update_telegram_id=AsyncMock(return_value=user),
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
patch("routes.link_codes.get_link_code_by_code", new=AsyncMock(return_value=link_code)),
|
||||||
|
patch("routes.link_codes.UserRepository", return_value=repository),
|
||||||
|
):
|
||||||
|
response = client.post(
|
||||||
|
"/link-codes/consume", json={"code": "link-code", "telegram_id": 12345}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 200
|
||||||
|
assert response.json() == {
|
||||||
|
"username": "alice",
|
||||||
|
"telegram_id": 12345,
|
||||||
|
"referal_code": "ref-code",
|
||||||
|
"bonus_balance": 100,
|
||||||
|
}
|
||||||
|
repository.get_user_by_id.assert_awaited_once_with(42)
|
||||||
|
repository.update_telegram_id.assert_awaited_once_with(user, 12345)
|
||||||
Reference in New Issue
Block a user