fix(auth): handle all PyJWT errors gracefully and enforce link code status checks

This commit is contained in:
2026-08-19 21:12:45 +07:00
parent 2e01b6502c
commit 77ef4aaa57
5 changed files with 22 additions and 5 deletions

View File

@@ -45,7 +45,7 @@ def decode_jwt(
) -> dict[str, Any] | None: ) -> dict[str, Any] | None:
try: try:
return jwt.decode(token, public_key, algo) return jwt.decode(token, public_key, algo)
except jwt.ExpiredSignatureError: except jwt.exceptions.PyJWTError:
return return

View File

@@ -32,3 +32,10 @@ async def get_link_code_by_code(session: AsyncSession, code: str) -> LinkCode |
r = await session.execute(stmt) r = await session.execute(stmt)
return r.scalar_one_or_none() 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

View File

@@ -2,10 +2,15 @@ from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from db.models.service_signatures import ServiceSignature 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: 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) r = await session.execute(stmt)
return r.scalar_one_or_none() return r.scalar_one_or_none()

View File

@@ -7,7 +7,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from config import cfg from config import cfg
from core.deps import get_auth_context, get_service_identity from core.deps import get_auth_context, get_service_identity
from db.session import get_db 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 repositories.users import UserRepository
from schemas.dto import AuthContext from schemas.dto import AuthContext
from schemas.enums import LinkCodeStatus from schemas.enums import LinkCodeStatus
@@ -41,6 +41,9 @@ async def consume_link_code(
if not link_code: if not link_code:
raise HTTPException(404, detail="Code not found") 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) users_repo = UserRepository(session)
user = await users_repo.get_user_by_id(link_code.user_id) 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) user = await users_repo.update_telegram_id(user, payload.telegram_id)
await use_link_code(session, code=link_code)
return UserInfo( return UserInfo(
username=user.username, username=user.username,
telegram_id=user.telegram_id, telegram_id=user.telegram_id,

View File

@@ -93,7 +93,7 @@ def test_consume_link_code_rejects_unknown_code(client):
def test_consume_link_code_rejects_code_for_deleted_user(client): def test_consume_link_code_rejects_code_for_deleted_user(client):
client.app.dependency_overrides[link_codes.get_service_identity] = service_identity 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)) repository = SimpleNamespace(get_user_by_id=AsyncMock(return_value=None))
with ( with (
patch("routes.link_codes.get_link_code_by_code", new=AsyncMock(return_value=link_code)), 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): def test_consume_link_code_updates_telegram_id_and_returns_user(client):
client.app.dependency_overrides[link_codes.get_service_identity] = service_identity 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( user = SimpleNamespace(
username="alice", telegram_id=12345, referal_code="ref-code", balance=100 username="alice", telegram_id=12345, referal_code="ref-code", balance=100
) )