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:
try:
return jwt.decode(token, public_key, algo)
except jwt.ExpiredSignatureError:
except jwt.exceptions.PyJWTError:
return

View File

@@ -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

View File

@@ -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()

View File

@@ -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,

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):
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
)