fix(avatar): prevent takeover loops and isolate settings

This commit is contained in:
stefanfeng
2026-08-27 11:30:12 +08:00
parent ef58c5f2d2
commit 46d42b7d98
14 changed files with 848 additions and 62 deletions
@@ -2,7 +2,7 @@ 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 models import Authorization, Avatar, TakeoverCursor, TakeoverReplyTask
from responses import fail, ok
from routers.avatars import _require_owned_avatar
@@ -14,6 +14,10 @@ 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",
@@ -91,9 +95,62 @@ 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)
@@ -144,10 +201,19 @@ def update_permission_settings(
db: Session = Depends(get_db),
):
avatar = _require_owned_avatar(db, avatar_id, authorization)
if "permissions" not in payload:
return fail("缺少 permissions", 400)
if "permissions" not in payload and TAKEOVER_DELAY_KEY not in payload:
return fail("缺少授权设置", 400)
try:
permissions = _normalize_avatar_permissions(payload["permissions"])
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)
@@ -155,7 +221,9 @@ def update_permission_settings(
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
@@ -179,7 +247,9 @@ def update_permission_settings(
task.locked_at = None
db.commit()
db.refresh(avatar)
return ok(_permission_settings_payload(avatar), "授权设置已保存")
response = _permission_settings_payload(avatar)
response["disabledAvatarIds"] = disabled_avatar_ids
return ok(response, "授权设置已保存")
@router.get("/avatar/{avatar_id}/authorizations")
@@ -240,7 +310,7 @@ def create_auth(
status="active",
takeover_enabled=False,
takeover_mode="immediate",
takeover_delay_seconds=30,
takeover_delay_seconds=DEFAULT_TAKEOVER_DELAY_SECONDS,
)
db.add(item)
db.commit()
+42 -16
View File
@@ -6,7 +6,17 @@ from sqlalchemy.orm import Session
from database import get_db
from routers.knowledge import UPLOAD_DIR
from models import Avatar, KnowledgeDoc, KnowledgeChunk, QAPair, Authorization, User
from models import (
Authorization,
Avatar,
KnowledgeChunk,
KnowledgeDoc,
QAPair,
TakeoverCursor,
TakeoverMessage,
TakeoverReplyTask,
User,
)
from responses import ok, fail
router = APIRouter(tags=["分身"])
@@ -74,18 +84,21 @@ def list_avatars(page: int = 1, limit: int = 20, authorization: str = Header(Non
@router.get("/avatar/{avatar_id}")
def get_avatar(avatar_id: str, db: Session = Depends(get_db)):
a = db.query(Avatar).filter(Avatar.id == avatar_id).first()
if not a:
return fail("分身不存在", 404)
return ok(a.to_dict())
def get_avatar(
avatar_id: str,
authorization: str = Header(None),
db: Session = Depends(get_db),
):
return ok(_require_owned_avatar(db, avatar_id, authorization).to_dict())
@router.post("/avatar")
def create_avatar(payload: dict = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)):
user = _resolve_user(authorization, db)
if not user:
raise HTTPException(status_code=401, detail="未登录")
a = Avatar(
owner_id=user.huihui_user_id if user else "",
owner_id=user.huihui_user_id,
name=payload.get("name", "未命名分身"),
display_name=payload.get("displayName", "") or payload.get("display_name", ""),
description=payload.get("description", ""),
@@ -102,10 +115,13 @@ def create_avatar(payload: dict = Body(...), authorization: str = Header(None),
@router.put("/avatar/{avatar_id}")
def update_avatar(avatar_id: str, payload: dict = Body(...), db: Session = Depends(get_db)):
a = db.query(Avatar).filter(Avatar.id == avatar_id).first()
if not a:
return fail("分身不存在", 404)
def update_avatar(
avatar_id: str,
payload: dict = Body(...),
authorization: str = Header(None),
db: Session = Depends(get_db),
):
a = _require_owned_avatar(db, avatar_id, authorization)
mapping = {
"displayName": "display_name",
"photoUrl": "photo_url",
@@ -114,22 +130,32 @@ def update_avatar(avatar_id: str, payload: dict = Body(...), db: Session = Depen
for key in ("name", "displayName", "description", "photoUrl", "emoji", "status", "tokenBalance", "config"):
if key in payload:
col = mapping.get(key, key)
setattr(a, col, payload[key])
value = payload[key]
if key == "config":
if not isinstance(value, dict):
return fail("分身配置格式不正确", 400)
value = {**(a.config or {}), **value}
setattr(a, col, value)
db.commit()
db.refresh(a)
return ok(a.to_dict())
@router.delete("/avatar/{avatar_id}")
def delete_avatar(avatar_id: str, db: Session = Depends(get_db)):
a = db.query(Avatar).filter(Avatar.id == avatar_id).first()
if not a:
return fail("分身不存在", 404)
def delete_avatar(
avatar_id: str,
authorization: str = Header(None),
db: Session = Depends(get_db),
):
a = _require_owned_avatar(db, avatar_id, authorization)
# 级联清理关联数据,避免孤儿记录
db.query(KnowledgeDoc).filter(KnowledgeDoc.avatar_id == avatar_id).delete()
db.query(KnowledgeChunk).filter(KnowledgeChunk.avatar_id == avatar_id).delete()
db.query(QAPair).filter(QAPair.avatar_id == avatar_id).delete()
db.query(Authorization).filter(Authorization.avatar_id == avatar_id).delete()
db.query(TakeoverReplyTask).filter(TakeoverReplyTask.avatar_id == avatar_id).delete()
db.query(TakeoverMessage).filter(TakeoverMessage.avatar_id == avatar_id).delete()
db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar_id).delete()
db.delete(a)
db.commit()
return ok({"success": True})
+26 -5
View File
@@ -8,13 +8,25 @@ from sqlalchemy.orm import Session
from database import get_db
from models import TakeoverCursor, TakeoverReplyTask, User
from responses import fail, ok
from routers.authorizations import _require_authorization
from routers.authorizations import (
DEFAULT_TAKEOVER_DELAY_SECONDS,
MAX_TAKEOVER_DELAY_SECONDS,
MIN_TAKEOVER_DELAY_SECONDS,
_require_authorization,
_stored_takeover_delay,
)
from routers.avatars import _require_owned_avatar
router = APIRouter(tags=["分身接管"])
BOXIM_STATUS_FRESH_SECONDS = 60
def _delay_label(seconds: int) -> str:
if seconds % 60 == 0:
return f"{seconds // 60} 分钟"
return f"{seconds} 秒"
@router.get("/avatar/{avatar_id}/takeover/status")
def get_takeover_status(
avatar_id: str,
@@ -24,6 +36,7 @@ def get_takeover_status(
avatar = _require_owned_avatar(db, avatar_id, authorization)
permissions = (avatar.config or {}).get("authorizationPermissions", [])
enabled = isinstance(permissions, list) and "takeover" in permissions
reply_delay_seconds = _stored_takeover_delay(avatar)
user = db.query(User).filter(User.huihui_user_id == avatar.owner_id).first()
cursor = db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar.id).first()
pending_count = (
@@ -49,7 +62,10 @@ def get_takeover_status(
and cursor.last_polled_at
>= datetime.utcnow() - timedelta(seconds=BOXIM_STATUS_FRESH_SECONDS)
):
status, message = "ready", "BOXIM 已连接,收到私聊消息 3 秒后自动回复"
status, message = (
"ready",
f"BOXIM 已连接,收到私聊消息 {_delay_label(reply_delay_seconds)}后自动回复",
)
else:
status, message = "connecting", "正在连接 BOXIM"
@@ -59,6 +75,7 @@ def get_takeover_status(
"status": status,
"message": message,
"pendingCount": pending_count,
"takeoverReplyDelaySeconds": reply_delay_seconds,
"lastPolledAt": cursor.last_polled_at.isoformat() if cursor and cursor.last_polled_at else None,
}
)
@@ -91,7 +108,7 @@ def update_takeover_config(
auth = _require_authorization(db, avatar_id, str(auth_id))
enabled = bool(auth.takeover_enabled)
mode = auth.takeover_mode or "immediate"
delay = auth.takeover_delay_seconds or 30
delay = auth.takeover_delay_seconds or DEFAULT_TAKEOVER_DELAY_SECONDS
if _has(payload, "takeoverEnabled", "takeover_enabled"):
raw_enabled = _read(payload, "takeoverEnabled", "takeover_enabled")
@@ -106,8 +123,12 @@ def update_takeover_config(
if _has(payload, "takeoverDelaySeconds", "takeover_delay_seconds"):
delay = _read(payload, "takeoverDelaySeconds", "takeover_delay_seconds")
if isinstance(delay, bool) or not isinstance(delay, int) or not 5 <= delay <= 3600:
return fail("延迟时间需在 5 到 3600 秒之间", 400)
if (
isinstance(delay, bool)
or not isinstance(delay, int)
or not MIN_TAKEOVER_DELAY_SECONDS <= delay <= MAX_TAKEOVER_DELAY_SECONDS
):
return fail("延迟时间需在 3 秒到 24 小时之间", 400)
if enabled and auth.target_type != "user":
return fail("本期仅支持对会会用户开启单聊接管", 400)