Codex/avatar integrated 20260819 #2

Merged
stefanfeng merged 18 commits from codex/avatar-integrated-20260819 into main 2026-08-21 09:31:44 +08:00
10 changed files with 2106 additions and 614 deletions
Showing only changes of commit 2e2adeb9e2 - Show all commits
@@ -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}, "授权已删除")
+58 -20
View File
@@ -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)
+72 -1
View File
@@ -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
+27 -11
View File
@@ -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