Compare commits
10
Commits
bd5f64d000
...
e2b928273c
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e2b928273c | ||
|
|
64b7680ec4 | ||
|
|
89f52963b7 | ||
|
|
76bd22c24b | ||
|
|
51a317ccd9 | ||
|
|
08590bf9ea | ||
|
|
25fb8fbee5 | ||
|
|
cfcfe7146e | ||
|
|
2e2adeb9e2 | ||
|
|
25d2494616 |
@@ -6,7 +6,6 @@ import logging
|
||||
|
||||
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
||||
from apscheduler.triggers.interval import IntervalTrigger
|
||||
import redis as redis_lib
|
||||
|
||||
from database import init_db, SessionLocal
|
||||
from models import Avatar, Authorization, Organization, TokenAccount, TokenPlan
|
||||
@@ -24,7 +23,6 @@ from responses import ok
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
takeover_scheduler = None
|
||||
takeover_db = None
|
||||
|
||||
app = FastAPI(title="会会数字分身 API", version="1.0.0")
|
||||
|
||||
@@ -113,7 +111,7 @@ def seed():
|
||||
|
||||
@app.on_event("startup")
|
||||
def on_startup():
|
||||
global takeover_scheduler, takeover_db
|
||||
global takeover_scheduler
|
||||
|
||||
init_db()
|
||||
seed()
|
||||
@@ -123,47 +121,43 @@ def on_startup():
|
||||
|
||||
# --- Takeover scheduler ---
|
||||
try:
|
||||
# Initialize Redis (optional)
|
||||
redis_client = None
|
||||
redis_url = os.getenv("REDIS_URL", "")
|
||||
if redis_url:
|
||||
try:
|
||||
redis_client = redis_lib.from_url(redis_url)
|
||||
redis_client.ping()
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis connection failed, delayed takeover will degrade to immediate: {e}")
|
||||
|
||||
# Initialize Box IM client
|
||||
# BOXIM production endpoints are intentionally separate from the login API.
|
||||
from services.boxim_client import BoxIMClient
|
||||
boxim_config = {
|
||||
"HUIHUI_IM_BASE_URL": os.getenv("HUIHUI_IM_BASE_URL", "http://192.168.1.200:60040"),
|
||||
"HUIHUI_PLATFORM_BASE_URL": os.getenv(
|
||||
"HUIHUI_PLATFORM_BASE_URL", "https://open.99hui.com/api"
|
||||
),
|
||||
"BOXIM_API_BASE_URL": os.getenv(
|
||||
"BOXIM_API_BASE_URL", "https://im.99hui.com/api"
|
||||
),
|
||||
"HUIHUI_APP_ID": os.getenv("HUIHUI_APP_ID", ""),
|
||||
"HUIHUI_ACCESS_ID": os.getenv("HUIHUI_ACCESS_ID", ""),
|
||||
"HUIHUI_ACCESS_SECRET": os.getenv("HUIHUI_ACCESS_SECRET", ""),
|
||||
"BOXIM_TIMEOUT_SECONDS": os.getenv("BOXIM_TIMEOUT_SECONDS", "20"),
|
||||
}
|
||||
boxim_client = BoxIMClient(boxim_config)
|
||||
|
||||
# Initialize takeover service
|
||||
from services.takeover_service import TakeoverService
|
||||
takeover_db = SessionLocal()
|
||||
takeover_service = TakeoverService(takeover_db, boxim_client, redis_client)
|
||||
takeover_service = TakeoverService(SessionLocal, boxim_client)
|
||||
|
||||
# AsyncIOScheduler awaits the service coroutine instead of dropping it.
|
||||
poll_interval = max(0.5, float(os.getenv("BOXIM_POLL_INTERVAL_SECONDS", "1")))
|
||||
takeover_scheduler = AsyncIOScheduler()
|
||||
takeover_scheduler.add_job(
|
||||
takeover_service.poll_and_process_messages,
|
||||
trigger=IntervalTrigger(seconds=10),
|
||||
trigger=IntervalTrigger(seconds=poll_interval),
|
||||
id="takeover_message_poll",
|
||||
max_instances=1,
|
||||
coalesce=True,
|
||||
)
|
||||
takeover_scheduler.start()
|
||||
logger.info("Takeover message polling scheduler started (interval=10s)")
|
||||
logger.info("BOXIM takeover scheduler started (interval=%ss)", poll_interval)
|
||||
except Exception as e:
|
||||
stop_takeover_scheduler()
|
||||
logger.warning(f"Failed to initialize takeover scheduler, app will continue without it: {e}")
|
||||
|
||||
|
||||
def stop_takeover_scheduler():
|
||||
global takeover_scheduler, takeover_db
|
||||
global takeover_scheduler
|
||||
|
||||
if takeover_scheduler is not None:
|
||||
try:
|
||||
@@ -174,13 +168,6 @@ def stop_takeover_scheduler():
|
||||
finally:
|
||||
takeover_scheduler = None
|
||||
|
||||
if takeover_db is not None:
|
||||
try:
|
||||
takeover_db.close()
|
||||
finally:
|
||||
takeover_db = None
|
||||
|
||||
|
||||
@app.on_event("shutdown")
|
||||
def on_shutdown():
|
||||
stop_takeover_scheduler()
|
||||
|
||||
@@ -1,6 +1,17 @@
|
||||
import uuid
|
||||
|
||||
from sqlalchemy import Column, String, Integer, Float, DateTime, Text, JSON, Boolean
|
||||
from sqlalchemy import (
|
||||
Boolean,
|
||||
Column,
|
||||
DateTime,
|
||||
Float,
|
||||
Index,
|
||||
Integer,
|
||||
JSON,
|
||||
String,
|
||||
Text,
|
||||
UniqueConstraint,
|
||||
)
|
||||
from sqlalchemy.sql import func
|
||||
|
||||
from database import Base
|
||||
@@ -74,6 +85,76 @@ class Authorization(Base):
|
||||
}
|
||||
|
||||
|
||||
class TakeoverCursor(Base):
|
||||
"""Durable BOXIM polling cursor for one avatar owner."""
|
||||
|
||||
__tablename__ = "takeover_cursors"
|
||||
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
|
||||
avatar_id = Column(String, nullable=False, unique=True, index=True)
|
||||
owner_id = Column(String, nullable=False, default="", index=True)
|
||||
boxim_owner_id = Column(String, default="")
|
||||
last_message_id = Column(String, default="0")
|
||||
initialized = Column(Boolean, default=False)
|
||||
last_polled_at = Column(DateTime)
|
||||
last_error = Column(Text, default="")
|
||||
created_at = Column(DateTime, server_default=func.now())
|
||||
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
|
||||
|
||||
|
||||
class TakeoverMessage(Base):
|
||||
"""BOXIM message receipt used for audit, deduplication, and chat context."""
|
||||
|
||||
__tablename__ = "takeover_messages"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("owner_id", "boxim_message_id", name="uq_takeover_message_owner_boxim"),
|
||||
Index("ix_takeover_message_conversation", "owner_id", "peer_id", "send_time"),
|
||||
)
|
||||
|
||||
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
|
||||
avatar_id = Column(String, nullable=False, index=True)
|
||||
owner_id = Column(String, nullable=False, index=True)
|
||||
boxim_message_id = Column(String, nullable=False)
|
||||
boxim_local_id = Column(String, nullable=True)
|
||||
peer_id = Column(String, nullable=False, index=True)
|
||||
direction = Column(String, nullable=False) # incoming | outgoing
|
||||
message_type = Column(Integer, default=0)
|
||||
content = Column(Text, default="")
|
||||
is_avatar = Column(Boolean, default=False)
|
||||
send_time = Column(DateTime, nullable=False)
|
||||
created_at = Column(DateTime, server_default=func.now())
|
||||
|
||||
|
||||
class TakeoverReplyTask(Base):
|
||||
"""Restart-safe three-second BOXIM reply task."""
|
||||
|
||||
__tablename__ = "takeover_reply_tasks"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("owner_id", "trigger_message_id", name="uq_takeover_task_owner_trigger"),
|
||||
Index("ix_takeover_task_due", "status", "scheduled_at"),
|
||||
Index("ix_takeover_task_conversation", "owner_id", "peer_id", "status"),
|
||||
)
|
||||
|
||||
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
|
||||
avatar_id = Column(String, nullable=False, index=True)
|
||||
owner_id = Column(String, nullable=False, index=True)
|
||||
peer_id = Column(String, nullable=False, index=True)
|
||||
trigger_message_id = Column(String, nullable=False)
|
||||
source_message_ids = Column(JSON, default=list)
|
||||
prompt = Column(Text, default="")
|
||||
response_text = Column(Text, default="")
|
||||
status = Column(String, default="pending")
|
||||
scheduled_at = Column(DateTime, nullable=False)
|
||||
locked_at = Column(DateTime)
|
||||
sent_at = Column(DateTime)
|
||||
attempts = Column(Integer, default=0)
|
||||
last_error = Column(Text, default="")
|
||||
cancel_reason = Column(String, default="")
|
||||
boxim_local_id = Column(String, nullable=False)
|
||||
boxim_sent_message_id = Column(String, default="")
|
||||
created_at = Column(DateTime, server_default=func.now())
|
||||
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
|
||||
|
||||
|
||||
class Organization(Base):
|
||||
__tablename__ = "organizations"
|
||||
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
|
||||
|
||||
@@ -7,5 +7,4 @@ httpx
|
||||
pypdf
|
||||
python-docx
|
||||
openpyxl
|
||||
redis>=5.0
|
||||
apscheduler>=3.10
|
||||
|
||||
@@ -1,32 +1,334 @@
|
||||
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 models import Authorization, TakeoverCursor, TakeoverReplyTask
|
||||
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)
|
||||
AVATAR_PERMISSION_ORDER = PERMISSION_ORDER
|
||||
AVATAR_PERMISSION_KEY = "authorizationPermissions"
|
||||
DEFAULT_AVATAR_PERMISSIONS = ["friend", "chat"]
|
||||
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 _normalize_avatar_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 AVATAR_PERMISSION_ORDER:
|
||||
raise ValueError(f"不支持的权限:{raw}")
|
||||
if permission not in normalized:
|
||||
normalized.append(permission)
|
||||
return sorted(normalized, key=AVATAR_PERMISSION_ORDER.index)
|
||||
|
||||
|
||||
def _stored_avatar_permissions(avatar) -> list[str]:
|
||||
config = avatar.config or {}
|
||||
if AVATAR_PERMISSION_KEY not in config:
|
||||
return list(DEFAULT_AVATAR_PERMISSIONS)
|
||||
|
||||
stored = config.get(AVATAR_PERMISSION_KEY)
|
||||
if not isinstance(stored, list):
|
||||
return list(DEFAULT_AVATAR_PERMISSIONS)
|
||||
|
||||
permissions = []
|
||||
for raw in stored:
|
||||
permission = LEGACY_PERMISSION_MAP.get(str(raw).strip(), str(raw).strip())
|
||||
if permission in AVATAR_PERMISSION_ORDER and permission not in permissions:
|
||||
permissions.append(permission)
|
||||
return sorted(permissions, key=AVATAR_PERMISSION_ORDER.index)
|
||||
|
||||
|
||||
def _permission_settings_payload(avatar) -> dict:
|
||||
return {
|
||||
"avatarId": avatar.id,
|
||||
"permissions": _stored_avatar_permissions(avatar),
|
||||
}
|
||||
|
||||
|
||||
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}/permission-settings")
|
||||
def get_permission_settings(
|
||||
avatar_id: str,
|
||||
authorization: str = Header(None),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
avatar = _require_owned_avatar(db, avatar_id, authorization)
|
||||
return ok(_permission_settings_payload(avatar))
|
||||
|
||||
|
||||
@router.put("/avatar/{avatar_id}/permission-settings")
|
||||
def update_permission_settings(
|
||||
avatar_id: str,
|
||||
payload: dict = Body(...),
|
||||
authorization: str = Header(None),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
avatar = _require_owned_avatar(db, avatar_id, authorization)
|
||||
if "permissions" not in payload:
|
||||
return fail("缺少 permissions", 400)
|
||||
try:
|
||||
permissions = _normalize_avatar_permissions(payload["permissions"])
|
||||
except ValueError as exc:
|
||||
return fail(str(exc), 400)
|
||||
|
||||
previous_permissions = _stored_avatar_permissions(avatar)
|
||||
avatar.config = {
|
||||
**(avatar.config or {}),
|
||||
AVATAR_PERMISSION_KEY: permissions,
|
||||
}
|
||||
cursor = db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar.id).first()
|
||||
if cursor and "takeover" in permissions and "takeover" not in previous_permissions:
|
||||
cursor.initialized = False
|
||||
cursor.last_message_id = "0"
|
||||
cursor.last_error = ""
|
||||
elif cursor and "takeover" not in permissions:
|
||||
cursor.last_error = ""
|
||||
|
||||
if "takeover" not in permissions:
|
||||
tasks = (
|
||||
db.query(TakeoverReplyTask)
|
||||
.filter(
|
||||
TakeoverReplyTask.avatar_id == avatar.id,
|
||||
TakeoverReplyTask.status.in_(("pending", "generating", "ready", "sending")),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
for task in tasks:
|
||||
task.status = "cancelled"
|
||||
task.cancel_reason = "takeover_disabled"
|
||||
task.locked_at = None
|
||||
db.commit()
|
||||
db.refresh(avatar)
|
||||
return ok(_permission_settings_payload(avatar), "授权设置已保存")
|
||||
|
||||
|
||||
@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)
|
||||
|
||||
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:
|
||||
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:
|
||||
a.status = payload["status"]
|
||||
if "permissions" in payload:
|
||||
a.permissions = payload["permissions"]
|
||||
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}, "授权已删除")
|
||||
|
||||
@@ -27,7 +27,7 @@ from sqlalchemy.orm import Session
|
||||
_CN_TZ = timezone(timedelta(hours=8))
|
||||
|
||||
from database import get_db
|
||||
from models import User
|
||||
from models import Avatar, TakeoverCursor, TakeoverMessage, TakeoverReplyTask, User
|
||||
from responses import ok, fail
|
||||
|
||||
router = APIRouter(tags=["会会账号"])
|
||||
@@ -281,12 +281,61 @@ def pwd_login(body: dict = Body(...), db: Session = Depends(get_db)):
|
||||
})
|
||||
|
||||
|
||||
def _transfer_avatar_ownership(db: Session, old_owner_id: str, new_owner_id: str) -> int:
|
||||
"""Move one user's avatar-owned data to a replacement Huihui identity."""
|
||||
if not old_owner_id or old_owner_id == new_owner_id:
|
||||
return 0
|
||||
|
||||
avatar_ids = [
|
||||
avatar_id
|
||||
for (avatar_id,) in db.query(Avatar.id).filter(Avatar.owner_id == old_owner_id).all()
|
||||
]
|
||||
if not avatar_ids:
|
||||
return 0
|
||||
|
||||
db.query(Avatar).filter(Avatar.id.in_(avatar_ids)).update(
|
||||
{Avatar.owner_id: new_owner_id}, synchronize_session="fetch"
|
||||
)
|
||||
for model in (TakeoverCursor, TakeoverMessage, TakeoverReplyTask):
|
||||
db.query(model).filter(model.avatar_id.in_(avatar_ids)).update(
|
||||
{model.owner_id: new_owner_id}, synchronize_session="fetch"
|
||||
)
|
||||
return len(avatar_ids)
|
||||
|
||||
|
||||
def _find_or_link_user(db: Session, phone: str, huihui_user_id: str) -> User:
|
||||
"""Resolve an account and safely retain avatars across Huihui environments."""
|
||||
user = db.query(User).filter(User.huihui_user_id == huihui_user_id).first()
|
||||
if not phone:
|
||||
return user or User(huihui_user_id=huihui_user_id)
|
||||
|
||||
same_phone_users = db.query(User).filter(User.phone == phone).all()
|
||||
|
||||
if user is None:
|
||||
# A unique verified-phone match is the same person whose upstream ID changed.
|
||||
if len(same_phone_users) == 1:
|
||||
user = same_phone_users[0]
|
||||
old_owner_id = user.huihui_user_id
|
||||
_transfer_avatar_ownership(db, old_owner_id, huihui_user_id)
|
||||
user.huihui_user_id = huihui_user_id
|
||||
return user
|
||||
return User(huihui_user_id=huihui_user_id)
|
||||
|
||||
legacy_users = [candidate for candidate in same_phone_users if candidate.id != user.id]
|
||||
current_avatar_count = db.query(Avatar).filter(Avatar.owner_id == huihui_user_id).count()
|
||||
if len(legacy_users) == 1 and current_avatar_count == 0:
|
||||
legacy_user = legacy_users[0]
|
||||
_transfer_avatar_ownership(db, legacy_user.huihui_user_id, huihui_user_id)
|
||||
legacy_user.app_token = ""
|
||||
legacy_user.huihui_token = ""
|
||||
db.add(legacy_user)
|
||||
return user
|
||||
|
||||
|
||||
def _issue_session(db: Session, phone: str, info: dict):
|
||||
"""建/链本地用户并签发本系统会话 token"""
|
||||
huihui_user_id = info.get("userId", "")
|
||||
user = db.query(User).filter(User.huihui_user_id == huihui_user_id).first()
|
||||
if not user:
|
||||
user = User(huihui_user_id=huihui_user_id)
|
||||
user = _find_or_link_user(db, phone, huihui_user_id)
|
||||
if phone:
|
||||
user.phone = phone
|
||||
if info.get("nickname"):
|
||||
|
||||
@@ -32,6 +32,14 @@ class EnabledIn(BaseModel):
|
||||
enabled: bool = True
|
||||
|
||||
|
||||
def _doc_payload(doc: KnowledgeDoc) -> dict:
|
||||
payload = doc.to_dict()
|
||||
stored_name = os.path.basename(doc.file_url or "")
|
||||
stored_path = os.path.join(UPLOAD_DIR, doc.avatar_id, stored_name)
|
||||
payload["filePresent"] = bool(stored_name and os.path.isfile(stored_path))
|
||||
return payload
|
||||
|
||||
|
||||
def _resolve_user(authorization: str | None, db: Session):
|
||||
if not authorization:
|
||||
return None
|
||||
@@ -61,7 +69,7 @@ def list_docs(avatar_id: str, authorization: str = Header(None), db: Session = D
|
||||
.order_by(KnowledgeDoc.created_at.desc())
|
||||
.all()
|
||||
)
|
||||
return ok([d.to_dict() for d in docs])
|
||||
return ok([_doc_payload(d) for d in docs])
|
||||
|
||||
|
||||
@router.post("/avatar/{avatar_id}/knowledge/docs")
|
||||
@@ -121,7 +129,7 @@ async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization
|
||||
db.commit()
|
||||
db.refresh(doc)
|
||||
|
||||
return ok(doc.to_dict())
|
||||
return ok(_doc_payload(doc))
|
||||
|
||||
|
||||
@router.delete("/avatar/{avatar_id}/knowledge/docs/{doc_id}")
|
||||
|
||||
@@ -1,41 +1,132 @@
|
||||
"""分身接管配置 API"""
|
||||
from fastapi import APIRouter, Depends, Body
|
||||
"""数字分身 BOXIM 单聊接管 API。"""
|
||||
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
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 models import TakeoverCursor, TakeoverReplyTask, User
|
||||
from responses import fail, ok
|
||||
from routers.authorizations import _require_authorization
|
||||
from routers.avatars import _require_owned_avatar
|
||||
|
||||
router = APIRouter(tags=["分身接管"])
|
||||
BOXIM_STATUS_FRESH_SECONDS = 60
|
||||
|
||||
|
||||
@router.get("/avatar/{avatar_id}/takeover/status")
|
||||
def get_takeover_status(
|
||||
avatar_id: str,
|
||||
authorization: str = Header(None),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
avatar = _require_owned_avatar(db, avatar_id, authorization)
|
||||
permissions = (avatar.config or {}).get("authorizationPermissions", [])
|
||||
enabled = isinstance(permissions, list) and "takeover" in permissions
|
||||
user = db.query(User).filter(User.huihui_user_id == avatar.owner_id).first()
|
||||
cursor = db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar.id).first()
|
||||
pending_count = (
|
||||
db.query(TakeoverReplyTask)
|
||||
.filter(
|
||||
TakeoverReplyTask.avatar_id == avatar.id,
|
||||
TakeoverReplyTask.status.in_(("pending", "generating", "ready", "sending")),
|
||||
)
|
||||
.count()
|
||||
)
|
||||
|
||||
if cursor and cursor.last_error:
|
||||
status, message = "error", cursor.last_error
|
||||
elif not enabled:
|
||||
status, message = "disabled", "主动接管未开启"
|
||||
elif not user or not user.huihui_token:
|
||||
status, message = "needs_login", "请重新登录会会生产账号以连接 BOXIM"
|
||||
elif (
|
||||
cursor
|
||||
and cursor.initialized
|
||||
and cursor.last_polled_at
|
||||
# BOXIM offline-message reads can long-poll for about 20 seconds.
|
||||
and cursor.last_polled_at
|
||||
>= datetime.utcnow() - timedelta(seconds=BOXIM_STATUS_FRESH_SECONDS)
|
||||
):
|
||||
status, message = "ready", "BOXIM 已连接,收到私聊消息 3 秒后自动回复"
|
||||
else:
|
||||
status, message = "connecting", "正在连接 BOXIM"
|
||||
|
||||
return ok(
|
||||
{
|
||||
"enabled": enabled,
|
||||
"status": status,
|
||||
"message": message,
|
||||
"pendingCount": pending_count,
|
||||
"lastPolledAt": cursor.last_polled_at.isoformat() if cursor and cursor.last_polled_at else None,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
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(), "接管配置已保存")
|
||||
|
||||
@@ -1,76 +1,210 @@
|
||||
"""盒子 IM 客户端 — 封装网易云信 IM 接口调用"""
|
||||
"""Client for Huihui's self-hosted BOXIM production APIs."""
|
||||
|
||||
import hashlib
|
||||
import random
|
||||
import secrets
|
||||
import string
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
import time
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
_CN_TZ = timezone(timedelta(hours=8))
|
||||
|
||||
|
||||
class BoxIMError(RuntimeError):
|
||||
def __init__(self, message: str, *, code: Any = None, auth_error: bool = False):
|
||||
super().__init__(message)
|
||||
self.code = code
|
||||
self.auth_error = auth_error
|
||||
|
||||
|
||||
class BoxIMClient:
|
||||
"""盒子 IM 客户端,通过会会平台网关调用网易云信 IM"""
|
||||
"""Exchange Huihui credentials and call BOXIM's private-message API."""
|
||||
|
||||
def __init__(self, config: dict):
|
||||
self.base_url = config.get("HUIHUI_IM_BASE_URL", "http://192.168.1.200:60040")
|
||||
self.platform_base_url = config.get(
|
||||
"HUIHUI_PLATFORM_BASE_URL", "https://open.99hui.com/api"
|
||||
).rstrip("/")
|
||||
self.im_base_url = config.get(
|
||||
"BOXIM_API_BASE_URL", "https://im.99hui.com/api"
|
||||
).rstrip("/")
|
||||
self.app_id = config.get("HUIHUI_APP_ID", "")
|
||||
self.access_id = config.get("HUIHUI_ACCESS_ID", "")
|
||||
self.access_secret = config.get("HUIHUI_ACCESS_SECRET", "")
|
||||
self.timeout = float(config.get("BOXIM_TIMEOUT_SECONDS", 20))
|
||||
|
||||
def _build_sign_params(self, extra: dict) -> dict:
|
||||
"""构建带签名的请求参数(复用 news_service 签名模式)"""
|
||||
nonce = "".join(random.choices(string.ascii_lowercase + string.digits, k=12))
|
||||
timestamp = datetime.now().strftime("%Y%m%d%H%M%S") # 24小时制
|
||||
def _build_sign_params(self, extra: dict | None = None) -> dict:
|
||||
"""Build the same signed form used by Huihui's current production app."""
|
||||
params = {
|
||||
"appId": self.app_id,
|
||||
"accessId": self.access_id,
|
||||
"nonce": nonce,
|
||||
"timestamp": timestamp,
|
||||
**extra,
|
||||
"nonce": "".join(random.choices(string.ascii_lowercase + string.digits, k=12)),
|
||||
"timestamp": datetime.now(_CN_TZ).strftime("%Y%m%d%H%M%S"),
|
||||
"signType": "MD5",
|
||||
"signVersion": "1.0",
|
||||
**(extra or {}),
|
||||
}
|
||||
# 计算签名 — 排序 key, 过滤空值, 拼接后加 accessSecret, MD5 大写
|
||||
keys = sorted(params.keys())
|
||||
params.pop("accessSecret", None)
|
||||
params.pop("signature", None)
|
||||
sign_parts = []
|
||||
for k in keys:
|
||||
if k in ("signature", "accessSecret"):
|
||||
for key in sorted(params):
|
||||
value = params[key]
|
||||
if value in (None, "", []):
|
||||
continue
|
||||
v = params.get(k)
|
||||
if v and v != "" and v != []:
|
||||
sign_parts.append(f"{k}={v}")
|
||||
sign_str = "&".join(sign_parts) + f"&accessSecret={self.access_secret}"
|
||||
signature = hashlib.md5(sign_str.encode("utf-8")).hexdigest().upper()
|
||||
params["signature"] = signature
|
||||
params["signType"] = "MD5"
|
||||
params["signVersion"] = "1.0"
|
||||
if isinstance(value, list):
|
||||
continue
|
||||
sign_parts.append(f"{key}={value}")
|
||||
sign_source = "&".join(sign_parts) + f"&accessSecret={self.access_secret}"
|
||||
params["signature"] = hashlib.md5(sign_source.encode("utf-8")).hexdigest().upper()
|
||||
return params
|
||||
|
||||
async def get_credentials(self, user_id: str) -> Optional[dict]:
|
||||
"""获取用户的网易云信 IM 凭证 (accid, token)"""
|
||||
params = self._build_sign_params({"userId": user_id})
|
||||
async with httpx.AsyncClient(timeout=10) as client:
|
||||
r = await client.post(
|
||||
f"{self.base_url}/box/netease",
|
||||
params=params,
|
||||
)
|
||||
data = r.json()
|
||||
if data.get("code") in (0, 200):
|
||||
return data.get("data", {})
|
||||
return None
|
||||
@staticmethod
|
||||
def _response_payload(response: httpx.Response) -> dict:
|
||||
try:
|
||||
payload = response.json()
|
||||
except ValueError as exc:
|
||||
raise BoxIMError("BOXIM 返回了无效响应") from exc
|
||||
if not isinstance(payload, dict):
|
||||
raise BoxIMError("BOXIM 返回格式不正确")
|
||||
return payload
|
||||
|
||||
async def send_p2p_message(
|
||||
self, from_accid: str, to_accid: str, content: str
|
||||
) -> bool:
|
||||
"""发送单聊消息(文本)"""
|
||||
params = self._build_sign_params({
|
||||
"from": from_accid,
|
||||
"to": to_accid,
|
||||
"msgType": "text",
|
||||
"content": content,
|
||||
})
|
||||
async with httpx.AsyncClient(timeout=10) as client:
|
||||
r = await client.post(
|
||||
f"{self.base_url}/box/message/send/p2p",
|
||||
params=params,
|
||||
async def exchange_access_token(self, huihui_token: str) -> dict:
|
||||
"""Exchange a production Huihui token for a BOXIM access token."""
|
||||
if not huihui_token:
|
||||
raise BoxIMError("缺少会会登录凭证", auth_error=True)
|
||||
if not (self.app_id and self.access_id and self.access_secret):
|
||||
raise BoxIMError("会会开放平台凭证未配置", auth_error=True)
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {huihui_token}",
|
||||
"appId": self.app_id,
|
||||
"windowAppId": self.app_id,
|
||||
}
|
||||
async with httpx.AsyncClient(timeout=self.timeout, follow_redirects=True) as client:
|
||||
response = await client.post(
|
||||
f"{self.platform_base_url}/im/box/netease",
|
||||
headers=headers,
|
||||
data=self._build_sign_params(),
|
||||
)
|
||||
data = r.json()
|
||||
return data.get("code") in (0, 200)
|
||||
payload = self._response_payload(response)
|
||||
data = payload.get("data") or {}
|
||||
code = payload.get("code")
|
||||
if response.status_code >= 400 or code not in (0, 200, "0", "200"):
|
||||
raise BoxIMError(
|
||||
payload.get("message") or "BOXIM 授权失败",
|
||||
code=code or response.status_code,
|
||||
auth_error=response.status_code in (400, 401, 403)
|
||||
or code in (
|
||||
400,
|
||||
401,
|
||||
40100,
|
||||
40101,
|
||||
403,
|
||||
"400",
|
||||
"401",
|
||||
"40100",
|
||||
"40101",
|
||||
"403",
|
||||
),
|
||||
)
|
||||
if not data.get("accessToken"):
|
||||
raise BoxIMError("会会未返回 BOXIM 访问凭证", auth_error=True)
|
||||
return data
|
||||
|
||||
async def _request(
|
||||
self,
|
||||
method: str,
|
||||
path: str,
|
||||
access_token: str,
|
||||
*,
|
||||
params: dict | None = None,
|
||||
json: dict | None = None,
|
||||
) -> Any:
|
||||
headers = {"accessToken": access_token}
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.request(
|
||||
method,
|
||||
f"{self.im_base_url}{path}",
|
||||
headers=headers,
|
||||
params=params,
|
||||
json=json,
|
||||
)
|
||||
payload = self._response_payload(response)
|
||||
code = payload.get("code")
|
||||
if response.status_code >= 400 or code not in (200, "200"):
|
||||
raise BoxIMError(
|
||||
payload.get("message") or "BOXIM 请求失败",
|
||||
code=code or response.status_code,
|
||||
auth_error=response.status_code in (400, 401, 403)
|
||||
or code in (400, 401, 40100, 40101, 403, "400", "401", "40100", "40101", "403"),
|
||||
)
|
||||
return payload.get("data")
|
||||
|
||||
async def get_self(self, access_token: str) -> dict:
|
||||
data = await self._request("GET", "/user/self", access_token)
|
||||
if not isinstance(data, dict) or data.get("id") is None:
|
||||
raise BoxIMError("BOXIM 未返回当前用户信息")
|
||||
return data
|
||||
|
||||
async def fetch_private_messages(self, access_token: str, min_id: str = "0") -> list[dict]:
|
||||
data = await self._request(
|
||||
"GET",
|
||||
"/message/private/loadOfflineMessage",
|
||||
access_token,
|
||||
params={"minId": str(min_id or "0")},
|
||||
)
|
||||
if data is None:
|
||||
return []
|
||||
if not isinstance(data, list):
|
||||
raise BoxIMError("BOXIM 私聊消息格式不正确")
|
||||
return [item for item in data if isinstance(item, dict)]
|
||||
|
||||
async def mark_private_messages_read(
|
||||
self,
|
||||
access_token: str,
|
||||
friend_id: int | str,
|
||||
message_id: int | str,
|
||||
) -> None:
|
||||
"""Mark one private conversation read through its latest received message."""
|
||||
friend_id_text = str(friend_id).strip()
|
||||
message_id_text = str(message_id).strip()
|
||||
if not friend_id_text.isdigit() or not message_id_text.isdigit():
|
||||
raise BoxIMError("BOXIM 已读回执参数不正确")
|
||||
await self._request(
|
||||
"PUT",
|
||||
"/message/private/readed",
|
||||
access_token,
|
||||
params={
|
||||
"friendId": int(friend_id_text),
|
||||
"messageId": int(message_id_text),
|
||||
},
|
||||
)
|
||||
|
||||
async def send_private_message(
|
||||
self,
|
||||
access_token: str,
|
||||
peer_id: str,
|
||||
content: str,
|
||||
*,
|
||||
local_id: int | str | None = None,
|
||||
) -> dict:
|
||||
local_id = int(local_id or (int(time.time() * 1000) * 1000 + secrets.randbelow(1000)))
|
||||
data = await self._request(
|
||||
"POST",
|
||||
"/message/private/send",
|
||||
access_token,
|
||||
json={
|
||||
"localId": local_id,
|
||||
"recvId": int(peer_id) if str(peer_id).isdigit() else peer_id,
|
||||
"content": content,
|
||||
"type": 0,
|
||||
"receipt": False,
|
||||
"atUserIds": [],
|
||||
},
|
||||
)
|
||||
if not isinstance(data, dict):
|
||||
raise BoxIMError("BOXIM 未返回发送结果")
|
||||
return data
|
||||
|
||||
@@ -1,174 +1,613 @@
|
||||
"""Takeover service — message listening, decision, reply execution."""
|
||||
import json
|
||||
"""Restart-safe automatic replies over Huihui's self-hosted BOXIM."""
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import logging
|
||||
import os
|
||||
from typing import Optional
|
||||
import re
|
||||
import secrets
|
||||
import time
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Callable
|
||||
|
||||
import httpx
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from models import Avatar, Authorization
|
||||
from services.boxim_client import BoxIMClient
|
||||
from models import (
|
||||
Avatar,
|
||||
TakeoverCursor,
|
||||
TakeoverMessage,
|
||||
TakeoverReplyTask,
|
||||
User,
|
||||
)
|
||||
from services.boxim_client import BoxIMClient, BoxIMError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
ACTIVE_TASK_STATUSES = ("pending", "generating", "ready", "sending")
|
||||
GENERATABLE_TASK_STATUSES = ("pending",)
|
||||
MAX_PROMPT_LENGTH = 4000
|
||||
MAX_STALE_SECONDS = 120
|
||||
STUCK_LOCK_SECONDS = 90
|
||||
TAKEOVER_PERMISSION = "takeover"
|
||||
|
||||
|
||||
def _utcnow() -> datetime:
|
||||
return datetime.utcnow()
|
||||
|
||||
|
||||
def _takeover_enabled(avatar: Avatar | None) -> bool:
|
||||
if not avatar or avatar.status != "active":
|
||||
return False
|
||||
permissions = (avatar.config or {}).get("authorizationPermissions", [])
|
||||
return isinstance(permissions, list) and TAKEOVER_PERMISSION in permissions
|
||||
|
||||
|
||||
def _boxim_time(value, fallback: datetime) -> datetime:
|
||||
try:
|
||||
timestamp = float(value)
|
||||
if timestamp > 10_000_000_000:
|
||||
timestamp /= 1000
|
||||
return datetime.utcfromtimestamp(timestamp)
|
||||
except (TypeError, ValueError, OSError, OverflowError):
|
||||
return fallback
|
||||
|
||||
|
||||
def _numeric_id(value) -> int:
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError):
|
||||
return 0
|
||||
|
||||
|
||||
def _plain_text_reply(value: str) -> str:
|
||||
"""BOXIM is plain text, so remove Markdown markers without damaging paragraphs."""
|
||||
text = (value or "").replace("\r\n", "\n").replace("\r", "\n")
|
||||
text = re.sub(r"```(?:\w+)?\n?(.*?)```", r"\1", text, flags=re.S)
|
||||
text = re.sub(r"\*\*(.*?)\*\*|__(.*?)__", lambda m: m.group(1) or m.group(2), text)
|
||||
text = re.sub(r"(?<!\*)\*([^*\n]+)\*(?!\*)", r"\1", text)
|
||||
text = re.sub(r"`([^`]+)`", r"\1", text)
|
||||
text = re.sub(r"^\s{0,3}#{1,6}\s*", "", text, flags=re.M)
|
||||
lines = [line.strip() for line in text.split("\n")]
|
||||
return "\n".join(line for line in lines if line).strip()
|
||||
|
||||
|
||||
class TakeoverService:
|
||||
"""Service for handling avatar takeover — generating replies and sending them via IM."""
|
||||
"""Poll BOXIM, prepare replies during the grace period, then send at +3s."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
db: Session,
|
||||
session_factory: Callable[[], Session],
|
||||
boxim_client: BoxIMClient,
|
||||
redis_client=None,
|
||||
*,
|
||||
reply_delay_seconds: int = 3,
|
||||
now: Callable[[], datetime] = _utcnow,
|
||||
):
|
||||
self.db = db
|
||||
self.session_factory = session_factory
|
||||
self.boxim = boxim_client
|
||||
self.redis = redis_client
|
||||
self._chat_api_base = os.getenv(
|
||||
"TAKEOVER_CHAT_API_BASE", "http://localhost:8000/api"
|
||||
)
|
||||
|
||||
def check_takeover_enabled(
|
||||
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()
|
||||
if not avatar:
|
||||
return None
|
||||
|
||||
auth = (
|
||||
self.db.query(Authorization)
|
||||
.filter(Authorization.avatar_id == avatar.id)
|
||||
.filter(Authorization.target_id == from_user_id)
|
||||
.filter(Authorization.takeover_enabled == True)
|
||||
.first()
|
||||
)
|
||||
return auth if auth and auth.takeover_enabled else None
|
||||
|
||||
async def generate_reply(self, avatar_id: str, message: str) -> str:
|
||||
"""Call the avatar chat endpoint to generate a reply."""
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=30) as client:
|
||||
r = await client.post(
|
||||
f"{self._chat_api_base}/avatar/{avatar_id}/chat",
|
||||
json={"message": message, "history": []},
|
||||
)
|
||||
data = r.json()
|
||||
if data.get("code") in (0, 200):
|
||||
return data.get("data", {}).get("answer", "")
|
||||
logger.warning(f"Avatar chat API returned error code: {data}")
|
||||
return ""
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to call avatar chat API: {e}")
|
||||
return ""
|
||||
|
||||
async def execute_takeover(self, auth: Authorization, message: dict) -> bool:
|
||||
"""Execute takeover: generate a reply and send it as the owner via IM."""
|
||||
try:
|
||||
# Resolve owner through Avatar model
|
||||
avatar = self.db.query(Avatar).filter(Avatar.id == auth.avatar_id).first()
|
||||
if not avatar:
|
||||
logger.warning(f"Avatar not found: {auth.avatar_id}")
|
||||
return False
|
||||
|
||||
owner_huihui_id = avatar.owner_id
|
||||
credentials = await self.boxim.get_credentials(owner_huihui_id)
|
||||
if not credentials:
|
||||
logger.warning(f"Cannot obtain IM credentials for owner: {owner_huihui_id}")
|
||||
return False
|
||||
|
||||
reply = await self.generate_reply(auth.avatar_id, message.get("content", ""))
|
||||
if not reply:
|
||||
logger.warning("Avatar did not generate a reply")
|
||||
return False
|
||||
|
||||
success = await self.boxim.send_p2p_message(
|
||||
from_accid=credentials["accid"],
|
||||
to_accid=message.get("from_accid", ""),
|
||||
content=reply,
|
||||
)
|
||||
if success:
|
||||
logger.info(f"Takeover reply sent successfully: {reply[:50]}...")
|
||||
return success
|
||||
except Exception as e:
|
||||
logger.error(f"Takeover execution failed: {e}")
|
||||
return False
|
||||
|
||||
def enqueue_delayed_message(self, auth: Authorization, message: dict):
|
||||
"""Write a message into the Redis delayed queue (TTL = delay + 10s buffer)."""
|
||||
if not self.redis:
|
||||
logger.warning("Redis not configured, degrading to immediate takeover")
|
||||
return
|
||||
|
||||
avatar = self.db.query(Avatar).filter(Avatar.id == auth.avatar_id).first()
|
||||
owner_huihui_id = avatar.owner_id if avatar else ""
|
||||
key = f"takeover:delayed:{auth.target_id}:{message.get('msg_id', '')}"
|
||||
value = json.dumps({
|
||||
"avatar_id": auth.avatar_id,
|
||||
"from_accid": message.get("from_accid", ""),
|
||||
"content": message.get("content", ""),
|
||||
"owner_huihui_id": owner_huihui_id,
|
||||
})
|
||||
self.redis.setex(key, auth.takeover_delay_seconds + 10, value)
|
||||
logger.info(f"Message enqueued to delayed queue: {key}")
|
||||
|
||||
async def process_delayed_queue(self):
|
||||
"""Process expired messages from the delayed queue.
|
||||
|
||||
Scans Redis keys matching the takeover:delayed: pattern and dispatches
|
||||
each to execute_takeover after resolving the Authorization.
|
||||
"""
|
||||
if not self.redis:
|
||||
return
|
||||
try:
|
||||
pattern = "takeover:delayed:*"
|
||||
keys = self.redis.keys(pattern)
|
||||
for key in keys:
|
||||
raw = self.redis.get(key)
|
||||
if not raw:
|
||||
continue
|
||||
data = json.loads(raw)
|
||||
auth = (
|
||||
self.db.query(Authorization)
|
||||
.filter(Authorization.target_id == key.split(":")[2])
|
||||
.first()
|
||||
)
|
||||
if auth:
|
||||
message = {
|
||||
"msg_id": key.split(":")[-1],
|
||||
"from_accid": data.get("from_accid", ""),
|
||||
"content": data.get("content", ""),
|
||||
}
|
||||
await self.execute_takeover(auth, message)
|
||||
self.redis.delete(key)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to process delayed queue: {e}")
|
||||
self.reply_delay_seconds = reply_delay_seconds
|
||||
self.now = now
|
||||
self._sessions: dict[str, dict] = {}
|
||||
self._run_lock = asyncio.Lock()
|
||||
|
||||
async def poll_and_process_messages(self):
|
||||
"""Periodic polling job: fetch unread messages and process each."""
|
||||
"""Run one complete cycle; polling always happens before reply dispatch."""
|
||||
if self._run_lock.locked():
|
||||
return
|
||||
async with self._run_lock:
|
||||
self._recover_stuck_tasks()
|
||||
avatar_ids = self._enabled_avatar_ids()
|
||||
self._cancel_disabled_tasks(set(avatar_ids))
|
||||
for avatar_id in avatar_ids:
|
||||
await self._sync_avatar(avatar_id)
|
||||
|
||||
generated = await self._prepare_replies()
|
||||
if generated:
|
||||
# Catch a human reply sent while the model was preparing its answer.
|
||||
for avatar_id in avatar_ids:
|
||||
await self._sync_avatar(avatar_id)
|
||||
await self._dispatch_ready_replies()
|
||||
|
||||
def _enabled_avatar_ids(self) -> list[str]:
|
||||
db = self.session_factory()
|
||||
try:
|
||||
messages = await self.fetch_unread_messages()
|
||||
for msg in messages:
|
||||
await self.process_message(msg)
|
||||
except Exception as e:
|
||||
logger.error(f"poll_and_process_messages failed: {e}")
|
||||
return [
|
||||
avatar.id
|
||||
for avatar in db.query(Avatar).filter(Avatar.status == "active").all()
|
||||
if _takeover_enabled(avatar)
|
||||
]
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
async def fetch_unread_messages(self) -> list:
|
||||
"""Fetch unread messages from Box IM. Stub — replace with real API call."""
|
||||
logger.debug("fetch_unread_messages: no real API wired yet")
|
||||
return []
|
||||
def _cancel_disabled_tasks(self, enabled_avatar_ids: set[str]):
|
||||
db = self.session_factory()
|
||||
try:
|
||||
tasks = (
|
||||
db.query(TakeoverReplyTask)
|
||||
.filter(TakeoverReplyTask.status.in_(ACTIVE_TASK_STATUSES))
|
||||
.all()
|
||||
)
|
||||
changed = False
|
||||
for task in tasks:
|
||||
if task.avatar_id not in enabled_avatar_ids:
|
||||
task.status = "cancelled"
|
||||
task.cancel_reason = "takeover_disabled"
|
||||
task.locked_at = None
|
||||
changed = True
|
||||
if changed:
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
async def process_message(self, message: dict):
|
||||
"""Process a single message: check takeover, dispatch immediate or delayed."""
|
||||
owner_id = message.get("owner_huihui_id", "")
|
||||
from_id = message.get("from_accid", "")
|
||||
def _recover_stuck_tasks(self):
|
||||
db = self.session_factory()
|
||||
try:
|
||||
threshold = self.now() - timedelta(seconds=STUCK_LOCK_SECONDS)
|
||||
tasks = (
|
||||
db.query(TakeoverReplyTask)
|
||||
.filter(
|
||||
TakeoverReplyTask.status.in_(("generating", "sending")),
|
||||
TakeoverReplyTask.locked_at.isnot(None),
|
||||
TakeoverReplyTask.locked_at < threshold,
|
||||
)
|
||||
.all()
|
||||
)
|
||||
for task in tasks:
|
||||
task.status = "pending" if task.status == "generating" else "ready"
|
||||
task.locked_at = None
|
||||
task.last_error = "上次处理意外中断,已自动恢复"
|
||||
if tasks:
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
auth = self.check_takeover_enabled(owner_id, from_id)
|
||||
if not auth:
|
||||
async def _boxim_session(self, user: User) -> dict:
|
||||
token_fingerprint = hashlib.sha256((user.huihui_token or "").encode()).hexdigest()
|
||||
cached = self._sessions.get(user.id)
|
||||
if (
|
||||
cached
|
||||
and cached["expires_at"] > time.monotonic()
|
||||
and cached["token_fingerprint"] == token_fingerprint
|
||||
):
|
||||
return cached
|
||||
|
||||
token_data = await self.boxim.exchange_access_token(user.huihui_token)
|
||||
access_token = token_data["accessToken"]
|
||||
profile = await self.boxim.get_self(access_token)
|
||||
try:
|
||||
expires_in = int(token_data.get("accessTokenExpiresIn") or 3600)
|
||||
except (TypeError, ValueError):
|
||||
expires_in = 3600
|
||||
if expires_in > 86_400:
|
||||
expires_in //= 1000
|
||||
cache_for = max(60, min(expires_in - 60, 3600))
|
||||
cached = {
|
||||
"access_token": access_token,
|
||||
"boxim_owner_id": str(profile["id"]),
|
||||
"expires_at": time.monotonic() + cache_for,
|
||||
"token_fingerprint": token_fingerprint,
|
||||
}
|
||||
self._sessions[user.id] = cached
|
||||
return cached
|
||||
|
||||
def _forget_boxim_session(self, user_id: str):
|
||||
self._sessions.pop(user_id, None)
|
||||
|
||||
def _disable_after_connection_failure(
|
||||
self,
|
||||
db: Session,
|
||||
avatar: Avatar,
|
||||
cursor: TakeoverCursor,
|
||||
message: str,
|
||||
):
|
||||
permissions = (avatar.config or {}).get("authorizationPermissions", [])
|
||||
avatar.config = {
|
||||
**(avatar.config or {}),
|
||||
"authorizationPermissions": [
|
||||
permission
|
||||
for permission in permissions
|
||||
if permission != TAKEOVER_PERMISSION
|
||||
],
|
||||
}
|
||||
cursor.last_error = message
|
||||
cursor.last_polled_at = self.now()
|
||||
tasks = (
|
||||
db.query(TakeoverReplyTask)
|
||||
.filter(
|
||||
TakeoverReplyTask.avatar_id == avatar.id,
|
||||
TakeoverReplyTask.status.in_(ACTIVE_TASK_STATUSES),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
for task in tasks:
|
||||
task.status = "cancelled"
|
||||
task.cancel_reason = "connection_failed"
|
||||
task.locked_at = None
|
||||
|
||||
async def _sync_avatar(self, avatar_id: str) -> bool:
|
||||
db = self.session_factory()
|
||||
try:
|
||||
avatar = db.query(Avatar).filter(Avatar.id == avatar_id).first()
|
||||
if not _takeover_enabled(avatar):
|
||||
return False
|
||||
user = db.query(User).filter(User.huihui_user_id == avatar.owner_id).first()
|
||||
cursor = db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar.id).first()
|
||||
if not cursor:
|
||||
cursor = TakeoverCursor(avatar_id=avatar.id, owner_id=avatar.owner_id)
|
||||
db.add(cursor)
|
||||
db.flush()
|
||||
if not user or not user.huihui_token:
|
||||
self._disable_after_connection_failure(
|
||||
db,
|
||||
avatar,
|
||||
cursor,
|
||||
"请重新登录会会生产账号后再开启主动接管",
|
||||
)
|
||||
db.commit()
|
||||
return False
|
||||
|
||||
try:
|
||||
session = await self._boxim_session(user)
|
||||
owner_boxim_id = session["boxim_owner_id"]
|
||||
if cursor.boxim_owner_id and cursor.boxim_owner_id != owner_boxim_id:
|
||||
cursor.initialized = False
|
||||
cursor.last_message_id = "0"
|
||||
cursor.boxim_owner_id = owner_boxim_id
|
||||
messages = await self.boxim.fetch_private_messages(
|
||||
session["access_token"], cursor.last_message_id or "0"
|
||||
)
|
||||
except Exception as exc:
|
||||
if isinstance(exc, BoxIMError) and exc.auth_error:
|
||||
self._forget_boxim_session(user.id)
|
||||
message = "BOXIM 授权已失效,请重新登录会会生产账号"
|
||||
else:
|
||||
message = f"BOXIM 暂时连接失败:{str(exc)[:160]}"
|
||||
self._disable_after_connection_failure(db, avatar, cursor, message)
|
||||
db.commit()
|
||||
logger.warning("BOXIM sync failed for avatar %s: %s", avatar.id, exc)
|
||||
return False
|
||||
|
||||
messages.sort(key=lambda item: (_numeric_id(item.get("id")), item.get("sendTime") or 0))
|
||||
priming = not bool(cursor.initialized)
|
||||
max_message_id = _numeric_id(cursor.last_message_id)
|
||||
read_receipts: dict[str, int] = {}
|
||||
for message in messages:
|
||||
self._record_message(
|
||||
db,
|
||||
avatar,
|
||||
cursor.boxim_owner_id,
|
||||
message,
|
||||
schedule_reply=not priming,
|
||||
)
|
||||
message_id = _numeric_id(message.get("id"))
|
||||
max_message_id = max(max_message_id, message_id)
|
||||
send_id = str(message.get("sendId") or "")
|
||||
recv_id = str(message.get("recvId") or "")
|
||||
if recv_id == cursor.boxim_owner_id and send_id and message_id:
|
||||
read_receipts[send_id] = max(read_receipts.get(send_id, 0), message_id)
|
||||
|
||||
# BOXIM publishes this HTTP state change to connected socket clients.
|
||||
# Do it before advancing the cursor so a failed receipt is retried.
|
||||
for peer_id, message_id in read_receipts.items():
|
||||
await self.boxim.mark_private_messages_read(
|
||||
session["access_token"], peer_id, message_id
|
||||
)
|
||||
|
||||
cursor.last_message_id = str(max_message_id)
|
||||
cursor.initialized = True
|
||||
cursor.last_polled_at = self.now()
|
||||
cursor.last_error = ""
|
||||
db.commit()
|
||||
return True
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to persist BOXIM messages for avatar %s", avatar_id)
|
||||
return False
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
def _record_message(
|
||||
self,
|
||||
db: Session,
|
||||
avatar: Avatar,
|
||||
boxim_owner_id: str,
|
||||
message: dict,
|
||||
*,
|
||||
schedule_reply: bool,
|
||||
):
|
||||
message_id = str(message.get("id") or "").strip()
|
||||
if not message_id:
|
||||
return
|
||||
local_id = str(message.get("localId") or "").strip() or None
|
||||
if (
|
||||
db.query(TakeoverMessage)
|
||||
.filter(
|
||||
TakeoverMessage.owner_id == avatar.owner_id,
|
||||
TakeoverMessage.boxim_message_id == message_id,
|
||||
)
|
||||
.first()
|
||||
):
|
||||
return
|
||||
|
||||
if auth.takeover_mode == "immediate":
|
||||
await self.execute_takeover(auth, message)
|
||||
send_id = str(message.get("sendId") or "")
|
||||
recv_id = str(message.get("recvId") or "")
|
||||
if send_id == boxim_owner_id:
|
||||
direction, peer_id = "outgoing", recv_id
|
||||
elif recv_id == boxim_owner_id:
|
||||
direction, peer_id = "incoming", send_id
|
||||
else:
|
||||
self.enqueue_delayed_message(auth, message)
|
||||
return
|
||||
if not peer_id:
|
||||
return
|
||||
|
||||
now = self.now()
|
||||
send_time = _boxim_time(message.get("sendTime"), now)
|
||||
is_avatar = False
|
||||
if direction == "outgoing" and local_id:
|
||||
is_avatar = bool(
|
||||
db.query(TakeoverReplyTask)
|
||||
.filter(
|
||||
TakeoverReplyTask.owner_id == avatar.owner_id,
|
||||
TakeoverReplyTask.boxim_local_id == local_id,
|
||||
TakeoverReplyTask.status == "sent",
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
event = TakeoverMessage(
|
||||
avatar_id=avatar.id,
|
||||
owner_id=avatar.owner_id,
|
||||
boxim_message_id=message_id,
|
||||
boxim_local_id=local_id,
|
||||
peer_id=peer_id,
|
||||
direction=direction,
|
||||
message_type=int(message.get("type") or 0),
|
||||
content=str(message.get("content") or ""),
|
||||
is_avatar=is_avatar,
|
||||
send_time=send_time,
|
||||
)
|
||||
db.add(event)
|
||||
db.flush()
|
||||
|
||||
if direction == "outgoing":
|
||||
if not is_avatar:
|
||||
self._cancel_conversation(db, avatar.owner_id, peer_id, "owner_replied")
|
||||
return
|
||||
if not schedule_reply or event.message_type != 0 or not event.content.strip():
|
||||
return
|
||||
if (now - send_time).total_seconds() > MAX_STALE_SECONDS:
|
||||
return
|
||||
self._schedule_reply(db, avatar, event)
|
||||
|
||||
@staticmethod
|
||||
def _cancel_conversation(db: Session, owner_id: str, peer_id: str, reason: str):
|
||||
tasks = (
|
||||
db.query(TakeoverReplyTask)
|
||||
.filter(
|
||||
TakeoverReplyTask.owner_id == owner_id,
|
||||
TakeoverReplyTask.peer_id == peer_id,
|
||||
TakeoverReplyTask.status.in_(ACTIVE_TASK_STATUSES),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
for task in tasks:
|
||||
task.status = "cancelled"
|
||||
task.cancel_reason = reason
|
||||
task.locked_at = None
|
||||
|
||||
def _schedule_reply(self, db: Session, avatar: Avatar, event: TakeoverMessage):
|
||||
active_tasks = (
|
||||
db.query(TakeoverReplyTask)
|
||||
.filter(
|
||||
TakeoverReplyTask.owner_id == avatar.owner_id,
|
||||
TakeoverReplyTask.peer_id == event.peer_id,
|
||||
TakeoverReplyTask.status.in_(("pending", "generating", "ready")),
|
||||
)
|
||||
.order_by(TakeoverReplyTask.created_at.desc())
|
||||
.all()
|
||||
)
|
||||
prompt_parts = []
|
||||
source_ids = []
|
||||
if active_tasks:
|
||||
latest = active_tasks[0]
|
||||
prompt_parts.append(latest.prompt)
|
||||
source_ids.extend(latest.source_message_ids or [])
|
||||
for task in active_tasks:
|
||||
task.status = "cancelled"
|
||||
task.cancel_reason = "newer_incoming_message"
|
||||
task.locked_at = None
|
||||
prompt_parts.append(event.content.strip())
|
||||
source_ids.append(event.boxim_message_id)
|
||||
prompt = "\n".join(part for part in prompt_parts if part).strip()[-MAX_PROMPT_LENGTH:]
|
||||
due_at = event.send_time + timedelta(seconds=self.reply_delay_seconds)
|
||||
task_id = secrets.token_hex(16)
|
||||
local_id = int(time.time() * 1000) * 1000 + secrets.randbelow(1000)
|
||||
db.add(
|
||||
TakeoverReplyTask(
|
||||
id=task_id,
|
||||
avatar_id=avatar.id,
|
||||
owner_id=avatar.owner_id,
|
||||
peer_id=event.peer_id,
|
||||
trigger_message_id=event.boxim_message_id,
|
||||
source_message_ids=source_ids,
|
||||
prompt=prompt,
|
||||
status="pending",
|
||||
scheduled_at=due_at,
|
||||
boxim_local_id=str(local_id),
|
||||
)
|
||||
)
|
||||
|
||||
async def _prepare_replies(self) -> int:
|
||||
db = self.session_factory()
|
||||
try:
|
||||
task_ids = [
|
||||
row[0]
|
||||
for row in (
|
||||
db.query(TakeoverReplyTask.id)
|
||||
.filter(
|
||||
TakeoverReplyTask.status.in_(GENERATABLE_TASK_STATUSES),
|
||||
TakeoverReplyTask.response_text == "",
|
||||
)
|
||||
.order_by(TakeoverReplyTask.created_at.asc())
|
||||
.limit(10)
|
||||
.all()
|
||||
)
|
||||
]
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
generated = 0
|
||||
for task_id in task_ids:
|
||||
if await asyncio.to_thread(self._generate_reply, task_id):
|
||||
generated += 1
|
||||
return generated
|
||||
|
||||
def _generate_reply(self, task_id: str) -> bool:
|
||||
db = self.session_factory()
|
||||
try:
|
||||
task = db.query(TakeoverReplyTask).filter(TakeoverReplyTask.id == task_id).first()
|
||||
if not task or task.status != "pending":
|
||||
return False
|
||||
avatar = db.query(Avatar).filter(Avatar.id == task.avatar_id).first()
|
||||
if not _takeover_enabled(avatar):
|
||||
task.status = "cancelled"
|
||||
task.cancel_reason = "takeover_disabled"
|
||||
db.commit()
|
||||
return False
|
||||
|
||||
task.status = "generating"
|
||||
task.locked_at = self.now()
|
||||
db.commit()
|
||||
|
||||
excluded_ids = set(task.source_message_ids or [])
|
||||
events = (
|
||||
db.query(TakeoverMessage)
|
||||
.filter(
|
||||
TakeoverMessage.owner_id == task.owner_id,
|
||||
TakeoverMessage.peer_id == task.peer_id,
|
||||
)
|
||||
.order_by(TakeoverMessage.send_time.desc())
|
||||
.limit(30)
|
||||
.all()
|
||||
)
|
||||
history = []
|
||||
for event in reversed(events):
|
||||
if event.boxim_message_id in excluded_ids or not event.content.strip():
|
||||
continue
|
||||
history.append(
|
||||
{
|
||||
"role": "user" if event.direction == "incoming" else "assistant",
|
||||
"content": event.content.strip(),
|
||||
}
|
||||
)
|
||||
history = history[-10:]
|
||||
|
||||
from routers.chat import _resolve_reply
|
||||
|
||||
result = _resolve_reply(db, avatar, task.prompt, history)
|
||||
answer = _plain_text_reply(result.get("answer", ""))
|
||||
db.refresh(task)
|
||||
if task.status != "generating":
|
||||
return False
|
||||
if not answer:
|
||||
raise RuntimeError("分身没有生成有效回复")
|
||||
task.response_text = answer
|
||||
task.status = "ready"
|
||||
task.locked_at = None
|
||||
task.last_error = ""
|
||||
db.commit()
|
||||
return True
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
task = db.query(TakeoverReplyTask).filter(TakeoverReplyTask.id == task_id).first()
|
||||
if task and task.status in ("pending", "generating"):
|
||||
task.attempts = (task.attempts or 0) + 1
|
||||
task.status = "pending" if task.attempts < 3 else "failed"
|
||||
task.locked_at = None
|
||||
task.last_error = str(exc)[:300]
|
||||
db.commit()
|
||||
logger.warning("Failed to prepare takeover reply %s: %s", task_id, exc)
|
||||
return False
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
async def _dispatch_ready_replies(self):
|
||||
db = self.session_factory()
|
||||
try:
|
||||
task_ids = [
|
||||
row[0]
|
||||
for row in (
|
||||
db.query(TakeoverReplyTask.id)
|
||||
.filter(
|
||||
TakeoverReplyTask.status == "ready",
|
||||
TakeoverReplyTask.scheduled_at <= self.now(),
|
||||
)
|
||||
.order_by(TakeoverReplyTask.scheduled_at.asc())
|
||||
.limit(10)
|
||||
.all()
|
||||
)
|
||||
]
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
for task_id in task_ids:
|
||||
await self._send_task(task_id)
|
||||
|
||||
async def _send_task(self, task_id: str) -> bool:
|
||||
db = self.session_factory()
|
||||
user = None
|
||||
try:
|
||||
task = db.query(TakeoverReplyTask).filter(TakeoverReplyTask.id == task_id).first()
|
||||
if not task or task.status != "ready":
|
||||
return False
|
||||
avatar = db.query(Avatar).filter(Avatar.id == task.avatar_id).first()
|
||||
if not _takeover_enabled(avatar):
|
||||
task.status = "cancelled"
|
||||
task.cancel_reason = "takeover_disabled"
|
||||
db.commit()
|
||||
return False
|
||||
if (self.now() - task.scheduled_at).total_seconds() > MAX_STALE_SECONDS:
|
||||
task.status = "cancelled"
|
||||
task.cancel_reason = "stale_reply"
|
||||
db.commit()
|
||||
return False
|
||||
user = db.query(User).filter(User.huihui_user_id == task.owner_id).first()
|
||||
if not user or not user.huihui_token:
|
||||
raise BoxIMError("缺少会会登录凭证", auth_error=True)
|
||||
|
||||
task.status = "sending"
|
||||
task.locked_at = self.now()
|
||||
db.commit()
|
||||
session = await self._boxim_session(user)
|
||||
result = await self.boxim.send_private_message(
|
||||
session["access_token"],
|
||||
task.peer_id,
|
||||
task.response_text,
|
||||
local_id=task.boxim_local_id,
|
||||
)
|
||||
db.refresh(task)
|
||||
if task.status != "sending":
|
||||
return False
|
||||
task.status = "sent"
|
||||
task.sent_at = self.now()
|
||||
task.locked_at = None
|
||||
task.last_error = ""
|
||||
task.boxim_sent_message_id = str(result.get("id") or "")
|
||||
db.commit()
|
||||
logger.info("BOXIM takeover reply sent for task %s", task.id)
|
||||
return True
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
if user and isinstance(exc, BoxIMError) and exc.auth_error:
|
||||
self._forget_boxim_session(user.id)
|
||||
task = db.query(TakeoverReplyTask).filter(TakeoverReplyTask.id == task_id).first()
|
||||
if task and task.status in ("ready", "sending"):
|
||||
task.attempts = (task.attempts or 0) + 1
|
||||
task.status = "ready" if task.attempts < 3 else "failed"
|
||||
task.locked_at = None
|
||||
task.last_error = str(exc)[:300]
|
||||
if task.status == "ready":
|
||||
task.scheduled_at = self.now() + timedelta(seconds=2 ** task.attempts)
|
||||
db.commit()
|
||||
logger.warning("Failed to send takeover reply %s: %s", task_id, exc)
|
||||
return False
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@@ -1,6 +1,15 @@
|
||||
import uuid
|
||||
|
||||
import pytest
|
||||
from database import init_db, SessionLocal
|
||||
from models import Authorization
|
||||
from models import (
|
||||
Authorization,
|
||||
Avatar,
|
||||
TakeoverCursor,
|
||||
TakeoverMessage,
|
||||
TakeoverReplyTask,
|
||||
User,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="session", autouse=True)
|
||||
@@ -24,3 +33,82 @@ 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()
|
||||
avatar_ids = [avatar.id, other_avatar.id]
|
||||
db.query(TakeoverReplyTask).filter(
|
||||
TakeoverReplyTask.avatar_id.in_(avatar_ids)
|
||||
).delete(synchronize_session=False)
|
||||
db.query(TakeoverMessage).filter(
|
||||
TakeoverMessage.avatar_id.in_(avatar_ids)
|
||||
).delete(synchronize_session=False)
|
||||
db.query(TakeoverCursor).filter(
|
||||
TakeoverCursor.avatar_id.in_(avatar_ids)
|
||||
).delete(synchronize_session=False)
|
||||
db.query(Authorization).filter(
|
||||
Authorization.avatar_id.in_(avatar_ids)
|
||||
).delete(synchronize_session=False)
|
||||
db.query(Avatar).filter(Avatar.id.in_(avatar_ids)).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,158 @@
|
||||
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
|
||||
|
||||
|
||||
def test_avatar_permission_settings_default_and_persist(authorization_context):
|
||||
context = authorization_context
|
||||
endpoint = f"/api/avatar/{context['avatar'].id}/permission-settings"
|
||||
|
||||
initial = client.get(endpoint, headers=context["owner_headers"]).json()
|
||||
assert initial["code"] == 200
|
||||
assert initial["data"] == {
|
||||
"avatarId": context["avatar"].id,
|
||||
"permissions": ["friend", "chat"],
|
||||
}
|
||||
|
||||
updated = client.put(
|
||||
endpoint,
|
||||
headers=context["owner_headers"],
|
||||
json={"permissions": ["interact", "takeover", "publish", "friend", "friend"]},
|
||||
).json()
|
||||
assert updated["code"] == 200
|
||||
assert updated["data"]["permissions"] == ["friend", "publish", "interact", "takeover"]
|
||||
|
||||
reloaded = client.get(endpoint, headers=context["owner_headers"]).json()
|
||||
assert reloaded["data"]["permissions"] == ["friend", "publish", "interact", "takeover"]
|
||||
|
||||
|
||||
def test_avatar_permission_settings_allow_all_disabled(authorization_context):
|
||||
context = authorization_context
|
||||
endpoint = f"/api/avatar/{context['avatar'].id}/permission-settings"
|
||||
|
||||
response = client.put(
|
||||
endpoint,
|
||||
headers=context["owner_headers"],
|
||||
json={"permissions": []},
|
||||
).json()
|
||||
assert response["code"] == 200
|
||||
assert response["data"]["permissions"] == []
|
||||
|
||||
|
||||
def test_avatar_permission_settings_validate_owner_and_permissions(authorization_context):
|
||||
context = authorization_context
|
||||
endpoint = f"/api/avatar/{context['avatar'].id}/permission-settings"
|
||||
|
||||
invalid = client.put(
|
||||
endpoint,
|
||||
headers=context["owner_headers"],
|
||||
json={"permissions": ["admin"]},
|
||||
).json()
|
||||
assert invalid["code"] == 400
|
||||
|
||||
missing = client.put(
|
||||
endpoint,
|
||||
headers=context["owner_headers"],
|
||||
json={},
|
||||
).json()
|
||||
assert missing["code"] == 400
|
||||
|
||||
forbidden = client.get(
|
||||
f"/api/avatar/{context['other_avatar'].id}/permission-settings",
|
||||
headers=context["owner_headers"],
|
||||
)
|
||||
assert forbidden.status_code == 403
|
||||
|
||||
unauthenticated = client.get(endpoint)
|
||||
assert unauthenticated.status_code == 401
|
||||
@@ -1,134 +1,130 @@
|
||||
"""Tests for the Box IM client (Netease Yunxin gateway wrapper)."""
|
||||
import pytest
|
||||
"""Contract tests for the self-hosted BOXIM client."""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from services.boxim_client import BoxIMClient, BoxIMError
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_config():
|
||||
def config():
|
||||
return {
|
||||
"HUIHUI_IM_BASE_URL": "http://192.168.1.200:60040",
|
||||
"HUIHUI_PLATFORM_BASE_URL": "https://open.example/api",
|
||||
"BOXIM_API_BASE_URL": "https://im.example/api",
|
||||
"HUIHUI_APP_ID": "test_app",
|
||||
"HUIHUI_ACCESS_ID": "test_access",
|
||||
"HUIHUI_ACCESS_SECRET": "test_secret",
|
||||
}
|
||||
|
||||
|
||||
def _make_mock_response(json_data: dict):
|
||||
"""Create a properly configured mock for httpx.Response."""
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = json_data
|
||||
return mock_response
|
||||
def _response(payload: dict, status_code: int = 200):
|
||||
response = MagicMock()
|
||||
response.status_code = status_code
|
||||
response.json.return_value = payload
|
||||
return response
|
||||
|
||||
|
||||
def _patch_httpx_client(json_data: dict):
|
||||
"""Patch httpx.AsyncClient so that `async with httpx.AsyncClient() as c: await c.post(...)` returns json_data."""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = _make_mock_response(json_data)
|
||||
|
||||
mock_cm = AsyncMock()
|
||||
mock_cm.__aenter__.return_value = mock_client
|
||||
mock_cm.__aexit__.return_value = None
|
||||
|
||||
return patch("httpx.AsyncClient", return_value=mock_cm)
|
||||
def _client_patch(*, post_payload=None, request_payload=None, status_code=200):
|
||||
client = AsyncMock()
|
||||
if post_payload is not None:
|
||||
client.post.return_value = _response(post_payload, status_code)
|
||||
if request_payload is not None:
|
||||
client.request.return_value = _response(request_payload, status_code)
|
||||
context = AsyncMock()
|
||||
context.__aenter__.return_value = client
|
||||
context.__aexit__.return_value = None
|
||||
return patch("services.boxim_client.httpx.AsyncClient", return_value=context), client
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_credentials(mock_config):
|
||||
"""get_credentials should return accid and token from the gateway response."""
|
||||
with _patch_httpx_client({"code": 200, "data": {"accid": "user123", "token": "tok_xyz"}}):
|
||||
from services.boxim_client import BoxIMClient
|
||||
async def test_exchange_access_token_uses_huihui_bearer_and_signed_form(config):
|
||||
mocked, client = _client_patch(
|
||||
post_payload={"code": 0, "data": {"accessToken": "box-token", "accessTokenExpiresIn": 3600}}
|
||||
)
|
||||
with mocked:
|
||||
result = await BoxIMClient(config).exchange_access_token("huihui-token")
|
||||
|
||||
client = BoxIMClient(mock_config)
|
||||
result = await client.get_credentials("user123")
|
||||
|
||||
assert result["accid"] == "user123"
|
||||
assert result["token"] == "tok_xyz"
|
||||
assert result["accessToken"] == "box-token"
|
||||
call = client.post.await_args
|
||||
assert call.args[0] == "https://open.example/api/im/box/netease"
|
||||
assert call.kwargs["headers"]["Authorization"] == "Bearer huihui-token"
|
||||
assert call.kwargs["data"]["appId"] == "test_app"
|
||||
assert len(call.kwargs["data"]["signature"]) == 32
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_p2p_message_success(mock_config):
|
||||
"""send_p2p_message should return True when the gateway responds with code 200."""
|
||||
with _patch_httpx_client({"code": 200}):
|
||||
from services.boxim_client import BoxIMClient
|
||||
async def test_get_self_and_incremental_private_messages_use_boxim_header(config):
|
||||
client_instance = BoxIMClient(config)
|
||||
mocked, client = _client_patch(
|
||||
request_payload={"code": 200, "data": {"id": 42, "nickName": "Owner"}}
|
||||
)
|
||||
with mocked:
|
||||
profile = await client_instance.get_self("box-token")
|
||||
assert profile["id"] == 42
|
||||
assert client.request.await_args.kwargs["headers"] == {"accessToken": "box-token"}
|
||||
|
||||
client = BoxIMClient(mock_config)
|
||||
result = await client.send_p2p_message("owner_acc", "target_acc", "Hello")
|
||||
|
||||
assert result is True
|
||||
mocked, client = _client_patch(
|
||||
request_payload={"code": 200, "data": [{"id": 101, "sendId": 7, "recvId": 42}]}
|
||||
)
|
||||
with mocked:
|
||||
messages = await client_instance.fetch_private_messages("box-token", "100")
|
||||
assert messages[0]["id"] == 101
|
||||
assert client.request.await_args.kwargs["params"] == {"minId": "100"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_p2p_message_failure(mock_config):
|
||||
"""send_p2p_message should return False when the gateway responds with a non-200 code."""
|
||||
with _patch_httpx_client({"code": 500, "message": "error"}):
|
||||
from services.boxim_client import BoxIMClient
|
||||
async def test_send_private_message_matches_boxim_payload(config):
|
||||
mocked, client = _client_patch(
|
||||
request_payload={"code": 200, "data": {"id": 88, "localId": 12345}}
|
||||
)
|
||||
with mocked:
|
||||
result = await BoxIMClient(config).send_private_message(
|
||||
"box-token", "77", "你好", local_id="12345"
|
||||
)
|
||||
|
||||
client = BoxIMClient(mock_config)
|
||||
result = await client.send_p2p_message("owner_acc", "target_acc", "Hello")
|
||||
|
||||
assert result is False
|
||||
assert result["id"] == 88
|
||||
call = client.request.await_args
|
||||
assert call.args[:2] == ("POST", "https://im.example/api/message/private/send")
|
||||
assert call.kwargs["json"] == {
|
||||
"localId": 12345,
|
||||
"recvId": 77,
|
||||
"content": "你好",
|
||||
"type": 0,
|
||||
"receipt": False,
|
||||
"atUserIds": [],
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_credentials_returns_none_on_error(mock_config):
|
||||
"""get_credentials should return None when the gateway responds with an error code."""
|
||||
with _patch_httpx_client({"code": 500, "message": "user not found"}):
|
||||
from services.boxim_client import BoxIMClient
|
||||
async def test_mark_private_messages_read_uses_latest_message_id(config):
|
||||
mocked, client = _client_patch(request_payload={"code": 200, "data": None})
|
||||
with mocked:
|
||||
await BoxIMClient(config).mark_private_messages_read("box-token", "77", "101")
|
||||
|
||||
client = BoxIMClient(mock_config)
|
||||
result = await client.get_credentials("nonexistent")
|
||||
|
||||
assert result is None
|
||||
call = client.request.await_args
|
||||
assert call.args[:2] == ("PUT", "https://im.example/api/message/private/readed")
|
||||
assert call.kwargs["headers"] == {"accessToken": "box-token"}
|
||||
assert call.kwargs["params"] == {"friendId": 77, "messageId": 101}
|
||||
|
||||
|
||||
def test_build_sign_params_contains_required_fields(mock_config):
|
||||
"""_build_sign_params should produce appId, accessId, nonce, timestamp, signature, signType, signVersion."""
|
||||
from services.boxim_client import BoxIMClient
|
||||
@pytest.mark.asyncio
|
||||
async def test_boxim_auth_error_is_explicit(config):
|
||||
mocked, _ = _client_patch(
|
||||
request_payload={"code": 400, "message": "未登录"}, status_code=200
|
||||
)
|
||||
with mocked, pytest.raises(BoxIMError) as exc_info:
|
||||
await BoxIMClient(config).get_self("expired")
|
||||
assert exc_info.value.auth_error is True
|
||||
|
||||
client = BoxIMClient(mock_config)
|
||||
params = client._build_sign_params({"userId": "u1"})
|
||||
|
||||
assert "appId" in params
|
||||
assert "accessId" in params
|
||||
assert "nonce" in params
|
||||
assert "timestamp" in params
|
||||
assert "signature" in params
|
||||
def test_sign_params_include_production_required_fields(config):
|
||||
params = BoxIMClient(config)._build_sign_params()
|
||||
assert params["appId"] == "test_app"
|
||||
assert params["accessId"] == "test_access"
|
||||
assert params["signType"] == "MD5"
|
||||
assert params["signVersion"] == "1.0"
|
||||
assert len(params["nonce"]) == 12
|
||||
|
||||
|
||||
def test_build_sign_params_excludes_signature_and_accessSecret_from_signing_string(mock_config):
|
||||
"""signature and accessSecret must be excluded from the signing string to match news_service.py."""
|
||||
from services.boxim_client import BoxIMClient
|
||||
|
||||
client = BoxIMClient(mock_config)
|
||||
|
||||
# Pass params that already contain a stale "signature" value
|
||||
params_with_stale_sig = client._build_sign_params({
|
||||
"userId": "u1",
|
||||
"signature": "OLD_STALE_SIG",
|
||||
})
|
||||
|
||||
# The returned signature must be freshly computed (32-char MD5 uppercase),
|
||||
# NOT the stale value we passed in.
|
||||
assert params_with_stale_sig["signature"] != "OLD_STALE_SIG"
|
||||
assert len(params_with_stale_sig["signature"]) == 32
|
||||
|
||||
# Calling with the same extra params but no stale signature should also work.
|
||||
params_clean = client._build_sign_params({"userId": "u1"})
|
||||
assert len(params_clean["signature"]) == 32
|
||||
|
||||
|
||||
def test_build_sign_params_signature_is_deterministic(mock_config):
|
||||
"""Same inputs should produce valid MD5 signatures."""
|
||||
from services.boxim_client import BoxIMClient
|
||||
|
||||
client = BoxIMClient(mock_config)
|
||||
|
||||
params1 = client._build_sign_params({"userId": "u1"})
|
||||
params2 = client._build_sign_params({"userId": "u1"})
|
||||
|
||||
assert params1["signature"] is not None
|
||||
assert params2["signature"] is not None
|
||||
assert len(params1["signature"]) == 32 # MD5 hex length
|
||||
assert len(params["timestamp"]) == 14
|
||||
assert len(params["signature"]) == 32
|
||||
assert "accessSecret" not in params
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
"""Tests for preserving local avatar ownership when Huihui IDs change."""
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
from sqlalchemy.pool import StaticPool
|
||||
|
||||
from database import Base
|
||||
from models import Avatar, TakeoverCursor, TakeoverMessage, TakeoverReplyTask, User
|
||||
from routers.huihui_auth import _issue_session
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def db():
|
||||
engine = create_engine(
|
||||
"sqlite://",
|
||||
connect_args={"check_same_thread": False},
|
||||
poolclass=StaticPool,
|
||||
)
|
||||
Base.metadata.create_all(engine)
|
||||
session = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)()
|
||||
try:
|
||||
yield session
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def _add_avatar_data(db, owner_id: str, suffix: str = "1") -> Avatar:
|
||||
avatar = Avatar(id=f"avatar-{suffix}", owner_id=owner_id, name="冯医生")
|
||||
db.add_all(
|
||||
[
|
||||
avatar,
|
||||
TakeoverCursor(id=f"cursor-{suffix}", avatar_id=avatar.id, owner_id=owner_id),
|
||||
TakeoverMessage(
|
||||
id=f"message-{suffix}",
|
||||
avatar_id=avatar.id,
|
||||
owner_id=owner_id,
|
||||
boxim_message_id=f"box-{suffix}",
|
||||
peer_id="peer",
|
||||
direction="incoming",
|
||||
send_time=datetime(2026, 8, 20, 12, 0, 0),
|
||||
),
|
||||
TakeoverReplyTask(
|
||||
id=f"task-{suffix}",
|
||||
avatar_id=avatar.id,
|
||||
owner_id=owner_id,
|
||||
peer_id="peer",
|
||||
trigger_message_id=f"trigger-{suffix}",
|
||||
scheduled_at=datetime(2026, 8, 20, 12, 0, 3),
|
||||
boxim_local_id=f"local-{suffix}",
|
||||
),
|
||||
]
|
||||
)
|
||||
db.commit()
|
||||
return avatar
|
||||
|
||||
|
||||
def _assert_avatar_data_owner(db, avatar_id: str, owner_id: str):
|
||||
assert db.query(Avatar).filter_by(id=avatar_id).one().owner_id == owner_id
|
||||
assert db.query(TakeoverCursor).filter_by(avatar_id=avatar_id).one().owner_id == owner_id
|
||||
assert db.query(TakeoverMessage).filter_by(avatar_id=avatar_id).one().owner_id == owner_id
|
||||
assert db.query(TakeoverReplyTask).filter_by(avatar_id=avatar_id).one().owner_id == owner_id
|
||||
|
||||
|
||||
def test_unique_phone_user_is_reused_when_huihui_id_changes(db):
|
||||
legacy = User(
|
||||
id="legacy-local",
|
||||
huihui_user_id="fat-user-id",
|
||||
phone="18500000000",
|
||||
app_token="old-session",
|
||||
)
|
||||
db.add(legacy)
|
||||
db.commit()
|
||||
avatar = _add_avatar_data(db, legacy.huihui_user_id)
|
||||
|
||||
response = _issue_session(
|
||||
db,
|
||||
"18500000000",
|
||||
{"userId": "prod-user-id", "nickname": "用户", "token": "prod-token"},
|
||||
)
|
||||
|
||||
users = db.query(User).all()
|
||||
assert len(users) == 1
|
||||
assert users[0].id == "legacy-local"
|
||||
assert users[0].huihui_user_id == "prod-user-id"
|
||||
assert response["data"]["token"] == users[0].app_token
|
||||
_assert_avatar_data_owner(db, avatar.id, "prod-user-id")
|
||||
|
||||
|
||||
def test_existing_production_user_claims_one_legacy_phone_account(db):
|
||||
current = User(
|
||||
id="prod-local",
|
||||
huihui_user_id="prod-user-id",
|
||||
phone="18500000000",
|
||||
)
|
||||
legacy = User(
|
||||
id="legacy-local",
|
||||
huihui_user_id="fat-user-id",
|
||||
phone="18500000000",
|
||||
app_token="old-session",
|
||||
huihui_token="fat-token",
|
||||
)
|
||||
db.add_all([current, legacy])
|
||||
db.commit()
|
||||
avatar = _add_avatar_data(db, legacy.huihui_user_id)
|
||||
|
||||
_issue_session(
|
||||
db,
|
||||
"18500000000",
|
||||
{"userId": "prod-user-id", "nickname": "用户", "token": "prod-token"},
|
||||
)
|
||||
|
||||
db.refresh(legacy)
|
||||
assert legacy.app_token == ""
|
||||
assert legacy.huihui_token == ""
|
||||
_assert_avatar_data_owner(db, avatar.id, "prod-user-id")
|
||||
|
||||
|
||||
def test_ambiguous_phone_matches_do_not_move_existing_avatars(db):
|
||||
first = User(id="first", huihui_user_id="fat-1", phone="18500000000")
|
||||
second = User(id="second", huihui_user_id="fat-2", phone="18500000000")
|
||||
db.add_all([first, second])
|
||||
db.commit()
|
||||
first_avatar = _add_avatar_data(db, first.huihui_user_id, "1")
|
||||
second_avatar = _add_avatar_data(db, second.huihui_user_id, "2")
|
||||
|
||||
_issue_session(
|
||||
db,
|
||||
"18500000000",
|
||||
{"userId": "prod-user-id", "nickname": "用户", "token": "prod-token"},
|
||||
)
|
||||
|
||||
assert db.query(User).count() == 3
|
||||
_assert_avatar_data_owner(db, first_avatar.id, "fat-1")
|
||||
_assert_avatar_data_owner(db, second_avatar.id, "fat-2")
|
||||
@@ -0,0 +1,23 @@
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from routers.knowledge import _doc_payload
|
||||
|
||||
|
||||
def test_doc_payload_reports_whether_the_persisted_file_exists(tmp_path: Path):
|
||||
avatar_id = "avatar-1"
|
||||
stored_name = "knowledge.md"
|
||||
doc = SimpleNamespace(
|
||||
avatar_id=avatar_id,
|
||||
file_url=f"/api/files/{avatar_id}/{stored_name}",
|
||||
to_dict=lambda: {"id": "doc-1", "fileUrl": f"/api/files/{avatar_id}/{stored_name}"},
|
||||
)
|
||||
stored_dir = tmp_path / avatar_id
|
||||
stored_dir.mkdir()
|
||||
stored_file = stored_dir / stored_name
|
||||
|
||||
with patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)):
|
||||
assert _doc_payload(doc)["filePresent"] is False
|
||||
stored_file.write_text("knowledge", encoding="utf-8")
|
||||
assert _doc_payload(doc)["filePresent"] is True
|
||||
@@ -1,105 +1,253 @@
|
||||
"""Tests for PUT /api/avatar/{avatar_id}/authorizations/takeover endpoint."""
|
||||
"""Tests for takeover configuration and BOXIM connection status."""
|
||||
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
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, Avatar, TakeoverCursor, TakeoverReplyTask, User
|
||||
|
||||
|
||||
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",
|
||||
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},
|
||||
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={
|
||||
"authorizationId": context["authorization"].id,
|
||||
"takeoverEnabled": True,
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["code"] == 400
|
||||
|
||||
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_not_found():
|
||||
client = TestClient(app)
|
||||
response = client.put(
|
||||
f"/api/avatar/test/authorizations/takeover",
|
||||
json={"authorization_id": "nonexistent"},
|
||||
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": 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"] == 404
|
||||
assert forbidden.status_code == 403
|
||||
|
||||
|
||||
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"]
|
||||
|
||||
|
||||
def test_takeover_status_reports_disabled_and_requires_owner_login(authorization_context):
|
||||
context = authorization_context
|
||||
endpoint = f"/api/avatar/{context['avatar'].id}/takeover/status"
|
||||
|
||||
disabled = client.get(endpoint, headers=context["owner_headers"])
|
||||
assert disabled.status_code == 200
|
||||
assert disabled.json()["data"]["status"] == "disabled"
|
||||
|
||||
client.put(
|
||||
f"/api/avatar/{context['avatar'].id}/permission-settings",
|
||||
headers=context["owner_headers"],
|
||||
json={"permissions": ["chat", "takeover"]},
|
||||
)
|
||||
needs_login = client.get(endpoint, headers=context["owner_headers"]).json()["data"]
|
||||
assert needs_login["enabled"] is True
|
||||
assert needs_login["status"] == "needs_login"
|
||||
assert "BOXIM" in needs_login["message"]
|
||||
|
||||
assert client.get(endpoint).status_code == 401
|
||||
assert client.get(endpoint, headers=context["other_headers"]).status_code == 403
|
||||
|
||||
|
||||
def test_takeover_status_reports_ready_pending_count_and_errors(authorization_context):
|
||||
context = authorization_context
|
||||
avatar_id = context["avatar"].id
|
||||
endpoint = f"/api/avatar/{avatar_id}/takeover/status"
|
||||
client.put(
|
||||
f"/api/avatar/{avatar_id}/permission-settings",
|
||||
headers=context["owner_headers"],
|
||||
json={"permissions": ["chat", "takeover"]},
|
||||
)
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
owner = db.query(User).filter(User.id == context["owner"].id).one()
|
||||
owner.huihui_token = "production-login-token"
|
||||
cursor = TakeoverCursor(
|
||||
avatar_id=avatar_id,
|
||||
owner_id=owner.huihui_user_id,
|
||||
boxim_owner_id="100",
|
||||
last_message_id="10",
|
||||
initialized=True,
|
||||
last_polled_at=datetime.utcnow(),
|
||||
)
|
||||
task = TakeoverReplyTask(
|
||||
avatar_id=avatar_id,
|
||||
owner_id=owner.huihui_user_id,
|
||||
peer_id="200",
|
||||
trigger_message_id="11",
|
||||
source_message_ids=["11"],
|
||||
prompt="你好",
|
||||
status="pending",
|
||||
scheduled_at=datetime.utcnow(),
|
||||
boxim_local_id="123",
|
||||
)
|
||||
db.add_all([cursor, task])
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
ready = client.get(endpoint, headers=context["owner_headers"]).json()["data"]
|
||||
assert ready["status"] == "ready"
|
||||
assert ready["pendingCount"] == 1
|
||||
assert ready["lastPolledAt"]
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
cursor = db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar_id).one()
|
||||
cursor.last_polled_at = datetime.utcnow() - timedelta(seconds=30)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
long_polling = client.get(endpoint, headers=context["owner_headers"]).json()["data"]
|
||||
assert long_polling["status"] == "ready"
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
cursor = db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar_id).one()
|
||||
cursor.last_polled_at = datetime.utcnow() - timedelta(seconds=61)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
stale = client.get(endpoint, headers=context["owner_headers"]).json()["data"]
|
||||
assert stale["status"] == "connecting"
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
cursor = db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar_id).one()
|
||||
cursor.last_error = "BOXIM 暂时不可用"
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
failed = client.get(endpoint, headers=context["owner_headers"]).json()["data"]
|
||||
assert failed["status"] == "error"
|
||||
assert failed["message"] == "BOXIM 暂时不可用"
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
avatar = db.query(Avatar).filter(Avatar.id == avatar_id).one()
|
||||
avatar.config = {"authorizationPermissions": ["chat"]}
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
auto_disabled = client.get(endpoint, headers=context["owner_headers"]).json()["data"]
|
||||
assert auto_disabled["enabled"] is False
|
||||
assert auto_disabled["status"] == "error"
|
||||
|
||||
client.put(
|
||||
f"/api/avatar/{avatar_id}/permission-settings",
|
||||
headers=context["owner_headers"],
|
||||
json={"permissions": ["chat", "takeover"]},
|
||||
)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
cursor = db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar_id).one()
|
||||
assert cursor.initialized is False
|
||||
assert cursor.last_message_id == "0"
|
||||
assert cursor.last_error == ""
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@@ -1,221 +1,83 @@
|
||||
"""Tests for the scheduled takeover message polling."""
|
||||
import json
|
||||
import pytest
|
||||
from unittest.mock import MagicMock, patch, AsyncMock
|
||||
"""Tests for the BOXIM takeover scheduler lifecycle."""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
||||
def test_app_has_startup_event():
|
||||
"""Verify the app has a startup event configured."""
|
||||
def test_app_has_startup_and_shutdown_events():
|
||||
from main import app
|
||||
startup_handlers = [handler for handler in app.router.on_startup]
|
||||
assert len(startup_handlers) > 0
|
||||
|
||||
|
||||
@patch("services.takeover_service.TakeoverService")
|
||||
@patch("services.boxim_client.BoxIMClient")
|
||||
@patch("main.redis_lib.from_url")
|
||||
@patch("main.AsyncIOScheduler")
|
||||
def test_scheduler_initialized_with_redis(mock_scheduler_class, mock_redis_from_url, mock_boxim_cls, mock_takeover_cls):
|
||||
"""Verify scheduler is initialized when Redis is available."""
|
||||
mock_redis = MagicMock()
|
||||
mock_redis.ping.return_value = None
|
||||
mock_redis_from_url.return_value = mock_redis
|
||||
|
||||
mock_boxim = MagicMock()
|
||||
mock_boxim_cls.return_value = mock_boxim
|
||||
|
||||
mock_takeover = MagicMock()
|
||||
mock_takeover.poll_and_process_messages = AsyncMock()
|
||||
mock_takeover_cls.return_value = mock_takeover
|
||||
|
||||
with patch("main.init_db"), patch("main.seed"), patch.dict("os.environ", {"REDIS_URL": "redis://localhost:6379"}):
|
||||
from main import on_startup
|
||||
on_startup()
|
||||
|
||||
mock_scheduler_class.return_value.add_job.assert_called_once()
|
||||
scheduled_callable = mock_scheduler_class.return_value.add_job.call_args.args[0]
|
||||
call_kwargs = mock_scheduler_class.return_value.add_job.call_args[1]
|
||||
assert scheduled_callable is mock_takeover.poll_and_process_messages
|
||||
assert call_kwargs["id"] == "takeover_message_poll"
|
||||
mock_scheduler_class.return_value.start.assert_called_once_with()
|
||||
assert app.router.on_startup
|
||||
assert app.router.on_shutdown
|
||||
|
||||
|
||||
@patch("services.takeover_service.TakeoverService")
|
||||
@patch("services.boxim_client.BoxIMClient")
|
||||
@patch("main.AsyncIOScheduler")
|
||||
def test_scheduler_starts_without_redis(mock_scheduler_class, mock_boxim_cls, mock_takeover_cls):
|
||||
"""App should start even when REDIS_URL is not set."""
|
||||
mock_boxim = MagicMock()
|
||||
mock_boxim_cls.return_value = mock_boxim
|
||||
|
||||
mock_takeover = MagicMock()
|
||||
mock_takeover.poll_and_process_messages = AsyncMock()
|
||||
mock_takeover_cls.return_value = mock_takeover
|
||||
|
||||
with patch("main.init_db"), patch("main.seed"), patch.dict("os.environ", {"REDIS_URL": ""}, clear=False):
|
||||
from main import on_startup
|
||||
on_startup()
|
||||
|
||||
mock_scheduler_class.return_value.add_job.assert_called_once()
|
||||
|
||||
|
||||
@patch("services.takeover_service.TakeoverService")
|
||||
@patch("services.boxim_client.BoxIMClient")
|
||||
@patch("main.redis_lib.from_url")
|
||||
@patch("main.AsyncIOScheduler")
|
||||
def test_scheduler_starts_when_redis_fails(mock_scheduler_class, mock_redis_from_url, mock_boxim_cls, mock_takeover_cls):
|
||||
"""App should start even when Redis ping fails."""
|
||||
mock_redis_from_url.side_effect = ConnectionError("Connection refused")
|
||||
|
||||
mock_boxim = MagicMock()
|
||||
mock_boxim_cls.return_value = mock_boxim
|
||||
|
||||
mock_takeover = MagicMock()
|
||||
mock_takeover.poll_and_process_messages = AsyncMock()
|
||||
mock_takeover_cls.return_value = mock_takeover
|
||||
|
||||
with patch("main.init_db"), patch("main.seed"), patch.dict("os.environ", {"REDIS_URL": "redis://badhost:6379"}):
|
||||
from main import on_startup
|
||||
on_startup()
|
||||
|
||||
mock_scheduler_class.return_value.add_job.assert_called_once()
|
||||
|
||||
|
||||
@patch("main.AsyncIOScheduler")
|
||||
def test_scheduler_fails_gracefully(mock_scheduler_class):
|
||||
"""If scheduler init raises, the app should still start (exception caught)."""
|
||||
mock_scheduler_class.side_effect = RuntimeError("Scheduler crash")
|
||||
|
||||
with patch("main.init_db"), patch("main.seed"):
|
||||
from main import on_startup
|
||||
on_startup()
|
||||
|
||||
# No exception should propagate
|
||||
|
||||
|
||||
def test_scheduler_shutdown_releases_resources():
|
||||
"""Shutdown should stop polling and close its dedicated database session."""
|
||||
def test_scheduler_uses_boxim_and_restart_safe_service(
|
||||
mock_scheduler_class,
|
||||
mock_boxim_class,
|
||||
mock_takeover_class,
|
||||
):
|
||||
import main
|
||||
|
||||
mock_scheduler = MagicMock()
|
||||
mock_scheduler.running = True
|
||||
mock_db = MagicMock()
|
||||
main.takeover_scheduler = mock_scheduler
|
||||
main.takeover_db = mock_db
|
||||
scheduler = MagicMock()
|
||||
mock_scheduler_class.return_value = scheduler
|
||||
boxim = MagicMock()
|
||||
mock_boxim_class.return_value = boxim
|
||||
takeover = MagicMock()
|
||||
takeover.poll_and_process_messages = AsyncMock()
|
||||
mock_takeover_class.return_value = takeover
|
||||
|
||||
environment = {
|
||||
"HUIHUI_PLATFORM_BASE_URL": "https://open.example/api",
|
||||
"BOXIM_API_BASE_URL": "https://im.example/api",
|
||||
"HUIHUI_APP_ID": "app-id",
|
||||
"HUIHUI_ACCESS_ID": "access-id",
|
||||
"HUIHUI_ACCESS_SECRET": "secret",
|
||||
"BOXIM_POLL_INTERVAL_SECONDS": "1",
|
||||
}
|
||||
with patch("main.init_db"), patch("main.seed"), patch.dict(
|
||||
"os.environ", environment, clear=False
|
||||
):
|
||||
main.on_startup()
|
||||
|
||||
config = mock_boxim_class.call_args.args[0]
|
||||
assert config["HUIHUI_PLATFORM_BASE_URL"] == "https://open.example/api"
|
||||
assert config["BOXIM_API_BASE_URL"] == "https://im.example/api"
|
||||
mock_takeover_class.assert_called_once_with(main.SessionLocal, boxim)
|
||||
|
||||
scheduler.add_job.assert_called_once()
|
||||
scheduled_callable = scheduler.add_job.call_args.args[0]
|
||||
job_options = scheduler.add_job.call_args.kwargs
|
||||
assert scheduled_callable is takeover.poll_and_process_messages
|
||||
assert job_options["id"] == "takeover_message_poll"
|
||||
assert job_options["trigger"].interval.total_seconds() == 1
|
||||
assert job_options["max_instances"] == 1
|
||||
assert job_options["coalesce"] is True
|
||||
scheduler.start.assert_called_once_with()
|
||||
|
||||
main.takeover_scheduler = None
|
||||
|
||||
|
||||
@patch("main.AsyncIOScheduler")
|
||||
def test_scheduler_failure_does_not_stop_the_api(mock_scheduler_class):
|
||||
import main
|
||||
|
||||
mock_scheduler_class.side_effect = RuntimeError("scheduler crash")
|
||||
with patch("main.init_db"), patch("main.seed"):
|
||||
main.on_startup()
|
||||
|
||||
assert main.takeover_scheduler is None
|
||||
|
||||
|
||||
def test_shutdown_stops_only_the_scheduler():
|
||||
import main
|
||||
|
||||
scheduler = MagicMock()
|
||||
scheduler.running = True
|
||||
main.takeover_scheduler = scheduler
|
||||
|
||||
main.on_shutdown()
|
||||
|
||||
mock_scheduler.shutdown.assert_called_once_with(wait=False)
|
||||
mock_db.close.assert_called_once_with()
|
||||
scheduler.shutdown.assert_called_once_with(wait=False)
|
||||
assert main.takeover_scheduler is None
|
||||
assert main.takeover_db is None
|
||||
|
||||
|
||||
# --- poll_and_process_messages ---
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_db():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_boxim():
|
||||
return AsyncMock()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_poll_and_process_messages_calls_fetch_and_process(mock_db, mock_boxim):
|
||||
"""poll_and_process_messages should fetch messages and process each."""
|
||||
from services.takeover_service import TakeoverService
|
||||
|
||||
service = TakeoverService(mock_db, mock_boxim)
|
||||
service.fetch_unread_messages = AsyncMock(return_value=[
|
||||
{"owner_huihui_id": "owner_1", "from_accid": "user_1", "content": "hi"},
|
||||
{"owner_huihui_id": "owner_2", "from_accid": "user_2", "content": "hello"},
|
||||
])
|
||||
service.process_message = AsyncMock()
|
||||
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
service.fetch_unread_messages.assert_awaited_once()
|
||||
assert service.process_message.await_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_poll_and_process_messages_handles_errors(mock_db, mock_boxim):
|
||||
"""poll_and_process_messages should not crash on fetch failure."""
|
||||
from services.takeover_service import TakeoverService
|
||||
|
||||
service = TakeoverService(mock_db, mock_boxim)
|
||||
service.fetch_unread_messages = AsyncMock(side_effect=ConnectionError("Box IM down"))
|
||||
|
||||
await service.poll_and_process_messages()
|
||||
# No exception should propagate
|
||||
|
||||
|
||||
# --- process_message ---
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_auth():
|
||||
auth = MagicMock()
|
||||
auth.takeover_enabled = True
|
||||
auth.takeover_mode = "immediate"
|
||||
auth.takeover_delay_seconds = 30
|
||||
auth.avatar_id = "avatar_123"
|
||||
auth.target_id = "target_user_123"
|
||||
return auth
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_message_immediate_mode(mock_db, mock_boxim, mock_auth):
|
||||
"""When takeover_mode is 'immediate', execute_takeover should be called."""
|
||||
from services.takeover_service import TakeoverService
|
||||
|
||||
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_delayed_mode(mock_db, mock_boxim, mock_auth):
|
||||
"""When takeover_mode is not 'immediate', message should be enqueued."""
|
||||
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()
|
||||
service.enqueue_delayed_message = MagicMock()
|
||||
|
||||
message = {"owner_huihui_id": "owner_1", "from_accid": "user_1", "content": "hi"}
|
||||
await service.process_message(message)
|
||||
|
||||
service.enqueue_delayed_message.assert_called_once_with(mock_auth, message)
|
||||
service.execute_takeover.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_message_no_takeover(mock_db, mock_boxim):
|
||||
"""When takeover is not enabled, nothing should happen."""
|
||||
from services.takeover_service import TakeoverService
|
||||
|
||||
service = TakeoverService(mock_db, mock_boxim)
|
||||
service.check_takeover_enabled = MagicMock(return_value=None)
|
||||
service.execute_takeover = AsyncMock()
|
||||
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_not_awaited()
|
||||
service.enqueue_delayed_message.assert_not_called()
|
||||
|
||||
@@ -1,318 +1,263 @@
|
||||
"""Tests for the TakeoverService — message listening, decision, reply execution."""
|
||||
"""End-to-end service tests for BOXIM takeover timing and human priority."""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, patch, MagicMock
|
||||
from services.takeover_service import TakeoverService
|
||||
from models import Authorization, Avatar
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
from sqlalchemy.pool import StaticPool
|
||||
|
||||
from database import Base
|
||||
from models import Avatar, TakeoverCursor, TakeoverMessage, TakeoverReplyTask, User
|
||||
from services.boxim_client import BoxIMError
|
||||
from services.takeover_service import TakeoverService, _plain_text_reply
|
||||
|
||||
|
||||
class Clock:
|
||||
def __init__(self):
|
||||
self.value = datetime(2026, 8, 19, 10, 0, 0)
|
||||
|
||||
def now(self):
|
||||
return self.value
|
||||
|
||||
def advance(self, seconds: int):
|
||||
self.value += timedelta(seconds=seconds)
|
||||
|
||||
def millis(self):
|
||||
return int(self.value.replace(tzinfo=timezone.utc).timestamp() * 1000)
|
||||
|
||||
|
||||
class FakeBoxIM:
|
||||
def __init__(self):
|
||||
self.messages = []
|
||||
self.sent = []
|
||||
self.read_receipts = []
|
||||
|
||||
async def exchange_access_token(self, huihui_token):
|
||||
assert huihui_token == "prod-huihui-token"
|
||||
return {"accessToken": "box-token", "accessTokenExpiresIn": 3600}
|
||||
|
||||
async def get_self(self, access_token):
|
||||
assert access_token == "box-token"
|
||||
return {"id": 100}
|
||||
|
||||
async def fetch_private_messages(self, access_token, min_id="0"):
|
||||
assert access_token == "box-token"
|
||||
return [item.copy() for item in self.messages if int(item["id"]) > int(min_id)]
|
||||
|
||||
async def mark_private_messages_read(self, access_token, friend_id, message_id):
|
||||
assert access_token == "box-token"
|
||||
self.read_receipts.append(
|
||||
{"friendId": str(friend_id), "messageId": str(message_id)}
|
||||
)
|
||||
|
||||
async def send_private_message(self, access_token, peer_id, content, *, local_id=None):
|
||||
self.sent.append({"peerId": str(peer_id), "content": content, "localId": str(local_id)})
|
||||
return {"id": 900 + len(self.sent), "localId": int(local_id)}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_db():
|
||||
db = MagicMock()
|
||||
return db
|
||||
def service_context():
|
||||
engine = create_engine(
|
||||
"sqlite://",
|
||||
connect_args={"check_same_thread": False},
|
||||
poolclass=StaticPool,
|
||||
)
|
||||
session_factory = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
|
||||
Base.metadata.create_all(engine)
|
||||
db = session_factory()
|
||||
user = User(
|
||||
id="owner-local",
|
||||
huihui_user_id="owner-huihui",
|
||||
huihui_token="prod-huihui-token",
|
||||
app_token="app-token",
|
||||
)
|
||||
avatar = Avatar(
|
||||
id="avatar-1",
|
||||
owner_id=user.huihui_user_id,
|
||||
name="分身",
|
||||
status="active",
|
||||
config={"authorizationPermissions": ["chat", "takeover"]},
|
||||
)
|
||||
db.add_all([user, avatar])
|
||||
db.commit()
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_boxim():
|
||||
client = AsyncMock()
|
||||
client.get_credentials.return_value = {"accid": "owner_acc", "token": "tok"}
|
||||
client.send_p2p_message.return_value = True
|
||||
return client
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_auth():
|
||||
auth = MagicMock(spec=Authorization)
|
||||
auth.takeover_enabled = True
|
||||
auth.takeover_mode = "immediate"
|
||||
auth.takeover_delay_seconds = 30
|
||||
auth.avatar_id = "avatar_123"
|
||||
auth.target_id = "target_user_123"
|
||||
return auth
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_avatar():
|
||||
avatar = MagicMock(spec=Avatar)
|
||||
avatar.id = "avatar_123"
|
||||
avatar.owner_id = "owner_huihui_123"
|
||||
return avatar
|
||||
|
||||
|
||||
# --- check_takeover_enabled ---
|
||||
|
||||
|
||||
def test_check_takeover_enabled_returns_auth_when_enabled(mock_db, mock_auth, mock_boxim, mock_avatar):
|
||||
avatar_filter = MagicMock()
|
||||
avatar_filter.first.return_value = mock_avatar
|
||||
avatar_query = MagicMock()
|
||||
avatar_query.filter.return_value = avatar_filter
|
||||
|
||||
auth_filter = MagicMock()
|
||||
auth_filter.filter.return_value = auth_filter
|
||||
auth_filter.first.return_value = mock_auth
|
||||
|
||||
def query_side_effect(model):
|
||||
if model == Avatar:
|
||||
return avatar_query
|
||||
return auth_filter
|
||||
|
||||
mock_db.query.side_effect = query_side_effect
|
||||
|
||||
service = TakeoverService(mock_db, mock_boxim)
|
||||
result = service.check_takeover_enabled("owner_huihui_123", "target_user_123")
|
||||
assert result == mock_auth
|
||||
|
||||
|
||||
def test_check_takeover_enabled_returns_none_when_no_avatar(mock_db, mock_boxim):
|
||||
avatar_filter = MagicMock()
|
||||
avatar_filter.first.return_value = None
|
||||
avatar_query = MagicMock()
|
||||
avatar_query.filter.return_value = avatar_filter
|
||||
mock_db.query.return_value = avatar_query
|
||||
|
||||
service = TakeoverService(mock_db, mock_boxim)
|
||||
result = service.check_takeover_enabled("owner_123", "target_123")
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_check_takeover_enabled_returns_none_when_disabled(mock_db, mock_boxim, mock_avatar):
|
||||
avatar_filter = MagicMock()
|
||||
avatar_filter.first.return_value = mock_avatar
|
||||
avatar_query = MagicMock()
|
||||
avatar_query.filter.return_value = avatar_filter
|
||||
|
||||
disabled_auth = MagicMock(spec=Authorization)
|
||||
disabled_auth.takeover_enabled = False
|
||||
auth_filter = MagicMock()
|
||||
auth_filter.filter.return_value = auth_filter
|
||||
auth_filter.first.return_value = disabled_auth
|
||||
|
||||
def query_side_effect(model):
|
||||
if model == Avatar:
|
||||
return avatar_query
|
||||
return auth_filter
|
||||
|
||||
mock_db.query.side_effect = query_side_effect
|
||||
|
||||
service = TakeoverService(mock_db, mock_boxim)
|
||||
result = service.check_takeover_enabled("owner_123", "target_123")
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_check_takeover_enabled_filters_by_owner_and_target(mock_db, mock_boxim, mock_avatar, mock_auth):
|
||||
"""Verify that queries use the correct filter arguments."""
|
||||
avatar_filter = MagicMock()
|
||||
avatar_filter.first.return_value = mock_avatar
|
||||
avatar_query = MagicMock()
|
||||
avatar_query.filter.return_value = avatar_filter
|
||||
|
||||
auth_filter = MagicMock()
|
||||
auth_filter.filter.return_value = auth_filter
|
||||
auth_filter.first.return_value = mock_auth
|
||||
|
||||
call_order = []
|
||||
|
||||
def query_side_effect(model):
|
||||
if model == Avatar:
|
||||
call_order.append("Avatar")
|
||||
return avatar_query
|
||||
call_order.append("Authorization")
|
||||
return auth_filter
|
||||
|
||||
mock_db.query.side_effect = query_side_effect
|
||||
|
||||
service = TakeoverService(mock_db, mock_boxim)
|
||||
service.check_takeover_enabled("owner_huihui_123", "target_user_123")
|
||||
|
||||
assert "Avatar" in call_order
|
||||
assert "Authorization" in call_order
|
||||
|
||||
|
||||
# --- generate_reply ---
|
||||
clock = Clock()
|
||||
boxim = FakeBoxIM()
|
||||
service = TakeoverService(session_factory, boxim, now=clock.now)
|
||||
return session_factory, service, boxim, clock
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_reply_returns_answer(mock_boxim):
|
||||
mock_db = MagicMock()
|
||||
with patch("services.takeover_service.httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {"code": 200, "data": {"answer": "Hello back"}}
|
||||
mock_client_class.return_value.__aenter__.return_value.post.return_value = mock_response
|
||||
async def test_first_sync_primes_cursor_without_replying_to_history(service_context):
|
||||
session_factory, service, boxim, clock = service_context
|
||||
boxim.messages = [
|
||||
{"id": 10, "localId": 1, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "旧消息"}
|
||||
]
|
||||
|
||||
service = TakeoverService(mock_db, mock_boxim)
|
||||
result = await service.generate_reply("avatar_123", "Hello")
|
||||
assert result == "Hello back"
|
||||
with patch("routers.chat._resolve_reply", return_value={"answer": "不应发送"}):
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
db = session_factory()
|
||||
try:
|
||||
cursor = db.query(TakeoverCursor).one()
|
||||
assert cursor.initialized is True
|
||||
assert cursor.last_message_id == "10"
|
||||
assert db.query(TakeoverMessage).count() == 1
|
||||
assert db.query(TakeoverReplyTask).count() == 0
|
||||
assert boxim.sent == []
|
||||
assert boxim.read_receipts == [{"friendId": "200", "messageId": "10"}]
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_reply_handles_empty_answer(mock_boxim):
|
||||
"""generate_reply should return empty string when answer is missing."""
|
||||
mock_db = MagicMock()
|
||||
with patch("services.takeover_service.httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {"code": 200, "data": {}}
|
||||
mock_client_class.return_value.__aenter__.return_value.post.return_value = mock_response
|
||||
async def test_incoming_message_is_prepared_then_sent_at_three_seconds(service_context):
|
||||
session_factory, service, boxim, clock = service_context
|
||||
await service.poll_and_process_messages()
|
||||
boxim.messages.append(
|
||||
{"id": 11, "localId": 2, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "你好"}
|
||||
)
|
||||
|
||||
service = TakeoverService(mock_db, mock_boxim)
|
||||
result = await service.generate_reply("avatar_123", "Hello")
|
||||
assert result == ""
|
||||
with patch("routers.chat._resolve_reply", return_value={"answer": "**你好**\n\n很高兴见到你"}):
|
||||
await service.poll_and_process_messages()
|
||||
assert boxim.sent == []
|
||||
assert boxim.read_receipts == [{"friendId": "200", "messageId": "11"}]
|
||||
|
||||
clock.advance(2)
|
||||
await service.poll_and_process_messages()
|
||||
assert boxim.sent == []
|
||||
|
||||
clock.advance(1)
|
||||
await service.poll_and_process_messages()
|
||||
assert boxim.sent == [{"peerId": "200", "content": "你好\n很高兴见到你", "localId": boxim.sent[0]["localId"]}]
|
||||
|
||||
db = session_factory()
|
||||
try:
|
||||
task = db.query(TakeoverReplyTask).one()
|
||||
assert task.status == "sent"
|
||||
assert task.sent_at == clock.now()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_reply_handles_error_code(mock_boxim):
|
||||
"""generate_reply should return empty string when API returns error code."""
|
||||
mock_db = MagicMock()
|
||||
with patch("services.takeover_service.httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {"code": 500, "message": "Internal error"}
|
||||
mock_client_class.return_value.__aenter__.return_value.post.return_value = mock_response
|
||||
async def test_read_receipt_failure_does_not_advance_cursor(service_context):
|
||||
session_factory, service, boxim, clock = service_context
|
||||
await service.poll_and_process_messages()
|
||||
boxim.messages.append(
|
||||
{"id": 12, "localId": 3, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "未读消息"}
|
||||
)
|
||||
boxim.mark_private_messages_read = AsyncMock(side_effect=BoxIMError("回执失败"))
|
||||
|
||||
service = TakeoverService(mock_db, mock_boxim)
|
||||
result = await service.generate_reply("avatar_123", "Hello")
|
||||
assert result == ""
|
||||
with patch("routers.chat._resolve_reply", return_value={"answer": "稍后回复"}):
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
db = session_factory()
|
||||
try:
|
||||
cursor = db.query(TakeoverCursor).one()
|
||||
assert cursor.last_message_id == "0"
|
||||
assert db.query(TakeoverMessage).count() == 0
|
||||
assert db.query(TakeoverReplyTask).count() == 0
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- execute_takeover ---
|
||||
boxim.mark_private_messages_read = AsyncMock(return_value=None)
|
||||
with patch("routers.chat._resolve_reply", return_value={"answer": "稍后回复"}):
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
db = session_factory()
|
||||
try:
|
||||
assert db.query(TakeoverCursor).one().last_message_id == "12"
|
||||
assert db.query(TakeoverMessage).count() == 1
|
||||
assert db.query(TakeoverReplyTask).count() == 1
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_takeover_success(mock_db, mock_boxim, mock_auth, mock_avatar):
|
||||
avatar_filter = MagicMock()
|
||||
avatar_filter.first.return_value = mock_avatar
|
||||
avatar_query = MagicMock()
|
||||
avatar_query.filter.return_value = avatar_filter
|
||||
mock_db.query.return_value = avatar_query
|
||||
async def test_owner_message_cancels_pending_reply(service_context):
|
||||
session_factory, service, boxim, clock = service_context
|
||||
await service.poll_and_process_messages()
|
||||
boxim.messages.append(
|
||||
{"id": 21, "localId": 3, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "在吗"}
|
||||
)
|
||||
with patch("routers.chat._resolve_reply", return_value={"answer": "在的"}):
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
with patch("services.takeover_service.httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {"code": 200, "data": {"answer": "Hello back"}}
|
||||
mock_client_class.return_value.__aenter__.return_value.post.return_value = mock_response
|
||||
clock.advance(2)
|
||||
boxim.messages.append(
|
||||
{"id": 22, "localId": 4, "sendId": 100, "recvId": 200, "sendTime": clock.millis(), "type": 0, "content": "我来回复"}
|
||||
)
|
||||
await service.poll_and_process_messages()
|
||||
clock.advance(2)
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
service = TakeoverService(mock_db, mock_boxim)
|
||||
message = {"from_accid": "user_acc", "content": "Hello"}
|
||||
|
||||
result = await service.execute_takeover(mock_auth, message)
|
||||
|
||||
assert result is True
|
||||
mock_boxim.get_credentials.assert_called_once_with("owner_huihui_123")
|
||||
mock_boxim.send_p2p_message.assert_called_once()
|
||||
db = session_factory()
|
||||
try:
|
||||
task = db.query(TakeoverReplyTask).filter(TakeoverReplyTask.trigger_message_id == "21").one()
|
||||
assert task.status == "cancelled"
|
||||
assert task.cancel_reason == "owner_replied"
|
||||
assert boxim.sent == []
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_takeover_fails_when_avatar_not_found(mock_db, mock_boxim, mock_auth):
|
||||
"""execute_takeover should return False when Avatar is not found."""
|
||||
avatar_filter = MagicMock()
|
||||
avatar_filter.first.return_value = None
|
||||
avatar_query = MagicMock()
|
||||
avatar_query.filter.return_value = avatar_filter
|
||||
mock_db.query.return_value = avatar_query
|
||||
async def test_quick_successive_messages_are_coalesced_into_one_reply(service_context):
|
||||
session_factory, service, boxim, clock = service_context
|
||||
await service.poll_and_process_messages()
|
||||
boxim.messages.append(
|
||||
{"id": 31, "localId": 5, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "第一句"}
|
||||
)
|
||||
with patch("routers.chat._resolve_reply", return_value={"answer": "第一版"}):
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
service = TakeoverService(mock_db, mock_boxim)
|
||||
message = {"from_accid": "user_acc", "content": "Hello"}
|
||||
clock.advance(1)
|
||||
boxim.messages.append(
|
||||
{"id": 32, "localId": 6, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "第二句"}
|
||||
)
|
||||
with patch("routers.chat._resolve_reply", return_value={"answer": "合并回复"}) as resolver:
|
||||
await service.poll_and_process_messages()
|
||||
assert resolver.call_args.args[2] == "第一句\n第二句"
|
||||
|
||||
result = await service.execute_takeover(mock_auth, message)
|
||||
clock.advance(3)
|
||||
await service.poll_and_process_messages()
|
||||
assert [item["content"] for item in boxim.sent] == ["合并回复"]
|
||||
|
||||
assert result is False
|
||||
mock_boxim.get_credentials.assert_not_called()
|
||||
db = session_factory()
|
||||
try:
|
||||
tasks = db.query(TakeoverReplyTask).order_by(TakeoverReplyTask.created_at).all()
|
||||
assert [task.status for task in tasks] == ["cancelled", "sent"]
|
||||
assert tasks[0].cancel_reason == "newer_incoming_message"
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_takeover_fails_when_no_credentials(mock_db, mock_boxim, mock_auth, mock_avatar):
|
||||
"""execute_takeover should return False when boxim.get_credentials returns None."""
|
||||
avatar_filter = MagicMock()
|
||||
avatar_filter.first.return_value = mock_avatar
|
||||
avatar_query = MagicMock()
|
||||
avatar_query.filter.return_value = avatar_filter
|
||||
mock_db.query.return_value = avatar_query
|
||||
async def test_connection_failure_disables_takeover_and_stops_retrying(service_context):
|
||||
session_factory, service, boxim, _ = service_context
|
||||
boxim.exchange_access_token = AsyncMock(
|
||||
side_effect=BoxIMError("无效的访问令牌", code=40101, auth_error=True)
|
||||
)
|
||||
|
||||
mock_boxim.get_credentials.return_value = None
|
||||
service = TakeoverService(mock_db, mock_boxim)
|
||||
message = {"from_accid": "user_acc", "content": "Hello"}
|
||||
await service.poll_and_process_messages()
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
result = await service.execute_takeover(mock_auth, message)
|
||||
|
||||
assert result is False
|
||||
db = session_factory()
|
||||
try:
|
||||
avatar = db.query(Avatar).one()
|
||||
cursor = db.query(TakeoverCursor).one()
|
||||
assert "takeover" not in avatar.config["authorizationPermissions"]
|
||||
assert cursor.initialized is False
|
||||
assert "重新登录" in cursor.last_error
|
||||
assert db.query(TakeoverReplyTask).count() == 0
|
||||
finally:
|
||||
db.close()
|
||||
boxim.exchange_access_token.assert_awaited_once_with("prod-huihui-token")
|
||||
|
||||
|
||||
# --- enqueue_delayed_message ---
|
||||
|
||||
|
||||
def test_enqueue_delayed_message_with_redis(mock_db, mock_boxim, mock_auth, mock_avatar):
|
||||
mock_redis = MagicMock()
|
||||
avatar_filter = MagicMock()
|
||||
avatar_filter.first.return_value = mock_avatar
|
||||
avatar_query = MagicMock()
|
||||
avatar_query.filter.return_value = avatar_filter
|
||||
mock_db.query.return_value = avatar_query
|
||||
|
||||
service = TakeoverService(mock_db, mock_boxim, mock_redis)
|
||||
message = {"msg_id": "msg_1", "from_accid": "user_acc", "content": "Hello"}
|
||||
|
||||
service.enqueue_delayed_message(mock_auth, message)
|
||||
|
||||
mock_redis.setex.assert_called_once()
|
||||
call_args = mock_redis.setex.call_args
|
||||
value = call_args[0][1]
|
||||
import json
|
||||
payload = json.loads(call_args[0][2])
|
||||
assert payload["owner_huihui_id"] == "owner_huihui_123"
|
||||
|
||||
|
||||
def test_enqueue_delayed_message_without_redis_logs_warning(mock_db, mock_boxim, mock_auth, mock_avatar):
|
||||
"""When Redis is not configured, enqueue_delayed_message should log a warning and not crash."""
|
||||
avatar_filter = MagicMock()
|
||||
avatar_filter.first.return_value = mock_avatar
|
||||
avatar_query = MagicMock()
|
||||
avatar_query.filter.return_value = avatar_filter
|
||||
mock_db.query.return_value = avatar_query
|
||||
|
||||
service = TakeoverService(mock_db, mock_boxim)
|
||||
message = {"msg_id": "msg_1", "from_accid": "user_acc", "content": "Hello"}
|
||||
|
||||
service.enqueue_delayed_message(mock_auth, message)
|
||||
|
||||
|
||||
# --- process_delayed_queue ---
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_delayed_queue_no_redis(mock_db, mock_boxim):
|
||||
"""process_delayed_queue should return immediately without Redis."""
|
||||
service = TakeoverService(mock_db, mock_boxim)
|
||||
await service.process_delayed_queue()
|
||||
mock_db.query.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_delayed_queue_processes_messages(mock_db, mock_boxim, mock_auth, mock_avatar):
|
||||
"""process_delayed_queue should read from Redis, resolve auth, and execute takeover."""
|
||||
mock_redis = MagicMock()
|
||||
mock_redis.keys.return_value = ["takeover:delayed:target_user_123:msg_1"]
|
||||
mock_redis.get.return_value = '{"from_accid": "user_acc", "content": "Hello"}'
|
||||
|
||||
avatar_filter = MagicMock()
|
||||
avatar_filter.first.return_value = mock_avatar
|
||||
avatar_query = MagicMock()
|
||||
avatar_query.filter.return_value = avatar_filter
|
||||
|
||||
auth_filter = MagicMock()
|
||||
auth_filter.filter.return_value = auth_filter
|
||||
auth_filter.first.return_value = mock_auth
|
||||
|
||||
def query_side_effect(model):
|
||||
if model == Avatar:
|
||||
return avatar_query
|
||||
return auth_filter
|
||||
|
||||
mock_db.query.side_effect = query_side_effect
|
||||
|
||||
with patch("services.takeover_service.httpx.AsyncClient") as mock_client_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {"code": 200, "data": {"answer": "Hello back"}}
|
||||
mock_client_class.return_value.__aenter__.return_value.post.return_value = mock_response
|
||||
|
||||
service = TakeoverService(mock_db, mock_boxim, mock_redis)
|
||||
await service.process_delayed_queue()
|
||||
|
||||
mock_boxim.send_p2p_message.assert_called_once()
|
||||
mock_redis.delete.assert_called_once()
|
||||
def test_plain_text_reply_removes_markdown_and_empty_lines():
|
||||
assert _plain_text_reply("## 建议\n\n**不能自行用药**\n`必要时就医`") == "建议\n不能自行用药\n必要时就医"
|
||||
|
||||
@@ -43,6 +43,7 @@ const showNav = ref<boolean>(shouldShowNav(route.path))
|
||||
|
||||
function shouldShowNav(path: string) {
|
||||
return path !== '/'
|
||||
&& path !== '/authorization'
|
||||
&& path !== '/avatar/create'
|
||||
&& path !== '/login/sms'
|
||||
&& !path.startsWith('/avatar/edit')
|
||||
|
||||
@@ -145,6 +145,30 @@ export const chargeToken = (planId: string) =>
|
||||
|
||||
// ==================== 授权管理 API ====================
|
||||
|
||||
export type AvatarPermission = 'friend' | 'chat' | 'publish' | 'browse' | 'interact' | 'takeover'
|
||||
|
||||
export interface AvatarPermissionSettings {
|
||||
avatarId: string
|
||||
permissions: AvatarPermission[]
|
||||
}
|
||||
|
||||
export const getAvatarPermissionSettings = (avatarId: string) =>
|
||||
request.get<AvatarPermissionSettings>(`/avatar/${avatarId}/permission-settings`)
|
||||
|
||||
export const updateAvatarPermissionSettings = (avatarId: string, permissions: AvatarPermission[]) =>
|
||||
request.put<AvatarPermissionSettings>(`/avatar/${avatarId}/permission-settings`, { permissions })
|
||||
|
||||
export interface TakeoverStatus {
|
||||
enabled: boolean
|
||||
status: 'disabled' | 'connecting' | 'ready' | 'needs_login' | 'error'
|
||||
message: string
|
||||
pendingCount: number
|
||||
lastPolledAt: string | null
|
||||
}
|
||||
|
||||
export const getTakeoverStatus = (avatarId: string) =>
|
||||
request.get<TakeoverStatus>(`/avatar/${avatarId}/takeover/status`)
|
||||
|
||||
export interface Authorization {
|
||||
id: string
|
||||
avatarId: string
|
||||
@@ -153,16 +177,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 ====================
|
||||
|
||||
@@ -202,6 +251,7 @@ export interface KnowledgeDoc {
|
||||
fileSize: number
|
||||
fileUrl: string
|
||||
status: string
|
||||
filePresent?: boolean
|
||||
vectorized?: boolean
|
||||
embeddingModel?: string
|
||||
chunkCount?: number
|
||||
@@ -382,13 +432,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
@@ -37,10 +37,10 @@
|
||||
<div class="card-content">
|
||||
<div class="card-title-row">
|
||||
<strong>{{ doc.filename }}</strong>
|
||||
<span class="status-pill" :class="{ pending: !doc.vectorized }">{{ doc.vectorized ? '已入库' : '处理中' }}</span>
|
||||
<span class="status-pill" :class="{ pending: !doc.vectorized && doc.filePresent !== false, missing: doc.filePresent === false }">{{ doc.filePresent === false ? '文件缺失' : (doc.vectorized ? '已入库' : '处理中') }}</span>
|
||||
</div>
|
||||
<p class="card-meta">{{ doc.fileType.toUpperCase() }} · {{ formatSize(doc.fileSize) }} · {{ formatDate(doc.createdAt) }}</p>
|
||||
<p class="card-detail">{{ doc.vectorized ? `已切分 ${doc.chunkCount || 0} 段,可用于对话` : '正在解析并建立知识索引' }}</p>
|
||||
<p class="card-detail">{{ doc.filePresent === false ? '原文件不可用,请删除后重新上传' : (doc.vectorized ? `已切分 ${doc.chunkCount || 0} 段,可用于对话` : '正在解析并建立知识索引') }}</p>
|
||||
</div>
|
||||
<button class="card-delete" @click="removeDoc(doc.id)">删除</button>
|
||||
</article>
|
||||
@@ -268,6 +268,7 @@ onMounted(async () => {
|
||||
.card-title-row { display: flex; align-items: center; gap: 8px; min-width: 0; }
|
||||
.card-title-row strong { min-width: 0; flex: 1; overflow: hidden; color: #27201C; font-size: 14px; text-overflow: ellipsis; white-space: nowrap; }
|
||||
.status-pill { flex: 0 0 auto; display: inline-flex; padding: 4px 7px; border-radius: 999px; color: #15803D; background: #ECFDF3; font-size: 10px; white-space: nowrap; }.status-pill.pending { color: #B45309; background: #FFFBEB; }
|
||||
.status-pill.missing { color: #B91C1C; background: #FEF2F2; }
|
||||
.card-meta, .card-detail { margin: 5px 0 0; color: #9398AE; font-size: 11px; line-height: 1.4; }.card-detail { color: #8B6B58; }
|
||||
.card-delete { flex: 0 0 auto; align-self: center; border: 0; color: #EF4444; background: #FEF2F2; border-radius: 8px; padding: 7px 9px; font-size: 12px; cursor: pointer; }
|
||||
.card-empty { padding: 42px 16px; border: 1px dashed #F1D9C3; border-radius: 16px; color: #9398AE; background: #fff; font-size: 14px; text-align: center; }
|
||||
|
||||
Reference in New Issue
Block a user