Compare commits

..

1 Commits

Author SHA1 Message Date
xiaoxia b203d84956 fix: WebCodecs 解码失败时自动 fallback 到原生 video 播放
- useCanvasPlayer 新增 hasDecodeError/errorMessage 状态和 onError 回调
- decodeSegment catch 块不再静默失败,更新状态并通知上层
- VideoDecoder error 回调同步报告错误状态
- 连续 3 次解码失败时报告 hasDecodeError
- FrontendPreviewPlayer 检测 hasDecodeError 后自动切换 video fallback
- 解码失败时显示具体错误信息而非静默'暂无可播放素材'
- 新增 normalizeCodecString 函数规范化 codec 字符串
- configure 调用增加 codec 字符码调试日志
2026-08-21 14:43:01 +08:00
13 changed files with 601 additions and 1083 deletions
+39 -169
View File
@@ -31,6 +31,8 @@ logger = logging.getLogger(__name__)
router = APIRouter(tags=["Generation"])
# ── Schemas ──────────────────────────────────────────────────────────────
@@ -63,44 +65,6 @@ class GenerateCoverResponse(BaseModel):
# ── Route ────────────────────────────────────────────────────────────────
def _persist_cover_frame(frame_url: str, plan_id: str) -> str:
"""下载 MediaKit 返回的临时帧图并转存到 OSS covers/ 路径。"""
import tempfile
import uuid
from pathlib import Path
tmp_path: str | None = None
try:
import httpx
resp = httpx.get(frame_url, timeout=30, follow_redirects=True)
resp.raise_for_status()
if not resp.content:
return frame_url
with tempfile.NamedTemporaryFile(suffix=".jpg", delete=False) as tmp:
tmp.write(resp.content)
tmp_path = tmp.name
from packages.shared.storage import get_shared_storage_service
storage = get_shared_storage_service()
cover_key = f"covers/{plan_id}/cover_{uuid.uuid4().hex[:8]}.jpg"
storage.upload_file(
file_or_path=tmp_path,
storage_key=cover_key,
content_type="image/jpeg",
)
public_url = storage.get_url(cover_key)
return public_url or frame_url
except Exception:
logger.warning("封面帧转存失败,返回原始 URL: plan_id=%s", plan_id, exc_info=True)
return frame_url
finally:
if tmp_path:
Path(tmp_path).unlink(missing_ok=True)
@router.post("/generate-cover", response_model=GenerateCoverResponse)
def generate_cover(
body: GenerateCoverRequest,
@@ -234,31 +198,44 @@ def generate_cover(
exc_info=True,
)
# 使用裸 URLrendered/* 已配置公开读);找不到渲染视频时不立即报错,
# 因为步骤 E 可以直接从源素材抽帧(历史数据或 Worker 抽帧失败时的兜底)
# 仍然找不到才报 400
if not rendered_storage_key:
logger.error("[封面生成] ❌ 找不到预览视频: plan_id=%s", plan_id)
raise HTTPException(
status_code=400,
detail="请先生成预览视频,再生成封面",
)
# 回写到 plan.config
plan_svc.update_plan_config(plan_id, {"rendered_storage_key": rendered_storage_key})
# 使用裸 URLrendered/* 已配置公开读)
primary_video_url = None
if rendered_storage_key:
plan_svc.update_plan_config(plan_id, {"rendered_storage_key": rendered_storage_key})
try:
if rendered_storage_key.startswith("http"):
primary_video_url = rendered_storage_key
else:
from packages.shared.storage import get_shared_storage_service
try:
if rendered_storage_key.startswith("http"):
primary_video_url = rendered_storage_key
else:
from packages.shared.storage import get_shared_storage_service
storage_svc = get_shared_storage_service()
primary_video_url = storage_svc.get_url(rendered_storage_key)
if primary_video_url:
import re as _re
storage_svc = get_shared_storage_service()
primary_video_url = storage_svc.get_url(rendered_storage_key)
# 防御性规范化:合并路径中的双斜杠(// -> /),但保留协议头的 ://
# 历史数据中 project_id 为空时会产生 projects//tasks/ 路径,
# MediaKit 的 HTTP 客户端会规范化 URL 导致 404
if primary_video_url:
import re as _re
primary_video_url = _re.sub(r"(?<!:)//", "/", primary_video_url)
logger.info(
"获取预览视频URL用于封面生成: plan_id=%s url=%s",
plan_id,
primary_video_url[:80] if primary_video_url else "",
)
except Exception as e:
logger.warning("获取预览视频URL失败: plan_id=%s err=%s", plan_id, e)
primary_video_url = None
primary_video_url = _re.sub(r"(?<!:)//", "/", primary_video_url)
logger.info(
"获取预览视频URL用于封面生成: plan_id=%s url=%s",
plan_id,
primary_video_url[:80] if primary_video_url else "",
)
except Exception as e:
raise HTTPException(
status_code=500,
detail=f"获取预览视频URL失败: {e}",
) from e
# 统一封面管道:优先从 GenerationTask.cover_url 读取渲染后视频抽帧的封面
# 多步查找 cover_url,和查找视频 URL 一样的 fallback 逻辑
@@ -333,113 +310,6 @@ def generate_cover(
exc_info=True,
)
# 步骤 D:从 plan.config.cover_candidates 读取(Worker 渲染时写入)
if not cover_url_from_task:
_candidates = (plan.config or {}).get("cover_candidates") or []
if isinstance(_candidates, list) and _candidates:
_first = _candidates[0]
if isinstance(_first, dict):
cover_url_from_task = _first.get("image_url") or _first.get("url") or ""
if cover_url_from_task:
logger.info(
"[封面生成] 统一管道封面(步骤D-cover_candidates): plan_id=%s url=%s",
plan_id,
cover_url_from_task[:80],
)
# 步骤 E1:如果有已渲染的预览视频 URL 但 cover_url 未持久化(历史数据),
# 直接从渲染视频抽帧
if not cover_url_from_task and primary_video_url:
try:
from packages.shared.mediakit_client import get_mediakit_client
mk_client = get_mediakit_client()
if mk_client.is_available:
logger.info(
"[封面生成] 步骤E1-从渲染视频抽帧: plan_id=%s url=%s",
plan_id,
primary_video_url[:80],
)
snapshots = mk_client.extract_frames(
video_url=primary_video_url,
strategy="SpecifiedFrames",
max_frames=1,
poll_interval=2.0,
max_poll_attempts=5,
max_retries=0,
)
if snapshots:
raw = snapshots[0].get("image_url") or snapshots[0].get("url") or ""
if raw:
cover_url_from_task = _persist_cover_frame(raw, plan_id)
logger.info(
"[封面生成] 统一管道封面(步骤E1-rendered-video): plan_id=%s url=%s",
plan_id,
cover_url_from_task[:80],
)
except Exception:
logger.warning(
"[封面生成] 步骤E1从渲染视频抽帧失败: plan_id=%s",
plan_id,
exc_info=True,
)
# 步骤 E2:当 A/B/C/D/E1 均未命中(如历史预览任务无 cover_url)时,
# 直接从用户选择的第一个视频素材中抽取封面帧作为兜底。API 请求内短超时,不阻塞。
if not cover_url_from_task and body.asset_ids:
from packages.adapters.sqlalchemy_impl.asset_repository import (
SQLAlchemyAssetRepository,
)
from packages.shared.mediakit_client import get_mediakit_client
from packages.shared.storage import get_shared_storage_service
asset_repo = SQLAlchemyAssetRepository(db)
storage_svc = get_shared_storage_service()
mk_client = get_mediakit_client()
if mk_client.is_available:
for aid in body.asset_ids:
try:
asset = asset_repo.get(aid)
if not asset or asset.file_type != "video":
continue
sk = asset.storage_key or ""
if not sk:
continue
src_url = sk if sk.startswith("http") else storage_svc.get_url(sk)
if not src_url:
continue
logger.info(
"[封面生成] 步骤E-从素材抽帧: plan_id=%s asset_id=%s url=%s",
plan_id,
aid,
src_url[:80],
)
snapshots = mk_client.extract_frames(
video_url=src_url,
strategy="SpecifiedFrames",
max_frames=1,
poll_interval=2.0,
max_poll_attempts=5,
max_retries=0,
)
if snapshots:
raw = snapshots[0].get("image_url") or snapshots[0].get("url") or ""
if raw:
cover_url_from_task = _persist_cover_frame(raw, plan_id)
logger.info(
"[封面生成] 统一管道封面(步骤E-source-asset): plan_id=%s url=%s",
plan_id,
cover_url_from_task[:80],
)
break
except Exception:
logger.warning(
"[封面生成] 步骤E从素材抽帧失败: plan_id=%s asset_id=%s",
plan_id,
aid,
exc_info=True,
)
if cover_url_from_task:
# 标题已在预览视频渲染时烧录(ASS字幕),封面帧自然包含标题
cover_data = {
@@ -455,13 +325,13 @@ def generate_cover(
return GenerateCoverResponse(plan_id=plan_id, cover=cover_data)
logger.warning(
"[封面生成] 统一管道未找到 cover_url (A/B/C/D均未命中): plan_id=%s",
"[封面生成] 统一管道未找到 cover_url: plan_id=%s",
plan_id,
)
# ai_frame/ai_regenerate 类型必须从渲染管道获取,不再回退到 AI 服务
raise HTTPException(
status_code=400,
detail="封面生成失败:未找到可抽帧的视频素材,请确认已上传视频素材后重试",
detail="封面尚未生成,请先重新生成预览视频以触发封面自动提取",
)
from packages.shared.ai_service import run_generate_cover
+4 -1
View File
@@ -231,7 +231,10 @@ test.describe("Core generation flow", () => {
(response) => {
const url = response.url()
const path = new URL(url).pathname
return response.request().method() === "POST" && path.endsWith("/generation/tasks")
return (
response.request().method() === "POST" &&
path.endsWith("/generation/tasks")
)
},
{ timeout: 30_000 },
)
@@ -17,7 +17,7 @@ import {
import type { AssetItem } from "@/api/assets"
import type { EditingTemplate } from "@/api/editing-planner"
import { useSegmentScheduler, type PlaybackSegment } from "../hooks/useSegmentScheduler"
import { useCanvasPlayer } from "../hooks/useCanvasPlayer"
import { useCanvasPlayer, isWebCodecsSupported } from "../hooks/useCanvasPlayer"
interface FrontendPreviewPlayerProps {
assets: AssetItem[]
@@ -82,9 +82,7 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
titleSettings,
}) => {
const segments = useMemo(() => buildPlaybackSegments(assets, template), [assets, template])
// 默认走原生 video 播放(浏览器硬件解码,独立线程,不阻塞 UI)
// WebCodecs 仅在明确需要时启用(保留代码作为兜底)
const useWebCodecs = false
const useWebCodecs = isWebCodecsSupported()
// ── 两条路径共用同一个 canvas ref(fallback 路径不使用) ──
const canvasRef = useRef<HTMLCanvasElement>(null)
@@ -124,10 +122,9 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
const { state: canvasState, controls: canvasControls } = useCanvasPlayer(
canvasRef,
useWebCodecs && !forceVideoFallback ? canvasSegments : [],
canvasSegments,
useWebCodecs && !forceVideoFallback ? canvasTitle : undefined,
handleCanvasError,
useWebCodecs && !forceVideoFallback,
)
// WebCodecs 报告解码失败时自动切换到 video fallback
@@ -198,9 +195,7 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
const audio = audioRef.current
if (!audio || !audio.src || !isPlaying) return
audio.currentTime = currentTime
// 注意:不要把 currentTime 放进依赖数组,否则每200ms会重置音频位置导致卡顿
// eslint-disable-next-line react-hooks/exhaustive-deps
}, [segmentSyncKey, isPlaying])
}, [segmentSyncKey, isPlaying, currentTime])
const handleSeekTo = useCallback(
(time: number) => {
@@ -270,8 +265,32 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
const progressPercent = totalDuration > 0 ? (currentTime / totalDuration) * 100 : 0
// ── Canvas 容器 ref(保留声明,WebCodecs 兜底路径仍引用) ──
// ── Canvas ResizeObserver ──
const canvasContainerRef = useRef<HTMLDivElement>(null)
useEffect(() => {
if (!effectiveUseWebCodecs || !canPlay) return
const container = canvasContainerRef.current
const canvas = canvasRef.current
if (!container || !canvas) return
// 立即设置一次 canvas 像素分辨率,避免默认 300×150 导致首帧变形
const initRect = container.getBoundingClientRect()
if (initRect.width > 0 && initRect.height > 0) {
const dpr = window.devicePixelRatio || 1
canvas.width = initRect.width * dpr
canvas.height = initRect.height * dpr
}
const ro = new ResizeObserver((entries) => {
for (const entry of entries) {
const { width, height } = entry.contentRect
if (width > 0 && height > 0) {
canvas.width = width * window.devicePixelRatio
canvas.height = height * window.devicePixelRatio
}
}
})
ro.observe(container)
return () => ro.disconnect()
}, [effectiveUseWebCodecs, canPlay])
// ── 未就绪 ──
if (!ready || !assets.length) {
@@ -366,7 +385,7 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
</div>
)}
{/* ── Video 渲染层(默认路径,浏览器原生硬件解码 ── */}
{/* ── Video 渲染层(fallback 路径,或 WebCodecs 解码失败时自动切换 ── */}
{!effectiveUseWebCodecs &&
segments.map((seg, i) => (
<video
@@ -375,7 +394,9 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
ref={(el) => {
videoRefs.current[i] = el
}}
preload="auto"
preload={
i === videoCurrentSegIdx ? "auto" : i === videoCurrentSegIdx + 1 ? "metadata" : "none"
}
src={seg.videoUrl}
style={{
position: "absolute",
@@ -434,7 +455,11 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
zIndex: 10,
}}
>
{`片段 ${videoCurrentSegIdx + 1}/${segments.length}`}
{effectiveUseWebCodecs
? "Canvas"
: forceVideoFallback
? "Canvas 解码失败,已切换原生播放"
: `片段 ${videoCurrentSegIdx + 1}/${segments.length}`}
</div>
{/* 控制条 */}
@@ -16,7 +16,6 @@ import { LoadingOutlined } from "@ant-design/icons"
import type { AssetItem } from "@/api/assets"
import type { EditingTemplate } from "@/api/editing-planner"
import type { TitleSettings } from "../types"
import { getFontFamily } from "../constants"
import FrontendPreviewPlayer from "./FrontendPreviewPlayer"
interface PreviewVideoPanelProps {
@@ -88,7 +87,7 @@ function buildTitleStyle(settings: TitleSettings, containerHeight: number): Reac
: (Math.min(settings.size, 96) / ASS_VIDEO_HEIGHT) * 400 // fallback
const base: React.CSSProperties = {
fontFamily: getFontFamily(settings.font),
fontFamily: settings.font || "思源黑体",
fontSize: `${fontSizePx}px`,
color: settings.color || "#ffffff",
fontWeight: settings.bold ? 700 : 400,
@@ -177,12 +176,7 @@ const TitleOverlay: React.FC<{ titleSettings: TitleSettings }> = ({ titleSetting
position: "absolute",
}}
>
{displayTitle.split("/").map((part, i) => (
<span key={i}>
{i > 0 && <br />}
{part}
</span>
))}
{displayTitle}
</div>
</div>
)
@@ -17,16 +17,6 @@ interface Step5VoiceSelectProps {
}
/** 格式化时长 mm:ss */
/** 获取素材实际时长(优先顶层 durationfallback 到 metadata.duration */
const getDuration = (item: AssetItem): number => {
return item.duration ?? (item.metadata?.duration as number) ?? 0
}
/** 获取素材实际文件大小 */
const getFileSize = (item: AssetItem): number => {
return item.file_size ?? (item.metadata?.file_size as number) ?? 0
}
const formatDuration = (seconds?: number): string => {
if (!seconds || seconds <= 0) return "00:00"
const m = Math.floor(seconds / 60)
@@ -99,7 +89,7 @@ const Step5VoiceSelect: React.FC<Step5VoiceSelectProps> = ({
// 如果启用了时长校验,且配音时长不足
if (totalVideoDuration > 0) {
const material = materials.find((m) => m.id === id)
if (material && getDuration(material) < totalVideoDuration) {
if (material && (material.duration || 0) < totalVideoDuration) {
setPendingVoiceId(id)
setDurationWarningOpen(true)
return
@@ -281,24 +271,25 @@ const Step5VoiceSelect: React.FC<Step5VoiceSelectProps> = ({
}}
>
<span style={{ display: "flex", alignItems: "center", gap: 4 }}>
{formatDuration(getDuration(item))}
{totalVideoDuration > 0 && getDuration(item) < Number(totalVideoDuration) && (
<span
style={{
color: "#ff4d4f",
fontSize: 11,
fontWeight: 500,
display: "inline-flex",
alignItems: "center",
gap: 2,
}}
>
<WarningOutlined />
</span>
)}
{formatDuration(item.duration)}
{totalVideoDuration > 0 &&
(Number(item.duration) || 0) < Number(totalVideoDuration) && (
<span
style={{
color: "#ff4d4f",
fontSize: 11,
fontWeight: 500,
display: "inline-flex",
alignItems: "center",
gap: 2,
}}
>
<WarningOutlined />
</span>
)}
</span>
<span>{formatFileSize(getFileSize(item))}</span>
<span>{formatFileSize(item.file_size)}</span>
</div>
</div>
)
@@ -327,9 +318,7 @@ const Step5VoiceSelect: React.FC<Step5VoiceSelectProps> = ({
return (
<p>
<strong>
{pendingMaterial ? formatDuration(getDuration(pendingMaterial)) : "--"}
</strong>
<strong>{pendingMaterial ? formatDuration(pendingMaterial.duration) : "--"}</strong>
<strong>{formatDuration(totalVideoDuration)}</strong>
@@ -2,7 +2,6 @@
* 标题预设样式网格
*/
import React from "react"
import { getFontFamily } from "../../constants"
interface TitlePresetItem {
key: string
@@ -36,7 +35,7 @@ const TitlePresetsGrid: React.FC<TitlePresetsGridProps> = ({
>
<span
className="xx-title-preset-preview-text"
style={{ ...p.previewStyle, fontFamily: getFontFamily(fontFamily || "思源黑体") }}
style={{ ...p.previewStyle, ...(fontFamily ? { fontFamily } : {}) }}
>
</span>
-15
View File
@@ -56,21 +56,6 @@ export const FONT_OPTIONS = [
"华康俪金黑",
]
/* ── 标题字体 CSS font-family 映射(中文显示名 → 浏览器可识别的字体栈) ── */
export const FONT_FAMILY_MAP: Record<string, string> = {
: '"Source Han Sans SC", "Noto Sans SC", "PingFang SC", "Microsoft YaHei", sans-serif',
: '"Source Han Serif SC", "Noto Serif SC", "Songti SC", "SimSun", serif',
: '"PingFang SC", -apple-system, "Helvetica Neue", sans-serif',
PingFang: '"PingFang SC", -apple-system, "Helvetica Neue", sans-serif',
: '"Microsoft YaHei", "PingFang SC", sans-serif',
: '"KaiTi", "STKaiti", "DFKai-SB", serif',
: '"华康俪金黑", "DFLiJinHei-W8", "Source Han Sans SC", "Microsoft YaHei", sans-serif',
}
export function getFontFamily(font: string): string {
return FONT_FAMILY_MAP[font] || FONT_FAMILY_MAP["思源黑体"]
}
/* ── 标题样式预设 ── */
export const TITLE_PRESETS = [
{
-2
View File
@@ -2309,8 +2309,6 @@
.xx-cover-preview-box {
position: relative;
aspect-ratio: 9 / 16;
max-width: 180px;
margin: 0 auto;
background: var(--bg-tertiary);
border-radius: var(--radius-md);
overflow: hidden;
@@ -9,6 +9,8 @@ import { createFile } from "mp4box"
import type { Movie, Sample } from "mp4box"
// ── 常量 ──
/** 初始化预解码最大帧数(约 2 秒 @30fps),后续帧通过 decodeAroundPosition 按需解码 */
const MAX_INIT_FRAMES = 60
/**
* 规范化 mp4box 提取的 codec 字符串为 WebCodecs 兼容格式
@@ -109,7 +111,7 @@ class FrameQueue {
private frames: FrameEntry[] = []
private maxSize: number
constructor(maxSize = 200) {
constructor(maxSize = 5) {
this.maxSize = maxSize
}
@@ -121,36 +123,24 @@ class FrameQueue {
this.frames.push(entry)
}
/** 获取当前时间戳应显示的帧(二分查找,O(log n) */
/** 获取当前时间戳应显示的帧 */
getCurrentFrame(timestamp: number): VideoFrame | null {
if (this.frames.length === 0) return null
const target = timestamp + 0.01
// 找到最后一个 pts <= target 的帧(右边界)
let lo = 0,
hi = this.frames.length - 1,
bestIdx = -1
while (lo <= hi) {
const mid = (lo + hi) >> 1
if (this.frames[mid].pts <= target) {
bestIdx = mid
lo = mid + 1
} else {
hi = mid - 1
let best: FrameEntry | null = null
let bestIdx = -1
for (let i = 0; i < this.frames.length; i++) {
const f = this.frames[i]
if (f.pts <= timestamp + 0.01) {
best = f
bestIdx = i
}
}
if (bestIdx < 0) return null
// 关闭并移除 bestIdx 之前的所有已播放帧
for (let i = 0; i < bestIdx; i++) {
this.frames[i].frame.close()
}
this.frames.splice(0, bestIdx)
// 此时 bestIdx 对应帧已在索引 0
return this.frames[0]?.frame ?? null
if (bestIdx >= 0) {
this.frames = this.frames.slice(bestIdx)
}
return best?.frame ?? null
}
clear() {
@@ -239,10 +229,9 @@ export function useCanvasPlayer(
shadow?: boolean
},
onError?: (error: Error) => void,
enabled: boolean = true,
) {
const [state, setState] = useState<CanvasPlayerState>({
hasSupport: enabled && isWebCodecsSupported(),
hasSupport: isWebCodecsSupported(),
isPlaying: false,
currentTime: 0,
duration: 0,
@@ -254,13 +243,9 @@ export function useCanvasPlayer(
// ── 内部引用 ──
const decoderRef = useRef<VideoDecoder | null>(null)
const frameQueueRef = useRef(new FrameQueue(200))
/** 每个片段持久化解码器,避免每次新建导致关键帧错误 */
const segmentDecodersRef = useRef(new Map<number, VideoDecoder>())
/** 每个片段已送入解码器的 sample 游标(用于续解码) */
const segmentSampleCursorRef = useRef(new Map<number, number>())
/** 后台补充解码是否正在运行(防重入) */
const isFeedingRef = useRef(false)
const frameQueueRef = useRef(new FrameQueue(600))
/** 已解码的片段索引集合,用于按需解码(先标记防重入,失败时移除允许重试) */
const decodedSegmentsRef = useRef(new Set<number>())
/** 解码代数计数器,seek 时递增以作废正在进行的异步解码 */
const decodeGenerationRef = useRef(0)
const rafRef = useRef<number>(0)
@@ -490,91 +475,115 @@ export function useCanvasPlayer(
[segments, extractCodecDescription],
)
// ── 解码片段的一批帧(使用持久化解码器,支持从断点续解码) ──
/**
* @param segIdx 片段索引
* @param maxFrames 本次最多解码多少帧
*/
const decodeSegmentBatch = useCallback(
async (segIdx: number, maxFrames: number = 60): Promise<number> => {
if (isDestroyedRef.current) return 0
const metas = segmentMetaRef.current
const meta = metas[segIdx]
if (!meta) return 0
const buffer = segmentDataRef.current.get(meta.assetId)
if (!buffer) return 0
// ── 初始化 VideoDecoder 并解码指定片段 ──
const decodeSegment = useCallback(
async (_buffer: ArrayBuffer, meta: SegmentMeta, maxFrames?: number): Promise<void> => {
if (isDestroyedRef.current) return
const gen = decodeGenerationRef.current
let decoder = segmentDecodersRef.current.get(segIdx)
let cursor = segmentSampleCursorRef.current.get(segIdx) ?? 0
const samples = meta.samples
let decoderReady = false
// 如果还没有解码器,新建一个(从关键帧开始,不会报 key frame 错误
if (!decoder || decoder.state === "closed") {
decoder = new VideoDecoder({
output: (frame: VideoFrame) => {
if (videoDimRef.current.width === 0 || videoDimRef.current.height === 0) {
videoDimRef.current = { width: frame.codedWidth, height: frame.codedHeight }
}
const localTime = frame.timestamp / 1_000_000
const globalTime = localTime + meta.globalStartTime
frameQueueRef.current.push({
frame,
pts: globalTime,
duration: (frame.duration ?? 0) / 1_000_000,
})
},
error: (e: DOMException) => {
console.error(`[useCanvasPlayer] Segment ${segIdx} decoder error:`, e)
// 重置该片段的解码器和游标,允许重试
segmentDecodersRef.current.delete(segIdx)
segmentSampleCursorRef.current.set(segIdx, 0)
},
})
try {
await decoder.configure({
codec: meta.codec,
...(meta.description ? { description: meta.description } : {}),
// 配置解码器(每个片段可能需要不同的 codec/分辨率
const decoder = new VideoDecoder({
output: (frame: VideoFrame) => {
// 从第一帧获取实际尺寸
if (videoDimRef.current.width === 0 || videoDimRef.current.height === 0) {
videoDimRef.current = { width: frame.codedWidth, height: frame.codedHeight }
console.log(
`[useCanvasPlayer] Actual frame size: ${frame.codedWidth}x${frame.codedHeight}`,
)
}
const localTime = frame.timestamp / 1_000_000
const globalTime = localTime + meta.globalStartTime
frameQueueRef.current.push({
frame,
pts: globalTime,
duration: (frame.duration ?? 0) / 1_000_000,
})
} catch (err) {
console.error(`[useCanvasPlayer] Segment ${segIdx} configure failed:`, err)
const error = err instanceof Error ? err : new Error(String(err))
},
error: (e: DOMException) => {
console.error("[useCanvasPlayer] Decoder error callback:", e)
const error = new Error(`VideoDecoder error: ${e.message || e.name || "unknown"}`)
setState((s) => ({
...s,
isBuffering: false,
hasDecodeError: true,
errorMessage: `视频解码失败: ${error.message || "不支持的编解码器"}`,
errorMessage: `视频解码器错误: ${e.message || "解码异常"}`,
}))
onErrorRef.current?.(error)
return 0
}
},
})
segmentDecodersRef.current.set(segIdx, decoder)
cursor = 0
// 标记缓冲结束
console.log("[useCanvasPlayer] configure:", {
codec: meta.codec,
description: meta.description,
descriptionByteLength: meta.description?.byteLength,
videoWidth: meta.videoWidth,
videoHeight: meta.videoHeight,
codecCharCodes: meta.codec.split("").map((c) => c.charCodeAt(0)),
})
try {
await decoder.configure({
codec: meta.codec,
...(meta.description ? { description: meta.description } : {}),
})
decoderRef.current = decoder
decoderReady = true
// 标记缓冲结束,让 UI 开始渲染
setState((s) => ({ ...s, isBuffering: false }))
} catch (err) {
const error = err instanceof Error ? err : new Error(String(err))
console.error(
"[useCanvasPlayer] Decoder configure failed for segment:",
meta.assetId,
error,
)
console.error("[useCanvasPlayer] Failed codec config:", {
codec: meta.codec,
descriptionByteLength: meta.description?.byteLength,
videoWidth: meta.videoWidth,
videoHeight: meta.videoHeight,
})
setState((s) => ({
...s,
isBuffering: false,
isReady: false,
hasDecodeError: true,
errorMessage: `视频解码失败: ${error.message || "不支持的编解码器"}`,
}))
onErrorRef.current?.(error)
return
}
if (decoder.state !== "configured") return 0
if (!decoderReady) return
// 从 cursor 继续喂 sample(流水线批量提交,不 await 单个 decode
let decoded = 0
let si = cursor
while (si < samples.length && decoded < maxFrames) {
if (decodeGenerationRef.current !== gen || isDestroyedRef.current) break
if ((decoder.state as string) === "closed") break
// 帧队列快满时停止提交(这才是真正的背压)
if (frameQueueRef.current.size >= 180) break
// 解码器内部队列积压过多时短暂让出线程(阈值64,给硬件足够流水线深度)
if (decoder.decodeQueueSize > 64) {
await new Promise((r) => setTimeout(r, 5))
// 使用 demuxSegment 中已提取并过滤的 samples(前端切片
const samplesCollected = meta.samples
console.log(
`[useCanvasPlayer] Segment ${meta.assetId}: ${samplesCollected.length} samples to decode`,
)
if (samplesCollected.length === 0) {
console.warn("[useCanvasPlayer] No samples to decode for segment", meta.assetId)
return
}
// 送入解码器
let decodedCount = 0
let skippedCount = 0
let decodeErrors = 0
for (const sample of samplesCollected) {
if (!sample.data || isDestroyedRef.current) {
skippedCount++
continue
}
const sample = samples[si]
si++
if (!sample.data) continue
if (decoder.state === "closed") break
// 初始化阶段限制解码帧数,避免帧缓冲溢出
if (maxFrames && decodedCount >= maxFrames) {
console.log(
`[useCanvasPlayer] Segment ${meta.assetId}: init decode limited to ${maxFrames} frames`,
)
break
}
const chunk = new EncodedVideoChunk({
type: sample.is_sync ? "key" : "delta",
@@ -584,73 +593,93 @@ export function useCanvasPlayer(
})
try {
decoder.decode(chunk)
decoded++
await decoder.decode(chunk) // 修复:await 捕获异步错误
decodedCount++
} catch (e) {
console.warn(`[useCanvasPlayer] Segment ${segIdx} decode error:`, e)
break
}
}
cursor = si
segmentSampleCursorRef.current.set(segIdx, cursor)
// 等待解码器输出帧(最多500ms)
if (decoded > 0 && (decoder.state as string) === "configured") {
let waited = 0
while (frameQueueRef.current.size < Math.min(decoded, 10) && waited < 500) {
await new Promise((r) => setTimeout(r, 20))
waited += 20
if (isDestroyedRef.current) break
decodeErrors++
console.warn(`[useCanvasPlayer] Decode chunk error (${decodeErrors}):`, e)
// 连续 3 次解码失败,放弃当前片段并报告错误
if (decodeErrors >= 3) {
console.error("[useCanvasPlayer] Too many decode errors, aborting segment")
const error = new Error(`视频解码连续失败 ${decodeErrors} 次,片段: ${meta.assetId}`)
setState((s) => ({
...s,
isBuffering: false,
hasDecodeError: true,
errorMessage: `视频解码失败: 连续 ${decodeErrors} 次错误`,
}))
onErrorRef.current?.(error)
break
}
}
}
console.log(
`[useCanvasPlayer] Segment ${segIdx} decoded ${decoded} frames, queue size: ${frameQueueRef.current.size}`,
`[useCanvasPlayer] Segment ${meta.assetId}: decoded ${decodedCount}, skipped ${skippedCount}, errors ${decodeErrors}, decoder.state=${decoder.state}`,
)
return decoded
// flush 仅在解码器状态正常时执行
if (decoder.state === "configured") {
try {
await decoder.flush()
console.log(`[useCanvasPlayer] Segment ${meta.assetId}: flush complete`)
} catch (e) {
console.warn("[useCanvasPlayer] Decoder flush error:", e)
}
}
},
[],
)
// ── 后台持续补充帧 ──
/**
* 根据当前播放时间,确保队列中有足够缓冲
* 播放循环每 200ms 调用一次
* 按需解码当前播放位置 ±1 个片段。
* 在渲染循环中定期调用,避免一次性解码所有片段导致环形缓冲区溢出丢帧。
* 使用"先标记再解码"模式防止并发重复解码,失败时移除标记允许重试。
*/
const feedFrames = useCallback(
const decodeAroundPosition = useCallback(
async (currentTime: number) => {
if (isFeedingRef.current) return
isFeedingRef.current = true
try {
const metas = segmentMetaRef.current
if (!metas || metas.length === 0) return
const metas = segmentMetaRef.current
if (!metas || metas.length === 0) return
// 队列帧数充足时不解码(目标:保持 >= 80 帧缓冲)
if (frameQueueRef.current.size >= 80) return
// 记录当前代数,seek 后代数变化则中止
const gen = decodeGenerationRef.current
// 找到当前播放的片段
let targetIdx = 0
let acc = 0
for (let i = 0; i < metas.length; i++) {
const dur = metas[i].globalEndTime - metas[i].globalStartTime
if (currentTime < acc + dur) {
targetIdx = i
break
}
acc += dur
let targetIdx = -1
let acc = 0
for (let i = 0; i < metas.length; i++) {
const dur = metas[i].globalEndTime - metas[i].globalStartTime
if (currentTime < acc + dur) {
targetIdx = i
break
}
acc += dur
}
if (targetIdx === -1) targetIdx = metas.length - 1
// 依次补充:当前片段 → 下一个片段 → 再下一个
for (let offset = 0; offset <= 2; offset++) {
const idx = targetIdx + offset
if (idx >= metas.length) break
if (frameQueueRef.current.size >= 180) break
await decodeSegmentBatch(idx, 60)
for (
let i = Math.max(0, targetIdx - 1);
i <= Math.min(metas.length - 1, targetIdx + 1);
i++
) {
// seek 已作废当前解码任务
if (decodeGenerationRef.current !== gen) return
if (decodedSegmentsRef.current.has(i)) continue
const meta = metas[i]
const buffer = segmentDataRef.current.get(meta.assetId)
if (!buffer) continue
// 先标记为解码中,防止下一帧渲染时重复发起解码
decodedSegmentsRef.current.add(i)
try {
await decodeSegment(buffer, meta, 300)
} catch (e) {
// 解码失败则移除标记,允许后续重试
decodedSegmentsRef.current.delete(i)
console.warn(`[useCanvasPlayer] 按需解码片段 ${i} 失败:`, e)
}
} finally {
isFeedingRef.current = false
// await 后再次检查代数,seek 期间不更新标记
if (decodeGenerationRef.current !== gen) return
}
},
[decodeSegmentBatch],
[decodeSegment],
)
// ── 标题绘制 ──
@@ -667,6 +696,7 @@ export function useCanvasPlayer(
// 按 "/" 分割为多行("/" 作为手动换行符)
const lines = title.text.split(/[//⁄∕]/)
console.log("[drawTitle] 原始标题:", JSON.stringify(title.text), "分割后:", lines)
const lineHeight = fontSize * 1.3
const totalHeight = lines.length * lineHeight
@@ -778,10 +808,8 @@ export function useCanvasPlayer(
}
return s
})
// 后台补充帧:队列不足时自动续解码
if (frameQueueRef.current.size < 80) {
void feedFrames(currentTime)
}
// 按需解码当前 ±1 片段
decodeAroundPosition(currentTime)
}
if (currentTime >= totalDuration) {
@@ -790,46 +818,29 @@ export function useCanvasPlayer(
}
rafRef.current = requestAnimationFrame(renderFrame)
}, [canvasRef, totalDuration, titleSettings, drawTitle, computeDrawRect, feedFrames])
}, [canvasRef, totalDuration, titleSettings, drawTitle, computeDrawRect, decodeAroundPosition])
// ── 播放控制 ──
const play = useCallback(async () => {
if (!state.hasSupport || isDestroyedRef.current) return
// 重播:必须关闭旧解码器、清空队列、重置游标,从头重新解码
if (state.currentTime >= totalDuration - 0.1 || state.currentTime <= 0.1) {
// 重播场景:currentTime 已回到起点但 decodedSegmentsRef 仍有旧标记
// 此时 FrameQueue 中旧帧已被淘汰,需清空标记让 decodeAroundPosition 重新解码
if (state.currentTime <= 0.1 && decodedSegmentsRef.current.size > 0) {
decodeGenerationRef.current++
// 关闭所有持久化解码器
for (const d of segmentDecodersRef.current.values()) {
try {
if (d.state !== "closed") d.close()
} catch {
/* noop */
}
}
segmentDecodersRef.current.clear()
segmentSampleCursorRef.current.clear()
decodedSegmentsRef.current.clear()
// 同步清空帧缓冲,避免旧帧残留导致 getCurrentFrame 返回 null
frameQueueRef.current.clear()
playStartOffsetRef.current = 0
setState((s) => ({ ...s, currentTime: 0 }))
// 重新初始化解码
const metas = segmentMetaRef.current
const initialDecodeCount = Math.min(metas.length, 2)
for (let i = 0; i < initialDecodeCount; i++) {
await decodeSegmentBatch(i, 60)
}
}
setState((s) => ({ ...s, isPlaying: true }))
playStartRef.current = performance.now()
if (state.currentTime < 0.1) {
playStartOffsetRef.current = 0
} else {
playStartOffsetRef.current = state.currentTime
}
playStartOffsetRef.current = state.currentTime
lastProgressUpdateRef.current = 0
rafRef.current = requestAnimationFrame(renderFrame)
}, [state.hasSupport, state.currentTime, totalDuration, renderFrame, decodeSegmentBatch])
// 立即触发一次按需解码,不等渲染循环 200ms 节流
decodeAroundPosition(state.currentTime)
}, [state.hasSupport, state.currentTime, renderFrame, decodeAroundPosition])
const pause = useCallback(() => {
setState((s) => ({ ...s, isPlaying: false }))
@@ -839,61 +850,33 @@ export function useCanvasPlayer(
const seek = useCallback(
async (time: number) => {
const clampedTime = Math.max(0, Math.min(time, totalDuration))
decodeGenerationRef.current++
// 关闭所有解码器、清空队列、重置游标
for (const d of segmentDecodersRef.current.values()) {
try {
if (d.state !== "closed") d.close()
} catch {
/* noop */
}
}
segmentDecodersRef.current.clear()
segmentSampleCursorRef.current.clear()
frameQueueRef.current.clear()
setState((s) => ({ ...s, currentTime: clampedTime }))
playStartOffsetRef.current = clampedTime
playStartRef.current = performance.now()
// 找到 seek 目标片段,从该片段开始解码
const metas = segmentMetaRef.current
let targetIdx = 0,
acc = 0
for (let i = 0; i < metas.length; i++) {
const dur = metas[i].globalEndTime - metas[i].globalStartTime
if (clampedTime < acc + dur) {
targetIdx = i
break
}
acc += dur
}
await decodeSegmentBatch(targetIdx, 60)
await decodeSegmentBatch(Math.min(targetIdx + 1, metas.length - 1), 60)
// seek 时递增解码代数,作废正在进行的异步解码
decodeGenerationRef.current++
// 清空帧队列(clear 内部会 close 所有帧)+ 清空已解码标记
frameQueueRef.current.clear()
decodedSegmentsRef.current.clear()
await decodeAroundPosition(clampedTime)
},
[totalDuration, decodeSegmentBatch],
[totalDuration, decodeAroundPosition],
)
const destroy = useCallback(() => {
isDestroyedRef.current = true
cancelAnimationFrame(rafRef.current)
// 关闭所有持久化解码器
for (const d of segmentDecodersRef.current.values()) {
try {
if (d.state !== "closed") d.close()
} catch {
/* noop */
}
}
segmentDecodersRef.current.clear()
segmentSampleCursorRef.current.clear()
if (decoderRef.current && decoderRef.current.state !== "closed") {
decoderRef.current.close()
}
// 递增代数中止进行中的异步解码,清空帧队列(clear 内部 close 所有帧)
decodeGenerationRef.current++
frameQueueRef.current.clear()
segmentDataRef.current.clear()
segmentMetaRef.current = []
decodedSegmentsRef.current.clear()
}, [])
// ── 预加载下一个片段的数据 ──
@@ -910,7 +893,6 @@ export function useCanvasPlayer(
// ── 初始化:加载并解码所有片段 ──
useEffect(() => {
if (!enabled) return
if (!state.hasSupport || segments.length === 0) {
console.log("[useCanvasPlayer] Skip init:", {
hasSupport: state.hasSupport,
@@ -920,24 +902,10 @@ export function useCanvasPlayer(
}
let cancelled = false
console.log("[useCanvasPlayer] Init start, segments:", segments.length)
const init = async () => {
// ✅ 关键修复:重置销毁标记,允许新的 init 周期正常工作
// destroy() 在 useEffect cleanup 中被调用,将 isDestroyedRef 设为 true
// 如果不重置,后续的 loadSegment / decodeSegment 会立即 return
isDestroyedRef.current = false
// ✅ Strict Mode 修复:init 不再递增 generation
// seek() 和 play() 仍保留 generation 递增用于中止异步解码
// 重置错误状态,避免上一轮的解码错误影响新的 init 周期
setState((s) => ({
...s,
isBuffering: true,
hasDecodeError: false,
errorMessage: "",
isReady: false,
}))
console.log("[useCanvasPlayer] Init start v2_DIAG, segments:", segments.length)
setState((s) => ({ ...s, isBuffering: true }))
// 1. 加载所有片段数据
for (const seg of segments) {
@@ -979,44 +947,32 @@ export function useCanvasPlayer(
segmentMetaRef.current = metas
// 3. 初始化解码:关闭旧解码器,2 个片段各解 60 帧
// 后续由 feedFrames 后台补充
for (const d of segmentDecodersRef.current.values()) {
try {
if (d.state !== "closed") d.close()
} catch {
/* noop */
}
}
segmentDecodersRef.current.clear()
segmentSampleCursorRef.current.clear()
frameQueueRef.current.clear()
// 3. 按需解码:初始只解码3 个片段,后续通过 decodeAroundPosition 动态加载
// 避免一次性全量解码导致 frameQueue 环形缓冲区旧帧被丢弃引发黑屏
decodedSegmentsRef.current.clear()
const initGen = decodeGenerationRef.current
const initialDecodeCount = Math.min(metas.length, 2)
console.log(
`[useCanvasPlayer] Starting init decode: ${initialDecodeCount} segments, metas: ${metas.length}`,
)
const initialDecodeCount = Math.min(metas.length, 3)
for (let i = 0; i < initialDecodeCount; i++) {
if (cancelled) break
// seek 或 destroy 已作废当前初始化
if (decodeGenerationRef.current !== initGen) break
console.log(`[DIAG_v2] Init decode segment ${i}...`)
const meta = metas[i]
const buffer = segmentDataRef.current.get(meta.assetId)
if (!buffer) continue
// 先标记为解码中,防止重复解码
decodedSegmentsRef.current.add(i)
try {
await decodeSegmentBatch(i, 60)
console.log(`[DIAG_v2] Init decode segment ${i} done`)
await decodeSegment(buffer, meta, MAX_INIT_FRAMES)
} catch (e) {
// 解码失败则移除标记,允许后续重试
decodedSegmentsRef.current.delete(i)
console.warn(`[useCanvasPlayer] 初始化解码片段 ${i} 失败:`, e)
}
if (cancelled) break
}
console.log(`[useCanvasPlayer] Init decode finished, cancelled:`, cancelled)
if (!cancelled) {
console.log("[useCanvasPlayer] Init complete, isReady = true, duration:", totalDuration)
console.log("[useCanvasPlayer] Init complete, isReady = true")
setState((s) => ({ ...s, duration: totalDuration, isReady: true, isBuffering: false }))
} else {
console.warn("[useCanvasPlayer] Init was cancelled before completion")
}
}
@@ -1,37 +1,55 @@
/**
* 素材片段调度器 Hook(多 video 元素方案 v3
*
* v3 修复:
* - 所有动态状态存入 ref,tick 为稳定函数,彻底消除 RAF 闭包陷阱
* - 片段切换时先启动下一个 video 再切可见性,消除冻屏间隔
* - 进度更新 200ms 节流
* 素材片段调度器 Hook(多 video 元素方案 v2
* 每个片段对应一个独立 <video> 元素,全部预加载,通过 display 切换实现无缝播放
* 替代单 video + 切 src 方案,消除片段切换延迟
*/
import { useState, useRef, useCallback, useEffect, useMemo } from "react"
/** 单个播放片段 */
export interface PlaybackSegment {
/** 素材 ID */
assetId: string
/** 素材视频 URL */
videoUrl: string
/** 片段在素材中的入点(秒) */
startTime: number
/** 片段在素材中的出点(秒) */
endTime: number
/** 片段在时间线中的顺序 */
order: number
}
/** 调度器返回 */
export interface SegmentSchedulerState {
/** 是否正在播放 */
isPlaying: boolean
/** 当前播放的全局时间(秒) */
currentTime: number
/** 总时长(秒) */
totalDuration: number
/** 当前片段索引 */
currentSegmentIndex: number
/** 当前片段的本地播放时间 */
segmentLocalTime: number
/** 是否已播完 */
isEnded: boolean
/** 是否可以播放(至少有 1 个片段) */
canPlay: boolean
/** 播放 */
play: () => void
/** 暂停 */
pause: () => void
/** 切换播放/暂停 */
togglePlayPause: () => void
/** 跳转到全局时间 */
seekTo: (time: number) => void
/** 每个片段对应的 video 元素 ref 数组 */
videoRefs: React.MutableRefObject<(HTMLVideoElement | null)[]>
}
/**
* 根据全局时间定位对应的片段和本地时间
*/
function findSegmentAtTime(
segments: PlaybackSegment[],
globalTime: number,
@@ -48,6 +66,9 @@ function findSegmentAtTime(
return { index: segments.length - 1, localTime: segments[segments.length - 1].endTime }
}
/**
* 计算每个片段的全局起始时间
*/
function buildTimeline(segments: PlaybackSegment[]): number[] {
const starts: number[] = []
let acc = 0
@@ -58,242 +79,224 @@ function buildTimeline(segments: PlaybackSegment[]): number[] {
return starts
}
/**
* useSegmentScheduler — 多 video 元素版素材片段调度器
*
* 核心改变:
* - 每个片段对应一个独立 <video> 元素(由组件渲染,ref 传入)
* - 所有 video 在挂载时即设置 src + preload="auto",浏览器自动预加载
* - 切换片段仅改 currentSegmentIndex + display,无需重新 load
* - 实现无缝切换,无加载延迟
*/
export function useSegmentScheduler(segments: PlaybackSegment[]): SegmentSchedulerState {
/** 每个片段对应的 video 元素 ref(由组件 JSX 渲染并绑定) */
const videoRefs = useRef<(HTMLVideoElement | null)[]>([])
const [isPlaying, setIsPlaying] = useState(false)
const [currentTime, setCurrentTime] = useState(0)
const [currentSegmentIndex, setCurrentSegmentIndex] = useState(0)
const [isEnded, setIsEnded] = useState(false)
const rafRef = useRef(0)
const rafRef = useRef<number>(0)
const isSeekingRef = useRef(false)
const lastTimeUpdateRef = useRef(0)
// 所有动态值存入 ref,tick 始终读取最新值,不依赖闭包
const segIdxRef = useRef(0)
const segmentsRef = useRef(segments)
const timelineStartsData = useMemo(() => buildTimeline(segments), [segments])
const totalDurationData = useMemo(
// 计算时间线
const timelineStarts = useMemo(() => buildTimeline(segments), [segments])
const totalDuration = useMemo(
() => segments.reduce((sum, seg) => sum + (seg.endTime - seg.startTime), 0),
[segments],
)
const timelineStartsRef = useRef(timelineStartsData)
const totalDurationRef = useRef(totalDurationData)
const isPlayingRef = useRef(false)
segmentsRef.current = segments
timelineStartsRef.current = timelineStartsData
totalDurationRef.current = totalDurationData
const canPlay = segments.length > 0
useEffect(() => {
segIdxRef.current = currentSegmentIndex
}, [currentSegmentIndex])
useEffect(() => {
isPlayingRef.current = isPlaying
}, [isPlaying])
const waitForReady = useCallback((video: HTMLVideoElement, timeout = 3000): Promise<void> => {
if (video.readyState >= 3) return Promise.resolve()
return new Promise((resolve) => {
const onCanPlay = () => {
video.removeEventListener("canplay", onCanPlay)
clearTimeout(timer)
resolve()
}
const timer = setTimeout(() => {
video.removeEventListener("canplay", onCanPlay)
resolve()
}, timeout)
video.addEventListener("canplay", onCanPlay)
})
}, [])
// 当前片段信息
const currentSegment = segments[currentSegmentIndex] || null
const segmentLocalTime = currentSegment
? currentTime - (timelineStarts[currentSegmentIndex] || 0) + currentSegment.startTime
: 0
/**
* 切换到指定片段
* 不改变 src(video 已在 JSX 中设置),仅 seek + 等待可播
*/
const switchToSegment = useCallback(
async (index: number, seekToLocalTime?: number) => {
const segs = segmentsRef.current
const video = videoRefs.current[index]
if (!video || index >= segs.length) return
(index: number, seekToLocalTime?: number): Promise<void> => {
return new Promise((resolve) => {
// 暂停当前视频
const prevVideo = videoRefs.current[currentSegmentIndex]
if (prevVideo) prevVideo.pause()
const seg = segs[index]
const localTime = seekToLocalTime ?? seg.startTime
const oldIdx = segIdxRef.current
const oldVideo = videoRefs.current[oldIdx]
const video = videoRefs.current[index]
if (!video || index >= segments.length) {
resolve()
return
}
if (oldVideo && oldVideo !== video) oldVideo.pause()
const seg = segments[index]
const localTime = seekToLocalTime ?? seg.startTime
if (!video.src && seg.videoUrl) {
video.src = seg.videoUrl
video.load()
}
if (Math.abs(video.currentTime - localTime) > 0.05) {
// 设置播放位置
video.currentTime = localTime
}
segIdxRef.current = index
setCurrentSegmentIndex(index)
// 如果已有足够帧数据,直接 resolve
if (video.readyState >= 2) {
setCurrentSegmentIndex(index)
resolve()
return
}
await waitForReady(video)
// 等待 canplay 事件
const onCanPlay = () => {
video.removeEventListener("canplay", onCanPlay)
clearTimeout(timeoutId)
setCurrentSegmentIndex(index)
resolve()
}
// 10 秒超时保护
const timeoutId = setTimeout(() => {
video.removeEventListener("canplay", onCanPlay)
console.warn(
`[useSegmentScheduler] 片段 ${index} 预加载超时 (10s), readyState=${video.readyState}`,
)
setCurrentSegmentIndex(index)
resolve()
}, 10000)
video.addEventListener("canplay", onCanPlay)
})
},
[waitForReady],
[segments, currentSegmentIndex],
)
// 稳定的 tick 函数,空依赖,所有值从 ref 读取
/** 播放循环 — 检测片段边界并切换 */
const tick = useCallback(() => {
const segs = segmentsRef.current
const idx = segIdxRef.current
const video = videoRefs.current[idx]
const video = videoRefs.current[currentSegmentIndex]
if (!video || isSeekingRef.current) {
rafRef.current = requestAnimationFrame(tick)
return
}
const seg = segs[idx]
const seg = segments[currentSegmentIndex]
if (!seg) return
// 预加载下一个片段
const nextIndex = idx + 1
if (nextIndex < segs.length) {
const nextVideo = videoRefs.current[nextIndex]
if (nextVideo) {
const timeToEnd = seg.endTime - video.currentTime
if (timeToEnd <= 2 && nextVideo.readyState < 3) {
const nextSeg = segs[nextIndex]
if (Math.abs(nextVideo.currentTime - nextSeg.startTime) > 0.5) {
nextVideo.currentTime = nextSeg.startTime
// 检查是否到达出点(容差 0.15s
if (video.currentTime >= seg.endTime - 0.15) {
video.pause()
const nextIndex = currentSegmentIndex + 1
if (nextIndex < segments.length) {
switchToSegment(nextIndex).then(() => {
setIsPlaying(true)
rafRef.current = requestAnimationFrame(tick)
const nextVideo = videoRefs.current[nextIndex]
if (nextVideo) {
const canPlay = () => {
nextVideo
.play()
.catch((e) =>
console.warn("[useSegmentScheduler] auto-play next segment failed:", e),
)
}
if (nextVideo.readyState >= 3) {
canPlay()
} else {
const timeout = setTimeout(canPlay, 300)
nextVideo.addEventListener(
"canplay",
() => {
clearTimeout(timeout)
canPlay()
},
{ once: true },
)
}
}
}
}
}
// 检测片段边界
if (video.currentTime >= seg.endTime - 0.1) {
if (nextIndex < segs.length) {
const nextVideo = videoRefs.current[nextIndex]
const nextSeg = segs[nextIndex]
const accumulatedTime =
(timelineStartsRef.current[idx] || 0) + (seg.endTime - seg.startTime)
if (nextVideo) {
if (Math.abs(nextVideo.currentTime - nextSeg.startTime) > 0.1) {
nextVideo.currentTime = nextSeg.startTime
}
// 先启动下一个视频(muted,可安全同时播放)
nextVideo
.play()
.catch((e) => console.warn("[useSegmentScheduler] next segment play failed:", e))
}
// 立即切换可见性
segIdxRef.current = nextIndex
setCurrentSegmentIndex(nextIndex)
setCurrentTime(accumulatedTime)
lastTimeUpdateRef.current = 0
setIsPlaying(true)
// 下一帧暂停旧视频(让新视频先渲染,避免冻屏)
const oldVideo = video
requestAnimationFrame(() => {
oldVideo.pause()
})
rafRef.current = requestAnimationFrame(tick)
return
const accumulatedTime =
(timelineStarts[currentSegmentIndex] || 0) + (seg.endTime - seg.startTime)
setCurrentTime(accumulatedTime)
} else {
video.pause()
setIsPlaying(false)
setIsEnded(true)
setCurrentTime(totalDurationRef.current)
setCurrentTime(totalDuration)
return
}
}
const globalTime = (timelineStartsRef.current[idx] || 0) + (video.currentTime - seg.startTime)
const now = performance.now()
if (now - lastTimeUpdateRef.current >= 200) {
lastTimeUpdateRef.current = now
setCurrentTime(Math.max(0, Math.min(globalTime, totalDurationRef.current)))
} else {
const globalTime =
(timelineStarts[currentSegmentIndex] || 0) + (video.currentTime - seg.startTime)
setCurrentTime(Math.max(0, Math.min(globalTime, totalDuration)))
}
rafRef.current = requestAnimationFrame(tick)
}, [])
}, [segments, currentSegmentIndex, timelineStarts, totalDuration, switchToSegment])
/** 播放 */
const play = useCallback(async () => {
if (!canPlay) return
setIsEnded(false)
const idx = segIdxRef.current
const video = videoRefs.current[idx]
if (!video) return
if (idx === 0 && video.readyState < 2) {
if (!video.src && segmentsRef.current[0]?.videoUrl) {
video.src = segmentsRef.current[0].videoUrl
video.load()
}
await waitForReady(video)
setIsEnded(false)
// 确保第一段可播放
const firstVideo = videoRefs.current[0]
if (firstVideo && currentSegmentIndex === 0 && firstVideo.readyState < 2) {
await switchToSegment(0)
}
const video = videoRefs.current[currentSegmentIndex]
if (!video) return
try {
await video.play()
const playPromise = video.play()
if (playPromise !== undefined) {
await playPromise
}
setIsPlaying(true)
cancelAnimationFrame(rafRef.current)
rafRef.current = requestAnimationFrame(tick)
} catch (err) {
console.warn("[useSegmentScheduler] 播放失败:", err)
}
}, [canPlay, waitForReady, tick])
}, [canPlay, switchToSegment, tick, currentSegmentIndex])
/** 暂停 */
const pause = useCallback(() => {
const video = videoRefs.current[segIdxRef.current]
const video = videoRefs.current[currentSegmentIndex]
if (video) video.pause()
setIsPlaying(false)
cancelAnimationFrame(rafRef.current)
}, [])
}, [currentSegmentIndex])
/** 切换播放/暂停 */
const togglePlayPause = useCallback(() => {
if (isPlayingRef.current) {
if (isPlaying) {
pause()
} else {
if (isEnded) {
// 播放结束后再次播放,从头开始
setIsEnded(false)
lastTimeUpdateRef.current = 0
const firstVideo = videoRefs.current[0]
if (firstVideo) {
videoRefs.current.forEach((v, i) => {
if (v && i !== 0) v.pause()
})
firstVideo.currentTime = segmentsRef.current[0]?.startTime || 0
segIdxRef.current = 0
setCurrentSegmentIndex(0)
setCurrentTime(0)
firstVideo
.play()
.then(() => {
setIsPlaying(true)
cancelAnimationFrame(rafRef.current)
rafRef.current = requestAnimationFrame(tick)
})
.catch((e) => console.warn("[useSegmentScheduler] restart failed:", e))
}
switchToSegment(0, segments[0]?.startTime).then(() => {
const video = videoRefs.current[0]
if (video) {
video.play().catch((e) => console.warn("[useSegmentScheduler] restart play failed:", e))
setIsPlaying(true)
setCurrentTime(0)
rafRef.current = requestAnimationFrame(tick)
}
})
} else {
play()
}
}
}, [isEnded, pause, play, tick])
}, [isPlaying, isEnded, pause, play, switchToSegment, segments, tick])
/** 跳转到指定全局时间 */
const seekTo = useCallback(
async (time: number) => {
if (!canPlay) return
const clampedTime = Math.max(0, Math.min(time, totalDurationRef.current))
const { index, localTime } = findSegmentAtTime(segmentsRef.current, clampedTime)
const clampedTime = Math.max(0, Math.min(time, totalDuration))
const { index, localTime } = findSegmentAtTime(segments, clampedTime)
isSeekingRef.current = true
cancelAnimationFrame(rafRef.current)
if (index !== segIdxRef.current) {
if (index !== currentSegmentIndex) {
await switchToSegment(index, localTime)
} else {
const video = videoRefs.current[index]
@@ -302,54 +305,48 @@ export function useSegmentScheduler(segments: PlaybackSegment[]): SegmentSchedul
setCurrentTime(clampedTime)
setIsEnded(false)
lastTimeUpdateRef.current = 0
if (isPlayingRef.current) {
const video = videoRefs.current[index]
if (video) {
video.play().catch(() => {})
}
rafRef.current = requestAnimationFrame(tick)
}
setTimeout(() => {
isSeekingRef.current = false
}, 200)
},
[canPlay, switchToSegment, tick],
[canPlay, totalDuration, segments, currentSegmentIndex, switchToSegment],
)
// 确保 videoRefs 数组长度与 segments 一致 + 强制预加载
useEffect(() => {
videoRefs.current = videoRefs.current.slice(0, segments.length)
while (videoRefs.current.length < segments.length) {
videoRefs.current.push(null)
}
// 强制预加载:所有 video 元素挂载后,调用 load() 确保浏览器真正开始加载数据
videoRefs.current.forEach((video) => {
if (video) {
video.load()
}
})
}, [segments])
// 组件卸载时清理
useEffect(() => {
return () => {
cancelAnimationFrame(rafRef.current)
}
}, [])
// 片段列表变化时重置
useEffect(() => {
cancelAnimationFrame(rafRef.current)
segIdxRef.current = 0
setIsPlaying(false)
setCurrentTime(0)
setCurrentSegmentIndex(0)
setIsEnded(false)
}, [segments])
const currentSegment = segments[currentSegmentIndex] || null
const segmentLocalTime = currentSegment
? currentTime - (timelineStartsRef.current[currentSegmentIndex] || 0) + currentSegment.startTime
: 0
return {
isPlaying,
currentTime,
totalDuration: totalDurationData,
totalDuration,
currentSegmentIndex,
segmentLocalTime,
isEnded,
+1 -294
View File
@@ -182,7 +182,6 @@ class TestUnifiedCoverPipelineEndpoint:
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch("packages.shared.storage.get_shared_storage_service") as mock_storage_getter,
patch("packages.shared.mediakit_client.get_mediakit_client") as mock_mk_getter,
):
mock_repo = MagicMock()
mock_repo.get.return_value = mock_task
@@ -195,13 +194,6 @@ class TestUnifiedCoverPipelineEndpoint:
mock_storage_svc.get_url.return_value = "https://oss.example.com/rendered/plan-2/video.mp4"
mock_storage_getter.return_value = mock_storage_svc
# MediaKit 抽帧也返回 None,模拟最终失败
mock_mk = MagicMock()
mock_mk.is_available = True
mock_mk.extract_frames.return_value = None
mock_mk_getter.return_value = mock_mk
# body 不传 asset_ids,步骤 E2 不会进入
from app.api.routes.generation_cover import generate_cover
with pytest.raises(HTTPException) as exc_info:
@@ -215,6 +207,7 @@ class TestUnifiedCoverPipelineEndpoint:
)
assert exc_info.value.status_code == 400
assert "封面尚未生成" in exc_info.value.detail
def test_cover_url_found_via_source_edit_plan(self):
"""步骤B:通过 source_edit_plan_id 找到预览任务的 cover_url。"""
@@ -325,157 +318,6 @@ class TestUnifiedCoverPipelineEndpoint:
template_id="template-y",
)
def test_cover_url_found_via_cover_candidates_image_url(self):
"""步骤Dplan.config.cover_candidates 有 image_url 时,直接使用第一个候选封面。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
mock_plan = MagicMock()
# 步骤A/B/C 都找不到,进入步骤D
mock_plan.config = {
"rendered_storage_key": "rendered/plan-z/video.mp4", # 必须有预览视频才能通过前置检查
"cover_candidates": [
{"image_url": "https://oss.example.com/candidates/cover-1.jpg", "score": 0.95},
{"image_url": "https://oss.example.com/candidates/cover-2.jpg", "score": 0.80},
],
}
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_db = MagicMock()
body = GenerateCoverRequest(cover_type="ai_frame")
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch("app.api.routes.generation_cover.normalize_plan_config") as mock_normalize,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_repo.list_by_source_edit_plan.return_value = []
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
mock_normalize.return_value = {
"cover": {"type": "ai_frame", "image_url": "https://oss.example.com/candidates/cover-1.jpg"}
}
from app.api.routes.generation_cover import generate_cover
result = generate_cover(
body=body,
template_id="template-z",
plan_id="plan-z",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=MagicMock(),
)
# 步骤D从 cover_candidates 第一个元素的 image_url 提取封面
assert result.cover["image_url"] == "https://oss.example.com/candidates/cover-1.jpg"
# 验证 plan.config 被更新(至少调用一次:rendered_storage_key + cover
assert mock_plan_svc.update_plan_config.call_count >= 1
def test_cover_url_found_via_cover_candidates_url_key(self):
"""步骤Dcover_candidates 用 url 键(非 image_url)时,也能正确提取。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
mock_plan = MagicMock()
mock_plan.config = {
"rendered_storage_key": "rendered/plan-w/video.mp4", # 必须有预览视频才能通过前置检查
"cover_candidates": [
{"url": "https://oss.example.com/candidates/alt-cover.jpg"},
],
}
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_db = MagicMock()
body = GenerateCoverRequest(cover_type="ai_frame")
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch("app.api.routes.generation_cover.normalize_plan_config") as mock_normalize,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_repo.list_by_source_edit_plan.return_value = []
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
mock_normalize.return_value = {
"cover": {"type": "ai_frame", "image_url": "https://oss.example.com/candidates/alt-cover.jpg"}
}
from app.api.routes.generation_cover import generate_cover
result = generate_cover(
body=body,
template_id="template-w",
plan_id="plan-w",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=MagicMock(),
)
# 步骤D fallback 到 url 键
assert result.cover["image_url"] == "https://oss.example.com/candidates/alt-cover.jpg"
def test_cover_candidates_skips_non_dict_first_element(self):
"""步骤Dcover_candidates 第一个元素不是 dict 时,安全跳过不崩溃。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
from fastapi import HTTPException
mock_plan = MagicMock()
mock_plan.config = {
"rendered_storage_key": "rendered/plan-skip/video.mp4", # 必须有预览视频才能通过前置检查
"cover_candidates": ["not-a-dict", 42, None],
}
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_db = MagicMock()
body = GenerateCoverRequest(cover_type="ai_frame")
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch("packages.shared.storage.get_shared_storage_service") as mock_storage_getter,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_repo.list_by_source_edit_plan.return_value = []
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
# storage fallback 也找不到封面
mock_storage_svc = MagicMock()
mock_storage_svc.get_url.return_value = ""
mock_storage_getter.return_value = mock_storage_svc
from app.api.routes.generation_cover import generate_cover
# 所有步骤都失败,应返回 400
with pytest.raises(HTTPException) as exc_info:
generate_cover(
body=body,
template_id="template-skip",
plan_id="plan-skip",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=MagicMock(),
)
assert exc_info.value.status_code == 400
class TestSourceEditPlanFallback:
"""测试步骤 2.5:通过 source_edit_plan_id 查找预览视频兜底逻辑。"""
@@ -845,138 +687,3 @@ class TestUploadCoverType:
assert result.cover["image_url"] == "https://oss.example.com/uploaded/my-cover.png"
# 验证没有调用任何预览视频查找逻辑
# (normalize_plan_config 是唯一被调用的外部函数)
def test_cover_extracted_from_source_asset_when_no_preview(self):
"""步骤E2:无后端渲染产物时,直接从用户选择的视频素材抽帧。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
mock_plan = MagicMock()
mock_plan.config = {} # 无 rendered_storage_key,无 generation_task_id
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_db = MagicMock()
# 模拟视频素材
mock_asset = MagicMock()
mock_asset.file_type = "video"
mock_asset.storage_key = "uploads/source-clip.mp4"
mock_asset_repo = MagicMock()
mock_asset_repo.get.return_value = mock_asset
mock_mk = MagicMock()
mock_mk.is_available = True
mock_mk.extract_frames.return_value = [{"image_url": "https://mediakit.internal/frame-abc.jpg"}]
mock_storage = MagicMock()
mock_storage.get_url.return_value = "https://oss.example.com/uploads/source-clip.mp4"
body = GenerateCoverRequest(
cover_type="ai_frame",
asset_ids=["asset-video-1"],
)
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch(
"packages.adapters.sqlalchemy_impl.asset_repository.SQLAlchemyAssetRepository",
return_value=mock_asset_repo,
),
patch("packages.shared.mediakit_client.get_mediakit_client", return_value=mock_mk),
patch("packages.shared.storage.get_shared_storage_service", return_value=mock_storage),
patch(
"app.api.routes.generation_cover._persist_cover_frame",
return_value="https://oss.example.com/covers/final.jpg",
),
patch("app.api.routes.generation_cover.normalize_plan_config") as mock_normalize,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_repo.list_by_source_edit_plan.return_value = []
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
mock_normalize.return_value = {
"cover": {"type": "ai_frame", "image_url": "https://oss.example.com/covers/final.jpg"}
}
from app.api.routes.generation_cover import generate_cover
result = generate_cover(
body=body,
template_id="tpl-source",
plan_id="plan-source",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=MagicMock(),
)
assert result.cover["image_url"] == "https://oss.example.com/covers/final.jpg"
mock_mk.extract_frames.assert_called_once()
# 确保用的是源素材 URL
call_kwargs = mock_mk.extract_frames.call_args.kwargs
assert "source-clip.mp4" in call_kwargs["video_url"]
def test_step_e_skips_non_video_assets(self):
"""步骤E2:asset_ids 里只有图片素材时,不调用 MediaKit 并返回 400。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
from fastapi import HTTPException
mock_plan = MagicMock()
mock_plan.config = {}
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_db = MagicMock()
mock_image_asset = MagicMock()
mock_image_asset.file_type = "image"
mock_image_asset.storage_key = "uploads/photo.png"
mock_asset_repo = MagicMock()
mock_asset_repo.get.return_value = mock_image_asset
mock_mk = MagicMock()
mock_mk.is_available = True
body = GenerateCoverRequest(
cover_type="ai_frame",
asset_ids=["asset-img-1"],
)
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch(
"packages.adapters.sqlalchemy_impl.asset_repository.SQLAlchemyAssetRepository",
return_value=mock_asset_repo,
),
patch("packages.shared.mediakit_client.get_mediakit_client", return_value=mock_mk),
patch("packages.shared.storage.get_shared_storage_service") as mock_storage_getter,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_repo.list_by_source_edit_plan.return_value = []
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
mock_storage_getter.return_value = MagicMock()
from app.api.routes.generation_cover import generate_cover
with pytest.raises(HTTPException) as exc_info:
generate_cover(
body=body,
template_id="tpl-img",
plan_id="plan-img",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=MagicMock(),
)
assert exc_info.value.status_code == 400
mock_mk.extract_frames.assert_not_called()
+5 -24
View File
@@ -13,18 +13,6 @@ os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret")
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
def _patch_session_local(mock_session):
"""Patch worker_app.db.SessionLocal robustly even when other tests
have pre-registered a MagicMock for worker_app.db in sys.modules.
Uses patch.dict to inject a clean module so that
'from worker_app.db import SessionLocal' resolves correctly."""
from types import ModuleType
_fresh_db = ModuleType("worker_app.db")
_fresh_db.SessionLocal = lambda *a, **kw: mock_session
return patch.dict(sys.modules, {"worker_app.db": _fresh_db})
class TestLoadTemplateSegmentDurations:
"""_load_template_segment_durations 单元测试 (covers lines 198-226)."""
@@ -52,7 +40,8 @@ class TestLoadTemplateSegmentDurations:
mock_session = MagicMock()
mock_session.query.return_value = mock_query
with _patch_session_local(mock_session):
# Patch at the source module since it's imported inside the function
with patch("worker_app.db.SessionLocal", return_value=mock_session):
result = _load_template_segment_durations("tpl_123")
assert result == [5.0, 8.0, 3.0]
@@ -74,24 +63,16 @@ class TestLoadTemplateSegmentDurations:
mock_session = MagicMock()
mock_session.query.return_value = mock_query
with _patch_session_local(mock_session):
with patch("worker_app.db.SessionLocal", return_value=mock_session):
result = _load_template_segment_durations("tpl_456")
assert result == [5.0]
def test_db_error_returns_empty(self):
"""数据库异常返回空列表,不抛出。"""
from types import ModuleType
from worker_app.tasks.generation import _load_template_segment_durations
_err_db = ModuleType("worker_app.db")
def _raise(*a, **kw):
raise Exception("DB down")
_err_db.SessionLocal = _raise
with patch.dict(sys.modules, {"worker_app.db": _err_db}):
with patch("worker_app.db.SessionLocal", side_effect=Exception("DB down")):
result = _load_template_segment_durations("tpl_789")
assert result == []
@@ -105,7 +86,7 @@ class TestLoadTemplateSegmentDurations:
mock_session = MagicMock()
mock_session.query.return_value = mock_query
with _patch_session_local(mock_session):
with patch("worker_app.db.SessionLocal", return_value=mock_session):
result = _load_template_segment_durations("tpl_empty")
assert result == []
+92 -78
View File
@@ -1,6 +1,7 @@
"""Unit tests for apps/api/app/api/routes/health.py
覆盖 _check_database() 和 _check_migrations() 中 psycopg3 连接逻辑。
确保增量覆盖率 ≥ 60%(目标覆盖 lines 52, 127)。
"""
from unittest.mock import AsyncMock, MagicMock, patch
@@ -8,70 +9,63 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
def _make_cursor(fetchone_result=None):
"""Create a mock cursor with context manager support."""
cur = MagicMock()
cur.__enter__ = MagicMock(return_value=cur)
cur.__exit__ = MagicMock(return_value=False)
if fetchone_result is not None:
cur.fetchone.return_value = fetchone_result
return cur
def _make_conn(cursor_result=None):
conn = MagicMock()
conn.cursor.return_value = cursor_result or _make_cursor()
return conn
@pytest.mark.asyncio
class TestCheckDatabase:
"""Tests for _check_database() health check function."""
async def test_check_database_success(self):
mock_cur = _make_cursor(fetchone_result=(1,))
mock_conn = _make_conn(mock_cur)
mock_settings = MagicMock()
@patch("apps.api.app.api.routes.health.settings")
@patch("apps.api.app.api.routes.health.psycopg.connect")
async def test_check_database_success(self, mock_connect, mock_settings):
"""PostgreSQL 连接成功时返回 healthy。"""
mock_settings.USE_IN_MEMORY_DB = False
mock_settings.DATABASE_URL = "postgresql+psycopg://test:test@localhost/test"
from apps.api.app.api.routes import health
# Mock connection and cursor
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_cursor.__enter__ = MagicMock(return_value=mock_cursor)
mock_cursor.__exit__ = MagicMock(return_value=False)
mock_cursor.fetchone.return_value = (1,)
mock_conn.cursor.return_value = mock_cursor
mock_connect.return_value = mock_conn
with patch.object(health, "psycopg") as mock_psycopg, patch.object(health, "settings", mock_settings):
mock_psycopg.connect.return_value = mock_conn
result = await health._check_database()
from apps.api.app.api.routes.health import _check_database
result = await _check_database()
assert result["status"] == "healthy"
assert result["type"] == "postgresql"
assert result["message"] == "Database connection successful"
mock_psycopg.connect.assert_called_once_with(
mock_connect.assert_called_once_with(
"postgresql+psycopg://test:test@localhost/test", connect_timeout=3
)
mock_cur.execute.assert_called_once_with("SELECT 1")
mock_cursor.execute.assert_called_once_with("SELECT 1")
mock_conn.close.assert_called_once()
async def test_check_database_connection_failure(self):
mock_settings = MagicMock()
@patch("apps.api.app.api.routes.health.settings")
@patch("apps.api.app.api.routes.health.psycopg.connect")
async def test_check_database_connection_failure(self, mock_connect, mock_settings):
"""PostgreSQL 连接失败时返回 unhealthy。"""
mock_settings.USE_IN_MEMORY_DB = False
mock_settings.DATABASE_URL = "postgresql+psycopg://test:test@localhost/test"
mock_connect.side_effect = Exception("connection refused")
from apps.api.app.api.routes import health
from apps.api.app.api.routes.health import _check_database
with patch.object(health, "psycopg") as mock_psycopg, patch.object(health, "settings", mock_settings):
mock_psycopg.connect.side_effect = Exception("connection refused")
result = await health._check_database()
result = await _check_database()
assert result["status"] == "unhealthy"
assert result["type"] == "postgresql"
assert "connection refused" in result["message"]
async def test_check_database_in_memory(self):
mock_settings = MagicMock()
@patch("apps.api.app.api.routes.health.settings")
async def test_check_database_in_memory(self, mock_settings):
"""使用内存数据库时跳过 PostgreSQL 检查。"""
mock_settings.USE_IN_MEMORY_DB = True
from apps.api.app.api.routes import health
from apps.api.app.api.routes.health import _check_database
with patch.object(health, "settings", mock_settings):
result = await health._check_database()
result = await _check_database()
assert result["status"] == "healthy"
assert result["type"] == "in_memory"
@@ -79,66 +73,80 @@ class TestCheckDatabase:
@pytest.mark.asyncio
class TestCheckMigrations:
"""Tests for _check_migrations() health check function."""
async def test_check_migrations_success(self):
mock_cur = _make_cursor(fetchone_result=(5,))
mock_conn = _make_conn(mock_cur)
mock_settings = MagicMock()
@patch("apps.api.app.api.routes.health.settings")
@patch("apps.api.app.api.routes.health.psycopg.connect")
async def test_check_migrations_success(self, mock_connect, mock_settings):
"""所有迁移表存在时返回 healthy。"""
mock_settings.USE_IN_MEMORY_DB = False
mock_settings.DATABASE_URL = "postgresql+psycopg://test:test@localhost/test"
from apps.api.app.api.routes import health
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_cursor.__enter__ = MagicMock(return_value=mock_cursor)
mock_cursor.__exit__ = MagicMock(return_value=False)
mock_cursor.fetchone.return_value = (5,) # 5 tables found
mock_conn.cursor.return_value = mock_cursor
mock_connect.return_value = mock_conn
with patch.object(health, "psycopg") as mock_psycopg, patch.object(health, "settings", mock_settings):
mock_psycopg.connect.return_value = mock_conn
result = await health._check_migrations()
from apps.api.app.api.routes.health import _check_migrations
result = await _check_migrations()
assert result["status"] == "healthy"
assert result["message"] == "Database migrations applied"
mock_psycopg.connect.assert_called_once_with(
mock_connect.assert_called_once_with(
"postgresql+psycopg://test:test@localhost/test", connect_timeout=3
)
mock_conn.close.assert_called_once()
async def test_check_migrations_missing_tables(self):
mock_cur = _make_cursor(fetchone_result=(2,))
mock_conn = _make_conn(mock_cur)
mock_settings = MagicMock()
@patch("apps.api.app.api.routes.health.settings")
@patch("apps.api.app.api.routes.health.psycopg.connect")
async def test_check_migrations_missing_tables(self, mock_connect, mock_settings):
"""迁移表不完整时返回 unhealthy。"""
mock_settings.USE_IN_MEMORY_DB = False
mock_settings.DATABASE_URL = "postgresql+psycopg://test:test@localhost/test"
from apps.api.app.api.routes import health
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_cursor.__enter__ = MagicMock(return_value=mock_cursor)
mock_cursor.__exit__ = MagicMock(return_value=False)
mock_cursor.fetchone.return_value = (2,) # Only 2 of 5 tables
mock_conn.cursor.return_value = mock_cursor
mock_connect.return_value = mock_conn
with patch.object(health, "psycopg") as mock_psycopg, patch.object(health, "settings", mock_settings):
mock_psycopg.connect.return_value = mock_conn
result = await health._check_migrations()
from apps.api.app.api.routes.health import _check_migrations
result = await _check_migrations()
assert result["status"] == "unhealthy"
assert "Missing tables" in result["message"]
assert "2/5" in result["message"]
async def test_check_migrations_connection_failure(self):
mock_settings = MagicMock()
@patch("apps.api.app.api.routes.health.settings")
@patch("apps.api.app.api.routes.health.psycopg.connect")
async def test_check_migrations_connection_failure(self, mock_connect, mock_settings):
"""数据库连接失败时返回 unhealthy。"""
mock_settings.USE_IN_MEMORY_DB = False
mock_settings.DATABASE_URL = "postgresql+psycopg://test:test@localhost/test"
mock_connect.side_effect = Exception("connection refused")
from apps.api.app.api.routes import health
from apps.api.app.api.routes.health import _check_migrations
with patch.object(health, "psycopg") as mock_psycopg, patch.object(health, "settings", mock_settings):
mock_psycopg.connect.side_effect = Exception("connection refused")
result = await health._check_migrations()
result = await _check_migrations()
assert result["status"] == "unhealthy"
assert "Migration check failed" in result["message"]
async def test_check_migrations_in_memory(self):
mock_settings = MagicMock()
@patch("apps.api.app.api.routes.health.settings")
async def test_check_migrations_in_memory(self, mock_settings):
"""使用内存数据库时跳过迁移检查。"""
mock_settings.USE_IN_MEMORY_DB = True
from apps.api.app.api.routes import health
from apps.api.app.api.routes.health import _check_migrations
with patch.object(health, "settings", mock_settings):
result = await health._check_migrations()
result = await _check_migrations()
assert result["status"] == "healthy"
assert "no migrations needed" in result["message"]
@@ -146,27 +154,33 @@ class TestCheckMigrations:
@pytest.mark.asyncio
class TestStartupCheck:
"""Tests for startup_check() endpoint."""
async def test_startup_all_healthy(self):
from apps.api.app.api.routes import health
@patch("apps.api.app.api.routes.health._check_migrations")
@patch("apps.api.app.api.routes.health._check_database")
async def test_startup_all_healthy(self, mock_db, mock_mig):
"""所有检查通过时返回 started。"""
mock_db.return_value = {"status": "healthy"}
mock_mig.return_value = {"status": "healthy"}
with patch.object(health, "_check_migrations", new_callable=AsyncMock) as mock_mig, patch.object(health, "_check_database", new_callable=AsyncMock) as mock_db:
mock_db.return_value = {"status": "healthy"}
mock_mig.return_value = {"status": "healthy"}
result = await health.startup_check()
from apps.api.app.api.routes.health import startup_check
result = await startup_check()
assert result["status"] == "started"
async def test_startup_db_unhealthy(self):
import json
@patch("apps.api.app.api.routes.health._check_migrations")
@patch("apps.api.app.api.routes.health._check_database")
async def test_startup_db_unhealthy(self, mock_db, mock_mig):
"""数据库不健康时返回 starting + 503。"""
mock_db.return_value = {"status": "unhealthy", "message": "fail"}
mock_mig.return_value = {"status": "healthy"}
from apps.api.app.api.routes import health
from apps.api.app.api.routes.health import startup_check
with patch.object(health, "_check_migrations", new_callable=AsyncMock) as mock_mig, patch.object(health, "_check_database", new_callable=AsyncMock) as mock_db:
mock_db.return_value = {"status": "unhealthy", "message": "fail"}
mock_mig.return_value = {"status": "healthy"}
result = await health.startup_check()
result = await startup_check()
assert result.status_code == 503
import json
body = json.loads(result.body)
assert body["status"] == "starting"