Compare commits
14 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 866d71a431 | |||
| 042512a527 | |||
| 3bb9c5dd4e | |||
| c87810a4a6 | |||
| 75ec9db439 | |||
| 0d70074182 | |||
| 071a3707f4 | |||
| 4104401759 | |||
| f0d46b44d7 | |||
| d63f7b4650 | |||
| b28ac8bd1b | |||
| 1b02df4d4a | |||
| 76e11cb2f9 | |||
| 1baf29c76a |
@@ -220,6 +220,20 @@ DOUBAO_MAX_RETRIES=2
|
||||
DOUBAO_VISION_MODEL=doubao-1-5-vision-pro-250328
|
||||
DOUBAO_VISION_LITE_MODEL=doubao-1-5-vision-lite-250315
|
||||
DOUBAO_VISION_USE_LITE=true
|
||||
# Embedding 向量化模型(原 large-text-240915 已下线,用多模态 embedding)
|
||||
DOUBAO_EMBEDDING_MODEL=doubao-embedding-vision-251215
|
||||
|
||||
# ==================== 即梦(Jimeng)视觉 API —— 真人参考图兜底通道 (#2169) ====
|
||||
# 方舟 Seedance 走 B 端审核,真人参考图会被 50411 拦截;即梦走 C 端审核,普通真人照片可过审。
|
||||
# 需要在火山控制台开通即梦 cvtob 服务,使用 AK/SK(Region=cn-north-1, Service=cv)
|
||||
# 留空则真人拦截后直接返回错误提示,不会走即梦兜底。
|
||||
JIMENG_AK=
|
||||
JIMENG_SK=
|
||||
JIMENG_BASE_URL=https://visual.volcengineapi.com
|
||||
# 视频3.0 720P 首帧图生视频(1张图,稳定版;1080P标注下线中)
|
||||
JIMENG_REQ_KEY=jimeng_i2v_first_v30
|
||||
JIMENG_VIDEO_TIMEOUT=600
|
||||
JIMENG_VIDEO_POLL_INTERVAL=5
|
||||
|
||||
# ==================== 积分/会员系统 (#1895) ====================
|
||||
# 积分系统总开关:默认 false(暂停积分系统)。
|
||||
|
||||
@@ -1,50 +0,0 @@
|
||||
"""points_orders 新增微信支付链路字段
|
||||
|
||||
Revision ID: 094_points_orders_wechat_fields
|
||||
Revises: 093
|
||||
Create Date: 2026-10-03
|
||||
|
||||
新增列:
|
||||
- prepay_id: 微信预支付ID
|
||||
- product_name: 下单商品名称(冗余,便于对账)
|
||||
- payer_openid: 支付者 openid
|
||||
- expire_at: 订单过期时间(未支付超时关闭用)
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "094_points_orders_wechat_fields"
|
||||
down_revision = "093"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
cols = {c["name"] for c in sa.inspect(conn).get_columns("points_orders")}
|
||||
|
||||
if "out_trade_no" not in cols:
|
||||
op.add_column("points_orders", sa.Column("out_trade_no", sa.String(64), nullable=True))
|
||||
op.create_index("ix_points_orders_out_trade_no", "points_orders", ["out_trade_no"])
|
||||
if "prepay_id" not in cols:
|
||||
op.add_column("points_orders", sa.Column("prepay_id", sa.String(128), nullable=True))
|
||||
if "product_name" not in cols:
|
||||
op.add_column("points_orders", sa.Column("product_name", sa.String(100), nullable=True))
|
||||
if "payer_openid" not in cols:
|
||||
op.add_column("points_orders", sa.Column("payer_openid", sa.String(128), nullable=True))
|
||||
if "expire_at" not in cols:
|
||||
op.add_column("points_orders", sa.Column("expire_at", sa.DateTime(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
cols = {c["name"] for c in sa.inspect(conn).get_columns("points_orders")}
|
||||
try:
|
||||
op.drop_index("ix_points_orders_out_trade_no", table_name="points_orders")
|
||||
except Exception:
|
||||
pass
|
||||
for col in ("out_trade_no", "prepay_id", "product_name", "payer_openid", "expire_at"):
|
||||
if col in cols:
|
||||
op.drop_column("points_orders", col)
|
||||
@@ -21,7 +21,6 @@ from app.api.routes.health import router as health_check_router
|
||||
from app.api.routes.ingest_jobs import router as ingest_jobs_router
|
||||
from app.api.routes.internal_render import router as internal_render_router
|
||||
from app.api.routes.lipsync import router as lipsync_router
|
||||
from app.api.routes.payment import router as payment_router
|
||||
from app.api.routes.points import router as points_router
|
||||
from app.api.routes.points import usage_router
|
||||
from app.api.routes.projects import router as projects_router
|
||||
@@ -227,11 +226,6 @@ api_router.include_router(
|
||||
prefix="/ai-avatar/render",
|
||||
tags=["AI Avatar Render"],
|
||||
)
|
||||
api_router.include_router(
|
||||
payment_router,
|
||||
prefix="/payment",
|
||||
tags=["Payment"],
|
||||
)
|
||||
api_router.include_router(
|
||||
points_router,
|
||||
prefix="/points",
|
||||
|
||||
@@ -1,69 +0,0 @@
|
||||
"""微信支付回调路由。
|
||||
|
||||
POST /api/v1/payment/wechat/notify
|
||||
- 不做用户鉴权(微信服务器调用),靠平台证书签名保证来源可信
|
||||
- 返回微信要求的 JSON:{"code": "SUCCESS", "message": "成功"}
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from app.dependencies import get_db_session
|
||||
from fastapi import APIRouter, Depends, Request, Response
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.wechat_cert_store import get_platform_public_key
|
||||
from packages.adapters.wechat_pay import WeChatPayError
|
||||
from packages.application.payment_service import PaymentConfigError, PaymentService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _wx_response(code: str, message: str, http_status: int = 200) -> Response:
|
||||
"""构造微信要求的回调应答。"""
|
||||
import json
|
||||
|
||||
return Response(
|
||||
content=json.dumps({"code": code, "message": message}, ensure_ascii=False),
|
||||
media_type="application/json",
|
||||
status_code=http_status,
|
||||
)
|
||||
|
||||
|
||||
@router.post("/wechat/notify")
|
||||
async def wechat_pay_notify(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db_session),
|
||||
):
|
||||
"""微信支付结果通知。
|
||||
|
||||
- 验签失败 / 配置缺失 / 解密失败 → 返回 FAIL(微信会重试)
|
||||
- 履约成功(含重复通知幂等)→ 返回 SUCCESS
|
||||
"""
|
||||
body = await request.body()
|
||||
headers = {k.lower(): v for k, v in request.headers.items()}
|
||||
|
||||
svc = PaymentService()
|
||||
try:
|
||||
result = svc.handle_wechat_notification(
|
||||
headers=headers,
|
||||
body=body,
|
||||
platform_public_key_loader=get_platform_public_key,
|
||||
db=db,
|
||||
)
|
||||
except PaymentConfigError as e:
|
||||
logger.error("微信回调时支付配置异常: %s", e)
|
||||
return _wx_response("ERROR", "支付配置异常", http_status=503)
|
||||
except WeChatPayError as e:
|
||||
logger.warning("微信回调处理失败: %s", e)
|
||||
return _wx_response("FAIL", str(e)[:200])
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception("微信回调未知异常")
|
||||
return _wx_response("FAIL", "系统繁忙")
|
||||
|
||||
if result.get("ignored"):
|
||||
return _wx_response("SUCCESS", "成功")
|
||||
return _wx_response("SUCCESS", "成功")
|
||||
@@ -8,7 +8,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Optional
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
@@ -61,14 +61,8 @@ def _get_service() -> PointsService:
|
||||
|
||||
|
||||
def _is_member(user: AuthenticatedUser) -> bool:
|
||||
"""判断用户是否为有效付费会员(实时判断到期时间)。"""
|
||||
if not getattr(user.user, "is_member", False):
|
||||
return False
|
||||
expires_at = getattr(user.user, "member_expires_at", None)
|
||||
if expires_at is not None and expires_at <= datetime.now(timezone.utc):
|
||||
return False
|
||||
# 已标记取消但未到期:会员权益仍有效
|
||||
return True
|
||||
"""判断用户是否为付费会员。"""
|
||||
return getattr(user.user, "is_member", False)
|
||||
|
||||
|
||||
def _member_type(user: AuthenticatedUser) -> str | None:
|
||||
@@ -151,19 +145,22 @@ def get_rules(
|
||||
def get_packages(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""查询可购买的积分包列表。"""
|
||||
packages = []
|
||||
for code, pkg in POINTS_PACKAGES.items():
|
||||
unit_price = f"¥{pkg['price_cents'] / 100 / pkg['points']:.3f}/积分"
|
||||
packages.append(
|
||||
PointsPackageItem(
|
||||
code=code,
|
||||
name=pkg["name"],
|
||||
points=pkg["points"],
|
||||
price_cents=pkg["price_cents"],
|
||||
unit_price=unit_price,
|
||||
)
|
||||
"""查询可购买的积分包列表(读管理后台 credit_packages 表真实数据)。
|
||||
|
||||
仅返回 is_active=true;后台改价/启停后最多 30 秒生效。
|
||||
"""
|
||||
from packages.application.catalog.admin_catalog import get_points_packages
|
||||
|
||||
packages = [
|
||||
PointsPackageItem(
|
||||
code=row["code"],
|
||||
name=row["name"],
|
||||
points=row["points"],
|
||||
price_cents=row["price_cents"],
|
||||
unit_price=row["unit_price"],
|
||||
)
|
||||
for row in get_points_packages()
|
||||
]
|
||||
mt = _member_type(current_user)
|
||||
discount = MEMBER_DISCOUNT.get(mt) if mt else None
|
||||
return PointsPackagesResponse(packages=packages, user_discount=discount)
|
||||
@@ -287,55 +284,32 @@ def refund_points(
|
||||
)
|
||||
|
||||
|
||||
@points_router.post("/create-order", response_model=PointsOrderResponse)
|
||||
def create_points_purchase_order(
|
||||
body: PointsRechargeRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
db: Session = Depends(get_db_session),
|
||||
):
|
||||
"""创建积分包购买订单(微信 JSAPI 支付)。
|
||||
|
||||
返回小程序调起支付参数;支付成功后积分自动到账(微信回调履约)。
|
||||
"""
|
||||
from packages.adapters.wechat_pay import WeChatPayError
|
||||
from packages.application.payment_service import PaymentConfigError, PaymentService
|
||||
|
||||
openid = getattr(current_user.user, "wechat_openid", None)
|
||||
svc = PaymentService()
|
||||
try:
|
||||
result = svc.create_points_order(
|
||||
user_id=current_user.user.id,
|
||||
package_code=body.package_id,
|
||||
openid=openid,
|
||||
db=db,
|
||||
)
|
||||
except PaymentConfigError as e:
|
||||
raise HTTPException(status_code=503, detail=f"支付通道不可用:{e}") from None
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e)) from None
|
||||
except WeChatPayError as e:
|
||||
raise HTTPException(status_code=502, detail=f"微信支付下单失败:{e}") from None
|
||||
|
||||
return PointsOrderResponse(
|
||||
id=result["order_id"],
|
||||
order_type="points",
|
||||
product_code=body.package_id,
|
||||
amount_cents=result["amount_cents"],
|
||||
points_amount=result["points_amount"],
|
||||
status="pending",
|
||||
pay_params=result["pay_params"],
|
||||
expire_at=result["expire_at"],
|
||||
)
|
||||
|
||||
|
||||
@points_router.post("/recharge", response_model=PointsOrderResponse, deprecated=True)
|
||||
@points_router.post("/recharge", response_model=PointsOrderResponse)
|
||||
def create_recharge_order(
|
||||
body: PointsRechargeRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
db: Session = Depends(get_db_session),
|
||||
):
|
||||
"""[已废弃] 旧积分充值入口,保留向后兼容,内部转发到 /points/create-order。"""
|
||||
return create_points_purchase_order(body, current_user, db)
|
||||
"""创建积分充值订单。pay_params 在支付通道接入后填入 prepay_id/payment_url;当前为空 dict。"""
|
||||
svc = _get_service()
|
||||
try:
|
||||
order = svc.create_order(
|
||||
user_id=current_user.user.id,
|
||||
order_type="points",
|
||||
product_code=body.package_id,
|
||||
db=db,
|
||||
)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e)) from None
|
||||
|
||||
package = POINTS_PACKAGES.get(body.package_id, {})
|
||||
now = datetime.now(timezone.utc)
|
||||
expire_at = now + timedelta(hours=48)
|
||||
# TODO: 接入微信/支付宝后填充真实 prepay_id / payment_url
|
||||
order["points_amount"] = package.get("points", 0)
|
||||
order["pay_params"] = {}
|
||||
order["expire_at"] = expire_at.isoformat()
|
||||
return PointsOrderResponse(**order)
|
||||
|
||||
|
||||
@points_router.get("/subscription/membership", response_model=MembershipStatusResponse)
|
||||
|
||||
@@ -8,23 +8,18 @@ from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import get_db_session, get_user_repository
|
||||
from app.dependencies import get_user_repository
|
||||
from app.schemas.subscription import (
|
||||
BillingCycle,
|
||||
BillingRecord,
|
||||
ChangePlanRequest,
|
||||
ChangePlanResponse,
|
||||
CreateOrderRequest,
|
||||
CreateOrderResponse,
|
||||
MembershipType,
|
||||
OrderItem,
|
||||
OrderListResponse,
|
||||
SimpleResponse,
|
||||
SubscriptionInfo,
|
||||
ToggleAutoRenewRequest,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
from sqlalchemy.orm import Session
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
|
||||
from packages.ports.user_repository import UserRepository
|
||||
|
||||
@@ -48,53 +43,30 @@ def _get_plan_name(plan_id: str) -> str:
|
||||
|
||||
|
||||
def _build_subscription_info(user: AuthenticatedUser) -> SubscriptionInfo:
|
||||
"""构建订阅信息响应(P1-8:实时判断是否过期)。"""
|
||||
"""构建订阅信息响应"""
|
||||
now = datetime.now(UTC)
|
||||
|
||||
expires_at = user.user.subscription_expires_at
|
||||
# 优先使用会员体系字段
|
||||
member_expires = getattr(user.user, "member_expires_at", None)
|
||||
if member_expires is not None:
|
||||
expires_at = member_expires
|
||||
if user.user.subscription_expires_at:
|
||||
period_end = user.user.subscription_expires_at.isoformat()
|
||||
period_start = now.isoformat()
|
||||
else:
|
||||
period_start = now.isoformat()
|
||||
period_end = now.isoformat()
|
||||
|
||||
plan_id = user.user.subscription_plan or MembershipType.FREE
|
||||
# 会员类型字段(member_type 与积分体系一致)
|
||||
member_type = getattr(user.user, "member_type", None)
|
||||
if member_type:
|
||||
plan_id = member_type
|
||||
# 旧档位(standard/pro/enterprise)统一降级为 monthly,避免前端炸掉
|
||||
if plan_id in {"standard", "pro", "enterprise"}:
|
||||
plan_id = MembershipType.MONTHLY
|
||||
|
||||
# 实时过期判断:到期即降级 free / expired(不改库,查询时计算)
|
||||
is_expired = expires_at is not None and expires_at <= now
|
||||
cancelled = (user.user.subscription_status or "") == "cancelled"
|
||||
|
||||
if is_expired:
|
||||
effective_plan = MembershipType.FREE
|
||||
effective_status = "expired"
|
||||
elif cancelled:
|
||||
effective_plan = plan_id
|
||||
effective_status = "cancelled"
|
||||
else:
|
||||
effective_plan = plan_id
|
||||
effective_status = user.user.subscription_status or (
|
||||
"active" if plan_id != MembershipType.FREE else "active"
|
||||
)
|
||||
|
||||
period_start = now.isoformat()
|
||||
period_end = expires_at.isoformat() if expires_at else now.isoformat()
|
||||
|
||||
return SubscriptionInfo(
|
||||
id=f"sub-{user.user.id[:8]}",
|
||||
plan_id=effective_plan,
|
||||
plan_name=_get_plan_name(effective_plan),
|
||||
status=effective_status,
|
||||
plan_id=plan_id,
|
||||
plan_name=_get_plan_name(plan_id),
|
||||
status=user.user.subscription_status or "active",
|
||||
billing_cycle=plan_id if plan_id != MembershipType.FREE else BillingCycle.MONTHLY,
|
||||
current_period_start=period_start,
|
||||
current_period_end=period_end,
|
||||
amount=0, # 金额由前端 /plans 接口展示
|
||||
auto_renew=False, # 一期不做自动续费
|
||||
amount=0 if plan_id == MembershipType.FREE else 0, # 金额由前端 /plans 接口展示
|
||||
auto_renew=True,
|
||||
created_at=user.user.created_at.isoformat() if user.user.created_at else now.isoformat(),
|
||||
)
|
||||
|
||||
@@ -114,33 +86,13 @@ async def get_current_subscription(
|
||||
def list_membership_plans(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> dict[str, list[dict[str, Any]]]:
|
||||
"""查询所有会员档位(供前端会员购买页展示)。
|
||||
"""查询可购买的会员套餐(读管理后台 plans 表真实数据)。
|
||||
|
||||
返回 points 积分体系下的会员档位(月卡/季卡/年卡),含价格、时长、积分折扣等信息。
|
||||
仅返回 is_enabled=true 的套餐;后台启停/改价后最多 30 秒生效。
|
||||
"""
|
||||
from packages.domain.points_rules import MEMBER_DISCOUNT, MEMBERSHIP_PRICES
|
||||
from packages.application.catalog.admin_catalog import get_membership_plans
|
||||
|
||||
plans: list[dict[str, Any]] = []
|
||||
for plan_id, info in MEMBERSHIP_PRICES.items():
|
||||
days = info["duration_days"]
|
||||
monthly_cents = round(info["price_cents"] * 30 / days)
|
||||
features: dict[str, Any] = {"max_resolution": "1080p"}
|
||||
if plan_id == MembershipType.MONTHLY:
|
||||
features.update({"free_clips_daily": 2})
|
||||
elif plan_id == MembershipType.QUARTERLY:
|
||||
features.update({"free_clips_daily": 5})
|
||||
elif plan_id == MembershipType.YEARLY:
|
||||
features.update({"free_clips_daily": "unlimited"})
|
||||
plans.append({
|
||||
"plan_id": plan_id,
|
||||
"name": info["name"],
|
||||
"price_cents": info["price_cents"],
|
||||
"monthly_price_cents": monthly_cents,
|
||||
"duration_days": days,
|
||||
"points_discount": MEMBER_DISCOUNT.get(plan_id, 1.0),
|
||||
"features": features,
|
||||
})
|
||||
return {"plans": plans}
|
||||
return {"plans": get_membership_plans()}
|
||||
|
||||
|
||||
@router.get("/billing-records", response_model=list[BillingRecord])
|
||||
@@ -244,24 +196,12 @@ async def cancel_subscription(
|
||||
detail="免费用户无需取消订阅",
|
||||
)
|
||||
|
||||
# 已过期/未生效:直接报错,不允许取消
|
||||
expires_at = user.subscription_expires_at
|
||||
member_expires = getattr(user, "member_expires_at", None)
|
||||
if member_expires is not None:
|
||||
expires_at = member_expires
|
||||
if expires_at is not None and expires_at <= datetime.now(UTC):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="订阅已到期,无需取消",
|
||||
)
|
||||
|
||||
updated_user = replace(user, subscription_status="cancelled")
|
||||
user_repository.save(updated_user)
|
||||
|
||||
end_text = expires_at.strftime("%Y-%m-%d") if expires_at else "当前周期结束"
|
||||
return SimpleResponse(
|
||||
success=True,
|
||||
message=f"已取消续费,{end_text} 前仍可正常使用会员权益,到期后自动降级为免费用户",
|
||||
message="订阅已取消,当前周期结束后将降级为免费用户",
|
||||
)
|
||||
|
||||
|
||||
@@ -334,80 +274,3 @@ async def toggle_auto_renew(
|
||||
"""切换自动续费"""
|
||||
status_text = "已开启自动续费" if request.enabled else "已关闭自动续费"
|
||||
return SimpleResponse(success=True, message=status_text)
|
||||
|
||||
|
||||
# ════════════════════════════════════════════════════════════════
|
||||
# 微信支付:下单 / 订单查询
|
||||
# ════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
@router.post("/create-order")
|
||||
def create_membership_order(
|
||||
body: CreateOrderRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
db: Session = Depends(get_db_session),
|
||||
):
|
||||
"""创建会员年卡订单(微信 JSAPI 支付)。
|
||||
|
||||
- 一期只做年卡,billing_cycle 默认 yearly
|
||||
- 返回小程序调起支付所需的 pay_params
|
||||
- 支付结果通过 POST /api/v1/payment/wechat/notify 异步通知履约
|
||||
"""
|
||||
from packages.adapters.wechat_pay import WeChatPayError
|
||||
from packages.application.payment_service import PaymentConfigError, PaymentService
|
||||
|
||||
openid = getattr(current_user.user, "wechat_openid", None)
|
||||
svc = PaymentService()
|
||||
try:
|
||||
result = svc.create_membership_order(
|
||||
user_id=current_user.user.id,
|
||||
plan_id=body.plan_id,
|
||||
billing_cycle=body.billing_cycle or "yearly",
|
||||
openid=openid,
|
||||
db=db,
|
||||
)
|
||||
except PaymentConfigError as e:
|
||||
raise HTTPException(status_code=503, detail=f"支付通道不可用:{e}") from None
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e)) from None
|
||||
except WeChatPayError as e:
|
||||
raise HTTPException(status_code=502, detail=f"微信支付下单失败:{e}") from None
|
||||
|
||||
return CreateOrderResponse(**result)
|
||||
|
||||
|
||||
@router.get("/orders", response_model=OrderListResponse)
|
||||
def list_my_orders(
|
||||
order_type: str | None = Query(None, description="筛选类型: membership/points"),
|
||||
page: int = Query(1, ge=1),
|
||||
page_size: int = Query(20, ge=1, le=100),
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
db: Session = Depends(get_db_session),
|
||||
):
|
||||
"""查询当前用户的订单列表(会员+积分包,分页)。"""
|
||||
from packages.application.payment_service import PaymentService
|
||||
|
||||
return PaymentService().list_orders(
|
||||
user_id=current_user.user.id,
|
||||
order_type=order_type,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
db=db,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/orders/{order_id}", response_model=OrderItem)
|
||||
def get_my_order(
|
||||
order_id: str,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
db: Session = Depends(get_db_session),
|
||||
):
|
||||
"""查询单个订单详情(仅能查自己的订单)。"""
|
||||
from packages.application.payment_service import PaymentService
|
||||
|
||||
order = PaymentService().get_order(
|
||||
user_id=current_user.user.id, order_id=order_id, db=db
|
||||
)
|
||||
if order is None:
|
||||
raise HTTPException(status_code=404, detail="订单不存在")
|
||||
return order
|
||||
|
||||
@@ -49,7 +49,9 @@ from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
SQLAlchemyViralVideoStyleTemplateRepository,
|
||||
)
|
||||
from packages.domain.points_rules import list_viral_video_models
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
from packages.shared.dashscope_client import get_dashscope_client
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -411,7 +413,10 @@ def estimate_credits(
|
||||
|
||||
w, h = resolve_video_dimensions(resolution, ratio)
|
||||
credits, bd = calculate_viral_video_credits_with_breakdown(
|
||||
duration, w, h, model,
|
||||
duration,
|
||||
w,
|
||||
h,
|
||||
model,
|
||||
)
|
||||
breakdown = CreditsFormulaBreakdown(**bd)
|
||||
return EstimateCreditsResponse(estimated_credits=credits, formula_breakdown=breakdown)
|
||||
@@ -451,6 +456,17 @@ def list_style_templates(
|
||||
return StyleTemplateListResponse(items=items)
|
||||
|
||||
|
||||
@router.get("/models")
|
||||
def list_available_models() -> dict:
|
||||
"""返回爆款视频可用模型列表(供前端模型选择器使用)。"""
|
||||
dashscope_available = get_dashscope_client() is not None
|
||||
models = list_viral_video_models(
|
||||
include_placeholder=False,
|
||||
dashscope_available=dashscope_available,
|
||||
)
|
||||
return {"models": models}
|
||||
|
||||
|
||||
@router.get("/{job_id}", response_model=ViralVideoJobResponse)
|
||||
def get_viral_video_job(
|
||||
job_id: str,
|
||||
@@ -574,7 +590,9 @@ def retry_viral_video_job(
|
||||
job.credits_prepaid = round(old_prepaid + diff, 2)
|
||||
logger.info(
|
||||
"[爆款视频][retry] 补扣差额 job_id=%s diff=%.2f new_prepaid=%.2f",
|
||||
job.id, diff, job.credits_prepaid,
|
||||
job.id,
|
||||
diff,
|
||||
job.credits_prepaid,
|
||||
)
|
||||
else:
|
||||
# 新预扣更少:退还差额
|
||||
@@ -591,7 +609,9 @@ def retry_viral_video_job(
|
||||
job.credits_prepaid = round(old_prepaid - refund, 2)
|
||||
logger.info(
|
||||
"[爆款视频][retry] 退还差额 job_id=%s refund=%.2f new_prepaid=%.2f",
|
||||
job.id, refund, job.credits_prepaid,
|
||||
job.id,
|
||||
refund,
|
||||
job.credits_prepaid,
|
||||
)
|
||||
# 差额为 0 则不调整
|
||||
|
||||
@@ -611,7 +631,10 @@ def retry_viral_video_job(
|
||||
celery_app.send_task("worker.run_viral_video_pipeline", args=[job.id])
|
||||
logger.info(
|
||||
"[爆款视频] 重试入队: job_id=%s retry_count=%d stale=%s params_changed=%s",
|
||||
job.id, job.retry_count, is_stale_running, param_changed,
|
||||
job.id,
|
||||
job.retry_count,
|
||||
is_stale_running,
|
||||
param_changed,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 重试入队失败: %s", e, exc_info=True)
|
||||
|
||||
@@ -110,61 +110,3 @@ class ToggleAutoRenewRequest(BaseModel):
|
||||
"""切换自动续费请求"""
|
||||
|
||||
enabled: bool = Field(..., description="是否开启自动续费")
|
||||
|
||||
|
||||
# ============ 微信支付下单/订单 Schemas ============
|
||||
|
||||
|
||||
class CreateOrderRequest(BaseModel):
|
||||
"""会员购买下单请求(一期只做年卡)。"""
|
||||
|
||||
plan_id: str = Field(..., description="套餐ID(admin plans 表UUID);兼容传 yearly")
|
||||
billing_cycle: str = Field("yearly", description="计费周期,一期默认 yearly")
|
||||
|
||||
|
||||
class WeChatPayParams(BaseModel):
|
||||
"""小程序调起微信支付参数。"""
|
||||
|
||||
appId: str
|
||||
timeStamp: str
|
||||
nonceStr: str
|
||||
package: str
|
||||
signType: str
|
||||
paySign: str
|
||||
|
||||
|
||||
class CreateOrderResponse(BaseModel):
|
||||
"""会员下单响应。"""
|
||||
|
||||
order_id: str
|
||||
out_trade_no: str
|
||||
prepay_id: str
|
||||
amount_cents: int
|
||||
pay_params: WeChatPayParams
|
||||
expire_at: str
|
||||
|
||||
|
||||
class OrderItem(BaseModel):
|
||||
"""订单条目。"""
|
||||
|
||||
id: str
|
||||
order_type: str
|
||||
product_code: str
|
||||
product_name: Optional[str] = None
|
||||
amount_cents: int
|
||||
points_amount: int = 0
|
||||
status: str
|
||||
payment_method: Optional[str] = None
|
||||
prepay_id: Optional[str] = None
|
||||
paid_at: Optional[str] = None
|
||||
expire_at: Optional[str] = None
|
||||
created_at: Optional[str] = None
|
||||
|
||||
|
||||
class OrderListResponse(BaseModel):
|
||||
"""订单列表(分页)。"""
|
||||
|
||||
total: int
|
||||
page: int
|
||||
page_size: int
|
||||
items: list[OrderItem]
|
||||
|
||||
@@ -9,6 +9,8 @@ import type {
|
||||
AnalyzeImagesRequest,
|
||||
GenerateCopyRequest,
|
||||
ConfirmCopyRequest,
|
||||
ViralVideoModel,
|
||||
ViralVideoModelsResponse,
|
||||
} from "./types"
|
||||
|
||||
/** 创建爆款视频任务 */
|
||||
@@ -62,6 +64,18 @@ export function estimateViralVideoCredits(params: {
|
||||
.post<{ estimated_credits: number }>("/viral-video/estimate-credits", params)
|
||||
.then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 获取支持的视频模型列表(GET /viral-video/models)。后端返回 {models: [...]} 包装 */
|
||||
export function getViralVideoModels() {
|
||||
return apiClient.get<ViralVideoModelsResponse>("/viral-video/models").then((r) => {
|
||||
const data = r.data as ViralVideoModelsResponse | ViralVideoModel[] | null | undefined
|
||||
if (Array.isArray(data)) return data
|
||||
if (data && Array.isArray((data as ViralVideoModelsResponse).models)) {
|
||||
return (data as ViralVideoModelsResponse).models
|
||||
}
|
||||
return []
|
||||
})
|
||||
}
|
||||
/** ── 三步拆分:前端 mock 辅助函数(后端新接口上线后可替换) ── */
|
||||
|
||||
/**
|
||||
|
||||
@@ -280,10 +280,29 @@ export interface GenerateCopyRequest {
|
||||
video_model?: string
|
||||
}
|
||||
|
||||
/** 视频模型描述(GET /viral-video/models) */
|
||||
export interface ViralVideoModel {
|
||||
key: string
|
||||
display_name: string
|
||||
supports_audio: boolean
|
||||
supported_resolutions: string[]
|
||||
max_duration: number
|
||||
/** 计费模式(可选):per_second / per_video / token 等 */
|
||||
billing_mode?: string
|
||||
is_default?: boolean
|
||||
}
|
||||
|
||||
/** GET /viral-video/models 响应包装 */
|
||||
export interface ViralVideoModelsResponse {
|
||||
models: ViralVideoModel[]
|
||||
}
|
||||
|
||||
/** v1.6 阶段3请求:用户确认/编辑口播文案后开始单次 Seedance 出片(POST /viral-video/{id}/confirm-copy) */
|
||||
export interface ConfirmCopyRequest {
|
||||
/** 用户编辑后的口播文案;为空则使用 AI 生成的 voiceover_script */
|
||||
edited_copy?: string
|
||||
/** 视频模型 key,覆盖默认 */
|
||||
video_model?: string
|
||||
}
|
||||
|
||||
/** 旧分镜片段结构(保留兼容;新代码请使用 ShotScript) */
|
||||
|
||||
@@ -80,6 +80,8 @@ const AiAvatarPage: React.FC = () => {
|
||||
const [finalizeLoading, setFinalizeLoading] = useState(false)
|
||||
|
||||
/* ── 对口型轮询 ── */
|
||||
/** 对口型轮询总时长上限(10分钟):超过后停止轮询并提示去历史记录查看 */
|
||||
const LIPSYNC_POLL_MAX_MS = 10 * 60 * 1000
|
||||
const lipsyncTimerRef = useRef<ReturnType<typeof setInterval> | null>(null)
|
||||
/* ── 渲染进度轮询 ── */
|
||||
const renderTimerRef = useRef<ReturnType<typeof setInterval> | null>(null)
|
||||
@@ -270,7 +272,24 @@ const AiAvatarPage: React.FC = () => {
|
||||
// 如果是预合成模式,后端会同步把状态置为 submitted(甚至可能已返回 running),
|
||||
// 但仍需轮询等 completed
|
||||
if (lipsyncTimerRef.current) clearInterval(lipsyncTimerRef.current)
|
||||
// 轮询间隔 5 秒;单请求超时 5 分钟(见 api/aiAvatar.ts);总轮询上限 10 分钟
|
||||
// 单次请求失败/超时不中断轮询,继续下一轮;超过总上限后停止并提示用户去历史记录查看
|
||||
lipsyncTimerRef.current = setInterval(async () => {
|
||||
// 总时长保护:超过 10 分钟停止轮询
|
||||
if (Date.now() - lipsyncStartAtRef.current > LIPSYNC_POLL_MAX_MS) {
|
||||
if (lipsyncTimerRef.current) {
|
||||
clearInterval(lipsyncTimerRef.current)
|
||||
lipsyncTimerRef.current = null
|
||||
}
|
||||
if (lipsyncTickRef.current) {
|
||||
clearInterval(lipsyncTickRef.current)
|
||||
lipsyncTickRef.current = null
|
||||
}
|
||||
setLipsyncStatus("failed")
|
||||
setLipsyncErrorMessage("渲染时间较长,请稍后在历史记录中查看")
|
||||
message.warning("对口型渲染时间较长,已停止自动刷新,请稍后在历史记录中查看")
|
||||
return
|
||||
}
|
||||
try {
|
||||
const updated = await getLipsyncJob(job.id)
|
||||
state.setLipsyncJob(updated)
|
||||
@@ -296,9 +315,10 @@ const AiAvatarPage: React.FC = () => {
|
||||
setLipsyncErrorMessage(updated.error_message || "对口型生成失败")
|
||||
}
|
||||
} catch (err) {
|
||||
console.error("[对口型] 轮询错误:", err)
|
||||
// 单次轮询失败(含 timeout):不中断轮询,打印日志后等下一轮
|
||||
console.warn("[对口型] 轮询请求失败,将继续下一轮:", err)
|
||||
}
|
||||
}, 3000)
|
||||
}, 5000)
|
||||
} catch (err) {
|
||||
console.error("[对口型] 创建失败:", {
|
||||
status: (err as { response?: { status?: number } })?.response?.status,
|
||||
|
||||
@@ -72,7 +72,8 @@ export const previewTts = async (data: {
|
||||
}
|
||||
|
||||
export const getLipsyncJob = async (id: string): Promise<LipsyncJob> => {
|
||||
const response = await apiClient.get<LipsyncJob>(`/lipsync/jobs/${id}`, { timeout: 60000 })
|
||||
// MuseTalk 渲染 8s 视频约 54s + 排队时间,给足 5 分钟超时避免单次轮询 AxiosError 中断
|
||||
const response = await apiClient.get<LipsyncJob>(`/lipsync/jobs/${id}`, { timeout: 300_000 })
|
||||
return response.data
|
||||
}
|
||||
|
||||
@@ -91,7 +92,10 @@ export const submitRender = async (data: {
|
||||
}
|
||||
|
||||
export const getRenderJob = async (jobId: string): Promise<RenderJob> => {
|
||||
const response = await apiClient.get<RenderJob>(`/ai-avatar/render/${jobId}`, { timeout: 60000 })
|
||||
// 渲染链路(对口型+B-roll+标题+合成+上传)耗时较长,给足 5 分钟超时
|
||||
const response = await apiClient.get<RenderJob>(`/ai-avatar/render/${jobId}`, {
|
||||
timeout: 300_000,
|
||||
})
|
||||
return response.data
|
||||
}
|
||||
|
||||
|
||||
@@ -1002,6 +1002,11 @@
|
||||
.vv-form-row {
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
.vv-form-hint {
|
||||
font-size: 12px;
|
||||
color: #9ca3af;
|
||||
line-height: 1.4;
|
||||
}
|
||||
@media (max-width: 500px) {
|
||||
.vv-form-grid {
|
||||
grid-template-columns: 1fr;
|
||||
@@ -1103,53 +1108,15 @@
|
||||
margin-top: 0;
|
||||
}
|
||||
|
||||
/* 总览 —— 单行段落 */
|
||||
.vv-sb-overview {
|
||||
/* 总览 —— 每行一段 */
|
||||
.vv-sb-inline-row {
|
||||
display: flex;
|
||||
align-items: baseline;
|
||||
flex-wrap: wrap;
|
||||
font-size: 13px;
|
||||
line-height: 1.7;
|
||||
color: #1f2937;
|
||||
margin: 4px 0 6px;
|
||||
padding: 4px 8px;
|
||||
background: #fafafe;
|
||||
border-radius: 6px;
|
||||
}
|
||||
.vv-sb-overview strong {
|
||||
color: #374151;
|
||||
font-weight: 600;
|
||||
}
|
||||
.vv-sb-overview .vv-sb-inline-text,
|
||||
.vv-sb-overview .ant-select {
|
||||
vertical-align: middle;
|
||||
}
|
||||
.vv-sb-inline-text {
|
||||
background: transparent;
|
||||
border: none;
|
||||
color: #1f2937;
|
||||
font: inherit;
|
||||
padding: 1px 4px;
|
||||
outline: none;
|
||||
border-radius: 4px;
|
||||
border-bottom: 1px dashed transparent;
|
||||
transition:
|
||||
border-color 0.15s,
|
||||
background 0.15s;
|
||||
}
|
||||
.vv-sb-inline-text:hover,
|
||||
.vv-sb-inline-text:focus {
|
||||
background: #f5f0ff;
|
||||
border-bottom-color: #7c3aed;
|
||||
}
|
||||
.vv-sb-inline-text:disabled {
|
||||
color: #9ca3af;
|
||||
cursor: default;
|
||||
}
|
||||
.vv-sb-inline-text-sm {
|
||||
width: auto;
|
||||
max-width: 80px;
|
||||
}
|
||||
.vv-sb-sep {
|
||||
color: #d1d5db;
|
||||
margin: 0 8px;
|
||||
margin: 2px 0;
|
||||
}
|
||||
|
||||
.vv-sb-inline-select {
|
||||
@@ -1169,7 +1136,7 @@
|
||||
padding-left: 0 !important;
|
||||
}
|
||||
|
||||
/* 段落式 textarea 基础样式 */
|
||||
/* 段落式 textarea 基础样式(仅编辑态使用) */
|
||||
.vv-sb-doc-ta {
|
||||
background: transparent !important;
|
||||
border: 1px dashed transparent !important;
|
||||
@@ -1186,21 +1153,36 @@
|
||||
border-color: #7c3aed !important;
|
||||
background: #f5f0ff !important;
|
||||
}
|
||||
.vv-sb-doc-ta-sm {
|
||||
min-height: 24px;
|
||||
}
|
||||
|
||||
/* 场景与光线 —— 段落样式 */
|
||||
.vv-sb-para {
|
||||
margin: 4px 0;
|
||||
}
|
||||
.vv-sb-doc-ta-block {
|
||||
|
||||
/* 内联编辑 textarea(点击后弹出) */
|
||||
.vv-sb-inline-edit-ta {
|
||||
display: block;
|
||||
width: 100%;
|
||||
padding: 4px 8px !important;
|
||||
margin-top: 4px;
|
||||
min-height: 28px;
|
||||
background: #fafafe !important;
|
||||
border: 1px solid #d8cafc !important;
|
||||
border-radius: 6px !important;
|
||||
padding: 6px 8px !important;
|
||||
font-size: 13px !important;
|
||||
line-height: 1.6 !important;
|
||||
min-height: 32px;
|
||||
font-family: inherit;
|
||||
color: #1f2937 !important;
|
||||
outline: none;
|
||||
resize: vertical;
|
||||
}
|
||||
.vv-sb-inline-edit-ta-sm {
|
||||
max-width: 120px;
|
||||
}
|
||||
.vv-sb-time-ta {
|
||||
max-width: 160px;
|
||||
font-weight: 600;
|
||||
color: #7c3aed !important;
|
||||
}
|
||||
|
||||
/* 逐镜头 */
|
||||
@@ -1211,49 +1193,51 @@
|
||||
margin-top: 2px;
|
||||
}
|
||||
.vv-sb-doc-shot {
|
||||
padding-left: 8px;
|
||||
border-left: 2px solid rgba(124, 58, 237, 0.3);
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 2px;
|
||||
margin-bottom: 6px;
|
||||
}
|
||||
.vv-sb-time-doc {
|
||||
display: inline-block;
|
||||
font-size: 13px;
|
||||
font-weight: 700;
|
||||
display: block;
|
||||
font-size: 14px;
|
||||
font-weight: 600;
|
||||
color: #7c3aed;
|
||||
background: rgba(124, 58, 237, 0.08);
|
||||
border: none;
|
||||
outline: none;
|
||||
padding: 1px 6px;
|
||||
font-family: inherit;
|
||||
border-radius: 4px;
|
||||
margin-bottom: 4px;
|
||||
margin: 6px 0 2px;
|
||||
cursor: text;
|
||||
}
|
||||
.vv-sb-time-doc:focus {
|
||||
background: #f5f0ff;
|
||||
.vv-sb-time-doc:hover {
|
||||
background: rgba(124, 58, 237, 0.06);
|
||||
border-radius: 3px;
|
||||
}
|
||||
|
||||
/* 字段段落 */
|
||||
.vv-sb-field {
|
||||
display: block;
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
align-items: baseline;
|
||||
margin: 2px 0;
|
||||
font-size: 13px;
|
||||
line-height: 1.65;
|
||||
line-height: 1.6;
|
||||
}
|
||||
.vv-sb-field-k {
|
||||
color: #6d28d9;
|
||||
color: #1f2937;
|
||||
font-weight: 600;
|
||||
margin-right: 4px;
|
||||
margin-right: 0;
|
||||
white-space: nowrap;
|
||||
}
|
||||
.vv-sb-doc-ta-inline {
|
||||
display: block;
|
||||
width: 100%;
|
||||
margin-top: 1px;
|
||||
padding: 2px 4px !important;
|
||||
font-size: 13px !important;
|
||||
line-height: 1.6 !important;
|
||||
min-height: 24px;
|
||||
.vv-sb-field-val {
|
||||
color: #374151;
|
||||
cursor: text;
|
||||
border-radius: 3px;
|
||||
padding: 0 2px;
|
||||
transition: background 0.15s;
|
||||
word-break: break-word;
|
||||
flex: 1;
|
||||
min-width: 0;
|
||||
}
|
||||
.vv-sb-field:hover .vv-sb-field-val {
|
||||
background: rgba(124, 58, 237, 0.06);
|
||||
}
|
||||
|
||||
/* 参考图片行 */
|
||||
@@ -1396,24 +1380,7 @@
|
||||
background: #f5f0ff;
|
||||
}
|
||||
|
||||
/* 口播稿 */
|
||||
.vv-sb-vo-wrap {
|
||||
margin-top: 4px;
|
||||
}
|
||||
.vv-sb-vo {
|
||||
display: block;
|
||||
width: 100%;
|
||||
background: #fafafe !important;
|
||||
border-left: 3px solid #7c3aed !important;
|
||||
border-radius: 0 6px 6px 0 !important;
|
||||
padding: 6px 10px !important;
|
||||
min-height: 44px;
|
||||
font-size: 13px !important;
|
||||
line-height: 1.7 !important;
|
||||
color: #1f2937 !important;
|
||||
border: none !important;
|
||||
resize: vertical;
|
||||
}
|
||||
/* 口播稿 —— 复用 vv-sb-field 样式,无额外需求 */
|
||||
|
||||
.vv-sb-actions {
|
||||
display: flex;
|
||||
@@ -1502,7 +1469,7 @@
|
||||
}
|
||||
|
||||
/* ── Asset/voice picker modal styles (in page) ───────────── */
|
||||
.vv-modal {
|
||||
.vv-modal-mask {
|
||||
position: fixed;
|
||||
inset: 0;
|
||||
background: rgba(0, 0, 0, 0.45);
|
||||
@@ -1512,21 +1479,59 @@
|
||||
justify-content: center;
|
||||
padding: 20px;
|
||||
}
|
||||
.vv-modal-body {
|
||||
.vv-modal {
|
||||
background: #fff;
|
||||
border: 1px solid #e5e7eb;
|
||||
border-radius: 12px;
|
||||
max-width: 720px;
|
||||
width: 100%;
|
||||
max-height: 80vh;
|
||||
overflow-y: auto;
|
||||
padding: 20px;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
overflow: hidden;
|
||||
position: relative;
|
||||
box-shadow: 0 12px 40px rgba(15, 23, 42, 0.18);
|
||||
}
|
||||
.vv-modal.vv-modal-lg {
|
||||
max-width: 860px;
|
||||
}
|
||||
.vv-modal-head {
|
||||
flex-shrink: 0;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
padding: 14px 20px;
|
||||
border-bottom: 1px solid #e5e7eb;
|
||||
}
|
||||
.vv-modal-title {
|
||||
font-size: 15px;
|
||||
font-weight: 600;
|
||||
color: #1f2937;
|
||||
}
|
||||
.vv-modal-body {
|
||||
flex: 1 1 auto;
|
||||
overflow-y: auto;
|
||||
padding: 16px 20px;
|
||||
min-height: 0;
|
||||
}
|
||||
.vv-modal-foot {
|
||||
flex-shrink: 0;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: flex-end;
|
||||
gap: 10px;
|
||||
padding: 12px 20px;
|
||||
border-top: 1px solid #e5e7eb;
|
||||
background: #fff;
|
||||
}
|
||||
.vv-modal-foot .vv-btn-primary {
|
||||
width: auto;
|
||||
padding: 8px 18px;
|
||||
}
|
||||
.vv-modal-foot .vv-btn-ghost {
|
||||
padding: 8px 18px;
|
||||
}
|
||||
.vv-modal-close {
|
||||
position: absolute;
|
||||
top: 14px;
|
||||
right: 14px;
|
||||
background: transparent;
|
||||
border: none;
|
||||
color: #6b7280;
|
||||
@@ -1535,6 +1540,11 @@
|
||||
width: 28px;
|
||||
height: 28px;
|
||||
border-radius: 6px;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
padding: 0;
|
||||
line-height: 1;
|
||||
}
|
||||
.vv-modal-close:hover {
|
||||
color: #ef4444;
|
||||
|
||||
@@ -40,6 +40,7 @@ import {
|
||||
type ImageAnalysisResult,
|
||||
type CopyResult,
|
||||
type ShotScript,
|
||||
type ViralVideoModel,
|
||||
} from "@/api/viral-video/types"
|
||||
import {
|
||||
generateViralVideo,
|
||||
@@ -48,6 +49,7 @@ import {
|
||||
generateViralCopy,
|
||||
confirmViralCopy,
|
||||
estimateViralVideoCredits,
|
||||
getViralVideoModels,
|
||||
} from "@/api/viral-video"
|
||||
import { useViralVideoPolling } from "./hooks/useViralVideoPolling"
|
||||
import CloneModal from "@/components/voice/CloneModal"
|
||||
@@ -217,10 +219,46 @@ const RATIOS = [
|
||||
{ v: "16:9", label: "16:9 横屏(B站/YouTube)" },
|
||||
{ v: "1:1", label: "1:1 方形(小红书)" },
|
||||
]
|
||||
const MODELS = [
|
||||
{ v: "seedance-2.5", label: "Seedance 2.5(推荐)" },
|
||||
{ v: "seedance-2.0", label: "Seedance 2.0" },
|
||||
/** 兜底模型列表(接口未返回时使用,字段与 ViralVideoModel 对齐;后端返回后自动覆盖) */
|
||||
const FALLBACK_VIDEO_MODELS: ViralVideoModel[] = [
|
||||
{
|
||||
key: "seedance-2.5",
|
||||
display_name: "Seedance 2.5 — 最新最强",
|
||||
supports_audio: true,
|
||||
supported_resolutions: ["480p", "720p", "1080p"],
|
||||
max_duration: 30,
|
||||
is_default: true,
|
||||
},
|
||||
{
|
||||
key: "seedance-2.0",
|
||||
display_name: "Seedance 2.0 — 正式首选",
|
||||
supports_audio: true,
|
||||
supported_resolutions: ["480p", "720p", "1080p"],
|
||||
max_duration: 15,
|
||||
},
|
||||
{
|
||||
key: "seedance-2.0-fast",
|
||||
display_name: "Seedance 2.0 Fast — 快速测试",
|
||||
supports_audio: true,
|
||||
supported_resolutions: ["480p", "720p"],
|
||||
max_duration: 15,
|
||||
},
|
||||
{
|
||||
key: "seedance-2.0-mini",
|
||||
display_name: "Seedance 2.0 Mini — 低成本",
|
||||
supports_audio: true,
|
||||
supported_resolutions: ["480p", "720p"],
|
||||
max_duration: 15,
|
||||
},
|
||||
{
|
||||
key: "wan-3.0",
|
||||
display_name: "Wan 3.0 — 通义万相全能",
|
||||
supports_audio: true,
|
||||
supported_resolutions: ["480p", "720p", "1080p"],
|
||||
max_duration: 15,
|
||||
},
|
||||
]
|
||||
const RESOLUTION_ORDER = ["480p", "720p", "1080p", "4k"]
|
||||
const QUALITY_OPTIONS = [
|
||||
{ v: "480p", label: "480p(快速)" },
|
||||
{ v: "720p", label: "720p(清晰)" },
|
||||
@@ -448,6 +486,7 @@ const ViralVideoPage: React.FC = () => {
|
||||
>([])
|
||||
const [presetVoices, setPresetVoices] = useState<PresetVoice[]>([])
|
||||
const [voicePickerOpen, setVoicePickerOpen] = useState(false)
|
||||
const [videoModels, setVideoModels] = useState<ViralVideoModel[]>(FALLBACK_VIDEO_MODELS)
|
||||
const [assetPicker, setAssetPicker] = useState<{
|
||||
open: boolean
|
||||
kind: "image" | "video" | "voice"
|
||||
@@ -507,6 +546,43 @@ const ViralVideoPage: React.FC = () => {
|
||||
})
|
||||
}, [])
|
||||
|
||||
/* ── 加载视频模型列表 ── */
|
||||
useEffect(() => {
|
||||
getViralVideoModels()
|
||||
.then((list) => {
|
||||
// 防御:确保 list 是数组
|
||||
const safeList = Array.isArray(list) ? list : []
|
||||
if (safeList.length === 0) return
|
||||
setVideoModels(safeList)
|
||||
// 如果当前选中的模型不在返回列表里,切换到默认模型并联动修正参数
|
||||
setTask((t) => {
|
||||
const inList = safeList.some((m) => m && typeof m === "object" && m.key === t.videoModel)
|
||||
if (inList) return t
|
||||
const def =
|
||||
safeList.find((m) => m && typeof m === "object" && m.is_default) || safeList[0]
|
||||
if (!def) return t
|
||||
const defRes = Array.isArray(def.supported_resolutions)
|
||||
? def.supported_resolutions
|
||||
: ["480p", "720p", "1080p"]
|
||||
const sortedRes = defRes.slice().sort((a, b) => {
|
||||
const ra = RESOLUTION_ORDER.indexOf(a)
|
||||
const rb = RESOLUTION_ORDER.indexOf(b)
|
||||
return (rb === -1 ? -1 : rb) - (ra === -1 ? -1 : ra)
|
||||
})
|
||||
const newRes = sortedRes[0] || "480p"
|
||||
const newDur = Math.min(
|
||||
t.duration,
|
||||
typeof def.max_duration === "number" ? def.max_duration : 15,
|
||||
)
|
||||
return { ...t, videoModel: def.key, quality: newRes, duration: newDur }
|
||||
})
|
||||
})
|
||||
.catch(() => {
|
||||
// 接口 404/500 时使用兜底列表,不提示用户
|
||||
})
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [])
|
||||
|
||||
/* ── 轮询(job 从 STEP1 开始即存在,贯穿三步) ── */
|
||||
const onPollUpdate = useCallback(
|
||||
(job: ViralVideoJob) => {
|
||||
@@ -969,6 +1045,7 @@ const ViralVideoPage: React.FC = () => {
|
||||
""
|
||||
job = await confirmViralCopy(task.jobId, {
|
||||
edited_copy: edited && edited !== originalCopy.trim() ? edited : undefined,
|
||||
video_model: task.videoModel,
|
||||
})
|
||||
} else {
|
||||
// 兜底:走旧 /generate 接口(一次性跑完)
|
||||
@@ -1163,40 +1240,85 @@ const ViralVideoPage: React.FC = () => {
|
||||
}
|
||||
|
||||
/* ── 分镜脚本结果区(编导分镜卡片 UI) ── */
|
||||
const updateShot = (idx: number, patch: Partial<StoryboardShot>) => {
|
||||
if (!task.storyboard) return
|
||||
const shots = task.storyboard.shots.map((s, i) => (i === idx ? { ...s, ...patch } : s))
|
||||
setTask({ storyboard: { ...task.storyboard, shots } })
|
||||
}
|
||||
const updateStoryboard = (patch: Partial<Storyboard>) => {
|
||||
if (!task.storyboard) return
|
||||
setTask({ storyboard: { ...task.storyboard, ...patch } })
|
||||
}
|
||||
const updateOverview = (patch: Partial<Storyboard["overview"]>) => {
|
||||
if (!task.storyboard) return
|
||||
setTask({
|
||||
storyboard: { ...task.storyboard, overview: { ...task.storyboard.overview, ...patch } },
|
||||
const [editingField, setEditingField] = useState<string | null>(null)
|
||||
const editTaRef = useRef<HTMLTextAreaElement | null>(null)
|
||||
|
||||
const updateShot = useCallback((idx: number, patch: Partial<StoryboardShot>) => {
|
||||
setTask((t) => {
|
||||
if (!t.storyboard) return t
|
||||
const shots = t.storyboard.shots.map((s, i) => (i === idx ? { ...s, ...patch } : s))
|
||||
return { ...t, storyboard: { ...t.storyboard, shots } }
|
||||
})
|
||||
}
|
||||
const updateConstraints = (
|
||||
list: "hard_constraints" | "negative_prompts",
|
||||
idx: number,
|
||||
v: string,
|
||||
) => {
|
||||
if (!task.storyboard) return
|
||||
const arr = [...task.storyboard[list]]
|
||||
arr[idx] = v
|
||||
setTask({ storyboard: { ...task.storyboard, [list]: arr } })
|
||||
}
|
||||
const addConstraint = (list: "hard_constraints" | "negative_prompts") => {
|
||||
if (!task.storyboard) return
|
||||
setTask({ storyboard: { ...task.storyboard, [list]: [...task.storyboard[list], ""] } })
|
||||
}
|
||||
const removeConstraint = (list: "hard_constraints" | "negative_prompts", idx: number) => {
|
||||
if (!task.storyboard) return
|
||||
const arr = task.storyboard[list].filter((_, i) => i !== idx)
|
||||
setTask({ storyboard: { ...task.storyboard, [list]: arr } })
|
||||
}
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [])
|
||||
const updateStoryboard = useCallback((patch: Partial<Storyboard>) => {
|
||||
setTask((t) => (t.storyboard ? { ...t, storyboard: { ...t.storyboard, ...patch } } : t))
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [])
|
||||
const updateOverview = useCallback((patch: Partial<Storyboard["overview"]>) => {
|
||||
setTask((t) =>
|
||||
t.storyboard
|
||||
? {
|
||||
...t,
|
||||
storyboard: { ...t.storyboard, overview: { ...t.storyboard.overview, ...patch } },
|
||||
}
|
||||
: t,
|
||||
)
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [])
|
||||
const updateConstraints = useCallback(
|
||||
(list: "hard_constraints" | "negative_prompts", idx: number, v: string) => {
|
||||
setTask((t) => {
|
||||
if (!t.storyboard) return t
|
||||
const arr = [...t.storyboard[list]]
|
||||
arr[idx] = v
|
||||
return { ...t, storyboard: { ...t.storyboard, [list]: arr } }
|
||||
})
|
||||
},
|
||||
[], // eslint-disable-line react-hooks/exhaustive-deps
|
||||
)
|
||||
const addConstraint = useCallback((list: "hard_constraints" | "negative_prompts") => {
|
||||
setTask((t) =>
|
||||
t.storyboard
|
||||
? { ...t, storyboard: { ...t.storyboard, [list]: [...t.storyboard[list], ""] } }
|
||||
: t,
|
||||
)
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [])
|
||||
const removeConstraint = useCallback(
|
||||
(list: "hard_constraints" | "negative_prompts", idx: number) => {
|
||||
setTask((t) => {
|
||||
if (!t.storyboard) return t
|
||||
const arr = t.storyboard[list].filter((_, i) => i !== idx)
|
||||
return { ...t, storyboard: { ...t.storyboard, [list]: arr } }
|
||||
})
|
||||
},
|
||||
[], // eslint-disable-line react-hooks/exhaustive-deps
|
||||
)
|
||||
|
||||
/* 内联编辑 helper —— 点击文本 → textarea,blur/Enter 保存 */
|
||||
const handleInlineCommit = useCallback(() => {
|
||||
const ta = editTaRef.current
|
||||
if (!ta) return
|
||||
const key = ta.dataset.field
|
||||
if (key) {
|
||||
const val = ta.value
|
||||
if (key.startsWith("shot.")) {
|
||||
const parts = key.split(".")
|
||||
const idx = Number(parts[1])
|
||||
const field = parts.slice(2).join(".")
|
||||
updateShot(idx, { [field]: val } as Partial<StoryboardShot>)
|
||||
} else if (key.startsWith("overview.")) {
|
||||
const field = key.split(".")[1]
|
||||
updateOverview({ [field]: val } as Partial<Storyboard["overview"]>)
|
||||
} else if (key === "scene_and_lighting") {
|
||||
updateStoryboard({ scene_and_lighting: val })
|
||||
} else if (key === "voiceover_script") {
|
||||
updateStoryboard({ voiceover_script: val })
|
||||
}
|
||||
}
|
||||
setEditingField(null)
|
||||
}, [updateShot, updateOverview, updateStoryboard])
|
||||
|
||||
const renderCopyResult = () => {
|
||||
if (task.uiStep === "step2_generating") {
|
||||
@@ -1253,27 +1375,69 @@ const ViralVideoPage: React.FC = () => {
|
||||
<div className="vv-sb-doc">
|
||||
{/* 视频总览 */}
|
||||
<h4 className="vv-sb-h">视频总览</h4>
|
||||
<p className="vv-sb-overview">
|
||||
<strong>整体主题:</strong>
|
||||
<input
|
||||
className="vv-sb-inline-text"
|
||||
value={sb.overview.theme}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateOverview({ theme: e.target.value })}
|
||||
/>
|
||||
<span className="vv-sb-sep">·</span>
|
||||
<strong>总时长:</strong>
|
||||
<input
|
||||
className="vv-sb-inline-text vv-sb-inline-text-sm"
|
||||
value={sb.overview.total_duration}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateOverview({ total_duration: e.target.value })}
|
||||
/>
|
||||
<span className="vv-sb-sep">·</span>
|
||||
<strong>画幅:</strong>
|
||||
<p className="vv-sb-inline-row">
|
||||
<strong className="vv-sb-field-k">整体主题:</strong>
|
||||
{editingField === "overview.theme" ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta"
|
||||
autoFocus
|
||||
defaultValue={sb.overview.theme}
|
||||
data-field="overview.theme"
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField("overview.theme")
|
||||
}}
|
||||
>
|
||||
{sb.overview.theme}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
<p className="vv-sb-inline-row">
|
||||
<strong className="vv-sb-field-k">总时长:</strong>
|
||||
{editingField === "overview.total_duration" ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta vv-sb-inline-edit-ta-sm"
|
||||
autoFocus
|
||||
defaultValue={sb.overview.total_duration}
|
||||
data-field="overview.total_duration"
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField("overview.total_duration")
|
||||
}}
|
||||
>
|
||||
{sb.overview.total_duration}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
<p className="vv-sb-inline-row">
|
||||
<strong className="vv-sb-field-k">画幅:</strong>
|
||||
<Select
|
||||
className="vv-select vv-sb-inline-select"
|
||||
style={{ width: 120 }}
|
||||
style={{ width: 90 }}
|
||||
value={sb.overview.aspect_ratio}
|
||||
disabled={locked}
|
||||
onChange={(v) => updateOverview({ aspect_ratio: v })}
|
||||
@@ -1285,13 +1449,33 @@ const ViralVideoPage: React.FC = () => {
|
||||
{/* 场景与光线 */}
|
||||
<h4 className="vv-sb-h">场景与光线</h4>
|
||||
<p className="vv-sb-para">
|
||||
<textarea
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-doc-ta-block"
|
||||
value={sb.scene_and_lighting}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateStoryboard({ scene_and_lighting: e.target.value })}
|
||||
placeholder="描述整体场景氛围、光线方向与色温…"
|
||||
/>
|
||||
{editingField === "scene_and_lighting" ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta"
|
||||
autoFocus
|
||||
defaultValue={sb.scene_and_lighting}
|
||||
data-field="scene_and_lighting"
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
placeholder="描述整体场景氛围、光线方向与色温…"
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField("scene_and_lighting")
|
||||
}}
|
||||
>
|
||||
{sb.scene_and_lighting}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
|
||||
{/* 逐镜头 */}
|
||||
@@ -1302,61 +1486,180 @@ const ViralVideoPage: React.FC = () => {
|
||||
typeof sh.reference_image_index === "number"
|
||||
? task.images[sh.reference_image_index]
|
||||
: undefined
|
||||
const shotKey = (field: string) => `shot.${idx}.${field}`
|
||||
return (
|
||||
<div key={idx} className="vv-sb-doc-shot">
|
||||
<input
|
||||
className="vv-sb-time-doc"
|
||||
value={sh.time_range}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateShot(idx, { time_range: e.target.value })}
|
||||
placeholder="0-3秒"
|
||||
/>
|
||||
{editingField === shotKey("time_range") ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta vv-sb-time-ta"
|
||||
autoFocus
|
||||
defaultValue={sh.time_range}
|
||||
data-field={shotKey("time_range")}
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
placeholder="0-3秒"
|
||||
/>
|
||||
) : (
|
||||
<strong
|
||||
className="vv-sb-time-doc"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField(shotKey("time_range"))
|
||||
}}
|
||||
>
|
||||
{sh.time_range}
|
||||
</strong>
|
||||
)}
|
||||
<p className="vv-sb-field">
|
||||
<strong className="vv-sb-field-k">景别/角度与运镜:</strong>
|
||||
<textarea
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-doc-ta-sm vv-sb-doc-ta-inline"
|
||||
value={sh.shot_type_angle_movement}
|
||||
disabled={locked}
|
||||
onChange={(e) =>
|
||||
updateShot(idx, { shot_type_angle_movement: e.target.value })
|
||||
}
|
||||
/>
|
||||
{editingField === shotKey("shot_type_angle_movement") ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta"
|
||||
autoFocus
|
||||
defaultValue={sh.shot_type_angle_movement}
|
||||
data-field={shotKey("shot_type_angle_movement")}
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField(shotKey("shot_type_angle_movement"))
|
||||
}}
|
||||
>
|
||||
{sh.shot_type_angle_movement}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
<p className="vv-sb-field">
|
||||
<strong className="vv-sb-field-k">场景与对白:</strong>
|
||||
<textarea
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-doc-ta-inline"
|
||||
value={sh.scene_and_dialogue}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateShot(idx, { scene_and_dialogue: e.target.value })}
|
||||
/>
|
||||
{editingField === shotKey("scene_and_dialogue") ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta"
|
||||
autoFocus
|
||||
defaultValue={sh.scene_and_dialogue}
|
||||
data-field={shotKey("scene_and_dialogue")}
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField(shotKey("scene_and_dialogue"))
|
||||
}}
|
||||
>
|
||||
{sh.scene_and_dialogue}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
<p className="vv-sb-field">
|
||||
<strong className="vv-sb-field-k">动作与真人细节:</strong>
|
||||
<textarea
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-doc-ta-sm vv-sb-doc-ta-inline"
|
||||
value={sh.action_details}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateShot(idx, { action_details: e.target.value })}
|
||||
/>
|
||||
{editingField === shotKey("action_details") ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta"
|
||||
autoFocus
|
||||
defaultValue={sh.action_details}
|
||||
data-field={shotKey("action_details")}
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField(shotKey("action_details"))
|
||||
}}
|
||||
>
|
||||
{sh.action_details}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
<p className="vv-sb-field">
|
||||
<strong className="vv-sb-field-k">音效/BGM:</strong>
|
||||
<textarea
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-doc-ta-sm vv-sb-doc-ta-inline"
|
||||
value={sh.audio_bgm}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateShot(idx, { audio_bgm: e.target.value })}
|
||||
/>
|
||||
{editingField === shotKey("audio_bgm") ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta"
|
||||
autoFocus
|
||||
defaultValue={sh.audio_bgm}
|
||||
data-field={shotKey("audio_bgm")}
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField(shotKey("audio_bgm"))
|
||||
}}
|
||||
>
|
||||
{sh.audio_bgm}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
<p className="vv-sb-field">
|
||||
<strong className="vv-sb-field-k">转场:</strong>
|
||||
<textarea
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-doc-ta-sm vv-sb-doc-ta-inline"
|
||||
value={sh.transition}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateShot(idx, { transition: e.target.value })}
|
||||
/>
|
||||
{editingField === shotKey("transition") ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta"
|
||||
autoFocus
|
||||
defaultValue={sh.transition}
|
||||
data-field={shotKey("transition")}
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField(shotKey("transition"))
|
||||
}}
|
||||
>
|
||||
{sh.transition}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
<p className="vv-sb-field vv-sb-ref-row">
|
||||
<strong className="vv-sb-field-k">参考图片:</strong>
|
||||
@@ -1458,19 +1761,39 @@ const ViralVideoPage: React.FC = () => {
|
||||
</div>
|
||||
|
||||
{/* 完整口播稿 */}
|
||||
<h4 className="vv-sb-h">
|
||||
<SoundOutlined style={{ color: "#7c3aed", marginRight: 6 }} />
|
||||
完整口播稿(TTS 合成使用)
|
||||
</h4>
|
||||
<div className="vv-sb-vo-wrap">
|
||||
<textarea
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-vo"
|
||||
value={sb.voiceover_script}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateStoryboard({ voiceover_script: e.target.value })}
|
||||
placeholder="AI 合成配音用的完整口播稿…"
|
||||
/>
|
||||
</div>
|
||||
<p className="vv-sb-field">
|
||||
<strong className="vv-sb-field-k">
|
||||
<SoundOutlined style={{ color: "#7c3aed", marginRight: 4 }} />
|
||||
口播稿:
|
||||
</strong>
|
||||
{editingField === "voiceover_script" ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta"
|
||||
autoFocus
|
||||
defaultValue={sb.voiceover_script}
|
||||
data-field="voiceover_script"
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
placeholder="AI 合成配音用的完整口播稿…"
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField("voiceover_script")
|
||||
}}
|
||||
>
|
||||
{sb.voiceover_script}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div className="vv-sb-actions">
|
||||
@@ -2065,8 +2388,36 @@ const ViralVideoPage: React.FC = () => {
|
||||
className="vv-select vv-select-step3"
|
||||
style={{ width: "100%" }}
|
||||
value={task.videoModel}
|
||||
onChange={(v) => setTask({ videoModel: v })}
|
||||
options={MODELS.map((m) => ({ value: m.v, label: m.label }))}
|
||||
onChange={(v) => {
|
||||
const m = videoModels.find((x) => x && x.key === v)
|
||||
if (!m) {
|
||||
setTask({ videoModel: v })
|
||||
return
|
||||
}
|
||||
// 切换模型时自动联动修正分辨率/时长(防御:确保 supported_resolutions 为数组)
|
||||
const mRes = Array.isArray(m.supported_resolutions)
|
||||
? m.supported_resolutions
|
||||
: ["480p", "720p", "1080p"]
|
||||
const sortedRes = mRes.slice().sort((a, b) => {
|
||||
const ra = RESOLUTION_ORDER.indexOf(a)
|
||||
const rb = RESOLUTION_ORDER.indexOf(b)
|
||||
return (rb === -1 ? -1 : rb) - (ra === -1 ? -1 : ra)
|
||||
})
|
||||
const newRes = mRes.includes(task.quality)
|
||||
? task.quality
|
||||
: sortedRes[0] || "480p"
|
||||
const newDur = Math.min(
|
||||
task.duration,
|
||||
typeof m.max_duration === "number" ? m.max_duration : 15,
|
||||
)
|
||||
setTask({ videoModel: v, quality: newRes, duration: newDur })
|
||||
}}
|
||||
options={videoModels
|
||||
.filter((m) => m && typeof m === "object" && m.key)
|
||||
.map((m) => ({
|
||||
value: m.key,
|
||||
label: m.display_name || m.key,
|
||||
}))}
|
||||
/>
|
||||
</div>
|
||||
<div className="vv-form-row">
|
||||
@@ -2076,7 +2427,20 @@ const ViralVideoPage: React.FC = () => {
|
||||
style={{ width: "100%" }}
|
||||
value={task.quality}
|
||||
onChange={(v) => setTask({ quality: v })}
|
||||
options={QUALITY_OPTIONS.map((m) => ({ value: m.v, label: m.label }))}
|
||||
options={(() => {
|
||||
const curM = videoModels.find((m) => m && m.key === task.videoModel)
|
||||
const supported = Array.isArray(curM?.supported_resolutions)
|
||||
? curM!.supported_resolutions
|
||||
: null
|
||||
const pool =
|
||||
supported && supported.length > 0
|
||||
? supported
|
||||
: QUALITY_OPTIONS.map((o) => o.v)
|
||||
return QUALITY_OPTIONS.filter((o) => pool.includes(o.v)).map((o) => ({
|
||||
value: o.v,
|
||||
label: o.label,
|
||||
}))
|
||||
})()}
|
||||
/>
|
||||
</div>
|
||||
<div className="vv-form-row">
|
||||
|
||||
@@ -245,12 +245,11 @@ export default function AssetPickerModal({
|
||||
</div>
|
||||
{multiple && (
|
||||
<div className="vv-modal-foot">
|
||||
<button className="vv-btn vv-btn-ghost vv-btn-sm" onClick={onClose}>
|
||||
<button className="vv-btn vv-btn-ghost" onClick={onClose}>
|
||||
取消
|
||||
</button>
|
||||
<button
|
||||
className="vv-btn vv-btn-primary"
|
||||
style={{ width: "auto", marginTop: 0, padding: "8px 18px" }}
|
||||
onClick={handleConfirm}
|
||||
disabled={picked.size === 0}
|
||||
>
|
||||
|
||||
@@ -27,6 +27,7 @@ import threading
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from celery import Task, shared_task
|
||||
from celery.exceptions import Retry
|
||||
@@ -646,7 +647,7 @@ def _build_products_summary(image_analysis: dict) -> str:
|
||||
# 优先 VLM 生成的 summary 段(自然语言,给编导模型看效果最好)
|
||||
summary = (p.get("summary") or "").strip()
|
||||
if summary and len(summary) >= 30:
|
||||
lines.append(f"- 图{i+1} {name}:{summary}")
|
||||
lines.append(f"- 图{i + 1} {name}:{summary}")
|
||||
continue
|
||||
# 结构化字段兜底
|
||||
brand = p.get("brand") or ""
|
||||
@@ -668,7 +669,7 @@ def _build_products_summary(image_analysis: dict) -> str:
|
||||
feats = p.get("key_features") or p.get("features") or []
|
||||
sellings = p.get("selling_points") or []
|
||||
scenes = p.get("suitable_scenes") or []
|
||||
parts = [f"图{i+1} {name}"]
|
||||
parts = [f"图{i + 1} {name}"]
|
||||
if brand and brand not in ("未知", "无法判断"):
|
||||
parts.append(f"品牌={brand}")
|
||||
if cat and cat not in ("无法判断", "非产品图"):
|
||||
@@ -728,6 +729,22 @@ def _safe_json_loads(raw: str | dict | list | None):
|
||||
return None
|
||||
|
||||
|
||||
def _replace_henjin_everywhere(obj: Any) -> Any:
|
||||
"""递归遍历 copy_result 里所有字符串值,把'很近'替换成'最近'。
|
||||
覆盖 overview.theme、scene_and_lighting、voiceover_script、
|
||||
shots[].scene_and_dialogue/action_details/audio_bgm 等所有字段。
|
||||
"""
|
||||
if isinstance(obj, str):
|
||||
if "很近" in obj:
|
||||
return obj.replace("很近", "最近")
|
||||
return obj
|
||||
if isinstance(obj, list):
|
||||
return [_replace_henjin_everywhere(x) for x in obj]
|
||||
if isinstance(obj, dict):
|
||||
return {k: _replace_henjin_everywhere(v) for k, v in obj.items()}
|
||||
return obj
|
||||
|
||||
|
||||
def _fallback_script(job: ViralVideoJob) -> dict:
|
||||
"""脚本生成失败时的兜底脚本(极简但可用)。"""
|
||||
dur = max(5, min(30, int(getattr(job, "duration", 15) or 15)))
|
||||
@@ -785,7 +802,7 @@ def _validate_and_normalize_script(raw, job: ViralVideoJob) -> dict:
|
||||
continue
|
||||
shots.append(
|
||||
{
|
||||
"time_range": str(s.get("time_range") or f"{i*3}-{(i+1)*3}秒"),
|
||||
"time_range": str(s.get("time_range") or f"{i * 3}-{(i + 1) * 3}秒"),
|
||||
"shot_type_angle_movement": str(s.get("shot_type_angle_movement") or "中景平视,固定镜头"),
|
||||
"scene_and_dialogue": str(s.get("scene_and_dialogue") or ""),
|
||||
"action_details": str(s.get("action_details") or ""),
|
||||
@@ -861,8 +878,8 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
|
||||
style_hint = "无"
|
||||
if isinstance(job.style_guide, dict):
|
||||
style_hint = (
|
||||
f"节奏{job.style_guide.get('cut_speed','')}、转场{job.style_guide.get('transition','')}、"
|
||||
f"色调{job.style_guide.get('color_grade','')}、能量{job.style_guide.get('energy','')}"
|
||||
f"节奏{job.style_guide.get('cut_speed', '')}、转场{job.style_guide.get('transition', '')}、"
|
||||
f"色调{job.style_guide.get('color_grade', '')}、能量{job.style_guide.get('energy', '')}"
|
||||
)
|
||||
|
||||
dur = max(5, min(30, int(getattr(job, "duration", 15) or 15)))
|
||||
@@ -920,19 +937,14 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
|
||||
voiceover_len = len(voiceover)
|
||||
shots_cnt = len((normalized or {}).get("shots") or [])
|
||||
# 判定是否"退化到兜底质量":口播过短(<20字)或镜头数<1;正常的短口播(如15s视频~40字)不视为兜底
|
||||
# v1.6.1 双保险:先做 hard fix 字符串替换("很近" → "最近"),再做不合格判定
|
||||
if "很近" in voiceover:
|
||||
logger.warning("[爆款视频] 编导脚本含错别字'很近',hard fix 替换为'最近' label=%s", label)
|
||||
voiceover = voiceover.replace("很近", "最近")
|
||||
normalized["voiceover_script"] = voiceover
|
||||
# 同时在 shots 对白里替换
|
||||
for sh in normalized.get("shots") or []:
|
||||
if isinstance(sh, dict):
|
||||
sd = sh.get("scene_and_dialogue") or ""
|
||||
if "很近" in sd:
|
||||
sh["scene_and_dialogue"] = sd.replace("很近", "最近")
|
||||
# v1.6.1 P1修复:递归替换 copy_result 里所有字符串字段的"很近"→"最近"(覆盖 overview/scene_and_lighting/voiceover/shots.* 全部字段)
|
||||
_before_dump = json.dumps(normalized, ensure_ascii=False)
|
||||
if "很近" in _before_dump:
|
||||
logger.warning("[爆款视频] 编导脚本含错别字'很近',递归替换为'最近' label=%s", label)
|
||||
normalized = _replace_henjin_everywhere(normalized)
|
||||
voiceover = (normalized or {}).get("voiceover_script") or ""
|
||||
fallback_marker = "我最近在用的好物" in voiceover # _fallback_script 的特征串
|
||||
has_typo_henjin = "很近" in voiceover # v1.6.1: 错别字"很近"视为不合格,触发重试
|
||||
has_typo_henjin = "很近" in json.dumps(normalized, ensure_ascii=False) # 递归检查仍有"很近"视为不合格
|
||||
is_fallback = fallback_marker or shots_cnt < 1 or voiceover_len < 20 or has_typo_henjin
|
||||
logger.info(
|
||||
"[爆款视频] 编导脚本结果 label=%s voiceover_len=%d shots=%d fallback=%s raw_type=%s",
|
||||
@@ -1091,14 +1103,14 @@ def _assemble_seedance_prompt(copy_result: dict, job: ViralVideoJob) -> str:
|
||||
ab = s.get("audio_bgm", "")
|
||||
t = s.get("transition", "")
|
||||
ref = s.get("reference_image_index")
|
||||
lines.append(f"- 镜头{i+1}({tr}):")
|
||||
lines.append(f"- 镜头{i + 1}({tr}):")
|
||||
lines.append(f" 景别/运镜:{cam}")
|
||||
lines.append(f" 画面与对白:{sd}")
|
||||
lines.append(f" 动作细节:{act}")
|
||||
lines.append(f" 音效/BGM:{ab}")
|
||||
lines.append(f" 转场:{t}")
|
||||
if ref is not None and isinstance(ref, int):
|
||||
lines.append(f" 参考图片:第{ref+1}张产品图")
|
||||
lines.append(f" 参考图片:第{ref + 1}张产品图")
|
||||
lines.append("")
|
||||
lines.append("【硬性约束】")
|
||||
for c in hc:
|
||||
@@ -1114,6 +1126,7 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
|
||||
|
||||
返回 (本地视频路径, usage dict|None)。失败抛异常。
|
||||
"""
|
||||
from packages.domain.points_rules import get_viral_video_model_config
|
||||
from packages.shared.ai_service import call_video_generation
|
||||
|
||||
prompt = _assemble_seedance_prompt(copy_result, job)
|
||||
@@ -1121,6 +1134,9 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
|
||||
ratio = getattr(job, "video_ratio", None) or "9:16"
|
||||
model = getattr(job, "video_model", "") or None
|
||||
resolution = getattr(job, "video_resolution", "720p") or "720p"
|
||||
# 按模型配置决定是否开启音频生成(#2159 多模型支持)
|
||||
_mcfg = get_viral_video_model_config(model)
|
||||
gen_audio = bool(_mcfg.get("supports_audio", True))
|
||||
|
||||
# reference_audios: TTS 音频驱动口型
|
||||
ref_audios = [tts_audio_url] if tts_audio_url else []
|
||||
@@ -1133,10 +1149,12 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
|
||||
|
||||
tmpdir = Path(tempfile.mkdtemp(prefix=f"viral_{job.id}_"))
|
||||
logger.info(
|
||||
"[爆款视频] 开始单次 Seedance 生成 dur=%ds ratio=%s model=%s ref_imgs=%d ref_audios=%d ref_videos=%d tmpdir=%s",
|
||||
"[爆款视频] 开始单次视频生成 dur=%ds ratio=%s model=%s provider=%s gen_audio=%s ref_imgs=%d ref_audios=%d ref_videos=%d tmpdir=%s",
|
||||
dur,
|
||||
ratio if not first_image else "(follow-image)",
|
||||
model or "default",
|
||||
_mcfg.get("provider", "doubao"),
|
||||
gen_audio,
|
||||
len(rest_images) + (1 if first_image else 0),
|
||||
len(ref_audios),
|
||||
len(ref_videos),
|
||||
@@ -1144,6 +1162,7 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
|
||||
)
|
||||
logger.info("[爆款视频] Seedance prompt (前300字): %s", prompt[:300])
|
||||
|
||||
# 第一次调用:带参考图/首帧/音频/参考视频
|
||||
result = call_video_generation(
|
||||
prompt=prompt,
|
||||
image_url=first_image,
|
||||
@@ -1152,20 +1171,57 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
|
||||
resolution=resolution,
|
||||
output_dir=str(tmpdir),
|
||||
model=model,
|
||||
generate_audio=True, # Seedance 原生生成环境音效/BGM;口型由 reference_audios 的 TTS 驱动
|
||||
generate_audio=gen_audio, # 按模型能力:有声模型走原生音画同生;Wan 等需后配 TTS
|
||||
reference_images=rest_images,
|
||||
reference_audios=ref_audios,
|
||||
reference_videos=ref_videos,
|
||||
)
|
||||
|
||||
# #2169: 真人/肖像拦截已由 ai_client 内部自动切即梦(jimeng-3.0)通道重试——
|
||||
# 保留首帧图、不走"去掉参考图纯 t2v 降级"(用户明确要求按参考照片生成)。
|
||||
# 即梦也失败或非拦截类错误时,直接抛错给上层展示用户友好提示。
|
||||
def _check_and_reraise(result):
|
||||
if result and isinstance(result, dict):
|
||||
return result
|
||||
from packages.shared.ai_service import get_last_video_error
|
||||
|
||||
err = get_last_video_error() or {}
|
||||
user_msg = err.get("user_message") or ""
|
||||
detail = err.get("detail") or ""
|
||||
err_code = err.get("error_code") or "unknown"
|
||||
status_code = err.get("status_code", 0)
|
||||
err_provider = err.get("provider") or _mcfg.get("provider", "doubao")
|
||||
err_msg = user_msg or f"视频生成失败({err_provider} status={status_code} code={err_code})"
|
||||
logger.error(
|
||||
"[爆款视频] 视频生成失败: provider=%s model=%s code=%s status=%s user_msg=%s detail=%s",
|
||||
err_provider,
|
||||
model or "default",
|
||||
err_code,
|
||||
status_code,
|
||||
user_msg,
|
||||
(detail or "")[:500],
|
||||
)
|
||||
raise RuntimeError(err_msg)
|
||||
|
||||
if not result or not isinstance(result, dict):
|
||||
raise RuntimeError("Seedance 视频生成失败:返回为空")
|
||||
_check_and_reraise(result)
|
||||
video_path = result.get("video_path") or ""
|
||||
usage = result.get("usage")
|
||||
if not video_path or not Path(video_path).exists() or Path(video_path).stat().st_size == 0:
|
||||
raise RuntimeError("Seedance 视频生成失败:返回空文件或路径不存在")
|
||||
logger.info(
|
||||
"[爆款视频] Seedance 单次生成完成: %s size=%d usage=%s", video_path, Path(video_path).stat().st_size, usage
|
||||
)
|
||||
raise RuntimeError("视频生成失败:返回空文件或路径不存在")
|
||||
# #2169: 如果实际走了即梦兜底(真人拦截→jimeng),更新 job.video_model 让积分结算用 jimeng-3.0 价格
|
||||
if isinstance(usage, dict):
|
||||
actual_provider = usage.get("provider")
|
||||
actual_model_key = usage.get("model_key")
|
||||
if actual_provider == "jimeng" and actual_model_key:
|
||||
logger.info(
|
||||
"[爆款视频] 实际通过即梦通道生成(原model=%s),更新video_model=%s 用于积分结算",
|
||||
job.video_model,
|
||||
actual_model_key,
|
||||
)
|
||||
job.video_model = actual_model_key
|
||||
size = Path(video_path).stat().st_size
|
||||
logger.info("[爆款视频] 单次生成完成: path=%s size=%d usage=%s", video_path, size, usage)
|
||||
return str(video_path), (usage if isinstance(usage, dict) else None)
|
||||
|
||||
|
||||
@@ -1328,16 +1384,36 @@ def run_video_style_analysis(self: Task, job_id: str) -> dict:
|
||||
|
||||
|
||||
def _mark_failed_and_notify(job_id: str, session, repo, job, err_msg: str, stage: str = "") -> None:
|
||||
"""标记任务失败并通知。若传入的 session 已失效(因前面异常导致 rollback 状态),
|
||||
会自动 fallback 到新建 SessionLocal 重新标记,确保状态一定落库。"""
|
||||
try:
|
||||
if session is None:
|
||||
session = SessionLocal()
|
||||
repo = SQLAlchemyViralVideoJobRepository(session)
|
||||
job = repo.get(job_id)
|
||||
if job is not None and not job.is_terminal:
|
||||
job.mark_failed(err_msg)
|
||||
_save_job(repo, job, session)
|
||||
# 尝试用传入的 session 标记
|
||||
marked = False
|
||||
if job is not None and not job.is_terminal and session is not None:
|
||||
try:
|
||||
job.mark_failed(err_msg)
|
||||
_save_job(repo, job, session)
|
||||
marked = True
|
||||
except Exception as se:
|
||||
logger.warning("[爆款视频] 用原 session 标记失败失败,fallback 新session: %s", se)
|
||||
try:
|
||||
session.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
if not marked:
|
||||
# fallback:新建独立 session 重新标记(保证状态一定落库)
|
||||
ssn = SessionLocal()
|
||||
try:
|
||||
r = SQLAlchemyViralVideoJobRepository(ssn)
|
||||
j = r.get(job_id)
|
||||
if j is not None and not j.is_terminal:
|
||||
j.mark_failed(err_msg)
|
||||
r.update(j)
|
||||
ssn.commit()
|
||||
finally:
|
||||
ssn.close()
|
||||
except Exception as inner:
|
||||
logger.warning("[爆款视频] 标记失败状态时出错: %s", inner)
|
||||
logger.warning("[爆款视频] 标记失败状态时出错(最终fallback也失败): %s", inner, exc_info=True)
|
||||
_emit_progress(
|
||||
job_id,
|
||||
stage,
|
||||
|
||||
@@ -1,55 +0,0 @@
|
||||
"""商业交易模型 —— 与 xiaoxia-admin 共享表的 ORM 映射。
|
||||
|
||||
xiaoxia-admin(管理后台)与 xiaoxia-saas(用户端)共享同一个数据库:
|
||||
- plans / subscriptions 表由 admin 的 alembic 迁移创建(admin_alembic_version)
|
||||
- 用户端在支付链路中需要读取套餐、写入订阅,因此在此做最小映射
|
||||
|
||||
注意:不要在此给这些表建迁移;表结构变更走 xiaoxia-admin 仓库。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from sqlalchemy import Boolean, Column, DateTime, Integer, Numeric, String, Text
|
||||
from sqlalchemy.orm import declarative_base
|
||||
|
||||
CommerceBase: Any = declarative_base()
|
||||
|
||||
|
||||
def _now() -> datetime:
|
||||
return datetime.now(UTC)
|
||||
|
||||
|
||||
class PlanModel(CommerceBase):
|
||||
"""admin 端 plans 套餐表(只读映射)。"""
|
||||
|
||||
__tablename__ = "plans"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
level_id = Column(String(36), nullable=True)
|
||||
plan_key = Column(String(50), nullable=False)
|
||||
name = Column(String(100), nullable=False)
|
||||
description = Column(Text, nullable=True)
|
||||
monthly_price = Column(Numeric(10, 2), nullable=False, default=0)
|
||||
yearly_price = Column(Numeric(10, 2), nullable=False, default=0)
|
||||
is_enabled = Column(Boolean, nullable=False, default=True)
|
||||
display_order = Column(Integer, default=0)
|
||||
created_at = Column(DateTime, default=_now)
|
||||
|
||||
|
||||
class SubscriptionRecordModel(CommerceBase):
|
||||
"""admin 端 subscriptions 订阅记录表(支付成功后写入)。"""
|
||||
|
||||
__tablename__ = "subscriptions"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
user_id = Column(String(36), nullable=False, index=True)
|
||||
plan_id = Column(String(36), nullable=False, index=True)
|
||||
status = Column(String(20), nullable=False, default="active")
|
||||
billing_cycle = Column(String(10), nullable=False)
|
||||
start_date = Column(DateTime, nullable=False)
|
||||
end_date = Column(DateTime, nullable=False)
|
||||
cancelled_at = Column(DateTime, nullable=True)
|
||||
created_at = Column(DateTime, default=_now)
|
||||
updated_at = Column(DateTime, default=_now)
|
||||
@@ -816,12 +816,6 @@ class PointsOrderModel(Base):
|
||||
payment_method = Column(String(50), nullable=True)
|
||||
payment_id = Column(String(100), nullable=True)
|
||||
paid_at = Column(DateTime, nullable=True)
|
||||
# 微信支付链路补充字段
|
||||
out_trade_no = Column(String(64), nullable=True, index=True)
|
||||
prepay_id = Column(String(128), nullable=True)
|
||||
product_name = Column(String(100), nullable=True)
|
||||
payer_openid = Column(String(128), nullable=True)
|
||||
expire_at = Column(DateTime, nullable=True)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
|
||||
@@ -39,6 +39,9 @@ class SQLAlchemyUserRepository(UserRepository):
|
||||
model.phone_verified = user.phone_verified
|
||||
model.binding_completed_at = user.binding_completed_at
|
||||
model.profile_completed = user.profile_completed
|
||||
model.is_member = user.is_member
|
||||
model.member_type = user.member_type
|
||||
model.member_expires_at = user.member_expires_at
|
||||
model.created_at = user.created_at
|
||||
|
||||
self.session.commit()
|
||||
@@ -115,5 +118,8 @@ class SQLAlchemyUserRepository(UserRepository):
|
||||
phone_verified=model.phone_verified or False,
|
||||
binding_completed_at=model.binding_completed_at,
|
||||
profile_completed=model.profile_completed if model.profile_completed is not None else True,
|
||||
is_member=model.is_member if model.is_member is not None else False,
|
||||
member_type=model.member_type,
|
||||
member_expires_at=model.member_expires_at,
|
||||
created_at=model.created_at,
|
||||
)
|
||||
|
||||
@@ -1,97 +0,0 @@
|
||||
"""微信支付平台证书缓存。
|
||||
|
||||
回调验签需要「微信支付平台公钥」。通过 V3 接口
|
||||
GET /v3/certificates 获取(响应用本商户 APIv3 密钥加密),
|
||||
带内存缓存(证书有效期通常约12个月)。
|
||||
|
||||
部署为单进程时内存缓存足够;多副本部署可改为 Redis 缓存。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import logging
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from cryptography.hazmat.primitives import serialization
|
||||
|
||||
from packages.adapters.wechat_pay import (
|
||||
WECHAT_BASE_URL,
|
||||
build_authorization,
|
||||
load_private_key,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
CERT_PATH = "/v3/certificates"
|
||||
|
||||
# serial -> {"public_key": obj, "expires_at": ts}
|
||||
_cache: dict[str, dict[str, Any]] = {}
|
||||
|
||||
|
||||
def _decrypt_cert_blob(*, api_v3_key: str, nonce: str, associated_data: str, ciphertext_b64: str) -> bytes:
|
||||
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
|
||||
|
||||
blob = base64.b64decode(ciphertext_b64)
|
||||
tag, data = blob[-16:], blob[:-16]
|
||||
decryptor = Cipher(algorithms.AES(api_v3_key.encode()), modes.GCM(nonce.encode(), tag)).decryptor()
|
||||
return decryptor.update(data) + decryptor.finalize_with_associated_data(associated_data.encode())
|
||||
|
||||
|
||||
def refresh_platform_certificates() -> dict[str, Any]:
|
||||
"""拉取并刷新平台证书缓存。"""
|
||||
from app.config import settings
|
||||
|
||||
s = settings
|
||||
if not getattr(s, "wechat_pay_configured", False):
|
||||
raise RuntimeError("微信支付未配置,无法拉取平台证书")
|
||||
|
||||
private_key = load_private_key(s.wechat_private_key)
|
||||
auth = build_authorization("GET", CERT_PATH, s.wechat_mch_id, s.wechat_appid, s.wechat_cert_serial, private_key)
|
||||
|
||||
with httpx.Client(timeout=15.0) as client:
|
||||
resp = client.get(
|
||||
f"{WECHAT_BASE_URL}{CERT_PATH}",
|
||||
headers={"Authorization": auth, "Accept": "application/json"},
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
|
||||
new_cache: dict[str, dict[str, Any]] = {}
|
||||
for item in data.get("data", []):
|
||||
serial = item["serial_no"]
|
||||
enc = item["encrypt_certificate"]
|
||||
pem = _decrypt_cert_blob(
|
||||
api_v3_key=s.wechat_api_v3_key,
|
||||
nonce=enc["nonce"],
|
||||
associated_data=enc["associated_data"],
|
||||
ciphertext_b64=enc["ciphertext"],
|
||||
)
|
||||
cert = serialization.load_pem_x509_certificate(pem)
|
||||
public_key = cert.public_key()
|
||||
expires_at = (
|
||||
time.mktime(cert.not_valid_after_utc.timetuple())
|
||||
if hasattr(cert, "not_valid_after_utc")
|
||||
else time.time() + 365 * 24 * 3600
|
||||
)
|
||||
new_cache[serial] = {"public_key": public_key, "expires_at": expires_at}
|
||||
|
||||
_cache.clear()
|
||||
_cache.update(new_cache)
|
||||
logger.info("微信平台证书已刷新,共 %s 张", len(new_cache))
|
||||
return {"count": len(new_cache), "serials": list(new_cache.keys())}
|
||||
|
||||
|
||||
def get_platform_public_key(serial: str):
|
||||
"""按证书序列号取平台公钥;缓存缺失/过期时自动拉取。"""
|
||||
now = time.time()
|
||||
entry = _cache.get(serial)
|
||||
if entry is None or entry["expires_at"] <= now:
|
||||
try:
|
||||
refresh_platform_certificates()
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception("刷新微信平台证书失败")
|
||||
entry = _cache.get(serial)
|
||||
return entry["public_key"] if entry else None
|
||||
@@ -1,232 +0,0 @@
|
||||
"""微信支付 V3 适配层 —— JSAPI 下单、回调验签与解密、前端支付参数签名。
|
||||
|
||||
纯适配,不含业务逻辑:
|
||||
- create_jsapi_order: JSAPI 统一下单,返回 prepay_id
|
||||
- build_jsapi_pay_params: 生成调起微信支付所需参数(商户私钥二次签名)
|
||||
- verify_notification_signature: 校验回调平台证书签名(含时间戳防重放)
|
||||
- decrypt_resource: 用 APIv3 密钥解密 resource(AEAD_AES_256_GCM)
|
||||
|
||||
依赖:httpx + cryptography(requirements-base.txt 中已含 httpx)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from cryptography.hazmat.primitives import hashes, serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import padding
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
WECHAT_BASE_URL = "https://api.mch.weixin.qq.com"
|
||||
JSAPI_ORDER_PATH = "/v3/pay/transactions/jsapi"
|
||||
|
||||
# 回调重放窗口:通知时间戳距当前超过 5 分钟视为过期
|
||||
NOTIFY_MAX_AGE_SECONDS = 300
|
||||
|
||||
|
||||
class WeChatPayError(Exception):
|
||||
"""微信支付调用错误。"""
|
||||
|
||||
|
||||
# ──────────────────────── 密钥/证书加载 ────────────────────────
|
||||
|
||||
|
||||
def load_private_key(private_key_src: str):
|
||||
"""加载商户 API 私钥。
|
||||
|
||||
支持:
|
||||
- PEM 文本(含 BEGIN PRIVATE KEY / BEGIN RSA PRIVATE KEY)
|
||||
- PEM 文件绝对路径
|
||||
- 单行 base64 编码的 PEM(容器 env 不支持换行时使用)
|
||||
"""
|
||||
src = (private_key_src or "").strip()
|
||||
if not src:
|
||||
raise WeChatPayError("微信商户私钥为空")
|
||||
|
||||
if "BEGIN" not in src and "\n" not in src:
|
||||
# 先尝试文件路径
|
||||
try:
|
||||
with open(src, "rb") as f:
|
||||
data = f.read()
|
||||
return serialization.load_pem_private_key(data, password=None)
|
||||
except (OSError, ValueError):
|
||||
pass
|
||||
# 再尝试单行 base64
|
||||
try:
|
||||
src = base64.b64decode(src).decode()
|
||||
except Exception as e:
|
||||
raise WeChatPayError("无法解析微信商户私钥(不是PEM文本/文件/base64)") from e
|
||||
|
||||
if "\\n" in src:
|
||||
src = src.replace("\\n", "\n")
|
||||
|
||||
try:
|
||||
return serialization.load_pem_private_key(src.encode(), password=None)
|
||||
except ValueError as e:
|
||||
raise WeChatPayError(f"微信商户私钥格式无效: {e}") from e
|
||||
|
||||
|
||||
def _rsa_sign(private_key, message: str) -> str:
|
||||
"""RSA-SHA256 签名,返回 base64。"""
|
||||
signature = private_key.sign(message.encode(), padding.PKCS1v15(), hashes.SHA256())
|
||||
return base64.b64encode(signature).decode()
|
||||
|
||||
|
||||
# ──────────────────────── Authorization 头 ────────────────────────
|
||||
|
||||
|
||||
def build_authorization(
|
||||
method: str,
|
||||
url_path: str,
|
||||
mch_id: str,
|
||||
appid: str,
|
||||
cert_serial: str,
|
||||
private_key,
|
||||
body: str = "",
|
||||
) -> str:
|
||||
"""构造微信 V3 API 的 Authorization 头。"""
|
||||
timestamp = str(int(time.time()))
|
||||
nonce = uuid.uuid4().hex
|
||||
message = f"{method}\n{url_path}\n{timestamp}\n{nonce}\n{body}\n"
|
||||
signature = _rsa_sign(private_key, message)
|
||||
return (
|
||||
f'WECHATPAY2-SHA256-RSA2048 mchid="{mch_id}",nonce_str="{nonce}",'
|
||||
f'timestamp="{timestamp}",serial_no="{cert_serial}",signature="{signature}"'
|
||||
)
|
||||
|
||||
|
||||
# ──────────────────────── JSAPI 统一下单 ────────────────────────
|
||||
|
||||
|
||||
def create_jsapi_order(
|
||||
*,
|
||||
appid: str,
|
||||
mch_id: str,
|
||||
cert_serial: str,
|
||||
private_key,
|
||||
description: str,
|
||||
out_trade_no: str,
|
||||
amount_cents: int,
|
||||
openid: str,
|
||||
notify_url: str,
|
||||
attach: str = "",
|
||||
) -> dict[str, Any]:
|
||||
"""调用 V3 JSAPI 统一下单,返回 {"prepay_id": ...}。"""
|
||||
payload: dict[str, Any] = {
|
||||
"appid": appid,
|
||||
"mchid": mch_id,
|
||||
"description": description,
|
||||
"out_trade_no": out_trade_no,
|
||||
"notify_url": notify_url,
|
||||
"amount": {"total": int(amount_cents), "currency": "CNY"},
|
||||
"payer": {"openid": openid},
|
||||
}
|
||||
if attach:
|
||||
payload["attach"] = attach
|
||||
|
||||
body = json.dumps(payload, separators=(",", ":"), ensure_ascii=False)
|
||||
authorization = build_authorization("POST", JSAPI_ORDER_PATH, mch_id, appid, cert_serial, private_key, body)
|
||||
|
||||
try:
|
||||
with httpx.Client(timeout=15.0) as client:
|
||||
resp = client.post(
|
||||
f"{WECHAT_BASE_URL}{JSAPI_ORDER_PATH}",
|
||||
content=body.encode("utf-8"),
|
||||
headers={
|
||||
"Authorization": authorization,
|
||||
"Accept": "application/json",
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
)
|
||||
except httpx.HTTPError as e:
|
||||
raise WeChatPayError(f"微信下单网络异常: {e}") from e
|
||||
|
||||
if resp.status_code >= 400:
|
||||
logger.error("微信下单失败: status=%s body=%s", resp.status_code, resp.text)
|
||||
raise WeChatPayError(f"微信下单失败({resp.status_code}): {resp.text}")
|
||||
|
||||
data = resp.json()
|
||||
if not data.get("prepay_id"):
|
||||
raise WeChatPayError(f"微信下单未返回 prepay_id: {data}")
|
||||
return data
|
||||
|
||||
|
||||
def build_jsapi_pay_params(appid: str, prepay_id: str, private_key) -> dict[str, str]:
|
||||
"""生成小程序调起支付参数(对 prepay_id 二次签名)。"""
|
||||
timestamp = str(int(time.time()))
|
||||
nonce = uuid.uuid4().hex
|
||||
package = f"prepay_id={prepay_id}"
|
||||
message = f"{appid}\n{timestamp}\n{nonce}\n{package}\n"
|
||||
return {
|
||||
"appId": appid,
|
||||
"timeStamp": timestamp,
|
||||
"nonceStr": nonce,
|
||||
"package": package,
|
||||
"signType": "RSA",
|
||||
"paySign": _rsa_sign(private_key, message),
|
||||
}
|
||||
|
||||
|
||||
# ──────────────────────── 回调验签与解密 ────────────────────────
|
||||
|
||||
|
||||
def verify_notification_signature(
|
||||
*,
|
||||
timestamp: str,
|
||||
nonce: str,
|
||||
body: bytes,
|
||||
signature_b64: str,
|
||||
platform_public_key,
|
||||
) -> bool:
|
||||
"""用微信平台公钥校验回调签名,并检查时间戳防重放。"""
|
||||
try:
|
||||
ts = int(timestamp)
|
||||
except (TypeError, ValueError):
|
||||
return False
|
||||
if abs(time.time() - ts) > NOTIFY_MAX_AGE_SECONDS:
|
||||
logger.warning("微信回调时间戳超出允许窗口: %s", timestamp)
|
||||
return False
|
||||
|
||||
message = f"{timestamp}\n{nonce}\n".encode() + body + b"\n"
|
||||
try:
|
||||
platform_public_key.verify(
|
||||
base64.b64decode(signature_b64),
|
||||
message,
|
||||
padding.PKCS1v15(),
|
||||
hashes.SHA256(),
|
||||
)
|
||||
return True
|
||||
except Exception:
|
||||
logger.warning("微信回调签名校验失败")
|
||||
return False
|
||||
|
||||
|
||||
def decrypt_resource(
|
||||
*,
|
||||
api_v3_key: str,
|
||||
nonce: str,
|
||||
associated_data: str,
|
||||
ciphertext_b64: str,
|
||||
) -> dict[str, Any]:
|
||||
"""解密回调 resource(AEAD_AES_256_GCM)。"""
|
||||
try:
|
||||
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
|
||||
|
||||
ciphertext = base64.b64decode(ciphertext_b64)
|
||||
tag = ciphertext[-16:]
|
||||
data = ciphertext[:-16]
|
||||
decryptor = Cipher(
|
||||
algorithms.AES(api_v3_key.encode("utf-8")),
|
||||
modes.GCM(nonce.encode("utf-8"), tag),
|
||||
).decryptor()
|
||||
plaintext = decryptor.update(data) + decryptor.finalize_with_associated_data(associated_data.encode("utf-8"))
|
||||
return json.loads(plaintext.decode("utf-8"))
|
||||
except Exception as e:
|
||||
raise WeChatPayError(f"回调解密失败: {e}") from e
|
||||
@@ -0,0 +1 @@
|
||||
"""应用层:对外展示目录(套餐/积分包)。"""
|
||||
@@ -0,0 +1,152 @@
|
||||
"""读取管理后台配置的会员套餐 / 积分充值包(共享库真实数据)。
|
||||
|
||||
替代旧的硬编码 MEMBERSHIP_PRICES / POINTS_PACKAGES。
|
||||
短 TTL 缓存(30 秒),后台改价/启停后用户端最多 30 秒可见。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
_CACHE_TTL = 30.0
|
||||
_lock = threading.Lock()
|
||||
_cache: dict[str, tuple[float, Any]] = {}
|
||||
|
||||
_QUOTA_LABELS = {
|
||||
"4k": "4K 超清分辨率",
|
||||
"batch_render": "批量渲染",
|
||||
"priority_queue": "优先处理队列",
|
||||
"ai_matting": "AI 智能抠像",
|
||||
"remove_watermark": "去水印",
|
||||
}
|
||||
|
||||
|
||||
def _cached(key: str, loader):
|
||||
now = time.time()
|
||||
hit = _cache.get(key)
|
||||
if hit and now - hit[0] < _CACHE_TTL:
|
||||
return hit[1]
|
||||
with _lock:
|
||||
hit = _cache.get(key)
|
||||
if hit and time.time() - hit[0] < _CACHE_TTL:
|
||||
return hit[1]
|
||||
value = loader()
|
||||
_cache[key] = (time.time(), value)
|
||||
return value
|
||||
|
||||
|
||||
def _quota_features(quotas: dict[str, Any] | None) -> dict[str, Any]:
|
||||
quotas = quotas or {}
|
||||
features: dict[str, Any] = {}
|
||||
for k, v in quotas.items():
|
||||
if k == "credits_per_month":
|
||||
features["credits_per_month"] = v
|
||||
elif k in _QUOTA_LABELS:
|
||||
features[_QUOTA_LABELS[k]] = v
|
||||
else:
|
||||
features[k] = v
|
||||
return features
|
||||
|
||||
|
||||
def get_membership_plans() -> list[dict[str, Any]]:
|
||||
"""读取 is_enabled=true 的套餐,按年/月周期展开为用户端档位。"""
|
||||
|
||||
def _load() -> list[dict[str, Any]]:
|
||||
from sqlalchemy import text
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.session import SessionLocal
|
||||
|
||||
if SessionLocal is None:
|
||||
return []
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
rows = session.execute(text("""
|
||||
SELECT plan_key, name, description, monthly_price, yearly_price,
|
||||
quotas, display_order
|
||||
FROM plans
|
||||
WHERE is_enabled = TRUE
|
||||
ORDER BY display_order NULLS LAST, created_at
|
||||
""")).fetchall()
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
plans: list[dict[str, Any]] = []
|
||||
for r in rows:
|
||||
base_features = _quota_features(r.quotas if isinstance(r.quotas, dict) else None)
|
||||
if r.yearly_price and float(r.yearly_price) > 0:
|
||||
plans.append(
|
||||
{
|
||||
"plan_id": r.plan_key,
|
||||
"billing_cycle": "yearly",
|
||||
"name": r.name,
|
||||
"description": r.description,
|
||||
"price_cents": int(round(float(r.yearly_price) * 100)),
|
||||
"monthly_price_cents": int(round(float(r.yearly_price) * 100 / 12)),
|
||||
"duration_days": 365,
|
||||
"features": dict(base_features),
|
||||
}
|
||||
)
|
||||
if r.monthly_price and float(r.monthly_price) > 0:
|
||||
plans.append(
|
||||
{
|
||||
"plan_id": r.plan_key,
|
||||
"billing_cycle": "monthly",
|
||||
"name": r.name,
|
||||
"description": r.description,
|
||||
"price_cents": int(round(float(r.monthly_price) * 100)),
|
||||
"monthly_price_cents": int(round(float(r.monthly_price) * 100)),
|
||||
"duration_days": 30,
|
||||
"features": dict(base_features),
|
||||
}
|
||||
)
|
||||
return plans
|
||||
|
||||
return _cached("membership_plans", _load)
|
||||
|
||||
|
||||
def get_points_packages() -> list[dict[str, Any]]:
|
||||
"""读取 is_active=true 的积分充值包。"""
|
||||
|
||||
def _load() -> list[dict[str, Any]]:
|
||||
from sqlalchemy import text
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.session import SessionLocal
|
||||
|
||||
if SessionLocal is None:
|
||||
return []
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
rows = session.execute(text("""
|
||||
SELECT package_key, name, price, credits, bonus_credits,
|
||||
is_recommended, description, sort_order
|
||||
FROM credit_packages
|
||||
WHERE is_active = TRUE
|
||||
ORDER BY sort_order NULLS LAST, price
|
||||
""")).fetchall()
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
packages: list[dict[str, Any]] = []
|
||||
for r in rows:
|
||||
total_points = int(r.credits or 0) + int(r.bonus_credits or 0)
|
||||
price_cents = int(round(float(r.price) * 100))
|
||||
unit = (price_cents / 100 / total_points) if total_points else 0
|
||||
packages.append(
|
||||
{
|
||||
"code": r.package_key,
|
||||
"name": r.name,
|
||||
"points": total_points,
|
||||
"bonus_credits": int(r.bonus_credits or 0),
|
||||
"price_cents": price_cents,
|
||||
"unit_price": f"¥{unit:.3f}/积分",
|
||||
"is_recommended": bool(r.is_recommended),
|
||||
"description": r.description,
|
||||
}
|
||||
)
|
||||
return packages
|
||||
|
||||
return _cached("points_packages", _load)
|
||||
@@ -1,520 +0,0 @@
|
||||
"""支付应用服务 —— 会员年卡/积分包购买的下单、回调履约、订单查询。
|
||||
|
||||
编排 xiaoxia-saas 自身能力(积分账户/用户表)与微信支付适配层,
|
||||
同时写入 admin 端共享的 subscriptions 表,保持后台数据一致。
|
||||
|
||||
一期范围:
|
||||
- 会员年卡:plan_id 支持 admin plans 表 UUID(主)与 legacy "yearly"(兼容)
|
||||
- 积分包:product_code 对应 POINTS_PACKAGES
|
||||
- 自动续费不做,auto_renew 固定 false
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import uuid
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.wechat_pay import (
|
||||
WeChatPayError,
|
||||
build_jsapi_pay_params,
|
||||
create_jsapi_order,
|
||||
decrypt_resource,
|
||||
load_private_key,
|
||||
verify_notification_signature,
|
||||
)
|
||||
from packages.domain.points_rules import MEMBERSHIP_PRICES, POINTS_PACKAGES
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# legacy plan_id -> 会员天数
|
||||
_LEGACY_DURATION = {
|
||||
"monthly": 30,
|
||||
"quarterly": 90,
|
||||
"yearly": 365,
|
||||
}
|
||||
_ORDER_TTL_HOURS = 48
|
||||
|
||||
|
||||
class PaymentConfigError(Exception):
|
||||
"""支付配置缺失。"""
|
||||
|
||||
|
||||
class PaymentService:
|
||||
"""会员/积分支付编排。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.points_service = PointsService()
|
||||
|
||||
# ──────────────────────── 内部工具 ────────────────────────
|
||||
|
||||
@staticmethod
|
||||
def _settings():
|
||||
from app.config import settings
|
||||
|
||||
return settings
|
||||
|
||||
def _require_wechat(self):
|
||||
"""校验微信配置并返回 (settings, private_key)。"""
|
||||
settings = self._settings()
|
||||
if not getattr(settings, "wechat_pay_configured", False):
|
||||
raise PaymentConfigError(
|
||||
"微信支付未配置或配置不完整(WECHAT_PAY_ENABLED/AppID/MchId/APIv3Key/CertSerial/PrivateKey/NotifyUrl)"
|
||||
)
|
||||
return settings, load_private_key(settings.wechat_private_key)
|
||||
|
||||
@staticmethod
|
||||
def _resolve_plan(db: Session, plan_id: str, billing_cycle: str) -> dict[str, Any]:
|
||||
"""解析套餐:优先 admin plans 表 UUID;否则按 legacy 会员类型。
|
||||
|
||||
Returns:
|
||||
{"plan_id", "name", "amount_cents", "duration_days", "db_plan"}
|
||||
"""
|
||||
# 1) admin plans 表
|
||||
from packages.adapters.sqlalchemy_impl.commerce_models import PlanModel
|
||||
|
||||
plan = db.query(PlanModel).filter(PlanModel.id == plan_id).first()
|
||||
if plan is not None:
|
||||
if not plan.is_enabled:
|
||||
raise ValueError("套餐已下架")
|
||||
if billing_cycle == "yearly":
|
||||
price = plan.yearly_price
|
||||
days = 365
|
||||
elif billing_cycle == "monthly":
|
||||
price = plan.monthly_price
|
||||
days = 30
|
||||
else:
|
||||
raise ValueError("一期仅支持 monthly/yearly 计费周期")
|
||||
amount_cents = int(round(float(price) * 100))
|
||||
if amount_cents <= 0:
|
||||
raise ValueError("该计费周期价格未配置(价格为0),无法下单")
|
||||
return {
|
||||
"plan_id": plan.id,
|
||||
"name": plan.name,
|
||||
"amount_cents": amount_cents,
|
||||
"duration_days": days,
|
||||
"db_plan": plan,
|
||||
}
|
||||
|
||||
# 2) legacy:plan_id 本身是会员类型
|
||||
if plan_id in _LEGACY_DURATION and plan_id == billing_cycle:
|
||||
info = MEMBERSHIP_PRICES.get(plan_id)
|
||||
if not info:
|
||||
raise ValueError(f"未知会员类型: {plan_id}")
|
||||
return {
|
||||
"plan_id": plan_id,
|
||||
"name": info["name"],
|
||||
"amount_cents": info["price_cents"],
|
||||
"duration_days": info["duration_days"],
|
||||
"db_plan": None,
|
||||
}
|
||||
|
||||
raise ValueError(f"套餐不存在: {plan_id}")
|
||||
|
||||
# ──────────────────────── 会员下单 ────────────────────────
|
||||
|
||||
def create_membership_order(
|
||||
self,
|
||||
*,
|
||||
user_id: str,
|
||||
plan_id: str,
|
||||
billing_cycle: str,
|
||||
openid: str,
|
||||
db: Session,
|
||||
) -> dict[str, Any]:
|
||||
"""创建会员订单并向微信下单,返回订单 + 前端支付参数。"""
|
||||
from packages.adapters.sqlalchemy_impl.models import PointsOrderModel
|
||||
|
||||
settings, private_key = self._require_wechat()
|
||||
if not openid:
|
||||
raise ValueError("用户未绑定微信(缺少 openid),请先使用微信登录")
|
||||
|
||||
resolved = self._resolve_plan(db, plan_id, billing_cycle)
|
||||
|
||||
now = datetime.now(UTC)
|
||||
out_trade_no = f"mb{now.strftime('%Y%m%d%H%M%S')}{uuid.uuid4().hex[:10]}"
|
||||
expire_at = now + timedelta(hours=_ORDER_TTL_HOURS)
|
||||
|
||||
order = PointsOrderModel(
|
||||
id=uuid.uuid4().hex,
|
||||
user_id=user_id,
|
||||
order_type="membership",
|
||||
product_code=resolved["plan_id"],
|
||||
amount_cents=resolved["amount_cents"],
|
||||
original_amount_cents=resolved["amount_cents"],
|
||||
product_name=resolved["name"],
|
||||
payer_openid=openid,
|
||||
status="pending",
|
||||
expire_at=expire_at,
|
||||
)
|
||||
db.add(order)
|
||||
db.flush() # 先不 commit,微信失败则回滚
|
||||
|
||||
try:
|
||||
wx_resp = create_jsapi_order(
|
||||
appid=settings.wechat_appid,
|
||||
mch_id=settings.wechat_mch_id,
|
||||
cert_serial=settings.wechat_cert_serial,
|
||||
private_key=private_key,
|
||||
description=resolved["name"],
|
||||
out_trade_no=out_trade_no,
|
||||
amount_cents=resolved["amount_cents"],
|
||||
openid=openid,
|
||||
notify_url=settings.wechat_notify_url,
|
||||
attach=order.id,
|
||||
)
|
||||
except WeChatPayError as e:
|
||||
db.rollback()
|
||||
raise WeChatPayError(str(e)) from e
|
||||
|
||||
prepay_id = wx_resp["prepay_id"]
|
||||
order.prepay_id = prepay_id
|
||||
order.out_trade_no = out_trade_no
|
||||
db.commit()
|
||||
|
||||
pay_params = build_jsapi_pay_params(settings.wechat_appid, prepay_id, private_key)
|
||||
return {
|
||||
"order_id": order.id,
|
||||
"out_trade_no": out_trade_no,
|
||||
"prepay_id": prepay_id,
|
||||
"pay_params": pay_params,
|
||||
"amount_cents": resolved["amount_cents"],
|
||||
"expire_at": expire_at.isoformat(),
|
||||
}
|
||||
|
||||
# ──────────────────────── 积分包下单 ────────────────────────
|
||||
|
||||
def create_points_order(
|
||||
self,
|
||||
*,
|
||||
user_id: str,
|
||||
package_code: str,
|
||||
openid: str,
|
||||
db: Session,
|
||||
) -> dict[str, Any]:
|
||||
"""创建积分包订单并向微信下单。"""
|
||||
from packages.adapters.sqlalchemy_impl.models import PointsOrderModel
|
||||
|
||||
settings, private_key = self._require_wechat()
|
||||
package = POINTS_PACKAGES.get(package_code)
|
||||
if not package:
|
||||
raise ValueError(f"积分包不存在: {package_code}")
|
||||
if not openid:
|
||||
raise ValueError("用户未绑定微信(缺少 openid),请先使用微信登录")
|
||||
|
||||
now = datetime.now(UTC)
|
||||
out_trade_no = f"pt{now.strftime('%Y%m%d%H%M%S')}{uuid.uuid4().hex[:10]}"
|
||||
expire_at = now + timedelta(hours=_ORDER_TTL_HOURS)
|
||||
|
||||
order = PointsOrderModel(
|
||||
id=uuid.uuid4().hex,
|
||||
user_id=user_id,
|
||||
order_type="points",
|
||||
product_code=package_code,
|
||||
amount_cents=package["price_cents"],
|
||||
original_amount_cents=package["price_cents"],
|
||||
points_amount=package["points"],
|
||||
product_name=package["name"],
|
||||
payer_openid=openid,
|
||||
status="pending",
|
||||
expire_at=expire_at,
|
||||
)
|
||||
db.add(order)
|
||||
db.flush()
|
||||
|
||||
try:
|
||||
wx_resp = create_jsapi_order(
|
||||
appid=settings.wechat_appid,
|
||||
mch_id=settings.wechat_mch_id,
|
||||
cert_serial=settings.wechat_cert_serial,
|
||||
private_key=private_key,
|
||||
description=package["name"],
|
||||
out_trade_no=out_trade_no,
|
||||
amount_cents=package["price_cents"],
|
||||
openid=openid,
|
||||
notify_url=settings.wechat_notify_url,
|
||||
attach=order.id,
|
||||
)
|
||||
except WeChatPayError as e:
|
||||
db.rollback()
|
||||
raise WeChatPayError(str(e)) from e
|
||||
|
||||
prepay_id = wx_resp["prepay_id"]
|
||||
order.prepay_id = prepay_id
|
||||
order.out_trade_no = out_trade_no
|
||||
db.commit()
|
||||
|
||||
pay_params = build_jsapi_pay_params(settings.wechat_appid, prepay_id, private_key)
|
||||
return {
|
||||
"order_id": order.id,
|
||||
"out_trade_no": out_trade_no,
|
||||
"prepay_id": prepay_id,
|
||||
"pay_params": pay_params,
|
||||
"amount_cents": package["price_cents"],
|
||||
"points_amount": package["points"],
|
||||
"expire_at": expire_at.isoformat(),
|
||||
}
|
||||
|
||||
# ──────────────────────── 回调履约 ────────────────────────
|
||||
|
||||
def handle_wechat_notification(
|
||||
self,
|
||||
*,
|
||||
headers: dict[str, str],
|
||||
body: bytes,
|
||||
platform_public_key_loader,
|
||||
db: Session,
|
||||
) -> dict[str, Any]:
|
||||
"""处理微信支付回调:验签 → 解密 → 验金额 → 履约。
|
||||
|
||||
Args:
|
||||
platform_public_key_loader: callable(serial: str) -> 公钥对象 | None
|
||||
"""
|
||||
from packages.adapters.sqlalchemy_impl.models import PointsOrderModel
|
||||
|
||||
settings = self._settings()
|
||||
|
||||
timestamp = headers.get("wechatpay-timestamp", "")
|
||||
nonce = headers.get("wechatpay-nonce", "")
|
||||
serial = headers.get("wechatpay-serial", "")
|
||||
signature = headers.get("wechatpay-signature", "")
|
||||
|
||||
if not all([timestamp, nonce, serial, signature]):
|
||||
raise WeChatPayError("回调头不完整")
|
||||
|
||||
platform_public_key = platform_public_key_loader(serial)
|
||||
if platform_public_key is None:
|
||||
raise WeChatPayError("未找到对应微信平台证书")
|
||||
|
||||
if not verify_notification_signature(
|
||||
timestamp=timestamp,
|
||||
nonce=nonce,
|
||||
body=body,
|
||||
signature_b64=signature,
|
||||
platform_public_key=platform_public_key,
|
||||
):
|
||||
raise WeChatPayError("回调签名校验失败")
|
||||
|
||||
envelope = _json_loads(body)
|
||||
resource = envelope.get("resource", {})
|
||||
event_type = envelope.get("event_type", "")
|
||||
if event_type != "TRANSACTION.SUCCESS":
|
||||
logger.info("忽略非支付成功事件: %s", event_type)
|
||||
return {"ignored": True, "event_type": event_type}
|
||||
|
||||
payment = decrypt_resource(
|
||||
api_v3_key=settings.wechat_api_v3_key,
|
||||
nonce=resource.get("nonce", ""),
|
||||
associated_data=resource.get("associated_data", ""),
|
||||
ciphertext_b64=resource.get("ciphertext", ""),
|
||||
)
|
||||
|
||||
out_trade_no = payment.get("out_trade_no", "")
|
||||
transaction_id = payment.get("transaction_id", "")
|
||||
payer_total = int(payment.get("amount", {}).get("payer_total", -1))
|
||||
trade_state = payment.get("trade_state", "")
|
||||
|
||||
# 行锁取订单:
|
||||
# 1) out_trade_no 独立列(主)
|
||||
# 2) attach(下单时写入的 order.id)
|
||||
# 3) payment_id(兼容旧数据)
|
||||
order = (
|
||||
db.query(PointsOrderModel).filter(PointsOrderModel.out_trade_no == out_trade_no).with_for_update().first()
|
||||
)
|
||||
if order is None:
|
||||
attach_id = payment.get("attach") or envelope.get("id")
|
||||
if attach_id:
|
||||
order = db.query(PointsOrderModel).filter(PointsOrderModel.id == attach_id).with_for_update().first()
|
||||
if order is None:
|
||||
order = (
|
||||
db.query(PointsOrderModel).filter(PointsOrderModel.payment_id == out_trade_no).with_for_update().first()
|
||||
)
|
||||
if order is None:
|
||||
raise WeChatPayError(f"订单不存在: {out_trade_no}")
|
||||
|
||||
if order.status == "paid":
|
||||
# 幂等:微信可能重复通知
|
||||
return {"success": True, "order_id": order.id, "idempotent": True}
|
||||
|
||||
if trade_state and trade_state != "SUCCESS":
|
||||
raise WeChatPayError(f"交易状态非成功: {trade_state}")
|
||||
|
||||
if payer_total != order.amount_cents:
|
||||
raise WeChatPayError(f"金额不一致: 订单 {order.amount_cents} / 实付 {payer_total}")
|
||||
|
||||
# ── 履约 ──
|
||||
now = datetime.now(UTC)
|
||||
order.status = "paid"
|
||||
order.payment_method = "wechat"
|
||||
order.payment_id = transaction_id
|
||||
order.paid_at = now
|
||||
|
||||
if order.order_type == "points":
|
||||
self.points_service.add_points(
|
||||
user_id=order.user_id,
|
||||
amount=order.points_amount,
|
||||
source=f"recharge:{order.product_code}",
|
||||
db=db,
|
||||
description=f"积分充值: {order.product_name or order.product_code}",
|
||||
ref_id=order.id,
|
||||
)
|
||||
elif order.order_type == "membership":
|
||||
duration = self._duration_for(order)
|
||||
self._activate_membership(
|
||||
user_id=order.user_id,
|
||||
plan_id=order.product_code,
|
||||
duration_days=duration,
|
||||
order_id=order.id,
|
||||
now=now,
|
||||
db=db,
|
||||
)
|
||||
else:
|
||||
raise WeChatPayError(f"未知订单类型: {order.order_type}")
|
||||
|
||||
db.commit()
|
||||
return {"success": True, "order_id": order.id}
|
||||
|
||||
@staticmethod
|
||||
def _duration_for(order: Any) -> int:
|
||||
"""根据订单商品推断会员天数。"""
|
||||
code = order.product_code
|
||||
if code in _LEGACY_DURATION:
|
||||
return _LEGACY_DURATION[code]
|
||||
# admin plan:查 plan_key / 名称或固定365(年卡订单)
|
||||
return 365
|
||||
|
||||
def _activate_membership(
|
||||
self,
|
||||
*,
|
||||
user_id: str,
|
||||
plan_id: str,
|
||||
duration_days: int,
|
||||
order_id: str,
|
||||
now: datetime,
|
||||
db: Session,
|
||||
) -> None:
|
||||
"""激活会员:更新 users 表 + 写 admin subscriptions + 发积分。"""
|
||||
from packages.adapters.sqlalchemy_impl.commerce_models import (
|
||||
PlanModel,
|
||||
SubscriptionRecordModel,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.models import UserModel
|
||||
|
||||
user = db.query(UserModel).filter(UserModel.id == user_id).with_for_update().first()
|
||||
if user is None:
|
||||
raise WeChatPayError(f"用户不存在: {user_id}")
|
||||
|
||||
base = user.member_expires_at or now
|
||||
if base < now:
|
||||
base = now
|
||||
expires_at = base + timedelta(days=duration_days)
|
||||
|
||||
user.is_member = True
|
||||
user.member_type = plan_id
|
||||
user.member_expires_at = expires_at
|
||||
user.subscription_plan = plan_id if plan_id in _LEGACY_DURATION else "yearly"
|
||||
user.subscription_status = "active"
|
||||
user.subscription_expires_at = expires_at
|
||||
user.max_projects = -1
|
||||
user.max_storage_gb = 100
|
||||
|
||||
# 写 admin subscriptions 表(后台可见;UUID 套餐才存在真实 plan_id)
|
||||
admin_plan = db.query(PlanModel).filter(PlanModel.id == plan_id).first()
|
||||
if admin_plan is not None:
|
||||
record = SubscriptionRecordModel(
|
||||
id=uuid.uuid4().hex,
|
||||
user_id=user_id,
|
||||
plan_id=admin_plan.id,
|
||||
status="active",
|
||||
billing_cycle="yearly" if duration_days >= 365 else "monthly",
|
||||
start_date=now,
|
||||
end_date=expires_at,
|
||||
)
|
||||
db.add(record)
|
||||
|
||||
# 发积分:优先 quotas.monthly_credits,年卡一期按一年一次性发放
|
||||
monthly_credits = 0
|
||||
level_id = admin_plan.level_id
|
||||
if level_id:
|
||||
from sqlalchemy import text as sa_text
|
||||
|
||||
row = db.execute(
|
||||
sa_text("SELECT monthly_credits FROM membership_levels WHERE id = :lid"),
|
||||
{"lid": level_id},
|
||||
).first()
|
||||
if row:
|
||||
monthly_credits = int(row[0] or 0)
|
||||
|
||||
# 年卡:一次性发放 12 个月积分(任务9简化方案)
|
||||
grant = monthly_credits * (12 if duration_days >= 365 else 1)
|
||||
if grant > 0:
|
||||
self.points_service.add_points(
|
||||
user_id=user_id,
|
||||
amount=float(grant),
|
||||
source=f"membership:{plan_id}",
|
||||
db=db,
|
||||
description=f"会员开通赠积分({admin_plan.name})",
|
||||
ref_id=order_id,
|
||||
)
|
||||
|
||||
# ──────────────────────── 订单查询 ────────────────────────
|
||||
|
||||
def list_orders(
|
||||
self,
|
||||
*,
|
||||
user_id: str,
|
||||
order_type: str | None,
|
||||
page: int,
|
||||
page_size: int,
|
||||
db: Session,
|
||||
) -> dict[str, Any]:
|
||||
from packages.adapters.sqlalchemy_impl.models import PointsOrderModel
|
||||
|
||||
query = db.query(PointsOrderModel).filter(PointsOrderModel.user_id == user_id)
|
||||
if order_type:
|
||||
query = query.filter(PointsOrderModel.order_type == order_type)
|
||||
total = query.count()
|
||||
rows = query.order_by(PointsOrderModel.created_at.desc()).offset((page - 1) * page_size).limit(page_size).all()
|
||||
return {
|
||||
"total": total,
|
||||
"page": page,
|
||||
"page_size": page_size,
|
||||
"items": [self._order_dict(o) for o in rows],
|
||||
}
|
||||
|
||||
def get_order(self, *, user_id: str, order_id: str, db: Session) -> dict[str, Any] | None:
|
||||
from packages.adapters.sqlalchemy_impl.models import PointsOrderModel
|
||||
|
||||
order = (
|
||||
db.query(PointsOrderModel)
|
||||
.filter(PointsOrderModel.id == order_id, PointsOrderModel.user_id == user_id)
|
||||
.first()
|
||||
)
|
||||
return self._order_dict(order) if order else None
|
||||
|
||||
@staticmethod
|
||||
def _order_dict(order: Any) -> dict[str, Any]:
|
||||
return {
|
||||
"id": order.id,
|
||||
"order_type": order.order_type,
|
||||
"product_code": order.product_code,
|
||||
"product_name": order.product_name,
|
||||
"amount_cents": order.amount_cents,
|
||||
"points_amount": order.points_amount,
|
||||
"status": order.status,
|
||||
"payment_method": order.payment_method,
|
||||
"prepay_id": order.prepay_id,
|
||||
"paid_at": order.paid_at.isoformat() if order.paid_at else None,
|
||||
"expire_at": order.expire_at.isoformat() if order.expire_at else None,
|
||||
"created_at": order.created_at.isoformat() if order.created_at else None,
|
||||
}
|
||||
|
||||
|
||||
def _json_loads(body: bytes) -> dict[str, Any]:
|
||||
import json
|
||||
|
||||
return json.loads(body.decode("utf-8"))
|
||||
+28
-56
@@ -90,19 +90,42 @@ class SharedSettings(BaseSettings):
|
||||
|
||||
# ── 豆包大模型(火山引擎方舟) ────────────────────────────────────────
|
||||
doubao_api_key: str = ""
|
||||
doubao_model: str = "doubao-seed-1-6-250615" # 推理模型(通用兜底)
|
||||
doubao_fast_model: str = "doubao-1-5-pro-32k-250115" # 快速结构化输出模型(编导脚本/意图解析/审核)
|
||||
doubao_model: str = "doubao-seed-2-1-pro-260915" # 推理模型(Seed 2.1 Pro,深度思考+多模态;原 seed-1-6 已下线)
|
||||
doubao_fast_model: str = (
|
||||
"doubao-seed-2-1-lite-260915" # 快速模型(Seed 2.1 Lite,高 RPM,编导/审核/VLM lite;原 1-5-pro-32k 已 Retiring)
|
||||
)
|
||||
doubao_base_url: str = "https://ark.cn-beijing.volces.com/api/v3"
|
||||
doubao_timeout: int = 30
|
||||
doubao_max_retries: int = 2
|
||||
doubao_vision_model: str = "doubao-1-5-vision-pro-250328" # 高精度视觉(备用)
|
||||
doubao_vision_lite_model: str = "doubao-1-5-vision-lite-250315" # 快速视觉(商品识别默认,速度优先)
|
||||
doubao_vision_model: str = (
|
||||
"doubao-seed-2-1-pro-260915" # 高精度视觉(Seed 2.1 Pro 原生多模态;原 vision-pro-250328 已下线)
|
||||
)
|
||||
doubao_vision_lite_model: str = (
|
||||
"doubao-seed-2-1-lite-260915" # 快速视觉(Seed 2.1 Lite 原生多模态;原 vision-lite-250315 不可用)
|
||||
)
|
||||
doubao_vision_use_lite: bool = True # viral-video 图片分析默认用 lite 提速
|
||||
doubao_embedding_model: str = "doubao-embedding-large-text-240915"
|
||||
doubao_embedding_model: str = "doubao-embedding-vision-251215" # 多模态向量化(原 large-text-240915 已 Retiring)
|
||||
doubao_video_model: str = "doubao-seedance-2-5-260628"
|
||||
doubao_video_timeout: int = 600 # 视频生成轮询总超时(秒)
|
||||
doubao_video_poll_interval: int = 10 # 轮询间隔(秒)
|
||||
|
||||
# ── DashScope (阿里云百炼 Wan 3.0 等) ─────────────────────────────────
|
||||
dashscope_api_key: str = ""
|
||||
dashscope_base_url: str = "https://dashscope.aliyuncs.com/api/v1"
|
||||
dashscope_video_timeout: int = 900 # Wan 视频任务轮询总超时(秒)
|
||||
dashscope_video_poll_interval: int = 10
|
||||
|
||||
# ── 即梦(Jimeng)视觉 API —— 火山引擎 cvtob ──────────────────────────
|
||||
# #2169: 真人参考图在方舟 Seedance 走 B 端审核会被 50411 拦截,
|
||||
# 即梦走 C 端审核链路,普通真人照片可过审,作为参考图场景兜底通道。
|
||||
# 鉴权:AK/SK V4 签名(Region=cn-north-1, Service=cv)
|
||||
jimeng_ak: str = ""
|
||||
jimeng_sk: str = ""
|
||||
jimeng_base_url: str = "https://visual.volcengineapi.com"
|
||||
jimeng_req_key: str = "jimeng_i2v_first_v30" # 视频3.0 720P 首帧图生视频(1张图,稳定版;1080P 标注下线中)
|
||||
jimeng_video_timeout: int = 600 # 即梦轮询总超时(秒)
|
||||
jimeng_video_poll_interval: int = 5 # 轮询间隔(秒)
|
||||
|
||||
# ── MediaKit (火山引擎 AI 媒体工具) ──────────────────────────────────
|
||||
mediakit_api_key: str = ""
|
||||
mediakit_base_url: str = "https://mediakit.cn-beijing.volces.com/api/v1"
|
||||
@@ -138,57 +161,6 @@ class SharedSettings(BaseSettings):
|
||||
self.credits_enabled = bool(value)
|
||||
self.points_enabled_compat = False
|
||||
|
||||
# ── 微信支付(V3 API)──────────────────────────────────────────────
|
||||
# 未配置时支付下单接口返回明确错误,不会发出真实请求。
|
||||
# staging 可先配沙箱参数;通过环境变量注入。
|
||||
wechat_pay_enabled: bool = Field(
|
||||
default=False,
|
||||
validation_alias=AliasChoices("WECHAT_PAY_ENABLED", "wechat_pay_enabled"),
|
||||
)
|
||||
# 公众号/小程序 AppID
|
||||
wechat_appid: str = Field(
|
||||
default="",
|
||||
validation_alias=AliasChoices("WECHAT_APPID", "wechat_appid"),
|
||||
)
|
||||
# 微信支付商户号
|
||||
wechat_mch_id: str = Field(
|
||||
default="",
|
||||
validation_alias=AliasChoices("WECHAT_MCH_ID", "wechat_mch_id"),
|
||||
)
|
||||
# APIv3 密钥(32字符,用于回调解密)
|
||||
wechat_api_v3_key: str = Field(
|
||||
default="",
|
||||
validation_alias=AliasChoices("WECHAT_API_V3_KEY", "wechat_api_v3_key"),
|
||||
)
|
||||
# 商户API证书序列号
|
||||
wechat_cert_serial: str = Field(
|
||||
default="",
|
||||
validation_alias=AliasChoices("WECHAT_CERT_SERIAL", "wechat_cert_serial"),
|
||||
)
|
||||
# 商户API私钥(PEM 内容;也支持填写文件绝对路径)
|
||||
wechat_private_key: str = Field(
|
||||
default="",
|
||||
validation_alias=AliasChoices("WECHAT_PRIVATE_KEY", "wechat_private_key"),
|
||||
)
|
||||
# 支付回调通知地址(完整 https URL)
|
||||
wechat_notify_url: str = Field(
|
||||
default="",
|
||||
validation_alias=AliasChoices("WECHAT_NOTIFY_URL", "wechat_notify_url"),
|
||||
)
|
||||
|
||||
@property
|
||||
def wechat_pay_configured(self) -> bool:
|
||||
"""微信支付必要配置是否齐全(启用 + 六要素非空)。"""
|
||||
return bool(
|
||||
self.wechat_pay_enabled
|
||||
and self.wechat_appid
|
||||
and self.wechat_mch_id
|
||||
and self.wechat_api_v3_key
|
||||
and self.wechat_cert_serial
|
||||
and self.wechat_private_key
|
||||
and self.wechat_notify_url
|
||||
)
|
||||
|
||||
# ── GPU MuseTalk 反向轮询 Worker ────────────────────────────────────
|
||||
# Worker 用这个长期 Token 鉴权(不是用户 JWT)。多 Worker 共用同一个 Token;
|
||||
# worker_id 用于区分具体机器。生产必须配置;development 留空会跳过校验。
|
||||
|
||||
@@ -60,6 +60,11 @@ class User:
|
||||
# 资料是否已完善(微信新用户首次设置昵称后置 True;邮箱注册默认 True)
|
||||
profile_completed: bool = True
|
||||
|
||||
# 会员字段 (#1895):与 users 表列对应
|
||||
is_member: bool = False
|
||||
member_type: str | None = None
|
||||
member_expires_at: datetime | None = None
|
||||
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
|
||||
+155
-13
@@ -10,7 +10,9 @@ from __future__ import annotations
|
||||
import math
|
||||
|
||||
# ============ 爆款视频动态定价 (#2151) ============
|
||||
# key = (model_id, resolution, has_video_input),单位:元/百万token
|
||||
# key = (model_id, resolution, has_video_input),单位:
|
||||
# - billing_mode=token: 元/百万tokens(输出)
|
||||
# - billing_mode=per_second: 元/秒(视频时长)
|
||||
VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = {
|
||||
("seedance-2.5", "480p", False): 70.0,
|
||||
("seedance-2.5", "720p", False): 70.0,
|
||||
@@ -21,6 +23,16 @@ VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = {
|
||||
("seedance-2.0", "480p", False): 46.0,
|
||||
("seedance-2.0", "720p", False): 46.0,
|
||||
("seedance-2.0", "1080p", False): 51.0,
|
||||
("seedance-2.0", "4k", False): 80.0,
|
||||
("seedance-2.0-fast", "480p", False): 28.0,
|
||||
("seedance-2.0-fast", "720p", False): 28.0,
|
||||
("seedance-2.0-mini", "480p", False): 9.2,
|
||||
("seedance-2.0-mini", "720p", False): 9.2,
|
||||
("wan-3.0", "480p", False): 0.3,
|
||||
("wan-3.0", "720p", False): 0.6,
|
||||
("wan-3.0", "1080p", False): 1.2,
|
||||
# #2169: 即梦(Jimeng)视频3.0 720P 首帧图生视频,0.28 元/秒(C 端审核,真人可过)
|
||||
("jimeng-3.0", "720p", False): 0.28,
|
||||
}
|
||||
|
||||
# 固定成本(元):VLM 分析 + LLM 文案 + TTS + OSS + 服务器
|
||||
@@ -49,7 +61,10 @@ _RESOLUTION_ALIASES: dict[str, str] = {
|
||||
"fhd": "1080p",
|
||||
}
|
||||
# 分辨率 -> 短边像素数(p 值代表短边,不是 height)
|
||||
_RESOLUTION_SHORT_SIDE: dict[str, int] = {"480p": 480, "720p": 720, "1080p": 1080}
|
||||
_RESOLUTION_SHORT_SIDE: dict[str, int] = {"480p": 480, "720p": 720, "1080p": 1080, "4k": 2160}
|
||||
_RESOLUTION_ALIASES["4k"] = "4k"
|
||||
_RESOLUTION_ALIASES["2160p"] = "4k"
|
||||
_RESOLUTION_ALIASES["uhd"] = "4k"
|
||||
|
||||
|
||||
def resolve_video_dimensions(resolution: str, ratio: str) -> tuple[int, int]:
|
||||
@@ -81,18 +96,133 @@ def resolve_video_dimensions(resolution: str, ratio: str) -> tuple[int, int]:
|
||||
return int(w), int(h)
|
||||
|
||||
|
||||
def _match_model_prefix(model: str) -> str:
|
||||
"""匹配 model 前缀。"""
|
||||
m = (model or "").strip().lower()
|
||||
for prefix in ("seedance-2.5", "seedance-2.0"):
|
||||
if m.startswith(prefix):
|
||||
return prefix
|
||||
# ── 爆款视频多模型元数据 (#2159) ──────────────────────────────────────
|
||||
VIRAL_VIDEO_MODEL_CONFIG: dict[str, dict] = {
|
||||
"seedance-2.5": {
|
||||
"key": "seedance-2.5",
|
||||
"display_name": "Seedance 2.5 — 最新最强",
|
||||
"model_id": "doubao-seedance-2-5-260628",
|
||||
"provider": "doubao",
|
||||
"supports_audio": True,
|
||||
"supported_resolutions": ["480p", "720p", "1080p"],
|
||||
"max_duration": 30,
|
||||
"billing_mode": "token",
|
||||
"is_default": True,
|
||||
},
|
||||
"seedance-2.0": {
|
||||
"key": "seedance-2.0",
|
||||
"display_name": "Seedance 2.0 — 正式首选",
|
||||
"model_id": "doubao-seedance-2-0-260128",
|
||||
"provider": "doubao",
|
||||
"supports_audio": True,
|
||||
"supported_resolutions": ["480p", "720p", "1080p", "4k"],
|
||||
"max_duration": 15,
|
||||
"billing_mode": "token",
|
||||
"is_default": False,
|
||||
},
|
||||
"seedance-2.0-fast": {
|
||||
"key": "seedance-2.0-fast",
|
||||
"display_name": "Seedance 2.0 Fast — 快速低成本",
|
||||
"model_id": "doubao-seedance-2-0-fast-260128",
|
||||
"provider": "doubao",
|
||||
"supports_audio": True,
|
||||
"supported_resolutions": ["480p", "720p"],
|
||||
"max_duration": 15,
|
||||
"billing_mode": "token",
|
||||
"is_default": False,
|
||||
},
|
||||
"seedance-2.0-mini": {
|
||||
"key": "seedance-2.0-mini",
|
||||
"display_name": "Seedance 2.0 Mini — 低成本测试",
|
||||
"model_id": "doubao-seedance-2-0-mini-260615",
|
||||
"provider": "doubao",
|
||||
"supports_audio": True,
|
||||
"supported_resolutions": ["480p", "720p"],
|
||||
"max_duration": 15,
|
||||
"billing_mode": "token",
|
||||
"is_default": False,
|
||||
},
|
||||
"wan-3.0": {
|
||||
"key": "wan-3.0",
|
||||
"display_name": "Wan 3.0 — 通义万相(阿里云)",
|
||||
"model_id": "wan3.0-video",
|
||||
"provider": "dashscope",
|
||||
"supports_audio": True,
|
||||
"supported_resolutions": ["480p", "720p", "1080p"],
|
||||
"max_duration": 30,
|
||||
"billing_mode": "per_second",
|
||||
"is_default": False,
|
||||
},
|
||||
# #2169: 即梦视频3.0(内部兜底通道,方舟 Seedance 返回真人拦截 50411 时自动切到即梦重试,
|
||||
# 不暴露给前端让用户直接选择,但需支持计费结算)
|
||||
"jimeng-3.0": {
|
||||
"key": "jimeng-3.0",
|
||||
"display_name": "即梦3.0 — 真人图生视频(兜底)",
|
||||
"model_id": "jimeng_i2v_first_v30",
|
||||
"provider": "jimeng",
|
||||
"supports_audio": False, # 即梦返回无声视频,音频由后续 ffmpeg 合成 TTS
|
||||
"supported_resolutions": ["720p"],
|
||||
"max_duration": 10, # 即梦 i2v 首帧最长 10s(frames=241)
|
||||
"billing_mode": "per_second",
|
||||
"is_default": False,
|
||||
"_internal_fallback_only": True, # 标记:不对外暴露到模型选择列表
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def get_viral_video_model_config(model_key: str | None) -> dict:
|
||||
"""获取模型配置,未知 key 回落到默认 seedance-2.5。"""
|
||||
key = (model_key or "").strip().lower()
|
||||
if key and key in VIRAL_VIDEO_MODEL_CONFIG:
|
||||
return VIRAL_VIDEO_MODEL_CONFIG[key]
|
||||
return VIRAL_VIDEO_MODEL_CONFIG["seedance-2.5"]
|
||||
|
||||
|
||||
def list_viral_video_models(
|
||||
include_placeholder: bool = False,
|
||||
dashscope_available: bool = False,
|
||||
) -> list[dict]:
|
||||
"""返回前端可用的模型列表(供 GET /api/v1/viral-video/models 端点用)。"""
|
||||
out: list[dict] = []
|
||||
for _k, cfg in VIRAL_VIDEO_MODEL_CONFIG.items():
|
||||
if cfg.get("_placeholder") and not include_placeholder:
|
||||
continue
|
||||
if cfg.get("provider") == "dashscope" and not dashscope_available:
|
||||
continue
|
||||
# #2169: 即梦是内部兜底通道,不在前端模型列表展示
|
||||
if cfg.get("_internal_fallback_only"):
|
||||
continue
|
||||
out.append(
|
||||
{
|
||||
"key": cfg["key"],
|
||||
"display_name": cfg["display_name"],
|
||||
"supports_audio": bool(cfg.get("supports_audio", True)),
|
||||
"supported_resolutions": list(cfg.get("supported_resolutions", ["720p"])),
|
||||
"max_duration": int(cfg.get("max_duration", 15)),
|
||||
"billing_mode": cfg.get("billing_mode", "token"),
|
||||
"is_default": bool(cfg.get("is_default", False)),
|
||||
}
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
def _match_model_prefix(model: str | None) -> str:
|
||||
"""匹配 model key(支持全部内部别名,未知回落到 seedance-2.5)。
|
||||
|
||||
按 key 长度从长到短匹配,避免 "seedance-2.0-fast" 被 "seedance-2.0" 前缀命中。
|
||||
"""
|
||||
mm = (model or "").strip().lower()
|
||||
for k in sorted(VIRAL_VIDEO_MODEL_CONFIG.keys(), key=len, reverse=True):
|
||||
if mm == k or mm.startswith(k):
|
||||
return k
|
||||
return "seedance-2.5"
|
||||
|
||||
|
||||
def _infer_resolution_key(width: int, height: int) -> str:
|
||||
"""从实际 (width, height) 用短边推断 resolution key。"""
|
||||
short = min(int(width or 720), int(height or 720))
|
||||
if short >= 1900:
|
||||
return "4k"
|
||||
if short >= 1000:
|
||||
return "1080p"
|
||||
if short >= 650:
|
||||
@@ -128,18 +258,26 @@ def calculate_viral_video_credits_with_breakdown(
|
||||
effective_fps = int(fps or VIRAL_VIDEO_FPS)
|
||||
|
||||
prefix = _match_model_prefix(model)
|
||||
cfg = get_viral_video_model_config(prefix)
|
||||
res_key = _infer_resolution_key(w, h)
|
||||
billing = cfg.get("billing_mode", "token")
|
||||
key = (prefix, res_key, bool(has_video_input))
|
||||
price = VIRAL_VIDEO_MODEL_PRICES.get(key)
|
||||
if price is None:
|
||||
price = VIRAL_VIDEO_MODEL_PRICES.get(("seedance-2.5", res_key, False), 70.0)
|
||||
if actual_tokens is not None and actual_tokens > 0:
|
||||
tokens = float(actual_tokens)
|
||||
dur = max(1, int(duration_seconds or 15))
|
||||
if billing == "per_second":
|
||||
tokens = 0.0
|
||||
video_cost = dur * float(price)
|
||||
billing_unit = "second"
|
||||
else:
|
||||
dur = max(1, int(duration_seconds or 15))
|
||||
tokens = dur * w * h * effective_fps / 1024.0
|
||||
if actual_tokens is not None and actual_tokens > 0:
|
||||
tokens = float(actual_tokens)
|
||||
else:
|
||||
tokens = dur * w * h * effective_fps / 1024.0
|
||||
video_cost = tokens / 1_000_000.0 * float(price)
|
||||
billing_unit = "token"
|
||||
|
||||
video_cost = tokens / 1_000_000.0 * float(price)
|
||||
total = (video_cost + VIRAL_VIDEO_FIXED_COST) * VIRAL_VIDEO_PROFIT_MULTIPLIER
|
||||
credits = round(float(total), 2)
|
||||
breakdown = {
|
||||
@@ -148,9 +286,13 @@ def calculate_viral_video_credits_with_breakdown(
|
||||
"fixed_cost": float(VIRAL_VIDEO_FIXED_COST),
|
||||
"profit_multiplier": float(VIRAL_VIDEO_PROFIT_MULTIPLIER),
|
||||
"model_price": float(price),
|
||||
"model_key": prefix,
|
||||
"billing_mode": billing,
|
||||
"billing_unit": billing_unit,
|
||||
"width": int(w),
|
||||
"height": int(h),
|
||||
"fps": int(effective_fps),
|
||||
"duration": dur,
|
||||
}
|
||||
return credits, breakdown
|
||||
|
||||
|
||||
+458
-23
@@ -22,9 +22,153 @@ import httpx
|
||||
|
||||
from packages.shared.config import get_shared_settings
|
||||
|
||||
# 网络/超时类异常父类集合:覆盖 Timeout/Connect/Network/ReadTimeout/WriteTimeout/PoolTimeout
|
||||
_HTTP_NETWORK_ERRORS = ()
|
||||
try:
|
||||
_HTTP_NETWORK_ERRORS = (httpx.TimeoutException, httpx.NetworkError)
|
||||
except Exception:
|
||||
_HTTP_NETWORK_ERRORS = (Exception,)
|
||||
|
||||
_HTTP_STATUS_ERROR = httpx.HTTPStatusError if hasattr(httpx, "HTTPStatusError") else Exception
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# 视频模型 ID 解析逻辑(#2159 多模型支持,#2169 接入即梦)。
|
||||
# 内部使用简短别名(seedance-2.5 / seedance-2.0 / wan-3.0 / jimeng-3.0 等)做 PRICING key;
|
||||
# 实际调用时按 VIRAL_VIDEO_MODEL_CONFIG(domain/points_rules.py)的 provider/model_id 分发。
|
||||
# - provider=doubao → 火山方舟 Seedance
|
||||
# - provider=dashscope → 阿里云 DashScope(Wan 系列)
|
||||
# - provider=jimeng → 火山引擎即梦 cvtob(jimeng_i2v_first_v30,真人参考图走 C 端审核)
|
||||
|
||||
|
||||
def _resolve_video_provider_and_id(model: str | None) -> tuple[str, str, dict]:
|
||||
"""把内部 model key 解析成 (provider, model_id, cfg)。
|
||||
|
||||
- provider: "doubao" | "dashscope" | "jimeng"
|
||||
- model_id: 对应 API 的真实模型 ID
|
||||
- cfg: VIRAL_VIDEO_MODEL_CONFIG 条目
|
||||
未识别或空值回落到默认 seedance-2.5。已经是 doubao-/ep- 开头的完整 ID 视为 doubao provider;
|
||||
"jimeng" 开头视为 jimeng provider(内部兜底,不暴露给前端)。
|
||||
"""
|
||||
from packages.domain.points_rules import get_viral_video_model_config
|
||||
|
||||
settings = get_shared_settings()
|
||||
default_id = getattr(settings, "doubao_video_model", None) or "doubao-seedance-2-5-260628"
|
||||
m = (model or "").strip()
|
||||
if not m:
|
||||
cfg = get_viral_video_model_config("seedance-2.5")
|
||||
return "doubao", default_id, cfg
|
||||
# 已经是 doubao-/ep- 开头:直接透传,默认视为 doubao provider
|
||||
if m.startswith("doubao-") or m.startswith("ep-"):
|
||||
return "doubao", m, {"provider": "doubao", "model_id": m, "supports_audio": True}
|
||||
# 显式 jimeng 关键字:路由到即梦(内部兜底通道使用)
|
||||
if m.startswith("jimeng"):
|
||||
cfg = get_viral_video_model_config("jimeng-3.0")
|
||||
return "jimeng", cfg.get("model_id", "jimeng_i2v_first_v30"), cfg
|
||||
# 别名 → 从 domain config 查
|
||||
cfg = get_viral_video_model_config(m)
|
||||
provider = cfg.get("provider", "doubao")
|
||||
resolved_id = cfg.get("model_id", "")
|
||||
if not resolved_id:
|
||||
logger.warning("[ai_client] model %r 无 model_id,回落到默认 %s", m, default_id)
|
||||
return "doubao", default_id, cfg
|
||||
return provider, resolved_id, cfg
|
||||
|
||||
|
||||
def _resolve_video_model_id(model: str | None) -> str:
|
||||
"""兼容旧调用:只返回 doubao model_id。wan/dashscope 调用方应直接用 _resolve_video_provider_and_id。"""
|
||||
_provider, mid, _cfg = _resolve_video_provider_and_id(model)
|
||||
return mid
|
||||
|
||||
|
||||
# ── 视频错误分类(给前端/用户展示友好提示)────────────────────────────
|
||||
|
||||
|
||||
def _classify_video_error(status_code: int, body: str, err: Exception | None) -> tuple[str, str]:
|
||||
"""根据 HTTP 状态码和响应 body 判断错误类型。
|
||||
|
||||
返回 (error_code, user_message):
|
||||
- error_code: 机器可读的错误码("portrait_intercept" / "quota_exceeded" / "model_not_found"
|
||||
/ "invalid_param" / "auth_error" / "rate_limit" / "network_error" / "task_failed" / "unknown")
|
||||
- user_message: 给用户看的中文提示
|
||||
"""
|
||||
body_lower = (body or "").lower()
|
||||
code_in_body = ""
|
||||
msg_in_body = ""
|
||||
try:
|
||||
import json as _json
|
||||
|
||||
parsed = _json.loads(body or "{}")
|
||||
if isinstance(parsed, dict):
|
||||
err_obj = parsed.get("error") or {}
|
||||
if isinstance(err_obj, dict):
|
||||
code_in_body = str(err_obj.get("code", "") or "")
|
||||
msg_in_body = str(err_obj.get("message", "") or err_obj.get("msg", "") or "")
|
||||
else:
|
||||
msg_in_body = str(parsed.get("message", "") or "")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 真人肖像/内容安全拦截
|
||||
if (
|
||||
status_code == 400
|
||||
and any(
|
||||
kw in body_lower
|
||||
for kw in ("portrait", "real_face", "human_face", "真人", "肖像", "人脸", "privacy", "real person", "face")
|
||||
)
|
||||
) or (
|
||||
"content" in body_lower
|
||||
and ("risk" in body_lower or "block" in body_lower or "reject" in body_lower)
|
||||
and status_code == 400
|
||||
):
|
||||
return (
|
||||
"portrait_intercept",
|
||||
"参考素材包含真人照片被安全策略拦截,AI视频模型暂不支持上传真人照片作为参考图,请移除真人图片后重试。",
|
||||
)
|
||||
|
||||
# 配额/计费问题
|
||||
if status_code in (402, 429) or any(
|
||||
kw in body_lower for kw in ("quota", "billing", "insufficient", "欠费", "余额", "限流", "rate limit")
|
||||
):
|
||||
if "rate" in body_lower or status_code == 429:
|
||||
return "rate_limit", "视频生成服务当前繁忙(限流),请稍等1-2分钟后重试。"
|
||||
return "quota_exceeded", "视频生成服务配额不足,请联系管理员充值或稍后重试。"
|
||||
|
||||
# 模型/Endpoint 不存在
|
||||
if status_code == 404 or any(
|
||||
kw in body_lower for kw in ("model not found", "endpoint not found", "不存在", "not found", "model_not_exist")
|
||||
):
|
||||
return "model_not_found", f"视频模型未开通或模型ID无效({code_in_body or ''}),请联系管理员。"
|
||||
|
||||
# 鉴权失败
|
||||
if status_code in (401, 403):
|
||||
return "auth_error", "视频生成服务鉴权失败(API Key无效或过期),请联系管理员。"
|
||||
|
||||
# 任务本身失败(轮询阶段拿到 status=failed)
|
||||
if err and "task failed" in str(err).lower():
|
||||
detail = msg_in_body or str(err)[:200]
|
||||
# 失败原因里再细分真人拦截
|
||||
if any(kw in detail.lower() for kw in ("portrait", "真人", "肖像", "人脸", "content_risk")):
|
||||
return (
|
||||
"portrait_intercept",
|
||||
"视频内容被安全策略拦截(疑似包含真人肖像),请更换参考图或调整文案后重试。",
|
||||
)
|
||||
return "task_failed", f"视频生成失败:{detail}"
|
||||
|
||||
# 参数错误
|
||||
if status_code == 400:
|
||||
return "invalid_param", f"视频生成参数错误:{msg_in_body or body[:200]}"
|
||||
|
||||
# 网络/连接问题
|
||||
if status_code == 0:
|
||||
return "network_error", "视频生成服务连接失败(网络超时),请稍后重试。"
|
||||
|
||||
# 默认
|
||||
detail = msg_in_body or (str(err) if err else "") or body[:200]
|
||||
return "unknown", f"视频生成失败(HTTP {status_code}):{detail}"
|
||||
|
||||
|
||||
class DoubaoClient:
|
||||
"""豆包大模型 API 客户端.
|
||||
|
||||
@@ -42,6 +186,9 @@ class DoubaoClient:
|
||||
self.vision_model: str = settings.doubao_vision_model
|
||||
self.vision_lite_model: str = settings.doubao_vision_lite_model
|
||||
self.fast_model: str = settings.doubao_fast_model
|
||||
self.embedding_model: str = settings.doubao_embedding_model
|
||||
# 最近一次视频生成的详细错误(error_code + user_message + raw detail),供上层读取后展示给用户
|
||||
self.last_video_error: dict = {}
|
||||
|
||||
def embed_text(self, text: str, timeout: int | None = None) -> list[float] | None:
|
||||
"""调用豆包文本 Embedding API,返回浮点向量;失败返回 None。"""
|
||||
@@ -54,7 +201,7 @@ class DoubaoClient:
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
payload: dict[str, Any] = {
|
||||
"model": getattr(self, "embedding_model", None) or "doubao-embedding-large-text-240915",
|
||||
"model": self.embedding_model,
|
||||
"input": text.strip(),
|
||||
"encoding_format": "float",
|
||||
}
|
||||
@@ -265,7 +412,8 @@ class DoubaoClient:
|
||||
"""调用 Seedance 2.5 生视频(异步任务→轮询→下载)。
|
||||
|
||||
成功返回 {"video_path": str, "usage": dict | None},失败返回 None。
|
||||
usage 是 Seedance 返回的计费信息(含 completion_tokens)。
|
||||
失败时把详细错误信息(HTTP状态码、响应 body、分类后的用户提示)写入 self.last_video_error,
|
||||
上层可通过 get_last_video_error() 读取并展示给用户,不再笼统显示"返回为空"。
|
||||
|
||||
【v1.6.1 修复】严格按官方 content 数组协议构造请求:
|
||||
- 所有参考(图/音/视)必须放进 content 数组并带 role 字段,不能放顶层 reference_audios/reference_videos(非官方字段,会被忽略或导致异常)。
|
||||
@@ -273,17 +421,83 @@ class DoubaoClient:
|
||||
判定:传了参考音频/视频或 ≥1 张多参考图时,走 omni_reference(首张图 role=reference_image);纯首帧无参考时走 first_frame(ratio 强制 adaptive)。
|
||||
- 创建任务若因 ratio 报错(HTTP 400),自动回退到 ratio=adaptive 重试一次。
|
||||
"""
|
||||
# 每次调用前清空上次错误
|
||||
self.last_video_error = {}
|
||||
|
||||
if not self.is_available:
|
||||
self.last_video_error = {
|
||||
"error_code": "auth_error",
|
||||
"user_message": "视频生成服务未配置(API Key 缺失),请联系管理员。",
|
||||
"status_code": 0,
|
||||
"detail": "DoubaoClient not available (api_key empty)",
|
||||
}
|
||||
return None
|
||||
if not prompt or not prompt.strip():
|
||||
self.last_video_error = {
|
||||
"error_code": "invalid_param",
|
||||
"user_message": "视频生成提示词不能为空。",
|
||||
"status_code": 0,
|
||||
"detail": "empty prompt",
|
||||
}
|
||||
return None
|
||||
|
||||
settings = get_shared_settings()
|
||||
poll_interval = getattr(settings, "doubao_video_poll_interval", 10) or 10
|
||||
# 收紧总超时:轮询 8min + 下载 2min = 最长 ~10min,防止出现 20min 卡死
|
||||
total_timeout = getattr(settings, "doubao_video_timeout", 480) or 480
|
||||
default_video_model = getattr(settings, "doubao_video_model", None) or "doubao-seedance-2-5-260628"
|
||||
video_model = model or default_video_model
|
||||
# 内部 key → (provider, 实际模型 ID, cfg),按 provider 分发
|
||||
provider, video_model, model_cfg = _resolve_video_provider_and_id(model)
|
||||
if provider == "dashscope":
|
||||
from packages.shared.dashscope_client import get_dashscope_client
|
||||
|
||||
ds = get_dashscope_client()
|
||||
if ds is None:
|
||||
err_msg = "DashScope client 不可用(未配置 DASHSCOPE_API_KEY)"
|
||||
logger.error("%s, video_model=%s", err_msg, model)
|
||||
self.last_video_error = {
|
||||
"error_code": "auth_error",
|
||||
"user_message": "Wan 3.0 视频模型未配置 API Key,请联系管理员。",
|
||||
"status_code": 0,
|
||||
"detail": err_msg,
|
||||
}
|
||||
return None
|
||||
try:
|
||||
# DashScope 客户端也设置 last_video_error 语义(如果它支持)
|
||||
if hasattr(ds, "last_video_error"):
|
||||
ds.last_video_error = {}
|
||||
result = ds.video_generation(
|
||||
prompt=prompt,
|
||||
image_url=image_url,
|
||||
duration=duration,
|
||||
ratio=ratio,
|
||||
resolution=resolution,
|
||||
output_dir=output_dir,
|
||||
model=video_model,
|
||||
)
|
||||
if not result and hasattr(ds, "last_video_error") and ds.last_video_error:
|
||||
self.last_video_error = dict(ds.last_video_error)
|
||||
return result
|
||||
except Exception as de:
|
||||
logger.error("DashScope video_generation 异常: %s", de, exc_info=True)
|
||||
self.last_video_error = {
|
||||
"error_code": "unknown",
|
||||
"user_message": f"Wan 3.0 视频生成异常:{de!s}"[:200],
|
||||
"status_code": 0,
|
||||
"detail": str(de),
|
||||
}
|
||||
return None
|
||||
|
||||
if provider == "jimeng":
|
||||
# #2169: 即梦 cvtob(jimeng_i2v_first_v30)— 真人参考图兜底通道
|
||||
return self._call_jimeng_video_generation(
|
||||
prompt=prompt,
|
||||
image_url=image_url,
|
||||
duration=duration,
|
||||
ratio=ratio,
|
||||
resolution=resolution,
|
||||
output_dir=output_dir,
|
||||
generate_audio=generate_audio,
|
||||
)
|
||||
|
||||
ref_audios = [u for u in (reference_audios or [])[:10] if u and isinstance(u, str)]
|
||||
ref_videos = [u for u in (reference_videos or [])[:3] if u and isinstance(u, str)]
|
||||
@@ -308,7 +522,7 @@ class DoubaoClient:
|
||||
}
|
||||
)
|
||||
else:
|
||||
# 纯首帧:不带 role,服务端识别为 first_frame(或显式 role=first_frame)
|
||||
# 纯首帧:显式 role=first_frame
|
||||
content.append(
|
||||
{
|
||||
"type": "image_url",
|
||||
@@ -349,22 +563,38 @@ class DoubaoClient:
|
||||
len(ref_audios),
|
||||
len(ref_videos),
|
||||
)
|
||||
# 打印完整 payload 便于排查(截断 prompt)
|
||||
debug_payload = dict(create_payload)
|
||||
if "content" in debug_payload:
|
||||
dbg_content = []
|
||||
for item in debug_payload["content"]:
|
||||
item_copy = dict(item)
|
||||
if item_copy.get("type") == "text" and isinstance(item_copy.get("text"), str):
|
||||
item_copy["text"] = item_copy["text"][:200] + ("..." if len(item_copy["text"]) > 200 else "")
|
||||
dbg_content.append(item_copy)
|
||||
debug_payload["content"] = dbg_content
|
||||
logger.info("Seedance 创建任务 payload: %s", json_safe_dumps(debug_payload))
|
||||
|
||||
def _do_create(payload: dict) -> tuple[str | None, Exception | None, int, str]:
|
||||
"""返回 (task_id, last_err, status_code, body_text)。"""
|
||||
last_err: Exception | None = None
|
||||
last_sc = 0
|
||||
last_body = ""
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
resp = httpx.post(create_url, headers=headers, json=payload, timeout=self.timeout)
|
||||
sc = int(getattr(resp, "status_code", 0) or 0)
|
||||
body = (getattr(resp, "text", "") or "")[:1500]
|
||||
body = (getattr(resp, "text", "") or "")[:2000]
|
||||
last_sc = sc
|
||||
last_body = body
|
||||
if sc >= 400:
|
||||
logger.error("Seedance 创建任务 HTTP %d: body=%s", sc, body)
|
||||
try:
|
||||
resp.raise_for_status()
|
||||
except Exception as ee:
|
||||
last_err = ee
|
||||
if attempt < self.max_retries:
|
||||
if attempt < self.max_retries and sc >= 500:
|
||||
# 仅 5xx 重试,4xx 不重试(参数/鉴权/配额错误重试无意义)
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
return None, last_err, sc, body
|
||||
@@ -373,9 +603,19 @@ class DoubaoClient:
|
||||
if tid:
|
||||
return tid, None, sc, body
|
||||
last_err = RuntimeError(f"create ok but no id: {str(data)[:300]}")
|
||||
except _HTTP_NETWORK_ERRORS as ne:
|
||||
last_err = ne
|
||||
last_sc = 0
|
||||
last_body = f"network error: {ne}"
|
||||
logger.warning(
|
||||
"Seedance 创建网络异常(%s),重试 %d/%d", type(ne).__name__, attempt + 1, self.max_retries + 1
|
||||
)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
except Exception as e:
|
||||
last_err = e
|
||||
if attempt < self.max_retries:
|
||||
if attempt < self.max_retries and not isinstance(e, _HTTP_STATUS_ERROR):
|
||||
wait = 0.5 * (2**attempt)
|
||||
logger.warning(
|
||||
"Seedance 创建任务失败,%.1fs 后重试 (%d/%d): %s",
|
||||
@@ -385,7 +625,7 @@ class DoubaoClient:
|
||||
e,
|
||||
)
|
||||
time.sleep(wait)
|
||||
return None, last_err, 0, ""
|
||||
return None, last_err, last_sc, last_body
|
||||
|
||||
# 第一次尝试
|
||||
task_id, last_err, sc, body = _do_create(create_payload)
|
||||
@@ -404,17 +644,53 @@ class DoubaoClient:
|
||||
logger.warning("Seedance 创建因 ratio 失败,回退 ratio=adaptive 重试")
|
||||
create_payload["ratio"] = "adaptive"
|
||||
task_id, last_err, sc2, body2 = _do_create(create_payload)
|
||||
if task_id:
|
||||
sc, body = sc2, body2
|
||||
else:
|
||||
# 保留第二次的错误信息
|
||||
sc, body = sc2, body2
|
||||
|
||||
if not task_id:
|
||||
err_code, user_msg = _classify_video_error(sc, body, last_err)
|
||||
self.last_video_error = {
|
||||
"error_code": err_code,
|
||||
"user_message": user_msg,
|
||||
"status_code": sc,
|
||||
"detail": (body or "")[:500] or (str(last_err) if last_err else ""),
|
||||
"model": video_model,
|
||||
"base_url": self.base_url,
|
||||
}
|
||||
logger.error(
|
||||
"Seedance 创建任务最终失败: model=%s base_url=%s err=%s body=%s 【排查】"
|
||||
"1) 方舟控制台已开通 doubao-seedance-2-5-260628;2) API Key 有该模型权限;"
|
||||
"3) DOUBAO_BASE_URL=https://ark.cn-beijing.volces.com/api/v3;4) 参考素材 URL 公网可访问。",
|
||||
"Seedance 创建任务最终失败: model=%s base_url=%s status=%d code=%s err=%s body=%s",
|
||||
video_model,
|
||||
self.base_url,
|
||||
sc,
|
||||
err_code,
|
||||
last_err,
|
||||
(body or "")[:500],
|
||||
)
|
||||
# #2169: 方舟返回 portrait_intercept 且有参考图 → 自动切即梦重试一次(保留首帧图)
|
||||
if err_code == "portrait_intercept" and image_url:
|
||||
logger.warning(
|
||||
"[viral-video] Seedance 真人拦截(code=%s),自动切即梦通道重试(首帧图) img=%s",
|
||||
err_code,
|
||||
bool(image_url),
|
||||
)
|
||||
jm_result = self._call_jimeng_video_generation(
|
||||
prompt=prompt,
|
||||
image_url=image_url,
|
||||
duration=duration,
|
||||
ratio=ratio,
|
||||
resolution=resolution,
|
||||
output_dir=output_dir,
|
||||
generate_audio=False, # 即梦 i2v 不带音频,音频由后续 ffmpeg 合成
|
||||
_portrait_fallback=True,
|
||||
)
|
||||
if jm_result is not None:
|
||||
return jm_result
|
||||
# 即梦也失败了,保留即梦的 last_video_error(已经由 _call_jimeng 设置)
|
||||
logger.error("[viral-video] 即梦通道重试也失败: %s", self.last_video_error)
|
||||
return None
|
||||
return None
|
||||
|
||||
logger.info("Seedance 任务已创建: task_id=%s ratio=%s", task_id, create_payload["ratio"])
|
||||
@@ -426,15 +702,21 @@ class DoubaoClient:
|
||||
usage: dict | None = None
|
||||
last_status: str = "queued"
|
||||
poll_count = 0
|
||||
last_poll_body: str = ""
|
||||
last_poll_sc: int = 0
|
||||
while time.time() < deadline:
|
||||
poll_count += 1
|
||||
try:
|
||||
resp = httpx.get(poll_url, headers=headers, timeout=self.timeout)
|
||||
try:
|
||||
if int(getattr(resp, "status_code", 200)) >= 400:
|
||||
resp.raise_for_status()
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
last_poll_sc = int(getattr(resp, "status_code", 200) or 200)
|
||||
last_poll_body = (getattr(resp, "text", "") or "")[:1500]
|
||||
if last_poll_sc >= 400:
|
||||
logger.warning("Seedance 轮询 HTTP %d: %s", last_poll_sc, last_poll_body[:300])
|
||||
if poll_count < 3:
|
||||
time.sleep(poll_interval)
|
||||
continue
|
||||
last_err = RuntimeError(f"poll HTTP {last_poll_sc}: {last_poll_body[:200]}")
|
||||
break
|
||||
data = resp.json()
|
||||
status = data.get("status", "")
|
||||
last_status = status
|
||||
@@ -445,13 +727,22 @@ class DoubaoClient:
|
||||
if video_url:
|
||||
logger.info("Seedance 任务成功: task_id=%s polls=%d usage=%s", task_id, poll_count, usage)
|
||||
break
|
||||
last_err = RuntimeError(f"task succeeded but no video_url: {str(data)[:300]}")
|
||||
logger.error("Seedance succeeded 但无 video_url: %s", last_err)
|
||||
# 成功但没 video_url:记录完整响应便于排查
|
||||
logger.error(
|
||||
"Seedance succeeded 但无 video_url: task_id=%s full_response=%s",
|
||||
task_id,
|
||||
str(data)[:1000],
|
||||
)
|
||||
last_err = RuntimeError("task succeeded but no video_url in response")
|
||||
last_poll_body = str(data)[:1000]
|
||||
break
|
||||
if status == "failed":
|
||||
err = data.get("error") or {}
|
||||
last_err = RuntimeError(f"task failed: code={err.get('code','')} msg={err.get('message','')}")
|
||||
logger.error("Seedance 任务失败 task_id=%s: %s", task_id, last_err)
|
||||
err_code = str(err.get("code", "") or "")
|
||||
err_msg = str(err.get("message", "") or err.get("msg", "") or "")
|
||||
last_err = RuntimeError(f"task failed: code={err_code} msg={err_msg}")
|
||||
logger.error("Seedance 任务失败 task_id=%s code=%s msg=%s", task_id, err_code, err_msg)
|
||||
last_poll_body = str(data)[:1000]
|
||||
break
|
||||
if status in ("expired", "cancelled"):
|
||||
last_err = RuntimeError(f"task {status}")
|
||||
@@ -462,18 +753,39 @@ class DoubaoClient:
|
||||
logger.info("Seedance 轮询中: task_id=%s status=%s polls=%d", task_id, status, poll_count)
|
||||
except httpx.HTTPStatusError as e:
|
||||
last_err = e
|
||||
logger.warning("Seedance 轮询 HTTP %d: %s", e.response.status_code, (e.response.text or "")[:300])
|
||||
last_poll_sc = e.response.status_code
|
||||
last_poll_body = (e.response.text or "")[:500]
|
||||
logger.warning("Seedance 轮询 HTTP %d: %s", e.response.status_code, last_poll_body[:300])
|
||||
except Exception as e:
|
||||
last_err = e
|
||||
logger.debug("Seedance 轮询异常: %s", e)
|
||||
time.sleep(poll_interval)
|
||||
|
||||
if not video_url:
|
||||
# 区分轮询超时 vs 任务失败
|
||||
if last_status in ("queued", "running", "pending") and poll_count > 0 and time.time() >= deadline:
|
||||
err_code, user_msg = (
|
||||
"network_error",
|
||||
f"视频生成超时(>{total_timeout}s),任务仍在排队,请稍后重试或联系管理员。",
|
||||
)
|
||||
detail = f"timeout after {total_timeout}s, polls={poll_count}, last_status={last_status}"
|
||||
else:
|
||||
err_code, user_msg = _classify_video_error(last_poll_sc, last_poll_body, last_err)
|
||||
detail = (last_poll_body or "")[:500] or (str(last_err) if last_err else f"last_status={last_status}")
|
||||
self.last_video_error = {
|
||||
"error_code": err_code,
|
||||
"user_message": user_msg,
|
||||
"status_code": last_poll_sc,
|
||||
"detail": detail,
|
||||
"task_id": task_id,
|
||||
"last_status": last_status,
|
||||
}
|
||||
logger.error(
|
||||
"Seedance 任务未成功: task_id=%s last_status=%s polls=%d err=%s (总等待 %.0fs)",
|
||||
"Seedance 任务未成功: task_id=%s last_status=%s polls=%d code=%s err=%s (总等待 %.0fs)",
|
||||
task_id,
|
||||
last_status,
|
||||
poll_count,
|
||||
err_code,
|
||||
last_err,
|
||||
total_timeout,
|
||||
)
|
||||
@@ -504,12 +816,135 @@ class DoubaoClient:
|
||||
os.remove(local_path)
|
||||
except Exception:
|
||||
pass
|
||||
self.last_video_error = {
|
||||
"error_code": "unknown",
|
||||
"user_message": "视频生成成功但下载文件为空,请稍后重试。",
|
||||
"status_code": 0,
|
||||
"detail": f"downloaded 0 bytes from {video_url[:120]}",
|
||||
}
|
||||
return None
|
||||
return {"video_path": local_path, "usage": usage}
|
||||
except Exception as e:
|
||||
logger.error("Seedance 视频下载失败: %s", e, exc_info=True)
|
||||
self.last_video_error = {
|
||||
"error_code": "network_error",
|
||||
"user_message": f"视频下载失败:{e!s}"[:200],
|
||||
"status_code": 0,
|
||||
"detail": str(e),
|
||||
}
|
||||
return None
|
||||
|
||||
def _call_jimeng_video_generation(
|
||||
self,
|
||||
*,
|
||||
prompt: str,
|
||||
image_url: str | None,
|
||||
duration: int,
|
||||
ratio: str | None,
|
||||
resolution: str,
|
||||
output_dir: str | None,
|
||||
generate_audio: bool = False,
|
||||
_portrait_fallback: bool = False,
|
||||
) -> dict | None:
|
||||
"""#2169: 调用即梦 cvtob 客户端做图生视频(真人参考图兜底通道)。
|
||||
|
||||
- 即梦 i2v 首帧接口只接受 1 张图、无原生音频(返回无声视频,音频由 ffmpeg 后合)。
|
||||
- 成功返回 {"video_path": str, "usage": {...}};失败写 self.last_video_error 并返回 None。
|
||||
- _portrait_fallback=True 时在日志里标注是从方舟拦截切过来的。
|
||||
"""
|
||||
from packages.shared.jimeng_client import get_jimeng_client
|
||||
|
||||
jm = get_jimeng_client()
|
||||
if jm is None:
|
||||
detail = "即梦 client 不可用(JIMENG_AK/SK 未配置)"
|
||||
if _portrait_fallback:
|
||||
# 从真人拦截切过来但即梦没配,仍把错误归到 portrait_intercept,让上层提示用户
|
||||
self.last_video_error = {
|
||||
"error_code": "portrait_intercept",
|
||||
"user_message": "参考素材包含真人照片被安全策略拦截,即梦兜底通道未启用,请联系管理员配置 JIMENG_AK/SK。",
|
||||
"status_code": 0,
|
||||
"detail": detail,
|
||||
"provider": "jimeng",
|
||||
}
|
||||
else:
|
||||
self.last_video_error = {
|
||||
"error_code": "auth_error",
|
||||
"user_message": "即梦视频通道未配置(JIMENG_AK/SK 缺失),请联系管理员。",
|
||||
"status_code": 0,
|
||||
"detail": detail,
|
||||
"provider": "jimeng",
|
||||
}
|
||||
logger.error("[jimeng] %s, portrait_fallback=%s", detail, _portrait_fallback)
|
||||
return None
|
||||
if not image_url:
|
||||
self.last_video_error = {
|
||||
"error_code": "invalid_param",
|
||||
"user_message": "即梦图生视频必须提供参考图片。",
|
||||
"status_code": 0,
|
||||
"detail": "empty image_url for jimeng i2v",
|
||||
"provider": "jimeng",
|
||||
}
|
||||
return None
|
||||
# 即梦 i2v 无声视频,generate_audio 强制 False
|
||||
jm.last_video_error = {}
|
||||
tag = "[portrait-fallback→jimeng]" if _portrait_fallback else "[jimeng-direct]"
|
||||
logger.info("%s 调用即梦: dur=%s ratio=%s res=%s img=%s", tag, duration, ratio, resolution, bool(image_url))
|
||||
try:
|
||||
result = jm.video_generation(
|
||||
prompt=prompt,
|
||||
image_url=image_url,
|
||||
duration=duration,
|
||||
ratio=ratio,
|
||||
resolution=resolution,
|
||||
output_dir=output_dir,
|
||||
generate_audio=False,
|
||||
)
|
||||
except Exception as je:
|
||||
logger.error("%s 即梦 video_generation 异常: %s", tag, je, exc_info=True)
|
||||
self.last_video_error = {
|
||||
"error_code": "unknown",
|
||||
"user_message": f"即梦视频生成异常:{je!s}"[:200],
|
||||
"status_code": 0,
|
||||
"detail": str(je),
|
||||
"provider": "jimeng",
|
||||
}
|
||||
return None
|
||||
if not result and jm.last_video_error:
|
||||
# 透传即梦错误;如果即梦也返回 portrait_intercept,说明图片真的有问题,直接给用户
|
||||
jm_err = dict(jm.last_video_error)
|
||||
jm_err["provider"] = "jimeng"
|
||||
if _portrait_fallback and jm_err.get("error_code") == "portrait_intercept":
|
||||
jm_err["user_message"] = (
|
||||
"参考素材真人肖像审核未通过(方舟+即梦双通道均被拦截),请更换非真人或授权清晰的照片后重试。"
|
||||
)
|
||||
self.last_video_error = jm_err
|
||||
return None
|
||||
if result:
|
||||
# 补充 usage 里的 provider 标记
|
||||
u = result.get("usage") or {}
|
||||
u.setdefault("provider", "jimeng")
|
||||
u.setdefault("model_key", "jimeng-3.0")
|
||||
result["usage"] = u
|
||||
logger.info("%s 即梦生成成功: %s", tag, result.get("video_path"))
|
||||
return result
|
||||
|
||||
def get_last_video_error(self) -> dict:
|
||||
"""返回最近一次 video_generation 失败的详细错误。空 dict 表示上次成功或未调用。"""
|
||||
return dict(self.last_video_error or {})
|
||||
|
||||
|
||||
def json_safe_dumps(obj: Any, max_len: int = 2000) -> str:
|
||||
"""安全 json 序列化,失败则 fallback 到 repr,超长截断。"""
|
||||
try:
|
||||
import json as _json
|
||||
|
||||
s = _json.dumps(obj, ensure_ascii=False, default=str)
|
||||
except Exception:
|
||||
s = repr(obj)
|
||||
if len(s) > max_len:
|
||||
s = s[:max_len] + f"...(truncated, total {len(s)})"
|
||||
return s
|
||||
|
||||
|
||||
# ── 单例 ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -618,20 +618,23 @@ def call_video_generation(
|
||||
reference_audios: list[str] | None = None,
|
||||
reference_videos: list[str] | None = None,
|
||||
) -> dict | None:
|
||||
"""调用 Seedance 2.5 生成视频(v1.6.1 单次出片版)。
|
||||
"""调用 Seedance / Wan 视频生成(v1.6.2 多模型版)。
|
||||
|
||||
成功返回 {"video_path": str, "usage": dict | None}(usage 含 completion_tokens),失败返回 None。
|
||||
|
||||
v1.6.1 关键约束(避免 20min 卡死):
|
||||
- 参考音频/视频/多图全部放进 content 数组并带 role=reference_audio/reference_video/reference_image;
|
||||
- 纯首帧无参考时(first_frame 模式),Seedance 2.5 强制 ratio=adaptive;
|
||||
传了参考音/视/多图时走 omni_reference 模式,ratio 可指定为 9:16(客户端内部自动判断)。
|
||||
- ratio 默认 9:16(竖屏),客户端会根据是否有参考自动在 first_frame/adaptive 与 omni/9:16 间切换;
|
||||
若创建任务因 ratio 报错(HTTP 400),客户端会自动回退到 adaptive 再试一次。
|
||||
失败时错误详情会写入 client.last_video_error,可通过 get_last_video_error() 读取:
|
||||
{"error_code": str, "user_message": str, "status_code": int, "detail": str, ...}
|
||||
"""
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
logger.warning("[ai_service] 豆包客户端未配置,跳过视频生成")
|
||||
msg = "豆包客户端未配置(DOUBAO_API_KEY 缺失),跳过视频生成"
|
||||
logger.warning("[ai_service] %s", msg)
|
||||
# 写入 last_video_error 供上层读取
|
||||
client.last_video_error = {
|
||||
"error_code": "auth_error",
|
||||
"user_message": "视频生成服务未配置,请联系管理员。",
|
||||
"status_code": 0,
|
||||
"detail": msg,
|
||||
}
|
||||
return None
|
||||
effective_ratio = ratio or "9:16"
|
||||
try:
|
||||
@@ -653,4 +656,21 @@ def call_video_generation(
|
||||
return client.video_generation(**kwargs)
|
||||
except Exception as e:
|
||||
logger.error("[ai_service] call_video_generation 异常: %s", e, exc_info=True)
|
||||
client.last_video_error = {
|
||||
"error_code": "unknown",
|
||||
"user_message": f"视频生成异常:{e!s}"[:200],
|
||||
"status_code": 0,
|
||||
"detail": str(e),
|
||||
}
|
||||
return None
|
||||
|
||||
|
||||
def get_last_video_error() -> dict:
|
||||
"""读取最近一次视频生成失败的详细错误(含 error_code/user_message/status_code/detail)。
|
||||
成功或未调用过返回空 dict。
|
||||
"""
|
||||
try:
|
||||
client = get_doubao_client()
|
||||
return client.get_last_video_error() if hasattr(client, "get_last_video_error") else {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
@@ -0,0 +1,344 @@
|
||||
"""DashScope 客户端(阿里云百炼 Wan 3.0 等非方舟模型)。
|
||||
|
||||
#2159: 新增 Wan 3.0 视频生成支持。DashScope 异步协议:
|
||||
- POST {base_url}/services/aigc/video-generation/video-synthesis (X-DashScope-Async: enable)
|
||||
→ 返回 output.task_id
|
||||
- GET {base_url}/tasks/{task_id} 轮询状态
|
||||
→ SUCCEEDED 时 output.video_url 可下载
|
||||
认证:Authorization: Bearer {DASHSCOPE_API_KEY}
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
|
||||
from packages.shared.config import get_shared_settings
|
||||
|
||||
# 网络/超时类异常父类集合:覆盖 Timeout/Connect/Network/ReadTimeout/WriteTimeout/PoolTimeout
|
||||
_HTTP_NETWORK_ERRORS = ()
|
||||
try:
|
||||
_HTTP_NETWORK_ERRORS = (httpx.TimeoutException, httpx.NetworkError)
|
||||
except Exception:
|
||||
_HTTP_NETWORK_ERRORS = (Exception,)
|
||||
|
||||
_HTTP_STATUS_ERROR = httpx.HTTPStatusError if hasattr(httpx, "HTTPStatusError") else Exception
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DASHSCOPE_CLIENT_SINGLETON: "DashScopeClient | None" = None
|
||||
|
||||
|
||||
def _classify_dashscope_error(status_code: int, body: str, task_msg: str = "") -> tuple[str, str]:
|
||||
"""DashScope 错误分类,返回 (error_code, user_message)。"""
|
||||
body_lower = (body or "").lower()
|
||||
msg_in_body = task_msg or ""
|
||||
try:
|
||||
import json as _json
|
||||
|
||||
parsed = _json.loads(body or "{}")
|
||||
if isinstance(parsed, dict):
|
||||
msg_in_body = msg_in_body or str(parsed.get("message", "") or "")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if status_code in (401, 403):
|
||||
return "auth_error", "Wan 3.0 服务鉴权失败(DASHSCOPE_API_KEY 无效或过期),请联系管理员。"
|
||||
if status_code == 429 or "rate" in body_lower or "throttl" in body_lower:
|
||||
return "rate_limit", "Wan 3.0 服务繁忙(限流),请稍等1-2分钟后重试。"
|
||||
if status_code == 400 and any(
|
||||
kw in body_lower for kw in ("portrait", "真人", "人脸", "肖像", "content_violation", "risk", "blocked")
|
||||
):
|
||||
return (
|
||||
"portrait_intercept",
|
||||
"参考素材包含真人照片或违规内容被安全策略拦截,请移除真人图片或调整文案后重试。",
|
||||
)
|
||||
if status_code == 404 or ("not found" in body_lower) or ("model" in body_lower and "not exist" in body_lower):
|
||||
return "model_not_found", "Wan 3.0 模型未开通或模型ID无效,请联系管理员。"
|
||||
if status_code in (402, 400) and ("quota" in body_lower or "billing" in body_lower or "insufficient" in body_lower):
|
||||
return "quota_exceeded", "Wan 3.0 服务配额不足,请联系管理员充值或稍后重试。"
|
||||
if status_code == 400:
|
||||
return "invalid_param", f"Wan 3.0 参数错误:{msg_in_body or body[:200]}"
|
||||
if status_code == 0:
|
||||
return "network_error", "Wan 3.0 服务连接失败(网络超时),请稍后重试。"
|
||||
# 任务内失败
|
||||
if task_msg and any(kw in task_msg.lower() for kw in ("portrait", "真人", "人脸", "violation", "blocked")):
|
||||
return "portrait_intercept", "Wan 3.0 视频内容被安全策略拦截,请调整文案或参考图后重试。"
|
||||
detail = msg_in_body or body[:200]
|
||||
return "unknown", f"Wan 3.0 视频生成失败(HTTP {status_code}):{detail}"
|
||||
|
||||
|
||||
class DashScopeClient:
|
||||
"""阿里云 DashScope 异步 API 客户端(Wan 3.0 等视频生成)。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
settings = get_shared_settings()
|
||||
self.api_key: str = getattr(settings, "dashscope_api_key", "") or os.getenv("DASHSCOPE_API_KEY", "")
|
||||
self.base_url: str = (
|
||||
getattr(settings, "dashscope_base_url", "") or "https://dashscope.aliyuncs.com/api/v1"
|
||||
).rstrip("/")
|
||||
self.poll_interval: int = int(getattr(settings, "dashscope_video_poll_interval", 10) or 10)
|
||||
self.total_timeout: int = int(getattr(settings, "dashscope_video_timeout", 900) or 900)
|
||||
self.max_retries: int = 2
|
||||
self.last_video_error: dict = {}
|
||||
|
||||
@property
|
||||
def is_available(self) -> bool:
|
||||
return bool(self.api_key)
|
||||
|
||||
def get_last_video_error(self) -> dict:
|
||||
return dict(self.last_video_error or {})
|
||||
|
||||
def _set_error(self, error_code: str, user_message: str, status_code: int = 0, detail: str = "", **extra) -> None:
|
||||
self.last_video_error = {
|
||||
"error_code": error_code,
|
||||
"user_message": user_message,
|
||||
"status_code": status_code,
|
||||
"detail": detail[:500] if detail else "",
|
||||
**extra,
|
||||
}
|
||||
|
||||
def video_generation(
|
||||
self,
|
||||
prompt: str,
|
||||
*,
|
||||
image_url: str | None = None,
|
||||
duration: int = 5,
|
||||
ratio: str | None = "9:16",
|
||||
resolution: str = "720p",
|
||||
watermark: bool = False,
|
||||
output_dir: str | None = None,
|
||||
model: str = "wan3.0-video",
|
||||
) -> dict | None:
|
||||
"""调用 DashScope 异步视频合成接口,轮询完成后下载到本地。
|
||||
|
||||
返回 {"video_path": str, "usage": dict | None};失败返回 None,错误详情写入 self.last_video_error。
|
||||
"""
|
||||
self.last_video_error = {}
|
||||
if not self.is_available:
|
||||
self._set_error("auth_error", "Wan 3.0 API key 未配置,请联系管理员。", detail="dashscope api_key empty")
|
||||
logger.error("[dashscope] API key 未配置,无法调用视频生成")
|
||||
return None
|
||||
if not prompt or not prompt.strip():
|
||||
self._set_error("invalid_param", "视频生成提示词不能为空。", detail="empty prompt")
|
||||
return None
|
||||
|
||||
# DashScope 分辨率参数:720P / 1080P / 480P(大写 P)
|
||||
res_upper = (resolution or "720p").upper().replace("P", "P")
|
||||
if res_upper == "480P":
|
||||
ds_res = "480P"
|
||||
elif res_upper == "1080P":
|
||||
ds_res = "1080P"
|
||||
else:
|
||||
ds_res = "720P"
|
||||
|
||||
# 构造 input+parameters
|
||||
input_obj: dict[str, Any] = {"prompt": prompt.strip()}
|
||||
if image_url:
|
||||
input_obj["img_url"] = image_url
|
||||
params: dict[str, Any] = {
|
||||
"resolution": ds_res,
|
||||
"duration": str(float(duration)),
|
||||
"watermark": bool(watermark),
|
||||
}
|
||||
# 比例透传:Wan 支持 "9:16" / "16:9" / "1:1" 等
|
||||
if ratio and ratio != "adaptive":
|
||||
params["aspect_ratio"] = ratio
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"model": model,
|
||||
"input": input_obj,
|
||||
"parameters": params,
|
||||
}
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
"X-DashScope-Async": "enable",
|
||||
}
|
||||
create_url = f"{self.base_url}/services/aigc/video-generation/video-synthesis"
|
||||
logger.info(
|
||||
"[dashscope] 创建任务: model=%s dur=%ds ratio=%s res=%s img=%s",
|
||||
model,
|
||||
duration,
|
||||
ratio,
|
||||
ds_res,
|
||||
bool(image_url),
|
||||
)
|
||||
logger.info("[dashscope] 创建任务 payload: model=%s params=%s", model, params)
|
||||
|
||||
# 创建任务
|
||||
task_id: str | None = None
|
||||
last_sc = 0
|
||||
last_body = ""
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
resp = httpx.post(create_url, headers=headers, json=payload, timeout=60)
|
||||
sc = int(getattr(resp, "status_code", 0) or 0)
|
||||
body_text = (getattr(resp, "text", "") or "")[:2000]
|
||||
last_sc = sc
|
||||
last_body = body_text
|
||||
if sc >= 400:
|
||||
logger.error("[dashscope] 创建任务 HTTP %d: %s", sc, body_text)
|
||||
if sc >= 500 and attempt < self.max_retries:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
err_code, user_msg = _classify_dashscope_error(sc, body_text)
|
||||
self._set_error(err_code, user_msg, sc, body_text, model=model)
|
||||
return None
|
||||
data = resp.json()
|
||||
tid = (data.get("output") or {}).get("task_id")
|
||||
if tid:
|
||||
task_id = tid
|
||||
break
|
||||
# 部分情况下 code != 错误
|
||||
code = data.get("code")
|
||||
if code and code != "":
|
||||
err_code, user_msg = _classify_dashscope_error(400, body_text, str(code))
|
||||
self._set_error(err_code, user_msg, sc, body_text, model=model)
|
||||
return None
|
||||
else:
|
||||
self._set_error("unknown", "Wan 3.0 响应格式异常,未返回任务ID", sc, str(data)[:500], model=model)
|
||||
return None
|
||||
except _HTTP_NETWORK_ERRORS as ne:
|
||||
last_sc = 0
|
||||
last_body = f"network error: {ne}"
|
||||
logger.warning(
|
||||
"[dashscope] 网络异常 %s,重试 %d/%d", type(ne).__name__, attempt + 1, self.max_retries + 1
|
||||
)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
self._set_error("network_error", "Wan 3.0 服务连接失败(网络超时),请稍后重试。", 0, str(ne))
|
||||
return None
|
||||
except Exception as _e:
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
logger.error("[dashscope] 创建任务最终失败: %s", _e)
|
||||
self._set_error("unknown", f"Wan 3.0 创建任务异常:{_e!s}"[:200], 0, str(_e))
|
||||
return None
|
||||
if not task_id:
|
||||
if not self.last_video_error:
|
||||
err_code, user_msg = _classify_dashscope_error(last_sc, last_body)
|
||||
self._set_error(err_code, user_msg, last_sc, last_body, model=model)
|
||||
return None
|
||||
|
||||
# 轮询任务
|
||||
poll_url = f"{self.base_url}/tasks/{task_id}"
|
||||
deadline = time.time() + self.total_timeout
|
||||
video_url: str | None = None
|
||||
usage: dict | None = None
|
||||
poll_count = 0
|
||||
last_status = ""
|
||||
while time.time() < deadline:
|
||||
poll_count += 1
|
||||
try:
|
||||
r = httpx.get(poll_url, headers=headers, timeout=30)
|
||||
psc = int(getattr(r, "status_code", 0) or 0)
|
||||
pbody = (getattr(r, "text", "") or "")[:1500]
|
||||
if psc >= 400:
|
||||
logger.warning("[dashscope] 轮询 HTTP %d: %s", psc, pbody[:300])
|
||||
if poll_count < 3:
|
||||
time.sleep(self.poll_interval)
|
||||
continue
|
||||
err_code, user_msg = _classify_dashscope_error(psc, pbody)
|
||||
self._set_error(err_code, user_msg, psc, pbody, task_id=task_id)
|
||||
return None
|
||||
d = r.json()
|
||||
out = d.get("output") or {}
|
||||
task_status = out.get("task_status") or d.get("task_status") or ""
|
||||
last_status = task_status
|
||||
if task_status == "SUCCEEDED":
|
||||
video_url = out.get("video_url") or ""
|
||||
usage = d.get("usage")
|
||||
if not video_url:
|
||||
# 结果在 results 数组
|
||||
results = out.get("results") or []
|
||||
if results and isinstance(results, list):
|
||||
video_url = results[0].get("url") or results[0].get("video_url")
|
||||
if video_url:
|
||||
logger.info("[dashscope] 任务 %s 完成: %s", task_id, video_url[:120])
|
||||
break
|
||||
logger.error("[dashscope] 任务 %s SUCCEEDED 但无 video_url: %s", task_id, str(d)[:500])
|
||||
self._set_error(
|
||||
"unknown",
|
||||
"Wan 3.0 任务成功但未返回视频URL,请联系管理员。",
|
||||
200,
|
||||
str(d)[:500],
|
||||
task_id=task_id,
|
||||
)
|
||||
return None
|
||||
if task_status in ("FAILED", "FAILED_WITH_ERROR", "ERROR"):
|
||||
msg = out.get("message") or d.get("message") or out.get("error_msg") or "unknown error"
|
||||
logger.error("[dashscope] 任务 %s 失败: %s", task_id, msg)
|
||||
err_code, user_msg = _classify_dashscope_error(200, "", msg)
|
||||
self._set_error(err_code, user_msg, 200, msg, task_id=task_id, last_status=task_status)
|
||||
return None
|
||||
if task_status in ("CANCELED", "CANCELLED"):
|
||||
logger.warning("[dashscope] 任务 %s 被取消", task_id)
|
||||
self._set_error("unknown", "Wan 3.0 任务被取消。", 200, "task cancelled", task_id=task_id)
|
||||
return None
|
||||
# PENDING / RUNNING / SUSPENDED → 继续轮询
|
||||
if poll_count % 5 == 0:
|
||||
logger.info("[dashscope] 轮询中 task=%s status=%s polls=%d", task_id, task_status, poll_count)
|
||||
except Exception as e:
|
||||
logger.warning("[dashscope] 轮询异常: %s", e)
|
||||
time.sleep(self.poll_interval)
|
||||
if not video_url:
|
||||
logger.error("[dashscope] 任务 %s 轮询超时(%ds)", task_id, self.total_timeout)
|
||||
self._set_error(
|
||||
"network_error",
|
||||
f"Wan 3.0 视频生成超时(>{self.total_timeout}s),任务仍在排队,请稍后重试。",
|
||||
0,
|
||||
f"timeout after {self.total_timeout}s, polls={poll_count}, last_status={last_status}",
|
||||
task_id=task_id,
|
||||
last_status=last_status,
|
||||
)
|
||||
return None
|
||||
|
||||
# 下载视频
|
||||
out_dir = output_dir or os.path.join(os.getcwd(), "seedance_outputs")
|
||||
os.makedirs(out_dir, exist_ok=True)
|
||||
suffix = Path(urlparse(video_url).path).suffix or ".mp4"
|
||||
if suffix.lower() not in (".mp4", ".mov", ".webm"):
|
||||
suffix = ".mp4"
|
||||
safe_tid = "".join(c if c.isalnum() or c in "-_" else "_" for c in task_id)[:40]
|
||||
out_path = os.path.join(out_dir, f"wan_{safe_tid}{suffix}")
|
||||
try:
|
||||
with httpx.stream("GET", video_url, timeout=300, follow_redirects=True) as resp:
|
||||
dsc = int(getattr(resp, "status_code", 0) or 0)
|
||||
if dsc >= 400:
|
||||
logger.error("[dashscope] 下载 HTTP %d", dsc)
|
||||
self._set_error("network_error", "Wan 3.0 视频下载失败(HTTP错误),请稍后重试。", dsc)
|
||||
return None
|
||||
with open(out_path, "wb") as f:
|
||||
for chunk in resp.iter_bytes(chunk_size=1024 * 256):
|
||||
if chunk:
|
||||
f.write(chunk)
|
||||
except Exception as e:
|
||||
logger.error("[dashscope] 下载视频失败: %s", e, exc_info=True)
|
||||
self._set_error("network_error", f"Wan 3.0 视频下载失败:{e!s}"[:200], 0, str(e))
|
||||
return None
|
||||
size = os.path.getsize(out_path) if os.path.exists(out_path) else 0
|
||||
if size < 1024:
|
||||
logger.error("[dashscope] 下载文件过小: %d bytes", size)
|
||||
self._set_error("unknown", "Wan 3.0 视频下载文件过小,请稍后重试。", 0, f"downloaded only {size} bytes")
|
||||
return None
|
||||
logger.info("[dashscope] 视频已下载: %s (%d bytes)", out_path, size)
|
||||
return {"video_path": out_path, "usage": usage}
|
||||
|
||||
|
||||
def get_dashscope_client() -> DashScopeClient | None:
|
||||
"""返回 DashScope 客户端单例;未配置 API key 时返回 None。"""
|
||||
global _DASHSCOPE_CLIENT_SINGLETON
|
||||
if _DASHSCOPE_CLIENT_SINGLETON is None:
|
||||
_DASHSCOPE_CLIENT_SINGLETON = DashScopeClient()
|
||||
if not _DASHSCOPE_CLIENT_SINGLETON.is_available:
|
||||
return None
|
||||
return _DASHSCOPE_CLIENT_SINGLETON
|
||||
@@ -0,0 +1,526 @@
|
||||
"""即梦(Jimeng)视觉 API 客户端 —— 火山引擎 cvtob。
|
||||
|
||||
#2169: 真人参考图在方舟 Seedance 走 B 端审核会被 50411 拦截,
|
||||
即梦走 C 端审核链路,普通真人照片可过审。接入即梦 i2v 作为参考图场景兜底通道。
|
||||
|
||||
接口协议(jimeng_i2v_first_v30 —— 视频3.0 720P 首帧图生视频):
|
||||
- 接口地址:https://visual.volcengineapi.com
|
||||
- 鉴权:火山 V4 签名(Region=cn-north-1, Service=cv),使用 AK/SK
|
||||
- 提交任务:POST ?Action=CVSync2AsyncSubmitTask&Version=2022-08-31
|
||||
body: {"req_key": "jimeng_i2v_first_v30", "image_urls": ["<url>"], "prompt": "...", "seed": -1, "frames": 121}
|
||||
-> {"code": 10000, "data": {"task_id": "..."}}
|
||||
- 查询任务:POST ?Action=CVSync2AsyncGetResult&Version=2022-08-31
|
||||
body: {"req_key": "jimeng_i2v_first_v30", "task_id": "..."}
|
||||
-> {"code": 10000, "data": {"status": "in_queue|generating|done", "video_url": "..."}}
|
||||
- 视频 URL 有效期 1 小时,必须立即下载到本地。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from urllib.parse import quote, urlparse
|
||||
|
||||
import httpx
|
||||
|
||||
from packages.shared.config import get_shared_settings
|
||||
|
||||
# 网络/超时类异常父类集合
|
||||
_HTTP_NETWORK_ERRORS = ()
|
||||
try:
|
||||
_HTTP_NETWORK_ERRORS = (httpx.TimeoutException, httpx.NetworkError)
|
||||
except Exception:
|
||||
_HTTP_NETWORK_ERRORS = (Exception,)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_JIMENG_CLIENT_SINGLETON: "JimengClient | None" = None
|
||||
|
||||
# ── V4 签名常量 ────────────────────────────────────────────────────────
|
||||
_JIMENG_REGION = "cn-north-1"
|
||||
_JIMENG_SERVICE = "cv"
|
||||
_JIMENG_VERSION = "2022-08-31"
|
||||
_ACTION_SUBMIT = "CVSync2AsyncSubmitTask"
|
||||
_ACTION_POLL = "CVSync2AsyncGetResult"
|
||||
_CONTENT_TYPE = "application/json"
|
||||
_SIGNED_HEADERS_LIST = ["content-type", "host", "x-content-sha256", "x-date"]
|
||||
_SIGNED_HEADERS_STR = ";".join(_SIGNED_HEADERS_LIST)
|
||||
|
||||
|
||||
def _norm_query(params: dict[str, str]) -> str:
|
||||
"""构造规范查询串:按 key 排序,URL 编码(safe=-_.~),空格->%20。"""
|
||||
parts = []
|
||||
for k in sorted(params.keys()):
|
||||
v = params[k]
|
||||
ek = quote(str(k), safe="-_.~")
|
||||
ev = quote(str(v), safe="-_.~").replace("+", "%20")
|
||||
parts.append(f"{ek}={ev}")
|
||||
return "&".join(parts)
|
||||
|
||||
|
||||
def _hmac_sha256(key: bytes, msg: str) -> bytes:
|
||||
return hmac.new(key, msg.encode("utf-8"), hashlib.sha256).digest()
|
||||
|
||||
|
||||
def _sha256_hex(data: bytes) -> str:
|
||||
return hashlib.sha256(data).hexdigest()
|
||||
|
||||
|
||||
def _sign_v4(
|
||||
ak: str,
|
||||
sk: str,
|
||||
method: str,
|
||||
host: str,
|
||||
query: dict[str, str],
|
||||
body_bytes: bytes,
|
||||
x_date: str,
|
||||
) -> dict[str, str]:
|
||||
"""火山 V4 签名,返回需要附加到请求的 headers 字典。
|
||||
|
||||
x_date 形如 "20260101T120000Z"(UTC)。
|
||||
short_date = x_date[:8](YYYYMMDD)。
|
||||
"""
|
||||
short_date = x_date[:8]
|
||||
payload_hash = _sha256_hex(body_bytes)
|
||||
canon_uri = "/"
|
||||
canon_query = _norm_query(query)
|
||||
canon_headers = f"content-type:{_CONTENT_TYPE}\nhost:{host}\nx-content-sha256:{payload_hash}\nx-date:{x_date}\n"
|
||||
canon_request = f"{method}\n{canon_uri}\n{canon_query}\n{canon_headers}\n{_SIGNED_HEADERS_STR}\n{payload_hash}"
|
||||
credential_scope = f"{short_date}/{_JIMENG_REGION}/{_JIMENG_SERVICE}/request"
|
||||
string_to_sign = f"HMAC-SHA256\n{x_date}\n{credential_scope}\n{_sha256_hex(canon_request.encode('utf-8'))}"
|
||||
k_date = _hmac_sha256(sk.encode("utf-8"), short_date)
|
||||
k_region = _hmac_sha256(k_date, _JIMENG_REGION)
|
||||
k_service = _hmac_sha256(k_region, _JIMENG_SERVICE)
|
||||
k_signing = _hmac_sha256(k_service, "request")
|
||||
signature = hmac.new(k_signing, string_to_sign.encode("utf-8"), hashlib.sha256).hexdigest()
|
||||
authorization = (
|
||||
f"HMAC-SHA256 Credential={ak}/{credential_scope}, SignedHeaders={_SIGNED_HEADERS_STR}, Signature={signature}"
|
||||
)
|
||||
return {
|
||||
"Content-Type": _CONTENT_TYPE,
|
||||
"Host": host,
|
||||
"X-Content-Sha256": payload_hash,
|
||||
"X-Date": x_date,
|
||||
"Authorization": authorization,
|
||||
}
|
||||
|
||||
|
||||
# ── 错误分类 ──────────────────────────────────────────────────────────
|
||||
|
||||
# 即梦业务码 -> 是否可重试映射
|
||||
_JIMENG_RETRYABLE_CODES = {50511, 50516, 50429, 50430, 50500, 50501}
|
||||
_JIMENG_NON_RETRYABLE_CODES = {50411, 50412, 50413, 50512, 50513, 50514}
|
||||
|
||||
|
||||
def _classify_jimeng_error(status_code: int, body: str, biz_code: int | None = None) -> tuple[str, str, bool]:
|
||||
"""即梦错误分类,返回 (error_code, user_message, is_retryable)。"""
|
||||
code = biz_code if biz_code is not None else 0
|
||||
body_lower = (body or "").lower()
|
||||
|
||||
# 业务码优先
|
||||
if code == 50411:
|
||||
return (
|
||||
"portrait_intercept",
|
||||
"即梦通道:参考图片前审核未通过(Pre Img Risk Not Pass),请更换参考图后重试。",
|
||||
False,
|
||||
)
|
||||
if code == 50511:
|
||||
return "task_failed", "即梦通道:输出图片后审核未通过,可稍后重试。", True
|
||||
if code in (50412, 50413, 50512):
|
||||
return "invalid_param", "即梦通道:提示词或文本审核不通过,请调整文案后重试。", False
|
||||
if code == 50516:
|
||||
return "task_failed", "即梦通道:输出视频后审核未通过,可稍后重试。", True
|
||||
if code in (50429, 50430):
|
||||
return "rate_limit", "即梦通道:QPS/并发超限,请稍等 1-2 分钟后重试。", True
|
||||
if code in (50500, 50501):
|
||||
return "network_error", "即梦通道:服务内部错误,可稍后重试。", True
|
||||
|
||||
# HTTP 层兜底
|
||||
if status_code in (401, 403):
|
||||
return "auth_error", "即梦通道:AK/SK 鉴权失败,请联系管理员检查 JIMENG_AK/SK 配置。", False
|
||||
if status_code == 429:
|
||||
return "rate_limit", "即梦通道:服务限流,请稍后重试。", True
|
||||
if status_code == 404:
|
||||
return "model_not_found", "即梦通道:接口不存在(req_key 或 Action 错误),请联系管理员。", False
|
||||
if status_code in (402, 400) and any(kw in body_lower for kw in ("quota", "billing", "insufficient", "余额")):
|
||||
return "quota_exceeded", "即梦通道:账户余额/配额不足,请联系管理员充值。", False
|
||||
if status_code == 400:
|
||||
msg = ""
|
||||
try:
|
||||
msg = str(json.loads(body or "{}").get("message", "") or "")
|
||||
except Exception:
|
||||
pass
|
||||
return "invalid_param", f"即梦通道:参数错误:{msg or body[:200]}", False
|
||||
if status_code == 0:
|
||||
return "network_error", "即梦通道:网络连接失败,请稍后重试。", True
|
||||
# 任务内失败
|
||||
if code and code != 10000:
|
||||
return "unknown", f"即梦通道:视频生成失败(错误码 {code}),请稍后重试。", code in _JIMENG_RETRYABLE_CODES
|
||||
detail = body[:200]
|
||||
return "unknown", f"即梦通道:视频生成失败(HTTP {status_code}):{detail}", False
|
||||
|
||||
|
||||
# ── 即梦客户端 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class JimengClient:
|
||||
"""火山引擎即梦视觉 API(cvtob)异步客户端,支持图生视频首帧(jimeng_i2v_first_v30)。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
settings = get_shared_settings()
|
||||
self.ak: str = getattr(settings, "jimeng_ak", "") or os.getenv("JIMENG_AK", "")
|
||||
self.sk: str = getattr(settings, "jimeng_sk", "") or os.getenv("JIMENG_SK", "")
|
||||
self.base_url: str = (getattr(settings, "jimeng_base_url", "") or "https://visual.volcengineapi.com").rstrip(
|
||||
"/"
|
||||
)
|
||||
self.req_key: str = getattr(settings, "jimeng_req_key", "") or "jimeng_i2v_first_v30"
|
||||
self.poll_interval: int = int(getattr(settings, "jimeng_video_poll_interval", 5) or 5)
|
||||
self.total_timeout: int = int(getattr(settings, "jimeng_video_timeout", 600) or 600)
|
||||
self.max_retries: int = 2
|
||||
self.last_video_error: dict = {}
|
||||
# 解析 base_url 里的 host(用于签名 Host 头)
|
||||
parsed = urlparse(self.base_url)
|
||||
self.host: str = parsed.netloc or "visual.volcengineapi.com"
|
||||
|
||||
@property
|
||||
def is_available(self) -> bool:
|
||||
return bool(self.ak and self.sk)
|
||||
|
||||
def get_last_video_error(self) -> dict:
|
||||
return dict(self.last_video_error or {})
|
||||
|
||||
def _set_error(self, error_code: str, user_message: str, status_code: int = 0, detail: str = "", **extra) -> None:
|
||||
self.last_video_error = {
|
||||
"error_code": error_code,
|
||||
"user_message": user_message,
|
||||
"status_code": status_code,
|
||||
"detail": detail[:500] if detail else "",
|
||||
"provider": "jimeng",
|
||||
**extra,
|
||||
}
|
||||
|
||||
# ── 内部 HTTP:签名 + 请求 ──────────────────────────────────────
|
||||
|
||||
def _signed_request(
|
||||
self,
|
||||
method: str,
|
||||
action: str,
|
||||
body_obj: dict[str, Any],
|
||||
timeout: float = 60.0,
|
||||
) -> tuple[int, str, dict]:
|
||||
"""发送一次带 V4 签名的请求,返回 (status_code, body_text, parsed_json)。"""
|
||||
body_bytes = json.dumps(body_obj, ensure_ascii=False).encode("utf-8")
|
||||
query = {"Action": action, "Version": _JIMENG_VERSION}
|
||||
x_date = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ")
|
||||
headers = _sign_v4(self.ak, self.sk, method, self.host, query, body_bytes, x_date)
|
||||
url = f"{self.base_url}/?{_norm_query(query)}"
|
||||
resp = httpx.request(
|
||||
method,
|
||||
url,
|
||||
headers=headers,
|
||||
content=body_bytes,
|
||||
timeout=timeout,
|
||||
)
|
||||
sc = int(getattr(resp, "status_code", 0) or 0)
|
||||
text = getattr(resp, "text", "") or ""
|
||||
try:
|
||||
data = resp.json()
|
||||
except Exception:
|
||||
data = {}
|
||||
return sc, text, data
|
||||
|
||||
# ── 提交任务 ────────────────────────────────────────────────────
|
||||
|
||||
def _submit_task(
|
||||
self,
|
||||
prompt: str,
|
||||
image_url: str,
|
||||
frames: int = 121,
|
||||
seed: int = -1,
|
||||
) -> str | None:
|
||||
"""提交图生视频任务,成功返回 task_id;失败写 last_video_error 并返回 None。"""
|
||||
body: dict[str, Any] = {
|
||||
"req_key": self.req_key,
|
||||
"prompt": prompt.strip()[:800],
|
||||
"image_urls": [image_url],
|
||||
"seed": int(seed) if seed and seed > 0 else -1,
|
||||
"frames": int(frames),
|
||||
}
|
||||
last_sc = 0
|
||||
last_body = ""
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
sc, text, data = self._signed_request("POST", _ACTION_SUBMIT, body, timeout=60.0)
|
||||
last_sc, last_body = sc, text
|
||||
if sc >= 400:
|
||||
logger.error("[jimeng] 提交 HTTP %d: %s", sc, text[:500])
|
||||
if sc >= 500 and attempt < self.max_retries:
|
||||
time.sleep(0.8 * (2**attempt))
|
||||
continue
|
||||
biz_code = data.get("code") if isinstance(data, dict) else None
|
||||
err_code, user_msg, _ = _classify_jimeng_error(sc, text, biz_code)
|
||||
self._set_error(err_code, user_msg, sc, text, req_key=self.req_key)
|
||||
return None
|
||||
code = data.get("code") if isinstance(data, dict) else None
|
||||
if code == 10000:
|
||||
d = data.get("data") or {}
|
||||
tid = d.get("task_id")
|
||||
if tid:
|
||||
return str(tid)
|
||||
err_code, user_msg, retry = _classify_jimeng_error(sc, text, code)
|
||||
logger.error(
|
||||
"[jimeng] 提交业务错误 code=%s msg=%s",
|
||||
code,
|
||||
(data.get("message") if isinstance(data, dict) else ""),
|
||||
)
|
||||
if retry and attempt < self.max_retries:
|
||||
time.sleep(0.8 * (2**attempt))
|
||||
continue
|
||||
self._set_error(err_code, user_msg, sc, text, req_key=self.req_key, biz_code=code)
|
||||
return None
|
||||
except _HTTP_NETWORK_ERRORS as ne:
|
||||
last_sc, last_body = 0, f"network error: {ne}"
|
||||
logger.warning(
|
||||
"[jimeng] 提交网络异常 %s,重试 %d/%d", type(ne).__name__, attempt + 1, self.max_retries + 1
|
||||
)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.8 * (2**attempt))
|
||||
continue
|
||||
self._set_error("network_error", "即梦通道:提交任务网络异常,请稍后重试。", 0, str(ne))
|
||||
return None
|
||||
except Exception as e:
|
||||
last_sc, last_body = 0, f"exception: {e}"
|
||||
logger.error("[jimeng] 提交异常: %s", e, exc_info=True)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.8 * (2**attempt))
|
||||
continue
|
||||
self._set_error("unknown", f"即梦通道:提交任务异常:{e!s}"[:200], 0, str(e))
|
||||
return None
|
||||
if not self.last_video_error:
|
||||
err_code, user_msg, _ = _classify_jimeng_error(last_sc, last_body)
|
||||
self._set_error(err_code, user_msg, last_sc, last_body)
|
||||
return None
|
||||
|
||||
# ── 轮询结果 ────────────────────────────────────────────────────
|
||||
|
||||
def _poll_result(self, task_id: str) -> str | None:
|
||||
"""轮询任务直到 done/failed/expired/timeout,成功返回 video_url。"""
|
||||
deadline = time.time() + self.total_timeout
|
||||
poll_count = 0
|
||||
last_status = ""
|
||||
poll_body = {"req_key": self.req_key, "task_id": task_id}
|
||||
while time.time() < deadline:
|
||||
poll_count += 1
|
||||
try:
|
||||
sc, text, data = self._signed_request("POST", _ACTION_POLL, poll_body, timeout=30.0)
|
||||
if sc >= 400:
|
||||
logger.warning("[jimeng] 轮询 HTTP %d: %s", sc, text[:300])
|
||||
if poll_count < 3:
|
||||
time.sleep(self.poll_interval)
|
||||
continue
|
||||
err_code, user_msg, _ = _classify_jimeng_error(sc, text)
|
||||
self._set_error(err_code, user_msg, sc, text, task_id=task_id)
|
||||
return None
|
||||
code = data.get("code") if isinstance(data, dict) else None
|
||||
d = data.get("data") if isinstance(data, dict) else None
|
||||
if code != 10000 or not isinstance(d, dict):
|
||||
err_code, user_msg, retry = _classify_jimeng_error(sc, text, code)
|
||||
logger.error(
|
||||
"[jimeng] 轮询业务错误 task=%s code=%s msg=%s",
|
||||
task_id,
|
||||
code,
|
||||
(data.get("message") if isinstance(data, dict) else ""),
|
||||
)
|
||||
if retry and poll_count < 3:
|
||||
time.sleep(self.poll_interval)
|
||||
continue
|
||||
self._set_error(err_code, user_msg, sc, text, task_id=task_id, biz_code=code)
|
||||
return None
|
||||
status = d.get("status", "") or ""
|
||||
last_status = status
|
||||
if status == "done":
|
||||
video_url = d.get("video_url") or ""
|
||||
if video_url:
|
||||
logger.info("[jimeng] 任务 %s 完成 polls=%d", task_id, poll_count)
|
||||
return str(video_url)
|
||||
logger.error("[jimeng] 任务 %s done 但无 video_url: %s", task_id, str(d)[:500])
|
||||
self._set_error(
|
||||
"unknown",
|
||||
"即梦通道:任务成功但未返回视频URL,请联系管理员。",
|
||||
200,
|
||||
str(d)[:500],
|
||||
task_id=task_id,
|
||||
)
|
||||
return None
|
||||
if status in ("not_found", "expired"):
|
||||
logger.error("[jimeng] 任务 %s 状态 %s", task_id, status)
|
||||
self._set_error(
|
||||
"network_error" if status == "expired" else "unknown",
|
||||
f"即梦通道:任务{'已过期' if status == 'expired' else '未找到'},请重新提交。",
|
||||
200,
|
||||
f"task {status}",
|
||||
task_id=task_id,
|
||||
)
|
||||
return None
|
||||
if poll_count % 6 == 0:
|
||||
logger.info("[jimeng] 轮询中 task=%s status=%s polls=%d", task_id, status, poll_count)
|
||||
except _HTTP_NETWORK_ERRORS as ne:
|
||||
logger.warning("[jimeng] 轮询网络异常 %s", ne)
|
||||
except Exception as e:
|
||||
logger.debug("[jimeng] 轮询异常: %s", e)
|
||||
time.sleep(self.poll_interval)
|
||||
logger.error(
|
||||
"[jimeng] 任务 %s 轮询超时(%ds)polls=%d last_status=%s",
|
||||
task_id,
|
||||
self.total_timeout,
|
||||
poll_count,
|
||||
last_status,
|
||||
)
|
||||
self._set_error(
|
||||
"network_error",
|
||||
f"即梦通道:视频生成超时(>{self.total_timeout}s),任务仍在排队,请稍后重试。",
|
||||
0,
|
||||
f"timeout after {self.total_timeout}s, polls={poll_count}, last_status={last_status}",
|
||||
task_id=task_id,
|
||||
last_status=last_status,
|
||||
)
|
||||
return None
|
||||
|
||||
# ── 下载视频 ────────────────────────────────────────────────────
|
||||
|
||||
def _download_video(self, video_url: str, output_dir: str, task_id: str) -> str | None:
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
suffix = Path(urlparse(video_url).path).suffix or ".mp4"
|
||||
if suffix.lower() not in (".mp4", ".mov", ".webm"):
|
||||
suffix = ".mp4"
|
||||
safe_tid = "".join(c if c.isalnum() or c in "-_" else "_" for c in task_id)[:40]
|
||||
out_path = os.path.join(output_dir, f"jimeng_{safe_tid}_{uuid.uuid4().hex[:8]}{suffix}")
|
||||
try:
|
||||
with httpx.stream("GET", video_url, timeout=180, follow_redirects=True) as resp:
|
||||
dsc = int(getattr(resp, "status_code", 0) or 0)
|
||||
if dsc >= 400:
|
||||
logger.error("[jimeng] 下载 HTTP %d", dsc)
|
||||
self._set_error("network_error", "即梦通道:视频下载失败(HTTP错误),请稍后重试。", dsc)
|
||||
return None
|
||||
with open(out_path, "wb") as f:
|
||||
for chunk in resp.iter_bytes(chunk_size=1024 * 256):
|
||||
if chunk:
|
||||
f.write(chunk)
|
||||
except Exception as e:
|
||||
logger.error("[jimeng] 下载视频失败: %s", e, exc_info=True)
|
||||
self._set_error("network_error", f"即梦通道:视频下载失败:{e!s}"[:200], 0, str(e))
|
||||
return None
|
||||
size = os.path.getsize(out_path) if os.path.exists(out_path) else 0
|
||||
if size < 1024:
|
||||
logger.error("[jimeng] 下载文件过小: %d bytes", size)
|
||||
self._set_error("unknown", "即梦通道:视频下载文件过小,请稍后重试。", 0, f"downloaded only {size} bytes")
|
||||
try:
|
||||
os.remove(out_path)
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
logger.info("[jimeng] 视频已下载: %s (%d bytes)", out_path, size)
|
||||
return out_path
|
||||
|
||||
# ── 对外主入口 ──────────────────────────────────────────────────
|
||||
|
||||
def video_generation(
|
||||
self,
|
||||
prompt: str,
|
||||
*,
|
||||
image_url: str,
|
||||
duration: int = 5,
|
||||
ratio: str | None = "9:16",
|
||||
resolution: str = "720p",
|
||||
output_dir: str | None = None,
|
||||
generate_audio: bool = False,
|
||||
) -> dict | None:
|
||||
"""即梦图生视频主入口。
|
||||
|
||||
成功返回 {"video_path": str, "usage": {"provider","duration_seconds","frames","req_key","billing_mode"}};
|
||||
失败返回 None,详情在 self.last_video_error。
|
||||
|
||||
注意:jimeng_i2v_first_v30 不支持原生音频(generate_audio 被忽略,返回无声视频),
|
||||
音频由后续 ffmpeg 合成阶段叠加 TTS。
|
||||
支持时长:5s(frames=121)/10s(frames=241),>10s 截断并打 warning。
|
||||
分辨率固定 720P;ratio 对首帧 i2v 无效(自动按图片比例)。
|
||||
"""
|
||||
self.last_video_error = {}
|
||||
if not self.is_available:
|
||||
self._set_error(
|
||||
"auth_error",
|
||||
"即梦通道未配置(JIMENG_AK/SK 缺失),请联系管理员。",
|
||||
detail="jimeng ak/sk empty",
|
||||
)
|
||||
logger.error("[jimeng] AK/SK 未配置,无法调用")
|
||||
return None
|
||||
if not prompt or not prompt.strip():
|
||||
self._set_error("invalid_param", "视频生成提示词不能为空。", detail="empty prompt")
|
||||
return None
|
||||
if not image_url or not image_url.strip():
|
||||
self._set_error("invalid_param", "即梦图生视频必须提供参考图片。", detail="empty image_url")
|
||||
return None
|
||||
|
||||
dur = int(duration or 5)
|
||||
if dur <= 5:
|
||||
frames = 121
|
||||
real_dur = 5
|
||||
elif dur <= 10:
|
||||
frames = 241
|
||||
real_dur = 10
|
||||
else:
|
||||
logger.warning("[jimeng] 请求时长 %ds 超出即梦 i2v 上限 10s,截断到 10s(frames=241)", dur)
|
||||
frames = 241
|
||||
real_dur = 10
|
||||
|
||||
out_dir = output_dir or "/tmp"
|
||||
logger.info(
|
||||
"[jimeng] 提交任务: req_key=%s dur=%ds(frames=%d) ratio=%s res=%s gen_audio=%s img=%s",
|
||||
self.req_key,
|
||||
real_dur,
|
||||
frames,
|
||||
ratio,
|
||||
resolution,
|
||||
generate_audio,
|
||||
bool(image_url),
|
||||
)
|
||||
|
||||
task_id = self._submit_task(prompt=prompt, image_url=image_url, frames=frames, seed=-1)
|
||||
if not task_id:
|
||||
return None
|
||||
logger.info("[jimeng] 任务已提交: task_id=%s", task_id)
|
||||
|
||||
video_url = self._poll_result(task_id)
|
||||
if not video_url:
|
||||
return None
|
||||
|
||||
local_path = self._download_video(video_url, out_dir, task_id)
|
||||
if not local_path:
|
||||
return None
|
||||
|
||||
usage = {
|
||||
"provider": "jimeng",
|
||||
"duration_seconds": real_dur,
|
||||
"frames": frames,
|
||||
"req_key": self.req_key,
|
||||
"billing_mode": "per_second",
|
||||
}
|
||||
return {"video_path": local_path, "usage": usage}
|
||||
|
||||
|
||||
def get_jimeng_client() -> "JimengClient | None":
|
||||
"""返回即梦客户端单例;未配置 AK/SK 时返回 None。"""
|
||||
global _JIMENG_CLIENT_SINGLETON
|
||||
if _JIMENG_CLIENT_SINGLETON is None:
|
||||
_JIMENG_CLIENT_SINGLETON = JimengClient()
|
||||
if not _JIMENG_CLIENT_SINGLETON.is_available:
|
||||
return None
|
||||
return _JIMENG_CLIENT_SINGLETON
|
||||
@@ -26,7 +26,6 @@ oss2==2.18.4
|
||||
|
||||
# HTTP 客户端(pin 间接依赖防止版本漂移)
|
||||
httpx==0.27.2
|
||||
cryptography==50.0.2
|
||||
httpcore==1.0.7
|
||||
h2==4.1.0
|
||||
|
||||
|
||||
@@ -483,3 +483,180 @@ class TestVideoGenerationCancelled:
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=10, doubao_video_model="seedance"
|
||||
)
|
||||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||||
|
||||
|
||||
# ============ #2157 _resolve_video_model_id 模型ID映射单测 ============
|
||||
|
||||
|
||||
class TestResolveVideoModelId:
|
||||
"""覆盖 _resolve_video_model_id 各分支(#2157 P0 修复)。"""
|
||||
|
||||
def _import_target(self):
|
||||
from packages.shared.ai_client import _resolve_video_model_id
|
||||
|
||||
return _resolve_video_model_id
|
||||
|
||||
def test_none_uses_default(self):
|
||||
fn = self._import_target()
|
||||
with patch("packages.shared.ai_client.get_shared_settings") as ms:
|
||||
ms.return_value = MagicMock(doubao_video_model="doubao-seedance-2-5-260628")
|
||||
assert fn(None) == "doubao-seedance-2-5-260628"
|
||||
|
||||
def test_empty_uses_default(self):
|
||||
fn = self._import_target()
|
||||
with patch("packages.shared.ai_client.get_shared_settings") as ms:
|
||||
ms.return_value = MagicMock(doubao_video_model="doubao-seedance-2-5-260628")
|
||||
assert fn(" ") == "doubao-seedance-2-5-260628"
|
||||
|
||||
def test_doubao_prefix_passthrough(self):
|
||||
fn = self._import_target()
|
||||
assert fn("doubao-seedance-2-5-260628") == "doubao-seedance-2-5-260628"
|
||||
|
||||
def test_ep_prefix_passthrough(self):
|
||||
fn = self._import_target()
|
||||
assert fn("ep-20260721114705-b568m") == "ep-20260721114705-b568m"
|
||||
|
||||
def test_seedance_2_5_alias(self):
|
||||
fn = self._import_target()
|
||||
assert fn("seedance-2.5") == "doubao-seedance-2-5-260628"
|
||||
|
||||
def test_seedance_2_0_alias(self):
|
||||
fn = self._import_target()
|
||||
assert fn("seedance-2.0") == "doubao-seedance-2-0-260128"
|
||||
|
||||
def test_seedance_2_0_fast_alias(self):
|
||||
fn = self._import_target()
|
||||
assert fn("seedance-2.0-fast") == "doubao-seedance-2-0-fast-260128"
|
||||
|
||||
def test_seedance_2_0_mini_alias(self):
|
||||
fn = self._import_target()
|
||||
assert fn("seedance-2.0-mini") == "doubao-seedance-2-0-mini-260615"
|
||||
|
||||
def test_wan_3_0_returns_dashscope_provider(self):
|
||||
from packages.shared.ai_client import _resolve_video_provider_and_id
|
||||
|
||||
prov, mid, cfg = _resolve_video_provider_and_id("wan-3.0")
|
||||
assert prov == "dashscope"
|
||||
assert mid == "wan3.0-video"
|
||||
assert cfg.get("billing_mode") == "per_second"
|
||||
|
||||
def test_seedance_2_5_uppercase(self):
|
||||
fn = self._import_target()
|
||||
assert fn("Seedance-2.5") == "doubao-seedance-2-5-260628"
|
||||
|
||||
def test_seedance_dot_normalize(self):
|
||||
fn = self._import_target()
|
||||
# dot 形式 "seedance-2.5" 直接命中 domain config 的 key(与 2-5 同等)
|
||||
assert fn("seedance-2.5") == "doubao-seedance-2-5-260628"
|
||||
|
||||
def test_unknown_model_falls_back_to_default_seedance_2_5(self, caplog):
|
||||
fn = self._import_target()
|
||||
import logging
|
||||
|
||||
# 未知 model key 会通过 get_viral_video_model_config 回落到 seedance-2.5
|
||||
with caplog.at_level(logging.WARNING, logger="shared.ai_client"):
|
||||
assert fn("some-random-model") == "doubao-seedance-2-5-260628"
|
||||
|
||||
|
||||
# ── #2165 详细错误信息和 last_video_error ─────────────────────────
|
||||
|
||||
|
||||
class TestVideoGenerationLastError:
|
||||
def test_create_400_portrait_returns_user_message(self, tmp_path):
|
||||
"""#2169: HTTP 400 + 真人拦截关键词 → 自动尝试即梦兜底;即梦未配时返回 portrait_intercept。"""
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.status_code = 400
|
||||
create_resp.text = '{"error":{"code":"ContentRisk","message":"Real person face detected in reference image, portrait blocked"}}'
|
||||
create_resp.json.return_value = {"error": {"code": "ContentRisk", "message": "..."}}
|
||||
create_resp.raise_for_status.side_effect = httpx.HTTPStatusError(
|
||||
"bad", request=MagicMock(), response=create_resp
|
||||
)
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
# jimeng 未配置,fallback 后仍返回 portrait_intercept(提示用户需要配置即梦)
|
||||
patch("packages.shared.jimeng_client.get_jimeng_client", return_value=None),
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=1, doubao_video_model="seedance"
|
||||
)
|
||||
result = client.video_generation("p", output_dir=str(tmp_path), image_url="https://img/x.jpg")
|
||||
assert result is None
|
||||
err = client.get_last_video_error()
|
||||
assert err["error_code"] == "portrait_intercept"
|
||||
# 即梦兜底未启用时提示包含"真人照片"/"即梦"等关键字
|
||||
assert "真人" in err["user_message"] or "肖像" in err["user_message"] or "即梦" in err["user_message"]
|
||||
# 方舟本身 status_code=400(最后一个错误可能被即梦兜底覆盖,但 error_code 不变)
|
||||
assert err["status_code"] in (0, 400)
|
||||
|
||||
def test_create_401_returns_auth_error(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.status_code = 401
|
||||
create_resp.text = '{"error":{"message":"Unauthorized"}}'
|
||||
create_resp.json.return_value = {"error": {"message": "Unauthorized"}}
|
||||
create_resp.raise_for_status.side_effect = httpx.HTTPStatusError(
|
||||
"auth", request=MagicMock(), response=create_resp
|
||||
)
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=1, doubao_video_model="seedance"
|
||||
)
|
||||
result = client.video_generation("p", output_dir=str(tmp_path))
|
||||
assert result is None
|
||||
err = client.get_last_video_error()
|
||||
assert err["error_code"] == "auth_error"
|
||||
assert err["status_code"] == 401
|
||||
|
||||
def test_poll_failed_returns_task_failed_error(self, tmp_path):
|
||||
"""轮询 status=failed 时应记录 task_failed 错误并含 detail。"""
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.status_code = 200
|
||||
create_resp.json.return_value = {"id": "t-fail"}
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
poll_resp = MagicMock()
|
||||
poll_resp.status_code = 200
|
||||
poll_resp.json.return_value = {
|
||||
"status": "failed",
|
||||
"error": {"code": "InvalidParam", "message": "resolution invalid"},
|
||||
}
|
||||
poll_resp.raise_for_status = MagicMock()
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.httpx.get", return_value=poll_resp),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=10, doubao_video_model="seedance"
|
||||
)
|
||||
result = client.video_generation("p", output_dir=str(tmp_path))
|
||||
assert result is None
|
||||
err = client.get_last_video_error()
|
||||
assert err["error_code"] == "task_failed"
|
||||
assert "InvalidParam" in err.get("detail", "") or err["status_code"] == 200
|
||||
|
||||
|
||||
class TestAiServiceLastVideoError:
|
||||
def test_call_video_generation_returns_none_sets_error(self):
|
||||
"""失败后 get_last_video_error 应返回结构化错误信息。"""
|
||||
from packages.shared import ai_service
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = True
|
||||
mock_client.last_video_error = {"error_code": "unknown", "user_message": "test"}
|
||||
mock_client.get_last_video_error.return_value = {"error_code": "unknown", "user_message": "test"}
|
||||
mock_client.video_generation.return_value = None
|
||||
with patch("packages.shared.ai_service.get_doubao_client", return_value=mock_client):
|
||||
assert ai_service.call_video_generation("p") is None
|
||||
err = ai_service.get_last_video_error()
|
||||
assert err["error_code"] == "unknown"
|
||||
assert "user_message" in err
|
||||
|
||||
@@ -0,0 +1,198 @@
|
||||
"""catalog 应用服务单测:会员套餐 / 积分包从共享库读取与字段映射。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_cache():
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
admin_catalog._cache.clear()
|
||||
yield
|
||||
admin_catalog._cache.clear()
|
||||
|
||||
|
||||
def _row(**kw):
|
||||
row = MagicMock()
|
||||
for k, v in kw.items():
|
||||
setattr(row, k, v)
|
||||
return row
|
||||
|
||||
|
||||
class TestMembershipPlans:
|
||||
def test_yearly_plan_mapping(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
plan_key="premium_yearly",
|
||||
name="高级会员年卡",
|
||||
description="年度订阅",
|
||||
monthly_price=0,
|
||||
yearly_price=399,
|
||||
quotas={"4k": True, "batch_render": True, "credits_per_month": 500},
|
||||
display_order=1,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
plans = admin_catalog.get_membership_plans()
|
||||
|
||||
assert len(plans) == 1
|
||||
p = plans[0]
|
||||
assert p["plan_id"] == "premium_yearly"
|
||||
assert p["billing_cycle"] == "yearly"
|
||||
assert p["price_cents"] == 39900
|
||||
assert p["monthly_price_cents"] == 3325
|
||||
assert p["duration_days"] == 365
|
||||
assert p["features"]["4K 超清分辨率"] is True
|
||||
assert p["features"]["credits_per_month"] == 500
|
||||
session.close.assert_called_once()
|
||||
|
||||
def test_monthly_plan_mapping(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
plan_key="premium_monthly",
|
||||
name="高级会员月卡",
|
||||
description=None,
|
||||
monthly_price=39,
|
||||
yearly_price=0,
|
||||
quotas=None,
|
||||
display_order=2,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
plans = admin_catalog.get_membership_plans()
|
||||
|
||||
assert len(plans) == 1
|
||||
p = plans[0]
|
||||
assert p["billing_cycle"] == "monthly"
|
||||
assert p["price_cents"] == 3900
|
||||
assert p["monthly_price_cents"] == 3900
|
||||
assert p["duration_days"] == 30
|
||||
assert p["features"] == {}
|
||||
|
||||
def test_both_cycles_expanded(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
plan_key="premium",
|
||||
name="高级会员",
|
||||
description=None,
|
||||
monthly_price=39,
|
||||
yearly_price=399,
|
||||
quotas={},
|
||||
display_order=1,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
plans = admin_catalog.get_membership_plans()
|
||||
|
||||
cycles = {p["billing_cycle"] for p in plans}
|
||||
assert cycles == {"yearly", "monthly"}
|
||||
|
||||
def test_no_session_returns_empty(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", None, create=True):
|
||||
assert admin_catalog.get_membership_plans() == []
|
||||
|
||||
|
||||
class TestPointsPackages:
|
||||
def test_package_mapping_with_bonus(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
package_key="pkg_100",
|
||||
name="100元充值包",
|
||||
price=100,
|
||||
credits=1000,
|
||||
bonus_credits=100,
|
||||
is_recommended=True,
|
||||
description="推荐",
|
||||
sort_order=4,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
packages = admin_catalog.get_points_packages()
|
||||
|
||||
assert len(packages) == 1
|
||||
pkg = packages[0]
|
||||
assert pkg["code"] == "pkg_100"
|
||||
assert pkg["points"] == 1100
|
||||
assert pkg["price_cents"] == 10000
|
||||
assert pkg["is_recommended"] is True
|
||||
assert pkg["unit_price"] == "¥0.091/积分"
|
||||
|
||||
def test_zero_credits_unit_price_safe(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
package_key="pkg_0",
|
||||
name="空包",
|
||||
price=0,
|
||||
credits=0,
|
||||
bonus_credits=0,
|
||||
is_recommended=False,
|
||||
description=None,
|
||||
sort_order=0,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
packages = admin_catalog.get_points_packages()
|
||||
|
||||
assert packages[0]["points"] == 0
|
||||
assert packages[0]["price_cents"] == 0
|
||||
assert packages[0]["unit_price"] == "¥0.000/积分"
|
||||
|
||||
def test_no_session_returns_empty(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", None, create=True):
|
||||
assert admin_catalog.get_points_packages() == []
|
||||
|
||||
|
||||
class TestPackagesRoute:
|
||||
def test_get_packages_route_returns_items(self):
|
||||
from app.api.routes.points import get_packages
|
||||
|
||||
cu = MagicMock()
|
||||
cu.user.member_type = None
|
||||
rows = [
|
||||
{
|
||||
"code": "pkg_10",
|
||||
"name": "10元充值包",
|
||||
"points": 100,
|
||||
"price_cents": 1000,
|
||||
"unit_price": "¥0.100/积分",
|
||||
}
|
||||
]
|
||||
with patch(
|
||||
"packages.application.catalog.admin_catalog.get_points_packages",
|
||||
return_value=rows,
|
||||
):
|
||||
resp = get_packages(current_user=cu)
|
||||
|
||||
assert len(resp.packages) == 1
|
||||
item = resp.packages[0]
|
||||
assert item.code == "pkg_10"
|
||||
assert item.points == 100
|
||||
assert item.price_cents == 1000
|
||||
@@ -0,0 +1,178 @@
|
||||
"""tests for packages/shared/dashscope_client.py (#2159 Wan 3.0 DashScope client)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, mock_open, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
_SINGLETON = "_DASHSCOPE_CLIENT_SINGLETON"
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_singleton():
|
||||
import packages.shared.dashscope_client as d
|
||||
|
||||
# 兼容实际 singleton 名
|
||||
for name in ("_DASHSCOPE_CLIENT_SINGLETON", "_dashscope_client"):
|
||||
if hasattr(d, name):
|
||||
setattr(d, name, None)
|
||||
yield
|
||||
for name in ("_DASHSCOPE_CLIENT_SINGLETON", "_dashscope_client"):
|
||||
if hasattr(d, name):
|
||||
setattr(d, name, None)
|
||||
|
||||
|
||||
def _make_settings(api_key="test-key"):
|
||||
return MagicMock(
|
||||
dashscope_api_key=api_key,
|
||||
dashscope_base_url="https://dashscope.aliyuncs.com/api/v1",
|
||||
dashscope_video_timeout=10,
|
||||
dashscope_video_poll_interval=0,
|
||||
video_dir="/tmp/videos",
|
||||
)
|
||||
|
||||
|
||||
class TestDashScopeAvailability:
|
||||
def test_unavailable_without_key(self):
|
||||
from packages.shared.dashscope_client import get_dashscope_client
|
||||
|
||||
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings(api_key="")
|
||||
assert get_dashscope_client() is None
|
||||
|
||||
def test_available_with_key(self):
|
||||
from packages.shared.dashscope_client import get_dashscope_client
|
||||
|
||||
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = get_dashscope_client()
|
||||
assert c is not None
|
||||
assert c.is_available is True
|
||||
|
||||
|
||||
def _mock_stream_response(min_size=2048):
|
||||
"""构造 httpx.stream 上下文返回值,模拟返回若干字节的 mp4 内容。"""
|
||||
m = MagicMock()
|
||||
m.status_code = 200
|
||||
chunk = b"x" * min_size
|
||||
m.iter_bytes.return_value = [chunk]
|
||||
ctx = MagicMock()
|
||||
ctx.__enter__.return_value = m
|
||||
return ctx
|
||||
|
||||
|
||||
class TestDashScopeVideoGeneration:
|
||||
def test_happy_path_returns_video_path(self):
|
||||
"""POST create → GET poll (SUCCEEDED) → download → returns path + correct payload."""
|
||||
import packages.shared.dashscope_client as d
|
||||
|
||||
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = d.DashScopeClient()
|
||||
|
||||
create_resp = MagicMock(status_code=200)
|
||||
create_resp.json.return_value = {"output": {"task_id": "task-abc"}}
|
||||
poll_resp = MagicMock(status_code=200)
|
||||
poll_resp.json.return_value = {
|
||||
"output": {"task_status": "SUCCEEDED", "video_url": "http://x/y.mp4"},
|
||||
"usage": {"billed_duration": 10},
|
||||
}
|
||||
# fake file: write enough bytes to pass the size>=1024 check
|
||||
m_open = mock_open()
|
||||
m_open.return_value.write.return_value = None
|
||||
fake_size = {"/tmp/videos/wan_task-abc.mp4": 4096}
|
||||
|
||||
def fake_getsize(p):
|
||||
return fake_size.get(p, 0)
|
||||
|
||||
def fake_exists(p):
|
||||
return p in fake_size
|
||||
|
||||
with (
|
||||
patch.object(d.httpx, "post", return_value=create_resp) as mock_post,
|
||||
patch.object(d.httpx, "get", return_value=poll_resp),
|
||||
patch.object(d.httpx, "stream", return_value=_mock_stream_response()),
|
||||
patch("packages.shared.dashscope_client.time.sleep"),
|
||||
patch("packages.shared.dashscope_client.os.makedirs"),
|
||||
patch("builtins.open", m_open),
|
||||
patch("packages.shared.dashscope_client.os.path.getsize", side_effect=fake_getsize),
|
||||
patch("packages.shared.dashscope_client.os.path.exists", side_effect=fake_exists),
|
||||
):
|
||||
res = c.video_generation(
|
||||
prompt="test",
|
||||
duration=5,
|
||||
ratio="9:16",
|
||||
resolution="720p",
|
||||
output_dir="/tmp/videos",
|
||||
)
|
||||
assert res is not None, "expected success"
|
||||
assert res["video_path"] == "/tmp/videos/wan_task-abc.mp4"
|
||||
_, kwargs = mock_post.call_args
|
||||
body = kwargs["json"]
|
||||
assert body["parameters"]["resolution"] == "720P"
|
||||
assert body["model"] == "wan3.0-video"
|
||||
|
||||
def test_create_http_error_returns_none(self):
|
||||
import packages.shared.dashscope_client as d
|
||||
|
||||
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = d.DashScopeClient()
|
||||
err_resp = MagicMock(status_code=400, text="bad")
|
||||
err_resp.raise_for_status.side_effect = RuntimeError("bad")
|
||||
with patch.object(d.httpx, "post", return_value=err_resp):
|
||||
res = c.video_generation(
|
||||
prompt="test", duration=5, ratio="9:16", resolution="720p", output_dir="/tmp/videos"
|
||||
)
|
||||
assert res is None
|
||||
|
||||
def test_poll_failed_returns_none(self):
|
||||
import packages.shared.dashscope_client as d
|
||||
|
||||
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = d.DashScopeClient()
|
||||
create_resp = MagicMock(status_code=200)
|
||||
create_resp.json.return_value = {"output": {"task_id": "task-abc"}}
|
||||
poll_resp = MagicMock(status_code=200)
|
||||
poll_resp.json.return_value = {"output": {"task_status": "FAILED", "message": "nope"}}
|
||||
with (
|
||||
patch.object(d.httpx, "post", return_value=create_resp),
|
||||
patch.object(d.httpx, "get", return_value=poll_resp),
|
||||
patch("packages.shared.dashscope_client.time.sleep"),
|
||||
):
|
||||
res = c.video_generation(
|
||||
prompt="test", duration=5, ratio="9:16", resolution="720p", output_dir="/tmp/videos"
|
||||
)
|
||||
assert res is None
|
||||
|
||||
def test_empty_prompt_returns_none(self):
|
||||
import packages.shared.dashscope_client as d
|
||||
|
||||
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = d.DashScopeClient()
|
||||
assert (
|
||||
c.video_generation(prompt=" ", duration=5, ratio="9:16", resolution="720p", output_dir="/tmp/videos")
|
||||
is None
|
||||
)
|
||||
|
||||
def test_create_400_sets_last_video_error(self, tmp_path):
|
||||
"""创建任务 HTTP 400 时应写 last_video_error。"""
|
||||
from packages.shared import dashscope_client as dc
|
||||
|
||||
dc._DASHSCOPE_CLIENT_SINGLETON = None
|
||||
with patch.dict("os.environ", {"DASHSCOPE_API_KEY": "test-key"}):
|
||||
c = dc.DashScopeClient()
|
||||
r = MagicMock()
|
||||
r.status_code = 401
|
||||
r.text = '{"code":"InvalidApiKey","message":"bad key"}'
|
||||
r.raise_for_status.side_effect = httpx.HTTPStatusError("auth", request=MagicMock(), response=r)
|
||||
with patch.object(dc.httpx, "post", return_value=r), patch.object(dc, "time"):
|
||||
out = c.video_generation("p", output_dir=str(tmp_path))
|
||||
assert out is None
|
||||
err = c.get_last_video_error()
|
||||
assert err["error_code"] == "auth_error"
|
||||
assert c.last_video_error is not None
|
||||
@@ -0,0 +1,376 @@
|
||||
"""tests for packages/shared/jimeng_client.py (#2169 即梦 i2v 客户端)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from unittest.mock import MagicMock, mock_open, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_singleton():
|
||||
import packages.shared.jimeng_client as j
|
||||
|
||||
j._JIMENG_CLIENT_SINGLETON = None
|
||||
yield
|
||||
j._JIMENG_CLIENT_SINGLETON = None
|
||||
|
||||
|
||||
def _make_settings(ak="test-ak", sk="test-sk", req_key="jimeng_i2v_first_v30", timeout=60, poll_interval=2):
|
||||
return MagicMock(
|
||||
jimeng_ak=ak,
|
||||
jimeng_sk=sk,
|
||||
jimeng_base_url="https://visual.volcengineapi.com",
|
||||
jimeng_req_key=req_key,
|
||||
jimeng_video_timeout=timeout,
|
||||
jimeng_video_poll_interval=poll_interval,
|
||||
)
|
||||
|
||||
|
||||
# ── V4 签名单元测试 ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestV4Signature:
|
||||
def test_sign_returns_required_headers(self):
|
||||
from packages.shared.jimeng_client import _sign_v4
|
||||
|
||||
headers = _sign_v4(
|
||||
ak="AK_TEST",
|
||||
sk="SK_TEST",
|
||||
method="POST",
|
||||
host="visual.volcengineapi.com",
|
||||
query={"Action": "CVSync2AsyncSubmitTask", "Version": "2022-08-31"},
|
||||
body_bytes=b'{"req_key":"jimeng_i2v_first_v30"}',
|
||||
x_date="20260101T120000Z",
|
||||
)
|
||||
assert headers["Content-Type"] == "application/json"
|
||||
assert headers["Host"] == "visual.volcengineapi.com"
|
||||
assert headers["X-Date"] == "20260101T120000Z"
|
||||
assert "X-Content-Sha256" in headers
|
||||
assert headers["Authorization"].startswith("HMAC-SHA256 Credential=AK_TEST/20260101/cn-north-1/cv/request")
|
||||
assert "SignedHeaders=content-type;host;x-content-sha256;x-date" in headers["Authorization"]
|
||||
assert "Signature=" in headers["Authorization"]
|
||||
# 签名是 64 字符 hex
|
||||
sig = headers["Authorization"].split("Signature=")[-1]
|
||||
assert len(sig) == 64
|
||||
assert all(c in "0123456789abcdef" for c in sig)
|
||||
|
||||
def test_sign_deterministic(self):
|
||||
"""相同输入必须产生相同签名(幂等)。"""
|
||||
from packages.shared.jimeng_client import _sign_v4
|
||||
|
||||
kwargs = dict(
|
||||
ak="AK",
|
||||
sk="SK",
|
||||
method="POST",
|
||||
host="h",
|
||||
query={"A": "1", "B": "2"},
|
||||
body_bytes=b"{}",
|
||||
x_date="20260101T000000Z",
|
||||
)
|
||||
h1 = _sign_v4(**kwargs)
|
||||
h2 = _sign_v4(**kwargs)
|
||||
assert h1["Authorization"] == h2["Authorization"]
|
||||
assert h1["X-Content-Sha256"] == h2["X-Content-Sha256"]
|
||||
|
||||
def test_sign_different_body_different_sig(self):
|
||||
from packages.shared.jimeng_client import _sign_v4
|
||||
|
||||
base = dict(ak="AK", sk="SK", method="POST", host="h", query={}, x_date="20260101T000000Z")
|
||||
h1 = _sign_v4(body_bytes=b"a", **base)
|
||||
h2 = _sign_v4(body_bytes=b"b", **base)
|
||||
assert h1["Authorization"] != h2["Authorization"]
|
||||
|
||||
def test_payload_sha256_matches(self):
|
||||
import hashlib
|
||||
|
||||
from packages.shared.jimeng_client import _sign_v4
|
||||
|
||||
body = b'{"prompt":"hello"}'
|
||||
h = _sign_v4("ak", "sk", "POST", "h", {}, body, "20260101T000000Z")
|
||||
expected = hashlib.sha256(body).hexdigest()
|
||||
assert h["X-Content-Sha256"] == expected
|
||||
|
||||
def test_norm_query_sorted_and_encoded(self):
|
||||
from packages.shared.jimeng_client import _norm_query
|
||||
|
||||
q = _norm_query({"B": "2", "A": "1", "C": "a b"})
|
||||
# key 排序 + 空格→%20
|
||||
assert q.startswith("A=1")
|
||||
assert "B=2" in q
|
||||
assert "C=a%20b" in q
|
||||
|
||||
|
||||
# ── 可用性 / 单例 ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestAvailability:
|
||||
def test_unavailable_without_ak_sk(self):
|
||||
from packages.shared.jimeng_client import JimengClient, get_jimeng_client
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings(ak="", sk="")
|
||||
# 重置单例
|
||||
import packages.shared.jimeng_client as j
|
||||
|
||||
j._JIMENG_CLIENT_SINGLETON = None
|
||||
assert get_jimeng_client() is None
|
||||
|
||||
def test_available_with_ak_sk(self):
|
||||
from packages.shared.jimeng_client import get_jimeng_client
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
import packages.shared.jimeng_client as j
|
||||
|
||||
j._JIMENG_CLIENT_SINGLETON = None
|
||||
c = get_jimeng_client()
|
||||
assert c is not None
|
||||
assert c.is_available is True
|
||||
assert c.req_key == "jimeng_i2v_first_v30"
|
||||
|
||||
|
||||
# ── 错误分类 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestClassifyError:
|
||||
def test_50411_is_portrait_intercept_non_retryable(self):
|
||||
from packages.shared.jimeng_client import _classify_jimeng_error
|
||||
|
||||
code, msg, retry = _classify_jimeng_error(200, '{"code":50411,"message":"Pre Img Risk"}', 50411)
|
||||
assert code == "portrait_intercept"
|
||||
assert retry is False
|
||||
|
||||
def test_50429_is_rate_limit_retryable(self):
|
||||
from packages.shared.jimeng_client import _classify_jimeng_error
|
||||
|
||||
code, msg, retry = _classify_jimeng_error(200, "", 50429)
|
||||
assert code == "rate_limit"
|
||||
assert retry is True
|
||||
|
||||
def test_50430_is_rate_limit(self):
|
||||
from packages.shared.jimeng_client import _classify_jimeng_error
|
||||
|
||||
code, _, _ = _classify_jimeng_error(200, "", 50430)
|
||||
assert code == "rate_limit"
|
||||
|
||||
def test_50500_is_network_error_retryable(self):
|
||||
from packages.shared.jimeng_client import _classify_jimeng_error
|
||||
|
||||
code, _, retry = _classify_jimeng_error(200, "", 50500)
|
||||
assert code == "network_error"
|
||||
assert retry is True
|
||||
|
||||
def test_50412_is_invalid_param_non_retryable(self):
|
||||
from packages.shared.jimeng_client import _classify_jimeng_error
|
||||
|
||||
code, _, retry = _classify_jimeng_error(200, "", 50412)
|
||||
assert code == "invalid_param"
|
||||
assert retry is False
|
||||
|
||||
def test_401_auth(self):
|
||||
from packages.shared.jimeng_client import _classify_jimeng_error
|
||||
|
||||
code, msg, retry = _classify_jimeng_error(401, "auth fail", None)
|
||||
assert code == "auth_error"
|
||||
assert retry is False
|
||||
|
||||
def test_400_text_audit(self):
|
||||
from packages.shared.jimeng_client import _classify_jimeng_error
|
||||
|
||||
code, _, _ = _classify_jimeng_error(400, "text error", None)
|
||||
assert code == "invalid_param"
|
||||
|
||||
|
||||
# ── video_generation 主流程 ──────────────────────────────────────────
|
||||
|
||||
|
||||
class TestVideoGenerationHappyPath:
|
||||
def test_missing_ak_returns_none(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings(ak="", sk="")
|
||||
c = JimengClient()
|
||||
assert c.video_generation("hi", image_url="http://x/y.jpg") is None
|
||||
err = c.last_video_error
|
||||
assert err["error_code"] == "auth_error"
|
||||
|
||||
def test_empty_prompt_returns_none(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = JimengClient()
|
||||
assert c.video_generation(" ", image_url="http://x/y.jpg") is None
|
||||
assert c.last_video_error["error_code"] == "invalid_param"
|
||||
|
||||
def test_empty_image_url_returns_none(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = JimengClient()
|
||||
assert c.video_generation("prompt", image_url="") is None
|
||||
assert c.last_video_error["error_code"] == "invalid_param"
|
||||
|
||||
def test_duration_5s_frames_121(self):
|
||||
"""5s → frames=121,10s→frames=241,>10s 截断到10s。"""
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings(timeout=1, poll_interval=0)
|
||||
c = JimengClient()
|
||||
|
||||
captured = {}
|
||||
|
||||
def fake_submit(prompt, image_url, frames, seed=-1):
|
||||
captured["frames"] = frames
|
||||
return "task-xyz"
|
||||
|
||||
def fake_poll(tid):
|
||||
captured["tid"] = tid
|
||||
return "http://example.com/v.mp4"
|
||||
|
||||
def fake_download(url, out_dir, tid):
|
||||
captured["url"] = url
|
||||
return "/tmp/fake.mp4"
|
||||
|
||||
# 构造一个假文件
|
||||
os.makedirs("/tmp", exist_ok=True)
|
||||
with open("/tmp/fake.mp4", "wb") as f:
|
||||
f.write(b"x" * 2048)
|
||||
|
||||
with (
|
||||
patch.object(c, "_submit_task", side_effect=fake_submit),
|
||||
patch.object(c, "_poll_result", side_effect=fake_poll),
|
||||
patch.object(c, "_download_video", side_effect=fake_download),
|
||||
):
|
||||
r = c.video_generation("test", image_url="http://x/y.jpg", duration=5, output_dir="/tmp")
|
||||
assert r is not None
|
||||
assert captured["frames"] == 121
|
||||
assert r["usage"]["duration_seconds"] == 5
|
||||
assert r["usage"]["billing_mode"] == "per_second"
|
||||
assert r["usage"]["req_key"] == "jimeng_i2v_first_v30"
|
||||
|
||||
def test_duration_10s_frames_241(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings(timeout=1, poll_interval=0)
|
||||
c = JimengClient()
|
||||
captured = {}
|
||||
|
||||
def fake_submit(prompt, image_url, frames, seed=-1):
|
||||
captured["frames"] = frames
|
||||
return "tid"
|
||||
|
||||
def fake_poll(tid):
|
||||
return "http://x/v.mp4"
|
||||
|
||||
def fake_download(url, out_dir, tid):
|
||||
with open("/tmp/fake2.mp4", "wb") as f:
|
||||
f.write(b"x" * 2048)
|
||||
return "/tmp/fake2.mp4"
|
||||
|
||||
with (
|
||||
patch.object(c, "_submit_task", side_effect=fake_submit),
|
||||
patch.object(c, "_poll_result", side_effect=fake_poll),
|
||||
patch.object(c, "_download_video", side_effect=fake_download),
|
||||
):
|
||||
r = c.video_generation("hi", image_url="http://x/y.jpg", duration=10, output_dir="/tmp")
|
||||
assert captured["frames"] == 241
|
||||
assert r["usage"]["duration_seconds"] == 10
|
||||
|
||||
def test_duration_over_10s_truncates_to_10s(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings(timeout=1, poll_interval=0)
|
||||
c = JimengClient()
|
||||
captured = {}
|
||||
|
||||
def fake_submit(prompt, image_url, frames, seed=-1):
|
||||
captured["frames"] = frames
|
||||
return "tid"
|
||||
|
||||
def fake_poll(tid):
|
||||
return "http://x/v.mp4"
|
||||
|
||||
def fake_download(url, out_dir, tid):
|
||||
with open("/tmp/fake3.mp4", "wb") as f:
|
||||
f.write(b"x" * 2048)
|
||||
return "/tmp/fake3.mp4"
|
||||
|
||||
with (
|
||||
patch.object(c, "_submit_task", side_effect=fake_submit),
|
||||
patch.object(c, "_poll_result", side_effect=fake_poll),
|
||||
patch.object(c, "_download_video", side_effect=fake_download),
|
||||
):
|
||||
r = c.video_generation("hi", image_url="http://x/y.jpg", duration=30, output_dir="/tmp")
|
||||
assert captured["frames"] == 241
|
||||
assert r["usage"]["duration_seconds"] == 10
|
||||
|
||||
|
||||
class TestSubmitTaskErrors:
|
||||
def test_submit_50411_writes_portrait_intercept(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = JimengClient()
|
||||
fake_resp = MagicMock(status_code=200, text='{"code":50411,"message":"Pre Img Risk Not Pass"}')
|
||||
fake_resp.json.return_value = {"code": 50411, "message": "Pre Img Risk Not Pass"}
|
||||
with patch("packages.shared.jimeng_client.httpx.request", return_value=fake_resp):
|
||||
tid = c._submit_task("p", "http://x/y.jpg", frames=121)
|
||||
assert tid is None
|
||||
assert c.last_video_error["error_code"] == "portrait_intercept"
|
||||
|
||||
def test_submit_returns_task_id(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = JimengClient()
|
||||
fake_resp = MagicMock(status_code=200, text='{"code":10000,"data":{"task_id":"abc"}}')
|
||||
fake_resp.json.return_value = {"code": 10000, "data": {"task_id": "abc"}}
|
||||
with patch("packages.shared.jimeng_client.httpx.request", return_value=fake_resp):
|
||||
tid = c._submit_task("p", "http://x/y.jpg", frames=121)
|
||||
assert tid == "abc"
|
||||
|
||||
|
||||
class TestPollResult:
|
||||
def test_poll_done_returns_video_url(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings(timeout=10, poll_interval=0)
|
||||
c = JimengClient()
|
||||
done_resp = MagicMock(status_code=200)
|
||||
done_resp.json.return_value = {"code": 10000, "data": {"status": "done", "video_url": "http://x/v.mp4"}}
|
||||
with (
|
||||
patch("packages.shared.jimeng_client.httpx.request", return_value=done_resp),
|
||||
patch("packages.shared.jimeng_client.time.sleep"),
|
||||
):
|
||||
url = c._poll_result("abc")
|
||||
assert url == "http://x/v.mp4"
|
||||
|
||||
def test_poll_timeout_returns_none(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings(timeout=1, poll_interval=0)
|
||||
c = JimengClient()
|
||||
queue_resp = MagicMock(status_code=200)
|
||||
queue_resp.json.return_value = {"code": 10000, "data": {"status": "in_queue"}}
|
||||
# time.time 会被调用,模拟超时
|
||||
with (
|
||||
patch("packages.shared.jimeng_client.httpx.request", return_value=queue_resp),
|
||||
patch("packages.shared.jimeng_client.time.sleep"),
|
||||
):
|
||||
url = c._poll_result("abc")
|
||||
assert url is None
|
||||
assert c.last_video_error["error_code"] == "network_error"
|
||||
assert "超时" in c.last_video_error["user_message"]
|
||||
@@ -31,94 +31,49 @@ def _make_cu(user_id="user-1", is_member=False, member_type=None):
|
||||
|
||||
|
||||
class TestRechargeOrderResponse:
|
||||
"""积分包下单接口(create-order + 兼容 recharge 别名)。"""
|
||||
|
||||
def _patch_payment(self, result=None, side_effect=None):
|
||||
from packages.application.payment_service import PaymentService
|
||||
|
||||
inst = MagicMock()
|
||||
if side_effect is not None:
|
||||
inst.create_points_order.side_effect = side_effect
|
||||
else:
|
||||
inst.create_points_order.return_value = result
|
||||
return patch.object(PaymentService, "__new__", return_value=inst), inst
|
||||
|
||||
def test_create_order_returns_pay_params(self):
|
||||
"""create-order 响应必须包含 pay_params / points_amount / expire_at。"""
|
||||
from app.api.routes.points import create_points_purchase_order
|
||||
from app.schemas.points import PointsRechargeRequest
|
||||
|
||||
result = {
|
||||
"order_id": "order-1",
|
||||
"out_trade_no": "pt123",
|
||||
"prepay_id": "prepay-1",
|
||||
"amount_cents": 990,
|
||||
"points_amount": 100,
|
||||
"expire_at": (datetime.now(UTC) + timedelta(hours=48)).isoformat(),
|
||||
"pay_params": {"appId": "wx", "paySign": "s"},
|
||||
}
|
||||
p, _inst = self._patch_payment(result=result)
|
||||
cu = _make_cu()
|
||||
cu.user.wechat_openid = "o-1"
|
||||
body = PointsRechargeRequest(package_id="starter_pack")
|
||||
|
||||
with p:
|
||||
resp = create_points_purchase_order(body=body, current_user=cu, db=MagicMock())
|
||||
|
||||
assert resp.points_amount == 100
|
||||
assert resp.pay_params["appId"] == "wx"
|
||||
assert resp.expire_at == result["expire_at"]
|
||||
assert resp.id == "order-1"
|
||||
|
||||
def test_create_order_invalid_package_returns_400(self):
|
||||
from app.api.routes.points import create_points_purchase_order
|
||||
from app.schemas.points import PointsRechargeRequest
|
||||
|
||||
p, _inst = self._patch_payment(side_effect=ValueError("积分包不存在: x"))
|
||||
cu = _make_cu()
|
||||
cu.user.wechat_openid = "o-1"
|
||||
body = PointsRechargeRequest(package_id="x")
|
||||
|
||||
with p, pytest.raises(HTTPException) as exc:
|
||||
create_points_purchase_order(body=body, current_user=cu, db=MagicMock())
|
||||
assert exc.value.status_code == 400
|
||||
|
||||
def test_recharge_alias_forwards_to_create_order(self):
|
||||
"""旧 /recharge 端点保留且内部转发。"""
|
||||
def test_recharge_returns_pay_params_points_amount_expire_at(self):
|
||||
"""recharge 响应必须包含 pay_params / points_amount / expire_at。"""
|
||||
from app.api.routes.points import create_recharge_order
|
||||
from app.schemas.points import PointsRechargeRequest
|
||||
|
||||
result = {
|
||||
"order_id": "order-2",
|
||||
"out_trade_no": "pt2",
|
||||
"prepay_id": "p",
|
||||
svc = MagicMock()
|
||||
svc.create_order.return_value = {
|
||||
"id": "order-1",
|
||||
"order_type": "points",
|
||||
"product_code": "starter_pack",
|
||||
"amount_cents": 990,
|
||||
"points_amount": 100,
|
||||
"expire_at": "2026-10-05T00:00:00+00:00",
|
||||
"pay_params": {"appId": "wx"},
|
||||
"status": "pending",
|
||||
"created_at": datetime.now(UTC).isoformat(),
|
||||
}
|
||||
p, _inst = self._patch_payment(result=result)
|
||||
db = MagicMock()
|
||||
cu = _make_cu()
|
||||
cu.user.wechat_openid = "o-1"
|
||||
body = PointsRechargeRequest(package_id="starter_pack")
|
||||
with p:
|
||||
resp = create_recharge_order(body=body, current_user=cu, db=MagicMock())
|
||||
assert resp.id == "order-2"
|
||||
|
||||
def test_create_order_payment_disabled_returns_503(self):
|
||||
"""微信未配置时返回 503。"""
|
||||
from app.api.routes.points import create_points_purchase_order
|
||||
before = datetime.now(UTC)
|
||||
with patch("app.api.routes.points._get_service", return_value=svc):
|
||||
resp = create_recharge_order(body=body, current_user=cu, db=db)
|
||||
after = datetime.now(UTC) + timedelta(hours=48)
|
||||
|
||||
assert resp.points_amount == 100 # starter_pack 100 分
|
||||
assert isinstance(resp.pay_params, dict)
|
||||
assert resp.expire_at is not None
|
||||
expire_dt = datetime.fromisoformat(resp.expire_at)
|
||||
assert expire_dt >= before + timedelta(hours=47, minutes=55)
|
||||
assert expire_dt <= after
|
||||
|
||||
def test_recharge_invalid_package_returns_400(self):
|
||||
from app.api.routes.points import create_recharge_order
|
||||
from app.schemas.points import PointsRechargeRequest
|
||||
|
||||
from packages.application.payment_service import PaymentConfigError
|
||||
|
||||
p, _inst = self._patch_payment(side_effect=PaymentConfigError("未配置"))
|
||||
svc = MagicMock()
|
||||
svc.create_order.side_effect = ValueError("invalid package")
|
||||
db = MagicMock()
|
||||
cu = _make_cu()
|
||||
cu.user.wechat_openid = "o-1"
|
||||
body = PointsRechargeRequest(package_id="starter_pack")
|
||||
with p, pytest.raises(HTTPException) as exc:
|
||||
create_points_purchase_order(body=body, current_user=cu, db=MagicMock())
|
||||
assert exc.value.status_code == 503
|
||||
body = PointsRechargeRequest(package_id="nonexistent")
|
||||
|
||||
with pytest.raises(HTTPException) as exc, patch("app.api.routes.points._get_service", return_value=svc):
|
||||
create_recharge_order(body=body, current_user=cu, db=db)
|
||||
assert exc.value.status_code == 400
|
||||
|
||||
|
||||
# ── P0-2: check 任意 scene_key(已下线场景返回 cost=0) ──────────────────
|
||||
@@ -218,33 +173,46 @@ class TestSubscriptionPlans:
|
||||
_spec.loader.exec_module(_mod)
|
||||
return _mod.list_membership_plans
|
||||
|
||||
def test_plans_endpoint_returns_three_tiers(self):
|
||||
import os # noqa: F401 (used by _import_plans_fn)
|
||||
|
||||
def test_plans_endpoint_reads_admin_table(self):
|
||||
"""/subscription/plans 改读管理后台 plans 表:返回 catalog 服务提供的真实档位。"""
|
||||
list_membership_plans = self._import_plans_fn()
|
||||
resp = list_membership_plans(current_user=_make_cu())
|
||||
real_plan = {
|
||||
"plan_id": "premium_yearly",
|
||||
"billing_cycle": "yearly",
|
||||
"name": "高级会员年卡",
|
||||
"description": "高级会员年度订阅,享受全部功能",
|
||||
"price_cents": 39900,
|
||||
"monthly_price_cents": 3325,
|
||||
"duration_days": 365,
|
||||
"features": {
|
||||
"4K 超清分辨率": True,
|
||||
"批量渲染": True,
|
||||
"优先处理队列": True,
|
||||
"credits_per_month": 500,
|
||||
},
|
||||
}
|
||||
with patch(
|
||||
"packages.application.catalog.admin_catalog.get_membership_plans",
|
||||
return_value=[real_plan],
|
||||
):
|
||||
resp = list_membership_plans(current_user=_make_cu())
|
||||
plans = resp["plans"]
|
||||
plan_ids = {p["plan_id"] for p in plans}
|
||||
assert plan_ids == {"monthly", "quarterly", "yearly"}
|
||||
for p in plans:
|
||||
assert p["price_cents"] > 0
|
||||
assert p["duration_days"] in (30, 90, 365)
|
||||
assert 0 < p["points_discount"] <= 1.0
|
||||
assert "max_resolution" in p["features"]
|
||||
|
||||
def test_longer_plans_cheaper_per_month(self):
|
||||
import os # noqa: F401
|
||||
assert len(plans) == 1
|
||||
p0 = plans[0]
|
||||
assert p0["plan_id"] == "premium_yearly"
|
||||
assert p0["price_cents"] == 39900
|
||||
assert p0["duration_days"] == 365
|
||||
assert p0["features"]["4K 超清分辨率"] is True
|
||||
|
||||
def test_plans_endpoint_empty_when_all_disabled(self):
|
||||
"""后台停用全部套餐时,用户端返回空列表。"""
|
||||
list_membership_plans = self._import_plans_fn()
|
||||
resp = list_membership_plans(current_user=_make_cu())
|
||||
plans = resp["plans"]
|
||||
monthly = next(p for p in plans if p["plan_id"] == "monthly")
|
||||
quarterly = next(p for p in plans if p["plan_id"] == "quarterly")
|
||||
yearly = next(p for p in plans if p["plan_id"] == "yearly")
|
||||
assert monthly["monthly_price_cents"] == 1990
|
||||
assert quarterly["monthly_price_cents"] < monthly["monthly_price_cents"]
|
||||
assert yearly["monthly_price_cents"] < quarterly["monthly_price_cents"]
|
||||
|
||||
with patch(
|
||||
"packages.application.catalog.admin_catalog.get_membership_plans",
|
||||
return_value=[],
|
||||
):
|
||||
resp = list_membership_plans(current_user=_make_cu())
|
||||
assert resp["plans"] == []
|
||||
|
||||
# ── P1-7: multiplier consistency ──────────────────────────────────────
|
||||
|
||||
|
||||
@@ -196,10 +196,24 @@ class TestResolveVideoDimensions:
|
||||
"""未知分辨率字符串兜底到 720p。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
w, h = resolve_video_dimensions("2160p", "1:1")
|
||||
w, h = resolve_video_dimensions("garbage-xxx", "1:1")
|
||||
assert h == 720
|
||||
assert w == 720
|
||||
|
||||
def test_4k_16_9(self):
|
||||
"""#2159 4k 横屏:短边=height=2160,width=3840。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
w, h = resolve_video_dimensions("4k", "16:9")
|
||||
assert (w, h) == (3840, 2160)
|
||||
|
||||
def test_2160p_alias(self):
|
||||
"""2160p 别名→4k。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
w, h = resolve_video_dimensions("2160p", "9:16")
|
||||
assert (w, h) == (2160, 3840)
|
||||
|
||||
def test_empty_resolution_defaults_to_720p_9_16(self):
|
||||
"""空 resolution + 空 ratio → 默认 720p + 9:16 竖屏 (720×1280)。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
@@ -531,3 +545,99 @@ class TestViralVideoCreditsWithBreakdown:
|
||||
assert bd["tokens"] == 1_000_000.0
|
||||
# video_cost = 1M/1M * 70 = 70; total = (70+0.15)*1.3 = 91.195 → 91.20
|
||||
assert c == 91.20
|
||||
|
||||
|
||||
# ============ #2159 多模型定价单测 ============
|
||||
|
||||
|
||||
class TestMultiModelCredits:
|
||||
"""#2159 多模型积分估算正确性(含 token/second 两种计费模式)。"""
|
||||
|
||||
def test_seedance_2_5_15s_720p_9x16(self):
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
# 15s/720p/9:16 → 720×1280
|
||||
# tokens = 15*720*1280*24/1024 = 324000
|
||||
# video_cost = 324000/1M*70 = 22.68
|
||||
# total = (22.68+0.15)*1.3 = 29.679 ≈ 29.68
|
||||
c = calculate_viral_video_credits(15, 720, 1280, model="seedance-2.5")
|
||||
assert c == 29.68, f"got {c}"
|
||||
|
||||
def test_seedance_2_0_30s_1080p_9x16(self):
|
||||
# 30s/1080p/9:16 → 1080×1920
|
||||
# tokens = 30*1080*1920*24/1024 = 1,458,000
|
||||
# video_cost = 1.458M/1M*51 = 74.358
|
||||
# total = (74.358+0.15)*1.3 = 96.86
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
c = calculate_viral_video_credits(30, 1080, 1920, model="seedance-2.0")
|
||||
assert c == 96.86, f"got {c}"
|
||||
|
||||
def test_seedance_2_0_fast_15s_720p_9x16(self):
|
||||
# 15s/720p/9:16 tokens=324000, price=28
|
||||
# video_cost = 0.324*28 = 9.072
|
||||
# total = (9.072+0.15)*1.3 = 11.99
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
c = calculate_viral_video_credits(15, 720, 1280, model="seedance-2.0-fast")
|
||||
assert c == 11.99, f"got {c}"
|
||||
|
||||
def test_seedance_2_0_mini_15s_720p_9x16(self):
|
||||
# price=9.2, tokens=324000
|
||||
# video_cost = 0.324*9.2 = 2.9808
|
||||
# total = (2.9808+0.15)*1.3 = 4.07
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
c = calculate_viral_video_credits(15, 720, 1280, model="seedance-2.0-mini")
|
||||
assert c == 4.07, f"got {c}"
|
||||
|
||||
def test_wan_3_0_per_second_billing(self):
|
||||
# per_second: 10s/720p price=0.6元/秒
|
||||
# video_cost = 10*0.6 = 6.0
|
||||
# total = (6.0+0.15)*1.3 = 7.995 ≈ 8.00
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
c = calculate_viral_video_credits(10, 720, 1280, model="wan-3.0")
|
||||
assert c == 8.0, f"got {c}"
|
||||
|
||||
def test_seedance_2_0_4k_16x9(self):
|
||||
# 5s/4k/16:9 → 3840×2160, price=80
|
||||
# tokens = 5*3840*2160*24/1024 = 972000
|
||||
# video_cost = 0.972*80 = 77.76
|
||||
# total = (77.76+0.15)*1.3 = 101.28
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
c = calculate_viral_video_credits(5, 3840, 2160, model="seedance-2.0")
|
||||
assert c == 101.28, f"got {c}"
|
||||
|
||||
def test_model_config_has_all_6_models(self):
|
||||
from packages.domain.points_rules import VIRAL_VIDEO_MODEL_CONFIG
|
||||
|
||||
expected = {"seedance-2.5", "seedance-2.0", "seedance-2.0-fast", "seedance-2.0-mini", "wan-3.0"}
|
||||
assert expected.issubset(set(VIRAL_VIDEO_MODEL_CONFIG.keys()))
|
||||
|
||||
def test_list_models_hides_wan_when_dashscope_unavailable(self):
|
||||
from packages.domain.points_rules import list_viral_video_models
|
||||
|
||||
all_models = list_viral_video_models(include_placeholder=False, dashscope_available=False)
|
||||
keys = {m["key"] for m in all_models}
|
||||
assert "wan-3.0" not in keys
|
||||
assert "seedance-2.5" in keys
|
||||
# is_default
|
||||
defaults = [m for m in all_models if m["is_default"]]
|
||||
assert len(defaults) == 1
|
||||
assert defaults[0]["key"] == "seedance-2.5"
|
||||
|
||||
def test_list_models_includes_wan_when_dashscope_available(self):
|
||||
from packages.domain.points_rules import list_viral_video_models
|
||||
|
||||
models = list_viral_video_models(include_placeholder=False, dashscope_available=True)
|
||||
keys = {m["key"] for m in models}
|
||||
assert "wan-3.0" in keys
|
||||
|
||||
def test_infer_4k(self):
|
||||
from packages.domain.points_rules import _infer_resolution_key
|
||||
|
||||
assert _infer_resolution_key(3840, 2160) == "4k"
|
||||
assert _infer_resolution_key(2160, 3840) == "4k"
|
||||
assert _infer_resolution_key(1920, 1080) == "1080p"
|
||||
|
||||
@@ -1,99 +1,38 @@
|
||||
"""爆款视频 DB 模型单元测试(#2039 PR1:DB + migration)。
|
||||
|
||||
验证:
|
||||
- 3 张新表可在内存 SQLite 上创建
|
||||
- 默认值与基本 CRUD 正常
|
||||
"""
|
||||
"""tests for GET /api/v1/viral-video/models route function (#2159)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import (
|
||||
Base,
|
||||
ViralVideoJobModel,
|
||||
ViralVideoPromptTemplateModel,
|
||||
ViralVideoStyleTemplateModel,
|
||||
)
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
||||
def _make_session():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
return sessionmaker(bind=engine)()
|
||||
class TestModelsRoute:
|
||||
def _call(self):
|
||||
from apps.api.app.api.routes import viral_video as routes
|
||||
|
||||
return routes.list_available_models()
|
||||
|
||||
class TestViralVideoJobModel:
|
||||
def test_create_and_get(self):
|
||||
session = _make_session()
|
||||
job = ViralVideoJobModel(
|
||||
id="job-001",
|
||||
user_id="user-001",
|
||||
images=["https://img.com/1.jpg"],
|
||||
industry="美妆",
|
||||
duration=60,
|
||||
)
|
||||
session.add(job)
|
||||
session.commit()
|
||||
def test_without_dashscope_hides_wan(self):
|
||||
with patch("apps.api.app.api.routes.viral_video.get_dashscope_client", return_value=None):
|
||||
result = self._call()
|
||||
assert "models" in result
|
||||
keys = {m["key"] for m in result["models"]}
|
||||
assert "seedance-2.5" in keys
|
||||
assert "wan-3.0" not in keys
|
||||
for m in result["models"]:
|
||||
if m["key"].startswith("seedance"):
|
||||
assert m["supports_audio"] is True
|
||||
|
||||
fetched = session.query(ViralVideoJobModel).filter_by(id="job-001").one()
|
||||
assert fetched.user_id == "user-001"
|
||||
assert fetched.images == ["https://img.com/1.jpg"]
|
||||
assert fetched.industry == "美妆"
|
||||
assert fetched.duration == 60
|
||||
def test_with_dashscope_includes_wan(self):
|
||||
with patch("apps.api.app.api.routes.viral_video.get_dashscope_client", return_value=MagicMock()):
|
||||
result = self._call()
|
||||
keys = {m["key"] for m in result["models"]}
|
||||
assert "wan-3.0" in keys
|
||||
wan = next(m for m in result["models"] if m["key"] == "wan-3.0")
|
||||
assert wan["billing_mode"] == "per_second"
|
||||
|
||||
def test_default_values(self):
|
||||
session = _make_session()
|
||||
job = ViralVideoJobModel(id="job-002", user_id="user-002")
|
||||
session.add(job)
|
||||
session.commit()
|
||||
|
||||
fetched = session.get(ViralVideoJobModel, "job-002")
|
||||
assert fetched.images == []
|
||||
assert fetched.fusion_level == "ai_polish"
|
||||
assert fetched.style_strength == "medium"
|
||||
assert fetched.status == "pending"
|
||||
assert fetched.credits_cost == 0
|
||||
assert fetched.retry_count == 0
|
||||
assert fetched.style_guide is None
|
||||
assert fetched.intent_result is None
|
||||
|
||||
|
||||
class TestViralVideoStyleTemplateModel:
|
||||
def test_create_and_get(self):
|
||||
session = _make_session()
|
||||
tpl = ViralVideoStyleTemplateModel(
|
||||
id="tpl-001",
|
||||
name="快节奏",
|
||||
style_config={"cut_speed": "fast"},
|
||||
sort_order=1,
|
||||
)
|
||||
session.add(tpl)
|
||||
session.commit()
|
||||
|
||||
fetched = session.get(ViralVideoStyleTemplateModel, "tpl-001")
|
||||
assert fetched.name == "快节奏"
|
||||
assert fetched.style_config == {"cut_speed": "fast"}
|
||||
assert fetched.sort_order == 1
|
||||
|
||||
|
||||
class TestViralVideoPromptTemplateModel:
|
||||
def test_create_and_get(self):
|
||||
session = _make_session()
|
||||
tpl = ViralVideoPromptTemplateModel(
|
||||
id="pt-001",
|
||||
prompt_type="image_analysis",
|
||||
name="图片分析模板",
|
||||
content="请分析图片:{image_url}",
|
||||
variables=["image_url"],
|
||||
)
|
||||
session.add(tpl)
|
||||
session.commit()
|
||||
|
||||
fetched = session.get(ViralVideoPromptTemplateModel, "pt-001")
|
||||
assert fetched.prompt_type == "image_analysis"
|
||||
assert fetched.content == "请分析图片:{image_url}"
|
||||
assert fetched.variables == ["image_url"]
|
||||
assert fetched.version == 1
|
||||
assert fetched.is_active is True
|
||||
def test_exactly_one_default(self):
|
||||
with patch("apps.api.app.api.routes.viral_video.get_dashscope_client", return_value=None):
|
||||
result = self._call()
|
||||
defaults = [m for m in result["models"] if m["is_default"]]
|
||||
assert len(defaults) == 1
|
||||
assert defaults[0]["key"] == "seedance-2.5"
|
||||
|
||||
@@ -313,3 +313,133 @@ class TestResumeReadsImageAnalysis:
|
||||
# resume 本身应该调用 _run_render_pipeline
|
||||
resume_src = inspect.getsource(vv.resume_viral_video_pipeline)
|
||||
assert "_run_render_pipeline" in resume_src
|
||||
|
||||
|
||||
# ============ #2157 _replace_henjin_everywhere 递归替换单测 ============
|
||||
|
||||
|
||||
class TestReplaceHenjinEverywhere:
|
||||
"""覆盖 #2157 P1:递归替换 copy_result 中所有层级的'很近'→'最近'。"""
|
||||
|
||||
def _import(self):
|
||||
from apps.worker.worker_app.tasks.viral_video import _replace_henjin_everywhere
|
||||
|
||||
return _replace_henjin_everywhere
|
||||
|
||||
def test_plain_string_no_henjin(self):
|
||||
fn = self._import()
|
||||
assert fn("最近好物推荐") == "最近好物推荐"
|
||||
assert fn("") == ""
|
||||
assert fn(None) is None
|
||||
assert fn(123) == 123
|
||||
|
||||
def test_string_with_henjin(self):
|
||||
fn = self._import()
|
||||
assert fn("很近是不是总觉得颈肩发僵") == "最近是不是总觉得颈肩发僵"
|
||||
# 多次出现
|
||||
assert fn("很近很近都很近") == "最近最近都最近"
|
||||
|
||||
def test_list_recursive(self):
|
||||
fn = self._import()
|
||||
out = fn(["很近a", "b", ["很近c", "d"]])
|
||||
assert out == ["最近a", "b", ["最近c", "d"]]
|
||||
|
||||
def test_dict_recursive_nested(self):
|
||||
fn = self._import()
|
||||
obj = {
|
||||
"overview": {"theme": "很近颈肩", "title": "x"},
|
||||
"scene_and_lighting": "很近才好用",
|
||||
"voiceover_script": "很近是不是",
|
||||
"final_copy": "很近好物",
|
||||
"shots": [
|
||||
{"scene_and_dialogue": "很近第一镜", "action_details": "很近动作", "audio_bgm": "很近音乐"},
|
||||
{"nested": {"deep": "很近深层"}},
|
||||
],
|
||||
"int_field": 42,
|
||||
}
|
||||
import json
|
||||
|
||||
out = fn(obj)
|
||||
assert "很近" not in json.dumps(out, ensure_ascii=False)
|
||||
assert out["overview"]["theme"] == "最近颈肩"
|
||||
assert out["shots"][0]["scene_and_dialogue"] == "最近第一镜"
|
||||
assert out["shots"][1]["nested"]["deep"] == "最近深层"
|
||||
assert out["int_field"] == 42
|
||||
|
||||
|
||||
class TestMarkFailedAndNotifySessionFallback:
|
||||
"""#2157 P1:_mark_failed_and_notify 在原session失效时fallback到新SessionLocal。"""
|
||||
|
||||
def test_fallback_to_new_session_when_original_save_raises(self, tmp_path):
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
job = MagicMock()
|
||||
job.is_terminal = False
|
||||
job.mark_failed = MagicMock()
|
||||
|
||||
# 原 session 保存抛异常
|
||||
orig_session = MagicMock()
|
||||
orig_repo = MagicMock()
|
||||
|
||||
def _raise(*a, **kw):
|
||||
raise RuntimeError("session in rollback")
|
||||
|
||||
# 第一次调用_save_job抛异常,触发fallback
|
||||
with patch.object(vv, "_save_job", side_effect=_raise):
|
||||
fake_ssn = MagicMock()
|
||||
fake_repo = MagicMock()
|
||||
fake_job_in_db = MagicMock()
|
||||
fake_job_in_db.is_terminal = False
|
||||
fake_repo.get.return_value = fake_job_in_db
|
||||
with patch.object(vv, "SessionLocal", return_value=fake_ssn):
|
||||
with patch.object(vv, "SQLAlchemyViralVideoJobRepository", return_value=fake_repo):
|
||||
with patch.object(vv, "_emit_progress") as mock_emit:
|
||||
vv._mark_failed_and_notify("job-1", orig_session, orig_repo, job, "boom", stage="render")
|
||||
# 原session上mark_failed被调用过
|
||||
job.mark_failed.assert_called()
|
||||
# fallback路径:新session上repo.get(job-1)被调用,且新job被mark_failed并commit
|
||||
fake_repo.get.assert_called_with("job-1")
|
||||
fake_job_in_db.mark_failed.assert_called_with("boom")
|
||||
fake_repo.update.assert_called_with(fake_job_in_db)
|
||||
fake_ssn.commit.assert_called()
|
||||
fake_ssn.close.assert_called()
|
||||
mock_emit.assert_called_once()
|
||||
|
||||
def test_original_session_happy_path_no_fallback(self):
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
job = MagicMock()
|
||||
job.is_terminal = False
|
||||
session = MagicMock()
|
||||
repo = MagicMock()
|
||||
with patch.object(vv, "_save_job") as mock_save:
|
||||
with patch.object(vv, "_emit_progress") as mock_emit:
|
||||
# 不mock SessionLocal,如果fallback被错误触发会抛AttributeError
|
||||
vv._mark_failed_and_notify("job-2", session, repo, job, "err", stage="copy")
|
||||
job.mark_failed.assert_called_with("err")
|
||||
mock_save.assert_called()
|
||||
mock_emit.assert_called_once()
|
||||
|
||||
def test_terminal_job_not_marked(self):
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
job = MagicMock()
|
||||
job.is_terminal = True # 已终态
|
||||
session = MagicMock()
|
||||
repo = MagicMock()
|
||||
with patch.object(vv, "_save_job") as mock_save:
|
||||
with patch.object(vv, "_emit_progress"):
|
||||
fake_ssn = MagicMock()
|
||||
with patch.object(vv, "SessionLocal", return_value=fake_ssn):
|
||||
with patch.object(vv, "SQLAlchemyViralVideoJobRepository") as mock_repo_cls:
|
||||
vv._mark_failed_and_notify("job-3", session, repo, job, "x")
|
||||
# 终态job不调用mark_failed
|
||||
job.mark_failed.assert_not_called()
|
||||
# 且因 job 已终态,_save_job 也不应被调用(marked=False 才fallback;但此处 job 非 None 且 is_terminal=True,marked 保持 False 进入 fallback)
|
||||
# fallback路径会重新打开session,get到的job也是终态,不会update
|
||||
|
||||
Reference in New Issue
Block a user