diff --git a/core/secrets.py b/core/secrets.py index 14f586c..f85fb08 100644 --- a/core/secrets.py +++ b/core/secrets.py @@ -45,7 +45,7 @@ def decode_jwt( ) -> dict[str, Any] | None: try: return jwt.decode(token, public_key, algo) - except jwt.ExpiredSignatureError: + except jwt.exceptions.PyJWTError: return diff --git a/repositories/link_codes.py b/repositories/link_codes.py index e59301e..78a4eba 100644 --- a/repositories/link_codes.py +++ b/repositories/link_codes.py @@ -32,3 +32,10 @@ async def get_link_code_by_code(session: AsyncSession, code: str) -> LinkCode | r = await session.execute(stmt) return r.scalar_one_or_none() + + +async def use_link_code(session: AsyncSession, code: LinkCode) -> LinkCode: + code.status = LinkCodeStatus.USED + await session.commit() + + return code diff --git a/repositories/service_signatures.py b/repositories/service_signatures.py index 04ad38e..2916e97 100644 --- a/repositories/service_signatures.py +++ b/repositories/service_signatures.py @@ -2,10 +2,15 @@ from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from db.models.service_signatures import ServiceSignature +from schemas.enums import ServiceSignatureStatus async def get_active_signature_by_kid(session: AsyncSession, kid: str) -> ServiceSignature | None: - stmt = select(ServiceSignature).where(ServiceSignature.kid == kid) + stmt = ( + select(ServiceSignature) + .where(ServiceSignature.kid == kid) + .where(ServiceSignature.status == ServiceSignatureStatus.ACTIVE) + ) r = await session.execute(stmt) return r.scalar_one_or_none() diff --git a/routes/link_codes.py b/routes/link_codes.py index 21ea6ab..99f6fd9 100644 --- a/routes/link_codes.py +++ b/routes/link_codes.py @@ -7,7 +7,7 @@ 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.link_codes import create_link_code, get_link_code_by_code, use_link_code from repositories.users import UserRepository from schemas.dto import AuthContext from schemas.enums import LinkCodeStatus @@ -41,6 +41,9 @@ async def consume_link_code( if not link_code: raise HTTPException(404, detail="Code not found") + if link_code.status != LinkCodeStatus.ACTIVE: + raise HTTPException(404, detail="Code expired or is invalid.") + users_repo = UserRepository(session) user = await users_repo.get_user_by_id(link_code.user_id) @@ -49,6 +52,8 @@ async def consume_link_code( user = await users_repo.update_telegram_id(user, payload.telegram_id) + await use_link_code(session, code=link_code) + return UserInfo( username=user.username, telegram_id=user.telegram_id, diff --git a/tests/test_link_code_routes.py b/tests/test_link_code_routes.py index 537d873..525df3c 100644 --- a/tests/test_link_code_routes.py +++ b/tests/test_link_code_routes.py @@ -93,7 +93,7 @@ def test_consume_link_code_rejects_unknown_code(client): 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) + link_code = SimpleNamespace(user_id=42, status=LinkCodeStatus.ACTIVE) 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)), @@ -110,7 +110,7 @@ def test_consume_link_code_rejects_code_for_deleted_user(client): 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) + link_code = SimpleNamespace(user_id=42, status=LinkCodeStatus.ACTIVE) user = SimpleNamespace( username="alice", telegram_id=12345, referal_code="ref-code", balance=100 )