更新
This commit is contained in:
@@ -0,0 +1,284 @@
|
||||
"""Persist and cache roles; seed built-ins on startup."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi import HTTPException, status
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from models.models import Role, User
|
||||
from .permissions import ALL_PERMISSIONS, normalize_permissions
|
||||
from .roles import (
|
||||
ROLE_ADMIN,
|
||||
RoleRecord,
|
||||
default_role_seeds,
|
||||
ensure_role,
|
||||
get_cached_role,
|
||||
is_admin,
|
||||
list_cached_roles,
|
||||
role_label,
|
||||
sanitize_role_permissions,
|
||||
set_role_cache,
|
||||
)
|
||||
|
||||
logger = logging.getLogger("auth.roles")
|
||||
|
||||
_ROLE_CODE_RE = re.compile(r"^[a-z][a-z0-9_]{1,49}$")
|
||||
|
||||
|
||||
def _encode_permissions(codes: list[str]) -> str:
|
||||
return json.dumps(codes, ensure_ascii=False)
|
||||
|
||||
|
||||
def _decode_permissions(raw: str | None) -> list[str]:
|
||||
if not raw:
|
||||
return []
|
||||
try:
|
||||
data = json.loads(raw)
|
||||
except Exception:
|
||||
return []
|
||||
if not isinstance(data, list):
|
||||
return []
|
||||
return normalize_permissions([str(item) for item in data])
|
||||
|
||||
|
||||
def role_to_record(row: Role) -> RoleRecord:
|
||||
perms = list(ALL_PERMISSIONS) if row.is_admin else _decode_permissions(row.permissions)
|
||||
return RoleRecord(
|
||||
code=row.code,
|
||||
label=row.label,
|
||||
description=row.description or "",
|
||||
is_system=bool(row.is_system),
|
||||
is_admin=bool(row.is_admin),
|
||||
permissions=perms,
|
||||
)
|
||||
|
||||
|
||||
async def refresh_role_cache(db: AsyncSession) -> list[RoleRecord]:
|
||||
result = await db.execute(select(Role).order_by(Role.id.asc()))
|
||||
rows = result.scalars().all()
|
||||
records = [role_to_record(row) for row in rows]
|
||||
if not records:
|
||||
records = default_role_seeds()
|
||||
set_role_cache(records)
|
||||
return records
|
||||
|
||||
|
||||
async def seed_builtin_roles(db: AsyncSession) -> None:
|
||||
"""Insert missing built-in roles and keep admin permissions complete."""
|
||||
seeds = {seed.code: seed for seed in default_role_seeds()}
|
||||
result = await db.execute(select(Role))
|
||||
existing = {row.code: row for row in result.scalars().all()}
|
||||
changed = False
|
||||
|
||||
for code, seed in seeds.items():
|
||||
row = existing.get(code)
|
||||
payload = _encode_permissions(seed.permissions)
|
||||
if row is None:
|
||||
db.add(
|
||||
Role(
|
||||
code=seed.code,
|
||||
label=seed.label,
|
||||
description=seed.description,
|
||||
is_system=True,
|
||||
is_admin=seed.is_admin,
|
||||
permissions=payload,
|
||||
)
|
||||
)
|
||||
changed = True
|
||||
continue
|
||||
# Keep system flags and admin full permission set in sync.
|
||||
if not row.is_system:
|
||||
row.is_system = True
|
||||
changed = True
|
||||
if seed.is_admin and (not row.is_admin or row.permissions != payload):
|
||||
row.is_admin = True
|
||||
row.permissions = payload
|
||||
row.label = seed.label
|
||||
changed = True
|
||||
elif not row.label:
|
||||
row.label = seed.label
|
||||
changed = True
|
||||
|
||||
if changed:
|
||||
await db.commit()
|
||||
await refresh_role_cache(db)
|
||||
logger.info("Role cache loaded: %s", ", ".join(r.code for r in list_cached_roles()))
|
||||
|
||||
|
||||
async def list_roles(db: AsyncSession) -> list[RoleRecord]:
|
||||
await refresh_role_cache(db)
|
||||
return list_cached_roles()
|
||||
|
||||
|
||||
async def get_role_or_404(db: AsyncSession, code: str) -> Role:
|
||||
result = await db.execute(select(Role).where(Role.code == code))
|
||||
row = result.scalar_one_or_none()
|
||||
if not row:
|
||||
raise HTTPException(status_code=404, detail="角色不存在")
|
||||
return row
|
||||
|
||||
|
||||
async def count_users_with_role(db: AsyncSession, code: str) -> int:
|
||||
result = await db.execute(
|
||||
select(func.count()).select_from(User).where(User.role == code)
|
||||
)
|
||||
return int(result.scalar() or 0)
|
||||
|
||||
|
||||
async def count_admin_users(db: AsyncSession) -> int:
|
||||
result = await db.execute(
|
||||
select(func.count()).select_from(User).where(User.role == ROLE_ADMIN)
|
||||
)
|
||||
return int(result.scalar() or 0)
|
||||
|
||||
|
||||
def validate_role_code(code: str) -> str:
|
||||
value = str(code or "").strip().lower()
|
||||
if not _ROLE_CODE_RE.match(value):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="角色码需为小写字母开头,仅含小写字母/数字/下划线,长度 2-50",
|
||||
)
|
||||
return value
|
||||
|
||||
|
||||
async def create_role(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
code: str,
|
||||
label: str,
|
||||
description: str | None,
|
||||
permissions: list[str] | None,
|
||||
) -> RoleRecord:
|
||||
role_code = validate_role_code(code)
|
||||
if get_cached_role(role_code) or (
|
||||
await db.execute(select(Role).where(Role.code == role_code))
|
||||
).scalar_one_or_none():
|
||||
raise HTTPException(status_code=400, detail="角色码已存在")
|
||||
name = (label or "").strip() or role_code
|
||||
perms = sanitize_role_permissions(permissions, force_all=False)
|
||||
row = Role(
|
||||
code=role_code,
|
||||
label=name,
|
||||
description=(description or "").strip() or None,
|
||||
is_system=False,
|
||||
is_admin=False,
|
||||
permissions=_encode_permissions(perms),
|
||||
)
|
||||
db.add(row)
|
||||
await db.commit()
|
||||
await db.refresh(row)
|
||||
await refresh_role_cache(db)
|
||||
return role_to_record(row)
|
||||
|
||||
|
||||
async def update_role(
|
||||
db: AsyncSession,
|
||||
code: str,
|
||||
*,
|
||||
label: str | None = None,
|
||||
description: str | None = None,
|
||||
permissions: list[str] | None = None,
|
||||
) -> RoleRecord:
|
||||
row = await get_role_or_404(db, code)
|
||||
if row.is_admin or row.code == ROLE_ADMIN:
|
||||
# Admin role always keeps full permissions; label/description may update.
|
||||
if label is not None:
|
||||
row.label = (label or "").strip() or row.label
|
||||
if description is not None:
|
||||
row.description = (description or "").strip() or None
|
||||
row.permissions = _encode_permissions(list(ALL_PERMISSIONS))
|
||||
row.is_admin = True
|
||||
row.is_system = True
|
||||
row.updated_at = datetime.utcnow()
|
||||
await db.commit()
|
||||
await db.refresh(row)
|
||||
await refresh_role_cache(db)
|
||||
return role_to_record(row)
|
||||
|
||||
if label is not None:
|
||||
row.label = (label or "").strip() or row.label
|
||||
if description is not None:
|
||||
row.description = (description or "").strip() or None
|
||||
if permissions is not None:
|
||||
row.permissions = _encode_permissions(
|
||||
sanitize_role_permissions(permissions, force_all=False)
|
||||
)
|
||||
row.is_admin = False
|
||||
row.updated_at = datetime.utcnow()
|
||||
await db.commit()
|
||||
await db.refresh(row)
|
||||
await refresh_role_cache(db)
|
||||
return role_to_record(row)
|
||||
|
||||
|
||||
async def delete_role(db: AsyncSession, code: str) -> None:
|
||||
row = await get_role_or_404(db, code)
|
||||
if row.is_system or row.is_admin or row.code == ROLE_ADMIN:
|
||||
raise HTTPException(status_code=400, detail="系统内置角色不可删除")
|
||||
used = await count_users_with_role(db, code)
|
||||
if used > 0:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"仍有 {used} 个用户使用该角色,请先调整用户角色后再删除",
|
||||
)
|
||||
await db.delete(row)
|
||||
await db.commit()
|
||||
await refresh_role_cache(db)
|
||||
|
||||
|
||||
async def ensure_role_assignable(db: AsyncSession, role_code: str) -> str:
|
||||
"""Validate role exists in DB (refresh cache if needed)."""
|
||||
code = str(role_code or "").strip()
|
||||
if not code:
|
||||
raise HTTPException(status_code=400, detail="角色不能为空")
|
||||
if get_cached_role(code) is None:
|
||||
await refresh_role_cache(db)
|
||||
try:
|
||||
return ensure_role(code)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
|
||||
async def guard_last_admin_change(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
user: User,
|
||||
new_role: str | None = None,
|
||||
deactivating: bool = False,
|
||||
deleting: bool = False,
|
||||
) -> None:
|
||||
"""Prevent removing the last admin user."""
|
||||
if not is_admin(user.role):
|
||||
return
|
||||
admin_count = await count_admin_users(db)
|
||||
if admin_count > 1:
|
||||
return
|
||||
if deleting or deactivating:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="不能删除或禁用最后一个管理员账号",
|
||||
)
|
||||
if new_role is not None and not is_admin(new_role):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="不能将最后一个管理员改为非管理员角色",
|
||||
)
|
||||
|
||||
|
||||
def user_permission_payload(role_code: str) -> dict:
|
||||
record = get_cached_role(role_code)
|
||||
admin = bool(record.is_admin) if record else is_admin(role_code)
|
||||
from .roles import permissions_for_role
|
||||
|
||||
return {
|
||||
"role_label": role_label(role_code),
|
||||
"is_admin": admin,
|
||||
"permissions": permissions_for_role(role_code),
|
||||
}
|
||||
Reference in New Issue
Block a user