146 lines
5.3 KiB
Python
146 lines
5.3 KiB
Python
"""认证 API"""
|
|
import hmac
|
|
import secrets
|
|
from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit
|
|
|
|
from fastapi import APIRouter, Depends, HTTPException, Header, Request, Response
|
|
from pydantic import BaseModel
|
|
from sqlalchemy.orm import Session
|
|
from app.core.database import get_db
|
|
from app.core.casdoor import casdoor_sdk
|
|
from app.core.errors import internal_error
|
|
from app.core.security import create_access_token, verify_casdoor_token, verify_token
|
|
from app.core.config import settings
|
|
from app.models.user import User
|
|
from app.schemas.auth import Token, UserInfo
|
|
from datetime import datetime
|
|
|
|
router = APIRouter(prefix="/api/auth", tags=["认证"])
|
|
OAUTH_STATE_COOKIE = "h3c_oauth_state"
|
|
OAUTH_STATE_TTL_SECONDS = 300
|
|
|
|
|
|
def _with_oauth_state(url: str, state: str) -> str:
|
|
"""Replace the SDK-generated state with the browser-bound state value."""
|
|
parts = urlsplit(url)
|
|
query = [(key, value) for key, value in parse_qsl(parts.query, keep_blank_values=True) if key != "state"]
|
|
query.append(("state", state))
|
|
return urlunsplit((parts.scheme, parts.netloc, parts.path, urlencode(query), parts.fragment))
|
|
|
|
|
|
@router.get("/login")
|
|
def login(response: Response):
|
|
"""获取 Casdoor 登录 URL"""
|
|
state = secrets.token_urlsafe(32)
|
|
response.set_cookie(
|
|
key=OAUTH_STATE_COOKIE,
|
|
value=state,
|
|
max_age=OAUTH_STATE_TTL_SECONDS,
|
|
httponly=True,
|
|
secure=not settings.DEBUG,
|
|
samesite="lax",
|
|
path="/api/auth",
|
|
)
|
|
login_url = casdoor_sdk.get_auth_link(settings.CASDOOR_REDIRECT_URL)
|
|
return {"url": _with_oauth_state(login_url, state)}
|
|
|
|
|
|
class CallbackRequest(BaseModel):
|
|
code: str
|
|
state: str
|
|
|
|
|
|
@router.post("/callback", response_model=Token)
|
|
def callback(
|
|
body: CallbackRequest,
|
|
request: Request,
|
|
response: Response,
|
|
db: Session = Depends(get_db),
|
|
):
|
|
"""Casdoor 登录回调"""
|
|
try:
|
|
expected_state = request.cookies.get(OAUTH_STATE_COOKIE)
|
|
response.delete_cookie(OAUTH_STATE_COOKIE, path="/api/auth")
|
|
if not expected_state or not hmac.compare_digest(body.state, expected_state):
|
|
raise HTTPException(status_code=400, detail="登录状态校验失败,请重新登录")
|
|
|
|
token_response = casdoor_sdk.get_oauth_token(code=body.code)
|
|
if isinstance(token_response, dict) and "error" in token_response:
|
|
raise HTTPException(status_code=400, detail=token_response.get("error_description", token_response["error"]))
|
|
|
|
access_token = token_response.get("access_token") if isinstance(token_response, dict) else token_response
|
|
identity_token = token_response.get("id_token") if isinstance(token_response, dict) else None
|
|
if not access_token:
|
|
raise HTTPException(status_code=400, detail="Casdoor 未返回 access_token")
|
|
|
|
casdoor_user = verify_casdoor_token(identity_token or access_token)
|
|
|
|
user = db.query(User).filter(User.casdoor_id == casdoor_user["sub"]).first()
|
|
if not user:
|
|
user = User(
|
|
casdoor_id=casdoor_user["sub"],
|
|
username=casdoor_user.get("preferred_username") or casdoor_user.get("name", ""),
|
|
display_name=casdoor_user.get("displayName") or casdoor_user.get("name", ""),
|
|
email=casdoor_user.get("email"),
|
|
role="user"
|
|
)
|
|
db.add(user)
|
|
else:
|
|
# 每次登录同步 Casdoor 信息(姓名、邮箱等可能更新)
|
|
user.display_name = casdoor_user.get("displayName") or casdoor_user.get("name", user.display_name or "")
|
|
user.email = casdoor_user.get("email", user.email)
|
|
|
|
user.last_login = datetime.utcnow()
|
|
db.commit()
|
|
|
|
jwt_token = create_access_token({
|
|
"sub": str(user.id),
|
|
"username": user.username,
|
|
"role": user.role or "user",
|
|
})
|
|
return {"access_token": jwt_token}
|
|
except HTTPException:
|
|
raise
|
|
except Exception as e:
|
|
raise internal_error("Casdoor login callback", e)
|
|
|
|
|
|
@router.get("/permissions")
|
|
def get_my_permissions(
|
|
authorization: str = Header(None, alias="Authorization"),
|
|
db: Session = Depends(get_db)
|
|
):
|
|
"""获取当前用户的权限码列表"""
|
|
from app.middleware.permission_middleware import get_role_permissions
|
|
if not authorization or not authorization.startswith("Bearer "):
|
|
raise HTTPException(status_code=401, detail="未授权")
|
|
token = authorization[7:]
|
|
payload = verify_token(token)
|
|
if not payload:
|
|
raise HTTPException(status_code=401, detail="无效的令牌")
|
|
role = payload.get('role', 'user')
|
|
perms = get_role_permissions(role, db)
|
|
return {"role": role, "permissions": perms}
|
|
|
|
|
|
@router.get("/profile")
|
|
def get_profile(
|
|
authorization: str = Header(None, alias="Authorization"),
|
|
db: Session = Depends(get_db)
|
|
):
|
|
"""获取当前用户信息"""
|
|
if not authorization or not authorization.startswith("Bearer "):
|
|
raise HTTPException(status_code=401, detail="未授权")
|
|
token = authorization[7:]
|
|
payload = verify_token(token)
|
|
if not payload:
|
|
raise HTTPException(status_code=401, detail="无效的令牌")
|
|
|
|
user = db.query(User).filter(User.id == int(payload["sub"])).first()
|
|
if not user:
|
|
raise HTTPException(status_code=404, detail="用户不存在")
|
|
|
|
return user
|
|
|
|
|