fix(auth): handle all PyJWT errors gracefully and enforce link code status checks
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user