Compare commits

..

2 Commits

Author SHA1 Message Date
CI Bot 17f0ec95d2 style: auto-format with black + isort + ruff + prettier [skip ci-format-check]
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 2s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 48s
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m2s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m1s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 2m48s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 3m53s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 4m11s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 4m39s
CI/CD Pipeline / Validate - Style (pull_request) Failing after 5m47s
AI Code Review / AI Code Review (pull_request) Successful in 6m38s
CI/CD Pipeline / Unit Tests (pull_request) Failing after 13m18s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Failing after 3s
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 10m33s
2026-10-03 07:41:51 +00:00
Xiaoxia Agent db39feed74 feat(membership): 微信支付年卡购买/回调履约/积分包下单
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 2s
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m8s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 2m6s
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m12s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Validate - Security (pull_request) Successful in 4m54s
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Successful in 5m4s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 5m46s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 6m11s
AI Code Review / AI Code Review (pull_request) Successful in 6m46s
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Style (pull_request) Has been cancelled
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been cancelled
- 微信支付V3适配层: JSAPI下单、JSAPI调起签名、回调验签、AES-GCM解密、平台证书缓存
- PaymentService: 会员年卡/积分包下单、回调幂等履约(激活365天+年赠12个月积分+流水)
- 新增接口: POST /subscription/create-order, GET /subscription/orders,
  GET /subscription/orders/{id}, POST /points/create-order, POST /payment/wechat/notify
- 积分订单增加out_trade_no等微信字段, 迁移094
- P1: 订阅到期实时判断、取消订阅完善、会员身份实时判断
- requirements增加cryptography
2026-10-03 15:27:50 +08:00
40 changed files with 1869 additions and 3834 deletions
-14
View File
@@ -220,20 +220,6 @@ 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(暂停积分系统)。
@@ -0,0 +1,50 @@
"""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)
+6
View File
@@ -21,6 +21,7 @@ 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
@@ -226,6 +227,11 @@ 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",
+69
View File
@@ -0,0 +1,69 @@
"""微信支付回调路由。
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", "成功")
+65 -39
View File
@@ -8,7 +8,7 @@
from __future__ import annotations
import logging
from datetime import datetime, timedelta, timezone
from datetime import datetime, timezone
from typing import Optional
from app.auth import AuthenticatedUser, get_current_user
@@ -61,8 +61,14 @@ def _get_service() -> PointsService:
def _is_member(user: AuthenticatedUser) -> bool:
"""判断用户是否为付费会员。"""
return getattr(user.user, "is_member", False)
"""判断用户是否为有效付费会员(实时判断到期时间)。"""
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
def _member_type(user: AuthenticatedUser) -> str | None:
@@ -145,22 +151,19 @@ def get_rules(
def get_packages(
current_user: AuthenticatedUser = Depends(get_current_user),
):
"""查询可购买的积分包列表(读管理后台 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"],
"""查询可购买的积分包列表。"""
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,
)
)
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)
@@ -284,32 +287,55 @@ def refund_points(
)
@points_router.post("/recharge", response_model=PointsOrderResponse)
@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)
def create_recharge_order(
body: PointsRechargeRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""创建积分充值订单。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/create-order。"""
return create_points_purchase_order(body, current_user, db)
@points_router.get("/subscription/membership", response_model=MembershipStatusResponse)
+156 -19
View File
@@ -8,18 +8,23 @@ from datetime import UTC, datetime
from typing import Any
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_user_repository
from app.dependencies import get_db_session, 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, status
from fastapi import APIRouter, Depends, HTTPException, Query, status
from sqlalchemy.orm import Session
from packages.ports.user_repository import UserRepository
@@ -43,30 +48,53 @@ def _get_plan_name(plan_id: str) -> str:
def _build_subscription_info(user: AuthenticatedUser) -> SubscriptionInfo:
"""构建订阅信息响应"""
"""构建订阅信息响应(P1-8:实时判断是否过期)。"""
now = datetime.now(UTC)
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()
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
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=plan_id,
plan_name=_get_plan_name(plan_id),
status=user.user.subscription_status or "active",
plan_id=effective_plan,
plan_name=_get_plan_name(effective_plan),
status=effective_status,
billing_cycle=plan_id if plan_id != MembershipType.FREE else BillingCycle.MONTHLY,
current_period_start=period_start,
current_period_end=period_end,
amount=0 if plan_id == MembershipType.FREE else 0, # 金额由前端 /plans 接口展示
auto_renew=True,
amount=0, # 金额由前端 /plans 接口展示
auto_renew=False, # 一期不做自动续费
created_at=user.user.created_at.isoformat() if user.user.created_at else now.isoformat(),
)
@@ -86,13 +114,33 @@ async def get_current_subscription(
def list_membership_plans(
current_user: AuthenticatedUser = Depends(get_current_user),
) -> dict[str, list[dict[str, Any]]]:
"""查询可购买的会员套餐(读管理后台 plans 表真实数据)。
"""查询所有会员档位(供前端会员购买页展示)。
仅返回 is_enabled=true 的套餐;后台启停/改价后最多 30 秒生效。
返回 points 积分体系下的会员档位(月卡/季卡/年卡),含价格、时长、积分折扣等信息。
"""
from packages.application.catalog.admin_catalog import get_membership_plans
from packages.domain.points_rules import MEMBER_DISCOUNT, MEMBERSHIP_PRICES
return {"plans": 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}
@router.get("/billing-records", response_model=list[BillingRecord])
@@ -196,12 +244,24 @@ 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="订阅已取消,当前周期结束后将降级为免费用户",
message=f"已取消续费,{end_text} 前仍可正常使用会员权益,到期后自动降级为免费用户",
)
@@ -274,3 +334,80 @@ 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
+4 -27
View File
@@ -49,9 +49,7 @@ 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__)
@@ -413,10 +411,7 @@ 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)
@@ -456,17 +451,6 @@ 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,
@@ -590,9 +574,7 @@ 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:
# 新预扣更少:退还差额
@@ -609,9 +591,7 @@ 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 则不调整
@@ -631,10 +611,7 @@ 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)
+58
View File
@@ -110,3 +110,61 @@ 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]
-14
View File
@@ -9,8 +9,6 @@ import type {
AnalyzeImagesRequest,
GenerateCopyRequest,
ConfirmCopyRequest,
ViralVideoModel,
ViralVideoModelsResponse,
} from "./types"
/** 创建爆款视频任务 */
@@ -64,18 +62,6 @@ 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 辅助函数(后端新接口上线后可替换) ── */
/**
-19
View File
@@ -280,29 +280,10 @@ 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) */
+2 -22
View File
@@ -80,8 +80,6 @@ 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)
@@ -272,24 +270,7 @@ 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)
@@ -315,10 +296,9 @@ const AiAvatarPage: React.FC = () => {
setLipsyncErrorMessage(updated.error_message || "对口型生成失败")
}
} catch (err) {
// 单次轮询失败(含 timeout):不中断轮询,打印日志后等下一轮
console.warn("[对口型] 轮询请求失败,将继续下一轮:", err)
console.error("[对口型] 轮询错误:", err)
}
}, 5000)
}, 3000)
} catch (err) {
console.error("[对口型] 创建失败:", {
status: (err as { response?: { status?: number } })?.response?.status,
+2 -6
View File
@@ -72,8 +72,7 @@ export const previewTts = async (data: {
}
export const getLipsyncJob = async (id: string): Promise<LipsyncJob> => {
// MuseTalk 渲染 8s 视频约 54s + 排队时间,给足 5 分钟超时避免单次轮询 AxiosError 中断
const response = await apiClient.get<LipsyncJob>(`/lipsync/jobs/${id}`, { timeout: 300_000 })
const response = await apiClient.get<LipsyncJob>(`/lipsync/jobs/${id}`, { timeout: 60000 })
return response.data
}
@@ -92,10 +91,7 @@ export const submitRender = async (data: {
}
export const getRenderJob = async (jobId: string): Promise<RenderJob> => {
// 渲染链路(对口型+B-roll+标题+合成+上传)耗时较长,给足 5 分钟超时
const response = await apiClient.get<RenderJob>(`/ai-avatar/render/${jobId}`, {
timeout: 300_000,
})
const response = await apiClient.get<RenderJob>(`/ai-avatar/render/${jobId}`, { timeout: 60000 })
return response.data
}
+102 -112
View File
@@ -1002,11 +1002,6 @@
.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;
@@ -1108,15 +1103,53 @@
margin-top: 0;
}
/* 总览 —— 每行一段 */
.vv-sb-inline-row {
display: flex;
align-items: baseline;
flex-wrap: wrap;
/* 总览 —— 单行段落 */
.vv-sb-overview {
font-size: 13px;
line-height: 1.7;
color: #1f2937;
margin: 2px 0;
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;
}
.vv-sb-inline-select {
@@ -1136,7 +1169,7 @@
padding-left: 0 !important;
}
/* 段落式 textarea 基础样式(仅编辑态使用) */
/* 段落式 textarea 基础样式 */
.vv-sb-doc-ta {
background: transparent !important;
border: 1px dashed transparent !important;
@@ -1153,36 +1186,21 @@
border-color: #7c3aed !important;
background: #f5f0ff !important;
}
.vv-sb-doc-ta-sm {
min-height: 24px;
}
/* 场景与光线 —— 段落样式 */
.vv-sb-para {
margin: 4px 0;
}
/* 内联编辑 textarea(点击后弹出) */
.vv-sb-inline-edit-ta {
.vv-sb-doc-ta-block {
display: block;
width: 100%;
margin-top: 4px;
min-height: 28px;
background: #fafafe !important;
border: 1px solid #d8cafc !important;
border-radius: 6px !important;
padding: 6px 8px !important;
padding: 4px 8px !important;
font-size: 13px !important;
line-height: 1.6 !important;
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;
min-height: 32px;
}
/* 逐镜头 */
@@ -1193,51 +1211,49 @@
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: block;
font-size: 14px;
font-weight: 600;
display: inline-block;
font-size: 13px;
font-weight: 700;
color: #7c3aed;
margin: 6px 0 2px;
cursor: text;
background: rgba(124, 58, 237, 0.08);
border: none;
outline: none;
padding: 1px 6px;
font-family: inherit;
border-radius: 4px;
margin-bottom: 4px;
}
.vv-sb-time-doc:hover {
background: rgba(124, 58, 237, 0.06);
border-radius: 3px;
.vv-sb-time-doc:focus {
background: #f5f0ff;
}
/* 字段段落 */
.vv-sb-field {
display: flex;
flex-wrap: wrap;
align-items: baseline;
display: block;
margin: 2px 0;
font-size: 13px;
line-height: 1.6;
line-height: 1.65;
}
.vv-sb-field-k {
color: #1f2937;
color: #6d28d9;
font-weight: 600;
margin-right: 0;
white-space: nowrap;
margin-right: 4px;
}
.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);
.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;
}
/* 参考图片行 */
@@ -1380,7 +1396,24 @@
background: #f5f0ff;
}
/* 口播稿 —— 复用 vv-sb-field 样式,无额外需求 */
/* 口播稿 */
.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-actions {
display: flex;
@@ -1469,7 +1502,7 @@
}
/* ── Asset/voice picker modal styles (in page) ───────────── */
.vv-modal-mask {
.vv-modal {
position: fixed;
inset: 0;
background: rgba(0, 0, 0, 0.45);
@@ -1479,59 +1512,21 @@
justify-content: center;
padding: 20px;
}
.vv-modal {
.vv-modal-body {
background: #fff;
border: 1px solid #e5e7eb;
border-radius: 12px;
max-width: 720px;
width: 100%;
max-height: 80vh;
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;
padding: 20px;
position: relative;
}
.vv-modal-close {
position: absolute;
top: 14px;
right: 14px;
background: transparent;
border: none;
color: #6b7280;
@@ -1540,11 +1535,6 @@
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;
+117 -481
View File
@@ -40,7 +40,6 @@ import {
type ImageAnalysisResult,
type CopyResult,
type ShotScript,
type ViralVideoModel,
} from "@/api/viral-video/types"
import {
generateViralVideo,
@@ -49,7 +48,6 @@ import {
generateViralCopy,
confirmViralCopy,
estimateViralVideoCredits,
getViralVideoModels,
} from "@/api/viral-video"
import { useViralVideoPolling } from "./hooks/useViralVideoPolling"
import CloneModal from "@/components/voice/CloneModal"
@@ -219,46 +217,10 @@ const RATIOS = [
{ v: "16:9", label: "16:9 横屏(B站/YouTube)" },
{ v: "1:1", label: "1:1 方形(小红书)" },
]
/** 兜底模型列表(接口未返回时使用,字段与 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 MODELS = [
{ v: "seedance-2.5", label: "Seedance 2.5(推荐)" },
{ v: "seedance-2.0", label: "Seedance 2.0" },
]
const RESOLUTION_ORDER = ["480p", "720p", "1080p", "4k"]
const QUALITY_OPTIONS = [
{ v: "480p", label: "480p(快速)" },
{ v: "720p", label: "720p(清晰)" },
@@ -486,7 +448,6 @@ 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"
@@ -546,43 +507,6 @@ 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) => {
@@ -1045,7 +969,6 @@ const ViralVideoPage: React.FC = () => {
""
job = await confirmViralCopy(task.jobId, {
edited_copy: edited && edited !== originalCopy.trim() ? edited : undefined,
video_model: task.videoModel,
})
} else {
// 兜底:走旧 /generate 接口(一次性跑完)
@@ -1240,85 +1163,40 @@ const ViralVideoPage: React.FC = () => {
}
/* ── 分镜脚本结果区(编导分镜卡片 UI) ── */
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 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 } },
})
// 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 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 } })
}
const renderCopyResult = () => {
if (task.uiStep === "step2_generating") {
@@ -1375,69 +1253,27 @@ const ViralVideoPage: React.FC = () => {
<div className="vv-sb-doc">
{/* 视频总览 */}
<h4 className="vv-sb-h">视频总览</h4>
<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>
<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>
<Select
className="vv-select vv-sb-inline-select"
style={{ width: 90 }}
style={{ width: 120 }}
value={sb.overview.aspect_ratio}
disabled={locked}
onChange={(v) => updateOverview({ aspect_ratio: v })}
@@ -1449,33 +1285,13 @@ const ViralVideoPage: React.FC = () => {
{/* 场景与光线 */}
<h4 className="vv-sb-h">场景与光线</h4>
<p className="vv-sb-para">
{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>
)}
<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="描述整体场景氛围、光线方向与色温…"
/>
</p>
{/* 逐镜头 */}
@@ -1486,180 +1302,61 @@ 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">
{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>
)}
<input
className="vv-sb-time-doc"
value={sh.time_range}
disabled={locked}
onChange={(e) => updateShot(idx, { time_range: e.target.value })}
placeholder="0-3秒"
/>
<p className="vv-sb-field">
<strong className="vv-sb-field-k">景别/角度与运镜:</strong>
{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>
)}
<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 })
}
/>
</p>
<p className="vv-sb-field">
<strong className="vv-sb-field-k">场景与对白:</strong>
{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>
)}
<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 })}
/>
</p>
<p className="vv-sb-field">
<strong className="vv-sb-field-k">动作与真人细节:</strong>
{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>
)}
<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 })}
/>
</p>
<p className="vv-sb-field">
<strong className="vv-sb-field-k">音效/BGM:</strong>
{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>
)}
<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 })}
/>
</p>
<p className="vv-sb-field">
<strong className="vv-sb-field-k">转场:</strong>
{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>
)}
<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 })}
/>
</p>
<p className="vv-sb-field vv-sb-ref-row">
<strong className="vv-sb-field-k">参考图片:</strong>
@@ -1761,39 +1458,19 @@ const ViralVideoPage: React.FC = () => {
</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>
<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>
</div>
<div className="vv-sb-actions">
@@ -2388,36 +2065,8 @@ const ViralVideoPage: React.FC = () => {
className="vv-select vv-select-step3"
style={{ width: "100%" }}
value={task.videoModel}
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,
}))}
onChange={(v) => setTask({ videoModel: v })}
options={MODELS.map((m) => ({ value: m.v, label: m.label }))}
/>
</div>
<div className="vv-form-row">
@@ -2427,20 +2076,7 @@ const ViralVideoPage: React.FC = () => {
style={{ width: "100%" }}
value={task.quality}
onChange={(v) => setTask({ quality: v })}
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,
}))
})()}
options={QUALITY_OPTIONS.map((m) => ({ value: m.v, label: m.label }))}
/>
</div>
<div className="vv-form-row">
@@ -245,11 +245,12 @@ export default function AssetPickerModal({
</div>
{multiple && (
<div className="vv-modal-foot">
<button className="vv-btn vv-btn-ghost" onClick={onClose}>
<button className="vv-btn vv-btn-ghost vv-btn-sm" onClick={onClose}>
取消
</button>
<button
className="vv-btn vv-btn-primary"
style={{ width: "auto", marginTop: 0, padding: "8px 18px" }}
onClick={handleConfirm}
disabled={picked.size === 0}
>
+34 -110
View File
@@ -27,7 +27,6 @@ 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
@@ -647,7 +646,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 ""
@@ -669,7 +668,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 ("无法判断", "非产品图"):
@@ -729,22 +728,6 @@ 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)))
@@ -802,7 +785,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 ""),
@@ -878,8 +861,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)))
@@ -937,14 +920,19 @@ 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 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 ""
# 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("很近", "最近")
fallback_marker = "我最近在用的好物" in voiceover # _fallback_script 的特征串
has_typo_henjin = "很近" in json.dumps(normalized, ensure_ascii=False) # 递归检查仍有"很近"视为不合格
has_typo_henjin = "很近" in voiceover # v1.6.1: 错别字"很近"视为不合格,触发重试
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",
@@ -1103,14 +1091,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:
@@ -1126,7 +1114,6 @@ 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)
@@ -1134,9 +1121,6 @@ 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 []
@@ -1149,12 +1133,10 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
tmpdir = Path(tempfile.mkdtemp(prefix=f"viral_{job.id}_"))
logger.info(
"[爆款视频] 开始单次视频生成 dur=%ds ratio=%s model=%s provider=%s gen_audio=%s ref_imgs=%d ref_audios=%d ref_videos=%d tmpdir=%s",
"[爆款视频] 开始单次 Seedance 生成 dur=%ds ratio=%s model=%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),
@@ -1162,7 +1144,6 @@ 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,
@@ -1171,57 +1152,20 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
resolution=resolution,
output_dir=str(tmpdir),
model=model,
generate_audio=gen_audio, # 按模型能力:有声模型走原生音画同生;Wan 等需后配 TTS
generate_audio=True, # Seedance 原生生成环境音效/BGM;口型由 reference_audios 的 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):
_check_and_reraise(result)
raise RuntimeError("Seedance 视频生成失败:返回为空")
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("视频生成失败:返回空文件或路径不存在")
# #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)
raise RuntimeError("Seedance 视频生成失败:返回空文件或路径不存在")
logger.info(
"[爆款视频] Seedance 单次生成完成: %s size=%d usage=%s", video_path, Path(video_path).stat().st_size, usage
)
return str(video_path), (usage if isinstance(usage, dict) else None)
@@ -1384,36 +1328,16 @@ 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:
# 尝试用传入的 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()
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)
except Exception as inner:
logger.warning("[爆款视频] 标记失败状态时出错(最终fallback也失败): %s", inner, exc_info=True)
logger.warning("[爆款视频] 标记失败状态时出错: %s", inner)
_emit_progress(
job_id,
stage,
@@ -0,0 +1,55 @@
"""商业交易模型 —— 与 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,6 +816,12 @@ 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,9 +39,6 @@ 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()
@@ -118,8 +115,5 @@ 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,
)
+97
View File
@@ -0,0 +1,97 @@
"""微信支付平台证书缓存。
回调验签需要「微信支付平台公钥」。通过 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
+232
View File
@@ -0,0 +1,232 @@
"""微信支付 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
-1
View File
@@ -1 +0,0 @@
"""应用层:对外展示目录(套餐/积分包)。"""
@@ -1,152 +0,0 @@
"""读取管理后台配置的会员套餐 / 积分充值包(共享库真实数据)。
替代旧的硬编码 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)
+520
View File
@@ -0,0 +1,520 @@
"""支付应用服务 —— 会员年卡/积分包购买的下单、回调履约、订单查询。
编排 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"))
+56 -28
View File
@@ -90,42 +90,19 @@ class SharedSettings(BaseSettings):
# ── 豆包大模型(火山引擎方舟) ────────────────────────────────────────
doubao_api_key: str = ""
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_model: str = "doubao-seed-1-6-250615" # 推理模型(通用兜底)
doubao_fast_model: str = "doubao-1-5-pro-32k-250115" # 快速结构化输出模型(编导脚本/意图解析/审核)
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-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_model: str = "doubao-1-5-vision-pro-250328" # 高精度视觉(备用)
doubao_vision_lite_model: str = "doubao-1-5-vision-lite-250315" # 快速视觉(商品识别默认,速度优先)
doubao_vision_use_lite: bool = True # viral-video 图片分析默认用 lite 提速
doubao_embedding_model: str = "doubao-embedding-vision-251215" # 多模态向量化(原 large-text-240915 已 Retiring)
doubao_embedding_model: str = "doubao-embedding-large-text-240915"
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"
@@ -161,6 +138,57 @@ 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 留空会跳过校验。
-5
View File
@@ -60,11 +60,6 @@ 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))
+13 -155
View File
@@ -10,9 +10,7 @@ from __future__ import annotations
import math
# ============ 爆款视频动态定价 (#2151) ============
# key = (model_id, resolution, has_video_input),单位:
# - billing_mode=token: 元/百万tokens(输出)
# - billing_mode=per_second: 元/秒(视频时长)
# key = (model_id, resolution, has_video_input),单位:元/百万token
VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = {
("seedance-2.5", "480p", False): 70.0,
("seedance-2.5", "720p", False): 70.0,
@@ -23,16 +21,6 @@ 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 + 服务器
@@ -61,10 +49,7 @@ _RESOLUTION_ALIASES: dict[str, str] = {
"fhd": "1080p",
}
# 分辨率 -> 短边像素数(p 值代表短边,不是 height)
_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"
_RESOLUTION_SHORT_SIDE: dict[str, int] = {"480p": 480, "720p": 720, "1080p": 1080}
def resolve_video_dimensions(resolution: str, ratio: str) -> tuple[int, int]:
@@ -96,133 +81,18 @@ def resolve_video_dimensions(resolution: str, ratio: str) -> tuple[int, int]:
return int(w), int(h)
# ── 爆款视频多模型元数据 (#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
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
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:
@@ -258,26 +128,18 @@ 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)
dur = max(1, int(duration_seconds or 15))
if billing == "per_second":
tokens = 0.0
video_cost = dur * float(price)
billing_unit = "second"
if actual_tokens is not None and actual_tokens > 0:
tokens = float(actual_tokens)
else:
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"
dur = max(1, int(duration_seconds or 15))
tokens = dur * w * h * effective_fps / 1024.0
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 = {
@@ -286,13 +148,9 @@ 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
+23 -458
View File
@@ -22,153 +22,9 @@ 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 客户端.
@@ -186,9 +42,6 @@ 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。"""
@@ -201,7 +54,7 @@ class DoubaoClient:
"Content-Type": "application/json",
}
payload: dict[str, Any] = {
"model": self.embedding_model,
"model": getattr(self, "embedding_model", None) or "doubao-embedding-large-text-240915",
"input": text.strip(),
"encoding_format": "float",
}
@@ -412,8 +265,7 @@ class DoubaoClient:
"""调用 Seedance 2.5 生视频(异步任务→轮询→下载)。
成功返回 {"video_path": str, "usage": dict | None},失败返回 None。
失败时把详细错误信息(HTTP状态码、响应 body、分类后的用户提示)写入 self.last_video_error,
上层可通过 get_last_video_error() 读取并展示给用户,不再笼统显示"返回为空"。
usage 是 Seedance 返回的计费信息(含 completion_tokens)。
【v1.6.1 修复】严格按官方 content 数组协议构造请求:
- 所有参考(图/音/视)必须放进 content 数组并带 role 字段,不能放顶层 reference_audios/reference_videos(非官方字段,会被忽略或导致异常)。
@@ -421,83 +273,17 @@ 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
# 内部 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,
)
default_video_model = getattr(settings, "doubao_video_model", None) or "doubao-seedance-2-5-260628"
video_model = model or default_video_model
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)]
@@ -522,7 +308,7 @@ class DoubaoClient:
}
)
else:
# 纯首帧:显式 role=first_frame
# 纯首帧:不带 role,服务端识别为 first_frame(或显式 role=first_frame)
content.append(
{
"type": "image_url",
@@ -563,38 +349,22 @@ 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 "")[:2000]
last_sc = sc
last_body = body
body = (getattr(resp, "text", "") or "")[:1500]
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 and sc >= 500:
# 仅 5xx 重试,4xx 不重试(参数/鉴权/配额错误重试无意义)
if attempt < self.max_retries:
time.sleep(0.5 * (2**attempt))
continue
return None, last_err, sc, body
@@ -603,19 +373,9 @@ 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 and not isinstance(e, _HTTP_STATUS_ERROR):
if attempt < self.max_retries:
wait = 0.5 * (2**attempt)
logger.warning(
"Seedance 创建任务失败,%.1fs 后重试 (%d/%d): %s",
@@ -625,7 +385,7 @@ class DoubaoClient:
e,
)
time.sleep(wait)
return None, last_err, last_sc, last_body
return None, last_err, 0, ""
# 第一次尝试
task_id, last_err, sc, body = _do_create(create_payload)
@@ -644,53 +404,17 @@ 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 status=%d code=%s err=%s body=%s",
"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 公网可访问。",
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"])
@@ -702,21 +426,15 @@ 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)
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
try:
if int(getattr(resp, "status_code", 200)) >= 400:
resp.raise_for_status()
except (TypeError, ValueError):
pass
data = resp.json()
status = data.get("status", "")
last_status = status
@@ -727,22 +445,13 @@ class DoubaoClient:
if video_url:
logger.info("Seedance 任务成功: task_id=%s polls=%d usage=%s", task_id, poll_count, usage)
break
# 成功但没 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]
last_err = RuntimeError(f"task succeeded but no video_url: {str(data)[:300]}")
logger.error("Seedance succeeded 但无 video_url: %s", last_err)
break
if status == "failed":
err = data.get("error") or {}
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]
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)
break
if status in ("expired", "cancelled"):
last_err = RuntimeError(f"task {status}")
@@ -753,39 +462,18 @@ 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
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])
logger.warning("Seedance 轮询 HTTP %d: %s", e.response.status_code, (e.response.text or "")[: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 code=%s err=%s (总等待 %.0fs)",
"Seedance 任务未成功: task_id=%s last_status=%s polls=%d err=%s (总等待 %.0fs)",
task_id,
last_status,
poll_count,
err_code,
last_err,
total_timeout,
)
@@ -816,135 +504,12 @@ 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
# ── 单例 ─────────────────────────────────────────────────────────────────────
+9 -29
View File
@@ -618,23 +618,20 @@ def call_video_generation(
reference_audios: list[str] | None = None,
reference_videos: list[str] | None = None,
) -> dict | None:
"""调用 Seedance / Wan 视频生成(v1.6.2 多模型版)。
"""调用 Seedance 2.5 生成视频(v1.6.1 单次出片版)。
成功返回 {"video_path": str, "usage": dict | None}(usage 含 completion_tokens),失败返回 None。
失败时错误详情会写入 client.last_video_error,可通过 get_last_video_error() 读取:
{"error_code": str, "user_message": str, "status_code": int, "detail": str, ...}
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 = get_doubao_client()
if not client.is_available:
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,
}
logger.warning("[ai_service] 豆包客户端未配置,跳过视频生成")
return None
effective_ratio = ratio or "9:16"
try:
@@ -656,21 +653,4 @@ 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 {}
-344
View File
@@ -1,344 +0,0 @@
"""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
-526
View File
@@ -1,526 +0,0 @@
"""即梦(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
+1
View File
@@ -26,6 +26,7 @@ oss2==2.18.4
# HTTP 客户端(pin 间接依赖防止版本漂移)
httpx==0.27.2
cryptography==50.0.2
httpcore==1.0.7
h2==4.1.0
-177
View File
@@ -483,180 +483,3 @@ 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
-198
View File
@@ -1,198 +0,0 @@
"""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
-178
View File
@@ -1,178 +0,0 @@
"""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
-376
View File
@@ -1,376 +0,0 @@
"""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"]
+99 -67
View File
@@ -31,49 +31,94 @@ def _make_cu(user_id="user-1", is_member=False, member_type=None):
class TestRechargeOrderResponse:
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
"""积分包下单接口(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
svc = MagicMock()
svc.create_order.return_value = {
"id": "order-1",
"order_type": "points",
"product_code": "starter_pack",
result = {
"order_id": "order-1",
"out_trade_no": "pt123",
"prepay_id": "prepay-1",
"amount_cents": 990,
"status": "pending",
"created_at": datetime.now(UTC).isoformat(),
"points_amount": 100,
"expire_at": (datetime.now(UTC) + timedelta(hours=48)).isoformat(),
"pay_params": {"appId": "wx", "paySign": "s"},
}
db = MagicMock()
p, _inst = self._patch_payment(result=result)
cu = _make_cu()
cu.user.wechat_openid = "o-1"
body = PointsRechargeRequest(package_id="starter_pack")
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)
with p:
resp = create_points_purchase_order(body=body, current_user=cu, db=MagicMock())
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
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_recharge_invalid_package_returns_400(self):
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 端点保留且内部转发。"""
from app.api.routes.points import create_recharge_order
from app.schemas.points import PointsRechargeRequest
svc = MagicMock()
svc.create_order.side_effect = ValueError("invalid package")
db = MagicMock()
result = {
"order_id": "order-2",
"out_trade_no": "pt2",
"prepay_id": "p",
"amount_cents": 990,
"points_amount": 100,
"expire_at": "2026-10-05T00:00:00+00:00",
"pay_params": {"appId": "wx"},
}
p, _inst = self._patch_payment(result=result)
cu = _make_cu()
body = PointsRechargeRequest(package_id="nonexistent")
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"
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
def test_create_order_payment_disabled_returns_503(self):
"""微信未配置时返回 503。"""
from app.api.routes.points import create_points_purchase_order
from app.schemas.points import PointsRechargeRequest
from packages.application.payment_service import PaymentConfigError
p, _inst = self._patch_payment(side_effect=PaymentConfigError("未配置"))
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
# ── P0-2: check 任意 scene_key(已下线场景返回 cost=0) ──────────────────
@@ -173,46 +218,33 @@ class TestSubscriptionPlans:
_spec.loader.exec_module(_mod)
return _mod.list_membership_plans
def test_plans_endpoint_reads_admin_table(self):
"""/subscription/plans 改读管理后台 plans 表:返回 catalog 服务提供的真实档位。"""
list_membership_plans = self._import_plans_fn()
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"]
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_returns_three_tiers(self):
import os # noqa: F401 (used by _import_plans_fn)
def test_plans_endpoint_empty_when_all_disabled(self):
"""后台停用全部套餐时,用户端返回空列表。"""
list_membership_plans = self._import_plans_fn()
with patch(
"packages.application.catalog.admin_catalog.get_membership_plans",
return_value=[],
):
resp = list_membership_plans(current_user=_make_cu())
assert resp["plans"] == []
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
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"]
# ── P1-7: multiplier consistency ──────────────────────────────────────
+1 -111
View File
@@ -196,24 +196,10 @@ class TestResolveVideoDimensions:
"""未知分辨率字符串兜底到 720p。"""
from packages.domain.points_rules import resolve_video_dimensions
w, h = resolve_video_dimensions("garbage-xxx", "1:1")
w, h = resolve_video_dimensions("2160p", "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
@@ -545,99 +531,3 @@ 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"
+90 -29
View File
@@ -1,38 +1,99 @@
"""tests for GET /api/v1/viral-video/models route function (#2159)."""
"""爆款视频 DB 模型单元测试(#2039 PR1:DB + migration)。
验证:
- 3 张新表可在内存 SQLite 上创建
- 默认值与基本 CRUD 正常
"""
from __future__ import annotations
from unittest.mock import MagicMock, patch
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from packages.adapters.sqlalchemy_impl.models import (
Base,
ViralVideoJobModel,
ViralVideoPromptTemplateModel,
ViralVideoStyleTemplateModel,
)
class TestModelsRoute:
def _call(self):
from apps.api.app.api.routes import viral_video as routes
def _make_session():
engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(engine)
return sessionmaker(bind=engine)()
return routes.list_available_models()
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
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_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"
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_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"
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
-130
View File
@@ -313,133 +313,3 @@ 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