Compare commits
60
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
48f3b9bddf | ||
|
|
2122615725 | ||
|
|
cd0c6e3162 | ||
|
|
b48e33fc6b | ||
|
|
d9f12685fa | ||
|
|
b3367eedaa | ||
|
|
e6a4988cf9 | ||
|
|
becf2c7c52 | ||
|
|
d421a9da72 | ||
|
|
f4389612d1 | ||
|
|
ac5a332bce | ||
|
|
9ba89ef2d1 | ||
|
|
b1023eb783 | ||
|
|
2f287ef538 | ||
|
|
f5cbbe9eef | ||
|
|
5fc56143ee | ||
|
|
d521585bf2 | ||
|
|
35e47bf0a1 | ||
|
|
f52e42d9c0 | ||
|
|
1ab56ad0b1 | ||
|
|
1889c8ebba | ||
|
|
910a05107a | ||
|
|
857d6f2562 | ||
|
|
5ce12771b9 | ||
|
|
4c152230aa | ||
|
|
9ce46cd883 | ||
|
|
140ac20281 | ||
|
|
9afc2d5a6c | ||
|
|
8585d101d5 | ||
|
|
62eb9578fd | ||
|
|
848657219e | ||
|
|
1bcdcead8d | ||
|
|
434caac056 | ||
|
|
bb9e9da1f3 | ||
|
|
28fcd5373b | ||
|
|
2a01a9946a | ||
|
|
2ce1079bb6 | ||
|
|
6fba6dbaaa | ||
|
|
e71267cf86 | ||
|
|
359e558dbe | ||
|
|
3edf92c7cc | ||
|
|
97c4c73b58 | ||
|
|
08c58fe0e6 | ||
|
|
6b7201e890 | ||
|
|
28553aba15 | ||
|
|
95f91450d0 | ||
|
|
b98a2b9507 | ||
|
|
59350fb41d | ||
|
|
6a4b35c49a | ||
|
|
207bbd02cf | ||
|
|
7cac96356d | ||
|
|
03c32309a8 | ||
|
|
0fc43908ae | ||
|
|
3d999f9472 | ||
|
|
540edb58c4 | ||
|
|
a7eb6ac2a5 | ||
|
|
016bc22c05 | ||
|
|
094f8cd40f | ||
|
|
6794e88d53 | ||
|
|
7884430b3d |
@@ -9,6 +9,8 @@ backend/logs/
|
|||||||
# Node
|
# Node
|
||||||
frontend/node_modules/
|
frontend/node_modules/
|
||||||
frontend/dist/
|
frontend/dist/
|
||||||
|
uniapp-avatar/node_modules/
|
||||||
|
uniapp-avatar/dist/
|
||||||
|
|
||||||
# macOS
|
# macOS
|
||||||
.DS_Store
|
.DS_Store
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
"""API路由汇总"""
|
"""API路由汇总"""
|
||||||
from fastapi import APIRouter
|
from fastapi import APIRouter
|
||||||
from app.api.endpoints import users, interactions, ai_models, dashboard, system, logs, avatars
|
from app.api.endpoints import users, interactions, ai_models, dashboard, system, logs, avatars, finance
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
|
|
||||||
@@ -11,3 +11,4 @@ router.include_router(dashboard.router, prefix="/dashboard", tags=["数据看板
|
|||||||
router.include_router(system.router, prefix="/system", tags=["系统设置"])
|
router.include_router(system.router, prefix="/system", tags=["系统设置"])
|
||||||
router.include_router(logs.router, prefix="/logs", tags=["日志管理"])
|
router.include_router(logs.router, prefix="/logs", tags=["日志管理"])
|
||||||
router.include_router(avatars.router, prefix="/avatars", tags=["数字分身管理"])
|
router.include_router(avatars.router, prefix="/avatars", tags=["数字分身管理"])
|
||||||
|
router.include_router(finance.router, prefix="/finance", tags=["财务管理"])
|
||||||
|
|||||||
@@ -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(),
|
||||||
|
|||||||
@@ -0,0 +1,130 @@
|
|||||||
|
"""Admin finance API for avatar Token orders, refunds and invoices."""
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Body, HTTPException, Query
|
||||||
|
|
||||||
|
from app.schemas import ApiResponse
|
||||||
|
from app.services.avatar_service import get_session, is_available
|
||||||
|
from app.services.finance_service import FinanceServiceError, finance_service
|
||||||
|
|
||||||
|
|
||||||
|
router = APIRouter()
|
||||||
|
|
||||||
|
|
||||||
|
def _session():
|
||||||
|
if not is_available():
|
||||||
|
raise HTTPException(status_code=503, detail="数字分身数据库尚未初始化")
|
||||||
|
return get_session()
|
||||||
|
|
||||||
|
|
||||||
|
def _raise(exc: FinanceServiceError):
|
||||||
|
raise HTTPException(status_code=exc.status_code, detail=str(exc))
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/summary")
|
||||||
|
def summary():
|
||||||
|
db = _session()
|
||||||
|
try:
|
||||||
|
return ApiResponse(data=finance_service.summary(db))
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/orders")
|
||||||
|
def orders(
|
||||||
|
page: int = Query(1, ge=1),
|
||||||
|
page_size: int = Query(20, ge=1, le=100),
|
||||||
|
keyword: str = Query(""),
|
||||||
|
status: str = Query(""),
|
||||||
|
provider: str = Query(""),
|
||||||
|
):
|
||||||
|
db = _session()
|
||||||
|
try:
|
||||||
|
total, items = finance_service.list_orders(
|
||||||
|
db, page=page, page_size=page_size, keyword=keyword.strip(), status=status, provider=provider
|
||||||
|
)
|
||||||
|
return ApiResponse(data={"total": total, "page": page, "page_size": page_size, "items": items})
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
@router.patch("/orders/{order_no}/status")
|
||||||
|
def update_order_status(order_no: str, body: dict = Body(...)):
|
||||||
|
db = _session()
|
||||||
|
try:
|
||||||
|
finance_service.close_order(
|
||||||
|
db, order_no, status=str(body.get("status") or ""), reason=str(body.get("reason") or "")
|
||||||
|
)
|
||||||
|
return ApiResponse(message="订单状态已更新")
|
||||||
|
except FinanceServiceError as exc:
|
||||||
|
_raise(exc)
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/orders/{order_no}/refund")
|
||||||
|
def request_refund(order_no: str, body: dict = Body(...)):
|
||||||
|
try:
|
||||||
|
data = finance_service.request_refund(
|
||||||
|
order_no,
|
||||||
|
reason=str(body.get("reason") or "").strip(),
|
||||||
|
operator=str(body.get("operator") or "后台管理员").strip(),
|
||||||
|
)
|
||||||
|
return ApiResponse(data=data, message="退款申请已提交")
|
||||||
|
except FinanceServiceError as exc:
|
||||||
|
_raise(exc)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/refunds")
|
||||||
|
def refunds(
|
||||||
|
page: int = Query(1, ge=1),
|
||||||
|
page_size: int = Query(20, ge=1, le=100),
|
||||||
|
status: str = Query(""),
|
||||||
|
):
|
||||||
|
db = _session()
|
||||||
|
try:
|
||||||
|
total, items = finance_service.list_refunds(db, page=page, page_size=page_size, status=status)
|
||||||
|
return ApiResponse(data={"total": total, "page": page, "page_size": page_size, "items": items})
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/refunds/{refund_no}/confirm")
|
||||||
|
def confirm_refund(refund_no: str, body: dict = Body(...)):
|
||||||
|
try:
|
||||||
|
data = finance_service.confirm_refund(refund_no, body)
|
||||||
|
return ApiResponse(data=data, message="退款结果已登记")
|
||||||
|
except FinanceServiceError as exc:
|
||||||
|
_raise(exc)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/invoices")
|
||||||
|
def invoices(
|
||||||
|
page: int = Query(1, ge=1),
|
||||||
|
page_size: int = Query(20, ge=1, le=100),
|
||||||
|
status: str = Query(""),
|
||||||
|
):
|
||||||
|
db = _session()
|
||||||
|
try:
|
||||||
|
total, items = finance_service.list_invoices(db, page=page, page_size=page_size, status=status)
|
||||||
|
return ApiResponse(data={"total": total, "page": page, "page_size": page_size, "items": items})
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
@router.patch("/invoices/{invoice_id}")
|
||||||
|
def update_invoice(invoice_id: str, body: dict = Body(...)):
|
||||||
|
db = _session()
|
||||||
|
try:
|
||||||
|
finance_service.update_invoice(
|
||||||
|
db,
|
||||||
|
invoice_id,
|
||||||
|
status=str(body.get("status") or ""),
|
||||||
|
invoice_no=str(body.get("invoiceNo") or ""),
|
||||||
|
invoice_url=str(body.get("invoiceUrl") or ""),
|
||||||
|
remark=str(body.get("remark") or ""),
|
||||||
|
)
|
||||||
|
return ApiResponse(message="发票申请已处理")
|
||||||
|
except FinanceServiceError as exc:
|
||||||
|
_raise(exc)
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
@@ -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:
|
||||||
result = await conn.execute(text(
|
columns = (
|
||||||
"SELECT COUNT(*) FROM information_schema.COLUMNS "
|
(
|
||||||
"WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'ai_model_configs' "
|
"usage_scope",
|
||||||
"AND COLUMN_NAME = 'usage_scope'"
|
|
||||||
))
|
|
||||||
if result.scalar_one() == 0:
|
|
||||||
await conn.execute(text(
|
|
||||||
"ALTER TABLE ai_model_configs ADD COLUMN usage_scope "
|
"ALTER TABLE ai_model_configs ADD COLUMN usage_scope "
|
||||||
"VARCHAR(16) NOT NULL DEFAULT 'general' AFTER provider"
|
"VARCHAR(16) NOT NULL DEFAULT 'general' AFTER provider",
|
||||||
))
|
),
|
||||||
logger.info("AI模型配置表已增加 usage_scope 字段")
|
(
|
||||||
|
"vision_model_version",
|
||||||
|
"ALTER TABLE ai_model_configs ADD COLUMN vision_model_version "
|
||||||
|
"VARCHAR(64) NULL AFTER model_version",
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"ocr_model_version",
|
||||||
|
"ALTER TABLE ai_model_configs ADD COLUMN ocr_model_version "
|
||||||
|
"VARCHAR(64) NULL AFTER vision_model_version",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
for column_name, ddl in columns:
|
||||||
|
result = await conn.execute(text(
|
||||||
|
"SELECT COUNT(*) FROM information_schema.COLUMNS "
|
||||||
|
"WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'ai_model_configs' "
|
||||||
|
"AND COLUMN_NAME = :column_name"
|
||||||
|
), {"column_name": column_name})
|
||||||
|
if result.scalar_one() == 0:
|
||||||
|
await conn.execute(text(ddl))
|
||||||
|
logger.info("AI模型配置表已增加 %s 字段", column_name)
|
||||||
finally:
|
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("✅ 数据库初始化完成")
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -1,16 +1,24 @@
|
|||||||
"""数字分身管理服务层 — 同步连接数字分身应用的 SQLite 数据库"""
|
"""数字分身管理服务层 — 同步连接数字分身应用的 SQLite 数据库"""
|
||||||
|
import json
|
||||||
import os
|
import os
|
||||||
|
from datetime import datetime, timedelta
|
||||||
from typing import Optional, Tuple
|
from typing import Optional, Tuple
|
||||||
|
|
||||||
from sqlalchemy import create_engine, text
|
from sqlalchemy import create_engine, select, text
|
||||||
from sqlalchemy.orm import sessionmaker, Session
|
from sqlalchemy.orm import sessionmaker, Session
|
||||||
|
|
||||||
from app.core.config import settings
|
from app.core.config import settings
|
||||||
|
from app.core.logger import logger
|
||||||
|
from app.models import UserPersonality, VirtualUser
|
||||||
|
|
||||||
|
|
||||||
_engine = None
|
_engine = None
|
||||||
_SessionLocal: Optional[sessionmaker] = None
|
_SessionLocal: Optional[sessionmaker] = None
|
||||||
|
|
||||||
|
AVATAR_ACCOUNT_PREFIX = "__avatar__:"
|
||||||
|
SQUARE_INTERACTION_PERMISSION = "interact"
|
||||||
|
SQUARE_INTERACTION_ACTIONS = frozenset({"like", "collect", "comment", "reply"})
|
||||||
|
|
||||||
|
|
||||||
def _get_engine_and_session():
|
def _get_engine_and_session():
|
||||||
global _engine, _SessionLocal
|
global _engine, _SessionLocal
|
||||||
@@ -27,7 +35,7 @@ def _get_engine_and_session():
|
|||||||
return None, None
|
return None, None
|
||||||
_engine = create_engine(
|
_engine = create_engine(
|
||||||
f"sqlite:///{db_path}",
|
f"sqlite:///{db_path}",
|
||||||
connect_args={"check_same_thread": False},
|
connect_args={"check_same_thread": False, "timeout": 30},
|
||||||
)
|
)
|
||||||
_SessionLocal = sessionmaker(bind=_engine, autoflush=False, expire_on_commit=False)
|
_SessionLocal = sessionmaker(bind=_engine, autoflush=False, expire_on_commit=False)
|
||||||
return _engine, _SessionLocal()
|
return _engine, _SessionLocal()
|
||||||
@@ -66,6 +74,185 @@ def _get_global_token_balance(db: Session) -> int:
|
|||||||
return 0
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
def _decode_config(value) -> dict:
|
||||||
|
if isinstance(value, dict):
|
||||||
|
return value
|
||||||
|
if isinstance(value, str):
|
||||||
|
try:
|
||||||
|
decoded = json.loads(value)
|
||||||
|
return decoded if isinstance(decoded, dict) else {}
|
||||||
|
except (json.JSONDecodeError, ValueError):
|
||||||
|
return {}
|
||||||
|
return {}
|
||||||
|
|
||||||
|
|
||||||
|
def is_delegated_avatar_user(user: VirtualUser | None) -> bool:
|
||||||
|
return bool(user and (user.account or "").startswith(AVATAR_ACCOUNT_PREFIX))
|
||||||
|
|
||||||
|
|
||||||
|
def delegated_avatar_id(user: VirtualUser | None) -> str:
|
||||||
|
if not is_delegated_avatar_user(user):
|
||||||
|
return ""
|
||||||
|
return (user.account or "")[len(AVATAR_ACCOUNT_PREFIX):]
|
||||||
|
|
||||||
|
|
||||||
|
def _list_square_interaction_authorizations(db: Session) -> list[dict]:
|
||||||
|
"""读取已明确授权分身参与广场互动的身份与会会令牌。"""
|
||||||
|
rows = db.execute(text("""
|
||||||
|
SELECT
|
||||||
|
a.id AS avatar_id,
|
||||||
|
a.name AS avatar_name,
|
||||||
|
a.display_name AS avatar_display_name,
|
||||||
|
a.description AS avatar_description,
|
||||||
|
a.photo_url AS avatar_photo_url,
|
||||||
|
a.config AS avatar_config,
|
||||||
|
u.huihui_user_id,
|
||||||
|
u.nickname AS owner_nickname,
|
||||||
|
u.avatar_url AS owner_avatar_url,
|
||||||
|
u.huihui_token
|
||||||
|
FROM avatars a
|
||||||
|
JOIN users u ON u.huihui_user_id = a.owner_id
|
||||||
|
WHERE a.status = 'active'
|
||||||
|
""")).fetchall()
|
||||||
|
|
||||||
|
authorized = []
|
||||||
|
for row in rows:
|
||||||
|
config = _decode_config(row.avatar_config)
|
||||||
|
permissions = config.get("authorizationPermissions", [])
|
||||||
|
if not isinstance(permissions, list) or SQUARE_INTERACTION_PERMISSION not in permissions:
|
||||||
|
continue
|
||||||
|
platform_uid = str(row.huihui_user_id or "").strip()
|
||||||
|
token = str(row.huihui_token or "").strip()
|
||||||
|
if not platform_uid or not token:
|
||||||
|
continue
|
||||||
|
authorized.append({
|
||||||
|
"avatar_id": str(row.avatar_id),
|
||||||
|
"avatar_name": row.avatar_display_name or row.avatar_name or row.owner_nickname or "数字分身",
|
||||||
|
"avatar_description": row.avatar_description or "",
|
||||||
|
"avatar_url": _resolve_photo_url(row.avatar_photo_url or row.owner_avatar_url or ""),
|
||||||
|
"config": config,
|
||||||
|
"platform_uid": platform_uid,
|
||||||
|
"token": token,
|
||||||
|
})
|
||||||
|
return authorized
|
||||||
|
|
||||||
|
|
||||||
|
def get_square_interaction_permissions(avatar_id: str) -> frozenset[str]:
|
||||||
|
"""实时复核授权;数据库不可用、令牌失效或撤权时一律拒绝执行。"""
|
||||||
|
avatar_db = get_session()
|
||||||
|
if avatar_db is None:
|
||||||
|
return frozenset()
|
||||||
|
try:
|
||||||
|
authorized_ids = {
|
||||||
|
item["avatar_id"] for item in _list_square_interaction_authorizations(avatar_db)
|
||||||
|
}
|
||||||
|
return SQUARE_INTERACTION_ACTIONS if avatar_id in authorized_ids else frozenset()
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error(f"读取数字分身广场互动授权失败: {exc}")
|
||||||
|
return frozenset()
|
||||||
|
finally:
|
||||||
|
avatar_db.close()
|
||||||
|
|
||||||
|
|
||||||
|
def _word_count_range(config: dict) -> tuple[int, int]:
|
||||||
|
ranges = {
|
||||||
|
"short": (10, 35),
|
||||||
|
"medium": (20, 60),
|
||||||
|
"long": (30, 80),
|
||||||
|
}
|
||||||
|
return ranges.get(str(config.get("responseLength") or "medium"), (20, 60))
|
||||||
|
|
||||||
|
|
||||||
|
async def sync_square_interaction_users(db) -> set[str]:
|
||||||
|
"""把已授权分身同步为调度器身份,并刷新其会会会话。"""
|
||||||
|
avatar_db = get_session()
|
||||||
|
if avatar_db is None:
|
||||||
|
logger.warning("数字分身数据库不可用,跳过广场互动授权同步")
|
||||||
|
return set()
|
||||||
|
try:
|
||||||
|
authorized = _list_square_interaction_authorizations(avatar_db)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error(f"同步数字分身广场互动授权失败: {exc}")
|
||||||
|
return set()
|
||||||
|
finally:
|
||||||
|
avatar_db.close()
|
||||||
|
|
||||||
|
from app.core.redis_client import delete_session, set_session
|
||||||
|
|
||||||
|
result = await db.execute(
|
||||||
|
select(VirtualUser).where(VirtualUser.account.like(f"{AVATAR_ACCOUNT_PREFIX}%"))
|
||||||
|
)
|
||||||
|
existing_users = {delegated_avatar_id(user): user for user in result.scalars().all()}
|
||||||
|
authorized_ids = {item["avatar_id"] for item in authorized}
|
||||||
|
|
||||||
|
for avatar_id, user in existing_users.items():
|
||||||
|
if avatar_id not in authorized_ids:
|
||||||
|
user.is_enabled = 0
|
||||||
|
user.status = 0
|
||||||
|
user.session_token = None
|
||||||
|
user.session_expires_at = None
|
||||||
|
await delete_session(user.id)
|
||||||
|
|
||||||
|
for item in authorized:
|
||||||
|
avatar_id = item["avatar_id"]
|
||||||
|
user = existing_users.get(avatar_id)
|
||||||
|
if user is None:
|
||||||
|
user = VirtualUser(
|
||||||
|
nickname=item["avatar_name"],
|
||||||
|
account=f"{AVATAR_ACCOUNT_PREFIX}{avatar_id}",
|
||||||
|
password_enc="",
|
||||||
|
status=2,
|
||||||
|
is_enabled=1,
|
||||||
|
platform_uid=item["platform_uid"],
|
||||||
|
remark="用户授权的数字分身广场互动身份",
|
||||||
|
)
|
||||||
|
db.add(user)
|
||||||
|
await db.flush()
|
||||||
|
|
||||||
|
expires_at = datetime.now() + timedelta(days=1)
|
||||||
|
user.nickname = item["avatar_name"]
|
||||||
|
user.real_name = item["avatar_name"]
|
||||||
|
user.avatar_url = item["avatar_url"]
|
||||||
|
user.platform_uid = item["platform_uid"]
|
||||||
|
user.session_token = item["token"]
|
||||||
|
user.session_expires_at = expires_at
|
||||||
|
user.last_login_at = datetime.now()
|
||||||
|
user.status = 2
|
||||||
|
user.is_enabled = 1
|
||||||
|
|
||||||
|
config = item["config"]
|
||||||
|
personality_result = await db.execute(
|
||||||
|
select(UserPersonality).where(UserPersonality.user_id == user.id)
|
||||||
|
)
|
||||||
|
personality = personality_result.scalar_one_or_none()
|
||||||
|
word_min, word_max = _word_count_range(config)
|
||||||
|
prompt_parts = [item["avatar_description"], str(config.get("systemPrompt") or "")]
|
||||||
|
style_prompt = "\n".join(part.strip() for part in prompt_parts if part and part.strip())
|
||||||
|
if personality is None:
|
||||||
|
personality = UserPersonality(user_id=user.id)
|
||||||
|
db.add(personality)
|
||||||
|
personality.language_style = str(config.get("replyStyle") or "professional")
|
||||||
|
personality.personality_desc = item["avatar_description"]
|
||||||
|
personality.comment_style_prompt = style_prompt
|
||||||
|
personality.word_count_min = word_min
|
||||||
|
personality.word_count_max = word_max
|
||||||
|
|
||||||
|
await set_session(user.id, {
|
||||||
|
"token": item["token"],
|
||||||
|
"session_id": f"avatar:{avatar_id}",
|
||||||
|
"platform_uid": item["platform_uid"],
|
||||||
|
"org_id": "",
|
||||||
|
"login_time": datetime.now().isoformat(),
|
||||||
|
"nickname": item["avatar_name"],
|
||||||
|
"real_name": item["avatar_name"],
|
||||||
|
"avatar": item["avatar_url"],
|
||||||
|
"delegated_avatar_id": avatar_id,
|
||||||
|
}, expire=86400)
|
||||||
|
|
||||||
|
await db.commit()
|
||||||
|
return authorized_ids
|
||||||
|
|
||||||
|
|
||||||
class AvatarService:
|
class AvatarService:
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@@ -0,0 +1,254 @@
|
|||||||
|
"""Finance operations for digital-avatar Token purchases."""
|
||||||
|
|
||||||
|
import os
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
from sqlalchemy import text
|
||||||
|
|
||||||
|
from app.core.config import settings
|
||||||
|
|
||||||
|
|
||||||
|
class FinanceServiceError(RuntimeError):
|
||||||
|
def __init__(self, message: str, status_code: int = 400):
|
||||||
|
super().__init__(message)
|
||||||
|
self.status_code = status_code
|
||||||
|
|
||||||
|
|
||||||
|
def _mapping(row):
|
||||||
|
return dict(row._mapping) if row is not None else None
|
||||||
|
|
||||||
|
|
||||||
|
def _iso(value):
|
||||||
|
return value.isoformat() if hasattr(value, "isoformat") else value
|
||||||
|
|
||||||
|
|
||||||
|
def _money(cents):
|
||||||
|
return round(int(cents or 0) / 100, 2)
|
||||||
|
|
||||||
|
|
||||||
|
class FinanceService:
|
||||||
|
@staticmethod
|
||||||
|
def _tables_ready(db) -> bool:
|
||||||
|
names = {
|
||||||
|
row[0]
|
||||||
|
for row in db.execute(text(
|
||||||
|
"SELECT name FROM sqlite_master WHERE type='table' AND "
|
||||||
|
"name IN ('token_payment_orders','payment_refunds','invoice_applications')"
|
||||||
|
)).fetchall()
|
||||||
|
}
|
||||||
|
return len(names) == 3
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def summary(cls, db) -> dict:
|
||||||
|
if not cls._tables_ready(db):
|
||||||
|
return {
|
||||||
|
"paid_revenue": 0,
|
||||||
|
"paid_orders": 0,
|
||||||
|
"pending_orders": 0,
|
||||||
|
"processing_refunds": 0,
|
||||||
|
"pending_invoices": 0,
|
||||||
|
}
|
||||||
|
row = db.execute(text("""
|
||||||
|
SELECT
|
||||||
|
COALESCE(SUM(CASE WHEN status='paid' THEN price_cents ELSE 0 END), 0) paid_revenue,
|
||||||
|
SUM(CASE WHEN status='paid' THEN 1 ELSE 0 END) paid_orders,
|
||||||
|
SUM(CASE WHEN status='pending' THEN 1 ELSE 0 END) pending_orders
|
||||||
|
FROM token_payment_orders
|
||||||
|
""")).fetchone()
|
||||||
|
processing_refunds = db.execute(text(
|
||||||
|
"SELECT COUNT(*) FROM payment_refunds WHERE status IN ('pending','processing')"
|
||||||
|
)).scalar() or 0
|
||||||
|
pending_invoices = db.execute(text(
|
||||||
|
"SELECT COUNT(*) FROM invoice_applications WHERE status='pending'"
|
||||||
|
)).scalar() or 0
|
||||||
|
return {
|
||||||
|
"paid_revenue": _money(row.paid_revenue),
|
||||||
|
"paid_orders": int(row.paid_orders or 0),
|
||||||
|
"pending_orders": int(row.pending_orders or 0),
|
||||||
|
"processing_refunds": int(processing_refunds),
|
||||||
|
"pending_invoices": int(pending_invoices),
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def list_orders(cls, db, *, page=1, page_size=20, keyword="", status="", provider=""):
|
||||||
|
if not cls._tables_ready(db):
|
||||||
|
return 0, []
|
||||||
|
clauses = ["1=1"]
|
||||||
|
params = {}
|
||||||
|
if keyword:
|
||||||
|
clauses.append("(o.order_no LIKE :keyword OR u.phone LIKE :keyword OR u.nickname LIKE :keyword)")
|
||||||
|
params["keyword"] = f"%{keyword}%"
|
||||||
|
if status:
|
||||||
|
clauses.append("o.status=:status")
|
||||||
|
params["status"] = status
|
||||||
|
if provider:
|
||||||
|
clauses.append("o.provider=:provider")
|
||||||
|
params["provider"] = provider
|
||||||
|
where = " AND ".join(clauses)
|
||||||
|
total = db.execute(text(f"""
|
||||||
|
SELECT COUNT(*) FROM token_payment_orders o
|
||||||
|
LEFT JOIN users u ON u.id=o.user_id WHERE {where}
|
||||||
|
"""), params).scalar() or 0
|
||||||
|
params.update({"limit": page_size, "offset": (page - 1) * page_size})
|
||||||
|
rows = db.execute(text(f"""
|
||||||
|
SELECT o.*, u.nickname user_nickname, u.phone user_phone,
|
||||||
|
i.status invoice_status, i.id invoice_id
|
||||||
|
FROM token_payment_orders o
|
||||||
|
LEFT JOIN users u ON u.id=o.user_id
|
||||||
|
LEFT JOIN invoice_applications i ON i.order_no=o.order_no
|
||||||
|
WHERE {where}
|
||||||
|
ORDER BY o.created_at DESC LIMIT :limit OFFSET :offset
|
||||||
|
"""), params).fetchall()
|
||||||
|
items = []
|
||||||
|
for row in rows:
|
||||||
|
item = _mapping(row)
|
||||||
|
item["price"] = _money(item.pop("price_cents"))
|
||||||
|
for key in ("created_at", "updated_at", "paid_at", "refunded_at"):
|
||||||
|
item[key] = _iso(item.get(key))
|
||||||
|
item.pop("pay_message", None)
|
||||||
|
items.append(item)
|
||||||
|
return int(total), items
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def list_refunds(cls, db, *, page=1, page_size=20, status=""):
|
||||||
|
if not cls._tables_ready(db):
|
||||||
|
return 0, []
|
||||||
|
where = "WHERE r.status=:status" if status else ""
|
||||||
|
params = {"status": status} if status else {}
|
||||||
|
total = db.execute(text(f"SELECT COUNT(*) FROM payment_refunds r {where}"), params).scalar() or 0
|
||||||
|
params.update({"limit": page_size, "offset": (page - 1) * page_size})
|
||||||
|
rows = db.execute(text(f"""
|
||||||
|
SELECT r.*, o.provider, o.payment_method, u.nickname user_nickname, u.phone user_phone
|
||||||
|
FROM payment_refunds r
|
||||||
|
JOIN token_payment_orders o ON o.order_no=r.order_no
|
||||||
|
LEFT JOIN users u ON u.id=o.user_id
|
||||||
|
{where}
|
||||||
|
ORDER BY r.created_at DESC LIMIT :limit OFFSET :offset
|
||||||
|
"""), params).fetchall()
|
||||||
|
items = []
|
||||||
|
for row in rows:
|
||||||
|
item = _mapping(row)
|
||||||
|
item["amount"] = _money(item.pop("amount_cents"))
|
||||||
|
for key in ("created_at", "updated_at", "completed_at"):
|
||||||
|
item[key] = _iso(item.get(key))
|
||||||
|
items.append(item)
|
||||||
|
return int(total), items
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def list_invoices(cls, db, *, page=1, page_size=20, status=""):
|
||||||
|
if not cls._tables_ready(db):
|
||||||
|
return 0, []
|
||||||
|
where = "WHERE i.status=:status" if status else ""
|
||||||
|
params = {"status": status} if status else {}
|
||||||
|
total = db.execute(text(f"SELECT COUNT(*) FROM invoice_applications i {where}"), params).scalar() or 0
|
||||||
|
params.update({"limit": page_size, "offset": (page - 1) * page_size})
|
||||||
|
rows = db.execute(text(f"""
|
||||||
|
SELECT i.*, u.nickname user_nickname, u.phone user_phone
|
||||||
|
FROM invoice_applications i
|
||||||
|
LEFT JOIN users u ON u.id=i.user_id
|
||||||
|
{where}
|
||||||
|
ORDER BY i.created_at DESC LIMIT :limit OFFSET :offset
|
||||||
|
"""), params).fetchall()
|
||||||
|
items = []
|
||||||
|
for row in rows:
|
||||||
|
item = _mapping(row)
|
||||||
|
item["amount"] = _money(item.pop("amount_cents"))
|
||||||
|
for key in ("created_at", "updated_at", "issued_at"):
|
||||||
|
item[key] = _iso(item.get(key))
|
||||||
|
items.append(item)
|
||||||
|
return int(total), items
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def close_order(db, order_no: str, *, status: str, reason: str):
|
||||||
|
if status not in {"closed", "failed"}:
|
||||||
|
raise FinanceServiceError("后台只能将待支付订单关闭或标记失败")
|
||||||
|
order = db.execute(text(
|
||||||
|
"SELECT status FROM token_payment_orders WHERE order_no=:order_no"
|
||||||
|
), {"order_no": order_no}).fetchone()
|
||||||
|
if not order:
|
||||||
|
raise FinanceServiceError("订单不存在", 404)
|
||||||
|
if order.status != "pending":
|
||||||
|
raise FinanceServiceError("只有待支付订单可以修改状态", 409)
|
||||||
|
db.execute(text("""
|
||||||
|
UPDATE token_payment_orders
|
||||||
|
SET status=:status, failure_reason=:reason, updated_at=:updated_at
|
||||||
|
WHERE order_no=:order_no
|
||||||
|
"""), {
|
||||||
|
"status": status,
|
||||||
|
"reason": (reason or "后台关闭订单")[:500],
|
||||||
|
"updated_at": datetime.utcnow(),
|
||||||
|
"order_no": order_no,
|
||||||
|
})
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _avatar_admin_call(path: str, payload: dict):
|
||||||
|
base_url = (settings.AVATAR_BACKEND_URL or os.getenv("AVATAR_BACKEND_URL", "")).rstrip("/")
|
||||||
|
secret = os.getenv("AVATAR_FINANCE_ADMIN_SECRET", "").strip()
|
||||||
|
if not base_url or len(secret) < 16:
|
||||||
|
raise FinanceServiceError("数字分身财务服务尚未完成配置", 503)
|
||||||
|
try:
|
||||||
|
response = httpx.post(
|
||||||
|
f"{base_url}/api{path}",
|
||||||
|
json=payload,
|
||||||
|
headers={"X-Avatar-Finance-Key": secret},
|
||||||
|
timeout=35,
|
||||||
|
)
|
||||||
|
data = response.json()
|
||||||
|
except (httpx.HTTPError, ValueError) as exc:
|
||||||
|
raise FinanceServiceError("数字分身财务服务暂时不可用", 502) from exc
|
||||||
|
if response.status_code >= 400 or data.get("code") not in (0, 200, "0", "200"):
|
||||||
|
raise FinanceServiceError(data.get("message") or data.get("detail") or "财务操作失败", response.status_code)
|
||||||
|
return data.get("data")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def request_refund(cls, order_no: str, *, reason: str, operator: str):
|
||||||
|
return cls._avatar_admin_call(
|
||||||
|
f"/token/admin/orders/{order_no}/refund",
|
||||||
|
{"reason": reason, "operator": operator},
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def confirm_refund(cls, refund_no: str, payload: dict):
|
||||||
|
return cls._avatar_admin_call(f"/token/admin/refunds/{refund_no}/confirm", payload)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def update_invoice(db, invoice_id: str, *, status: str, invoice_no="", invoice_url="", remark=""):
|
||||||
|
row = db.execute(text(
|
||||||
|
"SELECT * FROM invoice_applications WHERE id=:invoice_id"
|
||||||
|
), {"invoice_id": invoice_id}).fetchone()
|
||||||
|
if not row:
|
||||||
|
raise FinanceServiceError("发票申请不存在", 404)
|
||||||
|
if row.status != "pending":
|
||||||
|
raise FinanceServiceError("该发票申请已处理", 409)
|
||||||
|
if status == "issued":
|
||||||
|
if not invoice_no.strip():
|
||||||
|
raise FinanceServiceError("请填写发票号码")
|
||||||
|
if invoice_url.strip() and not invoice_url.strip().lower().startswith(("https://", "http://")):
|
||||||
|
raise FinanceServiceError("电子发票地址必须是 HTTP 或 HTTPS 链接")
|
||||||
|
issued_at = datetime.utcnow()
|
||||||
|
elif status == "rejected":
|
||||||
|
if not remark.strip():
|
||||||
|
raise FinanceServiceError("请填写驳回原因")
|
||||||
|
issued_at = None
|
||||||
|
else:
|
||||||
|
raise FinanceServiceError("发票状态只能是已开具或已驳回")
|
||||||
|
db.execute(text("""
|
||||||
|
UPDATE invoice_applications
|
||||||
|
SET status=:status, invoice_no=:invoice_no, invoice_url=:invoice_url,
|
||||||
|
remark=:remark, issued_at=:issued_at, updated_at=:updated_at
|
||||||
|
WHERE id=:invoice_id
|
||||||
|
"""), {
|
||||||
|
"status": status,
|
||||||
|
"invoice_no": invoice_no.strip()[:120],
|
||||||
|
"invoice_url": invoice_url.strip()[:500],
|
||||||
|
"remark": remark.strip()[:500],
|
||||||
|
"issued_at": issued_at,
|
||||||
|
"updated_at": datetime.utcnow(),
|
||||||
|
"invoice_id": invoice_id,
|
||||||
|
})
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
|
||||||
|
finance_service = FinanceService()
|
||||||
@@ -23,6 +23,7 @@ class SchedulerService:
|
|||||||
from app.core.database import AsyncSessionLocal
|
from app.core.database import AsyncSessionLocal
|
||||||
logger.info("⚡ 立即触发互动任务")
|
logger.info("⚡ 立即触发互动任务")
|
||||||
async with AsyncSessionLocal() as session:
|
async with AsyncSessionLocal() as session:
|
||||||
|
await self._sync_delegated_avatar_users(session)
|
||||||
try:
|
try:
|
||||||
max_concurrent = int(await self._get_config(session, "max_concurrent_users", "5"))
|
max_concurrent = int(await self._get_config(session, "max_concurrent_users", "5"))
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
@@ -146,7 +147,9 @@ class SchedulerService:
|
|||||||
async def _check_sessions(self):
|
async def _check_sessions(self):
|
||||||
"""定时校验登录状态"""
|
"""定时校验登录状态"""
|
||||||
from app.services.news_service import news_service
|
from app.services.news_service import news_service
|
||||||
|
from app.services.avatar_service import is_delegated_avatar_user
|
||||||
async with AsyncSessionLocal() as db:
|
async with AsyncSessionLocal() as db:
|
||||||
|
await self._sync_delegated_avatar_users(db)
|
||||||
result = await db.execute(
|
result = await db.execute(
|
||||||
select(VirtualUser).where(VirtualUser.status == 2, VirtualUser.is_enabled == 1)
|
select(VirtualUser).where(VirtualUser.status == 2, VirtualUser.is_enabled == 1)
|
||||||
)
|
)
|
||||||
@@ -154,7 +157,7 @@ class SchedulerService:
|
|||||||
for user in users:
|
for user in users:
|
||||||
try:
|
try:
|
||||||
valid = await news_service.check_session(db, user)
|
valid = await news_service.check_session(db, user)
|
||||||
if not valid:
|
if not valid and not is_delegated_avatar_user(user):
|
||||||
logger.warning(f"用户 {user.account} 会话失效,尝试重登")
|
logger.warning(f"用户 {user.account} 会话失效,尝试重登")
|
||||||
await news_service.login(db, user)
|
await news_service.login(db, user)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -163,6 +166,7 @@ class SchedulerService:
|
|||||||
async def _run_interactions(self):
|
async def _run_interactions(self):
|
||||||
"""执行互动任务"""
|
"""执行互动任务"""
|
||||||
async with AsyncSessionLocal() as db:
|
async with AsyncSessionLocal() as db:
|
||||||
|
await self._sync_delegated_avatar_users(db)
|
||||||
# 检查调度器开关
|
# 检查调度器开关
|
||||||
enabled = await self._get_config(db, "scheduler_enabled", "true")
|
enabled = await self._get_config(db, "scheduler_enabled", "true")
|
||||||
if enabled != "true":
|
if enabled != "true":
|
||||||
@@ -184,8 +188,11 @@ class SchedulerService:
|
|||||||
logger.debug(f"[调度] 当前北京时间 {now_time} 不在互动时段 {start_str}-{end_str}")
|
logger.debug(f"[调度] 当前北京时间 {now_time} 不在互动时段 {start_str}-{end_str}")
|
||||||
return
|
return
|
||||||
|
|
||||||
# 获取最小互动间隔(秒)
|
# 获取互动间隔范围(秒),与调度设置页面字段保持一致
|
||||||
min_interval = int(await self._get_config(db, "interact_min_interval", "300"))
|
min_interval = await self._get_int_config(db, "interact_interval_min", 300)
|
||||||
|
max_interval = await self._get_int_config(db, "interact_interval_max", min_interval)
|
||||||
|
min_interval = max(0, min_interval)
|
||||||
|
max_interval = max(min_interval, max_interval)
|
||||||
|
|
||||||
# 获取最大并发
|
# 获取最大并发
|
||||||
max_concurrent = int(await self._get_config(db, "max_concurrent_users", "5"))
|
max_concurrent = int(await self._get_config(db, "max_concurrent_users", "5"))
|
||||||
@@ -204,7 +211,7 @@ class SchedulerService:
|
|||||||
await self._try_login_users(db)
|
await self._try_login_users(db)
|
||||||
return
|
return
|
||||||
|
|
||||||
# 检查互动间隔:过滤掉最近 min_interval 秒内已互动的用户
|
# 每个用户在其最小/最大间隔内取得稳定随机值,直到下次互动后再变化
|
||||||
now_dt = datetime.now()
|
now_dt = datetime.now()
|
||||||
eligible = []
|
eligible = []
|
||||||
for u in all_users:
|
for u in all_users:
|
||||||
@@ -212,11 +219,17 @@ class SchedulerService:
|
|||||||
eligible.append(u)
|
eligible.append(u)
|
||||||
else:
|
else:
|
||||||
elapsed = (now_dt - u.last_interact_at).total_seconds()
|
elapsed = (now_dt - u.last_interact_at).total_seconds()
|
||||||
if elapsed >= min_interval:
|
interval = random.Random(
|
||||||
|
f"{u.id}:{u.last_interact_at.isoformat()}"
|
||||||
|
).randint(min_interval, max_interval)
|
||||||
|
if elapsed >= interval:
|
||||||
eligible.append(u)
|
eligible.append(u)
|
||||||
|
|
||||||
if not eligible:
|
if not eligible:
|
||||||
logger.debug(f"[调度] 所有 {len(all_users)} 个用户在 {min_interval}s 内已互动,跳过本次")
|
logger.debug(
|
||||||
|
f"[调度] 所有 {len(all_users)} 个用户尚未达到 "
|
||||||
|
f"{min_interval}-{max_interval}s 随机互动间隔,跳过本次"
|
||||||
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
# 按最后互动时间升序排序:最久没互动的用户优先
|
# 按最后互动时间升序排序:最久没互动的用户优先
|
||||||
@@ -257,10 +270,12 @@ class SchedulerService:
|
|||||||
async def _try_login_users(self, db):
|
async def _try_login_users(self, db):
|
||||||
"""尝试登录未登录的用户"""
|
"""尝试登录未登录的用户"""
|
||||||
from app.services.news_service import news_service
|
from app.services.news_service import news_service
|
||||||
|
from app.services.avatar_service import AVATAR_ACCOUNT_PREFIX
|
||||||
result = await db.execute(
|
result = await db.execute(
|
||||||
select(VirtualUser).where(
|
select(VirtualUser).where(
|
||||||
VirtualUser.status.in_([0, 3]),
|
VirtualUser.status.in_([0, 3]),
|
||||||
VirtualUser.is_enabled == 1
|
VirtualUser.is_enabled == 1,
|
||||||
|
~VirtualUser.account.like(f"{AVATAR_ACCOUNT_PREFIX}%"),
|
||||||
).limit(3)
|
).limit(3)
|
||||||
)
|
)
|
||||||
users = result.scalars().all()
|
users = result.scalars().all()
|
||||||
@@ -275,6 +290,11 @@ class SchedulerService:
|
|||||||
"""执行单用户互动 - 基于真实接口"""
|
"""执行单用户互动 - 基于真实接口"""
|
||||||
from app.services.news_service import news_service
|
from app.services.news_service import news_service
|
||||||
from app.services.ai_service import ai_service
|
from app.services.ai_service import ai_service
|
||||||
|
from app.services.avatar_service import (
|
||||||
|
delegated_avatar_id,
|
||||||
|
get_square_interaction_permissions,
|
||||||
|
is_delegated_avatar_user,
|
||||||
|
)
|
||||||
|
|
||||||
async with AsyncSessionLocal() as db:
|
async with AsyncSessionLocal() as db:
|
||||||
try:
|
try:
|
||||||
@@ -289,6 +309,23 @@ class SchedulerService:
|
|||||||
"interactions": [],
|
"interactions": [],
|
||||||
}
|
}
|
||||||
|
|
||||||
|
allowed_actions = {"like", "collect", "comment", "reply", "forward"}
|
||||||
|
if is_delegated_avatar_user(user):
|
||||||
|
allowed_actions = set(
|
||||||
|
get_square_interaction_permissions(delegated_avatar_id(user))
|
||||||
|
)
|
||||||
|
if not allowed_actions:
|
||||||
|
user.status = 0
|
||||||
|
user.is_enabled = 0
|
||||||
|
await db.commit()
|
||||||
|
return {
|
||||||
|
"user_id": user.id,
|
||||||
|
"account": user.account,
|
||||||
|
"status": "skipped",
|
||||||
|
"reason": "avatar_interaction_not_authorized",
|
||||||
|
"interactions": [],
|
||||||
|
}
|
||||||
|
|
||||||
# 检查今日评论限额
|
# 检查今日评论限额
|
||||||
can_comment = True
|
can_comment = True
|
||||||
if user.today_comment_count >= user.daily_comment_limit:
|
if user.today_comment_count >= user.daily_comment_limit:
|
||||||
@@ -398,14 +435,53 @@ class SchedulerService:
|
|||||||
interactions_done = []
|
interactions_done = []
|
||||||
action_failures = []
|
action_failures = []
|
||||||
|
|
||||||
# ① 先记录阅读(每次必做,模拟真实用户打开文章)
|
done_on_this = today_done.get(news_id, set())
|
||||||
|
wants = {
|
||||||
|
"like": (
|
||||||
|
"like" in allowed_actions
|
||||||
|
and "like" not in done_on_this
|
||||||
|
and random.random() < like_prob
|
||||||
|
),
|
||||||
|
"collect": (
|
||||||
|
"collect" in allowed_actions
|
||||||
|
and "collect" not in done_on_this
|
||||||
|
and random.random() < collect_prob
|
||||||
|
),
|
||||||
|
"forward": (
|
||||||
|
"forward" in allowed_actions
|
||||||
|
and "forward" not in done_on_this
|
||||||
|
and random.random() < forward_prob
|
||||||
|
),
|
||||||
|
"reply": (
|
||||||
|
"reply" in allowed_actions
|
||||||
|
and can_comment
|
||||||
|
and personality is not None
|
||||||
|
and random.random() < reply_prob
|
||||||
|
),
|
||||||
|
"comment": (
|
||||||
|
"comment" in allowed_actions
|
||||||
|
and can_comment
|
||||||
|
and personality is not None
|
||||||
|
and not already_commented_this
|
||||||
|
and random.random() < comment_prob
|
||||||
|
),
|
||||||
|
}
|
||||||
|
if not any(wants.values()):
|
||||||
|
return {
|
||||||
|
"user_id": user.id,
|
||||||
|
"account": user.account,
|
||||||
|
"status": "skipped",
|
||||||
|
"reason": "no_actions_triggered",
|
||||||
|
"interactions": [],
|
||||||
|
"article_id": news_id,
|
||||||
|
"article_title": news_title,
|
||||||
|
}
|
||||||
|
|
||||||
|
# 只有动作命中调度概率后才打开文章
|
||||||
await news_service.read_news(db, user, news_id)
|
await news_service.read_news(db, user, news_id)
|
||||||
|
|
||||||
# 今日已对此文章做过的互动类型
|
|
||||||
done_on_this = today_done.get(news_id, set())
|
|
||||||
|
|
||||||
# ② 点赞(每篇文章每用户每天只点赞一次)
|
# ② 点赞(每篇文章每用户每天只点赞一次)
|
||||||
if "like" not in done_on_this and random.random() < like_prob:
|
if wants["like"]:
|
||||||
success, err = await news_service.like_news(db, user, news_id, org_id=article_org_id, to_user_id=news_author, title=news_title)
|
success, err = await news_service.like_news(db, user, news_id, org_id=article_org_id, to_user_id=news_author, title=news_title)
|
||||||
await self._save_record(db, user, news_id, news_title, "like", None, 0, success, err)
|
await self._save_record(db, user, news_id, news_title, "like", None, 0, success, err)
|
||||||
if success:
|
if success:
|
||||||
@@ -415,16 +491,17 @@ class SchedulerService:
|
|||||||
action_failures.append({"type": "like", "error": err})
|
action_failures.append({"type": "like", "error": err})
|
||||||
|
|
||||||
# ③ 收藏(每篇文章每用户每天只收藏一次)
|
# ③ 收藏(每篇文章每用户每天只收藏一次)
|
||||||
if "collect" not in done_on_this and random.random() < collect_prob:
|
if wants["collect"]:
|
||||||
success, err = await news_service.collect_news(db, user, news_id, org_id=article_org_id, to_user_id=news_author, title=news_title)
|
success, err = await news_service.collect_news(db, user, news_id, org_id=article_org_id, to_user_id=news_author, title=news_title)
|
||||||
await self._save_record(db, user, news_id, news_title, "collect", None, 0, success, err)
|
await self._save_record(db, user, news_id, news_title, "collect", None, 0, success, err)
|
||||||
if success:
|
if success:
|
||||||
interactions_done.append("collect")
|
interactions_done.append("collect")
|
||||||
|
await self._incr_total(db, user_id)
|
||||||
else:
|
else:
|
||||||
action_failures.append({"type": "collect", "error": err})
|
action_failures.append({"type": "collect", "error": err})
|
||||||
|
|
||||||
# ④ 转发(每篇文章每用户每天只转发一次)
|
# ④ 转发(每篇文章每用户每天只转发一次)
|
||||||
if "forward" not in done_on_this and random.random() < forward_prob:
|
if wants["forward"]:
|
||||||
success, err = await news_service.forward_news(db, user, news_id)
|
success, err = await news_service.forward_news(db, user, news_id)
|
||||||
await self._save_record(db, user, news_id, news_title, "forward", None, 0, success, err)
|
await self._save_record(db, user, news_id, news_title, "forward", None, 0, success, err)
|
||||||
if success:
|
if success:
|
||||||
@@ -438,7 +515,7 @@ class SchedulerService:
|
|||||||
style_prompt = personality.comment_style_prompt or ""
|
style_prompt = personality.comment_style_prompt or ""
|
||||||
safe_word_max = min(personality.word_count_max, 80)
|
safe_word_max = min(personality.word_count_max, 80)
|
||||||
|
|
||||||
if random.random() < reply_prob:
|
if wants["reply"]:
|
||||||
reply_actions, reply_failures = await self._run_reply_interaction_chain(
|
reply_actions, reply_failures = await self._run_reply_interaction_chain(
|
||||||
db=db,
|
db=db,
|
||||||
starter=user,
|
starter=user,
|
||||||
@@ -455,7 +532,7 @@ class SchedulerService:
|
|||||||
action_failures.extend(reply_failures)
|
action_failures.extend(reply_failures)
|
||||||
|
|
||||||
# 每篇文章每个用户每天只发一条顶层评论;回复不再要求先评论
|
# 每篇文章每个用户每天只发一条顶层评论;回复不再要求先评论
|
||||||
if not already_commented_this and random.random() < comment_prob:
|
if wants["comment"]:
|
||||||
comment_text, tokens = await ai_service.generate_comment(
|
comment_text, tokens = await ai_service.generate_comment(
|
||||||
db, news_title, news_content,
|
db, news_title, news_content,
|
||||||
style_prompt, personality.word_count_min, safe_word_max
|
style_prompt, personality.word_count_min, safe_word_max
|
||||||
@@ -679,6 +756,7 @@ class SchedulerService:
|
|||||||
|
|
||||||
async with AsyncSessionLocal() as db:
|
async with AsyncSessionLocal() as db:
|
||||||
try:
|
try:
|
||||||
|
await self._sync_delegated_avatar_users(db)
|
||||||
now = datetime.now()
|
now = datetime.now()
|
||||||
await db.execute(
|
await db.execute(
|
||||||
update(PendingReplyTask)
|
update(PendingReplyTask)
|
||||||
@@ -706,6 +784,12 @@ class SchedulerService:
|
|||||||
logger.error(f"待发送回复队列处理异常: {e}")
|
logger.error(f"待发送回复队列处理异常: {e}")
|
||||||
|
|
||||||
async def _process_pending_reply_task(self, db, task: PendingReplyTask, news_service, ai_service):
|
async def _process_pending_reply_task(self, db, task: PendingReplyTask, news_service, ai_service):
|
||||||
|
from app.services.avatar_service import (
|
||||||
|
delegated_avatar_id,
|
||||||
|
get_square_interaction_permissions,
|
||||||
|
is_delegated_avatar_user,
|
||||||
|
)
|
||||||
|
|
||||||
task.status = 1
|
task.status = 1
|
||||||
task.locked_at = datetime.now()
|
task.locked_at = datetime.now()
|
||||||
task.attempts = (task.attempts or 0) + 1
|
task.attempts = (task.attempts or 0) + 1
|
||||||
@@ -716,6 +800,13 @@ class SchedulerService:
|
|||||||
task.status = 3
|
task.status = 3
|
||||||
task.last_error = "用户未登录或已禁用"
|
task.last_error = "用户未登录或已禁用"
|
||||||
return
|
return
|
||||||
|
if (
|
||||||
|
is_delegated_avatar_user(actor)
|
||||||
|
and "reply" not in get_square_interaction_permissions(delegated_avatar_id(actor))
|
||||||
|
):
|
||||||
|
task.status = 3
|
||||||
|
task.last_error = "数字分身广场互动授权已撤销"
|
||||||
|
return
|
||||||
|
|
||||||
reply_result = await self._post_contextual_reply(
|
reply_result = await self._post_contextual_reply(
|
||||||
db=db,
|
db=db,
|
||||||
@@ -858,6 +949,16 @@ class SchedulerService:
|
|||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
return default
|
return default
|
||||||
|
|
||||||
|
async def _sync_delegated_avatar_users(self, db):
|
||||||
|
from app.services.avatar_service import sync_square_interaction_users
|
||||||
|
|
||||||
|
try:
|
||||||
|
return await sync_square_interaction_users(db)
|
||||||
|
except Exception as exc:
|
||||||
|
await db.rollback()
|
||||||
|
logger.error(f"数字分身广场互动身份同步异常: {exc}")
|
||||||
|
return set()
|
||||||
|
|
||||||
async def _incr_total(self, db, user_id: int):
|
async def _incr_total(self, db, user_id: int):
|
||||||
await db.execute(
|
await db.execute(
|
||||||
update(VirtualUser).where(VirtualUser.id == user_id).values(
|
update(VirtualUser).where(VirtualUser.id == user_id).values(
|
||||||
|
|||||||
@@ -0,0 +1,126 @@
|
|||||||
|
import json
|
||||||
|
import os
|
||||||
|
import sqlite3
|
||||||
|
import tempfile
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from app.services import avatar_service
|
||||||
|
|
||||||
|
|
||||||
|
class AvatarSquareAuthorizationTests(unittest.TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
fd, self.db_path = tempfile.mkstemp(suffix=".db")
|
||||||
|
os.close(fd)
|
||||||
|
connection = sqlite3.connect(self.db_path)
|
||||||
|
connection.executescript("""
|
||||||
|
CREATE TABLE users (
|
||||||
|
huihui_user_id TEXT,
|
||||||
|
nickname TEXT,
|
||||||
|
avatar_url TEXT,
|
||||||
|
huihui_token TEXT
|
||||||
|
);
|
||||||
|
CREATE TABLE avatars (
|
||||||
|
id TEXT,
|
||||||
|
owner_id TEXT,
|
||||||
|
name TEXT,
|
||||||
|
display_name TEXT,
|
||||||
|
description TEXT,
|
||||||
|
photo_url TEXT,
|
||||||
|
config TEXT,
|
||||||
|
status TEXT
|
||||||
|
);
|
||||||
|
""")
|
||||||
|
connection.execute(
|
||||||
|
"INSERT INTO users VALUES (?, ?, ?, ?)",
|
||||||
|
("huihui-7", "主人", "/owner.jpg", "huihui-token"),
|
||||||
|
)
|
||||||
|
connection.commit()
|
||||||
|
connection.close()
|
||||||
|
avatar_service._engine = None
|
||||||
|
avatar_service._SessionLocal = None
|
||||||
|
self.path_patch = patch.object(avatar_service.settings, "AVATAR_DB_PATH", self.db_path)
|
||||||
|
self.path_patch.start()
|
||||||
|
|
||||||
|
def tearDown(self):
|
||||||
|
self.path_patch.stop()
|
||||||
|
if avatar_service._engine is not None:
|
||||||
|
avatar_service._engine.dispose()
|
||||||
|
avatar_service._engine = None
|
||||||
|
avatar_service._SessionLocal = None
|
||||||
|
os.unlink(self.db_path)
|
||||||
|
|
||||||
|
def _insert_avatar(self, permissions, *, status="active", token=None):
|
||||||
|
connection = sqlite3.connect(self.db_path)
|
||||||
|
connection.execute(
|
||||||
|
"INSERT INTO avatars VALUES (?, ?, ?, ?, ?, ?, ?, ?)",
|
||||||
|
(
|
||||||
|
"avatar-7",
|
||||||
|
"huihui-7",
|
||||||
|
"avatar",
|
||||||
|
"小会",
|
||||||
|
"语气友好,表达简洁",
|
||||||
|
"/avatar.jpg",
|
||||||
|
json.dumps({
|
||||||
|
"authorizationPermissions": permissions,
|
||||||
|
"replyStyle": "warm",
|
||||||
|
"responseLength": "short",
|
||||||
|
}),
|
||||||
|
status,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
if token is not None:
|
||||||
|
connection.execute(
|
||||||
|
"UPDATE users SET huihui_token = ? WHERE huihui_user_id = ?",
|
||||||
|
(token, "huihui-7"),
|
||||||
|
)
|
||||||
|
connection.commit()
|
||||||
|
connection.close()
|
||||||
|
|
||||||
|
def test_interact_permission_exposes_only_requested_square_actions(self):
|
||||||
|
self._insert_avatar(["chat", "interact"])
|
||||||
|
|
||||||
|
permissions = avatar_service.get_square_interaction_permissions("avatar-7")
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
permissions,
|
||||||
|
frozenset({"like", "collect", "comment", "reply"}),
|
||||||
|
)
|
||||||
|
self.assertNotIn("forward", permissions)
|
||||||
|
|
||||||
|
def test_missing_permission_inactive_avatar_or_missing_token_denies_execution(self):
|
||||||
|
scenarios = [
|
||||||
|
(["chat"], "active", "huihui-token"),
|
||||||
|
(["interact"], "inactive", "huihui-token"),
|
||||||
|
(["interact"], "active", ""),
|
||||||
|
]
|
||||||
|
for permissions, status, token in scenarios:
|
||||||
|
with self.subTest(permissions=permissions, status=status, token=token):
|
||||||
|
connection = sqlite3.connect(self.db_path)
|
||||||
|
connection.execute("DELETE FROM avatars")
|
||||||
|
connection.commit()
|
||||||
|
connection.close()
|
||||||
|
self._insert_avatar(permissions, status=status, token=token)
|
||||||
|
self.assertEqual(
|
||||||
|
avatar_service.get_square_interaction_permissions("avatar-7"),
|
||||||
|
frozenset(),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_delegated_avatar_identity_is_recognized_without_matching_normal_users(self):
|
||||||
|
delegated = SimpleNamespace(account="__avatar__:avatar-7")
|
||||||
|
normal = SimpleNamespace(account="13800000000")
|
||||||
|
|
||||||
|
self.assertTrue(avatar_service.is_delegated_avatar_user(delegated))
|
||||||
|
self.assertEqual(avatar_service.delegated_avatar_id(delegated), "avatar-7")
|
||||||
|
self.assertFalse(avatar_service.is_delegated_avatar_user(normal))
|
||||||
|
self.assertEqual(avatar_service.delegated_avatar_id(normal), "")
|
||||||
|
|
||||||
|
def test_response_length_maps_to_scheduler_comment_limits(self):
|
||||||
|
self.assertEqual(avatar_service._word_count_range({"responseLength": "short"}), (10, 35))
|
||||||
|
self.assertEqual(avatar_service._word_count_range({"responseLength": "long"}), (30, 80))
|
||||||
|
self.assertEqual(avatar_service._word_count_range({"responseLength": "unknown"}), (20, 60))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,91 @@
|
|||||||
|
from sqlalchemy import create_engine, text
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
from app.services.finance_service import FinanceService
|
||||||
|
|
||||||
|
|
||||||
|
def _db():
|
||||||
|
engine = create_engine("sqlite:///:memory:")
|
||||||
|
db = Session(engine)
|
||||||
|
db.execute(text("""
|
||||||
|
CREATE TABLE users (id TEXT PRIMARY KEY, nickname TEXT, phone TEXT)
|
||||||
|
"""))
|
||||||
|
db.execute(text("""
|
||||||
|
CREATE TABLE token_payment_orders (
|
||||||
|
id TEXT, order_no TEXT PRIMARY KEY, user_id TEXT, plan_id TEXT,
|
||||||
|
payment_method TEXT, pay_type TEXT, pay_way TEXT, points_amount INTEGER,
|
||||||
|
price_cents INTEGER, status TEXT, provider TEXT, provider_order_id TEXT,
|
||||||
|
provider_order_no TEXT, provider_status TEXT, pay_message TEXT,
|
||||||
|
failure_reason TEXT, refund_status TEXT, created_at TEXT, updated_at TEXT,
|
||||||
|
paid_at TEXT, refunded_at TEXT
|
||||||
|
)
|
||||||
|
"""))
|
||||||
|
db.execute(text("""
|
||||||
|
CREATE TABLE payment_refunds (
|
||||||
|
id TEXT, refund_no TEXT, order_no TEXT, amount_cents INTEGER,
|
||||||
|
points_amount INTEGER, reason TEXT, status TEXT, provider_refund_no TEXT,
|
||||||
|
requested_by TEXT, failure_reason TEXT, created_at TEXT, updated_at TEXT,
|
||||||
|
completed_at TEXT
|
||||||
|
)
|
||||||
|
"""))
|
||||||
|
db.execute(text("""
|
||||||
|
CREATE TABLE invoice_applications (
|
||||||
|
id TEXT, order_no TEXT, user_id TEXT, amount_cents INTEGER, title TEXT,
|
||||||
|
invoice_type TEXT, tax_number TEXT, email TEXT, status TEXT,
|
||||||
|
invoice_no TEXT, invoice_url TEXT, remark TEXT, created_at TEXT,
|
||||||
|
updated_at TEXT, issued_at TEXT
|
||||||
|
)
|
||||||
|
"""))
|
||||||
|
db.execute(text("INSERT INTO users VALUES ('u1','测试用户','13800000000')"))
|
||||||
|
db.execute(text("""
|
||||||
|
INSERT INTO token_payment_orders VALUES (
|
||||||
|
'o1','AV1','u1','1','wechat','WECHAT','APP',2000000,1000,'paid','huihui',
|
||||||
|
'','','SUCCESS','secret-payment-message','','none','2026-09-08 12:00:00',
|
||||||
|
'2026-09-08 12:01:00','2026-09-08 12:01:00',NULL
|
||||||
|
)
|
||||||
|
"""))
|
||||||
|
db.execute(text("""
|
||||||
|
INSERT INTO invoice_applications VALUES (
|
||||||
|
'i1','AV1','u1',1000,'测试用户','personal','','u@example.com','pending',
|
||||||
|
'','','','2026-09-08 12:02:00','2026-09-08 12:02:00',NULL
|
||||||
|
)
|
||||||
|
"""))
|
||||||
|
db.commit()
|
||||||
|
return db
|
||||||
|
|
||||||
|
|
||||||
|
def test_finance_summary_and_orders_hide_provider_payment_payload():
|
||||||
|
db = _db()
|
||||||
|
try:
|
||||||
|
summary = FinanceService.summary(db)
|
||||||
|
assert summary == {
|
||||||
|
"paid_revenue": 10.0,
|
||||||
|
"paid_orders": 1,
|
||||||
|
"pending_orders": 0,
|
||||||
|
"processing_refunds": 0,
|
||||||
|
"pending_invoices": 1,
|
||||||
|
}
|
||||||
|
total, orders = FinanceService.list_orders(db, keyword="测试用户")
|
||||||
|
assert total == 1
|
||||||
|
assert orders[0]["price"] == 10.0
|
||||||
|
assert orders[0]["invoice_status"] == "pending"
|
||||||
|
assert "pay_message" not in orders[0]
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
def test_invoice_can_be_issued_and_pending_order_can_be_closed():
|
||||||
|
db = _db()
|
||||||
|
try:
|
||||||
|
FinanceService.update_invoice(db, "i1", status="issued", invoice_no="FP-001", invoice_url="", remark="")
|
||||||
|
assert db.execute(text("SELECT status, invoice_no FROM invoice_applications WHERE id='i1'" )).fetchone() == ("issued", "FP-001")
|
||||||
|
db.execute(text("""
|
||||||
|
INSERT INTO token_payment_orders
|
||||||
|
(id,order_no,user_id,plan_id,payment_method,pay_type,pay_way,points_amount,price_cents,status,provider,refund_status)
|
||||||
|
VALUES ('o2','AV2','u1','1','alipay','ALIPAY','H5',1,100,'pending','huihui','none')
|
||||||
|
"""))
|
||||||
|
db.commit()
|
||||||
|
FinanceService.close_order(db, "AV2", status="closed", reason="超时")
|
||||||
|
assert db.execute(text("SELECT status FROM token_payment_orders WHERE order_no='AV2'" )).scalar() == "closed"
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
@@ -1,12 +1,16 @@
|
|||||||
# 构建阶段:安装依赖并打包 H5
|
# 构建阶段:安装依赖并打包 H5
|
||||||
FROM node:18-alpine AS build
|
FROM node:18-alpine AS build
|
||||||
|
|
||||||
|
ARG APP_GIT_SHA=unknown
|
||||||
|
ARG APP_BUILD_TIME=unknown
|
||||||
|
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
COPY package*.json ./
|
COPY package*.json ./
|
||||||
RUN npm ci
|
RUN npm ci
|
||||||
|
|
||||||
COPY . .
|
COPY . .
|
||||||
|
RUN printf '{"gitSha":"%s","buildTime":"%s"}\n' "$APP_GIT_SHA" "$APP_BUILD_TIME" > public/version.json
|
||||||
RUN npm run build
|
RUN npm run build
|
||||||
|
|
||||||
# 运行阶段:nginx 托管静态资源并反向代理 /api 到后端
|
# 运行阶段:nginx 托管静态资源并反向代理 /api 到后端
|
||||||
@@ -14,6 +18,11 @@ RUN npm run build
|
|||||||
# 新版 nginx(>=1.31) 用 pwrite 写 pid 文件会被拦导致致命退出;1.28 用 write() 可正常启动。
|
# 新版 nginx(>=1.31) 用 pwrite 写 pid 文件会被拦导致致命退出;1.28 用 write() 可正常启动。
|
||||||
FROM nginx:1.28-alpine
|
FROM nginx:1.28-alpine
|
||||||
|
|
||||||
|
ARG APP_GIT_SHA=unknown
|
||||||
|
ARG APP_BUILD_TIME=unknown
|
||||||
|
LABEL org.opencontainers.image.revision=${APP_GIT_SHA} \
|
||||||
|
org.opencontainers.image.created=${APP_BUILD_TIME}
|
||||||
|
|
||||||
COPY --from=build /app/dist /usr/share/nginx/html
|
COPY --from=build /app/dist /usr/share/nginx/html
|
||||||
# 覆盖 nginx 默认主配置(含唯一可写的 pid /tmp/nginx.pid,规避受限容器内 /run 不可写导致反复重启)
|
# 覆盖 nginx 默认主配置(含唯一可写的 pid /tmp/nginx.pid,规避受限容器内 /run 不可写导致反复重启)
|
||||||
COPY nginx.conf /etc/nginx/nginx.conf
|
COPY nginx.conf /etc/nginx/nginx.conf
|
||||||
|
|||||||
@@ -7,6 +7,13 @@ WORKDIR /app
|
|||||||
COPY requirements.txt .
|
COPY requirements.txt .
|
||||||
RUN pip install --no-cache-dir --timeout 120 --retries 10 -i https://pypi.tuna.tsinghua.edu.cn/simple -r requirements.txt
|
RUN pip install --no-cache-dir --timeout 120 --retries 10 -i https://pypi.tuna.tsinghua.edu.cn/simple -r requirements.txt
|
||||||
|
|
||||||
|
ARG APP_GIT_SHA=unknown
|
||||||
|
ARG APP_BUILD_TIME=unknown
|
||||||
|
ENV APP_GIT_SHA=${APP_GIT_SHA} \
|
||||||
|
APP_BUILD_TIME=${APP_BUILD_TIME}
|
||||||
|
LABEL org.opencontainers.image.revision=${APP_GIT_SHA} \
|
||||||
|
org.opencontainers.image.created=${APP_BUILD_TIME}
|
||||||
|
|
||||||
COPY . .
|
COPY . .
|
||||||
|
|
||||||
# 后端使用 SQLite(avatar.db 落在 /app 内),平铺结构以 `uvicorn main:app` 启动
|
# 后端使用 SQLite(avatar.db 落在 /app 内),平铺结构以 `uvicorn main:app` 启动
|
||||||
|
|||||||
@@ -1,16 +1,30 @@
|
|||||||
import os
|
import os
|
||||||
|
|
||||||
from sqlalchemy import create_engine
|
from sqlalchemy import create_engine, event
|
||||||
from sqlalchemy.orm import sessionmaker, declarative_base, Session
|
from sqlalchemy.orm import sessionmaker, declarative_base, Session
|
||||||
|
|
||||||
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
|
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||||
DB_FILE = os.path.join(BASE_DIR, "avatar.db")
|
DB_FILE = os.path.join(BASE_DIR, "avatar.db")
|
||||||
DATABASE_URL = os.getenv("DATABASE_URL", f"sqlite:///{DB_FILE}")
|
DATABASE_URL = os.getenv("DATABASE_URL", f"sqlite:///{DB_FILE}")
|
||||||
|
|
||||||
|
IS_SQLITE = DATABASE_URL.startswith("sqlite:")
|
||||||
engine = create_engine(
|
engine = create_engine(
|
||||||
DATABASE_URL,
|
DATABASE_URL,
|
||||||
connect_args={"check_same_thread": False} if DATABASE_URL.startswith("sqlite:") else {},
|
connect_args={"check_same_thread": False, "timeout": 30} if IS_SQLITE else {},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if IS_SQLITE:
|
||||||
|
@event.listens_for(engine, "connect")
|
||||||
|
def _configure_sqlite_connection(dbapi_connection, _connection_record):
|
||||||
|
cursor = dbapi_connection.cursor()
|
||||||
|
try:
|
||||||
|
cursor.execute("PRAGMA synchronous=NORMAL")
|
||||||
|
cursor.execute("PRAGMA busy_timeout=30000")
|
||||||
|
finally:
|
||||||
|
cursor.close()
|
||||||
|
|
||||||
|
|
||||||
SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
|
SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
|
||||||
Base = declarative_base()
|
Base = declarative_base()
|
||||||
|
|
||||||
@@ -26,6 +40,10 @@ def get_db():
|
|||||||
def init_db():
|
def init_db():
|
||||||
import models
|
import models
|
||||||
|
|
||||||
|
if IS_SQLITE:
|
||||||
|
with engine.connect() as conn:
|
||||||
|
conn.exec_driver_sql("PRAGMA journal_mode=WAL")
|
||||||
|
conn.commit()
|
||||||
Base.metadata.create_all(bind=engine)
|
Base.metadata.create_all(bind=engine)
|
||||||
|
|
||||||
# 轻量迁移:为已存在的表补充新列(SQLite 不支持自动 ALTER,逐列尝试)
|
# 轻量迁移:为已存在的表补充新列(SQLite 不支持自动 ALTER,逐列尝试)
|
||||||
@@ -35,6 +53,9 @@ def init_db():
|
|||||||
("knowledge_docs", "embedding_model", "VARCHAR DEFAULT ''"),
|
("knowledge_docs", "embedding_model", "VARCHAR DEFAULT ''"),
|
||||||
("knowledge_docs", "chunk_count", "INTEGER DEFAULT 0"),
|
("knowledge_docs", "chunk_count", "INTEGER DEFAULT 0"),
|
||||||
("knowledge_docs", "vectorized_at", "TIMESTAMP"),
|
("knowledge_docs", "vectorized_at", "TIMESTAMP"),
|
||||||
|
("knowledge_docs", "error_message", "VARCHAR DEFAULT ''"),
|
||||||
|
("knowledge_docs", "index_stage", "VARCHAR DEFAULT ''"),
|
||||||
|
("knowledge_docs", "index_progress", "INTEGER DEFAULT 0"),
|
||||||
("avatars", "owner_id", "VARCHAR DEFAULT ''"),
|
("avatars", "owner_id", "VARCHAR DEFAULT ''"),
|
||||||
("authorizations", "takeover_enabled", "BOOLEAN DEFAULT 0"),
|
("authorizations", "takeover_enabled", "BOOLEAN DEFAULT 0"),
|
||||||
("authorizations", "takeover_mode", "VARCHAR DEFAULT 'immediate'"),
|
("authorizations", "takeover_mode", "VARCHAR DEFAULT 'immediate'"),
|
||||||
@@ -45,10 +66,18 @@ def init_db():
|
|||||||
("token_account", "total_consumed", "BIGINT DEFAULT 0"),
|
("token_account", "total_consumed", "BIGINT DEFAULT 0"),
|
||||||
("token_account", "created_at", "TIMESTAMP"),
|
("token_account", "created_at", "TIMESTAMP"),
|
||||||
("token_account", "updated_at", "TIMESTAMP"),
|
("token_account", "updated_at", "TIMESTAMP"),
|
||||||
|
("token_plans", "virtual_product_id", "VARCHAR DEFAULT ''"),
|
||||||
|
("token_payment_orders", "provider", "VARCHAR DEFAULT 'huihui'"),
|
||||||
|
("token_payment_orders", "refund_status", "VARCHAR DEFAULT 'none'"),
|
||||||
|
("token_payment_orders", "refunded_at", "TIMESTAMP"),
|
||||||
|
("users", "wechat_mp_openid", "VARCHAR DEFAULT ''"),
|
||||||
|
("users", "wechat_mp_session_key", "VARCHAR DEFAULT ''"),
|
||||||
|
("takeover_messages", "attachment_id", "VARCHAR DEFAULT NULL"),
|
||||||
)
|
)
|
||||||
_normalize_optional_unique_values()
|
_normalize_optional_unique_values()
|
||||||
_normalize_takeover_delays()
|
_normalize_takeover_delays()
|
||||||
_create_token_indexes()
|
_create_token_indexes()
|
||||||
|
_create_payment_indexes()
|
||||||
|
|
||||||
|
|
||||||
def _try_add_columns(*cols):
|
def _try_add_columns(*cols):
|
||||||
@@ -82,3 +111,19 @@ def _create_token_indexes():
|
|||||||
"CREATE UNIQUE INDEX IF NOT EXISTS ux_token_account_user_id "
|
"CREATE UNIQUE INDEX IF NOT EXISTS ux_token_account_user_id "
|
||||||
"ON token_account(user_id) WHERE user_id <> ''"
|
"ON token_account(user_id) WHERE user_id <> ''"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _create_payment_indexes():
|
||||||
|
with engine.begin() as conn:
|
||||||
|
conn.exec_driver_sql(
|
||||||
|
"CREATE INDEX IF NOT EXISTS ix_token_payment_orders_provider "
|
||||||
|
"ON token_payment_orders(provider)"
|
||||||
|
)
|
||||||
|
conn.exec_driver_sql(
|
||||||
|
"CREATE INDEX IF NOT EXISTS ix_token_payment_orders_refund_status "
|
||||||
|
"ON token_payment_orders(refund_status)"
|
||||||
|
)
|
||||||
|
conn.exec_driver_sql(
|
||||||
|
"CREATE INDEX IF NOT EXISTS ix_users_wechat_mp_openid "
|
||||||
|
"ON users(wechat_mp_openid)"
|
||||||
|
)
|
||||||
|
|||||||
@@ -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 逐字(中文无空格,需拆到字级才能命中子词)
|
||||||
@@ -43,11 +51,11 @@ def _hash_embedding(texts, dim=EMBED_DIM):
|
|||||||
return vecs
|
return vecs
|
||||||
|
|
||||||
|
|
||||||
def embed(texts):
|
def embed(texts, on_progress=None):
|
||||||
"""返回 list[list[float]],与输入顺序一致。"""
|
"""返回 list[list[float]],与输入顺序一致。"""
|
||||||
if not texts:
|
if not texts:
|
||||||
return []
|
return []
|
||||||
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")
|
||||||
@@ -56,6 +64,7 @@ def embed(texts):
|
|||||||
except ValueError:
|
except ValueError:
|
||||||
batch_size = 10
|
batch_size = 10
|
||||||
embeddings = []
|
embeddings = []
|
||||||
|
total = len(texts)
|
||||||
for start in range(0, len(texts), batch_size):
|
for start in range(0, len(texts), batch_size):
|
||||||
batch = texts[start:start + batch_size]
|
batch = texts[start:start + batch_size]
|
||||||
payload = json.dumps({"input": batch, "model": model}).encode("utf-8")
|
payload = json.dumps({"input": batch, "model": model}).encode("utf-8")
|
||||||
@@ -76,8 +85,13 @@ def embed(texts):
|
|||||||
if len(items) != len(batch):
|
if len(items) != len(batch):
|
||||||
raise ValueError("embedding response count does not match request")
|
raise ValueError("embedding response count does not match request")
|
||||||
embeddings.extend(item["embedding"] for item in items)
|
embeddings.extend(item["embedding"] for item in items)
|
||||||
|
if on_progress:
|
||||||
|
on_progress(len(embeddings), total)
|
||||||
return embeddings
|
return embeddings
|
||||||
return _hash_embedding(texts)
|
vectors = _hash_embedding(texts)
|
||||||
|
if on_progress:
|
||||||
|
on_progress(len(vectors), len(texts))
|
||||||
|
return vectors
|
||||||
|
|
||||||
|
|
||||||
def cosine(a, b):
|
def cosine(a, b):
|
||||||
|
|||||||
@@ -1,13 +1,14 @@
|
|||||||
from fastapi import FastAPI
|
from fastapi import FastAPI
|
||||||
from fastapi.middleware.cors import CORSMiddleware
|
from fastapi.middleware.cors import CORSMiddleware
|
||||||
|
|
||||||
import os
|
import importlib.util
|
||||||
import logging
|
import logging
|
||||||
|
import os
|
||||||
|
|
||||||
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
||||||
from apscheduler.triggers.interval import IntervalTrigger
|
from apscheduler.triggers.interval import IntervalTrigger
|
||||||
|
|
||||||
from database import init_db, SessionLocal
|
from database import engine, init_db, SessionLocal
|
||||||
from models import Avatar, Authorization, Organization, TokenAccount, TokenPlan, User
|
from models import Avatar, Authorization, Organization, TokenAccount, TokenPlan, User
|
||||||
from fastapi.staticfiles import StaticFiles
|
from fastapi.staticfiles import StaticFiles
|
||||||
import routers.avatars
|
import routers.avatars
|
||||||
@@ -19,11 +20,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")
|
||||||
|
|
||||||
@@ -51,7 +55,31 @@ app.mount("/api/files", StaticFiles(directory=UPLOAD_DIR), name="knowledge-files
|
|||||||
|
|
||||||
@app.get("/api/health")
|
@app.get("/api/health")
|
||||||
def health():
|
def health():
|
||||||
return ok({"status": "ok"})
|
checks = _runtime_checks()
|
||||||
|
return ok({
|
||||||
|
"status": "ok" if all(checks.values()) else "degraded",
|
||||||
|
"gitSha": os.getenv("APP_GIT_SHA", "unknown"),
|
||||||
|
"buildTime": os.getenv("APP_BUILD_TIME", "unknown"),
|
||||||
|
"checks": checks,
|
||||||
|
})
|
||||||
|
|
||||||
|
|
||||||
|
def _runtime_checks():
|
||||||
|
return {
|
||||||
|
"database": _database_is_ready(),
|
||||||
|
"uploads": os.path.isdir(UPLOAD_DIR) and os.access(UPLOAD_DIR, os.W_OK),
|
||||||
|
"pdfOcr": importlib.util.find_spec("pymupdf") is not None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _database_is_ready():
|
||||||
|
try:
|
||||||
|
with engine.connect() as connection:
|
||||||
|
connection.exec_driver_sql("SELECT 1")
|
||||||
|
return True
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Database readiness check failed")
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
def seed():
|
def seed():
|
||||||
@@ -129,9 +157,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 +190,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 +241,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()
|
||||||
|
|||||||
@@ -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__ = (
|
||||||
@@ -189,6 +190,9 @@ class KnowledgeDoc(Base):
|
|||||||
file_size = Column(Integer, default=0)
|
file_size = Column(Integer, default=0)
|
||||||
file_url = Column(String, default="")
|
file_url = Column(String, default="")
|
||||||
status = Column(String, default="uploaded") # uploaded | parsing | ready | failed
|
status = Column(String, default="uploaded") # uploaded | parsing | ready | failed
|
||||||
|
error_message = Column(String, default="") # 建立索引失败原因
|
||||||
|
index_stage = Column(String, default="") # queued | extracting | chunking | embedding | ready | failed
|
||||||
|
index_progress = Column(Integer, default=0) # 0-100
|
||||||
vectorized = Column(Boolean, default=False) # 是否已向量化
|
vectorized = Column(Boolean, default=False) # 是否已向量化
|
||||||
embedding_model = Column(String, default="") # 向量模型标识
|
embedding_model = Column(String, default="") # 向量模型标识
|
||||||
chunk_count = Column(Integer, default=0) # 切片数量
|
chunk_count = Column(Integer, default=0) # 切片数量
|
||||||
@@ -204,6 +208,9 @@ class KnowledgeDoc(Base):
|
|||||||
"fileSize": self.file_size,
|
"fileSize": self.file_size,
|
||||||
"fileUrl": self.file_url,
|
"fileUrl": self.file_url,
|
||||||
"status": self.status,
|
"status": self.status,
|
||||||
|
"errorMessage": self.error_message or "",
|
||||||
|
"indexStage": self.index_stage or "",
|
||||||
|
"indexProgress": int(self.index_progress or 0),
|
||||||
"vectorized": bool(self.vectorized),
|
"vectorized": bool(self.vectorized),
|
||||||
"embeddingModel": self.embedding_model,
|
"embeddingModel": self.embedding_model,
|
||||||
"chunkCount": self.chunk_count,
|
"chunkCount": self.chunk_count,
|
||||||
@@ -257,6 +264,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)
|
||||||
@@ -299,6 +344,7 @@ class TokenPlan(Base):
|
|||||||
price = Column(Float, default=0)
|
price = Column(Float, default=0)
|
||||||
badge = Column(String, default="")
|
badge = Column(String, default="")
|
||||||
desc = Column(String, default="")
|
desc = Column(String, default="")
|
||||||
|
virtual_product_id = Column(String, default="")
|
||||||
|
|
||||||
def to_dict(self):
|
def to_dict(self):
|
||||||
return {
|
return {
|
||||||
@@ -308,6 +354,7 @@ class TokenPlan(Base):
|
|||||||
"price": self.price,
|
"price": self.price,
|
||||||
"badge": self.badge,
|
"badge": self.badge,
|
||||||
"desc": self.desc,
|
"desc": self.desc,
|
||||||
|
"virtualProductId": self.virtual_product_id,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -324,14 +371,17 @@ class TokenPaymentOrder(Base):
|
|||||||
points_amount = Column(BigInteger, nullable=False)
|
points_amount = Column(BigInteger, nullable=False)
|
||||||
price_cents = Column(Integer, nullable=False)
|
price_cents = Column(Integer, nullable=False)
|
||||||
status = Column(String, nullable=False, default="pending", index=True)
|
status = Column(String, nullable=False, default="pending", index=True)
|
||||||
|
provider = Column(String, nullable=False, default="huihui", index=True)
|
||||||
provider_order_id = Column(String, default="")
|
provider_order_id = Column(String, default="")
|
||||||
provider_order_no = Column(String, default="")
|
provider_order_no = Column(String, default="")
|
||||||
provider_status = Column(String, default="")
|
provider_status = Column(String, default="")
|
||||||
pay_message = Column(Text, default="")
|
pay_message = Column(Text, default="")
|
||||||
failure_reason = Column(String, default="")
|
failure_reason = Column(String, default="")
|
||||||
|
refund_status = Column(String, nullable=False, default="none", index=True)
|
||||||
created_at = Column(DateTime, server_default=func.now())
|
created_at = Column(DateTime, server_default=func.now())
|
||||||
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
|
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
|
||||||
paid_at = Column(DateTime)
|
paid_at = Column(DateTime)
|
||||||
|
refunded_at = Column(DateTime)
|
||||||
|
|
||||||
def to_dict(self):
|
def to_dict(self):
|
||||||
return {
|
return {
|
||||||
@@ -344,11 +394,117 @@ class TokenPaymentOrder(Base):
|
|||||||
"pointsAmount": self.points_amount,
|
"pointsAmount": self.points_amount,
|
||||||
"price": self.price_cents / 100,
|
"price": self.price_cents / 100,
|
||||||
"status": self.status,
|
"status": self.status,
|
||||||
|
"provider": self.provider,
|
||||||
"providerStatus": self.provider_status,
|
"providerStatus": self.provider_status,
|
||||||
"payMessage": self.pay_message,
|
"payMessage": self.pay_message,
|
||||||
"failureReason": self.failure_reason,
|
"failureReason": self.failure_reason,
|
||||||
|
"refundStatus": self.refund_status,
|
||||||
"createdAt": _iso(self.created_at),
|
"createdAt": _iso(self.created_at),
|
||||||
"paidAt": _iso(self.paid_at),
|
"paidAt": _iso(self.paid_at),
|
||||||
|
"refundedAt": _iso(self.refunded_at),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class PaymentTransaction(Base):
|
||||||
|
"""Auditable provider event for one Token purchase order."""
|
||||||
|
|
||||||
|
__tablename__ = "payment_transactions"
|
||||||
|
__table_args__ = (
|
||||||
|
Index("ix_payment_transactions_order_created", "order_no", "created_at"),
|
||||||
|
)
|
||||||
|
|
||||||
|
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
|
||||||
|
order_no = Column(String, nullable=False, index=True)
|
||||||
|
provider = Column(String, nullable=False, default="huihui")
|
||||||
|
transaction_no = Column(String, nullable=False, default="")
|
||||||
|
event_type = Column(String, nullable=False, default="payment")
|
||||||
|
status = Column(String, nullable=False, default="pending")
|
||||||
|
amount_cents = Column(Integer, nullable=False, default=0)
|
||||||
|
raw_summary = Column(Text, default="")
|
||||||
|
created_at = Column(DateTime, server_default=func.now())
|
||||||
|
|
||||||
|
def to_dict(self):
|
||||||
|
return {
|
||||||
|
"id": self.id,
|
||||||
|
"orderNo": self.order_no,
|
||||||
|
"provider": self.provider,
|
||||||
|
"transactionNo": self.transaction_no,
|
||||||
|
"eventType": self.event_type,
|
||||||
|
"status": self.status,
|
||||||
|
"amount": self.amount_cents / 100,
|
||||||
|
"createdAt": _iso(self.created_at),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class PaymentRefund(Base):
|
||||||
|
__tablename__ = "payment_refunds"
|
||||||
|
|
||||||
|
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
|
||||||
|
refund_no = Column(String, nullable=False, unique=True, index=True)
|
||||||
|
order_no = Column(String, nullable=False, index=True)
|
||||||
|
amount_cents = Column(Integer, nullable=False)
|
||||||
|
points_amount = Column(BigInteger, nullable=False)
|
||||||
|
reason = Column(String, default="")
|
||||||
|
status = Column(String, nullable=False, default="pending", index=True)
|
||||||
|
provider_refund_no = Column(String, default="")
|
||||||
|
requested_by = Column(String, default="admin")
|
||||||
|
failure_reason = Column(String, default="")
|
||||||
|
created_at = Column(DateTime, server_default=func.now())
|
||||||
|
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
|
||||||
|
completed_at = Column(DateTime)
|
||||||
|
|
||||||
|
def to_dict(self):
|
||||||
|
return {
|
||||||
|
"id": self.id,
|
||||||
|
"refundNo": self.refund_no,
|
||||||
|
"orderNo": self.order_no,
|
||||||
|
"amount": self.amount_cents / 100,
|
||||||
|
"pointsAmount": self.points_amount,
|
||||||
|
"reason": self.reason,
|
||||||
|
"status": self.status,
|
||||||
|
"providerRefundNo": self.provider_refund_no,
|
||||||
|
"requestedBy": self.requested_by,
|
||||||
|
"failureReason": self.failure_reason,
|
||||||
|
"createdAt": _iso(self.created_at),
|
||||||
|
"completedAt": _iso(self.completed_at),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class InvoiceApplication(Base):
|
||||||
|
__tablename__ = "invoice_applications"
|
||||||
|
|
||||||
|
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
|
||||||
|
order_no = Column(String, nullable=False, unique=True, index=True)
|
||||||
|
user_id = Column(String, nullable=False, index=True)
|
||||||
|
amount_cents = Column(Integer, nullable=False)
|
||||||
|
title = Column(String, nullable=False)
|
||||||
|
invoice_type = Column(String, nullable=False, default="personal")
|
||||||
|
tax_number = Column(String, default="")
|
||||||
|
email = Column(String, default="")
|
||||||
|
status = Column(String, nullable=False, default="pending", index=True)
|
||||||
|
invoice_no = Column(String, default="")
|
||||||
|
invoice_url = Column(String, default="")
|
||||||
|
remark = Column(String, default="")
|
||||||
|
created_at = Column(DateTime, server_default=func.now())
|
||||||
|
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
|
||||||
|
issued_at = Column(DateTime)
|
||||||
|
|
||||||
|
def to_dict(self):
|
||||||
|
return {
|
||||||
|
"id": self.id,
|
||||||
|
"orderNo": self.order_no,
|
||||||
|
"userId": self.user_id,
|
||||||
|
"amount": self.amount_cents / 100,
|
||||||
|
"title": self.title,
|
||||||
|
"invoiceType": self.invoice_type,
|
||||||
|
"taxNumber": self.tax_number,
|
||||||
|
"email": self.email,
|
||||||
|
"status": self.status,
|
||||||
|
"invoiceNo": self.invoice_no,
|
||||||
|
"invoiceUrl": self.invoice_url,
|
||||||
|
"remark": self.remark,
|
||||||
|
"createdAt": _iso(self.created_at),
|
||||||
|
"issuedAt": _iso(self.issued_at),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -363,6 +519,9 @@ class User(Base):
|
|||||||
avatar_url = Column(String, default="")
|
avatar_url = Column(String, default="")
|
||||||
huihui_token = Column(String, default="") # 会会 access_token
|
huihui_token = Column(String, default="") # 会会 access_token
|
||||||
app_token = Column(String, default="") # 本系统会话 token
|
app_token = Column(String, default="") # 本系统会话 token
|
||||||
|
wechat_mp_openid = Column(String, default="", index=True)
|
||||||
|
# 微信 session_key 仅保存在服务端,用于虚拟支付用户态签名,绝不下发客户端。
|
||||||
|
wechat_mp_session_key = Column(String, default="")
|
||||||
last_login_at = Column(DateTime)
|
last_login_at = Column(DateTime)
|
||||||
created_at = Column(DateTime, server_default=func.now())
|
created_at = Column(DateTime, server_default=func.now())
|
||||||
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
|
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
|
||||||
|
|||||||
@@ -5,6 +5,8 @@ pydantic
|
|||||||
python-multipart
|
python-multipart
|
||||||
httpx
|
httpx
|
||||||
pypdf
|
pypdf
|
||||||
|
PyMuPDF>=1.24,<2
|
||||||
python-docx
|
python-docx
|
||||||
openpyxl
|
openpyxl
|
||||||
apscheduler>=3.10
|
apscheduler>=3.10
|
||||||
|
Pillow>=10.4
|
||||||
|
|||||||
@@ -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,6 +47,25 @@ QA_SEMANTIC_THRESHOLD = 0.72
|
|||||||
QA_MATCH_MARGIN = 0.06
|
QA_MATCH_MARGIN = 0.06
|
||||||
KNOWLEDGE_MIN_SCORE = float(os.getenv("KNOWLEDGE_MIN_SCORE", "0.42"))
|
KNOWLEDGE_MIN_SCORE = float(os.getenv("KNOWLEDGE_MIN_SCORE", "0.42"))
|
||||||
|
|
||||||
|
_IMAGE_ACCESS_DENIAL_PATTERNS = (
|
||||||
|
re.compile(
|
||||||
|
r"(?:我|目前|暂时|这里|本身|系统)?\s*(?:无法|不能|没法|不支持)\s*"
|
||||||
|
r"(?:直接)?\s*(?:查看|看到|看见|识别|读取|访问|打开|分析|理解)"
|
||||||
|
r"(?:\s*(?:或|、|/)\s*(?:查看|看到|看见|识别|读取|访问|打开|分析|理解))*\s*"
|
||||||
|
r"(?:你(?:发|提供|上传)的|这张|该|当前)?\s*(?:图片|图像|照片|影像|文件)"
|
||||||
|
),
|
||||||
|
re.compile(
|
||||||
|
r"(?:我|这里|目前|暂时)?\s*(?:看不到|看不见|未看到|没有看到|没收到|未收到)\s*"
|
||||||
|
r"(?:你(?:发|提供|上传)的|这张|该|当前)?\s*(?:图片|图像|照片|影像)"
|
||||||
|
),
|
||||||
|
re.compile(
|
||||||
|
r"\b(?:i\s+)?(?:can(?:not|'t)|am\s+unable\s+to)\s+(?:directly\s+)?"
|
||||||
|
r"(?:view|see|access|read|analy[sz]e|recogni[sz]e)\s+"
|
||||||
|
r"(?:the\s+|this\s+|your\s+)?(?:image|photo|picture|scan)\b",
|
||||||
|
re.IGNORECASE,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
_WRITING_SYSTEM_PATTERNS = {
|
_WRITING_SYSTEM_PATTERNS = {
|
||||||
"han": re.compile(r"[\u3400-\u4dbf\u4e00-\u9fff]"),
|
"han": re.compile(r"[\u3400-\u4dbf\u4e00-\u9fff]"),
|
||||||
"latin": re.compile(r"[A-Za-z\u00c0-\u024f]"),
|
"latin": re.compile(r"[A-Za-z\u00c0-\u024f]"),
|
||||||
@@ -49,14 +81,35 @@ _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:
|
||||||
@@ -77,6 +130,324 @@ 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)
|
||||||
@@ -107,6 +478,24 @@ def _qa_requires_language_adaptation(question: str, answer: str) -> bool:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _qa_requires_per_turn_rendering(
|
||||||
|
question: str,
|
||||||
|
answer: str,
|
||||||
|
history: list[Any],
|
||||||
|
) -> bool:
|
||||||
|
"""Keep the direct QA fast path only when no conversation can bias language."""
|
||||||
|
return bool(history) or _qa_requires_language_adaptation(question, answer)
|
||||||
|
|
||||||
|
|
||||||
|
def _per_turn_language_instruction() -> str:
|
||||||
|
return (
|
||||||
|
"本轮语言覆盖指令:只根据紧随其后的最新用户消息判断本轮回答语言。"
|
||||||
|
"即使此前整段对话一直使用另一种语言,只要最新消息切换了语言,本轮就必须立即切换到相同语言;"
|
||||||
|
"不要沿用上一轮语言。若最新消息明确指定回答语言,以该指定为准;若混用多种语言,使用其中占主导的"
|
||||||
|
"自然语言。不要说明你检测、切换或翻译了语言。"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _canonicalize_question(value: str) -> str:
|
def _canonicalize_question(value: str) -> str:
|
||||||
value = _normalize_question(value)
|
value = _normalize_question(value)
|
||||||
replacements = (
|
replacements = (
|
||||||
@@ -233,6 +622,7 @@ def _build_prompt(
|
|||||||
knowledge_hits: list[dict],
|
knowledge_hits: list[dict],
|
||||||
*,
|
*,
|
||||||
standard_answer: str = "",
|
standard_answer: str = "",
|
||||||
|
image_contexts: list[dict] | None = None,
|
||||||
) -> list[dict]:
|
) -> list[dict]:
|
||||||
config = _config(avatar)
|
config = _config(avatar)
|
||||||
description = (getattr(avatar, "description", "") or "").strip()
|
description = (getattr(avatar, "description", "") or "").strip()
|
||||||
@@ -241,6 +631,7 @@ def _build_prompt(
|
|||||||
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 (
|
||||||
@@ -265,6 +656,26 @@ def _build_prompt(
|
|||||||
)
|
)
|
||||||
if config["systemPrompt"]:
|
if config["systemPrompt"]:
|
||||||
system += f"\n额外系统提示词:{config['systemPrompt']}"
|
system += f"\n额外系统提示词:{config['systemPrompt']}"
|
||||||
|
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:
|
if standard_answer:
|
||||||
system += (
|
system += (
|
||||||
f"\n以下是本次问题命中的已确认标准答案:\n{standard_answer.strip()}"
|
f"\n以下是本次问题命中的已确认标准答案:\n{standard_answer.strip()}"
|
||||||
@@ -277,6 +688,11 @@ def _build_prompt(
|
|||||||
"\n涉及事实、专业判断、地址、流程、数据或建议时,只能依据本人资料、标准问答形成的上下文"
|
"\n涉及事实、专业判断、地址、流程、数据或建议时,只能依据本人资料、标准问答形成的上下文"
|
||||||
"和以上可靠资料作答,不要补充资料之外的通用知识或自行推测。"
|
"和以上可靠资料作答,不要补充资料之外的通用知识或自行推测。"
|
||||||
)
|
)
|
||||||
|
elif image_contexts:
|
||||||
|
system += (
|
||||||
|
"\n本次没有命中标准答题对或文件知识库,但已提供图片识别资料。只能围绕图片中的可确认内容、"
|
||||||
|
"本人资料和当前对话作答;不要补充图片之外的事实、专业判断或具体建议。"
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
system += (
|
system += (
|
||||||
"\n本次问题没有检索到可靠资料。除自然寒暄和基于本人资料的回答外,不要凭通用知识给出事实、"
|
"\n本次问题没有检索到可靠资料。除自然寒暄和基于本人资料的回答外,不要凭通用知识给出事实、"
|
||||||
@@ -311,6 +727,8 @@ def _build_prompt(
|
|||||||
messages = [{"role": "system", "content": system}]
|
messages = [{"role": "system", "content": system}]
|
||||||
for item in history[-MAX_HISTORY_MESSAGES:]:
|
for item in history[-MAX_HISTORY_MESSAGES:]:
|
||||||
messages.append({"role": item.role, "content": item.content} if hasattr(item, "role") else item)
|
messages.append({"role": item.role, "content": item.content} if hasattr(item, "role") else item)
|
||||||
|
# Keep the language instruction adjacent to the current turn so long histories cannot override it.
|
||||||
|
messages.append({"role": "system", "content": _per_turn_language_instruction()})
|
||||||
messages.append({"role": "user", "content": question.strip()})
|
messages.append({"role": "user", "content": question.strip()})
|
||||||
return messages
|
return messages
|
||||||
|
|
||||||
@@ -439,14 +857,17 @@ 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)
|
||||||
adapt_qa_language = bool(
|
adapt_qa_language = bool(
|
||||||
matched and _qa_requires_language_adaptation(question, matched.answer)
|
matched and _qa_requires_per_turn_rendering(question, matched.answer, history)
|
||||||
)
|
)
|
||||||
if matched and not adapt_qa_language:
|
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:
|
if matched:
|
||||||
@@ -457,11 +878,19 @@ def _resolve_reply(
|
|||||||
question,
|
question,
|
||||||
hits,
|
hits,
|
||||||
standard_answer=matched.answer,
|
standard_answer=matched.answer,
|
||||||
|
image_contexts=image_contexts,
|
||||||
)
|
)
|
||||||
else:
|
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 = 0.0 if matched else min(
|
temperature = 0.0 if matched else min(
|
||||||
0.45 if hits else 0.25,
|
0.45 if hits else 0.25,
|
||||||
@@ -496,9 +925,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": "qa" if matched else ("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:
|
||||||
@@ -514,14 +953,17 @@ 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)
|
||||||
adapt_qa_language = bool(
|
adapt_qa_language = bool(
|
||||||
matched and _qa_requires_language_adaptation(question, matched.answer)
|
matched and _qa_requires_per_turn_rendering(question, matched.answer, history)
|
||||||
)
|
)
|
||||||
messages, reservation = [], None
|
messages, reservation = [], None
|
||||||
if matched and not adapt_qa_language:
|
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:
|
||||||
if matched:
|
if matched:
|
||||||
@@ -533,11 +975,21 @@ def _stream_reply(
|
|||||||
question,
|
question,
|
||||||
references,
|
references,
|
||||||
standard_answer=matched.answer,
|
standard_answer=matched.answer,
|
||||||
|
image_contexts=image_contexts,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
references = _search_knowledge(db, avatar.id, question)
|
retrieval_question = _image_retrieval_question(question, image_contexts)
|
||||||
source = "knowledge" if references else "qwen"
|
references = _search_knowledge(db, avatar.id, retrieval_question)
|
||||||
messages = _build_prompt(avatar, history, question, references)
|
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 = 0.0 if matched else min(
|
temperature = 0.0 if matched else min(
|
||||||
0.45 if references else 0.25,
|
0.45 if references else 0.25,
|
||||||
@@ -641,11 +1093,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"
|
||||||
@@ -660,8 +1163,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:
|
||||||
@@ -671,7 +1183,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
|
||||||
|
|
||||||
@@ -679,13 +1201,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
|
||||||
|
|||||||
@@ -1,8 +1,8 @@
|
|||||||
import os
|
import os
|
||||||
import json
|
import json
|
||||||
import logging
|
import shutil
|
||||||
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
from datetime import datetime, timezone
|
|
||||||
|
|
||||||
from fastapi import APIRouter, UploadFile, File, Depends, Header, HTTPException
|
from fastapi import APIRouter, UploadFile, File, Depends, Header, HTTPException
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
@@ -12,16 +12,19 @@ from database import get_db
|
|||||||
from models import KnowledgeDoc, QAPair, KnowledgeChunk, Avatar, User
|
from models import KnowledgeDoc, QAPair, KnowledgeChunk, Avatar, User
|
||||||
from responses import ok, fail
|
from responses import ok, fail
|
||||||
import embeddings
|
import embeddings
|
||||||
|
from services.knowledge_vectorizer import knowledge_vectorizer
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
|
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||||
UPLOAD_DIR = os.path.abspath(os.getenv("UPLOAD_DIR", os.path.join(BASE_DIR, "uploads")))
|
UPLOAD_DIR = os.path.abspath(os.getenv("UPLOAD_DIR", os.path.join(BASE_DIR, "uploads")))
|
||||||
os.makedirs(UPLOAD_DIR, exist_ok=True)
|
os.makedirs(UPLOAD_DIR, exist_ok=True)
|
||||||
|
|
||||||
ALLOWED_EXT = {".md", ".txt", ".pdf", ".doc", ".docx", ".xlsx"}
|
ALLOWED_EXT = {".md", ".txt", ".pdf", ".doc", ".docx", ".xlsx"}
|
||||||
MAX_UPLOAD_BYTES = 10 * 1024 * 1024
|
MAX_UPLOAD_BYTES = 50 * 1024 * 1024
|
||||||
|
UPLOAD_CHUNK_BYTES = 1024 * 1024
|
||||||
|
MULTIPART_CHUNK_BYTES = 5 * 1024 * 1024
|
||||||
|
MULTIPART_ROOT = ".multipart"
|
||||||
|
MULTIPART_TTL_SECONDS = 24 * 60 * 60
|
||||||
|
|
||||||
|
|
||||||
class QAIn(BaseModel):
|
class QAIn(BaseModel):
|
||||||
@@ -34,6 +37,74 @@ class EnabledIn(BaseModel):
|
|||||||
enabled: bool = True
|
enabled: bool = True
|
||||||
|
|
||||||
|
|
||||||
|
class MultipartUploadIn(BaseModel):
|
||||||
|
filename: str
|
||||||
|
fileSize: int
|
||||||
|
totalChunks: int
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_document(filename: str, file_size: int):
|
||||||
|
ext = os.path.splitext(filename or "")[1].lower()
|
||||||
|
if ext not in ALLOWED_EXT:
|
||||||
|
return None, f"不支持的文件类型:{ext or '空'},仅支持 md/txt/pdf/doc/docx/xlsx"
|
||||||
|
if file_size <= 0:
|
||||||
|
return None, "文件内容不能为空"
|
||||||
|
if file_size > MAX_UPLOAD_BYTES:
|
||||||
|
return None, "文件不能超过 50MB"
|
||||||
|
return ext, ""
|
||||||
|
|
||||||
|
|
||||||
|
def _multipart_dir(avatar_id: str, upload_id: str) -> str:
|
||||||
|
safe_avatar_id = os.path.basename(avatar_id)
|
||||||
|
safe_upload_id = os.path.basename(upload_id)
|
||||||
|
if (
|
||||||
|
safe_avatar_id != avatar_id
|
||||||
|
or safe_upload_id != upload_id
|
||||||
|
or len(upload_id) != 32
|
||||||
|
or any(character not in "0123456789abcdef" for character in upload_id)
|
||||||
|
):
|
||||||
|
raise HTTPException(status_code=400, detail="上传标识无效")
|
||||||
|
return os.path.join(UPLOAD_DIR, MULTIPART_ROOT, safe_avatar_id, safe_upload_id)
|
||||||
|
|
||||||
|
|
||||||
|
def _purge_stale_multipart_uploads(avatar_id: str):
|
||||||
|
avatar_upload_root = os.path.join(UPLOAD_DIR, MULTIPART_ROOT, os.path.basename(avatar_id))
|
||||||
|
if not os.path.isdir(avatar_upload_root):
|
||||||
|
return
|
||||||
|
cutoff = time.time() - MULTIPART_TTL_SECONDS
|
||||||
|
for entry in os.scandir(avatar_upload_root):
|
||||||
|
if entry.is_dir(follow_symlinks=False) and entry.stat(follow_symlinks=False).st_mtime < cutoff:
|
||||||
|
shutil.rmtree(entry.path, ignore_errors=True)
|
||||||
|
|
||||||
|
|
||||||
|
def _read_multipart_metadata(avatar_id: str, upload_id: str) -> tuple[str, dict]:
|
||||||
|
upload_dir = _multipart_dir(avatar_id, upload_id)
|
||||||
|
metadata_path = os.path.join(upload_dir, "metadata.json")
|
||||||
|
if not os.path.isfile(metadata_path):
|
||||||
|
raise HTTPException(status_code=404, detail="上传任务不存在或已过期")
|
||||||
|
with open(metadata_path, "r", encoding="utf-8") as stream:
|
||||||
|
return upload_dir, json.load(stream)
|
||||||
|
|
||||||
|
|
||||||
|
def _create_knowledge_doc(db: Session, avatar_id: str, filename: str, ext: str, file_size: int, stored: str):
|
||||||
|
doc = KnowledgeDoc(
|
||||||
|
id=uuid.uuid4().hex,
|
||||||
|
avatar_id=avatar_id,
|
||||||
|
filename=filename,
|
||||||
|
file_type=ext.lstrip("."),
|
||||||
|
file_size=file_size,
|
||||||
|
file_url=f"/api/files/{avatar_id}/{stored}",
|
||||||
|
status="parsing",
|
||||||
|
index_stage="queued",
|
||||||
|
index_progress=0,
|
||||||
|
)
|
||||||
|
db.add(doc)
|
||||||
|
db.commit()
|
||||||
|
db.refresh(doc)
|
||||||
|
knowledge_vectorizer.enqueue(doc.id)
|
||||||
|
return doc
|
||||||
|
|
||||||
|
|
||||||
def _doc_payload(doc: KnowledgeDoc) -> dict:
|
def _doc_payload(doc: KnowledgeDoc) -> dict:
|
||||||
payload = doc.to_dict()
|
payload = doc.to_dict()
|
||||||
stored_name = os.path.basename(doc.file_url or "")
|
stored_name = os.path.basename(doc.file_url or "")
|
||||||
@@ -71,84 +142,178 @@ def list_docs(avatar_id: str, authorization: str = Header(None), db: Session = D
|
|||||||
.order_by(KnowledgeDoc.created_at.desc())
|
.order_by(KnowledgeDoc.created_at.desc())
|
||||||
.all()
|
.all()
|
||||||
)
|
)
|
||||||
# Older synchronous uploads could be interrupted after persisting "parsing".
|
|
||||||
# New uploads are committed only after indexing finishes, so these rows are stale.
|
|
||||||
stale_docs = [doc for doc in docs if doc.status == "parsing"]
|
|
||||||
if stale_docs:
|
|
||||||
for doc in stale_docs:
|
|
||||||
doc.status = "failed"
|
|
||||||
doc.vectorized = False
|
|
||||||
doc.chunk_count = 0
|
|
||||||
db.commit()
|
|
||||||
return ok([_doc_payload(d) for d in docs])
|
return ok([_doc_payload(d) for d in docs])
|
||||||
|
|
||||||
|
|
||||||
@router.post("/avatar/{avatar_id}/knowledge/docs")
|
@router.post("/avatar/{avatar_id}/knowledge/docs")
|
||||||
async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization: str = Header(None), db: Session = Depends(get_db)):
|
async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization: str = Header(None), db: Session = Depends(get_db)):
|
||||||
_require_owned_avatar(db, avatar_id, authorization)
|
_require_owned_avatar(db, avatar_id, authorization)
|
||||||
ext = os.path.splitext(file.filename or "")[1].lower()
|
ext, validation_error = _validate_document(file.filename or "", 1)
|
||||||
if ext not in ALLOWED_EXT:
|
if validation_error:
|
||||||
return fail(f"不支持的文件类型:{ext or '空'},仅支持 md/txt/pdf/doc/docx/xlsx", code=400)
|
return fail(validation_error, code=400)
|
||||||
avatar_dir = os.path.join(UPLOAD_DIR, avatar_id)
|
avatar_dir = os.path.join(UPLOAD_DIR, avatar_id)
|
||||||
os.makedirs(avatar_dir, exist_ok=True)
|
os.makedirs(avatar_dir, exist_ok=True)
|
||||||
stored = f"{uuid.uuid4().hex}{ext}"
|
stored = f"{uuid.uuid4().hex}{ext}"
|
||||||
path = os.path.join(avatar_dir, stored)
|
path = os.path.join(avatar_dir, stored)
|
||||||
content = await file.read()
|
file_size = 0
|
||||||
if len(content) > MAX_UPLOAD_BYTES:
|
|
||||||
return fail("文件不能超过 10MB", code=400)
|
|
||||||
with open(path, "wb") as f:
|
|
||||||
f.write(content)
|
|
||||||
doc = KnowledgeDoc(
|
|
||||||
id=uuid.uuid4().hex,
|
|
||||||
avatar_id=avatar_id,
|
|
||||||
filename=file.filename,
|
|
||||||
file_type=ext.lstrip("."),
|
|
||||||
file_size=len(content),
|
|
||||||
file_url=f"/api/files/{avatar_id}/{stored}",
|
|
||||||
status="parsing",
|
|
||||||
)
|
|
||||||
|
|
||||||
# Complete extraction and embedding before the first database commit so a
|
|
||||||
# process restart cannot leave a permanent "parsing" row behind.
|
|
||||||
try:
|
try:
|
||||||
text = embeddings.extract_text(path, ext)
|
# Stream large files to disk so a 100MB upload does not occupy 100MB RAM.
|
||||||
chunks = embeddings.chunk_text(text)
|
with open(path, "wb") as f:
|
||||||
if not chunks:
|
while chunk := await file.read(UPLOAD_CHUNK_BYTES):
|
||||||
raise ValueError("文档没有可建立索引的文字内容")
|
file_size += len(chunk)
|
||||||
vectors = embeddings.embed(chunks)
|
if file_size > MAX_UPLOAD_BYTES:
|
||||||
if len(vectors) != len(chunks):
|
raise ValueError("文件不能超过 50MB")
|
||||||
raise ValueError("向量服务返回数量与文档分段不一致")
|
f.write(chunk)
|
||||||
doc.vectorized = True
|
except ValueError as exc:
|
||||||
doc.embedding_model = embeddings.MODEL
|
if os.path.exists(path):
|
||||||
doc.chunk_count = len(chunks)
|
os.remove(path)
|
||||||
doc.vectorized_at = datetime.now(timezone.utc)
|
return fail(str(exc), code=400)
|
||||||
doc.status = "ready"
|
if file_size == 0:
|
||||||
db.add(doc)
|
if os.path.exists(path):
|
||||||
for i, (chunk, vector) in enumerate(zip(chunks, vectors)):
|
os.remove(path)
|
||||||
db.add(
|
return fail("文件内容不能为空", code=400)
|
||||||
KnowledgeChunk(
|
|
||||||
doc_id=doc.id,
|
|
||||||
avatar_id=avatar_id,
|
|
||||||
content=chunk,
|
|
||||||
vector=json.dumps(vector),
|
|
||||||
chunk_index=i,
|
|
||||||
embedding_model=embeddings.MODEL,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
db.commit()
|
|
||||||
db.refresh(doc)
|
|
||||||
except Exception as exc:
|
|
||||||
db.rollback()
|
|
||||||
doc.status = "failed"
|
|
||||||
doc.vectorized = False
|
|
||||||
doc.embedding_model = ""
|
|
||||||
doc.chunk_count = 0
|
|
||||||
doc.vectorized_at = None
|
|
||||||
db.add(doc)
|
|
||||||
db.commit()
|
|
||||||
db.refresh(doc)
|
|
||||||
logger.exception("knowledge vectorization failed for %s: %s", doc.id, exc)
|
|
||||||
|
|
||||||
|
doc = _create_knowledge_doc(db, avatar_id, file.filename or stored, ext, file_size, stored)
|
||||||
|
return ok(_doc_payload(doc))
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/avatar/{avatar_id}/knowledge/uploads")
|
||||||
|
def create_multipart_upload(
|
||||||
|
avatar_id: str,
|
||||||
|
body: MultipartUploadIn,
|
||||||
|
authorization: str = Header(None),
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
):
|
||||||
|
_require_owned_avatar(db, avatar_id, authorization)
|
||||||
|
ext, validation_error = _validate_document(body.filename, body.fileSize)
|
||||||
|
if validation_error:
|
||||||
|
return fail(validation_error, code=400)
|
||||||
|
expected_chunks = (body.fileSize + MULTIPART_CHUNK_BYTES - 1) // MULTIPART_CHUNK_BYTES
|
||||||
|
if body.totalChunks != expected_chunks:
|
||||||
|
return fail("文件分片数量不正确", code=400)
|
||||||
|
|
||||||
|
_purge_stale_multipart_uploads(avatar_id)
|
||||||
|
upload_id = uuid.uuid4().hex
|
||||||
|
upload_dir = _multipart_dir(avatar_id, upload_id)
|
||||||
|
os.makedirs(upload_dir, exist_ok=False)
|
||||||
|
metadata = {
|
||||||
|
"filename": body.filename,
|
||||||
|
"fileSize": body.fileSize,
|
||||||
|
"totalChunks": body.totalChunks,
|
||||||
|
"extension": ext,
|
||||||
|
}
|
||||||
|
with open(os.path.join(upload_dir, "metadata.json"), "w", encoding="utf-8") as stream:
|
||||||
|
json.dump(metadata, stream, ensure_ascii=False)
|
||||||
|
return ok({"uploadId": upload_id, "chunkSize": MULTIPART_CHUNK_BYTES})
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/avatar/{avatar_id}/knowledge/uploads/{upload_id}/chunks/{chunk_index}")
|
||||||
|
async def upload_multipart_chunk(
|
||||||
|
avatar_id: str,
|
||||||
|
upload_id: str,
|
||||||
|
chunk_index: int,
|
||||||
|
file: UploadFile = File(...),
|
||||||
|
authorization: str = Header(None),
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
):
|
||||||
|
_require_owned_avatar(db, avatar_id, authorization)
|
||||||
|
upload_dir, metadata = _read_multipart_metadata(avatar_id, upload_id)
|
||||||
|
total_chunks = int(metadata["totalChunks"])
|
||||||
|
if chunk_index < 0 or chunk_index >= total_chunks:
|
||||||
|
return fail("文件分片序号不正确", code=400)
|
||||||
|
|
||||||
|
expected_size = min(
|
||||||
|
MULTIPART_CHUNK_BYTES,
|
||||||
|
int(metadata["fileSize"]) - chunk_index * MULTIPART_CHUNK_BYTES,
|
||||||
|
)
|
||||||
|
part_path = os.path.join(upload_dir, f"{chunk_index}.part")
|
||||||
|
temporary_path = f"{part_path}.uploading"
|
||||||
|
received = 0
|
||||||
|
try:
|
||||||
|
with open(temporary_path, "wb") as stream:
|
||||||
|
while chunk := await file.read(UPLOAD_CHUNK_BYTES):
|
||||||
|
received += len(chunk)
|
||||||
|
if received > expected_size:
|
||||||
|
raise ValueError("文件分片大小不正确")
|
||||||
|
stream.write(chunk)
|
||||||
|
if received != expected_size:
|
||||||
|
raise ValueError("文件分片大小不正确")
|
||||||
|
os.replace(temporary_path, part_path)
|
||||||
|
except ValueError as exc:
|
||||||
|
if os.path.exists(temporary_path):
|
||||||
|
os.remove(temporary_path)
|
||||||
|
return fail(str(exc), code=400)
|
||||||
|
return ok({"chunkIndex": chunk_index, "uploadedBytes": received})
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/avatar/{avatar_id}/knowledge/uploads/{upload_id}/complete")
|
||||||
|
def complete_multipart_upload(
|
||||||
|
avatar_id: str,
|
||||||
|
upload_id: str,
|
||||||
|
authorization: str = Header(None),
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
):
|
||||||
|
_require_owned_avatar(db, avatar_id, authorization)
|
||||||
|
upload_dir, metadata = _read_multipart_metadata(avatar_id, upload_id)
|
||||||
|
total_chunks = int(metadata["totalChunks"])
|
||||||
|
part_paths = [os.path.join(upload_dir, f"{index}.part") for index in range(total_chunks)]
|
||||||
|
if not all(os.path.isfile(path) for path in part_paths):
|
||||||
|
return fail("文件分片尚未上传完整", code=400)
|
||||||
|
if sum(os.path.getsize(path) for path in part_paths) != int(metadata["fileSize"]):
|
||||||
|
return fail("文件分片总大小不正确", code=400)
|
||||||
|
|
||||||
|
avatar_dir = os.path.join(UPLOAD_DIR, avatar_id)
|
||||||
|
os.makedirs(avatar_dir, exist_ok=True)
|
||||||
|
stored = f"{uuid.uuid4().hex}{metadata['extension']}"
|
||||||
|
final_path = os.path.join(avatar_dir, stored)
|
||||||
|
temporary_path = f"{final_path}.assembling"
|
||||||
|
try:
|
||||||
|
with open(temporary_path, "wb") as output:
|
||||||
|
for part_path in part_paths:
|
||||||
|
with open(part_path, "rb") as source:
|
||||||
|
shutil.copyfileobj(source, output, UPLOAD_CHUNK_BYTES)
|
||||||
|
os.replace(temporary_path, final_path)
|
||||||
|
doc = _create_knowledge_doc(
|
||||||
|
db,
|
||||||
|
avatar_id,
|
||||||
|
metadata["filename"],
|
||||||
|
metadata["extension"],
|
||||||
|
int(metadata["fileSize"]),
|
||||||
|
stored,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
if os.path.exists(temporary_path):
|
||||||
|
os.remove(temporary_path)
|
||||||
|
raise
|
||||||
|
shutil.rmtree(upload_dir, ignore_errors=True)
|
||||||
|
return ok(_doc_payload(doc))
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/avatar/{avatar_id}/knowledge/docs/{doc_id}/retry")
|
||||||
|
def retry_doc(avatar_id: str, doc_id: str, authorization: str = Header(None), db: Session = Depends(get_db)):
|
||||||
|
_require_owned_avatar(db, avatar_id, authorization)
|
||||||
|
doc = db.query(KnowledgeDoc).filter(
|
||||||
|
KnowledgeDoc.id == doc_id, KnowledgeDoc.avatar_id == avatar_id
|
||||||
|
).first()
|
||||||
|
if not doc:
|
||||||
|
return fail("文档不存在", code=404)
|
||||||
|
if doc.vectorized and doc.status == "ready":
|
||||||
|
return ok(_doc_payload(doc))
|
||||||
|
stored_name = os.path.basename(doc.file_url or "")
|
||||||
|
if not stored_name or not os.path.isfile(os.path.join(UPLOAD_DIR, avatar_id, stored_name)):
|
||||||
|
return fail("原文件不可用,请重新上传", code=400)
|
||||||
|
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == doc.id).delete()
|
||||||
|
doc.status = "parsing"
|
||||||
|
doc.vectorized = False
|
||||||
|
doc.embedding_model = ""
|
||||||
|
doc.chunk_count = 0
|
||||||
|
doc.vectorized_at = None
|
||||||
|
doc.error_message = ""
|
||||||
|
doc.index_stage = "queued"
|
||||||
|
doc.index_progress = 0
|
||||||
|
db.commit()
|
||||||
|
db.refresh(doc)
|
||||||
|
knowledge_vectorizer.enqueue(doc.id)
|
||||||
return ok(_doc_payload(doc))
|
return ok(_doc_payload(doc))
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -7,20 +7,43 @@ from datetime import datetime
|
|||||||
from decimal import Decimal, InvalidOperation, ROUND_HALF_UP
|
from decimal import Decimal, InvalidOperation, ROUND_HALF_UP
|
||||||
from urllib.parse import parse_qs
|
from urllib.parse import parse_qs
|
||||||
|
|
||||||
from fastapi import APIRouter, Body, Depends, Header, HTTPException, Request
|
from fastapi import APIRouter, Body, Depends, Header, HTTPException, Query, Request, Response
|
||||||
from sqlalchemy import func
|
from sqlalchemy import func
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from database import get_db
|
from database import get_db
|
||||||
from models import TokenAccount, TokenPaymentOrder, TokenPlan, TokenUsage, User
|
from models import (
|
||||||
|
InvoiceApplication,
|
||||||
|
PaymentRefund,
|
||||||
|
PaymentTransaction,
|
||||||
|
TokenAccount,
|
||||||
|
TokenPaymentOrder,
|
||||||
|
TokenPlan,
|
||||||
|
TokenUsage,
|
||||||
|
User,
|
||||||
|
)
|
||||||
from responses import fail, ok
|
from responses import fail, ok
|
||||||
from services.huihui_payment import HuihuiPaymentClient, HuihuiPaymentError
|
from services.huihui_payment import HuihuiPaymentClient, HuihuiPaymentError
|
||||||
from services.token_billing import DEFAULT_TOKEN_GRANT, get_or_create_account
|
from services.token_billing import get_or_create_account
|
||||||
|
from services.wechat_virtual_payment import (
|
||||||
|
PAYMENT_EVENTS as WECHAT_PAYMENT_EVENTS,
|
||||||
|
REFUND_EVENTS as WECHAT_REFUND_EVENTS,
|
||||||
|
WechatVirtualPaymentError,
|
||||||
|
build_payment_params as build_wechat_virtual_payment_params,
|
||||||
|
callback_value as wechat_callback_value,
|
||||||
|
exchange_code as exchange_wechat_code,
|
||||||
|
parse_callback_body as parse_wechat_callback_body,
|
||||||
|
product_id_for_plan,
|
||||||
|
query_order as query_wechat_virtual_order,
|
||||||
|
request_refund as request_wechat_virtual_refund,
|
||||||
|
verify_callback_signature as verify_wechat_callback_signature,
|
||||||
|
virtual_env as wechat_virtual_env,
|
||||||
|
)
|
||||||
|
|
||||||
router = APIRouter(tags=["Token"])
|
router = APIRouter(tags=["Token"])
|
||||||
|
|
||||||
PAYMENT_METHODS = {"wechat": "WECHAT", "alipay": "ALIPAY"}
|
PAYMENT_METHODS = {"wechat": "WECHAT", "alipay": "ALIPAY"}
|
||||||
PAYMENT_SCENES = {"APP", "LITE", "JSAPI"}
|
PAYMENT_SCENES = {"APP", "H5", "LITE", "JSAPI"}
|
||||||
SUCCESS_STATUSES = {"SUCCESS", "SUCCEEDED", "PAID", "COMPLETED", "TRADE_SUCCESS"}
|
SUCCESS_STATUSES = {"SUCCESS", "SUCCEEDED", "PAID", "COMPLETED", "TRADE_SUCCESS"}
|
||||||
FAILED_STATUSES = {"FAIL", "FAILED", "CLOSED", "CANCELLED", "CANCELED", "EXPIRED"}
|
FAILED_STATUSES = {"FAIL", "FAILED", "CLOSED", "CANCELLED", "CANCELED", "EXPIRED"}
|
||||||
|
|
||||||
@@ -35,6 +58,13 @@ def _require_user(authorization: str | None, db: Session) -> User:
|
|||||||
return user
|
return user
|
||||||
|
|
||||||
|
|
||||||
|
def _require_finance_admin(value: str | None):
|
||||||
|
expected = os.getenv("AVATAR_FINANCE_ADMIN_SECRET", "").strip()
|
||||||
|
provided = str(value or "").strip()
|
||||||
|
if len(expected) < 16 or not hmac.compare_digest(provided, expected):
|
||||||
|
raise HTTPException(status_code=403, detail="财务管理凭证无效")
|
||||||
|
|
||||||
|
|
||||||
def _payment_client() -> HuihuiPaymentClient:
|
def _payment_client() -> HuihuiPaymentClient:
|
||||||
return HuihuiPaymentClient({
|
return HuihuiPaymentClient({
|
||||||
"HUIHUI_PAYMENT_BASE_URL": os.getenv(
|
"HUIHUI_PAYMENT_BASE_URL": os.getenv(
|
||||||
@@ -70,6 +100,132 @@ def _payment_payload(order: TokenPaymentOrder, account: TokenAccount) -> dict:
|
|||||||
return {**order.to_dict(), "balance": account.balance}
|
return {**order.to_dict(), "balance": account.balance}
|
||||||
|
|
||||||
|
|
||||||
|
def _safe_event_summary(payload: dict) -> str:
|
||||||
|
"""Persist only reconciliation fields, never signatures, tokens or session keys."""
|
||||||
|
summary = {}
|
||||||
|
for key in (
|
||||||
|
"Event", "OutTradeNo", "OpenId", "Env", "MchOrderId", "MchRefundId",
|
||||||
|
"WxRefundId", "RefundFee", "RetCode", "RetMsg",
|
||||||
|
):
|
||||||
|
value = wechat_callback_value(payload, key)
|
||||||
|
if value not in (None, ""):
|
||||||
|
summary[key] = value
|
||||||
|
goods = wechat_callback_value(payload, "GoodsInfo")
|
||||||
|
if isinstance(goods, dict):
|
||||||
|
summary["GoodsInfo"] = {
|
||||||
|
key: goods.get(key)
|
||||||
|
for key in ("ProductId", "Quantity", "OrigPrice", "ActualPrice")
|
||||||
|
if goods.get(key) not in (None, "")
|
||||||
|
}
|
||||||
|
return json.dumps(summary, ensure_ascii=False, separators=(",", ":"))[:2000]
|
||||||
|
|
||||||
|
|
||||||
|
def _record_transaction(
|
||||||
|
db: Session,
|
||||||
|
*,
|
||||||
|
order: TokenPaymentOrder,
|
||||||
|
provider: str,
|
||||||
|
status: str,
|
||||||
|
amount_cents: int,
|
||||||
|
event_type: str = "payment",
|
||||||
|
transaction_no: str = "",
|
||||||
|
raw_summary: str = "",
|
||||||
|
):
|
||||||
|
if transaction_no:
|
||||||
|
duplicate = db.query(PaymentTransaction).filter(
|
||||||
|
PaymentTransaction.provider == provider,
|
||||||
|
PaymentTransaction.transaction_no == transaction_no,
|
||||||
|
PaymentTransaction.event_type == event_type,
|
||||||
|
).first()
|
||||||
|
if duplicate:
|
||||||
|
return duplicate
|
||||||
|
row = PaymentTransaction(
|
||||||
|
order_no=order.order_no,
|
||||||
|
provider=provider,
|
||||||
|
transaction_no=transaction_no,
|
||||||
|
event_type=event_type,
|
||||||
|
status=status,
|
||||||
|
amount_cents=amount_cents,
|
||||||
|
raw_summary=raw_summary,
|
||||||
|
)
|
||||||
|
db.add(row)
|
||||||
|
return row
|
||||||
|
|
||||||
|
|
||||||
|
def _settle_paid_order(
|
||||||
|
db: Session,
|
||||||
|
order: TokenPaymentOrder,
|
||||||
|
*,
|
||||||
|
provider_status: str,
|
||||||
|
transaction_no: str = "",
|
||||||
|
raw_summary: str = "",
|
||||||
|
) -> bool:
|
||||||
|
if order.status in {"paid", "refunded"}:
|
||||||
|
return False
|
||||||
|
updated = db.query(TokenPaymentOrder).filter(
|
||||||
|
TokenPaymentOrder.id == order.id,
|
||||||
|
TokenPaymentOrder.status.in_(["pending", "failed", "closed"]),
|
||||||
|
).update({
|
||||||
|
TokenPaymentOrder.status: "paid",
|
||||||
|
TokenPaymentOrder.provider_status: provider_status,
|
||||||
|
TokenPaymentOrder.paid_at: datetime.utcnow(),
|
||||||
|
TokenPaymentOrder.failure_reason: "",
|
||||||
|
}, synchronize_session=False)
|
||||||
|
if not updated:
|
||||||
|
return False
|
||||||
|
account = get_or_create_account(db, order.user_id)
|
||||||
|
account.balance = int(account.balance or 0) + order.points_amount
|
||||||
|
account.total_granted = int(account.total_granted or 0) + order.points_amount
|
||||||
|
_record_transaction(
|
||||||
|
db,
|
||||||
|
order=order,
|
||||||
|
provider=order.provider,
|
||||||
|
status="paid",
|
||||||
|
amount_cents=order.price_cents,
|
||||||
|
transaction_no=transaction_no,
|
||||||
|
raw_summary=raw_summary,
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
def _complete_refund(
|
||||||
|
db: Session,
|
||||||
|
order: TokenPaymentOrder,
|
||||||
|
refund: PaymentRefund,
|
||||||
|
*,
|
||||||
|
provider_refund_no: str = "",
|
||||||
|
failure_reason: str = "",
|
||||||
|
):
|
||||||
|
if failure_reason:
|
||||||
|
refund.status = "failed"
|
||||||
|
refund.failure_reason = failure_reason[:500]
|
||||||
|
order.refund_status = "failed"
|
||||||
|
return
|
||||||
|
if refund.status == "succeeded":
|
||||||
|
return
|
||||||
|
account = get_or_create_account(db, order.user_id)
|
||||||
|
# Provider-confirmed refunds must claw back the full grant. A negative
|
||||||
|
# balance records consumed refunded points and blocks further usage.
|
||||||
|
account.balance = int(account.balance or 0) - int(refund.points_amount or 0)
|
||||||
|
account.total_granted = max(0, int(account.total_granted or 0) - int(refund.points_amount or 0))
|
||||||
|
refund.status = "succeeded"
|
||||||
|
refund.provider_refund_no = provider_refund_no[:128]
|
||||||
|
refund.failure_reason = ""
|
||||||
|
refund.completed_at = datetime.utcnow()
|
||||||
|
order.status = "refunded"
|
||||||
|
order.refund_status = "succeeded"
|
||||||
|
order.refunded_at = datetime.utcnow()
|
||||||
|
_record_transaction(
|
||||||
|
db,
|
||||||
|
order=order,
|
||||||
|
provider=order.provider,
|
||||||
|
status="succeeded",
|
||||||
|
amount_cents=refund.amount_cents,
|
||||||
|
event_type="refund",
|
||||||
|
transaction_no=provider_refund_no or refund.refund_no,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _nested_payload(value):
|
def _nested_payload(value):
|
||||||
if isinstance(value, str):
|
if isinstance(value, str):
|
||||||
text = value.strip()
|
text = value.strip()
|
||||||
@@ -159,8 +315,11 @@ def charge(payload: dict = Body(...), authorization: str = Header(None), db: Ses
|
|||||||
pay_way = str(payload.get("payScene") or "APP").upper()
|
pay_way = str(payload.get("payScene") or "APP").upper()
|
||||||
if pay_way not in PAYMENT_SCENES:
|
if pay_way not in PAYMENT_SCENES:
|
||||||
return fail("当前支付场景不受支持", 400)
|
return fail("当前支付场景不受支持", 400)
|
||||||
|
if pay_way == "LITE" and payment_method != "wechat":
|
||||||
|
return fail("微信小程序虚拟支付仅支持微信支付", 400)
|
||||||
|
|
||||||
cents = _price_cents(plan.price)
|
cents = _price_cents(plan.price)
|
||||||
|
provider = "wechat_virtual" if pay_way == "LITE" else "huihui"
|
||||||
order = TokenPaymentOrder(
|
order = TokenPaymentOrder(
|
||||||
order_no=f"AV{datetime.utcnow().strftime('%Y%m%d%H%M%S')}{uuid.uuid4().hex[:12].upper()}",
|
order_no=f"AV{datetime.utcnow().strftime('%Y%m%d%H%M%S')}{uuid.uuid4().hex[:12].upper()}",
|
||||||
user_id=user.id,
|
user_id=user.id,
|
||||||
@@ -171,10 +330,35 @@ def charge(payload: dict = Body(...), authorization: str = Header(None), db: Ses
|
|||||||
points_amount=plan.amount,
|
points_amount=plan.amount,
|
||||||
price_cents=cents,
|
price_cents=cents,
|
||||||
status="pending",
|
status="pending",
|
||||||
|
provider=provider,
|
||||||
)
|
)
|
||||||
db.add(order)
|
db.add(order)
|
||||||
db.commit()
|
db.commit()
|
||||||
|
|
||||||
|
if provider == "wechat_virtual":
|
||||||
|
if not user.wechat_mp_openid or not user.wechat_mp_session_key:
|
||||||
|
order.status = "failed"
|
||||||
|
order.failure_reason = "微信小程序登录态尚未准备好,请重新进入支付页"
|
||||||
|
db.commit()
|
||||||
|
return fail(order.failure_reason, 409)
|
||||||
|
try:
|
||||||
|
result = build_wechat_virtual_payment_params(
|
||||||
|
order=order,
|
||||||
|
plan=plan,
|
||||||
|
session_key=user.wechat_mp_session_key,
|
||||||
|
)
|
||||||
|
except WechatVirtualPaymentError as exc:
|
||||||
|
order.status = "failed"
|
||||||
|
order.failure_reason = str(exc)[:500]
|
||||||
|
db.commit()
|
||||||
|
return fail(str(exc), 503)
|
||||||
|
order.provider_order_id = order.order_no
|
||||||
|
order.provider_order_no = order.order_no
|
||||||
|
order.provider_status = "CREATED"
|
||||||
|
order.pay_message = json.dumps(result, ensure_ascii=False, separators=(",", ":"))
|
||||||
|
db.commit()
|
||||||
|
return ok(_payment_payload(order, get_or_create_account(db, user.id)))
|
||||||
|
|
||||||
try:
|
try:
|
||||||
callback_url = _callback_url(order.order_no)
|
callback_url = _callback_url(order.order_no)
|
||||||
except HuihuiPaymentError as exc:
|
except HuihuiPaymentError as exc:
|
||||||
@@ -229,9 +413,267 @@ def payment_status(order_id: str, authorization: str = Header(None), db: Session
|
|||||||
).first()
|
).first()
|
||||||
if not order:
|
if not order:
|
||||||
return fail("支付订单不存在", 404)
|
return fail("支付订单不存在", 404)
|
||||||
|
if order.provider == "wechat_virtual" and order.status == "pending" and user.wechat_mp_openid:
|
||||||
|
try:
|
||||||
|
provider_data = query_wechat_virtual_order(
|
||||||
|
openid=user.wechat_mp_openid,
|
||||||
|
order_no=order.order_no,
|
||||||
|
)
|
||||||
|
provider_order = provider_data.get("order") or {}
|
||||||
|
provider_status = int(provider_order.get("status", 0) or 0)
|
||||||
|
paid_cents = int(provider_order.get("paid_fee") or provider_order.get("order_fee") or 0)
|
||||||
|
order.provider_status = str(provider_status)
|
||||||
|
if provider_status in {2, 3, 4} and paid_cents == order.price_cents:
|
||||||
|
_settle_paid_order(
|
||||||
|
db,
|
||||||
|
order,
|
||||||
|
provider_status=f"XPAY_{provider_status}",
|
||||||
|
transaction_no=str(
|
||||||
|
provider_order.get("wxpay_order_id")
|
||||||
|
or provider_order.get("channel_order_id")
|
||||||
|
or order.order_no
|
||||||
|
),
|
||||||
|
)
|
||||||
|
elif provider_status == 6:
|
||||||
|
order.status = "failed"
|
||||||
|
order.failure_reason = "微信虚拟支付订单已关闭"
|
||||||
|
db.commit()
|
||||||
|
db.refresh(order)
|
||||||
|
except WechatVirtualPaymentError:
|
||||||
|
# 回调仍是首选确认路径;短暂查询失败不覆盖订单状态。
|
||||||
|
pass
|
||||||
return ok(_payment_payload(order, get_or_create_account(db, user.id)))
|
return ok(_payment_payload(order, get_or_create_account(db, user.id)))
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/token/wechat/session")
|
||||||
|
def bind_wechat_session(
|
||||||
|
payload: dict = Body(...),
|
||||||
|
authorization: str = Header(None),
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
):
|
||||||
|
user = _require_user(authorization, db)
|
||||||
|
code = str(payload.get("code") or "").strip()
|
||||||
|
if not code or len(code) > 256:
|
||||||
|
return fail("微信登录凭证无效", 400)
|
||||||
|
try:
|
||||||
|
session = exchange_wechat_code(code)
|
||||||
|
except WechatVirtualPaymentError as exc:
|
||||||
|
return fail(str(exc), 502)
|
||||||
|
|
||||||
|
conflict = db.query(User).filter(
|
||||||
|
User.wechat_mp_openid == session["openid"],
|
||||||
|
User.id != user.id,
|
||||||
|
).first()
|
||||||
|
if conflict:
|
||||||
|
return fail("该微信账号已绑定其他会会账号", 409)
|
||||||
|
user.wechat_mp_openid = session["openid"]
|
||||||
|
user.wechat_mp_session_key = session["session_key"]
|
||||||
|
db.commit()
|
||||||
|
return ok({"ready": True})
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/token/orders")
|
||||||
|
def list_user_orders(
|
||||||
|
page: int = Query(1, ge=1),
|
||||||
|
page_size: int = Query(20, ge=1, le=100),
|
||||||
|
authorization: str = Header(None),
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
):
|
||||||
|
user = _require_user(authorization, db)
|
||||||
|
query = db.query(TokenPaymentOrder).filter(TokenPaymentOrder.user_id == user.id)
|
||||||
|
total = query.count()
|
||||||
|
orders = query.order_by(TokenPaymentOrder.created_at.desc()).offset((page - 1) * page_size).limit(page_size).all()
|
||||||
|
invoice_by_order = {
|
||||||
|
item.order_no: item.to_dict()
|
||||||
|
for item in db.query(InvoiceApplication).filter(
|
||||||
|
InvoiceApplication.order_no.in_([order.order_no for order in orders])
|
||||||
|
).all()
|
||||||
|
} if orders else {}
|
||||||
|
return ok({
|
||||||
|
"total": total,
|
||||||
|
"page": page,
|
||||||
|
"pageSize": page_size,
|
||||||
|
"items": [
|
||||||
|
{**_payment_payload(order, get_or_create_account(db, user.id)), "invoice": invoice_by_order.get(order.order_no)}
|
||||||
|
for order in orders
|
||||||
|
],
|
||||||
|
})
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/token/orders/{order_no}/invoice")
|
||||||
|
def apply_invoice(
|
||||||
|
order_no: str,
|
||||||
|
payload: dict = Body(...),
|
||||||
|
authorization: str = Header(None),
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
):
|
||||||
|
user = _require_user(authorization, db)
|
||||||
|
order = db.query(TokenPaymentOrder).filter(
|
||||||
|
TokenPaymentOrder.order_no == order_no,
|
||||||
|
TokenPaymentOrder.user_id == user.id,
|
||||||
|
).first()
|
||||||
|
if not order:
|
||||||
|
return fail("订单不存在", 404)
|
||||||
|
if order.status != "paid" or order.refund_status not in {"", "none"}:
|
||||||
|
return fail("只有已支付且未退款的订单可以申请发票", 409)
|
||||||
|
title = str(payload.get("title") or "").strip()
|
||||||
|
invoice_type = str(payload.get("invoiceType") or "personal").strip().lower()
|
||||||
|
tax_number = str(payload.get("taxNumber") or "").strip().upper()
|
||||||
|
email = str(payload.get("email") or "").strip()
|
||||||
|
if not title or len(title) > 120:
|
||||||
|
return fail("请填写正确的发票抬头", 400)
|
||||||
|
if invoice_type not in {"personal", "company"}:
|
||||||
|
return fail("发票类型不正确", 400)
|
||||||
|
if invoice_type == "company" and (len(tax_number) < 15 or len(tax_number) > 20):
|
||||||
|
return fail("请填写正确的企业税号", 400)
|
||||||
|
if email and ("@" not in email or len(email) > 160):
|
||||||
|
return fail("请填写正确的接收邮箱", 400)
|
||||||
|
|
||||||
|
invoice = db.query(InvoiceApplication).filter(InvoiceApplication.order_no == order_no).first()
|
||||||
|
if invoice and invoice.status not in {"rejected", "cancelled"}:
|
||||||
|
return fail("该订单已申请发票", 409)
|
||||||
|
if invoice is None:
|
||||||
|
invoice = InvoiceApplication(order_no=order_no, user_id=user.id, amount_cents=order.price_cents)
|
||||||
|
db.add(invoice)
|
||||||
|
invoice.title = title
|
||||||
|
invoice.invoice_type = invoice_type
|
||||||
|
invoice.tax_number = tax_number if invoice_type == "company" else ""
|
||||||
|
invoice.email = email
|
||||||
|
invoice.status = "pending"
|
||||||
|
invoice.remark = ""
|
||||||
|
db.commit()
|
||||||
|
db.refresh(invoice)
|
||||||
|
return ok(invoice.to_dict())
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/token/admin/orders/{order_no}/refund")
|
||||||
|
def admin_request_refund(
|
||||||
|
order_no: str,
|
||||||
|
payload: dict = Body(...),
|
||||||
|
finance_key: str = Header(None, alias="X-Avatar-Finance-Key"),
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
):
|
||||||
|
_require_finance_admin(finance_key)
|
||||||
|
order = db.query(TokenPaymentOrder).filter(TokenPaymentOrder.order_no == order_no).first()
|
||||||
|
if not order:
|
||||||
|
return fail("订单不存在", 404)
|
||||||
|
if order.status != "paid" or order.refund_status not in {"", "none", "failed"}:
|
||||||
|
return fail("该订单当前不可退款", 409)
|
||||||
|
account = get_or_create_account(db, order.user_id)
|
||||||
|
if int(account.balance or 0) < int(order.points_amount or 0):
|
||||||
|
return fail("该订单发放的积分已使用,不能执行全额退款", 409)
|
||||||
|
invoice = db.query(InvoiceApplication).filter(InvoiceApplication.order_no == order.order_no).first()
|
||||||
|
if invoice and invoice.status == "issued":
|
||||||
|
return fail("该订单发票已开具,请先完成红冲再退款", 409)
|
||||||
|
reason = str(payload.get("reason") or "后台退款").strip()
|
||||||
|
if not reason or len(reason) > 200:
|
||||||
|
return fail("请填写 200 字以内的退款原因", 400)
|
||||||
|
|
||||||
|
refund = PaymentRefund(
|
||||||
|
refund_no=f"RF{datetime.utcnow().strftime('%Y%m%d%H%M%S')}{uuid.uuid4().hex[:10].upper()}",
|
||||||
|
order_no=order.order_no,
|
||||||
|
amount_cents=order.price_cents,
|
||||||
|
points_amount=order.points_amount,
|
||||||
|
reason=reason,
|
||||||
|
status="processing",
|
||||||
|
requested_by=str(payload.get("operator") or "admin")[:80],
|
||||||
|
)
|
||||||
|
claimed = db.query(TokenPaymentOrder).filter(
|
||||||
|
TokenPaymentOrder.id == order.id,
|
||||||
|
TokenPaymentOrder.status == "paid",
|
||||||
|
TokenPaymentOrder.refund_status.in_(["", "none", "failed"]),
|
||||||
|
).update({TokenPaymentOrder.refund_status: "processing"}, synchronize_session=False)
|
||||||
|
if not claimed:
|
||||||
|
db.rollback()
|
||||||
|
return fail("该订单已有退款任务正在处理", 409)
|
||||||
|
db.add(refund)
|
||||||
|
if invoice and invoice.status == "pending":
|
||||||
|
invoice.status = "cancelled"
|
||||||
|
invoice.remark = "订单已申请退款,发票申请自动取消"
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
user = db.query(User).filter(User.id == order.user_id).first()
|
||||||
|
try:
|
||||||
|
if order.provider == "wechat_virtual":
|
||||||
|
if not user or not user.wechat_mp_openid:
|
||||||
|
raise WechatVirtualPaymentError("订单缺少微信 OpenID,无法退款")
|
||||||
|
provider_result = request_wechat_virtual_refund(
|
||||||
|
openid=user.wechat_mp_openid,
|
||||||
|
order_no=order.order_no,
|
||||||
|
refund_no=refund.refund_no,
|
||||||
|
amount_cents=refund.amount_cents,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
provider_result = _payment_client().request_refund(
|
||||||
|
huihui_token=user.huihui_token if user else "",
|
||||||
|
huihui_user_id=user.huihui_user_id if user else "",
|
||||||
|
order_no=order.order_no,
|
||||||
|
refund_no=refund.refund_no,
|
||||||
|
amount=f"{refund.amount_cents / 100:.2f}",
|
||||||
|
reason=reason,
|
||||||
|
)
|
||||||
|
except (WechatVirtualPaymentError, HuihuiPaymentError) as exc:
|
||||||
|
_complete_refund(db, order, refund, failure_reason=str(exc))
|
||||||
|
db.commit()
|
||||||
|
return fail(str(exc), 502)
|
||||||
|
|
||||||
|
provider_status = str(
|
||||||
|
provider_result.get("status")
|
||||||
|
or provider_result.get("refundStatus")
|
||||||
|
or provider_result.get("result")
|
||||||
|
or "PROCESSING"
|
||||||
|
).upper()
|
||||||
|
provider_refund_no = str(
|
||||||
|
provider_result.get("refundNo")
|
||||||
|
or provider_result.get("refundId")
|
||||||
|
or provider_result.get("wx_refund_id")
|
||||||
|
or ""
|
||||||
|
)
|
||||||
|
refund.provider_refund_no = provider_refund_no[:128]
|
||||||
|
if provider_status in {"SUCCESS", "SUCCEEDED", "REFUNDED", "COMPLETED"}:
|
||||||
|
_complete_refund(db, order, refund, provider_refund_no=provider_refund_no)
|
||||||
|
db.commit()
|
||||||
|
db.refresh(refund)
|
||||||
|
return ok(refund.to_dict())
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/token/admin/refunds/{refund_no}/confirm")
|
||||||
|
def admin_confirm_refund(
|
||||||
|
refund_no: str,
|
||||||
|
payload: dict = Body(...),
|
||||||
|
finance_key: str = Header(None, alias="X-Avatar-Finance-Key"),
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
):
|
||||||
|
"""Record a provider-console reconciliation result for asynchronous refunds."""
|
||||||
|
_require_finance_admin(finance_key)
|
||||||
|
refund = db.query(PaymentRefund).filter(PaymentRefund.refund_no == refund_no).first()
|
||||||
|
if not refund:
|
||||||
|
return fail("退款单不存在", 404)
|
||||||
|
order = db.query(TokenPaymentOrder).filter(TokenPaymentOrder.order_no == refund.order_no).first()
|
||||||
|
if not order:
|
||||||
|
return fail("原支付订单不存在", 404)
|
||||||
|
status = str(payload.get("status") or "").lower()
|
||||||
|
if status == "succeeded":
|
||||||
|
_complete_refund(
|
||||||
|
db,
|
||||||
|
order,
|
||||||
|
refund,
|
||||||
|
provider_refund_no=str(payload.get("providerRefundNo") or refund.provider_refund_no or ""),
|
||||||
|
)
|
||||||
|
elif status == "failed":
|
||||||
|
_complete_refund(
|
||||||
|
db,
|
||||||
|
order,
|
||||||
|
refund,
|
||||||
|
failure_reason=str(payload.get("failureReason") or "供应商退款失败"),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
return fail("退款确认状态只能是 succeeded 或 failed", 400)
|
||||||
|
db.commit()
|
||||||
|
db.refresh(refund)
|
||||||
|
return ok(refund.to_dict())
|
||||||
|
|
||||||
|
|
||||||
@router.post("/token/payment/callback/{order_no}/{callback_signature}")
|
@router.post("/token/payment/callback/{order_no}/{callback_signature}")
|
||||||
async def payment_callback(
|
async def payment_callback(
|
||||||
order_no: str,
|
order_no: str,
|
||||||
@@ -271,6 +713,8 @@ async def payment_callback(
|
|||||||
return fail("支付订单不存在", 404)
|
return fail("支付订单不存在", 404)
|
||||||
if order.status == "paid":
|
if order.status == "paid":
|
||||||
return ok({"received": True, "duplicate": True})
|
return ok({"received": True, "duplicate": True})
|
||||||
|
if order.status == "refunded":
|
||||||
|
return ok({"received": True, "duplicate": True, "refunded": True})
|
||||||
|
|
||||||
provider_status = str(_find_value(
|
provider_status = str(_find_value(
|
||||||
payload, "status", "payStatus", "tradeStatus", "paymentStatus"
|
payload, "status", "payStatus", "tradeStatus", "paymentStatus"
|
||||||
@@ -282,6 +726,14 @@ async def payment_callback(
|
|||||||
order.failure_reason = str(
|
order.failure_reason = str(
|
||||||
_find_value(payload, "message", "errorMsg", "failReason") or "支付失败"
|
_find_value(payload, "message", "errorMsg", "failReason") or "支付失败"
|
||||||
)[:500]
|
)[:500]
|
||||||
|
_record_transaction(
|
||||||
|
db,
|
||||||
|
order=order,
|
||||||
|
provider="huihui",
|
||||||
|
status="failed",
|
||||||
|
amount_cents=order.price_cents,
|
||||||
|
transaction_no=str(_find_value(payload, "transactionId", "tradeNo") or ""),
|
||||||
|
)
|
||||||
db.commit()
|
db.commit()
|
||||||
return ok({"received": True, "paid": False})
|
return ok({"received": True, "paid": False})
|
||||||
|
|
||||||
@@ -291,32 +743,148 @@ async def payment_callback(
|
|||||||
db.commit()
|
db.commit()
|
||||||
return fail("支付金额不匹配", 422)
|
return fail("支付金额不匹配", 422)
|
||||||
|
|
||||||
updated = db.query(TokenPaymentOrder).filter(
|
_settle_paid_order(
|
||||||
TokenPaymentOrder.id == order.id,
|
db,
|
||||||
TokenPaymentOrder.status != "paid",
|
order,
|
||||||
).update({
|
provider_status=provider_status,
|
||||||
TokenPaymentOrder.status: "paid",
|
transaction_no=str(_find_value(payload, "transactionId", "tradeNo", "paymentNo") or ""),
|
||||||
TokenPaymentOrder.provider_status: provider_status,
|
)
|
||||||
TokenPaymentOrder.paid_at: datetime.utcnow(),
|
|
||||||
TokenPaymentOrder.failure_reason: "",
|
|
||||||
}, synchronize_session=False)
|
|
||||||
if updated:
|
|
||||||
account = db.query(TokenAccount).filter(TokenAccount.user_id == order.user_id).first()
|
|
||||||
if account is None:
|
|
||||||
account = TokenAccount(
|
|
||||||
user_id=order.user_id,
|
|
||||||
balance=DEFAULT_TOKEN_GRANT,
|
|
||||||
total_granted=DEFAULT_TOKEN_GRANT,
|
|
||||||
total_consumed=0,
|
|
||||||
)
|
|
||||||
db.add(account)
|
|
||||||
db.flush()
|
|
||||||
account.balance = int(account.balance or 0) + order.points_amount
|
|
||||||
account.total_granted = int(account.total_granted or 0) + order.points_amount
|
|
||||||
db.commit()
|
db.commit()
|
||||||
return ok({"received": True, "paid": True})
|
return ok({"received": True, "paid": True})
|
||||||
|
|
||||||
|
|
||||||
|
def _wechat_notify_response(request: Request, *, success: bool, message: str = ""):
|
||||||
|
code = 0 if success else 1
|
||||||
|
text = "success" if success else (message or "fail")[:200].replace("]]>", "")
|
||||||
|
if "xml" in (request.headers.get("content-type") or "").lower():
|
||||||
|
return Response(
|
||||||
|
content=f"<xml><ErrCode>{code}</ErrCode><ErrMsg><![CDATA[{text}]]></ErrMsg></xml>",
|
||||||
|
media_type="application/xml",
|
||||||
|
)
|
||||||
|
return {"ErrCode": code, "ErrMsg": text}
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/token/payment/wechat/virtual/notify")
|
||||||
|
def validate_wechat_virtual_notify(
|
||||||
|
signature: str = Query(""),
|
||||||
|
timestamp: str = Query(""),
|
||||||
|
nonce: str = Query(""),
|
||||||
|
echostr: str = Query(""),
|
||||||
|
):
|
||||||
|
if not verify_wechat_callback_signature(signature, timestamp, nonce):
|
||||||
|
raise HTTPException(status_code=403, detail="invalid signature")
|
||||||
|
return Response(content=echostr or "ok", media_type="text/plain")
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/token/payment/wechat/virtual/notify")
|
||||||
|
async def wechat_virtual_notify(
|
||||||
|
request: Request,
|
||||||
|
signature: str = Query(""),
|
||||||
|
timestamp: str = Query(""),
|
||||||
|
nonce: str = Query(""),
|
||||||
|
db: Session = Depends(get_db),
|
||||||
|
):
|
||||||
|
if not verify_wechat_callback_signature(signature, timestamp, nonce):
|
||||||
|
return _wechat_notify_response(request, success=False, message="invalid signature")
|
||||||
|
try:
|
||||||
|
payload = _nested_payload(parse_wechat_callback_body(await request.body()))
|
||||||
|
except WechatVirtualPaymentError as exc:
|
||||||
|
return _wechat_notify_response(request, success=False, message=str(exc))
|
||||||
|
|
||||||
|
event = str(wechat_callback_value(payload, "Event") or "").lower()
|
||||||
|
if event in WECHAT_PAYMENT_EVENTS:
|
||||||
|
order_no = str(wechat_callback_value(payload, "OutTradeNo") or "").strip()
|
||||||
|
order = db.query(TokenPaymentOrder).filter(TokenPaymentOrder.order_no == order_no).first()
|
||||||
|
if not order or order.provider != "wechat_virtual":
|
||||||
|
return _wechat_notify_response(request, success=False, message="order not found")
|
||||||
|
user = db.query(User).filter(User.id == order.user_id).first()
|
||||||
|
openid = str(wechat_callback_value(payload, "OpenId") or "").strip()
|
||||||
|
if not user or not openid or openid != user.wechat_mp_openid:
|
||||||
|
return _wechat_notify_response(request, success=False, message="openid mismatch")
|
||||||
|
try:
|
||||||
|
callback_env = int(wechat_callback_value(payload, "Env"))
|
||||||
|
actual_price = int(wechat_callback_value(payload, "GoodsInfo", "ActualPrice"))
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return _wechat_notify_response(request, success=False, message="invalid payment amount")
|
||||||
|
plan = db.query(TokenPlan).filter(TokenPlan.id == order.plan_id).first()
|
||||||
|
product_id = str(wechat_callback_value(payload, "GoodsInfo", "ProductId") or "")
|
||||||
|
try:
|
||||||
|
expected_product_id = product_id_for_plan(plan) if plan else ""
|
||||||
|
except WechatVirtualPaymentError:
|
||||||
|
expected_product_id = ""
|
||||||
|
if (
|
||||||
|
callback_env != wechat_virtual_env()
|
||||||
|
or actual_price != order.price_cents
|
||||||
|
or not expected_product_id
|
||||||
|
or product_id != expected_product_id
|
||||||
|
):
|
||||||
|
return _wechat_notify_response(request, success=False, message="payment verification failed")
|
||||||
|
transaction_no = str(
|
||||||
|
wechat_callback_value(payload, "WeChatPayInfo", "TransactionId")
|
||||||
|
or wechat_callback_value(payload, "WeChatPayInfo", "MchOrderNo")
|
||||||
|
or order_no
|
||||||
|
)
|
||||||
|
_settle_paid_order(
|
||||||
|
db,
|
||||||
|
order,
|
||||||
|
provider_status=event,
|
||||||
|
transaction_no=transaction_no,
|
||||||
|
raw_summary=_safe_event_summary(payload),
|
||||||
|
)
|
||||||
|
db.commit()
|
||||||
|
return _wechat_notify_response(request, success=True)
|
||||||
|
|
||||||
|
if event in WECHAT_REFUND_EVENTS:
|
||||||
|
order_no = str(wechat_callback_value(payload, "MchOrderId") or "").strip()
|
||||||
|
refund_no = str(wechat_callback_value(payload, "MchRefundId") or "").strip()
|
||||||
|
order = db.query(TokenPaymentOrder).filter(TokenPaymentOrder.order_no == order_no).first()
|
||||||
|
if not order or order.provider != "wechat_virtual":
|
||||||
|
return _wechat_notify_response(request, success=False, message="order not found")
|
||||||
|
if order.status == "refunded" or order.refund_status == "succeeded":
|
||||||
|
return _wechat_notify_response(request, success=True)
|
||||||
|
try:
|
||||||
|
refund_cents = int(wechat_callback_value(payload, "RefundFee") or 0)
|
||||||
|
result_code_value = wechat_callback_value(payload, "RetCode")
|
||||||
|
if result_code_value in (None, ""):
|
||||||
|
raise ValueError("missing RetCode")
|
||||||
|
result_code = int(result_code_value)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return _wechat_notify_response(request, success=False, message="invalid refund")
|
||||||
|
refund = db.query(PaymentRefund).filter(PaymentRefund.refund_no == refund_no).first()
|
||||||
|
if refund is None:
|
||||||
|
refund = PaymentRefund(
|
||||||
|
refund_no=refund_no or f"WR{uuid.uuid4().hex[:20].upper()}",
|
||||||
|
order_no=order.order_no,
|
||||||
|
amount_cents=refund_cents,
|
||||||
|
points_amount=order.points_amount,
|
||||||
|
reason="微信侧退款",
|
||||||
|
status="processing",
|
||||||
|
requested_by="wechat",
|
||||||
|
)
|
||||||
|
db.add(refund)
|
||||||
|
if refund_cents != refund.amount_cents:
|
||||||
|
return _wechat_notify_response(request, success=False, message="refund amount mismatch")
|
||||||
|
if result_code == 0:
|
||||||
|
_complete_refund(
|
||||||
|
db,
|
||||||
|
order,
|
||||||
|
refund,
|
||||||
|
provider_refund_no=str(wechat_callback_value(payload, "WxRefundId") or refund_no),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
_complete_refund(
|
||||||
|
db,
|
||||||
|
order,
|
||||||
|
refund,
|
||||||
|
failure_reason=str(wechat_callback_value(payload, "RetMsg") or "微信退款失败"),
|
||||||
|
)
|
||||||
|
db.commit()
|
||||||
|
return _wechat_notify_response(request, success=True)
|
||||||
|
|
||||||
|
# Irrelevant official-account events should not be retried as payment failures.
|
||||||
|
return _wechat_notify_response(request, success=True)
|
||||||
|
|
||||||
|
|
||||||
@router.get("/token/usage")
|
@router.get("/token/usage")
|
||||||
def usage(authorization: str = Header(None), db: Session = Depends(get_db)):
|
def usage(authorization: str = Header(None), db: Session = Depends(get_db)):
|
||||||
user = _require_user(authorization, db)
|
user = _require_user(authorization, db)
|
||||||
|
|||||||
@@ -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",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
"""Signed client for Huihui's production payment-v3 service."""
|
"""Signed client for Huihui's production payment-v3 service."""
|
||||||
|
|
||||||
import hashlib
|
import hashlib
|
||||||
|
import os
|
||||||
import random
|
import random
|
||||||
import string
|
import string
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
@@ -122,3 +123,59 @@ class HuihuiPaymentClient:
|
|||||||
if not isinstance(data, dict):
|
if not isinstance(data, dict):
|
||||||
raise HuihuiPaymentError("会会支付未返回订单信息")
|
raise HuihuiPaymentError("会会支付未返回订单信息")
|
||||||
return data
|
return data
|
||||||
|
|
||||||
|
def request_refund(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
huihui_token: str,
|
||||||
|
huihui_user_id: str,
|
||||||
|
order_no: str,
|
||||||
|
refund_no: str,
|
||||||
|
amount: str,
|
||||||
|
reason: str,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Submit a full refund to payment-v3.
|
||||||
|
|
||||||
|
The refund path remains configurable because private Huihui deployments
|
||||||
|
may expose the same contract below a different gateway route.
|
||||||
|
"""
|
||||||
|
if not self.configured:
|
||||||
|
raise HuihuiPaymentError("会会支付服务未配置")
|
||||||
|
if not huihui_token or not huihui_user_id:
|
||||||
|
raise HuihuiPaymentError("当前会会登录凭证无法发起退款")
|
||||||
|
|
||||||
|
path = os.getenv("HUIHUI_PAYMENT_REFUND_PATH", "/payment/refund").strip()
|
||||||
|
if not path.startswith("/"):
|
||||||
|
path = f"/{path}"
|
||||||
|
if ".." in path:
|
||||||
|
raise HuihuiPaymentError("会会退款接口路径配置不正确")
|
||||||
|
body = {
|
||||||
|
"appId": self.app_id,
|
||||||
|
"masterOrderNo": order_no,
|
||||||
|
"refundOrderNo": refund_no,
|
||||||
|
"refundAmt": float(amount),
|
||||||
|
"refundReason": reason or "后台退款",
|
||||||
|
}
|
||||||
|
headers = {
|
||||||
|
"Authorization": f"Bearer {huihui_token}",
|
||||||
|
"appId": self.app_id,
|
||||||
|
"windowAppId": self.app_id,
|
||||||
|
}
|
||||||
|
try:
|
||||||
|
response = httpx.post(
|
||||||
|
f"{self.base_url}{path}",
|
||||||
|
headers=headers,
|
||||||
|
params=self._signed_params(huihui_user_id),
|
||||||
|
json=body,
|
||||||
|
timeout=self.timeout,
|
||||||
|
follow_redirects=True,
|
||||||
|
)
|
||||||
|
except httpx.HTTPError as exc:
|
||||||
|
raise HuihuiPaymentError("会会退款连接失败,请稍后重试") from exc
|
||||||
|
|
||||||
|
payload = self._json(response)
|
||||||
|
code = payload.get("code")
|
||||||
|
if response.status_code >= 400 or code not in (0, 200, "0", "200"):
|
||||||
|
raise HuihuiPaymentError(payload.get("message") or "会会退款申请失败")
|
||||||
|
data = payload.get("data") or {}
|
||||||
|
return data if isinstance(data, dict) else {"result": data}
|
||||||
|
|||||||
@@ -0,0 +1,160 @@
|
|||||||
|
"""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 Avatar, KnowledgeChunk, KnowledgeDoc
|
||||||
|
from services.pdf_ocr_service import extract_scanned_pdf_text
|
||||||
|
import embeddings
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
BACKEND_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||||
|
UPLOAD_DIR = os.path.abspath(
|
||||||
|
os.getenv("UPLOAD_DIR", os.path.join(BACKEND_DIR, "routers", "uploads"))
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class KnowledgeVectorizer:
|
||||||
|
"""Indexes one document at a time so slow providers cannot block uploads."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._queue: queue.Queue[str] = queue.Queue()
|
||||||
|
self._queued: set[str] = set()
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
self._thread: threading.Thread | None = None
|
||||||
|
|
||||||
|
def start(self):
|
||||||
|
if self._thread and self._thread.is_alive():
|
||||||
|
return
|
||||||
|
self._thread = threading.Thread(
|
||||||
|
target=self._run, name="knowledge-vectorizer", daemon=True
|
||||||
|
)
|
||||||
|
self._thread.start()
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
# A process restart must not abandon documents already accepted by upload.
|
||||||
|
for (doc_id,) in db.query(KnowledgeDoc.id).filter(KnowledgeDoc.status == "parsing"):
|
||||||
|
self.enqueue(doc_id)
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
def enqueue(self, doc_id: str):
|
||||||
|
with self._lock:
|
||||||
|
if doc_id in self._queued:
|
||||||
|
return
|
||||||
|
self._queued.add(doc_id)
|
||||||
|
self._queue.put(doc_id)
|
||||||
|
|
||||||
|
def _run(self):
|
||||||
|
while True:
|
||||||
|
doc_id = self._queue.get()
|
||||||
|
try:
|
||||||
|
self.vectorize_document(doc_id)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Unexpected knowledge vectorizer failure for %s", doc_id)
|
||||||
|
finally:
|
||||||
|
with self._lock:
|
||||||
|
self._queued.discard(doc_id)
|
||||||
|
self._queue.task_done()
|
||||||
|
|
||||||
|
def vectorize_document(self, doc_id: str):
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
doc = db.get(KnowledgeDoc, doc_id)
|
||||||
|
if not doc or doc.status != "parsing":
|
||||||
|
return
|
||||||
|
|
||||||
|
stored_name = os.path.basename(doc.file_url or "")
|
||||||
|
path = os.path.join(UPLOAD_DIR, doc.avatar_id, stored_name)
|
||||||
|
if not stored_name or not os.path.isfile(path):
|
||||||
|
raise FileNotFoundError("原文件不可用,请重新上传")
|
||||||
|
|
||||||
|
self._set_progress(db, doc, "extracting", 8)
|
||||||
|
text = embeddings.extract_text(path, f".{doc.file_type}")
|
||||||
|
if doc.file_type == "pdf" and not text.strip():
|
||||||
|
avatar = db.get(Avatar, doc.avatar_id)
|
||||||
|
if not avatar:
|
||||||
|
raise ValueError("文档所属分身不存在")
|
||||||
|
|
||||||
|
def ocr_progress(done: int, total: int):
|
||||||
|
percent = 8 + int((done / max(1, total)) * 20)
|
||||||
|
self._set_progress(db, doc, "ocr", min(percent, 28))
|
||||||
|
|
||||||
|
self._set_progress(db, doc, "ocr", 8)
|
||||||
|
text = extract_scanned_pdf_text(
|
||||||
|
db,
|
||||||
|
avatar,
|
||||||
|
path,
|
||||||
|
on_progress=ocr_progress,
|
||||||
|
)
|
||||||
|
self._set_progress(db, doc, "chunking", 29)
|
||||||
|
chunks = embeddings.chunk_text(text)
|
||||||
|
if not chunks:
|
||||||
|
raise ValueError("文档没有可建立索引的文字内容")
|
||||||
|
self._set_progress(db, doc, "embedding", 30)
|
||||||
|
|
||||||
|
def embedding_progress(done: int, total: int):
|
||||||
|
percent = 30 + int((done / max(1, total)) * 65)
|
||||||
|
self._set_progress(db, doc, "embedding", min(percent, 95))
|
||||||
|
|
||||||
|
vectors = embeddings.embed(chunks, on_progress=embedding_progress)
|
||||||
|
if len(vectors) != len(chunks):
|
||||||
|
raise ValueError("向量服务返回数量与文档分段不一致")
|
||||||
|
|
||||||
|
# Commit the document and every chunk together. Chat only sees complete indexes.
|
||||||
|
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == doc.id).delete()
|
||||||
|
db.add_all(
|
||||||
|
[
|
||||||
|
KnowledgeChunk(
|
||||||
|
doc_id=doc.id,
|
||||||
|
avatar_id=doc.avatar_id,
|
||||||
|
content=chunk,
|
||||||
|
vector=json.dumps(vector),
|
||||||
|
chunk_index=index,
|
||||||
|
embedding_model=embeddings.MODEL,
|
||||||
|
)
|
||||||
|
for index, (chunk, vector) in enumerate(zip(chunks, vectors))
|
||||||
|
]
|
||||||
|
)
|
||||||
|
doc.vectorized = True
|
||||||
|
doc.embedding_model = embeddings.MODEL
|
||||||
|
doc.chunk_count = len(chunks)
|
||||||
|
doc.vectorized_at = datetime.now(timezone.utc)
|
||||||
|
doc.status = "ready"
|
||||||
|
doc.error_message = ""
|
||||||
|
doc.index_stage = "ready"
|
||||||
|
doc.index_progress = 100
|
||||||
|
db.commit()
|
||||||
|
logger.info("Knowledge document %s indexed with %s chunks", doc.id, len(chunks))
|
||||||
|
except Exception as exc:
|
||||||
|
db.rollback()
|
||||||
|
failed_doc = db.get(KnowledgeDoc, doc_id)
|
||||||
|
if failed_doc:
|
||||||
|
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == failed_doc.id).delete()
|
||||||
|
failed_doc.status = "failed"
|
||||||
|
failed_doc.vectorized = False
|
||||||
|
failed_doc.embedding_model = ""
|
||||||
|
failed_doc.chunk_count = 0
|
||||||
|
failed_doc.vectorized_at = None
|
||||||
|
failed_doc.error_message = str(exc)[:300] or "建立知识索引失败"
|
||||||
|
failed_doc.index_stage = "failed"
|
||||||
|
failed_doc.index_progress = 0
|
||||||
|
db.commit()
|
||||||
|
logger.exception("Knowledge vectorization failed for %s: %s", doc_id, exc)
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _set_progress(db, doc, stage: str, progress: int):
|
||||||
|
doc.index_stage = stage
|
||||||
|
doc.index_progress = progress
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
|
||||||
|
knowledge_vectorizer = KnowledgeVectorizer()
|
||||||
@@ -0,0 +1,130 @@
|
|||||||
|
"""OCR fallback for image-only PDF knowledge documents."""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import time
|
||||||
|
from typing import Callable
|
||||||
|
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
from models import Avatar
|
||||||
|
from services.chat_model_config import get_chat_model_config
|
||||||
|
from services.token_billing import (
|
||||||
|
estimate_fallback_usage,
|
||||||
|
release_reservation,
|
||||||
|
reserve_avatar_tokens,
|
||||||
|
settle_reservation,
|
||||||
|
)
|
||||||
|
from services.vision_service import call_vision_model, prepare_image
|
||||||
|
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
PDF_OCR_PROMPT = (
|
||||||
|
"请逐字转录这一页扫描文档中的全部可见文字和表格,只输出转录内容,不要解释,不要使用 Markdown 代码块。"
|
||||||
|
"保留标题、段落、项目编号、数值和自然换行;看不清的内容写作[无法辨认],不要猜测、纠错或补全。"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _positive_int(name: str, default: int, minimum: int, maximum: int) -> int:
|
||||||
|
try:
|
||||||
|
value = int(os.getenv(name, str(default)))
|
||||||
|
except ValueError:
|
||||||
|
value = default
|
||||||
|
return max(minimum, min(maximum, value))
|
||||||
|
|
||||||
|
|
||||||
|
def extract_scanned_pdf_text(
|
||||||
|
db: Session,
|
||||||
|
avatar: Avatar,
|
||||||
|
path: str,
|
||||||
|
*,
|
||||||
|
on_progress: Callable[[int, int], None] | None = None,
|
||||||
|
) -> str:
|
||||||
|
"""Render and OCR an image-only PDF while preserving page order."""
|
||||||
|
try:
|
||||||
|
import pymupdf
|
||||||
|
except ImportError as exc:
|
||||||
|
raise RuntimeError("扫描型 PDF 识别组件未安装") from exc
|
||||||
|
|
||||||
|
max_pages = _positive_int("KNOWLEDGE_PDF_OCR_MAX_PAGES", 80, 1, 300)
|
||||||
|
render_dpi = _positive_int("KNOWLEDGE_PDF_OCR_DPI", 144, 96, 200)
|
||||||
|
max_attempts = _positive_int("KNOWLEDGE_PDF_OCR_ATTEMPTS", 3, 1, 5)
|
||||||
|
model_config = get_chat_model_config()
|
||||||
|
model = model_config.ocr_model or model_config.vision_model
|
||||||
|
if not model_config.api_key or not model:
|
||||||
|
raise RuntimeError("扫描型 PDF 需要配置视觉 OCR 模型")
|
||||||
|
|
||||||
|
texts: list[str] = []
|
||||||
|
with pymupdf.open(path) as document:
|
||||||
|
total_pages = document.page_count
|
||||||
|
if total_pages <= 0:
|
||||||
|
raise ValueError("PDF 没有可识别页面")
|
||||||
|
if total_pages > max_pages:
|
||||||
|
raise ValueError(
|
||||||
|
f"扫描型 PDF 共 {total_pages} 页,超过单次 OCR 上限 {max_pages} 页,请拆分后上传"
|
||||||
|
)
|
||||||
|
|
||||||
|
scale = render_dpi / 72
|
||||||
|
for page_index in range(total_pages):
|
||||||
|
page = document.load_page(page_index)
|
||||||
|
pixmap = page.get_pixmap(
|
||||||
|
matrix=pymupdf.Matrix(scale, scale),
|
||||||
|
colorspace=pymupdf.csRGB,
|
||||||
|
alpha=False,
|
||||||
|
)
|
||||||
|
prepared = prepare_image(pixmap.tobytes("jpeg", jpg_quality=88))
|
||||||
|
estimate_messages = [{
|
||||||
|
"role": "user",
|
||||||
|
"content": f"[扫描 PDF 第 {page_index + 1}/{total_pages} 页]\n{PDF_OCR_PROMPT}",
|
||||||
|
}]
|
||||||
|
reservation = reserve_avatar_tokens(
|
||||||
|
db,
|
||||||
|
avatar,
|
||||||
|
"knowledge_pdf_ocr",
|
||||||
|
model,
|
||||||
|
estimate_messages,
|
||||||
|
model_config.vision_max_tokens,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
result = None
|
||||||
|
for attempt in range(1, max_attempts + 1):
|
||||||
|
try:
|
||||||
|
result = call_vision_model(
|
||||||
|
prepared,
|
||||||
|
model_config,
|
||||||
|
model=model,
|
||||||
|
prompt=PDF_OCR_PROMPT,
|
||||||
|
json_output=False,
|
||||||
|
)
|
||||||
|
break
|
||||||
|
except RuntimeError:
|
||||||
|
if attempt == max_attempts:
|
||||||
|
raise
|
||||||
|
time.sleep(min(4, attempt))
|
||||||
|
content = str((result or {}).get("content") or "").strip()
|
||||||
|
if not content:
|
||||||
|
raise RuntimeError("扫描型 PDF 页面识别结果为空")
|
||||||
|
settle_reservation(
|
||||||
|
db,
|
||||||
|
reservation,
|
||||||
|
(result or {}).get("usage"),
|
||||||
|
fallback_total=estimate_fallback_usage(estimate_messages, content),
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
release_reservation(db, reservation, str(exc))
|
||||||
|
raise RuntimeError(
|
||||||
|
f"扫描型 PDF 第 {page_index + 1}/{total_pages} 页识别失败:{exc}"
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
texts.append(f"[第 {page_index + 1} 页]\n{content}")
|
||||||
|
if on_progress:
|
||||||
|
on_progress(page_index + 1, total_pages)
|
||||||
|
logger.info(
|
||||||
|
"Scanned PDF OCR completed avatar=%s page=%s/%s",
|
||||||
|
avatar.id,
|
||||||
|
page_index + 1,
|
||||||
|
total_pages,
|
||||||
|
)
|
||||||
|
|
||||||
|
return "\n\n".join(texts).strip()
|
||||||
@@ -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,19 +14,27 @@ 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"
|
TAKEOVER_DELAY_KEY = "takeoverReplyDelaySeconds"
|
||||||
@@ -36,6 +45,19 @@ HUMAN_PAUSE_SECONDS = 600
|
|||||||
RATE_LIMIT_WINDOW_SECONDS = 300
|
RATE_LIMIT_WINDOW_SECONDS = 300
|
||||||
RATE_LIMIT_MAX_REPLIES = 5
|
RATE_LIMIT_MAX_REPLIES = 5
|
||||||
AVATAR_LOCAL_ID_PREFIX = "880"
|
AVATAR_LOCAL_ID_PREFIX = "880"
|
||||||
|
BOXIM_TEXT_MESSAGE_TYPE = 0
|
||||||
|
BOXIM_IMAGE_MESSAGE_TYPE = 1
|
||||||
|
BOXIM_IMAGE_PROMPT = "请看看这张图片。"
|
||||||
|
BOXIM_IMAGE_UNAVAILABLE_REPLY = "这张图片我暂时没看清,麻烦重新发送一张清晰的原图。"
|
||||||
|
IMAGE_CONTEXT_LOOKBACK_SECONDS = 1800
|
||||||
|
IMAGE_REFERENCE_LOOKBACK_SECONDS = 172_800
|
||||||
|
MAX_RECENT_IMAGE_CONTEXTS = 3
|
||||||
|
|
||||||
|
_IMAGE_REFERENCE_PATTERN = re.compile(
|
||||||
|
r"(?:图片|图像|照片|截图|这张图|刚才.{0,8}图|病例|病历|检查单|检验单|化验单|报告|影像|"
|
||||||
|
r"\b(?:image|photo|picture|screenshot|scan|report)\b)",
|
||||||
|
re.IGNORECASE,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _utcnow() -> datetime:
|
def _utcnow() -> datetime:
|
||||||
@@ -106,6 +128,18 @@ def _configured_reply_delay(avatar: Avatar, fallback: int | None = None) -> int:
|
|||||||
return delay
|
return delay
|
||||||
|
|
||||||
|
|
||||||
|
def _event_prompt(event: TakeoverMessage) -> str:
|
||||||
|
if event.message_type == BOXIM_TEXT_MESSAGE_TYPE:
|
||||||
|
return event.content.strip()
|
||||||
|
if event.message_type == BOXIM_IMAGE_MESSAGE_TYPE:
|
||||||
|
return BOXIM_IMAGE_PROMPT
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
def _references_recent_image(value: str) -> bool:
|
||||||
|
return bool(_IMAGE_REFERENCE_PATTERN.search(value or ""))
|
||||||
|
|
||||||
|
|
||||||
class TakeoverService:
|
class TakeoverService:
|
||||||
"""Poll BOXIM, honor the owner grace period, then generate and send one reply."""
|
"""Poll BOXIM, honor the owner grace period, then generate and send one reply."""
|
||||||
|
|
||||||
@@ -115,15 +149,20 @@ class TakeoverService:
|
|||||||
boxim_client: BoxIMClient,
|
boxim_client: BoxIMClient,
|
||||||
*,
|
*,
|
||||||
reply_delay_seconds: int | None = None,
|
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."""
|
||||||
@@ -138,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."""
|
||||||
@@ -290,7 +369,10 @@ 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._record_connection_failure(
|
self._record_connection_failure(
|
||||||
db,
|
db,
|
||||||
@@ -341,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 "")
|
||||||
@@ -362,11 +437,22 @@ class TakeoverService:
|
|||||||
session["access_token"], peer_id, message_id
|
session["access_token"], peer_id, message_id
|
||||||
)
|
)
|
||||||
|
|
||||||
cursor.last_message_id = str(max_message_id)
|
# Keep SQLite write transactions short. The read-receipt request above
|
||||||
cursor.initialized = True
|
# can block on the network and must not hold the database write lock.
|
||||||
cursor.last_polled_at = self.now()
|
async with self._persist_lock:
|
||||||
cursor.last_error = ""
|
for message in messages:
|
||||||
db.commit()
|
self._record_message(
|
||||||
|
db,
|
||||||
|
avatar,
|
||||||
|
cursor.boxim_owner_id,
|
||||||
|
message,
|
||||||
|
schedule_reply=not priming,
|
||||||
|
)
|
||||||
|
cursor.last_message_id = str(max_message_id)
|
||||||
|
cursor.initialized = True
|
||||||
|
cursor.last_polled_at = self.now()
|
||||||
|
cursor.last_error = ""
|
||||||
|
db.commit()
|
||||||
return True
|
return True
|
||||||
except Exception:
|
except Exception:
|
||||||
db.rollback()
|
db.rollback()
|
||||||
@@ -450,18 +536,58 @@ 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
|
return
|
||||||
if is_avatar:
|
if is_avatar:
|
||||||
self._cancel_conversation(db, avatar.owner_id, peer_id, "peer_avatar_message")
|
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
|
return
|
||||||
if self._human_pause_active(db, avatar.owner_id, peer_id, now):
|
if self._human_pause_active(db, avatar.owner_id, peer_id, now):
|
||||||
self._cancel_conversation(db, avatar.owner_id, peer_id, "owner_active")
|
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
|
return
|
||||||
if self._conversation_rate_limited(db, avatar.owner_id, peer_id, now):
|
if self._conversation_rate_limited(db, avatar.owner_id, peer_id, now):
|
||||||
self._cancel_conversation(db, avatar.owner_id, peer_id, "rate_limited")
|
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)
|
||||||
|
|
||||||
@@ -537,11 +663,22 @@ 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(
|
due_at = max(
|
||||||
seconds=_configured_reply_delay(avatar, self.reply_delay_seconds)
|
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 = _avatar_local_id(avatar.owner_id, event.boxim_message_id)
|
local_id = _avatar_local_id(avatar.owner_id, event.boxim_message_id)
|
||||||
@@ -560,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:
|
||||||
@@ -594,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:
|
||||||
@@ -612,6 +855,21 @@ class TakeoverService:
|
|||||||
db.commit()
|
db.commit()
|
||||||
|
|
||||||
excluded_ids = set(task.source_message_ids or [])
|
excluded_ids = set(task.source_message_ids or [])
|
||||||
|
source_events = {
|
||||||
|
event.boxim_message_id: event
|
||||||
|
for event in (
|
||||||
|
db.query(TakeoverMessage)
|
||||||
|
.filter(
|
||||||
|
TakeoverMessage.owner_id == task.owner_id,
|
||||||
|
TakeoverMessage.peer_id == task.peer_id,
|
||||||
|
TakeoverMessage.avatar_id == task.avatar_id,
|
||||||
|
TakeoverMessage.boxim_message_id.in_(excluded_ids),
|
||||||
|
)
|
||||||
|
.all()
|
||||||
|
if excluded_ids
|
||||||
|
else []
|
||||||
|
)
|
||||||
|
}
|
||||||
events = (
|
events = (
|
||||||
db.query(TakeoverMessage)
|
db.query(TakeoverMessage)
|
||||||
.filter(
|
.filter(
|
||||||
@@ -623,9 +881,31 @@ class TakeoverService:
|
|||||||
.limit(30)
|
.limit(30)
|
||||||
.all()
|
.all()
|
||||||
)
|
)
|
||||||
|
image_attachments = []
|
||||||
|
image_failed = False
|
||||||
|
for message_id in (task.source_message_ids or [])[-3:]:
|
||||||
|
event = source_events.get(message_id)
|
||||||
|
if not event or event.message_type != BOXIM_IMAGE_MESSAGE_TYPE:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
image_attachments.append(
|
||||||
|
self._takeover_image_attachment(db, avatar, event)
|
||||||
|
)
|
||||||
|
except (BoxIMImageError, ImageValidationError) as exc:
|
||||||
|
image_failed = True
|
||||||
|
logger.warning(
|
||||||
|
"BOXIM image unavailable message=%s avatar=%s: %s",
|
||||||
|
event.boxim_message_id,
|
||||||
|
avatar.id,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
history = []
|
history = []
|
||||||
for event in reversed(events):
|
for event in reversed(events):
|
||||||
if event.boxim_message_id in excluded_ids or not event.content.strip():
|
if (
|
||||||
|
event.boxim_message_id in excluded_ids
|
||||||
|
or event.message_type != BOXIM_TEXT_MESSAGE_TYPE
|
||||||
|
or not event.content.strip()
|
||||||
|
):
|
||||||
continue
|
continue
|
||||||
if event.direction == "incoming" and event.is_avatar:
|
if event.direction == "incoming" and event.is_avatar:
|
||||||
continue
|
continue
|
||||||
@@ -637,10 +917,21 @@ class TakeoverService:
|
|||||||
)
|
)
|
||||||
history = history[-10:]
|
history = history[-10:]
|
||||||
|
|
||||||
from routers.chat import _resolve_reply
|
from routers.chat import _attachment_contexts, _resolve_reply
|
||||||
|
|
||||||
result = _resolve_reply(db, avatar, task.prompt, history, usage_source="takeover")
|
image_contexts = _attachment_contexts(image_attachments)
|
||||||
answer = _plain_text_reply(result.get("answer", ""))
|
if image_failed and not image_contexts:
|
||||||
|
answer = BOXIM_IMAGE_UNAVAILABLE_REPLY
|
||||||
|
else:
|
||||||
|
result = _resolve_reply(
|
||||||
|
db,
|
||||||
|
avatar,
|
||||||
|
task.prompt,
|
||||||
|
history,
|
||||||
|
usage_source="takeover",
|
||||||
|
image_contexts=image_contexts,
|
||||||
|
)
|
||||||
|
answer = _plain_text_reply(result.get("answer", ""))
|
||||||
db.refresh(task)
|
db.refresh(task)
|
||||||
if task.status != "generating":
|
if task.status != "generating":
|
||||||
return False
|
return False
|
||||||
@@ -701,7 +992,7 @@ 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()
|
||||||
|
|||||||
@@ -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()]
|
||||||
@@ -0,0 +1,245 @@
|
|||||||
|
"""WeChat mini-program virtual-payment signing and server API adapter.
|
||||||
|
|
||||||
|
The AppKey and session_key never leave the backend. The JSON string returned as
|
||||||
|
``signData`` is exactly the string used for both HMAC signatures.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import hmac
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import time
|
||||||
|
import xml.etree.ElementTree as ET
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
|
||||||
|
REQUEST_VIRTUAL_PAYMENT_URI = "requestVirtualPayment"
|
||||||
|
PAYMENT_EVENTS = {"xpay_goods_deliver_notify"}
|
||||||
|
REFUND_EVENTS = {"xpay_refund_notify"}
|
||||||
|
|
||||||
|
|
||||||
|
class WechatVirtualPaymentError(RuntimeError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def json_compact(payload: dict[str, Any]) -> str:
|
||||||
|
return json.dumps(payload, ensure_ascii=False, separators=(",", ":"))
|
||||||
|
|
||||||
|
|
||||||
|
def hmac_sha256_hex(key: str, message: str) -> str:
|
||||||
|
return hmac.new(key.encode("utf-8"), message.encode("utf-8"), hashlib.sha256).hexdigest()
|
||||||
|
|
||||||
|
|
||||||
|
def virtual_env() -> int:
|
||||||
|
value = os.getenv("WECHAT_VIRTUAL_ENV", "sandbox").strip().lower()
|
||||||
|
return 0 if value in {"0", "prod", "production", "live", "online"} else 1
|
||||||
|
|
||||||
|
|
||||||
|
def _app_key(env: int) -> str:
|
||||||
|
name = "WECHAT_VIRTUAL_APP_KEY" if env == 0 else "WECHAT_VIRTUAL_SANDBOX_APP_KEY"
|
||||||
|
return os.getenv(name, "").strip()
|
||||||
|
|
||||||
|
|
||||||
|
def _offer_id() -> str:
|
||||||
|
return os.getenv("WECHAT_VIRTUAL_OFFER_ID", "").strip()
|
||||||
|
|
||||||
|
|
||||||
|
def product_id_for_plan(plan) -> str:
|
||||||
|
configured = str(getattr(plan, "virtual_product_id", "") or "").strip()
|
||||||
|
if not configured:
|
||||||
|
configured = os.getenv(f"WECHAT_VIRTUAL_PRODUCT_{plan.id}", "").strip()
|
||||||
|
if not configured:
|
||||||
|
raise WechatVirtualPaymentError(f"套餐 {plan.id} 尚未配置微信虚拟支付商品 ID")
|
||||||
|
if len(configured) > 64 or not all(ch.isalnum() or ch in "_-" for ch in configured):
|
||||||
|
raise WechatVirtualPaymentError("微信虚拟支付商品 ID 格式不正确")
|
||||||
|
return configured
|
||||||
|
|
||||||
|
|
||||||
|
def build_payment_params(*, order, plan, session_key: str) -> dict[str, Any]:
|
||||||
|
env = virtual_env()
|
||||||
|
offer_id = _offer_id()
|
||||||
|
app_key = _app_key(env)
|
||||||
|
if not offer_id or not app_key or not session_key:
|
||||||
|
raise WechatVirtualPaymentError("微信小程序虚拟支付配置不完整")
|
||||||
|
|
||||||
|
sign_data = json_compact({
|
||||||
|
"offerId": offer_id,
|
||||||
|
"buyQuantity": 1,
|
||||||
|
"env": env,
|
||||||
|
"currencyType": "CNY",
|
||||||
|
"productId": product_id_for_plan(plan),
|
||||||
|
"goodsPrice": int(order.price_cents),
|
||||||
|
"outTradeNo": order.order_no,
|
||||||
|
"attach": json_compact({"orderNo": order.order_no, "planId": order.plan_id}),
|
||||||
|
})
|
||||||
|
return {
|
||||||
|
"provider": "wechat_virtual",
|
||||||
|
"payment_channel": "virtual",
|
||||||
|
"payment_method": "wechat",
|
||||||
|
"mode": "short_series_goods",
|
||||||
|
"signData": sign_data,
|
||||||
|
"paySig": hmac_sha256_hex(app_key, f"{REQUEST_VIRTUAL_PAYMENT_URI}&{sign_data}"),
|
||||||
|
"signature": hmac_sha256_hex(session_key, sign_data),
|
||||||
|
"env": env,
|
||||||
|
"offerId": offer_id,
|
||||||
|
"outTradeNo": order.order_no,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def exchange_code(code: str) -> dict[str, str]:
|
||||||
|
app_id = os.getenv("WECHAT_MP_APP_ID", "").strip()
|
||||||
|
app_secret = os.getenv("WECHAT_MP_APP_SECRET", "").strip()
|
||||||
|
if not app_id or not app_secret:
|
||||||
|
raise WechatVirtualPaymentError("微信小程序登录配置不完整")
|
||||||
|
try:
|
||||||
|
response = httpx.get(
|
||||||
|
"https://api.weixin.qq.com/sns/jscode2session",
|
||||||
|
params={
|
||||||
|
"appid": app_id,
|
||||||
|
"secret": app_secret,
|
||||||
|
"js_code": code,
|
||||||
|
"grant_type": "authorization_code",
|
||||||
|
},
|
||||||
|
timeout=15,
|
||||||
|
)
|
||||||
|
data = response.json()
|
||||||
|
except (httpx.HTTPError, ValueError) as exc:
|
||||||
|
raise WechatVirtualPaymentError("微信登录态交换失败,请稍后重试") from exc
|
||||||
|
if response.status_code >= 400 or data.get("errcode"):
|
||||||
|
raise WechatVirtualPaymentError(data.get("errmsg") or "微信登录态交换失败")
|
||||||
|
openid = str(data.get("openid") or "").strip()
|
||||||
|
session_key = str(data.get("session_key") or "").strip()
|
||||||
|
if not openid or not session_key:
|
||||||
|
raise WechatVirtualPaymentError("微信未返回完整登录态")
|
||||||
|
return {"openid": openid, "session_key": session_key}
|
||||||
|
|
||||||
|
|
||||||
|
def verify_callback_signature(signature: str, timestamp: str, nonce: str) -> bool:
|
||||||
|
token = os.getenv("WECHAT_VIRTUAL_CALLBACK_TOKEN", "").strip()
|
||||||
|
if not token or not signature or not timestamp or not nonce:
|
||||||
|
return False
|
||||||
|
source = "".join(sorted([token, timestamp, nonce]))
|
||||||
|
expected = hashlib.sha1(source.encode("utf-8")).hexdigest()
|
||||||
|
return hmac.compare_digest(signature, expected)
|
||||||
|
|
||||||
|
|
||||||
|
def _xml_value(element: ET.Element) -> Any:
|
||||||
|
children = list(element)
|
||||||
|
if not children:
|
||||||
|
return element.text or ""
|
||||||
|
return {child.tag: _xml_value(child) for child in children}
|
||||||
|
|
||||||
|
|
||||||
|
def parse_callback_body(body: bytes) -> dict[str, Any]:
|
||||||
|
text = body.decode("utf-8", errors="replace").strip()
|
||||||
|
if not text:
|
||||||
|
return {}
|
||||||
|
try:
|
||||||
|
payload = json.loads(text)
|
||||||
|
if isinstance(payload, dict):
|
||||||
|
return payload
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
pass
|
||||||
|
try:
|
||||||
|
parsed = _xml_value(ET.fromstring(text))
|
||||||
|
except ET.ParseError as exc:
|
||||||
|
raise WechatVirtualPaymentError("微信虚拟支付回调格式不正确") from exc
|
||||||
|
return parsed if isinstance(parsed, dict) else {}
|
||||||
|
|
||||||
|
|
||||||
|
def case_get(payload: Any, key: str) -> Any:
|
||||||
|
if not isinstance(payload, dict):
|
||||||
|
return None
|
||||||
|
lowered = key.lower()
|
||||||
|
for current, value in payload.items():
|
||||||
|
if str(current).lower() == lowered:
|
||||||
|
return value
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def callback_value(payload: dict[str, Any], *path: str) -> Any:
|
||||||
|
current: Any = payload
|
||||||
|
for key in path:
|
||||||
|
current = case_get(current, key)
|
||||||
|
if current is None:
|
||||||
|
break
|
||||||
|
return current
|
||||||
|
|
||||||
|
|
||||||
|
_access_token_cache: tuple[str, float] = ("", 0)
|
||||||
|
|
||||||
|
|
||||||
|
def _access_token() -> str:
|
||||||
|
global _access_token_cache
|
||||||
|
token, expires_at = _access_token_cache
|
||||||
|
if token and expires_at > time.monotonic() + 60:
|
||||||
|
return token
|
||||||
|
app_id = os.getenv("WECHAT_MP_APP_ID", "").strip()
|
||||||
|
app_secret = os.getenv("WECHAT_MP_APP_SECRET", "").strip()
|
||||||
|
if not app_id or not app_secret:
|
||||||
|
raise WechatVirtualPaymentError("微信小程序服务端配置不完整")
|
||||||
|
try:
|
||||||
|
response = httpx.get(
|
||||||
|
"https://api.weixin.qq.com/cgi-bin/token",
|
||||||
|
params={"grant_type": "client_credential", "appid": app_id, "secret": app_secret},
|
||||||
|
timeout=15,
|
||||||
|
)
|
||||||
|
data = response.json()
|
||||||
|
except (httpx.HTTPError, ValueError) as exc:
|
||||||
|
raise WechatVirtualPaymentError("微信 access_token 获取失败") from exc
|
||||||
|
if response.status_code >= 400 or data.get("errcode"):
|
||||||
|
raise WechatVirtualPaymentError(data.get("errmsg") or "微信 access_token 获取失败")
|
||||||
|
token = str(data.get("access_token") or "")
|
||||||
|
if not token:
|
||||||
|
raise WechatVirtualPaymentError("微信未返回 access_token")
|
||||||
|
_access_token_cache = (token, time.monotonic() + int(data.get("expires_in") or 7200))
|
||||||
|
return token
|
||||||
|
|
||||||
|
|
||||||
|
def call_xpay(uri: str, payload: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
env = int(payload.get("env", virtual_env()))
|
||||||
|
app_key = _app_key(env)
|
||||||
|
if not app_key:
|
||||||
|
raise WechatVirtualPaymentError("微信虚拟支付 AppKey 未配置")
|
||||||
|
body = json_compact(payload)
|
||||||
|
pay_sig = hmac_sha256_hex(app_key, f"{uri}&{body}")
|
||||||
|
try:
|
||||||
|
response = httpx.post(
|
||||||
|
f"https://api.weixin.qq.com{uri}",
|
||||||
|
params={"access_token": _access_token(), "pay_sig": pay_sig},
|
||||||
|
content=body.encode("utf-8"),
|
||||||
|
headers={"Content-Type": "application/json"},
|
||||||
|
timeout=20,
|
||||||
|
)
|
||||||
|
data = response.json()
|
||||||
|
except (httpx.HTTPError, ValueError) as exc:
|
||||||
|
raise WechatVirtualPaymentError("微信虚拟支付服务暂时不可用") from exc
|
||||||
|
if response.status_code >= 400 or data.get("errcode") not in (None, 0):
|
||||||
|
raise WechatVirtualPaymentError(data.get("errmsg") or "微信虚拟支付请求失败")
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
|
def request_refund(*, openid: str, order_no: str, refund_no: str, amount_cents: int, reason: int = 3) -> dict[str, Any]:
|
||||||
|
return call_xpay("/xpay/refund_order", {
|
||||||
|
"openid": openid,
|
||||||
|
"order_id": order_no,
|
||||||
|
"refund_order_id": refund_no,
|
||||||
|
"left_fee": amount_cents,
|
||||||
|
"refund_fee": amount_cents,
|
||||||
|
"biz_meta": json_compact({"orderNo": order_no}),
|
||||||
|
"refund_reason": int(reason),
|
||||||
|
"req_from": 1,
|
||||||
|
"env": virtual_env(),
|
||||||
|
})
|
||||||
|
|
||||||
|
|
||||||
|
def query_order(*, openid: str, order_no: str) -> dict[str, Any]:
|
||||||
|
return call_xpay("/xpay/query_order", {
|
||||||
|
"openid": openid,
|
||||||
|
"order_id": order_no,
|
||||||
|
"env": virtual_env(),
|
||||||
|
})
|
||||||
@@ -5,6 +5,10 @@ from database import init_db, SessionLocal
|
|||||||
from models import (
|
from models import (
|
||||||
Authorization,
|
Authorization,
|
||||||
Avatar,
|
Avatar,
|
||||||
|
ChatAttachment,
|
||||||
|
InvoiceApplication,
|
||||||
|
PaymentRefund,
|
||||||
|
PaymentTransaction,
|
||||||
TakeoverCursor,
|
TakeoverCursor,
|
||||||
TakeoverMessage,
|
TakeoverMessage,
|
||||||
TakeoverReplyTask,
|
TakeoverReplyTask,
|
||||||
@@ -95,6 +99,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)
|
||||||
@@ -111,6 +118,21 @@ def authorization_context():
|
|||||||
synchronize_session=False
|
synchronize_session=False
|
||||||
)
|
)
|
||||||
user_ids = [owner.id, other.id]
|
user_ids = [owner.id, other.id]
|
||||||
|
order_numbers = [
|
||||||
|
row[0] for row in db.query(TokenPaymentOrder.order_no).filter(
|
||||||
|
TokenPaymentOrder.user_id.in_(user_ids)
|
||||||
|
).all()
|
||||||
|
]
|
||||||
|
if order_numbers:
|
||||||
|
db.query(InvoiceApplication).filter(InvoiceApplication.order_no.in_(order_numbers)).delete(
|
||||||
|
synchronize_session=False
|
||||||
|
)
|
||||||
|
db.query(PaymentRefund).filter(PaymentRefund.order_no.in_(order_numbers)).delete(
|
||||||
|
synchronize_session=False
|
||||||
|
)
|
||||||
|
db.query(PaymentTransaction).filter(PaymentTransaction.order_no.in_(order_numbers)).delete(
|
||||||
|
synchronize_session=False
|
||||||
|
)
|
||||||
db.query(TokenPaymentOrder).filter(TokenPaymentOrder.user_id.in_(user_ids)).delete(
|
db.query(TokenPaymentOrder).filter(TokenPaymentOrder.user_id.in_(user_ids)).delete(
|
||||||
synchronize_session=False
|
synchronize_session=False
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -0,0 +1,71 @@
|
|||||||
|
import ipaddress
|
||||||
|
import json
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from services.boxim_image_service import (
|
||||||
|
BoxIMImageError,
|
||||||
|
download_boxim_image,
|
||||||
|
parse_boxim_image_url,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_parse_boxim_image_prefers_origin_and_supports_relative_url():
|
||||||
|
content = json.dumps({"originUrl": "/files/original.png", "thumbUrl": "/thumb.png"})
|
||||||
|
assert parse_boxim_image_url(content, base_url="https://im.example/api") == (
|
||||||
|
"https://im.example/files/original.png"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_download_boxim_image_streams_public_https(monkeypatch):
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"services.boxim_image_service._resolved_addresses",
|
||||||
|
lambda _host, _port: {ipaddress.ip_address("8.8.8.8")},
|
||||||
|
)
|
||||||
|
transport = httpx.MockTransport(
|
||||||
|
lambda request: httpx.Response(
|
||||||
|
200,
|
||||||
|
headers={"content-type": "image/png"},
|
||||||
|
content=b"png-bytes",
|
||||||
|
request=request,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
image = download_boxim_image(
|
||||||
|
json.dumps({"originUrl": "https://cdn.example/case%20photo.png"}),
|
||||||
|
transport=transport,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert image.content == b"png-bytes"
|
||||||
|
assert image.filename == "case photo.png"
|
||||||
|
assert image.mime_type == "image/png"
|
||||||
|
|
||||||
|
|
||||||
|
def test_download_boxim_image_rejects_private_network_url():
|
||||||
|
with pytest.raises(BoxIMImageError, match="受限网络"):
|
||||||
|
download_boxim_image(
|
||||||
|
json.dumps({"originUrl": "https://127.0.0.1/private.png"}),
|
||||||
|
transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request)),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_download_boxim_image_stops_oversized_stream(monkeypatch):
|
||||||
|
monkeypatch.setenv("CHAT_IMAGE_MAX_BYTES", "1024")
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"services.boxim_image_service._resolved_addresses",
|
||||||
|
lambda _host, _port: {ipaddress.ip_address("8.8.8.8")},
|
||||||
|
)
|
||||||
|
transport = httpx.MockTransport(
|
||||||
|
lambda request: httpx.Response(
|
||||||
|
200,
|
||||||
|
headers={"content-length": "2048"},
|
||||||
|
request=request,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(BoxIMImageError, match="超过大小限制"):
|
||||||
|
download_boxim_image(
|
||||||
|
json.dumps({"originUrl": "https://cdn.example/large.png"}),
|
||||||
|
transport=transport,
|
||||||
|
)
|
||||||
@@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ from routers.chat import (
|
|||||||
_match_standard_qa,
|
_match_standard_qa,
|
||||||
_public_avatar_payload,
|
_public_avatar_payload,
|
||||||
_qa_requires_language_adaptation,
|
_qa_requires_language_adaptation,
|
||||||
|
_qa_requires_per_turn_rendering,
|
||||||
_require_owned_avatar,
|
_require_owned_avatar,
|
||||||
_resolve_reply,
|
_resolve_reply,
|
||||||
)
|
)
|
||||||
@@ -86,6 +87,41 @@ class ChatOrchestrationTests(unittest.TestCase):
|
|||||||
self.assertTrue(_qa_requires_language_adaptation("안녕하세요", "你好"))
|
self.assertTrue(_qa_requires_language_adaptation("안녕하세요", "你好"))
|
||||||
self.assertFalse(_qa_requires_language_adaptation("你好", "您好"))
|
self.assertFalse(_qa_requires_language_adaptation("你好", "您好"))
|
||||||
|
|
||||||
|
def test_conversation_qa_is_rendered_for_the_current_turn_language(self):
|
||||||
|
history = [SimpleNamespace(role="user", content="Please answer in English.")]
|
||||||
|
self.assertTrue(_qa_requires_per_turn_rendering("Quelle est votre adresse ?", "Our address is Test Road 1.", history))
|
||||||
|
|
||||||
|
fake_model = Mock(return_value="Notre adresse est Test Road 1.")
|
||||||
|
result = _resolve_reply(
|
||||||
|
None,
|
||||||
|
self.avatar,
|
||||||
|
"Quelle est votre adresse ?",
|
||||||
|
history,
|
||||||
|
qa_pairs=[SimpleNamespace(question="Quelle est votre adresse ?", answer="Our address is Test Road 1.", enabled=True)],
|
||||||
|
search_fn=Mock(),
|
||||||
|
model_client=fake_model,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(result["source"], "qa")
|
||||||
|
self.assertEqual(result["answer"], "Notre adresse est Test Road 1.")
|
||||||
|
messages = fake_model.call_args.kwargs["messages"]
|
||||||
|
self.assertEqual(messages[-1], {"role": "user", "content": "Quelle est votre adresse ?"})
|
||||||
|
self.assertEqual(messages[-2]["role"], "system")
|
||||||
|
self.assertIn("本轮语言覆盖指令", messages[-2]["content"])
|
||||||
|
self.assertIn("不要沿用上一轮语言", messages[-2]["content"])
|
||||||
|
|
||||||
|
def test_latest_user_message_has_an_adjacent_language_override(self):
|
||||||
|
history = [
|
||||||
|
SimpleNamespace(role="user", content="请用中文回答"),
|
||||||
|
SimpleNamespace(role="assistant", content="好的,请问有什么可以帮你?"),
|
||||||
|
]
|
||||||
|
messages = _build_prompt(self.avatar, history, "What can you help me with?", [])
|
||||||
|
|
||||||
|
self.assertEqual(messages[-1], {"role": "user", "content": "What can you help me with?"})
|
||||||
|
self.assertEqual(messages[-2]["role"], "system")
|
||||||
|
self.assertIn("最新用户消息", messages[-2]["content"])
|
||||||
|
self.assertIn("立即切换到相同语言", messages[-2]["content"])
|
||||||
|
|
||||||
def test_conversational_paraphrase_matches_standard_qa(self):
|
def test_conversational_paraphrase_matches_standard_qa(self):
|
||||||
for question in ("请问一下,你们公司在哪里呀?", "请问去你们那边怎么走"):
|
for question in ("请问一下,你们公司在哪里呀?", "请问去你们那边怎么走"):
|
||||||
with self.subTest(question=question):
|
with self.subTest(question=question):
|
||||||
|
|||||||
@@ -48,9 +48,12 @@ 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 = []
|
||||||
|
progress_updates = []
|
||||||
|
|
||||||
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,15 +64,29 @@ 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",
|
||||||
}), patch("embeddings.urllib.request.urlopen", side_effect=fake_urlopen):
|
}), patch("embeddings.urllib.request.urlopen", side_effect=fake_urlopen):
|
||||||
result = embeddings.embed(texts)
|
result = embeddings.embed(
|
||||||
|
texts,
|
||||||
|
on_progress=lambda completed, total: progress_updates.append((completed, total)),
|
||||||
|
)
|
||||||
|
|
||||||
self.assertEqual(batch_sizes, [10, 4])
|
self.assertEqual(batch_sizes, [10, 4])
|
||||||
|
self.assertEqual(requested_urls, [
|
||||||
|
"https://embedding.example/v1/embeddings",
|
||||||
|
"https://embedding.example/v1/embeddings",
|
||||||
|
])
|
||||||
self.assertEqual(result, [[float(index)] for index in range(14)])
|
self.assertEqual(result, [[float(index)] for index in range(14)])
|
||||||
|
self.assertEqual(progress_updates, [(10, 14), (14, 14)])
|
||||||
|
|
||||||
|
def test_full_embedding_endpoint_is_not_modified(self):
|
||||||
|
self.assertEqual(
|
||||||
|
embeddings._embedding_endpoint("https://embedding.example/v1/embeddings/"),
|
||||||
|
"https://embedding.example/v1/embeddings",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
@@ -0,0 +1,33 @@
|
|||||||
|
import main
|
||||||
|
|
||||||
|
|
||||||
|
def test_health_reports_release_and_runtime_capabilities(monkeypatch):
|
||||||
|
monkeypatch.setenv("APP_GIT_SHA", "test-sha")
|
||||||
|
monkeypatch.setenv("APP_BUILD_TIME", "2026-09-09T00:00:00Z")
|
||||||
|
monkeypatch.setattr(main, "_runtime_checks", lambda: {
|
||||||
|
"database": True,
|
||||||
|
"uploads": True,
|
||||||
|
"pdfOcr": True,
|
||||||
|
})
|
||||||
|
|
||||||
|
response = main.health()
|
||||||
|
|
||||||
|
assert response["code"] == 200
|
||||||
|
assert response["data"]["status"] == "ok"
|
||||||
|
assert response["data"]["gitSha"] == "test-sha"
|
||||||
|
assert response["data"]["buildTime"] == "2026-09-09T00:00:00Z"
|
||||||
|
assert response["data"]["checks"] == {
|
||||||
|
"database": True,
|
||||||
|
"uploads": True,
|
||||||
|
"pdfOcr": True,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_health_is_degraded_when_a_required_capability_is_missing(monkeypatch):
|
||||||
|
monkeypatch.setattr(main, "_runtime_checks", lambda: {
|
||||||
|
"database": True,
|
||||||
|
"uploads": True,
|
||||||
|
"pdfOcr": False,
|
||||||
|
})
|
||||||
|
|
||||||
|
assert main.health()["data"]["status"] == "degraded"
|
||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import os
|
||||||
from unittest.mock import Mock, patch
|
from unittest.mock import Mock, patch
|
||||||
|
|
||||||
from services.huihui_payment import HuihuiPaymentClient
|
from services.huihui_payment import HuihuiPaymentClient
|
||||||
@@ -48,3 +49,35 @@ def test_create_payment_uses_huihui_payment_v3_contract():
|
|||||||
assert body["payWay"] == "APP"
|
assert body["payWay"] == "APP"
|
||||||
assert body["masterOrderAmt"] == "10.00"
|
assert body["masterOrderAmt"] == "10.00"
|
||||||
assert body["payAmt"] == 10.0
|
assert body["payAmt"] == 10.0
|
||||||
|
|
||||||
|
|
||||||
|
def test_request_refund_uses_configured_huihui_endpoint_without_exposing_secret():
|
||||||
|
client = HuihuiPaymentClient({
|
||||||
|
"HUIHUI_PAYMENT_BASE_URL": "https://open.example/api/payment-v3",
|
||||||
|
"HUIHUI_APP_ID": "app-id",
|
||||||
|
"HUIHUI_ACCESS_ID": "access-id",
|
||||||
|
"HUIHUI_ACCESS_SECRET": "access-secret",
|
||||||
|
})
|
||||||
|
response = Mock(status_code=200)
|
||||||
|
response.json.return_value = {"code": 200, "data": {"status": "PROCESSING", "refundNo": "provider-rf"}}
|
||||||
|
with patch.dict(os.environ, {"HUIHUI_PAYMENT_REFUND_PATH": "/payment/refund"}), patch(
|
||||||
|
"services.huihui_payment.httpx.post", return_value=response
|
||||||
|
) as post:
|
||||||
|
result = client.request_refund(
|
||||||
|
huihui_token="user-token",
|
||||||
|
huihui_user_id="user-id",
|
||||||
|
order_no="AV1",
|
||||||
|
refund_no="RF1",
|
||||||
|
amount="10.00",
|
||||||
|
reason="用户申请",
|
||||||
|
)
|
||||||
|
assert result["refundNo"] == "provider-rf"
|
||||||
|
assert post.call_args.args[0] == "https://open.example/api/payment-v3/payment/refund"
|
||||||
|
assert post.call_args.kwargs["json"] == {
|
||||||
|
"appId": "app-id",
|
||||||
|
"masterOrderNo": "AV1",
|
||||||
|
"refundOrderNo": "RF1",
|
||||||
|
"refundAmt": 10.0,
|
||||||
|
"refundReason": "用户申请",
|
||||||
|
}
|
||||||
|
assert "accessSecret" not in post.call_args.kwargs["params"]
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from database import SessionLocal
|
|||||||
from main import app
|
from main import app
|
||||||
from models import Avatar, KnowledgeChunk, KnowledgeDoc, QAPair
|
from models import Avatar, KnowledgeChunk, KnowledgeDoc, QAPair
|
||||||
from routers.knowledge import _doc_payload
|
from routers.knowledge import _doc_payload
|
||||||
|
from services.knowledge_vectorizer import knowledge_vectorizer
|
||||||
|
|
||||||
|
|
||||||
client = TestClient(app)
|
client = TestClient(app)
|
||||||
@@ -31,14 +32,14 @@ def test_doc_payload_reports_whether_the_persisted_file_exists(tmp_path: Path):
|
|||||||
assert _doc_payload(doc)["filePresent"] is True
|
assert _doc_payload(doc)["filePresent"] is True
|
||||||
|
|
||||||
|
|
||||||
def test_upload_marks_vectorization_failure_instead_of_staying_processing(
|
def test_upload_returns_before_background_vectorization(
|
||||||
tmp_path: Path,
|
tmp_path: Path,
|
||||||
authorization_context,
|
authorization_context,
|
||||||
):
|
):
|
||||||
context = authorization_context
|
context = authorization_context
|
||||||
with (
|
with (
|
||||||
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
||||||
patch("routers.knowledge.embeddings.embed", side_effect=RuntimeError("provider unavailable")),
|
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
|
||||||
):
|
):
|
||||||
response = client.post(
|
response = client.post(
|
||||||
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
|
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
|
||||||
@@ -47,14 +48,15 @@ def test_upload_marks_vectorization_failure_instead_of_staying_processing(
|
|||||||
)
|
)
|
||||||
|
|
||||||
payload = response.json()["data"]
|
payload = response.json()["data"]
|
||||||
assert payload["status"] == "failed"
|
assert payload["status"] == "parsing"
|
||||||
assert payload["vectorized"] is False
|
assert payload["vectorized"] is False
|
||||||
assert payload["chunkCount"] == 0
|
assert payload["chunkCount"] == 0
|
||||||
|
enqueue.assert_called_once_with(payload["id"])
|
||||||
|
|
||||||
db = SessionLocal()
|
db = SessionLocal()
|
||||||
try:
|
try:
|
||||||
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
|
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
|
||||||
assert stored.status == "failed"
|
assert stored.status == "parsing"
|
||||||
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 0
|
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 0
|
||||||
db.delete(stored)
|
db.delete(stored)
|
||||||
db.commit()
|
db.commit()
|
||||||
@@ -62,14 +64,115 @@ def test_upload_marks_vectorization_failure_instead_of_staying_processing(
|
|||||||
db.close()
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
def test_markdown_upload_commits_ready_document_and_chunks_together(
|
def test_upload_rejects_oversize_file_before_queuing_indexing(
|
||||||
tmp_path: Path,
|
tmp_path: Path,
|
||||||
authorization_context,
|
authorization_context,
|
||||||
):
|
):
|
||||||
context = authorization_context
|
context = authorization_context
|
||||||
with (
|
with (
|
||||||
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
||||||
patch("routers.knowledge.embeddings.embed", return_value=[[1.0, 0.0]]),
|
patch("routers.knowledge.MAX_UPLOAD_BYTES", 4),
|
||||||
|
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
|
||||||
|
):
|
||||||
|
response = client.post(
|
||||||
|
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
files={"file": ("oversize.md", b"12345", "text/markdown")},
|
||||||
|
)
|
||||||
|
|
||||||
|
payload = response.json()
|
||||||
|
assert payload["code"] == 400
|
||||||
|
assert payload["message"] == "文件不能超过 50MB"
|
||||||
|
enqueue.assert_not_called()
|
||||||
|
assert not list((tmp_path / context["avatar"].id).glob("*"))
|
||||||
|
|
||||||
|
|
||||||
|
def test_multipart_upload_reassembles_file_before_queuing_indexing(
|
||||||
|
tmp_path: Path,
|
||||||
|
authorization_context,
|
||||||
|
):
|
||||||
|
context = authorization_context
|
||||||
|
avatar_id = context["avatar"].id
|
||||||
|
content = b"0123456789"
|
||||||
|
with (
|
||||||
|
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
||||||
|
patch("routers.knowledge.MULTIPART_CHUNK_BYTES", 4),
|
||||||
|
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
|
||||||
|
):
|
||||||
|
created = client.post(
|
||||||
|
f"/api/avatar/{avatar_id}/knowledge/uploads",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
json={"filename": "large.pdf", "fileSize": len(content), "totalChunks": 3},
|
||||||
|
).json()["data"]
|
||||||
|
|
||||||
|
for index, chunk in enumerate((content[:4], content[4:8], content[8:])):
|
||||||
|
response = client.post(
|
||||||
|
f"/api/avatar/{avatar_id}/knowledge/uploads/{created['uploadId']}/chunks/{index}",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
files={"file": (f"chunk-{index}", chunk, "application/octet-stream")},
|
||||||
|
)
|
||||||
|
assert response.json()["code"] == 200
|
||||||
|
|
||||||
|
completed = client.post(
|
||||||
|
f"/api/avatar/{avatar_id}/knowledge/uploads/{created['uploadId']}/complete",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
).json()["data"]
|
||||||
|
|
||||||
|
assert completed["status"] == "parsing"
|
||||||
|
assert completed["fileSize"] == len(content)
|
||||||
|
enqueue.assert_called_once_with(completed["id"])
|
||||||
|
stored_path = tmp_path / avatar_id / Path(completed["fileUrl"]).name
|
||||||
|
assert stored_path.read_bytes() == content
|
||||||
|
assert not (tmp_path / ".multipart" / avatar_id / created["uploadId"]).exists()
|
||||||
|
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == completed["id"]).one()
|
||||||
|
db.delete(stored)
|
||||||
|
db.commit()
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
def test_multipart_upload_rejects_incomplete_parts(
|
||||||
|
tmp_path: Path,
|
||||||
|
authorization_context,
|
||||||
|
):
|
||||||
|
context = authorization_context
|
||||||
|
avatar_id = context["avatar"].id
|
||||||
|
with (
|
||||||
|
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
||||||
|
patch("routers.knowledge.MULTIPART_CHUNK_BYTES", 4),
|
||||||
|
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
|
||||||
|
):
|
||||||
|
created = client.post(
|
||||||
|
f"/api/avatar/{avatar_id}/knowledge/uploads",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
json={"filename": "large.pdf", "fileSize": 6, "totalChunks": 2},
|
||||||
|
).json()["data"]
|
||||||
|
client.post(
|
||||||
|
f"/api/avatar/{avatar_id}/knowledge/uploads/{created['uploadId']}/chunks/0",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
files={"file": ("chunk-0", b"0123", "application/octet-stream")},
|
||||||
|
)
|
||||||
|
response = client.post(
|
||||||
|
f"/api/avatar/{avatar_id}/knowledge/uploads/{created['uploadId']}/complete",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.json()["code"] == 400
|
||||||
|
assert response.json()["message"] == "文件分片尚未上传完整"
|
||||||
|
enqueue.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
def test_background_vectorizer_commits_ready_document_and_chunks_together(
|
||||||
|
tmp_path: Path,
|
||||||
|
authorization_context,
|
||||||
|
):
|
||||||
|
context = authorization_context
|
||||||
|
with (
|
||||||
|
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
||||||
|
patch("routers.knowledge.knowledge_vectorizer.enqueue"),
|
||||||
):
|
):
|
||||||
response = client.post(
|
response = client.post(
|
||||||
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
|
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
|
||||||
@@ -78,14 +181,21 @@ def test_markdown_upload_commits_ready_document_and_chunks_together(
|
|||||||
)
|
)
|
||||||
|
|
||||||
payload = response.json()["data"]
|
payload = response.json()["data"]
|
||||||
assert payload["status"] == "ready"
|
assert payload["status"] == "parsing"
|
||||||
assert payload["vectorized"] is True
|
with (
|
||||||
assert payload["chunkCount"] == 1
|
patch("services.knowledge_vectorizer.UPLOAD_DIR", str(tmp_path)),
|
||||||
|
patch("services.knowledge_vectorizer.embeddings.embed", return_value=[[1.0, 0.0]]),
|
||||||
|
):
|
||||||
|
knowledge_vectorizer.vectorize_document(payload["id"])
|
||||||
|
|
||||||
db = SessionLocal()
|
db = SessionLocal()
|
||||||
try:
|
try:
|
||||||
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
|
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
|
||||||
assert stored.status == "ready"
|
assert stored.status == "ready"
|
||||||
|
assert stored.vectorized is True
|
||||||
|
assert stored.chunk_count == 1
|
||||||
|
assert stored.index_stage == "ready"
|
||||||
|
assert stored.index_progress == 100
|
||||||
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 1
|
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 1
|
||||||
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).delete()
|
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).delete()
|
||||||
db.delete(stored)
|
db.delete(stored)
|
||||||
@@ -94,6 +204,134 @@ def test_markdown_upload_commits_ready_document_and_chunks_together(
|
|||||||
db.close()
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
def test_background_vectorizer_keeps_failure_reason_for_retry(
|
||||||
|
tmp_path: Path,
|
||||||
|
authorization_context,
|
||||||
|
):
|
||||||
|
context = authorization_context
|
||||||
|
with (
|
||||||
|
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
||||||
|
patch("routers.knowledge.knowledge_vectorizer.enqueue"),
|
||||||
|
):
|
||||||
|
response = client.post(
|
||||||
|
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
files={"file": ("knowledge.md", b"# Knowledge\n\nTest content", "text/markdown")},
|
||||||
|
)
|
||||||
|
|
||||||
|
payload = response.json()["data"]
|
||||||
|
with (
|
||||||
|
patch("services.knowledge_vectorizer.UPLOAD_DIR", str(tmp_path)),
|
||||||
|
patch("services.knowledge_vectorizer.embeddings.embed", side_effect=RuntimeError("provider unavailable")),
|
||||||
|
):
|
||||||
|
knowledge_vectorizer.vectorize_document(payload["id"])
|
||||||
|
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
|
||||||
|
assert stored.status == "failed"
|
||||||
|
assert stored.error_message == "provider unavailable"
|
||||||
|
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 0
|
||||||
|
db.delete(stored)
|
||||||
|
db.commit()
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
def test_background_vectorizer_uses_ocr_for_image_only_pdf(
|
||||||
|
tmp_path: Path,
|
||||||
|
authorization_context,
|
||||||
|
):
|
||||||
|
context = authorization_context
|
||||||
|
with (
|
||||||
|
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
||||||
|
patch("routers.knowledge.knowledge_vectorizer.enqueue"),
|
||||||
|
):
|
||||||
|
response = client.post(
|
||||||
|
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
files={"file": ("scanned.pdf", b"image-only-pdf", "application/pdf")},
|
||||||
|
)
|
||||||
|
|
||||||
|
payload = response.json()["data"]
|
||||||
|
progress = []
|
||||||
|
with (
|
||||||
|
patch("services.knowledge_vectorizer.UPLOAD_DIR", str(tmp_path)),
|
||||||
|
patch("services.knowledge_vectorizer.embeddings.extract_text", return_value=""),
|
||||||
|
patch(
|
||||||
|
"services.knowledge_vectorizer.extract_scanned_pdf_text",
|
||||||
|
side_effect=lambda _db, _avatar, _path, on_progress: (
|
||||||
|
on_progress(1, 2), on_progress(2, 2), "扫描页文字"
|
||||||
|
)[-1],
|
||||||
|
) as ocr,
|
||||||
|
patch("services.knowledge_vectorizer.embeddings.embed", return_value=[[1.0, 0.0]]),
|
||||||
|
patch.object(knowledge_vectorizer, "_set_progress", wraps=knowledge_vectorizer._set_progress) as set_progress,
|
||||||
|
):
|
||||||
|
knowledge_vectorizer.vectorize_document(payload["id"])
|
||||||
|
progress = [(call.args[2], call.args[3]) for call in set_progress.call_args_list]
|
||||||
|
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
|
||||||
|
assert stored.status == "ready"
|
||||||
|
assert stored.chunk_count == 1
|
||||||
|
assert ("ocr", 18) in progress
|
||||||
|
assert ("ocr", 28) in progress
|
||||||
|
ocr.assert_called_once()
|
||||||
|
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).delete()
|
||||||
|
db.delete(stored)
|
||||||
|
db.commit()
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
def test_retry_queues_a_failed_document_again(
|
||||||
|
tmp_path: Path,
|
||||||
|
authorization_context,
|
||||||
|
):
|
||||||
|
context = authorization_context
|
||||||
|
document_id = f"retry-doc-{context['suffix']}"
|
||||||
|
avatar_dir = tmp_path / context["avatar"].id
|
||||||
|
avatar_dir.mkdir()
|
||||||
|
(avatar_dir / "retry.md").write_text("retry content", encoding="utf-8")
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
db.add(
|
||||||
|
KnowledgeDoc(
|
||||||
|
id=document_id,
|
||||||
|
avatar_id=context["avatar"].id,
|
||||||
|
filename="retry.md",
|
||||||
|
file_type="md",
|
||||||
|
file_url=f"/api/files/{context['avatar'].id}/retry.md",
|
||||||
|
status="failed",
|
||||||
|
error_message="provider unavailable",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
db.commit()
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
||||||
|
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
|
||||||
|
):
|
||||||
|
response = client.post(
|
||||||
|
f"/api/avatar/{context['avatar'].id}/knowledge/docs/{document_id}/retry",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
)
|
||||||
|
|
||||||
|
payload = response.json()["data"]
|
||||||
|
assert payload["status"] == "parsing"
|
||||||
|
assert payload["errorMessage"] == ""
|
||||||
|
enqueue.assert_called_once_with(document_id)
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
db.query(KnowledgeDoc).filter(KnowledgeDoc.id == document_id).delete()
|
||||||
|
db.commit()
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
def test_each_avatar_has_an_independent_document_and_qa_scope(authorization_context):
|
def test_each_avatar_has_an_independent_document_and_qa_scope(authorization_context):
|
||||||
context = authorization_context
|
context = authorization_context
|
||||||
first_avatar_id = context["avatar"].id
|
first_avatar_id = context["avatar"].id
|
||||||
|
|||||||
@@ -0,0 +1,104 @@
|
|||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from services.pdf_ocr_service import extract_scanned_pdf_text
|
||||||
|
|
||||||
|
|
||||||
|
class FakePixmap:
|
||||||
|
def tobytes(self, *_args, **_kwargs):
|
||||||
|
return b"jpeg-page"
|
||||||
|
|
||||||
|
|
||||||
|
class FakePage:
|
||||||
|
def get_pixmap(self, **_kwargs):
|
||||||
|
return FakePixmap()
|
||||||
|
|
||||||
|
|
||||||
|
class FakeDocument:
|
||||||
|
page_count = 2
|
||||||
|
|
||||||
|
def __enter__(self):
|
||||||
|
return self
|
||||||
|
|
||||||
|
def __exit__(self, *_args):
|
||||||
|
return None
|
||||||
|
|
||||||
|
def load_page(self, _index):
|
||||||
|
return FakePage()
|
||||||
|
|
||||||
|
|
||||||
|
def test_scanned_pdf_ocr_preserves_page_order_and_reports_progress(monkeypatch):
|
||||||
|
fake_pymupdf = SimpleNamespace(
|
||||||
|
open=lambda _path: FakeDocument(),
|
||||||
|
Matrix=lambda x, y: (x, y),
|
||||||
|
csRGB="rgb",
|
||||||
|
)
|
||||||
|
monkeypatch.setitem(__import__("sys").modules, "pymupdf", fake_pymupdf)
|
||||||
|
progress = []
|
||||||
|
reservation = SimpleNamespace()
|
||||||
|
config = SimpleNamespace(
|
||||||
|
api_key="configured",
|
||||||
|
ocr_model="qwen-vl-ocr",
|
||||||
|
vision_model="vision",
|
||||||
|
vision_max_tokens=2048,
|
||||||
|
)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch("services.pdf_ocr_service.get_chat_model_config", return_value=config),
|
||||||
|
patch("services.pdf_ocr_service.prepare_image", return_value=SimpleNamespace()),
|
||||||
|
patch(
|
||||||
|
"services.pdf_ocr_service.call_vision_model",
|
||||||
|
side_effect=[
|
||||||
|
{"content": "第一页文字", "usage": {"total_tokens": 10}},
|
||||||
|
{"content": "第二页文字", "usage": {"total_tokens": 12}},
|
||||||
|
],
|
||||||
|
),
|
||||||
|
patch("services.pdf_ocr_service.reserve_avatar_tokens", return_value=reservation) as reserve,
|
||||||
|
patch("services.pdf_ocr_service.settle_reservation") as settle,
|
||||||
|
):
|
||||||
|
text = extract_scanned_pdf_text(
|
||||||
|
MagicMock(),
|
||||||
|
SimpleNamespace(id="avatar-1"),
|
||||||
|
"/tmp/scanned.pdf",
|
||||||
|
on_progress=lambda done, total: progress.append((done, total)),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert text == "[第 1 页]\n第一页文字\n\n[第 2 页]\n第二页文字"
|
||||||
|
assert progress == [(1, 2), (2, 2)]
|
||||||
|
assert reserve.call_count == 2
|
||||||
|
assert settle.call_count == 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_scanned_pdf_ocr_releases_tokens_after_retries_fail(monkeypatch):
|
||||||
|
fake_document = FakeDocument()
|
||||||
|
fake_document.page_count = 1
|
||||||
|
fake_pymupdf = SimpleNamespace(
|
||||||
|
open=lambda _path: fake_document,
|
||||||
|
Matrix=lambda x, y: (x, y),
|
||||||
|
csRGB="rgb",
|
||||||
|
)
|
||||||
|
monkeypatch.setitem(__import__("sys").modules, "pymupdf", fake_pymupdf)
|
||||||
|
monkeypatch.setenv("KNOWLEDGE_PDF_OCR_ATTEMPTS", "2")
|
||||||
|
reservation = SimpleNamespace()
|
||||||
|
config = SimpleNamespace(
|
||||||
|
api_key="configured",
|
||||||
|
ocr_model="qwen-vl-ocr",
|
||||||
|
vision_model="vision",
|
||||||
|
vision_max_tokens=2048,
|
||||||
|
)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch("services.pdf_ocr_service.get_chat_model_config", return_value=config),
|
||||||
|
patch("services.pdf_ocr_service.prepare_image", return_value=SimpleNamespace()),
|
||||||
|
patch("services.pdf_ocr_service.call_vision_model", side_effect=RuntimeError("timeout")) as call,
|
||||||
|
patch("services.pdf_ocr_service.reserve_avatar_tokens", return_value=reservation),
|
||||||
|
patch("services.pdf_ocr_service.release_reservation") as release,
|
||||||
|
patch("services.pdf_ocr_service.time.sleep"),
|
||||||
|
):
|
||||||
|
with pytest.raises(RuntimeError, match="第 1/1 页识别失败"):
|
||||||
|
extract_scanned_pdf_text(MagicMock(), SimpleNamespace(id="avatar-1"), "/tmp/scanned.pdf")
|
||||||
|
|
||||||
|
assert call.call_count == 2
|
||||||
|
release.assert_called_once()
|
||||||
@@ -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,8 +11,9 @@ from sqlalchemy import create_engine
|
|||||||
from sqlalchemy.orm import sessionmaker
|
from sqlalchemy.orm import sessionmaker
|
||||||
|
|
||||||
from database import Base
|
from database import Base
|
||||||
from models import Avatar, TakeoverCursor, TakeoverMessage, TakeoverReplyTask, User
|
from models import Avatar, ChatAttachment, TakeoverCursor, TakeoverMessage, TakeoverReplyTask, User
|
||||||
from services.boxim_client import BoxIMError
|
from services.boxim_client import BoxIMError
|
||||||
|
from services.boxim_image_service import DownloadedBoxIMImage
|
||||||
from services.takeover_service import (
|
from services.takeover_service import (
|
||||||
AVATAR_LOCAL_ID_PREFIX,
|
AVATAR_LOCAL_ID_PREFIX,
|
||||||
TakeoverService,
|
TakeoverService,
|
||||||
@@ -62,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(
|
||||||
@@ -150,6 +196,247 @@ async def test_incoming_message_is_prepared_then_sent_at_three_seconds(service_c
|
|||||||
db.close()
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_incoming_image_is_analyzed_and_used_in_takeover_reply(service_context):
|
||||||
|
session_factory, service, boxim, clock = service_context
|
||||||
|
await service.poll_and_process_messages()
|
||||||
|
boxim.messages.append(
|
||||||
|
{
|
||||||
|
"id": 111,
|
||||||
|
"localId": 111,
|
||||||
|
"sendId": 200,
|
||||||
|
"recvId": 100,
|
||||||
|
"sendTime": clock.millis(),
|
||||||
|
"type": 1,
|
||||||
|
"content": json.dumps(
|
||||||
|
{
|
||||||
|
"originUrl": "https://cdn.example/case.png",
|
||||||
|
"thumbUrl": "https://cdn.example/case-thumb.png",
|
||||||
|
}
|
||||||
|
),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
await service.poll_and_process_messages()
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
scheduled = db.query(TakeoverReplyTask).filter_by(trigger_message_id="111").one()
|
||||||
|
assert scheduled.status == "pending"
|
||||||
|
assert scheduled.prompt == "请看看这张图片。"
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
clock.advance(3)
|
||||||
|
|
||||||
|
def analyze(db, avatar, content, **kwargs):
|
||||||
|
assert content == b"image-content"
|
||||||
|
attachment = ChatAttachment(
|
||||||
|
avatar_id=avatar.id,
|
||||||
|
uploader_kind=kwargs["uploader_kind"],
|
||||||
|
filename=kwargs["filename"],
|
||||||
|
mime_type="image/jpeg",
|
||||||
|
file_size=len(content),
|
||||||
|
status="ready",
|
||||||
|
category="medical_document",
|
||||||
|
summary="一张门诊病例",
|
||||||
|
extracted_text="主诉:咳嗽三天",
|
||||||
|
structured_data={"medical": {"chief_complaint": "咳嗽三天"}},
|
||||||
|
warning="请核对原始资料",
|
||||||
|
expires_at=clock.now() + timedelta(hours=24),
|
||||||
|
)
|
||||||
|
db.add(attachment)
|
||||||
|
db.commit()
|
||||||
|
db.refresh(attachment)
|
||||||
|
return attachment
|
||||||
|
|
||||||
|
downloaded = DownloadedBoxIMImage(
|
||||||
|
content=b"image-content",
|
||||||
|
filename="case.png",
|
||||||
|
mime_type="image/png",
|
||||||
|
source_url="https://cdn.example/case.png",
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
patch("services.takeover_service.download_boxim_image", return_value=downloaded),
|
||||||
|
patch("routers.chat._analyze_image_bytes", side_effect=analyze) as analyzer,
|
||||||
|
patch("routers.chat._resolve_reply", return_value={"answer": "这份资料里写的是咳嗽三天。"}) as resolver,
|
||||||
|
):
|
||||||
|
await service.poll_and_process_messages()
|
||||||
|
|
||||||
|
analyzer.assert_called_once()
|
||||||
|
assert resolver.call_args.args[2] == "请看看这张图片。"
|
||||||
|
image_contexts = resolver.call_args.kwargs["image_contexts"]
|
||||||
|
assert image_contexts[0]["summary"] == "一张门诊病例"
|
||||||
|
assert image_contexts[0]["extractedText"] == "主诉:咳嗽三天"
|
||||||
|
assert [item["content"] for item in boxim.sent] == ["这份资料里写的是咳嗽三天。"]
|
||||||
|
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
event = db.query(TakeoverMessage).filter_by(boxim_message_id="111").one()
|
||||||
|
task = db.query(TakeoverReplyTask).filter_by(trigger_message_id="111").one()
|
||||||
|
assert event.attachment_id
|
||||||
|
assert db.get(ChatAttachment, event.attachment_id).uploader_kind == "boxim"
|
||||||
|
assert task.status == "sent"
|
||||||
|
with patch(
|
||||||
|
"services.takeover_service.download_boxim_image",
|
||||||
|
side_effect=AssertionError("cached image must not be downloaded again"),
|
||||||
|
):
|
||||||
|
cached = service._takeover_image_attachment(db, db.get(Avatar, "avatar-1"), event)
|
||||||
|
assert cached.id == event.attachment_id
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_followup_text_recovers_recent_image_recorded_without_task(service_context):
|
||||||
|
session_factory, service, boxim, clock = service_context
|
||||||
|
await service.poll_and_process_messages()
|
||||||
|
image_message = {
|
||||||
|
"id": 113,
|
||||||
|
"localId": 113,
|
||||||
|
"sendId": 200,
|
||||||
|
"recvId": 100,
|
||||||
|
"sendTime": clock.millis(),
|
||||||
|
"type": 1,
|
||||||
|
"content": json.dumps(
|
||||||
|
{
|
||||||
|
"originUrl": "https://cdn.example/case.png",
|
||||||
|
"thumbUrl": "https://cdn.example/case-thumb.png",
|
||||||
|
}
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
avatar = db.get(Avatar, "avatar-1")
|
||||||
|
service._record_message(db, avatar, "100", image_message, schedule_reply=False)
|
||||||
|
cursor = db.query(TakeoverCursor).one()
|
||||||
|
cursor.last_message_id = "113"
|
||||||
|
db.commit()
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
clock.advance(60)
|
||||||
|
boxim.messages.extend(
|
||||||
|
[
|
||||||
|
image_message,
|
||||||
|
{
|
||||||
|
"id": 114,
|
||||||
|
"localId": 114,
|
||||||
|
"sendId": 200,
|
||||||
|
"recvId": 100,
|
||||||
|
"sendTime": clock.millis(),
|
||||||
|
"type": 0,
|
||||||
|
"content": "请帮我看看这张图",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
await service.poll_messages()
|
||||||
|
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
task = db.query(TakeoverReplyTask).filter_by(trigger_message_id="114").one()
|
||||||
|
assert task.source_message_ids == ["113", "114"]
|
||||||
|
assert task.prompt == "请看看这张图片。\n请帮我看看这张图"
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
clock.advance(1)
|
||||||
|
boxim.messages.append(
|
||||||
|
{
|
||||||
|
"id": 115,
|
||||||
|
"localId": 115,
|
||||||
|
"sendId": 200,
|
||||||
|
"recvId": 100,
|
||||||
|
"sendTime": clock.millis(),
|
||||||
|
"type": 0,
|
||||||
|
"content": "图里写了什么",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
await service.poll_messages()
|
||||||
|
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
latest = db.query(TakeoverReplyTask).filter_by(trigger_message_id="115").one()
|
||||||
|
assert latest.source_message_ids == ["113", "114", "115"]
|
||||||
|
assert latest.source_message_ids.count("113") == 1
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_explicit_followup_reuses_handled_image_within_two_days(service_context):
|
||||||
|
session_factory, service, boxim, clock = service_context
|
||||||
|
await service.poll_and_process_messages()
|
||||||
|
image_message = {
|
||||||
|
"id": 116,
|
||||||
|
"localId": 116,
|
||||||
|
"sendId": 200,
|
||||||
|
"recvId": 100,
|
||||||
|
"sendTime": clock.millis(),
|
||||||
|
"type": 1,
|
||||||
|
"content": json.dumps({"originUrl": "https://cdn.example/handled-case.png"}),
|
||||||
|
}
|
||||||
|
boxim.messages.append(image_message)
|
||||||
|
await service.poll_messages()
|
||||||
|
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
image_task = db.query(TakeoverReplyTask).filter_by(trigger_message_id="116").one()
|
||||||
|
image_task.status = "sent"
|
||||||
|
image_task.sent_at = clock.now()
|
||||||
|
db.commit()
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
clock.advance(47 * 60 * 60)
|
||||||
|
boxim.messages.append(
|
||||||
|
{
|
||||||
|
"id": 117,
|
||||||
|
"localId": 117,
|
||||||
|
"sendId": 200,
|
||||||
|
"recvId": 100,
|
||||||
|
"sendTime": clock.millis(),
|
||||||
|
"type": 0,
|
||||||
|
"content": "重新看一下刚才那张病例图片",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
await service.poll_messages()
|
||||||
|
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
task = db.query(TakeoverReplyTask).filter_by(trigger_message_id="117").one()
|
||||||
|
assert task.source_message_ids == ["116", "117"]
|
||||||
|
assert task.prompt == "请看看这张图片。\n重新看一下刚才那张病例图片"
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_invalid_image_message_is_recorded_but_not_scheduled(service_context):
|
||||||
|
session_factory, service, boxim, clock = service_context
|
||||||
|
await service.poll_and_process_messages()
|
||||||
|
boxim.messages.append(
|
||||||
|
{
|
||||||
|
"id": 112,
|
||||||
|
"localId": 112,
|
||||||
|
"sendId": 200,
|
||||||
|
"recvId": 100,
|
||||||
|
"sendTime": clock.millis(),
|
||||||
|
"type": 1,
|
||||||
|
"content": json.dumps({"width": 100, "height": 100}),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
await service.poll_and_process_messages()
|
||||||
|
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
assert db.query(TakeoverMessage).filter_by(boxim_message_id="112").one()
|
||||||
|
assert db.query(TakeoverReplyTask).count() == 0
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_default_reply_delay_is_three_minutes(service_context):
|
async def test_default_reply_delay_is_three_minutes(service_context):
|
||||||
session_factory, service, boxim, clock = service_context
|
session_factory, service, boxim, clock = service_context
|
||||||
@@ -183,6 +470,87 @@ async def test_default_reply_delay_is_three_minutes(service_context):
|
|||||||
assert [item["content"] for item in boxim.sent] == ["好的"]
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_avatar_origin_message_never_schedules_a_reply(service_context):
|
async def test_avatar_origin_message_never_schedules_a_reply(service_context):
|
||||||
session_factory, service, boxim, clock = service_context
|
session_factory, service, boxim, clock = service_context
|
||||||
|
|||||||
@@ -0,0 +1,104 @@
|
|||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from fastapi.testclient import TestClient
|
||||||
|
|
||||||
|
from database import SessionLocal
|
||||||
|
from main import app, seed
|
||||||
|
from models import InvoiceApplication, PaymentRefund, TokenAccount, TokenPaymentOrder, TokenPlan, User
|
||||||
|
from services.token_billing import DEFAULT_TOKEN_GRANT
|
||||||
|
|
||||||
|
|
||||||
|
client = TestClient(app)
|
||||||
|
|
||||||
|
|
||||||
|
def _signature(token, timestamp, nonce):
|
||||||
|
return hashlib.sha1("".join(sorted([token, timestamp, nonce])).encode()).hexdigest()
|
||||||
|
|
||||||
|
|
||||||
|
def test_virtual_payment_callback_and_refund_are_idempotent(authorization_context):
|
||||||
|
seed()
|
||||||
|
context = authorization_context
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
user = db.query(User).filter(User.id == context["owner"].id).one()
|
||||||
|
user.wechat_mp_openid = "openid-flow"
|
||||||
|
user.wechat_mp_session_key = "session-flow"
|
||||||
|
plan = db.query(TokenPlan).filter(TokenPlan.id == "1").one()
|
||||||
|
plan.virtual_product_id = "points_plan_1"
|
||||||
|
db.commit()
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
env = {
|
||||||
|
"WECHAT_VIRTUAL_ENV": "sandbox",
|
||||||
|
"WECHAT_VIRTUAL_SANDBOX_APP_KEY": "sandbox-key",
|
||||||
|
"WECHAT_VIRTUAL_OFFER_ID": "offer-1",
|
||||||
|
"WECHAT_VIRTUAL_CALLBACK_TOKEN": "callback-token",
|
||||||
|
"AVATAR_FINANCE_ADMIN_SECRET": "finance-admin-secret-123",
|
||||||
|
}
|
||||||
|
with patch.dict(os.environ, env):
|
||||||
|
created = client.post(
|
||||||
|
"/api/token/charge",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
json={"planId": "1", "paymentMethod": "wechat", "payScene": "LITE"},
|
||||||
|
).json()["data"]
|
||||||
|
assert created["provider"] == "wechat_virtual"
|
||||||
|
params = json.loads(created["payMessage"])
|
||||||
|
assert params["mode"] == "short_series_goods"
|
||||||
|
assert "session-flow" not in created["payMessage"]
|
||||||
|
|
||||||
|
notify = {
|
||||||
|
"Event": "xpay_goods_deliver_notify",
|
||||||
|
"OutTradeNo": created["orderNo"],
|
||||||
|
"OpenId": "openid-flow",
|
||||||
|
"Env": 1,
|
||||||
|
"GoodsInfo": json.dumps({"ProductId": "points_plan_1", "ActualPrice": 1000}),
|
||||||
|
"WeChatPayInfo": json.dumps({"TransactionId": "wx-transaction-1"}),
|
||||||
|
}
|
||||||
|
query = {"timestamp": "100", "nonce": "nonce", "signature": _signature("callback-token", "100", "nonce")}
|
||||||
|
assert client.post("/api/token/payment/wechat/virtual/notify", params=query, json=notify).json()["ErrCode"] == 0
|
||||||
|
assert client.post("/api/token/payment/wechat/virtual/notify", params=query, json=notify).json()["ErrCode"] == 0
|
||||||
|
|
||||||
|
invoice = client.post(
|
||||||
|
f"/api/token/orders/{created['orderNo']}/invoice",
|
||||||
|
headers=context["owner_headers"],
|
||||||
|
json={"title": "测试用户", "invoiceType": "personal", "email": "test@example.com"},
|
||||||
|
).json()["data"]
|
||||||
|
assert invoice["status"] == "pending"
|
||||||
|
|
||||||
|
with patch("routers.tokens.request_wechat_virtual_refund", return_value={"errcode": 0}):
|
||||||
|
refund_response = client.post(
|
||||||
|
f"/api/token/admin/orders/{created['orderNo']}/refund",
|
||||||
|
headers={"X-Avatar-Finance-Key": "finance-admin-secret-123"},
|
||||||
|
json={"reason": "用户申请退款", "operator": "tester"},
|
||||||
|
)
|
||||||
|
assert refund_response.json()["data"]["status"] == "processing"
|
||||||
|
refund_no = refund_response.json()["data"]["refundNo"]
|
||||||
|
|
||||||
|
refund_notify = {
|
||||||
|
"Event": "xpay_refund_notify",
|
||||||
|
"MchOrderId": created["orderNo"],
|
||||||
|
"MchRefundId": refund_no,
|
||||||
|
"WxRefundId": "wx-refund-1",
|
||||||
|
"RefundFee": 1000,
|
||||||
|
"RetCode": 0,
|
||||||
|
}
|
||||||
|
assert client.post("/api/token/payment/wechat/virtual/notify", params=query, json=refund_notify).json()["ErrCode"] == 0
|
||||||
|
assert client.post("/api/token/payment/wechat/virtual/notify", params=query, json=refund_notify).json()["ErrCode"] == 0
|
||||||
|
assert client.post("/api/token/payment/wechat/virtual/notify", params=query, json=notify).json()["ErrCode"] == 0
|
||||||
|
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
order = db.query(TokenPaymentOrder).filter(TokenPaymentOrder.order_no == created["orderNo"]).one()
|
||||||
|
account = db.query(TokenAccount).filter(TokenAccount.user_id == context["owner"].id).one()
|
||||||
|
refund = db.query(PaymentRefund).filter(PaymentRefund.refund_no == refund_no).one()
|
||||||
|
invoice = db.query(InvoiceApplication).filter(InvoiceApplication.order_no == created["orderNo"]).one()
|
||||||
|
assert order.status == "refunded"
|
||||||
|
assert refund.status == "succeeded"
|
||||||
|
assert invoice.status == "cancelled"
|
||||||
|
assert account.balance == DEFAULT_TOKEN_GRANT
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
@@ -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
|
||||||
@@ -0,0 +1,72 @@
|
|||||||
|
import hashlib
|
||||||
|
import hmac
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import Mock, patch
|
||||||
|
|
||||||
|
from services import wechat_virtual_payment as virtual
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_payment_params_signs_the_exact_compact_payload():
|
||||||
|
order = SimpleNamespace(order_no="AV202609080001", plan_id="plan-1", price_cents=1000)
|
||||||
|
plan = SimpleNamespace(id="plan-1", virtual_product_id="points_plan_1")
|
||||||
|
env = {
|
||||||
|
"WECHAT_VIRTUAL_ENV": "sandbox",
|
||||||
|
"WECHAT_VIRTUAL_SANDBOX_APP_KEY": "sandbox-key",
|
||||||
|
"WECHAT_VIRTUAL_OFFER_ID": "offer-1",
|
||||||
|
}
|
||||||
|
with patch.dict(os.environ, env, clear=False):
|
||||||
|
result = virtual.build_payment_params(order=order, plan=plan, session_key="session-key")
|
||||||
|
|
||||||
|
sign_data = result["signData"]
|
||||||
|
assert sign_data == json.dumps({
|
||||||
|
"offerId": "offer-1",
|
||||||
|
"buyQuantity": 1,
|
||||||
|
"env": 1,
|
||||||
|
"currencyType": "CNY",
|
||||||
|
"productId": "points_plan_1",
|
||||||
|
"goodsPrice": 1000,
|
||||||
|
"outTradeNo": "AV202609080001",
|
||||||
|
"attach": '{"orderNo":"AV202609080001","planId":"plan-1"}',
|
||||||
|
}, ensure_ascii=False, separators=(",", ":"))
|
||||||
|
assert result["paySig"] == hmac.new(
|
||||||
|
b"sandbox-key", f"requestVirtualPayment&{sign_data}".encode(), hashlib.sha256
|
||||||
|
).hexdigest()
|
||||||
|
assert result["signature"] == hmac.new(
|
||||||
|
b"session-key", sign_data.encode(), hashlib.sha256
|
||||||
|
).hexdigest()
|
||||||
|
|
||||||
|
|
||||||
|
def test_callback_signature_and_xml_body_are_supported():
|
||||||
|
with patch.dict(os.environ, {"WECHAT_VIRTUAL_CALLBACK_TOKEN": "callback-token"}):
|
||||||
|
signature = hashlib.sha1("".join(sorted(["callback-token", "100", "nonce"])).encode()).hexdigest()
|
||||||
|
assert virtual.verify_callback_signature(signature, "100", "nonce")
|
||||||
|
payload = virtual.parse_callback_body(
|
||||||
|
b"<xml><Event>xpay_refund_notify</Event><GoodsInfo><ActualPrice>1000</ActualPrice></GoodsInfo></xml>"
|
||||||
|
)
|
||||||
|
assert virtual.callback_value(payload, "event") == "xpay_refund_notify"
|
||||||
|
assert virtual.callback_value(payload, "goodsinfo", "actualprice") == "1000"
|
||||||
|
|
||||||
|
|
||||||
|
def test_xpay_request_uses_server_access_token_and_pay_signature():
|
||||||
|
token_response = Mock(status_code=200)
|
||||||
|
token_response.json.return_value = {"access_token": "server-token", "expires_in": 7200}
|
||||||
|
pay_response = Mock(status_code=200)
|
||||||
|
pay_response.json.return_value = {"errcode": 0, "order": {"status": 2}}
|
||||||
|
virtual._access_token_cache = ("", 0)
|
||||||
|
env = {
|
||||||
|
"WECHAT_MP_APP_ID": "wx-app",
|
||||||
|
"WECHAT_MP_APP_SECRET": "wx-secret",
|
||||||
|
"WECHAT_VIRTUAL_SANDBOX_APP_KEY": "sandbox-key",
|
||||||
|
"WECHAT_VIRTUAL_ENV": "sandbox",
|
||||||
|
}
|
||||||
|
with patch.dict(os.environ, env), patch.object(virtual.httpx, "get", return_value=token_response), patch.object(
|
||||||
|
virtual.httpx, "post", return_value=pay_response
|
||||||
|
) as post:
|
||||||
|
result = virtual.query_order(openid="openid", order_no="AV1")
|
||||||
|
assert result["order"]["status"] == 2
|
||||||
|
body = '{"openid":"openid","order_id":"AV1","env":1}'
|
||||||
|
expected = hmac.new(b"sandbox-key", f"/xpay/query_order&{body}".encode(), hashlib.sha256).hexdigest()
|
||||||
|
assert post.call_args.args[0] == "https://api.weixin.qq.com/xpay/query_order"
|
||||||
|
assert post.call_args.kwargs["params"] == {"access_token": "server-token", "pay_sig": expected}
|
||||||
@@ -1,42 +1,56 @@
|
|||||||
# 会会数字分身 —— Docker 测试实例(独立端口,不干扰现有 :8088 huihui 部署)
|
# 会会数字分身 —— Docker 测试实例(独立端口,不干扰现有 :8088 huihui 部署)
|
||||||
services:
|
services:
|
||||||
avatar-backend:
|
avatar-backend:
|
||||||
build: ./backend
|
build:
|
||||||
image: avatar-test-backend:latest
|
context: ./backend
|
||||||
|
args:
|
||||||
|
APP_GIT_SHA: ${APP_GIT_SHA:?APP_GIT_SHA must be the full release commit}
|
||||||
|
APP_BUILD_TIME: ${APP_BUILD_TIME:?APP_BUILD_TIME must be set}
|
||||||
|
image: avatar-test-backend:${APP_GIT_SHA}
|
||||||
container_name: avatar-test-backend
|
container_name: avatar-test-backend
|
||||||
restart: unless-stopped
|
restart: unless-stopped
|
||||||
env_file:
|
env_file:
|
||||||
- .env
|
- .env
|
||||||
environment:
|
environment:
|
||||||
DATABASE_URL: sqlite:////data/avatar.db
|
DATABASE_URL: sqlite:////data/db/avatar.db
|
||||||
UPLOAD_DIR: /data/uploads
|
UPLOAD_DIR: /data/uploads
|
||||||
CHAT_MODEL_CONFIG_URL: http://host.docker.internal:8000/api/ai-models/runtime/digital-avatar
|
CHAT_MODEL_CONFIG_URL: http://host.docker.internal:8000/api/ai-models/runtime/digital-avatar
|
||||||
extra_hosts:
|
extra_hosts:
|
||||||
- "host.docker.internal:host-gateway"
|
- "host.docker.internal:host-gateway"
|
||||||
volumes:
|
volumes:
|
||||||
- avatar-data:/data
|
# Mount the directory, not only avatar.db: SQLite WAL/SHM files must survive recreation.
|
||||||
|
- ${AVATAR_DB_DIR:?AVATAR_DB_DIR must contain the persistent avatar.db}:/data/db
|
||||||
|
- ${AVATAR_UPLOAD_DIR:?AVATAR_UPLOAD_DIR must point to persistent uploads}:/data/uploads
|
||||||
expose:
|
expose:
|
||||||
- "8000"
|
- "8000"
|
||||||
ports:
|
ports:
|
||||||
- "8011:8000" # 仅用于直接调试 API;前端经内部网络访问,不走 host 端口
|
- "8011:8000" # 仅用于直接调试 API;前端经内部网络访问,不走 host 端口
|
||||||
|
healthcheck:
|
||||||
|
test: ["CMD", "python", "-c", "import json,urllib.request; d=json.load(urllib.request.urlopen('http://127.0.0.1:8000/api/health', timeout=5))['data']; assert d['status']=='ok' and all(d['checks'].values())"]
|
||||||
|
interval: 10s
|
||||||
|
timeout: 8s
|
||||||
|
retries: 12
|
||||||
|
start_period: 20s
|
||||||
networks:
|
networks:
|
||||||
- avatar-net
|
- avatar-net
|
||||||
|
|
||||||
avatar-frontend:
|
avatar-frontend:
|
||||||
build: .
|
build:
|
||||||
image: avatar-test-frontend:latest
|
context: .
|
||||||
|
args:
|
||||||
|
APP_GIT_SHA: ${APP_GIT_SHA:?APP_GIT_SHA must be the full release commit}
|
||||||
|
APP_BUILD_TIME: ${APP_BUILD_TIME:?APP_BUILD_TIME must be set}
|
||||||
|
image: avatar-test-frontend:${APP_GIT_SHA}
|
||||||
container_name: avatar-test-frontend
|
container_name: avatar-test-frontend
|
||||||
restart: unless-stopped
|
restart: unless-stopped
|
||||||
ports:
|
ports:
|
||||||
- "8099:80" # 浏览器访问 http://<host>:8099
|
- "8099:80" # 浏览器访问 http://<host>:8099
|
||||||
depends_on:
|
depends_on:
|
||||||
- avatar-backend
|
avatar-backend:
|
||||||
|
condition: service_healthy
|
||||||
networks:
|
networks:
|
||||||
- avatar-net
|
- avatar-net
|
||||||
|
|
||||||
networks:
|
networks:
|
||||||
avatar-net:
|
avatar-net:
|
||||||
driver: bridge
|
driver: bridge
|
||||||
|
|
||||||
volumes:
|
|
||||||
avatar-data:
|
|
||||||
|
|||||||
@@ -39,19 +39,62 @@ 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位随机密钥>
|
||||||
HUIHUI_PAYMENT_TIMEOUT_SECONDS=30
|
HUIHUI_PAYMENT_TIMEOUT_SECONDS=30
|
||||||
|
HUIHUI_PAYMENT_REFUND_PATH=/payment/refund
|
||||||
|
AVATAR_FINANCE_ADMIN_SECRET=<至少32位随机密钥,与管理后台一致>
|
||||||
|
|
||||||
DATABASE_URL=sqlite:////data/avatar.db
|
# 微信小程序虚拟支付;联调先使用 sandbox
|
||||||
|
WECHAT_MP_APP_ID=<小程序AppID>
|
||||||
|
WECHAT_MP_APP_SECRET=<小程序AppSecret>
|
||||||
|
WECHAT_VIRTUAL_ENV=sandbox
|
||||||
|
WECHAT_VIRTUAL_SANDBOX_APP_KEY=<沙箱AppKey>
|
||||||
|
WECHAT_VIRTUAL_APP_KEY=<正式AppKey>
|
||||||
|
WECHAT_VIRTUAL_OFFER_ID=<offer-id>
|
||||||
|
WECHAT_VIRTUAL_CALLBACK_TOKEN=<回调校验Token>
|
||||||
|
WECHAT_VIRTUAL_PRODUCT_1=<10元套餐商品ID>
|
||||||
|
WECHAT_VIRTUAL_PRODUCT_2=<100元套餐商品ID>
|
||||||
|
WECHAT_VIRTUAL_PRODUCT_3=<1000元套餐商品ID>
|
||||||
|
WECHAT_VIRTUAL_PRODUCT_4=<10000元套餐商品ID>
|
||||||
|
|
||||||
|
AVATAR_DB_DIR=/srv/digital-avatar/data/db
|
||||||
|
AVATAR_UPLOAD_DIR=/srv/digital-avatar/data/uploads
|
||||||
|
DATABASE_URL=sqlite:////data/db/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
|
||||||
|
APP_GIT_SHA=<本次发布的完整提交SHA>
|
||||||
|
APP_BUILD_TIME=<UTC ISO-8601构建时间>
|
||||||
|
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` 等兜底配置。数据库文件与上传目录必须从宿主机显式挂载,不能存放在容器临时层。
|
||||||
|
|
||||||
积分充值使用会会支付体系的 `payment-v3/payment/pay`,渠道值为 `WECHAT` / `ALIPAY`,端内支付场景为 `APP`,微信内 H5 使用 `JSAPI`。`HUIHUI_PAYMENT_CALLBACK_SECRET` 只用于为每笔订单生成 HMAC 回调签名,不会发送到前端或直接出现在回调地址中。支付回调确认状态成功且金额与套餐价格完全一致后才增加积分,重复回调不会重复到账。
|
`AVATAR_DB_DIR` 和 `AVATAR_UPLOAD_DIR` 必须是已备份的宿主机绝对路径,编排缺少任一变量都会直接拒绝构建或启动,防止误挂空卷造成用户、分身或知识库“丢失”的假象。SQLite 必须挂载整个数据库目录,不能只挂载 `avatar.db` 单文件,否则 `avatar.db-wal` 和 `avatar.db-shm` 会留在容器临时层,换容器后可能出现数据状态回退。
|
||||||
|
|
||||||
|
`EMBEDDING_API_URL` 同时支持 OpenAI 兼容基础地址(如上面的 `/v1`)和完整的 `/v1/embeddings` 地址,后端会统一请求 `/embeddings`。发布后必须在后端容器内执行一次最小向量探针,确认返回向量数量和维度,而不能只检查 `/api/health`。
|
||||||
|
|
||||||
|
App 与 H5 积分充值使用会会支付体系的 `payment-v3/payment/pay`,渠道值为 `WECHAT` / `ALIPAY`;App 场景为 `APP`,普通浏览器为 `H5`,微信内 H5 为 `JSAPI`。`HUIHUI_PAYMENT_CALLBACK_SECRET` 只用于为每笔订单生成 HMAC 回调签名,不会发送到前端或直接出现在回调地址中。支付回调确认状态成功且金额与套餐价格完全一致后才增加积分,重复回调不会重复到账。
|
||||||
|
|
||||||
|
微信小程序使用微信虚拟支付:小程序先通过 `POST /api/token/wechat/session` 交换临时登录码,再由 `POST /api/token/charge`(`payScene=LITE`)返回已签名的 `requestVirtualPayment` 参数。微信回调地址配置为 `https://digital.99hui.com/api/token/payment/wechat/virtual/notify`。回调会复核签名、OpenID、环境、商品 ID 与实付金额,退款回调确认后才扣回积分。AppKey、AppSecret、session_key 均不得下发前端或写日志。
|
||||||
|
|
||||||
|
管理后台需要配置相同的 `AVATAR_FINANCE_ADMIN_SECRET` 和 `AVATAR_BACKEND_URL=https://digital.99hui.com`。退款只支持整单原路退款;供应商受理后显示“处理中”,收到渠道成功回调(或经渠道后台核对后人工确认)才将订单置为已退款。已消费掉本订单积分时,后台会拒绝主动退款;若渠道外部退款先发生,积分账户允许形成负数以记录欠额并阻止继续消费。
|
||||||
|
|
||||||
## 3. 构建与发布
|
## 3. 构建与发布
|
||||||
|
|
||||||
@@ -60,7 +103,7 @@ CHAT_MODEL_CONFIG_URL=http://<huihuisquare-api>/api/ai-models/runtime/digital-av
|
|||||||
```bash
|
```bash
|
||||||
BACKUP_DIR="backups/$(date +%Y%m%d-%H%M%S)"
|
BACKUP_DIR="backups/$(date +%Y%m%d-%H%M%S)"
|
||||||
mkdir -p "$BACKUP_DIR"
|
mkdir -p "$BACKUP_DIR"
|
||||||
cp /srv/digital-avatar/data/avatar.db "$BACKUP_DIR/"
|
cp /srv/digital-avatar/data/db/avatar.db "$BACKUP_DIR/"
|
||||||
tar -C /srv/digital-avatar/data -czf "$BACKUP_DIR/uploads.tgz" uploads
|
tar -C /srv/digital-avatar/data -czf "$BACKUP_DIR/uploads.tgz" uploads
|
||||||
```
|
```
|
||||||
|
|
||||||
@@ -70,12 +113,22 @@ tar -C /srv/digital-avatar/data -czf "$BACKUP_DIR/uploads.tgz" uploads
|
|||||||
git fetch origin
|
git fetch origin
|
||||||
git checkout <已验收的提交SHA>
|
git checkout <已验收的提交SHA>
|
||||||
cd digital-avatar-app
|
cd digital-avatar-app
|
||||||
docker compose build --pull avatar-backend avatar-frontend
|
export APP_GIT_SHA="$(git rev-parse HEAD)"
|
||||||
docker compose up -d avatar-backend avatar-frontend
|
export APP_BUILD_TIME="$(date -u +%Y-%m-%dT%H:%M:%SZ)"
|
||||||
|
docker compose build --pull --no-cache avatar-backend avatar-frontend
|
||||||
|
docker compose up -d --force-recreate --wait avatar-backend avatar-frontend
|
||||||
docker compose ps
|
docker compose ps
|
||||||
curl -fsS http://127.0.0.1:8099/api/health
|
python3 scripts/verify-deployment.py \
|
||||||
|
https://digital.99hui.com "$APP_GIT_SHA" \
|
||||||
|
--backend-container avatar-backend \
|
||||||
|
--frontend-container avatar-frontend \
|
||||||
|
--expected-db-source /srv/digital-avatar/data/db \
|
||||||
|
--expected-upload-source /srv/digital-avatar/data/uploads
|
||||||
|
docker compose exec avatar-backend python -c 'import embeddings; v=embeddings.embed(["部署向量探针"]); print(len(v), len(v[0]))'
|
||||||
```
|
```
|
||||||
|
|
||||||
|
Jenkins 必须以 `verify-deployment.py` 返回成功作为发布成功条件,不能只以镜像构建或容器启动成功作为条件。脚本会同时核对公网前后端 Git SHA、数据库可读、上传目录可写、PDF OCR 依赖和宿主机数据挂载;任意一项不一致都会返回非零状态并阻止发布标绿。镜像使用 Git SHA 标签,不再依赖可被旧缓存覆盖的 `latest`。
|
||||||
|
|
||||||
生产编排应把示例中的测试端口改为内网暴露,由统一 HTTPS 网关接入。后端暂时使用 SQLite,必须保持单实例写入;若扩展为多后端实例,应先迁移到 PostgreSQL,并把延迟接管任务改为共享队列。
|
生产编排应把示例中的测试端口改为内网暴露,由统一 HTTPS 网关接入。后端暂时使用 SQLite,必须保持单实例写入;若扩展为多后端实例,应先迁移到 PostgreSQL,并把延迟接管任务改为共享队列。
|
||||||
|
|
||||||
## 4. 网关要求
|
## 4. 网关要求
|
||||||
@@ -98,11 +151,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. 发布验收
|
||||||
|
|
||||||
@@ -113,9 +168,14 @@ location /api/ {
|
|||||||
5. 使用过期或伪造 token 时进入登录页并显示凭证失效,不得继续访问旧用户数据。
|
5. 使用过期或伪造 token 时进入登录页并显示凭证失效,不得继续访问旧用户数据。
|
||||||
6. 分身聊天 SSE 逐段输出正常,Markdown 正常渲染,知识库优先级和积分扣费正常。
|
6. 分身聊天 SSE 逐段输出正常,Markdown 正常渲染,知识库优先级和积分扣费正常。
|
||||||
7. 开启 BOXIM 主动接管后保持在线,默认三分钟回复、自定义等待时间、已读回执、分身防回环和主人发言暂停均正常。
|
7. 开启 BOXIM 主动接管后保持在线,默认三分钟回复、自定义等待时间、已读回执、分身防回环和主人发言暂停均正常。
|
||||||
8. 重建容器后数据库、头像、知识库文档仍存在,`/api/health` 返回成功。
|
8. 重建容器后数据库、头像、知识库文档仍存在,`/api/health` 的 `gitSha` 与发布 SHA 一致,`database`、`uploads`、`pdfOcr` 三项检查均为 `true`。
|
||||||
9. `https://digital.99hui.com/api/health` 可访问,证书域名和有效期正确,HTTP 自动跳转 HTTPS。
|
9. `https://digital.99hui.com/api/health` 可访问,证书域名和有效期正确,HTTP 自动跳转 HTTPS。
|
||||||
10. 微信和支付宝各创建一笔最小套餐订单,未付款时积分不变;支付成功后回调到账一次,重复回调积分不重复增加。
|
10. 微信和支付宝各创建一笔最小套餐订单,未付款时积分不变;支付成功后回调到账一次,重复回调积分不重复增加。
|
||||||
|
11. 微信虚拟支付在沙箱环境完成下单、支付回调、查单兜底和退款回调;错误 OpenID、商品、环境或金额均被拒绝。
|
||||||
|
12. 财务后台能筛选订单、关闭待支付订单、发起整单退款、登记退款对账结果,并处理个人/企业电子发票申请。
|
||||||
|
13. 私聊和公开分享各上传 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 光回答不作确定诊断,并显示人工复核提示。
|
||||||
@@ -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 / {
|
||||||
|
|||||||
@@ -0,0 +1,107 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""Fail a deployment unless frontend and backend run the expected release."""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import subprocess
|
||||||
|
import sys
|
||||||
|
import urllib.request
|
||||||
|
|
||||||
|
|
||||||
|
def fetch_json(url):
|
||||||
|
with urllib.request.urlopen(url, timeout=20) as response:
|
||||||
|
if response.status != 200:
|
||||||
|
raise RuntimeError(f"{url} returned HTTP {response.status}")
|
||||||
|
return json.load(response)
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("base_url", help="Public site URL, for example https://digital.99hui.com")
|
||||||
|
parser.add_argument("expected_sha", help="Full Git commit SHA being deployed")
|
||||||
|
parser.add_argument("--backend-container", help="Backend container name for image and mount checks")
|
||||||
|
parser.add_argument("--frontend-container", help="Frontend container name for image checks")
|
||||||
|
parser.add_argument("--expected-db-source", help="Required host directory mounted for SQLite and its WAL files")
|
||||||
|
parser.add_argument("--expected-upload-source", help="Required host source mounted as the upload directory")
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
base_url = args.base_url.rstrip("/")
|
||||||
|
errors = []
|
||||||
|
try:
|
||||||
|
health = fetch_json(f"{base_url}/api/health").get("data") or {}
|
||||||
|
except Exception as exc:
|
||||||
|
errors.append(f"cannot read backend release metadata: {exc}")
|
||||||
|
health = {}
|
||||||
|
try:
|
||||||
|
frontend = fetch_json(f"{base_url}/version.json")
|
||||||
|
except Exception as exc:
|
||||||
|
errors.append(f"cannot read frontend release metadata: {exc}")
|
||||||
|
frontend = {}
|
||||||
|
|
||||||
|
if health.get("status") != "ok":
|
||||||
|
errors.append(f"backend status is {health.get('status')!r}")
|
||||||
|
failed_checks = [name for name, passed in (health.get("checks") or {}).items() if not passed]
|
||||||
|
if failed_checks:
|
||||||
|
errors.append("backend checks failed: " + ", ".join(failed_checks))
|
||||||
|
if health.get("gitSha") != args.expected_sha:
|
||||||
|
errors.append(f"backend SHA is {health.get('gitSha')!r}")
|
||||||
|
if frontend.get("gitSha") != args.expected_sha:
|
||||||
|
errors.append(f"frontend SHA is {frontend.get('gitSha')!r}")
|
||||||
|
|
||||||
|
if args.backend_container:
|
||||||
|
backend = inspect_container(args.backend_container, errors)
|
||||||
|
check_container_revision(backend, args.expected_sha, "backend", errors)
|
||||||
|
check_mount(backend, args.expected_db_source, "database", errors)
|
||||||
|
check_mount(backend, args.expected_upload_source, "uploads", errors)
|
||||||
|
elif args.expected_db_source or args.expected_upload_source:
|
||||||
|
errors.append("--backend-container is required when checking data mounts")
|
||||||
|
|
||||||
|
if args.frontend_container:
|
||||||
|
frontend_container = inspect_container(args.frontend_container, errors)
|
||||||
|
check_container_revision(frontend_container, args.expected_sha, "frontend", errors)
|
||||||
|
|
||||||
|
if errors:
|
||||||
|
print("Deployment verification failed:", file=sys.stderr)
|
||||||
|
for error in errors:
|
||||||
|
print(f"- {error}", file=sys.stderr)
|
||||||
|
return 1
|
||||||
|
|
||||||
|
print(f"Deployment verified: {args.expected_sha}")
|
||||||
|
print("Backend checks: database, uploads, pdfOcr")
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
def inspect_container(name, errors):
|
||||||
|
try:
|
||||||
|
output = subprocess.check_output(
|
||||||
|
["docker", "inspect", name], universal_newlines=True
|
||||||
|
)
|
||||||
|
return json.loads(output)[0]
|
||||||
|
except Exception as exc:
|
||||||
|
errors.append(f"cannot inspect container {name!r}: {exc}")
|
||||||
|
return {}
|
||||||
|
|
||||||
|
|
||||||
|
def check_container_revision(container, expected_sha, label, errors):
|
||||||
|
actual = ((container.get("Config") or {}).get("Labels") or {}).get(
|
||||||
|
"org.opencontainers.image.revision"
|
||||||
|
)
|
||||||
|
if actual != expected_sha:
|
||||||
|
errors.append(f"{label} container image SHA is {actual!r}")
|
||||||
|
|
||||||
|
|
||||||
|
def check_mount(container, expected_source, label, errors):
|
||||||
|
if not expected_source:
|
||||||
|
return
|
||||||
|
expected = os.path.realpath(expected_source)
|
||||||
|
sources = {
|
||||||
|
os.path.realpath(mount.get("Source", ""))
|
||||||
|
for mount in container.get("Mounts") or []
|
||||||
|
}
|
||||||
|
if expected not in sources:
|
||||||
|
errors.append(f"{label} mount source {expected!r} is not attached")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
raise SystemExit(main())
|
||||||
@@ -152,16 +152,35 @@ export interface TokenPaymentOrder {
|
|||||||
planId: string
|
planId: string
|
||||||
paymentMethod: 'wechat' | 'alipay'
|
paymentMethod: 'wechat' | 'alipay'
|
||||||
payType: 'WECHAT' | 'ALIPAY'
|
payType: 'WECHAT' | 'ALIPAY'
|
||||||
payWay: 'APP' | 'LITE' | 'JSAPI'
|
payWay: 'APP' | 'H5' | 'LITE' | 'JSAPI'
|
||||||
pointsAmount: number
|
pointsAmount: number
|
||||||
price: number
|
price: number
|
||||||
status: 'pending' | 'paid' | 'failed'
|
status: 'pending' | 'paid' | 'failed' | 'closed' | 'refunded'
|
||||||
|
provider: 'huihui' | 'wechat_virtual'
|
||||||
providerStatus: string
|
providerStatus: string
|
||||||
payMessage: string
|
payMessage: string
|
||||||
failureReason: string
|
failureReason: string
|
||||||
|
refundStatus: 'none' | 'pending' | 'processing' | 'succeeded' | 'failed'
|
||||||
|
createdAt: string | null
|
||||||
|
paidAt: string | null
|
||||||
|
refundedAt: string | null
|
||||||
balance: number
|
balance: number
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export interface TokenInvoice {
|
||||||
|
id: string
|
||||||
|
orderNo: string
|
||||||
|
title: string
|
||||||
|
invoiceType: 'personal' | 'company'
|
||||||
|
taxNumber: string
|
||||||
|
email: string
|
||||||
|
amount: number
|
||||||
|
status: 'pending' | 'issued' | 'rejected' | 'cancelled'
|
||||||
|
invoiceNo: string
|
||||||
|
invoiceUrl: string
|
||||||
|
remark: string
|
||||||
|
}
|
||||||
|
|
||||||
// 获取 Token 余额
|
// 获取 Token 余额
|
||||||
export const getTokenBalance = () =>
|
export const getTokenBalance = () =>
|
||||||
request.get<TokenBalance>('/token/balance')
|
request.get<TokenBalance>('/token/balance')
|
||||||
@@ -174,12 +193,25 @@ export const getRechargePlans = () =>
|
|||||||
export const chargeToken = (
|
export const chargeToken = (
|
||||||
planId: string,
|
planId: string,
|
||||||
paymentMethod: 'wechat' | 'alipay',
|
paymentMethod: 'wechat' | 'alipay',
|
||||||
payScene: 'APP' | 'LITE' | 'JSAPI'
|
payScene: 'APP' | 'H5' | 'LITE' | 'JSAPI'
|
||||||
) => request.post<TokenPaymentOrder>('/token/charge', { planId, paymentMethod, payScene })
|
) => request.post<TokenPaymentOrder>('/token/charge', { planId, paymentMethod, payScene })
|
||||||
|
|
||||||
export const getTokenPaymentStatus = (orderId: string) =>
|
export const getTokenPaymentStatus = (orderId: string) =>
|
||||||
request.get<TokenPaymentOrder>(`/token/payment/${orderId}`)
|
request.get<TokenPaymentOrder>(`/token/payment/${orderId}`)
|
||||||
|
|
||||||
|
export const bindWechatVirtualSession = (code: string) =>
|
||||||
|
request.post<{ ready: boolean }>('/token/wechat/session', { code })
|
||||||
|
|
||||||
|
export const getTokenOrders = (page = 1, pageSize = 20) =>
|
||||||
|
request.get<{ total: number; page: number; pageSize: number; items: Array<TokenPaymentOrder & { invoice?: TokenInvoice }> }>(
|
||||||
|
'/token/orders', { params: { page, page_size: pageSize } }
|
||||||
|
)
|
||||||
|
|
||||||
|
export const applyTokenInvoice = (
|
||||||
|
orderNo: string,
|
||||||
|
payload: { title: string; invoiceType: 'personal' | 'company'; taxNumber?: string; email?: string }
|
||||||
|
) => request.post<TokenInvoice>(`/token/orders/${orderNo}/invoice`, payload)
|
||||||
|
|
||||||
// 按分身和使用场景汇总 Token 消耗
|
// 按分身和使用场景汇总 Token 消耗
|
||||||
export const getTokenUsage = () =>
|
export const getTokenUsage = () =>
|
||||||
request.get<TokenUsageSummary[]>('/token/usage')
|
request.get<TokenUsageSummary[]>('/token/usage')
|
||||||
@@ -305,6 +337,9 @@ export interface KnowledgeDoc {
|
|||||||
vectorized?: boolean
|
vectorized?: boolean
|
||||||
embeddingModel?: string
|
embeddingModel?: string
|
||||||
chunkCount?: number
|
chunkCount?: number
|
||||||
|
errorMessage?: string
|
||||||
|
indexStage?: string
|
||||||
|
indexProgress?: number
|
||||||
createdAt: string
|
createdAt: string
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -330,12 +365,83 @@ export interface SearchResult {
|
|||||||
export const getKnowledgeDocs = (avatarId: string) =>
|
export const getKnowledgeDocs = (avatarId: string) =>
|
||||||
request.get<KnowledgeDoc[]>(`/avatar/${avatarId}/knowledge/docs`)
|
request.get<KnowledgeDoc[]>(`/avatar/${avatarId}/knowledge/docs`)
|
||||||
|
|
||||||
// 上传文档(支持 md/txt/pdf/doc/docx/xlsx)
|
const KNOWLEDGE_UPLOAD_CHUNK_SIZE = 5 * 1024 * 1024
|
||||||
export const uploadKnowledgeDoc = (avatarId: string, file: File) => {
|
|
||||||
|
const uploadKnowledgeChunk = async (
|
||||||
|
avatarId: string,
|
||||||
|
uploadId: string,
|
||||||
|
chunkIndex: number,
|
||||||
|
chunk: Blob,
|
||||||
|
onProgress?: (loaded: number) => void
|
||||||
|
) => {
|
||||||
|
const form = new FormData()
|
||||||
|
form.append('file', chunk, `chunk-${chunkIndex}`)
|
||||||
|
let reportedLoaded = 0
|
||||||
|
for (let attempt = 1; attempt <= 3; attempt += 1) {
|
||||||
|
try {
|
||||||
|
await request.post(
|
||||||
|
`/avatar/${avatarId}/knowledge/uploads/${uploadId}/chunks/${chunkIndex}`,
|
||||||
|
form,
|
||||||
|
{
|
||||||
|
headers: { 'Content-Type': 'multipart/form-data' },
|
||||||
|
timeout: 2 * 60 * 1000,
|
||||||
|
onUploadProgress: (event) => {
|
||||||
|
reportedLoaded = Math.max(reportedLoaded, Math.min(event.loaded, chunk.size))
|
||||||
|
onProgress?.(reportedLoaded)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return
|
||||||
|
} catch (error: any) {
|
||||||
|
const status = Number(error?.response?.status || 0)
|
||||||
|
const retryable = !status || status === 408 || status === 429 || status >= 500
|
||||||
|
if (!retryable || attempt === 3) throw error
|
||||||
|
await new Promise((resolve) => window.setTimeout(resolve, attempt * 800))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 大文件拆成 5MB 分片,避免生产代理的请求体限制拦截整个文件。
|
||||||
|
export const uploadKnowledgeDoc = async (
|
||||||
|
avatarId: string,
|
||||||
|
file: File,
|
||||||
|
onUploadProgress?: (loaded: number, total: number) => void
|
||||||
|
) => {
|
||||||
|
if (file.size > KNOWLEDGE_UPLOAD_CHUNK_SIZE) {
|
||||||
|
const totalChunks = Math.ceil(file.size / KNOWLEDGE_UPLOAD_CHUNK_SIZE)
|
||||||
|
const upload: any = await request.post(`/avatar/${avatarId}/knowledge/uploads`, {
|
||||||
|
filename: file.name,
|
||||||
|
fileSize: file.size,
|
||||||
|
totalChunks
|
||||||
|
})
|
||||||
|
let uploadedBytes = 0
|
||||||
|
for (let index = 0; index < totalChunks; index += 1) {
|
||||||
|
const start = index * KNOWLEDGE_UPLOAD_CHUNK_SIZE
|
||||||
|
const chunk = file.slice(start, Math.min(start + KNOWLEDGE_UPLOAD_CHUNK_SIZE, file.size))
|
||||||
|
await uploadKnowledgeChunk(
|
||||||
|
avatarId,
|
||||||
|
upload.uploadId,
|
||||||
|
index,
|
||||||
|
chunk,
|
||||||
|
(chunkLoaded) => onUploadProgress?.(uploadedBytes + chunkLoaded, file.size)
|
||||||
|
)
|
||||||
|
uploadedBytes += chunk.size
|
||||||
|
onUploadProgress?.(uploadedBytes, file.size)
|
||||||
|
}
|
||||||
|
return request.post<KnowledgeDoc>(
|
||||||
|
`/avatar/${avatarId}/knowledge/uploads/${upload.uploadId}/complete`,
|
||||||
|
undefined,
|
||||||
|
{ timeout: 2 * 60 * 1000 }
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
const form = new FormData()
|
const form = new FormData()
|
||||||
form.append('file', file)
|
form.append('file', file)
|
||||||
return request.post<KnowledgeDoc>(`/avatar/${avatarId}/knowledge/docs`, form, {
|
return request.post<KnowledgeDoc>(`/avatar/${avatarId}/knowledge/docs`, form, {
|
||||||
headers: { 'Content-Type': 'multipart/form-data' }
|
headers: { 'Content-Type': 'multipart/form-data' },
|
||||||
|
// A slow mobile uplink must not be mistaken for a failed upload.
|
||||||
|
timeout: 10 * 60 * 1000,
|
||||||
|
onUploadProgress: (event) => onUploadProgress?.(event.loaded, event.total || file.size)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -343,6 +449,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`)
|
||||||
@@ -374,15 +483,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 {
|
||||||
@@ -401,19 +530,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()
|
||||||
@@ -436,10 +587,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 ====================
|
||||||
|
|||||||
@@ -179,7 +179,7 @@ const permissionItems: Array<{
|
|||||||
},
|
},
|
||||||
{
|
{
|
||||||
key: 'publish',
|
key: 'publish',
|
||||||
title: '发布微博内容',
|
title: '发布微播内容',
|
||||||
description: '允许分身自动发布动态内容',
|
description: '允许分身自动发布动态内容',
|
||||||
tone: 'green',
|
tone: 'green',
|
||||||
},
|
},
|
||||||
@@ -192,7 +192,7 @@ const permissionItems: Array<{
|
|||||||
{
|
{
|
||||||
key: 'interact',
|
key: 'interact',
|
||||||
title: '广场互动操作',
|
title: '广场互动操作',
|
||||||
description: '点赞、收藏、评论、回复等操作',
|
description: '允许分身代表你点赞、收藏、评论和回复;时段、间隔与触发概率由广场调度设置统一控制',
|
||||||
tone: 'pink',
|
tone: 'pink',
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -691,7 +691,9 @@ svg {
|
|||||||
font-weight: 400;
|
font-weight: 400;
|
||||||
line-height: 1.45;
|
line-height: 1.45;
|
||||||
text-overflow: ellipsis;
|
text-overflow: ellipsis;
|
||||||
white-space: nowrap;
|
white-space: normal;
|
||||||
|
overflow-wrap: anywhere;
|
||||||
|
word-break: break-word;
|
||||||
}
|
}
|
||||||
|
|
||||||
.permission-row.takeover .permission-copy small {
|
.permission-row.takeover .permission-copy small {
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -48,7 +54,7 @@
|
|||||||
</template>
|
</template>
|
||||||
<template v-else>{{ message.content }}</template>
|
<template v-else>{{ message.content }}</template>
|
||||||
</div>
|
</div>
|
||||||
<div v-if="message.source || message.references?.length" class="message-source">
|
<div v-if="sourceLabel(message.source) || message.references?.length" class="message-source">
|
||||||
{{ sourceLabel(message.source) }}
|
{{ sourceLabel(message.source) }}
|
||||||
<span v-if="message.references?.length"> · {{ message.references.map((item) => item.filename).filter(Boolean).join('、') }}</span>
|
<span v-if="message.references?.length"> · {{ message.references.map((item) => item.filename).filter(Boolean).join('、') }}</span>
|
||||||
</div>
|
</div>
|
||||||
@@ -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,26 @@ 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: '参考文件知识库',
|
||||||
qwen: '智能回答',
|
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 +277,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 +389,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 +412,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 +420,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 +440,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>
|
||||||
|
|||||||
@@ -15,7 +15,7 @@
|
|||||||
|
|
||||||
<template v-else>
|
<template v-else>
|
||||||
<div class="tab-switcher" role="tablist" aria-label="知识库类型">
|
<div class="tab-switcher" role="tablist" aria-label="知识库类型">
|
||||||
<button class="tab-btn" :class="{ active: activeTab === 'docs' }" role="tab" :aria-selected="activeTab === 'docs'" @click="activeTab = 'docs'">文档知识库 <b>{{ docs.length }}</b></button>
|
<button class="tab-btn" :class="{ active: activeTab === 'docs' }" role="tab" :aria-selected="activeTab === 'docs'" @click="activeTab = 'docs'">文档知识库 <b>{{ displayDocs.length }}</b></button>
|
||||||
<button class="tab-btn" :class="{ active: activeTab === 'qa' }" role="tab" :aria-selected="activeTab === 'qa'" @click="activeTab = 'qa'">标准问答对 <b>{{ qaPairs.length }}</b></button>
|
<button class="tab-btn" :class="{ active: activeTab === 'qa' }" role="tab" :aria-selected="activeTab === 'qa'" @click="activeTab = 'qa'">标准问答对 <b>{{ qaPairs.length }}</b></button>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
@@ -25,14 +25,14 @@
|
|||||||
<div class="upload-icon">📥</div>
|
<div class="upload-icon">📥</div>
|
||||||
<p class="upload-title"><span class="upload-link">点击上传</span></p>
|
<p class="upload-title"><span class="upload-link">点击上传</span></p>
|
||||||
<p class="upload-hint">支持 MD / TXT / PDF / DOC / DOCX / XLSX,上传后自动向量化</p>
|
<p class="upload-hint">支持 MD / TXT / PDF / DOC / DOCX / XLSX,上传后自动向量化</p>
|
||||||
<input ref="fileInput" type="file" accept=".md,.txt,.pdf,.doc,.docx,.xlsx" class="hidden-input" @change="onFileChange" />
|
<input ref="fileInput" type="file" multiple accept=".md,.txt,.pdf,.doc,.docx,.xlsx" class="hidden-input" @change="onFileChange" />
|
||||||
</div>
|
</div>
|
||||||
<p v-if="uploading" class="uploading-text">上传并向量化中…</p>
|
<p v-if="uploading" class="uploading-text">{{ pendingUploads.length }} 个文件正在上传</p>
|
||||||
<p v-if="uploadError" class="error-text">{{ uploadError }}</p>
|
<p v-if="uploadError" class="error-text">{{ uploadError }}</p>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<div v-if="docs.length" class="mobile-card-list">
|
<div v-if="displayDocs.length" class="mobile-card-list">
|
||||||
<article v-for="doc in docs" :key="doc.id" class="knowledge-card">
|
<article v-for="doc in displayDocs" :key="doc.id" class="knowledge-card document-card">
|
||||||
<div class="card-icon">{{ fileEmoji(doc.fileType) }}</div>
|
<div class="card-icon">{{ fileEmoji(doc.fileType) }}</div>
|
||||||
<div class="card-content">
|
<div class="card-content">
|
||||||
<div class="card-title-row">
|
<div class="card-title-row">
|
||||||
@@ -41,8 +41,19 @@
|
|||||||
</div>
|
</div>
|
||||||
<p class="card-meta">{{ doc.fileType.toUpperCase() }} · {{ formatSize(doc.fileSize) }} · {{ formatDate(doc.createdAt) }}</p>
|
<p class="card-meta">{{ doc.fileType.toUpperCase() }} · {{ formatSize(doc.fileSize) }} · {{ formatDate(doc.createdAt) }}</p>
|
||||||
<p class="card-detail">{{ documentState(doc).detail }}</p>
|
<p class="card-detail">{{ documentState(doc).detail }}</p>
|
||||||
|
<div v-if="documentState(doc).progress !== undefined" class="progress-track" :aria-label="`${documentState(doc).label} ${documentState(doc).progress}%`">
|
||||||
|
<span class="progress-fill" :style="{ width: `${documentState(doc).progress}%` }"></span>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
<div class="card-actions">
|
||||||
|
<button v-if="!doc.localUploading" class="card-delete" @click="removeDoc(doc.id)">{{ doc.localOnly ? '移除' : '删除' }}</button>
|
||||||
|
</div>
|
||||||
|
<div v-if="canRetryDoc(doc)" class="card-retry-area">
|
||||||
|
<span v-if="retryErrors[doc.id]" class="card-retry-error">{{ retryErrors[doc.id] }}</span>
|
||||||
|
<button class="card-retry" :disabled="retryingDocs[doc.id]" @click="retryDoc(doc)">
|
||||||
|
{{ retryingDocs[doc.id] ? '重新索引中…' : '重新索引' }}
|
||||||
|
</button>
|
||||||
</div>
|
</div>
|
||||||
<button class="card-delete" @click="removeDoc(doc.id)">删除</button>
|
|
||||||
</article>
|
</article>
|
||||||
</div>
|
</div>
|
||||||
<div v-else class="card-empty">📂 暂无文档,先上传一个知识文件</div>
|
<div v-else class="card-empty">📂 暂无文档,先上传一个知识文件</div>
|
||||||
@@ -78,7 +89,7 @@
|
|||||||
</template>
|
</template>
|
||||||
|
|
||||||
<script setup lang="ts">
|
<script setup lang="ts">
|
||||||
import { ref, onMounted, computed } from 'vue'
|
import { ref, onMounted, onUnmounted, computed } from 'vue'
|
||||||
import { useRoute, useRouter } from 'vue-router'
|
import { useRoute, useRouter } from 'vue-router'
|
||||||
import { useAvatarStore } from '@/store/avatar'
|
import { useAvatarStore } from '@/store/avatar'
|
||||||
import { pickScopedAvatarId, unwrapListData } from '@/utils/avatar-page-data.js'
|
import { pickScopedAvatarId, unwrapListData } from '@/utils/avatar-page-data.js'
|
||||||
@@ -87,6 +98,7 @@ import {
|
|||||||
getKnowledgeDocs,
|
getKnowledgeDocs,
|
||||||
uploadKnowledgeDoc,
|
uploadKnowledgeDoc,
|
||||||
deleteKnowledgeDoc,
|
deleteKnowledgeDoc,
|
||||||
|
retryKnowledgeDoc,
|
||||||
getQAPairs,
|
getQAPairs,
|
||||||
deleteQAPair,
|
deleteQAPair,
|
||||||
searchKnowledge,
|
searchKnowledge,
|
||||||
@@ -102,18 +114,30 @@ const avatarId = computed(() => pickScopedAvatarId(route.params.avatarId, store.
|
|||||||
const activeTab = ref<'docs' | 'qa'>('docs')
|
const activeTab = ref<'docs' | 'qa'>('docs')
|
||||||
|
|
||||||
const docs = ref<any[]>([])
|
const docs = ref<any[]>([])
|
||||||
|
const pendingUploads = ref<any[]>([])
|
||||||
const qaPairs = ref<any[]>([])
|
const qaPairs = ref<any[]>([])
|
||||||
const uploading = ref(false)
|
const uploading = computed(() => pendingUploads.value.some((doc) => doc.localUploading))
|
||||||
const uploadError = ref('')
|
const uploadError = ref('')
|
||||||
const dragOver = ref(false)
|
const dragOver = ref(false)
|
||||||
const fileInput = ref<HTMLInputElement | null>(null)
|
const fileInput = ref<HTMLInputElement | null>(null)
|
||||||
|
const retryingDocs = ref<Record<string, boolean>>({})
|
||||||
|
const retryErrors = ref<Record<string, string>>({})
|
||||||
|
let documentPollingTimer: ReturnType<typeof setInterval> | undefined
|
||||||
|
|
||||||
const query = ref('')
|
const query = ref('')
|
||||||
const searching = ref(false)
|
const searching = ref(false)
|
||||||
const searched = ref(false)
|
const searched = ref(false)
|
||||||
const searchResults = ref<any[]>([])
|
const searchResults = ref<any[]>([])
|
||||||
|
|
||||||
|
const displayDocs = computed(() => [...pendingUploads.value, ...docs.value])
|
||||||
|
|
||||||
const documentState = (doc: any) => {
|
const documentState = (doc: any) => {
|
||||||
|
if (doc.localUploading) {
|
||||||
|
return { tone: 'pending', label: '上传中', detail: `正在上传 ${doc.uploadProgress || 0}%`, progress: doc.uploadProgress || 0 }
|
||||||
|
}
|
||||||
|
if (doc.localOnly) {
|
||||||
|
return { tone: 'failed', label: '上传失败', detail: doc.errorMessage || '文件未上传成功,请移除后重试' }
|
||||||
|
}
|
||||||
if (doc.filePresent === false) {
|
if (doc.filePresent === false) {
|
||||||
return { tone: 'missing', label: '文件缺失', detail: '原文件不可用,请删除后重新上传' }
|
return { tone: 'missing', label: '文件缺失', detail: '原文件不可用,请删除后重新上传' }
|
||||||
}
|
}
|
||||||
@@ -121,9 +145,33 @@ const documentState = (doc: any) => {
|
|||||||
return { tone: 'ready', label: '已入库', detail: `已切分 ${doc.chunkCount || 0} 段,可用于对话` }
|
return { tone: 'ready', label: '已入库', detail: `已切分 ${doc.chunkCount || 0} 段,可用于对话` }
|
||||||
}
|
}
|
||||||
if (['uploaded', 'parsing'].includes(String(doc.status || '').toLowerCase())) {
|
if (['uploaded', 'parsing'].includes(String(doc.status || '').toLowerCase())) {
|
||||||
return { tone: 'pending', label: '处理中', detail: '正在解析并建立知识索引' }
|
const stage = String(doc.indexStage || 'queued').toLowerCase()
|
||||||
|
const labels: Record<string, string> = {
|
||||||
|
queued: '等待处理', extracting: '解析文档', ocr: '扫描件识别', chunking: '切分文本', embedding: '向量化中'
|
||||||
|
}
|
||||||
|
const progress = Math.max(0, Math.min(99, Number(doc.indexProgress || 0)))
|
||||||
|
return { tone: 'pending', label: labels[stage] || '处理中', detail: `${labels[stage] || '正在建立知识索引'} ${progress}%`, progress }
|
||||||
}
|
}
|
||||||
return { tone: 'failed', label: '处理失败', detail: '未能建立知识索引,请删除后重新上传' }
|
return { tone: 'failed', label: '处理失败', detail: doc.errorMessage || '未能建立知识索引,请重新索引或重新上传' }
|
||||||
|
}
|
||||||
|
|
||||||
|
const hasPendingDocuments = () => docs.value.some((doc) =>
|
||||||
|
['uploaded', 'parsing'].includes(String(doc.status || '').toLowerCase())
|
||||||
|
)
|
||||||
|
|
||||||
|
const stopDocumentPolling = () => {
|
||||||
|
if (documentPollingTimer) {
|
||||||
|
clearInterval(documentPollingTimer)
|
||||||
|
documentPollingTimer = undefined
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
const startDocumentPolling = () => {
|
||||||
|
if (documentPollingTimer || !hasPendingDocuments()) return
|
||||||
|
documentPollingTimer = setInterval(async () => {
|
||||||
|
await loadDocs()
|
||||||
|
if (!hasPendingDocuments()) stopDocumentPolling()
|
||||||
|
}, 2000)
|
||||||
}
|
}
|
||||||
|
|
||||||
const loadDocs = async () => {
|
const loadDocs = async () => {
|
||||||
@@ -131,6 +179,7 @@ const loadDocs = async () => {
|
|||||||
try {
|
try {
|
||||||
const res: any = await getKnowledgeDocs(avatarId.value)
|
const res: any = await getKnowledgeDocs(avatarId.value)
|
||||||
docs.value = unwrapListData(res)
|
docs.value = unwrapListData(res)
|
||||||
|
startDocumentPolling()
|
||||||
} catch (e) {
|
} catch (e) {
|
||||||
console.error(e)
|
console.error(e)
|
||||||
}
|
}
|
||||||
@@ -149,40 +198,92 @@ const loadQA = async () => {
|
|||||||
const triggerFile = () => fileInput.value?.click()
|
const triggerFile = () => fileInput.value?.click()
|
||||||
|
|
||||||
const onFileChange = (e: Event) => {
|
const onFileChange = (e: Event) => {
|
||||||
const f = (e.target as HTMLInputElement).files?.[0]
|
const files = Array.from((e.target as HTMLInputElement).files || [])
|
||||||
if (f) doUpload(f)
|
if (files.length) uploadFiles(files)
|
||||||
;(e.target as HTMLInputElement).value = ''
|
;(e.target as HTMLInputElement).value = ''
|
||||||
}
|
}
|
||||||
|
|
||||||
const onDrop = (e: DragEvent) => {
|
const onDrop = (e: DragEvent) => {
|
||||||
dragOver.value = false
|
dragOver.value = false
|
||||||
const f = e.dataTransfer?.files?.[0]
|
const files = Array.from(e.dataTransfer?.files || [])
|
||||||
if (f) doUpload(f)
|
if (files.length) uploadFiles(files)
|
||||||
}
|
}
|
||||||
|
|
||||||
const doUpload = async (file: File) => {
|
const uploadFiles = (files: File[]) => {
|
||||||
uploadError.value = ''
|
uploadError.value = ''
|
||||||
const ext = '.' + (file.name.split('.').pop() || '').toLowerCase()
|
|
||||||
if (!['.md', '.txt', '.pdf', '.doc', '.docx', '.xlsx'].includes(ext)) {
|
|
||||||
uploadError.value = `不支持的类型:${ext},仅支持 md/txt/pdf/doc/docx/xlsx`
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if (!avatarId.value) {
|
if (!avatarId.value) {
|
||||||
uploadError.value = '请先创建数字分身'
|
uploadError.value = '请先创建数字分身'
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
uploading.value = true
|
for (const file of files) {
|
||||||
|
const ext = '.' + (file.name.split('.').pop() || '').toLowerCase()
|
||||||
|
if (!['.md', '.txt', '.pdf', '.doc', '.docx', '.xlsx'].includes(ext)) {
|
||||||
|
uploadError.value = `不支持的类型:${ext},仅支持 md/txt/pdf/doc/docx/xlsx`
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
void uploadOne(file, ext)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
const uploadOne = async (file: File, ext: string) => {
|
||||||
|
if (!avatarId.value) return
|
||||||
|
const localId = `upload-${Date.now()}-${Math.random().toString(16).slice(2)}`
|
||||||
|
const card = {
|
||||||
|
id: localId,
|
||||||
|
filename: file.name,
|
||||||
|
fileType: ext.slice(1),
|
||||||
|
fileSize: file.size,
|
||||||
|
createdAt: new Date().toISOString(),
|
||||||
|
localUploading: true,
|
||||||
|
localOnly: true,
|
||||||
|
uploadProgress: 0,
|
||||||
|
errorMessage: ''
|
||||||
|
}
|
||||||
|
pendingUploads.value.unshift(card)
|
||||||
try {
|
try {
|
||||||
await uploadKnowledgeDoc(avatarId.value, file)
|
const created: any = await uploadKnowledgeDoc(avatarId.value, file, (loaded, total) => {
|
||||||
await loadDocs()
|
const current = pendingUploads.value.find((doc) => doc.id === localId)
|
||||||
|
if (current) current.uploadProgress = Math.min(99, Math.round((loaded / Math.max(1, total)) * 100))
|
||||||
|
})
|
||||||
|
pendingUploads.value = pendingUploads.value.filter((doc) => doc.id !== localId)
|
||||||
|
docs.value = [created, ...docs.value.filter((doc) => doc.id !== created.id)]
|
||||||
|
startDocumentPolling()
|
||||||
} catch (e: any) {
|
} catch (e: any) {
|
||||||
uploadError.value = e?.message || '上传失败'
|
const current = pendingUploads.value.find((doc) => doc.id === localId)
|
||||||
|
if (current) {
|
||||||
|
current.localUploading = false
|
||||||
|
current.errorMessage = e?.message || '上传失败'
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
const canRetryDoc = (doc: any) =>
|
||||||
|
!doc.localOnly && doc.filePresent !== false && documentState(doc).tone === 'failed'
|
||||||
|
|
||||||
|
const retryDoc = async (doc: any) => {
|
||||||
|
if (!avatarId.value || !canRetryDoc(doc) || retryingDocs.value[doc.id]) return
|
||||||
|
retryingDocs.value = { ...retryingDocs.value, [doc.id]: true }
|
||||||
|
retryErrors.value = { ...retryErrors.value, [doc.id]: '' }
|
||||||
|
try {
|
||||||
|
const updated: any = await retryKnowledgeDoc(avatarId.value, doc.id)
|
||||||
|
Object.assign(doc, updated)
|
||||||
|
startDocumentPolling()
|
||||||
|
} catch (e: any) {
|
||||||
|
retryErrors.value = {
|
||||||
|
...retryErrors.value,
|
||||||
|
[doc.id]: e?.response?.data?.message || e?.response?.data?.detail || e?.message || '重新索引失败'
|
||||||
|
}
|
||||||
} finally {
|
} finally {
|
||||||
uploading.value = false
|
retryingDocs.value = { ...retryingDocs.value, [doc.id]: false }
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
const removeDoc = async (id: string) => {
|
const removeDoc = async (id: string) => {
|
||||||
|
const local = pendingUploads.value.find((doc) => doc.id === id)
|
||||||
|
if (local?.localOnly) {
|
||||||
|
pendingUploads.value = pendingUploads.value.filter((doc) => doc.id !== id)
|
||||||
|
return
|
||||||
|
}
|
||||||
if (!avatarId.value) return
|
if (!avatarId.value) return
|
||||||
await deleteKnowledgeDoc(avatarId.value, id)
|
await deleteKnowledgeDoc(avatarId.value, id)
|
||||||
await loadDocs()
|
await loadDocs()
|
||||||
@@ -257,6 +358,8 @@ onMounted(async () => {
|
|||||||
if (avatarId.value) store.currentAvatarId = avatarId.value
|
if (avatarId.value) store.currentAvatarId = avatarId.value
|
||||||
await Promise.all([loadDocs(), loadQA()])
|
await Promise.all([loadDocs(), loadQA()])
|
||||||
})
|
})
|
||||||
|
|
||||||
|
onUnmounted(stopDocumentPolling)
|
||||||
</script>
|
</script>
|
||||||
|
|
||||||
<style scoped>
|
<style scoped>
|
||||||
@@ -292,6 +395,7 @@ onMounted(async () => {
|
|||||||
.panel-heading p { margin: -5px 0 0; color: #9398AE; font-size: 12px; }
|
.panel-heading p { margin: -5px 0 0; color: #9398AE; font-size: 12px; }
|
||||||
.mobile-card-list { display: grid; grid-template-columns: minmax(0, 1fr); width: 100%; min-width: 0; gap: 10px; }
|
.mobile-card-list { display: grid; grid-template-columns: minmax(0, 1fr); width: 100%; min-width: 0; gap: 10px; }
|
||||||
.knowledge-card { display: flex; align-items: center; width: 100%; min-width: 0; box-sizing: border-box; gap: 11px; padding: 14px; background: #fff; border: 1px solid #F1E1D3; border-radius: 16px; box-shadow: 0 5px 16px rgba(112, 62, 22, .04); }
|
.knowledge-card { display: flex; align-items: center; width: 100%; min-width: 0; box-sizing: border-box; gap: 11px; padding: 14px; background: #fff; border: 1px solid #F1E1D3; border-radius: 16px; box-shadow: 0 5px 16px rgba(112, 62, 22, .04); }
|
||||||
|
.document-card { display: grid; grid-template-columns: 42px minmax(0, 1fr) auto; align-items: center; }
|
||||||
.card-icon { flex: 0 0 auto; width: 42px; height: 42px; display: grid; place-items: center; border-radius: 13px; background: #FFF3E6; font-size: 22px; }
|
.card-icon { flex: 0 0 auto; width: 42px; height: 42px; display: grid; place-items: center; border-radius: 13px; background: #FFF3E6; font-size: 22px; }
|
||||||
.card-content { min-width: 0; flex: 1; overflow: hidden; }
|
.card-content { min-width: 0; flex: 1; overflow: hidden; }
|
||||||
.card-title-row { display: flex; align-items: center; gap: 8px; min-width: 0; }
|
.card-title-row { display: flex; align-items: center; gap: 8px; min-width: 0; }
|
||||||
@@ -300,7 +404,15 @@ onMounted(async () => {
|
|||||||
.status-pill.missing { color: #B91C1C; background: #FEF2F2; }
|
.status-pill.missing { color: #B91C1C; background: #FEF2F2; }
|
||||||
.status-pill.failed { color: #B91C1C; background: #FEF2F2; }
|
.status-pill.failed { color: #B91C1C; background: #FEF2F2; }
|
||||||
.card-meta, .card-detail { margin: 5px 0 0; color: #9398AE; font-size: 11px; line-height: 1.4; }.card-detail { color: #8B6B58; }
|
.card-meta, .card-detail { margin: 5px 0 0; color: #9398AE; font-size: 11px; line-height: 1.4; }.card-detail { color: #8B6B58; }
|
||||||
.card-delete { flex: 0 0 auto; align-self: center; border: 0; color: #EF4444; background: #FEF2F2; border-radius: 8px; padding: 7px 9px; font-size: 12px; cursor: pointer; }
|
.progress-track { width: 100%; height: 4px; margin-top: 8px; overflow: hidden; border-radius: 999px; background: #FDE7D1; }
|
||||||
|
.progress-fill { display: block; height: 100%; border-radius: inherit; background: linear-gradient(90deg, #FB923C, #F97316); transition: width .25s ease; }
|
||||||
|
.card-actions { flex: 0 0 auto; display: flex; align-items: center; }
|
||||||
|
.card-delete, .card-retry { align-self: center; border: 0; border-radius: 8px; padding: 7px 9px; font-size: 12px; cursor: pointer; white-space: nowrap; }
|
||||||
|
.card-delete { color: #EF4444; background: #FEF2F2; }
|
||||||
|
.card-retry { color: #C15F18; background: #FFF3E6; }
|
||||||
|
.card-retry:disabled { cursor: wait; opacity: .65; }
|
||||||
|
.card-retry-area { grid-column: 1 / -1; display: flex; align-items: center; justify-content: flex-end; gap: 10px; min-width: 0; }
|
||||||
|
.card-retry-error { min-width: 0; overflow: hidden; color: #DC2626; font-size: 11px; line-height: 1.35; text-overflow: ellipsis; white-space: nowrap; }
|
||||||
.card-empty { padding: 42px 16px; border: 1px dashed #F1D9C3; border-radius: 16px; color: #9398AE; background: #fff; font-size: 14px; text-align: center; }
|
.card-empty { padding: 42px 16px; border: 1px dashed #F1D9C3; border-radius: 16px; color: #9398AE; background: #fff; font-size: 14px; text-align: center; }
|
||||||
.qa-card { align-items: stretch; text-align: left; }.qa-card.qa-disabled { opacity: .58; }
|
.qa-card { align-items: stretch; text-align: left; }.qa-card.qa-disabled { opacity: .58; }
|
||||||
.qa-card .card-content,
|
.qa-card .card-content,
|
||||||
@@ -526,8 +638,11 @@ onMounted(async () => {
|
|||||||
@media (max-width: 520px) {
|
@media (max-width: 520px) {
|
||||||
.knowledge-panel { padding: 0 12px; }
|
.knowledge-panel { padding: 0 12px; }
|
||||||
.knowledge-card { display: grid; grid-template-columns: 42px minmax(0, 1fr); align-items: start; gap: 10px; padding: 13px; }
|
.knowledge-card { display: grid; grid-template-columns: 42px minmax(0, 1fr); align-items: start; gap: 10px; padding: 13px; }
|
||||||
|
.document-card { grid-template-columns: 42px minmax(0, 1fr) auto; }
|
||||||
.card-content { grid-column: 2; }
|
.card-content { grid-column: 2; }
|
||||||
.card-delete { grid-column: 2; justify-self: end; margin-top: -2px; }
|
.card-actions { grid-column: 3; grid-row: 1; }
|
||||||
|
.card-delete { justify-self: end; margin-top: -2px; }
|
||||||
|
.card-retry-area { grid-column: 1 / -1; }
|
||||||
.qa-card { display: block; }
|
.qa-card { display: block; }
|
||||||
.qa-card .card-content { width: 100%; grid-column: 1; }
|
.qa-card .card-content { width: 100%; grid-column: 1; }
|
||||||
.card-title-row { align-items: flex-start; flex-wrap: wrap; gap: 5px 7px; }
|
.card-title-row { align-items: flex-start; flex-wrap: wrap; gap: 5px 7px; }
|
||||||
|
|||||||
@@ -74,6 +74,35 @@
|
|||||||
{{ checkoutLabel }}
|
{{ checkoutLabel }}
|
||||||
</button>
|
</button>
|
||||||
</section>
|
</section>
|
||||||
|
|
||||||
|
<section v-if="recentOrders.length" class="orders-section">
|
||||||
|
<h3 class="section-title">充值记录</h3>
|
||||||
|
<div v-for="order in recentOrders" :key="order.id" class="order-card">
|
||||||
|
<div>
|
||||||
|
<strong>{{ order.pointsAmount.toLocaleString() }} 积分</strong>
|
||||||
|
<p>{{ order.orderNo }} · {{ order.createdAt ? new Date(order.createdAt).toLocaleDateString('zh-CN') : '' }}</p>
|
||||||
|
</div>
|
||||||
|
<div class="order-side">
|
||||||
|
<strong>¥{{ order.price.toFixed(2) }}</strong>
|
||||||
|
<button v-if="canInvoice(order)" class="text-btn" @click="openInvoice(order)">申请发票</button>
|
||||||
|
<span v-else class="order-status">{{ orderStatus(order) }}</span>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</section>
|
||||||
|
|
||||||
|
<div v-if="invoiceOrder" class="modal-mask" @click.self="invoiceOrder = null">
|
||||||
|
<form class="invoice-modal" @submit.prevent="submitInvoice">
|
||||||
|
<h3>申请电子发票</h3>
|
||||||
|
<label>发票类型
|
||||||
|
<select v-model="invoiceType"><option value="personal">个人</option><option value="company">企业</option></select>
|
||||||
|
</label>
|
||||||
|
<label>发票抬头<input v-model.trim="invoiceTitle" maxlength="120" required /></label>
|
||||||
|
<label v-if="invoiceType === 'company'">企业税号<input v-model.trim="invoiceTaxNumber" minlength="15" maxlength="20" required /></label>
|
||||||
|
<label>接收邮箱<input v-model.trim="invoiceEmail" type="email" placeholder="选填" /></label>
|
||||||
|
<p v-if="invoiceError" class="invoice-error">{{ invoiceError }}</p>
|
||||||
|
<div class="modal-actions"><button type="button" @click="invoiceOrder = null">取消</button><button class="primary" :disabled="invoiceSubmitting">{{ invoiceSubmitting ? '提交中…' : '提交申请' }}</button></div>
|
||||||
|
</form>
|
||||||
|
</div>
|
||||||
</div>
|
</div>
|
||||||
</template>
|
</template>
|
||||||
|
|
||||||
@@ -82,8 +111,10 @@ import { computed, onMounted, onUnmounted, ref } from 'vue'
|
|||||||
import { useRouter } from 'vue-router'
|
import { useRouter } from 'vue-router'
|
||||||
import {
|
import {
|
||||||
chargeToken,
|
chargeToken,
|
||||||
|
applyTokenInvoice,
|
||||||
getRechargePlans,
|
getRechargePlans,
|
||||||
getTokenBalance,
|
getTokenBalance,
|
||||||
|
getTokenOrders,
|
||||||
getTokenPaymentStatus,
|
getTokenPaymentStatus,
|
||||||
type TokenPaymentOrder
|
type TokenPaymentOrder
|
||||||
} from '@/api'
|
} from '@/api'
|
||||||
@@ -114,6 +145,14 @@ const paymentMethod = ref<'wechat' | 'alipay'>('wechat')
|
|||||||
const paymentNotice = ref('')
|
const paymentNotice = ref('')
|
||||||
const paymentNoticeTone = ref<'pending' | 'success' | 'error'>('pending')
|
const paymentNoticeTone = ref<'pending' | 'success' | 'error'>('pending')
|
||||||
const pendingOrderId = ref(sessionStorage.getItem('hh_pending_payment_order') || '')
|
const pendingOrderId = ref(sessionStorage.getItem('hh_pending_payment_order') || '')
|
||||||
|
const recentOrders = ref<Array<TokenPaymentOrder & { invoice?: any }>>([])
|
||||||
|
const invoiceOrder = ref<(TokenPaymentOrder & { invoice?: any }) | null>(null)
|
||||||
|
const invoiceType = ref<'personal' | 'company'>('personal')
|
||||||
|
const invoiceTitle = ref('')
|
||||||
|
const invoiceTaxNumber = ref('')
|
||||||
|
const invoiceEmail = ref('')
|
||||||
|
const invoiceError = ref('')
|
||||||
|
const invoiceSubmitting = ref(false)
|
||||||
let pollTimer: number | undefined
|
let pollTimer: number | undefined
|
||||||
let pollDeadline = 0
|
let pollDeadline = 0
|
||||||
let removeNativeListener: (() => void) | undefined
|
let removeNativeListener: (() => void) | undefined
|
||||||
@@ -133,6 +172,51 @@ const loadData = async () => {
|
|||||||
} catch (e) {
|
} catch (e) {
|
||||||
console.error('加载套餐失败', e)
|
console.error('加载套餐失败', e)
|
||||||
}
|
}
|
||||||
|
try {
|
||||||
|
const result = await getTokenOrders(1, 10)
|
||||||
|
recentOrders.value = result?.items || []
|
||||||
|
} catch (e) {
|
||||||
|
console.error('加载充值记录失败', e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
const canInvoice = (order: TokenPaymentOrder & { invoice?: any }) =>
|
||||||
|
order.status === 'paid' && (!order.refundStatus || order.refundStatus === 'none') &&
|
||||||
|
(!order.invoice || ['rejected', 'cancelled'].includes(order.invoice.status))
|
||||||
|
const orderStatus = (order: TokenPaymentOrder & { invoice?: any }) => {
|
||||||
|
if (order.invoice?.status === 'issued') return '发票已开具'
|
||||||
|
if (order.invoice?.status === 'pending') return '发票处理中'
|
||||||
|
if (order.invoice?.status === 'rejected') return '发票已驳回'
|
||||||
|
return ({ pending: '待支付', paid: '已支付', failed: '支付失败', closed: '已关闭', refunded: '已退款' } as Record<string, string>)[order.status] || order.status
|
||||||
|
}
|
||||||
|
const openInvoice = (order: TokenPaymentOrder & { invoice?: any }) => {
|
||||||
|
invoiceOrder.value = order
|
||||||
|
invoiceType.value = 'personal'
|
||||||
|
invoiceTitle.value = ''
|
||||||
|
invoiceTaxNumber.value = ''
|
||||||
|
invoiceEmail.value = ''
|
||||||
|
invoiceError.value = ''
|
||||||
|
}
|
||||||
|
const submitInvoice = async () => {
|
||||||
|
if (!invoiceOrder.value || invoiceSubmitting.value) return
|
||||||
|
invoiceSubmitting.value = true
|
||||||
|
invoiceError.value = ''
|
||||||
|
try {
|
||||||
|
await applyTokenInvoice(invoiceOrder.value.orderNo, {
|
||||||
|
title: invoiceTitle.value,
|
||||||
|
invoiceType: invoiceType.value,
|
||||||
|
taxNumber: invoiceTaxNumber.value,
|
||||||
|
email: invoiceEmail.value
|
||||||
|
})
|
||||||
|
paymentNoticeTone.value = 'success'
|
||||||
|
paymentNotice.value = '发票申请已提交,请等待财务处理'
|
||||||
|
invoiceOrder.value = null
|
||||||
|
await loadData()
|
||||||
|
} catch (error: any) {
|
||||||
|
invoiceError.value = error?.message || '发票申请提交失败'
|
||||||
|
} finally {
|
||||||
|
invoiceSubmitting.value = false
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 会会支付订单创建与到账确认
|
// 会会支付订单创建与到账确认
|
||||||
@@ -151,8 +235,9 @@ const checkoutLabel = computed(() => {
|
|||||||
})
|
})
|
||||||
|
|
||||||
const payScene = () => {
|
const payScene = () => {
|
||||||
|
if (isInUniWebView()) return 'APP' as const
|
||||||
if (paymentMethod.value === 'wechat' && /MicroMessenger/i.test(navigator.userAgent)) return 'JSAPI' as const
|
if (paymentMethod.value === 'wechat' && /MicroMessenger/i.test(navigator.userAgent)) return 'JSAPI' as const
|
||||||
return 'APP' as const
|
return 'H5' as const
|
||||||
}
|
}
|
||||||
|
|
||||||
const parsePayMessage = (message: string) => {
|
const parsePayMessage = (message: string) => {
|
||||||
@@ -231,6 +316,7 @@ const pollPayment = async () => {
|
|||||||
paymentNoticeTone.value = 'success'
|
paymentNoticeTone.value = 'success'
|
||||||
paymentNotice.value = `支付成功,${order.pointsAmount.toLocaleString()} 积分已到账`
|
paymentNotice.value = `支付成功,${order.pointsAmount.toLocaleString()} 积分已到账`
|
||||||
clearPendingOrder()
|
clearPendingOrder()
|
||||||
|
void loadData()
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if (order.status === 'failed') {
|
if (order.status === 'failed') {
|
||||||
@@ -562,6 +648,17 @@ onUnmounted(() => {
|
|||||||
padding: 0 20px;
|
padding: 0 20px;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
.orders-section { padding: 24px 20px 0; }
|
||||||
|
.order-card { display:flex; align-items:center; justify-content:space-between; padding:14px 16px; margin-bottom:10px; background:#fff; border:1px solid #EDEEF1; border-radius:12px; }
|
||||||
|
.order-card strong { color:#18191C; font-size:14px; }.order-card p,.order-status { color:#9398AE; font-size:11px; margin:5px 0 0; }
|
||||||
|
.order-side { text-align:right; }.text-btn { display:block; margin-top:5px; padding:0; border:0; background:none; color:#F97316; font-size:12px; cursor:pointer; }
|
||||||
|
.modal-mask { position:fixed; inset:0; z-index:20; display:grid; place-items:center; padding:20px; background:rgba(15,23,42,.45); }
|
||||||
|
.invoice-modal { width:min(100%,420px); padding:22px; border-radius:16px; background:#fff; box-shadow:0 18px 50px rgba(15,23,42,.2); }
|
||||||
|
.invoice-modal h3 { margin:0 0 18px; }.invoice-modal label { display:grid; gap:7px; margin:12px 0; color:#4B5563; font-size:13px; }
|
||||||
|
.invoice-modal input,.invoice-modal select { width:100%; height:42px; padding:0 12px; border:1px solid #D9DCE3; border-radius:9px; background:#fff; color:#18191C; font-size:14px; }
|
||||||
|
.invoice-error { color:#B42318; font-size:12px; }.modal-actions { display:flex; justify-content:flex-end; gap:10px; margin-top:20px; }
|
||||||
|
.modal-actions button { padding:9px 18px; border:1px solid #D9DCE3; border-radius:9px; background:#fff; }.modal-actions .primary { border-color:#F97316; background:#F97316; color:#fff; }
|
||||||
|
|
||||||
.checkout-btn {
|
.checkout-btn {
|
||||||
width: 100%;
|
width: 100%;
|
||||||
padding: 16px;
|
padding: 16px;
|
||||||
|
|||||||
+3
-1
@@ -20,11 +20,13 @@ services:
|
|||||||
- AVATAR_MODEL_CONFIG_TOKEN=${AVATAR_MODEL_CONFIG_TOKEN:-}
|
- AVATAR_MODEL_CONFIG_TOKEN=${AVATAR_MODEL_CONFIG_TOKEN:-}
|
||||||
- TZ=Asia/Shanghai
|
- TZ=Asia/Shanghai
|
||||||
- AVATAR_DB_PATH=/app/avatar.db
|
- AVATAR_DB_PATH=/app/avatar.db
|
||||||
|
- AVATAR_BACKEND_URL=${AVATAR_BACKEND_URL:-}
|
||||||
|
- AVATAR_FINANCE_ADMIN_SECRET=${AVATAR_FINANCE_ADMIN_SECRET:-}
|
||||||
volumes:
|
volumes:
|
||||||
- ./backend/app:/app/app # ← 核心:代码目录直接挂载,改文件无需重建
|
- ./backend/app:/app/app # ← 核心:代码目录直接挂载,改文件无需重建
|
||||||
- ./backend/logs:/app/logs
|
- ./backend/logs:/app/logs
|
||||||
- ./backend/config:/app/config
|
- ./backend/config:/app/config
|
||||||
- ./digital-avatar-app/backend/avatar.db:/app/avatar.db:ro # 数字分身 SQLite(只读)
|
- ./digital-avatar-app/backend/avatar.db:/app/avatar.db # 财务管理需要写入退款与开票状态
|
||||||
depends_on:
|
depends_on:
|
||||||
- ai-virtual-mysql
|
- ai-virtual-mysql
|
||||||
- ai-virtual-redis
|
- ai-virtual-redis
|
||||||
|
|||||||
@@ -94,5 +94,15 @@ export const uploadAvatarPhoto = (id, formData) => request.post(`/avatars/${id}/
|
|||||||
headers: { 'Content-Type': 'multipart/form-data' }
|
headers: { 'Content-Type': 'multipart/form-data' }
|
||||||
})
|
})
|
||||||
|
|
||||||
|
// Finance (数字分身积分订单)
|
||||||
|
export const getFinanceSummary = () => request.get('/finance/summary')
|
||||||
|
export const getFinanceOrders = (params) => request.get('/finance/orders', { params })
|
||||||
|
export const updateFinanceOrderStatus = (orderNo, data) => request.patch(`/finance/orders/${orderNo}/status`, data)
|
||||||
|
export const requestFinanceRefund = (orderNo, data) => request.post(`/finance/orders/${orderNo}/refund`, data)
|
||||||
|
export const getFinanceRefunds = (params) => request.get('/finance/refunds', { params })
|
||||||
|
export const confirmFinanceRefund = (refundNo, data) => request.post(`/finance/refunds/${refundNo}/confirm`, data)
|
||||||
|
export const getFinanceInvoices = (params) => request.get('/finance/invoices', { params })
|
||||||
|
export const updateFinanceInvoice = (invoiceId, data) => request.patch(`/finance/invoices/${invoiceId}`, data)
|
||||||
|
|
||||||
export default request
|
export default request
|
||||||
export const uploadAvatar = (userId, formData) => request.post(`/users/${userId}/upload-avatar`, formData, { headers: { "Content-Type": "multipart/form-data" } })
|
export const uploadAvatar = (userId, formData) => request.post(`/users/${userId}/upload-avatar`, formData, { headers: { "Content-Type": "multipart/form-data" } })
|
||||||
|
|||||||
@@ -15,6 +15,10 @@
|
|||||||
<el-icon><UserFilled /></el-icon>
|
<el-icon><UserFilled /></el-icon>
|
||||||
<span>数字分身管理</span>
|
<span>数字分身管理</span>
|
||||||
</el-menu-item>
|
</el-menu-item>
|
||||||
|
<el-menu-item index="/finance">
|
||||||
|
<el-icon><WalletFilled /></el-icon>
|
||||||
|
<span>财务管理</span>
|
||||||
|
</el-menu-item>
|
||||||
<el-menu-item index="/users">
|
<el-menu-item index="/users">
|
||||||
<el-icon><User /></el-icon>
|
<el-icon><User /></el-icon>
|
||||||
<span>虚拟用户</span>
|
<span>虚拟用户</span>
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ const routes = [
|
|||||||
{ path: '', redirect: '/dashboard' },
|
{ path: '', redirect: '/dashboard' },
|
||||||
{ path: 'dashboard', component: () => import('@/views/Dashboard.vue'), meta: { title: '数据看板' } },
|
{ path: 'dashboard', component: () => import('@/views/Dashboard.vue'), meta: { title: '数据看板' } },
|
||||||
{ path: 'avatars', component: () => import('@/views/Avatars.vue'), meta: { title: '数字分身管理' } },
|
{ path: 'avatars', component: () => import('@/views/Avatars.vue'), meta: { title: '数字分身管理' } },
|
||||||
|
{ path: 'finance', component: () => import('@/views/Finance.vue'), meta: { title: '财务管理' } },
|
||||||
{ path: 'users', component: () => import('@/views/Users.vue'), meta: { title: '虚拟用户管理' } },
|
{ path: 'users', component: () => import('@/views/Users.vue'), meta: { title: '虚拟用户管理' } },
|
||||||
{ path: 'interactions', component: () => import('@/views/Interactions.vue'), meta: { title: '互动记录' } },
|
{ path: 'interactions', component: () => import('@/views/Interactions.vue'), meta: { title: '互动记录' } },
|
||||||
{ path: 'ai-models', component: () => import('@/views/AIModels.vue'), meta: { title: 'AI模型配置' } },
|
{ path: 'ai-models', component: () => import('@/views/AIModels.vue'), meta: { title: 'AI模型配置' } },
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,142 @@
|
|||||||
|
<template>
|
||||||
|
<div class="page-container finance-page">
|
||||||
|
<div class="page-header">
|
||||||
|
<div>
|
||||||
|
<div class="page-title">财务管理</div>
|
||||||
|
<div class="subtitle">数字分身积分订单、退款与发票处理</div>
|
||||||
|
</div>
|
||||||
|
<el-button :loading="loading" @click="loadAll"><el-icon><Refresh /></el-icon>刷新</el-button>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<div class="summary-grid">
|
||||||
|
<div class="stat-card"><div class="stat-label">实收金额</div><div class="stat-value">¥{{ money(summary.paid_revenue) }}</div><small>{{ summary.paid_orders }} 笔已支付</small></div>
|
||||||
|
<div class="stat-card"><div class="stat-label">待支付订单</div><div class="stat-value warning">{{ summary.pending_orders }}</div><small>可关闭或标记失败</small></div>
|
||||||
|
<div class="stat-card"><div class="stat-label">处理中退款</div><div class="stat-value danger">{{ summary.processing_refunds }}</div><small>等待支付渠道确认</small></div>
|
||||||
|
<div class="stat-card"><div class="stat-label">待开发票</div><div class="stat-value success">{{ summary.pending_invoices }}</div><small>等待财务处理</small></div>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<el-tabs v-model="activeTab" class="finance-tabs" @tab-change="loadActive">
|
||||||
|
<el-tab-pane label="支付订单" name="orders">
|
||||||
|
<div class="filter-bar">
|
||||||
|
<el-input v-model="filters.keyword" placeholder="订单号、昵称或手机号" clearable style="width:240px" @keyup.enter="loadOrders" />
|
||||||
|
<el-select v-model="filters.status" placeholder="订单状态" clearable style="width:140px" @change="loadOrders">
|
||||||
|
<el-option v-for="item in orderStatuses" :key="item.value" :label="item.label" :value="item.value" />
|
||||||
|
</el-select>
|
||||||
|
<el-select v-model="filters.provider" placeholder="支付渠道" clearable style="width:150px" @change="loadOrders">
|
||||||
|
<el-option label="会会支付" value="huihui" /><el-option label="微信虚拟支付" value="wechat_virtual" />
|
||||||
|
</el-select>
|
||||||
|
<el-button type="primary" @click="loadOrders">查询</el-button>
|
||||||
|
</div>
|
||||||
|
<el-table :data="orders" v-loading="loading" class="data-table">
|
||||||
|
<el-table-column label="订单 / 用户" min-width="230">
|
||||||
|
<template #default="{ row }"><b>{{ row.order_no }}</b><div class="muted">{{ row.user_nickname || '会会用户' }} · {{ row.user_phone || '--' }}</div></template>
|
||||||
|
</el-table-column>
|
||||||
|
<el-table-column label="套餐积分" width="130"><template #default="{ row }">{{ Number(row.points_amount || 0).toLocaleString() }}</template></el-table-column>
|
||||||
|
<el-table-column label="金额" width="100"><template #default="{ row }">¥{{ money(row.price) }}</template></el-table-column>
|
||||||
|
<el-table-column label="渠道" width="130"><template #default="{ row }">{{ providerLabel[row.provider] || row.provider }}</template></el-table-column>
|
||||||
|
<el-table-column label="支付状态" width="110"><template #default="{ row }"><el-tag :type="tagType(row.status)">{{ statusLabel[row.status] || row.status }}</el-tag></template></el-table-column>
|
||||||
|
<el-table-column label="退款" width="100"><template #default="{ row }"><el-tag v-if="row.refund_status && row.refund_status !== 'none'" :type="tagType(row.refund_status)">{{ statusLabel[row.refund_status] || row.refund_status }}</el-tag><span v-else>--</span></template></el-table-column>
|
||||||
|
<el-table-column label="下单时间" min-width="165"><template #default="{ row }">{{ dateTime(row.created_at) }}</template></el-table-column>
|
||||||
|
<el-table-column label="操作" width="165" fixed="right">
|
||||||
|
<template #default="{ row }">
|
||||||
|
<el-button v-if="row.status === 'paid' && ['none','',null].includes(row.refund_status)" link type="danger" @click="refund(row)">退款</el-button>
|
||||||
|
<el-dropdown v-if="row.status === 'pending'" @command="command => closeOrder(row, command)">
|
||||||
|
<el-button link type="warning">状态处理<el-icon><ArrowDown /></el-icon></el-button>
|
||||||
|
<template #dropdown><el-dropdown-menu><el-dropdown-item command="closed">关闭订单</el-dropdown-item><el-dropdown-item command="failed">标记失败</el-dropdown-item></el-dropdown-menu></template>
|
||||||
|
</el-dropdown>
|
||||||
|
<span v-if="row.status !== 'pending' && !(row.status === 'paid' && ['none','',null].includes(row.refund_status))" class="muted">已处理</span>
|
||||||
|
</template>
|
||||||
|
</el-table-column>
|
||||||
|
</el-table>
|
||||||
|
<el-pagination v-model:current-page="orderPage" :page-size="20" :total="orderTotal" layout="total, prev, pager, next" @change="loadOrders" />
|
||||||
|
</el-tab-pane>
|
||||||
|
|
||||||
|
<el-tab-pane label="退款管理" name="refunds">
|
||||||
|
<el-table :data="refunds" v-loading="loading" class="data-table">
|
||||||
|
<el-table-column label="退款单号" min-width="220" prop="refund_no" />
|
||||||
|
<el-table-column label="原订单" min-width="210" prop="order_no" />
|
||||||
|
<el-table-column label="用户" width="150"><template #default="{ row }">{{ row.user_nickname || row.user_phone || '--' }}</template></el-table-column>
|
||||||
|
<el-table-column label="金额" width="100"><template #default="{ row }">¥{{ money(row.amount) }}</template></el-table-column>
|
||||||
|
<el-table-column label="状态" width="110"><template #default="{ row }"><el-tag :type="tagType(row.status)">{{ statusLabel[row.status] || row.status }}</el-tag></template></el-table-column>
|
||||||
|
<el-table-column label="原因" min-width="160" prop="reason" show-overflow-tooltip />
|
||||||
|
<el-table-column label="创建时间" min-width="165"><template #default="{ row }">{{ dateTime(row.created_at) }}</template></el-table-column>
|
||||||
|
<el-table-column label="对账确认" width="155" fixed="right">
|
||||||
|
<template #default="{ row }"><template v-if="['pending','processing'].includes(row.status)"><el-button link type="success" @click="confirmRefund(row, 'succeeded')">已退款</el-button><el-button link type="danger" @click="confirmRefund(row, 'failed')">失败</el-button></template><span v-else class="muted">{{ row.failure_reason || '已完成' }}</span></template>
|
||||||
|
</el-table-column>
|
||||||
|
</el-table>
|
||||||
|
</el-tab-pane>
|
||||||
|
|
||||||
|
<el-tab-pane label="发票管理" name="invoices">
|
||||||
|
<el-table :data="invoices" v-loading="loading" class="data-table">
|
||||||
|
<el-table-column label="订单号" min-width="210" prop="order_no" />
|
||||||
|
<el-table-column label="抬头" min-width="180" prop="title" />
|
||||||
|
<el-table-column label="类型 / 税号" min-width="190"><template #default="{ row }">{{ row.invoice_type === 'company' ? '企业' : '个人' }}<div class="muted">{{ row.tax_number || '--' }}</div></template></el-table-column>
|
||||||
|
<el-table-column label="金额" width="100"><template #default="{ row }">¥{{ money(row.amount) }}</template></el-table-column>
|
||||||
|
<el-table-column label="邮箱" min-width="175" prop="email" />
|
||||||
|
<el-table-column label="状态" width="100"><template #default="{ row }"><el-tag :type="tagType(row.status)">{{ invoiceStatusLabel[row.status] || row.status }}</el-tag></template></el-table-column>
|
||||||
|
<el-table-column label="操作" width="145" fixed="right"><template #default="{ row }"><template v-if="row.status === 'pending'"><el-button link type="primary" @click="openInvoice(row)">开具</el-button><el-button link type="danger" @click="rejectInvoice(row)">驳回</el-button></template><a v-else-if="row.invoice_url" :href="row.invoice_url" target="_blank">查看发票</a><span v-else class="muted">{{ row.invoice_no || row.remark || '已处理' }}</span></template></el-table-column>
|
||||||
|
</el-table>
|
||||||
|
</el-tab-pane>
|
||||||
|
</el-tabs>
|
||||||
|
|
||||||
|
<el-dialog v-model="invoiceDialog" title="登记已开发票" width="480px">
|
||||||
|
<el-form label-width="92px">
|
||||||
|
<el-form-item label="发票号码"><el-input v-model="invoiceForm.invoiceNo" maxlength="120" /></el-form-item>
|
||||||
|
<el-form-item label="发票地址"><el-input v-model="invoiceForm.invoiceUrl" placeholder="电子发票下载地址(选填)" /></el-form-item>
|
||||||
|
<el-form-item label="备注"><el-input v-model="invoiceForm.remark" type="textarea" /></el-form-item>
|
||||||
|
</el-form>
|
||||||
|
<template #footer><el-button @click="invoiceDialog=false">取消</el-button><el-button type="primary" :loading="saving" @click="issueInvoice">确认开具</el-button></template>
|
||||||
|
</el-dialog>
|
||||||
|
</div>
|
||||||
|
</template>
|
||||||
|
|
||||||
|
<script setup>
|
||||||
|
import { onMounted, ref } from 'vue'
|
||||||
|
import { ElMessage, ElMessageBox } from 'element-plus'
|
||||||
|
import { confirmFinanceRefund, getFinanceInvoices, getFinanceOrders, getFinanceRefunds, getFinanceSummary, requestFinanceRefund, updateFinanceInvoice, updateFinanceOrderStatus } from '@/api'
|
||||||
|
|
||||||
|
const activeTab = ref('orders')
|
||||||
|
const loading = ref(false)
|
||||||
|
const saving = ref(false)
|
||||||
|
const summary = ref({ paid_revenue: 0, paid_orders: 0, pending_orders: 0, processing_refunds: 0, pending_invoices: 0 })
|
||||||
|
const orders = ref([]), refunds = ref([]), invoices = ref([])
|
||||||
|
const orderPage = ref(1), orderTotal = ref(0)
|
||||||
|
const filters = ref({ keyword: '', status: '', provider: '' })
|
||||||
|
const invoiceDialog = ref(false), currentInvoice = ref(null)
|
||||||
|
const invoiceForm = ref({ invoiceNo: '', invoiceUrl: '', remark: '' })
|
||||||
|
const orderStatuses = [{ label: '待支付', value: 'pending' }, { label: '已支付', value: 'paid' }, { label: '失败', value: 'failed' }, { label: '已关闭', value: 'closed' }, { label: '已退款', value: 'refunded' }]
|
||||||
|
const providerLabel = { huihui: '会会支付', wechat_virtual: '微信虚拟支付' }
|
||||||
|
const statusLabel = { pending: '待处理', processing: '处理中', paid: '已支付', succeeded: '已成功', failed: '失败', closed: '已关闭', refunded: '已退款' }
|
||||||
|
const invoiceStatusLabel = { pending: '待开具', issued: '已开具', rejected: '已驳回' }
|
||||||
|
const money = value => Number(value || 0).toFixed(2)
|
||||||
|
const dateTime = value => value ? new Date(value).toLocaleString('zh-CN', { hour12: false }) : '--'
|
||||||
|
const tagType = value => ({ paid: 'success', succeeded: 'success', issued: 'success', pending: 'warning', processing: 'warning', failed: 'danger', rejected: 'danger', refunded: 'info', closed: 'info' }[value] || 'info')
|
||||||
|
|
||||||
|
async function loadSummary() { const res = await getFinanceSummary(); summary.value = res.data || summary.value }
|
||||||
|
async function loadOrders() { const res = await getFinanceOrders({ page: orderPage.value, page_size: 20, ...filters.value }); orders.value = res.data?.items || []; orderTotal.value = res.data?.total || 0 }
|
||||||
|
async function loadRefunds() { const res = await getFinanceRefunds({ page: 1, page_size: 100 }); refunds.value = res.data?.items || [] }
|
||||||
|
async function loadInvoices() { const res = await getFinanceInvoices({ page: 1, page_size: 100 }); invoices.value = res.data?.items || [] }
|
||||||
|
async function loadActive() { loading.value = true; try { if (activeTab.value === 'orders') await loadOrders(); if (activeTab.value === 'refunds') await loadRefunds(); if (activeTab.value === 'invoices') await loadInvoices() } finally { loading.value = false } }
|
||||||
|
async function loadAll() { loading.value = true; try { await Promise.all([loadSummary(), loadOrders(), loadRefunds(), loadInvoices()]) } finally { loading.value = false } }
|
||||||
|
|
||||||
|
async function refund(row) {
|
||||||
|
try { const { value } = await ElMessageBox.prompt('退款成功后会扣回本订单发放的全部积分。', `退款 ¥${money(row.price)}`, { inputPlaceholder: '请输入退款原因', inputValidator: value => value?.trim() ? true : '退款原因不能为空', confirmButtonText: '提交退款', cancelButtonText: '取消' }); await requestFinanceRefund(row.order_no, { reason: value, operator: '后台管理员' }); ElMessage.success('退款已提交渠道处理'); await loadAll() } catch (error) { if (error !== 'cancel' && error !== 'close') console.error(error) }
|
||||||
|
}
|
||||||
|
async function closeOrder(row, status) { try { await ElMessageBox.confirm(`确认将订单标记为“${status === 'closed' ? '已关闭' : '失败'}”?`, '订单状态处理', { type: 'warning' }); await updateFinanceOrderStatus(row.order_no, { status, reason: '后台管理员处理' }); ElMessage.success('订单状态已更新'); await loadAll() } catch {} }
|
||||||
|
async function confirmRefund(row, status) { try { const message = status === 'succeeded' ? '仅在支付渠道后台已确认退款到账时操作。' : '确认供应商退款失败?'; await ElMessageBox.confirm(message, '退款对账确认', { type: 'warning' }); await confirmFinanceRefund(row.refund_no, { status, failureReason: status === 'failed' ? '供应商退款失败' : '' }); ElMessage.success('退款结果已登记'); await loadAll() } catch {} }
|
||||||
|
function openInvoice(row) { currentInvoice.value = row; invoiceForm.value = { invoiceNo: '', invoiceUrl: '', remark: '' }; invoiceDialog.value = true }
|
||||||
|
async function issueInvoice() { if (!invoiceForm.value.invoiceNo.trim()) return ElMessage.warning('请填写发票号码'); saving.value = true; try { await updateFinanceInvoice(currentInvoice.value.id, { status: 'issued', ...invoiceForm.value }); ElMessage.success('发票已登记'); invoiceDialog.value = false; await loadAll() } finally { saving.value = false } }
|
||||||
|
async function rejectInvoice(row) { try { const { value } = await ElMessageBox.prompt('请输入驳回原因', '驳回发票申请', { inputValidator: value => value?.trim() ? true : '驳回原因不能为空' }); await updateFinanceInvoice(row.id, { status: 'rejected', remark: value }); ElMessage.success('发票申请已驳回'); await loadAll() } catch {} }
|
||||||
|
onMounted(loadAll)
|
||||||
|
</script>
|
||||||
|
|
||||||
|
<style scoped>
|
||||||
|
.finance-page { overflow-y: auto; }
|
||||||
|
.subtitle, .muted, small { color: var(--color-text-muted); font-size: 12px; margin-top: 5px; }
|
||||||
|
.summary-grid { display:grid; grid-template-columns:repeat(4,minmax(0,1fr)); gap:16px; margin-bottom:20px; }
|
||||||
|
.stat-value { margin:8px 0 4px; }.stat-value.warning{color:var(--color-accent-orange)}.stat-value.danger{color:var(--color-accent-red)}.stat-value.success{color:var(--color-accent-green)}
|
||||||
|
.finance-tabs { background:#fff; border:1px solid var(--color-border); border-radius:12px; padding:8px 16px 16px; box-shadow:var(--shadow-sm); }
|
||||||
|
.filter-bar { display:flex; gap:10px; margin:8px 0 16px; }.el-pagination{justify-content:flex-end;margin-top:16px}
|
||||||
|
a { color:var(--color-accent); text-decoration:none; }
|
||||||
|
@media (max-width: 1100px) { .summary-grid{grid-template-columns:repeat(2,1fr)} }
|
||||||
|
</style>
|
||||||
@@ -61,7 +61,7 @@ H5 引入 uniapp web-view bridge 后调用:
|
|||||||
壳通过 `web-view.evalJS` 调用 H5 全局函数 `window.__uniBridgeHandle__(message)`:
|
壳通过 `web-view.evalJS` 调用 H5 全局函数 `window.__uniBridgeHandle__(message)`:
|
||||||
| type | payload | 含义 |
|
| type | payload | 含义 |
|
||||||
|------|---------|------|
|
|------|---------|------|
|
||||||
| `context` | `platform, version` | 注入运行环境信息 |
|
| `context` | `surface, version` | 注入运行环境;`surface` 为 `app` / `mp-weixin` / `h5` |
|
||||||
| `tokenRefresh` | `token` | 登录刷新后下发新 token |
|
| `tokenRefresh` | `token` | 登录刷新后下发新 token |
|
||||||
| `userUpdate` | `user` | 会会资料变更 |
|
| `userUpdate` | `user` | 会会资料变更 |
|
||||||
| `paymentResult` | `orderId,status` | 原生支付结束通知;`status` 为 `success/cancelled/failed` |
|
| `paymentResult` | `orderId,status` | 原生支付结束通知;`status` 为 `success/cancelled/failed` |
|
||||||
@@ -70,6 +70,8 @@ H5 引入 uniapp web-view bridge 后调用:
|
|||||||
|
|
||||||
壳收到 `payment` 后应调用会会 App 已有的微信/支付宝支付能力(或 `uni.requestPayment`),把 `payMessage/paymentParams` 原样交给对应渠道。原生 SDK 返回后再发送 `paymentResult`;H5 不以原生返回作为到账依据,只会轮询本地订单,最终由会会服务端支付回调确认并增加积分。
|
壳收到 `payment` 后应调用会会 App 已有的微信/支付宝支付能力(或 `uni.requestPayment`),把 `payMessage/paymentParams` 原样交给对应渠道。原生 SDK 返回后再发送 `paymentResult`;H5 不以原生返回作为到账依据,只会轮询本地订单,最终由会会服务端支付回调确认并增加积分。
|
||||||
|
|
||||||
|
微信小程序虚拟支付本期只交付后端能力(登录态交换、签名下单参数、服务端查单/退款和回调验收)。小程序原生充值页接入 `requestVirtualPayment` 后,应把后端返回的 `signData/paySig/signature/mode/env/offerId` 原样传入微信 API;不要在 web-view 中发起虚拟支付。
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## 3. 项目结构(uni CLI / src 布局,已验证可编译)
|
## 3. 项目结构(uni CLI / src 布局,已验证可编译)
|
||||||
@@ -84,13 +86,13 @@ uniapp-avatar/
|
|||||||
├── manifest.json # 应用配置(名称/AppID/模块)
|
├── manifest.json # 应用配置(名称/AppID/模块)
|
||||||
├── pages.json # 页面路由
|
├── pages.json # 页面路由
|
||||||
├── uni.scss # 全局样式变量
|
├── uni.scss # 全局样式变量
|
||||||
├── App.vue # 启动即做会会登录(onLaunch → userStore.init)
|
├── App.vue # 启动时恢复会会登录态(onLaunch → userStore.init)
|
||||||
├── main.js # createSSRApp + pinia
|
├── main.js # createSSRApp + pinia
|
||||||
├── pages/index/index.vue # web-view 容器(内嵌 digital-avatar-app H5)
|
├── pages/index/index.vue # web-view 容器(内嵌 digital-avatar-app H5)
|
||||||
├── store/user.js # 会会会话(token/资料,本地缓存)
|
├── store/user.js # 会会会话(宿主调用 applySession 注入并缓存)
|
||||||
└── utils/
|
└── utils/
|
||||||
├── bridge.js # H5 URL 构造 + 原生→H5 推送
|
├── bridge.js # H5 URL 构造 + 原生→H5 推送
|
||||||
└── huihui.js # 会会登录(MOCK,留真实接入位)
|
└── payment.js # App 微信/支付宝原生支付适配
|
||||||
```
|
```
|
||||||
> 构建产物:`npm run build:h5` → `dist/build/h5/`(含 index.html + assets)。
|
> 构建产物:`npm run build:h5` → `dist/build/h5/`(含 index.html + assets)。
|
||||||
|
|
||||||
@@ -114,10 +116,9 @@ npm run build:h5 # 生产构建 → dist/build/h5/
|
|||||||
> 若 npm 依赖版本与本地 HBuilderX 不一致,执行 `npx @dcloudio/uvm` 对齐。
|
> 若 npm 依赖版本与本地 HBuilderX 不一致,执行 `npx @dcloudio/uvm` 对齐。
|
||||||
|
|
||||||
### 会会登录接入
|
### 会会登录接入
|
||||||
- 当前 `utils/huihui.js` 为 **MOCK**(`MOCK_AUTH = true`),便于联调。
|
- 壳只恢复会会宿主已经持有的登录态,不内置演示账号,也不会伪造会会 token。
|
||||||
- 生产接入:把 `MOCK_AUTH` 改为 `false`,在 `loginHuihui()` 接入会会开放平台授权,
|
- 会会主 App 完成登录或刷新后调用 `userStore.applySession({ token, userId, nickname, avatarUrl })`;数字分身 H5 会把一次性会会 token 换成本系统会话并立即从地址中清除。
|
||||||
换取 `access_token`、`userId`;`getUserInfo()` 请求会会 `usercenter` 真实资料接口
|
- 独立打开且没有宿主会话时,H5 会进入已有的短信登录流程。
|
||||||
(接口基址见 `docs/production-interface-inventory.md`)。
|
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|||||||
Generated
+7937
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,28 @@
|
|||||||
|
{
|
||||||
|
"name": "uniapp-avatar",
|
||||||
|
"version": "1.1.0",
|
||||||
|
"private": true,
|
||||||
|
"type": "module",
|
||||||
|
"scripts": {
|
||||||
|
"dev:h5": "uni -p h5",
|
||||||
|
"build:h5": "uni build -p h5",
|
||||||
|
"dev:mp-weixin": "uni -p mp-weixin",
|
||||||
|
"build:mp-weixin": "uni build -p mp-weixin",
|
||||||
|
"dev:app": "uni -p app",
|
||||||
|
"build:app": "uni build -p app"
|
||||||
|
},
|
||||||
|
"dependencies": {
|
||||||
|
"@dcloudio/uni-app": "3.0.0-4010520240507001",
|
||||||
|
"@dcloudio/uni-app-plus": "3.0.0-4010520240507001",
|
||||||
|
"@dcloudio/uni-components": "3.0.0-4010520240507001",
|
||||||
|
"@dcloudio/uni-h5": "3.0.0-4010520240507001",
|
||||||
|
"@dcloudio/uni-mp-weixin": "3.0.0-4010520240507001",
|
||||||
|
"pinia": "^2.0.36",
|
||||||
|
"vue": "^3.4.21"
|
||||||
|
},
|
||||||
|
"devDependencies": {
|
||||||
|
"@dcloudio/uni-cli-shared": "3.0.0-4010520240507001",
|
||||||
|
"@dcloudio/vite-plugin-uni": "3.0.0-4010520240507001",
|
||||||
|
"vite": "^5.2.0"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,10 @@
|
|||||||
|
<script setup>
|
||||||
|
import { onLaunch } from '@dcloudio/uni-app'
|
||||||
|
import { useUserStore } from '@/store/user'
|
||||||
|
|
||||||
|
onLaunch(() => useUserStore().init())
|
||||||
|
</script>
|
||||||
|
|
||||||
|
<style>
|
||||||
|
page { background:#f5f6f8; }
|
||||||
|
</style>
|
||||||
@@ -0,0 +1,9 @@
|
|||||||
|
import { createSSRApp } from 'vue'
|
||||||
|
import { createPinia } from 'pinia'
|
||||||
|
import App from './App.vue'
|
||||||
|
|
||||||
|
export function createApp() {
|
||||||
|
const app = createSSRApp(App)
|
||||||
|
app.use(createPinia())
|
||||||
|
return { app }
|
||||||
|
}
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
{
|
||||||
|
"name": "会会数字分身",
|
||||||
|
"appid": "",
|
||||||
|
"description": "会会数字分身 uni-app 混合壳",
|
||||||
|
"versionName": "1.1.0",
|
||||||
|
"versionCode": "110",
|
||||||
|
"transformPx": false,
|
||||||
|
"app-plus": {
|
||||||
|
"usingComponents": true,
|
||||||
|
"compilerVersion": 3,
|
||||||
|
"modules": { "Payment": {} },
|
||||||
|
"distribute": {
|
||||||
|
"android": { "permissions": ["<uses-permission android:name=\"android.permission.INTERNET\"/>"] },
|
||||||
|
"ios": {},
|
||||||
|
"sdkConfigs": {}
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"mp-weixin": { "appid": "", "setting": { "urlCheck": true }, "usingComponents": true },
|
||||||
|
"h5": { "publicPath": "./", "router": { "base": "./" } },
|
||||||
|
"vueVersion": "3"
|
||||||
|
}
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
{
|
||||||
|
"pages": [{
|
||||||
|
"path": "pages/index/index",
|
||||||
|
"style": {
|
||||||
|
"navigationBarTitleText": "会会数字分身",
|
||||||
|
"navigationBarBackgroundColor": "#0F2B4C",
|
||||||
|
"navigationBarTextStyle": "white"
|
||||||
|
}
|
||||||
|
}],
|
||||||
|
"globalStyle": {
|
||||||
|
"navigationBarTextStyle": "white",
|
||||||
|
"navigationBarTitleText": "会会数字分身",
|
||||||
|
"navigationBarBackgroundColor": "#0F2B4C",
|
||||||
|
"backgroundColor": "#F5F6F8"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,55 @@
|
|||||||
|
<template>
|
||||||
|
<view class="shell">
|
||||||
|
<web-view ref="webview" :src="h5Url" @message="onH5Message" />
|
||||||
|
</view>
|
||||||
|
</template>
|
||||||
|
|
||||||
|
<script setup>
|
||||||
|
import { computed, ref } from 'vue'
|
||||||
|
import { useUserStore } from '@/store/user'
|
||||||
|
import { buildH5Url, parseH5Message, postToH5 } from '@/utils/bridge'
|
||||||
|
import { requestAppPayment } from '@/utils/payment'
|
||||||
|
|
||||||
|
const userStore = useUserStore()
|
||||||
|
const webview = ref(null)
|
||||||
|
const H5_BASE_URL = import.meta.env.VITE_AVATAR_H5_URL || 'https://digital.99hui.com/'
|
||||||
|
const h5Url = computed(() => buildH5Url(H5_BASE_URL, userStore.$state))
|
||||||
|
|
||||||
|
function runtimeSurface() {
|
||||||
|
let value = 'h5'
|
||||||
|
// #ifdef APP-PLUS
|
||||||
|
value = 'app'
|
||||||
|
// #endif
|
||||||
|
// #ifdef MP-WEIXIN
|
||||||
|
value = 'mp-weixin'
|
||||||
|
// #endif
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
|
async function handlePayment(payment) {
|
||||||
|
try {
|
||||||
|
await requestAppPayment(payment)
|
||||||
|
postToH5(webview.value, { type: 'paymentResult', orderId: payment.orderId, status: 'success' })
|
||||||
|
} catch (error) {
|
||||||
|
const message = error?.message || '支付未完成'
|
||||||
|
const status = /cancel/i.test(message) || /取消/.test(message) ? 'cancelled' : 'failed'
|
||||||
|
postToH5(webview.value, { type: 'paymentResult', orderId: payment.orderId, status, message })
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
function onH5Message(event) {
|
||||||
|
const message = parseH5Message(event)
|
||||||
|
if (!message?.type) return
|
||||||
|
if (message.type === 'ready') {
|
||||||
|
postToH5(webview.value, { type: 'context', surface: runtimeSurface(), version: '1.1.0' })
|
||||||
|
} else if (message.type === 'payment') {
|
||||||
|
void handlePayment(message.payment)
|
||||||
|
} else if (message.type === 'setTitle' && message.title) {
|
||||||
|
uni.setNavigationBarTitle({ title: message.title })
|
||||||
|
} else if (message.type === 'back') {
|
||||||
|
uni.navigateBack({ delta: 1 })
|
||||||
|
}
|
||||||
|
}
|
||||||
|
</script>
|
||||||
|
|
||||||
|
<style>.shell{width:100%;height:100vh}</style>
|
||||||
@@ -0,0 +1,24 @@
|
|||||||
|
import { defineStore } from 'pinia'
|
||||||
|
|
||||||
|
export const useUserStore = defineStore('user', {
|
||||||
|
state: () => ({ token: '', userId: '', nickname: '', avatarUrl: '', initialized: false }),
|
||||||
|
getters: { isLogin: state => Boolean(state.token) },
|
||||||
|
actions: {
|
||||||
|
init() {
|
||||||
|
this.$patch({
|
||||||
|
token: uni.getStorageSync('hh_token') || '',
|
||||||
|
userId: uni.getStorageSync('hh_userId') || '',
|
||||||
|
nickname: uni.getStorageSync('hh_nickname') || '',
|
||||||
|
avatarUrl: uni.getStorageSync('hh_avatarUrl') || '',
|
||||||
|
initialized: true
|
||||||
|
})
|
||||||
|
},
|
||||||
|
applySession(session) {
|
||||||
|
this.$patch(session)
|
||||||
|
uni.setStorageSync('hh_token', this.token)
|
||||||
|
uni.setStorageSync('hh_userId', this.userId)
|
||||||
|
uni.setStorageSync('hh_nickname', this.nickname)
|
||||||
|
uni.setStorageSync('hh_avatarUrl', this.avatarUrl)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
$brand-navy: #0f2b4c;
|
||||||
|
$brand-orange: #f97316;
|
||||||
|
$bg-page: #f5f6f8;
|
||||||
@@ -0,0 +1,33 @@
|
|||||||
|
export function buildH5Url(base, session) {
|
||||||
|
const url = new URL(base)
|
||||||
|
if (session.token) url.searchParams.set('token', session.token)
|
||||||
|
if (session.userId) url.searchParams.set('userId', session.userId)
|
||||||
|
if (session.nickname) url.searchParams.set('nickname', session.nickname)
|
||||||
|
if (session.avatarUrl) url.searchParams.set('avatar', session.avatarUrl)
|
||||||
|
url.searchParams.set('ts', String(Date.now()))
|
||||||
|
return url.toString()
|
||||||
|
}
|
||||||
|
|
||||||
|
export function postToH5(webviewRef, message) {
|
||||||
|
if (!webviewRef) return false
|
||||||
|
const js = `window.__uniBridgeHandle__&&window.__uniBridgeHandle__(${JSON.stringify(message)})`
|
||||||
|
try {
|
||||||
|
if (typeof webviewRef.evalJS === 'function') {
|
||||||
|
webviewRef.evalJS(js)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
// #ifdef APP-PLUS
|
||||||
|
const child = plus.webview.currentWebview().children()[0]
|
||||||
|
if (child) { child.evalJS(js); return true }
|
||||||
|
// #endif
|
||||||
|
} catch (error) {
|
||||||
|
console.error('[bridge] postToH5 failed', error)
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
export function parseH5Message(event) {
|
||||||
|
const items = event?.detail?.data
|
||||||
|
if (Array.isArray(items) && items.length) return items[items.length - 1]
|
||||||
|
return event?.detail || null
|
||||||
|
}
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
function parseOrderInfo(payment) {
|
||||||
|
if (payment.paymentParams && typeof payment.paymentParams === 'object') return payment.paymentParams
|
||||||
|
if (typeof payment.payMessage !== 'string') return payment.payMessage
|
||||||
|
try { return JSON.parse(payment.payMessage) } catch { return payment.payMessage }
|
||||||
|
}
|
||||||
|
|
||||||
|
export function requestAppPayment(payment) {
|
||||||
|
if (!payment || payment.payWay !== 'APP') {
|
||||||
|
return Promise.reject(new Error('当前订单不是 App 支付订单'))
|
||||||
|
}
|
||||||
|
const provider = payment.paymentMethod === 'alipay' ? 'alipay' : 'wxpay'
|
||||||
|
const orderInfo = parseOrderInfo(payment)
|
||||||
|
if (!orderInfo) return Promise.reject(new Error('支付参数为空'))
|
||||||
|
return new Promise((resolve, reject) => {
|
||||||
|
uni.requestPayment({
|
||||||
|
provider,
|
||||||
|
orderInfo,
|
||||||
|
success: resolve,
|
||||||
|
fail: error => reject(new Error(error?.errMsg || '支付未完成'))
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -0,0 +1,5 @@
|
|||||||
|
import { defineConfig } from 'vite'
|
||||||
|
import uni from '@dcloudio/vite-plugin-uni'
|
||||||
|
|
||||||
|
const uniPlugin = typeof uni === 'function' ? uni : uni.default
|
||||||
|
export default defineConfig({ plugins: [uniPlugin()] })
|
||||||
Reference in New Issue
Block a user