from fastapi import APIRouter, Body, Depends, Header, HTTPException from sqlalchemy.orm import Session from database import get_db from models import Authorization, Avatar, 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"] TAKEOVER_DELAY_KEY = "takeoverReplyDelaySeconds" DEFAULT_TAKEOVER_DELAY_SECONDS = 180 MIN_TAKEOVER_DELAY_SECONDS = 3 MAX_TAKEOVER_DELAY_SECONDS = 86_400 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), "takeoverReplyDelaySeconds": _stored_takeover_delay(avatar), } def _stored_takeover_delay(avatar) -> int: raw = (avatar.config or {}).get(TAKEOVER_DELAY_KEY, DEFAULT_TAKEOVER_DELAY_SECONDS) if isinstance(raw, bool): return DEFAULT_TAKEOVER_DELAY_SECONDS try: delay = int(raw) except (TypeError, ValueError): return DEFAULT_TAKEOVER_DELAY_SECONDS if not MIN_TAKEOVER_DELAY_SECONDS <= delay <= MAX_TAKEOVER_DELAY_SECONDS: return DEFAULT_TAKEOVER_DELAY_SECONDS return delay def _validate_takeover_delay(value) -> int: if isinstance(value, bool) or not isinstance(value, int): raise ValueError("自动回复等待时间必须是整数秒") if not MIN_TAKEOVER_DELAY_SECONDS <= value <= MAX_TAKEOVER_DELAY_SECONDS: raise ValueError("自动回复等待时间需在 3 秒到 24 小时之间") return value def _disable_other_takeovers(db: Session, avatar) -> list[str]: disabled_ids = [] others = ( db.query(Avatar) .filter(Avatar.owner_id == avatar.owner_id, Avatar.id != avatar.id) .all() ) for other in others: permissions = _stored_avatar_permissions(other) if "takeover" not in permissions: continue other.config = { **(other.config or {}), AVATAR_PERMISSION_KEY: [item for item in permissions if item != "takeover"], } disabled_ids.append(other.id) tasks = ( db.query(TakeoverReplyTask) .filter( TakeoverReplyTask.avatar_id == other.id, TakeoverReplyTask.status.in_(("pending", "generating", "ready", "sending")), ) .all() ) for task in tasks: task.status = "cancelled" task.cancel_reason = "another_avatar_takeover_enabled" task.locked_at = None return disabled_ids 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 and TAKEOVER_DELAY_KEY not in payload: return fail("缺少授权设置", 400) try: permissions = ( _normalize_avatar_permissions(payload["permissions"]) if "permissions" in payload else _stored_avatar_permissions(avatar) ) takeover_delay = ( _validate_takeover_delay(payload[TAKEOVER_DELAY_KEY]) if TAKEOVER_DELAY_KEY in payload else _stored_takeover_delay(avatar) ) except ValueError as exc: return fail(str(exc), 400) previous_permissions = _stored_avatar_permissions(avatar) avatar.config = { **(avatar.config or {}), AVATAR_PERMISSION_KEY: permissions, TAKEOVER_DELAY_KEY: takeover_delay, } disabled_avatar_ids = _disable_other_takeovers(db, avatar) if "takeover" in permissions else [] 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) response = _permission_settings_payload(avatar) response["disabledAvatarIds"] = disabled_avatar_ids return ok(response, "授权设置已保存") @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=DEFAULT_TAKEOVER_DELAY_SECONDS, ) 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}, "授权已删除")