Compare commits

..
12 changed files with 792 additions and 23 deletions
+188 -1
View File
@@ -1,16 +1,24 @@
"""数字分身管理服务层 — 同步连接数字分身应用的 SQLite 数据库""" """数字分身管理服务层 — 同步连接数字分身应用的 SQLite 数据库"""
import json
import os import os
from datetime import datetime, timedelta
from typing import Optional, Tuple 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 sqlalchemy.orm import sessionmaker, Session
from app.core.config import settings from app.core.config import settings
from app.core.logger import logger
from app.models import UserPersonality, VirtualUser
_engine = None _engine = None
_SessionLocal: Optional[sessionmaker] = 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(): def _get_engine_and_session():
global _engine, _SessionLocal global _engine, _SessionLocal
@@ -66,6 +74,185 @@ def _get_global_token_balance(db: Session) -> int:
return 0 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: class AvatarService:
@staticmethod @staticmethod
+117 -16
View File
@@ -23,6 +23,7 @@ class SchedulerService:
from app.core.database import AsyncSessionLocal from app.core.database import AsyncSessionLocal
logger.info("⚡ 立即触发互动任务") logger.info("⚡ 立即触发互动任务")
async with AsyncSessionLocal() as session: async with AsyncSessionLocal() as session:
await self._sync_delegated_avatar_users(session)
try: try:
max_concurrent = int(await self._get_config(session, "max_concurrent_users", "5")) max_concurrent = int(await self._get_config(session, "max_concurrent_users", "5"))
except (TypeError, ValueError): except (TypeError, ValueError):
@@ -146,7 +147,9 @@ class SchedulerService:
async def _check_sessions(self): async def _check_sessions(self):
"""定时校验登录状态""" """定时校验登录状态"""
from app.services.news_service import news_service from app.services.news_service import news_service
from app.services.avatar_service import is_delegated_avatar_user
async with AsyncSessionLocal() as db: async with AsyncSessionLocal() as db:
await self._sync_delegated_avatar_users(db)
result = await db.execute( result = await db.execute(
select(VirtualUser).where(VirtualUser.status == 2, VirtualUser.is_enabled == 1) select(VirtualUser).where(VirtualUser.status == 2, VirtualUser.is_enabled == 1)
) )
@@ -154,7 +157,7 @@ class SchedulerService:
for user in users: for user in users:
try: try:
valid = await news_service.check_session(db, user) 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} 会话失效,尝试重登") logger.warning(f"用户 {user.account} 会话失效,尝试重登")
await news_service.login(db, user) await news_service.login(db, user)
except Exception as e: except Exception as e:
@@ -163,6 +166,7 @@ class SchedulerService:
async def _run_interactions(self): async def _run_interactions(self):
"""执行互动任务""" """执行互动任务"""
async with AsyncSessionLocal() as db: async with AsyncSessionLocal() as db:
await self._sync_delegated_avatar_users(db)
# 检查调度器开关 # 检查调度器开关
enabled = await self._get_config(db, "scheduler_enabled", "true") enabled = await self._get_config(db, "scheduler_enabled", "true")
if enabled != "true": if enabled != "true":
@@ -184,8 +188,11 @@ class SchedulerService:
logger.debug(f"[调度] 当前北京时间 {now_time} 不在互动时段 {start_str}-{end_str}") logger.debug(f"[调度] 当前北京时间 {now_time} 不在互动时段 {start_str}-{end_str}")
return 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")) max_concurrent = int(await self._get_config(db, "max_concurrent_users", "5"))
@@ -204,7 +211,7 @@ class SchedulerService:
await self._try_login_users(db) await self._try_login_users(db)
return return
# 检查互动间隔:过滤掉最近 min_interval 秒内已互动的用户 # 每个用户在其最小/最大间隔内取得稳定随机值,直到下次互动后再变化
now_dt = datetime.now() now_dt = datetime.now()
eligible = [] eligible = []
for u in all_users: for u in all_users:
@@ -212,11 +219,17 @@ class SchedulerService:
eligible.append(u) eligible.append(u)
else: else:
elapsed = (now_dt - u.last_interact_at).total_seconds() 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) eligible.append(u)
if not eligible: 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 return
# 按最后互动时间升序排序:最久没互动的用户优先 # 按最后互动时间升序排序:最久没互动的用户优先
@@ -257,10 +270,12 @@ class SchedulerService:
async def _try_login_users(self, db): async def _try_login_users(self, db):
"""尝试登录未登录的用户""" """尝试登录未登录的用户"""
from app.services.news_service import news_service from app.services.news_service import news_service
from app.services.avatar_service import AVATAR_ACCOUNT_PREFIX
result = await db.execute( result = await db.execute(
select(VirtualUser).where( select(VirtualUser).where(
VirtualUser.status.in_([0, 3]), VirtualUser.status.in_([0, 3]),
VirtualUser.is_enabled == 1 VirtualUser.is_enabled == 1,
~VirtualUser.account.like(f"{AVATAR_ACCOUNT_PREFIX}%"),
).limit(3) ).limit(3)
) )
users = result.scalars().all() users = result.scalars().all()
@@ -275,6 +290,11 @@ class SchedulerService:
"""执行单用户互动 - 基于真实接口""" """执行单用户互动 - 基于真实接口"""
from app.services.news_service import news_service from app.services.news_service import news_service
from app.services.ai_service import ai_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: async with AsyncSessionLocal() as db:
try: try:
@@ -289,6 +309,23 @@ class SchedulerService:
"interactions": [], "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 can_comment = True
if user.today_comment_count >= user.daily_comment_limit: if user.today_comment_count >= user.daily_comment_limit:
@@ -398,14 +435,53 @@ class SchedulerService:
interactions_done = [] interactions_done = []
action_failures = [] 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) 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) 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) await self._save_record(db, user, news_id, news_title, "like", None, 0, success, err)
if success: if success:
@@ -415,16 +491,17 @@ class SchedulerService:
action_failures.append({"type": "like", "error": err}) 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) 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) await self._save_record(db, user, news_id, news_title, "collect", None, 0, success, err)
if success: if success:
interactions_done.append("collect") interactions_done.append("collect")
await self._incr_total(db, user_id)
else: else:
action_failures.append({"type": "collect", "error": err}) 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) 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) await self._save_record(db, user, news_id, news_title, "forward", None, 0, success, err)
if success: if success:
@@ -438,7 +515,7 @@ class SchedulerService:
style_prompt = personality.comment_style_prompt or "" style_prompt = personality.comment_style_prompt or ""
safe_word_max = min(personality.word_count_max, 80) 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( reply_actions, reply_failures = await self._run_reply_interaction_chain(
db=db, db=db,
starter=user, starter=user,
@@ -455,7 +532,7 @@ class SchedulerService:
action_failures.extend(reply_failures) 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( comment_text, tokens = await ai_service.generate_comment(
db, news_title, news_content, db, news_title, news_content,
style_prompt, personality.word_count_min, safe_word_max style_prompt, personality.word_count_min, safe_word_max
@@ -679,6 +756,7 @@ class SchedulerService:
async with AsyncSessionLocal() as db: async with AsyncSessionLocal() as db:
try: try:
await self._sync_delegated_avatar_users(db)
now = datetime.now() now = datetime.now()
await db.execute( await db.execute(
update(PendingReplyTask) update(PendingReplyTask)
@@ -706,6 +784,12 @@ class SchedulerService:
logger.error(f"待发送回复队列处理异常: {e}") logger.error(f"待发送回复队列处理异常: {e}")
async def _process_pending_reply_task(self, db, task: PendingReplyTask, news_service, ai_service): 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.status = 1
task.locked_at = datetime.now() task.locked_at = datetime.now()
task.attempts = (task.attempts or 0) + 1 task.attempts = (task.attempts or 0) + 1
@@ -716,6 +800,13 @@ class SchedulerService:
task.status = 3 task.status = 3
task.last_error = "用户未登录或已禁用" task.last_error = "用户未登录或已禁用"
return 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( reply_result = await self._post_contextual_reply(
db=db, db=db,
@@ -858,6 +949,16 @@ class SchedulerService:
except (TypeError, ValueError): except (TypeError, ValueError):
return default 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): async def _incr_total(self, db, user_id: int):
await db.execute( await db.execute(
update(VirtualUser).where(VirtualUser.id == user_id).values( 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()
@@ -5,6 +5,7 @@ pydantic
python-multipart python-multipart
httpx httpx
pypdf pypdf
PyMuPDF>=1.24,<2
python-docx python-docx
openpyxl openpyxl
apscheduler>=3.10 apscheduler>=3.10
+22 -2
View File
@@ -478,6 +478,24 @@ 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: def _canonicalize_question(value: str) -> str:
value = _normalize_question(value) value = _normalize_question(value)
replacements = ( replacements = (
@@ -709,6 +727,8 @@ def _build_prompt(
messages = [{"role": "system", "content": system}] messages = [{"role": "system", "content": system}]
for item in history[-MAX_HISTORY_MESSAGES:]: for item in history[-MAX_HISTORY_MESSAGES:]:
messages.append({"role": item.role, "content": item.content} if hasattr(item, "role") else item) 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()}) messages.append({"role": "user", "content": question.strip()})
return messages return messages
@@ -845,7 +865,7 @@ def _resolve_reply(
qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all() qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
matched = _match_standard_qa(question, qa_pairs) matched = _match_standard_qa(question, qa_pairs)
adapt_qa_language = bool( adapt_qa_language = bool(
matched and _qa_requires_language_adaptation(question, matched.answer) matched and _qa_requires_per_turn_rendering(question, matched.answer, history)
) )
if matched and not adapt_qa_language and not image_contexts: if matched and not adapt_qa_language and not image_contexts:
return {"answer": matched.answer, "source": "qa", "references": []} return {"answer": matched.answer, "source": "qa", "references": []}
@@ -940,7 +960,7 @@ def _stream_reply(
qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all() qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
matched = _match_standard_qa(question, qa_pairs) matched = _match_standard_qa(question, qa_pairs)
adapt_qa_language = bool( adapt_qa_language = bool(
matched and _qa_requires_language_adaptation(question, matched.answer) matched and _qa_requires_per_turn_rendering(question, matched.answer, history)
) )
messages, reservation = [], None messages, reservation = [], None
if matched and not adapt_qa_language and not image_contexts: if matched and not adapt_qa_language and not image_contexts:
@@ -8,7 +8,8 @@ import threading
from datetime import datetime, timezone from datetime import datetime, timezone
from database import SessionLocal from database import SessionLocal
from models import KnowledgeChunk, KnowledgeDoc from models import Avatar, KnowledgeChunk, KnowledgeDoc
from services.pdf_ocr_service import extract_scanned_pdf_text
import embeddings import embeddings
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -76,7 +77,23 @@ class KnowledgeVectorizer:
self._set_progress(db, doc, "extracting", 8) self._set_progress(db, doc, "extracting", 8)
text = embeddings.extract_text(path, f".{doc.file_type}") text = embeddings.extract_text(path, f".{doc.file_type}")
self._set_progress(db, doc, "chunking", 22) if doc.file_type == "pdf" and not text.strip():
avatar = db.get(Avatar, doc.avatar_id)
if not avatar:
raise ValueError("文档所属分身不存在")
def ocr_progress(done: int, total: int):
percent = 8 + int((done / max(1, total)) * 20)
self._set_progress(db, doc, "ocr", min(percent, 28))
self._set_progress(db, doc, "ocr", 8)
text = extract_scanned_pdf_text(
db,
avatar,
path,
on_progress=ocr_progress,
)
self._set_progress(db, doc, "chunking", 29)
chunks = embeddings.chunk_text(text) chunks = embeddings.chunk_text(text)
if not chunks: if not chunks:
raise ValueError("文档没有可建立索引的文字内容") raise ValueError("文档没有可建立索引的文字内容")
@@ -0,0 +1,130 @@
"""OCR fallback for image-only PDF knowledge documents."""
import logging
import os
import time
from typing import Callable
from sqlalchemy.orm import Session
from models import Avatar
from services.chat_model_config import get_chat_model_config
from services.token_billing import (
estimate_fallback_usage,
release_reservation,
reserve_avatar_tokens,
settle_reservation,
)
from services.vision_service import call_vision_model, prepare_image
logger = logging.getLogger(__name__)
PDF_OCR_PROMPT = (
"请逐字转录这一页扫描文档中的全部可见文字和表格,只输出转录内容,不要解释,不要使用 Markdown 代码块。"
"保留标题、段落、项目编号、数值和自然换行;看不清的内容写作[无法辨认],不要猜测、纠错或补全。"
)
def _positive_int(name: str, default: int, minimum: int, maximum: int) -> int:
try:
value = int(os.getenv(name, str(default)))
except ValueError:
value = default
return max(minimum, min(maximum, value))
def extract_scanned_pdf_text(
db: Session,
avatar: Avatar,
path: str,
*,
on_progress: Callable[[int, int], None] | None = None,
) -> str:
"""Render and OCR an image-only PDF while preserving page order."""
try:
import pymupdf
except ImportError as exc:
raise RuntimeError("扫描型 PDF 识别组件未安装") from exc
max_pages = _positive_int("KNOWLEDGE_PDF_OCR_MAX_PAGES", 80, 1, 300)
render_dpi = _positive_int("KNOWLEDGE_PDF_OCR_DPI", 144, 96, 200)
max_attempts = _positive_int("KNOWLEDGE_PDF_OCR_ATTEMPTS", 3, 1, 5)
model_config = get_chat_model_config()
model = model_config.ocr_model or model_config.vision_model
if not model_config.api_key or not model:
raise RuntimeError("扫描型 PDF 需要配置视觉 OCR 模型")
texts: list[str] = []
with pymupdf.open(path) as document:
total_pages = document.page_count
if total_pages <= 0:
raise ValueError("PDF 没有可识别页面")
if total_pages > max_pages:
raise ValueError(
f"扫描型 PDF 共 {total_pages} 页,超过单次 OCR 上限 {max_pages} 页,请拆分后上传"
)
scale = render_dpi / 72
for page_index in range(total_pages):
page = document.load_page(page_index)
pixmap = page.get_pixmap(
matrix=pymupdf.Matrix(scale, scale),
colorspace=pymupdf.csRGB,
alpha=False,
)
prepared = prepare_image(pixmap.tobytes("jpeg", jpg_quality=88))
estimate_messages = [{
"role": "user",
"content": f"[扫描 PDF 第 {page_index + 1}/{total_pages} 页]\n{PDF_OCR_PROMPT}",
}]
reservation = reserve_avatar_tokens(
db,
avatar,
"knowledge_pdf_ocr",
model,
estimate_messages,
model_config.vision_max_tokens,
)
try:
result = None
for attempt in range(1, max_attempts + 1):
try:
result = call_vision_model(
prepared,
model_config,
model=model,
prompt=PDF_OCR_PROMPT,
json_output=False,
)
break
except RuntimeError:
if attempt == max_attempts:
raise
time.sleep(min(4, attempt))
content = str((result or {}).get("content") or "").strip()
if not content:
raise RuntimeError("扫描型 PDF 页面识别结果为空")
settle_reservation(
db,
reservation,
(result or {}).get("usage"),
fallback_total=estimate_fallback_usage(estimate_messages, content),
)
except Exception as exc:
release_reservation(db, reservation, str(exc))
raise RuntimeError(
f"扫描型 PDF 第 {page_index + 1}/{total_pages} 页识别失败:{exc}"
) from exc
texts.append(f"[第 {page_index + 1} 页]\n{content}")
if on_progress:
on_progress(page_index + 1, total_pages)
logger.info(
"Scanned PDF OCR completed avatar=%s page=%s/%s",
avatar.id,
page_index + 1,
total_pages,
)
return "\n\n".join(texts).strip()
@@ -11,6 +11,7 @@ from routers.chat import (
_match_standard_qa, _match_standard_qa,
_public_avatar_payload, _public_avatar_payload,
_qa_requires_language_adaptation, _qa_requires_language_adaptation,
_qa_requires_per_turn_rendering,
_require_owned_avatar, _require_owned_avatar,
_resolve_reply, _resolve_reply,
) )
@@ -86,6 +87,41 @@ class ChatOrchestrationTests(unittest.TestCase):
self.assertTrue(_qa_requires_language_adaptation("안녕하세요", "你好")) self.assertTrue(_qa_requires_language_adaptation("안녕하세요", "你好"))
self.assertFalse(_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): def test_conversational_paraphrase_matches_standard_qa(self):
for question in ("请问一下,你们公司在哪里呀?", "请问去你们那边怎么走"): for question in ("请问一下,你们公司在哪里呀?", "请问去你们那边怎么走"):
with self.subTest(question=question): with self.subTest(question=question):
@@ -238,6 +238,53 @@ def test_background_vectorizer_keeps_failure_reason_for_retry(
db.close() db.close()
def test_background_vectorizer_uses_ocr_for_image_only_pdf(
tmp_path: Path,
authorization_context,
):
context = authorization_context
with (
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
patch("routers.knowledge.knowledge_vectorizer.enqueue"),
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
headers=context["owner_headers"],
files={"file": ("scanned.pdf", b"image-only-pdf", "application/pdf")},
)
payload = response.json()["data"]
progress = []
with (
patch("services.knowledge_vectorizer.UPLOAD_DIR", str(tmp_path)),
patch("services.knowledge_vectorizer.embeddings.extract_text", return_value=""),
patch(
"services.knowledge_vectorizer.extract_scanned_pdf_text",
side_effect=lambda _db, _avatar, _path, on_progress: (
on_progress(1, 2), on_progress(2, 2), "扫描页文字"
)[-1],
) as ocr,
patch("services.knowledge_vectorizer.embeddings.embed", return_value=[[1.0, 0.0]]),
patch.object(knowledge_vectorizer, "_set_progress", wraps=knowledge_vectorizer._set_progress) as set_progress,
):
knowledge_vectorizer.vectorize_document(payload["id"])
progress = [(call.args[2], call.args[3]) for call in set_progress.call_args_list]
db = SessionLocal()
try:
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
assert stored.status == "ready"
assert stored.chunk_count == 1
assert ("ocr", 18) in progress
assert ("ocr", 28) in progress
ocr.assert_called_once()
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).delete()
db.delete(stored)
db.commit()
finally:
db.close()
def test_retry_queues_a_failed_document_again( def test_retry_queues_a_failed_document_again(
tmp_path: Path, tmp_path: Path,
authorization_context, authorization_context,
@@ -0,0 +1,104 @@
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from services.pdf_ocr_service import extract_scanned_pdf_text
class FakePixmap:
def tobytes(self, *_args, **_kwargs):
return b"jpeg-page"
class FakePage:
def get_pixmap(self, **_kwargs):
return FakePixmap()
class FakeDocument:
page_count = 2
def __enter__(self):
return self
def __exit__(self, *_args):
return None
def load_page(self, _index):
return FakePage()
def test_scanned_pdf_ocr_preserves_page_order_and_reports_progress(monkeypatch):
fake_pymupdf = SimpleNamespace(
open=lambda _path: FakeDocument(),
Matrix=lambda x, y: (x, y),
csRGB="rgb",
)
monkeypatch.setitem(__import__("sys").modules, "pymupdf", fake_pymupdf)
progress = []
reservation = SimpleNamespace()
config = SimpleNamespace(
api_key="configured",
ocr_model="qwen-vl-ocr",
vision_model="vision",
vision_max_tokens=2048,
)
with (
patch("services.pdf_ocr_service.get_chat_model_config", return_value=config),
patch("services.pdf_ocr_service.prepare_image", return_value=SimpleNamespace()),
patch(
"services.pdf_ocr_service.call_vision_model",
side_effect=[
{"content": "第一页文字", "usage": {"total_tokens": 10}},
{"content": "第二页文字", "usage": {"total_tokens": 12}},
],
),
patch("services.pdf_ocr_service.reserve_avatar_tokens", return_value=reservation) as reserve,
patch("services.pdf_ocr_service.settle_reservation") as settle,
):
text = extract_scanned_pdf_text(
MagicMock(),
SimpleNamespace(id="avatar-1"),
"/tmp/scanned.pdf",
on_progress=lambda done, total: progress.append((done, total)),
)
assert text == "[第 1 页]\n第一页文字\n\n[第 2 页]\n第二页文字"
assert progress == [(1, 2), (2, 2)]
assert reserve.call_count == 2
assert settle.call_count == 2
def test_scanned_pdf_ocr_releases_tokens_after_retries_fail(monkeypatch):
fake_document = FakeDocument()
fake_document.page_count = 1
fake_pymupdf = SimpleNamespace(
open=lambda _path: fake_document,
Matrix=lambda x, y: (x, y),
csRGB="rgb",
)
monkeypatch.setitem(__import__("sys").modules, "pymupdf", fake_pymupdf)
monkeypatch.setenv("KNOWLEDGE_PDF_OCR_ATTEMPTS", "2")
reservation = SimpleNamespace()
config = SimpleNamespace(
api_key="configured",
ocr_model="qwen-vl-ocr",
vision_model="vision",
vision_max_tokens=2048,
)
with (
patch("services.pdf_ocr_service.get_chat_model_config", return_value=config),
patch("services.pdf_ocr_service.prepare_image", return_value=SimpleNamespace()),
patch("services.pdf_ocr_service.call_vision_model", side_effect=RuntimeError("timeout")) as call,
patch("services.pdf_ocr_service.reserve_avatar_tokens", return_value=reservation),
patch("services.pdf_ocr_service.release_reservation") as release,
patch("services.pdf_ocr_service.time.sleep"),
):
with pytest.raises(RuntimeError, match="第 1/1 页识别失败"):
extract_scanned_pdf_text(MagicMock(), SimpleNamespace(id="avatar-1"), "/tmp/scanned.pdf")
assert call.call_count == 2
release.assert_called_once()
@@ -192,7 +192,7 @@ const permissionItems: Array<{
{ {
key: 'interact', key: 'interact',
title: '广场互动操作', title: '广场互动操作',
description: '点赞、收藏、评论、回复等操作', description: '允许分身代表你点赞、收藏、评论和回复;时段、间隔与触发概率由广场调度设置统一控制',
tone: 'pink', tone: 'pink',
}, },
{ {
@@ -147,7 +147,7 @@ const documentState = (doc: any) => {
if (['uploaded', 'parsing'].includes(String(doc.status || '').toLowerCase())) { if (['uploaded', 'parsing'].includes(String(doc.status || '').toLowerCase())) {
const stage = String(doc.indexStage || 'queued').toLowerCase() const stage = String(doc.indexStage || 'queued').toLowerCase()
const labels: Record<string, string> = { const labels: Record<string, string> = {
queued: '等待处理', extracting: '解析文档', chunking: '切分文本', embedding: '向量化中' queued: '等待处理', extracting: '解析文档', ocr: '扫描件识别', chunking: '切分文本', embedding: '向量化中'
} }
const progress = Math.max(0, Math.min(99, Number(doc.indexProgress || 0))) const progress = Math.max(0, Math.min(99, Number(doc.indexProgress || 0)))
return { tone: 'pending', label: labels[stage] || '处理中', detail: `${labels[stage] || '正在建立知识索引'} ${progress}%`, progress } return { tone: 'pending', label: labels[stage] || '处理中', detail: `${labels[stage] || '正在建立知识索引'} ${progress}%`, progress }