from fastapi import APIRouter, Body, Depends, Header, HTTPException from sqlalchemy.orm import Session from database import get_db from models import Authorization, TakeoverCursor, TakeoverReplyTask from responses import fail, ok from routers.avatars import _require_owned_avatar router = APIRouter(tags=["授权"]) TARGET_TYPES = {"user", "organization", "application"} PERMISSION_ORDER = ("friend", "chat", "publish", "browse", "interact", "takeover") ALLOWED_PERMISSIONS = set(PERMISSION_ORDER) AVATAR_PERMISSION_ORDER = PERMISSION_ORDER AVATAR_PERMISSION_KEY = "authorizationPermissions" DEFAULT_AVATAR_PERMISSIONS = ["friend", "chat"] LEGACY_PERMISSION_MAP = { "read": "browse", "reply": "chat", "write": "publish", "edit": "publish", } def _read(payload: dict, camel_key: str, snake_key: str | None = None, default=None): if camel_key in payload: return payload[camel_key] if snake_key and snake_key in payload: return payload[snake_key] return default def _clean_text(value, field_name: str, *, max_length: int) -> str: text = str(value or "").strip() if not text: raise ValueError(f"{field_name}不能为空") if len(text) > max_length: raise ValueError(f"{field_name}不能超过 {max_length} 个字符") return text def _normalize_permissions(value) -> list[str]: if not isinstance(value, list): raise ValueError("权限格式不正确") normalized = [] for raw in value: permission = LEGACY_PERMISSION_MAP.get(str(raw).strip(), str(raw).strip()) if permission not in ALLOWED_PERMISSIONS: raise ValueError(f"不支持的权限:{raw}") if permission not in normalized: normalized.append(permission) if not [item for item in normalized if item != "takeover"]: raise ValueError("请至少选择一项权限") return sorted(normalized, key=PERMISSION_ORDER.index) def _normalize_avatar_permissions(value) -> list[str]: if not isinstance(value, list): raise ValueError("权限格式不正确") normalized = [] for raw in value: permission = LEGACY_PERMISSION_MAP.get(str(raw).strip(), str(raw).strip()) if permission not in AVATAR_PERMISSION_ORDER: raise ValueError(f"不支持的权限:{raw}") if permission not in normalized: normalized.append(permission) return sorted(normalized, key=AVATAR_PERMISSION_ORDER.index) def _stored_avatar_permissions(avatar) -> list[str]: config = avatar.config or {} if AVATAR_PERMISSION_KEY not in config: return list(DEFAULT_AVATAR_PERMISSIONS) stored = config.get(AVATAR_PERMISSION_KEY) if not isinstance(stored, list): return list(DEFAULT_AVATAR_PERMISSIONS) permissions = [] for raw in stored: permission = LEGACY_PERMISSION_MAP.get(str(raw).strip(), str(raw).strip()) if permission in AVATAR_PERMISSION_ORDER and permission not in permissions: permissions.append(permission) return sorted(permissions, key=AVATAR_PERMISSION_ORDER.index) def _permission_settings_payload(avatar) -> dict: return { "avatarId": avatar.id, "permissions": _stored_avatar_permissions(avatar), } def _require_authorization(db: Session, avatar_id: str, authorization_id: str) -> Authorization: authorization = ( db.query(Authorization) .filter( Authorization.id == authorization_id, Authorization.avatar_id == avatar_id, ) .first() ) if not authorization: raise HTTPException(status_code=404, detail="授权不存在") return authorization def _duplicate_target( db: Session, avatar_id: str, target_type: str, target_id: str, *, exclude_id: str | None = None, ): query = db.query(Authorization).filter( Authorization.avatar_id == avatar_id, Authorization.target_type == target_type, Authorization.target_id == target_id, ) if exclude_id: query = query.filter(Authorization.id != exclude_id) return query.first() @router.get("/avatar/{avatar_id}/permission-settings") def get_permission_settings( avatar_id: str, authorization: str = Header(None), db: Session = Depends(get_db), ): avatar = _require_owned_avatar(db, avatar_id, authorization) return ok(_permission_settings_payload(avatar)) @router.put("/avatar/{avatar_id}/permission-settings") def update_permission_settings( avatar_id: str, payload: dict = Body(...), authorization: str = Header(None), db: Session = Depends(get_db), ): avatar = _require_owned_avatar(db, avatar_id, authorization) if "permissions" not in payload: return fail("缺少 permissions", 400) try: permissions = _normalize_avatar_permissions(payload["permissions"]) except ValueError as exc: return fail(str(exc), 400) previous_permissions = _stored_avatar_permissions(avatar) avatar.config = { **(avatar.config or {}), AVATAR_PERMISSION_KEY: permissions, } cursor = db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar.id).first() if cursor and "takeover" in permissions and "takeover" not in previous_permissions: cursor.initialized = False cursor.last_message_id = "0" cursor.last_error = "" elif cursor and "takeover" not in permissions: cursor.last_error = "" if "takeover" not in permissions: tasks = ( db.query(TakeoverReplyTask) .filter( TakeoverReplyTask.avatar_id == avatar.id, TakeoverReplyTask.status.in_(("pending", "generating", "ready", "sending")), ) .all() ) for task in tasks: task.status = "cancelled" task.cancel_reason = "takeover_disabled" task.locked_at = None db.commit() db.refresh(avatar) return ok(_permission_settings_payload(avatar), "授权设置已保存") @router.get("/avatar/{avatar_id}/authorizations") def list_auth( avatar_id: str, authorization: str = Header(None), db: Session = Depends(get_db), ): _require_owned_avatar(db, avatar_id, authorization) items = ( db.query(Authorization) .filter(Authorization.avatar_id == avatar_id) .order_by(Authorization.created_at.desc()) .all() ) return ok([item.to_dict() for item in items]) @router.post("/avatar/{avatar_id}/authorizations") def create_auth( avatar_id: str, payload: dict = Body(...), authorization: str = Header(None), db: Session = Depends(get_db), ): _require_owned_avatar(db, avatar_id, authorization) try: target_type = _clean_text( _read(payload, "targetType", "target_type", "user"), "授权类型", max_length=24, ) if target_type not in TARGET_TYPES: return fail("授权类型不正确", 400) target_id = _clean_text( _read(payload, "targetId", "target_id"), "对象标识", max_length=120, ) target_name = _clean_text( _read(payload, "targetName", "target_name"), "对象名称", max_length=50, ) permissions = _normalize_permissions(payload.get("permissions", [])) except ValueError as exc: return fail(str(exc), 400) if _duplicate_target(db, avatar_id, target_type, target_id): return fail("该对象已在授权列表中,可直接编辑现有授权", 409) item = Authorization( avatar_id=avatar_id, target_type=target_type, target_id=target_id, target_name=target_name, permissions=permissions, status="active", takeover_enabled=False, takeover_mode="immediate", takeover_delay_seconds=30, ) db.add(item) db.commit() db.refresh(item) return ok(item.to_dict(), "授权已添加") @router.put("/avatar/{avatar_id}/authorizations") def update_auth( avatar_id: str, payload: dict = Body(...), authorization: str = Header(None), db: Session = Depends(get_db), ): _require_owned_avatar(db, avatar_id, authorization) auth_id = payload.get("id") or _read(payload, "authorizationId", "authorization_id") if not auth_id: return fail("缺少授权 id", 400) item = _require_authorization(db, avatar_id, str(auth_id)) try: target_type = item.target_type target_id = item.target_id if "targetType" in payload or "target_type" in payload: target_type = _clean_text( _read(payload, "targetType", "target_type"), "授权类型", max_length=24, ) if target_type not in TARGET_TYPES: return fail("授权类型不正确", 400) if "targetId" in payload or "target_id" in payload: target_id = _clean_text( _read(payload, "targetId", "target_id"), "对象标识", max_length=120, ) if "targetName" in payload or "target_name" in payload: item.target_name = _clean_text( _read(payload, "targetName", "target_name"), "对象名称", max_length=50, ) if "permissions" in payload: item.permissions = _normalize_permissions(payload["permissions"]) except ValueError as exc: return fail(str(exc), 400) if _duplicate_target( db, avatar_id, target_type, target_id, exclude_id=item.id, ): return fail("该对象已在授权列表中", 409) if "status" in payload: status = str(payload["status"] or "") if status not in ("active", "inactive"): return fail("授权状态不正确", 400) item.status = status item.target_type = target_type item.target_id = target_id permissions = list(item.permissions or []) chat_allowed = "chat" in permissions or "reply" in permissions if item.status != "active" or item.target_type != "user" or not chat_allowed: item.takeover_enabled = False item.permissions = [permission for permission in permissions if permission != "takeover"] elif item.takeover_enabled and "takeover" not in permissions: item.permissions = permissions + ["takeover"] db.commit() db.refresh(item) return ok(item.to_dict(), "授权已更新") @router.delete("/avatar/{avatar_id}/authorizations/{authorization_id}") def delete_auth( avatar_id: str, authorization_id: str, authorization: str = Header(None), db: Session = Depends(get_db), ): _require_owned_avatar(db, avatar_id, authorization) item = _require_authorization(db, avatar_id, authorization_id) db.delete(item) db.commit() return ok({"id": authorization_id}, "授权已删除")