Compare commits
20
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
434caac056 | ||
|
|
2a01a9946a | ||
|
|
2ce1079bb6 | ||
|
|
6fba6dbaaa | ||
|
|
e71267cf86 | ||
|
|
359e558dbe | ||
|
|
3edf92c7cc | ||
|
|
97c4c73b58 | ||
|
|
08c58fe0e6 | ||
|
|
6b7201e890 | ||
|
|
28553aba15 | ||
|
|
95f91450d0 | ||
|
|
b98a2b9507 | ||
|
|
59350fb41d | ||
|
|
6a4b35c49a | ||
|
|
207bbd02cf | ||
|
|
7cac96356d | ||
|
|
03c32309a8 | ||
|
|
0fc43908ae | ||
|
|
3d999f9472 |
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -1,16 +1,30 @@
|
|||||||
import os
|
import os
|
||||||
|
|
||||||
from sqlalchemy import create_engine
|
from sqlalchemy import create_engine, event
|
||||||
from sqlalchemy.orm import sessionmaker, declarative_base, Session
|
from sqlalchemy.orm import sessionmaker, declarative_base, Session
|
||||||
|
|
||||||
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
|
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||||
DB_FILE = os.path.join(BASE_DIR, "avatar.db")
|
DB_FILE = os.path.join(BASE_DIR, "avatar.db")
|
||||||
DATABASE_URL = os.getenv("DATABASE_URL", f"sqlite:///{DB_FILE}")
|
DATABASE_URL = os.getenv("DATABASE_URL", f"sqlite:///{DB_FILE}")
|
||||||
|
|
||||||
|
IS_SQLITE = DATABASE_URL.startswith("sqlite:")
|
||||||
engine = create_engine(
|
engine = create_engine(
|
||||||
DATABASE_URL,
|
DATABASE_URL,
|
||||||
connect_args={"check_same_thread": False} if DATABASE_URL.startswith("sqlite:") else {},
|
connect_args={"check_same_thread": False, "timeout": 30} if IS_SQLITE else {},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if IS_SQLITE:
|
||||||
|
@event.listens_for(engine, "connect")
|
||||||
|
def _configure_sqlite_connection(dbapi_connection, _connection_record):
|
||||||
|
cursor = dbapi_connection.cursor()
|
||||||
|
try:
|
||||||
|
cursor.execute("PRAGMA synchronous=NORMAL")
|
||||||
|
cursor.execute("PRAGMA busy_timeout=30000")
|
||||||
|
finally:
|
||||||
|
cursor.close()
|
||||||
|
|
||||||
|
|
||||||
SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
|
SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
|
||||||
Base = declarative_base()
|
Base = declarative_base()
|
||||||
|
|
||||||
@@ -26,6 +40,10 @@ def get_db():
|
|||||||
def init_db():
|
def init_db():
|
||||||
import models
|
import models
|
||||||
|
|
||||||
|
if IS_SQLITE:
|
||||||
|
with engine.connect() as conn:
|
||||||
|
conn.exec_driver_sql("PRAGMA journal_mode=WAL")
|
||||||
|
conn.commit()
|
||||||
Base.metadata.create_all(bind=engine)
|
Base.metadata.create_all(bind=engine)
|
||||||
|
|
||||||
# 轻量迁移:为已存在的表补充新列(SQLite 不支持自动 ALTER,逐列尝试)
|
# 轻量迁移:为已存在的表补充新列(SQLite 不支持自动 ALTER,逐列尝试)
|
||||||
@@ -35,6 +53,9 @@ def init_db():
|
|||||||
("knowledge_docs", "embedding_model", "VARCHAR DEFAULT ''"),
|
("knowledge_docs", "embedding_model", "VARCHAR DEFAULT ''"),
|
||||||
("knowledge_docs", "chunk_count", "INTEGER DEFAULT 0"),
|
("knowledge_docs", "chunk_count", "INTEGER DEFAULT 0"),
|
||||||
("knowledge_docs", "vectorized_at", "TIMESTAMP"),
|
("knowledge_docs", "vectorized_at", "TIMESTAMP"),
|
||||||
|
("knowledge_docs", "error_message", "VARCHAR DEFAULT ''"),
|
||||||
|
("knowledge_docs", "index_stage", "VARCHAR DEFAULT ''"),
|
||||||
|
("knowledge_docs", "index_progress", "INTEGER DEFAULT 0"),
|
||||||
("avatars", "owner_id", "VARCHAR DEFAULT ''"),
|
("avatars", "owner_id", "VARCHAR DEFAULT ''"),
|
||||||
("authorizations", "takeover_enabled", "BOOLEAN DEFAULT 0"),
|
("authorizations", "takeover_enabled", "BOOLEAN DEFAULT 0"),
|
||||||
("authorizations", "takeover_mode", "VARCHAR DEFAULT 'immediate'"),
|
("authorizations", "takeover_mode", "VARCHAR DEFAULT 'immediate'"),
|
||||||
@@ -45,6 +66,7 @@ def init_db():
|
|||||||
("token_account", "total_consumed", "BIGINT DEFAULT 0"),
|
("token_account", "total_consumed", "BIGINT DEFAULT 0"),
|
||||||
("token_account", "created_at", "TIMESTAMP"),
|
("token_account", "created_at", "TIMESTAMP"),
|
||||||
("token_account", "updated_at", "TIMESTAMP"),
|
("token_account", "updated_at", "TIMESTAMP"),
|
||||||
|
("takeover_messages", "attachment_id", "VARCHAR DEFAULT NULL"),
|
||||||
)
|
)
|
||||||
_normalize_optional_unique_values()
|
_normalize_optional_unique_values()
|
||||||
_normalize_takeover_delays()
|
_normalize_takeover_delays()
|
||||||
|
|||||||
@@ -51,7 +51,7 @@ def _hash_embedding(texts, dim=EMBED_DIM):
|
|||||||
return vecs
|
return vecs
|
||||||
|
|
||||||
|
|
||||||
def embed(texts):
|
def embed(texts, on_progress=None):
|
||||||
"""返回 list[list[float]],与输入顺序一致。"""
|
"""返回 list[list[float]],与输入顺序一致。"""
|
||||||
if not texts:
|
if not texts:
|
||||||
return []
|
return []
|
||||||
@@ -64,6 +64,7 @@ def embed(texts):
|
|||||||
except ValueError:
|
except ValueError:
|
||||||
batch_size = 10
|
batch_size = 10
|
||||||
embeddings = []
|
embeddings = []
|
||||||
|
total = len(texts)
|
||||||
for start in range(0, len(texts), batch_size):
|
for start in range(0, len(texts), batch_size):
|
||||||
batch = texts[start:start + batch_size]
|
batch = texts[start:start + batch_size]
|
||||||
payload = json.dumps({"input": batch, "model": model}).encode("utf-8")
|
payload = json.dumps({"input": batch, "model": model}).encode("utf-8")
|
||||||
@@ -84,8 +85,13 @@ def embed(texts):
|
|||||||
if len(items) != len(batch):
|
if len(items) != len(batch):
|
||||||
raise ValueError("embedding response count does not match request")
|
raise ValueError("embedding response count does not match request")
|
||||||
embeddings.extend(item["embedding"] for item in items)
|
embeddings.extend(item["embedding"] for item in items)
|
||||||
|
if on_progress:
|
||||||
|
on_progress(len(embeddings), total)
|
||||||
return embeddings
|
return embeddings
|
||||||
return _hash_embedding(texts)
|
vectors = _hash_embedding(texts)
|
||||||
|
if on_progress:
|
||||||
|
on_progress(len(vectors), len(texts))
|
||||||
|
return vectors
|
||||||
|
|
||||||
|
|
||||||
def cosine(a, b):
|
def cosine(a, b):
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ import routers.chat
|
|||||||
import routers.takeover
|
import routers.takeover
|
||||||
from responses import ok
|
from responses import ok
|
||||||
from services.chat_attachment_service import purge_expired_chat_attachments
|
from services.chat_attachment_service import purge_expired_chat_attachments
|
||||||
|
from services.knowledge_vectorizer import knowledge_vectorizer
|
||||||
from services.token_billing import DEFAULT_TOKEN_GRANT, release_stale_reservations
|
from services.token_billing import DEFAULT_TOKEN_GRANT, release_stale_reservations
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -131,6 +132,7 @@ def on_startup():
|
|||||||
|
|
||||||
init_db()
|
init_db()
|
||||||
seed()
|
seed()
|
||||||
|
knowledge_vectorizer.start()
|
||||||
|
|
||||||
# Release stale resources when startup is invoked again by a reload/test.
|
# Release stale resources when startup is invoked again by a reload/test.
|
||||||
stop_takeover_scheduler()
|
stop_takeover_scheduler()
|
||||||
|
|||||||
@@ -120,6 +120,7 @@ class TakeoverMessage(Base):
|
|||||||
direction = Column(String, nullable=False) # incoming | outgoing
|
direction = Column(String, nullable=False) # incoming | outgoing
|
||||||
message_type = Column(Integer, default=0)
|
message_type = Column(Integer, default=0)
|
||||||
content = Column(Text, default="")
|
content = Column(Text, default="")
|
||||||
|
attachment_id = Column(String, nullable=True)
|
||||||
is_avatar = Column(Boolean, default=False)
|
is_avatar = Column(Boolean, default=False)
|
||||||
send_time = Column(DateTime, nullable=False)
|
send_time = Column(DateTime, nullable=False)
|
||||||
created_at = Column(DateTime, server_default=func.now())
|
created_at = Column(DateTime, server_default=func.now())
|
||||||
@@ -189,6 +190,9 @@ class KnowledgeDoc(Base):
|
|||||||
file_size = Column(Integer, default=0)
|
file_size = Column(Integer, default=0)
|
||||||
file_url = Column(String, default="")
|
file_url = Column(String, default="")
|
||||||
status = Column(String, default="uploaded") # uploaded | parsing | ready | failed
|
status = Column(String, default="uploaded") # uploaded | parsing | ready | failed
|
||||||
|
error_message = Column(String, default="") # 建立索引失败原因
|
||||||
|
index_stage = Column(String, default="") # queued | extracting | chunking | embedding | ready | failed
|
||||||
|
index_progress = Column(Integer, default=0) # 0-100
|
||||||
vectorized = Column(Boolean, default=False) # 是否已向量化
|
vectorized = Column(Boolean, default=False) # 是否已向量化
|
||||||
embedding_model = Column(String, default="") # 向量模型标识
|
embedding_model = Column(String, default="") # 向量模型标识
|
||||||
chunk_count = Column(Integer, default=0) # 切片数量
|
chunk_count = Column(Integer, default=0) # 切片数量
|
||||||
@@ -204,6 +208,9 @@ class KnowledgeDoc(Base):
|
|||||||
"fileSize": self.file_size,
|
"fileSize": self.file_size,
|
||||||
"fileUrl": self.file_url,
|
"fileUrl": self.file_url,
|
||||||
"status": self.status,
|
"status": self.status,
|
||||||
|
"errorMessage": self.error_message or "",
|
||||||
|
"indexStage": self.index_stage or "",
|
||||||
|
"indexProgress": int(self.index_progress or 0),
|
||||||
"vectorized": bool(self.vectorized),
|
"vectorized": bool(self.vectorized),
|
||||||
"embeddingModel": self.embedding_model,
|
"embeddingModel": self.embedding_model,
|
||||||
"chunkCount": self.chunk_count,
|
"chunkCount": self.chunk_count,
|
||||||
@@ -263,7 +270,7 @@ class ChatAttachment(Base):
|
|||||||
__tablename__ = "chat_attachments"
|
__tablename__ = "chat_attachments"
|
||||||
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
|
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
|
||||||
avatar_id = Column(String, nullable=False, default="", index=True)
|
avatar_id = Column(String, nullable=False, default="", index=True)
|
||||||
uploader_kind = Column(String, default="owner") # owner | public
|
uploader_kind = Column(String, default="owner") # owner | public | boxim
|
||||||
filename = Column(String, default="")
|
filename = Column(String, default="")
|
||||||
mime_type = Column(String, default="")
|
mime_type = Column(String, default="")
|
||||||
file_size = Column(Integer, default=0)
|
file_size = Column(Integer, default=0)
|
||||||
|
|||||||
@@ -47,6 +47,25 @@ QA_SEMANTIC_THRESHOLD = 0.72
|
|||||||
QA_MATCH_MARGIN = 0.06
|
QA_MATCH_MARGIN = 0.06
|
||||||
KNOWLEDGE_MIN_SCORE = float(os.getenv("KNOWLEDGE_MIN_SCORE", "0.42"))
|
KNOWLEDGE_MIN_SCORE = float(os.getenv("KNOWLEDGE_MIN_SCORE", "0.42"))
|
||||||
|
|
||||||
|
_IMAGE_ACCESS_DENIAL_PATTERNS = (
|
||||||
|
re.compile(
|
||||||
|
r"(?:我|目前|暂时|这里|本身|系统)?\s*(?:无法|不能|没法|不支持)\s*"
|
||||||
|
r"(?:直接)?\s*(?:查看|看到|看见|识别|读取|访问|打开|分析|理解)"
|
||||||
|
r"(?:\s*(?:或|、|/)\s*(?:查看|看到|看见|识别|读取|访问|打开|分析|理解))*\s*"
|
||||||
|
r"(?:你(?:发|提供|上传)的|这张|该|当前)?\s*(?:图片|图像|照片|影像|文件)"
|
||||||
|
),
|
||||||
|
re.compile(
|
||||||
|
r"(?:我|这里|目前|暂时)?\s*(?:看不到|看不见|未看到|没有看到|没收到|未收到)\s*"
|
||||||
|
r"(?:你(?:发|提供|上传)的|这张|该|当前)?\s*(?:图片|图像|照片|影像)"
|
||||||
|
),
|
||||||
|
re.compile(
|
||||||
|
r"\b(?:i\s+)?(?:can(?:not|'t)|am\s+unable\s+to)\s+(?:directly\s+)?"
|
||||||
|
r"(?:view|see|access|read|analy[sz]e|recogni[sz]e)\s+"
|
||||||
|
r"(?:the\s+|this\s+|your\s+)?(?:image|photo|picture|scan)\b",
|
||||||
|
re.IGNORECASE,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
_WRITING_SYSTEM_PATTERNS = {
|
_WRITING_SYSTEM_PATTERNS = {
|
||||||
"han": re.compile(r"[\u3400-\u4dbf\u4e00-\u9fff]"),
|
"han": re.compile(r"[\u3400-\u4dbf\u4e00-\u9fff]"),
|
||||||
"latin": re.compile(r"[A-Za-z\u00c0-\u024f]"),
|
"latin": re.compile(r"[A-Za-z\u00c0-\u024f]"),
|
||||||
@@ -178,6 +197,75 @@ def _image_retrieval_question(question: str, image_contexts: list[dict]) -> str:
|
|||||||
return "\n".join(part for part in parts if part).strip()
|
return "\n".join(part for part in parts if part).strip()
|
||||||
|
|
||||||
|
|
||||||
|
def _answer_denies_available_image(answer: str) -> bool:
|
||||||
|
"""Reject only whole-image access denials, not uncertainty about one field."""
|
||||||
|
value = re.sub(r"\s+", " ", answer or "").strip()
|
||||||
|
return any(pattern.search(value) for pattern in _IMAGE_ACCESS_DENIAL_PATTERNS)
|
||||||
|
|
||||||
|
|
||||||
|
def _compact_context_text(value: Any, limit: int) -> str:
|
||||||
|
lines = [re.sub(r"\s+", " ", line).strip() for line in str(value or "").splitlines()]
|
||||||
|
text = "\n".join(line for line in lines if line).strip()
|
||||||
|
return text[:limit].rstrip()
|
||||||
|
|
||||||
|
|
||||||
|
def _grounded_image_fallback(question: str, image_contexts: list[dict]) -> str:
|
||||||
|
"""Build a safe answer from completed vision data when the chat model contradicts it."""
|
||||||
|
summaries: list[str] = []
|
||||||
|
facts: list[str] = []
|
||||||
|
excerpts: list[str] = []
|
||||||
|
warnings: list[str] = []
|
||||||
|
for context in image_contexts:
|
||||||
|
summary = _compact_context_text(context.get("summary"), 500)
|
||||||
|
if summary:
|
||||||
|
summaries.append(summary)
|
||||||
|
structured = context.get("structuredData") or {}
|
||||||
|
if isinstance(structured, dict):
|
||||||
|
for fact in structured.get("key_facts") or []:
|
||||||
|
value = _compact_context_text(fact, 300)
|
||||||
|
if value:
|
||||||
|
facts.append(value)
|
||||||
|
extracted = _compact_context_text(context.get("extractedText"), 900)
|
||||||
|
if extracted:
|
||||||
|
excerpts.append(extracted)
|
||||||
|
warning = _compact_context_text(context.get("warning"), 300)
|
||||||
|
if warning:
|
||||||
|
warnings.append(warning)
|
||||||
|
|
||||||
|
summaries = list(dict.fromkeys(summaries))
|
||||||
|
facts = list(dict.fromkeys(facts))[:6]
|
||||||
|
excerpts = list(dict.fromkeys(excerpts))
|
||||||
|
warnings = list(dict.fromkeys(warnings))
|
||||||
|
writing_system = _dominant_writing_system(question)
|
||||||
|
|
||||||
|
if writing_system == "latin":
|
||||||
|
parts = []
|
||||||
|
if summaries:
|
||||||
|
parts.append("From the image, I can confirm: " + " ".join(summaries))
|
||||||
|
if facts:
|
||||||
|
parts.append("Key details:\n" + "\n".join(
|
||||||
|
f"{index}. {fact}" for index, fact in enumerate(facts, 1)
|
||||||
|
))
|
||||||
|
elif excerpts:
|
||||||
|
parts.append("Visible text:\n" + excerpts[0])
|
||||||
|
if warnings:
|
||||||
|
parts.append("Please note: " + " ".join(warnings))
|
||||||
|
return "\n".join(parts).strip() or "The image is available, but there is not enough clear detail to confirm more."
|
||||||
|
|
||||||
|
parts = []
|
||||||
|
if summaries:
|
||||||
|
parts.append("从这张图中可以确认:" + ";".join(summaries).rstrip("。;") + "。")
|
||||||
|
if facts:
|
||||||
|
parts.append("其中比较明确的信息有:\n" + "\n".join(
|
||||||
|
f"{index}. {fact}" for index, fact in enumerate(facts, 1)
|
||||||
|
))
|
||||||
|
elif excerpts:
|
||||||
|
parts.append("图中可见的主要文字是:\n" + excerpts[0])
|
||||||
|
if warnings:
|
||||||
|
parts.append("需要注意:" + ";".join(warnings).rstrip("。;") + "。")
|
||||||
|
return "\n".join(parts).strip() or "这张图已经看到了,但目前能确认的清晰信息比较有限。"
|
||||||
|
|
||||||
|
|
||||||
def _run_billed_vision_call(
|
def _run_billed_vision_call(
|
||||||
db: Session,
|
db: Session,
|
||||||
avatar: Avatar,
|
avatar: Avatar,
|
||||||
@@ -236,11 +324,40 @@ async def _analyze_uploaded_image(
|
|||||||
max_bytes = max(1024, int(os.getenv("CHAT_IMAGE_MAX_BYTES", str(8 * 1024 * 1024))))
|
max_bytes = max(1024, int(os.getenv("CHAT_IMAGE_MAX_BYTES", str(8 * 1024 * 1024))))
|
||||||
content = await file.read(max_bytes + 1)
|
content = await file.read(max_bytes + 1)
|
||||||
filename = os.path.basename(file.filename or "图片")[:255]
|
filename = os.path.basename(file.filename or "图片")[:255]
|
||||||
|
try:
|
||||||
|
return _analyze_image_bytes(
|
||||||
|
db,
|
||||||
|
avatar,
|
||||||
|
content,
|
||||||
|
filename=filename,
|
||||||
|
mime_type=file.content_type or "",
|
||||||
|
uploader_kind=uploader_kind,
|
||||||
|
)
|
||||||
|
except ImageValidationError as exc:
|
||||||
|
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||||
|
except InsufficientTokensError:
|
||||||
|
raise
|
||||||
|
except RuntimeError as exc:
|
||||||
|
raise HTTPException(status_code=502, detail=str(exc)) from exc
|
||||||
|
finally:
|
||||||
|
content = b""
|
||||||
|
|
||||||
|
|
||||||
|
def _analyze_image_bytes(
|
||||||
|
db: Session,
|
||||||
|
avatar: Avatar,
|
||||||
|
content: bytes,
|
||||||
|
*,
|
||||||
|
filename: str,
|
||||||
|
mime_type: str,
|
||||||
|
uploader_kind: str,
|
||||||
|
) -> ChatAttachment:
|
||||||
|
"""Analyze image bytes from either HTTP upload or BOXIM without persisting raw data."""
|
||||||
attachment = ChatAttachment(
|
attachment = ChatAttachment(
|
||||||
avatar_id=avatar.id,
|
avatar_id=avatar.id,
|
||||||
uploader_kind=uploader_kind,
|
uploader_kind=uploader_kind,
|
||||||
filename=filename,
|
filename=filename,
|
||||||
mime_type=(file.content_type or "")[:100],
|
mime_type=(mime_type or "")[:100],
|
||||||
file_size=len(content),
|
file_size=len(content),
|
||||||
status="processing",
|
status="processing",
|
||||||
expires_at=_attachment_expiry(),
|
expires_at=_attachment_expiry(),
|
||||||
@@ -312,7 +429,7 @@ async def _analyze_uploaded_image(
|
|||||||
attachment.status = "failed"
|
attachment.status = "failed"
|
||||||
attachment.warning = str(exc)
|
attachment.warning = str(exc)
|
||||||
db.commit()
|
db.commit()
|
||||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
raise
|
||||||
except InsufficientTokensError:
|
except InsufficientTokensError:
|
||||||
attachment.status = "failed"
|
attachment.status = "failed"
|
||||||
attachment.warning = "积分余额不足"
|
attachment.warning = "积分余额不足"
|
||||||
@@ -328,9 +445,7 @@ async def _analyze_uploaded_image(
|
|||||||
avatar.id,
|
avatar.id,
|
||||||
type(exc).__name__,
|
type(exc).__name__,
|
||||||
)
|
)
|
||||||
raise HTTPException(status_code=502, detail=str(exc)) from exc
|
raise
|
||||||
finally:
|
|
||||||
content = b""
|
|
||||||
|
|
||||||
|
|
||||||
def _normalize_question(value: str) -> str:
|
def _normalize_question(value: str) -> str:
|
||||||
@@ -526,8 +641,10 @@ def _build_prompt(
|
|||||||
if image_contexts:
|
if image_contexts:
|
||||||
image_material = json.dumps(image_contexts, ensure_ascii=False, default=str)
|
image_material = json.dumps(image_contexts, ensure_ascii=False, default=str)
|
||||||
system += (
|
system += (
|
||||||
"\n以下是当前会话图片经过视觉识别后得到的资料:\n"
|
"\n当前会话图片已经成功读取并完成内容识别,以下资料就是可直接使用的图片内容:\n"
|
||||||
f"{image_material}"
|
f"{image_material}"
|
||||||
|
"\n必须直接依据这些图片内容回答当前问题。禁止声称无法查看、看不到、未收到、无法识别、"
|
||||||
|
"无法读取或不能访问图片,也不要要求对方重新上传;只有资料明确标记读取失败时才可以请对方重发。"
|
||||||
"\n图片资料可能包含 OCR 错字、模糊内容或用户尚未确认的信息,只能按可见内容谨慎表达。"
|
"\n图片资料可能包含 OCR 错字、模糊内容或用户尚未确认的信息,只能按可见内容谨慎表达。"
|
||||||
"标准答题对中的事实优先级高于图片资料,知识库事实优先级高于模型推测;发生冲突时遵循更高优先级资料,"
|
"标准答题对中的事实优先级高于图片资料,知识库事实优先级高于模型推测;发生冲突时遵循更高优先级资料,"
|
||||||
"并自然提醒对方核对原图。不得声称看到了图片中不存在的内容。"
|
"并自然提醒对方核对原图。不得声称看到了图片中不存在的内容。"
|
||||||
@@ -788,6 +905,14 @@ def _resolve_reply(
|
|||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
release_reservation(db, reservation, str(exc))
|
release_reservation(db, reservation, str(exc))
|
||||||
raise
|
raise
|
||||||
|
answer = str(answer or "").strip()
|
||||||
|
if image_contexts and _answer_denies_available_image(answer):
|
||||||
|
logger.warning(
|
||||||
|
"chat model contradicted ready image context avatar=%s source=%s",
|
||||||
|
avatar.id,
|
||||||
|
usage_source,
|
||||||
|
)
|
||||||
|
answer = _grounded_image_fallback(question, image_contexts)
|
||||||
result = {
|
result = {
|
||||||
"answer": answer,
|
"answer": answer,
|
||||||
"source": "qa" if matched else (
|
"source": "qa" if matched else (
|
||||||
|
|||||||
@@ -1,8 +1,8 @@
|
|||||||
import os
|
import os
|
||||||
import json
|
import json
|
||||||
import logging
|
import shutil
|
||||||
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
from datetime import datetime, timezone
|
|
||||||
|
|
||||||
from fastapi import APIRouter, UploadFile, File, Depends, Header, HTTPException
|
from fastapi import APIRouter, UploadFile, File, Depends, Header, HTTPException
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
@@ -12,16 +12,19 @@ from database import get_db
|
|||||||
from models import KnowledgeDoc, QAPair, KnowledgeChunk, Avatar, User
|
from models import KnowledgeDoc, QAPair, KnowledgeChunk, Avatar, User
|
||||||
from responses import ok, fail
|
from responses import ok, fail
|
||||||
import embeddings
|
import embeddings
|
||||||
|
from services.knowledge_vectorizer import knowledge_vectorizer
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
|
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||||
UPLOAD_DIR = os.path.abspath(os.getenv("UPLOAD_DIR", os.path.join(BASE_DIR, "uploads")))
|
UPLOAD_DIR = os.path.abspath(os.getenv("UPLOAD_DIR", os.path.join(BASE_DIR, "uploads")))
|
||||||
os.makedirs(UPLOAD_DIR, exist_ok=True)
|
os.makedirs(UPLOAD_DIR, exist_ok=True)
|
||||||
|
|
||||||
ALLOWED_EXT = {".md", ".txt", ".pdf", ".doc", ".docx", ".xlsx"}
|
ALLOWED_EXT = {".md", ".txt", ".pdf", ".doc", ".docx", ".xlsx"}
|
||||||
MAX_UPLOAD_BYTES = 10 * 1024 * 1024
|
MAX_UPLOAD_BYTES = 50 * 1024 * 1024
|
||||||
|
UPLOAD_CHUNK_BYTES = 1024 * 1024
|
||||||
|
MULTIPART_CHUNK_BYTES = 5 * 1024 * 1024
|
||||||
|
MULTIPART_ROOT = ".multipart"
|
||||||
|
MULTIPART_TTL_SECONDS = 24 * 60 * 60
|
||||||
|
|
||||||
|
|
||||||
class QAIn(BaseModel):
|
class QAIn(BaseModel):
|
||||||
@@ -34,6 +37,74 @@ class EnabledIn(BaseModel):
|
|||||||
enabled: bool = True
|
enabled: bool = True
|
||||||
|
|
||||||
|
|
||||||
|
class MultipartUploadIn(BaseModel):
|
||||||
|
filename: str
|
||||||
|
fileSize: int
|
||||||
|
totalChunks: int
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_document(filename: str, file_size: int):
|
||||||
|
ext = os.path.splitext(filename or "")[1].lower()
|
||||||
|
if ext not in ALLOWED_EXT:
|
||||||
|
return None, f"不支持的文件类型:{ext or '空'},仅支持 md/txt/pdf/doc/docx/xlsx"
|
||||||
|
if file_size <= 0:
|
||||||
|
return None, "文件内容不能为空"
|
||||||
|
if file_size > MAX_UPLOAD_BYTES:
|
||||||
|
return None, "文件不能超过 50MB"
|
||||||
|
return ext, ""
|
||||||
|
|
||||||
|
|
||||||
|
def _multipart_dir(avatar_id: str, upload_id: str) -> str:
|
||||||
|
safe_avatar_id = os.path.basename(avatar_id)
|
||||||
|
safe_upload_id = os.path.basename(upload_id)
|
||||||
|
if (
|
||||||
|
safe_avatar_id != avatar_id
|
||||||
|
or safe_upload_id != upload_id
|
||||||
|
or len(upload_id) != 32
|
||||||
|
or any(character not in "0123456789abcdef" for character in upload_id)
|
||||||
|
):
|
||||||
|
raise HTTPException(status_code=400, detail="上传标识无效")
|
||||||
|
return os.path.join(UPLOAD_DIR, MULTIPART_ROOT, safe_avatar_id, safe_upload_id)
|
||||||
|
|
||||||
|
|
||||||
|
def _purge_stale_multipart_uploads(avatar_id: str):
|
||||||
|
avatar_upload_root = os.path.join(UPLOAD_DIR, MULTIPART_ROOT, os.path.basename(avatar_id))
|
||||||
|
if not os.path.isdir(avatar_upload_root):
|
||||||
|
return
|
||||||
|
cutoff = time.time() - MULTIPART_TTL_SECONDS
|
||||||
|
for entry in os.scandir(avatar_upload_root):
|
||||||
|
if entry.is_dir(follow_symlinks=False) and entry.stat(follow_symlinks=False).st_mtime < cutoff:
|
||||||
|
shutil.rmtree(entry.path, ignore_errors=True)
|
||||||
|
|
||||||
|
|
||||||
|
def _read_multipart_metadata(avatar_id: str, upload_id: str) -> tuple[str, dict]:
|
||||||
|
upload_dir = _multipart_dir(avatar_id, upload_id)
|
||||||
|
metadata_path = os.path.join(upload_dir, "metadata.json")
|
||||||
|
if not os.path.isfile(metadata_path):
|
||||||
|
raise HTTPException(status_code=404, detail="上传任务不存在或已过期")
|
||||||
|
with open(metadata_path, "r", encoding="utf-8") as stream:
|
||||||
|
return upload_dir, json.load(stream)
|
||||||
|
|
||||||
|
|
||||||
|
def _create_knowledge_doc(db: Session, avatar_id: str, filename: str, ext: str, file_size: int, stored: str):
|
||||||
|
doc = KnowledgeDoc(
|
||||||
|
id=uuid.uuid4().hex,
|
||||||
|
avatar_id=avatar_id,
|
||||||
|
filename=filename,
|
||||||
|
file_type=ext.lstrip("."),
|
||||||
|
file_size=file_size,
|
||||||
|
file_url=f"/api/files/{avatar_id}/{stored}",
|
||||||
|
status="parsing",
|
||||||
|
index_stage="queued",
|
||||||
|
index_progress=0,
|
||||||
|
)
|
||||||
|
db.add(doc)
|
||||||
|
db.commit()
|
||||||
|
db.refresh(doc)
|
||||||
|
knowledge_vectorizer.enqueue(doc.id)
|
||||||
|
return doc
|
||||||
|
|
||||||
|
|
||||||
def _doc_payload(doc: KnowledgeDoc) -> dict:
|
def _doc_payload(doc: KnowledgeDoc) -> dict:
|
||||||
payload = doc.to_dict()
|
payload = doc.to_dict()
|
||||||
stored_name = os.path.basename(doc.file_url or "")
|
stored_name = os.path.basename(doc.file_url or "")
|
||||||
@@ -71,84 +142,178 @@ def list_docs(avatar_id: str, authorization: str = Header(None), db: Session = D
|
|||||||
.order_by(KnowledgeDoc.created_at.desc())
|
.order_by(KnowledgeDoc.created_at.desc())
|
||||||
.all()
|
.all()
|
||||||
)
|
)
|
||||||
# Older synchronous uploads could be interrupted after persisting "parsing".
|
|
||||||
# New uploads are committed only after indexing finishes, so these rows are stale.
|
|
||||||
stale_docs = [doc for doc in docs if doc.status == "parsing"]
|
|
||||||
if stale_docs:
|
|
||||||
for doc in stale_docs:
|
|
||||||
doc.status = "failed"
|
|
||||||
doc.vectorized = False
|
|
||||||
doc.chunk_count = 0
|
|
||||||
db.commit()
|
|
||||||
return ok([_doc_payload(d) for d in docs])
|
return ok([_doc_payload(d) for d in docs])
|
||||||
|
|
||||||
|
|
||||||
@router.post("/avatar/{avatar_id}/knowledge/docs")
|
@router.post("/avatar/{avatar_id}/knowledge/docs")
|
||||||
async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization: str = Header(None), db: Session = Depends(get_db)):
|
async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization: str = Header(None), db: Session = Depends(get_db)):
|
||||||
_require_owned_avatar(db, avatar_id, authorization)
|
_require_owned_avatar(db, avatar_id, authorization)
|
||||||
ext = os.path.splitext(file.filename or "")[1].lower()
|
ext, validation_error = _validate_document(file.filename or "", 1)
|
||||||
if ext not in ALLOWED_EXT:
|
if validation_error:
|
||||||
return fail(f"不支持的文件类型:{ext or '空'},仅支持 md/txt/pdf/doc/docx/xlsx", code=400)
|
return fail(validation_error, code=400)
|
||||||
avatar_dir = os.path.join(UPLOAD_DIR, avatar_id)
|
avatar_dir = os.path.join(UPLOAD_DIR, avatar_id)
|
||||||
os.makedirs(avatar_dir, exist_ok=True)
|
os.makedirs(avatar_dir, exist_ok=True)
|
||||||
stored = f"{uuid.uuid4().hex}{ext}"
|
stored = f"{uuid.uuid4().hex}{ext}"
|
||||||
path = os.path.join(avatar_dir, stored)
|
path = os.path.join(avatar_dir, stored)
|
||||||
content = await file.read()
|
file_size = 0
|
||||||
if len(content) > MAX_UPLOAD_BYTES:
|
|
||||||
return fail("文件不能超过 10MB", code=400)
|
|
||||||
with open(path, "wb") as f:
|
|
||||||
f.write(content)
|
|
||||||
doc = KnowledgeDoc(
|
|
||||||
id=uuid.uuid4().hex,
|
|
||||||
avatar_id=avatar_id,
|
|
||||||
filename=file.filename,
|
|
||||||
file_type=ext.lstrip("."),
|
|
||||||
file_size=len(content),
|
|
||||||
file_url=f"/api/files/{avatar_id}/{stored}",
|
|
||||||
status="parsing",
|
|
||||||
)
|
|
||||||
|
|
||||||
# Complete extraction and embedding before the first database commit so a
|
|
||||||
# process restart cannot leave a permanent "parsing" row behind.
|
|
||||||
try:
|
try:
|
||||||
text = embeddings.extract_text(path, ext)
|
# Stream large files to disk so a 100MB upload does not occupy 100MB RAM.
|
||||||
chunks = embeddings.chunk_text(text)
|
with open(path, "wb") as f:
|
||||||
if not chunks:
|
while chunk := await file.read(UPLOAD_CHUNK_BYTES):
|
||||||
raise ValueError("文档没有可建立索引的文字内容")
|
file_size += len(chunk)
|
||||||
vectors = embeddings.embed(chunks)
|
if file_size > MAX_UPLOAD_BYTES:
|
||||||
if len(vectors) != len(chunks):
|
raise ValueError("文件不能超过 50MB")
|
||||||
raise ValueError("向量服务返回数量与文档分段不一致")
|
f.write(chunk)
|
||||||
doc.vectorized = True
|
except ValueError as exc:
|
||||||
doc.embedding_model = embeddings.MODEL
|
if os.path.exists(path):
|
||||||
doc.chunk_count = len(chunks)
|
os.remove(path)
|
||||||
doc.vectorized_at = datetime.now(timezone.utc)
|
return fail(str(exc), code=400)
|
||||||
doc.status = "ready"
|
if file_size == 0:
|
||||||
db.add(doc)
|
if os.path.exists(path):
|
||||||
for i, (chunk, vector) in enumerate(zip(chunks, vectors)):
|
os.remove(path)
|
||||||
db.add(
|
return fail("文件内容不能为空", code=400)
|
||||||
KnowledgeChunk(
|
|
||||||
doc_id=doc.id,
|
|
||||||
avatar_id=avatar_id,
|
|
||||||
content=chunk,
|
|
||||||
vector=json.dumps(vector),
|
|
||||||
chunk_index=i,
|
|
||||||
embedding_model=embeddings.MODEL,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
db.commit()
|
|
||||||
db.refresh(doc)
|
|
||||||
except Exception as exc:
|
|
||||||
db.rollback()
|
|
||||||
doc.status = "failed"
|
|
||||||
doc.vectorized = False
|
|
||||||
doc.embedding_model = ""
|
|
||||||
doc.chunk_count = 0
|
|
||||||
doc.vectorized_at = None
|
|
||||||
db.add(doc)
|
|
||||||
db.commit()
|
|
||||||
db.refresh(doc)
|
|
||||||
logger.exception("knowledge vectorization failed for %s: %s", doc.id, exc)
|
|
||||||
|
|
||||||
|
doc = _create_knowledge_doc(db, avatar_id, file.filename or stored, ext, file_size, stored)
|
||||||
|
return ok(_doc_payload(doc))
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/avatar/{avatar_id}/knowledge/uploads")
|
||||||
|
def create_multipart_upload(
|
||||||
|
avatar_id: str,
|
||||||
|
body: MultipartUploadIn,
|
||||||
|
authorization: str = Header(None),
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
):
|
||||||
|
_require_owned_avatar(db, avatar_id, authorization)
|
||||||
|
ext, validation_error = _validate_document(body.filename, body.fileSize)
|
||||||
|
if validation_error:
|
||||||
|
return fail(validation_error, code=400)
|
||||||
|
expected_chunks = (body.fileSize + MULTIPART_CHUNK_BYTES - 1) // MULTIPART_CHUNK_BYTES
|
||||||
|
if body.totalChunks != expected_chunks:
|
||||||
|
return fail("文件分片数量不正确", code=400)
|
||||||
|
|
||||||
|
_purge_stale_multipart_uploads(avatar_id)
|
||||||
|
upload_id = uuid.uuid4().hex
|
||||||
|
upload_dir = _multipart_dir(avatar_id, upload_id)
|
||||||
|
os.makedirs(upload_dir, exist_ok=False)
|
||||||
|
metadata = {
|
||||||
|
"filename": body.filename,
|
||||||
|
"fileSize": body.fileSize,
|
||||||
|
"totalChunks": body.totalChunks,
|
||||||
|
"extension": ext,
|
||||||
|
}
|
||||||
|
with open(os.path.join(upload_dir, "metadata.json"), "w", encoding="utf-8") as stream:
|
||||||
|
json.dump(metadata, stream, ensure_ascii=False)
|
||||||
|
return ok({"uploadId": upload_id, "chunkSize": MULTIPART_CHUNK_BYTES})
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/avatar/{avatar_id}/knowledge/uploads/{upload_id}/chunks/{chunk_index}")
|
||||||
|
async def upload_multipart_chunk(
|
||||||
|
avatar_id: str,
|
||||||
|
upload_id: str,
|
||||||
|
chunk_index: int,
|
||||||
|
file: UploadFile = File(...),
|
||||||
|
authorization: str = Header(None),
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
):
|
||||||
|
_require_owned_avatar(db, avatar_id, authorization)
|
||||||
|
upload_dir, metadata = _read_multipart_metadata(avatar_id, upload_id)
|
||||||
|
total_chunks = int(metadata["totalChunks"])
|
||||||
|
if chunk_index < 0 or chunk_index >= total_chunks:
|
||||||
|
return fail("文件分片序号不正确", code=400)
|
||||||
|
|
||||||
|
expected_size = min(
|
||||||
|
MULTIPART_CHUNK_BYTES,
|
||||||
|
int(metadata["fileSize"]) - chunk_index * MULTIPART_CHUNK_BYTES,
|
||||||
|
)
|
||||||
|
part_path = os.path.join(upload_dir, f"{chunk_index}.part")
|
||||||
|
temporary_path = f"{part_path}.uploading"
|
||||||
|
received = 0
|
||||||
|
try:
|
||||||
|
with open(temporary_path, "wb") as stream:
|
||||||
|
while chunk := await file.read(UPLOAD_CHUNK_BYTES):
|
||||||
|
received += len(chunk)
|
||||||
|
if received > expected_size:
|
||||||
|
raise ValueError("文件分片大小不正确")
|
||||||
|
stream.write(chunk)
|
||||||
|
if received != expected_size:
|
||||||
|
raise ValueError("文件分片大小不正确")
|
||||||
|
os.replace(temporary_path, part_path)
|
||||||
|
except ValueError as exc:
|
||||||
|
if os.path.exists(temporary_path):
|
||||||
|
os.remove(temporary_path)
|
||||||
|
return fail(str(exc), code=400)
|
||||||
|
return ok({"chunkIndex": chunk_index, "uploadedBytes": received})
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/avatar/{avatar_id}/knowledge/uploads/{upload_id}/complete")
|
||||||
|
def complete_multipart_upload(
|
||||||
|
avatar_id: str,
|
||||||
|
upload_id: str,
|
||||||
|
authorization: str = Header(None),
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
):
|
||||||
|
_require_owned_avatar(db, avatar_id, authorization)
|
||||||
|
upload_dir, metadata = _read_multipart_metadata(avatar_id, upload_id)
|
||||||
|
total_chunks = int(metadata["totalChunks"])
|
||||||
|
part_paths = [os.path.join(upload_dir, f"{index}.part") for index in range(total_chunks)]
|
||||||
|
if not all(os.path.isfile(path) for path in part_paths):
|
||||||
|
return fail("文件分片尚未上传完整", code=400)
|
||||||
|
if sum(os.path.getsize(path) for path in part_paths) != int(metadata["fileSize"]):
|
||||||
|
return fail("文件分片总大小不正确", code=400)
|
||||||
|
|
||||||
|
avatar_dir = os.path.join(UPLOAD_DIR, avatar_id)
|
||||||
|
os.makedirs(avatar_dir, exist_ok=True)
|
||||||
|
stored = f"{uuid.uuid4().hex}{metadata['extension']}"
|
||||||
|
final_path = os.path.join(avatar_dir, stored)
|
||||||
|
temporary_path = f"{final_path}.assembling"
|
||||||
|
try:
|
||||||
|
with open(temporary_path, "wb") as output:
|
||||||
|
for part_path in part_paths:
|
||||||
|
with open(part_path, "rb") as source:
|
||||||
|
shutil.copyfileobj(source, output, UPLOAD_CHUNK_BYTES)
|
||||||
|
os.replace(temporary_path, final_path)
|
||||||
|
doc = _create_knowledge_doc(
|
||||||
|
db,
|
||||||
|
avatar_id,
|
||||||
|
metadata["filename"],
|
||||||
|
metadata["extension"],
|
||||||
|
int(metadata["fileSize"]),
|
||||||
|
stored,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
if os.path.exists(temporary_path):
|
||||||
|
os.remove(temporary_path)
|
||||||
|
raise
|
||||||
|
shutil.rmtree(upload_dir, ignore_errors=True)
|
||||||
|
return ok(_doc_payload(doc))
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/avatar/{avatar_id}/knowledge/docs/{doc_id}/retry")
|
||||||
|
def retry_doc(avatar_id: str, doc_id: str, authorization: str = Header(None), db: Session = Depends(get_db)):
|
||||||
|
_require_owned_avatar(db, avatar_id, authorization)
|
||||||
|
doc = db.query(KnowledgeDoc).filter(
|
||||||
|
KnowledgeDoc.id == doc_id, KnowledgeDoc.avatar_id == avatar_id
|
||||||
|
).first()
|
||||||
|
if not doc:
|
||||||
|
return fail("文档不存在", code=404)
|
||||||
|
if doc.vectorized and doc.status == "ready":
|
||||||
|
return ok(_doc_payload(doc))
|
||||||
|
stored_name = os.path.basename(doc.file_url or "")
|
||||||
|
if not stored_name or not os.path.isfile(os.path.join(UPLOAD_DIR, avatar_id, stored_name)):
|
||||||
|
return fail("原文件不可用,请重新上传", code=400)
|
||||||
|
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == doc.id).delete()
|
||||||
|
doc.status = "parsing"
|
||||||
|
doc.vectorized = False
|
||||||
|
doc.embedding_model = ""
|
||||||
|
doc.chunk_count = 0
|
||||||
|
doc.vectorized_at = None
|
||||||
|
doc.error_message = ""
|
||||||
|
doc.index_stage = "queued"
|
||||||
|
doc.index_progress = 0
|
||||||
|
db.commit()
|
||||||
|
db.refresh(doc)
|
||||||
|
knowledge_vectorizer.enqueue(doc.id)
|
||||||
return ok(_doc_payload(doc))
|
return ok(_doc_payload(doc))
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,151 @@
|
|||||||
|
"""Parse and safely download image payloads from BOXIM private messages."""
|
||||||
|
|
||||||
|
import ipaddress
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import socket
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import PurePosixPath
|
||||||
|
from urllib.parse import unquote, urljoin, urlsplit
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
|
||||||
|
MAX_REDIRECTS = 3
|
||||||
|
|
||||||
|
|
||||||
|
class BoxIMImageError(RuntimeError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class DownloadedBoxIMImage:
|
||||||
|
content: bytes
|
||||||
|
filename: str
|
||||||
|
mime_type: str
|
||||||
|
source_url: str
|
||||||
|
|
||||||
|
|
||||||
|
def parse_boxim_image_url(content: str, *, base_url: str = "") -> str:
|
||||||
|
try:
|
||||||
|
payload = json.loads(content or "")
|
||||||
|
except (TypeError, ValueError) as exc:
|
||||||
|
raise BoxIMImageError("BOXIM 图片消息格式无效") from exc
|
||||||
|
if not isinstance(payload, dict):
|
||||||
|
raise BoxIMImageError("BOXIM 图片消息格式无效")
|
||||||
|
|
||||||
|
value = payload.get("originUrl") or payload.get("thumbUrl") or payload.get("url")
|
||||||
|
if not isinstance(value, str) or not value.strip():
|
||||||
|
raise BoxIMImageError("BOXIM 图片消息缺少图片地址")
|
||||||
|
value = value.strip()
|
||||||
|
if value.startswith("/"):
|
||||||
|
if not base_url:
|
||||||
|
raise BoxIMImageError("BOXIM 图片地址不完整")
|
||||||
|
value = urljoin(f"{base_url.rstrip('/')}/", value)
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _configured_hosts(name: str) -> set[str]:
|
||||||
|
return {
|
||||||
|
value.strip().lower().rstrip(".")
|
||||||
|
for value in os.getenv(name, "").split(",")
|
||||||
|
if value.strip()
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _host_matches(host: str, configured: set[str]) -> bool:
|
||||||
|
return any(host == value or host.endswith(f".{value}") for value in configured)
|
||||||
|
|
||||||
|
|
||||||
|
def _resolved_addresses(host: str, port: int) -> set[ipaddress.IPv4Address | ipaddress.IPv6Address]:
|
||||||
|
try:
|
||||||
|
return {
|
||||||
|
ipaddress.ip_address(item[4][0])
|
||||||
|
for item in socket.getaddrinfo(host, port, type=socket.SOCK_STREAM)
|
||||||
|
}
|
||||||
|
except (OSError, ValueError) as exc:
|
||||||
|
raise BoxIMImageError("BOXIM 图片地址无法解析") from exc
|
||||||
|
|
||||||
|
|
||||||
|
def _is_safe_remote_url(url: str) -> None:
|
||||||
|
parsed = urlsplit(url)
|
||||||
|
scheme = parsed.scheme.lower()
|
||||||
|
allow_http = os.getenv("BOXIM_IMAGE_ALLOW_HTTP", "").lower() in {"1", "true", "yes"}
|
||||||
|
if scheme not in ({"https", "http"} if allow_http else {"https"}):
|
||||||
|
raise BoxIMImageError("BOXIM 图片地址必须使用 HTTPS")
|
||||||
|
if parsed.username or parsed.password or not parsed.hostname:
|
||||||
|
raise BoxIMImageError("BOXIM 图片地址无效")
|
||||||
|
|
||||||
|
host = parsed.hostname.lower().rstrip(".")
|
||||||
|
allowed_hosts = _configured_hosts("BOXIM_IMAGE_ALLOWED_HOSTS")
|
||||||
|
if allowed_hosts and not _host_matches(host, allowed_hosts):
|
||||||
|
raise BoxIMImageError("BOXIM 图片地址不在允许的域名范围内")
|
||||||
|
|
||||||
|
private_hosts = _configured_hosts("BOXIM_IMAGE_PRIVATE_HOSTS")
|
||||||
|
try:
|
||||||
|
addresses = {ipaddress.ip_address(host)}
|
||||||
|
except ValueError:
|
||||||
|
addresses = _resolved_addresses(host, parsed.port or (443 if scheme == "https" else 80))
|
||||||
|
if not addresses:
|
||||||
|
raise BoxIMImageError("BOXIM 图片地址无法解析")
|
||||||
|
if _host_matches(host, private_hosts):
|
||||||
|
return
|
||||||
|
if any(not address.is_global for address in addresses):
|
||||||
|
raise BoxIMImageError("BOXIM 图片地址指向受限网络")
|
||||||
|
|
||||||
|
|
||||||
|
def _filename_from_url(url: str) -> str:
|
||||||
|
value = unquote(PurePosixPath(urlsplit(url).path).name).strip()
|
||||||
|
value = value.replace("\x00", "")
|
||||||
|
return (value or "boxim-image")[:255]
|
||||||
|
|
||||||
|
|
||||||
|
def download_boxim_image(
|
||||||
|
content: str,
|
||||||
|
*,
|
||||||
|
base_url: str = "",
|
||||||
|
transport: httpx.BaseTransport | None = None,
|
||||||
|
) -> DownloadedBoxIMImage:
|
||||||
|
"""Download one BOXIM image without redirects or oversized responses escaping checks."""
|
||||||
|
url = parse_boxim_image_url(content, base_url=base_url)
|
||||||
|
max_bytes = max(1024, int(os.getenv("CHAT_IMAGE_MAX_BYTES", str(8 * 1024 * 1024))))
|
||||||
|
timeout = max(1.0, min(float(os.getenv("BOXIM_IMAGE_TIMEOUT_SECONDS", "15")), 60.0))
|
||||||
|
|
||||||
|
with httpx.Client(
|
||||||
|
timeout=timeout,
|
||||||
|
follow_redirects=False,
|
||||||
|
trust_env=False,
|
||||||
|
transport=transport,
|
||||||
|
) as client:
|
||||||
|
for _ in range(MAX_REDIRECTS + 1):
|
||||||
|
_is_safe_remote_url(url)
|
||||||
|
try:
|
||||||
|
with client.stream("GET", url, headers={"Accept": "image/*"}) as response:
|
||||||
|
if response.status_code in {301, 302, 303, 307, 308}:
|
||||||
|
location = response.headers.get("location", "").strip()
|
||||||
|
if not location:
|
||||||
|
raise BoxIMImageError("BOXIM 图片跳转地址无效")
|
||||||
|
url = urljoin(url, location)
|
||||||
|
continue
|
||||||
|
response.raise_for_status()
|
||||||
|
raw_length = response.headers.get("content-length", "")
|
||||||
|
if raw_length.isdigit() and int(raw_length) > max_bytes:
|
||||||
|
raise BoxIMImageError("BOXIM 图片超过大小限制")
|
||||||
|
chunks = bytearray()
|
||||||
|
for chunk in response.iter_bytes():
|
||||||
|
chunks.extend(chunk)
|
||||||
|
if len(chunks) > max_bytes:
|
||||||
|
raise BoxIMImageError("BOXIM 图片超过大小限制")
|
||||||
|
if not chunks:
|
||||||
|
raise BoxIMImageError("BOXIM 图片内容为空")
|
||||||
|
return DownloadedBoxIMImage(
|
||||||
|
content=bytes(chunks),
|
||||||
|
filename=_filename_from_url(url),
|
||||||
|
mime_type=response.headers.get("content-type", "").split(";", 1)[0][:100],
|
||||||
|
source_url=url,
|
||||||
|
)
|
||||||
|
except BoxIMImageError:
|
||||||
|
raise
|
||||||
|
except (httpx.HTTPError, OSError) as exc:
|
||||||
|
raise BoxIMImageError("BOXIM 图片下载失败") from exc
|
||||||
|
raise BoxIMImageError("BOXIM 图片跳转次数过多")
|
||||||
@@ -0,0 +1,143 @@
|
|||||||
|
"""Durable, serial knowledge-document indexing for the avatar knowledge base."""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import queue
|
||||||
|
import threading
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
|
from database import SessionLocal
|
||||||
|
from models import KnowledgeChunk, KnowledgeDoc
|
||||||
|
import embeddings
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
BACKEND_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||||
|
UPLOAD_DIR = os.path.abspath(
|
||||||
|
os.getenv("UPLOAD_DIR", os.path.join(BACKEND_DIR, "routers", "uploads"))
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class KnowledgeVectorizer:
|
||||||
|
"""Indexes one document at a time so slow providers cannot block uploads."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._queue: queue.Queue[str] = queue.Queue()
|
||||||
|
self._queued: set[str] = set()
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
self._thread: threading.Thread | None = None
|
||||||
|
|
||||||
|
def start(self):
|
||||||
|
if self._thread and self._thread.is_alive():
|
||||||
|
return
|
||||||
|
self._thread = threading.Thread(
|
||||||
|
target=self._run, name="knowledge-vectorizer", daemon=True
|
||||||
|
)
|
||||||
|
self._thread.start()
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
# A process restart must not abandon documents already accepted by upload.
|
||||||
|
for (doc_id,) in db.query(KnowledgeDoc.id).filter(KnowledgeDoc.status == "parsing"):
|
||||||
|
self.enqueue(doc_id)
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
def enqueue(self, doc_id: str):
|
||||||
|
with self._lock:
|
||||||
|
if doc_id in self._queued:
|
||||||
|
return
|
||||||
|
self._queued.add(doc_id)
|
||||||
|
self._queue.put(doc_id)
|
||||||
|
|
||||||
|
def _run(self):
|
||||||
|
while True:
|
||||||
|
doc_id = self._queue.get()
|
||||||
|
try:
|
||||||
|
self.vectorize_document(doc_id)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Unexpected knowledge vectorizer failure for %s", doc_id)
|
||||||
|
finally:
|
||||||
|
with self._lock:
|
||||||
|
self._queued.discard(doc_id)
|
||||||
|
self._queue.task_done()
|
||||||
|
|
||||||
|
def vectorize_document(self, doc_id: str):
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
doc = db.get(KnowledgeDoc, doc_id)
|
||||||
|
if not doc or doc.status != "parsing":
|
||||||
|
return
|
||||||
|
|
||||||
|
stored_name = os.path.basename(doc.file_url or "")
|
||||||
|
path = os.path.join(UPLOAD_DIR, doc.avatar_id, stored_name)
|
||||||
|
if not stored_name or not os.path.isfile(path):
|
||||||
|
raise FileNotFoundError("原文件不可用,请重新上传")
|
||||||
|
|
||||||
|
self._set_progress(db, doc, "extracting", 8)
|
||||||
|
text = embeddings.extract_text(path, f".{doc.file_type}")
|
||||||
|
self._set_progress(db, doc, "chunking", 22)
|
||||||
|
chunks = embeddings.chunk_text(text)
|
||||||
|
if not chunks:
|
||||||
|
raise ValueError("文档没有可建立索引的文字内容")
|
||||||
|
self._set_progress(db, doc, "embedding", 30)
|
||||||
|
|
||||||
|
def embedding_progress(done: int, total: int):
|
||||||
|
percent = 30 + int((done / max(1, total)) * 65)
|
||||||
|
self._set_progress(db, doc, "embedding", min(percent, 95))
|
||||||
|
|
||||||
|
vectors = embeddings.embed(chunks, on_progress=embedding_progress)
|
||||||
|
if len(vectors) != len(chunks):
|
||||||
|
raise ValueError("向量服务返回数量与文档分段不一致")
|
||||||
|
|
||||||
|
# Commit the document and every chunk together. Chat only sees complete indexes.
|
||||||
|
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == doc.id).delete()
|
||||||
|
db.add_all(
|
||||||
|
[
|
||||||
|
KnowledgeChunk(
|
||||||
|
doc_id=doc.id,
|
||||||
|
avatar_id=doc.avatar_id,
|
||||||
|
content=chunk,
|
||||||
|
vector=json.dumps(vector),
|
||||||
|
chunk_index=index,
|
||||||
|
embedding_model=embeddings.MODEL,
|
||||||
|
)
|
||||||
|
for index, (chunk, vector) in enumerate(zip(chunks, vectors))
|
||||||
|
]
|
||||||
|
)
|
||||||
|
doc.vectorized = True
|
||||||
|
doc.embedding_model = embeddings.MODEL
|
||||||
|
doc.chunk_count = len(chunks)
|
||||||
|
doc.vectorized_at = datetime.now(timezone.utc)
|
||||||
|
doc.status = "ready"
|
||||||
|
doc.error_message = ""
|
||||||
|
doc.index_stage = "ready"
|
||||||
|
doc.index_progress = 100
|
||||||
|
db.commit()
|
||||||
|
logger.info("Knowledge document %s indexed with %s chunks", doc.id, len(chunks))
|
||||||
|
except Exception as exc:
|
||||||
|
db.rollback()
|
||||||
|
failed_doc = db.get(KnowledgeDoc, doc_id)
|
||||||
|
if failed_doc:
|
||||||
|
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == failed_doc.id).delete()
|
||||||
|
failed_doc.status = "failed"
|
||||||
|
failed_doc.vectorized = False
|
||||||
|
failed_doc.embedding_model = ""
|
||||||
|
failed_doc.chunk_count = 0
|
||||||
|
failed_doc.vectorized_at = None
|
||||||
|
failed_doc.error_message = str(exc)[:300] or "建立知识索引失败"
|
||||||
|
failed_doc.index_stage = "failed"
|
||||||
|
failed_doc.index_progress = 0
|
||||||
|
db.commit()
|
||||||
|
logger.exception("Knowledge vectorization failed for %s: %s", doc_id, exc)
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _set_progress(db, doc, stage: str, progress: int):
|
||||||
|
doc.index_stage = stage
|
||||||
|
doc.index_progress = progress
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
|
||||||
|
knowledge_vectorizer = KnowledgeVectorizer()
|
||||||
@@ -3,6 +3,7 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import hashlib
|
import hashlib
|
||||||
import logging
|
import logging
|
||||||
|
import os
|
||||||
import re
|
import re
|
||||||
import secrets
|
import secrets
|
||||||
import time
|
import time
|
||||||
@@ -13,12 +14,19 @@ from sqlalchemy.orm import Session
|
|||||||
|
|
||||||
from models import (
|
from models import (
|
||||||
Avatar,
|
Avatar,
|
||||||
|
ChatAttachment,
|
||||||
TakeoverCursor,
|
TakeoverCursor,
|
||||||
TakeoverMessage,
|
TakeoverMessage,
|
||||||
TakeoverReplyTask,
|
TakeoverReplyTask,
|
||||||
User,
|
User,
|
||||||
)
|
)
|
||||||
from services.boxim_client import BoxIMClient, BoxIMError
|
from services.boxim_client import BoxIMClient, BoxIMError
|
||||||
|
from services.boxim_image_service import (
|
||||||
|
BoxIMImageError,
|
||||||
|
download_boxim_image,
|
||||||
|
parse_boxim_image_url,
|
||||||
|
)
|
||||||
|
from services.vision_service import ImageValidationError
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -37,6 +45,19 @@ HUMAN_PAUSE_SECONDS = 600
|
|||||||
RATE_LIMIT_WINDOW_SECONDS = 300
|
RATE_LIMIT_WINDOW_SECONDS = 300
|
||||||
RATE_LIMIT_MAX_REPLIES = 5
|
RATE_LIMIT_MAX_REPLIES = 5
|
||||||
AVATAR_LOCAL_ID_PREFIX = "880"
|
AVATAR_LOCAL_ID_PREFIX = "880"
|
||||||
|
BOXIM_TEXT_MESSAGE_TYPE = 0
|
||||||
|
BOXIM_IMAGE_MESSAGE_TYPE = 1
|
||||||
|
BOXIM_IMAGE_PROMPT = "请看看这张图片。"
|
||||||
|
BOXIM_IMAGE_UNAVAILABLE_REPLY = "这张图片我暂时没看清,麻烦重新发送一张清晰的原图。"
|
||||||
|
IMAGE_CONTEXT_LOOKBACK_SECONDS = 1800
|
||||||
|
IMAGE_REFERENCE_LOOKBACK_SECONDS = 172_800
|
||||||
|
MAX_RECENT_IMAGE_CONTEXTS = 3
|
||||||
|
|
||||||
|
_IMAGE_REFERENCE_PATTERN = re.compile(
|
||||||
|
r"(?:图片|图像|照片|截图|这张图|刚才.{0,8}图|病例|病历|检查单|检验单|化验单|报告|影像|"
|
||||||
|
r"\b(?:image|photo|picture|screenshot|scan|report)\b)",
|
||||||
|
re.IGNORECASE,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _utcnow() -> datetime:
|
def _utcnow() -> datetime:
|
||||||
@@ -107,6 +128,18 @@ def _configured_reply_delay(avatar: Avatar, fallback: int | None = None) -> int:
|
|||||||
return delay
|
return delay
|
||||||
|
|
||||||
|
|
||||||
|
def _event_prompt(event: TakeoverMessage) -> str:
|
||||||
|
if event.message_type == BOXIM_TEXT_MESSAGE_TYPE:
|
||||||
|
return event.content.strip()
|
||||||
|
if event.message_type == BOXIM_IMAGE_MESSAGE_TYPE:
|
||||||
|
return BOXIM_IMAGE_PROMPT
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
def _references_recent_image(value: str) -> bool:
|
||||||
|
return bool(_IMAGE_REFERENCE_PATTERN.search(value or ""))
|
||||||
|
|
||||||
|
|
||||||
class TakeoverService:
|
class TakeoverService:
|
||||||
"""Poll BOXIM, honor the owner grace period, then generate and send one reply."""
|
"""Poll BOXIM, honor the owner grace period, then generate and send one reply."""
|
||||||
|
|
||||||
@@ -129,6 +162,7 @@ class TakeoverService:
|
|||||||
self._sessions: dict[str, dict] = {}
|
self._sessions: dict[str, dict] = {}
|
||||||
self._poll_lock = asyncio.Lock()
|
self._poll_lock = asyncio.Lock()
|
||||||
self._process_lock = asyncio.Lock()
|
self._process_lock = asyncio.Lock()
|
||||||
|
self._persist_lock = asyncio.Lock()
|
||||||
|
|
||||||
async def poll_and_process_messages(self):
|
async def poll_and_process_messages(self):
|
||||||
"""Run one complete cycle for callers that do not use the split scheduler."""
|
"""Run one complete cycle for callers that do not use the split scheduler."""
|
||||||
@@ -389,13 +423,6 @@ class TakeoverService:
|
|||||||
max_message_id = _numeric_id(cursor.last_message_id)
|
max_message_id = _numeric_id(cursor.last_message_id)
|
||||||
read_receipts: dict[str, int] = {}
|
read_receipts: dict[str, int] = {}
|
||||||
for message in messages:
|
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"))
|
message_id = _numeric_id(message.get("id"))
|
||||||
max_message_id = max(max_message_id, message_id)
|
max_message_id = max(max_message_id, message_id)
|
||||||
send_id = str(message.get("sendId") or "")
|
send_id = str(message.get("sendId") or "")
|
||||||
@@ -410,11 +437,22 @@ class TakeoverService:
|
|||||||
session["access_token"], peer_id, message_id
|
session["access_token"], peer_id, message_id
|
||||||
)
|
)
|
||||||
|
|
||||||
cursor.last_message_id = str(max_message_id)
|
# Keep SQLite write transactions short. The read-receipt request above
|
||||||
cursor.initialized = True
|
# can block on the network and must not hold the database write lock.
|
||||||
cursor.last_polled_at = self.now()
|
async with self._persist_lock:
|
||||||
cursor.last_error = ""
|
for message in messages:
|
||||||
db.commit()
|
self._record_message(
|
||||||
|
db,
|
||||||
|
avatar,
|
||||||
|
cursor.boxim_owner_id,
|
||||||
|
message,
|
||||||
|
schedule_reply=not priming,
|
||||||
|
)
|
||||||
|
cursor.last_message_id = str(max_message_id)
|
||||||
|
cursor.initialized = True
|
||||||
|
cursor.last_polled_at = self.now()
|
||||||
|
cursor.last_error = ""
|
||||||
|
db.commit()
|
||||||
return True
|
return True
|
||||||
except Exception:
|
except Exception:
|
||||||
db.rollback()
|
db.rollback()
|
||||||
@@ -498,8 +536,27 @@ class TakeoverService:
|
|||||||
if not is_avatar:
|
if not is_avatar:
|
||||||
self._cancel_conversation(db, avatar.owner_id, peer_id, "owner_replied")
|
self._cancel_conversation(db, avatar.owner_id, peer_id, "owner_replied")
|
||||||
return
|
return
|
||||||
if not schedule_reply or event.message_type != 0 or not event.content.strip():
|
if not schedule_reply or event.message_type not in {
|
||||||
|
BOXIM_TEXT_MESSAGE_TYPE,
|
||||||
|
BOXIM_IMAGE_MESSAGE_TYPE,
|
||||||
|
}:
|
||||||
return
|
return
|
||||||
|
if event.message_type == BOXIM_TEXT_MESSAGE_TYPE and not event.content.strip():
|
||||||
|
return
|
||||||
|
if event.message_type == BOXIM_IMAGE_MESSAGE_TYPE:
|
||||||
|
try:
|
||||||
|
parse_boxim_image_url(
|
||||||
|
event.content,
|
||||||
|
base_url=getattr(self.boxim, "im_base_url", ""),
|
||||||
|
)
|
||||||
|
except BoxIMImageError as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Ignored invalid BOXIM image message %s for avatar %s: %s",
|
||||||
|
message_id,
|
||||||
|
avatar.id,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
return
|
||||||
if (now - send_time).total_seconds() > self.max_message_age_seconds:
|
if (now - send_time).total_seconds() > self.max_message_age_seconds:
|
||||||
logger.info(
|
logger.info(
|
||||||
"Ignored stale BOXIM message %s for avatar %s (age=%ss)",
|
"Ignored stale BOXIM message %s for avatar %s (age=%ss)",
|
||||||
@@ -606,7 +663,16 @@ class TakeoverService:
|
|||||||
task.status = "cancelled"
|
task.status = "cancelled"
|
||||||
task.cancel_reason = "newer_incoming_message"
|
task.cancel_reason = "newer_incoming_message"
|
||||||
task.locked_at = None
|
task.locked_at = None
|
||||||
prompt_parts.append(event.content.strip())
|
if event.message_type == BOXIM_TEXT_MESSAGE_TYPE:
|
||||||
|
for image_event in self._recent_unhandled_images(
|
||||||
|
db,
|
||||||
|
avatar,
|
||||||
|
event,
|
||||||
|
source_ids,
|
||||||
|
):
|
||||||
|
prompt_parts.append(_event_prompt(image_event))
|
||||||
|
source_ids.append(image_event.boxim_message_id)
|
||||||
|
prompt_parts.append(_event_prompt(event))
|
||||||
source_ids.append(event.boxim_message_id)
|
source_ids.append(event.boxim_message_id)
|
||||||
prompt = "\n".join(part for part in prompt_parts if part).strip()[-MAX_PROMPT_LENGTH:]
|
prompt = "\n".join(part for part in prompt_parts if part).strip()[-MAX_PROMPT_LENGTH:]
|
||||||
due_at = max(
|
due_at = max(
|
||||||
@@ -631,6 +697,68 @@ class TakeoverService:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _recent_unhandled_images(
|
||||||
|
db: Session,
|
||||||
|
avatar: Avatar,
|
||||||
|
event: TakeoverMessage,
|
||||||
|
current_source_ids: list[str],
|
||||||
|
) -> list[TakeoverMessage]:
|
||||||
|
"""Recover missed images, or reuse a referenced image from the last two days."""
|
||||||
|
references_image = _references_recent_image(event.content)
|
||||||
|
lookback_seconds = (
|
||||||
|
IMAGE_REFERENCE_LOOKBACK_SECONDS
|
||||||
|
if references_image
|
||||||
|
else IMAGE_CONTEXT_LOOKBACK_SECONDS
|
||||||
|
)
|
||||||
|
threshold = event.send_time - timedelta(seconds=lookback_seconds)
|
||||||
|
candidates = (
|
||||||
|
db.query(TakeoverMessage)
|
||||||
|
.filter(
|
||||||
|
TakeoverMessage.avatar_id == avatar.id,
|
||||||
|
TakeoverMessage.owner_id == avatar.owner_id,
|
||||||
|
TakeoverMessage.peer_id == event.peer_id,
|
||||||
|
TakeoverMessage.direction == "incoming",
|
||||||
|
TakeoverMessage.message_type == BOXIM_IMAGE_MESSAGE_TYPE,
|
||||||
|
TakeoverMessage.is_avatar.is_(False),
|
||||||
|
TakeoverMessage.send_time >= threshold,
|
||||||
|
TakeoverMessage.send_time <= event.send_time,
|
||||||
|
)
|
||||||
|
.order_by(TakeoverMessage.send_time.desc())
|
||||||
|
.limit(MAX_RECENT_IMAGE_CONTEXTS)
|
||||||
|
.all()
|
||||||
|
)
|
||||||
|
if not candidates:
|
||||||
|
return []
|
||||||
|
|
||||||
|
current_ids = set(current_source_ids)
|
||||||
|
if references_image:
|
||||||
|
return [
|
||||||
|
image
|
||||||
|
for image in reversed(candidates)
|
||||||
|
if image.boxim_message_id not in current_ids
|
||||||
|
]
|
||||||
|
|
||||||
|
handled_ids = set(current_ids)
|
||||||
|
task_sources = (
|
||||||
|
db.query(TakeoverReplyTask.source_message_ids)
|
||||||
|
.filter(
|
||||||
|
TakeoverReplyTask.avatar_id == avatar.id,
|
||||||
|
TakeoverReplyTask.owner_id == avatar.owner_id,
|
||||||
|
TakeoverReplyTask.peer_id == event.peer_id,
|
||||||
|
TakeoverReplyTask.created_at >= threshold,
|
||||||
|
)
|
||||||
|
.all()
|
||||||
|
)
|
||||||
|
for (source_message_ids,) in task_sources:
|
||||||
|
handled_ids.update(source_message_ids or [])
|
||||||
|
|
||||||
|
return [
|
||||||
|
image
|
||||||
|
for image in reversed(candidates)
|
||||||
|
if image.boxim_message_id not in handled_ids
|
||||||
|
]
|
||||||
|
|
||||||
async def _prepare_replies(self) -> int:
|
async def _prepare_replies(self) -> int:
|
||||||
db = self.session_factory()
|
db = self.session_factory()
|
||||||
try:
|
try:
|
||||||
@@ -665,6 +793,50 @@ class TakeoverService:
|
|||||||
results = await asyncio.gather(*(generate(task_id) for task_id in task_ids))
|
results = await asyncio.gather(*(generate(task_id) for task_id in task_ids))
|
||||||
return sum(bool(result) for result in results)
|
return sum(bool(result) for result in results)
|
||||||
|
|
||||||
|
def _takeover_image_attachment(
|
||||||
|
self,
|
||||||
|
db: Session,
|
||||||
|
avatar: Avatar,
|
||||||
|
event: TakeoverMessage,
|
||||||
|
) -> ChatAttachment:
|
||||||
|
now = self.now()
|
||||||
|
if event.attachment_id:
|
||||||
|
cached = db.get(ChatAttachment, event.attachment_id)
|
||||||
|
if cached and cached.status == "ready" and cached.expires_at > now:
|
||||||
|
cached.used_at = now
|
||||||
|
db.commit()
|
||||||
|
return cached
|
||||||
|
|
||||||
|
downloaded = download_boxim_image(
|
||||||
|
event.content,
|
||||||
|
base_url=getattr(
|
||||||
|
self.boxim,
|
||||||
|
"im_base_url",
|
||||||
|
os.getenv("BOXIM_API_BASE_URL", "https://im.99hui.com/api"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
from routers.chat import _analyze_image_bytes
|
||||||
|
|
||||||
|
attachment = _analyze_image_bytes(
|
||||||
|
db,
|
||||||
|
avatar,
|
||||||
|
downloaded.content,
|
||||||
|
filename=downloaded.filename,
|
||||||
|
mime_type=downloaded.mime_type,
|
||||||
|
uploader_kind="boxim",
|
||||||
|
)
|
||||||
|
event.attachment_id = attachment.id
|
||||||
|
attachment.used_at = now
|
||||||
|
db.commit()
|
||||||
|
logger.info(
|
||||||
|
"BOXIM image analyzed message=%s attachment=%s avatar=%s category=%s",
|
||||||
|
event.boxim_message_id,
|
||||||
|
attachment.id,
|
||||||
|
avatar.id,
|
||||||
|
attachment.category,
|
||||||
|
)
|
||||||
|
return attachment
|
||||||
|
|
||||||
def _generate_reply(self, task_id: str) -> bool:
|
def _generate_reply(self, task_id: str) -> bool:
|
||||||
db = self.session_factory()
|
db = self.session_factory()
|
||||||
try:
|
try:
|
||||||
@@ -683,6 +855,21 @@ class TakeoverService:
|
|||||||
db.commit()
|
db.commit()
|
||||||
|
|
||||||
excluded_ids = set(task.source_message_ids or [])
|
excluded_ids = set(task.source_message_ids or [])
|
||||||
|
source_events = {
|
||||||
|
event.boxim_message_id: event
|
||||||
|
for event in (
|
||||||
|
db.query(TakeoverMessage)
|
||||||
|
.filter(
|
||||||
|
TakeoverMessage.owner_id == task.owner_id,
|
||||||
|
TakeoverMessage.peer_id == task.peer_id,
|
||||||
|
TakeoverMessage.avatar_id == task.avatar_id,
|
||||||
|
TakeoverMessage.boxim_message_id.in_(excluded_ids),
|
||||||
|
)
|
||||||
|
.all()
|
||||||
|
if excluded_ids
|
||||||
|
else []
|
||||||
|
)
|
||||||
|
}
|
||||||
events = (
|
events = (
|
||||||
db.query(TakeoverMessage)
|
db.query(TakeoverMessage)
|
||||||
.filter(
|
.filter(
|
||||||
@@ -694,9 +881,31 @@ class TakeoverService:
|
|||||||
.limit(30)
|
.limit(30)
|
||||||
.all()
|
.all()
|
||||||
)
|
)
|
||||||
|
image_attachments = []
|
||||||
|
image_failed = False
|
||||||
|
for message_id in (task.source_message_ids or [])[-3:]:
|
||||||
|
event = source_events.get(message_id)
|
||||||
|
if not event or event.message_type != BOXIM_IMAGE_MESSAGE_TYPE:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
image_attachments.append(
|
||||||
|
self._takeover_image_attachment(db, avatar, event)
|
||||||
|
)
|
||||||
|
except (BoxIMImageError, ImageValidationError) as exc:
|
||||||
|
image_failed = True
|
||||||
|
logger.warning(
|
||||||
|
"BOXIM image unavailable message=%s avatar=%s: %s",
|
||||||
|
event.boxim_message_id,
|
||||||
|
avatar.id,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
history = []
|
history = []
|
||||||
for event in reversed(events):
|
for event in reversed(events):
|
||||||
if event.boxim_message_id in excluded_ids or not event.content.strip():
|
if (
|
||||||
|
event.boxim_message_id in excluded_ids
|
||||||
|
or event.message_type != BOXIM_TEXT_MESSAGE_TYPE
|
||||||
|
or not event.content.strip()
|
||||||
|
):
|
||||||
continue
|
continue
|
||||||
if event.direction == "incoming" and event.is_avatar:
|
if event.direction == "incoming" and event.is_avatar:
|
||||||
continue
|
continue
|
||||||
@@ -708,10 +917,21 @@ class TakeoverService:
|
|||||||
)
|
)
|
||||||
history = history[-10:]
|
history = history[-10:]
|
||||||
|
|
||||||
from routers.chat import _resolve_reply
|
from routers.chat import _attachment_contexts, _resolve_reply
|
||||||
|
|
||||||
result = _resolve_reply(db, avatar, task.prompt, history, usage_source="takeover")
|
image_contexts = _attachment_contexts(image_attachments)
|
||||||
answer = _plain_text_reply(result.get("answer", ""))
|
if image_failed and not image_contexts:
|
||||||
|
answer = BOXIM_IMAGE_UNAVAILABLE_REPLY
|
||||||
|
else:
|
||||||
|
result = _resolve_reply(
|
||||||
|
db,
|
||||||
|
avatar,
|
||||||
|
task.prompt,
|
||||||
|
history,
|
||||||
|
usage_source="takeover",
|
||||||
|
image_contexts=image_contexts,
|
||||||
|
)
|
||||||
|
answer = _plain_text_reply(result.get("answer", ""))
|
||||||
db.refresh(task)
|
db.refresh(task)
|
||||||
if task.status != "generating":
|
if task.status != "generating":
|
||||||
return False
|
return False
|
||||||
|
|||||||
@@ -0,0 +1,71 @@
|
|||||||
|
import ipaddress
|
||||||
|
import json
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from services.boxim_image_service import (
|
||||||
|
BoxIMImageError,
|
||||||
|
download_boxim_image,
|
||||||
|
parse_boxim_image_url,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_parse_boxim_image_prefers_origin_and_supports_relative_url():
|
||||||
|
content = json.dumps({"originUrl": "/files/original.png", "thumbUrl": "/thumb.png"})
|
||||||
|
assert parse_boxim_image_url(content, base_url="https://im.example/api") == (
|
||||||
|
"https://im.example/files/original.png"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_download_boxim_image_streams_public_https(monkeypatch):
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"services.boxim_image_service._resolved_addresses",
|
||||||
|
lambda _host, _port: {ipaddress.ip_address("8.8.8.8")},
|
||||||
|
)
|
||||||
|
transport = httpx.MockTransport(
|
||||||
|
lambda request: httpx.Response(
|
||||||
|
200,
|
||||||
|
headers={"content-type": "image/png"},
|
||||||
|
content=b"png-bytes",
|
||||||
|
request=request,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
image = download_boxim_image(
|
||||||
|
json.dumps({"originUrl": "https://cdn.example/case%20photo.png"}),
|
||||||
|
transport=transport,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert image.content == b"png-bytes"
|
||||||
|
assert image.filename == "case photo.png"
|
||||||
|
assert image.mime_type == "image/png"
|
||||||
|
|
||||||
|
|
||||||
|
def test_download_boxim_image_rejects_private_network_url():
|
||||||
|
with pytest.raises(BoxIMImageError, match="受限网络"):
|
||||||
|
download_boxim_image(
|
||||||
|
json.dumps({"originUrl": "https://127.0.0.1/private.png"}),
|
||||||
|
transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request)),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_download_boxim_image_stops_oversized_stream(monkeypatch):
|
||||||
|
monkeypatch.setenv("CHAT_IMAGE_MAX_BYTES", "1024")
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"services.boxim_image_service._resolved_addresses",
|
||||||
|
lambda _host, _port: {ipaddress.ip_address("8.8.8.8")},
|
||||||
|
)
|
||||||
|
transport = httpx.MockTransport(
|
||||||
|
lambda request: httpx.Response(
|
||||||
|
200,
|
||||||
|
headers={"content-length": "2048"},
|
||||||
|
request=request,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(BoxIMImageError, match="超过大小限制"):
|
||||||
|
download_boxim_image(
|
||||||
|
json.dumps({"originUrl": "https://cdn.example/large.png"}),
|
||||||
|
transport=transport,
|
||||||
|
)
|
||||||
@@ -12,11 +12,13 @@ from main import app
|
|||||||
from models import ChatAttachment
|
from models import ChatAttachment
|
||||||
from routers.chat import (
|
from routers.chat import (
|
||||||
ChatIn,
|
ChatIn,
|
||||||
|
_answer_denies_available_image,
|
||||||
_attachment_contexts,
|
_attachment_contexts,
|
||||||
_load_chat_attachments,
|
_load_chat_attachments,
|
||||||
_resolve_reply,
|
_resolve_reply,
|
||||||
)
|
)
|
||||||
from services.chat_attachment_service import purge_expired_chat_attachments
|
from services.chat_attachment_service import purge_expired_chat_attachments
|
||||||
|
from services.token_billing import InsufficientTokensError
|
||||||
from services.vision_service import PreparedImage
|
from services.vision_service import PreparedImage
|
||||||
|
|
||||||
|
|
||||||
@@ -75,6 +77,22 @@ def test_non_owner_cannot_upload_chat_image(authorization_context):
|
|||||||
assert response.status_code == 403
|
assert response.status_code == 403
|
||||||
|
|
||||||
|
|
||||||
|
def test_image_upload_preserves_insufficient_points_response(authorization_context):
|
||||||
|
context = authorization_context
|
||||||
|
with patch(
|
||||||
|
"routers.chat._analyze_image_bytes",
|
||||||
|
side_effect=InsufficientTokensError("积分余额不足"),
|
||||||
|
):
|
||||||
|
response = client.post(
|
||||||
|
f"/api/avatar/{context['avatar'].id}/chat/images",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
files={"file": ("private.png", b"image-bytes", "image/png")},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 402
|
||||||
|
assert response.json()["detail"] == "积分余额不足"
|
||||||
|
|
||||||
|
|
||||||
def test_public_share_can_upload_without_exposing_analysis_details(authorization_context):
|
def test_public_share_can_upload_without_exposing_analysis_details(authorization_context):
|
||||||
context = authorization_context
|
context = authorization_context
|
||||||
db = SessionLocal()
|
db = SessionLocal()
|
||||||
@@ -257,6 +275,47 @@ def test_image_context_keeps_standard_answer_authoritative():
|
|||||||
assert "标准答题对中的事实优先级高于图片资料" in system
|
assert "标准答题对中的事实优先级高于图片资料" in system
|
||||||
|
|
||||||
|
|
||||||
|
def test_ready_image_context_never_returns_whole_image_access_denial():
|
||||||
|
avatar = SimpleNamespace(
|
||||||
|
id="avatar-vision",
|
||||||
|
name="测试分身",
|
||||||
|
description="产品顾问",
|
||||||
|
config={},
|
||||||
|
)
|
||||||
|
model = Mock(return_value="抱歉,我无法查看或识别图片,请重新上传。")
|
||||||
|
result = _resolve_reply(
|
||||||
|
None,
|
||||||
|
avatar,
|
||||||
|
"请看看这张图片",
|
||||||
|
[],
|
||||||
|
qa_pairs=[],
|
||||||
|
search_fn=Mock(return_value=[]),
|
||||||
|
model_client=model,
|
||||||
|
image_contexts=[{
|
||||||
|
"id": "attachment",
|
||||||
|
"filename": "report.jpg",
|
||||||
|
"category": "medical_document",
|
||||||
|
"summary": "一份耳鼻喉科门诊记录",
|
||||||
|
"extractedText": "主诉:咽痛三天",
|
||||||
|
"structuredData": {"key_facts": ["主诉为咽痛三天"]},
|
||||||
|
"warning": "请核对原始资料",
|
||||||
|
}],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result["source"] == "vision"
|
||||||
|
assert "一份耳鼻喉科门诊记录" in result["answer"]
|
||||||
|
assert "主诉为咽痛三天" in result["answer"]
|
||||||
|
assert "无法查看" not in result["answer"]
|
||||||
|
system = model.call_args.kwargs["messages"][0]["content"]
|
||||||
|
assert "当前会话图片已经成功读取" in system
|
||||||
|
assert "禁止声称无法查看" in system
|
||||||
|
|
||||||
|
|
||||||
|
def test_image_denial_detector_allows_uncertain_field_in_ready_image():
|
||||||
|
assert _answer_denies_available_image("我无法查看这张图片") is True
|
||||||
|
assert _answer_denies_available_image("图片中患者姓名无法辨认,主诉为咽痛三天。") is False
|
||||||
|
|
||||||
|
|
||||||
def test_attachment_context_does_not_expose_internal_fields():
|
def test_attachment_context_does_not_expose_internal_fields():
|
||||||
row = SimpleNamespace(
|
row = SimpleNamespace(
|
||||||
id="attachment",
|
id="attachment",
|
||||||
|
|||||||
@@ -49,6 +49,7 @@ class RemoteEmbeddingTests(unittest.TestCase):
|
|||||||
texts = [f"chunk-{index}" for index in range(14)]
|
texts = [f"chunk-{index}" for index in range(14)]
|
||||||
batch_sizes = []
|
batch_sizes = []
|
||||||
requested_urls = []
|
requested_urls = []
|
||||||
|
progress_updates = []
|
||||||
|
|
||||||
def fake_urlopen(request, timeout):
|
def fake_urlopen(request, timeout):
|
||||||
self.assertEqual(timeout, 30)
|
self.assertEqual(timeout, 30)
|
||||||
@@ -68,7 +69,10 @@ class RemoteEmbeddingTests(unittest.TestCase):
|
|||||||
"EMBEDDING_MODEL": "text-embedding-v4",
|
"EMBEDDING_MODEL": "text-embedding-v4",
|
||||||
"EMBEDDING_BATCH_SIZE": "10",
|
"EMBEDDING_BATCH_SIZE": "10",
|
||||||
}), patch("embeddings.urllib.request.urlopen", side_effect=fake_urlopen):
|
}), patch("embeddings.urllib.request.urlopen", side_effect=fake_urlopen):
|
||||||
result = embeddings.embed(texts)
|
result = embeddings.embed(
|
||||||
|
texts,
|
||||||
|
on_progress=lambda completed, total: progress_updates.append((completed, total)),
|
||||||
|
)
|
||||||
|
|
||||||
self.assertEqual(batch_sizes, [10, 4])
|
self.assertEqual(batch_sizes, [10, 4])
|
||||||
self.assertEqual(requested_urls, [
|
self.assertEqual(requested_urls, [
|
||||||
@@ -76,6 +80,7 @@ class RemoteEmbeddingTests(unittest.TestCase):
|
|||||||
"https://embedding.example/v1/embeddings",
|
"https://embedding.example/v1/embeddings",
|
||||||
])
|
])
|
||||||
self.assertEqual(result, [[float(index)] for index in range(14)])
|
self.assertEqual(result, [[float(index)] for index in range(14)])
|
||||||
|
self.assertEqual(progress_updates, [(10, 14), (14, 14)])
|
||||||
|
|
||||||
def test_full_embedding_endpoint_is_not_modified(self):
|
def test_full_embedding_endpoint_is_not_modified(self):
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from database import SessionLocal
|
|||||||
from main import app
|
from main import app
|
||||||
from models import Avatar, KnowledgeChunk, KnowledgeDoc, QAPair
|
from models import Avatar, KnowledgeChunk, KnowledgeDoc, QAPair
|
||||||
from routers.knowledge import _doc_payload
|
from routers.knowledge import _doc_payload
|
||||||
|
from services.knowledge_vectorizer import knowledge_vectorizer
|
||||||
|
|
||||||
|
|
||||||
client = TestClient(app)
|
client = TestClient(app)
|
||||||
@@ -31,14 +32,14 @@ def test_doc_payload_reports_whether_the_persisted_file_exists(tmp_path: Path):
|
|||||||
assert _doc_payload(doc)["filePresent"] is True
|
assert _doc_payload(doc)["filePresent"] is True
|
||||||
|
|
||||||
|
|
||||||
def test_upload_marks_vectorization_failure_instead_of_staying_processing(
|
def test_upload_returns_before_background_vectorization(
|
||||||
tmp_path: Path,
|
tmp_path: Path,
|
||||||
authorization_context,
|
authorization_context,
|
||||||
):
|
):
|
||||||
context = authorization_context
|
context = authorization_context
|
||||||
with (
|
with (
|
||||||
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
||||||
patch("routers.knowledge.embeddings.embed", side_effect=RuntimeError("provider unavailable")),
|
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
|
||||||
):
|
):
|
||||||
response = client.post(
|
response = client.post(
|
||||||
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
|
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
|
||||||
@@ -47,14 +48,15 @@ def test_upload_marks_vectorization_failure_instead_of_staying_processing(
|
|||||||
)
|
)
|
||||||
|
|
||||||
payload = response.json()["data"]
|
payload = response.json()["data"]
|
||||||
assert payload["status"] == "failed"
|
assert payload["status"] == "parsing"
|
||||||
assert payload["vectorized"] is False
|
assert payload["vectorized"] is False
|
||||||
assert payload["chunkCount"] == 0
|
assert payload["chunkCount"] == 0
|
||||||
|
enqueue.assert_called_once_with(payload["id"])
|
||||||
|
|
||||||
db = SessionLocal()
|
db = SessionLocal()
|
||||||
try:
|
try:
|
||||||
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
|
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
|
||||||
assert stored.status == "failed"
|
assert stored.status == "parsing"
|
||||||
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 0
|
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 0
|
||||||
db.delete(stored)
|
db.delete(stored)
|
||||||
db.commit()
|
db.commit()
|
||||||
@@ -62,14 +64,115 @@ def test_upload_marks_vectorization_failure_instead_of_staying_processing(
|
|||||||
db.close()
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
def test_markdown_upload_commits_ready_document_and_chunks_together(
|
def test_upload_rejects_oversize_file_before_queuing_indexing(
|
||||||
tmp_path: Path,
|
tmp_path: Path,
|
||||||
authorization_context,
|
authorization_context,
|
||||||
):
|
):
|
||||||
context = authorization_context
|
context = authorization_context
|
||||||
with (
|
with (
|
||||||
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
||||||
patch("routers.knowledge.embeddings.embed", return_value=[[1.0, 0.0]]),
|
patch("routers.knowledge.MAX_UPLOAD_BYTES", 4),
|
||||||
|
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
|
||||||
|
):
|
||||||
|
response = client.post(
|
||||||
|
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
files={"file": ("oversize.md", b"12345", "text/markdown")},
|
||||||
|
)
|
||||||
|
|
||||||
|
payload = response.json()
|
||||||
|
assert payload["code"] == 400
|
||||||
|
assert payload["message"] == "文件不能超过 50MB"
|
||||||
|
enqueue.assert_not_called()
|
||||||
|
assert not list((tmp_path / context["avatar"].id).glob("*"))
|
||||||
|
|
||||||
|
|
||||||
|
def test_multipart_upload_reassembles_file_before_queuing_indexing(
|
||||||
|
tmp_path: Path,
|
||||||
|
authorization_context,
|
||||||
|
):
|
||||||
|
context = authorization_context
|
||||||
|
avatar_id = context["avatar"].id
|
||||||
|
content = b"0123456789"
|
||||||
|
with (
|
||||||
|
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
||||||
|
patch("routers.knowledge.MULTIPART_CHUNK_BYTES", 4),
|
||||||
|
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
|
||||||
|
):
|
||||||
|
created = client.post(
|
||||||
|
f"/api/avatar/{avatar_id}/knowledge/uploads",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
json={"filename": "large.pdf", "fileSize": len(content), "totalChunks": 3},
|
||||||
|
).json()["data"]
|
||||||
|
|
||||||
|
for index, chunk in enumerate((content[:4], content[4:8], content[8:])):
|
||||||
|
response = client.post(
|
||||||
|
f"/api/avatar/{avatar_id}/knowledge/uploads/{created['uploadId']}/chunks/{index}",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
files={"file": (f"chunk-{index}", chunk, "application/octet-stream")},
|
||||||
|
)
|
||||||
|
assert response.json()["code"] == 200
|
||||||
|
|
||||||
|
completed = client.post(
|
||||||
|
f"/api/avatar/{avatar_id}/knowledge/uploads/{created['uploadId']}/complete",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
).json()["data"]
|
||||||
|
|
||||||
|
assert completed["status"] == "parsing"
|
||||||
|
assert completed["fileSize"] == len(content)
|
||||||
|
enqueue.assert_called_once_with(completed["id"])
|
||||||
|
stored_path = tmp_path / avatar_id / Path(completed["fileUrl"]).name
|
||||||
|
assert stored_path.read_bytes() == content
|
||||||
|
assert not (tmp_path / ".multipart" / avatar_id / created["uploadId"]).exists()
|
||||||
|
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == completed["id"]).one()
|
||||||
|
db.delete(stored)
|
||||||
|
db.commit()
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
def test_multipart_upload_rejects_incomplete_parts(
|
||||||
|
tmp_path: Path,
|
||||||
|
authorization_context,
|
||||||
|
):
|
||||||
|
context = authorization_context
|
||||||
|
avatar_id = context["avatar"].id
|
||||||
|
with (
|
||||||
|
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
||||||
|
patch("routers.knowledge.MULTIPART_CHUNK_BYTES", 4),
|
||||||
|
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
|
||||||
|
):
|
||||||
|
created = client.post(
|
||||||
|
f"/api/avatar/{avatar_id}/knowledge/uploads",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
json={"filename": "large.pdf", "fileSize": 6, "totalChunks": 2},
|
||||||
|
).json()["data"]
|
||||||
|
client.post(
|
||||||
|
f"/api/avatar/{avatar_id}/knowledge/uploads/{created['uploadId']}/chunks/0",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
files={"file": ("chunk-0", b"0123", "application/octet-stream")},
|
||||||
|
)
|
||||||
|
response = client.post(
|
||||||
|
f"/api/avatar/{avatar_id}/knowledge/uploads/{created['uploadId']}/complete",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.json()["code"] == 400
|
||||||
|
assert response.json()["message"] == "文件分片尚未上传完整"
|
||||||
|
enqueue.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
def test_background_vectorizer_commits_ready_document_and_chunks_together(
|
||||||
|
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(
|
response = client.post(
|
||||||
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
|
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
|
||||||
@@ -78,14 +181,21 @@ def test_markdown_upload_commits_ready_document_and_chunks_together(
|
|||||||
)
|
)
|
||||||
|
|
||||||
payload = response.json()["data"]
|
payload = response.json()["data"]
|
||||||
assert payload["status"] == "ready"
|
assert payload["status"] == "parsing"
|
||||||
assert payload["vectorized"] is True
|
with (
|
||||||
assert payload["chunkCount"] == 1
|
patch("services.knowledge_vectorizer.UPLOAD_DIR", str(tmp_path)),
|
||||||
|
patch("services.knowledge_vectorizer.embeddings.embed", return_value=[[1.0, 0.0]]),
|
||||||
|
):
|
||||||
|
knowledge_vectorizer.vectorize_document(payload["id"])
|
||||||
|
|
||||||
db = SessionLocal()
|
db = SessionLocal()
|
||||||
try:
|
try:
|
||||||
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
|
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
|
||||||
assert stored.status == "ready"
|
assert stored.status == "ready"
|
||||||
|
assert stored.vectorized is True
|
||||||
|
assert stored.chunk_count == 1
|
||||||
|
assert stored.index_stage == "ready"
|
||||||
|
assert stored.index_progress == 100
|
||||||
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 1
|
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 1
|
||||||
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).delete()
|
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).delete()
|
||||||
db.delete(stored)
|
db.delete(stored)
|
||||||
@@ -94,6 +204,87 @@ def test_markdown_upload_commits_ready_document_and_chunks_together(
|
|||||||
db.close()
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
def test_background_vectorizer_keeps_failure_reason_for_retry(
|
||||||
|
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": ("knowledge.md", b"# Knowledge\n\nTest content", "text/markdown")},
|
||||||
|
)
|
||||||
|
|
||||||
|
payload = response.json()["data"]
|
||||||
|
with (
|
||||||
|
patch("services.knowledge_vectorizer.UPLOAD_DIR", str(tmp_path)),
|
||||||
|
patch("services.knowledge_vectorizer.embeddings.embed", side_effect=RuntimeError("provider unavailable")),
|
||||||
|
):
|
||||||
|
knowledge_vectorizer.vectorize_document(payload["id"])
|
||||||
|
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
|
||||||
|
assert stored.status == "failed"
|
||||||
|
assert stored.error_message == "provider unavailable"
|
||||||
|
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 0
|
||||||
|
db.delete(stored)
|
||||||
|
db.commit()
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
def test_retry_queues_a_failed_document_again(
|
||||||
|
tmp_path: Path,
|
||||||
|
authorization_context,
|
||||||
|
):
|
||||||
|
context = authorization_context
|
||||||
|
document_id = f"retry-doc-{context['suffix']}"
|
||||||
|
avatar_dir = tmp_path / context["avatar"].id
|
||||||
|
avatar_dir.mkdir()
|
||||||
|
(avatar_dir / "retry.md").write_text("retry content", encoding="utf-8")
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
db.add(
|
||||||
|
KnowledgeDoc(
|
||||||
|
id=document_id,
|
||||||
|
avatar_id=context["avatar"].id,
|
||||||
|
filename="retry.md",
|
||||||
|
file_type="md",
|
||||||
|
file_url=f"/api/files/{context['avatar'].id}/retry.md",
|
||||||
|
status="failed",
|
||||||
|
error_message="provider unavailable",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
db.commit()
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
||||||
|
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
|
||||||
|
):
|
||||||
|
response = client.post(
|
||||||
|
f"/api/avatar/{context['avatar'].id}/knowledge/docs/{document_id}/retry",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
)
|
||||||
|
|
||||||
|
payload = response.json()["data"]
|
||||||
|
assert payload["status"] == "parsing"
|
||||||
|
assert payload["errorMessage"] == ""
|
||||||
|
enqueue.assert_called_once_with(document_id)
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
db.query(KnowledgeDoc).filter(KnowledgeDoc.id == document_id).delete()
|
||||||
|
db.commit()
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
def test_each_avatar_has_an_independent_document_and_qa_scope(authorization_context):
|
def test_each_avatar_has_an_independent_document_and_qa_scope(authorization_context):
|
||||||
context = authorization_context
|
context = authorization_context
|
||||||
first_avatar_id = context["avatar"].id
|
first_avatar_id = context["avatar"].id
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
"""End-to-end service tests for BOXIM takeover timing and human priority."""
|
"""End-to-end service tests for BOXIM takeover timing and human priority."""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import json
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
from threading import Barrier
|
from threading import Barrier
|
||||||
from unittest.mock import AsyncMock, patch
|
from unittest.mock import AsyncMock, patch
|
||||||
@@ -10,8 +11,9 @@ from sqlalchemy import create_engine
|
|||||||
from sqlalchemy.orm import sessionmaker
|
from sqlalchemy.orm import sessionmaker
|
||||||
|
|
||||||
from database import Base
|
from database import Base
|
||||||
from models import Avatar, TakeoverCursor, TakeoverMessage, TakeoverReplyTask, User
|
from models import Avatar, ChatAttachment, TakeoverCursor, TakeoverMessage, TakeoverReplyTask, User
|
||||||
from services.boxim_client import BoxIMError
|
from services.boxim_client import BoxIMError
|
||||||
|
from services.boxim_image_service import DownloadedBoxIMImage
|
||||||
from services.takeover_service import (
|
from services.takeover_service import (
|
||||||
AVATAR_LOCAL_ID_PREFIX,
|
AVATAR_LOCAL_ID_PREFIX,
|
||||||
TakeoverService,
|
TakeoverService,
|
||||||
@@ -83,6 +85,29 @@ class ConcurrentPollingBoxIM(FakeBoxIM):
|
|||||||
return []
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
class ConcurrentMessagePollingBoxIM(ConcurrentPollingBoxIM):
|
||||||
|
async def fetch_private_messages(self, access_token, min_id="0"):
|
||||||
|
await super().fetch_private_messages(access_token, min_id)
|
||||||
|
owner_id = 100 if access_token == "prod-huihui-token" else 101
|
||||||
|
return [
|
||||||
|
{
|
||||||
|
"id": owner_id,
|
||||||
|
"localId": owner_id,
|
||||||
|
"sendId": owner_id + 100,
|
||||||
|
"recvId": owner_id,
|
||||||
|
"sendTime": 1_700_000_000_000,
|
||||||
|
"type": 0,
|
||||||
|
"content": "并发写入测试",
|
||||||
|
}
|
||||||
|
]
|
||||||
|
|
||||||
|
async def mark_private_messages_read(self, access_token, friend_id, message_id):
|
||||||
|
await asyncio.sleep(0.05)
|
||||||
|
self.read_receipts.append(
|
||||||
|
{"friendId": str(friend_id), "messageId": str(message_id)}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def service_context(tmp_path):
|
def service_context(tmp_path):
|
||||||
engine = create_engine(
|
engine = create_engine(
|
||||||
@@ -171,6 +196,247 @@ async def test_incoming_message_is_prepared_then_sent_at_three_seconds(service_c
|
|||||||
db.close()
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_incoming_image_is_analyzed_and_used_in_takeover_reply(service_context):
|
||||||
|
session_factory, service, boxim, clock = service_context
|
||||||
|
await service.poll_and_process_messages()
|
||||||
|
boxim.messages.append(
|
||||||
|
{
|
||||||
|
"id": 111,
|
||||||
|
"localId": 111,
|
||||||
|
"sendId": 200,
|
||||||
|
"recvId": 100,
|
||||||
|
"sendTime": clock.millis(),
|
||||||
|
"type": 1,
|
||||||
|
"content": json.dumps(
|
||||||
|
{
|
||||||
|
"originUrl": "https://cdn.example/case.png",
|
||||||
|
"thumbUrl": "https://cdn.example/case-thumb.png",
|
||||||
|
}
|
||||||
|
),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
await service.poll_and_process_messages()
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
scheduled = db.query(TakeoverReplyTask).filter_by(trigger_message_id="111").one()
|
||||||
|
assert scheduled.status == "pending"
|
||||||
|
assert scheduled.prompt == "请看看这张图片。"
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
clock.advance(3)
|
||||||
|
|
||||||
|
def analyze(db, avatar, content, **kwargs):
|
||||||
|
assert content == b"image-content"
|
||||||
|
attachment = ChatAttachment(
|
||||||
|
avatar_id=avatar.id,
|
||||||
|
uploader_kind=kwargs["uploader_kind"],
|
||||||
|
filename=kwargs["filename"],
|
||||||
|
mime_type="image/jpeg",
|
||||||
|
file_size=len(content),
|
||||||
|
status="ready",
|
||||||
|
category="medical_document",
|
||||||
|
summary="一张门诊病例",
|
||||||
|
extracted_text="主诉:咳嗽三天",
|
||||||
|
structured_data={"medical": {"chief_complaint": "咳嗽三天"}},
|
||||||
|
warning="请核对原始资料",
|
||||||
|
expires_at=clock.now() + timedelta(hours=24),
|
||||||
|
)
|
||||||
|
db.add(attachment)
|
||||||
|
db.commit()
|
||||||
|
db.refresh(attachment)
|
||||||
|
return attachment
|
||||||
|
|
||||||
|
downloaded = DownloadedBoxIMImage(
|
||||||
|
content=b"image-content",
|
||||||
|
filename="case.png",
|
||||||
|
mime_type="image/png",
|
||||||
|
source_url="https://cdn.example/case.png",
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
patch("services.takeover_service.download_boxim_image", return_value=downloaded),
|
||||||
|
patch("routers.chat._analyze_image_bytes", side_effect=analyze) as analyzer,
|
||||||
|
patch("routers.chat._resolve_reply", return_value={"answer": "这份资料里写的是咳嗽三天。"}) as resolver,
|
||||||
|
):
|
||||||
|
await service.poll_and_process_messages()
|
||||||
|
|
||||||
|
analyzer.assert_called_once()
|
||||||
|
assert resolver.call_args.args[2] == "请看看这张图片。"
|
||||||
|
image_contexts = resolver.call_args.kwargs["image_contexts"]
|
||||||
|
assert image_contexts[0]["summary"] == "一张门诊病例"
|
||||||
|
assert image_contexts[0]["extractedText"] == "主诉:咳嗽三天"
|
||||||
|
assert [item["content"] for item in boxim.sent] == ["这份资料里写的是咳嗽三天。"]
|
||||||
|
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
event = db.query(TakeoverMessage).filter_by(boxim_message_id="111").one()
|
||||||
|
task = db.query(TakeoverReplyTask).filter_by(trigger_message_id="111").one()
|
||||||
|
assert event.attachment_id
|
||||||
|
assert db.get(ChatAttachment, event.attachment_id).uploader_kind == "boxim"
|
||||||
|
assert task.status == "sent"
|
||||||
|
with patch(
|
||||||
|
"services.takeover_service.download_boxim_image",
|
||||||
|
side_effect=AssertionError("cached image must not be downloaded again"),
|
||||||
|
):
|
||||||
|
cached = service._takeover_image_attachment(db, db.get(Avatar, "avatar-1"), event)
|
||||||
|
assert cached.id == event.attachment_id
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_followup_text_recovers_recent_image_recorded_without_task(service_context):
|
||||||
|
session_factory, service, boxim, clock = service_context
|
||||||
|
await service.poll_and_process_messages()
|
||||||
|
image_message = {
|
||||||
|
"id": 113,
|
||||||
|
"localId": 113,
|
||||||
|
"sendId": 200,
|
||||||
|
"recvId": 100,
|
||||||
|
"sendTime": clock.millis(),
|
||||||
|
"type": 1,
|
||||||
|
"content": json.dumps(
|
||||||
|
{
|
||||||
|
"originUrl": "https://cdn.example/case.png",
|
||||||
|
"thumbUrl": "https://cdn.example/case-thumb.png",
|
||||||
|
}
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
avatar = db.get(Avatar, "avatar-1")
|
||||||
|
service._record_message(db, avatar, "100", image_message, schedule_reply=False)
|
||||||
|
cursor = db.query(TakeoverCursor).one()
|
||||||
|
cursor.last_message_id = "113"
|
||||||
|
db.commit()
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
clock.advance(60)
|
||||||
|
boxim.messages.extend(
|
||||||
|
[
|
||||||
|
image_message,
|
||||||
|
{
|
||||||
|
"id": 114,
|
||||||
|
"localId": 114,
|
||||||
|
"sendId": 200,
|
||||||
|
"recvId": 100,
|
||||||
|
"sendTime": clock.millis(),
|
||||||
|
"type": 0,
|
||||||
|
"content": "请帮我看看这张图",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
await service.poll_messages()
|
||||||
|
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
task = db.query(TakeoverReplyTask).filter_by(trigger_message_id="114").one()
|
||||||
|
assert task.source_message_ids == ["113", "114"]
|
||||||
|
assert task.prompt == "请看看这张图片。\n请帮我看看这张图"
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
clock.advance(1)
|
||||||
|
boxim.messages.append(
|
||||||
|
{
|
||||||
|
"id": 115,
|
||||||
|
"localId": 115,
|
||||||
|
"sendId": 200,
|
||||||
|
"recvId": 100,
|
||||||
|
"sendTime": clock.millis(),
|
||||||
|
"type": 0,
|
||||||
|
"content": "图里写了什么",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
await service.poll_messages()
|
||||||
|
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
latest = db.query(TakeoverReplyTask).filter_by(trigger_message_id="115").one()
|
||||||
|
assert latest.source_message_ids == ["113", "114", "115"]
|
||||||
|
assert latest.source_message_ids.count("113") == 1
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_explicit_followup_reuses_handled_image_within_two_days(service_context):
|
||||||
|
session_factory, service, boxim, clock = service_context
|
||||||
|
await service.poll_and_process_messages()
|
||||||
|
image_message = {
|
||||||
|
"id": 116,
|
||||||
|
"localId": 116,
|
||||||
|
"sendId": 200,
|
||||||
|
"recvId": 100,
|
||||||
|
"sendTime": clock.millis(),
|
||||||
|
"type": 1,
|
||||||
|
"content": json.dumps({"originUrl": "https://cdn.example/handled-case.png"}),
|
||||||
|
}
|
||||||
|
boxim.messages.append(image_message)
|
||||||
|
await service.poll_messages()
|
||||||
|
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
image_task = db.query(TakeoverReplyTask).filter_by(trigger_message_id="116").one()
|
||||||
|
image_task.status = "sent"
|
||||||
|
image_task.sent_at = clock.now()
|
||||||
|
db.commit()
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
clock.advance(47 * 60 * 60)
|
||||||
|
boxim.messages.append(
|
||||||
|
{
|
||||||
|
"id": 117,
|
||||||
|
"localId": 117,
|
||||||
|
"sendId": 200,
|
||||||
|
"recvId": 100,
|
||||||
|
"sendTime": clock.millis(),
|
||||||
|
"type": 0,
|
||||||
|
"content": "重新看一下刚才那张病例图片",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
await service.poll_messages()
|
||||||
|
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
task = db.query(TakeoverReplyTask).filter_by(trigger_message_id="117").one()
|
||||||
|
assert task.source_message_ids == ["116", "117"]
|
||||||
|
assert task.prompt == "请看看这张图片。\n重新看一下刚才那张病例图片"
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_invalid_image_message_is_recorded_but_not_scheduled(service_context):
|
||||||
|
session_factory, service, boxim, clock = service_context
|
||||||
|
await service.poll_and_process_messages()
|
||||||
|
boxim.messages.append(
|
||||||
|
{
|
||||||
|
"id": 112,
|
||||||
|
"localId": 112,
|
||||||
|
"sendId": 200,
|
||||||
|
"recvId": 100,
|
||||||
|
"sendTime": clock.millis(),
|
||||||
|
"type": 1,
|
||||||
|
"content": json.dumps({"width": 100, "height": 100}),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
await service.poll_and_process_messages()
|
||||||
|
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
assert db.query(TakeoverMessage).filter_by(boxim_message_id="112").one()
|
||||||
|
assert db.query(TakeoverReplyTask).count() == 0
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_default_reply_delay_is_three_minutes(service_context):
|
async def test_default_reply_delay_is_three_minutes(service_context):
|
||||||
session_factory, service, boxim, clock = service_context
|
session_factory, service, boxim, clock = service_context
|
||||||
@@ -230,7 +496,7 @@ async def test_multiple_avatar_owners_are_polled_concurrently(service_context):
|
|||||||
finally:
|
finally:
|
||||||
db.close()
|
db.close()
|
||||||
|
|
||||||
boxim = ConcurrentPollingBoxIM()
|
boxim = ConcurrentMessagePollingBoxIM()
|
||||||
service = TakeoverService(
|
service = TakeoverService(
|
||||||
session_factory,
|
session_factory,
|
||||||
boxim,
|
boxim,
|
||||||
@@ -244,6 +510,8 @@ async def test_multiple_avatar_owners_are_polled_concurrently(service_context):
|
|||||||
db = session_factory()
|
db = session_factory()
|
||||||
try:
|
try:
|
||||||
assert db.query(TakeoverCursor).filter(TakeoverCursor.initialized.is_(True)).count() == 2
|
assert db.query(TakeoverCursor).filter(TakeoverCursor.initialized.is_(True)).count() == 2
|
||||||
|
assert db.query(TakeoverMessage).count() == 2
|
||||||
|
assert len(boxim.read_receipts) == 2
|
||||||
finally:
|
finally:
|
||||||
db.close()
|
db.close()
|
||||||
|
|
||||||
|
|||||||
@@ -117,7 +117,7 @@ location /api/ {
|
|||||||
proxy_set_header X-Forwarded-Proto $scheme;
|
proxy_set_header X-Forwarded-Proto $scheme;
|
||||||
proxy_buffering off;
|
proxy_buffering off;
|
||||||
proxy_read_timeout 300s;
|
proxy_read_timeout 300s;
|
||||||
client_max_body_size 20m;
|
client_max_body_size 100m;
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|||||||
@@ -23,6 +23,10 @@ http {
|
|||||||
|
|
||||||
root /usr/share/nginx/html;
|
root /usr/share/nginx/html;
|
||||||
index index.html;
|
index index.html;
|
||||||
|
# Keep the application gateway aligned with the production edge gateway.
|
||||||
|
# Without this Nginx rejects ordinary PDF uploads with HTTP 413 before
|
||||||
|
# FastAPI can return its user-facing file-size validation message.
|
||||||
|
client_max_body_size 100m;
|
||||||
|
|
||||||
# SPA 兜底(hash 路由下深链接也可正常加载)
|
# SPA 兜底(hash 路由下深链接也可正常加载)
|
||||||
location / {
|
location / {
|
||||||
|
|||||||
@@ -305,6 +305,9 @@ export interface KnowledgeDoc {
|
|||||||
vectorized?: boolean
|
vectorized?: boolean
|
||||||
embeddingModel?: string
|
embeddingModel?: string
|
||||||
chunkCount?: number
|
chunkCount?: number
|
||||||
|
errorMessage?: string
|
||||||
|
indexStage?: string
|
||||||
|
indexProgress?: number
|
||||||
createdAt: string
|
createdAt: string
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -330,12 +333,83 @@ export interface SearchResult {
|
|||||||
export const getKnowledgeDocs = (avatarId: string) =>
|
export const getKnowledgeDocs = (avatarId: string) =>
|
||||||
request.get<KnowledgeDoc[]>(`/avatar/${avatarId}/knowledge/docs`)
|
request.get<KnowledgeDoc[]>(`/avatar/${avatarId}/knowledge/docs`)
|
||||||
|
|
||||||
// 上传文档(支持 md/txt/pdf/doc/docx/xlsx)
|
const KNOWLEDGE_UPLOAD_CHUNK_SIZE = 5 * 1024 * 1024
|
||||||
export const uploadKnowledgeDoc = (avatarId: string, file: File) => {
|
|
||||||
|
const uploadKnowledgeChunk = async (
|
||||||
|
avatarId: string,
|
||||||
|
uploadId: string,
|
||||||
|
chunkIndex: number,
|
||||||
|
chunk: Blob,
|
||||||
|
onProgress?: (loaded: number) => void
|
||||||
|
) => {
|
||||||
|
const form = new FormData()
|
||||||
|
form.append('file', chunk, `chunk-${chunkIndex}`)
|
||||||
|
let reportedLoaded = 0
|
||||||
|
for (let attempt = 1; attempt <= 3; attempt += 1) {
|
||||||
|
try {
|
||||||
|
await request.post(
|
||||||
|
`/avatar/${avatarId}/knowledge/uploads/${uploadId}/chunks/${chunkIndex}`,
|
||||||
|
form,
|
||||||
|
{
|
||||||
|
headers: { 'Content-Type': 'multipart/form-data' },
|
||||||
|
timeout: 2 * 60 * 1000,
|
||||||
|
onUploadProgress: (event) => {
|
||||||
|
reportedLoaded = Math.max(reportedLoaded, Math.min(event.loaded, chunk.size))
|
||||||
|
onProgress?.(reportedLoaded)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return
|
||||||
|
} catch (error: any) {
|
||||||
|
const status = Number(error?.response?.status || 0)
|
||||||
|
const retryable = !status || status === 408 || status === 429 || status >= 500
|
||||||
|
if (!retryable || attempt === 3) throw error
|
||||||
|
await new Promise((resolve) => window.setTimeout(resolve, attempt * 800))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 大文件拆成 5MB 分片,避免生产代理的请求体限制拦截整个文件。
|
||||||
|
export const uploadKnowledgeDoc = async (
|
||||||
|
avatarId: string,
|
||||||
|
file: File,
|
||||||
|
onUploadProgress?: (loaded: number, total: number) => void
|
||||||
|
) => {
|
||||||
|
if (file.size > KNOWLEDGE_UPLOAD_CHUNK_SIZE) {
|
||||||
|
const totalChunks = Math.ceil(file.size / KNOWLEDGE_UPLOAD_CHUNK_SIZE)
|
||||||
|
const upload: any = await request.post(`/avatar/${avatarId}/knowledge/uploads`, {
|
||||||
|
filename: file.name,
|
||||||
|
fileSize: file.size,
|
||||||
|
totalChunks
|
||||||
|
})
|
||||||
|
let uploadedBytes = 0
|
||||||
|
for (let index = 0; index < totalChunks; index += 1) {
|
||||||
|
const start = index * KNOWLEDGE_UPLOAD_CHUNK_SIZE
|
||||||
|
const chunk = file.slice(start, Math.min(start + KNOWLEDGE_UPLOAD_CHUNK_SIZE, file.size))
|
||||||
|
await uploadKnowledgeChunk(
|
||||||
|
avatarId,
|
||||||
|
upload.uploadId,
|
||||||
|
index,
|
||||||
|
chunk,
|
||||||
|
(chunkLoaded) => onUploadProgress?.(uploadedBytes + chunkLoaded, file.size)
|
||||||
|
)
|
||||||
|
uploadedBytes += chunk.size
|
||||||
|
onUploadProgress?.(uploadedBytes, file.size)
|
||||||
|
}
|
||||||
|
return request.post<KnowledgeDoc>(
|
||||||
|
`/avatar/${avatarId}/knowledge/uploads/${upload.uploadId}/complete`,
|
||||||
|
undefined,
|
||||||
|
{ timeout: 2 * 60 * 1000 }
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
const form = new FormData()
|
const form = new FormData()
|
||||||
form.append('file', file)
|
form.append('file', file)
|
||||||
return request.post<KnowledgeDoc>(`/avatar/${avatarId}/knowledge/docs`, form, {
|
return request.post<KnowledgeDoc>(`/avatar/${avatarId}/knowledge/docs`, form, {
|
||||||
headers: { 'Content-Type': 'multipart/form-data' }
|
headers: { 'Content-Type': 'multipart/form-data' },
|
||||||
|
// A slow mobile uplink must not be mistaken for a failed upload.
|
||||||
|
timeout: 10 * 60 * 1000,
|
||||||
|
onUploadProgress: (event) => onUploadProgress?.(event.loaded, event.total || file.size)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -343,6 +417,9 @@ export const uploadKnowledgeDoc = (avatarId: string, file: File) => {
|
|||||||
export const deleteKnowledgeDoc = (avatarId: string, docId: string) =>
|
export const deleteKnowledgeDoc = (avatarId: string, docId: string) =>
|
||||||
request.delete(`/avatar/${avatarId}/knowledge/docs/${docId}`)
|
request.delete(`/avatar/${avatarId}/knowledge/docs/${docId}`)
|
||||||
|
|
||||||
|
export const retryKnowledgeDoc = (avatarId: string, docId: string) =>
|
||||||
|
request.post<KnowledgeDoc>(`/avatar/${avatarId}/knowledge/docs/${docId}/retry`)
|
||||||
|
|
||||||
// 标准问答对列表
|
// 标准问答对列表
|
||||||
export const getQAPairs = (avatarId: string) =>
|
export const getQAPairs = (avatarId: string) =>
|
||||||
request.get<QAPair[]>(`/avatar/${avatarId}/knowledge/qa`)
|
request.get<QAPair[]>(`/avatar/${avatarId}/knowledge/qa`)
|
||||||
|
|||||||
@@ -192,7 +192,7 @@ const permissionItems: Array<{
|
|||||||
{
|
{
|
||||||
key: 'interact',
|
key: 'interact',
|
||||||
title: '广场互动操作',
|
title: '广场互动操作',
|
||||||
description: '点赞、收藏、评论、回复等操作',
|
description: '允许分身代表你点赞、收藏、评论和回复;时段、间隔与触发概率由广场调度设置统一控制',
|
||||||
tone: 'pink',
|
tone: 'pink',
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -15,7 +15,7 @@
|
|||||||
|
|
||||||
<template v-else>
|
<template v-else>
|
||||||
<div class="tab-switcher" role="tablist" aria-label="知识库类型">
|
<div class="tab-switcher" role="tablist" aria-label="知识库类型">
|
||||||
<button class="tab-btn" :class="{ active: activeTab === 'docs' }" role="tab" :aria-selected="activeTab === 'docs'" @click="activeTab = 'docs'">文档知识库 <b>{{ docs.length }}</b></button>
|
<button class="tab-btn" :class="{ active: activeTab === 'docs' }" role="tab" :aria-selected="activeTab === 'docs'" @click="activeTab = 'docs'">文档知识库 <b>{{ displayDocs.length }}</b></button>
|
||||||
<button class="tab-btn" :class="{ active: activeTab === 'qa' }" role="tab" :aria-selected="activeTab === 'qa'" @click="activeTab = 'qa'">标准问答对 <b>{{ qaPairs.length }}</b></button>
|
<button class="tab-btn" :class="{ active: activeTab === 'qa' }" role="tab" :aria-selected="activeTab === 'qa'" @click="activeTab = 'qa'">标准问答对 <b>{{ qaPairs.length }}</b></button>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
@@ -25,14 +25,14 @@
|
|||||||
<div class="upload-icon">📥</div>
|
<div class="upload-icon">📥</div>
|
||||||
<p class="upload-title"><span class="upload-link">点击上传</span></p>
|
<p class="upload-title"><span class="upload-link">点击上传</span></p>
|
||||||
<p class="upload-hint">支持 MD / TXT / PDF / DOC / DOCX / XLSX,上传后自动向量化</p>
|
<p class="upload-hint">支持 MD / TXT / PDF / DOC / DOCX / XLSX,上传后自动向量化</p>
|
||||||
<input ref="fileInput" type="file" accept=".md,.txt,.pdf,.doc,.docx,.xlsx" class="hidden-input" @change="onFileChange" />
|
<input ref="fileInput" type="file" multiple accept=".md,.txt,.pdf,.doc,.docx,.xlsx" class="hidden-input" @change="onFileChange" />
|
||||||
</div>
|
</div>
|
||||||
<p v-if="uploading" class="uploading-text">上传并向量化中…</p>
|
<p v-if="uploading" class="uploading-text">{{ pendingUploads.length }} 个文件正在上传</p>
|
||||||
<p v-if="uploadError" class="error-text">{{ uploadError }}</p>
|
<p v-if="uploadError" class="error-text">{{ uploadError }}</p>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<div v-if="docs.length" class="mobile-card-list">
|
<div v-if="displayDocs.length" class="mobile-card-list">
|
||||||
<article v-for="doc in docs" :key="doc.id" class="knowledge-card">
|
<article v-for="doc in displayDocs" :key="doc.id" class="knowledge-card document-card">
|
||||||
<div class="card-icon">{{ fileEmoji(doc.fileType) }}</div>
|
<div class="card-icon">{{ fileEmoji(doc.fileType) }}</div>
|
||||||
<div class="card-content">
|
<div class="card-content">
|
||||||
<div class="card-title-row">
|
<div class="card-title-row">
|
||||||
@@ -41,8 +41,19 @@
|
|||||||
</div>
|
</div>
|
||||||
<p class="card-meta">{{ doc.fileType.toUpperCase() }} · {{ formatSize(doc.fileSize) }} · {{ formatDate(doc.createdAt) }}</p>
|
<p class="card-meta">{{ doc.fileType.toUpperCase() }} · {{ formatSize(doc.fileSize) }} · {{ formatDate(doc.createdAt) }}</p>
|
||||||
<p class="card-detail">{{ documentState(doc).detail }}</p>
|
<p class="card-detail">{{ documentState(doc).detail }}</p>
|
||||||
|
<div v-if="documentState(doc).progress !== undefined" class="progress-track" :aria-label="`${documentState(doc).label} ${documentState(doc).progress}%`">
|
||||||
|
<span class="progress-fill" :style="{ width: `${documentState(doc).progress}%` }"></span>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
<div class="card-actions">
|
||||||
|
<button v-if="!doc.localUploading" class="card-delete" @click="removeDoc(doc.id)">{{ doc.localOnly ? '移除' : '删除' }}</button>
|
||||||
|
</div>
|
||||||
|
<div v-if="canRetryDoc(doc)" class="card-retry-area">
|
||||||
|
<span v-if="retryErrors[doc.id]" class="card-retry-error">{{ retryErrors[doc.id] }}</span>
|
||||||
|
<button class="card-retry" :disabled="retryingDocs[doc.id]" @click="retryDoc(doc)">
|
||||||
|
{{ retryingDocs[doc.id] ? '重新索引中…' : '重新索引' }}
|
||||||
|
</button>
|
||||||
</div>
|
</div>
|
||||||
<button class="card-delete" @click="removeDoc(doc.id)">删除</button>
|
|
||||||
</article>
|
</article>
|
||||||
</div>
|
</div>
|
||||||
<div v-else class="card-empty">📂 暂无文档,先上传一个知识文件</div>
|
<div v-else class="card-empty">📂 暂无文档,先上传一个知识文件</div>
|
||||||
@@ -78,7 +89,7 @@
|
|||||||
</template>
|
</template>
|
||||||
|
|
||||||
<script setup lang="ts">
|
<script setup lang="ts">
|
||||||
import { ref, onMounted, computed } from 'vue'
|
import { ref, onMounted, onUnmounted, computed } from 'vue'
|
||||||
import { useRoute, useRouter } from 'vue-router'
|
import { useRoute, useRouter } from 'vue-router'
|
||||||
import { useAvatarStore } from '@/store/avatar'
|
import { useAvatarStore } from '@/store/avatar'
|
||||||
import { pickScopedAvatarId, unwrapListData } from '@/utils/avatar-page-data.js'
|
import { pickScopedAvatarId, unwrapListData } from '@/utils/avatar-page-data.js'
|
||||||
@@ -87,6 +98,7 @@ import {
|
|||||||
getKnowledgeDocs,
|
getKnowledgeDocs,
|
||||||
uploadKnowledgeDoc,
|
uploadKnowledgeDoc,
|
||||||
deleteKnowledgeDoc,
|
deleteKnowledgeDoc,
|
||||||
|
retryKnowledgeDoc,
|
||||||
getQAPairs,
|
getQAPairs,
|
||||||
deleteQAPair,
|
deleteQAPair,
|
||||||
searchKnowledge,
|
searchKnowledge,
|
||||||
@@ -102,18 +114,30 @@ const avatarId = computed(() => pickScopedAvatarId(route.params.avatarId, store.
|
|||||||
const activeTab = ref<'docs' | 'qa'>('docs')
|
const activeTab = ref<'docs' | 'qa'>('docs')
|
||||||
|
|
||||||
const docs = ref<any[]>([])
|
const docs = ref<any[]>([])
|
||||||
|
const pendingUploads = ref<any[]>([])
|
||||||
const qaPairs = ref<any[]>([])
|
const qaPairs = ref<any[]>([])
|
||||||
const uploading = ref(false)
|
const uploading = computed(() => pendingUploads.value.some((doc) => doc.localUploading))
|
||||||
const uploadError = ref('')
|
const uploadError = ref('')
|
||||||
const dragOver = ref(false)
|
const dragOver = ref(false)
|
||||||
const fileInput = ref<HTMLInputElement | null>(null)
|
const fileInput = ref<HTMLInputElement | null>(null)
|
||||||
|
const retryingDocs = ref<Record<string, boolean>>({})
|
||||||
|
const retryErrors = ref<Record<string, string>>({})
|
||||||
|
let documentPollingTimer: ReturnType<typeof setInterval> | undefined
|
||||||
|
|
||||||
const query = ref('')
|
const query = ref('')
|
||||||
const searching = ref(false)
|
const searching = ref(false)
|
||||||
const searched = ref(false)
|
const searched = ref(false)
|
||||||
const searchResults = ref<any[]>([])
|
const searchResults = ref<any[]>([])
|
||||||
|
|
||||||
|
const displayDocs = computed(() => [...pendingUploads.value, ...docs.value])
|
||||||
|
|
||||||
const documentState = (doc: any) => {
|
const documentState = (doc: any) => {
|
||||||
|
if (doc.localUploading) {
|
||||||
|
return { tone: 'pending', label: '上传中', detail: `正在上传 ${doc.uploadProgress || 0}%`, progress: doc.uploadProgress || 0 }
|
||||||
|
}
|
||||||
|
if (doc.localOnly) {
|
||||||
|
return { tone: 'failed', label: '上传失败', detail: doc.errorMessage || '文件未上传成功,请移除后重试' }
|
||||||
|
}
|
||||||
if (doc.filePresent === false) {
|
if (doc.filePresent === false) {
|
||||||
return { tone: 'missing', label: '文件缺失', detail: '原文件不可用,请删除后重新上传' }
|
return { tone: 'missing', label: '文件缺失', detail: '原文件不可用,请删除后重新上传' }
|
||||||
}
|
}
|
||||||
@@ -121,9 +145,33 @@ const documentState = (doc: any) => {
|
|||||||
return { tone: 'ready', label: '已入库', detail: `已切分 ${doc.chunkCount || 0} 段,可用于对话` }
|
return { tone: 'ready', label: '已入库', detail: `已切分 ${doc.chunkCount || 0} 段,可用于对话` }
|
||||||
}
|
}
|
||||||
if (['uploaded', 'parsing'].includes(String(doc.status || '').toLowerCase())) {
|
if (['uploaded', 'parsing'].includes(String(doc.status || '').toLowerCase())) {
|
||||||
return { tone: 'pending', label: '处理中', detail: '正在解析并建立知识索引' }
|
const stage = String(doc.indexStage || 'queued').toLowerCase()
|
||||||
|
const labels: Record<string, string> = {
|
||||||
|
queued: '等待处理', extracting: '解析文档', chunking: '切分文本', embedding: '向量化中'
|
||||||
|
}
|
||||||
|
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: 'failed', label: '处理失败', detail: '未能建立知识索引,请删除后重新上传' }
|
return { tone: 'failed', label: '处理失败', detail: doc.errorMessage || '未能建立知识索引,请重新索引或重新上传' }
|
||||||
|
}
|
||||||
|
|
||||||
|
const hasPendingDocuments = () => docs.value.some((doc) =>
|
||||||
|
['uploaded', 'parsing'].includes(String(doc.status || '').toLowerCase())
|
||||||
|
)
|
||||||
|
|
||||||
|
const stopDocumentPolling = () => {
|
||||||
|
if (documentPollingTimer) {
|
||||||
|
clearInterval(documentPollingTimer)
|
||||||
|
documentPollingTimer = undefined
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
const startDocumentPolling = () => {
|
||||||
|
if (documentPollingTimer || !hasPendingDocuments()) return
|
||||||
|
documentPollingTimer = setInterval(async () => {
|
||||||
|
await loadDocs()
|
||||||
|
if (!hasPendingDocuments()) stopDocumentPolling()
|
||||||
|
}, 2000)
|
||||||
}
|
}
|
||||||
|
|
||||||
const loadDocs = async () => {
|
const loadDocs = async () => {
|
||||||
@@ -131,6 +179,7 @@ const loadDocs = async () => {
|
|||||||
try {
|
try {
|
||||||
const res: any = await getKnowledgeDocs(avatarId.value)
|
const res: any = await getKnowledgeDocs(avatarId.value)
|
||||||
docs.value = unwrapListData(res)
|
docs.value = unwrapListData(res)
|
||||||
|
startDocumentPolling()
|
||||||
} catch (e) {
|
} catch (e) {
|
||||||
console.error(e)
|
console.error(e)
|
||||||
}
|
}
|
||||||
@@ -149,40 +198,92 @@ const loadQA = async () => {
|
|||||||
const triggerFile = () => fileInput.value?.click()
|
const triggerFile = () => fileInput.value?.click()
|
||||||
|
|
||||||
const onFileChange = (e: Event) => {
|
const onFileChange = (e: Event) => {
|
||||||
const f = (e.target as HTMLInputElement).files?.[0]
|
const files = Array.from((e.target as HTMLInputElement).files || [])
|
||||||
if (f) doUpload(f)
|
if (files.length) uploadFiles(files)
|
||||||
;(e.target as HTMLInputElement).value = ''
|
;(e.target as HTMLInputElement).value = ''
|
||||||
}
|
}
|
||||||
|
|
||||||
const onDrop = (e: DragEvent) => {
|
const onDrop = (e: DragEvent) => {
|
||||||
dragOver.value = false
|
dragOver.value = false
|
||||||
const f = e.dataTransfer?.files?.[0]
|
const files = Array.from(e.dataTransfer?.files || [])
|
||||||
if (f) doUpload(f)
|
if (files.length) uploadFiles(files)
|
||||||
}
|
}
|
||||||
|
|
||||||
const doUpload = async (file: File) => {
|
const uploadFiles = (files: File[]) => {
|
||||||
uploadError.value = ''
|
uploadError.value = ''
|
||||||
const ext = '.' + (file.name.split('.').pop() || '').toLowerCase()
|
|
||||||
if (!['.md', '.txt', '.pdf', '.doc', '.docx', '.xlsx'].includes(ext)) {
|
|
||||||
uploadError.value = `不支持的类型:${ext},仅支持 md/txt/pdf/doc/docx/xlsx`
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if (!avatarId.value) {
|
if (!avatarId.value) {
|
||||||
uploadError.value = '请先创建数字分身'
|
uploadError.value = '请先创建数字分身'
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
uploading.value = true
|
for (const file of files) {
|
||||||
|
const ext = '.' + (file.name.split('.').pop() || '').toLowerCase()
|
||||||
|
if (!['.md', '.txt', '.pdf', '.doc', '.docx', '.xlsx'].includes(ext)) {
|
||||||
|
uploadError.value = `不支持的类型:${ext},仅支持 md/txt/pdf/doc/docx/xlsx`
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
void uploadOne(file, ext)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
const uploadOne = async (file: File, ext: string) => {
|
||||||
|
if (!avatarId.value) return
|
||||||
|
const localId = `upload-${Date.now()}-${Math.random().toString(16).slice(2)}`
|
||||||
|
const card = {
|
||||||
|
id: localId,
|
||||||
|
filename: file.name,
|
||||||
|
fileType: ext.slice(1),
|
||||||
|
fileSize: file.size,
|
||||||
|
createdAt: new Date().toISOString(),
|
||||||
|
localUploading: true,
|
||||||
|
localOnly: true,
|
||||||
|
uploadProgress: 0,
|
||||||
|
errorMessage: ''
|
||||||
|
}
|
||||||
|
pendingUploads.value.unshift(card)
|
||||||
try {
|
try {
|
||||||
await uploadKnowledgeDoc(avatarId.value, file)
|
const created: any = await uploadKnowledgeDoc(avatarId.value, file, (loaded, total) => {
|
||||||
await loadDocs()
|
const current = pendingUploads.value.find((doc) => doc.id === localId)
|
||||||
|
if (current) current.uploadProgress = Math.min(99, Math.round((loaded / Math.max(1, total)) * 100))
|
||||||
|
})
|
||||||
|
pendingUploads.value = pendingUploads.value.filter((doc) => doc.id !== localId)
|
||||||
|
docs.value = [created, ...docs.value.filter((doc) => doc.id !== created.id)]
|
||||||
|
startDocumentPolling()
|
||||||
} catch (e: any) {
|
} catch (e: any) {
|
||||||
uploadError.value = e?.message || '上传失败'
|
const current = pendingUploads.value.find((doc) => doc.id === localId)
|
||||||
|
if (current) {
|
||||||
|
current.localUploading = false
|
||||||
|
current.errorMessage = e?.message || '上传失败'
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
const canRetryDoc = (doc: any) =>
|
||||||
|
!doc.localOnly && doc.filePresent !== false && documentState(doc).tone === 'failed'
|
||||||
|
|
||||||
|
const retryDoc = async (doc: any) => {
|
||||||
|
if (!avatarId.value || !canRetryDoc(doc) || retryingDocs.value[doc.id]) return
|
||||||
|
retryingDocs.value = { ...retryingDocs.value, [doc.id]: true }
|
||||||
|
retryErrors.value = { ...retryErrors.value, [doc.id]: '' }
|
||||||
|
try {
|
||||||
|
const updated: any = await retryKnowledgeDoc(avatarId.value, doc.id)
|
||||||
|
Object.assign(doc, updated)
|
||||||
|
startDocumentPolling()
|
||||||
|
} catch (e: any) {
|
||||||
|
retryErrors.value = {
|
||||||
|
...retryErrors.value,
|
||||||
|
[doc.id]: e?.response?.data?.message || e?.response?.data?.detail || e?.message || '重新索引失败'
|
||||||
|
}
|
||||||
} finally {
|
} finally {
|
||||||
uploading.value = false
|
retryingDocs.value = { ...retryingDocs.value, [doc.id]: false }
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
const removeDoc = async (id: string) => {
|
const removeDoc = async (id: string) => {
|
||||||
|
const local = pendingUploads.value.find((doc) => doc.id === id)
|
||||||
|
if (local?.localOnly) {
|
||||||
|
pendingUploads.value = pendingUploads.value.filter((doc) => doc.id !== id)
|
||||||
|
return
|
||||||
|
}
|
||||||
if (!avatarId.value) return
|
if (!avatarId.value) return
|
||||||
await deleteKnowledgeDoc(avatarId.value, id)
|
await deleteKnowledgeDoc(avatarId.value, id)
|
||||||
await loadDocs()
|
await loadDocs()
|
||||||
@@ -257,6 +358,8 @@ onMounted(async () => {
|
|||||||
if (avatarId.value) store.currentAvatarId = avatarId.value
|
if (avatarId.value) store.currentAvatarId = avatarId.value
|
||||||
await Promise.all([loadDocs(), loadQA()])
|
await Promise.all([loadDocs(), loadQA()])
|
||||||
})
|
})
|
||||||
|
|
||||||
|
onUnmounted(stopDocumentPolling)
|
||||||
</script>
|
</script>
|
||||||
|
|
||||||
<style scoped>
|
<style scoped>
|
||||||
@@ -292,6 +395,7 @@ onMounted(async () => {
|
|||||||
.panel-heading p { margin: -5px 0 0; color: #9398AE; font-size: 12px; }
|
.panel-heading p { margin: -5px 0 0; color: #9398AE; font-size: 12px; }
|
||||||
.mobile-card-list { display: grid; grid-template-columns: minmax(0, 1fr); width: 100%; min-width: 0; gap: 10px; }
|
.mobile-card-list { display: grid; grid-template-columns: minmax(0, 1fr); width: 100%; min-width: 0; gap: 10px; }
|
||||||
.knowledge-card { display: flex; align-items: center; width: 100%; min-width: 0; box-sizing: border-box; gap: 11px; padding: 14px; background: #fff; border: 1px solid #F1E1D3; border-radius: 16px; box-shadow: 0 5px 16px rgba(112, 62, 22, .04); }
|
.knowledge-card { display: flex; align-items: center; width: 100%; min-width: 0; box-sizing: border-box; gap: 11px; padding: 14px; background: #fff; border: 1px solid #F1E1D3; border-radius: 16px; box-shadow: 0 5px 16px rgba(112, 62, 22, .04); }
|
||||||
|
.document-card { display: grid; grid-template-columns: 42px minmax(0, 1fr) auto; align-items: center; }
|
||||||
.card-icon { flex: 0 0 auto; width: 42px; height: 42px; display: grid; place-items: center; border-radius: 13px; background: #FFF3E6; font-size: 22px; }
|
.card-icon { flex: 0 0 auto; width: 42px; height: 42px; display: grid; place-items: center; border-radius: 13px; background: #FFF3E6; font-size: 22px; }
|
||||||
.card-content { min-width: 0; flex: 1; overflow: hidden; }
|
.card-content { min-width: 0; flex: 1; overflow: hidden; }
|
||||||
.card-title-row { display: flex; align-items: center; gap: 8px; min-width: 0; }
|
.card-title-row { display: flex; align-items: center; gap: 8px; min-width: 0; }
|
||||||
@@ -300,7 +404,15 @@ onMounted(async () => {
|
|||||||
.status-pill.missing { color: #B91C1C; background: #FEF2F2; }
|
.status-pill.missing { color: #B91C1C; background: #FEF2F2; }
|
||||||
.status-pill.failed { color: #B91C1C; background: #FEF2F2; }
|
.status-pill.failed { 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-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; }
|
.progress-track { width: 100%; height: 4px; margin-top: 8px; overflow: hidden; border-radius: 999px; background: #FDE7D1; }
|
||||||
|
.progress-fill { display: block; height: 100%; border-radius: inherit; background: linear-gradient(90deg, #FB923C, #F97316); transition: width .25s ease; }
|
||||||
|
.card-actions { flex: 0 0 auto; display: flex; align-items: center; }
|
||||||
|
.card-delete, .card-retry { align-self: center; border: 0; border-radius: 8px; padding: 7px 9px; font-size: 12px; cursor: pointer; white-space: nowrap; }
|
||||||
|
.card-delete { color: #EF4444; background: #FEF2F2; }
|
||||||
|
.card-retry { color: #C15F18; background: #FFF3E6; }
|
||||||
|
.card-retry:disabled { cursor: wait; opacity: .65; }
|
||||||
|
.card-retry-area { grid-column: 1 / -1; display: flex; align-items: center; justify-content: flex-end; gap: 10px; min-width: 0; }
|
||||||
|
.card-retry-error { min-width: 0; overflow: hidden; color: #DC2626; font-size: 11px; line-height: 1.35; text-overflow: ellipsis; white-space: nowrap; }
|
||||||
.card-empty { padding: 42px 16px; border: 1px dashed #F1D9C3; border-radius: 16px; color: #9398AE; background: #fff; font-size: 14px; text-align: center; }
|
.card-empty { padding: 42px 16px; border: 1px dashed #F1D9C3; border-radius: 16px; color: #9398AE; background: #fff; font-size: 14px; text-align: center; }
|
||||||
.qa-card { align-items: stretch; text-align: left; }.qa-card.qa-disabled { opacity: .58; }
|
.qa-card { align-items: stretch; text-align: left; }.qa-card.qa-disabled { opacity: .58; }
|
||||||
.qa-card .card-content,
|
.qa-card .card-content,
|
||||||
@@ -526,8 +638,11 @@ onMounted(async () => {
|
|||||||
@media (max-width: 520px) {
|
@media (max-width: 520px) {
|
||||||
.knowledge-panel { padding: 0 12px; }
|
.knowledge-panel { padding: 0 12px; }
|
||||||
.knowledge-card { display: grid; grid-template-columns: 42px minmax(0, 1fr); align-items: start; gap: 10px; padding: 13px; }
|
.knowledge-card { display: grid; grid-template-columns: 42px minmax(0, 1fr); align-items: start; gap: 10px; padding: 13px; }
|
||||||
|
.document-card { grid-template-columns: 42px minmax(0, 1fr) auto; }
|
||||||
.card-content { grid-column: 2; }
|
.card-content { grid-column: 2; }
|
||||||
.card-delete { grid-column: 2; justify-self: end; margin-top: -2px; }
|
.card-actions { grid-column: 3; grid-row: 1; }
|
||||||
|
.card-delete { justify-self: end; margin-top: -2px; }
|
||||||
|
.card-retry-area { grid-column: 1 / -1; }
|
||||||
.qa-card { display: block; }
|
.qa-card { display: block; }
|
||||||
.qa-card .card-content { width: 100%; grid-column: 1; }
|
.qa-card .card-content { width: 100%; grid-column: 1; }
|
||||||
.card-title-row { align-items: flex-start; flex-wrap: wrap; gap: 5px 7px; }
|
.card-title-row { align-items: flex-start; flex-wrap: wrap; gap: 5px 7px; }
|
||||||
|
|||||||
Reference in New Issue
Block a user