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:
|
) -> 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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user