87 lines
2.9 KiB
Python
87 lines
2.9 KiB
Python
import logging
|
|
from typing import Optional
|
|
|
|
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
|
from pydantic import BaseModel
|
|
from sqlalchemy import text
|
|
from sqlalchemy.orm import Session
|
|
|
|
try:
|
|
from app.db import get_db
|
|
from app.security import create_access_token
|
|
except ModuleNotFoundError:
|
|
from db import get_db
|
|
from security import create_access_token
|
|
|
|
|
|
router = APIRouter(prefix="/auth", tags=["auth"])
|
|
logger = logging.getLogger("app.auth")
|
|
|
|
|
|
class LoginRequest(BaseModel):
|
|
username: Optional[str] = None
|
|
password: Optional[str] = None
|
|
code: Optional[str] = None
|
|
nickname: Optional[str] = None
|
|
avatar_url: Optional[str] = None
|
|
|
|
|
|
class LoginResponse(BaseModel):
|
|
access_token: str
|
|
expires_in: int
|
|
user_id: int
|
|
|
|
|
|
@router.post("/login", response_model=LoginResponse)
|
|
def login(
|
|
payload: LoginRequest,
|
|
request: Request,
|
|
db: Session = Depends(get_db),
|
|
) -> LoginResponse:
|
|
logger.info(f"/auth/login called with code={payload.code[:20] if payload.code else None}...")
|
|
identifier = payload.code or payload.username
|
|
if not identifier:
|
|
raise HTTPException(status_code=400, detail="username or code is required")
|
|
|
|
with db.begin():
|
|
row = (
|
|
db.execute(
|
|
text("SELECT user_id FROM parents WHERE openid = :openid"),
|
|
{"openid": identifier},
|
|
).mappings().first()
|
|
)
|
|
|
|
if row:
|
|
user_id = int(row["user_id"])
|
|
logger.info(f"Existing user found: user_id={user_id}")
|
|
if payload.nickname or payload.avatar_url:
|
|
db.execute(
|
|
text(
|
|
"""
|
|
UPDATE parents
|
|
SET nickname = COALESCE(:nickname, nickname),
|
|
avatar_url = COALESCE(:avatar_url, avatar_url)
|
|
WHERE user_id = :user_id
|
|
"""
|
|
),
|
|
{"user_id": user_id, "nickname": payload.nickname, "avatar_url": payload.avatar_url},
|
|
)
|
|
else:
|
|
logger.info("Creating new user...")
|
|
db.execute(
|
|
text(
|
|
"INSERT INTO parents (openid, nickname, avatar_url, status) VALUES (:openid, :nickname, :avatar_url, 1)"
|
|
),
|
|
{"openid": identifier, "nickname": payload.nickname, "avatar_url": payload.avatar_url},
|
|
)
|
|
user_id = int(db.execute(text("SELECT last_insert_rowid()")).scalar_one())
|
|
logger.info(f"New user created: user_id={user_id}")
|
|
|
|
access_token, expires_in = create_access_token(user_id=user_id)
|
|
logger.info(f"Login succeeded: user_id={user_id}, expires_in={expires_in}")
|
|
return LoginResponse(access_token=access_token, expires_in=expires_in, user_id=user_id)
|
|
|
|
|
|
@router.post("/logout")
|
|
def logout(request: Request):
|
|
return {"message": "logged out"} |