307 lines
9.6 KiB
Python
307 lines
9.6 KiB
Python
"""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, expand_paired_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 expand_paired_permissions(_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 system role permissions in sync."""
|
|
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 built-in permission sets in sync with code.
|
|
if not row.is_system:
|
|
row.is_system = True
|
|
changed = True
|
|
if seed.is_admin:
|
|
if not row.is_admin or row.permissions != payload or row.label != seed.label:
|
|
row.is_admin = True
|
|
row.permissions = payload
|
|
row.label = seed.label
|
|
if seed.description and row.description != seed.description:
|
|
row.description = seed.description
|
|
changed = True
|
|
else:
|
|
# operator / viewer: resync catalog so new menu/action codes ship.
|
|
desired = _encode_permissions(expand_paired_permissions(seed.permissions))
|
|
if row.permissions != desired or row.label != seed.label:
|
|
row.permissions = desired
|
|
row.label = seed.label
|
|
row.is_admin = False
|
|
changed = True
|
|
|
|
# Repair custom roles that have unpaired menu/action selections.
|
|
for row in existing.values():
|
|
if row.code in seeds or row.is_admin:
|
|
continue
|
|
repaired = expand_paired_permissions(_decode_permissions(row.permissions))
|
|
encoded = _encode_permissions(repaired)
|
|
if row.permissions != encoded:
|
|
row.permissions = encoded
|
|
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),
|
|
}
|