diff --git a/digital-avatar-app/backend/main.py b/digital-avatar-app/backend/main.py index 341b751..bfb230f 100644 --- a/digital-avatar-app/backend/main.py +++ b/digital-avatar-app/backend/main.py @@ -18,6 +18,7 @@ import routers.organizations import routers.knowledge import routers.huihui_auth import routers.chat +import routers.takeover from responses import ok logger = logging.getLogger(__name__) @@ -39,6 +40,7 @@ app.include_router(routers.organizations.router, prefix="/api") app.include_router(routers.knowledge.router, prefix="/api") app.include_router(routers.huihui_auth.router, prefix="/api") app.include_router(routers.chat.router, prefix="/api") +app.include_router(routers.takeover.router, prefix="/api") UPLOAD_DIR = routers.knowledge.UPLOAD_DIR os.makedirs(UPLOAD_DIR, exist_ok=True) diff --git a/digital-avatar-app/backend/routers/takeover.py b/digital-avatar-app/backend/routers/takeover.py new file mode 100644 index 0000000..a9e8021 --- /dev/null +++ b/digital-avatar-app/backend/routers/takeover.py @@ -0,0 +1,41 @@ +"""分身接管配置 API""" +from fastapi import APIRouter, Depends, Body +from sqlalchemy.orm import Session + +from database import get_db +from models import Authorization +from responses import ok, fail + +router = APIRouter(tags=["分身接管"]) + + +@router.put("/avatar/{avatar_id}/authorizations/takeover") +def update_takeover_config( + avatar_id: str, + payload: dict = Body(...), + db: Session = Depends(get_db), +): + """更新分身接管配置""" + auth_id = payload.get("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) + + if "takeover_enabled" in payload: + auth.takeover_enabled = payload["takeover_enabled"] + if "takeover_mode" in payload: + mode = payload["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 + + db.commit() + return ok(auth.to_dict()) diff --git a/digital-avatar-app/backend/tests/test_takeover_api.py b/digital-avatar-app/backend/tests/test_takeover_api.py new file mode 100644 index 0000000..b7158d3 --- /dev/null +++ b/digital-avatar-app/backend/tests/test_takeover_api.py @@ -0,0 +1,105 @@ +"""Tests for PUT /api/avatar/{avatar_id}/authorizations/takeover endpoint.""" +from fastapi.testclient import TestClient +from main import app +from database import SessionLocal, Base, engine +from models import Authorization, Avatar + + +def setup_test_db(): + Base.metadata.create_all(bind=engine) + 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 + 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", + json={ + "authorization_id": auth_id, + "takeover_mode": "invalid_mode", + }, + ) + 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() + 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