Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
434caac056 |
@@ -1,16 +1,24 @@
|
||||
"""数字分身管理服务层 — 同步连接数字分身应用的 SQLite 数据库"""
|
||||
import json
|
||||
import os
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Optional, Tuple
|
||||
|
||||
from sqlalchemy import create_engine, text
|
||||
from sqlalchemy import create_engine, select, text
|
||||
from sqlalchemy.orm import sessionmaker, Session
|
||||
|
||||
from app.core.config import settings
|
||||
from app.core.logger import logger
|
||||
from app.models import UserPersonality, VirtualUser
|
||||
|
||||
|
||||
_engine = None
|
||||
_SessionLocal: Optional[sessionmaker] = None
|
||||
|
||||
AVATAR_ACCOUNT_PREFIX = "__avatar__:"
|
||||
SQUARE_INTERACTION_PERMISSION = "interact"
|
||||
SQUARE_INTERACTION_ACTIONS = frozenset({"like", "collect", "comment", "reply"})
|
||||
|
||||
|
||||
def _get_engine_and_session():
|
||||
global _engine, _SessionLocal
|
||||
@@ -66,6 +74,185 @@ def _get_global_token_balance(db: Session) -> int:
|
||||
return 0
|
||||
|
||||
|
||||
def _decode_config(value) -> dict:
|
||||
if isinstance(value, dict):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
decoded = json.loads(value)
|
||||
return decoded if isinstance(decoded, dict) else {}
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return {}
|
||||
return {}
|
||||
|
||||
|
||||
def is_delegated_avatar_user(user: VirtualUser | None) -> bool:
|
||||
return bool(user and (user.account or "").startswith(AVATAR_ACCOUNT_PREFIX))
|
||||
|
||||
|
||||
def delegated_avatar_id(user: VirtualUser | None) -> str:
|
||||
if not is_delegated_avatar_user(user):
|
||||
return ""
|
||||
return (user.account or "")[len(AVATAR_ACCOUNT_PREFIX):]
|
||||
|
||||
|
||||
def _list_square_interaction_authorizations(db: Session) -> list[dict]:
|
||||
"""读取已明确授权分身参与广场互动的身份与会会令牌。"""
|
||||
rows = db.execute(text("""
|
||||
SELECT
|
||||
a.id AS avatar_id,
|
||||
a.name AS avatar_name,
|
||||
a.display_name AS avatar_display_name,
|
||||
a.description AS avatar_description,
|
||||
a.photo_url AS avatar_photo_url,
|
||||
a.config AS avatar_config,
|
||||
u.huihui_user_id,
|
||||
u.nickname AS owner_nickname,
|
||||
u.avatar_url AS owner_avatar_url,
|
||||
u.huihui_token
|
||||
FROM avatars a
|
||||
JOIN users u ON u.huihui_user_id = a.owner_id
|
||||
WHERE a.status = 'active'
|
||||
""")).fetchall()
|
||||
|
||||
authorized = []
|
||||
for row in rows:
|
||||
config = _decode_config(row.avatar_config)
|
||||
permissions = config.get("authorizationPermissions", [])
|
||||
if not isinstance(permissions, list) or SQUARE_INTERACTION_PERMISSION not in permissions:
|
||||
continue
|
||||
platform_uid = str(row.huihui_user_id or "").strip()
|
||||
token = str(row.huihui_token or "").strip()
|
||||
if not platform_uid or not token:
|
||||
continue
|
||||
authorized.append({
|
||||
"avatar_id": str(row.avatar_id),
|
||||
"avatar_name": row.avatar_display_name or row.avatar_name or row.owner_nickname or "数字分身",
|
||||
"avatar_description": row.avatar_description or "",
|
||||
"avatar_url": _resolve_photo_url(row.avatar_photo_url or row.owner_avatar_url or ""),
|
||||
"config": config,
|
||||
"platform_uid": platform_uid,
|
||||
"token": token,
|
||||
})
|
||||
return authorized
|
||||
|
||||
|
||||
def get_square_interaction_permissions(avatar_id: str) -> frozenset[str]:
|
||||
"""实时复核授权;数据库不可用、令牌失效或撤权时一律拒绝执行。"""
|
||||
avatar_db = get_session()
|
||||
if avatar_db is None:
|
||||
return frozenset()
|
||||
try:
|
||||
authorized_ids = {
|
||||
item["avatar_id"] for item in _list_square_interaction_authorizations(avatar_db)
|
||||
}
|
||||
return SQUARE_INTERACTION_ACTIONS if avatar_id in authorized_ids else frozenset()
|
||||
except Exception as exc:
|
||||
logger.error(f"读取数字分身广场互动授权失败: {exc}")
|
||||
return frozenset()
|
||||
finally:
|
||||
avatar_db.close()
|
||||
|
||||
|
||||
def _word_count_range(config: dict) -> tuple[int, int]:
|
||||
ranges = {
|
||||
"short": (10, 35),
|
||||
"medium": (20, 60),
|
||||
"long": (30, 80),
|
||||
}
|
||||
return ranges.get(str(config.get("responseLength") or "medium"), (20, 60))
|
||||
|
||||
|
||||
async def sync_square_interaction_users(db) -> set[str]:
|
||||
"""把已授权分身同步为调度器身份,并刷新其会会会话。"""
|
||||
avatar_db = get_session()
|
||||
if avatar_db is None:
|
||||
logger.warning("数字分身数据库不可用,跳过广场互动授权同步")
|
||||
return set()
|
||||
try:
|
||||
authorized = _list_square_interaction_authorizations(avatar_db)
|
||||
except Exception as exc:
|
||||
logger.error(f"同步数字分身广场互动授权失败: {exc}")
|
||||
return set()
|
||||
finally:
|
||||
avatar_db.close()
|
||||
|
||||
from app.core.redis_client import delete_session, set_session
|
||||
|
||||
result = await db.execute(
|
||||
select(VirtualUser).where(VirtualUser.account.like(f"{AVATAR_ACCOUNT_PREFIX}%"))
|
||||
)
|
||||
existing_users = {delegated_avatar_id(user): user for user in result.scalars().all()}
|
||||
authorized_ids = {item["avatar_id"] for item in authorized}
|
||||
|
||||
for avatar_id, user in existing_users.items():
|
||||
if avatar_id not in authorized_ids:
|
||||
user.is_enabled = 0
|
||||
user.status = 0
|
||||
user.session_token = None
|
||||
user.session_expires_at = None
|
||||
await delete_session(user.id)
|
||||
|
||||
for item in authorized:
|
||||
avatar_id = item["avatar_id"]
|
||||
user = existing_users.get(avatar_id)
|
||||
if user is None:
|
||||
user = VirtualUser(
|
||||
nickname=item["avatar_name"],
|
||||
account=f"{AVATAR_ACCOUNT_PREFIX}{avatar_id}",
|
||||
password_enc="",
|
||||
status=2,
|
||||
is_enabled=1,
|
||||
platform_uid=item["platform_uid"],
|
||||
remark="用户授权的数字分身广场互动身份",
|
||||
)
|
||||
db.add(user)
|
||||
await db.flush()
|
||||
|
||||
expires_at = datetime.now() + timedelta(days=1)
|
||||
user.nickname = item["avatar_name"]
|
||||
user.real_name = item["avatar_name"]
|
||||
user.avatar_url = item["avatar_url"]
|
||||
user.platform_uid = item["platform_uid"]
|
||||
user.session_token = item["token"]
|
||||
user.session_expires_at = expires_at
|
||||
user.last_login_at = datetime.now()
|
||||
user.status = 2
|
||||
user.is_enabled = 1
|
||||
|
||||
config = item["config"]
|
||||
personality_result = await db.execute(
|
||||
select(UserPersonality).where(UserPersonality.user_id == user.id)
|
||||
)
|
||||
personality = personality_result.scalar_one_or_none()
|
||||
word_min, word_max = _word_count_range(config)
|
||||
prompt_parts = [item["avatar_description"], str(config.get("systemPrompt") or "")]
|
||||
style_prompt = "\n".join(part.strip() for part in prompt_parts if part and part.strip())
|
||||
if personality is None:
|
||||
personality = UserPersonality(user_id=user.id)
|
||||
db.add(personality)
|
||||
personality.language_style = str(config.get("replyStyle") or "professional")
|
||||
personality.personality_desc = item["avatar_description"]
|
||||
personality.comment_style_prompt = style_prompt
|
||||
personality.word_count_min = word_min
|
||||
personality.word_count_max = word_max
|
||||
|
||||
await set_session(user.id, {
|
||||
"token": item["token"],
|
||||
"session_id": f"avatar:{avatar_id}",
|
||||
"platform_uid": item["platform_uid"],
|
||||
"org_id": "",
|
||||
"login_time": datetime.now().isoformat(),
|
||||
"nickname": item["avatar_name"],
|
||||
"real_name": item["avatar_name"],
|
||||
"avatar": item["avatar_url"],
|
||||
"delegated_avatar_id": avatar_id,
|
||||
}, expire=86400)
|
||||
|
||||
await db.commit()
|
||||
return authorized_ids
|
||||
|
||||
|
||||
class AvatarService:
|
||||
|
||||
@staticmethod
|
||||
|
||||
@@ -23,6 +23,7 @@ class SchedulerService:
|
||||
from app.core.database import AsyncSessionLocal
|
||||
logger.info("⚡ 立即触发互动任务")
|
||||
async with AsyncSessionLocal() as session:
|
||||
await self._sync_delegated_avatar_users(session)
|
||||
try:
|
||||
max_concurrent = int(await self._get_config(session, "max_concurrent_users", "5"))
|
||||
except (TypeError, ValueError):
|
||||
@@ -146,7 +147,9 @@ class SchedulerService:
|
||||
async def _check_sessions(self):
|
||||
"""定时校验登录状态"""
|
||||
from app.services.news_service import news_service
|
||||
from app.services.avatar_service import is_delegated_avatar_user
|
||||
async with AsyncSessionLocal() as db:
|
||||
await self._sync_delegated_avatar_users(db)
|
||||
result = await db.execute(
|
||||
select(VirtualUser).where(VirtualUser.status == 2, VirtualUser.is_enabled == 1)
|
||||
)
|
||||
@@ -154,7 +157,7 @@ class SchedulerService:
|
||||
for user in users:
|
||||
try:
|
||||
valid = await news_service.check_session(db, user)
|
||||
if not valid:
|
||||
if not valid and not is_delegated_avatar_user(user):
|
||||
logger.warning(f"用户 {user.account} 会话失效,尝试重登")
|
||||
await news_service.login(db, user)
|
||||
except Exception as e:
|
||||
@@ -163,6 +166,7 @@ class SchedulerService:
|
||||
async def _run_interactions(self):
|
||||
"""执行互动任务"""
|
||||
async with AsyncSessionLocal() as db:
|
||||
await self._sync_delegated_avatar_users(db)
|
||||
# 检查调度器开关
|
||||
enabled = await self._get_config(db, "scheduler_enabled", "true")
|
||||
if enabled != "true":
|
||||
@@ -184,8 +188,11 @@ class SchedulerService:
|
||||
logger.debug(f"[调度] 当前北京时间 {now_time} 不在互动时段 {start_str}-{end_str}")
|
||||
return
|
||||
|
||||
# 获取最小互动间隔(秒)
|
||||
min_interval = int(await self._get_config(db, "interact_min_interval", "300"))
|
||||
# 获取互动间隔范围(秒),与调度设置页面字段保持一致
|
||||
min_interval = await self._get_int_config(db, "interact_interval_min", 300)
|
||||
max_interval = await self._get_int_config(db, "interact_interval_max", min_interval)
|
||||
min_interval = max(0, min_interval)
|
||||
max_interval = max(min_interval, max_interval)
|
||||
|
||||
# 获取最大并发
|
||||
max_concurrent = int(await self._get_config(db, "max_concurrent_users", "5"))
|
||||
@@ -204,7 +211,7 @@ class SchedulerService:
|
||||
await self._try_login_users(db)
|
||||
return
|
||||
|
||||
# 检查互动间隔:过滤掉最近 min_interval 秒内已互动的用户
|
||||
# 每个用户在其最小/最大间隔内取得稳定随机值,直到下次互动后再变化
|
||||
now_dt = datetime.now()
|
||||
eligible = []
|
||||
for u in all_users:
|
||||
@@ -212,11 +219,17 @@ class SchedulerService:
|
||||
eligible.append(u)
|
||||
else:
|
||||
elapsed = (now_dt - u.last_interact_at).total_seconds()
|
||||
if elapsed >= min_interval:
|
||||
interval = random.Random(
|
||||
f"{u.id}:{u.last_interact_at.isoformat()}"
|
||||
).randint(min_interval, max_interval)
|
||||
if elapsed >= interval:
|
||||
eligible.append(u)
|
||||
|
||||
if not eligible:
|
||||
logger.debug(f"[调度] 所有 {len(all_users)} 个用户在 {min_interval}s 内已互动,跳过本次")
|
||||
logger.debug(
|
||||
f"[调度] 所有 {len(all_users)} 个用户尚未达到 "
|
||||
f"{min_interval}-{max_interval}s 随机互动间隔,跳过本次"
|
||||
)
|
||||
return
|
||||
|
||||
# 按最后互动时间升序排序:最久没互动的用户优先
|
||||
@@ -257,10 +270,12 @@ class SchedulerService:
|
||||
async def _try_login_users(self, db):
|
||||
"""尝试登录未登录的用户"""
|
||||
from app.services.news_service import news_service
|
||||
from app.services.avatar_service import AVATAR_ACCOUNT_PREFIX
|
||||
result = await db.execute(
|
||||
select(VirtualUser).where(
|
||||
VirtualUser.status.in_([0, 3]),
|
||||
VirtualUser.is_enabled == 1
|
||||
VirtualUser.is_enabled == 1,
|
||||
~VirtualUser.account.like(f"{AVATAR_ACCOUNT_PREFIX}%"),
|
||||
).limit(3)
|
||||
)
|
||||
users = result.scalars().all()
|
||||
@@ -275,6 +290,11 @@ class SchedulerService:
|
||||
"""执行单用户互动 - 基于真实接口"""
|
||||
from app.services.news_service import news_service
|
||||
from app.services.ai_service import ai_service
|
||||
from app.services.avatar_service import (
|
||||
delegated_avatar_id,
|
||||
get_square_interaction_permissions,
|
||||
is_delegated_avatar_user,
|
||||
)
|
||||
|
||||
async with AsyncSessionLocal() as db:
|
||||
try:
|
||||
@@ -289,6 +309,23 @@ class SchedulerService:
|
||||
"interactions": [],
|
||||
}
|
||||
|
||||
allowed_actions = {"like", "collect", "comment", "reply", "forward"}
|
||||
if is_delegated_avatar_user(user):
|
||||
allowed_actions = set(
|
||||
get_square_interaction_permissions(delegated_avatar_id(user))
|
||||
)
|
||||
if not allowed_actions:
|
||||
user.status = 0
|
||||
user.is_enabled = 0
|
||||
await db.commit()
|
||||
return {
|
||||
"user_id": user.id,
|
||||
"account": user.account,
|
||||
"status": "skipped",
|
||||
"reason": "avatar_interaction_not_authorized",
|
||||
"interactions": [],
|
||||
}
|
||||
|
||||
# 检查今日评论限额
|
||||
can_comment = True
|
||||
if user.today_comment_count >= user.daily_comment_limit:
|
||||
@@ -398,14 +435,53 @@ class SchedulerService:
|
||||
interactions_done = []
|
||||
action_failures = []
|
||||
|
||||
# ① 先记录阅读(每次必做,模拟真实用户打开文章)
|
||||
done_on_this = today_done.get(news_id, set())
|
||||
wants = {
|
||||
"like": (
|
||||
"like" in allowed_actions
|
||||
and "like" not in done_on_this
|
||||
and random.random() < like_prob
|
||||
),
|
||||
"collect": (
|
||||
"collect" in allowed_actions
|
||||
and "collect" not in done_on_this
|
||||
and random.random() < collect_prob
|
||||
),
|
||||
"forward": (
|
||||
"forward" in allowed_actions
|
||||
and "forward" not in done_on_this
|
||||
and random.random() < forward_prob
|
||||
),
|
||||
"reply": (
|
||||
"reply" in allowed_actions
|
||||
and can_comment
|
||||
and personality is not None
|
||||
and random.random() < reply_prob
|
||||
),
|
||||
"comment": (
|
||||
"comment" in allowed_actions
|
||||
and can_comment
|
||||
and personality is not None
|
||||
and not already_commented_this
|
||||
and random.random() < comment_prob
|
||||
),
|
||||
}
|
||||
if not any(wants.values()):
|
||||
return {
|
||||
"user_id": user.id,
|
||||
"account": user.account,
|
||||
"status": "skipped",
|
||||
"reason": "no_actions_triggered",
|
||||
"interactions": [],
|
||||
"article_id": news_id,
|
||||
"article_title": news_title,
|
||||
}
|
||||
|
||||
# 只有动作命中调度概率后才打开文章
|
||||
await news_service.read_news(db, user, news_id)
|
||||
|
||||
# 今日已对此文章做过的互动类型
|
||||
done_on_this = today_done.get(news_id, set())
|
||||
|
||||
# ② 点赞(每篇文章每用户每天只点赞一次)
|
||||
if "like" not in done_on_this and random.random() < like_prob:
|
||||
if wants["like"]:
|
||||
success, err = await news_service.like_news(db, user, news_id, org_id=article_org_id, to_user_id=news_author, title=news_title)
|
||||
await self._save_record(db, user, news_id, news_title, "like", None, 0, success, err)
|
||||
if success:
|
||||
@@ -415,16 +491,17 @@ class SchedulerService:
|
||||
action_failures.append({"type": "like", "error": err})
|
||||
|
||||
# ③ 收藏(每篇文章每用户每天只收藏一次)
|
||||
if "collect" not in done_on_this and random.random() < collect_prob:
|
||||
if wants["collect"]:
|
||||
success, err = await news_service.collect_news(db, user, news_id, org_id=article_org_id, to_user_id=news_author, title=news_title)
|
||||
await self._save_record(db, user, news_id, news_title, "collect", None, 0, success, err)
|
||||
if success:
|
||||
interactions_done.append("collect")
|
||||
await self._incr_total(db, user_id)
|
||||
else:
|
||||
action_failures.append({"type": "collect", "error": err})
|
||||
|
||||
# ④ 转发(每篇文章每用户每天只转发一次)
|
||||
if "forward" not in done_on_this and random.random() < forward_prob:
|
||||
if wants["forward"]:
|
||||
success, err = await news_service.forward_news(db, user, news_id)
|
||||
await self._save_record(db, user, news_id, news_title, "forward", None, 0, success, err)
|
||||
if success:
|
||||
@@ -438,7 +515,7 @@ class SchedulerService:
|
||||
style_prompt = personality.comment_style_prompt or ""
|
||||
safe_word_max = min(personality.word_count_max, 80)
|
||||
|
||||
if random.random() < reply_prob:
|
||||
if wants["reply"]:
|
||||
reply_actions, reply_failures = await self._run_reply_interaction_chain(
|
||||
db=db,
|
||||
starter=user,
|
||||
@@ -455,7 +532,7 @@ class SchedulerService:
|
||||
action_failures.extend(reply_failures)
|
||||
|
||||
# 每篇文章每个用户每天只发一条顶层评论;回复不再要求先评论
|
||||
if not already_commented_this and random.random() < comment_prob:
|
||||
if wants["comment"]:
|
||||
comment_text, tokens = await ai_service.generate_comment(
|
||||
db, news_title, news_content,
|
||||
style_prompt, personality.word_count_min, safe_word_max
|
||||
@@ -679,6 +756,7 @@ class SchedulerService:
|
||||
|
||||
async with AsyncSessionLocal() as db:
|
||||
try:
|
||||
await self._sync_delegated_avatar_users(db)
|
||||
now = datetime.now()
|
||||
await db.execute(
|
||||
update(PendingReplyTask)
|
||||
@@ -706,6 +784,12 @@ class SchedulerService:
|
||||
logger.error(f"待发送回复队列处理异常: {e}")
|
||||
|
||||
async def _process_pending_reply_task(self, db, task: PendingReplyTask, news_service, ai_service):
|
||||
from app.services.avatar_service import (
|
||||
delegated_avatar_id,
|
||||
get_square_interaction_permissions,
|
||||
is_delegated_avatar_user,
|
||||
)
|
||||
|
||||
task.status = 1
|
||||
task.locked_at = datetime.now()
|
||||
task.attempts = (task.attempts or 0) + 1
|
||||
@@ -716,6 +800,13 @@ class SchedulerService:
|
||||
task.status = 3
|
||||
task.last_error = "用户未登录或已禁用"
|
||||
return
|
||||
if (
|
||||
is_delegated_avatar_user(actor)
|
||||
and "reply" not in get_square_interaction_permissions(delegated_avatar_id(actor))
|
||||
):
|
||||
task.status = 3
|
||||
task.last_error = "数字分身广场互动授权已撤销"
|
||||
return
|
||||
|
||||
reply_result = await self._post_contextual_reply(
|
||||
db=db,
|
||||
@@ -858,6 +949,16 @@ class SchedulerService:
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
|
||||
async def _sync_delegated_avatar_users(self, db):
|
||||
from app.services.avatar_service import sync_square_interaction_users
|
||||
|
||||
try:
|
||||
return await sync_square_interaction_users(db)
|
||||
except Exception as exc:
|
||||
await db.rollback()
|
||||
logger.error(f"数字分身广场互动身份同步异常: {exc}")
|
||||
return set()
|
||||
|
||||
async def _incr_total(self, db, user_id: int):
|
||||
await db.execute(
|
||||
update(VirtualUser).where(VirtualUser.id == user_id).values(
|
||||
|
||||
@@ -0,0 +1,126 @@
|
||||
import json
|
||||
import os
|
||||
import sqlite3
|
||||
import tempfile
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from app.services import avatar_service
|
||||
|
||||
|
||||
class AvatarSquareAuthorizationTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
fd, self.db_path = tempfile.mkstemp(suffix=".db")
|
||||
os.close(fd)
|
||||
connection = sqlite3.connect(self.db_path)
|
||||
connection.executescript("""
|
||||
CREATE TABLE users (
|
||||
huihui_user_id TEXT,
|
||||
nickname TEXT,
|
||||
avatar_url TEXT,
|
||||
huihui_token TEXT
|
||||
);
|
||||
CREATE TABLE avatars (
|
||||
id TEXT,
|
||||
owner_id TEXT,
|
||||
name TEXT,
|
||||
display_name TEXT,
|
||||
description TEXT,
|
||||
photo_url TEXT,
|
||||
config TEXT,
|
||||
status TEXT
|
||||
);
|
||||
""")
|
||||
connection.execute(
|
||||
"INSERT INTO users VALUES (?, ?, ?, ?)",
|
||||
("huihui-7", "主人", "/owner.jpg", "huihui-token"),
|
||||
)
|
||||
connection.commit()
|
||||
connection.close()
|
||||
avatar_service._engine = None
|
||||
avatar_service._SessionLocal = None
|
||||
self.path_patch = patch.object(avatar_service.settings, "AVATAR_DB_PATH", self.db_path)
|
||||
self.path_patch.start()
|
||||
|
||||
def tearDown(self):
|
||||
self.path_patch.stop()
|
||||
if avatar_service._engine is not None:
|
||||
avatar_service._engine.dispose()
|
||||
avatar_service._engine = None
|
||||
avatar_service._SessionLocal = None
|
||||
os.unlink(self.db_path)
|
||||
|
||||
def _insert_avatar(self, permissions, *, status="active", token=None):
|
||||
connection = sqlite3.connect(self.db_path)
|
||||
connection.execute(
|
||||
"INSERT INTO avatars VALUES (?, ?, ?, ?, ?, ?, ?, ?)",
|
||||
(
|
||||
"avatar-7",
|
||||
"huihui-7",
|
||||
"avatar",
|
||||
"小会",
|
||||
"语气友好,表达简洁",
|
||||
"/avatar.jpg",
|
||||
json.dumps({
|
||||
"authorizationPermissions": permissions,
|
||||
"replyStyle": "warm",
|
||||
"responseLength": "short",
|
||||
}),
|
||||
status,
|
||||
),
|
||||
)
|
||||
if token is not None:
|
||||
connection.execute(
|
||||
"UPDATE users SET huihui_token = ? WHERE huihui_user_id = ?",
|
||||
(token, "huihui-7"),
|
||||
)
|
||||
connection.commit()
|
||||
connection.close()
|
||||
|
||||
def test_interact_permission_exposes_only_requested_square_actions(self):
|
||||
self._insert_avatar(["chat", "interact"])
|
||||
|
||||
permissions = avatar_service.get_square_interaction_permissions("avatar-7")
|
||||
|
||||
self.assertEqual(
|
||||
permissions,
|
||||
frozenset({"like", "collect", "comment", "reply"}),
|
||||
)
|
||||
self.assertNotIn("forward", permissions)
|
||||
|
||||
def test_missing_permission_inactive_avatar_or_missing_token_denies_execution(self):
|
||||
scenarios = [
|
||||
(["chat"], "active", "huihui-token"),
|
||||
(["interact"], "inactive", "huihui-token"),
|
||||
(["interact"], "active", ""),
|
||||
]
|
||||
for permissions, status, token in scenarios:
|
||||
with self.subTest(permissions=permissions, status=status, token=token):
|
||||
connection = sqlite3.connect(self.db_path)
|
||||
connection.execute("DELETE FROM avatars")
|
||||
connection.commit()
|
||||
connection.close()
|
||||
self._insert_avatar(permissions, status=status, token=token)
|
||||
self.assertEqual(
|
||||
avatar_service.get_square_interaction_permissions("avatar-7"),
|
||||
frozenset(),
|
||||
)
|
||||
|
||||
def test_delegated_avatar_identity_is_recognized_without_matching_normal_users(self):
|
||||
delegated = SimpleNamespace(account="__avatar__:avatar-7")
|
||||
normal = SimpleNamespace(account="13800000000")
|
||||
|
||||
self.assertTrue(avatar_service.is_delegated_avatar_user(delegated))
|
||||
self.assertEqual(avatar_service.delegated_avatar_id(delegated), "avatar-7")
|
||||
self.assertFalse(avatar_service.is_delegated_avatar_user(normal))
|
||||
self.assertEqual(avatar_service.delegated_avatar_id(normal), "")
|
||||
|
||||
def test_response_length_maps_to_scheduler_comment_limits(self):
|
||||
self.assertEqual(avatar_service._word_count_range({"responseLength": "short"}), (10, 35))
|
||||
self.assertEqual(avatar_service._word_count_range({"responseLength": "long"}), (30, 80))
|
||||
self.assertEqual(avatar_service._word_count_range({"responseLength": "unknown"}), (20, 60))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -478,24 +478,6 @@ def _qa_requires_language_adaptation(question: str, answer: str) -> bool:
|
||||
)
|
||||
|
||||
|
||||
def _qa_requires_per_turn_rendering(
|
||||
question: str,
|
||||
answer: str,
|
||||
history: list[Any],
|
||||
) -> bool:
|
||||
"""Keep the direct QA fast path only when no conversation can bias language."""
|
||||
return bool(history) or _qa_requires_language_adaptation(question, answer)
|
||||
|
||||
|
||||
def _per_turn_language_instruction() -> str:
|
||||
return (
|
||||
"本轮语言覆盖指令:只根据紧随其后的最新用户消息判断本轮回答语言。"
|
||||
"即使此前整段对话一直使用另一种语言,只要最新消息切换了语言,本轮就必须立即切换到相同语言;"
|
||||
"不要沿用上一轮语言。若最新消息明确指定回答语言,以该指定为准;若混用多种语言,使用其中占主导的"
|
||||
"自然语言。不要说明你检测、切换或翻译了语言。"
|
||||
)
|
||||
|
||||
|
||||
def _canonicalize_question(value: str) -> str:
|
||||
value = _normalize_question(value)
|
||||
replacements = (
|
||||
@@ -727,8 +709,6 @@ def _build_prompt(
|
||||
messages = [{"role": "system", "content": system}]
|
||||
for item in history[-MAX_HISTORY_MESSAGES:]:
|
||||
messages.append({"role": item.role, "content": item.content} if hasattr(item, "role") else item)
|
||||
# Keep the language instruction adjacent to the current turn so long histories cannot override it.
|
||||
messages.append({"role": "system", "content": _per_turn_language_instruction()})
|
||||
messages.append({"role": "user", "content": question.strip()})
|
||||
return messages
|
||||
|
||||
@@ -865,7 +845,7 @@ def _resolve_reply(
|
||||
qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
|
||||
matched = _match_standard_qa(question, qa_pairs)
|
||||
adapt_qa_language = bool(
|
||||
matched and _qa_requires_per_turn_rendering(question, matched.answer, history)
|
||||
matched and _qa_requires_language_adaptation(question, matched.answer)
|
||||
)
|
||||
if matched and not adapt_qa_language and not image_contexts:
|
||||
return {"answer": matched.answer, "source": "qa", "references": []}
|
||||
@@ -960,7 +940,7 @@ def _stream_reply(
|
||||
qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
|
||||
matched = _match_standard_qa(question, qa_pairs)
|
||||
adapt_qa_language = bool(
|
||||
matched and _qa_requires_per_turn_rendering(question, matched.answer, history)
|
||||
matched and _qa_requires_language_adaptation(question, matched.answer)
|
||||
)
|
||||
messages, reservation = [], None
|
||||
if matched and not adapt_qa_language and not image_contexts:
|
||||
|
||||
@@ -11,7 +11,6 @@ from routers.chat import (
|
||||
_match_standard_qa,
|
||||
_public_avatar_payload,
|
||||
_qa_requires_language_adaptation,
|
||||
_qa_requires_per_turn_rendering,
|
||||
_require_owned_avatar,
|
||||
_resolve_reply,
|
||||
)
|
||||
@@ -87,41 +86,6 @@ class ChatOrchestrationTests(unittest.TestCase):
|
||||
self.assertTrue(_qa_requires_language_adaptation("안녕하세요", "你好"))
|
||||
self.assertFalse(_qa_requires_language_adaptation("你好", "您好"))
|
||||
|
||||
def test_conversation_qa_is_rendered_for_the_current_turn_language(self):
|
||||
history = [SimpleNamespace(role="user", content="Please answer in English.")]
|
||||
self.assertTrue(_qa_requires_per_turn_rendering("Quelle est votre adresse ?", "Our address is Test Road 1.", history))
|
||||
|
||||
fake_model = Mock(return_value="Notre adresse est Test Road 1.")
|
||||
result = _resolve_reply(
|
||||
None,
|
||||
self.avatar,
|
||||
"Quelle est votre adresse ?",
|
||||
history,
|
||||
qa_pairs=[SimpleNamespace(question="Quelle est votre adresse ?", answer="Our address is Test Road 1.", enabled=True)],
|
||||
search_fn=Mock(),
|
||||
model_client=fake_model,
|
||||
)
|
||||
|
||||
self.assertEqual(result["source"], "qa")
|
||||
self.assertEqual(result["answer"], "Notre adresse est Test Road 1.")
|
||||
messages = fake_model.call_args.kwargs["messages"]
|
||||
self.assertEqual(messages[-1], {"role": "user", "content": "Quelle est votre adresse ?"})
|
||||
self.assertEqual(messages[-2]["role"], "system")
|
||||
self.assertIn("本轮语言覆盖指令", messages[-2]["content"])
|
||||
self.assertIn("不要沿用上一轮语言", messages[-2]["content"])
|
||||
|
||||
def test_latest_user_message_has_an_adjacent_language_override(self):
|
||||
history = [
|
||||
SimpleNamespace(role="user", content="请用中文回答"),
|
||||
SimpleNamespace(role="assistant", content="好的,请问有什么可以帮你?"),
|
||||
]
|
||||
messages = _build_prompt(self.avatar, history, "What can you help me with?", [])
|
||||
|
||||
self.assertEqual(messages[-1], {"role": "user", "content": "What can you help me with?"})
|
||||
self.assertEqual(messages[-2]["role"], "system")
|
||||
self.assertIn("最新用户消息", messages[-2]["content"])
|
||||
self.assertIn("立即切换到相同语言", messages[-2]["content"])
|
||||
|
||||
def test_conversational_paraphrase_matches_standard_qa(self):
|
||||
for question in ("请问一下,你们公司在哪里呀?", "请问去你们那边怎么走"):
|
||||
with self.subTest(question=question):
|
||||
|
||||
@@ -192,7 +192,7 @@ const permissionItems: Array<{
|
||||
{
|
||||
key: 'interact',
|
||||
title: '广场互动操作',
|
||||
description: '点赞、收藏、评论、回复等操作',
|
||||
description: '允许分身代表你点赞、收藏、评论和回复;时段、间隔与触发概率由广场调度设置统一控制',
|
||||
tone: 'pink',
|
||||
},
|
||||
{
|
||||
|
||||
Reference in New Issue
Block a user