Compare commits

..

3 Commits

Author SHA1 Message Date
build-ops 7843b92ef8 style(vision): rename ambiguous l/r vars to lb/rb (ruff E741)
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 3s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 54s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 1m45s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m19s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m2s
AI Code Review / AI Code Review (pull_request) Successful in 7m3s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 10m3s
CI/CD Pipeline / Validate - Style (pull_request) Successful in 11m0s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 12m44s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 10m45s
CI/CD Pipeline / Unit Tests (pull_request) Failing after 17m30s
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Security (pull_request) Has been cancelled
2026-10-05 19:22:03 +08:00
CI Bot 12f74b8d52 style: auto-format with black + isort + ruff + prettier [skip ci-format-check]
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 2s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 1m2s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 1m13s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m48s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m21s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 4m0s
CI/CD Pipeline / Validate - Style (pull_request) Failing after 4m18s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 5m0s
AI Code Review / AI Code Review (pull_request) Successful in 7m7s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 10m25s
CI/CD Pipeline / Unit Tests (pull_request) Failing after 12m4s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Failing after 2s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 8m53s
2026-10-05 11:11:29 +00:00
xiaoxia cdd343131e fix(vision): #2205 thinking参数互斥修复——只传thinking=disabled,去掉reasoning_effort
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 1s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 1s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 1m10s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m52s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 1m59s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 2m58s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 4m8s
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Security (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Style (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Has been cancelled
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
AI Code Review / AI Code Review (pull_request) Has been cancelled
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been cancelled
根因:上版同时传 thinking={type:disabled} + reasoning_effort=low,方舟返回400
Invalid combination of reasoning_effort and thinking type。降级重试分支
把两个参数都pop掉,模型回到默认thinking开启→响应9-12s超时,fast全败。

修复:
1. 只传 thinking={"type":"disabled"},去掉互斥的 reasoning_effort
2. 400降级只pop thinking,保留精简payload重试(不pop多个)
3. 先单独验证 lite 单次HTTP 200且reasoning_tokens=0再跑E2E,省一轮部署
2026-10-05 19:06:02 +08:00
20 changed files with 499 additions and 2553 deletions
@@ -1,61 +0,0 @@
"""功能计费积分字段(爆款/对口型/智能剪辑 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
+1 -48
View File
@@ -44,7 +44,6 @@ 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 维度)
@@ -164,6 +163,7 @@ def _infer_expected_categories(script_tags: set[str] | None) -> set[str] | None:
return matched or None
logger = logging.getLogger(__name__)
router = APIRouter()
@@ -700,17 +700,6 @@ 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% → 短素材禁复用);
@@ -941,42 +930,6 @@ 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,8 +228,6 @@ 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)
@@ -242,8 +240,6 @@ 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,
)
-157
View File
@@ -38,7 +38,6 @@ 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,
@@ -369,8 +368,6 @@ 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",
@@ -418,121 +415,6 @@ 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(
@@ -584,35 +466,6 @@ 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(
@@ -629,8 +482,6 @@ 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()
@@ -826,8 +677,6 @@ class LipsyncService:
job.completed_at = _now
job.updated_at = _now
self.db.commit()
# lip_sync 超时全额退款
self._refund_lip_sync(job)
return job
# 未提交的任务不轮询
@@ -853,8 +702,6 @@ 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
@@ -872,8 +719,6 @@ 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:
@@ -967,8 +812,6 @@ 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
-29
View File
@@ -104,7 +104,6 @@ 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":
@@ -142,7 +141,6 @@ 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:
@@ -159,33 +157,6 @@ 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:
@@ -387,41 +387,6 @@ 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 重渲。
@@ -1202,10 +1167,6 @@ 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,
@@ -1244,7 +1205,6 @@ def generate_video(self, task_id: str) -> dict:
)
# ── 自动重试逻辑 ──────────────────────────────────────────────────
will_retry = False
try:
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
@@ -1257,7 +1217,6 @@ 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,
@@ -1291,13 +1250,6 @@ 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,
+29 -19
View File
@@ -3,7 +3,7 @@
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. _step_image_analysis 图片分析(V2: OCR+lite VLM并行 + pro兜底)
1.5 _step_video_analysis 参考视频风格分析(可选)
2. _step_intent_parsing 用户文案意图解析
3. _step_script_generation 编导分镜脚本生成(融合原 copy_fusion+storyboard+review,输出 copy_result 结构 + voiceover_script)
@@ -358,10 +358,10 @@ def _step_image_analysis(job: ViralVideoJob) -> dict:
"""步骤 1: 图片分析(V2 主路径)。
架构:
- 主力:火山 MediaKit OCR(专用API,未配置时自动跳过)+ qwen3.8-flash 强约束 JSON,每图2路并行,目标<3s;
- 主力:火山 MediaKit OCR(专用API)+ doubao-seed-2.1-lite 强约束 JSON(弥补火山云端缺失的
人体属性/商品检测/图像标签专用HTTP API),每图2路并行,目标<3s;
- 外层全并发(workers=8),目标8图<15s;
- 兜底:fast 结果不可用时单次调用 qwen3.7-plus(简单、无竞速)。
- 唯一后端:阿里云百炼 DashScope,API Key 从环境变量 DASHSCOPE_API_KEY 读取。
- 兜底:fast 结果不可用时单次调用 doubao-seed-2.1-pro VLM(简单、无竞速)。
输出 dict 字段(name/brand/category/appearance/key_features/scene/mood/portrait_prompt/summary/_source)
与旧版格式完全一致,下游信任链/t2i/intent_parsing/script_generation 零改动。
"""
@@ -383,8 +383,26 @@ def _step_image_analysis(job: ViralVideoJob) -> dict:
logger.error("[爆款视频] vision 模块导入失败: %s", e)
return {"products": [_vision_fallback(0, f"vision_import_error:{e}")]}
# V2 内部 httpx 直连 dashscope,单次调用无重试,无需调整全局 client
results = _aiv2(normalized_urls)
# 整个阶段关闭底层 httpx 重试,避免线程里出现不可控等待
try:
from packages.shared.ai_client import get_doubao_client as _gdc
_cli = _gdc()
_orig_retries = _cli.max_retries
_cli.max_retries = 0
except Exception:
_cli = None
_orig_retries = 0
try:
results = _aiv2(normalized_urls)
finally:
if _cli is not None:
try:
_cli.max_retries = _orig_retries
except Exception:
pass
return {"products": list(results)}
@@ -421,14 +439,10 @@ def _step_intent_parsing(job: ViralVideoJob, image_analysis: dict) -> dict:
render_system_prompt,
render_user_prompt,
)
from packages.shared.ai_client import get_doubao_client
from packages.shared.ai_service import call_llm
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:
@@ -483,13 +497,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 = _llm_client.chat_completion(
raw = call_llm(
[{"role": "system", "content": system}, {"role": "user", "content": user}],
temperature=0.4,
max_tokens=1024,
model=_m,
timeout=60,
) # #2180/#2215: 直接用 client.chat_completion 传 messages list,不再走 call_llm 字符串包装
) # #2180: 意图解析 LLM 实测需更长响应,原25s太紧
if not raw:
continue
parsed = _parse(raw)
@@ -829,14 +843,10 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
GLOBAL_CONSTRAINTS,
NEGATIVE_RULES,
)
from packages.shared.ai_client import get_doubao_client
from packages.shared.ai_service import call_llm
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)))
@@ -872,7 +882,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 = _llm_client2.chat_completion(
raw = call_llm(
[{"role": "system", "content": system_tpl}, {"role": "user", "content": user}],
temperature=temp,
max_tokens=max_tok,
@@ -1,208 +0,0 @@
# -*- 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()
+100 -439
View File
@@ -1,8 +1,5 @@
# -*- 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兼容两种格式。
"""把 fast_json VLM 输出 + OCR 文本组装为与旧 _normalize() 完全一致的 dict。
目标:下游(信任链t2i/intent_parsing/script_generation)零改动。
必出字段:name, brand, category, appearance, packaging, text_on_package,
@@ -13,16 +10,29 @@ from __future__ import annotations
from typing import Any
# ---------- portrait_prompt 模板 ----------
# 目标:60-100 字的人物穿搭描述,用于 Seedream 纯文生图。要求具体、风格化、视觉细节丰富。
# 旧 VLM 输出格式参考:"一位25岁左右的亚洲女性,身穿白色V领短袖T恤,黑色高腰阔腿裤,
# 搭配银色项链,长发披肩,表情自信,街拍风格,阳光明媚的城市街头"
def _join_parts(*parts: str | None) -> str:
return "".join(p for p in parts if p)
_AGE_PREFIX = {"青年": "年轻", "中年": "中年", "老年": "老年"}
_AGE_PREFIX = {
"青年": "年轻",
"中年": "中年",
"老年": "老年",
}
# gender 后缀
_GENDER_WORD = {"男": "男性", "女": "女性"}
def _person_subject(gender: str, age: str) -> str:
def _person_subject(fj: dict[str, Any]) -> str:
"""人物主语:年轻女性 / 中年男性 / 少女 / 小男孩 / 人物 等。"""
gender = fj.get("gender") or ""
age = fj.get("age_range") or ""
gw = _GENDER_WORD.get(gender, "")
if age == "儿童":
if gender == "女":
@@ -42,147 +52,8 @@ def _person_subject(gender: str, age: str) -> str:
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:
def _build_wear_sentence(fj: dict[str, Any]) -> str:
"""穿搭段:上装+下装/连衣裙,带颜色+材质+图案。"""
upper = fj.get("upper_wear") or ""
upper_color = fj.get("upper_color") or ""
lower = fj.get("lower_wear") or ""
@@ -190,6 +61,7 @@ def _build_wear_sentence_old(fj: dict) -> str:
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
@@ -199,22 +71,25 @@ def _build_wear_sentence_old(fj: dict) -> str:
if pattern and pattern not in wear and pattern != "纯色":
wear += f",{pattern}图案"
return f"身穿{wear}"
parts = []
parts: list[str] = []
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}")
parts.append(f"上身{up}" if up else "")
if lower:
lo = f"{lower_color}{lower}" if lower_color else lower
parts.append(f"下身{lo}")
parts.append(f"下身{lo}" if lo else "")
return ",".join(p for p in parts if p)
def _build_portrait_prompt_old(fj: dict) -> str:
def _build_portrait_prompt(fj: dict[str, Any]) -> str:
"""组装最终 portrait_prompt(目标 60-100 字,用于 Seedream 纯文生图)。"""
if not fj.get("has_person"):
# 非人像:用商品+场景+mood 拼一段
name = fj.get("product_name") or "商品"
brand = fj.get("brand") or ""
colors = fj.get("colors") or []
@@ -226,15 +101,7 @@ def _build_portrait_prompt_old(fj: dict) -> str:
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) + "配色")
pieces.append("、".join(colors[:3]) + "配色")
if style:
pieces.append(style + "风格")
if mood:
@@ -244,32 +111,40 @@ def _build_portrait_prompt_old(fj: dict) -> str:
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)
subject = _person_subject(fj)
wear = _build_wear_sentence(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 ""
acc_str = ""
if accessories:
acc_str = ",佩戴" + "、".join(str(a) for a in accessories if a)
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 = []
detail_parts: list[str] = []
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 = []
style_parts: list[str] = []
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)
@@ -277,40 +152,53 @@ def _build_portrait_prompt_old(fj: dict) -> str:
pieces.append(acc_str.lstrip(","))
if detail_parts:
pieces.append(",".join(detail_parts))
pieces.append("".join(style_parts) + "风格" if style_parts else "人像写真")
if style_parts:
# 风格词之间不用逗号,用空格紧凑
pieces.append("".join(style_parts) + "风格")
else:
pieces.append("人像写真")
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:
# ---------- 商品字段 ----------
def _infer_name(fj: dict[str, Any], ocr_texts: list[str]) -> str:
pname = fj.get("product_name")
if pname and pname != "未识别":
return str(pname)
# 人物图 → name 用穿搭主件
if fj.get("has_person"):
up = fj.get("upper_wear") or ""
if "连衣裙" in up:
return up
return up or "人物穿搭"
if ocr_texts:
# 商品名可能是 OCR 最长的一行(品牌/产品名)
return max(ocr_texts, key=len)
return "未识别"
def _infer_brand_old(fj: dict, ocr_texts: list[str]) -> str:
def _infer_brand(fj: dict[str, Any], ocr_texts: list[str]) -> str:
brand = fj.get("brand")
if brand:
return str(brand)
# OCR 里短的、纯字母/汉字短串可能是 brand
for t in ocr_texts:
if 1 < len(t) <= 12:
return t
return "无法判断"
def _infer_category_old(fj: dict) -> str:
def _infer_category(fj: dict[str, Any]) -> str:
cat = fj.get("category")
if cat:
return str(cat)
@@ -319,19 +207,27 @@ def _infer_category_old(fj: dict) -> str:
return "非产品图"
def _build_appearance_old(fj: dict) -> str:
parts = []
for key in ("upper_color", "upper_wear", "material", "pattern"):
def _build_appearance(fj: dict[str, Any]) -> str:
"""外观描述:颜色+款式+材质+图案 拼成一段。"""
parts: list[str] = []
for key, _label 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 "无法判断"
if fj.get("has_person"):
return "人像穿搭整体造型"
return "无法判断"
return "、".join(parts)
def _build_key_features_old(fj: dict, ocr_texts: list[str]) -> list[str]:
feats = []
def _build_key_features(fj: dict[str, Any], ocr_texts: list[str]) -> list[str]:
feats: list[str] = []
for key in (
"upper_wear",
"lower_wear",
@@ -352,7 +248,9 @@ def _build_key_features_old(fj: dict, ocr_texts: list[str]) -> list[str]:
feats.append(v)
if ocr_texts:
feats.append(f"画面文字: {'/'.join(ocr_texts[:3])}")
out, seen = [], set()
# 去重
out: list[str] = []
seen: set[str] = set()
for f in feats:
f = f.strip()
if f and f not in seen and len(f) <= 30:
@@ -361,275 +259,27 @@ def _build_key_features_old(fj: dict, ocr_texts: list[str]) -> list[str]:
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]:
def assemble_result(
idx: int,
fast_json: dict[str, Any] | None,
ocr_texts: list[str],
) -> dict[str, Any]:
"""把 fast_json 结果 + OCR 文本组装成下游兼容的 product dict。"""
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 []
portrait_prompt = _build_portrait_prompt(fj)
name = _infer_name(fj, ocr_texts)
brand = _infer_brand(fj, ocr_texts)
category = _infer_category(fj)
appearance = _build_appearance(fj)
key_features = _build_key_features(fj, ocr_texts)
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 = "无法判断"
packaging = "无法判断" # 包装细节专用API无,保留占位
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
summary = _build_summary(fj, name, brand, category)
return {
"name": name,
"brand": brand,
@@ -644,3 +294,14 @@ def _assemble_old(idx: int, fj: dict, ocr_texts: list[str]) -> dict[str, Any]:
"summary": summary,
"_source": "v2_fast_json",
}
def _build_summary(fj: dict, name: str, brand: str, category: str) -> str:
if fj.get("has_person"):
up = fj.get("upper_wear") or "穿搭"
style = fj.get("style") or ""
base = f"{style}{up}" if style and style not in up else up
return base
if brand != "无法判断" and name != brand:
return f"{brand} {name}"
return name
@@ -1,11 +1,9 @@
# -*- coding: utf-8 -*-
"""V2 图片分析主路径:每图并行 OCR(火山MediaKit,未配置时自动跳过)+ qwen3.8-flash JSON VLM,
失败时单次 qwen3.7-plus 兜底。
"""V2 图片分析主路径:每图并行 OCR(火山专用API)+ lite JSON VLM,失败时单次 pro VLM 兜底。
架构(灵应10-05确认):
- 唯一后端:阿里云百炼 DashScope,qwen3.8-flash 做快速路径、qwen3.7-plus 做兜底
- 主力:单图2路并行(OCR + fast VLM),外层N图全并发(workers=8)
- 兜底:单次 pro VLM 调用,无竞速/重试/复杂超时
设计原则(灵应10-05要求):
- 主力路径简洁:单图2路并行,外层N图全并发
- 兜底简单:单次 pro VLM 调用,无竞速/重试/复杂超时
- 输出 dict 格式与旧版完全一致,下游零改动
"""
@@ -21,12 +19,12 @@ 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"))
_FAST_TIMEOUT = float(os.environ.get("VISION_V2_FAST_TIMEOUT", "8"))
_FAST_JSON_TIMEOUT = float(os.environ.get("VISION_V2_FAST_JSON_TIMEOUT", "8"))
_OCR_TIMEOUT = float(os.environ.get("VISION_V2_OCR_TIMEOUT", "6"))
_PRO_TIMEOUT = float(os.environ.get("VISION_V2_PRO_TIMEOUT", "30"))
_PRO_TIMEOUT = float(os.environ.get("VISION_V2_PRO_TIMEOUT", "45"))
_FALLBACK_RESULT = {
"name": "未识别",
@@ -44,6 +42,7 @@ _FALLBACK_RESULT = {
def _is_usable(r: dict[str, Any]) -> bool:
"""结果可用判定:portrait_prompt 是核心,有效就算 usable。"""
pp = (r.get("portrait_prompt") or "").strip()
if pp and pp not in ("无人像", "无法判断", "未识别"):
return True
@@ -54,8 +53,10 @@ def _is_usable(r: dict[str, Any]) -> bool:
def analyze_image_v2(idx: int, img_url: str) -> dict[str, Any]:
"""单张图片 V2 分析。"""
t0 = time.time()
# 第1层:OCR + lite JSON VLM 并行
fj_result: dict[str, Any] | None = None
ocr_result: list[str] = []
with ThreadPoolExecutor(max_workers=2) as pool:
@@ -73,6 +74,7 @@ def analyze_image_v2(idx: int, img_url: str) -> dict[str, Any]:
elif fut is f_ocr and isinstance(res, list):
ocr_result = res
except TimeoutError:
# fast 整体超时,取消还没跑完的子任务,继续走 pro 兜底
for f in (f_fj, f_ocr):
if not f.done():
f.cancel()
@@ -80,6 +82,7 @@ def analyze_image_v2(idx: int, img_url: str) -> dict[str, Any]:
fast_elapsed = time.time() - t0
# 组装 fast 结果
if fj_result:
assembled = assembler.assemble_result(idx, fj_result, ocr_result)
if _is_usable(assembled):
@@ -92,6 +95,7 @@ def analyze_image_v2(idx: int, img_url: str) -> dict[str, Any]:
)
return assembled
# 第2层:pro VLM 单次兜底
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):
@@ -103,6 +107,7 @@ def analyze_image_v2(idx: int, img_url: str) -> dict[str, Any]:
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"
@@ -112,18 +117,13 @@ def analyze_image_v2(idx: int, img_url: str) -> dict[str, Any]:
def analyze_images_v2(img_urls: list[str]) -> list[dict[str, Any]]:
"""批量图片 V2 分析,外层全并发。"""
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,
)
logger.info("[vision.v2] 开始图片分析 n=%d workers=%d fast_timeout=%.0fs", len(img_urls), workers, _FAST_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)}
@@ -1,126 +1,226 @@
# -*- coding: utf-8 -*-
"""V2 兜底路径:qwen3.7-plus(阿里云百炼/DashScope)单图调用。
"""VLM 兜底:专用API路径失败时的最后一道防线,单次调用 doubao-seed-2.1-pro。
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 路径输出格式完全一致
设计原则:简单、直接、无竞速、无复杂超时逻辑。只在 fast_json 结果不可用时调用。
"""
from __future__ import annotations
import json
import logging
import os
import re
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
DEFAULT_PRO_MODEL = "doubao-seed-2-1-pro-260915"
DEFAULT_TIMEOUT = 45
DEFAULT_MAX_TOKENS = 800
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 _xml_text(tag: str, xml: str) -> str:
m = re.search(rf"<{tag}[^>]*>(.*?)</{tag}>", xml, re.S)
return (m.group(1) if m else "").strip()
def _xml_attr(tag: str, attr: str, xml: str) -> str:
m = re.search(rf"<{tag}[^>]*\b{attr}\s*=\s*[\"']([^\"']*)[\"']", xml)
return (m.group(1) if m else "").strip()
def _xml_to_product(raw: str, idx: int) -> dict[str, Any]:
"""解析 VLM 输出的 XML 格式(简化版)。"""
scene = _xml_text("scene", raw) or "通用"
mood = _xml_text("mood", raw) or ""
portrait_prompt = "无人像"
p_has = _xml_attr("people", "has_person", raw)
if p_has and p_has.lower() != "false":
gender = _xml_attr("people", "gender", raw) or ""
age = _xml_attr("people", "age_range", raw) or ""
outfit = _xml_attr("people", "outfit", raw) or ""
hair = _xml_attr("people", "hair", raw) or "自然发型"
pose = _xml_attr("people", "pose", raw) or ""
expr = _xml_attr("people", "expression", raw) or "自然"
parts: list[str] = []
if gender:
parts.append(gender + ("性" if not gender.endswith("性") else ""))
if age:
parts.append(age)
parts.append("人物")
parts.append(hair)
if outfit:
parts.append(f"身着{outfit}")
if pose:
parts.append(f"姿态{pose}")
parts.append(f"表情{expr}")
portrait_prompt = ",".join(parts)
m = re.search(r"<product[^>]*>(.*?)</product>", raw, re.S)
if m:
pbody = m.group(1)
name = _xml_attr("product", "name", raw) or _xml_text("name", pbody) or "未识别"
brand = _xml_attr("product", "brand", raw) or _xml_text("brand", pbody) or "无法判断"
category = _xml_attr("product", "category", raw) or _xml_text("category", pbody) or "无法判断"
appearance = _xml_attr("product", "appearance", raw) or _xml_text("appearance", pbody) or "无法判断"
packaging = _xml_attr("product", "packaging", raw) or _xml_text("packaging", pbody) or "无法判断"
feat = _xml_attr("product", "features", raw) or _xml_text("features", pbody) or ""
feat_list = [x.strip() for x in re.split(r"[,,;;]", feat) if x.strip()] if feat else ["无法判断"]
top_text = _xml_attr("product", "text_on_package", raw) or _xml_text("text_on_package", pbody) or ""
text_list = [x.strip() for x in re.split(r"[,,;;]", top_text) if x.strip()] if top_text else []
summary = _xml_attr("product", "summary", raw) or _xml_text("summary", pbody) or f"{brand} {name}"
pp_attr = _xml_attr("product", "portrait_prompt", raw)
if pp_attr and pp_attr != "无人像":
portrait_prompt = pp_attr
return {
"name": name,
"brand": brand,
"category": category,
"appearance": appearance,
"packaging": packaging,
"text_on_package": text_list,
"key_features": feat_list,
"scene": scene,
"mood": mood,
"portrait_prompt": portrait_prompt,
"summary": summary,
"_source": "vlm_pro_xml",
}
if portrait_prompt != "无人像":
return {
"name": "未识别",
"brand": "无法判断",
"category": "无法判断",
"appearance": "无法判断",
"packaging": "无法判断",
"text_on_package": [],
"key_features": ["无法判断"],
"scene": scene,
"mood": mood,
"portrait_prompt": portrait_prompt,
"summary": "未识别",
"_source": "vlm_pro_no_product",
}
return {
"name": "未识别",
"brand": "无法判断",
"category": "无法判断",
"appearance": "无法判断",
"packaging": "无法判断",
"text_on_package": [],
"key_features": ["无法判断"],
"scene": scene,
"mood": mood,
"portrait_prompt": "无人像",
"summary": "未识别",
"_source": "vlm_pro_no_tag",
}
def call_pro_vlm(
img_url: str,
idx: int,
*,
timeout: int = _DEFAULT_TIMEOUT,
model: str | None = None,
timeout: int = DEFAULT_TIMEOUT,
) -> dict[str, Any] | None:
"""单次调用 pro VLM,解析后返回 product dict;失败返回 None。"""
t0 = time.time()
import httpx
api_key = _api_key()
if not api_key:
logger.warning("[vision.v2] pro DASHSCOPE_API_KEY 未配置,跳过")
try:
from packages.application.viral_video.prompt_loader import (
get_template,
render_system_prompt,
render_user_prompt,
)
from packages.shared.ai_client import get_doubao_client
except ImportError as e:
logger.warning("[vision.vlm] 导入失败: %s", e)
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
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}")
except Exception as e:
logger.warning("[vision.vlm] 模板加载失败: %s", e)
return None
# 通过assembler统一组装,兼容v4嵌套schema和旧扁平schema
result = assembler.assemble_result(idx, obj, [])
result["_source"] = "vlm_pro"
client = get_doubao_client()
if not client.is_available:
return None
use_model = model or DEFAULT_PRO_MODEL
_orig_retries = client.max_retries
client.max_retries = 0
try:
raw = client.vision_completion(
messages=[{"role": "system", "content": system}, {"role": "user", "content": user}],
images=[img_url],
temperature=0.3,
max_tokens=DEFAULT_MAX_TOKENS,
timeout=timeout,
model=use_model,
)
except Exception as e:
logger.warning("[vision.vlm] 图片 #%d pro VLM 调用失败 elapsed=%.1fs err=%s", idx, time.time() - t0, e)
client.max_retries = _orig_retries
return None
client.max_retries = _orig_retries
elapsed = time.time() - t0
if not raw:
logger.warning("[vision.vlm] 图片 #%d pro VLM 返回空 elapsed=%.1fs", idx, elapsed)
return None
text = _strip_code_fence(raw)
lb, rb = text.find("{"), text.rfind("}")
if lb >= 0 and rb > lb:
try:
obj = json.loads(text[lb : rb + 1])
if isinstance(obj, dict):
logger.info("[vision.vlm] 图片 #%d pro VLM JSON 完成 elapsed=%.1fs", idx, elapsed)
return {
"name": obj.get("name") or "未识别",
"brand": obj.get("brand") or "无法判断",
"category": obj.get("category") or "无法判断",
"appearance": obj.get("appearance") or "无法判断",
"packaging": obj.get("packaging") or "无法判断",
"text_on_package": obj.get("text_on_package") or [],
"key_features": obj.get("key_features") or obj.get("features") or ["无法判断"],
"scene": obj.get("scene") or "通用",
"mood": obj.get("mood") or "",
"portrait_prompt": obj.get("portrait_prompt") or "无人像",
"summary": obj.get("summary") or f"{obj.get('brand','')} {obj.get('name','')}",
"_source": "vlm_pro_json",
}
except json.JSONDecodeError:
pass
try:
result = _xml_to_product(text, idx)
result["_fallback_used"] = True
result["_pro_elapsed"] = round(elapsed, 2)
logger.info(
"[vision.vlm] 图片 #%d pro VLM XML 完成 elapsed=%.2fs pp=%s",
idx,
elapsed,
(result.get("portrait_prompt") or "")[:40],
)
return result
except Exception as e:
elapsed = time.time() - t0
logger.warning("[vision.v2] pro 异常 elapsed=%.1fs err=%s", elapsed, e, exc_info=True)
logger.warning("[vision.vlm] 图片 #%d 解析失败 elapsed=%.1fs err=%s head=%s", idx, elapsed, e, raw[:200])
return None
@@ -1,46 +1,71 @@
# -*- coding: utf-8 -*-
"""V2 快速路径:qwen3.8-flash(阿里云百炼/DashScope)强约束 JSON-only 调用。
"""doubao-seed-2.1-lite 强约束 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 读取
- system prompt 极致精简,只给字段 schema 和强约束(禁止自然语言、禁止 markdown)
- max_tokens=350(比旧 VLM 的 1200 小很多,降低延迟)
- temperature=0.1(极低,稳定输出 JSON)
- timeout=8s(够快,失败则由外层走 pro VLM 兜底)
- 期望返回纯 JSON object(无 ```json 包裹、无解释文字)
"""
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
# 极简 system prompt:只给字段定义 + 硬性输出要求
_FAST_SYSTEM = (
"你是图片结构化识别器。严格按下方 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'
"}"
)
_FAST_USER = "识别这张图片的人物穿搭与主体信息,只返回JSON对象。"
def _api_key() -> str | None:
return os.environ.get("DASHSCOPE_API_KEY")
# 默认模型
DEFAULT_LITE_MODEL = "doubao-seed-2-1-lite-260915"
DEFAULT_TIMEOUT = 8
DEFAULT_MAX_TOKENS = 350
def _strip_code_fence(s: str) -> str:
"""剥离 ```json ... ``` 包裹(即使要求纯 JSON,模型偶尔仍会包代码块)。"""
s = s.strip()
if s.startswith("```"):
lines = s.split("\n")
# 去掉首行 ```json
if lines and lines[0].startswith("```"):
lines = lines[1:]
# 去掉尾行 ```
if lines and lines[-1].strip().startswith("```"):
lines = lines[:-1]
s = "\n".join(lines).strip()
@@ -50,40 +75,52 @@ def _strip_code_fence(s: str) -> str:
def call_fast_json(
img_url: str,
*,
timeout: int = _DEFAULT_TIMEOUT,
max_tokens: int = _DEFAULT_MAX_TOKENS,
model: str | None = None,
timeout: int = DEFAULT_TIMEOUT,
max_tokens: int = DEFAULT_MAX_TOKENS,
) -> dict[str, Any] | None:
"""调用 qwen3.8-flash 返回结构化 dict;失败/非 JSON 返回 None。"""
"""调用 lite VLM 返回结构化 dict;失败/非 JSON 返回 None。
直接用 httpx 发最小 payload(关闭 thinking),不走 ai_client 包装:
- 关闭 thinking/推理链(reasoning_tokens 是延迟主因,单次要10-12s)
- 单次调用不重试(失败由外层走 pro 兜底)
- 温度=0.1 稳定输出 JSON
"""
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:
from packages.shared import get_shared_settings
settings = get_shared_settings()
api_key = settings.doubao_api_key
base_url = (settings.doubao_base_url or "https://ark.cn-beijing.volces.com/api/v3").rstrip("/")
if not api_key:
logger.warning("[vision.v2] doubao api_key 未配置,跳过 fast_json")
return None
use_model = model or DEFAULT_LITE_MODEL
url = f"{base_url}/chat/completions"
payload: dict[str, Any] = {
"model": use_model,
"messages": [
{"role": "system", "content": _FAST_SYSTEM},
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": img_url}},
{"type": "text", "text": _FAST_USER},
],
},
],
"temperature": 0.1,
"max_tokens": max_tokens,
"stream": False,
}
# 关键:关闭 thinking(reasoning_tokens 是延迟主因,单次要10-12s)
# 方舟/豆包 Seed 2.x 支持 thinking={type:"disabled"},且不要和 reasoning_effort 同时传(两者互斥会400)
payload["thinking"] = {"type": "disabled"}
resp = httpx.post(
url,
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
@@ -91,53 +128,69 @@ def call_fast_json(
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:
# 400 说明模型不支持 thinking 参数(极少数旧模型),重试一次不带 thinking
if resp.status_code == 400:
body_preview = resp.text[:300].lower()
logger.warning("[vision.v2] fast_json HTTP 400 elapsed=%.1fs body=%s", elapsed, resp.text[:200])
if "thinking" in body_preview or "reasoning" in body_preview:
payload.pop("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", resp.status_code, elapsed)
return None
else:
return None
elif 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)
if raw is None:
logger.warning("[vision.v2] fast_json 返回 None 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,
"[vision.v2] fast_json 直连完成 model=%s elapsed=%.1fs in=%d out=%d reasoning=%d",
use_model,
elapsed,
usage.get("prompt_tokens", 0),
usage.get("completion_tokens", 0),
reasoning_tokens,
usage.get("reasoning_tokens", 0),
)
elapsed = time.time() - t0
if raw is None:
logger.warning("[vision.v2] fast_json 返回 None elapsed=%.1fs model=%s", elapsed, use_model)
return None
text = _strip_code_fence(raw)
lpos, r = text.find("{"), text.rfind("}")
if lpos >= 0 and r > lpos:
text = text[lpos : r + 1]
# 截到第一个 { 和最后一个 } 之间,容忍前后偶发文字
lb = text.find("{")
rb = text.rfind("}")
if lb >= 0 and rb > lb:
text = text[lb : rb + 1]
try:
obj = json.loads(text)
except json.JSONDecodeError:
logger.warning("[vision.v2] fast_json JSON 解析失败 elapsed=%.1fs head=%s", elapsed, raw[:200])
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",
"[vision.v2] fast_json 完成 model=%s elapsed=%.1fs has_person=%s has_product=%s category=%s",
use_model,
elapsed,
obj.get("has_person"),
obj.get("has_product"),
@@ -335,10 +335,6 @@ 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,
@@ -731,11 +727,6 @@ 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)
@@ -914,11 +905,6 @@ 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 注册表 — 反向轮询模式下用于心跳与监控."""
-376
View File
@@ -1,376 +0,0 @@
"""功能计费配置服务:从 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
+14 -77
View File
@@ -2,21 +2,17 @@
v1.6.1: 按产品决策,智能混剪/AI数字人/AI配音/抖音解析/改写/标题/封面 全部免费,
仅保留声音克隆合成(voice_clone_synth)的扣点逻辑;声音克隆训练保持免费。
爆款视频(viral_video)走动态定价,计费参数 DB 化(feature_pricing_configs,
见 feature_pricing_service),calculate_viral_video_credits 从配置读取单价/
固定成本/利润系数/封顶,DB 不可用时回落兜底配置。
爆款视频(viral_video)走动态定价,见本文件 VIRAL_VIDEO_MODEL_PRICES + calculate_viral_video_credits。
"""
from __future__ import annotations
import math
from packages.domain import feature_pricing_service
# ============ 爆款视频动态定价 ============
# 单价/固定成本/利润系数已 DB 化(feature_pricing_configs,feature_key=viral_video),
# 由 feature_pricing_service 读取(300s 缓存),DB 不可用时回落内置兜底配置。
# 以下三个常量仅为向后兼容保留(旧引用方/兜底场景),值取自兜底配置。
# ============ 爆款视频动态定价 (#2151) ============
# key = (model_id, resolution, has_video_input),单位:
# - billing_mode=token: 元/百万tokens(输出)
# - billing_mode=per_second: 元/秒(视频时长)
VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = {
("seedance-2.5", "480p", False): 70.0,
("seedance-2.5", "720p", False): 70.0,
@@ -37,9 +33,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
@@ -226,22 +222,17 @@ 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 / feature_enabled / charged / price_cap
字段,便于前端展示计费明细。功能关闭时 credits=0、charged=False。
model_price / width / height / fps 字段,便于前端展示计费明细。
"""
w = max(1, int(width or 1))
h = max(1, int(height or 1))
@@ -251,36 +242,11 @@ 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")
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)
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)
@@ -293,40 +259,13 @@ def calculate_viral_video_credits_with_breakdown(
video_cost = tokens / 1_000_000.0 * float(price)
billing_unit = "token"
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)
total = (video_cost + VIRAL_VIDEO_FIXED_COST) * VIRAL_VIDEO_PROFIT_MULTIPLIER
credits = round(float(total), 2)
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),
"fixed_cost": float(VIRAL_VIDEO_FIXED_COST),
"profit_multiplier": float(VIRAL_VIDEO_PROFIT_MULTIPLIER),
"model_price": float(price),
"model_key": prefix,
"billing_mode": billing,
@@ -335,8 +274,6 @@ def calculate_viral_video_credits_with_breakdown(
"height": int(h),
"fps": int(effective_fps),
"duration": dur,
"feature_enabled": True,
"charged": True,
}
return credits, breakdown
@@ -1,221 +0,0 @@
"""功能计费改造测试:爆款读配置、对口型/智能剪辑预扣逻辑。
策略:
- 爆款:通过修改缓存中的 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
-235
View File
@@ -1,235 +0,0 @@
"""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
+18 -51
View File
@@ -97,43 +97,21 @@ def invalidate_loader_cache():
class TestImageAnalysisWiring:
def test_step_image_analysis_uses_v2_batch_path(self, job):
"""#2200/#2207 后图片分析走 V2 批处理(OCR+lite JSON 并行),
_step_image_analysis 归一化 URL 后调用 analyze_images_v2。"""
def test_uses_loader_template_and_xml_parse(self, job):
from apps.worker.worker_app.tasks import viral_video as vv
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)
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)
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": []}
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"]
# ── 2) 意图解析走模板 ───────────────────────────────────────────────
@@ -274,29 +252,18 @@ 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.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_vision", return_value=IMAGE_XML),
patch("packages.shared.ai_service.call_llm", return_value=INTENT_XML),
):
# 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 模板)
# 1) image
img_res = vv._analyze_single_image(0, "https://img/1.jpg", "vlm", 15)
# 2) intent
intent_res = vv._step_intent_parsing(job, {"products": [img_res]})
# V2 图片分析不再调用 loader;意图解析调用 intent_parsing 模板
assert "image_analysis" not in called_types
# 前两步分别调用了 image_analysis 和 intent_parsing
assert "image_analysis" in called_types
assert "intent_parsing" in called_types
# script 和 review 单独验证(需要不同的 LLM 返回)
-291
View File
@@ -1,291 +0,0 @@
# -*- 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