Compare commits
56 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 3b9a3dd426 | |||
| 617c40e1d4 | |||
| 2597962528 | |||
| 8c56694599 | |||
| cd618c3f29 | |||
| 015fd2c381 | |||
| 8caf3ac8c3 | |||
| 3577108e29 | |||
| 0bf4f359a7 | |||
| 35e7789c81 | |||
| 70c526ffc3 | |||
| e2de75c9f6 | |||
| d949e90051 | |||
| 7adbb7d331 | |||
| 8624896379 | |||
| 65343473d8 | |||
| d7fa9d9e8d | |||
| 90004cced4 | |||
| e672c17eb2 | |||
| fb0e4989cd | |||
| 702f09e6b7 | |||
| 12b0d15473 | |||
| 9344314eac | |||
| 3f49867384 | |||
| dd420c556f | |||
| 2d823a9255 | |||
| ed24c7cd68 | |||
| 69f88434bd | |||
| 599388d9e0 | |||
| 9ffe909dc0 | |||
| 6243196408 | |||
| 3d8f2c2ce2 | |||
| 3eef497dfe | |||
| 61c15eb987 | |||
| 0012ecad30 | |||
| e1994ada0a | |||
| f9f6c53ef4 | |||
| 9699a1fcde | |||
| 305e2bd9d5 | |||
| 30cc58441f | |||
| b3b5dbd459 | |||
| 9f5948dd2b | |||
| fbc1df36e0 | |||
| d36cc09be5 | |||
| fa4dbc6761 | |||
| 77133eb9a0 | |||
| 35207ae040 | |||
| 5af16f1aea | |||
| c232721fff | |||
| fc27d4e81e | |||
| 3c782f89d1 | |||
| 2b4f11036f | |||
| 2881da65cf | |||
| 1037e218bb | |||
| f0514d7487 | |||
| bf8f62ec5b |
@@ -0,0 +1,2 @@
|
||||
Mon Oct 5 04:09:11 PM CST 2026
|
||||
2198 lite/pro并行竞速 (commit 9699a1f) — CI rebuild trigger Mon Oct 5 08:09:11 AM UTC 2026
|
||||
+61
@@ -0,0 +1,61 @@
|
||||
"""功能计费积分字段(爆款/对口型/智能剪辑 DB 化计费)。
|
||||
|
||||
给 gpu_lipsync_tasks / generation_tasks / lipsync_jobs 三张表加积分字段:
|
||||
- credits_prepaid: 提交任务时预扣积分
|
||||
- credits_cost: 最终结算积分
|
||||
- credits_transaction_id: 预扣流水 ID
|
||||
|
||||
注意:feature_pricing_configs 配置表由 xiaoxia-admin 侧 migration 建立,
|
||||
本仓库只读,不在此创建。
|
||||
|
||||
Revision ID: 096_feature_billing_fields
|
||||
Revises: 095_viral_video_prompt_templates
|
||||
Create Date: 2026-10-05
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "096_feature_billing_fields"
|
||||
down_revision = "095_viral_video_prompt_templates"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
_TABLES = ("gpu_lipsync_tasks", "generation_tasks", "lipsync_jobs")
|
||||
_COLUMNS = (
|
||||
("credits_prepaid", sa.Float(), "0"),
|
||||
("credits_cost", sa.Float(), "0"),
|
||||
("credits_transaction_id", sa.String(36), ""),
|
||||
)
|
||||
|
||||
|
||||
def _table_exists(conn, name: str) -> bool:
|
||||
return name in sa.inspect(conn).get_table_names()
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
for table in _TABLES:
|
||||
if not _table_exists(conn, table):
|
||||
continue
|
||||
existing = {c["name"] for c in sa.inspect(conn).get_columns(table)}
|
||||
for col_name, col_type, default in _COLUMNS:
|
||||
if col_name in existing:
|
||||
continue
|
||||
op.add_column(
|
||||
table,
|
||||
sa.Column(col_name, col_type, nullable=False, server_default=default),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
for table in _TABLES:
|
||||
if not _table_exists(conn, table):
|
||||
continue
|
||||
existing = {c["name"] for c in sa.inspect(conn).get_columns(table)}
|
||||
for col_name, _col_type, _default in _COLUMNS:
|
||||
if col_name not in existing:
|
||||
continue
|
||||
op.drop_column(table, col_name)
|
||||
File diff suppressed because one or more lines are too long
@@ -44,6 +44,7 @@ from packages.application import (
|
||||
GetGenerationTaskUseCase,
|
||||
ListGeneratedVideosByTaskUseCase,
|
||||
)
|
||||
from packages.domain import feature_pricing_service
|
||||
from packages.domain.smart_match import smart_select_assets
|
||||
|
||||
# #2035:文案关键词 → 素材分类 映射表(用于 smart_match category_match 维度)
|
||||
@@ -163,7 +164,6 @@ def _infer_expected_categories(script_tags: set[str] | None) -> set[str] | None:
|
||||
return matched or None
|
||||
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
@@ -700,6 +700,17 @@ def create_generation_task(
|
||||
logger.info("画中画已下线,strategy_id %s → one_take", effective_strategy_id)
|
||||
effective_strategy_id = "one_take"
|
||||
|
||||
# ── smart_edit 计费预扣(全局 points 开关 + 功能开关均开才扣) ──
|
||||
# 首期固定价:dynamic_cost=0,price=(0+fixed_cost)×multiplier,price_cap 封顶。
|
||||
# 预览任务不扣费;按任务条数扣费,任一任务预扣失败(余额不足)整体拒绝。
|
||||
smart_edit_charge = 0.0
|
||||
charged_task_count = 0
|
||||
if not request.is_preview and feature_pricing_service.is_feature_enabled("smart_edit"):
|
||||
unit_credits, _bd = feature_pricing_service.calculate_price("smart_edit", 0.0)
|
||||
if unit_credits > 0:
|
||||
smart_edit_charge = round(unit_credits * count, 2)
|
||||
charged_task_count = count
|
||||
|
||||
# 批量生成(count>1):每个变体必须走与单视频完全相同的独立选片流程(#1743/#1749)。
|
||||
# - 变体 0:clone 源 plan(不污染源 plan),变体 1..N-1 用 reselect_plan_for_variant
|
||||
# 完整重跑选片(素材级去重:fresh 优先 → 受控复用 overlap≤20% → 短素材禁复用);
|
||||
@@ -930,6 +941,42 @@ def create_generation_task(
|
||||
)
|
||||
# 变体序号写入 extra_meta(响应/排查时可辨识)
|
||||
task.extra_meta["variant_index"] = task_index
|
||||
|
||||
# smart_edit 逐条预扣(首期固定价,credits_cost=prepaid,不做结算)
|
||||
task_txn_id = ""
|
||||
if charged_task_count > 0:
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
unit_credits = round(smart_edit_charge / count, 2)
|
||||
res = PointsService().deduct_points(
|
||||
user_id=user_id,
|
||||
amount=unit_credits,
|
||||
source="smart_edit",
|
||||
db=db,
|
||||
description="智能剪辑生成预扣",
|
||||
ref_id=task.id,
|
||||
)
|
||||
if not res.get("success"):
|
||||
# 余额不足:退还本次请求已扣积分后整体拒绝
|
||||
already_charged = round(unit_credits * task_index, 2)
|
||||
if already_charged > 0:
|
||||
PointsService().refund_points(
|
||||
user_id=user_id,
|
||||
amount=already_charged,
|
||||
source="smart_edit",
|
||||
db=db,
|
||||
ref_id=task.id,
|
||||
description="智能剪辑批量提交失败退回",
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail=(f"积分不足:智能剪辑每条需 {unit_credits:.2f} 积分,当前余额 {res.get('balance', 0)}"),
|
||||
)
|
||||
task_txn_id = str(res.get("transaction_id") or "")
|
||||
task.credits_prepaid = unit_credits
|
||||
task.credits_cost = unit_credits
|
||||
task.credits_transaction_id = task_txn_id
|
||||
generation_task_repository.update(task)
|
||||
try:
|
||||
# 兜底关联编辑计划:前端未传 source_edit_plan_id 时,
|
||||
# 通过 template_id + user_id 在 DB 层直接查找最新的 plan。
|
||||
|
||||
@@ -228,6 +228,8 @@ class GpuLipsyncService:
|
||||
lipsync_job_id: str = "",
|
||||
user_id: str = "",
|
||||
project_id: str = "",
|
||||
credits_prepaid: float = 0.0,
|
||||
credits_transaction_id: str = "",
|
||||
) -> GpuLipsyncTaskModel:
|
||||
task_id = str(uuid.uuid4())
|
||||
now = datetime.now(UTC)
|
||||
@@ -240,6 +242,8 @@ class GpuLipsyncService:
|
||||
audio_url=audio_url,
|
||||
status="pending",
|
||||
attempt=0,
|
||||
credits_prepaid=float(credits_prepaid or 0.0),
|
||||
credits_transaction_id=str(credits_transaction_id or ""),
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
|
||||
@@ -38,6 +38,7 @@ from sqlalchemy.orm import Session
|
||||
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
|
||||
from packages.application.cosyvoice_service import CosyVoiceError
|
||||
from packages.config import get_api_settings
|
||||
from packages.domain import feature_pricing_service
|
||||
from packages.domain.sentence_timings import (
|
||||
compute_sentence_timings,
|
||||
probe_audio_duration,
|
||||
@@ -368,6 +369,8 @@ class LipsyncService:
|
||||
lipsync_job_id=job.id,
|
||||
user_id=job.user_id,
|
||||
project_id=job.project_id,
|
||||
credits_prepaid=float(getattr(job, "credits_prepaid", 0) or 0),
|
||||
credits_transaction_id=str(getattr(job, "credits_transaction_id", "") or ""),
|
||||
)
|
||||
logger.info(
|
||||
"[lipsync] 已创建 GPU 任务(异步): job_id=%s gpu_task=%s",
|
||||
@@ -415,6 +418,121 @@ class LipsyncService:
|
||||
job.output_duration,
|
||||
)
|
||||
|
||||
# ── lip_sync 计费辅助 ────────────────────────────────────────────────
|
||||
|
||||
@staticmethod
|
||||
def _estimate_duration(
|
||||
*,
|
||||
audio_duration: Optional[float] = None,
|
||||
sentence_timings: Optional[list] = None,
|
||||
script_text: str = "",
|
||||
) -> float:
|
||||
"""预估音频/成片秒数。
|
||||
|
||||
优先级:audio_duration(预合成前端已 ffprobe)> timings 末句 end_time >
|
||||
脚本字数 / 5 字每秒 > 默认 10 秒。
|
||||
"""
|
||||
if audio_duration and float(audio_duration) > 0:
|
||||
return float(audio_duration)
|
||||
if sentence_timings:
|
||||
max_end = 0.0
|
||||
for item in sentence_timings:
|
||||
if isinstance(item, dict):
|
||||
end = item.get("end_time") or item.get("end") or 0.0
|
||||
else:
|
||||
end = 0.0
|
||||
try:
|
||||
max_end = max(max_end, float(end))
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
if max_end > 0:
|
||||
return max_end
|
||||
text = (script_text or "").strip()
|
||||
if text:
|
||||
return max(1.0, len(text) / 5.0)
|
||||
return 10.0
|
||||
|
||||
def _settle_lip_sync(self, job: LipsyncJobModel, actual_duration: float) -> None:
|
||||
"""按实际时长结算(首期只退不补:final < prepaid 退差额,> 不补)。
|
||||
|
||||
幂等:credits_cost 已 > 0 说明结算过,直接跳过。
|
||||
结算失败不阻塞业务(结果已产出),仅记录日志。
|
||||
"""
|
||||
try:
|
||||
prepaid = float(getattr(job, "credits_prepaid", 0) or 0)
|
||||
if prepaid <= 0:
|
||||
return
|
||||
if float(getattr(job, "credits_cost", 0) or 0) > 0:
|
||||
return
|
||||
feature_cfg = feature_pricing_service.get_feature_config("lip_sync")
|
||||
unit_cost = float(feature_cfg.dynamic_unit_cost) if feature_cfg is not None else 0.0
|
||||
duration = float(actual_duration or 0.0)
|
||||
if duration <= 0:
|
||||
duration = self._estimate_duration(
|
||||
sentence_timings=job.sentence_timings,
|
||||
script_text=job.script_text,
|
||||
)
|
||||
final_price, _bd = feature_pricing_service.calculate_price("lip_sync", duration * unit_cost)
|
||||
final_price = round(float(final_price), 2)
|
||||
job.credits_cost = final_price
|
||||
if final_price < prepaid - 0.009:
|
||||
refund = round(prepaid - final_price, 2)
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
res = PointsService().refund_points(
|
||||
user_id=job.user_id,
|
||||
amount=refund,
|
||||
source="lip_sync",
|
||||
db=self.db,
|
||||
ref_id=str(job.credits_transaction_id or job.id),
|
||||
description="对口型结算退费",
|
||||
)
|
||||
if not res.get("success"):
|
||||
logger.warning(
|
||||
"[lip_sync] 结算退费失败 job_id=%s refund=%.2f(不阻塞)",
|
||||
job.id,
|
||||
refund,
|
||||
)
|
||||
# final > prepaid:首期只退不补,不补扣
|
||||
self.db.commit()
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception("[lip_sync] 结算异常 job_id=%s(不阻塞结果)", job.id)
|
||||
try:
|
||||
self.db.rollback()
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
def _refund_lip_sync(self, job: LipsyncJobModel) -> None:
|
||||
"""任务失败/取消时全额退还预扣积分(credits_cost 已结算则退实际未消耗部分)。"""
|
||||
try:
|
||||
prepaid = float(getattr(job, "credits_prepaid", 0) or 0)
|
||||
if prepaid <= 0:
|
||||
return
|
||||
txn_id = str(getattr(job, "credits_transaction_id", "") or "")
|
||||
cost = float(getattr(job, "credits_cost", 0) or 0)
|
||||
refund = round(prepaid - cost, 2) if cost > 0 else round(prepaid, 2)
|
||||
if refund <= 0:
|
||||
return
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
res = PointsService().refund_points(
|
||||
user_id=job.user_id,
|
||||
amount=refund,
|
||||
source="lip_sync",
|
||||
db=self.db,
|
||||
ref_id=txn_id or job.id,
|
||||
description="对口型失败/取消退款",
|
||||
)
|
||||
if res.get("success"):
|
||||
job.credits_cost = prepaid # 标记已全额退回,防重复退
|
||||
self.db.commit()
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception("[lip_sync] 退款异常 job_id=%s", job.id)
|
||||
try:
|
||||
self.db.rollback()
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
# ── 创建任务 ──────────────────────────────────────────────────────────
|
||||
|
||||
def create_job(
|
||||
@@ -466,6 +584,35 @@ class LipsyncService:
|
||||
if not isinstance(sentence_timings, list) or len(sentence_timings) == 0:
|
||||
raise MediaKitError("预合成模式 sentence_timings 不能为空", code="InvalidInput")
|
||||
|
||||
# 0.5 lip_sync 计费预扣(全局 points 开关 + 功能开关均开才扣)
|
||||
prepaid_credits = 0.0
|
||||
prepaid_txn_id = ""
|
||||
if feature_pricing_service.is_feature_enabled("lip_sync"):
|
||||
est_duration = self._estimate_duration(
|
||||
audio_duration=audio_duration,
|
||||
sentence_timings=sentence_timings,
|
||||
script_text=script_text,
|
||||
)
|
||||
feature_cfg = feature_pricing_service.get_feature_config("lip_sync")
|
||||
unit_cost = float(feature_cfg.dynamic_unit_cost) if feature_cfg is not None else 0.0
|
||||
dynamic_cost = est_duration * unit_cost
|
||||
prepaid_credits, _bd = feature_pricing_service.calculate_price("lip_sync", dynamic_cost)
|
||||
if prepaid_credits > 0:
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
res = PointsService().deduct_points(
|
||||
user_id=user_id,
|
||||
amount=prepaid_credits,
|
||||
source="lip_sync",
|
||||
db=self.db,
|
||||
description="对口型生成预扣",
|
||||
)
|
||||
if not res.get("success"):
|
||||
raise ValueError(
|
||||
f"积分不足:本次对口型需 {prepaid_credits:.2f} 积分,当前余额 {res.get('balance', 0)}"
|
||||
)
|
||||
prepaid_txn_id = str(res.get("transaction_id") or "")
|
||||
|
||||
# 1. 创建数据库记录
|
||||
job_id = str(uuid.uuid4())
|
||||
job = LipsyncJobModel(
|
||||
@@ -482,6 +629,8 @@ class LipsyncService:
|
||||
emotion=emotion or "",
|
||||
# 音频直传(含预合成)直接进入 pending(后续同步改为 submitted);TTS 模式进入 tts_processing
|
||||
status="tts_processing" if is_tts_mode else "pending",
|
||||
credits_prepaid=prepaid_credits,
|
||||
credits_transaction_id=prepaid_txn_id,
|
||||
)
|
||||
self.db.add(job)
|
||||
self.db.flush()
|
||||
@@ -677,6 +826,8 @@ class LipsyncService:
|
||||
job.completed_at = _now
|
||||
job.updated_at = _now
|
||||
self.db.commit()
|
||||
# lip_sync 超时全额退款
|
||||
self._refund_lip_sync(job)
|
||||
return job
|
||||
|
||||
# 未提交的任务不轮询
|
||||
@@ -702,6 +853,8 @@ class LipsyncService:
|
||||
job.completed_at = datetime.now(UTC)
|
||||
job.updated_at = datetime.now(UTC)
|
||||
self.db.commit()
|
||||
# lip_sync 结算(只退不补)
|
||||
self._settle_lip_sync(job, float(job.output_duration or 0.0))
|
||||
# 异步转存自家 OSS
|
||||
try:
|
||||
from app.tasks.lipsync_tts import persist_output_video_task
|
||||
@@ -719,6 +872,8 @@ class LipsyncService:
|
||||
job.error_message = error.get("message", "任务执行失败")
|
||||
job.error_code = error.get("code", "TaskFailed")
|
||||
job.completed_at = datetime.now(UTC)
|
||||
# lip_sync 失败全额退款(先退款再统一 commit)
|
||||
self._refund_lip_sync(job)
|
||||
else:
|
||||
# 中间状态(running/processing/queued 等)同步到 DB,避免前端永远卡在 submitted
|
||||
if isinstance(mk_status, str) and mk_status:
|
||||
@@ -812,6 +967,8 @@ class LipsyncService:
|
||||
job.status = "cancelled"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
self.db.commit()
|
||||
# lip_sync 取消全额退款
|
||||
self._refund_lip_sync(job)
|
||||
self.db.refresh(job)
|
||||
|
||||
return job
|
||||
|
||||
@@ -104,6 +104,7 @@ def lipsync_gpu_process_async(self, job_id: str, user_id: str, gpu_task_id: str)
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
logger.info("[lipsync_gpu_async] GPU 任务已被用户取消: job_id=%s", job_id)
|
||||
_refund_lip_sync(db, job)
|
||||
return
|
||||
|
||||
if final_task.status != "done":
|
||||
@@ -141,6 +142,7 @@ def lipsync_gpu_process_async(self, job_id: str, user_id: str, gpu_task_id: str)
|
||||
job_id,
|
||||
job.output_duration,
|
||||
)
|
||||
_settle_lip_sync(db, job, final_task)
|
||||
except Exception as exc:
|
||||
logger.exception("[lipsync_gpu_async] 异常: job_id=%s err=%s", job_id, exc)
|
||||
try:
|
||||
@@ -157,6 +159,33 @@ def lipsync_gpu_process_async(self, job_id: str, user_id: str, gpu_task_id: str)
|
||||
db.close()
|
||||
|
||||
|
||||
def _settle_lip_sync(db: Session, job: LipsyncJobModel, gpu_task) -> None:
|
||||
"""GPU 成功后结算:同步 credits_cost 到 gpu 任务并按实际时长多退少不补。"""
|
||||
try:
|
||||
from app.services.lipsync_service import LipsyncService
|
||||
|
||||
# GPU 任务表先同步结算结果(标记用)
|
||||
LipsyncService._settle_lip_sync(job, float(getattr(gpu_task, "result_duration", 0) or 0.0))
|
||||
gpu_task.credits_cost = float(job.credits_cost or 0.0)
|
||||
db.commit()
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception("[lipsync_gpu_async] lip_sync 结算异常 job_id=%s(不阻塞)", job.id)
|
||||
try:
|
||||
db.rollback()
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
|
||||
def _refund_lip_sync(db: Session, job: LipsyncJobModel) -> None:
|
||||
"""GPU 取消/失败路径全额退款。"""
|
||||
try:
|
||||
from app.services.lipsync_service import LipsyncService
|
||||
|
||||
LipsyncService(db)._refund_lip_sync(job)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception("[lipsync_gpu_async] lip_sync 退款异常 job_id=%s", job.id)
|
||||
|
||||
|
||||
def _fallback_to_mediakit(db: Session, job: LipsyncJobModel) -> None:
|
||||
"""GPU 失败时回退到 MediaKit 云端渲染。"""
|
||||
try:
|
||||
|
||||
@@ -176,7 +176,7 @@ export interface ViralVideoJob {
|
||||
fusion_level?: FusionLevel
|
||||
voice_id?: string
|
||||
voice_mode?: "global" | "per_video"
|
||||
voice_source?: "preset" | "library" | "clone" | "upload"
|
||||
voice_source?: "preset" | "library" | "clone" | "upload" | "my_voice"
|
||||
bgm_preference?: string
|
||||
intent_result?: IntentResult
|
||||
intent_text?: string
|
||||
@@ -212,7 +212,7 @@ export interface GenerateViralVideoRequest {
|
||||
user_copy_text?: string
|
||||
fusion_level?: FusionLevel
|
||||
voice_id?: string
|
||||
voice_source?: "preset" | "library" | "clone" | "upload"
|
||||
voice_source?: "preset" | "library" | "clone" | "upload" | "my_voice"
|
||||
bgm_preference?: string
|
||||
industry?: string
|
||||
target_customer?: string
|
||||
@@ -244,7 +244,7 @@ export interface AnalyzeImagesRequest {
|
||||
/** TTS 音色 ID(STEP1 已选音色时传) */
|
||||
voice_id?: string
|
||||
/** 音色来源:preset | library | clone | upload */
|
||||
voice_source?: "preset" | "library" | "clone" | "upload"
|
||||
voice_source?: "preset" | "library" | "clone" | "upload" | "my_voice"
|
||||
/** Seedance 视频比例:9:16 | 16:9 | 1:1 */
|
||||
video_ratio?: string
|
||||
/** Seedance 模型 ID(空则使用服务端默认) */
|
||||
@@ -273,7 +273,7 @@ export interface GenerateCopyRequest {
|
||||
/** TTS 音色 ID(优先级高于 persona_id) */
|
||||
voice_id?: string
|
||||
/** 音色来源:preset | library | clone | upload */
|
||||
voice_source?: "preset" | "library" | "clone" | "upload"
|
||||
voice_source?: "preset" | "library" | "clone" | "upload" | "my_voice"
|
||||
/** Seedance 视频比例(9:16/16:9/1:1 等) */
|
||||
video_ratio?: string
|
||||
/** Seedance 模型 ID(空则使用服务端默认) */
|
||||
|
||||
@@ -18,6 +18,8 @@ export interface VoiceClone {
|
||||
language: string
|
||||
gender: string
|
||||
error_message: string | null
|
||||
/** CosyVoice 实际使用的音色 ID(status=ready 时由后端填充,用于 TTS 调用) */
|
||||
voice_id?: string | null
|
||||
created_at: string
|
||||
updated_at: string
|
||||
}
|
||||
|
||||
@@ -18,6 +18,7 @@ export const toVoiceClone = (profile: VoiceCloneProfile): VoiceClone => ({
|
||||
language: profile.language || "",
|
||||
gender: profile.gender || "",
|
||||
error_message: profile.error_message || null,
|
||||
voice_id: profile.voice_id,
|
||||
created_at: profile.created_at,
|
||||
updated_at: profile.updated_at,
|
||||
})
|
||||
|
||||
@@ -0,0 +1,182 @@
|
||||
/* DurationWheelPicker —— 弹层式滚轮选择器(样式与表单一致) */
|
||||
|
||||
/* 触发按钮:外观复用 .vv-select 风格 */
|
||||
.dw-trigger {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
width: 100%;
|
||||
height: 36px;
|
||||
padding: 0 12px;
|
||||
background: #fff;
|
||||
border: 1px solid #e0e0e8;
|
||||
border-radius: 8px;
|
||||
font-size: 13px;
|
||||
color: #1f2937;
|
||||
cursor: pointer;
|
||||
box-sizing: border-box;
|
||||
transition: all 0.15s;
|
||||
user-select: none;
|
||||
}
|
||||
.dw-trigger:hover {
|
||||
border-color: #c0c0d0;
|
||||
}
|
||||
.dw-trigger-open,
|
||||
.dw-trigger:focus-within {
|
||||
border-color: #7c3aed !important;
|
||||
box-shadow: 0 0 0 2px rgba(124, 58, 237, 0.12);
|
||||
}
|
||||
.dw-trigger-disabled {
|
||||
opacity: 0.5;
|
||||
pointer-events: none;
|
||||
cursor: not-allowed;
|
||||
}
|
||||
.dw-trigger-val {
|
||||
flex: 1;
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
white-space: nowrap;
|
||||
}
|
||||
.dw-trigger-placeholder {
|
||||
color: #9ca3af;
|
||||
}
|
||||
.dw-trigger-arrow {
|
||||
font-size: 10px;
|
||||
color: #9ca3af;
|
||||
margin-left: 8px;
|
||||
transition: transform 0.2s;
|
||||
}
|
||||
.dw-trigger-arrow-up {
|
||||
transform: rotate(180deg);
|
||||
}
|
||||
|
||||
/* 弹层容器 */
|
||||
.dw-popup {
|
||||
padding: 8px;
|
||||
min-width: 140px;
|
||||
}
|
||||
|
||||
/* 滚轮 */
|
||||
.dw-picker {
|
||||
position: relative;
|
||||
width: 100%;
|
||||
overflow: hidden;
|
||||
border-radius: 8px;
|
||||
background: #fafafe;
|
||||
border: 1px solid #e5e7eb;
|
||||
}
|
||||
.dw-picker-list {
|
||||
margin: 0;
|
||||
padding: 0;
|
||||
list-style: none;
|
||||
height: 100%;
|
||||
overflow-y: scroll;
|
||||
scroll-snap-type: y mandatory;
|
||||
-webkit-overflow-scrolling: touch;
|
||||
scrollbar-width: none;
|
||||
}
|
||||
.dw-picker-list::-webkit-scrollbar {
|
||||
display: none;
|
||||
}
|
||||
.dw-picker-item {
|
||||
display: flex;
|
||||
align-items: baseline;
|
||||
justify-content: center;
|
||||
gap: 3px;
|
||||
scroll-snap-align: center;
|
||||
cursor: pointer;
|
||||
font-size: 15px;
|
||||
color: #9ca3af;
|
||||
font-weight: 400;
|
||||
transition:
|
||||
color 0.15s,
|
||||
transform 0.15s,
|
||||
font-weight 0.15s;
|
||||
}
|
||||
.dw-picker-item-val {
|
||||
font-variant-numeric: tabular-nums;
|
||||
}
|
||||
.dw-picker-item-unit {
|
||||
font-size: 13px;
|
||||
color: inherit;
|
||||
}
|
||||
.dw-picker-item-active {
|
||||
color: #7c3aed;
|
||||
font-weight: 600;
|
||||
}
|
||||
.dw-picker-item-active .dw-picker-item-val {
|
||||
font-size: 18px;
|
||||
}
|
||||
.dw-picker-item-active .dw-picker-item-unit {
|
||||
font-size: 14px;
|
||||
}
|
||||
|
||||
/* 中心选中条 */
|
||||
.dw-picker-mask {
|
||||
position: absolute;
|
||||
left: 6px;
|
||||
right: 6px;
|
||||
pointer-events: none;
|
||||
background: #f5f0ff;
|
||||
border-radius: 6px;
|
||||
z-index: 1;
|
||||
}
|
||||
.dw-picker-mask::before,
|
||||
.dw-picker-mask::after {
|
||||
content: "";
|
||||
position: absolute;
|
||||
left: 0;
|
||||
right: 0;
|
||||
height: 1px;
|
||||
background: #d8c4ff;
|
||||
}
|
||||
.dw-picker-mask::before {
|
||||
top: 0;
|
||||
}
|
||||
.dw-picker-mask::after {
|
||||
bottom: 0;
|
||||
}
|
||||
|
||||
/* 上下渐变 */
|
||||
.dw-picker-fade {
|
||||
position: absolute;
|
||||
left: 0;
|
||||
right: 0;
|
||||
height: 40%;
|
||||
pointer-events: none;
|
||||
z-index: 2;
|
||||
}
|
||||
.dw-picker-fade-top {
|
||||
top: 0;
|
||||
background: linear-gradient(to bottom, #fafafe 25%, rgba(250, 250, 254, 0));
|
||||
}
|
||||
.dw-picker-fade-bottom {
|
||||
bottom: 0;
|
||||
background: linear-gradient(to top, #fafafe 25%, rgba(250, 250, 254, 0));
|
||||
}
|
||||
|
||||
/* 弹层按钮区 */
|
||||
.dw-popup-actions {
|
||||
display: flex;
|
||||
gap: 8px;
|
||||
justify-content: flex-end;
|
||||
margin-top: 8px;
|
||||
}
|
||||
.dw-popup-actions .ant-btn {
|
||||
border-radius: 6px;
|
||||
}
|
||||
.dw-popup-actions .ant-btn-primary {
|
||||
background: #7c3aed;
|
||||
}
|
||||
.dw-popup-actions .ant-btn-primary:hover {
|
||||
background: #6d28d9 !important;
|
||||
}
|
||||
|
||||
/* 覆盖 antd Popover 默认内边距 */
|
||||
.dw-popover .ant-popover-inner {
|
||||
padding: 0 !important;
|
||||
overflow: hidden;
|
||||
}
|
||||
.dw-popover .ant-popover-arrow {
|
||||
display: none;
|
||||
}
|
||||
@@ -0,0 +1,180 @@
|
||||
/**
|
||||
* DurationWheelPicker —— 竖屏滚轮式时长选择器(弹层版)
|
||||
*
|
||||
* 设计:
|
||||
* - 外观是和其他表单 Select 一致的输入框(白色底+1px灰边+紫色focus ring)
|
||||
* - 点击输入框弹出 Popover,内部是滚轮 picker(原生 scroll-snap,零依赖)
|
||||
* - 滚轮样式:白底容器,选中行 #7c3aed 紫字加粗+浅紫背景条
|
||||
* - 支持触摸/鼠标滚轮/点击;松手吸附;底部"确认/取消"按钮
|
||||
* - 默认范围 15–30 秒,步长 1 秒
|
||||
*/
|
||||
import React, { useEffect, useMemo, useRef, useState, useCallback } from "react"
|
||||
import { Popover, Button } from "antd"
|
||||
import { DownOutlined } from "@ant-design/icons"
|
||||
import "./DurationWheelPicker.css"
|
||||
|
||||
export interface DurationWheelPickerProps {
|
||||
value?: number
|
||||
min?: number
|
||||
max?: number
|
||||
step?: number
|
||||
unit?: string
|
||||
onChange?: (value: number) => void
|
||||
placeholder?: string
|
||||
disabled?: boolean
|
||||
/** 弹层宽度,默认 160px */
|
||||
popupWidth?: number
|
||||
/** 弹层内滚轮高度,默认 180px */
|
||||
wheelHeight?: number
|
||||
}
|
||||
|
||||
const ITEM_HEIGHT = 36
|
||||
|
||||
const DurationWheelPicker: React.FC<DurationWheelPickerProps> = ({
|
||||
value = 20,
|
||||
min = 15,
|
||||
max = 30,
|
||||
step = 1,
|
||||
unit = "秒",
|
||||
onChange,
|
||||
placeholder = "请选择时长",
|
||||
disabled = false,
|
||||
popupWidth = 160,
|
||||
wheelHeight = 180,
|
||||
}) => {
|
||||
const options = useMemo(() => {
|
||||
const arr: number[] = []
|
||||
for (let v = min; v <= max; v += step) arr.push(v)
|
||||
return arr
|
||||
}, [min, max, step])
|
||||
|
||||
const [open, setOpen] = useState(false)
|
||||
// 弹层内暂存值,点确认才提交
|
||||
const [draft, setDraft] = useState<number>(value)
|
||||
const listRef = useRef<HTMLUListElement>(null)
|
||||
const scrollTimerRef = useRef<ReturnType<typeof setTimeout> | null>(null)
|
||||
|
||||
useEffect(() => {
|
||||
if (open) {
|
||||
setDraft(value)
|
||||
// 下一帧滚到当前值
|
||||
requestAnimationFrame(() => scrollToValue(value, false))
|
||||
}
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [open])
|
||||
|
||||
const scrollToValue = useCallback(
|
||||
(v: number, smooth = true) => {
|
||||
const list = listRef.current
|
||||
if (!list) return
|
||||
const idx = options.indexOf(v)
|
||||
if (idx < 0) return
|
||||
list.scrollTo({ top: idx * ITEM_HEIGHT, behavior: smooth ? "smooth" : "auto" })
|
||||
},
|
||||
[options],
|
||||
)
|
||||
|
||||
const handleScroll = () => {
|
||||
if (scrollTimerRef.current) clearTimeout(scrollTimerRef.current)
|
||||
scrollTimerRef.current = setTimeout(() => {
|
||||
const list = listRef.current
|
||||
if (!list) return
|
||||
const idx = Math.round(list.scrollTop / ITEM_HEIGHT)
|
||||
const clamped = Math.max(0, Math.min(options.length - 1, idx))
|
||||
const targetTop = clamped * ITEM_HEIGHT
|
||||
if (Math.abs(list.scrollTop - targetTop) > 1) {
|
||||
list.scrollTo({ top: targetTop, behavior: "smooth" })
|
||||
}
|
||||
setDraft(options[clamped])
|
||||
}, 100)
|
||||
}
|
||||
|
||||
const handleConfirm = () => {
|
||||
onChange?.(draft)
|
||||
setOpen(false)
|
||||
}
|
||||
|
||||
const handleCancel = () => {
|
||||
setOpen(false)
|
||||
}
|
||||
|
||||
const handleItemClick = (v: number) => {
|
||||
setDraft(v)
|
||||
scrollToValue(v, true)
|
||||
}
|
||||
|
||||
const maskTop = wheelHeight / 2 - ITEM_HEIGHT / 2
|
||||
|
||||
const wheel = (
|
||||
<div className="dw-popup">
|
||||
<div
|
||||
className="dw-picker"
|
||||
style={{ height: wheelHeight, width: popupWidth - 24 /* padding */ }}
|
||||
>
|
||||
<div className="dw-picker-mask" style={{ top: maskTop, height: ITEM_HEIGHT }} aria-hidden />
|
||||
<div className="dw-picker-fade dw-picker-fade-top" aria-hidden />
|
||||
<div className="dw-picker-fade dw-picker-fade-bottom" aria-hidden />
|
||||
<ul
|
||||
ref={listRef}
|
||||
className="dw-picker-list"
|
||||
onScroll={handleScroll}
|
||||
style={{
|
||||
paddingTop: wheelHeight / 2 - ITEM_HEIGHT / 2,
|
||||
paddingBottom: wheelHeight / 2 - ITEM_HEIGHT / 2,
|
||||
}}
|
||||
>
|
||||
{options.map((v) => {
|
||||
const isActive = v === draft
|
||||
return (
|
||||
<li
|
||||
key={v}
|
||||
className={`dw-picker-item${isActive ? " dw-picker-item-active" : ""}`}
|
||||
style={{ height: ITEM_HEIGHT, lineHeight: `${ITEM_HEIGHT}px` }}
|
||||
onClick={() => handleItemClick(v)}
|
||||
aria-selected={isActive}
|
||||
role="option"
|
||||
>
|
||||
<span className="dw-picker-item-val">{v}</span>
|
||||
<span className="dw-picker-item-unit">{unit}</span>
|
||||
</li>
|
||||
)
|
||||
})}
|
||||
</ul>
|
||||
</div>
|
||||
<div className="dw-popup-actions">
|
||||
<Button size="small" onClick={handleCancel}>
|
||||
取消
|
||||
</Button>
|
||||
<Button size="small" type="primary" onClick={handleConfirm}>
|
||||
确认
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
|
||||
return (
|
||||
<Popover
|
||||
open={!disabled && open}
|
||||
onOpenChange={(v) => setOpen(v)}
|
||||
content={wheel}
|
||||
trigger="click"
|
||||
placement="bottomLeft"
|
||||
overlayClassName="dw-popover"
|
||||
overlayStyle={{ padding: 0 }}
|
||||
overlayInnerStyle={{ padding: 0, borderRadius: 10 }}
|
||||
destroyTooltipOnHide
|
||||
>
|
||||
<div
|
||||
className={`dw-trigger${disabled ? " dw-trigger-disabled" : ""}${open ? " dw-trigger-open" : ""}`}
|
||||
style={{ height: 36 }}
|
||||
>
|
||||
<span className={`dw-trigger-val${value != null ? "" : " dw-trigger-placeholder"}`}>
|
||||
{value != null ? `${value}${unit}` : placeholder}
|
||||
</span>
|
||||
<DownOutlined className={`dw-trigger-arrow${open ? " dw-trigger-arrow-up" : ""}`} />
|
||||
</div>
|
||||
</Popover>
|
||||
)
|
||||
}
|
||||
|
||||
export default DurationWheelPicker
|
||||
@@ -10,4 +10,4 @@
|
||||
* 功能流程不做积分预校验,直接走生成。
|
||||
* - true:展示完整积分系统 UI。
|
||||
*/
|
||||
export const ENABLE_CREDIT_SYSTEM = false
|
||||
export const ENABLE_CREDIT_SYSTEM = true
|
||||
|
||||
@@ -1834,3 +1834,117 @@
|
||||
color: #7c3aed;
|
||||
font-weight: 500;
|
||||
}
|
||||
|
||||
/* ── v1.6 我的音色(默认主路径) ── */
|
||||
.vv-voice-section {
|
||||
margin-top: 4px;
|
||||
border: 1px solid #e5e7eb;
|
||||
border-radius: 8px;
|
||||
background: #fff;
|
||||
padding: 10px;
|
||||
}
|
||||
.vv-voice-section-head {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
margin-bottom: 8px;
|
||||
}
|
||||
.vv-voice-section-title {
|
||||
font-size: 13px;
|
||||
font-weight: 600;
|
||||
color: #1f2937;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
}
|
||||
.vv-voice-section-title .anticon {
|
||||
color: #7c3aed;
|
||||
}
|
||||
.vv-link-btn-sm {
|
||||
font-size: 12px;
|
||||
padding: 2px 6px;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 4px;
|
||||
}
|
||||
.vv-my-voice-list {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 4px;
|
||||
max-height: 220px;
|
||||
overflow-y: auto;
|
||||
}
|
||||
.vv-my-voice-item {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 10px;
|
||||
padding: 8px 10px;
|
||||
border-radius: 6px;
|
||||
cursor: pointer;
|
||||
transition: background 0.15s;
|
||||
}
|
||||
.vv-my-voice-item:hover {
|
||||
background: #f5f0ff;
|
||||
}
|
||||
.vv-my-voice-item.selected {
|
||||
background: #f5f0ff;
|
||||
}
|
||||
.vv-my-voice-item.selected .vv-voice-radio {
|
||||
border-color: #7c3aed;
|
||||
background: #7c3aed;
|
||||
box-shadow: inset 0 0 0 2px #fff;
|
||||
}
|
||||
.vv-voice-empty {
|
||||
padding: 16px 12px;
|
||||
text-align: center;
|
||||
color: #9ca3af;
|
||||
font-size: 12px;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
}
|
||||
.vv-voice-empty-ic {
|
||||
font-size: 28px;
|
||||
color: #d1d5db;
|
||||
}
|
||||
.vv-voice-empty-text {
|
||||
font-size: 12px;
|
||||
}
|
||||
.vv-voice-alt-row {
|
||||
display: flex;
|
||||
gap: 8px;
|
||||
margin-top: 8px;
|
||||
}
|
||||
.vv-voice-alt-btn {
|
||||
flex: 1;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
gap: 6px;
|
||||
padding: 8px 10px;
|
||||
background: #fafafe;
|
||||
border: 1px solid #e5e7eb;
|
||||
border-radius: 6px;
|
||||
color: #6b7280;
|
||||
font-size: 12px;
|
||||
cursor: pointer;
|
||||
transition: all 0.15s;
|
||||
}
|
||||
.vv-voice-alt-btn:hover {
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
background: #f5f0ff;
|
||||
}
|
||||
.vv-voice-alt-btn.selected {
|
||||
background: #f5f0ff;
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
}
|
||||
.vv-voice-panel-actions {
|
||||
display: flex;
|
||||
gap: 12px;
|
||||
margin-bottom: 6px;
|
||||
padding-bottom: 6px;
|
||||
border-bottom: 1px dashed #e5e7eb;
|
||||
}
|
||||
|
||||
@@ -23,10 +23,13 @@ import {
|
||||
CheckCircleFilled,
|
||||
EditOutlined,
|
||||
HistoryOutlined,
|
||||
UserOutlined,
|
||||
} from "@ant-design/icons"
|
||||
import { Select, Input, message } from "antd"
|
||||
import { uploadAssetDirect, getAssetLibraries, getAssetsByKind, type AssetItem } from "@/api/assets"
|
||||
import { fetchPresetVoices } from "@/api/voices"
|
||||
import { getVoiceClones, getVoiceClonePreview } from "@/api/voice-clone"
|
||||
import type { VoiceClone } from "@/api/voice-clone"
|
||||
import {
|
||||
FUSION_LEVELS,
|
||||
STYLE_STRENGTHS,
|
||||
@@ -94,13 +97,15 @@ type RefVideo = {
|
||||
} | null
|
||||
|
||||
type RefAudio = {
|
||||
source: "preset" | "library" | "upload" | "clone"
|
||||
source: "preset" | "library" | "upload" | "clone" | "my_voice"
|
||||
id?: string
|
||||
name: string
|
||||
url?: string
|
||||
/** 克隆音色的 CosyVoice voice_id(source=my_voice 时使用) */
|
||||
ttsVoiceId?: string
|
||||
} | null
|
||||
|
||||
type VoicePanel = "preset" | "upload" | "library" | "record" | "douyin" | null
|
||||
type VoicePanel = "my_voices" | "preset" | "upload" | "library" | "record" | "douyin" | null
|
||||
|
||||
type TabTask = {
|
||||
id: string
|
||||
@@ -213,12 +218,12 @@ const PURPOSES = [
|
||||
"悬念短剧",
|
||||
"情绪短片",
|
||||
]
|
||||
const DURATIONS = [15, 20, 30, 45, 60]
|
||||
const RATIOS = [
|
||||
{ v: "9:16", label: "9:16 竖屏(抖音/视频号)" },
|
||||
{ v: "16:9", label: "16:9 横屏(B站/YouTube)" },
|
||||
{ v: "1:1", label: "1:1 方形(小红书)" },
|
||||
]
|
||||
const DURATIONS = Array.from({ length: 16 }, (_, i) => 15 + i)
|
||||
/** 兜底模型列表(接口未返回时使用,字段与 ViralVideoModel 对齐;后端返回后自动覆盖) */
|
||||
const FALLBACK_VIDEO_MODELS: ViralVideoModel[] = [
|
||||
{
|
||||
@@ -432,7 +437,7 @@ const emptyTask = (id: string, title: string): TabTask => ({
|
||||
language: "中文(普通话)",
|
||||
viralStructure: STRUCTURES[0],
|
||||
marketingPurpose: "",
|
||||
duration: 15,
|
||||
duration: 20,
|
||||
persona: "",
|
||||
videoRatio: "9:16",
|
||||
videoModel: "seedance-2.5",
|
||||
@@ -485,6 +490,9 @@ const ViralVideoPage: React.FC = () => {
|
||||
{ id: string; name: string; url: string; desc?: string }[]
|
||||
>([])
|
||||
const [presetVoices, setPresetVoices] = useState<PresetVoice[]>([])
|
||||
/** 我的克隆音色(GET /voice-clones?status=ready) */
|
||||
const [myVoices, setMyVoices] = useState<VoiceClone[]>([])
|
||||
const [myVoicesLoading, setMyVoicesLoading] = useState(false)
|
||||
const [voicePickerOpen, setVoicePickerOpen] = useState(false)
|
||||
const [videoModels, setVideoModels] = useState<ViralVideoModel[]>(FALLBACK_VIDEO_MODELS)
|
||||
const [assetPicker, setAssetPicker] = useState<{
|
||||
@@ -514,8 +522,18 @@ const ViralVideoPage: React.FC = () => {
|
||||
[activeId],
|
||||
)
|
||||
|
||||
/* ── 加载音色 ── */
|
||||
/* ── 加载音色(我的克隆音色 + 上传的配音素材 + 预设音色) ── */
|
||||
const reloadMyVoices = useCallback(() => {
|
||||
setMyVoicesLoading(true)
|
||||
getVoiceClones({ status: "ready", limit: 50 })
|
||||
.then((list) => {
|
||||
setMyVoices(list.filter((v) => v.status === "ready"))
|
||||
})
|
||||
.catch(() => setMyVoices([]))
|
||||
.finally(() => setMyVoicesLoading(false))
|
||||
}, [])
|
||||
useEffect(() => {
|
||||
reloadMyVoices()
|
||||
getAssetsByKind("voice", { limit: 50 })
|
||||
.then((list) =>
|
||||
setVoiceAssets(
|
||||
@@ -523,7 +541,7 @@ const ViralVideoPage: React.FC = () => {
|
||||
id: a.id,
|
||||
name: a.name,
|
||||
url: a.file_url || "",
|
||||
desc: a.duration ? `${Math.round(a.duration)}s` : "我的音色",
|
||||
desc: a.duration ? `${Math.round(a.duration)}s` : "配音文件",
|
||||
})),
|
||||
),
|
||||
)
|
||||
@@ -544,7 +562,7 @@ const ViralVideoPage: React.FC = () => {
|
||||
// 接口失败兜底:弹窗内 MOCK_VOICES 会生效(voice_id 与后端 CosyVoice 一致)
|
||||
setPresetVoices([])
|
||||
})
|
||||
}, [])
|
||||
}, [reloadMyVoices])
|
||||
|
||||
/* ── 加载视频模型列表 ── */
|
||||
useEffect(() => {
|
||||
@@ -914,21 +932,36 @@ const ViralVideoPage: React.FC = () => {
|
||||
[setTask, uploadFile],
|
||||
)
|
||||
|
||||
const onCloneSuccess = (voice: { id?: string; name?: string }) => {
|
||||
const onCloneSuccess = (
|
||||
voice: VoiceClone | { id?: string; name?: string; voice_id?: string | null },
|
||||
) => {
|
||||
setCloneModalOpen(false)
|
||||
// 克隆成功后刷新我的音色列表
|
||||
reloadMyVoices()
|
||||
const vid = "voice_id" in voice ? voice.voice_id : undefined
|
||||
if (voice?.id) {
|
||||
setTask({
|
||||
refAudio: { source: "clone", id: voice.id, name: voice.name || "我录制的音色" },
|
||||
refAudio: {
|
||||
source: "my_voice",
|
||||
id: voice.id,
|
||||
name: voice.name || "我录制的音色",
|
||||
ttsVoiceId: vid || undefined,
|
||||
},
|
||||
voicePanel: null,
|
||||
})
|
||||
getAssetsByKind("voice", { limit: 50 })
|
||||
.then((list) =>
|
||||
setVoiceAssets(
|
||||
list.map((a) => ({ id: a.id, name: a.name, url: a.file_url || "", desc: "我的音色" })),
|
||||
list.map((a) => ({
|
||||
id: a.id,
|
||||
name: a.name,
|
||||
url: a.file_url || "",
|
||||
desc: a.duration ? `${Math.round(a.duration)}s` : "配音文件",
|
||||
})),
|
||||
),
|
||||
)
|
||||
.catch(() => {})
|
||||
message.success("录音已完成,已自动选中")
|
||||
message.success("音色克隆完成,已自动选中")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1014,7 +1047,12 @@ const ViralVideoPage: React.FC = () => {
|
||||
: undefined,
|
||||
fusion_level: task.fusionLevel,
|
||||
style_strength: task.refVideo?.ossUrl ? task.styleStrength : undefined,
|
||||
voice_id: task.refAudio?.source === "preset" ? task.refAudio.id : undefined,
|
||||
voice_id:
|
||||
task.refAudio?.source === "preset"
|
||||
? task.refAudio.id
|
||||
: task.refAudio?.source === "my_voice"
|
||||
? task.refAudio.ttsVoiceId || task.refAudio.id
|
||||
: undefined,
|
||||
voice_source: task.refAudio?.source || undefined,
|
||||
video_ratio: task.videoRatio,
|
||||
video_model: task.videoModel,
|
||||
@@ -1058,7 +1096,12 @@ const ViralVideoPage: React.FC = () => {
|
||||
style_strength: task.refVideo?.ossUrl ? task.styleStrength : undefined,
|
||||
user_copy_text: task.storyboard?.voiceover_script || task.userCopy || undefined,
|
||||
fusion_level: task.fusionLevel,
|
||||
voice_id: task.refAudio?.id || undefined,
|
||||
voice_id:
|
||||
task.refAudio?.source === "preset"
|
||||
? task.refAudio.id
|
||||
: task.refAudio?.source === "my_voice"
|
||||
? task.refAudio.ttsVoiceId || task.refAudio.id
|
||||
: undefined,
|
||||
voice_source: task.refAudio?.source,
|
||||
industry: task.industry || undefined,
|
||||
target_customer: task.targetCustomer || undefined,
|
||||
@@ -2033,43 +2076,119 @@ const ViralVideoPage: React.FC = () => {
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* 参考音频 2x2 */}
|
||||
{/* 配音选择(默认"我的音色") */}
|
||||
<div className="vv-subblock">
|
||||
<div className="vv-subblock-head">
|
||||
<span className="vv-subblock-title">参考音频</span>
|
||||
<span className="vv-subblock-title">AI 配音</span>
|
||||
<span className="vv-opt-tag">选填</span>
|
||||
<span className="vv-subblock-hint">最大 10MB · 最长 30 秒</span>
|
||||
<span className="vv-subblock-hint">
|
||||
默认使用我的克隆音色,也可选择内置音色或上传配音
|
||||
</span>
|
||||
</div>
|
||||
<div className="vv-audio-grid">
|
||||
|
||||
{/* 我的音色(主路径) */}
|
||||
<div className="vv-voice-section">
|
||||
<div className="vv-voice-section-head">
|
||||
<span className="vv-voice-section-title">
|
||||
<UserOutlined /> 我的音色
|
||||
</span>
|
||||
<button
|
||||
className="vv-link-btn vv-link-btn-sm"
|
||||
onClick={() => setCloneModalOpen(true)}
|
||||
>
|
||||
<PlusOutlined /> 录制新音色
|
||||
</button>
|
||||
</div>
|
||||
{myVoicesLoading ? (
|
||||
<div className="vv-muted vv-voice-empty">加载中…</div>
|
||||
) : myVoices.length === 0 ? (
|
||||
<div className="vv-voice-empty">
|
||||
<AudioMutedOutlined className="vv-voice-empty-ic" />
|
||||
<div className="vv-voice-empty-text">还没有克隆音色</div>
|
||||
<button
|
||||
className="vv-btn vv-btn-primary vv-btn-sm"
|
||||
onClick={() => setCloneModalOpen(true)}
|
||||
>
|
||||
去录制一个
|
||||
</button>
|
||||
</div>
|
||||
) : (
|
||||
<div className="vv-my-voice-list">
|
||||
{myVoices.map((v) => {
|
||||
const selected =
|
||||
task.refAudio?.source === "my_voice" && task.refAudio.id === v.id
|
||||
const playing = task.playingVoiceId === `my-${v.id}`
|
||||
return (
|
||||
<div
|
||||
key={v.id}
|
||||
className={`vv-my-voice-item ${selected ? "selected" : ""}`}
|
||||
onClick={() =>
|
||||
setTask({
|
||||
refAudio: {
|
||||
source: "my_voice",
|
||||
id: v.id,
|
||||
name: v.name,
|
||||
ttsVoiceId: v.voice_id || undefined,
|
||||
},
|
||||
voicePanel: null,
|
||||
})
|
||||
}
|
||||
>
|
||||
<div className="vv-voice-radio" />
|
||||
<div className="vv-voice-info">
|
||||
<div className="vv-voice-name">{v.name}</div>
|
||||
<div className="vv-voice-desc">
|
||||
AI 克隆声音
|
||||
{v.gender
|
||||
? ` · ${v.gender === "male" ? "男" : v.gender === "female" ? "女" : ""}`
|
||||
: ""}
|
||||
</div>
|
||||
</div>
|
||||
<button
|
||||
className={`vv-voice-play ${playing ? "playing" : ""}`}
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
// 克隆音色试听:优先 sample_url,否则调用 preview 接口
|
||||
if (v.sample_url) {
|
||||
toggleVoice({ id: `my-${v.id}`, url: v.sample_url })
|
||||
} else {
|
||||
message.info("正在获取试听音频…")
|
||||
getVoiceClonePreview(v.id)
|
||||
.then((r) => {
|
||||
if (r.audio_url)
|
||||
toggleVoice({ id: `my-${v.id}`, url: r.audio_url })
|
||||
})
|
||||
.catch(() => message.error("试听获取失败"))
|
||||
}
|
||||
}}
|
||||
>
|
||||
{playing ? <PauseCircleOutlined /> : <PlayCircleOutlined />}
|
||||
</button>
|
||||
</div>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* 次要选项:内置音色 / 上传配音 */}
|
||||
<div className="vv-voice-alt-row">
|
||||
<button
|
||||
className={`vv-audio-cell ${task.refAudio?.source === "preset" ? "selected" : ""}`}
|
||||
className={`vv-voice-alt-btn ${task.refAudio?.source === "preset" ? "selected" : ""}`}
|
||||
onClick={() => setVoicePickerOpen(true)}
|
||||
>
|
||||
<FileTextOutlined className="vv-audio-ic" />
|
||||
<span>选择内置音色</span>
|
||||
<FileTextOutlined /> 内置音色
|
||||
</button>
|
||||
<button
|
||||
className={`vv-audio-cell ${task.voicePanel === "upload" ? "active" : ""} ${task.refAudio?.source === "upload" ? "selected" : ""}`}
|
||||
onClick={() => voiceInputRef.current?.click()}
|
||||
className={`vv-voice-alt-btn ${task.refAudio?.source === "upload" || task.refAudio?.source === "library" ? "selected" : ""}`}
|
||||
onClick={() =>
|
||||
setTask({ voicePanel: task.voicePanel === "library" ? null : "library" })
|
||||
}
|
||||
>
|
||||
<UploadOutlined className="vv-audio-ic" />
|
||||
<span>本地上传</span>
|
||||
</button>
|
||||
<button
|
||||
className={`vv-audio-cell ${task.voicePanel === "library" ? "active" : ""} ${task.refAudio?.source === "library" ? "selected" : ""}`}
|
||||
onClick={() => setAssetPicker({ open: true, kind: "voice", multiple: false })}
|
||||
>
|
||||
<FolderOpenOutlined className="vv-audio-ic" />
|
||||
<span>从素材库选择</span>
|
||||
</button>
|
||||
<button
|
||||
className={`vv-audio-cell ${task.voicePanel === "record" ? "active" : ""} ${task.refAudio?.source === "clone" ? "selected" : ""}`}
|
||||
onClick={() => setCloneModalOpen(true)}
|
||||
>
|
||||
<AudioMutedOutlined className="vv-audio-ic" />
|
||||
<span>直接录音</span>
|
||||
<FolderOpenOutlined /> 上传配音
|
||||
</button>
|
||||
</div>
|
||||
|
||||
<input
|
||||
ref={voiceInputRef}
|
||||
type="file"
|
||||
@@ -2083,7 +2202,14 @@ const ViralVideoPage: React.FC = () => {
|
||||
{task.refAudio && (
|
||||
<div className="vv-selected-audio">
|
||||
<SoundOutlined style={{ color: "#7c3aed", marginRight: 6 }} />
|
||||
<span style={{ flex: 1 }}>已选:{task.refAudio.name}</span>
|
||||
<span style={{ flex: 1 }}>
|
||||
已选:
|
||||
{task.refAudio.source === "my_voice"
|
||||
? `我的音色 · ${task.refAudio.name}`
|
||||
: task.refAudio.source === "preset"
|
||||
? `内置 · ${task.refAudio.name}`
|
||||
: task.refAudio.name}
|
||||
</span>
|
||||
<button className="vv-link-btn" onClick={() => setTask({ refAudio: null })}>
|
||||
移除
|
||||
</button>
|
||||
@@ -2091,16 +2217,32 @@ const ViralVideoPage: React.FC = () => {
|
||||
)}
|
||||
{task.voicePanel === "library" && (
|
||||
<div className="vv-voice-panel">
|
||||
<div className="vv-voice-panel-actions">
|
||||
<button
|
||||
className="vv-link-btn vv-link-btn-sm"
|
||||
onClick={() => voiceInputRef.current?.click()}
|
||||
>
|
||||
<UploadOutlined /> 本地上传音频文件
|
||||
</button>
|
||||
<button
|
||||
className="vv-link-btn vv-link-btn-sm"
|
||||
onClick={() =>
|
||||
setAssetPicker({ open: true, kind: "voice", multiple: false })
|
||||
}
|
||||
>
|
||||
<FolderOpenOutlined /> 从素材库选择
|
||||
</button>
|
||||
</div>
|
||||
{voiceAssets.length === 0 ? (
|
||||
<div className="vv-muted" style={{ padding: 12, textAlign: "center" }}>
|
||||
暂无已上传的音频
|
||||
暂无已上传的配音文件
|
||||
</div>
|
||||
) : (
|
||||
<div className="vv-voice-list">
|
||||
{voiceAssets.map((v) => (
|
||||
<div
|
||||
key={v.id}
|
||||
className={`vv-voice-item ${task.refAudio?.id === v.id && task.refAudio.source === "library" ? "selected" : ""}`}
|
||||
className={`vv-voice-item ${task.refAudio?.id === v.id && (task.refAudio.source === "library" || task.refAudio.source === "upload") ? "selected" : ""}`}
|
||||
onClick={() =>
|
||||
pickVoiceAsset({
|
||||
id: v.id,
|
||||
@@ -2289,7 +2431,7 @@ const ViralVideoPage: React.FC = () => {
|
||||
options={PURPOSES.map((i) => ({ value: i, label: i }))}
|
||||
/>
|
||||
</div>
|
||||
<div className="vv-form-row" style={{ gridColumn: "1 / -1" }}>
|
||||
<div className="vv-form-row">
|
||||
<label className="vv-label">文案视频时长</label>
|
||||
<Select
|
||||
className="vv-select"
|
||||
|
||||
@@ -387,6 +387,41 @@ BATCH_RENDER_SIMILARITY_LIMIT = 0.20
|
||||
"""批次内成片查重相似度阈值:超过则重选独立 plan 重渲一次(20%)。"""
|
||||
|
||||
|
||||
def _refund_smart_edit_prepaid(task_id: str) -> None:
|
||||
"""智能剪辑任务最终失败时退还预扣积分(幂等)。"""
|
||||
session = SessionLocal()
|
||||
try:
|
||||
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
|
||||
SQLAlchemyGenerationTaskRepository,
|
||||
)
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
repo = SQLAlchemyGenerationTaskRepository(session)
|
||||
task = repo.get(task_id)
|
||||
if not task:
|
||||
return
|
||||
prepaid = float(getattr(task, "credits_prepaid", 0) or 0)
|
||||
if prepaid <= 0:
|
||||
return
|
||||
txn_id = getattr(task, "credits_transaction_id", "") or ""
|
||||
res = PointsService().refund_points(
|
||||
user_id=task.user_id,
|
||||
amount=prepaid,
|
||||
source="smart_edit",
|
||||
db=session,
|
||||
ref_id=task.id,
|
||||
related_transaction_id=txn_id or None,
|
||||
description="智能剪辑任务失败退回",
|
||||
)
|
||||
task.credits_cost = 0.0
|
||||
task.credits_prepaid = 0.0
|
||||
repo.update(task)
|
||||
if not res.get("success"):
|
||||
logger.warning("[task_id=%s] 失败退积分未成功: %s", task_id, res)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def should_rerender_for_batch_dedup(*, batch_id: str, render_attempt: int, batch_similarity) -> bool:
|
||||
"""批次内查重后判定是否需要重选 plan 重渲。
|
||||
|
||||
@@ -1167,6 +1202,10 @@ def generate_video(self, task_id: str) -> dict:
|
||||
"mark_failed",
|
||||
error_message="source_edit_plan_id is required. Please create a preview task first.",
|
||||
)
|
||||
try:
|
||||
_refund_smart_edit_prepaid(task_id)
|
||||
except Exception:
|
||||
logger.warning("[task_id=%s] 失败退积分异常", task_id, exc_info=True)
|
||||
return {
|
||||
"status": "failed",
|
||||
"task_id": task_id,
|
||||
@@ -1205,6 +1244,7 @@ def generate_video(self, task_id: str) -> dict:
|
||||
)
|
||||
|
||||
# ── 自动重试逻辑 ──────────────────────────────────────────────────
|
||||
will_retry = False
|
||||
try:
|
||||
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
|
||||
SQLAlchemyGenerationTaskRepository,
|
||||
@@ -1217,6 +1257,7 @@ def generate_video(self, task_id: str) -> dict:
|
||||
if _task and _task.auto_retry_enabled and _task.auto_retry_max > 0:
|
||||
current_retry = _task.retry_count or 0
|
||||
if current_retry < _task.auto_retry_max:
|
||||
will_retry = True
|
||||
logger.info(
|
||||
"[task_id=%s] 触发自动重试: 当前重试次数=%d, 最大重试次数=%d",
|
||||
task_id,
|
||||
@@ -1250,6 +1291,13 @@ def generate_video(self, task_id: str) -> dict:
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
# 最终失败(不再重试):退还 smart_edit 预扣积分
|
||||
if not will_retry:
|
||||
try:
|
||||
_refund_smart_edit_prepaid(task_id)
|
||||
except Exception:
|
||||
logger.warning("[task_id=%s] 失败退积分异常", task_id, exc_info=True)
|
||||
|
||||
return {
|
||||
"status": "failed",
|
||||
"task_id": task_id,
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
"""爆款视频 Celery 编排器 — ViralVideoOrchestrator (v1.6 单次 Seedance 出片版).
|
||||
"""爆款视频 Celery 编排器 — ViralVideoOrchestrator.
|
||||
|
||||
v1.6 重大简化(Seedance 2.5 单次最长 30 秒,直接出片):
|
||||
1. _step_image_analysis 图片 VLM 分析(保留)
|
||||
V2 图片分析(10-05):火山OCR专用API + doubao-lite强约束JSON并行,单图<3s,8图<15s;pro VLM单次兜底。输出字段兼容旧格式,下游信任链/t2i零改动。
|
||||
|
||||
流水线步骤:
|
||||
1. _step_image_analysis 图片分析(V2: OCR+qwen3.8-flash并行 + qwen3.7-plus兜底)
|
||||
1.5 _step_video_analysis 参考视频风格分析(可选)
|
||||
2. _step_intent_parsing 用户文案意图解析
|
||||
3. _step_script_generation 编导分镜脚本生成(融合原 copy_fusion+storyboard+review,输出 copy_result 结构 + voiceover_script)
|
||||
@@ -22,11 +24,9 @@ from __future__ import annotations
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import tempfile
|
||||
import threading
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
@@ -330,269 +330,62 @@ def _vision_fallback(idx: int, reason: str, extra: dict | None = None) -> dict:
|
||||
return d
|
||||
|
||||
|
||||
def _is_vision_result_usable(result: dict) -> bool:
|
||||
"""判断 VLM 返回是否有效:name/summary 不能为未识别/无法判断/空,summary 要够长。"""
|
||||
if not isinstance(result, dict):
|
||||
return False
|
||||
name = (result.get("name") or "").strip()
|
||||
if not name or name in ("未识别", "无法判断", "未知"):
|
||||
return False
|
||||
summary = (result.get("summary") or "").strip()
|
||||
if len(summary) < 30 or summary in ("无法判断", "未识别"):
|
||||
return False
|
||||
category = (result.get("category") or "").strip()
|
||||
if category == "非产品图":
|
||||
return True
|
||||
feats = result.get("key_features") or []
|
||||
if not isinstance(feats, list) or len(feats) == 0:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _analyze_single_image(
|
||||
idx: int,
|
||||
img_url: str,
|
||||
vision_model: str,
|
||||
timeout: int,
|
||||
*,
|
||||
pro_fallback_model: str | None = None,
|
||||
) -> dict:
|
||||
"""单张图片 VLM 分析(#2040:改为从 prompt_loader 读模板 + XML 解析)。
|
||||
|
||||
lite 失败/不可用时用 pro 降级重试 1 次。失败/None 最终返回含默认字段的 dict。
|
||||
def _normalize_image_url(raw: str, idx: int) -> str:
|
||||
"""将 job.images 中的 storage_key/相对路径/空值统一归一化为可公网访问 URL。
|
||||
- http(s):// → 直接用
|
||||
- 其他 → storage_key,通过 SharedStorageService.get_url() 转公网 URL
|
||||
- 空值/非字符串 → 抛 ValueError
|
||||
"""
|
||||
if not raw or not isinstance(raw, str):
|
||||
raise ValueError(f"图片 #{idx} URL 为空或类型错误: {type(raw).__name__}={raw!r}")
|
||||
url = raw.strip()
|
||||
if not url:
|
||||
raise ValueError(f"图片 #{idx} URL 为空白字符串")
|
||||
if url.startswith("http://") or url.startswith("https://"):
|
||||
return url
|
||||
storage_key = url.lstrip("/")
|
||||
try:
|
||||
from packages.application.viral_video import xml_parser as xp
|
||||
from packages.application.viral_video.prompt_loader import (
|
||||
get_template,
|
||||
render_system_prompt,
|
||||
render_user_prompt,
|
||||
)
|
||||
from packages.shared.ai_service import call_vision
|
||||
except ImportError as e:
|
||||
logger.warning("[爆款视频] prompt 模板/解析模块不可用: %s", e)
|
||||
return _vision_fallback(idx, f"fallback_import_error:{e}")
|
||||
from packages.shared.storage import get_storage_service
|
||||
|
||||
if not img_url or not isinstance(img_url, str):
|
||||
return _vision_fallback(idx, "invalid_url")
|
||||
|
||||
template = get_template("image_analysis")
|
||||
system = render_system_prompt(template)
|
||||
user = render_user_prompt(
|
||||
template,
|
||||
image_count=1,
|
||||
industry="通用",
|
||||
image_urls=f"第1张:{img_url}",
|
||||
)
|
||||
|
||||
def _call(model: str, tmo: int):
|
||||
try:
|
||||
return call_vision(
|
||||
image_url=img_url,
|
||||
prompt=user,
|
||||
model=model,
|
||||
max_tokens=2048,
|
||||
temperature=0.3,
|
||||
timeout=tmo,
|
||||
system_prompt=system,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 图片 #%d call_vision(%s) 异常 err=%s", idx, model, e)
|
||||
return None
|
||||
|
||||
def _xml_to_product(nodes: list, raw_text: str) -> dict:
|
||||
product_nodes = [n for n in nodes if n["tag"] == "product"]
|
||||
scene = xp.text_of(raw_text, "scene") or "通用"
|
||||
mood = xp.text_of(raw_text, "mood") or ""
|
||||
# #2184: #2177 XML 重构后人物信息放在顶层 <people has_person count gender age_range pose expression/>,
|
||||
# 不再是 <product> 的 portrait_prompt 属性。需从顶层 people 标签提取并拼装 portrait_prompt。
|
||||
portrait_prompt = "无人像"
|
||||
try:
|
||||
people_node = xp.find_first(raw_text, "people")
|
||||
if people_node:
|
||||
_pa = people_node.get("attrs") or {}
|
||||
_has_person = xp.attr_bool(_pa.get("has_person"), False)
|
||||
if _has_person:
|
||||
_gender = _pa.get("gender", "无法判断") or "无法判断"
|
||||
_age = _pa.get("age_range", "无法判断") or "无法判断"
|
||||
_hair = _pa.get("hair", "无法判断") or "无法判断"
|
||||
_skin = _pa.get("skin_tone", "无法判断") or "无法判断"
|
||||
_face = _pa.get("face_shape", "无法判断") or "无法判断"
|
||||
_outfit = _pa.get("outfit", "无法判断") or "无法判断"
|
||||
_pose = _pa.get("pose", "无法判断") or "无法判断"
|
||||
_expr = _pa.get("expression", "无法判断") or "无法判断"
|
||||
_count = xp.attr_int(_pa.get("count"), 1)
|
||||
_parts = []
|
||||
if _gender != "无法判断":
|
||||
_g = _gender + ("性" if not _gender.endswith("性") else "")
|
||||
_parts.append(_g)
|
||||
if _age != "无法判断":
|
||||
_parts.append(_age)
|
||||
_parts.append("人物")
|
||||
if _hair != "无法判断":
|
||||
_parts.append(_hair)
|
||||
if _skin != "无法判断":
|
||||
_parts.append(f"{_skin}肤色")
|
||||
if _face != "无法判断":
|
||||
_parts.append(f"{_face}脸型")
|
||||
if _outfit != "无法判断":
|
||||
_parts.append(f"身着{_outfit}")
|
||||
if _pose != "无法判断":
|
||||
_parts.append(f"姿态{_pose}")
|
||||
if _expr != "无法判断":
|
||||
_parts.append(f"表情{_expr}")
|
||||
portrait_prompt = ",".join(_parts)
|
||||
logger.info(
|
||||
"[爆款视频] 图片 #%d 解析<people>: count=%d gender=%s age=%s hair=%s skin=%s face=%s outfit=%s pose=%s expr=%s → %s",
|
||||
idx,
|
||||
_count,
|
||||
_gender,
|
||||
_age,
|
||||
_hair,
|
||||
_skin,
|
||||
_face,
|
||||
_outfit,
|
||||
_pose,
|
||||
_expr,
|
||||
portrait_prompt,
|
||||
)
|
||||
except Exception as _pe:
|
||||
logger.warning("[爆款视频] 图片 #%d 解析<people>标签异常: %s,回退无人像", idx, _pe)
|
||||
for p in product_nodes:
|
||||
a = p["attrs"]
|
||||
text_on_pkg = a.get("text_on_package", "")
|
||||
p_body = p.get("text", "") or ""
|
||||
if not text_on_pkg and p_body:
|
||||
text_on_pkg = xp.text_of(p_body, "text_on_package") or ""
|
||||
text_list = [x.strip() for x in re.split(r"[,,;;]", text_on_pkg) if x.strip()] if text_on_pkg else []
|
||||
features = a.get("features", "")
|
||||
feat_list = [x.strip() for x in re.split(r"[,,;;]", features) if x.strip()] if features else []
|
||||
name = a.get("name", "") or "未识别"
|
||||
brand = a.get("brand", "") or "无法判断"
|
||||
category = a.get("category", "") or "无法判断"
|
||||
appearance = a.get("appearance", "") or "无法判断"
|
||||
packaging = a.get("packaging", "") or "无法判断"
|
||||
summary = a.get("summary", "") or f"{brand} {name}"
|
||||
# 优先取 product 属性上的 portrait_prompt(兼容旧schema),否则用顶层 <people> 解析结果
|
||||
_pp_from_attr = a.get("portrait_prompt", "")
|
||||
if _pp_from_attr and _pp_from_attr != "无人像":
|
||||
portrait_prompt = _pp_from_attr
|
||||
return {
|
||||
"name": name,
|
||||
"brand": brand,
|
||||
"category": category,
|
||||
"appearance": appearance,
|
||||
"packaging": packaging,
|
||||
"text_on_package": text_list,
|
||||
"key_features": feat_list or [features] if features else ["无法判断"],
|
||||
"scene": scene,
|
||||
"mood": mood,
|
||||
"portrait_prompt": portrait_prompt,
|
||||
"summary": summary,
|
||||
"_source": "xml",
|
||||
}
|
||||
# 没有 product 标签但有 <people has_person="true"> 也要能取到人物描述(兜底)
|
||||
if portrait_prompt != "无人像":
|
||||
return {
|
||||
"name": "未识别",
|
||||
"brand": "无法判断",
|
||||
"category": "无法判断",
|
||||
"appearance": "无法判断",
|
||||
"packaging": "无法判断",
|
||||
"text_on_package": [],
|
||||
"key_features": ["无法判断"],
|
||||
"scene": scene,
|
||||
"mood": mood,
|
||||
"portrait_prompt": portrait_prompt,
|
||||
"summary": "未识别",
|
||||
"_source": "xml_no_product",
|
||||
}
|
||||
return _vision_fallback(idx, "no_product_tag")
|
||||
|
||||
def _normalize(raw, source: str) -> dict:
|
||||
if raw is None:
|
||||
return _vision_fallback(idx, f"{source}_none")
|
||||
if not isinstance(raw, str):
|
||||
return _vision_fallback(idx, f"{source}_badtype")
|
||||
nodes = xp.parse_tags(raw)
|
||||
if not nodes:
|
||||
logger.warning("[爆款视频] 图片 #%d XML 解析失败 source=%s", idx, source)
|
||||
return _vision_fallback(idx, f"{source}_xml_fail", {"_raw": raw[:500]})
|
||||
product = _xml_to_product(nodes, raw)
|
||||
product.setdefault("_source", source)
|
||||
product["raw"] = raw[:500]
|
||||
return product
|
||||
|
||||
first_raw = _call(vision_model, timeout)
|
||||
tag1 = vision_model.split("/")[-1] if "/" in vision_model else vision_model
|
||||
first_result = _normalize(first_raw, tag1)
|
||||
if _is_vision_result_usable(first_result):
|
||||
return first_result
|
||||
|
||||
if pro_fallback_model and pro_fallback_model != vision_model:
|
||||
pro_raw = _call(pro_fallback_model, 60) # #2180: pro VLM 实测也需25-38s,原25s太短,提到60s
|
||||
pro_result = _normalize(pro_raw, "pro_fallback")
|
||||
if _is_vision_result_usable(pro_result):
|
||||
pro_result["_fallback_used"] = True
|
||||
return pro_result
|
||||
return pro_result
|
||||
return first_result
|
||||
url = get_storage_service().get_url(storage_key)
|
||||
except Exception as _e:
|
||||
raise ValueError(f"图片 #{idx} storage_key={storage_key!r} 转公网URL失败: {_e}") from _e
|
||||
logger.info("[爆款视频] 图片 #%d storage_key → 公网URL: %s", idx, url[:120])
|
||||
return url
|
||||
|
||||
|
||||
def _step_image_analysis(job: ViralVideoJob) -> dict:
|
||||
"""步骤 1: 图片 VLM 分析 — 识别产品特征(v1.6 优化:并行 + lite 模型提速)。"""
|
||||
try:
|
||||
from packages.shared.ai_service import call_vision # noqa: F401
|
||||
except ImportError:
|
||||
logger.warning("[爆款视频] ai_service.call_vision 不可用,使用占位结果")
|
||||
return {"products": [_vision_fallback(0, "fallback_import_error")]}
|
||||
"""步骤 1: 图片分析(V2 主路径)。
|
||||
|
||||
架构:
|
||||
- 主力:火山 MediaKit OCR(专用API,未配置时自动跳过)+ qwen3.8-flash 强约束 JSON,每图2路并行,目标<3s;
|
||||
- 外层全并发(workers=8),目标8图<15s;
|
||||
- 兜底:fast 结果不可用时单次调用 qwen3.7-plus(简单、无竞速)。
|
||||
- 唯一后端:阿里云百炼 DashScope,API Key 从环境变量 DASHSCOPE_API_KEY 读取。
|
||||
输出 dict 字段(name/brand/category/appearance/key_features/scene/mood/portrait_prompt/summary/_source)
|
||||
与旧版格式完全一致,下游信任链/t2i/intent_parsing/script_generation 零改动。
|
||||
"""
|
||||
if not job.images:
|
||||
logger.warning("[爆款视频] 任务无 images,跳过图片分析")
|
||||
return {"products": []}
|
||||
|
||||
# 选择视觉模型:lite 速度优先(默认),pro 作为降级备用
|
||||
# URL 归一化(storage_key→公网URL;空值直接400)
|
||||
normalized_urls: list[str] = []
|
||||
for idx, raw in enumerate(job.images):
|
||||
normalized_urls.append(_normalize_image_url(raw, idx))
|
||||
|
||||
try:
|
||||
_s = get_shared_settings()
|
||||
if _s.doubao_vision_use_lite:
|
||||
vision_model = _s.doubao_vision_lite_model
|
||||
pro_model = _s.doubao_vision_model
|
||||
vision_timeout = 45 # #2180: 方舟 VLM 实测服务端处理24-38s,原15s必超时3次重试全挂,提到45s
|
||||
else:
|
||||
vision_model = _s.doubao_vision_model
|
||||
pro_model = None # 已经是 pro,不再降级
|
||||
vision_timeout = 60
|
||||
except Exception:
|
||||
vision_model = "doubao-1-5-vision-lite-250315"
|
||||
pro_model = "doubao-1-5-vision-pro-250328"
|
||||
vision_timeout = 45
|
||||
from worker_app.tasks.vision import analyze_images_v2 as _aiv2
|
||||
except ImportError:
|
||||
try:
|
||||
from tasks.vision import analyze_images_v2 as _aiv2 # type: ignore
|
||||
except ImportError as e:
|
||||
logger.error("[爆款视频] vision 模块导入失败: %s", e)
|
||||
return {"products": [_vision_fallback(0, f"vision_import_error:{e}")]}
|
||||
|
||||
results: list[dict] = [None] * len(job.images) # type: ignore
|
||||
max_workers = min(4, max(1, len(job.images)))
|
||||
logger.info(
|
||||
"[爆款视频] 开始并行图片分析 n=%d model=%s pro_fallback=%s timeout=%d workers=%d",
|
||||
len(job.images),
|
||||
vision_model,
|
||||
pro_model,
|
||||
vision_timeout,
|
||||
max_workers,
|
||||
)
|
||||
with ThreadPoolExecutor(max_workers=max_workers) as pool:
|
||||
future_to_idx = {
|
||||
pool.submit(
|
||||
_analyze_single_image, idx, url, vision_model, vision_timeout, pro_fallback_model=pro_model
|
||||
): idx
|
||||
for idx, url in enumerate(job.images)
|
||||
}
|
||||
for fut in as_completed(future_to_idx):
|
||||
idx = future_to_idx[fut]
|
||||
try:
|
||||
results[idx] = fut.result()
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 图片 #%d future 异常 err=%s", idx, e, exc_info=True)
|
||||
results[idx] = _vision_fallback(idx, "future_exception", {"_error": str(e)[:200]})
|
||||
|
||||
return {"products": results}
|
||||
# V2 内部 httpx 直连 dashscope,单次调用无重试,无需调整全局 client
|
||||
results = _aiv2(normalized_urls)
|
||||
return {"products": list(results)}
|
||||
|
||||
|
||||
def _step_video_analysis(job: ViralVideoJob) -> dict | None:
|
||||
@@ -628,10 +421,14 @@ def _step_intent_parsing(job: ViralVideoJob, image_analysis: dict) -> dict:
|
||||
render_system_prompt,
|
||||
render_user_prompt,
|
||||
)
|
||||
from packages.shared.ai_service import call_llm
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
except ImportError:
|
||||
return {"intent": "推广产品", "key_messages": ["产品亮点"], "tone": "专业", "suggested_title": ""}
|
||||
|
||||
_llm_client = get_doubao_client()
|
||||
if not _llm_client.is_available:
|
||||
return {"intent": "推广产品", "key_messages": ["产品亮点"], "tone": "专业", "suggested_title": ""}
|
||||
|
||||
products_summary = ""
|
||||
products = (image_analysis or {}).get("products", []) or []
|
||||
for p in products:
|
||||
@@ -686,13 +483,13 @@ def _step_intent_parsing(job: ViralVideoJob, image_analysis: dict) -> dict:
|
||||
for _m, _lbl in [(_fast, "fast"), (_pro, "pro-fallback")]:
|
||||
try:
|
||||
logger.info("[爆款视频] 意图解析 model=%s label=%s", _m, _lbl)
|
||||
raw = call_llm(
|
||||
raw = _llm_client.chat_completion(
|
||||
[{"role": "system", "content": system}, {"role": "user", "content": user}],
|
||||
temperature=0.4,
|
||||
max_tokens=1024,
|
||||
model=_m,
|
||||
timeout=60,
|
||||
) # #2180: 意图解析 LLM 实测需更长响应,原25s太紧
|
||||
) # #2180/#2215: 直接用 client.chat_completion 传 messages list,不再走 call_llm 字符串包装
|
||||
if not raw:
|
||||
continue
|
||||
parsed = _parse(raw)
|
||||
@@ -1032,10 +829,14 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
|
||||
GLOBAL_CONSTRAINTS,
|
||||
NEGATIVE_RULES,
|
||||
)
|
||||
from packages.shared.ai_service import call_llm
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
except ImportError:
|
||||
return _fallback_script(job)
|
||||
|
||||
_llm_client2 = get_doubao_client()
|
||||
if not _llm_client2.is_available:
|
||||
return _fallback_script(job)
|
||||
|
||||
products_summary = _build_products_summary(image_analysis)
|
||||
dur = max(5, min(30, int(getattr(job, "duration", 15) or 15)))
|
||||
|
||||
@@ -1071,7 +872,7 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
|
||||
|
||||
def _try_gen(model: str, temp: float, max_tok: int, label: str, tmo: int = 25):
|
||||
logger.info("[爆款视频] 编导脚本生成 model=%s label=%s timeout=%d", model, label, tmo)
|
||||
raw = call_llm(
|
||||
raw = _llm_client2.chat_completion(
|
||||
[{"role": "system", "content": system_tpl}, {"role": "user", "content": user}],
|
||||
temperature=temp,
|
||||
max_tokens=max_tok,
|
||||
@@ -1109,18 +910,19 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
|
||||
_s = get_shared_settings()
|
||||
_fast = _s.doubao_fast_model
|
||||
_pro = getattr(_s, "doubao_model", None) or _fast
|
||||
_script_fast_tmo = int(os.environ.get("VIRAL_VIDEO_SCRIPT_FAST_TIMEOUT", "150"))
|
||||
_script_pro_tmo = int(os.environ.get("VIRAL_VIDEO_SCRIPT_PRO_TIMEOUT", "150"))
|
||||
try:
|
||||
# 第一次:快模型 25s
|
||||
normalized = _try_gen(_fast, 0.8, 2500, "fast-first", tmo=90)
|
||||
# #2217: doubao-seed-2-1-pro生成长编导脚本高峰期>90s,上调到150s,支持ENV覆盖
|
||||
normalized = _try_gen(_fast, 0.8, 2500, "fast-first", tmo=_script_fast_tmo)
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
# #2183: 实测pro 1500tok输出需75.8s,单次timeout提到90s
|
||||
normalized = _try_gen(_fast, 0.6, 3200, "fast-retry", tmo=90)
|
||||
normalized = _try_gen(_fast, 0.6, 3200, "fast-retry", tmo=_script_fast_tmo)
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
# 第三次:用主力模型兜底,给 120s
|
||||
# 第三次:用主力模型兜底
|
||||
if _pro and _pro != _fast:
|
||||
normalized = _try_gen(_pro, 0.7, 3500, "pro-fallback", tmo=120)
|
||||
normalized = _try_gen(_pro, 0.7, 3500, "pro-fallback", tmo=_script_pro_tmo)
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
logger.warning("[爆款视频] 编导脚本三次都未生成合格结果,使用兜底脚本")
|
||||
@@ -1221,15 +1023,79 @@ def _step_review(job: ViralVideoJob, copy_result: dict) -> dict:
|
||||
return {"passed": True, "score": 75, "details": {}, "issues": []}
|
||||
|
||||
|
||||
def _resolve_tts_voice_id(job: ViralVideoJob) -> str:
|
||||
"""#2188: 根据 voice_source + voice_id 解析真正传给 CosyVoice 的 voice_id。
|
||||
|
||||
- voice_source 在 ("my_voice", "clone"):voice_id 是 VoiceCloneProfile.id,
|
||||
需要从 DB 查 profile.voice_id(CosyVoice 返回的音色 ID)。
|
||||
- 其他/空:voice_id 直接视为 CosyVoice preset 音色名(longxiaochun_v3 等)。
|
||||
- 任何解析失败都回退默认 longxiaochun_v3,保证任务不崩。
|
||||
"""
|
||||
default_voice = "longxiaochun_v3"
|
||||
raw_voice_id = (getattr(job, "voice_id", "") or "").strip()
|
||||
voice_source = (getattr(job, "voice_source", "") or "").strip().lower()
|
||||
|
||||
if not raw_voice_id:
|
||||
return default_voice
|
||||
|
||||
# 克隆音色:前端传 profile.id,需查 DB 取 cosyvoice_voice_id
|
||||
if voice_source in ("my_voice", "clone"):
|
||||
session = None
|
||||
try:
|
||||
from packages.adapters.sqlalchemy_impl.voice_clone_profile_repository import (
|
||||
SQLAlchemyVoiceCloneProfileRepository,
|
||||
)
|
||||
|
||||
session = SessionLocal()
|
||||
_repo = SQLAlchemyVoiceCloneProfileRepository(session)
|
||||
_profile = _repo.get(raw_voice_id)
|
||||
if _profile and _profile.status.value == "ready" and (_profile.voice_id or "").strip():
|
||||
logger.info(
|
||||
"[爆款视频] 克隆音色解析: profile_id=%s → cosyvoice_voice_id=%s",
|
||||
raw_voice_id,
|
||||
_profile.voice_id,
|
||||
)
|
||||
return _profile.voice_id.strip()
|
||||
# 找不到/未就绪/无voice_id
|
||||
if _profile is None:
|
||||
logger.warning("[爆款视频] 克隆音色 profile_id=%s 不存在,回退默认音色", raw_voice_id)
|
||||
elif _profile.status.value != "ready":
|
||||
logger.warning(
|
||||
"[爆款视频] 克隆音色 profile_id=%s 状态=%s 未就绪,回退默认音色",
|
||||
raw_voice_id,
|
||||
_profile.status.value,
|
||||
)
|
||||
else:
|
||||
logger.warning("[爆款视频] 克隆音色 profile_id=%s 就绪但 voice_id 为空,回退默认音色", raw_voice_id)
|
||||
return default_voice
|
||||
except Exception as _e:
|
||||
logger.warning("[爆款视频] 克隆音色解析异常 profile_id=%s err=%s,回退默认音色", raw_voice_id, _e)
|
||||
return default_voice
|
||||
finally:
|
||||
if session is not None:
|
||||
try:
|
||||
session.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# preset/my-voice 直传/空 source:voice_id 就是 CosyVoice 音色名
|
||||
return raw_voice_id or default_voice
|
||||
|
||||
|
||||
def _step_tts(job: ViralVideoJob, voiceover_script: str):
|
||||
"""步骤 5: CosyVoice 整段配音 → 返回本地 MP3 Path;失败返回 None。"""
|
||||
"""步骤 5: CosyVoice 整段配音 → 返回本地 MP3 Path;失败返回 None。
|
||||
|
||||
#2188: 支持 voice_source='my_voice'/'clone',先把前端传的 profile.id
|
||||
解析成 CosyVoice 真正的克隆音色 voice_id,再走同一条 synthesize 链路。
|
||||
"""
|
||||
try:
|
||||
from pathlib import Path as _Path
|
||||
|
||||
from apps.worker.services.tts_service_factory import get_tts_service
|
||||
|
||||
tts_service = get_tts_service()
|
||||
voice_id = (getattr(job, "voice_id", "") or "").strip()
|
||||
# #2188: 根据 voice_source 解析实际传给 CosyVoice 的 voice_id
|
||||
voice_id = _resolve_tts_voice_id(job)
|
||||
text = (voiceover_script or "").strip()
|
||||
if not text:
|
||||
logger.warning("[爆款视频] voiceover_script 为空,跳过 TTS")
|
||||
@@ -1237,21 +1103,19 @@ def _step_tts(job: ViralVideoJob, voiceover_script: str):
|
||||
try:
|
||||
result = tts_service.synthesize(
|
||||
text=text,
|
||||
voice_id=voice_id or "longxiaochun_v3",
|
||||
voice_id=voice_id,
|
||||
format="mp3",
|
||||
)
|
||||
except TypeError:
|
||||
try:
|
||||
result = tts_service.synthesize(text=text, voice_id=voice_id or "longxiaochun_v3")
|
||||
result = tts_service.synthesize(text=text, voice_id=voice_id)
|
||||
except TypeError:
|
||||
result = tts_service.synthesize(text=text)
|
||||
if result is None:
|
||||
return None
|
||||
p = _Path(result) if not isinstance(result, _Path) else result
|
||||
if p.exists() and p.stat().st_size > 0:
|
||||
logger.info(
|
||||
"[爆款视频] TTS 合成完成: voice=%s path=%s size=%d", voice_id or "longxiaochun_v3", p, p.stat().st_size
|
||||
)
|
||||
logger.info("[爆款视频] TTS 合成完成: voice=%s path=%s size=%d", voice_id, p, p.stat().st_size)
|
||||
return p
|
||||
logger.warning("[爆款视频] TTS 返回路径不存在或空文件: %s", p)
|
||||
return None
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""V2 图片分析:火山OCR专用API + doubao-lite强约束JSON并行,单次pro VLM兜底。"""
|
||||
|
||||
from .fast_path import analyze_image_v2, analyze_images_v2 # noqa: F401
|
||||
@@ -0,0 +1,208 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""V2 prompt 解析:优先读后台 viral_video_prompt_templates 表(prompt_type='image_analysis'
|
||||
且 is_active=true),30s TTL 热加载;DB 无有效记录/异常时,fallback 到纯硬编码 JSON schema prompt。
|
||||
|
||||
规则(简单直接,不做字符串匹配判断):
|
||||
- DB 有 is_active=true 的 image_analysis 记录(含种子版本和用户修改后的版本):
|
||||
* system = DB.system_prompt(DB prompt 自带完整输出格式,不追加硬编码 schema,
|
||||
避免 DB 写 XML、调用强制 json_object 造成的格式冲突)
|
||||
* user = DB.user_prompt_template 渲染后使用;渲染后为空则用硬编码默认
|
||||
- DB 无记录/连接异常/返回空:system/user 全部用纯硬编码 JSON schema prompt
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---- 纯硬编码 JSON schema(DB 无有效配置时全量使用) ----
|
||||
|
||||
_FAST_JSON_SCHEMA = (
|
||||
"你是图片结构化识别器。严格按下方 JSON schema 返回一个对象,不要任何解释、"
|
||||
"不要markdown、不要代码块、不要前后缀文字。字段值不确定时填 null 或空数组。\n"
|
||||
"{\n"
|
||||
' "has_person": true/false,\n'
|
||||
' "gender": "男"/"女"/null,\n'
|
||||
' "age_range": "儿童"/"青少年"/"青年"/"中年"/"老年"/null,\n'
|
||||
' "upper_wear": "上装款式,如T恤/衬衫/卫衣/毛衣/西装/夹克/连衣裙/吊带/背心/外套等",\n'
|
||||
' "upper_color": "上装主色",\n'
|
||||
' "lower_wear": "下装款式;穿连衣裙时填null",\n'
|
||||
' "lower_color": "下装主色",\n'
|
||||
' "dress_color": "连衣裙主色(穿连衣裙时填)",\n'
|
||||
' "accessories": ["眼镜"/"帽子"/"项链"/"耳环"/"背包"/"手表"等数组],\n'
|
||||
' "hairstyle": "发型,如短发/长发/马尾/卷发/丸子头/光头等",\n'
|
||||
' "expression": "表情,如微笑/严肃/酷/开心等",\n'
|
||||
' "pose": "姿势,如站立/坐姿/侧身/行走等",\n'
|
||||
' "scene": "场景,如室内/街拍/户外/办公室/家居/海边/雪景/森林等",\n'
|
||||
' "style": "风格,如休闲/商务/运动/复古/潮流/甜美/酷飒/优雅/街头/法式等",\n'
|
||||
' "has_product": true/false,\n'
|
||||
' "category": "产品类目:服饰/鞋包/美妆/数码/食品/家居/配饰/母婴/非产品图",\n'
|
||||
' "product_name": "产品名称,非产品图填null",\n'
|
||||
' "brand": "品牌或文字标识,无则null",\n'
|
||||
' "material": "材质,如棉质/牛仔/皮革/真丝/针织/涤纶等",\n'
|
||||
' "pattern": "图案,如纯色/条纹/波点/格子/印花/碎花/Logo等",\n'
|
||||
' "colors": ["主色数组"],\n'
|
||||
' "mood": "整体氛围/情绪,如清新/活力/高级/温暖/冷峻/甜美/复古等"\n'
|
||||
"}\n\n"
|
||||
"你必须只返回一个合法的JSON对象,不要输出任何其他文字、解释、XML标签或markdown。"
|
||||
)
|
||||
DEFAULT_FAST_USER = "识别这张图片的人物穿搭与主体信息,只返回JSON对象。"
|
||||
|
||||
_PRO_JSON_SCHEMA = (
|
||||
"你是图片分析专家。严格按下方 JSON schema 返回一个对象,不要解释、不要markdown、不要代码块、不要XML标签。\n"
|
||||
"{\n"
|
||||
' "has_person": true/false,\n'
|
||||
' "gender": "男"/"女"/null,\n'
|
||||
' "age_range": "儿童"/"青少年"/"青年"/"中年"/"老年"/null,\n'
|
||||
' "outfit": "整体穿着描述(含颜色款式)",\n'
|
||||
' "hair": "发型发色",\n'
|
||||
' "pose": "姿势",\n'
|
||||
' "expression": "表情",\n'
|
||||
' "scene": "场景",\n'
|
||||
' "mood": "氛围",\n'
|
||||
' "has_product": true/false,\n'
|
||||
' "category": "类目:服饰/鞋包/美妆/数码/食品/家居/配饰/母婴/非产品图",\n'
|
||||
' "product_name": "产品名,非产品图填null",\n'
|
||||
' "brand": "品牌,无则null",\n'
|
||||
' "key_features": ["核心特征数组,3-6个短语"]\n'
|
||||
"}\n\n"
|
||||
"你必须只返回一个合法的JSON对象,不要输出任何其他文字、解释、XML标签或markdown。"
|
||||
)
|
||||
DEFAULT_PRO_USER = "分析这张图片,返回符合schema的JSON。"
|
||||
|
||||
# 保留旧 JSON schema 追加文本作为常量(DB prompt 完全控制输出格式后不再使用,
|
||||
# 保留以便排查历史行为)。
|
||||
_FAST_JSON_APPEND = (
|
||||
"\n\n【输出格式要求】无论上文如何要求,最终你必须只返回一个合法的JSON对象,"
|
||||
"严格包含以下字段(字段值不确定时填null或空数组):\n"
|
||||
"{\n"
|
||||
' "has_person": true/false,\n'
|
||||
' "gender": "男"/"女"/null,\n'
|
||||
' "age_range": "儿童"/"青少年"/"青年"/"中年"/"老年"/null,\n'
|
||||
' "upper_wear": "上装款式字符串",\n'
|
||||
' "upper_color": "上装主色",\n'
|
||||
' "lower_wear": "下装款式(穿连衣裙时填null)",\n'
|
||||
' "lower_color": "下装主色",\n'
|
||||
' "dress_color": "连衣裙主色(穿连衣裙时填)",\n'
|
||||
' "accessories": ["配饰数组"],\n'
|
||||
' "hairstyle": "发型",\n'
|
||||
' "expression": "表情",\n'
|
||||
' "pose": "姿势",\n'
|
||||
' "scene": "场景",\n'
|
||||
' "style": "风格",\n'
|
||||
' "has_product": true/false,\n'
|
||||
' "category": "产品类目:服饰/鞋包/美妆/数码/食品/家居/配饰/母婴/非产品图",\n'
|
||||
' "product_name": "产品名称,非产品图填null",\n'
|
||||
' "brand": "品牌或文字标识,无则null",\n'
|
||||
' "material": "材质",\n'
|
||||
' "pattern": "图案",\n'
|
||||
' "colors": ["主色数组"],\n'
|
||||
' "mood": "整体氛围"\n'
|
||||
"}\n"
|
||||
"不要输出任何其他文字、解释、XML标签或markdown。"
|
||||
)
|
||||
|
||||
_PRO_JSON_APPEND = (
|
||||
"\n\n【输出格式要求】无论上文如何要求,最终你必须只返回一个合法的JSON对象,"
|
||||
"严格包含以下字段(字段值不确定时填null或空数组):\n"
|
||||
"{\n"
|
||||
' "has_person": true/false,\n'
|
||||
' "gender": "男"/"女"/null,\n'
|
||||
' "age_range": "儿童"/"青少年"/"青年"/"中年"/"老年"/null,\n'
|
||||
' "outfit": "整体穿着描述(含颜色款式)",\n'
|
||||
' "hair": "发型发色",\n'
|
||||
' "pose": "姿势",\n'
|
||||
' "expression": "表情",\n'
|
||||
' "scene": "场景",\n'
|
||||
' "mood": "氛围",\n'
|
||||
' "has_product": true/false,\n'
|
||||
' "category": "类目:服饰/鞋包/美妆/数码/食品/家居/配饰/母婴/非产品图",\n'
|
||||
' "product_name": "产品名,非产品图填null",\n'
|
||||
' "brand": "品牌,无则null",\n'
|
||||
' "key_features": ["核心特征3-6个短语"]\n'
|
||||
"}\n"
|
||||
"不要输出任何其他文字、解释、XML标签或markdown。"
|
||||
)
|
||||
|
||||
_cache_lock = threading.Lock()
|
||||
_cache: dict[str, tuple[float, Any]] = {}
|
||||
_CACHE_TTL = 30.0
|
||||
|
||||
|
||||
def _load_db_template() -> Any | None:
|
||||
"""直接查DB viral_video_prompt_templates 中 is_active=true 的 image_analysis 记录;
|
||||
DB不可达/无记录/异常返回None。
|
||||
复用 prompt_loader._load_from_db,它只查DB不做DEFAULT_TEMPLATES fallback,
|
||||
返回None表示DB无记录或异常。"""
|
||||
try:
|
||||
from packages.application.viral_video.prompt_loader import _load_from_db
|
||||
|
||||
return _load_from_db("image_analysis")
|
||||
except Exception as e:
|
||||
logger.warning("[vision.v2] 查询DB prompt配置失败: %s", e)
|
||||
return None
|
||||
|
||||
|
||||
def _render_user(tpl: Any | None, default_user: str) -> str:
|
||||
if not tpl:
|
||||
return default_user
|
||||
tpl_str = getattr(tpl, "user_prompt_template", "") or ""
|
||||
if not tpl_str.strip():
|
||||
return default_user
|
||||
rendered = tpl_str.replace("{image_count}", "1").replace("{industry}", "通用").replace("{image_urls}", "").strip()
|
||||
return rendered or default_user
|
||||
|
||||
|
||||
def resolve_fast_prompt() -> tuple[str, str]:
|
||||
return _resolve("fast")
|
||||
|
||||
|
||||
def resolve_pro_prompt() -> tuple[str, str]:
|
||||
return _resolve("pro")
|
||||
|
||||
|
||||
def _resolve(kind: str) -> tuple[str, str]:
|
||||
now = time.time()
|
||||
cache_key = f"prompt_{kind}"
|
||||
with _cache_lock:
|
||||
hit = _cache.get(cache_key)
|
||||
if hit and now - hit[0] < _CACHE_TTL:
|
||||
return hit[1]
|
||||
|
||||
default_sys = _FAST_JSON_SCHEMA if kind == "fast" else _PRO_JSON_SCHEMA
|
||||
default_user = DEFAULT_FAST_USER if kind == "fast" else DEFAULT_PRO_USER
|
||||
|
||||
sys_prompt = default_sys
|
||||
usr_prompt = default_user
|
||||
try:
|
||||
tpl = _load_db_template()
|
||||
if tpl is not None:
|
||||
db_sys = (getattr(tpl, "system_prompt", "") or "").strip()
|
||||
if db_sys:
|
||||
sys_prompt = db_sys # DB prompt自带完整输出格式,不追加硬编码schema避免冲突
|
||||
usr_prompt = _render_user(tpl, default_user)
|
||||
logger.info(
|
||||
"[vision.v2] 使用DB image_analysis prompt (kind=%s version=%s sys_len=%d)",
|
||||
kind,
|
||||
getattr(tpl, "version", "?"),
|
||||
len(db_sys),
|
||||
)
|
||||
else:
|
||||
logger.debug("[vision.v2] DB image_analysis system_prompt为空,使用默认JSON (kind=%s)", kind)
|
||||
else:
|
||||
logger.debug("[vision.v2] DB无image_analysis记录/不可达,使用默认JSON prompt (kind=%s)", kind)
|
||||
except Exception as e:
|
||||
logger.warning("[vision.v2] 解析DB prompt异常,使用默认: %s", e)
|
||||
|
||||
with _cache_lock:
|
||||
_cache[cache_key] = (now, (sys_prompt, usr_prompt))
|
||||
return sys_prompt, usr_prompt
|
||||
|
||||
|
||||
def invalidate_cache() -> None:
|
||||
with _cache_lock:
|
||||
_cache.clear()
|
||||
@@ -0,0 +1,646 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""把 fast_json VLM 输出 + OCR 文本组装为下游兼容的 product dict。
|
||||
|
||||
v4 schema: DB prompt完全控制输出格式,可能是v4嵌套schema(type/products/people/store_info)
|
||||
或旧扁平schema(has_person/upper_wear/product_name/brand等)。assembler兼容两种格式。
|
||||
|
||||
目标:下游(信任链t2i/intent_parsing/script_generation)零改动。
|
||||
必出字段:name, brand, category, appearance, packaging, text_on_package,
|
||||
key_features, scene, mood, portrait_prompt, summary, _source
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
|
||||
def _join_parts(*parts: str | None) -> str:
|
||||
return "".join(p for p in parts if p)
|
||||
|
||||
|
||||
_AGE_PREFIX = {"青年": "年轻", "中年": "中年", "老年": "老年"}
|
||||
_GENDER_WORD = {"男": "男性", "女": "女性"}
|
||||
|
||||
|
||||
def _person_subject(gender: str, age: str) -> str:
|
||||
gw = _GENDER_WORD.get(gender, "")
|
||||
if age == "儿童":
|
||||
if gender == "女":
|
||||
return "小女孩"
|
||||
if gender == "男":
|
||||
return "小男孩"
|
||||
return "儿童"
|
||||
if age == "青少年":
|
||||
if gender == "女":
|
||||
return "少女"
|
||||
if gender == "男":
|
||||
return "少年"
|
||||
return "青少年"
|
||||
prefix = _AGE_PREFIX.get(age, "")
|
||||
if gw:
|
||||
return f"{prefix}{gw}" if prefix else gw
|
||||
return f"{prefix}人物" if prefix else "人物"
|
||||
|
||||
|
||||
def _build_wear_from_v4(p: dict) -> str:
|
||||
"""v4 person schema: upper_wear/upper_color/lower_wear/lower_color/dress_color"""
|
||||
upper = p.get("upper_wear") or ""
|
||||
upper_color = p.get("upper_color") or ""
|
||||
lower = p.get("lower_wear") or ""
|
||||
lower_color = p.get("lower_color") or ""
|
||||
dress_color = p.get("dress_color") or ""
|
||||
is_dress = ("连衣裙" in upper) or ("裙" in upper and not lower)
|
||||
if is_dress:
|
||||
c = dress_color or upper_color
|
||||
return f"身穿{c}{upper}" if c else f"身穿{upper}"
|
||||
parts = []
|
||||
if upper:
|
||||
up = f"{upper_color}{upper}" if upper_color else upper
|
||||
parts.append(f"上身{up}")
|
||||
if lower:
|
||||
lo = f"{lower_color}{lower}" if lower_color else lower
|
||||
parts.append(f"下身{lo}")
|
||||
return ",".join(parts)
|
||||
|
||||
|
||||
def _build_portrait_prompt_from_v4(p: dict) -> str:
|
||||
"""v4 person: 直接用portrait_prompt字段;没有就拼"""
|
||||
direct = p.get("portrait_prompt")
|
||||
if direct and len(direct) >= 10:
|
||||
return direct
|
||||
subject = _person_subject(p.get("gender", ""), p.get("age_range", ""))
|
||||
wear = _build_wear_from_v4(p)
|
||||
acc = p.get("accessories") or []
|
||||
if isinstance(acc, str):
|
||||
acc = [acc]
|
||||
acc_str = ",佩戴" + "、".join(str(a) for a in acc if a) if acc else ""
|
||||
hair = p.get("hairstyle") or ""
|
||||
expr = p.get("expression") or ""
|
||||
pose = p.get("pose") or ""
|
||||
style = p.get("outfit_style") or p.get("style") or ""
|
||||
scene = p.get("scene") or ""
|
||||
mood = p.get("mood") or ""
|
||||
details = []
|
||||
if hair:
|
||||
details.append(hair)
|
||||
if expr and expr not in ("自然", "平静"):
|
||||
details.append(f"神情{expr}")
|
||||
if pose and pose not in ("站立",):
|
||||
details.append(pose)
|
||||
style_parts = []
|
||||
if style:
|
||||
style_parts.append(style)
|
||||
if mood:
|
||||
style_parts.append(mood)
|
||||
if scene and scene not in ("通用",):
|
||||
style_parts.append(scene)
|
||||
pieces = [f"一位{subject}"]
|
||||
if wear:
|
||||
pieces.append(wear)
|
||||
if acc_str:
|
||||
pieces.append(acc_str.lstrip(","))
|
||||
if details:
|
||||
pieces.append(",".join(details))
|
||||
pieces.append(("".join(style_parts) + "风格") if style_parts else "人像写真")
|
||||
full = ",".join(p for p in pieces if p)
|
||||
if len(full) < 40:
|
||||
full += ",自然光线下人像特写,画面清晰"
|
||||
if len(full) > 120:
|
||||
full = full[:120].rstrip(",") + "。"
|
||||
return full
|
||||
|
||||
|
||||
def _build_product_prompt_from_v4(prod: dict, top: dict) -> str:
|
||||
"""v4 product: 拼商品视觉描述prompt(用于AI生图参考)"""
|
||||
name = prod.get("product_name") or "商品"
|
||||
brand = prod.get("brand") or ""
|
||||
lead = f"{brand} {name}" if brand and brand not in name else name
|
||||
pkg_color = prod.get("package_color") or ""
|
||||
pkg_type = prod.get("package_type") or ""
|
||||
cap = prod.get("cap_type") or ""
|
||||
body = prod.get("body_shape") or ""
|
||||
features = prod.get("product_features") or []
|
||||
sell = prod.get("key_selling_points") or []
|
||||
colors = top.get("colors") or []
|
||||
style = top.get("style") or ""
|
||||
scene = top.get("scene") or ""
|
||||
mood = top.get("mood") or ""
|
||||
|
||||
parts = [lead]
|
||||
desc = []
|
||||
if pkg_color:
|
||||
desc.append(pkg_color)
|
||||
if pkg_type:
|
||||
desc.append(pkg_type)
|
||||
if cap and len(desc) < 3:
|
||||
desc.append(f"配{cap}")
|
||||
if body and len(desc) < 3:
|
||||
desc.append(body)
|
||||
if desc:
|
||||
parts.append(",".join(desc))
|
||||
if features:
|
||||
core = [str(f) for f in features[:3] if f and len(str(f)) <= 25]
|
||||
if core:
|
||||
parts.append(";".join(core))
|
||||
if sell:
|
||||
s = [str(x) for x in sell[:2] if x]
|
||||
if s:
|
||||
parts.append("突出" + "、".join(s))
|
||||
cnames = []
|
||||
for cc in colors:
|
||||
if isinstance(cc, dict) and cc.get("name"):
|
||||
cnames.append(cc["name"])
|
||||
elif isinstance(cc, str):
|
||||
cnames.append(cc)
|
||||
cnames = cnames[:3]
|
||||
if cnames:
|
||||
parts.append("、".join(cnames) + "主色")
|
||||
if style:
|
||||
parts.append(style)
|
||||
if mood:
|
||||
parts.append(mood)
|
||||
if scene and not any(k in scene for k in ("白色背景", "纯色", "通用")):
|
||||
parts.append(scene)
|
||||
parts.append("产品特写,画面清晰")
|
||||
prompt = ",".join(p for p in parts if p)
|
||||
return prompt if len(prompt) >= 10 else "产品展示图,特写镜头"
|
||||
|
||||
|
||||
def _is_v4_schema(fj: dict) -> bool:
|
||||
"""判断是v4嵌套schema还是旧扁平schema"""
|
||||
return (
|
||||
isinstance(fj.get("products"), list)
|
||||
or fj.get("type") in ("product", "store", "person", "other")
|
||||
or isinstance(fj.get("people"), dict)
|
||||
)
|
||||
|
||||
|
||||
# ---------- 旧扁平schema兼容(保留原逻辑) ----------
|
||||
|
||||
|
||||
def _person_subject_old(fj: dict) -> str:
|
||||
return _person_subject(fj.get("gender", ""), fj.get("age_range", ""))
|
||||
|
||||
|
||||
def _build_wear_sentence_old(fj: dict) -> str:
|
||||
upper = fj.get("upper_wear") or ""
|
||||
upper_color = fj.get("upper_color") or ""
|
||||
lower = fj.get("lower_wear") or ""
|
||||
lower_color = fj.get("lower_color") or ""
|
||||
dress_color = fj.get("dress_color") or ""
|
||||
material = fj.get("material") or ""
|
||||
pattern = fj.get("pattern") or ""
|
||||
is_dress = ("连衣裙" in upper) or ("裙" in upper and not lower)
|
||||
if is_dress:
|
||||
c = dress_color or upper_color
|
||||
wear = f"{c}{upper}" if c else upper
|
||||
if material and material not in wear:
|
||||
wear = f"{material}{wear}"
|
||||
if pattern and pattern not in wear and pattern != "纯色":
|
||||
wear += f",{pattern}图案"
|
||||
return f"身穿{wear}"
|
||||
parts = []
|
||||
if upper:
|
||||
up = f"{upper_color}{upper}" if upper_color else upper
|
||||
if material and material not in up:
|
||||
up = f"{material}{up}"
|
||||
if pattern and pattern != "纯色" and pattern not in up:
|
||||
up += f"({pattern})"
|
||||
parts.append(f"上身{up}")
|
||||
if lower:
|
||||
lo = f"{lower_color}{lower}" if lower_color else lower
|
||||
parts.append(f"下身{lo}")
|
||||
return ",".join(p for p in parts if p)
|
||||
|
||||
|
||||
def _build_portrait_prompt_old(fj: dict) -> str:
|
||||
if not fj.get("has_person"):
|
||||
name = fj.get("product_name") or "商品"
|
||||
brand = fj.get("brand") or ""
|
||||
colors = fj.get("colors") or []
|
||||
style = fj.get("style") or ""
|
||||
scene = fj.get("scene") or ""
|
||||
mood = fj.get("mood") or ""
|
||||
pieces = []
|
||||
if brand:
|
||||
pieces.append(brand)
|
||||
pieces.append(name)
|
||||
if colors:
|
||||
cnames = []
|
||||
for c in colors:
|
||||
if isinstance(c, dict):
|
||||
cnames.append(c.get("name", ""))
|
||||
elif isinstance(c, str):
|
||||
cnames.append(c)
|
||||
cnames = [c for c in cnames if c][:3]
|
||||
if cnames:
|
||||
pieces.append("、".join(cnames) + "配色")
|
||||
if style:
|
||||
pieces.append(style + "风格")
|
||||
if mood:
|
||||
pieces.append(mood + "氛围")
|
||||
if scene and scene not in ("通用",):
|
||||
pieces.append(scene + "场景")
|
||||
pieces.append("产品特写")
|
||||
prompt = ",".join(p for p in pieces if p)
|
||||
return prompt if len(prompt) >= 10 else "产品展示图,特写镜头"
|
||||
subject = _person_subject_old(fj)
|
||||
wear = _build_wear_sentence_old(fj)
|
||||
accessories = fj.get("accessories") or []
|
||||
if isinstance(accessories, str):
|
||||
accessories = [accessories]
|
||||
acc_str = ",佩戴" + "、".join(str(a) for a in accessories if a) if accessories else ""
|
||||
hairstyle = fj.get("hairstyle") or ""
|
||||
expression = fj.get("expression") or ""
|
||||
pose = fj.get("pose") or ""
|
||||
style = fj.get("style") or ""
|
||||
scene = fj.get("scene") or ""
|
||||
mood = fj.get("mood") or ""
|
||||
detail_parts = []
|
||||
if hairstyle:
|
||||
detail_parts.append(hairstyle)
|
||||
if expression and expression not in ("自然", "平静"):
|
||||
detail_parts.append(f"神情{expression}")
|
||||
if pose and pose not in ("站立",):
|
||||
detail_parts.append(pose)
|
||||
style_parts = []
|
||||
if style:
|
||||
style_parts.append(style)
|
||||
if mood:
|
||||
style_parts.append(mood)
|
||||
if scene and scene not in ("通用",):
|
||||
style_parts.append(scene)
|
||||
pieces = [f"一位{subject}"]
|
||||
if wear:
|
||||
pieces.append(wear)
|
||||
if acc_str:
|
||||
pieces.append(acc_str.lstrip(","))
|
||||
if detail_parts:
|
||||
pieces.append(",".join(detail_parts))
|
||||
pieces.append("".join(style_parts) + "风格" if style_parts else "人像写真")
|
||||
full = ",".join(p for p in pieces if p)
|
||||
if len(full) < 40:
|
||||
full += ",自然光线下人像特写,画面清晰"
|
||||
if len(full) > 120:
|
||||
full = full[:120].rstrip(",") + "。"
|
||||
return full
|
||||
|
||||
|
||||
def _infer_name_old(fj: dict, ocr_texts: list[str]) -> str:
|
||||
pname = fj.get("product_name")
|
||||
if pname and pname != "未识别":
|
||||
return str(pname)
|
||||
if fj.get("has_person"):
|
||||
up = fj.get("upper_wear") or ""
|
||||
if "连衣裙" in up:
|
||||
return up
|
||||
return up or "人物穿搭"
|
||||
if ocr_texts:
|
||||
return max(ocr_texts, key=len)
|
||||
return "未识别"
|
||||
|
||||
|
||||
def _infer_brand_old(fj: dict, ocr_texts: list[str]) -> str:
|
||||
brand = fj.get("brand")
|
||||
if brand:
|
||||
return str(brand)
|
||||
for t in ocr_texts:
|
||||
if 1 < len(t) <= 12:
|
||||
return t
|
||||
return "无法判断"
|
||||
|
||||
|
||||
def _infer_category_old(fj: dict) -> str:
|
||||
cat = fj.get("category")
|
||||
if cat:
|
||||
return str(cat)
|
||||
if fj.get("has_person"):
|
||||
return "服饰"
|
||||
return "非产品图"
|
||||
|
||||
|
||||
def _build_appearance_old(fj: dict) -> str:
|
||||
parts = []
|
||||
for key in ("upper_color", "upper_wear", "material", "pattern"):
|
||||
v = fj.get(key)
|
||||
if v and v not in ("无法判断", "未知", "纯色"):
|
||||
parts.append(str(v))
|
||||
if not parts:
|
||||
return "人像穿搭整体造型" if fj.get("has_person") else "无法判断"
|
||||
return "、".join(parts)
|
||||
|
||||
|
||||
def _build_key_features_old(fj: dict, ocr_texts: list[str]) -> list[str]:
|
||||
feats = []
|
||||
for key in (
|
||||
"upper_wear",
|
||||
"lower_wear",
|
||||
"upper_color",
|
||||
"lower_color",
|
||||
"dress_color",
|
||||
"material",
|
||||
"pattern",
|
||||
"style",
|
||||
"accessories",
|
||||
):
|
||||
v = fj.get(key)
|
||||
if not v:
|
||||
continue
|
||||
if isinstance(v, list):
|
||||
feats.extend(str(x) for x in v if x)
|
||||
elif isinstance(v, str) and v not in ("无法判断", "未知", "纯色"):
|
||||
feats.append(v)
|
||||
if ocr_texts:
|
||||
feats.append(f"画面文字: {'/'.join(ocr_texts[:3])}")
|
||||
out, seen = [], set()
|
||||
for f in feats:
|
||||
f = f.strip()
|
||||
if f and f not in seen and len(f) <= 30:
|
||||
seen.add(f)
|
||||
out.append(f)
|
||||
return out[:6] if out else ["无法判断"]
|
||||
|
||||
|
||||
def _flatten_colors(c) -> list[str]:
|
||||
"""colors可能是字符串数组或[{hex,name,coverage}],统一返回名字数组"""
|
||||
if not c:
|
||||
return []
|
||||
out = []
|
||||
for item in c:
|
||||
if isinstance(item, dict):
|
||||
n = item.get("name")
|
||||
if n:
|
||||
out.append(n)
|
||||
elif isinstance(item, str):
|
||||
out.append(item)
|
||||
return out
|
||||
|
||||
|
||||
def assemble_result(idx: int, fast_json: dict | None, ocr_texts: list[str]) -> dict[str, Any]:
|
||||
fj = fast_json or {}
|
||||
ocr_texts = ocr_texts or []
|
||||
|
||||
if _is_v4_schema(fj):
|
||||
return _assemble_v4(idx, fj, ocr_texts)
|
||||
else:
|
||||
return _assemble_old(idx, fj, ocr_texts)
|
||||
|
||||
|
||||
def _assemble_v4(idx: int, fj: dict, ocr_texts: list[str]) -> dict[str, Any]:
|
||||
"""v4嵌套schema → 下游product dict"""
|
||||
vtype = fj.get("type") or "other"
|
||||
products = fj.get("products") or []
|
||||
scene = fj.get("scene") or "通用"
|
||||
mood = fj.get("mood") or ""
|
||||
colors = fj.get("colors") or []
|
||||
visible_text = fj.get("visible_text") or []
|
||||
color_names = _flatten_colors(colors)
|
||||
|
||||
# 合并OCR文字和visible_text
|
||||
pkg_texts = []
|
||||
for vt in visible_text:
|
||||
if isinstance(vt, dict):
|
||||
t = vt.get("text")
|
||||
if t:
|
||||
pkg_texts.append(str(t))
|
||||
elif isinstance(vt, str):
|
||||
pkg_texts.append(vt)
|
||||
pkg_texts.extend(ocr_texts[:5])
|
||||
# 去重
|
||||
seen_t = set()
|
||||
text_on_package = []
|
||||
for t in pkg_texts:
|
||||
t = str(t).strip()
|
||||
if t and t not in seen_t and len(t) <= 50:
|
||||
seen_t.add(t)
|
||||
text_on_package.append(t)
|
||||
text_on_package = text_on_package[:8]
|
||||
|
||||
has_person = fj.get("has_person", False)
|
||||
|
||||
# 人物类
|
||||
if vtype == "person" or has_person:
|
||||
# 取第一个人物信息(v4 schema人物信息在顶层)
|
||||
person_info = fj
|
||||
# 兼容people嵌套
|
||||
ppl = fj.get("people")
|
||||
if isinstance(ppl, dict) and ppl.get("has_person"):
|
||||
person_info = {**fj, **ppl}
|
||||
has_person = True
|
||||
|
||||
portrait_prompt = _build_portrait_prompt_from_v4(person_info)
|
||||
name = person_info.get("upper_wear") or "人物穿搭"
|
||||
if "连衣裙" in name:
|
||||
pass
|
||||
else:
|
||||
lower = person_info.get("lower_wear") or ""
|
||||
if lower:
|
||||
name = f"{name}+{lower}"
|
||||
brand = "无法判断"
|
||||
category = "服饰"
|
||||
outfit_parts = []
|
||||
for k in ("upper_wear", "lower_wear", "dress_color", "upper_color", "lower_color", "outfit_style"):
|
||||
v = person_info.get(k)
|
||||
if v and v not in ("null", None):
|
||||
outfit_parts.append(str(v))
|
||||
appearance = "、".join(outfit_parts) if outfit_parts else "人像穿搭整体造型"
|
||||
# key_features: 穿搭特征+配饰
|
||||
kf = []
|
||||
for k in (
|
||||
"upper_wear",
|
||||
"lower_wear",
|
||||
"upper_color",
|
||||
"lower_color",
|
||||
"hairstyle",
|
||||
"expression",
|
||||
"pose",
|
||||
"outfit_style",
|
||||
):
|
||||
v = person_info.get(k)
|
||||
if v and v not in ("null", None, "无法判断"):
|
||||
kf.append(str(v))
|
||||
acc = person_info.get("accessories") or []
|
||||
if isinstance(acc, list):
|
||||
kf.extend(str(a) for a in acc if a)
|
||||
if text_on_package:
|
||||
kf.append(f"画面文字: {'/'.join(text_on_package[:3])}")
|
||||
kf = kf[:6] or ["无法判断"]
|
||||
summary = (person_info.get("outfit_style") or "") + (person_info.get("upper_wear") or "穿搭")
|
||||
if not summary or summary == "穿搭":
|
||||
summary = "人物穿搭"
|
||||
return {
|
||||
"name": name[:30],
|
||||
"brand": brand,
|
||||
"category": category,
|
||||
"appearance": appearance,
|
||||
"packaging": "人物形象无包装",
|
||||
"text_on_package": text_on_package,
|
||||
"key_features": kf,
|
||||
"scene": scene,
|
||||
"mood": mood,
|
||||
"portrait_prompt": portrait_prompt,
|
||||
"summary": summary[:40],
|
||||
"_source": "v2_fast_json_v4",
|
||||
}
|
||||
|
||||
# 商品类
|
||||
if vtype == "product" and products:
|
||||
# 主商品(第一个position=main或第一个)
|
||||
main = products[0]
|
||||
for p in products:
|
||||
if p.get("position") == "main":
|
||||
main = p
|
||||
break
|
||||
name = main.get("product_name") or "未识别"
|
||||
brand = main.get("brand") or "无法判断"
|
||||
category = main.get("category") or "非产品图"
|
||||
# appearance: 包装外观
|
||||
app_parts = []
|
||||
for k in ("package_color", "package_type", "cap_type", "body_shape", "label_design"):
|
||||
v = main.get(k)
|
||||
if v and v not in ("null", None):
|
||||
app_parts.append(str(v))
|
||||
appearance = ";".join(app_parts) if app_parts else "无法判断"
|
||||
# packaging: 包装信息(直接用package_type+package_color)
|
||||
pkg_parts = []
|
||||
if main.get("package_type"):
|
||||
pkg_parts.append(str(main["package_type"]))
|
||||
if main.get("package_color"):
|
||||
pkg_parts.append(str(main["package_color"]))
|
||||
if main.get("cap_type"):
|
||||
pkg_parts.append(f"配{main['cap_type']}")
|
||||
packaging = ",".join(pkg_parts) if pkg_parts else "无法判断"
|
||||
# key_features: product_features字段
|
||||
feats = main.get("product_features") or []
|
||||
if not isinstance(feats, list):
|
||||
feats = [str(feats)]
|
||||
kf = [str(f) for f in feats if f and len(str(f)) <= 40][:6]
|
||||
# 补充卖点
|
||||
sell = main.get("key_selling_points") or []
|
||||
if isinstance(sell, list):
|
||||
for s in sell[:2]:
|
||||
if s and len(str(s)) <= 30 and str(s) not in kf:
|
||||
kf.append(f"卖点:{s}")
|
||||
if text_on_package:
|
||||
kf.append(f"文字: {'/'.join(text_on_package[:3])}")
|
||||
kf = kf[:6] or ["无法判断"]
|
||||
portrait_prompt = _build_product_prompt_from_v4(main, fj)
|
||||
if brand != "无法判断" and brand not in name:
|
||||
summary = f"{brand} {name}"
|
||||
else:
|
||||
summary = name
|
||||
return {
|
||||
"name": str(name)[:50],
|
||||
"brand": str(brand)[:30],
|
||||
"category": str(category)[:20],
|
||||
"appearance": appearance[:200],
|
||||
"packaging": packaging[:100],
|
||||
"text_on_package": text_on_package,
|
||||
"key_features": kf,
|
||||
"scene": scene,
|
||||
"mood": mood,
|
||||
"portrait_prompt": portrait_prompt[:200],
|
||||
"summary": str(summary)[:60],
|
||||
"_source": "v2_fast_json_v4",
|
||||
}
|
||||
|
||||
# 门店类或其他
|
||||
if vtype == "store":
|
||||
store_type = fj.get("store_type") or "店铺"
|
||||
name = store_type
|
||||
brand = fj.get("brand_signage") or "无法判断"
|
||||
category = "门店场景"
|
||||
visual = fj.get("visual_elements") or []
|
||||
if isinstance(visual, str):
|
||||
visual = [visual]
|
||||
atmosphere = fj.get("atmosphere") or mood
|
||||
appearance_parts = []
|
||||
if fj.get("store_layout"):
|
||||
appearance_parts.append(str(fj["store_layout"]))
|
||||
if visual:
|
||||
appearance_parts.append("、".join(str(v) for v in visual[:3]))
|
||||
if fj.get("cleanliness"):
|
||||
appearance_parts.append(str(fj["cleanliness"]))
|
||||
appearance = ";".join(appearance_parts) if appearance_parts else "门店环境"
|
||||
kf = []
|
||||
if isinstance(visual, list):
|
||||
kf.extend(str(v) for v in visual if v and len(str(v)) <= 30)
|
||||
prods_vis = fj.get("product_categories_visible") or []
|
||||
if isinstance(prods_vis, list):
|
||||
kf.extend(str(c) for c in prods_vis[:3] if c)
|
||||
promo = fj.get("promotion_elements") or []
|
||||
if isinstance(promo, list) and promo:
|
||||
kf.append("促销活动:" + "、".join(str(p) for p in promo[:2]))
|
||||
if text_on_package:
|
||||
kf.append(f"文字: {'/'.join(text_on_package[:3])}")
|
||||
kf = kf[:6] or ["门店场景"]
|
||||
portrait_prompt = f"{brand if brand!='无法判断' else ''}{store_type},{atmosphere},{scene}场景,{('、'.join(color_names[:3])+'配色,') if color_names else ''}产品陈列丰富,门店实拍"
|
||||
portrait_prompt = portrait_prompt.strip(",")
|
||||
summary = f"{store_type}场景"
|
||||
return {
|
||||
"name": name[:30],
|
||||
"brand": str(brand)[:30],
|
||||
"category": category,
|
||||
"appearance": appearance[:200],
|
||||
"packaging": "门店场景无包装",
|
||||
"text_on_package": text_on_package,
|
||||
"key_features": kf,
|
||||
"scene": scene,
|
||||
"mood": atmosphere or mood,
|
||||
"portrait_prompt": portrait_prompt[:200],
|
||||
"summary": summary[:40],
|
||||
"_source": "v2_fast_json_v4",
|
||||
}
|
||||
|
||||
# other 兜底
|
||||
desc = fj.get("description") or "未识别"
|
||||
return {
|
||||
"name": desc[:30],
|
||||
"brand": "无法判断",
|
||||
"category": "非产品图",
|
||||
"appearance": desc[:200],
|
||||
"packaging": "无法判断",
|
||||
"text_on_package": text_on_package,
|
||||
"key_features": [desc[:30]] if desc != "未识别" else ["无法判断"],
|
||||
"scene": scene,
|
||||
"mood": mood,
|
||||
"portrait_prompt": f"{scene},{mood}氛围,{desc}"[:200],
|
||||
"summary": desc[:40],
|
||||
"_source": "v2_fast_json_v4_other",
|
||||
}
|
||||
|
||||
|
||||
def _assemble_old(idx: int, fj: dict, ocr_texts: list[str]) -> dict[str, Any]:
|
||||
"""旧扁平schema(兼容存量prompt或pro兜底输出)"""
|
||||
portrait_prompt = _build_portrait_prompt_old(fj)
|
||||
name = _infer_name_old(fj, ocr_texts)
|
||||
brand = _infer_brand_old(fj, ocr_texts)
|
||||
category = _infer_category_old(fj)
|
||||
appearance = _build_appearance_old(fj)
|
||||
key_features = _build_key_features_old(fj, ocr_texts)
|
||||
scene = fj.get("scene") or "通用"
|
||||
mood = fj.get("mood") or ""
|
||||
packaging = "无法判断"
|
||||
text_on_package = ocr_texts[:8]
|
||||
if fj.get("has_person"):
|
||||
up = fj.get("upper_wear") or "穿搭"
|
||||
style = fj.get("style") or ""
|
||||
summary = f"{style}{up}" if style and style not in up else up
|
||||
elif brand != "无法判断" and name != brand:
|
||||
summary = f"{brand} {name}"
|
||||
else:
|
||||
summary = name
|
||||
return {
|
||||
"name": name,
|
||||
"brand": brand,
|
||||
"category": category,
|
||||
"appearance": appearance,
|
||||
"packaging": packaging,
|
||||
"text_on_package": text_on_package,
|
||||
"key_features": key_features,
|
||||
"scene": scene,
|
||||
"mood": mood,
|
||||
"portrait_prompt": portrait_prompt,
|
||||
"summary": summary,
|
||||
"_source": "v2_fast_json",
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""V2 图片分析主路径:每图并行 OCR(火山MediaKit,未配置时自动跳过)+ qwen3.8-flash JSON VLM,
|
||||
失败时单次 qwen3.7-plus 兜底。
|
||||
|
||||
架构(灵应10-05确认):
|
||||
- 唯一后端:阿里云百炼 DashScope,qwen3.8-flash 做快速路径、qwen3.7-plus 做兜底
|
||||
- 主力:单图2路并行(OCR + fast VLM),外层N图全并发(workers=8)
|
||||
- 兜底:单次 pro VLM 调用,无竞速/重试/复杂超时
|
||||
- 输出 dict 格式与旧版完全一致,下游零改动
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from typing import Any
|
||||
|
||||
from . import assembler, ocr_volc, vlm_fallback, vlm_fast_json
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 超时(可通过环境变量覆盖)
|
||||
_IMG_WORKERS = int(os.environ.get("VISION_V2_IMG_WORKERS", "8"))
|
||||
_FAST_TIMEOUT = float(os.environ.get("VISION_V2_FAST_TIMEOUT", "15"))
|
||||
_FAST_JSON_TIMEOUT = float(os.environ.get("VISION_V2_FAST_JSON_TIMEOUT", "15"))
|
||||
_OCR_TIMEOUT = float(os.environ.get("VISION_V2_OCR_TIMEOUT", "6"))
|
||||
_PRO_TIMEOUT = float(os.environ.get("VISION_V2_PRO_TIMEOUT", "30"))
|
||||
|
||||
_FALLBACK_RESULT = {
|
||||
"name": "未识别",
|
||||
"brand": "无法判断",
|
||||
"category": "非产品图",
|
||||
"appearance": "无法判断",
|
||||
"packaging": "无法判断",
|
||||
"text_on_package": [],
|
||||
"key_features": ["无法判断"],
|
||||
"scene": "通用",
|
||||
"mood": "",
|
||||
"portrait_prompt": "无法判断",
|
||||
"summary": "未识别",
|
||||
}
|
||||
|
||||
|
||||
def _is_usable(r: dict[str, Any]) -> bool:
|
||||
pp = (r.get("portrait_prompt") or "").strip()
|
||||
if pp and pp not in ("无人像", "无法判断", "未识别"):
|
||||
return True
|
||||
name = (r.get("name") or "").strip()
|
||||
if name and name not in ("未识别", "无法判断", "未知"):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def analyze_image_v2(idx: int, img_url: str) -> dict[str, Any]:
|
||||
t0 = time.time()
|
||||
|
||||
fj_result: dict[str, Any] | None = None
|
||||
ocr_result: list[str] = []
|
||||
with ThreadPoolExecutor(max_workers=2) as pool:
|
||||
f_fj = pool.submit(vlm_fast_json.call_fast_json, img_url, timeout=_FAST_JSON_TIMEOUT)
|
||||
f_ocr = pool.submit(ocr_volc.call_ocr, img_url, timeout=_OCR_TIMEOUT)
|
||||
try:
|
||||
for fut in as_completed([f_fj, f_ocr], timeout=_FAST_TIMEOUT):
|
||||
try:
|
||||
res = fut.result(timeout=1)
|
||||
except Exception as e:
|
||||
logger.warning("[vision.v2] 图片 #%d 子任务异常: %s", idx, e)
|
||||
continue
|
||||
if fut is f_fj and isinstance(res, dict):
|
||||
fj_result = res
|
||||
elif fut is f_ocr and isinstance(res, list):
|
||||
ocr_result = res
|
||||
except TimeoutError:
|
||||
for f in (f_fj, f_ocr):
|
||||
if not f.done():
|
||||
f.cancel()
|
||||
logger.warning("[vision.v2] 图片 #%d fast路径超时(%.0fs),走pro兜底", idx, _FAST_TIMEOUT)
|
||||
|
||||
fast_elapsed = time.time() - t0
|
||||
|
||||
if fj_result:
|
||||
assembled = assembler.assemble_result(idx, fj_result, ocr_result)
|
||||
if _is_usable(assembled):
|
||||
assembled["_fast_elapsed"] = round(fast_elapsed, 2)
|
||||
logger.info(
|
||||
"[vision.v2] 图片 #%d fast命中 elapsed=%.2fs pp=%s",
|
||||
idx,
|
||||
fast_elapsed,
|
||||
(assembled.get("portrait_prompt") or "")[:40],
|
||||
)
|
||||
return assembled
|
||||
|
||||
pro_t0 = time.time()
|
||||
pro_result = vlm_fallback.call_pro_vlm(img_url, idx, timeout=_PRO_TIMEOUT)
|
||||
if pro_result and _is_usable(pro_result):
|
||||
pro_result["_fallback_used"] = True
|
||||
pro_result["_fast_elapsed"] = round(fast_elapsed, 2)
|
||||
pro_result["_pro_elapsed"] = round(time.time() - pro_t0, 2)
|
||||
if ocr_result and not pro_result.get("text_on_package"):
|
||||
pro_result["text_on_package"] = ocr_result[:8]
|
||||
logger.info("[vision.v2] 图片 #%d pro兜底命中 total=%.2fs", idx, time.time() - t0)
|
||||
return pro_result
|
||||
|
||||
logger.warning("[vision.v2] 图片 #%d 全路径失败 elapsed=%.2fs", idx, time.time() - t0)
|
||||
out = dict(_FALLBACK_RESULT)
|
||||
out["_source"] = "v2_all_failed"
|
||||
out["text_on_package"] = ocr_result[:8]
|
||||
out["_fast_elapsed"] = round(fast_elapsed, 2)
|
||||
return out
|
||||
|
||||
|
||||
def analyze_images_v2(img_urls: list[str]) -> list[dict[str, Any]]:
|
||||
if not img_urls:
|
||||
return []
|
||||
workers = min(_IMG_WORKERS, len(img_urls), 16)
|
||||
results: list[dict[str, Any] | None] = [None] * len(img_urls)
|
||||
|
||||
logger.info(
|
||||
"[vision.v2] 开始图片分析 n=%d workers=%d fast_timeout=%.0fs pro_timeout=%.0fs",
|
||||
len(img_urls),
|
||||
workers,
|
||||
_FAST_TIMEOUT,
|
||||
_PRO_TIMEOUT,
|
||||
)
|
||||
t0 = time.time()
|
||||
with ThreadPoolExecutor(max_workers=workers) as pool:
|
||||
future_to_idx = {pool.submit(analyze_image_v2, idx, url): idx for idx, url in enumerate(img_urls)}
|
||||
for fut in as_completed(future_to_idx):
|
||||
idx = future_to_idx[fut]
|
||||
try:
|
||||
results[idx] = fut.result()
|
||||
except Exception as e:
|
||||
logger.warning("[vision.v2] 图片 #%d future异常: %s", idx, e, exc_info=True)
|
||||
r = dict(_FALLBACK_RESULT)
|
||||
r["_source"] = "v2_future_exception"
|
||||
results[idx] = r
|
||||
|
||||
elapsed = time.time() - t0
|
||||
succ = sum(1 for r in results if r and _is_usable(r))
|
||||
fb = sum(1 for r in results if r and r.get("_fallback_used"))
|
||||
logger.info("[vision.v2] 完成 n=%d usable=%d pro_fallback=%d elapsed=%.2fs", len(img_urls), succ, fb, elapsed)
|
||||
return [r for r in results if r is not None]
|
||||
@@ -0,0 +1,109 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""火山引擎 AI MediaKit OCR(同步)调用封装。
|
||||
|
||||
接口:POST {mediakit_base_url}/tools-sync/ocr
|
||||
鉴权:Bearer {mediakit_api_key}
|
||||
请求体:{"image_url": "<公网可访问URL>"} (部分版本也支持 image_base64)
|
||||
响应:{"code":0,"data":{"texts":[{"text":"...","bbox":[x,y,w,h],...},...],...}}
|
||||
|
||||
目标:识别商品包装/Logo/水印上的文字,作为 fast_json VLM 的补充。
|
||||
返回值:识别到的文本字符串列表(失败返回 [])。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_TIMEOUT = 8 # OCR 秒级返回,8s 绰绰有余
|
||||
|
||||
|
||||
def call_ocr(img_url: str, *, timeout: int = DEFAULT_TIMEOUT) -> list[str]:
|
||||
"""调用 MediaKit 同步 OCR,返回去重后的纯文本列表。
|
||||
|
||||
不做重试(外层降级逻辑负责)。失败/未配置返回空列表,不抛异常。
|
||||
"""
|
||||
t0 = time.time()
|
||||
try:
|
||||
import httpx
|
||||
|
||||
from packages.shared.mediakit_client import get_mediakit_client
|
||||
|
||||
client = get_mediakit_client()
|
||||
if not client.is_available:
|
||||
logger.info("[vision.v2] mediakit 未配置,跳过 OCR")
|
||||
return []
|
||||
|
||||
url = f"{client.base_url}/tools-sync/ocr"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {client.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
payload: dict[str, Any] = {"image_url": img_url}
|
||||
# 部分文档版本用 image_base64,但公网 URL 场景下 image_url 最简
|
||||
resp = httpx.post(url, headers=headers, json=payload, timeout=timeout)
|
||||
elapsed = time.time() - t0
|
||||
if resp.status_code != 200:
|
||||
logger.warning(
|
||||
"[vision.v2] OCR HTTP %d elapsed=%.1fs body=%s",
|
||||
resp.status_code,
|
||||
elapsed,
|
||||
resp.text[:200],
|
||||
)
|
||||
return []
|
||||
data = resp.json()
|
||||
# 兼容几种可能的响应结构
|
||||
code = data.get("code", data.get("status", 0))
|
||||
if code not in (0, "OK", "success", 200):
|
||||
logger.warning("[vision.v2] OCR 业务错误 code=%s elapsed=%.1fs resp=%s", code, elapsed, str(data)[:200])
|
||||
return []
|
||||
texts = _extract_texts(data)
|
||||
# 去重 + 过滤空
|
||||
seen: set[str] = set()
|
||||
out: list[str] = []
|
||||
for t in texts:
|
||||
t = (t or "").strip()
|
||||
if t and t not in seen and len(t) <= 100: # 过滤过长的误识别
|
||||
seen.add(t)
|
||||
out.append(t)
|
||||
logger.info("[vision.v2] OCR 完成 elapsed=%.1fs n=%d texts=%s", elapsed, len(out), out[:5])
|
||||
return out
|
||||
except Exception as e:
|
||||
elapsed = time.time() - t0
|
||||
logger.warning("[vision.v2] OCR 异常 elapsed=%.1fs err=%s", elapsed, e, exc_info=True)
|
||||
return []
|
||||
|
||||
|
||||
def _extract_texts(data: dict) -> list[str]:
|
||||
"""从 OCR 响应中抽取文本,兼容多种结构。"""
|
||||
out: list[str] = []
|
||||
# 常见结构1: data.texts = [{"text": "..."}, ...]
|
||||
d = data.get("data") or data
|
||||
if isinstance(d, dict):
|
||||
for key in ("texts", "lines", "words", "items", "result"):
|
||||
items = d.get(key)
|
||||
if isinstance(items, list):
|
||||
for it in items:
|
||||
if isinstance(it, dict):
|
||||
txt = it.get("text") or it.get("content") or it.get("word")
|
||||
if txt:
|
||||
out.append(str(txt))
|
||||
elif isinstance(it, str):
|
||||
out.append(it)
|
||||
break
|
||||
# 结构2: data.text = "..."
|
||||
if not out:
|
||||
t = d.get("text")
|
||||
if isinstance(t, str):
|
||||
out.append(t)
|
||||
# 结构3: data.ocr_text / data.content
|
||||
if not out:
|
||||
for key in ("ocr_text", "content", "raw_text"):
|
||||
v = d.get(key)
|
||||
if isinstance(v, str) and v.strip():
|
||||
out.append(v)
|
||||
break
|
||||
return out
|
||||
@@ -0,0 +1,126 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""V2 兜底路径:qwen3.7-plus(阿里云百炼/DashScope)单图调用。
|
||||
|
||||
fast_json 超时/返回非 JSON/识别为空时,本路径单次调用兜底。
|
||||
设计要点:
|
||||
- 直接 httpx 直连 DashScope,不走 ai_client
|
||||
- enable_thinking=false + response_format=json_object
|
||||
- system prompt 优先读后台 viral_video_prompt_templates 配置,DB不可用时fallback到硬编码JSON schema
|
||||
- timeout=25s
|
||||
- API Key 从环境变量 DASHSCOPE_API_KEY 读取
|
||||
- 返回 dict 统一走 assembler.assemble_result 组装,与 fast 路径输出格式完全一致
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from . import _prompt, assembler
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_BASE_URL = "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||
_PRO_MODEL = "qwen3.7-plus"
|
||||
_DEFAULT_TIMEOUT = 30
|
||||
_DEFAULT_MAX_TOKENS = 800
|
||||
|
||||
|
||||
def _api_key() -> str | None:
|
||||
return os.environ.get("DASHSCOPE_API_KEY")
|
||||
|
||||
|
||||
def call_pro_vlm(
|
||||
img_url: str,
|
||||
idx: int,
|
||||
*,
|
||||
timeout: int = _DEFAULT_TIMEOUT,
|
||||
) -> dict[str, Any] | None:
|
||||
t0 = time.time()
|
||||
import httpx
|
||||
|
||||
api_key = _api_key()
|
||||
if not api_key:
|
||||
logger.warning("[vision.v2] pro DASHSCOPE_API_KEY 未配置,跳过")
|
||||
return None
|
||||
|
||||
system_prompt, user_prompt = _prompt.resolve_pro_prompt()
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"model": _PRO_MODEL,
|
||||
"messages": [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "image_url", "image_url": {"url": img_url}},
|
||||
{"type": "text", "text": user_prompt},
|
||||
],
|
||||
},
|
||||
],
|
||||
"temperature": 0.3,
|
||||
"max_tokens": _DEFAULT_MAX_TOKENS,
|
||||
"stream": False,
|
||||
"enable_thinking": False,
|
||||
"response_format": {"type": "json_object"},
|
||||
}
|
||||
try:
|
||||
r = httpx.post(
|
||||
f"{_BASE_URL}/chat/completions",
|
||||
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
|
||||
json=payload,
|
||||
timeout=timeout,
|
||||
)
|
||||
elapsed = time.time() - t0
|
||||
if r.status_code != 200:
|
||||
logger.warning("[vision.v2] pro HTTP %d elapsed=%.1fs body=%s", r.status_code, elapsed, r.text[:200])
|
||||
return None
|
||||
data = r.json()
|
||||
raw = (data.get("choices") or [{}])[0].get("message", {}).get("content")
|
||||
if not raw:
|
||||
logger.warning("[vision.v2] pro 返回空 elapsed=%.1fs", elapsed)
|
||||
return None
|
||||
usage = data.get("usage") or {}
|
||||
reasoning_tokens = usage.get("reasoning_tokens", 0)
|
||||
ctd = usage.get("completion_tokens_details") or {}
|
||||
if not reasoning_tokens:
|
||||
reasoning_tokens = ctd.get("reasoning_tokens", 0)
|
||||
logger.info(
|
||||
"[vision.v2] pro 完成 model=%s elapsed=%.1fs in=%d out=%d reasoning=%d",
|
||||
_PRO_MODEL,
|
||||
elapsed,
|
||||
usage.get("prompt_tokens", 0),
|
||||
usage.get("completion_tokens", 0),
|
||||
reasoning_tokens,
|
||||
)
|
||||
s = raw.strip()
|
||||
if s.startswith("```"):
|
||||
lines = s.split("\n")
|
||||
if lines and lines[0].startswith("```"):
|
||||
lines = lines[1:]
|
||||
if lines and lines[-1].strip().startswith("```"):
|
||||
lines = lines[:-1]
|
||||
s = "\n".join(lines).strip()
|
||||
lpos, rr = s.find("{"), s.rfind("}")
|
||||
if lpos >= 0 and rr > lpos:
|
||||
s = s[lpos : rr + 1]
|
||||
try:
|
||||
obj = json.loads(s)
|
||||
except json.JSONDecodeError:
|
||||
logger.warning("[vision.v2] pro JSON 解析失败 head=%s", raw[:200])
|
||||
return None
|
||||
if not isinstance(obj, dict):
|
||||
return None
|
||||
|
||||
# 通过assembler统一组装,兼容v4嵌套schema和旧扁平schema
|
||||
result = assembler.assemble_result(idx, obj, [])
|
||||
result["_source"] = "vlm_pro"
|
||||
result["_fallback_used"] = True
|
||||
return result
|
||||
except Exception as e:
|
||||
elapsed = time.time() - t0
|
||||
logger.warning("[vision.v2] pro 异常 elapsed=%.1fs err=%s", elapsed, e, exc_info=True)
|
||||
return None
|
||||
@@ -0,0 +1,150 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""V2 快速路径:qwen3.8-flash(阿里云百炼/DashScope)强约束 JSON-only 调用。
|
||||
|
||||
目标:替代"人体属性/商品检测/图像标签"三个火山不存在的专用云端 API。
|
||||
设计要点:
|
||||
- 直接用 httpx 发最小 payload 到 DashScope OpenAI 兼容 endpoint,不走 ai_client 包装
|
||||
- enable_thinking=false 关闭推理链(reasoning 是延迟主因)
|
||||
- response_format=json_object 强约束JSON输出
|
||||
- system prompt 优先读后台 viral_video_prompt_templates 配置,DB不可用时fallback到硬编码JSON schema
|
||||
- max_tokens=350、temperature=0.1(稳定输出 JSON)
|
||||
- timeout=12s(失败由外层走 pro 兜底)
|
||||
- API Key 从环境变量 DASHSCOPE_API_KEY 读取
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from . import _prompt
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# DashScope OpenAI 兼容 endpoint
|
||||
_BASE_URL = "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||
_FAST_MODEL = "qwen3.8-flash"
|
||||
_DEFAULT_TIMEOUT = 15
|
||||
_DEFAULT_MAX_TOKENS = 350
|
||||
|
||||
|
||||
def _api_key() -> str | None:
|
||||
return os.environ.get("DASHSCOPE_API_KEY")
|
||||
|
||||
|
||||
def _strip_code_fence(s: str) -> str:
|
||||
s = s.strip()
|
||||
if s.startswith("```"):
|
||||
lines = s.split("\n")
|
||||
if lines and lines[0].startswith("```"):
|
||||
lines = lines[1:]
|
||||
if lines and lines[-1].strip().startswith("```"):
|
||||
lines = lines[:-1]
|
||||
s = "\n".join(lines).strip()
|
||||
return s
|
||||
|
||||
|
||||
def call_fast_json(
|
||||
img_url: str,
|
||||
*,
|
||||
timeout: int = _DEFAULT_TIMEOUT,
|
||||
max_tokens: int = _DEFAULT_MAX_TOKENS,
|
||||
) -> dict[str, Any] | None:
|
||||
"""调用 qwen3.8-flash 返回结构化 dict;失败/非 JSON 返回 None。"""
|
||||
t0 = time.time()
|
||||
import httpx
|
||||
|
||||
api_key = _api_key()
|
||||
if not api_key:
|
||||
logger.warning("[vision.v2] DASHSCOPE_API_KEY 未配置,跳过 fast_json")
|
||||
return None
|
||||
|
||||
system_prompt, user_prompt = _prompt.resolve_fast_prompt()
|
||||
|
||||
url = f"{_BASE_URL}/chat/completions"
|
||||
payload: dict[str, Any] = {
|
||||
"model": _FAST_MODEL,
|
||||
"messages": [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "image_url", "image_url": {"url": img_url}},
|
||||
{"type": "text", "text": user_prompt},
|
||||
],
|
||||
},
|
||||
],
|
||||
"temperature": 0.1,
|
||||
"max_tokens": max_tokens,
|
||||
"stream": False,
|
||||
"enable_thinking": False,
|
||||
"response_format": {"type": "json_object"},
|
||||
}
|
||||
try:
|
||||
resp = httpx.post(
|
||||
url,
|
||||
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
|
||||
json=payload,
|
||||
timeout=timeout,
|
||||
)
|
||||
elapsed = time.time() - t0
|
||||
if resp.status_code == 400 and "enable_thinking" in resp.text[:300].lower():
|
||||
logger.warning("[vision.v2] fast_json HTTP 400 thinking 参数不兼容,重试 elapsed=%.1fs", elapsed)
|
||||
payload.pop("enable_thinking", None)
|
||||
resp = httpx.post(
|
||||
url,
|
||||
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
|
||||
json=payload,
|
||||
timeout=timeout,
|
||||
)
|
||||
elapsed = time.time() - t0
|
||||
if resp.status_code != 200:
|
||||
logger.warning(
|
||||
"[vision.v2] fast_json HTTP %d elapsed=%.1fs body=%s", resp.status_code, elapsed, resp.text[:200]
|
||||
)
|
||||
return None
|
||||
data = resp.json()
|
||||
raw = (data.get("choices") or [{}])[0].get("message", {}).get("content")
|
||||
if not raw:
|
||||
logger.warning("[vision.v2] fast_json 返回空 elapsed=%.1fs", elapsed)
|
||||
return None
|
||||
usage = data.get("usage") or {}
|
||||
reasoning_tokens = usage.get("reasoning_tokens", 0)
|
||||
ctd = usage.get("completion_tokens_details") or {}
|
||||
if not reasoning_tokens:
|
||||
reasoning_tokens = ctd.get("reasoning_tokens", 0)
|
||||
logger.info(
|
||||
"[vision.v2] fast_json 完成 model=%s elapsed=%.1fs in=%d out=%d reasoning=%d",
|
||||
_FAST_MODEL,
|
||||
elapsed,
|
||||
usage.get("prompt_tokens", 0),
|
||||
usage.get("completion_tokens", 0),
|
||||
reasoning_tokens,
|
||||
)
|
||||
text = _strip_code_fence(raw)
|
||||
lpos, r = text.find("{"), text.rfind("}")
|
||||
if lpos >= 0 and r > lpos:
|
||||
text = text[lpos : r + 1]
|
||||
try:
|
||||
obj = json.loads(text)
|
||||
except json.JSONDecodeError:
|
||||
logger.warning("[vision.v2] fast_json JSON 解析失败 elapsed=%.1fs head=%s", elapsed, raw[:200])
|
||||
return None
|
||||
if not isinstance(obj, dict):
|
||||
logger.warning("[vision.v2] fast_json 非 dict: %s", type(obj))
|
||||
return None
|
||||
logger.info(
|
||||
"[vision.v2] fast_json 完成 elapsed=%.1fs has_person=%s has_product=%s category=%s",
|
||||
elapsed,
|
||||
obj.get("has_person"),
|
||||
obj.get("has_product"),
|
||||
obj.get("category"),
|
||||
)
|
||||
return obj
|
||||
except Exception as e:
|
||||
elapsed = time.time() - t0
|
||||
logger.warning("[vision.v2] fast_json 异常 elapsed=%.1fs err=%s", elapsed, e, exc_info=True)
|
||||
return None
|
||||
@@ -335,6 +335,10 @@ class GenerationTaskModel(Base):
|
||||
bgm_config = Column(JSON, nullable=False, default=dict)
|
||||
extra_meta = Column("metadata", JSON, nullable=False, default=dict)
|
||||
logs = Column(Text, nullable=False, default="[]", server_default="[]")
|
||||
# 功能计费(smart_edit):预扣积分 / 最终积分 / 预扣流水 ID
|
||||
credits_prepaid = Column(Float, nullable=False, default=0.0, server_default="0")
|
||||
credits_cost = Column(Float, nullable=False, default=0.0, server_default="0")
|
||||
credits_transaction_id = Column(String(36), nullable=False, default="", server_default="")
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(
|
||||
DateTime,
|
||||
@@ -727,6 +731,11 @@ class LipsyncJobModel(Base):
|
||||
# 精确句子时间戳(TTS 合成后由 silencedetect 计算,用于 B-roll 精确定位)
|
||||
sentence_timings = Column(JSON, nullable=True) # list[{index,text,start_time,end_time}]
|
||||
|
||||
# 功能计费(lip_sync):预扣积分 / 最终积分 / 预扣流水 ID
|
||||
credits_prepaid = Column(Float, nullable=False, default=0.0, server_default="0")
|
||||
credits_cost = Column(Float, nullable=False, default=0.0, server_default="0")
|
||||
credits_transaction_id = Column(String(36), nullable=False, default="", server_default="")
|
||||
|
||||
# 时间戳
|
||||
submitted_at = Column(DateTime, nullable=True)
|
||||
completed_at = Column(DateTime, nullable=True)
|
||||
@@ -905,6 +914,11 @@ class GpuLipsyncTaskModel(Base):
|
||||
# 心跳:worker 最近一次 poll/result 的时间,用于判定 worker 失联
|
||||
last_heartbeat_at = Column(DateTime, nullable=True)
|
||||
|
||||
# 功能计费(lip_sync):预扣积分 / 最终积分 / 预扣流水 ID
|
||||
credits_prepaid = Column(Float, nullable=False, default=0.0, server_default="0")
|
||||
credits_cost = Column(Float, nullable=False, default=0.0, server_default="0")
|
||||
credits_transaction_id = Column(String(36), nullable=False, default="", server_default="")
|
||||
|
||||
|
||||
class GpuWorkerModel(Base):
|
||||
"""GPU Worker 注册表 — 反向轮询模式下用于心跳与监控."""
|
||||
|
||||
@@ -57,7 +57,18 @@ _IMAGE_ANALYSIS_SYSTEM = f"""你是电商商品视觉分析师,负责从商品
|
||||
<quality> 用一个标签,属性 resolution、lighting、composition、blur 描述画质。
|
||||
<key_selling_points> 下面每个卖点用一个 <point> 标签。
|
||||
|
||||
看不到或无法判断的内容,属性值填“无法判断”,布尔值填 false,不要留空标签。"""
|
||||
【人物属性硬性要求(has_person=true时必须遵守)】
|
||||
hair/skin_tone/face_shape/outfit四项绝对禁止填“无法判断”,必须基于图片可见特征给出具体中文描述:
|
||||
- hair:必须描述发型+发色,如“黑色齐肩直发”“棕色微卷中长发”“深棕色短发”
|
||||
- skin_tone:必须描述肤色,如“暖调自然肤色”“白皙肤色”“小麦色”
|
||||
- face_shape:必须描述脸型,如“鹅蛋脸”“圆脸”“瓜子脸”“方脸”
|
||||
- outfit:必须描述可见穿着,如“米色翻领衬衫”“白色T恤”“黑色连衣裙”
|
||||
即使局部被遮挡也要根据可见部分合理推断;确实看不清时按最接近的直观印象描述。
|
||||
|
||||
其他非人物属性看不到或无法判断时填“无法判断”,布尔值填false,不要留空标签。
|
||||
|
||||
【有人物场景输出参考(女性手持商品示例,必须写全10个属性,禁止省略)】
|
||||
<people has_person="true" count="1" gender="女" age_range="青年" hair="黑色齐肩直发" skin_tone="暖调自然肤色" face_shape="鹅蛋脸" outfit="米色翻领衬衫" pose="正面半身,手持商品" expression="面带微笑"/>"""
|
||||
|
||||
_IMAGE_ANALYSIS_USER = """请分析以下商品图片,共 {image_count} 张。
|
||||
所属行业:{industry}
|
||||
|
||||
@@ -103,7 +103,7 @@ class SharedSettings(BaseSettings):
|
||||
doubao_vision_lite_model: str = (
|
||||
"doubao-seed-2-1-lite-260915" # 快速视觉(Seed 2.1 Lite 原生多模态;原 vision-lite-250315 不可用)
|
||||
)
|
||||
doubao_vision_use_lite: bool = False # #2181: lite视觉模型100%超时,默认关闭走pro(25-38s稳定返回)
|
||||
doubao_vision_use_lite: bool = True # #2188: lite恢复稳定,爆款视频默认lite-first提速(20-30s)
|
||||
doubao_embedding_model: str = "doubao-embedding-vision-251215" # 多模态向量化(原 large-text-240915 已 Retiring)
|
||||
doubao_video_model: str = "doubao-seedance-2-5-260628"
|
||||
doubao_video_timeout: int = 600 # 视频生成轮询总超时(秒)
|
||||
|
||||
Executable
+376
@@ -0,0 +1,376 @@
|
||||
"""功能计费配置服务:从 feature_pricing_configs 读配置,300 秒 TTL 内存缓存。
|
||||
|
||||
配置表由 xiaoxia-admin 侧维护(同库 PostgreSQL),本服务只读。
|
||||
DB 不可用 / 表不存在 / 无数据时自动回落到内置兜底配置,保证业务不崩。
|
||||
|
||||
计费公式:最终积分 = (动态成本 + 固定成本) × 利润系数,price_cap 封顶。
|
||||
启用条件:全局 points_enabled 总开关 AND 功能 is_enabled 同时为 true。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from packages.adapters.sqlalchemy_impl import session as _session_mod
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
CACHE_TTL_SECONDS = 300.0
|
||||
|
||||
# ── 爆款视频兜底模型单价(与旧硬编码表/现状一致;DB 不可用时使用) ───────
|
||||
# 结构:models[model_key][resolution]["true"/"false"] = 单价
|
||||
# token 模式:元/百万输出 tokens;per_second 模式:元/秒
|
||||
# 注意:仅 seedance-2.5 配置 true(图生视频)单价;其余模型只有 false,
|
||||
# 精确 key 缺失时由 points_rules 回落到 seedance-2.5/false(与旧现状一致)。
|
||||
_FALLBACK_VIRAL_MODEL_PRICING: dict = {
|
||||
"seedance-2.5": {
|
||||
"480p": {"false": 70.0, "true": 42.0},
|
||||
"720p": {"false": 70.0, "true": 42.0},
|
||||
"1080p": {"false": 77.0, "true": 46.0},
|
||||
},
|
||||
"seedance-2.0": {
|
||||
"480p": {"false": 46.0},
|
||||
"720p": {"false": 46.0},
|
||||
"1080p": {"false": 51.0},
|
||||
"4k": {"false": 80.0},
|
||||
},
|
||||
"seedance-2.0-fast": {
|
||||
"480p": {"false": 28.0},
|
||||
"720p": {"false": 28.0},
|
||||
},
|
||||
"seedance-2.0-mini": {
|
||||
"480p": {"false": 9.2},
|
||||
"720p": {"false": 9.2},
|
||||
},
|
||||
"wan-3.0": {
|
||||
"480p": {"false": 0.3},
|
||||
"720p": {"false": 0.6},
|
||||
"1080p": {"false": 1.2},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class FeatureConfig:
|
||||
"""功能计费配置快照。"""
|
||||
|
||||
feature_key: str
|
||||
name: str = ""
|
||||
emoji: str = ""
|
||||
is_enabled: bool = False
|
||||
fixed_cost: float = 0.0
|
||||
profit_multiplier: float = 1.0
|
||||
dynamic_unit_cost: float = 0.0
|
||||
billing_mode: str = "model_based"
|
||||
price_cap: float = 0.0
|
||||
model_pricing: dict = field(default_factory=dict)
|
||||
description: str = ""
|
||||
|
||||
|
||||
# ── 进程内缓存:(loaded_monotonic, {feature_key: FeatureConfig}) ──────────
|
||||
_lock = threading.Lock()
|
||||
_cache: Optional[tuple[float, dict[str, FeatureConfig]]] = None
|
||||
|
||||
|
||||
def _fallback_configs() -> dict[str, FeatureConfig]:
|
||||
"""内置兜底配置:爆款启用(与现状一致),其余两个关闭。"""
|
||||
return {
|
||||
"viral_video": FeatureConfig(
|
||||
feature_key="viral_video",
|
||||
name="爆款视频",
|
||||
emoji="🎬",
|
||||
is_enabled=True,
|
||||
fixed_cost=0.15,
|
||||
profit_multiplier=1.3,
|
||||
dynamic_unit_cost=0.0,
|
||||
billing_mode="model_based",
|
||||
price_cap=0.0,
|
||||
model_pricing=json.loads(json.dumps(_FALLBACK_VIRAL_MODEL_PRICING)),
|
||||
description="爆款视频动态定价(兜底配置)",
|
||||
),
|
||||
"lip_sync": FeatureConfig(
|
||||
feature_key="lip_sync",
|
||||
name="对口型",
|
||||
emoji="🎙️",
|
||||
is_enabled=False,
|
||||
fixed_cost=0.0,
|
||||
profit_multiplier=1.0,
|
||||
dynamic_unit_cost=0.0,
|
||||
billing_mode="per_second",
|
||||
price_cap=0.0,
|
||||
description="对口型计费(兜底配置,默认关闭)",
|
||||
),
|
||||
"smart_edit": FeatureConfig(
|
||||
feature_key="smart_edit",
|
||||
name="智能剪辑",
|
||||
emoji="✂️",
|
||||
is_enabled=False,
|
||||
fixed_cost=0.0,
|
||||
profit_multiplier=1.0,
|
||||
dynamic_unit_cost=0.0,
|
||||
billing_mode="model_based",
|
||||
price_cap=0.0,
|
||||
description="智能剪辑固定价计费(兜底配置,默认关闭)",
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
_lazy_session = None
|
||||
|
||||
|
||||
def _get_session():
|
||||
"""优先用全局 SessionLocal(worker);否则按应用配置懒建同步引擎(api)。"""
|
||||
global _lazy_session
|
||||
if _session_mod.SessionLocal is not None:
|
||||
return _session_mod.SessionLocal()
|
||||
if _lazy_session is not None:
|
||||
return _lazy_session()
|
||||
try:
|
||||
from packages.config import get_shared_settings
|
||||
|
||||
url = str(get_shared_settings().database_url)
|
||||
except Exception: # noqa: BLE001
|
||||
return None
|
||||
if not url:
|
||||
return None
|
||||
url = url.replace("postgresql+asyncpg://", "postgresql+psycopg://")
|
||||
if url.startswith("postgresql://"):
|
||||
url = url.replace("postgresql://", "postgresql+psycopg://")
|
||||
engine = sa.create_engine(url, pool_pre_ping=True, pool_size=2, max_overflow=2)
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
_lazy_session = sessionmaker(bind=engine)
|
||||
return _lazy_session()
|
||||
|
||||
|
||||
def _parse_model_pricing(raw) -> dict:
|
||||
"""解析 model_pricing_json(Text JSON),空/失败 → {}。"""
|
||||
if raw is None:
|
||||
return {}
|
||||
if isinstance(raw, dict):
|
||||
return raw
|
||||
text = str(raw).strip()
|
||||
if not text:
|
||||
return {}
|
||||
try:
|
||||
data = json.loads(text)
|
||||
except (ValueError, TypeError):
|
||||
logger.warning("model_pricing_json 解析失败,按空配置处理: %r", text[:200])
|
||||
return {}
|
||||
return data if isinstance(data, dict) else {}
|
||||
|
||||
|
||||
def _to_float(value, default: float = 0.0) -> float:
|
||||
try:
|
||||
if value is None:
|
||||
return default
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
|
||||
|
||||
def _load_all() -> dict[str, FeatureConfig]:
|
||||
"""SELECT * FROM feature_pricing_configs,返回 {feature_key: FeatureConfig}。
|
||||
|
||||
表不存在 / DB 异常由调用方捕获并回落兜底配置。
|
||||
"""
|
||||
session = None
|
||||
try:
|
||||
session = _get_session()
|
||||
if session is None:
|
||||
raise RuntimeError("no db session available")
|
||||
sql = sa.text("""
|
||||
SELECT feature_key, name, emoji, is_enabled, fixed_cost,
|
||||
profit_multiplier, dynamic_unit_cost, billing_mode,
|
||||
price_cap, model_pricing_json, description
|
||||
FROM feature_pricing_configs
|
||||
""")
|
||||
rows = session.execute(sql).mappings().all()
|
||||
configs: dict[str, FeatureConfig] = {}
|
||||
for row in rows:
|
||||
key = str(row["feature_key"] or "").strip()
|
||||
if not key:
|
||||
continue
|
||||
configs[key] = FeatureConfig(
|
||||
feature_key=key,
|
||||
name=str(row["name"] or key),
|
||||
emoji=str(row["emoji"] or ""),
|
||||
is_enabled=bool(row["is_enabled"]),
|
||||
fixed_cost=_to_float(row["fixed_cost"]),
|
||||
profit_multiplier=_to_float(row["profit_multiplier"], 1.0),
|
||||
dynamic_unit_cost=_to_float(row["dynamic_unit_cost"]),
|
||||
billing_mode=str(row["billing_mode"] or "model_based"),
|
||||
price_cap=_to_float(row["price_cap"]),
|
||||
model_pricing=_parse_model_pricing(row["model_pricing_json"]),
|
||||
description=str(row["description"] or ""),
|
||||
)
|
||||
return configs
|
||||
finally:
|
||||
if session is not None:
|
||||
try:
|
||||
session.close()
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
|
||||
def _get_cache() -> dict[str, FeatureConfig]:
|
||||
"""TTL 内返回缓存,否则重新 load;DB 异常/表不存在时返回内置兜底配置。"""
|
||||
global _cache
|
||||
now = time.monotonic()
|
||||
with _lock:
|
||||
if _cache is not None and now - _cache[0] < CACHE_TTL_SECONDS:
|
||||
return _cache[1]
|
||||
|
||||
try:
|
||||
loaded = _load_all()
|
||||
except Exception: # noqa: BLE001 - 表不存在/DB 不可用时静默回落
|
||||
logger.info("feature_pricing_configs 读取失败,使用内置兜底配置", exc_info=True)
|
||||
return _fallback_configs()
|
||||
|
||||
# DB 可用但表为空:同样回落兜底(保证爆款现状不被改变)
|
||||
if not loaded:
|
||||
fallback = _fallback_configs()
|
||||
with _lock:
|
||||
_cache = (now, fallback)
|
||||
return fallback
|
||||
|
||||
# 以兜底为底(DB 未配置的 feature_key 仍有兜底),DB 行覆盖
|
||||
merged = _fallback_configs()
|
||||
merged.update(loaded)
|
||||
with _lock:
|
||||
_cache = (now, merged)
|
||||
return merged
|
||||
|
||||
|
||||
def get_feature_config(feature_key: str) -> Optional[FeatureConfig]:
|
||||
"""获取指定功能配置,未知 key 返回 None。"""
|
||||
key = str(feature_key or "").strip()
|
||||
if not key:
|
||||
return None
|
||||
return _get_cache().get(key)
|
||||
|
||||
|
||||
def _global_points_enabled() -> bool:
|
||||
"""全局积分总开关(兼容 api / worker 运行时),取不到时默认关闭。"""
|
||||
try:
|
||||
from packages.shared import get_shared_settings
|
||||
|
||||
return bool(get_shared_settings().points_enabled)
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
try:
|
||||
from app.config import settings
|
||||
|
||||
return bool(getattr(settings, "points_enabled", False))
|
||||
except Exception: # noqa: BLE001
|
||||
return False
|
||||
|
||||
|
||||
def is_feature_enabled(feature_key: str) -> bool:
|
||||
"""功能是否启用并扣费:全局 points_enabled AND 功能 is_enabled。"""
|
||||
cfg = get_feature_config(feature_key)
|
||||
if cfg is None:
|
||||
return False
|
||||
return bool(cfg.is_enabled) and _global_points_enabled()
|
||||
|
||||
|
||||
def calculate_price(feature_key: str, dynamic_cost: float = 0.0) -> tuple[float, dict]:
|
||||
"""按公式计算最终积分并返回明细。
|
||||
|
||||
price = (dynamic_cost + fixed_cost) × profit_multiplier
|
||||
price_cap > 0 时封顶(取 min)。
|
||||
功能未启用 → (0.0, breakdown{is_enabled: False, charged: False})。
|
||||
"""
|
||||
cfg = get_feature_config(feature_key)
|
||||
dynamic = max(0.0, _to_float(dynamic_cost))
|
||||
if cfg is None or not cfg.is_enabled:
|
||||
return 0.0, {
|
||||
"feature_key": feature_key,
|
||||
"is_enabled": False,
|
||||
"charged": False,
|
||||
"dynamic_cost": dynamic,
|
||||
"fixed_cost": 0.0,
|
||||
"profit_multiplier": 1.0,
|
||||
"price_cap": 0.0,
|
||||
"final_price": 0.0,
|
||||
}
|
||||
|
||||
fixed = max(0.0, cfg.fixed_cost)
|
||||
multiplier = cfg.profit_multiplier if cfg.profit_multiplier > 0 else 1.0
|
||||
raw_price = (dynamic + fixed) * multiplier
|
||||
cap = cfg.price_cap if cfg.price_cap and cfg.price_cap > 0 else 0.0
|
||||
final_price = min(raw_price, cap) if cap else raw_price
|
||||
final_price = round(float(final_price), 2)
|
||||
breakdown = {
|
||||
"feature_key": cfg.feature_key,
|
||||
"is_enabled": True,
|
||||
"charged": True,
|
||||
"dynamic_cost": round(dynamic, 4),
|
||||
"fixed_cost": float(fixed),
|
||||
"profit_multiplier": float(multiplier),
|
||||
"price_cap": float(cap),
|
||||
"raw_price": round(float(raw_price), 4),
|
||||
"final_price": final_price,
|
||||
}
|
||||
return final_price, breakdown
|
||||
|
||||
|
||||
def lookup_model_price(
|
||||
model_pricing: dict,
|
||||
model_key: str,
|
||||
resolution: str,
|
||||
has_video_input: bool,
|
||||
) -> Optional[float]:
|
||||
"""从 model_pricing dict 取模型单价,兼容两种常见 JSON 结构。
|
||||
|
||||
1. 嵌套:{model: {resolution: {"true"/"false": price}}}
|
||||
(内层 bool key 也兼容直接 bool / 省略)
|
||||
2. 扁平:{"model|resolution|true_or_false": price}
|
||||
(分隔符支持 | / : / , / 空格;bool 段可省略)
|
||||
取不到返回 None。
|
||||
"""
|
||||
if not isinstance(model_pricing, dict):
|
||||
return None
|
||||
model = str(model_key or "").strip()
|
||||
res = str(resolution or "").strip()
|
||||
flag = "true" if has_video_input else "false"
|
||||
|
||||
# 1. 嵌套
|
||||
model_node = model_pricing.get(model)
|
||||
if isinstance(model_node, dict):
|
||||
res_node = model_node.get(res)
|
||||
if isinstance(res_node, dict):
|
||||
# 精确 bool key 命中才返回;不做“只有一个值就取”的模糊匹配
|
||||
# (否则缺失 true 时会错误地取到 false 价,破坏旧版回落规则)
|
||||
if flag in res_node:
|
||||
return _to_float(res_node[flag]) if res_node[flag] is not None else None
|
||||
if has_video_input in res_node:
|
||||
val = res_node[has_video_input]
|
||||
return _to_float(val) if val is not None else None
|
||||
elif isinstance(res_node, (int, float)):
|
||||
return float(res_node)
|
||||
|
||||
# 2. 扁平
|
||||
for sep in ("|", ":", ",", " "):
|
||||
for key in (
|
||||
f"{model}{sep}{res}{sep}{flag}",
|
||||
f"{model}{sep}{res}",
|
||||
):
|
||||
if key in model_pricing:
|
||||
value = model_pricing[key]
|
||||
return _to_float(value) if value is not None else None
|
||||
return None
|
||||
|
||||
|
||||
def refresh_feature_configs() -> None:
|
||||
"""清空缓存(下次读取重新 load DB;测试/admin 改配置后可手动调)。"""
|
||||
global _cache
|
||||
with _lock:
|
||||
_cache = None
|
||||
@@ -2,17 +2,21 @@
|
||||
|
||||
v1.6.1: 按产品决策,智能混剪/AI数字人/AI配音/抖音解析/改写/标题/封面 全部免费,
|
||||
仅保留声音克隆合成(voice_clone_synth)的扣点逻辑;声音克隆训练保持免费。
|
||||
爆款视频(viral_video)走动态定价,见本文件 VIRAL_VIDEO_MODEL_PRICES + calculate_viral_video_credits。
|
||||
爆款视频(viral_video)走动态定价,计费参数 DB 化(feature_pricing_configs,
|
||||
见 feature_pricing_service),calculate_viral_video_credits 从配置读取单价/
|
||||
固定成本/利润系数/封顶,DB 不可用时回落兜底配置。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
# ============ 爆款视频动态定价 (#2151) ============
|
||||
# key = (model_id, resolution, has_video_input),单位:
|
||||
# - billing_mode=token: 元/百万tokens(输出)
|
||||
# - billing_mode=per_second: 元/秒(视频时长)
|
||||
from packages.domain import feature_pricing_service
|
||||
|
||||
# ============ 爆款视频动态定价 ============
|
||||
# 单价/固定成本/利润系数已 DB 化(feature_pricing_configs,feature_key=viral_video),
|
||||
# 由 feature_pricing_service 读取(300s 缓存),DB 不可用时回落内置兜底配置。
|
||||
# 以下三个常量仅为向后兼容保留(旧引用方/兜底场景),值取自兜底配置。
|
||||
VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = {
|
||||
("seedance-2.5", "480p", False): 70.0,
|
||||
("seedance-2.5", "720p", False): 70.0,
|
||||
@@ -33,9 +37,9 @@ VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = {
|
||||
("wan-3.0", "1080p", False): 1.2,
|
||||
}
|
||||
|
||||
# 固定成本(元):VLM 分析 + LLM 文案 + TTS + OSS + 服务器
|
||||
# 固定成本(元):VLM 分析 + LLM 文案 + TTS + OSS + 服务器(兜底默认值)
|
||||
VIRAL_VIDEO_FIXED_COST = 0.15
|
||||
# 利润系数
|
||||
# 利润系数(兜底默认值)
|
||||
VIRAL_VIDEO_PROFIT_MULTIPLIER = 1.3
|
||||
# Seedance 输出帧率
|
||||
VIRAL_VIDEO_FPS = 24
|
||||
@@ -222,17 +226,22 @@ def calculate_viral_video_credits_with_breakdown(
|
||||
) -> tuple[float, dict]:
|
||||
"""计算爆款视频所需积分(1 积分 = 1 元),并返回计费公式明细。
|
||||
|
||||
单价/固定成本/利润系数/封顶从 feature_pricing_configs(viral_video)读取;
|
||||
DB 不可用时回落与现状一致的内置兜底配置。
|
||||
|
||||
公式:
|
||||
tokens = duration * width * height * fps / 1024
|
||||
video_cost = tokens / 1_000_000 * model_token_price
|
||||
total = round((video_cost + fixed_cost) * profit_multiplier, 2)
|
||||
price_cap > 0 时封顶取 min
|
||||
若传入 actual_tokens 则用它替代计算值。
|
||||
|
||||
Returns:
|
||||
(credits, breakdown) 二元组:
|
||||
- credits: 四舍五入保留两位小数的最终积分
|
||||
- breakdown: dict,包含 tokens / video_cost / fixed_cost / profit_multiplier /
|
||||
model_price / width / height / fps 字段,便于前端展示计费明细。
|
||||
model_price / width / height / fps / feature_enabled / charged / price_cap
|
||||
字段,便于前端展示计费明细。功能关闭时 credits=0、charged=False。
|
||||
"""
|
||||
w = max(1, int(width or 1))
|
||||
h = max(1, int(height or 1))
|
||||
@@ -242,11 +251,36 @@ def calculate_viral_video_credits_with_breakdown(
|
||||
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)
|
||||
dur = max(1, int(duration_seconds or 15))
|
||||
|
||||
# ── 从 DB 配置(兜底内置)取计费参数 ──
|
||||
feature_cfg = feature_pricing_service.get_feature_config("viral_video")
|
||||
# 注意:此处 feature_enabled 只表示“功能自身开关”,不并入全局 points_enabled
|
||||
# 总开关(保持与旧版计费函数行为一致:价格照常计算)。全局总开关由业务层
|
||||
# (route/worker)通过 feature_pricing_service.is_feature_enabled 统一把关。
|
||||
feature_enabled = bool(feature_cfg.is_enabled) if feature_cfg is not None else True
|
||||
model_pricing = feature_cfg.model_pricing if feature_cfg is not None else {}
|
||||
fixed_cost = float(feature_cfg.fixed_cost) if feature_cfg is not None else float(VIRAL_VIDEO_FIXED_COST)
|
||||
multiplier = (
|
||||
float(feature_cfg.profit_multiplier)
|
||||
if feature_cfg is not None and feature_cfg.profit_multiplier > 0
|
||||
else float(VIRAL_VIDEO_PROFIT_MULTIPLIER)
|
||||
)
|
||||
price_cap = float(feature_cfg.price_cap) if feature_cfg is not None else 0.0
|
||||
|
||||
# 单价:优先配置 dict;复刻旧版回落规则——精确 key 取不到时,回落
|
||||
# seedance-2.5 同分辨率 False 单价;最终兜底 70.0。
|
||||
price = feature_pricing_service.lookup_model_price(model_pricing, prefix, res_key, bool(has_video_input))
|
||||
if price is None:
|
||||
# 配置表未命中:先尝试配置里的 seedance-2.5/False
|
||||
if prefix != "seedance-2.5" or bool(has_video_input):
|
||||
price = feature_pricing_service.lookup_model_price(model_pricing, "seedance-2.5", res_key, False)
|
||||
if price is None:
|
||||
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)
|
||||
@@ -259,13 +293,40 @@ def calculate_viral_video_credits_with_breakdown(
|
||||
video_cost = tokens / 1_000_000.0 * float(price)
|
||||
billing_unit = "token"
|
||||
|
||||
total = (video_cost + VIRAL_VIDEO_FIXED_COST) * VIRAL_VIDEO_PROFIT_MULTIPLIER
|
||||
if not feature_enabled:
|
||||
# 功能关闭(is_enabled=false 或全局 points 关闭):不扣费,明细照旧返回
|
||||
credits = 0.0
|
||||
raw_total = (video_cost + fixed_cost) * multiplier
|
||||
breakdown = {
|
||||
"tokens": float(tokens),
|
||||
"video_cost": float(video_cost),
|
||||
"fixed_cost": float(fixed_cost),
|
||||
"profit_multiplier": float(multiplier),
|
||||
"price_cap": float(price_cap or 0.0),
|
||||
"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,
|
||||
"feature_enabled": False,
|
||||
"charged": False,
|
||||
"raw_price": round(float(raw_total), 4),
|
||||
}
|
||||
return credits, breakdown
|
||||
|
||||
total = (video_cost + fixed_cost) * multiplier
|
||||
if price_cap and price_cap > 0:
|
||||
total = min(total, price_cap)
|
||||
credits = round(float(total), 2)
|
||||
breakdown = {
|
||||
"tokens": float(tokens),
|
||||
"video_cost": float(video_cost),
|
||||
"fixed_cost": float(VIRAL_VIDEO_FIXED_COST),
|
||||
"profit_multiplier": float(VIRAL_VIDEO_PROFIT_MULTIPLIER),
|
||||
"fixed_cost": float(fixed_cost),
|
||||
"profit_multiplier": float(multiplier),
|
||||
"price_cap": float(price_cap or 0.0),
|
||||
"model_price": float(price),
|
||||
"model_key": prefix,
|
||||
"billing_mode": billing,
|
||||
@@ -274,6 +335,8 @@ def calculate_viral_video_credits_with_breakdown(
|
||||
"height": int(h),
|
||||
"fps": int(effective_fps),
|
||||
"duration": dur,
|
||||
"feature_enabled": True,
|
||||
"charged": True,
|
||||
}
|
||||
return credits, breakdown
|
||||
|
||||
|
||||
@@ -291,6 +291,7 @@ class DoubaoClient:
|
||||
data.get("usage", {}).get("completion_tokens", 0),
|
||||
_elapsed,
|
||||
attempt + 1,
|
||||
_req_timeout,
|
||||
)
|
||||
return content.strip()
|
||||
except Exception as e:
|
||||
|
||||
+221
@@ -0,0 +1,221 @@
|
||||
"""功能计费改造测试:爆款读配置、对口型/智能剪辑预扣逻辑。
|
||||
|
||||
策略:
|
||||
- 爆款:通过修改缓存中的 FeatureConfig(multiplier/model_pricing)验证价格随配置变化
|
||||
- lip_sync / smart_edit:直接测 LipsyncService 的预扣/结算/退款辅助方法,
|
||||
PointsService 用 mock,避免依赖真实积分账户。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain import feature_pricing_service as fps
|
||||
from packages.domain.feature_pricing_service import FeatureConfig, refresh_feature_configs
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_cache():
|
||||
refresh_feature_configs()
|
||||
yield
|
||||
refresh_feature_configs()
|
||||
|
||||
|
||||
def _seed_cache(configs: dict) -> None:
|
||||
import time
|
||||
|
||||
fps._cache = (time.monotonic(), configs)
|
||||
|
||||
|
||||
class TestViralVideoReadsConfig:
|
||||
def test_multiplier_change_changes_price(self):
|
||||
"""配置里 multiplier 改大后,爆款价格随之变大(证明不再读死常量)。"""
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
# 基线兜底
|
||||
base = calculate_viral_video_credits(15, 1280, 720)
|
||||
assert base == 29.68
|
||||
|
||||
fallback = fps._fallback_configs()
|
||||
vv = fallback["viral_video"]
|
||||
vv.profit_multiplier = 2.0
|
||||
_seed_cache(fallback)
|
||||
|
||||
changed = calculate_viral_video_credits(15, 1280, 720)
|
||||
assert changed > base
|
||||
# 精确校验:video_cost 相同,仅系数从 1.3 → 2.0
|
||||
_, bd = __import__(
|
||||
"packages.domain.points_rules", fromlist=["calculate_viral_video_credits_with_breakdown"]
|
||||
).calculate_viral_video_credits_with_breakdown(15, 1280, 720)
|
||||
assert bd["profit_multiplier"] == 2.0
|
||||
|
||||
def test_model_price_from_config(self):
|
||||
"""model_pricing 改单价后,token 成本按新单价计算。"""
|
||||
from packages.domain.points_rules import calculate_viral_video_credits_with_breakdown
|
||||
|
||||
fallback = fps._fallback_configs()
|
||||
vv = fallback["viral_video"]
|
||||
# seedance-2.5/720p/false 从 70 改成 100
|
||||
vv.model_pricing["seedance-2.5"]["720p"]["false"] = 100.0
|
||||
_seed_cache(fallback)
|
||||
|
||||
_, bd = calculate_viral_video_credits_with_breakdown(15, 1280, 720)
|
||||
assert bd["model_price"] == 100.0
|
||||
|
||||
def test_price_cap_from_config(self):
|
||||
from packages.domain.points_rules import calculate_viral_video_credits_with_breakdown
|
||||
|
||||
fallback = fps._fallback_configs()
|
||||
vv = fallback["viral_video"]
|
||||
vv.price_cap = 5.0
|
||||
_seed_cache(fallback)
|
||||
|
||||
credits, bd = calculate_viral_video_credits_with_breakdown(15, 1280, 720)
|
||||
assert credits == 5.0
|
||||
assert bd["price_cap"] == 5.0
|
||||
|
||||
def test_disabled_feature_returns_zero_credits(self):
|
||||
"""功能 is_enabled=false 时计费函数返回 0(纯计费层语义)。"""
|
||||
from packages.domain.points_rules import calculate_viral_video_credits_with_breakdown
|
||||
|
||||
fallback = fps._fallback_configs()
|
||||
fallback["viral_video"].is_enabled = False
|
||||
_seed_cache(fallback)
|
||||
|
||||
credits, bd = calculate_viral_video_credits_with_breakdown(15, 1280, 720)
|
||||
assert credits == 0.0
|
||||
assert bd["feature_enabled"] is False
|
||||
assert bd["charged"] is False
|
||||
|
||||
|
||||
class TestLipSyncPricing:
|
||||
def _make_service(self):
|
||||
from app.services.lipsync_service import LipsyncService
|
||||
|
||||
svc = LipsyncService.__new__(LipsyncService)
|
||||
svc.db = MagicMock()
|
||||
return svc
|
||||
|
||||
def _lip_cfg(self, **kw):
|
||||
base = dict(
|
||||
feature_key="lip_sync",
|
||||
name="对口型",
|
||||
is_enabled=True,
|
||||
fixed_cost=0.1,
|
||||
profit_multiplier=1.0,
|
||||
dynamic_unit_cost=0.05,
|
||||
billing_mode="per_second",
|
||||
price_cap=0.0,
|
||||
model_pricing={},
|
||||
description="",
|
||||
)
|
||||
base.update(kw)
|
||||
return FeatureConfig(**base)
|
||||
|
||||
def test_estimate_duration_from_script(self):
|
||||
svc = self._make_service()
|
||||
# 10 个字 / 5 = 2 秒,下限 1
|
||||
assert svc._estimate_duration(script_text="一二三四五六七八九十") == 2.0
|
||||
# 无任何信息 → 默认 10 秒
|
||||
assert svc._estimate_duration() == 10.0
|
||||
|
||||
def test_calculate_lipsync_price_per_second(self):
|
||||
_seed_cache({"lip_sync": self._lip_cfg()})
|
||||
price, bd = fps.calculate_price("lip_sync", dynamic_cost=20.0 * 0.05)
|
||||
# dynamic 1.0 + fixed 0.1 = 1.1
|
||||
assert price == 1.1
|
||||
assert bd["charged"] is True
|
||||
|
||||
def test_settle_refunds_overcharge(self):
|
||||
"""实际时长短 → 只退不补,退还差额。"""
|
||||
svc = self._make_service()
|
||||
_seed_cache({"lip_sync": self._lip_cfg()})
|
||||
|
||||
job = MagicMock()
|
||||
job.credits_prepaid = 2.0
|
||||
job.credits_cost = 0.0 # 未结算
|
||||
job.user_id = "u1"
|
||||
job.credits_transaction_id = "txn-old"
|
||||
|
||||
with patch("packages.domain.points_service.PointsService") as MockPS:
|
||||
inst = MockPS.return_value
|
||||
inst.refund_points.return_value = {"success": True}
|
||||
svc._settle_lip_sync(job, actual_duration=10.0)
|
||||
|
||||
# final: (10*0.05 + 0.1)*1.0 = 0.6;退 2.0-0.6=1.4
|
||||
assert round(job.credits_cost, 2) == 0.6
|
||||
inst.refund_points.assert_called_once()
|
||||
kwargs = inst.refund_points.call_args.kwargs
|
||||
assert kwargs["amount"] == 1.4
|
||||
|
||||
def test_settle_no_refund_when_longer(self):
|
||||
"""首期只退不补:实际更贵不补扣。"""
|
||||
svc = self._make_service()
|
||||
_seed_cache({"lip_sync": self._lip_cfg()})
|
||||
|
||||
job = MagicMock()
|
||||
job.credits_prepaid = 0.5
|
||||
job.credits_cost = 0.0
|
||||
|
||||
with patch("packages.domain.points_service.PointsService") as MockPS:
|
||||
inst = MockPS.return_value
|
||||
svc._settle_lip_sync(job, actual_duration=60.0)
|
||||
|
||||
assert round(job.credits_cost, 2) > 0.5
|
||||
inst.refund_points.assert_not_called()
|
||||
|
||||
def test_refund_on_failure_full(self):
|
||||
svc = self._make_service()
|
||||
job = MagicMock()
|
||||
job.credits_prepaid = 3.0
|
||||
job.credits_cost = 0.0
|
||||
job.user_id = "u1"
|
||||
job.credits_transaction_id = "t1"
|
||||
|
||||
with patch("packages.domain.points_service.PointsService") as MockPS:
|
||||
inst = MockPS.return_value
|
||||
inst.refund_points.return_value = {"success": True}
|
||||
svc._refund_lip_sync(job)
|
||||
|
||||
kwargs = inst.refund_points.call_args.kwargs
|
||||
assert kwargs["amount"] == 3.0
|
||||
|
||||
|
||||
class TestSmartEditFixedPrice:
|
||||
def test_fixed_price_formula(self):
|
||||
"""首期固定价:dynamic=0,price=fixed*multiplier,cap 封顶。"""
|
||||
cfg = FeatureConfig(
|
||||
feature_key="smart_edit",
|
||||
name="智能剪辑",
|
||||
is_enabled=True,
|
||||
fixed_cost=2.0,
|
||||
profit_multiplier=1.5,
|
||||
billing_mode="model_based",
|
||||
price_cap=0.0,
|
||||
)
|
||||
_seed_cache({"smart_edit": cfg})
|
||||
price, bd = fps.calculate_price("smart_edit", dynamic_cost=0.0)
|
||||
# (0+2)*1.5 = 3.0
|
||||
assert price == 3.0
|
||||
assert bd["dynamic_cost"] == 0.0
|
||||
|
||||
def test_fixed_price_with_cap(self):
|
||||
cfg = FeatureConfig(
|
||||
feature_key="smart_edit",
|
||||
is_enabled=True,
|
||||
fixed_cost=10.0,
|
||||
profit_multiplier=2.0,
|
||||
price_cap=8.0,
|
||||
)
|
||||
_seed_cache({"smart_edit": cfg})
|
||||
price, _ = fps.calculate_price("smart_edit", dynamic_cost=0.0)
|
||||
assert price == 8.0
|
||||
|
||||
def test_disabled_smart_edit_free(self):
|
||||
cfg = FeatureConfig(feature_key="smart_edit", is_enabled=False, fixed_cost=2.0)
|
||||
_seed_cache({"smart_edit": cfg})
|
||||
price, bd = fps.calculate_price("smart_edit", dynamic_cost=0.0)
|
||||
assert price == 0.0
|
||||
assert bd["charged"] is False
|
||||
Executable
+235
@@ -0,0 +1,235 @@
|
||||
"""feature_pricing_service 单元测试。
|
||||
|
||||
覆盖:
|
||||
- 300s TTL 内存缓存(命中不重复 load / 过期重新 load / refresh 强制刷新)
|
||||
- calculate_price 公式 (dynamic+fixed)*multiplier、price_cap 封顶、round
|
||||
- disabled / 未知 key 返回 0
|
||||
- DB 异常 / 空表 → 内置兜底配置(爆款启用且价格与现状一致)
|
||||
- lookup_model_price 嵌套/扁平结构与旧版回落语义
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain import feature_pricing_service as fps
|
||||
from packages.domain.feature_pricing_service import (
|
||||
CACHE_TTL_SECONDS,
|
||||
FeatureConfig,
|
||||
calculate_price,
|
||||
get_feature_config,
|
||||
is_feature_enabled,
|
||||
lookup_model_price,
|
||||
refresh_feature_configs,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_cache():
|
||||
"""每个用例前后清空模块缓存,避免相互污染。"""
|
||||
refresh_feature_configs()
|
||||
yield
|
||||
refresh_feature_configs()
|
||||
|
||||
|
||||
def _cfg(key="x", **kw) -> FeatureConfig:
|
||||
base = dict(
|
||||
feature_key=key,
|
||||
name=key,
|
||||
is_enabled=True,
|
||||
fixed_cost=0.2,
|
||||
profit_multiplier=2.0,
|
||||
dynamic_unit_cost=0.0,
|
||||
billing_mode="per_second",
|
||||
price_cap=0.0,
|
||||
model_pricing={},
|
||||
description="",
|
||||
)
|
||||
base.update(kw)
|
||||
return FeatureConfig(**base)
|
||||
|
||||
|
||||
class TestCacheTTL:
|
||||
def test_cache_hit_avoids_reload(self, monkeypatch):
|
||||
"""TTL 内第二次读取不再调 _load_all。"""
|
||||
calls = {"n": 0}
|
||||
|
||||
def fake_load():
|
||||
calls["n"] += 1
|
||||
return {"x": _cfg()}
|
||||
|
||||
monkeypatch.setattr(fps, "_load_all", fake_load)
|
||||
get_feature_config("x")
|
||||
get_feature_config("x")
|
||||
get_feature_config("x")
|
||||
assert calls["n"] == 1
|
||||
|
||||
def test_expired_cache_reloads(self, monkeypatch):
|
||||
"""超过 TTL 后重新 load。"""
|
||||
calls = {"n": 0}
|
||||
|
||||
def fake_load():
|
||||
calls["n"] += 1
|
||||
return {"x": _cfg()}
|
||||
|
||||
monkeypatch.setattr(fps, "_load_all", fake_load)
|
||||
get_feature_config("x")
|
||||
assert calls["n"] == 1
|
||||
|
||||
# 把缓存时间戳回拨到 TTL 之前
|
||||
ts, data = fps._cache
|
||||
fps._cache = (ts - CACHE_TTL_SECONDS - 1, data)
|
||||
get_feature_config("x")
|
||||
assert calls["n"] == 2
|
||||
|
||||
def test_refresh_forces_reload(self, monkeypatch):
|
||||
calls = {"n": 0}
|
||||
|
||||
def fake_load():
|
||||
calls["n"] += 1
|
||||
return {"x": _cfg()}
|
||||
|
||||
monkeypatch.setattr(fps, "_load_all", fake_load)
|
||||
get_feature_config("x")
|
||||
refresh_feature_configs()
|
||||
get_feature_config("x")
|
||||
assert calls["n"] == 2
|
||||
|
||||
def test_ttl_constant_is_300(self):
|
||||
assert CACHE_TTL_SECONDS == 300.0
|
||||
|
||||
|
||||
class TestCalculatePrice:
|
||||
def test_basic_formula(self, monkeypatch):
|
||||
# (dynamic 1.0 + fixed 0.2) * 2.0 = 2.4
|
||||
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(dynamic_unit_cost=1.0)})
|
||||
price, bd = calculate_price("x", dynamic_cost=1.0)
|
||||
assert price == 2.4
|
||||
assert bd["dynamic_cost"] == 1.0
|
||||
assert bd["fixed_cost"] == 0.2
|
||||
assert bd["profit_multiplier"] == 2.0
|
||||
assert bd["final_price"] == 2.4
|
||||
assert bd["charged"] is True
|
||||
|
||||
def test_price_cap_clamps(self, monkeypatch):
|
||||
# raw = (1+0.2)*2 = 2.4,cap=1.0 → 1.0
|
||||
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(price_cap=1.0)})
|
||||
price, bd = calculate_price("x", dynamic_cost=1.0)
|
||||
assert price == 1.0
|
||||
assert bd["price_cap"] == 1.0
|
||||
|
||||
def test_no_cap_keeps_raw(self, monkeypatch):
|
||||
# cap=0 视为不封顶
|
||||
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(price_cap=0.0)})
|
||||
price, _ = calculate_price("x", dynamic_cost=1.0)
|
||||
assert price == 2.4
|
||||
|
||||
def test_rounded_two_decimals(self, monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
fps,
|
||||
"_load_all",
|
||||
lambda: {"x": _cfg(fixed_cost=0.1, profit_multiplier=1.0)},
|
||||
)
|
||||
price, _ = calculate_price("x", dynamic_cost=1.0 / 3.0)
|
||||
# 0.3333... + 0.1 = 0.4333 → 0.43
|
||||
assert price == 0.43
|
||||
|
||||
def test_negative_dynamic_treated_as_zero(self, monkeypatch):
|
||||
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg()})
|
||||
price, _ = calculate_price("x", dynamic_cost=-5.0)
|
||||
# (0 + 0.2) * 2 = 0.4
|
||||
assert price == 0.4
|
||||
|
||||
def test_disabled_returns_zero(self, monkeypatch):
|
||||
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(is_enabled=False)})
|
||||
price, bd = calculate_price("x", dynamic_cost=1.0)
|
||||
assert price == 0.0
|
||||
assert bd["is_enabled"] is False
|
||||
assert bd["charged"] is False
|
||||
|
||||
def test_unknown_key_returns_zero(self, monkeypatch):
|
||||
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg()})
|
||||
price, bd = calculate_price("nope", dynamic_cost=1.0)
|
||||
assert price == 0.0
|
||||
assert bd["charged"] is False
|
||||
|
||||
|
||||
class TestDBFailureFallback:
|
||||
def test_load_exception_uses_fallback(self, monkeypatch):
|
||||
def boom():
|
||||
raise RuntimeError("table does not exist")
|
||||
|
||||
monkeypatch.setattr(fps, "_load_all", boom)
|
||||
cfg = get_feature_config("viral_video")
|
||||
assert cfg is not None
|
||||
assert cfg.is_enabled is True
|
||||
assert cfg.fixed_cost == 0.15
|
||||
assert cfg.profit_multiplier == 1.3
|
||||
|
||||
def test_empty_table_uses_fallback(self, monkeypatch):
|
||||
monkeypatch.setattr(fps, "_load_all", lambda: {})
|
||||
assert get_feature_config("viral_video").is_enabled is True
|
||||
assert get_feature_config("lip_sync").is_enabled is False
|
||||
assert get_feature_config("smart_edit").is_enabled is False
|
||||
|
||||
def test_fallback_viral_price_matches_current(self, monkeypatch):
|
||||
"""兜底爆款价格与旧硬编码现状一致:seedance-2.5/720p/false=70。"""
|
||||
monkeypatch.setattr(fps, "_load_all", lambda: {})
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
# 默认全局开关关闭,但纯计费函数价格照常算
|
||||
assert calculate_viral_video_credits(15, 1280, 720) == 29.68
|
||||
|
||||
def test_db_row_overrides_fallback(self, monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
fps,
|
||||
"_load_all",
|
||||
lambda: {"viral_video": _cfg("viral_video", fixed_cost=0.5, profit_multiplier=2.0, price_cap=50.0)},
|
||||
)
|
||||
cfg = get_feature_config("viral_video")
|
||||
assert cfg.fixed_cost == 0.5
|
||||
assert cfg.profit_multiplier == 2.0
|
||||
assert cfg.price_cap == 50.0
|
||||
|
||||
|
||||
class TestIsFeatureEnabled:
|
||||
def test_disabled_feature(self, monkeypatch):
|
||||
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(is_enabled=False)})
|
||||
assert is_feature_enabled("x") is False
|
||||
|
||||
def test_global_switch_off_blocks_enabled_feature(self, monkeypatch):
|
||||
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(is_enabled=True)})
|
||||
monkeypatch.setattr(fps, "_global_points_enabled", lambda: False)
|
||||
assert is_feature_enabled("x") is False
|
||||
|
||||
def test_both_switches_on(self, monkeypatch):
|
||||
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(is_enabled=True)})
|
||||
monkeypatch.setattr(fps, "_global_points_enabled", lambda: True)
|
||||
assert is_feature_enabled("x") is True
|
||||
|
||||
|
||||
class TestLookupModelPrice:
|
||||
NESTED = {
|
||||
"seedance-2.5": {
|
||||
"720p": {"false": 70.0, "true": 42.0},
|
||||
},
|
||||
"wan-3.0": {"480p": {"false": 0.3}},
|
||||
}
|
||||
|
||||
def test_nested_exact_hit(self):
|
||||
assert lookup_model_price(self.NESTED, "seedance-2.5", "720p", False) == 70.0
|
||||
assert lookup_model_price(self.NESTED, "seedance-2.5", "720p", True) == 42.0
|
||||
|
||||
def test_missing_bool_key_returns_none(self):
|
||||
# wan-3.0/480p 只有 false,请求 true → None(由调用方回落)
|
||||
assert lookup_model_price(self.NESTED, "wan-3.0", "480p", True) is None
|
||||
|
||||
def test_unknown_model_returns_none(self):
|
||||
assert lookup_model_price(self.NESTED, "nope", "720p", False) is None
|
||||
|
||||
def test_flat_structure(self):
|
||||
flat = {"m|720p|false": 12.5}
|
||||
assert lookup_model_price(flat, "m", "720p", False) == 12.5
|
||||
assert lookup_model_price(flat, "m", "720p", True) is None
|
||||
@@ -97,21 +97,43 @@ def invalidate_loader_cache():
|
||||
|
||||
|
||||
class TestImageAnalysisWiring:
|
||||
def test_uses_loader_template_and_xml_parse(self, job):
|
||||
def test_step_image_analysis_uses_v2_batch_path(self, job):
|
||||
"""#2200/#2207 后图片分析走 V2 批处理(OCR+lite JSON 并行),
|
||||
_step_image_analysis 归一化 URL 后调用 analyze_images_v2。"""
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
with patch("packages.shared.ai_service.call_vision", return_value=IMAGE_XML) as mock_v:
|
||||
result = vv._analyze_single_image(0, "https://img/1.jpg", "vlm-lite", 15)
|
||||
fake_product = {
|
||||
"name": "lipstick",
|
||||
"brand": "品牌X",
|
||||
"category": "唇部彩妆",
|
||||
"key_features": ["显白", "持久"],
|
||||
"text_on_package": ["品牌X", "211"],
|
||||
"_source": "v2",
|
||||
}
|
||||
with patch.object(vv, "_normalize_image_url", side_effect=lambda raw, idx: raw):
|
||||
with patch(
|
||||
"worker_app.tasks.vision.analyze_images_v2",
|
||||
return_value=[fake_product, fake_product],
|
||||
create=True,
|
||||
) as mock_v2:
|
||||
result = vv._step_image_analysis(job)
|
||||
|
||||
mock_v.assert_called_once()
|
||||
# 验证调用时传入了 system_prompt(说明走了 loader 渲染的模板)
|
||||
call_kwargs = mock_v.call_args.kwargs
|
||||
assert "system_prompt" in call_kwargs and call_kwargs["system_prompt"]
|
||||
# 结果包含从 XML 解析出的产品信息
|
||||
assert result["name"] == "lipstick"
|
||||
assert result["brand"] == "品牌X"
|
||||
assert "显白" in result["key_features"]
|
||||
assert result["text_on_package"] == ["品牌X", "211"]
|
||||
mock_v2.assert_called_once()
|
||||
# 传入的是归一化后的图片 URL 列表
|
||||
assert mock_v2.call_args.args[0] == job.images
|
||||
products = result["products"]
|
||||
assert len(products) == 2
|
||||
assert products[0]["name"] == "lipstick"
|
||||
assert products[0]["brand"] == "品牌X"
|
||||
assert "显白" in products[0]["key_features"]
|
||||
assert products[0]["text_on_package"] == ["品牌X", "211"]
|
||||
|
||||
def test_step_image_analysis_empty_images(self, job):
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
job.images = []
|
||||
result = vv._step_image_analysis(job)
|
||||
assert result == {"products": []}
|
||||
|
||||
|
||||
# ── 2) 意图解析走模板 ───────────────────────────────────────────────
|
||||
@@ -252,18 +274,29 @@ class TestEndToEndLoaderUsed:
|
||||
called_types.append(prompt_type)
|
||||
return real_get(prompt_type, **kwargs)
|
||||
|
||||
v2_product = {
|
||||
"name": "lipstick",
|
||||
"brand": "品牌X",
|
||||
"key_features": ["显白", "持久"],
|
||||
}
|
||||
with (
|
||||
patch.object(pl, "get_template", side_effect=spy_get),
|
||||
patch("packages.shared.ai_service.call_vision", return_value=IMAGE_XML),
|
||||
patch.object(vv, "_normalize_image_url", side_effect=lambda raw, idx: raw),
|
||||
patch(
|
||||
"worker_app.tasks.vision.analyze_images_v2",
|
||||
return_value=[v2_product],
|
||||
create=True,
|
||||
),
|
||||
patch("packages.shared.ai_service.call_llm", return_value=INTENT_XML),
|
||||
):
|
||||
# 1) image
|
||||
img_res = vv._analyze_single_image(0, "https://img/1.jpg", "vlm", 15)
|
||||
# 2) intent
|
||||
# 1) image(V2 路径,不再经过 prompt_loader)
|
||||
img_step = vv._step_image_analysis(job)
|
||||
img_res = img_step["products"][0]
|
||||
# 2) intent(走 loader image_analysis? 否——intent_parsing 模板)
|
||||
intent_res = vv._step_intent_parsing(job, {"products": [img_res]})
|
||||
|
||||
# 前两步分别调用了 image_analysis 和 intent_parsing
|
||||
assert "image_analysis" in called_types
|
||||
# V2 图片分析不再调用 loader;意图解析调用 intent_parsing 模板
|
||||
assert "image_analysis" not in called_types
|
||||
assert "intent_parsing" in called_types
|
||||
|
||||
# script 和 review 单独验证(需要不同的 LLM 返回)
|
||||
|
||||
Executable
+291
@@ -0,0 +1,291 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""vision v4 prompt / assembler 单元测试:
|
||||
|
||||
- assembler 正确识别 v4 嵌套 schema 与旧扁平 schema
|
||||
- v4 product/person/store/other 四类输出组装出下游必出字段
|
||||
- 旧扁平 schema 行为不变
|
||||
- _prompt._resolve:DB 有 active prompt 时原样使用(不追加硬编码 schema);
|
||||
DB 无记录时回落到硬编码 JSON schema
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import types
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from worker_app.tasks.vision import _prompt, assembler
|
||||
|
||||
REQUIRED_KEYS = {
|
||||
"name",
|
||||
"brand",
|
||||
"category",
|
||||
"appearance",
|
||||
"packaging",
|
||||
"text_on_package",
|
||||
"key_features",
|
||||
"scene",
|
||||
"mood",
|
||||
"portrait_prompt",
|
||||
"summary",
|
||||
"_source",
|
||||
}
|
||||
|
||||
|
||||
# ---------- schema 识别 ----------
|
||||
|
||||
|
||||
def test_is_v4_schema_products_list() -> None:
|
||||
assert assembler._is_v4_schema({"type": "product", "products": []})
|
||||
|
||||
|
||||
def test_is_v4_schema_type_only() -> None:
|
||||
assert assembler._is_v4_schema({"type": "person"})
|
||||
|
||||
|
||||
def test_is_v4_schema_people_dict() -> None:
|
||||
assert assembler._is_v4_schema({"people": {"has_person": True}})
|
||||
|
||||
|
||||
def test_is_not_v4_schema_flat() -> None:
|
||||
assert not assembler._is_v4_schema({"has_person": True, "upper_wear": "T恤"})
|
||||
|
||||
|
||||
# ---------- v4 product ----------
|
||||
|
||||
V4_PRODUCT: dict[str, Any] = {
|
||||
"type": "product",
|
||||
"scene": "白色背景产品图",
|
||||
"mood": "清新专业",
|
||||
"style": "商业产品摄影",
|
||||
"colors": [{"hex": "#E60012", "name": "亮红色", "coverage": 0.6}],
|
||||
"visible_text": [{"text": "OMO奥妙除菌除螨", "position": "瓶身正面"}],
|
||||
"products": [
|
||||
{
|
||||
"product_name": "OMO奥妙除菌除螨洗衣液",
|
||||
"brand": "OMO奥妙",
|
||||
"category": "洗护",
|
||||
"package_type": "瓶装",
|
||||
"package_color": "亮红色瓶身",
|
||||
"cap_type": "透明翻盖式按压瓶口",
|
||||
"body_shape": "带侧面握持把手的竖款瓶身",
|
||||
"label_design": "瓶身印十字盾牌图案",
|
||||
"product_features": ["亮红色瓶装", "按压式瓶口", "十字盾牌标签"],
|
||||
"key_selling_points": ["天然除菌除螨"],
|
||||
"position": "main",
|
||||
}
|
||||
],
|
||||
"has_person": False,
|
||||
}
|
||||
|
||||
|
||||
def test_assemble_v4_product_fields() -> None:
|
||||
r = assembler.assemble_result(0, V4_PRODUCT, ["OMO奥妙"])
|
||||
assert REQUIRED_KEYS <= set(r.keys())
|
||||
assert r["name"] == "OMO奥妙除菌除螨洗衣液"
|
||||
assert r["brand"] == "OMO奥妙"
|
||||
assert r["category"] == "洗护"
|
||||
assert "瓶装" in r["packaging"]
|
||||
assert isinstance(r["key_features"], list) and r["key_features"]
|
||||
assert any("除菌" in str(t) for t in r["text_on_package"])
|
||||
assert len(r["portrait_prompt"]) >= 10
|
||||
assert r["_source"] == "v2_fast_json_v4"
|
||||
|
||||
|
||||
def test_assemble_v4_product_multi_selects_main() -> None:
|
||||
fj = {
|
||||
"type": "product",
|
||||
"products": [
|
||||
{"product_name": "次要商品", "brand": "B"},
|
||||
{"product_name": "主商品", "brand": "A", "position": "main"},
|
||||
],
|
||||
}
|
||||
r = assembler.assemble_result(1, fj, [])
|
||||
assert r["name"] == "主商品"
|
||||
|
||||
|
||||
# ---------- v4 person ----------
|
||||
|
||||
V4_PERSON: dict[str, Any] = {
|
||||
"type": "person",
|
||||
"scene": "户外街拍",
|
||||
"mood": "自信",
|
||||
"style": "街拍",
|
||||
"colors": [],
|
||||
"visible_text": [],
|
||||
"has_person": True,
|
||||
"gender": "女",
|
||||
"age_range": "青年",
|
||||
"upper_wear": "白色V领短袖T恤",
|
||||
"upper_color": "白色",
|
||||
"lower_wear": "黑色高腰阔腿裤",
|
||||
"lower_color": "黑色",
|
||||
"dress_color": None,
|
||||
"accessories": ["银色项链"],
|
||||
"hairstyle": "黑色长直发",
|
||||
"expression": "自信",
|
||||
"pose": "侧身站立",
|
||||
"outfit_style": "休闲日常",
|
||||
"portrait_prompt": (
|
||||
"一位年轻女性,身穿白色V领短袖T恤、黑色高腰阔腿裤,佩戴银色项链,"
|
||||
"黑色长直发,神情自信,侧身站立,休闲日常风格,城市街拍场景"
|
||||
),
|
||||
"products": [],
|
||||
}
|
||||
|
||||
|
||||
def test_assemble_v4_person() -> None:
|
||||
r = assembler.assemble_result(0, V4_PERSON, [])
|
||||
assert REQUIRED_KEYS <= set(r.keys())
|
||||
assert r["category"] == "服饰"
|
||||
assert "T恤" in r["name"]
|
||||
assert "阔腿裤" in r["name"]
|
||||
assert "年轻女性" in r["portrait_prompt"]
|
||||
assert "项链" in r["portrait_prompt"]
|
||||
assert isinstance(r["key_features"], list) and len(r["key_features"]) <= 6
|
||||
|
||||
|
||||
def test_assemble_v4_person_people_nested() -> None:
|
||||
fj = {"type": "person", "people": {**V4_PERSON, "has_person": True}}
|
||||
r = assembler.assemble_result(0, fj, [])
|
||||
assert r["category"] == "服饰"
|
||||
assert "年轻女性" in r["portrait_prompt"]
|
||||
|
||||
|
||||
# ---------- v4 store ----------
|
||||
|
||||
|
||||
def test_assemble_v4_store() -> None:
|
||||
fj = {
|
||||
"type": "store",
|
||||
"scene": "便利店内部",
|
||||
"mood": "日常便民",
|
||||
"style": "门店实拍",
|
||||
"store_type": "社区便利店",
|
||||
"store_layout": "纵深货架布局",
|
||||
"brand_signage": "全家FamilyMart",
|
||||
"visual_elements": ["红白主色调", "促销海报"],
|
||||
"product_categories_visible": ["饮料", "零食"],
|
||||
"promotion_elements": ["第二件半价海报"],
|
||||
"atmosphere": "亲民生活化",
|
||||
"has_person": False,
|
||||
}
|
||||
r = assembler.assemble_result(0, fj, [])
|
||||
assert REQUIRED_KEYS <= set(r.keys())
|
||||
assert r["name"] == "社区便利店"
|
||||
assert r["brand"] == "全家FamilyMart"
|
||||
assert r["category"] == "门店场景"
|
||||
assert any("饮料" in str(f) for f in r["key_features"])
|
||||
assert "门店实拍" in r["portrait_prompt"]
|
||||
|
||||
|
||||
# ---------- v4 other ----------
|
||||
|
||||
|
||||
def test_assemble_v4_other() -> None:
|
||||
fj = {"type": "other", "description": "海边日落风景", "scene": "海边", "mood": "宁静"}
|
||||
r = assembler.assemble_result(0, fj, [])
|
||||
assert REQUIRED_KEYS <= set(r.keys())
|
||||
assert r["name"] == "海边日落风景"
|
||||
assert r["category"] == "非产品图"
|
||||
|
||||
|
||||
# ---------- 旧扁平 schema 兼容 ----------
|
||||
|
||||
|
||||
def test_assemble_old_flat_person() -> None:
|
||||
fj = {
|
||||
"has_person": True,
|
||||
"gender": "男",
|
||||
"age_range": "中年",
|
||||
"upper_wear": "西装",
|
||||
"upper_color": "深灰色",
|
||||
"lower_wear": "西裤",
|
||||
"lower_color": "黑色",
|
||||
"accessories": ["手表"],
|
||||
"hairstyle": "短发",
|
||||
"expression": "严肃",
|
||||
"scene": "办公室",
|
||||
"style": "商务",
|
||||
"mood": "专业",
|
||||
}
|
||||
r = assembler.assemble_result(0, fj, [])
|
||||
assert REQUIRED_KEYS <= set(r.keys())
|
||||
assert "中年男性" in r["portrait_prompt"]
|
||||
assert r["_source"] == "v2_fast_json"
|
||||
|
||||
|
||||
def test_assemble_old_flat_product() -> None:
|
||||
fj = {
|
||||
"has_person": False,
|
||||
"product_name": "口红",
|
||||
"brand": "Dior",
|
||||
"category": "美妆",
|
||||
"colors": ["红色"],
|
||||
"scene": "通用",
|
||||
"style": "商业",
|
||||
"mood": "高级",
|
||||
}
|
||||
r = assembler.assemble_result(0, fj, ["Dior"])
|
||||
assert r["name"] == "口红"
|
||||
assert r["brand"] == "Dior"
|
||||
assert r["text_on_package"] == ["Dior"]
|
||||
|
||||
|
||||
def test_assemble_none_input() -> None:
|
||||
r = assembler.assemble_result(0, None, [])
|
||||
assert REQUIRED_KEYS <= set(r.keys())
|
||||
|
||||
|
||||
# ---------- _prompt 解析 ----------
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_prompt_cache() -> Any:
|
||||
_prompt.invalidate_cache()
|
||||
yield
|
||||
_prompt.invalidate_cache()
|
||||
|
||||
|
||||
def _fake_tpl(system_prompt: str = "v4 system prompt 只返回JSON") -> Any:
|
||||
return types.SimpleNamespace(
|
||||
system_prompt=system_prompt,
|
||||
user_prompt_template="分析 {image_count} 张图",
|
||||
version=4,
|
||||
)
|
||||
|
||||
|
||||
def test_resolve_uses_db_prompt_without_append(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(_prompt, "_load_db_template", lambda: _fake_tpl("DB_V4_PROMPT_XYZ"))
|
||||
sys_prompt, user_prompt = _prompt.resolve_fast_prompt()
|
||||
assert sys_prompt == "DB_V4_PROMPT_XYZ"
|
||||
assert "DB_V4_PROMPT_XYZ" not in _prompt._FAST_JSON_APPEND # sanity: 旧append是另一段文本
|
||||
assert "分析 1 张图" in user_prompt
|
||||
|
||||
|
||||
def test_resolve_pro_uses_db_prompt_without_append(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(_prompt, "_load_db_template", lambda: _fake_tpl("DB_V4_PRO_PROMPT"))
|
||||
sys_prompt, _ = _prompt.resolve_pro_prompt()
|
||||
assert sys_prompt == "DB_V4_PRO_PROMPT"
|
||||
assert "【输出格式要求】" not in sys_prompt
|
||||
|
||||
|
||||
def test_resolve_falls_back_when_no_db(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(_prompt, "_load_db_template", lambda: None)
|
||||
sys_prompt, user_prompt = _prompt.resolve_fast_prompt()
|
||||
assert sys_prompt == _prompt._FAST_JSON_SCHEMA
|
||||
assert user_prompt == _prompt.DEFAULT_FAST_USER
|
||||
|
||||
|
||||
def test_resolve_caches(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
calls = {"n": 0}
|
||||
|
||||
def _load() -> Any:
|
||||
calls["n"] += 1
|
||||
return _fake_tpl("CACHED_PROMPT")
|
||||
|
||||
monkeypatch.setattr(_prompt, "_load_db_template", _load)
|
||||
s1, _ = _prompt.resolve_fast_prompt()
|
||||
s2, _ = _prompt.resolve_fast_prompt()
|
||||
assert s1 == s2 == "CACHED_PROMPT"
|
||||
assert calls["n"] == 1
|
||||
Reference in New Issue
Block a user