feat(#2035): 语义标签增强 + 质量评分自动计算 + atom_clip caption/embedding #2036

Merged
auto-approve-bot merged 6 commits from feature/semantic-tags-quality-2035 into develop 2026-09-25 12:44:39 +08:00
23 changed files with 1327 additions and 102 deletions
@@ -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")
+84 -2
View File
@@ -46,6 +46,37 @@ from packages.application import (
ListGeneratedVideosByTaskUseCase,
)
from packages.domain.smart_match import smart_select_assets
# #2035:文案关键词 → 素材分类 映射表(用于 smart_match category_match 维度)
# AssetClassification 枚举: scenic / product / person / animal / food / tech / sport / music / other
_CATEGORY_KEYWORDS: dict[str, set[str]] = {
"scenic": {"风景", "自然", "山水", "大海", "天空", "日落", "日出", "森林", "城市", "建筑", "夜景", "街道", "公园", "景区", "旅行", "旅游", "户外"},
"product": {"产品", "商品", "展示", "演示", "开箱", "评测", "好物", "推荐", "种草", "购物", "电商", "带货", "品牌", "广告", "包装"},
"person": {"人物", "人物采访", "对话", "说话", "讲解", "演讲", "采访", "聊天", "开会", "工作", "办公室", "团队", "员工", "老板", "女性", "男性", "美女", "帅哥"},
"animal": {"动物", "宠物", "狗", "猫", "鸟", "鱼", "马", "牛", "羊", "野生动物", "动物园"},
"food": {"美食", "食物", "餐饮", "餐厅", "做饭", "烹饪", "厨房", "菜品", "饮料", "水果", "甜点", "蛋糕", "咖啡", "茶", "零食", "吃"},
"tech": {"科技", "数码", "电脑", "手机", "屏幕", "软件", "APP", "互联网", "AI", "人工智能", "机器人", "办公", "程序员", "代码", "屏幕录制"},
"sport": {"运动", "健身", "跑步", "篮球", "足球", "游泳", "瑜伽", "户外", "锻炼", "体育", "比赛", "球场"},
"music": {"音乐", "歌曲", "演唱会", "乐器", "唱歌", "跳舞", "舞蹈", "MV", "演出", "乐队", "钢琴", "吉他", "节奏"},
}
def _infer_expected_categories(script_tags: set[str] | None) -> set[str] | None:
"""从文案标签集合推断期望的素材分类(可能命中多个)。标签为空返回 None。"""
if not script_tags:
return None
matched: set[str] = set()
for cat, kws in _CATEGORY_KEYWORDS.items():
for tag in script_tags:
tag.lower()
for kw in kws:
if kw in tag or tag in kw:
matched.add(cat)
break
if cat in matched:
break
return matched or None
from packages.middleware.points_gate import points_gate
logger = logging.getLogger(__name__)
@@ -134,10 +165,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 +190,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 +235,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 +246,18 @@ 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()}
expected_categories = _infer_expected_categories(norm_script)
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,
expected_categories=expected_categories,
)
return [r.asset.id for r in results]
# 默认 all 模式:返回全部 ready 视频素材
@@ -396,6 +476,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 +491,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(
+1 -2
View File
@@ -3722,7 +3722,6 @@
z-index: -1;
}
/* Cover template selected check */
.xx-cover-template-check {
position: absolute;
@@ -3739,7 +3738,7 @@
font-size: 14px;
font-weight: 700;
z-index: 2;
box-shadow: 0 2px 6px rgba(124,58,237,0.4);
box-shadow: 0 2px 6px rgba(124, 58, 237, 0.4);
}
.xx-cover-template-thumb {
position: relative;
+1
View File
@@ -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",
@@ -445,22 +445,3 @@ def classify_asset_real(video_path: str) -> tuple[str, float]:
except Exception as e:
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,143 @@
"""素材质量评分 Celery 任务 — #2035.
视频素材 READY 入库后异步触发:下载视频到临时文件,运行 FFmpeg+NumPy 质量分析,
将 0-100 总分写入 assets.quality_score 字段。同时复用已下载的视频,调用 AssetAnalyzer
完成 9 类素材分类(写入 asset.metadata.classification / classification_confidence),
供 smart_match 选片打分使用。任一环节失败均不阻断主流程(质量分兜底 50,分类降级 "other")。
任务名:worker.calculate_asset_quality
"""
from __future__ import annotations
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;
4. 复用同一临时文件,调用 AssetAnalyzer.classify() 做 9 类素材分类,
结果写入 asset.metadata.classification / classification_confidence;
如已有分类结果则幂等跳过(避免重复计算)。
失败/非视频/无文件等情况均静默降级,返回 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
# #2035:自动触发 9 类分类(复用已下载的临时文件,避免重复下载)
existing_meta = dict(asset.metadata or {})
existing_classification = existing_meta.get("classification")
classification = None
confidence = None
if not existing_classification or existing_classification == "other":
try:
from worker_app.tasks.asset_analyzer import AssetAnalyzer as _AA
# 重新构造analyzer可能会重复抽帧,但classify()会复用临时帧
_analyzer = _AA(str(local_path), temp_dir=tmp_dir)
_cls_result = _analyzer.classify()
classification = getattr(_cls_result, "category", None) or "other"
confidence = float(getattr(_cls_result, "confidence", 0.0) or 0.0)
if confidence < 0:
confidence = 0.0
if confidence > 1:
confidence = 1.0
existing_meta["classification"] = classification
existing_meta["classification_confidence"] = confidence
asset.classification_status = "completed"
asset.metadata = existing_meta
logger.info(
"[quality_score] asset=%s 自动分类完成: category=%s confidence=%.2f",
asset_id,
classification,
confidence,
)
except Exception as cls_err: # noqa: BLE001
logger.warning(
"[quality_score] asset=%s 自动分类失败(不影响质量分): %s",
asset_id,
cls_err,
)
asset_repo.update(asset)
db.commit()
logger.info(
"[quality_score] asset=%s score=%.1f classification=%s",
asset_id,
total,
classification or existing_classification,
)
return {
"status": "completed",
"asset_id": asset_id,
"quality_score": total,
"classification": classification or existing_classification or "other",
}
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
@@ -1,7 +1,8 @@
"""片段级 AI 标签 Celery 任务 — #1970 智能剪辑流程重构 P2.
为单个 atom_clip 调用视觉 AI 生成结构化标签,并更新到 ai_tags 字段。
失败不阻断流程(降级为仅继承素材标签)。
为单个 atom_clip 调用视觉 AI 生成结构化标签(含 caption),再调用
豆包 embedding 接口为 caption 生成向量,一并写入数据库。
失败不阻断流程(降级为仅继承素材标签 / caption 留空 / embedding 留空)。
任务名:worker.tag_atom_clip
"""
@@ -26,16 +27,16 @@ logger = get_task_logger(__name__)
@celery_app.task(name="worker.tag_atom_clip", bind=True, max_retries=2, default_retry_delay=10)
def tag_atom_clip_task(self, atom_clip_id: str, force: bool = False) -> dict:
"""为单个原子片段生成 AI 标签.
"""为单个原子片段生成 AI 标签 + caption + embedding.
Args:
atom_clip_id: 原子片段 ID。
force: True 时允许覆盖只有 inherited_tags 的降级记录
(视觉 API 曾失败写入的占位标签,#1970)。
已有完整标签(含 has_text)始终跳过,保证幂等。
已有完整标签且有 caption 始终跳过,保证幂等。
Returns:
任务结果 dict:status / clip_id / ai_tags(部分字段)。
任务结果 dict:status / clip_id / has_ai_tags / caption / embedding_dim。
"""
db = SessionLocal()
try:
@@ -46,11 +47,27 @@ def tag_atom_clip_task(self, atom_clip_id: str, force: bool = False) -> dict:
if clip is None:
return {"status": "skipped", "reason": "clip not found", "clip_id": atom_clip_id}
# 已有完整标签则跳过(幂等);force 仅放行缺失 has_text 的降级记录
if clip.ai_tags is not None:
has_real_tags = isinstance(clip.ai_tags, dict) and "has_text" in clip.ai_tags
if has_real_tags or not force:
return {"status": "skipped", "reason": "already tagged", "clip_id": atom_clip_id}
# 幂等:已有任意 ai_tags(含降级占位)则按 force 策略跳过;person_count/text_content 为附加字段不单独触发重跑
# - 无 force:只要 ai_tags 非 None 就跳过(与旧逻辑一致)
# - force=True 且 ai_tags 是完整标签(含 has_text)且 caption 已存在才跳过
existing_tags = clip.ai_tags
has_real_tags = isinstance(existing_tags, dict) and "has_text" in existing_tags
bool(getattr(clip, "caption", None))
if existing_tags is not None:
if not force:
return {
"status": "skipped",
"reason": "already tagged",
"clip_id": atom_clip_id,
}
# force=True:有完整标签(has_text)就跳过;caption 是 #2035 新增的
# 字段,对已有完整标签的历史数据不强制重跑
if has_real_tags:
return {
"status": "skipped",
"reason": "already tagged",
"clip_id": atom_clip_id,
}
# 获取素材信息
asset = asset_repo.find_by_id(clip.asset_id)
@@ -65,7 +82,7 @@ def tag_atom_clip_task(self, atom_clip_id: str, force: bool = False) -> dict:
doubao_client = get_doubao_client()
mediakit_client = get_mediakit_client()
# 调用 tagger
# 调用 tagger(视觉 API → ai_tags + caption)
ai_tags = tag_atom_clip(
clip=clip,
video_url=video_url,
@@ -74,18 +91,54 @@ def tag_atom_clip_task(self, atom_clip_id: str, force: bool = False) -> dict:
storage=storage,
)
# 更新数据库
# 先写入 AI 标签(含 caption 字段在 ai_tags 字典里)
atom_repo.update_ai_tags(atom_clip_id, ai_tags)
# 提取 caption 并生成 embedding(失败降级,不阻断主流程)
caption = (ai_tags or {}).get("caption", "") or ""
embedding = None
try:
if caption.strip() and doubao_client.is_available:
embedding = doubao_client.embed_text(caption)
except Exception as exc: # noqa: BLE001
logger.warning(
"[atom_clip_tagging] clip_id=%s embedding 生成失败,降级为空: %s",
atom_clip_id,
exc,
)
embedding = None
# 写入 caption + embedding(caption 冗余写一次到独立列,便于查询)
try:
atom_repo.update_caption_embedding(atom_clip_id, caption, embedding)
except Exception as exc: # noqa: BLE001
logger.warning(
"[atom_clip_tagging] clip_id=%s caption/embedding 写入失败: %s",
atom_clip_id,
exc,
)
person_count = (ai_tags or {}).get("person_count", 0)
(ai_tags or {}).get("text_content", "") or ""
logger.info(
"[atom_clip_tagging] clip_id=%s ai_tags=%s",
"[atom_clip_tagging] clip_id=%s ai_tags=%s caption=%r person_count=%s has_text=%s embedding_dim=%s",
atom_clip_id,
{k: v for k, v in ai_tags.items() if k != "inherited_tags"},
{k: v for k, v in ai_tags.items() if k not in ("inherited_tags", "caption", "text_content")},
caption,
person_count,
bool((ai_tags or {}).get("has_text")),
len(embedding) if embedding else 0,
)
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),
"has_ai_tags": any(
v for k, v in ai_tags.items() if k not in ("inherited_tags", "caption", "text_content") and v
),
"caption": caption,
"person_count": (ai_tags or {}).get("person_count", 0),
"text_content": (ai_tags or {}).get("text_content", "") or "",
"embedding_dim": len(embedding) if embedding else 0,
}
except Exception as exc:
db.rollback()
@@ -19,6 +19,7 @@ from packages.adapters.sqlalchemy_impl.asset_atom_clip_repository import (
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
from packages.domain.atom_clip_service import compute_atom_clips
from packages.domain.plan_generator_utils import extract_scene_points_from_metadata
from packages.shared.mediakit_client import get_mediakit_client
logger = get_task_logger(__name__)
@@ -56,6 +57,35 @@ def generate_atom_clips(asset_id: str) -> dict:
}
scene_points = extract_scene_points_from_metadata(asset.metadata)
# #2035:metadata 中没有 scene_change_points 时,按需调用 MediaKit 检测
# (templates_editor 路由会主动写 metadata,ingest 流程此前未触发检测导致切点无法对齐)
if not scene_points:
try:
mk = get_mediakit_client()
video_url = getattr(asset, "file_url", "") or ""
if mk.is_available and video_url:
timestamps = mk.detect_scene_changes(video_url)
if timestamps:
scene_points = timestamps
# 持久化到 metadata,避免下次重复检测
new_meta = dict(asset.metadata or {})
new_meta["scene_change_points"] = list(timestamps)
asset.metadata = new_meta
asset_repo.update(asset)
db.commit()
logger.info(
"[atom_clips] asset_id=%s 自动检测到 %d 个场景切换点并写回metadata",
asset_id,
len(timestamps),
)
except Exception as detect_err: # noqa: BLE001
logger.warning(
"[atom_clips] asset_id=%s scene_change自动检测失败,降级为均匀切片: %s",
asset_id,
detect_err,
)
db.rollback() # 回滚metadata写失败,不影响后续切片
# P1 阶段继承素材的标签 ID;片段级语义标签是 P2 功能
tags = list(getattr(asset, "tag_ids", []) or [])
+8 -2
View File
@@ -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,19 @@ 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 +131,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 +147,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,
@@ -841,6 +841,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))
+1
View File
@@ -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 = ""
+6
View File
@@ -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,
)
+112 -18
View File
@@ -23,36 +23,53 @@ from typing import Any, Optional
logger = logging.getLogger(__name__)
# AI 标签结构的键
AI_TAG_KEYS = ("scene", "objects", "action", "shot", "has_text")
AI_TAG_KEYS = ("scene", "objects", "action", "shot", "has_text", "person_count", "text_content", "caption")
def build_vision_prompt() -> str:
"""返回结构化标签提取 prompt.
要求 AI 以 JSON 格式返回片段内容标签,包含:
- scene: 场景类型列表(如 "工厂", "办公室", "户外")
- objects: 出现的物体列表(如 "产品", "手机", "电脑")
- action: 动作类型列表(如 "演示", "说话", "操作")
- scene: 场景类型列表(如 "工厂", "办公室", "户外", "家庭", "商店")
- objects: 画面中出现的主要物体/人物/动物类别,详细列出,常见类别包括:
人物类:"人物"/"男性"/"女性"/"儿童"
食物类:"食物"/"水果"/"饮料"/"菜肴"
电子设备类:"手机"/"电脑"/"笔记本"/"平板"/"电视"/"相机"
交通类:"汽车"/"自行车"/"公交车"/"飞机"
建筑类:"建筑"/"房屋"/"桥梁"/"道路"
动物类:"狗"/"猫"/"鸟"/"马"
其他常见:"桌子"/"椅子"/"书本"/"花草"/"产品"等
尽可能列出所有可识别的主要物体,3-8个
- action: 动作类型列表(如 "演示", "说话", "操作", "行走", "奔跑", "进食")
- shot: 景别("特写" / "中景" / "远景" 之一)
- has_text: 画面中是否有显著文字(true/false)
- has_text: 画面中是否有显著文字(标题/字幕/标语/海报文字)
- person_count: 画面中可见的人数,0/1/2/3(3代表3人及以上)
- text_content: 若 has_text=true,提取画面中最显著的文字内容(不超过30字,概括即可);否则为空字符串
- caption: 一句中文画面描述(15-30字),简洁概括这段视频的人物、动作、场景和主体内容
"""
return """请分析这段视频片段的关键帧,识别内容并返回 JSON 格式标签。
要求返回以下 JSON 结构(严格 JSON,不要添加其他文字):
{
"scene": ["场景1", "场景2"],
"objects": ["物体1", "物体2"],
"objects": ["物体1", "物体2", "物体3"],
"action": ["动作1"],
"shot": "特写|中景|远景",
"has_text": true/false
"has_text": true/false,
"person_count": 0,
"text_content": "",
"caption": "一句中文描述"
}
规则:
- scene: 场景类型,如"工厂"、"办公室"、"户外"、"商店"、"家庭"等,1-3个
- objects: 画面中可见的主要物体,如"产品"、"手机"、"电脑"、"食品"等,1-5个
- action: 人物或物体正在进行的动作,如"演示"、"说话"、"操作"、"展示"等,1-3个
- scene: 场景类型,如"工厂"、"办公室"、"户外"、"商店"、"家庭"、"街道"等,1-3个
- objects: 画面中可见的所有主要物体/人物/动物/食物/设备等,详细列出(3-8个)。人物算作"人物",不要写具体人名。
- action: 人物或物体正在进行的动作,如"演示"、"说话"、"操作"、"展示"、"行走"等,1-3个
- shot: 景别判断,只能是"特写"、"中景"或"远景"之一
- has_text: 画面中是否有显著可读文字(标题、字幕、标语等)
- has_text: 画面中是否有显著可读文字(标题、字幕、标语、海报文字等)
- person_count: 画面中可见的清晰人物数量,0=无人/远景人物不计数,1=1人,2=2人,3=3人及以上
- text_content: 仅当 has_text=true 时填写,提取画面中最显眼的文字内容(不要超过30字);has_text=false 时填空字符串
- caption: 一句简洁的中文画面描述(15-30字),概括主体人物、动作、场景和物体,例如"一名女性在办公室中讲解产品展示,桌上放有笔记本电脑"
请只返回 JSON,不要有其他说明文字。"""
@@ -65,7 +82,7 @@ def parse_vision_response(text: str) -> dict:
Returns:
结构化标签 dict,格式如:
{"scene": [...], "objects": [...], "action": [...], "shot": "...", "has_text": bool}
{"scene": [...], "objects": [...], "action": [...], "shot": "...", "has_text": bool, "person_count": int, "text_content": str, "caption": "..."}
解析失败时返回空 dict。
"""
@@ -127,6 +144,33 @@ 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) > 80:
cap_val = cap_val[:80]
else:
cap_val = ""
result["caption"] = cap_val
# person_count: 0/1/2/3
pc_val = data.get("person_count", 0)
try:
pc = int(pc_val)
result["person_count"] = max(0, min(3, pc))
except (TypeError, ValueError):
result["person_count"] = 0
# text_content: OCR 文字
tc_val = data.get("text_content", "")
if isinstance(tc_val, str):
tc_val = tc_val.strip()
if len(tc_val) > 100:
tc_val = tc_val[:100]
else:
tc_val = ""
result["text_content"] = tc_val
return result
@@ -237,14 +281,24 @@ 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,
"person_count": 0,
"text_content": "",
"caption": "",
"inherited_tags": inherited,
}
# 提取帧图片
frame_urls: Optional[list[str]] = None
@@ -261,7 +315,17 @@ 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,
"person_count": 0,
"text_content": "",
"caption": "",
"inherited_tags": inherited,
}
# 调用视觉 API
prompt = build_vision_prompt()
@@ -275,17 +339,47 @@ 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,
"person_count": 0,
"text_content": "",
"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,
"person_count": 0,
"text_content": "",
"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,
"person_count": 0,
"text_content": "",
"caption": "",
"inherited_tags": inherited,
}
# 合并 inherited_tags
ai_tags["inherited_tags"] = inherited
+28 -5
View File
@@ -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
+83 -11
View File
@@ -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,24 @@ 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,
expected_categories: set[str] | 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 (28%):素材质量分(0-100),无质量分按 50 计
- duration_fitness (22%):时长适配度,5-30s 为最优区间
- recency (12%):新鲜度,30 天内衰减
- unused (8%):未被/少被使用过的素材加分
- ai_semantic (20%):AI 标签(scene/objects/action)与文案标签重合度;无数据给 50 中性分
- category_match (10%):FFmpeg 自动分类结果(scenic/product/person/animal/food/tech/sport/music)
与期望类别重合度;无分类或无期望类别时给 60 中性分
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 +91,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.28
breakdown["quality"] = round(quality_component, 2)
# 2. 时长适配度 (0-100) → 权重 30%
@@ -90,7 +108,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.22
breakdown["duration"] = round(duration_component, 2)
# 3. 新鲜度 (0-100) → 权重 20%
@@ -103,7 +121,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.12
breakdown["recency"] = round(recency_component, 2)
# 4. 未使用偏好 (0-100) → 权重 10%
@@ -118,10 +136,55 @@ def score_asset(
unused_score = 70.0
else:
unused_score = 30.0
unused_component = unused_score * 0.1
unused_component = unused_score * 0.08
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)
# 6. 分类匹配 (0-100) → 权重 10%
asset_meta = getattr(asset, "metadata", None) or {}
asset_cat = (asset_meta.get("classification") or "").strip().lower()
if expected_categories and asset_cat:
norm_expected = {c.strip().lower() for c in expected_categories if c and c.strip()}
if asset_cat == "other":
cat_score = 50.0 # other 类不给额外加分也不扣分
elif asset_cat in norm_expected:
cat_score = 100.0
else:
cat_score = 30.0 # 分类明确但不匹配,略扣分
elif expected_categories:
cat_score = 60.0 # 无分类结果,中性
else:
cat_score = 60.0 # 无期望类别,中性
cat_component = cat_score * 0.10
breakdown["category_match"] = round(cat_component, 2)
total = quality_component + duration_component + recency_component + unused_component + ai_component + cat_component
return round(total, 2), breakdown
@@ -132,6 +195,9 @@ 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,
expected_categories: set[str] | None = None,
) -> list[SmartMatchResult]:
"""从素材列表中智能选取素材。
@@ -159,7 +225,13 @@ 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,
expected_categories=expected_categories,
)
scored.append(SmartMatchResult(asset=a, score=total, breakdown=breakdown))
# Step 4: 按「得分 + 随机噪声」降序排序
+41
View File
@@ -39,6 +39,47 @@ class DoubaoClient:
self.max_retries: int = settings.doubao_max_retries
self.vision_model: str = settings.doubao_vision_model
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
@property
def is_available(self) -> bool:
"""是否可用(配置了 API Key)."""
+32 -6
View File
@@ -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__":
+262
View File
@@ -0,0 +1,262 @@
"""Additional unit tests to hit uncovered lines for diff-coverage >=60%."""
from __future__ import annotations
import json
from dataclasses import dataclass
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
from packages.shared.ai_client import DoubaoClient
class _FakeSettings:
doubao_api_key = "test-key"
doubao_model = "test-model"
doubao_base_url = "https://ark.cn-beijing.volces.com/api/v3"
doubao_timeout = 10
doubao_max_retries = 0
doubao_vision_model = "test-vision"
doubao_embedding_model = "test-embedding"
def _make_client(api_key: str = "test-key") -> DoubaoClient:
with patch("packages.shared.ai_client.get_shared_settings", return_value=_FakeSettings()):
c = DoubaoClient()
c.api_key = api_key
c.max_retries = 0
return c
class TestDoubaoClientEmbedText:
def test_no_api_key_returns_none(self):
c = _make_client(api_key="")
assert c.embed_text("hello") is None
def test_empty_text_returns_none(self):
c = _make_client()
assert c.embed_text("") is None
assert c.embed_text(" ") is None
def test_none_text_returns_none(self):
c = _make_client()
assert c.embed_text(None) is None
@patch("packages.shared.ai_client.httpx.post")
def test_successful_embedding(self, mock_post):
mock_resp = MagicMock()
mock_resp.json.return_value = {"data": [{"embedding": [0.1, 0.2, 0.3]}]}
mock_resp.raise_for_status = MagicMock()
mock_post.return_value = mock_resp
c = _make_client()
result = c.embed_text("hello world")
assert result == [0.1, 0.2, 0.3]
mock_post.assert_called_once()
@patch("packages.shared.ai_client.httpx.post")
def test_malformed_response_returns_none(self, mock_post):
mock_resp = MagicMock()
mock_resp.json.return_value = {"data": []}
mock_resp.raise_for_status = MagicMock()
mock_post.return_value = mock_resp
c = _make_client()
assert c.embed_text("hello") is None
@patch("packages.shared.ai_client.httpx.post", side_effect=Exception("network error"))
def test_network_error_returns_none(self, mock_post):
c = _make_client()
assert c.embed_text("hello") is None
def test_is_available_with_key(self):
c = _make_client(api_key="sk-xxx")
assert c.is_available is True
def test_is_available_without_key(self):
c = _make_client(api_key="")
assert c.is_available is False
# --- 2. _infer_expected_categories ---
_GEN_TASKS_PATH = Path(__file__).resolve().parents[2] / "apps/api/app/api/routes/generation_tasks.py"
def _load_infer_func():
src = _GEN_TASKS_PATH.read_text()
start = src.index("# #2035:文案关键词")
end = src.index("from packages.middleware")
code = src[start:end]
ns: dict = {}
exec(code, ns)
return ns["_infer_expected_categories"]
_infer_expected_categories = _load_infer_func()
class TestInferExpectedCategories:
def test_none_returns_none(self):
assert _infer_expected_categories(None) is None
assert _infer_expected_categories(set()) is None
def test_product_keyword_matches(self):
cats = _infer_expected_categories({"产品展示"})
assert cats is not None
assert "product" in cats
def test_scenic_keyword_matches(self):
cats = _infer_expected_categories({"户外风景"})
assert cats is not None
assert "scenic" in cats
def test_food_keyword_matches(self):
cats = _infer_expected_categories({"美食制作"})
assert cats is not None
assert "food" in cats
def test_no_match_returns_none(self):
assert _infer_expected_categories({"抽象概念xyz"}) is None
# --- 3. parse_vision_response edge cases ---
from packages.domain.atom_clip_tagger import parse_vision_response
class TestParseVisionResponseEdgeCases:
def test_person_count_type_error_defaults_zero(self):
text = json.dumps({
"scene": [], "objects": [], "action": [], "shot": "", "has_text": False,
"person_count": "not-an-int", "text_content": "", "caption": "x",
})
r = parse_vision_response(text)
assert r["person_count"] == 0
def test_person_count_out_of_range_clamped(self):
text = json.dumps({
"scene": [], "objects": [], "action": [], "shot": "", "has_text": False,
"person_count": 10, "text_content": "", "caption": "x",
})
r = parse_vision_response(text)
assert r["person_count"] == 3
def test_person_count_negative_clamped(self):
text = json.dumps({
"scene": [], "objects": [], "action": [], "shot": "", "has_text": False,
"person_count": -5, "text_content": "", "caption": "x",
})
r = parse_vision_response(text)
assert r["person_count"] == 0
def test_text_content_non_string_defaults_empty(self):
text = '{"scene":[],"objects":[],"action":[],"shot":"","has_text":true,"person_count":0,"text_content":123,"caption":"x"}'
r = parse_vision_response(text)
assert r["text_content"] == ""
def test_caption_truncation_at_80(self):
long_caption = "描" * 100
text = json.dumps({
"scene": [], "objects": [], "action": [], "shot": "", "has_text": False,
"person_count": 0, "text_content": "", "caption": long_caption,
})
r = parse_vision_response(text)
assert len(r["caption"]) == 80
# --- 4. smart_match normalize_tag ---
from packages.domain.smart_match import normalize_tag
class TestNormalizeTagEdge:
def test_none_returns_empty(self):
assert normalize_tag(None) == ""
def test_non_string_converted(self):
assert normalize_tag(123) == "123"
def test_strip_and_lower(self):
assert normalize_tag(" FOO Bar ") == "foo bar"
# --- 5. narrative_match non-dict clip_tags skip ---
from packages.domain.narrative_match import match_assets_by_script_tags
@dataclass
class _FA:
id: str
tags: list
class TestNarrativeMatchNonDictClipTags:
def test_non_dict_clip_tags_are_skipped(self):
a1 = _FA("a1", tags=[])
clip_map = {"a1": [None, "bad", {"scene": ["工厂"], "objects": [], "action": []}, 123]}
matched, unmatched = match_assets_by_script_tags(
[a1], script_tags=["工厂"], clip_ai_tags_by_asset=clip_map
)
assert [a.id for a in matched] == ["a1"]
# --- 6. update_caption_embedding ---
class _FakeSession:
def __init__(self, rows_found: int = 1):
self.rows_found = rows_found
self.commits = 0
self.updates = []
def query(self, model):
return _FQuery(self)
def commit(self):
self.commits += 1
class _FQuery:
def __init__(self, session):
self.session = session
def filter(self, *a, **kw):
return self
def update(self, upd):
self.session.updates.append(upd)
return self.session.rows_found
class TestUpdateCaptionEmbedding:
def _make_repo(self, session):
from packages.adapters.sqlalchemy_impl.asset_atom_clip_repository import SQLAlchemyAssetAtomClipRepository
repo = SQLAlchemyAssetAtomClipRepository.__new__(SQLAlchemyAssetAtomClipRepository)
repo.session = session
return repo
def test_updates_both_caption_and_embedding(self):
s = _FakeSession(rows_found=1)
repo = self._make_repo(s)
ok = repo.update_caption_embedding("c1", "new caption", [0.1, 0.2])
assert ok is True
assert s.commits == 1
assert s.updates[0]["caption"] == "new caption"
assert s.updates[0]["embedding"] == [0.1, 0.2]
def test_only_caption_update(self):
s = _FakeSession(rows_found=1)
repo = self._make_repo(s)
ok = repo.update_caption_embedding("c1", "cap", None)
assert ok is True
assert "embedding" not in s.updates[0]
assert s.updates[0]["caption"] == "cap"
def test_no_update_when_both_none(self):
s = _FakeSession()
repo = self._make_repo(s)
ok = repo.update_caption_embedding("c1", None, None)
assert ok is False
assert s.commits == 0
assert s.updates == []
def test_returns_false_when_row_not_found(self):
s = _FakeSession(rows_found=0)
repo = self._make_repo(s)
ok = repo.update_caption_embedding("c1", "x", [0.1])
assert ok is False
+352
View File
@@ -0,0 +1,352 @@
"""#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_80(self):
long = "A" * 100
text = '{"scene":[],"objects":[],"action":[],"shot":"中景","has_text":false,"caption":"' + long + '"}'
result = parse_vision_response(text)
assert len(result["caption"]) == 80
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;category 中性分 60 * 0.10 = 6
assert breakdown["ai_semantic"] == 10.0
assert breakdown["category_match"] == 6.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=28, duration=22, recency=~12 (new), unused=8, ai=10(neutral), cat=6(neutral)
# 总和应该 ~86
assert 80 <= 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
# ── parse_vision_response: person_count / text_content ───────────
class TestParseVisionResponseEnhanced:
def test_person_count_parsed(self):
text = '{"scene":["办公室"],"objects":["人物","电脑"],"action":["说话"],"shot":"中景","has_text":false,"person_count":1,"text_content":"","caption":"职场女性在办公室讲解产品功能,桌上有笔记本电脑"}'
result = parse_vision_response(text)
assert result["person_count"] == 1
assert result["text_content"] == ""
def test_person_count_multi_people(self):
text = '{"scene":["会议室"],"objects":["人物","桌子","椅子"],"action":["开会"],"shot":"中景","has_text":false,"person_count":3,"text_content":"","caption":"多人在会议室开会讨论项目方案"}'
result = parse_vision_response(text)
assert result["person_count"] == 3 # 3人及以上
def test_person_count_non_int_defaults_zero(self):
text = '{"scene":[],"objects":[],"action":[],"shot":"特写","has_text":false,"person_count":"abc","text_content":"","caption":""}'
result = parse_vision_response(text)
assert result["person_count"] == 0
def test_text_content_extracted_when_has_text(self):
text = '{"scene":["街道"],"objects":["招牌","建筑"],"action":[],"shot":"远景","has_text":true,"person_count":0,"text_content":"欢迎光临","caption":"街道上有一家店铺招牌写着欢迎光临"}'
result = parse_vision_response(text)
assert result["text_content"] == "欢迎光临"
assert result["has_text"] is True
def test_text_content_truncated_at_100(self):
long_text = "X" * 200
text = (
'{"scene":[],"objects":[],"action":[],"shot":"特写","has_text":true,"person_count":0,"text_content":"'
+ long_text
+ '","caption":""}'
)
result = parse_vision_response(text)
assert len(result["text_content"]) == 100
def test_missing_person_count_defaults_zero(self):
text = '{"scene":[],"objects":[],"action":[],"shot":"特写","has_text":false,"caption":"一个苹果"}'
result = parse_vision_response(text)
assert result["person_count"] == 0
assert result["text_content"] == ""
def test_fallback_returns_person_count_zero(self):
"""非 JSON 输入应返回空 dict(不是 fallback tags)。"""
result = parse_vision_response("not json at all")
assert result == {}
def test_objects_list_merged(self):
"""objects 应该被保留并转为列表。"""
text = '{"scene":["厨房"],"objects":["食物","锅","蔬菜","刀具"],"action":["烹饪"],"shot":"中景","has_text":false,"person_count":1,"text_content":"","caption":"厨师在厨房烹饪食物,食材摆放整齐"}'
result = parse_vision_response(text)
assert "食物" in result["objects"]
assert "锅" in result["objects"]
assert "蔬菜" in result["objects"]
assert len(result["objects"]) >= 3
# ── score_asset category_match 维度 ────────────────────────────
class TestScoreAssetCategoryMatch:
def test_no_category_gives_neutral(self):
a = FakeAsset("a1", quality_score=50, metadata={})
_, bd = score_asset(a, now=datetime.now(UTC))
assert bd["category_match"] == 6.0 # 60 * 0.10 = 6
def test_category_hit_gives_10(self):
a = FakeAsset("a1", quality_score=50, metadata={"classification": "product"})
_, bd = score_asset(a, now=datetime.now(UTC), expected_categories={"product", "person"})
assert bd["category_match"] == 10.0 # 100 * 0.10 = 10
def test_category_miss_gives_low(self):
a = FakeAsset("a1", quality_score=50, metadata={"classification": "scenic"})
_, bd_hit = score_asset(a, now=datetime.now(UTC), expected_categories={"product"})
_, bd_neutral = score_asset(a, now=datetime.now(UTC))
assert bd_hit["category_match"] == 3.0 # 30 * 0.10 = 3
assert bd_neutral["category_match"] == 6.0
def test_other_category_neutral(self):
"""other 类不给额外加分。"""
a = FakeAsset("a1", quality_score=50, metadata={"classification": "other"})
_, bd = score_asset(a, now=datetime.now(UTC), expected_categories={"product"})
assert bd["category_match"] == 5.0 # 50 * 0.10 = 5
def test_category_affects_ranking(self):
a_product = FakeAsset("a_product", quality_score=50, duration=15, metadata={"classification": "product"})
a_scenic = FakeAsset("a_scenic", quality_score=50, duration=15, metadata={"classification": "scenic"})
a_none = FakeAsset("a_none", quality_score=50, duration=15, metadata={})
rng = random.Random(42)
results = smart_select_assets(
[a_product, a_scenic, a_none],
kind="video",
rng=rng,
expected_categories={"product"},
)
assert results[0].asset.id == "a_product"
+13 -13
View File
@@ -93,57 +93,57 @@ 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(14.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(22.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"] < 22.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"] < 22.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(6.6, abs=0.1)
def test_unused_asset_gets_full_bonus(self):
asset = FakeAsset(id="a1", quality_score=50, duration=15, metadata={})
_, breakdown = score_asset(asset, now=NOW)
assert breakdown["unused"] == pytest.approx(10.0, abs=0.1)
assert breakdown["unused"] == pytest.approx(8.0, abs=0.1)
def test_used_asset_gets_reduced_bonus(self):
asset = FakeAsset(id="a1", quality_score=50, duration=15, metadata={"generation_use_count": 5})
_, breakdown = score_asset(asset, now=NOW)
assert breakdown["unused"] == pytest.approx(3.0, abs=0.1)
assert breakdown["unused"] == pytest.approx(2.4, abs=0.1)
def test_dirty_metadata_use_count_string_does_not_crash(self):
"""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(8.0, abs=0.1) # use_count=0 → unused_score=100 → 100*0.08=8
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"] > 8.5 # > 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"] < 3.5 # 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 = [
+3 -3
View File
@@ -155,10 +155,10 @@ class TestSmartMatchAvailabilityFallback:
project = Project(id="proj-1", name="Test", owner_user_id="user-1")
assets = [
_video_asset("top-exhausted.mp4", quality=100, used_ranges=_exhausted_ranges(15)),
# second 质量分显著高于 third(质量项差 (90-30)*0.4=24 > 噪声上限 20),
# second 质量分显著高于 third(质量项差 (95-20)*0.28=21 > 噪声上限 20),
# 排除耗尽素材后 second 稳定排首位回补(噪声不影响大分差排名)
_video_asset("second-fresh.mp4", quality=90, used_ranges=None),
_video_asset("third-fresh.mp4", quality=30, used_ranges=None),
_video_asset("second-fresh.mp4", quality=95, used_ranges=None),
_video_asset("third-fresh.mp4", quality=20, used_ranges=None),
]
repo = _StubAssetRepo(assets)
app = _make_app(repo, _StubAssetLibraryRepo({"lib-1": _library()}), _StubProjectRepo({"proj-1": project}))
+6 -6
View File
@@ -97,12 +97,12 @@ class TestScoreAssetUnusedDiminsh:
_, low_bd = score_asset(low)
_, high_bd = score_asset(high)
# use_count=0 → unused_score=100 → component=10.0
assert fresh_bd["unused"] == 10.0
# use_count=2 → unused_score=70 → component=7.0
assert low_bd["unused"] == 7.0
# use_count=10 → unused_score=30 → component=3.0
assert high_bd["unused"] == 3.0
# use_count=0 → unused_score=100 → component=8.0 (weight 0.08)
assert fresh_bd["unused"] == pytest.approx(8.0, abs=0.01)
# use_count=2 → unused_score=70 → component=5.6
assert low_bd["unused"] == pytest.approx(5.6, abs=0.01)
# use_count=10 → unused_score=30 → component=2.4
assert high_bd["unused"] == pytest.approx(2.4, abs=0.01)
def test_monotonically_decreasing_scores(self):
"""使用次数递增时,总评分单调不增。"""