feat: add login, session cookies, and get_current_user
This commit is contained in:
30
backend/app/deps.py
Normal file
30
backend/app/deps.py
Normal file
@@ -0,0 +1,30 @@
|
|||||||
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
|
from fastapi import Cookie, Depends, HTTPException, status
|
||||||
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from app.db import get_db
|
||||||
|
from app.models.auth_session import AuthSession, hash_token
|
||||||
|
from app.models.user import User
|
||||||
|
|
||||||
|
SESSION_COOKIE_NAME = "qm_session"
|
||||||
|
|
||||||
|
|
||||||
|
async def get_current_user(
|
||||||
|
qm_session: str | None = Cookie(default=None),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
) -> User:
|
||||||
|
if qm_session is None:
|
||||||
|
raise HTTPException(status.HTTP_401_UNAUTHORIZED, "not authenticated")
|
||||||
|
|
||||||
|
token_hash = hash_token(qm_session)
|
||||||
|
result = await db.execute(select(AuthSession).where(AuthSession.token_hash == token_hash))
|
||||||
|
session = result.scalar_one_or_none()
|
||||||
|
if session is None or session.expires_at < datetime.now(timezone.utc):
|
||||||
|
raise HTTPException(status.HTTP_401_UNAUTHORIZED, "session expired")
|
||||||
|
|
||||||
|
user = await db.get(User, session.user_id)
|
||||||
|
if user is None:
|
||||||
|
raise HTTPException(status.HTTP_401_UNAUTHORIZED, "user not found")
|
||||||
|
return user
|
||||||
@@ -1,3 +1,4 @@
|
|||||||
|
from app.models.auth_session import AuthSession
|
||||||
from app.models.user import User
|
from app.models.user import User
|
||||||
|
|
||||||
__all__ = ["User"]
|
__all__ = ["User", "AuthSession"]
|
||||||
|
|||||||
32
backend/app/models/auth_session.py
Normal file
32
backend/app/models/auth_session.py
Normal file
@@ -0,0 +1,32 @@
|
|||||||
|
import hashlib
|
||||||
|
import secrets
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime, timedelta, timezone
|
||||||
|
|
||||||
|
from sqlalchemy import DateTime, ForeignKey, String
|
||||||
|
from sqlalchemy.orm import Mapped, mapped_column
|
||||||
|
|
||||||
|
from app.db import Base
|
||||||
|
|
||||||
|
SESSION_TTL = timedelta(days=14)
|
||||||
|
|
||||||
|
|
||||||
|
class AuthSession(Base):
|
||||||
|
__tablename__ = "auth_sessions"
|
||||||
|
|
||||||
|
id: Mapped[uuid.UUID] = mapped_column(primary_key=True, default=uuid.uuid4)
|
||||||
|
user_id: Mapped[uuid.UUID] = mapped_column(ForeignKey("users.id"))
|
||||||
|
token_hash: Mapped[str] = mapped_column(String(64), unique=True, index=True)
|
||||||
|
created_at: Mapped[datetime] = mapped_column(
|
||||||
|
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc)
|
||||||
|
)
|
||||||
|
expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True))
|
||||||
|
|
||||||
|
|
||||||
|
def generate_session_token() -> tuple[str, str]:
|
||||||
|
raw = secrets.token_urlsafe(32)
|
||||||
|
return raw, hash_token(raw)
|
||||||
|
|
||||||
|
|
||||||
|
def hash_token(raw: str) -> str:
|
||||||
|
return hashlib.sha256(raw.encode()).hexdigest()
|
||||||
@@ -1,11 +1,15 @@
|
|||||||
from fastapi import APIRouter, Depends, HTTPException, status
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Depends, HTTPException, Response, status
|
||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
from app.db import get_db
|
from app.db import get_db
|
||||||
|
from app.deps import SESSION_COOKIE_NAME, get_current_user
|
||||||
|
from app.models.auth_session import AuthSession, SESSION_TTL, generate_session_token
|
||||||
from app.models.user import User
|
from app.models.user import User
|
||||||
from app.schemas import RegisterRequest, UserOut
|
from app.schemas import LoginRequest, RegisterRequest, UserOut
|
||||||
from app.security import hash_password
|
from app.security import hash_password, verify_password
|
||||||
|
|
||||||
router = APIRouter(prefix="/auth", tags=["auth"])
|
router = APIRouter(prefix="/auth", tags=["auth"])
|
||||||
|
|
||||||
@@ -25,3 +29,38 @@ async def register(payload: RegisterRequest, db: AsyncSession = Depends(get_db))
|
|||||||
await db.commit()
|
await db.commit()
|
||||||
await db.refresh(user)
|
await db.refresh(user)
|
||||||
return user
|
return user
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/login", response_model=UserOut)
|
||||||
|
async def login(payload: LoginRequest, response: Response, db: AsyncSession = Depends(get_db)):
|
||||||
|
user = await db.scalar(select(User).where(User.username == payload.username))
|
||||||
|
if user is None or not verify_password(payload.password, user.password_hash):
|
||||||
|
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid credentials")
|
||||||
|
|
||||||
|
raw_token, token_hash = generate_session_token()
|
||||||
|
session = AuthSession(
|
||||||
|
user_id=user.id,
|
||||||
|
token_hash=token_hash,
|
||||||
|
expires_at=datetime.now(timezone.utc) + SESSION_TTL,
|
||||||
|
)
|
||||||
|
db.add(session)
|
||||||
|
await db.commit()
|
||||||
|
|
||||||
|
response.set_cookie(
|
||||||
|
SESSION_COOKIE_NAME,
|
||||||
|
raw_token,
|
||||||
|
httponly=True,
|
||||||
|
samesite="lax",
|
||||||
|
max_age=int(SESSION_TTL.total_seconds()),
|
||||||
|
)
|
||||||
|
return user
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/logout", status_code=status.HTTP_204_NO_CONTENT)
|
||||||
|
async def logout(response: Response):
|
||||||
|
response.delete_cookie(SESSION_COOKIE_NAME)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/me", response_model=UserOut)
|
||||||
|
async def me(user: User = Depends(get_current_user)):
|
||||||
|
return user
|
||||||
|
|||||||
@@ -15,3 +15,8 @@ class UserOut(BaseModel):
|
|||||||
|
|
||||||
class Config:
|
class Config:
|
||||||
from_attributes = True
|
from_attributes = True
|
||||||
|
|
||||||
|
|
||||||
|
class LoginRequest(BaseModel):
|
||||||
|
username: str
|
||||||
|
password: str
|
||||||
|
|||||||
@@ -19,3 +19,28 @@ async def test_register_duplicate_username_rejected(client):
|
|||||||
await client.post("/auth/register", json={"username": "medium1", "password": "spookyspooky"})
|
await client.post("/auth/register", json={"username": "medium1", "password": "spookyspooky"})
|
||||||
response = await client.post("/auth/register", json={"username": "medium1", "password": "anotherpass"})
|
response = await client.post("/auth/register", json={"username": "medium1", "password": "anotherpass"})
|
||||||
assert response.status_code == 409
|
assert response.status_code == 409
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_login_sets_cookie_and_me_returns_user(client):
|
||||||
|
await client.post("/auth/register", json={"username": "medium2", "password": "spookyspooky"})
|
||||||
|
login_resp = await client.post("/auth/login", json={"username": "medium2", "password": "spookyspooky"})
|
||||||
|
assert login_resp.status_code == 200
|
||||||
|
assert "qm_session" in login_resp.cookies
|
||||||
|
|
||||||
|
me_resp = await client.get("/auth/me")
|
||||||
|
assert me_resp.status_code == 200
|
||||||
|
assert me_resp.json()["username"] == "medium2"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_login_wrong_password_rejected(client):
|
||||||
|
await client.post("/auth/register", json={"username": "medium3", "password": "spookyspooky"})
|
||||||
|
response = await client.post("/auth/login", json={"username": "medium3", "password": "wrongpass"})
|
||||||
|
assert response.status_code == 401
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_me_without_cookie_rejected(client):
|
||||||
|
response = await client.get("/auth/me")
|
||||||
|
assert response.status_code == 401
|
||||||
|
|||||||
Reference in New Issue
Block a user