Codex/avatar integrated 20260819 #2
@@ -1,32 +1,237 @@
|
||||
from fastapi import APIRouter, Depends, Body
|
||||
from fastapi import APIRouter, Body, Depends, Header, HTTPException
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from database import get_db
|
||||
from models import Authorization
|
||||
from responses import ok, fail
|
||||
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)
|
||||
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")
|
||||
def list_auth(avatar_id: str, db: Session = Depends(get_db)):
|
||||
# demo:返回全部授权(忽略具体 avatar 绑定,便于联调)
|
||||
items = db.query(Authorization).order_by(Authorization.created_at.desc()).all()
|
||||
return ok([a.to_dict() for a in items])
|
||||
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(...), db: Session = Depends(get_db)):
|
||||
auth_id = payload.get("id")
|
||||
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)
|
||||
a = db.query(Authorization).filter(Authorization.id == auth_id).first()
|
||||
if not a:
|
||||
return fail("授权不存在", 404)
|
||||
if "status" in payload:
|
||||
a.status = payload["status"]
|
||||
|
||||
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:
|
||||
a.permissions = payload["permissions"]
|
||||
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()
|
||||
items = db.query(Authorization).order_by(Authorization.created_at.desc()).all()
|
||||
return ok([x.to_dict() for x in items])
|
||||
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}, "授权已删除")
|
||||
|
||||
@@ -1,41 +1,79 @@
|
||||
"""分身接管配置 API"""
|
||||
from fastapi import APIRouter, Depends, Body
|
||||
"""数字分身单聊接管配置 API。"""
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, Header
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from database import get_db
|
||||
from models import Authorization
|
||||
from responses import ok, fail
|
||||
from responses import fail, ok
|
||||
from routers.authorizations import _require_authorization
|
||||
from routers.avatars import _require_owned_avatar
|
||||
|
||||
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")
|
||||
def update_takeover_config(
|
||||
avatar_id: str,
|
||||
payload: dict = Body(...),
|
||||
authorization: str = Header(None),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""更新分身接管配置"""
|
||||
auth_id = payload.get("authorizationId") or payload.get("authorization_id")
|
||||
_require_owned_avatar(db, avatar_id, authorization)
|
||||
auth_id = _read(payload, "authorizationId", "authorization_id")
|
||||
if not auth_id:
|
||||
return fail("缺少 authorization_id", 400)
|
||||
|
||||
auth = db.query(Authorization).filter(Authorization.id == auth_id).first()
|
||||
if not auth:
|
||||
return fail("授权不存在", 404)
|
||||
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
|
||||
|
||||
if "takeover_enabled" in payload:
|
||||
auth.takeover_enabled = payload["takeover_enabled"]
|
||||
if "takeover_mode" in payload:
|
||||
mode = payload["takeover_mode"]
|
||||
if _has(payload, "takeoverEnabled", "takeover_enabled"):
|
||||
raw_enabled = _read(payload, "takeoverEnabled", "takeover_enabled")
|
||||
if not isinstance(raw_enabled, bool):
|
||||
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"):
|
||||
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()
|
||||
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
|
||||
) -> Optional[Authorization]:
|
||||
"""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:
|
||||
return None
|
||||
|
||||
@@ -41,10 +45,13 @@ class TakeoverService:
|
||||
self.db.query(Authorization)
|
||||
.filter(Authorization.avatar_id == avatar.id)
|
||||
.filter(Authorization.target_id == from_user_id)
|
||||
.filter(Authorization.target_type == "user")
|
||||
.filter(Authorization.status == "active")
|
||||
.filter(Authorization.takeover_enabled == True)
|
||||
.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:
|
||||
"""Call the avatar chat endpoint to generate a reply."""
|
||||
@@ -170,5 +177,8 @@ class TakeoverService:
|
||||
|
||||
if auth.takeover_mode == "immediate":
|
||||
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:
|
||||
self.enqueue_delayed_message(auth, message)
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
import uuid
|
||||
|
||||
import pytest
|
||||
from database import init_db, SessionLocal
|
||||
from models import Authorization
|
||||
from models import Authorization, Avatar, User
|
||||
|
||||
|
||||
@pytest.fixture(scope="session", autouse=True)
|
||||
@@ -24,3 +26,72 @@ def setup_database():
|
||||
db.commit()
|
||||
finally:
|
||||
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 database import SessionLocal
|
||||
from main import app
|
||||
from database import SessionLocal, Base, engine
|
||||
from models import Authorization, Avatar
|
||||
from models import Authorization
|
||||
|
||||
|
||||
def setup_test_db():
|
||||
Base.metadata.create_all(bind=engine)
|
||||
client = TestClient(app)
|
||||
|
||||
|
||||
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()
|
||||
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:
|
||||
client = TestClient(app)
|
||||
response = client.put(
|
||||
f"/api/avatar/test_avatar_id/authorizations/takeover",
|
||||
json={
|
||||
"authorization_id": auth_id,
|
||||
"takeover_enabled": True,
|
||||
"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
|
||||
stored = db.query(Authorization).filter(
|
||||
Authorization.id == context["authorization"].id
|
||||
).first()
|
||||
assert stored.takeover_enabled is True
|
||||
assert stored.takeover_mode == "delayed"
|
||||
assert stored.takeover_delay_seconds == 60
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_update_takeover_invalid_mode():
|
||||
db, auth_id = setup_test_db()
|
||||
try:
|
||||
client = TestClient(app)
|
||||
response = client.put(
|
||||
f"/api/avatar/test/authorizations/takeover",
|
||||
def test_disabling_authorization_also_disables_takeover(authorization_context):
|
||||
context = authorization_context
|
||||
endpoint = f"/api/avatar/{context['avatar'].id}/authorizations/takeover"
|
||||
client.put(
|
||||
endpoint,
|
||||
headers=context["owner_headers"],
|
||||
json={
|
||||
"authorization_id": auth_id,
|
||||
"takeover_mode": "invalid_mode",
|
||||
"authorizationId": context["authorization"].id,
|
||||
"takeoverEnabled": True,
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["code"] == 400
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
updated = client.put(
|
||||
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_invalid_delay():
|
||||
db, auth_id = setup_test_db()
|
||||
try:
|
||||
client = TestClient(app)
|
||||
response = client.put(
|
||||
f"/api/avatar/test/authorizations/takeover",
|
||||
def test_takeover_rejects_invalid_values_and_cross_avatar_access(authorization_context):
|
||||
context = authorization_context
|
||||
endpoint = f"/api/avatar/{context['avatar'].id}/authorizations/takeover"
|
||||
|
||||
invalid_mode = client.put(
|
||||
endpoint,
|
||||
headers=context["owner_headers"],
|
||||
json={
|
||||
"authorization_id": auth_id,
|
||||
"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
|
||||
data = response.json()
|
||||
assert data["code"] == 400
|
||||
finally:
|
||||
db.close()
|
||||
assert forbidden.status_code == 403
|
||||
|
||||
|
||||
def test_update_takeover_missing_auth_id():
|
||||
client = TestClient(app)
|
||||
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/test/authorizations/takeover",
|
||||
json={"takeover_enabled": True},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["code"] == 400
|
||||
|
||||
|
||||
def test_update_takeover_not_found():
|
||||
client = TestClient(app)
|
||||
response = client.put(
|
||||
f"/api/avatar/test/authorizations/takeover",
|
||||
json={"authorization_id": "nonexistent"},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["code"] == 404
|
||||
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"
|
||||
|
||||
service = TakeoverService(mock_db, mock_boxim)
|
||||
service = TakeoverService(mock_db, mock_boxim, MagicMock())
|
||||
service.check_takeover_enabled = MagicMock(return_value=mock_auth)
|
||||
service.execute_takeover = AsyncMock()
|
||||
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()
|
||||
|
||||
|
||||
@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
|
||||
async def test_process_message_no_takeover(mock_db, mock_boxim):
|
||||
"""When takeover is not enabled, nothing should happen."""
|
||||
|
||||
@@ -27,6 +27,9 @@ def mock_auth():
|
||||
auth.takeover_delay_seconds = 30
|
||||
auth.avatar_id = "avatar_123"
|
||||
auth.target_id = "target_user_123"
|
||||
auth.target_type = "user"
|
||||
auth.status = "active"
|
||||
auth.permissions = ["chat", "takeover"]
|
||||
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.takeover_enabled = False
|
||||
disabled_auth.permissions = []
|
||||
auth_filter = MagicMock()
|
||||
auth_filter.filter.return_value = auth_filter
|
||||
auth_filter.first.return_value = disabled_auth
|
||||
|
||||
@@ -153,16 +153,41 @@ export interface Authorization {
|
||||
targetName: string
|
||||
permissions: string[]
|
||||
status: 'active' | 'inactive'
|
||||
takeoverEnabled: boolean
|
||||
takeoverMode: 'immediate' | 'delayed'
|
||||
takeoverDelaySeconds: number
|
||||
createdAt: string
|
||||
}
|
||||
|
||||
export type AuthorizationInput = Pick<
|
||||
Authorization,
|
||||
'targetType' | 'targetId' | 'targetName' | 'permissions'
|
||||
>
|
||||
|
||||
// 获取授权列表
|
||||
export const getAuthorizationList = (avatarId: string) =>
|
||||
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>) =>
|
||||
request.put(`/avatar/${avatarId}/authorizations`, data)
|
||||
export const updateAuthorization = (avatarId: string, data: Partial<Authorization> & { id: string }) =>
|
||||
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 ====================
|
||||
|
||||
@@ -383,13 +408,4 @@ export const getCurrentUser = () =>
|
||||
export const logoutUser = () =>
|
||||
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
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user