Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0fc43908ae | ||
|
|
3d999f9472 | ||
|
|
540edb58c4 | ||
|
|
a7eb6ac2a5 | ||
|
|
016bc22c05 | ||
|
|
094f8cd40f | ||
|
|
6794e88d53 | ||
|
|
7884430b3d | ||
|
|
46d42b7d98 | ||
|
|
ef58c5f2d2 | ||
|
|
6e3fe5a616 | ||
|
|
67b6bd1b48 | ||
|
|
0752001d85 | ||
|
|
0c6419f37e | ||
|
|
f768e7648f | ||
|
|
e30ab2b889 |
@@ -37,6 +37,8 @@ async def create_model(req: AIModelCreateRequest, db=Depends(get_db)):
|
||||
api_base_url=req.api_base_url,
|
||||
api_key_enc=encrypt(req.api_key) if req.api_key else None,
|
||||
model_version=req.model_version,
|
||||
vision_model_version=req.vision_model_version,
|
||||
ocr_model_version=req.ocr_model_version,
|
||||
temperature=req.temperature,
|
||||
max_tokens=req.max_tokens,
|
||||
timeout_seconds=req.timeout_seconds,
|
||||
@@ -100,6 +102,8 @@ async def get_digital_avatar_runtime_model(
|
||||
"api_base_url": model.api_base_url or "https://api.openai.com/v1",
|
||||
"api_key": decrypt(model.api_key_enc) if model.api_key_enc else "",
|
||||
"model": model.model_version or model.model_name,
|
||||
"vision_model": model.vision_model_version or "qwen3.6-flash",
|
||||
"ocr_model": model.ocr_model_version or "qwen-vl-ocr",
|
||||
"temperature": model.temperature,
|
||||
"max_tokens": model.max_tokens,
|
||||
"timeout_seconds": model.timeout_seconds,
|
||||
@@ -129,6 +133,8 @@ def _format_model(m: AIModelConfig) -> dict:
|
||||
"usage_scope": m.usage_scope,
|
||||
"api_base_url": m.api_base_url, "has_api_key": bool(m.api_key_enc),
|
||||
"model_version": m.model_version, "temperature": m.temperature,
|
||||
"vision_model_version": m.vision_model_version,
|
||||
"ocr_model_version": m.ocr_model_version,
|
||||
"max_tokens": m.max_tokens, "timeout_seconds": m.timeout_seconds,
|
||||
"is_default": m.is_default, "is_enabled": m.is_enabled,
|
||||
"created_at": m.created_at.isoformat(),
|
||||
|
||||
@@ -66,21 +66,36 @@ async def init_db():
|
||||
PendingReplyTask, TokenStat, AIModelConfig, SystemConfig, LoginLog
|
||||
)
|
||||
async with engine.begin() as conn:
|
||||
await conn.execute(text("SELECT GET_LOCK('ai_model_usage_scope_migration', 30)"))
|
||||
await conn.execute(text("SELECT GET_LOCK('ai_model_config_migration', 30)"))
|
||||
try:
|
||||
result = await conn.execute(text(
|
||||
"SELECT COUNT(*) FROM information_schema.COLUMNS "
|
||||
"WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'ai_model_configs' "
|
||||
"AND COLUMN_NAME = 'usage_scope'"
|
||||
))
|
||||
if result.scalar_one() == 0:
|
||||
await conn.execute(text(
|
||||
columns = (
|
||||
(
|
||||
"usage_scope",
|
||||
"ALTER TABLE ai_model_configs ADD COLUMN usage_scope "
|
||||
"VARCHAR(16) NOT NULL DEFAULT 'general' AFTER provider"
|
||||
))
|
||||
logger.info("AI模型配置表已增加 usage_scope 字段")
|
||||
"VARCHAR(16) NOT NULL DEFAULT 'general' AFTER provider",
|
||||
),
|
||||
(
|
||||
"vision_model_version",
|
||||
"ALTER TABLE ai_model_configs ADD COLUMN vision_model_version "
|
||||
"VARCHAR(64) NULL AFTER model_version",
|
||||
),
|
||||
(
|
||||
"ocr_model_version",
|
||||
"ALTER TABLE ai_model_configs ADD COLUMN ocr_model_version "
|
||||
"VARCHAR(64) NULL AFTER vision_model_version",
|
||||
),
|
||||
)
|
||||
for column_name, ddl in columns:
|
||||
result = await conn.execute(text(
|
||||
"SELECT COUNT(*) FROM information_schema.COLUMNS "
|
||||
"WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'ai_model_configs' "
|
||||
"AND COLUMN_NAME = :column_name"
|
||||
), {"column_name": column_name})
|
||||
if result.scalar_one() == 0:
|
||||
await conn.execute(text(ddl))
|
||||
logger.info("AI模型配置表已增加 %s 字段", column_name)
|
||||
finally:
|
||||
await conn.execute(text("SELECT RELEASE_LOCK('ai_model_usage_scope_migration')"))
|
||||
await conn.execute(text("SELECT RELEASE_LOCK('ai_model_config_migration')"))
|
||||
logger.info("✅ 数据库模型注册成功")
|
||||
logger.info("✅ 数据库初始化完成")
|
||||
|
||||
|
||||
@@ -126,6 +126,8 @@ class AIModelConfig(Base):
|
||||
api_base_url: Mapped[str | None] = mapped_column(String(256))
|
||||
api_key_enc: Mapped[str | None] = mapped_column(String(512))
|
||||
model_version: Mapped[str | None] = mapped_column(String(64))
|
||||
vision_model_version: Mapped[str | None] = mapped_column(String(64))
|
||||
ocr_model_version: Mapped[str | None] = mapped_column(String(64))
|
||||
temperature: Mapped[float] = mapped_column(Float, default=0.7)
|
||||
max_tokens: Mapped[int] = mapped_column(Integer, default=1000)
|
||||
timeout_seconds: Mapped[int] = mapped_column(Integer, default=30)
|
||||
|
||||
@@ -158,6 +158,8 @@ class AIModelCreateRequest(BaseModel):
|
||||
api_base_url: Optional[str] = None
|
||||
api_key: Optional[str] = None
|
||||
model_version: Optional[str] = None
|
||||
vision_model_version: Optional[str] = Field(None, max_length=64)
|
||||
ocr_model_version: Optional[str] = Field(None, max_length=64)
|
||||
temperature: float = Field(default=0.7, ge=0.0, le=2.0)
|
||||
max_tokens: int = Field(default=1000, ge=1, le=32000)
|
||||
timeout_seconds: int = Field(default=30, ge=5, le=300)
|
||||
@@ -171,6 +173,8 @@ class AIModelUpdateRequest(BaseModel):
|
||||
api_base_url: Optional[str] = None
|
||||
api_key: Optional[str] = None
|
||||
model_version: Optional[str] = None
|
||||
vision_model_version: Optional[str] = Field(None, max_length=64)
|
||||
ocr_model_version: Optional[str] = Field(None, max_length=64)
|
||||
temperature: Optional[float] = Field(None, ge=0.0, le=2.0)
|
||||
max_tokens: Optional[int] = Field(None, ge=1, le=32000)
|
||||
timeout_seconds: Optional[int] = Field(None, ge=5, le=300)
|
||||
@@ -186,6 +190,8 @@ class AIModelResponse(BaseModel):
|
||||
api_base_url: Optional[str]
|
||||
has_api_key: bool
|
||||
model_version: Optional[str]
|
||||
vision_model_version: Optional[str]
|
||||
ocr_model_version: Optional[str]
|
||||
temperature: float
|
||||
max_tokens: int
|
||||
timeout_seconds: int
|
||||
|
||||
@@ -38,15 +38,17 @@ def init_db():
|
||||
("avatars", "owner_id", "VARCHAR DEFAULT ''"),
|
||||
("authorizations", "takeover_enabled", "BOOLEAN DEFAULT 0"),
|
||||
("authorizations", "takeover_mode", "VARCHAR DEFAULT 'immediate'"),
|
||||
("authorizations", "takeover_delay_seconds", "INTEGER DEFAULT 30"),
|
||||
("authorizations", "takeover_delay_seconds", "INTEGER DEFAULT 180"),
|
||||
("avatars", "share_token", "VARCHAR DEFAULT NULL"),
|
||||
("token_account", "user_id", "VARCHAR DEFAULT ''"),
|
||||
("token_account", "total_granted", "BIGINT DEFAULT 0"),
|
||||
("token_account", "total_consumed", "BIGINT DEFAULT 0"),
|
||||
("token_account", "created_at", "TIMESTAMP"),
|
||||
("token_account", "updated_at", "TIMESTAMP"),
|
||||
("takeover_messages", "attachment_id", "VARCHAR DEFAULT NULL"),
|
||||
)
|
||||
_normalize_optional_unique_values()
|
||||
_normalize_takeover_delays()
|
||||
_create_token_indexes()
|
||||
|
||||
|
||||
@@ -66,6 +68,15 @@ def _normalize_optional_unique_values():
|
||||
conn.exec_driver_sql("UPDATE avatars SET share_token = NULL WHERE share_token = ''")
|
||||
|
||||
|
||||
def _normalize_takeover_delays():
|
||||
with engine.begin() as conn:
|
||||
# The old 30-second column default was never wired into the scheduler.
|
||||
conn.exec_driver_sql(
|
||||
"UPDATE authorizations SET takeover_delay_seconds = 180 "
|
||||
"WHERE takeover_delay_seconds IS NULL OR takeover_delay_seconds = 30"
|
||||
)
|
||||
|
||||
|
||||
def _create_token_indexes():
|
||||
with engine.begin() as conn:
|
||||
conn.exec_driver_sql(
|
||||
|
||||
@@ -18,6 +18,14 @@ EMBED_DIM = 256
|
||||
MODEL = os.getenv("EMBEDDING_MODEL", "mock-hash-embed-v1")
|
||||
|
||||
|
||||
def _embedding_endpoint(api_url):
|
||||
"""Accept either an OpenAI-compatible base URL or its full endpoint."""
|
||||
api_url = (api_url or "").strip().rstrip("/")
|
||||
if not api_url or api_url.endswith("/embeddings"):
|
||||
return api_url
|
||||
return f"{api_url}/embeddings"
|
||||
|
||||
|
||||
def _tokenize(text):
|
||||
text = (text or "").lower()
|
||||
# 英文/数字按词,CJK 逐字(中文无空格,需拆到字级才能命中子词)
|
||||
@@ -47,7 +55,7 @@ def embed(texts):
|
||||
"""返回 list[list[float]],与输入顺序一致。"""
|
||||
if not texts:
|
||||
return []
|
||||
api_url = os.getenv("EMBEDDING_API_URL")
|
||||
api_url = _embedding_endpoint(os.getenv("EMBEDDING_API_URL"))
|
||||
if api_url:
|
||||
api_key = os.getenv("EMBEDDING_API_KEY", "")
|
||||
model = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small")
|
||||
|
||||
@@ -19,11 +19,13 @@ import routers.huihui_auth
|
||||
import routers.chat
|
||||
import routers.takeover
|
||||
from responses import ok
|
||||
from services.chat_attachment_service import purge_expired_chat_attachments
|
||||
from services.token_billing import DEFAULT_TOKEN_GRANT, release_stale_reservations
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
takeover_scheduler = None
|
||||
maintenance_scheduler = None
|
||||
|
||||
app = FastAPI(title="会会数字分身 API", version="1.0.0")
|
||||
|
||||
@@ -132,6 +134,15 @@ def on_startup():
|
||||
|
||||
# Release stale resources when startup is invoked again by a reload/test.
|
||||
stop_takeover_scheduler()
|
||||
stop_maintenance_scheduler()
|
||||
try:
|
||||
start_maintenance_scheduler()
|
||||
except Exception as exc:
|
||||
stop_maintenance_scheduler()
|
||||
logger.warning(
|
||||
"Failed to initialize chat attachment cleanup, app will continue: %s",
|
||||
exc,
|
||||
)
|
||||
|
||||
# --- Takeover scheduler ---
|
||||
try:
|
||||
@@ -152,7 +163,14 @@ def on_startup():
|
||||
boxim_client = BoxIMClient(boxim_config)
|
||||
|
||||
from services.takeover_service import TakeoverService
|
||||
takeover_service = TakeoverService(SessionLocal, boxim_client)
|
||||
takeover_service = TakeoverService(
|
||||
SessionLocal,
|
||||
boxim_client,
|
||||
poll_concurrency=int(os.getenv("BOXIM_POLL_CONCURRENCY", "8")),
|
||||
max_message_age_seconds=int(
|
||||
os.getenv("BOXIM_MAX_MESSAGE_AGE_SECONDS", "600")
|
||||
),
|
||||
)
|
||||
|
||||
poll_interval = max(0.5, float(os.getenv("BOXIM_POLL_INTERVAL_SECONDS", "1")))
|
||||
takeover_scheduler = AsyncIOScheduler()
|
||||
@@ -196,6 +214,51 @@ def stop_takeover_scheduler():
|
||||
finally:
|
||||
takeover_scheduler = None
|
||||
|
||||
|
||||
def purge_expired_chat_attachments_job():
|
||||
db = SessionLocal()
|
||||
try:
|
||||
count = purge_expired_chat_attachments(db)
|
||||
if count:
|
||||
logger.info("Purged %s expired chat image attachment(s)", count)
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
logger.warning("Failed to purge expired chat image attachments: %s", exc)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def start_maintenance_scheduler():
|
||||
global maintenance_scheduler
|
||||
|
||||
purge_expired_chat_attachments_job()
|
||||
interval_minutes = max(
|
||||
5, min(1440, int(os.getenv("CHAT_ATTACHMENT_CLEANUP_MINUTES", "60")))
|
||||
)
|
||||
maintenance_scheduler = AsyncIOScheduler()
|
||||
maintenance_scheduler.add_job(
|
||||
purge_expired_chat_attachments_job,
|
||||
trigger=IntervalTrigger(minutes=interval_minutes),
|
||||
id="chat_attachment_cleanup",
|
||||
max_instances=1,
|
||||
coalesce=True,
|
||||
)
|
||||
maintenance_scheduler.start()
|
||||
|
||||
|
||||
def stop_maintenance_scheduler():
|
||||
global maintenance_scheduler
|
||||
|
||||
if maintenance_scheduler is not None:
|
||||
try:
|
||||
if maintenance_scheduler.running:
|
||||
maintenance_scheduler.shutdown(wait=False)
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to stop maintenance scheduler cleanly: %s", exc)
|
||||
finally:
|
||||
maintenance_scheduler = None
|
||||
|
||||
@app.on_event("shutdown")
|
||||
def on_shutdown():
|
||||
stop_takeover_scheduler()
|
||||
stop_maintenance_scheduler()
|
||||
|
||||
@@ -67,7 +67,7 @@ class Authorization(Base):
|
||||
status = Column(String, default="active") # active | inactive
|
||||
takeover_enabled = Column(Boolean, default=False) # 是否开启分身接管
|
||||
takeover_mode = Column(String, default="immediate") # immediate | delayed
|
||||
takeover_delay_seconds = Column(Integer, default=30) # 延迟秒数
|
||||
takeover_delay_seconds = Column(Integer, default=180) # 延迟秒数,默认 3 分钟
|
||||
created_at = Column(DateTime, server_default=func.now())
|
||||
|
||||
def to_dict(self):
|
||||
@@ -120,13 +120,14 @@ class TakeoverMessage(Base):
|
||||
direction = Column(String, nullable=False) # incoming | outgoing
|
||||
message_type = Column(Integer, default=0)
|
||||
content = Column(Text, default="")
|
||||
attachment_id = Column(String, nullable=True)
|
||||
is_avatar = Column(Boolean, default=False)
|
||||
send_time = Column(DateTime, nullable=False)
|
||||
created_at = Column(DateTime, server_default=func.now())
|
||||
|
||||
|
||||
class TakeoverReplyTask(Base):
|
||||
"""Restart-safe three-second BOXIM reply task."""
|
||||
"""Restart-safe delayed BOXIM reply task."""
|
||||
|
||||
__tablename__ = "takeover_reply_tasks"
|
||||
__table_args__ = (
|
||||
@@ -188,7 +189,7 @@ class KnowledgeDoc(Base):
|
||||
file_type = Column(String, default="") # pdf | doc | docx | xlsx
|
||||
file_size = Column(Integer, default=0)
|
||||
file_url = Column(String, default="")
|
||||
status = Column(String, default="uploaded") # uploaded | parsing | ready
|
||||
status = Column(String, default="uploaded") # uploaded | parsing | ready | failed
|
||||
vectorized = Column(Boolean, default=False) # 是否已向量化
|
||||
embedding_model = Column(String, default="") # 向量模型标识
|
||||
chunk_count = Column(Integer, default=0) # 切片数量
|
||||
@@ -257,6 +258,44 @@ class KnowledgeChunk(Base):
|
||||
}
|
||||
|
||||
|
||||
class ChatAttachment(Base):
|
||||
"""Private, avatar-scoped result of one chat image analysis."""
|
||||
|
||||
__tablename__ = "chat_attachments"
|
||||
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
|
||||
avatar_id = Column(String, nullable=False, default="", index=True)
|
||||
uploader_kind = Column(String, default="owner") # owner | public | boxim
|
||||
filename = Column(String, default="")
|
||||
mime_type = Column(String, default="")
|
||||
file_size = Column(Integer, default=0)
|
||||
status = Column(String, default="processing") # processing | ready | failed
|
||||
category = Column(String, default="general_image")
|
||||
summary = Column(Text, default="")
|
||||
extracted_text = Column(Text, default="")
|
||||
structured_data = Column(JSON, default=dict)
|
||||
warning = Column(Text, default="")
|
||||
vision_model = Column(String, default="")
|
||||
ocr_model = Column(String, default="")
|
||||
used_at = Column(DateTime)
|
||||
expires_at = Column(DateTime, nullable=False)
|
||||
created_at = Column(DateTime, server_default=func.now())
|
||||
|
||||
def to_dict(self):
|
||||
return {
|
||||
"id": self.id,
|
||||
"avatarId": self.avatar_id,
|
||||
"filename": self.filename,
|
||||
"mimeType": self.mime_type,
|
||||
"fileSize": self.file_size,
|
||||
"status": self.status,
|
||||
"category": self.category,
|
||||
"summary": self.summary,
|
||||
"warning": self.warning,
|
||||
"expiresAt": _iso(self.expires_at),
|
||||
"createdAt": _iso(self.created_at),
|
||||
}
|
||||
|
||||
|
||||
class TokenAccount(Base):
|
||||
__tablename__ = "token_account"
|
||||
id = Column(Integer, primary_key=True)
|
||||
|
||||
@@ -8,3 +8,4 @@ pypdf
|
||||
python-docx
|
||||
openpyxl
|
||||
apscheduler>=3.10
|
||||
Pillow>=10.4
|
||||
|
||||
@@ -2,7 +2,7 @@ from fastapi import APIRouter, Body, Depends, Header, HTTPException
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from database import get_db
|
||||
from models import Authorization, TakeoverCursor, TakeoverReplyTask
|
||||
from models import Authorization, Avatar, TakeoverCursor, TakeoverReplyTask
|
||||
from responses import fail, ok
|
||||
from routers.avatars import _require_owned_avatar
|
||||
|
||||
@@ -14,6 +14,10 @@ ALLOWED_PERMISSIONS = set(PERMISSION_ORDER)
|
||||
AVATAR_PERMISSION_ORDER = PERMISSION_ORDER
|
||||
AVATAR_PERMISSION_KEY = "authorizationPermissions"
|
||||
DEFAULT_AVATAR_PERMISSIONS = ["friend", "chat"]
|
||||
TAKEOVER_DELAY_KEY = "takeoverReplyDelaySeconds"
|
||||
DEFAULT_TAKEOVER_DELAY_SECONDS = 180
|
||||
MIN_TAKEOVER_DELAY_SECONDS = 3
|
||||
MAX_TAKEOVER_DELAY_SECONDS = 86_400
|
||||
LEGACY_PERMISSION_MAP = {
|
||||
"read": "browse",
|
||||
"reply": "chat",
|
||||
@@ -91,9 +95,62 @@ def _permission_settings_payload(avatar) -> dict:
|
||||
return {
|
||||
"avatarId": avatar.id,
|
||||
"permissions": _stored_avatar_permissions(avatar),
|
||||
"takeoverReplyDelaySeconds": _stored_takeover_delay(avatar),
|
||||
}
|
||||
|
||||
|
||||
def _stored_takeover_delay(avatar) -> int:
|
||||
raw = (avatar.config or {}).get(TAKEOVER_DELAY_KEY, DEFAULT_TAKEOVER_DELAY_SECONDS)
|
||||
if isinstance(raw, bool):
|
||||
return DEFAULT_TAKEOVER_DELAY_SECONDS
|
||||
try:
|
||||
delay = int(raw)
|
||||
except (TypeError, ValueError):
|
||||
return DEFAULT_TAKEOVER_DELAY_SECONDS
|
||||
if not MIN_TAKEOVER_DELAY_SECONDS <= delay <= MAX_TAKEOVER_DELAY_SECONDS:
|
||||
return DEFAULT_TAKEOVER_DELAY_SECONDS
|
||||
return delay
|
||||
|
||||
|
||||
def _validate_takeover_delay(value) -> int:
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise ValueError("自动回复等待时间必须是整数秒")
|
||||
if not MIN_TAKEOVER_DELAY_SECONDS <= value <= MAX_TAKEOVER_DELAY_SECONDS:
|
||||
raise ValueError("自动回复等待时间需在 3 秒到 24 小时之间")
|
||||
return value
|
||||
|
||||
|
||||
def _disable_other_takeovers(db: Session, avatar) -> list[str]:
|
||||
disabled_ids = []
|
||||
others = (
|
||||
db.query(Avatar)
|
||||
.filter(Avatar.owner_id == avatar.owner_id, Avatar.id != avatar.id)
|
||||
.all()
|
||||
)
|
||||
for other in others:
|
||||
permissions = _stored_avatar_permissions(other)
|
||||
if "takeover" not in permissions:
|
||||
continue
|
||||
other.config = {
|
||||
**(other.config or {}),
|
||||
AVATAR_PERMISSION_KEY: [item for item in permissions if item != "takeover"],
|
||||
}
|
||||
disabled_ids.append(other.id)
|
||||
tasks = (
|
||||
db.query(TakeoverReplyTask)
|
||||
.filter(
|
||||
TakeoverReplyTask.avatar_id == other.id,
|
||||
TakeoverReplyTask.status.in_(("pending", "generating", "ready", "sending")),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
for task in tasks:
|
||||
task.status = "cancelled"
|
||||
task.cancel_reason = "another_avatar_takeover_enabled"
|
||||
task.locked_at = None
|
||||
return disabled_ids
|
||||
|
||||
|
||||
def _require_authorization(db: Session, avatar_id: str, authorization_id: str) -> Authorization:
|
||||
authorization = (
|
||||
db.query(Authorization)
|
||||
@@ -144,10 +201,19 @@ def update_permission_settings(
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
avatar = _require_owned_avatar(db, avatar_id, authorization)
|
||||
if "permissions" not in payload:
|
||||
return fail("缺少 permissions", 400)
|
||||
if "permissions" not in payload and TAKEOVER_DELAY_KEY not in payload:
|
||||
return fail("缺少授权设置", 400)
|
||||
try:
|
||||
permissions = _normalize_avatar_permissions(payload["permissions"])
|
||||
permissions = (
|
||||
_normalize_avatar_permissions(payload["permissions"])
|
||||
if "permissions" in payload
|
||||
else _stored_avatar_permissions(avatar)
|
||||
)
|
||||
takeover_delay = (
|
||||
_validate_takeover_delay(payload[TAKEOVER_DELAY_KEY])
|
||||
if TAKEOVER_DELAY_KEY in payload
|
||||
else _stored_takeover_delay(avatar)
|
||||
)
|
||||
except ValueError as exc:
|
||||
return fail(str(exc), 400)
|
||||
|
||||
@@ -155,7 +221,9 @@ def update_permission_settings(
|
||||
avatar.config = {
|
||||
**(avatar.config or {}),
|
||||
AVATAR_PERMISSION_KEY: permissions,
|
||||
TAKEOVER_DELAY_KEY: takeover_delay,
|
||||
}
|
||||
disabled_avatar_ids = _disable_other_takeovers(db, avatar) if "takeover" in permissions else []
|
||||
cursor = db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar.id).first()
|
||||
if cursor and "takeover" in permissions and "takeover" not in previous_permissions:
|
||||
cursor.initialized = False
|
||||
@@ -179,7 +247,9 @@ def update_permission_settings(
|
||||
task.locked_at = None
|
||||
db.commit()
|
||||
db.refresh(avatar)
|
||||
return ok(_permission_settings_payload(avatar), "授权设置已保存")
|
||||
response = _permission_settings_payload(avatar)
|
||||
response["disabledAvatarIds"] = disabled_avatar_ids
|
||||
return ok(response, "授权设置已保存")
|
||||
|
||||
|
||||
@router.get("/avatar/{avatar_id}/authorizations")
|
||||
@@ -240,7 +310,7 @@ def create_auth(
|
||||
status="active",
|
||||
takeover_enabled=False,
|
||||
takeover_mode="immediate",
|
||||
takeover_delay_seconds=30,
|
||||
takeover_delay_seconds=DEFAULT_TAKEOVER_DELAY_SECONDS,
|
||||
)
|
||||
db.add(item)
|
||||
db.commit()
|
||||
|
||||
@@ -6,7 +6,17 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from database import get_db
|
||||
from routers.knowledge import UPLOAD_DIR
|
||||
from models import Avatar, KnowledgeDoc, KnowledgeChunk, QAPair, Authorization, User
|
||||
from models import (
|
||||
Authorization,
|
||||
Avatar,
|
||||
KnowledgeChunk,
|
||||
KnowledgeDoc,
|
||||
QAPair,
|
||||
TakeoverCursor,
|
||||
TakeoverMessage,
|
||||
TakeoverReplyTask,
|
||||
User,
|
||||
)
|
||||
from responses import ok, fail
|
||||
|
||||
router = APIRouter(tags=["分身"])
|
||||
@@ -74,18 +84,21 @@ def list_avatars(page: int = 1, limit: int = 20, authorization: str = Header(Non
|
||||
|
||||
|
||||
@router.get("/avatar/{avatar_id}")
|
||||
def get_avatar(avatar_id: str, db: Session = Depends(get_db)):
|
||||
a = db.query(Avatar).filter(Avatar.id == avatar_id).first()
|
||||
if not a:
|
||||
return fail("分身不存在", 404)
|
||||
return ok(a.to_dict())
|
||||
def get_avatar(
|
||||
avatar_id: str,
|
||||
authorization: str = Header(None),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
return ok(_require_owned_avatar(db, avatar_id, authorization).to_dict())
|
||||
|
||||
|
||||
@router.post("/avatar")
|
||||
def create_avatar(payload: dict = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)):
|
||||
user = _resolve_user(authorization, db)
|
||||
if not user:
|
||||
raise HTTPException(status_code=401, detail="未登录")
|
||||
a = Avatar(
|
||||
owner_id=user.huihui_user_id if user else "",
|
||||
owner_id=user.huihui_user_id,
|
||||
name=payload.get("name", "未命名分身"),
|
||||
display_name=payload.get("displayName", "") or payload.get("display_name", ""),
|
||||
description=payload.get("description", ""),
|
||||
@@ -102,10 +115,13 @@ def create_avatar(payload: dict = Body(...), authorization: str = Header(None),
|
||||
|
||||
|
||||
@router.put("/avatar/{avatar_id}")
|
||||
def update_avatar(avatar_id: str, payload: dict = Body(...), db: Session = Depends(get_db)):
|
||||
a = db.query(Avatar).filter(Avatar.id == avatar_id).first()
|
||||
if not a:
|
||||
return fail("分身不存在", 404)
|
||||
def update_avatar(
|
||||
avatar_id: str,
|
||||
payload: dict = Body(...),
|
||||
authorization: str = Header(None),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
a = _require_owned_avatar(db, avatar_id, authorization)
|
||||
mapping = {
|
||||
"displayName": "display_name",
|
||||
"photoUrl": "photo_url",
|
||||
@@ -114,22 +130,32 @@ def update_avatar(avatar_id: str, payload: dict = Body(...), db: Session = Depen
|
||||
for key in ("name", "displayName", "description", "photoUrl", "emoji", "status", "tokenBalance", "config"):
|
||||
if key in payload:
|
||||
col = mapping.get(key, key)
|
||||
setattr(a, col, payload[key])
|
||||
value = payload[key]
|
||||
if key == "config":
|
||||
if not isinstance(value, dict):
|
||||
return fail("分身配置格式不正确", 400)
|
||||
value = {**(a.config or {}), **value}
|
||||
setattr(a, col, value)
|
||||
db.commit()
|
||||
db.refresh(a)
|
||||
return ok(a.to_dict())
|
||||
|
||||
|
||||
@router.delete("/avatar/{avatar_id}")
|
||||
def delete_avatar(avatar_id: str, db: Session = Depends(get_db)):
|
||||
a = db.query(Avatar).filter(Avatar.id == avatar_id).first()
|
||||
if not a:
|
||||
return fail("分身不存在", 404)
|
||||
def delete_avatar(
|
||||
avatar_id: str,
|
||||
authorization: str = Header(None),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
a = _require_owned_avatar(db, avatar_id, authorization)
|
||||
# 级联清理关联数据,避免孤儿记录
|
||||
db.query(KnowledgeDoc).filter(KnowledgeDoc.avatar_id == avatar_id).delete()
|
||||
db.query(KnowledgeChunk).filter(KnowledgeChunk.avatar_id == avatar_id).delete()
|
||||
db.query(QAPair).filter(QAPair.avatar_id == avatar_id).delete()
|
||||
db.query(Authorization).filter(Authorization.avatar_id == avatar_id).delete()
|
||||
db.query(TakeoverReplyTask).filter(TakeoverReplyTask.avatar_id == avatar_id).delete()
|
||||
db.query(TakeoverMessage).filter(TakeoverMessage.avatar_id == avatar_id).delete()
|
||||
db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar_id).delete()
|
||||
db.delete(a)
|
||||
db.commit()
|
||||
return ok({"success": True})
|
||||
|
||||
@@ -1,21 +1,32 @@
|
||||
import difflib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import secrets
|
||||
import string
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any, Callable
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Body, Depends, Header, HTTPException
|
||||
from fastapi import APIRouter, Body, Depends, File, Header, HTTPException, UploadFile
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
import embeddings
|
||||
from database import get_db
|
||||
from models import Avatar, KnowledgeChunk, KnowledgeDoc, QAPair, User
|
||||
from models import Avatar, ChatAttachment, KnowledgeChunk, KnowledgeDoc, QAPair, User
|
||||
from responses import ok, fail
|
||||
from services.vision_service import (
|
||||
GENERAL_VISION_PROMPT,
|
||||
MEDICAL_OCR_PROMPT,
|
||||
ImageValidationError,
|
||||
build_attachment_warning,
|
||||
call_vision_model,
|
||||
parse_vision_analysis,
|
||||
prepare_image,
|
||||
)
|
||||
from services.token_billing import (
|
||||
InsufficientTokensError,
|
||||
estimate_fallback_usage,
|
||||
@@ -24,8 +35,10 @@ from services.token_billing import (
|
||||
settle_reservation,
|
||||
)
|
||||
from services.chat_model_config import ChatModelConfig, get_chat_model_config
|
||||
from services.chat_attachment_service import purge_expired_chat_attachments
|
||||
|
||||
router = APIRouter(tags=["数字分身聊天"])
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
MAX_MESSAGE_LENGTH = 4000
|
||||
MAX_HISTORY_MESSAGES = 10
|
||||
@@ -34,16 +47,50 @@ QA_SEMANTIC_THRESHOLD = 0.72
|
||||
QA_MATCH_MARGIN = 0.06
|
||||
KNOWLEDGE_MIN_SCORE = float(os.getenv("KNOWLEDGE_MIN_SCORE", "0.42"))
|
||||
|
||||
_WRITING_SYSTEM_PATTERNS = {
|
||||
"han": re.compile(r"[\u3400-\u4dbf\u4e00-\u9fff]"),
|
||||
"latin": re.compile(r"[A-Za-z\u00c0-\u024f]"),
|
||||
"cyrillic": re.compile(r"[\u0400-\u052f]"),
|
||||
"arabic": re.compile(r"[\u0600-\u06ff]"),
|
||||
"hebrew": re.compile(r"[\u0590-\u05ff]"),
|
||||
"devanagari": re.compile(r"[\u0900-\u097f]"),
|
||||
"thai": re.compile(r"[\u0e00-\u0e7f]"),
|
||||
"greek": re.compile(r"[\u0370-\u03ff]"),
|
||||
}
|
||||
_JAPANESE_KANA = re.compile(r"[\u3040-\u30ff]")
|
||||
_KOREAN_HANGUL = re.compile(r"[\uac00-\ud7af\u1100-\u11ff]")
|
||||
|
||||
|
||||
class ChatMessage(BaseModel):
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
|
||||
role: str = Field(pattern="^(user|assistant)$")
|
||||
content: str = Field(min_length=1, max_length=MAX_MESSAGE_LENGTH)
|
||||
attachment_ids: list[str] = Field(
|
||||
default_factory=list,
|
||||
alias="attachmentIds",
|
||||
max_length=3,
|
||||
)
|
||||
|
||||
|
||||
class ChatIn(BaseModel):
|
||||
message: str = Field(min_length=1, max_length=MAX_MESSAGE_LENGTH)
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
|
||||
message: str = Field(default="", max_length=MAX_MESSAGE_LENGTH)
|
||||
attachment_ids: list[str] = Field(
|
||||
default_factory=list,
|
||||
alias="attachmentIds",
|
||||
max_length=3,
|
||||
)
|
||||
history: list[ChatMessage] = Field(default_factory=list, max_length=MAX_HISTORY_MESSAGES)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def require_message_or_image(self):
|
||||
self.message = self.message.strip()
|
||||
if not self.message and not self.attachment_ids:
|
||||
raise ValueError("请输入消息或选择图片")
|
||||
return self
|
||||
|
||||
|
||||
def _resolve_user(authorization: str | None, db: Session):
|
||||
if not authorization:
|
||||
@@ -64,12 +111,285 @@ def _require_owned_avatar(db: Session, avatar_id: str, authorization: str | None
|
||||
return avatar
|
||||
|
||||
|
||||
def _attachment_expiry() -> datetime:
|
||||
retention_hours = max(
|
||||
1, min(168, int(os.getenv("CHAT_ATTACHMENT_RETENTION_HOURS", "24")))
|
||||
)
|
||||
return datetime.utcnow() + timedelta(hours=retention_hours)
|
||||
|
||||
|
||||
def _chat_attachment_ids(body: ChatIn) -> list[str]:
|
||||
values = list(body.attachment_ids)
|
||||
for message in body.history[-MAX_HISTORY_MESSAGES:]:
|
||||
values.extend(message.attachment_ids)
|
||||
unique = list(dict.fromkeys(str(value).strip() for value in values if str(value).strip()))
|
||||
if len(unique) > 3:
|
||||
raise HTTPException(status_code=400, detail="一次会话最多引用 3 张图片")
|
||||
return unique
|
||||
|
||||
|
||||
def _load_chat_attachments(db: Session, avatar_id: str, body: ChatIn) -> list[ChatAttachment]:
|
||||
attachment_ids = _chat_attachment_ids(body)
|
||||
if not attachment_ids:
|
||||
return []
|
||||
purge_expired_chat_attachments(db)
|
||||
rows = db.query(ChatAttachment).filter(
|
||||
ChatAttachment.avatar_id == avatar_id,
|
||||
ChatAttachment.id.in_(attachment_ids),
|
||||
).all()
|
||||
by_id = {row.id: row for row in rows}
|
||||
if len(by_id) != len(attachment_ids):
|
||||
raise HTTPException(status_code=400, detail="图片资料不存在、已过期或不属于当前分身")
|
||||
ordered = [by_id[attachment_id] for attachment_id in attachment_ids]
|
||||
if any(row.status != "ready" for row in ordered):
|
||||
raise HTTPException(status_code=409, detail="图片尚未识别完成,请稍后重试")
|
||||
now = datetime.utcnow()
|
||||
for row in ordered:
|
||||
row.used_at = now
|
||||
db.commit()
|
||||
return ordered
|
||||
|
||||
|
||||
def _attachment_contexts(rows: list[ChatAttachment]) -> list[dict]:
|
||||
contexts = []
|
||||
remaining_text = 12000
|
||||
for row in rows:
|
||||
extracted = (row.extracted_text or "")[:remaining_text]
|
||||
remaining_text = max(0, remaining_text - len(extracted))
|
||||
contexts.append({
|
||||
"id": row.id,
|
||||
"filename": row.filename,
|
||||
"category": row.category,
|
||||
"summary": row.summary,
|
||||
"extractedText": extracted,
|
||||
"structuredData": row.structured_data or {},
|
||||
"warning": row.warning,
|
||||
})
|
||||
return contexts
|
||||
|
||||
|
||||
def _image_retrieval_question(question: str, image_contexts: list[dict]) -> str:
|
||||
parts = [question.strip()]
|
||||
for context in image_contexts:
|
||||
parts.extend([
|
||||
str(context.get("summary") or "")[:600],
|
||||
str(context.get("extractedText") or "")[:1200],
|
||||
])
|
||||
return "\n".join(part for part in parts if part).strip()
|
||||
|
||||
|
||||
def _run_billed_vision_call(
|
||||
db: Session,
|
||||
avatar: Avatar,
|
||||
prepared,
|
||||
*,
|
||||
model: str,
|
||||
prompt: str,
|
||||
source: str,
|
||||
json_output: bool,
|
||||
model_config: ChatModelConfig,
|
||||
) -> dict:
|
||||
estimate_messages = [{
|
||||
"role": "user",
|
||||
"content": f"[一张待识别图片]\n{prompt}",
|
||||
}]
|
||||
reservation = reserve_avatar_tokens(
|
||||
db,
|
||||
avatar,
|
||||
source,
|
||||
model,
|
||||
estimate_messages,
|
||||
model_config.vision_max_tokens,
|
||||
minimum_reserve_tokens=max(
|
||||
1000, int(os.getenv("VISION_TOKEN_RESERVE", "12000"))
|
||||
),
|
||||
)
|
||||
try:
|
||||
result = call_vision_model(
|
||||
prepared,
|
||||
model_config,
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
json_output=json_output,
|
||||
)
|
||||
settle_reservation(
|
||||
db,
|
||||
reservation,
|
||||
result.get("usage"),
|
||||
fallback_total=estimate_fallback_usage(
|
||||
estimate_messages, result.get("content") or ""
|
||||
),
|
||||
)
|
||||
return result
|
||||
except Exception as exc:
|
||||
release_reservation(db, reservation, str(exc))
|
||||
raise
|
||||
|
||||
|
||||
async def _analyze_uploaded_image(
|
||||
db: Session,
|
||||
avatar: Avatar,
|
||||
file: UploadFile,
|
||||
*,
|
||||
uploader_kind: str,
|
||||
) -> ChatAttachment:
|
||||
max_bytes = max(1024, int(os.getenv("CHAT_IMAGE_MAX_BYTES", str(8 * 1024 * 1024))))
|
||||
content = await file.read(max_bytes + 1)
|
||||
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(
|
||||
avatar_id=avatar.id,
|
||||
uploader_kind=uploader_kind,
|
||||
filename=filename,
|
||||
mime_type=(mime_type or "")[:100],
|
||||
file_size=len(content),
|
||||
status="processing",
|
||||
expires_at=_attachment_expiry(),
|
||||
)
|
||||
db.add(attachment)
|
||||
db.commit()
|
||||
db.refresh(attachment)
|
||||
|
||||
try:
|
||||
prepared = prepare_image(content)
|
||||
model_config = get_chat_model_config()
|
||||
vision_result = _run_billed_vision_call(
|
||||
db,
|
||||
avatar,
|
||||
prepared,
|
||||
model=model_config.vision_model,
|
||||
prompt=GENERAL_VISION_PROMPT,
|
||||
source="vision_image",
|
||||
json_output=True,
|
||||
model_config=model_config,
|
||||
)
|
||||
analysis = parse_vision_analysis(vision_result["content"])
|
||||
extracted_text = analysis.get("visible_text") or ""
|
||||
ocr_model = ""
|
||||
ocr_failed = False
|
||||
if analysis["category"] == "medical_document" and model_config.ocr_model:
|
||||
try:
|
||||
ocr_result = _run_billed_vision_call(
|
||||
db,
|
||||
avatar,
|
||||
prepared,
|
||||
model=model_config.ocr_model,
|
||||
prompt=MEDICAL_OCR_PROMPT,
|
||||
source="vision_medical_ocr",
|
||||
json_output=False,
|
||||
model_config=model_config,
|
||||
)
|
||||
extracted_text = ocr_result["content"]
|
||||
ocr_model = model_config.ocr_model
|
||||
except (RuntimeError, InsufficientTokensError):
|
||||
ocr_failed = True
|
||||
logger.warning(
|
||||
"medical OCR degraded for attachment %s avatar %s",
|
||||
attachment.id,
|
||||
avatar.id,
|
||||
)
|
||||
|
||||
attachment.mime_type = prepared.mime_type
|
||||
attachment.status = "ready"
|
||||
attachment.category = analysis["category"]
|
||||
attachment.summary = analysis.get("summary") or "图片内容已识别"
|
||||
attachment.extracted_text = extracted_text
|
||||
attachment.structured_data = analysis
|
||||
attachment.warning = build_attachment_warning(analysis, ocr_failed=ocr_failed)
|
||||
attachment.vision_model = model_config.vision_model
|
||||
attachment.ocr_model = ocr_model
|
||||
db.commit()
|
||||
db.refresh(attachment)
|
||||
logger.info(
|
||||
"chat image ready attachment=%s avatar=%s category=%s model=%s ocr=%s",
|
||||
attachment.id,
|
||||
avatar.id,
|
||||
attachment.category,
|
||||
attachment.vision_model,
|
||||
bool(attachment.ocr_model),
|
||||
)
|
||||
return attachment
|
||||
except ImageValidationError as exc:
|
||||
attachment.status = "failed"
|
||||
attachment.warning = str(exc)
|
||||
db.commit()
|
||||
raise
|
||||
except InsufficientTokensError:
|
||||
attachment.status = "failed"
|
||||
attachment.warning = "积分余额不足"
|
||||
db.commit()
|
||||
raise
|
||||
except RuntimeError as exc:
|
||||
attachment.status = "failed"
|
||||
attachment.warning = str(exc)
|
||||
db.commit()
|
||||
logger.warning(
|
||||
"chat image failed attachment=%s avatar=%s error=%s",
|
||||
attachment.id,
|
||||
avatar.id,
|
||||
type(exc).__name__,
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
def _normalize_question(value: str) -> str:
|
||||
value = (value or "").strip().lower()
|
||||
value = re.sub(r"\s+", "", value)
|
||||
return value.translate(str.maketrans("", "", string.punctuation + ",。!?;:、()【】「」‘’“”《》"))
|
||||
|
||||
|
||||
def _dominant_writing_system(value: str) -> str:
|
||||
value = value or ""
|
||||
if _JAPANESE_KANA.search(value):
|
||||
return "japanese"
|
||||
if _KOREAN_HANGUL.search(value):
|
||||
return "korean"
|
||||
counts = {
|
||||
name: len(pattern.findall(value))
|
||||
for name, pattern in _WRITING_SYSTEM_PATTERNS.items()
|
||||
}
|
||||
name, count = max(counts.items(), key=lambda item: item[1])
|
||||
return name if count else "unknown"
|
||||
|
||||
|
||||
def _qa_requires_language_adaptation(question: str, answer: str) -> bool:
|
||||
question_system = _dominant_writing_system(question)
|
||||
answer_system = _dominant_writing_system(answer)
|
||||
return (
|
||||
question_system != "unknown"
|
||||
and answer_system != "unknown"
|
||||
and question_system != answer_system
|
||||
)
|
||||
|
||||
|
||||
def _canonicalize_question(value: str) -> str:
|
||||
value = _normalize_question(value)
|
||||
replacements = (
|
||||
@@ -189,7 +509,15 @@ def _config(avatar: Avatar) -> dict:
|
||||
}
|
||||
|
||||
|
||||
def _build_prompt(avatar: Avatar, history: list[Any], question: str, knowledge_hits: list[dict]) -> list[dict]:
|
||||
def _build_prompt(
|
||||
avatar: Avatar,
|
||||
history: list[Any],
|
||||
question: str,
|
||||
knowledge_hits: list[dict],
|
||||
*,
|
||||
standard_answer: str = "",
|
||||
image_contexts: list[dict] | None = None,
|
||||
) -> list[dict]:
|
||||
config = _config(avatar)
|
||||
description = (getattr(avatar, "description", "") or "").strip()
|
||||
knowledge = "\n".join(
|
||||
@@ -197,6 +525,7 @@ def _build_prompt(avatar: Avatar, history: list[Any], question: str, knowledge_h
|
||||
for hit in knowledge_hits
|
||||
if hit.get("snippet")
|
||||
)
|
||||
image_contexts = image_contexts or []
|
||||
profile_items = [
|
||||
(label, config[key])
|
||||
for label, key in (
|
||||
@@ -210,7 +539,7 @@ def _build_prompt(avatar: Avatar, history: list[Any], question: str, knowledge_h
|
||||
profile = ";".join(f"{label}:{value}" for label, value in profile_items)
|
||||
system = (
|
||||
f"你的专业或服务范围是:「{description or '未设置'}」。"
|
||||
"请基于已提供的知识库回答,不要编造事实;"
|
||||
"请基于已提供的可靠资料回答,不要编造事实;"
|
||||
f"回复风格:{config['replyStyle']};严谨度:{config['rigor']}/100;"
|
||||
f"幽默感:{config['humor']}/100;回复长度:{config['responseLength']}。"
|
||||
)
|
||||
@@ -221,17 +550,47 @@ def _build_prompt(avatar: Avatar, history: list[Any], question: str, knowledge_h
|
||||
)
|
||||
if config["systemPrompt"]:
|
||||
system += f"\n额外系统提示词:{config['systemPrompt']}"
|
||||
if knowledge:
|
||||
if image_contexts:
|
||||
image_material = json.dumps(image_contexts, ensure_ascii=False, default=str)
|
||||
system += (
|
||||
"\n以下是当前会话图片经过视觉识别后得到的资料:\n"
|
||||
f"{image_material}"
|
||||
"\n图片资料可能包含 OCR 错字、模糊内容或用户尚未确认的信息,只能按可见内容谨慎表达。"
|
||||
"标准答题对中的事实优先级高于图片资料,知识库事实优先级高于模型推测;发生冲突时遵循更高优先级资料,"
|
||||
"并自然提醒对方核对原图。不得声称看到了图片中不存在的内容。"
|
||||
)
|
||||
if any(
|
||||
context.get("category") in {"medical_document", "medical_image"}
|
||||
for context in image_contexts
|
||||
):
|
||||
system += (
|
||||
"\n本次包含医疗资料。可以整理病例原文、解释指标含义和提示需要关注的异常,但不能仅凭图片作出"
|
||||
"确定诊断、疾病分期、处方、停药或治疗决定。医学影像只能客观描述,并提醒结合正规报告和医生意见。"
|
||||
"回答结尾用与用户相同的语言简短说明图片识别结果仅供辅助,不能替代医生诊断。"
|
||||
)
|
||||
if standard_answer:
|
||||
system += (
|
||||
f"\n以下是本次问题命中的已确认标准答案:\n{standard_answer.strip()}"
|
||||
"\n必须保持标准答案中的事实、数字、专有名词和结论不变,只允许为匹配用户当前语言进行忠实转换"
|
||||
"和必要的自然表达,不得补充、删减或改写其含义。不要提及标准答案或转换过程。"
|
||||
)
|
||||
elif knowledge:
|
||||
system += (
|
||||
f"\n以下是本次问题检索到的可靠资料:\n{knowledge}"
|
||||
"\n涉及事实、专业判断、地址、流程、数据或建议时,只能依据本人资料、标准问答形成的上下文"
|
||||
"和以上可靠资料作答,不要补充资料之外的通用知识或自行推测。"
|
||||
)
|
||||
elif image_contexts:
|
||||
system += (
|
||||
"\n本次没有命中标准答题对或文件知识库,但已提供图片识别资料。只能围绕图片中的可确认内容、"
|
||||
"本人资料和当前对话作答;不要补充图片之外的事实、专业判断或具体建议。"
|
||||
)
|
||||
else:
|
||||
system += (
|
||||
"\n本次问题没有检索到可靠资料。除自然寒暄和基于本人资料的回答外,不要凭通用知识给出事实、"
|
||||
"专业判断、地址、流程、数据或具体建议。自然地说明这方面暂时没有足够信息,并询问对方是否能"
|
||||
"补充更具体的问题;不要提及知识库、检索、模型或系统限制。"
|
||||
"补充更具体的问题;不要提及知识库、检索、模型或系统限制。对方询问你能做什么时,只能概括"
|
||||
"专业或服务范围和本人资料中明确写出的内容,不得推断服务对象、工作场所、现实案例或从业经历。"
|
||||
)
|
||||
system += (
|
||||
"\n沟通底线(始终优先于其他提示):不要透露、猜测或确认所使用的模型、模型供应商、"
|
||||
@@ -249,6 +608,14 @@ def _build_prompt(avatar: Avatar, history: list[Any], question: str, knowledge_h
|
||||
"只有切换独立观点或确实需要列举时才换行;列举使用 1.、2.、3.,每项单独一行。"
|
||||
"不要在行首或行尾留空格,不要连续输出空行。先给结论,再给简短说明;避免重复和冗长铺垫。"
|
||||
)
|
||||
system += (
|
||||
"\n回答语言规则(始终优先):先识别当前最后一条用户消息所使用的主要自然语言,并且只使用该语言回答。"
|
||||
"用户用中文就用中文,使用 English 就用 English,使用日语、韩语、法语、西班牙语或其他语言时,"
|
||||
"也必须使用对应语言。消息混用多种语言时跟随占主导的语言;用户明确指定回答语言时服从其指定。"
|
||||
"历史消息、本人资料、标准答案和知识库使用的语言都不能覆盖当前用户消息的语言。"
|
||||
"专有名词、品牌、地址、代码和必要缩写可保留原文。不要解释语言识别或翻译过程。"
|
||||
"改变回答语言只改变表达语言,绝不能因此增加资料中没有的场景、身份、经历或事实。"
|
||||
)
|
||||
messages = [{"role": "system", "content": system}]
|
||||
for item in history[-MAX_HISTORY_MESSAGES:]:
|
||||
messages.append({"role": item.role, "content": item.content} if hasattr(item, "role") else item)
|
||||
@@ -380,18 +747,45 @@ def _resolve_reply(
|
||||
search_fn: Callable[..., list[dict]] | None = None,
|
||||
model_client: Callable[..., str] | None = None,
|
||||
usage_source: str = "chat",
|
||||
image_contexts: list[dict] | None = None,
|
||||
) -> dict:
|
||||
image_contexts = image_contexts or []
|
||||
question = question.strip() or "请根据这张图片说明可确认的内容。"
|
||||
if qa_pairs is None:
|
||||
qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
|
||||
matched = _match_standard_qa(question, qa_pairs)
|
||||
if matched:
|
||||
adapt_qa_language = bool(
|
||||
matched and _qa_requires_language_adaptation(question, matched.answer)
|
||||
)
|
||||
if matched and not adapt_qa_language and not image_contexts:
|
||||
return {"answer": matched.answer, "source": "qa", "references": []}
|
||||
|
||||
search_fn = search_fn or (lambda query, avatar_id: _search_knowledge(db, avatar_id, query))
|
||||
hits = search_fn(question, avatar.id)
|
||||
messages = _build_prompt(avatar, history, question, hits)
|
||||
if matched:
|
||||
hits = []
|
||||
messages = _build_prompt(
|
||||
avatar,
|
||||
history,
|
||||
question,
|
||||
hits,
|
||||
standard_answer=matched.answer,
|
||||
image_contexts=image_contexts,
|
||||
)
|
||||
else:
|
||||
search_fn = search_fn or (lambda query, avatar_id: _search_knowledge(db, avatar_id, query))
|
||||
retrieval_question = _image_retrieval_question(question, image_contexts)
|
||||
hits = search_fn(retrieval_question, avatar.id)
|
||||
messages = _build_prompt(
|
||||
avatar,
|
||||
history,
|
||||
question,
|
||||
hits,
|
||||
image_contexts=image_contexts,
|
||||
)
|
||||
config = _config(avatar)
|
||||
temperature = min(0.45 if hits else 0.25, 0.2 + config["creativity"] / 100 * 0.6)
|
||||
temperature = 0.0 if matched else min(
|
||||
0.45 if hits else 0.25,
|
||||
0.2 + config["creativity"] / 100 * 0.6,
|
||||
)
|
||||
token_usage = None
|
||||
if model_client is not None:
|
||||
answer = model_client(messages=messages, temperature=temperature)
|
||||
@@ -423,7 +817,9 @@ def _resolve_reply(
|
||||
raise
|
||||
result = {
|
||||
"answer": answer,
|
||||
"source": "knowledge" if hits else "qwen",
|
||||
"source": "qa" if matched else (
|
||||
"knowledge" if hits else ("vision" if image_contexts else "qwen")
|
||||
),
|
||||
"references": hits,
|
||||
}
|
||||
if token_usage:
|
||||
@@ -439,17 +835,48 @@ def _stream_reply(
|
||||
*,
|
||||
public: bool = False,
|
||||
usage_source: str = "chat_stream",
|
||||
image_contexts: list[dict] | None = None,
|
||||
):
|
||||
image_contexts = image_contexts or []
|
||||
question = question.strip() or "请根据这张图片说明可确认的内容。"
|
||||
qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
|
||||
matched = _match_standard_qa(question, qa_pairs)
|
||||
if matched:
|
||||
adapt_qa_language = bool(
|
||||
matched and _qa_requires_language_adaptation(question, matched.answer)
|
||||
)
|
||||
messages, reservation = [], None
|
||||
if matched and not adapt_qa_language and not image_contexts:
|
||||
source, references, chunks = "qa", [], _iter_text_chunks(matched.answer)
|
||||
else:
|
||||
references = _search_knowledge(db, avatar.id, question)
|
||||
source = "knowledge" if references else "qwen"
|
||||
if matched:
|
||||
references = []
|
||||
source = "qa"
|
||||
messages = _build_prompt(
|
||||
avatar,
|
||||
history,
|
||||
question,
|
||||
references,
|
||||
standard_answer=matched.answer,
|
||||
image_contexts=image_contexts,
|
||||
)
|
||||
else:
|
||||
retrieval_question = _image_retrieval_question(question, image_contexts)
|
||||
references = _search_knowledge(db, avatar.id, retrieval_question)
|
||||
source = "knowledge" if references else (
|
||||
"vision" if image_contexts else "qwen"
|
||||
)
|
||||
messages = _build_prompt(
|
||||
avatar,
|
||||
history,
|
||||
question,
|
||||
references,
|
||||
image_contexts=image_contexts,
|
||||
)
|
||||
config = _config(avatar)
|
||||
temperature = min(0.45 if references else 0.25, 0.2 + config["creativity"] / 100 * 0.6)
|
||||
messages = _build_prompt(avatar, history, question, references)
|
||||
temperature = 0.0 if matched else min(
|
||||
0.45 if references else 0.25,
|
||||
0.2 + config["creativity"] / 100 * 0.6,
|
||||
)
|
||||
model_config = get_chat_model_config()
|
||||
reservation = reserve_avatar_tokens(
|
||||
db,
|
||||
@@ -460,8 +887,6 @@ def _stream_reply(
|
||||
model_config.max_tokens,
|
||||
)
|
||||
chunks = _iter_qwen_stream(messages, temperature, model_config)
|
||||
if matched:
|
||||
messages, reservation = [], None
|
||||
if public:
|
||||
source, references = "public", []
|
||||
|
||||
@@ -550,11 +975,62 @@ def get_shared_avatar(share_token: str, db: Session = Depends(get_db)):
|
||||
return ok(_public_avatar_payload(_require_shared_avatar(db, share_token)))
|
||||
|
||||
|
||||
@router.post("/avatar/{avatar_id}/chat/images")
|
||||
async def upload_chat_image(
|
||||
avatar_id: str,
|
||||
file: UploadFile = File(...),
|
||||
authorization: str = Header(None),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
avatar = _require_owned_avatar(db, avatar_id, authorization)
|
||||
purge_expired_chat_attachments(db)
|
||||
try:
|
||||
attachment = await _analyze_uploaded_image(
|
||||
db,
|
||||
avatar,
|
||||
file,
|
||||
uploader_kind="owner",
|
||||
)
|
||||
return ok(attachment.to_dict())
|
||||
except InsufficientTokensError as exc:
|
||||
raise HTTPException(status_code=402, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/public/avatar/{share_token}/chat/images")
|
||||
async def upload_public_chat_image(
|
||||
share_token: str,
|
||||
file: UploadFile = File(...),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
avatar = _require_shared_avatar(db, share_token)
|
||||
purge_expired_chat_attachments(db)
|
||||
try:
|
||||
attachment = await _analyze_uploaded_image(
|
||||
db,
|
||||
avatar,
|
||||
file,
|
||||
uploader_kind="public",
|
||||
)
|
||||
return ok(attachment.to_dict())
|
||||
except InsufficientTokensError as exc:
|
||||
raise HTTPException(status_code=402, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/public/avatar/{share_token}/chat")
|
||||
def public_chat(share_token: str, body: ChatIn = Body(...), db: Session = Depends(get_db)):
|
||||
avatar = _require_shared_avatar(db, share_token)
|
||||
image_contexts = _attachment_contexts(
|
||||
_load_chat_attachments(db, avatar.id, body)
|
||||
)
|
||||
try:
|
||||
result = _resolve_reply(db, avatar, body.message, body.history, usage_source="public_chat")
|
||||
result = _resolve_reply(
|
||||
db,
|
||||
avatar,
|
||||
body.message,
|
||||
body.history,
|
||||
usage_source="public_chat",
|
||||
image_contexts=image_contexts,
|
||||
)
|
||||
# 公开访客无需获知知识文件名、检索分数或内部答复来源。
|
||||
result["references"] = []
|
||||
result["source"] = "public"
|
||||
@@ -569,8 +1045,17 @@ def public_chat(share_token: str, body: ChatIn = Body(...), db: Session = Depend
|
||||
@router.post("/avatar/{avatar_id}/chat")
|
||||
def chat(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)):
|
||||
avatar = _require_owned_avatar(db, avatar_id, authorization)
|
||||
image_contexts = _attachment_contexts(
|
||||
_load_chat_attachments(db, avatar.id, body)
|
||||
)
|
||||
try:
|
||||
return ok(_resolve_reply(db, avatar, body.message, body.history))
|
||||
return ok(_resolve_reply(
|
||||
db,
|
||||
avatar,
|
||||
body.message,
|
||||
body.history,
|
||||
image_contexts=image_contexts,
|
||||
))
|
||||
except InsufficientTokensError as exc:
|
||||
return fail(str(exc), code=402)
|
||||
except RuntimeError as exc:
|
||||
@@ -580,7 +1065,17 @@ def chat(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(N
|
||||
@router.post("/avatar/{avatar_id}/chat/stream")
|
||||
def chat_stream(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)):
|
||||
try:
|
||||
return _stream_reply(db, _require_owned_avatar(db, avatar_id, authorization), body.message, body.history)
|
||||
avatar = _require_owned_avatar(db, avatar_id, authorization)
|
||||
image_contexts = _attachment_contexts(
|
||||
_load_chat_attachments(db, avatar.id, body)
|
||||
)
|
||||
return _stream_reply(
|
||||
db,
|
||||
avatar,
|
||||
body.message,
|
||||
body.history,
|
||||
image_contexts=image_contexts,
|
||||
)
|
||||
except InsufficientTokensError as exc:
|
||||
raise HTTPException(status_code=402, detail=str(exc)) from exc
|
||||
|
||||
@@ -588,13 +1083,18 @@ def chat_stream(avatar_id: str, body: ChatIn = Body(...), authorization: str = H
|
||||
@router.post("/public/avatar/{share_token}/chat/stream")
|
||||
def public_chat_stream(share_token: str, body: ChatIn = Body(...), db: Session = Depends(get_db)):
|
||||
try:
|
||||
avatar = _require_shared_avatar(db, share_token)
|
||||
image_contexts = _attachment_contexts(
|
||||
_load_chat_attachments(db, avatar.id, body)
|
||||
)
|
||||
return _stream_reply(
|
||||
db,
|
||||
_require_shared_avatar(db, share_token),
|
||||
avatar,
|
||||
body.message,
|
||||
body.history,
|
||||
public=True,
|
||||
usage_source="public_chat_stream",
|
||||
image_contexts=image_contexts,
|
||||
)
|
||||
except InsufficientTokensError as exc:
|
||||
raise HTTPException(status_code=402, detail=str(exc)) from exc
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import os
|
||||
import json
|
||||
import logging
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
|
||||
@@ -13,6 +14,7 @@ from responses import ok, fail
|
||||
import embeddings
|
||||
|
||||
router = APIRouter()
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||
UPLOAD_DIR = os.path.abspath(os.getenv("UPLOAD_DIR", os.path.join(BASE_DIR, "uploads")))
|
||||
@@ -69,6 +71,15 @@ def list_docs(avatar_id: str, authorization: str = Header(None), db: Session = D
|
||||
.order_by(KnowledgeDoc.created_at.desc())
|
||||
.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])
|
||||
|
||||
|
||||
@@ -88,6 +99,7 @@ async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization
|
||||
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("."),
|
||||
@@ -95,39 +107,47 @@ async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization
|
||||
file_url=f"/api/files/{avatar_id}/{stored}",
|
||||
status="parsing",
|
||||
)
|
||||
db.add(doc)
|
||||
db.commit()
|
||||
db.refresh(doc)
|
||||
|
||||
# 向量化:抽取文本 -> 分块 -> 调第三方/本地嵌入 -> 存切片
|
||||
# Complete extraction and embedding before the first database commit so a
|
||||
# process restart cannot leave a permanent "parsing" row behind.
|
||||
try:
|
||||
text = embeddings.extract_text(path, ext)
|
||||
chunks = embeddings.chunk_text(text)
|
||||
if chunks:
|
||||
vectors = embeddings.embed(chunks)
|
||||
for i, (c, v) in enumerate(zip(chunks, vectors)):
|
||||
db.add(
|
||||
KnowledgeChunk(
|
||||
doc_id=doc.id,
|
||||
avatar_id=avatar_id,
|
||||
content=c,
|
||||
vector=json.dumps(v),
|
||||
chunk_index=i,
|
||||
embedding_model=embeddings.MODEL,
|
||||
)
|
||||
)
|
||||
doc.vectorized = True
|
||||
doc.embedding_model = embeddings.MODEL
|
||||
doc.chunk_count = len(chunks)
|
||||
doc.vectorized_at = datetime.now(timezone.utc)
|
||||
if not chunks:
|
||||
raise ValueError("文档没有可建立索引的文字内容")
|
||||
vectors = embeddings.embed(chunks)
|
||||
if len(vectors) != len(chunks):
|
||||
raise ValueError("向量服务返回数量与文档分段不一致")
|
||||
doc.vectorized = True
|
||||
doc.embedding_model = embeddings.MODEL
|
||||
doc.chunk_count = len(chunks)
|
||||
doc.vectorized_at = datetime.now(timezone.utc)
|
||||
doc.status = "ready"
|
||||
db.add(doc)
|
||||
for i, (chunk, vector) in enumerate(zip(chunks, vectors)):
|
||||
db.add(
|
||||
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 e:
|
||||
print("vectorize failed:", e)
|
||||
doc.status = "ready" # 上传成功但向量化失败,仍可展示
|
||||
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)
|
||||
|
||||
return ok(_doc_payload(doc))
|
||||
|
||||
|
||||
@@ -8,13 +8,25 @@ from sqlalchemy.orm import Session
|
||||
from database import get_db
|
||||
from models import TakeoverCursor, TakeoverReplyTask, User
|
||||
from responses import fail, ok
|
||||
from routers.authorizations import _require_authorization
|
||||
from routers.authorizations import (
|
||||
DEFAULT_TAKEOVER_DELAY_SECONDS,
|
||||
MAX_TAKEOVER_DELAY_SECONDS,
|
||||
MIN_TAKEOVER_DELAY_SECONDS,
|
||||
_require_authorization,
|
||||
_stored_takeover_delay,
|
||||
)
|
||||
from routers.avatars import _require_owned_avatar
|
||||
|
||||
router = APIRouter(tags=["分身接管"])
|
||||
BOXIM_STATUS_FRESH_SECONDS = 60
|
||||
|
||||
|
||||
def _delay_label(seconds: int) -> str:
|
||||
if seconds % 60 == 0:
|
||||
return f"{seconds // 60} 分钟"
|
||||
return f"{seconds} 秒"
|
||||
|
||||
|
||||
@router.get("/avatar/{avatar_id}/takeover/status")
|
||||
def get_takeover_status(
|
||||
avatar_id: str,
|
||||
@@ -24,6 +36,7 @@ def get_takeover_status(
|
||||
avatar = _require_owned_avatar(db, avatar_id, authorization)
|
||||
permissions = (avatar.config or {}).get("authorizationPermissions", [])
|
||||
enabled = isinstance(permissions, list) and "takeover" in permissions
|
||||
reply_delay_seconds = _stored_takeover_delay(avatar)
|
||||
user = db.query(User).filter(User.huihui_user_id == avatar.owner_id).first()
|
||||
cursor = db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar.id).first()
|
||||
pending_count = (
|
||||
@@ -49,7 +62,10 @@ def get_takeover_status(
|
||||
and cursor.last_polled_at
|
||||
>= datetime.utcnow() - timedelta(seconds=BOXIM_STATUS_FRESH_SECONDS)
|
||||
):
|
||||
status, message = "ready", "BOXIM 已连接,收到私聊消息 3 秒后自动回复"
|
||||
status, message = (
|
||||
"ready",
|
||||
f"BOXIM 已连接,收到私聊消息 {_delay_label(reply_delay_seconds)}后自动回复",
|
||||
)
|
||||
else:
|
||||
status, message = "connecting", "正在连接 BOXIM"
|
||||
|
||||
@@ -59,6 +75,7 @@ def get_takeover_status(
|
||||
"status": status,
|
||||
"message": message,
|
||||
"pendingCount": pending_count,
|
||||
"takeoverReplyDelaySeconds": reply_delay_seconds,
|
||||
"lastPolledAt": cursor.last_polled_at.isoformat() if cursor and cursor.last_polled_at else None,
|
||||
}
|
||||
)
|
||||
@@ -91,7 +108,7 @@ def update_takeover_config(
|
||||
auth = _require_authorization(db, avatar_id, str(auth_id))
|
||||
enabled = bool(auth.takeover_enabled)
|
||||
mode = auth.takeover_mode or "immediate"
|
||||
delay = auth.takeover_delay_seconds or 30
|
||||
delay = auth.takeover_delay_seconds or DEFAULT_TAKEOVER_DELAY_SECONDS
|
||||
|
||||
if _has(payload, "takeoverEnabled", "takeover_enabled"):
|
||||
raw_enabled = _read(payload, "takeoverEnabled", "takeover_enabled")
|
||||
@@ -106,8 +123,12 @@ def update_takeover_config(
|
||||
|
||||
if _has(payload, "takeoverDelaySeconds", "takeover_delay_seconds"):
|
||||
delay = _read(payload, "takeoverDelaySeconds", "takeover_delay_seconds")
|
||||
if isinstance(delay, bool) or not isinstance(delay, int) or not 5 <= delay <= 3600:
|
||||
return fail("延迟时间需在 5 到 3600 秒之间", 400)
|
||||
if (
|
||||
isinstance(delay, bool)
|
||||
or not isinstance(delay, int)
|
||||
or not MIN_TAKEOVER_DELAY_SECONDS <= delay <= MAX_TAKEOVER_DELAY_SECONDS
|
||||
):
|
||||
return fail("延迟时间需在 3 秒到 24 小时之间", 400)
|
||||
|
||||
if enabled and auth.target_type != "user":
|
||||
return fail("本期仅支持对会会用户开启单聊接管", 400)
|
||||
|
||||
@@ -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,20 @@
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from models import ChatAttachment
|
||||
|
||||
|
||||
def purge_expired_chat_attachments(
|
||||
db: Session,
|
||||
*,
|
||||
now: datetime | None = None,
|
||||
) -> int:
|
||||
"""Remove expired derived image data; raw image bytes are never persisted."""
|
||||
count = db.query(ChatAttachment).filter(
|
||||
ChatAttachment.expires_at < (now or datetime.utcnow())
|
||||
).delete(synchronize_session=False)
|
||||
if count:
|
||||
db.commit()
|
||||
db.expire_all()
|
||||
return count
|
||||
@@ -16,6 +16,10 @@ class ChatModelConfig:
|
||||
model: str
|
||||
max_tokens: int
|
||||
timeout_seconds: float
|
||||
vision_model: str
|
||||
ocr_model: str
|
||||
vision_max_tokens: int
|
||||
vision_timeout_seconds: float
|
||||
source: str
|
||||
|
||||
|
||||
@@ -33,6 +37,10 @@ def _environment_config() -> ChatModelConfig:
|
||||
model=os.getenv("CHAT_MODEL", "qwen-plus"),
|
||||
max_tokens=max(128, int(os.getenv("CHAT_MAX_OUTPUT_TOKENS", "1024"))),
|
||||
timeout_seconds=max(5.0, float(os.getenv("CHAT_TIMEOUT_SECONDS", "30"))),
|
||||
vision_model=os.getenv("VISION_MODEL", "qwen3.6-flash"),
|
||||
ocr_model=os.getenv("VISION_OCR_MODEL", "qwen-vl-ocr"),
|
||||
vision_max_tokens=max(256, int(os.getenv("VISION_MAX_OUTPUT_TOKENS", "2048"))),
|
||||
vision_timeout_seconds=max(10.0, float(os.getenv("VISION_TIMEOUT_SECONDS", "90"))),
|
||||
source="environment",
|
||||
)
|
||||
|
||||
@@ -60,6 +68,20 @@ def _fetch_runtime_config() -> ChatModelConfig | None:
|
||||
model=model,
|
||||
max_tokens=max(128, int(payload.get("max_tokens") or 1024)),
|
||||
timeout_seconds=max(5.0, float(payload.get("timeout_seconds") or 30)),
|
||||
vision_model=str(
|
||||
payload.get("vision_model")
|
||||
or os.getenv("VISION_MODEL", "qwen3.6-flash")
|
||||
),
|
||||
ocr_model=str(
|
||||
payload.get("ocr_model")
|
||||
or os.getenv("VISION_OCR_MODEL", "qwen-vl-ocr")
|
||||
),
|
||||
vision_max_tokens=max(
|
||||
256, int(os.getenv("VISION_MAX_OUTPUT_TOKENS", "2048"))
|
||||
),
|
||||
vision_timeout_seconds=max(
|
||||
10.0, float(os.getenv("VISION_TIMEOUT_SECONDS", "90"))
|
||||
),
|
||||
source="admin",
|
||||
)
|
||||
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
import asyncio
|
||||
import hashlib
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import secrets
|
||||
import time
|
||||
@@ -13,21 +14,41 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from models import (
|
||||
Avatar,
|
||||
ChatAttachment,
|
||||
TakeoverCursor,
|
||||
TakeoverMessage,
|
||||
TakeoverReplyTask,
|
||||
User,
|
||||
)
|
||||
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__)
|
||||
|
||||
ACTIVE_TASK_STATUSES = ("pending", "generating", "ready", "sending")
|
||||
GENERATABLE_TASK_STATUSES = ("pending",)
|
||||
MAX_PROMPT_LENGTH = 4000
|
||||
MAX_STALE_SECONDS = 120
|
||||
DEFAULT_MAX_MESSAGE_AGE_SECONDS = 600
|
||||
MAX_SEND_OVERDUE_SECONDS = 120
|
||||
STUCK_LOCK_SECONDS = 90
|
||||
TAKEOVER_PERMISSION = "takeover"
|
||||
TAKEOVER_DELAY_KEY = "takeoverReplyDelaySeconds"
|
||||
DEFAULT_REPLY_DELAY_SECONDS = 180
|
||||
MIN_REPLY_DELAY_SECONDS = 3
|
||||
MAX_REPLY_DELAY_SECONDS = 86_400
|
||||
HUMAN_PAUSE_SECONDS = 600
|
||||
RATE_LIMIT_WINDOW_SECONDS = 300
|
||||
RATE_LIMIT_MAX_REPLIES = 5
|
||||
AVATAR_LOCAL_ID_PREFIX = "880"
|
||||
BOXIM_TEXT_MESSAGE_TYPE = 0
|
||||
BOXIM_IMAGE_MESSAGE_TYPE = 1
|
||||
BOXIM_IMAGE_PROMPT = "请看看这张图片。"
|
||||
BOXIM_IMAGE_UNAVAILABLE_REPLY = "这张图片我暂时没看清,麻烦重新发送一张清晰的原图。"
|
||||
|
||||
|
||||
def _utcnow() -> datetime:
|
||||
@@ -70,20 +91,60 @@ def _plain_text_reply(value: str) -> str:
|
||||
return "\n".join(line for line in lines if line).strip()
|
||||
|
||||
|
||||
def _avatar_local_id(owner_id: str, trigger_message_id: str) -> str:
|
||||
"""Build a deterministic BOXIM idempotency key that also marks avatar traffic."""
|
||||
digest = hashlib.sha256(f"{owner_id}:{trigger_message_id}".encode("utf-8")).digest()
|
||||
suffix = int.from_bytes(digest[:8], "big") % (10**15)
|
||||
return f"{AVATAR_LOCAL_ID_PREFIX}{suffix:015d}"
|
||||
|
||||
|
||||
def _is_avatar_local_id(value: str | None) -> bool:
|
||||
local_id = str(value or "").strip()
|
||||
return len(local_id) == 18 and local_id.isdigit() and local_id.startswith(AVATAR_LOCAL_ID_PREFIX)
|
||||
|
||||
|
||||
def _configured_reply_delay(avatar: Avatar, fallback: int | None = None) -> int:
|
||||
raw = (avatar.config or {}).get(
|
||||
TAKEOVER_DELAY_KEY,
|
||||
fallback if fallback is not None else DEFAULT_REPLY_DELAY_SECONDS,
|
||||
)
|
||||
if isinstance(raw, bool):
|
||||
return DEFAULT_REPLY_DELAY_SECONDS
|
||||
try:
|
||||
delay = int(raw)
|
||||
except (TypeError, ValueError):
|
||||
return DEFAULT_REPLY_DELAY_SECONDS
|
||||
if not MIN_REPLY_DELAY_SECONDS <= delay <= MAX_REPLY_DELAY_SECONDS:
|
||||
return DEFAULT_REPLY_DELAY_SECONDS
|
||||
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 ""
|
||||
|
||||
|
||||
class TakeoverService:
|
||||
"""Poll BOXIM, prepare replies during the grace period, then send at +3s."""
|
||||
"""Poll BOXIM, honor the owner grace period, then generate and send one reply."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
session_factory: Callable[[], Session],
|
||||
boxim_client: BoxIMClient,
|
||||
*,
|
||||
reply_delay_seconds: int = 3,
|
||||
reply_delay_seconds: int | None = None,
|
||||
poll_concurrency: int = 8,
|
||||
max_message_age_seconds: int = DEFAULT_MAX_MESSAGE_AGE_SECONDS,
|
||||
now: Callable[[], datetime] = _utcnow,
|
||||
):
|
||||
self.session_factory = session_factory
|
||||
self.boxim = boxim_client
|
||||
self.reply_delay_seconds = reply_delay_seconds
|
||||
self.poll_concurrency = max(1, min(int(poll_concurrency), 64))
|
||||
self.max_message_age_seconds = max(60, int(max_message_age_seconds))
|
||||
self.now = now
|
||||
self._sessions: dict[str, dict] = {}
|
||||
self._poll_lock = asyncio.Lock()
|
||||
@@ -102,8 +163,48 @@ class TakeoverService:
|
||||
self._recover_stuck_tasks()
|
||||
avatar_ids = self._enabled_avatar_ids()
|
||||
self._cancel_disabled_tasks(set(avatar_ids))
|
||||
for avatar_id in avatar_ids:
|
||||
await self._sync_avatar(avatar_id)
|
||||
self._ensure_takeover_cursors(avatar_ids)
|
||||
semaphore = asyncio.Semaphore(self.poll_concurrency)
|
||||
|
||||
async def sync(avatar_id: str):
|
||||
async with semaphore:
|
||||
return await self._sync_avatar(avatar_id)
|
||||
|
||||
results = await asyncio.gather(
|
||||
*(sync(avatar_id) for avatar_id in avatar_ids),
|
||||
return_exceptions=True,
|
||||
)
|
||||
for avatar_id, result in zip(avatar_ids, results):
|
||||
if isinstance(result, Exception):
|
||||
logger.warning("BOXIM poll crashed for avatar %s: %s", avatar_id, result)
|
||||
|
||||
def _ensure_takeover_cursors(self, avatar_ids: list[str]):
|
||||
"""Create durable cursors before concurrent network polling starts."""
|
||||
if not avatar_ids:
|
||||
return
|
||||
db = self.session_factory()
|
||||
try:
|
||||
existing = {
|
||||
row[0]
|
||||
for row in db.query(TakeoverCursor.avatar_id)
|
||||
.filter(TakeoverCursor.avatar_id.in_(avatar_ids))
|
||||
.all()
|
||||
}
|
||||
avatars = (
|
||||
db.query(Avatar.id, Avatar.owner_id)
|
||||
.filter(
|
||||
Avatar.id.in_(
|
||||
[avatar_id for avatar_id in avatar_ids if avatar_id not in existing]
|
||||
)
|
||||
)
|
||||
.all()
|
||||
)
|
||||
for avatar_id, owner_id in avatars:
|
||||
db.add(TakeoverCursor(avatar_id=avatar_id, owner_id=owner_id))
|
||||
if avatars:
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
async def process_reply_tasks(self):
|
||||
"""Generate and send replies independently from BOXIM's long poll."""
|
||||
@@ -119,11 +220,17 @@ class TakeoverService:
|
||||
def _enabled_avatar_ids(self) -> list[str]:
|
||||
db = self.session_factory()
|
||||
try:
|
||||
return [
|
||||
avatar.id
|
||||
for avatar in db.query(Avatar).filter(Avatar.status == "active").all()
|
||||
if _takeover_enabled(avatar)
|
||||
]
|
||||
avatars = (
|
||||
db.query(Avatar)
|
||||
.filter(Avatar.status == "active")
|
||||
.order_by(Avatar.updated_at.desc(), Avatar.created_at.desc())
|
||||
.all()
|
||||
)
|
||||
selected = {}
|
||||
for avatar in avatars:
|
||||
if _takeover_enabled(avatar) and avatar.owner_id not in selected:
|
||||
selected[avatar.owner_id] = avatar.id
|
||||
return list(selected.values())
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@@ -201,13 +308,20 @@ class TakeoverService:
|
||||
def _forget_boxim_session(self, user_id: str):
|
||||
self._sessions.pop(user_id, None)
|
||||
|
||||
def _disable_after_connection_failure(
|
||||
def _record_connection_failure(
|
||||
self,
|
||||
db: Session,
|
||||
avatar: Avatar,
|
||||
cursor: TakeoverCursor,
|
||||
message: str,
|
||||
*,
|
||||
disable_takeover: bool,
|
||||
):
|
||||
cursor.last_error = message
|
||||
cursor.last_polled_at = self.now()
|
||||
if not disable_takeover:
|
||||
return
|
||||
|
||||
permissions = (avatar.config or {}).get("authorizationPermissions", [])
|
||||
avatar.config = {
|
||||
**(avatar.config or {}),
|
||||
@@ -217,8 +331,6 @@ class TakeoverService:
|
||||
if permission != TAKEOVER_PERMISSION
|
||||
],
|
||||
}
|
||||
cursor.last_error = message
|
||||
cursor.last_polled_at = self.now()
|
||||
tasks = (
|
||||
db.query(TakeoverReplyTask)
|
||||
.filter(
|
||||
@@ -243,13 +355,17 @@ class TakeoverService:
|
||||
if not cursor:
|
||||
cursor = TakeoverCursor(avatar_id=avatar.id, owner_id=avatar.owner_id)
|
||||
db.add(cursor)
|
||||
db.flush()
|
||||
db.commit()
|
||||
else:
|
||||
# Release SQLite's read transaction before the long network poll.
|
||||
db.commit()
|
||||
if not user or not user.huihui_token:
|
||||
self._disable_after_connection_failure(
|
||||
self._record_connection_failure(
|
||||
db,
|
||||
avatar,
|
||||
cursor,
|
||||
"请重新登录会会生产账号后再开启主动接管",
|
||||
disable_takeover=True,
|
||||
)
|
||||
db.commit()
|
||||
return False
|
||||
@@ -268,11 +384,24 @@ class TakeoverService:
|
||||
if isinstance(exc, BoxIMError) and exc.auth_error:
|
||||
self._forget_boxim_session(user.id)
|
||||
message = "BOXIM 授权已失效,请重新登录会会生产账号"
|
||||
disable_takeover = True
|
||||
else:
|
||||
message = f"BOXIM 暂时连接失败:{str(exc)[:160]}"
|
||||
self._disable_after_connection_failure(db, avatar, cursor, message)
|
||||
disable_takeover = False
|
||||
self._record_connection_failure(
|
||||
db,
|
||||
avatar,
|
||||
cursor,
|
||||
message,
|
||||
disable_takeover=disable_takeover,
|
||||
)
|
||||
db.commit()
|
||||
logger.warning("BOXIM sync failed for avatar %s: %s", avatar.id, exc)
|
||||
logger.warning(
|
||||
"BOXIM sync failed for avatar %s (will_retry=%s): %s",
|
||||
avatar.id,
|
||||
not disable_takeover,
|
||||
exc,
|
||||
)
|
||||
return False
|
||||
|
||||
messages.sort(key=lambda item: (_numeric_id(item.get("id")), item.get("sendTime") or 0))
|
||||
@@ -350,13 +479,21 @@ class TakeoverService:
|
||||
|
||||
now = self.now()
|
||||
send_time = _boxim_time(message.get("sendTime"), now)
|
||||
is_avatar = False
|
||||
if direction == "outgoing" and local_id:
|
||||
is_avatar = _is_avatar_local_id(local_id)
|
||||
if not is_avatar and local_id:
|
||||
is_avatar = bool(
|
||||
db.query(TakeoverReplyTask)
|
||||
.filter(
|
||||
TakeoverReplyTask.owner_id == avatar.owner_id,
|
||||
TakeoverReplyTask.boxim_local_id == local_id,
|
||||
TakeoverReplyTask.status.in_(("ready", "sending", "sent")),
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if not is_avatar:
|
||||
is_avatar = bool(
|
||||
db.query(TakeoverReplyTask)
|
||||
.filter(
|
||||
TakeoverReplyTask.boxim_sent_message_id == message_id,
|
||||
TakeoverReplyTask.status == "sent",
|
||||
)
|
||||
.first()
|
||||
@@ -381,12 +518,96 @@ class TakeoverService:
|
||||
if not is_avatar:
|
||||
self._cancel_conversation(db, avatar.owner_id, peer_id, "owner_replied")
|
||||
return
|
||||
if not schedule_reply or event.message_type != 0 or not event.content.strip():
|
||||
if not schedule_reply or event.message_type not in {
|
||||
BOXIM_TEXT_MESSAGE_TYPE,
|
||||
BOXIM_IMAGE_MESSAGE_TYPE,
|
||||
}:
|
||||
return
|
||||
if (now - send_time).total_seconds() > MAX_STALE_SECONDS:
|
||||
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:
|
||||
logger.info(
|
||||
"Ignored stale BOXIM message %s for avatar %s (age=%ss)",
|
||||
message_id,
|
||||
avatar.id,
|
||||
int((now - send_time).total_seconds()),
|
||||
)
|
||||
return
|
||||
if is_avatar:
|
||||
self._cancel_conversation(db, avatar.owner_id, peer_id, "peer_avatar_message")
|
||||
logger.info(
|
||||
"Skipped BOXIM reply for avatar %s message %s: peer_avatar_message",
|
||||
avatar.id,
|
||||
message_id,
|
||||
)
|
||||
return
|
||||
if self._human_pause_active(db, avatar.owner_id, peer_id, now):
|
||||
self._cancel_conversation(db, avatar.owner_id, peer_id, "owner_active")
|
||||
logger.info(
|
||||
"Skipped BOXIM reply for avatar %s message %s: owner_active",
|
||||
avatar.id,
|
||||
message_id,
|
||||
)
|
||||
return
|
||||
if self._conversation_rate_limited(db, avatar.owner_id, peer_id, now):
|
||||
self._cancel_conversation(db, avatar.owner_id, peer_id, "rate_limited")
|
||||
logger.info(
|
||||
"Skipped BOXIM reply for avatar %s message %s: rate_limited",
|
||||
avatar.id,
|
||||
message_id,
|
||||
)
|
||||
return
|
||||
self._schedule_reply(db, avatar, event)
|
||||
|
||||
@staticmethod
|
||||
def _human_pause_active(db: Session, owner_id: str, peer_id: str, now: datetime) -> bool:
|
||||
threshold = now - timedelta(seconds=HUMAN_PAUSE_SECONDS)
|
||||
return bool(
|
||||
db.query(TakeoverMessage.id)
|
||||
.filter(
|
||||
TakeoverMessage.owner_id == owner_id,
|
||||
TakeoverMessage.peer_id == peer_id,
|
||||
TakeoverMessage.direction == "outgoing",
|
||||
TakeoverMessage.is_avatar.is_(False),
|
||||
TakeoverMessage.send_time >= threshold,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _conversation_rate_limited(
|
||||
db: Session,
|
||||
owner_id: str,
|
||||
peer_id: str,
|
||||
now: datetime,
|
||||
) -> bool:
|
||||
threshold = now - timedelta(seconds=RATE_LIMIT_WINDOW_SECONDS)
|
||||
return (
|
||||
db.query(TakeoverReplyTask.id)
|
||||
.filter(
|
||||
TakeoverReplyTask.owner_id == owner_id,
|
||||
TakeoverReplyTask.peer_id == peer_id,
|
||||
TakeoverReplyTask.status == "sent",
|
||||
TakeoverReplyTask.sent_at >= threshold,
|
||||
)
|
||||
.count()
|
||||
>= RATE_LIMIT_MAX_REPLIES
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _cancel_conversation(db: Session, owner_id: str, peer_id: str, reason: str):
|
||||
tasks = (
|
||||
@@ -424,12 +645,16 @@ class TakeoverService:
|
||||
task.status = "cancelled"
|
||||
task.cancel_reason = "newer_incoming_message"
|
||||
task.locked_at = None
|
||||
prompt_parts.append(event.content.strip())
|
||||
prompt_parts.append(_event_prompt(event))
|
||||
source_ids.append(event.boxim_message_id)
|
||||
prompt = "\n".join(part for part in prompt_parts if part).strip()[-MAX_PROMPT_LENGTH:]
|
||||
due_at = event.send_time + timedelta(seconds=self.reply_delay_seconds)
|
||||
due_at = max(
|
||||
event.send_time
|
||||
+ timedelta(seconds=_configured_reply_delay(avatar, self.reply_delay_seconds)),
|
||||
self.now(),
|
||||
)
|
||||
task_id = secrets.token_hex(16)
|
||||
local_id = int(time.time() * 1000) * 1000 + secrets.randbelow(1000)
|
||||
local_id = _avatar_local_id(avatar.owner_id, event.boxim_message_id)
|
||||
db.add(
|
||||
TakeoverReplyTask(
|
||||
id=task_id,
|
||||
@@ -455,6 +680,7 @@ class TakeoverService:
|
||||
.filter(
|
||||
TakeoverReplyTask.status.in_(GENERATABLE_TASK_STATUSES),
|
||||
TakeoverReplyTask.response_text == "",
|
||||
TakeoverReplyTask.scheduled_at <= self.now(),
|
||||
)
|
||||
.order_by(TakeoverReplyTask.created_at.asc())
|
||||
.limit(10)
|
||||
@@ -478,6 +704,50 @@ class TakeoverService:
|
||||
results = await asyncio.gather(*(generate(task_id) for task_id in task_ids))
|
||||
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:
|
||||
db = self.session_factory()
|
||||
try:
|
||||
@@ -501,14 +771,44 @@ class TakeoverService:
|
||||
.filter(
|
||||
TakeoverMessage.owner_id == task.owner_id,
|
||||
TakeoverMessage.peer_id == task.peer_id,
|
||||
TakeoverMessage.avatar_id == task.avatar_id,
|
||||
)
|
||||
.order_by(TakeoverMessage.send_time.desc())
|
||||
.limit(30)
|
||||
.all()
|
||||
)
|
||||
source_events = {
|
||||
event.boxim_message_id: event
|
||||
for event in events
|
||||
if event.boxim_message_id in excluded_ids
|
||||
}
|
||||
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 = []
|
||||
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
|
||||
if event.direction == "incoming" and event.is_avatar:
|
||||
continue
|
||||
history.append(
|
||||
{
|
||||
@@ -518,10 +818,25 @@ class TakeoverService:
|
||||
)
|
||||
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")
|
||||
answer = _plain_text_reply(result.get("answer", ""))
|
||||
image_contexts = _attachment_contexts(image_attachments)
|
||||
has_source_text = any(
|
||||
event.message_type == BOXIM_TEXT_MESSAGE_TYPE and event.content.strip()
|
||||
for event in source_events.values()
|
||||
)
|
||||
if image_failed and not image_contexts and not has_source_text:
|
||||
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)
|
||||
if task.status != "generating":
|
||||
return False
|
||||
@@ -582,11 +897,20 @@ class TakeoverService:
|
||||
task.cancel_reason = "takeover_disabled"
|
||||
db.commit()
|
||||
return False
|
||||
if (self.now() - task.scheduled_at).total_seconds() > MAX_STALE_SECONDS:
|
||||
if (self.now() - task.scheduled_at).total_seconds() > MAX_SEND_OVERDUE_SECONDS:
|
||||
task.status = "cancelled"
|
||||
task.cancel_reason = "stale_reply"
|
||||
db.commit()
|
||||
return False
|
||||
cursor = (
|
||||
db.query(TakeoverCursor)
|
||||
.filter(TakeoverCursor.avatar_id == task.avatar_id)
|
||||
.first()
|
||||
)
|
||||
if not cursor or not cursor.last_polled_at or cursor.last_polled_at < task.scheduled_at:
|
||||
# Do not race the owner's final seconds of the grace period. A
|
||||
# completed poll at/after the due time must confirm no human reply.
|
||||
return False
|
||||
user = db.query(User).filter(User.huihui_user_id == task.owner_id).first()
|
||||
if not user or not user.huihui_token:
|
||||
raise BoxIMError("缺少会会登录凭证", auth_error=True)
|
||||
|
||||
@@ -79,12 +79,17 @@ def reserve_avatar_tokens(
|
||||
model: str,
|
||||
messages: list[dict],
|
||||
max_output_tokens: int,
|
||||
*,
|
||||
minimum_reserve_tokens: int = 0,
|
||||
) -> TokenReservation:
|
||||
user = avatar_owner_user(db, avatar)
|
||||
if not user:
|
||||
raise InsufficientTokensError("分身尚未关联有效用户,暂时无法使用积分")
|
||||
account = get_or_create_account(db, user.id)
|
||||
reserved = estimate_request_tokens(messages, max_output_tokens)
|
||||
reserved = max(
|
||||
estimate_request_tokens(messages, max_output_tokens),
|
||||
max(0, int(minimum_reserve_tokens or 0)),
|
||||
)
|
||||
updated = (
|
||||
db.query(TokenAccount)
|
||||
.filter(TokenAccount.id == account.id, TokenAccount.balance >= reserved)
|
||||
|
||||
@@ -0,0 +1,196 @@
|
||||
"""Private image normalization and OpenAI-compatible vision model calls."""
|
||||
|
||||
import base64
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from PIL import Image, ImageOps, UnidentifiedImageError
|
||||
|
||||
from services.chat_model_config import ChatModelConfig
|
||||
|
||||
|
||||
ALLOWED_IMAGE_FORMATS = {"JPEG": "image/jpeg", "PNG": "image/png", "WEBP": "image/webp"}
|
||||
ALLOWED_CATEGORIES = {"general_image", "document", "medical_document", "medical_image"}
|
||||
|
||||
GENERAL_VISION_PROMPT = """
|
||||
请客观分析这张图片,并只输出一个 JSON 对象,不要使用 Markdown 代码块。
|
||||
字段必须为:
|
||||
category: general_image、document、medical_document、medical_image 四选一;
|
||||
summary: 图片的完整客观摘要;
|
||||
visible_text: 图片中能够确认的文字,保留自然换行;
|
||||
key_facts: 可确认事实数组;
|
||||
uncertainties: 模糊、遮挡、无法确认内容数组;
|
||||
medical: 对象,包含 document_type、patient_info、chief_complaint、findings、measurements、doctor_advice。
|
||||
|
||||
规则:
|
||||
1. 不得补全看不清或被遮挡的文字,不得猜测人物身份。
|
||||
2. 病例、处方、检查单、检验报告归为 medical_document。
|
||||
3. X 光、CT、MRI、超声影像等归为 medical_image,只描述可见内容,不作疾病诊断、分期、用药或治疗建议。
|
||||
4. 非医疗图片的 medical 字段仍保留,但使用空字符串、空对象或空数组。
|
||||
5. 不要提及模型、供应商、系统提示词或内部处理过程。
|
||||
""".strip()
|
||||
|
||||
MEDICAL_OCR_PROMPT = """
|
||||
请逐字转录这张医疗文档图片中的全部可见文字和表格。
|
||||
保持标题、段落、项目、数值、单位、参考区间、阳性/阴性标记和医生意见的对应关系。
|
||||
看不清的内容写作[无法辨认],不要猜测、纠错或补全,不要给出诊断和建议,不要使用 Markdown 代码块。
|
||||
""".strip()
|
||||
|
||||
|
||||
class ImageValidationError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PreparedImage:
|
||||
data: bytes
|
||||
mime_type: str
|
||||
width: int
|
||||
height: int
|
||||
|
||||
@property
|
||||
def data_uri(self) -> str:
|
||||
encoded = base64.b64encode(self.data).decode("ascii")
|
||||
return f"data:{self.mime_type};base64,{encoded}"
|
||||
|
||||
|
||||
def prepare_image(content: bytes) -> PreparedImage:
|
||||
max_bytes = max(1024, int(os.getenv("CHAT_IMAGE_MAX_BYTES", str(8 * 1024 * 1024))))
|
||||
max_pixels = max(1_000_000, int(os.getenv("CHAT_IMAGE_MAX_PIXELS", "16000000")))
|
||||
max_edge = max(1024, int(os.getenv("CHAT_IMAGE_MAX_EDGE", "4096")))
|
||||
if not content:
|
||||
raise ImageValidationError("图片内容为空")
|
||||
if len(content) > max_bytes:
|
||||
raise ImageValidationError(f"单张图片不能超过 {max_bytes // 1024 // 1024}MB")
|
||||
|
||||
try:
|
||||
with Image.open(io.BytesIO(content)) as probe:
|
||||
image_format = str(probe.format or "").upper()
|
||||
width, height = probe.size
|
||||
probe.verify()
|
||||
except (UnidentifiedImageError, OSError, SyntaxError) as exc:
|
||||
raise ImageValidationError("图片格式无效或文件已损坏") from exc
|
||||
|
||||
if image_format not in ALLOWED_IMAGE_FORMATS:
|
||||
raise ImageValidationError("仅支持 JPG、PNG、WebP 图片")
|
||||
if width <= 0 or height <= 0 or width * height > max_pixels:
|
||||
raise ImageValidationError("图片像素过大,请压缩后重新上传")
|
||||
|
||||
try:
|
||||
with Image.open(io.BytesIO(content)) as original:
|
||||
image = ImageOps.exif_transpose(original)
|
||||
image.load()
|
||||
if max(image.size) > max_edge:
|
||||
image.thumbnail((max_edge, max_edge), Image.Resampling.LANCZOS)
|
||||
if image.mode in {"RGBA", "LA"}:
|
||||
canvas = Image.new("RGB", image.size, "white")
|
||||
alpha = image.getchannel("A")
|
||||
canvas.paste(image.convert("RGB"), mask=alpha)
|
||||
image = canvas
|
||||
elif image.mode != "RGB":
|
||||
image = image.convert("RGB")
|
||||
output = io.BytesIO()
|
||||
image.save(output, format="JPEG", quality=92, optimize=True)
|
||||
normalized = output.getvalue()
|
||||
normalized_width, normalized_height = image.size
|
||||
except (OSError, ValueError) as exc:
|
||||
raise ImageValidationError("图片解码失败,请重新选择图片") from exc
|
||||
|
||||
return PreparedImage(
|
||||
data=normalized,
|
||||
mime_type="image/jpeg",
|
||||
width=normalized_width,
|
||||
height=normalized_height,
|
||||
)
|
||||
|
||||
|
||||
def call_vision_model(
|
||||
prepared: PreparedImage,
|
||||
model_config: ChatModelConfig,
|
||||
*,
|
||||
model: str,
|
||||
prompt: str,
|
||||
json_output: bool,
|
||||
) -> dict:
|
||||
if not model_config.api_key:
|
||||
raise RuntimeError("视觉模型服务未配置")
|
||||
payload: dict[str, Any] = {
|
||||
"model": model,
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "image_url", "image_url": {"url": prepared.data_uri}},
|
||||
{"type": "text", "text": prompt},
|
||||
],
|
||||
}
|
||||
],
|
||||
"temperature": 0,
|
||||
"max_tokens": model_config.vision_max_tokens,
|
||||
}
|
||||
if json_output:
|
||||
payload["response_format"] = {"type": "json_object"}
|
||||
try:
|
||||
response = httpx.post(
|
||||
f"{model_config.api_base_url}/chat/completions",
|
||||
headers={"Authorization": f"Bearer {model_config.api_key}"},
|
||||
json=payload,
|
||||
timeout=model_config.vision_timeout_seconds,
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
content = data.get("choices", [{}])[0].get("message", {}).get("content", "")
|
||||
except (httpx.HTTPError, ValueError, KeyError, IndexError) as exc:
|
||||
raise RuntimeError("图片识别服务暂时不可用") from exc
|
||||
if not isinstance(content, str) or not content.strip():
|
||||
raise RuntimeError("图片识别服务没有返回有效结果")
|
||||
return {"content": content.strip(), "usage": data.get("usage") or {}}
|
||||
|
||||
|
||||
def parse_vision_analysis(content: str) -> dict:
|
||||
value = (content or "").strip()
|
||||
fenced = re.match(r"^```(?:json)?\s*(.*?)\s*```$", value, re.DOTALL | re.IGNORECASE)
|
||||
if fenced:
|
||||
value = fenced.group(1).strip()
|
||||
try:
|
||||
payload = json.loads(value)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise RuntimeError("图片识别结果格式无效") from exc
|
||||
if not isinstance(payload, dict):
|
||||
raise RuntimeError("图片识别结果格式无效")
|
||||
|
||||
category = str(payload.get("category") or "general_image").strip().lower()
|
||||
if category not in ALLOWED_CATEGORIES:
|
||||
category = "general_image"
|
||||
medical = payload.get("medical") if isinstance(payload.get("medical"), dict) else {}
|
||||
return {
|
||||
"category": category,
|
||||
"summary": str(payload.get("summary") or "").strip(),
|
||||
"visible_text": str(payload.get("visible_text") or "").strip(),
|
||||
"key_facts": _string_list(payload.get("key_facts")),
|
||||
"uncertainties": _string_list(payload.get("uncertainties")),
|
||||
"medical": medical,
|
||||
}
|
||||
|
||||
|
||||
def build_attachment_warning(analysis: dict, *, ocr_failed: bool = False) -> str:
|
||||
warnings = list(analysis.get("uncertainties") or [])
|
||||
category = analysis.get("category")
|
||||
if ocr_failed:
|
||||
warnings.append("精确文字识别暂时不可用,请人工核对图片原文")
|
||||
if category == "medical_document":
|
||||
warnings.append("病例识别结果仅供辅助,不能替代医生诊断,请核对原始文档")
|
||||
elif category == "medical_image":
|
||||
warnings.append("医学影像仅作客观描述,不能替代影像报告和医生诊断")
|
||||
return ";".join(dict.fromkeys(item for item in warnings if item))
|
||||
|
||||
|
||||
def _string_list(value: Any) -> list[str]:
|
||||
if not isinstance(value, list):
|
||||
return []
|
||||
return [str(item).strip() for item in value if str(item).strip()]
|
||||
@@ -5,6 +5,7 @@ from database import init_db, SessionLocal
|
||||
from models import (
|
||||
Authorization,
|
||||
Avatar,
|
||||
ChatAttachment,
|
||||
TakeoverCursor,
|
||||
TakeoverMessage,
|
||||
TakeoverReplyTask,
|
||||
@@ -95,6 +96,9 @@ def authorization_context():
|
||||
finally:
|
||||
db.rollback()
|
||||
avatar_ids = [avatar.id, other_avatar.id]
|
||||
db.query(ChatAttachment).filter(
|
||||
ChatAttachment.avatar_id.in_(avatar_ids)
|
||||
).delete(synchronize_session=False)
|
||||
db.query(TakeoverReplyTask).filter(
|
||||
TakeoverReplyTask.avatar_id.in_(avatar_ids)
|
||||
).delete(synchronize_session=False)
|
||||
|
||||
@@ -103,6 +103,7 @@ def test_avatar_permission_settings_default_and_persist(authorization_context):
|
||||
assert initial["data"] == {
|
||||
"avatarId": context["avatar"].id,
|
||||
"permissions": ["friend", "chat"],
|
||||
"takeoverReplyDelaySeconds": 180,
|
||||
}
|
||||
|
||||
updated = client.put(
|
||||
@@ -115,6 +116,7 @@ def test_avatar_permission_settings_default_and_persist(authorization_context):
|
||||
|
||||
reloaded = client.get(endpoint, headers=context["owner_headers"]).json()
|
||||
assert reloaded["data"]["permissions"] == ["friend", "publish", "interact", "takeover"]
|
||||
assert reloaded["data"]["takeoverReplyDelaySeconds"] == 180
|
||||
|
||||
|
||||
def test_avatar_permission_settings_allow_all_disabled(authorization_context):
|
||||
@@ -156,3 +158,57 @@ def test_avatar_permission_settings_validate_owner_and_permissions(authorization
|
||||
|
||||
unauthenticated = client.get(endpoint)
|
||||
assert unauthenticated.status_code == 401
|
||||
|
||||
|
||||
def test_takeover_delay_minimum_and_single_active_avatar_per_owner(authorization_context):
|
||||
from database import SessionLocal
|
||||
from models import Avatar
|
||||
|
||||
context = authorization_context
|
||||
endpoint = f"/api/avatar/{context['avatar'].id}/permission-settings"
|
||||
invalid = client.put(
|
||||
endpoint,
|
||||
headers=context["owner_headers"],
|
||||
json={"permissions": ["chat"], "takeoverReplyDelaySeconds": 2},
|
||||
).json()
|
||||
assert invalid["code"] == 400
|
||||
|
||||
second_avatar_id = f"second-{context['suffix']}"
|
||||
db = SessionLocal()
|
||||
try:
|
||||
db.add(
|
||||
Avatar(
|
||||
id=second_avatar_id,
|
||||
owner_id=context["owner"].huihui_user_id,
|
||||
name="第二个分身",
|
||||
status="active",
|
||||
config={"authorizationPermissions": ["chat", "takeover"]},
|
||||
)
|
||||
)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
try:
|
||||
updated = client.put(
|
||||
endpoint,
|
||||
headers=context["owner_headers"],
|
||||
json={"permissions": ["chat", "takeover"], "takeoverReplyDelaySeconds": 3},
|
||||
).json()
|
||||
assert updated["code"] == 200
|
||||
assert updated["data"]["takeoverReplyDelaySeconds"] == 3
|
||||
assert updated["data"]["disabledAvatarIds"] == [second_avatar_id]
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
second = db.query(Avatar).filter(Avatar.id == second_avatar_id).one()
|
||||
assert "takeover" not in second.config["authorizationPermissions"]
|
||||
finally:
|
||||
db.close()
|
||||
finally:
|
||||
db = SessionLocal()
|
||||
try:
|
||||
db.query(Avatar).filter(Avatar.id == second_avatar_id).delete()
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
"""Ownership and configuration-isolation tests for digital avatars."""
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from database import SessionLocal
|
||||
from main import app
|
||||
from models import Avatar
|
||||
|
||||
|
||||
client = TestClient(app)
|
||||
|
||||
|
||||
def test_avatar_detail_and_update_require_the_owner(authorization_context):
|
||||
context = authorization_context
|
||||
avatar_id = context["avatar"].id
|
||||
|
||||
assert client.get(f"/api/avatar/{avatar_id}").status_code == 401
|
||||
assert client.get(
|
||||
f"/api/avatar/{avatar_id}", headers=context["other_headers"]
|
||||
).status_code == 403
|
||||
|
||||
updated = client.put(
|
||||
f"/api/avatar/{avatar_id}",
|
||||
headers=context["owner_headers"],
|
||||
json={
|
||||
"description": "独立描述",
|
||||
"config": {"replyStyle": "concise"},
|
||||
},
|
||||
)
|
||||
assert updated.status_code == 200
|
||||
assert updated.json()["data"]["description"] == "独立描述"
|
||||
|
||||
forbidden = client.put(
|
||||
f"/api/avatar/{avatar_id}",
|
||||
headers=context["other_headers"],
|
||||
json={"description": "越权修改"},
|
||||
)
|
||||
assert forbidden.status_code == 403
|
||||
|
||||
|
||||
def test_avatar_config_updates_do_not_erase_takeover_or_knowledge_scope(authorization_context):
|
||||
context = authorization_context
|
||||
avatar_id = context["avatar"].id
|
||||
db = SessionLocal()
|
||||
try:
|
||||
avatar = db.query(Avatar).filter(Avatar.id == avatar_id).one()
|
||||
avatar.config = {
|
||||
"authorizationPermissions": ["chat", "takeover"],
|
||||
"takeoverReplyDelaySeconds": 180,
|
||||
}
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
response = client.put(
|
||||
f"/api/avatar/{avatar_id}",
|
||||
headers=context["owner_headers"],
|
||||
json={"config": {"replyStyle": "warm", "creativity": 25}},
|
||||
).json()
|
||||
config = response["data"]["config"]
|
||||
assert config["replyStyle"] == "warm"
|
||||
assert config["creativity"] == 25
|
||||
assert config["authorizationPermissions"] == ["chat", "takeover"]
|
||||
assert config["takeoverReplyDelaySeconds"] == 180
|
||||
|
||||
|
||||
def test_avatar_create_and_delete_require_login_and_ownership(authorization_context):
|
||||
context = authorization_context
|
||||
assert client.post("/api/avatar", json={"name": "匿名分身"}).status_code == 401
|
||||
|
||||
created = client.post(
|
||||
"/api/avatar",
|
||||
headers=context["owner_headers"],
|
||||
json={"name": "待删除分身"},
|
||||
)
|
||||
assert created.status_code == 200
|
||||
avatar_id = created.json()["data"]["id"]
|
||||
|
||||
try:
|
||||
assert client.delete(
|
||||
f"/api/avatar/{avatar_id}", headers=context["other_headers"]
|
||||
).status_code == 403
|
||||
deleted = client.delete(
|
||||
f"/api/avatar/{avatar_id}", headers=context["owner_headers"]
|
||||
).json()
|
||||
assert deleted["code"] == 200
|
||||
finally:
|
||||
db = SessionLocal()
|
||||
try:
|
||||
db.query(Avatar).filter(Avatar.id == avatar_id).delete()
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
@@ -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,
|
||||
)
|
||||
@@ -0,0 +1,290 @@
|
||||
import json
|
||||
from datetime import datetime, timedelta
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from database import SessionLocal
|
||||
from main import app
|
||||
from models import ChatAttachment
|
||||
from routers.chat import (
|
||||
ChatIn,
|
||||
_attachment_contexts,
|
||||
_load_chat_attachments,
|
||||
_resolve_reply,
|
||||
)
|
||||
from services.chat_attachment_service import purge_expired_chat_attachments
|
||||
from services.token_billing import InsufficientTokensError
|
||||
from services.vision_service import PreparedImage
|
||||
|
||||
|
||||
client = TestClient(app)
|
||||
|
||||
|
||||
GENERAL_RESULT = {
|
||||
"content": json.dumps({
|
||||
"category": "general_image",
|
||||
"summary": "一张包含产品路线图的截图",
|
||||
"visible_text": "产品路线图",
|
||||
"key_facts": ["包含三个阶段"],
|
||||
"uncertainties": [],
|
||||
"medical": {},
|
||||
}, ensure_ascii=False),
|
||||
"usage": {"total_tokens": 120},
|
||||
}
|
||||
|
||||
|
||||
def test_owner_can_upload_and_cache_image_analysis(authorization_context):
|
||||
context = authorization_context
|
||||
prepared = PreparedImage(b"jpeg", "image/jpeg", 100, 80)
|
||||
with (
|
||||
patch("routers.chat.prepare_image", return_value=prepared),
|
||||
patch("routers.chat._run_billed_vision_call", return_value=GENERAL_RESULT),
|
||||
):
|
||||
response = client.post(
|
||||
f"/api/avatar/{context['avatar'].id}/chat/images",
|
||||
headers=context["owner_headers"],
|
||||
files={"file": ("roadmap.png", b"image-bytes", "image/png")},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
payload = response.json()["data"]
|
||||
assert payload["status"] == "ready"
|
||||
assert payload["category"] == "general_image"
|
||||
assert payload["summary"] == "一张包含产品路线图的截图"
|
||||
db = SessionLocal()
|
||||
try:
|
||||
stored = db.query(ChatAttachment).filter(ChatAttachment.id == payload["id"]).one()
|
||||
assert stored.avatar_id == context["avatar"].id
|
||||
assert stored.extracted_text == "产品路线图"
|
||||
assert stored.structured_data["key_facts"] == ["包含三个阶段"]
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_non_owner_cannot_upload_chat_image(authorization_context):
|
||||
context = authorization_context
|
||||
response = client.post(
|
||||
f"/api/avatar/{context['avatar'].id}/chat/images",
|
||||
headers=context["other_headers"],
|
||||
files={"file": ("private.png", b"image-bytes", "image/png")},
|
||||
)
|
||||
|
||||
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):
|
||||
context = authorization_context
|
||||
db = SessionLocal()
|
||||
try:
|
||||
avatar = db.get(type(context["avatar"]), context["avatar"].id)
|
||||
avatar.share_token = f"share-{context['suffix']}"
|
||||
db.commit()
|
||||
share_token = avatar.share_token
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routers.chat.prepare_image",
|
||||
return_value=PreparedImage(b"jpeg", "image/jpeg", 100, 80),
|
||||
),
|
||||
patch("routers.chat._run_billed_vision_call", return_value=GENERAL_RESULT),
|
||||
):
|
||||
response = client.post(
|
||||
f"/api/public/avatar/{share_token}/chat/images",
|
||||
files={"file": ("visitor.png", b"image-bytes", "image/png")},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
payload = response.json()["data"]
|
||||
assert payload["status"] == "ready"
|
||||
assert "structuredData" not in payload
|
||||
assert "extractedText" not in payload
|
||||
assert "visionModel" not in payload
|
||||
assert "ocrModel" not in payload
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
stored = db.get(ChatAttachment, payload["id"])
|
||||
assert stored.uploader_kind == "public"
|
||||
assert stored.avatar_id == context["avatar"].id
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_medical_document_uses_ocr_result(authorization_context):
|
||||
context = authorization_context
|
||||
general = {
|
||||
"content": json.dumps({
|
||||
"category": "medical_document",
|
||||
"summary": "血常规报告",
|
||||
"visible_text": "初步文字",
|
||||
"key_facts": [],
|
||||
"uncertainties": [],
|
||||
"medical": {"document_type": "检验报告"},
|
||||
}, ensure_ascii=False),
|
||||
"usage": {},
|
||||
}
|
||||
ocr = {"content": "白细胞 11.2 x10^9/L", "usage": {}}
|
||||
with (
|
||||
patch("routers.chat.prepare_image", return_value=PreparedImage(b"jpeg", "image/jpeg", 100, 80)),
|
||||
patch("routers.chat._run_billed_vision_call", side_effect=[general, ocr]) as model,
|
||||
):
|
||||
response = client.post(
|
||||
f"/api/avatar/{context['avatar'].id}/chat/images",
|
||||
headers=context["owner_headers"],
|
||||
files={"file": ("report.jpg", b"image-bytes", "image/jpeg")},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
attachment_id = response.json()["data"]["id"]
|
||||
assert model.call_count == 2
|
||||
assert model.call_args_list[1].kwargs["source"] == "vision_medical_ocr"
|
||||
db = SessionLocal()
|
||||
try:
|
||||
stored = db.query(ChatAttachment).filter(ChatAttachment.id == attachment_id).one()
|
||||
assert stored.extracted_text == "白细胞 11.2 x10^9/L"
|
||||
assert stored.ocr_model == "qwen-vl-ocr"
|
||||
assert "不能替代医生诊断" in stored.warning
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_attachment_cannot_cross_avatar_boundary(authorization_context):
|
||||
context = authorization_context
|
||||
db = SessionLocal()
|
||||
try:
|
||||
attachment = ChatAttachment(
|
||||
avatar_id=context["avatar"].id,
|
||||
filename="private.jpg",
|
||||
status="ready",
|
||||
expires_at=datetime.utcnow() + timedelta(hours=1),
|
||||
)
|
||||
db.add(attachment)
|
||||
db.commit()
|
||||
body = ChatIn(message="看看图片", attachmentIds=[attachment.id])
|
||||
with pytest.raises(HTTPException, match="不属于当前分身") as caught:
|
||||
_load_chat_attachments(db, context["other_avatar"].id, body)
|
||||
assert caught.value.status_code == 400
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_expired_attachment_is_removed(authorization_context):
|
||||
context = authorization_context
|
||||
db = SessionLocal()
|
||||
try:
|
||||
attachment = ChatAttachment(
|
||||
avatar_id=context["avatar"].id,
|
||||
filename="expired.jpg",
|
||||
status="ready",
|
||||
expires_at=datetime.utcnow() - timedelta(seconds=1),
|
||||
)
|
||||
db.add(attachment)
|
||||
db.commit()
|
||||
attachment_id = attachment.id
|
||||
body = ChatIn(message="看看图片", attachmentIds=[attachment_id])
|
||||
with pytest.raises(HTTPException):
|
||||
_load_chat_attachments(db, context["avatar"].id, body)
|
||||
assert db.query(ChatAttachment).filter(ChatAttachment.id == attachment_id).first() is None
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_cleanup_keeps_unexpired_attachment(authorization_context):
|
||||
context = authorization_context
|
||||
now = datetime.utcnow()
|
||||
db = SessionLocal()
|
||||
try:
|
||||
expired = ChatAttachment(
|
||||
avatar_id=context["avatar"].id,
|
||||
filename="expired.jpg",
|
||||
status="ready",
|
||||
expires_at=now - timedelta(seconds=1),
|
||||
)
|
||||
active = ChatAttachment(
|
||||
avatar_id=context["avatar"].id,
|
||||
filename="active.jpg",
|
||||
status="ready",
|
||||
expires_at=now + timedelta(hours=1),
|
||||
)
|
||||
db.add_all([expired, active])
|
||||
db.commit()
|
||||
expired_id, active_id = expired.id, active.id
|
||||
|
||||
assert purge_expired_chat_attachments(db, now=now) == 1
|
||||
assert db.get(ChatAttachment, expired_id) is None
|
||||
assert db.get(ChatAttachment, active_id) is not None
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_image_context_keeps_standard_answer_authoritative():
|
||||
avatar = SimpleNamespace(
|
||||
id="avatar-vision",
|
||||
name="测试分身",
|
||||
description="产品顾问",
|
||||
config={},
|
||||
)
|
||||
model = Mock(return_value="标准退款期限是七天;图片显示的是商品包装。")
|
||||
result = _resolve_reply(
|
||||
None,
|
||||
avatar,
|
||||
"退款期限是多少?",
|
||||
[],
|
||||
qa_pairs=[SimpleNamespace(question="退款期限是多少?", answer="七天", enabled=True)],
|
||||
search_fn=Mock(return_value=[]),
|
||||
model_client=model,
|
||||
image_contexts=[{
|
||||
"id": "attachment",
|
||||
"filename": "product.jpg",
|
||||
"category": "general_image",
|
||||
"summary": "商品包装",
|
||||
"extractedText": "",
|
||||
"structuredData": {},
|
||||
"warning": "",
|
||||
}],
|
||||
)
|
||||
|
||||
assert result["source"] == "qa"
|
||||
system = model.call_args.kwargs["messages"][0]["content"]
|
||||
assert "已确认标准答案" in system
|
||||
assert "七天" in system
|
||||
assert "商品包装" in system
|
||||
assert "标准答题对中的事实优先级高于图片资料" in system
|
||||
|
||||
|
||||
def test_attachment_context_does_not_expose_internal_fields():
|
||||
row = SimpleNamespace(
|
||||
id="attachment",
|
||||
filename="case.jpg",
|
||||
category="medical_document",
|
||||
summary="门诊病例",
|
||||
extracted_text="主诉:咳嗽",
|
||||
structured_data={"medical": {"chief_complaint": "咳嗽"}},
|
||||
warning="请核对原文",
|
||||
)
|
||||
context = _attachment_contexts([row])[0]
|
||||
assert context["filename"] == "case.jpg"
|
||||
assert "avatar_id" not in context
|
||||
assert "vision_model" not in context
|
||||
@@ -26,6 +26,8 @@ def test_admin_runtime_config_takes_priority(monkeypatch):
|
||||
"api_base_url": "https://model.test/v1/",
|
||||
"api_key": "runtime-key",
|
||||
"model": "avatar-model",
|
||||
"vision_model": "avatar-vision-model",
|
||||
"ocr_model": "avatar-ocr-model",
|
||||
"max_tokens": 2048,
|
||||
"timeout_seconds": 42,
|
||||
}
|
||||
@@ -37,6 +39,8 @@ def test_admin_runtime_config_takes_priority(monkeypatch):
|
||||
assert config.source == "admin"
|
||||
assert config.api_base_url == "https://model.test/v1"
|
||||
assert config.model == "avatar-model"
|
||||
assert config.vision_model == "avatar-vision-model"
|
||||
assert config.ocr_model == "avatar-ocr-model"
|
||||
assert config.max_tokens == 2048
|
||||
request.assert_called_once_with(
|
||||
"http://config.test/runtime",
|
||||
@@ -51,6 +55,8 @@ def test_runtime_failure_falls_back_to_environment(monkeypatch):
|
||||
monkeypatch.setenv("CHAT_API_URL", "https://fallback.test/v1/")
|
||||
monkeypatch.setenv("CHAT_API_KEY", "fallback-key")
|
||||
monkeypatch.setenv("CHAT_MODEL", "fallback-model")
|
||||
monkeypatch.setenv("VISION_MODEL", "fallback-vision")
|
||||
monkeypatch.setenv("VISION_OCR_MODEL", "fallback-ocr")
|
||||
monkeypatch.setenv("CHAT_MAX_OUTPUT_TOKENS", "1536")
|
||||
|
||||
request = httpx.Request("GET", "http://config.test/runtime")
|
||||
@@ -64,6 +70,8 @@ def test_runtime_failure_falls_back_to_environment(monkeypatch):
|
||||
assert config.api_base_url == "https://fallback.test/v1"
|
||||
assert config.api_key == "fallback-key"
|
||||
assert config.model == "fallback-model"
|
||||
assert config.vision_model == "fallback-vision"
|
||||
assert config.ocr_model == "fallback-ocr"
|
||||
assert config.max_tokens == 1536
|
||||
|
||||
|
||||
|
||||
@@ -5,7 +5,15 @@ from unittest.mock import Mock
|
||||
from fastapi import HTTPException
|
||||
|
||||
from models import Avatar, User
|
||||
from routers.chat import _build_prompt, _iter_text_chunks, _match_standard_qa, _public_avatar_payload, _require_owned_avatar, _resolve_reply
|
||||
from routers.chat import (
|
||||
_build_prompt,
|
||||
_iter_text_chunks,
|
||||
_match_standard_qa,
|
||||
_public_avatar_payload,
|
||||
_qa_requires_language_adaptation,
|
||||
_require_owned_avatar,
|
||||
_resolve_reply,
|
||||
)
|
||||
|
||||
|
||||
class ChatOrchestrationTests(unittest.TestCase):
|
||||
@@ -50,6 +58,34 @@ class ChatOrchestrationTests(unittest.TestCase):
|
||||
self.assertEqual(result["answer"], "标准地址")
|
||||
fake_model.assert_not_called()
|
||||
|
||||
def test_cross_language_qa_is_faithfully_adapted_by_model(self):
|
||||
fake_model = Mock(return_value="Our address is Test Road 1.")
|
||||
fake_search = Mock(return_value=[])
|
||||
result = _resolve_reply(
|
||||
None,
|
||||
self.avatar,
|
||||
"Where is your office?",
|
||||
[],
|
||||
qa_pairs=[SimpleNamespace(question="Where is your office?", answer="地址是测试路1号。", enabled=True)],
|
||||
search_fn=fake_search,
|
||||
model_client=fake_model,
|
||||
)
|
||||
|
||||
self.assertEqual(result["source"], "qa")
|
||||
self.assertEqual(result["answer"], "Our address is Test Road 1.")
|
||||
self.assertEqual(fake_model.call_args.kwargs["temperature"], 0.0)
|
||||
system = fake_model.call_args.kwargs["messages"][0]["content"]
|
||||
self.assertIn("已确认标准答案", system)
|
||||
self.assertIn("地址是测试路1号", system)
|
||||
self.assertIn("只使用该语言回答", system)
|
||||
fake_search.assert_not_called()
|
||||
|
||||
def test_qa_language_adaptation_detects_common_writing_system_changes(self):
|
||||
self.assertTrue(_qa_requires_language_adaptation("Hello", "你好"))
|
||||
self.assertTrue(_qa_requires_language_adaptation("こんにちは", "你好"))
|
||||
self.assertTrue(_qa_requires_language_adaptation("안녕하세요", "你好"))
|
||||
self.assertFalse(_qa_requires_language_adaptation("你好", "您好"))
|
||||
|
||||
def test_conversational_paraphrase_matches_standard_qa(self):
|
||||
for question in ("请问一下,你们公司在哪里呀?", "请问去你们那边怎么走"):
|
||||
with self.subTest(question=question):
|
||||
@@ -106,6 +142,9 @@ class ChatOrchestrationTests(unittest.TestCase):
|
||||
self.assertIn("像熟人之间微信聊天一样", messages[0]["content"])
|
||||
self.assertIn("不隶属于任何机构", messages[0]["content"])
|
||||
self.assertIn("不要连续输出空行", messages[0]["content"])
|
||||
self.assertIn("回答语言规则", messages[0]["content"])
|
||||
self.assertIn("当前最后一条用户消息", messages[0]["content"])
|
||||
self.assertIn("历史消息", messages[0]["content"])
|
||||
|
||||
def test_prompt_blocks_ungrounded_factual_answers(self):
|
||||
messages = _build_prompt(self.avatar, [], "聊聊国际新闻", [])
|
||||
@@ -113,6 +152,8 @@ class ChatOrchestrationTests(unittest.TestCase):
|
||||
self.assertIn("没有检索到可靠资料", system)
|
||||
self.assertIn("不要凭通用知识", system)
|
||||
self.assertIn("不要提及知识库", system)
|
||||
self.assertIn("不得推断服务对象", system)
|
||||
self.assertIn("工作场所", system)
|
||||
|
||||
def test_public_avatar_payload_excludes_internal_configuration(self):
|
||||
payload = _public_avatar_payload(self.avatar)
|
||||
|
||||
@@ -48,9 +48,11 @@ class RemoteEmbeddingTests(unittest.TestCase):
|
||||
def test_large_input_is_split_into_provider_safe_batches(self):
|
||||
texts = [f"chunk-{index}" for index in range(14)]
|
||||
batch_sizes = []
|
||||
requested_urls = []
|
||||
|
||||
def fake_urlopen(request, timeout):
|
||||
self.assertEqual(timeout, 30)
|
||||
requested_urls.append(request.full_url)
|
||||
payload = json.loads(request.data.decode("utf-8"))
|
||||
batch_sizes.append(len(payload["input"]))
|
||||
return FakeResponse({
|
||||
@@ -61,7 +63,7 @@ class RemoteEmbeddingTests(unittest.TestCase):
|
||||
})
|
||||
|
||||
with patch.dict(os.environ, {
|
||||
"EMBEDDING_API_URL": "https://embedding.example/v1/embeddings",
|
||||
"EMBEDDING_API_URL": "https://embedding.example/v1",
|
||||
"EMBEDDING_API_KEY": "test-key",
|
||||
"EMBEDDING_MODEL": "text-embedding-v4",
|
||||
"EMBEDDING_BATCH_SIZE": "10",
|
||||
@@ -69,8 +71,18 @@ class RemoteEmbeddingTests(unittest.TestCase):
|
||||
result = embeddings.embed(texts)
|
||||
|
||||
self.assertEqual(batch_sizes, [10, 4])
|
||||
self.assertEqual(requested_urls, [
|
||||
"https://embedding.example/v1/embeddings",
|
||||
"https://embedding.example/v1/embeddings",
|
||||
])
|
||||
self.assertEqual(result, [[float(index)] for index in range(14)])
|
||||
|
||||
def test_full_embedding_endpoint_is_not_modified(self):
|
||||
self.assertEqual(
|
||||
embeddings._embedding_endpoint("https://embedding.example/v1/embeddings/"),
|
||||
"https://embedding.example/v1/embeddings",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -2,9 +2,17 @@ from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from database import SessionLocal
|
||||
from main import app
|
||||
from models import Avatar, KnowledgeChunk, KnowledgeDoc, QAPair
|
||||
from routers.knowledge import _doc_payload
|
||||
|
||||
|
||||
client = TestClient(app)
|
||||
|
||||
|
||||
def test_doc_payload_reports_whether_the_persisted_file_exists(tmp_path: Path):
|
||||
avatar_id = "avatar-1"
|
||||
stored_name = "knowledge.md"
|
||||
@@ -21,3 +29,157 @@ def test_doc_payload_reports_whether_the_persisted_file_exists(tmp_path: Path):
|
||||
assert _doc_payload(doc)["filePresent"] is False
|
||||
stored_file.write_text("knowledge", encoding="utf-8")
|
||||
assert _doc_payload(doc)["filePresent"] is True
|
||||
|
||||
|
||||
def test_upload_marks_vectorization_failure_instead_of_staying_processing(
|
||||
tmp_path: Path,
|
||||
authorization_context,
|
||||
):
|
||||
context = authorization_context
|
||||
with (
|
||||
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
||||
patch("routers.knowledge.embeddings.embed", side_effect=RuntimeError("provider unavailable")),
|
||||
):
|
||||
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"]
|
||||
assert payload["status"] == "failed"
|
||||
assert payload["vectorized"] is False
|
||||
assert payload["chunkCount"] == 0
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
|
||||
assert stored.status == "failed"
|
||||
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 0
|
||||
db.delete(stored)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_markdown_upload_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.embeddings.embed", return_value=[[1.0, 0.0]]),
|
||||
):
|
||||
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"]
|
||||
assert payload["status"] == "ready"
|
||||
assert payload["vectorized"] is True
|
||||
assert payload["chunkCount"] == 1
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
|
||||
assert stored.status == "ready"
|
||||
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 1
|
||||
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).delete()
|
||||
db.delete(stored)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_each_avatar_has_an_independent_document_and_qa_scope(authorization_context):
|
||||
context = authorization_context
|
||||
first_avatar_id = context["avatar"].id
|
||||
second_avatar_id = f"knowledge-second-{context['suffix']}"
|
||||
first_doc_id = f"knowledge-first-doc-{context['suffix']}"
|
||||
second_doc_id = f"knowledge-second-doc-{context['suffix']}"
|
||||
first_qa_id = f"knowledge-first-qa-{context['suffix']}"
|
||||
second_qa_id = f"knowledge-second-qa-{context['suffix']}"
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
db.add_all(
|
||||
[
|
||||
Avatar(
|
||||
id=second_avatar_id,
|
||||
owner_id=context["owner"].huihui_user_id,
|
||||
name="独立知识库分身",
|
||||
status="active",
|
||||
config={},
|
||||
),
|
||||
KnowledgeDoc(
|
||||
id=first_doc_id,
|
||||
avatar_id=first_avatar_id,
|
||||
filename="first.md",
|
||||
status="ready",
|
||||
vectorized=True,
|
||||
),
|
||||
KnowledgeDoc(
|
||||
id=second_doc_id,
|
||||
avatar_id=second_avatar_id,
|
||||
filename="second.md",
|
||||
status="ready",
|
||||
vectorized=True,
|
||||
),
|
||||
QAPair(
|
||||
id=first_qa_id,
|
||||
avatar_id=first_avatar_id,
|
||||
question="第一个分身问题",
|
||||
answer="第一个分身答案",
|
||||
),
|
||||
QAPair(
|
||||
id=second_qa_id,
|
||||
avatar_id=second_avatar_id,
|
||||
question="第二个分身问题",
|
||||
answer="第二个分身答案",
|
||||
),
|
||||
]
|
||||
)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
try:
|
||||
first_docs = client.get(
|
||||
f"/api/avatar/{first_avatar_id}/knowledge/docs",
|
||||
headers=context["owner_headers"],
|
||||
).json()["data"]
|
||||
second_docs = client.get(
|
||||
f"/api/avatar/{second_avatar_id}/knowledge/docs",
|
||||
headers=context["owner_headers"],
|
||||
).json()["data"]
|
||||
first_qa = client.get(
|
||||
f"/api/avatar/{first_avatar_id}/knowledge/qa",
|
||||
headers=context["owner_headers"],
|
||||
).json()["data"]
|
||||
second_qa = client.get(
|
||||
f"/api/avatar/{second_avatar_id}/knowledge/qa",
|
||||
headers=context["owner_headers"],
|
||||
).json()["data"]
|
||||
|
||||
assert [item["id"] for item in first_docs if item["id"] == first_doc_id] == [first_doc_id]
|
||||
assert second_doc_id not in {item["id"] for item in first_docs}
|
||||
assert [item["id"] for item in second_docs] == [second_doc_id]
|
||||
assert first_qa_id in {item["id"] for item in first_qa}
|
||||
assert second_qa_id not in {item["id"] for item in first_qa}
|
||||
assert [item["id"] for item in second_qa] == [second_qa_id]
|
||||
finally:
|
||||
db = SessionLocal()
|
||||
try:
|
||||
db.query(QAPair).filter(QAPair.id.in_([first_qa_id, second_qa_id])).delete(
|
||||
synchronize_session=False
|
||||
)
|
||||
db.query(KnowledgeDoc).filter(
|
||||
KnowledgeDoc.id.in_([first_doc_id, second_doc_id])
|
||||
).delete(synchronize_session=False)
|
||||
db.query(Avatar).filter(Avatar.id == second_avatar_id).delete()
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@@ -13,7 +13,7 @@ def test_authorization_takeover_fields():
|
||||
assert hasattr(auth, 'takeover_delay_seconds')
|
||||
assert auth.takeover_enabled == False
|
||||
assert auth.takeover_mode == 'immediate'
|
||||
assert auth.takeover_delay_seconds == 30
|
||||
assert auth.takeover_delay_seconds == 180
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
@@ -20,8 +20,9 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
|
||||
):
|
||||
import main
|
||||
|
||||
maintenance_scheduler = MagicMock()
|
||||
scheduler = MagicMock()
|
||||
mock_scheduler_class.return_value = scheduler
|
||||
mock_scheduler_class.side_effect = [maintenance_scheduler, scheduler]
|
||||
boxim = MagicMock()
|
||||
mock_boxim_class.return_value = boxim
|
||||
takeover = MagicMock()
|
||||
@@ -45,7 +46,16 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
|
||||
config = mock_boxim_class.call_args.args[0]
|
||||
assert config["HUIHUI_PLATFORM_BASE_URL"] == "https://open.example/api"
|
||||
assert config["BOXIM_API_BASE_URL"] == "https://im.example/api"
|
||||
mock_takeover_class.assert_called_once_with(main.SessionLocal, boxim)
|
||||
mock_takeover_class.assert_called_once_with(
|
||||
main.SessionLocal,
|
||||
boxim,
|
||||
poll_concurrency=8,
|
||||
max_message_age_seconds=600,
|
||||
)
|
||||
|
||||
maintenance_scheduler.add_job.assert_called_once()
|
||||
assert maintenance_scheduler.add_job.call_args.kwargs["id"] == "chat_attachment_cleanup"
|
||||
maintenance_scheduler.start.assert_called_once_with()
|
||||
|
||||
assert scheduler.add_job.call_count == 2
|
||||
poll_call, process_call = scheduler.add_job.call_args_list
|
||||
@@ -62,6 +72,7 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
|
||||
scheduler.start.assert_called_once_with()
|
||||
|
||||
main.takeover_scheduler = None
|
||||
main.maintenance_scheduler = None
|
||||
|
||||
|
||||
@patch("main.AsyncIOScheduler")
|
||||
@@ -73,6 +84,7 @@ def test_scheduler_failure_does_not_stop_the_api(mock_scheduler_class):
|
||||
main.on_startup()
|
||||
|
||||
assert main.takeover_scheduler is None
|
||||
assert main.maintenance_scheduler is None
|
||||
|
||||
|
||||
def test_shutdown_stops_only_the_scheduler():
|
||||
@@ -80,9 +92,14 @@ def test_shutdown_stops_only_the_scheduler():
|
||||
|
||||
scheduler = MagicMock()
|
||||
scheduler.running = True
|
||||
maintenance_scheduler = MagicMock()
|
||||
maintenance_scheduler.running = True
|
||||
main.takeover_scheduler = scheduler
|
||||
main.maintenance_scheduler = maintenance_scheduler
|
||||
|
||||
main.on_shutdown()
|
||||
|
||||
scheduler.shutdown.assert_called_once_with(wait=False)
|
||||
maintenance_scheduler.shutdown.assert_called_once_with(wait=False)
|
||||
assert main.takeover_scheduler is None
|
||||
assert main.maintenance_scheduler is None
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""End-to-end service tests for BOXIM takeover timing and human priority."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from threading import Barrier
|
||||
from unittest.mock import AsyncMock, patch
|
||||
@@ -9,9 +11,15 @@ from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
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.takeover_service import TakeoverService, _plain_text_reply
|
||||
from services.boxim_image_service import DownloadedBoxIMImage
|
||||
from services.takeover_service import (
|
||||
AVATAR_LOCAL_ID_PREFIX,
|
||||
TakeoverService,
|
||||
_avatar_local_id,
|
||||
_plain_text_reply,
|
||||
)
|
||||
|
||||
|
||||
class Clock:
|
||||
@@ -57,6 +65,26 @@ class FakeBoxIM:
|
||||
return {"id": 900 + len(self.sent), "localId": int(local_id)}
|
||||
|
||||
|
||||
class ConcurrentPollingBoxIM(FakeBoxIM):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.active_polls = 0
|
||||
self.peak_active_polls = 0
|
||||
|
||||
async def exchange_access_token(self, huihui_token):
|
||||
return {"accessToken": huihui_token, "accessTokenExpiresIn": 3600}
|
||||
|
||||
async def get_self(self, access_token):
|
||||
return {"id": 100 if access_token == "prod-huihui-token" else 101}
|
||||
|
||||
async def fetch_private_messages(self, access_token, min_id="0"):
|
||||
self.active_polls += 1
|
||||
self.peak_active_polls = max(self.peak_active_polls, self.active_polls)
|
||||
await asyncio.sleep(0.05)
|
||||
self.active_polls -= 1
|
||||
return []
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def service_context(tmp_path):
|
||||
engine = create_engine(
|
||||
@@ -77,7 +105,10 @@ def service_context(tmp_path):
|
||||
owner_id=user.huihui_user_id,
|
||||
name="分身",
|
||||
status="active",
|
||||
config={"authorizationPermissions": ["chat", "takeover"]},
|
||||
config={
|
||||
"authorizationPermissions": ["chat", "takeover"],
|
||||
"takeoverReplyDelaySeconds": 3,
|
||||
},
|
||||
)
|
||||
db.add_all([user, avatar])
|
||||
db.commit()
|
||||
@@ -120,8 +151,7 @@ async def test_incoming_message_is_prepared_then_sent_at_three_seconds(service_c
|
||||
{"id": 11, "localId": 2, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "你好"}
|
||||
)
|
||||
|
||||
with patch("routers.chat._resolve_reply", return_value={"answer": "**你好**\n\n很高兴见到你"}):
|
||||
await service.poll_and_process_messages()
|
||||
await service.poll_and_process_messages()
|
||||
assert boxim.sent == []
|
||||
assert boxim.read_receipts == [{"friendId": "200", "messageId": "11"}]
|
||||
|
||||
@@ -130,7 +160,8 @@ async def test_incoming_message_is_prepared_then_sent_at_three_seconds(service_c
|
||||
assert boxim.sent == []
|
||||
|
||||
clock.advance(1)
|
||||
await service.poll_and_process_messages()
|
||||
with patch("routers.chat._resolve_reply", return_value={"answer": "**你好**\n\n很高兴见到你"}):
|
||||
await service.poll_and_process_messages()
|
||||
assert boxim.sent == [{"peerId": "200", "content": "你好\n很高兴见到你", "localId": boxim.sent[0]["localId"]}]
|
||||
|
||||
db = session_factory()
|
||||
@@ -142,6 +173,325 @@ async def test_incoming_message_is_prepared_then_sent_at_three_seconds(service_c
|
||||
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_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
|
||||
async def test_default_reply_delay_is_three_minutes(service_context):
|
||||
session_factory, service, boxim, clock = service_context
|
||||
db = session_factory()
|
||||
try:
|
||||
avatar = db.query(Avatar).one()
|
||||
avatar.config = {"authorizationPermissions": ["chat", "takeover"]}
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
await service.poll_and_process_messages()
|
||||
boxim.messages.append(
|
||||
{"id": 12, "localId": 12, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "三分钟后回复"}
|
||||
)
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
db = session_factory()
|
||||
try:
|
||||
task = db.query(TakeoverReplyTask).one()
|
||||
assert task.scheduled_at == clock.now() + timedelta(seconds=180)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
clock.advance(179)
|
||||
await service.process_reply_tasks()
|
||||
assert boxim.sent == []
|
||||
clock.advance(1)
|
||||
with patch("routers.chat._resolve_reply", return_value={"answer": "好的"}):
|
||||
await service.poll_and_process_messages()
|
||||
assert [item["content"] for item in boxim.sent] == ["好的"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multiple_avatar_owners_are_polled_concurrently(service_context):
|
||||
session_factory, _service, _boxim, clock = service_context
|
||||
db = session_factory()
|
||||
try:
|
||||
db.add_all(
|
||||
[
|
||||
User(
|
||||
id="owner-local-2",
|
||||
huihui_user_id="owner-huihui-2",
|
||||
huihui_token="prod-huihui-token-2",
|
||||
app_token="app-token-2",
|
||||
),
|
||||
Avatar(
|
||||
id="avatar-2",
|
||||
owner_id="owner-huihui-2",
|
||||
name="分身二",
|
||||
status="active",
|
||||
config={"authorizationPermissions": ["chat", "takeover"]},
|
||||
),
|
||||
]
|
||||
)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
boxim = ConcurrentPollingBoxIM()
|
||||
service = TakeoverService(
|
||||
session_factory,
|
||||
boxim,
|
||||
poll_concurrency=2,
|
||||
now=clock.now,
|
||||
)
|
||||
|
||||
await service.poll_messages()
|
||||
|
||||
assert boxim.peak_active_polls == 2
|
||||
db = session_factory()
|
||||
try:
|
||||
assert db.query(TakeoverCursor).filter(TakeoverCursor.initialized.is_(True)).count() == 2
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delayed_poll_still_schedules_recent_message(service_context):
|
||||
session_factory, service, boxim, clock = service_context
|
||||
await service.poll_messages()
|
||||
delayed_send_time = int(
|
||||
(clock.value - timedelta(seconds=150)).replace(tzinfo=timezone.utc).timestamp()
|
||||
* 1000
|
||||
)
|
||||
boxim.messages.append(
|
||||
{
|
||||
"id": 13,
|
||||
"localId": 13,
|
||||
"sendId": 200,
|
||||
"recvId": 100,
|
||||
"sendTime": delayed_send_time,
|
||||
"type": 0,
|
||||
"content": "排队后仍需回复",
|
||||
}
|
||||
)
|
||||
|
||||
await service.poll_messages()
|
||||
|
||||
db = session_factory()
|
||||
try:
|
||||
task = db.query(TakeoverReplyTask).one()
|
||||
assert task.status == "pending"
|
||||
assert task.scheduled_at == clock.now()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
with patch("routers.chat._resolve_reply", return_value={"answer": "已经收到"}):
|
||||
await service.process_reply_tasks()
|
||||
assert [item["content"] for item in boxim.sent] == ["已经收到"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_avatar_origin_message_never_schedules_a_reply(service_context):
|
||||
session_factory, service, boxim, clock = service_context
|
||||
await service.poll_and_process_messages()
|
||||
local_id = _avatar_local_id("peer-owner", "peer-trigger")
|
||||
boxim.messages.append(
|
||||
{"id": 15, "localId": local_id, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "另一端分身回复"}
|
||||
)
|
||||
|
||||
with patch("routers.chat._resolve_reply") as resolver:
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
resolver.assert_not_called()
|
||||
db = session_factory()
|
||||
try:
|
||||
event = db.query(TakeoverMessage).filter(TakeoverMessage.boxim_message_id == "15").one()
|
||||
assert event.is_avatar is True
|
||||
assert db.query(TakeoverReplyTask).count() == 0
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_peer_avatar_messages_are_excluded_from_later_human_context(service_context):
|
||||
_session_factory, service, boxim, clock = service_context
|
||||
await service.poll_and_process_messages()
|
||||
boxim.messages.append(
|
||||
{
|
||||
"id": 16,
|
||||
"localId": _avatar_local_id("peer-owner", "peer-trigger"),
|
||||
"sendId": 200,
|
||||
"recvId": 100,
|
||||
"sendTime": clock.millis(),
|
||||
"type": 0,
|
||||
"content": "分身生成的夸张长文",
|
||||
}
|
||||
)
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
clock.advance(1)
|
||||
boxim.messages.append(
|
||||
{
|
||||
"id": 17,
|
||||
"localId": 17,
|
||||
"sendId": 200,
|
||||
"recvId": 100,
|
||||
"sendTime": clock.millis(),
|
||||
"type": 0,
|
||||
"content": "真人的新问题",
|
||||
}
|
||||
)
|
||||
await service.poll_and_process_messages()
|
||||
clock.advance(3)
|
||||
with patch("routers.chat._resolve_reply", return_value={"answer": "正常回复"}) as resolver:
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
assert resolver.call_args.args[3] == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_owner_message_pauses_future_takeover_for_ten_minutes(service_context):
|
||||
session_factory, service, boxim, clock = service_context
|
||||
await service.poll_and_process_messages()
|
||||
boxim.messages.append(
|
||||
{"id": 18, "localId": 18, "sendId": 100, "recvId": 200, "sendTime": clock.millis(), "type": 0, "content": "我先来回复"}
|
||||
)
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
clock.advance(30)
|
||||
boxim.messages.append(
|
||||
{"id": 19, "localId": 19, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "收到"}
|
||||
)
|
||||
with patch("routers.chat._resolve_reply") as resolver:
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
resolver.assert_not_called()
|
||||
db = session_factory()
|
||||
try:
|
||||
assert db.query(TakeoverReplyTask).count() == 0
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_avatar_local_id_is_deterministic_and_self_describing():
|
||||
first = _avatar_local_id("owner", "message-1")
|
||||
assert first == _avatar_local_id("owner", "message-1")
|
||||
assert first != _avatar_local_id("owner", "message-2")
|
||||
assert first.startswith(AVATAR_LOCAL_ID_PREFIX)
|
||||
assert len(first) == 18
|
||||
assert first.isdigit()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_different_contacts_generate_without_blocking_each_other(service_context):
|
||||
session_factory, service, boxim, clock = service_context
|
||||
@@ -158,11 +508,11 @@ async def test_different_contacts_generate_without_blocking_each_other(service_c
|
||||
both_generating.wait()
|
||||
return {"answer": f"回复{prompt[-1]}"}
|
||||
|
||||
with patch("routers.chat._resolve_reply", side_effect=resolve):
|
||||
await service.poll_and_process_messages()
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
clock.advance(3)
|
||||
await service.process_reply_tasks()
|
||||
with patch("routers.chat._resolve_reply", side_effect=resolve):
|
||||
await service.poll_and_process_messages()
|
||||
assert {(item["peerId"], item["content"]) for item in boxim.sent} == {
|
||||
("200", "回复甲"),
|
||||
("300", "回复乙"),
|
||||
@@ -237,6 +587,34 @@ async def test_owner_message_cancels_pending_reply(service_context):
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_owner_message_in_final_second_wins_before_generation(service_context):
|
||||
session_factory, service, boxim, clock = service_context
|
||||
await service.poll_and_process_messages()
|
||||
boxim.messages.append(
|
||||
{"id": 23, "localId": 23, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "还在吗"}
|
||||
)
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
clock.advance(2)
|
||||
boxim.messages.append(
|
||||
{"id": 24, "localId": 24, "sendId": 100, "recvId": 200, "sendTime": clock.millis(), "type": 0, "content": "我来处理"}
|
||||
)
|
||||
clock.advance(1)
|
||||
with patch("routers.chat._resolve_reply") as resolver:
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
resolver.assert_not_called()
|
||||
assert boxim.sent == []
|
||||
db = session_factory()
|
||||
try:
|
||||
task = db.query(TakeoverReplyTask).filter(TakeoverReplyTask.trigger_message_id == "23").one()
|
||||
assert task.status == "cancelled"
|
||||
assert task.cancel_reason == "owner_replied"
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_quick_successive_messages_are_coalesced_into_one_reply(service_context):
|
||||
session_factory, service, boxim, clock = service_context
|
||||
@@ -244,19 +622,18 @@ async def test_quick_successive_messages_are_coalesced_into_one_reply(service_co
|
||||
boxim.messages.append(
|
||||
{"id": 31, "localId": 5, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "第一句"}
|
||||
)
|
||||
with patch("routers.chat._resolve_reply", return_value={"answer": "第一版"}):
|
||||
await service.poll_and_process_messages()
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
clock.advance(1)
|
||||
boxim.messages.append(
|
||||
{"id": 32, "localId": 6, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "第二句"}
|
||||
)
|
||||
await service.poll_and_process_messages()
|
||||
|
||||
clock.advance(3)
|
||||
with patch("routers.chat._resolve_reply", return_value={"answer": "合并回复"}) as resolver:
|
||||
await service.poll_and_process_messages()
|
||||
assert resolver.call_args.args[2] == "第一句\n第二句"
|
||||
|
||||
clock.advance(3)
|
||||
await service.poll_and_process_messages()
|
||||
assert [item["content"] for item in boxim.sent] == ["合并回复"]
|
||||
|
||||
db = session_factory()
|
||||
@@ -291,5 +668,41 @@ async def test_connection_failure_disables_takeover_and_stops_retrying(service_c
|
||||
boxim.exchange_access_token.assert_awaited_once_with("prod-huihui-token")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transient_connection_failure_keeps_takeover_and_recovers(service_context):
|
||||
session_factory, service, boxim, _ = service_context
|
||||
boxim.exchange_access_token = AsyncMock(
|
||||
side_effect=[
|
||||
BoxIMError("连接超时"),
|
||||
{"accessToken": "box-token", "accessTokenExpiresIn": 3600},
|
||||
]
|
||||
)
|
||||
|
||||
await service.poll_messages()
|
||||
|
||||
db = session_factory()
|
||||
try:
|
||||
avatar = db.query(Avatar).one()
|
||||
cursor = db.query(TakeoverCursor).one()
|
||||
assert "takeover" in avatar.config["authorizationPermissions"]
|
||||
assert cursor.initialized is False
|
||||
assert "暂时连接失败" in cursor.last_error
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
await service.poll_messages()
|
||||
|
||||
db = session_factory()
|
||||
try:
|
||||
avatar = db.query(Avatar).one()
|
||||
cursor = db.query(TakeoverCursor).one()
|
||||
assert "takeover" in avatar.config["authorizationPermissions"]
|
||||
assert cursor.initialized is True
|
||||
assert cursor.last_error == ""
|
||||
finally:
|
||||
db.close()
|
||||
assert boxim.exchange_access_token.await_count == 2
|
||||
|
||||
|
||||
def test_plain_text_reply_removes_markdown_and_empty_lines():
|
||||
assert _plain_text_reply("## 建议\n\n**不能自行用药**\n`必要时就医`") == "建议\n不能自行用药\n必要时就医"
|
||||
|
||||
@@ -0,0 +1,98 @@
|
||||
import io
|
||||
import json
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
from PIL import Image
|
||||
|
||||
from services.chat_model_config import ChatModelConfig
|
||||
from services.vision_service import (
|
||||
ImageValidationError,
|
||||
build_attachment_warning,
|
||||
call_vision_model,
|
||||
parse_vision_analysis,
|
||||
prepare_image,
|
||||
)
|
||||
|
||||
|
||||
def _image_bytes(fmt="PNG", size=(120, 80)):
|
||||
output = io.BytesIO()
|
||||
Image.new("RGB", size, "#f97316").save(output, format=fmt)
|
||||
return output.getvalue()
|
||||
|
||||
|
||||
def _config():
|
||||
return ChatModelConfig(
|
||||
api_base_url="https://model.test/v1",
|
||||
api_key="secret-key",
|
||||
model="chat-model",
|
||||
max_tokens=1024,
|
||||
timeout_seconds=30,
|
||||
vision_model="vision-model",
|
||||
ocr_model="ocr-model",
|
||||
vision_max_tokens=2048,
|
||||
vision_timeout_seconds=90,
|
||||
source="test",
|
||||
)
|
||||
|
||||
|
||||
def test_prepare_image_validates_and_reencodes_without_metadata():
|
||||
prepared = prepare_image(_image_bytes())
|
||||
|
||||
assert prepared.mime_type == "image/jpeg"
|
||||
assert prepared.width == 120
|
||||
assert prepared.height == 80
|
||||
with Image.open(io.BytesIO(prepared.data)) as image:
|
||||
assert image.format == "JPEG"
|
||||
assert not image.getexif()
|
||||
|
||||
|
||||
def test_prepare_image_rejects_non_image_content():
|
||||
with pytest.raises(ImageValidationError, match="格式无效"):
|
||||
prepare_image(b"not-an-image")
|
||||
|
||||
|
||||
def test_vision_request_uses_openai_compatible_image_content():
|
||||
response = Mock()
|
||||
response.raise_for_status.return_value = None
|
||||
response.json.return_value = {
|
||||
"choices": [{"message": {"content": '{"category":"general_image"}'}}],
|
||||
"usage": {"total_tokens": 88},
|
||||
}
|
||||
prepared = prepare_image(_image_bytes())
|
||||
|
||||
with patch("services.vision_service.httpx.post", return_value=response) as request:
|
||||
result = call_vision_model(
|
||||
prepared,
|
||||
_config(),
|
||||
model="vision-model",
|
||||
prompt="describe",
|
||||
json_output=True,
|
||||
)
|
||||
|
||||
payload = request.call_args.kwargs["json"]
|
||||
content = payload["messages"][0]["content"]
|
||||
assert payload["model"] == "vision-model"
|
||||
assert payload["response_format"] == {"type": "json_object"}
|
||||
assert content[0]["type"] == "image_url"
|
||||
assert content[0]["image_url"]["url"].startswith("data:image/jpeg;base64,")
|
||||
assert content[1] == {"type": "text", "text": "describe"}
|
||||
assert result["usage"]["total_tokens"] == 88
|
||||
|
||||
|
||||
def test_parse_medical_analysis_and_build_warning():
|
||||
analysis = parse_vision_analysis(json.dumps({
|
||||
"category": "medical_document",
|
||||
"summary": "血常规报告",
|
||||
"visible_text": "白细胞 11.2",
|
||||
"key_facts": ["白细胞偏高"],
|
||||
"uncertainties": ["日期模糊"],
|
||||
"medical": {"document_type": "检验报告"},
|
||||
}, ensure_ascii=False))
|
||||
|
||||
assert analysis["category"] == "medical_document"
|
||||
assert analysis["medical"]["document_type"] == "检验报告"
|
||||
warning = build_attachment_warning(analysis, ocr_failed=True)
|
||||
assert "日期模糊" in warning
|
||||
assert "人工核对" in warning
|
||||
assert "不能替代医生诊断" in warning
|
||||
@@ -39,6 +39,8 @@ HUIHUI_ACCESS_ID=<production-access-id>
|
||||
HUIHUI_ACCESS_SECRET=<production-access-secret>
|
||||
HUIHUI_CLIENT_CODE=<production-client-code>
|
||||
BOXIM_TIMEOUT_SECONDS=20
|
||||
BOXIM_POLL_CONCURRENCY=8
|
||||
BOXIM_MAX_MESSAGE_AGE_SECONDS=600
|
||||
HUIHUI_PAYMENT_BASE_URL=https://open.99hui.com/api/payment-v3
|
||||
HUIHUI_PAYMENT_CALLBACK_BASE_URL=https://digital.99hui.com
|
||||
HUIHUI_PAYMENT_CALLBACK_SECRET=<至少32位随机密钥>
|
||||
@@ -47,10 +49,26 @@ HUIHUI_PAYMENT_TIMEOUT_SECONDS=30
|
||||
DATABASE_URL=sqlite:////data/avatar.db
|
||||
UPLOAD_DIR=/data/uploads
|
||||
CHAT_MODEL_CONFIG_URL=http://<huihuisquare-api>/api/ai-models/runtime/digital-avatar
|
||||
EMBEDDING_API_URL=https://dashscope.aliyuncs.com/compatible-mode/v1
|
||||
EMBEDDING_API_KEY=<production-embedding-api-key>
|
||||
EMBEDDING_MODEL=text-embedding-v3
|
||||
EMBEDDING_BATCH_SIZE=10
|
||||
|
||||
VISION_MODEL=qwen3.6-flash
|
||||
VISION_OCR_MODEL=qwen-vl-ocr
|
||||
VISION_MAX_OUTPUT_TOKENS=2048
|
||||
VISION_TIMEOUT_SECONDS=90
|
||||
VISION_TOKEN_RESERVE=12000
|
||||
CHAT_IMAGE_MAX_BYTES=8388608
|
||||
CHAT_IMAGE_MAX_PIXELS=16000000
|
||||
CHAT_ATTACHMENT_RETENTION_HOURS=24
|
||||
CHAT_ATTACHMENT_CLEANUP_MINUTES=60
|
||||
```
|
||||
|
||||
如生产 AI 配置中心不可用,还应提供当前项目支持的 `OPENAI_API_KEY`、`OPENAI_BASE_URL`、`CHAT_MODEL` 等兜底配置。`/data` 必须挂载持久卷,数据库与知识库文件不可存放在容器临时层。
|
||||
|
||||
`EMBEDDING_API_URL` 同时支持 OpenAI 兼容基础地址(如上面的 `/v1`)和完整的 `/v1/embeddings` 地址,后端会统一请求 `/embeddings`。发布后必须在后端容器内执行一次最小向量探针,确认返回向量数量和维度,而不能只检查 `/api/health`。
|
||||
|
||||
积分充值使用会会支付体系的 `payment-v3/payment/pay`,渠道值为 `WECHAT` / `ALIPAY`,端内支付场景为 `APP`,微信内 H5 使用 `JSAPI`。`HUIHUI_PAYMENT_CALLBACK_SECRET` 只用于为每笔订单生成 HMAC 回调签名,不会发送到前端或直接出现在回调地址中。支付回调确认状态成功且金额与套餐价格完全一致后才增加积分,重复回调不会重复到账。
|
||||
|
||||
## 3. 构建与发布
|
||||
@@ -74,6 +92,7 @@ docker compose build --pull avatar-backend avatar-frontend
|
||||
docker compose up -d avatar-backend avatar-frontend
|
||||
docker compose ps
|
||||
curl -fsS http://127.0.0.1:8099/api/health
|
||||
docker compose exec avatar-backend python -c 'import embeddings; v=embeddings.embed(["部署向量探针"]); print(len(v), len(v[0]))'
|
||||
```
|
||||
|
||||
生产编排应把示例中的测试端口改为内网暴露,由统一 HTTPS 网关接入。后端暂时使用 SQLite,必须保持单实例写入;若扩展为多后端实例,应先迁移到 PostgreSQL,并把延迟接管任务改为共享队列。
|
||||
@@ -102,7 +121,9 @@ location /api/ {
|
||||
}
|
||||
```
|
||||
|
||||
`proxy_buffering off` 用于数字分身 SSE 流式吐字,`client_max_body_size` 用于知识库文件上传。网关和应用日志必须关闭完整 URL 查询参数记录,任何异常日志都不得输出 token、Authorization 或平台密钥。建议同时设置严格的 `Referrer-Policy: no-referrer`。
|
||||
`proxy_buffering off` 用于数字分身 SSE 流式吐字,`client_max_body_size` 同时用于知识库文件和聊天图片上传。应用只保存图片识别结果,不保存原图;识别结果 24 小时失效,后台默认每小时清理一次。公开分享图片识别会消耗分身所有者积分,生产网关应针对 `/api/public/avatar/*/chat/images` 设置每 IP 和每分享令牌的上传频率限制,防止恶意消耗。
|
||||
|
||||
网关和应用日志必须关闭完整 URL 查询参数记录,任何异常日志都不得输出 token、Authorization、图片 Base64、病例正文或平台密钥。建议同时设置严格的 `Referrer-Policy: no-referrer`。
|
||||
|
||||
## 5. 发布验收
|
||||
|
||||
@@ -112,10 +133,13 @@ location /api/ {
|
||||
4. A、B 两个会会用户分别进入时只能看到各自的数字分身与知识库,不会继承上一用户缓存。
|
||||
5. 使用过期或伪造 token 时进入登录页并显示凭证失效,不得继续访问旧用户数据。
|
||||
6. 分身聊天 SSE 逐段输出正常,Markdown 正常渲染,知识库优先级和积分扣费正常。
|
||||
7. 开启 BOXIM 主动接管后保持在线,收到消息、三秒回复、已读回执和主人发言暂停均正常。
|
||||
7. 开启 BOXIM 主动接管后保持在线,默认三分钟回复、自定义等待时间、已读回执、分身防回环和主人发言暂停均正常。
|
||||
8. 重建容器后数据库、头像、知识库文档仍存在,`/api/health` 返回成功。
|
||||
9. `https://digital.99hui.com/api/health` 可访问,证书域名和有效期正确,HTTP 自动跳转 HTTPS。
|
||||
10. 微信和支付宝各创建一笔最小套餐订单,未付款时积分不变;支付成功后回调到账一次,重复回调积分不重复增加。
|
||||
11. 私聊和公开分享各上传 JPG、PNG、WebP 图片并完成追问;上传非图片、超过 8MB 或跨分身附件时必须拒绝。
|
||||
12. 病例图片可以提取可见文字并标记待核对内容,医学影像不作确定诊断;视觉与 OCR 调用分别扣减积分。
|
||||
13. 检查服务器上传目录不残留聊天原图,数据库过期图片识别记录在清理周期后删除,日志不出现 Base64 或病例正文。
|
||||
|
||||
## 6. 回滚
|
||||
|
||||
|
||||
@@ -0,0 +1,184 @@
|
||||
# 数字分身图片与病例理解详细设计
|
||||
|
||||
## 1. 目标与边界
|
||||
|
||||
本功能让数字分身在私聊和公开分享聊天中接收图片,并围绕图片内容继续使用现有的“标准答题对 -> 分身独立知识库 -> Qwen 兼容模型”链路回答。
|
||||
|
||||
第一期支持 JPEG、PNG、WebP,覆盖以下场景:
|
||||
|
||||
1. 普通照片、截图、图表和界面图片的内容理解。
|
||||
2. 病例、处方、检查单、检验报告等图片文档的文字和表格提取。
|
||||
3. X 光、CT、MRI 等医学影像的客观可见内容描述。
|
||||
|
||||
第一期不把通用视觉模型的输出当作医学诊断,不自动把图片或病例写入知识库,不保存原图供长期访问,也不支持 DICOM 原始影像。
|
||||
|
||||
## 2. 核心原则
|
||||
|
||||
- **资料优先级不变**:标准答题对最高,分身独立知识库其次,图片识别结果属于待核对的会话资料,最后才由模型组织表达。
|
||||
- **病例最小留存**:应用不把原图写入业务存储,上传内容在内存中归一化并调用视觉服务;数据库只保存结构化结果和必要元数据。
|
||||
- **严格隔离**:每条图片记录必须绑定 `avatar_id`,私聊校验分身所有者,公开聊天校验分享令牌对应的分身。
|
||||
- **不确定性显式化**:OCR 看不清、表格列错位、医学影像无法确认时必须指出待核对项,不允许补齐缺失内容。
|
||||
- **可计量**:视觉理解和病例 OCR 分别计入分身所有者的积分消耗,失败时释放预留积分。
|
||||
- **可降级**:OCR 失败但通用视觉结果有效时仍可回答;视觉主调用失败则不进入聊天发送。
|
||||
|
||||
## 3. 总体流程
|
||||
|
||||
```text
|
||||
用户选择图片
|
||||
-> 前端本地预览
|
||||
-> 私聊/公开图片上传接口
|
||||
-> 文件大小、MIME、真实格式、像素数校验
|
||||
-> 自动旋转、缩放、去 EXIF、统一 JPEG
|
||||
-> 通用视觉模型分类并输出结构化 JSON
|
||||
-> 若为病例/检查单,再调用 OCR 模型精确转录
|
||||
-> 保存结构化结果,不持久化原图
|
||||
-> 返回 attachmentId
|
||||
-> 用户发送文字 + attachmentIds
|
||||
-> 标准答题对匹配
|
||||
-> 用文字 + 图片提取结果检索独立知识库
|
||||
-> 把标准答案、知识片段、图片资料注入系统上下文
|
||||
-> Qwen SSE 流式回答
|
||||
```
|
||||
|
||||
## 4. 模型编排
|
||||
|
||||
### 4.1 通用视觉模型
|
||||
|
||||
默认 `qwen3.6-flash`,可在后台数字分身专用模型配置中修改。输入为归一化后的 Base64 Data URL,要求返回 JSON:
|
||||
|
||||
```json
|
||||
{
|
||||
"category": "general_image|document|medical_document|medical_image",
|
||||
"summary": "客观、完整的图片描述",
|
||||
"visible_text": "图片中可确认的文字",
|
||||
"key_facts": ["事实1", "事实2"],
|
||||
"uncertainties": ["无法确认的内容"],
|
||||
"medical": {
|
||||
"document_type": "",
|
||||
"patient_info": {},
|
||||
"chief_complaint": "",
|
||||
"findings": [],
|
||||
"measurements": [],
|
||||
"doctor_advice": ""
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
模型提示词禁止诊断、补全被遮挡文字、猜测患者身份和输出模型信息。
|
||||
|
||||
### 4.2 病例 OCR
|
||||
|
||||
当 `category=medical_document` 时追加调用 `qwen-vl-ocr`,按原布局转录文字和表格。OCR 文本优先替换通用视觉输出中的 `visible_text`,但保留通用视觉模型提供的分类、摘要和不确定项。
|
||||
|
||||
### 4.3 医学影像
|
||||
|
||||
当 `category=medical_image` 时只保存客观描述,不输出疾病结论、分期、用药或治疗方案。聊天提示词必须要求结合正规影像报告和医生意见,并显示“图片识别结果仅供辅助,不能替代医生诊断”。
|
||||
|
||||
## 5. 数据模型
|
||||
|
||||
新增 `chat_attachments`:
|
||||
|
||||
| 字段 | 说明 |
|
||||
|---|---|
|
||||
| `id` | 不可猜测的附件 ID |
|
||||
| `avatar_id` | 所属数字分身,强制隔离 |
|
||||
| `filename` | 原文件名,去除路径 |
|
||||
| `mime_type` / `file_size` | 上传元数据 |
|
||||
| `status` | `processing / ready / failed` |
|
||||
| `category` | 图片分类 |
|
||||
| `summary` | 通用视觉摘要 |
|
||||
| `extracted_text` | 可确认文字/OCR 结果 |
|
||||
| `structured_data` | 结构化 JSON |
|
||||
| `warning` | 不确定项和医学提示 |
|
||||
| `vision_model` / `ocr_model` | 实际调用模型 |
|
||||
| `created_at` / `used_at` | 创建和最近使用时间 |
|
||||
|
||||
不保存公开原图 URL。应用层不落盘原图;框架上传缓冲在请求结束时关闭,处理结果在 24 小时后自动清理。
|
||||
|
||||
## 6. API 设计
|
||||
|
||||
### 6.1 上传并解析
|
||||
|
||||
- `POST /api/avatar/{avatar_id}/chat/images`
|
||||
- `POST /api/public/avatar/{share_token}/chat/images`
|
||||
- `multipart/form-data: file`
|
||||
|
||||
成功返回:
|
||||
|
||||
```json
|
||||
{
|
||||
"id": "attachment-id",
|
||||
"filename": "病例.jpg",
|
||||
"status": "ready",
|
||||
"category": "medical_document",
|
||||
"summary": "门诊检查单",
|
||||
"warning": "部分手写内容需要人工核对"
|
||||
}
|
||||
```
|
||||
|
||||
### 6.2 聊天
|
||||
|
||||
原聊天接口增加:
|
||||
|
||||
```json
|
||||
{
|
||||
"message": "请帮我看看异常指标",
|
||||
"attachmentIds": ["attachment-id"],
|
||||
"history": [
|
||||
{"role": "user", "content": "上一条问题", "attachmentIds": ["attachment-id"]}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
当前消息最多 3 张图,历史最多引用最近 3 个不同附件。后端只读取与当前 `avatar_id` 相同且状态为 `ready` 的记录。
|
||||
|
||||
## 7. 安全与隐私
|
||||
|
||||
- 单图最大 8MB,解码后最大 1600 万像素,最长边归一化到 4096 像素以内。
|
||||
- 使用 Pillow 验证真实图片格式并防止解压炸弹;重新编码时清除 EXIF、GPS 和其他元数据。
|
||||
- 图片不会写入 FastAPI `StaticFiles` 或知识库目录,模型请求和日志不得输出 Base64 内容。
|
||||
- 日志只记录附件 ID、分身 ID、状态、耗时和模型,不记录图片 Base64、OCR 全文、病例内容或 API Key。
|
||||
- 公开分享上传仍消耗分身所有者积分;余额不足时拒绝视觉调用。
|
||||
- 生产环境需要补充用户授权、数据处理协议、存储地域和模型供应商留存策略确认。
|
||||
|
||||
## 8. 前端交互
|
||||
|
||||
- 输入框左侧增加图片按钮,支持相册选择和移动端拍照。
|
||||
- 选择后显示本地缩略图和“正在识别图片”,识别完成前禁止发送。
|
||||
- 用户可删除待发送图片;发送后图片保留在当前会话气泡中,但刷新页面后不恢复原图。
|
||||
- 病例和医学影像在输入区及回答下方显示辅助提示,不使用恐吓式红色告警。
|
||||
- 上传或识别失败时保留文字输入,明确提示重新选择图片,不产生空白消息。
|
||||
|
||||
## 9. 配置
|
||||
|
||||
数字分身专用模型配置新增:
|
||||
|
||||
- `vision_model_version`,默认 `qwen3.6-flash`
|
||||
- `ocr_model_version`,默认 `qwen-vl-ocr`
|
||||
|
||||
环境变量兜底:
|
||||
|
||||
```dotenv
|
||||
VISION_MODEL=qwen3.6-flash
|
||||
VISION_OCR_MODEL=qwen-vl-ocr
|
||||
VISION_MAX_OUTPUT_TOKENS=2048
|
||||
VISION_TIMEOUT_SECONDS=90
|
||||
VISION_TOKEN_RESERVE=12000
|
||||
CHAT_IMAGE_MAX_BYTES=8388608
|
||||
CHAT_IMAGE_MAX_PIXELS=16000000
|
||||
CHAT_ATTACHMENT_RETENTION_HOURS=24
|
||||
CHAT_ATTACHMENT_CLEANUP_MINUTES=60
|
||||
```
|
||||
|
||||
视觉调用复用数字分身专用配置的 `api_base_url` 和 `api_key`,不额外复制密钥。
|
||||
|
||||
## 10. 验收标准
|
||||
|
||||
1. 普通照片、截图和图表能够返回与图片一致的描述并支持追问。
|
||||
2. 病例图片可以提取标题、患者字段、检查结果、异常指标和医生意见,模糊内容明确标记待核对。
|
||||
3. 上传后服务器业务目录不残留原图,响应和日志不包含 Base64 或完整病例正文。
|
||||
4. A 分身无法引用 B 分身附件;公开分享令牌无法访问其他分身附件。
|
||||
5. 有图片时标准答题对仍作为最高优先级事实,知识库命中次之。
|
||||
6. 视觉与 OCR 积分分别结算,失败调用释放预留积分。
|
||||
7. SSE 打字效果、Markdown、用户头像、公开分享和纯文本聊天均无回归。
|
||||
8. CT、MRI、X 光回答不作确定诊断,并显示人工复核提示。
|
||||
@@ -191,19 +191,28 @@ export type AvatarPermission = 'friend' | 'chat' | 'publish' | 'browse' | 'inter
|
||||
export interface AvatarPermissionSettings {
|
||||
avatarId: string
|
||||
permissions: AvatarPermission[]
|
||||
takeoverReplyDelaySeconds: number
|
||||
disabledAvatarIds?: string[]
|
||||
}
|
||||
|
||||
export const getAvatarPermissionSettings = (avatarId: string) =>
|
||||
request.get<AvatarPermissionSettings>(`/avatar/${avatarId}/permission-settings`)
|
||||
|
||||
export const updateAvatarPermissionSettings = (avatarId: string, permissions: AvatarPermission[]) =>
|
||||
request.put<AvatarPermissionSettings>(`/avatar/${avatarId}/permission-settings`, { permissions })
|
||||
export const updateAvatarPermissionSettings = (
|
||||
avatarId: string,
|
||||
permissions: AvatarPermission[],
|
||||
takeoverReplyDelaySeconds: number
|
||||
) => request.put<AvatarPermissionSettings>(`/avatar/${avatarId}/permission-settings`, {
|
||||
permissions,
|
||||
takeoverReplyDelaySeconds,
|
||||
})
|
||||
|
||||
export interface TakeoverStatus {
|
||||
enabled: boolean
|
||||
status: 'disabled' | 'connecting' | 'ready' | 'needs_login' | 'error'
|
||||
message: string
|
||||
pendingCount: number
|
||||
takeoverReplyDelaySeconds: number
|
||||
lastPolledAt: string | null
|
||||
}
|
||||
|
||||
@@ -365,15 +374,35 @@ export const searchKnowledge = (avatarId: string, q: string, topK = 5) =>
|
||||
export interface ChatMessage {
|
||||
role: 'user' | 'assistant'
|
||||
content: string
|
||||
attachmentIds?: string[]
|
||||
}
|
||||
|
||||
export interface ChatResponse {
|
||||
answer: string
|
||||
source: 'qa' | 'knowledge' | 'qwen'
|
||||
source: 'qa' | 'knowledge' | 'vision' | 'qwen'
|
||||
references?: Array<{ docId?: string; filename?: string; fileType?: string; snippet?: string; score?: number }>
|
||||
}
|
||||
|
||||
export const sendAvatarChat = (avatarId: string, payload: { message: string; history?: ChatMessage[] }) =>
|
||||
export interface ChatAttachment {
|
||||
id: string
|
||||
avatarId: string
|
||||
filename: string
|
||||
mimeType: string
|
||||
fileSize: number
|
||||
status: 'processing' | 'ready' | 'failed'
|
||||
category: 'general_image' | 'document' | 'medical_document' | 'medical_image'
|
||||
summary: string
|
||||
warning: string
|
||||
expiresAt: string
|
||||
}
|
||||
|
||||
export interface ChatPayload {
|
||||
message: string
|
||||
attachmentIds?: string[]
|
||||
history?: ChatMessage[]
|
||||
}
|
||||
|
||||
export const sendAvatarChat = (avatarId: string, payload: ChatPayload) =>
|
||||
request.post<ChatResponse>(`/avatar/${avatarId}/chat`, payload)
|
||||
|
||||
export interface PublicAvatar {
|
||||
@@ -392,19 +421,41 @@ export const createAvatarShareLink = (avatarId: string) =>
|
||||
export const getPublicAvatar = (shareToken: string) =>
|
||||
request.get<PublicAvatar>(`/public/avatar/${shareToken}`)
|
||||
|
||||
export const sendPublicAvatarChat = (shareToken: string, payload: { message: string; history?: ChatMessage[] }) =>
|
||||
export const sendPublicAvatarChat = (shareToken: string, payload: ChatPayload) =>
|
||||
request.post<ChatResponse>(`/public/avatar/${shareToken}/chat`, payload)
|
||||
|
||||
const imageForm = (file: File) => {
|
||||
const form = new FormData()
|
||||
form.append('file', file)
|
||||
return form
|
||||
}
|
||||
|
||||
export const uploadAvatarChatImage = (avatarId: string, file: File) =>
|
||||
request.post<ChatAttachment>(`/avatar/${avatarId}/chat/images`, imageForm(file), {
|
||||
headers: { 'Content-Type': 'multipart/form-data' },
|
||||
timeout: 120000
|
||||
})
|
||||
|
||||
export const uploadPublicAvatarChatImage = (shareToken: string, file: File) =>
|
||||
request.post<ChatAttachment>(`/public/avatar/${shareToken}/chat/images`, imageForm(file), {
|
||||
headers: { 'Content-Type': 'multipart/form-data' },
|
||||
timeout: 120000
|
||||
})
|
||||
|
||||
type ChatStreamHandlers = {
|
||||
onMeta: (meta: Pick<ChatResponse, 'source' | 'references'>) => void
|
||||
onDelta: (content: string) => void
|
||||
}
|
||||
|
||||
const streamChat = async (path: string, payload: { message: string; history?: ChatMessage[] }, handlers: ChatStreamHandlers) => {
|
||||
const streamChat = async (path: string, payload: ChatPayload, handlers: ChatStreamHandlers) => {
|
||||
const headers: Record<string, string> = { 'Content-Type': 'application/json', Accept: 'text/event-stream' }
|
||||
if (_authToken) headers.Authorization = `Bearer ${_authToken}`
|
||||
const response = await fetch(`${resolveBaseURL()}${path}`, { method: 'POST', headers, body: JSON.stringify(payload) })
|
||||
if (!response.ok || !response.body) throw new Error(`对话请求失败(${response.status})`)
|
||||
if (!response.ok) {
|
||||
const errorBody = await response.json().catch(() => null)
|
||||
throw new Error(errorBody?.detail || errorBody?.message || `对话请求失败(${response.status})`)
|
||||
}
|
||||
if (!response.body) throw new Error('对话响应为空,请稍后重试')
|
||||
|
||||
const reader = response.body.getReader()
|
||||
const decoder = new TextDecoder()
|
||||
@@ -427,10 +478,10 @@ const streamChat = async (path: string, payload: { message: string; history?: Ch
|
||||
}
|
||||
}
|
||||
|
||||
export const streamAvatarChat = (avatarId: string, payload: { message: string; history?: ChatMessage[] }, handlers: ChatStreamHandlers) =>
|
||||
export const streamAvatarChat = (avatarId: string, payload: ChatPayload, handlers: ChatStreamHandlers) =>
|
||||
streamChat(`/avatar/${avatarId}/chat/stream`, payload, handlers)
|
||||
|
||||
export const streamPublicAvatarChat = (shareToken: string, payload: { message: string; history?: ChatMessage[] }, handlers: ChatStreamHandlers) =>
|
||||
export const streamPublicAvatarChat = (shareToken: string, payload: ChatPayload, handlers: ChatStreamHandlers) =>
|
||||
streamChat(`/public/avatar/${shareToken}/chat/stream`, payload, handlers)
|
||||
|
||||
// ==================== 会会用户资料 API ====================
|
||||
|
||||
@@ -10,6 +10,7 @@ import {
|
||||
type SmsLoginResult,
|
||||
type UserProfile
|
||||
} from '@/api'
|
||||
import { clearHuihuiEmbeddedMode, markHuihuiEmbeddedMode } from '@/utils/embed-mode'
|
||||
|
||||
const TOKEN_KEY = 'hh_app_token'
|
||||
const USER_KEY = 'hh_app_user'
|
||||
@@ -70,16 +71,23 @@ export const useUserStore = defineStore('smsuser', () => {
|
||||
|
||||
// 短信登录
|
||||
const login = async (phone: string, code: string) => {
|
||||
return acceptLogin(await loginBySms(phone, code))
|
||||
const result = await loginBySms(phone, code)
|
||||
clearHuihuiEmbeddedMode()
|
||||
return acceptLogin(result)
|
||||
}
|
||||
|
||||
// 账号密码登录
|
||||
const loginByPwd = async (account: string, password: string) => {
|
||||
return acceptLogin(await loginByPassword(account, password))
|
||||
const result = await loginByPassword(account, password)
|
||||
clearHuihuiEmbeddedMode()
|
||||
return acceptLogin(result)
|
||||
}
|
||||
|
||||
const loginByToken = async (huihuiToken: string) =>
|
||||
acceptLogin(await loginByHuihuiToken(huihuiToken))
|
||||
const loginByToken = async (huihuiToken: string) => {
|
||||
const result = await loginByHuihuiToken(huihuiToken)
|
||||
markHuihuiEmbeddedMode()
|
||||
return acceptLogin(result)
|
||||
}
|
||||
|
||||
// 退出
|
||||
const logout = async () => {
|
||||
@@ -88,6 +96,7 @@ export const useUserStore = defineStore('smsuser', () => {
|
||||
} catch {
|
||||
/* 忽略网络错误,本地清除即可 */
|
||||
}
|
||||
clearHuihuiEmbeddedMode()
|
||||
clearSession()
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
const HUIHUI_EMBED_MODE_KEY = 'hh_huihui_embed_mode'
|
||||
|
||||
export function markHuihuiEmbeddedMode(): void {
|
||||
sessionStorage.setItem(HUIHUI_EMBED_MODE_KEY, '1')
|
||||
}
|
||||
|
||||
export function clearHuihuiEmbeddedMode(): void {
|
||||
sessionStorage.removeItem(HUIHUI_EMBED_MODE_KEY)
|
||||
}
|
||||
|
||||
export function isHuihuiEmbeddedMode(): boolean {
|
||||
return sessionStorage.getItem(HUIHUI_EMBED_MODE_KEY) === '1'
|
||||
}
|
||||
@@ -1,6 +1,6 @@
|
||||
<template>
|
||||
<div class="authorization-page">
|
||||
<header class="page-header">
|
||||
<div class="authorization-page" :class="{ embedded: isEmbedded }">
|
||||
<header v-if="!isEmbedded" class="page-header">
|
||||
<button class="back-button" type="button" aria-label="返回数字分身管理" @click="goBack">
|
||||
<svg viewBox="0 0 24 24" aria-hidden="true">
|
||||
<path d="m15 18-6-6 6-6" />
|
||||
@@ -64,7 +64,7 @@
|
||||
<span class="permission-copy">
|
||||
<strong>{{ item.title }}</strong>
|
||||
<small>
|
||||
{{ item.description }}
|
||||
{{ item.key === 'takeover' ? takeoverDescription : item.description }}
|
||||
<span
|
||||
v-if="item.key === 'takeover' && takeoverConnectionLabel"
|
||||
class="connection-state"
|
||||
@@ -79,6 +79,33 @@
|
||||
</button>
|
||||
</section>
|
||||
|
||||
<section v-if="permissionState.takeover" class="takeover-delay-card" aria-label="自动回复等待时间">
|
||||
<div class="delay-heading">
|
||||
<div>
|
||||
<strong>自动回复等待时间</strong>
|
||||
<small>等待期间主人发言会取消本次回复,最短 3 秒</small>
|
||||
</div>
|
||||
<span>{{ formattedTakeoverDelay }}</span>
|
||||
</div>
|
||||
<div class="delay-control">
|
||||
<input
|
||||
v-model.number="takeoverDelayValue"
|
||||
type="number"
|
||||
inputmode="numeric"
|
||||
step="1"
|
||||
:min="takeoverDelayUnit === 'minutes' ? 1 : 3"
|
||||
:max="takeoverDelayUnit === 'minutes' ? 1440 : 86400"
|
||||
aria-label="等待时间"
|
||||
:disabled="loading || saving"
|
||||
@blur="normalizeTakeoverDelay"
|
||||
/>
|
||||
<select v-model="takeoverDelayUnit" aria-label="等待时间单位" :disabled="loading || saving">
|
||||
<option value="seconds">秒</option>
|
||||
<option value="minutes">分钟</option>
|
||||
</select>
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<p v-if="errorMessage" class="error-message" role="alert">{{ errorMessage }}</p>
|
||||
</template>
|
||||
|
||||
@@ -95,6 +122,9 @@
|
||||
</main>
|
||||
|
||||
<footer v-if="activeAvatarId" class="save-area">
|
||||
<button v-if="isEmbedded" class="footer-back-button" type="button" :disabled="saving" @click="goBack">
|
||||
返回
|
||||
</button>
|
||||
<button class="save-button" type="button" :disabled="loading || saving" @click="saveSettings()">
|
||||
<span v-if="saving" class="saving-spinner" aria-hidden="true"></span>
|
||||
{{ saving ? '保存中...' : '保存授权设置' }}
|
||||
@@ -119,6 +149,7 @@ import {
|
||||
} from '@/api'
|
||||
import { useAvatarStore } from '@/store/avatar'
|
||||
import { pickScopedAvatarId } from '@/utils/avatar-page-data.js'
|
||||
import { isHuihuiEmbeddedMode } from '@/utils/embed-mode'
|
||||
|
||||
type PermissionState = Record<AvatarPermission, boolean>
|
||||
|
||||
@@ -126,6 +157,7 @@ const router = useRouter()
|
||||
const route = useRoute()
|
||||
const avatarStore = useAvatarStore()
|
||||
const avatarId = computed(() => pickScopedAvatarId(route.params.avatarId, avatarStore.currentAvatarId, avatarStore.avatars))
|
||||
const isEmbedded = isHuihuiEmbeddedMode()
|
||||
|
||||
const permissionItems: Array<{
|
||||
key: AvatarPermission
|
||||
@@ -166,7 +198,7 @@ const permissionItems: Array<{
|
||||
{
|
||||
key: 'takeover',
|
||||
title: '分身主动接管聊天回复',
|
||||
description: '收到私聊消息 3 秒后回复,主人发言时暂停',
|
||||
description: '收到私聊消息后按设定时间回复,主人发言时暂停',
|
||||
tone: 'cyan',
|
||||
},
|
||||
]
|
||||
@@ -185,6 +217,8 @@ const saving = ref(false)
|
||||
const errorMessage = ref('')
|
||||
const toastMessage = ref('')
|
||||
const takeoverStatus = ref<TakeoverStatus | null>(null)
|
||||
const takeoverDelayValue = ref(3)
|
||||
const takeoverDelayUnit = ref<'seconds' | 'minutes'>('minutes')
|
||||
let toastTimer: number | undefined
|
||||
let takeoverStatusTimer: number | undefined
|
||||
|
||||
@@ -205,6 +239,38 @@ const takeoverConnectionTone = computed(() => {
|
||||
return 'connecting'
|
||||
})
|
||||
|
||||
const takeoverDelaySeconds = computed(() => {
|
||||
const value = Math.trunc(Number(takeoverDelayValue.value) || 0)
|
||||
return takeoverDelayUnit.value === 'minutes' ? value * 60 : value
|
||||
})
|
||||
|
||||
const formattedTakeoverDelay = computed(() => {
|
||||
const seconds = takeoverDelaySeconds.value
|
||||
if (seconds > 0 && seconds % 60 === 0) return `${seconds / 60} 分钟`
|
||||
return `${seconds} 秒`
|
||||
})
|
||||
|
||||
const takeoverDescription = computed(() =>
|
||||
`收到私聊消息 ${formattedTakeoverDelay.value}后回复,主人发言时暂停`
|
||||
)
|
||||
|
||||
const applyTakeoverDelay = (seconds: number) => {
|
||||
const normalized = Number.isFinite(seconds) && seconds >= 3 ? Math.trunc(seconds) : 180
|
||||
if (normalized % 60 === 0) {
|
||||
takeoverDelayUnit.value = 'minutes'
|
||||
takeoverDelayValue.value = normalized / 60
|
||||
} else {
|
||||
takeoverDelayUnit.value = 'seconds'
|
||||
takeoverDelayValue.value = normalized
|
||||
}
|
||||
}
|
||||
|
||||
const normalizeTakeoverDelay = () => {
|
||||
const min = takeoverDelayUnit.value === 'minutes' ? 1 : 3
|
||||
const max = takeoverDelayUnit.value === 'minutes' ? 1440 : 86400
|
||||
takeoverDelayValue.value = Math.min(max, Math.max(min, Math.trunc(Number(takeoverDelayValue.value) || min)))
|
||||
}
|
||||
|
||||
const setPermissions = (permissions: AvatarPermission[]) => {
|
||||
const enabled = new Set(permissions)
|
||||
for (const item of permissionItems) permissionState[item.key] = enabled.has(item.key)
|
||||
@@ -260,6 +326,7 @@ const loadSettings = async () => {
|
||||
try {
|
||||
const settings = await getAvatarPermissionSettings(activeAvatarId.value)
|
||||
setPermissions(settings.permissions || [])
|
||||
applyTakeoverDelay(settings.takeoverReplyDelaySeconds || 180)
|
||||
await loadTakeoverStatus()
|
||||
scheduleTakeoverStatusRefresh()
|
||||
} catch (error: any) {
|
||||
@@ -286,8 +353,18 @@ const saveSettings = async (takeoverToggle = false): Promise<boolean> => {
|
||||
saving.value = true
|
||||
errorMessage.value = ''
|
||||
try {
|
||||
const settings = await updateAvatarPermissionSettings(activeAvatarId.value, selectedPermissions())
|
||||
normalizeTakeoverDelay()
|
||||
if (takeoverDelaySeconds.value < 3 || takeoverDelaySeconds.value > 86400) {
|
||||
errorMessage.value = '自动回复等待时间需在 3 秒到 24 小时之间'
|
||||
return false
|
||||
}
|
||||
const settings = await updateAvatarPermissionSettings(
|
||||
activeAvatarId.value,
|
||||
selectedPermissions(),
|
||||
takeoverDelaySeconds.value,
|
||||
)
|
||||
setPermissions(settings.permissions || [])
|
||||
applyTakeoverDelay(settings.takeoverReplyDelaySeconds || 180)
|
||||
await loadTakeoverStatus()
|
||||
scheduleTakeoverStatusRefresh()
|
||||
if (takeoverToggle) {
|
||||
@@ -394,6 +471,10 @@ svg {
|
||||
padding: 0 20px;
|
||||
}
|
||||
|
||||
.authorization-page.embedded .page-content {
|
||||
padding-top: 16px;
|
||||
}
|
||||
|
||||
.permission-intro {
|
||||
min-height: 96px;
|
||||
padding: 15px 16px 14px;
|
||||
@@ -459,6 +540,75 @@ svg {
|
||||
min-height: 76px;
|
||||
}
|
||||
|
||||
.takeover-delay-card {
|
||||
margin-top: 12px;
|
||||
padding: 16px;
|
||||
border: 1px solid #dff1ef;
|
||||
border-radius: 15px;
|
||||
background: linear-gradient(135deg, #f5fcfb 0%, #fff 100%);
|
||||
box-shadow: 0 8px 24px rgba(53, 166, 162, .06);
|
||||
}
|
||||
|
||||
.delay-heading {
|
||||
display: flex;
|
||||
align-items: flex-start;
|
||||
justify-content: space-between;
|
||||
gap: 12px;
|
||||
}
|
||||
|
||||
.delay-heading strong,
|
||||
.delay-heading small {
|
||||
display: block;
|
||||
}
|
||||
|
||||
.delay-heading strong {
|
||||
font-size: 14px;
|
||||
line-height: 1.4;
|
||||
}
|
||||
|
||||
.delay-heading small {
|
||||
margin-top: 5px;
|
||||
color: #8c929f;
|
||||
font-size: 11px;
|
||||
line-height: 1.55;
|
||||
}
|
||||
|
||||
.delay-heading > span {
|
||||
flex: none;
|
||||
padding: 4px 8px;
|
||||
border-radius: 999px;
|
||||
color: #258e8a;
|
||||
background: #e8f8f6;
|
||||
font-size: 11px;
|
||||
font-weight: 700;
|
||||
}
|
||||
|
||||
.delay-control {
|
||||
margin-top: 14px;
|
||||
display: grid;
|
||||
grid-template-columns: minmax(0, 1fr) 88px;
|
||||
gap: 10px;
|
||||
}
|
||||
|
||||
.delay-control input,
|
||||
.delay-control select {
|
||||
min-width: 0;
|
||||
height: 42px;
|
||||
padding: 0 12px;
|
||||
border: 1px solid #dfe5e8;
|
||||
border-radius: 11px;
|
||||
outline: none;
|
||||
color: #222528;
|
||||
background: #fff;
|
||||
font: inherit;
|
||||
}
|
||||
|
||||
.delay-control input:focus,
|
||||
.delay-control select:focus {
|
||||
border-color: #35a6a2;
|
||||
box-shadow: 0 0 0 3px rgba(53, 166, 162, .1);
|
||||
}
|
||||
|
||||
.permission-icon {
|
||||
width: 34px;
|
||||
height: 34px;
|
||||
@@ -599,12 +749,15 @@ svg {
|
||||
bottom: 0;
|
||||
width: min(100%, 390px);
|
||||
padding: 12px 20px calc(20px + env(safe-area-inset-bottom));
|
||||
display: flex;
|
||||
gap: 10px;
|
||||
background: linear-gradient(to bottom, rgba(250, 250, 250, 0), #fafafa 20%, #fafafa 100%);
|
||||
transform: translateX(-50%);
|
||||
}
|
||||
|
||||
.save-button {
|
||||
width: 100%;
|
||||
min-width: 0;
|
||||
flex: 1;
|
||||
height: 48px;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
@@ -620,6 +773,20 @@ svg {
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.footer-back-button {
|
||||
flex: 0 0 96px;
|
||||
height: 48px;
|
||||
border: 1px solid #eadfd6;
|
||||
border-radius: 24px;
|
||||
color: #6f665f;
|
||||
background: #fff;
|
||||
font-size: 14px;
|
||||
font-weight: 500;
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.footer-back-button:disabled { opacity: .58; }
|
||||
|
||||
.save-button:disabled {
|
||||
opacity: .68;
|
||||
}
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
<template>
|
||||
<div class="chat-page">
|
||||
<div class="chat-page" :class="{ 'has-pending-images': pendingImages.length }">
|
||||
<header class="chat-header">
|
||||
<button v-if="!isPublic" class="back-btn" @click="router.back()">‹</button>
|
||||
<div class="avatar-heading">
|
||||
@@ -31,6 +31,12 @@
|
||||
<span v-else>{{ avatar?.emoji || '🤖' }}</span>
|
||||
</div>
|
||||
<div class="message-column">
|
||||
<div v-if="message.attachments?.length" class="message-images">
|
||||
<figure v-for="attachment in message.attachments" :key="attachment.id" class="message-image-card">
|
||||
<img :src="attachment.previewUrl" :alt="attachment.filename" />
|
||||
<figcaption v-if="attachment.warning">{{ attachment.warning }}</figcaption>
|
||||
</figure>
|
||||
</div>
|
||||
<div class="message-bubble" :class="{ streaming: sending && message.role === 'assistant' && index === messages.length - 1 }">
|
||||
<template v-if="message.role === 'assistant'">
|
||||
<span
|
||||
@@ -70,24 +76,66 @@
|
||||
</main>
|
||||
|
||||
<form class="composer" @submit.prevent="sendMessage(inputText)">
|
||||
<textarea v-model="inputText" rows="1" :disabled="sending" placeholder="输入你想聊的内容…" @keydown.enter.exact.prevent="sendMessage(inputText)"></textarea>
|
||||
<button class="send-btn" type="submit" :disabled="sending || !inputText.trim()">发送</button>
|
||||
<div v-if="pendingImages.length" class="pending-images">
|
||||
<div v-for="image in pendingImages" :key="image.localId" class="pending-image" :class="image.status">
|
||||
<img :src="image.previewUrl" :alt="image.filename" />
|
||||
<div class="pending-image-copy">
|
||||
<strong>{{ image.status === 'uploading' ? '正在识别图片…' : image.summary || image.filename }}</strong>
|
||||
<span>{{ image.status === 'uploading' ? '正在提取图片中的可见内容' : categoryLabel(image.category) }}</span>
|
||||
</div>
|
||||
<button type="button" aria-label="移除图片" :disabled="sending" @click="removePendingImage(image.localId)">×</button>
|
||||
</div>
|
||||
<p v-if="hasPendingMedicalImage" class="medical-note">病例与医学影像识别仅供辅助,请以原始资料和医生意见为准。</p>
|
||||
</div>
|
||||
<div class="composer-row">
|
||||
<button class="image-btn" type="button" :disabled="sending || uploadingImage || pendingImages.length >= 3" aria-label="选择图片" @click="imageInput?.click()">
|
||||
<svg viewBox="0 0 24 24" aria-hidden="true"><path d="M4 5.5A2.5 2.5 0 0 1 6.5 3h11A2.5 2.5 0 0 1 20 5.5v13a2.5 2.5 0 0 1-2.5 2.5h-11A2.5 2.5 0 0 1 4 18.5v-13Zm2 12.7 3.8-4.2 2.7 2.8 1.7-1.8 3.8 3.2V5.5a.5.5 0 0 0-.5-.5h-11a.5.5 0 0 0-.5.5v12.7Zm8.3-7.8a1.7 1.7 0 1 0 0-3.4 1.7 1.7 0 0 0 0 3.4Z"/></svg>
|
||||
</button>
|
||||
<input ref="imageInput" class="image-input" type="file" accept="image/jpeg,image/png,image/webp" multiple @change="selectImages" />
|
||||
<textarea v-model="inputText" rows="1" :disabled="sending" placeholder="输入问题,或选择一张图片…" @keydown.enter.exact.prevent="sendMessage(inputText)"></textarea>
|
||||
<button class="send-btn" type="submit" :disabled="sending || uploadingImage || (!inputText.trim() && !readyPendingImages.length)">发送</button>
|
||||
</div>
|
||||
</form>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { computed, nextTick, onMounted, reactive, ref } from 'vue'
|
||||
import { computed, nextTick, onBeforeUnmount, onMounted, reactive, ref } from 'vue'
|
||||
import { useRoute, useRouter } from 'vue-router'
|
||||
import { getAvatarDetail, getPublicAvatar, streamAvatarChat, streamPublicAvatarChat, type ChatMessage } from '@/api'
|
||||
import {
|
||||
getAvatarDetail,
|
||||
getPublicAvatar,
|
||||
streamAvatarChat,
|
||||
streamPublicAvatarChat,
|
||||
uploadAvatarChatImage,
|
||||
uploadPublicAvatarChatImage,
|
||||
type ChatAttachment,
|
||||
type ChatMessage
|
||||
} from '@/api'
|
||||
import { useAvatarStore } from '@/store/avatar'
|
||||
import { useUserStore } from '@/store/user'
|
||||
import { renderChatMarkdownCharacters } from '@/utils/chat-markdown.js'
|
||||
|
||||
type DisplayMessage = ChatMessage & {
|
||||
source?: 'qa' | 'knowledge' | 'qwen' | 'public'
|
||||
source?: 'qa' | 'knowledge' | 'vision' | 'qwen' | 'public'
|
||||
references?: Array<{ filename?: string }>
|
||||
characters?: string[]
|
||||
attachments?: MessageAttachment[]
|
||||
}
|
||||
|
||||
type MessageAttachment = {
|
||||
id: string
|
||||
filename: string
|
||||
previewUrl: string
|
||||
category?: ChatAttachment['category']
|
||||
summary?: string
|
||||
warning?: string
|
||||
}
|
||||
|
||||
type PendingImage = MessageAttachment & {
|
||||
localId: string
|
||||
attachmentId?: string
|
||||
status: 'uploading' | 'ready'
|
||||
}
|
||||
|
||||
const route = useRoute()
|
||||
@@ -103,8 +151,11 @@ const inputText = ref('')
|
||||
const sending = ref(false)
|
||||
const thinking = ref(false)
|
||||
const errorMessage = ref('')
|
||||
const lastQuestion = ref('')
|
||||
const lastRequest = ref<{ question: string; attachments: MessageAttachment[] } | null>(null)
|
||||
const messageList = ref<HTMLElement | null>(null)
|
||||
const imageInput = ref<HTMLInputElement | null>(null)
|
||||
const pendingImages = ref<PendingImage[]>([])
|
||||
const previewUrls = new Set<string>()
|
||||
let scrollFrame: number | null = null
|
||||
|
||||
const userAvatarUrl = computed(() => userStore.user?.avatarUrl || store.userProfile?.avatarUrl || '')
|
||||
@@ -115,14 +166,25 @@ const avatarStatus = computed(() => {
|
||||
if (status === 'training') return { tone: 'training', label: '知识训练中' }
|
||||
return { tone: 'active', label: '在线,随时可以和我聊聊' }
|
||||
})
|
||||
const readyPendingImages = computed(() => pendingImages.value.filter((image) => image.status === 'ready' && image.attachmentId))
|
||||
const uploadingImage = computed(() => pendingImages.value.some((image) => image.status === 'uploading'))
|
||||
const hasPendingMedicalImage = computed(() => readyPendingImages.value.some((image) => ['medical_document', 'medical_image'].includes(image.category || '')))
|
||||
|
||||
const sourceLabels: Record<NonNullable<DisplayMessage['source']>, string> = {
|
||||
qa: '标准问答对',
|
||||
knowledge: '参考文件知识库',
|
||||
vision: '图片理解',
|
||||
qwen: '智能回答',
|
||||
public: ''
|
||||
}
|
||||
const sourceLabel = (source?: DisplayMessage['source']) => source ? sourceLabels[source] : ''
|
||||
const categoryLabels: Record<ChatAttachment['category'], string> = {
|
||||
general_image: '图片内容已识别',
|
||||
document: '文档图片已识别',
|
||||
medical_document: '病例文字已提取,请核对原文',
|
||||
medical_image: '医学影像已作客观描述'
|
||||
}
|
||||
const categoryLabel = (category?: ChatAttachment['category']) => category ? categoryLabels[category] : '图片内容已识别'
|
||||
|
||||
const scrollToBottom = async () => {
|
||||
await nextTick()
|
||||
@@ -214,20 +276,90 @@ const loadAvatar = async () => {
|
||||
document.title = avatar.value?.displayName || avatar.value?.name || '会会数字分身'
|
||||
}
|
||||
|
||||
const sendMessage = async (value: string) => {
|
||||
const question = value.trim()
|
||||
if (!question || sending.value) return
|
||||
lastQuestion.value = question
|
||||
const removePendingImage = (localId: string) => {
|
||||
const target = pendingImages.value.find((image) => image.localId === localId)
|
||||
if (target) {
|
||||
URL.revokeObjectURL(target.previewUrl)
|
||||
previewUrls.delete(target.previewUrl)
|
||||
}
|
||||
pendingImages.value = pendingImages.value.filter((image) => image.localId !== localId)
|
||||
}
|
||||
|
||||
const selectImages = async (event: Event) => {
|
||||
const input = event.target as HTMLInputElement
|
||||
const slots = Math.max(0, 3 - pendingImages.value.length)
|
||||
const files = Array.from(input.files || []).slice(0, slots)
|
||||
input.value = ''
|
||||
for (const file of files) {
|
||||
if (!['image/jpeg', 'image/png', 'image/webp'].includes(file.type)) {
|
||||
errorMessage.value = '仅支持 JPG、PNG、WebP 图片'
|
||||
continue
|
||||
}
|
||||
if (file.size > 8 * 1024 * 1024) {
|
||||
errorMessage.value = '单张图片不能超过 8MB'
|
||||
continue
|
||||
}
|
||||
const previewUrl = URL.createObjectURL(file)
|
||||
previewUrls.add(previewUrl)
|
||||
const localId = `local-${Date.now()}-${Math.random().toString(16).slice(2)}`
|
||||
pendingImages.value.push({
|
||||
id: localId,
|
||||
localId,
|
||||
filename: file.name,
|
||||
previewUrl,
|
||||
status: 'uploading'
|
||||
})
|
||||
errorMessage.value = ''
|
||||
try {
|
||||
const result = isPublic
|
||||
? await uploadPublicAvatarChatImage(shareToken, file)
|
||||
: await uploadAvatarChatImage(avatarId.value, file)
|
||||
const pending = pendingImages.value.find((image) => image.localId === localId)
|
||||
if (!pending) continue
|
||||
Object.assign(pending, {
|
||||
id: result.id,
|
||||
attachmentId: result.id,
|
||||
status: 'ready',
|
||||
category: result.category,
|
||||
summary: result.summary,
|
||||
warning: result.warning
|
||||
})
|
||||
} catch (error: any) {
|
||||
removePendingImage(localId)
|
||||
errorMessage.value = error?.response?.data?.detail || error?.message || '图片识别失败,请重新选择图片'
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const sendMessage = async (value: string, retryAttachments?: MessageAttachment[]) => {
|
||||
const selectedAttachments = retryAttachments || readyPendingImages.value.map((image) => ({
|
||||
id: image.attachmentId || image.id,
|
||||
filename: image.filename,
|
||||
previewUrl: image.previewUrl,
|
||||
category: image.category,
|
||||
summary: image.summary,
|
||||
warning: image.warning
|
||||
}))
|
||||
const question = value.trim() || (selectedAttachments.length ? '请帮我看看这张图片。' : '')
|
||||
if (!question || sending.value || (!retryAttachments && uploadingImage.value)) return
|
||||
const history = messages.value.slice(-10).map(({ role, content, attachments }) => ({
|
||||
role,
|
||||
content,
|
||||
attachmentIds: attachments?.map((attachment) => attachment.id) || []
|
||||
}))
|
||||
lastRequest.value = { question, attachments: selectedAttachments }
|
||||
inputText.value = ''
|
||||
errorMessage.value = ''
|
||||
messages.value.push({ role: 'user', content: question })
|
||||
if (!retryAttachments) pendingImages.value = []
|
||||
messages.value.push({ role: 'user', content: question, attachments: selectedAttachments })
|
||||
sending.value = true
|
||||
thinking.value = true
|
||||
await scrollToBottom()
|
||||
try {
|
||||
const payload = {
|
||||
message: question,
|
||||
history: messages.value.slice(-10).map(({ role, content }) => ({ role, content }))
|
||||
attachmentIds: selectedAttachments.map((attachment) => attachment.id),
|
||||
history
|
||||
}
|
||||
const streamed = createStreamReply()
|
||||
const handlers = {
|
||||
@@ -256,13 +388,17 @@ const sendMessage = async (value: string) => {
|
||||
}
|
||||
|
||||
const retryLast = () => {
|
||||
if (!lastQuestion.value || sending.value) return
|
||||
const last = messages.value[messages.value.length - 1]
|
||||
if (last?.role === 'user') messages.value.pop()
|
||||
sendMessage(lastQuestion.value)
|
||||
if (!lastRequest.value || sending.value) return
|
||||
while (messages.value[messages.value.length - 1]?.role === 'assistant') messages.value.pop()
|
||||
if (messages.value[messages.value.length - 1]?.role === 'user') messages.value.pop()
|
||||
void sendMessage(lastRequest.value.question, lastRequest.value.attachments)
|
||||
}
|
||||
|
||||
onMounted(loadAvatar)
|
||||
onBeforeUnmount(() => {
|
||||
previewUrls.forEach((url) => URL.revokeObjectURL(url))
|
||||
previewUrls.clear()
|
||||
})
|
||||
</script>
|
||||
|
||||
<style scoped>
|
||||
@@ -275,6 +411,7 @@ onMounted(loadAvatar)
|
||||
.avatar-heading h1 { margin: 0; font-size: 17px; }
|
||||
.online-state { display: flex; align-items: center; gap: 4px; margin-top: 3px; font-size: 11px; opacity: .9; }.online-state i { width: 7px; height: 7px; border-radius: 50%; background: #86EFAC; box-shadow: 0 0 0 2px rgba(255,255,255,.22); }.online-state.training i { background: #FDE68A; }.online-state.inactive i { background: #FDA4AF; }
|
||||
.message-list { min-height: 0; flex: 1 1 auto; width: min(760px, 100%); box-sizing: border-box; margin: 0 auto; padding: 24px 18px 120px; overflow-y: auto; overscroll-behavior: contain; }
|
||||
.chat-page.has-pending-images .message-list { padding-bottom: min(330px, 42vh); }
|
||||
.welcome-card { padding: 28px 20px; text-align: center; background: rgba(255,255,255,.72); border: 1px solid #FFE1C2; border-radius: 22px; box-shadow: 0 10px 28px rgba(181, 99, 35, .08); }
|
||||
.welcome-avatar { width: 64px; height: 64px; display: grid; place-items: center; margin: 0 auto 14px; overflow: hidden; border: 3px solid #fff; border-radius: 50%; background: #FFE4C7; box-shadow: 0 7px 16px rgba(181, 99, 35, .18); font-size: 32px; }.welcome-avatar img { width: 100%; height: 100%; object-fit: cover; }
|
||||
.welcome-card h2 { margin: 0 0 8px; font-size: 20px; }.welcome-description { max-width: 340px; margin: 0 auto; color: #8B6B58; font-size: 14px; line-height: 1.65; }
|
||||
@@ -282,6 +419,10 @@ onMounted(loadAvatar)
|
||||
.message-row.user { justify-content: flex-end; }
|
||||
.message-avatar { flex: 0 0 auto; width: 42px; height: 42px; display: grid; place-items: center; overflow: hidden; border: 2px solid rgba(255,255,255,.9); border-radius: 14px; background: #FFE4C7; box-shadow: 0 3px 10px rgba(96, 52, 21, .12); font-size: 16px; }.message-avatar img { width: 100%; height: 100%; object-fit: cover; }.user-message-face { color: #fff; background: #D97706; }
|
||||
.message-column { max-width: min(78%, 560px); }
|
||||
.message-images { display: grid; grid-template-columns: repeat(2, minmax(0, 150px)); gap: 8px; margin-bottom: 8px; }
|
||||
.message-image-card { margin: 0; overflow: hidden; border: 1px solid #F4D4B8; border-radius: 14px; background: #fff; box-shadow: 0 4px 14px rgba(96, 52, 21, .08); }
|
||||
.message-image-card img { display: block; width: 100%; max-height: 210px; object-fit: cover; }
|
||||
.message-image-card figcaption { padding: 7px 9px; color: #8A5A3B; background: #FFF6ED; font-size: 10px; line-height: 1.45; }
|
||||
.message-bubble { padding: 12px 14px; white-space: pre-wrap; line-height: 1.6; font-size: 15px; border-radius: 4px 16px 16px 16px; background: white; box-shadow: 0 3px 12px rgba(96, 52, 21, .07); }
|
||||
.message-bubble.streaming::after { content: ''; display: inline-block; width: 2px; height: 1.05em; margin-left: 3px; vertical-align: -0.16em; background: currentColor; animation: type-cursor .75s step-end infinite; }
|
||||
.typing-character { display: inline-block; animation: character-in .24s cubic-bezier(.2,.72,.25,1) both; }.typing-character.newline { display: block; height: 0; }
|
||||
@@ -298,7 +439,27 @@ onMounted(loadAvatar)
|
||||
@keyframes type-cursor { 50% { opacity: 0; } }
|
||||
@keyframes character-in { from { opacity: 0; transform: translateY(3px); } to { opacity: 1; transform: translateY(0); } }
|
||||
.chat-error { margin: 4px auto; color: #B42318; font-size: 13px; }.chat-error button { border: 0; background: none; color: #C15F18; cursor: pointer; text-decoration: underline; }
|
||||
.composer { position: fixed; left: 0; right: 0; bottom: 0; display: flex; gap: 10px; padding: 12px max(18px, calc((100vw - 760px) / 2 + 18px)); background: rgba(255,255,255,.92); border-top: 1px solid #F4DCC7; backdrop-filter: blur(12px); }
|
||||
.composer { position: fixed; left: 0; right: 0; bottom: 0; display: flex; flex-direction: column; gap: 9px; padding: 10px max(18px, calc((100vw - 760px) / 2 + 18px)) 12px; background: rgba(255,255,255,.94); border-top: 1px solid #F4DCC7; backdrop-filter: blur(14px); }
|
||||
.composer-row { display: flex; align-items: flex-end; gap: 9px; }
|
||||
.composer textarea { flex: 1; resize: none; min-height: 22px; max-height: 100px; padding: 11px 13px; border: 1px solid #EED8C5; border-radius: 13px; font: inherit; color: #3B2417; outline: none; }.composer textarea:focus { border-color: #F97316; }
|
||||
.image-input { display: none; }
|
||||
.image-btn { flex: 0 0 auto; width: 44px; height: 44px; display: grid; place-items: center; border: 1px solid #EED8C5; border-radius: 13px; color: #C65A11; background: #FFF8F1; cursor: pointer; }
|
||||
.image-btn svg { width: 22px; height: 22px; fill: currentColor; }
|
||||
.image-btn:disabled { opacity: .4; cursor: not-allowed; }
|
||||
.pending-images { display: grid; gap: 7px; }
|
||||
.pending-image { display: grid; grid-template-columns: 48px minmax(0, 1fr) 30px; align-items: center; gap: 9px; min-height: 48px; padding: 6px 8px; border: 1px solid #F1D4BB; border-radius: 14px; background: #FFF9F3; }
|
||||
.pending-image img { width: 48px; height: 48px; object-fit: cover; border-radius: 10px; }
|
||||
.pending-image-copy { min-width: 0; display: flex; flex-direction: column; gap: 2px; }
|
||||
.pending-image-copy strong { overflow: hidden; color: #4B2B19; font-size: 12px; text-overflow: ellipsis; white-space: nowrap; }
|
||||
.pending-image-copy span { color: #9A7159; font-size: 10px; }
|
||||
.pending-image.uploading strong::after { content: ''; display: inline-block; width: 7px; height: 7px; margin-left: 7px; border: 2px solid #F6B889; border-top-color: #F97316; border-radius: 50%; animation: image-spin .7s linear infinite; }
|
||||
.pending-image > button { width: 28px; height: 28px; border: 0; border-radius: 9px; color: #9A7159; background: #F8E8D9; font-size: 19px; cursor: pointer; }
|
||||
.medical-note { margin: 0; padding: 0 2px; color: #9A5A2E; font-size: 10px; line-height: 1.45; }
|
||||
.send-btn { align-self: flex-end; padding: 11px 18px; border: 0; border-radius: 12px; color: white; background: #F97316; cursor: pointer; }.send-btn:disabled { opacity: .45; cursor: not-allowed; }
|
||||
@keyframes image-spin { to { transform: rotate(360deg); } }
|
||||
@media (max-width: 520px) {
|
||||
.message-images { grid-template-columns: minmax(0, 220px); }
|
||||
.message-column { max-width: 80%; }
|
||||
.send-btn { padding-inline: 14px; }
|
||||
}
|
||||
</style>
|
||||
|
||||
@@ -1,12 +1,10 @@
|
||||
<template>
|
||||
<div class="edit-avatar-page">
|
||||
<!-- 顶部导航 -->
|
||||
<header class="page-header">
|
||||
<header v-if="!isEmbedded" class="page-header">
|
||||
<button class="back-btn" @click="goBack">‹</button>
|
||||
<h1 class="page-title">分身微调</h1>
|
||||
<button class="save-btn" :disabled="loading || saving || uploadingPhoto" @click="saveChanges">
|
||||
{{ saving ? '保存中...' : '保存' }}
|
||||
</button>
|
||||
<span class="header-spacer" aria-hidden="true"></span>
|
||||
</header>
|
||||
|
||||
<div v-if="loading" class="status-banner">加载中...</div>
|
||||
@@ -159,6 +157,13 @@
|
||||
{{ deleting ? '删除中...' : '删除数字分身' }}
|
||||
</button>
|
||||
</section>
|
||||
|
||||
<footer class="edit-action-bar">
|
||||
<button class="action-back-btn" type="button" :disabled="saving" @click="goBack">返回</button>
|
||||
<button class="action-save-btn" type="button" :disabled="loading || saving || uploadingPhoto" @click="saveChanges">
|
||||
{{ saving ? '保存中...' : '保存修改' }}
|
||||
</button>
|
||||
</footer>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
@@ -168,11 +173,13 @@ import { useRoute, useRouter } from 'vue-router'
|
||||
import { deleteAvatar as apiDeleteAvatar, getAvatarDetail, updateAvatar, uploadAvatarPhoto } from '@/api'
|
||||
import { useAvatarStore } from '@/store/avatar'
|
||||
import { buildAvatarUpdatePayload, normalizeAvatarEditForm } from '@/utils/avatar-page-data.js'
|
||||
import { isHuihuiEmbeddedMode } from '@/utils/embed-mode'
|
||||
|
||||
const router = useRouter()
|
||||
const route = useRoute()
|
||||
const avatarStore = useAvatarStore()
|
||||
const avatarId = route.params.id as string
|
||||
const isEmbedded = isHuihuiEmbeddedMode()
|
||||
|
||||
// 表单数据
|
||||
const formData = reactive({
|
||||
@@ -288,7 +295,7 @@ onMounted(async () => {
|
||||
.edit-avatar-page {
|
||||
min-height: 100vh;
|
||||
background: #F8F9FA;
|
||||
padding-bottom: 40px;
|
||||
padding-bottom: calc(104px + env(safe-area-inset-bottom));
|
||||
}
|
||||
|
||||
/* 顶部导航 */
|
||||
@@ -331,16 +338,7 @@ onMounted(async () => {
|
||||
color: #B91C1C;
|
||||
}
|
||||
|
||||
.save-btn {
|
||||
background: #F97316;
|
||||
color: white;
|
||||
border: none;
|
||||
padding: 8px 20px;
|
||||
border-radius: 8px;
|
||||
font-size: 14px;
|
||||
font-weight: 600;
|
||||
cursor: pointer;
|
||||
}
|
||||
.header-spacer { width: 40px; }
|
||||
|
||||
/* 头像上传 */
|
||||
.photo-section {
|
||||
@@ -580,4 +578,47 @@ onMounted(async () => {
|
||||
background: #EF4444;
|
||||
color: white;
|
||||
}
|
||||
|
||||
.edit-action-bar {
|
||||
position: fixed;
|
||||
z-index: 30;
|
||||
left: 0;
|
||||
right: 0;
|
||||
bottom: 0;
|
||||
display: flex;
|
||||
gap: 12px;
|
||||
padding: 12px 20px calc(14px + env(safe-area-inset-bottom));
|
||||
border-top: 1px solid rgba(229, 231, 235, .9);
|
||||
background: rgba(248, 249, 250, .96);
|
||||
box-shadow: 0 -8px 24px rgba(56, 38, 24, .06);
|
||||
backdrop-filter: blur(12px);
|
||||
}
|
||||
|
||||
.action-back-btn,
|
||||
.action-save-btn {
|
||||
height: 48px;
|
||||
border-radius: 14px;
|
||||
font-size: 15px;
|
||||
font-weight: 600;
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.action-back-btn {
|
||||
flex: 0 0 104px;
|
||||
border: 1px solid #E4E0DC;
|
||||
color: #655E58;
|
||||
background: #fff;
|
||||
}
|
||||
|
||||
.action-save-btn {
|
||||
min-width: 0;
|
||||
flex: 1;
|
||||
border: 0;
|
||||
color: #fff;
|
||||
background: linear-gradient(105deg, #F79A38, #F97316);
|
||||
box-shadow: 0 8px 18px rgba(249, 115, 22, .18);
|
||||
}
|
||||
|
||||
.action-back-btn:disabled,
|
||||
.action-save-btn:disabled { opacity: .6; cursor: not-allowed; }
|
||||
</style>
|
||||
|
||||
@@ -1,15 +1,11 @@
|
||||
<template>
|
||||
<div class="avatar-manage-page">
|
||||
<!-- 顶部导航 -->
|
||||
<header class="page-header">
|
||||
<header v-if="!isEmbedded" class="page-header">
|
||||
<div class="header-left">
|
||||
<button class="back-btn" @click="goBack">‹</button>
|
||||
<h1 class="page-title">数字分身管理</h1>
|
||||
</div>
|
||||
<!-- 右上角创建入口 -->
|
||||
<div class="header-right">
|
||||
<button class="icon-btn" @click="goCreate" title="创建数字分身">➕</button>
|
||||
</div>
|
||||
</header>
|
||||
|
||||
<!-- 用户资料头(会会登录账号的头像 / 昵称) -->
|
||||
@@ -39,8 +35,13 @@
|
||||
<!-- 数字分身列表(只放分身相关) -->
|
||||
<section class="avatar-list-section">
|
||||
<div class="section-head">
|
||||
<h3 class="section-title">我的数字分身</h3>
|
||||
<span class="count-badge">{{ avatars.length }}</span>
|
||||
<div class="section-heading-copy">
|
||||
<h3 class="section-title">我的数字分身</h3>
|
||||
<span class="count-badge">{{ avatars.length }}</span>
|
||||
</div>
|
||||
<button class="section-create-btn" type="button" @click="goCreate">
|
||||
<span aria-hidden="true">+</span> 添加分身
|
||||
</button>
|
||||
</div>
|
||||
|
||||
<div v-if="avatars.length" class="avatar-list">
|
||||
@@ -87,10 +88,12 @@ import { useRouter } from 'vue-router'
|
||||
import { useAvatarStore } from '@/store/avatar'
|
||||
import { useUserStore } from '@/store/user'
|
||||
import { createAvatarShareLink } from '@/api'
|
||||
import { isHuihuiEmbeddedMode } from '@/utils/embed-mode'
|
||||
|
||||
const router = useRouter()
|
||||
const avatarStore = useAvatarStore()
|
||||
const userStore = useUserStore()
|
||||
const isEmbedded = isHuihuiEmbeddedMode()
|
||||
|
||||
// 临时产品开关:余额卡片代码保留,后续改为 true 即可恢复展示。
|
||||
const SHOW_POINTS_BALANCE_CARD = false
|
||||
@@ -354,10 +357,36 @@ onMounted(() => {
|
||||
.section-head {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
justify-content: space-between;
|
||||
gap: 12px;
|
||||
margin: 8px 0 12px;
|
||||
}
|
||||
|
||||
.section-heading-copy {
|
||||
min-width: 0;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
}
|
||||
|
||||
.section-create-btn {
|
||||
flex: 0 0 auto;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 3px;
|
||||
padding: 8px 12px;
|
||||
border: 1px solid #FED7B5;
|
||||
border-radius: 999px;
|
||||
color: #E9650C;
|
||||
background: #FFF7ED;
|
||||
font-size: 12px;
|
||||
font-weight: 650;
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.section-create-btn span { font-size: 17px; line-height: 1; }
|
||||
.section-create-btn:active { background: #FFEDD5; }
|
||||
|
||||
.section-title {
|
||||
font-size: 16px;
|
||||
font-weight: 600;
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
<template>
|
||||
<div class="knowledge-page">
|
||||
<div class="knowledge-page" :class="{ embedded: isEmbedded }">
|
||||
<!-- 顶部导航 -->
|
||||
<header class="page-header">
|
||||
<header v-if="!isEmbedded" class="page-header">
|
||||
<div class="header-left">
|
||||
<button class="back-btn" @click="goBack">‹</button>
|
||||
<h1 class="page-title">知识库管理</h1>
|
||||
@@ -23,7 +23,7 @@
|
||||
<div class="upload-section">
|
||||
<div class="upload-zone" :class="{ 'drag-over': dragOver }" @click="triggerFile" @dragover.prevent="dragOver = true" @dragleave.prevent="dragOver = false" @drop.prevent="onDrop">
|
||||
<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>
|
||||
<input ref="fileInput" type="file" accept=".md,.txt,.pdf,.doc,.docx,.xlsx" class="hidden-input" @change="onFileChange" />
|
||||
</div>
|
||||
@@ -37,10 +37,10 @@
|
||||
<div class="card-content">
|
||||
<div class="card-title-row">
|
||||
<strong>{{ doc.filename }}</strong>
|
||||
<span class="status-pill" :class="{ pending: !doc.vectorized && doc.filePresent !== false, missing: doc.filePresent === false }">{{ doc.filePresent === false ? '文件缺失' : (doc.vectorized ? '已入库' : '处理中') }}</span>
|
||||
<span class="status-pill" :class="documentState(doc).tone">{{ documentState(doc).label }}</span>
|
||||
</div>
|
||||
<p class="card-meta">{{ doc.fileType.toUpperCase() }} · {{ formatSize(doc.fileSize) }} · {{ formatDate(doc.createdAt) }}</p>
|
||||
<p class="card-detail">{{ doc.filePresent === false ? '原文件不可用,请删除后重新上传' : (doc.vectorized ? `已切分 ${doc.chunkCount || 0} 段,可用于对话` : '正在解析并建立知识索引') }}</p>
|
||||
<p class="card-detail">{{ documentState(doc).detail }}</p>
|
||||
</div>
|
||||
<button class="card-delete" @click="removeDoc(doc.id)">删除</button>
|
||||
</article>
|
||||
@@ -82,6 +82,7 @@ import { ref, onMounted, computed } from 'vue'
|
||||
import { useRoute, useRouter } from 'vue-router'
|
||||
import { useAvatarStore } from '@/store/avatar'
|
||||
import { pickScopedAvatarId, unwrapListData } from '@/utils/avatar-page-data.js'
|
||||
import { isHuihuiEmbeddedMode } from '@/utils/embed-mode'
|
||||
import {
|
||||
getKnowledgeDocs,
|
||||
uploadKnowledgeDoc,
|
||||
@@ -95,6 +96,7 @@ import {
|
||||
const router = useRouter()
|
||||
const route = useRoute()
|
||||
const store = useAvatarStore()
|
||||
const isEmbedded = isHuihuiEmbeddedMode()
|
||||
|
||||
const avatarId = computed(() => pickScopedAvatarId(route.params.avatarId, store.currentAvatarId, store.avatars))
|
||||
const activeTab = ref<'docs' | 'qa'>('docs')
|
||||
@@ -111,6 +113,19 @@ const searching = ref(false)
|
||||
const searched = ref(false)
|
||||
const searchResults = ref<any[]>([])
|
||||
|
||||
const documentState = (doc: any) => {
|
||||
if (doc.filePresent === false) {
|
||||
return { tone: 'missing', label: '文件缺失', detail: '原文件不可用,请删除后重新上传' }
|
||||
}
|
||||
if (doc.vectorized) {
|
||||
return { tone: 'ready', label: '已入库', detail: `已切分 ${doc.chunkCount || 0} 段,可用于对话` }
|
||||
}
|
||||
if (['uploaded', 'parsing'].includes(String(doc.status || '').toLowerCase())) {
|
||||
return { tone: 'pending', label: '处理中', detail: '正在解析并建立知识索引' }
|
||||
}
|
||||
return { tone: 'failed', label: '处理失败', detail: '未能建立知识索引,请删除后重新上传' }
|
||||
}
|
||||
|
||||
const loadDocs = async () => {
|
||||
if (!avatarId.value) return
|
||||
try {
|
||||
@@ -283,6 +298,7 @@ onMounted(async () => {
|
||||
.card-title-row strong { min-width: 0; flex: 1; overflow: hidden; color: #27201C; font-size: 14px; text-overflow: ellipsis; white-space: nowrap; }
|
||||
.status-pill { flex: 0 0 auto; display: inline-flex; padding: 4px 7px; border-radius: 999px; color: #15803D; background: #ECFDF3; font-size: 10px; white-space: nowrap; }.status-pill.pending { color: #B45309; background: #FFFBEB; }
|
||||
.status-pill.missing { color: #B91C1C; background: #FEF2F2; }
|
||||
.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-delete { flex: 0 0 auto; align-self: center; border: 0; color: #EF4444; background: #FEF2F2; border-radius: 8px; padding: 7px 9px; font-size: 12px; cursor: pointer; }
|
||||
.card-empty { padding: 42px 16px; border: 1px dashed #F1D9C3; border-radius: 16px; color: #9398AE; background: #fff; font-size: 14px; text-align: center; }
|
||||
|
||||
@@ -10,10 +10,12 @@
|
||||
<section class="form-section">
|
||||
<label class="field-label">问题</label>
|
||||
<textarea
|
||||
ref="questionInput"
|
||||
v-model="form.question"
|
||||
class="field-input"
|
||||
rows="3"
|
||||
class="field-input question-input"
|
||||
rows="1"
|
||||
placeholder="例如:你们的退款政策是什么?"
|
||||
@input="resizeQuestion"
|
||||
></textarea>
|
||||
|
||||
<label class="field-label">标准答案</label>
|
||||
@@ -46,7 +48,7 @@
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { ref, reactive, computed, onMounted } from 'vue'
|
||||
import { ref, reactive, computed, nextTick, onMounted } from 'vue'
|
||||
import { useRouter, useRoute } from 'vue-router'
|
||||
import { useAvatarStore } from '@/store/avatar'
|
||||
import { pickScopedAvatarId, unwrapListData } from '@/utils/avatar-page-data.js'
|
||||
@@ -63,6 +65,14 @@ const isEdit = computed(() => !!qaId.value)
|
||||
const form = reactive({ question: '', answer: '', enabled: true })
|
||||
const saving = ref(false)
|
||||
const error = ref('')
|
||||
const questionInput = ref<HTMLTextAreaElement | null>(null)
|
||||
|
||||
const resizeQuestion = (event?: Event) => {
|
||||
const element = (event?.target as HTMLTextAreaElement | null) || questionInput.value
|
||||
if (!element) return
|
||||
element.style.height = 'auto'
|
||||
element.style.height = `${element.scrollHeight}px`
|
||||
}
|
||||
|
||||
const goBack = () => router.back()
|
||||
|
||||
@@ -124,6 +134,8 @@ onMounted(async () => {
|
||||
if (isEdit.value) {
|
||||
await loadForEdit()
|
||||
}
|
||||
await nextTick()
|
||||
resizeQuestion()
|
||||
})
|
||||
</script>
|
||||
|
||||
@@ -194,6 +206,13 @@ onMounted(async () => {
|
||||
border-color: #F97316;
|
||||
}
|
||||
|
||||
.question-input {
|
||||
min-height: 44px;
|
||||
overflow: hidden;
|
||||
resize: none;
|
||||
line-height: 1.55;
|
||||
}
|
||||
|
||||
.switch-row {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
|
||||
@@ -24,6 +24,8 @@
|
||||
</div>
|
||||
<div class="model-meta">
|
||||
<span>版本: {{ m.model_version || '--' }}</span>
|
||||
<span v-if="m.usage_scope === 'digital_avatar'">视觉: {{ m.vision_model_version || 'qwen3.6-flash' }}</span>
|
||||
<span v-if="m.usage_scope === 'digital_avatar'">病例OCR: {{ m.ocr_model_version || 'qwen-vl-ocr' }}</span>
|
||||
<span>温度: {{ m.temperature }}</span>
|
||||
<span>Max Tokens: {{ m.max_tokens }}</span>
|
||||
<span>超时: {{ m.timeout_seconds }}s</span>
|
||||
@@ -68,6 +70,15 @@
|
||||
<el-form-item label="模型版本">
|
||||
<el-input v-model="form.model_version" placeholder="如: gpt-4-turbo, glm-4" />
|
||||
</el-form-item>
|
||||
<template v-if="form.usage_scope === 'digital_avatar'">
|
||||
<el-form-item label="视觉模型">
|
||||
<el-input v-model="form.vision_model_version" placeholder="如: qwen3.6-flash" />
|
||||
</el-form-item>
|
||||
<el-form-item label="病例OCR模型">
|
||||
<el-input v-model="form.ocr_model_version" placeholder="如: qwen-vl-ocr" />
|
||||
<div class="scope-tip">识别为病例、处方、检查单后自动调用,普通图片不会重复调用。</div>
|
||||
</el-form-item>
|
||||
</template>
|
||||
<el-row :gutter="16">
|
||||
<el-col :span="12">
|
||||
<el-form-item label="温度">
|
||||
@@ -142,7 +153,7 @@ const testing = ref(false)
|
||||
|
||||
const providerLabels = { openai: 'OpenAI', zhipu: '智谱GLM', wenxin: '文心一言', qianwen: '通义千问', local: '本地模型' }
|
||||
const scopeLabels = { general: '通用业务', digital_avatar: '数字分身专用' }
|
||||
const form = reactive({ model_name: '', provider: 'openai', usage_scope: 'general', api_base_url: '', api_key: '', model_version: '', temperature: 0.7, max_tokens: 1000, timeout_seconds: 30, is_default: 0 })
|
||||
const form = reactive({ model_name: '', provider: 'openai', usage_scope: 'general', api_base_url: '', api_key: '', model_version: '', vision_model_version: 'qwen3.6-flash', ocr_model_version: 'qwen-vl-ocr', temperature: 0.7, max_tokens: 1000, timeout_seconds: 30, is_default: 0 })
|
||||
const rules = { model_name: [{ required: true, message: '请输入模型名称' }], provider: [{ required: true }], usage_scope: [{ required: true }] }
|
||||
|
||||
async function load() {
|
||||
@@ -166,13 +177,13 @@ function onProviderChange(provider) {
|
||||
|
||||
function openCreate() {
|
||||
editModel.value = null
|
||||
Object.assign(form, { model_name: '', provider: 'openai', usage_scope: 'general', api_base_url: PROVIDER_DEFAULTS.openai.api_base_url, api_key: '', model_version: PROVIDER_DEFAULTS.openai.model_version, temperature: 0.7, max_tokens: 1000, timeout_seconds: 30, is_default: 0 })
|
||||
Object.assign(form, { model_name: '', provider: 'openai', usage_scope: 'general', api_base_url: PROVIDER_DEFAULTS.openai.api_base_url, api_key: '', model_version: PROVIDER_DEFAULTS.openai.model_version, vision_model_version: 'qwen3.6-flash', ocr_model_version: 'qwen-vl-ocr', temperature: 0.7, max_tokens: 1000, timeout_seconds: 30, is_default: 0 })
|
||||
dialogVisible.value = true
|
||||
}
|
||||
|
||||
function openEdit(m) {
|
||||
editModel.value = m
|
||||
Object.assign(form, { model_name: m.model_name, provider: m.provider, usage_scope: m.usage_scope || 'general', api_base_url: m.api_base_url || '', api_key: '', model_version: m.model_version || '', temperature: m.temperature, max_tokens: m.max_tokens, timeout_seconds: m.timeout_seconds, is_default: m.is_default })
|
||||
Object.assign(form, { model_name: m.model_name, provider: m.provider, usage_scope: m.usage_scope || 'general', api_base_url: m.api_base_url || '', api_key: '', model_version: m.model_version || '', vision_model_version: m.vision_model_version || 'qwen3.6-flash', ocr_model_version: m.ocr_model_version || 'qwen-vl-ocr', temperature: m.temperature, max_tokens: m.max_tokens, timeout_seconds: m.timeout_seconds, is_default: m.is_default })
|
||||
dialogVisible.value = true
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user