Files
dy/backend/auth/role_service.py
T
2026-08-07 17:51:57 +08:00

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),
}