Compare commits

...
Author SHA1 Message Date
stefanfeng 08c58fe0e6 fix(avatar): cap knowledge files at 50MB 2026-09-04 15:18:39 +08:00
stefanfeng 6b7201e890 fix(avatar): align knowledge upload limit with production 2026-09-04 13:56:29 +08:00
stefanfeng 28553aba15 fix(avatar): allow knowledge uploads up to 20MB 2026-09-04 13:47:46 +08:00
stefanfeng b98a2b9507 fix(avatar): index knowledge documents asynchronously 2026-09-04 11:53:56 +08:00
stefanfeng 59350fb41d Merge pull request 'fix(avatar): ground replies in recognized images' (#14) from codex/avatar-image-answer-hotfix-20260902 into main
Reviewed-on: #14
2026-09-02 14:49:15 +08:00
stefanfeng 6a4b35c49a fix(avatar): ground replies in recognized images 2026-09-02 14:47:27 +08:00
stefanfeng 207bbd02cf Merge pull request 'fix(avatar): recover BOXIM image replies' (#13) from codex/avatar-boxim-vision-hotfix-20260902 into main
Reviewed-on: #13
2026-09-02 14:13:19 +08:00
stefanfeng 7cac96356d fix(avatar): recover BOXIM image replies 2026-09-02 14:12:04 +08:00
stefanfeng 03c32309a8 Merge pull request 'feat(avatar): understand BOXIM image messages' (#12) from codex/avatar-boxim-vision-20260901 into main
Merge pull request #12: BOXIM image understanding
2026-09-01 15:36:19 +08:00
stefanfeng 0fc43908ae feat(avatar): understand BOXIM image messages 2026-09-01 15:35:41 +08:00
stefanfeng 3d999f9472 Merge pull request 'fix(avatar): prevent BOXIM polling starvation' (#11) from codex/avatar-takeover-poll-20260901 into main 2026-09-01 14:04:18 +08:00
stefanfeng 540edb58c4 fix(avatar): prevent BOXIM polling starvation 2026-09-01 14:03:35 +08:00
stefanfeng a7eb6ac2a5 Merge pull request #10 from codex/avatar-vision-chat-20260831
feat(avatar): 支持图片与病例理解对话
2026-09-01 10:41:51 +08:00
stefanfeng 016bc22c05 feat(avatar): add private vision chat support 2026-08-31 15:10:50 +08:00
stefanfeng 094f8cd40f Merge pull request 'fix(avatar): normalize embedding API endpoint' (#9) from codex/avatar-embedding-endpoint-20260828 into main
Reviewed-on: #9
2026-08-28 17:00:16 +08:00
stefanfeng 6794e88d53 fix(avatar): normalize embedding API endpoint 2026-08-28 16:59:22 +08:00
stefanfeng 7884430b3d Merge pull request 'feat(avatar): 优化会会 H5 嵌入管理流程' (#8) from codex/avatar-h5-embedded-layout-20260827 into main
Reviewed-on: #8
2026-08-27 11:41:21 +08:00
stefanfeng 46d42b7d98 fix(avatar): prevent takeover loops and isolate settings 2026-08-27 11:30:12 +08:00
stefanfeng ef58c5f2d2 feat(avatar): optimize embedded H5 management flow 2026-08-27 09:28:05 +08:00
stefanfeng 6e3fe5a616 Merge pull request 'feat(avatar): 接入会会支付并统一积分展示' (#7) from codex/avatar-token-copy-to-points-20260826 into main
Reviewed-on: #7
2026-08-26 14:53:58 +08:00
stefanfeng 67b6bd1b48 Merge pull request 'fix(avatar): 部署重启后自动恢复 BOXIM 接管' (#6) from codex/avatar-takeover-restart-safe-20260826 into main
Reviewed-on: #6
2026-08-26 14:53:49 +08:00
stefanfeng 0752001d85 Merge pull request 'feat(avatar): 数字分身自动跟随用户语言回答' (#5) from codex/avatar-auto-reply-language-20260826 into main
Reviewed-on: #5
2026-08-26 14:53:40 +08:00
stefanfeng 0c6419f37e fix(avatar): restore BOXIM takeover after restart 2026-08-26 13:25:16 +08:00
stefanfeng f768e7648f fix(avatar): prevent inferred reply scenarios 2026-08-26 11:56:54 +08:00
stefanfeng e30ab2b889 feat(avatar): follow user language in replies 2026-08-26 11:52:32 +08:00
47 changed files with 4431 additions and 252 deletions
+6
View File
@@ -37,6 +37,8 @@ async def create_model(req: AIModelCreateRequest, db=Depends(get_db)):
api_base_url=req.api_base_url, api_base_url=req.api_base_url,
api_key_enc=encrypt(req.api_key) if req.api_key else None, api_key_enc=encrypt(req.api_key) if req.api_key else None,
model_version=req.model_version, model_version=req.model_version,
vision_model_version=req.vision_model_version,
ocr_model_version=req.ocr_model_version,
temperature=req.temperature, temperature=req.temperature,
max_tokens=req.max_tokens, max_tokens=req.max_tokens,
timeout_seconds=req.timeout_seconds, 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_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 "", "api_key": decrypt(model.api_key_enc) if model.api_key_enc else "",
"model": model.model_version or model.model_name, "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, "temperature": model.temperature,
"max_tokens": model.max_tokens, "max_tokens": model.max_tokens,
"timeout_seconds": model.timeout_seconds, "timeout_seconds": model.timeout_seconds,
@@ -129,6 +133,8 @@ def _format_model(m: AIModelConfig) -> dict:
"usage_scope": m.usage_scope, "usage_scope": m.usage_scope,
"api_base_url": m.api_base_url, "has_api_key": bool(m.api_key_enc), "api_base_url": m.api_base_url, "has_api_key": bool(m.api_key_enc),
"model_version": m.model_version, "temperature": m.temperature, "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, "max_tokens": m.max_tokens, "timeout_seconds": m.timeout_seconds,
"is_default": m.is_default, "is_enabled": m.is_enabled, "is_default": m.is_default, "is_enabled": m.is_enabled,
"created_at": m.created_at.isoformat(), "created_at": m.created_at.isoformat(),
+24 -9
View File
@@ -66,21 +66,36 @@ async def init_db():
PendingReplyTask, TokenStat, AIModelConfig, SystemConfig, LoginLog PendingReplyTask, TokenStat, AIModelConfig, SystemConfig, LoginLog
) )
async with engine.begin() as conn: 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: try:
columns = (
(
"usage_scope",
"ALTER TABLE ai_model_configs ADD COLUMN 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( result = await conn.execute(text(
"SELECT COUNT(*) FROM information_schema.COLUMNS " "SELECT COUNT(*) FROM information_schema.COLUMNS "
"WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'ai_model_configs' " "WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'ai_model_configs' "
"AND COLUMN_NAME = 'usage_scope'" "AND COLUMN_NAME = :column_name"
)) ), {"column_name": column_name})
if result.scalar_one() == 0: if result.scalar_one() == 0:
await conn.execute(text( await conn.execute(text(ddl))
"ALTER TABLE ai_model_configs ADD COLUMN usage_scope " logger.info("AI模型配置表已增加 %s 字段", column_name)
"VARCHAR(16) NOT NULL DEFAULT 'general' AFTER provider"
))
logger.info("AI模型配置表已增加 usage_scope 字段")
finally: 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("✅ 数据库模型注册成功")
logger.info("✅ 数据库初始化完成") logger.info("✅ 数据库初始化完成")
+2
View File
@@ -126,6 +126,8 @@ class AIModelConfig(Base):
api_base_url: Mapped[str | None] = mapped_column(String(256)) api_base_url: Mapped[str | None] = mapped_column(String(256))
api_key_enc: Mapped[str | None] = mapped_column(String(512)) api_key_enc: Mapped[str | None] = mapped_column(String(512))
model_version: Mapped[str | None] = mapped_column(String(64)) 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) temperature: Mapped[float] = mapped_column(Float, default=0.7)
max_tokens: Mapped[int] = mapped_column(Integer, default=1000) max_tokens: Mapped[int] = mapped_column(Integer, default=1000)
timeout_seconds: Mapped[int] = mapped_column(Integer, default=30) timeout_seconds: Mapped[int] = mapped_column(Integer, default=30)
+6
View File
@@ -158,6 +158,8 @@ class AIModelCreateRequest(BaseModel):
api_base_url: Optional[str] = None api_base_url: Optional[str] = None
api_key: Optional[str] = None api_key: Optional[str] = None
model_version: 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) temperature: float = Field(default=0.7, ge=0.0, le=2.0)
max_tokens: int = Field(default=1000, ge=1, le=32000) max_tokens: int = Field(default=1000, ge=1, le=32000)
timeout_seconds: int = Field(default=30, ge=5, le=300) timeout_seconds: int = Field(default=30, ge=5, le=300)
@@ -171,6 +173,8 @@ class AIModelUpdateRequest(BaseModel):
api_base_url: Optional[str] = None api_base_url: Optional[str] = None
api_key: Optional[str] = None api_key: Optional[str] = None
model_version: 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) temperature: Optional[float] = Field(None, ge=0.0, le=2.0)
max_tokens: Optional[int] = Field(None, ge=1, le=32000) max_tokens: Optional[int] = Field(None, ge=1, le=32000)
timeout_seconds: Optional[int] = Field(None, ge=5, le=300) timeout_seconds: Optional[int] = Field(None, ge=5, le=300)
@@ -186,6 +190,8 @@ class AIModelResponse(BaseModel):
api_base_url: Optional[str] api_base_url: Optional[str]
has_api_key: bool has_api_key: bool
model_version: Optional[str] model_version: Optional[str]
vision_model_version: Optional[str]
ocr_model_version: Optional[str]
temperature: float temperature: float
max_tokens: int max_tokens: int
timeout_seconds: int timeout_seconds: int
+33 -3
View File
@@ -1,16 +1,30 @@
import os import os
from sqlalchemy import create_engine from sqlalchemy import create_engine, event
from sqlalchemy.orm import sessionmaker, declarative_base, Session from sqlalchemy.orm import sessionmaker, declarative_base, Session
BASE_DIR = os.path.dirname(os.path.abspath(__file__)) BASE_DIR = os.path.dirname(os.path.abspath(__file__))
DB_FILE = os.path.join(BASE_DIR, "avatar.db") DB_FILE = os.path.join(BASE_DIR, "avatar.db")
DATABASE_URL = os.getenv("DATABASE_URL", f"sqlite:///{DB_FILE}") DATABASE_URL = os.getenv("DATABASE_URL", f"sqlite:///{DB_FILE}")
IS_SQLITE = DATABASE_URL.startswith("sqlite:")
engine = create_engine( engine = create_engine(
DATABASE_URL, DATABASE_URL,
connect_args={"check_same_thread": False} if DATABASE_URL.startswith("sqlite:") else {}, connect_args={"check_same_thread": False, "timeout": 30} if IS_SQLITE else {},
) )
if IS_SQLITE:
@event.listens_for(engine, "connect")
def _configure_sqlite_connection(dbapi_connection, _connection_record):
cursor = dbapi_connection.cursor()
try:
cursor.execute("PRAGMA synchronous=NORMAL")
cursor.execute("PRAGMA busy_timeout=30000")
finally:
cursor.close()
SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False) SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
Base = declarative_base() Base = declarative_base()
@@ -26,6 +40,10 @@ def get_db():
def init_db(): def init_db():
import models import models
if IS_SQLITE:
with engine.connect() as conn:
conn.exec_driver_sql("PRAGMA journal_mode=WAL")
conn.commit()
Base.metadata.create_all(bind=engine) Base.metadata.create_all(bind=engine)
# 轻量迁移:为已存在的表补充新列(SQLite 不支持自动 ALTER,逐列尝试) # 轻量迁移:为已存在的表补充新列(SQLite 不支持自动 ALTER,逐列尝试)
@@ -35,18 +53,21 @@ def init_db():
("knowledge_docs", "embedding_model", "VARCHAR DEFAULT ''"), ("knowledge_docs", "embedding_model", "VARCHAR DEFAULT ''"),
("knowledge_docs", "chunk_count", "INTEGER DEFAULT 0"), ("knowledge_docs", "chunk_count", "INTEGER DEFAULT 0"),
("knowledge_docs", "vectorized_at", "TIMESTAMP"), ("knowledge_docs", "vectorized_at", "TIMESTAMP"),
("knowledge_docs", "error_message", "VARCHAR DEFAULT ''"),
("avatars", "owner_id", "VARCHAR DEFAULT ''"), ("avatars", "owner_id", "VARCHAR DEFAULT ''"),
("authorizations", "takeover_enabled", "BOOLEAN DEFAULT 0"), ("authorizations", "takeover_enabled", "BOOLEAN DEFAULT 0"),
("authorizations", "takeover_mode", "VARCHAR DEFAULT 'immediate'"), ("authorizations", "takeover_mode", "VARCHAR DEFAULT 'immediate'"),
("authorizations", "takeover_delay_seconds", "INTEGER DEFAULT 30"), ("authorizations", "takeover_delay_seconds", "INTEGER DEFAULT 180"),
("avatars", "share_token", "VARCHAR DEFAULT NULL"), ("avatars", "share_token", "VARCHAR DEFAULT NULL"),
("token_account", "user_id", "VARCHAR DEFAULT ''"), ("token_account", "user_id", "VARCHAR DEFAULT ''"),
("token_account", "total_granted", "BIGINT DEFAULT 0"), ("token_account", "total_granted", "BIGINT DEFAULT 0"),
("token_account", "total_consumed", "BIGINT DEFAULT 0"), ("token_account", "total_consumed", "BIGINT DEFAULT 0"),
("token_account", "created_at", "TIMESTAMP"), ("token_account", "created_at", "TIMESTAMP"),
("token_account", "updated_at", "TIMESTAMP"), ("token_account", "updated_at", "TIMESTAMP"),
("takeover_messages", "attachment_id", "VARCHAR DEFAULT NULL"),
) )
_normalize_optional_unique_values() _normalize_optional_unique_values()
_normalize_takeover_delays()
_create_token_indexes() _create_token_indexes()
@@ -66,6 +87,15 @@ def _normalize_optional_unique_values():
conn.exec_driver_sql("UPDATE avatars SET share_token = NULL WHERE share_token = ''") 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(): def _create_token_indexes():
with engine.begin() as conn: with engine.begin() as conn:
conn.exec_driver_sql( conn.exec_driver_sql(
+9 -1
View File
@@ -18,6 +18,14 @@ EMBED_DIM = 256
MODEL = os.getenv("EMBEDDING_MODEL", "mock-hash-embed-v1") 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): def _tokenize(text):
text = (text or "").lower() text = (text or "").lower()
# 英文/数字按词,CJK 逐字(中文无空格,需拆到字级才能命中子词) # 英文/数字按词,CJK 逐字(中文无空格,需拆到字级才能命中子词)
@@ -47,7 +55,7 @@ def embed(texts):
"""返回 list[list[float]],与输入顺序一致。""" """返回 list[list[float]],与输入顺序一致。"""
if not texts: if not texts:
return [] return []
api_url = os.getenv("EMBEDDING_API_URL") api_url = _embedding_endpoint(os.getenv("EMBEDDING_API_URL"))
if api_url: if api_url:
api_key = os.getenv("EMBEDDING_API_KEY", "") api_key = os.getenv("EMBEDDING_API_KEY", "")
model = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small") model = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small")
+66 -1
View File
@@ -19,11 +19,14 @@ import routers.huihui_auth
import routers.chat import routers.chat
import routers.takeover import routers.takeover
from responses import ok from responses import ok
from services.chat_attachment_service import purge_expired_chat_attachments
from services.knowledge_vectorizer import knowledge_vectorizer
from services.token_billing import DEFAULT_TOKEN_GRANT, release_stale_reservations from services.token_billing import DEFAULT_TOKEN_GRANT, release_stale_reservations
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
takeover_scheduler = None takeover_scheduler = None
maintenance_scheduler = None
app = FastAPI(title="会会数字分身 API", version="1.0.0") app = FastAPI(title="会会数字分身 API", version="1.0.0")
@@ -129,9 +132,19 @@ def on_startup():
init_db() init_db()
seed() seed()
knowledge_vectorizer.start()
# Release stale resources when startup is invoked again by a reload/test. # Release stale resources when startup is invoked again by a reload/test.
stop_takeover_scheduler() stop_takeover_scheduler()
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 --- # --- Takeover scheduler ---
try: try:
@@ -152,7 +165,14 @@ def on_startup():
boxim_client = BoxIMClient(boxim_config) boxim_client = BoxIMClient(boxim_config)
from services.takeover_service import TakeoverService 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"))) poll_interval = max(0.5, float(os.getenv("BOXIM_POLL_INTERVAL_SECONDS", "1")))
takeover_scheduler = AsyncIOScheduler() takeover_scheduler = AsyncIOScheduler()
@@ -196,6 +216,51 @@ def stop_takeover_scheduler():
finally: finally:
takeover_scheduler = None 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") @app.on_event("shutdown")
def on_shutdown(): def on_shutdown():
stop_takeover_scheduler() stop_takeover_scheduler()
stop_maintenance_scheduler()
+44 -3
View File
@@ -67,7 +67,7 @@ class Authorization(Base):
status = Column(String, default="active") # active | inactive status = Column(String, default="active") # active | inactive
takeover_enabled = Column(Boolean, default=False) # 是否开启分身接管 takeover_enabled = Column(Boolean, default=False) # 是否开启分身接管
takeover_mode = Column(String, default="immediate") # immediate | delayed 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()) created_at = Column(DateTime, server_default=func.now())
def to_dict(self): def to_dict(self):
@@ -120,13 +120,14 @@ class TakeoverMessage(Base):
direction = Column(String, nullable=False) # incoming | outgoing direction = Column(String, nullable=False) # incoming | outgoing
message_type = Column(Integer, default=0) message_type = Column(Integer, default=0)
content = Column(Text, default="") content = Column(Text, default="")
attachment_id = Column(String, nullable=True)
is_avatar = Column(Boolean, default=False) is_avatar = Column(Boolean, default=False)
send_time = Column(DateTime, nullable=False) send_time = Column(DateTime, nullable=False)
created_at = Column(DateTime, server_default=func.now()) created_at = Column(DateTime, server_default=func.now())
class TakeoverReplyTask(Base): class TakeoverReplyTask(Base):
"""Restart-safe three-second BOXIM reply task.""" """Restart-safe delayed BOXIM reply task."""
__tablename__ = "takeover_reply_tasks" __tablename__ = "takeover_reply_tasks"
__table_args__ = ( __table_args__ = (
@@ -188,7 +189,8 @@ class KnowledgeDoc(Base):
file_type = Column(String, default="") # pdf | doc | docx | xlsx file_type = Column(String, default="") # pdf | doc | docx | xlsx
file_size = Column(Integer, default=0) file_size = Column(Integer, default=0)
file_url = Column(String, default="") file_url = Column(String, default="")
status = Column(String, default="uploaded") # uploaded | parsing | ready status = Column(String, default="uploaded") # uploaded | parsing | ready | failed
error_message = Column(String, default="") # 建立索引失败原因
vectorized = Column(Boolean, default=False) # 是否已向量化 vectorized = Column(Boolean, default=False) # 是否已向量化
embedding_model = Column(String, default="") # 向量模型标识 embedding_model = Column(String, default="") # 向量模型标识
chunk_count = Column(Integer, default=0) # 切片数量 chunk_count = Column(Integer, default=0) # 切片数量
@@ -204,6 +206,7 @@ class KnowledgeDoc(Base):
"fileSize": self.file_size, "fileSize": self.file_size,
"fileUrl": self.file_url, "fileUrl": self.file_url,
"status": self.status, "status": self.status,
"errorMessage": self.error_message or "",
"vectorized": bool(self.vectorized), "vectorized": bool(self.vectorized),
"embeddingModel": self.embedding_model, "embeddingModel": self.embedding_model,
"chunkCount": self.chunk_count, "chunkCount": self.chunk_count,
@@ -257,6 +260,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): class TokenAccount(Base):
__tablename__ = "token_account" __tablename__ = "token_account"
id = Column(Integer, primary_key=True) id = Column(Integer, primary_key=True)
@@ -8,3 +8,4 @@ pypdf
python-docx python-docx
openpyxl openpyxl
apscheduler>=3.10 apscheduler>=3.10
Pillow>=10.4
@@ -2,7 +2,7 @@ from fastapi import APIRouter, Body, Depends, Header, HTTPException
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from database import get_db from database import get_db
from models import Authorization, TakeoverCursor, TakeoverReplyTask from models import Authorization, Avatar, TakeoverCursor, TakeoverReplyTask
from responses import fail, ok from responses import fail, ok
from routers.avatars import _require_owned_avatar from routers.avatars import _require_owned_avatar
@@ -14,6 +14,10 @@ ALLOWED_PERMISSIONS = set(PERMISSION_ORDER)
AVATAR_PERMISSION_ORDER = PERMISSION_ORDER AVATAR_PERMISSION_ORDER = PERMISSION_ORDER
AVATAR_PERMISSION_KEY = "authorizationPermissions" AVATAR_PERMISSION_KEY = "authorizationPermissions"
DEFAULT_AVATAR_PERMISSIONS = ["friend", "chat"] 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 = { LEGACY_PERMISSION_MAP = {
"read": "browse", "read": "browse",
"reply": "chat", "reply": "chat",
@@ -91,9 +95,62 @@ def _permission_settings_payload(avatar) -> dict:
return { return {
"avatarId": avatar.id, "avatarId": avatar.id,
"permissions": _stored_avatar_permissions(avatar), "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: def _require_authorization(db: Session, avatar_id: str, authorization_id: str) -> Authorization:
authorization = ( authorization = (
db.query(Authorization) db.query(Authorization)
@@ -144,10 +201,19 @@ def update_permission_settings(
db: Session = Depends(get_db), db: Session = Depends(get_db),
): ):
avatar = _require_owned_avatar(db, avatar_id, authorization) avatar = _require_owned_avatar(db, avatar_id, authorization)
if "permissions" not in payload: if "permissions" not in payload and TAKEOVER_DELAY_KEY not in payload:
return fail("缺少 permissions", 400) return fail("缺少授权设置", 400)
try: 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: except ValueError as exc:
return fail(str(exc), 400) return fail(str(exc), 400)
@@ -155,7 +221,9 @@ def update_permission_settings(
avatar.config = { avatar.config = {
**(avatar.config or {}), **(avatar.config or {}),
AVATAR_PERMISSION_KEY: permissions, 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() cursor = db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar.id).first()
if cursor and "takeover" in permissions and "takeover" not in previous_permissions: if cursor and "takeover" in permissions and "takeover" not in previous_permissions:
cursor.initialized = False cursor.initialized = False
@@ -179,7 +247,9 @@ def update_permission_settings(
task.locked_at = None task.locked_at = None
db.commit() db.commit()
db.refresh(avatar) 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") @router.get("/avatar/{avatar_id}/authorizations")
@@ -240,7 +310,7 @@ def create_auth(
status="active", status="active",
takeover_enabled=False, takeover_enabled=False,
takeover_mode="immediate", takeover_mode="immediate",
takeover_delay_seconds=30, takeover_delay_seconds=DEFAULT_TAKEOVER_DELAY_SECONDS,
) )
db.add(item) db.add(item)
db.commit() db.commit()
+42 -16
View File
@@ -6,7 +6,17 @@ from sqlalchemy.orm import Session
from database import get_db from database import get_db
from routers.knowledge import UPLOAD_DIR 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 from responses import ok, fail
router = APIRouter(tags=["分身"]) 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}") @router.get("/avatar/{avatar_id}")
def get_avatar(avatar_id: str, db: Session = Depends(get_db)): def get_avatar(
a = db.query(Avatar).filter(Avatar.id == avatar_id).first() avatar_id: str,
if not a: authorization: str = Header(None),
return fail("分身不存在", 404) db: Session = Depends(get_db),
return ok(a.to_dict()) ):
return ok(_require_owned_avatar(db, avatar_id, authorization).to_dict())
@router.post("/avatar") @router.post("/avatar")
def create_avatar(payload: dict = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)): def create_avatar(payload: dict = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)):
user = _resolve_user(authorization, db) user = _resolve_user(authorization, db)
if not user:
raise HTTPException(status_code=401, detail="未登录")
a = Avatar( a = Avatar(
owner_id=user.huihui_user_id if user else "", owner_id=user.huihui_user_id,
name=payload.get("name", "未命名分身"), name=payload.get("name", "未命名分身"),
display_name=payload.get("displayName", "") or payload.get("display_name", ""), display_name=payload.get("displayName", "") or payload.get("display_name", ""),
description=payload.get("description", ""), description=payload.get("description", ""),
@@ -102,10 +115,13 @@ def create_avatar(payload: dict = Body(...), authorization: str = Header(None),
@router.put("/avatar/{avatar_id}") @router.put("/avatar/{avatar_id}")
def update_avatar(avatar_id: str, payload: dict = Body(...), db: Session = Depends(get_db)): def update_avatar(
a = db.query(Avatar).filter(Avatar.id == avatar_id).first() avatar_id: str,
if not a: payload: dict = Body(...),
return fail("分身不存在", 404) authorization: str = Header(None),
db: Session = Depends(get_db),
):
a = _require_owned_avatar(db, avatar_id, authorization)
mapping = { mapping = {
"displayName": "display_name", "displayName": "display_name",
"photoUrl": "photo_url", "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"): for key in ("name", "displayName", "description", "photoUrl", "emoji", "status", "tokenBalance", "config"):
if key in payload: if key in payload:
col = mapping.get(key, key) 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.commit()
db.refresh(a) db.refresh(a)
return ok(a.to_dict()) return ok(a.to_dict())
@router.delete("/avatar/{avatar_id}") @router.delete("/avatar/{avatar_id}")
def delete_avatar(avatar_id: str, db: Session = Depends(get_db)): def delete_avatar(
a = db.query(Avatar).filter(Avatar.id == avatar_id).first() avatar_id: str,
if not a: authorization: str = Header(None),
return fail("分身不存在", 404) 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(KnowledgeDoc).filter(KnowledgeDoc.avatar_id == avatar_id).delete()
db.query(KnowledgeChunk).filter(KnowledgeChunk.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(QAPair).filter(QAPair.avatar_id == avatar_id).delete()
db.query(Authorization).filter(Authorization.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.delete(a)
db.commit() db.commit()
return ok({"success": True}) return ok({"success": True})
+622 -24
View File
@@ -1,21 +1,32 @@
import difflib import difflib
import json import json
import logging
import os import os
import re import re
import secrets import secrets
import string import string
from datetime import datetime, timedelta
from typing import Any, Callable from typing import Any, Callable
import httpx 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 fastapi.responses import StreamingResponse
from pydantic import BaseModel, Field from pydantic import BaseModel, ConfigDict, Field, model_validator
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
import embeddings import embeddings
from database import get_db 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 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 ( from services.token_billing import (
InsufficientTokensError, InsufficientTokensError,
estimate_fallback_usage, estimate_fallback_usage,
@@ -24,8 +35,10 @@ from services.token_billing import (
settle_reservation, settle_reservation,
) )
from services.chat_model_config import ChatModelConfig, get_chat_model_config from services.chat_model_config import ChatModelConfig, get_chat_model_config
from services.chat_attachment_service import purge_expired_chat_attachments
router = APIRouter(tags=["数字分身聊天"]) router = APIRouter(tags=["数字分身聊天"])
logger = logging.getLogger(__name__)
MAX_MESSAGE_LENGTH = 4000 MAX_MESSAGE_LENGTH = 4000
MAX_HISTORY_MESSAGES = 10 MAX_HISTORY_MESSAGES = 10
@@ -34,16 +47,69 @@ QA_SEMANTIC_THRESHOLD = 0.72
QA_MATCH_MARGIN = 0.06 QA_MATCH_MARGIN = 0.06
KNOWLEDGE_MIN_SCORE = float(os.getenv("KNOWLEDGE_MIN_SCORE", "0.42")) KNOWLEDGE_MIN_SCORE = float(os.getenv("KNOWLEDGE_MIN_SCORE", "0.42"))
_IMAGE_ACCESS_DENIAL_PATTERNS = (
re.compile(
r"(?:我|目前|暂时|这里|本身|系统)?\s*(?:无法|不能|没法|不支持)\s*"
r"(?:直接)?\s*(?:查看|看到|看见|识别|读取|访问|打开|分析|理解)"
r"(?:\s*(?:或|、|/)\s*(?:查看|看到|看见|识别|读取|访问|打开|分析|理解))*\s*"
r"(?:你(?:发|提供|上传)的|这张|该|当前)?\s*(?:图片|图像|照片|影像|文件)"
),
re.compile(
r"(?:我|这里|目前|暂时)?\s*(?:看不到|看不见|未看到|没有看到|没收到|未收到)\s*"
r"(?:你(?:发|提供|上传)的|这张|该|当前)?\s*(?:图片|图像|照片|影像)"
),
re.compile(
r"\b(?:i\s+)?(?:can(?:not|'t)|am\s+unable\s+to)\s+(?:directly\s+)?"
r"(?:view|see|access|read|analy[sz]e|recogni[sz]e)\s+"
r"(?:the\s+|this\s+|your\s+)?(?:image|photo|picture|scan)\b",
re.IGNORECASE,
),
)
_WRITING_SYSTEM_PATTERNS = {
"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): class ChatMessage(BaseModel):
model_config = ConfigDict(populate_by_name=True)
role: str = Field(pattern="^(user|assistant)$") role: str = Field(pattern="^(user|assistant)$")
content: str = Field(min_length=1, max_length=MAX_MESSAGE_LENGTH) 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): 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) 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): def _resolve_user(authorization: str | None, db: Session):
if not authorization: if not authorization:
@@ -64,12 +130,354 @@ def _require_owned_avatar(db: Session, avatar_id: str, authorization: str | None
return avatar 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 _answer_denies_available_image(answer: str) -> bool:
"""Reject only whole-image access denials, not uncertainty about one field."""
value = re.sub(r"\s+", " ", answer or "").strip()
return any(pattern.search(value) for pattern in _IMAGE_ACCESS_DENIAL_PATTERNS)
def _compact_context_text(value: Any, limit: int) -> str:
lines = [re.sub(r"\s+", " ", line).strip() for line in str(value or "").splitlines()]
text = "\n".join(line for line in lines if line).strip()
return text[:limit].rstrip()
def _grounded_image_fallback(question: str, image_contexts: list[dict]) -> str:
"""Build a safe answer from completed vision data when the chat model contradicts it."""
summaries: list[str] = []
facts: list[str] = []
excerpts: list[str] = []
warnings: list[str] = []
for context in image_contexts:
summary = _compact_context_text(context.get("summary"), 500)
if summary:
summaries.append(summary)
structured = context.get("structuredData") or {}
if isinstance(structured, dict):
for fact in structured.get("key_facts") or []:
value = _compact_context_text(fact, 300)
if value:
facts.append(value)
extracted = _compact_context_text(context.get("extractedText"), 900)
if extracted:
excerpts.append(extracted)
warning = _compact_context_text(context.get("warning"), 300)
if warning:
warnings.append(warning)
summaries = list(dict.fromkeys(summaries))
facts = list(dict.fromkeys(facts))[:6]
excerpts = list(dict.fromkeys(excerpts))
warnings = list(dict.fromkeys(warnings))
writing_system = _dominant_writing_system(question)
if writing_system == "latin":
parts = []
if summaries:
parts.append("From the image, I can confirm: " + " ".join(summaries))
if facts:
parts.append("Key details:\n" + "\n".join(
f"{index}. {fact}" for index, fact in enumerate(facts, 1)
))
elif excerpts:
parts.append("Visible text:\n" + excerpts[0])
if warnings:
parts.append("Please note: " + " ".join(warnings))
return "\n".join(parts).strip() or "The image is available, but there is not enough clear detail to confirm more."
parts = []
if summaries:
parts.append("从这张图中可以确认:" + ";".join(summaries).rstrip("。;") + "。")
if facts:
parts.append("其中比较明确的信息有:\n" + "\n".join(
f"{index}. {fact}" for index, fact in enumerate(facts, 1)
))
elif excerpts:
parts.append("图中可见的主要文字是:\n" + excerpts[0])
if warnings:
parts.append("需要注意:" + ";".join(warnings).rstrip("。;") + "。")
return "\n".join(parts).strip() or "这张图已经看到了,但目前能确认的清晰信息比较有限。"
def _run_billed_vision_call(
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: def _normalize_question(value: str) -> str:
value = (value or "").strip().lower() value = (value or "").strip().lower()
value = re.sub(r"\s+", "", value) value = re.sub(r"\s+", "", value)
return value.translate(str.maketrans("", "", string.punctuation + ",。!?;:、()【】「」‘’“”《》")) 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: def _canonicalize_question(value: str) -> str:
value = _normalize_question(value) value = _normalize_question(value)
replacements = ( replacements = (
@@ -189,7 +597,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) config = _config(avatar)
description = (getattr(avatar, "description", "") or "").strip() description = (getattr(avatar, "description", "") or "").strip()
knowledge = "\n".join( knowledge = "\n".join(
@@ -197,6 +613,7 @@ def _build_prompt(avatar: Avatar, history: list[Any], question: str, knowledge_h
for hit in knowledge_hits for hit in knowledge_hits
if hit.get("snippet") if hit.get("snippet")
) )
image_contexts = image_contexts or []
profile_items = [ profile_items = [
(label, config[key]) (label, config[key])
for label, key in ( for label, key in (
@@ -210,7 +627,7 @@ def _build_prompt(avatar: Avatar, history: list[Any], question: str, knowledge_h
profile = ";".join(f"{label}:{value}" for label, value in profile_items) profile = ";".join(f"{label}:{value}" for label, value in profile_items)
system = ( system = (
f"你的专业或服务范围是:「{description or '未设置'}」。" f"你的专业或服务范围是:「{description or '未设置'}」。"
"请基于已提供的知识库回答,不要编造事实;" "请基于已提供的可靠资料回答,不要编造事实;"
f"回复风格:{config['replyStyle']};严谨度:{config['rigor']}/100;" f"回复风格:{config['replyStyle']};严谨度:{config['rigor']}/100;"
f"幽默感:{config['humor']}/100;回复长度:{config['responseLength']}。" f"幽默感:{config['humor']}/100;回复长度:{config['responseLength']}。"
) )
@@ -221,17 +638,49 @@ def _build_prompt(avatar: Avatar, history: list[Any], question: str, knowledge_h
) )
if config["systemPrompt"]: if config["systemPrompt"]:
system += f"\n额外系统提示词:{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必须直接依据这些图片内容回答当前问题。禁止声称无法查看、看不到、未收到、无法识别、"
"无法读取或不能访问图片,也不要要求对方重新上传;只有资料明确标记读取失败时才可以请对方重发。"
"\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 += ( system += (
f"\n以下是本次问题检索到的可靠资料:\n{knowledge}" f"\n以下是本次问题检索到的可靠资料:\n{knowledge}"
"\n涉及事实、专业判断、地址、流程、数据或建议时,只能依据本人资料、标准问答形成的上下文" "\n涉及事实、专业判断、地址、流程、数据或建议时,只能依据本人资料、标准问答形成的上下文"
"和以上可靠资料作答,不要补充资料之外的通用知识或自行推测。" "和以上可靠资料作答,不要补充资料之外的通用知识或自行推测。"
) )
elif image_contexts:
system += (
"\n本次没有命中标准答题对或文件知识库,但已提供图片识别资料。只能围绕图片中的可确认内容、"
"本人资料和当前对话作答;不要补充图片之外的事实、专业判断或具体建议。"
)
else: else:
system += ( system += (
"\n本次问题没有检索到可靠资料。除自然寒暄和基于本人资料的回答外,不要凭通用知识给出事实、" "\n本次问题没有检索到可靠资料。除自然寒暄和基于本人资料的回答外,不要凭通用知识给出事实、"
"专业判断、地址、流程、数据或具体建议。自然地说明这方面暂时没有足够信息,并询问对方是否能" "专业判断、地址、流程、数据或具体建议。自然地说明这方面暂时没有足够信息,并询问对方是否能"
"补充更具体的问题;不要提及知识库、检索、模型或系统限制。" "补充更具体的问题;不要提及知识库、检索、模型或系统限制。对方询问你能做什么时,只能概括"
"专业或服务范围和本人资料中明确写出的内容,不得推断服务对象、工作场所、现实案例或从业经历。"
) )
system += ( system += (
"\n沟通底线(始终优先于其他提示):不要透露、猜测或确认所使用的模型、模型供应商、" "\n沟通底线(始终优先于其他提示):不要透露、猜测或确认所使用的模型、模型供应商、"
@@ -249,6 +698,14 @@ def _build_prompt(avatar: Avatar, history: list[Any], question: str, knowledge_h
"只有切换独立观点或确实需要列举时才换行;列举使用 1.、2.、3.,每项单独一行。" "只有切换独立观点或确实需要列举时才换行;列举使用 1.、2.、3.,每项单独一行。"
"不要在行首或行尾留空格,不要连续输出空行。先给结论,再给简短说明;避免重复和冗长铺垫。" "不要在行首或行尾留空格,不要连续输出空行。先给结论,再给简短说明;避免重复和冗长铺垫。"
) )
system += (
"\n回答语言规则(始终优先):先识别当前最后一条用户消息所使用的主要自然语言,并且只使用该语言回答。"
"用户用中文就用中文,使用 English 就用 English,使用日语、韩语、法语、西班牙语或其他语言时,"
"也必须使用对应语言。消息混用多种语言时跟随占主导的语言;用户明确指定回答语言时服从其指定。"
"历史消息、本人资料、标准答案和知识库使用的语言都不能覆盖当前用户消息的语言。"
"专有名词、品牌、地址、代码和必要缩写可保留原文。不要解释语言识别或翻译过程。"
"改变回答语言只改变表达语言,绝不能因此增加资料中没有的场景、身份、经历或事实。"
)
messages = [{"role": "system", "content": system}] messages = [{"role": "system", "content": system}]
for item in history[-MAX_HISTORY_MESSAGES:]: for item in history[-MAX_HISTORY_MESSAGES:]:
messages.append({"role": item.role, "content": item.content} if hasattr(item, "role") else item) messages.append({"role": item.role, "content": item.content} if hasattr(item, "role") else item)
@@ -380,18 +837,45 @@ def _resolve_reply(
search_fn: Callable[..., list[dict]] | None = None, search_fn: Callable[..., list[dict]] | None = None,
model_client: Callable[..., str] | None = None, model_client: Callable[..., str] | None = None,
usage_source: str = "chat", usage_source: str = "chat",
image_contexts: list[dict] | None = None,
) -> dict: ) -> dict:
image_contexts = image_contexts or []
question = question.strip() or "请根据这张图片说明可确认的内容。"
if qa_pairs is None: if qa_pairs is None:
qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all() qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
matched = _match_standard_qa(question, qa_pairs) matched = _match_standard_qa(question, qa_pairs)
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": []} return {"answer": matched.answer, "source": "qa", "references": []}
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)) search_fn = search_fn or (lambda query, avatar_id: _search_knowledge(db, avatar_id, query))
hits = search_fn(question, avatar.id) retrieval_question = _image_retrieval_question(question, image_contexts)
messages = _build_prompt(avatar, history, question, hits) hits = search_fn(retrieval_question, avatar.id)
messages = _build_prompt(
avatar,
history,
question,
hits,
image_contexts=image_contexts,
)
config = _config(avatar) 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 token_usage = None
if model_client is not None: if model_client is not None:
answer = model_client(messages=messages, temperature=temperature) answer = model_client(messages=messages, temperature=temperature)
@@ -421,9 +905,19 @@ def _resolve_reply(
except Exception as exc: except Exception as exc:
release_reservation(db, reservation, str(exc)) release_reservation(db, reservation, str(exc))
raise raise
answer = str(answer or "").strip()
if image_contexts and _answer_denies_available_image(answer):
logger.warning(
"chat model contradicted ready image context avatar=%s source=%s",
avatar.id,
usage_source,
)
answer = _grounded_image_fallback(question, image_contexts)
result = { result = {
"answer": answer, "answer": answer,
"source": "knowledge" if hits else "qwen", "source": "qa" if matched else (
"knowledge" if hits else ("vision" if image_contexts else "qwen")
),
"references": hits, "references": hits,
} }
if token_usage: if token_usage:
@@ -439,17 +933,48 @@ def _stream_reply(
*, *,
public: bool = False, public: bool = False,
usage_source: str = "chat_stream", 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() qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
matched = _match_standard_qa(question, qa_pairs) matched = _match_standard_qa(question, qa_pairs)
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) source, references, chunks = "qa", [], _iter_text_chunks(matched.answer)
else: else:
references = _search_knowledge(db, avatar.id, question) if matched:
source = "knowledge" if references else "qwen" 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) config = _config(avatar)
temperature = min(0.45 if references else 0.25, 0.2 + config["creativity"] / 100 * 0.6) temperature = 0.0 if matched else min(
messages = _build_prompt(avatar, history, question, references) 0.45 if references else 0.25,
0.2 + config["creativity"] / 100 * 0.6,
)
model_config = get_chat_model_config() model_config = get_chat_model_config()
reservation = reserve_avatar_tokens( reservation = reserve_avatar_tokens(
db, db,
@@ -460,8 +985,6 @@ def _stream_reply(
model_config.max_tokens, model_config.max_tokens,
) )
chunks = _iter_qwen_stream(messages, temperature, model_config) chunks = _iter_qwen_stream(messages, temperature, model_config)
if matched:
messages, reservation = [], None
if public: if public:
source, references = "public", [] source, references = "public", []
@@ -550,11 +1073,62 @@ def get_shared_avatar(share_token: str, db: Session = Depends(get_db)):
return ok(_public_avatar_payload(_require_shared_avatar(db, share_token))) 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") @router.post("/public/avatar/{share_token}/chat")
def public_chat(share_token: str, body: ChatIn = Body(...), db: Session = Depends(get_db)): def public_chat(share_token: str, body: ChatIn = Body(...), db: Session = Depends(get_db)):
avatar = _require_shared_avatar(db, share_token) avatar = _require_shared_avatar(db, share_token)
image_contexts = _attachment_contexts(
_load_chat_attachments(db, avatar.id, body)
)
try: 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["references"] = []
result["source"] = "public" result["source"] = "public"
@@ -569,8 +1143,17 @@ def public_chat(share_token: str, body: ChatIn = Body(...), db: Session = Depend
@router.post("/avatar/{avatar_id}/chat") @router.post("/avatar/{avatar_id}/chat")
def chat(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)): 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) avatar = _require_owned_avatar(db, avatar_id, authorization)
image_contexts = _attachment_contexts(
_load_chat_attachments(db, avatar.id, body)
)
try: 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: except InsufficientTokensError as exc:
return fail(str(exc), code=402) return fail(str(exc), code=402)
except RuntimeError as exc: except RuntimeError as exc:
@@ -580,7 +1163,17 @@ def chat(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(N
@router.post("/avatar/{avatar_id}/chat/stream") @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)): def chat_stream(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)):
try: 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: except InsufficientTokensError as exc:
raise HTTPException(status_code=402, detail=str(exc)) from exc raise HTTPException(status_code=402, detail=str(exc)) from exc
@@ -588,13 +1181,18 @@ def chat_stream(avatar_id: str, body: ChatIn = Body(...), authorization: str = H
@router.post("/public/avatar/{share_token}/chat/stream") @router.post("/public/avatar/{share_token}/chat/stream")
def public_chat_stream(share_token: str, body: ChatIn = Body(...), db: Session = Depends(get_db)): def public_chat_stream(share_token: str, body: ChatIn = Body(...), db: Session = Depends(get_db)):
try: try:
avatar = _require_shared_avatar(db, share_token)
image_contexts = _attachment_contexts(
_load_chat_attachments(db, avatar.id, body)
)
return _stream_reply( return _stream_reply(
db, db,
_require_shared_avatar(db, share_token), avatar,
body.message, body.message,
body.history, body.history,
public=True, public=True,
usage_source="public_chat_stream", usage_source="public_chat_stream",
image_contexts=image_contexts,
) )
except InsufficientTokensError as exc: except InsufficientTokensError as exc:
raise HTTPException(status_code=402, detail=str(exc)) from exc raise HTTPException(status_code=402, detail=str(exc)) from exc
+45 -37
View File
@@ -1,7 +1,5 @@
import os import os
import json
import uuid import uuid
from datetime import datetime, timezone
from fastapi import APIRouter, UploadFile, File, Depends, Header, HTTPException from fastapi import APIRouter, UploadFile, File, Depends, Header, HTTPException
from pydantic import BaseModel from pydantic import BaseModel
@@ -11,15 +9,16 @@ from database import get_db
from models import KnowledgeDoc, QAPair, KnowledgeChunk, Avatar, User from models import KnowledgeDoc, QAPair, KnowledgeChunk, Avatar, User
from responses import ok, fail from responses import ok, fail
import embeddings import embeddings
from services.knowledge_vectorizer import knowledge_vectorizer
router = APIRouter() router = APIRouter()
BASE_DIR = os.path.dirname(os.path.abspath(__file__)) BASE_DIR = os.path.dirname(os.path.abspath(__file__))
UPLOAD_DIR = os.path.abspath(os.getenv("UPLOAD_DIR", os.path.join(BASE_DIR, "uploads"))) UPLOAD_DIR = os.path.abspath(os.getenv("UPLOAD_DIR", os.path.join(BASE_DIR, "uploads")))
os.makedirs(UPLOAD_DIR, exist_ok=True) os.makedirs(UPLOAD_DIR, exist_ok=True)
ALLOWED_EXT = {".md", ".txt", ".pdf", ".doc", ".docx", ".xlsx"} ALLOWED_EXT = {".md", ".txt", ".pdf", ".doc", ".docx", ".xlsx"}
MAX_UPLOAD_BYTES = 10 * 1024 * 1024 MAX_UPLOAD_BYTES = 50 * 1024 * 1024
UPLOAD_CHUNK_BYTES = 1024 * 1024
class QAIn(BaseModel): class QAIn(BaseModel):
@@ -82,53 +81,62 @@ async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization
os.makedirs(avatar_dir, exist_ok=True) os.makedirs(avatar_dir, exist_ok=True)
stored = f"{uuid.uuid4().hex}{ext}" stored = f"{uuid.uuid4().hex}{ext}"
path = os.path.join(avatar_dir, stored) path = os.path.join(avatar_dir, stored)
content = await file.read() file_size = 0
if len(content) > MAX_UPLOAD_BYTES: try:
return fail("文件不能超过 10MB", code=400) # Stream large files to disk so a 100MB upload does not occupy 100MB RAM.
with open(path, "wb") as f: with open(path, "wb") as f:
f.write(content) while chunk := await file.read(UPLOAD_CHUNK_BYTES):
file_size += len(chunk)
if file_size > MAX_UPLOAD_BYTES:
raise ValueError("文件不能超过 50MB")
f.write(chunk)
except ValueError as exc:
if os.path.exists(path):
os.remove(path)
return fail(str(exc), code=400)
doc = KnowledgeDoc( doc = KnowledgeDoc(
id=uuid.uuid4().hex,
avatar_id=avatar_id, avatar_id=avatar_id,
filename=file.filename, filename=file.filename,
file_type=ext.lstrip("."), file_type=ext.lstrip("."),
file_size=len(content), file_size=file_size,
file_url=f"/api/files/{avatar_id}/{stored}", file_url=f"/api/files/{avatar_id}/{stored}",
status="parsing", status="parsing",
) )
# Persist and acknowledge the upload first. Extraction and embeddings may take
# minutes for a PDF and must never consume the browser request timeout.
db.add(doc) db.add(doc)
db.commit() db.commit()
db.refresh(doc) db.refresh(doc)
knowledge_vectorizer.enqueue(doc.id)
# 向量化:抽取文本 -> 分块 -> 调第三方/本地嵌入 -> 存切片 return ok(_doc_payload(doc))
try:
text = embeddings.extract_text(path, ext)
chunks = embeddings.chunk_text(text) @router.post("/avatar/{avatar_id}/knowledge/docs/{doc_id}/retry")
if chunks: def retry_doc(avatar_id: str, doc_id: str, authorization: str = Header(None), db: Session = Depends(get_db)):
vectors = embeddings.embed(chunks) _require_owned_avatar(db, avatar_id, authorization)
for i, (c, v) in enumerate(zip(chunks, vectors)): doc = db.query(KnowledgeDoc).filter(
db.add( KnowledgeDoc.id == doc_id, KnowledgeDoc.avatar_id == avatar_id
KnowledgeChunk( ).first()
doc_id=doc.id, if not doc:
avatar_id=avatar_id, return fail("文档不存在", code=404)
content=c, if doc.vectorized and doc.status == "ready":
vector=json.dumps(v), return ok(_doc_payload(doc))
chunk_index=i, stored_name = os.path.basename(doc.file_url or "")
embedding_model=embeddings.MODEL, if not stored_name or not os.path.isfile(os.path.join(UPLOAD_DIR, avatar_id, stored_name)):
) return fail("原文件不可用,请重新上传", code=400)
) db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == doc.id).delete()
doc.vectorized = True doc.status = "parsing"
doc.embedding_model = embeddings.MODEL doc.vectorized = False
doc.chunk_count = len(chunks) doc.embedding_model = ""
doc.vectorized_at = datetime.now(timezone.utc) doc.chunk_count = 0
doc.status = "ready" doc.vectorized_at = None
doc.error_message = ""
db.commit() db.commit()
db.refresh(doc) db.refresh(doc)
except Exception as e: knowledge_vectorizer.enqueue(doc.id)
print("vectorize failed:", e)
doc.status = "ready" # 上传成功但向量化失败,仍可展示
db.commit()
db.refresh(doc)
return ok(_doc_payload(doc)) return ok(_doc_payload(doc))
+26 -5
View File
@@ -8,13 +8,25 @@ from sqlalchemy.orm import Session
from database import get_db from database import get_db
from models import TakeoverCursor, TakeoverReplyTask, User from models import TakeoverCursor, TakeoverReplyTask, User
from responses import fail, ok 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 from routers.avatars import _require_owned_avatar
router = APIRouter(tags=["分身接管"]) router = APIRouter(tags=["分身接管"])
BOXIM_STATUS_FRESH_SECONDS = 60 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") @router.get("/avatar/{avatar_id}/takeover/status")
def get_takeover_status( def get_takeover_status(
avatar_id: str, avatar_id: str,
@@ -24,6 +36,7 @@ def get_takeover_status(
avatar = _require_owned_avatar(db, avatar_id, authorization) avatar = _require_owned_avatar(db, avatar_id, authorization)
permissions = (avatar.config or {}).get("authorizationPermissions", []) permissions = (avatar.config or {}).get("authorizationPermissions", [])
enabled = isinstance(permissions, list) and "takeover" in permissions 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() user = db.query(User).filter(User.huihui_user_id == avatar.owner_id).first()
cursor = db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar.id).first() cursor = db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar.id).first()
pending_count = ( pending_count = (
@@ -49,7 +62,10 @@ def get_takeover_status(
and cursor.last_polled_at and cursor.last_polled_at
>= datetime.utcnow() - timedelta(seconds=BOXIM_STATUS_FRESH_SECONDS) >= datetime.utcnow() - timedelta(seconds=BOXIM_STATUS_FRESH_SECONDS)
): ):
status, message = "ready", "BOXIM 已连接,收到私聊消息 3 秒后自动回复" status, message = (
"ready",
f"BOXIM 已连接,收到私聊消息 {_delay_label(reply_delay_seconds)}后自动回复",
)
else: else:
status, message = "connecting", "正在连接 BOXIM" status, message = "connecting", "正在连接 BOXIM"
@@ -59,6 +75,7 @@ def get_takeover_status(
"status": status, "status": status,
"message": message, "message": message,
"pendingCount": pending_count, "pendingCount": pending_count,
"takeoverReplyDelaySeconds": reply_delay_seconds,
"lastPolledAt": cursor.last_polled_at.isoformat() if cursor and cursor.last_polled_at else None, "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)) auth = _require_authorization(db, avatar_id, str(auth_id))
enabled = bool(auth.takeover_enabled) enabled = bool(auth.takeover_enabled)
mode = auth.takeover_mode or "immediate" 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"): if _has(payload, "takeoverEnabled", "takeover_enabled"):
raw_enabled = _read(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"): if _has(payload, "takeoverDelaySeconds", "takeover_delay_seconds"):
delay = _read(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: if (
return fail("延迟时间需在 5 到 3600 秒之间", 400) 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": if enabled and auth.target_type != "user":
return fail("本期仅支持对会会用户开启单聊接管", 400) 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 model: str
max_tokens: int max_tokens: int
timeout_seconds: float timeout_seconds: float
vision_model: str
ocr_model: str
vision_max_tokens: int
vision_timeout_seconds: float
source: str source: str
@@ -33,6 +37,10 @@ def _environment_config() -> ChatModelConfig:
model=os.getenv("CHAT_MODEL", "qwen-plus"), model=os.getenv("CHAT_MODEL", "qwen-plus"),
max_tokens=max(128, int(os.getenv("CHAT_MAX_OUTPUT_TOKENS", "1024"))), max_tokens=max(128, int(os.getenv("CHAT_MAX_OUTPUT_TOKENS", "1024"))),
timeout_seconds=max(5.0, float(os.getenv("CHAT_TIMEOUT_SECONDS", "30"))), 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", source="environment",
) )
@@ -60,6 +68,20 @@ def _fetch_runtime_config() -> ChatModelConfig | None:
model=model, model=model,
max_tokens=max(128, int(payload.get("max_tokens") or 1024)), max_tokens=max(128, int(payload.get("max_tokens") or 1024)),
timeout_seconds=max(5.0, float(payload.get("timeout_seconds") or 30)), 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", source="admin",
) )
@@ -0,0 +1,125 @@
"""Durable, serial knowledge-document indexing for the avatar knowledge base."""
import json
import logging
import os
import queue
import threading
from datetime import datetime, timezone
from database import SessionLocal
from models import KnowledgeChunk, KnowledgeDoc
import embeddings
logger = logging.getLogger(__name__)
BACKEND_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
UPLOAD_DIR = os.path.abspath(
os.getenv("UPLOAD_DIR", os.path.join(BACKEND_DIR, "routers", "uploads"))
)
class KnowledgeVectorizer:
"""Indexes one document at a time so slow providers cannot block uploads."""
def __init__(self):
self._queue: queue.Queue[str] = queue.Queue()
self._queued: set[str] = set()
self._lock = threading.Lock()
self._thread: threading.Thread | None = None
def start(self):
if self._thread and self._thread.is_alive():
return
self._thread = threading.Thread(
target=self._run, name="knowledge-vectorizer", daemon=True
)
self._thread.start()
db = SessionLocal()
try:
# A process restart must not abandon documents already accepted by upload.
for (doc_id,) in db.query(KnowledgeDoc.id).filter(KnowledgeDoc.status == "parsing"):
self.enqueue(doc_id)
finally:
db.close()
def enqueue(self, doc_id: str):
with self._lock:
if doc_id in self._queued:
return
self._queued.add(doc_id)
self._queue.put(doc_id)
def _run(self):
while True:
doc_id = self._queue.get()
try:
self.vectorize_document(doc_id)
except Exception:
logger.exception("Unexpected knowledge vectorizer failure for %s", doc_id)
finally:
with self._lock:
self._queued.discard(doc_id)
self._queue.task_done()
def vectorize_document(self, doc_id: str):
db = SessionLocal()
try:
doc = db.get(KnowledgeDoc, doc_id)
if not doc or doc.status != "parsing":
return
stored_name = os.path.basename(doc.file_url or "")
path = os.path.join(UPLOAD_DIR, doc.avatar_id, stored_name)
if not stored_name or not os.path.isfile(path):
raise FileNotFoundError("原文件不可用,请重新上传")
text = embeddings.extract_text(path, f".{doc.file_type}")
chunks = embeddings.chunk_text(text)
if not chunks:
raise ValueError("文档没有可建立索引的文字内容")
vectors = embeddings.embed(chunks)
if len(vectors) != len(chunks):
raise ValueError("向量服务返回数量与文档分段不一致")
# Commit the document and every chunk together. Chat only sees complete indexes.
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == doc.id).delete()
db.add_all(
[
KnowledgeChunk(
doc_id=doc.id,
avatar_id=doc.avatar_id,
content=chunk,
vector=json.dumps(vector),
chunk_index=index,
embedding_model=embeddings.MODEL,
)
for index, (chunk, vector) in enumerate(zip(chunks, vectors))
]
)
doc.vectorized = True
doc.embedding_model = embeddings.MODEL
doc.chunk_count = len(chunks)
doc.vectorized_at = datetime.now(timezone.utc)
doc.status = "ready"
doc.error_message = ""
db.commit()
logger.info("Knowledge document %s indexed with %s chunks", doc.id, len(chunks))
except Exception as exc:
db.rollback()
failed_doc = db.get(KnowledgeDoc, doc_id)
if failed_doc:
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == failed_doc.id).delete()
failed_doc.status = "failed"
failed_doc.vectorized = False
failed_doc.embedding_model = ""
failed_doc.chunk_count = 0
failed_doc.vectorized_at = None
failed_doc.error_message = str(exc)[:300] or "建立知识索引失败"
db.commit()
logger.exception("Knowledge vectorization failed for %s: %s", doc_id, exc)
finally:
db.close()
knowledge_vectorizer = KnowledgeVectorizer()
@@ -3,6 +3,7 @@
import asyncio import asyncio
import hashlib import hashlib
import logging import logging
import os
import re import re
import secrets import secrets
import time import time
@@ -13,21 +14,50 @@ from sqlalchemy.orm import Session
from models import ( from models import (
Avatar, Avatar,
ChatAttachment,
TakeoverCursor, TakeoverCursor,
TakeoverMessage, TakeoverMessage,
TakeoverReplyTask, TakeoverReplyTask,
User, User,
) )
from services.boxim_client import BoxIMClient, BoxIMError from services.boxim_client import BoxIMClient, BoxIMError
from services.boxim_image_service import (
BoxIMImageError,
download_boxim_image,
parse_boxim_image_url,
)
from services.vision_service import ImageValidationError
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
ACTIVE_TASK_STATUSES = ("pending", "generating", "ready", "sending") ACTIVE_TASK_STATUSES = ("pending", "generating", "ready", "sending")
GENERATABLE_TASK_STATUSES = ("pending",) GENERATABLE_TASK_STATUSES = ("pending",)
MAX_PROMPT_LENGTH = 4000 MAX_PROMPT_LENGTH = 4000
MAX_STALE_SECONDS = 120 DEFAULT_MAX_MESSAGE_AGE_SECONDS = 600
MAX_SEND_OVERDUE_SECONDS = 120
STUCK_LOCK_SECONDS = 90 STUCK_LOCK_SECONDS = 90
TAKEOVER_PERMISSION = "takeover" 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 = "这张图片我暂时没看清,麻烦重新发送一张清晰的原图。"
IMAGE_CONTEXT_LOOKBACK_SECONDS = 1800
IMAGE_REFERENCE_LOOKBACK_SECONDS = 172_800
MAX_RECENT_IMAGE_CONTEXTS = 3
_IMAGE_REFERENCE_PATTERN = re.compile(
r"(?:图片|图像|照片|截图|这张图|刚才.{0,8}图|病例|病历|检查单|检验单|化验单|报告|影像|"
r"\b(?:image|photo|picture|screenshot|scan|report)\b)",
re.IGNORECASE,
)
def _utcnow() -> datetime: def _utcnow() -> datetime:
@@ -70,24 +100,69 @@ def _plain_text_reply(value: str) -> str:
return "\n".join(line for line in lines if line).strip() 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 ""
def _references_recent_image(value: str) -> bool:
return bool(_IMAGE_REFERENCE_PATTERN.search(value or ""))
class TakeoverService: 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__( def __init__(
self, self,
session_factory: Callable[[], Session], session_factory: Callable[[], Session],
boxim_client: BoxIMClient, 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, now: Callable[[], datetime] = _utcnow,
): ):
self.session_factory = session_factory self.session_factory = session_factory
self.boxim = boxim_client self.boxim = boxim_client
self.reply_delay_seconds = reply_delay_seconds 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.now = now
self._sessions: dict[str, dict] = {} self._sessions: dict[str, dict] = {}
self._poll_lock = asyncio.Lock() self._poll_lock = asyncio.Lock()
self._process_lock = asyncio.Lock() self._process_lock = asyncio.Lock()
self._persist_lock = asyncio.Lock()
async def poll_and_process_messages(self): async def poll_and_process_messages(self):
"""Run one complete cycle for callers that do not use the split scheduler.""" """Run one complete cycle for callers that do not use the split scheduler."""
@@ -102,8 +177,48 @@ class TakeoverService:
self._recover_stuck_tasks() self._recover_stuck_tasks()
avatar_ids = self._enabled_avatar_ids() avatar_ids = self._enabled_avatar_ids()
self._cancel_disabled_tasks(set(avatar_ids)) self._cancel_disabled_tasks(set(avatar_ids))
for avatar_id in avatar_ids: self._ensure_takeover_cursors(avatar_ids)
await self._sync_avatar(avatar_id) 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): async def process_reply_tasks(self):
"""Generate and send replies independently from BOXIM's long poll.""" """Generate and send replies independently from BOXIM's long poll."""
@@ -119,11 +234,17 @@ class TakeoverService:
def _enabled_avatar_ids(self) -> list[str]: def _enabled_avatar_ids(self) -> list[str]:
db = self.session_factory() db = self.session_factory()
try: try:
return [ avatars = (
avatar.id db.query(Avatar)
for avatar in db.query(Avatar).filter(Avatar.status == "active").all() .filter(Avatar.status == "active")
if _takeover_enabled(avatar) .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: finally:
db.close() db.close()
@@ -201,13 +322,20 @@ class TakeoverService:
def _forget_boxim_session(self, user_id: str): def _forget_boxim_session(self, user_id: str):
self._sessions.pop(user_id, None) self._sessions.pop(user_id, None)
def _disable_after_connection_failure( def _record_connection_failure(
self, self,
db: Session, db: Session,
avatar: Avatar, avatar: Avatar,
cursor: TakeoverCursor, cursor: TakeoverCursor,
message: str, 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", []) permissions = (avatar.config or {}).get("authorizationPermissions", [])
avatar.config = { avatar.config = {
**(avatar.config or {}), **(avatar.config or {}),
@@ -217,8 +345,6 @@ class TakeoverService:
if permission != TAKEOVER_PERMISSION if permission != TAKEOVER_PERMISSION
], ],
} }
cursor.last_error = message
cursor.last_polled_at = self.now()
tasks = ( tasks = (
db.query(TakeoverReplyTask) db.query(TakeoverReplyTask)
.filter( .filter(
@@ -243,13 +369,17 @@ class TakeoverService:
if not cursor: if not cursor:
cursor = TakeoverCursor(avatar_id=avatar.id, owner_id=avatar.owner_id) cursor = TakeoverCursor(avatar_id=avatar.id, owner_id=avatar.owner_id)
db.add(cursor) 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: if not user or not user.huihui_token:
self._disable_after_connection_failure( self._record_connection_failure(
db, db,
avatar, avatar,
cursor, cursor,
"请重新登录会会生产账号后再开启主动接管", "请重新登录会会生产账号后再开启主动接管",
disable_takeover=True,
) )
db.commit() db.commit()
return False return False
@@ -268,11 +398,24 @@ class TakeoverService:
if isinstance(exc, BoxIMError) and exc.auth_error: if isinstance(exc, BoxIMError) and exc.auth_error:
self._forget_boxim_session(user.id) self._forget_boxim_session(user.id)
message = "BOXIM 授权已失效,请重新登录会会生产账号" message = "BOXIM 授权已失效,请重新登录会会生产账号"
disable_takeover = True
else: else:
message = f"BOXIM 暂时连接失败:{str(exc)[:160]}" 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() 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 return False
messages.sort(key=lambda item: (_numeric_id(item.get("id")), item.get("sendTime") or 0)) messages.sort(key=lambda item: (_numeric_id(item.get("id")), item.get("sendTime") or 0))
@@ -280,13 +423,6 @@ class TakeoverService:
max_message_id = _numeric_id(cursor.last_message_id) max_message_id = _numeric_id(cursor.last_message_id)
read_receipts: dict[str, int] = {} read_receipts: dict[str, int] = {}
for message in messages: for message in messages:
self._record_message(
db,
avatar,
cursor.boxim_owner_id,
message,
schedule_reply=not priming,
)
message_id = _numeric_id(message.get("id")) message_id = _numeric_id(message.get("id"))
max_message_id = max(max_message_id, message_id) max_message_id = max(max_message_id, message_id)
send_id = str(message.get("sendId") or "") send_id = str(message.get("sendId") or "")
@@ -301,6 +437,17 @@ class TakeoverService:
session["access_token"], peer_id, message_id session["access_token"], peer_id, message_id
) )
# Keep SQLite write transactions short. The read-receipt request above
# can block on the network and must not hold the database write lock.
async with self._persist_lock:
for message in messages:
self._record_message(
db,
avatar,
cursor.boxim_owner_id,
message,
schedule_reply=not priming,
)
cursor.last_message_id = str(max_message_id) cursor.last_message_id = str(max_message_id)
cursor.initialized = True cursor.initialized = True
cursor.last_polled_at = self.now() cursor.last_polled_at = self.now()
@@ -350,13 +497,21 @@ class TakeoverService:
now = self.now() now = self.now()
send_time = _boxim_time(message.get("sendTime"), now) send_time = _boxim_time(message.get("sendTime"), now)
is_avatar = False is_avatar = _is_avatar_local_id(local_id)
if direction == "outgoing" and local_id: if not is_avatar and local_id:
is_avatar = bool( is_avatar = bool(
db.query(TakeoverReplyTask) db.query(TakeoverReplyTask)
.filter( .filter(
TakeoverReplyTask.owner_id == avatar.owner_id,
TakeoverReplyTask.boxim_local_id == local_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", TakeoverReplyTask.status == "sent",
) )
.first() .first()
@@ -381,12 +536,96 @@ class TakeoverService:
if not is_avatar: if not is_avatar:
self._cancel_conversation(db, avatar.owner_id, peer_id, "owner_replied") self._cancel_conversation(db, avatar.owner_id, peer_id, "owner_replied")
return return
if not schedule_reply or event.message_type != 0 or not event.content.strip(): if not schedule_reply or event.message_type not in {
BOXIM_TEXT_MESSAGE_TYPE,
BOXIM_IMAGE_MESSAGE_TYPE,
}:
return return
if (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 return
self._schedule_reply(db, avatar, event) 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 @staticmethod
def _cancel_conversation(db: Session, owner_id: str, peer_id: str, reason: str): def _cancel_conversation(db: Session, owner_id: str, peer_id: str, reason: str):
tasks = ( tasks = (
@@ -424,12 +663,25 @@ class TakeoverService:
task.status = "cancelled" task.status = "cancelled"
task.cancel_reason = "newer_incoming_message" task.cancel_reason = "newer_incoming_message"
task.locked_at = None task.locked_at = None
prompt_parts.append(event.content.strip()) if event.message_type == BOXIM_TEXT_MESSAGE_TYPE:
for image_event in self._recent_unhandled_images(
db,
avatar,
event,
source_ids,
):
prompt_parts.append(_event_prompt(image_event))
source_ids.append(image_event.boxim_message_id)
prompt_parts.append(_event_prompt(event))
source_ids.append(event.boxim_message_id) source_ids.append(event.boxim_message_id)
prompt = "\n".join(part for part in prompt_parts if part).strip()[-MAX_PROMPT_LENGTH:] prompt = "\n".join(part for part in prompt_parts if part).strip()[-MAX_PROMPT_LENGTH:]
due_at = 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) 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( db.add(
TakeoverReplyTask( TakeoverReplyTask(
id=task_id, id=task_id,
@@ -445,6 +697,68 @@ class TakeoverService:
) )
) )
@staticmethod
def _recent_unhandled_images(
db: Session,
avatar: Avatar,
event: TakeoverMessage,
current_source_ids: list[str],
) -> list[TakeoverMessage]:
"""Recover missed images, or reuse a referenced image from the last two days."""
references_image = _references_recent_image(event.content)
lookback_seconds = (
IMAGE_REFERENCE_LOOKBACK_SECONDS
if references_image
else IMAGE_CONTEXT_LOOKBACK_SECONDS
)
threshold = event.send_time - timedelta(seconds=lookback_seconds)
candidates = (
db.query(TakeoverMessage)
.filter(
TakeoverMessage.avatar_id == avatar.id,
TakeoverMessage.owner_id == avatar.owner_id,
TakeoverMessage.peer_id == event.peer_id,
TakeoverMessage.direction == "incoming",
TakeoverMessage.message_type == BOXIM_IMAGE_MESSAGE_TYPE,
TakeoverMessage.is_avatar.is_(False),
TakeoverMessage.send_time >= threshold,
TakeoverMessage.send_time <= event.send_time,
)
.order_by(TakeoverMessage.send_time.desc())
.limit(MAX_RECENT_IMAGE_CONTEXTS)
.all()
)
if not candidates:
return []
current_ids = set(current_source_ids)
if references_image:
return [
image
for image in reversed(candidates)
if image.boxim_message_id not in current_ids
]
handled_ids = set(current_ids)
task_sources = (
db.query(TakeoverReplyTask.source_message_ids)
.filter(
TakeoverReplyTask.avatar_id == avatar.id,
TakeoverReplyTask.owner_id == avatar.owner_id,
TakeoverReplyTask.peer_id == event.peer_id,
TakeoverReplyTask.created_at >= threshold,
)
.all()
)
for (source_message_ids,) in task_sources:
handled_ids.update(source_message_ids or [])
return [
image
for image in reversed(candidates)
if image.boxim_message_id not in handled_ids
]
async def _prepare_replies(self) -> int: async def _prepare_replies(self) -> int:
db = self.session_factory() db = self.session_factory()
try: try:
@@ -455,6 +769,7 @@ class TakeoverService:
.filter( .filter(
TakeoverReplyTask.status.in_(GENERATABLE_TASK_STATUSES), TakeoverReplyTask.status.in_(GENERATABLE_TASK_STATUSES),
TakeoverReplyTask.response_text == "", TakeoverReplyTask.response_text == "",
TakeoverReplyTask.scheduled_at <= self.now(),
) )
.order_by(TakeoverReplyTask.created_at.asc()) .order_by(TakeoverReplyTask.created_at.asc())
.limit(10) .limit(10)
@@ -478,6 +793,50 @@ class TakeoverService:
results = await asyncio.gather(*(generate(task_id) for task_id in task_ids)) results = await asyncio.gather(*(generate(task_id) for task_id in task_ids))
return sum(bool(result) for result in results) return sum(bool(result) for result in results)
def _takeover_image_attachment(
self,
db: Session,
avatar: Avatar,
event: TakeoverMessage,
) -> ChatAttachment:
now = self.now()
if event.attachment_id:
cached = db.get(ChatAttachment, event.attachment_id)
if cached and cached.status == "ready" and cached.expires_at > now:
cached.used_at = now
db.commit()
return cached
downloaded = download_boxim_image(
event.content,
base_url=getattr(
self.boxim,
"im_base_url",
os.getenv("BOXIM_API_BASE_URL", "https://im.99hui.com/api"),
),
)
from routers.chat import _analyze_image_bytes
attachment = _analyze_image_bytes(
db,
avatar,
downloaded.content,
filename=downloaded.filename,
mime_type=downloaded.mime_type,
uploader_kind="boxim",
)
event.attachment_id = attachment.id
attachment.used_at = now
db.commit()
logger.info(
"BOXIM image analyzed message=%s attachment=%s avatar=%s category=%s",
event.boxim_message_id,
attachment.id,
avatar.id,
attachment.category,
)
return attachment
def _generate_reply(self, task_id: str) -> bool: def _generate_reply(self, task_id: str) -> bool:
db = self.session_factory() db = self.session_factory()
try: try:
@@ -496,19 +855,59 @@ class TakeoverService:
db.commit() db.commit()
excluded_ids = set(task.source_message_ids or []) excluded_ids = set(task.source_message_ids or [])
source_events = {
event.boxim_message_id: event
for event in (
db.query(TakeoverMessage)
.filter(
TakeoverMessage.owner_id == task.owner_id,
TakeoverMessage.peer_id == task.peer_id,
TakeoverMessage.avatar_id == task.avatar_id,
TakeoverMessage.boxim_message_id.in_(excluded_ids),
)
.all()
if excluded_ids
else []
)
}
events = ( events = (
db.query(TakeoverMessage) db.query(TakeoverMessage)
.filter( .filter(
TakeoverMessage.owner_id == task.owner_id, TakeoverMessage.owner_id == task.owner_id,
TakeoverMessage.peer_id == task.peer_id, TakeoverMessage.peer_id == task.peer_id,
TakeoverMessage.avatar_id == task.avatar_id,
) )
.order_by(TakeoverMessage.send_time.desc()) .order_by(TakeoverMessage.send_time.desc())
.limit(30) .limit(30)
.all() .all()
) )
image_attachments = []
image_failed = False
for message_id in (task.source_message_ids or [])[-3:]:
event = source_events.get(message_id)
if not event or event.message_type != BOXIM_IMAGE_MESSAGE_TYPE:
continue
try:
image_attachments.append(
self._takeover_image_attachment(db, avatar, event)
)
except (BoxIMImageError, ImageValidationError) as exc:
image_failed = True
logger.warning(
"BOXIM image unavailable message=%s avatar=%s: %s",
event.boxim_message_id,
avatar.id,
exc,
)
history = [] history = []
for event in reversed(events): for event in reversed(events):
if event.boxim_message_id in excluded_ids or not event.content.strip(): if (
event.boxim_message_id in excluded_ids
or event.message_type != BOXIM_TEXT_MESSAGE_TYPE
or not event.content.strip()
):
continue
if event.direction == "incoming" and event.is_avatar:
continue continue
history.append( history.append(
{ {
@@ -518,9 +917,20 @@ class TakeoverService:
) )
history = history[-10:] history = history[-10:]
from routers.chat import _resolve_reply from routers.chat import _attachment_contexts, _resolve_reply
result = _resolve_reply(db, avatar, task.prompt, history, usage_source="takeover") image_contexts = _attachment_contexts(image_attachments)
if image_failed and not image_contexts:
answer = BOXIM_IMAGE_UNAVAILABLE_REPLY
else:
result = _resolve_reply(
db,
avatar,
task.prompt,
history,
usage_source="takeover",
image_contexts=image_contexts,
)
answer = _plain_text_reply(result.get("answer", "")) answer = _plain_text_reply(result.get("answer", ""))
db.refresh(task) db.refresh(task)
if task.status != "generating": if task.status != "generating":
@@ -582,11 +992,20 @@ class TakeoverService:
task.cancel_reason = "takeover_disabled" task.cancel_reason = "takeover_disabled"
db.commit() db.commit()
return False 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.status = "cancelled"
task.cancel_reason = "stale_reply" task.cancel_reason = "stale_reply"
db.commit() db.commit()
return False 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() user = db.query(User).filter(User.huihui_user_id == task.owner_id).first()
if not user or not user.huihui_token: if not user or not user.huihui_token:
raise BoxIMError("缺少会会登录凭证", auth_error=True) raise BoxIMError("缺少会会登录凭证", auth_error=True)
@@ -79,12 +79,17 @@ def reserve_avatar_tokens(
model: str, model: str,
messages: list[dict], messages: list[dict],
max_output_tokens: int, max_output_tokens: int,
*,
minimum_reserve_tokens: int = 0,
) -> TokenReservation: ) -> TokenReservation:
user = avatar_owner_user(db, avatar) user = avatar_owner_user(db, avatar)
if not user: if not user:
raise InsufficientTokensError("分身尚未关联有效用户,暂时无法使用积分") raise InsufficientTokensError("分身尚未关联有效用户,暂时无法使用积分")
account = get_or_create_account(db, user.id) 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 = ( updated = (
db.query(TokenAccount) db.query(TokenAccount)
.filter(TokenAccount.id == account.id, TokenAccount.balance >= reserved) .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 ( from models import (
Authorization, Authorization,
Avatar, Avatar,
ChatAttachment,
TakeoverCursor, TakeoverCursor,
TakeoverMessage, TakeoverMessage,
TakeoverReplyTask, TakeoverReplyTask,
@@ -95,6 +96,9 @@ def authorization_context():
finally: finally:
db.rollback() db.rollback()
avatar_ids = [avatar.id, other_avatar.id] 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( db.query(TakeoverReplyTask).filter(
TakeoverReplyTask.avatar_id.in_(avatar_ids) TakeoverReplyTask.avatar_id.in_(avatar_ids)
).delete(synchronize_session=False) ).delete(synchronize_session=False)
@@ -103,6 +103,7 @@ def test_avatar_permission_settings_default_and_persist(authorization_context):
assert initial["data"] == { assert initial["data"] == {
"avatarId": context["avatar"].id, "avatarId": context["avatar"].id,
"permissions": ["friend", "chat"], "permissions": ["friend", "chat"],
"takeoverReplyDelaySeconds": 180,
} }
updated = client.put( 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() reloaded = client.get(endpoint, headers=context["owner_headers"]).json()
assert reloaded["data"]["permissions"] == ["friend", "publish", "interact", "takeover"] assert reloaded["data"]["permissions"] == ["friend", "publish", "interact", "takeover"]
assert reloaded["data"]["takeoverReplyDelaySeconds"] == 180
def test_avatar_permission_settings_allow_all_disabled(authorization_context): 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) unauthenticated = client.get(endpoint)
assert unauthenticated.status_code == 401 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,332 @@
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,
_answer_denies_available_image,
_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_ready_image_context_never_returns_whole_image_access_denial():
avatar = SimpleNamespace(
id="avatar-vision",
name="测试分身",
description="产品顾问",
config={},
)
model = Mock(return_value="抱歉,我无法查看或识别图片,请重新上传。")
result = _resolve_reply(
None,
avatar,
"请看看这张图片",
[],
qa_pairs=[],
search_fn=Mock(return_value=[]),
model_client=model,
image_contexts=[{
"id": "attachment",
"filename": "report.jpg",
"category": "medical_document",
"summary": "一份耳鼻喉科门诊记录",
"extractedText": "主诉:咽痛三天",
"structuredData": {"key_facts": ["主诉为咽痛三天"]},
"warning": "请核对原始资料",
}],
)
assert result["source"] == "vision"
assert "一份耳鼻喉科门诊记录" in result["answer"]
assert "主诉为咽痛三天" in result["answer"]
assert "无法查看" not in result["answer"]
system = model.call_args.kwargs["messages"][0]["content"]
assert "当前会话图片已经成功读取" in system
assert "禁止声称无法查看" in system
def test_image_denial_detector_allows_uncertain_field_in_ready_image():
assert _answer_denies_available_image("我无法查看这张图片") is True
assert _answer_denies_available_image("图片中患者姓名无法辨认,主诉为咽痛三天。") is False
def test_attachment_context_does_not_expose_internal_fields():
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_base_url": "https://model.test/v1/",
"api_key": "runtime-key", "api_key": "runtime-key",
"model": "avatar-model", "model": "avatar-model",
"vision_model": "avatar-vision-model",
"ocr_model": "avatar-ocr-model",
"max_tokens": 2048, "max_tokens": 2048,
"timeout_seconds": 42, "timeout_seconds": 42,
} }
@@ -37,6 +39,8 @@ def test_admin_runtime_config_takes_priority(monkeypatch):
assert config.source == "admin" assert config.source == "admin"
assert config.api_base_url == "https://model.test/v1" assert config.api_base_url == "https://model.test/v1"
assert config.model == "avatar-model" assert config.model == "avatar-model"
assert config.vision_model == "avatar-vision-model"
assert config.ocr_model == "avatar-ocr-model"
assert config.max_tokens == 2048 assert config.max_tokens == 2048
request.assert_called_once_with( request.assert_called_once_with(
"http://config.test/runtime", "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_URL", "https://fallback.test/v1/")
monkeypatch.setenv("CHAT_API_KEY", "fallback-key") monkeypatch.setenv("CHAT_API_KEY", "fallback-key")
monkeypatch.setenv("CHAT_MODEL", "fallback-model") 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") monkeypatch.setenv("CHAT_MAX_OUTPUT_TOKENS", "1536")
request = httpx.Request("GET", "http://config.test/runtime") 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_base_url == "https://fallback.test/v1"
assert config.api_key == "fallback-key" assert config.api_key == "fallback-key"
assert config.model == "fallback-model" assert config.model == "fallback-model"
assert config.vision_model == "fallback-vision"
assert config.ocr_model == "fallback-ocr"
assert config.max_tokens == 1536 assert config.max_tokens == 1536
@@ -5,7 +5,15 @@ from unittest.mock import Mock
from fastapi import HTTPException from fastapi import HTTPException
from models import Avatar, User 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): class ChatOrchestrationTests(unittest.TestCase):
@@ -50,6 +58,34 @@ class ChatOrchestrationTests(unittest.TestCase):
self.assertEqual(result["answer"], "标准地址") self.assertEqual(result["answer"], "标准地址")
fake_model.assert_not_called() 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): def test_conversational_paraphrase_matches_standard_qa(self):
for question in ("请问一下,你们公司在哪里呀?", "请问去你们那边怎么走"): for question in ("请问一下,你们公司在哪里呀?", "请问去你们那边怎么走"):
with self.subTest(question=question): 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"])
self.assertIn("回答语言规则", messages[0]["content"])
self.assertIn("当前最后一条用户消息", messages[0]["content"])
self.assertIn("历史消息", messages[0]["content"])
def test_prompt_blocks_ungrounded_factual_answers(self): def test_prompt_blocks_ungrounded_factual_answers(self):
messages = _build_prompt(self.avatar, [], "聊聊国际新闻", []) 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) self.assertIn("不要提及知识库", system)
self.assertIn("不得推断服务对象", system)
self.assertIn("工作场所", system)
def test_public_avatar_payload_excludes_internal_configuration(self): def test_public_avatar_payload_excludes_internal_configuration(self):
payload = _public_avatar_payload(self.avatar) 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): def test_large_input_is_split_into_provider_safe_batches(self):
texts = [f"chunk-{index}" for index in range(14)] texts = [f"chunk-{index}" for index in range(14)]
batch_sizes = [] batch_sizes = []
requested_urls = []
def fake_urlopen(request, timeout): def fake_urlopen(request, timeout):
self.assertEqual(timeout, 30) self.assertEqual(timeout, 30)
requested_urls.append(request.full_url)
payload = json.loads(request.data.decode("utf-8")) payload = json.loads(request.data.decode("utf-8"))
batch_sizes.append(len(payload["input"])) batch_sizes.append(len(payload["input"]))
return FakeResponse({ return FakeResponse({
@@ -61,7 +63,7 @@ class RemoteEmbeddingTests(unittest.TestCase):
}) })
with patch.dict(os.environ, { 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_API_KEY": "test-key",
"EMBEDDING_MODEL": "text-embedding-v4", "EMBEDDING_MODEL": "text-embedding-v4",
"EMBEDDING_BATCH_SIZE": "10", "EMBEDDING_BATCH_SIZE": "10",
@@ -69,8 +71,18 @@ class RemoteEmbeddingTests(unittest.TestCase):
result = embeddings.embed(texts) result = embeddings.embed(texts)
self.assertEqual(batch_sizes, [10, 4]) 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)]) 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__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -2,7 +2,16 @@ from pathlib import Path
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import patch 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 from routers.knowledge import _doc_payload
from services.knowledge_vectorizer import knowledge_vectorizer
client = TestClient(app)
def test_doc_payload_reports_whether_the_persisted_file_exists(tmp_path: Path): def test_doc_payload_reports_whether_the_persisted_file_exists(tmp_path: Path):
@@ -21,3 +30,267 @@ def test_doc_payload_reports_whether_the_persisted_file_exists(tmp_path: Path):
assert _doc_payload(doc)["filePresent"] is False assert _doc_payload(doc)["filePresent"] is False
stored_file.write_text("knowledge", encoding="utf-8") stored_file.write_text("knowledge", encoding="utf-8")
assert _doc_payload(doc)["filePresent"] is True assert _doc_payload(doc)["filePresent"] is True
def test_upload_returns_before_background_vectorization(
tmp_path: Path,
authorization_context,
):
context = authorization_context
with (
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
headers=context["owner_headers"],
files={"file": ("knowledge.md", b"# Knowledge\n\nTest content", "text/markdown")},
)
payload = response.json()["data"]
assert payload["status"] == "parsing"
assert payload["vectorized"] is False
assert payload["chunkCount"] == 0
enqueue.assert_called_once_with(payload["id"])
db = SessionLocal()
try:
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
assert stored.status == "parsing"
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 0
db.delete(stored)
db.commit()
finally:
db.close()
def test_upload_rejects_oversize_file_before_queuing_indexing(
tmp_path: Path,
authorization_context,
):
context = authorization_context
with (
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
patch("routers.knowledge.MAX_UPLOAD_BYTES", 4),
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
headers=context["owner_headers"],
files={"file": ("oversize.md", b"12345", "text/markdown")},
)
payload = response.json()
assert payload["code"] == 400
assert payload["message"] == "文件不能超过 50MB"
enqueue.assert_not_called()
assert not list((tmp_path / context["avatar"].id).glob("*"))
def test_background_vectorizer_commits_ready_document_and_chunks_together(
tmp_path: Path,
authorization_context,
):
context = authorization_context
with (
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
patch("routers.knowledge.knowledge_vectorizer.enqueue"),
):
response = client.post(
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"] == "parsing"
with (
patch("services.knowledge_vectorizer.UPLOAD_DIR", str(tmp_path)),
patch("services.knowledge_vectorizer.embeddings.embed", return_value=[[1.0, 0.0]]),
):
knowledge_vectorizer.vectorize_document(payload["id"])
db = SessionLocal()
try:
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
assert stored.status == "ready"
assert stored.vectorized is True
assert stored.chunk_count == 1
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 1
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).delete()
db.delete(stored)
db.commit()
finally:
db.close()
def test_background_vectorizer_keeps_failure_reason_for_retry(
tmp_path: Path,
authorization_context,
):
context = authorization_context
with (
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
patch("routers.knowledge.knowledge_vectorizer.enqueue"),
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
headers=context["owner_headers"],
files={"file": ("knowledge.md", b"# Knowledge\n\nTest content", "text/markdown")},
)
payload = response.json()["data"]
with (
patch("services.knowledge_vectorizer.UPLOAD_DIR", str(tmp_path)),
patch("services.knowledge_vectorizer.embeddings.embed", side_effect=RuntimeError("provider unavailable")),
):
knowledge_vectorizer.vectorize_document(payload["id"])
db = SessionLocal()
try:
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
assert stored.status == "failed"
assert stored.error_message == "provider unavailable"
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 0
db.delete(stored)
db.commit()
finally:
db.close()
def test_retry_queues_a_failed_document_again(
tmp_path: Path,
authorization_context,
):
context = authorization_context
document_id = f"retry-doc-{context['suffix']}"
avatar_dir = tmp_path / context["avatar"].id
avatar_dir.mkdir()
(avatar_dir / "retry.md").write_text("retry content", encoding="utf-8")
db = SessionLocal()
try:
db.add(
KnowledgeDoc(
id=document_id,
avatar_id=context["avatar"].id,
filename="retry.md",
file_type="md",
file_url=f"/api/files/{context['avatar'].id}/retry.md",
status="failed",
error_message="provider unavailable",
)
)
db.commit()
finally:
db.close()
with (
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/knowledge/docs/{document_id}/retry",
headers=context["owner_headers"],
)
payload = response.json()["data"]
assert payload["status"] == "parsing"
assert payload["errorMessage"] == ""
enqueue.assert_called_once_with(document_id)
db = SessionLocal()
try:
db.query(KnowledgeDoc).filter(KnowledgeDoc.id == document_id).delete()
db.commit()
finally:
db.close()
def test_each_avatar_has_an_independent_document_and_qa_scope(authorization_context):
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 hasattr(auth, 'takeover_delay_seconds')
assert auth.takeover_enabled == False assert auth.takeover_enabled == False
assert auth.takeover_mode == 'immediate' assert auth.takeover_mode == 'immediate'
assert auth.takeover_delay_seconds == 30 assert auth.takeover_delay_seconds == 180
finally: finally:
db.close() db.close()
@@ -20,8 +20,9 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
): ):
import main import main
maintenance_scheduler = MagicMock()
scheduler = MagicMock() scheduler = MagicMock()
mock_scheduler_class.return_value = scheduler mock_scheduler_class.side_effect = [maintenance_scheduler, scheduler]
boxim = MagicMock() boxim = MagicMock()
mock_boxim_class.return_value = boxim mock_boxim_class.return_value = boxim
takeover = MagicMock() takeover = MagicMock()
@@ -45,7 +46,16 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
config = mock_boxim_class.call_args.args[0] config = mock_boxim_class.call_args.args[0]
assert config["HUIHUI_PLATFORM_BASE_URL"] == "https://open.example/api" assert config["HUIHUI_PLATFORM_BASE_URL"] == "https://open.example/api"
assert config["BOXIM_API_BASE_URL"] == "https://im.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 assert scheduler.add_job.call_count == 2
poll_call, process_call = scheduler.add_job.call_args_list 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() scheduler.start.assert_called_once_with()
main.takeover_scheduler = None main.takeover_scheduler = None
main.maintenance_scheduler = None
@patch("main.AsyncIOScheduler") @patch("main.AsyncIOScheduler")
@@ -73,6 +84,7 @@ def test_scheduler_failure_does_not_stop_the_api(mock_scheduler_class):
main.on_startup() main.on_startup()
assert main.takeover_scheduler is None assert main.takeover_scheduler is None
assert main.maintenance_scheduler is None
def test_shutdown_stops_only_the_scheduler(): def test_shutdown_stops_only_the_scheduler():
@@ -80,9 +92,14 @@ def test_shutdown_stops_only_the_scheduler():
scheduler = MagicMock() scheduler = MagicMock()
scheduler.running = True scheduler.running = True
maintenance_scheduler = MagicMock()
maintenance_scheduler.running = True
main.takeover_scheduler = scheduler main.takeover_scheduler = scheduler
main.maintenance_scheduler = maintenance_scheduler
main.on_shutdown() main.on_shutdown()
scheduler.shutdown.assert_called_once_with(wait=False) 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.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.""" """End-to-end service tests for BOXIM takeover timing and human priority."""
import asyncio
import json
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
from threading import Barrier from threading import Barrier
from unittest.mock import AsyncMock, patch from unittest.mock import AsyncMock, patch
@@ -9,9 +11,15 @@ from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker from sqlalchemy.orm import sessionmaker
from database import Base from database import Base
from models import Avatar, TakeoverCursor, TakeoverMessage, TakeoverReplyTask, User from models import Avatar, ChatAttachment, TakeoverCursor, TakeoverMessage, TakeoverReplyTask, User
from services.boxim_client import BoxIMError from services.boxim_client import BoxIMError
from services.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: class Clock:
@@ -57,6 +65,49 @@ class FakeBoxIM:
return {"id": 900 + len(self.sent), "localId": int(local_id)} 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 []
class ConcurrentMessagePollingBoxIM(ConcurrentPollingBoxIM):
async def fetch_private_messages(self, access_token, min_id="0"):
await super().fetch_private_messages(access_token, min_id)
owner_id = 100 if access_token == "prod-huihui-token" else 101
return [
{
"id": owner_id,
"localId": owner_id,
"sendId": owner_id + 100,
"recvId": owner_id,
"sendTime": 1_700_000_000_000,
"type": 0,
"content": "并发写入测试",
}
]
async def mark_private_messages_read(self, access_token, friend_id, message_id):
await asyncio.sleep(0.05)
self.read_receipts.append(
{"friendId": str(friend_id), "messageId": str(message_id)}
)
@pytest.fixture @pytest.fixture
def service_context(tmp_path): def service_context(tmp_path):
engine = create_engine( engine = create_engine(
@@ -77,7 +128,10 @@ def service_context(tmp_path):
owner_id=user.huihui_user_id, owner_id=user.huihui_user_id,
name="分身", name="分身",
status="active", status="active",
config={"authorizationPermissions": ["chat", "takeover"]}, config={
"authorizationPermissions": ["chat", "takeover"],
"takeoverReplyDelaySeconds": 3,
},
) )
db.add_all([user, avatar]) db.add_all([user, avatar])
db.commit() db.commit()
@@ -120,7 +174,6 @@ 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": "你好"} {"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.sent == []
assert boxim.read_receipts == [{"friendId": "200", "messageId": "11"}] assert boxim.read_receipts == [{"friendId": "200", "messageId": "11"}]
@@ -130,6 +183,7 @@ async def test_incoming_message_is_prepared_then_sent_at_three_seconds(service_c
assert boxim.sent == [] assert boxim.sent == []
clock.advance(1) clock.advance(1)
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 == [{"peerId": "200", "content": "你好\n很高兴见到你", "localId": boxim.sent[0]["localId"]}] assert boxim.sent == [{"peerId": "200", "content": "你好\n很高兴见到你", "localId": boxim.sent[0]["localId"]}]
@@ -142,6 +196,453 @@ async def test_incoming_message_is_prepared_then_sent_at_three_seconds(service_c
db.close() db.close()
@pytest.mark.asyncio
async def test_incoming_image_is_analyzed_and_used_in_takeover_reply(service_context):
session_factory, service, boxim, clock = service_context
await service.poll_and_process_messages()
boxim.messages.append(
{
"id": 111,
"localId": 111,
"sendId": 200,
"recvId": 100,
"sendTime": clock.millis(),
"type": 1,
"content": json.dumps(
{
"originUrl": "https://cdn.example/case.png",
"thumbUrl": "https://cdn.example/case-thumb.png",
}
),
}
)
await service.poll_and_process_messages()
db = session_factory()
try:
scheduled = db.query(TakeoverReplyTask).filter_by(trigger_message_id="111").one()
assert scheduled.status == "pending"
assert scheduled.prompt == "请看看这张图片。"
finally:
db.close()
clock.advance(3)
def analyze(db, avatar, content, **kwargs):
assert content == b"image-content"
attachment = ChatAttachment(
avatar_id=avatar.id,
uploader_kind=kwargs["uploader_kind"],
filename=kwargs["filename"],
mime_type="image/jpeg",
file_size=len(content),
status="ready",
category="medical_document",
summary="一张门诊病例",
extracted_text="主诉:咳嗽三天",
structured_data={"medical": {"chief_complaint": "咳嗽三天"}},
warning="请核对原始资料",
expires_at=clock.now() + timedelta(hours=24),
)
db.add(attachment)
db.commit()
db.refresh(attachment)
return attachment
downloaded = DownloadedBoxIMImage(
content=b"image-content",
filename="case.png",
mime_type="image/png",
source_url="https://cdn.example/case.png",
)
with (
patch("services.takeover_service.download_boxim_image", return_value=downloaded),
patch("routers.chat._analyze_image_bytes", side_effect=analyze) as analyzer,
patch("routers.chat._resolve_reply", return_value={"answer": "这份资料里写的是咳嗽三天。"}) as resolver,
):
await service.poll_and_process_messages()
analyzer.assert_called_once()
assert resolver.call_args.args[2] == "请看看这张图片。"
image_contexts = resolver.call_args.kwargs["image_contexts"]
assert image_contexts[0]["summary"] == "一张门诊病例"
assert image_contexts[0]["extractedText"] == "主诉:咳嗽三天"
assert [item["content"] for item in boxim.sent] == ["这份资料里写的是咳嗽三天。"]
db = session_factory()
try:
event = db.query(TakeoverMessage).filter_by(boxim_message_id="111").one()
task = db.query(TakeoverReplyTask).filter_by(trigger_message_id="111").one()
assert event.attachment_id
assert db.get(ChatAttachment, event.attachment_id).uploader_kind == "boxim"
assert task.status == "sent"
with patch(
"services.takeover_service.download_boxim_image",
side_effect=AssertionError("cached image must not be downloaded again"),
):
cached = service._takeover_image_attachment(db, db.get(Avatar, "avatar-1"), event)
assert cached.id == event.attachment_id
finally:
db.close()
@pytest.mark.asyncio
async def test_followup_text_recovers_recent_image_recorded_without_task(service_context):
session_factory, service, boxim, clock = service_context
await service.poll_and_process_messages()
image_message = {
"id": 113,
"localId": 113,
"sendId": 200,
"recvId": 100,
"sendTime": clock.millis(),
"type": 1,
"content": json.dumps(
{
"originUrl": "https://cdn.example/case.png",
"thumbUrl": "https://cdn.example/case-thumb.png",
}
),
}
db = session_factory()
try:
avatar = db.get(Avatar, "avatar-1")
service._record_message(db, avatar, "100", image_message, schedule_reply=False)
cursor = db.query(TakeoverCursor).one()
cursor.last_message_id = "113"
db.commit()
finally:
db.close()
clock.advance(60)
boxim.messages.extend(
[
image_message,
{
"id": 114,
"localId": 114,
"sendId": 200,
"recvId": 100,
"sendTime": clock.millis(),
"type": 0,
"content": "请帮我看看这张图",
},
]
)
await service.poll_messages()
db = session_factory()
try:
task = db.query(TakeoverReplyTask).filter_by(trigger_message_id="114").one()
assert task.source_message_ids == ["113", "114"]
assert task.prompt == "请看看这张图片。\n请帮我看看这张图"
finally:
db.close()
clock.advance(1)
boxim.messages.append(
{
"id": 115,
"localId": 115,
"sendId": 200,
"recvId": 100,
"sendTime": clock.millis(),
"type": 0,
"content": "图里写了什么",
}
)
await service.poll_messages()
db = session_factory()
try:
latest = db.query(TakeoverReplyTask).filter_by(trigger_message_id="115").one()
assert latest.source_message_ids == ["113", "114", "115"]
assert latest.source_message_ids.count("113") == 1
finally:
db.close()
@pytest.mark.asyncio
async def test_explicit_followup_reuses_handled_image_within_two_days(service_context):
session_factory, service, boxim, clock = service_context
await service.poll_and_process_messages()
image_message = {
"id": 116,
"localId": 116,
"sendId": 200,
"recvId": 100,
"sendTime": clock.millis(),
"type": 1,
"content": json.dumps({"originUrl": "https://cdn.example/handled-case.png"}),
}
boxim.messages.append(image_message)
await service.poll_messages()
db = session_factory()
try:
image_task = db.query(TakeoverReplyTask).filter_by(trigger_message_id="116").one()
image_task.status = "sent"
image_task.sent_at = clock.now()
db.commit()
finally:
db.close()
clock.advance(47 * 60 * 60)
boxim.messages.append(
{
"id": 117,
"localId": 117,
"sendId": 200,
"recvId": 100,
"sendTime": clock.millis(),
"type": 0,
"content": "重新看一下刚才那张病例图片",
}
)
await service.poll_messages()
db = session_factory()
try:
task = db.query(TakeoverReplyTask).filter_by(trigger_message_id="117").one()
assert task.source_message_ids == ["116", "117"]
assert task.prompt == "请看看这张图片。\n重新看一下刚才那张病例图片"
finally:
db.close()
@pytest.mark.asyncio
async def test_invalid_image_message_is_recorded_but_not_scheduled(service_context):
session_factory, service, boxim, clock = service_context
await service.poll_and_process_messages()
boxim.messages.append(
{
"id": 112,
"localId": 112,
"sendId": 200,
"recvId": 100,
"sendTime": clock.millis(),
"type": 1,
"content": json.dumps({"width": 100, "height": 100}),
}
)
await service.poll_and_process_messages()
db = session_factory()
try:
assert db.query(TakeoverMessage).filter_by(boxim_message_id="112").one()
assert db.query(TakeoverReplyTask).count() == 0
finally:
db.close()
@pytest.mark.asyncio
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 = ConcurrentMessagePollingBoxIM()
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
assert db.query(TakeoverMessage).count() == 2
assert len(boxim.read_receipts) == 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 @pytest.mark.asyncio
async def test_different_contacts_generate_without_blocking_each_other(service_context): async def test_different_contacts_generate_without_blocking_each_other(service_context):
session_factory, service, boxim, clock = service_context session_factory, service, boxim, clock = service_context
@@ -158,11 +659,11 @@ async def test_different_contacts_generate_without_blocking_each_other(service_c
both_generating.wait() both_generating.wait()
return {"answer": f"回复{prompt[-1]}"} 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) 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} == { assert {(item["peerId"], item["content"]) for item in boxim.sent} == {
("200", "回复甲"), ("200", "回复甲"),
("300", "回复乙"), ("300", "回复乙"),
@@ -237,6 +738,34 @@ async def test_owner_message_cancels_pending_reply(service_context):
db.close() 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 @pytest.mark.asyncio
async def test_quick_successive_messages_are_coalesced_into_one_reply(service_context): async def test_quick_successive_messages_are_coalesced_into_one_reply(service_context):
session_factory, service, boxim, clock = service_context session_factory, service, boxim, clock = service_context
@@ -244,19 +773,18 @@ async def test_quick_successive_messages_are_coalesced_into_one_reply(service_co
boxim.messages.append( boxim.messages.append(
{"id": 31, "localId": 5, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "第一句"} {"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) clock.advance(1)
boxim.messages.append( boxim.messages.append(
{"id": 32, "localId": 6, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "第二句"} {"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: with patch("routers.chat._resolve_reply", return_value={"answer": "合并回复"}) as resolver:
await service.poll_and_process_messages() await service.poll_and_process_messages()
assert resolver.call_args.args[2] == "第一句\n第二句" assert resolver.call_args.args[2] == "第一句\n第二句"
clock.advance(3)
await service.poll_and_process_messages()
assert [item["content"] for item in boxim.sent] == ["合并回复"] assert [item["content"] for item in boxim.sent] == ["合并回复"]
db = session_factory() db = session_factory()
@@ -291,5 +819,41 @@ async def test_connection_failure_disables_takeover_and_stops_retrying(service_c
boxim.exchange_access_token.assert_awaited_once_with("prod-huihui-token") 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(): def test_plain_text_reply_removes_markdown_and_empty_lines():
assert _plain_text_reply("## 建议\n\n**不能自行用药**\n`必要时就医`") == "建议\n不能自行用药\n必要时就医" 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_ACCESS_SECRET=<production-access-secret>
HUIHUI_CLIENT_CODE=<production-client-code> HUIHUI_CLIENT_CODE=<production-client-code>
BOXIM_TIMEOUT_SECONDS=20 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_BASE_URL=https://open.99hui.com/api/payment-v3
HUIHUI_PAYMENT_CALLBACK_BASE_URL=https://digital.99hui.com HUIHUI_PAYMENT_CALLBACK_BASE_URL=https://digital.99hui.com
HUIHUI_PAYMENT_CALLBACK_SECRET=<至少32位随机密钥> HUIHUI_PAYMENT_CALLBACK_SECRET=<至少32位随机密钥>
@@ -47,10 +49,26 @@ HUIHUI_PAYMENT_TIMEOUT_SECONDS=30
DATABASE_URL=sqlite:////data/avatar.db DATABASE_URL=sqlite:////data/avatar.db
UPLOAD_DIR=/data/uploads UPLOAD_DIR=/data/uploads
CHAT_MODEL_CONFIG_URL=http://<huihuisquare-api>/api/ai-models/runtime/digital-avatar 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` 必须挂载持久卷,数据库与知识库文件不可存放在容器临时层。 如生产 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 回调签名,不会发送到前端或直接出现在回调地址中。支付回调确认状态成功且金额与套餐价格完全一致后才增加积分,重复回调不会重复到账。 积分充值使用会会支付体系的 `payment-v3/payment/pay`,渠道值为 `WECHAT` / `ALIPAY`,端内支付场景为 `APP`,微信内 H5 使用 `JSAPI`。`HUIHUI_PAYMENT_CALLBACK_SECRET` 只用于为每笔订单生成 HMAC 回调签名,不会发送到前端或直接出现在回调地址中。支付回调确认状态成功且金额与套餐价格完全一致后才增加积分,重复回调不会重复到账。
## 3. 构建与发布 ## 3. 构建与发布
@@ -74,6 +92,7 @@ docker compose build --pull avatar-backend avatar-frontend
docker compose up -d avatar-backend avatar-frontend docker compose up -d avatar-backend avatar-frontend
docker compose ps docker compose ps
curl -fsS http://127.0.0.1:8099/api/health 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,并把延迟接管任务改为共享队列。 生产编排应把示例中的测试端口改为内网暴露,由统一 HTTPS 网关接入。后端暂时使用 SQLite,必须保持单实例写入;若扩展为多后端实例,应先迁移到 PostgreSQL,并把延迟接管任务改为共享队列。
@@ -98,11 +117,13 @@ location /api/ {
proxy_set_header X-Forwarded-Proto $scheme; proxy_set_header X-Forwarded-Proto $scheme;
proxy_buffering off; proxy_buffering off;
proxy_read_timeout 300s; proxy_read_timeout 300s;
client_max_body_size 20m; client_max_body_size 100m;
} }
``` ```
`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. 发布验收 ## 5. 发布验收
@@ -112,10 +133,13 @@ location /api/ {
4. A、B 两个会会用户分别进入时只能看到各自的数字分身与知识库,不会继承上一用户缓存。 4. A、B 两个会会用户分别进入时只能看到各自的数字分身与知识库,不会继承上一用户缓存。
5. 使用过期或伪造 token 时进入登录页并显示凭证失效,不得继续访问旧用户数据。 5. 使用过期或伪造 token 时进入登录页并显示凭证失效,不得继续访问旧用户数据。
6. 分身聊天 SSE 逐段输出正常,Markdown 正常渲染,知识库优先级和积分扣费正常。 6. 分身聊天 SSE 逐段输出正常,Markdown 正常渲染,知识库优先级和积分扣费正常。
7. 开启 BOXIM 主动接管后保持在线,收到消息、三秒回复、已读回执和主人发言暂停均正常。 7. 开启 BOXIM 主动接管后保持在线,默认三分钟回复、自定义等待时间、已读回执、分身防回环和主人发言暂停均正常。
8. 重建容器后数据库、头像、知识库文档仍存在,`/api/health` 返回成功。 8. 重建容器后数据库、头像、知识库文档仍存在,`/api/health` 返回成功。
9. `https://digital.99hui.com/api/health` 可访问,证书域名和有效期正确,HTTP 自动跳转 HTTPS。 9. `https://digital.99hui.com/api/health` 可访问,证书域名和有效期正确,HTTP 自动跳转 HTTPS。
10. 微信和支付宝各创建一笔最小套餐订单,未付款时积分不变;支付成功后回调到账一次,重复回调积分不重复增加。 10. 微信和支付宝各创建一笔最小套餐订单,未付款时积分不变;支付成功后回调到账一次,重复回调积分不重复增加。
11. 私聊和公开分享各上传 JPG、PNG、WebP 图片并完成追问;上传非图片、超过 8MB 或跨分身附件时必须拒绝。
12. 病例图片可以提取可见文字并标记待核对内容,医学影像不作确定诊断;视觉与 OCR 调用分别扣减积分。
13. 检查服务器上传目录不残留聊天原图,数据库过期图片识别记录在清理周期后删除,日志不出现 Base64 或病例正文。
## 6. 回滚 ## 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 光回答不作确定诊断,并显示人工复核提示。
+4
View File
@@ -23,6 +23,10 @@ http {
root /usr/share/nginx/html; root /usr/share/nginx/html;
index index.html; index index.html;
# Keep the application gateway aligned with the production edge gateway.
# Without this Nginx rejects ordinary PDF uploads with HTTP 413 before
# FastAPI can return its user-facing file-size validation message.
client_max_body_size 100m;
# SPA 兜底(hash 路由下深链接也可正常加载) # SPA 兜底(hash 路由下深链接也可正常加载)
location / { location / {
+66 -10
View File
@@ -191,19 +191,28 @@ export type AvatarPermission = 'friend' | 'chat' | 'publish' | 'browse' | 'inter
export interface AvatarPermissionSettings { export interface AvatarPermissionSettings {
avatarId: string avatarId: string
permissions: AvatarPermission[] permissions: AvatarPermission[]
takeoverReplyDelaySeconds: number
disabledAvatarIds?: string[]
} }
export const getAvatarPermissionSettings = (avatarId: string) => export const getAvatarPermissionSettings = (avatarId: string) =>
request.get<AvatarPermissionSettings>(`/avatar/${avatarId}/permission-settings`) request.get<AvatarPermissionSettings>(`/avatar/${avatarId}/permission-settings`)
export const updateAvatarPermissionSettings = (avatarId: string, permissions: AvatarPermission[]) => export const updateAvatarPermissionSettings = (
request.put<AvatarPermissionSettings>(`/avatar/${avatarId}/permission-settings`, { permissions }) avatarId: string,
permissions: AvatarPermission[],
takeoverReplyDelaySeconds: number
) => request.put<AvatarPermissionSettings>(`/avatar/${avatarId}/permission-settings`, {
permissions,
takeoverReplyDelaySeconds,
})
export interface TakeoverStatus { export interface TakeoverStatus {
enabled: boolean enabled: boolean
status: 'disabled' | 'connecting' | 'ready' | 'needs_login' | 'error' status: 'disabled' | 'connecting' | 'ready' | 'needs_login' | 'error'
message: string message: string
pendingCount: number pendingCount: number
takeoverReplyDelaySeconds: number
lastPolledAt: string | null lastPolledAt: string | null
} }
@@ -296,6 +305,7 @@ export interface KnowledgeDoc {
vectorized?: boolean vectorized?: boolean
embeddingModel?: string embeddingModel?: string
chunkCount?: number chunkCount?: number
errorMessage?: string
createdAt: string createdAt: string
} }
@@ -326,7 +336,8 @@ export const uploadKnowledgeDoc = (avatarId: string, file: File) => {
const form = new FormData() const form = new FormData()
form.append('file', file) form.append('file', file)
return request.post<KnowledgeDoc>(`/avatar/${avatarId}/knowledge/docs`, form, { return request.post<KnowledgeDoc>(`/avatar/${avatarId}/knowledge/docs`, form, {
headers: { 'Content-Type': 'multipart/form-data' } headers: { 'Content-Type': 'multipart/form-data' },
timeout: 120000
}) })
} }
@@ -334,6 +345,9 @@ export const uploadKnowledgeDoc = (avatarId: string, file: File) => {
export const deleteKnowledgeDoc = (avatarId: string, docId: string) => export const deleteKnowledgeDoc = (avatarId: string, docId: string) =>
request.delete(`/avatar/${avatarId}/knowledge/docs/${docId}`) request.delete(`/avatar/${avatarId}/knowledge/docs/${docId}`)
export const retryKnowledgeDoc = (avatarId: string, docId: string) =>
request.post<KnowledgeDoc>(`/avatar/${avatarId}/knowledge/docs/${docId}/retry`)
// 标准问答对列表 // 标准问答对列表
export const getQAPairs = (avatarId: string) => export const getQAPairs = (avatarId: string) =>
request.get<QAPair[]>(`/avatar/${avatarId}/knowledge/qa`) request.get<QAPair[]>(`/avatar/${avatarId}/knowledge/qa`)
@@ -365,15 +379,35 @@ export const searchKnowledge = (avatarId: string, q: string, topK = 5) =>
export interface ChatMessage { export interface ChatMessage {
role: 'user' | 'assistant' role: 'user' | 'assistant'
content: string content: string
attachmentIds?: string[]
} }
export interface ChatResponse { export interface ChatResponse {
answer: string answer: string
source: 'qa' | 'knowledge' | 'qwen' source: 'qa' | 'knowledge' | 'vision' | 'qwen'
references?: Array<{ docId?: string; filename?: string; fileType?: string; snippet?: string; score?: number }> 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) request.post<ChatResponse>(`/avatar/${avatarId}/chat`, payload)
export interface PublicAvatar { export interface PublicAvatar {
@@ -392,19 +426,41 @@ export const createAvatarShareLink = (avatarId: string) =>
export const getPublicAvatar = (shareToken: string) => export const getPublicAvatar = (shareToken: string) =>
request.get<PublicAvatar>(`/public/avatar/${shareToken}`) 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) 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 = { type ChatStreamHandlers = {
onMeta: (meta: Pick<ChatResponse, 'source' | 'references'>) => void onMeta: (meta: Pick<ChatResponse, 'source' | 'references'>) => void
onDelta: (content: string) => 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' } const headers: Record<string, string> = { 'Content-Type': 'application/json', Accept: 'text/event-stream' }
if (_authToken) headers.Authorization = `Bearer ${_authToken}` if (_authToken) headers.Authorization = `Bearer ${_authToken}`
const response = await fetch(`${resolveBaseURL()}${path}`, { method: 'POST', headers, body: JSON.stringify(payload) }) 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 reader = response.body.getReader()
const decoder = new TextDecoder() const decoder = new TextDecoder()
@@ -427,10 +483,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) 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) streamChat(`/public/avatar/${shareToken}/chat/stream`, payload, handlers)
// ==================== 会会用户资料 API ==================== // ==================== 会会用户资料 API ====================
+13 -4
View File
@@ -10,6 +10,7 @@ import {
type SmsLoginResult, type SmsLoginResult,
type UserProfile type UserProfile
} from '@/api' } from '@/api'
import { clearHuihuiEmbeddedMode, markHuihuiEmbeddedMode } from '@/utils/embed-mode'
const TOKEN_KEY = 'hh_app_token' const TOKEN_KEY = 'hh_app_token'
const USER_KEY = 'hh_app_user' const USER_KEY = 'hh_app_user'
@@ -70,16 +71,23 @@ export const useUserStore = defineStore('smsuser', () => {
// 短信登录 // 短信登录
const login = async (phone: string, code: string) => { 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) => { 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) => const loginByToken = async (huihuiToken: string) => {
acceptLogin(await loginByHuihuiToken(huihuiToken)) const result = await loginByHuihuiToken(huihuiToken)
markHuihuiEmbeddedMode()
return acceptLogin(result)
}
// 退出 // 退出
const logout = async () => { const logout = async () => {
@@ -88,6 +96,7 @@ export const useUserStore = defineStore('smsuser', () => {
} catch { } catch {
/* 忽略网络错误,本地清除即可 */ /* 忽略网络错误,本地清除即可 */
} }
clearHuihuiEmbeddedMode()
clearSession() 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> <template>
<div class="authorization-page"> <div class="authorization-page" :class="{ embedded: isEmbedded }">
<header class="page-header"> <header v-if="!isEmbedded" class="page-header">
<button class="back-button" type="button" aria-label="返回数字分身管理" @click="goBack"> <button class="back-button" type="button" aria-label="返回数字分身管理" @click="goBack">
<svg viewBox="0 0 24 24" aria-hidden="true"> <svg viewBox="0 0 24 24" aria-hidden="true">
<path d="m15 18-6-6 6-6" /> <path d="m15 18-6-6 6-6" />
@@ -64,7 +64,7 @@
<span class="permission-copy"> <span class="permission-copy">
<strong>{{ item.title }}</strong> <strong>{{ item.title }}</strong>
<small> <small>
{{ item.description }} {{ item.key === 'takeover' ? takeoverDescription : item.description }}
<span <span
v-if="item.key === 'takeover' && takeoverConnectionLabel" v-if="item.key === 'takeover' && takeoverConnectionLabel"
class="connection-state" class="connection-state"
@@ -79,6 +79,33 @@
</button> </button>
</section> </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> <p v-if="errorMessage" class="error-message" role="alert">{{ errorMessage }}</p>
</template> </template>
@@ -95,6 +122,9 @@
</main> </main>
<footer v-if="activeAvatarId" class="save-area"> <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()"> <button class="save-button" type="button" :disabled="loading || saving" @click="saveSettings()">
<span v-if="saving" class="saving-spinner" aria-hidden="true"></span> <span v-if="saving" class="saving-spinner" aria-hidden="true"></span>
{{ saving ? '保存中...' : '保存授权设置' }} {{ saving ? '保存中...' : '保存授权设置' }}
@@ -119,6 +149,7 @@ import {
} from '@/api' } from '@/api'
import { useAvatarStore } from '@/store/avatar' import { useAvatarStore } from '@/store/avatar'
import { pickScopedAvatarId } from '@/utils/avatar-page-data.js' import { pickScopedAvatarId } from '@/utils/avatar-page-data.js'
import { isHuihuiEmbeddedMode } from '@/utils/embed-mode'
type PermissionState = Record<AvatarPermission, boolean> type PermissionState = Record<AvatarPermission, boolean>
@@ -126,6 +157,7 @@ const router = useRouter()
const route = useRoute() const route = useRoute()
const avatarStore = useAvatarStore() const avatarStore = useAvatarStore()
const avatarId = computed(() => pickScopedAvatarId(route.params.avatarId, avatarStore.currentAvatarId, avatarStore.avatars)) const avatarId = computed(() => pickScopedAvatarId(route.params.avatarId, avatarStore.currentAvatarId, avatarStore.avatars))
const isEmbedded = isHuihuiEmbeddedMode()
const permissionItems: Array<{ const permissionItems: Array<{
key: AvatarPermission key: AvatarPermission
@@ -166,7 +198,7 @@ const permissionItems: Array<{
{ {
key: 'takeover', key: 'takeover',
title: '分身主动接管聊天回复', title: '分身主动接管聊天回复',
description: '收到私聊消息 3 秒后回复,主人发言时暂停', description: '收到私聊消息后按设定时间回复,主人发言时暂停',
tone: 'cyan', tone: 'cyan',
}, },
] ]
@@ -185,6 +217,8 @@ const saving = ref(false)
const errorMessage = ref('') const errorMessage = ref('')
const toastMessage = ref('') const toastMessage = ref('')
const takeoverStatus = ref<TakeoverStatus | null>(null) const takeoverStatus = ref<TakeoverStatus | null>(null)
const takeoverDelayValue = ref(3)
const takeoverDelayUnit = ref<'seconds' | 'minutes'>('minutes')
let toastTimer: number | undefined let toastTimer: number | undefined
let takeoverStatusTimer: number | undefined let takeoverStatusTimer: number | undefined
@@ -205,6 +239,38 @@ const takeoverConnectionTone = computed(() => {
return 'connecting' 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 setPermissions = (permissions: AvatarPermission[]) => {
const enabled = new Set(permissions) const enabled = new Set(permissions)
for (const item of permissionItems) permissionState[item.key] = enabled.has(item.key) for (const item of permissionItems) permissionState[item.key] = enabled.has(item.key)
@@ -260,6 +326,7 @@ const loadSettings = async () => {
try { try {
const settings = await getAvatarPermissionSettings(activeAvatarId.value) const settings = await getAvatarPermissionSettings(activeAvatarId.value)
setPermissions(settings.permissions || []) setPermissions(settings.permissions || [])
applyTakeoverDelay(settings.takeoverReplyDelaySeconds || 180)
await loadTakeoverStatus() await loadTakeoverStatus()
scheduleTakeoverStatusRefresh() scheduleTakeoverStatusRefresh()
} catch (error: any) { } catch (error: any) {
@@ -286,8 +353,18 @@ const saveSettings = async (takeoverToggle = false): Promise<boolean> => {
saving.value = true saving.value = true
errorMessage.value = '' errorMessage.value = ''
try { 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 || []) setPermissions(settings.permissions || [])
applyTakeoverDelay(settings.takeoverReplyDelaySeconds || 180)
await loadTakeoverStatus() await loadTakeoverStatus()
scheduleTakeoverStatusRefresh() scheduleTakeoverStatusRefresh()
if (takeoverToggle) { if (takeoverToggle) {
@@ -394,6 +471,10 @@ svg {
padding: 0 20px; padding: 0 20px;
} }
.authorization-page.embedded .page-content {
padding-top: 16px;
}
.permission-intro { .permission-intro {
min-height: 96px; min-height: 96px;
padding: 15px 16px 14px; padding: 15px 16px 14px;
@@ -459,6 +540,75 @@ svg {
min-height: 76px; 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 { .permission-icon {
width: 34px; width: 34px;
height: 34px; height: 34px;
@@ -599,12 +749,15 @@ svg {
bottom: 0; bottom: 0;
width: min(100%, 390px); width: min(100%, 390px);
padding: 12px 20px calc(20px + env(safe-area-inset-bottom)); 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%); background: linear-gradient(to bottom, rgba(250, 250, 250, 0), #fafafa 20%, #fafafa 100%);
transform: translateX(-50%); transform: translateX(-50%);
} }
.save-button { .save-button {
width: 100%; min-width: 0;
flex: 1;
height: 48px; height: 48px;
display: flex; display: flex;
align-items: center; align-items: center;
@@ -620,6 +773,20 @@ svg {
cursor: pointer; 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 { .save-button:disabled {
opacity: .68; opacity: .68;
} }
+179 -18
View File
@@ -1,5 +1,5 @@
<template> <template>
<div class="chat-page"> <div class="chat-page" :class="{ 'has-pending-images': pendingImages.length }">
<header class="chat-header"> <header class="chat-header">
<button v-if="!isPublic" class="back-btn" @click="router.back()">‹</button> <button v-if="!isPublic" class="back-btn" @click="router.back()">‹</button>
<div class="avatar-heading"> <div class="avatar-heading">
@@ -31,6 +31,12 @@
<span v-else>{{ avatar?.emoji || '🤖' }}</span> <span v-else>{{ avatar?.emoji || '🤖' }}</span>
</div> </div>
<div class="message-column"> <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 }"> <div class="message-bubble" :class="{ streaming: sending && message.role === 'assistant' && index === messages.length - 1 }">
<template v-if="message.role === 'assistant'"> <template v-if="message.role === 'assistant'">
<span <span
@@ -70,24 +76,66 @@
</main> </main>
<form class="composer" @submit.prevent="sendMessage(inputText)"> <form class="composer" @submit.prevent="sendMessage(inputText)">
<textarea v-model="inputText" rows="1" :disabled="sending" placeholder="输入你想聊的内容…" @keydown.enter.exact.prevent="sendMessage(inputText)"></textarea> <div v-if="pendingImages.length" class="pending-images">
<button class="send-btn" type="submit" :disabled="sending || !inputText.trim()">发送</button> <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> </form>
</div> </div>
</template> </template>
<script setup lang="ts"> <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 { 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 { useAvatarStore } from '@/store/avatar'
import { useUserStore } from '@/store/user' import { useUserStore } from '@/store/user'
import { renderChatMarkdownCharacters } from '@/utils/chat-markdown.js' import { renderChatMarkdownCharacters } from '@/utils/chat-markdown.js'
type DisplayMessage = ChatMessage & { type DisplayMessage = ChatMessage & {
source?: 'qa' | 'knowledge' | 'qwen' | 'public' source?: 'qa' | 'knowledge' | 'vision' | 'qwen' | 'public'
references?: Array<{ filename?: string }> references?: Array<{ filename?: string }>
characters?: 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() const route = useRoute()
@@ -103,8 +151,11 @@ const inputText = ref('')
const sending = ref(false) const sending = ref(false)
const thinking = ref(false) const thinking = ref(false)
const errorMessage = ref('') const errorMessage = ref('')
const lastQuestion = ref('') const lastRequest = ref<{ question: string; attachments: MessageAttachment[] } | null>(null)
const messageList = ref<HTMLElement | 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 let scrollFrame: number | null = null
const userAvatarUrl = computed(() => userStore.user?.avatarUrl || store.userProfile?.avatarUrl || '') const userAvatarUrl = computed(() => userStore.user?.avatarUrl || store.userProfile?.avatarUrl || '')
@@ -115,14 +166,25 @@ const avatarStatus = computed(() => {
if (status === 'training') return { tone: 'training', label: '知识训练中' } if (status === 'training') return { tone: 'training', label: '知识训练中' }
return { tone: 'active', 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> = { const sourceLabels: Record<NonNullable<DisplayMessage['source']>, string> = {
qa: '标准问答对', qa: '标准问答对',
knowledge: '参考文件知识库', knowledge: '参考文件知识库',
vision: '图片理解',
qwen: '智能回答', qwen: '智能回答',
public: '' public: ''
} }
const sourceLabel = (source?: DisplayMessage['source']) => source ? sourceLabels[source] : '' 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 () => { const scrollToBottom = async () => {
await nextTick() await nextTick()
@@ -214,20 +276,90 @@ const loadAvatar = async () => {
document.title = avatar.value?.displayName || avatar.value?.name || '会会数字分身' document.title = avatar.value?.displayName || avatar.value?.name || '会会数字分身'
} }
const sendMessage = async (value: string) => { const removePendingImage = (localId: string) => {
const question = value.trim() const target = pendingImages.value.find((image) => image.localId === localId)
if (!question || sending.value) return if (target) {
lastQuestion.value = question 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 = '' inputText.value = ''
errorMessage.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 sending.value = true
thinking.value = true thinking.value = true
await scrollToBottom() await scrollToBottom()
try { try {
const payload = { const payload = {
message: question, message: question,
history: messages.value.slice(-10).map(({ role, content }) => ({ role, content })) attachmentIds: selectedAttachments.map((attachment) => attachment.id),
history
} }
const streamed = createStreamReply() const streamed = createStreamReply()
const handlers = { const handlers = {
@@ -256,13 +388,17 @@ const sendMessage = async (value: string) => {
} }
const retryLast = () => { const retryLast = () => {
if (!lastQuestion.value || sending.value) return if (!lastRequest.value || sending.value) return
const last = messages.value[messages.value.length - 1] while (messages.value[messages.value.length - 1]?.role === 'assistant') messages.value.pop()
if (last?.role === 'user') messages.value.pop() if (messages.value[messages.value.length - 1]?.role === 'user') messages.value.pop()
sendMessage(lastQuestion.value) void sendMessage(lastRequest.value.question, lastRequest.value.attachments)
} }
onMounted(loadAvatar) onMounted(loadAvatar)
onBeforeUnmount(() => {
previewUrls.forEach((url) => URL.revokeObjectURL(url))
previewUrls.clear()
})
</script> </script>
<style scoped> <style scoped>
@@ -275,6 +411,7 @@ onMounted(loadAvatar)
.avatar-heading h1 { margin: 0; font-size: 17px; } .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; } .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; } .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-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-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; } .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-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-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-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 { 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; } .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; } .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 type-cursor { 50% { opacity: 0; } }
@keyframes character-in { from { opacity: 0; transform: translateY(3px); } to { opacity: 1; transform: translateY(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; } .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; } .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; } .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> </style>
+56 -15
View File
@@ -1,12 +1,10 @@
<template> <template>
<div class="edit-avatar-page"> <div class="edit-avatar-page">
<!-- 顶部导航 --> <!-- 顶部导航 -->
<header class="page-header"> <header v-if="!isEmbedded" class="page-header">
<button class="back-btn" @click="goBack">‹</button> <button class="back-btn" @click="goBack">‹</button>
<h1 class="page-title">分身微调</h1> <h1 class="page-title">分身微调</h1>
<button class="save-btn" :disabled="loading || saving || uploadingPhoto" @click="saveChanges"> <span class="header-spacer" aria-hidden="true"></span>
{{ saving ? '保存中...' : '保存' }}
</button>
</header> </header>
<div v-if="loading" class="status-banner">加载中...</div> <div v-if="loading" class="status-banner">加载中...</div>
@@ -159,6 +157,13 @@
{{ deleting ? '删除中...' : '删除数字分身' }} {{ deleting ? '删除中...' : '删除数字分身' }}
</button> </button>
</section> </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> </div>
</template> </template>
@@ -168,11 +173,13 @@ import { useRoute, useRouter } from 'vue-router'
import { deleteAvatar as apiDeleteAvatar, getAvatarDetail, updateAvatar, uploadAvatarPhoto } from '@/api' import { deleteAvatar as apiDeleteAvatar, getAvatarDetail, updateAvatar, uploadAvatarPhoto } from '@/api'
import { useAvatarStore } from '@/store/avatar' import { useAvatarStore } from '@/store/avatar'
import { buildAvatarUpdatePayload, normalizeAvatarEditForm } from '@/utils/avatar-page-data.js' import { buildAvatarUpdatePayload, normalizeAvatarEditForm } from '@/utils/avatar-page-data.js'
import { isHuihuiEmbeddedMode } from '@/utils/embed-mode'
const router = useRouter() const router = useRouter()
const route = useRoute() const route = useRoute()
const avatarStore = useAvatarStore() const avatarStore = useAvatarStore()
const avatarId = route.params.id as string const avatarId = route.params.id as string
const isEmbedded = isHuihuiEmbeddedMode()
// 表单数据 // 表单数据
const formData = reactive({ const formData = reactive({
@@ -288,7 +295,7 @@ onMounted(async () => {
.edit-avatar-page { .edit-avatar-page {
min-height: 100vh; min-height: 100vh;
background: #F8F9FA; background: #F8F9FA;
padding-bottom: 40px; padding-bottom: calc(104px + env(safe-area-inset-bottom));
} }
/* 顶部导航 */ /* 顶部导航 */
@@ -331,16 +338,7 @@ onMounted(async () => {
color: #B91C1C; color: #B91C1C;
} }
.save-btn { .header-spacer { width: 40px; }
background: #F97316;
color: white;
border: none;
padding: 8px 20px;
border-radius: 8px;
font-size: 14px;
font-weight: 600;
cursor: pointer;
}
/* 头像上传 */ /* 头像上传 */
.photo-section { .photo-section {
@@ -580,4 +578,47 @@ onMounted(async () => {
background: #EF4444; background: #EF4444;
color: white; 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> </style>
+35 -6
View File
@@ -1,15 +1,11 @@
<template> <template>
<div class="avatar-manage-page"> <div class="avatar-manage-page">
<!-- 顶部导航 --> <!-- 顶部导航 -->
<header class="page-header"> <header v-if="!isEmbedded" class="page-header">
<div class="header-left"> <div class="header-left">
<button class="back-btn" @click="goBack">‹</button> <button class="back-btn" @click="goBack">‹</button>
<h1 class="page-title">数字分身管理</h1> <h1 class="page-title">数字分身管理</h1>
</div> </div>
<!-- 右上角创建入口 -->
<div class="header-right">
<button class="icon-btn" @click="goCreate" title="创建数字分身">➕</button>
</div>
</header> </header>
<!-- 用户资料头(会会登录账号的头像 / 昵称) --> <!-- 用户资料头(会会登录账号的头像 / 昵称) -->
@@ -39,9 +35,14 @@
<!-- 数字分身列表(只放分身相关) --> <!-- 数字分身列表(只放分身相关) -->
<section class="avatar-list-section"> <section class="avatar-list-section">
<div class="section-head"> <div class="section-head">
<div class="section-heading-copy">
<h3 class="section-title">我的数字分身</h3> <h3 class="section-title">我的数字分身</h3>
<span class="count-badge">{{ avatars.length }}</span> <span class="count-badge">{{ avatars.length }}</span>
</div> </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"> <div v-if="avatars.length" class="avatar-list">
<div class="avatar-card" v-for="a in avatars" :key="a.id"> <div class="avatar-card" v-for="a in avatars" :key="a.id">
@@ -87,10 +88,12 @@ import { useRouter } from 'vue-router'
import { useAvatarStore } from '@/store/avatar' import { useAvatarStore } from '@/store/avatar'
import { useUserStore } from '@/store/user' import { useUserStore } from '@/store/user'
import { createAvatarShareLink } from '@/api' import { createAvatarShareLink } from '@/api'
import { isHuihuiEmbeddedMode } from '@/utils/embed-mode'
const router = useRouter() const router = useRouter()
const avatarStore = useAvatarStore() const avatarStore = useAvatarStore()
const userStore = useUserStore() const userStore = useUserStore()
const isEmbedded = isHuihuiEmbeddedMode()
// 临时产品开关:余额卡片代码保留,后续改为 true 即可恢复展示。 // 临时产品开关:余额卡片代码保留,后续改为 true 即可恢复展示。
const SHOW_POINTS_BALANCE_CARD = false const SHOW_POINTS_BALANCE_CARD = false
@@ -354,10 +357,36 @@ onMounted(() => {
.section-head { .section-head {
display: flex; display: flex;
align-items: center; align-items: center;
gap: 8px; justify-content: space-between;
gap: 12px;
margin: 8px 0 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 { .section-title {
font-size: 16px; font-size: 16px;
font-weight: 600; font-weight: 600;
@@ -1,7 +1,7 @@
<template> <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"> <div class="header-left">
<button class="back-btn" @click="goBack">‹</button> <button class="back-btn" @click="goBack">‹</button>
<h1 class="page-title">知识库管理</h1> <h1 class="page-title">知识库管理</h1>
@@ -23,11 +23,11 @@
<div class="upload-section"> <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-zone" :class="{ 'drag-over': dragOver }" @click="triggerFile" @dragover.prevent="dragOver = true" @dragleave.prevent="dragOver = false" @drop.prevent="onDrop">
<div class="upload-icon">📥</div> <div class="upload-icon">📥</div>
<p class="upload-title">拖拽文件到此处,或<span class="upload-link">点击上传</span></p> <p class="upload-title"><span class="upload-link">点击上传</span></p>
<p class="upload-hint">支持 MD / TXT / PDF / DOC / DOCX / XLSX,上传后自动向量化</p> <p class="upload-hint">支持 MD / TXT / PDF / DOC / DOCX / XLSX,上传后自动向量化</p>
<input ref="fileInput" type="file" accept=".md,.txt,.pdf,.doc,.docx,.xlsx" class="hidden-input" @change="onFileChange" /> <input ref="fileInput" type="file" accept=".md,.txt,.pdf,.doc,.docx,.xlsx" class="hidden-input" @change="onFileChange" />
</div> </div>
<p v-if="uploading" class="uploading-text">上传并向量化中…</p> <p v-if="uploading" class="uploading-text">文件上传中…</p>
<p v-if="uploadError" class="error-text">{{ uploadError }}</p> <p v-if="uploadError" class="error-text">{{ uploadError }}</p>
</div> </div>
@@ -37,12 +37,15 @@
<div class="card-content"> <div class="card-content">
<div class="card-title-row"> <div class="card-title-row">
<strong>{{ doc.filename }}</strong> <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> </div>
<p class="card-meta">{{ doc.fileType.toUpperCase() }} · {{ formatSize(doc.fileSize) }} · {{ formatDate(doc.createdAt) }}</p> <p class="card-meta">{{ doc.fileType.toUpperCase() }} · {{ formatSize(doc.fileSize) }} · {{ formatDate(doc.createdAt) }}</p>
<p class="card-detail">{{ doc.filePresent === false ? '原文件不可用,请删除后重新上传' : (doc.vectorized ? `已切分 ${doc.chunkCount || 0} 段,可用于对话` : '正在解析并建立知识索引') }}</p> <p class="card-detail">{{ documentState(doc).detail }}</p>
</div> </div>
<div class="card-actions">
<button v-if="documentState(doc).tone === 'failed'" class="card-retry" @click="retryDoc(doc.id)">重新索引</button>
<button class="card-delete" @click="removeDoc(doc.id)">删除</button> <button class="card-delete" @click="removeDoc(doc.id)">删除</button>
</div>
</article> </article>
</div> </div>
<div v-else class="card-empty">📂 暂无文档,先上传一个知识文件</div> <div v-else class="card-empty">📂 暂无文档,先上传一个知识文件</div>
@@ -78,14 +81,16 @@
</template> </template>
<script setup lang="ts"> <script setup lang="ts">
import { ref, onMounted, computed } from 'vue' import { ref, onMounted, onUnmounted, computed } from 'vue'
import { useRoute, useRouter } from 'vue-router' import { useRoute, useRouter } from 'vue-router'
import { useAvatarStore } from '@/store/avatar' import { useAvatarStore } from '@/store/avatar'
import { pickScopedAvatarId, unwrapListData } from '@/utils/avatar-page-data.js' import { pickScopedAvatarId, unwrapListData } from '@/utils/avatar-page-data.js'
import { isHuihuiEmbeddedMode } from '@/utils/embed-mode'
import { import {
getKnowledgeDocs, getKnowledgeDocs,
uploadKnowledgeDoc, uploadKnowledgeDoc,
deleteKnowledgeDoc, deleteKnowledgeDoc,
retryKnowledgeDoc,
getQAPairs, getQAPairs,
deleteQAPair, deleteQAPair,
searchKnowledge, searchKnowledge,
@@ -95,6 +100,7 @@ import {
const router = useRouter() const router = useRouter()
const route = useRoute() const route = useRoute()
const store = useAvatarStore() const store = useAvatarStore()
const isEmbedded = isHuihuiEmbeddedMode()
const avatarId = computed(() => pickScopedAvatarId(route.params.avatarId, store.currentAvatarId, store.avatars)) const avatarId = computed(() => pickScopedAvatarId(route.params.avatarId, store.currentAvatarId, store.avatars))
const activeTab = ref<'docs' | 'qa'>('docs') const activeTab = ref<'docs' | 'qa'>('docs')
@@ -105,17 +111,51 @@ const uploading = ref(false)
const uploadError = ref('') const uploadError = ref('')
const dragOver = ref(false) const dragOver = ref(false)
const fileInput = ref<HTMLInputElement | null>(null) const fileInput = ref<HTMLInputElement | null>(null)
let documentPollingTimer: ReturnType<typeof setInterval> | undefined
const query = ref('') const query = ref('')
const searching = ref(false) const searching = ref(false)
const searched = ref(false) const searched = ref(false)
const searchResults = ref<any[]>([]) const searchResults = ref<any[]>([])
const 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: doc.errorMessage || '未能建立知识索引,请重新索引或重新上传' }
}
const hasPendingDocuments = () => docs.value.some((doc) =>
['uploaded', 'parsing'].includes(String(doc.status || '').toLowerCase())
)
const stopDocumentPolling = () => {
if (documentPollingTimer) {
clearInterval(documentPollingTimer)
documentPollingTimer = undefined
}
}
const startDocumentPolling = () => {
if (documentPollingTimer || !hasPendingDocuments()) return
documentPollingTimer = setInterval(async () => {
await loadDocs()
if (!hasPendingDocuments()) stopDocumentPolling()
}, 2000)
}
const loadDocs = async () => { const loadDocs = async () => {
if (!avatarId.value) return if (!avatarId.value) return
try { try {
const res: any = await getKnowledgeDocs(avatarId.value) const res: any = await getKnowledgeDocs(avatarId.value)
docs.value = unwrapListData(res) docs.value = unwrapListData(res)
startDocumentPolling()
} catch (e) { } catch (e) {
console.error(e) console.error(e)
} }
@@ -167,6 +207,17 @@ const doUpload = async (file: File) => {
} }
} }
const retryDoc = async (id: string) => {
if (!avatarId.value) return
uploadError.value = ''
try {
await retryKnowledgeDoc(avatarId.value, id)
await loadDocs()
} catch (e: any) {
uploadError.value = e?.message || '重新索引失败'
}
}
const removeDoc = async (id: string) => { const removeDoc = async (id: string) => {
if (!avatarId.value) return if (!avatarId.value) return
await deleteKnowledgeDoc(avatarId.value, id) await deleteKnowledgeDoc(avatarId.value, id)
@@ -242,6 +293,8 @@ onMounted(async () => {
if (avatarId.value) store.currentAvatarId = avatarId.value if (avatarId.value) store.currentAvatarId = avatarId.value
await Promise.all([loadDocs(), loadQA()]) await Promise.all([loadDocs(), loadQA()])
}) })
onUnmounted(stopDocumentPolling)
</script> </script>
<style scoped> <style scoped>
@@ -283,8 +336,12 @@ onMounted(async () => {
.card-title-row strong { min-width: 0; flex: 1; overflow: hidden; color: #27201C; font-size: 14px; text-overflow: ellipsis; white-space: nowrap; } .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 { 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.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-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-actions { flex: 0 0 auto; display: flex; flex-direction: column; align-items: stretch; gap: 6px; }
.card-delete, .card-retry { align-self: center; border: 0; border-radius: 8px; padding: 7px 9px; font-size: 12px; cursor: pointer; white-space: nowrap; }
.card-delete { color: #EF4444; background: #FEF2F2; }
.card-retry { color: #C15F18; background: #FFF3E6; }
.card-empty { padding: 42px 16px; border: 1px dashed #F1D9C3; border-radius: 16px; color: #9398AE; background: #fff; font-size: 14px; text-align: center; } .card-empty { padding: 42px 16px; border: 1px dashed #F1D9C3; border-radius: 16px; color: #9398AE; background: #fff; font-size: 14px; text-align: center; }
.qa-card { align-items: stretch; text-align: left; }.qa-card.qa-disabled { opacity: .58; } .qa-card { align-items: stretch; text-align: left; }.qa-card.qa-disabled { opacity: .58; }
.qa-card .card-content, .qa-card .card-content,
+22 -3
View File
@@ -10,10 +10,12 @@
<section class="form-section"> <section class="form-section">
<label class="field-label">问题</label> <label class="field-label">问题</label>
<textarea <textarea
ref="questionInput"
v-model="form.question" v-model="form.question"
class="field-input" class="field-input question-input"
rows="3" rows="1"
placeholder="例如:你们的退款政策是什么?" placeholder="例如:你们的退款政策是什么?"
@input="resizeQuestion"
></textarea> ></textarea>
<label class="field-label">标准答案</label> <label class="field-label">标准答案</label>
@@ -46,7 +48,7 @@
</template> </template>
<script setup lang="ts"> <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 { useRouter, useRoute } from 'vue-router'
import { useAvatarStore } from '@/store/avatar' import { useAvatarStore } from '@/store/avatar'
import { pickScopedAvatarId, unwrapListData } from '@/utils/avatar-page-data.js' import { pickScopedAvatarId, unwrapListData } from '@/utils/avatar-page-data.js'
@@ -63,6 +65,14 @@ const isEdit = computed(() => !!qaId.value)
const form = reactive({ question: '', answer: '', enabled: true }) const form = reactive({ question: '', answer: '', enabled: true })
const saving = ref(false) const saving = ref(false)
const error = ref('') 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() const goBack = () => router.back()
@@ -124,6 +134,8 @@ onMounted(async () => {
if (isEdit.value) { if (isEdit.value) {
await loadForEdit() await loadForEdit()
} }
await nextTick()
resizeQuestion()
}) })
</script> </script>
@@ -194,6 +206,13 @@ onMounted(async () => {
border-color: #F97316; border-color: #F97316;
} }
.question-input {
min-height: 44px;
overflow: hidden;
resize: none;
line-height: 1.55;
}
.switch-row { .switch-row {
display: flex; display: flex;
align-items: center; align-items: center;
+14 -3
View File
@@ -24,6 +24,8 @@
</div> </div>
<div class="model-meta"> <div class="model-meta">
<span>版本: {{ m.model_version || '--' }}</span> <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>温度: {{ m.temperature }}</span>
<span>Max Tokens: {{ m.max_tokens }}</span> <span>Max Tokens: {{ m.max_tokens }}</span>
<span>超时: {{ m.timeout_seconds }}s</span> <span>超时: {{ m.timeout_seconds }}s</span>
@@ -68,6 +70,15 @@
<el-form-item label="模型版本"> <el-form-item label="模型版本">
<el-input v-model="form.model_version" placeholder="如: gpt-4-turbo, glm-4" /> <el-input v-model="form.model_version" placeholder="如: gpt-4-turbo, glm-4" />
</el-form-item> </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-row :gutter="16">
<el-col :span="12"> <el-col :span="12">
<el-form-item label="温度"> <el-form-item label="温度">
@@ -142,7 +153,7 @@ const testing = ref(false)
const providerLabels = { openai: 'OpenAI', zhipu: '智谱GLM', wenxin: '文心一言', qianwen: '通义千问', local: '本地模型' } const providerLabels = { openai: 'OpenAI', zhipu: '智谱GLM', wenxin: '文心一言', qianwen: '通义千问', local: '本地模型' }
const scopeLabels = { general: '通用业务', digital_avatar: '数字分身专用' } 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 }] } const rules = { model_name: [{ required: true, message: '请输入模型名称' }], provider: [{ required: true }], usage_scope: [{ required: true }] }
async function load() { async function load() {
@@ -166,13 +177,13 @@ function onProviderChange(provider) {
function openCreate() { function openCreate() {
editModel.value = null 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 dialogVisible.value = true
} }
function openEdit(m) { function openEdit(m) {
editModel.value = 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 dialogVisible.value = true
} }