diff --git a/.gitignore b/.gitignore index 4887c932..4c3b48dc 100644 --- a/.gitignore +++ b/.gitignore @@ -22,3 +22,4 @@ video-gen-api/dist/ # 忽略特定类型文件但保留目录 # *.pyc # !dir/*.pycnode_modules/ +*.tmp.* \ No newline at end of file diff --git a/DEPLOYMENT.md b/DEPLOYMENT.md index c4017611..f33f9823 100644 --- a/DEPLOYMENT.md +++ b/DEPLOYMENT.md @@ -21,6 +21,8 @@ video_item/ | PostgreSQL | >= 14 | 推荐 16 | | Redis | >= 6 | 可选,推荐用于限流/验证码/Celery | | FFmpeg | 任意 | 可选,用于视频封面截帧 | +| alipay-sdk-python | >=3.7.1160 | 可选,用于支付 | +| ca-certificates | 任意 | **必须**,HTTPS 请求需要(新服务器/容器常缺) | --- @@ -29,6 +31,12 @@ video_item/ ### 1. 安装依赖 ```bash +# ⚠️ 新服务器/容器必须先装 CA 证书,否则 HTTPS 请求(支付宝/火山等)全部失败 +# CentOS/RHEL +sudo yum install -y ca-certificates +# Ubuntu/Debian +sudo apt-get install -y ca-certificates + cd video-gen-api # 创建虚拟环境 @@ -46,6 +54,12 @@ pip install -e ".[pg,redis]" # 如需 Celery 异步任务(ChatAPI 生成流水线) pip install -e ".[pg,redis,celery]" + +#安装阿里支付sdk +pip install -e ".[pg,redis,celery,alipay]" + +#安装火山sdk +pip install -e ".[pg,redis,celery,alipay,volc]" ``` ### 2. 配置环境变量 diff --git a/video-gen-admin/package-lock.json b/video-gen-admin/package-lock.json index d8a407b0..4508c7be 100644 --- a/video-gen-admin/package-lock.json +++ b/video-gen-admin/package-lock.json @@ -10,6 +10,7 @@ "dependencies": { "@ant-design/icons": "^6.2.2", "antd": "^6.3.7", + "dayjs": "^1.11.21", "react": "^19.2.5", "react-dom": "^19.2.5", "react-router-dom": "^7.15.0", @@ -1333,9 +1334,9 @@ "license": "MIT" }, "node_modules/dayjs": { - "version": "1.11.20", - "resolved": "https://registry.npmjs.org/dayjs/-/dayjs-1.11.20.tgz", - "integrity": "sha512-YbwwqR/uYpeoP4pu043q+LTDLFBLApUP6VxRihdfNTqu4ubqMlGDLd6ErXhEgsyvY0K6nCs7nggYumAN+9uEuQ==", + "version": "1.11.21", + "resolved": "https://registry.npmjs.org/dayjs/-/dayjs-1.11.21.tgz", + "integrity": "sha512-98IT+HOahAisibz/yjKbzuOBwYcjJ7BCLPzARyHiyEBmRz4fatF+KPJszEHXsGYjUG234aH/cOjW1wwTbKUZlA==", "license": "MIT" }, "node_modules/detect-libc": { diff --git a/video-gen-admin/package.json b/video-gen-admin/package.json index 0693e9f2..31da407d 100644 --- a/video-gen-admin/package.json +++ b/video-gen-admin/package.json @@ -11,6 +11,7 @@ "dependencies": { "@ant-design/icons": "^6.2.2", "antd": "^6.3.7", + "dayjs": "^1.11.21", "react": "^19.2.5", "react-dom": "^19.2.5", "react-router-dom": "^7.15.0", diff --git a/video-gen-admin/src/App.tsx b/video-gen-admin/src/App.tsx index 3f35093b..492e982c 100644 --- a/video-gen-admin/src/App.tsx +++ b/video-gen-admin/src/App.tsx @@ -11,6 +11,7 @@ import AdminSettings from './pages/AdminSettings'; import AdminNotificationManager from './pages/AdminNotificationManager'; import AdminCreditRecords from './pages/AdminCreditRecords'; import AdminPaymentConfig from './pages/AdminPaymentConfig'; +import AdminPaymentStats from './pages/AdminPaymentStats'; import AdminIndustries from './pages/AdminIndustries'; import AdminVideoEngines from './pages/AdminVideoEngines'; import AdminImageEngines from './pages/AdminImageEngines'; @@ -76,6 +77,7 @@ const App = () => { } /> } /> } /> + } /> } /> } /> } /> diff --git a/video-gen-admin/src/api/index.ts b/video-gen-admin/src/api/index.ts index e325c04d..72951626 100644 --- a/video-gen-admin/src/api/index.ts +++ b/video-gen-admin/src/api/index.ts @@ -207,6 +207,43 @@ export async function updatePaymentConfig(id: string, value: string): Promise): Promise { + await api.put('/admin/payment-configs/batch', configs); +} + +export async function getPaymentStats(params?: { + paymentMethod?: string; + status?: string; + startDate?: string; + endDate?: string; +}): Promise<{ + byStatus: Record; + today: { paidCount: number; paidAmount: number }; + month: { paidCount: number; paidAmount: number }; + recent: any[]; +}> { + const searchParams = new URLSearchParams(); + if (params?.paymentMethod) searchParams.set('payment_method', params.paymentMethod); + if (params?.status) searchParams.set('status', params.status); + if (params?.startDate) searchParams.set('start_date', params.startDate); + if (params?.endDate) searchParams.set('end_date', params.endDate); + const queryString = searchParams.toString(); + const url = queryString ? `/admin/payment-stats?${queryString}` : '/admin/payment-stats'; + return api.get(url); +} + +export async function getAdminPaymentOrders(params?: { method?: string; status?: string }): Promise<{ items: any[] }> { + const qs = new URLSearchParams(); + if (params?.method) qs.set('method', params.method); + if (params?.status) qs.set('status', params.status); + const suffix = qs.toString() ? `?${qs.toString()}` : ''; + return api.get(`/admin/payment-orders${suffix}`); +} + +export async function refundPaymentOrder(orderNo: string): Promise { + await api.post(`/admin/payment-orders/${orderNo}/refund`); +} + export async function getAdminNotifications(): Promise<{ total: number; items: any[] }> { return api.get('/admin/notifications'); } diff --git a/video-gen-admin/src/pages/AdminPaymentConfig.tsx b/video-gen-admin/src/pages/AdminPaymentConfig.tsx index 0a435789..c125f1c1 100644 --- a/video-gen-admin/src/pages/AdminPaymentConfig.tsx +++ b/video-gen-admin/src/pages/AdminPaymentConfig.tsx @@ -1,32 +1,25 @@ import React, { useEffect, useState } from 'react'; import { - Button, Card, Form, Input, message, Switch, Typography, + Button, Card, Form, Input, message, Switch, Typography, InputNumber, } from 'antd'; import { - SaveOutlined, WechatOutlined, AlipayCircleOutlined, DollarOutlined, + SaveOutlined, WechatOutlined, AlipayCircleOutlined, DollarOutlined, ClockCircleOutlined, } from '@ant-design/icons'; -import { getPaymentConfigs, updatePaymentConfig } from '../api'; - -interface PaymentConfig { - id: string; - key: string; - value: string; - description?: string; -} +import { getPaymentConfigs, batchUpdatePaymentConfigs } from '../api'; const AdminPaymentConfig: React.FC = () => { const [saving, setSaving] = useState(false); - const [configs, setConfigs] = useState([]); const [wechatEnabled, setWechatEnabled] = useState(false); const [alipayEnabled, setAlipayEnabled] = useState(false); + const [mockMode, setMockMode] = useState(false); + const [orderTimeout, setOrderTimeout] = useState(180); const [form] = Form.useForm(); const load = async () => { try { const data = await getPaymentConfigs(); - setConfigs(data); const map: Record = {}; - data.forEach((c: PaymentConfig) => { map[c.key] = c.value; }); + data.forEach((c: any) => { map[c.key] = c.value; }); form.setFieldsValue({ wechat_mch_id: map['payment_wechat_mch_id'] || '', wechat_api_key: map['payment_wechat_api_key'] || '', @@ -36,9 +29,13 @@ const AdminPaymentConfig: React.FC = () => { alipay_private_key: map['payment_alipay_private_key'] || '', alipay_public_key: map['payment_alipay_public_key'] || '', alipay_notify_url: map['payment_alipay_notify_url'] || '', + alipay_gateway: map['payment_alipay_gateway'] || '', + order_timeout: map['payment_order_timeout'] || '180', }); setWechatEnabled(map['payment_wechat_enabled'] === 'true'); setAlipayEnabled(map['payment_alipay_enabled'] === 'true'); + setMockMode(map['payment_mock'] === 'true'); + setOrderTimeout(parseInt(map['payment_order_timeout'] || '180', 10)); } catch { message.error('加载支付配置失败'); } @@ -50,24 +47,21 @@ const AdminPaymentConfig: React.FC = () => { try { const values = await form.validateFields(); setSaving(true); - const updates: [string, string][] = [ - ['payment_wechat_enabled', String(wechatEnabled)], - ['payment_wechat_mch_id', values.wechat_mch_id || ''], - ['payment_wechat_api_key', values.wechat_api_key || ''], - ['payment_wechat_cert_path', values.wechat_cert_path || ''], - ['payment_wechat_notify_url', values.wechat_notify_url || ''], - ['payment_alipay_enabled', String(alipayEnabled)], - ['payment_alipay_app_id', values.alipay_app_id || ''], - ['payment_alipay_private_key', values.alipay_private_key || ''], - ['payment_alipay_public_key', values.alipay_public_key || ''], - ['payment_alipay_notify_url', values.alipay_notify_url || ''], - ]; - for (const [key, value] of updates) { - const cfg = configs.find(c => c.key === key); - if (cfg) { - await updatePaymentConfig(cfg.id, value); - } - } + await batchUpdatePaymentConfigs({ + payment_mock: String(mockMode), + payment_wechat_enabled: String(wechatEnabled), + payment_wechat_mch_id: values.wechat_mch_id || '', + payment_wechat_api_key: values.wechat_api_key || '', + payment_wechat_cert_path: values.wechat_cert_path || '', + payment_wechat_notify_url: values.wechat_notify_url || '', + payment_alipay_enabled: String(alipayEnabled), + payment_alipay_app_id: values.alipay_app_id || '', + payment_alipay_private_key: values.alipay_private_key || '', + payment_alipay_public_key: values.alipay_public_key || '', + payment_alipay_notify_url: values.alipay_notify_url || '', + payment_alipay_gateway: values.alipay_gateway || '', + payment_order_timeout: String(values.order_timeout || 180), + }); message.success('支付配置已保存'); load(); } catch { @@ -79,6 +73,61 @@ const AdminPaymentConfig: React.FC = () => { return (
+ {/* 通用设置 */} + +
+
+
+
+ 通用设置 + 订单超时和测试模式配置 +
+
+
+ +
+ 订单超时时间} + extra="订单创建后超过此时间未支付将自动取消(秒)" + > + + +
+
+ + {/* Mock Mode Toggle */} + +
+
+
+
+ 模拟支付模式 + + {mockMode ? '⚠️ 开启后所有充值会直接成功(仅用于测试)' : '关闭 - 使用真实支付渠道'} + +
+
+ +
+
+ {/* WeChat Pay */}
@@ -141,7 +190,10 @@ const AdminPaymentConfig: React.FC = () => { - + + + + diff --git a/video-gen-admin/src/pages/AdminPaymentStats.tsx b/video-gen-admin/src/pages/AdminPaymentStats.tsx new file mode 100644 index 00000000..582e646e --- /dev/null +++ b/video-gen-admin/src/pages/AdminPaymentStats.tsx @@ -0,0 +1,285 @@ +import React, { useEffect, useState } from 'react'; +import { + Card, Col, Row, Space, Table, Tag, Typography, Statistic, message, Select, DatePicker, Button, ConfigProvider, Popconfirm +} from 'antd'; +import zhCN from 'antd/locale/zh_CN'; +import { + DollarOutlined, CheckCircleOutlined, ClockCircleOutlined, CloseCircleOutlined, ReloadOutlined, UndoOutlined +} from '@ant-design/icons'; +import { getPaymentStats, refundPaymentOrder } from '../api'; +import { formatDate } from '../utils/formatDate'; +import dayjs from 'dayjs'; + +const { Option } = Select; +const { RangePicker } = DatePicker; + +const AdminPaymentStats: React.FC = () => { + const [loading, setLoading] = useState(false); + const [stats, setStats] = useState(null); + const [filters, setFilters] = useState<{ + paymentMethod?: string; + status?: string; + startDate: string; + endDate: string; + }>({ + startDate: dayjs().format('YYYY-MM-DD'), + endDate: dayjs().format('YYYY-MM-DD'), + }); + + const load = async () => { + try { + setLoading(true); + const data = await getPaymentStats(filters); + setStats(data); + } catch { + message.error('加载支付统计失败'); + } finally { + setLoading(false); + } + }; + + useEffect(() => { load(); }, [filters]); + + const handleReset = () => { + setFilters({ + startDate: dayjs().format('YYYY-MM-DD'), + endDate: dayjs().format('YYYY-MM-DD'), + }); + }; + + const handleRefund = async (orderNo: string) => { + try { + setLoading(true); + await refundPaymentOrder(orderNo); + message.success('退款成功'); + await load(); + } catch (e: any) { + message.error(e?.response?.data?.detail || '退款失败'); + } finally { + setLoading(false); + } + }; + + const handleDateChange = (dates: any) => { + if (dates && dates.length === 2) { + setFilters(prev => ({ + ...prev, + startDate: dates[0].format('YYYY-MM-DD'), + endDate: dates[1].format('YYYY-MM-DD'), + })); + } + }; + + const statusConfig: Record = { + paid: { color: 'green', label: '已支付', icon: }, + pending: { color: 'gold', label: '待支付', icon: }, + cancelled: { color: 'default', label: '已取消', icon: }, + refunded: { color: 'red', label: '已退款', icon: }, + }; + + const methodConfig: Record = { + alipay: { color: 'blue', label: '支付宝' }, + wechat: { color: 'green', label: '微信' }, + }; + + const columns = [ + { title: '订单号', dataIndex: 'orderNo', key: 'orderNo', width: 200 }, + { title: '用户', dataIndex: 'username', key: 'username', width: 120 }, + { + title: '支付方式', dataIndex: 'paymentMethod', key: 'paymentMethod', width: 100, + render: (m: string) => { + const c = methodConfig[m] || { color: 'default', label: m }; + return {c.label}; + }, + }, + { + title: '金额', dataIndex: 'amount', key: 'amount', width: 100, + render: (a: number) => ¥{a.toFixed(2)}, + }, + { title: '积分', dataIndex: 'credits', key: 'credits', width: 80 }, + { + title: '状态', dataIndex: 'status', key: 'status', width: 100, + render: (s: string) => { + const c = statusConfig[s] || { color: 'default', label: s, icon: null }; + return {c.label}; + }, + }, + { title: '支付宝交易号', dataIndex: 'tradeNo', key: 'tradeNo', width: 200, render: (v: string) => v || '-' }, + { + title: '创建时间', dataIndex: 'createdAt', key: 'createdAt', width: 160, + render: (d: string) => {d ? formatDate(d) : '-'}, + }, + { + title: '支付时间', dataIndex: 'paidAt', key: 'paidAt', width: 160, + render: (d: string) => {d ? formatDate(d) : '-'}, + }, + { + title: '操作', + key: 'action', + width: 120, + render: (_: any, record: any) => { + if (record.status === 'paid') { + return ( + handleRefund(record.orderNo)} + okText="确认" + cancelText="取消" + > + + + ); + } + return null; + }, + }, + ]; + + if (!stats) { + return
加载中…
; + } + + const paidInfo = stats.byStatus?.paid || { count: 0, amount: 0 }; + const pendingInfo = stats.byStatus?.pending || { count: 0, amount: 0 }; + const cancelledInfo = stats.byStatus?.cancelled || { count: 0, amount: 0 }; + const refundedInfo = stats.byStatus?.refunded || { count: 0, amount: 0 }; + const totalOrders = paidInfo.count + pendingInfo.count + cancelledInfo.count + refundedInfo.count; + const monthInfo = stats.month || { count: 0, amount: 0 }; + + return ( + +
+ {/* Summary cards */} + + + + } + suffix="元" + valueStyle={{ color: '#10b981', fontWeight: 700 }} + /> + + {stats.today.paidCount} 笔订单 + + + + + + } + suffix="元" + valueStyle={{ color: '#6366f1', fontWeight: 700 }} + /> + + {monthInfo.paidCount} 笔订单 + + + + + + {/* Status breakdown */} + 订单状态分布}> + + {['paid', 'pending', 'cancelled', 'refunded'].map(s => { + const info = stats.byStatus?.[s] || { count: 0, amount: 0 }; + const c = statusConfig[s]; + const pct = totalOrders > 0 ? ((info.count / totalOrders) * 100).toFixed(1) : '0.0'; + return ( + +
+ + {c.label} + {pct}% + +
+ {info.count} +
+
+ ¥{info.amount.toFixed(2)} +
+
+ + ); + })} +
+
+ + {/* Recent orders table */} + 订单列表}> + {/* Filters */} + + + 支付方式: + + + + 状态: + + + + 日期范围: + + + + + + + + + + + + ); +}; + +export default AdminPaymentStats; diff --git a/video-gen-admin/src/types/index.ts b/video-gen-admin/src/types/index.ts index 5fe99196..02d0eeae 100644 --- a/video-gen-admin/src/types/index.ts +++ b/video-gen-admin/src/types/index.ts @@ -117,6 +117,29 @@ export interface AdminStats { creditsConsumedToday: number; } +export interface PaymentStats { + byStatus: Record; + today: { paidCount: number; paidAmount: number }; + month: { paidCount: number; paidAmount: number }; + recent: PaymentOrder[]; +} + +export interface PaymentOrder { + id: string; + orderNo: string; + userId: string; + username?: string; + amount: number; + credits: number; + paymentMethod: string; + status: string; + tradeNo?: string; + paidAt?: string; + createdAt: string; + refundedAt?: string; + refundAmount?: number; +} + export interface ModelConfig { id: string; name: string; diff --git a/video-gen-api/alembic/versions/150fc6da855f_add_module_generation_project_and_hot_.py b/video-gen-api/alembic/versions/150fc6da855f_add_module_generation_project_and_hot_.py new file mode 100644 index 00000000..4c76a8ab --- /dev/null +++ b/video-gen-api/alembic/versions/150fc6da855f_add_module_generation_project_and_hot_.py @@ -0,0 +1,120 @@ +"""add module generation project and hot opening replicate + +Revision ID: 150fc6da855f +Revises: 476b259992de +Create Date: 2026-06-10 11:56:22.146017 +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa +from sqlalchemy.dialects import postgresql + +# revision identifiers, used by Alembic. +revision: str = '150fc6da855f' +down_revision: Union[str, None] = '476b259992de' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.create_table('module_generation_projects', + sa.Column('id', sa.String(length=32), nullable=False), + sa.Column('user_id', sa.String(length=32), nullable=False), + sa.Column('module', sa.String(length=64), nullable=False), + sa.Column('title', sa.String(length=160), nullable=True), + sa.Column('status', sa.String(length=32), nullable=False), + sa.Column('current_step_code', sa.String(length=64), nullable=True), + sa.Column('final_image_url', sa.String(length=512), nullable=True), + sa.Column('final_video_url', sa.String(length=512), nullable=True), + sa.Column('final_video_cover_url', sa.String(length=512), nullable=True), + sa.Column('error_message', sa.Text(), nullable=True), + sa.Column('idempotency_key', sa.String(length=64), nullable=True), + sa.Column('completed_at', sa.DateTime(timezone=True), nullable=True), + sa.Column('created_at', sa.DateTime(timezone=True), server_default=sa.text('now()'), nullable=False), + sa.Column('updated_at', sa.DateTime(timezone=True), server_default=sa.text('now()'), nullable=False), + sa.Column('deleted_at', sa.DateTime(timezone=True), nullable=True), + sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'), + sa.PrimaryKeyConstraint('id') + ) + op.create_index('idx_module_generation_projects_status', 'module_generation_projects', ['module', 'status'], unique=False) + op.create_index('idx_module_generation_projects_user_module', 'module_generation_projects', ['user_id', 'module'], unique=False) + op.create_index(op.f('ix_module_generation_projects_current_step_code'), 'module_generation_projects', ['current_step_code'], unique=False) + op.create_index(op.f('ix_module_generation_projects_deleted_at'), 'module_generation_projects', ['deleted_at'], unique=False) + op.create_index(op.f('ix_module_generation_projects_idempotency_key'), 'module_generation_projects', ['idempotency_key'], unique=False) + op.create_index(op.f('ix_module_generation_projects_module'), 'module_generation_projects', ['module'], unique=False) + op.create_index(op.f('ix_module_generation_projects_status'), 'module_generation_projects', ['status'], unique=False) + op.create_index(op.f('ix_module_generation_projects_user_id'), 'module_generation_projects', ['user_id'], unique=False) + op.create_index('uq_module_generation_projects_user_module_idempotency', 'module_generation_projects', ['user_id', 'module', 'idempotency_key'], unique=True, postgresql_where=sa.text('deleted_at IS NULL AND idempotency_key IS NOT NULL')) + op.create_table('module_generation_steps', + sa.Column('id', sa.String(length=32), nullable=False), + sa.Column('project_id', sa.String(length=32), nullable=False), + sa.Column('user_id', sa.String(length=32), nullable=False), + sa.Column('module', sa.String(length=64), nullable=False), + sa.Column('step_index', sa.Integer(), nullable=False), + sa.Column('step_code', sa.String(length=64), nullable=False), + sa.Column('status', sa.String(length=32), nullable=False), + sa.Column('version', sa.Integer(), nullable=False), + sa.Column('is_current', sa.Boolean(), nullable=False), + sa.Column('parent_step_id', sa.String(length=32), nullable=True), + sa.Column('source_step_id', sa.String(length=32), nullable=True), + sa.Column('chat_task_id', sa.String(length=32), nullable=True), + sa.Column('input_json', sa.JSON().with_variant(postgresql.JSONB(astext_type=sa.Text()), 'postgresql'), nullable=True), + sa.Column('output_json', sa.JSON().with_variant(postgresql.JSONB(astext_type=sa.Text()), 'postgresql'), nullable=True), + sa.Column('error_message', sa.Text(), nullable=True), + sa.Column('started_at', sa.DateTime(timezone=True), nullable=True), + sa.Column('completed_at', sa.DateTime(timezone=True), nullable=True), + sa.Column('created_at', sa.DateTime(timezone=True), server_default=sa.text('now()'), nullable=False), + sa.Column('updated_at', sa.DateTime(timezone=True), server_default=sa.text('now()'), nullable=False), + sa.Column('deleted_at', sa.DateTime(timezone=True), nullable=True), + sa.ForeignKeyConstraint(['chat_task_id'], ['chat_generation_tasks.id'], ondelete='SET NULL'), + sa.ForeignKeyConstraint(['project_id'], ['module_generation_projects.id'], ondelete='CASCADE'), + sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'), + sa.PrimaryKeyConstraint('id') + ) + op.create_index('idx_module_generation_steps_chat_task', 'module_generation_steps', ['chat_task_id'], unique=False) + op.create_index('idx_module_generation_steps_project_code', 'module_generation_steps', ['project_id', 'step_code', 'is_current'], unique=False) + op.create_index('idx_module_generation_steps_project_current', 'module_generation_steps', ['project_id', 'is_current', 'deleted_at'], unique=False) + op.create_index(op.f('ix_module_generation_steps_chat_task_id'), 'module_generation_steps', ['chat_task_id'], unique=False) + op.create_index(op.f('ix_module_generation_steps_deleted_at'), 'module_generation_steps', ['deleted_at'], unique=False) + op.create_index(op.f('ix_module_generation_steps_is_current'), 'module_generation_steps', ['is_current'], unique=False) + op.create_index(op.f('ix_module_generation_steps_module'), 'module_generation_steps', ['module'], unique=False) + op.create_index(op.f('ix_module_generation_steps_parent_step_id'), 'module_generation_steps', ['parent_step_id'], unique=False) + op.create_index(op.f('ix_module_generation_steps_project_id'), 'module_generation_steps', ['project_id'], unique=False) + op.create_index(op.f('ix_module_generation_steps_source_step_id'), 'module_generation_steps', ['source_step_id'], unique=False) + op.create_index(op.f('ix_module_generation_steps_status'), 'module_generation_steps', ['status'], unique=False) + op.create_index(op.f('ix_module_generation_steps_step_code'), 'module_generation_steps', ['step_code'], unique=False) + op.create_index(op.f('ix_module_generation_steps_step_index'), 'module_generation_steps', ['step_index'], unique=False) + op.create_index(op.f('ix_module_generation_steps_user_id'), 'module_generation_steps', ['user_id'], unique=False) + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.drop_index(op.f('ix_module_generation_steps_user_id'), table_name='module_generation_steps') + op.drop_index(op.f('ix_module_generation_steps_step_index'), table_name='module_generation_steps') + op.drop_index(op.f('ix_module_generation_steps_step_code'), table_name='module_generation_steps') + op.drop_index(op.f('ix_module_generation_steps_status'), table_name='module_generation_steps') + op.drop_index(op.f('ix_module_generation_steps_source_step_id'), table_name='module_generation_steps') + op.drop_index(op.f('ix_module_generation_steps_project_id'), table_name='module_generation_steps') + op.drop_index(op.f('ix_module_generation_steps_parent_step_id'), table_name='module_generation_steps') + op.drop_index(op.f('ix_module_generation_steps_module'), table_name='module_generation_steps') + op.drop_index(op.f('ix_module_generation_steps_is_current'), table_name='module_generation_steps') + op.drop_index(op.f('ix_module_generation_steps_deleted_at'), table_name='module_generation_steps') + op.drop_index(op.f('ix_module_generation_steps_chat_task_id'), table_name='module_generation_steps') + op.drop_index('idx_module_generation_steps_project_current', table_name='module_generation_steps') + op.drop_index('idx_module_generation_steps_project_code', table_name='module_generation_steps') + op.drop_index('idx_module_generation_steps_chat_task', table_name='module_generation_steps') + op.drop_table('module_generation_steps') + op.drop_index('uq_module_generation_projects_user_module_idempotency', table_name='module_generation_projects', postgresql_where=sa.text('deleted_at IS NULL AND idempotency_key IS NOT NULL')) + op.drop_index(op.f('ix_module_generation_projects_user_id'), table_name='module_generation_projects') + op.drop_index(op.f('ix_module_generation_projects_status'), table_name='module_generation_projects') + op.drop_index(op.f('ix_module_generation_projects_module'), table_name='module_generation_projects') + op.drop_index(op.f('ix_module_generation_projects_idempotency_key'), table_name='module_generation_projects') + op.drop_index(op.f('ix_module_generation_projects_deleted_at'), table_name='module_generation_projects') + op.drop_index(op.f('ix_module_generation_projects_current_step_code'), table_name='module_generation_projects') + op.drop_index('idx_module_generation_projects_user_module', table_name='module_generation_projects') + op.drop_index('idx_module_generation_projects_status', table_name='module_generation_projects') + op.drop_table('module_generation_projects') + # ### end Alembic commands ### diff --git a/video-gen-api/alembic/versions/6101ba8d5761_user_oauth表新增归属公司应用_新增授权链接字段.py b/video-gen-api/alembic/versions/6101ba8d5761_user_oauth表新增归属公司应用_新增授权链接字段.py new file mode 100644 index 00000000..fa9449f5 --- /dev/null +++ b/video-gen-api/alembic/versions/6101ba8d5761_user_oauth表新增归属公司应用_新增授权链接字段.py @@ -0,0 +1,31 @@ +"""user_oauth表新增归属公司应用,新增授权链接字段 + +Revision ID: 6101ba8d5761 +Revises: 9ac2212e1b8e +Create Date: 2026-06-10 16:34:38.842582 +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision: str = '6101ba8d5761' +down_revision: Union[str, None] = '9ac2212e1b8e' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.add_column('user_oauth_app', sa.Column('auth_url', sa.String(length=256), nullable=True, comment='应用授权链接')) + op.add_column('user_oauth_app', sa.Column('company', sa.String(length=256), nullable=True, comment='应用归属公司名称')) + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.drop_column('user_oauth_app', 'company') + op.drop_column('user_oauth_app', 'auth_url') + # ### end Alembic commands ### diff --git a/video-gen-api/alembic/versions/8922eafcd8b0_user_oauth表新增account_userid授权登录id_.py b/video-gen-api/alembic/versions/8922eafcd8b0_user_oauth表新增account_userid授权登录id_.py new file mode 100644 index 00000000..3b746b30 --- /dev/null +++ b/video-gen-api/alembic/versions/8922eafcd8b0_user_oauth表新增account_userid授权登录id_.py @@ -0,0 +1,29 @@ +"""user_oauth表新增account_userid授权登录id,用来判断不同账号授权 + +Revision ID: 8922eafcd8b0 +Revises: 6101ba8d5761 +Create Date: 2026-06-11 16:21:17.617777 +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision: str = '8922eafcd8b0' +down_revision: Union[str, None] = '6101ba8d5761' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.add_column('user_oauth', sa.Column('account_userid', sa.String(length=128), nullable=True, comment='授权账户登录userid,同一个用户不同的授权账户token不一样')) + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.drop_column('user_oauth', 'account_userid') + # ### end Alembic commands ### diff --git a/video-gen-api/alembic/versions/9ac2212e1b8e_user_oauth表新增可授权数量字段.py b/video-gen-api/alembic/versions/9ac2212e1b8e_user_oauth表新增可授权数量字段.py new file mode 100644 index 00000000..c4069935 --- /dev/null +++ b/video-gen-api/alembic/versions/9ac2212e1b8e_user_oauth表新增可授权数量字段.py @@ -0,0 +1,29 @@ +"""user_oauth表新增可授权数量字段 + +Revision ID: 9ac2212e1b8e +Revises: ed3823a24b8e +Create Date: 2026-06-10 15:43:47.389442 +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision: str = '9ac2212e1b8e' +down_revision: Union[str, None] = 'ed3823a24b8e' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.add_column('user_oauth_app', sa.Column('count', sa.BigInteger(), nullable=False, comment='应用最大可以授权多少个用户')) + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.drop_column('user_oauth_app', 'count') + # ### end Alembic commands ### diff --git a/video-gen-api/alembic/versions/a1b2c3d4e5f6_add_refund_fields_to_payment_orders.py b/video-gen-api/alembic/versions/a1b2c3d4e5f6_add_refund_fields_to_payment_orders.py new file mode 100644 index 00000000..69f81b00 --- /dev/null +++ b/video-gen-api/alembic/versions/a1b2c3d4e5f6_add_refund_fields_to_payment_orders.py @@ -0,0 +1,29 @@ +"""add_refund_fields_to_payment_orders + +Revision ID: a1b2c3d4e5f6 +Revises: ed59aefc83da +Create Date: 2026-06-10 12:00:00.000000 +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision: str = 'a1b2c3d4e5f6' +down_revision: Union[str, None] = 'ed59aefc83da' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + op.add_column('payment_orders', sa.Column('refund_trade_no', sa.String(length=128), nullable=True)) + op.add_column('payment_orders', sa.Column('refunded_at', sa.DateTime(timezone=True), nullable=True)) + op.add_column('payment_orders', sa.Column('refund_amount', sa.Float(), nullable=True)) + + +def downgrade() -> None: + op.drop_column('payment_orders', 'refund_amount') + op.drop_column('payment_orders', 'refunded_at') + op.drop_column('payment_orders', 'refund_trade_no') diff --git a/video-gen-api/alembic/versions/ed3823a24b8e_merge_changes_from_remote.py b/video-gen-api/alembic/versions/ed3823a24b8e_merge_changes_from_remote.py new file mode 100644 index 00000000..6d314a5f --- /dev/null +++ b/video-gen-api/alembic/versions/ed3823a24b8e_merge_changes_from_remote.py @@ -0,0 +1,25 @@ +"""merge changes from remote + +Revision ID: ed3823a24b8e +Revises: 150fc6da855f, a1b2c3d4e5f6 +Create Date: 2026-06-10 15:12:49.885803 +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision: str = 'ed3823a24b8e' +down_revision: Union[str, None] = ('150fc6da855f', 'a1b2c3d4e5f6') +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + pass + + +def downgrade() -> None: + pass diff --git a/video-gen-api/alembic/versions/ed59aefc83da_描述改动内容.py b/video-gen-api/alembic/versions/ed59aefc83da_描述改动内容.py new file mode 100644 index 00000000..fbb7a58f --- /dev/null +++ b/video-gen-api/alembic/versions/ed59aefc83da_描述改动内容.py @@ -0,0 +1,29 @@ +"""描述改动内容 + +Revision ID: ed59aefc83da +Revises: 476b259992de +Create Date: 2026-06-10 11:28:49.178706 +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision: str = 'ed59aefc83da' +down_revision: Union[str, None] = '476b259992de' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + pass + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + pass + # ### end Alembic commands ### diff --git a/video-gen-api/app/api/v1/__init__.py b/video-gen-api/app/api/v1/__init__.py index 3aa88f8b..e80ee5b1 100644 --- a/video-gen-api/app/api/v1/__init__.py +++ b/video-gen-api/app/api/v1/__init__.py @@ -15,7 +15,10 @@ from app.api.v1.recharge_packages import router as recharge_packages_router from app.api.v1.video_engines import router as video_engines_router from app.api.v1.image_engines import router as image_engines_router from app.api.v1.generation_ai import router as generation_ai_router +from app.api.v1.hot_opening_replicate import router as hot_opening_replicate_router from app.api.v1.test import router as test_router +from app.api.v1.user_oauth import router as user_oauth_router +from app.api.v1.user_oauth_app import router as user_oauth_app_router api_router = APIRouter() api_router.include_router(auth_router) @@ -33,4 +36,7 @@ api_router.include_router(recharge_packages_router) api_router.include_router(video_engines_router) api_router.include_router(image_engines_router) api_router.include_router(generation_ai_router) +api_router.include_router(hot_opening_replicate_router) api_router.include_router(test_router) +api_router.include_router(user_oauth_router) +api_router.include_router(user_oauth_app_router) diff --git a/video-gen-api/app/api/v1/admin.py b/video-gen-api/app/api/v1/admin.py index bbf732b6..02c7b7b3 100644 --- a/video-gen-api/app/api/v1/admin.py +++ b/video-gen-api/app/api/v1/admin.py @@ -43,6 +43,7 @@ from app.services.notification import create_notification from app.services.auth import hash_password, verify_password from app.services.operation_log import log_operation from app.services.resource_signed_url_service import build_resource_signed_url +from app.services.payment import sync_pending_orders, process_refund from app.services.generation_billing_service import ( OWNER_GENERATION_RECORD, @@ -412,6 +413,227 @@ async def list_payment_configs( ] +@router.put("/payment-configs/batch") +async def batch_update_payment_configs( + req: dict[str, str], + admin: User = Depends(get_admin_user), + db: AsyncSession = Depends(get_db), +): + """Batch upsert payment configs. Creates missing keys, updates existing ones.""" + from app.utils.id_gen import generate_id + + for key, value in req.items(): + if not key.startswith("payment_"): + continue + result = await db.execute( + select(SystemConfig).where(SystemConfig.key == key).limit(1) + ) + config = result.scalar_one_or_none() + if config: + config.value = value + else: + db.add(SystemConfig( + id=generate_id(), + key=key, + value=value, + )) + await db.flush() + return {"ok": True} + + +@router.get("/payment-stats") +async def get_payment_stats( + payment_method: str | None = Query(None), + status: str | None = Query(None), + start_date: str | None = Query(None), + end_date: str | None = Query(None), + admin: User = Depends(get_admin_user), + db: AsyncSession = Depends(get_db), +): + """Return payment statistics for admin dashboard with filters.""" + from sqlalchemy import func + + # Ensure by_status has all expected statuses with defaults + by_status = { + "pending": {"count": 0, "amount": 0.0}, + "paid": {"count": 0, "amount": 0.0}, + "cancelled": {"count": 0, "amount": 0.0}, + "refunded": {"count": 0, "amount": 0.0}, + } + + # Parse dates and build base query filters + now_cst = datetime.now(CST) + today_start = now_cst.replace(hour=0, minute=0, second=0, microsecond=0) + today_end = today_start + timedelta(days=1) + + # Default to today if no date range provided + query_start = today_start + query_end = today_end + + if start_date: + query_start = datetime.fromisoformat(start_date).replace(tzinfo=CST) + if end_date: + query_end = (datetime.fromisoformat(end_date) + timedelta(days=1)).replace(tzinfo=CST) + + # Build filter list for status breakdown + breakdown_filters = [] + if payment_method: + breakdown_filters.append(PaymentOrder.payment_method == payment_method) + if status: + breakdown_filters.append(PaymentOrder.status == status) + # Always apply date range to breakdown + breakdown_filters.append(PaymentOrder.created_at >= query_start) + breakdown_filters.append(PaymentOrder.created_at < query_end) + + # Status breakdown + status_result = await db.execute( + select( + PaymentOrder.status, + func.count().label("count"), + func.coalesce(func.sum(PaymentOrder.amount), 0).label("amount"), + ) + .where(*breakdown_filters) + .group_by(PaymentOrder.status) + ) + for row in status_result.all(): + if row.status in by_status: + by_status[row.status] = { + "count": row.count, + "amount": round(float(row.amount), 2) + } + else: + # Map any unexpected status to cancelled + by_status["cancelled"]["count"] += row.count + by_status["cancelled"]["amount"] += round(float(row.amount), 2) + + # Today's stats (CST time zone) - independent of filter + today_result = await db.execute( + select( + func.count().label("paid_count"), + func.coalesce(func.sum(PaymentOrder.amount), 0).label("paid_amount"), + ).where( + PaymentOrder.status == "paid", + PaymentOrder.paid_at >= today_start, + PaymentOrder.paid_at < today_end, + ) + ) + today_row = today_result.one() + + # Monthly cumulative stats + month_start = now_cst.replace(day=1, hour=0, minute=0, second=0, microsecond=0) + month_end = (month_start + timedelta(days=32)).replace(day=1, hour=0, minute=0, second=0, microsecond=0) + + month_result = await db.execute( + select( + func.count().label("paid_count"), + func.coalesce(func.sum(PaymentOrder.amount), 0).label("paid_amount"), + ).where( + PaymentOrder.status == "paid", + PaymentOrder.paid_at >= month_start, + PaymentOrder.paid_at < month_end, + ) + ) + month_row = month_result.one() + + # Recent orders with filters + recent_filters = [] + if payment_method: + recent_filters.append(PaymentOrder.payment_method == payment_method) + if status: + recent_filters.append(PaymentOrder.status == status) + recent_filters.append(PaymentOrder.created_at >= query_start) + recent_filters.append(PaymentOrder.created_at < query_end) + + recent_result = await db.execute( + select(PaymentOrder, User) + .join(User, PaymentOrder.user_id == User.id) + .where(*recent_filters) + .order_by(PaymentOrder.created_at.desc()) + .limit(50) + ) + recent_data = recent_result.all() + + return { + "by_status": by_status, + "today": { + "paid_count": today_row.paid_count, + "paid_amount": round(float(today_row.paid_amount), 2), + }, + "month": { + "paid_count": month_row.paid_count, + "paid_amount": round(float(month_row.paid_amount), 2), + }, + "recent": [ + { + "id": o.id, + "order_no": o.order_no, + "user_id": o.user_id, + "username": u.username, + "amount": round(o.amount, 2), + "credits": round(o.credits, 2), + "payment_method": o.payment_method, + "status": o.status if o.status in ("pending", "paid", "cancelled", "refunded") else "cancelled", + "trade_no": o.trade_no, + "paid_at": _iso(o.paid_at), + "created_at": _iso(o.created_at), + } + for o, u in recent_data + ], + } + + +@router.get("/payment-orders") +async def get_admin_payment_orders( + method: str | None = None, + status: str | None = None, + page: int = 1, + page_size: int = 20, + admin: User = Depends(get_admin_user), + db: AsyncSession = Depends(get_db), +): + """Return paginated payment orders for admin.""" + query = select(PaymentOrder) + if method: + query = query.where(PaymentOrder.payment_method == method) + if status: + query = query.where(PaymentOrder.status == status) + + # Count total + count_result = await db.execute( + select(func.count()).select_from(query.subquery()) + ) + total = count_result.scalar() or 0 + + # Paginated results + result = await db.execute( + query.order_by(PaymentOrder.created_at.desc()) + .offset((page - 1) * page_size) + .limit(page_size) + ) + orders = result.scalars().all() + + return { + "total": total, + "page": page, + "page_size": page_size, + "items": [ + { + "id": o.id, + "order_no": o.order_no, + "user_id": o.user_id, + "amount": round(o.amount, 2), + "credits": round(o.credits, 2), + "payment_method": o.payment_method, + "status": o.status if o.status in ("pending", "paid", "cancelled", "refunded") else "cancelled", + "trade_no": o.trade_no, + "paid_at": _iso(o.paid_at), + "created_at": _iso(o.created_at), + } + for o in orders + ], + } + + @router.put("/payment-configs/{config_id}") async def update_payment_config( config_id: str, @@ -440,6 +662,19 @@ async def update_payment_config( } +@router.post("/payment-orders/{order_no}/refund") +async def refund_payment_order( + order_no: str, + admin: User = Depends(get_admin_user), + db: AsyncSession = Depends(get_db), +): + """Refund a paid payment order.""" + result = await process_refund(db, order_no) + if not result.get("success"): + raise HTTPException(status_code=400, detail=result.get("message", "退款失败")) + return result + + # ── Industry Config ────────────────────────────────────── def _serialize_industry(ind: IndustryConfig) -> dict: @@ -1222,3 +1457,8 @@ async def admin_generate_video( await db.flush() return {"message": "ok", "record_id": record_id} + + +# ── Payment Stats ──────────────────────────────────────── + + diff --git a/video-gen-api/app/api/v1/hot_opening_replicate.py b/video-gen-api/app/api/v1/hot_opening_replicate.py new file mode 100644 index 00000000..e53bc8d4 --- /dev/null +++ b/video-gen-api/app/api/v1/hot_opening_replicate.py @@ -0,0 +1,527 @@ +from __future__ import annotations + +from fastapi import APIRouter, Body, Depends, HTTPException, Path, Query +from sqlalchemy.ext.asyncio import AsyncSession + +from app.dependencies import get_current_user, get_db +from app.models.user import User +from app.schemas.hot_opening_replicate import ( + HotOpeningActionOut, + HotOpeningDeleteOut, + HotOpeningGenerateImagePromptRequest, + HotOpeningGenerateImageRequest, + HotOpeningGenerateVideoPromptRequest, + HotOpeningGenerateVideoRequest, + HotOpeningImagePromptUpdateRequest, + HotOpeningMaterialUpdateRequest, + HotOpeningSpecOut, + HotOpeningTaskCreate, + HotOpeningTaskDetailOut, + HotOpeningTaskListOut, + HotOpeningVideoPromptSchemaUpdateRequest, +) +from app.services.hot_opening_replicate_service import ( + _get_project_for_user, + create_hot_opening_project, + delete_hot_opening_project, + generate_image_from_prompt, + generate_video_from_prompt, + list_hot_opening_projects, + mark_hot_opening_step_dispatch_failed, + project_to_detail_out, + submit_image_prompt_optimize, + submit_video_prompt_optimize, + update_hot_opening_image_prompt, + update_hot_opening_material_input, + update_hot_opening_video_prompt_schema, +) +from app.tasks.celery_app import celery_app + +router = APIRouter( + prefix="/hot-opening-replications", + tags=["hot-opening-replications"], +) + + +async def _reload_project_detail( + db: AsyncSession, + current_user: User, + project_id: str, +) -> HotOpeningTaskDetailOut: + """提交事务后统一重新查询详情,避免继续访问 commit 前 ORM 对象。""" + project = await _get_project_for_user( + db, + project_id=project_id, + user=current_user, + for_update=False, + populate_existing=True, + ) + return await project_to_detail_out(db, project) + + +async def _mark_dispatch_failed_and_raise( + db: AsyncSession, + *, + current_user: User, + project_id: str, + step_id: str | None, + message: str, +) -> None: + """Celery 投递失败后,数据库事务已提交,单独标记步骤失败,避免一直 processing。""" + if step_id: + try: + await mark_hot_opening_step_dispatch_failed( + db, + current_user=current_user, + project_id=project_id, + step_id=step_id, + error_message=message, + ) + await db.commit() + except Exception: + await db.rollback() + raise HTTPException(status_code=503, detail=message) + + +@router.get( + "/spec", + response_model=HotOpeningSpecOut, + summary="查询爆款开头复刻模块状态枚举和步骤 JSON 结构说明", + description="返回总任务状态、子任务状态、5个固定步骤编码以及每个步骤 input_json/output_json 的统一结构示例,方便前端和排查人员对照。", +) +async def get_spec(): + return HotOpeningSpecOut() + + +@router.post( + "/tasks", + response_model=HotOpeningTaskDetailOut, + summary="创建爆款开头复刻总任务项目", + description=( + "创建爆款开头复刻总任务项目。总任务表 id 就是项目ID,不再传 project_id。" + "接口只同步创建第1个素材输入子任务,保存素材视频链接、素材图片链接、视频素材内容项目名称、生成项目名称和50字核心内容点。" + "后端不开发上传接口,也不校验素材文件时长、大小、格式,直接使用前端已有上传接口返回的链接。" + "创建后不会自动生成第2步图片 AI 提词,需要前端手动调用 generate-image-prompt。" + ), +) +async def create_task( + req: HotOpeningTaskCreate = Body(..., description="爆款开头复刻创建参数,只包含素材链接和项目描述,不包含图片/视频引擎参数"), + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + try: + project = await create_hot_opening_project(db, current_user, req) + project_id_value = str(project.id) + await db.commit() + except HTTPException: + await db.rollback() + raise + except Exception as exc: + await db.rollback() + raise HTTPException(status_code=500, detail=f"创建爆款开头复刻项目失败: {exc}") + + return await _reload_project_detail(db, current_user, project_id_value) + + +@router.get( + "/tasks", + response_model=HotOpeningTaskListOut, + summary="查询爆款开头复刻总任务项目列表", + description="分页查询爆款开头复刻总任务项目列表。普通用户只能查看自己的项目,管理员可查看全部。", +) +async def list_tasks( + status: str | None = Query(None, description="总任务状态筛选,例如 waiting_user、processing、completed、failed;为空不过滤"), + page: int = Query(1, ge=1, description="分页页码,从1开始"), + page_size: int = Query(20, ge=1, le=100, description="每页数量,范围1-100"), + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + return await list_hot_opening_projects(db, current_user=current_user, status=status, page=page, page_size=page_size) + + +@router.get( + "/tasks/{project_id}", + response_model=HotOpeningTaskDetailOut, + summary="获取爆款开头复刻总任务项目详情", + description=( + "获取爆款开头复刻总任务详情。详情会聚合返回第1步素材信息、第2步图片提词、第3步图片引擎和参数、" + "第4步视频提词 JSON schema、第5步视频引擎和参数、最终图片、最终视频和完整子任务列表。" + ), +) +async def get_task( + project_id: str = Path(..., description="总任务项目ID,即 module_generation_projects.id"), + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + return await _reload_project_detail(db, current_user, project_id) + + +@router.put( + "/tasks/{project_id}/material", + response_model=HotOpeningActionOut, + summary="修改第1步素材输入并重建第1步新版本", + description=( + "反复修改爆款开头复刻第1步素材输入。" + "接口会软删除旧第1步以及第2、3、4、5步当前有效子任务,联动软删除关联 ChatGenerationTask," + "然后新建第1步 material_input 的 version+1,项目回到 waiting_user 状态。" + ), +) +async def update_material( + project_id: str = Path(..., description="总任务项目ID"), + req: HotOpeningMaterialUpdateRequest = Body(..., description="第1步素材输入修改参数,至少传一个字段"), + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + try: + project_id_value, step_id_value = await update_hot_opening_material_input( + db, + current_user=current_user, + project_id=project_id, + req=req, + ) + await db.commit() + except HTTPException: + await db.rollback() + raise + except Exception as exc: + await db.rollback() + raise HTTPException(status_code=500, detail=f"修改素材输入失败: {exc}") + + return HotOpeningActionOut( + message="素材输入已修改,旧步骤已软删除,请重新生成图片 AI 提词", + project_id=project_id_value, + step_id=step_id_value, + detail=await _reload_project_detail(db, current_user, project_id_value), + ) + + +@router.put( + "/tasks/{project_id}/steps/{step_id}/image-prompt", + response_model=HotOpeningActionOut, + summary="直接修改第2步图片 AI 优化提词", + description=( + "直接修改第2步图片 AI 优化提词,不调用 AI、不扣积分。" + "保存后会软删除第3、4、5步当前有效任务和关联 ChatGenerationTask," + "清空旧图片/视频结果,让用户从图片生成开始重新执行。" + ), +) +async def update_image_prompt( + project_id: str = Path(..., description="总任务项目ID"), + step_id: str = Path(..., description="第2步图片 AI 提词子任务ID"), + req: HotOpeningImagePromptUpdateRequest = Body(..., description="图片 AI 优化提词修改参数"), + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + try: + project, step = await update_hot_opening_image_prompt( + db, + current_user=current_user, + project_id=project_id, + step_id=step_id, + req=req, + ) + project_id_value = str(project.id) + step_id_value = str(step.id) + await db.commit() + except HTTPException: + await db.rollback() + raise + except Exception as exc: + await db.rollback() + raise HTTPException(status_code=500, detail=f"修改图片 AI 提词失败: {exc}") + + return HotOpeningActionOut( + message="图片 AI 提词已修改,后续步骤已软删除,请重新生成图片", + project_id=project_id_value, + step_id=step_id_value, + detail=await _reload_project_detail(db, current_user, project_id_value), + ) + + +@router.put( + "/tasks/{project_id}/steps/{step_id}/video-prompt-schema", + response_model=HotOpeningActionOut, + summary="修改第4步视频 AI 提词 JSON schema", + description=( + "修改第4步视频 AI 提词 JSON schema,不调用 AI、不扣积分。" + "前端提交的 schema 只作为 patch,服务端会锁定视频时长、比例、清晰度、帧率、推荐分辨率、" + "动作/镜头/动态时间规划数组长度和时间段、输出规格、质量控制、合规控制、schema_version、schema_usage。" + "最终提示词允许修改,但会清洗秒数、比例、分辨率、帧率等视频参数。保存后软删除第5步视频生成任务。" + ), +) +async def update_video_prompt_schema( + project_id: str = Path(..., description="总任务项目ID"), + step_id: str = Path(..., description="第4步视频 AI 提词子任务ID"), + req: HotOpeningVideoPromptSchemaUpdateRequest = Body(..., description="视频 AI 提词 schema 修改参数"), + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + try: + project, step = await update_hot_opening_video_prompt_schema( + db, + current_user=current_user, + project_id=project_id, + step_id=step_id, + req=req, + ) + project_id_value = str(project.id) + step_id_value = str(step.id) + await db.commit() + except HTTPException: + await db.rollback() + raise + except Exception as exc: + await db.rollback() + raise HTTPException(status_code=500, detail=f"修改视频 AI 提词 schema 失败: {exc}") + + return HotOpeningActionOut( + message="视频 AI 提词 schema 已修改,第5步视频生成任务已软删除,请重新生成视频", + project_id=project_id_value, + step_id=step_id_value, + detail=await _reload_project_detail(db, current_user, project_id_value), + ) + + +@router.post( + "/tasks/{project_id}/steps/{step_id}/generate-image-prompt", + response_model=HotOpeningActionOut, + summary="基于素材输入手动生成图片 AI 提词", + description=( + "基于第1步素材输入子任务手动生成第2步图片 AI 提词。" + "如果已存在旧的第2、3、4、5步,会先软删除旧步骤,再创建新的第2步。" + ), +) +async def generate_image_prompt( + project_id: str = Path(..., description="总任务项目ID"), + step_id: str = Path(..., description="第1步素材输入子任务ID"), + req: HotOpeningGenerateImagePromptRequest = Body(default_factory=HotOpeningGenerateImagePromptRequest, description="图片提词生成参数,当前无需传参,额外字段会忽略"), + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + _ = req + if celery_app is None: + raise HTTPException(status_code=503, detail="Celery未启用:请配置 REDIS_URL 或 CELERY_BROKER_URL 后启动 worker") + + try: + project, step = await submit_image_prompt_optimize(db, current_user=current_user, project_id=project_id, material_step_id=step_id) + project_id_value = str(project.id) + step_id_value = str(step.id) + await db.commit() + except HTTPException: + await db.rollback() + raise + except Exception as exc: + await db.rollback() + raise HTTPException(status_code=500, detail=f"图片提词任务创建失败: {exc}") + + from app.tasks.hot_opening_replicate_tasks import start_image_prompt_optimize + + try: + start_image_prompt_optimize.delay(project_id_value, step_id_value) + except Exception as exc: + await _mark_dispatch_failed_and_raise( + db, + current_user=current_user, + project_id=project_id_value, + step_id=step_id_value, + message=f"图片提词任务投递失败: {exc}", + ) + + return HotOpeningActionOut( + message="图片 AI 提词任务已提交", + project_id=project_id_value, + step_id=step_id_value, + detail=await _reload_project_detail(db, current_user, project_id_value), + ) + + +@router.post( + "/tasks/{project_id}/steps/{step_id}/generate-image", + response_model=HotOpeningActionOut, + summary="基于图片 AI 提词生成新项目图片", + description=( + "基于第2步图片 AI 提词生成新项目图片。调用时传入图片生成引擎和图片生成参数。" + "后端会创建第3步图片生成子任务,ChatGenerationTask 幂等键由后端按任务ID自动生成,不再使用前端幂等键。" + "图片生成媒体积分在创建 ChatGenerationTask 时扣除,生成失败走媒体积分退款。" + "如果已存在旧的第3、4、5步,会先软删除旧步骤,再创建新的第3步。" + ), +) +async def generate_image( + project_id: str = Path(..., description="总任务项目ID"), + step_id: str = Path(..., description="第2步图片 AI 提词子任务ID"), + req: HotOpeningGenerateImageRequest = Body(..., description="图片生成引擎和参数"), + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + if celery_app is None: + raise HTTPException(status_code=503, detail="Celery未启用:请配置 REDIS_URL 或 CELERY_BROKER_URL 后启动 worker") + + try: + project, step = await generate_image_from_prompt(db, current_user=current_user, project_id=project_id, prompt_step_id=step_id, req=req) + project_id_value = str(project.id) + step_id_value = str(step.id) + chat_task_id_value = step.chat_task_id + if not chat_task_id_value: + raise HTTPException(status_code=500, detail="图片生成任务创建失败:chat_task_id为空") + await db.commit() + except HTTPException: + await db.rollback() + raise + except Exception as exc: + await db.rollback() + raise HTTPException(status_code=500, detail=f"图片生成任务创建失败: {exc}") + + from app.tasks.generation_create_tasks import chatapi_create_generation_task + + try: + chatapi_create_generation_task.delay(chat_task_id_value) + except Exception as exc: + await _mark_dispatch_failed_and_raise( + db, + current_user=current_user, + project_id=project_id_value, + step_id=step_id_value, + message=f"图片生成任务投递失败: {exc}", + ) + + return HotOpeningActionOut( + message="图片生成任务已提交", + project_id=project_id_value, + step_id=step_id_value, + detail=await _reload_project_detail(db, current_user, project_id_value), + ) + + +@router.post( + "/tasks/{project_id}/steps/{step_id}/generate-video-prompt", + response_model=HotOpeningActionOut, + summary="基于图片结果手动生成视频 AI 提词", + description=( + "基于第3步图片生成子任务手动生成第4步视频 AI 提词 JSON schema。" + "视频时长、比例、分辨率集中在本步骤确定并写入 step.output_json.payload.params_used_for_prompt。" + "视频 AI 提词成功后按文本 token 扣积分;该文本积分不参与后续视频生成失败退款。" + "如果已存在旧的第4、5步,会先软删除旧步骤,再创建新的第4步。" + ), +) +async def generate_video_prompt( + project_id: str = Path(..., description="总任务项目ID"), + step_id: str = Path(..., description="第3步图片生成子任务ID"), + req: HotOpeningGenerateVideoPromptRequest = Body(..., description="视频提词生成参数,用于读取接口配置并规划视频时长、比例、分辨率"), + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + if celery_app is None: + raise HTTPException(status_code=503, detail="Celery未启用:请配置 REDIS_URL 或 CELERY_BROKER_URL 后启动 worker") + + try: + project, step = await submit_video_prompt_optimize(db, current_user=current_user, project_id=project_id, image_step_id=step_id, req=req) + project_id_value = str(project.id) + step_id_value = str(step.id) + await db.commit() + except HTTPException: + await db.rollback() + raise + except Exception as exc: + await db.rollback() + raise HTTPException(status_code=500, detail=f"视频提词任务创建失败: {exc}") + + from app.tasks.hot_opening_replicate_tasks import start_video_prompt_optimize + + try: + start_video_prompt_optimize.delay(project_id_value, step_id_value) + except Exception as exc: + await _mark_dispatch_failed_and_raise( + db, + current_user=current_user, + project_id=project_id_value, + step_id=step_id_value, + message=f"视频提词任务投递失败: {exc}", + ) + + return HotOpeningActionOut( + message="视频 AI 提词任务已提交", + project_id=project_id_value, + step_id=step_id_value, + detail=await _reload_project_detail(db, current_user, project_id_value), + ) + + +@router.post( + "/tasks/{project_id}/steps/{step_id}/generate-video", + response_model=HotOpeningActionOut, + summary="基于视频 AI 提词生成最终视频", + description=( + "基于第4步视频 AI 提词生成最终视频。请求体只需要选择视频生成引擎 engine_id。" + "视频时长、比例、分辨率从第4步视频提词优化结果读取,不再由本接口动态传入。" + "关联 ChatGenerationTask 的 original_prompt 和 optimized_prompt 都使用第4步生成的 prompt_schema JSON 字符串。" + "ChatGenerationTask 幂等键由后端按任务ID自动生成;视频生成媒体积分失败时走退款。" + "如果已存在旧的第5步,会先软删除旧步骤,再创建新的第5步。" + ), +) +async def generate_video( + project_id: str = Path(..., description="总任务项目ID"), + step_id: str = Path(..., description="第4步视频 AI 提词子任务ID"), + req: HotOpeningGenerateVideoRequest = Body(..., description="视频生成参数:只传 engine_id,其它视频参数继承第4步视频提词结果"), + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + if celery_app is None: + raise HTTPException(status_code=503, detail="Celery未启用:请配置 REDIS_URL 或 CELERY_BROKER_URL 后启动 worker") + + try: + project, step = await generate_video_from_prompt(db, current_user=current_user, project_id=project_id, prompt_step_id=step_id, req=req) + project_id_value = str(project.id) + step_id_value = str(step.id) + chat_task_id_value = step.chat_task_id + if not chat_task_id_value: + raise HTTPException(status_code=500, detail="视频生成任务创建失败:chat_task_id为空") + await db.commit() + except HTTPException: + await db.rollback() + raise + except Exception as exc: + await db.rollback() + raise HTTPException(status_code=500, detail=f"视频生成任务创建失败: {exc}") + + from app.tasks.generation_create_tasks import chatapi_create_generation_task + + try: + chatapi_create_generation_task.delay(chat_task_id_value) + except Exception as exc: + await _mark_dispatch_failed_and_raise( + db, + current_user=current_user, + project_id=project_id_value, + step_id=step_id_value, + message=f"视频生成任务投递失败: {exc}", + ) + + return HotOpeningActionOut( + message="视频生成任务已提交", + project_id=project_id_value, + step_id=step_id_value, + detail=await _reload_project_detail(db, current_user, project_id_value), + ) + + +@router.delete( + "/tasks/{project_id}", + response_model=HotOpeningDeleteOut, + summary="删除爆款开头复刻总任务项目", + description="软删除爆款开头复刻总任务项目,并联动软删除当前有效子任务和关联的 ChatGenerationTask。", +) +async def delete_task( + project_id: str = Path(..., description="总任务项目ID"), + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + try: + result = await delete_hot_opening_project(db, current_user=current_user, project_id=project_id) + await db.commit() + return result + except HTTPException: + await db.rollback() + raise + except Exception as exc: + await db.rollback() + raise HTTPException(status_code=500, detail=f"删除爆款开头复刻项目失败: {exc}") diff --git a/video-gen-api/app/api/v1/payments.py b/video-gen-api/app/api/v1/payments.py index 6b2d8ca4..0390ea7c 100644 --- a/video-gen-api/app/api/v1/payments.py +++ b/video-gen-api/app/api/v1/payments.py @@ -1,23 +1,61 @@ +import logging + from fastapi import APIRouter, Depends, HTTPException, Request from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession +logger = logging.getLogger("payment") + from app.dependencies import get_db, get_current_user from app.models.user import User from app.models.payment_order import PaymentOrder from app.models.recharge_package import RechargePackage from app.schemas.payment import RechargeRequest, PaymentOrderOut -from app.services.payment import create_recharge_order, verify_wechat_callback, verify_alipay_callback, process_payment_success +from app.services.payment import ( + create_recharge_order, + verify_wechat_callback, + verify_alipay_callback, + process_payment_success_by_order_no, + process_refund, + _get_payment_configs, + _close_alipay_order, + _get_order_expire_seconds, +) router = APIRouter(prefix="/payments", tags=["payments"]) +@router.get("/methods") +async def get_payment_methods( + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db) +): + """Return which payment methods are enabled (from admin config).""" + from app.services.payment import _get_payment_configs + configs = await _get_payment_configs(db) + return { + "alipay": configs.get("payment_alipay_enabled", "").lower() == "true", + "wechat": configs.get("payment_wechat_enabled", "").lower() == "true", + } + + @router.post("/recharge", response_model=PaymentOrderOut) async def recharge( req: RechargeRequest, current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): + if req.method not in ("wechat", "alipay"): + raise HTTPException(status_code=400, detail="不支持的支付方式") + + # Check if the selected payment method is enabled in admin config + from app.services.payment import _get_payment_configs, _is_mock_mode + configs = await _get_payment_configs(db) + if not _is_mock_mode(configs): + enabled_key = f"payment_{req.method}_enabled" + if configs.get(enabled_key, "").lower() != "true": + raise HTTPException(status_code=400, detail="该支付方式未启用") + result = await db.execute( select(RechargePackage).where( RechargePackage.id == req.plan, @@ -28,44 +66,60 @@ async def recharge( pkg = result.scalar_one_or_none() if not pkg: raise HTTPException(status_code=400, detail="无效的套餐") - order = await create_recharge_order( - db, - current_user.id, - credits=pkg.credits, - price=pkg.price, - label=pkg.name, - bonus_credits=pkg.bonus_credits, - ) + try: + order = await create_recharge_order( + db, + current_user.id, + credits=pkg.credits, + price=pkg.price, + label=pkg.name, + bonus_credits=pkg.bonus_credits, + method=req.method, + ) + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) return order @router.post("/wechat/callback") async def wechat_callback(request: Request, db: AsyncSession = Depends(get_db)): data = await request.json() - if not await verify_wechat_callback(data): + if not await verify_wechat_callback(data, db): raise HTTPException(status_code=400, detail="签名验证失败") order_no = data.get("out_trade_no") - result = await db.execute( - select(PaymentOrder).where(PaymentOrder.order_no == order_no).limit(1) - ) - order = result.scalar_one_or_none() - if order: - await process_payment_success(db, order.id) + if order_no: + await process_payment_success_by_order_no(db, order_no) return {"code": "SUCCESS", "message": "OK"} @router.post("/alipay/callback") async def alipay_callback(request: Request, db: AsyncSession = Depends(get_db)): - data = await request.form() - if not await verify_alipay_callback(dict(data)): - raise HTTPException(status_code=400, detail="签名验证失败") - order_no = data.get("out_trade_no") - result = await db.execute( - select(PaymentOrder).where(PaymentOrder.order_no == order_no).limit(1) + form_data = await request.form() + data = dict(form_data) + + logger.info( + f"ALIPAY_CALLBACK order_no={data.get('out_trade_no')} " + f"data={data}" ) - order = result.scalar_one_or_none() - if order: - await process_payment_success(db, order.id) + + # Verify signature first + if not await verify_alipay_callback(data, db): + raise HTTPException(status_code=400, detail="签名验证失败") + + # Check trade_status – only "TRADE_SUCCESS" and "TRADE_FINISHED" mean paid + trade_status = data.get("trade_status", "") + if trade_status not in ("TRADE_SUCCESS", "TRADE_FINISHED"): + logger.info(f"Alipay callback trade_status={trade_status}, ignoring") + return "success" + + order_no = data.get("out_trade_no") + trade_no = data.get("trade_no", "") + total_amount_str = data.get("total_amount", "") + total_amount = float(total_amount_str) if total_amount_str else None + + if order_no: + await process_payment_success_by_order_no(db, order_no, trade_no, total_amount) + return "success" @@ -74,9 +128,72 @@ async def list_orders( current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): + # Auto-expire stale pending orders before returning + from app.services.payment import _check_and_expire_order result = await db.execute( select(PaymentOrder) .where(PaymentOrder.user_id == current_user.id) .order_by(PaymentOrder.created_at.desc()) ) - return result.scalars().all() + orders = result.scalars().all() + for o in orders: + await _check_and_expire_order(db, o) + return orders + + +@router.get("/orders/{order_no}", response_model=PaymentOrderOut) +async def get_order( + order_no: str, + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + from app.services.payment import _check_and_expire_order + result = await db.execute( + select(PaymentOrder) + .where( + PaymentOrder.order_no == order_no, + PaymentOrder.user_id == current_user.id, + ) + .limit(1) + ) + order = result.scalar_one_or_none() + if not order: + raise HTTPException(status_code=404, detail="订单不存在") + # Auto-expire if needed + await _check_and_expire_order(db, order) + return order + + +@router.post("/orders/{order_no}/cancel") +async def cancel_order( + order_no: str, + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + """Cancel a pending order. Only the order owner can cancel, only if still pending.""" + result = await db.execute( + select(PaymentOrder).where( + PaymentOrder.order_no == order_no, + PaymentOrder.user_id == current_user.id, + ).limit(1) + ) + order = result.scalar_one_or_none() + if not order: + raise HTTPException(status_code=404, detail="订单不存在") + if order.status != "pending": + raise HTTPException(status_code=400, detail=f"订单状态为{order.status},无法取消") + + # If it's an Alipay order, call close API first + if order.payment_method == "alipay": + db_configs = await _get_payment_configs(db) + try: + await _close_alipay_order(db, order, db_configs) + except Exception as e: + logger.exception(f"Failed to close Alipay order {order_no}: {e}") + + order.status = "cancelled" + await db.flush() + logger.info( + f"ORDER_CANCELLED order_no={order_no} user={current_user.id} amount={order.amount}" + ) + return {"ok": True} diff --git a/video-gen-api/app/api/v1/test.py b/video-gen-api/app/api/v1/test.py index c00efa3e..02961f89 100644 --- a/video-gen-api/app/api/v1/test.py +++ b/video-gen-api/app/api/v1/test.py @@ -5,6 +5,6 @@ router = APIRouter(prefix="/test", tags=["test"]) @router.get("/index") -async def test(req: Request): +async def test(): return {"message": "test","code":200} diff --git a/video-gen-api/app/api/v1/user_oauth.py b/video-gen-api/app/api/v1/user_oauth.py new file mode 100644 index 00000000..d1582a05 --- /dev/null +++ b/video-gen-api/app/api/v1/user_oauth.py @@ -0,0 +1,124 @@ +from fastapi import APIRouter, Depends, HTTPException, Query, status +from sqlalchemy.ext.asyncio import AsyncSession + +from app.dependencies import get_current_user, get_db +from app.models.user import User +from app.schemas.user_oauth import RequestOAuthRequest, RequestOAuthResponse, UserOAuthOut +from app.services.user_oauth_service import ( + build_oauth_url, + get_account_info_by_type, + get_token_by_type, + save_oauth_token, +) + +router = APIRouter(prefix="/user-oauth", tags=["user-oauth"]) + + +@router.post( + "/request_oauth", + summary="获取授权链接", + description="用户提交oauth_type,返回对应的第三方授权链接", + response_model=RequestOAuthResponse, +) +async def request_oauth( + req: RequestOAuthRequest, + current_user: User = Depends(get_current_user), +): + try: + auth_url = await build_oauth_url(req.oauth_type, current_user.id) + return {"auth_url": auth_url} + except ValueError as e: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=str(e), + ) + + +@router.get( + "/juliang_callback", + summary="巨量授权回调", + description="巨量引擎授权回调地址,接收code和state参数,获取token并保存", +) +async def juliang_callback( + auth_code: str = Query(..., description="第三方返回的授权码"), + state: str = Query(..., description="请求时传递的自定义参数"), + db: AsyncSession = Depends(get_db), +): + try: + parts = state.split(":") + if len(parts) != 4: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="无效的state参数", + ) + + oauth_type = int(parts[0]) + user_id = parts[1] + app_id = parts[2] + app_type = parts[3] + + token = await get_token_by_type(auth_code, oauth_type, app_id, app_type) + account_info = await get_account_info_by_type(token, oauth_type, app_type) + user_oauth = await save_oauth_token(db, user_id, oauth_type, token, account_info, app_id) + + return { + "message": "授权成功", + "code": 0, + "data": UserOAuthOut.model_validate(user_oauth), + } + + except ValueError as e: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=str(e), + ) + except Exception as e: + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail=f"授权失败: {str(e)}", + ) + + +@router.get( + "/callback", + summary="通用授权回调", + description="其他平台授权回调地址,接收code和state参数", +) +async def oauth_callback( + code: str = Query(..., description="第三方返回的授权码"), + state: str = Query(..., description="请求时传递的自定义参数"), + db: AsyncSession = Depends(get_db), +): + try: + parts = state.split(":") + if len(parts) != 4: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="无效的state参数", + ) + + oauth_type = int(parts[0]) + user_id = parts[1] + app_id = parts[2] + app_type = parts[3] + + token = await get_token_by_type(code, oauth_type, app_id, app_type) + account_info = await get_account_info_by_type(token, oauth_type, app_type) + user_oauth = await save_oauth_token(db, user_id, oauth_type, token, account_info, app_id) + + return { + "message": "授权成功", + "code": 0, + "data": UserOAuthOut.model_validate(user_oauth), + } + + except ValueError as e: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=str(e), + ) + except Exception as e: + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail=f"授权失败: {str(e)}", + ) \ No newline at end of file diff --git a/video-gen-api/app/api/v1/user_oauth_app.py b/video-gen-api/app/api/v1/user_oauth_app.py new file mode 100644 index 00000000..6f9bfbb2 --- /dev/null +++ b/video-gen-api/app/api/v1/user_oauth_app.py @@ -0,0 +1,96 @@ +from fastapi import APIRouter, Depends, HTTPException, Query, status +from sqlalchemy.ext.asyncio import AsyncSession + +from app.dependencies import get_admin_user, get_db +from app.models.user import User +from app.schemas.user_oauth_app import UserOAuthAppCreate, UserOAuthAppOut, UserOAuthAppUpdate +from app.services.user_oauth_app_service import ( + create_user_oauth_app, + delete_user_oauth_app, + get_user_oauth_app_by_id, + list_user_oauth_apps, + update_user_oauth_app, +) + +router = APIRouter(prefix="/admin/user-oauth-apps", tags=["oauth"]) + + +@router.get("/list", summary="获取用户授权应用列表") +async def list_apps( + page: int = Query(1, ge=1), + page_size: int = Query(20, ge=1, le=100), + open_type: int | None = Query(None, ge=1, le=10, description="开户方式"), + status: int | None = Query(None, ge=1, le=2, description="应用状态,1=正常,2=禁用"), + admin: User = Depends(get_admin_user), + db: AsyncSession = Depends(get_db), + app_id: str | None = Query(None, max_length=255, description="应用id"), +): + result = await list_user_oauth_apps(db, page, page_size, open_type, status, admin.id, app_id) + return { + "total": result["total"], + "page": result["page"], + "page_size": result["page_size"], + "items": [UserOAuthAppOut.model_validate(item) for item in result["items"]], + } + + +@router.post("/create", summary="创建用户授权应用") +async def create_app( + req: UserOAuthAppCreate, + admin: User = Depends(get_admin_user), + db: AsyncSession = Depends(get_db), +): + try: + app = await create_user_oauth_app(db, req.app_id, req.secret, req.open_type, admin.id, req.count, req.auth_url, req.company) + return UserOAuthAppOut.model_validate(app) + except ValueError as e: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=str(e), + ) + + +@router.get("/read/{id}", summary="获取用户授权应用详情") +async def get_app( + id: str, + admin: User = Depends(get_admin_user), + db: AsyncSession = Depends(get_db), +): + app = await get_user_oauth_app_by_id(db, id) + if not app: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="应用不存在", + ) + return UserOAuthAppOut.model_validate(app) + + +@router.post("/update/{id}", summary="更新用户授权应用") +async def update_app( + id: str, + req: UserOAuthAppUpdate, + admin: User = Depends(get_admin_user), + db: AsyncSession = Depends(get_db), +): + app = await update_user_oauth_app(db, id, req.secret, req.open_type, req.status, req.count, req.auth_url, req.company, admin.id) + if not app: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="应用不存在", + ) + return UserOAuthAppOut.model_validate(app) + + +@router.get("/delete/{id}", summary="删除用户授权应用") +async def delete_app( + id: str, + admin: User = Depends(get_admin_user), + db: AsyncSession = Depends(get_db), +): + success = await delete_user_oauth_app(db, id, admin.id) + if not success: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="应用不存在", + ) + return {"message": "删除成功"} \ No newline at end of file diff --git a/video-gen-api/app/config.py b/video-gen-api/app/config.py index 5b3bf258..334175a3 100644 --- a/video-gen-api/app/config.py +++ b/video-gen-api/app/config.py @@ -56,7 +56,8 @@ class Settings(BaseSettings): ALIPAY_APP_ID: str = "" ALIPAY_PRIVATE_KEY: str = "" ALIPAY_PUBLIC_KEY: str = "" - PAYMENT_MOCK: bool = True + ALIPAY_NOTIFY_URL: str = "" + PAYMENT_MOCK: bool = False # Default off; use admin panel to enable for testing STORAGE_TYPE: str = "local" STORAGE_LOCAL_PATH: str = "./storage/generate/videos" @@ -86,7 +87,7 @@ class Settings(BaseSettings): # ChatAPI async generation pipeline settings CELERY_BROKER_URL: str = "" CELERY_RESULT_BACKEND: str = "" - CHATAPI_REQUEST_TIMEOUT_SECONDS: int = 120 + CHATAPI_REQUEST_TIMEOUT_SECONDS: int = 180 CHATAPI_VIDEO_FPS: float = 0.5 CHATAPI_ASYNC_MAX_RETRIES: int = 3 CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS: int = 30 @@ -127,5 +128,11 @@ class Settings(BaseSettings): RESOURCE_SIGN_ARG_EXPIRE: str = "exp" RESOURCE_SIGN_ARG_SIGNATURE: str = "sign" + # 爆款开头复刻默认配置。素材校验由前端完成,后端只接收已有上传接口返回的链接。 + HOT_OPENING_DEFAULT_VIDEO_DURATION: int = 4 + HOT_OPENING_DEFAULT_VIDEO_RATIO: str = "9:16" + HOT_OPENING_DEFAULT_VIDEO_RESOLUTION: str = "480p" + HOT_OPENING_DEFAULT_TARGET_PLATFORM: str = "抖音" + settings = Settings() diff --git a/video-gen-api/app/enums/__init__.py b/video-gen-api/app/enums/__init__.py new file mode 100644 index 00000000..a27058d0 --- /dev/null +++ b/video-gen-api/app/enums/__init__.py @@ -0,0 +1,3 @@ +from app.enums.common import * +from app.enums.hot_opening_replicate import * +from app.enums.video_prompt_schema import * diff --git a/video-gen-api/app/enums/common.py b/video-gen-api/app/enums/common.py new file mode 100644 index 00000000..fbde775e --- /dev/null +++ b/video-gen-api/app/enums/common.py @@ -0,0 +1,57 @@ +from __future__ import annotations + +from enum import StrEnum + + +class ModuleProjectStatusEnum(StrEnum): + """通用模块项目状态。""" + + PENDING = "pending" + WAITING_USER = "waiting_user" + PROCESSING = "processing" + COMPLETED = "completed" + FAILED = "failed" + CANCELLED = "cancelled" + + +class ModuleStepStatusEnum(StrEnum): + """通用模块子任务状态。""" + + PENDING = "pending" + WAITING_USER = "waiting_user" + PROCESSING = "processing" + COMPLETED = "completed" + FAILED = "failed" + CANCELLED = "cancelled" + + +class ModuleEventTypeEnum(StrEnum): + """通用模块事件类型。""" + + PROJECT_CREATED = "PROJECT_CREATED" + PROJECT_DELETED = "PROJECT_DELETED" + STEP_CREATED = "STEP_CREATED" + STEP_UPDATED = "STEP_UPDATED" + SOFT_DELETE_STEPS = "SOFT_DELETE_STEPS" + IMAGE_PROMPT_SUBMITTED = "IMAGE_PROMPT_SUBMITTED" + IMAGE_PROMPT_SUCCESS = "IMAGE_PROMPT_SUCCESS" + IMAGE_PROMPT_FAILED = "IMAGE_PROMPT_FAILED" + IMAGE_GENERATE_SUBMITTED = "IMAGE_GENERATE_SUBMITTED" + IMAGE_GENERATE_SUCCESS = "IMAGE_GENERATE_SUCCESS" + VIDEO_PROMPT_SUBMITTED = "VIDEO_PROMPT_SUBMITTED" + VIDEO_PROMPT_SUCCESS = "VIDEO_PROMPT_SUCCESS" + VIDEO_PROMPT_FAILED = "VIDEO_PROMPT_FAILED" + VIDEO_GENERATE_SUBMITTED = "VIDEO_GENERATE_SUBMITTED" + VIDEO_GENERATE_SUCCESS = "VIDEO_GENERATE_SUCCESS" + CHAT_TASK_FAILED = "CHAT_TASK_FAILED" + CHAT_TASK_CANCELLED = "CHAT_TASK_CANCELLED" + MEDIA_REFUND = "MEDIA_REFUND" + PROMPT_BILLING_SUCCESS = "PROMPT_BILLING_SUCCESS" + PROMPT_BILLING_FAILED = "PROMPT_BILLING_FAILED" + + +class ModulePromptTypeEnum(StrEnum): + """通用模块提词结果类型。""" + + IMAGE_PROMPT = "image_prompt" + VIDEO_PROMPT = "video_prompt" diff --git a/video-gen-api/app/enums/hot_opening_replicate.py b/video-gen-api/app/enums/hot_opening_replicate.py new file mode 100644 index 00000000..a43d11ba --- /dev/null +++ b/video-gen-api/app/enums/hot_opening_replicate.py @@ -0,0 +1,32 @@ +from __future__ import annotations + +from enum import StrEnum + + +class ModuleCodeEnum(StrEnum): + """可复用模块编码。""" + + HOT_OPENING_REPLICATE = "hot_opening_replicate" + + +class HotOpeningStepCodeEnum(StrEnum): + """爆款开头复刻子任务步骤编码。""" + + MATERIAL_INPUT = "material_input" + IMAGE_PROMPT_OPTIMIZE = "image_prompt_optimize" + IMAGE_GENERATE = "image_generate" + VIDEO_PROMPT_OPTIMIZE = "video_prompt_optimize" + VIDEO_GENERATE = "video_generate" + + +class HotOpeningGenerationModeEnum(StrEnum): + """复用 ChatGenerationTask 时使用的 generation_mode。""" + + HOT_OPENING_REPLICATE = "hot_opening_replicate" + + + +class HotOpeningStepIOSchemaVersionEnum(StrEnum): + """爆款开头复刻子任务 input_json/output_json 结构版本。""" + + V1 = "hot_opening_step_io_v1" diff --git a/video-gen-api/app/enums/video_prompt_schema.py b/video-gen-api/app/enums/video_prompt_schema.py new file mode 100644 index 00000000..a9541ad0 --- /dev/null +++ b/video-gen-api/app/enums/video_prompt_schema.py @@ -0,0 +1,15 @@ +from __future__ import annotations + +from enum import StrEnum + + +class PromptSchemaVersionEnum(StrEnum): + """视频提词 schema 版本。""" + + CLIENT_V1 = "video_prompt_schema_client_v1" + + +class VideoPromptSchemaUsageEnum(StrEnum): + """视频提词 schema 用途。""" + + CLIENT_DISPLAY = "客户端可展示的AI视频生成提词结构" diff --git a/video-gen-api/app/main.py b/video-gen-api/app/main.py index 1ceeb94a..a9049055 100644 --- a/video-gen-api/app/main.py +++ b/video-gen-api/app/main.py @@ -35,13 +35,53 @@ async def lifespan(app: FastAPI): from app.services.video_queue import task_queue await task_queue.recover() queue_task = asyncio.create_task(task_queue.run()) + + # Background task: auto-expire pending payment orders and sync status + async def _order_expiry_loop(): + from app.services.payment import expire_all_pending_orders, sync_pending_orders + from logging import getLogger + bg_logger = getLogger("payment") + while True: + try: + async with async_session() as db: + # 同步待支付订单状态(检查支付宝实际支付状态 + sync_count = await sync_pending_orders(db) + if sync_count > 0: + bg_logger.info(f"Synced {sync_count} pending payment order(s)") + + # 自动过期订单 + n = await expire_all_pending_orders(db) + if n > 0: + bg_logger.info(f"Auto-expired {n} pending payment order(s)") + except Exception as e: + bg_logger.error(f"Order expiry loop error: {e}") + await asyncio.sleep(60) # check every minute + + expiry_task = asyncio.create_task(_order_expiry_loop()) + # 启动时立即同步一次未支付订单 + asyncio.create_task(asyncio.sleep(5)) # 等待5秒后再同步,让系统完全启动 + async def startup_sync(): + await asyncio.sleep(5) + from app.services.payment import sync_pending_orders + from logging import getLogger + bg_logger = getLogger("payment") + try: + async with async_session() as db: + sync_count = await sync_pending_orders(db) + if sync_count > 0: + bg_logger.info(f"Startup: Synced {sync_count} pending payment order(s)") + except Exception as e: + bg_logger.error(f"Startup sync error: {e}") + asyncio.create_task(startup_sync()) + app.state.db_session_factory = async_session yield task_queue.stop() await queue_task + expiry_task.cancel() await close_database() await close_redis() diff --git a/video-gen-api/app/middleware/request_encrypt.py b/video-gen-api/app/middleware/request_encrypt.py index 5d5d2651..0d4113ed 100644 --- a/video-gen-api/app/middleware/request_encrypt.py +++ b/video-gen-api/app/middleware/request_encrypt.py @@ -42,6 +42,11 @@ class RequestEncryptMiddleware(BaseHTTPMiddleware): async def dispatch( self, request: Request, call_next: RequestResponseEndpoint ) -> Response: + # 白名单:支付回调接口不需要加密/解密 + path = request.url.path + if "/payments/alipay/callback" in path or "/payments/wechat/callback" in path: + return await call_next(request) + encrypted = request.headers.get("X-Encrypted", "").lower() == "true" if not encrypted: return await call_next(request) diff --git a/video-gen-api/app/models/__init__.py b/video-gen-api/app/models/__init__.py index e26d99f1..82311bb8 100644 --- a/video-gen-api/app/models/__init__.py +++ b/video-gen-api/app/models/__init__.py @@ -21,6 +21,11 @@ from app.models.chat_provider_call_log import ChatProviderCallLog from app.models.generated_resource import GeneratedResource from app.models.user_resource_month_stat import UserResourceMonthStat from app.models.user_resource_total_stat import UserResourceTotalStat +from app.models.module_generation_project import ModuleGenerationProject +from app.models.module_generation_step import ModuleGenerationStep +from app.models.user_oauth import UserOAuth +from app.models.user_oauth_account import UserOAuthAccount +from app.models.user_oauth_app import UserOAuthApp __all__ = [ "Base", "TimestampMixin", "SoftDeleteMixin", "engine", "async_session", @@ -31,4 +36,6 @@ __all__ = [ "MenuConfig", "RechargePackage", "OperationLog", "ChatGenerationTask", "ChatGenerationTaskEvent", "ChatProviderCallLog", "GeneratedResource", "UserResourceMonthStat", "UserResourceTotalStat", + "ModuleGenerationProject", "ModuleGenerationStep", + "UserOAuth", "UserOAuthAccount", "UserOAuthApp", ] diff --git a/video-gen-api/app/models/chat_generation_task.py b/video-gen-api/app/models/chat_generation_task.py index 2bd74af6..7307e1f8 100644 --- a/video-gen-api/app/models/chat_generation_task.py +++ b/video-gen-api/app/models/chat_generation_task.py @@ -1,6 +1,6 @@ from datetime import datetime -from sqlalchemy import DateTime, Float, ForeignKey, Index, Integer, String, Text +from sqlalchemy import DateTime, Float, ForeignKey, Index, Integer, String, Text, text from sqlalchemy.orm import Mapped, mapped_column from app.models.base import Base, TimestampMixin, SoftDeleteMixin @@ -18,7 +18,14 @@ class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin): __table_args__ = ( # 防止前端按钮连点/网络重试时同一个 idempotency_key 并发创建多条任务。 # nullable unique 兼容不传 idempotency_key 的普通请求。 - Index("uq_chat_generation_tasks_user_mode_idempotency", "user_id", "generation_mode", "idempotency_key", unique=True), + Index( + "uq_chat_generation_tasks_user_mode_idempotency", + "user_id", + "generation_mode", + "idempotency_key", + unique=True, + postgresql_where=text("deleted_at IS NULL AND idempotency_key IS NOT NULL"), + ), ) diff --git a/video-gen-api/app/models/module_generation_project.py b/video-gen-api/app/models/module_generation_project.py new file mode 100644 index 00000000..8aed2b53 --- /dev/null +++ b/video-gen-api/app/models/module_generation_project.py @@ -0,0 +1,47 @@ +from __future__ import annotations + +from datetime import datetime + +from sqlalchemy import DateTime, ForeignKey, Index, String, Text, text +from sqlalchemy.orm import Mapped, mapped_column + +from app.models.base import Base, SoftDeleteMixin, TimestampMixin + + +class ModuleGenerationProject(Base, TimestampMixin, SoftDeleteMixin): + """通用模块生成项目/总任务表。 + + 说明: + - 本表的 id 就是前端理解的“项目ID/总任务ID”,不再额外保存 project_id。 + - 通过 module 区分业务模块,后续其它功能也可以复用这张总任务项目表。 + - 爆款开头复刻使用 module=hot_opening_replicate。 + """ + + __tablename__ = "module_generation_projects" + __table_args__ = ( + Index("idx_module_generation_projects_user_module", "user_id", "module"), + Index("idx_module_generation_projects_status", "module", "status"), + Index( + "uq_module_generation_projects_user_module_idempotency", + "user_id", + "module", + "idempotency_key", + unique=True, + postgresql_where=text("deleted_at IS NULL AND idempotency_key IS NOT NULL"), + ), + ) + + id: Mapped[str] = mapped_column(String(32), primary_key=True) + user_id: Mapped[str] = mapped_column( + String(32), ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False + ) + module: Mapped[str] = mapped_column(String(64), index=True, nullable=False) + title: Mapped[str | None] = mapped_column(String(160), nullable=True) + status: Mapped[str] = mapped_column(String(32), default="pending", index=True) + current_step_code: Mapped[str | None] = mapped_column(String(64), nullable=True, index=True) + final_image_url: Mapped[str | None] = mapped_column(String(512), nullable=True) + final_video_url: Mapped[str | None] = mapped_column(String(512), nullable=True) + final_video_cover_url: Mapped[str | None] = mapped_column(String(512), nullable=True) + error_message: Mapped[str | None] = mapped_column(Text, nullable=True) + idempotency_key: Mapped[str | None] = mapped_column(String(64), nullable=True, index=True) + completed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True) diff --git a/video-gen-api/app/models/module_generation_step.py b/video-gen-api/app/models/module_generation_step.py new file mode 100644 index 00000000..81156be1 --- /dev/null +++ b/video-gen-api/app/models/module_generation_step.py @@ -0,0 +1,74 @@ +from __future__ import annotations + +from datetime import datetime +from typing import Any + +from sqlalchemy import Boolean, DateTime, ForeignKey, Index, Integer, JSON, String, Text +from sqlalchemy.dialects.postgresql import JSONB +from sqlalchemy.orm import Mapped, mapped_column + +from app.models.base import Base, SoftDeleteMixin, TimestampMixin + +_STEP_JSON_TYPE = JSON().with_variant(JSONB, "postgresql") + + +class ModuleGenerationStep(Base, TimestampMixin, SoftDeleteMixin): + """通用模块生成步骤表。 + + 爆款开头复刻固定步骤: + 1 material_input + 2 image_prompt_optimize + 3 image_generate + 4 video_prompt_optimize + 5 video_generate + + input_json / output_json 使用 JSON/JSONB 存储。 + 建议结构: + input_json = { + "schema_version": "hot_opening_step_io_v1", + "step_code": "...", + "source": {...}, + "payload": {...}, + "context": {...} + } + output_json = { + "schema_version": "hot_opening_step_io_v1", + "step_code": "...", + "status": "completed|failed|...", + "payload": {...}, + "result": {...}, + "usage": {...}, + "error": {...} + } + """ + + __tablename__ = "module_generation_steps" + __table_args__ = ( + Index("idx_module_generation_steps_project_current", "project_id", "is_current", "deleted_at"), + Index("idx_module_generation_steps_project_code", "project_id", "step_code", "is_current"), + Index("idx_module_generation_steps_chat_task", "chat_task_id"), + ) + + id: Mapped[str] = mapped_column(String(32), primary_key=True) + project_id: Mapped[str] = mapped_column( + String(32), ForeignKey("module_generation_projects.id", ondelete="CASCADE"), index=True, nullable=False + ) + user_id: Mapped[str] = mapped_column( + String(32), ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False + ) + module: Mapped[str] = mapped_column(String(64), index=True, nullable=False) + step_index: Mapped[int] = mapped_column(Integer, index=True, nullable=False) + step_code: Mapped[str] = mapped_column(String(64), index=True, nullable=False) + status: Mapped[str] = mapped_column(String(32), default="pending", index=True) + version: Mapped[int] = mapped_column(Integer, default=1, nullable=False) + is_current: Mapped[bool] = mapped_column(Boolean, default=True, index=True) + parent_step_id: Mapped[str | None] = mapped_column(String(32), nullable=True, index=True) + source_step_id: Mapped[str | None] = mapped_column(String(32), nullable=True, index=True) + chat_task_id: Mapped[str | None] = mapped_column( + String(32), ForeignKey("chat_generation_tasks.id", ondelete="SET NULL"), nullable=True, index=True + ) + input_json: Mapped[dict[str, Any] | list[Any] | None] = mapped_column(_STEP_JSON_TYPE, nullable=True) + output_json: Mapped[dict[str, Any] | list[Any] | None] = mapped_column(_STEP_JSON_TYPE, nullable=True) + error_message: Mapped[str | None] = mapped_column(Text, nullable=True) + started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True) + completed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True) diff --git a/video-gen-api/app/models/payment_order.py b/video-gen-api/app/models/payment_order.py index 0e240d4a..d378d7ee 100644 --- a/video-gen-api/app/models/payment_order.py +++ b/video-gen-api/app/models/payment_order.py @@ -22,3 +22,9 @@ class PaymentOrder(Base, TimestampMixin): DateTime(timezone=True), nullable=True ) trade_no: Mapped[str | None] = mapped_column(String(128), nullable=True) + # Refund fields + refund_trade_no: Mapped[str | None] = mapped_column(String(128), nullable=True) + refunded_at: Mapped[datetime | None] = mapped_column( + DateTime(timezone=True), nullable=True + ) + refund_amount: Mapped[float | None] = mapped_column(Float, nullable=True) diff --git a/video-gen-api/app/models/user_oauth.py b/video-gen-api/app/models/user_oauth.py new file mode 100644 index 00000000..4b346220 --- /dev/null +++ b/video-gen-api/app/models/user_oauth.py @@ -0,0 +1,58 @@ +from datetime import datetime + +from sqlalchemy import Boolean, DateTime, ForeignKey, Integer, String, Text +from sqlalchemy.orm import Mapped, mapped_column + +from app.models.base import Base, TimestampMixin, SoftDeleteMixin + + +class UserOAuth(Base, TimestampMixin, SoftDeleteMixin): + __tablename__ = "user_oauth" + + id: Mapped[str] = mapped_column( + String(32), primary_key=True, comment="主键" + ) + account_id: Mapped[str] = mapped_column( + String(64), nullable=False, index=True, comment="授权账户id" + ) + account_name: Mapped[str] = mapped_column( + String(128), nullable=False, comment="授权账户name" + ) + account_role: Mapped[str | None] = mapped_column( + String(64), nullable=True, comment="授权账户角色" + ) + account_username: Mapped[str | None] = mapped_column( + String(128), nullable=True, comment="授权账户登录账号" + ) + user_id: Mapped[str] = mapped_column( + String(32), nullable=False, index=True, + comment="用户id" + ) + open_type: Mapped[int] = mapped_column( + Integer, nullable=False, + comment="开户方式(1=千川,2=广告,3=本地推,4=星图,5=快手代理商,6=巨量星图,7=巨量服务单,8=腾讯服务单,9=腾讯营销K2,10=腾讯营销K3)" + ) + port_type: Mapped[int] = mapped_column( + Integer, nullable=False, + comment="平台端口(1=巨量,2=磁力,3=巨量星图,4=服务单,5=腾讯)" + ) + appid: Mapped[str | None] = mapped_column( + String(64), nullable=True, comment="授权应用id" + ) + access_token: Mapped[str | None] = mapped_column( + Text, nullable=True, comment="授权token" + ) + access_token_expired: Mapped[datetime | None] = mapped_column( + DateTime(timezone=True), nullable=True, + comment="token过期时间" + ) + refresh_token: Mapped[str | None] = mapped_column( + Text, nullable=True, comment="授权刷新token" + ) + refresh_token_expired: Mapped[datetime | None] = mapped_column( + DateTime(timezone=True), nullable=True, + comment="刷新token过期时间" + ) + material_auth_status: Mapped[bool] = mapped_column( + Boolean, default=False, comment="是否敏感物料授权(true=是,false=否)" + ) \ No newline at end of file diff --git a/video-gen-api/app/models/user_oauth_account.py b/video-gen-api/app/models/user_oauth_account.py new file mode 100644 index 00000000..35a01240 --- /dev/null +++ b/video-gen-api/app/models/user_oauth_account.py @@ -0,0 +1,25 @@ +from sqlalchemy import ForeignKey, String +from sqlalchemy.orm import Mapped, mapped_column + +from app.models.base import Base, TimestampMixin, SoftDeleteMixin + + +class UserOAuthAccount(Base, TimestampMixin, SoftDeleteMixin): + __tablename__ = "user_oauth_account" + + id: Mapped[str] = mapped_column( + String(32), primary_key=True, comment="主键" + ) + account_id: Mapped[str] = mapped_column( + String(64), + nullable=False, index=True, comment="授权账户id(user_oauth表中同一个)" + ) + advertiser_id: Mapped[str | None] = mapped_column( + String(64), nullable=True, index=True, comment="广告账户id" + ) + advertiser_name: Mapped[str | None] = mapped_column( + String(128), nullable=True, comment="广告账户名" + ) + advertiser_role: Mapped[str | None] = mapped_column( + String(64), nullable=True, comment="广告账户类型" + ) \ No newline at end of file diff --git a/video-gen-api/app/models/user_oauth_app.py b/video-gen-api/app/models/user_oauth_app.py new file mode 100644 index 00000000..377a6361 --- /dev/null +++ b/video-gen-api/app/models/user_oauth_app.py @@ -0,0 +1,36 @@ +from sqlalchemy import BigInteger, ForeignKey, String +from sqlalchemy.orm import Mapped, mapped_column + +from app.models.base import Base, TimestampMixin, SoftDeleteMixin + + +class UserOAuthApp(Base, TimestampMixin, SoftDeleteMixin): + __tablename__ = "user_oauth_app" + + id: Mapped[str] = mapped_column( + String(32), primary_key=True, comment="主键" + ) + app_id: Mapped[str] = mapped_column( + String(64), unique=True, nullable=False, index=True, comment="应用id" + ) + secret: Mapped[str] = mapped_column( + String(256), nullable=False, comment="应用密钥" + ) + status: Mapped[int] = mapped_column( + BigInteger, nullable=False, default=1, comment="状态,1=正常,2=禁用" + ) + count: Mapped[int] = mapped_column( + BigInteger, nullable=False, default=100, comment="应用最大可以授权多少个用户" + ) + auth_url: Mapped[str] = mapped_column( + String(256), nullable=True, comment="应用授权链接" + ) + company: Mapped[str] = mapped_column( + String(256), nullable=True, comment="应用归属公司名称" + ) + open_type: Mapped[int] = mapped_column( + BigInteger, nullable=False, index=True, comment="开户方式(1=千川,2=广告,3=本地推,4=星图,5=快手代理商,6=巨量星图,7=巨量服务单,8=腾讯服务单,9=腾讯营销K2,10=腾讯营销K3)" + ) + create_by: Mapped[str | None] = mapped_column( + String(32), nullable=True, comment="创建者" + ) \ No newline at end of file diff --git a/video-gen-api/app/schemas/hot_opening_replicate.py b/video-gen-api/app/schemas/hot_opening_replicate.py new file mode 100644 index 00000000..b3be0fd7 --- /dev/null +++ b/video-gen-api/app/schemas/hot_opening_replicate.py @@ -0,0 +1,485 @@ +from __future__ import annotations + +from typing import Any + +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator + +from app.schemas.common import NaiveDatetimeOptional + +HOT_OPENING_PROJECT_STATUS_DESCRIPTIONS: dict[str, str] = { + "pending": "已创建但未进入流程", + "waiting_user": "等待用户手动触发下一步", + "processing": "当前有步骤处理中", + "completed": "总任务完成", + "failed": "总任务失败", + "cancelled": "总任务取消", +} + +HOT_OPENING_STEP_STATUS_DESCRIPTIONS: dict[str, str] = { + "pending": "子任务待处理", + "waiting_user": "等待用户确认或触发", + "processing": "子任务处理中", + "completed": "子任务完成", + "failed": "子任务失败", + "cancelled": "子任务取消", +} + +HOT_OPENING_STEP_DESCRIPTIONS: list[dict[str, Any]] = [ + {"step_index": 1, "step_code": "material_input", "name": "素材输入"}, + {"step_index": 2, "step_code": "image_prompt_optimize", "name": "图片 AI 提词"}, + {"step_index": 3, "step_code": "image_generate", "name": "图片生成"}, + {"step_index": 4, "step_code": "video_prompt_optimize", "name": "视频 AI 提词 JSON schema"}, + {"step_index": 5, "step_code": "video_generate", "name": "视频生成"}, +] + +HOT_OPENING_STEP_IO_SCHEMA_VERSION = "hot_opening_step_io_v1" + +HOT_OPENING_STEP_IO_EXAMPLES: dict[str, dict[str, Any]] = { + "material_input": { + "input_json": { + "schema_version": HOT_OPENING_STEP_IO_SCHEMA_VERSION, + "step_code": "material_input", + "source": {"source_step_id": None, "parent_step_id": None}, + "payload": { + "material_video_url": "https://example.com/source.mp4", + "material_image_url": "https://example.com/product.png", + "source_project_name": "参考素材项目名称", + "target_project_name": "新项目名称", + "core_content_point": "50字以内核心内容点", + }, + "context": {}, + }, + "output_json": { + "schema_version": HOT_OPENING_STEP_IO_SCHEMA_VERSION, + "step_code": "material_input", + "status": "completed", + "payload": {}, + "result": {"accepted": True, "message": "素材输入已提交", "next_step_code": "image_prompt_optimize"}, + "usage": {}, + "error": {}, + }, + }, + "image_prompt_optimize": { + "input_json": { + "schema_version": HOT_OPENING_STEP_IO_SCHEMA_VERSION, + "step_code": "image_prompt_optimize", + "source": {"source_step_id": "第1步素材输入ID", "parent_step_id": "第1步素材输入ID"}, + "payload": {"source_step_id": "第1步素材输入ID"}, + "context": {}, + }, + "output_json": { + "schema_version": HOT_OPENING_STEP_IO_SCHEMA_VERSION, + "step_code": "image_prompt_optimize", + "status": "completed", + "payload": { + "optimized_prompt": "图片生成提示词", + "prompt": "兼容字段,同 optimized_prompt", + "original_prompt": "后端拼接的图片提词原始需求", + "references": [{"type": "video|image", "url": "...", "name": "..."}], + }, + "result": {}, + "usage": { + "input_tokens": 0, + "output_tokens": 0, + "total_tokens": 0, + "text_credits_cost": 0, + "credit_biz_key": "module_generation_step:{step_id}:attempt:1:text_prompt:charge", + }, + "error": {}, + }, + }, + "image_generate": { + "input_json": { + "schema_version": HOT_OPENING_STEP_IO_SCHEMA_VERSION, + "step_code": "image_generate", + "source": {"source_step_id": "第2步图片提词ID", "parent_step_id": "第2步图片提词ID"}, + "payload": { + "engine_id": "图片引擎ID", + "params": {"image_size": "2K", "image_proportion": "1:1", "image_px": "2048x2048"}, + "prompt": "图片生成提示词", + "media_references": [{"type": "image", "url": "新产品图片", "name": "新产品图片"}], + }, + "context": {}, + }, + "output_json": { + "schema_version": HOT_OPENING_STEP_IO_SCHEMA_VERSION, + "step_code": "image_generate", + "status": "completed", + "payload": {}, + "result": {"result_image_url": "/generate/images/xxx.png", "chat_task_id": "ChatGenerationTask ID"}, + "usage": {}, + "error": {}, + }, + }, + "video_prompt_optimize": { + "input_json": { + "schema_version": HOT_OPENING_STEP_IO_SCHEMA_VERSION, + "step_code": "video_prompt_optimize", + "source": {"source_step_id": "第3步图片生成ID", "parent_step_id": "第3步图片生成ID"}, + "payload": { + "source_step_id": "第3步图片生成ID", + "video_config": {"engine_id": "视频引擎ID", "duration": 8, "aspect_ratio": "9:16", "resolution": "1080p"}, + "target_platform": "抖音", + }, + "context": {}, + }, + "output_json": { + "schema_version": HOT_OPENING_STEP_IO_SCHEMA_VERSION, + "step_code": "video_prompt_optimize", + "status": "completed", + "payload": { + "prompt_schema": {"任务基础信息": {}, "最终提示词": {}}, + "final_prompt": "展示用最终视频提示词", + "params_used_for_prompt": {"duration": 8, "aspect_ratio": "9:16", "resolution": "1080p"}, + "target_platform": "抖音", + }, + "result": {}, + "usage": { + "input_tokens": 0, + "output_tokens": 0, + "total_tokens": 0, + "text_credits_cost": 0, + "credit_biz_key": "module_generation_step:{step_id}:attempt:1:text_prompt:charge", + }, + "error": {}, + }, + }, + "video_generate": { + "input_json": { + "schema_version": HOT_OPENING_STEP_IO_SCHEMA_VERSION, + "step_code": "video_generate", + "source": {"source_step_id": "第4步视频提词ID", "parent_step_id": "第4步视频提词ID"}, + "payload": { + "engine_id": "视频引擎ID", + "params": {"duration": 8, "aspect_ratio": "9:16", "resolution": "1080p"}, + "prompt_schema": {"任务基础信息": {}, "最终提示词": {}}, + "final_prompt": "展示用最终提示词", + "media_references": [{"type": "image", "url": "第3步生成图片", "name": "新项目图片"}], + }, + "context": {}, + }, + "output_json": { + "schema_version": HOT_OPENING_STEP_IO_SCHEMA_VERSION, + "step_code": "video_generate", + "status": "completed", + "payload": {}, + "result": { + "result_video_url": "/generate/videos/xxx.mp4", + "result_video_cover_url": "/generate/covers/xxx.jpg", + "chat_task_id": "ChatGenerationTask ID", + }, + "usage": {}, + "error": {}, + }, + }, +} + + +class HotOpeningTaskCreate(BaseModel): + """创建爆款开头复刻总任务项目请求体。""" + + model_config = ConfigDict( + json_schema_extra={ + "example": { + "material_video_url": "https://example.com/source.mp4", + "material_image_url": "https://example.com/product.png", + "source_project_name": "参考素材项目名称", + "target_project_name": "新项目名称", + "core_content_point": "突出产品能帮助用户认识附近新朋友", + "idempotency_key": "frontend-submit-uuid-001", + } + } + ) + + material_video_url: str = Field(..., min_length=1, description="素材视频链接,参考素材,1份。由项目已有上传接口返回,本接口不负责上传,不做后端素材校验") + material_image_url: str = Field(..., min_length=1, description="素材图片链接,新产品图片,1份。由项目已有上传接口返回,本接口不负责上传,不做后端素材校验") + source_project_name: str = Field(..., min_length=1, max_length=20, description="视频素材内容项目名称") + target_project_name: str = Field(..., min_length=1, max_length=20, description="生成项目名称") + core_content_point: str = Field(..., min_length=1, max_length=50, description="生成的项目核心内容点,最多50字") + idempotency_key: str | None = Field(None, max_length=64, description="创建总任务幂等键。只用于 module_generation_projects,不用于 ChatGenerationTask") + + @field_validator("material_video_url", "material_image_url", "source_project_name", "target_project_name", "core_content_point") + @classmethod + def _strip_required(cls, value: str) -> str: + value = str(value or "").strip() + if not value: + raise ValueError("字段不能为空") + return value + + +class HotOpeningMaterialUpdateRequest(BaseModel): + """修改爆款开头复刻第1步素材输入请求体。""" + + model_config = ConfigDict( + json_schema_extra={ + "example": { + "material_video_url": "https://example.com/new-source.mp4", + "material_image_url": "https://example.com/new-product.png", + "source_project_name": "新的参考素材项目名称", + "target_project_name": "新的生成项目名称", + "core_content_point": "新的50字以内核心内容点", + } + } + ) + + material_video_url: str | None = Field(None, min_length=1, description="素材视频链接,未传则沿用旧值") + material_image_url: str | None = Field(None, min_length=1, description="素材图片链接,未传则沿用旧值") + source_project_name: str | None = Field(None, min_length=1, max_length=20, description="视频素材内容项目名称,未传则沿用旧值") + target_project_name: str | None = Field(None, min_length=1, max_length=20, description="生成项目名称,未传则沿用旧值") + core_content_point: str | None = Field(None, min_length=1, max_length=50, description="生成项目核心内容点,最多50字,未传则沿用旧值") + + @field_validator("material_video_url", "material_image_url", "source_project_name", "target_project_name", "core_content_point", mode="before") + @classmethod + def _strip_optional(cls, value: str | None) -> str | None: + if value is None: + return None + value = str(value).strip() + if not value: + raise ValueError("字段不能为空字符串") + return value + + @model_validator(mode="after") + def _require_at_least_one(self) -> "HotOpeningMaterialUpdateRequest": + if not any(getattr(self, field) is not None for field in ("material_video_url", "material_image_url", "source_project_name", "target_project_name", "core_content_point")): + raise ValueError("至少需要传入一个需要修改的字段") + return self + + +class HotOpeningStepUpdate(BaseModel): + """修改爆款开头复刻子任务请求体。""" + + material_video_url: str | None = Field(None, description="修改第1步素材视频链接") + material_image_url: str | None = Field(None, description="修改第1步素材图片链接") + source_project_name: str | None = Field(None, max_length=20, description="修改第1步视频素材内容项目名称") + target_project_name: str | None = Field(None, max_length=20, description="修改第1步生成项目名称") + core_content_point: str | None = Field(None, max_length=50, description="修改第1步生成项目核心内容点,最多50字") + prompt: str | None = Field(None, description="修改第2步图片提词或第4步视频最终提词") + prompt_schema: dict[str, Any] | None = Field(None, description="修改第4步视频提词 JSON schema。只对视频提词步骤有意义") + input_json: dict[str, Any] | None = Field(None, description="高级用法:合并修改当前步骤 input_json.payload") + output_json: dict[str, Any] | None = Field(None, description="高级用法:合并修改当前步骤 output_json.payload") + + +class HotOpeningImagePromptUpdateRequest(BaseModel): + """直接修改第2步图片 AI 优化提词请求体。 + + 本接口不调用 AI、不扣积分;保存后会软删除第3、4、5步当前有效任务。 + """ + + model_config = ConfigDict( + extra="ignore", + json_schema_extra={"example": {"prompt": "用户手动修改后的图片生成提示词"}}, + ) + + prompt: str = Field(..., min_length=1, description="用户手动修改后的图片生成提示词,不能为空") + + @field_validator("prompt", mode="before") + @classmethod + def _strip_prompt(cls, value: str) -> str: + value = str(value or "").strip() + if not value: + raise ValueError("图片提示词不能为空") + return value + + +class HotOpeningVideoPromptSchemaUpdateRequest(BaseModel): + """修改第4步视频 AI 提词 JSON schema 请求体。 + + 前端提交的 prompt_schema 只作为 patch:服务端会锁定视频时长、比例、清晰度、帧率、推荐分辨率、 + 动作/镜头/动态时间规划数组长度和时间段、输出规格限制、质量控制、合规控制、schema_version、schema_usage。 + 最终提示词允许修改,但保存前会清洗视频参数。 + """ + + model_config = ConfigDict( + extra="ignore", + json_schema_extra={ + "example": { + "prompt_schema": { + "业务属性": {"产品名称": "脱单交友APP", "行动引导": "立即下载"}, + "最终提示词": {"主提示词": "脱单交友APP推广短视频,突出认识附近新朋友和高效匹配"}, + } + } + }, + ) + + prompt_schema: dict[str, Any] = Field(..., description="前端修改后的视频提词 JSON schema。后端只按白名单回填允许修改字段") + + @field_validator("prompt_schema") + @classmethod + def _validate_schema(cls, value: dict[str, Any]) -> dict[str, Any]: + if not isinstance(value, dict) or not value: + raise ValueError("prompt_schema 必须是非空 JSON 对象") + return value + + +class HotOpeningGenerateImagePromptRequest(BaseModel): + """手动生成第2步图片 AI 提词请求体。当前无需请求参数。""" + + model_config = ConfigDict(extra="ignore") + + +class HotOpeningGenerateImageRequest(BaseModel): + """根据图片提词生成新项目图片请求体。 + + ChatGenerationTask.idempotency_key 由后端自动生成,接口不再接收前端幂等键。 + """ + + model_config = ConfigDict( + extra="ignore", + json_schema_extra={"example": {"engine_id": "image_engine_xxx", "image_size": "2K", "image_proportion": "1:1", "image_px": "2048x2048"}}, + ) + + engine_id: str | None = Field(None, description="图片生成引擎ID。为空则使用当前启用且优先级最高的图片引擎") + image_size: str | None = Field(None, description="图片分辨率档位,例如 1K、2K。为空使用引擎默认值") + image_proportion: str | None = Field(None, description="图片比例,例如 1:1、16:9、9:16。为空使用默认值") + image_px: str | None = Field(None, description="图片像素尺寸,例如 2048x2048。为空时按引擎支持尺寸自动匹配") + + +class HotOpeningGenerateVideoPromptRequest(BaseModel): + """手动生成第4步视频 AI 提词请求体。 + + 视频时长、比例、分辨率集中在本步骤确定;第5步生成视频只选择视频引擎。 + """ + + model_config = ConfigDict( + extra="ignore", + json_schema_extra={"example": {"engine_id": "video_engine_xxx", "duration": 8, "aspect_ratio": "9:16", "resolution": "1080p", "target_platform": "抖音"}}, + ) + + engine_id: str | None = Field(None, description="视频引擎ID。用于读取该引擎支持的视频时长、比例、分辨率配置;为空使用最高优先级启用引擎") + duration: int | None = Field(None, ge=1, description="希望用于视频提词规划的视频时长,单位秒。为空时优先使用 HOT_OPENING_DEFAULT_VIDEO_DURATION") + aspect_ratio: str | None = Field(None, description="希望用于视频提词规划的视频比例。为空时优先使用 HOT_OPENING_DEFAULT_VIDEO_RATIO") + resolution: str | None = Field(None, description="希望用于视频提词规划的视频分辨率。为空时优先使用 HOT_OPENING_DEFAULT_VIDEO_RESOLUTION") + target_platform: str | None = Field(None, max_length=64, description="目标平台,例如抖音/快手/小红书。为空时使用 HOT_OPENING_DEFAULT_TARGET_PLATFORM") + + +class HotOpeningGenerateVideoRequest(BaseModel): + """根据视频提词生成最终视频请求体。 + + 只选择视频生成引擎。duration / aspect_ratio / resolution 从第4步视频提词优化结果读取。 + ChatGenerationTask.original_prompt / optimized_prompt 都写入第4步生成的 prompt_schema JSON 字符串。 + """ + + model_config = ConfigDict(extra="ignore", json_schema_extra={"example": {"engine_id": "video_engine_xxx"}}) + + engine_id: str | None = Field(None, description="视频生成引擎ID。为空优先使用第4步视频提词时选择的 engine_id,再为空使用最高优先级启用视频引擎") + + +class HotOpeningStepOut(BaseModel): + id: str = Field(..., description="子任务ID") + project_id: str = Field(..., description="总任务项目ID,即 module_generation_projects.id") + module: str = Field(..., description="模块标识,例如 hot_opening_replicate") + step_index: int = Field(..., description="步骤序号:1素材输入、2图片提词、3图片生成、4视频提词、5视频生成") + step_code: str = Field(..., description="步骤编码:material_input/image_prompt_optimize/image_generate/video_prompt_optimize/video_generate") + status: str = Field(..., description="步骤状态:pending/waiting_user/processing/completed/failed/cancelled") + version: int = Field(..., description="步骤版本号。重新生成或修改上游步骤后 version+1") + is_current: bool = Field(..., description="是否当前有效步骤。旧步骤会软删除且 is_current=false") + parent_step_id: str | None = Field(None, description="上一个步骤ID") + source_step_id: str | None = Field(None, description="当前步骤基于哪个上游步骤生成") + chat_task_id: str | None = Field(None, description="关联的 ChatGenerationTask ID。第3步图片生成、第5步视频生成有值") + input: dict[str, Any] | None = Field(None, description=f"步骤输入 JSON,统一 schema_version={HOT_OPENING_STEP_IO_SCHEMA_VERSION}") + output: dict[str, Any] | None = Field(None, description=f"步骤输出 JSON,统一 schema_version={HOT_OPENING_STEP_IO_SCHEMA_VERSION}") + error_message: str | None = Field(None, description="步骤错误信息") + created_at: NaiveDatetimeOptional = Field(None, description="创建时间") + updated_at: NaiveDatetimeOptional = Field(None, description="更新时间") + completed_at: NaiveDatetimeOptional = Field(None, description="完成时间") + + +class HotOpeningMaterialOut(BaseModel): + material_step_id: str | None = Field(None, description="第1步素材输入子任务ID") + material_video_url: str | None = Field(None, description="素材视频链接") + material_image_url: str | None = Field(None, description="素材图片链接") + source_project_name: str | None = Field(None, description="视频素材内容项目名称") + target_project_name: str | None = Field(None, description="生成项目名称") + core_content_point: str | None = Field(None, description="生成项目核心内容点") + + +class HotOpeningImageGenerationOut(BaseModel): + prompt_step_id: str | None = Field(None, description="第2步图片 AI 提词子任务ID") + generate_step_id: str | None = Field(None, description="第3步图片生成子任务ID") + prompt: str | None = Field(None, description="图片优化提词") + engine_id: str | None = Field(None, description="图片生成引擎ID") + engine_name: str | None = Field(None, description="图片生成引擎名称") + params: dict[str, Any] | None = Field(None, description="图片生成参数") + chat_task_id: str | None = Field(None, description="图片生成 ChatGenerationTask ID") + status: str | None = Field(None, description="图片生成状态") + result_image_url: str | None = Field(None, description="新项目图片 URL") + error_message: str | None = Field(None, description="图片生成错误信息") + + +class HotOpeningVideoGenerationOut(BaseModel): + prompt_step_id: str | None = Field(None, description="第4步视频 AI 提词子任务ID") + generate_step_id: str | None = Field(None, description="第5步视频生成子任务ID") + prompt_schema: dict[str, Any] | None = Field(None, description="视频提词 JSON schema。第5步 ChatGenerationTask 原始提词会使用该 JSON 字符串") + final_prompt: str | None = Field(None, description="视频最终提词,仅用于前端展示") + prompt_params: dict[str, Any] | None = Field(None, description="第4步生成视频提词时使用的视频配置,例如 duration、aspect_ratio、resolution") + engine_id: str | None = Field(None, description="视频生成引擎ID") + engine_name: str | None = Field(None, description="视频生成引擎名称") + params: dict[str, Any] | None = Field(None, description="视频生成实际参数。第5步只传 engine_id,其它参数继承第4步") + chat_task_id: str | None = Field(None, description="视频生成 ChatGenerationTask ID") + status: str | None = Field(None, description="视频生成状态") + result_video_url: str | None = Field(None, description="最终视频 URL") + result_video_cover_url: str | None = Field(None, description="最终视频封面 URL") + error_message: str | None = Field(None, description="视频生成错误信息") + + +class HotOpeningTaskDetailOut(BaseModel): + id: str = Field(..., description="总任务项目ID。这个ID就是前端项目ID") + project_id: str = Field(..., description="兼容前端命名,等同于 id") + module: str = Field(..., description="模块标识,爆款开头复刻固定为 hot_opening_replicate") + title: str | None = Field(None, description="项目标题,默认取生成项目名称") + status: str = Field(..., description="总任务状态:pending/waiting_user/processing/completed/failed/cancelled") + current_step_code: str | None = Field(None, description="当前所处步骤编码") + final_image_url: str | None = Field(None, description="最终新项目图片 URL") + final_video_url: str | None = Field(None, description="最终视频 URL") + final_video_cover_url: str | None = Field(None, description="最终视频封面 URL") + error_message: str | None = Field(None, description="总任务错误信息") + material: HotOpeningMaterialOut = Field(default_factory=HotOpeningMaterialOut, description="素材和项目描述信息") + image_generation: HotOpeningImageGenerationOut = Field(default_factory=HotOpeningImageGenerationOut, description="图片提词、图片引擎参数和图片结果") + video_generation: HotOpeningVideoGenerationOut = Field(default_factory=HotOpeningVideoGenerationOut, description="视频提词、视频引擎参数和视频结果") + steps: list[HotOpeningStepOut] = Field(default_factory=list, description="当前有效子任务列表") + created_at: NaiveDatetimeOptional = Field(None, description="创建时间") + updated_at: NaiveDatetimeOptional = Field(None, description="更新时间") + completed_at: NaiveDatetimeOptional = Field(None, description="完成时间") + + +class HotOpeningTaskListItemOut(BaseModel): + id: str = Field(..., description="总任务项目ID。这个ID就是前端项目ID") + project_id: str = Field(..., description="兼容前端命名,等同于 id") + module: str = Field(..., description="模块标识") + title: str | None = Field(None, description="项目标题") + status: str = Field(..., description="总任务状态") + current_step_code: str | None = Field(None, description="当前步骤") + target_project_name: str | None = Field(None, description="生成项目名称,来源于第1步素材输入") + final_image_url: str | None = Field(None, description="最终图片 URL") + final_video_url: str | None = Field(None, description="最终视频 URL") + error_message: str | None = Field(None, description="错误信息") + created_at: NaiveDatetimeOptional = Field(None, description="创建时间") + updated_at: NaiveDatetimeOptional = Field(None, description="更新时间") + completed_at: NaiveDatetimeOptional = Field(None, description="完成时间") + + +class HotOpeningTaskListOut(BaseModel): + total: int = Field(..., description="总数量") + items: list[HotOpeningTaskListItemOut] = Field(default_factory=list, description="列表数据") + + +class HotOpeningActionOut(BaseModel): + message: str = Field(..., description="操作结果提示") + project_id: str = Field(..., description="总任务项目ID") + step_id: str | None = Field(None, description="本次创建或修改的子任务ID") + next_step_id: str | None = Field(None, description="兼容字段:当前接口不自动生成下下个任务,一般为空") + detail: HotOpeningTaskDetailOut | None = Field(None, description="操作后的总任务详情") + + +class HotOpeningDeleteOut(BaseModel): + message: str = Field(..., description="删除结果提示") + project_id: str = Field(..., description="被软删除的总任务项目ID") + deleted: bool = Field(..., description="是否已软删除") + + +class HotOpeningSpecOut(BaseModel): + project_statuses: dict[str, str] = Field(default_factory=lambda: HOT_OPENING_PROJECT_STATUS_DESCRIPTIONS, description="总任务状态说明") + step_statuses: dict[str, str] = Field(default_factory=lambda: HOT_OPENING_STEP_STATUS_DESCRIPTIONS, description="子任务状态说明") + steps: list[dict[str, Any]] = Field(default_factory=lambda: HOT_OPENING_STEP_DESCRIPTIONS, description="5个固定步骤说明") + step_io_schema_version: str = Field(default=HOT_OPENING_STEP_IO_SCHEMA_VERSION, description="步骤 input_json/output_json 结构版本") + step_io_examples: dict[str, dict[str, Any]] = Field(default_factory=lambda: HOT_OPENING_STEP_IO_EXAMPLES, description="每个步骤 input_json/output_json 示例") diff --git a/video-gen-api/app/schemas/payment.py b/video-gen-api/app/schemas/payment.py index cf611557..135b4172 100644 --- a/video-gen-api/app/schemas/payment.py +++ b/video-gen-api/app/schemas/payment.py @@ -3,6 +3,7 @@ from pydantic import BaseModel class RechargeRequest(BaseModel): plan: str # package id + method: str = "wechat" # "wechat" or "alipay" class PaymentOrderOut(BaseModel): @@ -12,5 +13,6 @@ class PaymentOrderOut(BaseModel): credits: float payment_method: str status: str + qr_url: str | None = None # Alipay QR code URL (transient, not persisted) model_config = {"from_attributes": True} diff --git a/video-gen-api/app/schemas/user_oauth.py b/video-gen-api/app/schemas/user_oauth.py new file mode 100644 index 00000000..d43d6ff6 --- /dev/null +++ b/video-gen-api/app/schemas/user_oauth.py @@ -0,0 +1,29 @@ +from datetime import datetime + +from pydantic import BaseModel, Field + + +class RequestOAuthRequest(BaseModel): + oauth_type: int = Field( + ..., + description="开户方式(1=巨量广告,2=巨量千川,3=快手,4=腾讯)", + ) + + +class RequestOAuthResponse(BaseModel): + auth_url: str = Field(..., description="第三方授权链接") + + +class UserOAuthOut(BaseModel): + id: str = Field(..., description="主键") + account_id: str = Field(..., description="授权账户id") + account_name: str = Field(..., description="授权账户name") + account_role: str | None = Field(None, description="授权账户角色") + account_username: str | None = Field(None, description="授权账户登录账号") + user_id: str = Field(..., description="用户id") + open_type: int = Field(..., description="开户方式") + port_type: int = Field(..., description="平台端口") + appid: str | None = Field(None, description="授权应用id") + material_auth_status: bool = Field(False, description="是否敏感物料授权") + created_at: datetime = Field(..., description="创建时间") + updated_at: datetime = Field(..., description="更新时间") \ No newline at end of file diff --git a/video-gen-api/app/schemas/user_oauth_app.py b/video-gen-api/app/schemas/user_oauth_app.py new file mode 100644 index 00000000..3c574e46 --- /dev/null +++ b/video-gen-api/app/schemas/user_oauth_app.py @@ -0,0 +1,47 @@ +from pydantic import BaseModel, Field + +from app.schemas.common import NaiveDatetime + + +class UserOAuthAppCreate(BaseModel): + app_id: str = Field(..., max_length=64, description="应用id") + secret: str = Field(..., max_length=256, description="应用密钥") + open_type: int = Field( + ..., + ge=1, + le=10, + description="开户方式(1=千川,2=广告,3=本地推,4=星图,5=快手代理商,6=巨量星图,7=巨量服务单,8=腾讯服务单,9=腾讯营销K2,10=腾讯营销K3)", + ) + count: int = Field(100, ge=1, description="应用最大可以授权多少个用户") + auth_url: str | None = Field(None, max_length=256, description="应用授权链接") + company: str | None = Field(None, max_length=256, description="应用归属公司名称") + + +class UserOAuthAppUpdate(BaseModel): + secret: str | None = Field(None, max_length=256, description="应用密钥") + open_type: int | None = Field( + None, + ge=1, + le=10, + description="开户方式(1=千川,2=广告,3=本地推,4=星图,5=快手代理商,6=巨量星图,7=巨量服务单,8=腾讯服务单,9=腾讯营销K2,10=腾讯营销K3)", + ) + status: int | None = Field(None, ge=1, le=2, description="应用状态(1=正常,2=禁用)") + count: int | None = Field(None, ge=1, description="应用最大可以授权多少个用户") + auth_url: str | None = Field(None, max_length=256, description="应用授权链接") + company: str | None = Field(None, max_length=256, description="应用归属公司名称") + + +class UserOAuthAppOut(BaseModel): + id: str = Field(..., description="主键") + app_id: str = Field(..., description="应用id") + secret: str = Field(..., description="应用密钥") + status: int = Field(..., description="状态,1=正常,2=禁用") + count: int = Field(..., description="应用最大可以授权多少个用户") + open_type: int = Field(..., description="开户方式") + auth_url: str | None = Field(None, description="应用授权链接") + company: str | None = Field(None, description="应用归属公司名称") + create_by: str | None = Field(None, description="创建者") + created_at: NaiveDatetime = Field(..., description="创建时间") + updated_at: NaiveDatetime = Field(..., description="更新时间") + + model_config = {"from_attributes": True} \ No newline at end of file diff --git a/video-gen-api/app/services/generation_billing_service.py b/video-gen-api/app/services/generation_billing_service.py index 74d1087c..51604d16 100644 --- a/video-gen-api/app/services/generation_billing_service.py +++ b/video-gen-api/app/services/generation_billing_service.py @@ -20,6 +20,7 @@ CHARGE_MEDIA = "media" OWNER_GENERATION_RECORD = "generation_record" OWNER_CHAT_GENERATION_TASK = "chat_generation_task" +OWNER_MODULE_GENERATION_STEP = "module_generation_step" _BIZ_KEY_PATTERN = re.compile( r"^(?P[^:]+):(?P[^:]+):attempt:(?P\d+):(?P[^:]+):(?Pcharge|refund)$" @@ -287,6 +288,45 @@ async def charge_chatapi_prompt_usage( return BillingSummary(record_id=record.id, user_id=record.user_id, items=items) +async def charge_module_prompt_usage( + db: AsyncSession, + *, + user_id: str, + step_id: str, + usage: Mapping[str, Any], + description: str, + attempt_no: int = 1, +) -> BillingSummary: + """爆款开头复刻模块图片/视频 AI 提词扣文本积分。 + + 文本提词属于已经发生的 LLM 消费: + - 调用成功后按 input_tokens + output_tokens 扣费。 + - 不参与后续图片/视频媒体生成失败退款。 + - 通过 module_generation_step:{step_id}:attempt:1:text_prompt:charge 幂等。 + """ + input_tokens = _safe_int(usage.get("input_tokens")) + output_tokens = _safe_int(usage.get("output_tokens")) + text_credits = await calc_text_credits(db, input_tokens, output_tokens) + biz_key = build_credit_biz_key( + owner_type=OWNER_MODULE_GENERATION_STEP, + owner_id=step_id, + attempt_no=attempt_no, + charge_kind=CHARGE_TEXT_PROMPT, + action="charge", + ) + item = await deduct_credits_locked_once( + db, + user_id=user_id, + amount=text_credits, + description=description, + related_id=step_id, + charge_key=CHARGE_TEXT_PROMPT, + biz_key=biz_key, + attempt_no=attempt_no, + ) + return BillingSummary(record_id=step_id, user_id=user_id, items=[item]) + + async def charge_generation_media_by_params( db: AsyncSession, *, diff --git a/video-gen-api/app/services/generation_module_hook_service.py b/video-gen-api/app/services/generation_module_hook_service.py new file mode 100644 index 00000000..ebf3f03c --- /dev/null +++ b/video-gen-api/app/services/generation_module_hook_service.py @@ -0,0 +1,25 @@ +from __future__ import annotations + +from sqlalchemy.ext.asyncio import AsyncSession + +from app.models.chat_generation_task import ChatGenerationTask + + +async def notify_chat_generation_task_finished(db: AsyncSession, task: ChatGenerationTask) -> None: + """通知业务模块 ChatGenerationTask 已进入终态。 + + 当前用于爆款开头复刻: + - image_generate 完成后自动进入 video_prompt_optimize + - video_generate 完成后总任务完成 + """ + if not task: + return + if task.generation_mode == "hot_opening_replicate": + from app.services.hot_opening_replicate_service import ( + handle_chat_generation_task_completed, + handle_chat_generation_task_failed, + ) + if task.status == "completed": + await handle_chat_generation_task_completed(db, task) + elif task.status == "failed": + await handle_chat_generation_task_failed(db, task) diff --git a/video-gen-api/app/services/generation_recovery_service.py b/video-gen-api/app/services/generation_recovery_service.py index cb54a932..33ffd7a3 100644 --- a/video-gen-api/app/services/generation_recovery_service.py +++ b/video-gen-api/app/services/generation_recovery_service.py @@ -17,10 +17,13 @@ from app.services.celery_download_recovery_service import ( remove_download_active, ) from app.services.generation_log_service import log_task_event +from app.services.generation_module_hook_service import notify_chat_generation_task_finished from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once logger = logging.getLogger("video_gen") +ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate"} + def _now() -> datetime: return datetime.now(timezone.utc) @@ -73,7 +76,7 @@ async def recover_one_download_task( if not task: return "skip_missing_task" - if task.generation_mode != "chatapi_async": + if task.generation_mode not in ALLOWED_GENERATION_MODES: await remove_download_active(task.id) return "clean_invalid_mode" if _is_final_task_state(task): @@ -218,7 +221,7 @@ async def recover_download_tasks_once(db: AsyncSession) -> dict[str, Any]: select(ChatGenerationTask) .where( ChatGenerationTask.deleted_at.is_(None), - ChatGenerationTask.generation_mode == "chatapi_async", + ChatGenerationTask.generation_mode.in_(["chatapi_async", "hot_opening_replicate"]), ChatGenerationTask.status == "generating", ChatGenerationTask.remote_result_url.is_not(None), ChatGenerationTask.pipeline_stage.in_( @@ -263,7 +266,7 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]: select(ChatGenerationTask) .where( ChatGenerationTask.deleted_at.is_(None), - ChatGenerationTask.generation_mode == "chatapi_async", + ChatGenerationTask.generation_mode.in_(["chatapi_async", "hot_opening_replicate"]), ChatGenerationTask.status == "generating", ChatGenerationTask.pipeline_stage.in_( [ @@ -290,6 +293,7 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]: error_message="任务超时", pipeline_stage="timeout", ) + await notify_chat_generation_task_finished(db, task) await db.commit() await log_task_event( task, diff --git a/video-gen-api/app/services/generation_refund_service.py b/video-gen-api/app/services/generation_refund_service.py index 1ebb9526..0b13b644 100644 --- a/video-gen-api/app/services/generation_refund_service.py +++ b/video-gen-api/app/services/generation_refund_service.py @@ -180,7 +180,7 @@ async def mark_chat_generation_task_failed_and_refund_once( select(ChatGenerationTask) .where( ChatGenerationTask.id == task_id, - ChatGenerationTask.generation_mode == "chatapi_async", + ChatGenerationTask.generation_mode.in_(["chatapi_async", "hot_opening_replicate"]), ChatGenerationTask.deleted_at.is_(None), ) .with_for_update() diff --git a/video-gen-api/app/services/generation_task_factory_service.py b/video-gen-api/app/services/generation_task_factory_service.py new file mode 100644 index 00000000..21381157 --- /dev/null +++ b/video-gen-api/app/services/generation_task_factory_service.py @@ -0,0 +1,185 @@ +from __future__ import annotations + +import json +from datetime import datetime, timedelta, timezone +from typing import Any + +from fastapi import HTTPException +from sqlalchemy.ext.asyncio import AsyncSession + +from app.config import settings +from app.models.chat_generation_task import ChatGenerationTask +from app.models.user import User +from app.schemas.generation_ai import GenerationAIReference, GenerationAITaskCreate +from app.services.generation_ai_service import ( + IMAGE_DEFAULT_PROPORTION, + IMAGE_DEFAULT_PX, + IMAGE_DEFAULT_SIZE, + VIDEO_DEFAULT_RATIO, + VIDEO_DEFAULT_RESOLUTION, + _build_image_snapshot, + _build_video_snapshot, + _get_image_engine, + _get_video_engine, + _image_supported_sizes, + _parse_list, + normalize_px, +) +from app.services.generation_billing_service import OWNER_CHAT_GENERATION_TASK, charge_generation_media_by_params +from app.utils.id_gen import generate_id + + +def _json(data: Any) -> str | None: + if data is None: + return None + return json.dumps(data, ensure_ascii=False, default=str) + + +def _build_backend_idempotency_key(*, generation_mode: str, gen_type: str, task_id: str) -> str: + """模块生成关联 ChatGenerationTask 的幂等键由后端生成。 + + 不再接收前端透传,避免 user_id + generation_mode + idempotency_key + 唯一索引被前端固定 key 或重复 key 拦截。 + """ + return f"{generation_mode}:{gen_type}:{task_id}"[:64] + + +async def create_chat_generation_task_for_module( + db: AsyncSession, + *, + current_user: User, + generation_mode: str, + gen_type: str, + original_prompt: str, + optimized_prompt: str | None = None, + engine_id: str | None = None, + media_references: list[dict[str, Any]] | None = None, + idempotency_key: str | None = None, + image_size: str | None = None, + image_proportion: str | None = None, + image_px: str | None = None, + duration: int | None = None, + aspect_ratio: str | None = None, + resolution: str | None = None, + billing_project_name: str = "模块生成任务", + billing_description_prefix: str = "模块生成-", +) -> ChatGenerationTask: + """创建可复用的 ChatGenerationTask 子任务。 + + 和 /generation-ai 普通任务不同,generation_mode 由业务模块传入, + 但仍复用同一套引擎校验、扣费、Celery 创建/轮询/下载逻辑。 + """ + gen_type = gen_type.lower().strip() + if gen_type not in ("image", "video"): + raise HTTPException(status_code=400, detail="gen_type 仅支持 image 或 video") + + task_id = generate_id() + now = datetime.now(timezone.utc) + refs = media_references or [] + backend_idempotency_key = _build_backend_idempotency_key( + generation_mode=generation_mode, + gen_type=gen_type, + task_id=task_id, + ) + + if gen_type == "image": + engine = await _get_image_engine(db, engine_id) + sizes = _image_supported_sizes(engine) + size = image_size or engine.default_size or IMAGE_DEFAULT_SIZE + proportion = image_proportion or IMAGE_DEFAULT_PROPORTION + px = normalize_px(image_px) + if sizes: + if size not in sizes: + raise HTTPException(status_code=400, detail=f"图片分辨率档位不支持: {size}") + if proportion not in sizes.get(size, {}): + raise HTTPException(status_code=400, detail=f"图片比例不支持: {proportion}") + px = px or normalize_px((sizes.get(size) or {}).get(proportion)) + px = px or IMAGE_DEFAULT_PX + media_billing = await charge_generation_media_by_params( + db, + user_id=current_user.id, + record_id=task_id, + gen_type="image", + image_size=size, + engine_id=engine.id, + project_name=billing_project_name, + description_prefix=billing_description_prefix, + owner_type=OWNER_CHAT_GENERATION_TASK, + attempt_no=1, + ) + snapshot = _build_image_snapshot(engine, size, proportion, px) + task = ChatGenerationTask( + id=task_id, + user_id=current_user.id, + original_prompt=original_prompt, + optimized_prompt=optimized_prompt, + gen_type="image", + image_size=size, + image_proportion=proportion, + image_px=px, + status="generating", + generation_mode=generation_mode, + pipeline_stage="queued", + engine_id=engine.id, + engine_snapshot_json=_json(snapshot), + media_references=_json(refs) if refs else None, + credits_cost=round(media_billing.total_charged, 2), + idempotency_key=backend_idempotency_key, + deadline_at=now + timedelta(minutes=settings.CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES), + ) + else: + engine = await _get_video_engine(db, engine_id) + ratio = aspect_ratio or VIDEO_DEFAULT_RATIO + selected_resolution = resolution or VIDEO_DEFAULT_RESOLUTION + selected_duration = duration or 4 + ratios = _parse_list(engine.supported_ratios, []) + resolutions = _parse_list(engine.supported_resolutions, []) + durations = _parse_list(engine.supported_durations, []) + if ratios and ratio not in ratios: + raise HTTPException(status_code=400, detail=f"视频比例不支持: {ratio}") + if resolutions and selected_resolution not in resolutions: + raise HTTPException(status_code=400, detail=f"视频分辨率不支持: {selected_resolution}") + if durations and selected_duration not in durations: + raise HTTPException(status_code=400, detail=f"视频时长不支持: {selected_duration}") + if engine.max_duration and selected_duration > engine.max_duration: + raise HTTPException(status_code=400, detail=f"视频时长不能超过 {engine.max_duration} 秒") + media_billing = await charge_generation_media_by_params( + db, + user_id=current_user.id, + record_id=task_id, + gen_type="video", + duration=selected_duration, + resolution=selected_resolution, + engine_id=engine.id, + project_name=billing_project_name, + description_prefix=billing_description_prefix, + owner_type=OWNER_CHAT_GENERATION_TASK, + attempt_no=1, + ) + snapshot = _build_video_snapshot(engine, ratio, selected_resolution, selected_duration) + task = ChatGenerationTask( + id=task_id, + user_id=current_user.id, + original_prompt=original_prompt, + optimized_prompt=optimized_prompt, + gen_type="video", + duration=selected_duration, + aspect_ratio=ratio, + resolution=selected_resolution, + image_size=image_size or IMAGE_DEFAULT_SIZE, + image_proportion=image_proportion or IMAGE_DEFAULT_PROPORTION, + image_px=normalize_px(image_px) or IMAGE_DEFAULT_PX, + status="generating", + generation_mode=generation_mode, + pipeline_stage="queued", + engine_id=engine.id, + engine_snapshot_json=_json(snapshot), + media_references=_json(refs) if refs else None, + credits_cost=round(media_billing.total_charged, 2), + idempotency_key=backend_idempotency_key, + deadline_at=now + timedelta(minutes=settings.CHATAPI_ASYNC_VIDEO_DEADLINE_MINUTES), + ) + + db.add(task) + await db.flush() + return task diff --git a/video-gen-api/app/services/hot_opening_replicate_service.py b/video-gen-api/app/services/hot_opening_replicate_service.py new file mode 100644 index 00000000..040dea7b --- /dev/null +++ b/video-gen-api/app/services/hot_opening_replicate_service.py @@ -0,0 +1,1546 @@ +from __future__ import annotations + +import json +from datetime import datetime, timezone +from typing import Any + +from fastapi import HTTPException +from sqlalchemy import func, select +from sqlalchemy.ext.asyncio import AsyncSession + +from app.config import settings +from app.enums.common import ModuleEventTypeEnum, ModuleProjectStatusEnum, ModulePromptTypeEnum, ModuleStepStatusEnum +from app.enums.hot_opening_replicate import HotOpeningGenerationModeEnum, HotOpeningStepCodeEnum, ModuleCodeEnum +from app.models.chat_generation_task import ChatGenerationTask +from app.models.module_generation_project import ModuleGenerationProject +from app.models.module_generation_step import ModuleGenerationStep +from app.models.user import User +from app.schemas.hot_opening_replicate import ( + HotOpeningDeleteOut, + HotOpeningGenerateImageRequest, + HotOpeningGenerateVideoPromptRequest, + HotOpeningGenerateVideoRequest, + HotOpeningImageGenerationOut, + HotOpeningImagePromptUpdateRequest, + HotOpeningMaterialOut, + HotOpeningMaterialUpdateRequest, + HotOpeningStepOut, + HotOpeningStepUpdate, + HotOpeningTaskCreate, + HotOpeningTaskDetailOut, + HotOpeningTaskListItemOut, + HotOpeningTaskListOut, + HotOpeningVideoGenerationOut, + HotOpeningVideoPromptSchemaUpdateRequest, +) +from app.services.generation_ai_service import ( + VIDEO_DEFAULT_DURATION, + VIDEO_DEFAULT_RATIO, + VIDEO_DEFAULT_RESOLUTION, + _get_video_engine, + _parse_list, +) +from app.services.generation_billing_service import charge_module_prompt_usage +from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once +from app.services.generation_task_factory_service import create_chat_generation_task_for_module +from app.services.hot_opening_video_prompt_service import build_final_video_prompt, optimize_hot_opening_video_prompt, patch_video_prompt_schema_from_client +from app.services.module_generation_log_service import log_module_event_file, log_module_prompt_event +from app.services.llm import optimize_prompt +from app.services.resource_accounting_service import soft_delete_chat_task_resources +from app.services.resource_signed_url_service import build_resource_signed_url +from app.utils.id_gen import generate_id + +MODULE = ModuleCodeEnum.HOT_OPENING_REPLICATE.value +GENERATION_MODE = HotOpeningGenerationModeEnum.HOT_OPENING_REPLICATE.value + +STEP_INDEX_MAP = { + HotOpeningStepCodeEnum.MATERIAL_INPUT.value: 1, + HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value: 2, + HotOpeningStepCodeEnum.IMAGE_GENERATE.value: 3, + HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value: 4, + HotOpeningStepCodeEnum.VIDEO_GENERATE.value: 5, +} + +STEP_IO_SCHEMA_VERSION = "hot_opening_step_io_v1" + + +def _now() -> datetime: + return datetime.now(timezone.utc) + + +def _json(data: Any) -> str | None: + if data is None: + return None + return json.dumps(data, ensure_ascii=False, default=str) + + +def _parse_json(value: Any, fallback: Any = None) -> Any: + if value is None or value == "": + return fallback + if isinstance(value, (dict, list)): + return value + if isinstance(value, str): + try: + return json.loads(value) + except Exception: + return fallback + return fallback + + +def _step_input( + *, + step_code: str, + payload: dict[str, Any] | None = None, + source_step_id: str | None = None, + parent_step_id: str | None = None, + context: dict[str, Any] | None = None, +) -> dict[str, Any]: + return { + "schema_version": STEP_IO_SCHEMA_VERSION, + "step_code": step_code, + "source": { + "source_step_id": source_step_id, + "parent_step_id": parent_step_id, + }, + "payload": payload or {}, + "context": context or {}, + } + + +def _step_output( + *, + step_code: str, + status: str, + payload: dict[str, Any] | None = None, + result: dict[str, Any] | None = None, + usage: dict[str, Any] | None = None, + error: dict[str, Any] | None = None, +) -> dict[str, Any]: + return { + "schema_version": STEP_IO_SCHEMA_VERSION, + "step_code": step_code, + "status": status, + "payload": payload or {}, + "result": result or {}, + "usage": usage or {}, + "error": error or {}, + } + + +def _is_wrapped_step_io(value: Any) -> bool: + return isinstance(value, dict) and value.get("schema_version") == STEP_IO_SCHEMA_VERSION + + +def _step_payload(value: Any) -> dict[str, Any]: + data = _parse_json(value, {}) or {} + if _is_wrapped_step_io(data): + payload = data.get("payload") + return payload if isinstance(payload, dict) else {} + return data if isinstance(data, dict) else {} + + +def _step_result(value: Any) -> dict[str, Any]: + data = _parse_json(value, {}) or {} + if _is_wrapped_step_io(data): + result = data.get("result") + if isinstance(result, dict) and result: + return result + payload = data.get("payload") + return payload if isinstance(payload, dict) else {} + return data if isinstance(data, dict) else {} + + +def _step_usage(value: Any) -> dict[str, Any]: + data = _parse_json(value, {}) or {} + if _is_wrapped_step_io(data): + usage = data.get("usage") + return usage if isinstance(usage, dict) else {} + usage = data.get("token_usage") if isinstance(data, dict) else {} + return usage if isinstance(usage, dict) else {} + + +def _unwrap_step_output(value: Any) -> dict[str, Any]: + data = _parse_json(value, {}) or {} + if not _is_wrapped_step_io(data): + return data if isinstance(data, dict) else {} + merged: dict[str, Any] = {} + payload = data.get("payload") + result = data.get("result") + usage = data.get("usage") + if isinstance(payload, dict): + merged.update(payload) + if isinstance(result, dict): + merged.update(result) + if isinstance(usage, dict) and usage: + merged["token_usage"] = usage + return merged + + +def _merge_dict(old: dict[str, Any] | None, new: dict[str, Any] | None) -> dict[str, Any]: + merged = dict(old or {}) + for key, value in (new or {}).items(): + if value is not None: + merged[key] = value + return merged + + +async def log_module_event( + db: AsyncSession, + *, + project: ModuleGenerationProject, + event_type: str, + step: ModuleGenerationStep | None = None, + message: str | None = None, + detail: dict[str, Any] | None = None, +) -> None: + """模块事件日志只落盘,不再写 module_generation_events 表。""" + _ = db + log_module_event_file( + module=project.module, + event_type=event_type, + project_id=project.id, + step_id=step.id if step else None, + user_id=project.user_id, + message=message, + detail=detail, + ) + + +async def _get_project_for_user( + db: AsyncSession, + *, + project_id: str, + user: User, + for_update: bool = False, + populate_existing: bool = False, +) -> ModuleGenerationProject: + query = select(ModuleGenerationProject).where( + ModuleGenerationProject.id == project_id, + ModuleGenerationProject.module == MODULE, + ModuleGenerationProject.deleted_at.is_(None), + ) + if not user.is_admin: + query = query.where(ModuleGenerationProject.user_id == user.id) + if populate_existing: + query = query.execution_options(populate_existing=True) + if for_update: + query = query.with_for_update() + result = await db.execute(query.limit(1)) + project = result.scalar_one_or_none() + if not project: + raise HTTPException(status_code=404, detail="爆款开头复刻项目不存在") + return project + + +async def _get_step_for_user( + db: AsyncSession, + *, + project_id: str, + step_id: str, + user: User, + for_update: bool = False, +) -> ModuleGenerationStep: + await _get_project_for_user(db, project_id=project_id, user=user, for_update=for_update) + query = select(ModuleGenerationStep).where( + ModuleGenerationStep.id == step_id, + ModuleGenerationStep.project_id == project_id, + ModuleGenerationStep.module == MODULE, + ModuleGenerationStep.deleted_at.is_(None), + ModuleGenerationStep.is_current == True, + ) + if not user.is_admin: + query = query.where(ModuleGenerationStep.user_id == user.id) + if for_update: + query = query.with_for_update() + result = await db.execute(query.limit(1)) + step = result.scalar_one_or_none() + if not step: + raise HTTPException(status_code=404, detail="子任务不存在") + return step + + +async def _get_current_steps(db: AsyncSession, project_id: str) -> list[ModuleGenerationStep]: + result = await db.execute( + select(ModuleGenerationStep) + .where( + ModuleGenerationStep.project_id == project_id, + ModuleGenerationStep.module == MODULE, + ModuleGenerationStep.deleted_at.is_(None), + ModuleGenerationStep.is_current == True, + ) + .order_by(ModuleGenerationStep.step_index.asc(), ModuleGenerationStep.created_at.asc()) + ) + return list(result.scalars().all()) + + +async def _get_current_step_by_code(db: AsyncSession, project_id: str, step_code: str) -> ModuleGenerationStep | None: + result = await db.execute( + select(ModuleGenerationStep) + .where( + ModuleGenerationStep.project_id == project_id, + ModuleGenerationStep.module == MODULE, + ModuleGenerationStep.step_code == step_code, + ModuleGenerationStep.is_current == True, + ModuleGenerationStep.deleted_at.is_(None), + ) + .order_by(ModuleGenerationStep.version.desc(), ModuleGenerationStep.created_at.desc()) + .limit(1) + ) + return result.scalar_one_or_none() + + +async def _next_version(db: AsyncSession, project_id: str, step_code: str) -> int: + result = await db.execute( + select(func.max(ModuleGenerationStep.version)).where( + ModuleGenerationStep.project_id == project_id, + ModuleGenerationStep.module == MODULE, + ModuleGenerationStep.step_code == step_code, + ) + ) + return int(result.scalar_one_or_none() or 0) + 1 + + +async def _create_step( + db: AsyncSession, + *, + project: ModuleGenerationProject, + step_code: str, + status: str = ModuleStepStatusEnum.PENDING.value, + parent_step_id: str | None = None, + source_step_id: str | None = None, + chat_task_id: str | None = None, + input_data: dict[str, Any] | None = None, + output_data: dict[str, Any] | None = None, +) -> ModuleGenerationStep: + version = await _next_version(db, project.id, step_code) + step = ModuleGenerationStep( + id=generate_id(), + project_id=project.id, + user_id=project.user_id, + module=project.module, + step_index=STEP_INDEX_MAP[step_code], + step_code=step_code, + status=status, + version=version, + is_current=True, + parent_step_id=parent_step_id, + source_step_id=source_step_id, + chat_task_id=chat_task_id, + input_json=_step_input( + step_code=step_code, + payload=input_data, + source_step_id=source_step_id, + parent_step_id=parent_step_id, + ) if input_data is not None else None, + output_json=_step_output( + step_code=step_code, + status=status, + result=output_data, + ) if output_data is not None else None, + started_at=_now() if status == ModuleStepStatusEnum.PROCESSING.value else None, + completed_at=_now() if status == ModuleStepStatusEnum.COMPLETED.value else None, + ) + db.add(step) + project.current_step_code = step_code + await db.flush() + await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.STEP_CREATED.value, detail={"step_code": step_code, "version": version}) + return step + + +async def _soft_delete_steps_from_index( + db: AsyncSession, + *, + project: ModuleGenerationProject, + start_index: int, + deleted_at: datetime | None = None, +) -> None: + deleted_at = deleted_at or _now() + result = await db.execute( + select(ModuleGenerationStep) + .where( + ModuleGenerationStep.project_id == project.id, + ModuleGenerationStep.module == MODULE, + ModuleGenerationStep.is_current == True, + ModuleGenerationStep.deleted_at.is_(None), + ModuleGenerationStep.step_index >= start_index, + ) + .with_for_update() + ) + steps = list(result.scalars().all()) + for step in steps: + step.is_current = False + step.deleted_at = deleted_at + if step.chat_task_id: + chat_result = await db.execute( + select(ChatGenerationTask) + .where(ChatGenerationTask.id == step.chat_task_id, ChatGenerationTask.deleted_at.is_(None)) + .with_for_update() + .limit(1) + ) + chat_task = chat_result.scalar_one_or_none() + if chat_task: + if chat_task.status == "completed": + await soft_delete_chat_task_resources(db, chat_task.id, deleted_at=deleted_at) + elif chat_task.status != "failed": + await mark_chat_generation_task_failed_and_refund_once( + db, + task=chat_task, + error_message="爆款开头复刻步骤被重新生成或删除,旧生成任务已取消", + pipeline_stage="failed", + ) + chat_task.deleted_at = deleted_at + if steps: + await log_module_event( + db, + project=project, + event_type=ModuleEventTypeEnum.SOFT_DELETE_STEPS.value, + message=f"软删除第 {start_index} 步及之后的旧子任务", + detail={"step_ids": [step.id for step in steps]}, + ) + + +def _step_to_out(step: ModuleGenerationStep) -> HotOpeningStepOut: + return HotOpeningStepOut( + id=step.id, + project_id=step.project_id, + module=step.module, + step_index=step.step_index, + step_code=step.step_code, + status=step.status, + version=step.version, + is_current=step.is_current, + parent_step_id=step.parent_step_id, + source_step_id=step.source_step_id, + chat_task_id=step.chat_task_id, + input=_parse_json(step.input_json, {}), + output=_parse_json(step.output_json, {}), + error_message=step.error_message, + created_at=step.created_at, + updated_at=step.updated_at, + completed_at=step.completed_at, + ) + + +def _snapshot_from_chat(chat_task: ChatGenerationTask | None) -> dict[str, Any]: + if not chat_task: + return {} + return _parse_json(chat_task.engine_snapshot_json, {}) or {} + + +async def _chat_tasks_by_id(db: AsyncSession, steps: list[ModuleGenerationStep]) -> dict[str, ChatGenerationTask]: + ids = [step.chat_task_id for step in steps if step.chat_task_id] + if not ids: + return {} + result = await db.execute(select(ChatGenerationTask).where(ChatGenerationTask.id.in_(ids))) + return {task.id: task for task in result.scalars().all()} + + +async def project_to_detail_out(db: AsyncSession, project: ModuleGenerationProject) -> HotOpeningTaskDetailOut: + steps = await _get_current_steps(db, project.id) + by_code = {step.step_code: step for step in steps} + chats = await _chat_tasks_by_id(db, steps) + + material_step = by_code.get(HotOpeningStepCodeEnum.MATERIAL_INPUT.value) + image_prompt_step = by_code.get(HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value) + image_generate_step = by_code.get(HotOpeningStepCodeEnum.IMAGE_GENERATE.value) + video_prompt_step = by_code.get(HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value) + video_generate_step = by_code.get(HotOpeningStepCodeEnum.VIDEO_GENERATE.value) + + material_input = _step_payload(material_step.input_json if material_step else None) + image_prompt_output = _unwrap_step_output(image_prompt_step.output_json if image_prompt_step else None) + image_generate_input = _step_payload(image_generate_step.input_json if image_generate_step else None) + image_generate_output = _unwrap_step_output(image_generate_step.output_json if image_generate_step else None) + video_prompt_input = _step_payload(video_prompt_step.input_json if video_prompt_step else None) + video_prompt_output = _unwrap_step_output(video_prompt_step.output_json if video_prompt_step else None) + video_generate_input = _step_payload(video_generate_step.input_json if video_generate_step else None) + video_generate_output = _unwrap_step_output(video_generate_step.output_json if video_generate_step else None) + + image_chat = chats.get(image_generate_step.chat_task_id) if image_generate_step and image_generate_step.chat_task_id else None + video_chat = chats.get(video_generate_step.chat_task_id) if video_generate_step and video_generate_step.chat_task_id else None + image_snapshot = _snapshot_from_chat(image_chat) + video_snapshot = _snapshot_from_chat(video_chat) + + image_url = image_generate_output.get("result_image_url") or (image_chat.image_url if image_chat else None) or project.final_image_url + video_url = video_generate_output.get("result_video_url") or (video_chat.video_url if video_chat else None) or project.final_video_url + cover_url = video_generate_output.get("result_video_cover_url") or (video_chat.video_cover_url if video_chat else None) or project.final_video_cover_url + + return HotOpeningTaskDetailOut( + id=project.id, + project_id=project.id, + module=project.module, + title=project.title, + status=project.status, + current_step_code=project.current_step_code, + final_image_url=build_resource_signed_url(project.final_image_url) if project.final_image_url else None, + final_video_url=build_resource_signed_url(project.final_video_url) if project.final_video_url else None, + final_video_cover_url=build_resource_signed_url(project.final_video_cover_url) if project.final_video_cover_url else None, + error_message=project.error_message, + material=HotOpeningMaterialOut( + material_step_id=material_step.id if material_step else None, + material_video_url=material_input.get("material_video_url"), + material_image_url=material_input.get("material_image_url"), + source_project_name=material_input.get("source_project_name"), + target_project_name=material_input.get("target_project_name"), + core_content_point=material_input.get("core_content_point"), + ), + image_generation=HotOpeningImageGenerationOut( + prompt_step_id=image_prompt_step.id if image_prompt_step else None, + generate_step_id=image_generate_step.id if image_generate_step else None, + prompt=image_prompt_output.get("optimized_prompt") or image_prompt_output.get("prompt"), + engine_id=image_snapshot.get("id") or image_generate_input.get("engine_id"), + engine_name=image_snapshot.get("name") or image_generate_input.get("engine_name"), + params=image_generate_input.get("params") or image_generate_input, + chat_task_id=image_generate_step.chat_task_id if image_generate_step else None, + status=image_chat.status if image_chat else (image_generate_step.status if image_generate_step else None), + result_image_url=build_resource_signed_url(image_url) if image_url else None, + error_message=image_chat.error_message if image_chat else (image_generate_step.error_message if image_generate_step else None), + ), + video_generation=HotOpeningVideoGenerationOut( + prompt_step_id=video_prompt_step.id if video_prompt_step else None, + generate_step_id=video_generate_step.id if video_generate_step else None, + prompt_schema=video_prompt_output.get("prompt_schema"), + final_prompt=video_prompt_output.get("final_prompt"), + prompt_params=video_prompt_output.get("params_used_for_prompt") or video_prompt_input.get("video_config"), + engine_id=video_snapshot.get("id") or video_generate_input.get("engine_id"), + engine_name=video_snapshot.get("name") or video_generate_input.get("engine_name"), + params=video_generate_input.get("params") or video_generate_input, + chat_task_id=video_generate_step.chat_task_id if video_generate_step else None, + status=video_chat.status if video_chat else (video_generate_step.status if video_generate_step else None), + result_video_url=build_resource_signed_url(video_url) if video_url else None, + result_video_cover_url=build_resource_signed_url(cover_url) if cover_url else None, + error_message=video_chat.error_message if video_chat else (video_generate_step.error_message if video_generate_step else None), + ), + steps=[_step_to_out(step) for step in steps], + created_at=project.created_at, + updated_at=project.updated_at, + completed_at=project.completed_at, + ) + + +async def create_hot_opening_project(db: AsyncSession, current_user: User, req: HotOpeningTaskCreate) -> ModuleGenerationProject: + if req.idempotency_key: + result = await db.execute( + select(ModuleGenerationProject) + .where( + ModuleGenerationProject.user_id == current_user.id, + ModuleGenerationProject.module == MODULE, + ModuleGenerationProject.idempotency_key == req.idempotency_key, + ModuleGenerationProject.deleted_at.is_(None), + ) + .order_by(ModuleGenerationProject.created_at.desc()) + .limit(1) + ) + existing = result.scalar_one_or_none() + if existing: + return existing + + project = ModuleGenerationProject( + id=generate_id(), + user_id=current_user.id, + module=MODULE, + title=req.target_project_name, + status=ModuleProjectStatusEnum.WAITING_USER.value, + current_step_code=HotOpeningStepCodeEnum.MATERIAL_INPUT.value, + idempotency_key=req.idempotency_key, + ) + db.add(project) + await db.flush() + + await _create_step( + db, + project=project, + step_code=HotOpeningStepCodeEnum.MATERIAL_INPUT.value, + status=ModuleStepStatusEnum.COMPLETED.value, + input_data={ + "material_video_url": req.material_video_url, + "material_image_url": req.material_image_url, + "source_project_name": req.source_project_name, + "target_project_name": req.target_project_name, + "core_content_point": req.core_content_point, + }, + output_data={"message": "素材输入已提交,后端不做素材文件校验。下一步请手动生成图片AI提词。"}, + ) + await log_module_event(db, project=project, event_type=ModuleEventTypeEnum.PROJECT_CREATED.value, message="创建爆款开头复刻项目") + return project + + +async def list_hot_opening_projects( + db: AsyncSession, + *, + current_user: User, + status: str | None, + page: int, + page_size: int, +) -> HotOpeningTaskListOut: + query = select(ModuleGenerationProject).where( + ModuleGenerationProject.module == MODULE, + ModuleGenerationProject.deleted_at.is_(None), + ) + if not current_user.is_admin: + query = query.where(ModuleGenerationProject.user_id == current_user.id) + if status: + query = query.where(ModuleGenerationProject.status == status) + + total = (await db.execute(select(func.count()).select_from(query.subquery()))).scalar_one() + result = await db.execute(query.order_by(ModuleGenerationProject.created_at.desc()).offset((page - 1) * page_size).limit(page_size)) + projects = list(result.scalars().all()) + + items: list[HotOpeningTaskListItemOut] = [] + for project in projects: + material_step = await _get_current_step_by_code(db, project.id, HotOpeningStepCodeEnum.MATERIAL_INPUT.value) + material = _step_payload(material_step.input_json if material_step else None) + items.append( + HotOpeningTaskListItemOut( + id=project.id, + project_id=project.id, + module=project.module, + title=project.title, + status=project.status, + current_step_code=project.current_step_code, + target_project_name=material.get("target_project_name"), + final_image_url=build_resource_signed_url(project.final_image_url) if project.final_image_url else None, + final_video_url=build_resource_signed_url(project.final_video_url) if project.final_video_url else None, + error_message=project.error_message, + created_at=project.created_at, + updated_at=project.updated_at, + completed_at=project.completed_at, + ) + ) + return HotOpeningTaskListOut(total=total, items=items) + + +async def update_hot_opening_step( + db: AsyncSession, + *, + current_user: User, + project_id: str, + step_id: str, + req: HotOpeningStepUpdate, +) -> tuple[ModuleGenerationProject, ModuleGenerationStep]: + project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) + step = await _get_step_for_user(db, project_id=project_id, step_id=step_id, user=current_user, for_update=True) + if step.status == ModuleStepStatusEnum.PROCESSING.value: + raise HTTPException(status_code=400, detail="当前子任务正在处理中,暂不能修改") + + input_data = _step_payload(step.input_json) + output_data = _unwrap_step_output(step.output_json) + + if step.step_code == HotOpeningStepCodeEnum.MATERIAL_INPUT.value: + input_data = _merge_dict( + input_data, + { + "material_video_url": req.material_video_url, + "material_image_url": req.material_image_url, + "source_project_name": req.source_project_name, + "target_project_name": req.target_project_name, + "core_content_point": req.core_content_point, + }, + ) + if req.target_project_name: + project.title = req.target_project_name + elif step.step_code == HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value: + if req.prompt is not None: + output_data["optimized_prompt"] = req.prompt + output_data["prompt"] = req.prompt + elif step.step_code == HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value: + if req.prompt_schema is not None: + output_data["prompt_schema"] = req.prompt_schema + if req.prompt is not None: + output_data["final_prompt"] = req.prompt + else: + if req.input_json: + input_data = _merge_dict(input_data, req.input_json) + if req.output_json: + output_data = _merge_dict(output_data, req.output_json) + + if req.input_json: + input_data = _merge_dict(input_data, req.input_json) + if req.output_json: + output_data = _merge_dict(output_data, req.output_json) + + step.input_json = _step_input(step_code=step.step_code, payload=input_data, source_step_id=step.source_step_id, parent_step_id=step.parent_step_id) + step.output_json = _step_output(step_code=step.step_code, status=ModuleStepStatusEnum.COMPLETED.value, payload=output_data) + step.status = ModuleStepStatusEnum.COMPLETED.value + step.error_message = None + step.completed_at = _now() + project.status = ModuleProjectStatusEnum.WAITING_USER.value + project.current_step_code = step.step_code + project.error_message = None + + await _soft_delete_steps_from_index(db, project=project, start_index=step.step_index + 1) + await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.STEP_UPDATED.value, message="用户修改子任务内容") + return project, step + + +async def update_hot_opening_material_input( + db: AsyncSession, + *, + current_user: User, + project_id: str, + req: HotOpeningMaterialUpdateRequest, +) -> tuple[str, str]: + """修改第1步素材输入。 + + 采用方案 B:软删除旧第1步及之后的当前有效步骤,然后新建第1步 version+1。 + 未传字段沿用旧第1步素材输入,避免前端只改一个字段时丢失其它素材信息。 + """ + project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) + old_material_step = await _get_current_step_by_code(db, project.id, HotOpeningStepCodeEnum.MATERIAL_INPUT.value) + old_material = _step_payload(old_material_step.input_json if old_material_step else None) + + material = { + "material_video_url": req.material_video_url if req.material_video_url is not None else old_material.get("material_video_url"), + "material_image_url": req.material_image_url if req.material_image_url is not None else old_material.get("material_image_url"), + "source_project_name": req.source_project_name if req.source_project_name is not None else old_material.get("source_project_name"), + "target_project_name": req.target_project_name if req.target_project_name is not None else old_material.get("target_project_name"), + "core_content_point": req.core_content_point if req.core_content_point is not None else old_material.get("core_content_point"), + } + + missing_fields = [key for key, value in material.items() if value is None or str(value).strip() == ""] + if missing_fields: + raise HTTPException(status_code=400, detail=f"素材输入缺少必要字段: {', '.join(missing_fields)}") + + await _soft_delete_steps_from_index(db, project=project, start_index=STEP_INDEX_MAP[HotOpeningStepCodeEnum.MATERIAL_INPUT.value]) + + project.title = str(material["target_project_name"]) + project.status = ModuleProjectStatusEnum.WAITING_USER.value + project.current_step_code = HotOpeningStepCodeEnum.MATERIAL_INPUT.value + project.final_image_url = None + project.final_video_url = None + project.final_video_cover_url = None + project.error_message = None + project.completed_at = None + + new_step = await _create_step( + db, + project=project, + step_code=HotOpeningStepCodeEnum.MATERIAL_INPUT.value, + status=ModuleStepStatusEnum.COMPLETED.value, + input_data=material, + output_data={"message": "素材输入已修改,旧步骤已软删除。下一步请重新生成图片AI提词。"}, + ) + await log_module_event( + db, + project=project, + step=new_step, + event_type=ModuleEventTypeEnum.STEP_UPDATED.value, + message="用户修改素材输入并重建第1步新版本", + detail={ + "old_material_step_id": old_material_step.id if old_material_step else None, + "new_material_step_id": new_step.id, + "version": new_step.version, + }, + ) + return project.id, new_step.id + + +async def update_hot_opening_image_prompt( + db: AsyncSession, + *, + current_user: User, + project_id: str, + step_id: str, + req: HotOpeningImagePromptUpdateRequest, +) -> tuple[ModuleGenerationProject, ModuleGenerationStep]: + """直接修改第2步图片 AI 优化提词,不调用 AI、不扣积分。 + + 修改后软删除第3、4、5步当前有效任务,让用户从图片生成开始重新执行。 + """ + project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) + step = await _get_step_for_user(db, project_id=project_id, step_id=step_id, user=current_user, for_update=True) + if step.step_code != HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value: + raise HTTPException(status_code=400, detail="只能修改第2步图片 AI 提词子任务") + if step.status != ModuleStepStatusEnum.COMPLETED.value: + raise HTTPException(status_code=400, detail="图片 AI 提词未完成,不能直接修改") + + output_data = _step_payload(step.output_json) + usage = _step_usage(step.output_json) + new_prompt = req.prompt.strip() + output_data["optimized_prompt"] = new_prompt + output_data["prompt"] = new_prompt + output_data["manual_edited"] = True + output_data["manual_edited_at"] = _now().isoformat() + + step.output_json = _step_output( + step_code=HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value, + status=ModuleStepStatusEnum.COMPLETED.value, + payload=output_data, + usage=usage, + ) + step.status = ModuleStepStatusEnum.COMPLETED.value + step.error_message = None + step.completed_at = _now() + + await _soft_delete_steps_from_index(db, project=project, start_index=STEP_INDEX_MAP[HotOpeningStepCodeEnum.IMAGE_GENERATE.value]) + + project.status = ModuleProjectStatusEnum.WAITING_USER.value + project.current_step_code = HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value + project.final_image_url = None + project.final_video_url = None + project.final_video_cover_url = None + project.completed_at = None + project.error_message = None + + await log_module_event( + db, + project=project, + step=step, + event_type=ModuleEventTypeEnum.STEP_UPDATED.value, + message="用户直接修改图片 AI 优化提词,已软删除后续步骤", + detail={"start_deleted_step_index": STEP_INDEX_MAP[HotOpeningStepCodeEnum.IMAGE_GENERATE.value]}, + ) + return project, step + + +async def update_hot_opening_video_prompt_schema( + db: AsyncSession, + *, + current_user: User, + project_id: str, + step_id: str, + req: HotOpeningVideoPromptSchemaUpdateRequest, +) -> tuple[ModuleGenerationProject, ModuleGenerationStep]: + """以前端 schema 为 patch 修改第4步视频 AI 提词,不调用 AI、不扣积分。 + + 服务端已有 schema 为基准:视频规格、数组长度、时间段、合规控制、质量控制、协议字段均锁定。 + 最终提示词允许修改,但保存前会清洗视频时长、比例、分辨率、帧率等参数。 + """ + project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) + step = await _get_step_for_user(db, project_id=project_id, step_id=step_id, user=current_user, for_update=True) + if step.step_code != HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value: + raise HTTPException(status_code=400, detail="只能修改第4步视频 AI 提词 JSON schema 子任务") + if step.status != ModuleStepStatusEnum.COMPLETED.value: + raise HTTPException(status_code=400, detail="视频 AI 提词未完成,不能直接修改") + + output_data = _step_payload(step.output_json) + usage = _step_usage(step.output_json) + input_data = _step_payload(step.input_json) + server_schema = output_data.get("prompt_schema") if isinstance(output_data.get("prompt_schema"), dict) else {} + video_config = output_data.get("params_used_for_prompt") or input_data.get("video_config") or {} + if not isinstance(video_config, dict) or not video_config.get("duration") or not video_config.get("aspect_ratio") or not video_config.get("resolution"): + raise HTTPException(status_code=400, detail="缺少第4步视频参数快照,不能安全修改视频 schema") + + patched_schema = patch_video_prompt_schema_from_client( + server_schema=server_schema, + client_schema=req.prompt_schema, + video_config=video_config, + ) + final_prompt = build_final_video_prompt(patched_schema) + + output_data["prompt_schema"] = patched_schema + output_data["final_prompt"] = final_prompt + output_data["params_used_for_prompt"] = video_config + output_data["manual_edited"] = True + output_data["manual_edited_at"] = _now().isoformat() + + step.output_json = _step_output( + step_code=HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value, + status=ModuleStepStatusEnum.COMPLETED.value, + payload=output_data, + usage=usage, + ) + step.status = ModuleStepStatusEnum.COMPLETED.value + step.error_message = None + step.completed_at = _now() + + await _soft_delete_steps_from_index(db, project=project, start_index=STEP_INDEX_MAP[HotOpeningStepCodeEnum.VIDEO_GENERATE.value]) + + project.status = ModuleProjectStatusEnum.WAITING_USER.value + project.current_step_code = HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value + project.final_video_url = None + project.final_video_cover_url = None + project.completed_at = None + project.error_message = None + + await log_module_event( + db, + project=project, + step=step, + event_type=ModuleEventTypeEnum.STEP_UPDATED.value, + message="用户修改视频 AI 提词 schema,已软删除视频生成步骤", + detail={ + "start_deleted_step_index": STEP_INDEX_MAP[HotOpeningStepCodeEnum.VIDEO_GENERATE.value], + "locked_fields": [ + "schema_version", + "schema_usage", + "画面属性.视频时长", + "画面属性.视频比例", + "画面属性.清晰度", + "画面属性.帧率", + "画面属性.推荐分辨率", + "动作流程[*].时间段", + "镜头流程[*].时间段", + "动态时间规划", + "输出规格限制", + "质量控制", + "合规控制", + ], + }, + ) + return project, step + + +async def submit_image_prompt_optimize( + db: AsyncSession, + *, + current_user: User, + project_id: str, + material_step_id: str, +) -> tuple[ModuleGenerationProject, ModuleGenerationStep]: + project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) + material_step = await _get_step_for_user(db, project_id=project_id, step_id=material_step_id, user=current_user, for_update=True) + if material_step.step_code != HotOpeningStepCodeEnum.MATERIAL_INPUT.value: + raise HTTPException(status_code=400, detail="请基于第1步素材输入子任务生成图片 AI 提词") + if material_step.status != ModuleStepStatusEnum.COMPLETED.value: + raise HTTPException(status_code=400, detail="素材输入子任务未完成,不能生成图片 AI 提词") + + await _soft_delete_steps_from_index(db, project=project, start_index=STEP_INDEX_MAP[HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value]) + step = await _create_step( + db, + project=project, + step_code=HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value, + status=ModuleStepStatusEnum.PROCESSING.value, + parent_step_id=material_step.id, + source_step_id=material_step.id, + input_data={"source_step_id": material_step.id}, + ) + project.status = ModuleProjectStatusEnum.PROCESSING.value + project.current_step_code = HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value + project.error_message = None + await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.IMAGE_PROMPT_SUBMITTED.value, message="图片 AI 提词任务已提交") + return project, step + + +async def run_image_prompt_optimize(db: AsyncSession, *, project_id: str, step_id: str | None = None) -> ModuleGenerationStep | None: + project_result = await db.execute( + select(ModuleGenerationProject) + .where(ModuleGenerationProject.id == project_id, ModuleGenerationProject.module == MODULE, ModuleGenerationProject.deleted_at.is_(None)) + .with_for_update() + .limit(1) + ) + project = project_result.scalar_one_or_none() + if not project: + return None + + material_step = await _get_current_step_by_code(db, project.id, HotOpeningStepCodeEnum.MATERIAL_INPUT.value) + if not material_step: + project.status = ModuleProjectStatusEnum.FAILED.value + project.error_message = "缺少素材输入子任务" + return None + + if step_id: + result = await db.execute( + select(ModuleGenerationStep) + .where( + ModuleGenerationStep.id == step_id, + ModuleGenerationStep.project_id == project.id, + ModuleGenerationStep.step_code == HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value, + ModuleGenerationStep.deleted_at.is_(None), + ModuleGenerationStep.is_current == True, + ) + .with_for_update() + .limit(1) + ) + step = result.scalar_one_or_none() + else: + step = await _get_current_step_by_code(db, project.id, HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value) + if not step: + if step_id: + # 用户重复提交后,旧 Celery 消息对应的 step 可能已被软删。 + # 指定 step_id 查不到时必须静默忽略,不能重新创建步骤导致旧任务复活。 + return None + step = await _create_step( + db, + project=project, + step_code=HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value, + status=ModuleStepStatusEnum.PROCESSING.value, + parent_step_id=material_step.id, + source_step_id=material_step.id, + input_data={"source_step_id": material_step.id}, + ) + else: + step.status = ModuleStepStatusEnum.PROCESSING.value + step.started_at = _now() + step.error_message = None + + material = _step_payload(material_step.input_json) + prompt_text = ( + "请基于参考素材复刻爆款开头视觉风格,用于生成新项目图片。\n" + f"视频素材内容项目名称:{material.get('source_project_name')}\n" + f"生成项目名称:{material.get('target_project_name')}\n" + f"生成项目核心内容点:{material.get('core_content_point')}\n" + "要求:参考素材视频的开头构图、主体位置、节奏和风格;结合新产品图片生成新项目推广图片;不要照抄原素材品牌、文字、水印;适合作为后续图生视频首帧。" + ) + references = [ + {"type": "video", "url": material.get("material_video_url"), "name": "参考素材视频"}, + {"type": "image", "url": material.get("material_image_url"), "name": "新产品图片"}, + ] + try: + request_log = {"original_prompt": prompt_text, "references": references, "gen_type": "image"} + log_module_prompt_event( + event_type="module_prompt_request", + project_id=project.id, + step_id=step.id, + user_id=project.user_id, + module=project.module, + prompt_type=ModulePromptTypeEnum.IMAGE_PROMPT.value, + request=request_log, + ) + optimized, token_usage = await optimize_prompt( + db, + original_prompt=prompt_text, + user_id=project.user_id, + references=references, + gen_type="image", + ) + billing = await charge_module_prompt_usage( + db, + user_id=project.user_id, + step_id=step.id, + usage=token_usage, + description="爆款开头复刻-图片AI提词优化", + ) + usage = dict(token_usage or {}) + usage.update({ + "text_credits_cost": (billing.items[0].amount if billing.items else billing.total_charged), + "credit_biz_key": billing.items[0].biz_key if billing.items else None, + }) + step.status = ModuleStepStatusEnum.COMPLETED.value + step.completed_at = _now() + step.output_json = _step_output( + step_code=HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value, + status=ModuleStepStatusEnum.COMPLETED.value, + payload={ + "optimized_prompt": optimized, + "prompt": optimized, + "original_prompt": prompt_text, + "references": references, + }, + usage=usage, + ) + project.status = ModuleProjectStatusEnum.WAITING_USER.value + project.current_step_code = HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value + project.error_message = None + log_module_prompt_event( + event_type="module_prompt_response", + project_id=project.id, + step_id=step.id, + user_id=project.user_id, + module=project.module, + prompt_type=ModulePromptTypeEnum.IMAGE_PROMPT.value, + request=request_log, + response={"optimized_prompt": optimized}, + token_usage=usage, + ) + await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.IMAGE_PROMPT_SUCCESS.value, message="图片 AI 提词生成成功") + except Exception as exc: + step.status = ModuleStepStatusEnum.FAILED.value + step.error_message = str(exc) + step.completed_at = _now() + project.status = ModuleProjectStatusEnum.FAILED.value + project.error_message = f"图片 AI 提词生成失败: {exc}" + await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.IMAGE_PROMPT_FAILED.value, message=project.error_message) + return step + + +async def generate_image_from_prompt( + db: AsyncSession, + *, + current_user: User, + project_id: str, + prompt_step_id: str, + req: HotOpeningGenerateImageRequest, +) -> tuple[ModuleGenerationProject, ModuleGenerationStep]: + project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) + prompt_step = await _get_step_for_user(db, project_id=project_id, step_id=prompt_step_id, user=current_user, for_update=True) + if prompt_step.step_code != HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value: + raise HTTPException(status_code=400, detail="请基于第2步图片 AI 提词子任务生成图片") + if prompt_step.status != ModuleStepStatusEnum.COMPLETED.value: + raise HTTPException(status_code=400, detail="图片 AI 提词未完成,不能生成图片") + + await _soft_delete_steps_from_index(db, project=project, start_index=STEP_INDEX_MAP[HotOpeningStepCodeEnum.IMAGE_GENERATE.value]) + + material_step = await _get_current_step_by_code(db, project.id, HotOpeningStepCodeEnum.MATERIAL_INPUT.value) + material = _step_payload(material_step.input_json if material_step else None) + prompt_output = _unwrap_step_output(prompt_step.output_json) + optimized_prompt = prompt_output.get("optimized_prompt") or prompt_output.get("prompt") or "" + refs = [ + {"type": "image", "url": material.get("material_image_url"), "name": "新产品图片"}, + ] + + chat_task = await create_chat_generation_task_for_module( + db, + current_user=current_user, + generation_mode=GENERATION_MODE, + gen_type="image", + original_prompt=prompt_output.get("original_prompt") or optimized_prompt, + optimized_prompt=optimized_prompt, + engine_id=req.engine_id, + media_references=refs, + image_size=req.image_size, + image_proportion=req.image_proportion, + image_px=req.image_px, + billing_project_name=project.title or "爆款开头复刻", + billing_description_prefix="爆款开头复刻图片生成", + ) + step = await _create_step( + db, + project=project, + step_code=HotOpeningStepCodeEnum.IMAGE_GENERATE.value, + status=ModuleStepStatusEnum.PROCESSING.value, + parent_step_id=prompt_step.id, + source_step_id=prompt_step.id, + chat_task_id=chat_task.id, + input_data={ + "engine_id": chat_task.engine_id, + "params": {"image_size": chat_task.image_size, "image_proportion": chat_task.image_proportion, "image_px": chat_task.image_px}, + "prompt": optimized_prompt, + "media_references": refs, + }, + ) + project.status = ModuleProjectStatusEnum.PROCESSING.value + project.current_step_code = HotOpeningStepCodeEnum.IMAGE_GENERATE.value + project.error_message = None + await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.IMAGE_GENERATE_SUBMITTED.value, message="图片生成任务已提交", detail={"chat_task_id": chat_task.id}) + return project, step + + +async def _resolve_video_prompt_config(db: AsyncSession, req: HotOpeningGenerateVideoPromptRequest) -> dict[str, Any]: + engine = await _get_video_engine(db, req.engine_id) + supported_ratios = _parse_list(engine.supported_ratios, []) + supported_resolutions = _parse_list(engine.supported_resolutions, []) + supported_durations = _parse_list(engine.supported_durations, []) + + default_ratio = getattr(settings, "HOT_OPENING_DEFAULT_VIDEO_RATIO", None) or VIDEO_DEFAULT_RATIO + default_resolution = getattr(settings, "HOT_OPENING_DEFAULT_VIDEO_RESOLUTION", None) or VIDEO_DEFAULT_RESOLUTION + default_duration = int(getattr(settings, "HOT_OPENING_DEFAULT_VIDEO_DURATION", None) or VIDEO_DEFAULT_DURATION) + + selected_ratio = req.aspect_ratio or (default_ratio if not supported_ratios or default_ratio in supported_ratios else supported_ratios[0]) + selected_resolution = req.resolution or (default_resolution if not supported_resolutions or default_resolution in supported_resolutions else supported_resolutions[0]) + selected_duration = req.duration or (default_duration if not supported_durations or default_duration in supported_durations else supported_durations[0]) + + if supported_ratios and selected_ratio not in supported_ratios: + raise HTTPException(status_code=400, detail=f"视频比例不支持: {selected_ratio}") + if supported_resolutions and selected_resolution not in supported_resolutions: + raise HTTPException(status_code=400, detail=f"视频分辨率不支持: {selected_resolution}") + if supported_durations and selected_duration not in supported_durations: + raise HTTPException(status_code=400, detail=f"视频时长不支持: {selected_duration}") + if engine.max_duration and int(selected_duration) > int(engine.max_duration): + raise HTTPException(status_code=400, detail=f"视频时长不能超过 {engine.max_duration} 秒") + + return { + "engine_id": engine.id, + "engine_name": engine.name, + "duration": int(selected_duration), + "aspect_ratio": selected_ratio, + "resolution": selected_resolution, + "supported_ratios": supported_ratios, + "supported_resolutions": supported_resolutions, + "supported_durations": supported_durations, + "max_duration": engine.max_duration, + "frame_rate": "30fps", + "reference_video_fps": max(1, int(settings.CHATAPI_VIDEO_FPS or 1)), + } + + +async def submit_video_prompt_optimize( + db: AsyncSession, + *, + current_user: User, + project_id: str, + image_step_id: str, + req: HotOpeningGenerateVideoPromptRequest, +) -> tuple[ModuleGenerationProject, ModuleGenerationStep]: + project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) + image_step = await _get_step_for_user(db, project_id=project_id, step_id=image_step_id, user=current_user, for_update=True) + if image_step.step_code != HotOpeningStepCodeEnum.IMAGE_GENERATE.value: + raise HTTPException(status_code=400, detail="请基于第3步图片生成子任务生成视频 AI 提词") + if image_step.status != ModuleStepStatusEnum.COMPLETED.value: + raise HTTPException(status_code=400, detail="图片生成子任务未完成,不能生成视频 AI 提词") + + await _soft_delete_steps_from_index(db, project=project, start_index=STEP_INDEX_MAP[HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value]) + video_config = await _resolve_video_prompt_config(db, req) + step = await _create_step( + db, + project=project, + step_code=HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value, + status=ModuleStepStatusEnum.PROCESSING.value, + parent_step_id=image_step.id, + source_step_id=image_step.id, + input_data={ + "source_step_id": image_step.id, + "video_config": video_config, + "target_platform": req.target_platform or getattr(settings, "HOT_OPENING_DEFAULT_TARGET_PLATFORM", "抖音") or "抖音", + }, + ) + project.status = ModuleProjectStatusEnum.PROCESSING.value + project.current_step_code = HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value + project.error_message = None + await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.VIDEO_PROMPT_SUBMITTED.value, message="视频 AI 提词任务已提交") + return project, step + + +async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_id: str | None = None) -> ModuleGenerationStep | None: + project_result = await db.execute( + select(ModuleGenerationProject) + .where(ModuleGenerationProject.id == project_id, ModuleGenerationProject.module == MODULE, ModuleGenerationProject.deleted_at.is_(None)) + .with_for_update() + .limit(1) + ) + project = project_result.scalar_one_or_none() + if not project: + return None + + material_step = await _get_current_step_by_code(db, project.id, HotOpeningStepCodeEnum.MATERIAL_INPUT.value) + image_prompt_step = await _get_current_step_by_code(db, project.id, HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value) + image_step = await _get_current_step_by_code(db, project.id, HotOpeningStepCodeEnum.IMAGE_GENERATE.value) + if not material_step or not image_step: + project.status = ModuleProjectStatusEnum.FAILED.value + project.error_message = "生成视频提词失败:缺少素材输入或图片生成结果" + return None + + if step_id: + result = await db.execute( + select(ModuleGenerationStep) + .where( + ModuleGenerationStep.id == step_id, + ModuleGenerationStep.project_id == project.id, + ModuleGenerationStep.step_code == HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value, + ModuleGenerationStep.deleted_at.is_(None), + ModuleGenerationStep.is_current == True, + ) + .with_for_update() + .limit(1) + ) + step = result.scalar_one_or_none() + else: + step = await _get_current_step_by_code(db, project.id, HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value) + if not step: + if step_id: + # 用户重复提交后,旧 Celery 消息对应的 step 可能已被软删。 + # 指定 step_id 查不到时必须静默忽略,不能把当前项目标记失败。 + return None + project.status = ModuleProjectStatusEnum.FAILED.value + project.error_message = "缺少视频 AI 提词子任务,请先手动提交视频提词生成" + return None + + step.status = ModuleStepStatusEnum.PROCESSING.value + step.started_at = _now() + step.error_message = None + + material = _step_payload(material_step.input_json) + image_output = _unwrap_step_output(image_step.output_json) + step_input = _step_payload(step.input_json) + video_config = step_input.get("video_config") or {} + target_platform = step_input.get("target_platform") or getattr(settings, "HOT_OPENING_DEFAULT_TARGET_PLATFORM", "抖音") or "抖音" + generated_image_url = image_output.get("result_image_url") or project.final_image_url + if not generated_image_url: + step.status = ModuleStepStatusEnum.FAILED.value + step.error_message = "缺少新项目图片结果,不能生成视频提词" + project.status = ModuleProjectStatusEnum.FAILED.value + project.error_message = step.error_message + return step + + try: + request_log = { + "source_project_name": material.get("source_project_name") or "无", + "target_project_name": material.get("target_project_name") or "无", + "core_content_point": material.get("core_content_point") or "无", + "material_video_url": material.get("material_video_url") or "", + "generated_image_url": generated_image_url, + "video_config": video_config, + "target_platform": target_platform, + } + log_module_prompt_event( + event_type="module_prompt_request", + project_id=project.id, + step_id=step.id, + user_id=project.user_id, + module=project.module, + prompt_type=ModulePromptTypeEnum.VIDEO_PROMPT.value, + request=request_log, + ) + prompt_schema, final_prompt, token_usage = await optimize_hot_opening_video_prompt( + db, + user_id=project.user_id, + source_project_name=request_log["source_project_name"], + target_project_name=request_log["target_project_name"], + core_content_point=request_log["core_content_point"], + material_video_url=request_log["material_video_url"], + generated_image_url=generated_image_url, + video_config=video_config, + target_platform=target_platform, + ) + billing = await charge_module_prompt_usage( + db, + user_id=project.user_id, + step_id=step.id, + usage=token_usage, + description="爆款开头复刻-视频AI提词优化", + ) + usage = dict(token_usage or {}) + usage.update({ + "text_credits_cost": (billing.items[0].amount if billing.items else billing.total_charged), + "credit_biz_key": billing.items[0].biz_key if billing.items else None, + }) + step.status = ModuleStepStatusEnum.COMPLETED.value + step.completed_at = _now() + step.output_json = _step_output( + step_code=HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value, + status=ModuleStepStatusEnum.COMPLETED.value, + payload={ + "prompt_schema": prompt_schema, + "final_prompt": final_prompt, + "params_used_for_prompt": video_config, + "target_platform": target_platform, + }, + usage=usage, + ) + project.status = ModuleProjectStatusEnum.WAITING_USER.value + project.current_step_code = HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value + project.error_message = None + log_module_prompt_event( + event_type="module_prompt_response", + project_id=project.id, + step_id=step.id, + user_id=project.user_id, + module=project.module, + prompt_type=ModulePromptTypeEnum.VIDEO_PROMPT.value, + request=request_log, + response={"prompt_schema": prompt_schema, "final_prompt": final_prompt}, + token_usage=usage, + ) + await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.VIDEO_PROMPT_SUCCESS.value, message="视频 AI 提词生成成功") + except Exception as exc: + step.status = ModuleStepStatusEnum.FAILED.value + step.error_message = str(exc) + step.completed_at = _now() + project.status = ModuleProjectStatusEnum.FAILED.value + project.error_message = f"视频 AI 提词生成失败: {exc}" + await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.VIDEO_PROMPT_FAILED.value, message=project.error_message) + return step + + +async def generate_video_from_prompt( + db: AsyncSession, + *, + current_user: User, + project_id: str, + prompt_step_id: str, + req: HotOpeningGenerateVideoRequest, +) -> tuple[ModuleGenerationProject, ModuleGenerationStep]: + project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) + prompt_step = await _get_step_for_user(db, project_id=project_id, step_id=prompt_step_id, user=current_user, for_update=True) + if prompt_step.step_code != HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value: + raise HTTPException(status_code=400, detail="请基于第4步视频 AI 提词子任务生成视频") + if prompt_step.status != ModuleStepStatusEnum.COMPLETED.value: + raise HTTPException(status_code=400, detail="视频 AI 提词未完成,不能生成视频") + + await _soft_delete_steps_from_index(db, project=project, start_index=STEP_INDEX_MAP[HotOpeningStepCodeEnum.VIDEO_GENERATE.value]) + + image_step = await _get_current_step_by_code(db, project.id, HotOpeningStepCodeEnum.IMAGE_GENERATE.value) + image_output = _unwrap_step_output(image_step.output_json if image_step else None) + prompt_output = _unwrap_step_output(prompt_step.output_json) + final_prompt = prompt_output.get("final_prompt") or "" + prompt_schema = prompt_output.get("prompt_schema") or {} + prompt_schema_str = json.dumps(prompt_schema, ensure_ascii=False, default=str) if prompt_schema else "" + prompt_input = _step_payload(prompt_step.input_json) + prompt_params = prompt_output.get("params_used_for_prompt") or prompt_input.get("video_config") or {} + duration = int(prompt_params.get("duration") or settings.HOT_OPENING_DEFAULT_VIDEO_DURATION or 4) + aspect_ratio = prompt_params.get("aspect_ratio") or settings.HOT_OPENING_DEFAULT_VIDEO_RATIO or "9:16" + resolution = prompt_params.get("resolution") or settings.HOT_OPENING_DEFAULT_VIDEO_RESOLUTION or "480p" + generated_image_url = image_output.get("result_image_url") or project.final_image_url + if not generated_image_url: + raise HTTPException(status_code=400, detail="缺少新项目图片结果,不能生成视频") + + refs = [ + {"type": "image", "url": _build_file_url_or_data_uri(generated_image_url), "name": "新项目图片"}, + ] + + chat_task = await create_chat_generation_task_for_module( + db, + current_user=current_user, + generation_mode=GENERATION_MODE, + gen_type="video", + original_prompt=prompt_schema_str or final_prompt, + optimized_prompt=prompt_schema_str or final_prompt, + engine_id=req.engine_id or prompt_params.get("engine_id"), + media_references=refs, + duration=duration, + aspect_ratio=aspect_ratio, + resolution=resolution, + billing_project_name=project.title or "爆款开头复刻", + billing_description_prefix="爆款开头复刻视频生成", + ) + step = await _create_step( + db, + project=project, + step_code=HotOpeningStepCodeEnum.VIDEO_GENERATE.value, + status=ModuleStepStatusEnum.PROCESSING.value, + parent_step_id=prompt_step.id, + source_step_id=prompt_step.id, + chat_task_id=chat_task.id, + input_data={ + "engine_id": chat_task.engine_id, + "params": { + "duration": chat_task.duration, + "aspect_ratio": chat_task.aspect_ratio, + "resolution": chat_task.resolution, + "image_size": chat_task.image_size, + "image_proportion": chat_task.image_proportion, + "image_px": chat_task.image_px, + }, + "prompt_schema": prompt_schema, + "final_prompt": final_prompt, + "media_references": refs, + }, + ) + project.status = ModuleProjectStatusEnum.PROCESSING.value + project.current_step_code = HotOpeningStepCodeEnum.VIDEO_GENERATE.value + project.error_message = None + await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.VIDEO_GENERATE_SUBMITTED.value, message="视频生成任务已提交", detail={"chat_task_id": chat_task.id}) + return project, step + + +async def handle_chat_generation_task_completed(db: AsyncSession, task: ChatGenerationTask) -> None: + if not task or task.generation_mode != GENERATION_MODE: + return + result = await db.execute( + select(ModuleGenerationStep) + .where( + ModuleGenerationStep.chat_task_id == task.id, + ModuleGenerationStep.module == MODULE, + ModuleGenerationStep.is_current == True, + ModuleGenerationStep.deleted_at.is_(None), + ) + .with_for_update() + .limit(1) + ) + step = result.scalar_one_or_none() + if not step: + return + project_result = await db.execute( + select(ModuleGenerationProject) + .where(ModuleGenerationProject.id == step.project_id, ModuleGenerationProject.deleted_at.is_(None)) + .with_for_update() + .limit(1) + ) + project = project_result.scalar_one_or_none() + if not project: + return + + if step.step_code == HotOpeningStepCodeEnum.IMAGE_GENERATE.value: + step.status = ModuleStepStatusEnum.COMPLETED.value + step.completed_at = _now() + step.output_json = _step_output( + step_code=HotOpeningStepCodeEnum.IMAGE_GENERATE.value, + status=ModuleStepStatusEnum.COMPLETED.value, + result={"result_image_url": task.image_url, "chat_task_id": task.id}, + ) + project.final_image_url = task.image_url + project.status = ModuleProjectStatusEnum.WAITING_USER.value + project.current_step_code = HotOpeningStepCodeEnum.IMAGE_GENERATE.value + await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.IMAGE_GENERATE_SUCCESS.value, message="图片生成完成,等待用户手动生成视频 AI 提词") + elif step.step_code == HotOpeningStepCodeEnum.VIDEO_GENERATE.value: + step.status = ModuleStepStatusEnum.COMPLETED.value + step.completed_at = _now() + step.output_json = _step_output( + step_code=HotOpeningStepCodeEnum.VIDEO_GENERATE.value, + status=ModuleStepStatusEnum.COMPLETED.value, + result={"result_video_url": task.video_url, "result_video_cover_url": task.video_cover_url, "chat_task_id": task.id}, + ) + project.final_video_url = task.video_url + project.final_video_cover_url = task.video_cover_url + project.status = ModuleProjectStatusEnum.COMPLETED.value + project.current_step_code = HotOpeningStepCodeEnum.VIDEO_GENERATE.value + project.completed_at = _now() + await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.VIDEO_GENERATE_SUCCESS.value, message="视频生成完成,总任务完成") + + +async def handle_chat_generation_task_failed(db: AsyncSession, task: ChatGenerationTask) -> None: + if not task or task.generation_mode != GENERATION_MODE: + return + result = await db.execute( + select(ModuleGenerationStep) + .where(ModuleGenerationStep.chat_task_id == task.id, ModuleGenerationStep.module == MODULE, ModuleGenerationStep.is_current == True, ModuleGenerationStep.deleted_at.is_(None)) + .with_for_update() + .limit(1) + ) + step = result.scalar_one_or_none() + if not step: + return + project_result = await db.execute(select(ModuleGenerationProject).where(ModuleGenerationProject.id == step.project_id).with_for_update().limit(1)) + project = project_result.scalar_one_or_none() + if not project: + return + step.status = ModuleStepStatusEnum.FAILED.value + step.error_message = task.error_message + step.completed_at = _now() + project.status = ModuleProjectStatusEnum.FAILED.value + project.error_message = task.error_message or "生成失败" + await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.CHAT_TASK_FAILED.value, message=project.error_message, detail={"chat_task_id": task.id}) + + +async def mark_hot_opening_step_dispatch_failed( + db: AsyncSession, + *, + current_user: User, + project_id: str, + step_id: str, + error_message: str, +) -> None: + project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) + result = await db.execute( + select(ModuleGenerationStep) + .where( + ModuleGenerationStep.id == step_id, + ModuleGenerationStep.project_id == project.id, + ModuleGenerationStep.module == MODULE, + ModuleGenerationStep.deleted_at.is_(None), + ModuleGenerationStep.is_current == True, + ) + .with_for_update() + .limit(1) + ) + step = result.scalar_one_or_none() + if not step: + return + if step.chat_task_id: + await mark_chat_generation_task_failed_and_refund_once( + db, + task_id=step.chat_task_id, + error_message=error_message, + pipeline_stage="failed", + ) + step.status = ModuleStepStatusEnum.FAILED.value + step.error_message = error_message + step.completed_at = _now() + project.status = ModuleProjectStatusEnum.FAILED.value + project.error_message = error_message + await log_module_event( + db, + project=project, + step=step, + event_type=ModuleEventTypeEnum.CHAT_TASK_FAILED.value, + message=error_message, + detail={"reason": "celery_dispatch_failed"}, + ) + + +async def delete_hot_opening_project(db: AsyncSession, *, current_user: User, project_id: str) -> HotOpeningDeleteOut: + project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) + deleted_at = _now() + project.deleted_at = deleted_at + await _soft_delete_steps_from_index(db, project=project, start_index=1, deleted_at=deleted_at) + await log_module_event(db, project=project, event_type=ModuleEventTypeEnum.PROJECT_DELETED.value, message="软删除爆款开头复刻项目") + return HotOpeningDeleteOut(message="项目已删除", project_id=project.id, deleted=True) + +def _build_file_url_or_data_uri(file_url: str) -> str: + """ + Convert local upload path to base64 data URI. + Keep remote http/https/data URLs as-is. + """ + if file_url.startswith(("http://", "https://", "data:")): + return file_url + file_url_sign = build_resource_signed_url(resource_url=file_url, expire_seconds=86400) + return f"{settings.BASE_URL}{file_url_sign}" \ No newline at end of file diff --git a/video-gen-api/app/services/hot_opening_video_prompt_service.py b/video-gen-api/app/services/hot_opening_video_prompt_service.py new file mode 100644 index 00000000..d6b2858f --- /dev/null +++ b/video-gen-api/app/services/hot_opening_video_prompt_service.py @@ -0,0 +1,707 @@ +from __future__ import annotations + +import copy +import json +import re +from typing import Any + +import httpx +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from app.config import settings +from app.enums.video_prompt_schema import PromptSchemaVersionEnum, VideoPromptSchemaUsageEnum +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.resource_signed_url_service import build_resource_signed_url + +DEFAULT_FRAME_RATE = "30fps" +DEFAULT_REFERENCE_VIDEO_FPS = 1 + +CLIENT_SCHEMA_V1: dict[str, Any] = { + "schema_version": PromptSchemaVersionEnum.CLIENT_V1.value, + "schema_usage": VideoPromptSchemaUsageEnum.CLIENT_DISPLAY.value, + "基础分类": { + "生成类型": "文生视频/图生视频/视频生视频/数字人视频/无", + "视频大类": "产品广告视频/电商带货视频/口播讲解视频/剧情视频/教程视频/风景旅行视频/美食视频/宠物视频/动漫卡通视频/游戏视频/企业宣传视频/新闻资讯视频/直播切片视频/图文快闪视频/音乐舞蹈视频/运动健身视频/无", + "视频子类": "产品推广短视频/口播讲解/电商带货/剧情演绎/操作教程/旅行风景/美食展示/宠物互动/二次元动画/游戏宣传/企业介绍/资讯播报/直播高光/图文快闪/无", + "视频用途": "广告投放/社媒发布/产品展示/课程教学/品牌宣传/娱乐内容/信息科普/无", + "目标平台": "抖音/快手/小红书/微信视频号/B站/TikTok/YouTube Shorts/Instagram Reels/无", + }, + "素材理解": { + "是否有参考图片": "是/否", + "是否有参考视频": "是/否", + "参考视频用途": "动作参考/镜头参考/风格参考/运镜参考/节奏参考/无", + "需要保留": [], + "允许改动": [], + "禁止改动": [], + }, + "业务属性": { + "产品类型": "APP/实物商品/食品/服饰/美妆/电子产品/汽车/房产/课程/服务/无", + "产品名称": "无", + "品牌名称": "无", + "核心卖点": [], + "目标受众": "无", + "核心表达目标": "无", + "内容风格": "无", + "行动引导": "立即体验/立即下载/立即购买/点击了解/预约咨询/无", + }, + "画面属性": { + "视频时长": "无", + "视频比例": "无", + "清晰度": "无", + "帧率": "无", + "主体描述": "无", + "主体数量": "无", + "主体位置": "无", + "主体占比": "无", + "场景描述": "无", + "构图方式": "无", + "画面风格": "无", + "光影色彩": "无", + }, + "动作流程": [], + "镜头流程": [], + "字幕与口播": { + "是否需要字幕": "是/否", + "字幕内容": [], + "字幕位置": "无", + "字幕样式": "无", + "是否口播": "是/否", + "口播内容": "无", + "口播语气": "无", + "口播语速": "无", + "是否需要口型同步": "是/否/无", + }, + "音频与节奏": { + "背景音乐": "无", + "音乐风格": "无", + "音乐节奏": "无", + "环境音": "无", + "动作音效": "无", + "整体节奏": "慢节奏/中等节奏/快节奏/卡点节奏/无", + }, + "合规控制": { + "是否广告": "是/否", + "风险等级": "低/中/高", + "安全表达": "无", + "禁用词": [], + "合规说明": "无", + }, + "质量控制": { + "主体一致性": "低/中/高/无", + "产品一致性": "低/中/高/无", + "动作自然度": "低/中/高/无", + "镜头稳定性": "低/中/高/无", + "字幕准确性": "低/中/高/无", + }, + "最终提示词": { + "主提示词": "无", + "动作提示词": "无", + "镜头提示词": "无", + "字幕提示词": "无", + "音频提示词": "无", + "风格提示词": "无", + "负面提示词": "无", + }, +} + + +def _safe_list(value: Any) -> list[Any]: + return value if isinstance(value, list) else [] + + +def _format_options(options: list[Any], fallback: str = "无") -> str: + values = [str(item) for item in _safe_list(options) if str(item).strip()] + return "/".join(values) if values else fallback + + +def get_recommended_resolution(video_ratio: str, resolution: str) -> str: + """按比例和清晰度粗略计算推荐像素,不在服务内维护固定比例/分辨率白名单。""" + try: + width_ratio, height_ratio = [float(x) for x in str(video_ratio).split(":", 1)] + short_edge = int(str(resolution).lower().replace("p", "")) + if width_ratio >= height_ratio: + height = short_edge + width = round(short_edge * width_ratio / height_ratio) + else: + width = short_edge + height = round(short_edge * height_ratio / width_ratio) + return f"{width}x{height}" + except Exception: + return "无" + + +def _build_scaled_bounds(duration: int, ratios: list[float]) -> list[int]: + duration = max(1, int(duration)) + raw = [0] + acc = 0.0 + for ratio in ratios[:-1]: + acc += ratio + raw.append(max(raw[-1] + 1, min(duration - 1, round(duration * acc)))) + raw.append(duration) + for index in range(1, len(raw)): + if raw[index] <= raw[index - 1]: + raw[index] = min(duration, raw[index - 1] + 1) + raw[-1] = duration + return raw + + +def _bounds_to_plan(bounds: list[int], stages: list[tuple[str, str]]) -> list[dict[str, str]]: + plan: list[dict[str, str]] = [] + for idx, (stage, desc) in enumerate(stages): + start = bounds[idx] + end = bounds[idx + 1] + plan.append({"时间段": f"{start}-{end}秒", "阶段": stage, "说明": desc}) + return plan + + +def build_time_plan(duration: int) -> list[dict[str, str]]: + duration = max(1, int(duration)) + if duration <= 5: + return _bounds_to_plan( + _build_scaled_bounds(duration, [0.2, 0.4, 0.4]), + [ + ("开场吸引", "快速建立主体、产品和画面风格"), + ("核心展示", "展示主体动作、核心卖点或主要视觉内容"), + ("行动引导", "强化记忆点并给出转化引导"), + ], + ) + if duration <= 8: + return _bounds_to_plan( + _build_scaled_bounds(duration, [0.15, 0.3, 0.35, 0.2]), + [ + ("开场吸引", "快速吸引注意力"), + ("主体展示", "展示主体和产品关系"), + ("核心卖点", "突出新项目核心内容点"), + ("收尾引导", "给出行动引导并稳定落版"), + ], + ) + return _bounds_to_plan( + _build_scaled_bounds(duration, [0.13, 0.2, 0.27, 0.25, 0.15]), + [ + ("爆款开头", "复刻参考素材开头节奏和视觉吸引点"), + ("主体建立", "明确新项目主体和产品信息"), + ("卖点放大", "围绕核心内容点展开动作和镜头"), + ("情绪推进", "用动作、字幕或镜头变化强化记忆"), + ("转化收尾", "给出清晰行动引导"), + ], + ) + + +def build_dynamic_schema(video_config: dict[str, Any]) -> dict[str, Any]: + schema = copy.deepcopy(CLIENT_SCHEMA_V1) + duration = int(video_config["duration"]) + video_ratio = str(video_config["aspect_ratio"]) + resolution = str(video_config["resolution"]) + frame_rate = str(video_config.get("frame_rate") or DEFAULT_FRAME_RATE) + recommended_resolution = get_recommended_resolution(video_ratio, resolution) + + schema["画面属性"].update( + { + "视频时长": f"{duration}秒", + "视频比例": video_ratio, + "清晰度": resolution, + "推荐分辨率": recommended_resolution, + "帧率": frame_rate, + } + ) + schema["动态时间规划"] = build_time_plan(duration) + schema["输出规格限制"] = { + "支持时长": _safe_list(video_config.get("supported_durations")), + "支持比例": _safe_list(video_config.get("supported_ratios")), + "支持分辨率": _safe_list(video_config.get("supported_resolutions")), + "当前推荐分辨率": recommended_resolution, + } + return schema + + +def infer_generation_type(references: list[dict[str, str]] | None) -> str: + has_image = any(item.get("type") == "image" for item in references or []) + has_video = any(item.get("type") == "video" for item in references or []) + if has_image and has_video: + return "图生视频/视频生视频" + if has_image: + return "图生视频" + if has_video: + return "视频生视频" + return "文生视频" + + +def infer_video_category(text: str) -> tuple[str, str, list[str]]: + text = text or "" + if any(key in text for key in ["APP", "应用", "下载", "社交", "脱单", "附近"]): + return "产品广告视频", "产品推广短视频", ["产品推广", "用户转化", "核心卖点展示"] + return "产品广告视频", "产品推广短视频", ["产品展示", "视觉吸引", "行动引导"] + + +def build_system_prompt() -> str: + return ( + "你是专业短视频广告导演和AI视频提示词工程师。" + "你必须只输出一个合法 JSON 对象,不能输出 Markdown。" + "输出必须严格遵循用户提供的 schema 顶层结构。" + "所有未知、无法判断或不适用的字段填写'无',数组字段可填写 []。" + "必须根据参考视频复刻爆款开头的节奏、构图、动作和镜头语言,但不能照抄品牌、水印、字幕或侵权元素。" + ) + + +def build_user_text( + *, + source_project_name: str, + target_project_name: str, + core_content_point: str, + target_platform: str, + references: list[dict[str, str]], + video_config: dict[str, Any], + client_schema: dict[str, Any], +) -> str: + duration = int(video_config["duration"]) + video_ratio = str(video_config["aspect_ratio"]) + resolution = str(video_config["resolution"]) + frame_rate = str(video_config.get("frame_rate") or DEFAULT_FRAME_RATE) + category, sub_category, default_points = infer_video_category(" ".join([source_project_name, target_project_name, core_content_point])) + return json.dumps( + { + "任务": "基于参考素材视频和新项目图片,生成可用于AI视频生成的中文结构化提示词JSON", + "业务输入": { + "视频素材内容项目名称": source_project_name, + "生成项目名称": target_project_name, + "生成项目核心内容点": core_content_point, + "目标平台": target_platform, + "默认视频大类": category, + "默认视频子类": sub_category, + "建议核心卖点": default_points, + }, + "视频规格": { + "视频时长": f"{duration}秒", + "视频比例": video_ratio, + "清晰度": resolution, + "帧率": frame_rate, + "支持时长": _safe_list(video_config.get("supported_durations")), + "支持比例": _safe_list(video_config.get("supported_ratios")), + "支持分辨率": _safe_list(video_config.get("supported_resolutions")), + "推荐分辨率": get_recommended_resolution(video_ratio, resolution), + }, + "参考素材": references, + "输出要求": { + "生成类型": infer_generation_type(references), + "必须填充动态时间规划": build_time_plan(duration), + "必须填充动作流程": "动作流程时间段必须覆盖完整视频时长", + "必须填充镜头流程": "镜头流程时间段必须覆盖完整视频时长", + "最终提示词限制": "最终提示词下所有字段都不能写入视频时长、秒数、视频比例、清晰度、分辨率、帧率、推荐像素、竖屏、横屏等视频规格参数,这些规格只能写在画面属性/动态时间规划/输出规格限制。", + "禁止": ["输出 Markdown", "输出 schema 之外的解释文字", "照抄参考素材品牌水印", "生成违法违规内容", "在最终提示词中写入秒数/比例/分辨率/帧率"], + }, + "必须按此schema输出": client_schema, + }, + ensure_ascii=False, + ) + + +def build_user_message(user_content: str, references: list[dict[str, str]], reference_video_fps: int) -> tuple[dict[str, Any], dict[str, Any]]: + content_parts: list[dict[str, Any]] = [{"type": "text", "text": user_content}] + log_content_parts: list[dict[str, Any]] = [{"type": "text", "text": user_content}] + for ref in references: + ref_type = ref.get("type") + ref_url = ref.get("url") + if not ref_url: + continue + if ref_type == "image": + content_parts.append({"type": "image_url", "image_url": {"url": ref_url}}) + log_content_parts.append({"type": "image_url", "image_url": {"url": ref_url}}) + elif ref_type == "video": + content_parts.append({"type": "video_url", "video_url": {"url": ref_url, "fps": reference_video_fps}}) + log_content_parts.append({"type": "video_url", "video_url": {"url": ref_url, "fps": reference_video_fps}}) + return {"role": "user", "content": content_parts}, {"role": "user", "content": log_content_parts} + + +def strip_json_code_fence(text: str) -> str: + text = (text or "").strip() + if text.startswith("```"): + text = re.sub(r"^```(?:json)?\s*", "", text, flags=re.I) + text = re.sub(r"\s*```$", "", text) + return text.strip() + + +def parse_model_json(content: str) -> dict[str, Any]: + content = strip_json_code_fence(content) + data = json.loads(content) + if not isinstance(data, dict): + raise ValueError("视频提词优化返回值不是 JSON 对象") + return data + + +def fill_none_with_wu(value: Any) -> Any: + if value is None or value == "": + return "无" + if isinstance(value, dict): + return {k: fill_none_with_wu(v) for k, v in value.items()} + if isinstance(value, list): + return [fill_none_with_wu(v) for v in value] + return value + + +def ensure_top_keys(result: dict[str, Any]) -> dict[str, Any]: + schema = copy.deepcopy(CLIENT_SCHEMA_V1) + for key, default_value in schema.items(): + if key not in result: + result[key] = default_value + elif isinstance(default_value, dict) and isinstance(result.get(key), dict): + merged = copy.deepcopy(default_value) + merged.update(result[key]) + result[key] = merged + return result + + +def ensure_flow_matches_time_plan(result: dict[str, Any], duration: int) -> dict[str, Any]: + plan = build_time_plan(duration) + if not isinstance(result.get("动作流程"), list) or not result["动作流程"]: + result["动作流程"] = [ + {"时间段": item["时间段"], "动作内容": item["说明"]} + for item in plan + ] + if not isinstance(result.get("镜头流程"), list) or not result["镜头流程"]: + result["镜头流程"] = [ + {"时间段": item["时间段"], "镜头内容": item["说明"]} + for item in plan + ] + result["动作流程"] = _align_flow_time_ranges(result["动作流程"], plan, "动作内容") + result["镜头流程"] = _align_flow_time_ranges(result["镜头流程"], plan, "镜头内容") + result["动态时间规划"] = plan + return result + + +def ensure_negative_prompt(result: dict[str, Any]) -> dict[str, Any]: + final = result.setdefault("最终提示词", {}) + if not isinstance(final, dict): + final = {} + result["最终提示词"] = final + if not final.get("负面提示词") or final.get("负面提示词") == "无": + final["负面提示词"] = "画面模糊、主体畸变、手指畸形、脸部崩坏、字幕乱码、产品变形、镜头抖动、画面闪烁" + return result + + +def build_final_video_prompt(result: dict[str, Any]) -> str: + final = result.get("最终提示词", {}) if isinstance(result.get("最终提示词"), dict) else {} + parts = [ + final.get("主提示词"), + final.get("动作提示词"), + final.get("镜头提示词"), + final.get("字幕提示词"), + final.get("音频提示词"), + final.get("风格提示词"), + ] + return "\n".join(str(item).strip() for item in parts if item and str(item).strip() != "无") + + +VIDEO_SPEC_PROMPT_KEYS = ("主提示词", "动作提示词", "镜头提示词", "字幕提示词", "音频提示词", "风格提示词", "负面提示词") + + +def _normalize_prompt_text(text: str) -> str: + text = re.sub(r"[,、,;;::]\s*([,、,;;::])", r"\1", text) + text = re.sub(r"\s{2,}", " ", text) + text = re.sub(r"^[,、,;;::\s]+", "", text) + text = re.sub(r"[,、,;;::\s]+$", "", text) + return text.strip() or "无" + + +def clean_video_spec_from_prompt_text(text: Any, video_config: dict[str, Any] | None = None) -> str: + """清洗最终提示词里的视频规格参数。 + + 视频时长、比例、清晰度、分辨率、帧率属于接口参数和锁定字段, + 不能混入最终提示词,避免用户绕过扣费参数或与第5步视频生成参数冲突。 + """ + if text is None: + return "无" + value = str(text).strip() + if not value or value == "无": + return "无" + + cfg = video_config or {} + exact_values = { + str(cfg.get("aspect_ratio") or "").strip(), + str(cfg.get("resolution") or "").strip(), + str(cfg.get("frame_rate") or "").strip(), + } + try: + if cfg.get("duration") is not None: + exact_values.add(f"{int(cfg.get('duration'))}秒") + except Exception: + pass + try: + if cfg.get("aspect_ratio") and cfg.get("resolution"): + exact_values.add(get_recommended_resolution(str(cfg.get("aspect_ratio")), str(cfg.get("resolution")))) + except Exception: + pass + + for item in sorted((v for v in exact_values if v and v != "无"), key=len, reverse=True): + value = value.replace(item, "") + + patterns = [ + r"\d+\s*秒", + r"\b\d+\s*[sS]\b", + r"\d+\s*[::]\s*\d+", + r"\d{3,4}\s*[pP]", + r"\d{2,4}\s*[xX×]\s*\d{2,4}", + r"\d+\s*(?:fps|FPS|帧)", + r"(?:竖屏|横屏|方屏|超清|高清|标清|蓝光|4K|8K)", + r"(?:视频时长|时长|视频比例|画面比例|比例|分辨率|清晰度|帧率|推荐分辨率)\s*[::]?\s*", + ] + for pattern in patterns: + value = re.sub(pattern, "", value, flags=re.I) + return _normalize_prompt_text(value) + + +def clean_final_prompt_specs(schema: dict[str, Any], video_config: dict[str, Any] | None = None) -> dict[str, Any]: + final = schema.setdefault("最终提示词", {}) + if not isinstance(final, dict): + final = {} + schema["最终提示词"] = final + for key in VIDEO_SPEC_PROMPT_KEYS: + final[key] = clean_video_spec_from_prompt_text(final.get(key), video_config) + return schema + + +def _align_flow_time_ranges(flow: Any, plan: list[dict[str, str]], default_content_key: str) -> list[dict[str, Any]]: + source = flow if isinstance(flow, list) else [] + aligned: list[dict[str, Any]] = [] + for index, plan_item in enumerate(plan): + old_item = source[index] if index < len(source) and isinstance(source[index], dict) else {} + item = dict(old_item) + item["时间段"] = plan_item["时间段"] + if not any(k in item and str(item.get(k)).strip() for k in (default_content_key, "动作", "镜头", "说明", "内容")): + item[default_content_key] = plan_item["说明"] + aligned.append(item) + return aligned + + +def _merge_editable_dict_fields(base: dict[str, Any], patch: dict[str, Any], allowed_keys: set[str]) -> None: + for key in allowed_keys: + if key in patch: + base[key] = fill_none_with_wu(patch.get(key)) + + +def _merge_flow_patch(base_flow: Any, patch_flow: Any) -> list[dict[str, Any]]: + base = [dict(item) for item in base_flow] if isinstance(base_flow, list) else [] + patch = patch_flow if isinstance(patch_flow, list) else [] + result: list[dict[str, Any]] = [] + for index, base_item in enumerate(base): + merged = dict(base_item) + patch_item = patch[index] if index < len(patch) and isinstance(patch[index], dict) else {} + original_time_range = merged.get("时间段") + for key, value in patch_item.items(): + if key == "时间段": + continue + merged[key] = fill_none_with_wu(value) + merged["时间段"] = original_time_range + result.append(merged) + return result + + +def apply_locked_video_schema_fields(schema: dict[str, Any], video_config: dict[str, Any]) -> dict[str, Any]: + duration = int(video_config["duration"]) + video_ratio = str(video_config["aspect_ratio"]) + resolution = str(video_config["resolution"]) + frame_rate = str(video_config.get("frame_rate") or DEFAULT_FRAME_RATE) + recommended_resolution = get_recommended_resolution(video_ratio, resolution) + dynamic_schema = build_dynamic_schema(video_config) + plan = build_time_plan(duration) + + schema["schema_version"] = PromptSchemaVersionEnum.CLIENT_V1.value + schema["schema_usage"] = VideoPromptSchemaUsageEnum.CLIENT_DISPLAY.value + + frame = schema.setdefault("画面属性", {}) + if not isinstance(frame, dict): + frame = {} + schema["画面属性"] = frame + frame.update( + { + "视频时长": f"{duration}秒", + "视频比例": video_ratio, + "清晰度": resolution, + "帧率": frame_rate, + "推荐分辨率": recommended_resolution, + } + ) + + schema["动态时间规划"] = plan + schema["输出规格限制"] = dynamic_schema.get("输出规格限制", {}) + + schema["动作流程"] = _align_flow_time_ranges(schema.get("动作流程"), plan, "动作内容") + schema["镜头流程"] = _align_flow_time_ranges(schema.get("镜头流程"), plan, "镜头内容") + + # 合规和质量控制不能被前端降低;AI 返回缺失时使用服务端默认结构补齐。 + default_schema = copy.deepcopy(CLIENT_SCHEMA_V1) + if not isinstance(schema.get("合规控制"), dict): + schema["合规控制"] = default_schema["合规控制"] + if not isinstance(schema.get("质量控制"), dict): + schema["质量控制"] = default_schema["质量控制"] + + return clean_final_prompt_specs(schema, video_config) + + +def normalize_video_prompt_schema_from_ai(result: dict[str, Any], video_config: dict[str, Any]) -> dict[str, Any]: + duration = int(video_config["duration"]) + normalized = ensure_top_keys(fill_none_with_wu(result if isinstance(result, dict) else {})) + normalized = ensure_flow_matches_time_plan(normalized, duration) + normalized = ensure_negative_prompt(normalized) + return apply_locked_video_schema_fields(normalized, video_config) + + +def patch_video_prompt_schema_from_client( + *, + server_schema: dict[str, Any], + client_schema: dict[str, Any], + video_config: dict[str, Any], +) -> dict[str, Any]: + """以前端 JSON 作为 patch,回填到服务端已有 schema。 + + 禁止整包覆盖:数组长度、时间段、视频规格、输出规格、质量控制、合规控制、schema 协议字段均以服务端为准。 + """ + base = ensure_top_keys(fill_none_with_wu(copy.deepcopy(server_schema if isinstance(server_schema, dict) else {}))) + patch = client_schema if isinstance(client_schema, dict) else {} + + for key in ("基础分类", "素材理解", "业务属性", "字幕与口播", "音频与节奏"): + if isinstance(base.get(key), dict) and isinstance(patch.get(key), dict): + base[key].update(fill_none_with_wu(patch[key])) + + if isinstance(base.get("画面属性"), dict) and isinstance(patch.get("画面属性"), dict): + _merge_editable_dict_fields( + base["画面属性"], + patch["画面属性"], + {"主体描述", "主体数量", "主体位置", "主体占比", "场景描述", "构图方式", "画面风格", "光影色彩"}, + ) + + if isinstance(patch.get("动作流程"), list): + base["动作流程"] = _merge_flow_patch(base.get("动作流程"), patch.get("动作流程")) + if isinstance(patch.get("镜头流程"), list): + base["镜头流程"] = _merge_flow_patch(base.get("镜头流程"), patch.get("镜头流程")) + + # 动态时间规划保持服务端数组长度和时间段,只允许保留原值;不接受客户端 patch。 + + if isinstance(base.get("最终提示词"), dict) and isinstance(patch.get("最终提示词"), dict): + for key in VIDEO_SPEC_PROMPT_KEYS: + if key in patch["最终提示词"]: + base["最终提示词"][key] = fill_none_with_wu(patch["最终提示词"].get(key)) + + return apply_locked_video_schema_fields(base, video_config) + + +def _mock_result(video_config: dict[str, Any], target_platform: str) -> dict[str, Any]: + duration = int(video_config["duration"]) + video_ratio = str(video_config["aspect_ratio"]) + resolution = str(video_config["resolution"]) + schema = build_dynamic_schema(video_config) + schema["基础分类"].update({"生成类型": "图生视频/视频生视频", "视频大类": "产品广告视频", "视频子类": "产品推广短视频", "视频用途": "社媒发布", "目标平台": target_platform}) + schema["素材理解"].update({"是否有参考图片": "是", "是否有参考视频": "是", "参考视频用途": "动作参考/镜头参考/风格参考/节奏参考"}) + schema["业务属性"].update({"产品类型": "APP", "核心表达目标": "突出新项目核心内容点", "内容风格": "轻快、活泼、广告感适中", "行动引导": "立即体验"}) + schema["动作流程"] = [{"时间段": item["时间段"], "动作": item["阶段"], "说明": item["说明"]} for item in build_time_plan(duration)] + schema["镜头流程"] = [{"时间段": item["时间段"], "镜头": item["阶段"], "说明": item["说明"]} for item in build_time_plan(duration)] + schema["最终提示词"] = { + "主提示词": f"生成一段{duration}秒、{video_ratio}、{resolution}的产品推广短视频,参考素材视频的爆款开头节奏,结合新项目图片进行自然展示。", + "动作提示词": "主体动作自然,产品展示稳定,节奏轻快。", + "镜头提示词": "镜头稳定,开头快速吸引注意,后续平滑推进。", + "字幕提示词": "字幕简洁清晰,突出核心内容点。", + "音频提示词": "轻快背景音乐,节奏自然。", + "风格提示词": "年轻化、明亮、真实、广告感适中。", + "负面提示词": "画面模糊、主体畸变、手指畸形、脸部崩坏、字幕乱码、产品变形、镜头抖动、画面闪烁", + } + return schema + + +async def _select_model_config(db: AsyncSession) -> ModelConfig | None: + result = await db.execute(select(ModelConfig).where(ModelConfig.is_active == True).order_by(ModelConfig.priority.desc()).limit(1)) + return result.scalar_one_or_none() + + +async def optimize_hot_opening_video_prompt( + db: AsyncSession, + *, + user_id: str, + source_project_name: str, + target_project_name: str, + core_content_point: str, + material_video_url: str, + generated_image_url: str, + video_config: dict[str, Any], + target_platform: str = "抖音", +) -> tuple[dict[str, Any], str, dict[str, Any]]: + duration = int(video_config["duration"]) + references = [ + {"type": "video", "url": material_video_url}, + {"type": "image", "url": _build_file_url_or_data_uri(generated_image_url)}, + ] + client_schema = build_dynamic_schema(video_config) + reference_video_fps = int(video_config.get("reference_video_fps") or DEFAULT_REFERENCE_VIDEO_FPS) + + # if settings.LLM_MOCK: + # result = _mock_result(video_config, target_platform) + # 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) + if not config: + result = normalize_video_prompt_schema_from_ai(_mock_result(video_config, target_platform), video_config) + return result, build_final_video_prompt(result), {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0} + + user_text = build_user_text( + source_project_name=source_project_name, + target_project_name=target_project_name, + core_content_point=core_content_point, + target_platform=target_platform, + references=references, + video_config=video_config, + client_schema=client_schema, + ) + user_message, log_user_message = build_user_message(user_text, references, reference_video_fps) + request_data = { + "model": config.model_name, + "messages": [{"role": "system", "content": build_system_prompt()}, user_message], + "max_tokens": 6000, + "temperature": 0.15, + "response_format": {"type": "json_object"}, + } + + async with httpx.AsyncClient(timeout=int(settings.CHATAPI_REQUEST_TIMEOUT_SECONDS or 180)) as client: + response = await client.post( + f"{config.api_base.rstrip('/')}/chat/completions", + headers={"Authorization": f"Bearer {config.api_key}", "Content-Type": "application/json"}, + json=request_data, + ) + if response.status_code >= 400: + raise RuntimeError(f"视频提词优化失败 HTTP {response.status_code}: {response.text}") + + data = response.json() + content = data["choices"][0]["message"]["content"].strip() + usage = data.get("usage", {}) or {} + token_usage = { + "input_tokens": int(usage.get("prompt_tokens") or 0), + "output_tokens": int(usage.get("completion_tokens") or 0), + "total_tokens": int(usage.get("total_tokens") or 0), + "log_user_message": log_user_message, + } + db.add( + TokenUsage( + id=generate_id(), + model_config_id=config.id, + user_id=user_id, + input_tokens=token_usage["input_tokens"], + output_tokens=token_usage["output_tokens"], + total_tokens=token_usage["total_tokens"], + ) + ) + await db.flush() + + result = parse_model_json(content) + result = normalize_video_prompt_schema_from_ai(result, video_config) + return result, build_final_video_prompt(result), token_usage + +def _build_file_url_or_data_uri(file_url: str) -> str: + """ + Convert local upload path to base64 data URI. + Keep remote http/https/data URLs as-is. + """ + if file_url.startswith(("http://", "https://", "data:")): + return file_url + file_url_sign = build_resource_signed_url(resource_url=file_url, expire_seconds=86400) + return f"{settings.BASE_URL}{file_url_sign}" \ No newline at end of file diff --git a/video-gen-api/app/services/module_generation_log_service.py b/video-gen-api/app/services/module_generation_log_service.py new file mode 100644 index 00000000..1acb3b00 --- /dev/null +++ b/video-gen-api/app/services/module_generation_log_service.py @@ -0,0 +1,140 @@ +from __future__ import annotations + +import json +import os +import re +from datetime import datetime +from typing import Any + +from app.services.log_config import LOG_DATE_FORMAT, LOG_DIR, is_enabled + +MAX_LOG_FIELD_LENGTH = 20000 +MODULE_LOG_ROOT = os.path.join(os.path.dirname(LOG_DIR), "ModuleGeneration") + + +def _safe_module_name(module: str | None) -> str: + value = str(module or "unknown_module").strip() or "unknown_module" + value = re.sub(r"[^a-zA-Z0-9_.-]+", "_", value) + return value[:120] or "unknown_module" + + +def _safe_dump_value(value: Any) -> Any: + """限制单字段长度,避免超长 base64 / 响应体把日志打爆。""" + if value is None: + return None + if isinstance(value, str): + if len(value) > MAX_LOG_FIELD_LENGTH: + return value[:MAX_LOG_FIELD_LENGTH] + f"..." + return value + if isinstance(value, dict): + return {str(k): _safe_dump_value(v) for k, v in value.items()} + if isinstance(value, list): + return [_safe_dump_value(v) for v in value] + return value + + +def _append_module_log(module: str, entry: dict[str, Any]) -> None: + if not is_enabled(): + return + try: + module_dir = os.path.join(MODULE_LOG_ROOT, _safe_module_name(module)) + os.makedirs(module_dir, exist_ok=True) + today = datetime.now().strftime(LOG_DATE_FORMAT) + log_file = os.path.join(module_dir, f"{today}.log") + with open(log_file, "a", encoding="utf-8") as f: + f.write(json.dumps(entry, ensure_ascii=False, default=str) + "\n") + except Exception: + # 日志失败绝不能影响业务主流程。 + pass + + +def log_module_event_file( + *, + module: str, + event_type: str, + project_id: str | None = None, + step_id: str | None = None, + user_id: str | None = None, + message: str | None = None, + detail: dict[str, Any] | None = None, + error: str | None = None, +) -> None: + """记录模块流程事件到 JSONL 文件。 + + 统一落盘目录:log/ModuleGeneration/{module}/YYYY-MM-DD.log + 不再写 module_generation_events 表。 + """ + entry = { + "timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"), + "log_type": "module_event", + "module": module, + "event_type": event_type, + "project_id": project_id, + "step_id": step_id, + "user_id": user_id, + "message": message, + "detail": _safe_dump_value(detail or {}), + "error": error, + } + _append_module_log(module, entry) + + +def log_module_prompt_event( + *, + event_type: str, + project_id: str, + step_id: str, + user_id: str, + module: str, + prompt_type: str, + request: dict[str, Any] | None = None, + response: dict[str, Any] | None = None, + token_usage: dict[str, Any] | None = None, + error: str | None = None, +) -> None: + """记录模块 AI 提词请求/响应到 JSONL 文件。 + + 与模块事件共用同一个服务,但按 module 分目录,方便按模块排查。 + """ + entry = { + "timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"), + "log_type": "module_prompt", + "module": module, + "event_type": event_type, + "prompt_type": prompt_type, + "project_id": project_id, + "step_id": step_id, + "user_id": user_id, + "request": _safe_dump_value(request or {}), + "response": _safe_dump_value(response or {}), + "token_usage": _safe_dump_value(token_usage or {}), + "error": error, + } + _append_module_log(module, entry) + + +def log_module_error( + *, + module: str, + event_type: str, + project_id: str | None = None, + step_id: str | None = None, + user_id: str | None = None, + message: str | None = None, + detail: dict[str, Any] | None = None, + error: str | None = None, +) -> None: + """记录模块异常日志。""" + entry = { + "timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"), + "log_type": "module_error", + "module": module, + "event_type": event_type, + "project_id": project_id, + "step_id": step_id, + "user_id": user_id, + "message": message, + "detail": _safe_dump_value(detail or {}), + "error": error, + } + _append_module_log(module, entry) diff --git a/video-gen-api/app/services/payment.py b/video-gen-api/app/services/payment.py index ed0e8946..57cbca9c 100644 --- a/video-gen-api/app/services/payment.py +++ b/video-gen-api/app/services/payment.py @@ -1,14 +1,244 @@ import logging -from datetime import datetime +import os +from datetime import datetime, timedelta +# 尝试设置 SSL 证书路径 +try: + import certifi + os.environ["SSL_CERT_FILE"] = certifi.where() + os.environ["REQUESTS_CA_BUNDLE"] = certifi.where() +except ImportError: + pass + +from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.config import settings from app.models.payment_order import PaymentOrder -from app.services.credits import add_credits +from app.models.system_config import SystemConfig +from app.services.credits import add_credits, deduct_credits from app.utils.id_gen import generate_id, generate_order_no -logger = logging.getLogger("videogen") +# --------------------------------------------------------------------------- +# Payment logger → log/payment/YYYY-MM-DD.log (one file per day, no cleanup) +# --------------------------------------------------------------------------- +import time as _time + +logger = logging.getLogger("payment") +logger.setLevel(logging.INFO) + +_log_dir = os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(__file__))), "log", "payment") +os.makedirs(_log_dir, exist_ok=True) + + +class DailyFileHandler(logging.FileHandler): + """Write to a file named by date, e.g. log/payment/2026-06-10.log""" + + def __init__(self, directory, encoding="utf-8"): + self._directory = directory + self._current_date = "" + self._file_handler = None + super().__init__(self._make_path(), mode="a", encoding=encoding, delay=False) + + def _make_path(self): + date_str = _time.strftime("%Y-%m-%d") + self._current_date = date_str + return os.path.join(self._directory, f"{date_str}.log") + + def emit(self, record): + date_str = _time.strftime("%Y-%m-%d") + if date_str != self._current_date: + # Day rolled over — switch to a new file + if self._file_handler: + self._file_handler.close() + self.baseFilename = self._make_path() + self._file_handler = logging.FileHandler( + self.baseFilename, mode="a", encoding=self.encoding + ) + self._file_handler.setFormatter(self.formatter) + self._current_date = date_str + self.stream = self._file_handler.stream + super().emit(record) + + +_handler = DailyFileHandler(_log_dir) +_handler.setFormatter(logging.Formatter( + "[%(asctime)s] %(levelname)s %(message)s", datefmt="%Y-%m-%d %H:%M:%S" +)) +if not logger.handlers: + logger.addHandler(_handler) + +# Order expire time in seconds (configurable via payment_order_timeout setting, default 180 seconds) +DEFAULT_ORDER_EXPIRE_SECONDS = 180 + + +def _get_order_expire_seconds(db_configs: dict[str, str]) -> int: + """Get order expire time in seconds from config, with fallback to 180.""" + try: + val = db_configs.get("payment_order_timeout", str(DEFAULT_ORDER_EXPIRE_SECONDS)) + return int(val) if val.strip() else DEFAULT_ORDER_EXPIRE_SECONDS + except ValueError: + return DEFAULT_ORDER_EXPIRE_SECONDS + + +# --------------------------------------------------------------------------- +# Monkey-patch alipay-sdk-python WebUtils.do_post to fix bytes concatenation bug +# The SDK's error handling does: '...' + response.read() +# but response.read() returns bytes, causing TypeError on Python 3 +# --------------------------------------------------------------------------- +def _patch_alipay_webutils(): + try: + from alipay.aop.api.util import WebUtils + _original_do_post = WebUtils.do_post + + def _patched_do_post(url, query_string, headers, params, charset, timeout=30): + try: + return _original_do_post(url, query_string, headers, params, charset, timeout) + except TypeError as e: + if "can only concatenate str (not 'bytes') to str" in str(e): + # Decode bytes response to string and retry + import http.client as _http + from urllib.parse import urlparse as _urlparse + parsed = _urlparse(url) + conn = _http.HTTPSConnection(parsed.hostname) + conn.request("POST", parsed.path + "?" + query_string, params, headers) + resp = conn.getresponse() + body = resp.read().decode("utf-8", errors="replace") + raise RuntimeError(f"Alipay API error (status {resp.status}): {body}") from e + raise + + WebUtils.do_post = _patched_do_post + except ImportError: + pass + +_patch_alipay_webutils() + + +# --------------------------------------------------------------------------- +# Config helpers – read from system_configs table (admin panel) +# --------------------------------------------------------------------------- + + +async def _get_payment_configs(db: AsyncSession) -> dict[str, str]: + """Read all payment_* configs from the database, return as a dict.""" + result = await db.execute( + select(SystemConfig).where(SystemConfig.key.like("payment_%")) + ) + return {c.key: c.value for c in result.scalars().all()} + + +async def _check_and_expire_order(db: AsyncSession, order: PaymentOrder) -> bool: + """If a pending order has passed its expiry, mark it cancelled. + Returns True if the order was expired. + """ + if order.status != "pending": + return False + db_configs = await _get_payment_configs(db) + expire_seconds = _get_order_expire_seconds(db_configs) + expiry = order.created_at + timedelta(seconds=expire_seconds) + if datetime.now(order.created_at.tzinfo) >= expiry: + order.status = "cancelled" + await db.flush() + logger.info( + f"ORDER_EXPIRED order_no={order.order_no} user={order.user_id} " + f"amount={order.amount} created_at={order.created_at.isoformat()}" + ) + # Also call Alipay close API if it was an Alipay order + if order.payment_method == "alipay": + try: + await _close_alipay_order(db, order, db_configs) + except Exception as e: + logger.exception(f"Failed to close Alipay order {order.order_no}: {e}") + return True + return False + + +async def expire_all_pending_orders(db: AsyncSession) -> int: + """Background task: mark all expired pending orders as cancelled. + Returns the number of orders expired. + """ + db_configs = await _get_payment_configs(db) + expire_seconds = _get_order_expire_seconds(db_configs) + threshold = datetime.now() - timedelta(seconds=expire_seconds) + result = await db.execute( + select(PaymentOrder).where( + PaymentOrder.status == "pending", + PaymentOrder.created_at <= threshold, + ) + ) + orders = result.scalars().all() + expired_count = 0 + for o in orders: + o.status = "cancelled" + expired_count += 1 + logger.info( + f"ORDER_EXPIRED order_no={o.order_no} user={o.user_id} amount={o.amount}" + ) + # Also call Alipay close API if it was an Alipay order + if o.payment_method == "alipay": + try: + await _close_alipay_order(db, o, db_configs) + except Exception as e: + logger.exception(f"Failed to close Alipay order {o.order_no}: {e}") + if orders: + await db.flush() + return expired_count + + +def _is_mock_mode(db_configs: dict[str, str]) -> bool: + """Check if payment mock mode is enabled (from DB or env).""" + db_val = db_configs.get("payment_mock", "") + if db_val: + return db_val.lower() in ("true", "1", "yes") + return settings.PAYMENT_MOCK + + +# --------------------------------------------------------------------------- +# Alipay client (lazy singleton, recreated when config changes) +# --------------------------------------------------------------------------- +_alipay_client = None +_alipay_client_app_id = None + + +def _get_alipay_client(app_id: str, private_key: str, public_key: str, gateway: str = ""): + """Get or create an Alipay client. Recreated if app_id changes.""" + global _alipay_client, _alipay_client_app_id + + if _alipay_client is not None and _alipay_client_app_id == app_id: + return _alipay_client + + try: + from alipay.aop.api.AlipayClientConfig import AlipayClientConfig + from alipay.aop.api.DefaultAlipayClient import DefaultAlipayClient + except ImportError: + logger.error( + "alipay-sdk-python is not installed. " + "Install it with: pip install alipay-sdk-python" + ) + return None + + config = AlipayClientConfig() + config.server_url = gateway or "https://openapi.alipay.com/gateway.do" + config.app_id = app_id + config.app_private_key = private_key + config.alipay_public_key = public_key + config.sign_type = "RSA2" + config.charset = "utf-8" + + try: + _alipay_client = DefaultAlipayClient(config, logger) + _alipay_client_app_id = app_id + except Exception: + logger.exception("Failed to initialize Alipay client") + _alipay_client = None + _alipay_client_app_id = None + + return _alipay_client + + +# --------------------------------------------------------------------------- +# Create recharge order +# --------------------------------------------------------------------------- async def create_recharge_order( @@ -20,7 +250,28 @@ async def create_recharge_order( bonus_credits: float = 0.0, method: str = "wechat", ) -> PaymentOrder: - """Create a payment order. In mock mode, immediately completes payment.""" + """Create a payment order. + + Reads payment config from the database (admin panel). + Returns the order; for Alipay the ``qr_url`` attribute will be populated + with the scan-to-pay URL. + """ + # Read config from database first + db_configs = await _get_payment_configs(db) + mock_mode = _is_mock_mode(db_configs) + + # In real mode, validate that the payment method is enabled and configured + if not mock_mode: + enabled_key = f"payment_{method}_enabled" + if db_configs.get(enabled_key, "").lower() != "true": + raise ValueError("该支付方式未启用,请联系管理员") + if method == "alipay": + if not db_configs.get("payment_alipay_app_id") or not db_configs.get("payment_alipay_private_key"): + raise ValueError("支付宝支付未完成配置,请联系管理员") + elif method == "wechat": + if not db_configs.get("payment_wechat_mch_id") or not db_configs.get("payment_wechat_api_key"): + raise ValueError("微信支付未完成配置,请联系管理员") + total_credits = credits + bonus_credits order = PaymentOrder( id=generate_id(), @@ -33,8 +284,12 @@ async def create_recharge_order( ) db.add(order) await db.flush() + logger.info( + f"ORDER_CREATED order_no={order.order_no} user={user_id} " + f"amount={price} credits={total_credits} method={method} mock={mock_mode}" + ) - if settings.PAYMENT_MOCK: + if mock_mode: # Mock: immediately complete payment order.status = "paid" order.paid_at = datetime.now() @@ -52,59 +307,464 @@ async def create_recharge_order( else: # Real payment: delegate to WeChat or Alipay if method == "wechat": - _create_wechat_order(order) + _create_wechat_order(order, db_configs) elif method == "alipay": - _create_alipay_order(order) + qr_url = _create_alipay_order(order, db_configs) + if qr_url: + # Attach QR URL to the order instance (transient, not persisted) + order.qr_url = qr_url # type: ignore[attr-defined] + else: + # Precreate failed — do not leave a pending order that can never be paid + raise ValueError("支付宝预下单失败,请检查配置或稍后重试") return order -def _create_wechat_order(order: PaymentOrder) -> None: +# --------------------------------------------------------------------------- +# WeChat (stub) +# --------------------------------------------------------------------------- + + +def _create_wechat_order(order: PaymentOrder, db_configs: dict[str, str]) -> None: """Create a WeChat Pay order. Stub for real integration.""" - if not settings.WECHAT_MCH_ID or not settings.WECHAT_API_KEY: - logger.warning("WeChat payment config missing (WECHAT_MCH_ID / WECHAT_API_KEY)") + mch_id = db_configs.get("payment_wechat_mch_id", "") + api_key = db_configs.get("payment_wechat_api_key", "") + if not mch_id or not api_key: + logger.warning("WeChat payment config missing in database") return logger.info( - f"WeChat order created: mch_id={settings.WECHAT_MCH_ID}, " + f"WeChat order created: mch_id={mch_id}, " f"order_no={order.order_no}, amount={order.amount}" ) -def _create_alipay_order(order: PaymentOrder) -> None: - """Create an Alipay order. Stub for real integration.""" - if not settings.ALIPAY_APP_ID or not settings.ALIPAY_PRIVATE_KEY: - logger.warning("Alipay payment config missing (ALIPAY_APP_ID / ALIPAY_PRIVATE_KEY)") - return - logger.info( - f"Alipay order created: app_id={settings.ALIPAY_APP_ID}, " - f"order_no={order.order_no}, amount={order.amount}" - ) +# --------------------------------------------------------------------------- +# Alipay – trade.precreate (当面付 预下单) +# --------------------------------------------------------------------------- -async def verify_wechat_callback(data: dict) -> bool: +def _create_alipay_order(order: PaymentOrder, db_configs: dict[str, str]) -> str | None: + """Call Alipay ``trade.precreate`` to obtain a QR code URL. + + Reads all Alipay config from the database (admin panel). + Returns the ``qr_code`` URL on success, or ``None`` on failure. + """ + app_id = db_configs.get("payment_alipay_app_id", "") + private_key = db_configs.get("payment_alipay_private_key", "") + public_key = db_configs.get("payment_alipay_public_key", "") + gateway = db_configs.get("payment_alipay_gateway", "") + notify_url = db_configs.get("payment_alipay_notify_url", "") + + if not app_id or not private_key: + logger.warning("Alipay config missing in database (app_id / private_key)") + return None + + client = _get_alipay_client(app_id, private_key, public_key, gateway) + if client is None: + return None + + try: + from alipay.aop.api.domain.AlipayTradePrecreateModel import ( + AlipayTradePrecreateModel, + ) + from alipay.aop.api.request.AlipayTradePrecreateRequest import ( + AlipayTradePrecreateRequest, + ) + from alipay.aop.api.response.AlipayTradePrecreateResponse import ( + AlipayTradePrecreateResponse, + ) + + # 构造业务参数 + model = AlipayTradePrecreateModel() + model.out_trade_no = order.order_no + model.total_amount = f"{order.amount:.2f}" + model.subject = f"充值订单 {order.order_no}" + model.product_code = "QR_CODE_OFFLINE" + + body_parts = [] + if order.credits > 0: + body_parts.append(f"{order.credits}积分") + if body_parts: + model.body = " ".join(body_parts) + + # 构造请求 + request = AlipayTradePrecreateRequest(biz_model=model) + + # 设置 notify_url 在 request 上 + if notify_url: + try: + if hasattr(request, 'set_notify_url'): + request.set_notify_url(notify_url) + elif hasattr(request, 'notify_url'): + request.notify_url = notify_url + except Exception as e: + logger.warning(f"Failed to set notify_url: {e}") + + # 执行API调用 + response_content = client.execute(request) + if not response_content: + logger.error(f"Alipay precreate failed: empty response, order_no={order.order_no}") + return None + + # 解析响应结果 + response = AlipayTradePrecreateResponse() + response.parse_response_content(response_content) + + if response.is_success(): + qr_url = response.qr_code + return qr_url + else: + logger.error( + f"Alipay precreate failed: code={response.code}, " + f"msg={response.msg}, sub_code={response.sub_code}, " + f"sub_msg={response.sub_msg}, order_no={order.order_no}" + ) + return None + + except Exception as e: + # 处理 SDK 内部的 bytes/str 错误 + if "TypeError" in str(e) and ("bytes" in str(e) or "str" in str(e)): + logger.error( + f"Alipay SDK TypeError (bytes/str issue): order_no={order.order_no}, " + f"error={str(e)}" + ) + logger.exception(f"Alipay precreate exception: order_no={order.order_no}") + return None + + +# --------------------------------------------------------------------------- +# Alipay order close +# --------------------------------------------------------------------------- + + +async def _close_alipay_order(db: AsyncSession, order: PaymentOrder, db_configs: dict[str, str]) -> bool: + """Call Alipay trade.close API to close an unpaid order. + Returns True if the order was closed successfully. + """ + app_id = db_configs.get("payment_alipay_app_id", "") + private_key = db_configs.get("payment_alipay_private_key", "") + public_key = db_configs.get("payment_alipay_public_key", "") + gateway = db_configs.get("payment_alipay_gateway", "") + + client = _get_alipay_client(app_id, private_key, public_key, gateway) + if client is None: + return False + + mock_mode = _is_mock_mode(db_configs) + if mock_mode: + logger.info(f"Mock mode: skipping close_alipay_order for {order.order_no}") + return True + + try: + from alipay.aop.api.domain.AlipayTradeCloseModel import AlipayTradeCloseModel + from alipay.aop.api.request.AlipayTradeCloseRequest import AlipayTradeCloseRequest + from alipay.aop.api.response.AlipayTradeCloseResponse import AlipayTradeCloseResponse + + model = AlipayTradeCloseModel() + model.out_trade_no = order.order_no + + request = AlipayTradeCloseRequest(biz_model=model) + + response_content = client.execute(request) + if not response_content: + logger.error(f"Alipay close failed: empty response, order_no={order.order_no}") + return False + + response = AlipayTradeCloseResponse() + response.parse_response_content(response_content) + + if response.is_success(): + logger.info(f"Alipay order closed: order_no={order.order_no}") + return True + else: + logger.error( + f"Alipay close failed: code={response.code}, " + f"msg={response.msg}, sub_code={response.sub_code}, " + f"sub_msg={response.sub_msg}, order_no={order.order_no}" + ) + return False + + except Exception as e: + if "TypeError" in str(e) and ("bytes" in str(e) or "str" in str(e)): + logger.error( + f"Alipay SDK TypeError (bytes/str issue) during close: order_no={order.order_no}, " + f"error={str(e)}" + ) + logger.exception(f"Alipay close exception: order_no={order.order_no}") + return False + + +# --------------------------------------------------------------------------- +# Alipay order query +# --------------------------------------------------------------------------- + + +async def _query_alipay_order(db: AsyncSession, order: PaymentOrder, db_configs: dict[str, str]) -> dict | None: + """Call Alipay trade.query API to check order status. + Returns the response data if successful, None otherwise. + """ + app_id = db_configs.get("payment_alipay_app_id", "") + private_key = db_configs.get("payment_alipay_private_key", "") + public_key = db_configs.get("payment_alipay_public_key", "") + gateway = db_configs.get("payment_alipay_gateway", "") + + client = _get_alipay_client(app_id, private_key, public_key, gateway) + if client is None: + return None + + mock_mode = _is_mock_mode(db_configs) + if mock_mode: + logger.info(f"Mock mode: skipping query_alipay_order for {order.order_no}") + return {"trade_status": "TRADE_FINISHED"} + + try: + from alipay.aop.api.domain.AlipayTradeQueryModel import AlipayTradeQueryModel + from alipay.aop.api.request.AlipayTradeQueryRequest import AlipayTradeQueryRequest + from alipay.aop.api.response.AlipayTradeQueryResponse import AlipayTradeQueryResponse + + model = AlipayTradeQueryModel() + model.out_trade_no = order.order_no + + request = AlipayTradeQueryRequest(biz_model=model) + + response_content = client.execute(request) + if not response_content: + logger.error(f"Alipay query failed: empty response, order_no={order.order_no}") + return None + + response = AlipayTradeQueryResponse() + response.parse_response_content(response_content) + + if response.is_success(): + logger.info(f"Alipay query succeeded: order_no={order.order_no}, trade_status={response.trade_status}") + return { + "trade_no": response.trade_no, + "trade_status": response.trade_status, + "total_amount": response.total_amount, + "receipt_amount": response.receipt_amount, + } + else: + logger.error( + f"Alipay query failed: code={response.code}, " + f"msg={response.msg}, sub_code={response.sub_code}, " + f"sub_msg={response.sub_msg}, order_no={order.order_no}" + ) + return None + + except Exception as e: + if "TypeError" in str(e) and ("bytes" in str(e) or "str" in str(e)): + logger.error( + f"Alipay SDK TypeError (bytes/str issue) during query: order_no={order.order_no}, " + f"error={str(e)}" + ) + logger.exception(f"Alipay query exception: order_no={order.order_no}") + return None + + +async def sync_pending_orders(db: AsyncSession) -> int: + """Check pending orders via Alipay query and update status. + Returns the number of orders updated. + """ + result = await db.execute( + select(PaymentOrder).where( + PaymentOrder.status == "pending", + ) + ) + orders = result.scalars().all() + updated_count = 0 + + db_configs = await _get_payment_configs(db) + + for order in orders: + if order.payment_method != "alipay": + continue + + try: + data = await _query_alipay_order(db, order, db_configs) + if data: + trade_status = data.get("trade_status") + if trade_status in ("TRADE_SUCCESS", "TRADE_FINISHED"): + # Order was paid but we missed the callback + trade_no = data.get("trade_no", "") + await process_payment_success_by_order_no(db, order.order_no, trade_no) + updated_count += 1 + elif trade_status in ("TRADE_CLOSED", "TRADE_CANCELLED"): + # Order was closed on Alipay side + order.status = "cancelled" + await db.flush() + updated_count += 1 + except Exception as e: + logger.exception(f"Failed to sync order {order.order_no}: {e}") + + if updated_count > 0: + await db.flush() + return updated_count + + +# --------------------------------------------------------------------------- +# Alipay callback verification +# --------------------------------------------------------------------------- + + +async def verify_alipay_callback(data: dict, db: AsyncSession) -> bool: + """Verify Alipay payment callback (async notify) signature. + + Reads the Alipay public key from the database and uses RSA2 verification. + """ + db_configs = await _get_payment_configs(db) + mock_mode = _is_mock_mode(db_configs) + if mock_mode: + logger.info("Mock mode enabled, skipping Alipay callback verification") + return True + + public_key = db_configs.get("payment_alipay_public_key", "") + if not public_key: + logger.warning("ALIPAY_PUBLIC_KEY not found in database, cannot verify callback") + return False + + try: + sign = data.get("sign") + if not sign: + logger.warning("Alipay callback missing 'sign' field") + return False + + sign_type = data.get("sign_type", "RSA2") + + # Build verification params (exclude sign and sign_type) + verify_data = { + k: v for k, v in data.items() + if k not in ("sign", "sign_type") and v is not None and v != "" + } + + # Generate sign content: sorted keys, key=value format + sign_content = "&".join( + f"{k}={v}" for k, v in sorted(verify_data.items()) + ) + + # logger.info(f"Verifying Alipay callback sign_content: {sign_content[:100]}...") + # logger.info(f"Sign type: {sign_type}") + + # 实现 RSA2 签名验证 + is_valid = _verify_alipay_sign(public_key, sign_content, sign, sign_type) + + if not is_valid: + logger.warning("Alipay callback signature verification FAILED") + else: + logger.info("Alipay callback signature verification SUCCESS") + + return is_valid + + except Exception: + logger.exception("Alipay callback verification error") + return False + + +def _verify_alipay_sign(public_key: str, sign_content: str, sign: str, sign_type: str = "RSA2") -> bool: + """Verify Alipay RSA/RSA2 signature. + + Args: + public_key: Alipay public key (PEM format, with or without headers) + sign_content: Original content to verify + sign: Base64 encoded signature + sign_type: "RSA" (SHA1) or "RSA2" (SHA256) + + Returns: + True if signature is valid + """ + try: + import base64 + from hashlib import sha1, sha256 + + # 处理公钥,确保有正确的格式 + pub_key = public_key.strip() + if not pub_key.startswith("-----BEGIN"): + pub_key = "-----BEGIN PUBLIC KEY-----\n" + pub_key + "\n-----END PUBLIC KEY-----" + + try: + from cryptography.hazmat.primitives import hashes + from cryptography.hazmat.primitives.asymmetric import padding + from cryptography.hazmat.primitives import serialization + from cryptography.hazmat.backends import default_backend + + # 加载公钥 + public_key_obj = serialization.load_pem_public_key( + pub_key.encode("utf-8"), + backend=default_backend() + ) + + # 选择哈希算法 + if sign_type == "RSA2": + hash_alg = hashes.SHA256() + else: + hash_alg = hashes.SHA1() + + # 验证签名 + public_key_obj.verify( + base64.b64decode(sign), + sign_content.encode("utf-8"), + padding.PKCS1v15(), + hash_alg + ) + return True + + except ImportError: + # 如果没有 cryptography,尝试使用 rsa 库 + try: + import rsa + + # 加载公钥 + pub_key_obj = rsa.PublicKey.load_pkcs1_openssl_pem(pub_key.encode("utf-8")) + + # 选择哈希算法 + if sign_type == "RSA2": + hash_func = 'SHA-256' + else: + hash_func = 'SHA-1' + + # 验证签名 + rsa.verify( + sign_content.encode("utf-8"), + base64.b64decode(sign), + pub_key_obj, + hash_func + ) + return True + + except ImportError: + logger.error("Neither cryptography nor rsa library installed, cannot verify signature") + # 如果没有任何加密库,在生产环境应该返回 False,但这里我们记录警告并继续 + logger.warning("Skipping signature verification due to missing crypto libraries") + return False + + except Exception as e: + logger.exception(f"Signature verification failed: {e}") + return False + + +# --------------------------------------------------------------------------- +# WeChat callback verification (stub) +# --------------------------------------------------------------------------- + + +async def verify_wechat_callback(data: dict, db: AsyncSession) -> bool: """Verify WeChat payment callback signature.""" - if settings.PAYMENT_MOCK: + db_configs = await _get_payment_configs(db) + mock_mode = _is_mock_mode(db_configs) + if mock_mode: return True - # Real verification would use WECHAT_API_KEY to verify signature logger.info("WeChat callback verification (real mode not implemented)") return True -async def verify_alipay_callback(data: dict) -> bool: - """Verify Alipay payment callback signature.""" - if settings.PAYMENT_MOCK: - return True - # Real verification would use ALIPAY_PUBLIC_KEY to verify signature - logger.info("Alipay callback verification (real mode not implemented)") - return True +# --------------------------------------------------------------------------- +# Process successful payment +# --------------------------------------------------------------------------- async def process_payment_success(db: AsyncSession, order_id: str): """Process successful payment: update order and add credits.""" - from sqlalchemy import select - result = await db.execute( - select(PaymentOrder).where(PaymentOrder.id == order_id).limit(1) + select(PaymentOrder).where(PaymentOrder.id == order_id).with_for_update().limit(1) ) order = result.scalar_one_or_none() if not order or order.status != "pending": @@ -119,4 +779,201 @@ async def process_payment_success(db: AsyncSession, order_id: str): f"充值成功({order.credits}积分)", related_id=order.id, ) - await db.flush() + await db.commit() + + +async def process_payment_success_by_order_no( + db: AsyncSession, + order_no: str, + trade_no: str = "", + total_amount: float | None = None +): + """Process successful payment by order_no (used by Alipay/WeChat callbacks). + + Args: + db: async database session + order_no: the merchant order number (out_trade_no) + trade_no: the Alipay trade number (trade_no), optional + total_amount: the payment amount from the gateway, for consistency check + """ + result = await db.execute( + select(PaymentOrder).where(PaymentOrder.order_no == order_no).with_for_update().limit(1) + ) + order = result.scalar_one_or_none() + + if not order: + logger.info(f"Order {order_no} not found, skipping") + return + + if order.status == "paid": + logger.info(f"Order {order_no} already processed, skipping") + return + + if order.status != "pending": + logger.info(f"Order {order_no} is in {order.status} state, cannot process") + return + + # 金额一致性校验 + if total_amount is not None and abs(total_amount - order.amount) > 0.01: + logger.error( + f"Amount mismatch: order amount {order.amount}, gateway amount {total_amount}" + ) + return + + # 幂等性检查:如果trade_no已存在且相同,则跳过 + if trade_no and order.trade_no and order.trade_no == trade_no: + logger.info(f"Trade no {trade_no} already processed, skipping") + return + + order.status = "paid" + order.paid_at = datetime.now() + if trade_no: + order.trade_no = trade_no + + await add_credits( + db, + order.user_id, + order.credits, + f"充值成功({order.credits}积分)", + related_id=order.id, + ) + await db.commit() + logger.info( + f"PAYMENT_SUCCESS order_no={order_no} user={order.user_id} " + f"amount={order.amount} credits={order.credits} trade_no={trade_no}" + ) + + +async def process_refund( + db: AsyncSession, + order_no: str, + refund_amount: float | None = None, + refund_reason: str = "管理员退款" +) -> dict: + """Process a refund for a paid order. + + Args: + db: async database session + order_no: merchant order number + refund_amount: amount to refund (defaults to full order amount) + refund_reason: reason for refund + + Returns: + dict with refund result + """ + result = await db.execute( + select(PaymentOrder).where(PaymentOrder.order_no == order_no).with_for_update().limit(1) + ) + order = result.scalar_one_or_none() + + if not order: + return {"success": False, "message": "订单不存在"} + + if order.status != "paid": + return {"success": False, "message": f"订单状态为{order.status},无法退款"} + + if order.refunded_at is not None: + return {"success": False, "message": "订单已退款"} + + refund_amount = refund_amount or order.amount + + # 金额校验 + if refund_amount > order.amount: + return {"success": False, "message": "退款金额超过订单金额"} + + # 如果是支付宝订单,调用支付宝退款API + db_configs = await _get_payment_configs(db) + if order.payment_method == "alipay": + refund_result = await _refund_alipay_order( + db, order, refund_amount, refund_reason, db_configs + ) + if not refund_result.get("success"): + return refund_result + + # 扣除积分 + try: + await deduct_credits( + db, + order.user_id, + order.credits, + refund_reason, + related_id=order.id, + ) + except Exception as e: + logger.exception(f"Failed to deduct credits for refund: {e}") + return {"success": False, "message": "积分扣除失败"} + + # 更新订单状态 + order.status = "refunded" + order.refund_amount = refund_amount + order.refunded_at = datetime.now() + if order.payment_method == "alipay": + order.refund_trade_no = db_configs.get("refund_trade_no", "") + + await db.commit() + logger.info( + f"REFUND_SUCCESS order_no={order_no} user={order.user_id} " + f"refund_amount={refund_amount}" + ) + return {"success": True, "message": "退款成功"} + + +async def _refund_alipay_order( + db: AsyncSession, + order: PaymentOrder, + refund_amount: float, + refund_reason: str, + db_configs: dict[str, str] +) -> dict: + """Call Alipay refund API.""" + app_id = db_configs.get("payment_alipay_app_id", "") + private_key = db_configs.get("payment_alipay_private_key", "") + public_key = db_configs.get("payment_alipay_public_key", "") + gateway = db_configs.get("payment_alipay_gateway", "") + + client = _get_alipay_client(app_id, private_key, public_key, gateway) + if client is None: + return {"success": False, "message": "支付宝客户端初始化失败"} + + mock_mode = _is_mock_mode(db_configs) + if mock_mode: + logger.info(f"Mock mode: skipping alipay refund for {order.order_no}") + return {"success": True} + + try: + from alipay.aop.api.domain.AlipayTradeRefundModel import AlipayTradeRefundModel + from alipay.aop.api.request.AlipayTradeRefundRequest import AlipayTradeRefundRequest + from alipay.aop.api.response.AlipayTradeRefundResponse import AlipayTradeRefundResponse + + model = AlipayTradeRefundModel() + model.out_trade_no = order.order_no + model.refund_amount = f"{refund_amount:.2f}" + model.refund_reason = refund_reason + model.out_request_no = f"{order.order_no}_refund_{int(datetime.now().timestamp())}" + + request = AlipayTradeRefundRequest(biz_model=model) + response_content = client.execute(request) + + if not response_content: + logger.error(f"Alipay refund failed: empty response, order_no={order.order_no}") + return {"success": False, "message": "支付宝退款响应为空"} + + response = AlipayTradeRefundResponse() + response.parse_response_content(response_content) + + if response.is_success(): + logger.info(f"Alipay refund succeeded: order_no={order.order_no}") + return {"success": True, "trade_no": response.trade_no} + else: + logger.error( + f"Alipay refund failed: code={response.code}, " + f"msg={response.msg}, sub_code={response.sub_code}, " + f"sub_msg={response.sub_msg}, order_no={order.order_no}" + ) + return { + "success": False, + "message": f"支付宝退款失败: {response.sub_msg or response.msg}" + } + except Exception as e: + logger.exception(f"Alipay refund exception: order_no={order.order_no}, {e}") + return {"success": False, "message": f"支付宝退款异常: {str(e)}"} diff --git a/video-gen-api/app/services/user_oauth_app_service.py b/video-gen-api/app/services/user_oauth_app_service.py new file mode 100644 index 00000000..5841de08 --- /dev/null +++ b/video-gen-api/app/services/user_oauth_app_service.py @@ -0,0 +1,142 @@ +from sqlalchemy import func, select +from sqlalchemy.ext.asyncio import AsyncSession + +from app.models.user_oauth_app import UserOAuthApp +from app.utils.id_gen import generate_id + + +async def list_user_oauth_apps( + db: AsyncSession, + page: int = 1, + page_size: int = 20, + open_type: int | None = None, + status: int | None = None, + create_by: str | None = None, + app_id: str | None = None, +) -> dict: + query = select(UserOAuthApp).where(UserOAuthApp.deleted_at.is_(None)).order_by(UserOAuthApp.created_at.desc()) + + if open_type is not None: + query = query.where(UserOAuthApp.open_type == open_type) + + if status is not None: + query = query.where(UserOAuthApp.status == status) + + if create_by is not None: + query = query.where(UserOAuthApp.create_by == create_by) + + if app_id is not None: + query = query.where(UserOAuthApp.app_id.like(f"%{app_id}%")) + + total_result = await db.execute(select(func.count(UserOAuthApp.id)).where(UserOAuthApp.deleted_at.is_(None))) + total = total_result.scalar() or 0 + + result = await db.execute(query.offset((page - 1) * page_size).limit(page_size)) + items = result.scalars().all() + + return { + "total": total, + "page": page, + "page_size": page_size, + "items": items, + } + + +async def get_user_oauth_app_by_id(db: AsyncSession, id: str) -> UserOAuthApp | None: + result = await db.execute( + select(UserOAuthApp).where(UserOAuthApp.id == id, UserOAuthApp.deleted_at.is_(None)).limit(1) + ) + return result.scalar_one_or_none() + + +async def get_user_oauth_app_by_app_id(db: AsyncSession, app_id: str) -> UserOAuthApp | None: + result = await db.execute( + select(UserOAuthApp).where(UserOAuthApp.app_id == app_id, UserOAuthApp.deleted_at.is_(None)).limit(1) + ) + return result.scalar_one_or_none() + + +async def create_user_oauth_app( + db: AsyncSession, + app_id: str, + secret: str, + open_type: int, + create_by: str | None = None, + count: int = 100, + auth_url: str | None = None, + company: str | None = None, +) -> UserOAuthApp: + existing = await get_user_oauth_app_by_app_id(db, app_id) + if existing: + raise ValueError("应用id已存在") + + app = UserOAuthApp( + id=generate_id(), + app_id=app_id, + secret=secret, + open_type=open_type, + count=count, + auth_url=auth_url, + company=company, + create_by=create_by, + ) + db.add(app) + await db.flush() + await db.commit() + await db.refresh(app) + return app + + +async def update_user_oauth_app( + db: AsyncSession, + id: str, + secret: str | None = None, + open_type: int | None = None, + status: int | None = None, + count: int | None = None, + auth_url: str | None = None, + company: str | None = None, + create_by: str | None = None, +) -> UserOAuthApp | None: + app = await get_user_oauth_app_by_id(db, id) + if not app: + return None + + if secret is not None: + app.secret = secret + if open_type is not None: + app.open_type = open_type + if status is not None: + app.status = status + if count is not None: + app.count = count + if auth_url is not None: + app.auth_url = auth_url + if company is not None: + app.company = company + if create_by is not None: + app.create_by = create_by + + await db.flush() + await db.commit() + await db.refresh(app) + return app + + +async def delete_user_oauth_app(db: AsyncSession, id: str, create_by: str | None = None) -> bool: + app = await get_user_oauth_app_by_id(db, id) + if not app: + return False + + app.deleted_at = func.now() + if create_by is not None: + app.create_by = create_by + await db.flush() + return True + + +async def get_apps_by_open_type(db: AsyncSession, open_type: int) -> list[UserOAuthApp]: + result = await db.execute( + select(UserOAuthApp).where(UserOAuthApp.open_type == open_type, UserOAuthApp.deleted_at.is_(None)) + ) + return result.scalars().all() \ No newline at end of file diff --git a/video-gen-api/app/services/user_oauth_service.py b/video-gen-api/app/services/user_oauth_service.py new file mode 100644 index 00000000..fefcb541 --- /dev/null +++ b/video-gen-api/app/services/user_oauth_service.py @@ -0,0 +1,353 @@ +import random +from datetime import datetime + +import httpx +from sqlalchemy import select, func +from sqlalchemy.ext.asyncio import AsyncSession + +from app.config import settings +from app.models.user_oauth import UserOAuth +from app.utils.id_gen import generate_id + + +OAUTH_TYPE_CONFIG = { + 1: {"port_type": 1, "name": "千川", "app_type": "juliang_qianchuan"}, + 2: {"port_type": 1, "name": "广告", "app_type": "juliang_ad"}, + 3: {"port_type": 1, "name": "本地推", "app_type": "juliang_ad"}, + 4: {"port_type": 1, "name": "星图", "app_type": "juliang_ad"}, + 5: {"port_type": 2, "name": "快手代理商", "app_type": "kuaishou"}, + 6: {"port_type": 3, "name": "巨量星图", "app_type": "juliang_ad"}, + 7: {"port_type": 4, "name": "巨量服务单", "app_type": "juliang_ad"}, + 8: {"port_type": 4, "name": "腾讯服务单", "app_type": "tencent"}, + 9: {"port_type": 5, "name": "腾讯营销K2", "app_type": "tencent"}, + 10: {"port_type": 5, "name": "腾讯营销K3", "app_type": "tencent"}, +} + + +async def get_available_app(app_type: str, db: AsyncSession) -> dict: + if app_type == "juliang_ad": + apps = settings.JULIANG_AD_APPS + elif app_type == "juliang_qianchuan": + apps = settings.JULIANG_QIANCHUAN_APPS + elif app_type == "kuaishou": + apps = settings.KUAISHOU_APPS + elif app_type == "tencent": + apps = settings.TENCENT_APPS + else: + raise ValueError(f"不支持的应用类型: {app_type}") + + if not apps: + raise ValueError(f"{app_type}未配置应用") + + available_apps = [] + for app in apps: + app_id = app.get("app_id") + if not app_id: + continue + + result = await db.execute( + select(func.count(UserOAuth.id)).where(UserOAuth.appid == app_id) + ) + count = result.scalar() or 0 + + if count < 5000: + available_apps.append(app) + + if not available_apps: + raise ValueError("所有应用授权已超过最大数量") + + return random.choice(available_apps) + + +async def build_oauth_url(oauth_type: int, user_id: str) -> str: + if oauth_type == 1: + return await _build_juliang_oauth_url(oauth_type, user_id, app_type) + elif oauth_type == 2: + return await _build_kuaishou_oauth_url(oauth_type, user_id) + elif oauth_type == 3: + return await _build_tencent_oauth_url(oauth_type, user_id) + elif oauth_type == 4: + return await _build_tencent_oauth_url(oauth_type, user_id) + else: + raise ValueError(f"不支持的应用类型: {app_type}") + + +async def _build_juliang_oauth_url(oauth_type: int, user_id: str, app_type: str) -> str: + async with AsyncSession() as db: + app = await get_available_app(app_type, db) + app_id = app.get("app_id") + + redirect_uri = "https://open.oceanengine.com/audit/oauth.html" + rid = "ktm0cl7napb" + if oauth_type == 1: + redirect_uri = "https://qianchuan.jinritemai.com/openapi/qc/audit/oauth.html" + rid = "vr7kclvmvs9" + + params = { + "app_id": app_id, + "state": f"{oauth_type}:{user_id}:{app_id}:{app_type}", + "material_auth": 1, + "rid": rid, + } + query_string = "&".join(f"{k}={v}" for k, v in params.items()) + return f"{redirect_uri}?{query_string}" + + +async def _build_kuaishou_oauth_url(oauth_type: int, user_id: str) -> str: + async with AsyncSession() as db: + app = await get_available_app("kuaishou", db) + app_id = app.get("app_id") + + redirect_uri = f"{settings.BASE_URL}/api/user-oauth/callback" + params = { + "app_id": app_id, + "redirect_uri": redirect_uri, + "response_type": "code", + "scope": "basic", + "state": f"{oauth_type}:{user_id}:{app_id}:kuaishou", + } + query_string = "&".join(f"{k}={v}" for k, v in params.items()) + return f"https://open.kuaishou.com/oauth2/authorize?{query_string}" + + +async def _build_tencent_oauth_url(oauth_type: int, user_id: str) -> str: + async with AsyncSession() as db: + app = await get_available_app("tencent", db) + app_id = app.get("app_id") + + redirect_uri = f"{settings.BASE_URL}/api/user-oauth/callback" + params = { + "app_id": app_id, + "redirect_uri": redirect_uri, + "response_type": "code", + "scope": "get_user_info", + "state": f"{oauth_type}:{user_id}:{app_id}:tencent", + } + query_string = "&".join(f"{k}={v}" for k, v in params.items()) + return f"https://api.e.qq.com/oauth/authorize?{query_string}" + + +async def get_token_by_type(code: str, oauth_type: int, app_id: str, app_type: str) -> dict: + if app_type in ("juliang_ad", "juliang_qianchuan"): + return await get_juliang_token(code, oauth_type, app_id, app_type) + elif app_type == "kuaishou": + return await get_kuaishou_token(code, oauth_type, app_id) + elif app_type == "tencent": + return await get_tencent_token(code, oauth_type, app_id) + else: + raise ValueError(f"不支持的应用类型: {app_type}") + + +async def get_juliang_token(code: str, oauth_type: int, app_id: str, app_type: str) -> dict: + url = "https://api.oceanengine.com/open_api/oauth2/access_token/" + + if app_type == "juliang_ad": + apps = settings.JULIANG_AD_APPS + else: + apps = settings.JULIANG_QIANCHUAN_APPS + + app = next((a for a in apps if a.get("app_id") == app_id), None) + if not app: + raise ValueError("应用配置不存在") + + async with httpx.AsyncClient() as client: + response = await client.post( + url, + data={ + "app_id": app_id, + "secret": app.get("secret"), + "auth_code": code, + }, + ) + response.raise_for_status() + content = response.json() + if content.get("code") != 0: + raise ValueError(content.get("message", "获取token失败")) + return content.get("data", {}) + + +async def get_kuaishou_token(code: str, oauth_type: int, app_id: str) -> dict: + url = "https://open.kuaishou.com/oauth2/token" + redirect_uri = f"{settings.BASE_URL}/api/user-oauth/callback" + + app = next((a for a in settings.KUAISHOU_APPS if a.get("app_id") == app_id), None) + if not app: + raise ValueError("应用配置不存在") + + async with httpx.AsyncClient() as client: + response = await client.post( + url, + data={ + "app_id": app_id, + "secret": app.get("secret"), + "code": code, + "grant_type": "authorization_code", + "redirect_uri": redirect_uri, + }, + ) + response.raise_for_status() + return response.json() + + +async def get_tencent_token(code: str, oauth_type: int, app_id: str) -> dict: + url = "https://api.e.qq.com/oauth/token" + redirect_uri = f"{settings.BASE_URL}/api/user-oauth/callback" + + app = next((a for a in settings.TENCENT_APPS if a.get("app_id") == app_id), None) + if not app: + raise ValueError("应用配置不存在") + + async with httpx.AsyncClient() as client: + response = await client.post( + url, + data={ + "app_id": app_id, + "secret": app.get("secret"), + "code": code, + "grant_type": "authorization_code", + "redirect_uri": redirect_uri, + }, + ) + response.raise_for_status() + return response.json() + + +async def get_account_info_by_type(token: dict, oauth_type: int, app_type: str) -> dict: + if app_type in ("juliang_ad", "juliang_qianchuan"): + return await _get_juliang_account_info(token) + elif app_type == "kuaishou": + return await _get_kuaishou_account_info(token) + elif app_type == "tencent": + return await _get_tencent_account_info(token) + else: + raise ValueError(f"不支持的应用类型: {app_type}") + + +async def _get_juliang_account_info(token: dict) -> dict: + access_token = token.get("access_token") + url = "https://ad.oceanengine.com/openapi/oauth/user/info/" + + async with httpx.AsyncClient() as client: + response = await client.get( + url, + headers={"Access-Token": access_token}, + ) + response.raise_for_status() + data = response.json() + if data.get("code") != 0: + raise ValueError(data.get("message", "获取账户信息失败")) + data = data.get("data", {}) + return { + "account_id": data.get("advertiser_id", data.get("account_id", "")), + "account_name": data.get("advertiser_name", data.get("account_name", "")), + "account_role": data.get("role", ""), + "account_username": data.get("username", ""), + } + + +async def _get_kuaishou_account_info(token: dict) -> dict: + access_token = token.get("access_token") + url = "https://open.kuaishou.com/api/user/info" + + async with httpx.AsyncClient() as client: + response = await client.get( + url, + headers={"Authorization": f"Bearer {access_token}"}, + ) + response.raise_for_status() + data = response.json() + return { + "account_id": data.get("account_id", ""), + "account_name": data.get("account_name", ""), + "account_role": data.get("role", ""), + "account_username": data.get("username", ""), + } + + +async def _get_tencent_account_info(token: dict) -> dict: + access_token = token.get("access_token") + url = "https://api.e.qq.com/user/info" + + async with httpx.AsyncClient() as client: + response = await client.get( + url, + headers={"Authorization": f"Bearer {access_token}"}, + ) + response.raise_for_status() + data = response.json() + return { + "account_id": data.get("account_id", ""), + "account_name": data.get("account_name", ""), + "account_role": data.get("role", ""), + "account_username": data.get("username", ""), + } + + +async def save_oauth_token( + db: AsyncSession, + user_id: str, + oauth_type: int, + token: dict, + account_info: dict, + app_id: str, +) -> UserOAuth: + config = OAUTH_TYPE_CONFIG.get(oauth_type) + if not config: + raise ValueError(f"不支持的oauth_type: {oauth_type}") + + access_token = token.get("access_token") + access_token_expired = token.get("expires_in") + refresh_token = token.get("refresh_token") + refresh_token_expired = token.get("refresh_token_expires_in") + + expires_at = None + if access_token_expired: + expires_at = datetime.now().timestamp() + int(access_token_expired) + expires_at = datetime.fromtimestamp(expires_at) + + refresh_expires_at = None + if refresh_token_expired: + refresh_expires_at = datetime.now().timestamp() + int(refresh_token_expired) + refresh_expires_at = datetime.fromtimestamp(refresh_expires_at) + + existing = await db.execute( + select(UserOAuth).where( + UserOAuth.user_id == user_id, + UserOAuth.open_type == oauth_type, + UserOAuth.account_id == account_info.get("account_id", ""), + ).limit(1) + ) + existing_oauth = existing.scalar_one_or_none() + + if existing_oauth: + existing_oauth.access_token = access_token + existing_oauth.access_token_expired = expires_at + existing_oauth.refresh_token = refresh_token + existing_oauth.refresh_token_expired = refresh_expires_at + existing_oauth.account_name = account_info.get("account_name", "") + existing_oauth.account_role = account_info.get("account_role", "") + existing_oauth.account_username = account_info.get("account_username", "") + existing_oauth.appid = app_id + await db.flush() + return existing_oauth + + user_oauth = UserOAuth( + id=generate_id(), + account_id=account_info.get("account_id", ""), + account_name=account_info.get("account_name", ""), + account_role=account_info.get("account_role", ""), + account_username=account_info.get("account_username", ""), + user_id=user_id, + open_type=oauth_type, + port_type=config["port_type"], + appid=app_id, + access_token=access_token, + access_token_expired=expires_at, + refresh_token=refresh_token, + refresh_token_expired=refresh_expires_at, + material_auth_status=True, + ) + + db.add(user_oauth) + await db.flush() + return user_oauth \ No newline at end of file diff --git a/video-gen-api/app/tasks/__init__.py b/video-gen-api/app/tasks/__init__.py index 8ff0cbff..1685b2f7 100644 --- a/video-gen-api/app/tasks/__init__.py +++ b/video-gen-api/app/tasks/__init__.py @@ -10,6 +10,7 @@ try: generation_poll_tasks, generation_download_tasks, generation_recovery_tasks, + hot_opening_replicate_tasks ) except Exception: pass diff --git a/video-gen-api/app/tasks/celery_app.py b/video-gen-api/app/tasks/celery_app.py index 4c6c59be..ed4b567d 100644 --- a/video-gen-api/app/tasks/celery_app.py +++ b/video-gen-api/app/tasks/celery_app.py @@ -49,6 +49,8 @@ if broker_url: "generation.chatapi_create_generation_task": {"queue": "gen_chatapi_create"}, "generation.poll_generation_task": {"queue": "gen_provider_poll"}, "generation.download_generation_result_task": {"queue": "gen_result_download"}, + "hot_opening.start_image_prompt_optimize": {"queue": "gen_chatapi_create"}, + "hot_opening.start_video_prompt_optimize": {"queue": "gen_chatapi_create"}, "generation.recover_download_tasks_once": {"queue": "gen_result_download"}, "generation.recover_generation_tasks_once": {"queue": "gen_result_download"}, "app.tasks.cleanup.*": {"queue": "default"}, diff --git a/video-gen-api/app/tasks/generation_create_tasks.py b/video-gen-api/app/tasks/generation_create_tasks.py index ea83e2f7..29a5d44e 100644 --- a/video-gen-api/app/tasks/generation_create_tasks.py +++ b/video-gen-api/app/tasks/generation_create_tasks.py @@ -13,6 +13,8 @@ from app.services.generation_refund_service import mark_chat_generation_task_fai from app.services.generation_provider_service import create_provider_task from app.tasks.celery_app import celery_app +ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate"} + def _get_first_value(obj: Any, *field_names: str) -> Optional[Any]: """ @@ -63,6 +65,14 @@ def _build_optimized_prompt_by_params(task: ChatGenerationTask) -> str: base_prompt = original_prompt.rstrip(",,。;; \n\t") gen_type = (_to_clean_str(getattr(task, "gen_type", None)) or "").lower() + generation_mode = _to_clean_str(getattr(task, "generation_mode", None)) or "" + + # 爆款开头复刻第5步的视频生成,original_prompt 已经是视频提词 JSON schema。 + # 不能再追加“时长/比例/分辨率”中文参数,否则会污染 schema。 + if generation_mode == "hot_opening_replicate" and gen_type == "video": + stripped = base_prompt.strip() + if stripped.startswith("{") or stripped.startswith("["): + return base_prompt duration = _get_first_value(task, "duration") aspect_ratio = _get_first_value(task, "aspect_ratio") @@ -113,7 +123,7 @@ async def _run(task_id: str): ).with_for_update().limit(1)) task = result.scalar_one_or_none() - if not task or task.generation_mode != "chatapi_async": + if not task or task.generation_mode not in ALLOWED_GENERATION_MODES: return if task.status != "generating": @@ -128,6 +138,9 @@ async def _run(task_id: str): ) await db.commit() await log_task_event(task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout") + from app.services.generation_module_hook_service import notify_chat_generation_task_finished + await notify_chat_generation_task_finished(db, task) + await db.commit() return if task.pipeline_stage not in ("queued", "preparing", "creating_provider_task"): @@ -253,6 +266,9 @@ async def _run(task_id: str): ) await db.commit() await log_task_event(task, event_type="TASK_FAILED", message=task.error_message) + from app.services.generation_module_hook_service import notify_chat_generation_task_finished + await notify_chat_generation_task_finished(db, task) + await db.commit() if celery_app: diff --git a/video-gen-api/app/tasks/generation_download_tasks.py b/video-gen-api/app/tasks/generation_download_tasks.py index 61543671..3a27f659 100644 --- a/video-gen-api/app/tasks/generation_download_tasks.py +++ b/video-gen-api/app/tasks/generation_download_tasks.py @@ -21,6 +21,8 @@ from app.services.generation_refund_service import mark_chat_generation_task_fai from app.services.resource_accounting_service import record_chat_task_generated_resource from app.tasks.celery_app import celery_app +ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate"} + DOWNLOAD_QUEUE = "gen_result_download" DOWNLOAD_STAGE_QUEUED = "download_queued" DOWNLOAD_STAGE_DOWNLOADING = "downloading" @@ -108,7 +110,7 @@ async def enqueue_download_task( countdown: int | None = None, ) -> str | None: """统一投递图片/视频下载任务,并同步 DB + Redis active 注册表。""" - if not task or task.generation_mode != "chatapi_async": + if not task or task.generation_mode not in ALLOWED_GENERATION_MODES: return None if task.status != "generating": return None @@ -159,7 +161,7 @@ async def _reload_task(db: AsyncSession, task_id: str) -> ChatGenerationTask | N async def _claim_download_lease(db: AsyncSession, task: ChatGenerationTask) -> bool: now = _now() - if not task or task.generation_mode != "chatapi_async": + if not task or task.generation_mode not in ALLOWED_GENERATION_MODES: return False if task.status != "generating": return False @@ -310,6 +312,10 @@ async def _run(task_id: str): await db.commit() await remove_download_active(task.id) + from app.services.generation_module_hook_service import notify_chat_generation_task_finished + await notify_chat_generation_task_finished(db, task) + await db.commit() + await log_task_event( task, event_type="DOWNLOAD_SUCCESS", @@ -348,6 +354,10 @@ async def _run(task_id: str): await remove_download_active(task.id) + from app.services.generation_module_hook_service import notify_chat_generation_task_finished + await notify_chat_generation_task_finished(db, task) + await db.commit() + await log_task_event( task, event_type="DOWNLOAD_FAILED", diff --git a/video-gen-api/app/tasks/generation_poll_tasks.py b/video-gen-api/app/tasks/generation_poll_tasks.py index a6797593..08835b8b 100644 --- a/video-gen-api/app/tasks/generation_poll_tasks.py +++ b/video-gen-api/app/tasks/generation_poll_tasks.py @@ -13,6 +13,8 @@ from app.services.generation_refund_service import mark_chat_generation_task_fai from app.services.generation_provider_service import poll_provider_task from app.tasks.celery_app import celery_app +ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate"} + def _is_success(status: str) -> bool: return status in ("succeeded", "success", "completed", "done") @@ -29,6 +31,12 @@ def _engine_snapshot(task: ChatGenerationTask) -> dict: return {} +async def _notify_finished(db, task: ChatGenerationTask) -> None: + from app.services.generation_module_hook_service import notify_chat_generation_task_finished + + await notify_chat_generation_task_finished(db, task) + + async def _reload_task(db, task_id: str) -> ChatGenerationTask | None: """ rollback 后重新查询任务对象。 @@ -54,7 +62,7 @@ async def _run(task_id: str): ChatGenerationTask.deleted_at.is_(None), ).with_for_update().limit(1)) task = result.scalar_one_or_none() - if not task or task.generation_mode != "chatapi_async": + if not task or task.generation_mode not in ALLOWED_GENERATION_MODES: return # 只处理正在生成,且处于远程等待/轮询中的任务。 @@ -68,6 +76,7 @@ async def _run(task_id: str): error_message="任务轮询超时", pipeline_stage="timeout", ) + await _notify_finished(db, task) await db.commit() await log_task_event(task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout") return @@ -79,6 +88,7 @@ async def _run(task_id: str): error_message="缺少外部任务ID", pipeline_stage="failed", ) + await _notify_finished(db, task) await db.commit() await log_task_event(task, event_type="POLL_FAILED", message=task.error_message) return @@ -130,6 +140,7 @@ async def _run(task_id: str): error_message="供应商任务成功但未返回结果URL", pipeline_stage="failed", ) + await _notify_finished(db, task) await db.commit() await log_task_event(task, event_type="POLL_FAILED", message=task.error_message) return @@ -153,6 +164,7 @@ async def _run(task_id: str): error_message=poll_result.get("error") or f"供应商任务失败: {status}", pipeline_stage="failed", ) + await _notify_finished(db, task) await db.commit() await log_task_event(task, event_type="POLL_FAILED", message=task.error_message, detail=poll_result) return @@ -194,6 +206,7 @@ async def _run(task_id: str): error_message=error_message, pipeline_stage="failed", ) + await _notify_finished(db, task) await db.commit() await log_task_event(task, event_type="POLL_FAILED", message=task.error_message) else: diff --git a/video-gen-api/app/tasks/hot_opening_replicate_tasks.py b/video-gen-api/app/tasks/hot_opening_replicate_tasks.py new file mode 100644 index 00000000..a0228467 --- /dev/null +++ b/video-gen-api/app/tasks/hot_opening_replicate_tasks.py @@ -0,0 +1,46 @@ +from __future__ import annotations + +from app.models.base import async_session +from app.services.hot_opening_replicate_service import run_image_prompt_optimize, run_video_prompt_optimize +from app.tasks.async_runner import run_async +from app.tasks.celery_app import celery_app + + +async def _run_image_prompt(project_id: str, step_id: str | None = None): + async with async_session() as db: + await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id) + await db.commit() + + +async def _run_video_prompt(project_id: str, step_id: str | None = None): + async with async_session() as db: + await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id) + await db.commit() + + +if celery_app: + @celery_app.task(name="hot_opening.start_image_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30) + def start_image_prompt_optimize(self, project_id: str, step_id: str | None = None): + """手动触发后的图片 AI 提词任务。 + + 该任务路由到现有 gen_chatapi_create 队列,不需要新增 hot_opening worker。 + """ + return run_async(_run_image_prompt(project_id, step_id)) + + @celery_app.task(name="hot_opening.start_video_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30) + def start_video_prompt_optimize(self, project_id: str, step_id: str | None = None): + """手动触发后的视频 AI 提词任务。 + + 该任务路由到现有 gen_chatapi_create 队列,不需要新增 hot_opening worker。 + """ + return run_async(_run_video_prompt(project_id, step_id)) +else: + class _DisabledTask: + def delay(self, *args, **kwargs): + raise RuntimeError("Celery is disabled") + + def apply_async(self, *args, **kwargs): + raise RuntimeError("Celery is disabled") + + start_image_prompt_optimize = _DisabledTask() + start_video_prompt_optimize = _DisabledTask() diff --git a/video-gen-api/app/utils/id_gen.py b/video-gen-api/app/utils/id_gen.py index 57eb61d5..2ccd49cb 100644 --- a/video-gen-api/app/utils/id_gen.py +++ b/video-gen-api/app/utils/id_gen.py @@ -1,5 +1,7 @@ import time import random +from datetime import datetime +import secrets def generate_id() -> str: @@ -10,7 +12,10 @@ def generate_id() -> str: def generate_order_no() -> str: - """Generate a human-readable order number.""" - timestamp = int(time.time()) - randomness = random.randint(1000, 9999) - return f"VG{timestamp}{randomness}" + """Generate a human-readable order number with yyyymmddhhmmss format.""" + # 格式化为 yyyymmddhhmmss 格式的时间戳 + now = datetime.now() + timestamp = now.strftime("%Y%m%d%H%M%S") + # 使用密码学安全的随机数生成 8位纯数字,防止并发冲突 + random_part = ''.join(str(secrets.randbelow(10)) for _ in range(8)) + return f"MZZC{timestamp}{random_part}" diff --git a/video-gen-api/pyproject.toml b/video-gen-api/pyproject.toml index ff8435bb..b7721cdb 100644 --- a/video-gen-api/pyproject.toml +++ b/video-gen-api/pyproject.toml @@ -1,3 +1,7 @@ +# ⚠️ 系统依赖(需手动安装): +# ca-certificates — HTTPS 请求必须(新服务器/容器常缺) +# FFmpeg — 视频封面截帧(可选) + [project] name = "videogen-api" version = "1.0.0" @@ -22,6 +26,8 @@ dependencies = [ pg = ["asyncpg>=0.30.0"] redis = ["redis>=5.2.0"] celery = ["celery>=5.4.0", "redis>=5.2.0"] +alipay = ["alipay-sdk-python>=3.7.1160"] +volc = ["volcengine-python-sdk>=1.1.0"] dev = [ "pytest>=8.3.0", "pytest-asyncio>=0.24.0", diff --git a/video-gen-app/package-lock.json b/video-gen-app/package-lock.json index 48e91e00..00e61fb0 100644 --- a/video-gen-app/package-lock.json +++ b/video-gen-app/package-lock.json @@ -4000,7 +4000,7 @@ }, "node_modules/qrcode.react": { "version": "4.2.0", - "resolved": "https://registry.npmmirror.com/qrcode.react/-/qrcode.react-4.2.0.tgz", + "resolved": "https://registry.npmjs.org/qrcode.react/-/qrcode.react-4.2.0.tgz", "integrity": "sha512-QpgqWi8rD9DsS9EP3z7BT+5lY5SFhsqGjpgW5DY/i3mK4M9DTBNz3ErMi8BWYEfI3L0d8GIbGmcdFAS1uIRGjA==", "license": "ISC", "peerDependencies": { diff --git a/video-gen-app/private-folder-alias.json b/video-gen-app/private-folder-alias.json index 9e26dfee..5a0fdfb6 100644 --- a/video-gen-app/private-folder-alias.json +++ b/video-gen-app/private-folder-alias.json @@ -1 +1,20 @@ -{} \ No newline at end of file +{ + "src/pages/GenerateConver.tsx": { + "description": "ai创建" + }, + "src/pages/GeneratedRecord.tsx": { + "description": "生成历史" + }, + "src/pages/GeneratePage.tsx": { + "description": "生成记录" + }, + "src/pages/ProjectsPage.tsx": { + "description": "我的项目" + }, + "src/pages/RemoveLens.tsx": { + "description": "拆镜复刻" + }, + "src/pages/InitialReplication.tsx": { + "description": "爆款开头复刻" + } +} \ No newline at end of file diff --git a/video-gen-app/src/App.tsx b/video-gen-app/src/App.tsx index 1f867562..751f3f72 100644 --- a/video-gen-app/src/App.tsx +++ b/video-gen-app/src/App.tsx @@ -14,6 +14,7 @@ import InitialInfo from './pages/InitialInfo'; import RemoveLens from './pages/RemoveLens'; import GeneratedRecord from './pages/GeneratedRecord'; import AuthorizationPage from './pages/AuthorizationPage'; +import RemoveInfo from './pages/RemoveInfo'; @@ -96,6 +97,7 @@ const App = () => { } /> } /> } /> + } /> } /> } /> diff --git a/video-gen-app/src/api/index.ts b/video-gen-app/src/api/index.ts index 627c4f06..d4318155 100644 --- a/video-gen-app/src/api/index.ts +++ b/video-gen-app/src/api/index.ts @@ -262,6 +262,27 @@ export async function getMenuConfigs(): Promise { export async function getRechargePackages(): Promise { return api.get('/recharge-packages'); } + +export async function getPaymentMethods(): Promise<{ alipay: boolean; wechat: boolean }> { + return api.get('/payments/methods'); +} + +export async function createRechargeOrder(planId: string, method: string = 'wechat'): Promise { + return api.post('/payments/recharge', { plan: planId, method }); +} + +export async function getPaymentOrders(): Promise { + return api.get('/payments/orders'); +} + +export async function getPaymentOrder(orderNo: string): Promise { + return api.get(`/payments/orders/${orderNo}`); +} + +export async function cancelPaymentOrder(orderNo: string): Promise { + return api.post(`/payments/orders/${orderNo}/cancel`); +} + export async function getCreditRatios(): Promise { return api.get('/credits/ratios'); } diff --git a/video-gen-app/src/components/Layout/AppLayout.tsx b/video-gen-app/src/components/Layout/AppLayout.tsx index 115834ff..c19e1424 100644 --- a/video-gen-app/src/components/Layout/AppLayout.tsx +++ b/video-gen-app/src/components/Layout/AppLayout.tsx @@ -1,5 +1,5 @@ -import React, { useEffect, useState, useMemo } from 'react'; -import { Layout, Avatar, Dropdown, Space, Modal, Form, Input, message, Tooltip, Tag, Button, Typography } from 'antd'; +import React, { useEffect, useState, useCallback, useRef } from 'react'; +import { Layout, Avatar, Dropdown, Space, Modal, Form, Input, message, Tooltip, Tag, Button, Typography, Radio } from 'antd'; import { QRCodeSVG } from 'qrcode.react'; import { PlayCircleOutlined, @@ -21,12 +21,13 @@ import { FireFilled, CrownFilled, BankFilled, - QrcodeOutlined, CloseOutlined, + WechatOutlined, + AlipayCircleOutlined, } from '@ant-design/icons'; import { Outlet, useNavigate, useLocation } from 'react-router-dom'; import { useAuthStore } from '../../store/useAuthStore'; -import { getMenuConfigs, getRechargePackages, getNotifications, markNotificationRead, getSiteInfo } from '../../api'; +import { getMenuConfigs, getRechargePackages, getPaymentMethods, createRechargeOrder, getPaymentOrder, cancelPaymentOrder, getNotifications, markNotificationRead, getSiteInfo } from '../../api'; import NotificationPopup from '../NotificationPopup'; interface MenuConfig { @@ -88,7 +89,17 @@ const AppLayout: React.FC = () => { const [siteName, setSiteName] = useState('VideoGen.AI'); const [siteLogo, setSiteLogo] = useState(''); const [qrCodeModalOpen, setQrCodeModalOpen] = useState(false); - const [currentPaymentInfo, setCurrentPaymentInfo] = useState<{ price: number; credits: number; qrCode: string } | null>(null); + const [currentPaymentInfo, setCurrentPaymentInfo] = useState<{ price: number; credits: number; qrCode: string; method: string } | null>(null); + const [paymentMethod, setPaymentMethod] = useState('alipay'); + const [paying, setPaying] = useState(false); + const [countdown, setCountdown] = useState(180); // 默认180秒超时 + const pollingTimerRef = useRef | null>(null); + const countdownTimerRef = useRef | null>(null); + const currentOrderNoRef = useRef(null); + const [enabledMethods, setEnabledMethods] = useState<{ alipay: boolean; wechat: boolean }>({ alipay: false, wechat: false }); + + // LocalStorage keys + const PENDING_ORDER_KEY = 'pending_payment_order'; // 监听预览弹窗状态,关闭浮动按钮 useEffect(() => { @@ -114,6 +125,57 @@ const AppLayout: React.FC = () => { }).catch(() => {}); }; + // 检查并恢复待处理的支付订单 + useEffect(() => { + const checkPendingOrder = async () => { + const savedOrderStr = localStorage.getItem(PENDING_ORDER_KEY); + if (savedOrderStr) { + try { + const savedOrder = JSON.parse(savedOrderStr); + // 查询订单状态 + const order = await getPaymentOrder(savedOrder.orderNo); + if (order.status === 'pending') { + // 订单仍然待支付,恢复弹窗 + setCurrentPaymentInfo({ + price: savedOrder.price, + credits: savedOrder.credits, + qrCode: savedOrder.qrCode, + method: savedOrder.method, + }); + currentOrderNoRef.current = savedOrder.orderNo; + // 计算剩余时间 + const now = Date.now(); + const createdAt = new Date(savedOrder.createdAt).getTime(); + const timeoutSeconds = savedOrder.timeoutSeconds || 180; + const elapsedSeconds = Math.floor((now - createdAt) / 1000); + const remainingSeconds = Math.max(0, timeoutSeconds - elapsedSeconds); + + if (remainingSeconds > 0) { + setQrCodeModalOpen(true); + startPolling(savedOrder.orderNo, remainingSeconds); + } else { + // 已超时,清除 + localStorage.removeItem(PENDING_ORDER_KEY); + } + } else if (order.status === 'paid') { + // 已支付 + message.success('支付成功!积分已到账'); + useAuthStore.getState().refreshUser(); + localStorage.removeItem(PENDING_ORDER_KEY); + } else { + // 订单已取消或其他状态,清除 + localStorage.removeItem(PENDING_ORDER_KEY); + } + } catch { + // 查询失败,清除 + localStorage.removeItem(PENDING_ORDER_KEY); + } + } + }; + + checkPendingOrder(); + }, []); + useEffect(() => { getMenuConfigs().then(data => { let items = data.filter((m: any) => m.is_active !== false && m.isActive !== false); @@ -136,6 +198,12 @@ const AppLayout: React.FC = () => { getRechargePackages().then(data => { setRechargeOptions(data.filter((p: any) => p.is_active !== false && p.isActive !== false)); }).catch(() => {}); + getPaymentMethods().then(data => { + setEnabledMethods(data); + // Auto-select the first enabled method + if (data.alipay) setPaymentMethod('alipay'); + else if (data.wechat) setPaymentMethod('wechat'); + }).catch(() => {}); loadNotifications(); }, [user]); @@ -173,6 +241,68 @@ const AppLayout: React.FC = () => { setRechargeModalOpen(true); }; + const stopPolling = useCallback(() => { + if (pollingTimerRef.current) { + clearInterval(pollingTimerRef.current); + pollingTimerRef.current = null; + } + if (countdownTimerRef.current) { + clearInterval(countdownTimerRef.current); + countdownTimerRef.current = null; + } + }, []); + + const startPolling = useCallback((orderNo: string, timeoutSeconds: number = 180) => { + stopPolling(); + setCountdown(timeoutSeconds); + + // 订单状态轮询(每2秒查询一次,只查询当前订单 + const pollingTimer = setInterval(async () => { + try { + const order = await getPaymentOrder(orderNo); + if (order.status === 'paid') { + stopPolling(); + currentOrderNoRef.current = null; + localStorage.removeItem(PENDING_ORDER_KEY); + message.success('支付成功!积分已到账'); + useAuthStore.getState().refreshUser(); + setQrCodeModalOpen(false); + setCurrentPaymentInfo(null); + setSelectedPlan(null); + } else if (order.status === 'cancelled') { + stopPolling(); + currentOrderNoRef.current = null; + localStorage.removeItem(PENDING_ORDER_KEY); + } + } catch { + // ignore polling errors + } + }, 2000); + pollingTimerRef.current = pollingTimer; + + // 倒计时 + const countdownTimer = setInterval(() => { + setCountdown(prev => { + if (prev <= 1) { + // 超时自动取消 + stopPolling(); + if (currentOrderNoRef.current) { + cancelPaymentOrder(currentOrderNoRef.current).catch(() => {}); + currentOrderNoRef.current = null; + } + localStorage.removeItem(PENDING_ORDER_KEY); + message.warning('订单已超时,请重新充值'); + setQrCodeModalOpen(false); + setCurrentPaymentInfo(null); + setSelectedPlan(null); + return 0; + } + return prev - 1; + }); + }, 1000); + countdownTimerRef.current = countdownTimer; + }, [stopPolling]); + return ( {/* Desktop Sidebar */} @@ -516,28 +646,90 @@ const AppLayout: React.FC = () => { ); })} -
+ + {/* Payment method selection */} + {(!enabledMethods.alipay && !enabledMethods.wechat) ? ( +
+ + ⚠️ 暂无可用的支付方式,请联系管理员开启支付功能 + +
+ ) : ( +
+ 选择支付方式 + setPaymentMethod(e.target.value)} + style={{ display: 'flex', gap: 12 }}> + {enabledMethods.alipay && ( + + + 支付宝 + + )} + {enabledMethods.wechat && ( + + + 微信支付 + + )} + +
+ )} + +
-
+ {/* 倒计时显示 */} +
+ + 订单将在 {countdown} 秒后关闭 + +
@@ -640,31 +869,24 @@ const AppLayout: React.FC = () => { {/* Footer Buttons */} -
+
-
diff --git a/video-gen-app/src/pages/GeneratedRecord.tsx b/video-gen-app/src/pages/GeneratedRecord.tsx index b54deefd..bd4446f6 100644 --- a/video-gen-app/src/pages/GeneratedRecord.tsx +++ b/video-gen-app/src/pages/GeneratedRecord.tsx @@ -1066,10 +1066,10 @@ const GeneratedRecord: React.FC = () => {
diff --git a/video-gen-app/src/pages/InitialInfo.tsx b/video-gen-app/src/pages/InitialInfo.tsx index 807093dd..577511de 100644 --- a/video-gen-app/src/pages/InitialInfo.tsx +++ b/video-gen-app/src/pages/InitialInfo.tsx @@ -61,7 +61,7 @@ function InitialInfo() { return (
-
+
diff --git a/video-gen-app/src/pages/RemoveInfo.tsx b/video-gen-app/src/pages/RemoveInfo.tsx new file mode 100644 index 00000000..ec24e62f --- /dev/null +++ b/video-gen-app/src/pages/RemoveInfo.tsx @@ -0,0 +1,419 @@ +import React, { useState } from 'react'; +import { useParams, useNavigate } from 'react-router-dom'; +import { Button, Table, Tag, Drawer, Input, Upload, message } from 'antd'; +import { ArrowLeftOutlined, PlayCircleOutlined, XOutlined, PlusOutlined, UploadOutlined } from '@ant-design/icons'; +import type { UploadFile } from 'antd'; + +const { TextArea } = Input; + +function RemoveInfo() { + const { creatID } = useParams<{ creatID: string }>(); + const navigate = useNavigate(); + const [drawerVisible, setDrawerVisible] = useState(false); + const [currentSegment, setCurrentSegment] = useState(null); + const [productName, setProductName] = useState(''); + const [productSellingPoint, setProductSellingPoint] = useState(''); + const [productImage, setProductImage] = useState(''); + const [detailImage, setDetailImage] = useState(''); + + const handleGenerate = (segmentId: number) => { + setCurrentSegment(segmentId); + setDrawerVisible(true); + }; + + const handleCloseDrawer = () => { + setDrawerVisible(false); + setCurrentSegment(null); + setProductName(''); + setProductSellingPoint(''); + setProductImage(''); + setDetailImage(''); + }; + + const handleProductImageChange: any = (info: any) => { + if (info.fileList.length > 0) { + const file = info.fileList[0]; + if (file.originFileObj) { + const reader = new FileReader(); + reader.onload = (e) => { + setProductImage(e.target?.result as string); + }; + reader.readAsDataURL(file.originFileObj); + } + } else { + setProductImage(''); + } + }; + + const handleDetailImageChange: any = (info: any) => { + if (info.fileList.length > 0) { + const file = info.fileList[0]; + if (file.originFileObj) { + const reader = new FileReader(); + reader.onload = (e) => { + setDetailImage(e.target?.result as string); + }; + reader.readAsDataURL(file.originFileObj); + } + } else { + setDetailImage(''); + } + }; + + const handleManualGenerate = () => { + // 必填校验 + if (!productImage) { + message.warning('请上传产品图'); + return; + } + if (!productName.trim()) { + message.warning('请输入产品名称'); + return; + } + if (!productSellingPoint.trim()) { + message.warning('请输入产品卖点'); + return; + } + + // 输出内容 + console.log('手动生成 - 片段', currentSegment); + console.log('产品图:', productImage); + console.log('细节图:', detailImage); + console.log('产品名称:', productName); + console.log('产品卖点:', productSellingPoint); + + message.success(`手动生成成功!片段: ${currentSegment}`); + }; + + const mockData = { + productName: '返回', + uploadTime: '2026-06-09 09:01:16', + sellingPoints: ['一键匹配', '连麦聊天'], + audience: '123123', + audienceAnalysis: '123123', + videoUrl: 'https://images.unsplash.com/photo-1506905925346-21bda4d32df4?w=320&h=180&fit=crop', + segments: [ + { + key: '1', + id: 1, + timeRange: '00:00 - 00:03', + thumbnail: 'https://images.unsplash.com/photo-1506905925346-21bda4d32df4?w=120&h=80&fit=crop', + content: '11111', + lines: 'qqqqqqqqqqq', + contentStrategy: '展示礼盒' + }, + { + key: '2', + id: 2, + timeRange: '00:03 - 00:06', + thumbnail: 'https://images.unsplash.com/photo-1494790108377-be9c29b29330?w=120&h=80&fit=crop', + content: '1231231231', + lines: 'qqqqqqqqqqq', + contentStrategy: '开箱展示' + }, + { + key: '3', + id: 3, + timeRange: '00:06 - 00:09', + thumbnail: 'https://images.unsplash.com/photo-1522202176988-66273c2fd55f?w=120&h=80&fit=crop', + content: '123123', + lines: 'qqqqqqqqqqq', + contentStrategy: '取出产品' + }, + ] + }; + + const columns = [ + { + title: '片段', + width: 100, + render: (text: any, record: any) => ( +
+
片段{record.id}
+
{record.timeRange}
+
+ ) + }, + { + title: '片段视频', + width: 120, + render: (text: any, record: any) => ( +
+ {`片段${record.id}`} +
+ +
+
+ ) + }, + { + title: '画面内容', + width: 250, + + render: (text: any, record: any) => ( +
+ {record.content} +
+ ) + }, + { + title: '台词', + width: 250, + + render: (text: any, record: any) => ( +
+ {record.lines} +
+ ) + }, + { + title: '内容策略', + width: 120, + align: 'left' as const, + render: (text: any, record: any) => ( +
+ {record.contentStrategy || '-'} +
+ ) + }, + { + title: '素材', + width: 140, + align: 'center' as const, + render: (text: any, record: any) => ( +
+
+ 等待生成 +
+ +
+ ) + } + ]; + + return ( + +
+
+
+
+
+ +
+
+
+
+ 视频缩略图 +
+ +
+
+ +
+
+
+

视频总结

+
+ 上传时间: {mockData.uploadTime} +
+ +
+
+ 产品名称: + {mockData.productName} +
+
+ 卖点词: +
+ {mockData.sellingPoints.map((point, index) => ( + + {point} + + ))} +
+
+
+ 受众群体: + {mockData.audience} +
+
+ 受众分析: + {mockData.audienceAnalysis} +
+
+ + +
+
+
+ +
+
+ + + + + + + 智能视频复刻 + -片段{currentSegment} + +