feat(avatar): complete authorization management
This commit is contained in:
@@ -1,32 +1,237 @@
|
|||||||
from fastapi import APIRouter, Depends, Body
|
from fastapi import APIRouter, Body, Depends, Header, HTTPException
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from database import get_db
|
from database import get_db
|
||||||
from models import Authorization
|
from models import Authorization
|
||||||
from responses import ok, fail
|
from responses import fail, ok
|
||||||
|
from routers.avatars import _require_owned_avatar
|
||||||
|
|
||||||
router = APIRouter(tags=["授权"])
|
router = APIRouter(tags=["授权"])
|
||||||
|
|
||||||
|
TARGET_TYPES = {"user", "organization", "application"}
|
||||||
|
PERMISSION_ORDER = ("friend", "chat", "publish", "browse", "interact", "takeover")
|
||||||
|
ALLOWED_PERMISSIONS = set(PERMISSION_ORDER)
|
||||||
|
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 _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}/authorizations")
|
@router.get("/avatar/{avatar_id}/authorizations")
|
||||||
def list_auth(avatar_id: str, db: Session = Depends(get_db)):
|
def list_auth(
|
||||||
# demo:返回全部授权(忽略具体 avatar 绑定,便于联调)
|
avatar_id: str,
|
||||||
items = db.query(Authorization).order_by(Authorization.created_at.desc()).all()
|
authorization: str = Header(None),
|
||||||
return ok([a.to_dict() for a in items])
|
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")
|
@router.put("/avatar/{avatar_id}/authorizations")
|
||||||
def update_auth(avatar_id: str, payload: dict = Body(...), db: Session = Depends(get_db)):
|
def update_auth(
|
||||||
auth_id = payload.get("id")
|
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:
|
if not auth_id:
|
||||||
return fail("缺少授权 id", 400)
|
return fail("缺少授权 id", 400)
|
||||||
a = db.query(Authorization).filter(Authorization.id == auth_id).first()
|
|
||||||
if not a:
|
item = _require_authorization(db, avatar_id, str(auth_id))
|
||||||
return fail("授权不存在", 404)
|
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:
|
if "status" in payload:
|
||||||
a.status = payload["status"]
|
status = str(payload["status"] or "")
|
||||||
if "permissions" in payload:
|
if status not in ("active", "inactive"):
|
||||||
a.permissions = payload["permissions"]
|
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.commit()
|
||||||
items = db.query(Authorization).order_by(Authorization.created_at.desc()).all()
|
db.refresh(item)
|
||||||
return ok([x.to_dict() for x in items])
|
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}, "授权已删除")
|
||||||
|
|||||||
@@ -1,41 +1,79 @@
|
|||||||
"""分身接管配置 API"""
|
"""数字分身单聊接管配置 API。"""
|
||||||
from fastapi import APIRouter, Depends, Body
|
|
||||||
|
from fastapi import APIRouter, Body, Depends, Header
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from database import get_db
|
from database import get_db
|
||||||
from models import Authorization
|
from responses import fail, ok
|
||||||
from responses import ok, fail
|
from routers.authorizations import _require_authorization
|
||||||
|
from routers.avatars import _require_owned_avatar
|
||||||
|
|
||||||
router = APIRouter(tags=["分身接管"])
|
router = APIRouter(tags=["分身接管"])
|
||||||
|
|
||||||
|
|
||||||
|
def _has(payload: dict, camel_key: str, snake_key: str) -> bool:
|
||||||
|
return camel_key in payload or snake_key in payload
|
||||||
|
|
||||||
|
|
||||||
|
def _read(payload: dict, camel_key: str, snake_key: str, default=None):
|
||||||
|
if camel_key in payload:
|
||||||
|
return payload[camel_key]
|
||||||
|
if snake_key in payload:
|
||||||
|
return payload[snake_key]
|
||||||
|
return default
|
||||||
|
|
||||||
|
|
||||||
@router.put("/avatar/{avatar_id}/authorizations/takeover")
|
@router.put("/avatar/{avatar_id}/authorizations/takeover")
|
||||||
def update_takeover_config(
|
def update_takeover_config(
|
||||||
avatar_id: str,
|
avatar_id: str,
|
||||||
payload: dict = Body(...),
|
payload: dict = Body(...),
|
||||||
|
authorization: str = Header(None),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""更新分身接管配置"""
|
_require_owned_avatar(db, avatar_id, authorization)
|
||||||
auth_id = payload.get("authorizationId") or payload.get("authorization_id")
|
auth_id = _read(payload, "authorizationId", "authorization_id")
|
||||||
if not auth_id:
|
if not auth_id:
|
||||||
return fail("缺少 authorization_id", 400)
|
return fail("缺少 authorization_id", 400)
|
||||||
|
|
||||||
auth = db.query(Authorization).filter(Authorization.id == auth_id).first()
|
auth = _require_authorization(db, avatar_id, str(auth_id))
|
||||||
if not auth:
|
enabled = bool(auth.takeover_enabled)
|
||||||
return fail("授权不存在", 404)
|
mode = auth.takeover_mode or "immediate"
|
||||||
|
delay = auth.takeover_delay_seconds or 30
|
||||||
|
|
||||||
if "takeover_enabled" in payload:
|
if _has(payload, "takeoverEnabled", "takeover_enabled"):
|
||||||
auth.takeover_enabled = payload["takeover_enabled"]
|
raw_enabled = _read(payload, "takeoverEnabled", "takeover_enabled")
|
||||||
if "takeover_mode" in payload:
|
if not isinstance(raw_enabled, bool):
|
||||||
mode = payload["takeover_mode"]
|
return fail("takeover_enabled 必须是布尔值", 400)
|
||||||
|
enabled = raw_enabled
|
||||||
|
|
||||||
|
if _has(payload, "takeoverMode", "takeover_mode"):
|
||||||
|
mode = _read(payload, "takeoverMode", "takeover_mode")
|
||||||
if mode not in ("immediate", "delayed"):
|
if mode not in ("immediate", "delayed"):
|
||||||
return fail("takeover_mode 必须是 immediate 或 delayed", 400)
|
return fail("takeover_mode 必须是 immediate 或 delayed", 400)
|
||||||
auth.takeover_mode = mode
|
|
||||||
if "takeover_delay_seconds" in payload:
|
|
||||||
delay = payload["takeover_delay_seconds"]
|
|
||||||
if not isinstance(delay, int) or delay < 5:
|
|
||||||
return fail("takeover_delay_seconds 必须 >= 5", 400)
|
|
||||||
auth.takeover_delay_seconds = delay
|
|
||||||
|
|
||||||
|
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 enabled and auth.target_type != "user":
|
||||||
|
return fail("本期仅支持对会会用户开启单聊接管", 400)
|
||||||
|
if enabled and auth.status != "active":
|
||||||
|
return fail("请先启用该授权,再开启聊天接管", 400)
|
||||||
|
|
||||||
|
permissions = list(auth.permissions or [])
|
||||||
|
if enabled:
|
||||||
|
if "chat" not in permissions and "reply" not in permissions:
|
||||||
|
permissions.append("chat")
|
||||||
|
if "takeover" not in permissions:
|
||||||
|
permissions.append("takeover")
|
||||||
|
else:
|
||||||
|
permissions = [permission for permission in permissions if permission != "takeover"]
|
||||||
|
|
||||||
|
auth.permissions = permissions
|
||||||
|
auth.takeover_enabled = enabled
|
||||||
|
auth.takeover_mode = mode
|
||||||
|
auth.takeover_delay_seconds = delay
|
||||||
db.commit()
|
db.commit()
|
||||||
return ok(auth.to_dict())
|
db.refresh(auth)
|
||||||
|
return ok(auth.to_dict(), "接管配置已保存")
|
||||||
|
|||||||
@@ -33,7 +33,11 @@ class TakeoverService:
|
|||||||
self, owner_huihui_id: str, from_user_id: str
|
self, owner_huihui_id: str, from_user_id: str
|
||||||
) -> Optional[Authorization]:
|
) -> Optional[Authorization]:
|
||||||
"""Check whether takeover is enabled for the given target user."""
|
"""Check whether takeover is enabled for the given target user."""
|
||||||
avatar = self.db.query(Avatar).filter(Avatar.owner_id == owner_huihui_id).first()
|
avatar = (
|
||||||
|
self.db.query(Avatar)
|
||||||
|
.filter(Avatar.owner_id == owner_huihui_id, Avatar.status == "active")
|
||||||
|
.first()
|
||||||
|
)
|
||||||
if not avatar:
|
if not avatar:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -41,10 +45,13 @@ class TakeoverService:
|
|||||||
self.db.query(Authorization)
|
self.db.query(Authorization)
|
||||||
.filter(Authorization.avatar_id == avatar.id)
|
.filter(Authorization.avatar_id == avatar.id)
|
||||||
.filter(Authorization.target_id == from_user_id)
|
.filter(Authorization.target_id == from_user_id)
|
||||||
|
.filter(Authorization.target_type == "user")
|
||||||
|
.filter(Authorization.status == "active")
|
||||||
.filter(Authorization.takeover_enabled == True)
|
.filter(Authorization.takeover_enabled == True)
|
||||||
.first()
|
.first()
|
||||||
)
|
)
|
||||||
return auth if auth and auth.takeover_enabled else None
|
permissions = set(auth.permissions or []) if auth else set()
|
||||||
|
return auth if auth and "takeover" in permissions else None
|
||||||
|
|
||||||
async def generate_reply(self, avatar_id: str, message: str) -> str:
|
async def generate_reply(self, avatar_id: str, message: str) -> str:
|
||||||
"""Call the avatar chat endpoint to generate a reply."""
|
"""Call the avatar chat endpoint to generate a reply."""
|
||||||
@@ -170,5 +177,8 @@ class TakeoverService:
|
|||||||
|
|
||||||
if auth.takeover_mode == "immediate":
|
if auth.takeover_mode == "immediate":
|
||||||
await self.execute_takeover(auth, message)
|
await self.execute_takeover(auth, message)
|
||||||
|
elif not self.redis:
|
||||||
|
logger.warning("Redis not configured, executing delayed takeover immediately")
|
||||||
|
await self.execute_takeover(auth, message)
|
||||||
else:
|
else:
|
||||||
self.enqueue_delayed_message(auth, message)
|
self.enqueue_delayed_message(auth, message)
|
||||||
|
|||||||
@@ -1,6 +1,8 @@
|
|||||||
|
import uuid
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from database import init_db, SessionLocal
|
from database import init_db, SessionLocal
|
||||||
from models import Authorization
|
from models import Authorization, Avatar, User
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(scope="session", autouse=True)
|
@pytest.fixture(scope="session", autouse=True)
|
||||||
@@ -24,3 +26,72 @@ def setup_database():
|
|||||||
db.commit()
|
db.commit()
|
||||||
finally:
|
finally:
|
||||||
db.close()
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def authorization_context():
|
||||||
|
"""Create isolated users, avatars, and one authorization for API tests."""
|
||||||
|
suffix = uuid.uuid4().hex
|
||||||
|
owner = User(
|
||||||
|
id=f"owner-{suffix}",
|
||||||
|
huihui_user_id=f"huihui-owner-{suffix}",
|
||||||
|
nickname="授权测试用户",
|
||||||
|
app_token=f"owner-token-{suffix}",
|
||||||
|
)
|
||||||
|
other = User(
|
||||||
|
id=f"other-{suffix}",
|
||||||
|
huihui_user_id=f"huihui-other-{suffix}",
|
||||||
|
nickname="其他用户",
|
||||||
|
app_token=f"other-token-{suffix}",
|
||||||
|
)
|
||||||
|
avatar = Avatar(
|
||||||
|
id=f"avatar-{suffix}",
|
||||||
|
owner_id=owner.huihui_user_id,
|
||||||
|
name="授权测试分身",
|
||||||
|
status="active",
|
||||||
|
config={},
|
||||||
|
)
|
||||||
|
other_avatar = Avatar(
|
||||||
|
id=f"other-avatar-{suffix}",
|
||||||
|
owner_id=other.huihui_user_id,
|
||||||
|
name="其他分身",
|
||||||
|
status="active",
|
||||||
|
config={},
|
||||||
|
)
|
||||||
|
authorization = Authorization(
|
||||||
|
id=f"authorization-{suffix}",
|
||||||
|
avatar_id=avatar.id,
|
||||||
|
target_type="user",
|
||||||
|
target_id=f"contact-{suffix}",
|
||||||
|
target_name="测试联系人",
|
||||||
|
permissions=["chat", "browse"],
|
||||||
|
status="active",
|
||||||
|
)
|
||||||
|
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
db.add_all([owner, other, avatar, other_avatar, authorization])
|
||||||
|
db.commit()
|
||||||
|
yield {
|
||||||
|
"owner": owner,
|
||||||
|
"other": other,
|
||||||
|
"avatar": avatar,
|
||||||
|
"other_avatar": other_avatar,
|
||||||
|
"authorization": authorization,
|
||||||
|
"owner_headers": {"Authorization": f"Bearer {owner.app_token}"},
|
||||||
|
"other_headers": {"Authorization": f"Bearer {other.app_token}"},
|
||||||
|
"suffix": suffix,
|
||||||
|
}
|
||||||
|
finally:
|
||||||
|
db.rollback()
|
||||||
|
db.query(Authorization).filter(
|
||||||
|
Authorization.avatar_id.in_([avatar.id, other_avatar.id])
|
||||||
|
).delete(synchronize_session=False)
|
||||||
|
db.query(Avatar).filter(Avatar.id.in_([avatar.id, other_avatar.id])).delete(
|
||||||
|
synchronize_session=False
|
||||||
|
)
|
||||||
|
db.query(User).filter(User.id.in_([owner.id, other.id])).delete(
|
||||||
|
synchronize_session=False
|
||||||
|
)
|
||||||
|
db.commit()
|
||||||
|
db.close()
|
||||||
|
|||||||
@@ -0,0 +1,94 @@
|
|||||||
|
from fastapi.testclient import TestClient
|
||||||
|
|
||||||
|
from main import app
|
||||||
|
|
||||||
|
|
||||||
|
client = TestClient(app)
|
||||||
|
|
||||||
|
|
||||||
|
def test_authorization_list_is_scoped_to_owned_avatar(authorization_context):
|
||||||
|
context = authorization_context
|
||||||
|
response = client.get(
|
||||||
|
f"/api/avatar/{context['avatar'].id}/authorizations",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
)
|
||||||
|
assert response.status_code == 200
|
||||||
|
payload = response.json()
|
||||||
|
assert payload["code"] == 200
|
||||||
|
assert [item["id"] for item in payload["data"]] == [context["authorization"].id]
|
||||||
|
|
||||||
|
forbidden = client.get(
|
||||||
|
f"/api/avatar/{context['other_avatar'].id}/authorizations",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
)
|
||||||
|
assert forbidden.status_code == 403
|
||||||
|
|
||||||
|
|
||||||
|
def test_create_update_and_delete_authorization(authorization_context):
|
||||||
|
context = authorization_context
|
||||||
|
avatar_id = context["avatar"].id
|
||||||
|
target_id = f"new-contact-{context['suffix']}"
|
||||||
|
created = client.post(
|
||||||
|
f"/api/avatar/{avatar_id}/authorizations",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
json={
|
||||||
|
"targetType": "user",
|
||||||
|
"targetId": target_id,
|
||||||
|
"targetName": "新联系人",
|
||||||
|
"permissions": ["friend", "chat", "browse"],
|
||||||
|
},
|
||||||
|
).json()
|
||||||
|
assert created["code"] == 200
|
||||||
|
authorization_id = created["data"]["id"]
|
||||||
|
assert created["data"]["permissions"] == ["friend", "chat", "browse"]
|
||||||
|
|
||||||
|
duplicate = client.post(
|
||||||
|
f"/api/avatar/{avatar_id}/authorizations",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
json={
|
||||||
|
"targetType": "user",
|
||||||
|
"targetId": target_id,
|
||||||
|
"targetName": "重复联系人",
|
||||||
|
"permissions": ["chat"],
|
||||||
|
},
|
||||||
|
).json()
|
||||||
|
assert duplicate["code"] == 409
|
||||||
|
|
||||||
|
updated = client.put(
|
||||||
|
f"/api/avatar/{avatar_id}/authorizations",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
json={
|
||||||
|
"id": authorization_id,
|
||||||
|
"targetName": "联系人新名称",
|
||||||
|
"permissions": ["interact", "publish"],
|
||||||
|
},
|
||||||
|
).json()
|
||||||
|
assert updated["code"] == 200
|
||||||
|
assert updated["data"]["targetName"] == "联系人新名称"
|
||||||
|
assert updated["data"]["permissions"] == ["publish", "interact"]
|
||||||
|
|
||||||
|
deleted = client.delete(
|
||||||
|
f"/api/avatar/{avatar_id}/authorizations/{authorization_id}",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
).json()
|
||||||
|
assert deleted["code"] == 200
|
||||||
|
assert deleted["data"]["id"] == authorization_id
|
||||||
|
|
||||||
|
|
||||||
|
def test_authorization_requires_login_and_rejects_unknown_permissions(authorization_context):
|
||||||
|
context = authorization_context
|
||||||
|
avatar_id = context["avatar"].id
|
||||||
|
no_session = client.get(f"/api/avatar/{avatar_id}/authorizations")
|
||||||
|
assert no_session.status_code == 401
|
||||||
|
|
||||||
|
invalid = client.post(
|
||||||
|
f"/api/avatar/{avatar_id}/authorizations",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
json={
|
||||||
|
"targetType": "user",
|
||||||
|
"targetId": "invalid-target",
|
||||||
|
"targetName": "无效权限",
|
||||||
|
"permissions": ["admin"],
|
||||||
|
},
|
||||||
|
).json()
|
||||||
|
assert invalid["code"] == 400
|
||||||
@@ -1,105 +1,125 @@
|
|||||||
"""Tests for PUT /api/avatar/{avatar_id}/authorizations/takeover endpoint."""
|
"""Tests for the authorization takeover configuration endpoint."""
|
||||||
|
|
||||||
from fastapi.testclient import TestClient
|
from fastapi.testclient import TestClient
|
||||||
|
|
||||||
|
from database import SessionLocal
|
||||||
from main import app
|
from main import app
|
||||||
from database import SessionLocal, Base, engine
|
from models import Authorization
|
||||||
from models import Authorization, Avatar
|
|
||||||
|
|
||||||
|
|
||||||
def setup_test_db():
|
client = TestClient(app)
|
||||||
Base.metadata.create_all(bind=engine)
|
|
||||||
|
|
||||||
|
def test_update_takeover_accepts_camel_case_and_persists(authorization_context):
|
||||||
|
context = authorization_context
|
||||||
|
response = client.put(
|
||||||
|
f"/api/avatar/{context['avatar'].id}/authorizations/takeover",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
json={
|
||||||
|
"authorizationId": context["authorization"].id,
|
||||||
|
"takeoverEnabled": True,
|
||||||
|
"takeoverMode": "delayed",
|
||||||
|
"takeoverDelaySeconds": 60,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
assert response.status_code == 200
|
||||||
|
payload = response.json()
|
||||||
|
assert payload["code"] == 200
|
||||||
|
assert payload["data"]["takeoverEnabled"] is True
|
||||||
|
assert payload["data"]["takeoverMode"] == "delayed"
|
||||||
|
assert payload["data"]["takeoverDelaySeconds"] == 60
|
||||||
|
assert "takeover" in payload["data"]["permissions"]
|
||||||
|
|
||||||
db = SessionLocal()
|
db = SessionLocal()
|
||||||
avatar = Avatar(name="test", status="active", config={})
|
|
||||||
db.add(avatar)
|
|
||||||
db.commit()
|
|
||||||
db.refresh(avatar)
|
|
||||||
auth = Authorization(avatar_id=avatar.id, target_id="user1", target_name="测试用户")
|
|
||||||
db.add(auth)
|
|
||||||
db.commit()
|
|
||||||
db.refresh(auth)
|
|
||||||
return db, auth.id
|
|
||||||
|
|
||||||
|
|
||||||
def test_update_takeover_config():
|
|
||||||
db, auth_id = setup_test_db()
|
|
||||||
try:
|
try:
|
||||||
client = TestClient(app)
|
stored = db.query(Authorization).filter(
|
||||||
response = client.put(
|
Authorization.id == context["authorization"].id
|
||||||
f"/api/avatar/test_avatar_id/authorizations/takeover",
|
).first()
|
||||||
json={
|
assert stored.takeover_enabled is True
|
||||||
"authorization_id": auth_id,
|
assert stored.takeover_mode == "delayed"
|
||||||
"takeover_enabled": True,
|
assert stored.takeover_delay_seconds == 60
|
||||||
"takeover_mode": "delayed",
|
|
||||||
"takeover_delay_seconds": 60,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
assert response.status_code == 200
|
|
||||||
data = response.json()
|
|
||||||
assert data["code"] == 200
|
|
||||||
assert data["data"]["takeoverEnabled"] is True
|
|
||||||
assert data["data"]["takeoverMode"] == "delayed"
|
|
||||||
assert data["data"]["takeoverDelaySeconds"] == 60
|
|
||||||
# 验证数据库已更新
|
|
||||||
auth = db.query(Authorization).filter(Authorization.id == auth_id).first()
|
|
||||||
assert auth.takeover_enabled is True
|
|
||||||
assert auth.takeover_mode == "delayed"
|
|
||||||
assert auth.takeover_delay_seconds == 60
|
|
||||||
finally:
|
finally:
|
||||||
db.close()
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
def test_update_takeover_invalid_mode():
|
def test_disabling_authorization_also_disables_takeover(authorization_context):
|
||||||
db, auth_id = setup_test_db()
|
context = authorization_context
|
||||||
try:
|
endpoint = f"/api/avatar/{context['avatar'].id}/authorizations/takeover"
|
||||||
client = TestClient(app)
|
client.put(
|
||||||
response = client.put(
|
endpoint,
|
||||||
f"/api/avatar/test/authorizations/takeover",
|
headers=context["owner_headers"],
|
||||||
json={
|
json={
|
||||||
"authorization_id": auth_id,
|
"authorizationId": context["authorization"].id,
|
||||||
"takeover_mode": "invalid_mode",
|
"takeoverEnabled": True,
|
||||||
},
|
},
|
||||||
)
|
|
||||||
assert response.status_code == 200
|
|
||||||
data = response.json()
|
|
||||||
assert data["code"] == 400
|
|
||||||
finally:
|
|
||||||
db.close()
|
|
||||||
|
|
||||||
|
|
||||||
def test_update_takeover_invalid_delay():
|
|
||||||
db, auth_id = setup_test_db()
|
|
||||||
try:
|
|
||||||
client = TestClient(app)
|
|
||||||
response = client.put(
|
|
||||||
f"/api/avatar/test/authorizations/takeover",
|
|
||||||
json={
|
|
||||||
"authorization_id": auth_id,
|
|
||||||
"takeover_delay_seconds": 2,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
assert response.status_code == 200
|
|
||||||
data = response.json()
|
|
||||||
assert data["code"] == 400
|
|
||||||
finally:
|
|
||||||
db.close()
|
|
||||||
|
|
||||||
|
|
||||||
def test_update_takeover_missing_auth_id():
|
|
||||||
client = TestClient(app)
|
|
||||||
response = client.put(
|
|
||||||
f"/api/avatar/test/authorizations/takeover",
|
|
||||||
json={"takeover_enabled": True},
|
|
||||||
)
|
)
|
||||||
assert response.status_code == 200
|
|
||||||
data = response.json()
|
updated = client.put(
|
||||||
assert data["code"] == 400
|
f"/api/avatar/{context['avatar'].id}/authorizations",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
json={"id": context["authorization"].id, "status": "inactive"},
|
||||||
|
).json()
|
||||||
|
assert updated["code"] == 200
|
||||||
|
assert updated["data"]["status"] == "inactive"
|
||||||
|
assert updated["data"]["takeoverEnabled"] is False
|
||||||
|
assert "takeover" not in updated["data"]["permissions"]
|
||||||
|
|
||||||
|
|
||||||
def test_update_takeover_not_found():
|
def test_takeover_rejects_invalid_values_and_cross_avatar_access(authorization_context):
|
||||||
client = TestClient(app)
|
context = authorization_context
|
||||||
response = client.put(
|
endpoint = f"/api/avatar/{context['avatar'].id}/authorizations/takeover"
|
||||||
f"/api/avatar/test/authorizations/takeover",
|
|
||||||
json={"authorization_id": "nonexistent"},
|
invalid_mode = client.put(
|
||||||
|
endpoint,
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
json={
|
||||||
|
"authorization_id": context["authorization"].id,
|
||||||
|
"takeover_mode": "invalid",
|
||||||
|
},
|
||||||
|
).json()
|
||||||
|
assert invalid_mode["code"] == 400
|
||||||
|
|
||||||
|
invalid_delay = client.put(
|
||||||
|
endpoint,
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
json={
|
||||||
|
"authorization_id": context["authorization"].id,
|
||||||
|
"takeover_delay_seconds": 2,
|
||||||
|
},
|
||||||
|
).json()
|
||||||
|
assert invalid_delay["code"] == 400
|
||||||
|
|
||||||
|
forbidden = client.put(
|
||||||
|
endpoint,
|
||||||
|
headers=context["other_headers"],
|
||||||
|
json={
|
||||||
|
"authorizationId": context["authorization"].id,
|
||||||
|
"takeoverEnabled": True,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
assert response.status_code == 200
|
assert forbidden.status_code == 403
|
||||||
data = response.json()
|
|
||||||
assert data["code"] == 404
|
|
||||||
|
def test_takeover_is_limited_to_active_user_authorizations(authorization_context):
|
||||||
|
context = authorization_context
|
||||||
|
avatar_id = context["avatar"].id
|
||||||
|
created = client.post(
|
||||||
|
f"/api/avatar/{avatar_id}/authorizations",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
json={
|
||||||
|
"targetType": "organization",
|
||||||
|
"targetId": f"org-{context['suffix']}",
|
||||||
|
"targetName": "测试组织",
|
||||||
|
"permissions": ["chat"],
|
||||||
|
},
|
||||||
|
).json()
|
||||||
|
response = client.put(
|
||||||
|
f"/api/avatar/{avatar_id}/authorizations/takeover",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
json={
|
||||||
|
"authorizationId": created["data"]["id"],
|
||||||
|
"takeoverEnabled": True,
|
||||||
|
},
|
||||||
|
).json()
|
||||||
|
assert response["code"] == 400
|
||||||
|
assert "单聊接管" in response["message"]
|
||||||
|
|||||||
@@ -192,7 +192,7 @@ async def test_process_message_delayed_mode(mock_db, mock_boxim, mock_auth):
|
|||||||
|
|
||||||
mock_auth.takeover_mode = "delayed"
|
mock_auth.takeover_mode = "delayed"
|
||||||
|
|
||||||
service = TakeoverService(mock_db, mock_boxim)
|
service = TakeoverService(mock_db, mock_boxim, MagicMock())
|
||||||
service.check_takeover_enabled = MagicMock(return_value=mock_auth)
|
service.check_takeover_enabled = MagicMock(return_value=mock_auth)
|
||||||
service.execute_takeover = AsyncMock()
|
service.execute_takeover = AsyncMock()
|
||||||
service.enqueue_delayed_message = MagicMock()
|
service.enqueue_delayed_message = MagicMock()
|
||||||
@@ -204,6 +204,24 @@ async def test_process_message_delayed_mode(mock_db, mock_boxim, mock_auth):
|
|||||||
service.execute_takeover.assert_not_awaited()
|
service.execute_takeover.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_process_message_delayed_mode_without_redis_falls_back_immediately(mock_db, mock_boxim, mock_auth):
|
||||||
|
"""A missing Redis connection must not silently drop delayed replies."""
|
||||||
|
from services.takeover_service import TakeoverService
|
||||||
|
|
||||||
|
mock_auth.takeover_mode = "delayed"
|
||||||
|
service = TakeoverService(mock_db, mock_boxim)
|
||||||
|
service.check_takeover_enabled = MagicMock(return_value=mock_auth)
|
||||||
|
service.execute_takeover = AsyncMock(return_value=True)
|
||||||
|
service.enqueue_delayed_message = MagicMock()
|
||||||
|
|
||||||
|
message = {"owner_huihui_id": "owner_1", "from_accid": "user_1", "content": "hi"}
|
||||||
|
await service.process_message(message)
|
||||||
|
|
||||||
|
service.execute_takeover.assert_awaited_once_with(mock_auth, message)
|
||||||
|
service.enqueue_delayed_message.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_process_message_no_takeover(mock_db, mock_boxim):
|
async def test_process_message_no_takeover(mock_db, mock_boxim):
|
||||||
"""When takeover is not enabled, nothing should happen."""
|
"""When takeover is not enabled, nothing should happen."""
|
||||||
|
|||||||
@@ -27,6 +27,9 @@ def mock_auth():
|
|||||||
auth.takeover_delay_seconds = 30
|
auth.takeover_delay_seconds = 30
|
||||||
auth.avatar_id = "avatar_123"
|
auth.avatar_id = "avatar_123"
|
||||||
auth.target_id = "target_user_123"
|
auth.target_id = "target_user_123"
|
||||||
|
auth.target_type = "user"
|
||||||
|
auth.status = "active"
|
||||||
|
auth.permissions = ["chat", "takeover"]
|
||||||
return auth
|
return auth
|
||||||
|
|
||||||
|
|
||||||
@@ -83,6 +86,7 @@ def test_check_takeover_enabled_returns_none_when_disabled(mock_db, mock_boxim,
|
|||||||
|
|
||||||
disabled_auth = MagicMock(spec=Authorization)
|
disabled_auth = MagicMock(spec=Authorization)
|
||||||
disabled_auth.takeover_enabled = False
|
disabled_auth.takeover_enabled = False
|
||||||
|
disabled_auth.permissions = []
|
||||||
auth_filter = MagicMock()
|
auth_filter = MagicMock()
|
||||||
auth_filter.filter.return_value = auth_filter
|
auth_filter.filter.return_value = auth_filter
|
||||||
auth_filter.first.return_value = disabled_auth
|
auth_filter.first.return_value = disabled_auth
|
||||||
|
|||||||
@@ -153,16 +153,41 @@ export interface Authorization {
|
|||||||
targetName: string
|
targetName: string
|
||||||
permissions: string[]
|
permissions: string[]
|
||||||
status: 'active' | 'inactive'
|
status: 'active' | 'inactive'
|
||||||
|
takeoverEnabled: boolean
|
||||||
|
takeoverMode: 'immediate' | 'delayed'
|
||||||
|
takeoverDelaySeconds: number
|
||||||
createdAt: string
|
createdAt: string
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export type AuthorizationInput = Pick<
|
||||||
|
Authorization,
|
||||||
|
'targetType' | 'targetId' | 'targetName' | 'permissions'
|
||||||
|
>
|
||||||
|
|
||||||
// 获取授权列表
|
// 获取授权列表
|
||||||
export const getAuthorizationList = (avatarId: string) =>
|
export const getAuthorizationList = (avatarId: string) =>
|
||||||
request.get<Authorization[]>(`/avatar/${avatarId}/authorizations`)
|
request.get<Authorization[]>(`/avatar/${avatarId}/authorizations`)
|
||||||
|
|
||||||
|
// 添加授权
|
||||||
|
export const createAuthorization = (avatarId: string, data: AuthorizationInput) =>
|
||||||
|
request.post<Authorization>(`/avatar/${avatarId}/authorizations`, data)
|
||||||
|
|
||||||
// 更新授权
|
// 更新授权
|
||||||
export const updateAuthorization = (avatarId: string, data: Partial<Authorization>) =>
|
export const updateAuthorization = (avatarId: string, data: Partial<Authorization> & { id: string }) =>
|
||||||
request.put(`/avatar/${avatarId}/authorizations`, data)
|
request.put<Authorization>(`/avatar/${avatarId}/authorizations`, data)
|
||||||
|
|
||||||
|
// 删除授权
|
||||||
|
export const deleteAuthorization = (avatarId: string, authorizationId: string) =>
|
||||||
|
request.delete<{ id: string }>(`/avatar/${avatarId}/authorizations/${authorizationId}`)
|
||||||
|
|
||||||
|
// 更新单聊接管配置
|
||||||
|
export const updateTakeoverConfig = (avatarId: string, data: {
|
||||||
|
authorizationId: string
|
||||||
|
takeoverEnabled: boolean
|
||||||
|
takeoverMode?: 'immediate' | 'delayed'
|
||||||
|
takeoverDelaySeconds?: number
|
||||||
|
}) =>
|
||||||
|
request.put<Authorization>(`/avatar/${avatarId}/authorizations/takeover`, data)
|
||||||
|
|
||||||
// ==================== 组织管理 API ====================
|
// ==================== 组织管理 API ====================
|
||||||
|
|
||||||
@@ -383,13 +408,4 @@ export const getCurrentUser = () =>
|
|||||||
export const logoutUser = () =>
|
export const logoutUser = () =>
|
||||||
request.post('/huihui/logout')
|
request.post('/huihui/logout')
|
||||||
|
|
||||||
// 更新接管配置
|
|
||||||
export const updateTakeoverConfig = (avatarId: string, data: {
|
|
||||||
authorizationId: string
|
|
||||||
takeoverEnabled: boolean
|
|
||||||
takeoverMode?: 'immediate' | 'delayed'
|
|
||||||
takeoverDelaySeconds?: number
|
|
||||||
}) =>
|
|
||||||
request.put(`/avatar/${avatarId}/authorizations/takeover`, data)
|
|
||||||
|
|
||||||
export default request
|
export default request
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user