This commit is contained in:
2026-07-15 13:51:21 +08:00
parent a9190ba4e1
commit 6db989dc42
65 changed files with 4499 additions and 1515 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-tEZsB6Nw.js"></script> <script type="module" crossorigin src="/assets/index-CKURqRU_.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>
@@ -0,0 +1,66 @@
import React from 'react';
import { Empty, Spin, Tag, Typography } from 'antd';
import { PlayCircleFilled } from '@ant-design/icons';
import type { GenerationAITaskOut } from '../../types';
interface Props {
task: GenerationAITaskOut;
resolveUrl: (url?: string | null) => string;
onPreview: (url: string, type: 'image' | 'video', title: string) => void;
}
const spanByCount = (count: number, index: number): number => {
if (count <= 1) return 6;
if (count === 2 || count === 4) return 3;
if (count === 3) return index < 2 ? 3 : 6;
return index < 3 ? 2 : 3;
};
const LABELS: Record<string, string> = {
pending: '待处理', queued: '已入队', preparing: '准备中', generating: '生成中',
creating_provider_task: '创建任务中', waiting_remote: '等待生成', polling: '轮询中',
result_ready: '结果就绪', download_queued: '等待下载', downloading: '下载中',
retry_waiting: '等待重试', completed: '已完成', failed: '生成失败',
download_failed: '下载失败', deleted: '已删除',
};
const GenerationTaskResourceGrid: React.FC<Props> = ({ task, resolveUrl, onPreview }) => {
const count = Math.max(1, Math.min(5, Number(task.generationCount || task.childItems?.length || 1)));
const sortedChildren = [...(task.childItems || [])].sort((a, b) => Number(a.generationIndex || 0) - Number(b.generationIndex || 0));
const items: GenerationAITaskOut[] = sortedChildren.length
? sortedChildren
: (count > 1
? Array.from({ length: count }, (_, index) => ({ ...task, id: `${task.id}-${index + 1}`, generationIndex: index + 1, childItems: [] }))
: [task]);
return (
<div style={{ width: '100%', height: 430, display: 'grid', gridTemplateColumns: 'repeat(6, minmax(0,1fr))', gridAutoRows: 'minmax(0,1fr)', gap: items.length > 1 ? 8 : 0 }}>
{items.map((item, index) => {
const status = item.displayStatus || item.pipelineStage || item.status || 'pending';
const isVideo = item.genType === 'video';
const resultUrl = resolveUrl(isVideo ? item.videoUrl : item.imageUrl);
const coverUrl = resolveUrl(item.videoCoverUrl);
const active = ['pending', 'queued', 'preparing', 'generating', 'creating_provider_task', 'waiting_remote', 'polling', 'result_ready', 'download_queued', 'downloading', 'retry_waiting'].includes(status);
return (
<div key={item.id} style={{ gridColumn: `span ${spanByCount(items.length, index)}`, minWidth: 0, minHeight: 0, border: '1px solid #edf0f5', borderRadius: 10, overflow: 'hidden', position: 'relative', background: '#f8f9fc' }}>
{resultUrl && status !== 'deleted' ? (
<button type="button" onClick={() => onPreview(resultUrl, isVideo ? 'video' : 'image', `生成结果 ${item.generationIndex || index + 1}`)} style={{ width: '100%', height: '100%', padding: 0, border: 0, background: 'transparent', cursor: 'pointer', position: 'relative' }}>
{isVideo ? (coverUrl ? <img src={coverUrl} alt="视频封面" style={{ width: '100%', height: '100%', objectFit: 'contain' }} /> : <video src={resultUrl} muted preload="metadata" style={{ width: '100%', height: '100%', objectFit: 'contain' }} />) : <img src={resultUrl} alt="生成图片" style={{ width: '100%', height: '100%', objectFit: 'contain' }} />}
{isVideo ? <PlayCircleFilled style={{ position: 'absolute', left: '50%', top: '50%', transform: 'translate(-50%,-50%)', color: '#fff', fontSize: 38, filter: 'drop-shadow(0 3px 8px rgba(0,0,0,.35))' }} /> : null}
</button>
) : (
<div style={{ width: '100%', height: '100%', display: 'flex', flexDirection: 'column', alignItems: 'center', justifyContent: 'center', gap: 9, padding: 12, textAlign: 'center' }}>
{active ? <Spin size="small" /> : <Empty image={Empty.PRESENTED_IMAGE_SIMPLE} description={null} />}
<Tag color={status === 'deleted' ? 'default' : (status === 'download_failed' || status === 'failed' ? 'error' : 'processing')}>{LABELS[status] || status}</Tag>
{item.errorMessage && !active ? <Typography.Text type="danger" style={{ fontSize: 11 }}>{item.errorMessage}</Typography.Text> : null}
</div>
)}
{items.length > 1 ? <span style={{ position: 'absolute', top: 6, left: 6, padding: '1px 7px', borderRadius: 10, color: '#fff', background: 'rgba(17,24,39,.58)', fontSize: 11 }}>#{item.generationIndex || index + 1}</span> : null}
</div>
);
})}
</div>
);
};
export default GenerationTaskResourceGrid;
@@ -32,6 +32,7 @@ import dayjs from 'dayjs';
import { getAdminGenerationAiTasks } from '../api'; import { getAdminGenerationAiTasks } from '../api';
import type { GenerationAIMediaReference, GenerationAITaskOut } from '../types'; import type { GenerationAIMediaReference, GenerationAITaskOut } from '../types';
import { formatDate } from '../utils/formatDate'; import { formatDate } from '../utils/formatDate';
import GenerationTaskResourceGrid from '../components/generation/GenerationTaskResourceGrid';
const { RangePicker } = DatePicker; const { RangePicker } = DatePicker;
@@ -71,6 +72,8 @@ const STATUS_MAP: Record<string, { color: string; text: string; icon: React.Reac
generating: { color: 'warning', text: '生成中', icon: <LoadingOutlined spin /> }, generating: { color: 'warning', text: '生成中', icon: <LoadingOutlined spin /> },
completed: { color: 'success', text: '已完成', icon: <CheckCircleOutlined /> }, completed: { color: 'success', text: '已完成', icon: <CheckCircleOutlined /> },
failed: { color: 'error', text: '失败', icon: <CloseCircleOutlined /> }, failed: { color: 'error', text: '失败', icon: <CloseCircleOutlined /> },
download_failed: { color: 'error', text: '下载失败', icon: <CloseCircleOutlined /> },
deleted: { color: 'default', text: '已删除', icon: <CloseCircleOutlined /> },
}; };
const PIPELINE_STAGE_MAP: Record<string, string> = { const PIPELINE_STAGE_MAP: Record<string, string> = {
@@ -425,6 +428,25 @@ const AdminGenerationAiRecords: React.FC = () => {
return <Tag color={cfg.color} icon={cfg.icon}>{cfg.text}</Tag>; return <Tag color={cfg.color} icon={cfg.icon}>{cfg.text}</Tag>;
}, },
}, },
{
title: '生成数量', key: 'generationCount', width: 150,
render: (_: any, r: GenerationAITaskOut) => {
const count = Math.max(1, Number(r.generationCount || 1));
if (count === 1) return <Tag>1</Tag>;
const children = r.childItems || [];
const completed = children.filter((item) => (item.displayStatus || item.status) === 'completed').length;
const failed = children.filter((item) => ['failed', 'download_failed'].includes(item.displayStatus || item.status)).length;
const deleted = children.filter((item) => (item.displayStatus || item.status) === 'deleted').length;
return (
<Space size={4} wrap>
<Tag color="purple">{count}</Tag>
<Typography.Text style={{ fontSize: 11, color: '#64748b' }}>
{completed}{failed ? ` / ${failed}失败` : ''}{deleted ? ` / ${deleted}删除` : ''}
</Typography.Text>
</Space>
);
},
},
{ {
title: '引擎', key: 'engine', width: 160, title: '引擎', key: 'engine', width: 160,
render: (_: any, r: GenerationAITaskOut) => { render: (_: any, r: GenerationAITaskOut) => {
@@ -1045,6 +1067,7 @@ const AdminGenerationAiRecords: React.FC = () => {
<InfoItem label="用户名称" value={preview.userName || '未知用户'} /> <InfoItem label="用户名称" value={preview.userName || '未知用户'} />
<InfoItem label="用户ID" value={preview.userId || '-'} /> <InfoItem label="用户ID" value={preview.userId || '-'} />
<InfoItem label="任务ID" value={preview.id} /> <InfoItem label="任务ID" value={preview.id} />
<InfoItem label="生成数量" value={`${preview.generationCount || 1}`} />
</div> </div>
<div> <div>
@@ -1122,38 +1145,16 @@ const AdminGenerationAiRecords: React.FC = () => {
</div> </div>
) : null} ) : null}
{preview.status === 'completed' ? ( <div>
<div> <Typography.Text style={{ fontSize: 12, color: '#94a3b8', display: 'block', marginBottom: 6 }}>
<div style={{ display: 'flex', justifyContent: 'space-between', alignItems: 'center', marginBottom: 6 }}> {preview.generationCount || 1}
<Typography.Text style={{ fontSize: 12, color: '#94a3b8', display: 'block' }}> </Typography.Text>
{preview.genType === 'video' ? '生成视频' : '生成图片'} <GenerationTaskResourceGrid
</Typography.Text> task={preview}
{preview.genType === 'video' && preview.videoUrl ? ( resolveUrl={apiUrl}
<Button onPreview={handlePreviewResource}
size="small" />
type="link" </div>
icon={<PlayCircleOutlined />}
onClick={() => handlePreviewResource(preview.videoUrl!, 'video', '生成视频')}
style={{ padding: 0 }}
>
</Button>
) : null}
{preview.genType === 'image' && preview.imageUrl ? (
<Button
size="small"
type="link"
icon={<FileImageOutlined />}
onClick={() => handlePreviewResource(preview.imageUrl!, 'image', '生成图片')}
style={{ padding: 0 }}
>
</Button>
) : null}
</div>
{preview.genType === 'video' ? renderResultVideo() : renderResultImage()}
</div>
) : null}
{preview.status === 'failed' && preview.errorMessage ? ( {preview.status === 'failed' && preview.errorMessage ? (
<div style={{ padding: 12, borderRadius: 10, background: 'rgba(239,68,68,0.04)', border: '1px solid rgba(239,68,68,0.15)' }}> <div style={{ padding: 12, borderRadius: 10, background: 'rgba(239,68,68,0.04)', border: '1px solid rgba(239,68,68,0.15)' }}>
@@ -1,6 +1,6 @@
import React, { useEffect, useState } from 'react'; import React, { useEffect, useState } from 'react';
import { import {
Button, Card, Checkbox, Form, Input, message, Modal, Popconfirm, Select, Space, Switch, Table, Tag, Typography, Button, Card, Checkbox, Form, Input, InputNumber, message, Modal, Popconfirm, Select, Space, Switch, Table, Tag, Typography,
} from 'antd'; } from 'antd';
import { import {
PictureOutlined, PlusOutlined, EditOutlined, DeleteOutlined, PictureOutlined, PlusOutlined, EditOutlined, DeleteOutlined,
@@ -21,6 +21,11 @@ interface ImageEngine {
generateUrl: string; generateUrl: string;
isActive: boolean; isActive: boolean;
priority: number; priority: number;
multiGenerationEnabled: boolean;
maxGenerationCount: number;
multiImageMaxImages: number;
maxReferenceImageCount: number;
outputFormat: '' | 'png' | 'jpeg';
} }
function parseJsonArray(val: unknown): any[] { function parseJsonArray(val: unknown): any[] {
@@ -80,6 +85,7 @@ const AdminImageEngines: React.FC = () => {
const [loading, setLoading] = useState(false); const [loading, setLoading] = useState(false);
const [modal, setModal] = useState<{ open: boolean; engine: ImageEngine | null }>({ open: false, engine: null }); const [modal, setModal] = useState<{ open: boolean; engine: ImageEngine | null }>({ open: false, engine: null });
const [form] = Form.useForm(); const [form] = Form.useForm();
const multiGenerationEnabled = Form.useWatch('multiGenerationEnabled', form) ?? false;
const load = async () => { const load = async () => {
setLoading(true); setLoading(true);
@@ -126,6 +132,11 @@ const AdminImageEngines: React.FC = () => {
generate_url: values.generateUrl || '', generate_url: values.generateUrl || '',
is_active: values.isActive ?? true, is_active: values.isActive ?? true,
priority: values.priority ?? 0, priority: values.priority ?? 0,
multi_generation_enabled: values.multiGenerationEnabled ?? false,
max_generation_count: values.maxGenerationCount ?? 1,
multi_image_max_images: values.multiImageMaxImages ?? 15,
max_reference_image_count: values.maxReferenceImageCount ?? 14,
output_format: values.outputFormat ?? '',
}; };
if (modal.engine) { if (modal.engine) {
await saveImageEngine({ id: modal.engine.id, ...payload }); await saveImageEngine({ id: modal.engine.id, ...payload });
@@ -168,6 +179,8 @@ const AdminImageEngines: React.FC = () => {
form.resetFields(); form.resetFields();
form.setFieldsValue({ form.setFieldsValue({
isActive: true, priority: 0, isActive: true, priority: 0,
multiGenerationEnabled: false, maxGenerationCount: 1, multiImageMaxImages: 15,
maxReferenceImageCount: 14, outputFormat: '',
supportedModels: ['doubao-seedream-5-0-260128'], supportedModels: ['doubao-seedream-5-0-260128'],
defaultSize: '2K', defaultSize: '2K',
maxImageCount: 0, maxImageCount: 0,
@@ -232,6 +245,18 @@ const AdminImageEngines: React.FC = () => {
title: '最大图片', dataIndex: 'maxImageCount', width: 100, title: '最大图片', dataIndex: 'maxImageCount', width: 100,
render: (v: number) => <Tag color="purple">{v} </Tag>, render: (v: number) => <Tag color="purple">{v} </Tag>,
}, },
{
title: '多份生成', dataIndex: 'multiGenerationEnabled', width: 100,
render: (v: boolean) => <Tag color={v ? 'blue' : 'default'}>{v ? '开启' : '关闭'}</Tag>,
},
{
title: '数量上限', dataIndex: 'maxGenerationCount', width: 100,
render: (v: number, r: ImageEngine) => (
<Tag color={r.multiGenerationEnabled && Number(v || 1) > 1 ? 'magenta' : 'default'}>
{r.multiGenerationEnabled ? (v || 1) : 1}
</Tag>
),
},
{ {
title: '状态', dataIndex: 'isActive', width: 80, title: '状态', dataIndex: 'isActive', width: 80,
render: (v: boolean) => <Tag color={v ? 'green' : 'default'}>{v ? '启用' : '停用'}</Tag>, render: (v: boolean) => <Tag color={v ? 'green' : 'default'}>{v ? '启用' : '停用'}</Tag>,
@@ -350,6 +375,33 @@ const AdminImageEngines: React.FC = () => {
<Form.Item name="generateUrl" label="生成接口地址"> <Form.Item name="generateUrl" label="生成接口地址">
<Input placeholder="https://ark.cn-beijing.volces.com/api/v3/images/generations" size="large" /> <Input placeholder="https://ark.cn-beijing.volces.com/api/v3/images/generations" size="large" />
</Form.Item> </Form.Item>
<div style={{ background: '#f8f9fc', borderRadius: 10, padding: 16, marginBottom: 12 }}>
<Typography.Text strong></Typography.Text>
<Typography.Paragraph style={{ margin: '6px 0 0', color: '#64748b', fontSize: 12 }}>
2-5 API
</Typography.Paragraph>
</div>
<div style={{ display: 'grid', gridTemplateColumns: 'repeat(2, minmax(0, 1fr))', gap: 16 }}>
<Form.Item name="multiGenerationEnabled" label="允许客户端多份生成" valuePropName="checked">
<Switch checkedChildren="开启" unCheckedChildren="关闭" />
</Form.Item>
<Form.Item name="maxGenerationCount" label="客户端最大生成数量" rules={[{ required: true }]}>
<InputNumber min={1} max={5} precision={0} size="large" style={{ width: '100%' }} disabled={!multiGenerationEnabled} />
</Form.Item>
<Form.Item name="multiImageMaxImages" label="组图输入输出总上限" rules={[{ required: true }]}>
<InputNumber min={1} max={15} precision={0} size="large" style={{ width: '100%' }} />
</Form.Item>
<Form.Item name="maxReferenceImageCount" label="最大参考图数量" rules={[{ required: true }]}>
<InputNumber min={0} max={14} precision={0} size="large" style={{ width: '100%' }} />
</Form.Item>
<Form.Item name="outputFormat" label="供应商输出格式">
<Select size="large" options={[
{ value: '', label: '不传(兼容不支持 output_format 的模型)' },
{ value: 'png', label: 'PNG' },
{ value: 'jpeg', label: 'JPEG' },
]} />
</Form.Item>
</div>
<div style={{ display: 'flex', gap: 16 }}> <div style={{ display: 'flex', gap: 16 }}>
<Form.Item name="priority" label="优先级"> <Form.Item name="priority" label="优先级">
<Select size="large" options={[ <Select size="large" options={[
@@ -1,6 +1,6 @@
import React, { useEffect, useState } from 'react'; import React, { useEffect, useState } from 'react';
import { import {
Button, Card, Form, Input, message, Modal, Popconfirm, Select, Space, Switch, Table, Tag, Typography, Button, Card, Form, Input, InputNumber, message, Modal, Popconfirm, Select, Space, Switch, Table, Tag, Typography,
} from 'antd'; } from 'antd';
import { import {
PlayCircleOutlined, PlusOutlined, EditOutlined, DeleteOutlined, PlayCircleOutlined, PlusOutlined, EditOutlined, DeleteOutlined,
@@ -25,6 +25,8 @@ interface VideoEngine {
supportsUniversalReference: boolean; supportsUniversalReference: boolean;
isActive: boolean; isActive: boolean;
priority: number; priority: number;
multiGenerationEnabled: boolean;
maxGenerationCount: number;
} }
function parseJsonArray(val: unknown): any[] { function parseJsonArray(val: unknown): any[] {
@@ -40,6 +42,7 @@ const AdminVideoEngines: React.FC = () => {
const [loading, setLoading] = useState(false); const [loading, setLoading] = useState(false);
const [modal, setModal] = useState<{ open: boolean; engine: VideoEngine | null }>({ open: false, engine: null }); const [modal, setModal] = useState<{ open: boolean; engine: VideoEngine | null }>({ open: false, engine: null });
const [form] = Form.useForm(); const [form] = Form.useForm();
const multiGenerationEnabled = Form.useWatch('multiGenerationEnabled', form) ?? false;
const load = async () => { const load = async () => {
setLoading(true); setLoading(true);
@@ -80,6 +83,8 @@ const AdminVideoEngines: React.FC = () => {
supports_universal_reference: values.supportsUniversalReference ?? true, supports_universal_reference: values.supportsUniversalReference ?? true,
is_active: values.isActive ?? true, is_active: values.isActive ?? true,
priority: values.priority ?? 0, priority: values.priority ?? 0,
multi_generation_enabled: values.multiGenerationEnabled ?? false,
max_generation_count: values.maxGenerationCount ?? 1,
}; };
if (modal.engine) { if (modal.engine) {
await saveVideoEngine({ id: modal.engine.id, ...payload }); await saveVideoEngine({ id: modal.engine.id, ...payload });
@@ -115,6 +120,7 @@ const AdminVideoEngines: React.FC = () => {
form.resetFields(); form.resetFields();
form.setFieldsValue({ form.setFieldsValue({
isActive: true, priority: 0, isActive: true, priority: 0,
multiGenerationEnabled: false, maxGenerationCount: 1,
maxDuration: 30, maxDuration: 30,
maxImageCount: 2, maxImageCount: 2,
maxVideoCount: 0, maxVideoCount: 0,
@@ -180,6 +186,18 @@ const AdminVideoEngines: React.FC = () => {
title: '全能参考', dataIndex: 'supportsUniversalReference', width: 100, title: '全能参考', dataIndex: 'supportsUniversalReference', width: 100,
render: (v: boolean) => <Tag color={v ? 'purple' : 'default'}>{v ? '支持' : '不支持'}</Tag>, render: (v: boolean) => <Tag color={v ? 'purple' : 'default'}>{v ? '支持' : '不支持'}</Tag>,
}, },
{
title: '多份生成', dataIndex: 'multiGenerationEnabled', width: 100,
render: (v: boolean) => <Tag color={v ? 'blue' : 'default'}>{v ? '开启' : '关闭'}</Tag>,
},
{
title: '数量上限', dataIndex: 'maxGenerationCount', width: 100,
render: (v: number, r: VideoEngine) => (
<Tag color={r.multiGenerationEnabled && Number(v || 1) > 1 ? 'magenta' : 'default'}>
{r.multiGenerationEnabled ? (v || 1) : 1}
</Tag>
),
},
{ {
title: '状态', dataIndex: 'isActive', width: 80, title: '状态', dataIndex: 'isActive', width: 80,
render: (v: boolean) => <Tag color={v ? 'green' : 'default'}>{v ? '启用' : '停用'}</Tag>, render: (v: boolean) => <Tag color={v ? 'green' : 'default'}>{v ? '启用' : '停用'}</Tag>,
@@ -315,7 +333,19 @@ const AdminVideoEngines: React.FC = () => {
<Switch /> <Switch />
</Form.Item> </Form.Item>
</div> </div>
<div style={{ background: '#f8f9fc', borderRadius: 10, padding: 16, marginBottom: 12 }}>
<Typography.Text strong></Typography.Text>
<Typography.Paragraph style={{ margin: '6px 0 0', color: '#64748b', fontSize: 12 }}>
1
</Typography.Paragraph>
</div>
<div style={{ display: 'flex', gap: 16 }}> <div style={{ display: 'flex', gap: 16 }}>
<Form.Item name="multiGenerationEnabled" label="允许客户端多份生成" valuePropName="checked" style={{ flex: 1 }}>
<Switch checkedChildren="开启" unCheckedChildren="关闭" />
</Form.Item>
<Form.Item name="maxGenerationCount" label="客户端最大生成数量" style={{ flex: 1 }} rules={[{ required: true }]}>
<InputNumber min={1} max={5} precision={0} size="large" style={{ width: '100%' }} disabled={!multiGenerationEnabled} />
</Form.Item>
<Form.Item name="priority" label="优先级" style={{ flex: 1 }}> <Form.Item name="priority" label="优先级" style={{ flex: 1 }}>
<Select size="large" options={[ <Select size="large" options={[
{ value: 0, label: '0 (默认)' }, { value: 0, label: '0 (默认)' },
+15
View File
@@ -281,6 +281,10 @@ export interface GenerationAiImageEngine {
supportedSizes: Record<string, Record<string, string>>; supportedSizes: Record<string, Record<string, string>>;
defaultSize: string; defaultSize: string;
priority: number; priority: number;
multiGenerationEnabled: boolean;
maxGenerationCount: number;
multiImageMaxImages: number;
maxReferenceImageCount: number;
} }
export interface GenerationAiVideoEngine { export interface GenerationAiVideoEngine {
@@ -299,6 +303,8 @@ export interface GenerationAiVideoEngine {
supportsFirstLastFrame?: boolean; supportsFirstLastFrame?: boolean;
supportsUniversalReference?: boolean; supportsUniversalReference?: boolean;
priority: number; priority: number;
multiGenerationEnabled: boolean;
maxGenerationCount: number;
} }
export interface GenerationAiEnginesResponse { export interface GenerationAiEnginesResponse {
@@ -326,6 +332,10 @@ export interface GenerationAiEngineOption {
supportsFirstLastFrame?: boolean; supportsFirstLastFrame?: boolean;
supportsUniversalReference?: boolean; supportsUniversalReference?: boolean;
priority: number; priority: number;
multiGenerationEnabled?: boolean;
maxGenerationCount?: number;
multiImageMaxImages?: number;
maxReferenceImageCount?: number;
genType: GenerationAiGenType; genType: GenerationAiGenType;
} }
@@ -399,6 +409,10 @@ export interface GenerationAITaskOut {
projectId?: string | null; projectId?: string | null;
genType: GenerationAiGenType | string; genType: GenerationAiGenType | string;
generationMode?: string | null; generationMode?: string | null;
parentTaskId?: string | null;
generationCount: number;
generationIndex?: number | null;
displayStatus?: string | null;
pipelineStage?: string | null; pipelineStage?: string | null;
status: GenerationAITaskStatus; status: GenerationAITaskStatus;
originalPrompt: string; originalPrompt: string;
@@ -428,6 +442,7 @@ export interface GenerationAITaskOut {
errorMessage?: string | null; errorMessage?: string | null;
createdAt?: string | null; createdAt?: string | null;
generatedAt?: string | null; generatedAt?: string | null;
childItems: GenerationAITaskOut[];
} }
export interface GenerationAITaskListOut { export interface GenerationAITaskListOut {
+1 -1
View File
@@ -1 +1 @@
{"root":["./src/app.tsx","./src/env.d.ts","./src/main.tsx","./src/api/client.ts","./src/api/crypto.ts","./src/api/index.ts","./src/components/preresultdisplay.tsx","./src/pages/adminauthoriz.tsx","./src/pages/adminconsume.tsx","./src/pages/admincontactrequests.tsx","./src/pages/admincreditratios.tsx","./src/pages/admincreditrecords.tsx","./src/pages/admindashboard.tsx","./src/pages/admingenerationairecords.tsx","./src/pages/admingenerationrecords.tsx","./src/pages/adminhomematerials.tsx","./src/pages/adminhotopeningreplicationdetail.tsx","./src/pages/adminhotopeningreplications.tsx","./src/pages/adminimageengines.tsx","./src/pages/adminindustries.tsx","./src/pages/adminlayout.tsx","./src/pages/adminloginpage.tsx","./src/pages/adminmateriallist.tsx","./src/pages/adminmenuconfig.tsx","./src/pages/adminmodels.tsx","./src/pages/adminnotificationmanager.tsx","./src/pages/adminoauthlist.tsx","./src/pages/adminoauthapplist.tsx","./src/pages/adminoperationlogs.tsx","./src/pages/adminpaymentconfig.tsx","./src/pages/adminpaymentstats.tsx","./src/pages/adminplatform.tsx","./src/pages/adminpretesttemplates.tsx","./src/pages/adminprivateportraitprojects.tsx","./src/pages/adminrechargepackages.tsx","./src/pages/adminreplicationprojectdetail.tsx","./src/pages/adminsettings.tsx","./src/pages/adminshotreplications.tsx","./src/pages/adminshottasksetdetail.tsx","./src/pages/adminteams.tsx","./src/pages/adminusers.tsx","./src/pages/adminvideoengines.tsx","./src/pages/adminvideopromptschemaconfig.tsx","./src/pages/adminreplication/components/jsoncollapse.tsx","./src/pages/adminreplication/components/mediapreview.tsx","./src/pages/adminreplication/components/statustag.tsx","./src/pages/adminreplication/components/videopromptschemaviewer.tsx","./src/pages/homematerials/homematerialassettable.tsx","./src/pages/homematerials/homematerialcategorypanel.tsx","./src/pages/homematerials/homematerialuploadmodal.tsx","./src/pages/homematerials/mediareferenceseditor.tsx","./src/pages/homematerials/watermarkeditor.tsx","./src/pages/homematerials/watermarklibrarymodal.tsx","./src/pages/homematerials/watermarkpreview.tsx","./src/store/index.ts","./src/types/index.ts","./src/types/xlsx-js-style.d.ts","./src/utils/excelexport.ts","./src/utils/formatdate.ts","./src/utils/resourceurl.ts","./src/utils/videopromptschema.ts"],"version":"6.0.3"} {"root":["./src/app.tsx","./src/env.d.ts","./src/main.tsx","./src/api/client.ts","./src/api/crypto.ts","./src/api/index.ts","./src/components/preresultdisplay.tsx","./src/components/generation/generationtaskresourcegrid.tsx","./src/pages/adminauthoriz.tsx","./src/pages/adminconsume.tsx","./src/pages/admincontactrequests.tsx","./src/pages/admincreditratios.tsx","./src/pages/admincreditrecords.tsx","./src/pages/admindashboard.tsx","./src/pages/admingenerationairecords.tsx","./src/pages/admingenerationrecords.tsx","./src/pages/adminhomematerials.tsx","./src/pages/adminhotopeningreplicationdetail.tsx","./src/pages/adminhotopeningreplications.tsx","./src/pages/adminimageengines.tsx","./src/pages/adminindustries.tsx","./src/pages/adminlayout.tsx","./src/pages/adminloginpage.tsx","./src/pages/adminmateriallist.tsx","./src/pages/adminmenuconfig.tsx","./src/pages/adminmodels.tsx","./src/pages/adminnotificationmanager.tsx","./src/pages/adminoauthlist.tsx","./src/pages/adminoauthapplist.tsx","./src/pages/adminoperationlogs.tsx","./src/pages/adminpaymentconfig.tsx","./src/pages/adminpaymentstats.tsx","./src/pages/adminplatform.tsx","./src/pages/adminpretesttemplates.tsx","./src/pages/adminprivateportraitprojects.tsx","./src/pages/adminrechargepackages.tsx","./src/pages/adminreplicationprojectdetail.tsx","./src/pages/adminsettings.tsx","./src/pages/adminshotreplications.tsx","./src/pages/adminshottasksetdetail.tsx","./src/pages/adminteams.tsx","./src/pages/adminusers.tsx","./src/pages/adminvideoengines.tsx","./src/pages/adminvideopromptschemaconfig.tsx","./src/pages/adminreplication/components/jsoncollapse.tsx","./src/pages/adminreplication/components/mediapreview.tsx","./src/pages/adminreplication/components/statustag.tsx","./src/pages/adminreplication/components/videopromptschemaviewer.tsx","./src/pages/homematerials/homematerialassettable.tsx","./src/pages/homematerials/homematerialcategorypanel.tsx","./src/pages/homematerials/homematerialuploadmodal.tsx","./src/pages/homematerials/mediareferenceseditor.tsx","./src/pages/homematerials/watermarkeditor.tsx","./src/pages/homematerials/watermarklibrarymodal.tsx","./src/pages/homematerials/watermarkpreview.tsx","./src/store/index.ts","./src/types/index.ts","./src/types/xlsx-js-style.d.ts","./src/utils/clipboard.ts","./src/utils/excelexport.ts","./src/utils/formatdate.ts","./src/utils/resourceurl.ts","./src/utils/videopromptschema.ts"],"version":"6.0.3"}
@@ -0,0 +1,230 @@
"""add client-selectable multi generation and image batch claim
Revision ID: abae3e1c70f7
Revises: 2026070902
Create Date: 2026-07-15 10:49:31.803342
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = "abae3e1c70f7"
down_revision: Union[str, None] = "2026070902"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
FK_CHAT_TASK_PARENT = "fk_chat_generation_tasks_parent_task_id"
CK_CHAT_TASK_GENERATION_COUNT = "ck_chat_generation_tasks_generation_count"
CK_CHAT_TASK_GENERATION_INDEX = "ck_chat_generation_tasks_generation_index"
CK_IMAGE_ENGINE_MAX_GENERATION_COUNT = "ck_image_engines_max_generation_count"
CK_IMAGE_ENGINE_MULTI_IMAGE_MAX = "ck_image_engines_multi_image_max_images"
CK_IMAGE_ENGINE_MAX_REFERENCE = "ck_image_engines_max_reference_image_count"
CK_VIDEO_ENGINE_MAX_GENERATION_COUNT = "ck_video_engines_max_generation_count"
def upgrade() -> None:
# ChatGenerationTask:任务级实际生成数量、主子关联和图片批次执行租约。
op.add_column(
"chat_generation_tasks",
sa.Column("parent_task_id", sa.String(length=32), nullable=True),
)
op.add_column(
"chat_generation_tasks",
sa.Column("generation_count", sa.Integer(), server_default=sa.text("1"), nullable=False),
)
op.add_column(
"chat_generation_tasks",
sa.Column("generation_index", sa.Integer(), nullable=True),
)
op.add_column(
"chat_generation_tasks",
sa.Column("provider_create_claim_token", sa.String(length=64), nullable=True),
)
op.add_column(
"chat_generation_tasks",
sa.Column("provider_create_lease_until", sa.DateTime(timezone=True), nullable=True),
)
op.add_column(
"chat_generation_tasks",
sa.Column("provider_create_started_at", sa.DateTime(timezone=True), nullable=True),
)
op.create_check_constraint(
CK_CHAT_TASK_GENERATION_COUNT,
"chat_generation_tasks",
"generation_count BETWEEN 1 AND 5",
)
op.create_check_constraint(
CK_CHAT_TASK_GENERATION_INDEX,
"chat_generation_tasks",
"generation_index IS NULL OR generation_index > 0",
)
op.create_foreign_key(
FK_CHAT_TASK_PARENT,
"chat_generation_tasks",
"chat_generation_tasks",
["parent_task_id"],
["id"],
ondelete="RESTRICT",
)
op.create_index(
"idx_chat_generation_tasks_parent",
"chat_generation_tasks",
["parent_task_id"],
unique=False,
)
op.create_index(
"idx_chat_generation_tasks_user_mode_created",
"chat_generation_tasks",
["user_id", "generation_mode", "created_at"],
unique=False,
)
op.create_index(
"ix_chat_generation_tasks_provider_create_claim_token",
"chat_generation_tasks",
["provider_create_claim_token"],
unique=False,
)
op.create_index(
"ix_chat_generation_tasks_provider_create_lease_until",
"chat_generation_tasks",
["provider_create_lease_until"],
unique=False,
)
op.create_index(
"uq_chat_generation_tasks_parent_index",
"chat_generation_tasks",
["parent_task_id", "generation_index"],
unique=True,
postgresql_where=sa.text(
"parent_task_id IS NOT NULL AND generation_index IS NOT NULL"
),
)
op.create_index(
"uq_chat_generation_tasks_user_chat_idempotency",
"chat_generation_tasks",
["user_id", "idempotency_key"],
unique=True,
postgresql_where=sa.text(
"deleted_at IS NULL "
"AND idempotency_key IS NOT NULL "
"AND generation_mode IN ('chatapi_async', 'chatapi_main')"
),
)
# ImageEngine:管理后台只配置是否允许客户端多份生成和数量上限。
op.add_column(
"image_engines",
sa.Column(
"multi_generation_enabled",
sa.Boolean(),
server_default=sa.text("false"),
nullable=False,
),
)
op.add_column(
"image_engines",
sa.Column("max_generation_count", sa.Integer(), server_default=sa.text("1"), nullable=False),
)
op.add_column(
"image_engines",
sa.Column("multi_image_max_images", sa.Integer(), server_default=sa.text("15"), nullable=False),
)
op.add_column(
"image_engines",
sa.Column("max_reference_image_count", sa.Integer(), server_default=sa.text("14"), nullable=False),
)
op.add_column(
"image_engines",
sa.Column("output_format", sa.String(length=16), server_default=sa.text("''"), nullable=False),
)
op.create_check_constraint(
CK_IMAGE_ENGINE_MAX_GENERATION_COUNT,
"image_engines",
"max_generation_count BETWEEN 1 AND 5",
)
op.create_check_constraint(
CK_IMAGE_ENGINE_MULTI_IMAGE_MAX,
"image_engines",
"multi_image_max_images BETWEEN 1 AND 15",
)
op.create_check_constraint(
CK_IMAGE_ENGINE_MAX_REFERENCE,
"image_engines",
"max_reference_image_count BETWEEN 0 AND 14",
)
# VideoEngine:管理后台只配置是否允许客户端多份生成和数量上限。
op.add_column(
"video_engines",
sa.Column(
"multi_generation_enabled",
sa.Boolean(),
server_default=sa.text("false"),
nullable=False,
),
)
op.add_column(
"video_engines",
sa.Column("max_generation_count", sa.Integer(), server_default=sa.text("1"), nullable=False),
)
op.create_check_constraint(
CK_VIDEO_ENGINE_MAX_GENERATION_COUNT,
"video_engines",
"max_generation_count BETWEEN 1 AND 5",
)
def downgrade() -> None:
op.drop_constraint(CK_VIDEO_ENGINE_MAX_GENERATION_COUNT, "video_engines", type_="check")
op.drop_column("video_engines", "max_generation_count")
op.drop_column("video_engines", "multi_generation_enabled")
op.drop_constraint(CK_IMAGE_ENGINE_MAX_REFERENCE, "image_engines", type_="check")
op.drop_constraint(CK_IMAGE_ENGINE_MULTI_IMAGE_MAX, "image_engines", type_="check")
op.drop_constraint(CK_IMAGE_ENGINE_MAX_GENERATION_COUNT, "image_engines", type_="check")
op.drop_column("image_engines", "output_format")
op.drop_column("image_engines", "max_reference_image_count")
op.drop_column("image_engines", "multi_image_max_images")
op.drop_column("image_engines", "max_generation_count")
op.drop_column("image_engines", "multi_generation_enabled")
op.drop_index(
"uq_chat_generation_tasks_user_chat_idempotency",
table_name="chat_generation_tasks",
postgresql_where=sa.text(
"deleted_at IS NULL "
"AND idempotency_key IS NOT NULL "
"AND generation_mode IN ('chatapi_async', 'chatapi_main')"
),
)
op.drop_index(
"uq_chat_generation_tasks_parent_index",
table_name="chat_generation_tasks",
postgresql_where=sa.text(
"parent_task_id IS NOT NULL AND generation_index IS NOT NULL"
),
)
op.drop_index(
"ix_chat_generation_tasks_provider_create_lease_until",
table_name="chat_generation_tasks",
)
op.drop_index(
"ix_chat_generation_tasks_provider_create_claim_token",
table_name="chat_generation_tasks",
)
op.drop_index("idx_chat_generation_tasks_user_mode_created", table_name="chat_generation_tasks")
op.drop_index("idx_chat_generation_tasks_parent", table_name="chat_generation_tasks")
op.drop_constraint(FK_CHAT_TASK_PARENT, "chat_generation_tasks", type_="foreignkey")
op.drop_constraint(CK_CHAT_TASK_GENERATION_INDEX, "chat_generation_tasks", type_="check")
op.drop_constraint(CK_CHAT_TASK_GENERATION_COUNT, "chat_generation_tasks", type_="check")
op.drop_column("chat_generation_tasks", "provider_create_started_at")
op.drop_column("chat_generation_tasks", "provider_create_lease_until")
op.drop_column("chat_generation_tasks", "provider_create_claim_token")
op.drop_column("chat_generation_tasks", "generation_index")
op.drop_column("chat_generation_tasks", "generation_count")
op.drop_column("chat_generation_tasks", "parent_task_id")
+2 -2
View File
@@ -55,12 +55,12 @@ from app.services.payment import sync_pending_orders, 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 ( from app.services.generation.billing_service import (
OWNER_GENERATION_RECORD, OWNER_GENERATION_RECORD,
charge_generation_media_by_params, charge_generation_media_by_params,
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.utils.id_gen import generate_id from app.utils.id_gen import generate_id
from app.schemas.generation import GenerationType, ASPECT_RATIOS, RESOLUTIONS from app.schemas.generation import GenerationType, ASPECT_RATIOS, RESOLUTIONS
+2 -2
View File
@@ -41,7 +41,7 @@ from app.services.resource_capacity_service import assert_user_resource_capacity
from app.services.upload_resource import delete_unbound_upload_resource, upload_reference_file, cleanup_upload_resource_files_after_commit from app.services.upload_resource import delete_unbound_upload_resource, upload_reference_file, cleanup_upload_resource_files_after_commit
from app.services.upload_resource.log_service import log_upload_resource_exception, safe_rollback_with_log from app.services.upload_resource.log_service import log_upload_resource_exception, safe_rollback_with_log
from app.enums.upload_resource import UploadResourceEventEnum, UploadResourceModuleEnum, UploadResourceTypeEnum from app.enums.upload_resource import UploadResourceEventEnum, UploadResourceModuleEnum, UploadResourceTypeEnum
from app.services.generation_billing_service import ( from app.services.generation.billing_service import (
CHARGE_TEXT_PROMPT, CHARGE_TEXT_PROMPT,
OWNER_GENERATION_RECORD, OWNER_GENERATION_RECORD,
build_credit_biz_key, build_credit_biz_key,
@@ -49,7 +49,7 @@ from app.services.generation_billing_service import (
charge_generation_media_for_record, charge_generation_media_for_record,
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.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
+266 -139
View File
@@ -1,7 +1,8 @@
from datetime import datetime, timezone from datetime import datetime
from fastapi import APIRouter, Body, Depends, HTTPException, Path, Query from fastapi import APIRouter, Body, Depends, HTTPException, Path, Query
from sqlalchemy import and_, select from sqlalchemy import select
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_current_user, get_db from app.dependencies import get_current_user, get_db
@@ -19,25 +20,35 @@ from app.schemas.generation_ai import (
GenerationAITaskListOut, GenerationAITaskListOut,
GenerationAITaskOut, GenerationAITaskOut,
) )
from app.services.generation_ai_service import ( from app.services.generation.ai.service import (
create_async_generation_task, build_task_out_list,
list_generation_ai_engine_options, list_generation_ai_engine_options,
list_async_generation_tasks, list_async_generation_tasks,
list_generation_history_day_items, list_generation_history_day_items,
list_generation_history_grouped_days, list_generation_history_grouped_days,
record_to_out,
soft_delete_chat_generation_task,
) )
from app.services.generation_billing_service import ( from app.enums.generation_task import ChatGenerationPipelineStage, ChatGenerationTaskStatus, GenerationMode
from app.services.generation.ai.task_create_service import (
GenerationTaskCreateResult,
create_generation_task_group,
enqueue_created_generation_tasks,
find_existing_top_level_task,
)
from app.services.generation.ai.task_group_service import (
aggregate_main_task_status,
load_children_map,
soft_delete_child_task,
soft_delete_top_level_task_group,
)
from app.services.generation.billing_service import (
OWNER_CHAT_GENERATION_TASK, OWNER_CHAT_GENERATION_TASK,
charge_generation_media_by_params, charge_generation_media_by_params,
get_next_credit_attempt_no, get_next_credit_attempt_no,
) )
from app.services.generation_history_delete_service import batch_delete_generation_history_items from app.services.generation.history_delete_service import batch_delete_generation_history_items
from app.services.generation_log_service import log_task_event from app.services.generation.log_service import log_task_event
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls, resolve_private_portrait_reference_display_urls
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.operation_log_service import log_operation_event
from app.tasks.celery_app import celery_app from app.tasks.celery_app import celery_app
router = APIRouter( router = APIRouter(
@@ -146,6 +157,7 @@ async def create_task(
..., ...,
description=( description=(
"AI生成任务创建参数。gen_type=image 时使用图片参数;gen_type=video 时使用视频参数。" "AI生成任务创建参数。gen_type=image 时使用图片参数;gen_type=video 时使用视频参数。"
"generation_count 为客户端本次选择的生成数量,默认1,后端会按引擎开关和数量上限校验。"
"枚举:gen_type=image/videomedia_references[].type=image/video/audio" "枚举:gen_type=image/videomedia_references[].type=image/video/audio"
"media_references[].source=upload_resource/private_portrait_asset/空;" "media_references[].source=upload_resource/private_portrait_asset/空;"
"media_references[].role=first_frame/last_frame/reference_image/reference_video/reference_audio。" "media_references[].role=first_frame/last_frame/reference_image/reference_video/reference_audio。"
@@ -157,34 +169,88 @@ async def create_task(
if celery_app is None: if celery_app is None:
raise HTTPException(status_code=503, detail="Celery未启用:请配置 REDIS_URL 或 CELERY_BROKER_URL 后启动 worker") raise HTTPException(status_code=503, detail="Celery未启用:请配置 REDIS_URL 或 CELERY_BROKER_URL 后启动 worker")
task = await create_async_generation_task(db, current_user, req) try:
await db.commit() create_result = await create_generation_task_group(db, current_user, req)
top_level_task_id = str(create_result.top_level_task_id)
enqueue_task_ids = list(create_result.enqueue_task_ids)
await db.commit()
except IntegrityError:
# 并发重复请求可能同时通过预查询;唯一索引负责兜底。
# 回滚本次任务和计费后,按幂等键返回已经成功提交的顶层任务。
await db.rollback()
existing = await find_existing_top_level_task(
db,
user_id=current_user.id,
idempotency_key=req.idempotency_key,
)
if not existing:
raise
create_result = GenerationTaskCreateResult(
top_level_task_id=str(existing.id),
generation_count=int(existing.generation_count or 1),
gen_type=str(existing.gen_type),
created=False,
)
top_level_task_id = str(existing.id)
enqueue_task_ids = []
if create_result.created:
log_operation_event(
domain="generation_ai_batch",
event_type="BATCH_COMMIT_SUCCESS",
event_status="success",
source="api",
user_id=current_user.id,
group_id=top_level_task_id,
task_id=top_level_task_id,
detail={
"gen_type": create_result.gen_type,
"generation_count": create_result.generation_count,
"child_task_ids": create_result.child_task_ids,
"physical_files_deleted": False,
},
)
await log_task_event( await log_task_event(
task, task_id=top_level_task_id,
event_type="TASK_CREATED", event_type=(
to_status="generating", "TASK_CREATED" if create_result.created else "IDEMPOTENCY_HIT"
to_stage="queued", ),
detail={"gen_type": task.gen_type}, to_status="generating" if create_result.created else None,
to_stage="queued" if create_result.created else None,
detail={
"gen_type": create_result.gen_type,
"generation_count": create_result.generation_count,
"child_task_ids": create_result.child_task_ids,
"created": create_result.created,
},
) )
from app.tasks.generation_create_tasks import chatapi_create_generation_task failed_enqueue_ids: list[str] = []
if create_result.created and enqueue_task_ids:
try: failed_enqueue_ids = await enqueue_created_generation_tasks(
chatapi_create_generation_task.delay(task.id)
except Exception as exc:
await mark_chat_generation_task_failed_and_refund_once(
db, db,
task_id=task.id, task_ids=enqueue_task_ids,
error_message=f"任务队列投递失败: {exc}",
pipeline_stage="failed",
) )
await db.commit()
raise HTTPException(status_code=503, detail="任务队列投递失败,请稍后重试")
refs = await resolve_private_portrait_reference_display_urls(db, record_to_out(task).media_references, user_id=current_user.id)
return record_to_out(task, media_references=refs)
result = await db.execute(
select(ChatGenerationTask).where(
ChatGenerationTask.id == top_level_task_id,
ChatGenerationTask.user_id == current_user.id,
ChatGenerationTask.deleted_at.is_(None),
).limit(1)
)
task = result.scalar_one_or_none()
if not task:
raise HTTPException(status_code=404, detail="任务创建后未找到")
output = await build_task_out_list(
db,
[task],
viewer_user_id=current_user.id,
)
if failed_enqueue_ids and len(failed_enqueue_ids) == len(enqueue_task_ids):
raise HTTPException(status_code=503, detail="任务已创建,但任务队列投递失败,请稍后重试")
return output[0]
@router.get( @router.get(
"/tasks", "/tasks",
@@ -261,10 +327,8 @@ async def list_tasks(
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db), db: AsyncSession = Depends(get_db),
): ):
is_admin = False is_admin = current_user.user_type == "admin"
if current_user.user_type == 'admin': if not is_admin:
is_admin = True
else:
user_id = current_user.id user_id = current_user.id
total, items = await list_async_generation_tasks( total, items = await list_async_generation_tasks(
@@ -280,24 +344,17 @@ async def list_tasks(
created_start=created_start, created_start=created_start,
created_end=created_end, created_end=created_end,
) )
# 同一个 API 同时服务管理后台和客户端:
# ====================== 在这里加排序(最新在前)====================== # - 管理员保持数据库倒序,最新记录在列表上方;
if not is_admin: # - 普通用户先查询最新一页,再仅反转当前页,聊天消息从旧到新排列。
# 按 created_at 降序(没有则用 id 降序) items_for_output = items if is_admin else list(reversed(items))
items_sorted = sorted( out_items = await build_task_out_list(
items,
key=lambda x: x.created_at if x.created_at is not None else x.id,
reverse=False # 升序
)
else:
items_sorted = items
refs_map = await batch_resolve_private_portrait_reference_display_urls(
db, db,
{item.id: record_to_out(task=item, is_admin=is_admin).media_references for item in items_sorted}, items_for_output,
user_id=None if is_admin else current_user.id, is_admin=is_admin,
viewer_user_id=None if is_admin else current_user.id,
) )
return GenerationAITaskListOut(total=total, items=[record_to_out(task=i, is_admin=is_admin, media_references=refs_map.get(i.id)) for i in items_sorted]) return GenerationAITaskListOut(total=total, items=out_items)
@router.get( @router.get(
"/history", "/history",
@@ -513,8 +570,9 @@ async def list_history_day_items(
summary="获取AI生成任务详情", summary="获取AI生成任务详情",
description=( description=(
"根据任务ID获取当前登录用户的AI生成任务详情。" "根据任务ID获取当前登录用户的AI生成任务详情。"
"只能查询当前用户自己的任务,且只查询 generation_mode=chatapi_async 的任务" "支持 chatapi_async、chatapi_main 和未删除的 chatapi_child"
"如果任务不存在或不属于当前用户,返回404" "查询 chatapi_main 时返回按 generation_index 升序排列的 child_items"
"已软删除 child 只在父任务 child_items 中保留槽位,不能通过 child ID 单独查询。"
), ),
responses={ responses={
200: { 200: {
@@ -541,17 +599,19 @@ async def get_task(
select(ChatGenerationTask).where( select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id, ChatGenerationTask.id == task_id,
ChatGenerationTask.user_id == current_user.id, ChatGenerationTask.user_id == current_user.id,
ChatGenerationTask.generation_mode == "chatapi_async", ).limit(1)
ChatGenerationTask.deleted_at.is_(None),
)
.limit(1)
) )
task = result.scalar_one_or_none() task = result.scalar_one_or_none()
if not task: if not task:
raise HTTPException(status_code=404, detail="任务不存在") raise HTTPException(status_code=404, detail="任务不存在")
refs = await resolve_private_portrait_reference_display_urls(db, record_to_out(task).media_references, user_id=current_user.id) if task.deleted_at is not None:
return record_to_out(task, media_references=refs) raise HTTPException(status_code=404, detail="任务不存在")
output = await build_task_out_list(
db,
[task],
viewer_user_id=current_user.id,
)
return output[0]
@router.delete( @router.delete(
"/tasks/{task_id}", "/tasks/{task_id}",
@@ -587,38 +647,33 @@ async def delete_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),
): ):
result = await db.execute( mode_result = await db.execute(
select(ChatGenerationTask).where( select(ChatGenerationTask.generation_mode).where(
ChatGenerationTask.id == task_id, ChatGenerationTask.id == task_id,
ChatGenerationTask.user_id == current_user.id, ChatGenerationTask.user_id == current_user.id,
ChatGenerationTask.generation_mode == "chatapi_async", ).limit(1)
ChatGenerationTask.deleted_at.is_(None), )
generation_mode = mode_result.scalar_one_or_none()
if generation_mode == GenerationMode.CHATAPI_CHILD.value:
freed_size_bytes = await soft_delete_child_task(
db,
child_task_id=task_id,
user_id=current_user.id,
) )
.limit(1) else:
) freed_size_bytes = await soft_delete_top_level_task_group(
task = result.scalar_one_or_none() db,
if not task: task_id=task_id,
raise HTTPException(status_code=404, detail="任务不存在") user_id=current_user.id,
)
if task.status == "generating": await db.commit()
raise HTTPException(status_code=400, detail="当前任务正在生成中,暂不能删除")
deleted_at = datetime.now(timezone.utc)
freed_size_bytes = await soft_delete_chat_generation_task(
db,
task=task,
deleted_at=deleted_at,
)
await db.flush()
return GenerationAITaskDeleteOut( return GenerationAITaskDeleteOut(
message="任务已删除", message="任务已删除",
task_id=task.id, task_id=task_id,
deleted=True, deleted=True,
freed_size_bytes=freed_size_bytes, freed_size_bytes=freed_size_bytes,
) )
@router.post( @router.post(
"/tasks/{task_id}/retry", "/tasks/{task_id}/retry",
response_model=GenerationAIRetryOut, response_model=GenerationAIRetryOut,
@@ -664,74 +719,146 @@ async def retry_task(
select(ChatGenerationTask).where( select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id, ChatGenerationTask.id == task_id,
ChatGenerationTask.user_id == current_user.id, ChatGenerationTask.user_id == current_user.id,
ChatGenerationTask.generation_mode == "chatapi_async",
ChatGenerationTask.deleted_at.is_(None), ChatGenerationTask.deleted_at.is_(None),
) ).with_for_update().limit(1)
.with_for_update()
.limit(1)
) )
task = result.scalar_one_or_none() task = result.scalar_one_or_none()
if not task: if not task:
raise HTTPException(status_code=404, detail="任务不存在") raise HTTPException(status_code=404, detail="任务不存在")
if task.status != "failed":
raise HTTPException(status_code=400, detail="只有失败任务可以重试") retry_targets: list[ChatGenerationTask]
retrying_group_children = False
if task.generation_mode == GenerationMode.CHATAPI_MAIN.value:
children_map = await load_children_map(db, [task.id], include_deleted=False)
children = children_map.get(task.id, [])
if task.gen_type == "video":
retry_targets = [
child for child in children
if child.status == ChatGenerationTaskStatus.FAILED.value
]
retrying_group_children = True
if not retry_targets:
raise HTTPException(status_code=400, detail="当前视频任务组没有可重试的失败子任务")
elif children:
# 图片供应商全部成功后才会拆子任务;已有子任务时只允许重试下载,
# 不能再次扣费并覆盖原有生成序号。
retry_targets = [
child for child in children
if child.status == ChatGenerationTaskStatus.FAILED.value
and child.pipeline_stage == ChatGenerationPipelineStage.DOWNLOAD_FAILED.value
and bool(child.remote_result_url)
]
retrying_group_children = True
if not retry_targets:
raise HTTPException(status_code=400, detail="当前图片任务组没有可重试的下载失败子任务")
else:
# 图片批次在供应商阶段整批失败时尚未创建子任务,可整批重新生成并重新计费。
if task.status != ChatGenerationTaskStatus.FAILED.value:
raise HTTPException(status_code=400, detail="只有失败任务可以重试")
retry_targets = [task]
else:
if task.status != ChatGenerationTaskStatus.FAILED.value:
raise HTTPException(status_code=400, detail="只有失败任务可以重试")
retry_targets = [task]
await assert_user_resource_capacity_available(db, current_user.id) await assert_user_resource_capacity_available(db, current_user.id)
enqueue_ids: list[str] = []
download_retry_ids: list[str] = []
for target in retry_targets:
if int(target.retry_count or 0) >= 3:
raise HTTPException(status_code=400, detail=f"任务 {target.id} 已超过最大重试次数")
attempt_no = await get_next_credit_attempt_no( is_download_retry = bool(
db, target.remote_result_url
owner_type=OWNER_CHAT_GENERATION_TASK, and target.pipeline_stage == ChatGenerationPipelineStage.DOWNLOAD_FAILED.value
owner_id=task.id, )
) if not is_download_retry:
media_billing = await charge_generation_media_by_params( attempt_no = await get_next_credit_attempt_no(
db, db,
user_id=task.user_id, owner_type=OWNER_CHAT_GENERATION_TASK,
record_id=task.id, owner_id=target.id,
gen_type=task.gen_type, )
image_size=task.image_size, quantity = int(target.generation_count or 1) if (
duration=task.duration, target.generation_mode == GenerationMode.CHATAPI_MAIN.value and target.gen_type == "image"
resolution=task.resolution, ) else 1
engine_id=task.engine_id, media_billing = await charge_generation_media_by_params(
project_name="AI生成任务", db,
description_prefix="Chat任务重试", user_id=target.user_id,
owner_type=OWNER_CHAT_GENERATION_TASK, record_id=target.id,
attempt_no=attempt_no, gen_type=target.gen_type,
) image_size=target.image_size,
duration=target.duration,
resolution=target.resolution,
engine_id=target.engine_id,
project_name="AI生成任务",
description_prefix="Chat任务重试",
owner_type=OWNER_CHAT_GENERATION_TASK,
attempt_no=attempt_no,
quantity=quantity,
)
target.credits_cost = round(float(target.credits_cost or 0) + media_billing.total_charged, 2)
target.provider_task_id = None
target.seedance_task_id = None
target.remote_result_url = None
target.provider_response_json = None
target.provider_create_claim_token = None
target.provider_create_lease_until = None
target.provider_create_started_at = None
target.image_url = None
target.video_url = None
target.video_cover_url = None
target.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
enqueue_ids.append(str(target.id))
else:
target.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value
download_retry_ids.append(str(target.id))
task.status = "generating" target.status = ChatGenerationTaskStatus.GENERATING.value
task.pipeline_stage = "queued" target.error_message = None
task.error_message = None target.poll_count = 0
task.poll_count = 0 target.last_poll_at = None
task.last_poll_at = None target.generated_at = None
task.provider_task_id = None target.retry_count = int(target.retry_count or 0) + 1
task.seedance_task_id = None
task.remote_result_url = None
task.provider_response_json = None
task.image_url = None
task.video_url = None
task.video_cover_url = None
task.generated_at = None
task.credits_cost = round(float(task.credits_cost or 0) + media_billing.total_charged, 2)
if retrying_group_children:
await db.flush()
await aggregate_main_task_status(db, parent_task_id=str(task.id))
refreshed_task_id = str(task.id)
await db.commit() await db.commit()
from app.tasks.generation_create_tasks import chatapi_create_generation_task failed_enqueue_ids = await enqueue_created_generation_tasks(db, task_ids=enqueue_ids) if enqueue_ids else []
failed_download_enqueue_ids: list[str] = []
if download_retry_ids:
from app.tasks.generation_download_tasks import enqueue_download_task
for target_id in download_retry_ids:
target_result = await db.execute(
select(ChatGenerationTask).where(
ChatGenerationTask.id == target_id,
ChatGenerationTask.deleted_at.is_(None),
).limit(1)
)
target = target_result.scalar_one_or_none()
if not target or not await enqueue_download_task(db, target, recover=True, reason="manual_retry"):
failed_download_enqueue_ids.append(target_id)
try: requested_enqueue_count = len(enqueue_ids) + len(download_retry_ids)
chatapi_create_generation_task.delay(task.id) failed_total_count = len(failed_enqueue_ids) + len(failed_download_enqueue_ids)
except Exception as exc: if requested_enqueue_count and failed_total_count == requested_enqueue_count:
await mark_chat_generation_task_failed_and_refund_once( raise HTTPException(status_code=503, detail="任务状态已重置,但任务队列投递全部失败,将由恢复任务继续处理")
db,
task_id=task.id,
error_message=f"任务队列投递失败: {exc}",
pipeline_stage="failed",
)
await db.commit()
raise HTTPException(status_code=503, detail="任务队列投递失败,请稍后重试")
refreshed = await db.execute(
select(ChatGenerationTask).where(ChatGenerationTask.id == refreshed_task_id).limit(1)
)
refreshed_task = refreshed.scalar_one_or_none()
if not refreshed_task:
raise HTTPException(status_code=404, detail="任务不存在")
return GenerationAIRetryOut( return GenerationAIRetryOut(
id=task.id, id=refreshed_task.id,
status=task.status, status=refreshed_task.status,
pipeline_stage=task.pipeline_stage, pipeline_stage=refreshed_task.pipeline_stage,
message="任务已重新扣费并重新投递", message=(
f"请求重试 {len(retry_targets)} 个任务,成功投递 {max(0, requested_enqueue_count - failed_total_count)} 个,"
f"投递失败 {failed_total_count}"
),
) )
@@ -45,5 +45,9 @@ async def list_active_engines(
"supported_sizes": sizes, "supported_sizes": sizes,
"default_size": e.default_size, "default_size": e.default_size,
"max_image_count": e.max_image_count, "max_image_count": e.max_image_count,
"multi_generation_enabled": bool(getattr(e, "multi_generation_enabled", False)),
"max_generation_count": int(getattr(e, "max_generation_count", 1) or 1),
"multi_image_max_images": int(getattr(e, "multi_image_max_images", 15) or 15),
"max_reference_image_count": int(getattr(e, "max_reference_image_count", 14) or 0),
}) })
return {"items": items} return {"items": items}
@@ -52,6 +52,8 @@ async def list_active_engines(
"max_image_count": e.max_image_count, "max_image_count": e.max_image_count,
"max_video_count": e.max_video_count, "max_video_count": e.max_video_count,
"max_audio_count": e.max_audio_count, "max_audio_count": e.max_audio_count,
"multi_generation_enabled": bool(getattr(e, "multi_generation_enabled", False)),
"max_generation_count": int(getattr(e, "max_generation_count", 1) or 1),
"supports_first_last_frame": e.supports_first_last_frame, "supports_first_last_frame": e.supports_first_last_frame,
"supports_universal_reference": e.supports_universal_reference, "supports_universal_reference": e.supports_universal_reference,
}) })
+1
View File
@@ -18,3 +18,4 @@ from app.enums.celery_queue import *
from app.enums.audio_reference import * from app.enums.audio_reference import *
from app.enums.private_portrait import * from app.enums.private_portrait import *
from app.enums.generation_provider import *
+4
View File
@@ -100,3 +100,7 @@ VIDEO_SCHEMA_MAX_SECTION_COUNT = 40
VIDEO_SCHEMA_MAX_FIELD_COUNT_PER_SECTION = 80 VIDEO_SCHEMA_MAX_FIELD_COUNT_PER_SECTION = 80
VIDEO_SCHEMA_MAX_TIME_RULE_COUNT = 30 VIDEO_SCHEMA_MAX_TIME_RULE_COUNT = 30
VIDEO_SCHEMA_MAX_SEGMENT_COUNT_PER_RULE = 12 VIDEO_SCHEMA_MAX_SEGMENT_COUNT_PER_RULE = 12
MIN_GENERATION_COUNT = 1
MAX_GENERATION_COUNT = 5
+28 -7
View File
@@ -45,17 +45,28 @@ GENERATION_HISTORY_MODULE_SOURCES: tuple[GenerationHistorySourceEnum, ...] = (
"""需要回填 module_generation_projects/module_generation_steps 的模块来源集合。""" """需要回填 module_generation_projects/module_generation_steps 的模块来源集合。"""
GENERATION_HISTORY_SOURCE_TO_TASK_MODE: dict[GenerationHistorySourceEnum, GenerationMode] = { GENERATION_HISTORY_SOURCE_TO_TASK_MODES: dict[GenerationHistorySourceEnum, tuple[GenerationMode, ...]] = {
GenerationHistorySourceEnum.CHAT_TASK: GenerationMode.CHATAPI_ASYNC, GenerationHistorySourceEnum.CHAT_TASK: (
GenerationHistorySourceEnum.HOT_OPENING_REPLICATE: GenerationMode.HOT_OPENING_REPLICATE, GenerationMode.CHATAPI_ASYNC,
GenerationHistorySourceEnum.SHOT_REPLICATE: GenerationMode.SHOT_REPLICATE, GenerationMode.CHATAPI_CHILD,
),
GenerationHistorySourceEnum.HOT_OPENING_REPLICATE: (GenerationMode.HOT_OPENING_REPLICATE,),
GenerationHistorySourceEnum.SHOT_REPLICATE: (GenerationMode.SHOT_REPLICATE,),
} }
"""history_source 到 ChatGenerationTask.generation_mode 的映射。""" """history_source 到 ChatGenerationTask.generation_mode 集合的映射。"""
GENERATION_HISTORY_SOURCE_TO_TASK_MODE: dict[GenerationHistorySourceEnum, GenerationMode] = {
history_source: task_modes[0]
for history_source, task_modes in GENERATION_HISTORY_SOURCE_TO_TASK_MODES.items()
}
"""兼容旧调用的单一模式映射;新查询应使用 GENERATION_HISTORY_SOURCE_TO_TASK_MODES。"""
GENERATION_HISTORY_TASK_MODE_VALUE_TO_SOURCE: dict[str, GenerationHistorySourceEnum] = { GENERATION_HISTORY_TASK_MODE_VALUE_TO_SOURCE: dict[str, GenerationHistorySourceEnum] = {
task_mode.value: history_source task_mode.value: history_source
for history_source, task_mode in GENERATION_HISTORY_SOURCE_TO_TASK_MODE.items() for history_source, task_modes in GENERATION_HISTORY_SOURCE_TO_TASK_MODES.items()
for task_mode in task_modes
} }
"""ChatGenerationTask.generation_mode 字符串值到 history_source 的映射。""" """ChatGenerationTask.generation_mode 字符串值到 history_source 的映射。"""
@@ -103,11 +114,17 @@ def get_generation_history_source_label(source: GenerationHistorySourceEnum | st
def get_generation_history_task_mode(source: GenerationHistorySourceEnum) -> GenerationMode | None: def get_generation_history_task_mode(source: GenerationHistorySourceEnum) -> GenerationMode | None:
"""获取 history_source 对应的 ChatGenerationTask.generation_mode""" """兼容旧调用:返回 history_source 对应的第一个任务模式"""
return GENERATION_HISTORY_SOURCE_TO_TASK_MODE.get(source) return GENERATION_HISTORY_SOURCE_TO_TASK_MODE.get(source)
def get_generation_history_task_modes(source: GenerationHistorySourceEnum) -> tuple[GenerationMode, ...]:
"""获取 history_source 对应的全部 ChatGenerationTask.generation_mode。"""
return GENERATION_HISTORY_SOURCE_TO_TASK_MODES.get(source, ())
def is_generation_history_chat_task_source(source: GenerationHistorySourceEnum) -> bool: def is_generation_history_chat_task_source(source: GenerationHistorySourceEnum) -> bool:
"""判断当前来源是否走 chat_generation_tasks 表。""" """判断当前来源是否走 chat_generation_tasks 表。"""
@@ -121,3 +138,7 @@ def is_generation_history_module_source(source: GenerationHistorySourceEnum) ->
MAX_BATCH_DELETE_COUNT = 30 MAX_BATCH_DELETE_COUNT = 30
HISTORY_DAY_PAGE_SIZE_MAX = 10
HISTORY_GROUP_ITEM_LIMIT = 10
@@ -0,0 +1,40 @@
from enum import StrEnum
class GenerationProviderResultType(StrEnum):
IMAGE = "image"
VIDEO = "video"
class GenerationProviderTaskPhase(StrEnum):
SUBMITTED = "submitted"
POLLING = "polling"
RESULT_READY = "result_ready"
DOWNLOAD_PENDING = "download_pending"
COMPLETED = "completed"
FAILED = "failed"
class ImageProviderErrorType(StrEnum):
TIMEOUT = "timeout"
NETWORK = "network"
RATE_LIMIT = "rate_limit"
AUTH = "auth"
INVALID_REQUEST = "invalid_request"
CAPABILITY_MISMATCH = "capability_mismatch"
CONTENT_REJECTED = "content_rejected"
PROVIDER_INTERNAL = "provider_internal"
INVALID_RESPONSE = "invalid_response"
UNKNOWN = "unknown"
IMAGE_MULTI_OUTPUT_MIN = 1
IMAGE_MULTI_OUTPUT_MAX = 15
IMAGE_MULTI_REFERENCE_MAX = 14
IMAGE_PROVIDER_CLAIM_LEASE_SECONDS = 10 * 60
MULTI_IMAGE_PROMPT_TEMPLATE = (
"请严格生成恰好{count}张内容相关但画面具有明显差异的图片。"
"每张图片必须作为独立图片分别输出,不要把多个画面拼接到同一张图片中,"
"不要生成九宫格、分镜图、组合图或包含多张子图的单张图片。"
)
+52 -1
View File
@@ -3,6 +3,8 @@ from enum import Enum
class GenerationMode(str, Enum): class GenerationMode(str, Enum):
CHATAPI_ASYNC = "chatapi_async" CHATAPI_ASYNC = "chatapi_async"
CHATAPI_MAIN = "chatapi_main"
CHATAPI_CHILD = "chatapi_child"
HOT_OPENING_REPLICATE = "hot_opening_replicate" HOT_OPENING_REPLICATE = "hot_opening_replicate"
SHOT_REPLICATE = "shot_replicate" SHOT_REPLICATE = "shot_replicate"
@@ -19,6 +21,15 @@ class ChatGenerationTaskStatus(str, Enum):
FAILED = "failed" FAILED = "failed"
class ChatGenerationDisplayStatus(str, Enum):
PENDING = "pending"
GENERATING = "generating"
COMPLETED = "completed"
FAILED = "failed"
DOWNLOAD_FAILED = "download_failed"
DELETED = "deleted"
class ChatGenerationPipelineStage(str, Enum): class ChatGenerationPipelineStage(str, Enum):
QUEUED = "queued" QUEUED = "queued"
PREPARING = "preparing" PREPARING = "preparing"
@@ -36,6 +47,31 @@ class ChatGenerationPipelineStage(str, Enum):
class ChatGenerationTaskEventType(str, Enum): class ChatGenerationTaskEventType(str, Enum):
TASK_CREATED = "TASK_CREATED"
IDEMPOTENCY_HIT = "IDEMPOTENCY_HIT"
BATCH_CREATE_START = "BATCH_CREATE_START"
BATCH_MAIN_CREATED = "BATCH_MAIN_CREATED"
BATCH_CHILDREN_CREATED = "BATCH_CHILDREN_CREATED"
BATCH_BILLING_SUCCESS = "BATCH_BILLING_SUCCESS"
BATCH_COMMIT_SUCCESS = "BATCH_COMMIT_SUCCESS"
CHILD_ENQUEUE_START = "CHILD_ENQUEUE_START"
CHILD_ENQUEUE_SUCCESS = "CHILD_ENQUEUE_SUCCESS"
CHILD_ENQUEUE_FAILED = "CHILD_ENQUEUE_FAILED"
IMAGE_MAIN_CLAIM_ACQUIRED = "IMAGE_MAIN_CLAIM_ACQUIRED"
IMAGE_MAIN_CLAIM_REJECTED = "IMAGE_MAIN_CLAIM_REJECTED"
IMAGE_MAIN_CLAIM_EXPIRED = "IMAGE_MAIN_CLAIM_EXPIRED"
IMAGE_BATCH_PROVIDER_START = "IMAGE_BATCH_PROVIDER_START"
IMAGE_BATCH_PROVIDER_SUCCESS = "IMAGE_BATCH_PROVIDER_SUCCESS"
IMAGE_BATCH_PROVIDER_FAILED = "IMAGE_BATCH_PROVIDER_FAILED"
IMAGE_BATCH_SPLIT_START = "IMAGE_BATCH_SPLIT_START"
IMAGE_BATCH_SPLIT_SUCCESS = "IMAGE_BATCH_SPLIT_SUCCESS"
IMAGE_BATCH_SPLIT_FAILED = "IMAGE_BATCH_SPLIT_FAILED"
MAIN_STATUS_AGGREGATED = "MAIN_STATUS_AGGREGATED"
CHILD_RESOURCE_DELETE_START = "CHILD_RESOURCE_DELETE_START"
CHILD_RESOURCE_DELETE_SUCCESS = "CHILD_RESOURCE_DELETE_SUCCESS"
BATCH_GROUP_DELETE_SUCCESS = "BATCH_GROUP_DELETE_SUCCESS"
BATCH_RECOVERY_RECONCILED = "BATCH_RECOVERY_RECONCILED"
PROMPT_CONCAT_START = "PROMPT_CONCAT_START" PROMPT_CONCAT_START = "PROMPT_CONCAT_START"
PROMPT_CONCAT_SUCCESS = "PROMPT_CONCAT_SUCCESS" PROMPT_CONCAT_SUCCESS = "PROMPT_CONCAT_SUCCESS"
@@ -87,8 +123,23 @@ class ChatGenerationTaskEventType(str, Enum):
TASK_FAILED = "TASK_FAILED" TASK_FAILED = "TASK_FAILED"
ALLOWED_GENERATION_MODES = { CHAT_TOP_LEVEL_MODES = {
GenerationMode.CHATAPI_ASYNC.value, GenerationMode.CHATAPI_ASYNC.value,
GenerationMode.CHATAPI_MAIN.value,
}
CHAT_RESOURCE_MODES = {
GenerationMode.CHATAPI_ASYNC.value,
GenerationMode.CHATAPI_CHILD.value,
}
CHAT_EXECUTABLE_MODES = {
GenerationMode.CHATAPI_ASYNC.value,
GenerationMode.CHATAPI_CHILD.value,
}
ALLOWED_GENERATION_MODES = {
*CHAT_EXECUTABLE_MODES,
GenerationMode.HOT_OPENING_REPLICATE.value, GenerationMode.HOT_OPENING_REPLICATE.value,
GenerationMode.SHOT_REPLICATE.value, GenerationMode.SHOT_REPLICATE.value,
} }
+17 -6
View File
@@ -38,17 +38,28 @@ RECENT_GENERATION_CHAT_TASK_MODULES: tuple[RecentGenerationModuleEnum, ...] = (
"""来自 chat_generation_tasks 表的模块集合。""" """来自 chat_generation_tasks 表的模块集合。"""
RECENT_GENERATION_MODULE_TO_TASK_MODE: dict[RecentGenerationModuleEnum, GenerationMode] = { RECENT_GENERATION_MODULE_TO_TASK_MODES: dict[RecentGenerationModuleEnum, tuple[GenerationMode, ...]] = {
RecentGenerationModuleEnum.CHAT_AI: GenerationMode.CHATAPI_ASYNC, RecentGenerationModuleEnum.CHAT_AI: (
RecentGenerationModuleEnum.HOT_OPENING_REPLICATE: GenerationMode.HOT_OPENING_REPLICATE, GenerationMode.CHATAPI_ASYNC,
RecentGenerationModuleEnum.SHOT_REPLICATE: GenerationMode.SHOT_REPLICATE, GenerationMode.CHATAPI_CHILD,
),
RecentGenerationModuleEnum.HOT_OPENING_REPLICATE: (GenerationMode.HOT_OPENING_REPLICATE,),
RecentGenerationModuleEnum.SHOT_REPLICATE: (GenerationMode.SHOT_REPLICATE,),
} }
"""最近生成记录模块枚举到 ChatGenerationTask.generation_mode 的映射。""" """最近生成记录模块枚举到 ChatGenerationTask.generation_mode 集合的映射。"""
RECENT_GENERATION_MODULE_TO_TASK_MODE: dict[RecentGenerationModuleEnum, GenerationMode] = {
module: task_modes[0]
for module, task_modes in RECENT_GENERATION_MODULE_TO_TASK_MODES.items()
}
"""兼容旧调用的单一任务模式映射。"""
RECENT_GENERATION_TASK_MODE_VALUE_TO_MODULE: dict[str, RecentGenerationModuleEnum] = { RECENT_GENERATION_TASK_MODE_VALUE_TO_MODULE: dict[str, RecentGenerationModuleEnum] = {
task_mode.value: module task_mode.value: module
for module, task_mode in RECENT_GENERATION_MODULE_TO_TASK_MODE.items() for module, task_modes in RECENT_GENERATION_MODULE_TO_TASK_MODES.items()
for task_mode in task_modes
} }
"""ChatGenerationTask.generation_mode 字符串值到最近生成记录模块枚举的映射。""" """ChatGenerationTask.generation_mode 字符串值到最近生成记录模块枚举的映射。"""
@@ -1,6 +1,6 @@
from datetime import datetime from datetime import datetime
from sqlalchemy import DateTime, Float, ForeignKey, Index, Integer, String, Text, text from sqlalchemy import CheckConstraint, DateTime, Float, ForeignKey, Index, Integer, String, Text, text
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base, TimestampMixin, SoftDeleteMixin from app.models.base import Base, TimestampMixin, SoftDeleteMixin
@@ -26,6 +26,19 @@ class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin):
unique=True, unique=True,
postgresql_where=text("deleted_at IS NULL AND idempotency_key IS NOT NULL"), postgresql_where=text("deleted_at IS NULL AND idempotency_key IS NOT NULL"),
), ),
# AI 创作顶层任务在 chatapi_async/chatapi_main 之间切换时,
# 同一个前端幂等键也只能创建一组任务。
Index(
"uq_chat_generation_tasks_user_chat_idempotency",
"user_id",
"idempotency_key",
unique=True,
postgresql_where=text(
"deleted_at IS NULL "
"AND idempotency_key IS NOT NULL "
"AND generation_mode IN ('chatapi_async', 'chatapi_main')"
),
),
# 视频 24 小时降频轮询调度使用。 # 视频 24 小时降频轮询调度使用。
Index( Index(
"idx_chat_generation_tasks_next_poll_at", "idx_chat_generation_tasks_next_poll_at",
@@ -37,6 +50,17 @@ class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin):
"AND next_poll_at IS NOT NULL" "AND next_poll_at IS NOT NULL"
), ),
), ),
Index(
"uq_chat_generation_tasks_parent_index",
"parent_task_id",
"generation_index",
unique=True,
postgresql_where=text("parent_task_id IS NOT NULL AND generation_index IS NOT NULL"),
),
Index("idx_chat_generation_tasks_parent", "parent_task_id"),
Index("idx_chat_generation_tasks_user_mode_created", "user_id", "generation_mode", "created_at"),
CheckConstraint("generation_count BETWEEN 1 AND 5", name="ck_chat_generation_tasks_generation_count"),
CheckConstraint("generation_index IS NULL OR generation_index > 0", name="ck_chat_generation_tasks_generation_index"),
) )
@@ -59,6 +83,17 @@ class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin):
status: Mapped[str] = mapped_column(String(32), default="generating", index=True) status: Mapped[str] = mapped_column(String(32), default="generating", index=True)
pipeline_stage: Mapped[str | None] = mapped_column(String(32), nullable=True, index=True) pipeline_stage: Mapped[str | None] = mapped_column(String(32), nullable=True, index=True)
generation_mode: Mapped[str] = mapped_column(String(32), default="chatapi_async", index=True) generation_mode: Mapped[str] = mapped_column(String(32), default="chatapi_async", index=True)
parent_task_id: Mapped[str | None] = mapped_column(
String(32), ForeignKey("chat_generation_tasks.id", ondelete="RESTRICT"), nullable=True
)
generation_count: Mapped[int] = mapped_column(Integer, default=1, server_default="1", nullable=False)
generation_index: Mapped[int | None] = mapped_column(Integer, nullable=True)
# 图片主任务同步调用供应商时的分布式执行租约。
# 防止重复 Celery 消息或恢复任务同时触发多次组图请求。
provider_create_claim_token: Mapped[str | None] = mapped_column(String(64), nullable=True, index=True)
provider_create_lease_until: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True, index=True)
provider_create_started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
media_references: Mapped[str | None] = mapped_column(Text, nullable=True) media_references: Mapped[str | None] = mapped_column(Text, nullable=True)
provider_task_id: Mapped[str | None] = mapped_column(String(128), nullable=True, index=True) provider_task_id: Mapped[str | None] = mapped_column(String(128), nullable=True, index=True)
+26 -1
View File
@@ -1,4 +1,4 @@
from sqlalchemy import Boolean, Integer, String, Text from sqlalchemy import Boolean, CheckConstraint, Integer, String, Text
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base, TimestampMixin from app.models.base import Base, TimestampMixin
@@ -6,6 +6,11 @@ from app.models.base import Base, TimestampMixin
class ImageEngine(Base, TimestampMixin): class ImageEngine(Base, TimestampMixin):
__tablename__ = "image_engines" __tablename__ = "image_engines"
__table_args__ = (
CheckConstraint("max_generation_count BETWEEN 1 AND 5", name="ck_image_engines_max_generation_count"),
CheckConstraint("multi_image_max_images BETWEEN 1 AND 15", name="ck_image_engines_multi_image_max_images"),
CheckConstraint("max_reference_image_count BETWEEN 0 AND 14", name="ck_image_engines_max_reference_image_count"),
)
id: Mapped[str] = mapped_column(String(32), primary_key=True) id: Mapped[str] = mapped_column(String(32), primary_key=True)
name: Mapped[str] = mapped_column(String(64), nullable=False) name: Mapped[str] = mapped_column(String(64), nullable=False)
@@ -18,6 +23,26 @@ class ImageEngine(Base, TimestampMixin):
supported_sizes: Mapped[str] = mapped_column(Text, default='{}') supported_sizes: Mapped[str] = mapped_column(Text, default='{}')
default_size: Mapped[str] = mapped_column(String(32), default="2K") default_size: Mapped[str] = mapped_column(String(32), default="2K")
max_image_count: Mapped[int] = mapped_column(Integer, default=0) max_image_count: Mapped[int] = mapped_column(Integer, default=0)
# 管理后台只配置能力开关与数量上限;本次实际生成数量保存在 ChatGenerationTask.generation_count。
multi_generation_enabled: Mapped[bool] = mapped_column(
Boolean,
default=False,
server_default="false",
nullable=False,
)
max_generation_count: Mapped[int] = mapped_column(
Integer,
default=1,
server_default="1",
nullable=False,
)
# 火山组图接口能力约束。多份图片始终只调用一次 sequential_image_generation=auto 接口。
multi_image_max_images: Mapped[int] = mapped_column(Integer, default=15, server_default="15", nullable=False)
max_reference_image_count: Mapped[int] = mapped_column(Integer, default=14, server_default="14", nullable=False)
# 留空表示不向供应商传 output_format;用于兼容不支持该参数的模型。
output_format: Mapped[str] = mapped_column(String(16), default="", server_default="", nullable=False)
generate_url: Mapped[str | None] = mapped_column(String(512), nullable=True, default="") generate_url: Mapped[str | None] = mapped_column(String(512), nullable=True, default="")
is_active: Mapped[bool] = mapped_column(Boolean, default=True) is_active: Mapped[bool] = mapped_column(Boolean, default=True)
priority: Mapped[int] = mapped_column(Integer, default=0) priority: Mapped[int] = mapped_column(Integer, default=0)
+19 -1
View File
@@ -1,4 +1,4 @@
from sqlalchemy import Boolean, Integer, String from sqlalchemy import Boolean, CheckConstraint, Integer, String
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base, TimestampMixin from app.models.base import Base, TimestampMixin
@@ -6,6 +6,9 @@ from app.models.base import Base, TimestampMixin
class VideoEngine(Base, TimestampMixin): class VideoEngine(Base, TimestampMixin):
__tablename__ = "video_engines" __tablename__ = "video_engines"
__table_args__ = (
CheckConstraint("max_generation_count BETWEEN 1 AND 5", name="ck_video_engines_max_generation_count"),
)
id: Mapped[str] = mapped_column(String(32), primary_key=True) id: Mapped[str] = mapped_column(String(32), primary_key=True)
name: Mapped[str] = mapped_column(String(64), nullable=False) name: Mapped[str] = mapped_column(String(64), nullable=False)
@@ -20,6 +23,21 @@ class VideoEngine(Base, TimestampMixin):
max_image_count: Mapped[int] = mapped_column(Integer, default=2) max_image_count: Mapped[int] = mapped_column(Integer, default=2)
max_video_count: Mapped[int] = mapped_column(Integer, default=0) max_video_count: Mapped[int] = mapped_column(Integer, default=0)
max_audio_count: Mapped[int] = mapped_column(Integer, default=0) max_audio_count: Mapped[int] = mapped_column(Integer, default=0)
# 管理后台只配置能力开关与数量上限;本次实际生成数量保存在 ChatGenerationTask.generation_count。
multi_generation_enabled: Mapped[bool] = mapped_column(
Boolean,
default=False,
server_default="false",
nullable=False,
)
max_generation_count: Mapped[int] = mapped_column(
Integer,
default=1,
server_default="1",
nullable=False,
)
supports_first_last_frame: Mapped[bool] = mapped_column(Boolean, default=False) supports_first_last_frame: Mapped[bool] = mapped_column(Boolean, default=False)
supports_universal_reference: Mapped[bool] = mapped_column(Boolean, default=True) supports_universal_reference: Mapped[bool] = mapped_column(Boolean, default=True)
generate_url: Mapped[str | None] = mapped_column(String(512), nullable=True, default="") generate_url: Mapped[str | None] = mapped_column(String(512), nullable=True, default="")
+25 -2
View File
@@ -108,6 +108,7 @@ class GenerationAITaskCreate(BaseModel):
} }
], ],
"idempotency_key": "frontend-submit-uuid-001", "idempotency_key": "frontend-submit-uuid-001",
"generation_count": 3,
"image_size": "2K", "image_size": "2K",
"image_proportion": "1:1", "image_proportion": "1:1",
"image_px": "2048x2048", "image_px": "2048x2048",
@@ -122,6 +123,7 @@ class GenerationAITaskCreate(BaseModel):
"engine_id": None, "engine_id": None,
"media_references": None, "media_references": None,
"idempotency_key": "frontend-submit-uuid-002", "idempotency_key": "frontend-submit-uuid-002",
"generation_count": 2,
"image_size": None, "image_size": None,
"image_proportion": None, "image_proportion": None,
"image_px": None, "image_px": None,
@@ -169,11 +171,21 @@ class GenerationAITaskCreate(BaseModel):
max_length=64, max_length=64,
description=( description=(
"幂等键,用于防止前端重复提交、网络重试导致重复创建任务和重复扣费。" "幂等键,用于防止前端重复提交、网络重试导致重复创建任务和重复扣费。"
"同一用户、同一 idempotency_key、同一 generation_mode 下重复请求会返回已有任务" "同一用户、同一 idempotency_key 的 AI 创作顶层请求会返回已有任务,即使客户端再次传入不同生成数量也不会重复创建"
"建议前端每次点击生成时生成 UUID;同一次请求失败重试时复用同一个 UUID。" "建议前端每次点击生成时生成 UUID;同一次请求失败重试时复用同一个 UUID。"
), ),
examples=["frontend-submit-uuid-001"], examples=["frontend-submit-uuid-001"],
) )
generation_count: int = Field(
1,
ge=1,
le=5,
description=(
"客户端本次实际选择的生成数量,默认 1。后端会按当前引擎的多份生成开关、"
"最大生成数量以及图片参考图总量限制再次校验。"
),
examples=[3],
)
# image params # image params
image_size: str | None = Field( image_size: str | None = Field(
@@ -227,6 +239,10 @@ class GenerationAIImageEngineOptionOut(BaseModel):
default_size: str | None = Field(None, description="默认图片分辨率档位,例如 2K") default_size: str | None = Field(None, description="默认图片分辨率档位,例如 2K")
priority: int = Field(0, description="引擎优先级,数值越大越优先") priority: int = Field(0, description="引擎优先级,数值越大越优先")
max_image_count: int = Field(0, description="最大图片数量") max_image_count: int = Field(0, description="最大图片数量")
multi_generation_enabled: bool = Field(False, description="是否允许客户端选择生成多份图片")
max_generation_count: int = Field(1, ge=1, le=5, description="客户端本次最多可选择的图片生成数量")
multi_image_max_images: int = Field(15, ge=1, le=15, description="单次组图输入与输出总图片上限")
max_reference_image_count: int = Field(14, ge=0, le=14, description="允许的最大参考图片数量")
class GenerationAIVideoEngineOptionOut(BaseModel): class GenerationAIVideoEngineOptionOut(BaseModel):
@@ -244,6 +260,8 @@ class GenerationAIVideoEngineOptionOut(BaseModel):
max_image_count: int | None = Field(None, description="最大图片数量") max_image_count: int | None = Field(None, description="最大图片数量")
max_video_count: int | None = Field(None, description="最大视频数量") max_video_count: int | None = Field(None, description="最大视频数量")
max_audio_count: int | None = Field(None, description="最大参考音频数量,0 表示不支持音频参考") max_audio_count: int | None = Field(None, description="最大参考音频数量,0 表示不支持音频参考")
multi_generation_enabled: bool = Field(False, description="是否允许客户端选择生成多个视频")
max_generation_count: int = Field(1, ge=1, le=5, description="客户端本次最多可选择的视频生成数量")
supports_first_last_frame: bool = Field(False, description="是否支持首帧和最后一帧") supports_first_last_frame: bool = Field(False, description="是否支持首帧和最后一帧")
supports_universal_reference: bool = Field(False, description="是否支持通用参考") supports_universal_reference: bool = Field(False, description="是否支持通用参考")
@@ -416,8 +434,12 @@ class GenerationAITaskOut(BaseModel):
gen_type: str = Field(..., description="生成类型:image=图片,video=视频") gen_type: str = Field(..., description="生成类型:image=图片,video=视频")
generation_mode: str | None = Field( generation_mode: str | None = Field(
None, None,
description="生成模式。当前异步Chat生成任务一般为 chatapi_async", description="生成模式chatapi_async=单份任务,chatapi_main=多份主任务,chatapi_child=多份子任务",
) )
parent_task_id: str | None = Field(None, description="多份生成子任务关联的主任务ID")
generation_count: int = Field(1, ge=1, le=5, description="本次实际生成数量快照")
generation_index: int | None = Field(None, ge=1, le=5, description="子任务生成序号,从1开始")
display_status: str | None = Field(None, description="前端展示状态,例如 download_failed、deleted")
pipeline_stage: str | None = Field( pipeline_stage: str | None = Field(
None, None,
description=( description=(
@@ -468,6 +490,7 @@ class GenerationAITaskOut(BaseModel):
error_message: str | None = Field(None, description="错误信息。成功任务一般为 null") error_message: str | None = Field(None, description="错误信息。成功任务一般为 null")
created_at: NaiveDatetimeOptional = Field(None, description="任务创建时间") created_at: NaiveDatetimeOptional = Field(None, description="任务创建时间")
generated_at: NaiveDatetimeOptional = Field(None, description="任务生成完成时间") generated_at: NaiveDatetimeOptional = Field(None, description="任务生成完成时间")
child_items: list["GenerationAITaskOut"] = Field(default_factory=list, description="多份生成子任务列表,按 generation_index 升序")
class GenerationAITaskListOut(BaseModel): class GenerationAITaskListOut(BaseModel):
+44 -1
View File
@@ -1,5 +1,6 @@
from pydantic import BaseModel, Field from pydantic import BaseModel, Field, model_validator
from app.enums.generation_provider import IMAGE_MULTI_OUTPUT_MAX, IMAGE_MULTI_REFERENCE_MAX
from app.schemas.common import NaiveDatetime from app.schemas.common import NaiveDatetime
@@ -13,10 +14,48 @@ class ImageEngineCreate(BaseModel):
supported_sizes: str = Field(default='{}') supported_sizes: str = Field(default='{}')
default_size: str = Field(default="2K", max_length=32) default_size: str = Field(default="2K", max_length=32)
max_image_count: int = Field(default=0) max_image_count: int = Field(default=0)
multi_generation_enabled: bool = Field(
default=False,
description="是否允许客户端选择生成多份图片;关闭时客户端只能选择 1 份",
)
max_generation_count: int = Field(
default=1,
ge=1,
le=5,
description="客户端单次最多可选择的生成数量,范围 1-5",
)
multi_image_max_images: int = Field(
default=IMAGE_MULTI_OUTPUT_MAX,
ge=1,
le=IMAGE_MULTI_OUTPUT_MAX,
description="火山组图接口输入参考图与输出图片总上限",
)
max_reference_image_count: int = Field(
default=IMAGE_MULTI_REFERENCE_MAX,
ge=0,
le=IMAGE_MULTI_REFERENCE_MAX,
description="图片引擎允许的最大参考图片数量",
)
output_format: str = Field(
default="",
max_length=16,
description="供应商输出格式;留空表示不传该参数,用于兼容不支持 output_format 的模型",
)
generate_url: str = Field(default="", max_length=512) generate_url: str = Field(default="", max_length=512)
is_active: bool = True is_active: bool = True
priority: int = 0 priority: int = 0
@model_validator(mode="after")
def validate_multi_generation_capability(self):
if self.max_generation_count > self.multi_image_max_images:
raise ValueError("max_generation_count 不能大于 multi_image_max_images")
normalized_output_format = (self.output_format or "").lower().strip()
if normalized_output_format not in {"", "png", "jpeg"}:
raise ValueError("output_format 仅支持留空、png 或 jpeg")
self.output_format = normalized_output_format
return self
class ImageEngineOut(ImageEngineCreate): class ImageEngineOut(ImageEngineCreate):
id: str id: str
@@ -33,6 +72,10 @@ class ImageEnginePublic(BaseModel):
supported_sizes: dict[str, dict[str, str]] = {} supported_sizes: dict[str, dict[str, str]] = {}
default_size: str = "2K" default_size: str = "2K"
max_image_count: int = 0 max_image_count: int = 0
multi_generation_enabled: bool = False
max_generation_count: int = 1
multi_image_max_images: int = IMAGE_MULTI_OUTPUT_MAX
max_reference_image_count: int = IMAGE_MULTI_REFERENCE_MAX
class ImageEngineListResponse(BaseModel): class ImageEngineListResponse(BaseModel):
+12
View File
@@ -16,6 +16,16 @@ class VideoEngineCreate(BaseModel):
max_image_count: int = Field(default=2) max_image_count: int = Field(default=2)
max_video_count: int = Field(default=0) max_video_count: int = Field(default=0)
max_audio_count: int = Field(default=0, ge=0, le=3, description="最大参考音频数量,0 表示不支持音频参考") max_audio_count: int = Field(default=0, ge=0, le=3, description="最大参考音频数量,0 表示不支持音频参考")
multi_generation_enabled: bool = Field(
default=False,
description="是否允许客户端选择生成多个视频;关闭时客户端只能选择 1 份",
)
max_generation_count: int = Field(
default=1,
ge=1,
le=5,
description="客户端单次最多可选择的生成数量,范围 1-5",
)
supports_first_last_frame: bool = Field(default=False, description="是否支持首尾帧模式") supports_first_last_frame: bool = Field(default=False, description="是否支持首尾帧模式")
supports_universal_reference: bool = Field(default=True, description="是否支持全能参考模式") supports_universal_reference: bool = Field(default=True, description="是否支持全能参考模式")
generate_url: str = Field(default="", max_length=512) generate_url: str = Field(default="", max_length=512)
@@ -41,6 +51,8 @@ class VideoEnginePublic(BaseModel):
max_image_count: int = 2 max_image_count: int = 2
max_video_count: int = 0 max_video_count: int = 0
max_audio_count: int = 0 max_audio_count: int = 0
multi_generation_enabled: bool = False
max_generation_count: int = 1
supports_first_last_frame: bool = False supports_first_last_frame: bool = False
supports_universal_reference: bool = True supports_universal_reference: bool = True
@@ -0,0 +1 @@
"""生成任务领域服务。"""
@@ -0,0 +1 @@
"""AI 创作生成编排服务。"""
@@ -0,0 +1,122 @@
from __future__ import annotations
import json
from fastapi import HTTPException
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.common import MAX_GENERATION_COUNT, MIN_GENERATION_COUNT
from app.enums.generation_provider import IMAGE_MULTI_OUTPUT_MAX, IMAGE_MULTI_REFERENCE_MAX
from app.models.image_engine import ImageEngine
from app.models.video_engine import VideoEngine
IMAGE_DEFAULT_SIZE = "2K"
IMAGE_DEFAULT_PROPORTION = "1:1"
IMAGE_DEFAULT_PX = "2048x2048"
VIDEO_DEFAULT_DURATION = 4
VIDEO_DEFAULT_RATIO = "16:9"
VIDEO_DEFAULT_RESOLUTION = "480p"
def normalize_px(value: str | None) -> str | None:
if not value:
return value
return value.replace("×", "x").replace("X", "x").replace("×x", "x").replace("x×", "x")
def parse_json_list(value: str | None, fallback: list):
try:
parsed = json.loads(value or "")
return parsed if isinstance(parsed, list) else fallback
except Exception:
return fallback
def image_supported_sizes(engine: ImageEngine) -> dict:
try:
data = json.loads(engine.supported_sizes or "{}")
return data if isinstance(data, dict) else {}
except Exception:
return {}
def normalize_generation_count(value: int | None) -> int:
try:
count = int(value or MIN_GENERATION_COUNT)
except (TypeError, ValueError):
count = MIN_GENERATION_COUNT
return min(MAX_GENERATION_COUNT, max(MIN_GENERATION_COUNT, count))
async def get_image_engine(db: AsyncSession, engine_id: str | None) -> ImageEngine:
query = select(ImageEngine).where(ImageEngine.is_active == True)
if engine_id:
query = query.where(ImageEngine.id == engine_id)
else:
query = query.order_by(ImageEngine.priority.desc()).limit(1)
result = await db.execute(query)
engine = result.scalar_one_or_none()
if not engine:
raise HTTPException(status_code=400, detail="没有可用的图片引擎")
return engine
async def get_video_engine(db: AsyncSession, engine_id: str | None) -> VideoEngine:
query = select(VideoEngine).where(VideoEngine.is_active == True)
if engine_id:
query = query.where(VideoEngine.id == engine_id)
else:
query = query.order_by(VideoEngine.priority.desc())
result = await db.execute(query.limit(1))
engine = result.scalar_one_or_none()
if not engine:
raise HTTPException(status_code=400, detail="没有可用的视频引擎")
return engine
def build_image_snapshot(engine: ImageEngine, size: str, proportion: str, px: str) -> dict:
return {
"engine_type": "image",
"id": engine.id,
"name": engine.name,
"provider": engine.provider,
"api_base": engine.api_base,
"api_key_masked": "****" if engine.api_key else "",
"model_name": engine.model_name,
"generate_url": engine.generate_url,
"supported_models": parse_json_list(engine.supported_models, []),
"default_size": engine.default_size,
"multi_generation_enabled": bool(getattr(engine, "multi_generation_enabled", False)),
"max_generation_count": normalize_generation_count(getattr(engine, "max_generation_count", 1)),
"multi_image_max_images": int(getattr(engine, "multi_image_max_images", IMAGE_MULTI_OUTPUT_MAX) or IMAGE_MULTI_OUTPUT_MAX),
"max_reference_image_count": int(getattr(engine, "max_reference_image_count", IMAGE_MULTI_REFERENCE_MAX) or 0),
"output_format": (getattr(engine, "output_format", "") or "").lower().strip(),
"selected_size": size,
"selected_proportion": proportion,
"selected_px": px,
}
def build_video_snapshot(engine: VideoEngine, ratio: str, resolution: str, duration: int) -> dict:
return {
"engine_type": "video",
"id": engine.id,
"name": engine.name,
"provider": engine.provider,
"api_base": engine.api_base,
"api_key_masked": "****" if engine.api_key else "",
"model_name": engine.model_name,
"generate_url": engine.generate_url,
"query_url": engine.query_url,
"supported_ratios": parse_json_list(engine.supported_ratios, []),
"supported_resolutions": parse_json_list(engine.supported_resolutions, []),
"supported_durations": parse_json_list(engine.supported_durations, []),
"max_duration": engine.max_duration,
"max_audio_count": engine.max_audio_count,
"multi_generation_enabled": bool(getattr(engine, "multi_generation_enabled", False)),
"max_generation_count": normalize_generation_count(getattr(engine, "max_generation_count", 1)),
"selected_ratio": ratio,
"selected_resolution": resolution,
"selected_duration": duration,
}
@@ -0,0 +1,532 @@
from __future__ import annotations
import json
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from uuid import uuid4
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.generation_provider import IMAGE_PROVIDER_CLAIM_LEASE_SECONDS
from app.enums.generation_task import (
ChatGenerationPipelineStage,
ChatGenerationTaskEventType,
ChatGenerationTaskStatus,
GenerationMode,
GenerationType,
)
from app.models.chat_generation_task import ChatGenerationTask
from app.services.generation.ai.task_group_service import aggregate_main_task_status, load_children_map
from app.services.generation.log_service import log_task_event
from app.services.generation.provider_service import (
create_image_sync_batch_result_with_engine,
get_runtime_engine,
)
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.image_gen import ImageProviderError
from app.services.operation_log_service import build_exception_detail, log_operation_event
from app.utils.id_gen import generate_id
@dataclass(slots=True)
class ImageBatchClaim:
acquired: bool
main_task_id: str
claim_token: str | None = None
task_snapshot: SimpleNamespace | None = None
runtime_engine: SimpleNamespace | None = None
existing_child_ids: list[str] | None = None
reason: str | None = None
def _now() -> datetime:
return datetime.now(timezone.utc)
def _json(value) -> str | None:
if value is None:
return None
return json.dumps(value, ensure_ascii=False, default=str)
def _aware(value: datetime | None) -> datetime | None:
if value is None:
return None
if value.tzinfo is None:
return value.replace(tzinfo=timezone.utc)
return value.astimezone(timezone.utc)
def _lease_alive(task: ChatGenerationTask, now: datetime | None = None) -> bool:
lease_until = _aware(task.provider_create_lease_until)
return bool(task.provider_create_claim_token and lease_until and lease_until > (now or _now()))
def _task_snapshot(main: ChatGenerationTask) -> SimpleNamespace:
return SimpleNamespace(
id=str(main.id),
user_id=str(main.user_id),
generation_mode=str(main.generation_mode),
generation_count=int(main.generation_count or 1),
original_prompt=main.original_prompt,
optimized_prompt=main.optimized_prompt,
media_references=main.media_references,
gen_type=main.gen_type,
duration=main.duration,
aspect_ratio=main.aspect_ratio,
resolution=main.resolution,
image_size=main.image_size,
image_proportion=main.image_proportion,
image_px=main.image_px,
engine_id=main.engine_id,
)
async def _claim_image_main_batch(db: AsyncSession, main_task_id: str) -> ImageBatchClaim:
result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == main_task_id,
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
ChatGenerationTask.gen_type == GenerationType.IMAGE.value,
ChatGenerationTask.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
main = result.scalar_one_or_none()
if not main:
await db.rollback()
return ImageBatchClaim(False, main_task_id, reason="main_missing")
children_map = await load_children_map(db, [main.id], include_deleted=True)
existing_children = children_map.get(main.id, [])
if existing_children:
child_ids = [str(child.id) for child in existing_children if child.deleted_at is None]
main.provider_create_claim_token = None
main.provider_create_lease_until = None
await db.commit()
return ImageBatchClaim(False, main_task_id, existing_child_ids=child_ids, reason="already_split")
if main.status != ChatGenerationTaskStatus.GENERATING.value:
status = str(main.status)
await db.rollback()
return ImageBatchClaim(False, main_task_id, reason=f"status_{status}")
now = _now()
if _lease_alive(main, now):
user_id = str(main.user_id)
group_id = str(main.id)
lease_until = main.provider_create_lease_until
await db.rollback()
log_operation_event(
domain="generation_ai_batch",
event_type="IMAGE_MAIN_CLAIM_REJECTED",
event_status="skipped",
source="celery",
user_id=user_id,
group_id=group_id,
task_id=group_id,
detail={"reason": "lease_alive", "lease_until": lease_until},
)
return ImageBatchClaim(False, main_task_id, reason="lease_alive")
deadline = _aware(main.deadline_at)
if deadline and deadline <= now:
main.provider_create_claim_token = None
main.provider_create_lease_until = None
await mark_chat_generation_task_failed_and_refund_once(
db,
task=main,
error_message="图片批量生成任务超时",
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
)
await db.commit()
return ImageBatchClaim(False, main_task_id, reason="deadline_expired")
claim_token = uuid4().hex
main.provider_create_claim_token = claim_token
main.provider_create_started_at = now
main.provider_create_lease_until = now + timedelta(seconds=IMAGE_PROVIDER_CLAIM_LEASE_SECONDS)
main.pipeline_stage = ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value
runtime_engine = await get_runtime_engine(db, main)
snapshot = _task_snapshot(main)
user_id = str(main.user_id)
generation_count = int(main.generation_count or 1)
lease_until = main.provider_create_lease_until
await db.commit()
log_operation_event(
domain="generation_ai_batch",
event_type="IMAGE_MAIN_CLAIM_ACQUIRED",
event_status="success",
source="celery",
user_id=user_id,
group_id=main_task_id,
task_id=main_task_id,
detail={
"generation_count": generation_count,
"claim_token_suffix": claim_token[-8:],
"lease_until": lease_until,
},
)
return ImageBatchClaim(
True,
main_task_id,
claim_token=claim_token,
task_snapshot=snapshot,
runtime_engine=runtime_engine,
)
def _validate_provider_batch(provider_result: dict, generation_count: int) -> list[dict]:
items = provider_result.get("items") or []
if not isinstance(items, list):
raise RuntimeError("图片供应商返回 items 结构异常")
success_items: list[dict] = []
errors: list[str] = []
for position, item in enumerate(items, start=1):
if not isinstance(item, dict):
errors.append(f"{position}项返回结构无效")
continue
if item.get("error_message") or item.get("error_code"):
errors.append(
f"{position}项: {item.get('error_message') or item.get('error_code') or '生成失败'}"
)
continue
remote_url = str(item.get("remote_result_url") or "").strip()
if not remote_url:
errors.append(f"{position}项: 供应商未返回图片地址")
continue
normalized = dict(item)
normalized["generation_index"] = position
success_items.append(normalized)
generated_images = int(provider_result.get("generated_images") or 0)
if generated_images and generated_images != len(success_items):
errors.append(
f"usage.generated_images={generated_images} 与有效图片数 {len(success_items)} 不一致"
)
if len(items) != generation_count:
errors.append(f"返回条目数应为 {generation_count},实际 {len(items)}")
if len(success_items) != generation_count:
errors.append(f"成功图片数应为 {generation_count},实际 {len(success_items)}")
if errors:
raise RuntimeError("图片组图未全部成功;" + "".join(errors))
return success_items
async def _fail_claimed_main(
db: AsyncSession,
*,
main_task_id: str,
claim_token: str,
error_message: str,
event_type: ChatGenerationTaskEventType,
exception: Exception | None = None,
) -> bool:
try:
await db.rollback()
except Exception:
pass
result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == main_task_id,
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
ChatGenerationTask.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
main = result.scalar_one_or_none()
if not main or main.provider_create_claim_token != claim_token:
await db.rollback()
return False
existing_map = await load_children_map(db, [main.id], include_deleted=True)
if existing_map.get(main.id):
# child 已经落库后不再允许图片生成退款。
main.provider_create_claim_token = None
main.provider_create_lease_until = None
await db.commit()
return False
main.provider_create_claim_token = None
main.provider_create_lease_until = None
await mark_chat_generation_task_failed_and_refund_once(
db,
task=main,
error_message=error_message,
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
)
task_id = str(main.id)
user_id = str(main.user_id)
await db.commit()
await log_task_event(
task_id=task_id,
event_type=event_type.value,
to_status=ChatGenerationTaskStatus.FAILED.value,
to_stage=ChatGenerationPipelineStage.FAILED.value,
message=error_message,
)
log_operation_event(
domain="generation_ai_batch",
event_type=event_type.value,
event_status="failed",
source="celery",
user_id=user_id,
group_id=task_id,
task_id=task_id,
message=error_message,
detail=build_exception_detail(exception) if exception else {"message": error_message},
error=error_message,
)
return True
async def _split_children(
db: AsyncSession,
*,
main_task_id: str,
claim_token: str,
provider_result: dict,
provider_items: list[dict],
) -> list[str]:
result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == main_task_id,
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
ChatGenerationTask.gen_type == GenerationType.IMAGE.value,
ChatGenerationTask.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
main = result.scalar_one_or_none()
if not main:
raise RuntimeError("图片主任务不存在或已删除")
if main.provider_create_claim_token != claim_token:
raise RuntimeError("图片主任务执行租约已失效,拒绝拆分子任务")
if main.status != ChatGenerationTaskStatus.GENERATING.value:
raise RuntimeError(f"图片主任务当前状态不允许拆分: {main.status}")
existing_map = await load_children_map(db, [main.id], include_deleted=True)
existing = existing_map.get(main.id, [])
if existing:
main.provider_create_claim_token = None
main.provider_create_lease_until = None
await db.commit()
return [str(child.id) for child in existing if child.deleted_at is None]
expected_count = max(1, int(main.generation_count or 1))
if len(provider_items) != expected_count:
raise RuntimeError(f"图片批量拆分数量不一致,期望 {expected_count},实际 {len(provider_items)}")
log_operation_event(
domain="generation_ai_batch",
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_SPLIT_START.value,
event_status="started",
source="celery",
user_id=main.user_id,
group_id=main.id,
task_id=main.id,
detail={"generation_count": expected_count},
)
children: list[ChatGenerationTask] = []
for item in provider_items:
index = int(item.get("generation_index") or 0)
if index < 1 or index > expected_count:
raise RuntimeError(f"无效的图片生成序号: {index}")
child = ChatGenerationTask(
id=generate_id(),
user_id=main.user_id,
original_prompt=main.original_prompt,
optimized_prompt=main.optimized_prompt,
gen_type=main.gen_type,
image_size=main.image_size,
image_proportion=main.image_proportion,
image_px=main.image_px,
status=ChatGenerationTaskStatus.GENERATING.value,
pipeline_stage=ChatGenerationPipelineStage.RESULT_READY.value,
generation_mode=GenerationMode.CHATAPI_CHILD.value,
parent_task_id=main.id,
generation_count=expected_count,
generation_index=index,
media_references=main.media_references,
remote_result_url=item.get("remote_result_url"),
engine_id=main.engine_id,
engine_snapshot_json=main.engine_snapshot_json,
provider_response_json=_json(item.get("response_data") or {}),
# 图片生成计费和 token 都归属于 main;child 只负责下载和资源展示。
credits_cost=0,
image_tokens_used=0,
deadline_at=main.deadline_at,
)
children.append(child)
children.sort(key=lambda child: int(child.generation_index or 0))
db.add_all(children)
main.provider_response_json = _json(provider_result.get("response_data") or provider_result)
main.image_tokens_used = int(provider_result.get("image_tokens") or 0)
main.provider_create_claim_token = None
main.provider_create_lease_until = None
await db.flush()
child_ids = [str(child.id) for child in children]
main_id = str(main.id)
main_user_id = str(main.user_id)
await aggregate_main_task_status(db, parent_task_id=main_id)
await db.commit()
log_operation_event(
domain="generation_ai_batch",
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_SPLIT_SUCCESS.value,
event_status="success",
source="celery",
user_id=main_user_id,
group_id=main_id,
task_id=main_id,
detail={"child_task_ids": child_ids},
)
return child_ids
async def _enqueue_child_downloads(db: AsyncSession, child_ids: list[str]) -> dict[str, list[str]]:
if not child_ids:
return {"enqueued": [], "failed": []}
result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.id.in_(child_ids),
ChatGenerationTask.deleted_at.is_(None),
)
.order_by(ChatGenerationTask.generation_index.asc())
)
children = list(result.scalars().all())
from app.tasks.generation_download_tasks import enqueue_download_task
enqueued: list[str] = []
failed: list[str] = []
for child in children:
if child.status == ChatGenerationTaskStatus.COMPLETED.value:
continue
if child.pipeline_stage in {
ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value,
ChatGenerationPipelineStage.DOWNLOADING.value,
ChatGenerationPipelineStage.RETRY_WAITING.value,
}:
continue
celery_task_id = await enqueue_download_task(db, child, reason="image_batch_split")
if celery_task_id:
enqueued.append(str(child.id))
else:
failed.append(str(child.id))
if children and children[0].parent_task_id:
await aggregate_main_task_status(db, parent_task_id=str(children[0].parent_task_id))
await db.commit()
return {"enqueued": enqueued, "failed": failed}
async def run_image_main_batch(db: AsyncSession, main_task: ChatGenerationTask) -> list[str]:
"""单次同步组图,全部成功后原子拆分 child。
绝不在组图 API 失败后退化为 N 次单图请求。
"""
main_task_id = str(main_task.id)
claim = await _claim_image_main_batch(db, main_task_id)
if claim.existing_child_ids is not None:
await _enqueue_child_downloads(db, claim.existing_child_ids)
return claim.existing_child_ids
if not claim.acquired or not claim.claim_token or not claim.task_snapshot or not claim.runtime_engine:
return []
generation_count = max(1, int(claim.task_snapshot.generation_count or 1))
try:
log_operation_event(
domain="generation_ai_batch",
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_PROVIDER_START.value,
event_status="started",
source="celery",
user_id=claim.task_snapshot.user_id,
group_id=main_task_id,
task_id=main_task_id,
detail={"generation_count": generation_count},
)
provider_result = await create_image_sync_batch_result_with_engine(
claim.task_snapshot,
claim.runtime_engine,
generation_count=generation_count,
)
provider_items = _validate_provider_batch(provider_result, generation_count)
log_operation_event(
domain="generation_ai_batch",
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_PROVIDER_SUCCESS.value,
event_status="success",
source="celery",
user_id=claim.task_snapshot.user_id,
group_id=main_task_id,
task_id=main_task_id,
detail={
"generation_count": generation_count,
"result_count": len(provider_items),
"image_tokens": int(provider_result.get("image_tokens") or 0),
"single_provider_request": True,
"fallback_to_single_requests": False,
},
)
except Exception as exc:
message = exc.safe_message if isinstance(exc, ImageProviderError) else str(exc)
await _fail_claimed_main(
db,
main_task_id=main_task_id,
claim_token=claim.claim_token,
error_message=message or "图片批量生成失败",
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_PROVIDER_FAILED,
exception=exc,
)
return []
try:
child_ids = await _split_children(
db,
main_task_id=main_task_id,
claim_token=claim.claim_token,
provider_result=provider_result,
provider_items=provider_items,
)
except Exception as exc:
await _fail_claimed_main(
db,
main_task_id=main_task_id,
claim_token=claim.claim_token,
error_message=f"图片批量结果拆分失败: {exc}",
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_SPLIT_FAILED,
exception=exc,
)
return []
# child 已提交后,下载投递失败不属于图片生成失败,不退款、不重新请求供应商。
enqueue_result = await _enqueue_child_downloads(db, child_ids)
if enqueue_result["failed"]:
log_operation_event(
domain="generation_ai_batch",
event_type="DOWNLOAD_ENQUEUE_FAILED",
event_status="failed",
source="celery",
user_id=claim.task_snapshot.user_id,
group_id=main_task_id,
task_id=main_task_id,
detail={
"failed_child_task_ids": enqueue_result["failed"],
"enqueued_child_task_ids": enqueue_result["enqueued"],
"provider_regenerated": False,
"generation_refunded": False,
},
)
return child_ids
@@ -1,77 +1,54 @@
from __future__ import annotations from __future__ import annotations
import json import json
from datetime import datetime, timedelta, timezone, date from datetime import datetime, date
from typing import Any from typing import Any
from fastapi import HTTPException from fastapi import HTTPException
from sqlalchemy import and_, func, select from sqlalchemy import and_, func, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
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.project import Project from app.models.project import Project
from app.models.image_engine import ImageEngine from app.models.image_engine import ImageEngine
from app.models.user import User from app.models.user import User
from app.models.video_engine import VideoEngine from app.models.video_engine import VideoEngine
from app.enums.audio_reference import ( from app.enums.generation_task import CHAT_TOP_LEVEL_MODES, GenerationMode
AUDIO_ALLOWED_EXTENSIONS,
AUDIO_MAX_COUNT_LIMIT,
AUDIO_MAX_DURATION_SECONDS,
AUDIO_MAX_TOTAL_DURATION_SECONDS,
AUDIO_MIN_DURATION_SECONDS,
)
from app.enums.generation_history import ( from app.enums.generation_history import (
GenerationHistorySourceEnum, GenerationHistorySourceEnum,
get_generation_history_source_label, get_generation_history_source_label,
get_generation_history_task_mode, get_generation_history_task_modes,
normalize_generation_history_source, normalize_generation_history_source,
HISTORY_DAY_PAGE_SIZE_MAX,
HISTORY_GROUP_ITEM_LIMIT,
) )
from app.schemas.generation_ai import ( from app.schemas.generation_ai import (
GenerationAIEngineGroupOut, GenerationAIEngineGroupOut,
GenerationAIEngineOptionsOut, GenerationAIEngineOptionsOut,
GenerationAIImageEngineOptionOut, GenerationAIImageEngineOptionOut,
GenerationAIRecordHistoryItemOut, GenerationAIRecordHistoryItemOut,
GenerationAITaskCreate,
GenerationAITaskOut, GenerationAITaskOut,
GenerationAIVideoEngineOptionOut, GenerationAIVideoEngineOptionOut,
) )
from app.services.generation_billing_service import (
OWNER_CHAT_GENERATION_TASK,
charge_generation_media_by_params,
)
from app.services.resource_accounting_service import ( from app.services.resource_accounting_service import (
SOURCE_MODEL_CHAT_TASK, SOURCE_MODEL_CHAT_TASK,
SOURCE_MODEL_GENERATION_RECORD, SOURCE_MODEL_GENERATION_RECORD,
batch_get_generated_resource_info_map, batch_get_generated_resource_info_map,
soft_delete_chat_task_resources,
) )
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.generation_history_meta_service import ( from app.services.generation.history_meta_service import (
GenerationHistoryMeta, GenerationHistoryMeta,
batch_load_generation_history_meta_map, batch_load_generation_history_meta_map,
build_empty_history_meta, build_empty_history_meta,
) )
from app.services.resource_capacity_service import assert_user_resource_capacity_available from app.services.generation.ai.task_group_service import get_display_status, load_children_map
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls, resolve_private_portrait_reference_display_urls, resolve_private_portrait_references from app.services.generation.ai.engine_service import (
from app.utils.id_gen import generate_id image_supported_sizes,
normalize_generation_count,
IMAGE_DEFAULT_SIZE = "2K" parse_json_list,
IMAGE_DEFAULT_PROPORTION = "1:1" )
IMAGE_DEFAULT_PX = "2048x2048" from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls
VIDEO_DEFAULT_DURATION = 4
VIDEO_DEFAULT_RATIO = "16:9"
VIDEO_DEFAULT_RESOLUTION = "480p"
HISTORY_DAY_PAGE_SIZE_MAX = 10
HISTORY_GROUP_ITEM_LIMIT = 10
def normalize_px(value: str | None) -> str | None:
if not value:
return value
return value.replace("×", "x").replace("X", "x").replace("×x", "x").replace("x×", "x")
def _json(data: Any) -> str | None: def _json(data: Any) -> str | None:
if data is None: if data is None:
@@ -88,7 +65,12 @@ def _parse_json(text: str | None):
return None return None
async def _resolve_task_reference_display_map(db: AsyncSession, tasks: list[ChatGenerationTask], *, user_id: str | None = None) -> dict[str, list[dict] | None]: async def _resolve_task_reference_display_map(
db: AsyncSession,
tasks: list[ChatGenerationTask],
*,
user_id: str | None = None,
) -> dict[str, list[dict] | None]:
return await batch_resolve_private_portrait_reference_display_urls( return await batch_resolve_private_portrait_reference_display_urls(
db, db,
{task.id: _parse_json(task.media_references) for task in tasks}, {task.id: _parse_json(task.media_references) for task in tasks},
@@ -96,7 +78,12 @@ async def _resolve_task_reference_display_map(db: AsyncSession, tasks: list[Chat
) )
async def _resolve_generation_record_reference_display_map(db: AsyncSession, records: list[GenerationRecord], *, user_id: str | None = None) -> dict[str, list[dict] | None]: async def _resolve_generation_record_reference_display_map(
db: AsyncSession,
records: list[GenerationRecord],
*,
user_id: str | None = None,
) -> dict[str, list[dict] | None]:
return await batch_resolve_private_portrait_reference_display_urls( return await batch_resolve_private_portrait_reference_display_urls(
db, db,
{record.id: _parse_json(record.media_references) for record in records}, {record.id: _parse_json(record.media_references) for record in records},
@@ -104,90 +91,6 @@ async def _resolve_generation_record_reference_display_map(db: AsyncSession, rec
) )
async def _get_image_engine(db: AsyncSession, engine_id: str | None) -> ImageEngine:
query = select(ImageEngine).where(ImageEngine.is_active == True)
if engine_id:
query = query.where(ImageEngine.id == engine_id)
else:
query = query.order_by(ImageEngine.priority.desc()).limit(1)
result = await db.execute(query)
engine = result.scalar_one_or_none()
if not engine:
raise HTTPException(status_code=400, detail="没有可用的图片引擎")
return engine
async def _get_video_engine(db: AsyncSession, engine_id: str | None) -> VideoEngine:
query = select(VideoEngine).where(VideoEngine.is_active == True)
if engine_id:
query = query.where(VideoEngine.id == engine_id)
else:
query = query.order_by(VideoEngine.priority.desc())
query = query.limit(1)
result = await db.execute(query)
engine = result.scalar_one_or_none()
if not engine:
raise HTTPException(status_code=400, detail="没有可用的视频引擎")
return engine
def _image_supported_sizes(engine: ImageEngine) -> dict:
try:
data = json.loads(engine.supported_sizes or "{}")
return data if isinstance(data, dict) else {}
except Exception:
return {}
def _parse_list(value: str | None, fallback: list):
try:
parsed = json.loads(value or "")
return parsed if isinstance(parsed, list) else fallback
except Exception:
return fallback
def _build_image_snapshot(engine: ImageEngine, size: str, proportion: str, px: str) -> dict:
return {
"engine_type": "image",
"id": engine.id,
"name": engine.name,
"provider": engine.provider,
"api_base": engine.api_base,
"api_key_masked": "****" if engine.api_key else "",
"model_name": engine.model_name,
"generate_url": engine.generate_url,
"supported_models": _parse_list(engine.supported_models, []),
"default_size": engine.default_size,
"selected_size": size,
"selected_proportion": proportion,
"selected_px": px,
}
def _build_video_snapshot(engine: VideoEngine, ratio: str, resolution: str, duration: int) -> dict:
return {
"engine_type": "video",
"id": engine.id,
"name": engine.name,
"provider": engine.provider,
"api_base": engine.api_base,
"api_key_masked": "****" if engine.api_key else "",
"model_name": engine.model_name,
"generate_url": engine.generate_url,
"query_url": engine.query_url,
"supported_ratios": _parse_list(engine.supported_ratios, []),
"supported_resolutions": _parse_list(engine.supported_resolutions, []),
"supported_durations": _parse_list(engine.supported_durations, []),
"max_duration": engine.max_duration,
"max_audio_count": engine.max_audio_count,
"selected_ratio": ratio,
"selected_resolution": resolution,
"selected_duration": duration,
}
async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEngineOptionsOut: async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEngineOptionsOut:
"""获取当前启用的图片/视频生成引擎,供前端创建任务时选择 engine_id。""" """获取当前启用的图片/视频生成引擎,供前端创建任务时选择 engine_id。"""
image_result = await db.execute( image_result = await db.execute(
@@ -207,11 +110,15 @@ async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEng
name=engine.name, name=engine.name,
provider=engine.provider, provider=engine.provider,
model_name=engine.model_name, model_name=engine.model_name,
supported_models=_parse_list(engine.supported_models, []), supported_models=parse_json_list(engine.supported_models, []),
supported_sizes=_image_supported_sizes(engine), supported_sizes=image_supported_sizes(engine),
default_size=engine.default_size, default_size=engine.default_size,
priority=engine.priority or 0, priority=engine.priority or 0,
max_image_count=engine.max_image_count, max_image_count=engine.max_image_count,
multi_generation_enabled=bool(getattr(engine, "multi_generation_enabled", False)),
max_generation_count=normalize_generation_count(getattr(engine, "max_generation_count", 1)),
multi_image_max_images=int(getattr(engine, "multi_image_max_images", 15) or 15),
max_reference_image_count=int(getattr(engine, "max_reference_image_count", 14) or 0),
) )
for engine in image_result.scalars().all() for engine in image_result.scalars().all()
] ]
@@ -221,9 +128,9 @@ async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEng
name=engine.name, name=engine.name,
provider=engine.provider, provider=engine.provider,
model_name=engine.model_name, model_name=engine.model_name,
supported_ratios=_parse_list(engine.supported_ratios, []), supported_ratios=parse_json_list(engine.supported_ratios, []),
supported_resolutions=_parse_list(engine.supported_resolutions, []), supported_resolutions=parse_json_list(engine.supported_resolutions, []),
supported_durations=_parse_list(engine.supported_durations, []), supported_durations=parse_json_list(engine.supported_durations, []),
max_duration=engine.max_duration, max_duration=engine.max_duration,
priority=engine.priority or 0, priority=engine.priority or 0,
max_image_count=engine.max_image_count, max_image_count=engine.max_image_count,
@@ -231,6 +138,8 @@ async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEng
max_audio_count=engine.max_audio_count, max_audio_count=engine.max_audio_count,
supports_first_last_frame=engine.supports_first_last_frame, supports_first_last_frame=engine.supports_first_last_frame,
supports_universal_reference=engine.supports_universal_reference, supports_universal_reference=engine.supports_universal_reference,
multi_generation_enabled=bool(getattr(engine, "multi_generation_enabled", False)),
max_generation_count=normalize_generation_count(getattr(engine, "max_generation_count", 1)),
) )
for engine in video_result.scalars().all() for engine in video_result.scalars().all()
] ]
@@ -240,193 +149,6 @@ async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEng
) )
async def create_async_generation_task(db: AsyncSession, current_user: User, req: GenerationAITaskCreate) -> ChatGenerationTask:
"""Create a project-independent chat generation task.
Important: this writes chat_generation_tasks, not generation_records, so chat
image/video generation no longer needs or validates a project_id.
"""
gen_type = req.gen_type.lower().strip()
if gen_type not in ("image", "video"):
raise HTTPException(status_code=400, detail="gen_type 仅支持 image 或 video")
if req.idempotency_key:
result = await db.execute(
select(ChatGenerationTask).where(
ChatGenerationTask.user_id == current_user.id,
ChatGenerationTask.idempotency_key == req.idempotency_key,
ChatGenerationTask.generation_mode == "chatapi_async",
ChatGenerationTask.deleted_at.is_(None),
).order_by(ChatGenerationTask.created_at.desc()).limit(1)
)
existing = result.scalar_one_or_none()
if existing:
return existing
refs = [r.model_dump(exclude_none=True) for r in (req.media_references or [])]
refs = await resolve_private_portrait_references(
db,
user_id=current_user.id,
media_references=refs,
gen_type=gen_type,
)
now = datetime.now(timezone.utc)
task_id = generate_id()
await assert_user_resource_capacity_available(db, current_user.id)
if gen_type == "image":
if any((r.get("type") or "").lower() == "audio" for r in refs):
raise HTTPException(status_code=400, detail="图片生成不支持音频参考素材")
engine = await _get_image_engine(db, req.engine_id)
sizes = _image_supported_sizes(engine)
size = req.image_size or engine.default_size or IMAGE_DEFAULT_SIZE
proportion = req.image_proportion or IMAGE_DEFAULT_PROPORTION
px = normalize_px(req.image_px)
if sizes:
if size not in sizes:
raise HTTPException(status_code=400, detail=f"图片分辨率档位不支持: {size}")
if proportion not in sizes.get(size, {}):
raise HTTPException(status_code=400, detail=f"图片比例不支持: {proportion}")
px = px or normalize_px((sizes.get(size) or {}).get(proportion))
px = px or IMAGE_DEFAULT_PX
media_billing = await charge_generation_media_by_params(
db,
user_id=current_user.id,
record_id=task_id,
gen_type="image",
image_size=size,
engine_id=engine.id,
project_name="AI生成任务",
description_prefix="AI创作-",
owner_type=OWNER_CHAT_GENERATION_TASK,
attempt_no=1,
)
snapshot = _build_image_snapshot(engine, size, proportion, px)
task = ChatGenerationTask(
id=task_id,
user_id=current_user.id,
original_prompt=req.original_prompt,
gen_type="image",
image_size=size,
image_proportion=proportion,
image_px=px,
status="generating",
generation_mode="chatapi_async",
pipeline_stage="queued",
engine_id=engine.id,
engine_snapshot_json=_json(snapshot),
media_references=_json(refs) if refs else None,
credits_cost=round(media_billing.total_charged, 2),
idempotency_key=req.idempotency_key,
deadline_at=now + timedelta(minutes=settings.CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES),
)
else:
engine = await _get_video_engine(db, req.engine_id)
ratio = req.aspect_ratio or VIDEO_DEFAULT_RATIO
resolution = req.resolution or VIDEO_DEFAULT_RESOLUTION
duration = req.duration or VIDEO_DEFAULT_DURATION
ratios = _parse_list(engine.supported_ratios, [])
resolutions = _parse_list(engine.supported_resolutions, [])
durations = _parse_list(engine.supported_durations, [])
if ratios and ratio not in ratios:
raise HTTPException(status_code=400, detail=f"视频比例不支持: {ratio}")
if resolutions and resolution not in resolutions:
raise HTTPException(status_code=400, detail=f"视频分辨率不支持: {resolution}")
if durations and duration not in durations:
raise HTTPException(status_code=400, detail=f"视频时长不支持: {duration}")
if engine.max_duration and duration > engine.max_duration:
raise HTTPException(status_code=400, detail=f"视频时长不能超过 {engine.max_duration}")
input_video_duration = 0.0
if refs:
video_refs = [r for r in refs if (r.get("type") or "").lower() == "video"]
for ref in video_refs:
ref_duration = float(ref.get("duration") or 0)
if ref_duration < 2:
raise HTTPException(status_code=400, detail=f"视频素材最短不能少于 2 秒")
input_video_duration += ref_duration
if input_video_duration > 15:
raise HTTPException(status_code=400, detail=f"所有视频素材总时长不能超过 15 秒,当前 {input_video_duration:.1f}")
audio_refs = [r for r in refs if (r.get("type") or "").lower() == "audio"]
if audio_refs:
max_audio_count = int(engine.max_audio_count or 0)
if max_audio_count <= 0:
raise HTTPException(status_code=400, detail="当前视频引擎不支持音频参考素材")
if max_audio_count > AUDIO_MAX_COUNT_LIMIT:
max_audio_count = AUDIO_MAX_COUNT_LIMIT
if len(audio_refs) > max_audio_count:
raise HTTPException(
status_code=400,
detail=f"参考音频最多可传 {max_audio_count} 段,当前 {len(audio_refs)}",
)
input_audio_duration = 0.0
for ref in audio_refs:
raw_duration = ref.get("duration")
if raw_duration is None:
raw_duration = 0.0
try:
ref_duration = float(raw_duration)
except (TypeError, ValueError):
ref_duration = 0.0
if ref_duration < AUDIO_MIN_DURATION_SECONDS or ref_duration > AUDIO_MAX_DURATION_SECONDS:
raise HTTPException(
status_code=400,
detail=f"单段参考音频时长必须在 {AUDIO_MIN_DURATION_SECONDS}-{AUDIO_MAX_DURATION_SECONDS} 秒之间",
)
input_audio_duration += ref_duration
if input_audio_duration > AUDIO_MAX_TOTAL_DURATION_SECONDS:
raise HTTPException(
status_code=400,
detail=f"所有参考音频总时长不能超过 {AUDIO_MAX_TOTAL_DURATION_SECONDS} 秒,当前 {input_audio_duration:.1f}",
)
media_billing = await charge_generation_media_by_params(
db,
user_id=current_user.id,
record_id=task_id,
gen_type="video",
duration=duration,
resolution=resolution,
engine_id=engine.id,
input_video_duration=input_video_duration if input_video_duration > 0 else None,
project_name="AI生成任务",
description_prefix="AI创作-",
owner_type=OWNER_CHAT_GENERATION_TASK,
attempt_no=1,
)
snapshot = _build_video_snapshot(engine, ratio, resolution, duration)
task = ChatGenerationTask(
id=task_id,
user_id=current_user.id,
original_prompt=req.original_prompt,
gen_type="video",
duration=duration,
aspect_ratio=ratio,
resolution=resolution,
image_size=req.image_size or IMAGE_DEFAULT_SIZE,
image_proportion=req.image_proportion or IMAGE_DEFAULT_PROPORTION,
image_px=normalize_px(req.image_px) or IMAGE_DEFAULT_PX,
status="generating",
generation_mode="chatapi_async",
pipeline_stage="queued",
engine_id=engine.id,
engine_snapshot_json=_json(snapshot),
media_references=_json(refs) if refs else None,
credits_cost=round(media_billing.total_charged, 2),
idempotency_key=req.idempotency_key,
deadline_at=now + timedelta(hours=settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS),
)
db.add(task)
await db.flush()
return task
def _resolve_error_message(error_message: str | None) -> str | None: def _resolve_error_message(error_message: str | None) -> str | None:
"""匹配 ARK_ERRORS 字典,将原始错误码转换为友好提示。 """匹配 ARK_ERRORS 字典,将原始错误码转换为友好提示。
app/api/v1/generation.py _record_to_out 保持一致 app/api/v1/generation.py _record_to_out 保持一致
@@ -479,26 +201,36 @@ def record_to_out(
file_name: str | None = None, file_name: str | None = None,
history_meta: GenerationHistoryMeta | None = None, history_meta: GenerationHistoryMeta | None = None,
media_references: list[dict] | None = None, media_references: list[dict] | None = None,
child_items: list[GenerationAITaskOut] | None = None,
) -> GenerationAITaskOut: ) -> GenerationAITaskOut:
refs = media_references if media_references is not None else _parse_json(task.media_references) refs = media_references if media_references is not None else _parse_json(task.media_references)
snapshot = engine_snapshot_out(_parse_json(task.engine_snapshot_json)) snapshot = engine_snapshot_out(_parse_json(task.engine_snapshot_json))
source = GenerationHistorySourceEnum.CHAT_TASK source = GenerationHistorySourceEnum.CHAT_TASK
try: try:
source = GenerationHistorySourceEnum( if task.generation_mode in {
"chat_task" if task.generation_mode == "chatapi_async" else str(task.generation_mode or "chat_task") GenerationMode.CHATAPI_ASYNC.value,
) GenerationMode.CHATAPI_MAIN.value,
GenerationMode.CHATAPI_CHILD.value,
}:
source = GenerationHistorySourceEnum.CHAT_TASK
else:
source = GenerationHistorySourceEnum(str(task.generation_mode or "chat_task"))
except ValueError: except ValueError:
source = GenerationHistorySourceEnum.CHAT_TASK source = GenerationHistorySourceEnum.CHAT_TASK
meta = history_meta or build_empty_history_meta(source) meta = history_meta or build_empty_history_meta(source)
is_deleted = task.deleted_at is not None
is_main = task.generation_mode == GenerationMode.CHATAPI_MAIN.value
hide_resource = is_deleted or is_main
return GenerationAITaskOut( return GenerationAITaskOut(
id=task.id, id=task.id,
user_id=task.user_id if is_admin else None, user_id=task.user_id if is_admin else None,
user_name=getattr(task, "username", None) if is_admin else None, user_name=getattr(task, "username", None) if is_admin else None,
project_id=None, project_id=None,
generated_resource_id=generated_resource_id, generated_resource_id=None if hide_resource else generated_resource_id,
file_name=file_name, file_name=None if hide_resource else file_name,
history_source=meta.get("history_source"), history_source=meta.get("history_source"),
history_source_label=meta.get("history_source_label"), history_source_label=meta.get("history_source_label"),
module_project_id=meta.get("module_project_id"), module_project_id=meta.get("module_project_id"),
@@ -515,10 +247,14 @@ def record_to_out(
shot_segment_label=meta.get("shot_segment_label"), shot_segment_label=meta.get("shot_segment_label"),
gen_type=task.gen_type, gen_type=task.gen_type,
generation_mode=task.generation_mode, generation_mode=task.generation_mode,
parent_task_id=task.parent_task_id,
generation_count=max(1, min(5, int(task.generation_count or 1))),
generation_index=task.generation_index,
display_status=get_display_status(task),
pipeline_stage=task.pipeline_stage, pipeline_stage=task.pipeline_stage,
status=task.status, status=task.status,
original_prompt=task.original_prompt, original_prompt=task.original_prompt,
# optimized_prompt=task.optimized_prompt, optimized_prompt=task.optimized_prompt,
duration=task.duration, duration=task.duration,
aspect_ratio=task.aspect_ratio, aspect_ratio=task.aspect_ratio,
resolution=task.resolution, resolution=task.resolution,
@@ -528,10 +264,9 @@ def record_to_out(
media_references=refs, media_references=refs,
provider_task_id=task.provider_task_id, provider_task_id=task.provider_task_id,
seedance_task_id=task.seedance_task_id, seedance_task_id=task.seedance_task_id,
# remote_result_url=task.remote_result_url, image_url="" if hide_resource else (build_resource_signed_url(task.image_url) if task.image_url else ""),
image_url=build_resource_signed_url(task.image_url) if task.image_url else "", video_url="" if hide_resource else (build_resource_signed_url(task.video_url) if task.video_url else ""),
video_url=build_resource_signed_url(task.video_url) if task.video_url else "", video_cover_url="" if hide_resource else (build_resource_signed_url(task.video_cover_url) if task.video_cover_url else ""),
video_cover_url=build_resource_signed_url(task.video_cover_url) if task.video_cover_url else "",
engine_id=task.engine_id, engine_id=task.engine_id,
engine_snapshot=snapshot, engine_snapshot=snapshot,
credits_cost=task.credits_cost or 0.0, credits_cost=task.credits_cost or 0.0,
@@ -541,33 +276,90 @@ def record_to_out(
video_tokens_used=task.video_tokens_used or 0, video_tokens_used=task.video_tokens_used or 0,
retry_count=task.retry_count or 0, retry_count=task.retry_count or 0,
poll_count=task.poll_count or 0, poll_count=task.poll_count or 0,
error_message=_resolve_error_message(task.error_message), error_message=task.error_message if is_main else _resolve_error_message(task.error_message),
created_at=task.created_at, created_at=task.created_at,
generated_at=task.generated_at, generated_at=task.generated_at,
child_items=child_items or [],
) )
def engine_snapshot_out(snapshot: dict) -> dict: def engine_snapshot_out(snapshot: dict) -> dict:
""" """从完整引擎快照中过滤前端允许展示的字段。"""
从完整的 engine_snapshot 中过滤出需要返回的字段
"""
if not snapshot: if not snapshot:
return {} return {}
keys = (
"engine_type", "id", "name", "provider", "model_name",
"supported_models", "default_size", "selected_size",
"selected_proportion", "selected_px", "supported_ratios",
"supported_resolutions", "supported_durations", "max_duration",
"max_audio_count", "selected_ratio", "selected_resolution",
"selected_duration", "generation_count", "multi_generation_enabled",
"max_generation_count", "multi_image_max_images", "max_reference_image_count", "output_format",
)
result = {key: snapshot.get(key) for key in keys if key in snapshot}
result.setdefault("generation_count", 1)
return result
return {
"engine_type": snapshot.get("engine_type"), async def build_task_out_list(
"id": snapshot.get("id"), db: AsyncSession,
"name": snapshot.get("name"), tasks: list[ChatGenerationTask],
"provider": snapshot.get("provider"), *,
# "api_base": snapshot.get("api_base"), is_admin: bool = False,
# "api_key_masked": snapshot.get("api_key_masked"), viewer_user_id: str | None = None,
"model_name": snapshot.get("model_name"), ) -> list[GenerationAITaskOut]:
# "generate_url": snapshot.get("generate_url"), """批量回填主任务子项、资源账本和参考素材,避免列表 N+1。"""
"supported_models": snapshot.get("supported_models", []), if not tasks:
"default_size": snapshot.get("default_size"), return []
"selected_size": snapshot.get("selected_size"), parent_ids = [
"selected_proportion": snapshot.get("selected_proportion"), task.id for task in tasks
"selected_px": snapshot.get("selected_px") if task.generation_mode == GenerationMode.CHATAPI_MAIN.value
} ]
children_map = await load_children_map(db, parent_ids, include_deleted=True)
children = [child for items in children_map.values() for child in items]
resource_task_ids = [
task.id for task in [*tasks, *children]
if task.generation_mode != GenerationMode.CHATAPI_MAIN.value and task.deleted_at is None
]
resource_info_map = await batch_get_generated_resource_info_map(
db,
source_model=SOURCE_MODEL_CHAT_TASK,
source_ids=resource_task_ids,
)
reference_display_map = await _resolve_task_reference_display_map(
db,
tasks,
user_id=viewer_user_id,
)
output: list[GenerationAITaskOut] = []
for task in tasks:
refs = reference_display_map.get(task.id)
child_out: list[GenerationAITaskOut] = []
for child in children_map.get(task.id, []):
if is_admin:
child.username = getattr(task, "username", None)
resource = resource_info_map.get(child.id, {})
child_out.append(
record_to_out(
child,
is_admin=is_admin,
generated_resource_id=resource.get("resource_id"),
file_name=resource.get("file_name"),
media_references=refs,
)
)
resource = resource_info_map.get(task.id, {})
output.append(
record_to_out(
task,
is_admin=is_admin,
generated_resource_id=resource.get("resource_id"),
file_name=resource.get("file_name"),
media_references=refs,
child_items=child_out,
)
)
return output
async def list_async_generation_tasks( async def list_async_generation_tasks(
db: AsyncSession, db: AsyncSession,
@@ -594,7 +386,7 @@ async def list_async_generation_tasks(
query = select(ChatGenerationTask) query = select(ChatGenerationTask)
query = query.where( query = query.where(
ChatGenerationTask.generation_mode == "chatapi_async", ChatGenerationTask.generation_mode.in_(list(CHAT_TOP_LEVEL_MODES)),
ChatGenerationTask.deleted_at.is_(None), ChatGenerationTask.deleted_at.is_(None),
) )
@@ -623,7 +415,7 @@ async def list_async_generation_tasks(
total = (await db.execute(count_query)).scalar_one() total = (await db.execute(count_query)).scalar_one()
result = await db.execute( result = await db.execute(
query.order_by(ChatGenerationTask.created_at.desc()) query.order_by(ChatGenerationTask.created_at.desc(), ChatGenerationTask.id.desc())
.offset((page - 1) * page_size) .offset((page - 1) * page_size)
.limit(page_size) .limit(page_size)
) )
@@ -677,12 +469,12 @@ def _parse_history_date(value: str) -> date:
def _history_base_filters(user_id: str, gen_type: str, source: GenerationHistorySourceEnum): def _history_base_filters(user_id: str, gen_type: str, source: GenerationHistorySourceEnum):
task_mode = get_generation_history_task_mode(source) task_modes = get_generation_history_task_modes(source)
if not task_mode: if not task_modes:
raise HTTPException(status_code=400, detail="history_source 不支持查询 ChatGenerationTask 历史") raise HTTPException(status_code=400, detail="history_source 不支持查询 ChatGenerationTask 历史")
return [ return [
ChatGenerationTask.user_id == user_id, ChatGenerationTask.user_id == user_id,
ChatGenerationTask.generation_mode == task_mode.value, ChatGenerationTask.generation_mode.in_([mode.value for mode in task_modes]),
ChatGenerationTask.deleted_at.is_(None), ChatGenerationTask.deleted_at.is_(None),
ChatGenerationTask.status == "completed", ChatGenerationTask.status == "completed",
ChatGenerationTask.gen_type == gen_type, ChatGenerationTask.gen_type == gen_type,
@@ -1193,13 +985,3 @@ async def list_generation_history_day_items(
], ],
} }
async def soft_delete_chat_generation_task(
db: AsyncSession,
*,
task: ChatGenerationTask,
deleted_at: datetime | None = None,
) -> int:
"""软删 ChatGenerationTask 并联动软删资源账本,返回释放的 active 空间字节数。"""
deleted_at = deleted_at or datetime.now(timezone.utc)
task.deleted_at = deleted_at
return await soft_delete_chat_task_resources(db, task.id, deleted_at=deleted_at)
@@ -0,0 +1,587 @@
from __future__ import annotations
import json
from dataclasses import dataclass, field
from datetime import datetime, timedelta, timezone
from typing import Any
from fastapi import HTTPException
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.enums.audio_reference import (
AUDIO_MAX_COUNT_LIMIT,
AUDIO_MAX_DURATION_SECONDS,
AUDIO_MAX_TOTAL_DURATION_SECONDS,
AUDIO_MIN_DURATION_SECONDS,
)
from app.enums.generation_task import CHAT_TOP_LEVEL_MODES, GenerationMode, GenerationType
from app.models.chat_generation_task import ChatGenerationTask
from app.models.user import User
from app.schemas.generation_ai import GenerationAITaskCreate
from app.services.generation.ai.engine_service import (
IMAGE_DEFAULT_PROPORTION,
IMAGE_DEFAULT_PX,
IMAGE_DEFAULT_SIZE,
VIDEO_DEFAULT_DURATION,
VIDEO_DEFAULT_RATIO,
VIDEO_DEFAULT_RESOLUTION,
build_image_snapshot,
build_video_snapshot,
get_image_engine,
get_video_engine,
image_supported_sizes,
normalize_generation_count,
normalize_px,
parse_json_list,
)
from app.services.generation.billing_service import OWNER_CHAT_GENERATION_TASK, charge_generation_media_by_params
from app.services.operation_log_service import log_operation_event
from app.services.private_portrait.reference_resolver import resolve_private_portrait_references
from app.services.resource_capacity_service import assert_user_resource_capacity_available
from app.utils.id_gen import generate_id
@dataclass(slots=True)
class GenerationTaskCreateResult:
top_level_task_id: str
enqueue_task_ids: list[str] = field(default_factory=list)
child_task_ids: list[str] = field(default_factory=list)
generation_count: int = 1
gen_type: str = GenerationType.IMAGE.value
created: bool = True
def _json(data: Any) -> str | None:
if data is None:
return None
return json.dumps(data, ensure_ascii=False, default=str)
async def find_existing_top_level_task(
db: AsyncSession,
*,
user_id: str,
idempotency_key: str | None,
) -> ChatGenerationTask | None:
if not idempotency_key:
return None
result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.user_id == user_id,
ChatGenerationTask.idempotency_key == idempotency_key,
ChatGenerationTask.generation_mode.in_(list(CHAT_TOP_LEVEL_MODES)),
ChatGenerationTask.deleted_at.is_(None),
)
.order_by(ChatGenerationTask.created_at.desc())
.limit(1)
)
return result.scalar_one_or_none()
def _validate_video_references(refs: list[dict], *, max_audio_count: int) -> float:
input_video_duration = 0.0
for ref in refs:
if (ref.get("type") or "").lower() != GenerationType.VIDEO.value:
continue
try:
ref_duration = float(ref.get("duration") or 0)
except (TypeError, ValueError):
ref_duration = 0.0
if ref_duration < 2:
raise HTTPException(status_code=400, detail="视频素材最短不能少于 2 秒")
input_video_duration += ref_duration
if input_video_duration > 15:
raise HTTPException(status_code=400, detail=f"所有视频素材总时长不能超过 15 秒,当前 {input_video_duration:.1f}")
audio_refs = [ref for ref in refs if (ref.get("type") or "").lower() == "audio"]
if audio_refs:
allowed_count = min(AUDIO_MAX_COUNT_LIMIT, max(0, int(max_audio_count or 0)))
if allowed_count <= 0:
raise HTTPException(status_code=400, detail="当前视频引擎不支持音频参考素材")
if len(audio_refs) > allowed_count:
raise HTTPException(status_code=400, detail=f"参考音频最多可传 {allowed_count} 段,当前 {len(audio_refs)}")
input_audio_duration = 0.0
for ref in audio_refs:
try:
ref_duration = float(ref.get("duration") or 0)
except (TypeError, ValueError):
ref_duration = 0.0
if ref_duration < AUDIO_MIN_DURATION_SECONDS or ref_duration > AUDIO_MAX_DURATION_SECONDS:
raise HTTPException(
status_code=400,
detail=f"单段参考音频时长必须在 {AUDIO_MIN_DURATION_SECONDS}-{AUDIO_MAX_DURATION_SECONDS} 秒之间",
)
input_audio_duration += ref_duration
if input_audio_duration > AUDIO_MAX_TOTAL_DURATION_SECONDS:
raise HTTPException(
status_code=400,
detail=f"所有参考音频总时长不能超过 {AUDIO_MAX_TOTAL_DURATION_SECONDS} 秒,当前 {input_audio_duration:.1f}",
)
return input_video_duration
def _base_task_kwargs(
*,
task_id: str,
user_id: str,
req: GenerationAITaskCreate,
gen_type: str,
generation_mode: str,
generation_count: int,
engine_id: str,
engine_snapshot_json: str,
media_references_json: str | None,
deadline_at: datetime,
parent_task_id: str | None = None,
generation_index: int | None = None,
credits_cost: float = 0.0,
idempotency_key: str | None = None,
) -> dict[str, Any]:
return {
"id": task_id,
"user_id": user_id,
"original_prompt": req.original_prompt,
"gen_type": gen_type,
"status": "generating",
"generation_mode": generation_mode,
"pipeline_stage": "queued",
"parent_task_id": parent_task_id,
"generation_count": generation_count,
"generation_index": generation_index,
"engine_id": engine_id,
"engine_snapshot_json": engine_snapshot_json,
"media_references": media_references_json,
"credits_cost": round(float(credits_cost or 0), 2),
"idempotency_key": idempotency_key,
"deadline_at": deadline_at,
}
async def create_generation_task_group(
db: AsyncSession,
current_user: User,
req: GenerationAITaskCreate,
) -> GenerationTaskCreateResult:
"""创建单份 chatapi_async 或多份 chatapi_main/chatapi_child 任务组。
本函数只 flush,不主动 commit。调用方提交成功后才能投递 Celery。
"""
gen_type = (req.gen_type or "").lower().strip()
if gen_type not in (GenerationType.IMAGE.value, GenerationType.VIDEO.value):
raise HTTPException(status_code=400, detail="gen_type 仅支持 image 或 video")
existing = await find_existing_top_level_task(
db,
user_id=current_user.id,
idempotency_key=req.idempotency_key,
)
if existing:
return GenerationTaskCreateResult(
top_level_task_id=existing.id,
generation_count=int(existing.generation_count or 1),
gen_type=existing.gen_type,
created=False,
)
refs = [item.model_dump(exclude_none=True) for item in (req.media_references or [])]
refs = await resolve_private_portrait_references(
db,
user_id=current_user.id,
media_references=refs,
gen_type=gen_type,
)
media_references_json = _json(refs) if refs else None
await assert_user_resource_capacity_available(db, current_user.id)
now = datetime.now(timezone.utc)
main_id = generate_id()
child_ids: list[str] = []
enqueue_ids: list[str] = []
total_billed_credits = 0.0
log_operation_event(
domain="generation_ai_batch",
event_type="BATCH_CREATE_START",
event_status="started",
source="service",
user_id=current_user.id,
group_id=main_id,
detail={
"gen_type": gen_type,
"requested_generation_count": normalize_generation_count(req.generation_count),
"idempotency_key_present": bool(req.idempotency_key),
},
)
if gen_type == GenerationType.IMAGE.value:
if any((ref.get("type") or "").lower() == "audio" for ref in refs):
raise HTTPException(status_code=400, detail="图片生成不支持音频参考素材")
engine = await get_image_engine(db, req.engine_id)
generation_count = normalize_generation_count(req.generation_count)
multi_generation_enabled = bool(getattr(engine, "multi_generation_enabled", False))
max_generation_count = normalize_generation_count(getattr(engine, "max_generation_count", 1))
if generation_count > 1 and not multi_generation_enabled:
raise HTTPException(status_code=400, detail="当前图片引擎未开启多份生成,本次生成数量只能为 1")
if generation_count > max_generation_count:
raise HTTPException(
status_code=400,
detail=f"当前图片引擎本次最多允许生成 {max_generation_count}",
)
reference_image_count = sum(
1 for ref in refs if (ref.get("type") or "").lower() == GenerationType.IMAGE.value
)
max_reference_count = max(0, int(getattr(engine, "max_reference_image_count", 14) or 0))
multi_image_max_images = max(1, int(getattr(engine, "multi_image_max_images", 15) or 15))
if reference_image_count > max_reference_count:
raise HTTPException(
status_code=400,
detail=f"当前图片引擎最多支持 {max_reference_count} 张参考图,当前 {reference_image_count}",
)
if generation_count > 1 and reference_image_count + generation_count > multi_image_max_images:
raise HTTPException(
status_code=400,
detail=(
f"参考图数量与生成数量合计不能超过 {multi_image_max_images} 张,"
f"当前参考图 {reference_image_count} 张、生成 {generation_count}"
),
)
sizes = image_supported_sizes(engine)
size = req.image_size or engine.default_size or IMAGE_DEFAULT_SIZE
proportion = req.image_proportion or IMAGE_DEFAULT_PROPORTION
px = normalize_px(req.image_px)
if sizes:
if size not in sizes:
raise HTTPException(status_code=400, detail=f"图片分辨率档位不支持: {size}")
if proportion not in sizes.get(size, {}):
raise HTTPException(status_code=400, detail=f"图片比例不支持: {proportion}")
px = px or normalize_px((sizes.get(size) or {}).get(proportion))
px = px or IMAGE_DEFAULT_PX
mode = GenerationMode.CHATAPI_ASYNC.value if generation_count == 1 else GenerationMode.CHATAPI_MAIN.value
billing = await charge_generation_media_by_params(
db,
user_id=current_user.id,
record_id=main_id,
gen_type=GenerationType.IMAGE.value,
image_size=size,
engine_id=engine.id,
project_name="AI生成任务",
description_prefix="AI创作-",
owner_type=OWNER_CHAT_GENERATION_TASK,
attempt_no=1,
quantity=generation_count,
)
image_snapshot = build_image_snapshot(engine, size, proportion, px)
image_snapshot["generation_count"] = generation_count
snapshot_json = _json(image_snapshot) or "{}"
total_billed_credits = round(float(billing.total_charged or 0), 2)
task = ChatGenerationTask(
**_base_task_kwargs(
task_id=main_id,
user_id=current_user.id,
req=req,
gen_type=GenerationType.IMAGE.value,
generation_mode=mode,
generation_count=generation_count,
engine_id=engine.id,
engine_snapshot_json=snapshot_json,
media_references_json=media_references_json,
deadline_at=now + timedelta(minutes=settings.CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES),
credits_cost=billing.total_charged,
idempotency_key=req.idempotency_key,
),
image_size=size,
image_proportion=proportion,
image_px=px,
)
db.add(task)
enqueue_ids.append(task.id)
else:
engine = await get_video_engine(db, req.engine_id)
generation_count = normalize_generation_count(req.generation_count)
multi_generation_enabled = bool(getattr(engine, "multi_generation_enabled", False))
max_generation_count = normalize_generation_count(getattr(engine, "max_generation_count", 1))
if generation_count > 1 and not multi_generation_enabled:
raise HTTPException(status_code=400, detail="当前视频引擎未开启多份生成,本次生成数量只能为 1")
if generation_count > max_generation_count:
raise HTTPException(
status_code=400,
detail=f"当前视频引擎本次最多允许生成 {max_generation_count}",
)
ratio = req.aspect_ratio or VIDEO_DEFAULT_RATIO
resolution = req.resolution or VIDEO_DEFAULT_RESOLUTION
duration = req.duration or VIDEO_DEFAULT_DURATION
ratios = parse_json_list(engine.supported_ratios, [])
resolutions = parse_json_list(engine.supported_resolutions, [])
durations = parse_json_list(engine.supported_durations, [])
if ratios and ratio not in ratios:
raise HTTPException(status_code=400, detail=f"视频比例不支持: {ratio}")
if resolutions and resolution not in resolutions:
raise HTTPException(status_code=400, detail=f"视频分辨率不支持: {resolution}")
if durations and duration not in durations:
raise HTTPException(status_code=400, detail=f"视频时长不支持: {duration}")
if engine.max_duration and duration > engine.max_duration:
raise HTTPException(status_code=400, detail=f"视频时长不能超过 {engine.max_duration}")
input_video_duration = _validate_video_references(refs, max_audio_count=engine.max_audio_count)
video_snapshot = build_video_snapshot(engine, ratio, resolution, duration)
video_snapshot["generation_count"] = generation_count
snapshot_json = _json(video_snapshot) or "{}"
deadline_at = now + timedelta(hours=settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS)
if generation_count == 1:
billing = await charge_generation_media_by_params(
db,
user_id=current_user.id,
record_id=main_id,
gen_type=GenerationType.VIDEO.value,
duration=duration,
resolution=resolution,
engine_id=engine.id,
input_video_duration=input_video_duration if input_video_duration > 0 else None,
project_name="AI生成任务",
description_prefix="AI创作-",
owner_type=OWNER_CHAT_GENERATION_TASK,
attempt_no=1,
)
total_billed_credits = round(float(billing.total_charged or 0), 2)
task = ChatGenerationTask(
**_base_task_kwargs(
task_id=main_id,
user_id=current_user.id,
req=req,
gen_type=GenerationType.VIDEO.value,
generation_mode=GenerationMode.CHATAPI_ASYNC.value,
generation_count=1,
engine_id=engine.id,
engine_snapshot_json=snapshot_json,
media_references_json=media_references_json,
deadline_at=deadline_at,
credits_cost=billing.total_charged,
idempotency_key=req.idempotency_key,
),
duration=duration,
aspect_ratio=ratio,
resolution=resolution,
image_size=req.image_size or IMAGE_DEFAULT_SIZE,
image_proportion=req.image_proportion or IMAGE_DEFAULT_PROPORTION,
image_px=normalize_px(req.image_px) or IMAGE_DEFAULT_PX,
)
db.add(task)
enqueue_ids.append(task.id)
else:
main_task = ChatGenerationTask(
**_base_task_kwargs(
task_id=main_id,
user_id=current_user.id,
req=req,
gen_type=GenerationType.VIDEO.value,
generation_mode=GenerationMode.CHATAPI_MAIN.value,
generation_count=generation_count,
engine_id=engine.id,
engine_snapshot_json=snapshot_json,
media_references_json=media_references_json,
deadline_at=deadline_at,
idempotency_key=req.idempotency_key,
),
duration=duration,
aspect_ratio=ratio,
resolution=resolution,
image_size=req.image_size or IMAGE_DEFAULT_SIZE,
image_proportion=req.image_proportion or IMAGE_DEFAULT_PROPORTION,
image_px=normalize_px(req.image_px) or IMAGE_DEFAULT_PX,
)
db.add(main_task)
await db.flush()
total_credits = 0.0
children: list[ChatGenerationTask] = []
for generation_index in range(1, generation_count + 1):
child_id = generate_id()
billing = await charge_generation_media_by_params(
db,
user_id=current_user.id,
record_id=child_id,
gen_type=GenerationType.VIDEO.value,
duration=duration,
resolution=resolution,
engine_id=engine.id,
input_video_duration=input_video_duration if input_video_duration > 0 else None,
project_name="AI生成任务",
description_prefix=f"AI创作-第{generation_index}份-",
owner_type=OWNER_CHAT_GENERATION_TASK,
attempt_no=1,
)
child = ChatGenerationTask(
**_base_task_kwargs(
task_id=child_id,
user_id=current_user.id,
req=req,
gen_type=GenerationType.VIDEO.value,
generation_mode=GenerationMode.CHATAPI_CHILD.value,
generation_count=generation_count,
generation_index=generation_index,
parent_task_id=main_id,
engine_id=engine.id,
engine_snapshot_json=snapshot_json,
media_references_json=media_references_json,
deadline_at=deadline_at,
credits_cost=billing.total_charged,
),
duration=duration,
aspect_ratio=ratio,
resolution=resolution,
image_size=req.image_size or IMAGE_DEFAULT_SIZE,
image_proportion=req.image_proportion or IMAGE_DEFAULT_PROPORTION,
image_px=normalize_px(req.image_px) or IMAGE_DEFAULT_PX,
)
children.append(child)
child_ids.append(child_id)
enqueue_ids.append(child_id)
total_credits = round(total_credits + billing.total_charged, 2)
db.add_all(children)
main_task.credits_cost = total_credits
total_billed_credits = total_credits
await db.flush()
log_operation_event(
domain="generation_ai_batch",
event_type="BATCH_BILLING_SUCCESS",
event_status="success",
source="service",
user_id=current_user.id,
group_id=main_id,
detail={
"gen_type": gen_type,
"generation_count": generation_count,
"total_billed_credits": total_billed_credits,
},
)
log_operation_event(
domain="generation_ai_batch",
event_type="BATCH_CHILDREN_CREATED" if child_ids else "BATCH_MAIN_CREATED",
event_status="success",
source="service",
user_id=current_user.id,
group_id=main_id,
detail={
"gen_type": gen_type,
"generation_count": generation_count,
"child_task_ids": child_ids,
"enqueue_task_ids": enqueue_ids,
},
)
return GenerationTaskCreateResult(
top_level_task_id=main_id,
enqueue_task_ids=enqueue_ids,
child_task_ids=child_ids,
generation_count=generation_count,
gen_type=gen_type,
created=True,
)
async def enqueue_created_generation_tasks(
db: AsyncSession,
*,
task_ids: list[str],
) -> list[str]:
"""在业务事务提交后投递任务;返回投递失败的任务ID。
投递失败会在补偿事务中将对应任务置为失败并幂等退款,视频子任务
同时触发主任务状态汇总。调用方不应在初始事务提交前调用本函数。
"""
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
from app.services.generation.log_service import log_task_event
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
from app.tasks.generation_create_tasks import chatapi_create_generation_task
normalized_ids = list(dict.fromkeys(str(item) for item in task_ids if item))
meta_result = await db.execute(
select(
ChatGenerationTask.id,
ChatGenerationTask.user_id,
ChatGenerationTask.parent_task_id,
ChatGenerationTask.generation_index,
).where(ChatGenerationTask.id.in_(normalized_ids))
) if normalized_ids else None
task_meta = {
str(row.id): {
"user_id": str(row.user_id),
"parent_task_id": str(row.parent_task_id) if row.parent_task_id else None,
"generation_index": row.generation_index,
}
for row in (meta_result.all() if meta_result is not None else [])
}
failed_ids: list[str] = []
for task_id in normalized_ids:
meta = task_meta.get(task_id, {})
log_operation_event(
domain="generation_ai_batch",
event_type="CHILD_ENQUEUE_START",
event_status="started",
source="api",
user_id=meta.get("user_id"),
group_id=meta.get("parent_task_id") or task_id,
task_id=task_id,
detail={"generation_index": meta.get("generation_index")},
)
try:
chatapi_create_generation_task.delay(task_id)
await log_task_event(
task_id=task_id,
event_type="CHILD_ENQUEUE_SUCCESS",
to_status="generating",
to_stage="queued",
detail={"task_id": task_id},
)
log_operation_event(
domain="generation_ai_batch",
event_type="CHILD_ENQUEUE_SUCCESS",
event_status="success",
source="api",
user_id=meta.get("user_id"),
group_id=meta.get("parent_task_id") or task_id,
task_id=task_id,
detail={"generation_index": meta.get("generation_index")},
)
except Exception as exc:
failed_ids.append(task_id)
await db.rollback()
failed_task = await mark_chat_generation_task_failed_and_refund_once(
db,
task_id=task_id,
error_message=f"任务队列投递失败: {exc}",
pipeline_stage="failed",
)
await aggregate_parent_for_child(db, failed_task)
await db.commit()
await log_task_event(
task_id=task_id,
event_type="CHILD_ENQUEUE_FAILED",
to_status="failed",
to_stage="failed",
message=str(exc),
detail={"task_id": task_id},
)
log_operation_event(
domain="generation_ai_batch",
event_type="CHILD_ENQUEUE_FAILED",
event_status="failed",
source="api",
user_id=getattr(failed_task, "user_id", None),
group_id=getattr(failed_task, "parent_task_id", None) or task_id,
task_id=task_id,
message=str(exc),
detail={"physical_files_deleted": False},
error=str(exc),
)
return failed_ids
@@ -0,0 +1,426 @@
from __future__ import annotations
from collections import Counter, defaultdict
from datetime import datetime, timezone
from typing import Iterable, Sequence
from fastapi import HTTPException
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.generation_task import (
ChatGenerationPipelineStage,
ChatGenerationTaskStatus,
GenerationMode,
)
from app.models.chat_generation_task import ChatGenerationTask
from app.services.operation_log_service import log_operation_event
from app.services.resource_accounting_service import (
SOURCE_MODEL_CHAT_TASK,
soft_delete_resources_by_source,
)
ACTIVE_STAGES = {
ChatGenerationPipelineStage.QUEUED.value,
ChatGenerationPipelineStage.PREPARING.value,
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
ChatGenerationPipelineStage.WAITING_REMOTE.value,
ChatGenerationPipelineStage.POLLING.value,
ChatGenerationPipelineStage.RESULT_READY.value,
ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value,
ChatGenerationPipelineStage.DOWNLOADING.value,
ChatGenerationPipelineStage.RETRY_WAITING.value,
}
def is_task_active(task: ChatGenerationTask) -> bool:
return task.deleted_at is None and (
task.status == ChatGenerationTaskStatus.GENERATING.value
or (task.pipeline_stage or "") in ACTIVE_STAGES
)
def get_display_status(task: ChatGenerationTask) -> str:
if task.deleted_at is not None:
return "deleted"
if task.pipeline_stage == ChatGenerationPipelineStage.DOWNLOAD_FAILED.value:
return "download_failed"
return task.status or ChatGenerationTaskStatus.PENDING.value
async def load_children_map(
db: AsyncSession,
parent_ids: Sequence[str] | Iterable[str],
*,
include_deleted: bool = True,
) -> dict[str, list[ChatGenerationTask]]:
ids = list(dict.fromkeys(str(item) for item in parent_ids if item))
if not ids:
return {}
query = select(ChatGenerationTask).where(
ChatGenerationTask.parent_task_id.in_(ids),
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
)
if not include_deleted:
query = query.where(ChatGenerationTask.deleted_at.is_(None))
result = await db.execute(
query.order_by(
ChatGenerationTask.parent_task_id.asc(),
ChatGenerationTask.generation_index.asc(),
ChatGenerationTask.created_at.asc(),
)
)
grouped: dict[str, list[ChatGenerationTask]] = defaultdict(list)
for task in result.scalars().all():
if task.parent_task_id:
grouped[task.parent_task_id].append(task)
return dict(grouped)
async def load_task_and_children(
db: AsyncSession,
*,
task_id: str,
user_id: str | None = None,
include_deleted_children: bool = True,
) -> tuple[ChatGenerationTask | None, list[ChatGenerationTask]]:
query = select(ChatGenerationTask).where(ChatGenerationTask.id == task_id)
if user_id:
query = query.where(ChatGenerationTask.user_id == user_id)
result = await db.execute(query.limit(1))
task = result.scalar_one_or_none()
if not task:
return None, []
if task.generation_mode == GenerationMode.CHATAPI_CHILD.value and task.parent_task_id:
parent_result = await db.execute(
select(ChatGenerationTask).where(ChatGenerationTask.id == task.parent_task_id).limit(1)
)
parent = parent_result.scalar_one_or_none()
return parent or task, [task]
if task.generation_mode != GenerationMode.CHATAPI_MAIN.value:
return task, []
children_map = await load_children_map(
db,
[task.id],
include_deleted=include_deleted_children,
)
return task, children_map.get(task.id, [])
def _generation_result_status(task: ChatGenerationTask) -> str:
"""返回任务真实生成结果,不受资源软删除影响。"""
if task.pipeline_stage == ChatGenerationPipelineStage.DOWNLOAD_FAILED.value:
return "download_failed"
if task.status == ChatGenerationTaskStatus.FAILED.value or (task.pipeline_stage or "") in {
ChatGenerationPipelineStage.FAILED.value,
ChatGenerationPipelineStage.TIMEOUT.value,
}:
return "failed"
if is_task_active(task):
return "generating"
if task.status == ChatGenerationTaskStatus.COMPLETED.value:
return "completed"
return task.status or "pending"
def _build_summary(children: list[ChatGenerationTask]) -> str | None:
if not children:
return None
result_counters: Counter[str] = Counter(_generation_result_status(child) for child in children)
labels = {
"completed": "完成",
"failed": "生成失败",
"download_failed": "下载失败",
"generating": "生成中",
"pending": "待处理",
}
parts = [f"{count}{labels.get(status, status)}" for status, count in result_counters.items() if count]
deleted_count = sum(1 for child in children if child.deleted_at is not None)
if deleted_count:
parts.append(f"{deleted_count}项资源已删除")
return f"{len(children)}项中" + "".join(parts)
async def aggregate_main_task_status(
db: AsyncSession,
*,
parent_task_id: str,
) -> ChatGenerationTask | None:
result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == parent_task_id,
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
)
.with_for_update()
.limit(1)
)
main = result.scalar_one_or_none()
if not main or main.deleted_at is not None:
return main
children_map = await load_children_map(db, [parent_task_id], include_deleted=True)
children = children_map.get(parent_task_id, [])
if not children:
return main
previous_status = main.status
previous_stage = main.pipeline_stage
active_children = [child for child in children if is_task_active(child)]
failed_children = [
child
for child in children
if (
child.status == ChatGenerationTaskStatus.FAILED.value
or (child.pipeline_stage or "") in {
ChatGenerationPipelineStage.FAILED.value,
ChatGenerationPipelineStage.TIMEOUT.value,
ChatGenerationPipelineStage.DOWNLOAD_FAILED.value,
}
)
]
completed_children = [
child for child in children if child.status == ChatGenerationTaskStatus.COMPLETED.value
]
if active_children:
main.status = ChatGenerationTaskStatus.GENERATING.value
main.pipeline_stage = active_children[0].pipeline_stage or ChatGenerationPipelineStage.QUEUED.value
main.generated_at = None
main.error_message = _build_summary(children)
elif failed_children:
main.status = ChatGenerationTaskStatus.FAILED.value
main.pipeline_stage = (
ChatGenerationPipelineStage.DOWNLOAD_FAILED.value
if any(child.pipeline_stage == ChatGenerationPipelineStage.DOWNLOAD_FAILED.value for child in failed_children)
else ChatGenerationPipelineStage.FAILED.value
)
main.generated_at = max(
(child.generated_at for child in completed_children if child.generated_at),
default=datetime.now(timezone.utc),
)
main.error_message = _build_summary(children)
else:
# 所有子任务真实生成结果均成功;资源是否软删除不改变生成历史终态。
main.status = ChatGenerationTaskStatus.COMPLETED.value
main.pipeline_stage = ChatGenerationPipelineStage.DONE.value
main.generated_at = max(
(child.generated_at for child in children if child.generated_at),
default=main.generated_at or datetime.now(timezone.utc),
)
main.error_message = None
if main.gen_type == "video":
main.credits_cost = round(sum(float(child.credits_cost or 0) for child in children), 2)
main.text_credits_cost = round(sum(float(child.text_credits_cost or 0) for child in children), 2)
main.text_tokens_used = sum(int(child.text_tokens_used or 0) for child in children)
main.image_tokens_used = sum(int(child.image_tokens_used or 0) for child in children)
main.video_tokens_used = sum(int(child.video_tokens_used or 0) for child in children)
main.retry_count = sum(int(child.retry_count or 0) for child in children)
main.poll_count = sum(int(child.poll_count or 0) for child in children)
await db.flush()
log_operation_event(
domain="generation_ai_batch",
event_type="MAIN_STATUS_AGGREGATED",
event_status="success",
source="service",
user_id=main.user_id,
group_id=main.id,
task_id=main.id,
detail={
"before_status": previous_status,
"before_stage": previous_stage,
"after_status": main.status,
"after_stage": main.pipeline_stage,
"summary": _build_summary(children),
},
)
return main
async def aggregate_parent_for_child(db: AsyncSession, child: ChatGenerationTask | None) -> ChatGenerationTask | None:
if not child or child.generation_mode != GenerationMode.CHATAPI_CHILD.value or not child.parent_task_id:
return None
# 项目关闭了 autoflush,先显式 flush 子任务的终态,确保聚合查询读取到本事务最新状态。
await db.flush()
return await aggregate_main_task_status(db, parent_task_id=str(child.parent_task_id))
async def soft_delete_child_tasks_batch(
db: AsyncSession,
*,
child_task_ids: Sequence[str] | Iterable[str],
user_id: str,
deleted_at: datetime | None = None,
require_completed: bool = False,
) -> int:
ids = list(dict.fromkeys(str(item) for item in child_task_ids if item))
if not ids:
return 0
deleted_at = deleted_at or datetime.now(timezone.utc)
result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.id.in_(ids),
ChatGenerationTask.user_id == user_id,
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
)
.with_for_update()
)
children = list(result.scalars().all())
found_ids = {str(child.id) for child in children}
missing_ids = [item for item in ids if item not in found_ids]
if missing_ids:
raise HTTPException(status_code=404, detail=f"子任务不存在: {','.join(missing_ids)}")
active_children = [child for child in children if child.deleted_at is None]
running_ids = [child.id for child in active_children if is_task_active(child)]
if running_ids:
raise HTTPException(status_code=400, detail=f"仍有 {len(running_ids)} 个子任务生成中,暂不能删除")
if require_completed:
invalid_ids = [
child.id for child in active_children
if child.status != ChatGenerationTaskStatus.COMPLETED.value or child.generated_at is None
]
if invalid_ids:
raise HTTPException(status_code=409, detail=f"只有生成完成的资源才能从素材云删除: {','.join(invalid_ids)}")
source_ids = [str(child.id) for child in active_children]
freed_size = await soft_delete_resources_by_source(
db,
source_model=SOURCE_MODEL_CHAT_TASK,
source_ids=source_ids,
deleted_at=deleted_at,
)
parent_ids = list(dict.fromkeys(str(child.parent_task_id) for child in active_children if child.parent_task_id))
for child in active_children:
child.deleted_at = deleted_at
await db.flush()
for parent_id in parent_ids:
await aggregate_main_task_status(db, parent_task_id=parent_id)
return int(freed_size or 0)
async def soft_delete_child_task(
db: AsyncSession,
*,
child_task_id: str,
user_id: str,
deleted_at: datetime | None = None,
) -> int:
deleted_at = deleted_at or datetime.now(timezone.utc)
detail_result = await db.execute(
select(
ChatGenerationTask.parent_task_id,
ChatGenerationTask.generation_index,
).where(
ChatGenerationTask.id == child_task_id,
ChatGenerationTask.user_id == user_id,
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
).limit(1)
)
detail = detail_result.one_or_none()
if not detail:
raise HTTPException(status_code=404, detail="子任务不存在")
log_operation_event(
domain="generation_ai_batch",
event_type="CHILD_RESOURCE_DELETE_START",
event_status="started",
source="service",
user_id=user_id,
group_id=detail.parent_task_id,
task_id=child_task_id,
detail={"generation_index": detail.generation_index},
)
freed_size = await soft_delete_child_tasks_batch(
db,
child_task_ids=[child_task_id],
user_id=user_id,
deleted_at=deleted_at,
)
log_operation_event(
domain="generation_ai_batch",
event_type="CHILD_RESOURCE_DELETE_SUCCESS",
event_status="success",
source="service",
user_id=user_id,
group_id=detail.parent_task_id,
task_id=child_task_id,
detail={"generation_index": detail.generation_index, "freed_size_bytes": freed_size},
)
return freed_size
async def soft_delete_top_level_task_group(
db: AsyncSession,
*,
task_id: str,
user_id: str,
deleted_at: datetime | None = None,
) -> int:
deleted_at = deleted_at or datetime.now(timezone.utc)
result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.user_id == user_id,
ChatGenerationTask.generation_mode.in_(
[GenerationMode.CHATAPI_ASYNC.value, GenerationMode.CHATAPI_MAIN.value]
),
ChatGenerationTask.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
task = result.scalar_one_or_none()
if not task:
raise HTTPException(status_code=404, detail="任务不存在")
if task.generation_mode == GenerationMode.CHATAPI_ASYNC.value:
if is_task_active(task):
raise HTTPException(status_code=400, detail="当前任务正在生成中,暂不能删除")
freed_size = await soft_delete_resources_by_source(
db,
source_model=SOURCE_MODEL_CHAT_TASK,
source_ids=[task.id],
deleted_at=deleted_at,
)
task.deleted_at = deleted_at
await db.flush()
return int(freed_size or 0)
children_map = await load_children_map(db, [task.id], include_deleted=True)
children = children_map.get(task.id, [])
active_ids = [child.id for child in children if is_task_active(child)]
if active_ids:
raise HTTPException(status_code=400, detail=f"任务组仍有 {len(active_ids)} 个子任务生成中,暂不能删除")
active_children = [child for child in children if child.deleted_at is None]
child_ids = [child.id for child in active_children]
freed_size = await soft_delete_resources_by_source(
db,
source_model=SOURCE_MODEL_CHAT_TASK,
source_ids=child_ids,
deleted_at=deleted_at,
)
for child in active_children:
child.deleted_at = deleted_at
task.deleted_at = deleted_at
await db.flush()
log_operation_event(
domain="generation_ai_batch",
event_type="BATCH_GROUP_DELETE_SUCCESS",
event_status="success",
source="service",
user_id=user_id,
group_id=task.id,
task_id=task.id,
detail={
"child_task_ids": child_ids,
"freed_size_bytes": int(freed_size or 0),
"physical_files_deleted": False,
},
)
return int(freed_size or 0)
@@ -470,6 +470,7 @@ async def charge_generation_media_by_params(
source_step_id: str | None = None, source_step_id: str | None = None,
source_step_code: str | None = None, source_step_code: str | None = None,
billing_scene: str | None = None, billing_scene: str | None = None,
quantity: int = 1,
) -> BillingSummary: ) -> BillingSummary:
"""图片/视频媒体生成扣费。 """图片/视频媒体生成扣费。
@@ -478,6 +479,7 @@ async def charge_generation_media_by_params(
""" """
project_name = project_name or "AI生成任务" project_name = project_name or "AI生成任务"
gen_type = (gen_type or "").lower().strip() gen_type = (gen_type or "").lower().strip()
quantity = max(1, int(quantity or 1))
attempt_no = attempt_no or await get_next_credit_attempt_no( attempt_no = attempt_no or await get_next_credit_attempt_no(
db, db,
owner_type=owner_type, owner_type=owner_type,
@@ -508,13 +510,14 @@ 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"
amount = await calc_image_credits(db, size, engine_id=engine_id) unit_amount = await calc_image_credits(db, size, engine_id=engine_id)
amount = round(unit_amount * quantity, 2)
items.append( items.append(
await deduct_credits_locked_once( await deduct_credits_locked_once(
db, db,
user_id=user_id, user_id=user_id,
amount=amount, amount=amount,
description=f"{description_prefix}图片生成", description=f"{description_prefix}图片生成" + (f"×{quantity}" if quantity > 1 else ""),
related_id=record_id, related_id=record_id,
charge_key=CHARGE_MEDIA, charge_key=CHARGE_MEDIA,
biz_key=biz_key, biz_key=biz_key,
@@ -523,17 +526,18 @@ async def charge_generation_media_by_params(
) )
) )
elif gen_type == "video": elif gen_type == "video":
amount = await calc_video_credits( unit_amount = await calc_video_credits(
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,
) )
amount = round(unit_amount * quantity, 2)
items.append( items.append(
await deduct_credits_locked_once( await deduct_credits_locked_once(
db, db,
user_id=user_id, user_id=user_id,
amount=amount, amount=amount,
description=f"{description_prefix}视频生成", description=f"{description_prefix}视频生成" + (f"×{quantity}" if quantity > 1 else ""),
related_id=record_id, related_id=record_id,
charge_key=CHARGE_MEDIA, charge_key=CHARGE_MEDIA,
biz_key=biz_key, biz_key=biz_key,
@@ -1,6 +1,5 @@
from __future__ import annotations from __future__ import annotations
import json
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import Iterable, Sequence from typing import Iterable, Sequence
@@ -24,12 +23,13 @@ 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.user import User from app.models.user import User
from app.schemas.generation_ai import GenerationAIHistoryBatchDeleteOut from app.schemas.generation_ai import GenerationAIHistoryBatchDeleteOut
from app.services.generation.ai.task_group_service import soft_delete_child_tasks_batch
from app.services.module_generation_flow_base_service import is_active_chat_generation_task from app.services.module_generation_flow_base_service import is_active_chat_generation_task
from app.services.module_generation_log_service import log_module_event_file from app.services.module_generation_log_service import log_module_event_file
from app.services.operation_log_service import log_operation_event
# from app.services.operation_log import log_operation # from app.services.operation_log import log_operation
from app.services.resource_accounting_service import ( from app.services.resource_accounting_service import (
SOURCE_MODEL_CHAT_TASK, SOURCE_MODEL_CHAT_TASK,
SOURCE_MODEL_GENERATION_RECORD,
SOURCE_MODEL_SHOT_SEGMENT, SOURCE_MODEL_SHOT_SEGMENT,
soft_delete_generation_record_resources, soft_delete_generation_record_resources,
soft_delete_resources_by_source, soft_delete_resources_by_source,
@@ -238,14 +238,19 @@ async def _delete_chat_tasks(
.where( .where(
ChatGenerationTask.id.in_(ids), ChatGenerationTask.id.in_(ids),
ChatGenerationTask.user_id == current_user.id, ChatGenerationTask.user_id == current_user.id,
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_ASYNC.value, ChatGenerationTask.generation_mode.in_([
ChatGenerationTask.deleted_at.is_(None), GenerationMode.CHATAPI_ASYNC.value,
GenerationMode.CHATAPI_CHILD.value,
]),
) )
.with_for_update() .with_for_update()
) )
tasks = list(result.scalars().all()) tasks = list(result.scalars().all())
_raise_missing_if_any(ids=ids, found_ids=[task.id for task in tasks], message="AI 创作记录不存在或已删除") _raise_missing_if_any(ids=ids, found_ids=[task.id for task in tasks], message="AI 创作记录不存在或已删除")
already_deleted_ids = [str(task.id) for task in tasks if task.deleted_at is not None]
_raise_invalid_if_any(invalid_ids=already_deleted_ids, message="AI 创作记录不存在或已删除", status_code=404)
invalid_ids = [ invalid_ids = [
task.id task.id
for task in tasks for task in tasks
@@ -253,14 +258,49 @@ async def _delete_chat_tasks(
] ]
_raise_invalid_if_any(invalid_ids=invalid_ids, message="AI 创作记录只有生成完成后才能删除") _raise_invalid_if_any(invalid_ids=invalid_ids, message="AI 创作记录只有生成完成后才能删除")
freed_size = await soft_delete_resources_by_source( async_ids = [str(task.id) for task in tasks if task.generation_mode == GenerationMode.CHATAPI_ASYNC.value]
db, child_ids = [str(task.id) for task in tasks if task.generation_mode == GenerationMode.CHATAPI_CHILD.value]
source_model=SOURCE_MODEL_CHAT_TASK,
source_ids=[task.id for task in tasks], freed_size = 0
deleted_at=deleted_at, if async_ids:
freed_size += int(await soft_delete_resources_by_source(
db,
source_model=SOURCE_MODEL_CHAT_TASK,
source_ids=async_ids,
deleted_at=deleted_at,
) or 0)
async_id_set = set(async_ids)
for task in tasks:
if str(task.id) in async_id_set:
task.deleted_at = deleted_at
if child_ids:
freed_size += await soft_delete_child_tasks_batch(
db,
child_task_ids=child_ids,
user_id=current_user.id,
deleted_at=deleted_at,
require_completed=True,
)
parent_task_ids = list(dict.fromkeys(
str(task.parent_task_id) for task in tasks if task.parent_task_id
))
log_operation_event(
domain="generation_ai_batch",
event_type="CHILD_RESOURCE_DELETE_SUCCESS",
event_status="success",
source="service",
user_id=current_user.id,
group_id=parent_task_ids[0] if len(parent_task_ids) == 1 else None,
detail={
"batch": True,
"task_ids": [str(task.id) for task in tasks],
"parent_task_ids": parent_task_ids,
"freed_size_bytes": int(freed_size or 0),
"physical_files_deleted": False,
},
) )
for task in tasks:
task.deleted_at = deleted_at
return _build_out( return _build_out(
source=source, source=source,
@@ -14,7 +14,7 @@ from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask from app.models.chat_generation_task import ChatGenerationTask
from app.models.model_config import ModelConfig from app.models.model_config import ModelConfig
from app.models.token_usage import TokenUsage from app.models.token_usage import TokenUsage
from app.services.generation_log_service import log_provider_call from app.services.generation.log_service import log_provider_call
from app.services.provider_limit import provider_limit from app.services.provider_limit import provider_limit
from app.utils.id_gen import generate_id from app.utils.id_gen import generate_id
@@ -13,10 +13,11 @@ from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask from app.models.chat_generation_task import ChatGenerationTask
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
from app.services.image_gen import poll_image_task_status, submit_image_task from app.services.image_gen import ImageProviderError, poll_image_task_status, submit_image_task
from app.services.provider_limit import provider_limit from app.services.provider_limit import provider_limit
from app.services.video_gen import poll_task_status, submit_video_task from app.services.video_gen import poll_task_status, submit_video_task
from app.types.generation.provider import ImageProviderBatchResult
def _loads(data: str | None) -> dict: def _loads(data: str | None) -> dict:
@@ -29,8 +30,17 @@ def _loads(data: str | None) -> dict:
return {} return {}
def _try_json(value: Any) -> Any:
if not isinstance(value, str):
return value
try:
return json.loads(value)
except Exception:
return None
async def get_runtime_engine(db: AsyncSession, task: ChatGenerationTask) -> Any: async def get_runtime_engine(db: AsyncSession, task: ChatGenerationTask) -> Any:
"""Use frozen snapshot for historical params, current DB row only for secret api_key.""" """使用任务快照冻结历史参数,只从当前引擎记录读取密钥。"""
snapshot = _loads(task.engine_snapshot_json) snapshot = _loads(task.engine_snapshot_json)
if not task.engine_id: if not task.engine_id:
raise ValueError("缺少 engine_id") raise ValueError("缺少 engine_id")
@@ -51,6 +61,31 @@ async def get_runtime_engine(db: AsyncSession, task: ChatGenerationTask) -> Any:
generate_url=snapshot.get("generate_url") or getattr(engine, "generate_url", ""), generate_url=snapshot.get("generate_url") or getattr(engine, "generate_url", ""),
query_url=snapshot.get("query_url") or getattr(engine, "query_url", ""), query_url=snapshot.get("query_url") or getattr(engine, "query_url", ""),
default_size=snapshot.get("default_size") or getattr(engine, "default_size", "2K"), default_size=snapshot.get("default_size") or getattr(engine, "default_size", "2K"),
multi_generation_enabled=bool(
snapshot.get("multi_generation_enabled")
if snapshot.get("multi_generation_enabled") is not None
else getattr(engine, "multi_generation_enabled", False)
),
max_generation_count=int(
snapshot.get("max_generation_count")
or getattr(engine, "max_generation_count", 1)
or 1
),
multi_image_max_images=int(
snapshot.get("multi_image_max_images")
or getattr(engine, "multi_image_max_images", 15)
or 15
),
max_reference_image_count=int(
snapshot.get("max_reference_image_count")
if snapshot.get("max_reference_image_count") is not None
else getattr(engine, "max_reference_image_count", 14)
),
output_format=(
snapshot.get("output_format")
if snapshot.get("output_format") is not None
else getattr(engine, "output_format", "")
) or "",
) )
@@ -58,22 +93,16 @@ async def create_provider_task(db: AsyncSession, task: ChatGenerationTask) -> di
if task.gen_type == "video": if task.gen_type == "video":
return await _create_video_task(db, task) return await _create_video_task(db, task)
if task.gen_type == "image": if task.gen_type == "image":
return await _create_image_sync_task(db, task) return await create_image_sync_result(db, task)
raise ValueError(f"不支持的生成类型: {task.gen_type}") raise ValueError(f"不支持的生成类型: {task.gen_type}")
async def _create_video_task(db: AsyncSession, task: ChatGenerationTask) -> dict: async def _create_video_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
"""Create video provider task through the original Ark SDK async task API."""
engine = await get_runtime_engine(db, task) engine = await get_runtime_engine(db, task)
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( provider_task_id = await submit_video_task(None, engine, task, include_media_references=True)
db,
engine,
task,
include_media_references=True,
)
response = {"task_id": provider_task_id} response = {"task_id": provider_task_id}
await log_provider_call( await log_provider_call(
task, task,
@@ -101,65 +130,90 @@ async def _create_video_task(db: AsyncSession, task: ChatGenerationTask) -> dict
raise raise
async def _create_image_sync_task(db: AsyncSession, task: ChatGenerationTask) -> dict: async def create_image_sync_batch_result(
"""Run the original synchronous image generation SDK under Celery control. db: AsyncSession,
task: ChatGenerationTask,
The legacy image SDK returns a final remote image URL immediately. We do *,
NOT use image_generation.tasks.create here, so image generation stays aligned generation_count: int,
with the old working flow while no longer blocking the FastAPI request. ) -> ImageProviderBatchResult:
"""
engine = await get_runtime_engine(db, task) engine = await get_runtime_engine(db, task)
return await create_image_sync_batch_result_with_engine(
task,
engine,
generation_count=generation_count,
)
async def create_image_sync_batch_result_with_engine(
task: ChatGenerationTask,
engine: Any,
*,
generation_count: int,
) -> ImageProviderBatchResult:
"""执行一次同步图片请求。
generation_count > 1 时是一次组图 API 调用失败后绝不退化为多次单图调用
"""
count = max(1, int(generation_count or 1))
started = time.perf_counter() started = time.perf_counter()
api_type = "image_sync_batch_create" if count > 1 else "image_sync_create"
async with provider_limit("ark_image_sync_create", settings.ARK_IMAGE_CREATE_MAX_CONCURRENCY): async with provider_limit("ark_image_sync_create", settings.ARK_IMAGE_CREATE_MAX_CONCURRENCY):
try: try:
result = await asyncio.to_thread( result = await asyncio.to_thread(
submit_image_task, submit_image_task,
db, None,
engine, engine,
task, task,
include_media_references=True, include_media_references=True,
generation_count=count,
) )
if result.get("error"): response_data = result.get("response_data") or result
raise RuntimeError(result.get("error"))
response_data = _try_json(result.get("response_data")) or result
await log_provider_call( await log_provider_call(
task, task,
provider=engine.provider, provider=engine.provider,
api_type="image_sync_create", api_type=api_type,
model=engine.model_name, model=engine.model_name,
engine_id=task.engine_id, engine_id=task.engine_id,
status="success", status="success",
latency_ms=int((time.perf_counter() - started) * 1000), latency_ms=int((time.perf_counter() - started) * 1000),
provider_task_id=None, provider_task_id=None,
response_data=response_data, response_data=response_data,
total_tokens=int(result.get("image_tokens", 0) or 0),
) )
return { return result
"task_id": None,
"remote_result_url": result.get("image_url"),
"image_tokens": result.get("image_tokens", 0) or 0,
"response_data": response_data,
}
except Exception as exc: except Exception as exc:
error_message = exc.safe_message if isinstance(exc, ImageProviderError) else str(exc)
await log_provider_call( await log_provider_call(
task, task,
provider=engine.provider, provider=engine.provider,
api_type="image_sync_create", api_type=api_type,
model=engine.model_name, model=engine.model_name,
engine_id=task.engine_id, engine_id=task.engine_id,
status="failed", status="failed",
latency_ms=int((time.perf_counter() - started) * 1000), latency_ms=int((time.perf_counter() - started) * 1000),
error_message=str(exc), error_message=error_message,
response_data=exc.as_dict() if isinstance(exc, ImageProviderError) else None,
) )
raise raise
def _try_json(text: Any) -> Any: async def create_image_sync_result(db: AsyncSession, task: ChatGenerationTask) -> dict:
if not isinstance(text, str): result = await create_image_sync_batch_result(db, task, generation_count=1)
return text items = result.get("items") or []
try: if len(items) != 1:
return json.loads(text) raise RuntimeError(f"图片供应商单图返回数量异常,期望 1,实际 {len(items)}")
except Exception: item = items[0]
return None if item.get("error_message"):
raise RuntimeError(item.get("error_message") or "图片生成失败")
image_url = item.get("remote_result_url")
if not image_url:
raise RuntimeError("图片供应商未返回有效图片地址")
return {
"task_id": None,
"remote_result_url": image_url,
"image_tokens": int(result.get("image_tokens", 0) or 0),
"response_data": result.get("response_data") or {},
}
async def poll_provider_task(db: AsyncSession, task: ChatGenerationTask) -> dict: async def poll_provider_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
@@ -15,6 +15,7 @@ from app.enums.generation_task import (
ChatGenerationPipelineStage, ChatGenerationPipelineStage,
ChatGenerationTaskEventType, ChatGenerationTaskEventType,
ChatGenerationTaskStatus, ChatGenerationTaskStatus,
GenerationMode,
GenerationType, GenerationType,
) )
from app.models.chat_generation_task import ChatGenerationTask from app.models.chat_generation_task import ChatGenerationTask
@@ -25,10 +26,10 @@ from app.services.celery_download_recovery_service import (
postpone_download_active_check, postpone_download_active_check,
remove_download_active, remove_download_active,
) )
from app.services.generation_log_service import log_task_event from app.services.generation.log_service import log_task_event
from app.services.generation_module_hook_service import notify_chat_generation_task_finished from app.services.generation.module_hook_service import notify_chat_generation_task_finished
from app.services.generation_poll_schedule_service import ensure_video_poll_fields, is_poll_not_due, is_video_generation_task from app.services.generation.poll_schedule_service import ensure_video_poll_fields, is_poll_not_due, is_video_generation_task
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.redis_registry_service import ( from app.services.redis_registry_service import (
redis_get_due_registry_ids, redis_get_due_registry_ids,
redis_get_registry_payloads, redis_get_registry_payloads,
@@ -341,6 +342,8 @@ async def _mark_timeout(
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value, pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
) )
await notify_chat_generation_task_finished(db, task) await notify_chat_generation_task_finished(db, task)
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
await aggregate_parent_for_child(db, task)
await db.commit() await db.commit()
await _remove_poll_active(task.id) await _remove_poll_active(task.id)
await log_task_event( await log_task_event(
@@ -367,6 +370,8 @@ async def _mark_failed(
pipeline_stage=ChatGenerationPipelineStage.FAILED.value, pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
) )
await notify_chat_generation_task_finished(db, task) await notify_chat_generation_task_finished(db, task)
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
await aggregate_parent_for_child(db, task)
await db.commit() await db.commit()
await _remove_poll_active(task.id) await _remove_poll_active(task.id)
await log_task_event(task, event_type=event_type, message=task.error_message, detail=detail) await log_task_event(task, event_type=event_type, message=task.error_message, detail=detail)
@@ -590,6 +595,99 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
checked_ids: set[str] = set() checked_ids: set[str] = set()
results: dict[str, int] = {} results: dict[str, int] = {}
# 图片多份主任务只补投递,不在恢复服务内直接调用供应商。
# 有效 claim 未过期时必须跳过,防止与正在运行的 Worker 重复调用组图 API。
from app.tasks.generation_create_tasks import chatapi_create_generation_task
image_main_cursor: str | None = None
image_main_batch_size = max(1, int(settings.GENERATION_RECOVERY_BATCH_SIZE or 100))
while True:
image_main_query = select(ChatGenerationTask).where(
ChatGenerationTask.deleted_at.is_(None),
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
ChatGenerationTask.gen_type == GenerationType.IMAGE.value,
ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value,
ChatGenerationTask.pipeline_stage.in_([
ChatGenerationPipelineStage.QUEUED.value,
ChatGenerationPipelineStage.PREPARING.value,
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
]),
)
if image_main_cursor:
image_main_query = image_main_query.where(ChatGenerationTask.id > image_main_cursor)
image_main_result = await db.execute(
image_main_query.order_by(ChatGenerationTask.id.asc())
.limit(image_main_batch_size)
.with_for_update()
)
image_mains = list(image_main_result.scalars().all())
if not image_mains:
break
for main in image_mains:
main_id = str(main.id)
image_main_cursor = main_id
checked_ids.add(main_id)
child_result = await db.execute(
select(ChatGenerationTask.id)
.where(
ChatGenerationTask.parent_task_id == main_id,
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
)
.limit(1)
)
if child_result.scalar_one_or_none() is not None:
main.provider_create_claim_token = None
main.provider_create_lease_until = None
await db.commit()
results["image_main_already_split"] = results.get("image_main_already_split", 0) + 1
continue
now = _now()
lease_until = ensure_aware_utc(main.provider_create_lease_until)
lease_alive = bool(main.provider_create_claim_token and lease_until and lease_until > now)
if lease_alive:
await db.commit()
results["image_main_claim_alive"] = results.get("image_main_claim_alive", 0) + 1
continue
if _is_expired(main.deadline_at, now):
main.provider_create_claim_token = None
main.provider_create_lease_until = None
await mark_chat_generation_task_failed_and_refund_once(
db,
task=main,
error_message="图片批量生成任务超时",
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
)
await db.commit()
results["image_main_timeout"] = results.get("image_main_timeout", 0) + 1
continue
if main.provider_create_claim_token or main.provider_create_lease_until:
main.provider_create_claim_token = None
main.provider_create_lease_until = None
main.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
await log_task_event(
main,
event_type=ChatGenerationTaskEventType.IMAGE_MAIN_CLAIM_EXPIRED.value,
message="图片主任务供应商执行租约已过期,恢复重新投递",
)
await db.commit()
try:
chatapi_create_generation_task.apply_async(
args=[main_id],
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
countdown=0,
)
results["recover_image_main_create"] = results.get("recover_image_main_create", 0) + 1
except Exception as exc:
logger.exception("恢复投递图片主任务失败 task_id=%s: %s", main_id, exc)
results["recover_image_main_enqueue_failed"] = results.get("recover_image_main_enqueue_failed", 0) + 1
if len(image_mains) < image_main_batch_size:
break
due_poll_ids = await redis_get_due_registry_ids( due_poll_ids = await redis_get_due_registry_ids(
zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY, zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY,
limit=int(settings.POLL_RECOVERY_BATCH_SIZE or settings.GENERATION_RECOVERY_BATCH_SIZE or 100), limit=int(settings.POLL_RECOVERY_BATCH_SIZE or settings.GENERATION_RECOVERY_BATCH_SIZE or 100),
@@ -673,6 +771,31 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
if len(tasks) < batch_size or progressed_this_round <= 0: if len(tasks) < batch_size or progressed_this_round <= 0:
break break
# 子任务可能在 worker 中断前已进入终态但主任务尚未汇总,按稳定游标完整重算全部主任务。
from app.services.generation.ai.task_group_service import aggregate_main_task_status
reconciled = 0
main_cursor: str | None = None
while True:
main_query = select(ChatGenerationTask.id).where(
ChatGenerationTask.deleted_at.is_(None),
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
)
if main_cursor:
main_query = main_query.where(ChatGenerationTask.id > main_cursor)
main_result = await db.execute(main_query.order_by(ChatGenerationTask.id.asc()).limit(batch_size))
parent_ids = list(main_result.scalars().all())
if not parent_ids:
break
for parent_task_id in parent_ids:
main_cursor = str(parent_task_id)
await aggregate_main_task_status(db, parent_task_id=str(parent_task_id))
await db.commit()
reconciled += 1
if len(parent_ids) < batch_size:
break
if reconciled:
results["reconcile_main"] = reconciled
return { return {
"checked": len(checked_ids), "checked": len(checked_ids),
"db_checked": total_db_checked, "db_checked": total_db_checked,
@@ -1,7 +1,5 @@
from __future__ import annotations from __future__ import annotations
from datetime import datetime, timezone
from typing import Iterable
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
@@ -11,7 +9,7 @@ from app.models.credit_record import CreditRecord
from app.models.generation_record import GenerationRecord from app.models.generation_record import GenerationRecord
from app.services.credits import refund_credits from app.services.credits import refund_credits
from app.services.credit_record_meta_service import build_refund_meta_from_charge from app.services.credit_record_meta_service import build_refund_meta_from_charge
from app.services.generation_billing_service import ( from app.services.generation.billing_service import (
CHARGE_MEDIA, CHARGE_MEDIA,
OWNER_CHAT_GENERATION_TASK, OWNER_CHAT_GENERATION_TASK,
OWNER_GENERATION_RECORD, OWNER_GENERATION_RECORD,
@@ -182,7 +180,7 @@ async def mark_chat_generation_task_failed_and_refund_once(
select(ChatGenerationTask) select(ChatGenerationTask)
.where( .where(
ChatGenerationTask.id == task_id, ChatGenerationTask.id == task_id,
ChatGenerationTask.generation_mode.in_(["chatapi_async", "hot_opening_replicate", "shot_replicate"]), ChatGenerationTask.generation_mode.in_(["chatapi_async", "chatapi_main", "chatapi_child", "hot_opening_replicate", "shot_replicate"]),
ChatGenerationTask.deleted_at.is_(None), ChatGenerationTask.deleted_at.is_(None),
) )
.with_for_update() .with_for_update()
@@ -11,21 +11,21 @@ from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask from app.models.chat_generation_task import ChatGenerationTask
from app.models.user import User from app.models.user import User
from app.schemas.generation_ai import GenerationAIReference, GenerationAITaskCreate from app.schemas.generation_ai import GenerationAIReference, GenerationAITaskCreate
from app.services.generation_ai_service import ( from app.services.generation.ai.engine_service import (
IMAGE_DEFAULT_PROPORTION, IMAGE_DEFAULT_PROPORTION,
IMAGE_DEFAULT_PX, IMAGE_DEFAULT_PX,
IMAGE_DEFAULT_SIZE, IMAGE_DEFAULT_SIZE,
VIDEO_DEFAULT_RATIO, VIDEO_DEFAULT_RATIO,
VIDEO_DEFAULT_RESOLUTION, VIDEO_DEFAULT_RESOLUTION,
_build_image_snapshot, build_image_snapshot as _build_image_snapshot,
_build_video_snapshot, build_video_snapshot as _build_video_snapshot,
_get_image_engine, get_image_engine as _get_image_engine,
_get_video_engine, get_video_engine as _get_video_engine,
_image_supported_sizes, image_supported_sizes as _image_supported_sizes,
_parse_list,
normalize_px, normalize_px,
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.resource_capacity_service import assert_user_resource_capacity_available from app.services.resource_capacity_service import assert_user_resource_capacity_available
from app.services.private_portrait.reference_resolver import resolve_private_portrait_references from app.services.private_portrait.reference_resolver import resolve_private_portrait_references
from app.utils.id_gen import generate_id from app.utils.id_gen import generate_id
@@ -128,6 +128,7 @@ async def create_chat_generation_task_for_module(
billing_scene=billing_scene, billing_scene=billing_scene,
) )
snapshot = _build_image_snapshot(engine, size, proportion, px) snapshot = _build_image_snapshot(engine, size, proportion, px)
snapshot["generation_count"] = 1
task = ChatGenerationTask( task = ChatGenerationTask(
id=task_id, id=task_id,
user_id=current_user.id, user_id=current_user.id,
@@ -182,6 +183,7 @@ async def create_chat_generation_task_for_module(
billing_scene=billing_scene, billing_scene=billing_scene,
) )
snapshot = _build_video_snapshot(engine, ratio, selected_resolution, selected_duration) snapshot = _build_video_snapshot(engine, ratio, selected_resolution, selected_duration)
snapshot["generation_count"] = 1
task = ChatGenerationTask( task = ChatGenerationTask(
id=task_id, id=task_id,
user_id=current_user.id, user_id=current_user.id,
@@ -1,14 +1,12 @@
from __future__ import annotations from __future__ import annotations
import json import json
from copy import deepcopy from datetime import datetime
from datetime import datetime, timezone
from typing import Any from typing import Any
from fastapi import HTTPException from fastapi import HTTPException
from sqlalchemy import String, cast, func, or_, select from sqlalchemy import String, cast, func, or_, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm.attributes import flag_modified
from app.config import settings from app.config import settings
from app.enums.common import ModuleEventTypeEnum, ModuleProjectStatusEnum, ModulePromptTypeEnum, ModuleStepStatusEnum from app.enums.common import ModuleEventTypeEnum, ModuleProjectStatusEnum, ModulePromptTypeEnum, ModuleStepStatusEnum
@@ -35,20 +33,19 @@ from app.schemas.hot_opening_replicate import (
HotOpeningVideoGenerationOut, HotOpeningVideoGenerationOut,
HotOpeningVideoPromptSchemaUpdateRequest, HotOpeningVideoPromptSchemaUpdateRequest,
) )
from app.services.generation_ai_service import ( from app.services.generation.ai.engine_service import (
VIDEO_DEFAULT_DURATION, VIDEO_DEFAULT_DURATION,
VIDEO_DEFAULT_RATIO, VIDEO_DEFAULT_RATIO,
VIDEO_DEFAULT_RESOLUTION, VIDEO_DEFAULT_RESOLUTION,
_get_video_engine, get_video_engine,
_parse_list, parse_json_list,
) )
from app.services.generation_billing_service import charge_module_prompt_usage from app.services.generation.billing_service import charge_module_prompt_usage
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.generation_task_factory_service import create_chat_generation_task_for_module from app.services.generation.task_factory_service import create_chat_generation_task_for_module
from app.services.hot_opening_video_prompt_service import build_final_video_prompt, optimize_hot_opening_video_prompt, patch_video_prompt_schema_from_client from app.services.hot_opening_video_prompt_service import build_final_video_prompt, optimize_hot_opening_video_prompt, patch_video_prompt_schema_from_client
from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event
from app.services.llm import optimize_prompt from app.services.llm import optimize_prompt
from app.services.resource_accounting_service import soft_delete_chat_task_resources
from app.services.module_generation_flow_base_service import ( from app.services.module_generation_flow_base_service import (
chat_tasks_by_id as _base_chat_tasks_by_id, chat_tasks_by_id as _base_chat_tasks_by_id,
create_module_step as _base_create_step, create_module_step as _base_create_step,
@@ -1021,10 +1018,10 @@ async def generate_image_from_prompt(
async def _resolve_video_prompt_config(db: AsyncSession, req: HotOpeningGenerateVideoPromptRequest) -> dict[str, Any]: async def _resolve_video_prompt_config(db: AsyncSession, req: HotOpeningGenerateVideoPromptRequest) -> dict[str, Any]:
engine = await _get_video_engine(db, req.engine_id) engine = await get_video_engine(db, req.engine_id)
supported_ratios = _parse_list(engine.supported_ratios, []) supported_ratios = parse_json_list(engine.supported_ratios, [])
supported_resolutions = _parse_list(engine.supported_resolutions, []) supported_resolutions = parse_json_list(engine.supported_resolutions, [])
supported_durations = _parse_list(engine.supported_durations, []) supported_durations = parse_json_list(engine.supported_durations, [])
default_ratio = getattr(settings, "HOT_OPENING_DEFAULT_VIDEO_RATIO", None) or VIDEO_DEFAULT_RATIO default_ratio = getattr(settings, "HOT_OPENING_DEFAULT_VIDEO_RATIO", None) or VIDEO_DEFAULT_RATIO
default_resolution = getattr(settings, "HOT_OPENING_DEFAULT_VIDEO_RESOLUTION", None) or VIDEO_DEFAULT_RESOLUTION default_resolution = getattr(settings, "HOT_OPENING_DEFAULT_VIDEO_RESOLUTION", None) or VIDEO_DEFAULT_RESOLUTION
+305 -90
View File
@@ -1,9 +1,8 @@
import base64
import json import json
import logging import logging
import mimetypes
import os import os
from datetime import datetime from datetime import datetime
from typing import Any
import httpx import httpx
from sqlalchemy import select from sqlalchemy import select
@@ -11,10 +10,16 @@ from sqlalchemy.ext.asyncio import AsyncSession
from volcenginesdkarkruntime import AsyncArk from volcenginesdkarkruntime import AsyncArk
from app.config import settings from app.config import settings
from app.enums.generation_provider import (
MULTI_IMAGE_PROMPT_TEMPLATE,
ImageProviderErrorType,
)
from app.enums.private_portrait import PRIVATE_PORTRAIT_ASSET_URI_PREFIX from app.enums.private_portrait import PRIVATE_PORTRAIT_ASSET_URI_PREFIX
from app.models.image_engine import ImageEngine from app.models.image_engine import ImageEngine
from app.services.log_config import is_enabled, LOG_DIR, LOG_DATE_FORMAT, encrypt_data from app.services.log_config import LOG_DATE_FORMAT, LOG_DIR, encrypt_data, is_enabled
from app.services.generation_provider_types import ( from app.types.generation.provider import (
ImageProviderBatchResult,
ImageProviderItem,
ProviderGenerationRecordLike, ProviderGenerationRecordLike,
ProviderImageEngineLike, ProviderImageEngineLike,
) )
@@ -22,8 +27,39 @@ from app.services.generation_provider_types import (
logger = logging.getLogger("videogen") logger = logging.getLogger("videogen")
class ImageProviderError(RuntimeError):
"""可被生成任务状态机安全收敛的图片供应商异常。"""
def __init__(
self,
message: str,
*,
error_type: ImageProviderErrorType = ImageProviderErrorType.UNKNOWN,
error_code: str | None = None,
retryable: bool = False,
http_status: int | None = None,
provider_request_id: str | None = None,
):
super().__init__(message)
self.safe_message = message
self.error_type = error_type
self.error_code = error_code
self.retryable = retryable
self.http_status = http_status
self.provider_request_id = provider_request_id
def as_dict(self) -> dict[str, Any]:
return {
"error_type": self.error_type.value,
"error_code": self.error_code,
"message": self.safe_message,
"retryable": self.retryable,
"http_status": self.http_status,
"provider_request_id": self.provider_request_id,
}
def _log_image_request(engine: ProviderImageEngineLike, record_id: str, request_data: dict): def _log_image_request(engine: ProviderImageEngineLike, record_id: str, request_data: dict):
"""Log image generation request to log/AiModel/YYYY-MM-DD.log"""
if not is_enabled(): if not is_enabled():
return return
try: try:
@@ -41,14 +77,13 @@ def _log_image_request(engine: ProviderImageEngineLike, record_id: str, request_
"request": request_encrypted, "request": request_encrypted,
"request_length": len(request_str), "request_length": len(request_str),
} }
with open(log_file, "a", encoding="utf-8") as f: with open(log_file, "a", encoding="utf-8") as file:
f.write(json.dumps(entry, ensure_ascii=False) + "\n") file.write(json.dumps(entry, ensure_ascii=False) + "\n")
except Exception: except Exception:
pass pass
def _log_image_response(record_id: str, response_data: dict, error: str | None = None): def _log_image_response(record_id: str, response_data: dict, error: str | None = None):
"""Log image generation response to log/AiModel/YYYY-MM-DD.log"""
if not is_enabled(): if not is_enabled():
return return
try: try:
@@ -63,17 +98,13 @@ def _log_image_response(record_id: str, response_data: dict, error: str | None =
"response": response_encrypted, "response": response_encrypted,
"error": error, "error": error,
} }
with open(log_file, "a", encoding="utf-8") as f: with open(log_file, "a", encoding="utf-8") as file:
f.write(json.dumps(entry, ensure_ascii=False) + "\n") file.write(json.dumps(entry, ensure_ascii=False) + "\n")
except Exception: except Exception:
pass pass
async def get_active_image_engine(db: AsyncSession) -> ImageEngine: async def get_active_image_engine(db: AsyncSession) -> ImageEngine:
"""Get the active image engine with highest priority."""
result = await db.execute( result = await db.execute(
select(ImageEngine) select(ImageEngine)
.where(ImageEngine.is_active == True) .where(ImageEngine.is_active == True)
@@ -103,105 +134,291 @@ def _resolve_url(url: str) -> str:
return f"{settings.BASE_URL.rstrip('/')}/{url.lstrip('/')}" return f"{settings.BASE_URL.rstrip('/')}/{url.lstrip('/')}"
def _value(obj: Any, name: str, default: Any = None) -> Any:
if obj is None:
return default
if isinstance(obj, dict):
return obj.get(name, default)
return getattr(obj, name, default)
def _jsonable(value: Any) -> Any:
if value is None or isinstance(value, (str, int, float, bool)):
return value
if isinstance(value, dict):
return {str(key): _jsonable(item) for key, item in value.items()}
if isinstance(value, (list, tuple)):
return [_jsonable(item) for item in value]
if hasattr(value, "model_dump"):
try:
return _jsonable(value.model_dump())
except Exception:
pass
if hasattr(value, "to_dict"):
try:
return _jsonable(value.to_dict())
except Exception:
pass
result: dict[str, Any] = {}
for key in ("url", "b64_json", "size", "output_format", "error", "code", "message"):
item = getattr(value, key, None)
if item is not None:
result[key] = _jsonable(item)
return result or str(value)
def _safe_text(value: Any, *, limit: int = 1000) -> str:
text = str(value or "").strip()
return text[:limit]
def _classify_provider_exception(exc: Exception) -> ImageProviderError:
if isinstance(exc, ImageProviderError):
return exc
if isinstance(exc, (httpx.TimeoutException, TimeoutError)):
return ImageProviderError(
"图片生成请求超时,请稍后重试",
error_type=ImageProviderErrorType.TIMEOUT,
retryable=True,
)
status_code = getattr(exc, "status_code", None)
request_id = getattr(exc, "request_id", None) or getattr(exc, "x_request_id", None)
code = getattr(exc, "code", None)
raw_message = _safe_text(getattr(exc, "message", None) or exc)
lowered = raw_message.lower()
if status_code == 429 or "rate limit" in lowered or "限流" in raw_message:
error_type = ImageProviderErrorType.RATE_LIMIT
retryable = True
message = "图片生成请求过于频繁,请稍后重试"
elif status_code in {401, 403} or "api key" in lowered or "unauthorized" in lowered:
error_type = ImageProviderErrorType.AUTH
retryable = False
message = "图片引擎鉴权失败,请联系管理员检查配置"
elif status_code and int(status_code) >= 500:
error_type = ImageProviderErrorType.PROVIDER_INTERNAL
retryable = True
message = "图片供应商服务异常,请稍后重试"
elif "sequential_image_generation" in lowered or "not support" in lowered or "unsupported" in lowered:
error_type = ImageProviderErrorType.CAPABILITY_MISMATCH
retryable = False
message = "图片引擎组图能力配置与供应商实际能力不匹配,请联系管理员"
elif "content" in lowered and ("risk" in lowered or "moderation" in lowered or "policy" in lowered):
error_type = ImageProviderErrorType.CONTENT_REJECTED
retryable = False
message = "图片内容未通过供应商审核,请调整提示词后重试"
elif status_code and 400 <= int(status_code) < 500:
error_type = ImageProviderErrorType.INVALID_REQUEST
retryable = False
message = "图片生成参数不被供应商支持,请联系管理员检查引擎配置"
elif isinstance(exc, httpx.HTTPError):
error_type = ImageProviderErrorType.NETWORK
retryable = True
message = "图片供应商网络连接异常,请稍后重试"
else:
error_type = ImageProviderErrorType.UNKNOWN
retryable = False
message = raw_message or "图片生成失败"
return ImageProviderError(
message,
error_type=error_type,
error_code=_safe_text(code, limit=128) or None,
retryable=retryable,
http_status=int(status_code) if status_code is not None else None,
provider_request_id=_safe_text(request_id, limit=128) or None,
)
def build_multi_image_provider_prompt(prompt: str, generation_count: int) -> str:
base_prompt = (prompt or "").strip()
if generation_count <= 1:
return base_prompt
suffix = MULTI_IMAGE_PROMPT_TEMPLATE.format(count=generation_count)
return f"{base_prompt}\n\n{suffix}" if base_prompt else suffix
def submit_image_task( def submit_image_task(
db, db,
engine: ProviderImageEngineLike, engine: ProviderImageEngineLike,
record: ProviderGenerationRecordLike, record: ProviderGenerationRecordLike,
*, *,
include_media_references: bool, include_media_references: bool,
) -> dict: generation_count: int = 1,
"""Submit an image generation task via Ark SDK. Returns image_url.""" ) -> ImageProviderBatchResult:
from volcenginesdkarkruntime import Ark """通过 Ark 同步图片接口生成单图或单次组图。
client = Ark(
base_url=engine.api_base,
api_key=engine.api_key,
timeout=300,
)
prompt = record.optimized_prompt or record.original_prompt generation_count > 1 时只执行一次 sequential_auto 请求;任何失败都直接抛出,
image_urls = [] 绝不退化为多次单图请求。
"""
from volcenginesdkarkruntime import Ark
count = max(1, int(generation_count or 1))
multi_generation_enabled = bool(getattr(engine, "multi_generation_enabled", False))
max_generation_count = max(1, min(5, int(getattr(engine, "max_generation_count", 1) or 1)))
if count > 1 and not multi_generation_enabled:
raise ImageProviderError(
"当前图片引擎未开启多份生成",
error_type=ImageProviderErrorType.CAPABILITY_MISMATCH,
)
if count > max_generation_count:
raise ImageProviderError(
f"当前图片引擎最多允许生成 {max_generation_count}",
error_type=ImageProviderErrorType.CAPABILITY_MISMATCH,
)
client = Ark(base_url=engine.api_base, api_key=engine.api_key, timeout=300)
original_prompt = record.optimized_prompt or record.original_prompt
provider_prompt = build_multi_image_provider_prompt(original_prompt, count)
image_urls: list[str] = []
if include_media_references and record.media_references: if include_media_references and record.media_references:
try: try:
refs = json.loads(record.media_references) refs = json.loads(record.media_references)
for ref in refs: for ref in refs if isinstance(refs, list) else []:
ref_type = ref.get("type") if (ref.get("type") or "").lower() == "image" and ref.get("url"):
ref_url = ref.get("url", "") image_urls.append(_resolve_url(ref["url"]))
if ref_type == "image" and ref_url:
resolved = _resolve_url(ref_url)
image_urls.append(resolved)
except (json.JSONDecodeError, TypeError): except (json.JSONDecodeError, TypeError):
pass image_urls = []
request_payload = { request_log_payload: dict[str, Any] = {
"model": engine.model_name, "model": engine.model_name,
"prompt": prompt, "prompt": provider_prompt,
"size": record.image_size or engine.default_size, "size": record.image_size or engine.default_size,
"sequential_image_generation": "disabled",
"output_format": "png",
"response_format": "url", "response_format": "url",
"watermark": False, "watermark": False,
"include_media_references": include_media_references,
} }
request_sdk_payload: dict[str, Any] = dict(request_log_payload)
if image_urls: if image_urls:
request_payload["image"] = image_urls request_log_payload["image"] = image_urls
request_sdk_payload["image"] = image_urls
output_format = (getattr(engine, "output_format", "") or "").lower().strip()
if output_format:
request_log_payload["output_format"] = output_format
request_sdk_payload["output_format"] = output_format
if count > 1:
try:
from volcenginesdkarkruntime.types.images import SequentialImageGenerationOptions
except Exception:
try:
from volcenginesdkarkruntime.types.images.image_generate_params import (
SequentialImageGenerationOptions,
)
except Exception as import_exc:
raise ImageProviderError(
"当前图片引擎运行依赖缺少组图参数对象,请升级火山 Ark SDK 后重试",
error_type=ImageProviderErrorType.CAPABILITY_MISMATCH,
) from import_exc
_log_image_request(engine, record.id, request_payload) request_log_payload["sequential_image_generation"] = "auto"
request_log_payload["sequential_image_generation_options"] = {"max_images": count}
request_log_payload["stream"] = False
request_sdk_payload["sequential_image_generation"] = "auto"
request_sdk_payload["sequential_image_generation_options"] = SequentialImageGenerationOptions(
max_images=count,
)
request_sdk_payload["stream"] = False
_log_image_request(engine, record.id, request_log_payload)
try: try:
result = client.images.generate( result = client.images.generate(**request_sdk_payload)
model=engine.model_name, top_error = _value(result, "error")
prompt=prompt, if top_error:
size=record.image_size or engine.default_size, error_code = _value(top_error, "code")
output_format="png", error_message = _value(top_error, "message") or str(top_error)
response_format="url", raise ImageProviderError(
watermark=False, _safe_text(error_message) or "图片供应商返回失败",
image=image_urls if image_urls else None, error_type=ImageProviderErrorType.INVALID_REQUEST,
) error_code=_safe_text(error_code, limit=128) or None,
image_url = result.data[0].url )
response_data = { raw_data = _value(result, "data", []) or []
"model": result.model, if not isinstance(raw_data, (list, tuple)):
"created": result.created, raise ImageProviderError(
"data": [{"url": item.url, "size": item.size} for item in result.data] if result.data else [], "图片供应商返回 data 结构异常",
"usage": { error_type=ImageProviderErrorType.INVALID_RESPONSE,
"generated_images": result.usage.generated_images if hasattr(result.usage, 'generated_images') else 0, )
"output_tokens": result.usage.output_tokens if hasattr(result.usage, 'output_tokens') else 0,
"total_tokens": result.usage.total_tokens if hasattr(result.usage, 'total_tokens') else 0, items: list[ImageProviderItem] = []
response_items: list[dict[str, Any]] = []
for index, raw_item in enumerate(raw_data, start=1):
item_error = _value(raw_item, "error")
if item_error:
error_code = _safe_text(_value(item_error, "code"), limit=128)
error_message = _safe_text(_value(item_error, "message") or item_error)
items.append({
"generation_index": index,
"error_code": error_code,
"error_message": error_message or "单张图片生成失败",
"response_data": _jsonable(raw_item),
})
response_items.append(_jsonable(raw_item))
continue
url = _safe_text(_value(raw_item, "url"), limit=4000)
b64_json = _safe_text(_value(raw_item, "b64_json"), limit=100) if not url else ""
item: ImageProviderItem = {
"generation_index": index,
"remote_result_url": url,
"size": _safe_text(_value(raw_item, "size"), limit=64),
"output_format": _safe_text(_value(raw_item, "output_format"), limit=32),
"response_data": _jsonable(raw_item),
} }
if b64_json:
item["b64_json"] = b64_json
items.append(item)
response_items.append(_jsonable(raw_item))
usage = _value(result, "usage")
generated_images = int(_value(usage, "generated_images", 0) or 0)
total_tokens = int(_value(usage, "total_tokens", 0) or 0)
response_data = {
"model": _value(result, "model", engine.model_name),
"created": _value(result, "created"),
"data": response_items,
"usage": {
"generated_images": generated_images,
"input_images": int(_value(usage, "input_images", 0) or 0),
"output_tokens": int(_value(usage, "output_tokens", 0) or 0),
"total_tokens": total_tokens,
},
} }
except httpx.TimeoutException: _log_image_response(record.id, response_data)
error_msg = "图片生成超时,请稍后重试" return {
logger.error(f"Image generation timeout for record {record.id}") "items": items,
_log_image_response(record.id, {}, error_msg) "model": str(response_data["model"] or ""),
raise TimeoutError(error_msg) "created": int(response_data["created"] or 0),
except Exception as e: "generated_images": generated_images,
error_msg = str(e) "image_tokens": total_tokens,
logger.error(f"Image generation failed for record {record.id}: {error_msg}") "response_data": response_data,
_log_image_response(record.id, {}, error_msg) }
raise except Exception as exc:
provider_error = _classify_provider_exception(exc)
logger.error(
"Image generation failed for record %s: type=%s code=%s message=%s",
record.id,
provider_error.error_type.value,
provider_error.error_code,
provider_error.safe_message,
)
_log_image_response(record.id, provider_error.as_dict(), provider_error.safe_message)
raise provider_error from exc
finally: finally:
client.close() try:
client.close()
return { except Exception:
"image_url": image_url, pass
"image_tokens": getattr(result.usage, "total_tokens", 0),
"response_data": json.dumps(response_data, ensure_ascii=False, default=str),
"error": str(result.error) if result.error else "",
}
async def poll_image_task_status(engine: ImageEngine, task_id: str) -> dict: async def poll_image_task_status(engine: ImageEngine, task_id: str) -> dict:
"""Query image task status via Ark SDK. Returns {status, image_url, response_data}.""" client = AsyncArk(base_url=engine.api_base, api_key=engine.api_key)
client = AsyncArk( try:
base_url=engine.api_base, result = await client.image_generation.tasks.get(task_id=task_id)
api_key=engine.api_key, finally:
) await client.close()
result = await client.image_generation.tasks.get(task_id=task_id)
await client.close()
response_dict = { response_dict = {
"id": result.id, "id": result.id,
@@ -239,13 +456,11 @@ async def poll_image_task_status(engine: ImageEngine, task_id: str) -> dict:
async def download_image(image_url: str, dest_path: str) -> str: async def download_image(image_url: str, dest_path: str) -> str:
"""Download image to local storage."""
os.makedirs(os.path.dirname(dest_path), exist_ok=True) os.makedirs(os.path.dirname(dest_path), exist_ok=True)
async with httpx.AsyncClient(timeout=300) as client: async with httpx.AsyncClient(timeout=300) as client:
async with client.stream("GET", image_url) as response: async with client.stream("GET", image_url) as response:
response.raise_for_status() response.raise_for_status()
with open(dest_path, "wb") as f: with open(dest_path, "wb") as file:
async for chunk in response.aiter_bytes(chunk_size=8192): async for chunk in response.aiter_bytes(chunk_size=8192):
f.write(chunk) file.write(chunk)
return dest_path return dest_path
@@ -4,7 +4,7 @@ from collections.abc import Iterable
from datetime import datetime from datetime import datetime
from typing import Any, TypedDict from typing import Any, TypedDict
from sqlalchemy import and_, func, or_, select from sqlalchemy import and_, case, func, or_, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.generation_task import GenerationType from app.enums.generation_task import GenerationType
@@ -12,7 +12,7 @@ from app.enums.recent_generation import (
RECENT_GENERATION_ALL_MODULES, RECENT_GENERATION_ALL_MODULES,
RECENT_GENERATION_CHAT_TASK_MODULES, RECENT_GENERATION_CHAT_TASK_MODULES,
RECENT_GENERATION_COMPLETED_STATUS, RECENT_GENERATION_COMPLETED_STATUS,
RECENT_GENERATION_MODULE_TO_TASK_MODE, RECENT_GENERATION_MODULE_TO_TASK_MODES,
RECENT_GENERATION_TASK_MODE_VALUE_TO_MODULE, RECENT_GENERATION_TASK_MODE_VALUE_TO_MODULE,
RecentGenerationModuleEnum, RecentGenerationModuleEnum,
RecentGenerationResourceTypeEnum, RecentGenerationResourceTypeEnum,
@@ -175,11 +175,12 @@ async def _list_chat_task_recent_rows(
modules: list[RecentGenerationModuleEnum], modules: list[RecentGenerationModuleEnum],
limit: int, limit: int,
) -> list[dict[str, Any]]: ) -> list[dict[str, Any]]:
task_mode_values = [ task_mode_values = list(dict.fromkeys(
RECENT_GENERATION_MODULE_TO_TASK_MODE[module].value task_mode.value
for module in modules for module in modules
if module in RECENT_GENERATION_CHAT_TASK_MODULES if module in RECENT_GENERATION_CHAT_TASK_MODULES
] for task_mode in RECENT_GENERATION_MODULE_TO_TASK_MODES[module]
))
if not task_mode_values: if not task_mode_values:
return [] return []
@@ -189,10 +190,19 @@ async def _list_chat_task_recent_rows(
ChatGenerationTask.created_at, ChatGenerationTask.created_at,
) )
module_partition_expr = case(
(
ChatGenerationTask.generation_mode.in_(["chatapi_async", "chatapi_child"]),
RecentGenerationModuleEnum.CHAT_AI.value,
),
else_=ChatGenerationTask.generation_mode,
)
ranked_subquery = ( ranked_subquery = (
select( select(
ChatGenerationTask.id.label("generation_id"), ChatGenerationTask.id.label("generation_id"),
ChatGenerationTask.generation_mode.label("generation_mode"), ChatGenerationTask.generation_mode.label("generation_mode"),
module_partition_expr.label("module_key"),
ChatGenerationTask.gen_type.label("gen_type"), ChatGenerationTask.gen_type.label("gen_type"),
ChatGenerationTask.image_url.label("image_url"), ChatGenerationTask.image_url.label("image_url"),
ChatGenerationTask.video_url.label("video_url"), ChatGenerationTask.video_url.label("video_url"),
@@ -200,7 +210,7 @@ async def _list_chat_task_recent_rows(
generated_time_expr.label("generated_time"), generated_time_expr.label("generated_time"),
func.row_number() func.row_number()
.over( .over(
partition_by=ChatGenerationTask.generation_mode, partition_by=module_partition_expr,
order_by=(generated_time_expr.desc(), ChatGenerationTask.created_at.desc()), order_by=(generated_time_expr.desc(), ChatGenerationTask.created_at.desc()),
) )
.label("row_num"), .label("row_num"),
@@ -218,7 +228,7 @@ async def _list_chat_task_recent_rows(
stmt = ( stmt = (
select(ranked_subquery) select(ranked_subquery)
.where(ranked_subquery.c.row_num <= limit) .where(ranked_subquery.c.row_num <= limit)
.order_by(ranked_subquery.c.generation_mode.asc(), ranked_subquery.c.generated_time.desc()) .order_by(ranked_subquery.c.module_key.asc(), ranked_subquery.c.generated_time.desc())
) )
return [dict(row) for row in (await db.execute(stmt)).mappings().all()] return [dict(row) for row in (await db.execute(stmt)).mappings().all()]
@@ -1,14 +1,12 @@
from __future__ import annotations from __future__ import annotations
import json import json
from copy import deepcopy from datetime import datetime
from datetime import datetime, timezone
from typing import Any from typing import Any
from fastapi import HTTPException from fastapi import HTTPException
from sqlalchemy import func, select from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm.attributes import flag_modified
from app.config import settings from app.config import settings
from app.enums.common import ModuleEventTypeEnum, ModuleProjectStatusEnum, ModulePromptTypeEnum, ModuleStepStatusEnum from app.enums.common import ModuleEventTypeEnum, ModuleProjectStatusEnum, ModulePromptTypeEnum, ModuleStepStatusEnum
@@ -35,16 +33,16 @@ from app.schemas.shot_replicate import (
ShotReplicateVideoGenerationOut, ShotReplicateVideoGenerationOut,
ShotReplicateVideoPromptSchemaUpdateRequest, ShotReplicateVideoPromptSchemaUpdateRequest,
) )
from app.services.generation_ai_service import ( from app.services.generation.ai.engine_service import (
VIDEO_DEFAULT_DURATION, VIDEO_DEFAULT_DURATION,
VIDEO_DEFAULT_RATIO, VIDEO_DEFAULT_RATIO,
VIDEO_DEFAULT_RESOLUTION, VIDEO_DEFAULT_RESOLUTION,
_get_video_engine, get_video_engine,
_parse_list, parse_json_list,
) )
from app.services.generation_billing_service import charge_module_prompt_usage from app.services.generation.billing_service import charge_module_prompt_usage
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.generation_task_factory_service import create_chat_generation_task_for_module from app.services.generation.task_factory_service import create_chat_generation_task_for_module
from app.services.hot_opening_video_prompt_service import ( from app.services.hot_opening_video_prompt_service import (
build_final_video_prompt, build_final_video_prompt,
optimize_hot_opening_video_prompt as optimize_shot_replicate_video_prompt, optimize_hot_opening_video_prompt as optimize_shot_replicate_video_prompt,
@@ -52,7 +50,6 @@ from app.services.hot_opening_video_prompt_service import (
) )
from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event
from app.services.llm import optimize_prompt from app.services.llm import optimize_prompt
from app.services.resource_accounting_service import soft_delete_chat_task_resources
from app.services.module_generation_flow_base_service import ( from app.services.module_generation_flow_base_service import (
assert_project_has_no_active_chat_tasks as _base_assert_project_has_no_active_chat_tasks, assert_project_has_no_active_chat_tasks as _base_assert_project_has_no_active_chat_tasks,
chat_tasks_by_id as _base_chat_tasks_by_id, chat_tasks_by_id as _base_chat_tasks_by_id,
@@ -976,10 +973,10 @@ async def generate_image_from_prompt(
async def _resolve_video_prompt_config(db: AsyncSession, req: ShotReplicateGenerateVideoPromptRequest) -> dict[str, Any]: async def _resolve_video_prompt_config(db: AsyncSession, req: ShotReplicateGenerateVideoPromptRequest) -> dict[str, Any]:
engine = await _get_video_engine(db, req.engine_id) engine = await get_video_engine(db, req.engine_id)
supported_ratios = _parse_list(engine.supported_ratios, []) supported_ratios = parse_json_list(engine.supported_ratios, [])
supported_resolutions = _parse_list(engine.supported_resolutions, []) supported_resolutions = parse_json_list(engine.supported_resolutions, [])
supported_durations = _parse_list(engine.supported_durations, []) supported_durations = parse_json_list(engine.supported_durations, [])
default_ratio = getattr(settings, "SHOT_REPLICATE_DEFAULT_VIDEO_RATIO", None) or VIDEO_DEFAULT_RATIO default_ratio = getattr(settings, "SHOT_REPLICATE_DEFAULT_VIDEO_RATIO", None) or VIDEO_DEFAULT_RATIO
default_resolution = getattr(settings, "SHOT_REPLICATE_DEFAULT_VIDEO_RESOLUTION", None) or VIDEO_DEFAULT_RESOLUTION default_resolution = getattr(settings, "SHOT_REPLICATE_DEFAULT_VIDEO_RESOLUTION", None) or VIDEO_DEFAULT_RESOLUTION
+1 -1
View File
@@ -14,7 +14,7 @@ from app.config import settings
from app.enums.private_portrait import PRIVATE_PORTRAIT_ASSET_URI_PREFIX from app.enums.private_portrait import PRIVATE_PORTRAIT_ASSET_URI_PREFIX
from app.models.video_engine import VideoEngine from app.models.video_engine import VideoEngine
from app.services.log_config import is_enabled, LOG_DIR, LOG_DATE_FORMAT, encrypt_data from app.services.log_config import is_enabled, LOG_DIR, LOG_DATE_FORMAT, encrypt_data
from app.services.generation_provider_types import ( from app.types.generation.provider import (
ProviderGenerationRecordLike, ProviderGenerationRecordLike,
ProviderVideoEngineLike, ProviderVideoEngineLike,
) )
+1 -1
View File
@@ -17,7 +17,7 @@ from app.services.resource_accounting_service import (
) )
from app.services.video_cover_service import create_video_cover_for_local_video from app.services.video_cover_service import create_video_cover_for_local_video
from app.config import settings from app.config import settings
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
logger = logging.getLogger("videogen") logger = logging.getLogger("videogen")
@@ -18,10 +18,10 @@ from app.enums.generation_task import (
from app.models.base import async_session from app.models.base import async_session
from app.models.chat_generation_task import ChatGenerationTask from app.models.chat_generation_task import ChatGenerationTask
from app.services.error_codes import extract_error_message from app.services.error_codes import extract_error_message
from app.services.generation_log_service import log_task_event from app.services.generation.log_service import log_task_event
from app.services.generation_poll_schedule_service import ensure_video_poll_fields from app.services.generation.poll_schedule_service import ensure_video_poll_fields
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.generation_provider_service import create_provider_task from app.services.generation.provider_service import create_provider_task
from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
from app.services.redis_registry_service import ensure_aware_utc from app.services.redis_registry_service import ensure_aware_utc
from app.tasks.celery_app import celery_app from app.tasks.celery_app import celery_app
@@ -138,14 +138,22 @@ async def _run(task_id: str):
).with_for_update().limit(1)) ).with_for_update().limit(1))
task = result.scalar_one_or_none() task = result.scalar_one_or_none()
if not task or task.generation_mode not in ALLOWED_GENERATION_MODES: is_image_main = bool(
task
and task.generation_mode == GenerationMode.CHATAPI_MAIN.value
and task.gen_type == GenerationType.IMAGE.value
and int(task.generation_count or 1) > 1
)
if not task or (task.generation_mode not in ALLOWED_GENERATION_MODES and not is_image_main):
return return
if task.status != ChatGenerationTaskStatus.GENERATING.value: if task.status != ChatGenerationTaskStatus.GENERATING.value:
return return
deadline_at = ensure_aware_utc(task.deadline_at) deadline_at = ensure_aware_utc(task.deadline_at)
if deadline_at and datetime.now(timezone.utc) > deadline_at: # 图片 main 的 deadline 与 provider claim 由 image_batch_service 原子处理,
# 避免重复 Celery 消息在有效租约期间把正在执行的批次错误退款。
if not is_image_main and deadline_at and datetime.now(timezone.utc) > deadline_at:
await mark_chat_generation_task_failed_and_refund_once( await mark_chat_generation_task_failed_and_refund_once(
db, db,
task=task, task=task,
@@ -159,8 +167,10 @@ async def _run(task_id: str):
to_status=ChatGenerationTaskStatus.FAILED.value, to_status=ChatGenerationTaskStatus.FAILED.value,
to_stage=ChatGenerationPipelineStage.TIMEOUT.value, to_stage=ChatGenerationPipelineStage.TIMEOUT.value,
) )
from app.services.generation_module_hook_service import notify_chat_generation_task_finished from app.services.generation.module_hook_service import notify_chat_generation_task_finished
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
await notify_chat_generation_task_finished(db, task) await notify_chat_generation_task_finished(db, task)
await aggregate_parent_for_child(db, task)
await db.commit() await db.commit()
return return
@@ -202,6 +212,12 @@ async def _run(task_id: str):
}, },
) )
if is_image_main:
from app.services.generation.ai.image_batch_service import run_image_main_batch
await run_image_main_batch(db, task)
return
if task.seedance_task_id or task.provider_task_id: if task.seedance_task_id or task.provider_task_id:
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
if task.gen_type == GenerationType.VIDEO.value: if task.gen_type == GenerationType.VIDEO.value:
@@ -302,16 +318,41 @@ async def _run(task_id: str):
if task: if task:
error_message = extract_error_message(exc, "生成任务") if callable(extract_error_message) else str(exc) error_message = extract_error_message(exc, "生成任务") if callable(extract_error_message) else str(exc)
await mark_chat_generation_task_failed_and_refund_once( if is_image_main:
db, # image_batch_service 负责供应商/拆分失败退款。若 child 已落库,
task=task, # 顶层兜底绝不能再把 main 退款。
error_message=error_message, child_result = await db.execute(
pipeline_stage=ChatGenerationPipelineStage.FAILED.value, select(ChatGenerationTask.id).where(
) ChatGenerationTask.parent_task_id == task.id,
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
).limit(1)
)
has_children = child_result.scalar_one_or_none() is not None
if not has_children:
task.provider_create_claim_token = None
task.provider_create_lease_until = None
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=error_message,
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
)
else:
from app.services.generation.ai.task_group_service import aggregate_main_task_status
await aggregate_main_task_status(db, parent_task_id=str(task.id))
else:
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=error_message,
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
)
await db.commit() await db.commit()
await log_task_event(task, event_type=ChatGenerationTaskEventType.TASK_FAILED.value, message=task.error_message) await log_task_event(task, event_type=ChatGenerationTaskEventType.TASK_FAILED.value, message=error_message)
from app.services.generation_module_hook_service import notify_chat_generation_task_finished from app.services.generation.module_hook_service import notify_chat_generation_task_finished
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
await notify_chat_generation_task_finished(db, task) await notify_chat_generation_task_finished(db, task)
await aggregate_parent_for_child(db, task)
await db.commit() await db.commit()
@@ -17,6 +17,7 @@ from app.enums.generation_task import (
ChatGenerationPipelineStage, ChatGenerationPipelineStage,
ChatGenerationTaskEventType, ChatGenerationTaskEventType,
ChatGenerationTaskStatus, ChatGenerationTaskStatus,
GenerationMode,
GenerationType, GenerationType,
) )
from app.models.base import async_session from app.models.base import async_session
@@ -28,9 +29,9 @@ from app.services.celery_download_recovery_service import (
upsert_download_active, upsert_download_active,
) )
from app.services.error_codes import extract_error_message from app.services.error_codes import extract_error_message
from app.services.generation_download_service import download_generation_result from app.services.generation.download_service import download_generation_result
from app.services.generation_log_service import log_task_event from app.services.generation.log_service import log_task_event
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
from app.services.resource_accounting_service import record_chat_task_generated_resource from app.services.resource_accounting_service import record_chat_task_generated_resource
from app.tasks.celery_app import celery_app from app.tasks.celery_app import celery_app
@@ -275,9 +276,18 @@ async def enqueue_download_task(
await db.commit() await db.commit()
check_at = _queue_timeout_at(now) check_at = _queue_timeout_at(now)
await _register_active_from_task(task, check_at=check_at, priority=priority, reason=reason) try:
await _register_active_from_task(task, check_at=check_at, priority=priority, reason=reason)
except Exception as exc:
# Redis active 注册表只用于恢复,不应阻止真实 Celery 投递。
await _log_download_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE_FAILED,
message=f"下载恢复注册表写入失败: {exc}",
detail={"reason": reason, "celery_task_id": celery_task_id},
)
await _apply_download_async( applied = await _apply_download_async(
task, task,
priority=priority, priority=priority,
countdown=countdown, countdown=countdown,
@@ -285,6 +295,31 @@ async def enqueue_download_task(
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE if recover else ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE, event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE if recover else ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE,
failed_event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE_FAILED if recover else ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE_FAILED, failed_event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE_FAILED if recover else ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE_FAILED,
) )
if not applied:
# apply_async 失败不能伪装成已投递。保留远程结果,进入下载恢复等待。
try:
await remove_download_active(task.id)
except Exception:
pass
refreshed = await _reload_task(db, task.id)
if refreshed and refreshed.status == ChatGenerationTaskStatus.GENERATING.value:
retry_at = now + timedelta(seconds=int(settings.DOWNLOAD_TASK_RETRY_BACKOFF_SECONDS or 30))
refreshed.pipeline_stage = DOWNLOAD_STAGE_RETRY_WAITING
refreshed.download_next_retry_at = retry_at
refreshed.download_last_error = "Celery 下载任务投递失败,等待恢复重试"
refreshed.download_lease_until = None
await db.commit()
try:
await _register_active_from_task(
refreshed,
check_at=retry_at,
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
reason="enqueue_failed_wait_recovery",
)
except Exception:
pass
return None
if old_stage != DOWNLOAD_STAGE_QUEUED: if old_stage != DOWNLOAD_STAGE_QUEUED:
# 独立记录阶段变化的上下文,便于和真正投递事件对照。 # 独立记录阶段变化的上下文,便于和真正投递事件对照。
await _log_download_event( await _log_download_event(
@@ -466,19 +501,32 @@ async def _mark_download_failed(
non_retryable: bool = False, non_retryable: bool = False,
) -> None: ) -> None:
error_message = extract_error_message(exc, "下载") if callable(extract_error_message) else str(exc) error_message = extract_error_message(exc, "下载") if callable(extract_error_message) else str(exc)
await mark_chat_generation_task_failed_and_refund_once( is_image_child = (
db, task.gen_type == GenerationType.IMAGE.value
task=task, and task.generation_mode == GenerationMode.CHATAPI_CHILD.value
error_message=error_message,
pipeline_stage=DOWNLOAD_STAGE_FAILED,
) )
if is_image_child:
# 图片生成费用属于 main;child 下载失败只记录下载终态,不退图片生成积分。
task.status = ChatGenerationTaskStatus.FAILED.value
task.pipeline_stage = DOWNLOAD_STAGE_FAILED
task.error_message = error_message
await db.flush()
else:
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=error_message,
pipeline_stage=DOWNLOAD_STAGE_FAILED,
)
task.download_last_error = error_message task.download_last_error = error_message
task.download_lease_until = None task.download_lease_until = None
task.download_next_retry_at = None task.download_next_retry_at = None
from app.services.generation_module_hook_service import notify_chat_generation_task_finished from app.services.generation.module_hook_service import notify_chat_generation_task_finished
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
await notify_chat_generation_task_finished(db, task) await notify_chat_generation_task_finished(db, task)
await aggregate_parent_for_child(db, task)
await db.commit() await db.commit()
await remove_download_active(task.id) await remove_download_active(task.id)
@@ -549,9 +597,11 @@ async def _run(task_id: str):
) )
await sync_chat_generation_task_media_token_snapshot(db, task) await sync_chat_generation_task_media_token_snapshot(db, task)
from app.services.generation_module_hook_service import notify_chat_generation_task_finished from app.services.generation.module_hook_service import notify_chat_generation_task_finished
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
await notify_chat_generation_task_finished(db, task) await notify_chat_generation_task_finished(db, task)
await aggregate_parent_for_child(db, task)
await db.commit() await db.commit()
await remove_download_active(task.id) await remove_download_active(task.id)
@@ -19,8 +19,8 @@ from app.enums.generation_task import (
from app.models.base import async_session from app.models.base import async_session
from app.models.chat_generation_task import ChatGenerationTask from app.models.chat_generation_task import ChatGenerationTask
from app.services.error_codes import extract_error_message from app.services.error_codes import extract_error_message
from app.services.generation_log_service import log_task_event, log_provider_call from app.services.generation.log_service import log_task_event, log_provider_call
from app.services.generation_poll_schedule_service import ( from app.services.generation.poll_schedule_service import (
build_default_poll_schedule, build_default_poll_schedule,
build_video_pending_poll_schedule, build_video_pending_poll_schedule,
ensure_video_poll_fields, ensure_video_poll_fields,
@@ -28,8 +28,8 @@ from app.services.generation_poll_schedule_service import (
is_poll_not_due, is_poll_not_due,
is_video_generation_task, is_video_generation_task,
) )
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.generation_provider_service import poll_provider_task from app.services.generation.provider_service import poll_provider_task
from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
from app.services.redis_registry_service import ( from app.services.redis_registry_service import (
datetime_to_epoch, datetime_to_epoch,
@@ -144,9 +144,11 @@ async def remove_poll_active(task_id: str) -> None:
async def _notify_finished(db, task: ChatGenerationTask) -> None: async def _notify_finished(db, task: ChatGenerationTask) -> None:
from app.services.generation_module_hook_service import notify_chat_generation_task_finished from app.services.generation.module_hook_service import notify_chat_generation_task_finished
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
await notify_chat_generation_task_finished(db, task) await notify_chat_generation_task_finished(db, task)
await aggregate_parent_for_child(db, task)
async def _reload_task(db, task_id: str) -> ChatGenerationTask | None: async def _reload_task(db, task_id: str) -> ChatGenerationTask | None:
@@ -18,21 +18,21 @@ RecoveryRunner = Callable[[], Awaitable[Dict[str, Any]]]
async def _run_download_once() -> Dict[str, Any]: async def _run_download_once() -> Dict[str, Any]:
from app.services.generation_recovery_service import recover_download_tasks_once from app.services.generation.recovery_service import recover_download_tasks_once
async with async_session() as db: async with async_session() as db:
return await recover_download_tasks_once(db) return await recover_download_tasks_once(db)
async def _run_generation_once() -> Dict[str, Any]: async def _run_generation_once() -> Dict[str, Any]:
from app.services.generation_recovery_service import recover_generation_tasks_once from app.services.generation.recovery_service import recover_generation_tasks_once
async with async_session() as db: async with async_session() as db:
return await recover_generation_tasks_once(db) return await recover_generation_tasks_once(db)
async def _run_due_poll_dispatch_once() -> Dict[str, Any]: async def _run_due_poll_dispatch_once() -> Dict[str, Any]:
from app.services.generation_recovery_service import dispatch_due_poll_tasks_once from app.services.generation.recovery_service import dispatch_due_poll_tasks_once
async with async_session() as db: async with async_session() as db:
return await dispatch_due_poll_tasks_once(db) return await dispatch_due_poll_tasks_once(db)
@@ -38,7 +38,7 @@ from app.services.module_async_recovery_service import (
) )
from app.services.shot_replicate_taskset_service import refresh_task_set_split_summary from app.services.shot_replicate_taskset_service import refresh_task_set_split_summary
from app.services.shot_video_analysis_service import analyze_video_for_shot_split from app.services.shot_video_analysis_service import analyze_video_for_shot_split
from app.services.generation_billing_service import charge_shot_video_analysis_usage from app.services.generation.billing_service import charge_shot_video_analysis_usage
from app.services.shot_video_split_service import split_video_segment_async from app.services.shot_video_split_service import split_video_segment_async
from app.services.upload_video_asset_service import validate_split_range from app.services.upload_video_asset_service import validate_split_range
from app.services.upload_resource import record_shot_segment_upload_resource from app.services.upload_resource import record_shot_segment_upload_resource
+1
View File
@@ -0,0 +1 @@
"""应用级类型契约。"""
@@ -0,0 +1 @@
"""生成领域类型契约。"""
@@ -1,14 +1,10 @@
from __future__ import annotations from __future__ import annotations
from typing import Protocol from typing import Any, Protocol, TypedDict
class ProviderGenerationRecordLike(Protocol): class ProviderGenerationRecordLike(Protocol):
"""图片/视频供应商提交接口需要的任务字段协议。 """图片/视频供应商提交接口需要的任务字段协议。"""
GenerationRecord ChatGenerationTask 都具备这些字段但二者不是同一个 ORM 模型
使用 Protocol 可以避免把 submit_image_task / submit_video_task 错误限制为某一个具体模型
"""
id: str id: str
original_prompt: str original_prompt: str
@@ -21,16 +17,25 @@ class ProviderGenerationRecordLike(Protocol):
image_size: str | None image_size: str | None
image_proportion: str | None image_proportion: str | None
image_px: str | None image_px: str | None
generation_count: int
engine_id: str | None
class ProviderImageEngineLike(Protocol): class ProviderImageEngineLike(Protocol):
"""图片生成提交接口需要的引擎字段协议。""" """图片生成提交接口需要的引擎字段协议。"""
id: str
name: str name: str
provider: str
api_base: str api_base: str
api_key: str api_key: str
model_name: str model_name: str
default_size: str | None default_size: str | None
multi_generation_enabled: bool
max_generation_count: int
multi_image_max_images: int
max_reference_image_count: int
output_format: str
class ProviderVideoEngineLike(Protocol): class ProviderVideoEngineLike(Protocol):
@@ -40,3 +45,23 @@ class ProviderVideoEngineLike(Protocol):
api_base: str api_base: str
api_key: str api_key: str
model_name: str model_name: str
class ImageProviderItem(TypedDict, total=False):
generation_index: int
remote_result_url: str
b64_json: str
size: str
output_format: str
response_data: dict[str, Any]
error_code: str
error_message: str
class ImageProviderBatchResult(TypedDict, total=False):
items: list[ImageProviderItem]
model: str
created: int
generated_images: int
image_tokens: int
response_data: dict[str, Any]
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-nrhXZrQV.js"></script> <script type="module" crossorigin src="/assets/index-DN27BTjS.js"></script>
<link rel="stylesheet" crossorigin href="/assets/index-DviWdElm.css"> <link rel="stylesheet" crossorigin href="/assets/index-DviWdElm.css">
</head> </head>
<body> <body>
<div id="root"></div> <div id="root"></div>
</body> </body>
</html> </html>
@@ -0,0 +1,119 @@
import React from 'react';
import { LoadingOutlined, PlayCircleFilled, WarningOutlined } from '@ant-design/icons';
export interface GenerationTaskResourceItem {
id?: string;
genType?: string;
status?: string;
displayStatus?: string | null;
pipelineStage?: string | null;
imageUrl?: string | null;
videoUrl?: string | null;
videoCoverUrl?: string | null;
errorMessage?: string | null;
generationIndex?: number | null;
}
export interface GenerationTaskResourceGroup extends GenerationTaskResourceItem {
generationCount?: number | null;
childItems?: GenerationTaskResourceItem[] | null;
}
interface Props {
task: GenerationTaskResourceGroup;
onPreview: (url: string, type: 'image' | 'video') => void;
resolveUrl?: (url?: string | null) => string;
}
const spanByCount = (count: number, index: number): number => {
if (count <= 1) return 6;
if (count === 2 || count === 4) return 3;
if (count === 3) return index < 2 ? 3 : 6;
return index < 3 ? 2 : 3;
};
const statusText = (item: GenerationTaskResourceItem): string => {
const status = item.displayStatus || item.pipelineStage || item.status || 'generating';
const labels: Record<string, string> = {
pending: '待处理', queued: '排队中', preparing: '准备中', generating: '生成中',
creating_provider_task: '创建任务中', waiting_remote: '等待生成', polling: '轮询中',
result_ready: '结果就绪', download_queued: '等待下载', downloading: '下载中',
retry_waiting: '等待重试', completed: '已完成', failed: '生成失败',
download_failed: '下载失败', deleted: '已删除',
};
return labels[status] || status;
};
const isPending = (item: GenerationTaskResourceItem): boolean => {
const status = item.displayStatus || item.pipelineStage || item.status;
return !status || ['pending', 'queued', 'preparing', 'generating', 'creating_provider_task', 'waiting_remote', 'polling', 'result_ready', 'download_queued', 'downloading', 'retry_waiting'].includes(status);
};
const GenerationTaskResourceGrid: React.FC<Props> = ({ task, onPreview, resolveUrl = (url) => url || '' }) => {
const count = Math.max(1, Math.min(5, Number(task.generationCount || task.childItems?.length || 1)));
const children = [...(task.childItems || [])].sort((a, b) => Number(a.generationIndex || 0) - Number(b.generationIndex || 0));
const items: GenerationTaskResourceItem[] = children.length > 0
? children
: (count > 1 ? Array.from({ length: count }, (_, index) => ({
id: `${task.id || 'task'}-placeholder-${index + 1}`,
genType: task.genType,
status: task.status,
displayStatus: task.displayStatus,
pipelineStage: task.pipelineStage,
generationIndex: index + 1,
errorMessage: task.errorMessage,
})) : [task]);
return (
<div style={{ width: '100%', height: '100%', display: 'grid', gridTemplateColumns: 'repeat(6, minmax(0, 1fr))', gridAutoRows: 'minmax(0, 1fr)', gap: count > 1 ? 4 : 0 }}>
{items.slice(0, 5).map((item, index) => {
const displayStatus = item.displayStatus || item.pipelineStage || item.status || 'generating';
const imageUrl = resolveUrl(item.imageUrl);
const videoUrl = resolveUrl(item.videoUrl);
const coverUrl = resolveUrl(item.videoCoverUrl);
const isVideo = (item.genType || task.genType) === 'video';
const hasResource = isVideo ? !!videoUrl : !!imageUrl;
return (
<div
key={item.id || `${index}`}
style={{
gridColumn: `span ${spanByCount(items.length, index)}`,
minWidth: 0,
minHeight: 0,
position: 'relative',
overflow: 'hidden',
borderRadius: items.length === 1 ? 12 : 8,
background: 'linear-gradient(135deg, #ffffff 0%, #FAFBFC 100%)',
border: items.length === 1 ? 'none' : '1px solid #E7EAF0',
}}
>
{hasResource && displayStatus !== 'deleted' ? (
<button
type="button"
onClick={() => onPreview(isVideo ? videoUrl : imageUrl, isVideo ? 'video' : 'image')}
style={{ width: '100%', height: '100%', border: 0, padding: 0, background: 'transparent', cursor: 'pointer', position: 'relative' }}
>
{isVideo ? (
coverUrl ? <img src={coverUrl} alt={`生成结果${item.generationIndex || index + 1}`} style={{ width: '100%', height: '100%', objectFit: 'contain' }} />
: <video src={videoUrl} muted preload="metadata" style={{ width: '100%', height: '100%', objectFit: 'contain' }} />
) : (
<img src={imageUrl} alt={`生成结果${item.generationIndex || index + 1}`} style={{ width: '100%', height: '100%', objectFit: 'contain' }} />
)}
{isVideo ? <PlayCircleFilled style={{ position: 'absolute', left: '50%', top: '50%', transform: 'translate(-50%, -50%)', fontSize: items.length > 2 ? 28 : 46, color: 'rgba(255,255,255,.92)', filter: 'drop-shadow(0 4px 10px rgba(0,0,0,.28))' }} /> : null}
</button>
) : (
<div style={{ width: '100%', height: '100%', minHeight: 0, display: 'flex', flexDirection: 'column', alignItems: 'center', justifyContent: 'center', gap: 7, padding: 8, textAlign: 'center' }}>
{isPending(item) ? <LoadingOutlined spin style={{ color: '#8b5cf6', fontSize: items.length > 2 ? 20 : 34 }} /> : <WarningOutlined style={{ color: displayStatus === 'deleted' ? '#98A2B3' : '#A45B5B', fontSize: items.length > 2 ? 20 : 34 }} />}
<span style={{ fontSize: items.length > 2 ? 10 : 12, color: isPending(item) ? '#8b5cf6' : (displayStatus === 'deleted' ? '#98A2B3' : '#A45B5B'), fontWeight: 500 }}>{statusText(item)}</span>
{!isPending(item) && item.errorMessage && items.length <= 2 ? <span style={{ fontSize: 10, color: '#A45B5B', lineHeight: 1.3, maxHeight: 28, overflow: 'hidden' }}>{item.errorMessage}</span> : null}
</div>
)}
{items.length > 1 ? <span style={{ position: 'absolute', top: 5, left: 5, zIndex: 2, padding: '1px 6px', borderRadius: 10, background: 'rgba(17,24,39,.58)', color: '#fff', fontSize: 10 }}>#{item.generationIndex || index + 1}</span> : null}
</div>
);
})}
</div>
);
};
export default GenerationTaskResourceGrid;
+91 -53
View File
@@ -24,6 +24,7 @@ import bg3 from '../assets/bg3.png';
import text from '../assets/testb.png'; import text from '../assets/testb.png';
import UploadSelector from '../components/UploadSelector'; import UploadSelector from '../components/UploadSelector';
import GenerationTaskResourceGrid from '../components/generation/GenerationTaskResourceGrid';
@@ -68,6 +69,17 @@ const { Option } = Select;
// 解构Typography组件 // 解构Typography组件
const { Text } = Typography; const { Text } = Typography;
const GENERATION_RESOURCE_BASE = (import.meta.env.VITE_API_BASE || 'http://localhost:8000')
.replace(/\/api\/?$/i, '')
.replace(/\/$/, '');
const resolveGenerationResourceUrl = (url?: string | null): string => {
if (!url) return '';
const value = String(url).trim();
if (!value) return '';
if (/^(https?:)?\/\//i.test(value) || /^(blob|data):/i.test(value)) return value;
return `${GENERATION_RESOURCE_BASE}${value.startsWith('/') ? value : `/${value}`}`;
};
interface MediaReference { interface MediaReference {
name: string; name: string;
@@ -97,6 +109,7 @@ interface Message {
resolution?: string; resolution?: string;
timestamp?: string; timestamp?: string;
engine_id: string; engine_id: string;
generation_count: number;
} }
@@ -138,6 +151,7 @@ const AIChatPage: React.FC = () => {
const { const {
mediaType, mediaType,
countType, countType,
generationCount,
selectedRatio, selectedRatio,
selectedResolution, selectedResolution,
width, width,
@@ -150,6 +164,7 @@ const AIChatPage: React.FC = () => {
inputValue, inputValue,
setMediaType, setMediaType,
setCountType, setCountType,
setGenerationCount,
setSelectedRatio, setSelectedRatio,
setSelectedResolution, setSelectedResolution,
setWidth, setWidth,
@@ -170,6 +185,24 @@ const AIChatPage: React.FC = () => {
const currentEngine = currentEngineList?.find((e: any) => e.id === countType); const currentEngine = currentEngineList?.find((e: any) => e.id === countType);
const maxImageCount = currentEngine?.maxImageCount ?? 4; const maxImageCount = currentEngine?.maxImageCount ?? 4;
const maxVideoCount = currentEngine?.maxVideoCount ?? 1; const maxVideoCount = currentEngine?.maxVideoCount ?? 1;
const multiGenerationEnabled = Boolean(currentEngine?.multiGenerationEnabled);
const configuredMaxGenerationCount = multiGenerationEnabled
? Math.max(1, Math.min(5, Number(currentEngine?.maxGenerationCount || 1)))
: 1;
const referenceImageCount = mediaType === 'image'
? currentMedia.filter((item) => item.type === 'image').length
: 0;
const imageProviderRemainingCount = mediaType === 'image'
? Math.max(1, Number(currentEngine?.multiImageMaxImages || 15) - referenceImageCount)
: 5;
const effectiveMaxGenerationCount = Math.max(
1,
Math.min(
5,
configuredMaxGenerationCount,
mediaType === 'image' ? imageProviderRemainingCount : 5,
),
);
const [uploading, setUploading] = useState<boolean>(false); const [uploading, setUploading] = useState<boolean>(false);
@@ -444,7 +477,7 @@ const AIChatPage: React.FC = () => {
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;
} }
return Number(total.toFixed(2)); return Number((total * generationCount).toFixed(2));
} else { } else {
// 图片:baseCredits × ratio // 图片:baseCredits × ratio
let total = config.baseCredits * config.ratio; let total = config.baseCredits * config.ratio;
@@ -456,7 +489,7 @@ const AIChatPage: React.FC = () => {
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;
} }
return Number(total.toFixed(2)); return Number((total * generationCount).toFixed(2));
} }
}; };
@@ -558,6 +591,22 @@ const AIChatPage: React.FC = () => {
} }
}, [mediaType, enginesele, enginesLoaded]); }, [mediaType, enginesele, enginesLoaded]);
// 切换媒体类型或引擎时,默认回到最安全的单份生成。
useEffect(() => {
if (!enginesLoaded) return;
setGenerationCount(1);
}, [mediaType, countType, enginesLoaded, setGenerationCount]);
// 图片参考图数量变化后动态收敛本次可选数量;后端仍会再次校验。
useEffect(() => {
if (generationCount > effectiveMaxGenerationCount) {
setGenerationCount(effectiveMaxGenerationCount);
if (mediaType === 'image') {
antdMessage.info(`受当前引擎或参考图数量限制,本次最多生成 ${effectiveMaxGenerationCount}`);
}
}
}, [generationCount, effectiveMaxGenerationCount, mediaType, setGenerationCount, antdMessage]);
// 点击外部关闭弹窗 // 点击外部关闭弹窗
useEffect(() => { useEffect(() => {
const handleClickOutside = (e: MouseEvent) => { const handleClickOutside = (e: MouseEvent) => {
@@ -703,8 +752,9 @@ const AIChatPage: React.FC = () => {
setCreditCalculationData(data); setCreditCalculationData(data);
}) })
getgen_list(Pagebreak).then((data: any) => { getgen_list(Pagebreak).then((data: any) => {
let mess_list = data.items // API 按创建时间倒序返回;对话区按时间正序展示,最新消息保持在底部。
let total = data.total const mess_list = data.items
const total = data.total
setGen_list(mess_list) setGen_list(mess_list)
setTotalnumber(total) setTotalnumber(total)
}) })
@@ -1005,6 +1055,7 @@ const AIChatPage: React.FC = () => {
engine_id: countType, engine_id: countType,
idempotency_key: new Date().toLocaleString('zh-CN'), idempotency_key: new Date().toLocaleString('zh-CN'),
generation_count: generationCount,
media_references: mediaReferences, media_references: mediaReferences,
// 图片参数(仅图片模式时添加) // 图片参数(仅图片模式时添加)
...(mediaType === 'image' && { ...(mediaType === 'image' && {
@@ -1034,6 +1085,7 @@ const AIChatPage: React.FC = () => {
setCurrentMedia([]); setCurrentMedia([]);
setFirstFrame(null); setFirstFrame(null);
setLastFrame(null); setLastFrame(null);
setGenerationCount(1);
// 创建任务成功后,重置页数为1,获取最新列表 // 创建任务成功后,重置页数为1,获取最新列表
const newPagebreak = { ...Pagebreak, page: 1 }; const newPagebreak = { ...Pagebreak, page: 1 };
@@ -1041,7 +1093,12 @@ const AIChatPage: React.FC = () => {
getgen_list(newPagebreak).then((data: any) => { getgen_list(newPagebreak).then((data: any) => {
// 将data.items的最后一个元素添加到gen_list末尾 // 将data.items的最后一个元素添加到gen_list末尾
setGen_list((prev: any[]) => [...prev, data.items[data.items.length - 1]]); const newestItem = Array.isArray(data.items) && data.items.length > 0
? data.items[data.items.length - 1]
: null;
if (newestItem) {
setGen_list((prev: any[]) => [...prev, newestItem]);
}
setTotalnumber(data.total); setTotalnumber(data.total);
// 发送消息后滚动到底部 // 发送消息后滚动到底部
setTimeout(() => { setTimeout(() => {
@@ -1121,13 +1178,13 @@ const AIChatPage: React.FC = () => {
setGen_list((prev: any[]) => { setGen_list((prev: any[]) => {
const existingIds = new Set(prev.map((item: any) => item.id)); const existingIds = new Set(prev.map((item: any) => item.id));
// 只添加不存在的新数据,保持新数据的原有顺序 // 只添加不存在的新数据,保持新数据的原有顺序
const newItems = data.items.filter((item: any) => { const newItems = (data.items || []).filter((item: any) => {
if (!item.id) return false; if (!item.id) return false;
if (existingIds.has(item.id)) return false; if (existingIds.has(item.id)) return false;
existingIds.add(item.id); existingIds.add(item.id);
return true; return true;
}); }).reverse();
// 新数据在前,旧数据在后 // 加载的是更早一页,按时间正序放到现有消息前面。
return [...newItems, ...prev]; return [...newItems, ...prev];
}); });
@@ -2453,50 +2510,16 @@ const AIChatPage: React.FC = () => {
display: 'flex', gap: 16, marginBottom: 16, marginTop: 16, width: '100%', alignItems: 'flex-start', display: 'flex', gap: 16, marginBottom: 16, marginTop: 16, width: '100%', alignItems: 'flex-start',
}}> }}>
<div style={{ width: '50%', overflow: 'hidden', borderRadius: 12, position: 'relative', height: 220, border: '1px solid #E7EAF0', boxShadow: '0 4px 12px rgba(139, 92, 246, 0.08)', display: 'flex', alignItems: 'center', justifyContent: 'center', background: 'linear-gradient(135deg, #ffffff 0%, #FAFBFC 100%)' }}> <div style={{ width: '50%', overflow: 'hidden', borderRadius: 12, position: 'relative', height: 220, border: '1px solid #E7EAF0', boxShadow: '0 4px 12px rgba(139, 92, 246, 0.08)', background: 'linear-gradient(135deg, #ffffff 0%, #FAFBFC 100%)' }}>
{msg.status === 'generating' ? ( <GenerationTaskResourceGrid
<> task={msg}
<div style={{ position: 'absolute', top: 12, left: 12, display: 'flex', alignItems: 'center', gap: 8, zIndex: 2 }}> resolveUrl={resolveGenerationResourceUrl}
{/* <div style={{ width: 20, height: 20, border: '2px solid #ddd6fe', borderTopColor: '#8b5cf6', borderRadius: '50%', animation: 'spin 1s linear infinite' }} /> */} onPreview={(url, type) => {
</div> setPreviewUrl(url);
<div style={{ position: 'absolute', inset: 0, background: 'linear-gradient(90deg, transparent 0%, rgba(255,255,255,0.6) 50%, transparent 100%)', animation: 'shimmer 2s infinite' }} /> setPreviewType(type);
<div style={{ position: 'absolute', top: '50%', left: '50%', transform: 'translate(-50%, -50%)', display: 'flex', flexDirection: 'column', alignItems: 'center', gap: 16 }}> setPreviewVisible(true);
<div style={{ width: 64, height: 64, borderRadius: '50%', background: 'rgba(139, 92, 246, 0.1)', display: 'flex', alignItems: 'center', justifyContent: 'center', boxShadow: '0 0 30px rgba(139, 92, 246, 0.2)' }}> }}
<div style={{ width: 48, height: 48, border: '3px solid #ddd6fe', borderTopColor: '#8b5cf6', borderRadius: '50%', animation: 'spin 1s linear infinite' }} /> />
</div>
<span style={{ fontSize: 12, color: '#8b5cf6', fontWeight: 500 }}>...</span>
{/* <div style={{ display: 'flex', gap: 4 }}>
<div style={{ width: 6, height: 6, borderRadius: '50%', background: '#8b5cf6', animation: 'pulse 1.5s ease-in-out infinite' }} />
<div style={{ width: 6, height: 6, borderRadius: '50%', background: '#a8a1b6', animation: 'pulse 1.5s ease-in-out 0.2s infinite' }} />
<div style={{ width: 6, height: 6, borderRadius: '50%', background: '#ddd6fe', animation: 'pulse 1.5s ease-in-out 0.4s infinite' }} />
</div> */}
</div>
</>
) : msg.status === 'failed' ? (
<div style={{ display: 'flex', flexDirection: 'column', alignItems: 'center', justifyContent: 'center', gap: 12 }}>
<div style={{ width: 48, height: 48, borderRadius: '50%', background: 'rgba(168, 90, 106, 0.10)', display: 'flex', alignItems: 'center', justifyContent: 'center' }}>
<WarningOutlined style={{ color: '#A45B5B', fontSize: 22 }} />
</div>
<span style={{ fontSize: 14, color: '#A45B5B', fontWeight: 620 }}>退</span>
{msg.errorMessage && (
<span style={{ fontSize: 13, color: '#A45B5B', textAlign: 'center', padding: '0 8px', lineHeight: 1.5 }}>{msg.errorMessage}</span>
)}
</div>
) : (
<>
{msg.genType === 'image' ? (
<img src={`${import.meta.env.VITE_API_BASE || "http://localhost:8000"}/static${msg.imageUrl}&w=300&p=50`} alt={msg.name} style={{ width: '100%', height: '100%', borderRadius: 12, objectFit: 'contain', cursor: 'pointer', transition: 'transform 0.3s ease' }} onClick={() => { setPreviewUrl(msg.imageUrl); setPreviewType('image'); setPreviewVisible(true); }} onMouseEnter={(e) => { e.currentTarget.style.transform = 'scale(1.05)'; }} onMouseLeave={(e) => { e.currentTarget.style.transform = 'scale(1)'; }} />
) : (
<>
<img src={`${import.meta.env.VITE_API_BASE || "http://localhost:8000"}/static${msg.videoCoverUrl}&w=300&p=50`} alt={msg.name} style={{ width: '100%', height: '100%', borderRadius: 12, objectFit: 'contain', cursor: 'pointer', transition: 'transform 0.3s ease' }} onClick={() => { setPreviewUrl(msg.videoUrl); setPreviewType('video'); setPreviewVisible(true); }} onMouseEnter={(e) => { e.currentTarget.style.transform = 'scale(1.05)'; }} onMouseLeave={(e) => { e.currentTarget.style.transform = 'scale(1)'; }} />
<div style={{ position: 'absolute', top: '50%', left: '50%', transform: 'translate(-50%, -50%)', width: 56, height: 56, background: 'rgba(47, 52, 64, 0.72)', borderRadius: '50%', display: 'flex', alignItems: 'center', justifyContent: 'center', pointerEvents: 'none', boxShadow: '0 10px 24px rgba(47, 52, 64, 0.22)' }}>
<svg width="24" height="24" viewBox="0 0 24 24" fill="#fff"><path d="M8 5v14l11-7z" /></svg>
</div>
</>
)}
</>
)}
</div> </div>
<div style={{ width: '50%', display: 'flex', flexDirection: 'column', gap: 12 }}> <div style={{ width: '50%', display: 'flex', flexDirection: 'column', gap: 12 }}>
<div style={{ background: '#FFFFFF', borderRadius: 12, padding: 12, border: '1px solid #E7EAF0', flex: 1, display: 'flex', alignItems: 'center', justifyContent: 'center' }}> <div style={{ background: '#FFFFFF', borderRadius: 12, padding: 12, border: '1px solid #E7EAF0', flex: 1, display: 'flex', alignItems: 'center', justifyContent: 'center' }}>
@@ -2506,7 +2529,7 @@ const AIChatPage: React.FC = () => {
<span style={{ <span style={{
// background: 'rgba(139, 92, 246, 0.08)', // background: 'rgba(139, 92, 246, 0.08)',
borderRadius: 16, color: '#8b5cf6', fontWeight: 500 borderRadius: 16, color: '#8b5cf6', fontWeight: 500
}}>{msg.engineSnapshot.name}</span> }}>{msg.engineSnapshot?.name || '未知引擎'}</span>
<span style={{ <span style={{
// background: 'rgba(139, 92, 246, 0.08)', // background: 'rgba(139, 92, 246, 0.08)',
borderRadius: 16, color: '#667085' borderRadius: 16, color: '#667085'
@@ -4418,6 +4441,21 @@ const AIChatPage: React.FC = () => {
</Space> </Space>
<div style={{ display: 'flex', alignItems: 'center', gap: 12, flexShrink: 0 }}> <div style={{ display: 'flex', alignItems: 'center', gap: 12, flexShrink: 0 }}>
<div style={{ display: 'flex', alignItems: 'center', gap: 6, whiteSpace: 'nowrap', height: 34, padding: '0 8px 0 12px', borderRadius: 11, background: 'rgba(255, 255, 255, 0.92)', border: '1px solid rgba(231, 234, 240, 0.92)', boxShadow: '0 4px 12px rgba(47, 52, 64, 0.04)' }}>
<Text style={{ fontSize: 13, color: '#667085', fontWeight: 600 }}></Text>
<Select
value={generationCount}
onChange={(value) => setGenerationCount(Number(value || 1))}
disabled={effectiveMaxGenerationCount <= 1}
size="small"
variant="borderless"
style={{ width: 68 }}
options={Array.from({ length: effectiveMaxGenerationCount }, (_, index) => ({
value: index + 1,
label: `${index + 1}`,
}))}
/>
</div>
<div style={{ display: 'flex', alignItems: 'center', gap: 6, color: '#667085', whiteSpace: 'nowrap', height: 34, padding: '0 12px', borderRadius: 11, background: 'rgba(255, 255, 255, 0.92)', border: '1px solid rgba(231, 234, 240, 0.92)', boxShadow: '0 4px 12px rgba(47, 52, 64, 0.04)' }}> <div style={{ display: 'flex', alignItems: 'center', gap: 6, color: '#667085', whiteSpace: 'nowrap', height: 34, padding: '0 12px', borderRadius: 11, background: 'rgba(255, 255, 255, 0.92)', border: '1px solid rgba(231, 234, 240, 0.92)', boxShadow: '0 4px 12px rgba(47, 52, 64, 0.04)' }}>
<Text style={{ fontSize: 13, color: '#667085', fontWeight: 600 }}></Text> <Text style={{ fontSize: 13, color: '#667085', fontWeight: 600 }}></Text>
<Text style={{ fontSize: 13, color: '#2f3440', fontWeight: 800 }}>{getEstimatedCredits()}</Text> <Text style={{ fontSize: 13, color: '#2f3440', fontWeight: 800 }}>{getEstimatedCredits()}</Text>
+5
View File
@@ -18,6 +18,7 @@ interface AppState {
// 生成配置状态 - 页面跳转时保留,刷新时重置 // 生成配置状态 - 页面跳转时保留,刷新时重置
mediaType: string; mediaType: string;
countType: string; countType: string;
generationCount: number;
selectedRatio: string; selectedRatio: string;
selectedResolution: string; selectedResolution: string;
width: number; width: number;
@@ -46,6 +47,7 @@ interface AppState {
// 生成配置状态更新方法 // 生成配置状态更新方法
setMediaType: (mediaType: string) => void; setMediaType: (mediaType: string) => void;
setCountType: (countType: string) => void; setCountType: (countType: string) => void;
setGenerationCount: (generationCount: number) => void;
setImageSettings: (ratio: string, resolution: string, width: number, height: number) => void; setImageSettings: (ratio: string, resolution: string, width: number, height: number) => void;
setVideoSettings: (duration: number, aspectRatio: string, resolution: string) => void; setVideoSettings: (duration: number, aspectRatio: string, resolution: string) => void;
setEngineOptions: (options: { ratios: string[]; resolutions: string[]; durations: number[] }) => void; setEngineOptions: (options: { ratios: string[]; resolutions: string[]; durations: number[] }) => void;
@@ -71,6 +73,7 @@ export const useAppStore = create<AppState>((set, get) => ({
// 生成配置状态初始值 // 生成配置状态初始值
mediaType: 'video', mediaType: 'video',
countType: '请选择', countType: '请选择',
generationCount: 1,
selectedRatio: '1:1', selectedRatio: '1:1',
selectedResolution: '2K', selectedResolution: '2K',
width: 2048, width: 2048,
@@ -160,6 +163,7 @@ export const useAppStore = create<AppState>((set, get) => ({
// 生成配置状态更新方法 // 生成配置状态更新方法
setMediaType: (mediaType) => set({ mediaType }), setMediaType: (mediaType) => set({ mediaType }),
setCountType: (countType) => set({ countType }), setCountType: (countType) => set({ countType }),
setGenerationCount: (generationCount) => set({ generationCount }),
setImageSettings: (ratio, resolution, width, height) => setImageSettings: (ratio, resolution, width, height) =>
set({ selectedRatio: ratio, selectedResolution: resolution, width, height }), set({ selectedRatio: ratio, selectedResolution: resolution, width, height }),
setVideoSettings: (duration, aspectRatio, resolution) => setVideoSettings: (duration, aspectRatio, resolution) =>
@@ -179,6 +183,7 @@ export const useAppStore = create<AppState>((set, get) => ({
resetGenerationConfig: () => set({ resetGenerationConfig: () => set({
mediaType: 'image', mediaType: 'image',
countType: '请选择', countType: '请选择',
generationCount: 1,
selectedRatio: '1:1', selectedRatio: '1:1',
selectedResolution: '2K', selectedResolution: '2K',
width: 2048, width: 2048,