暂存修改

This commit is contained in:
sjy
2026-07-23 13:07:33 +08:00
55 changed files with 4363 additions and 1920 deletions
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
-1
View File
@@ -1 +0,0 @@
@import "https://fonts.googleapis.com/css2?family=Outfit:wght@300;400;500;600;700&display=swap";:root{--font-sans:"Outfit", -apple-system, BlinkMacSystemFont, "Segoe UI", sans-serif;--nav-bg:#08080c;--nav-surface:#ffffff0a;--nav-border:#ffffff0f;--nav-text:#e8e8ec;--nav-text-muted:#8b8fa3;--nav-hover:#ffffff14;--nav-active:#ffffff1f}*,:before,:after{box-sizing:border-box}html,body{font-family:var(--font-sans);-webkit-font-smoothing:antialiased;-moz-osx-font-smoothing:grayscale;color:#1a1a2e;letter-spacing:-.01em;background:#f8f9fc;margin:0;padding:0}#root{min-height:100vh}
+2 -2
View File
@@ -28,8 +28,8 @@
}
})();
</script>
<script type="module" crossorigin src="/assets/index-CnO3FgtX.js"></script>
<link rel="stylesheet" crossorigin href="/assets/index-D7ShJUt4.css">
<script type="module" crossorigin src="/assets/index-25NqXtGU.js"></script>
<link rel="stylesheet" crossorigin href="/assets/index-D3fwIbOp.css">
</head>
<body>
<div id="root"></div>
+82
View File
@@ -13,6 +13,7 @@
"dayjs": "^1.11.21",
"react": "^19.2.5",
"react-dom": "^19.2.5",
"react-quill-new": "^3.8.3",
"react-router-dom": "^7.15.0",
"xlsx-js-style": "^1.2.0",
"zustand": "^5.0.13"
@@ -1428,6 +1429,12 @@
"node": ">=8"
}
},
"node_modules/eventemitter3": {
"version": "5.0.4",
"resolved": "https://registry.npmjs.org/eventemitter3/-/eventemitter3-5.0.4.tgz",
"integrity": "sha512-mlsTRyGaPBjPedk6Bvw+aqbsXDtoAyAzm5MO7JgU+yVRyMQ5O8bD4Kcci7BS85f93veegeCPkL8R4GLClnjLFw==",
"license": "MIT"
},
"node_modules/exit-on-epipe": {
"version": "1.0.1",
"resolved": "https://registry.npmmirror.com/exit-on-epipe/-/exit-on-epipe-1.0.1.tgz",
@@ -1437,6 +1444,12 @@
"node": ">=0.8"
}
},
"node_modules/fast-diff": {
"version": "1.3.0",
"resolved": "https://registry.npmjs.org/fast-diff/-/fast-diff-1.3.0.tgz",
"integrity": "sha512-VxPP4NqbUjj6MaAOafWeUn2cXWLcCtljklUtZf0Ind4XQ+QPtmA0b18zZy0jIQx+ExRVCR/ZQpBmik5lXshNsw==",
"license": "Apache-2.0"
},
"node_modules/fdir": {
"version": "6.5.0",
"resolved": "https://registry.npmjs.org/fdir/-/fdir-6.5.0.tgz",
@@ -1761,6 +1774,25 @@
"url": "https://opencollective.com/parcel"
}
},
"node_modules/lodash-es": {
"version": "4.18.1",
"resolved": "https://registry.npmjs.org/lodash-es/-/lodash-es-4.18.1.tgz",
"integrity": "sha512-J8xewKD/Gk22OZbhpOVSwcs60zhd95ESDwezOFuA3/099925PdHJ7OFHNTGtajL3AlZkykD32HykiMo+BIBI8A==",
"license": "MIT"
},
"node_modules/lodash.clonedeep": {
"version": "4.5.0",
"resolved": "https://registry.npmjs.org/lodash.clonedeep/-/lodash.clonedeep-4.5.0.tgz",
"integrity": "sha512-H5ZhCF25riFd9uB5UCkVKo61m3S/xZk1x4wA6yp/L3RFP6Z/eHH1ymQcGLo7J3GMPfm0V/7m1tryHuGVxpqEBQ==",
"license": "MIT"
},
"node_modules/lodash.isequal": {
"version": "4.5.0",
"resolved": "https://registry.npmjs.org/lodash.isequal/-/lodash.isequal-4.5.0.tgz",
"integrity": "sha512-pDo3lu8Jhfjqls6GkMgpahsF9kCyayhgykjyLMNFTKWrpVdAQtYyB4muAMWozBB4ig/dtWAmsMxLEI8wuz+DYQ==",
"deprecated": "This package is deprecated. Use require('node:util').isDeepStrictEqual instead.",
"license": "MIT"
},
"node_modules/nanoid": {
"version": "3.3.12",
"resolved": "https://registry.npmjs.org/nanoid/-/nanoid-3.3.12.tgz",
@@ -1780,6 +1812,12 @@
"node": "^10 || ^12 || ^13.7 || ^14 || >=15.0.1"
}
},
"node_modules/parchment": {
"version": "3.0.0",
"resolved": "https://registry.npmjs.org/parchment/-/parchment-3.0.0.tgz",
"integrity": "sha512-HUrJFQ/StvgmXRcQ1ftY6VEZUq3jA2t9ncFN4F84J/vN0/FPpQF+8FKXb3l6fLces6q0uOHj6NJn+2xvZnxO6A==",
"license": "BSD-3-Clause"
},
"node_modules/picocolors": {
"version": "1.1.1",
"resolved": "https://registry.npmjs.org/picocolors/-/picocolors-1.1.1.tgz",
@@ -1841,6 +1879,35 @@
"node": ">=0.8"
}
},
"node_modules/quill": {
"version": "2.0.3",
"resolved": "https://registry.npmjs.org/quill/-/quill-2.0.3.tgz",
"integrity": "sha512-xEYQBqfYx/sfb33VJiKnSJp8ehloavImQ2A6564GAbqG55PGw1dAWUn1MUbQB62t0azawUS2CZZhWCjO8gRvTw==",
"license": "BSD-3-Clause",
"dependencies": {
"eventemitter3": "^5.0.1",
"lodash-es": "^4.17.21",
"parchment": "^3.0.0",
"quill-delta": "^5.1.0"
},
"engines": {
"npm": ">=8.2.3"
}
},
"node_modules/quill-delta": {
"version": "5.1.0",
"resolved": "https://registry.npmjs.org/quill-delta/-/quill-delta-5.1.0.tgz",
"integrity": "sha512-X74oCeRI4/p0ucjb5Ma8adTXd9Scumz367kkMK5V/IatcX6A0vlgLgKbzXWy5nZmCGeNJm2oQX0d2Eqj+ZIlCA==",
"license": "MIT",
"dependencies": {
"fast-diff": "^1.3.0",
"lodash.clonedeep": "^4.5.0",
"lodash.isequal": "^4.5.0"
},
"engines": {
"node": ">= 12.0.0"
}
},
"node_modules/react": {
"version": "19.2.6",
"resolved": "https://registry.npmjs.org/react/-/react-19.2.6.tgz",
@@ -1868,6 +1935,21 @@
"integrity": "sha512-/LLMVyas0ljjAtoYiPqYiL8VWXzUUdThrmU5+n20DZv+a+ClRoevUzw5JxU+Ieh5/c87ytoTBV9G1FiKfNJdmg==",
"license": "MIT"
},
"node_modules/react-quill-new": {
"version": "3.8.3",
"resolved": "https://registry.npmjs.org/react-quill-new/-/react-quill-new-3.8.3.tgz",
"integrity": "sha512-c96PYqFTo0pI4R3e79B3rH9LUIce1kIQbmTBu/imJQZk8305ogyLyBqKKjG2UoInDlquXqePSzmBo2aVia3ttw==",
"license": "MIT",
"dependencies": {
"lodash-es": "^4.17.21",
"quill": "~2.0.3"
},
"peerDependencies": {
"quill-delta": "^5.1.0",
"react": "^16 || ^17 || ^18 || ^19",
"react-dom": "^16 || ^17 || ^18 || ^19"
}
},
"node_modules/react-router": {
"version": "7.15.0",
"resolved": "https://registry.npmjs.org/react-router/-/react-router-7.15.0.tgz",
+1
View File
@@ -14,6 +14,7 @@
"dayjs": "^1.11.21",
"react": "^19.2.5",
"react-dom": "^19.2.5",
"react-quill-new": "^3.8.3",
"react-router-dom": "^7.15.0",
"xlsx-js-style": "^1.2.0",
"zustand": "^5.0.13"
+140 -17
View File
@@ -136,13 +136,37 @@ const AdminDashboard: React.FC = () => {
</Col>
</Row>
</div>
{/* 视频参数分布 */}
<div style={{ marginBottom: 20 }}>
<div style={{ marginBottom: 12 }}>
<Typography.Text strong style={{ fontSize: 15 }}></Typography.Text>
</div>
<Row gutter={[16, 16]}>
<Col xs={24} lg={8}>
<ChartCard title="分辨率分布" loading={loading}>
<PieBarChart data={stats?.videoResolutionUsage || []} />
</ChartCard>
</Col>
<Col xs={24} lg={8}>
<ChartCard title="画面比例分布" loading={loading}>
<PieBarChart data={stats?.videoRatioUsage || []} />
</ChartCard>
</Col>
<Col xs={24} lg={8}>
<ChartCard title="时长分布" loading={loading}>
<PieBarChart data={stats?.videoDurationUsage || []} />
</ChartCard>
</Col>
</Row>
</div>
</div>
);
};
// ── 图表卡片 ──
const ChartCard: React.FC<{ title: string; loading: boolean; children: React.ReactNode }> = ({ title, loading, children }) => (
<Card bordered={false} style={{ borderRadius: 12, border: '1px solid #f0f0f5', height: '100%' }}
<Card bordered={false} style={{ borderRadius: 16, border: '1px solid #f0f0f5', height: '100%', boxShadow: '0 4px 20px rgba(0,0,0,0.04)', transition: 'box-shadow 0.3s' }}
styles={{ body: { padding: '16px' } }}>
<div style={{ marginBottom: 12, display: 'flex', alignItems: 'center', justifyContent: 'space-between' }}>
<Typography.Text strong style={{ fontSize: 14 }}>{title}</Typography.Text>
@@ -186,19 +210,35 @@ const LineChart: React.FC<{ data: { date: string; module: string; credits: numbe
return (
<div style={{ height: 220, display: 'flex', flexDirection: 'column' }}>
<div style={{ flex: 1, display: 'flex', alignItems: 'flex-end', gap: 4, borderBottom: '1px solid #f1f5f9', paddingBottom: 4 }}>
<div style={{ flex: 1, display: 'flex', alignItems: 'flex-end', gap: 6, borderBottom: '1px solid #f1f5f9', paddingBottom: 4 }}>
{sevenDays.map(d => {
const val = dateMap.get(d) || 0;
const pct = (val / maxVal) * 100;
return (
<div key={d} style={{ flex: 1, display: 'flex', flexDirection: 'column', alignItems: 'center', height: '100%', justifyContent: 'flex-end' }}>
<div key={d} style={{ flex: 1, display: 'flex', flexDirection: 'column', alignItems: 'center', height: '100%', justifyContent: 'flex-end', position: 'relative' }}
onMouseEnter={e => {
const bar = e.currentTarget.querySelector('.bar') as HTMLElement;
const tip = e.currentTarget.querySelector('.tip') as HTMLElement;
if (bar) { bar.style.transform = 'scaleY(1.08)'; bar.style.filter = 'brightness(1.15) drop-shadow(0 4px 12px rgba(99,102,241,0.4))'; }
if (tip) { tip.style.opacity = '1'; }
}}
onMouseLeave={e => {
const bar = e.currentTarget.querySelector('.bar') as HTMLElement;
const tip = e.currentTarget.querySelector('.tip') as HTMLElement;
if (bar) { bar.style.transform = 'scaleY(1)'; bar.style.filter = 'none'; }
if (tip) { tip.style.opacity = '0'; }
}}
>
<div className="tip" style={{ position: 'absolute', bottom: '100%', marginBottom: 6, background: '#1e293b', color: '#fff', fontSize: 11, padding: '4px 10px', borderRadius: 6, whiteSpace: 'nowrap', opacity: 0, transition: 'opacity 0.2s', pointerEvents: 'none', zIndex: 10 }}>
{val > 0 ? `${val.toFixed(0)} 积分` : '无数据'}
</div>
<span style={{ fontSize: 9, color: '#6366f1', fontWeight: 600, marginBottom: 2 }}>{val > 0 ? val.toFixed(0) : ''}</span>
<div style={{ width: '65%', maxWidth: 32, height: `${Math.max(pct, 2)}%`, background: 'linear-gradient(180deg, #6366f1 0%, #8b5cf6 100%)', borderRadius: '3px 3px 0 0', minHeight: 3, transition: 'height 0.3s' }} />
<div className="bar" style={{ width: '65%', maxWidth: 32, height: `${Math.max(pct, 2)}%`, background: 'linear-gradient(180deg, #818cf8 0%, #6366f1 40%, #4f46e5 100%)', borderRadius: '4px 4px 0 0', minHeight: 4, transition: 'all 0.25s cubic-bezier(0.4, 0, 0.2, 1)', boxShadow: '0 2px 8px rgba(99,102,241,0.25)' }} />
</div>
);
})}
</div>
<div style={{ display: 'flex', gap: 4, marginTop: 4 }}>
<div style={{ display: 'flex', gap: 6, marginTop: 4 }}>
{sevenDays.map(d => (
<div key={d} style={{ flex: 1, textAlign: 'center' }}>
<span style={{ fontSize: 9, color: '#94a3b8' }}>{d.slice(5)}</span>
@@ -243,9 +283,27 @@ const ModulePie: React.FC<{ data: { module: string; credits: number }[] }> = ({
return (
<div style={{ height: 220, display: 'flex', alignItems: 'center', gap: 16 }}>
<svg width={size} height={size} viewBox={`0 0 ${size} ${size}`} style={{ flexShrink: 0 }}>
<svg width={size} height={size} viewBox={`0 0 ${size} ${size}`} style={{ flexShrink: 0, filter: 'drop-shadow(0 4px 12px rgba(0,0,0,0.08))' }}>
<defs>
{slices.map((s, i) => (
<linearGradient key={i} id={`pie-grad-${i}`} x1="0%" y1="0%" x2="100%" y2="100%">
<stop offset="0%" stopColor={s.color} stopOpacity={1} />
<stop offset="100%" stopColor={s.color} stopOpacity={0.7} />
</linearGradient>
))}
</defs>
{slices.map((s, i) => (
<path key={i} d={s.d} fill={s.color} stroke="#fff" strokeWidth={2} />
<path key={i} d={s.d} fill={`url(#pie-grad-${i})`} stroke="#fff" strokeWidth={2}
onMouseEnter={e => {
(e.target as SVGPathElement).style.transform = 'scale(1.06)';
(e.target as SVGPathElement).style.filter = 'brightness(1.1) drop-shadow(0 4px 8px rgba(0,0,0,0.2))';
}}
onMouseLeave={e => {
(e.target as SVGPathElement).style.transform = 'scale(1)';
(e.target as SVGPathElement).style.filter = 'none';
}}
style={{ transition: 'all 0.25s cubic-bezier(0.4, 0, 0.2, 1)', transformOrigin: `${cx}px ${cy}px`, cursor: 'pointer' }}
/>
))}
<circle cx={cx} cy={cy} r={36} fill="#fff" />
<text x={cx} y={cy - 4} textAnchor="middle" fontSize={11} fill="#64748b"></text>
@@ -253,8 +311,11 @@ const ModulePie: React.FC<{ data: { module: string; credits: number }[] }> = ({
</svg>
<div style={{ display: 'flex', flexDirection: 'column', gap: 6 }}>
{slices.map((s, i) => (
<div key={i} style={{ display: 'flex', alignItems: 'center', gap: 6, fontSize: 12, color: '#64748b' }}>
<div style={{ width: 10, height: 10, borderRadius: 3, background: s.color, flexShrink: 0 }} />
<div key={i} style={{ display: 'flex', alignItems: 'center', gap: 6, fontSize: 12, color: '#64748b', padding: '3px 6px', borderRadius: 6, transition: 'background-color 0.2s', cursor: 'default' }}
onMouseEnter={e => { e.currentTarget.style.backgroundColor = '#f8fafc'; }}
onMouseLeave={e => { e.currentTarget.style.backgroundColor = 'transparent'; }}
>
<div style={{ width: 10, height: 10, borderRadius: 3, background: s.color, flexShrink: 0, boxShadow: `0 2px 4px ${s.color}40` }} />
<span>{s.label}</span>
<span style={{ fontWeight: 600, color: '#1e293b' }}>{s.val.toFixed(0)}</span>
<span style={{ fontSize: 10 }}>{(s.pct * 100).toFixed(1)}%</span>
@@ -271,19 +332,26 @@ const HorizontalBarChart: React.FC<{ data: { teamName?: string; modelName?: stri
const maxVal = Math.max(...data.map(d => (d as any)[valueKey] || 0), 1);
return (
<div style={{ height: 220, display: 'flex', flexDirection: 'column', gap: 8, overflowY: 'auto' }}>
<div style={{ height: 220, display: 'flex', flexDirection: 'column', gap: 8, overflowY: 'auto', paddingRight: 4 }}>
{data.map((d, i) => {
const label = d.teamName || d.modelName || '-';
const val = (d as any)[valueKey] || 0;
const pct = (val / maxVal) * 100;
const c1 = COLORS[i % COLORS.length];
const c2 = COLORS[(i + 1) % COLORS.length];
return (
<div key={i}>
<div style={{ display: 'flex', justifyContent: 'space-between', marginBottom: 3 }}>
<span style={{ fontSize: 12, color: '#475569', overflow: 'hidden', textOverflow: 'ellipsis', whiteSpace: 'nowrap', maxWidth: '70%' }}>{label}</span>
<span style={{ fontSize: 12, fontWeight: 600, color: '#1e293b' }}>{val.toLocaleString()}</span>
<div key={i} style={{ padding: '4px 8px', borderRadius: 8, transition: 'background-color 0.2s, box-shadow 0.2s' }}
onMouseEnter={e => { e.currentTarget.style.backgroundColor = '#fafbff'; e.currentTarget.style.boxShadow = '0 2px 8px rgba(99,102,241,0.08)'; }}
onMouseLeave={e => { e.currentTarget.style.backgroundColor = 'transparent'; e.currentTarget.style.boxShadow = 'none'; }}
>
<div style={{ display: 'flex', justifyContent: 'space-between', marginBottom: 4 }}>
<span style={{ fontSize: 12, color: '#475569', overflow: 'hidden', textOverflow: 'ellipsis', whiteSpace: 'nowrap', maxWidth: '65%', fontWeight: 500 }}>{label}</span>
<span style={{ fontSize: 12, fontWeight: 700, color: '#1e293b' }}>{val.toLocaleString()}</span>
</div>
<div style={{ height: 16, background: '#f1f5f9', borderRadius: 4, overflow: 'hidden' }}>
<div style={{ height: '100%', width: `${pct}%`, background: `linear-gradient(90deg, ${COLORS[i % COLORS.length]}, ${COLORS[(i + 1) % COLORS.length]})`, borderRadius: 4, transition: 'width 0.3s' }} />
<div style={{ height: 18, background: '#f1f5f9', borderRadius: 6, overflow: 'hidden', boxShadow: 'inset 0 1px 2px rgba(0,0,0,0.06)' }}>
<div style={{ height: '100%', width: `${pct}%`, background: `linear-gradient(90deg, ${c1}, ${c2})`, borderRadius: 6, transition: 'width 0.35s cubic-bezier(0.4, 0, 0.2, 1)', boxShadow: `0 1px 3px ${c1}40`, position: 'relative' }}>
<div style={{ position: 'absolute', top: 0, left: 0, right: 0, height: '50%', background: 'linear-gradient(180deg, rgba(255,255,255,0.25) 0%, transparent 100%)', borderRadius: '6px 6px 0 0' }} />
</div>
</div>
</div>
);
@@ -292,7 +360,62 @@ const HorizontalBarChart: React.FC<{ data: { teamName?: string; modelName?: stri
);
};
const EmptyChart = () => (
// ── 视频参数分布(紧凑饼图+列表)──
const PieBarChart: React.FC<{ data: { model: string; label: string; count: number }[] }> = ({ data }) => {
if (!data.length) return <EmptyChart />;
// 按模型分组
const modelMap = new Map<string, { label: string; count: number }[]>();
data.forEach(d => {
if (!modelMap.has(d.model)) modelMap.set(d.model, []);
modelMap.get(d.model)!.push({ label: d.label, count: d.count });
});
const models = Array.from(modelMap.entries());
const total = data.reduce((s, d) => s + d.count, 0) || 1;
return (
<div style={{ height: 220, overflowY: 'auto' }}>
{models.map(([model, items], mi) => {
const modelTotal = items.reduce((s, it) => s + it.count, 0);
return (
<div key={model} style={{ marginBottom: 12 }}>
{/* 模型名称 + 总数 */}
<div style={{ display: 'flex', justifyContent: 'space-between', alignItems: 'center', marginBottom: 4 }}>
<span style={{ fontSize: 12, fontWeight: 600, color: COLORS[mi % COLORS.length] }}>{model}</span>
<span style={{ fontSize: 10, color: '#94a2b3' }}> {modelTotal} ({((modelTotal / total) * 100).toFixed(1)}%)</span>
</div>
{/* 各参数 */}
<div style={{ display: 'flex', flexWrap: 'wrap', gap: 4 }}>
{items.map((item, i) => (
<div key={i} style={{
flex: '0 0 calc(50% - 2px)',
padding: '4px 8px',
background: '#f8fafc',
borderRadius: 6,
border: '1px solid #f1f5f9',
transition: 'all 0.2s',
cursor: 'default',
}}
onMouseEnter={e => { Object.assign(e.currentTarget.style, { background: '#fafbff', boxShadow: '0 2px 8px rgba(99,102,241,0.1)', transform: 'translateY(-1px)' }); }}
onMouseLeave={e => { Object.assign(e.currentTarget.style, { background: '#f8fafc', boxShadow: 'none', transform: 'translateY(0)' }); }}
>
<div style={{ display: 'flex', justifyContent: 'space-between', alignItems: 'center' }}>
<span style={{ fontSize: 11, color: '#475569', fontWeight: 500 }}>{item.label}</span>
<span style={{ fontSize: 11, fontWeight: 700, color: '#1e293b' }}>{item.count}</span>
</div>
<div style={{ height: 4, background: '#e2e8f0', borderRadius: 2, marginTop: 3, overflow: 'hidden' }}>
<div style={{ height: '100%', width: `${(item.count / modelTotal) * 100}%`, background: COLORS[mi % COLORS.length], borderRadius: 2, transition: 'width 0.3s' }} />
</div>
</div>
))}
</div>
</div>
);
})}
</div>
);
};
const EmptyChart: React.FC = () => (
<div style={{ height: 220, display: 'flex', alignItems: 'center', justifyContent: 'center', color: '#94a3b8', fontSize: 13 }}>
</div>
@@ -5,6 +5,8 @@ import {
import {
BellOutlined, PlusOutlined, DeleteOutlined, SendOutlined, EyeOutlined, TeamOutlined,
} from '@ant-design/icons';
import ReactQuill from 'react-quill-new';
import 'react-quill-new/dist/quill.snow.css';
import { getAdminNotifications, createAdminNotification, deleteAdminNotification, getAdminUsers, getNotificationReadUsers } from '../api';
import { formatDate } from '../utils/formatDate';
@@ -131,6 +133,7 @@ const AdminNotificationManager: React.FC = () => {
},
{
title: '内容', dataIndex: 'content', ellipsis: true,
render: (v: string) => <div style={{ maxWidth: 300 }} dangerouslySetInnerHTML={{ __html: v }} />,
},
{
title: '类型', dataIndex: 'type', width: 80,
@@ -210,8 +213,8 @@ const AdminNotificationManager: React.FC = () => {
<Input placeholder="请输入消息标题" size="large" />
</Form.Item>
<Form.Item name="content" label="消息内容"
rules={[{ required: true, message: '请输入内容' }]}>
<Input.TextArea rows={4} placeholder="请输入消息内容" size="large" />
rules={[{ required: true, validator: (_, v) => v && v !== '<p><br></p>' ? Promise.resolve() : Promise.reject('请输入内容') }]}>
<ReactQuill theme="snow" placeholder="请输入消息内容(支持富文本:加粗、斜体、颜色、链接等)" style={{ height: 180, marginBottom: 40 }} />
</Form.Item>
<div style={{ display: 'flex', gap: 16 }}>
<Form.Item name="type" label="消息类型" style={{ flex: 1 }}
+9
View File
@@ -217,6 +217,12 @@ export interface ModelUsageOut {
count: number;
}
export interface VideoParamOut {
model: string;
label: string;
count: number;
}
export interface AdminStats {
totalUsers: number;
totalProjects: number;
@@ -236,6 +242,9 @@ export interface AdminStats {
periodCreditsByModule: DailyCredit[];
creditsByTeam: TeamCredit[];
modelUsage: ModelUsageOut[];
videoResolutionUsage: VideoParamOut[];
videoRatioUsage: VideoParamOut[];
videoDurationUsage: VideoParamOut[];
}
export interface PaymentStats {
@@ -2,6 +2,7 @@ export interface GenerationStatusLike {
status?: string | null;
displayStatus?: string | null;
pipelineStage?: string | null;
shouldPoll?: boolean | null;
}
export type GenerationUiColor = 'default' | 'processing' | 'warning' | 'success' | 'error' | 'blue' | 'orange' | 'purple';
@@ -19,60 +20,42 @@ export interface GenerationUiState {
isTerminal: boolean;
}
const ACTIVE_STATUS_KEYS = new Set(['pending', 'optimizing', 'prompt_optimized', 'generating']);
const ACTIVE_STATUS_KEYS = new Set(['pending', 'optimizing', 'generating']);
const ACTIVE_PIPELINE_STAGES = new Set([
'queued', 'preparing', 'creating_provider_task', 'provider_result_staged',
'waiting_remote', 'polling', 'result_ready', 'download_queued', 'downloading',
'retry_waiting', 'upscale_queued', 'upscale_processing', 'upscale_polling',
'upscale_downloading', 'upscale_finalizing', 'upscale_retry_waiting',
'retry_waiting', 'recovery_inconsistent', 'upscale_queued', 'upscale_processing',
'upscale_polling', 'upscale_downloading', 'upscale_finalizing', 'upscale_retry_waiting',
]);
const SUCCESS_KEYS = new Set(['completed', 'done']);
const FAILURE_KEYS = new Set(['failed', 'timeout', 'download_failed', 'upscale_failed']);
const TERMINAL_KEYS = new Set([...SUCCESS_KEYS, ...FAILURE_KEYS, 'deleted']);
const LABELS: Record<string, string> = {
pending: '待处理',
optimizing: '优化中',
prompt_optimized: '待生成',
generating: '生成中',
queued: '已入队',
preparing: '准备中',
creating_provider_task: '创建供应商任务',
provider_result_staged: '供应商结果已暂存',
waiting_remote: '等待供应商结果',
polling: '轮询供应商结果',
result_ready: '远程结果已就绪',
download_queued: '下载已入队',
downloading: '下载中',
retry_waiting: '下载等待重试',
upscale_queued: '超分已入队',
upscale_processing: '本地超分处理中',
upscale_polling: '轮询远程超分',
upscale_downloading: '下载超分结果',
upscale_finalizing: '超分结果最终化',
upscale_retry_waiting: '超分等待重试',
completed: '已完成',
done: '已完成',
timeout: '任务超时',
download_failed: '下载失败',
upscale_failed: '超分失败',
failed: '失败',
deleted: '已删除',
pending: '待处理', optimizing: '提词处理中', prompt_optimized: '待生成', generating: '生成中',
queued: '已入队', preparing: '准备中', creating_provider_task: '创建供应商任务',
provider_result_staged: '供应商结果已暂存', waiting_remote: '等待供应商结果', polling: '轮询供应商结果',
result_ready: '远程结果已就绪', download_queued: '下载已入队', downloading: '下载中',
retry_waiting: '下载等待重试', recovery_inconsistent: '恢复证据异常', upscale_queued: '超分已入队',
upscale_processing: '本地超分处理中', upscale_polling: '轮询远程超分',
upscale_downloading: '下载超分结果', upscale_finalizing: '超分结果最终化',
upscale_retry_waiting: '超分等待重试', completed: '已完成', done: '已完成', timeout: '任务超时',
download_failed: '下载失败', upscale_failed: '超分失败', failed: '失败', deleted: '已删除',
};
const COLOR_MAP: Record<string, GenerationUiColor> = {
pending: 'default', optimizing: 'processing', prompt_optimized: 'processing', generating: 'warning',
pending: 'default', optimizing: 'processing', prompt_optimized: 'blue', generating: 'warning',
queued: 'processing', preparing: 'processing', creating_provider_task: 'processing',
provider_result_staged: 'processing', waiting_remote: 'processing', polling: 'processing',
result_ready: 'processing', download_queued: 'processing', downloading: 'processing',
retry_waiting: 'orange', upscale_queued: 'purple', upscale_processing: 'purple',
result_ready: 'processing', download_queued: 'processing', downloading: 'processing', retry_waiting: 'orange',
recovery_inconsistent: 'orange', upscale_queued: 'purple', upscale_processing: 'purple',
upscale_polling: 'purple', upscale_downloading: 'purple', upscale_finalizing: 'purple',
upscale_retry_waiting: 'orange', completed: 'success', done: 'success', failed: 'error',
timeout: 'error', download_failed: 'error', upscale_failed: 'error', deleted: 'default',
};
const normalize = (value?: string | null): string => String(value || '').trim().toLowerCase();
const firstMatching = (values: string[], keys: Set<string>): string => values.find((value) => keys.has(value)) || '';
const firstMatching = (values: string[], keys: Set<string>): string => values.find((item) => keys.has(item)) || '';
export const getGenerationStageLabel = (key?: string | null): string => {
const normalized = normalize(key);
@@ -88,41 +71,22 @@ export const resolveGenerationUiState = (value: GenerationStatusLike): Generatio
const status = normalize(value.status);
const displayStatus = normalize(value.displayStatus);
const pipelineStage = normalize(value.pipelineStage);
const values = [displayStatus, status, pipelineStage].filter(Boolean);
const failureKey = FAILURE_KEYS.has(pipelineStage)
? pipelineStage
: firstMatching([displayStatus, status], FAILURE_KEYS);
const deletedKey = firstMatching(values, new Set(['deleted']));
const values = [pipelineStage, displayStatus, status].filter(Boolean);
const failureKey = firstMatching(values, FAILURE_KEYS);
const successKey = firstMatching(values, SUCCESS_KEYS);
const effectiveKey = failureKey
|| deletedKey
|| (ACTIVE_PIPELINE_STAGES.has(pipelineStage) ? pipelineStage : '')
|| successKey
|| pipelineStage
|| displayStatus
|| status
|| 'pending';
const deletedKey = firstMatching(values, new Set(['deleted']));
const effectiveKey = failureKey || deletedKey || successKey || pipelineStage || displayStatus || status || 'pending';
const isFailure = FAILURE_KEYS.has(effectiveKey);
const isSuccess = SUCCESS_KEYS.has(effectiveKey);
const isActive = !isFailure && !isSuccess && effectiveKey !== 'deleted' && (
ACTIVE_PIPELINE_STAGES.has(pipelineStage)
|| ACTIVE_STATUS_KEYS.has(displayStatus)
|| ACTIVE_STATUS_KEYS.has(status)
|| ACTIVE_PIPELINE_STAGES.has(effectiveKey)
);
const isActive = typeof value.shouldPoll === 'boolean'
? value.shouldPoll
: (!isFailure && !isSuccess && effectiveKey !== 'deleted' && (ACTIVE_PIPELINE_STAGES.has(pipelineStage) || ACTIVE_STATUS_KEYS.has(status) || ACTIVE_STATUS_KEYS.has(displayStatus)));
return {
status,
displayStatus,
pipelineStage,
effectiveKey,
status, displayStatus, pipelineStage, effectiveKey,
label: getGenerationStageLabel(effectiveKey),
color: getGenerationStatusColor(effectiveKey),
isActive,
isSuccess,
isFailure,
isActive, isSuccess, isFailure,
isTerminal: TERMINAL_KEYS.has(effectiveKey),
};
};
+52
View File
@@ -36,6 +36,7 @@ from app.schemas.admin import (
DailyCreditOut,
TeamCreditOut,
ModelUsageOut,
VideoParamOut,
CreateUserRequest,
UpdateMenusRequest,
ResetPasswordRequest,
@@ -1711,6 +1712,15 @@ async def list_operation_logs(
}
def _build_param_out(model_map: dict[str, dict[str, int]]) -> list[VideoParamOut]:
"""{模型: {标签: 数量}} 转为扁平列表,按模型+数量排序。"""
result: list[VideoParamOut] = []
for model, labels in model_map.items():
for label, count in sorted(labels.items(), key=lambda x: -x[1]):
result.append(VideoParamOut(model=model, label=label, count=count))
return result
# ── Stats ────────────────────────────────────────────────
@router.get("/stats", response_model=AdminStatsOut)
@@ -1950,6 +1960,45 @@ async def get_stats(
for row in model_usage_rows
]
# ── 视频分辨率/比例/时长使用分布(按模型分组)
_video_gen_q = (
select(
func.coalesce(CreditRecord.engine_name, '未知').label('model'),
ChatGenerationTask.resolution,
ChatGenerationTask.aspect_ratio,
ChatGenerationTask.duration,
)
.join(CreditRecord, CreditRecord.related_id == ChatGenerationTask.id)
.where(
ChatGenerationTask.gen_type == "video",
CreditRecord.type == "consume",
ChatGenerationTask.created_at >= date_start,
ChatGenerationTask.created_at <= date_end,
ChatGenerationTask.deleted_at.is_(None),
)
)
_video_rows = (await db.execute(_video_gen_q)).all()
_res_map: dict[str, dict[str, int]] = {}
_ratio_map: dict[str, dict[str, int]] = {}
_dur_map: dict[str, dict[str, int]] = {}
for row in _video_rows:
model = row.model or '未知'
if row.resolution:
_res_map.setdefault(model, {})
_res_map[model][row.resolution] = _res_map[model].get(row.resolution, 0) + 1
if row.aspect_ratio:
_ratio_map.setdefault(model, {})
_ratio_map[model][row.aspect_ratio] = _ratio_map[model].get(row.aspect_ratio, 0) + 1
if row.duration:
_k = f"{row.duration}"
_dur_map.setdefault(model, {})
_dur_map[model][_k] = _dur_map[model].get(_k, 0) + 1
video_resolution_usage = _build_param_out(_res_map)
video_ratio_usage = _build_param_out(_ratio_map)
video_duration_usage = _build_param_out(_dur_map)
return AdminStatsOut(
total_users=total_users,
total_projects=total_projects,
@@ -1969,6 +2018,9 @@ async def get_stats(
period_credits_by_module=period_credits_by_module,
credits_by_team=credits_by_team,
model_usage=model_usage,
video_resolution_usage=video_resolution_usage,
video_ratio_usage=video_ratio_usage,
video_duration_usage=video_duration_usage,
)
File diff suppressed because it is too large Load Diff
+2
View File
@@ -19,3 +19,5 @@ from app.enums.audio_reference import *
from app.enums.private_portrait import *
from app.enums.generation_provider import *
from app.enums.generation_record import *
@@ -0,0 +1,31 @@
from __future__ import annotations
from enum import StrEnum
class GenerationRecordConfigSourceEnum(StrEnum):
"""生成记录配置冻结来源。"""
PROMPT_OPTIMIZE = "prompt_optimize"
LEGACY_GENERATE_FALLBACK = "legacy_generate_fallback"
EXISTING_FROZEN_CONFIG = "existing_frozen_config"
class GenerationRecordEventTypeEnum(StrEnum):
"""GenerationRecord 用户生成链路事件。"""
PROMPT_CONFIG_VALIDATE_START = "PROMPT_CONFIG_VALIDATE_START"
PROMPT_CONFIG_VALIDATE_SUCCESS = "PROMPT_CONFIG_VALIDATE_SUCCESS"
PROMPT_CONFIG_VALIDATE_FAILED = "PROMPT_CONFIG_VALIDATE_FAILED"
PROMPT_CONFIG_FREEZE_START = "PROMPT_CONFIG_FREEZE_START"
PROMPT_CONFIG_FREEZE_SUCCESS = "PROMPT_CONFIG_FREEZE_SUCCESS"
PROMPT_CONFIG_FREEZE_FAILED = "PROMPT_CONFIG_FREEZE_FAILED"
LEGACY_CONFIG_FALLBACK_START = "LEGACY_CONFIG_FALLBACK_START"
LEGACY_CONFIG_FALLBACK_SUCCESS = "LEGACY_CONFIG_FALLBACK_SUCCESS"
LEGACY_CONFIG_FALLBACK_FAILED = "LEGACY_CONFIG_FALLBACK_FAILED"
LEGACY_CONFIG_FALLBACK_SKIPPED = "LEGACY_CONFIG_FALLBACK_SKIPPED"
GENERATION_SUBMIT_START = "GENERATION_SUBMIT_START"
GENERATION_SUBMIT_CONFIG_READY = "GENERATION_SUBMIT_CONFIG_READY"
GENERATION_SUBMIT_BILLING_SUCCESS = "GENERATION_SUBMIT_BILLING_SUCCESS"
GENERATION_SUBMIT_ENQUEUE_SUCCESS = "GENERATION_SUBMIT_ENQUEUE_SUCCESS"
GENERATION_SUBMIT_FAILED = "GENERATION_SUBMIT_FAILED"
@@ -22,6 +22,7 @@ class GenerationRecordPipelineStage(str, Enum):
DOWNLOAD_QUEUED = "download_queued"
DOWNLOADING = "downloading"
RETRY_WAITING = "retry_waiting"
RECOVERY_INCONSISTENT = "recovery_inconsistent"
UPSCALE_QUEUED = "upscale_queued"
UPSCALE_PROCESSING = "upscale_processing"
UPSCALE_POLLING = "upscale_polling"
@@ -47,6 +47,7 @@ class ChatGenerationPipelineStage(str, Enum):
DOWNLOAD_QUEUED = "download_queued"
DOWNLOADING = "downloading"
RETRY_WAITING = "retry_waiting"
RECOVERY_INCONSISTENT = "recovery_inconsistent"
UPSCALE_QUEUED = "upscale_queued"
UPSCALE_PROCESSING = "upscale_processing"
UPSCALE_POLLING = "upscale_polling"
@@ -107,6 +108,7 @@ class ChatGenerationTaskEventType(str, Enum):
FINAL_POLL_BEFORE_TIMEOUT_PENDING = "FINAL_POLL_BEFORE_TIMEOUT_PENDING"
GENERATION_RECOVERY_ENQUEUE = "GENERATION_RECOVERY_ENQUEUE"
GENERATION_RECOVERY_TIMEOUT = "GENERATION_RECOVERY_TIMEOUT"
GENERATION_RECOVERY_INCONSISTENT = "GENERATION_RECOVERY_INCONSISTENT"
DOWNLOAD_ENQUEUE = "DOWNLOAD_ENQUEUE"
DOWNLOAD_ENQUEUE_FAILED = "DOWNLOAD_ENQUEUE_FAILED"
+9
View File
@@ -127,6 +127,12 @@ class ModelUsageOut(BaseModel):
count: int
class VideoParamOut(BaseModel):
model: str
label: str
count: int
class AdminStatsOut(BaseModel):
total_users: int
total_projects: int
@@ -147,6 +153,9 @@ class AdminStatsOut(BaseModel):
period_credits_by_module: list[DailyCreditOut] = []
credits_by_team: list[TeamCreditOut] = []
model_usage: list[ModelUsageOut] = []
video_resolution_usage: list[VideoParamOut] = []
video_ratio_usage: list[VideoParamOut] = []
video_duration_usage: list[VideoParamOut] = []
class AdminCreditRecordSummaryOut(BaseModel):
+17 -21
View File
@@ -1,22 +1,18 @@
from pydantic import BaseModel, Field
from app.enums.generation_status import (
GenerationStatus,
GenerationType,
DURATIONS,
ASPECT_RATIOS,
RESOLUTIONS,
IMAGE_SIZES,
)
from app.enums.generation_status import GenerationType
from app.schemas.common import NaiveDatetime, NaiveDatetimeOptional
from app.services.operation_log import log_operation
class OptimizeParams(BaseModel):
project_id: str
prompt: str = Field(..., max_length=500)
gen_type: GenerationType = Field(GenerationType.video, description="生成类型:video-视频,image-图片")
engine_id: str = Field(..., min_length=1, max_length=32, description="提词阶段选定并冻结的生成引擎ID")
include_media_references: bool = Field(False, description="资源生成时是否携带本次提词附件;提词完成后不可修改")
duration: int | None = Field(None, description="视频时长(秒),视频生成必填")
aspect_ratio: str | None = Field(None, description="视频比例,视频生成必填")
resolution: str | None = Field(None, description="视频目标分辨率,视频生成必填")
image_size: str | None = Field(None, description="画面分辨率,图片生成使用")
image_proportion: str | None = Field(None, description="图片比例,图片生成使用")
image_px: str | None = Field(None, description="图片像素大小,图片生成使用")
@@ -24,18 +20,10 @@ class OptimizeParams(BaseModel):
idempotency_key: str | None = Field(None, max_length=64, description="幂等键,防止重复请求")
class GenerateParams(BaseModel):
engine_id: str | None = Field(None, description="生成引擎ID;为空时优先沿用记录引擎,再回退默认引擎")
include_media_references: bool = Field(False, description="最终生成时是否携带提词阶段保存的附件")
aspect_ratio: str | None = None
resolution: str | None = None
image_size: str | None = None
class OptimizeResult(BaseModel):
optimized_prompt: str
text_credits_cost: float
# text_tokens_used: int
text_tokens_used: int = 0
record: "GenerationRecordOut"
class GenerationRecordOut(BaseModel):
@@ -62,11 +50,19 @@ class GenerationRecordOut(BaseModel):
engine_name: str | None = None
engine_snapshot: dict | None = None
include_media_references: bool = False
config_complete: bool = False
config_recoverable: bool = False
config_fallback_hint: str | None = None
can_generate: bool = False
can_retry: bool = False
should_poll: bool = False
client_status: str = "ready"
operation_phase: str = "prompt"
text_credits_cost: float = 0.0
# text_tokens_used: int = 0
text_tokens_used: int = 0
credits_cost: float = 0.0
# video_tokens_used: int = 0
# image_tokens_used: int = 0
video_tokens_used: int = 0
image_tokens_used: int = 0
error_message: str | None = None
created_at: NaiveDatetime
generated_at: NaiveDatetimeOptional = None
@@ -8,9 +8,8 @@ from app.enums.generation_task import GenerationMode, GenerationOwnerType
from app.models.base import async_session
from app.models.chat_generation_task import ChatGenerationTask
from app.models.chat_generation_task_event import ChatGenerationTaskEvent
from app.models.chat_provider_call_log import ChatProviderCallLog
from app.models.generation_record import GenerationRecord
from app.services.operation_log_service import build_exception_detail, log_operation_event, sanitize_log_value
from app.services.operation_log_service import build_exception_detail, log_ai_model_event, log_operation_event, sanitize_log_value
from app.utils.id_gen import generate_id
MAX_EXCERPT_CHARS = 2000
@@ -64,8 +63,14 @@ def _owner_fields(
resolved_record_id = None
resolved_mode = generation_mode or getattr(obj, "generation_mode", GenerationMode.CHATAPI_ASYNC.value)
else:
resolved_owner_type = owner_type or GenerationOwnerType.CHAT_GENERATION_TASK.value
resolved_owner_id = owner_id
inferred_mode = generation_mode or getattr(obj, "generation_mode", None)
inferred_owner_type = (
GenerationOwnerType.GENERATION_RECORD.value
if inferred_mode == GenerationMode.GENERATION_RECORD.value
else GenerationOwnerType.CHAT_GENERATION_TASK.value
)
resolved_owner_type = owner_type or inferred_owner_type
resolved_owner_id = owner_id or getattr(obj, "id", None)
if resolved_owner_type == GenerationOwnerType.GENERATION_RECORD.value:
resolved_task_id = None
resolved_record_id = resolved_owner_id
@@ -201,8 +206,16 @@ async def log_provider_call(
total_tokens: int = 0,
error_code: str | None = None,
error_message: str | None = None,
) -> None:
"""Write an owner-scoped provider call log in a separate transaction."""
call_id: str | None = None,
module: str | None = None,
step_code: str | None = None,
) -> str | None:
"""Write provider audit events to the AiModel file log only.
``ChatProviderCallLog`` is intentionally no longer written. The model and
historical table remain registered for backward compatibility, so no schema
migration is required.
"""
obj = task or record
fields = _owner_fields(
obj,
@@ -214,34 +227,75 @@ async def log_provider_call(
generation_mode=generation_mode,
)
if not fields:
return
try:
async with async_session() as db:
db.add(ChatProviderCallLog(
id=generate_id(),
owner_type=fields["owner_type"],
task_id=fields["task_id"],
generation_record_id=fields["generation_record_id"],
generation_attempt_no=fields["generation_attempt_no"],
generation_mode=fields["generation_mode"],
provider=provider,
api_type=api_type,
model=model,
engine_id=engine_id,
status=status,
latency_ms=latency_ms,
http_status=http_status,
provider_task_id=provider_task_id,
request_hash=_hash(request_data),
response_hash=_hash(response_data),
request_excerpt=_excerpt(request_data),
response_excerpt=_excerpt(response_data),
prompt_tokens=prompt_tokens or 0,
completion_tokens=completion_tokens or 0,
total_tokens=total_tokens or 0,
error_code=error_code,
error_message=error_message,
))
await db.commit()
except Exception as exc:
_fallback_log(api_type, fields, exc)
return None
resolved_call_id = call_id or generate_id()
resolved_module = module or fields.get("generation_mode") or "generation_pipeline"
resolved_step = step_code or api_type
common = {
"module": resolved_module,
"step_code": resolved_step,
"call_id": resolved_call_id,
"source": "app.services.generation.log_service",
"task_id": fields.get("owner_id"),
"owner_type": fields.get("owner_type"),
"owner_id": fields.get("owner_id"),
"generation_attempt_no": fields.get("generation_attempt_no"),
"remote_action": api_type,
"remote_request_id": provider_task_id,
"model_config_id": engine_id,
"model_config_name": engine_id,
"model_name": model,
"provider": provider,
"http_status": http_status,
}
detail = {
"generation_mode": fields.get("generation_mode"),
"provider_task_id": provider_task_id,
"error_code": error_code,
}
token_usage = {
"prompt_tokens": int(prompt_tokens or 0),
"completion_tokens": int(completion_tokens or 0),
"total_tokens": int(total_tokens or 0),
}
if request_data is not None or str(status).lower() == "request":
log_ai_model_event(
event_type="REQUEST",
event_phase="REQUEST",
event_status="started",
request=request_data if request_data is not None else {},
detail=detail,
**common,
)
normalized_status = str(status or "").lower()
if response_data is not None or normalized_status in {"success", "completed", "succeeded"}:
log_ai_model_event(
event_type="RESPONSE",
event_phase="RESPONSE",
event_status="success" if normalized_status not in {"failed", "error"} else "failed",
latency_ms=latency_ms,
response=response_data,
token_usage=token_usage,
detail=detail,
error=error_message if normalized_status in {"failed", "error"} else None,
**common,
)
if normalized_status in {"failed", "error"} or error_message:
log_ai_model_event(
event_type="ERROR",
event_phase="ERROR",
event_status="failed",
latency_ms=latency_ms,
token_usage=token_usage,
detail=build_exception_detail(
RuntimeError(error_message or "provider call failed"),
detail,
),
error=error_message or "provider call failed",
**common,
)
return resolved_call_id
@@ -7,24 +7,50 @@ from app.services.generation.pipeline.owner_service import GenerationOwner, owne
async def enqueue_generation_create(
owner: GenerationOwner,
owner: GenerationOwner | None = None,
*,
reason: str,
owner_type: str | None = None,
owner_id: str | None = None,
generation_attempt_no: int | None = None,
generation_mode: str | None = None,
) -> None:
"""Commit caller-owned state before invoking this function."""
"""Commit caller-owned state before invoking this function.
Scalar owner fields are accepted so callers can avoid touching an ORM object after
commit. Existing callers may continue passing ``owner``.
"""
from app.tasks.generation_create_tasks import chatapi_create_generation_task
owner_type = owner_type_of(owner)
attempt_no = int(getattr(owner, "generation_attempt_no", 1) or 1)
resolved_owner_type = owner_type or (owner_type_of(owner) if owner is not None else None)
resolved_owner_id = owner_id or (str(owner.id) if owner is not None else None)
resolved_attempt_no = int(
generation_attempt_no
or (getattr(owner, "generation_attempt_no", 1) if owner is not None else 1)
or 1
)
if not resolved_owner_type or not resolved_owner_id:
raise ValueError("投递生成任务缺少 owner_type 或 owner_id")
try:
chatapi_create_generation_task.apply_async(
args=[str(owner.id)],
kwargs={"owner_type": owner_type, "generation_attempt_no": attempt_no},
args=[resolved_owner_id],
kwargs={
"owner_type": resolved_owner_type,
"generation_attempt_no": resolved_attempt_no,
},
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
task_id=f"generation-create:{owner_type}:{owner.id}:attempt:{attempt_no}",
task_id=(
f"generation-create:{resolved_owner_type}:{resolved_owner_id}:"
f"attempt:{resolved_attempt_no}"
),
)
await log_task_event(
owner,
owner_type=resolved_owner_type,
owner_id=resolved_owner_id,
generation_attempt_no=resolved_attempt_no,
generation_mode=generation_mode,
event_type=ChatGenerationTaskEventType.GENERATION_RECORD_ENQUEUE_SUCCESS.value,
message="资源生成创建任务已投递",
detail={"reason": reason, "queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
@@ -32,6 +58,10 @@ async def enqueue_generation_create(
except Exception as exc:
await log_task_event(
owner,
owner_type=resolved_owner_type,
owner_id=resolved_owner_id,
generation_attempt_no=resolved_attempt_no,
generation_mode=generation_mode,
event_type=ChatGenerationTaskEventType.GENERATION_RECORD_ENQUEUE_FAILED.value,
message=str(exc),
detail={"reason": reason, "queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
@@ -0,0 +1,491 @@
from __future__ import annotations
import json
from types import SimpleNamespace
from typing import Any, Iterable
from fastapi import HTTPException
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.common import LogEventStatusEnum
from app.enums.generation_record import (
GenerationRecordConfigSourceEnum,
GenerationRecordEventTypeEnum,
)
from app.enums.generation_status import (
GenerationType,
)
from app.models.generation_record import GenerationRecord
from app.models.image_engine import ImageEngine
from app.models.video_engine import VideoEngine
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,
image_supported_sizes,
normalize_px,
parse_json_list,
)
from app.services.generation.pipeline.generation_record_service import freeze_generation_record_config
from app.services.operation_log_service import log_operation_event, log_operation_error
from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot
from app.utils.exceptions import InvalidStatusError
_GENERATION_RECORD_LOG_DOMAIN = "generation_record"
_GENERATION_RECORD_LOG_MODULE = "generation_record"
def _json_loads_object(value: str | None) -> dict[str, Any] | None:
if not value:
return None
try:
data = json.loads(value)
except (TypeError, json.JSONDecodeError):
return None
return data if isinstance(data, dict) else None
def _json_loads_list(value: str | None) -> list[Any]:
if not value:
return []
try:
data = json.loads(value)
except (TypeError, json.JSONDecodeError):
return []
return data if isinstance(data, list) else []
def generation_record_engine_snapshot(record: GenerationRecord) -> dict[str, Any] | None:
return _json_loads_object(record.engine_snapshot_json)
def is_generation_record_config_complete(record: GenerationRecord) -> bool:
snapshot = generation_record_engine_snapshot(record)
if not record.engine_id or not snapshot:
return False
if record.gen_type == GenerationType.video.value:
return bool(record.duration and record.aspect_ratio and record.resolution)
if record.gen_type == GenerationType.image.value:
return bool(record.image_size and record.image_proportion and record.image_px)
return False
def is_generation_record_config_recoverable(record: GenerationRecord) -> bool:
"""Return whether a prompt_optimized legacy row can try server-side config fallback.
This check intentionally avoids extra DB reads for list pages. The actual engine
existence and capability validation is performed while the generate API holds a
row lock for the single target record.
"""
if is_generation_record_config_complete(record):
return False
if record.status != "prompt_optimized":
return False
if not record.optimized_prompt:
return False
return record.gen_type in {GenerationType.video.value, GenerationType.image.value}
def generation_record_config_fallback_hint(record: GenerationRecord) -> str | None:
if not is_generation_record_config_recoverable(record):
return None
return "旧版本记录缺少冻结配置,提交生成时将由后端按可用引擎权重自动补齐一次"
def frozen_generation_record_engine_view(record: GenerationRecord) -> SimpleNamespace:
snapshot = generation_record_engine_snapshot(record)
if not snapshot:
raise InvalidStatusError("该记录缺少冻结的引擎配置,请重新生成提词")
snapshot = dict(snapshot)
snapshot["id"] = record.engine_id
return SimpleNamespace(**snapshot)
def _engine_plain_namespace(engine: ImageEngine | VideoEngine | SimpleNamespace) -> SimpleNamespace:
if isinstance(engine, SimpleNamespace):
return engine
return SimpleNamespace(
**{
key: value
for key, value in vars(engine).items()
if key != "_sa_instance_state"
}
)
def _safe_int(value: Any) -> int | None:
try:
return int(value)
except (TypeError, ValueError):
return None
def _first_existing_or_default(value: str | None, supported: Iterable[str], default_value: str) -> str:
normalized_supported = [str(item).strip() for item in supported if str(item or "").strip()]
current = str(value or "").strip()
if current and (not normalized_supported or current in normalized_supported):
return current
if default_value in normalized_supported or not normalized_supported:
return default_value
return normalized_supported[0]
def _first_duration(value: int | None, supported: Iterable[Any], max_duration: int | None) -> int:
supported_ints = [int(item) for item in supported if str(item).isdigit()]
current = _safe_int(value)
if current and current > 0:
if (not supported_ints or current in supported_ints) and (not max_duration or current <= int(max_duration or 0)):
return current
for item in supported_ints:
if item > 0 and (not max_duration or item <= int(max_duration or 0)):
return item
if max_duration and int(max_duration) > 0:
return min(VIDEO_DEFAULT_DURATION, int(max_duration)) or int(max_duration)
return VIDEO_DEFAULT_DURATION
def _video_engine_supports_record_params(engine: VideoEngine, record: GenerationRecord) -> bool:
ratios = [str(item) for item in parse_json_list(engine.supported_ratios, [])]
resolutions = [str(item) for item in parse_json_list(engine.supported_resolutions, [])]
durations = [int(item) for item in parse_json_list(engine.supported_durations, []) if str(item).isdigit()]
duration = _safe_int(record.duration)
if record.aspect_ratio and ratios and record.aspect_ratio not in ratios:
return False
if record.resolution and resolutions and record.resolution not in resolutions:
return False
if duration and durations and duration not in durations:
return False
if duration and int(engine.max_duration or 0) > 0 and duration > int(engine.max_duration or 0):
return False
return True
def _image_engine_supports_record_params(engine: ImageEngine, record: GenerationRecord) -> bool:
sizes = image_supported_sizes(engine)
if record.image_size and sizes and record.image_size not in sizes:
return False
if record.image_size and record.image_proportion and sizes:
ratios = sizes.get(record.image_size) or {}
if ratios and record.image_proportion not in ratios:
return False
return True
async def _list_active_video_engines(db: AsyncSession) -> list[VideoEngine]:
result = await db.execute(
select(VideoEngine)
.where(VideoEngine.is_active == True, VideoEngine.deleted_at.is_(None))
.order_by(VideoEngine.priority.desc(), VideoEngine.created_at.asc(), VideoEngine.id.asc())
)
return list(result.scalars().all())
async def _list_active_image_engines(db: AsyncSession) -> list[ImageEngine]:
result = await db.execute(
select(ImageEngine)
.where(ImageEngine.is_active == True, ImageEngine.deleted_at.is_(None))
.order_by(ImageEngine.priority.desc(), ImageEngine.created_at.asc(), ImageEngine.id.asc())
)
return list(result.scalars().all())
def _select_video_engine(engines: list[VideoEngine], record: GenerationRecord) -> tuple[VideoEngine, str]:
if record.engine_id:
for engine in engines:
if engine.id == record.engine_id:
return engine, "existing_record"
for engine in engines:
if _video_engine_supports_record_params(engine, record):
return engine, "priority_param_match"
if engines:
return engines[0], "priority_fallback"
raise InvalidStatusError("没有可用的视频引擎,无法补齐历史生成配置")
def _select_image_engine(engines: list[ImageEngine], record: GenerationRecord) -> tuple[ImageEngine, str]:
if record.engine_id:
for engine in engines:
if engine.id == record.engine_id:
return engine, "existing_record"
for engine in engines:
if _image_engine_supports_record_params(engine, record):
return engine, "priority_param_match"
if engines:
return engines[0], "priority_fallback"
raise InvalidStatusError("没有可用的图片引擎,无法补齐历史生成配置")
def _normalize_video_record_params(record: GenerationRecord, engine: VideoEngine) -> None:
ratios = [str(item) for item in parse_json_list(engine.supported_ratios, [])]
resolutions = [str(item) for item in parse_json_list(engine.supported_resolutions, [])]
durations = parse_json_list(engine.supported_durations, [])
record.duration = _first_duration(record.duration, durations, int(engine.max_duration or 0) or None)
record.aspect_ratio = _first_existing_or_default(record.aspect_ratio, ratios, VIDEO_DEFAULT_RATIO)
record.resolution = _first_existing_or_default(record.resolution, resolutions, VIDEO_DEFAULT_RESOLUTION)
def _normalize_image_record_params(record: GenerationRecord, engine: ImageEngine) -> None:
sizes = image_supported_sizes(engine)
size_keys = [str(item) for item in sizes.keys() if str(item or "").strip()]
current_size = str(record.image_size or "").strip()
default_size = str(engine.default_size or IMAGE_DEFAULT_SIZE).strip() or IMAGE_DEFAULT_SIZE
if current_size and (not sizes or current_size in sizes):
image_size = current_size
elif default_size in size_keys:
image_size = default_size
elif IMAGE_DEFAULT_SIZE in size_keys:
image_size = IMAGE_DEFAULT_SIZE
elif size_keys:
image_size = size_keys[0]
else:
image_size = current_size or default_size or IMAGE_DEFAULT_SIZE
ratios = sizes.get(image_size) if sizes else {}
ratio_keys = [str(item) for item in (ratios or {}).keys() if str(item or "").strip()]
current_ratio = str(record.image_proportion or "").strip()
if current_ratio and (not ratio_keys or current_ratio in ratio_keys):
image_proportion = current_ratio
elif IMAGE_DEFAULT_PROPORTION in ratio_keys or not ratio_keys:
image_proportion = IMAGE_DEFAULT_PROPORTION
else:
image_proportion = ratio_keys[0]
px_map = ratios or {}
current_px = normalize_px(str(record.image_px or "").strip()) if record.image_px else ""
image_px = normalize_px(str(px_map.get(image_proportion) or "").strip()) or current_px or IMAGE_DEFAULT_PX
record.image_size = image_size
record.image_proportion = image_proportion
record.image_px = image_px
def _config_log_detail(
record: GenerationRecord,
*,
source: str,
engine_selected_by: str | None = None,
before: dict[str, Any] | None = None,
extra: dict[str, Any] | None = None,
) -> dict[str, Any]:
references = _json_loads_list(record.media_references)
detail: dict[str, Any] = {
"record_id": record.id,
"user_id": record.user_id,
"project_id": record.project_id,
"gen_type": record.gen_type,
"source": source,
"engine_selected_by": engine_selected_by,
"engine_id": record.engine_id,
"duration": record.duration,
"aspect_ratio": record.aspect_ratio,
"resolution": record.resolution,
"provider_generation_resolution": record.provider_generation_resolution,
"image_size": record.image_size,
"image_proportion": record.image_proportion,
"image_px": record.image_px,
"include_media_references": bool(record.include_media_references),
"reference_count": len(references),
"video_upscale_enabled": bool(record.video_upscale_enabled_snapshot),
"config_complete": is_generation_record_config_complete(record),
}
if before:
detail["before"] = before
if extra:
detail.update(extra)
return detail
def _record_config_before(record: GenerationRecord) -> dict[str, Any]:
return {
"engine_id": record.engine_id,
"has_engine_snapshot": bool(record.engine_snapshot_json),
"duration": record.duration,
"aspect_ratio": record.aspect_ratio,
"resolution": record.resolution,
"provider_generation_resolution": record.provider_generation_resolution,
"image_size": record.image_size,
"image_proportion": record.image_proportion,
"image_px": record.image_px,
"include_media_references": bool(record.include_media_references),
}
def log_generation_record_config_event(
*,
event_type: GenerationRecordEventTypeEnum,
event_status: LogEventStatusEnum = LogEventStatusEnum.SUCCESS,
source: GenerationRecordConfigSourceEnum | str,
record: GenerationRecord,
message: str | None = None,
detail: dict[str, Any] | None = None,
error: str | None = None,
) -> None:
log_operation_event(
domain=_GENERATION_RECORD_LOG_DOMAIN,
module=_GENERATION_RECORD_LOG_MODULE,
event_type=event_type.value,
event_status=event_status.value,
source=str(source.value if isinstance(source, GenerationRecordConfigSourceEnum) else source),
user_id=str(record.user_id) if record.user_id else None,
project_id=str(record.project_id) if record.project_id else None,
task_id=str(record.id) if record.id else None,
message=message,
detail=detail,
error=error,
)
def freeze_generation_record_config_with_log(
record: GenerationRecord,
*,
engine: ImageEngine | VideoEngine | SimpleNamespace,
source: GenerationRecordConfigSourceEnum,
) -> None:
before = _record_config_before(record)
log_generation_record_config_event(
event_type=GenerationRecordEventTypeEnum.PROMPT_CONFIG_FREEZE_START,
event_status=LogEventStatusEnum.STARTED,
source=source,
record=record,
detail=_config_log_detail(record, source=source.value, before=before),
)
try:
freeze_generation_record_config(record, engine=_engine_plain_namespace(engine))
except Exception as exc:
log_operation_error(
domain=_GENERATION_RECORD_LOG_DOMAIN,
event_type=GenerationRecordEventTypeEnum.PROMPT_CONFIG_FREEZE_FAILED.value,
module=_GENERATION_RECORD_LOG_MODULE,
source=source.value,
user_id=str(record.user_id) if record.user_id else None,
project_id=str(record.project_id) if record.project_id else None,
task_id=str(record.id) if record.id else None,
detail=_config_log_detail(record, source=source.value, before=before),
exc=exc,
)
raise
log_generation_record_config_event(
event_type=GenerationRecordEventTypeEnum.PROMPT_CONFIG_FREEZE_SUCCESS,
event_status=LogEventStatusEnum.SUCCESS,
source=source,
record=record,
detail=_config_log_detail(
record,
source=source.value,
before=before,
extra={"config_changed": before != _record_config_before(record)},
),
)
async def ensure_generation_record_config_frozen(
db: AsyncSession,
record: GenerationRecord,
*,
source: GenerationRecordConfigSourceEnum = GenerationRecordConfigSourceEnum.LEGACY_GENERATE_FALLBACK,
) -> bool:
"""Ensure one GenerationRecord has a complete frozen config.
New records should already be complete and are left untouched. Legacy
prompt_optimized rows may be missing engine_id, engine_snapshot_json or
selected parameters; those are completed server-side without accepting any
generate-time user input.
Returns True when the record was changed.
"""
if is_generation_record_config_complete(record):
log_generation_record_config_event(
event_type=GenerationRecordEventTypeEnum.LEGACY_CONFIG_FALLBACK_SKIPPED,
event_status=LogEventStatusEnum.SKIPPED,
source=GenerationRecordConfigSourceEnum.EXISTING_FROZEN_CONFIG,
record=record,
detail=_config_log_detail(record, source=GenerationRecordConfigSourceEnum.EXISTING_FROZEN_CONFIG.value),
)
return False
if record.gen_type not in {GenerationType.video.value, GenerationType.image.value}:
raise InvalidStatusError("不支持的生成类型,无法补齐历史生成配置")
before = _record_config_before(record)
log_generation_record_config_event(
event_type=GenerationRecordEventTypeEnum.LEGACY_CONFIG_FALLBACK_START,
event_status=LogEventStatusEnum.STARTED,
source=source,
record=record,
detail=_config_log_detail(record, source=source.value, before=before),
)
try:
engine_selected_by = "priority_fallback"
if record.gen_type == GenerationType.video.value:
engines = await _list_active_video_engines(db)
engine, engine_selected_by = _select_video_engine(engines, record)
_normalize_video_record_params(record, engine)
provider_resolution, upscale_enabled, upscale_snapshot_json = await build_video_upscale_snapshot(
db,
target_resolution=record.resolution or VIDEO_DEFAULT_RESOLUTION,
aspect_ratio=record.aspect_ratio or VIDEO_DEFAULT_RATIO,
supported_provider_resolutions=parse_json_list(engine.supported_resolutions, []),
)
record.provider_generation_resolution = provider_resolution
record.video_upscale_enabled_snapshot = upscale_enabled
record.video_upscale_snapshot_json = upscale_snapshot_json
else:
engines = await _list_active_image_engines(db)
engine, engine_selected_by = _select_image_engine(engines, record)
_normalize_image_record_params(record, engine)
record.provider_generation_resolution = None
record.video_upscale_enabled_snapshot = False
record.video_upscale_snapshot_json = None
# Historical rows had no explicit resource attachment switch. Missing
# values must stay false to avoid silently changing provider input and
# billing semantics.
record.include_media_references = bool(record.include_media_references)
freeze_generation_record_config(record, engine=engine)
if not is_generation_record_config_complete(record):
raise InvalidStatusError("历史生成记录配置自动补齐失败,请重新生成提词")
log_generation_record_config_event(
event_type=GenerationRecordEventTypeEnum.LEGACY_CONFIG_FALLBACK_SUCCESS,
event_status=LogEventStatusEnum.SUCCESS,
source=source,
record=record,
detail=_config_log_detail(
record,
source=source.value,
engine_selected_by=engine_selected_by,
before=before,
extra={"config_changed": before != _record_config_before(record)},
),
)
return before != _record_config_before(record)
except HTTPException as exc:
log_generation_record_config_event(
event_type=GenerationRecordEventTypeEnum.LEGACY_CONFIG_FALLBACK_FAILED,
event_status=LogEventStatusEnum.FAILED,
source=source,
record=record,
detail=_config_log_detail(record, source=source.value, before=before),
error=str(exc.detail),
)
raise
except Exception as exc:
log_operation_error(
domain=_GENERATION_RECORD_LOG_DOMAIN,
event_type=GenerationRecordEventTypeEnum.LEGACY_CONFIG_FALLBACK_FAILED.value,
module=_GENERATION_RECORD_LOG_MODULE,
source=source.value,
user_id=str(record.user_id) if record.user_id else None,
project_id=str(record.project_id) if record.project_id else None,
task_id=str(record.id) if record.id else None,
detail=_config_log_detail(record, source=source.value, before=before),
exc=exc,
)
raise
@@ -18,29 +18,48 @@ def _json(data: dict) -> str:
return json.dumps(data, ensure_ascii=False, default=str)
def prepare_generation_record_execution(
def freeze_generation_record_config(
record: GenerationRecord,
*,
engine: ImageEngine | VideoEngine,
attempt_no: int,
) -> None:
now = datetime.now(timezone.utc)
reset_execution_fields(record, started_at=now, attempt_no=attempt_no)
"""Freeze the provider capability and user-selected parameters at prompt time.
Runtime API keys are intentionally not stored in the snapshot. Provider execution
reads only the current secret from the engine row while all capability and selected
parameters continue to come from this immutable snapshot.
"""
record.engine_id = engine.id
if record.gen_type == "image":
record.engine_snapshot_json = _json(build_image_snapshot(
engine,
record.image_size or getattr(engine, "default_size", "2K") or "2K",
record.image_proportion or "1:1",
record.image_px or "2048x2048",
))
record.engine_snapshot_json = _json(
build_image_snapshot(
engine,
record.image_size or getattr(engine, "default_size", "2K") or "2K",
record.image_proportion or "1:1",
record.image_px or "2048x2048",
)
)
else:
record.engine_snapshot_json = _json(build_video_snapshot(
engine,
record.aspect_ratio or "16:9",
record.resolution or "480p",
int(record.duration or 4),
))
record.engine_snapshot_json = _json(
build_video_snapshot(
engine,
record.aspect_ratio or "16:9",
record.resolution or "480p",
int(record.duration or 4),
)
)
def prepare_generation_record_execution(
record: GenerationRecord,
*,
attempt_no: int,
) -> None:
"""Reset execution-only fields without changing the frozen prompt configuration."""
if not record.engine_id or not record.engine_snapshot_json:
raise ValueError("生成记录缺少冻结的引擎配置")
now = datetime.now(timezone.utc)
reset_execution_fields(record, started_at=now, attempt_no=attempt_no)
record.status = GenerationStatus.generating.value
record.pipeline_stage = GenerationRecordPipelineStage.QUEUED.value
@@ -51,9 +70,18 @@ async def commit_and_enqueue_generation_record(
*,
reason: str,
) -> None:
record_id = str(record.id)
attempt_no = int(record.generation_attempt_no or 1)
await db.commit()
try:
await enqueue_generation_create(record, reason=reason)
await enqueue_generation_create(
None,
reason=reason,
owner_type="generation_record",
owner_id=record_id,
generation_attempt_no=attempt_no,
generation_mode="generation_record",
)
except Exception:
# queued stage and all execution metadata are already committed; recovery will retry.
# Queued stage and execution metadata are committed; recovery will retry.
return
@@ -24,6 +24,7 @@ class GenerationRecordRecoveryBatch:
create: list[GenerationOwnerRef]
poll: list[GenerationOwnerRef]
download: list[GenerationOwnerRef]
inconsistent: list[GenerationOwnerRef]
next_cursor: GenerationRecordRecoveryCursor | None
@@ -44,6 +45,7 @@ async def find_generation_record_recovery_batch(
GenerationRecordPipelineStage.DOWNLOAD_QUEUED.value,
GenerationRecordPipelineStage.DOWNLOADING.value,
GenerationRecordPipelineStage.RETRY_WAITING.value,
GenerationRecordPipelineStage.RECOVERY_INCONSISTENT.value,
}
page_size = max(1, int(limit))
now = datetime.now(timezone.utc)
@@ -77,6 +79,18 @@ async def find_generation_record_recovery_batch(
create: list[GenerationOwnerRef] = []
poll: list[GenerationOwnerRef] = []
download: list[GenerationOwnerRef] = []
inconsistent: list[GenerationOwnerRef] = []
create_stages = {
GenerationRecordPipelineStage.QUEUED.value,
GenerationRecordPipelineStage.PREPARING.value,
GenerationRecordPipelineStage.CREATING_PROVIDER_TASK.value,
}
inconsistent_stages = {
GenerationRecordPipelineStage.WAITING_REMOTE.value,
GenerationRecordPipelineStage.POLLING.value,
GenerationRecordPipelineStage.RESULT_READY.value,
GenerationRecordPipelineStage.RECOVERY_INCONSISTENT.value,
}
for (
owner_id,
attempt_no,
@@ -97,7 +111,14 @@ async def find_generation_record_recovery_batch(
)
stage = str(pipeline_stage or "")
if str(remote_result_url or "").strip():
if stage == GenerationRecordPipelineStage.RESULT_READY.value:
if stage not in {
GenerationRecordPipelineStage.DOWNLOAD_QUEUED.value,
GenerationRecordPipelineStage.DOWNLOADING.value,
GenerationRecordPipelineStage.RETRY_WAITING.value,
}:
# The remote result URL is stronger recovery evidence than the
# persisted stage. Always continue from download instead of
# recreating or polling the provider task.
download.append(ref)
elif stage == GenerationRecordPipelineStage.DOWNLOAD_QUEUED.value:
checked_enqueued_at = ensure_aware_utc(download_enqueued_at)
@@ -130,7 +151,9 @@ async def find_generation_record_recovery_batch(
):
poll.append(ref)
else:
if (
if stage in inconsistent_stages:
inconsistent.append(ref)
elif stage in create_stages and (
ensure_aware_utc(provider_create_lease_until) is None
or ensure_aware_utc(provider_create_lease_until) <= now
):
@@ -144,5 +167,6 @@ async def find_generation_record_recovery_batch(
create=create,
poll=poll,
download=download,
inconsistent=inconsistent,
next_cursor=next_cursor,
)
@@ -1,9 +1,8 @@
from __future__ import annotations
import json
import mimetypes
import os
import time
from types import SimpleNamespace
from typing import Any
import httpx
@@ -91,9 +90,32 @@ async def _get_model_config(db: AsyncSession) -> ModelConfig:
async def build_prompt_with_chatapi(db: AsyncSession, record: ChatGenerationTask) -> tuple[str, dict]:
"""Call ChatAPI once with current request params and attachments. No history context."""
config = await _get_model_config(db)
if config.provider == "mock":
return record.original_prompt, {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}
config_row = await _get_model_config(db)
if config_row.provider == "mock":
original_prompt = str(record.original_prompt or "")
await db.commit()
return original_prompt, {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}
user_content = await _build_user_content(record, db)
config = SimpleNamespace(
id=str(config_row.id),
name=str(config_row.name or ""),
provider=str(config_row.provider or ""),
api_base=str(config_row.api_base or ""),
api_key=str(config_row.api_key or ""),
model_name=str(config_row.model_name or ""),
max_tokens=config_row.max_tokens,
temperature=config_row.temperature,
)
record = SimpleNamespace(
id=str(record.id),
user_id=str(record.user_id),
engine_id=str(record.engine_id or "") or None,
generation_mode=str(record.generation_mode or ""),
generation_attempt_no=int(record.generation_attempt_no or 1),
)
# Release all configuration/media lookup reads before the remote request.
await db.commit()
system_prompt = (
"你是图片/视频生成提示词整理助手。你的职责是根据用户文字、上传图片/视频和生成参数,"
@@ -104,14 +126,26 @@ async def build_prompt_with_chatapi(db: AsyncSession, record: ChatGenerationTask
"model": config.model_name,
"messages": [
{"role": "system", "content": system_prompt},
{"role": "user", "content": await _build_user_content(record, db)},
{"role": "user", "content": user_content},
],
"max_tokens": config.max_tokens,
"temperature": config.temperature,
}
started = time.perf_counter()
call_id = await log_provider_call(
record,
provider=config.provider,
api_type="chat_prompt",
model=config.model_name,
engine_id=record.engine_id,
status="request",
request_data=request_data,
module="generation_record",
step_code="prompt_optimize",
)
async with provider_limit("ark_chat_prompt", settings.ARK_CHAT_PROMPT_MAX_CONCURRENCY):
async with httpx.AsyncClient(timeout=settings.CHATAPI_REQUEST_TIMEOUT_SECONDS) as client:
response: httpx.Response | None = None
try:
response = await client.post(
f"{config.api_base.rstrip('/')}/chat/completions",
@@ -123,6 +157,7 @@ async def build_prompt_with_chatapi(db: AsyncSession, record: ChatGenerationTask
)
latency_ms = int((time.perf_counter() - started) * 1000)
if response.status_code >= 400:
message = response.text[:1000]
await log_provider_call(
record,
provider=config.provider,
@@ -132,13 +167,17 @@ async def build_prompt_with_chatapi(db: AsyncSession, record: ChatGenerationTask
status="failed",
latency_ms=latency_ms,
http_status=response.status_code,
request_data=request_data,
response_data=response.text,
error_message=response.text[:1000],
error_message=message,
call_id=call_id,
module="generation_record",
step_code="prompt_optimize",
)
raise RuntimeError(f"ChatAPI HTTP {response.status_code}: {response.text}")
raise RuntimeError(f"ChatAPI HTTP {response.status_code}: {message}")
data = response.json()
except Exception as exc:
if isinstance(exc, RuntimeError) and str(exc).startswith("ChatAPI HTTP "):
raise
latency_ms = int((time.perf_counter() - started) * 1000)
await log_provider_call(
record,
@@ -148,9 +187,12 @@ async def build_prompt_with_chatapi(db: AsyncSession, record: ChatGenerationTask
engine_id=record.engine_id,
status="failed",
latency_ms=latency_ms,
request_data=request_data,
response_data=None,
http_status=response.status_code if response is not None else None,
response_data=response.text if response is not None else None,
error_message=str(exc),
call_id=call_id,
module="generation_record",
step_code="prompt_optimize",
)
raise
@@ -158,23 +200,6 @@ async def build_prompt_with_chatapi(db: AsyncSession, record: ChatGenerationTask
input_tokens = int(usage.get("prompt_tokens", 0) or 0)
output_tokens = int(usage.get("completion_tokens", 0) or 0)
total_tokens = int(usage.get("total_tokens", input_tokens + output_tokens) or 0)
content = data.get("choices", [{}])[0].get("message", {}).get("content", "").strip()
if not content:
raise RuntimeError("ChatAPI未返回有效prompt")
token_usage_id = generate_id()
db.add(TokenUsage(
id=token_usage_id,
model_config_id=config.id,
user_id=record.user_id,
owner_type="generation_record",
owner_id=record.id,
input_tokens=input_tokens,
output_tokens=output_tokens,
total_tokens=total_tokens,
))
await db.flush()
await log_provider_call(
record,
provider=config.provider,
@@ -183,13 +208,59 @@ async def build_prompt_with_chatapi(db: AsyncSession, record: ChatGenerationTask
engine_id=record.engine_id,
status="success",
latency_ms=int((time.perf_counter() - started) * 1000),
http_status=200,
request_data=request_data,
http_status=response.status_code if response is not None else 200,
response_data=data,
prompt_tokens=input_tokens,
completion_tokens=output_tokens,
total_tokens=total_tokens,
call_id=call_id,
module="generation_record",
step_code="prompt_optimize",
)
content = data.get("choices", [{}])[0].get("message", {}).get("content", "").strip()
if not content:
await log_provider_call(
record,
provider=config.provider,
api_type="chat_prompt",
model=config.model_name,
engine_id=record.engine_id,
status="failed",
error_message="ChatAPI未返回有效prompt",
call_id=call_id,
module="generation_record",
step_code="prompt_optimize",
)
raise RuntimeError("ChatAPI未返回有效prompt")
token_usage_id = generate_id()
try:
db.add(TokenUsage(
id=token_usage_id,
model_config_id=config.id,
user_id=record.user_id,
owner_type="generation_record",
owner_id=record.id,
input_tokens=input_tokens,
output_tokens=output_tokens,
total_tokens=total_tokens,
))
await db.flush()
except Exception as exc:
await log_provider_call(
record,
provider=config.provider,
api_type="chat_prompt",
model=config.model_name,
engine_id=record.engine_id,
status="failed",
error_message=f"token usage写入失败: {exc}",
call_id=call_id,
module="generation_record",
step_code="prompt_optimize",
)
raise
return content, {
"token_usage_id": token_usage_id,
"model_config_id": config.id,
@@ -18,10 +18,11 @@ from app.services.generation.pipeline.owner_service import (
from app.models.image_engine import ImageEngine
from app.models.video_engine import VideoEngine
from app.services.generation.log_service import log_provider_call
from app.services.image_gen import ImageProviderError, poll_image_task_status, submit_image_task
from app.services.image_gen import poll_image_task_status, submit_image_task
from app.services.provider_limit import provider_limit
from app.services.video_gen import poll_task_status, submit_video_task
from app.types.generation.provider import ImageProviderBatchResult
from app.utils.id_gen import generate_id
def _loads(data: str | None) -> dict:
@@ -43,6 +44,17 @@ def _try_json(value: Any) -> Any:
return None
def _snapshot_owner(task: GenerationOwner) -> SimpleNamespace:
"""Copy loaded scalar fields before commit closes the current transaction."""
values = {
key: value
for key, value in vars(task).items()
if key != "_sa_instance_state"
}
values.setdefault("generation_mode", getattr(task, "generation_mode", None) or "generation_record")
return SimpleNamespace(**values)
async def get_runtime_engine(db: AsyncSession, task: GenerationOwner) -> Any:
"""使用任务快照冻结历史参数,只从当前引擎记录读取密钥。"""
snapshot = _loads(task.engine_snapshot_json)
@@ -103,37 +115,18 @@ async def create_provider_task(db: AsyncSession, task: GenerationOwner) -> dict:
async def _create_video_task(db: AsyncSession, task: GenerationOwner) -> dict:
engine = await get_runtime_engine(db, task)
task_snapshot = _snapshot_owner(task)
include_references = owner_include_media_references(task_snapshot)
# Close the engine lookup transaction before the long provider HTTP call.
await db.commit()
started = time.perf_counter()
async with provider_limit("ark_video_create", settings.ARK_VIDEO_CREATE_MAX_CONCURRENCY):
try:
provider_task_id = await submit_video_task(None, engine, task, include_media_references=owner_include_media_references(task))
response = {"task_id": provider_task_id}
await log_provider_call(
task,
provider=engine.provider,
api_type="video_create",
model=engine.model_name,
engine_id=task.engine_id,
status="success",
latency_ms=int((time.perf_counter() - started) * 1000),
provider_task_id=provider_task_id,
response_data=response,
)
return {"task_id": provider_task_id, "response_data": response}
except Exception as exc:
await log_provider_call(
task,
provider=engine.provider,
api_type="video_create",
model=engine.model_name,
engine_id=task.engine_id,
status="failed",
latency_ms=int((time.perf_counter() - started) * 1000),
error_message=str(exc),
)
raise
provider_task_id = await submit_video_task(
None,
engine,
task_snapshot,
include_media_references=include_references,
)
return {"task_id": provider_task_id, "response_data": {"task_id": provider_task_id}}
async def create_image_sync_batch_result(
@@ -143,10 +136,11 @@ async def create_image_sync_batch_result(
generation_count: int,
) -> ImageProviderBatchResult:
engine = await get_runtime_engine(db, task)
task_snapshot = _snapshot_owner(task)
# Do not keep a database transaction open while the synchronous provider call runs.
await db.commit()
return await create_image_sync_batch_result_with_engine(
task,
task_snapshot,
engine,
generation_count=generation_count,
)
@@ -163,46 +157,15 @@ async def create_image_sync_batch_result_with_engine(
generation_count > 1 时是一次组图 API 调用;失败后绝不退化为多次单图调用。
"""
count = max(1, int(generation_count or 1))
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):
try:
result = await asyncio.to_thread(
submit_image_task,
None,
engine,
task,
include_media_references=owner_include_media_references(task),
generation_count=count,
)
response_data = result.get("response_data") or result
await log_provider_call(
task,
provider=engine.provider,
api_type=api_type,
model=engine.model_name,
engine_id=task.engine_id,
status="success",
latency_ms=int((time.perf_counter() - started) * 1000),
provider_task_id=None,
response_data=response_data,
total_tokens=int(result.get("image_tokens", 0) or 0),
)
return result
except Exception as exc:
error_message = exc.safe_message if isinstance(exc, ImageProviderError) else str(exc)
await log_provider_call(
task,
provider=engine.provider,
api_type=api_type,
model=engine.model_name,
engine_id=task.engine_id,
status="failed",
latency_ms=int((time.perf_counter() - started) * 1000),
error_message=error_message,
response_data=exc.as_dict() if isinstance(exc, ImageProviderError) else None,
)
raise
return await asyncio.to_thread(
submit_image_task,
None,
engine,
task,
include_media_references=owner_include_media_references(task),
generation_count=count,
)
async def create_image_sync_result(db: AsyncSession, task: GenerationOwner) -> dict:
@@ -226,13 +189,59 @@ async def create_image_sync_result(db: AsyncSession, task: GenerationOwner) -> d
async def poll_provider_task(db: AsyncSession, task: GenerationOwner) -> dict:
engine = await get_runtime_engine(db, task)
task_snapshot = _snapshot_owner(task)
task_id = owner_provider_task_id(task_snapshot)
# Polling may block on the remote provider; release the lookup transaction first.
await db.commit()
task_id = owner_provider_task_id(task)
if not task_id:
raise ValueError("缺少供应商任务ID")
if task.gen_type == "video":
async with provider_limit("ark_video_poll", settings.ARK_VIDEO_POLL_MAX_CONCURRENCY):
return await poll_task_status(engine, task_id)
async with provider_limit("ark_image_poll", settings.ARK_IMAGE_POLL_MAX_CONCURRENCY):
return await poll_image_task_status(engine, task_id)
api_type = f"{task_snapshot.gen_type}_poll"
call_id = generate_id()
await log_provider_call(
task_snapshot,
provider=engine.provider,
api_type=api_type,
model=engine.model_name,
engine_id=task_snapshot.engine_id,
status="request",
provider_task_id=task_id,
request_data={"provider_task_id": task_id},
call_id=call_id,
)
started = time.perf_counter()
try:
if task_snapshot.gen_type == "video":
async with provider_limit("ark_video_poll", settings.ARK_VIDEO_POLL_MAX_CONCURRENCY):
result = await poll_task_status(engine, task_id)
else:
async with provider_limit("ark_image_poll", settings.ARK_IMAGE_POLL_MAX_CONCURRENCY):
result = await poll_image_task_status(engine, task_id)
await log_provider_call(
task_snapshot,
provider=engine.provider,
api_type=api_type,
model=engine.model_name,
engine_id=task_snapshot.engine_id,
status="success",
latency_ms=int((time.perf_counter() - started) * 1000),
provider_task_id=task_id,
response_data=_try_json(result.get("response_data")) or result,
total_tokens=int(result.get("video_tokens", 0) or result.get("image_tokens", 0) or 0),
call_id=call_id,
)
return result
except Exception as exc:
await log_provider_call(
task_snapshot,
provider=engine.provider,
api_type=api_type,
model=engine.model_name,
engine_id=task_snapshot.engine_id,
status="failed",
latency_ms=int((time.perf_counter() - started) * 1000),
provider_task_id=task_id,
error_message=str(exc),
call_id=call_id,
)
raise
@@ -95,6 +95,45 @@ async def _load_chat_task_for_update(
return owner if isinstance(owner, ChatGenerationTask) else None
def _chat_task_post_commit_snapshot(task: ChatGenerationTask) -> Any:
"""Capture fields used by Redis/Celery/logging before committing the ORM row."""
from types import SimpleNamespace
return SimpleNamespace(
id=str(task.id),
generation_attempt_no=int(task.generation_attempt_no or 1),
generation_mode=str(task.generation_mode or GenerationMode.CHATAPI_ASYNC.value),
provider_task_id=str(task.provider_task_id or "") or None,
seedance_task_id=str(task.seedance_task_id or "") or None,
gen_type=str(task.gen_type or ""),
pipeline_stage=str(task.pipeline_stage or ""),
poll_count=int(task.poll_count or 0),
poll_error_count=int(task.poll_error_count or 0),
manual_retry_count=int(task.manual_retry_count or 0),
poll_started_at=task.poll_started_at,
poll_interval_seconds=int(task.poll_interval_seconds or 0),
last_poll_at=task.last_poll_at,
next_poll_at=task.next_poll_at,
poll_lease_until=task.poll_lease_until,
deadline_at=task.deadline_at,
user_id=str(task.user_id or "") or None,
project_id=str(task.project_id or "") or None,
error_message=str(task.error_message or "") or None,
)
async def _reload_chat_task_after_commit(
db: AsyncSession, task_id: str
) -> ChatGenerationTask | None:
owner = await load_generation_owner(
db,
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
owner_id=str(task_id),
for_update=False,
)
return owner if isinstance(owner, ChatGenerationTask) else None
def _now() -> datetime:
return datetime.now(timezone.utc)
@@ -125,6 +164,7 @@ def _is_final_task_state(task: ChatGenerationTask) -> bool:
ChatGenerationPipelineStage.FAILED.value,
ChatGenerationPipelineStage.TIMEOUT.value,
ChatGenerationPipelineStage.DOWNLOAD_FAILED.value,
ChatGenerationPipelineStage.UPSCALE_FAILED.value,
)
@@ -430,11 +470,17 @@ async def _mark_timeout(
error_message=error_message,
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
)
snapshot = _chat_task_post_commit_snapshot(task)
await db.commit()
await notify_owner_finished(db, task)
await _remove_poll_active(_chat_registry_id(task))
fresh_task = await _reload_chat_task_after_commit(db, snapshot.id)
if fresh_task is not None:
await notify_owner_finished(db, fresh_task)
await _remove_poll_active(_chat_registry_id(snapshot))
await log_task_event(
task,
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
owner_id=snapshot.id,
generation_attempt_no=snapshot.generation_attempt_no,
generation_mode=snapshot.generation_mode,
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
to_status="failed",
to_stage=ChatGenerationPipelineStage.TIMEOUT.value,
@@ -456,10 +502,21 @@ async def _mark_failed(
error_message=error_message,
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
)
snapshot = _chat_task_post_commit_snapshot(task)
await db.commit()
await notify_owner_finished(db, task)
await _remove_poll_active(_chat_registry_id(task))
await log_task_event(task, event_type=event_type, message=task.error_message, detail=detail)
fresh_task = await _reload_chat_task_after_commit(db, snapshot.id)
if fresh_task is not None:
await notify_owner_finished(db, fresh_task)
await _remove_poll_active(_chat_registry_id(snapshot))
await log_task_event(
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
owner_id=snapshot.id,
generation_attempt_no=snapshot.generation_attempt_no,
generation_mode=snapshot.generation_mode,
event_type=event_type,
message=snapshot.error_message or error_message,
detail=detail,
)
return "mark_failed"
@@ -475,8 +532,8 @@ async def recover_one_generation_task(
分流原则:
1. 已有 remote_result_url:只恢复下载,不 poll,不重新 create。
2. 已有 provider_task_id/seedance_task_id:恢复 poll。
3. 无结果 URL、无供应商任务 ID:deadline 未过才恢复 create。
4. 无结果 URL、无供应商任务 ID:deadline 已过直接超时失败,不再补救生成
3. 仅 queued/preparing/creating_provider_task 且无远程证据时允许恢复 create。
4. waiting_remote/polling/result_ready 缺少对应证据时隔离,deadline 到期后失败退款
"""
from app.tasks.generation_create_tasks import chatapi_create_generation_task
from app.tasks.generation_download_tasks import enqueue_download_task
@@ -543,21 +600,25 @@ async def recover_one_generation_task(
if is_deadline_expired:
if has_provider_task_id:
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
snapshot = _chat_task_post_commit_snapshot(task)
await db.commit()
await log_task_event(
task,
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
owner_id=snapshot.id,
generation_attempt_no=snapshot.generation_attempt_no,
generation_mode=snapshot.generation_mode,
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
message=f"{source} 发现任务已到 deadline 且存在供应商任务ID,投递 poll 队列做最终查询",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
detail={"pipeline_stage": snapshot.pipeline_stage, "payload": redis_payload},
)
poll_generation_task.apply_async(
args=[task.id],
kwargs={"force_due": True, "owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
args=[snapshot.id],
kwargs={"force_due": True, "owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": snapshot.generation_attempt_no},
queue=POLL_QUEUE,
countdown=0,
)
await register_poll_active(
task,
snapshot,
check_at=_poll_queue_timeout_at(),
reason=f"{source}_deadline_final_poll",
)
@@ -575,21 +636,25 @@ async def recover_one_generation_task(
if is_video_generation_task(task):
ensure_video_poll_fields(task, now=current_time)
if is_poll_not_due(task, now=current_time):
snapshot = _chat_task_post_commit_snapshot(task)
await db.commit()
await register_poll_active(
task,
check_at=task.next_poll_at,
next_poll_at=task.next_poll_at,
snapshot,
check_at=snapshot.next_poll_at,
next_poll_at=snapshot.next_poll_at,
reason=f"{source}_video_poll_not_due",
)
await log_task_event(
task,
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
owner_id=snapshot.id,
generation_attempt_no=snapshot.generation_attempt_no,
generation_mode=snapshot.generation_mode,
event_type=ChatGenerationTaskEventType.POLL_SKIP_NOT_DUE.value,
message=f"{source} 发现视频任务尚未到下一次轮询时间,启动容灾不提前投递 poll",
detail={
"pipeline_stage": task.pipeline_stage,
"pipeline_stage": snapshot.pipeline_stage,
"payload": redis_payload,
"next_poll_at": task.next_poll_at,
"next_poll_at": snapshot.next_poll_at,
},
)
return "skip_video_poll_not_due"
@@ -600,28 +665,32 @@ async def recover_one_generation_task(
# 这里仍复用 next_poll_at 做短暂队列保护,避免启动容灾重复投递。
# 真正消费时通过 force_due=True 跳过“未到期”校验,避免保护时间反向阻塞本次 poll。
task.next_poll_at = queue_hold_until
snapshot = _chat_task_post_commit_snapshot(task)
await db.commit()
await log_task_event(
task,
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
owner_id=snapshot.id,
generation_attempt_no=snapshot.generation_attempt_no,
generation_mode=snapshot.generation_mode,
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
message=f"{source} 发现任务存在供应商任务ID,恢复投递轮询队列",
detail={
"pipeline_stage": task.pipeline_stage,
"pipeline_stage": snapshot.pipeline_stage,
"payload": redis_payload,
"due_next_poll_at": original_next_poll_at,
"queue_hold_until": queue_hold_until,
},
)
poll_generation_task.apply_async(
args=[task.id],
kwargs={"force_due": True, "owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
args=[snapshot.id],
kwargs={"force_due": True, "owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": snapshot.generation_attempt_no},
queue=POLL_QUEUE,
countdown=0,
)
await register_poll_active(
task,
check_at=task.next_poll_at,
next_poll_at=task.next_poll_at,
snapshot,
check_at=snapshot.next_poll_at,
next_poll_at=snapshot.next_poll_at,
reason=f"{source}_has_provider_task_id",
)
return "recover_poll_has_provider_id"
@@ -633,64 +702,84 @@ async def recover_one_generation_task(
ChatGenerationPipelineStage.QUEUED.value,
ChatGenerationPipelineStage.PREPARING.value,
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
ChatGenerationPipelineStage.WAITING_REMOTE.value,
ChatGenerationPipelineStage.POLLING.value,
}
if task.pipeline_stage in recoverable_create_stages:
if task.pipeline_stage not in (
ChatGenerationPipelineStage.QUEUED.value,
ChatGenerationPipelineStage.PREPARING.value,
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
):
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
# 刷新更新时间形成创建队列保护窗口,避免 Beat 在任务尚未消费时每轮重复补投。
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
task.updated_at = current_time
# Release the recovery row lock before writing an event through the
# independent logging session or talking to the broker.
task_id = str(task.id)
attempt_no = int(task.generation_attempt_no or 1)
generation_mode = str(task.generation_mode or "")
await db.commit()
await _remove_poll_active(_chat_registry_id(task))
await _remove_poll_active(
redis_owner_item_id(
GenerationOwnerType.CHAT_GENERATION_TASK.value,
task_id,
attempt_no,
)
)
await log_task_event(
task,
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
owner_id=task_id,
generation_attempt_no=attempt_no,
generation_mode=generation_mode,
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
message=f"{source} 发现任务未超时且缺少 remote_result_url/供应商任务ID,恢复投递创建队列",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
detail={
"pipeline_stage": ChatGenerationPipelineStage.QUEUED.value,
"payload": redis_payload,
},
)
chatapi_create_generation_task.apply_async(
args=[task.id],
kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
args=[task_id],
kwargs={
"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value,
"generation_attempt_no": attempt_no,
},
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
countdown=0,
task_id=(
f"generation-create:{GenerationOwnerType.CHAT_GENERATION_TASK.value}:"
f"{task.id}:attempt:{int(task.generation_attempt_no or 1)}"
f"{task_id}:attempt:{attempt_no}"
),
)
return "recover_create_no_remote_no_provider_before_deadline"
# result_ready 但没有 URL 是脏状态;未过 deadline 时回创建队列重新处理,过期上面已标记超时。
if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value:
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
task.updated_at = current_time
inconsistent_stages = {
ChatGenerationPipelineStage.WAITING_REMOTE.value,
ChatGenerationPipelineStage.POLLING.value,
ChatGenerationPipelineStage.RESULT_READY.value,
ChatGenerationPipelineStage.RECOVERY_INCONSISTENT.value,
}
if task.pipeline_stage in inconsistent_stages:
original_stage = str(task.pipeline_stage or "")
task_id = str(task.id)
attempt_no = int(task.generation_attempt_no or 1)
generation_mode = str(task.generation_mode or "")
task.pipeline_stage = ChatGenerationPipelineStage.RECOVERY_INCONSISTENT.value
task.error_message = (
f"{source} 恢复证据异常:阶段 {original_stage} 缺少 remote_result_url 和供应商任务ID"
)
await db.commit()
await _remove_poll_active(_chat_registry_id(task))
await _remove_poll_active(
redis_owner_item_id(
GenerationOwnerType.CHAT_GENERATION_TASK.value,
task_id,
attempt_no,
)
)
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 result_ready 但缺少 remote_result_url,未超时,恢复投递创建队列",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
owner_id=task_id,
generation_attempt_no=attempt_no,
generation_mode=generation_mode,
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_INCONSISTENT.value,
from_stage=original_stage,
to_stage=ChatGenerationPipelineStage.RECOVERY_INCONSISTENT.value,
message=f"{source} 发现恢复证据异常,已隔离且不重新创建供应商任务",
detail={"payload": redis_payload},
)
chatapi_create_generation_task.apply_async(
args=[task.id],
kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
countdown=0,
task_id=(
f"generation-create:{GenerationOwnerType.CHAT_GENERATION_TASK.value}:"
f"{task.id}:attempt:{int(task.generation_attempt_no or 1)}"
),
)
return "recover_create_result_ready_no_url_before_deadline"
return "quarantine_inconsistent_recovery_evidence"
return f"skip_stage_{task.pipeline_stage}"
@@ -919,6 +1008,7 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
"waiting_remote",
"polling",
"result_ready",
"recovery_inconsistent",
]
),
)
@@ -971,6 +971,13 @@ async def run_image_prompt_optimize(
user_id=user_id_value,
references=references,
gen_type="image",
log_module=module_value,
log_step="hot_opening_image_prompt_optimize",
log_project_id=project_id_value,
log_task_id=step_id_value,
log_owner_type="module_generation_step",
log_owner_id=step_id_value,
generation_attempt_no=expected_step_version,
)
if execution_guard is not None:
await execution_guard()
@@ -3,6 +3,7 @@ from __future__ import annotations
import copy
import json
import re
from types import SimpleNamespace
from typing import Any
import httpx
@@ -15,7 +16,6 @@ from app.enums.hot_opening_replicate import HotOpeningLogEventEnum, HotOpeningRe
from app.enums.shot_replicate import ModuleCodeEnum as ShotModuleCodeEnum, ShotReplicateLogEventEnum, ShotReplicateRemoteActionEnum
from app.services.operation_log_service import log_ai_model_event
from app.enums.common import (
VIDEO_SCHEMA_CONFIG_DATABASE_SOURCE,
VIDEO_SCHEMA_CONFIG_DEFAULT_SOURCE,
VIDEO_SCHEMA_CONFIG_VERSION,
VIDEO_SCHEMA_EDITABLE_TEXT_MAX_LEN,
@@ -1477,6 +1477,7 @@ def _log_video_prompt_ai_event(
event_status: str,
config: ModelConfig,
trace_id: str | None,
call_id: str,
user_id: str | None,
project_id: str | None,
step_id: str | None,
@@ -1489,30 +1490,61 @@ def _log_video_prompt_ai_event(
error: str | None = None,
detail: dict[str, Any] | None = None,
) -> None:
log_ai_model_event(
event_type=event_type,
event_status=event_status,
source=LogSourceEnum.REMOTE_API.value,
module=module,
trace_id=trace_id,
user_id=user_id,
project_id=project_id,
step_id=step_id,
remote_action=action,
remote_request_id=remote_request_id,
model_config_id=str(config.id),
model_config_name=config.name,
model_name=config.model_name,
provider=config.provider,
api_base=config.api_base,
http_status=http_status,
request=request_data,
response=response_data,
token_usage=token_usage,
message=message,
detail=detail,
error=error,
)
common = {
"source": LogSourceEnum.REMOTE_API.value,
"module": module,
"step_code": "video_prompt_generate",
"call_id": call_id,
"trace_id": trace_id,
"user_id": user_id,
"project_id": project_id,
"task_id": step_id,
"step_id": step_id,
"owner_type": "module_generation_step",
"owner_id": step_id or project_id,
"remote_action": action,
"remote_request_id": remote_request_id,
"model_config_id": str(config.id),
"model_config_name": config.name,
"model_name": config.model_name,
"provider": config.provider,
"api_base": config.api_base,
"http_status": http_status,
}
normalized_status = str(event_status or "").lower()
if normalized_status == str(LogEventStatusEnum.STARTED.value).lower():
log_ai_model_event(
event_type=event_type,
event_phase="REQUEST",
event_status=event_status,
request=request_data,
message=message,
detail=detail,
**common,
)
return
if response_data is not None:
log_ai_model_event(
event_type=event_type,
event_phase="RESPONSE",
event_status=event_status,
response=response_data,
token_usage=token_usage,
message=message,
detail=detail,
error=error if normalized_status == str(LogEventStatusEnum.FAILED.value).lower() else None,
**common,
)
if normalized_status == str(LogEventStatusEnum.FAILED.value).lower() or error:
log_ai_model_event(
event_type=event_type,
event_phase="ERROR",
event_status=LogEventStatusEnum.FAILED.value,
message=message,
detail=detail,
error=error or "AI model call failed",
**common,
)
async def _select_model_config(db: AsyncSession) -> ModelConfig | None:
result = await db.execute(select(ModelConfig).where(ModelConfig.is_active == True, ModelConfig.deleted_at.is_(None)).order_by(ModelConfig.priority.desc()).limit(1))
@@ -1536,6 +1568,7 @@ async def optimize_hot_opening_video_prompt(
step_id: str | None = None,
trace_id: str | None = None,
) -> tuple[dict[str, Any], str, dict[str, Any]]:
call_id = generate_id()
duration = int(video_config["duration"])
from app.utils.media import media_to_base64, get_llm_media_as_base64
use_base64 = await get_llm_media_as_base64(db)
@@ -1559,9 +1592,22 @@ async def optimize_hot_opening_video_prompt(
# result = ensure_negative_prompt(ensure_flow_matches_time_plan(ensure_top_keys(fill_none_with_wu(result)), duration))
# return result, build_final_video_prompt(result), {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}
config = await _select_model_config(db)
config_row = await _select_model_config(db)
config = (
SimpleNamespace(
id=str(config_row.id),
name=str(config_row.name or ""),
provider=str(config_row.provider or ""),
api_base=str(config_row.api_base or ""),
api_key=str(config_row.api_key or ""),
model_name=str(config_row.model_name or ""),
)
if config_row is not None
else None
)
# All module/project claims are committed by the caller. Release this
# configuration read transaction before the remote model request.
# configuration read transaction before the remote model request and use
# only the scalar snapshot afterwards.
await db.commit()
if not config:
result = normalize_video_prompt_schema_from_ai(_mock_result(video_config, target_platform), video_config, schema_config_snapshot)
@@ -1602,6 +1648,7 @@ async def optimize_hot_opening_video_prompt(
}
started_event, remote_action = _video_prompt_remote_event(module, started=True)
_log_video_prompt_ai_event(
call_id=call_id,
module=module,
event_type=started_event,
action=remote_action,
@@ -1625,6 +1672,7 @@ async def optimize_hot_opening_video_prompt(
except Exception as exc:
failed_event, remote_action = _video_prompt_remote_event(module)
_log_video_prompt_ai_event(
call_id=call_id,
module=module,
event_type=failed_event,
action=remote_action,
@@ -1645,6 +1693,7 @@ async def optimize_hot_opening_video_prompt(
if response.status_code >= 400:
failed_event, remote_action = _video_prompt_remote_event(module)
_log_video_prompt_ai_event(
call_id=call_id,
module=module,
event_type=failed_event,
action=remote_action,
@@ -1671,6 +1720,7 @@ async def optimize_hot_opening_video_prompt(
except Exception as exc:
parse_event, remote_action = _video_prompt_remote_event(module, empty="content 为空" in str(exc), parse_failed="content 为空" not in str(exc))
_log_video_prompt_ai_event(
call_id=call_id,
module=module,
event_type=parse_event,
action=remote_action,
@@ -1721,6 +1771,7 @@ async def optimize_hot_opening_video_prompt(
except Exception as exc:
parse_event, remote_action = _video_prompt_remote_event(module, parse_failed=True)
_log_video_prompt_ai_event(
call_id=call_id,
module=module,
event_type=parse_event,
action=remote_action,
@@ -1740,6 +1791,7 @@ async def optimize_hot_opening_video_prompt(
raise
success_event, remote_action = _video_prompt_remote_event(module, success=True)
_log_video_prompt_ai_event(
call_id=call_id,
module=module,
event_type=success_event,
action=remote_action,
+57 -47
View File
@@ -1,7 +1,7 @@
import json
import logging
import os
from datetime import datetime
import time
from typing import Any
import httpx
@@ -16,7 +16,8 @@ from app.enums.generation_provider import (
)
from app.enums.private_portrait import PRIVATE_PORTRAIT_ASSET_URI_PREFIX
from app.models.image_engine import ImageEngine
from app.services.log_config import LOG_DATE_FORMAT, LOG_DIR, encrypt_data, is_enabled
from app.services.operation_log_service import build_exception_detail, log_ai_model_event
from app.utils.id_gen import generate_id
from app.types.generation.provider import (
ImageProviderBatchResult,
ImageProviderItem,
@@ -59,49 +60,27 @@ class ImageProviderError(RuntimeError):
}
def _log_image_request(engine: ProviderImageEngineLike, record_id: str, request_data: dict):
if not is_enabled():
return
try:
os.makedirs(LOG_DIR, exist_ok=True)
today = datetime.now().strftime(LOG_DATE_FORMAT)
log_file = os.path.join(LOG_DIR, f"{today}.log")
request_str = json.dumps(request_data, ensure_ascii=False)
request_encrypted = encrypt_data(request_data, True)
entry = {
"timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
"type": "image_gen_request",
"engine": engine.name,
"model": engine.model_name,
"record_id": record_id,
"request": request_encrypted,
"request_length": len(request_str),
}
with open(log_file, "a", encoding="utf-8") as file:
file.write(json.dumps(entry, ensure_ascii=False) + "\n")
except Exception:
pass
def _log_image_response(record_id: str, response_data: dict, error: str | None = None):
if not is_enabled():
return
try:
os.makedirs(LOG_DIR, exist_ok=True)
today = datetime.now().strftime(LOG_DATE_FORMAT)
log_file = os.path.join(LOG_DIR, f"{today}.log")
response_encrypted = encrypt_data(response_data, True) if response_data else ""
entry = {
"timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
"type": "image_gen_response",
"record_id": record_id,
"response": response_encrypted,
"error": error,
}
with open(log_file, "a", encoding="utf-8") as file:
file.write(json.dumps(entry, ensure_ascii=False) + "\n")
except Exception:
pass
def _provider_log_context(engine, record, *, call_id: str, step_code: str) -> dict:
generation_mode = str(getattr(record, "generation_mode", "") or "generation_record")
owner_type = "chat_generation_task" if generation_mode != "generation_record" else "generation_record"
return {
"module": generation_mode,
"step_code": step_code,
"call_id": call_id,
"source": "app.services.image_gen",
"user_id": str(getattr(record, "user_id", "") or "") or None,
"project_id": str(getattr(record, "project_id", "") or "") or None,
"task_id": str(getattr(record, "id", "") or "") or None,
"owner_type": owner_type,
"owner_id": str(getattr(record, "id", "") or "") or None,
"generation_attempt_no": int(getattr(record, "generation_attempt_no", 1) or 1),
"model_config_id": str(getattr(engine, "id", "") or "") or None,
"model_config_name": str(getattr(engine, "name", "") or "") or None,
"model_name": str(getattr(engine, "model_name", "") or "") or None,
"provider": str(getattr(engine, "provider", "") or "") or None,
"api_base": str(getattr(engine, "api_base", "") or "") or None,
}
async def get_active_image_engine(db: AsyncSession) -> ImageEngine:
@@ -321,7 +300,18 @@ def submit_image_task(
)
request_sdk_payload["stream"] = False
_log_image_request(engine, record.id, request_log_payload)
call_id = generate_id()
started = time.perf_counter()
api_step = "image_sync_batch_create" if count > 1 else "image_sync_create"
log_context = _provider_log_context(engine, record, call_id=call_id, step_code=api_step)
log_ai_model_event(
event_type="REQUEST",
event_phase="REQUEST",
event_status="started",
remote_action=api_step,
request=request_log_payload,
**log_context,
)
try:
result = client.images.generate(**request_sdk_payload)
@@ -386,7 +376,16 @@ def submit_image_task(
"total_tokens": total_tokens,
},
}
_log_image_response(record.id, response_data)
log_ai_model_event(
event_type="RESPONSE",
event_phase="RESPONSE",
event_status="success",
remote_action=api_step,
latency_ms=int((time.perf_counter() - started) * 1000),
response=response_data,
token_usage=response_data.get("usage"),
**log_context,
)
return {
"items": items,
"model": str(response_data["model"] or ""),
@@ -404,7 +403,18 @@ def submit_image_task(
provider_error.error_code,
provider_error.safe_message,
)
_log_image_response(record.id, provider_error.as_dict(), provider_error.safe_message)
log_ai_model_event(
event_type="ERROR",
event_phase="ERROR",
event_status="failed",
remote_action=api_step,
http_status=provider_error.http_status,
remote_request_id=provider_error.provider_request_id,
latency_ms=int((time.perf_counter() - started) * 1000),
detail=build_exception_detail(exc, provider_error.as_dict()),
error=provider_error.safe_message,
**log_context,
)
raise provider_error from exc
finally:
try:
+170 -79
View File
@@ -1,6 +1,6 @@
import json
import os
from datetime import datetime
import time
from types import SimpleNamespace
import httpx
from sqlalchemy import select
@@ -10,47 +10,13 @@ from app.config import settings
from app.models.model_config import ModelConfig
from app.models.token_usage import TokenUsage
from app.utils.id_gen import generate_id
from app.services.log_config import is_enabled, LOG_DIR, LOG_DATE_FORMAT, encrypt_data
from app.services.operation_log_service import build_exception_detail, log_ai_model_event
def _sanitize_for_log(data):
"""Replace base64 data URIs with placeholder for readable logs."""
if isinstance(data, str):
if data.startswith("data:") and ";base64," in data:
return "[base64 image data]"
return data
if isinstance(data, dict):
return {k: _sanitize_for_log(v) for k, v in data.items()}
if isinstance(data, list):
return [_sanitize_for_log(item) for item in data]
return data
def _log_ai_request_response(config, request_data: dict, response_data: dict | None, error: str | None = None):
"""Log AI model request/response to log/AiModel/YYYY-MM-DD.log"""
if not is_enabled():
return
try:
os.makedirs(LOG_DIR, exist_ok=True)
today = datetime.now().strftime(LOG_DATE_FORMAT)
log_file = os.path.join(LOG_DIR, f"{today}.log")
request_encrypted = encrypt_data(_sanitize_for_log(request_data), True)
response_encrypted = encrypt_data(_sanitize_for_log(response_data), True) if response_data else ""
entry = {
"timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
"model_name": config.name,
"model_id": config.model_name,
"provider": config.provider,
"api_base": config.api_base,
"request": request_encrypted,
"response": response_encrypted,
"error": error,
}
with open(log_file, "a", encoding="utf-8") as f:
f.write(json.dumps(entry, ensure_ascii=False) + "\n")
except Exception:
pass
class LLMProviderCallError(RuntimeError):
"""Remote model call or response validation failed and may use fallback."""
MOCK_OPTIMIZED_PROMPTS = {
"直播": "专业直播间场景,45度斜角机位,暖色柔光打光,主播居中构图,背景虚化处理,产品特写切换流畅,镜头推进节奏感强,画面色彩饱和度高,适合电商直播推广视频。",
@@ -92,9 +58,17 @@ async def optimize_prompt(
duration: int | None = None,
image_size: str | None = None,
image_proportion: str | None = None,
image_px: str | None | None = None,
image_px: str | None = None,
references: list[dict] | None = None,
gen_type: str = "video",
*,
log_module: str = "generation_ai",
log_step: str = "prompt_optimize",
log_project_id: str | None = None,
log_task_id: str | None = None,
log_owner_type: str | None = None,
log_owner_id: str | None = None,
generation_attempt_no: int | None = None,
) -> tuple[str, dict]:
"""Optimize user prompt using LLM. Returns (optimized_text, token_usage_dict)."""
@@ -103,9 +77,22 @@ async def optimize_prompt(
.where(ModelConfig.is_active == True, ModelConfig.deleted_at.is_(None))
.order_by(ModelConfig.priority.desc())
)
configs = list(result.scalars().all())
# Release the read transaction before the external LLM request. Callers
# must commit their business claim before invoking optimize_prompt.
configs = [
SimpleNamespace(
id=item.id,
name=item.name,
provider=item.provider,
api_base=item.api_base,
api_key=item.api_key,
model_name=item.model_name,
max_tokens=item.max_tokens,
temperature=item.temperature,
)
for item in result.scalars().all()
]
# Release the read transaction before the external LLM request. Only
# plain scalar snapshots are used afterwards, so expire_on_commit does
# not trigger an ORM refresh while the provider request is in flight.
await db.commit()
if configs:
@@ -122,8 +109,15 @@ async def optimize_prompt(
image_size=image_size,
image_proportion=image_proportion,
image_px=image_px,
log_module=log_module,
log_step=log_step,
log_project_id=log_project_id,
log_task_id=log_task_id,
log_owner_type=log_owner_type,
log_owner_id=log_owner_id,
generation_attempt_no=generation_attempt_no,
)
except Exception:
except LLMProviderCallError:
continue
# 所有真实模型都失败,降级到 mock
@@ -156,7 +150,15 @@ async def _call_openai_compatible(
gen_type: str = "video",
image_size: str | None = None,
image_proportion: str | None = None,
image_px: str | None | None = None,
image_px: str | None = None,
*,
log_module: str = "generation_ai",
log_step: str = "prompt_optimize",
log_project_id: str | None = None,
log_task_id: str | None = None,
log_owner_type: str | None = None,
log_owner_id: str | None = None,
generation_attempt_no: int | None = None,
) -> tuple[str, dict]:
"""Call an OpenAI-compatible API to optimize the prompt. Returns (content, token_usage)."""
system_prompt = None
@@ -321,14 +323,33 @@ async def _call_openai_compatible(
"max_tokens": config.max_tokens,
"temperature": config.temperature,
}
# Build log-friendly request data (image paths instead of base64)
if log_user_message:
log_request_data = {**request_data, "messages": [
{"role": "system", "content": system_prompt},
log_user_message,
]}
else:
log_request_data = request_data
call_id = generate_id()
started = time.perf_counter()
common_log = {
"module": log_module,
"step_code": log_step,
"call_id": call_id,
"source": "app.services.llm",
"user_id": user_id,
"project_id": log_project_id,
"task_id": log_task_id,
"owner_type": log_owner_type,
"owner_id": log_owner_id,
"generation_attempt_no": generation_attempt_no,
"model_config_id": config.id,
"model_config_name": config.name,
"model_name": config.model_name,
"provider": config.provider,
"api_base": config.api_base,
"remote_action": "chat_completions",
}
log_ai_model_event(
event_type="REQUEST",
event_phase="REQUEST",
event_status="started",
request=request_data,
**common_log,
)
try:
response = await client.post(
f"{config.api_base}/chat/completions",
@@ -338,47 +359,117 @@ async def _call_openai_compatible(
},
json=request_data,
)
latency_ms = int((time.perf_counter() - started) * 1000)
if response.status_code >= 400:
error_body = response.text
_log_ai_request_response(config, log_request_data, None, error=f"HTTP {response.status_code}: {error_body}")
raise RuntimeError(f"HTTP {response.status_code}: {error_body}")
log_ai_model_event(
event_type="RESPONSE",
event_phase="RESPONSE",
event_status="failed",
http_status=response.status_code,
latency_ms=latency_ms,
response={"body": error_body},
error=f"HTTP {response.status_code}",
**common_log,
)
error = LLMProviderCallError(f"HTTP {response.status_code}: {error_body}")
log_ai_model_event(
event_type="ERROR",
event_phase="ERROR",
event_status="failed",
http_status=response.status_code,
latency_ms=latency_ms,
detail=build_exception_detail(error),
error=str(error),
**common_log,
)
raise error
data = response.json()
except RuntimeError:
log_ai_model_event(
event_type="RESPONSE",
event_phase="RESPONSE",
event_status="success",
http_status=response.status_code,
latency_ms=latency_ms,
response=data,
token_usage=data.get("usage") if isinstance(data, dict) else None,
**common_log,
)
except LLMProviderCallError:
raise
except Exception as e:
_log_ai_request_response(config, log_request_data, None, error=str(e))
raise RuntimeError(f"{type(e).__name__}: {e}")
except Exception as exc:
latency_ms = int((time.perf_counter() - started) * 1000)
log_ai_model_event(
event_type="ERROR",
event_phase="ERROR",
event_status="failed",
latency_ms=latency_ms,
detail=build_exception_detail(exc),
error=str(exc),
**common_log,
)
raise LLMProviderCallError(f"{type(exc).__name__}: {exc}") from exc
# Log request/response
_log_ai_request_response(config, log_request_data, data)
# Record token usage
usage = data.get("usage", {})
input_tokens = usage.get("prompt_tokens", 0)
output_tokens = usage.get("completion_tokens", 0)
total_tokens = usage.get("total_tokens", input_tokens + output_tokens)
try:
usage = data.get("usage", {})
input_tokens = int(usage.get("prompt_tokens", 0) or 0)
output_tokens = int(usage.get("completion_tokens", 0) or 0)
total_tokens = int(usage.get("total_tokens", input_tokens + output_tokens) or 0)
content = data["choices"][0]["message"]["content"].strip()
if not content:
raise ValueError("模型未返回有效提示词")
except Exception as exc:
log_ai_model_event(
event_type="ERROR",
event_phase="ERROR",
event_status="failed",
latency_ms=int((time.perf_counter() - started) * 1000),
detail=build_exception_detail(exc, {"stage": "response_validation"}),
error=str(exc),
**common_log,
)
raise LLMProviderCallError(f"模型响应解析失败: {exc}") from exc
token_usage_id = None
if db is not None:
token_usage_id = generate_id()
record = TokenUsage(
id=token_usage_id,
model_config_id=config.id,
user_id=user_id,
input_tokens=input_tokens,
output_tokens=output_tokens,
total_tokens=total_tokens,
)
db.add(record)
await db.flush()
try:
token_usage_id = generate_id()
record = TokenUsage(
id=token_usage_id,
model_config_id=config.id,
user_id=user_id,
input_tokens=input_tokens,
output_tokens=output_tokens,
total_tokens=total_tokens,
source_module=log_module,
source_step_code=log_step,
owner_type=log_owner_type,
owner_id=log_owner_id,
)
db.add(record)
await db.flush()
except Exception as exc:
log_ai_model_event(
event_type="ERROR",
event_phase="ERROR",
event_status="failed",
latency_ms=int((time.perf_counter() - started) * 1000),
detail=build_exception_detail(exc, {"stage": "token_usage_persistence"}),
error=str(exc),
**common_log,
)
# A local transaction failure must not call a second provider after
# the first provider has already returned a valid response.
raise
content = data["choices"][0]["message"]["content"].strip()
token_usage = {
"token_usage_id": token_usage_id,
"model_config_id": config.id,
"model_config_name": config.name,
"model_provider": config.provider,
"model_name": config.model_name,
"source_module": log_module,
"source_step_code": log_step,
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"total_tokens": total_tokens,
@@ -61,8 +61,13 @@ def log_module_prompt_event(
step_id=step_id,
user_id=user_id,
trace_id=trace_id,
message=f"模块 AI 请求{prompt_type}",
detail={"prompt_type": prompt_type, "request": request or {}, "response": response or {}, "token_usage": token_usage or {}},
message=f"模块 AI 步骤{prompt_type}",
detail={
"prompt_type": prompt_type,
"has_request": request is not None,
"has_response": response is not None,
"token_usage": token_usage or {},
},
error=error,
event_status=LogEventStatusEnum.FAILED.value if error else LogEventStatusEnum.SUCCESS.value,
source=LogSourceEnum.SERVICE.value,
@@ -21,7 +21,6 @@ MODULE_GENERATION_LOG_ROOT = os.path.join(LOG_BASE_DIR, "ModuleGeneration")
AI_MODEL_LOG_ROOT = LOG_DIR
SENSITIVE_KEY_PATTERNS = (
"secret",
"token",
"authorization",
"cookie",
"credential",
@@ -30,9 +29,28 @@ SENSITIVE_KEY_PATTERNS = (
"access_key",
"api_key",
"apikey",
"security-token",
"x-tos-security-token",
"security_token",
)
SENSITIVE_TOKEN_KEYS = {
"token",
"access_token",
"refresh_token",
"bearer_token",
"security_token",
"x_tos_security_token",
}
FILE_BASE64_KEYS = {
"b64_json",
"file_data",
"file_base64",
"content_base64",
"image_base64",
"video_base64",
"audio_base64",
}
FILE_DATA_URI_MIME_PREFIXES = ("image/", "video/", "audio/")
FILE_DATA_URI_MIME_TYPES = {"application/pdf", "application/octet-stream"}
FILE_BASE64_PREVIEW_CHARS = 30
def _safe_name(value: str | None, default: str = "unknown") -> str:
@@ -49,7 +67,41 @@ def _mask_string(value: str) -> str:
def _is_sensitive_key(key: str) -> bool:
lower = str(key).replace("-", "_").lower()
return lower == "sign" or any(pattern in lower for pattern in SENSITIVE_KEY_PATTERNS)
if lower == "sign" or lower in SENSITIVE_TOKEN_KEYS:
return True
return any(pattern in lower for pattern in SENSITIVE_KEY_PATTERNS)
def _decoded_base64_size(value: str) -> int:
compact = "".join(value.split())
if not compact:
return 0
padding = 2 if compact.endswith("==") else (1 if compact.endswith("=") else 0)
return max(0, (len(compact) * 3) // 4 - padding)
def _file_base64_preview(value: str, key_path: tuple[str, ...]) -> str | None:
data_uri = re.match(r"^data:([^;,]+);base64,(.*)$", value, flags=re.IGNORECASE | re.DOTALL)
prefix = ""
payload = value
is_file = False
if data_uri:
mime = str(data_uri.group(1) or "").lower()
is_file = mime.startswith(FILE_DATA_URI_MIME_PREFIXES) or mime in FILE_DATA_URI_MIME_TYPES
prefix = value[: value.find(",") + 1]
payload = data_uri.group(2)
elif key_path and key_path[-1].replace("-", "_").lower() in FILE_BASE64_KEYS:
# Raw base64 is treated as file content only for an explicit file field.
is_file = len(value) >= 64 and bool(re.fullmatch(r"[A-Za-z0-9+/=\s]+", value))
if not is_file:
return None
compact = "".join(payload.split())
preview = compact[:FILE_BASE64_PREVIEW_CHARS]
total_bytes = _decoded_base64_size(compact)
preview_bytes = min(total_bytes, (len(preview) * 3) // 4)
remaining_bytes = max(0, total_bytes - preview_bytes)
return f"{prefix}{preview}...<remaining_file_bytes:{remaining_bytes}>"
def _sanitize_url(value: str) -> str:
@@ -65,10 +117,13 @@ def _sanitize_url(value: str) -> str:
return value
def sanitize_log_value(value: Any) -> Any:
def sanitize_log_value(value: Any, *, key_path: tuple[str, ...] = ()) -> Any:
if value is None:
return None
if isinstance(value, str):
file_preview = _file_base64_preview(value, key_path)
if file_preview is not None:
return file_preview
text = _sanitize_url(value) if value.startswith(("http://", "https://")) else value
if len(text) > MAX_LOG_FIELD_LENGTH:
return text[:MAX_LOG_FIELD_LENGTH] + f"...<truncated:{len(text) - MAX_LOG_FIELD_LENGTH}>"
@@ -77,10 +132,14 @@ def sanitize_log_value(value: Any) -> Any:
output: dict[str, Any] = {}
for k, v in value.items():
key = str(k)
output[key] = "***" if _is_sensitive_key(key) else sanitize_log_value(v)
output[key] = (
"***"
if _is_sensitive_key(key)
else sanitize_log_value(v, key_path=(*key_path, key))
)
return output
if isinstance(value, list):
return [sanitize_log_value(v) for v in value]
if isinstance(value, (list, tuple)):
return [sanitize_log_value(v, key_path=(*key_path, str(index))) for index, v in enumerate(value)]
return value
@@ -262,6 +321,13 @@ def log_ai_model_event(
*,
event_type: str,
module: str | None = None,
step_code: str | None = None,
call_id: str | None = None,
event_phase: str | None = None,
owner_type: str | None = None,
owner_id: str | None = None,
generation_attempt_no: int | None = None,
latency_ms: int | None = None,
event_status: str = "success",
source: str | None = None,
trace_id: str | None = None,
@@ -314,8 +380,19 @@ def log_ai_model_event(
)
entry.update(
{
"call_id": call_id,
"event_phase": event_phase or event_type,
"step_code": step_code,
"owner_type": owner_type,
"owner_id": owner_id,
"generation_attempt_no": generation_attempt_no,
"latency_ms": latency_ms,
# Preserve the legacy fields for existing log readers while also
# exposing unambiguous configuration/provider model names.
"model_name": model_config_name,
"model_id": model_name,
"model_config_name": model_config_name,
"provider_model_name": model_name,
"model_config_id": model_config_id,
"provider": provider,
"api_base": api_base,
@@ -2,6 +2,8 @@ from __future__ import annotations
import asyncio
import json
import time
import uuid
from typing import Any
from fastapi import HTTPException
@@ -14,10 +16,8 @@ from app.enums.private_portrait import (
ARK_PRIVATE_PORTRAIT_VERSION,
ArkPrivatePortraitAction,
PrivatePortraitEventSource,
PrivatePortraitEventStatus,
PrivatePortraitEventType,
)
from app.services.operation_log_service import log_remote_api_event
from app.services.operation_log_service import log_ai_model_event
from app.services.private_portrait.rate_limiter import acquire_private_portrait_action_token
DOMAIN = "private_portrait"
@@ -117,39 +117,69 @@ class ArkPrivateAssetClient:
async def _call(self, action: ArkPrivatePortraitAction, payload: dict[str, Any]) -> dict[str, Any]:
action_value = action.value
await acquire_private_portrait_action_token(action=action_value, wait_timeout_seconds=2.0, for_celery=self.for_celery)
log_remote_api_event(
domain=DOMAIN,
call_id = uuid.uuid4().hex
source = PrivatePortraitEventSource.CELERY.value if self.for_celery else PrivatePortraitEventSource.SERVICE.value
started = time.perf_counter()
log_ai_model_event(
event_type="REQUEST",
event_phase="REQUEST",
event_status="started",
module=DOMAIN,
step_code=action_value,
call_id=call_id,
source=source,
remote_action=action_value,
event_type=PrivatePortraitEventType.ARK_API_CALL_START.value,
event_status=PrivatePortraitEventStatus.PENDING.value,
source=PrivatePortraitEventSource.CELERY.value if self.for_celery else PrivatePortraitEventSource.SERVICE.value,
provider="volcengine_ark",
request=payload,
)
try:
result = await asyncio.to_thread(self._call_sync, action, payload)
log_remote_api_event(
domain=DOMAIN,
log_ai_model_event(
event_type="RESPONSE",
event_phase="RESPONSE",
event_status="success",
module=DOMAIN,
step_code=action_value,
call_id=call_id,
source=source,
remote_action=action_value,
event_type=PrivatePortraitEventType.ARK_API_CALL_SUCCESS.value,
event_status=PrivatePortraitEventStatus.SUCCESS.value,
source=PrivatePortraitEventSource.CELERY.value if self.for_celery else PrivatePortraitEventSource.SERVICE.value,
request=payload,
response=result,
remote_request_id=result.get("RequestId") or result.get("request_id"),
provider="volcengine_ark",
latency_ms=int((time.perf_counter() - started) * 1000),
response=result,
)
return result
except ArkPrivateAssetRemoteError as exc:
log_remote_api_event(
domain=DOMAIN,
log_ai_model_event(
event_type="RESPONSE",
event_phase="RESPONSE",
event_status="failed",
module=DOMAIN,
step_code=action_value,
call_id=call_id,
source=source,
remote_action=action_value,
event_type=PrivatePortraitEventType.ARK_API_CALL_FAILED.value,
event_status=PrivatePortraitEventStatus.FAILED.value,
source=PrivatePortraitEventSource.CELERY.value if self.for_celery else PrivatePortraitEventSource.SERVICE.value,
request=payload,
response=exc.raw,
remote_request_id=exc.request_id,
remote_code=exc.code,
remote_message=exc.message,
provider="volcengine_ark",
latency_ms=int((time.perf_counter() - started) * 1000),
response=exc.raw,
detail={"remote_code": exc.code, "remote_message": exc.message},
error=exc.message,
)
log_ai_model_event(
event_type="ERROR",
event_phase="ERROR",
event_status="failed",
module=DOMAIN,
step_code=action_value,
call_id=call_id,
source=source,
remote_action=action_value,
remote_request_id=exc.request_id,
provider="volcengine_ark",
latency_ms=int((time.perf_counter() - started) * 1000),
detail={"remote_code": exc.code, "remote_message": exc.message},
error=str(exc),
)
if self.for_celery:
raise
@@ -157,14 +187,18 @@ class ArkPrivateAssetClient:
except HTTPException:
raise
except Exception as exc:
log_remote_api_event(
domain=DOMAIN,
log_ai_model_event(
event_type="ERROR",
event_phase="ERROR",
event_status="failed",
module=DOMAIN,
step_code=action_value,
call_id=call_id,
source=source,
remote_action=action_value,
event_type=PrivatePortraitEventType.ARK_API_CALL_FAILED.value,
event_status=PrivatePortraitEventStatus.FAILED.value,
source=PrivatePortraitEventSource.CELERY.value if self.for_celery else PrivatePortraitEventSource.SERVICE.value,
request=payload,
remote_message=str(exc),
provider="volcengine_ark",
latency_ms=int((time.perf_counter() - started) * 1000),
error=str(exc),
)
if self.for_celery:
raise
@@ -921,6 +921,13 @@ async def run_image_prompt_optimize(
user_id=user_id_value,
references=references,
gen_type="image",
log_module=module_value,
log_step="shot_replicate_image_prompt_optimize",
log_project_id=project_id_value,
log_task_id=step_id_value,
log_owner_type="module_generation_step",
log_owner_id=step_id_value,
generation_attempt_no=expected_step_version,
)
if execution_guard is not None:
await execution_guard()
@@ -459,6 +459,7 @@ def _log_shot_ai_model_event(
event_status: str,
config: ModelConfig,
trace_id: str,
call_id: str,
user_id: str | None,
task_set_id: str | None,
segment_id: str | None,
@@ -487,30 +488,61 @@ def _log_shot_ai_model_event(
"remote_message": remote_message,
"remote_param": remote_param,
})
log_ai_model_event(
event_type=event_type,
event_status=event_status,
source=LogSourceEnum.REMOTE_API.value,
module=ModuleCodeEnum.SHOT_REPLICATE.value,
trace_id=trace_id,
user_id=user_id,
project_id=task_set_id,
step_id=segment_id,
remote_action=action,
remote_request_id=remote_request_id,
model_config_id=str(config.id),
model_config_name=config.name,
model_name=config.model_name,
provider=config.provider,
api_base=config.api_base,
http_status=http_status,
request=request_data,
response=response_data,
token_usage=token_usage,
message=message,
detail=detail,
error=error,
)
common = {
"source": LogSourceEnum.REMOTE_API.value,
"module": ModuleCodeEnum.SHOT_REPLICATE.value,
"step_code": "source_video_analysis" if mode == "full_breakdown" else "segment_video_analysis",
"call_id": call_id,
"trace_id": trace_id,
"user_id": user_id,
"project_id": task_set_id,
"task_id": segment_id or task_set_id,
"step_id": segment_id,
"owner_type": "shot_replicate_segment" if segment_id else "shot_replicate_task_set",
"owner_id": segment_id or task_set_id,
"remote_action": action,
"remote_request_id": remote_request_id,
"model_config_id": str(config.id),
"model_config_name": config.name,
"model_name": config.model_name,
"provider": config.provider,
"api_base": config.api_base,
"http_status": http_status,
}
normalized_status = str(event_status or "").lower()
if normalized_status == str(LogEventStatusEnum.STARTED.value).lower():
log_ai_model_event(
event_type=event_type,
event_phase="REQUEST",
event_status=event_status,
request=request_data,
message=message,
detail=detail,
**common,
)
return
if response_data is not None:
log_ai_model_event(
event_type=event_type,
event_phase="RESPONSE",
event_status=event_status,
response=response_data,
token_usage=token_usage,
message=message,
detail=detail,
error=error if normalized_status == str(LogEventStatusEnum.FAILED.value).lower() else None,
**common,
)
if normalized_status == str(LogEventStatusEnum.FAILED.value).lower() or error:
log_ai_model_event(
event_type=event_type,
event_phase="ERROR",
event_status=LogEventStatusEnum.FAILED.value,
message=message,
detail=detail,
error=error or remote_message or "AI model call failed",
**common,
)
async def analyze_video_for_shot_split(
db: AsyncSession,
@@ -529,6 +561,7 @@ async def analyze_video_for_shot_split(
也不再 fallback SEEDANCE_*避免拆镜分析走错通道
"""
trace_id = trace_id or generate_id()
call_id = generate_id()
config_row = await _select_model_config(db)
if not config_row:
raise RuntimeError("拆镜分析模型未配置:请先在 model_configs 表启用可用模型")
@@ -585,6 +618,7 @@ async def analyze_video_for_shot_split(
await db.rollback()
_log_shot_ai_model_event(
call_id=call_id,
event_type=(
ShotReplicateLogEventEnum.ANALYSIS_REMOTE_API_STARTED.value
if mode == "full_breakdown"
@@ -610,6 +644,7 @@ async def analyze_video_for_shot_split(
)
except Exception as exc:
_log_shot_ai_model_event(
call_id=call_id,
event_type=(
ShotReplicateLogEventEnum.ANALYSIS_REMOTE_API_FAILED.value
if mode == "full_breakdown"
@@ -633,6 +668,7 @@ async def analyze_video_for_shot_split(
remote_request_id, remote_code, remote_message, remote_param = _extract_remote_error(response_data)
if response.status_code >= 400:
_log_shot_ai_model_event(
call_id=call_id,
event_type=(
ShotReplicateLogEventEnum.ANALYSIS_REMOTE_API_FAILED.value
if mode == "full_breakdown"
@@ -661,6 +697,7 @@ async def analyze_video_for_shot_split(
raw = response.json()
except Exception as exc:
_log_shot_ai_model_event(
call_id=call_id,
event_type=ShotReplicateLogEventEnum.ANALYSIS_RESPONSE_PARSE_FAILED.value,
event_status=LogEventStatusEnum.FAILED.value,
config=config,
@@ -683,6 +720,7 @@ async def analyze_video_for_shot_split(
except Exception as exc:
event_type = ShotReplicateLogEventEnum.ANALYSIS_RESPONSE_EMPTY.value if "content 为空" in str(exc) else ShotReplicateLogEventEnum.ANALYSIS_RESPONSE_PARSE_FAILED.value
_log_shot_ai_model_event(
call_id=call_id,
event_type=event_type,
event_status=LogEventStatusEnum.FAILED.value,
config=config,
@@ -718,7 +756,6 @@ async def analyze_video_for_shot_split(
"split_max_seconds": _split_max_seconds(),
"analysis_mode": mode,
"trace_id": trace_id,
"log_request": log_request_data,
}
if not token_usage["total_tokens"]:
token_usage["total_tokens"] = token_usage["input_tokens"] + token_usage["output_tokens"]
@@ -731,6 +768,7 @@ async def analyze_video_for_shot_split(
})
_log_shot_ai_model_event(
call_id=call_id,
event_type=(
ShotReplicateLogEventEnum.ANALYSIS_REMOTE_API_SUCCESS.value
if mode == "full_breakdown"
+58 -59
View File
@@ -1,9 +1,7 @@
import base64
import json
import logging
import mimetypes
import os
from datetime import datetime, timezone
import time
import httpx
from sqlalchemy import select
@@ -13,7 +11,8 @@ from volcenginesdkarkruntime import AsyncArk
from app.config import settings
from app.enums.private_portrait import PRIVATE_PORTRAIT_ASSET_URI_PREFIX
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.operation_log_service import build_exception_detail, log_ai_model_event
from app.utils.id_gen import generate_id
from app.types.generation.provider import (
ProviderGenerationRecordLike,
ProviderVideoEngineLike,
@@ -22,55 +21,27 @@ from app.types.generation.provider import (
logger = logging.getLogger("videogen")
def _log_video_request(engine: ProviderVideoEngineLike, record_id: str, request_data: dict):
"""Log video generation request to log/AiModel/YYYY-MM-DD.log"""
if not is_enabled():
return
try:
os.makedirs(LOG_DIR, exist_ok=True)
today = datetime.now().strftime(LOG_DATE_FORMAT)
log_file = os.path.join(LOG_DIR, f"{today}.log")
request_str = json.dumps(request_data, ensure_ascii=False)
request_encrypted = encrypt_data(request_data, True)
entry = {
"timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
"type": "video_gen_request",
"engine": engine.name,
"model": engine.model_name,
"record_id": record_id,
"request": request_encrypted,
"request_length": len(request_str),
}
with open(log_file, "a", encoding="utf-8") as f:
f.write(json.dumps(entry, ensure_ascii=False) + "\n")
except Exception:
pass
def _log_video_response(record_id: str, response_data: dict, error: str | None = None):
"""Log video generation response to log/AiModel/YYYY-MM-DD.log"""
if not is_enabled():
return
try:
os.makedirs(LOG_DIR, exist_ok=True)
today = datetime.now().strftime(LOG_DATE_FORMAT)
log_file = os.path.join(LOG_DIR, f"{today}.log")
response_encrypted = encrypt_data(response_data, True) if response_data else ""
entry = {
"timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
"type": "video_gen_response",
"record_id": record_id,
"response": response_encrypted,
"error": error,
}
with open(log_file, "a", encoding="utf-8") as f:
f.write(json.dumps(entry, ensure_ascii=False) + "\n")
except Exception:
pass
def _provider_log_context(engine, record, *, call_id: str, step_code: str) -> dict:
generation_mode = str(getattr(record, "generation_mode", "") or "generation_record")
owner_type = "chat_generation_task" if generation_mode != "generation_record" else "generation_record"
return {
"module": generation_mode,
"step_code": step_code,
"call_id": call_id,
"source": "app.services.video_gen",
"user_id": str(getattr(record, "user_id", "") or "") or None,
"project_id": str(getattr(record, "project_id", "") or "") or None,
"task_id": str(getattr(record, "id", "") or "") or None,
"owner_type": owner_type,
"owner_id": str(getattr(record, "id", "") or "") or None,
"generation_attempt_no": int(getattr(record, "generation_attempt_no", 1) or 1),
"model_config_id": str(getattr(engine, "id", "") or "") or None,
"model_config_name": str(getattr(engine, "name", "") or "") or None,
"model_name": str(getattr(engine, "model_name", "") or "") or None,
"provider": str(getattr(engine, "provider", "") or "") or None,
"api_base": str(getattr(engine, "api_base", "") or "") or None,
}
async def get_active_engine(db: AsyncSession) -> VideoEngine:
@@ -153,19 +124,47 @@ async def submit_video_task(
"watermark": False,
}
# Log request to AiModel log. include_media_references 只用于排查日志,不传给供应商 API。
_log_video_request(
call_id = generate_id()
started = time.perf_counter()
log_context = _provider_log_context(
engine,
record.id,
{**request_payload, "include_media_references": include_media_references},
record,
call_id=call_id,
step_code="video_create",
)
log_ai_model_event(
event_type="REQUEST",
event_phase="REQUEST",
event_status="started",
remote_action="video_create",
request={**request_payload, "include_media_references": include_media_references},
**log_context,
)
try:
result = await client.content_generation.tasks.create(**request_payload)
task_id = result.id
_log_video_response(record.id, {"task_id": task_id})
except Exception as e:
_log_video_response(record.id, {}, str(e))
log_ai_model_event(
event_type="RESPONSE",
event_phase="RESPONSE",
event_status="success",
remote_action="video_create",
remote_request_id=task_id,
latency_ms=int((time.perf_counter() - started) * 1000),
response={"task_id": task_id},
**log_context,
)
except Exception as exc:
log_ai_model_event(
event_type="ERROR",
event_phase="ERROR",
event_status="failed",
remote_action="video_create",
latency_ms=int((time.perf_counter() - started) * 1000),
detail=build_exception_detail(exc),
error=str(exc),
**log_context,
)
raise
finally:
await client.close()
@@ -6,6 +6,7 @@ import os
import uuid
from datetime import datetime, timedelta, timezone
from collections.abc import Awaitable, Callable
from types import SimpleNamespace
from typing import Any
from urllib.parse import urlsplit, urlunsplit
@@ -305,6 +306,28 @@ async def enqueue_upscale_task(db: AsyncSession, *, upscale: VideoUpscaleTask, r
upscale.status = VideoUpscaleTaskStatus.PENDING.value
upscale.stage = VideoUpscaleStage.QUEUED.value
upscale.next_retry_at = None
log_snapshot = SimpleNamespace(
id=upscale_id,
chat_generation_task_id=(
str(upscale.chat_generation_task_id)
if upscale.chat_generation_task_id
else None
),
generation_record_id=(
str(upscale.generation_record_id)
if upscale.generation_record_id
else None
),
processor_key=processor_key,
status=VideoUpscaleTaskStatus.PENDING.value,
stage=VideoUpscaleStage.QUEUED.value,
attempt_count=int(upscale.attempt_count or 0),
failure_count=int(upscale.failure_count or 0),
provider_task_id=(str(upscale.provider_task_id) if upscale.provider_task_id else None),
input_source_type=upscale.input_source_type,
target_width=upscale.target_width,
target_height=upscale.target_height,
)
await db.commit()
try:
@@ -324,7 +347,7 @@ async def enqueue_upscale_task(db: AsyncSession, *, upscale: VideoUpscaleTask, r
raise RuntimeError(f"未注册的超分处理器: {processor_key}")
log_video_upscale_event(
event_type="upscale_task_enqueued",
upscale_task=upscale,
upscale_task=log_snapshot,
detail={"reason": reason, "celery_task_id": celery_id, "processor_key": processor_key},
)
except Exception as exc:
@@ -333,7 +356,7 @@ async def enqueue_upscale_task(db: AsyncSession, *, upscale: VideoUpscaleTask, r
log_video_upscale_event(
event_type="upscale_task_enqueue_failed",
event_status="failed",
upscale_task=upscale,
upscale_task=log_snapshot,
message=str(exc),
detail={"reason": reason, "celery_task_id": celery_id, "processor_key": processor_key},
error=str(exc),
@@ -1,12 +1,15 @@
from __future__ import annotations
from dataclasses import dataclass
import time
import uuid
from typing import Any
import httpx
from app.config import settings
from app.enums.video_upscale import VideoUpscaleProcessorKey
from app.services.operation_log_service import log_ai_model_event
class VolcMediaKitError(RuntimeError):
@@ -171,7 +174,24 @@ async def submit_video_enhance(
processor=processor,
client_token=client_token,
)
call_id = uuid.uuid4().hex
started = time.perf_counter()
log_ai_model_event(
event_type="REQUEST",
event_phase="REQUEST",
event_status="started",
module="video_upscale",
step_code="provider_submit",
call_id=call_id,
source="app.services.video_upscale.volc_service",
remote_action=endpoint,
provider="volcengine_mediakit",
api_base=_base_url(),
request=payload,
)
timeout = max(3, int(processor.get("request_timeout_seconds") or settings.VIDEO_UPSCALE_REMOTE_REQUEST_TIMEOUT_SECONDS))
response: httpx.Response | None = None
data: dict[str, Any] = {}
try:
async with httpx.AsyncClient(timeout=timeout, follow_redirects=True) as client:
response = await client.post(f"{_base_url()}{endpoint}", headers=_headers(), json=payload)
@@ -179,6 +199,23 @@ async def submit_video_enhance(
data = response.json()
except Exception:
data = {"success": False, "error": {"message": response.text[:2000]}}
log_ai_model_event(
event_type="RESPONSE",
event_phase="RESPONSE",
event_status="success" if response.status_code < 400 else "failed",
module="video_upscale",
step_code="provider_submit",
call_id=call_id,
source="app.services.video_upscale.volc_service",
remote_action=endpoint,
remote_request_id=str(data.get("request_id") or "") or None,
provider="volcengine_mediakit",
api_base=_base_url(),
http_status=response.status_code,
latency_ms=int((time.perf_counter() - started) * 1000),
response=data,
error=None if response.status_code < 400 else str((data.get("error") or {}).get("message") or "remote error"),
)
if response.status_code >= 400:
raise _error_from_payload(
data,
@@ -186,18 +223,51 @@ async def submit_video_enhance(
http_status=response.status_code,
endpoint=endpoint,
)
except VolcMediaKitError:
if not bool(data.get("success")) or not data.get("task_id"):
raise _error_from_payload(data, "火山超分提交失败", endpoint=endpoint)
except VolcMediaKitError as exc:
log_ai_model_event(
event_type="ERROR",
event_phase="ERROR",
event_status="failed",
module="video_upscale",
step_code="provider_submit",
call_id=call_id,
source="app.services.video_upscale.volc_service",
remote_action=endpoint,
remote_request_id=exc.request_id,
provider="volcengine_mediakit",
api_base=_base_url(),
http_status=exc.http_status or (response.status_code if response is not None else None),
latency_ms=int((time.perf_counter() - started) * 1000),
detail=exc.log_detail(),
error=str(exc),
)
raise
except (httpx.TimeoutException, httpx.NetworkError) as exc:
raise VolcMediaKitError(
wrapped = VolcMediaKitError(
f"火山超分提交网络异常: {exc}",
code="NetworkError",
retryable=True,
endpoint=endpoint,
) from exc
)
log_ai_model_event(
event_type="ERROR",
event_phase="ERROR",
event_status="failed",
module="video_upscale",
step_code="provider_submit",
call_id=call_id,
source="app.services.video_upscale.volc_service",
remote_action=endpoint,
provider="volcengine_mediakit",
api_base=_base_url(),
latency_ms=int((time.perf_counter() - started) * 1000),
detail=wrapped.log_detail(),
error=str(wrapped),
)
raise wrapped from exc
if not bool(data.get("success")) or not data.get("task_id"):
raise _error_from_payload(data, "火山超分提交失败", endpoint=endpoint)
return VolcSubmitResult(
task_id=str(data["task_id"]),
request_id=str(data.get("request_id")) if data.get("request_id") else None,
@@ -209,7 +279,26 @@ async def submit_video_enhance(
async def query_task(task_id: str, *, request_timeout_seconds: int | None = None) -> VolcQueryResult:
endpoint = f"/api/v1/tasks/{task_id}"
call_id = uuid.uuid4().hex
started = time.perf_counter()
request_payload = {"task_id": task_id}
log_ai_model_event(
event_type="REQUEST",
event_phase="REQUEST",
event_status="started",
module="video_upscale",
step_code="provider_poll",
call_id=call_id,
source="app.services.video_upscale.volc_service",
remote_action=endpoint,
remote_request_id=task_id,
provider="volcengine_mediakit",
api_base=_base_url(),
request=request_payload,
)
timeout = max(3, int(request_timeout_seconds or settings.VIDEO_UPSCALE_REMOTE_REQUEST_TIMEOUT_SECONDS))
response: httpx.Response | None = None
data: dict[str, Any] = {}
try:
async with httpx.AsyncClient(timeout=timeout, follow_redirects=True) as client:
response = await client.get(f"{_base_url()}{endpoint}", headers=_headers())
@@ -217,6 +306,23 @@ async def query_task(task_id: str, *, request_timeout_seconds: int | None = None
data = response.json()
except Exception:
data = {"success": False, "error": {"message": response.text[:2000]}}
log_ai_model_event(
event_type="RESPONSE",
event_phase="RESPONSE",
event_status="success" if response.status_code < 400 else "failed",
module="video_upscale",
step_code="provider_poll",
call_id=call_id,
source="app.services.video_upscale.volc_service",
remote_action=endpoint,
remote_request_id=str(data.get("request_id") or task_id),
provider="volcengine_mediakit",
api_base=_base_url(),
http_status=response.status_code,
latency_ms=int((time.perf_counter() - started) * 1000),
response=data,
error=None if response.status_code < 400 else str((data.get("error") or {}).get("message") or "remote error"),
)
if response.status_code >= 400:
raise _error_from_payload(
data,
@@ -224,28 +330,62 @@ async def query_task(task_id: str, *, request_timeout_seconds: int | None = None
http_status=response.status_code,
endpoint=endpoint,
)
except VolcMediaKitError:
if not bool(data.get("success")):
raise _error_from_payload(data, "火山超分任务查询失败", endpoint=endpoint)
status = str(data.get("status") or "").strip().lower()
if status not in {"running", "completed", "failed"}:
raise VolcMediaKitError(
f"火山超分返回未知任务状态: {status}",
code="UnknownStatus",
retryable=True,
request_id=str(data.get("request_id") or "") or None,
endpoint=endpoint,
response_payload=data,
)
except VolcMediaKitError as exc:
log_ai_model_event(
event_type="ERROR",
event_phase="ERROR",
event_status="failed",
module="video_upscale",
step_code="provider_poll",
call_id=call_id,
source="app.services.video_upscale.volc_service",
remote_action=endpoint,
remote_request_id=exc.request_id or task_id,
provider="volcengine_mediakit",
api_base=_base_url(),
http_status=exc.http_status or (response.status_code if response is not None else None),
latency_ms=int((time.perf_counter() - started) * 1000),
detail=exc.log_detail(),
error=str(exc),
)
raise
except (httpx.TimeoutException, httpx.NetworkError) as exc:
raise VolcMediaKitError(
wrapped = VolcMediaKitError(
f"火山超分查询网络异常: {exc}",
code="NetworkError",
retryable=True,
endpoint=endpoint,
) from exc
if not bool(data.get("success")):
raise _error_from_payload(data, "火山超分任务查询失败", endpoint=endpoint)
status = str(data.get("status") or "").strip().lower()
if status not in {"running", "completed", "failed"}:
raise VolcMediaKitError(
f"火山超分返回未知任务状态: {status}",
code="UnknownStatus",
retryable=True,
request_id=str(data.get("request_id") or "") or None,
endpoint=endpoint,
response_payload=data,
)
log_ai_model_event(
event_type="ERROR",
event_phase="ERROR",
event_status="failed",
module="video_upscale",
step_code="provider_poll",
call_id=call_id,
source="app.services.video_upscale.volc_service",
remote_action=endpoint,
remote_request_id=task_id,
provider="volcengine_mediakit",
api_base=_base_url(),
latency_ms=int((time.perf_counter() - started) * 1000),
detail=wrapped.log_detail(),
error=str(wrapped),
)
raise wrapped from exc
expires_raw = data.get("expires_at")
try:
expires_at = int(expires_raw) if expires_raw is not None else None
-51
View File
@@ -1,51 +0,0 @@
import asyncio
from datetime import datetime, timedelta
from app.tasks.celery_app import celery_app
@celery_app.task
def cleanup_expired_video_urls():
"""Run hourly. Clear expired video URL tokens."""
asyncio.run(_cleanup_urls())
async def _cleanup_urls():
from app.models.base import async_session
from app.models.generation_record import GenerationRecord
from sqlalchemy import update
async with async_session() as db:
now = datetime.now()
await db.execute(
update(GenerationRecord)
.where(
GenerationRecord.video_url_expires_at.isnot(None),
GenerationRecord.video_url_expires_at < now,
GenerationRecord.deleted_at.is_(None),
)
.values(video_url_expires_at=None)
)
await db.commit()
@celery_app.task
def cleanup_old_notifications():
"""Run daily. Delete read notifications older than 30 days."""
asyncio.run(_cleanup_notifications())
async def _cleanup_notifications():
from app.models.base import async_session
from app.models.notification import Notification
from sqlalchemy import delete
async with async_session() as db:
cutoff = datetime.now() - timedelta(days=30)
await db.execute(
delete(Notification).where(
Notification.is_read == True,
Notification.created_at < cutoff,
)
)
await db.commit()
@@ -213,6 +213,24 @@ async def _remove_active(owner: GenerationOwner) -> None:
await remove_download_active(_registry_id(owner))
async def _reload_owner_after_commit(
db: AsyncSession,
*,
owner_type: str,
owner_id: str,
attempt_no: int,
) -> GenerationOwner | None:
owner = await load_generation_owner(
db,
owner_type=owner_type,
owner_id=owner_id,
for_update=False,
)
if owner is None or not is_attempt_current(owner, attempt_no):
return None
return owner
async def _apply(
owner: GenerationOwner,
*,
@@ -297,7 +315,19 @@ async def enqueue_download_task(
if not owner.download_storage_date_dir:
created_at = ensure_aware_utc(owner.created_at) or _now()
owner.download_storage_date_dir = created_at.strftime("%Y/%m/%d")
owner_type_snapshot = owner_type_of(owner)
owner_id_snapshot = str(owner.id)
attempt_snapshot = int(owner.generation_attempt_no or 1)
celery_task_id_snapshot = str(owner.download_celery_task_id or "") or None
await db.commit()
owner = await load_generation_owner(
db,
owner_type=owner_type_snapshot,
owner_id=owner_id_snapshot,
for_update=False,
)
if owner is None or not is_attempt_current(owner, attempt_snapshot):
return None
priority = int(
settings.DOWNLOAD_TASK_PRIORITY_RECOVER
@@ -325,7 +355,7 @@ async def enqueue_download_task(
detail={"reason": reason, "error": str(exc)},
)
return None
return owner.download_celery_task_id
return str(owner.download_celery_task_id or celery_task_id_snapshot or "") or None
async def _claim(
@@ -333,9 +363,9 @@ async def _claim(
owner: GenerationOwner,
*,
claim_token: str,
) -> bool:
) -> GenerationOwner | None:
if not owner_is_generating(owner) or owner_is_completed(owner):
return False
return None
allowed = {
_stage(owner, ChatGenerationPipelineStage.RESULT_READY),
_stage(owner, ChatGenerationPipelineStage.DOWNLOAD_QUEUED),
@@ -343,7 +373,7 @@ async def _claim(
_stage(owner, ChatGenerationPipelineStage.RETRY_WAITING),
}
if owner.pipeline_stage not in allowed:
return False
return None
now = _now()
# Redis execution lock is authoritative. A database lease left by a
@@ -356,7 +386,7 @@ async def _claim(
and next_retry
and next_retry > now
):
return False
return None
owner.pipeline_stage = _stage(
owner, ChatGenerationPipelineStage.DOWNLOADING
@@ -366,7 +396,18 @@ async def _claim(
owner.download_lease_until = _lease_until()
owner.download_attempt_count = int(owner.download_attempt_count or 0) + 1
owner.download_last_error = None
owner_type_snapshot = owner_type_of(owner)
owner_id_snapshot = str(owner.id)
attempt_snapshot = int(owner.generation_attempt_no or 1)
await db.commit()
owner = await _reload_owner_after_commit(
db,
owner_type=owner_type_snapshot,
owner_id=owner_id_snapshot,
attempt_no=attempt_snapshot,
)
if owner is None:
return None
await _register_active(
owner,
check_at=owner.download_lease_until,
@@ -382,7 +423,7 @@ async def _claim(
"claim_token_suffix": claim_token[-8:],
},
)
return True
return owner
async def _sync_snapshot(db: AsyncSession, owner: GenerationOwner) -> None:
@@ -446,12 +487,33 @@ async def _mark_failed(
owner.download_last_error = error_message
owner.download_lease_until = None
owner.download_next_retry_at = None
owner_type_snapshot = owner_type_of(owner)
owner_id_snapshot = str(owner.id)
attempt_snapshot = int(owner.generation_attempt_no or 1)
generation_mode_snapshot = str(getattr(owner, "generation_mode", "") or "") or None
await db.commit()
await notify_owner_finished(db, owner)
owner = await _reload_owner_after_commit(
db,
owner_type=owner_type_snapshot,
owner_id=owner_id_snapshot,
attempt_no=attempt_snapshot,
)
if owner is not None:
await notify_owner_finished(db, owner)
await db.commit()
await _remove_active(owner)
owner = await _reload_owner_after_commit(
db,
owner_type=owner_type_snapshot,
owner_id=owner_id_snapshot,
attempt_no=attempt_snapshot,
)
if owner is not None:
await _remove_active(owner)
await log_task_event(
owner,
owner_type=owner_type_snapshot,
owner_id=owner_id_snapshot,
generation_attempt_no=attempt_snapshot,
generation_mode=generation_mode_snapshot,
event_type=(
ChatGenerationTaskEventType.DOWNLOAD_FAILED_NON_RETRYABLE.value
if non_retryable
@@ -481,7 +543,18 @@ async def _schedule_retry(
owner.download_celery_task_id = _build_celery_task_id(
owner, reason="retry"
)
owner_type_snapshot = owner_type_of(owner)
owner_id_snapshot = str(owner.id)
attempt_snapshot = int(owner.generation_attempt_no or 1)
await db.commit()
owner = await _reload_owner_after_commit(
db,
owner_type=owner_type_snapshot,
owner_id=owner_id_snapshot,
attempt_no=attempt_snapshot,
)
if owner is None:
return
await _register_active(
owner,
check_at=owner.download_next_retry_at,
@@ -555,7 +628,18 @@ async def _restore_after_lock_error(
0, int(owner.download_attempt_count or 0) - 1
)
owner.download_enqueued_at = _now()
owner_type_snapshot = owner_type_of(owner)
owner_id_snapshot = str(owner.id)
attempt_snapshot = int(owner.generation_attempt_no or 1)
await db.commit()
owner = await _reload_owner_after_commit(
db,
owner_type=owner_type_snapshot,
owner_id=owner_id_snapshot,
attempt_no=attempt_snapshot,
)
if owner is None:
return
await _register_active(
owner,
check_at=_queue_timeout_at(),
@@ -631,7 +715,8 @@ async def _run(
if not is_attempt_current(owner, effective_attempt):
await _remove_active(owner)
return
if not await _claim(db, owner, claim_token=lease.token):
owner = await _claim(db, owner, claim_token=lease.token)
if owner is None:
return
claimed = True
@@ -705,15 +790,26 @@ async def _run(
owner.download_lease_until = None
owner.download_next_retry_at = None
owner.download_last_error = None
await db.commit()
await _remove_active(owner)
upscale_mode = str(getattr(owner, "generation_mode", "") or "") or None
upscale_stage = str(owner.pipeline_stage or "") or None
await enqueue_upscale_task(
db, upscale=upscale, reason="source_download_completed"
)
owner = await _reload_owner_after_commit(
db,
owner_type=normalized_owner_type,
owner_id=task_id,
attempt_no=effective_attempt,
)
if owner is not None:
await _remove_active(owner)
await log_task_event(
owner,
owner_type=normalized_owner_type,
owner_id=task_id,
generation_attempt_no=effective_attempt,
generation_mode=upscale_mode,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SUCCESS.value,
to_stage=owner.pipeline_stage,
to_stage=(str(owner.pipeline_stage or "") if owner is not None else upscale_stage),
detail={
"upscale_source_path": downloaded.storage_path
},
@@ -739,14 +835,33 @@ async def _run(
owner.download_last_error = None
await _record_resource(db, owner, downloaded)
await _sync_snapshot(db, owner)
completion_mode = str(getattr(owner, "generation_mode", "") or "") or None
completion_stage = str(owner.pipeline_stage or "") or None
await db.commit()
await notify_owner_finished(db, owner)
owner = await _reload_owner_after_commit(
db,
owner_type=normalized_owner_type,
owner_id=task_id,
attempt_no=effective_attempt,
)
if owner is not None:
await notify_owner_finished(db, owner)
await db.commit()
await _remove_active(owner)
owner = await _reload_owner_after_commit(
db,
owner_type=normalized_owner_type,
owner_id=task_id,
attempt_no=effective_attempt,
)
if owner is not None:
await _remove_active(owner)
await log_task_event(
owner,
owner_type=normalized_owner_type,
owner_id=task_id,
generation_attempt_no=effective_attempt,
generation_mode=completion_mode,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SUCCESS.value,
to_stage=owner.pipeline_stage,
to_stage=(str(owner.pipeline_stage or "") if owner is not None else completion_stage),
detail={
"resource_url": downloaded.url,
"file_size_bytes": downloaded.file_size_bytes,
+119 -58
View File
@@ -22,7 +22,7 @@ from app.enums.generation_task import (
from app.models.base import async_session
from app.models.chat_generation_task import ChatGenerationTask
from app.services.error_codes import extract_error_message
from app.services.generation.log_service import log_provider_call, log_task_event
from app.services.generation.log_service import log_task_event
from app.services.generation.pipeline.db_lock_service import DatabaseRowLockBusy
from app.services.generation.pipeline.lifecycle_service import (
mark_owner_failed_and_refund_once,
@@ -101,24 +101,6 @@ def _engine_snapshot(owner: GenerationOwner) -> dict:
return {}
async def _log_poll_provider_call_after_commit(
owner: GenerationOwner,
*,
provider_response: Any,
) -> None:
"""Provider logs use an independent session, so the owner row must be committed first."""
snapshot = _engine_snapshot(owner)
await log_provider_call(
owner,
provider=snapshot.get("provider") or "ark",
api_type=f"{owner.gen_type}_poll",
model=snapshot.get("model_name"),
engine_id=owner.engine_id,
status="success",
provider_task_id=owner_provider_task_id(owner),
response_data=provider_response,
)
def _registry_id(owner: GenerationOwner) -> str:
return redis_owner_item_id(
@@ -235,6 +217,24 @@ async def remove_poll_active(
)
async def _reload_owner_after_commit(
db,
*,
owner_type: str,
owner_id: str,
attempt_no: int,
) -> GenerationOwner | None:
fresh = await load_generation_owner(
db,
owner_type=owner_type,
owner_id=owner_id,
for_update=False,
)
if fresh is None or not is_attempt_current(fresh, attempt_no):
return None
return fresh
async def _sync_snapshot(
db, owner: GenerationOwner, provider_response: Any = None
) -> None:
@@ -266,16 +266,34 @@ async def _mark_failed(
owner.next_poll_at = None
owner.poll_claim_token = None
owner.poll_lease_until = None
owner_type_snapshot = owner_type_of(owner)
owner_id_snapshot = str(owner.id)
attempt_snapshot = int(owner.generation_attempt_no or 1)
mode_snapshot = owner_mode(owner)
stage_snapshot = str(owner.pipeline_stage or "")
await db.commit()
await notify_owner_finished(db, owner)
await db.commit()
await remove_poll_active(owner)
fresh_owner = await _reload_owner_after_commit(
db,
owner_type=owner_type_snapshot,
owner_id=owner_id_snapshot,
attempt_no=attempt_snapshot,
)
if fresh_owner is not None:
await notify_owner_finished(db, fresh_owner)
await remove_poll_active(
owner_type=owner_type_snapshot,
owner_id=owner_id_snapshot,
attempt_no=attempt_snapshot,
)
await log_task_event(
owner,
owner_type=owner_type_snapshot,
owner_id=owner_id_snapshot,
generation_attempt_no=attempt_snapshot,
generation_mode=mode_snapshot,
event_type=event_type,
message=message,
detail=detail,
to_stage=owner.pipeline_stage,
to_stage=stage_snapshot,
)
@@ -302,15 +320,30 @@ async def _schedule_next_poll(
owner.poll_interval_seconds = schedule.poll_interval_seconds
owner.poll_claim_token = None
owner.poll_lease_until = None
owner_type_snapshot = owner_type_of(owner)
owner_id_snapshot = str(owner.id)
attempt_snapshot = int(owner.generation_attempt_no or 1)
mode_snapshot = owner_mode(owner)
await db.commit()
fresh_owner = await _reload_owner_after_commit(
db,
owner_type=owner_type_snapshot,
owner_id=owner_id_snapshot,
attempt_no=attempt_snapshot,
)
if fresh_owner is None:
return
await register_poll_active(
owner,
fresh_owner,
check_at=schedule.next_poll_at,
next_poll_at=schedule.next_poll_at,
reason=schedule.reason,
)
await log_task_event(
owner,
owner_type=owner_type_snapshot,
owner_id=owner_id_snapshot,
generation_attempt_no=attempt_snapshot,
generation_mode=mode_snapshot,
event_type=ChatGenerationTaskEventType.POLL_SCHEDULED.value,
message=f"已登记下一次轮询。reason={schedule.reason}",
detail={
@@ -320,10 +353,10 @@ async def _schedule_next_poll(
)
if schedule.direct_countdown:
poll_generation_task.apply_async(
args=[owner.id],
args=[owner_id_snapshot],
kwargs={
"owner_type": owner_type_of(owner),
"generation_attempt_no": int(owner.generation_attempt_no or 1),
"owner_type": owner_type_snapshot,
"generation_attempt_no": attempt_snapshot,
"force_due": False,
},
queue=POLL_QUEUE,
@@ -381,13 +414,21 @@ async def _restore_after_lock_error(
owner.poll_claim_token = None
owner.poll_lease_until = None
owner.next_poll_at = _now()
next_poll_at = owner.next_poll_at
await db.commit()
await register_poll_active(
owner,
check_at=owner.next_poll_at,
next_poll_at=owner.next_poll_at,
reason="poll_execution_lock_error",
fresh_owner = await _reload_owner_after_commit(
db,
owner_type=owner_type,
owner_id=owner_id,
attempt_no=attempt_no,
)
if fresh_owner is not None:
await register_poll_active(
fresh_owner,
check_at=next_poll_at,
next_poll_at=next_poll_at,
reason="poll_execution_lock_error",
)
async def _run(
@@ -478,13 +519,21 @@ async def _run(
# A stale database lease left by a crashed worker must not block the
# worker that successfully acquired the current Redis lock.
if not force_due and is_poll_not_due(owner, now=current):
next_poll_at = owner.next_poll_at
await db.commit()
await register_poll_active(
owner,
check_at=owner.next_poll_at,
next_poll_at=owner.next_poll_at,
reason="poll_task_not_due",
owner = await _reload_owner_after_commit(
db,
owner_type=normalized_owner_type,
owner_id=task_id,
attempt_no=effective_attempt,
)
if owner is not None:
await register_poll_active(
owner,
check_at=next_poll_at,
next_poll_at=next_poll_at,
reason="poll_task_not_due",
)
return
final_poll = is_final_poll_due(owner, now=current)
@@ -507,12 +556,22 @@ async def _run(
owner.poll_lease_until = _poll_lease_until(current)
owner.poll_count = int(owner.poll_count or 0) + 1
owner.last_poll_at = current
poll_lease_until = owner.poll_lease_until
next_poll_at = owner.next_poll_at
await db.commit()
claim_started = True
owner = await _reload_owner_after_commit(
db,
owner_type=normalized_owner_type,
owner_id=task_id,
attempt_no=effective_attempt,
)
if owner is None:
return
await register_poll_active(
owner,
check_at=owner.poll_lease_until,
next_poll_at=owner.next_poll_at,
check_at=poll_lease_until,
next_poll_at=next_poll_at,
reason="polling_lease",
)
@@ -562,9 +621,6 @@ async def _run(
event_type=ChatGenerationTaskEventType.POLL_FAILED.value,
detail=poll_result,
)
await _log_poll_provider_call_after_commit(
owner, provider_response=provider_response
)
return
owner.pipeline_stage = _stage(
owner, ChatGenerationPipelineStage.RESULT_READY
@@ -573,16 +629,30 @@ async def _run(
owner.poll_claim_token = None
owner.poll_lease_until = None
owner.next_poll_at = None
success_stage = str(owner.pipeline_stage or "")
success_mode = owner_mode(owner)
await db.commit()
await _log_poll_provider_call_after_commit(
owner, provider_response=provider_response
owner = await _reload_owner_after_commit(
db,
owner_type=normalized_owner_type,
owner_id=task_id,
attempt_no=effective_attempt,
)
await remove_poll_active(
owner_type=normalized_owner_type,
owner_id=task_id,
attempt_no=effective_attempt,
)
await remove_poll_active(owner)
await log_task_event(
owner,
owner_type=normalized_owner_type,
owner_id=task_id,
generation_attempt_no=effective_attempt,
generation_mode=success_mode,
event_type=ChatGenerationTaskEventType.POLL_SUCCESS.value,
to_stage=owner.pipeline_stage,
to_stage=success_stage,
)
if owner is None:
return
from app.tasks.generation_download_tasks import (
enqueue_download_task,
)
@@ -603,9 +673,6 @@ async def _run(
event_type=ChatGenerationTaskEventType.POLL_FAILED.value,
detail=poll_result,
)
await _log_poll_provider_call_after_commit(
owner, provider_response=provider_response
)
return
if final_poll:
@@ -617,18 +684,12 @@ async def _run(
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
detail=poll_result,
)
await _log_poll_provider_call_after_commit(
owner, provider_response=provider_response
)
return
owner.poll_error_count = 0
await _schedule_next_poll(
db, owner, reason="poll_pending_next"
)
await _log_poll_provider_call_after_commit(
owner, provider_response=provider_response
)
await log_task_event(
owner,
event_type=ChatGenerationTaskEventType.POLL_PENDING.value,
@@ -27,12 +27,97 @@ async def _recover_generation_records_once(*, include_create: bool, include_poll
from app.tasks.generation_poll_tasks import poll_generation_task
from app.tasks.generation_download_tasks import enqueue_download_task
counts: dict[str, Any] = {"create": 0, "poll": 0, "download": 0, "errors": []}
counts: dict[str, Any] = {
"create": 0,
"poll": 0,
"download": 0,
"inconsistent": 0,
"inconsistent_timeout": 0,
"errors": [],
}
batch_size = max(1, int(settings.GENERATION_RECOVERY_BATCH_SIZE or 20))
cursor = None
async with async_session() as db:
while True:
batch = await find_generation_record_recovery_batch(db, limit=batch_size, cursor=cursor)
if include_create or include_poll:
from datetime import datetime, timezone
from app.enums.generation_status import GenerationRecordPipelineStage
from app.enums.generation_task import ChatGenerationTaskEventType, GenerationMode
from app.services.generation.log_service import log_task_event
from app.services.generation.pipeline.owner_service import load_generation_owner
from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once
from app.services.redis_registry_service import ensure_aware_utc
for ref in batch.inconsistent:
try:
owner = await load_generation_owner(
db,
owner_type=ref.owner_type,
owner_id=ref.owner_id,
for_update=True,
)
if (
owner is None
or int(owner.generation_attempt_no or 1)
!= int(ref.generation_attempt_no or 1)
):
await db.rollback()
continue
deadline_at = ensure_aware_utc(getattr(owner, "deadline_at", None))
owner_id_snapshot = str(owner.id)
attempt_snapshot = int(owner.generation_attempt_no or 1)
previous_stage = str(owner.pipeline_stage or "")
if deadline_at is not None and deadline_at <= datetime.now(timezone.utc):
owner.pipeline_stage = GenerationRecordPipelineStage.TIMEOUT.value
await mark_generation_record_failed_and_refund_once(
db,
record=owner,
generation_attempt_no=ref.generation_attempt_no,
error_message="恢复证据异常且已超过任务截止时间",
)
await db.commit()
await log_task_event(
owner_type=ref.owner_type,
owner_id=owner_id_snapshot,
generation_attempt_no=attempt_snapshot,
generation_mode=GenerationMode.GENERATION_RECORD.value,
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
from_stage=previous_stage,
to_stage=GenerationRecordPipelineStage.TIMEOUT.value,
message="恢复证据异常任务超过截止时间,已失败并幂等退款",
)
counts["inconsistent_timeout"] += 1
continue
if previous_stage != GenerationRecordPipelineStage.RECOVERY_INCONSISTENT.value:
owner.pipeline_stage = GenerationRecordPipelineStage.RECOVERY_INCONSISTENT.value
owner.error_message = (
f"恢复证据异常:阶段 {previous_stage} 缺少 remote_result_url 和供应商任务ID"
)
await db.commit()
await log_task_event(
owner_type=ref.owner_type,
owner_id=owner_id_snapshot,
generation_attempt_no=attempt_snapshot,
generation_mode=GenerationMode.GENERATION_RECORD.value,
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_INCONSISTENT.value,
from_stage=previous_stage,
to_stage=GenerationRecordPipelineStage.RECOVERY_INCONSISTENT.value,
message="恢复证据异常,已隔离且不重新创建供应商任务",
)
else:
await db.rollback()
counts["inconsistent"] += 1
except Exception as exc:
await db.rollback()
counts["errors"].append(
{
"owner_id": ref.owner_id,
"stage": "recovery_inconsistent",
"error": str(exc),
}
)
if include_create:
for ref in batch.create:
try:
File diff suppressed because one or more lines are too long
+3 -1
View File
@@ -1,3 +1,4 @@
<!doctype html>
<html lang="zh-CN">
<head>
@@ -28,10 +29,11 @@
}
})();
</script>
<script type="module" crossorigin src="/assets/index-M8UKiCMs.js"></script>
<script type="module" crossorigin src="/assets/index-BhfopyFy.js"></script>
<link rel="stylesheet" crossorigin href="/assets/index-Bsz_Xon-.css">
</head>
<body>
<div id="root"></div>
</body>
</html>
+82
View File
@@ -17,6 +17,7 @@
"qrcode.react": "^4.2.0",
"react": "^19.2.5",
"react-dom": "^19.2.5",
"react-quill-new": "^3.8.3",
"react-router-dom": "^7.15.0",
"three": "^0.184.0",
"uuid": "^14.0.0",
@@ -2826,6 +2827,12 @@
"node": ">=0.10.0"
}
},
"node_modules/eventemitter3": {
"version": "5.0.4",
"resolved": "https://registry.npmjs.org/eventemitter3/-/eventemitter3-5.0.4.tgz",
"integrity": "sha512-mlsTRyGaPBjPedk6Bvw+aqbsXDtoAyAzm5MO7JgU+yVRyMQ5O8bD4Kcci7BS85f93veegeCPkL8R4GLClnjLFw==",
"license": "MIT"
},
"node_modules/fast-deep-equal": {
"version": "3.1.3",
"resolved": "https://registry.npmjs.org/fast-deep-equal/-/fast-deep-equal-3.1.3.tgz",
@@ -2833,6 +2840,12 @@
"dev": true,
"license": "MIT"
},
"node_modules/fast-diff": {
"version": "1.3.0",
"resolved": "https://registry.npmjs.org/fast-diff/-/fast-diff-1.3.0.tgz",
"integrity": "sha512-VxPP4NqbUjj6MaAOafWeUn2cXWLcCtljklUtZf0Ind4XQ+QPtmA0b18zZy0jIQx+ExRVCR/ZQpBmik5lXshNsw==",
"license": "Apache-2.0"
},
"node_modules/fast-glob": {
"version": "3.3.3",
"resolved": "https://registry.npmjs.org/fast-glob/-/fast-glob-3.3.3.tgz",
@@ -3615,6 +3628,25 @@
"url": "https://github.com/sponsors/sindresorhus"
}
},
"node_modules/lodash-es": {
"version": "4.18.1",
"resolved": "https://registry.npmjs.org/lodash-es/-/lodash-es-4.18.1.tgz",
"integrity": "sha512-J8xewKD/Gk22OZbhpOVSwcs60zhd95ESDwezOFuA3/099925PdHJ7OFHNTGtajL3AlZkykD32HykiMo+BIBI8A==",
"license": "MIT"
},
"node_modules/lodash.clonedeep": {
"version": "4.5.0",
"resolved": "https://registry.npmjs.org/lodash.clonedeep/-/lodash.clonedeep-4.5.0.tgz",
"integrity": "sha512-H5ZhCF25riFd9uB5UCkVKo61m3S/xZk1x4wA6yp/L3RFP6Z/eHH1ymQcGLo7J3GMPfm0V/7m1tryHuGVxpqEBQ==",
"license": "MIT"
},
"node_modules/lodash.isequal": {
"version": "4.5.0",
"resolved": "https://registry.npmjs.org/lodash.isequal/-/lodash.isequal-4.5.0.tgz",
"integrity": "sha512-pDo3lu8Jhfjqls6GkMgpahsF9kCyayhgykjyLMNFTKWrpVdAQtYyB4muAMWozBB4ig/dtWAmsMxLEI8wuz+DYQ==",
"deprecated": "This package is deprecated. Use require('node:util').isDeepStrictEqual instead.",
"license": "MIT"
},
"node_modules/lru-cache": {
"version": "5.1.1",
"resolved": "https://registry.npmjs.org/lru-cache/-/lru-cache-5.1.1.tgz",
@@ -3822,6 +3854,12 @@
"integrity": "sha512-4hLB8Py4zZce5s4yd9XzopqwVv/yGNhV1Bl8NTmCq1763HeK2+EwVTv+leGeL13Dnh2wfbqowVPXCIO0z4taYw==",
"license": "(MIT AND Zlib)"
},
"node_modules/parchment": {
"version": "3.0.0",
"resolved": "https://registry.npmjs.org/parchment/-/parchment-3.0.0.tgz",
"integrity": "sha512-HUrJFQ/StvgmXRcQ1ftY6VEZUq3jA2t9ncFN4F84J/vN0/FPpQF+8FKXb3l6fLces6q0uOHj6NJn+2xvZnxO6A==",
"license": "BSD-3-Clause"
},
"node_modules/path-exists": {
"version": "4.0.0",
"resolved": "https://registry.npmjs.org/path-exists/-/path-exists-4.0.0.tgz",
@@ -4138,6 +4176,35 @@
],
"license": "MIT"
},
"node_modules/quill": {
"version": "2.0.3",
"resolved": "https://registry.npmjs.org/quill/-/quill-2.0.3.tgz",
"integrity": "sha512-xEYQBqfYx/sfb33VJiKnSJp8ehloavImQ2A6564GAbqG55PGw1dAWUn1MUbQB62t0azawUS2CZZhWCjO8gRvTw==",
"license": "BSD-3-Clause",
"dependencies": {
"eventemitter3": "^5.0.1",
"lodash-es": "^4.17.21",
"parchment": "^3.0.0",
"quill-delta": "^5.1.0"
},
"engines": {
"npm": ">=8.2.3"
}
},
"node_modules/quill-delta": {
"version": "5.1.0",
"resolved": "https://registry.npmjs.org/quill-delta/-/quill-delta-5.1.0.tgz",
"integrity": "sha512-X74oCeRI4/p0ucjb5Ma8adTXd9Scumz367kkMK5V/IatcX6A0vlgLgKbzXWy5nZmCGeNJm2oQX0d2Eqj+ZIlCA==",
"license": "MIT",
"dependencies": {
"fast-diff": "^1.3.0",
"lodash.clonedeep": "^4.5.0",
"lodash.isequal": "^4.5.0"
},
"engines": {
"node": ">= 12.0.0"
}
},
"node_modules/rangetouch": {
"version": "2.0.1",
"resolved": "https://registry.npmjs.org/rangetouch/-/rangetouch-2.0.1.tgz",
@@ -4189,6 +4256,21 @@
"integrity": "sha512-/LLMVyas0ljjAtoYiPqYiL8VWXzUUdThrmU5+n20DZv+a+ClRoevUzw5JxU+Ieh5/c87ytoTBV9G1FiKfNJdmg==",
"license": "MIT"
},
"node_modules/react-quill-new": {
"version": "3.8.3",
"resolved": "https://registry.npmjs.org/react-quill-new/-/react-quill-new-3.8.3.tgz",
"integrity": "sha512-c96PYqFTo0pI4R3e79B3rH9LUIce1kIQbmTBu/imJQZk8305ogyLyBqKKjG2UoInDlquXqePSzmBo2aVia3ttw==",
"license": "MIT",
"dependencies": {
"lodash-es": "^4.17.21",
"quill": "~2.0.3"
},
"peerDependencies": {
"quill-delta": "^5.1.0",
"react": "^16 || ^17 || ^18 || ^19",
"react-dom": "^16 || ^17 || ^18 || ^19"
}
},
"node_modules/react-router": {
"version": "7.15.0",
"resolved": "https://registry.npmjs.org/react-router/-/react-router-7.15.0.tgz",
+1
View File
@@ -19,6 +19,7 @@
"qrcode.react": "^4.2.0",
"react": "^19.2.5",
"react-dom": "^19.2.5",
"react-quill-new": "^3.8.3",
"react-router-dom": "^7.15.0",
"three": "^0.184.0",
"uuid": "^14.0.0",
+12 -9
View File
@@ -5,7 +5,7 @@
import { api, setToken, clearToken } from './client';
import * as mock from './mock';
import type {
User, CreditRecord, Project, GenerationRecord, OptimizeParams, GenerateParams, OptimizeResult,
User, CreditRecord, Project, GenerationRecord, OptimizeParams, OptimizeResult,
Industry, IndustryConfig, AdminUser, AdminStats, ModelConfig, SystemConfig, AdminNotification,
PrivatePortraitConfig, PrivatePortraitProjectListOut, PrivatePortraitProject, PrivatePortraitValidateSession,
PrivatePortraitProjectCreateWithValidateOut, PrivatePortraitAssetListOut, PrivatePortraitAsset, PrivatePortraitSelectableAssetListOut,
@@ -112,7 +112,11 @@ export async function optimizePrompt(
project_id: projectId,//项目id
gen_type:params.genType,//生成类型
prompt: params.prompt,
engine_id: params.engineId,
include_media_references: params.includeMediaReferences ?? false,
duration: params.duration,
aspect_ratio: params.aspectRatio || null,
resolution: params.resolution || null,
references: params.references || null,
idempotency_key: params.idempotencyKey || null,
image_size: params.image_size || null,
@@ -266,15 +270,14 @@ export async function deleteUpload(url: string): Promise<void> {
export async function updateRecordPrompt(recordId: string, optimizedPrompt: string): Promise<void> {
await api.put(`/generation-records/${recordId}/prompt`, { optimized_prompt: optimizedPrompt });
}
export async function generateVideo(recordId: string, params: GenerateParams): Promise<GenerationRecord> {
export async function generateVideo(recordId: string): Promise<GenerationRecord> {
if (USE_MOCK) return mock.mockGenerateVideo(recordId);
return api.post<GenerationRecord>(`/generation-records/${recordId}/generate`, {
engine_id: params.engineId,
include_media_references: params.includeMediaReferences ?? false,
aspect_ratio: params.aspectRatio,
resolution: params.resolution,
image_size: params.imageSize,
});
return api.post<GenerationRecord>(`/generation-records/${recordId}/generate`);
}
export async function retryGeneration(recordId: string): Promise<GenerationRecord> {
if (USE_MOCK) return mock.mockGenerateVideo(recordId);
return api.post<GenerationRecord>(`/generation-records/${recordId}/retry`);
}
// ── Credits ───────────────────────────────────────────────
export async function getCredits(page = 1, pageSize = 20): Promise<{ credits: number; records: CreditRecord[]; total: number }> {
+29 -9
View File
@@ -3,7 +3,6 @@ import type {
CreditRecord,
Project,
GenerationRecord,
GenerateParams,
OptimizeParams,
OptimizeResult,
LoginParams,
@@ -208,7 +207,7 @@ export async function mockOptimizePrompt(
await delay(1500);
const project = MOCK_PROJECTS.find((p) => p.id === projectId);
const cost = Math.round(80 + params.prompt.length * 0.5 + params.duration * 2);
const cost = Math.round(80 + params.prompt.length * 0.5 + (params.duration || 0) * 2);
if (currentUser) {
currentUser.credits -= cost;
@@ -239,18 +238,28 @@ export async function mockOptimizePrompt(
projectName: project?.name ?? '未知项目',
originalPrompt: params.prompt,
optimizedPrompt,
duration: params.duration,
aspectRatio: params.aspectRatio as any,
resolution: params.resolution as any,
duration: params.genType === 'video' ? params.duration : undefined,
aspectRatio: params.genType === 'video' ? params.aspectRatio as any : undefined,
resolution: params.genType === 'video' ? params.resolution as any : undefined,
status: 'prompt_optimized',
creditsCost: cost,
textCreditsCost: cost,
textTokensUsed: 0,
videoTokensUsed: 0,
imageSize: params.resolution || '1080p',
imageProportion: params.aspectRatio || '16:9',
imagePx: '1920x1080',
imageSize: params.genType === 'image' ? params.image_size : undefined,
imageProportion: params.genType === 'image' ? params.image_proportion : undefined,
imagePx: params.genType === 'image' ? params.image_px : undefined,
imageUrl: '',
engineId: params.engineId,
engineName: params.engineId,
engineSnapshot: { id: params.engineId, name: params.engineId },
includeMediaReferences: Boolean(params.includeMediaReferences),
configComplete: true,
canGenerate: true,
canRetry: false,
shouldPoll: false,
clientStatus: 'ready',
operationPhase: 'prompt',
createdAt: new Date().toLocaleString('zh-CN'),
};
@@ -265,7 +274,18 @@ export async function mockGenerateVideo(recordId: string): Promise<GenerationRec
if (!record) throw new Error('记录不存在');
record.status = 'completed';
record.videoUrl = `https://example.com/video-${recordId}.mp4`;
record.canGenerate = false;
record.canRetry = false;
record.shouldPoll = false;
record.clientStatus = 'success';
record.operationPhase = 'resource';
if (record.genType === 'image') {
record.imageUrl = `https://example.com/image-${recordId}.png`;
record.videoUrl = undefined;
} else {
record.videoUrl = `https://example.com/video-${recordId}.mp4`;
record.imageUrl = undefined;
}
record.generatedAt = new Date().toLocaleString('zh-CN');
return record;
@@ -1,6 +1,6 @@
import React, { useEffect, useState, useCallback } from 'react';
import { Tag, Button } from 'antd';
import { BellOutlined, ThunderboltOutlined, GiftOutlined, StarOutlined, CloseOutlined } from '@ant-design/icons';
import { BellOutlined, ThunderboltOutlined, GiftOutlined, StarOutlined, CloseOutlined, LeftOutlined, RightOutlined } from '@ant-design/icons';
import { getNotifications, markNotificationRead } from '../api';
import { useAuthStore } from '../store/useAuthStore';
@@ -105,94 +105,61 @@ const NotificationPopup: React.FC = () => {
background: 'rgba(0,0,0,0.45)', backdropFilter: 'blur(6px)',
}} onClick={handleClose}>
<div onClick={e => e.stopPropagation()} style={{
width: 420,
borderRadius: 24,
width: 520,
borderRadius: 24,
overflow: 'hidden',
boxShadow: '0 25px 80px rgba(0,0,0,0.15)',
boxShadow: '0 30px 100px rgba(0,0,0,0.2)',
animation: 'notificationPop 0.35s cubic-bezier(0.34, 1.56, 0.64, 1)',
background: '#ffffff',
}}>
{/* Header */}
<div style={{
background: tc.gradient,
padding: '24px 28px 20px',
position: 'relative',
padding: '16px 28px 14px',
position: 'relative',
overflow: 'hidden',
}}>
<div style={{
position: 'absolute',
right: -50,
top: -50,
width: 200,
height: 200,
borderRadius: '50%',
<div style={{
position: 'absolute',
right: -50,
top: -50,
width: 200,
height: 200,
borderRadius: '50%',
background: 'radial-gradient(circle, rgba(255,255,255,0.15) 0%, transparent 60%)',
filter: 'blur(30px)',
}} />
<div style={{
position: 'absolute',
right: 30,
top: 10,
width: 80,
height: 80,
borderRadius: '50%',
background: 'rgba(255,255,255,0.06)',
filter: 'blur(15px)',
}} />
<div style={{ display: 'flex', justifyContent: 'space-between', alignItems: 'flex-start', position: 'relative' }}>
<div style={{ display: 'flex', alignItems: 'center', gap: 12 }}>
<div style={{
width: 44,
height: 44,
borderRadius: 14,
background: 'rgba(255,255,255,0.15)',
display: 'flex',
alignItems: 'center',
justifyContent: 'center',
fontSize: 20,
color: '#fff',
}}>{tc.icon}</div>
<div>
<div style={{
color: 'rgba(255,255,255,0.7)',
fontSize: 12,
marginBottom: 4,
fontWeight: 500,
}}></div>
<div style={{
color: '#fff',
fontSize: 28,
fontWeight: 700,
letterSpacing: '-1px',
}}>{current.title}</div>
</div>
</div>
<button
onClick={handleClose}
{/* Top row: type badge + close */}
<div style={{ display: 'flex', justifyContent: 'space-between', alignItems: 'center', marginBottom: 10 }}>
<Tag style={{ background: 'rgba(255,255,255,0.2)', border: 'none', color: '#fff', borderRadius: 6, padding: '2px 10px', fontSize: 11, fontWeight: 500 }}>
{tc.icon} {tc.label}
</Tag>
<button
onClick={handleClose}
style={{
width: 30,
height: 30,
borderRadius: '50%',
background: 'rgba(255,255,255,0.1)',
border: 'none',
cursor: 'pointer',
display: 'flex',
alignItems: 'center',
justifyContent: 'center',
transition: 'all 0.2s ease',
outline: 'none',
}}
onMouseEnter={e => {
e.currentTarget.style.background = 'rgba(255,255,255,0.2)';
}}
onMouseLeave={e => {
e.currentTarget.style.background = 'rgba(255,255,255,0.1)';
width: 28, height: 28, borderRadius: '50%',
background: 'rgba(255,255,255,0.1)',
border: 'none', cursor: 'pointer',
display: 'flex', alignItems: 'center', justifyContent: 'center',
transition: 'all 0.2s ease', outline: 'none',
}}
onMouseEnter={e => { e.currentTarget.style.background = 'rgba(255,255,255,0.2)'; }}
onMouseLeave={e => { e.currentTarget.style.background = 'rgba(255,255,255,0.1)'; }}
>
<CloseOutlined style={{ color: '#fff', fontSize: 14 }} />
<CloseOutlined style={{ color: '#fff', fontSize: 12 }} />
</button>
</div>
{/* Title + time row */}
<div style={{ display: 'flex', justifyContent: 'space-between', alignItems: 'baseline', gap: 12 }}>
<div style={{ color: '#fff', fontSize: 22, fontWeight: 700, letterSpacing: '-0.5px' }}>
{current.title}
</div>
<div style={{ color: 'rgba(255,255,255,0.6)', fontSize: 11, whiteSpace: 'nowrap', flexShrink: 0 }}>
{current.createdAt ? new Date(current.createdAt).toLocaleString('zh-CN') : ''}
</div>
</div>
</div>
{/* Body */}
@@ -201,48 +168,67 @@ const NotificationPopup: React.FC = () => {
padding: '24px 28px 28px',
}}>
<div style={{
padding: '18px 20px',
background: '#f8fafc',
padding: '18px 20px',
background: '#f8fafc',
borderRadius: 16,
marginBottom: 20,
marginBottom: 20,
border: '1px solid #f1f5f9',
}}>
<div style={{
fontSize: 14,
color: '#475569',
<div style={{
fontSize: 14,
color: '#475569',
lineHeight: 1.7,
}}>
{current.content}
</div>
}} dangerouslySetInnerHTML={{ __html: current.content }} />
</div>
<div style={{ display: 'flex', justifyContent: 'space-between', alignItems: 'center' }}>
{/* Left: navigation with left/right buttons */}
<div style={{ display: 'flex', alignItems: 'center', gap: 8 }}>
<Tag
color={tc.gradient.includes('#6366f1') ? '#6366f1' : tc.gradient.includes('#f59e0b') ? '#f59e0b' : '#8b5cf6'}
style={{
borderRadius: 6,
border: 'none',
padding: '3px 10px',
fontSize: 11,
fontWeight: 500,
background: '#f1f5f9',
color: '#64748b',
}}
>
{tc.label}
</Tag>
{notifications.length > 1 && (
<>
<button
onClick={(e) => { e.stopPropagation(); setCurrentIndex(currentIndex > 0 ? currentIndex - 1 : notifications.length - 1); }}
style={{
width: 30, height: 30, borderRadius: '50%',
background: '#f1f5f9', border: 'none', cursor: 'pointer',
display: 'flex', alignItems: 'center', justifyContent: 'center',
transition: 'all 0.2s', outline: 'none',
}}
onMouseEnter={e => { e.currentTarget.style.background = '#e2e8f0'; }}
onMouseLeave={e => { e.currentTarget.style.background = '#f1f5f9'; }}
>
<LeftOutlined style={{ color: '#64748b', fontSize: 12 }} />
</button>
<span style={{ fontSize: 12, color: '#64748b', fontWeight: 500, minWidth: 36, textAlign: 'center' }}>
{currentIndex + 1}/{notifications.length}
</span>
<button
onClick={(e) => { e.stopPropagation(); setCurrentIndex(currentIndex < notifications.length - 1 ? currentIndex + 1 : 0); }}
style={{
width: 30, height: 30, borderRadius: '50%',
background: '#f1f5f9', border: 'none', cursor: 'pointer',
display: 'flex', alignItems: 'center', justifyContent: 'center',
transition: 'all 0.2s', outline: 'none',
}}
onMouseEnter={e => { e.currentTarget.style.background = '#e2e8f0'; }}
onMouseLeave={e => { e.currentTarget.style.background = '#f1f5f9'; }}
>
<RightOutlined style={{ color: '#64748b', fontSize: 12 }} />
</button>
</>
)}
</div>
<Button
type="primary"
onClick={handleAcknowledge}
{/* Right: acknowledge button */}
<Button
type="primary"
onClick={handleAcknowledge}
style={{
borderRadius: 12,
fontWeight: 600,
borderRadius: 12,
fontWeight: 600,
height: 40,
background: tc.gradient,
background: tc.gradient,
border: 'none',
paddingLeft: 28,
paddingLeft: 28,
paddingRight: 28,
boxShadow: `0 6px 16px ${tc.glowColor}`,
transition: 'all 0.2s ease',
@@ -261,28 +247,6 @@ const NotificationPopup: React.FC = () => {
</Button>
</div>
{notifications.length > 1 && (
<div style={{ display: 'flex', justifyContent: 'center', gap: 6, marginTop: 20 }}>
{notifications.map((_, i) => (
<button
key={i}
onClick={() => setCurrentIndex(i)}
style={{
width: i === currentIndex ? 20 : 5,
height: 5,
borderRadius: 3,
background: i === currentIndex
? (tc.gradient.includes('#6366f1') ? '#6366f1' : tc.gradient.includes('#f59e0b') ? '#f59e0b' : '#8b5cf6')
: '#e2e8f0',
border: 'none',
cursor: 'pointer',
transition: 'all 0.3s ease',
outline: 'none',
}}
/>
))}
</div>
)}
</div>
</div>
<style>{`
+181 -440
View File
@@ -502,6 +502,7 @@ const GeneratePage: React.FC = () => {
fetchRecords,
optimizePrompt,
generateVideo,
retryGeneration,
} = useAppStore();
const recordItems = records.items;
const { user } = useAuthStore();
@@ -635,11 +636,6 @@ const GeneratePage: React.FC = () => {
Record<string, string>
>({});
const [expandedGroup, setExpandedGroup] = useState<string | null>(null);
// Per-record param selections for history prompt_optimized records
const [historyParams, setHistoryParams] = useState<
Record<string, { aspectRatio?: AspectRatio; resolution?: Resolution; engineId?: string; includeMediaReferences?: boolean }>
>({});
// Video preview modal
const [previewVideoUrl, setPreviewVideoUrl] = useState<string | null>(null);
@@ -1301,7 +1297,11 @@ const GeneratePage: React.FC = () => {
const pending = localStorage.getItem("pending_optimize");
if (!pending) return;
try {
const { key, projectId: pId, prompt, duration, ts } = JSON.parse(pending);
const {
key, projectId: pId, prompt, duration, engineId, genType, aspectRatio,
resolution, imageSize, imageProportion, imagePx, includeMediaReferences: pendingIncludeReferences,
referenceIds, ts,
} = JSON.parse(pending);
if (Date.now() - ts > 5 * 60 * 1000 || pId !== projectId) {
localStorage.removeItem("pending_optimize");
return;
@@ -1314,6 +1314,15 @@ const GeneratePage: React.FC = () => {
r.projectId === pId &&
r.originalPrompt === prompt &&
r.duration === duration &&
r.engineId === engineId &&
String(r.genType || "video") === String(genType || "video") &&
String(r.aspectRatio || "") === String(aspectRatio || "") &&
String(r.resolution || "") === String(resolution || "") &&
String(r.imageSize || "") === String(imageSize || "") &&
String(r.imageProportion || "") === String(imageProportion || "") &&
String(r.imagePx || "") === String(imagePx || "") &&
Boolean(r.includeMediaReferences) === Boolean(pendingIncludeReferences) &&
JSON.stringify((r.references || []).map((item) => item.upload_resource_id || item.url).filter(Boolean).sort()) === JSON.stringify((referenceIds || []).slice().sort()) &&
r.status === "prompt_optimized" &&
new Date(r.createdAt).getTime() > ts - 10000,
);
@@ -1370,8 +1379,19 @@ const GeneratePage: React.FC = () => {
// Media credits estimate for step 2 (video or image)
const estimatedVideoCredits = mediaType === "image"
? getImageCreditsFromCimage(selectedResolution)
: calcVideoCredits(videoDuration, videoResolution);
? getImageCreditsFromCimage(
currentRecord?.imageSize || selectedResolution,
currentRecord?.engineId || selectedEngineId,
currentRecord?.includeMediaReferences ?? includeMediaReferences,
currentRecord?.references || references,
)
: calcVideoCredits(
currentRecord?.duration || videoDuration,
(currentRecord?.resolution || videoResolution) as Resolution,
currentRecord?.engineId || selectedEngineId,
currentRecord?.includeMediaReferences ?? includeMediaReferences,
currentRecord?.references || references,
);
const canAffordVideo = userCredits >= estimatedVideoCredits;
// Step 1: Optimize prompt (text credits)
@@ -1386,6 +1406,14 @@ const GeneratePage: React.FC = () => {
message.error("请输入视频/图片描述");
return;
}
if (!selectedEngineId) {
message.error("请选择生成引擎");
return;
}
if (mediaType === "video" && (!videoAspectRatio || !videoResolution)) {
message.error("请选择视频比例和分辨率");
return;
}
setOptimizing(true);
const optionEntries = Object.entries(selectedOptions).map(
@@ -1404,20 +1432,32 @@ const GeneratePage: React.FC = () => {
projectId,
prompt: fullPrompt,
duration: videoDuration,
engineId: selectedEngineId,
genType: mediaType,
aspectRatio: mediaType === "video" ? videoAspectRatio : "",
resolution: mediaType === "video" ? videoResolution : "",
imageSize: mediaType === "image" ? selectedResolution : "",
imageProportion: mediaType === "image" ? selectedRatio : "",
imagePx: mediaType === "image" ? `${width}x${height}` : "",
includeMediaReferences,
referenceIds: references.map((item) => item.upload_resource_id || item.url).filter(Boolean).sort(),
ts: Date.now(),
}),
);
inFlightOptimizeKey.current = idempotencyKey;
const result = await optimizePrompt(projectId, {
prompt: fullPrompt,
duration: videoDuration,
duration: mediaType === "video" ? videoDuration : undefined,
genType: mediaType,
resolution: selectedResolution,
engineId: selectedEngineId,
includeMediaReferences,
aspectRatio: mediaType === "video" ? videoAspectRatio : undefined,
resolution: mediaType === "video" ? videoResolution : undefined,
references: references.length > 0 ? references : undefined,
idempotencyKey,
image_size: selectedResolution,
image_proportion: selectedRatio,
image_px: width + "x" + height,
image_size: mediaType === "image" ? selectedResolution : undefined,
image_proportion: mediaType === "image" ? selectedRatio : undefined,
image_px: mediaType === "image" ? width + "x" + height : undefined,
});
// console.log("按钮触发", result);
@@ -1542,7 +1582,7 @@ const GeneratePage: React.FC = () => {
const frontendState = recordStates[item.id];
// 如果后端状态不是generating,说明任务已经完成(成功或失败),需要处理
// 不管前端当前是什么状态,都要处理完成的任务
if (latest.status !== "generating") {
if (!(latest.shouldPoll ?? latest.status === "generating")) {
completedIds.push(item.id);
} else if (frontendState !== "generating") {
// 如果后端状态是generating,但前端不是,更新为generating(处理页面刷新后状态丢失的情况)
@@ -1617,7 +1657,7 @@ const GeneratePage: React.FC = () => {
}, []);
useEffect(() => {
const generatingRecords = recordItems.filter((r) => r.status === "generating");
const generatingRecords = recordItems.filter((r) => r.shouldPoll ?? r.status === "generating");
generatingRecords.forEach((record) => {
setRecordStates((p) => ({ ...p, [record.id]: "generating" }));
startPolling(record.id);
@@ -1625,10 +1665,6 @@ const GeneratePage: React.FC = () => {
}, [recordItems]);
const handleGenerate = async (recordId: string) => {
if (!canAffordVideo) {
message.error("积分不足,请先充值");
return;
}
setGenerating((p) => ({ ...p, [recordId]: true }));
setRecordStates((p) => ({ ...p, [recordId]: "generating" }));
message.loading({
@@ -1637,13 +1673,15 @@ const GeneratePage: React.FC = () => {
key: recordId,
});
try {
const result = await generateVideo(recordId, {
engineId: selectedEngineId || undefined,
includeMediaReferences,
aspectRatio: videoAspectRatio,
resolution: videoResolution,
imageSize: currentRecord?.imageSize || selectedResolution,
});
const record = projectRecords.find((item) => item.id === recordId) || currentRecord;
if (!record || record.canGenerate === false) {
throw new Error(record?.configFallbackHint || "该记录当前不可生成,请重新生成提词");
}
const nextPrompt = record.id === currentRecord?.id ? editedPrompt : editablePrompts[recordId];
if (nextPrompt && nextPrompt !== record.optimizedPrompt) {
await updateRecordPrompt(recordId, nextPrompt);
}
const result = await generateVideo(recordId);
if (result.status === "failed") {
setRecordStates((p) => ({ ...p, [recordId]: "failed" }));
message.error({
@@ -1698,18 +1736,10 @@ const GeneratePage: React.FC = () => {
key: recordId,
});
try {
// Save edited prompt first if changed
const editedPrompt = editablePrompts[recordId];
if (editedPrompt && editedPrompt !== record.optimizedPrompt) {
await updateRecordPrompt(recordId, editedPrompt);
if (!record.configComplete || record.canRetry === false) {
throw new Error("该失败记录不能直接重试,请重新生成提词");
}
const result = await generateVideo(recordId, {
engineId: record.engineId || selectedEngineId || undefined,
includeMediaReferences: Boolean(record.includeMediaReferences),
aspectRatio: record.aspectRatio || "16:9",
resolution: record.resolution || "720p",
imageSize: record.imageSize,
});
const result = await retryGeneration(recordId);
if (result.status === "failed") {
setRecordStates((p) => ({ ...p, [recordId]: "failed" }));
@@ -1758,6 +1788,10 @@ const GeneratePage: React.FC = () => {
return resolveGenerationUiState({
status: localStatus || record.status,
pipelineStage: localState ? null : record.pipelineStage,
clientStatus: localState ? null : record.clientStatus,
shouldPoll: localState ? localState === "generating" : record.shouldPoll,
canGenerate: record.canGenerate,
canRetry: record.canRetry,
});
};
@@ -2792,6 +2826,36 @@ const GeneratePage: React.FC = () => {
)}
</div>
)}
<div
style={{
display: "flex",
alignItems: "center",
gap: 10,
flexWrap: "wrap",
marginRight: 16,
}}
>
<Select
value={selectedEngineId || undefined}
onChange={handleGenerationEngineChange}
placeholder="选择生成引擎"
style={{ minWidth: 180 }}
options={(mediaType === "image" ? imageEngines : videoEngines).map((item: any) => ({ value: item.id, label: item.name }))}
/>
{mediaType === "video" && (
<>
<Select value={videoAspectRatio} onChange={(value) => setVideoAspectRatio(value as AspectRatio)} style={{ width: 110 }} options={engineOptions.ratios.map((value) => ({ value, label: value }))} />
<Select value={videoResolution} onChange={(value) => setVideoResolution(value as Resolution)} style={{ width: 110 }} options={engineOptions.resolutions.map((value) => ({ value, label: value }))} />
</>
)}
{references.length > 0 && (
<div style={{ display: "flex", alignItems: "center", gap: 6 }}>
<Switch size="small" checked={includeMediaReferences} onChange={setIncludeMediaReferences} />
<Typography.Text style={{ fontSize: 12 }}></Typography.Text>
</div>
)}
</div>
<div
style={{
display: "flex",
@@ -3273,57 +3337,6 @@ const GeneratePage: React.FC = () => {
</div>
)}
{/* Video params selection */}
{/* {mediaType !== 'image' && (
<div style={{ marginTop: 24 }}>
<Typography.Text style={{ fontSize: 13, color: '#1a1a2e', fontWeight: 600, display: 'block', marginBottom: 12 }}></Typography.Text>
<div style={{ display: 'flex', alignItems: 'center', gap: 12, flexWrap: 'wrap' }}>
<div style={{ display: 'flex', alignItems: 'center', gap: 6, padding: '5px 12px', borderRadius: 8, background: '#f8f9fc', border: '1px solid #e2e8f0' }}>
<Typography.Text style={{ fontSize: 12, color: '#94a3b8' }}></Typography.Text>
<Typography.Text strong style={{ fontSize: 13, color: '#6366f1' }}>{videoDuration}s</Typography.Text>
</div>
<PortalDropdown label="比例" value={videoAspectRatio}
options={engineOptions.ratios}
expanded={expandedEngine === 'ratio'}
onToggle={() => setExpandedEngine(expandedEngine === 'ratio' ? null : 'ratio')}
onSelect={(v) => setVideoAspectRatio(v as AspectRatio)}
onClose={() => setExpandedEngine(null)}
/>
<PortalDropdown label="分辨率" value={videoResolution}
options={engineOptions.resolutions}
expanded={expandedEngine === 'resolution'}
onToggle={() => setExpandedEngine(expandedEngine === 'resolution' ? null : 'resolution')}
onSelect={(v) => setVideoResolution(v as Resolution)}
onClose={() => setExpandedEngine(null)}
/>
</div>
</div>
)} */}
{/* Image params selection
{mediaType === 'image' && currentRecord && (
<div style={{ marginTop: 24, padding: 16, borderRadius: 12, background: '#f8f9fc' }}>
<div style={{ display: 'flex', justifyContent: 'space-around' }}>
<div style={{ textAlign: 'center' }}>
<Typography.Text style={{ fontSize: 11, color: '#94a3b8', display: 'block', marginBottom: 4 }}></Typography.Text>
<Typography.Text strong style={{ fontSize: 14, color: '#1a1a2e' }}>{currentRecord.imageSize || '-'}</Typography.Text>
</div>
<div style={{ textAlign: 'center' }}>
<Typography.Text style={{ fontSize: 11, color: '#94a3b8', display: 'block', marginBottom: 4 }}></Typography.Text>
<Typography.Text strong style={{ fontSize: 14, color: '#1a1a2e' }}>{currentRecord.imageProportion || '-'}</Typography.Text>
</div>
<div style={{ textAlign: 'center' }}>
<Typography.Text style={{ fontSize: 11, color: '#94a3b8', display: 'block', marginBottom: 4 }}></Typography.Text>
<Typography.Text strong style={{ fontSize: 14, color: '#1a1a2e' }}>{currentRecord.imagePx || '-'}</Typography.Text>
</div>
<div style={{ textAlign: 'center' }}>
<Typography.Text style={{ fontSize: 11, color: '#94a3b8', display: 'block', marginBottom: 4 }}></Typography.Text>
<Typography.Text strong style={{ fontSize: 14, color: '#1a1a2e' }}>-</Typography.Text>
</div>
</div>
</div>
)} */}
{/* Video credits + generate button */}
<div
style={{
@@ -3340,17 +3353,27 @@ const GeneratePage: React.FC = () => {
<div style={{ minWidth: 220, marginRight: 16 }}>
<Typography.Text style={{ fontSize: 12, color: '#64748b', display: 'block', marginBottom: 6 }}></Typography.Text>
<Select
value={selectedEngineId || undefined}
onChange={handleGenerationEngineChange}
value={currentRecord.engineId || undefined}
disabled
style={{ width: '100%' }}
options={(mediaType === 'image' ? imageEngines : videoEngines).map((item: any) => ({ value: item.id, label: item.name }))}
/>
{currentRecord?.references?.length ? (
<div style={{ marginTop: 10, display: 'flex', alignItems: 'center', gap: 8 }}>
<Switch size="small" checked={includeMediaReferences} onChange={setIncludeMediaReferences} />
<Typography.Text style={{ fontSize: 12, color: '#64748b' }}></Typography.Text>
<Switch size="small" checked={Boolean(currentRecord.includeMediaReferences)} disabled />
<Typography.Text style={{ fontSize: 12, color: '#64748b' }}></Typography.Text>
</div>
) : null}
{currentRecord.configComplete === false && (
<Typography.Text
type={currentRecord.canGenerate === false ? "danger" : "warning"}
style={{ fontSize: 12, display: "block", marginTop: 8 }}
>
{currentRecord.configFallbackHint || (currentRecord.canGenerate === false
? "旧版本配置不完整,请重新生成提词"
: "旧版本配置缺失,提交生成时将由后端自动补齐")}
</Typography.Text>
)}
</div>
{/* Video params selection */}
@@ -3372,7 +3395,7 @@ const GeneratePage: React.FC = () => {
marginBottom: 12,
}}
>
</Typography.Text>
<div
style={{
@@ -3400,33 +3423,25 @@ const GeneratePage: React.FC = () => {
strong
style={{ fontSize: 13, color: "#6366f1" }}
>
{videoDuration}s
{currentRecord.duration || videoDuration}s
</Typography.Text>
</div>
<PortalDropdown
label="比例"
value={videoAspectRatio}
options={engineOptions.ratios}
expanded={expandedEngine === "ratio"}
onToggle={() =>
setExpandedEngine(
expandedEngine === "ratio" ? null : "ratio",
)
}
onSelect={(v) => setVideoAspectRatio(v as AspectRatio)}
value={(currentRecord.aspectRatio || videoAspectRatio) as string}
options={[String(currentRecord.aspectRatio || videoAspectRatio)]}
expanded={false}
onToggle={() => undefined}
onSelect={() => undefined}
onClose={() => setExpandedEngine(null)}
/>
<PortalDropdown
label="分辨率"
value={videoResolution}
options={engineOptions.resolutions}
expanded={expandedEngine === "resolution"}
onToggle={() =>
setExpandedEngine(
expandedEngine === "resolution" ? null : "resolution",
)
}
onSelect={(v) => setVideoResolution(v as Resolution)}
value={(currentRecord.resolution || videoResolution) as string}
options={[String(currentRecord.resolution || videoResolution)]}
expanded={false}
onToggle={() => undefined}
onSelect={() => undefined}
onClose={() => setExpandedEngine(null)}
/>
</div>
@@ -3533,9 +3548,7 @@ const GeneratePage: React.FC = () => {
strong
style={{ fontSize: 14, color: "#10b981" }}
>
{mediaType === "image" && currentRecord.imageSize
? getImageCreditsFromCimage(currentRecord.imageSize) || estimatedVideoCredits
: calcVideoCredits(videoDuration, videoResolution)}
{estimatedVideoCredits}
</Typography.Text>
</div>
<Typography.Text style={{ color: "#cbd5e1", fontSize: 18 }}>
@@ -3551,7 +3564,7 @@ const GeneratePage: React.FC = () => {
strong
style={{ fontSize: 16, color: "#1a1a2e" }}
>
{(lastTextCredits + (mediaType === "image" ? getImageCreditsFromCimage(currentRecord.imageSize) : calcVideoCredits(videoDuration, videoResolution))).toFixed(2)}
{(lastTextCredits + estimatedVideoCredits).toFixed(2)}
</Typography.Text>
</div>
</div>
@@ -3587,6 +3600,7 @@ const GeneratePage: React.FC = () => {
loading={recordStates[currentRecord.id] === "generating"}
disabled={
!canAffordVideo ||
currentRecord.canGenerate === false ||
recordStates[currentRecord.id] === "generating" ||
recordStates[currentRecord.id] === "done"
}
@@ -3611,7 +3625,20 @@ const GeneratePage: React.FC = () => {
? "生成完成"
: recordStates[currentRecord.id] === "generating"
? "生成中..."
: `生成${mediaType === "image" ? "图片" : "视频"} (${ (mediaType === "image" ? getImageCreditsFromCimage(currentRecord.imageSize) : calcVideoCredits(videoDuration, videoResolution))}积分)`}
: `生成${mediaType === "image" ? "图片" : "视频"} (${(mediaType === "image"
? getImageCreditsFromCimage(
currentRecord.imageSize || selectedResolution,
currentRecord.engineId,
currentRecord.includeMediaReferences,
currentRecord.references,
)
: calcVideoCredits(
currentRecord.duration || videoDuration,
(currentRecord.resolution || videoResolution) as Resolution,
currentRecord.engineId,
currentRecord.includeMediaReferences,
currentRecord.references,
))})`}
</Button>
</Tooltip>
</div>
@@ -3810,7 +3837,7 @@ const GeneratePage: React.FC = () => {
flexShrink: 0,
}}
>
</Tag>
)}
@@ -3824,7 +3851,7 @@ const GeneratePage: React.FC = () => {
flexShrink: 0,
}}
>
{record.duration ? `${record.imageSize}` : "-"} ·{" "}
{record.imageSize || "-"} ·{" "}
{record.imageProportion || "-"} ·{" "}
{record.imagePx || "-"} ·{" "}
<span className="date-display" translate="no">
@@ -4429,279 +4456,54 @@ const GeneratePage: React.FC = () => {
position: "relative",
}}
>
{/* 将参数选择和视频生成按钮放到视频视频框 */}
{/* 待生成记录只能读取提词阶段冻结的配置 */}
{status === "prompt_optimized" && (
<div
style={{
marginTop: 16,
position: "absolute",
top: 0,
left: 0,
width: "100%",
height: "100%",
backgroundColor: "transparent",
inset: 0,
padding: 16,
borderRadius: 12,
background: "#f8f9fc",
border: "1px solid #e2e8f0",
display: "flex",
flexDirection: "column",
justifyContent: "space-between",
}}
>
<div style={{ display: 'flex', alignItems: 'center', gap: 8, marginBottom: 12, flexWrap: 'wrap' }}>
<Select
size="small"
value={historyParams[record.id]?.engineId || record.engineId || (type === 'image' ? imageEngines[0]?.id : videoEngines[0]?.id)}
onChange={(value) => {
const options = type === 'video' ? getVideoEngineOptions(value) : null;
setHistoryParams((prev) => ({
...prev,
[record.id]: {
...prev[record.id],
engineId: value,
aspectRatio: type === 'video' ? options!.ratios[0] as AspectRatio : prev[record.id]?.aspectRatio,
resolution: type === 'video' ? options!.resolutions[0] as Resolution : prev[record.id]?.resolution,
},
}));
}}
options={(type === 'image' ? imageEngines : videoEngines).map((item: any) => ({ value: item.id, label: item.name }))}
style={{ minWidth: 160 }}
/>
{record.references?.length ? (
<><Switch size="small" checked={Boolean(historyParams[record.id]?.includeMediaReferences ?? record.includeMediaReferences)} onChange={(checked) => setHistoryParams((prev) => ({ ...prev, [record.id]: { ...prev[record.id], includeMediaReferences: checked, aspectRatio: prev[record.id]?.aspectRatio || '' as AspectRatio, resolution: prev[record.id]?.resolution || '' as Resolution } }))} /><Typography.Text style={{ fontSize: 12 }}></Typography.Text></>
) : null}
<div>
<Typography.Text strong style={{ display: "block", marginBottom: 10 }}></Typography.Text>
<Typography.Text style={{ display: "block", fontSize: 12, color: "#64748b" }}>{record.engineName || record.engineId || "配置缺失"}</Typography.Text>
{type === "video" ? (
<>
<Typography.Text style={{ display: "block", fontSize: 12, color: "#64748b" }}>{record.duration || "-"}</Typography.Text>
<Typography.Text style={{ display: "block", fontSize: 12, color: "#64748b" }}>{record.aspectRatio || "-"}</Typography.Text>
<Typography.Text style={{ display: "block", fontSize: 12, color: "#64748b" }}>{record.resolution || "-"}</Typography.Text>
</>
) : (
<>
<Typography.Text style={{ display: "block", fontSize: 12, color: "#64748b" }}>{record.imageSize || "-"}</Typography.Text>
<Typography.Text style={{ display: "block", fontSize: 12, color: "#64748b" }}>{record.imageProportion || "-"}</Typography.Text>
<Typography.Text style={{ display: "block", fontSize: 12, color: "#64748b" }}>{record.imagePx || "-"}</Typography.Text>
</>
)}
<Typography.Text style={{ display: "block", fontSize: 12, color: "#64748b" }}>{record.includeMediaReferences ? "生成时携带" : "仅用于提词或不携带"}</Typography.Text>
{record.configComplete === false && (
<Typography.Text type={record.canGenerate === false ? "danger" : "warning"} style={{ display: "block", marginTop: 8 }}>
{record.configFallbackHint || (record.canGenerate === false ? "旧版本配置不完整,请重新生成提词" : "旧版本配置缺失,提交生成时将由后端自动补齐")}
</Typography.Text>
)}
</div>
{type === "video" && (
<div
style={{
display: "flex",
alignItems: "center",
gap: 12,
flexWrap: "wrap",
marginBottom: 12,
}}
>
<div
translate="no"
style={{
display: "flex",
alignItems: "center",
gap: 6,
padding: "5px 12px",
borderRadius: 8,
background: "#f8f9fc",
border: "1px solid #e2e8f0",
marginLeft: 12,
}}
>
<Typography.Text
style={{ fontSize: 12, color: "#94a3b8" }}
>
</Typography.Text>
<Typography.Text
strong
style={{ fontSize: 13, color: "#6366f1" }}
>
{record.duration}
</Typography.Text>
</div>
<PortalDropdown
label="比例"
value={
historyParams[record.id]?.aspectRatio ||
"选择比例"
}
options={getVideoEngineOptions(historyParams[record.id]?.engineId || record.engineId || videoEngines[0]?.id).ratios}
expanded={
expandedEngine === `ratio-${record.id}`
}
onToggle={() =>
setExpandedEngine(
expandedEngine === `ratio-${record.id}`
? null
: `ratio-${record.id}`,
)
}
onSelect={(v) =>
setHistoryParams((prev) => ({
...prev,
[record.id]: {
...prev[record.id],
aspectRatio: v as AspectRatio,
resolution:
prev[record.id]?.resolution || "",
},
}))
}
onClose={() => setExpandedEngine(null)}
/>
<PortalDropdown
label="分辨率"
value={
historyParams[record.id]?.resolution ||
"选择分辨率"
}
options={getVideoEngineOptions(historyParams[record.id]?.engineId || record.engineId || videoEngines[0]?.id).resolutions}
expanded={
expandedEngine === `res-${record.id}`
}
onToggle={() =>
setExpandedEngine(
expandedEngine === `res-${record.id}`
? null
: `res-${record.id}`,
)
}
onSelect={(v) =>
setHistoryParams((prev) => ({
...prev,
[record.id]: {
...prev[record.id],
aspectRatio:
prev[record.id]?.aspectRatio || "",
resolution: v as Resolution,
},
}))
}
onClose={() => setExpandedEngine(null)}
/>
</div>
)}
<Button
type="primary"
size="large"
icon={<RocketOutlined />}
block
loading={generating[record.id]}
disabled={
type === "video" &&
(!historyParams[record.id]?.aspectRatio ||
!historyParams[record.id]?.resolution)
}
onClick={async () => {
const params = historyParams[record.id];
// 视频类型需要检查比例和分辨率
if (type === "video") {
if (
!params?.aspectRatio ||
!params?.resolution
) {
message.error("请选择比例和分辨率");
return;
}
if (
userCredits <
calcVideoCredits(
record.duration || 5,
params.resolution,
params.engineId || record.engineId || videoEngines[0]?.id,
Boolean(params.includeMediaReferences ?? record.includeMediaReferences),
record.references,
)
) {
message.error("积分不足,请先充值");
return;
}
}
setGenerating((p) => ({
...p,
[record.id]: true,
}));
setRecordStates((p) => ({
...p,
[record.id]: "generating",
}));
const mediaText =
type === "video" ? "视频" : "图片";
message.loading({
content: `${projectName}」正在生成${mediaText}...`,
duration: 0,
key: record.id,
});
try {
await generateVideo(record.id, {
engineId: params?.engineId || record.engineId || (type === 'image' ? imageEngines[0]?.id : videoEngines[0]?.id),
includeMediaReferences: Boolean(params?.includeMediaReferences ?? record.includeMediaReferences),
aspectRatio: params?.aspectRatio,
resolution: params?.resolution,
imageSize: record.imageSize,
});
startPolling(record.id);
message.success({
content: `${projectName}${mediaText}已提交,正在生成中...`,
key: record.id,
duration: 3,
});
} catch (error: any) {
setRecordStates((p) => ({
...p,
[record.id]: "failed",
}));
const errorMsg = error?.response?.data?.message || error?.message || `${mediaText}生成失败`;
message.error({
content: `${projectName}${errorMsg}`,
key: record.id,
duration: 3,
});
} finally {
setGenerating((p) => ({
...p,
[record.id]: false,
}));
}
}}
style={{
borderRadius: 12,
fontWeight: 600,
color: "#fff",
height: 48,
background:
"linear-gradient(135deg, #6366f1 0%, #8b5cf6 100%)",
border: "none",
position: "absolute",
bottom: 18,
left: 0,
}}
disabled={record.canGenerate === false}
onClick={() => handleGenerate(record.id)}
style={{ borderRadius: 12, fontWeight: 600, height: 48 }}
>
{type === "video" ? "视频" : "图片"}
{type === "video" ? (
<span
style={{
color: "#d2d2d2",
fontSize: "14px",
marginLeft: "6px",
}}
>
()
</span>
) : (
""
)}{" "}
{/* {type === "video" &&
historyParams[record.id]?.resolution
? `(${calcVideoCredits(
record.duration || 5,
historyParams[record.id].resolution as Resolution,
historyParams[record.id].engineId || record.engineId || videoEngines[0]?.id,
Boolean(historyParams[record.id].includeMediaReferences ?? record.includeMediaReferences),
record.references,
)})`
: ""} */}
{type === "video" &&
historyParams[record.id]?.resolution
? `(${calcVideoCredits(
record.duration || 5,
historyParams[record.id].resolution as Resolution,
historyParams[record.id].engineId || record.engineId || videoEngines[0]?.id,
Boolean(historyParams[record.id].includeMediaReferences ?? record.includeMediaReferences),
record.references,
)})`
: ""}
{type === "image" && record.imageSize
? `(${getImageCreditsFromCimage(
record.imageSize,
historyParams[record.id]?.engineId || record.engineId || imageEngines[0]?.id,
Boolean(historyParams[record.id]?.includeMediaReferences ?? record.includeMediaReferences),
record.references,
)})`
: type === "image"
? `(积分)`
: ""}
</Button>
</div>
)}
@@ -4925,73 +4727,12 @@ const GeneratePage: React.FC = () => {
<Typography.Text
style={{ color: "#94a3b8", fontSize: 13 }}
>
{type === "image"
? "待生成图片"
: status === "prompt_optimized"
? "待配置视频参数"
: "待生成视频"}
{type === "image" ? "待生成图片" : "待生成视频"}
</Typography.Text>
</div>
)}
</div>
{/* Generate buttons for history records */}
{/* {status === 'prompt_optimized' && (
<div style={{ marginTop: 16, width: '100%' }}>
<div style={{ display: 'flex', alignItems: 'center', gap: 12, flexWrap: 'wrap', marginBottom: 12 }}>
<div translate="no" style={{ display: 'flex', alignItems: 'center', gap: 6, padding: '5px 12px', borderRadius: 8, background: '#f8f9fc', border: '1px solid #e2e8f0' }}>
<Typography.Text style={{ fontSize: 12, color: '#94a3b8' }}></Typography.Text>
<Typography.Text strong style={{ fontSize: 13, color: '#6366f1' }}>{record.duration}</Typography.Text>
</div>
<PortalDropdown label="比例" value={historyParams[record.id]?.aspectRatio || '选择比例'}
options={engineOptions.ratios}
expanded={expandedEngine === `ratio-${record.id}`}
onToggle={() => setExpandedEngine(expandedEngine === `ratio-${record.id}` ? null : `ratio-${record.id}`)}
onSelect={(v) => setHistoryParams(prev => ({ ...prev, [record.id]: { ...prev[record.id], aspectRatio: v as AspectRatio, resolution: prev[record.id]?.resolution || '' } }))}
onClose={() => setExpandedEngine(null)}
/>
<PortalDropdown label="分辨率" value={historyParams[record.id]?.resolution || '选择分辨率'}
options={engineOptions.resolutions}
expanded={expandedEngine === `res-${record.id}`}
onToggle={() => setExpandedEngine(expandedEngine === `res-${record.id}` ? null : `res-${record.id}`)}
onSelect={(v) => setHistoryParams(prev => ({ ...prev, [record.id]: { ...prev[record.id], aspectRatio: prev[record.id]?.aspectRatio || '', resolution: v as Resolution } }))}
onClose={() => setExpandedEngine(null)}
/>
</div>
<Button type="primary" size="large" icon={<RocketOutlined />} block
loading={generating[record.id]}
disabled={!historyParams[record.id]?.aspectRatio || !historyParams[record.id]?.resolution}
onClick={async () => {
const params = historyParams[record.id];
if (!params?.aspectRatio || !params?.resolution) { message.error('请选择比例和分辨率'); return; }
if (userCredits < calcVideoCredits(record.duration || 5, params.resolution)) { message.error('积分不足,请先充值'); return; }
setGenerating((p) => ({ ...p, [record.id]: true }));
setRecordStates((p) => ({ ...p, [record.id]: 'generating' }));
message.loading({ content: `${projectName}」正在生成视频...`, duration: 0, key: record.id });
try {
await generateVideo(record.id, { engineId: params.engineId || record.engineId || videoEngines[0]?.id, includeMediaReferences: Boolean(params.includeMediaReferences ?? record.includeMediaReferences), aspectRatio: params.aspectRatio, resolution: params.resolution, imageSize: record.imageSize });
setRecordStates((p) => ({ ...p, [record.id]: 'done' }));
message.success({ content: `${projectName}」视频生成成功!`, key: record.id, duration: 3 });
} catch {
setRecordStates((p) => ({ ...p, [record.id]: 'failed' }));
message.error({ content: `${projectName}」视频生成失败`, key: record.id, duration: 3 });
} finally {
setGenerating((p) => ({ ...p, [record.id]: false }));
}
}}
style={{ borderRadius: 12, fontWeight: 600, height: 48, background: 'linear-gradient(135deg, #6366f1 0%, #8b5cf6 100%)', border: 'none' }}>
{historyParams[record.id]?.resolution ? `(${calcVideoCredits(record.duration || 5, historyParams[record.id].resolution)}积分)` : ''}
</Button>
</div>
)}
{statusState.isFailure && (
<div style={{ marginTop: 16 }}>
<Button type="primary" danger size="large" icon={<PlayCircleOutlined />} block
loading={generating[record.id]}
onClick={() => handleRetryGeneration(record.id)}
style={{ borderRadius: 12, fontWeight: 600, height: 48 }}></Button>
</div>
)} */}
</div>
</div>
)}
+1
View File
@@ -81,6 +81,7 @@ const MessagesPage: React.FC = () => {
dataIndex: 'content',
key: 'content',
ellipsis: true,
render: (v: string) => <div style={{ maxWidth: 300 }} dangerouslySetInnerHTML={{ __html: v }} />,
},
{
title: '状态',
+42 -96
View File
@@ -4,7 +4,6 @@ import {
Empty,
Input,
message,
Modal,
Pagination,
Select,
Space,
@@ -31,14 +30,14 @@ import {
} from '@ant-design/icons';
import { useNavigate } from 'react-router-dom';
import { useAppStore } from '../store/useAppStore';
import type { AspectRatio, Resolution } from '../types';
import { updateRecordPrompt } from '../api';
import { formatDate } from '../utils/formatDate';
import { copyToClipboard } from '../utils/clipboard';
import { resolveGenerationUiState } from '../utils/generationTaskStatus';
const RecordsPage: React.FC = () => {
const navigate = useNavigate();
const { records, projects, fetchRecords, fetchProjects, generateVideo } = useAppStore();
const { records, projects, fetchRecords, fetchProjects, generateVideo, retryGeneration } = useAppStore();
const recordItems = records.items;
const [currentPage, setCurrentPage] = useState(1);
const [pageSize, setPageSize] = useState(10);
@@ -48,8 +47,33 @@ const RecordsPage: React.FC = () => {
const [editablePrompts, setEditablePrompts] = useState<Record<string, string>>({});
const [editingRecordId, setEditingRecordId] = useState<string | null>(null);
const [expandedId, setExpandedId] = useState<string | null>(null);
// Generate modal for param selection
const [genModal, setGenModal] = useState<{ recordId: string; projectName: string; ratio: AspectRatio; resolution: Resolution } | null>(null);
const submitFrozenGeneration = async (record: any) => {
const type = record.genType || 'image';
const retry = record.canRetry === true;
if ((!retry && (record.canGenerate === false || !record.configComplete)) || (retry && !record.configComplete)) {
message.error('该记录生成配置不完整,请重新生成提词');
return;
}
setGenerating((prev) => ({ ...prev, [record.id]: true }));
message.loading({ content: `${record.projectName}」正在提交${type === 'video' ? '视频' : '图片'}生成...`, duration: 0, key: record.id });
try {
if (retry) {
await retryGeneration(record.id);
} else {
const editedPrompt = editablePrompts[record.id];
if (editedPrompt && editedPrompt !== record.optimizedPrompt) {
await updateRecordPrompt(record.id, editedPrompt);
}
await generateVideo(record.id);
}
message.success({ content: `${record.projectName}」已提交,正在生成中...`, key: record.id, duration: 3 });
} catch (error: any) {
message.error({ content: error?.response?.data?.detail || error?.message || '提交生成失败', key: record.id, duration: 3 });
} finally {
setGenerating((prev) => ({ ...prev, [record.id]: false }));
}
};
useEffect(() => {
fetchProjects();
@@ -64,48 +88,6 @@ const RecordsPage: React.FC = () => {
});
}, [fetchRecords, filterProject, filterStatus, currentPage, pageSize]);
const handleGenerate = async () => {
if (!genModal) return;
const { recordId, projectName, ratio, resolution } = genModal;
setGenerating((p) => ({ ...p, [recordId]: true }));
setGenModal(null);
message.loading({ content: `${projectName}」正在生成视频...`, duration: 0, key: recordId });
try {
await generateVideo(recordId, { aspectRatio: ratio, resolution });
message.success({ content: `${projectName}」视频生成成功!`, key: recordId, duration: 3 });
} catch {
message.error({ content: `${projectName}」视频生成失败`, key: recordId, duration: 3 });
} finally {
setGenerating((p) => ({ ...p, [recordId]: false }));
}
};
const openGenModal = (record: any) => {
const type: any = record.genType || 'image';
// 如果是图片类型,直接生成,不需要弹窗
if (type === 'image') {
setGenerating((p) => ({ ...p, [record.id]: true }));
message.loading({ content: `${record.projectName}」正在生成图片...`, duration: 0, key: record.id });
generateVideo(record.id, {}).then(() => {
message.success({ content: `${record.projectName}」图片生成成功!`, key: record.id, duration: 3 });
}).catch(() => {
message.error({ content: `${record.projectName}」图片生成失败`, key: record.id, duration: 3 });
}).finally(() => {
setGenerating((p) => ({ ...p, [record.id]: false }));
});
return;
}
// 如果是视频类型,显示参数选择弹窗
setGenModal({
recordId: record.id,
projectName: record.projectName,
ratio: (record.aspectRatio as AspectRatio) || '16:9',
resolution: (record.resolution as Resolution) || '720p',
});
};
const filtered = recordItems;
return (
@@ -222,7 +204,7 @@ const RecordsPage: React.FC = () => {
{/* Meta */}
{type === 'image' ? <span className="mobile-meta" translate="no" style={{ fontSize: 12, color: '#94a3b8', flexShrink: 0 }}>
{record.duration ? `${record.imageSize}` : '-'} · {record.imageProportion || '-'} · {record.imagePx || '-'} · <span className="date-display" translate="no">{formatDate(record.createdAt)}</span>
{record.imageSize || '-'} · {record.imageProportion || '-'} · {record.imagePx || '-'} · <span className="date-display" translate="no">{formatDate(record.createdAt)}</span>
</span> : <span className="mobile-meta" translate="no" style={{ fontSize: 12, color: '#94a3b8', flexShrink: 0 }}>
{record.duration ? `${record.duration}` : '-'} · {record.aspectRatio || '-'} · {record.resolution || '-'} · <span className="date-display" translate="no">{formatDate(record.createdAt)}</span>
</span>}
@@ -233,7 +215,8 @@ const RecordsPage: React.FC = () => {
{record.status === 'prompt_optimized' && (
<Button type="primary" size="small" icon={<PlayCircleOutlined />}
loading={isGenerating}
onClick={() => openGenModal(record)}
disabled={!record.configComplete || record.canGenerate === false}
onClick={() => submitFrozenGeneration(record)}
style={{
borderRadius: 8,
background: 'linear-gradient(135deg, #6366f1, #8b5cf6)',
@@ -250,7 +233,8 @@ const RecordsPage: React.FC = () => {
{uiState.isFailure && (
<Button type="primary" danger size="small" icon={<PlayCircleOutlined />}
loading={isGenerating}
onClick={() => openGenModal(record)}
disabled={!record.configComplete || record.canRetry === false}
onClick={() => submitFrozenGeneration(record)}
style={{ borderRadius: 8 }}>
</Button>
@@ -355,7 +339,7 @@ const RecordsPage: React.FC = () => {
{type === 'image' ? [
{ label: '分辨率', value: record.imageSize ? `${record.imageSize} ` : '-' },
{ label: '画面比例', value: record.imageProportion ? `${record.imageProportion} ` : '-' },
{ label: '画面尺寸', value: record.imageProportion ? `${record.imageProportion} ` : '-' },
{ label: '画面尺寸', value: record.imagePx ? `${record.imagePx} ` : '-' },
{ label: '消耗积分', value: record.creditsCost ? `${record.creditsCost} ` : '-', highlight: !!record.creditsCost },
].map((item, j) => (
<div key={j} style={{ flex: 1 }}>
@@ -364,8 +348,8 @@ const RecordsPage: React.FC = () => {
</div>
)) : [
{ label: '时长', value: record.duration ? `${record.duration}` : '-' },
{ label: '画面比例', value: record.aspectRatio || (record.status === 'prompt_optimized' ? '待选择' : '-') },
{ label: '分辨率', value: record.resolution || (record.status === 'prompt_optimized' ? '待选择' : '-') },
{ label: '画面比例', value: record.aspectRatio || (record.status === 'prompt_optimized' ? '配置缺失' : '-') },
{ label: '分辨率', value: record.resolution || (record.status === 'prompt_optimized' ? '配置缺失' : '-') },
{ label: '消耗积分', value: record.creditsCost ? `${record.creditsCost}` : '-', highlight: !!record.creditsCost },
].map((item, j) => (
<div key={j} style={{ flex: 1 }}>
@@ -380,13 +364,14 @@ const RecordsPage: React.FC = () => {
<div style={{ marginTop: 16, display: 'flex', justifyContent: 'flex-end' }}>
<Button type="primary" size="large" icon={<RocketOutlined />}
loading={isGenerating}
onClick={() => openGenModal(record)}
disabled={!record.configComplete || record.canGenerate === false}
onClick={() => submitFrozenGeneration(record)}
style={{
borderRadius: 12, fontWeight: 600, height: 44,
background: 'linear-gradient(135deg, #6366f1 0%, #8b5cf6 100%)',
border: 'none', boxShadow: '0 8px 24px rgba(99,102,241,0.3)',
}}>
{type === 'video' ? '选择参数生成视频' : '生成图片'}
{type === 'video' ? '生成视频' : '生成图片'}
</Button>
</div>
)}
@@ -396,7 +381,8 @@ const RecordsPage: React.FC = () => {
<div style={{ marginTop: 16, display: 'flex', justifyContent: 'flex-end' }}>
<Button type="primary" danger size="large" icon={<PlayCircleOutlined />}
loading={isGenerating}
onClick={() => openGenModal(record)}
disabled={!record.configComplete || record.canRetry === false}
onClick={() => submitFrozenGeneration(record)}
style={{ borderRadius: 12, fontWeight: 600, height: 44 }}>
{type === 'video' ? '视频' : '图片'}
</Button>
@@ -509,46 +495,6 @@ const RecordsPage: React.FC = () => {
</div>
)}
{/* Generate modal */}
<Modal
title={<Space><RocketOutlined /></Space>}
open={!!genModal}
onCancel={() => setGenModal(null)}
onOk={handleGenerate}
okText="提交生成"
cancelText="取消"
width={420}
>
{genModal && (
<div style={{ display: 'flex', flexDirection: 'column', gap: 16, marginTop: 16 }}>
{(() => {
const rec = recordItems.find(r => r.id === genModal.recordId);
return (
<>
<div style={{ padding: 12, borderRadius: 10, background: '#f8f9fc' }}>
<Typography.Text style={{ fontSize: 11, color: '#94a3b8', display: 'block' }}></Typography.Text>
<Typography.Text strong>{rec?.duration || 5}</Typography.Text>
</div>
<div>
<Typography.Text style={{ fontSize: 12, color: '#64748b', display: 'block', marginBottom: 6 }}></Typography.Text>
<Select value={genModal.ratio} onChange={(v) => setGenModal(prev => prev ? { ...prev, ratio: v } : null)}
style={{ width: '100%' }}
options={['16:9', '4:3', '1:1', '3:4', '9:16', '21:9'].map(r => ({ value: r, label: r }))}
/>
</div>
<div>
<Typography.Text style={{ fontSize: 12, color: '#64748b', display: 'block', marginBottom: 6 }}></Typography.Text>
<Select value={genModal.resolution} onChange={(v) => setGenModal(prev => prev ? { ...prev, resolution: v } : null)}
style={{ width: '100%' }}
options={['480p', '720p', '1080p'].map(r => ({ value: r, label: r }))}
/>
</div>
</>
);
})()}
</div>
)}
</Modal>
</div>
);
};
+25 -6
View File
@@ -1,5 +1,5 @@
import { create } from 'zustand';
import type { Project, GenerationRecord, OptimizeParams, GenerateParams, OptimizeResult, Industry, MediaReference } from '../types';
import type { Project, GenerationRecord, OptimizeParams, OptimizeResult, Industry, MediaReference } from '../types';
import * as api from '../api';
import { useAuthStore } from './useAuthStore';
@@ -41,7 +41,8 @@ interface AppState {
fetchRecords: (params?: api.GetRecordsPageParams) => Promise<void>;
optimizePrompt: (projectId: string, params: OptimizeParams) => Promise<OptimizeResult>;
generateVideo: (recordId: string, params: GenerateParams) => Promise<GenerationRecord>;
generateVideo: (recordId: string) => Promise<GenerationRecord>;
retryGeneration: (recordId: string) => Promise<GenerationRecord>;
updateRecordReferences: (recordId: string, references: MediaReference[]) => void;
// 生成配置状态更新方法
@@ -126,18 +127,36 @@ export const useAppStore = create<AppState>((set, get) => ({
try { await useAuthStore.getState().checkAuth(); } catch { /* */ }
const currentRecords = get().records;
const existingIndex = currentRecords.items.findIndex((item) => item.id === result.record.id);
const nextItems = existingIndex >= 0
? currentRecords.items.map((item) => item.id === result.record.id ? result.record : item)
: [result.record, ...currentRecords.items];
set({
records: {
...currentRecords,
total: currentRecords.total + 1,
items: [result.record, ...currentRecords.items],
total: existingIndex >= 0 ? currentRecords.total : currentRecords.total + 1,
items: nextItems,
},
});
return result;
},
generateVideo: async (recordId, params) => {
const record = await api.generateVideo(recordId, params);
generateVideo: async (recordId) => {
const record = await api.generateVideo(recordId);
try { await useAuthStore.getState().checkAuth(); } catch { /* */ }
const currentRecords = get().records;
set({
records: {
...currentRecords,
items: currentRecords.items.map((r) => (r.id === recordId ? record : r)),
},
});
return record;
},
retryGeneration: async (recordId) => {
const record = await api.retryGeneration(recordId);
try { await useAuthStore.getState().checkAuth(); } catch { /* */ }
const currentRecords = get().records;
+20 -21
View File
@@ -181,7 +181,7 @@ export interface GenerationRecord {
originalPrompt: string;
optimizedPrompt?: string;
duration?: number;
genType?: number;
genType?: 'video' | 'image';
aspectRatio?: AspectRatio;
resolution?: Resolution;
status: GenerationStatus;
@@ -197,38 +197,37 @@ export interface GenerationRecord {
errorMessage?: string;
createdAt: string;
generatedAt?: string;
imageSize: string;
imageProportion: string;
imagePx: string;
imageUrl: string;
imageSize?: string;
imageProportion?: string;
imagePx?: string;
imageUrl?: string;
engineId?: string;
engineName?: string;
engineSnapshot?: Record<string, any>;
includeMediaReferences?: boolean;
configComplete?: boolean;
configRecoverable?: boolean;
configFallbackHint?: string;
canGenerate?: boolean;
canRetry?: boolean;
shouldPoll?: boolean;
clientStatus?: 'prompt_processing' | 'ready' | 'generating' | 'success' | 'failure';
operationPhase?: 'prompt' | 'resource';
}
export interface OptimizeParams {
prompt: string;
duration: number;
genType?: any;
duration?: number;
genType: 'video' | 'image';
engineId: string;
includeMediaReferences?: boolean;
resolution?: string;
aspectRatio?: string;
references?: MediaReference[];
idempotencyKey?: string;
image_size?: any;
image_proportion?: any;
image_px?: any;
video_duration?: number;
video_ratio?: string;
video_resolution?: string;
}
export interface GenerateParams {
engineId?: string;
includeMediaReferences?: boolean;
aspectRatio?: AspectRatio;
resolution?: Resolution;
imageSize?: string;
image_size?: string;
image_proportion?: string;
image_px?: string;
}
export type ModuleGenerationFlowVersion = 'v1' | 'v2';
+50 -85
View File
@@ -2,6 +2,10 @@ export interface GenerationStatusLike {
status?: string | null;
displayStatus?: string | null;
pipelineStage?: string | null;
clientStatus?: string | null;
shouldPoll?: boolean | null;
canGenerate?: boolean | null;
canRetry?: boolean | null;
}
export type GenerationUiColor = 'default' | 'processing' | 'warning' | 'success' | 'error' | 'blue' | 'orange' | 'purple';
@@ -17,42 +21,34 @@ export interface GenerationUiState {
isSuccess: boolean;
isFailure: boolean;
isTerminal: boolean;
canGenerate: boolean;
canRetry: boolean;
shouldPoll: boolean;
}
const ACTIVE_STATUS_KEYS = new Set([
'pending',
'optimizing',
'prompt_optimized',
'generating',
]);
const CLIENT_STATUS_MAP: Record<string, string> = {
prompt_processing: 'optimizing',
ready: 'prompt_optimized',
generating: 'generating',
success: 'completed',
failure: 'failed',
};
const ACTIVE_STATUS_KEYS = new Set(['pending', 'optimizing', 'generating']);
const ACTIVE_PIPELINE_STAGES = new Set([
'queued',
'preparing',
'creating_provider_task',
'provider_result_staged',
'waiting_remote',
'polling',
'result_ready',
'download_queued',
'downloading',
'retry_waiting',
'upscale_queued',
'upscale_processing',
'upscale_polling',
'upscale_downloading',
'upscale_finalizing',
'upscale_retry_waiting',
'queued', 'preparing', 'creating_provider_task', 'provider_result_staged',
'waiting_remote', 'polling', 'result_ready', 'download_queued', 'downloading',
'retry_waiting', 'recovery_inconsistent', 'upscale_queued', 'upscale_processing',
'upscale_polling', 'upscale_downloading', 'upscale_finalizing', 'upscale_retry_waiting',
]);
const SUCCESS_KEYS = new Set(['completed', 'done']);
const FAILURE_KEYS = new Set(['failed', 'timeout', 'download_failed', 'upscale_failed']);
const TERMINAL_KEYS = new Set([...SUCCESS_KEYS, ...FAILURE_KEYS, 'deleted']);
const LABELS: Record<string, string> = {
pending: '生成中',
optimizing: '生成中',
prompt_optimized: '生成',
pending: '待处理',
optimizing: '提词处理中',
prompt_optimized: '生成',
generating: '生成中',
queued: '生成中',
preparing: '生成中',
@@ -64,15 +60,16 @@ const LABELS: Record<string, string> = {
download_queued: '生成中',
downloading: '生成中',
retry_waiting: '生成中',
recovery_inconsistent: '生成中',
upscale_queued: '生成中',
upscale_processing: '生成中',
upscale_polling: '生成中',
upscale_downloading: '生成中',
upscale_finalizing: '生成中',
upscale_retry_waiting: '生成中',
completed: '已完成',
done: '已完成',
timeout: '生成超时',
completed: '成',
done: '成',
timeout: '失败',
download_failed: '失败',
upscale_failed: '失败',
failed: '失败',
@@ -80,33 +77,15 @@ const LABELS: Record<string, string> = {
};
const COLOR_MAP: Record<string, GenerationUiColor> = {
pending: 'default',
optimizing: 'processing',
prompt_optimized: 'processing',
generating: 'warning',
queued: 'processing',
preparing: 'processing',
creating_provider_task: 'processing',
provider_result_staged: 'processing',
waiting_remote: 'processing',
polling: 'processing',
result_ready: 'processing',
download_queued: 'processing',
downloading: 'processing',
retry_waiting: 'orange',
upscale_queued: 'purple',
upscale_processing: 'purple',
upscale_polling: 'purple',
upscale_downloading: 'purple',
upscale_finalizing: 'purple',
upscale_retry_waiting: 'orange',
completed: 'success',
done: 'success',
failed: 'error',
timeout: 'error',
download_failed: 'error',
upscale_failed: 'error',
deleted: 'default',
pending: 'default', optimizing: 'processing', prompt_optimized: 'blue', generating: 'processing',
queued: 'processing', preparing: 'processing', creating_provider_task: 'processing',
provider_result_staged: 'processing', waiting_remote: 'processing', polling: 'processing',
result_ready: 'processing', download_queued: 'processing', downloading: 'processing',
retry_waiting: 'orange', recovery_inconsistent: 'orange', upscale_queued: 'purple',
upscale_processing: 'purple', upscale_polling: 'purple', upscale_downloading: 'purple',
upscale_finalizing: 'purple', upscale_retry_waiting: 'orange', completed: 'success',
done: 'success', failed: 'error', timeout: 'error', download_failed: 'error',
upscale_failed: 'error', deleted: 'default',
};
const normalize = (value?: string | null): string => String(value || '').trim().toLowerCase();
@@ -121,45 +100,28 @@ export const getGenerationStatusColor = (key?: string | null): GenerationUiColor
return COLOR_MAP[normalized] || 'default';
};
const firstMatching = (values: string[], keys: Set<string>): string => (
values.find((value) => keys.has(value)) || ''
);
export const resolveGenerationUiState = (value: GenerationStatusLike): GenerationUiState => {
const status = normalize(value.status);
const displayStatus = normalize(value.displayStatus);
const pipelineStage = normalize(value.pipelineStage);
const values = [displayStatus, status, pipelineStage].filter(Boolean);
const clientStatus = normalize(value.clientStatus);
const mappedClientStatus = CLIENT_STATUS_MAP[clientStatus] || '';
const failureKey = FAILURE_KEYS.has(pipelineStage)
? pipelineStage
: firstMatching([displayStatus, status], FAILURE_KEYS);
const deletedKey = firstMatching(values, new Set(['deleted']));
const successKey = firstMatching(values, SUCCESS_KEYS);
let effectiveKey = '';
if (failureKey) {
effectiveKey = failureKey;
} else if (deletedKey) {
effectiveKey = deletedKey;
} else if (ACTIVE_PIPELINE_STAGES.has(pipelineStage)) {
let effectiveKey = mappedClientStatus || displayStatus || status || pipelineStage || 'pending';
if (FAILURE_KEYS.has(pipelineStage) || FAILURE_KEYS.has(status) || FAILURE_KEYS.has(displayStatus)) {
effectiveKey = FAILURE_KEYS.has(pipelineStage) ? pipelineStage : (FAILURE_KEYS.has(status) ? status : displayStatus);
} else if (SUCCESS_KEYS.has(status) || SUCCESS_KEYS.has(displayStatus)) {
effectiveKey = SUCCESS_KEYS.has(status) ? status : displayStatus;
} else if (!mappedClientStatus && pipelineStage && ACTIVE_PIPELINE_STAGES.has(pipelineStage)) {
effectiveKey = pipelineStage;
} else if (successKey) {
effectiveKey = successKey;
} else if (pipelineStage) {
effectiveKey = pipelineStage;
} else {
effectiveKey = displayStatus || status || 'pending';
}
const isFailure = FAILURE_KEYS.has(effectiveKey);
const isSuccess = SUCCESS_KEYS.has(effectiveKey);
const isActive = !isFailure && !isSuccess && effectiveKey !== 'deleted' && (
ACTIVE_PIPELINE_STAGES.has(pipelineStage)
|| ACTIVE_STATUS_KEYS.has(displayStatus)
|| ACTIVE_STATUS_KEYS.has(status)
|| ACTIVE_PIPELINE_STAGES.has(effectiveKey)
);
const shouldPoll = typeof value.shouldPoll === 'boolean'
? value.shouldPoll
: (!isFailure && !isSuccess && (ACTIVE_STATUS_KEYS.has(status) || ACTIVE_PIPELINE_STAGES.has(pipelineStage)));
const isActive = shouldPoll;
return {
status,
@@ -172,6 +134,9 @@ export const resolveGenerationUiState = (value: GenerationStatusLike): Generatio
isSuccess,
isFailure,
isTerminal: TERMINAL_KEYS.has(effectiveKey),
canGenerate: typeof value.canGenerate === 'boolean' ? value.canGenerate : status === 'prompt_optimized',
canRetry: typeof value.canRetry === 'boolean' ? value.canRetry : isFailure,
shouldPoll,
};
};