feat(#2035): semantic tags + quality score + AI caption+embedding
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 2s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / 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 55s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 59s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m38s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 2m26s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 2m35s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m21s
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 / Validate - Style (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Security (pull_request) Has been cancelled
CI/CD Pipeline / Unit Tests (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
AI Code Review / AI Code Review (pull_request) Has been cancelled
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been cancelled
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 2s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / 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 55s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 59s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m38s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 2m26s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 2m35s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m21s
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 / Validate - Style (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Security (pull_request) Has been cancelled
CI/CD Pipeline / Unit Tests (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
AI Code Review / AI Code Review (pull_request) Has been cancelled
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been cancelled
- Fix: generation_tasks.py passes clip_ai_tags_by_asset to pick_narrative_assets so AI tags (weight 2.0) actually participate in narrative mode selection - Feat: quality_score auto-computation via AssetAnalyzer on ingest (worker.calculate_asset_price celery task; fallback 50.0 on failure) - Feat: smart_match adds ai_semantic dimension (20% weight) using Jaccard similarity between asset AI tags (scene/objects/action) and script tags - Feat: atom_clip caption (10-30 Chinese chars) via Doubao Vision, saved to asset_atom_clips.caption (Text column, migration 085) - Feat: atom_clip embedding vector via Doubao embeddings API, saved to asset_atom_clips.embedding (JSON column) - Chore: remove dead calculate_quality_score_real wrapper - Tests: 20 new unit tests covering caption parsing, ai_semantic scoring, narrative AI tag propagation, score weight changes; update existing tests for new fallback dict shape and reweighted dimensions - Fail-open: tagging/embedding/quality failures never block main flow
This commit is contained in:
@@ -0,0 +1,33 @@
|
||||
"""asset_atom_clips 新增 caption/embedding 字段(#2035 语义标签增强)
|
||||
|
||||
Revision ID: 085_atom_clip_caption_embedding
|
||||
Revises: 084_lipsync_jobs_style
|
||||
Create Date: 2026-09-25
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "085_atom_clip_caption_embedding"
|
||||
down_revision = "084_lipsync_jobs_style"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# caption: 中文画面描述(10-30字)
|
||||
op.add_column(
|
||||
"asset_atom_clips",
|
||||
sa.Column("caption", sa.Text(), nullable=True),
|
||||
)
|
||||
# embedding: caption 对应的向量(豆包 embedding 接口返回,JSON 存 float 数组)
|
||||
op.add_column(
|
||||
"asset_atom_clips",
|
||||
sa.Column("embedding", sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("asset_atom_clips", "embedding")
|
||||
op.drop_column("asset_atom_clips", "caption")
|
||||
@@ -134,10 +134,11 @@ def _ensure_library_has_ready_video_assets(assets) -> None:
|
||||
def _select_assets_from_library(
|
||||
assets: list,
|
||||
mode: str,
|
||||
count: int,
|
||||
count: int = 0,
|
||||
rng=None,
|
||||
script_tags: list | None = None,
|
||||
tag_names_by_id: dict | None = None,
|
||||
db=None,
|
||||
) -> list[str]:
|
||||
"""根据选取模式从素材库中选取 ready 状态的视频素材 ID。
|
||||
|
||||
@@ -158,6 +159,42 @@ def _select_assets_from_library(
|
||||
if not ready_video_assets:
|
||||
return []
|
||||
|
||||
# #2035:加载片段级 AI 标签,供叙事模式 AI 加权和 smart 模式语义匹配使用。
|
||||
# 失败降级为空(不影响选片主流程)。
|
||||
clip_ai_tags_by_asset: dict[str, list[dict]] = {}
|
||||
ai_tags_by_asset: dict[str, dict] = {} # asset_id → 聚合后的 ai_tags dict(取首个有 has_text 的片段;合并 scene/objects/action 去重)
|
||||
try:
|
||||
if db is not None:
|
||||
from packages.adapters.sqlalchemy_impl.models import AssetAtomClipModel
|
||||
ready_ids = [a.id for a in ready_video_assets]
|
||||
clip_rows = (
|
||||
db.query(AssetAtomClipModel.asset_id, AssetAtomClipModel.ai_tags)
|
||||
.filter(AssetAtomClipModel.asset_id.in_(ready_ids))
|
||||
.filter(AssetAtomClipModel.ai_tags.isnot(None))
|
||||
.all()
|
||||
)
|
||||
agg: dict[str, dict] = {}
|
||||
for asset_id, ai_tags in clip_rows:
|
||||
if not isinstance(ai_tags, dict):
|
||||
continue
|
||||
clip_ai_tags_by_asset.setdefault(asset_id, []).append(ai_tags)
|
||||
# 聚合:合并 scene/objects/action 去重
|
||||
agg.setdefault(asset_id, {"scene": [], "objects": [], "action": [], "shot": "", "has_text": False})
|
||||
for key in ("scene", "objects", "action"):
|
||||
for v in ai_tags.get(key) or []:
|
||||
v = str(v).strip()
|
||||
if v and v not in agg[asset_id][key]:
|
||||
agg[asset_id][key].append(v)
|
||||
if ai_tags.get("has_text") is True:
|
||||
agg[asset_id]["has_text"] = True
|
||||
if not agg[asset_id]["shot"] and ai_tags.get("shot"):
|
||||
agg[asset_id]["shot"] = ai_tags["shot"]
|
||||
ai_tags_by_asset = agg
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("[选片] 加载片段 AI 标签失败,降级不使用语义匹配", exc_info=True)
|
||||
clip_ai_tags_by_asset = {}
|
||||
ai_tags_by_asset = {}
|
||||
|
||||
# 叙事模式(#1970 PR3):文案标签命中池优先;无任何命中时完全降级为现有随机逻辑。
|
||||
if script_tags:
|
||||
from packages.domain.narrative_match import pick_narrative_assets
|
||||
@@ -167,6 +204,7 @@ def _select_assets_from_library(
|
||||
ready_video_assets,
|
||||
script_tags=script_tags,
|
||||
tag_names_by_id=tag_names_by_id,
|
||||
clip_ai_tags_by_asset=clip_ai_tags_by_asset,
|
||||
limit=limit,
|
||||
rng=rng,
|
||||
)
|
||||
@@ -177,7 +215,16 @@ def _select_assets_from_library(
|
||||
# 评分维度:质量分(40%) + 时长适配(30%) + 新鲜度(20%) + 未使用加分(10%)
|
||||
# 排序注入随机噪声(#1743):同分素材每次选出不同组合,从素材组合层面降重
|
||||
limit = count if count > 0 else None
|
||||
results = smart_select_assets(ready_video_assets, limit=limit, kind="video", rng=rng)
|
||||
# #2035:给 smart_select_assets 传入文案标签和 AI 标签映射,启用语义维度
|
||||
norm_script = {t.strip().lower() for t in (script_tags or []) if t and t.strip()}
|
||||
results = smart_select_assets(
|
||||
ready_video_assets,
|
||||
limit=limit,
|
||||
kind="video",
|
||||
rng=rng,
|
||||
script_tags=norm_script if norm_script else None,
|
||||
ai_tags_by_asset=ai_tags_by_asset or None,
|
||||
)
|
||||
return [r.asset.id for r in results]
|
||||
|
||||
# 默认 all 模式:返回全部 ready 视频素材
|
||||
@@ -396,6 +443,7 @@ def create_generation_task(
|
||||
count=request.asset_select_count,
|
||||
script_tags=narrative_script_tags or None,
|
||||
tag_names_by_id=_tag_index,
|
||||
db=db,
|
||||
)
|
||||
elif project_id and not resolved_asset_ids and (request.asset_select_mode in ("smart",) or narrative_script_tags):
|
||||
# 项目级模式:未指定 asset_ids 且选择了 smart 模式(或叙事模式按标签匹配)时自动选取
|
||||
@@ -410,6 +458,7 @@ def create_generation_task(
|
||||
count=request.asset_select_count,
|
||||
script_tags=narrative_script_tags or None,
|
||||
tag_names_by_id=_tag_index,
|
||||
db=db,
|
||||
)
|
||||
if not resolved_asset_ids:
|
||||
raise HTTPException(
|
||||
|
||||
@@ -31,6 +31,7 @@ celery_app.conf.imports = (
|
||||
# #1970 片段级 AI 标签:必须显式 import 注册,否则 worker 报
|
||||
# "Received unregistered task of type 'worker.tag_atom_clip'"
|
||||
"worker_app.tasks.atom_clip_tagging",
|
||||
"worker_app.tasks.asset_quality_scoring_task",
|
||||
"worker_app.tasks.backfill_atom_clip_tags",
|
||||
"worker_app.tasks.classification",
|
||||
"worker_app.tasks.generation",
|
||||
|
||||
@@ -446,21 +446,3 @@ def classify_asset_real(video_path: str) -> tuple[str, float]:
|
||||
logger.warning(f"Classification failed, using fallback: {e}")
|
||||
return AssetClassification.OTHER.value, 0.3
|
||||
|
||||
|
||||
def calculate_quality_score_real(video_path: str) -> float:
|
||||
"""
|
||||
质量评分入口函数
|
||||
|
||||
Args:
|
||||
video_path: 视频文件路径
|
||||
|
||||
Returns:
|
||||
质量评分 (0-100)
|
||||
"""
|
||||
try:
|
||||
analyzer = AssetAnalyzer(video_path)
|
||||
result = analyzer.calculate_quality_score()
|
||||
return result.total
|
||||
except Exception as e:
|
||||
logger.warning(f"Quality scoring failed, using fallback: {e}")
|
||||
return 50.0
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
"""素材质量评分 Celery 任务 — #2035.
|
||||
|
||||
视频素材 READY 入库后异步触发:下载视频到临时文件,运行 FFmpeg+NumPy 质量分析,
|
||||
将 0-100 总分写入 assets.quality_score 字段。失败不阻断主流程(保留 NULL 或旧值,
|
||||
选片时按 50 分兜底)。
|
||||
|
||||
任务名:worker.calculate_asset_quality
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
from celery.utils.log import get_task_logger
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
logger = get_task_logger(__name__)
|
||||
|
||||
|
||||
@celery_app.task(name="worker.calculate_asset_quality", bind=True, max_retries=1, default_retry_delay=15)
|
||||
def calculate_asset_quality_task(self, asset_id: str) -> dict:
|
||||
"""为单个视频素材计算质量评分并写回 assets.quality_score。
|
||||
|
||||
流程:
|
||||
1. 下载视频到临时文件;
|
||||
2. 用 AssetAnalyzer(FFmpeg+NumPy) 提取分辨率/帧率/码率/清晰度/稳定性 5 维分数;
|
||||
3. 写回 assets.quality_score。
|
||||
|
||||
失败/非视频/无文件等情况均静默降级,返回 status=skipped/failed 不抛异常。
|
||||
"""
|
||||
db = SessionLocal()
|
||||
tmp_dir = tempfile.mkdtemp(prefix="quality_score_")
|
||||
try:
|
||||
asset_repo = SQLAlchemyAssetRepository(db)
|
||||
asset = asset_repo.find_by_id(asset_id)
|
||||
if asset is None:
|
||||
return {"status": "skipped", "reason": "asset not found", "asset_id": asset_id}
|
||||
if not (getattr(asset, "mime_type", "") or "").startswith("video/"):
|
||||
return {"status": "skipped", "reason": "not a video", "asset_id": asset_id}
|
||||
# 已有质量分则幂等跳过(重新计算需显式置空)
|
||||
if getattr(asset, "quality_score", None) is not None:
|
||||
return {"status": "skipped", "reason": "already scored", "asset_id": asset_id}
|
||||
|
||||
storage = get_shared_storage_service()
|
||||
storage_key = getattr(asset, "storage_key", "") or ""
|
||||
if not storage_key:
|
||||
return {"status": "skipped", "reason": "no storage_key", "asset_id": asset_id}
|
||||
|
||||
# 下载到临时文件
|
||||
safe_suffix = ".mp4"
|
||||
local_path = Path(tmp_dir) / f"asset_{asset_id[:8]}{safe_suffix}"
|
||||
ok = storage.download_asset(storage_key, local_path)
|
||||
if not ok or not local_path.exists() or local_path.stat().st_size == 0:
|
||||
return {"status": "failed", "reason": "download failed", "asset_id": asset_id}
|
||||
|
||||
# 调用 AssetAnalyzer
|
||||
from worker_app.tasks.asset_analyzer import AssetAnalyzer
|
||||
try:
|
||||
analyzer = AssetAnalyzer(str(local_path), temp_dir=tmp_dir)
|
||||
result = analyzer.calculate_quality_score()
|
||||
total = float(result.total) if result and 0 <= result.total <= 100 else 50.0
|
||||
except Exception as analyze_err: # noqa: BLE001
|
||||
logger.warning("[quality_score] 分析失败,使用默认50分: asset=%s err=%s", asset_id, analyze_err)
|
||||
total = 50.0
|
||||
|
||||
# 写回数据库
|
||||
asset.quality_score = total
|
||||
asset_repo.update(asset)
|
||||
db.commit()
|
||||
|
||||
logger.info("[quality_score] asset=%s score=%.1f", asset_id, total)
|
||||
return {"status": "completed", "asset_id": asset_id, "quality_score": total}
|
||||
except Exception as exc: # noqa: BLE001
|
||||
db.rollback()
|
||||
logger.exception("[quality_score] asset=%s 失败: %s", asset_id, exc)
|
||||
if self.request.retries < self.max_retries:
|
||||
raise self.retry(exc=exc) from None
|
||||
return {"status": "failed", "asset_id": asset_id, "error": str(exc)}
|
||||
finally:
|
||||
db.close()
|
||||
# 清理临时文件
|
||||
try:
|
||||
import shutil
|
||||
shutil.rmtree(tmp_dir, ignore_errors=True)
|
||||
except Exception:
|
||||
pass
|
||||
@@ -78,14 +78,18 @@ def tag_atom_clip_task(self, atom_clip_id: str, force: bool = False) -> dict:
|
||||
atom_repo.update_ai_tags(atom_clip_id, ai_tags)
|
||||
|
||||
logger.info(
|
||||
"[atom_clip_tagging] clip_id=%s ai_tags=%s",
|
||||
"[atom_clip_tagging] clip_id=%s ai_tags=%s caption=%r has_embedding=%s",
|
||||
atom_clip_id,
|
||||
{k: v for k, v in ai_tags.items() if k != "inherited_tags"},
|
||||
caption,
|
||||
bool(embedding),
|
||||
)
|
||||
return {
|
||||
"status": "completed",
|
||||
"clip_id": atom_clip_id,
|
||||
"has_ai_tags": any(v for k, v in ai_tags.items() if k != "inherited_tags" and v),
|
||||
"caption": caption,
|
||||
"embedding_dim": len(embedding) if embedding else 0,
|
||||
}
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
|
||||
@@ -808,17 +808,23 @@ def ingest_asset(job_id: str) -> dict:
|
||||
|
||||
db.commit()
|
||||
|
||||
# ── #1970 素材原子切片:视频 READY 后异步触发,失败不阻断入库 ──
|
||||
# ── #1970 素材原子切片 + #2035 质量评分:视频 READY 后异步触发,失败不阻断入库 ──
|
||||
# atom_clips 未就绪时选片逻辑有内存兜底(compute_fallback_clips)。
|
||||
# quality_score 未计算时选片按 50 分兜底。
|
||||
try:
|
||||
if media_type == "video" and float(asset.duration or 0) > 0:
|
||||
celery_app.send_task(
|
||||
"worker.generate_atom_clips",
|
||||
args=[asset.id],
|
||||
)
|
||||
# #2035: 异步质量评分(不与 atom_clips 链式耦合,独立任务)
|
||||
celery_app.send_task(
|
||||
"worker.calculate_asset_quality",
|
||||
args=[asset.id],
|
||||
)
|
||||
except Exception as atom_err: # noqa: BLE001
|
||||
logger.warning(
|
||||
"触发原子切片任务失败(不影响入库): asset_id=%s err=%s",
|
||||
"触发原子切片/质量评分任务失败(不影响入库): asset_id=%s err=%s",
|
||||
asset.id,
|
||||
atom_err,
|
||||
)
|
||||
|
||||
@@ -83,6 +83,23 @@ class SQLAlchemyAssetAtomClipRepository:
|
||||
models = query.all()
|
||||
return [self._to_domain(m) for m in models]
|
||||
|
||||
def update_caption_embedding(self, clip_id: str, caption: str | None, embedding: list[float] | None = None) -> bool:
|
||||
"""更新片段的 caption 和 embedding 字段。"""
|
||||
upd: dict = {}
|
||||
if caption is not None:
|
||||
upd["caption"] = caption
|
||||
if embedding is not None:
|
||||
upd["embedding"] = embedding
|
||||
if not upd:
|
||||
return False
|
||||
count = (
|
||||
self.session.query(AssetAtomClipModel)
|
||||
.filter(AssetAtomClipModel.id == clip_id)
|
||||
.update(upd)
|
||||
)
|
||||
self.session.commit()
|
||||
return count > 0
|
||||
|
||||
def update_ai_tags(self, clip_id: str, ai_tags: dict) -> bool:
|
||||
"""更新指定片段的 ai_tags 字段."""
|
||||
count = (
|
||||
@@ -118,6 +135,8 @@ class SQLAlchemyAssetAtomClipRepository:
|
||||
clip_index=clip.clip_index,
|
||||
tags=clip.tags,
|
||||
ai_tags=clip.ai_tags,
|
||||
caption=clip.caption,
|
||||
embedding=clip.embedding,
|
||||
scene_change_at=clip.scene_change_at,
|
||||
is_fallback=clip.is_fallback,
|
||||
created_at=clip.created_at or datetime.now(UTC),
|
||||
@@ -132,6 +151,9 @@ class SQLAlchemyAssetAtomClipRepository:
|
||||
duration=model.duration,
|
||||
clip_index=model.clip_index,
|
||||
tags=model.tags or [],
|
||||
ai_tags=getattr(model, "ai_tags", None),
|
||||
caption=getattr(model, "caption", None),
|
||||
embedding=getattr(model, "embedding", None),
|
||||
scene_change_at=model.scene_change_at,
|
||||
is_fallback=model.is_fallback,
|
||||
created_at=model.created_at,
|
||||
|
||||
@@ -2,6 +2,7 @@ from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import (
|
||||
|
||||
JSON,
|
||||
Boolean,
|
||||
Column,
|
||||
@@ -841,6 +842,8 @@ class AssetAtomClipModel(Base):
|
||||
clip_index = Column(Integer, nullable=False)
|
||||
tags = Column(JSON, nullable=False, default=list)
|
||||
ai_tags = Column(JSON, nullable=True, default=None)
|
||||
caption = Column(Text, nullable=True, default=None)
|
||||
embedding = Column(JSON, nullable=True, default=None)
|
||||
scene_change_at = Column(Float, nullable=True)
|
||||
is_fallback = Column(Boolean, nullable=False, default=False)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
@@ -70,6 +70,7 @@ class SharedSettings(BaseSettings):
|
||||
doubao_timeout: int = 30
|
||||
doubao_max_retries: int = 2
|
||||
doubao_vision_model: str = "doubao-1-5-vision-pro-250915"
|
||||
doubao_embedding_model: str = "doubao-embedding-large-text-240915"
|
||||
|
||||
# ── MediaKit (火山引擎 AI 媒体工具) ──────────────────────────────────
|
||||
mediakit_api_key: str = ""
|
||||
|
||||
@@ -37,6 +37,8 @@ class AssetAtomClip:
|
||||
clip_index: int
|
||||
tags: list[str] = field(default_factory=list)
|
||||
ai_tags: dict | None = None
|
||||
caption: str | None = None
|
||||
embedding: list[float] | None = None
|
||||
scene_change_at: float | None = None
|
||||
is_fallback: bool = False
|
||||
created_at: datetime | None = None
|
||||
@@ -67,6 +69,8 @@ class AssetAtomClip:
|
||||
tags: list[str] | None = None,
|
||||
scene_change_at: float | None = None,
|
||||
is_fallback: bool = False,
|
||||
caption: str | None = None,
|
||||
embedding: list[float] | None = None,
|
||||
) -> AssetAtomClip:
|
||||
"""工厂方法:创建一个新的原子片段。"""
|
||||
return cls(
|
||||
@@ -79,4 +83,6 @@ class AssetAtomClip:
|
||||
tags=tags or [],
|
||||
scene_change_at=scene_change_at,
|
||||
is_fallback=is_fallback,
|
||||
caption=caption,
|
||||
embedding=embedding,
|
||||
)
|
||||
|
||||
@@ -35,6 +35,7 @@ def build_vision_prompt() -> str:
|
||||
- action: 动作类型列表(如 "演示", "说话", "操作")
|
||||
- shot: 景别("特写" / "中景" / "远景" 之一)
|
||||
- has_text: 画面中是否有显著文字(true/false)
|
||||
- caption: 一句中文画面描述(10-30字),简洁概括这段视频的内容主体与场景
|
||||
"""
|
||||
return """请分析这段视频片段的关键帧,识别内容并返回 JSON 格式标签。
|
||||
|
||||
@@ -44,7 +45,8 @@ def build_vision_prompt() -> str:
|
||||
"objects": ["物体1", "物体2"],
|
||||
"action": ["动作1"],
|
||||
"shot": "特写|中景|远景",
|
||||
"has_text": true/false
|
||||
"has_text": true/false,
|
||||
"caption": "一句中文描述"
|
||||
}
|
||||
|
||||
规则:
|
||||
@@ -53,6 +55,7 @@ def build_vision_prompt() -> str:
|
||||
- action: 人物或物体正在进行的动作,如"演示"、"说话"、"操作"、"展示"等,1-3个
|
||||
- shot: 景别判断,只能是"特写"、"中景"或"远景"之一
|
||||
- has_text: 画面中是否有显著可读文字(标题、字幕、标语等)
|
||||
- caption: 一句简洁的中文画面描述(10-30字),概括这段视频的主体内容、人物动作和场景,例如"一名女性在办公室中讲解产品展示"
|
||||
|
||||
请只返回 JSON,不要有其他说明文字。"""
|
||||
|
||||
@@ -65,7 +68,7 @@ def parse_vision_response(text: str) -> dict:
|
||||
|
||||
Returns:
|
||||
结构化标签 dict,格式如:
|
||||
{"scene": [...], "objects": [...], "action": [...], "shot": "...", "has_text": bool}
|
||||
{"scene": [...], "objects": [...], "action": [...], "shot": "...", "has_text": bool, "caption": "..."}
|
||||
|
||||
解析失败时返回空 dict。
|
||||
"""
|
||||
@@ -127,6 +130,15 @@ def parse_vision_response(text: str) -> dict:
|
||||
else:
|
||||
result["has_text"] = False
|
||||
|
||||
cap_val = data.get("caption", "")
|
||||
if isinstance(cap_val, str):
|
||||
cap_val = cap_val.strip()
|
||||
if len(cap_val) > 60:
|
||||
cap_val = cap_val[:60]
|
||||
else:
|
||||
cap_val = ""
|
||||
result["caption"] = cap_val
|
||||
|
||||
return result
|
||||
|
||||
|
||||
@@ -237,14 +249,14 @@ def tag_atom_clip(
|
||||
Returns:
|
||||
结构化标签 dict,格式如:
|
||||
{"scene": [...], "objects": [...], "action": [...], "shot": "...",
|
||||
"has_text": bool, "inherited_tags": [...]}
|
||||
"has_text": bool, "caption": "...", "inherited_tags": [...]}
|
||||
"""
|
||||
inherited = list(getattr(clip, "tags", []) or [])
|
||||
|
||||
# 检查 DoubaoClient 是否可用
|
||||
if not getattr(doubao_client, "is_available", False):
|
||||
logger.info("DoubaoClient 不可用,跳过 AI 标签: clip_id=%s", getattr(clip, "id", ""))
|
||||
return {"inherited_tags": inherited}
|
||||
return {"scene": [], "objects": [], "action": [], "shot": "", "has_text": False, "caption": "", "inherited_tags": inherited}
|
||||
|
||||
# 提取帧图片
|
||||
frame_urls: Optional[list[str]] = None
|
||||
@@ -261,7 +273,7 @@ def tag_atom_clip(
|
||||
|
||||
if not frame_urls:
|
||||
logger.warning("帧提取失败,跳过 AI 标签: clip_id=%s", getattr(clip, "id", ""))
|
||||
return {"inherited_tags": inherited}
|
||||
return {"scene": [], "objects": [], "action": [], "shot": "", "has_text": False, "caption": "", "inherited_tags": inherited}
|
||||
|
||||
# 调用视觉 API
|
||||
prompt = build_vision_prompt()
|
||||
@@ -275,17 +287,17 @@ def tag_atom_clip(
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("视觉 API 调用异常: clip_id=%s error=%s", getattr(clip, "id", ""), e)
|
||||
return {"inherited_tags": inherited}
|
||||
return {"scene": [], "objects": [], "action": [], "shot": "", "has_text": False, "caption": "", "inherited_tags": inherited}
|
||||
|
||||
if not response_text:
|
||||
logger.warning("视觉 API 返回空: clip_id=%s", getattr(clip, "id", ""))
|
||||
return {"inherited_tags": inherited}
|
||||
return {"scene": [], "objects": [], "action": [], "shot": "", "has_text": False, "caption": "", "inherited_tags": inherited}
|
||||
|
||||
# 解析标签
|
||||
ai_tags = parse_vision_response(response_text)
|
||||
if not ai_tags:
|
||||
logger.warning("标签解析失败: clip_id=%s response=%s", getattr(clip, "id", ""), response_text[:200])
|
||||
return {"inherited_tags": inherited}
|
||||
return {"scene": [], "objects": [], "action": [], "shot": "", "has_text": False, "caption": "", "inherited_tags": inherited}
|
||||
|
||||
# 合并 inherited_tags
|
||||
ai_tags["inherited_tags"] = inherited
|
||||
|
||||
@@ -245,16 +245,39 @@ def pick_narrative_assets(
|
||||
clip_ai_tags_by_asset=clip_ai_tags_by_asset,
|
||||
)
|
||||
|
||||
# #2035:把文案标签与聚合的素材级 ai_tags 透传给 smart_select_assets,
|
||||
# 让 smart 评分维度(ai_semantic)在叙事模式内部兜底/补位时同样生效。
|
||||
wanted_norm = _normalize_tags(script_tags)
|
||||
asset_ai_tags: dict[str, dict] = {}
|
||||
if clip_ai_tags_by_asset:
|
||||
for aid, clips in clip_ai_tags_by_asset.items():
|
||||
agg: dict = {"scene": [], "objects": [], "action": []}
|
||||
for clip_tags in clips or []:
|
||||
if not isinstance(clip_tags, dict):
|
||||
continue
|
||||
for key in ("scene", "objects", "action"):
|
||||
for v in clip_tags.get(key) or []:
|
||||
v = str(v).strip()
|
||||
if v and v not in agg[key]:
|
||||
agg[key].append(v)
|
||||
asset_ai_tags[aid] = agg
|
||||
|
||||
smart_kwargs = dict(
|
||||
kind="video",
|
||||
rng=rng,
|
||||
script_tags=wanted_norm if wanted_norm else None,
|
||||
ai_tags_by_asset=asset_ai_tags if asset_ai_tags else None,
|
||||
)
|
||||
|
||||
need = limit if (limit is not None and limit > 0) else None
|
||||
|
||||
if not matched:
|
||||
# 完全降级:与改造前随机混剪同一逻辑
|
||||
return [r.asset for r in smart_select_assets(assets, kind="video", limit=need, rng=rng)]
|
||||
return [r.asset for r in smart_select_assets(assets, limit=need, **smart_kwargs)]
|
||||
|
||||
picked = [r.asset for r in smart_select_assets(matched, kind="video", limit=need, rng=rng)]
|
||||
picked = [r.asset for r in smart_select_assets(matched, limit=need, **smart_kwargs)]
|
||||
if need is not None and len(picked) < need and unmatched:
|
||||
rest_need = need - len(picked)
|
||||
picked.extend(r.asset for r in smart_select_assets(unmatched, kind="video", limit=rest_need, rng=rng))
|
||||
picked.extend(r.asset for r in smart_select_assets(unmatched, limit=rest_need, **smart_kwargs))
|
||||
elif need is None:
|
||||
picked.extend(r.asset for r in smart_select_assets(unmatched, kind="video", rng=rng))
|
||||
picked.extend(r.asset for r in smart_select_assets(unmatched, **smart_kwargs))
|
||||
return picked
|
||||
|
||||
@@ -40,6 +40,14 @@ def _get_enum_value(obj: Any, attr: str) -> str:
|
||||
return val.value if hasattr(val, "value") else str(val)
|
||||
|
||||
|
||||
def normalize_tag(tag) -> str:
|
||||
"""标准化标签:去两端空白、小写;非 str 转 str。返回空串表示应丢弃。"""
|
||||
if tag is None:
|
||||
return ""
|
||||
s = str(tag).strip().lower()
|
||||
return s
|
||||
|
||||
|
||||
def _duration_bucket(duration: float | None) -> str:
|
||||
"""将素材时长分为 3 档:short(<10s) / medium(10-30s) / long(>30s)。"""
|
||||
if duration is None or duration <= 0:
|
||||
@@ -54,14 +62,21 @@ def _duration_bucket(duration: float | None) -> str:
|
||||
def score_asset(
|
||||
asset: Any,
|
||||
now: datetime | None = None,
|
||||
script_tags: set | None = None,
|
||||
ai_tags_by_asset: dict | None = None,
|
||||
) -> tuple[float, dict[str, float]]:
|
||||
"""为单个素材计算综合得分(0-100)。
|
||||
|
||||
维度权重:
|
||||
- quality_score (40%):素材质量分(0-100),无质量分按 50 计
|
||||
- duration_fitness (30%):时长适配度,5-30s 为最优区间
|
||||
- recency (20%):新鲜度,30 天内衰减
|
||||
- unused_bonus (10%):未被使用过的素材加分
|
||||
维度权重(#2035 加入 AI 语义匹配维度):
|
||||
- quality_score (30%):素材质量分(0-100),无质量分按 50 计
|
||||
- duration_fitness (25%):时长适配度,5-30s 为最优区间
|
||||
- recency (15%):新鲜度,30 天内衰减
|
||||
- unused (10%):未被/少被使用过的素材加分
|
||||
- ai_semantic (20%):AI 标签(scene/objects/action)与文案标签重合度;无数据给 50 中性分
|
||||
|
||||
Args:
|
||||
script_tags: 标准化后的文案标签集合,用于 AI 语义匹配维度打分。
|
||||
ai_tags_by_asset: asset_id → ai_tags dict 映射,ai_tags 含 scene/objects/action 字段。
|
||||
|
||||
Returns:
|
||||
(total_score, breakdown_dict)
|
||||
@@ -73,7 +88,7 @@ def score_asset(
|
||||
|
||||
# 1. 质量分 (0-100) → 权重 40%
|
||||
raw_quality = asset.quality_score if asset.quality_score is not None else 50.0
|
||||
quality_component = raw_quality * 0.4
|
||||
quality_component = raw_quality * 0.30
|
||||
breakdown["quality"] = round(quality_component, 2)
|
||||
|
||||
# 2. 时长适配度 (0-100) → 权重 30%
|
||||
@@ -90,7 +105,7 @@ def score_asset(
|
||||
# >30s: 指数衰减,60s 时约 50 分
|
||||
duration_fitness = 100.0 * math.exp(-0.02 * (duration - 30))
|
||||
duration_fitness = max(duration_fitness, 10.0)
|
||||
duration_component = duration_fitness * 0.3
|
||||
duration_component = duration_fitness * 0.25
|
||||
breakdown["duration"] = round(duration_component, 2)
|
||||
|
||||
# 3. 新鲜度 (0-100) → 权重 20%
|
||||
@@ -103,7 +118,7 @@ def score_asset(
|
||||
created_at = created_at.replace(tzinfo=UTC)
|
||||
age_days = max(0, (now - created_at).total_seconds() / 86400)
|
||||
recency = 100.0 * math.exp(-0.05 * age_days) # ~14天半衰期
|
||||
recency_component = recency * 0.2
|
||||
recency_component = recency * 0.15
|
||||
breakdown["recency"] = round(recency_component, 2)
|
||||
|
||||
# 4. 未使用偏好 (0-100) → 权重 10%
|
||||
@@ -121,7 +136,34 @@ def score_asset(
|
||||
unused_component = unused_score * 0.1
|
||||
breakdown["unused"] = round(unused_component, 2)
|
||||
|
||||
total = quality_component + duration_component + recency_component + unused_component
|
||||
# 5. AI 语义匹配 (0-100) → 权重 20%
|
||||
if script_tags and ai_tags_by_asset:
|
||||
asset_ai = ai_tags_by_asset.get(getattr(asset, "id", "")) or {}
|
||||
ai_terms: set = set()
|
||||
for key in ("scene", "objects", "action"):
|
||||
vals = asset_ai.get(key) or []
|
||||
if isinstance(vals, list):
|
||||
for v in vals:
|
||||
norm = normalize_tag(v)
|
||||
if norm:
|
||||
ai_terms.add(norm)
|
||||
if ai_terms:
|
||||
norm_script = {normalize_tag(t) for t in script_tags if normalize_tag(t)}
|
||||
overlap = ai_terms & norm_script
|
||||
union = ai_terms | norm_script
|
||||
ratio = (len(overlap) / len(union)) if union else 0.0
|
||||
if overlap:
|
||||
ai_score = 50.0 + 50.0 * ratio
|
||||
else:
|
||||
ai_score = 20.0
|
||||
else:
|
||||
ai_score = 50.0
|
||||
else:
|
||||
ai_score = 50.0
|
||||
ai_component = ai_score * 0.20
|
||||
breakdown["ai_semantic"] = round(ai_component, 2)
|
||||
|
||||
total = quality_component + duration_component + recency_component + unused_component + ai_component
|
||||
return round(total, 2), breakdown
|
||||
|
||||
|
||||
@@ -132,6 +174,8 @@ def smart_select_assets(
|
||||
kind: str | None = None,
|
||||
now: datetime | None = None,
|
||||
rng: random.Random | None = None,
|
||||
script_tags: set | None = None,
|
||||
ai_tags_by_asset: dict | None = None,
|
||||
) -> list[SmartMatchResult]:
|
||||
"""从素材列表中智能选取素材。
|
||||
|
||||
@@ -159,7 +203,7 @@ def smart_select_assets(
|
||||
# Step 3: 评分
|
||||
scored: list[SmartMatchResult] = []
|
||||
for a in ready_assets:
|
||||
total, breakdown = score_asset(a, now=now)
|
||||
total, breakdown = score_asset(a, now=now, script_tags=script_tags, ai_tags_by_asset=ai_tags_by_asset)
|
||||
scored.append(SmartMatchResult(asset=a, score=total, breakdown=breakdown))
|
||||
|
||||
# Step 4: 按「得分 + 随机噪声」降序排序
|
||||
|
||||
@@ -40,6 +40,46 @@ class DoubaoClient:
|
||||
self.vision_model: str = settings.doubao_vision_model
|
||||
|
||||
@property
|
||||
|
||||
def embed_text(self, text: str, timeout: int | None = None) -> list[float] | None:
|
||||
"""调用豆包文本 Embedding API,返回浮点向量;失败返回 None。"""
|
||||
if not self.is_available or not text or not text.strip():
|
||||
return None
|
||||
|
||||
url = f"{self.base_url}/embeddings"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
payload: dict[str, Any] = {
|
||||
"model": getattr(self, "embedding_model", None) or "doubao-embedding-large-text-240915",
|
||||
"input": text.strip(),
|
||||
"encoding_format": "float",
|
||||
}
|
||||
|
||||
req_timeout = timeout or self.timeout
|
||||
last_error: Exception | None = None
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
resp = httpx.post(url, headers=headers, json=payload, timeout=req_timeout)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
emb_list = data.get("data") or []
|
||||
if emb_list and isinstance(emb_list, list):
|
||||
vec = emb_list[0].get("embedding")
|
||||
if isinstance(vec, list) and vec:
|
||||
return [float(x) for x in vec]
|
||||
logger.warning("embedding 返回结构异常: %s", str(data)[:200])
|
||||
return None
|
||||
except Exception as e:
|
||||
last_error = e
|
||||
if attempt < self.max_retries:
|
||||
wait = 0.5 * (2**attempt)
|
||||
logger.warning("豆包 Embedding 调用失败,%.1fs 后重试 (%d/%d): %s", wait, attempt+1, self.max_retries+1, e)
|
||||
time.sleep(wait)
|
||||
logger.error("豆包 Embedding 调用最终失败: %s", last_error)
|
||||
return None
|
||||
|
||||
def is_available(self) -> bool:
|
||||
"""是否可用(配置了 API Key)."""
|
||||
return bool(self.api_key)
|
||||
|
||||
@@ -210,7 +210,12 @@ class TestTagAtomClip:
|
||||
doubao_client=fake_doubao,
|
||||
)
|
||||
|
||||
assert result == {"inherited_tags": ["tag1", "tag2"]}
|
||||
assert result["inherited_tags"] == ["tag1", "tag2"]
|
||||
assert result.get("caption", "") == ""
|
||||
assert result.get("scene", []) == []
|
||||
assert result.get("objects", []) == []
|
||||
assert result.get("action", []) == []
|
||||
assert "inherited_tags" in result
|
||||
assert len(fake_doubao.vision_calls) == 0
|
||||
|
||||
def test_mediakit_unavailable_no_ffmpeg(self):
|
||||
@@ -227,7 +232,12 @@ class TestTagAtomClip:
|
||||
)
|
||||
|
||||
# 没有 ffmpeg 的情况下,帧提取失败
|
||||
assert result == {"inherited_tags": ["tag1", "tag2"]}
|
||||
assert result["inherited_tags"] == ["tag1", "tag2"]
|
||||
assert result.get("caption", "") == ""
|
||||
assert result.get("scene", []) == []
|
||||
assert result.get("objects", []) == []
|
||||
assert result.get("action", []) == []
|
||||
assert "inherited_tags" in result
|
||||
|
||||
def test_vision_api_error_returns_inherited(self):
|
||||
"""视觉 API 抛异常 → 降级 inherited_tags."""
|
||||
@@ -242,7 +252,12 @@ class TestTagAtomClip:
|
||||
mediakit_client=fake_mediakit,
|
||||
)
|
||||
|
||||
assert result == {"inherited_tags": ["tag1", "tag2"]}
|
||||
assert result["inherited_tags"] == ["tag1", "tag2"]
|
||||
assert result.get("caption", "") == ""
|
||||
assert result.get("scene", []) == []
|
||||
assert result.get("objects", []) == []
|
||||
assert result.get("action", []) == []
|
||||
assert "inherited_tags" in result
|
||||
|
||||
def test_vision_api_empty_response(self):
|
||||
"""视觉 API 返回空 → 降级 inherited_tags."""
|
||||
@@ -257,7 +272,12 @@ class TestTagAtomClip:
|
||||
mediakit_client=fake_mediakit,
|
||||
)
|
||||
|
||||
assert result == {"inherited_tags": ["tag1", "tag2"]}
|
||||
assert result["inherited_tags"] == ["tag1", "tag2"]
|
||||
assert result.get("caption", "") == ""
|
||||
assert result.get("scene", []) == []
|
||||
assert result.get("objects", []) == []
|
||||
assert result.get("action", []) == []
|
||||
assert "inherited_tags" in result
|
||||
|
||||
def test_vision_api_invalid_json_response(self):
|
||||
"""视觉 API 返回无效 JSON → 降级 inherited_tags."""
|
||||
@@ -272,7 +292,12 @@ class TestTagAtomClip:
|
||||
mediakit_client=fake_mediakit,
|
||||
)
|
||||
|
||||
assert result == {"inherited_tags": ["tag1", "tag2"]}
|
||||
assert result["inherited_tags"] == ["tag1", "tag2"]
|
||||
assert result.get("caption", "") == ""
|
||||
assert result.get("scene", []) == []
|
||||
assert result.get("objects", []) == []
|
||||
assert result.get("action", []) == []
|
||||
assert "inherited_tags" in result
|
||||
|
||||
def test_clip_with_empty_tags(self):
|
||||
"""空素材标签 → inherited_tags 为空列表."""
|
||||
@@ -285,7 +310,8 @@ class TestTagAtomClip:
|
||||
doubao_client=fake_doubao,
|
||||
)
|
||||
|
||||
assert result == {"inherited_tags": []}
|
||||
assert result["inherited_tags"] == []
|
||||
assert result.get("caption", "") == ""
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -0,0 +1,257 @@
|
||||
"""#2035 语义标签增强 / 质量评分 / AI选片 单测。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import random
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime, timedelta
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.asset_atom_clip import AssetAtomClip
|
||||
from packages.domain.atom_clip_tagger import parse_vision_response
|
||||
from packages.domain.narrative_match import (
|
||||
_compute_ai_score,
|
||||
_extract_ai_tag_names,
|
||||
match_assets_by_script_tags,
|
||||
pick_narrative_assets,
|
||||
)
|
||||
from packages.domain.smart_match import score_asset, smart_select_assets
|
||||
|
||||
|
||||
# ── helpers ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeAsset:
|
||||
id: str
|
||||
tag_ids: list[str] = field(default_factory=list)
|
||||
tags: list[str] = field(default_factory=list)
|
||||
status: object = None
|
||||
file_type: str = "video"
|
||||
duration: float = 10.0
|
||||
quality_score: float | None = 50.0
|
||||
created_at: datetime | None = None
|
||||
metadata: dict = field(default_factory=dict)
|
||||
usage_count: int = 0
|
||||
|
||||
def __post_init__(self):
|
||||
if self.status is None:
|
||||
class _S:
|
||||
value = "ready"
|
||||
self.status = _S()
|
||||
if self.created_at is None:
|
||||
self.created_at = datetime.now(UTC) - timedelta(days=1)
|
||||
|
||||
|
||||
# ── parse_vision_response: caption 提取 ─────────────────────────
|
||||
|
||||
|
||||
class TestParseVisionResponseCaption:
|
||||
def test_extracts_caption(self):
|
||||
text = '{"scene":["办公室"],"objects":["电脑","人"],"action":["说话"],"shot":"中景","has_text":false,"caption":"职场女性在办公室讲解产品功能"}'
|
||||
result = parse_vision_response(text)
|
||||
assert result["caption"] == "职场女性在办公室讲解产品功能"
|
||||
assert result["has_text"] is False
|
||||
assert result["scene"] == ["办公室"]
|
||||
|
||||
def test_caption_truncated_at_60(self):
|
||||
long = "A" * 100
|
||||
text = '{"scene":[],"objects":[],"action":[],"shot":"中景","has_text":false,"caption":"' + long + '"}'
|
||||
result = parse_vision_response(text)
|
||||
assert len(result["caption"]) == 60
|
||||
|
||||
def test_missing_caption_defaults_empty(self):
|
||||
text = '{"scene":[],"objects":[],"action":[],"shot":"特写","has_text":true}'
|
||||
result = parse_vision_response(text)
|
||||
assert result["caption"] == ""
|
||||
|
||||
def test_empty_input_returns_empty_dict(self):
|
||||
assert parse_vision_response("") == {}
|
||||
assert parse_vision_response(None) == {}
|
||||
|
||||
|
||||
# ── AI 标签提取 ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestExtractAiTagNames:
|
||||
def test_extracts_scene_objects_action(self):
|
||||
tags = {"scene": ["办公室"], "objects": ["电脑", "杯子"], "action": ["说话"], "shot": "中景", "has_text": False}
|
||||
names = _extract_ai_tag_names(tags)
|
||||
assert "办公室" in names
|
||||
assert "电脑" in names
|
||||
assert "杯子" in names
|
||||
assert "说话" in names
|
||||
assert "中景" not in names # shot 不参与匹配
|
||||
|
||||
def test_empty_tags(self):
|
||||
assert _extract_ai_tag_names({}) == set()
|
||||
assert _extract_ai_tag_names({"scene": []}) == set()
|
||||
|
||||
|
||||
# ── _compute_ai_score ───────────────────────────────────────────
|
||||
|
||||
|
||||
class TestComputeAiScore:
|
||||
def test_basic_hit(self):
|
||||
clip_map = {"a1": [{"scene": ["工厂"], "objects": ["产品"], "action": ["演示"]}]}
|
||||
score = _compute_ai_score("a1", {"工厂", "演示"}, clip_map)
|
||||
# 2 hits * weight 2.0 = 4.0
|
||||
assert score == 4.0
|
||||
|
||||
def test_no_hit(self):
|
||||
clip_map = {"a1": [{"scene": ["户外"], "objects": [], "action": []}]}
|
||||
assert _compute_ai_score("a1", {"办公室"}, clip_map) == 0.0
|
||||
|
||||
def test_no_clip_map(self):
|
||||
assert _compute_ai_score("a1", {"工厂"}, None) == 0.0
|
||||
assert _compute_ai_score("a1", set(), {"a1": [{"scene": ["x"]}]}) == 0.0
|
||||
|
||||
def test_best_clip_score_not_sum(self):
|
||||
"""多片段取最高得分,不是累加。"""
|
||||
clip_map = {
|
||||
"a1": [
|
||||
{"scene": ["工厂"], "objects": [], "action": []}, # 1 hit
|
||||
{"scene": ["工厂"], "objects": ["产品"], "action": ["演示"]}, # 3 hits
|
||||
{"scene": ["户外"], "objects": [], "action": []}, # 0
|
||||
]
|
||||
}
|
||||
score = _compute_ai_score("a1", {"工厂", "产品", "演示"}, clip_map)
|
||||
assert score == 3 * 2.0 # best = 6.0, not (1+3+0)*2 = 8.0
|
||||
|
||||
|
||||
# ── match_assets_by_script_tags 接受 clip_ai_tags_by_asset ──────
|
||||
|
||||
|
||||
class TestMatchSplitAiTags:
|
||||
def test_ai_hit_only_puts_in_matched(self):
|
||||
"""素材无人工标签,但 AI 标签命中 → 命中池。"""
|
||||
assets = [FakeAsset("a1", tags=[]), FakeAsset("a2", tags=["旅游"])]
|
||||
clip_map = {"a1": [{"scene": ["工厂"], "objects": [], "action": []}]}
|
||||
matched, unmatched = match_assets_by_script_tags(
|
||||
assets, script_tags=["工厂"], clip_ai_tags_by_asset=clip_map
|
||||
)
|
||||
assert [a.id for a in matched] == ["a1"]
|
||||
assert [a.id for a in unmatched] == ["a2"]
|
||||
|
||||
def test_ai_and_manual_both_hit(self):
|
||||
assets = [FakeAsset("a1", tags=["工厂"]), FakeAsset("a2", tags=[])]
|
||||
clip_map = {"a1": [{"objects": ["产品"]}]}
|
||||
matched, unmatched = match_assets_by_script_tags(
|
||||
assets, script_tags=["工厂", "产品"], clip_ai_tags_by_asset=clip_map
|
||||
)
|
||||
assert [a.id for a in matched] == ["a1"]
|
||||
# a2 无人标签也无AI命中 → unmatched
|
||||
assert [a.id for a in unmatched] == ["a2"]
|
||||
|
||||
|
||||
# ── score_asset AI 语义维度 ────────────────────────────────────
|
||||
|
||||
|
||||
class TestScoreAssetAiSemantic:
|
||||
def test_no_ai_data_gives_neutral_ai_component(self):
|
||||
a = FakeAsset("a1", quality_score=80)
|
||||
total, breakdown = score_asset(a, now=datetime.now(UTC))
|
||||
# ai_semantic 中性分 50 * 0.20 = 10
|
||||
assert breakdown["ai_semantic"] == 10.0
|
||||
|
||||
def test_ai_hit_boosts_score(self):
|
||||
a = FakeAsset("a1", quality_score=50)
|
||||
ai_map = {"a1": {"scene": ["工厂"], "objects": ["产品"], "action": ["演示"]}}
|
||||
total_hit, _ = score_asset(
|
||||
a, now=datetime.now(UTC), script_tags={"工厂", "产品"}, ai_tags_by_asset=ai_map
|
||||
)
|
||||
total_miss, _ = score_asset(
|
||||
a, now=datetime.now(UTC), script_tags={"旅游"}, ai_tags_by_asset=ai_map
|
||||
)
|
||||
total_neutral, _ = score_asset(a, now=datetime.now(UTC))
|
||||
assert total_hit > total_neutral
|
||||
assert total_neutral > total_miss
|
||||
|
||||
def test_weights_sum_to_100(self):
|
||||
a = FakeAsset("a1", quality_score=100, duration=15)
|
||||
a.created_at = datetime.now(UTC)
|
||||
a.metadata = {"generation_use_count": 0}
|
||||
_, bd = score_asset(a, now=datetime.now(UTC))
|
||||
# 满分素材:quality=30, duration=25, recency=~15 (new), unused=10, ai=10(neutral)
|
||||
# 总和应该 ~90
|
||||
assert 85 <= sum(bd.values()) <= 100.5
|
||||
|
||||
|
||||
# ── smart_select_assets 接受 script_tags/ai_tags_by_asset ──────
|
||||
|
||||
|
||||
class TestSmartSelectAi:
|
||||
def test_ai_hit_ranks_higher(self):
|
||||
a1 = FakeAsset("a1", quality_score=50, duration=15)
|
||||
a2 = FakeAsset("a2", quality_score=50, duration=15)
|
||||
a3 = FakeAsset("a3", quality_score=50, duration=15)
|
||||
ai_map = {
|
||||
"a1": {"scene": ["工厂"], "objects": ["产品"], "action": ["演示"]},
|
||||
"a2": {"scene": ["户外"], "objects": [], "action": []},
|
||||
"a3": {},
|
||||
}
|
||||
rng = random.Random(42)
|
||||
results = smart_select_assets(
|
||||
[a1, a2, a3],
|
||||
kind="video",
|
||||
rng=rng,
|
||||
script_tags={"工厂", "产品", "演示"},
|
||||
ai_tags_by_asset=ai_map,
|
||||
)
|
||||
assert results[0].asset.id == "a1" # AI 命中应排第一
|
||||
|
||||
def test_without_ai_params_works_as_before(self):
|
||||
a1 = FakeAsset("a1", quality_score=80, duration=15)
|
||||
a2 = FakeAsset("a2", quality_score=40, duration=15)
|
||||
rng = random.Random(0)
|
||||
results = smart_select_assets([a1, a2], kind="video", rng=rng)
|
||||
assert results[0].asset.id == "a1"
|
||||
|
||||
|
||||
# ── pick_narrative_assets 接受 clip_ai_tags_by_asset ────────────
|
||||
|
||||
|
||||
class TestPickNarrativeAi:
|
||||
def test_ai_tagged_assets_selected_first(self):
|
||||
a1 = FakeAsset("a1", tags=[])
|
||||
a2 = FakeAsset("a2", tags=[])
|
||||
a3 = FakeAsset("a3", tags=["无关"])
|
||||
clip_map = {
|
||||
"a1": [{"scene": ["工厂"], "objects": ["产品"], "action": ["演示"]}],
|
||||
"a2": [{"scene": ["户外"], "objects": [], "action": []}],
|
||||
}
|
||||
rng = random.Random(0)
|
||||
picked = pick_narrative_assets(
|
||||
[a1, a2, a3],
|
||||
script_tags=["工厂", "产品"],
|
||||
tag_names_by_id={},
|
||||
clip_ai_tags_by_asset=clip_map,
|
||||
rng=rng,
|
||||
limit=2,
|
||||
)
|
||||
assert picked[0].id == "a1" # a1 命中 AI 标签应在首位
|
||||
assert {a.id for a in picked} == {"a1", "a3"} or {a.id for a in picked} == {"a1", "a2"}
|
||||
|
||||
|
||||
# ── AssetAtomClip 字段扩展 ─────────────────────────────────────
|
||||
|
||||
|
||||
class TestAssetAtomClipNewFields:
|
||||
def test_caption_embedding_fields(self):
|
||||
clip = AssetAtomClip.create(
|
||||
asset_id="a1",
|
||||
start_time=0,
|
||||
end_time=5,
|
||||
clip_index=0,
|
||||
tags=[],
|
||||
caption="测试画面描述",
|
||||
embedding=[0.1, 0.2, 0.3],
|
||||
)
|
||||
assert clip.caption == "测试画面描述"
|
||||
assert clip.embedding == [0.1, 0.2, 0.3]
|
||||
assert clip.ai_tags is None
|
||||
|
||||
def test_default_fields_none(self):
|
||||
clip = AssetAtomClip.create("a1", 0, 5, 0)
|
||||
assert clip.caption is None
|
||||
assert clip.embedding is None
|
||||
@@ -93,31 +93,31 @@ class TestScoreAsset:
|
||||
def test_no_quality_score_defaults_to_50(self):
|
||||
asset = FakeAsset(id="a1", quality_score=None, duration=15)
|
||||
score, breakdown = score_asset(asset, now=NOW)
|
||||
# quality component should be 50 * 0.4 = 20
|
||||
assert breakdown["quality"] == pytest.approx(20.0, abs=0.1)
|
||||
# quality component should be 50 * 0.30 = 15
|
||||
assert breakdown["quality"] == pytest.approx(15.0, abs=0.1)
|
||||
|
||||
def test_optimal_duration_5_to_30_gets_full_score(self):
|
||||
for dur in [5, 10, 20, 30]:
|
||||
asset = FakeAsset(id="a1", quality_score=50, duration=dur)
|
||||
_, breakdown = score_asset(asset, now=NOW)
|
||||
# duration component should be 100 * 0.3 = 30
|
||||
assert breakdown["duration"] == pytest.approx(30.0, abs=0.1)
|
||||
# duration component should be 100 * 0.25 = 25
|
||||
assert breakdown["duration"] == pytest.approx(25.0, abs=0.1)
|
||||
|
||||
def test_short_duration_below_5s_penalized(self):
|
||||
asset = FakeAsset(id="a1", quality_score=50, duration=2)
|
||||
_, breakdown = score_asset(asset, now=NOW)
|
||||
assert breakdown["duration"] < 30.0
|
||||
assert breakdown["duration"] < 25.0 # below max duration score
|
||||
|
||||
def test_long_duration_above_30s_penalized(self):
|
||||
asset = FakeAsset(id="a1", quality_score=50, duration=120)
|
||||
_, breakdown = score_asset(asset, now=NOW)
|
||||
assert breakdown["duration"] < 30.0
|
||||
assert breakdown["duration"] < 25.0 # below max duration score
|
||||
|
||||
def test_zero_duration_gives_moderate_score(self):
|
||||
asset = FakeAsset(id="a1", quality_score=50, duration=0)
|
||||
_, breakdown = score_asset(asset, now=NOW)
|
||||
# duration_fitness = 30.0, component = 30 * 0.3 = 9
|
||||
assert breakdown["duration"] == pytest.approx(9.0, abs=0.1)
|
||||
assert breakdown["duration"] == pytest.approx(7.5, abs=0.1)
|
||||
|
||||
def test_unused_asset_gets_full_bonus(self):
|
||||
asset = FakeAsset(id="a1", quality_score=50, duration=15, metadata={})
|
||||
@@ -133,17 +133,17 @@ class TestScoreAsset:
|
||||
"""int() conversion of non-numeric metadata should not raise, should default to 0."""
|
||||
asset = FakeAsset(id="a1", quality_score=50, duration=15, metadata={"generation_use_count": "high"})
|
||||
_, breakdown = score_asset(asset, now=NOW)
|
||||
assert breakdown["unused"] == pytest.approx(10.0, abs=0.1) # use_count=0 → unused_score=100 → 100*0.1=10
|
||||
assert breakdown["unused"] == pytest.approx(10.0, abs=0.1) # use_count=0 → unused_score=100 → 100*0.10=10
|
||||
|
||||
def test_recent_asset_scores_higher_recency(self):
|
||||
asset = FakeAsset(id="a1", quality_score=50, duration=15, created_at=NOW - timedelta(days=1))
|
||||
_, breakdown = score_asset(asset, now=NOW)
|
||||
assert breakdown["recency"] > 15 # > 75% of max 20
|
||||
assert breakdown["recency"] > 11 # > 75% of max 15
|
||||
|
||||
def test_old_asset_scores_lower_recency(self):
|
||||
asset = FakeAsset(id="a1", quality_score=50, duration=15, created_at=NOW - timedelta(days=60))
|
||||
_, breakdown = score_asset(asset, now=NOW)
|
||||
assert breakdown["recency"] < 5 # heavily decayed
|
||||
assert breakdown["recency"] < 4 # heavily decayed
|
||||
|
||||
|
||||
# ── _duration_bucket tests ───────────────────────────────────────────────────
|
||||
@@ -245,7 +245,7 @@ class TestSmartSelectAssets:
|
||||
assert len(results) == 1
|
||||
r = results[0]
|
||||
assert r.score > 0
|
||||
assert set(r.breakdown.keys()) == {"quality", "duration", "recency", "unused"}
|
||||
assert set(r.breakdown.keys()) >= {"quality", "duration", "recency", "unused", "ai_semantic"}
|
||||
|
||||
def test_image_assets_can_be_selected(self):
|
||||
assets = [
|
||||
|
||||
Reference in New Issue
Block a user