This commit is contained in:
2026-07-21 19:22:50 +08:00
parent b374a541ea
commit 7683ebb186
65 changed files with 5067 additions and 1741 deletions
File diff suppressed because one or more lines are too long
+36 -36
View File
@@ -1,37 +1,37 @@
<!doctype html> <!doctype html>
<html lang="zh-CN"> <html lang="zh-CN">
<head> <head>
<meta charset="UTF-8" /> <meta charset="UTF-8" />
<link rel="icon" type="image/svg+xml" href="/favicon.svg" /> <link rel="icon" type="image/svg+xml" href="/favicon.svg" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" /> <meta name="viewport" content="width=device-width, initial-scale=1.0" />
<link rel="preconnect" href="https://fonts.googleapis.com" /> <link rel="preconnect" href="https://fonts.googleapis.com" />
<link rel="preconnect" href="https://fonts.gstatic.com" crossorigin /> <link rel="preconnect" href="https://fonts.gstatic.com" crossorigin />
<link href="https://fonts.googleapis.com/css2?family=Outfit:wght@300;400;500;600;700&display=swap" rel="stylesheet" /> <link href="https://fonts.googleapis.com/css2?family=Outfit:wght@300;400;500;600;700&display=swap" rel="stylesheet" />
<title>后台管理</title> <title>后台管理</title>
<script> <script>
(function() { (function() {
var cached = localStorage.getItem('siteInfo'); var cached = localStorage.getItem('siteInfo');
if (cached) { if (cached) {
try { try {
var info = JSON.parse(cached); var info = JSON.parse(cached);
if (info.siteName) { if (info.siteName) {
document.title = info.siteName + ' - 管理后台'; document.title = info.siteName + ' - 管理后台';
} }
if (info.siteLogo) { if (info.siteLogo) {
var link = document.querySelector('link[rel="icon"]'); var link = document.querySelector('link[rel="icon"]');
if (link) { if (link) {
link.href = info.siteLogo; link.href = info.siteLogo;
link.type = 'image/png'; link.type = 'image/png';
} }
} }
} catch (e) {} } catch (e) {}
} }
})(); })();
</script> </script>
<script type="module" crossorigin src="/assets/index-Dhv0zQu5.js"></script> <script type="module" crossorigin src="/assets/index-BLkGzXo7.js"></script>
<link rel="stylesheet" crossorigin href="/assets/index-D7ShJUt4.css"> <link rel="stylesheet" crossorigin href="/assets/index-D7ShJUt4.css">
</head> </head>
<body> <body>
<div id="root"></div> <div id="root"></div>
</body> </body>
</html> </html>
+24 -16
View File
@@ -618,28 +618,26 @@ export async function getOpenTypeAll(): Promise<{ data: OpenTypeItem[] }> {
// ── Generation Records (Admin) ───────────────────────────── // ── Generation Records (Admin) ─────────────────────────────
export async function getAdminGenerationRecords(params?: { export async function getAdminGenerationRecords(params?: {
userId?: string; status?: string; page?: number; pageSize?: number; userId?: string;
status?: string;
engineId?: string;
includeMediaReferences?: boolean;
page?: number;
pageSize?: number;
}): Promise<{ total: number; items: any[] }> { }): Promise<{ total: number; items: any[] }> {
const q = new URLSearchParams(); const q = new URLSearchParams();
if (params?.userId) q.set('user_id', params.userId); if (params?.userId) q.set('user_id', params.userId);
if (params?.status) q.set('status', params.status); if (params?.status) q.set('status', params.status);
if (params?.engineId) q.set('engine_id', params.engineId);
if (params?.includeMediaReferences !== undefined) {
q.set('include_media_references', String(params.includeMediaReferences));
}
if (params?.page) q.set('page', String(params.page)); if (params?.page) q.set('page', String(params.page));
if (params?.pageSize) q.set('page_size', String(params.pageSize)); if (params?.pageSize) q.set('page_size', String(params.pageSize));
const qs = q.toString(); const qs = q.toString();
return api.get(`/admin/generation-records${qs ? `?${qs}` : ''}`); return api.get(`/admin/generation-records${qs ? `?${qs}` : ''}`);
} }
export async function adminUpdateGenerationStatus(
recordId: string, status: string, videoUrl?: string
): Promise<void> {
await api.put(`/admin/generation-records/${recordId}/status`, { status, video_url: videoUrl });
}
export async function adminGenerateVideo(
recordId: string, aspectRatio: string, resolution: string, image_size: string
): Promise<void> {
await api.post(`/admin/generation-records/${recordId}/generate`, { aspect_ratio: aspectRatio, resolution, image_size });
}
// ── Generation AI Engines (Admin) ───────────────────────────── // ── Generation AI Engines (Admin) ─────────────────────────────
@@ -685,8 +683,13 @@ export async function getAdminHotOpeningTasks(params?: AdminHotOpeningTaskQueryP
return api.get<HotOpeningTaskListOut>(`/hot-opening-replications/tasks${qs ? `?${qs}` : ''}`); return api.get<HotOpeningTaskListOut>(`/hot-opening-replications/tasks${qs ? `?${qs}` : ''}`);
} }
export async function getAdminHotOpeningTaskDetail(projectId: string): Promise<ReplicationProjectDetailOut> { export async function getAdminHotOpeningTaskDetail(
return api.get<ReplicationProjectDetailOut>(`/hot-opening-replications/tasks/${projectId}`); projectId: string,
flowVersion: 'v1' | 'v2',
): Promise<ReplicationProjectDetailOut> {
return flowVersion === 'v2'
? api.get<ReplicationProjectDetailOut>(`/v2/hot-opening-replications/tasks/${projectId}`)
: api.get<ReplicationProjectDetailOut>(`/hot-opening-replications/tasks/${projectId}`);
} }
export async function getAdminShotTaskSets(params?: AdminShotTaskSetQueryParams): Promise<ShotTaskSetListOut> { export async function getAdminShotTaskSets(params?: AdminShotTaskSetQueryParams): Promise<ShotTaskSetListOut> {
@@ -725,8 +728,13 @@ export async function getAdminShotSegmentDetail(segmentId: string): Promise<Shot
return api.get<ShotSegmentDetailOut>(`/shot-replications/segments/${segmentId}`); return api.get<ShotSegmentDetailOut>(`/shot-replications/segments/${segmentId}`);
} }
export async function getAdminShotProjectDetail(projectId: string): Promise<ReplicationProjectDetailOut> { export async function getAdminShotProjectDetail(
return api.get<ReplicationProjectDetailOut>(`/shot-replications/projects/${projectId}`); projectId: string,
flowVersion: 'v1' | 'v2',
): Promise<ReplicationProjectDetailOut> {
return flowVersion === 'v2'
? api.get<ReplicationProjectDetailOut>(`/v2/shot-replications/projects/${projectId}`)
: api.get<ReplicationProjectDetailOut>(`/shot-replications/projects/${projectId}`);
} }
@@ -22,10 +22,9 @@ import {
CloseCircleOutlined, CloseCircleOutlined,
SearchOutlined, SearchOutlined,
VideoCameraOutlined, VideoCameraOutlined,
ExclamationCircleOutlined,
FileImageOutlined, FileImageOutlined,
} from '@ant-design/icons'; } from '@ant-design/icons';
import { getAdminGenerationRecords, adminUpdateGenerationStatus, adminGenerateVideo } from '../api'; import { getAdminGenerationRecords, getVideoEngines, getImageEngines } from '../api';
import type { AdminGenerationRecord, GenerationAIMediaReference } from '../types'; import type { AdminGenerationRecord, GenerationAIMediaReference } from '../types';
import { formatDate } from '../utils/formatDate'; import { formatDate } from '../utils/formatDate';
@@ -194,13 +193,14 @@ const AdminGenerationRecords: React.FC = () => {
const [pageSize] = useState(20); const [pageSize] = useState(20);
const [filterStatus, setFilterStatus] = useState<string>(''); const [filterStatus, setFilterStatus] = useState<string>('');
const [filterUserId, setFilterUserId] = useState<string>(''); const [filterUserId, setFilterUserId] = useState<string>('');
const [filterEngineId, setFilterEngineId] = useState<string>('');
const [filterIncludeMedia, setFilterIncludeMedia] = useState<'' | 'true' | 'false'>('');
const [engineOptions, setEngineOptions] = useState<Array<{ value: string; label: string }>>([]);
const [reloadKey, setReloadKey] = useState(0); const [reloadKey, setReloadKey] = useState(0);
const [preview, setPreview] = useState<AdminGenerationRecord | null>(null); const [preview, setPreview] = useState<AdminGenerationRecord | null>(null);
const [resourceState, setResourceState] = useState<PreviewResourceState>(EMPTY_RESOURCE_STATE); const [resourceState, setResourceState] = useState<PreviewResourceState>(EMPTY_RESOURCE_STATE);
const [videoPlaying, setVideoPlaying] = useState(false); const [videoPlaying, setVideoPlaying] = useState(false);
const videoRef = useRef<HTMLVideoElement | null>(null); const videoRef = useRef<HTMLVideoElement | null>(null);
const [updating, setUpdating] = useState<string | null>(null);
const [genModal, setGenModal] = useState<{ record: AdminGenerationRecord; ratio: string; resolution: string; image_size: string } | null>(null);
const load = useCallback(async () => { const load = useCallback(async () => {
setLoading(true); setLoading(true);
@@ -208,6 +208,8 @@ const AdminGenerationRecords: React.FC = () => {
const res = await getAdminGenerationRecords({ const res = await getAdminGenerationRecords({
userId: filterUserId.trim() || undefined, userId: filterUserId.trim() || undefined,
status: filterStatus || undefined, status: filterStatus || undefined,
engineId: filterEngineId || undefined,
includeMediaReferences: filterIncludeMedia === '' ? undefined : filterIncludeMedia === 'true',
page, page,
pageSize, pageSize,
}); });
@@ -240,6 +242,10 @@ const AdminGenerationRecords: React.FC = () => {
imageTokensUsed: item.imageTokensUsed || 0, imageTokensUsed: item.imageTokensUsed || 0,
imageProportion: item.imageProportion, imageProportion: item.imageProportion,
imagePx: item.imagePx, imagePx: item.imagePx,
engineId: item.engineId,
engineName: item.engineName,
engineSnapshot: item.engineSnapshot,
includeMediaReferences: item.includeMediaReferences,
}))); })));
setTotal(res.total || 0); setTotal(res.total || 0);
} catch { } catch {
@@ -247,12 +253,32 @@ const AdminGenerationRecords: React.FC = () => {
} finally { } finally {
setLoading(false); setLoading(false);
} }
}, [filterStatus, filterUserId, page, pageSize]); }, [filterStatus, filterUserId, filterEngineId, filterIncludeMedia, page, pageSize]);
useEffect(() => { useEffect(() => {
load(); load();
}, [load, reloadKey]); }, [load, reloadKey]);
useEffect(() => {
Promise.all([
getImageEngines({ includeDeleted: true }),
getVideoEngines({ includeDeleted: true }),
])
.then(([imageEngines, videoEngines]) => {
const items = [...(imageEngines || []), ...(videoEngines || [])];
const seen = new Set<string>();
setEngineOptions(items.reduce<Array<{ value: string; label: string }>>((acc, item: any) => {
const id = String(item?.id || '');
if (!id || seen.has(id)) return acc;
seen.add(id);
const deletedSuffix = item?.deletedAt ? '(已删除)' : '';
acc.push({ value: id, label: item?.name ? `${item.name}${deletedSuffix} (${id})` : `${id}${deletedSuffix}` });
return acc;
}, []));
})
.catch(() => setEngineOptions([]));
}, []);
useEffect(() => { useEffect(() => {
if (!preview) { if (!preview) {
setResourceState(EMPTY_RESOURCE_STATE); setResourceState(EMPTY_RESOURCE_STATE);
@@ -349,33 +375,7 @@ const AdminGenerationRecords: React.FC = () => {
}, 0); }, 0);
}; };
const handleStatusUpdate = async (recordId: string, newStatus: string, videoUrl?: string) => {
setUpdating(recordId);
try {
await adminUpdateGenerationStatus(recordId, newStatus, videoUrl);
message.success('状态已更新');
load();
} catch (e: any) {
message.error(e?.message || '更新失败');
} finally {
setUpdating(null);
}
};
const handleGenerate = async () => {
if (!genModal) return;
setUpdating(genModal.record.id);
try {
await adminGenerateVideo(genModal.record.id, genModal.ratio, genModal.resolution, genModal.image_size);
message.success(`已提交${genModal.record.genType === 'video' ? '视频' : '图片'}生成`);
setGenModal(null);
load();
} catch (e: any) {
message.error(e?.message || '生成失败');
} finally {
setUpdating(null);
}
};
const columns = useMemo(() => [ const columns = useMemo(() => [
{ {
@@ -416,21 +416,27 @@ const AdminGenerationRecords: React.FC = () => {
title: '参数', key: 'params', width: 160, title: '参数', key: 'params', width: 160,
render: (_: any, r: AdminGenerationRecord) => ( render: (_: any, r: AdminGenerationRecord) => (
r.genType === 'video' ? ( r.genType === 'video' ? (
r.duration || r.aspectRatio || r.resolution ? ( <Space size={4} wrap>
<Space size={4} wrap> {r.duration ? <Tag>{r.duration}s</Tag> : null}
{r.duration ? <Tag>{r.duration}s</Tag> : null} {r.aspectRatio ? <Tag>{r.aspectRatio}</Tag> : null}
{r.aspectRatio ? <Tag>{r.aspectRatio}</Tag> : null} {r.resolution ? <Tag>{r.resolution}</Tag> : null}
{r.resolution ? <Tag>{r.resolution}</Tag> : null} {!r.duration && !r.aspectRatio && !r.resolution ? <Tag color="default"></Tag> : null}
</Space> {r.engineName || r.engineId ? <Tag color="purple">{r.engineName || truncateId(r.engineId || '')}</Tag> : null}
) : <Tag color="default"></Tag> <Tag color={r.includeMediaReferences ? 'green' : 'default'}>
{r.includeMediaReferences ? '携带附件' : '不携带附件'}{r.references?.length ? `(${r.references.length})` : ''}
</Tag>
</Space>
) : ( ) : (
r.imageSize || r.imageProportion || r.imagePx ? ( <Space size={4} wrap>
<Space size={4} wrap> {r.imageSize ? <Tag>{r.imageSize}</Tag> : null}
{r.imageSize ? <Tag>{r.imageSize}</Tag> : null} {r.imageProportion ? <Tag>{r.imageProportion}</Tag> : null}
{r.imageProportion ? <Tag>{r.imageProportion}</Tag> : null} {r.imagePx ? <Tag>{r.imagePx}</Tag> : null}
{r.imagePx ? <Tag>{r.imagePx}</Tag> : null} {!r.imageSize && !r.imageProportion && !r.imagePx ? <Tag color="default"></Tag> : null}
</Space> {r.engineName || r.engineId ? <Tag color="purple">{r.engineName || truncateId(r.engineId || '')}</Tag> : null}
) : <Tag color="default"></Tag> <Tag color={r.includeMediaReferences ? 'green' : 'default'}>
{r.includeMediaReferences ? '携带附件' : '不携带附件'}{r.references?.length ? `(${r.references.length})` : ''}
</Tag>
</Space>
) )
), ),
}, },
@@ -467,72 +473,14 @@ const AdminGenerationRecords: React.FC = () => {
), ),
}, },
{ {
title: '操作', key: 'action', width: 150, fixed: 'right' as const, title: '操作', key: 'action', width: 90, fixed: 'right' as const,
render: (_: any, r: AdminGenerationRecord) => ( render: (_: any, r: AdminGenerationRecord) => (
<Space size={4} wrap> <Button size="small" icon={<EyeOutlined />} onClick={() => handleOpenPreview(r)}>
<Button size="small" icon={<EyeOutlined />} onClick={() => handleOpenPreview(r)}>
</Button>
</Button>
{r.status === 'generating' ? (
<Button
size="small"
danger
loading={updating === r.id}
onClick={() => {
Modal.confirm({
title: '确认操作',
icon: <ExclamationCircleOutlined />,
content: '确定将此记录标记为失败?',
onOk: () => handleStatusUpdate(r.id, 'failed'),
});
}}
>
</Button>
) : null}
{r.status === 'failed' ? (
<Button
size="small"
type="primary"
danger
loading={updating === r.id}
onClick={() => setGenModal({ record: r, ratio: r.aspectRatio || '16:9', resolution: r.resolution || '720p', image_size: r.imageSize || '2K' })}
>
</Button>
) : null}
{r.status === 'prompt_optimized' ? (
<>
<Button
size="small"
type="primary"
loading={updating === r.id}
onClick={() => setGenModal({ record: r, ratio: r.aspectRatio || '16:9', resolution: r.resolution || '720p', image_size: r.imageSize || '2K' })}
style={{ background: '#6366f1', border: 'none' }}
>
{r.genType === 'video' ? '视频' : '图片'}
</Button>
<Button
size="small"
danger
loading={updating === r.id}
onClick={() => {
Modal.confirm({
title: '确认操作',
icon: <ExclamationCircleOutlined />,
content: '确定将此记录标记为失败?',
onOk: () => handleStatusUpdate(r.id, 'failed'),
});
}}
>
</Button>
</>
) : null}
</Space>
), ),
}, },
], [handleOpenPreview, updating]); ], [handleOpenPreview]);
const previewTypeConfig = preview ? (GEN_TYPE_MAP[preview.genType || ''] || { text: preview.genType || '-', color: 'default', icon: null }) : null; const previewTypeConfig = preview ? (GEN_TYPE_MAP[preview.genType || ''] || { text: preview.genType || '-', color: 'default', icon: null }) : null;
const previewStatusConfig = preview ? (STATUS_MAP[preview.status] || { color: 'default', text: preview.status || '-', icon: null }) : null; const previewStatusConfig = preview ? (STATUS_MAP[preview.status] || { color: 'default', text: preview.status || '-', icon: null }) : null;
@@ -894,6 +842,27 @@ const AdminGenerationRecords: React.FC = () => {
{ value: 'failed', label: '失败' }, { value: 'failed', label: '失败' },
]} ]}
/> />
<Select
placeholder="引擎筛选"
allowClear
showSearch
optionFilterProp="label"
style={{ width: 220 }}
value={filterEngineId || undefined}
onChange={(v) => { setFilterEngineId(v || ''); setPage(1); }}
options={engineOptions}
/>
<Select
placeholder="附件状态"
allowClear
style={{ width: 130 }}
value={filterIncludeMedia || undefined}
onChange={(v) => { setFilterIncludeMedia((v || '') as '' | 'true' | 'false'); setPage(1); }}
options={[
{ value: 'true', label: '携带附件' },
{ value: 'false', label: '不携带附件' },
]}
/>
<Input <Input
placeholder="用户ID搜索" placeholder="用户ID搜索"
prefix={<SearchOutlined style={{ color: '#94a3b8' }} />} prefix={<SearchOutlined style={{ color: '#94a3b8' }} />}
@@ -1062,73 +1031,6 @@ const AdminGenerationRecords: React.FC = () => {
)} )}
</Modal> </Modal>
{/* Generate modal */}
<Modal
title={(
<Space>
{genModal && (genModal.record.genType === 'video' ? <PlayCircleOutlined /> : <FileImageOutlined />)}
{genModal && (genModal.record.genType === 'video' ? '生成视频' : '生成图片')}
</Space>
)}
open={!!genModal}
onCancel={() => setGenModal(null)}
onOk={handleGenerate}
okText="提交生成"
cancelText="取消"
confirmLoading={genModal ? updating === genModal.record.id : false}
width={420}
>
{genModal ? (
<div style={{ display: 'flex', flexDirection: 'column', gap: 16, marginTop: 16 }}>
{genModal.record.genType === 'video' ? (
<>
<div style={{ padding: 12, borderRadius: 10, background: '#f8f9fc' }}>
<Typography.Text style={{ fontSize: 11, color: '#94a3b8', display: 'block' }}></Typography.Text>
<Typography.Text strong>{genModal.record.duration || 5}s</Typography.Text>
</div>
<div>
<Typography.Text style={{ fontSize: 12, color: '#64748b', display: 'block', marginBottom: 6 }}></Typography.Text>
<Select
value={genModal.ratio}
onChange={(v) => setGenModal((prev) => (prev ? { ...prev, ratio: v } : null))}
style={{ width: '100%' }}
options={['16:9', '4:3', '1:1', '3:4', '9:16', '21:9'].map((r) => ({ value: r, label: r }))}
/>
</div>
<div>
<Typography.Text style={{ fontSize: 12, color: '#64748b', display: 'block', marginBottom: 6 }}></Typography.Text>
<Select
value={genModal.resolution}
onChange={(v) => setGenModal((prev) => (prev ? { ...prev, resolution: v } : null))}
style={{ width: '100%' }}
options={['480p', '720p', '1080p'].map((r) => ({ value: r, label: r }))}
/>
</div>
</>
) : (
<>
<div style={{ padding: 12, borderRadius: 10, background: '#f8f9fc' }}>
<Typography.Text style={{ fontSize: 11, color: '#94a3b8', display: 'block' }}></Typography.Text>
<Typography.Text strong>{genModal.record.imagePx || '-'}</Typography.Text>
</div>
<div style={{ padding: 12, borderRadius: 10, background: '#f8f9fc' }}>
<Typography.Text style={{ fontSize: 11, color: '#94a3b8', display: 'block' }}></Typography.Text>
<Typography.Text strong>{genModal.record.imageProportion || '-'}</Typography.Text>
</div>
<div>
<Typography.Text style={{ fontSize: 12, color: '#64748b', display: 'block', marginBottom: 6 }}></Typography.Text>
<Select
value={genModal.image_size}
onChange={(v) => setGenModal((prev) => (prev ? { ...prev, image_size: v } : null))}
style={{ width: '100%' }}
options={['2K', '4K'].map((r) => ({ value: r, label: r }))}
/>
</div>
</>
)}
</div>
) : null}
</Modal>
</div> </div>
); );
}; };
@@ -207,6 +207,7 @@ const AdminHotOpeningReplications: React.FC = () => {
), ),
}, },
{ title: '状态', dataIndex: 'status', width: 130, render: (v: string) => <StatusTag status={v} /> }, { title: '状态', dataIndex: 'status', width: 130, render: (v: string) => <StatusTag status={v} /> },
{ title: '流程版本', dataIndex: 'flowVersion', width: 100, render: (v: string) => <Tag color={v === 'v2' ? 'blue' : 'default'}>{String(v || 'v1').toUpperCase()}</Tag> },
{ title: '当前步骤', dataIndex: 'currentStepCode', width: 140, render: (v: string) => STEP_MAP[v] || v || '-' }, { title: '当前步骤', dataIndex: 'currentStepCode', width: 140, render: (v: string) => STEP_MAP[v] || v || '-' },
{ title: '图片结果', dataIndex: 'finalImageUrl', width: 90, render: (v: string) => v ? <Tag color="success"></Tag> : <Tag></Tag> }, { title: '图片结果', dataIndex: 'finalImageUrl', width: 90, render: (v: string) => v ? <Tag color="success"></Tag> : <Tag></Tag> },
{ title: '视频结果', dataIndex: 'finalVideoUrl', width: 90, render: (v: string) => v ? <Tag color="success"></Tag> : <Tag></Tag> }, { title: '视频结果', dataIndex: 'finalVideoUrl', width: 90, render: (v: string) => v ? <Tag color="success"></Tag> : <Tag></Tag> },
@@ -218,7 +219,7 @@ const AdminHotOpeningReplications: React.FC = () => {
fixed: 'right', fixed: 'right',
width: 110, width: 110,
render: (_, record) => ( render: (_, record) => (
<Button type="link" icon={<EyeOutlined />} onClick={() => navigate(`/hot-opening-replications/${record.id}`)}></Button> <Button type="link" icon={<EyeOutlined />} onClick={() => navigate(`/hot-opening-replications/${record.id}?flow_version=${record.flowVersion === 'v2' ? 'v2' : 'v1'}`)}></Button>
), ),
}, },
]} ]}
@@ -21,7 +21,7 @@ import {
PlayCircleOutlined, PlayCircleOutlined,
VideoCameraOutlined, VideoCameraOutlined,
} from '@ant-design/icons'; } from '@ant-design/icons';
import { useNavigate, useParams } from 'react-router-dom'; import { useNavigate, useParams, useSearchParams } from 'react-router-dom';
import { getAdminHotOpeningTaskDetail, getAdminShotProjectDetail } from '../api'; import { getAdminHotOpeningTaskDetail, getAdminShotProjectDetail } from '../api';
import type { ReplicationProjectDetailOut, ReplicationStepOut } from '../types'; import type { ReplicationProjectDetailOut, ReplicationStepOut } from '../types';
import { formatDate } from '../utils/formatDate'; import { formatDate } from '../utils/formatDate';
@@ -36,7 +36,7 @@ interface AdminReplicationProjectDetailProps {
moduleType?: ReplicationModuleType; moduleType?: ReplicationModuleType;
} }
const STEP_ORDER = [ const V1_STEP_ORDER = [
'material_input', 'material_input',
'image_prompt_optimize', 'image_prompt_optimize',
'image_generate', 'image_generate',
@@ -44,6 +44,12 @@ const STEP_ORDER = [
'video_generate', 'video_generate',
]; ];
const V2_STEP_ORDER = [
'material_input',
'video_prompt_optimize',
'video_generate',
];
const STEP_DESCRIPTIONS: Record<string, string> = { const STEP_DESCRIPTIONS: Record<string, string> = {
material_input: '参考素材、项目名称和核心内容点', material_input: '参考素材、项目名称和核心内容点',
image_prompt_optimize: '图片 AI 提词优化结果', image_prompt_optimize: '图片 AI 提词优化结果',
@@ -104,6 +110,8 @@ const renderPromptText = (value?: string | null, empty = '暂无提词') => {
const AdminReplicationProjectDetail: React.FC<AdminReplicationProjectDetailProps> = ({ moduleType = 'shot_replicate' }) => { const AdminReplicationProjectDetail: React.FC<AdminReplicationProjectDetailProps> = ({ moduleType = 'shot_replicate' }) => {
const { projectId } = useParams<{ projectId: string }>(); const { projectId } = useParams<{ projectId: string }>();
const navigate = useNavigate(); const navigate = useNavigate();
const [searchParams] = useSearchParams();
const flowVersion: 'v1' | 'v2' = searchParams.get('flow_version') === 'v2' ? 'v2' : 'v1';
const [detail, setDetail] = useState<ReplicationProjectDetailOut | null>(null); const [detail, setDetail] = useState<ReplicationProjectDetailOut | null>(null);
const [loading, setLoading] = useState(false); const [loading, setLoading] = useState(false);
@@ -112,15 +120,15 @@ const AdminReplicationProjectDetail: React.FC<AdminReplicationProjectDetailProps
setLoading(true); setLoading(true);
try { try {
const res = moduleType === 'hot_opening_replicate' const res = moduleType === 'hot_opening_replicate'
? await getAdminHotOpeningTaskDetail(projectId) ? await getAdminHotOpeningTaskDetail(projectId, flowVersion)
: await getAdminShotProjectDetail(projectId); : await getAdminShotProjectDetail(projectId, flowVersion);
setDetail(res); setDetail(res);
} catch (e: any) { } catch (e: any) {
message.error(e?.message || '加载复刻项目详情失败'); message.error(e?.message || '加载复刻项目详情失败');
} finally { } finally {
setLoading(false); setLoading(false);
} }
}, [moduleType, projectId]); }, [flowVersion, moduleType, projectId]);
useEffect(() => { useEffect(() => {
load(); load();
@@ -134,8 +142,10 @@ const AdminReplicationProjectDetail: React.FC<AdminReplicationProjectDetailProps
const moduleValue = detail?.module || moduleType; const moduleValue = detail?.module || moduleType;
const moduleName = getModuleLabel(moduleValue); const moduleName = getModuleLabel(moduleValue);
const isV2 = detail?.flowVersion === 'v2';
const effectiveStepOrder = isV2 ? V2_STEP_ORDER : V1_STEP_ORDER;
const stepItems = useMemo(() => STEP_ORDER.map(code => { const stepItems = useMemo(() => effectiveStepOrder.map(code => {
const step = stepsByCode[code]; const step = stepsByCode[code];
let status: 'wait' | 'process' | 'finish' | 'error' = 'wait'; let status: 'wait' | 'process' | 'finish' | 'error' = 'wait';
if (step?.status === 'completed') status = 'finish'; if (step?.status === 'completed') status = 'finish';
@@ -146,7 +156,7 @@ const AdminReplicationProjectDetail: React.FC<AdminReplicationProjectDetailProps
description: step ? <StatusTag status={step.status} /> : '未创建', description: step ? <StatusTag status={step.status} /> : '未创建',
status, status,
}; };
}), [stepsByCode]); }), [effectiveStepOrder, stepsByCode]);
const defaultActiveKeys = useMemo(() => buildDefaultActiveKeys(detail, stepsByCode), [detail, stepsByCode]); const defaultActiveKeys = useMemo(() => buildDefaultActiveKeys(detail, stepsByCode), [detail, stepsByCode]);
@@ -242,7 +252,7 @@ const AdminReplicationProjectDetail: React.FC<AdminReplicationProjectDetailProps
}, },
{ {
key: 'video_prompt_optimize', key: 'video_prompt_optimize',
label: <StepHeader index={4} stepCode="video_prompt_optimize" step={videoPromptStep} current={detail.currentStepCode === 'video_prompt_optimize'} />, label: <StepHeader index={isV2 ? 2 : 4} stepCode="video_prompt_optimize" step={videoPromptStep} current={detail.currentStepCode === 'video_prompt_optimize'} />,
children: ( children: (
<Space direction="vertical" size={16} style={{ width: '100%' }}> <Space direction="vertical" size={16} style={{ width: '100%' }}>
<Descriptions column={3} bordered size="small"> <Descriptions column={3} bordered size="small">
@@ -265,7 +275,7 @@ const AdminReplicationProjectDetail: React.FC<AdminReplicationProjectDetailProps
}, },
{ {
key: 'video_generate', key: 'video_generate',
label: <StepHeader index={5} stepCode="video_generate" step={videoGenerateStep} current={detail.currentStepCode === 'video_generate'} />, label: <StepHeader index={isV2 ? 3 : 5} stepCode="video_generate" step={videoGenerateStep} current={detail.currentStepCode === 'video_generate'} />,
children: ( children: (
<Space direction="vertical" size={16} style={{ width: '100%' }}> <Space direction="vertical" size={16} style={{ width: '100%' }}>
<Descriptions column={3} bordered size="small"> <Descriptions column={3} bordered size="small">
@@ -300,6 +310,9 @@ const AdminReplicationProjectDetail: React.FC<AdminReplicationProjectDetailProps
), ),
}, },
]; ];
const visibleCollapseItems = isV2
? collapseItems.filter(item => !['image_prompt_optimize', 'image_generate'].includes(String(item.key)))
: collapseItems;
return ( return (
<div style={{ padding: 24 }}> <div style={{ padding: 24 }}>
@@ -322,6 +335,8 @@ const AdminReplicationProjectDetail: React.FC<AdminReplicationProjectDetailProps
<Descriptions.Item label="用户名">{detail.userName || '-'}</Descriptions.Item> <Descriptions.Item label="用户名">{detail.userName || '-'}</Descriptions.Item>
<Descriptions.Item label="标题">{detail.title || '-'}</Descriptions.Item> <Descriptions.Item label="标题">{detail.title || '-'}</Descriptions.Item>
<Descriptions.Item label="当前步骤"><Tooltip title={detail.currentStepCode || ''}>{getStepCodeLabel(detail.currentStepCode)}</Tooltip></Descriptions.Item> <Descriptions.Item label="当前步骤"><Tooltip title={detail.currentStepCode || ''}>{getStepCodeLabel(detail.currentStepCode)}</Tooltip></Descriptions.Item>
<Descriptions.Item label="流程版本"><Tag color={isV2 ? 'blue' : 'default'}>{String(detail.flowVersion || 'v1').toUpperCase()}</Tag></Descriptions.Item>
<Descriptions.Item label="步骤数量">{detail.stepCount || effectiveStepOrder.length}</Descriptions.Item>
<Descriptions.Item label="状态"><StatusTag status={detail.status} /></Descriptions.Item> <Descriptions.Item label="状态"><StatusTag status={detail.status} /></Descriptions.Item>
<Descriptions.Item label="创建时间">{safeDate(detail.createdAt)}</Descriptions.Item> <Descriptions.Item label="创建时间">{safeDate(detail.createdAt)}</Descriptions.Item>
<Descriptions.Item label="更新时间">{safeDate(detail.updatedAt)}</Descriptions.Item> <Descriptions.Item label="更新时间">{safeDate(detail.updatedAt)}</Descriptions.Item>
@@ -336,9 +351,11 @@ const AdminReplicationProjectDetail: React.FC<AdminReplicationProjectDetailProps
<Card title="最终结果预览"> <Card title="最终结果预览">
<div style={{ display: 'grid', gridTemplateColumns: 'repeat(auto-fit, minmax(280px, 1fr))', gap: 16 }}> <div style={{ display: 'grid', gridTemplateColumns: 'repeat(auto-fit, minmax(280px, 1fr))', gap: 16 }}>
<Card size="small" title="最终图片"> {!isV2 ? (
<MediaPreview type="image" url={detail.finalImageUrl || detail.imageGeneration?.resultImageUrl} height={220} emptyDescription="暂无最终图片" /> <Card size="small" title="最终图片">
</Card> <MediaPreview type="image" url={detail.finalImageUrl || detail.imageGeneration?.resultImageUrl} height={220} emptyDescription="暂无最终图片" />
</Card>
) : null}
<Card size="small" title="最终视频封面"> <Card size="small" title="最终视频封面">
<MediaPreview type="image" url={detail.finalVideoCoverUrl || detail.videoGeneration?.resultVideoCoverUrl} height={220} emptyDescription="暂无最终视频封面" /> <MediaPreview type="image" url={detail.finalVideoCoverUrl || detail.videoGeneration?.resultVideoCoverUrl} height={220} emptyDescription="暂无最终视频封面" />
</Card> </Card>
@@ -348,7 +365,7 @@ const AdminReplicationProjectDetail: React.FC<AdminReplicationProjectDetailProps
</div> </div>
</Card> </Card>
<Collapse defaultActiveKey={defaultActiveKeys} items={collapseItems} /> <Collapse defaultActiveKey={defaultActiveKeys} items={visibleCollapseItems} />
<Collapse <Collapse
items={[ items={[
@@ -284,9 +284,14 @@ const AdminShotTaskSetDetail: React.FC = () => {
width: 220, width: 220,
render: (_, record) => record.moduleProjectId ? ( render: (_, record) => record.moduleProjectId ? (
<Space direction="vertical" size={0}> <Space direction="vertical" size={0}>
<Button type="link" style={{ padding: 0 }} onClick={() => navigate(`/shot-replications/projects/${record.moduleProjectId}`)}>{shortId(record.moduleProjectId)}</Button> <Button type="link" style={{ padding: 0 }} onClick={() => navigate(`/shot-replications/projects/${record.moduleProjectId}?flow_version=${record.moduleProjectFlowVersion === 'v2' ? 'v2' : 'v1'}`)}>{shortId(record.moduleProjectId)}</Button>
<Typography.Text type="secondary" style={{ fontSize: 12 }}>{record.moduleProjectTitle || getStepCodeLabel(record.moduleProjectCurrentStepCode)}</Typography.Text> <Typography.Text type="secondary" style={{ fontSize: 12 }}>{record.moduleProjectTitle || getStepCodeLabel(record.moduleProjectCurrentStepCode)}</Typography.Text>
<StatusTag status={record.moduleProjectStatus} /> <Space size={4}>
<StatusTag status={record.moduleProjectStatus} />
<Tag color={record.moduleProjectFlowVersion === 'v2' ? 'blue' : 'default'}>
{String(record.moduleProjectFlowVersion || 'v1').toUpperCase()}
</Tag>
</Space>
</Space> </Space>
) : <Tag></Tag>, ) : <Tag></Tag>,
}, },
@@ -313,7 +318,8 @@ const AdminShotTaskSetDetail: React.FC = () => {
<Descriptions.Item label="切割状态"><StatusTag status={segmentDetail.splitStatus} /></Descriptions.Item> <Descriptions.Item label="切割状态"><StatusTag status={segmentDetail.splitStatus} /></Descriptions.Item>
<Descriptions.Item label="分析状态"><StatusTag status={segmentDetail.analysisStatus} /></Descriptions.Item> <Descriptions.Item label="分析状态"><StatusTag status={segmentDetail.analysisStatus} /></Descriptions.Item>
<Descriptions.Item label="复刻状态"><StatusTag status={segmentDetail.replicateStatus} /></Descriptions.Item> <Descriptions.Item label="复刻状态"><StatusTag status={segmentDetail.replicateStatus} /></Descriptions.Item>
<Descriptions.Item label="关联项目">{segmentDetail.moduleProjectId ? <Button type="link" onClick={() => navigate(`/shot-replications/projects/${segmentDetail.moduleProjectId}`)}>{segmentDetail.moduleProjectId}</Button> : '-'}</Descriptions.Item> <Descriptions.Item label="关联项目">{segmentDetail.moduleProjectId ? <Button type="link" onClick={() => navigate(`/shot-replications/projects/${segmentDetail.moduleProjectId}?flow_version=${segmentDetail.moduleProjectFlowVersion === 'v2' ? 'v2' : 'v1'}`)}>{segmentDetail.moduleProjectId}</Button> : '-'}</Descriptions.Item>
<Descriptions.Item label="项目流程版本">{segmentDetail.moduleProjectId ? <Tag color={segmentDetail.moduleProjectFlowVersion === 'v2' ? 'blue' : 'default'}>{String(segmentDetail.moduleProjectFlowVersion || 'v1').toUpperCase()}</Tag> : '-'}</Descriptions.Item>
<Descriptions.Item label="片段内容" span={2}>{segmentDetail.segmentContent || '-'}</Descriptions.Item> <Descriptions.Item label="片段内容" span={2}>{segmentDetail.segmentContent || '-'}</Descriptions.Item>
<Descriptions.Item label="片段分类">{segmentDetail.segmentCategory || '-'}</Descriptions.Item> <Descriptions.Item label="片段分类">{segmentDetail.segmentCategory || '-'}</Descriptions.Item>
<Descriptions.Item label="片段受众">{segmentDetail.segmentAudience || '-'}</Descriptions.Item> <Descriptions.Item label="片段受众">{segmentDetail.segmentAudience || '-'}</Descriptions.Item>
+10
View File
@@ -394,6 +394,10 @@ export interface AdminGenerationRecord {
imageTokensUsed?: number; imageTokensUsed?: number;
imageProportion?: string; imageProportion?: string;
imagePx?: string; imagePx?: string;
engineId?: string;
engineName?: string;
engineSnapshot?: Record<string, unknown> | null;
includeMediaReferences?: boolean;
} }
export type GenerationAITaskStatus = 'pending' | 'generating' | 'completed' | 'failed' | string; export type GenerationAITaskStatus = 'pending' | 'generating' | 'completed' | 'failed' | string;
@@ -572,6 +576,9 @@ export interface ReplicationProjectDetailOut {
title?: string | null; title?: string | null;
status: ModuleReplicationStatus; status: ModuleReplicationStatus;
currentStepCode?: string | null; currentStepCode?: string | null;
flowVersion?: 'v1' | 'v2' | string | null;
stepCount?: number;
stepIoSchemaVersion?: string | null;
finalImageUrl?: string | null; finalImageUrl?: string | null;
finalVideoUrl?: string | null; finalVideoUrl?: string | null;
finalVideoCoverUrl?: string | null; finalVideoCoverUrl?: string | null;
@@ -594,6 +601,8 @@ export interface HotOpeningTaskListItemOut {
title?: string | null; title?: string | null;
status: ModuleReplicationStatus; status: ModuleReplicationStatus;
currentStepCode?: string | null; currentStepCode?: string | null;
flowVersion?: 'v1' | 'v2' | string | null;
stepCount?: number;
sourceProjectName?: string | null; sourceProjectName?: string | null;
targetProjectName?: string | null; targetProjectName?: string | null;
coreContentPoint?: string | null; coreContentPoint?: string | null;
@@ -707,6 +716,7 @@ export interface ShotSegmentOut {
moduleProjectTitle?: string | null; moduleProjectTitle?: string | null;
moduleProjectStatus?: string | null; moduleProjectStatus?: string | null;
moduleProjectCurrentStepCode?: string | null; moduleProjectCurrentStepCode?: string | null;
moduleProjectFlowVersion?: 'v1' | 'v2' | string | null;
createdAt?: string | null; createdAt?: string | null;
updatedAt?: string | null; updatedAt?: string | null;
} }
@@ -0,0 +1,87 @@
"""add module flow version and generation reference option
Revision ID: d8ebe79ab575
Revises: e6eac828ff61
Create Date: 2026-07-21 09:28:04.318458
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
revision: str = "d8ebe79ab575"
down_revision: Union[str, None] = "e6eac828ff61"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def _column_names(table_name: str) -> set[str]:
inspector = sa.inspect(op.get_bind())
return {str(column["name"]) for column in inspector.get_columns(table_name)}
def upgrade() -> None:
generation_record_columns = _column_names("generation_records")
if "include_media_references" not in generation_record_columns:
op.add_column(
"generation_records",
sa.Column(
"include_media_references",
sa.Boolean(),
server_default=sa.text("false"),
nullable=False,
),
)
else:
op.execute(
sa.text(
"UPDATE generation_records "
"SET include_media_references = false "
"WHERE include_media_references IS NULL"
)
)
op.alter_column(
"generation_records",
"include_media_references",
existing_type=sa.Boolean(),
nullable=False,
server_default=sa.text("false"),
)
project_columns = _column_names("module_generation_projects")
if "flow_version" not in project_columns:
op.add_column(
"module_generation_projects",
sa.Column(
"flow_version",
sa.String(length=16),
server_default=sa.text("'v1'"),
nullable=False,
),
)
else:
op.execute(
sa.text(
"UPDATE module_generation_projects "
"SET flow_version = 'v1' "
"WHERE flow_version IS NULL OR btrim(flow_version) = ''"
)
)
op.alter_column(
"module_generation_projects",
"flow_version",
existing_type=sa.String(length=16),
nullable=False,
server_default=sa.text("'v1'"),
)
def downgrade() -> None:
project_columns = _column_names("module_generation_projects")
if "flow_version" in project_columns:
op.drop_column("module_generation_projects", "flow_version")
generation_record_columns = _column_names("generation_records")
if "include_media_references" in generation_record_columns:
op.drop_column("generation_records", "include_media_references")
+32 -258
View File
@@ -1,7 +1,7 @@
from datetime import datetime, timezone, timedelta from datetime import datetime, timezone, timedelta
import json import json
from fastapi import APIRouter, Depends, HTTPException, Query, status from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy import delete, func, select, update from sqlalchemy import delete, func, select, update
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
@@ -24,7 +24,6 @@ from app.models.credit_ratio import CreditRatio
from app.models.operation_log import OperationLog from app.models.operation_log import OperationLog
from app.enums.user import FrontendUserKind, UserType from app.enums.user import FrontendUserKind, UserType
from app.enums.team import TEAM_UNASSIGNED_VALUE from app.enums.team import TEAM_UNASSIGNED_VALUE
from app.enums.generation_status import GenerationRecordPipelineStage
from app.schemas.admin import ( from app.schemas.admin import (
CreditAdjustRequest, CreditAdjustRequest,
ModelConfigCreate, ModelConfigCreate,
@@ -41,17 +40,12 @@ from app.schemas.admin import (
UpdateMenusRequest, UpdateMenusRequest,
ResetPasswordRequest, ResetPasswordRequest,
UpdateFrontendUserKindRequest, UpdateFrontendUserKindRequest,
OperationLogOut,
) )
from app.schemas.team import UpdateUserTeamRequest from app.schemas.team import UpdateUserTeamRequest
from app.schemas.industry import IndustryConfigCreate, IndustryConfigOut from app.schemas.industry import IndustryConfigCreate
from app.schemas.video_engine import VideoEngineCreate, VideoEngineOut from app.schemas.video_engine import VideoEngineCreate, VideoEngineOut
from app.schemas.image_engine import ImageEngineCreate, ImageEngineOut from app.schemas.image_engine import ImageEngineCreate, ImageEngineOut
from app.schemas.credit_ratio import CreditRatioCreate, CreditRatioOut from app.schemas.credit_ratio import CreditRatioCreate, CreditRatioOut
from app.services.generation.pipeline.db_lock_service import (
DatabaseRowLockBusy,
execute_with_lock_timeout,
)
from app.services.credits import add_credits, deduct_credits from app.services.credits import add_credits, deduct_credits
from app.services.credit_record_meta_service import build_admin_adjust_meta from app.services.credit_record_meta_service import build_admin_adjust_meta
from app.services.admin_credit_record_service import list_admin_credit_records from app.services.admin_credit_record_service import list_admin_credit_records
@@ -60,23 +54,26 @@ from app.services.auth import hash_password, verify_password
from app.services.operation_log import log_operation from app.services.operation_log import log_operation
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls
from app.services.resource_signed_url_service import build_resource_signed_url from app.services.resource_signed_url_service import build_resource_signed_url
from app.services.payment import sync_pending_orders, process_refund from app.services.payment import process_refund
from app.services.resource_capacity_service import batch_get_user_resource_capacity_usage, get_user_resource_capacity_usage from app.services.resource_capacity_service import batch_get_user_resource_capacity_usage, get_user_resource_capacity_usage
from app.services.team_service import batch_get_team_name_map, set_frontend_user_team from app.services.team_service import batch_get_team_name_map, set_frontend_user_team
from app.services.generation.billing_service import (
OWNER_GENERATION_RECORD,
charge_generation_media_for_record,
get_next_credit_attempt_no,
)
from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once
from app.utils.id_gen import generate_id from app.utils.id_gen import generate_id
from app.schemas.generation import GenerationType, ASPECT_RATIOS, RESOLUTIONS
CST = timezone(timedelta(hours=8)) CST = timezone(timedelta(hours=8))
def _safe_json_object(value: str | None) -> dict | None:
if not value:
return None
try:
parsed = json.loads(value)
except (TypeError, json.JSONDecodeError):
return None
return parsed if isinstance(parsed, dict) else None
def _iso(dt): def _iso(dt):
"""Serialize datetime as naive ISO string (UTC→CST, strip tzinfo).""" """Serialize datetime as naive ISO string (UTC→CST, strip tzinfo)."""
if dt is None: if dt is None:
@@ -2021,6 +2018,8 @@ async def list_token_usage(
async def admin_list_generation_records( async def admin_list_generation_records(
user_id: str | None = Query(None), user_id: str | None = Query(None),
status: str | None = Query(None), status: str | None = Query(None),
engine_id: str | None = Query(None),
include_media_references: bool | None = Query(None),
page: int = Query(1, ge=1), page: int = Query(1, ge=1),
page_size: int = Query(20, ge=1, le=500), page_size: int = Query(20, ge=1, le=500),
admin: User = Depends(get_admin_user), admin: User = Depends(get_admin_user),
@@ -2039,13 +2038,25 @@ async def admin_list_generation_records(
query = query.where(GenerationRecord.user_id == user_id) query = query.where(GenerationRecord.user_id == user_id)
if status: if status:
query = query.where(GenerationRecord.status == status) query = query.where(GenerationRecord.status == status)
if engine_id:
query = query.where(GenerationRecord.engine_id == engine_id)
if include_media_references is not None:
query = query.where(GenerationRecord.include_media_references.is_(include_media_references))
# Count total # Count total
count_query = select(func.count(GenerationRecord.id)).where(GenerationRecord.deleted_at.is_(None)) count_query = (
select(func.count(GenerationRecord.id))
.join(Project, GenerationRecord.project_id == Project.id)
.where(GenerationRecord.deleted_at.is_(None), Project.deleted_at.is_(None))
)
if user_id: if user_id:
count_query = count_query.where(GenerationRecord.user_id == user_id) count_query = count_query.where(GenerationRecord.user_id == user_id)
if status: if status:
count_query = count_query.where(GenerationRecord.status == status) count_query = count_query.where(GenerationRecord.status == status)
if engine_id:
count_query = count_query.where(GenerationRecord.engine_id == engine_id)
if include_media_references is not None:
count_query = count_query.where(GenerationRecord.include_media_references.is_(include_media_references))
total_result = await db.execute(count_query) total_result = await db.execute(count_query)
total = total_result.scalar() or 0 total = total_result.scalar() or 0
@@ -2081,6 +2092,10 @@ async def admin_list_generation_records(
"video_url": build_resource_signed_url(record.video_url) if record.video_url else '', "video_url": build_resource_signed_url(record.video_url) if record.video_url else '',
"video_cover_url": build_resource_signed_url(record.video_cover_url) if record.video_cover_url else '', "video_cover_url": build_resource_signed_url(record.video_cover_url) if record.video_cover_url else '',
"references": refs, "references": refs,
"engine_id": record.engine_id,
"engine_name": (_safe_json_object(record.engine_snapshot_json) or {}).get("name"),
"engine_snapshot": _safe_json_object(record.engine_snapshot_json),
"include_media_references": bool(record.include_media_references),
"credits_cost": record.credits_cost or 0, "credits_cost": record.credits_cost or 0,
"text_credits_cost": record.text_credits_cost or 0, "text_credits_cost": record.text_credits_cost or 0,
"text_tokens_used": record.text_tokens_used or 0, "text_tokens_used": record.text_tokens_used or 0,
@@ -2101,245 +2116,6 @@ async def admin_list_generation_records(
return {"total": total, "items": items} return {"total": total, "items": items}
@router.put("/generation-records/{record_id}/status")
async def admin_update_generation_status(
record_id: str,
body: dict,
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""管理员只能终止正在执行或待生成的记录,禁止绕过流水线裸改生成/完成状态。"""
try:
result = await execute_with_lock_timeout(
db,
select(GenerationRecord).where(
GenerationRecord.id == record_id,
GenerationRecord.deleted_at.is_(None),
)
.with_for_update()
.limit(1),
)
except DatabaseRowLockBusy as exc:
raise HTTPException(status_code=409, detail=exc.detail) from exc
record = result.scalar_one_or_none()
if not record:
raise HTTPException(status_code=404, detail="记录不存在")
new_status = str(body.get("status") or "").strip()
if new_status in {"generating", "completed", "prompt_optimized"}:
raise HTTPException(
status_code=409,
detail="禁止直接修改为该状态;生成请调用生成接口,完成必须由下载/超分流水线落库",
)
if new_status != "failed":
raise HTTPException(status_code=400, detail="该接口仅允许管理员终止任务")
if record.status == "completed":
raise HTTPException(status_code=409, detail="已完成记录不能直接改为失败")
error_message = body.get("error_message") or record.error_message or "管理员终止生成任务"
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=error_message,
generation_attempt_no=int(record.generation_attempt_no or 1),
)
record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
record.provider_create_claim_token = None
record.provider_create_lease_until = None
record.poll_claim_token = None
record.poll_lease_until = None
record.next_poll_at = None
record.download_claim_token = None
record.download_lease_until = None
record.download_next_retry_at = None
# 若任务已进入超分,必须同时撤销超分数据库租约;执行中的超分 Worker
# 在回填前校验 lease_token,发现 token 被清除后会中止,不得覆盖管理员终止状态。
from app.enums.video_upscale import VideoUpscaleStage, VideoUpscaleTaskStatus
from app.models.video_upscale_task import VideoUpscaleTask
try:
upscale_result = await execute_with_lock_timeout(
db,
select(VideoUpscaleTask)
.where(VideoUpscaleTask.generation_record_id == record.id)
.with_for_update()
.limit(1),
)
except DatabaseRowLockBusy as exc:
raise HTTPException(status_code=409, detail=exc.detail) from exc
upscale = upscale_result.scalar_one_or_none()
if upscale and upscale.status not in {
VideoUpscaleTaskStatus.COMPLETED.value,
VideoUpscaleTaskStatus.FAILED.value,
}:
upscale.status = VideoUpscaleTaskStatus.FAILED.value
upscale.stage = VideoUpscaleStage.FAILED.value
upscale.last_error = error_message
upscale.failed_at = datetime.now(CST)
upscale.next_retry_at = None
upscale.lease_token = None
upscale.lease_until = None
await db.flush()
await log_operation(
db,
admin.id,
admin.username,
"管理员终止生成记录",
"PUT",
f"/admin/generation-records/{record_id}/status",
detail=json.dumps(
{
"record_id": record_id,
"new_status": new_status,
"generation_attempt_no": int(record.generation_attempt_no or 1),
"error_message": error_message,
},
ensure_ascii=False,
),
)
await db.commit()
# Redis 注册表只做调度加速;删除失败不回滚已提交的业务终止状态。
try:
from app.services.celery_download_recovery_service import remove_download_active
from app.services.generation.pipeline.owner_service import redis_owner_item_id
from app.services.redis_registry_service import redis_remove_registry_item
from app.config import settings
registry_id = redis_owner_item_id(
"generation_record",
record_id,
int(record.generation_attempt_no or 1),
)
await remove_download_active(registry_id)
await redis_remove_registry_item(
hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY,
zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY,
item_id=registry_id,
log_context="admin_generation_record_terminate",
)
except Exception:
pass
return {"message": "ok"}
@router.post("/generation-records/{record_id}/generate")
async def admin_generate_record_resource(
record_id: str,
body: dict,
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""管理员触发 GenerationRecord 图片或视频资源生成。"""
from app.services.generation.pipeline.generation_record_service import (
commit_and_enqueue_generation_record,
prepare_generation_record_execution,
)
try:
result = await execute_with_lock_timeout(
db,
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(
GenerationRecord.id == record_id,
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
)
.with_for_update(),
)
except DatabaseRowLockBusy as exc:
raise HTTPException(status_code=409, detail=exc.detail) from exc
row = result.first()
if not row:
raise HTTPException(status_code=404, detail="记录不存在")
record, project_name = row
type_str = "视频" if record.gen_type == GenerationType.video else "图片"
if record.status not in ("prompt_optimized", "failed"):
raise HTTPException(status_code=400, detail=f"当前状态不允许生成{type_str}")
if record.pipeline_stage == GenerationRecordPipelineStage.UPSCALE_FAILED.value:
raise HTTPException(status_code=409, detail="该任务为画质增强失败,请使用超分恢复命令处理")
attempt_no = await get_next_credit_attempt_no(
db,
owner_type=OWNER_GENERATION_RECORD,
owner_id=record.id,
)
if record.gen_type == GenerationType.video:
aspect_ratio = body.get("aspect_ratio", "16:9")
resolution = body.get("resolution", "720p")
if aspect_ratio not in ASPECT_RATIOS:
raise HTTPException(status_code=400, detail="不支持的画面比例")
if resolution not in RESOLUTIONS:
raise HTTPException(status_code=400, detail="不支持的分辨率")
from app.services.video_gen import get_active_engine
from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot
engine = await get_active_engine(db)
try:
supported_provider_resolutions = json.loads(engine.supported_resolutions or "[]")
except (TypeError, json.JSONDecodeError):
supported_provider_resolutions = []
provider_resolution, upscale_enabled, upscale_snapshot_json = await build_video_upscale_snapshot(
db,
target_resolution=resolution,
aspect_ratio=aspect_ratio,
supported_provider_resolutions=supported_provider_resolutions,
)
record.aspect_ratio = aspect_ratio
record.resolution = resolution
record.provider_generation_resolution = provider_resolution
record.video_upscale_enabled_snapshot = upscale_enabled
record.video_upscale_snapshot_json = upscale_snapshot_json
elif record.gen_type == GenerationType.image:
from app.services.image_gen import get_active_image_engine
engine = await get_active_image_engine(db)
record.image_size = body.get("image_size") or record.image_size or "2K"
record.provider_generation_resolution = None
record.video_upscale_enabled_snapshot = False
record.video_upscale_snapshot_json = None
else:
raise HTTPException(status_code=400, detail="不支持的生成类型")
media_billing = await charge_generation_media_for_record(
db,
record=record,
project_name=project_name,
description_prefix=f"{type_str}生成(管理后台)-",
attempt_no=attempt_no,
engine_id=engine.id,
)
record.credits_cost = round(float(record.credits_cost or 0) + float(media_billing.total_charged or 0), 2)
prepare_generation_record_execution(record, engine=engine, attempt_no=attempt_no)
await db.flush()
await commit_and_enqueue_generation_record(db, record, reason="generation_record_admin_generate")
await log_operation(
db,
admin.id,
admin.username,
f"管理员触发生成{type_str}: {record_id}",
"POST",
f"/admin/generation-records/{record_id}/generate",
detail=json.dumps(
{
"record_id": record_id,
"gen_type": record.gen_type,
"project_name": project_name,
"generation_attempt_no": record.generation_attempt_no,
},
ensure_ascii=False,
),
)
return {"message": "ok", "record_id": record_id}
# ── File Uploads ───────────────────────────────────────── # ── File Uploads ─────────────────────────────────────────
import os import os
@@ -2355,7 +2131,6 @@ async def upload_pdf(
): ):
"""Upload a PDF file and save URL to system config.""" """Upload a PDF file and save URL to system config."""
from app.config import settings from app.config import settings
from app.utils.id_gen import generate_id
if not file.filename: if not file.filename:
raise HTTPException(status_code=400, detail="请选择文件") raise HTTPException(status_code=400, detail="请选择文件")
@@ -2417,7 +2192,6 @@ async def upload_logo(
): ):
"""Upload a Logo image file and save URL to system config.""" """Upload a Logo image file and save URL to system config."""
from app.config import settings from app.config import settings
from app.utils.id_gen import generate_id
if not file.filename: if not file.filename:
raise HTTPException(status_code=400, detail="请选择文件") raise HTTPException(status_code=400, detail="请选择文件")
+158 -44
View File
@@ -55,6 +55,16 @@ from app.services.generation.billing_service import (
get_next_credit_attempt_no, get_next_credit_attempt_no,
) )
from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once
from app.services.generation.ai.engine_service import (
get_image_engine,
get_video_engine,
image_supported_sizes,
parse_json_list,
)
from app.services.generation.media_reference_service import (
calculate_media_reference_usage,
validate_media_reference_usage_for_engine,
)
from app.services.media_token_usage_snapshot_service import sync_generation_record_media_token_snapshot from app.services.media_token_usage_snapshot_service import sync_generation_record_media_token_snapshot
from app.services.credit_record_meta_service import build_generation_record_prompt_meta from app.services.credit_record_meta_service import build_generation_record_prompt_meta
from app.services.video_cover_service import async_create_video_cover_for_local_video from app.services.video_cover_service import async_create_video_cover_for_local_video
@@ -70,6 +80,36 @@ router = APIRouter(prefix="/generation-records", tags=["generation"])
logger = logging.getLogger("videogen") logger = logging.getLogger("videogen")
def _engine_snapshot(record: GenerationRecord) -> dict | None:
if not record.engine_snapshot_json:
return None
try:
value = json.loads(record.engine_snapshot_json)
except (TypeError, json.JSONDecodeError):
return None
return value if isinstance(value, dict) else None
def _validate_video_engine_selection(engine, *, aspect_ratio: str, resolution: str, duration: int) -> None:
ratios = [str(item) for item in parse_json_list(engine.supported_ratios, [])]
resolutions = [str(item) for item in parse_json_list(engine.supported_resolutions, [])]
durations = [int(item) for item in parse_json_list(engine.supported_durations, []) if str(item).isdigit()]
if ratios and aspect_ratio not in ratios:
raise HTTPException(status_code=400, detail="当前视频引擎不支持所选画面比例")
if resolutions and resolution not in resolutions:
raise HTTPException(status_code=400, detail="当前视频引擎不支持所选分辨率")
if durations and duration not in durations:
raise HTTPException(status_code=400, detail="当前视频引擎不支持所选时长")
if int(engine.max_duration or 0) > 0 and duration > int(engine.max_duration):
raise HTTPException(status_code=400, detail="生成时长超过当前视频引擎上限")
def _validate_image_engine_selection(engine, *, image_size: str) -> None:
sizes = image_supported_sizes(engine)
if sizes and image_size not in sizes:
raise HTTPException(status_code=400, detail="当前图片引擎不支持所选画面分辨率")
def _record_to_out(record: GenerationRecord, project_name: str, refs_override: list[dict] | None = None) -> GenerationRecordOut: def _record_to_out(record: GenerationRecord, project_name: str, refs_override: list[dict] | None = None) -> GenerationRecordOut:
refs = refs_override refs = refs_override
if refs is None and record.media_references: if refs is None and record.media_references:
@@ -112,6 +152,10 @@ def _record_to_out(record: GenerationRecord, project_name: str, refs_override: l
video_cover_url=build_resource_signed_url(record.video_cover_url) if record.video_cover_url else '', video_cover_url=build_resource_signed_url(record.video_cover_url) if record.video_cover_url else '',
image_url=build_resource_signed_url(record.image_url) if record.image_url else '', image_url=build_resource_signed_url(record.image_url) if record.image_url else '',
references=refs, references=refs,
engine_id=record.engine_id,
engine_name=(_engine_snapshot(record) or {}).get("name"),
engine_snapshot=_engine_snapshot(record),
include_media_references=bool(record.include_media_references),
text_credits_cost=round(record.text_credits_cost or 0.00, 2), text_credits_cost=round(record.text_credits_cost or 0.00, 2),
# text_tokens_used=record.text_tokens_used or 0, # text_tokens_used=record.text_tokens_used or 0,
credits_cost=round(record.credits_cost or 0.00, 2), credits_cost=round(record.credits_cost or 0.00, 2),
@@ -436,7 +480,11 @@ async def generate_record_resource(
raise InvalidStatusError("该任务生成失败,请联系客服进行修复") raise InvalidStatusError("该任务生成失败,请联系客服进行修复")
await assert_user_resource_capacity_available(db, current_user.id) await assert_user_resource_capacity_available(db, current_user.id)
attempt_no = await get_next_credit_attempt_no(db, owner_type=OWNER_GENERATION_RECORD, owner_id=record.id) attempt_no = await get_next_credit_attempt_no(
db, owner_type=OWNER_GENERATION_RECORD, owner_id=record.id
)
selected_engine_id = req.engine_id or record.engine_id
record.include_media_references = bool(req.include_media_references)
from app.services.generation.pipeline.generation_record_service import ( from app.services.generation.pipeline.generation_record_service import (
commit_and_enqueue_generation_record, commit_and_enqueue_generation_record,
@@ -444,56 +492,83 @@ async def generate_record_resource(
) )
if record.gen_type == GenerationType.video: if record.gen_type == GenerationType.video:
if req.aspect_ratio not in ASPECT_RATIOS: aspect_ratio = req.aspect_ratio or record.aspect_ratio
resolution = req.resolution or record.resolution
if aspect_ratio not in ASPECT_RATIOS:
raise HTTPException(status_code=400, detail="不支持的画面比例") raise HTTPException(status_code=400, detail="不支持的画面比例")
if req.resolution not in RESOLUTIONS: if resolution not in RESOLUTIONS:
raise HTTPException(status_code=400, detail="不支持的分辨率") raise HTTPException(status_code=400, detail="不支持的分辨率")
from app.services.video_gen import get_active_engine engine = await get_video_engine(db, selected_engine_id)
_validate_video_engine_selection(
engine,
aspect_ratio=aspect_ratio,
resolution=resolution,
duration=int(record.duration or 5),
)
from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot
engine = await get_active_engine(db) supported_provider_resolutions = parse_json_list(engine.supported_resolutions, [])
try:
supported_provider_resolutions = json.loads(engine.supported_resolutions or "[]")
except (TypeError, json.JSONDecodeError):
supported_provider_resolutions = []
provider_resolution, upscale_enabled, upscale_snapshot_json = await build_video_upscale_snapshot( provider_resolution, upscale_enabled, upscale_snapshot_json = await build_video_upscale_snapshot(
db, db,
target_resolution=req.resolution, target_resolution=resolution,
aspect_ratio=req.aspect_ratio, aspect_ratio=aspect_ratio,
supported_provider_resolutions=supported_provider_resolutions, supported_provider_resolutions=supported_provider_resolutions,
) )
billing = await charge_generation_media_by_params( record.aspect_ratio = aspect_ratio
db, user_id=current_user.id, record_id=record.id, gen_type="video", record.resolution = resolution
duration=record.duration or 5, resolution=req.resolution, engine_id=engine.id,
project_name=project_name, description_prefix=project_name + "-",
owner_type=OWNER_GENERATION_RECORD, attempt_no=attempt_no,
)
record.aspect_ratio = req.aspect_ratio
record.resolution = req.resolution
record.provider_generation_resolution = provider_resolution record.provider_generation_resolution = provider_resolution
record.video_upscale_enabled_snapshot = upscale_enabled record.video_upscale_enabled_snapshot = upscale_enabled
record.video_upscale_snapshot_json = upscale_snapshot_json record.video_upscale_snapshot_json = upscale_snapshot_json
else: else:
from app.services.image_gen import get_active_image_engine engine = await get_image_engine(db, selected_engine_id)
engine = await get_active_image_engine(db)
image_size = req.image_size or record.image_size or engine.default_size or "2K" image_size = req.image_size or record.image_size or engine.default_size or "2K"
billing = await charge_generation_media_by_params( _validate_image_engine_selection(engine, image_size=image_size)
db, user_id=current_user.id, record_id=record.id, gen_type="image",
image_size=image_size, engine_id=engine.id, project_name=project_name,
description_prefix=project_name + "-", owner_type=OWNER_GENERATION_RECORD, attempt_no=attempt_no,
)
record.image_size = image_size record.image_size = image_size
record.provider_generation_resolution = None record.provider_generation_resolution = None
record.video_upscale_enabled_snapshot = False record.video_upscale_enabled_snapshot = False
record.video_upscale_snapshot_json = None record.video_upscale_snapshot_json = None
record.credits_cost = round(float(record.credits_cost or 0) + float(billing.total_charged or 0), 2) reference_usage = calculate_media_reference_usage(
record.media_references,
include=bool(record.include_media_references),
)
validate_media_reference_usage_for_engine(
reference_usage,
gen_type=record.gen_type,
engine=engine,
)
billing = await charge_generation_media_for_record(
db,
record=record,
project_name=project_name,
description_prefix=project_name + "-",
attempt_no=attempt_no,
engine_id=engine.id,
)
record.credits_cost = round(
float(record.credits_cost or 0) + float(billing.total_charged or 0), 2
)
prepare_generation_record_execution(record, engine=engine, attempt_no=attempt_no) prepare_generation_record_execution(record, engine=engine, attempt_no=attempt_no)
await db.flush() await db.flush()
await commit_and_enqueue_generation_record(db, record, reason="generation_record_api_generate") record_id_snapshot = str(record.id)
await commit_and_enqueue_generation_record(
db, record, reason="generation_record_api_generate"
)
refreshed = await db.execute(
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(GenerationRecord.id == record_id_snapshot)
.limit(1)
)
refreshed_row = refreshed.first()
if not refreshed_row:
raise RecordNotFoundError()
record, project_name = refreshed_row
refs = await resolve_private_portrait_reference_display_urls( refs = await resolve_private_portrait_reference_display_urls(
db, json.loads(record.media_references) if record.media_references else None, user_id=current_user.id db,
json.loads(record.media_references) if record.media_references else None,
user_id=current_user.id,
) )
return _record_to_out(record, project_name, refs_override=refs) return _record_to_out(record, project_name, refs_override=refs)
@@ -529,43 +604,82 @@ async def retry_generation(
raise InvalidStatusError("该任务生成失败,请联系客服进行修复") raise InvalidStatusError("该任务生成失败,请联系客服进行修复")
await assert_user_resource_capacity_available(db, current_user.id) await assert_user_resource_capacity_available(db, current_user.id)
attempt_no = await get_next_credit_attempt_no(db, owner_type=OWNER_GENERATION_RECORD, owner_id=record.id) attempt_no = await get_next_credit_attempt_no(
db, owner_type=OWNER_GENERATION_RECORD, owner_id=record.id
)
from app.services.generation.pipeline.generation_record_service import ( from app.services.generation.pipeline.generation_record_service import (
commit_and_enqueue_generation_record, commit_and_enqueue_generation_record,
prepare_generation_record_execution, prepare_generation_record_execution,
) )
if record.gen_type == GenerationType.video: if record.gen_type == GenerationType.video:
from app.services.video_gen import get_active_engine engine = await get_video_engine(db, record.engine_id)
_validate_video_engine_selection(
engine,
aspect_ratio=record.aspect_ratio or "16:9",
resolution=record.resolution or "480p",
duration=int(record.duration or 5),
)
from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot
engine = await get_active_engine(db)
try:
supported_provider_resolutions = json.loads(engine.supported_resolutions or "[]")
except (TypeError, json.JSONDecodeError):
supported_provider_resolutions = []
provider_resolution, upscale_enabled, upscale_snapshot_json = await build_video_upscale_snapshot( provider_resolution, upscale_enabled, upscale_snapshot_json = await build_video_upscale_snapshot(
db, target_resolution=record.resolution or "480p", aspect_ratio=record.aspect_ratio or "16:9", db,
supported_provider_resolutions=supported_provider_resolutions, target_resolution=record.resolution or "480p",
aspect_ratio=record.aspect_ratio or "16:9",
supported_provider_resolutions=parse_json_list(engine.supported_resolutions, []),
) )
record.provider_generation_resolution = provider_resolution record.provider_generation_resolution = provider_resolution
record.video_upscale_enabled_snapshot = upscale_enabled record.video_upscale_enabled_snapshot = upscale_enabled
record.video_upscale_snapshot_json = upscale_snapshot_json record.video_upscale_snapshot_json = upscale_snapshot_json
else: else:
from app.services.image_gen import get_active_image_engine engine = await get_image_engine(db, record.engine_id)
engine = await get_active_image_engine(db) _validate_image_engine_selection(
engine, image_size=record.image_size or engine.default_size or "2K"
)
reference_usage = calculate_media_reference_usage(
record.media_references,
include=bool(record.include_media_references),
)
validate_media_reference_usage_for_engine(
reference_usage,
gen_type=record.gen_type,
engine=engine,
)
billing = await charge_generation_media_for_record( billing = await charge_generation_media_for_record(
db, record=record, project_name=project_name, description_prefix="资源生成重试-", attempt_no=attempt_no, engine_id=engine.id db,
record=record,
project_name=project_name,
description_prefix="资源生成重试-",
attempt_no=attempt_no,
engine_id=engine.id,
)
record.credits_cost = round(
float(record.credits_cost or 0) + float(billing.total_charged or 0), 2
) )
record.credits_cost = round(float(record.credits_cost or 0) + float(billing.total_charged or 0), 2)
record.manual_retry_count = int(record.manual_retry_count or 0) + 1 record.manual_retry_count = int(record.manual_retry_count or 0) + 1
record.retry_count = int(record.manual_retry_count or 0) record.retry_count = int(record.manual_retry_count or 0)
prepare_generation_record_execution(record, engine=engine, attempt_no=attempt_no) prepare_generation_record_execution(record, engine=engine, attempt_no=attempt_no)
await db.flush() await db.flush()
await commit_and_enqueue_generation_record(db, record, reason="generation_record_api_retry") record_id_snapshot = str(record.id)
await commit_and_enqueue_generation_record(
db, record, reason="generation_record_api_retry"
)
refreshed = await db.execute(
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(GenerationRecord.id == record_id_snapshot)
.limit(1)
)
refreshed_row = refreshed.first()
if not refreshed_row:
raise RecordNotFoundError()
record, project_name = refreshed_row
refs = await resolve_private_portrait_reference_display_urls( refs = await resolve_private_portrait_reference_display_urls(
db, json.loads(record.media_references) if record.media_references else None, user_id=current_user.id db,
json.loads(record.media_references) if record.media_references else None,
user_id=current_user.id,
) )
return _record_to_out(record, project_name, refs_override=refs) return _record_to_out(record, project_name, refs_override=refs)
@@ -9,7 +9,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_current_user, get_db from app.dependencies import get_current_user, get_db
from app.models.user import User from app.models.user import User
from app.enums.common import ModuleProjectStatusEnum from app.enums.common import ModuleProjectStatusEnum, ModuleEventTypeEnum
from app.enums.generation_task import GenerationOwnerType from app.enums.generation_task import GenerationOwnerType
from app.enums.hot_opening_replicate import HotOpeningLogEventEnum, HotOpeningStepCodeEnum, ModuleCodeEnum from app.enums.hot_opening_replicate import HotOpeningLogEventEnum, HotOpeningStepCodeEnum, ModuleCodeEnum
from app.schemas.hot_opening_replicate import ( from app.schemas.hot_opening_replicate import (
@@ -42,7 +42,7 @@ from app.services.hot_opening_replicate_service import (
update_hot_opening_material_input, update_hot_opening_material_input,
update_hot_opening_video_prompt_schema, update_hot_opening_video_prompt_schema,
) )
from app.services.module_generation_log_service import log_module_error from app.services.module_generation_log_service import log_module_error, log_module_event_file
from app.services.module_async_recovery_service import ( from app.services.module_async_recovery_service import (
TASK_HOT_IMAGE_PROMPT, TASK_HOT_IMAGE_PROMPT,
TASK_HOT_VIDEO_PROMPT, TASK_HOT_VIDEO_PROMPT,
@@ -283,29 +283,17 @@ async def create_task(
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db), db: AsyncSession = Depends(get_db),
): ):
try: log_module_event_file(
project = await create_hot_opening_project(db, current_user, req) module=MODULE,
project_id_value = str(project.id) event_type=ModuleEventTypeEnum.V1_CREATE_BLOCKED.value,
await bind_upload_resources( user_id=current_user.id,
db, message="拦截爆款开头复刻 V1 创建请求",
user_id=current_user.id, detail={"api_version": "v1", "flow_version": "v1"},
module=UploadResourceModuleEnum.HOT_OPENING_REPLICATE.value, )
source_model=UploadResourceSourceModelEnum.MODULE_GENERATION_PROJECT.value, raise HTTPException(
source_id=project_id_value, status_code=410,
resource_ids=[req.material_video_resource_id, req.material_image_resource_id], detail="V1 创建流程已停止,请使用 V2 API",
urls=[req.material_video_url, req.material_image_url], )
allow_common_migrate=True,
)
await db.commit()
except HTTPException:
await db.rollback()
raise
except Exception as exc:
await db.rollback()
_log_api_exception_from_locals(exc, locals(), f"创建爆款开头复刻项目失败: {exc}")
raise HTTPException(status_code=500, detail=f"创建爆款开头复刻项目失败: {exc}")
return await _reload_project_detail(db, current_user, project_id_value)
@router.get( @router.get(
+12 -18
View File
@@ -10,6 +10,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings from app.config import settings
from app.dependencies import get_current_user, get_db from app.dependencies import get_current_user, get_db
from app.models.user import User from app.models.user import User
from app.enums.common import ModuleEventTypeEnum
from app.enums.generation_task import GenerationOwnerType from app.enums.generation_task import GenerationOwnerType
from app.enums.shot_replicate import ( from app.enums.shot_replicate import (
ModuleCodeEnum, ModuleCodeEnum,
@@ -862,24 +863,17 @@ async def create_replication_project_from_segment(
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db), db: AsyncSession = Depends(get_db),
): ):
try: log_module_event_file(
segment = await get_segment_for_user(db, segment_id=segment_id, user=current_user, for_update=True) module=MODULE,
project = await create_shot_replicate_project_from_segment(db, current_user=current_user, segment=segment, req=req) event_type=ModuleEventTypeEnum.V1_CREATE_BLOCKED.value,
project_id = project.id user_id=current_user.id,
await db.commit() step_id=segment_id,
except HTTPException: message="拦截拆镜复刻 V1 创建请求",
await db.rollback() detail={"api_version": "v1", "flow_version": "v1", "segment_id": segment_id},
raise )
except Exception as exc: raise HTTPException(
await db.rollback() status_code=410,
_log_api_exception_from_locals(exc, locals(), f"创建拆镜复刻项目失败: {exc}") detail="V1 创建流程已停止,请使用 V2 API",
raise HTTPException(status_code=500, detail=f"创建拆镜复刻项目失败: {exc}")
return ShotReplicateActionOut(
message="已从拆镜片段创建复刻项目,素材视频已锁定",
project_id=project_id,
step_id=None,
detail=await _reload_project_detail(db, current_user, project_id),
) )
+8
View File
@@ -0,0 +1,8 @@
from fastapi import APIRouter
from app.api.v2.hot_opening_replicate import router as hot_opening_router
from app.api.v2.shot_replicate import router as shot_replicate_router
api_router_v2 = APIRouter()
api_router_v2.include_router(hot_opening_router)
api_router_v2.include_router(shot_replicate_router)
@@ -0,0 +1,255 @@
from __future__ import annotations
from fastapi import APIRouter, Body, Depends, HTTPException, Path
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_current_user, get_db
from app.models.chat_generation_task import ChatGenerationTask
from app.models.user import User
from app.schemas.hot_opening_replicate import HotOpeningActionOut, HotOpeningDeleteOut, HotOpeningTaskDetailOut
from app.schemas.module_generation_v2 import (
HotOpeningTaskCreateV2,
ModuleVideoPromptRetryV2,
ModuleVideoPromptSchemaUpdateV2,
)
from app.services.generation.pipeline.enqueue_service import enqueue_generation_create
from app.services.hot_opening_replicate_service import project_to_detail_out
from app.services.module_generation_v2.config import HOT_OPENING_V2
from app.services.module_generation_v2.dispatch_service import (
dispatch_video_prompt_v2,
ensure_v2_celery_enabled,
)
from app.services.module_generation_v2.flow_service import (
create_hot_opening_project_v2,
delete_project_v2,
generate_video_from_prompt_v2,
get_v2_project_for_user,
is_project_idempotency_conflict,
mark_video_prompt_dispatch_failed_v2,
rebuild_video_prompt_step_v2,
update_video_prompt_schema_v2,
)
from app.services.upload_resource import cleanup_upload_resource_files_after_commit
router = APIRouter(prefix="/hot-opening-replications", tags=["hot-opening-replications-v2"])
async def _detail(db: AsyncSession, current_user: User, project_id: str) -> HotOpeningTaskDetailOut:
project = await get_v2_project_for_user(
db,
config=HOT_OPENING_V2,
project_id=project_id,
current_user=current_user,
)
return await project_to_detail_out(db, project)
async def _dispatch_or_mark_failed(
db: AsyncSession,
*,
project_id: str,
step_id: str,
) -> None:
dispatch = await dispatch_video_prompt_v2(
config=HOT_OPENING_V2,
project_id=project_id,
step_id=step_id,
)
if dispatch.recoverable:
return
error_message = "视频提词任务的 Redis 注册和 Celery 投递均失败,请重新执行步骤2"
await mark_video_prompt_dispatch_failed_v2(
db,
config=HOT_OPENING_V2,
project_id=project_id,
step_id=step_id,
error_message=error_message,
)
raise HTTPException(status_code=503, detail=error_message)
@router.post("/tasks", response_model=HotOpeningTaskDetailOut)
async def create_task_v2(
req: HotOpeningTaskCreateV2 = Body(...),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
ensure_v2_celery_enabled()
try:
result = await create_hot_opening_project_v2(db, current_user=current_user, req=req)
project_id = str(result.project.id)
step_id = str(result.prompt_step.id)
created_new = bool(result.created_new)
await db.commit()
except IntegrityError as exc:
await db.rollback()
if not req.idempotency_key or not is_project_idempotency_conflict(exc):
raise HTTPException(status_code=500, detail="项目创建失败") from exc
# 同幂等键并发请求由唯一索引收敛;回查已提交项目并按幂等成功返回。
result = await create_hot_opening_project_v2(db, current_user=current_user, req=req)
project_id = str(result.project.id)
step_id = str(result.prompt_step.id)
created_new = bool(result.created_new)
await db.commit()
except HTTPException:
await db.rollback()
raise
except Exception as exc:
await db.rollback()
raise HTTPException(status_code=500, detail="创建爆款复刻 V2 项目失败") from exc
if created_new:
await _dispatch_or_mark_failed(db, project_id=project_id, step_id=step_id)
return await _detail(db, current_user, project_id)
@router.get("/tasks/{project_id}", response_model=HotOpeningTaskDetailOut)
async def get_task_v2(
project_id: str = Path(...),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
return await _detail(db, current_user, project_id)
@router.post(
"/tasks/{project_id}/steps/{step_id}/retry-video-prompt",
response_model=HotOpeningActionOut,
)
async def retry_video_prompt_v2(
project_id: str,
step_id: str,
req: ModuleVideoPromptRetryV2,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
ensure_v2_celery_enabled()
try:
project, new_step = await rebuild_video_prompt_step_v2(
db,
config=HOT_OPENING_V2,
current_user=current_user,
project_id=project_id,
source_prompt_step_id=step_id,
video_config=req.video_config,
)
project_id_value = str(project.id)
step_id_value = str(new_step.id)
await db.commit()
except HTTPException:
await db.rollback()
raise
await _dispatch_or_mark_failed(db, project_id=project_id_value, step_id=step_id_value)
return HotOpeningActionOut(
message="视频提词已重新提交",
project_id=project_id_value,
step_id=step_id_value,
detail=await _detail(db, current_user, project_id_value),
)
@router.put(
"/tasks/{project_id}/steps/{step_id}/video-prompt-schema",
response_model=HotOpeningActionOut,
)
async def update_video_prompt_schema_route_v2(
project_id: str,
step_id: str,
req: ModuleVideoPromptSchemaUpdateV2,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
try:
project, step = await update_video_prompt_schema_v2(
db,
config=HOT_OPENING_V2,
current_user=current_user,
project_id=project_id,
step_id=step_id,
req=req,
)
project_id_value = str(project.id)
step_id_value = str(step.id)
await db.commit()
except HTTPException:
await db.rollback()
raise
return HotOpeningActionOut(
message="视频提词已保存",
project_id=project_id_value,
step_id=step_id_value,
detail=await _detail(db, current_user, project_id_value),
)
@router.post(
"/tasks/{project_id}/steps/{step_id}/generate-video",
response_model=HotOpeningActionOut,
)
async def generate_video_v2(
project_id: str,
step_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
ensure_v2_celery_enabled()
try:
project, step, task = await generate_video_from_prompt_v2(
db,
config=HOT_OPENING_V2,
current_user=current_user,
project_id=project_id,
prompt_step_id=step_id,
)
project_id_value = str(project.id)
step_id_value = str(step.id)
task_id = str(task.id)
await db.commit()
except HTTPException:
await db.rollback()
raise
# commit 后重新读取,避免 ORM expire/lazy-load 风险。
queued_task = await db.get(ChatGenerationTask, task_id)
if queued_task is None:
raise HTTPException(status_code=500, detail="视频生成任务提交后无法重新读取")
try:
await enqueue_generation_create(queued_task, reason="hot_opening_v2_generate_video")
except Exception as exc:
# queued 状态已持久化,周期生成恢复任务会使用确定性 task_id 补投。
raise HTTPException(status_code=503, detail="视频生成任务暂未投递,将由恢复任务自动补投") from exc
return HotOpeningActionOut(
message="视频生成任务已提交",
project_id=project_id_value,
step_id=step_id_value,
detail=await _detail(db, current_user, project_id_value),
)
@router.delete("/tasks/{project_id}", response_model=HotOpeningDeleteOut)
async def delete_project_route_v2(
project_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
try:
payload = await delete_project_v2(
db,
config=HOT_OPENING_V2,
current_user=current_user,
project_id=project_id,
)
pending_ids = list(payload.get("pending_delete_resource_ids") or [])
await db.commit()
except HTTPException:
await db.rollback()
raise
if pending_ids:
try:
await cleanup_upload_resource_files_after_commit(db, resource_ids=pending_ids)
await db.commit()
except Exception:
await db.rollback()
return HotOpeningDeleteOut(**payload)
+272
View File
@@ -0,0 +1,272 @@
from __future__ import annotations
from fastapi import APIRouter, Body, Depends, HTTPException, Path
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_current_user, get_db
from app.models.chat_generation_task import ChatGenerationTask
from app.models.user import User
from app.schemas.module_generation_v2 import (
ModuleVideoPromptRetryV2,
ModuleVideoPromptSchemaUpdateV2,
ShotReplicateProjectCreateV2,
)
from app.schemas.shot_replicate import ShotReplicateActionOut, ShotReplicateDeleteOut, ShotReplicateTaskDetailOut
from app.services.generation.pipeline.enqueue_service import enqueue_generation_create
from app.services.module_generation_v2.config import SHOT_REPLICATE_V2
from app.services.module_generation_v2.dispatch_service import (
dispatch_video_prompt_v2,
ensure_v2_celery_enabled,
)
from app.services.module_generation_v2.flow_service import (
create_shot_replicate_project_v2,
delete_project_v2,
generate_video_from_prompt_v2,
get_v2_project_for_user,
is_project_idempotency_conflict,
mark_video_prompt_dispatch_failed_v2,
rebuild_video_prompt_step_v2,
update_video_prompt_schema_v2,
)
from app.services.shot_replicate_flow_service import project_to_detail_out
from app.services.shot_replicate_taskset_service import get_segment_for_user
from app.services.upload_resource import cleanup_upload_resource_files_after_commit
router = APIRouter(prefix="/shot-replications", tags=["shot-replications-v2"])
async def _detail(db: AsyncSession, current_user: User, project_id: str) -> ShotReplicateTaskDetailOut:
project = await get_v2_project_for_user(
db,
config=SHOT_REPLICATE_V2,
project_id=project_id,
current_user=current_user,
)
return await project_to_detail_out(db, project)
async def _dispatch_or_mark_failed(
db: AsyncSession,
*,
project_id: str,
step_id: str,
) -> None:
dispatch = await dispatch_video_prompt_v2(
config=SHOT_REPLICATE_V2,
project_id=project_id,
step_id=step_id,
)
if dispatch.recoverable:
return
error_message = "视频提词任务的 Redis 注册和 Celery 投递均失败,请重新执行步骤2"
await mark_video_prompt_dispatch_failed_v2(
db,
config=SHOT_REPLICATE_V2,
project_id=project_id,
step_id=step_id,
error_message=error_message,
)
raise HTTPException(status_code=503, detail=error_message)
@router.post(
"/segments/{segment_id}/replication-projects",
response_model=ShotReplicateActionOut,
)
async def create_project_v2(
segment_id: str = Path(...),
req: ShotReplicateProjectCreateV2 = Body(...),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
ensure_v2_celery_enabled()
try:
segment = await get_segment_for_user(
db, segment_id=segment_id, user=current_user, for_update=True
)
result = await create_shot_replicate_project_v2(
db, current_user=current_user, segment=segment, req=req
)
project_id = str(result.project.id)
step_id = str(result.prompt_step.id)
created_new = bool(result.created_new)
await db.commit()
except IntegrityError as exc:
await db.rollback()
if not req.idempotency_key or not is_project_idempotency_conflict(exc):
raise HTTPException(status_code=500, detail="项目创建失败") from exc
segment = await get_segment_for_user(
db, segment_id=segment_id, user=current_user, for_update=True
)
result = await create_shot_replicate_project_v2(
db, current_user=current_user, segment=segment, req=req
)
project_id = str(result.project.id)
step_id = str(result.prompt_step.id)
created_new = bool(result.created_new)
await db.commit()
except HTTPException:
await db.rollback()
raise
except Exception as exc:
await db.rollback()
raise HTTPException(status_code=500, detail="创建拆镜复刻 V2 项目失败") from exc
if created_new:
await _dispatch_or_mark_failed(db, project_id=project_id, step_id=step_id)
return ShotReplicateActionOut(
message="V2 项目已创建,视频提词已自动提交" if created_new else "已返回现有幂等项目",
project_id=project_id,
step_id=step_id,
detail=await _detail(db, current_user, project_id),
)
@router.get("/projects/{project_id}", response_model=ShotReplicateTaskDetailOut)
async def get_project_v2(
project_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
return await _detail(db, current_user, project_id)
@router.post(
"/projects/{project_id}/steps/{step_id}/retry-video-prompt",
response_model=ShotReplicateActionOut,
)
async def retry_video_prompt_v2(
project_id: str,
step_id: str,
req: ModuleVideoPromptRetryV2,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
ensure_v2_celery_enabled()
try:
project, new_step = await rebuild_video_prompt_step_v2(
db,
config=SHOT_REPLICATE_V2,
current_user=current_user,
project_id=project_id,
source_prompt_step_id=step_id,
video_config=req.video_config,
)
project_id_value = str(project.id)
step_id_value = str(new_step.id)
await db.commit()
except HTTPException:
await db.rollback()
raise
await _dispatch_or_mark_failed(db, project_id=project_id_value, step_id=step_id_value)
return ShotReplicateActionOut(
message="视频提词已重新提交",
project_id=project_id_value,
step_id=step_id_value,
detail=await _detail(db, current_user, project_id_value),
)
@router.put(
"/projects/{project_id}/steps/{step_id}/video-prompt-schema",
response_model=ShotReplicateActionOut,
)
async def update_video_prompt_schema_route_v2(
project_id: str,
step_id: str,
req: ModuleVideoPromptSchemaUpdateV2,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
try:
project, step = await update_video_prompt_schema_v2(
db,
config=SHOT_REPLICATE_V2,
current_user=current_user,
project_id=project_id,
step_id=step_id,
req=req,
)
project_id_value = str(project.id)
step_id_value = str(step.id)
await db.commit()
except HTTPException:
await db.rollback()
raise
return ShotReplicateActionOut(
message="视频提词已保存",
project_id=project_id_value,
step_id=step_id_value,
detail=await _detail(db, current_user, project_id_value),
)
@router.post(
"/projects/{project_id}/steps/{step_id}/generate-video",
response_model=ShotReplicateActionOut,
)
async def generate_video_v2(
project_id: str,
step_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
ensure_v2_celery_enabled()
try:
project, step, task = await generate_video_from_prompt_v2(
db,
config=SHOT_REPLICATE_V2,
current_user=current_user,
project_id=project_id,
prompt_step_id=step_id,
)
project_id_value = str(project.id)
step_id_value = str(step.id)
task_id = str(task.id)
await db.commit()
except HTTPException:
await db.rollback()
raise
queued_task = await db.get(ChatGenerationTask, task_id)
if queued_task is None:
raise HTTPException(status_code=500, detail="视频生成任务提交后无法重新读取")
try:
await enqueue_generation_create(queued_task, reason="shot_replicate_v2_generate_video")
except Exception as exc:
raise HTTPException(status_code=503, detail="视频生成任务暂未投递,将由恢复任务自动补投") from exc
return ShotReplicateActionOut(
message="视频生成任务已提交",
project_id=project_id_value,
step_id=step_id_value,
detail=await _detail(db, current_user, project_id_value),
)
@router.delete("/projects/{project_id}", response_model=ShotReplicateDeleteOut)
async def delete_project_route_v2(
project_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
try:
payload = await delete_project_v2(
db,
config=SHOT_REPLICATE_V2,
current_user=current_user,
project_id=project_id,
)
pending_ids = list(payload.get("pending_delete_resource_ids") or [])
await db.commit()
except HTTPException:
await db.rollback()
raise
if pending_ids:
try:
await cleanup_upload_resource_files_after_commit(db, resource_ids=pending_ids)
await db.commit()
except Exception:
await db.rollback()
return ShotReplicateDeleteOut(**payload)
+3
View File
@@ -206,6 +206,9 @@ class Settings(BaseSettings):
# - poll active 使用独立 Redis key,避免影响稳定的下载 active 注册表。 # - poll active 使用独立 Redis key,避免影响稳定的下载 active 注册表。
GENERATION_RECOVERY_BATCH_SIZE: int = 20 GENERATION_RECOVERY_BATCH_SIZE: int = 20
GENERATION_RECOVERY_MAX_ROUNDS: int = 1 GENERATION_RECOVERY_MAX_ROUNDS: int = 1
GENERATION_CREATE_RECOVERY_INTERVAL_SECONDS: int = 60
GENERATION_CREATE_QUEUE_TIMEOUT_SECONDS: int = 5 * 60
MODULE_ASYNC_RECOVERY_INTERVAL_SECONDS: int = 60
POLL_RECOVERY_BATCH_SIZE: int = 20 POLL_RECOVERY_BATCH_SIZE: int = 20
POLL_TASK_LEASE_SECONDS: int = 5 * 60 POLL_TASK_LEASE_SECONDS: int = 5 * 60
POLL_TASK_QUEUE_TIMEOUT_SECONDS: int = 2 * 60 POLL_TASK_QUEUE_TIMEOUT_SECONDS: int = 2 * 60
+1
View File
@@ -18,6 +18,7 @@ class CeleryTaskName(str, Enum):
DOWNLOAD_GENERATION_RESULT = "generation.download_generation_result_task" DOWNLOAD_GENERATION_RESULT = "generation.download_generation_result_task"
RECOVER_DOWNLOAD = "generation.recover_download_tasks_once" RECOVER_DOWNLOAD = "generation.recover_download_tasks_once"
RECOVER_GENERATION = "generation.recover_generation_tasks_once" RECOVER_GENERATION = "generation.recover_generation_tasks_once"
RECOVER_CREATE = "generation.recover_create_tasks_once"
VIDEO_UPSCALE_EXECUTE_LOCAL = "video_upscale.execute_local" VIDEO_UPSCALE_EXECUTE_LOCAL = "video_upscale.execute_local"
VIDEO_UPSCALE_SUBMIT_REMOTE = "video_upscale.submit_remote" VIDEO_UPSCALE_SUBMIT_REMOTE = "video_upscale.submit_remote"
VIDEO_UPSCALE_POLL_REMOTE = "video_upscale.poll_remote" VIDEO_UPSCALE_POLL_REMOTE = "video_upscale.poll_remote"
+20
View File
@@ -27,6 +27,13 @@ class LogSourceEnum(StrEnum):
RECOVERY = "recovery" RECOVERY = "recovery"
REMOTE_API = "remote_api" REMOTE_API = "remote_api"
class ModuleGenerationFlowVersionEnum(StrEnum):
"""模块生成项目流程版本。"""
V1 = "v1"
V2 = "v2"
class ModuleProjectStatusEnum(StrEnum): class ModuleProjectStatusEnum(StrEnum):
"""通用模块项目状态。""" """通用模块项目状态。"""
@@ -72,6 +79,19 @@ class ModuleEventTypeEnum(StrEnum):
MEDIA_REFUND = "MEDIA_REFUND" MEDIA_REFUND = "MEDIA_REFUND"
PROMPT_BILLING_SUCCESS = "PROMPT_BILLING_SUCCESS" PROMPT_BILLING_SUCCESS = "PROMPT_BILLING_SUCCESS"
PROMPT_BILLING_FAILED = "PROMPT_BILLING_FAILED" PROMPT_BILLING_FAILED = "PROMPT_BILLING_FAILED"
V1_CREATE_BLOCKED = "V1_CREATE_BLOCKED"
FLOW_VERSION_MISMATCH = "FLOW_VERSION_MISMATCH"
V2_PROJECT_CREATED = "V2_PROJECT_CREATED"
V2_VIDEO_PROMPT_AUTO_CREATED = "V2_VIDEO_PROMPT_AUTO_CREATED"
V2_VIDEO_PROMPT_REGENERATED = "V2_VIDEO_PROMPT_REGENERATED"
V2_VIDEO_PROMPT_DISPATCHED = "V2_VIDEO_PROMPT_DISPATCHED"
V2_VIDEO_PROMPT_REGISTRY_FAILED = "V2_VIDEO_PROMPT_REGISTRY_FAILED"
V2_VIDEO_PROMPT_DISPATCH_FAILED = "V2_VIDEO_PROMPT_DISPATCH_FAILED"
STEP_SUPERSEDED = "STEP_SUPERSEDED"
STALE_STEP_RESULT_DISCARDED = "STALE_STEP_RESULT_DISCARDED"
GENERATION_REFERENCE_OPTION_SAVED = "GENERATION_REFERENCE_OPTION_SAVED"
GENERATION_REFERENCE_INCLUDED = "GENERATION_REFERENCE_INCLUDED"
GENERATION_REFERENCE_EXCLUDED = "GENERATION_REFERENCE_EXCLUDED"
class ModulePromptTypeEnum(StrEnum): class ModulePromptTypeEnum(StrEnum):
@@ -30,6 +30,7 @@ class HotOpeningStepIOSchemaVersionEnum(StrEnum):
"""爆款开头复刻子任务 input_json/output_json 结构版本。""" """爆款开头复刻子任务 input_json/output_json 结构版本。"""
V1 = "hot_opening_step_io_v1" V1 = "hot_opening_step_io_v1"
V2 = "hot_opening_step_io_v2"
class HotOpeningLogEventEnum(StrEnum): class HotOpeningLogEventEnum(StrEnum):
@@ -24,3 +24,4 @@ class ModuleGenerationFlowConfig:
cancel_chat_task_error_message: str cancel_chat_task_error_message: str
material_video_url_editable: bool = True material_video_url_editable: bool = True
step_io_schema_version: str = "module_generation_step_io_v1" step_io_schema_version: str = "module_generation_step_io_v1"
expected_flow_version: str | None = None
@@ -29,6 +29,7 @@ class ShotReplicateStepIOSchemaVersionEnum(StrEnum):
"""拆镜复刻子任务 input_json/output_json 结构版本。""" """拆镜复刻子任务 input_json/output_json 结构版本。"""
V1 = "shot_replicate_step_io_v1" V1 = "shot_replicate_step_io_v1"
V2 = "shot_replicate_step_io_v2"
class ShotTaskSetStatusEnum(StrEnum): class ShotTaskSetStatusEnum(StrEnum):
+2
View File
@@ -12,6 +12,7 @@ from app.config import settings
from app.models import init_database, close_database from app.models import init_database, close_database
from app.utils.redis import init_redis, close_redis from app.utils.redis import init_redis, close_redis
from app.api.v1 import api_router from app.api.v1 import api_router
from app.api.v2 import api_router_v2
from app.middleware.logging import RequestLoggingMiddleware from app.middleware.logging import RequestLoggingMiddleware
from app.middleware.anti_crawler import AntiCrawlerMiddleware from app.middleware.anti_crawler import AntiCrawlerMiddleware
from app.middleware.rate_limit import RateLimitMiddleware from app.middleware.rate_limit import RateLimitMiddleware
@@ -545,6 +546,7 @@ def create_app() -> FastAPI:
# Routes # Routes
application.include_router(api_router, prefix="/api") application.include_router(api_router, prefix="/api")
application.include_router(api_router_v2, prefix="/api/v2")
# Static files for uploads # Static files for uploads
upload_dir = os.path.abspath(settings.UPLOAD_LOCAL_PATH) upload_dir = os.path.abspath(settings.UPLOAD_LOCAL_PATH)
@@ -37,6 +37,9 @@ class GenerationRecord(Base, TimestampMixin, SoftDeleteMixin):
video_cover_url: Mapped[str | None] = mapped_column(String(512), nullable=True) video_cover_url: Mapped[str | None] = mapped_column(String(512), nullable=True)
image_url: Mapped[str | None] = mapped_column(String(512), nullable=True) image_url: Mapped[str | None] = mapped_column(String(512), nullable=True)
media_references: Mapped[str | None] = mapped_column(Text, nullable=True) media_references: Mapped[str | None] = mapped_column(Text, nullable=True)
include_media_references: Mapped[bool] = mapped_column(
Boolean, nullable=False, default=False, server_default="false"
)
video_url_expires_at: Mapped[datetime | None] = mapped_column( video_url_expires_at: Mapped[datetime | None] = mapped_column(
DateTime(timezone=True), nullable=True DateTime(timezone=True), nullable=True
) )
@@ -5,6 +5,7 @@ from datetime import datetime
from sqlalchemy import DateTime, ForeignKey, Index, String, Text, text from sqlalchemy import DateTime, ForeignKey, Index, String, Text, text
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
from app.enums.common import ModuleGenerationFlowVersionEnum
from app.models.base import Base, SoftDeleteMixin, TimestampMixin from app.models.base import Base, SoftDeleteMixin, TimestampMixin
@@ -36,6 +37,12 @@ class ModuleGenerationProject(Base, TimestampMixin, SoftDeleteMixin):
String(32), ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False String(32), ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
) )
module: Mapped[str] = mapped_column(String(64), index=True, nullable=False) module: Mapped[str] = mapped_column(String(64), index=True, nullable=False)
flow_version: Mapped[str] = mapped_column(
String(16),
nullable=False,
default=ModuleGenerationFlowVersionEnum.V1.value,
server_default=ModuleGenerationFlowVersionEnum.V1.value,
)
title: Mapped[str | None] = mapped_column(String(160), nullable=True) title: Mapped[str | None] = mapped_column(String(160), nullable=True)
status: Mapped[str] = mapped_column(String(32), default="pending", index=True) status: Mapped[str] = mapped_column(String(32), default="pending", index=True)
current_step_code: Mapped[str | None] = mapped_column(String(64), nullable=True, index=True) current_step_code: Mapped[str | None] = mapped_column(String(64), nullable=True, index=True)
@@ -15,12 +15,12 @@ _STEP_JSON_TYPE = JSON().with_variant(JSONB, "postgresql")
class ModuleGenerationStep(Base, TimestampMixin, SoftDeleteMixin): class ModuleGenerationStep(Base, TimestampMixin, SoftDeleteMixin):
"""通用模块生成步骤表。 """通用模块生成步骤表。
爆款开头复刻固定步骤: V1 固定五步:material_input / image_prompt_optimize / image_generate /
1 material_input video_prompt_optimize / video_generate。
2 image_prompt_optimize
3 image_generate V2 固定三步:material_input / video_prompt_optimize / video_generate
4 video_prompt_optimize version 只表示同一步骤的重建版本,不表示项目流程版本;项目版本由
5 video_generate ModuleGenerationProject.flow_version 保存。
input_json / output_json 使用 JSON/JSONB 存储。 input_json / output_json 使用 JSON/JSONB 存储。
建议结构: 建议结构:
+6
View File
@@ -25,6 +25,8 @@ class OptimizeParams(BaseModel):
class GenerateParams(BaseModel): class GenerateParams(BaseModel):
engine_id: str | None = Field(None, description="生成引擎ID;为空时优先沿用记录引擎,再回退默认引擎")
include_media_references: bool = Field(False, description="最终生成时是否携带提词阶段保存的附件")
aspect_ratio: str | None = None aspect_ratio: str | None = None
resolution: str | None = None resolution: str | None = None
image_size: str | None = None image_size: str | None = None
@@ -56,6 +58,10 @@ class GenerationRecordOut(BaseModel):
video_cover_url: str | None = None video_cover_url: str | None = None
image_url: str | None = None image_url: str | None = None
references: list[dict] | None = None references: list[dict] | None = None
engine_id: str | None = None
engine_name: str | None = None
engine_snapshot: dict | None = None
include_media_references: bool = False
text_credits_cost: float = 0.0 text_credits_cost: float = 0.0
# text_tokens_used: int = 0 # text_tokens_used: int = 0
credits_cost: float = 0.0 credits_cost: float = 0.0
@@ -404,6 +404,8 @@ class HotOpeningMaterialOut(BaseModel):
source_project_name: str | None = Field(None, description="视频素材内容项目名称") source_project_name: str | None = Field(None, description="视频素材内容项目名称")
target_project_name: str | None = Field(None, description="生成项目名称") target_project_name: str | None = Field(None, description="生成项目名称")
core_content_point: str | None = Field(None, description="生成项目核心内容点") core_content_point: str | None = Field(None, description="生成项目核心内容点")
project_description: str | None = Field(None, description="V2 可选项目描述")
video_config: dict[str, Any] | None = Field(None, description="V2 创建时保存的视频引擎、时长、比例和分辨率")
class HotOpeningImageGenerationOut(BaseModel): class HotOpeningImageGenerationOut(BaseModel):
@@ -445,6 +447,9 @@ class HotOpeningTaskDetailOut(BaseModel):
user_id: str | None = Field(None, description="所属用户ID;管理员后台排查使用") user_id: str | None = Field(None, description="所属用户ID;管理员后台排查使用")
user_name: str | None = Field(None, description="所属用户名;管理员后台排查使用") user_name: str | None = Field(None, description="所属用户名;管理员后台排查使用")
module: str = Field(..., description="模块标识,爆款开头复刻固定为 hot_opening_replicate") module: str = Field(..., description="模块标识,爆款开头复刻固定为 hot_opening_replicate")
flow_version: str = Field("v1", description="流程版本:v1=历史五步骤,v2=简化三步骤")
step_count: int = Field(5, description="当前流程步骤数")
step_io_schema_version: str | None = Field(None, description="当前流程步骤 IO Schema 版本")
title: str | None = Field(None, description="项目标题,默认取生成项目名称") title: str | None = Field(None, description="项目标题,默认取生成项目名称")
status: str = Field(..., description="总任务状态:pending=已创建,waiting_user=等待用户操作,processing=处理中,completed=完成,failed=失败,cancelled=取消") status: str = Field(..., description="总任务状态:pending=已创建,waiting_user=等待用户操作,processing=处理中,completed=完成,failed=失败,cancelled=取消")
current_step_code: str | None = Field(None, description="当前所处步骤编码") current_step_code: str | None = Field(None, description="当前所处步骤编码")
@@ -467,6 +472,8 @@ class HotOpeningTaskListItemOut(BaseModel):
user_id: str | None = Field(None, description="所属用户ID;管理员后台排查使用") user_id: str | None = Field(None, description="所属用户ID;管理员后台排查使用")
user_name: str | None = Field(None, description="所属用户名;管理员后台排查使用") user_name: str | None = Field(None, description="所属用户名;管理员后台排查使用")
module: str = Field(..., description="模块标识") module: str = Field(..., description="模块标识")
flow_version: str = Field("v1", description="流程版本")
step_count: int = Field(5, description="流程步骤数")
title: str | None = Field(None, description="项目标题") title: str | None = Field(None, description="项目标题")
status: str = Field(..., description="总任务状态:pending=已创建,waiting_user=等待用户操作,processing=处理中,completed=完成,failed=失败,cancelled=取消") status: str = Field(..., description="总任务状态:pending=已创建,waiting_user=等待用户操作,processing=处理中,completed=完成,failed=失败,cancelled=取消")
current_step_code: str | None = Field(None, description="当前步骤") current_step_code: str | None = Field(None, description="当前步骤")
@@ -0,0 +1,113 @@
from __future__ import annotations
from typing import Any
from pydantic import BaseModel, ConfigDict, Field, field_validator
class ModuleGenerationVideoConfigV2(BaseModel):
model_config = ConfigDict(extra="ignore")
engine_id: str = Field(..., min_length=1, max_length=32, description="视频引擎ID")
duration: int = Field(..., ge=1, le=120, description="视频时长,单位秒")
aspect_ratio: str = Field(..., min_length=1, max_length=16, description="视频比例")
resolution: str = Field(..., min_length=1, max_length=16, description="目标分辨率")
@field_validator("engine_id", "aspect_ratio", "resolution", mode="before")
@classmethod
def _strip_required(cls, value: str) -> str:
value = str(value or "").strip()
if not value:
raise ValueError("字段不能为空")
return value
class HotOpeningTaskCreateV2(BaseModel):
model_config = ConfigDict(extra="ignore")
material_video_url: str = Field(..., min_length=1, description="爆款参考视频,步骤2分析使用,步骤3绝不携带")
material_video_resource_id: str | None = Field(None, max_length=32)
material_video_duration_seconds: float | None = Field(None, gt=0)
material_image_url: str | None = Field(None, description="可选参考图片;仅支持本地上传/历史素材")
material_image_resource_id: str | None = Field(None, max_length=32)
source_project_name: str | None = Field(None, max_length=40)
target_project_name: str | None = Field(None, max_length=40)
core_content_point: str | None = Field(None, max_length=200)
project_description: str | None = Field(None, max_length=1000)
target_platform: str | None = Field(None, max_length=64)
video_config: ModuleGenerationVideoConfigV2
idempotency_key: str | None = Field(None, max_length=64)
@field_validator(
"material_video_url",
"material_image_url",
"source_project_name",
"target_project_name",
"core_content_point",
"project_description",
"target_platform",
"idempotency_key",
mode="before",
)
@classmethod
def _strip_text(cls, value: str | None) -> str | None:
if value is None:
return None
value = str(value).strip()
return value or None
class ShotReplicateProjectCreateV2(BaseModel):
model_config = ConfigDict(extra="ignore")
material_image_url: str | None = Field(None, description="可选参考图片;仅支持本地上传/历史素材")
material_image_resource_id: str | None = Field(None, max_length=32)
target_project_name: str | None = Field(None, max_length=40)
core_content_point: str | None = Field(None, max_length=200)
project_description: str | None = Field(None, max_length=1000)
target_platform: str | None = Field(None, max_length=64)
video_config: ModuleGenerationVideoConfigV2
idempotency_key: str | None = Field(None, max_length=64)
@field_validator(
"material_image_url",
"target_project_name",
"core_content_point",
"project_description",
"target_platform",
"idempotency_key",
mode="before",
)
@classmethod
def _strip_text(cls, value: str | None) -> str | None:
if value is None:
return None
value = str(value).strip()
return value or None
class ModuleVideoPromptRetryV2(BaseModel):
model_config = ConfigDict(extra="forbid")
video_config: ModuleGenerationVideoConfigV2
class ModuleVideoPromptSchemaUpdateV2(BaseModel):
model_config = ConfigDict(extra="ignore")
prompt_schema: dict[str, Any] = Field(..., description="允许编辑字段的 JSON patch")
@field_validator("prompt_schema")
@classmethod
def _non_empty_schema(cls, value: dict[str, Any]) -> dict[str, Any]:
if not isinstance(value, dict) or not value:
raise ValueError("prompt_schema 必须为非空对象")
return value
class ModuleGenerationV2ActionOut(BaseModel):
message: str
project_id: str
step_id: str | None = None
flow_version: str = "v2"
detail: Any | None = None
@@ -56,6 +56,10 @@ class RecentGenerationItemOut(BaseModel):
None, None,
description="关联拆镜片段ID。仅 shot_replicate 模块可能有值,来源 shot_replicate_segments.id;其他模块返回 null。", description="关联拆镜片段ID。仅 shot_replicate 模块可能有值,来源 shot_replicate_segments.id;其他模块返回 null。",
) )
module_project_flow_version: str | None = Field(
None,
description="关联模块项目流程版本。hot_opening_replicate/shot_replicate 返回 v1/v2;其他模块返回 null。",
)
module_project_id: str | None = Field( module_project_id: str | None = Field(
None, None,
description="通用模块项目ID。hot_opening_replicate/shot_replicate 模块可能有值,来源 module_generation_steps.project_idproject/chat_ai 返回 null。", description="通用模块项目ID。hot_opening_replicate/shot_replicate 模块可能有值,来源 module_generation_steps.project_idproject/chat_ai 返回 null。",
@@ -395,6 +395,8 @@ class ShotReplicateMaterialOut(BaseModel):
source_project_name: str | None = Field(None, description="视频素材内容项目名称") source_project_name: str | None = Field(None, description="视频素材内容项目名称")
target_project_name: str | None = Field(None, description="生成项目名称") target_project_name: str | None = Field(None, description="生成项目名称")
core_content_point: str | None = Field(None, description="生成项目核心内容点") core_content_point: str | None = Field(None, description="生成项目核心内容点")
project_description: str | None = Field(None, description="V2 可选项目描述")
video_config: dict[str, Any] | None = Field(None, description="V2 创建时保存的视频引擎、时长、比例和分辨率")
class ShotReplicateImageGenerationOut(BaseModel): class ShotReplicateImageGenerationOut(BaseModel):
@@ -436,6 +438,9 @@ class ShotReplicateTaskDetailOut(BaseModel):
user_id: str | None = Field(None, description="所属用户ID;管理员后台排查使用") user_id: str | None = Field(None, description="所属用户ID;管理员后台排查使用")
user_name: str | None = Field(None, description="所属用户名;管理员后台排查使用") user_name: str | None = Field(None, description="所属用户名;管理员后台排查使用")
module: str = Field(..., description="模块标识,拆镜复刻固定为 shot_replicate") module: str = Field(..., description="模块标识,拆镜复刻固定为 shot_replicate")
flow_version: str = Field("v1", description="流程版本:v1=历史五步骤,v2=简化三步骤")
step_count: int = Field(5, description="当前流程步骤数")
step_io_schema_version: str | None = Field(None, description="当前流程步骤 IO Schema 版本")
title: str | None = Field(None, description="项目标题,默认取生成项目名称") title: str | None = Field(None, description="项目标题,默认取生成项目名称")
status: str = Field(..., description="总任务状态:pending/waiting_user/processing/completed/failed/cancelled") status: str = Field(..., description="总任务状态:pending/waiting_user/processing/completed/failed/cancelled")
current_step_code: str | None = Field(None, description="当前所处步骤编码") current_step_code: str | None = Field(None, description="当前所处步骤编码")
@@ -456,6 +461,8 @@ class ShotReplicateTaskListItemOut(BaseModel):
id: str = Field(..., description="总任务项目ID。这个ID就是前端项目ID") id: str = Field(..., description="总任务项目ID。这个ID就是前端项目ID")
project_id: str = Field(..., description="兼容前端命名,等同于 id") project_id: str = Field(..., description="兼容前端命名,等同于 id")
module: str = Field(..., description="模块标识") module: str = Field(..., description="模块标识")
flow_version: str = Field("v1", description="流程版本")
step_count: int = Field(5, description="流程步骤数")
title: str | None = Field(None, description="项目标题") title: str | None = Field(None, description="项目标题")
status: str = Field(..., description="总任务状态") status: str = Field(..., description="总任务状态")
current_step_code: str | None = Field(None, description="当前步骤") current_step_code: str | None = Field(None, description="当前步骤")
@@ -702,6 +709,7 @@ class ShotSegmentOut(BaseModel):
module_project_title: str | None = Field(None, description="关联拆镜复刻项目标题;后台片段列表展示使用") module_project_title: str | None = Field(None, description="关联拆镜复刻项目标题;后台片段列表展示使用")
module_project_status: str | None = Field(None, description="关联拆镜复刻项目状态;后台片段列表展示使用") module_project_status: str | None = Field(None, description="关联拆镜复刻项目状态;后台片段列表展示使用")
module_project_current_step_code: str | None = Field(None, description="关联拆镜复刻项目当前步骤;后台片段列表展示使用") module_project_current_step_code: str | None = Field(None, description="关联拆镜复刻项目当前步骤;后台片段列表展示使用")
module_project_flow_version: str | None = Field(None, description="关联复刻项目流程版本:v1/v2")
created_at: NaiveDatetimeOptional = Field(None, description="创建时间") created_at: NaiveDatetimeOptional = Field(None, description="创建时间")
updated_at: NaiveDatetimeOptional = Field(None, description="更新时间") updated_at: NaiveDatetimeOptional = Field(None, description="更新时间")
@@ -113,7 +113,11 @@ def build_video_snapshot(engine: VideoEngine, ratio: str, resolution: str, durat
"supported_resolutions": parse_json_list(engine.supported_resolutions, []), "supported_resolutions": parse_json_list(engine.supported_resolutions, []),
"supported_durations": parse_json_list(engine.supported_durations, []), "supported_durations": parse_json_list(engine.supported_durations, []),
"max_duration": engine.max_duration, "max_duration": engine.max_duration,
"max_audio_count": engine.max_audio_count, "max_image_count": int(getattr(engine, "max_image_count", 0) or 0),
"max_video_count": int(getattr(engine, "max_video_count", 0) or 0),
"max_audio_count": int(getattr(engine, "max_audio_count", 0) or 0),
"supports_universal_reference": bool(getattr(engine, "supports_universal_reference", False)),
"supports_first_last_frame": bool(getattr(engine, "supports_first_last_frame", False)),
"multi_generation_enabled": bool(getattr(engine, "multi_generation_enabled", False)), "multi_generation_enabled": bool(getattr(engine, "multi_generation_enabled", False)),
"max_generation_count": normalize_generation_count(getattr(engine, "max_generation_count", 1)), "max_generation_count": normalize_generation_count(getattr(engine, "max_generation_count", 1)),
"selected_ratio": ratio, "selected_ratio": ratio,
@@ -10,6 +10,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.credit_record import CreditRecordBillingScene, CreditRecordChargeKind, CreditRecordOwnerType, CreditRecordSourceModule from app.enums.credit_record import CreditRecordBillingScene, CreditRecordChargeKind, CreditRecordOwnerType, CreditRecordSourceModule
from app.models.credit_record import CreditRecord from app.models.credit_record import CreditRecord
from app.models.generation_record import GenerationRecord from app.models.generation_record import GenerationRecord
from app.services.generation.media_reference_service import calculate_media_reference_usage
from app.models.module_generation_step import ModuleGenerationStep from app.models.module_generation_step import ModuleGenerationStep
from app.models.token_usage import TokenUsage from app.models.token_usage import TokenUsage
from app.models.system_config import SystemConfig from app.models.system_config import SystemConfig
@@ -461,6 +462,7 @@ async def charge_generation_media_by_params(
resolution: str | None = None, resolution: str | None = None,
engine_id: str | None = None, engine_id: str | None = None,
input_video_duration: float | None = None, input_video_duration: float | None = None,
input_image_count: int | None = None,
project_name: str | None = None, project_name: str | None = None,
description_prefix: str = "AI创作-", description_prefix: str = "AI创作-",
owner_type: str = OWNER_CHAT_GENERATION_TASK, owner_type: str = OWNER_CHAT_GENERATION_TASK,
@@ -510,7 +512,9 @@ async def charge_generation_media_by_params(
if gen_type == "image": if gen_type == "image":
size = image_size or "2K" size = image_size or "2K"
unit_amount = await calc_image_credits(db, size, engine_id=engine_id) unit_amount = await calc_image_credits(
db, size, engine_id=engine_id, input_image_count=input_image_count
)
amount = round(unit_amount * quantity, 2) amount = round(unit_amount * quantity, 2)
items.append( items.append(
await deduct_credits_locked_once( await deduct_credits_locked_once(
@@ -530,6 +534,7 @@ async def charge_generation_media_by_params(
db, duration or 5, resolution or "720p", db, duration or 5, resolution or "720p",
engine_id=engine_id, engine_id=engine_id,
input_video_duration=input_video_duration, input_video_duration=input_video_duration,
input_image_count=input_image_count,
) )
amount = round(unit_amount * quantity, 2) amount = round(unit_amount * quantity, 2)
items.append( items.append(
@@ -560,6 +565,10 @@ async def charge_generation_media_for_record(
attempt_no: int | None = None, attempt_no: int | None = None,
engine_id: str | None = None, engine_id: str | None = None,
) -> BillingSummary: ) -> BillingSummary:
reference_usage = calculate_media_reference_usage(
record.media_references,
include=bool(record.include_media_references),
)
return await charge_generation_media_by_params( return await charge_generation_media_by_params(
db, db,
user_id=record.user_id, user_id=record.user_id,
@@ -569,6 +578,8 @@ async def charge_generation_media_for_record(
duration=record.duration, duration=record.duration,
resolution=record.resolution, resolution=record.resolution,
engine_id=engine_id or getattr(record, "engine_id", None), engine_id=engine_id or getattr(record, "engine_id", None),
input_video_duration=reference_usage.input_video_duration or None,
input_image_count=reference_usage.image_count or None,
project_name=project_name, project_name=project_name,
description_prefix=description_prefix, description_prefix=description_prefix,
owner_type=OWNER_GENERATION_RECORD, owner_type=OWNER_GENERATION_RECORD,
@@ -0,0 +1,148 @@
from __future__ import annotations
import json
from dataclasses import dataclass
from typing import Any, Iterable
from fastapi import HTTPException
@dataclass(frozen=True, slots=True)
class MediaReferenceUsage:
image_count: int = 0
video_count: int = 0
audio_count: int = 0
input_video_duration: float = 0.0
input_audio_duration: float = 0.0
def parse_media_references(value: str | list[dict] | None) -> list[dict]:
if not value:
return []
data: Any = value
if isinstance(value, str):
try:
data = json.loads(value)
except (TypeError, json.JSONDecodeError):
return []
if not isinstance(data, list):
return []
return [item for item in data if isinstance(item, dict)]
def _reference_type(item: dict) -> str:
value = str(item.get("type") or item.get("media_type") or "").strip().lower()
if value in {"image", "video", "audio"}:
return value
mime = str(item.get("mime_type") or item.get("content_type") or "").lower()
if mime.startswith("image/"):
return "image"
if mime.startswith("video/"):
return "video"
if mime.startswith("audio/"):
return "audio"
url = str(item.get("url") or item.get("file_url") or "").lower().split("?", 1)[0]
if url.endswith((".png", ".jpg", ".jpeg", ".webp", ".gif", ".bmp")):
return "image"
if url.endswith((".mp4", ".mov", ".webm", ".mkv", ".avi")):
return "video"
if url.endswith((".mp3", ".wav", ".m4a", ".aac", ".ogg", ".flac")):
return "audio"
return ""
def _duration(item: dict) -> float:
for key in ("duration", "duration_seconds", "video_duration", "audio_duration"):
try:
value = float(item.get(key) or 0)
except (TypeError, ValueError):
continue
if value > 0:
return value
return 0.0
def calculate_media_reference_usage(
references: str | list[dict] | None,
*,
include: bool,
) -> MediaReferenceUsage:
if not include:
return MediaReferenceUsage()
image_count = video_count = audio_count = 0
video_duration = audio_duration = 0.0
for item in parse_media_references(references):
media_type = _reference_type(item)
if media_type == "image":
image_count += 1
elif media_type == "video":
video_count += 1
video_duration += _duration(item)
elif media_type == "audio":
audio_count += 1
audio_duration += _duration(item)
return MediaReferenceUsage(
image_count=image_count,
video_count=video_count,
audio_count=audio_count,
input_video_duration=round(video_duration, 3),
input_audio_duration=round(audio_duration, 3),
)
def filter_references_by_type(
references: str | list[dict] | None,
*,
allowed_types: Iterable[str],
max_count: int | None = None,
) -> list[dict]:
allowed = {str(item).lower() for item in allowed_types}
result = [item for item in parse_media_references(references) if _reference_type(item) in allowed]
if max_count is not None:
return result[: max(0, int(max_count))]
return result
def validate_media_reference_usage_for_engine(
usage: MediaReferenceUsage,
*,
gen_type: str,
engine: Any,
) -> None:
"""按所选引擎能力校验最终实际会发送的附件。"""
normalized = str(gen_type or "").strip().lower()
if normalized == "video":
if not bool(getattr(engine, "supports_universal_reference", True)) and (
usage.image_count or usage.video_count or usage.audio_count
):
raise HTTPException(status_code=400, detail="当前视频引擎不支持参考附件")
limits = {
"图片": int(getattr(engine, "max_image_count", 0) or 0),
"视频": int(getattr(engine, "max_video_count", 0) or 0),
"音频": int(getattr(engine, "max_audio_count", 0) or 0),
}
counts = {"图片": usage.image_count, "视频": usage.video_count, "音频": usage.audio_count}
for label, count in counts.items():
limit = limits[label]
if count > limit:
raise HTTPException(
status_code=400,
detail=f"当前视频引擎最多支持 {limit}{label}附件,当前为 {count}",
)
return
if normalized == "image":
if usage.video_count or usage.audio_count:
raise HTTPException(status_code=400, detail="图片生成只能携带图片附件")
limit = int(
getattr(engine, "max_reference_image_count", None)
if getattr(engine, "max_reference_image_count", None) is not None
else getattr(engine, "max_image_count", 0)
or 0
)
if usage.image_count > limit:
raise HTTPException(
status_code=400,
detail=f"当前图片引擎最多支持 {limit} 张参考图,当前为 {usage.image_count}",
)
return
raise HTTPException(status_code=400, detail="不支持的生成类型")
@@ -58,6 +58,18 @@ def owner_type_of(owner: GenerationOwner) -> str:
raise TypeError(f"不支持的生成任务对象: {type(owner)!r}") raise TypeError(f"不支持的生成任务对象: {type(owner)!r}")
def owner_include_media_references(owner: GenerationOwner) -> bool:
"""返回本次供应商创建是否应携带附件。
ChatGenerationTask 延续原有行为;GenerationRecord 使用用户提交并持久化的开关。
"""
if isinstance(owner, ChatGenerationTask):
return True
if isinstance(owner, GenerationRecord):
return bool(owner.include_media_references)
raise TypeError(f"不支持的生成任务对象: {type(owner)!r}")
def owner_mode(owner: GenerationOwner) -> str: def owner_mode(owner: GenerationOwner) -> str:
if isinstance(owner, ChatGenerationTask): if isinstance(owner, ChatGenerationTask):
return str(owner.generation_mode or GenerationMode.CHATAPI_ASYNC.value) return str(owner.generation_mode or GenerationMode.CHATAPI_ASYNC.value)
@@ -10,8 +10,11 @@ from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask from app.services.generation.pipeline.owner_service import (
from app.services.generation.pipeline.owner_service import GenerationOwner, owner_provider_task_id GenerationOwner,
owner_include_media_references,
owner_provider_task_id,
)
from app.models.image_engine import ImageEngine from app.models.image_engine import ImageEngine
from app.models.video_engine import VideoEngine from app.models.video_engine import VideoEngine
from app.services.generation.log_service import log_provider_call from app.services.generation.log_service import log_provider_call
@@ -105,7 +108,7 @@ async def _create_video_task(db: AsyncSession, task: GenerationOwner) -> dict:
started = time.perf_counter() started = time.perf_counter()
async with provider_limit("ark_video_create", settings.ARK_VIDEO_CREATE_MAX_CONCURRENCY): async with provider_limit("ark_video_create", settings.ARK_VIDEO_CREATE_MAX_CONCURRENCY):
try: try:
provider_task_id = await submit_video_task(None, engine, task, include_media_references=isinstance(task, ChatGenerationTask)) provider_task_id = await submit_video_task(None, engine, task, include_media_references=owner_include_media_references(task))
response = {"task_id": provider_task_id} response = {"task_id": provider_task_id}
await log_provider_call( await log_provider_call(
task, task,
@@ -169,7 +172,7 @@ async def create_image_sync_batch_result_with_engine(
None, None,
engine, engine,
task, task,
include_media_references=isinstance(task, ChatGenerationTask), include_media_references=owner_include_media_references(task),
generation_count=count, generation_count=count,
) )
response_data = result.get("response_data") or result response_data = result.get("response_data") or result
@@ -606,6 +606,8 @@ async def recover_one_generation_task(
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value, ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
): ):
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
# 刷新更新时间形成创建队列保护窗口,避免 Beat 在任务尚未消费时每轮重复补投。
task.updated_at = current_time
# Release the recovery row lock before writing an event through the # Release the recovery row lock before writing an event through the
# independent logging session or talking to the broker. # independent logging session or talking to the broker.
await db.commit() await db.commit()
@@ -622,12 +624,17 @@ async def recover_one_generation_task(
kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)}, kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
queue=CeleryQueue.GEN_CHATAPI_CREATE.value, queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
countdown=0, countdown=0,
task_id=(
f"generation-create:{GenerationOwnerType.CHAT_GENERATION_TASK.value}:"
f"{task.id}:attempt:{int(task.generation_attempt_no or 1)}"
),
) )
return "recover_create_no_remote_no_provider_before_deadline" return "recover_create_no_remote_no_provider_before_deadline"
# result_ready 但没有 URL 是脏状态;未过 deadline 时回创建队列重新处理,过期上面已标记超时。 # result_ready 但没有 URL 是脏状态;未过 deadline 时回创建队列重新处理,过期上面已标记超时。
if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value: if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value:
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
task.updated_at = current_time
await db.commit() await db.commit()
await _remove_poll_active(_chat_registry_id(task)) await _remove_poll_active(_chat_registry_id(task))
await log_task_event( await log_task_event(
@@ -641,6 +648,10 @@ async def recover_one_generation_task(
kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)}, kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
queue=CeleryQueue.GEN_CHATAPI_CREATE.value, queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
countdown=0, countdown=0,
task_id=(
f"generation-create:{GenerationOwnerType.CHAT_GENERATION_TASK.value}:"
f"{task.id}:attempt:{int(task.generation_attempt_no or 1)}"
),
) )
return "recover_create_result_ready_no_url_before_deadline" return "recover_create_result_ready_no_url_before_deadline"
@@ -887,6 +898,61 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
"results": results, "results": results,
} }
async def recover_stale_create_tasks_once(db: AsyncSession) -> dict[str, Any]:
"""轻量恢复长时间未消费的创建阶段 ChatGenerationTask。
只扫描 queued/preparing/creating_provider_task,避免周期任务重复执行完整
provider poll、下载和主任务汇总逻辑。
"""
current_time = _now()
cutoff = current_time - timedelta(
seconds=max(1, int(settings.GENERATION_CREATE_QUEUE_TIMEOUT_SECONDS or 300))
)
lease_expired_at = current_time
batch_size = max(1, int(settings.GENERATION_RECOVERY_BATCH_SIZE or 20))
result = await db.execute(
select(ChatGenerationTask.id)
.where(
ChatGenerationTask.deleted_at.is_(None),
ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)),
ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value,
ChatGenerationTask.pipeline_stage.in_(
[
ChatGenerationPipelineStage.QUEUED.value,
ChatGenerationPipelineStage.PREPARING.value,
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
]
),
ChatGenerationTask.remote_result_url.is_(None),
ChatGenerationTask.provider_task_id.is_(None),
ChatGenerationTask.seedance_task_id.is_(None),
ChatGenerationTask.updated_at <= cutoff,
(
ChatGenerationTask.provider_create_lease_until.is_(None)
| (ChatGenerationTask.provider_create_lease_until <= lease_expired_at)
),
)
.order_by(ChatGenerationTask.updated_at.asc(), ChatGenerationTask.id.asc())
.limit(batch_size)
)
task_ids = [str(value) for value in result.scalars().all()]
counts: dict[str, int] = {}
for task_id in task_ids:
task = await _load_chat_task_for_update(db, task_id)
if task is None:
await db.rollback()
action = "skip_missing_task"
else:
action = await recover_one_generation_task(
db,
task,
payload=None,
source="periodic_create_recovery",
)
counts[action] = counts.get(action, 0) + 1
return {"checked": len(task_ids), "results": counts}
async def dispatch_due_poll_tasks_once(db: AsyncSession) -> dict[str, Any]: async def dispatch_due_poll_tasks_once(db: AsyncSession) -> dict[str, Any]:
"""周期性轻量到期轮询调度。 """周期性轻量到期轮询调度。
@@ -26,6 +26,10 @@ from app.services.generation.ai.engine_service import (
parse_json_list as _parse_list, parse_json_list as _parse_list,
) )
from app.services.generation.billing_service import OWNER_CHAT_GENERATION_TASK, charge_generation_media_by_params from app.services.generation.billing_service import OWNER_CHAT_GENERATION_TASK, charge_generation_media_by_params
from app.services.generation.media_reference_service import (
calculate_media_reference_usage,
validate_media_reference_usage_for_engine,
)
from app.services.resource_capacity_service import assert_user_resource_capacity_available from app.services.resource_capacity_service import assert_user_resource_capacity_available
from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot
from app.services.private_portrait.reference_resolver import resolve_private_portrait_references from app.services.private_portrait.reference_resolver import resolve_private_portrait_references
@@ -111,6 +115,8 @@ async def create_chat_generation_task_for_module(
raise HTTPException(status_code=400, detail=f"图片比例不支持: {proportion}") raise HTTPException(status_code=400, detail=f"图片比例不支持: {proportion}")
px = px or normalize_px((sizes.get(size) or {}).get(proportion)) px = px or normalize_px((sizes.get(size) or {}).get(proportion))
px = px or IMAGE_DEFAULT_PX px = px or IMAGE_DEFAULT_PX
reference_usage = calculate_media_reference_usage(refs, include=True)
validate_media_reference_usage_for_engine(reference_usage, gen_type="image", engine=engine)
media_billing = await charge_generation_media_by_params( media_billing = await charge_generation_media_by_params(
db, db,
user_id=current_user.id, user_id=current_user.id,
@@ -118,6 +124,7 @@ async def create_chat_generation_task_for_module(
gen_type="image", gen_type="image",
image_size=size, image_size=size,
engine_id=engine.id, engine_id=engine.id,
input_image_count=reference_usage.image_count or None,
project_name=billing_project_name, project_name=billing_project_name,
description_prefix=billing_description_prefix, description_prefix=billing_description_prefix,
owner_type=OWNER_CHAT_GENERATION_TASK, owner_type=OWNER_CHAT_GENERATION_TASK,
@@ -174,6 +181,8 @@ async def create_chat_generation_task_for_module(
aspect_ratio=ratio, aspect_ratio=ratio,
supported_provider_resolutions=resolutions, supported_provider_resolutions=resolutions,
) )
reference_usage = calculate_media_reference_usage(refs, include=True)
validate_media_reference_usage_for_engine(reference_usage, gen_type="video", engine=engine)
media_billing = await charge_generation_media_by_params( media_billing = await charge_generation_media_by_params(
db, db,
user_id=current_user.id, user_id=current_user.id,
@@ -182,6 +191,8 @@ async def create_chat_generation_task_for_module(
duration=selected_duration, duration=selected_duration,
resolution=selected_resolution, resolution=selected_resolution,
engine_id=engine.id, engine_id=engine.id,
input_video_duration=reference_usage.input_video_duration or None,
input_image_count=reference_usage.image_count or None,
project_name=billing_project_name, project_name=billing_project_name,
description_prefix=billing_description_prefix, description_prefix=billing_description_prefix,
owner_type=OWNER_CHAT_GENERATION_TASK, owner_type=OWNER_CHAT_GENERATION_TASK,
@@ -117,6 +117,7 @@ FLOW_CONFIG = ModuleGenerationFlowConfig(
cancel_chat_task_error_message="爆款开头复刻步骤被重新生成或删除,旧生成任务已取消", cancel_chat_task_error_message="爆款开头复刻步骤被重新生成或删除,旧生成任务已取消",
material_video_url_editable=True, material_video_url_editable=True,
step_io_schema_version=STEP_IO_SCHEMA_VERSION, step_io_schema_version=STEP_IO_SCHEMA_VERSION,
expected_flow_version="v1",
) )
@@ -422,12 +423,19 @@ async def project_to_detail_out(db: AsyncSession, project: ModuleGenerationProje
user_result = await db.execute(select(User.username).where(User.id == project.user_id).limit(1)) user_result = await db.execute(select(User.username).where(User.id == project.user_id).limit(1))
user_name = user_result.scalar_one_or_none() user_name = user_result.scalar_one_or_none()
flow_version = str(getattr(project, "flow_version", None) or "v1")
video_prompt_config = dict(video_prompt_input.get("video_config") or {})
video_prompt_engine_snapshot = dict(video_prompt_config.get("engine_snapshot") or {})
return HotOpeningTaskDetailOut( return HotOpeningTaskDetailOut(
id=project.id, id=project.id,
project_id=project.id, project_id=project.id,
user_id=project.user_id, user_id=project.user_id,
user_name=user_name, user_name=user_name,
module=project.module, module=project.module,
flow_version=flow_version,
step_count=3 if flow_version == "v2" else 5,
step_io_schema_version=("hot_opening_step_io_v2" if project.module == "hot_opening_replicate" else "shot_replicate_step_io_v2") if flow_version == "v2" else STEP_IO_SCHEMA_VERSION,
title=project.title, title=project.title,
status=project.status, status=project.status,
current_step_code=project.current_step_code, current_step_code=project.current_step_code,
@@ -442,6 +450,8 @@ async def project_to_detail_out(db: AsyncSession, project: ModuleGenerationProje
source_project_name=material_input.get("source_project_name"), source_project_name=material_input.get("source_project_name"),
target_project_name=material_input.get("target_project_name"), target_project_name=material_input.get("target_project_name"),
core_content_point=material_input.get("core_content_point"), core_content_point=material_input.get("core_content_point"),
project_description=material_input.get("project_description"),
video_config=None if flow_version == "v2" else material_input.get("video_config"),
), ),
image_generation=HotOpeningImageGenerationOut( image_generation=HotOpeningImageGenerationOut(
prompt_step_id=image_prompt_step.id if image_prompt_step else None, prompt_step_id=image_prompt_step.id if image_prompt_step else None,
@@ -465,9 +475,9 @@ async def project_to_detail_out(db: AsyncSession, project: ModuleGenerationProje
schema_config_source=schema_config_source, schema_config_source=schema_config_source,
schema_config_version=schema_config_version, schema_config_version=schema_config_version,
schema_config_is_fallback=schema_config_is_fallback, schema_config_is_fallback=schema_config_is_fallback,
engine_id=video_snapshot.get("id") or video_generate_input.get("engine_id"), engine_id=video_snapshot.get("id") or video_generate_input.get("engine_id") or video_prompt_config.get("engine_id"),
engine_name=video_snapshot.get("name") or video_generate_input.get("engine_name"), engine_name=video_snapshot.get("name") or video_generate_input.get("engine_name") or video_prompt_engine_snapshot.get("name"),
params=video_generate_input.get("params") or video_generate_input, params=video_generate_input.get("params") or video_prompt_config or video_generate_input,
chat_task_id=video_generate_step.chat_task_id if video_generate_step else None, chat_task_id=video_generate_step.chat_task_id if video_generate_step else None,
status=video_chat.status if video_chat else (video_generate_step.status if video_generate_step else None), status=video_chat.status if video_chat else (video_generate_step.status if video_generate_step else None),
result_video_url=build_resource_signed_url(video_url) if video_url else None, result_video_url=build_resource_signed_url(video_url) if video_url else None,
@@ -660,6 +670,8 @@ async def list_hot_opening_projects(
user_id=project.user_id, user_id=project.user_id,
user_name=user_name_map.get(project.user_id) if current_user.is_admin and project.user_id else None, user_name=user_name_map.get(project.user_id) if current_user.is_admin and project.user_id else None,
module=project.module, module=project.module,
flow_version=str(getattr(project, "flow_version", None) or "v1"),
step_count=3 if str(getattr(project, "flow_version", None) or "v1") == "v2" else 5,
title=project.title, title=project.title,
status=project.status, status=project.status,
current_step_code=project.current_step_code, current_step_code=project.current_step_code,
@@ -1532,6 +1544,9 @@ async def generate_video_from_prompt(
async def handle_chat_generation_task_completed(db: AsyncSession, task: ChatGenerationTask) -> None: async def handle_chat_generation_task_completed(db: AsyncSession, task: ChatGenerationTask) -> None:
if not task or task.generation_mode != GENERATION_MODE: if not task or task.generation_mode != GENERATION_MODE:
return return
from app.services.module_generation_v2.flow_service import handle_chat_generation_task_finished_v2
if await handle_chat_generation_task_finished_v2(db, task=task):
return
meta_result = await db.execute( meta_result = await db.execute(
select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where( select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where(
ModuleGenerationStep.chat_task_id == task.id, ModuleGenerationStep.chat_task_id == task.id,
@@ -1614,6 +1629,9 @@ async def handle_chat_generation_task_completed(db: AsyncSession, task: ChatGene
async def handle_chat_generation_task_failed(db: AsyncSession, task: ChatGenerationTask) -> None: async def handle_chat_generation_task_failed(db: AsyncSession, task: ChatGenerationTask) -> None:
if not task or task.generation_mode != GENERATION_MODE: if not task or task.generation_mode != GENERATION_MODE:
return return
from app.services.module_generation_v2.flow_service import handle_chat_generation_task_finished_v2
if await handle_chat_generation_task_finished_v2(db, task=task):
return
meta_result = await db.execute( meta_result = await db.execute(
select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where( select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where(
ModuleGenerationStep.chat_task_id == task.id, ModuleGenerationStep.chat_task_id == task.id,
@@ -1527,7 +1527,7 @@ async def optimize_hot_opening_video_prompt(
target_project_name: str, target_project_name: str,
core_content_point: str, core_content_point: str,
material_video_url: str, material_video_url: str,
generated_image_url: str, generated_image_url: str | None,
video_config: dict[str, Any], video_config: dict[str, Any],
target_platform: str = "抖音", target_platform: str = "抖音",
schema_config_snapshot: Any | None = None, schema_config_snapshot: Any | None = None,
@@ -1538,16 +1538,19 @@ async def optimize_hot_opening_video_prompt(
) -> tuple[dict[str, Any], str, dict[str, Any]]: ) -> tuple[dict[str, Any], str, dict[str, Any]]:
duration = int(video_config["duration"]) duration = int(video_config["duration"])
from app.utils.media import media_to_base64, get_llm_media_as_base64 from app.utils.media import media_to_base64, get_llm_media_as_base64
if await get_llm_media_as_base64(db): use_base64 = await get_llm_media_as_base64(db)
if use_base64:
video_url_final = await media_to_base64(material_video_url, "video/mp4") video_url_final = await media_to_base64(material_video_url, "video/mp4")
image_url_final = await media_to_base64(generated_image_url, "image/png")
else: else:
video_url_final = _build_file_url_or_data_uri(material_video_url) video_url_final = _build_file_url_or_data_uri(material_video_url)
image_url_final = _build_file_url_or_data_uri(generated_image_url) references = [{"type": "video", "url": video_url_final}]
references = [ if generated_image_url:
{"type": "video", "url": video_url_final}, image_url_final = (
{"type": "image", "url": image_url_final}, await media_to_base64(generated_image_url, "image/png")
] if use_base64
else _build_file_url_or_data_uri(generated_image_url)
)
references.append({"type": "image", "url": image_url_final})
client_schema = build_dynamic_schema(video_config, schema_config_snapshot) client_schema = build_dynamic_schema(video_config, schema_config_snapshot)
reference_video_fps = int(video_config.get("reference_video_fps") or DEFAULT_REFERENCE_VIDEO_FPS) reference_video_fps = int(video_config.get("reference_video_fps") or DEFAULT_REFERENCE_VIDEO_FPS)
@@ -1690,7 +1693,7 @@ async def optimize_hot_opening_video_prompt(
"input_tokens": int(usage.get("prompt_tokens") or 0), "input_tokens": int(usage.get("prompt_tokens") or 0),
"output_tokens": int(usage.get("completion_tokens") or 0), "output_tokens": int(usage.get("completion_tokens") or 0),
"total_tokens": int(usage.get("total_tokens") or 0), "total_tokens": int(usage.get("total_tokens") or 0),
"log_user_message": log_user_message, # "log_user_message": log_user_message,
} }
token_usage_id = generate_id() token_usage_id = generate_id()
db.add( db.add(
@@ -17,6 +17,7 @@ from app.enums.shot_replicate import (
ShotSegmentAnalysisStatusEnum, ShotSegmentAnalysisStatusEnum,
ShotSplitStatusEnum, ShotSplitStatusEnum,
) )
from app.models.module_generation_project import ModuleGenerationProject
from app.models.module_generation_step import ModuleGenerationStep from app.models.module_generation_step import ModuleGenerationStep
from app.models.shot_replicate_segment import ShotReplicateSegment from app.models.shot_replicate_segment import ShotReplicateSegment
from app.models.shot_replicate_task_set import ShotReplicateTaskSet from app.models.shot_replicate_task_set import ShotReplicateTaskSet
@@ -47,6 +48,7 @@ TASK_HOT_IMAGE_PROMPT = "hot_opening.start_image_prompt_optimize"
TASK_HOT_VIDEO_PROMPT = "hot_opening.start_video_prompt_optimize" TASK_HOT_VIDEO_PROMPT = "hot_opening.start_video_prompt_optimize"
TASK_SHOT_IMAGE_PROMPT = "shot_replicate.start_image_prompt_optimize" TASK_SHOT_IMAGE_PROMPT = "shot_replicate.start_image_prompt_optimize"
TASK_SHOT_VIDEO_PROMPT = "shot_replicate.start_video_prompt_optimize" TASK_SHOT_VIDEO_PROMPT = "shot_replicate.start_video_prompt_optimize"
TASK_MODULE_V2_VIDEO_PROMPT = "module_generation_v2.start_video_prompt_optimize"
TASK_SHOT_ANALYZE_ORIGINAL = "shot_replicate.analyze_original_video" TASK_SHOT_ANALYZE_ORIGINAL = "shot_replicate.analyze_original_video"
TASK_SHOT_ANALYZE_CUSTOM_SEGMENT = "shot_replicate.analyze_custom_segment_video" TASK_SHOT_ANALYZE_CUSTOM_SEGMENT = "shot_replicate.analyze_custom_segment_video"
TASK_SHOT_SPLIT_ONE = "shot_replicate.split_one_segment" TASK_SHOT_SPLIT_ONE = "shot_replicate.split_one_segment"
@@ -511,9 +513,24 @@ async def _recover_stale_module_steps(db: AsyncSession, *, limit: int) -> dict[s
.with_for_update(skip_locked=True) .with_for_update(skip_locked=True)
) )
steps = list(result.scalars().all()) steps = list(result.scalars().all())
project_flow_map: dict[str, str] = {}
project_ids = list({step.project_id for step in steps})
if project_ids:
project_result = await db.execute(
select(ModuleGenerationProject.id, ModuleGenerationProject.flow_version).where(
ModuleGenerationProject.id.in_(project_ids),
ModuleGenerationProject.deleted_at.is_(None),
)
)
project_flow_map = {
str(project_id): str(flow_version or "v1")
for project_id, flow_version in project_result.all()
}
results: dict[str, int] = {} results: dict[str, int] = {}
for step in steps: for step in steps:
if step.module == HOT_MODULE and step.step_code == HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value: if project_flow_map.get(step.project_id, "v1") == "v2" and step.step_code == HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value:
task_name = TASK_MODULE_V2_VIDEO_PROMPT
elif step.module == HOT_MODULE and step.step_code == HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value:
task_name = TASK_HOT_IMAGE_PROMPT task_name = TASK_HOT_IMAGE_PROMPT
elif step.module == HOT_MODULE and step.step_code == HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value: elif step.module == HOT_MODULE and step.step_code == HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value:
task_name = TASK_HOT_VIDEO_PROMPT task_name = TASK_HOT_VIDEO_PROMPT
@@ -59,6 +59,18 @@ async def get_project_for_user(
project = result.scalar_one_or_none() project = result.scalar_one_or_none()
if not project: if not project:
raise HTTPException(status_code=404, detail=config.project_not_found_message) raise HTTPException(status_code=404, detail=config.project_not_found_message)
if config.expected_flow_version:
actual_flow_version = str(getattr(project, "flow_version", None) or "v1")
if actual_flow_version != config.expected_flow_version:
raise HTTPException(
status_code=409,
detail={
"message": "项目流程版本与当前接口版本不匹配",
"project_id": project.id,
"flow_version": actual_flow_version,
"expected_flow_version": config.expected_flow_version,
},
)
return project return project
@@ -81,7 +81,19 @@ def build_step_output(
def is_wrapped_step_io(value: Any, *, schema_version: str = STEP_IO_SCHEMA_VERSION) -> bool: def is_wrapped_step_io(value: Any, *, schema_version: str = STEP_IO_SCHEMA_VERSION) -> bool:
return isinstance(value, dict) and value.get("schema_version") == schema_version if not isinstance(value, dict):
return False
current = str(value.get("schema_version") or "")
if current == schema_version:
return True
# V1/V2 查询接口共用同一组解析器。只要结构符合统一步骤 IO 包装,
# 即按包装结构解析,避免 V2 被当成普通 payload。
return bool(
current
and value.get("step_code")
and "payload" in value
and ("status" in value or "source" in value)
)
def step_payload(value: Any, *, schema_version: str = STEP_IO_SCHEMA_VERSION) -> dict[str, Any]: def step_payload(value: Any, *, schema_version: str = STEP_IO_SCHEMA_VERSION) -> dict[str, Any]:
@@ -0,0 +1 @@
"""爆款复刻/拆镜复刻 V2 三步骤公共流程。"""
@@ -0,0 +1,75 @@
from __future__ import annotations
from dataclasses import dataclass
from app.enums.hot_opening_replicate import (
HotOpeningGenerationModeEnum,
HotOpeningStepIOSchemaVersionEnum,
ModuleCodeEnum as HotModuleCodeEnum,
)
from app.enums.shot_replicate import (
ModuleCodeEnum as ShotModuleCodeEnum,
ShotReplicateGenerationModeEnum,
ShotReplicateStepIOSchemaVersionEnum,
)
from app.enums.module_generation_flow import ModuleGenerationFlowConfig
MATERIAL_INPUT = "material_input"
VIDEO_PROMPT_OPTIMIZE = "video_prompt_optimize"
VIDEO_GENERATE = "video_generate"
@dataclass(frozen=True, slots=True)
class ModuleGenerationV2Config:
module: str
generation_mode: str
io_schema_version: str
display_name: str
project_not_found_message: str
material_video_locked: bool
@property
def flow_config(self) -> ModuleGenerationFlowConfig:
return ModuleGenerationFlowConfig(
module=self.module,
step_index_map={
MATERIAL_INPUT: 1,
VIDEO_PROMPT_OPTIMIZE: 2,
VIDEO_GENERATE: 3,
},
material_step_code=MATERIAL_INPUT,
image_prompt_step_code="__v2_no_image_prompt__",
image_generate_step_code="__v2_no_image_generate__",
video_prompt_step_code=VIDEO_PROMPT_OPTIMIZE,
video_generate_step_code=VIDEO_GENERATE,
project_not_found_message=self.project_not_found_message,
step_not_found_message="V2 子任务不存在",
cancel_chat_task_error_message=f"{self.display_name}步骤被重建或删除,旧生成任务已取消",
material_video_url_editable=not self.material_video_locked,
step_io_schema_version=self.io_schema_version,
expected_flow_version="v2",
)
HOT_OPENING_V2 = ModuleGenerationV2Config(
module=HotModuleCodeEnum.HOT_OPENING_REPLICATE.value,
generation_mode=HotOpeningGenerationModeEnum.HOT_OPENING_REPLICATE.value,
io_schema_version=HotOpeningStepIOSchemaVersionEnum.V2.value,
display_name="爆款开头复刻",
project_not_found_message="爆款开头复刻 V2 项目不存在",
material_video_locked=True,
)
SHOT_REPLICATE_V2 = ModuleGenerationV2Config(
module=ShotModuleCodeEnum.SHOT_REPLICATE.value,
generation_mode=ShotReplicateGenerationModeEnum.SHOT_REPLICATE.value,
io_schema_version=ShotReplicateStepIOSchemaVersionEnum.V2.value,
display_name="拆镜复刻",
project_not_found_message="拆镜复刻 V2 项目不存在",
material_video_locked=True,
)
CONFIG_BY_MODULE = {
HOT_OPENING_V2.module: HOT_OPENING_V2,
SHOT_REPLICATE_V2.module: SHOT_REPLICATE_V2,
}
@@ -0,0 +1,105 @@
from __future__ import annotations
from dataclasses import dataclass
from fastapi import HTTPException
from app.enums.celery_queue import CeleryQueue
from app.enums.common import ModuleEventTypeEnum
from app.services.module_async_recovery_service import (
TASK_MODULE_V2_VIDEO_PROMPT,
register_module_step_task,
)
from app.services.module_generation_log_service import log_module_error, log_module_event_file
from app.services.module_generation_v2.config import VIDEO_PROMPT_OPTIMIZE, ModuleGenerationV2Config
from app.tasks.celery_app import celery_app
from app.tasks.module_generation_v2_tasks import start_video_prompt_optimize_v2
@dataclass(frozen=True)
class VideoPromptDispatchResult:
registry_success: bool
celery_success: bool
registry_error: str | None = None
celery_error: str | None = None
@property
def recoverable(self) -> bool:
return self.registry_success or self.celery_success
def ensure_v2_celery_enabled() -> None:
if celery_app is None:
raise HTTPException(status_code=503, detail="Celery未启用")
async def dispatch_video_prompt_v2(
*,
config: ModuleGenerationV2Config,
project_id: str,
step_id: str,
) -> VideoPromptDispatchResult:
"""注册并投递 V2 视频提词任务。
Redis 注册成功但 Celery 直投失败时由周期恢复任务补投Celery 成功但
Redis 注册失败时任务仍可正常执行只有两个通道都失败时由 API 补偿落库为失败
"""
registry_error: Exception | None = None
try:
await register_module_step_task(
module=config.module,
project_id=project_id,
step_id=step_id,
step_code=VIDEO_PROMPT_OPTIMIZE,
task_name=TASK_MODULE_V2_VIDEO_PROMPT,
)
except Exception as exc:
registry_error = exc
log_module_error(
module=config.module,
event_type=ModuleEventTypeEnum.V2_VIDEO_PROMPT_REGISTRY_FAILED.value,
project_id=project_id,
step_id=step_id,
message="V2 视频提词 Redis 活跃注册失败,将继续尝试 Celery 直投",
exc=exc,
)
celery_error: Exception | None = None
try:
start_video_prompt_optimize_v2.apply_async(
args=[project_id, step_id],
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
countdown=0,
task_id=f"module-v2-video-prompt:{step_id}",
)
except Exception as exc:
celery_error = exc
log_module_error(
module=config.module,
event_type=ModuleEventTypeEnum.V2_VIDEO_PROMPT_DISPATCH_FAILED.value,
project_id=project_id,
step_id=step_id,
message="V2 视频提词 Celery 投递失败",
detail={"redis_registry_available": registry_error is None},
exc=exc,
)
result = VideoPromptDispatchResult(
registry_success=registry_error is None,
celery_success=celery_error is None,
registry_error=str(registry_error) if registry_error else None,
celery_error=str(celery_error) if celery_error else None,
)
if result.celery_success:
log_module_event_file(
module=config.module,
event_type=ModuleEventTypeEnum.V2_VIDEO_PROMPT_DISPATCHED.value,
project_id=project_id,
step_id=step_id,
message="V2 视频提词任务已投递",
detail={
"queue": CeleryQueue.GEN_CHATAPI_CREATE.value,
"redis_registry_available": result.registry_success,
},
)
return result
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,35 @@
from __future__ import annotations
from typing import Any
from fastapi import HTTPException
from app.services.module_generation_step_common_service import build_file_url_or_data_uri
def build_v2_video_generation_references(material_payload: dict[str, Any]) -> list[dict[str, Any]]:
"""V2 最终视频只能携带可选素材图片,绝不携带素材视频/音频。"""
image_url = str(material_payload.get("material_image_url") or "").strip()
if not image_url:
return []
return [
{
"type": "image",
"url": build_file_url_or_data_uri(image_url),
"name": "素材参考图片",
"upload_resource_id": material_payload.get("material_image_resource_id"),
"role": "reference_image",
}
]
def assert_v2_video_generation_references(references: list[dict[str, Any]]) -> None:
image_count = 0
for item in references:
media_type = str(item.get("type") or item.get("media_type") or "").lower()
if media_type == "image":
image_count += 1
continue
raise HTTPException(status_code=400, detail="V2 视频生成只允许携带一张素材图片,禁止视频或音频附件")
if image_count > 1:
raise HTTPException(status_code=400, detail="V2 视频生成最多携带一张素材图片")
@@ -19,6 +19,7 @@ from app.enums.recent_generation import (
) )
from app.models.chat_generation_task import ChatGenerationTask from app.models.chat_generation_task import ChatGenerationTask
from app.models.generation_record import GenerationRecord from app.models.generation_record import GenerationRecord
from app.models.module_generation_project import ModuleGenerationProject
from app.models.module_generation_step import ModuleGenerationStep from app.models.module_generation_step import ModuleGenerationStep
from app.models.shot_replicate_segment import ShotReplicateSegment from app.models.shot_replicate_segment import ShotReplicateSegment
from app.schemas.recent_generation import RecentGenerationGroupOut, RecentGenerationItemOut from app.schemas.recent_generation import RecentGenerationGroupOut, RecentGenerationItemOut
@@ -30,6 +31,7 @@ MAX_RECENT_GENERATION_LIMIT = 100
class _StepLinkInfo(TypedDict): class _StepLinkInfo(TypedDict):
module_project_id: str | None module_project_id: str | None
module_project_flow_version: str | None
module_step_id: str | None module_step_id: str | None
module: str | None module: str | None
@@ -115,6 +117,7 @@ def _build_item(
module=module, module=module,
shot_task_set_id=shot_info["shot_task_set_id"] if shot_info else None, shot_task_set_id=shot_info["shot_task_set_id"] if shot_info else None,
shot_segment_id=shot_info["shot_segment_id"] if shot_info else None, shot_segment_id=shot_info["shot_segment_id"] if shot_info else None,
module_project_flow_version=step_info["module_project_flow_version"] if step_info else None,
module_project_id=step_info["module_project_id"] if step_info else None, module_project_id=step_info["module_project_id"] if step_info else None,
module_step_id=step_info["module_step_id"] if step_info else None, module_step_id=step_info["module_step_id"] if step_info else None,
generation_id=generation_id, generation_id=generation_id,
@@ -248,11 +251,14 @@ async def _load_step_link_map(
ModuleGenerationStep.id.label("module_step_id"), ModuleGenerationStep.id.label("module_step_id"),
ModuleGenerationStep.project_id.label("module_project_id"), ModuleGenerationStep.project_id.label("module_project_id"),
ModuleGenerationStep.module.label("module"), ModuleGenerationStep.module.label("module"),
ModuleGenerationProject.flow_version.label("module_project_flow_version"),
ModuleGenerationStep.is_current.label("is_current"), ModuleGenerationStep.is_current.label("is_current"),
ModuleGenerationStep.updated_at.label("updated_at"), ModuleGenerationStep.updated_at.label("updated_at"),
) )
.join(ModuleGenerationProject, ModuleGenerationProject.id == ModuleGenerationStep.project_id)
.where( .where(
ModuleGenerationStep.deleted_at.is_(None), ModuleGenerationStep.deleted_at.is_(None),
ModuleGenerationProject.deleted_at.is_(None),
ModuleGenerationStep.chat_task_id.in_(chat_task_ids), ModuleGenerationStep.chat_task_id.in_(chat_task_ids),
ModuleGenerationStep.module.in_( ModuleGenerationStep.module.in_(
[ [
@@ -276,6 +282,7 @@ async def _load_step_link_map(
continue continue
link_map[chat_task_id] = { link_map[chat_task_id] = {
"module_project_id": row["module_project_id"], "module_project_id": row["module_project_id"],
"module_project_flow_version": str(row["module_project_flow_version"] or "v1"),
"module_step_id": row["module_step_id"], "module_step_id": row["module_step_id"],
"module": row["module"], "module": row["module"],
} }
@@ -125,6 +125,7 @@ FLOW_CONFIG = ModuleGenerationFlowConfig(
cancel_chat_task_error_message="拆镜复刻步骤被重新生成或删除,旧生成任务已取消", cancel_chat_task_error_message="拆镜复刻步骤被重新生成或删除,旧生成任务已取消",
material_video_url_editable=False, material_video_url_editable=False,
step_io_schema_version=STEP_IO_SCHEMA_VERSION, step_io_schema_version=STEP_IO_SCHEMA_VERSION,
expected_flow_version="v1",
) )
@@ -434,12 +435,19 @@ async def project_to_detail_out(db: AsyncSession, project: ModuleGenerationProje
user_result = await db.execute(select(User.username).where(User.id == project.user_id).limit(1)) user_result = await db.execute(select(User.username).where(User.id == project.user_id).limit(1))
user_name = user_result.scalar_one_or_none() user_name = user_result.scalar_one_or_none()
flow_version = str(getattr(project, "flow_version", None) or "v1")
video_prompt_config = dict(video_prompt_input.get("video_config") or {})
video_prompt_engine_snapshot = dict(video_prompt_config.get("engine_snapshot") or {})
return ShotReplicateTaskDetailOut( return ShotReplicateTaskDetailOut(
id=project.id, id=project.id,
project_id=project.id, project_id=project.id,
user_id=project.user_id, user_id=project.user_id,
user_name=user_name, user_name=user_name,
module=project.module, module=project.module,
flow_version=flow_version,
step_count=3 if flow_version == "v2" else 5,
step_io_schema_version=("hot_opening_step_io_v2" if project.module == "hot_opening_replicate" else "shot_replicate_step_io_v2") if flow_version == "v2" else STEP_IO_SCHEMA_VERSION,
title=project.title, title=project.title,
status=project.status, status=project.status,
current_step_code=project.current_step_code, current_step_code=project.current_step_code,
@@ -454,6 +462,8 @@ async def project_to_detail_out(db: AsyncSession, project: ModuleGenerationProje
source_project_name=material_input.get("source_project_name"), source_project_name=material_input.get("source_project_name"),
target_project_name=material_input.get("target_project_name"), target_project_name=material_input.get("target_project_name"),
core_content_point=material_input.get("core_content_point"), core_content_point=material_input.get("core_content_point"),
project_description=material_input.get("project_description"),
video_config=None if flow_version == "v2" else material_input.get("video_config"),
), ),
image_generation=ShotReplicateImageGenerationOut( image_generation=ShotReplicateImageGenerationOut(
prompt_step_id=image_prompt_step.id if image_prompt_step else None, prompt_step_id=image_prompt_step.id if image_prompt_step else None,
@@ -477,9 +487,9 @@ async def project_to_detail_out(db: AsyncSession, project: ModuleGenerationProje
schema_config_source=schema_config_source, schema_config_source=schema_config_source,
schema_config_version=schema_config_version, schema_config_version=schema_config_version,
schema_config_is_fallback=schema_config_is_fallback, schema_config_is_fallback=schema_config_is_fallback,
engine_id=video_snapshot.get("id") or video_generate_input.get("engine_id"), engine_id=video_snapshot.get("id") or video_generate_input.get("engine_id") or video_prompt_config.get("engine_id"),
engine_name=video_snapshot.get("name") or video_generate_input.get("engine_name"), engine_name=video_snapshot.get("name") or video_generate_input.get("engine_name") or video_prompt_engine_snapshot.get("name"),
params=video_generate_input.get("params") or video_generate_input, params=video_generate_input.get("params") or video_prompt_config or video_generate_input,
chat_task_id=video_generate_step.chat_task_id if video_generate_step else None, chat_task_id=video_generate_step.chat_task_id if video_generate_step else None,
status=video_chat.status if video_chat else (video_generate_step.status if video_generate_step else None), status=video_chat.status if video_chat else (video_generate_step.status if video_generate_step else None),
result_video_url=build_resource_signed_url(video_url) if video_url else None, result_video_url=build_resource_signed_url(video_url) if video_url else None,
@@ -606,6 +616,8 @@ async def list_shot_replicate_projects(
id=project.id, id=project.id,
project_id=project.id, project_id=project.id,
module=project.module, module=project.module,
flow_version=str(getattr(project, "flow_version", None) or "v1"),
step_count=3 if str(getattr(project, "flow_version", None) or "v1") == "v2" else 5,
title=project.title, title=project.title,
status=project.status, status=project.status,
current_step_code=project.current_step_code, current_step_code=project.current_step_code,
@@ -1497,6 +1509,9 @@ async def generate_video_from_prompt(
async def handle_chat_generation_task_completed(db: AsyncSession, task: ChatGenerationTask) -> None: async def handle_chat_generation_task_completed(db: AsyncSession, task: ChatGenerationTask) -> None:
if not task or task.generation_mode != GENERATION_MODE: if not task or task.generation_mode != GENERATION_MODE:
return return
from app.services.module_generation_v2.flow_service import handle_chat_generation_task_finished_v2
if await handle_chat_generation_task_finished_v2(db, task=task):
return
meta_result = await db.execute( meta_result = await db.execute(
select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where( select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where(
ModuleGenerationStep.chat_task_id == task.id, ModuleGenerationStep.chat_task_id == task.id,
@@ -1579,6 +1594,9 @@ async def handle_chat_generation_task_completed(db: AsyncSession, task: ChatGene
async def handle_chat_generation_task_failed(db: AsyncSession, task: ChatGenerationTask) -> None: async def handle_chat_generation_task_failed(db: AsyncSession, task: ChatGenerationTask) -> None:
if not task or task.generation_mode != GENERATION_MODE: if not task or task.generation_mode != GENERATION_MODE:
return return
from app.services.module_generation_v2.flow_service import handle_chat_generation_task_finished_v2
if await handle_chat_generation_task_finished_v2(db, task=task):
return
meta_result = await db.execute( meta_result = await db.execute(
select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where( select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where(
ModuleGenerationStep.chat_task_id == task.id, ModuleGenerationStep.chat_task_id == task.id,
@@ -32,6 +32,7 @@ from app.schemas.shot_replicate import (
ShotSegmentListOut, ShotSegmentListOut,
ShotSegmentSplitRetryOut, ShotSegmentSplitRetryOut,
ShotReanalyzeOut, ShotReanalyzeOut,
ShotReplicateDeleteOut,
ShotSegmentOut, ShotSegmentOut,
ShotSplitByAIOut, ShotSplitByAIOut,
ShotSplitByAIRequest, ShotSplitByAIRequest,
@@ -116,6 +117,7 @@ def _segment_to_out(segment: ShotReplicateSegment, project: ModuleGenerationProj
data.module_project_title = project.title data.module_project_title = project.title
data.module_project_status = project.status data.module_project_status = project.status
data.module_project_current_step_code = project.current_step_code data.module_project_current_step_code = project.current_step_code
data.module_project_flow_version = str(getattr(project, "flow_version", None) or "v1")
return data return data
@@ -126,6 +128,7 @@ def _segment_to_detail_out(segment: ShotReplicateSegment, project: ModuleGenerat
data.module_project_title = project.title data.module_project_title = project.title
data.module_project_status = project.status data.module_project_status = project.status
data.module_project_current_step_code = project.current_step_code data.module_project_current_step_code = project.current_step_code
data.module_project_flow_version = str(getattr(project, "flow_version", None) or "v1")
return data return data
@@ -716,6 +719,44 @@ async def segment_detail(db: AsyncSession, *, current_user: User, segment_id: st
project = project_result.scalar_one_or_none() project = project_result.scalar_one_or_none()
return _segment_to_detail_out(segment, project) return _segment_to_detail_out(segment, project)
async def _delete_linked_replication_project(
db: AsyncSession,
*,
current_user: User,
project_id: str,
) -> ShotReplicateDeleteOut:
"""按项目保存的流程版本分发 V1/V2 删除,供片段和任务集联动删除复用。"""
result = await db.execute(
select(ModuleGenerationProject)
.where(
ModuleGenerationProject.id == project_id,
ModuleGenerationProject.deleted_at.is_(None),
)
.limit(1)
)
project = result.scalar_one_or_none()
if project and str(getattr(project, "flow_version", None) or "v1") == "v2":
from app.services.module_generation_v2.config import SHOT_REPLICATE_V2
from app.services.module_generation_v2.flow_service import delete_project_v2
payload = await delete_project_v2(
db,
config=SHOT_REPLICATE_V2,
current_user=current_user,
project_id=project_id,
)
return ShotReplicateDeleteOut(**payload)
from app.services.shot_replicate_flow_service import delete_shot_replicate_project
return await delete_shot_replicate_project(
db,
current_user=current_user,
project_id=project_id,
refund_unfinished=False,
)
async def delete_segment( async def delete_segment(
db: AsyncSession, db: AsyncSession,
*, *,
@@ -775,13 +816,10 @@ async def delete_segment(
deleted_module_project_id: str | None = None deleted_module_project_id: str | None = None
if module_project_id: if module_project_id:
from app.services.shot_replicate_flow_service import delete_shot_replicate_project project_delete_out = await _delete_linked_replication_project(
project_delete_out = await delete_shot_replicate_project(
db, db,
current_user=current_user, current_user=current_user,
project_id=module_project_id, project_id=module_project_id,
refund_unfinished=False,
) )
deleted_module_project_id = project_delete_out.project_id deleted_module_project_id = project_delete_out.project_id
released_size_bytes += int(project_delete_out.released_size_bytes or 0) released_size_bytes += int(project_delete_out.released_size_bytes or 0)
@@ -894,15 +932,12 @@ async def delete_task_set(
upload_resource_released += int(segment_upload_release.get("released") or 0) upload_resource_released += int(segment_upload_release.get("released") or 0)
pending_delete_resource_ids.extend(segment_upload_release.get("released_resource_ids") or []) pending_delete_resource_ids.extend(segment_upload_release.get("released_resource_ids") or [])
from app.services.shot_replicate_flow_service import delete_shot_replicate_project
deleted_module_project_count = 0 deleted_module_project_count = 0
for module_project_id in dict.fromkeys(module_project_ids): for module_project_id in dict.fromkeys(module_project_ids):
project_delete_out = await delete_shot_replicate_project( project_delete_out = await _delete_linked_replication_project(
db, db,
current_user=current_user, current_user=current_user,
project_id=module_project_id, project_id=module_project_id,
refund_unfinished=False,
) )
deleted_module_project_count += 1 deleted_module_project_count += 1
released_size_bytes += int(project_delete_out.released_size_bytes or 0) released_size_bytes += int(project_delete_out.released_size_bytes or 0)
+22
View File
@@ -23,6 +23,7 @@ CELERY_TASK_IMPORTS = (
"app.tasks.shot_replicate_tasks", "app.tasks.shot_replicate_tasks",
"app.tasks.shot_replicate_flow_tasks", "app.tasks.shot_replicate_flow_tasks",
"app.tasks.module_async_recovery_tasks", "app.tasks.module_async_recovery_tasks",
"app.tasks.module_generation_v2_tasks",
"app.tasks.user_oauth_tasks", "app.tasks.user_oauth_tasks",
"app.tasks.cleanup", "app.tasks.cleanup",
"app.tasks.private_portrait_asset_tasks", "app.tasks.private_portrait_asset_tasks",
@@ -53,6 +54,22 @@ def _beat_schedule() -> dict:
"priority": settings.DOWNLOAD_TASK_PRIORITY_RECOVER, "priority": settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
}, },
} }
schedule["generation-create-recovery"] = {
"task": CeleryTaskName.RECOVER_CREATE.value,
"schedule": max(1, int(settings.GENERATION_CREATE_RECOVERY_INTERVAL_SECONDS or 60)),
"options": {
"queue": RECOVERY_QUEUE,
"priority": settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
},
}
schedule["module-async-recovery"] = {
"task": CeleryTaskName.MODULE_ASYNC_RECOVERY.value,
"schedule": max(1, int(settings.MODULE_ASYNC_RECOVERY_INTERVAL_SECONDS or 60)),
"options": {
"queue": RECOVERY_QUEUE,
"priority": settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
},
}
schedule["video-upscale-recovery-every-minute"] = { schedule["video-upscale-recovery-every-minute"] = {
"task": CeleryTaskName.VIDEO_UPSCALE_RECOVER.value, "task": CeleryTaskName.VIDEO_UPSCALE_RECOVER.value,
"schedule": 60, "schedule": 60,
@@ -109,6 +126,7 @@ if broker_url:
"shot_replicate.split_one_segment": {"ignore_result": True}, "shot_replicate.split_one_segment": {"ignore_result": True},
"shot_replicate.start_image_prompt_optimize": {"ignore_result": True}, "shot_replicate.start_image_prompt_optimize": {"ignore_result": True},
"shot_replicate.start_video_prompt_optimize": {"ignore_result": True}, "shot_replicate.start_video_prompt_optimize": {"ignore_result": True},
"module_generation_v2.start_video_prompt_optimize": {"ignore_result": True},
CeleryTaskName.VIDEO_UPSCALE_EXECUTE_LOCAL.value: { CeleryTaskName.VIDEO_UPSCALE_EXECUTE_LOCAL.value: {
"ignore_result": True, "ignore_result": True,
"soft_time_limit": max(60, int(settings.VIDEO_UPSCALE_LOCAL_TIMEOUT_SECONDS or 3600)) + 60, "soft_time_limit": max(60, int(settings.VIDEO_UPSCALE_LOCAL_TIMEOUT_SECONDS or 3600)) + 60,
@@ -119,6 +137,8 @@ if broker_url:
CeleryTaskName.VIDEO_UPSCALE_DOWNLOAD_REMOTE_RESULT.value: {"ignore_result": True}, CeleryTaskName.VIDEO_UPSCALE_DOWNLOAD_REMOTE_RESULT.value: {"ignore_result": True},
CeleryTaskName.VIDEO_UPSCALE_FINALIZE.value: {"ignore_result": True}, CeleryTaskName.VIDEO_UPSCALE_FINALIZE.value: {"ignore_result": True},
CeleryTaskName.VIDEO_UPSCALE_RECOVER.value: {"ignore_result": True}, CeleryTaskName.VIDEO_UPSCALE_RECOVER.value: {"ignore_result": True},
CeleryTaskName.RECOVER_CREATE.value: {"ignore_result": True},
CeleryTaskName.MODULE_ASYNC_RECOVERY.value: {"ignore_result": True},
}, },
worker_prefetch_multiplier=1, worker_prefetch_multiplier=1,
broker_transport_options={ broker_transport_options={
@@ -159,11 +179,13 @@ if broker_url:
"shot_replicate.split_one_segment": {"queue": CeleryQueue.GEN_RESULT_DOWNLOAD.value}, "shot_replicate.split_one_segment": {"queue": CeleryQueue.GEN_RESULT_DOWNLOAD.value},
"shot_replicate.start_image_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value}, "shot_replicate.start_image_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
"shot_replicate.start_video_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value}, "shot_replicate.start_video_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
"module_generation_v2.start_video_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
# 恢复扫描统一走独立队列,避免占用下载/轮询/创建业务 worker。 # 恢复扫描统一走独立队列,避免占用下载/轮询/创建业务 worker。
CeleryTaskName.STARTUP_RECOVERY.value: {"queue": RECOVERY_QUEUE}, CeleryTaskName.STARTUP_RECOVERY.value: {"queue": RECOVERY_QUEUE},
CeleryTaskName.SHOT_SPLIT_RECOVERY.value: {"queue": RECOVERY_QUEUE}, CeleryTaskName.SHOT_SPLIT_RECOVERY.value: {"queue": RECOVERY_QUEUE},
CeleryTaskName.RECOVER_DOWNLOAD.value: {"queue": RECOVERY_QUEUE}, CeleryTaskName.RECOVER_DOWNLOAD.value: {"queue": RECOVERY_QUEUE},
CeleryTaskName.RECOVER_GENERATION.value: {"queue": RECOVERY_QUEUE}, CeleryTaskName.RECOVER_GENERATION.value: {"queue": RECOVERY_QUEUE},
CeleryTaskName.RECOVER_CREATE.value: {"queue": RECOVERY_QUEUE},
CeleryTaskName.MODULE_ASYNC_RECOVERY.value: {"queue": RECOVERY_QUEUE}, CeleryTaskName.MODULE_ASYNC_RECOVERY.value: {"queue": RECOVERY_QUEUE},
"user_oauth.update_oauth_accounts": {"queue": CeleryQueue.DEFAULT.value}, "user_oauth.update_oauth_accounts": {"queue": CeleryQueue.DEFAULT.value},
"app.tasks.cleanup.*": {"queue": CeleryQueue.DEFAULT.value}, "app.tasks.cleanup.*": {"queue": CeleryQueue.DEFAULT.value},
@@ -40,6 +40,7 @@ async def _recover_generation_records_once(*, include_create: bool, include_poll
args=[ref.owner_id], args=[ref.owner_id],
kwargs={"owner_type": ref.owner_type, "generation_attempt_no": ref.generation_attempt_no}, kwargs={"owner_type": ref.owner_type, "generation_attempt_no": ref.generation_attempt_no},
queue=CeleryQueue.GEN_CHATAPI_CREATE.value, queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
task_id=f"generation-create:{ref.owner_type}:{ref.owner_id}:attempt:{ref.generation_attempt_no}",
) )
counts["create"] += 1 counts["create"] += 1
except Exception as exc: except Exception as exc:
@@ -116,6 +117,17 @@ async def _run_generation_once() -> Dict[str, Any]:
return {"chat_generation_task": chat_result, "generation_record": record_result} return {"chat_generation_task": chat_result, "generation_record": record_result}
async def _run_create_once() -> Dict[str, Any]:
from app.services.generation.recovery_service import recover_stale_create_tasks_once
async with async_session() as db:
chat_result = await recover_stale_create_tasks_once(db)
record_result = await _recover_generation_records_once(
include_create=True, include_poll=False, include_download=False
)
return {"chat_generation_task": chat_result, "generation_record": record_result}
async def _run_due_poll_dispatch_once() -> Dict[str, Any]: async def _run_due_poll_dispatch_once() -> Dict[str, Any]:
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
@@ -396,6 +408,23 @@ if celery_app:
) )
@celery_app.task(
name="generation.recover_create_tasks_once",
bind=True,
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
)
def recover_create_tasks_once(self) -> Dict[str, Any]:
return run_async(
_run_with_execution_lock(
lock_key=f"{settings.GENERATION_RECOVERY_LOCK_KEY}:create",
log_context="generation_create_recovery",
runner=_run_create_once,
ttl_seconds=max(55, int(settings.GENERATION_CREATE_RECOVERY_INTERVAL_SECONDS or 60) - 5),
)
)
@celery_app.task( @celery_app.task(
name="generation.dispatch_due_poll_tasks", name="generation.dispatch_due_poll_tasks",
bind=True, bind=True,
@@ -424,4 +453,5 @@ else:
startup_recovery_once = _DisabledTask() startup_recovery_once = _DisabledTask()
recover_download_tasks_once = _DisabledTask() recover_download_tasks_once = _DisabledTask()
recover_generation_tasks_once = _DisabledTask() recover_generation_tasks_once = _DisabledTask()
recover_create_tasks_once = _DisabledTask()
dispatch_due_poll_tasks = _DisabledTask() dispatch_due_poll_tasks = _DisabledTask()
@@ -0,0 +1,81 @@
from __future__ import annotations
from typing import Any
from sqlalchemy import select
from app.models.base import async_session
from app.models.module_generation_project import ModuleGenerationProject
from app.services.module_async_recovery_service import (
OBJECT_MODULE_STEP,
TASK_MODULE_V2_VIDEO_PROMPT,
acquire_object_lock,
cleanup_active_if_terminal,
mark_active_started,
register_module_step_task,
release_object_lock,
)
from app.services.module_generation_v2.config import VIDEO_PROMPT_OPTIMIZE
from app.services.module_generation_v2.flow_service import run_video_prompt_optimize_v2
from app.tasks.async_runner import run_async
from app.tasks.celery_app import celery_app
async def _run_video_prompt(project_id: str, step_id: str) -> None:
lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id)
if not lock_token:
return None
try:
async with async_session() as db:
result = await db.execute(
select(ModuleGenerationProject.module)
.where(
ModuleGenerationProject.id == project_id,
ModuleGenerationProject.deleted_at.is_(None),
)
.limit(1)
)
module = result.scalar_one_or_none()
if not module:
return None
await register_module_step_task(
module=str(module),
project_id=project_id,
step_id=step_id,
step_code=VIDEO_PROMPT_OPTIMIZE,
task_name=TASK_MODULE_V2_VIDEO_PROMPT,
)
await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id)
await run_video_prompt_optimize_v2(db, project_id=project_id, step_id=step_id)
await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id)
finally:
await release_object_lock(
object_type=OBJECT_MODULE_STEP,
object_id=step_id,
token=lock_token,
)
if celery_app:
@celery_app.task(
name="module_generation_v2.start_video_prompt_optimize",
bind=True,
max_retries=3,
default_retry_delay=30,
ignore_result=True,
)
def start_video_prompt_optimize_v2(self, project_id: str, step_id: str):
try:
run_async(_run_video_prompt(project_id, step_id))
return None
except Exception as exc:
raise self.retry(exc=exc) from exc
else:
class _DisabledTask:
def delay(self, *args: Any, **kwargs: Any):
raise RuntimeError("Celery is disabled")
def apply_async(self, *args: Any, **kwargs: Any):
raise RuntimeError("Celery is disabled")
start_video_prompt_optimize_v2 = _DisabledTask()
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+36 -36
View File
@@ -1,37 +1,37 @@
<!doctype html> <!doctype html>
<html lang="zh-CN"> <html lang="zh-CN">
<head> <head>
<meta charset="UTF-8" /> <meta charset="UTF-8" />
<link rel="icon" type="image/svg+xml" href="/favicon.svg" /> <link rel="icon" type="image/svg+xml" href="/favicon.svg" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" /> <meta name="viewport" content="width=device-width, initial-scale=1.0" />
<link rel="preconnect" href="https://fonts.googleapis.com" /> <link rel="preconnect" href="https://fonts.googleapis.com" />
<link rel="preconnect" href="https://fonts.gstatic.com" crossorigin /> <link rel="preconnect" href="https://fonts.gstatic.com" crossorigin />
<link href="https://fonts.googleapis.com/css2?family=Outfit:wght@300;400;500;600;700&display=swap" rel="stylesheet" /> <link href="https://fonts.googleapis.com/css2?family=Outfit:wght@300;400;500;600;700&display=swap" rel="stylesheet" />
<title>民众智创</title> <title>民众智创</title>
<script> <script>
(function() { (function() {
var cached = localStorage.getItem('siteInfo'); var cached = localStorage.getItem('siteInfo');
if (cached) { if (cached) {
try { try {
var info = JSON.parse(cached); var info = JSON.parse(cached);
if (info.siteName) { if (info.siteName) {
document.title = info.siteName; document.title = info.siteName;
} }
if (info.siteLogo) { if (info.siteLogo) {
var link = document.querySelector('link[rel="icon"]'); var link = document.querySelector('link[rel="icon"]');
if (link) { if (link) {
link.href = info.siteLogo; link.href = info.siteLogo;
link.type = 'image/png'; link.type = 'image/png';
} }
} }
} catch (e) {} } catch (e) {}
} }
})(); })();
</script> </script>
<script type="module" crossorigin src="/assets/index-D7Lr6Blo.js"></script> <script type="module" crossorigin src="/assets/index-DZg1zviY.js"></script>
<link rel="stylesheet" crossorigin href="/assets/index-Bsz_Xon-.css"> <link rel="stylesheet" crossorigin href="/assets/index-Bsz_Xon-.css">
</head> </head>
<body> <body>
<div id="root"></div> <div id="root"></div>
</body> </body>
</html> </html>
+63 -9
View File
@@ -269,8 +269,11 @@ export async function updateRecordPrompt(recordId: string, optimizedPrompt: stri
export async function generateVideo(recordId: string, params: GenerateParams): Promise<GenerationRecord> { export async function generateVideo(recordId: string, params: GenerateParams): Promise<GenerationRecord> {
if (USE_MOCK) return mock.mockGenerateVideo(recordId); if (USE_MOCK) return mock.mockGenerateVideo(recordId);
return api.post<GenerationRecord>(`/generation-records/${recordId}/generate`, { return api.post<GenerationRecord>(`/generation-records/${recordId}/generate`, {
engine_id: params.engineId,
include_media_references: params.includeMediaReferences ?? false,
aspect_ratio: params.aspectRatio, aspect_ratio: params.aspectRatio,
resolution: params.resolution, resolution: params.resolution,
image_size: params.imageSize,
}); });
} }
// ── Credits ─────────────────────────────────────────────── // ── Credits ───────────────────────────────────────────────
@@ -485,7 +488,7 @@ export async function getAuthorizationList(params: OAuthAppParam): Promise<OAuth
// 爆款开头复刻 // 爆款开头复刻
export async function generateReplication(params: any): Promise<any> { export async function generateReplication(params: any): Promise<any> {
return api.post('/hot-opening-replications/tasks', params); return api.post('/v2/hot-opening-replications/tasks', params);
} }
// 获取爆款开头复刻任务列表 // 获取爆款开头复刻任务列表
export async function getReplicationList(page: number, page_size: number, keyword?: string): Promise<any[]> { export async function getReplicationList(page: number, page_size: number, keyword?: string): Promise<any[]> {
@@ -495,9 +498,37 @@ export async function getReplicationList(page: number, page_size: number, keywor
} }
return api.get(url); return api.get(url);
} }
// 获取爆款开头复刻任务详情 // 获取爆款开头复刻任务详情。版本必须由列表/创建结果/URL 明确传入,禁止错误降级。
export async function getReplicationDetail(id: string): Promise<any> { export async function getReplicationDetail(id: string, flowVersion: 'v1' | 'v2'): Promise<any> {
return api.get(`/hot-opening-replications/tasks/${id}`); return flowVersion === 'v2'
? api.get(`/v2/hot-opening-replications/tasks/${id}`)
: api.get(`/hot-opening-replications/tasks/${id}`);
}
export interface RetryVideoPromptV2Params {
video_config: {
engine_id: string;
duration: number;
aspect_ratio: string;
resolution: string;
};
}
export async function retryHotOpeningVideoPromptV2(
projectId: string,
stepId: string,
params: RetryVideoPromptV2Params,
): Promise<any> {
return api.post(`/v2/hot-opening-replications/tasks/${projectId}/steps/${stepId}/retry-video-prompt`, params);
}
export async function updateHotOpeningVideoPromptSchemaV2(projectId: string, stepId: string, params: { prompt_schema: Record<string, any> }): Promise<any> {
return api.put(`/v2/hot-opening-replications/tasks/${projectId}/steps/${stepId}/video-prompt-schema`, params);
}
export async function generateHotOpeningVideoV2(projectId: string, stepId: string): Promise<any> {
return api.post(`/v2/hot-opening-replications/tasks/${projectId}/steps/${stepId}/generate-video`);
} }
// 第一步,生成提示词 // 第一步,生成提示词
export async function getone(projectId: string, stepId: string): Promise<any> { export async function getone(projectId: string, stepId: string): Promise<any> {
@@ -645,11 +676,30 @@ export async function Removelist(taskSetId: string): Promise<any> {
} }
// 生成视频 // 生成视频
export async function removeCreate(recordId: string, params: any): Promise<any> { export async function removeCreate(recordId: string, params: any): Promise<any> {
return api.post(`/shot-replications/segments/${recordId}/replication-projects`, params); return api.post(`/v2/shot-replications/segments/${recordId}/replication-projects`, params);
} }
// 获取爆款开头复刻任务详情 // 获取拆镜复刻项目详情。版本必须明确传入,禁止任何异常回退 V1。
export async function removeDetail(id: string): Promise<any> { export async function removeDetail(id: string, flowVersion: 'v1' | 'v2'): Promise<any> {
return api.get(`/shot-replications/projects/${id}`); return flowVersion === 'v2'
? api.get(`/v2/shot-replications/projects/${id}`)
: api.get(`/shot-replications/projects/${id}`);
}
export async function retryShotVideoPromptV2(
projectId: string,
stepId: string,
params: RetryVideoPromptV2Params,
): Promise<any> {
return api.post(`/v2/shot-replications/projects/${projectId}/steps/${stepId}/retry-video-prompt`, params);
}
export async function updateShotVideoPromptSchemaV2(projectId: string, stepId: string, params: any): Promise<any> {
return api.put(`/v2/shot-replications/projects/${projectId}/steps/${stepId}/video-prompt-schema`, params);
}
export async function generateShotVideoV2(projectId: string, stepId: string): Promise<any> {
return api.post(`/v2/shot-replications/projects/${projectId}/steps/${stepId}/generate-video`);
} }
// 第一步,生成提示词 // 第一步,生成提示词
export async function removeone(projectId: string, stepId: string): Promise<any> { export async function removeone(projectId: string, stepId: string): Promise<any> {
@@ -697,7 +747,11 @@ export async function deleteShotReplicationProject(taskSetId: string): Promise<v
await api.delete(`/shot-replications/task-sets/${taskSetId}`); await api.delete(`/shot-replications/task-sets/${taskSetId}`);
} }
// 删除爆款开头复刻任务 // 删除爆款开头复刻任务
export async function deleteHotOpeningReplicationTask(taskId: string): Promise<void> { export async function deleteHotOpeningReplicationTask(taskId: string, flowVersion?: string): Promise<void> {
if (flowVersion === 'v2') {
await api.delete(`/v2/hot-opening-replications/tasks/${taskId}`);
return;
}
await api.delete(`/hot-opening-replications/tasks/${taskId}`); await api.delete(`/hot-opening-replications/tasks/${taskId}`);
} }
// 获取地区信息 // 获取地区信息
+215 -31
View File
@@ -15,6 +15,8 @@ import {
Typography, Typography,
Upload, Upload,
Image, Image,
Select,
Switch,
} from "antd"; } from "antd";
import { import {
ArrowLeftOutlined, ArrowLeftOutlined,
@@ -635,7 +637,7 @@ const GeneratePage: React.FC = () => {
const [expandedGroup, setExpandedGroup] = useState<string | null>(null); const [expandedGroup, setExpandedGroup] = useState<string | null>(null);
// Per-record param selections for history prompt_optimized records // Per-record param selections for history prompt_optimized records
const [historyParams, setHistoryParams] = useState< const [historyParams, setHistoryParams] = useState<
Record<string, { aspectRatio: AspectRatio; resolution: Resolution }> Record<string, { aspectRatio?: AspectRatio; resolution?: Resolution; engineId?: string; includeMediaReferences?: boolean }>
>({}); >({});
// Video preview modal // Video preview modal
@@ -645,6 +647,10 @@ const GeneratePage: React.FC = () => {
const [videoDuration, setVideoDuration] = useState(5); const [videoDuration, setVideoDuration] = useState(5);
const [videoAspectRatio, setVideoAspectRatio] = useState<AspectRatio>("16:9"); const [videoAspectRatio, setVideoAspectRatio] = useState<AspectRatio>("16:9");
const [videoResolution, setVideoResolution] = useState<Resolution>("720p"); const [videoResolution, setVideoResolution] = useState<Resolution>("720p");
const [videoEngines, setVideoEngines] = useState<any[]>([]);
const [imageEngines, setImageEngines] = useState<any[]>([]);
const [selectedEngineId, setSelectedEngineId] = useState("");
const [includeMediaReferences, setIncludeMediaReferences] = useState(false);
const [expandedEngine, setExpandedEngine] = useState<string | null>(null); const [expandedEngine, setExpandedEngine] = useState<string | null>(null);
const [engineOptions, setEngineOptions] = useState<{ const [engineOptions, setEngineOptions] = useState<{
ratios: string[]; ratios: string[];
@@ -864,29 +870,46 @@ const GeneratePage: React.FC = () => {
}, [showImageSettingsModal]); }, [showImageSettingsModal]);
const calcVideoCredits = (duration: number, resolution: Resolution): any => { const getReferenceUsage = (items?: MediaReference[]) => {
const refs = Array.isArray(items) ? items : [];
return {
inputImageCount: refs.filter((item) => item.type === 'image').length,
inputVideoDuration: refs
.filter((item) => item.type === 'video')
.reduce((sum, item: any) => sum + Number(item.duration || 0), 0),
};
};
// console.log(duration,resolution); const calcVideoCredits = (
for (let i = 0; i < creditRatios.length; i++) { duration: number,
const ratio = creditRatios[i]; resolution: Resolution,
if (ratio.resolution === resolution) { engineId: string = selectedEngineId,
const cfg = ratio; includeReferences: boolean = includeMediaReferences,
return Math.round((cfg.baseCredits + cfg.perSecondCredits * duration) * cfg.ratio); referenceItems: MediaReference[] | undefined = currentRecord?.references,
): number => {
const cfg = creditRatios.find((item: any) =>
item.resolution === resolution && (!engineId || item.modelConfigId === engineId),
) || creditRatios.find((item: any) => item.resolution === resolution);
if (!cfg) return 0;
let total = (Number(cfg.baseCredits || 0) + Number(cfg.perSecondCredits || 0) * duration) * Number(cfg.ratio || 1);
if (includeReferences) {
const usage = getReferenceUsage(referenceItems);
if (usage.inputVideoDuration > 0) {
total += (Number(cfg.inputVideoBaseCredits || 0) + Number(cfg.inputVideoPerSecondCredits || 0) * usage.inputVideoDuration) * Number(cfg.inputVideoRatio || 1);
}
if (usage.inputImageCount > 0) {
total += (Number(cfg.inputImageBaseCredits || 0) + Number(cfg.inputImagePerImageCredits || 0) * usage.inputImageCount) * Number(cfg.inputImageRatio || 1);
} }
} }
return Number(total.toFixed(2));
// if (condition) {
// }
// return Math.round((cfg.base + cfg.perSecond * duration) * cfg.ratio);
}; };
// 根据图片分辨率获取积分 // 根据图片分辨率获取积分
const getImageCredits = (imageSize: string): any => { const getImageCredits = (imageSize: string): any => {
// 将 imageSize 转换为 creditRatios 中的 resolution 格式 // 将 imageSize 转换为 creditRatios 中的 resolution 格式
const resolution = imageSize === '2k' ? '2048' : imageSize === '4k' ? '4096' : imageSize; const normalizedSize = String(imageSize || '').toLowerCase();
const resolution = normalizedSize === '2k' ? '2048' : normalizedSize === '4k' ? '4096' : imageSize;
for (let i = 0; i < cimage.length; i++) { for (let i = 0; i < cimage.length; i++) {
const ratio = cimage[i]; const ratio = cimage[i];
@@ -1005,7 +1028,9 @@ const GeneratePage: React.FC = () => {
getParameters() getParameters()
.then((data) => { .then((data) => {
const sizes = (data as any).items?.[0]?.supportedSizes || {}; const imageItems = (data as any).items || [];
setImageEngines(imageItems);
const sizes = imageItems?.[0]?.supportedSizes || {};
setSupportedSizes(sizes); setSupportedSizes(sizes);
const resolutionKeys = Object.keys(sizes); const resolutionKeys = Object.keys(sizes);
@@ -1037,8 +1062,10 @@ const GeneratePage: React.FC = () => {
.catch(() => { }); .catch(() => { });
getVideoEngines() getVideoEngines()
.then((data) => { .then((data) => {
setVideoEngines(data.items || []);
if (data.items?.length) { if (data.items?.length) {
const e = data.items[0]; const e = data.items[0];
setSelectedEngineId((prev) => prev || e.id);
setEngineOptions({ setEngineOptions({
ratios: e.supportedRatios?.length ratios: e.supportedRatios?.length
? e.supportedRatios ? e.supportedRatios
@@ -1055,6 +1082,87 @@ const GeneratePage: React.FC = () => {
.catch(() => { }); .catch(() => { });
}, []); }, []);
useEffect(() => {
if (!currentRecord) return;
const type = currentRecord.genType || mediaType;
const engines = type === 'image' ? imageEngines : videoEngines;
const savedEngineId = currentRecord.engineId || '';
const engine = engines.find((item: any) => item.id === savedEngineId) || engines[0];
if (engine) {
setSelectedEngineId(engine.id);
if (type === 'image') {
const sizes = engine.supportedSizes || {};
setSupportedSizes(sizes);
const resolutionKeys = Object.keys(sizes);
const nextResolution = resolutionKeys.includes(currentRecord.imageSize || '')
? currentRecord.imageSize
: resolutionKeys[0] || currentRecord.imageSize || '2K';
setSelectedResolution(nextResolution);
const ratioKeys = Object.keys(sizes[nextResolution] || {});
const nextRatio = ratioKeys.includes(currentRecord.imageProportion || '')
? currentRecord.imageProportion
: ratioKeys[0] || currentRecord.imageProportion || '1:1';
setSelectedRatio(nextRatio);
const pixelSize = sizes[nextResolution]?.[nextRatio] || currentRecord.imagePx || '';
if (pixelSize) {
const [nextWidth, nextHeight] = String(pixelSize).split(/x/i).map(Number);
if (nextWidth > 0 && nextHeight > 0) {
setWidth(nextWidth);
setHeight(nextHeight);
}
}
} else {
setEngineOptions({
ratios: engine.supportedRatios?.length ? engine.supportedRatios : ['16:9'],
resolutions: engine.supportedResolutions?.length ? engine.supportedResolutions : ['720p'],
durations: engine.supportedDurations?.length ? engine.supportedDurations : [5],
});
if (!engine.supportedRatios?.includes(videoAspectRatio)) setVideoAspectRatio((engine.supportedRatios?.[0] || '16:9') as AspectRatio);
if (!engine.supportedResolutions?.includes(videoResolution)) setVideoResolution((engine.supportedResolutions?.[0] || '720p') as Resolution);
if (!engine.supportedDurations?.includes(videoDuration)) setVideoDuration(engine.supportedDurations?.[0] || currentRecord.duration || 5);
}
}
setIncludeMediaReferences(Boolean(currentRecord.includeMediaReferences));
}, [currentRecord?.id, currentRecord?.engineId, videoEngines, imageEngines]);
const getVideoEngineOptions = (engineId?: string) => {
const engine = videoEngines.find((item: any) => item.id === engineId) || videoEngines[0];
return {
ratios: engine?.supportedRatios?.length ? engine.supportedRatios : ['16:9'],
resolutions: engine?.supportedResolutions?.length ? engine.supportedResolutions : ['720p'],
durations: engine?.supportedDurations?.length ? engine.supportedDurations : [5],
};
};
const handleGenerationEngineChange = (engineId: string) => {
setSelectedEngineId(engineId);
const engines = mediaType === 'image' ? imageEngines : videoEngines;
const engine = engines.find((item: any) => item.id === engineId);
if (!engine) return;
if (mediaType === 'image') {
const sizes = engine.supportedSizes || {};
setSupportedSizes(sizes);
const resolution = Object.keys(sizes)[0] || '2K';
const ratio = Object.keys(sizes[resolution] || {})[0] || '1:1';
setSelectedResolution(resolution);
setSelectedRatio(ratio);
const pixelSize = sizes[resolution]?.[ratio] || '';
if (pixelSize) {
const [nextWidth, nextHeight] = String(pixelSize).split(/x/i).map(Number);
if (nextWidth > 0 && nextHeight > 0) {
setWidth(nextWidth);
setHeight(nextHeight);
}
}
return;
}
const options = getVideoEngineOptions(engineId);
setEngineOptions(options);
setVideoAspectRatio(options.ratios[0] as AspectRatio);
setVideoResolution(options.resolutions[0] as Resolution);
setVideoDuration(options.durations[0]);
};
useEffect(() => { useEffect(() => {
fetchProjects(); fetchProjects();
fetchRecords({ projectId, page: 1, pageSize: 100 }); fetchRecords({ projectId, page: 1, pageSize: 100 });
@@ -1239,17 +1347,25 @@ const GeneratePage: React.FC = () => {
const userCredits = user?.credits ?? 0; const userCredits = user?.credits ?? 0;
// 根据图片分辨率从 cimage 获取积分 // 根据图片分辨率从 cimage 获取积分
const getImageCreditsFromCimage = (imageSize: string): number => { const getImageCreditsFromCimage = (
// 将 imageSize 转换为 cimage 中的 resolution 格式 imageSize: string,
const resolution = imageSize === '2k' ? '2048' : imageSize === '4k' ? '4096' : imageSize; engineId: string = selectedEngineId,
includeReferences: boolean = includeMediaReferences,
for (let i = 0; i < cimage.length; i++) { referenceItems: MediaReference[] | undefined = currentRecord?.references,
const item = cimage[i]; ): number => {
if (item.resolution === resolution) { const normalizedSize = String(imageSize || '').toLowerCase();
return item.baseCredits; const resolution = normalizedSize === '2k' ? '2048' : normalizedSize === '4k' ? '4096' : imageSize;
const item = cimage.find((rule: any) => rule.resolution === resolution && (!engineId || rule.modelConfigId === engineId))
|| cimage.find((rule: any) => rule.resolution === resolution);
if (!item) return 0;
let total = Number(item.baseCredits || 0) * Number(item.ratio || 1);
if (includeReferences) {
const inputImageCount = getReferenceUsage(referenceItems).inputImageCount;
if (inputImageCount > 0) {
total += (Number(item.inputImageBaseCredits || 0) + Number(item.inputImagePerImageCredits || 0) * inputImageCount) * Number(item.inputImageRatio || 1);
} }
} }
return 0; return Number(total.toFixed(2));
}; };
// Media credits estimate for step 2 (video or image) // Media credits estimate for step 2 (video or image)
@@ -1522,8 +1638,11 @@ const GeneratePage: React.FC = () => {
}); });
try { try {
const result = await generateVideo(recordId, { const result = await generateVideo(recordId, {
engineId: selectedEngineId || undefined,
includeMediaReferences,
aspectRatio: videoAspectRatio, aspectRatio: videoAspectRatio,
resolution: videoResolution, resolution: videoResolution,
imageSize: currentRecord?.imageSize || selectedResolution,
}); });
if (result.status === "failed") { if (result.status === "failed") {
setRecordStates((p) => ({ ...p, [recordId]: "failed" })); setRecordStates((p) => ({ ...p, [recordId]: "failed" }));
@@ -1585,8 +1704,11 @@ const GeneratePage: React.FC = () => {
await updateRecordPrompt(recordId, editedPrompt); await updateRecordPrompt(recordId, editedPrompt);
} }
const result = await generateVideo(recordId, { const result = await generateVideo(recordId, {
engineId: record.engineId || selectedEngineId || undefined,
includeMediaReferences: Boolean(record.includeMediaReferences),
aspectRatio: record.aspectRatio || "16:9", aspectRatio: record.aspectRatio || "16:9",
resolution: record.resolution || "720p", resolution: record.resolution || "720p",
imageSize: record.imageSize,
}); });
if (result.status === "failed") { if (result.status === "failed") {
@@ -3238,6 +3360,22 @@ const GeneratePage: React.FC = () => {
border: "1px solid #f0f0f5", border: "1px solid #f0f0f5",
}} }}
> >
<div style={{ minWidth: 220, marginRight: 16 }}>
<Typography.Text style={{ fontSize: 12, color: '#64748b', display: 'block', marginBottom: 6 }}></Typography.Text>
<Select
value={selectedEngineId || undefined}
onChange={handleGenerationEngineChange}
style={{ width: '100%' }}
options={(mediaType === 'image' ? imageEngines : videoEngines).map((item: any) => ({ value: item.id, label: item.name }))}
/>
{currentRecord?.references?.length ? (
<div style={{ marginTop: 10, display: 'flex', alignItems: 'center', gap: 8 }}>
<Switch size="small" checked={includeMediaReferences} onChange={setIncludeMediaReferences} />
<Typography.Text style={{ fontSize: 12, color: '#64748b' }}></Typography.Text>
</div>
) : null}
</div>
{/* Video params selection */} {/* Video params selection */}
{mediaType !== "image" && ( {mediaType !== "image" && (
<div <div
@@ -4326,6 +4464,29 @@ const GeneratePage: React.FC = () => {
backgroundColor: "transparent", backgroundColor: "transparent",
}} }}
> >
<div style={{ display: 'flex', alignItems: 'center', gap: 8, marginBottom: 12, flexWrap: 'wrap' }}>
<Select
size="small"
value={historyParams[record.id]?.engineId || record.engineId || (type === 'image' ? imageEngines[0]?.id : videoEngines[0]?.id)}
onChange={(value) => {
const options = type === 'video' ? getVideoEngineOptions(value) : null;
setHistoryParams((prev) => ({
...prev,
[record.id]: {
...prev[record.id],
engineId: value,
aspectRatio: type === 'video' ? options!.ratios[0] as AspectRatio : prev[record.id]?.aspectRatio,
resolution: type === 'video' ? options!.resolutions[0] as Resolution : prev[record.id]?.resolution,
},
}));
}}
options={(type === 'image' ? imageEngines : videoEngines).map((item: any) => ({ value: item.id, label: item.name }))}
style={{ minWidth: 160 }}
/>
{record.references?.length ? (
<><Switch size="small" checked={Boolean(historyParams[record.id]?.includeMediaReferences ?? record.includeMediaReferences)} onChange={(checked) => setHistoryParams((prev) => ({ ...prev, [record.id]: { ...prev[record.id], includeMediaReferences: checked, aspectRatio: prev[record.id]?.aspectRatio || '' as AspectRatio, resolution: prev[record.id]?.resolution || '' as Resolution } }))} /><Typography.Text style={{ fontSize: 12 }}></Typography.Text></>
) : null}
</div>
{type === "video" && ( {type === "video" && (
<div <div
style={{ style={{
@@ -4367,7 +4528,7 @@ const GeneratePage: React.FC = () => {
historyParams[record.id]?.aspectRatio || historyParams[record.id]?.aspectRatio ||
"选择比例" "选择比例"
} }
options={engineOptions.ratios} options={getVideoEngineOptions(historyParams[record.id]?.engineId || record.engineId || videoEngines[0]?.id).ratios}
expanded={ expanded={
expandedEngine === `ratio-${record.id}` expandedEngine === `ratio-${record.id}`
} }
@@ -4397,7 +4558,7 @@ const GeneratePage: React.FC = () => {
historyParams[record.id]?.resolution || historyParams[record.id]?.resolution ||
"选择分辨率" "选择分辨率"
} }
options={engineOptions.resolutions} options={getVideoEngineOptions(historyParams[record.id]?.engineId || record.engineId || videoEngines[0]?.id).resolutions}
expanded={ expanded={
expandedEngine === `res-${record.id}` expandedEngine === `res-${record.id}`
} }
@@ -4450,6 +4611,9 @@ const GeneratePage: React.FC = () => {
calcVideoCredits( calcVideoCredits(
record.duration || 5, record.duration || 5,
params.resolution, params.resolution,
params.engineId || record.engineId || videoEngines[0]?.id,
Boolean(params.includeMediaReferences ?? record.includeMediaReferences),
record.references,
) )
) { ) {
message.error("积分不足,请先充值"); message.error("积分不足,请先充值");
@@ -4473,8 +4637,11 @@ const GeneratePage: React.FC = () => {
}); });
try { try {
await generateVideo(record.id, { await generateVideo(record.id, {
engineId: params?.engineId || record.engineId || (type === 'image' ? imageEngines[0]?.id : videoEngines[0]?.id),
includeMediaReferences: Boolean(params?.includeMediaReferences ?? record.includeMediaReferences),
aspectRatio: params?.aspectRatio, aspectRatio: params?.aspectRatio,
resolution: params?.resolution, resolution: params?.resolution,
imageSize: record.imageSize,
}); });
startPolling(record.id); startPolling(record.id);
message.success({ message.success({
@@ -4529,14 +4696,31 @@ const GeneratePage: React.FC = () => {
)}{" "} )}{" "}
{/* {type === "video" && {/* {type === "video" &&
historyParams[record.id]?.resolution historyParams[record.id]?.resolution
? `(${calcVideoCredits(record.duration || 5, historyParams[record.id].resolution)}积分)` ? `(${calcVideoCredits(
record.duration || 5,
historyParams[record.id].resolution as Resolution,
historyParams[record.id].engineId || record.engineId || videoEngines[0]?.id,
Boolean(historyParams[record.id].includeMediaReferences ?? record.includeMediaReferences),
record.references,
)})`
: ""} */} : ""} */}
{type === "video" && {type === "video" &&
historyParams[record.id]?.resolution historyParams[record.id]?.resolution
? `(${calcVideoCredits(record.duration || 5, historyParams[record.id].resolution)}积分)` ? `(${calcVideoCredits(
record.duration || 5,
historyParams[record.id].resolution as Resolution,
historyParams[record.id].engineId || record.engineId || videoEngines[0]?.id,
Boolean(historyParams[record.id].includeMediaReferences ?? record.includeMediaReferences),
record.references,
)})`
: ""} : ""}
{type === "image" && record.imageSize {type === "image" && record.imageSize
? `(${getImageCredits(record.imageSize)}积分)` ? `(${getImageCreditsFromCimage(
record.imageSize,
historyParams[record.id]?.engineId || record.engineId || imageEngines[0]?.id,
Boolean(historyParams[record.id]?.includeMediaReferences ?? record.includeMediaReferences),
record.references,
)})`
: type === "image" : type === "image"
? `(积分)` ? `(积分)`
: ""} : ""}
@@ -4807,7 +4991,7 @@ const GeneratePage: React.FC = () => {
setRecordStates((p) => ({ ...p, [record.id]: 'generating' })); setRecordStates((p) => ({ ...p, [record.id]: 'generating' }));
message.loading({ content: `${projectName}」正在生成视频...`, duration: 0, key: record.id }); message.loading({ content: `${projectName}」正在生成视频...`, duration: 0, key: record.id });
try { try {
await generateVideo(record.id, { aspectRatio: params.aspectRatio, resolution: params.resolution }); await generateVideo(record.id, { engineId: params.engineId || record.engineId || videoEngines[0]?.id, includeMediaReferences: Boolean(params.includeMediaReferences ?? record.includeMediaReferences), aspectRatio: params.aspectRatio, resolution: params.resolution, imageSize: record.imageSize });
setRecordStates((p) => ({ ...p, [record.id]: 'done' })); setRecordStates((p) => ({ ...p, [record.id]: 'done' }));
message.success({ content: `${projectName}」视频生成成功!`, key: record.id, duration: 3 }); message.success({ content: `${projectName}」视频生成成功!`, key: record.id, duration: 3 });
} catch { } catch {
+2 -2
View File
@@ -901,9 +901,9 @@ const HomePage: React.FC = () => {
onClick={() => { onClick={() => {
const id = video.moduleProjectId; const id = video.moduleProjectId;
if (video.type === 'hotOpeningReplicate' && id != null) { if (video.type === 'hotOpeningReplicate' && id != null) {
navigate(`/initial/${id}/initialinfo`); navigate(`/initial/${id}/initialinfo?flow_version=${video.moduleProjectFlowVersion === 'v2' ? 'v2' : 'v1'}`);
} else if (video.type === 'shotReplicate' && id != null) { } else if (video.type === 'shotReplicate' && id != null) {
navigate(`/removelens/${id}/removefenbu`); navigate(`/removelens/${id}/removefenbu?flow_version=${video.moduleProjectFlowVersion === 'v2' ? 'v2' : 'v1'}`);
} else if (video.type === 'chatAi') { } else if (video.type === 'chatAi') {
navigate(`/conversation`); navigate(`/conversation`);
} else if (video.type === 'project') { } else if (video.type === 'project') {
+256 -67
View File
@@ -1,8 +1,8 @@
import React, { useState, useEffect } from 'react'; import React, { useState, useEffect, useRef } from 'react';
import { Button, Typography, Collapse, Space, Modal, Input, Table, message, Tooltip } from 'antd'; import { Button, Typography, Collapse, Space, Modal, Input, Table, message, Tooltip } from 'antd';
import { ArrowLeftOutlined, PlayCircleOutlined, CheckCircleOutlined, EditOutlined, DownloadOutlined, SettingOutlined, LayoutOutlined, WarningOutlined } from '@ant-design/icons'; import { ArrowLeftOutlined, PlayCircleOutlined, CheckCircleOutlined, EditOutlined, DownloadOutlined, SettingOutlined, LayoutOutlined, WarningOutlined } from '@ant-design/icons';
import { useNavigate, useParams } from 'react-router-dom'; import { useNavigate, useParams, useSearchParams } from 'react-router-dom';
import { getReplicationList, getReplicationDetail, gettwo, getthree, getfour, getEngine, updateHotOpeningVideoPromptSchema, updateImagePrompt, calculateCredits } from '../api/index'; import { getReplicationList, getReplicationDetail, gettwo, getthree, getfour, getEngine, updateHotOpeningVideoPromptSchema, updateImagePrompt, calculateCredits, retryHotOpeningVideoPromptV2, updateHotOpeningVideoPromptSchemaV2, generateHotOpeningVideoV2 } from '../api/index';
import { useAuthStore } from '../store/useAuthStore'; import { useAuthStore } from '../store/useAuthStore';
import VideoPromptSchemaEditor from '../components/VideoPromptSchemaEditor'; import VideoPromptSchemaEditor from '../components/VideoPromptSchemaEditor';
import { validateVideoPromptSchemaByConfig } from '../utils/videoPromptSchema'; import { validateVideoPromptSchemaByConfig } from '../utils/videoPromptSchema';
@@ -25,6 +25,8 @@ const buildMediaUrl = (url: string): string => {
function InitialInfo() { function InitialInfo() {
const navigate = useNavigate(); const navigate = useNavigate();
const { creatID } = useParams<{ creatID: string }>(); const { creatID } = useParams<{ creatID: string }>();
const [searchParams] = useSearchParams();
const flowVersion: 'v1' | 'v2' = searchParams.get('flow_version') === 'v2' ? 'v2' : 'v1';
const { user } = useAuthStore(); const { user } = useAuthStore();
const [modalVisible, setModalVisible] = useState(false); const [modalVisible, setModalVisible] = useState(false);
@@ -35,7 +37,10 @@ function InitialInfo() {
const [videoSchemaConfigSnapshot, setVideoSchemaConfigSnapshot] = useState<any>(null); const [videoSchemaConfigSnapshot, setVideoSchemaConfigSnapshot] = useState<any>(null);
const [editingPromptStepId, setEditingPromptStepId] = useState<string>(''); const [editingPromptStepId, setEditingPromptStepId] = useState<string>('');
const [promptSaving, setPromptSaving] = useState(false); const [promptSaving, setPromptSaving] = useState(false);
const [pollingTimer, setPollingTimer] = useState<any>(null); const pollingTimerRef = useRef<ReturnType<typeof setInterval> | null>(null);
const [retryPromptModalVisible, setRetryPromptModalVisible] = useState(false);
const [retryPromptStepId, setRetryPromptStepId] = useState<string>('');
const [retryPromptSubmitting, setRetryPromptSubmitting] = useState(false);
// 引擎和视频参数相关状态 // 引擎和视频参数相关状态
const [enginesele, setEnginesele] = useState<any>({}); const [enginesele, setEnginesele] = useState<any>({});
@@ -139,14 +144,21 @@ function InitialInfo() {
document.body.removeChild(link); document.body.removeChild(link);
}; };
// const isV2 = flowVersion === 'v2';
const baseSteps = [
{ id: 1, title: '原始素材', description: '上传原始视频素材', childId: 1 }, const baseSteps = isV2
{ id: 2, title: '生成提示词', description: '根据素材生成描述词', childId: 2 }, ? [
{ id: 3, title: '生成产品融合图', description: '生成产品与场景融合图', childId: 3 }, { id: 1, title: '素材与项目信息', description: '固定素材和项目描述', childId: 1 },
{ id: 4, title: '生成视频提示词', description: '生成视频生成提示词', childId: 4 }, { id: 2, title: '生成视频提示词', description: '生成并确认视频提示词', childId: 4 },
{ id: 5, title: '生成最终视频', description: '合成最终视频', childId: 5 }, { id: 3, title: '生成最终视频', description: '使用当前视频提示词生成视频', childId: 5 },
]; ]
: [
{ id: 1, title: '原始素材', description: '上传原始视频素材', childId: 1 },
{ id: 2, title: '生成提示词', description: '根据素材生成描述词', childId: 2 },
{ id: 3, title: '生成产品融合图', description: '生成产品与场景融合图', childId: 3 },
{ id: 4, title: '生成视频提示词', description: '生成视频生成提示词', childId: 4 },
{ id: 5, title: '生成最终视频', description: '合成最终视频', childId: 5 },
];
// 合并基础步骤和API返回的状态 // 合并基础步骤和API返回的状态
const steps = baseSteps.map((step, index) => ({ const steps = baseSteps.map((step, index) => ({
@@ -295,7 +307,9 @@ function InitialInfo() {
total += inputVideoCost; total += inputVideoCost;
} }
const inputImageCount = taskDetail?.videoGeneration?.inputMedia?.image?.length || 0; const inputImageCount = isV2
? (taskDetail?.material?.materialImageUrl ? 1 : 0)
: (taskDetail?.videoGeneration?.inputMedia?.image?.length || 0);
if (inputImageCount > 0) { if (inputImageCount > 0) {
const inputImageCost = ((config.inputImageBaseCredits || 0) + (config.inputImagePerImageCredits || 0) * inputImageCount) * (config.inputImageRatio || 1); const inputImageCost = ((config.inputImageBaseCredits || 0) + (config.inputImagePerImageCredits || 0) * inputImageCount) * (config.inputImageRatio || 1);
total += inputImageCost; total += inputImageCost;
@@ -313,12 +327,12 @@ function InitialInfo() {
// 组件卸载时清理定时器 // 组件卸载时清理定时器
useEffect(() => { useEffect(() => {
return () => { return () => {
if (pollingTimer) { if (pollingTimerRef.current) {
clearInterval(pollingTimer); clearInterval(pollingTimerRef.current);
setPollingTimer(null); pollingTimerRef.current = null;
} }
}; };
}, [pollingTimer]); }, []);
// 点击外部关闭弹窗 // 点击外部关闭弹窗
useEffect(() => { useEffect(() => {
@@ -338,39 +352,30 @@ function InitialInfo() {
// 获取任务详情数据 // 获取任务详情数据
useEffect(() => { useEffect(() => {
if (creatID) { if (creatID) {
getReplicationDetail(creatID).then((res: any) => { getReplicationDetail(creatID, flowVersion).then((res: any) => {
setTaskDetail(res); setTaskDetail(res);
if (res.steps) { if (res.steps) {
setApiSteps(res.steps); setApiSteps(res.steps);
} }
}).catch((error: any) => { }).catch((error: any) => {
message.error(error?.message || '加载任务详情失败');
}); });
} }
}, [creatID]); }, [creatID, flowVersion]);
// 监听合并后的 steps 数据,判断第五步状态并启动/停止轮询 // 仅在项目或任一步骤处于 processing 时轮询;waiting_user 必须停止。
useEffect(() => { useEffect(() => {
const shouldPoll = taskDetail?.status === 'processing'
|| apiSteps.some((step: any) => step?.status === 'processing');
// 检查第五步的状态(索引 4 if (shouldPoll && !pollingTimerRef.current) {
const fifthStep = steps[4]; pollingTimerRef.current = setInterval(pollTaskDetail, 15000);
} else if (!shouldPoll && pollingTimerRef.current) {
if (fifthStep && fifthStep.status !== 'completed' && fifthStep.status !== 'failed') { clearInterval(pollingTimerRef.current);
// 第五步未完成,启动轮询 pollingTimerRef.current = null;
if (!pollingTimer) {
const timer = setInterval(pollTaskDetail, 15000);
setPollingTimer(timer);
} else {
}
} else {
// 第五步已完成或失败,停止轮询
if (fifthStep) {
if (pollingTimer) {
clearInterval(pollingTimer);
setPollingTimer(null);
}
}
} }
}, [steps]); }, [taskDetail?.status, apiSteps, creatID, flowVersion]);
const handleOpenModal = (prompt?: any, type?: string, stepId?: string | number, schemaConfigSnapshot?: any) => { const handleOpenModal = (prompt?: any, type?: string, stepId?: string | number, schemaConfigSnapshot?: any) => {
setCurrentType(type || 'image'); setCurrentType(type || 'image');
@@ -436,7 +441,8 @@ function InitialInfo() {
setPromptSaving(true); setPromptSaving(true);
try { try {
await updateHotOpeningVideoPromptSchema(taskDetail.id, editingPromptStepId, { const updateVideoPrompt = isV2 ? updateHotOpeningVideoPromptSchemaV2 : updateHotOpeningVideoPromptSchema;
await updateVideoPrompt(taskDetail.id, editingPromptStepId, {
prompt_schema: formData, prompt_schema: formData,
}); });
message.success('视频提示词已保存'); message.success('视频提示词已保存');
@@ -456,12 +462,13 @@ function InitialInfo() {
return; return;
} }
getReplicationDetail(creatID).then((res: any) => { getReplicationDetail(creatID, flowVersion).then((res: any) => {
setTaskDetail(res); setTaskDetail(res);
if (res.steps) { if (res.steps) {
setApiSteps(res.steps); setApiSteps(res.steps);
} }
}).catch((error: any) => { }).catch((error: any) => {
console.error('轮询任务详情失败', error);
}); });
}; };
@@ -469,12 +476,13 @@ function InitialInfo() {
const refreshTaskDetail = () => { const refreshTaskDetail = () => {
if (!creatID) return; if (!creatID) return;
getReplicationDetail(creatID).then((res: any) => { getReplicationDetail(creatID, flowVersion).then((res: any) => {
setTaskDetail(res); setTaskDetail(res);
if (res.steps) { if (res.steps) {
setApiSteps(res.steps); setApiSteps(res.steps);
} }
}).catch((error: any) => { }).catch((error: any) => {
message.error(error?.message || '刷新任务详情失败');
}); });
}; };
@@ -539,35 +547,109 @@ function InitialInfo() {
// 这里可以添加下一步的逻辑,比如调用接口等 // 这里可以添加下一步的逻辑,比如调用接口等
}; };
const createvideo = (stepId: number, engineId: string) => { const createvideo = async (stepId: number, engineId: string) => {
let params = { try {
engine_id: '', if (isV2) {
} await generateHotOpeningVideoV2(taskDetail.id, String(stepId));
} else {
await getfour(taskDetail.id, String(stepId), { engine_id: engineId || '' });
getfour(taskDetail.id, stepId.toString(), params).then((res: any) => { }
// 重新获取任务详情以更新数据 message.info('正在生成视频,请稍候...');
message.info('正在生成视频,请稍候...');
refreshTaskDetail(); refreshTaskDetail();
}).catch((error: any) => { } catch (error: any) {
const errorMsg = error?.message?.split(': ')?.[1] || error?.message || '生成失败'; const errorMsg = error?.message?.split(': ')?.[1] || error?.message || '生成失败';
message.error(errorMsg); message.error(errorMsg);
});
}
const agincreatevideo = () => {
let params = {
engine_id: '',
} }
};
const agincreatevideo = async () => {
const promptStep = isV2 ? steps[1] : steps[3];
if (!promptStep?.id) return;
await createvideo(promptStep.id, promptStep.engineId || '');
};
getfour(taskDetail.id, steps[3].id.toString(), params).then((res: any) => { const applyRetryVideoEngine = (
// 重新获取任务详情以更新数据 engineId: string,
message.info('正在生成视频,请稍候...'); preferred?: { duration?: number; aspectRatio?: string; resolution?: string },
) => {
const engine = (enginesele.video || []).find((item: any) => String(item.id) === String(engineId));
if (!engine) {
return false;
}
const ratios = Array.isArray(engine.supportedRatios) && engine.supportedRatios.length > 0
? engine.supportedRatios.map((item: any) => String(item))
: ['16:9', '4:3', '1:1', '3:4', '9:16', '21:9'];
const resolutions = Array.isArray(engine.supportedResolutions) && engine.supportedResolutions.length > 0
? engine.supportedResolutions.map((item: any) => String(item))
: ['480p', '720p', '1080p'];
const parsedDurations = Array.isArray(engine.supportedDurations)
? engine.supportedDurations
.map((item: any) => Number(item))
.filter((item: number) => Number.isFinite(item) && item > 0)
: [];
const durations = parsedDurations.length > 0 ? parsedDurations : [5, 8, 10, 12, 15];
const preferredDuration = Number(preferred?.duration ?? videoDuration);
const preferredRatio = String(preferred?.aspectRatio || videoAspectRatio || '');
const preferredResolution = String(preferred?.resolution || videoResolution || '');
setCountType(String(engine.id));
setEngineOptions({ ratios, resolutions, durations });
setVideoDuration(durations.includes(preferredDuration) ? preferredDuration : durations[0]);
setVideoAspectRatio(ratios.includes(preferredRatio) ? preferredRatio : ratios[0]);
setVideoResolution(resolutions.includes(preferredResolution) ? preferredResolution : resolutions[0]);
return true;
};
const openRegenerateVideoPrompt = (stepId: number) => {
if (!isV2 || !taskDetail?.id) return;
const currentConfig = taskDetail?.videoGeneration?.promptParams || taskDetail?.videoGeneration?.params || {};
const requestedEngineId = String(
currentConfig.engineId
|| currentConfig.engine_id
|| taskDetail?.videoGeneration?.engineId
|| countType
|| '',
);
const currentEngineId = String(
(enginesele.video || []).some((engine: any) => String(engine.id) === requestedEngineId)
? requestedEngineId
: enginesele.video?.[0]?.id || '',
);
if (!currentEngineId || !applyRetryVideoEngine(currentEngineId, {
duration: Number(currentConfig.duration || videoDuration),
aspectRatio: String(currentConfig.aspectRatio || currentConfig.aspect_ratio || videoAspectRatio),
resolution: String(currentConfig.resolution || videoResolution),
})) {
message.warning('当前没有可用的视频生成引擎');
return;
}
setRetryPromptStepId(String(stepId));
setRetryPromptModalVisible(true);
};
const submitRegenerateVideoPrompt = async () => {
if (!isV2 || !taskDetail?.id || !retryPromptStepId || !countType) return;
setRetryPromptSubmitting(true);
try {
await retryHotOpeningVideoPromptV2(taskDetail.id, retryPromptStepId, {
video_config: {
engine_id: countType,
duration: videoDuration,
aspect_ratio: videoAspectRatio,
resolution: videoResolution,
},
});
message.info('正在按新视频参数重新生成视频提示词,请稍候...');
setRetryPromptModalVisible(false);
setRetryPromptStepId('');
refreshTaskDetail(); refreshTaskDetail();
}).catch((error: any) => { } catch (error: any) {
}); message.error(error?.message || '重新生成视频提示词失败');
} } finally {
setRetryPromptSubmitting(false);
}
};
@@ -896,6 +978,7 @@ function InitialInfo() {
</Button> </Button>
</div> </div>
{!isV2 && (<>
{/* 引擎选择器和视频参数设置 */} {/* 引擎选择器和视频参数设置 */}
<p style={{marginBottom:6,fontSize: 14, background: 'linear-gradient(135deg, #6366f1 0%, #8b5cf6 100%)', WebkitBackgroundClip: 'text', WebkitTextFillColor: 'transparent', backgroundClip: 'text' }}></p> <p style={{marginBottom:6,fontSize: 14, background: 'linear-gradient(135deg, #6366f1 0%, #8b5cf6 100%)', WebkitBackgroundClip: 'text', WebkitTextFillColor: 'transparent', backgroundClip: 'text' }}></p>
<div style={{ width:'100%', display: 'flex', justifyContent: 'space-between', marginBottom: 16 }}> <div style={{ width:'100%', display: 'flex', justifyContent: 'space-between', marginBottom: 16 }}>
@@ -1287,6 +1370,7 @@ function InitialInfo() {
</span> </span>
</Button> </Button>
</Tooltip> </Tooltip>
</>)}
</> </>
)} )}
{/* 步骤4: 生成视频提示词 */} {/* 步骤4: 生成视频提示词 */}
@@ -1307,6 +1391,16 @@ function InitialInfo() {
> >
/ /
</Button> </Button>
{isV2 && (
<Button
type="default"
onClick={() => openRegenerateVideoPrompt(step.id)}
style={{ flex: 1, borderRadius: 10, borderColor: 'rgba(99, 102, 241, 0.3)', color: '#6366f1', height: 36, fontWeight: 500 }}
disabled={step.status === 'processing' || steps[2]?.status === 'processing'}
>
</Button>
)}
<Button <Button
onClick={() => { createvideo(step.id, step.engineId); }} onClick={() => { createvideo(step.id, step.engineId); }}
type="primary" type="primary"
@@ -1341,9 +1435,9 @@ function InitialInfo() {
type="default" type="default"
icon={<EditOutlined />} icon={<EditOutlined />}
style={{ flex: 1, borderRadius: 10, borderColor: 'rgba(99, 102, 241, 0.3)', color: '#6366f1', height: 36, fontWeight: 500, background: 'rgba(99, 102, 241, 0.04)' }} style={{ flex: 1, borderRadius: 10, borderColor: 'rgba(99, 102, 241, 0.3)', color: '#6366f1', height: 36, fontWeight: 500, background: 'rgba(99, 102, 241, 0.04)' }}
disabled={step.status !== 'completed'} disabled={isV2 ? !['completed', 'failed'].includes(step.status) : step.status !== 'completed'}
> >
{step.status === 'failed' ? '重试生成' : '重新生成'}
</Button> </Button>
<Button <Button
type="primary" type="primary"
@@ -1373,6 +1467,101 @@ function InitialInfo() {
</div> </div>
</div> </div>
</div> </div>
<Modal
title="重新生成视频提词"
open={retryPromptModalVisible}
onCancel={() => {
if (!retryPromptSubmitting) {
setRetryPromptModalVisible(false);
setRetryPromptStepId('');
}
}}
onOk={submitRegenerateVideoPrompt}
confirmLoading={retryPromptSubmitting}
okText="按新参数生成提词"
cancelText="取消"
width={720}
destroyOnClose
>
<div style={{ display: 'flex', flexDirection: 'column', gap: 20 }}>
<div>
<Text strong style={{ display: 'block', marginBottom: 10 }}></Text>
<div style={{ display: 'grid', gridTemplateColumns: 'repeat(2, minmax(0, 1fr))', gap: 10 }}>
{(enginesele.video || []).map((engine: any) => (
<button
key={engine.id}
type="button"
onClick={() => applyRetryVideoEngine(String(engine.id), {
duration: videoDuration,
aspectRatio: videoAspectRatio,
resolution: videoResolution,
})}
style={{
minHeight: 46,
padding: '8px 12px',
borderRadius: 8,
border: countType === String(engine.id) ? '2px solid #6366f1' : '1px solid #e5e7eb',
background: countType === String(engine.id) ? 'rgba(99,102,241,0.08)' : '#fff',
color: countType === String(engine.id) ? '#4f46e5' : '#374151',
cursor: 'pointer',
textAlign: 'left',
fontWeight: countType === String(engine.id) ? 600 : 400,
}}
>
{engine.name || engine.id}
</button>
))}
</div>
</div>
<div>
<Text strong style={{ display: 'block', marginBottom: 10 }}></Text>
<div style={{ display: 'flex', flexWrap: 'wrap', gap: 8 }}>
{engineOptions.durations.map((duration) => (
<Button
key={duration}
type={videoDuration === duration ? 'primary' : 'default'}
onClick={() => setVideoDuration(duration)}
>
{duration}
</Button>
))}
</div>
</div>
<div>
<Text strong style={{ display: 'block', marginBottom: 10 }}></Text>
<div style={{ display: 'flex', flexWrap: 'wrap', gap: 8 }}>
{engineOptions.ratios.map((ratio) => (
<Button
key={ratio}
type={videoAspectRatio === ratio ? 'primary' : 'default'}
onClick={() => setVideoAspectRatio(ratio)}
>
{ratio}
</Button>
))}
</div>
</div>
<div>
<Text strong style={{ display: 'block', marginBottom: 10 }}></Text>
<div style={{ display: 'flex', flexWrap: 'wrap', gap: 8 }}>
{engineOptions.resolutions.map((resolution) => (
<Button
key={resolution}
type={videoResolution === resolution ? 'primary' : 'default'}
onClick={() => setVideoResolution(resolution)}
>
{resolution}
</Button>
))}
</div>
</div>
<div style={{ padding: '10px 12px', borderRadius: 8, background: '#f8fafc', color: '#64748b' }}>
{enginesele.video?.find((engine: any) => String(engine.id) === countType)?.name || countType}
{' · '}{videoDuration} · {videoAspectRatio} · {videoResolution}
<span style={{ marginLeft: 12 }}>{estimatedCredits}</span>
</div>
</div>
</Modal>
<Modal <Modal
title={ title={
<div style={{ display: 'flex', alignItems: 'center', gap: 10 }}> <div style={{ display: 'flex', alignItems: 'center', gap: 10 }}>
@@ -1586,7 +1775,7 @@ function InitialInfo() {
render: (_, record) => ( render: (_, record) => (
<button <button
onClick={() => navigate(`/initial/${record.id}/initialinfo`)} onClick={() => navigate(`/initial/${record.id}/initialinfo?flow_version=${record.flowVersion === 'v2' ? 'v2' : 'v1'}`)}
style={{ color: '#6366f1', textDecoration: 'none', fontSize: 12, border: 'none', background: 'rgba(99, 102, 241, 0.08)', padding: '4px 12px', borderRadius: 8, cursor: 'pointer', transition: 'all 0.2s', fontWeight: 500 }} style={{ color: '#6366f1', textDecoration: 'none', fontSize: 12, border: 'none', background: 'rgba(99, 102, 241, 0.08)', padding: '4px 12px', borderRadius: 8, cursor: 'pointer', transition: 'all 0.2s', fontWeight: 500 }}
> >
+120 -78
View File
@@ -11,6 +11,7 @@ import {
Space, Space,
Pagination, Pagination,
Popconfirm, Popconfirm,
Select,
} from 'antd'; } from 'antd';
import { import {
PlusOutlined, PlusOutlined,
@@ -19,8 +20,9 @@ import {
LoadingOutlined, LoadingOutlined,
} from '@ant-design/icons'; } from '@ant-design/icons';
import { useNavigate } from 'react-router-dom'; import { useNavigate } from 'react-router-dom';
import { uploadHotOpeningVideo, uploadHotOpeningImage, generateReplication, getReplicationList, getone, getReplicationDetail, deleteHotOpeningReplicationTask } from '../api'; import { uploadHotOpeningVideo, uploadHotOpeningImage, generateReplication, getReplicationList, deleteHotOpeningReplicationTask, getEngine, calculateCredits } from '../api';
import bg1 from '../assets/bg1.png'; import bg1 from '../assets/bg1.png';
import UploadSelector from '../components/UploadSelector';
@@ -29,6 +31,13 @@ const { TextArea } = Input;
const API_BASE = import.meta.env.VITE_API_BASE || 'http://localhost:8000'; const API_BASE = import.meta.env.VITE_API_BASE || 'http://localhost:8000';
const supportsReferenceImage = (engine: any): boolean => {
if (!engine) return false;
const supports = engine.supportsUniversalReference ?? engine.supports_universal_reference;
const maxImages = engine.maxImageCount ?? engine.max_image_count;
return supports !== false && (maxImages === undefined || maxImages === null || Number(maxImages) >= 1);
};
const buildAssetUrl = (url?: string): string => { const buildAssetUrl = (url?: string): string => {
if (!url) return ''; if (!url) return '';
if (/^https?:\/\//i.test(url) || url.startsWith('blob:')) return url; if (/^https?:\/\//i.test(url) || url.startsWith('blob:')) return url;
@@ -70,6 +79,12 @@ const GenerateConver: React.FC = () => {
const [imageResourceId, setImageResourceId] = useState<string>(''); const [imageResourceId, setImageResourceId] = useState<string>('');
const [videoUploading, setVideoUploading] = useState(false); const [videoUploading, setVideoUploading] = useState(false);
const [imageUploading, setImageUploading] = useState(false); const [imageUploading, setImageUploading] = useState(false);
const [videoEngines, setVideoEngines] = useState<any[]>([]);
const [engineId, setEngineId] = useState('');
const [videoDuration, setVideoDuration] = useState(5);
const [videoAspectRatio, setVideoAspectRatio] = useState('16:9');
const [videoResolution, setVideoResolution] = useState('480p');
const [creditRules, setCreditRules] = useState<any[]>([]);
// 弹窗状态 // 弹窗状态
const [isModalOpen, setIsModalOpen] = useState(false); const [isModalOpen, setIsModalOpen] = useState(false);
@@ -79,6 +94,35 @@ const GenerateConver: React.FC = () => {
// 表格轮询定时器(使用 ref 避免闭包问题) // 表格轮询定时器(使用 ref 避免闭包问题)
const tablePollingTimer = useRef<any>(null); const tablePollingTimer = useRef<any>(null);
// 同一次创建在网络重试时复用幂等键;成功后再生成新键。
const createIdempotencyKeyRef = useRef<string | null>(null);
useEffect(() => {
getEngine().then((data: any) => {
const engines = data?.engine?.video || [];
setVideoEngines(engines);
if (engines.length > 0) {
const first = engines[0];
setEngineId(first.id);
setVideoAspectRatio(first.supportedRatios?.[0] || '16:9');
setVideoResolution(first.supportedResolutions?.[0] || '480p');
setVideoDuration(first.supportedDurations?.[0] || 5);
}
}).catch(() => message.error('视频引擎加载失败'));
calculateCredits().then((rules: any) => setCreditRules(Array.isArray(rules) ? rules : [])).catch(() => setCreditRules([]));
}, []);
const selectedEngine = videoEngines.find((item: any) => item.id === engineId);
const selectedEngineSupportsImage = !imageUrl || supportsReferenceImage(selectedEngine);
const estimatedCredits = (() => {
const rule = creditRules.find((item: any) => item.modelConfigId === engineId && item.genType === 'video' && item.resolution === videoResolution);
if (!rule) return 0;
let total = (Number(rule.baseCredits || 0) + Number(rule.perSecondCredits || 0) * videoDuration) * Number(rule.ratio || 1);
if (imageUrl) {
total += (Number(rule.inputImageBaseCredits || 0) + Number(rule.inputImagePerImageCredits || 0)) * Number(rule.inputImageRatio || 1);
}
return Number(total.toFixed(2));
})();
// 获取卡片列表数据(独立于表格) // 获取卡片列表数据(独立于表格)
const fetchCardList = (page: number, size: number, isPolling = false) => { const fetchCardList = (page: number, size: number, isPolling = false) => {
@@ -355,90 +399,52 @@ const GenerateConver: React.FC = () => {
message.warning('请上传复刻视频'); message.warning('请上传复刻视频');
return; return;
} }
if (!imageUrl) { if (!engineId || !videoDuration || !videoAspectRatio || !videoResolution) {
message.warning('请上传产品图片'); message.warning('请选择完整的视频生成参数');
return; return;
} }
if (!originalProductName.trim()) { if (!selectedEngineSupportsImage) {
message.warning('请输入原视频产品名称'); message.warning('当前视频引擎不支持参考图片,请移除素材图片或切换引擎');
return; return;
} }
if (!ownProductName.trim()) { const idempotencyKey = createIdempotencyKeyRef.current
message.warning('请输入自有产品名称'); || `hot_v2_${Date.now()}_${Math.random().toString(16).slice(2)}`;
return; createIdempotencyKeyRef.current = idempotencyKey;
} const params = {
if (!productSellingPoints.trim()) {
message.warning('请输入产品卖点');
return;
}
// message.success('正在生成爆款开头复刻视频...');
// console.log('生成参数:', {
// videoUrl,
// imageUrl,
// originalProductName,
// ownProductName,
// productSellingPoints,
// });
let params = {
material_video_url: videoUrl, material_video_url: videoUrl,
material_image_url: imageUrl, material_image_url: imageUrl || undefined,
material_video_resource_id: videoResourceId || undefined, material_video_resource_id: videoResourceId || undefined,
material_image_resource_id: imageResourceId || undefined, material_image_resource_id: imageResourceId || undefined,
material_video_duration_seconds: videoDurationSeconds || undefined, material_video_duration_seconds: videoDurationSeconds || undefined,
source_project_name: originalProductName, source_project_name: originalProductName.trim() || undefined,
target_project_name: ownProductName, target_project_name: ownProductName.trim() || undefined,
core_content_point: productSellingPoints, project_description: productSellingPoints.trim() || undefined,
idempotency_key: Date.now().toString(), core_content_point: productSellingPoints.trim() || undefined,
} video_config: {
engine_id: engineId,
// 在这里调用生成视频的API,并传递上述参数 duration: videoDuration,
generateReplication(params).then((res) => { aspect_ratio: videoAspectRatio,
resolution: videoResolution,
},
let projectId = ''; idempotency_key: idempotencyKey,
let childId = ''; };
// console.log('生成视频成功:', res); generateReplication(params).then((res: any) => {
// message.success('任务开始'); const projectId = res.id || res.projectId || res.detail?.id;
createIdempotencyKeyRef.current = null;
// 清空上传的媒体和文本
setVideoUrl(''); setVideoUrl('');
setVideoResourceId(''); setVideoResourceId('');
setVideoDurationSeconds(null); setVideoDurationSeconds(null);
setImageUrl(''); setImageUrl('');
setImageResourceId(''); setImageResourceId('');
setImageFile(null);
setOriginalProductName(''); setOriginalProductName('');
setOwnProductName(''); setOwnProductName('');
setProductSellingPoints(''); setProductSellingPoints('');
fetchCardList(1, 8); fetchCardList(1, 8);
message.success('项目已创建,视频提词正在生成');
// 获取第一个的id直接进行下一步 if (projectId) navigate(`/initial/${projectId}/initialinfo?flow_version=v2`);
getReplicationList(1, 20).then((res: any) => { }).catch((error: any) => {
if (res.items) { message.error(error?.message || '视频开头复刻失败');
const firstId = res.items[0].id;
setTableData(res.items);
getReplicationDetail(res.items[0].id).then((res: any) => {
projectId = res.id;
childId = res.steps[0].id;
getone(projectId, childId).then((res: any) => {
message.loading('创建中...', 3);
setTimeout(() => {
navigate(`/initial/${firstId}/initialinfo`);
}, 3000);
})
}).catch((error: any) => {
});
}
}).catch((error: any) => {
});
}).catch((error) => {
message.error('视频开头复刻失败');
}); });
}; };
@@ -548,7 +554,7 @@ const GenerateConver: React.FC = () => {
transition: 'transform 0.25s cubic-bezier(0.4, 0, 0.2, 1), box-shadow 0.25s cubic-bezier(0.4, 0, 0.2, 1)', transition: 'transform 0.25s cubic-bezier(0.4, 0, 0.2, 1), box-shadow 0.25s cubic-bezier(0.4, 0, 0.2, 1)',
// border: '1px solid rgba(99, 102, 241, 0.06)', // border: '1px solid rgba(99, 102, 241, 0.06)',
}} }}
onClick={() => navigate(`/initial/${item.id}/initialinfo`)} onClick={() => navigate(`/initial/${item.id}/initialinfo?flow_version=${item.flowVersion === 'v2' ? 'v2' : 'v1'}`)}
onMouseEnter={(e) => { onMouseEnter={(e) => {
e.currentTarget.style.transform = 'translateY(-6px)'; e.currentTarget.style.transform = 'translateY(-6px)';
e.currentTarget.style.boxShadow = '0 12px 32px rgba(99, 102, 241, 0.12)'; e.currentTarget.style.boxShadow = '0 12px 32px rgba(99, 102, 241, 0.12)';
@@ -1021,12 +1027,22 @@ const GenerateConver: React.FC = () => {
</div> </div>
</div> </div>
) : ( ) : (
<Upload <UploadSelector
style={{ width: '100%' }} accept="image/*"
multiple={false}
beforeUpload={beforeImageUpload} mediaType="image"
showUploadList={false} uploading={imageUploading}
accept="image/jpeg,image/jpg,image/png,.jpg,.jpeg,.png" maxImageCount={1}
usedImageCount={0}
onLocalSelect={(files) => { if (files[0]) void beforeImageUpload(files[0]); }}
onHistorySelect={(items) => {
const item = items[0];
if (!item) return;
setImageFile(null);
setImageUrl(item.resourceUrl || item.previewUrl || item.displayUrl || '');
setImageResourceId(item.id);
message.success('历史图片已选择');
}}
> >
<div <div
style={{ style={{
@@ -1088,7 +1104,7 @@ const GenerateConver: React.FC = () => {
</> </>
)} )}
</div> </div>
</Upload> </UploadSelector>
)} )}
</div> </div>
@@ -1161,12 +1177,38 @@ const GenerateConver: React.FC = () => {
</div> </div>
</div> </div>
<div style={{ marginBottom: 20, display: 'grid', gap: 10 }}>
<p style={{ margin: 0, fontSize: 13, fontWeight: 600, color: '#475569' }}></p>
<Select value={engineId || undefined} placeholder="选择视频引擎" onChange={(value) => {
setEngineId(value);
const engine = videoEngines.find((item: any) => item.id === value);
setVideoAspectRatio(engine?.supportedRatios?.[0] || '16:9');
setVideoResolution(engine?.supportedResolutions?.[0] || '480p');
setVideoDuration(engine?.supportedDurations?.[0] || 5);
}} options={videoEngines.map((item: any) => ({
value: item.id,
label: imageUrl && !supportsReferenceImage(item) ? `${item.name}(不支持参考图)` : item.name,
disabled: Boolean(imageUrl) && !supportsReferenceImage(item),
}))} />
<div style={{ display: 'grid', gridTemplateColumns: 'repeat(3, 1fr)', gap: 8 }}>
<Select value={videoDuration} onChange={setVideoDuration} options={(selectedEngine?.supportedDurations || [5]).map((value: number) => ({ value, label: `${value}` }))} />
<Select value={videoAspectRatio} onChange={setVideoAspectRatio} options={(selectedEngine?.supportedRatios || ['16:9']).map((value: string) => ({ value, label: value }))} />
<Select value={videoResolution} onChange={setVideoResolution} options={(selectedEngine?.supportedResolutions || ['480p']).map((value: string) => ({ value, label: value }))} />
</div>
<span style={{ fontSize: 12, color: selectedEngineSupportsImage ? '#8b5cf6' : '#ef4444' }}>
{selectedEngineSupportsImage
? `预估视频积分:${estimatedCredits || '-'};素材图片可不传,最终生成只会携带图片,不会携带参考视频。`
: '当前引擎不支持参考图片,请移除素材图片或切换引擎。'}
</span>
</div>
{/* 立即生成按钮 */} {/* 立即生成按钮 */}
<Button <Button
type="primary" type="primary"
block block
size="large" size="large"
onClick={handleGenerate} onClick={handleGenerate}
disabled={!selectedEngineSupportsImage}
style={{ style={{
borderRadius: 12, borderRadius: 12,
height: 44, height: 44,
@@ -1382,7 +1424,7 @@ const GenerateConver: React.FC = () => {
render: (_, record) => ( render: (_, record) => (
<Space> <Space>
<button <button
onClick={() => navigate(`/initial/${record.id}/initialinfo`)} onClick={() => navigate(`/initial/${record.id}/initialinfo?flow_version=${record.flowVersion === 'v2' ? 'v2' : 'v1'}`)}
style={{ style={{
color: '#6366f1', color: '#6366f1',
textDecoration: 'none', textDecoration: 'none',
@@ -1402,7 +1444,7 @@ const GenerateConver: React.FC = () => {
title="确认删除这个爆款开头复刻任务吗?" title="确认删除这个爆款开头复刻任务吗?"
onConfirm={async () => { onConfirm={async () => {
try { try {
await deleteHotOpeningReplicationTask(record.id); await deleteHotOpeningReplicationTask(record.id, record.flowVersion);
message.success('删除成功'); message.success('删除成功');
fetchList(1, pageSize, false, searchKeyword); fetchList(1, pageSize, false, searchKeyword);
} catch (err) { } catch (err) {
+145 -81
View File
@@ -1,11 +1,13 @@
import { useCallback, useEffect, useMemo, useRef, useState } from 'react'; import { useCallback, useEffect, useMemo, useRef, useState } from 'react';
import { useNavigate, useParams } from 'react-router-dom'; import { useNavigate, useParams } from 'react-router-dom';
import { Button, Drawer, Input, Modal, Popconfirm, Spin, Table, Tag, Tooltip, Upload, message } from 'antd'; import { Button, Drawer, Input, Modal, Popconfirm, Select, Spin, Table, Tag, Tooltip, message } from 'antd';
import { ArrowLeftOutlined, DeleteOutlined, PlusOutlined, XOutlined } from '@ant-design/icons'; import { ArrowLeftOutlined, DeleteOutlined, PlusOutlined, XOutlined } from '@ant-design/icons';
import {uploadShotReplicateImage, createRemoveLens, deleteSegment, getShotReplicationDetail, reanalyzeSegment, retrySplit, Removelist, removeCreate, reanalyzeShotReplication, splitCustom, uploadImage } from '../api'; import { uploadShotReplicateImage, createRemoveLens, deleteSegment, getShotReplicationDetail, reanalyzeSegment, retrySplit, Removelist, removeCreate, reanalyzeShotReplication, splitCustom, getEngine, calculateCredits } from '../api';
import VideoTrimPicker from '../components/VideoTrimPicker'; import VideoTrimPicker from '../components/VideoTrimPicker';
import UploadSelector from '../components/UploadSelector';
import { useAuthStore } from '../store/useAuthStore';
const { TextArea } = Input; const { TextArea } = Input;
@@ -13,6 +15,13 @@ const API_BASE = import.meta.env.VITE_API_BASE || 'http://localhost:8000';
const MIN_TRIM_SECONDS = 2; const MIN_TRIM_SECONDS = 2;
const MAX_TRIM_SECONDS = 15; const MAX_TRIM_SECONDS = 15;
const supportsReferenceImage = (engine: any): boolean => {
if (!engine) return false;
const supports = engine.supportsUniversalReference ?? engine.supports_universal_reference;
const maxImages = engine.maxImageCount ?? engine.max_image_count;
return supports !== false && (maxImages === undefined || maxImages === null || Number(maxImages) >= 1);
};
function buildAssetUrl(url?: string): string { function buildAssetUrl(url?: string): string {
if (!url) return ''; if (!url) return '';
if (/^https?:\/\//i.test(url)) return url; if (/^https?:\/\//i.test(url)) return url;
@@ -22,6 +31,7 @@ function buildAssetUrl(url?: string): string {
function RemoveInfo() { function RemoveInfo() {
const { creatID } = useParams<{ creatID: string }>(); const { creatID } = useParams<{ creatID: string }>();
const navigate = useNavigate(); const navigate = useNavigate();
const { user } = useAuthStore();
const [drawerVisible, setDrawerVisible] = useState(false); const [drawerVisible, setDrawerVisible] = useState(false);
const [trimModalVisible, setTrimModalVisible] = useState(false); const [trimModalVisible, setTrimModalVisible] = useState(false);
const [currentSegment, setCurrentSegment] = useState<string | null>(null); const [currentSegment, setCurrentSegment] = useState<string | null>(null);
@@ -33,6 +43,12 @@ function RemoveInfo() {
const [productSellingPoint, setProductSellingPoint] = useState(''); const [productSellingPoint, setProductSellingPoint] = useState('');
const [productImage, setProductImage] = useState(''); const [productImage, setProductImage] = useState('');
const [productImageResourceId, setProductImageResourceId] = useState(''); const [productImageResourceId, setProductImageResourceId] = useState('');
const [videoEngines, setVideoEngines] = useState<any[]>([]);
const [engineId, setEngineId] = useState('');
const [videoDuration, setVideoDuration] = useState(5);
const [videoAspectRatio, setVideoAspectRatio] = useState('16:9');
const [videoResolution, setVideoResolution] = useState('480p');
const [creditRules, setCreditRules] = useState<any[]>([]);
const [taskDetail, setTaskDetail] = useState<any>(null); const [taskDetail, setTaskDetail] = useState<any>(null);
const [tableData, setTableData] = useState<any[]>([]); const [tableData, setTableData] = useState<any[]>([]);
const [loading, setLoading] = useState(false); const [loading, setLoading] = useState(false);
@@ -43,6 +59,45 @@ function RemoveInfo() {
const videoUrl = useMemo(() => buildAssetUrl(taskDetail?.videoUrl), [taskDetail?.videoUrl]); const videoUrl = useMemo(() => buildAssetUrl(taskDetail?.videoUrl), [taskDetail?.videoUrl]);
useEffect(() => {
getEngine().then((data: any) => {
const engines = data?.engine?.video || [];
setVideoEngines(engines);
if (engines.length > 0) {
const first = engines[0];
setEngineId(first.id);
setVideoAspectRatio(first.supportedRatios?.[0] || '16:9');
setVideoResolution(first.supportedResolutions?.[0] || '480p');
setVideoDuration(first.supportedDurations?.[0] || 5);
}
}).catch(() => message.error('视频引擎加载失败'));
calculateCredits().then((rules: any) => setCreditRules(Array.isArray(rules) ? rules : [])).catch(() => setCreditRules([]));
}, []);
const createIdempotencyRef = useRef<{ segmentId: string; key: string } | null>(null);
const selectedEngine = videoEngines.find((item: any) => item.id === engineId);
const selectedEngineSupportsImage = !productImage || supportsReferenceImage(selectedEngine);
const estimatedCredits = useMemo(() => {
const rule = creditRules.find((item: any) => item.modelConfigId === engineId && item.genType === 'video' && item.resolution === videoResolution);
if (!rule) return 0;
let total = (Number(rule.baseCredits || 0) + Number(rule.perSecondCredits || 0) * videoDuration) * Number(rule.ratio || 1);
if (productImage) {
total += (Number(rule.inputImageBaseCredits || 0) + Number(rule.inputImagePerImageCredits || 0)) * Number(rule.inputImageRatio || 1);
}
return Number(total.toFixed(2));
}, [creditRules, engineId, productImage, videoDuration, videoResolution]);
const creditsInsufficient = Number(user?.credits || 0) < estimatedCredits;
const handleEngineChange = (value: string) => {
setEngineId(value);
const engine = videoEngines.find((item: any) => item.id === value);
if (!engine) return;
setVideoAspectRatio(engine.supportedRatios?.[0] || '16:9');
setVideoResolution(engine.supportedResolutions?.[0] || '480p');
setVideoDuration(engine.supportedDurations?.[0] || 5);
};
const autoSplitButtonText = useMemo(() => { const autoSplitButtonText = useMemo(() => {
if (tableData.length === 0) { if (tableData.length === 0) {
return 'AI自动拆分'; return 'AI自动拆分';
@@ -292,13 +347,6 @@ function RemoveInfo() {
setProductImageResourceId(''); setProductImageResourceId('');
}; };
const handleProductImageChange: any = (info: any) => {
if (info.fileList.length === 0) {
setProductImage('');
setProductImageResourceId('');
}
};
const beforeUploadProductImage = async (file: File) => { const beforeUploadProductImage = async (file: File) => {
try { try {
const uploadResult = await uploadShotReplicateImage(file); const uploadResult = await uploadShotReplicateImage(file);
@@ -316,44 +364,53 @@ function RemoveInfo() {
message.warning('请先选择拆镜片段'); message.warning('请先选择拆镜片段');
return; return;
} }
if (!productImage) { if (!engineId) {
message.warning('请上传产品图'); message.warning('请选择视频引擎');
return; return;
} }
if (!productName.trim()) { if (!selectedEngineSupportsImage) {
message.warning('请输入产品名称'); message.warning('当前视频引擎不支持参考图片,请移除素材图片或切换引擎');
return; return;
} }
if (!productSellingPoint.trim()) { if (creditsInsufficient) {
message.warning('请输入产品卖点'); message.warning('积分不足,请充值后再创建复刻项目');
return; return;
} }
const existingIdempotency = createIdempotencyRef.current;
const idempotencyKey = existingIdempotency?.segmentId === currentSegment
? existingIdempotency.key
: `shot_v2_${currentSegment}_${Date.now()}_${Math.random().toString(16).slice(2)}`;
createIdempotencyRef.current = { segmentId: currentSegment, key: idempotencyKey };
const params = { const params = {
target_project_name: productName.trim(), target_project_name: productName.trim(),
core_content_point: productSellingPoint.trim(), core_content_point: productSellingPoint.trim(),
material_image_url: productImage, project_description: productSellingPoint.trim() || undefined,
material_image_url: productImage || undefined,
material_image_resource_id: productImageResourceId || undefined, material_image_resource_id: productImageResourceId || undefined,
idempotency_key: `replication_${Date.now()}`, video_config: {
engine_id: engineId,
duration: videoDuration,
aspect_ratio: videoAspectRatio,
resolution: videoResolution,
},
idempotency_key: idempotencyKey,
}; };
setLoading(true); setLoading(true);
try { try {
await removeCreate(currentSegment, params); const result = await removeCreate(currentSegment, params);
message.success('视频生成任务创建成功'); const projectId = result?.projectId || result?.id || result?.detail?.id;
createIdempotencyRef.current = null;
message.success('项目已创建,视频提词正在生成');
handleCloseDrawer(); handleCloseDrawer();
Removelist(creatID).then((res: any) => { await fetchSegments();
if (projectId) {
const targetItem = res.items.find((item: any) => item.id === currentSegment); navigate(`/removelens/${projectId}/removefenbu?flow_version=v2`);
message.loading('创建中...', 3); }
setTimeout(() => {
navigate(`/removelens/${targetItem.moduleProjectId}/removefenbu`);
}, 3000);
});
// fetchSegments().then((res: any) => { // fetchSegments().then((res: any) => {
// console.log('123123123123',res); // console.log('123123123123',res);
@@ -677,7 +734,7 @@ function RemoveInfo() {
{record.moduleProjectId ? ( {record.moduleProjectId ? (
<Button <Button
type="text" type="text"
onClick={() => navigate(`/removelens/${record.moduleProjectId}/removefenbu`)} onClick={() => navigate(`/removelens/${record.moduleProjectId}/removefenbu?flow_version=${record.moduleProjectFlowVersion === 'v2' ? 'v2' : 'v1'}`)}
style={{ color: '#10b981', fontSize: 12, padding: 0, display: 'flex', alignItems: 'center', gap: 4 }} style={{ color: '#10b981', fontSize: 12, padding: 0, display: 'flex', alignItems: 'center', gap: 4 }}
> >
@@ -1130,65 +1187,47 @@ function RemoveInfo() {
<div style={{ display: 'flex', flexDirection: 'column', gap: 20 }}> <div style={{ display: 'flex', flexDirection: 'column', gap: 20 }}>
<div style={{ background: '#fff', borderRadius: 12, padding: 16, border: '1px solid rgba(99, 102, 241, 0.1)' }}> <div style={{ background: '#fff', borderRadius: 12, padding: 16, border: '1px solid rgba(99, 102, 241, 0.1)' }}>
<label style={{ fontWeight: 500, color: '#334155', marginBottom: 8, display: 'block', fontSize: 14 }}> <label style={{ fontWeight: 500, color: '#334155', marginBottom: 8, display: 'block', fontSize: 14 }}>
<span style={{ color: '#ef4444' }}>*</span>
</label> </label>
<div style={{ display: 'flex', gap: 12 }}> {productImage ? (
<div style={{ width: 120, height: 120, borderRadius: 8, border: '2px dashed rgba(99, 102, 241, 0.3)', background: 'rgba(99, 102, 241, 0.02)', position: 'relative', overflow: 'hidden' }}> <div style={{ width: 120, height: 120, borderRadius: 8, position: 'relative', overflow: 'hidden' }}>
{productImage ? ( <img src={buildAssetUrl(productImage)} alt="素材图" style={{ width: '100%', height: '100%', objectFit: 'cover' }} />
<> <button onClick={() => { setProductImage(''); setProductImageResourceId(''); }} style={{ position: 'absolute', bottom: 4, right: 4, width: 24, height: 24, borderRadius: '50%', backgroundColor: 'rgba(0,0,0,.5)', border: 'none', cursor: 'pointer' }}>
<img <DeleteOutlined style={{ color: '#fff', fontSize: 12 }} />
src={buildAssetUrl(productImage)} </button>
alt="产品图"
style={{ width: '100%', height: '100%', objectFit: 'cover' }}
/>
<button
onClick={() => {
setProductImage('');
setProductImageResourceId('');
}}
style={{
position: 'absolute',
bottom: 4,
right: 4,
width: 24,
height: 24,
borderRadius: '50%',
backgroundColor: 'rgba(0, 0, 0, 0.5)',
border: 'none',
cursor: 'pointer',
display: 'flex',
alignItems: 'center',
justifyContent: 'center',
}}
>
<DeleteOutlined style={{ color: '#fff', fontSize: 12 }} />
</button>
</>
) : (
<Upload
beforeUpload={beforeUploadProductImage}
maxCount={1}
accept="image/*"
style={{ width: '100%', height: '100%' }}
>
<div style={{ display: 'flex', flexDirection: 'column', alignItems: 'center', justifyContent: 'center', gap: 6, width: '100%', height: '100%' }}>
<PlusOutlined style={{ fontSize: 20, color: '#6366f1' }} />
<span style={{ fontSize: 12, color: '#64748b' }}></span>
</div>
</Upload>
)}
</div> </div>
</div> ) : (
<UploadSelector
accept="image/*"
multiple={false}
mediaType="image"
maxImageCount={1}
usedImageCount={0}
onLocalSelect={(files) => { if (files[0]) void beforeUploadProductImage(files[0]); }}
onHistorySelect={(items) => {
const item = items[0];
if (!item) return;
setProductImage(item.resourceUrl || item.previewUrl || item.displayUrl || '');
setProductImageResourceId(item.id);
message.success('历史图片已选择');
}}
>
<div style={{ width: 120, height: 120, borderRadius: 8, border: '2px dashed rgba(99,102,241,.3)', display: 'flex', flexDirection: 'column', alignItems: 'center', justifyContent: 'center', cursor: 'pointer', gap: 6 }}>
<PlusOutlined style={{ fontSize: 20, color: '#6366f1' }} />
<span style={{ fontSize: 12, color: '#64748b' }}></span>
</div>
</UploadSelector>
)}
</div> </div>
<div style={{ background: '#fff', borderRadius: 12, padding: 16, border: '1px solid rgba(99, 102, 241, 0.1)' }}> <div style={{ background: '#fff', borderRadius: 12, padding: 16, border: '1px solid rgba(99, 102, 241, 0.1)' }}>
<label style={{ fontWeight: 500, color: '#334155', marginBottom: 8, display: 'block', fontSize: 14 }}> <label style={{ fontWeight: 500, color: '#334155', marginBottom: 8, display: 'block', fontSize: 14 }}>
<span style={{ color: '#ef4444' }}>*</span>
</label> </label>
<Input <Input
value={productName} value={productName}
onChange={(e) => setProductName(e.target.value)} onChange={(e) => setProductName(e.target.value)}
placeholder="请输入产品名称" placeholder="请输入项目名称"
style={{ height: 44, borderRadius: 8, border: '1px solid rgba(99, 102, 241, 0.2)' }} style={{ height: 44, borderRadius: 8, border: '1px solid rgba(99, 102, 241, 0.2)' }}
maxLength={10} maxLength={10}
showCount showCount
@@ -1197,12 +1236,12 @@ function RemoveInfo() {
<div style={{ background: '#fff', borderRadius: 12, padding: 16, border: '1px solid rgba(99, 102, 241, 0.1)' }}> <div style={{ background: '#fff', borderRadius: 12, padding: 16, border: '1px solid rgba(99, 102, 241, 0.1)' }}>
<label style={{ fontWeight: 500, color: '#334155', marginBottom: 8, display: 'block', fontSize: 14 }}> <label style={{ fontWeight: 500, color: '#334155', marginBottom: 8, display: 'block', fontSize: 14 }}>
<span style={{ color: '#ef4444' }}>*</span>
</label> </label>
<TextArea <TextArea
value={productSellingPoint} value={productSellingPoint}
onChange={(e) => setProductSellingPoint(e.target.value)} onChange={(e) => setProductSellingPoint(e.target.value)}
placeholder="请输入产品卖点" placeholder="请输入项目描述"
style={{ borderRadius: 8, border: '1px solid rgba(99, 102, 241, 0.2)' }} style={{ borderRadius: 8, border: '1px solid rgba(99, 102, 241, 0.2)' }}
maxLength={100} maxLength={100}
showCount showCount
@@ -1210,12 +1249,37 @@ function RemoveInfo() {
/> />
</div> </div>
<div style={{ background: '#fff', borderRadius: 12, padding: 16, border: '1px solid rgba(99, 102, 241, 0.1)' }}>
<label style={{ fontWeight: 500, color: '#334155', marginBottom: 8, display: 'block', fontSize: 14 }}></label>
<Select
value={engineId || undefined}
onChange={handleEngineChange}
placeholder="选择视频引擎"
style={{ width: '100%', marginBottom: 12 }}
options={videoEngines.map((item: any) => ({
value: item.id,
label: productImage && !supportsReferenceImage(item) ? `${item.name}(不支持参考图)` : item.name,
disabled: Boolean(productImage) && !supportsReferenceImage(item),
}))}
/>
<div style={{ display: 'flex', gap: 8, marginBottom: 12 }}>
<Select value={videoDuration} onChange={setVideoDuration} style={{ flex: 1 }} options={(selectedEngine?.supportedDurations || [5]).map((item: number) => ({ value: item, label: `${item}` }))} />
<Select value={videoAspectRatio} onChange={setVideoAspectRatio} style={{ flex: 1 }} options={(selectedEngine?.supportedRatios || ['16:9']).map((item: string) => ({ value: item, label: item }))} />
<Select value={videoResolution} onChange={setVideoResolution} style={{ flex: 1 }} options={(selectedEngine?.supportedResolutions || ['480p']).map((item: string) => ({ value: item, label: item }))} />
</div>
<div style={{ color: creditsInsufficient || !selectedEngineSupportsImage ? '#ef4444' : '#6366f1', fontSize: 13 }}>
{selectedEngineSupportsImage
? `预估积分:${estimatedCredits},当前积分:${Number(user?.credits || 0).toFixed(2)}`
: '当前引擎不支持参考图片,请移除素材图片或切换引擎。'}
</div>
</div>
<div style={{ display: 'flex', gap: 12, marginTop: 8 }}> <div style={{ display: 'flex', gap: 12, marginTop: 8 }}>
<Button <Button
type="primary" type="primary"
onClick={handleManualGenerate} onClick={handleManualGenerate}
loading={loading} loading={loading}
disabled={loading} disabled={loading || creditsInsufficient || !engineId || !selectedEngineSupportsImage}
style={{ style={{
flex: 1, flex: 1,
height: 48, height: 48,
+236 -65
View File
@@ -1,8 +1,8 @@
import React, { useState, useEffect } from 'react'; import React, { useState, useEffect, useRef } from 'react';
import { Button, Typography, Collapse, Space, Modal, Input, Table, message, Tooltip } from 'antd'; import { Button, Typography, Collapse, Space, Modal, Input, Table, message, Tooltip } from 'antd';
import { ArrowLeftOutlined, PlayCircleOutlined, CheckCircleOutlined, EditOutlined, DownloadOutlined, SettingOutlined, LayoutOutlined } from '@ant-design/icons'; import { ArrowLeftOutlined, PlayCircleOutlined, CheckCircleOutlined, EditOutlined, DownloadOutlined, SettingOutlined, LayoutOutlined } from '@ant-design/icons';
import { useNavigate, useParams } from 'react-router-dom'; import { useNavigate, useParams, useSearchParams } from 'react-router-dom';
import { getShotReplicationList, removeDetail, removeone, removetwo, removethree, removefour, getEngine, updateShotImagePrompt, updateShotVideoPromptSchema, calculateCredits } from '../api/index'; import { getShotReplicationList, removeDetail, removeone, removetwo, removethree, removefour, getEngine, updateShotImagePrompt, updateShotVideoPromptSchema, calculateCredits, retryShotVideoPromptV2, updateShotVideoPromptSchemaV2, generateShotVideoV2 } from '../api/index';
import { useAuthStore } from '../store/useAuthStore'; import { useAuthStore } from '../store/useAuthStore';
import VideoPromptSchemaEditor from '../components/VideoPromptSchemaEditor'; import VideoPromptSchemaEditor from '../components/VideoPromptSchemaEditor';
import { validateVideoPromptSchemaByConfig } from '../utils/videoPromptSchema'; import { validateVideoPromptSchemaByConfig } from '../utils/videoPromptSchema';
@@ -25,6 +25,8 @@ const buildMediaUrl = (url: string): string => {
function InitialInfo() { function InitialInfo() {
const navigate = useNavigate(); const navigate = useNavigate();
const { creatID } = useParams<{ creatID: string }>(); const { creatID } = useParams<{ creatID: string }>();
const [searchParams] = useSearchParams();
const flowVersion: 'v1' | 'v2' = searchParams.get('flow_version') === 'v2' ? 'v2' : 'v1';
const { user } = useAuthStore(); const { user } = useAuthStore();
const [modalVisible, setModalVisible] = useState(false); const [modalVisible, setModalVisible] = useState(false);
@@ -35,9 +37,12 @@ function InitialInfo() {
const [currentType, setCurrentType] = useState<string>('image'); const [currentType, setCurrentType] = useState<string>('image');
const [formData, setFormData] = useState<any>({}); const [formData, setFormData] = useState<any>({});
const [videoSchemaConfigSnapshot, setVideoSchemaConfigSnapshot] = useState<any>(null); const [videoSchemaConfigSnapshot, setVideoSchemaConfigSnapshot] = useState<any>(null);
const [pollingTimer, setPollingTimer] = useState<any>(null); const pollingTimerRef = useRef<ReturnType<typeof setInterval> | null>(null);
const [editingPromptStepId, setEditingPromptStepId] = useState(''); const [editingPromptStepId, setEditingPromptStepId] = useState('');
const [promptSaving, setPromptSaving] = useState(false); const [promptSaving, setPromptSaving] = useState(false);
const [retryPromptModalVisible, setRetryPromptModalVisible] = useState(false);
const [retryPromptStepId, setRetryPromptStepId] = useState<string>('');
const [retryPromptSubmitting, setRetryPromptSubmitting] = useState(false);
// 引擎和视频参数相关状态 // 引擎和视频参数相关状态
const [enginesele, setEnginesele] = useState<any>({}); const [enginesele, setEnginesele] = useState<any>({});
@@ -111,14 +116,21 @@ function InitialInfo() {
} }
}, [previewVisible, previewType]); }, [previewVisible, previewType]);
// const isV2 = flowVersion === 'v2';
const baseSteps = [
{ id: 1, title: '原始素材', description: '上传原始视频素材', childId: 1 }, const baseSteps = isV2
{ id: 2, title: '生成提示词', description: '根据素材生成描述词', childId: 2 }, ? [
{ id: 3, title: '生成产品融合图', description: '生成产品与场景融合图', childId: 3 }, { id: 1, title: '素材与项目信息', description: '固定片段素材和项目描述', childId: 1 },
{ id: 4, title: '生成视频提示词', description: '生成视频生成提示词', childId: 4 }, { id: 2, title: '生成视频提示词', description: '生成并确认视频提示词', childId: 4 },
{ id: 5, title: '生成最终视频', description: '合成最终视频', childId: 5 }, { id: 3, title: '生成最终视频', description: '使用当前视频提示词生成视频', childId: 5 },
]; ]
: [
{ id: 1, title: '原始素材', description: '上传原始视频素材', childId: 1 },
{ id: 2, title: '生成提示词', description: '根据素材生成描述词', childId: 2 },
{ id: 3, title: '生成产品融合图', description: '生成产品与场景融合图', childId: 3 },
{ id: 4, title: '生成视频提示词', description: '生成视频生成提示词', childId: 4 },
{ id: 5, title: '生成最终视频', description: '合成最终视频', childId: 5 },
];
// 合并基础步骤和API返回的状态 // 合并基础步骤和API返回的状态
const steps = baseSteps.map((step, index) => ({ const steps = baseSteps.map((step, index) => ({
@@ -223,12 +235,12 @@ function InitialInfo() {
// 组件卸载时清理定时器 // 组件卸载时清理定时器
useEffect(() => { useEffect(() => {
return () => { return () => {
if (pollingTimer) { if (pollingTimerRef.current) {
clearInterval(pollingTimer); clearInterval(pollingTimerRef.current);
setPollingTimer(null); pollingTimerRef.current = null;
} }
}; };
}, [pollingTimer]); }, []);
// 点击外部关闭弹窗 // 点击外部关闭弹窗
useEffect(() => { useEffect(() => {
@@ -250,39 +262,30 @@ function InitialInfo() {
if (creatID) { if (creatID) {
removeDetail(creatID).then((res: any) => { removeDetail(creatID, flowVersion).then((res: any) => {
setTaskDetail(res); setTaskDetail(res);
if (res.steps) { if (res.steps) {
setApiSteps(res.steps); setApiSteps(res.steps);
} }
}).catch((error: any) => { }).catch((error: any) => {
message.error(error?.message || '加载任务详情失败');
}); });
} }
}, [creatID]); }, [creatID, flowVersion]);
// 监听合并后的 steps 数据,判断第五步状态并启动/停止轮询 // 仅在项目或任一步骤处于 processing 时轮询;waiting_user 必须停止。
useEffect(() => { useEffect(() => {
const shouldPoll = taskDetail?.status === 'processing'
|| apiSteps.some((step: any) => step?.status === 'processing');
// 检查第五步的状态(索引 4 if (shouldPoll && !pollingTimerRef.current) {
const fifthStep = steps[4]; pollingTimerRef.current = setInterval(pollTaskDetail, 30000);
} else if (!shouldPoll && pollingTimerRef.current) {
if (fifthStep && fifthStep.status !== 'completed' && fifthStep.status !== 'failed') { clearInterval(pollingTimerRef.current);
// 第五步未完成,启动轮询 pollingTimerRef.current = null;
if (!pollingTimer) {
const timer = setInterval(pollTaskDetail, 30000);
setPollingTimer(timer);
} else {
}
} else {
// 第五步已完成或失败,停止轮询
if (fifthStep) {
if (pollingTimer) {
clearInterval(pollingTimer);
setPollingTimer(null);
}
}
} }
}, [steps]); }, [taskDetail?.status, apiSteps, creatID, flowVersion]);
useEffect(() => { useEffect(() => {
calculateCredits().then((data: any) => { calculateCredits().then((data: any) => {
@@ -320,7 +323,9 @@ function InitialInfo() {
total += inputVideoCost; total += inputVideoCost;
} }
const inputImageCount = taskDetail?.videoGeneration?.inputMedia?.image?.length || 0; const inputImageCount = isV2
? (taskDetail?.material?.materialImageUrl ? 1 : 0)
: (taskDetail?.videoGeneration?.inputMedia?.image?.length || 0);
if (inputImageCount > 0) { if (inputImageCount > 0) {
const inputImageCost = ((config.inputImageBaseCredits || 0) + (config.inputImagePerImageCredits || 0) * inputImageCount) * (config.inputImageRatio || 1); const inputImageCost = ((config.inputImageBaseCredits || 0) + (config.inputImagePerImageCredits || 0) * inputImageCount) * (config.inputImageRatio || 1);
total += inputImageCost; total += inputImageCost;
@@ -399,7 +404,8 @@ function InitialInfo() {
setPromptSaving(true); setPromptSaving(true);
try { try {
await updateShotVideoPromptSchema(taskDetail.id, editingPromptStepId, { const updateVideoPrompt = isV2 ? updateShotVideoPromptSchemaV2 : updateShotVideoPromptSchema;
await updateVideoPrompt(taskDetail.id, editingPromptStepId, {
prompt_schema: formData, prompt_schema: formData,
}); });
message.success('视频提示词已保存'); message.success('视频提示词已保存');
@@ -419,12 +425,13 @@ function InitialInfo() {
return; return;
} }
removeDetail(creatID).then((res: any) => { removeDetail(creatID, flowVersion).then((res: any) => {
setTaskDetail(res); setTaskDetail(res);
if (res.steps) { if (res.steps) {
setApiSteps(res.steps); setApiSteps(res.steps);
} }
}).catch((error: any) => { }).catch((error: any) => {
console.error('轮询任务详情失败', error);
}); });
}; };
@@ -432,12 +439,13 @@ function InitialInfo() {
const refreshTaskDetail = () => { const refreshTaskDetail = () => {
if (!creatID) return; if (!creatID) return;
removeDetail(creatID).then((res: any) => { removeDetail(creatID, flowVersion).then((res: any) => {
setTaskDetail(res); setTaskDetail(res);
if (res.steps) { if (res.steps) {
setApiSteps(res.steps); setApiSteps(res.steps);
} }
}).catch((error: any) => { }).catch((error: any) => {
message.error(error?.message || '刷新任务详情失败');
}); });
}; };
@@ -540,34 +548,109 @@ function InitialInfo() {
// 这里可以添加下一步的逻辑,比如调用接口等 // 这里可以添加下一步的逻辑,比如调用接口等
}; };
const createvideo = (stepId: number, engineId: string) => { const createvideo = async (stepId: number, engineId: string) => {
let params = { try {
engine_id: '', if (isV2) {
} await generateShotVideoV2(taskDetail.id, String(stepId));
} else {
await removefour(taskDetail.id, String(stepId), { engine_id: engineId || '' });
removefour(taskDetail.id, stepId.toString(), params).then((res: any) => { }
// 重新获取任务详情以更新数据
message.info('正在生成视频,请稍候...'); message.info('正在生成视频,请稍候...');
refreshTaskDetail(); refreshTaskDetail();
}).catch((error: any) => { } catch (error: any) {
const errorMsg = error?.message?.split(': ')?.[1] || error?.message || '生成失败'; const errorMsg = error?.message?.split(': ')?.[1] || error?.message || '生成失败';
message.error(errorMsg); message.error(errorMsg);
});
}
const agincreatevideo = () => {
let params = {
engine_id: '',
} }
};
const agincreatevideo = async () => {
const promptStep = isV2 ? steps[1] : steps[3];
if (!promptStep?.id) return;
await createvideo(promptStep.id, promptStep.engineId || '');
};
removefour(taskDetail.id, steps[3].id.toString(), params).then((res: any) => { const applyRetryVideoEngine = (
// 重新获取任务详情以更新数据 engineId: string,
preferred?: { duration?: number; aspectRatio?: string; resolution?: string },
) => {
const engine = (enginesele.video || []).find((item: any) => String(item.id) === String(engineId));
if (!engine) {
return false;
}
const ratios = Array.isArray(engine.supportedRatios) && engine.supportedRatios.length > 0
? engine.supportedRatios.map((item: any) => String(item))
: ['16:9', '4:3', '1:1', '3:4', '9:16', '21:9'];
const resolutions = Array.isArray(engine.supportedResolutions) && engine.supportedResolutions.length > 0
? engine.supportedResolutions.map((item: any) => String(item))
: ['480p', '720p', '1080p'];
const parsedDurations = Array.isArray(engine.supportedDurations)
? engine.supportedDurations
.map((item: any) => Number(item))
.filter((item: number) => Number.isFinite(item) && item > 0)
: [];
const durations = parsedDurations.length > 0 ? parsedDurations : [5, 8, 10, 12, 15];
const preferredDuration = Number(preferred?.duration ?? videoDuration);
const preferredRatio = String(preferred?.aspectRatio || videoAspectRatio || '');
const preferredResolution = String(preferred?.resolution || videoResolution || '');
setCountType(String(engine.id));
setEngineOptions({ ratios, resolutions, durations });
setVideoDuration(durations.includes(preferredDuration) ? preferredDuration : durations[0]);
setVideoAspectRatio(ratios.includes(preferredRatio) ? preferredRatio : ratios[0]);
setVideoResolution(resolutions.includes(preferredResolution) ? preferredResolution : resolutions[0]);
return true;
};
const openRegenerateVideoPrompt = (stepId: number) => {
if (!isV2 || !taskDetail?.id) return;
const currentConfig = taskDetail?.videoGeneration?.promptParams || taskDetail?.videoGeneration?.params || {};
const requestedEngineId = String(
currentConfig.engineId
|| currentConfig.engine_id
|| taskDetail?.videoGeneration?.engineId
|| countType
|| '',
);
const currentEngineId = String(
(enginesele.video || []).some((engine: any) => String(engine.id) === requestedEngineId)
? requestedEngineId
: enginesele.video?.[0]?.id || '',
);
if (!currentEngineId || !applyRetryVideoEngine(currentEngineId, {
duration: Number(currentConfig.duration || videoDuration),
aspectRatio: String(currentConfig.aspectRatio || currentConfig.aspect_ratio || videoAspectRatio),
resolution: String(currentConfig.resolution || videoResolution),
})) {
message.warning('当前没有可用的视频生成引擎');
return;
}
setRetryPromptStepId(String(stepId));
setRetryPromptModalVisible(true);
};
const submitRegenerateVideoPrompt = async () => {
if (!isV2 || !taskDetail?.id || !retryPromptStepId || !countType) return;
setRetryPromptSubmitting(true);
try {
await retryShotVideoPromptV2(taskDetail.id, retryPromptStepId, {
video_config: {
engine_id: countType,
duration: videoDuration,
aspect_ratio: videoAspectRatio,
resolution: videoResolution,
},
});
message.info('正在按新视频参数重新生成视频提示词,请稍候...');
setRetryPromptModalVisible(false);
setRetryPromptStepId('');
refreshTaskDetail(); refreshTaskDetail();
}).catch((error: any) => { } catch (error: any) {
}); message.error(error?.message || '重新生成视频提示词失败');
} } finally {
setRetryPromptSubmitting(false);
}
};
@@ -1298,6 +1381,11 @@ function InitialInfo() {
> >
/ /
</Button> </Button>
{isV2 && (
<Button type="default" onClick={() => openRegenerateVideoPrompt(step.id)} style={{ flex: 1, borderRadius: 10, borderColor: 'rgba(99, 102, 241, 0.3)', color: '#6366f1', height: 36, fontWeight: 500 }} disabled={step.status === 'processing' || steps[2]?.status === 'processing'}>
</Button>
)}
<Button onClick={() => { createvideo(step.id, step.engineId); }} type="primary" style={{ flex: 1, borderRadius: 10, background: 'linear-gradient(135deg, #6366f1 0%, #8b5cf6 100%)', border: 'none', height: 36, fontWeight: 500, boxShadow: '0 4px 12px rgba(99, 102, 241, 0.3)' }} disabled={step.status !== 'completed'}> <Button onClick={() => { createvideo(step.id, step.engineId); }} type="primary" style={{ flex: 1, borderRadius: 10, background: 'linear-gradient(135deg, #6366f1 0%, #8b5cf6 100%)', border: 'none', height: 36, fontWeight: 500, boxShadow: '0 4px 12px rgba(99, 102, 241, 0.3)' }} disabled={step.status !== 'completed'}>
</Button> </Button>
@@ -1322,8 +1410,8 @@ function InitialInfo() {
)} )}
</div> </div>
<Space style={{ width: '100%', gap: 12 }}> <Space style={{ width: '100%', gap: 12 }}>
<Button onClick={() => agincreatevideo()} type="default" icon={<EditOutlined />} style={{ flex: 1, borderRadius: 10, borderColor: 'rgba(99, 102, 241, 0.3)', color: '#6366f1', height: 36, fontWeight: 500, background: 'rgba(99, 102, 241, 0.04)' }} disabled={step.status !== 'completed'}> <Button onClick={() => agincreatevideo()} type="default" icon={<EditOutlined />} style={{ flex: 1, borderRadius: 10, borderColor: 'rgba(99, 102, 241, 0.3)', color: '#6366f1', height: 36, fontWeight: 500, background: 'rgba(99, 102, 241, 0.04)' }} disabled={isV2 ? !['completed', 'failed'].includes(step.status) : step.status !== 'completed'}>
{step.status === 'failed' ? '重试生成' : '重新生成'}
</Button> </Button>
<Button <Button
type="primary" type="primary"
@@ -1353,6 +1441,89 @@ function InitialInfo() {
</div> </div>
</div> </div>
</div> </div>
<Modal
title="重新生成视频提词"
open={retryPromptModalVisible}
onCancel={() => {
if (!retryPromptSubmitting) {
setRetryPromptModalVisible(false);
setRetryPromptStepId('');
}
}}
onOk={submitRegenerateVideoPrompt}
confirmLoading={retryPromptSubmitting}
okText="按新参数生成提词"
cancelText="取消"
width={720}
destroyOnClose
>
<div style={{ display: 'flex', flexDirection: 'column', gap: 20 }}>
<div>
<Text strong style={{ display: 'block', marginBottom: 10 }}></Text>
<div style={{ display: 'grid', gridTemplateColumns: 'repeat(2, minmax(0, 1fr))', gap: 10 }}>
{(enginesele.video || []).map((engine: any) => (
<button
key={engine.id}
type="button"
onClick={() => applyRetryVideoEngine(String(engine.id), {
duration: videoDuration,
aspectRatio: videoAspectRatio,
resolution: videoResolution,
})}
style={{
minHeight: 46,
padding: '8px 12px',
borderRadius: 8,
border: countType === String(engine.id) ? '2px solid #6366f1' : '1px solid #e5e7eb',
background: countType === String(engine.id) ? 'rgba(99,102,241,0.08)' : '#fff',
color: countType === String(engine.id) ? '#4f46e5' : '#374151',
cursor: 'pointer',
textAlign: 'left',
fontWeight: countType === String(engine.id) ? 600 : 400,
}}
>
{engine.name || engine.id}
</button>
))}
</div>
</div>
<div>
<Text strong style={{ display: 'block', marginBottom: 10 }}></Text>
<div style={{ display: 'flex', flexWrap: 'wrap', gap: 8 }}>
{engineOptions.durations.map((duration) => (
<Button key={duration} type={videoDuration === duration ? 'primary' : 'default'} onClick={() => setVideoDuration(duration)}>
{duration}
</Button>
))}
</div>
</div>
<div>
<Text strong style={{ display: 'block', marginBottom: 10 }}></Text>
<div style={{ display: 'flex', flexWrap: 'wrap', gap: 8 }}>
{engineOptions.ratios.map((ratio) => (
<Button key={ratio} type={videoAspectRatio === ratio ? 'primary' : 'default'} onClick={() => setVideoAspectRatio(ratio)}>
{ratio}
</Button>
))}
</div>
</div>
<div>
<Text strong style={{ display: 'block', marginBottom: 10 }}></Text>
<div style={{ display: 'flex', flexWrap: 'wrap', gap: 8 }}>
{engineOptions.resolutions.map((resolution) => (
<Button key={resolution} type={videoResolution === resolution ? 'primary' : 'default'} onClick={() => setVideoResolution(resolution)}>
{resolution}
</Button>
))}
</div>
</div>
<div style={{ padding: '10px 12px', borderRadius: 8, background: '#f8fafc', color: '#64748b' }}>
{enginesele.video?.find((engine: any) => String(engine.id) === countType)?.name || countType}
{' · '}{videoDuration} · {videoAspectRatio} · {videoResolution}
<span style={{ marginLeft: 12 }}>{estimatedCredits}</span>
</div>
</div>
</Modal>
<Modal <Modal
title={ title={
<div style={{ display: 'flex', alignItems: 'center', gap: 10 }}> <div style={{ display: 'flex', alignItems: 'center', gap: 10 }}>
+27 -1
View File
@@ -169,7 +169,10 @@ export interface GenerationRecord {
imageProportion: string; imageProportion: string;
imagePx: string; imagePx: string;
imageUrl: string; imageUrl: string;
engineId?: string;
engineName?: string;
engineSnapshot?: Record<string, any>;
includeMediaReferences?: boolean;
} }
export interface OptimizeParams { export interface OptimizeParams {
@@ -189,8 +192,31 @@ export interface OptimizeParams {
} }
export interface GenerateParams { export interface GenerateParams {
engineId?: string;
includeMediaReferences?: boolean;
aspectRatio?: AspectRatio; aspectRatio?: AspectRatio;
resolution?: Resolution; resolution?: Resolution;
imageSize?: string;
}
export type ModuleGenerationFlowVersion = 'v1' | 'v2';
export interface ModuleGenerationVideoConfigV2 {
engineId: string;
duration: number;
aspectRatio: string;
resolution: string;
}
export interface ModuleGenerationProjectV2Create {
materialImageUrl?: string;
materialImageResourceId?: string;
targetProjectName?: string;
coreContentPoint?: string;
projectDescription?: string;
targetPlatform?: string;
videoConfig: ModuleGenerationVideoConfigV2;
idempotencyKey?: string;
} }
export interface OptimizeResult { export interface OptimizeResult {