This commit is contained in:
2026-07-20 14:01:22 +08:00
parent 53e01f6f74
commit a2e98ef3a3
75 changed files with 5912 additions and 2666 deletions
File diff suppressed because one or more lines are too long
+1 -1
View File
@@ -28,7 +28,7 @@
}
})();
</script>
<script type="module" crossorigin src="/assets/index-CzqFKsZY.js"></script>
<script type="module" crossorigin src="/assets/index-CvhbmotR.js"></script>
<link rel="stylesheet" crossorigin href="/assets/index-D7ShJUt4.css">
</head>
<body>
+6 -4
View File
@@ -348,8 +348,9 @@ export async function deleteIndustryConfig(id: string): Promise<void> {
await api.delete(`/admin/industry-configs/${id}`);
}
export async function getVideoEngines(): Promise<any[]> {
return api.get('/admin/video-engines');
export async function getVideoEngines(options?: { includeDeleted?: boolean }): Promise<any[]> {
const query = options?.includeDeleted ? '?include_deleted=true' : '';
return api.get(`/admin/video-engines${query}`);
}
export async function saveVideoEngine(engine: any): Promise<any> {
@@ -361,8 +362,9 @@ export async function deleteVideoEngine(id: string): Promise<void> {
await api.delete(`/admin/video-engines/${id}`);
}
export async function getImageEngines(): Promise<any[]> {
return api.get('/admin/image-engines');
export async function getImageEngines(options?: { includeDeleted?: boolean }): Promise<any[]> {
const query = options?.includeDeleted ? '?include_deleted=true' : '';
return api.get(`/admin/image-engines${query}`);
}
export async function saveImageEngine(engine: any): Promise<any> {
@@ -29,8 +29,13 @@ import {
VideoCameraOutlined,
} from '@ant-design/icons';
import dayjs from 'dayjs';
import { getAdminGenerationAiTasks } from '../api';
import type { GenerationAIMediaReference, GenerationAITaskOut } from '../types';
import { getAdminGenerationAiTasks, getImageEngines, getVideoEngines } from '../api';
import type {
GenerationAiImageEngine,
GenerationAIMediaReference,
GenerationAITaskOut,
GenerationAiVideoEngine,
} from '../types';
import { formatDate } from '../utils/formatDate';
import GenerationTaskResourceGrid from '../components/generation/GenerationTaskResourceGrid';
@@ -240,6 +245,9 @@ const AdminGenerationAiRecords: React.FC = () => {
const [createdRange, setCreatedRange] = useState<any>([todayStart(), todayEnd()]);
const [queryCreatedRange, setQueryCreatedRange] = useState<any>([todayStart(), todayEnd()]);
const [reloadKey, setReloadKey] = useState(0);
const [engineListLoading, setEngineListLoading] = useState(false);
const [imageEngines, setImageEngines] = useState<GenerationAiImageEngine[]>([]);
const [videoEngines, setVideoEngines] = useState<GenerationAiVideoEngine[]>([]);
const [preview, setPreview] = useState<GenerationAITaskOut | null>(null);
const [resourceState, setResourceState] = useState<PreviewResourceState>(EMPTY_RESOURCE_STATE);
@@ -280,6 +288,64 @@ const AdminGenerationAiRecords: React.FC = () => {
load();
}, [load, reloadKey]);
useEffect(() => {
let cancelled = false;
const loadEngineOptions = async () => {
setEngineListLoading(true);
try {
const [images, videos] = await Promise.all([
getImageEngines({ includeDeleted: true }),
getVideoEngines({ includeDeleted: true }),
]);
if (!cancelled) {
setImageEngines(images || []);
setVideoEngines(videos || []);
}
} catch (error: any) {
if (!cancelled) {
message.error(error?.message || '加载模型引擎列表失败');
}
} finally {
if (!cancelled) setEngineListLoading(false);
}
};
void loadEngineOptions();
return () => {
cancelled = true;
};
}, []);
const engineOptions = useMemo(() => {
const toOption = (
engine: GenerationAiImageEngine | GenerationAiVideoEngine,
type: 'image' | 'video',
) => {
const deleted = Boolean(engine.deletedAt);
const typeText = type === 'image' ? '图片' : '视频';
const deletedText = deleted ? '[已删除]' : '';
const detailText = [engine.name, engine.modelName, engine.provider, engine.id]
.filter(Boolean)
.join(' / ');
const label = `[${typeText}]${deletedText} ${detailText}`;
return {
value: engine.id,
label,
searchText: [engine.id, engine.name, engine.modelName, engine.provider, typeText, deleted ? '已删除' : '']
.filter(Boolean)
.join(' ')
.toLowerCase(),
};
};
const merged = [
...imageEngines.map((engine) => toOption(engine, 'image')),
...videoEngines.map((engine) => toOption(engine, 'video')),
];
return Array.from(new Map(merged.map((option) => [option.value, option])).values());
}, [imageEngines, videoEngines]);
useEffect(() => {
if (!preview) {
setResourceState(EMPTY_RESOURCE_STATE);
@@ -969,21 +1035,17 @@ const AdminGenerationAiRecords: React.FC = () => {
<Select
allowClear
showSearch
loading={engineListLoading}
placeholder="引擎筛选"
value={filterEngineId || undefined}
style={{ width: 180 }}
style={{ width: 260 }}
onChange={(v) => { setFilterEngineId(v || ''); setPage(1); setQueryEngineId(v || ''); }}
optionFilterProp="label"
options={Array.from(
new Map(
records
.filter((r) => r.engineId)
.map((r) => [r.engineId, {
value: r.engineId,
label: getEngineName(r.engineSnapshot as any) || r.engineId,
}]),
).values(),
)}
filterOption={(input, option: any) =>
String(option?.searchText || option?.label || '')
.toLowerCase()
.includes(input.trim().toLowerCase())
}
options={engineOptions}
/>
<RangePicker
value={createdRange}
+4
View File
@@ -285,6 +285,8 @@ export interface GenerationAiImageEngine {
maxGenerationCount: number;
multiImageMaxImages: number;
maxReferenceImageCount: number;
isActive?: boolean;
deletedAt?: string | null;
}
export interface GenerationAiVideoEngine {
@@ -305,6 +307,8 @@ export interface GenerationAiVideoEngine {
priority: number;
multiGenerationEnabled: boolean;
maxGenerationCount: number;
isActive?: boolean;
deletedAt?: string | null;
}
export interface GenerationAiEnginesResponse {
@@ -0,0 +1,922 @@
"""unify generation owners and soft delete engines
Revision ID: e6eac828ff61
Revises: 3f47680a71d0
Create Date: 2026-07-20 09:19:07.410183
This revision intentionally contains only the generation-pipeline and engine
soft-delete changes. Autogenerate noise from unrelated modules is excluded.
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = "e6eac828ff61"
down_revision: Union[str, None] = "3f47680a71d0"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
FK_EVENT_GENERATION_RECORD = "fk_chat_generation_task_events_generation_record_id"
FK_CALL_LOG_GENERATION_RECORD = "fk_chat_provider_call_logs_generation_record_id"
CK_EVENT_OWNER = "ck_chat_generation_task_events_owner"
CK_CALL_LOG_OWNER = "ck_chat_provider_call_logs_owner"
def _add_engine_soft_delete_columns() -> None:
for table_name in ("image_engines", "model_configs", "video_engines"):
op.add_column(
table_name,
sa.Column("deleted_at", sa.DateTime(timezone=True), nullable=True),
)
op.create_index(
op.f(f"ix_{table_name}_deleted_at"),
table_name,
["deleted_at"],
unique=False,
)
def _add_shared_log_owner_columns() -> None:
op.add_column(
"chat_generation_task_events",
sa.Column(
"owner_type",
sa.String(length=32),
server_default="chat_generation_task",
nullable=False,
),
)
op.add_column(
"chat_generation_task_events",
sa.Column("generation_record_id", sa.String(length=32), nullable=True),
)
op.add_column(
"chat_generation_task_events",
sa.Column(
"generation_attempt_no",
sa.Integer(),
server_default="1",
nullable=False,
),
)
op.alter_column(
"chat_generation_task_events",
"task_id",
existing_type=sa.VARCHAR(length=32),
nullable=True,
)
op.alter_column(
"chat_generation_task_events",
"from_stage",
existing_type=sa.VARCHAR(length=32),
type_=sa.String(length=48),
existing_nullable=True,
)
op.alter_column(
"chat_generation_task_events",
"to_stage",
existing_type=sa.VARCHAR(length=32),
type_=sa.String(length=48),
existing_nullable=True,
)
op.alter_column(
"chat_generation_task_events",
"message",
existing_type=sa.VARCHAR(length=512),
type_=sa.Text(),
existing_nullable=True,
)
op.create_index(
"idx_chat_generation_task_events_attempt_created",
"chat_generation_task_events",
["owner_type", "generation_attempt_no", "created_at"],
unique=False,
)
op.create_index(
"idx_chat_generation_task_events_record_created",
"chat_generation_task_events",
["owner_type", "generation_record_id", "created_at"],
unique=False,
)
op.create_index(
"idx_chat_generation_task_events_task_created",
"chat_generation_task_events",
["owner_type", "task_id", "created_at"],
unique=False,
)
op.create_index(
op.f("ix_chat_generation_task_events_generation_attempt_no"),
"chat_generation_task_events",
["generation_attempt_no"],
unique=False,
)
op.create_index(
op.f("ix_chat_generation_task_events_generation_record_id"),
"chat_generation_task_events",
["generation_record_id"],
unique=False,
)
op.create_index(
op.f("ix_chat_generation_task_events_owner_type"),
"chat_generation_task_events",
["owner_type"],
unique=False,
)
op.create_foreign_key(
FK_EVENT_GENERATION_RECORD,
"chat_generation_task_events",
"generation_records",
["generation_record_id"],
["id"],
ondelete="CASCADE",
)
op.create_check_constraint(
CK_EVENT_OWNER,
"chat_generation_task_events",
"(owner_type = 'chat_generation_task' AND task_id IS NOT NULL AND generation_record_id IS NULL) "
"OR (owner_type = 'generation_record' AND task_id IS NULL AND generation_record_id IS NOT NULL)",
)
op.add_column(
"chat_provider_call_logs",
sa.Column(
"owner_type",
sa.String(length=32),
server_default="chat_generation_task",
nullable=False,
),
)
op.add_column(
"chat_provider_call_logs",
sa.Column("generation_record_id", sa.String(length=32), nullable=True),
)
op.add_column(
"chat_provider_call_logs",
sa.Column(
"generation_attempt_no",
sa.Integer(),
server_default="1",
nullable=False,
),
)
op.alter_column(
"chat_provider_call_logs",
"task_id",
existing_type=sa.VARCHAR(length=32),
nullable=True,
)
op.create_index(
"idx_chat_provider_call_logs_attempt_created",
"chat_provider_call_logs",
["owner_type", "generation_attempt_no", "created_at"],
unique=False,
)
op.create_index(
"idx_chat_provider_call_logs_record_created",
"chat_provider_call_logs",
["owner_type", "generation_record_id", "created_at"],
unique=False,
)
op.create_index(
"idx_chat_provider_call_logs_task_created",
"chat_provider_call_logs",
["owner_type", "task_id", "created_at"],
unique=False,
)
op.create_index(
op.f("ix_chat_provider_call_logs_generation_attempt_no"),
"chat_provider_call_logs",
["generation_attempt_no"],
unique=False,
)
op.create_index(
op.f("ix_chat_provider_call_logs_generation_record_id"),
"chat_provider_call_logs",
["generation_record_id"],
unique=False,
)
op.create_index(
op.f("ix_chat_provider_call_logs_owner_type"),
"chat_provider_call_logs",
["owner_type"],
unique=False,
)
op.create_foreign_key(
FK_CALL_LOG_GENERATION_RECORD,
"chat_provider_call_logs",
"generation_records",
["generation_record_id"],
["id"],
ondelete="CASCADE",
)
op.create_check_constraint(
CK_CALL_LOG_OWNER,
"chat_provider_call_logs",
"(owner_type = 'chat_generation_task' AND task_id IS NOT NULL AND generation_record_id IS NULL) "
"OR (owner_type = 'generation_record' AND task_id IS NULL AND generation_record_id IS NOT NULL)",
)
def _add_chat_generation_task_columns() -> None:
op.add_column(
"chat_generation_tasks",
sa.Column(
"generation_attempt_no",
sa.Integer(),
server_default="1",
nullable=False,
),
)
op.add_column(
"chat_generation_tasks",
sa.Column(
"resource_generation_started_at",
sa.DateTime(timezone=True),
nullable=True,
),
)
op.add_column(
"chat_generation_tasks",
sa.Column(
"manual_retry_count",
sa.Integer(),
server_default="0",
nullable=False,
),
)
op.add_column(
"chat_generation_tasks",
sa.Column(
"poll_error_count",
sa.Integer(),
server_default="0",
nullable=False,
),
)
op.add_column(
"chat_generation_tasks",
sa.Column("poll_claim_token", sa.String(length=64), nullable=True),
)
op.add_column(
"chat_generation_tasks",
sa.Column("poll_lease_until", sa.DateTime(timezone=True), nullable=True),
)
op.add_column(
"chat_generation_tasks",
sa.Column("download_claim_token", sa.String(length=64), nullable=True),
)
op.create_index(
op.f("ix_chat_generation_tasks_download_claim_token"),
"chat_generation_tasks",
["download_claim_token"],
unique=False,
)
op.create_index(
op.f("ix_chat_generation_tasks_poll_claim_token"),
"chat_generation_tasks",
["poll_claim_token"],
unique=False,
)
op.create_index(
op.f("ix_chat_generation_tasks_poll_lease_until"),
"chat_generation_tasks",
["poll_lease_until"],
unique=False,
)
op.create_index(
op.f("ix_chat_generation_tasks_resource_generation_started_at"),
"chat_generation_tasks",
["resource_generation_started_at"],
unique=False,
)
def _add_generation_record_columns() -> None:
columns = [
sa.Column("generation_attempt_no", sa.Integer(), server_default="1", nullable=False),
sa.Column("resource_generation_started_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("deadline_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("engine_id", sa.String(length=32), nullable=True),
sa.Column("engine_snapshot_json", sa.Text(), nullable=True),
sa.Column("provider_response_json", sa.Text(), nullable=True),
sa.Column("remote_result_url", sa.Text(), nullable=True),
sa.Column("provider_create_claim_token", sa.String(length=64), nullable=True),
sa.Column("provider_create_lease_until", sa.DateTime(timezone=True), nullable=True),
sa.Column("provider_create_started_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("retry_count", sa.Integer(), server_default="0", nullable=False),
sa.Column("manual_retry_count", sa.Integer(), server_default="0", nullable=False),
sa.Column("poll_error_count", sa.Integer(), server_default="0", nullable=False),
sa.Column("poll_count", sa.Integer(), server_default="0", nullable=False),
sa.Column("last_poll_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("poll_started_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("next_poll_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("poll_interval_seconds", sa.Integer(), server_default="0", nullable=False),
sa.Column("poll_claim_token", sa.String(length=64), nullable=True),
sa.Column("poll_lease_until", sa.DateTime(timezone=True), nullable=True),
sa.Column("download_celery_task_id", sa.String(length=160), nullable=True),
sa.Column("download_enqueued_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("download_started_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("download_claim_token", sa.String(length=64), nullable=True),
sa.Column("download_lease_until", sa.DateTime(timezone=True), nullable=True),
sa.Column("download_next_retry_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("download_attempt_count", sa.Integer(), server_default="0", nullable=False),
sa.Column("download_last_error", sa.Text(), nullable=True),
sa.Column("download_storage_date_dir", sa.String(length=16), nullable=True),
]
for column in columns:
op.add_column("generation_records", column)
op.create_index(
"idx_genrec_next_poll_at",
"generation_records",
["next_poll_at"],
unique=False,
postgresql_where=sa.text(
"deleted_at IS NULL AND status = 'generating' "
"AND gen_type = 'video' AND next_poll_at IS NOT NULL"
),
)
for column_name in (
"deadline_at",
"download_celery_task_id",
"download_claim_token",
"download_lease_until",
"download_next_retry_at",
"engine_id",
"poll_claim_token",
"poll_lease_until",
"provider_create_claim_token",
"provider_create_lease_until",
"resource_generation_started_at",
):
op.create_index(
op.f(f"ix_generation_records_{column_name}"),
"generation_records",
[column_name],
unique=False,
)
def _backfill_generation_attempts() -> None:
# Existing ChatGenerationTask rows started resource generation when the row
# was created. Do not derive this timestamp from poll/download timestamps.
op.execute(
"""
UPDATE chat_generation_tasks
SET resource_generation_started_at = created_at
WHERE resource_generation_started_at IS NULL
"""
)
# Prefer structured owner/attempt fields. The biz_key parser is retained
# for older credit rows that were written before those columns were filled.
op.execute(
"""
WITH credit_attempts AS (
SELECT
COALESCE(
NULLIF(owner_id, ''),
NULLIF(related_id, ''),
substring(biz_key FROM '^chat_generation_task:([^:]+):attempt:')
) AS resolved_owner_id,
MAX(
COALESCE(
attempt_no,
NULLIF(substring(biz_key FROM ':attempt:([0-9]+):media:charge$'), '')::integer
)
) AS max_attempt
FROM credit_records
WHERE type = 'consume'
AND (
(owner_type = 'chat_generation_task' AND charge_kind = 'media')
OR biz_key ~ '^chat_generation_task:[^:]+:attempt:[0-9]+:media:charge$'
)
GROUP BY 1
)
UPDATE chat_generation_tasks AS task
SET generation_attempt_no = GREATEST(1, attempts.max_attempt)
FROM credit_attempts AS attempts
WHERE attempts.resolved_owner_id = task.id
AND attempts.max_attempt IS NOT NULL
"""
)
op.execute(
"""
UPDATE chat_generation_tasks
SET manual_retry_count = GREATEST(generation_attempt_no - 1, 0),
retry_count = GREATEST(generation_attempt_no - 1, 0),
poll_error_count = 0
"""
)
op.execute(
"""
WITH generation_charges AS (
SELECT
resolved_owner_id,
resolved_attempt,
charge_created_at,
engine_id
FROM (
SELECT
COALESCE(
NULLIF(owner_id, ''),
NULLIF(related_id, ''),
substring(biz_key FROM '^generation_record:([^:]+):attempt:')
) AS resolved_owner_id,
COALESCE(
attempt_no,
NULLIF(substring(biz_key FROM ':attempt:([0-9]+):media:charge$'), '')::integer,
1
) AS resolved_attempt,
created_at AS charge_created_at,
engine_id,
ROW_NUMBER() OVER (
PARTITION BY COALESCE(
NULLIF(owner_id, ''),
NULLIF(related_id, ''),
substring(biz_key FROM '^generation_record:([^:]+):attempt:')
)
ORDER BY
COALESCE(
attempt_no,
NULLIF(substring(biz_key FROM ':attempt:([0-9]+):media:charge$'), '')::integer,
1
) DESC,
created_at DESC,
id DESC
) AS row_no
FROM credit_records
WHERE type = 'consume'
AND (
(owner_type = 'generation_record' AND charge_kind = 'media')
OR biz_key ~ '^generation_record:[^:]+:attempt:[0-9]+:media:charge$'
)
) AS ranked
WHERE row_no = 1
AND resolved_owner_id IS NOT NULL
)
UPDATE generation_records AS record
SET generation_attempt_no = GREATEST(1, charge.resolved_attempt),
resource_generation_started_at = COALESCE(
record.resource_generation_started_at,
charge.charge_created_at
),
engine_id = COALESCE(record.engine_id, charge.engine_id)
FROM generation_charges AS charge
WHERE charge.resolved_owner_id = record.id
"""
)
# Old GenerationRecord rows do not have a dedicated resource-start field.
# Only rows with clear resource-generation evidence are backfilled; prompt-
# optimized-only records intentionally remain NULL.
op.execute(
"""
UPDATE generation_records
SET resource_generation_started_at = COALESCE(updated_at, created_at)
WHERE resource_generation_started_at IS NULL
AND (
status IN ('generating', 'completed')
OR seedance_task_id IS NOT NULL
OR image_url IS NOT NULL
OR video_url IS NOT NULL
)
"""
)
op.execute(
"""
UPDATE generation_records
SET manual_retry_count = GREATEST(generation_attempt_no - 1, 0),
retry_count = GREATEST(generation_attempt_no - 1, 0),
poll_error_count = 0
"""
)
def _backfill_generation_record_engine() -> None:
# For recoverable active rows without a historical billing engine snapshot,
# fall back to the current highest-priority active engine of the same type.
# Completed/failed history is not assigned a guessed engine.
op.execute(
"""
UPDATE generation_records AS record
SET engine_id = engine.id
FROM (
SELECT id
FROM image_engines
WHERE is_active IS TRUE AND deleted_at IS NULL
ORDER BY priority DESC, id ASC
LIMIT 1
) AS engine
WHERE record.engine_id IS NULL
AND record.status = 'generating'
AND record.gen_type = 'image'
"""
)
op.execute(
"""
UPDATE generation_records AS record
SET engine_id = engine.id
FROM (
SELECT id
FROM video_engines
WHERE is_active IS TRUE AND deleted_at IS NULL
ORDER BY priority DESC, id ASC
LIMIT 1
) AS engine
WHERE record.engine_id IS NULL
AND record.status = 'generating'
AND record.gen_type = 'video'
"""
)
# Store a key-free execution snapshot. Runtime code still reads api_key
# from the engine row by engine_id, including soft-deleted historical rows.
op.execute(
"""
UPDATE generation_records AS record
SET engine_snapshot_json = jsonb_build_object(
'engine_type', 'image',
'id', engine.id,
'name', engine.name,
'provider', engine.provider,
'api_base', engine.api_base,
'api_key_masked', CASE WHEN COALESCE(engine.api_key, '') <> '' THEN '****' ELSE '' END,
'model_name', engine.model_name,
'generate_url', engine.generate_url,
'default_size', engine.default_size,
'multi_generation_enabled', engine.multi_generation_enabled,
'max_generation_count', engine.max_generation_count,
'multi_image_max_images', engine.multi_image_max_images,
'max_reference_image_count', engine.max_reference_image_count,
'output_format', engine.output_format,
'selected_size', COALESCE(record.image_size, engine.default_size, '2K'),
'selected_proportion', COALESCE(record.image_proportion, '1:1'),
'selected_px', COALESCE(record.image_px, '2048x2048')
)::text
FROM image_engines AS engine
WHERE record.gen_type = 'image'
AND record.engine_id = engine.id
AND record.engine_snapshot_json IS NULL
"""
)
op.execute(
"""
UPDATE generation_records AS record
SET engine_snapshot_json = jsonb_build_object(
'engine_type', 'video',
'id', engine.id,
'name', engine.name,
'provider', engine.provider,
'api_base', engine.api_base,
'api_key_masked', CASE WHEN COALESCE(engine.api_key, '') <> '' THEN '****' ELSE '' END,
'model_name', engine.model_name,
'generate_url', engine.generate_url,
'query_url', engine.query_url,
'max_duration', engine.max_duration,
'max_audio_count', engine.max_audio_count,
'multi_generation_enabled', engine.multi_generation_enabled,
'max_generation_count', engine.max_generation_count,
'selected_ratio', COALESCE(record.aspect_ratio, '16:9'),
'selected_resolution', COALESCE(record.resolution, '480p'),
'selected_duration', COALESCE(record.duration, 4)
)::text
FROM video_engines AS engine
WHERE record.gen_type = 'video'
AND record.engine_id = engine.id
AND record.engine_snapshot_json IS NULL
"""
)
def _backfill_generation_pipeline_state() -> None:
# Preserve absolute historical provider URLs when they are available.
op.execute(
"""
UPDATE generation_records
SET remote_result_url = CASE
WHEN gen_type = 'image' AND image_url ~ '^https?://' THEN image_url
WHEN gen_type = 'video' AND video_url ~ '^https?://' THEN video_url
ELSE remote_result_url
END
WHERE remote_result_url IS NULL
AND (
(gen_type = 'image' AND image_url ~ '^https?://')
OR (gen_type = 'video' AND video_url ~ '^https?://')
)
"""
)
op.execute(
"""
UPDATE generation_records
SET pipeline_stage = CASE
WHEN remote_result_url IS NOT NULL THEN 'result_ready'
WHEN seedance_task_id IS NOT NULL THEN 'waiting_remote'
ELSE 'queued'
END
WHERE status = 'generating'
AND pipeline_stage IS NULL
"""
)
op.execute(
"""
UPDATE generation_records
SET pipeline_stage = 'done'
WHERE status = 'completed'
AND pipeline_stage IS NULL
"""
)
op.execute(
"""
UPDATE generation_records
SET pipeline_stage = 'failed'
WHERE status = 'failed'
AND pipeline_stage IS NULL
"""
)
# Image deadline is now 30 minutes; video remains 24 hours. Only active
# tasks are rewritten so terminal historical audit values are preserved.
op.execute(
"""
UPDATE chat_generation_tasks
SET deadline_at = resource_generation_started_at + INTERVAL '30 minutes'
WHERE status = 'generating'
AND gen_type = 'image'
AND resource_generation_started_at IS NOT NULL
"""
)
op.execute(
"""
UPDATE chat_generation_tasks
SET deadline_at = resource_generation_started_at + INTERVAL '24 hours'
WHERE status = 'generating'
AND gen_type = 'video'
AND resource_generation_started_at IS NOT NULL
"""
)
op.execute(
"""
UPDATE generation_records
SET deadline_at = resource_generation_started_at
+ CASE WHEN gen_type = 'image' THEN INTERVAL '30 minutes' ELSE INTERVAL '24 hours' END
WHERE status = 'generating'
AND resource_generation_started_at IS NOT NULL
"""
)
# Active provider tasks are checked immediately after deployment. The
# recovery code still honours Redis execution locks and poll leases.
op.execute(
"""
UPDATE generation_records
SET poll_started_at = COALESCE(poll_started_at, resource_generation_started_at),
next_poll_at = CURRENT_TIMESTAMP
WHERE status = 'generating'
AND gen_type = 'video'
AND seedance_task_id IS NOT NULL
AND remote_result_url IS NULL
AND pipeline_stage IN ('waiting_remote', 'polling')
"""
)
def _create_generated_resource_idempotency_index() -> None:
# Do not silently soft-delete duplicates here: doing so without rebuilding
# user_resource_*_stats would corrupt capacity totals. Abort with a clear
# message so dirty data can be repaired and stats rebuilt deliberately.
op.execute(
"""
DO $$
BEGIN
IF EXISTS (
SELECT 1
FROM generated_resources
WHERE deleted_at IS NULL
GROUP BY source_model, source_id, resource_type
HAVING COUNT(*) > 1
) THEN
RAISE EXCEPTION
'Cannot create uq_generated_resources_active_source_type: duplicate active generated_resources exist'
USING HINT = 'Query duplicates by source_model/source_id/resource_type, keep one active row, soft-delete the others, then rebuild user resource stats before rerunning this revision.';
END IF;
END
$$;
"""
)
op.create_index(
"uq_generated_resources_active_source_type",
"generated_resources",
["source_model", "source_id", "resource_type"],
unique=True,
postgresql_where=sa.text("deleted_at IS NULL"),
)
def upgrade() -> None:
_add_engine_soft_delete_columns()
_add_shared_log_owner_columns()
_add_chat_generation_task_columns()
_add_generation_record_columns()
_backfill_generation_attempts()
_backfill_generation_record_engine()
_backfill_generation_pipeline_state()
_create_generated_resource_idempotency_index()
def downgrade() -> None:
op.drop_index(
"uq_generated_resources_active_source_type",
table_name="generated_resources",
postgresql_where=sa.text("deleted_at IS NULL"),
)
for column_name in (
"resource_generation_started_at",
"provider_create_lease_until",
"provider_create_claim_token",
"poll_lease_until",
"poll_claim_token",
"engine_id",
"download_next_retry_at",
"download_lease_until",
"download_claim_token",
"download_celery_task_id",
"deadline_at",
):
op.drop_index(
op.f(f"ix_generation_records_{column_name}"),
table_name="generation_records",
)
op.drop_index("idx_genrec_next_poll_at", table_name="generation_records")
for column_name in (
"download_storage_date_dir",
"download_last_error",
"download_attempt_count",
"download_next_retry_at",
"download_lease_until",
"download_claim_token",
"download_started_at",
"download_enqueued_at",
"download_celery_task_id",
"poll_lease_until",
"poll_claim_token",
"poll_interval_seconds",
"next_poll_at",
"poll_started_at",
"last_poll_at",
"poll_count",
"poll_error_count",
"manual_retry_count",
"retry_count",
"provider_create_started_at",
"provider_create_lease_until",
"provider_create_claim_token",
"remote_result_url",
"provider_response_json",
"engine_snapshot_json",
"engine_id",
"deadline_at",
"resource_generation_started_at",
"generation_attempt_no",
):
op.drop_column("generation_records", column_name)
# GenerationRecord log rows cannot be represented by the old task_id-only
# schema, so they are removed during downgrade before task_id becomes NOT NULL.
op.execute(
"DELETE FROM chat_provider_call_logs "
"WHERE owner_type = 'generation_record' OR generation_record_id IS NOT NULL"
)
op.drop_constraint(CK_CALL_LOG_OWNER, "chat_provider_call_logs", type_="check")
op.drop_constraint(
FK_CALL_LOG_GENERATION_RECORD,
"chat_provider_call_logs",
type_="foreignkey",
)
op.drop_index(
op.f("ix_chat_provider_call_logs_owner_type"),
table_name="chat_provider_call_logs",
)
op.drop_index(
op.f("ix_chat_provider_call_logs_generation_record_id"),
table_name="chat_provider_call_logs",
)
op.drop_index(
op.f("ix_chat_provider_call_logs_generation_attempt_no"),
table_name="chat_provider_call_logs",
)
op.drop_index("idx_chat_provider_call_logs_task_created", table_name="chat_provider_call_logs")
op.drop_index("idx_chat_provider_call_logs_record_created", table_name="chat_provider_call_logs")
op.drop_index("idx_chat_provider_call_logs_attempt_created", table_name="chat_provider_call_logs")
op.alter_column(
"chat_provider_call_logs",
"task_id",
existing_type=sa.VARCHAR(length=32),
nullable=False,
)
op.drop_column("chat_provider_call_logs", "generation_attempt_no")
op.drop_column("chat_provider_call_logs", "generation_record_id")
op.drop_column("chat_provider_call_logs", "owner_type")
op.execute(
"DELETE FROM chat_generation_task_events "
"WHERE owner_type = 'generation_record' OR generation_record_id IS NOT NULL"
)
op.drop_constraint(CK_EVENT_OWNER, "chat_generation_task_events", type_="check")
op.drop_constraint(
FK_EVENT_GENERATION_RECORD,
"chat_generation_task_events",
type_="foreignkey",
)
op.drop_index(
op.f("ix_chat_generation_task_events_owner_type"),
table_name="chat_generation_task_events",
)
op.drop_index(
op.f("ix_chat_generation_task_events_generation_record_id"),
table_name="chat_generation_task_events",
)
op.drop_index(
op.f("ix_chat_generation_task_events_generation_attempt_no"),
table_name="chat_generation_task_events",
)
op.drop_index("idx_chat_generation_task_events_task_created", table_name="chat_generation_task_events")
op.drop_index("idx_chat_generation_task_events_record_created", table_name="chat_generation_task_events")
op.drop_index("idx_chat_generation_task_events_attempt_created", table_name="chat_generation_task_events")
op.execute(
"UPDATE chat_generation_task_events "
"SET message = LEFT(message, 512), "
"from_stage = LEFT(from_stage, 32), "
"to_stage = LEFT(to_stage, 32)"
)
op.alter_column(
"chat_generation_task_events",
"message",
existing_type=sa.Text(),
type_=sa.VARCHAR(length=512),
existing_nullable=True,
)
op.alter_column(
"chat_generation_task_events",
"to_stage",
existing_type=sa.String(length=48),
type_=sa.VARCHAR(length=32),
existing_nullable=True,
)
op.alter_column(
"chat_generation_task_events",
"from_stage",
existing_type=sa.String(length=48),
type_=sa.VARCHAR(length=32),
existing_nullable=True,
)
op.alter_column(
"chat_generation_task_events",
"task_id",
existing_type=sa.VARCHAR(length=32),
nullable=False,
)
op.drop_column("chat_generation_task_events", "generation_attempt_no")
op.drop_column("chat_generation_task_events", "generation_record_id")
op.drop_column("chat_generation_task_events", "owner_type")
op.drop_index(
op.f("ix_chat_generation_tasks_resource_generation_started_at"),
table_name="chat_generation_tasks",
)
op.drop_index(
op.f("ix_chat_generation_tasks_poll_lease_until"),
table_name="chat_generation_tasks",
)
op.drop_index(
op.f("ix_chat_generation_tasks_poll_claim_token"),
table_name="chat_generation_tasks",
)
op.drop_index(
op.f("ix_chat_generation_tasks_download_claim_token"),
table_name="chat_generation_tasks",
)
for column_name in (
"download_claim_token",
"poll_lease_until",
"poll_claim_token",
"poll_error_count",
"manual_retry_count",
"resource_generation_started_at",
"generation_attempt_no",
):
op.drop_column("chat_generation_tasks", column_name)
for table_name in ("video_engines", "model_configs", "image_engines"):
op.drop_index(op.f(f"ix_{table_name}_deleted_at"), table_name=table_name)
op.drop_column(table_name, "deleted_at")
+172 -144
View File
@@ -45,6 +45,10 @@ from app.schemas.industry import IndustryConfigCreate, IndustryConfigOut
from app.schemas.video_engine import VideoEngineCreate, VideoEngineOut
from app.schemas.image_engine import ImageEngineCreate, ImageEngineOut
from app.schemas.credit_ratio import CreditRatioCreate, CreditRatioOut
from app.services.generation.pipeline.db_lock_service import (
DatabaseRowLockBusy,
execute_with_lock_timeout,
)
from app.services.credits import add_credits, deduct_credits
from app.services.credit_record_meta_service import build_admin_adjust_meta
from app.services.admin_credit_record_service import list_admin_credit_records
@@ -59,7 +63,7 @@ from app.services.team_service import batch_get_team_name_map, set_frontend_user
from app.services.generation.billing_service import (
OWNER_GENERATION_RECORD,
charge_generation_media_by_params,
charge_generation_media_for_record,
get_next_credit_attempt_no,
)
from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once
@@ -1117,11 +1121,15 @@ async def delete_industry_config(
@router.get("/video-engines", response_model=list[VideoEngineOut])
async def list_video_engines(
include_deleted: bool = Query(False),
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
query = select(VideoEngine)
if not include_deleted:
query = query.where(VideoEngine.deleted_at.is_(None))
result = await db.execute(
select(VideoEngine).order_by(VideoEngine.priority.desc())
query.order_by(VideoEngine.priority.desc(), VideoEngine.id.desc())
)
return result.scalars().all()
@@ -1161,7 +1169,7 @@ async def update_video_engine(
db: AsyncSession = Depends(get_db),
):
result = await db.execute(
select(VideoEngine).where(VideoEngine.id == engine_id).limit(1)
select(VideoEngine).where(VideoEngine.id == engine_id, VideoEngine.deleted_at.is_(None)).limit(1)
)
engine = result.scalar_one_or_none()
if not engine:
@@ -1194,24 +1202,26 @@ async def delete_video_engine(
db: AsyncSession = Depends(get_db),
):
result = await db.execute(
select(VideoEngine).where(VideoEngine.id == engine_id).limit(1)
select(VideoEngine).where(VideoEngine.id == engine_id, VideoEngine.deleted_at.is_(None)).limit(1)
)
engine = result.scalar_one_or_none()
if not engine:
raise HTTPException(status_code=404, detail="视频引擎不存在")
await db.delete(engine)
engine_name = engine.name
engine.deleted_at = datetime.now(timezone.utc)
engine.is_active = False
await db.flush()
await log_operation(
db,
admin.id,
admin.username,
f"删除视频引擎: {engine.name}",
f"删除视频引擎: {engine_name}",
"DELETE",
f"/admin/video-engines/{engine_id}",
detail=json.dumps(
{
"engine_id": engine_id,
"name": engine.name,
"name": engine_name,
},
ensure_ascii=False,
),
@@ -1223,11 +1233,15 @@ async def delete_video_engine(
@router.get("/image-engines", response_model=list[ImageEngineOut])
async def list_image_engines(
include_deleted: bool = Query(False),
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
query = select(ImageEngine)
if not include_deleted:
query = query.where(ImageEngine.deleted_at.is_(None))
result = await db.execute(
select(ImageEngine).order_by(ImageEngine.priority.desc())
query.order_by(ImageEngine.priority.desc(), ImageEngine.id.desc())
)
return result.scalars().all()
@@ -1267,7 +1281,7 @@ async def update_image_engine(
db: AsyncSession = Depends(get_db),
):
result = await db.execute(
select(ImageEngine).where(ImageEngine.id == engine_id).limit(1)
select(ImageEngine).where(ImageEngine.id == engine_id, ImageEngine.deleted_at.is_(None)).limit(1)
)
engine = result.scalar_one_or_none()
if not engine:
@@ -1300,24 +1314,26 @@ async def delete_image_engine(
db: AsyncSession = Depends(get_db),
):
result = await db.execute(
select(ImageEngine).where(ImageEngine.id == engine_id).limit(1)
select(ImageEngine).where(ImageEngine.id == engine_id, ImageEngine.deleted_at.is_(None)).limit(1)
)
engine = result.scalar_one_or_none()
if not engine:
raise HTTPException(status_code=404, detail="图片引擎不存在")
await db.delete(engine)
engine_name = engine.name
engine.deleted_at = datetime.now(timezone.utc)
engine.is_active = False
await db.flush()
await log_operation(
db,
admin.id,
admin.username,
f"删除图片引擎: {engine.name}",
f"删除图片引擎: {engine_name}",
"DELETE",
f"/admin/image-engines/{engine_id}",
detail=json.dumps(
{
"engine_id": engine_id,
"name": engine.name,
"name": engine_name,
},
ensure_ascii=False,
),
@@ -1342,7 +1358,7 @@ async def _validate_credit_ratio_engine(db: AsyncSession, req: CreditRatioCreate
raise HTTPException(status_code=400, detail="model_config_id 不能为空,当前字段用于保存图片/视频引擎ID")
model = ImageEngine if gen_type == "image" else VideoEngine
result = await db.execute(select(model).where(model.id == engine_id).limit(1))
result = await db.execute(select(model).where(model.id == engine_id, model.deleted_at.is_(None)).limit(1))
engine = result.scalar_one_or_none()
if not engine:
detail = "图片积分规则绑定的图片引擎不存在" if gen_type == "image" else "视频积分规则绑定的视频引擎不存在"
@@ -1488,10 +1504,14 @@ async def list_credit_ratios_grouped(
@router.get("/model-configs", response_model=list[ModelConfigOut])
async def list_model_configs(
include_deleted: bool = Query(False),
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(select(ModelConfig).order_by(ModelConfig.priority.desc()))
query = select(ModelConfig)
if not include_deleted:
query = query.where(ModelConfig.deleted_at.is_(None))
result = await db.execute(query.order_by(ModelConfig.priority.desc(), ModelConfig.id.desc()))
return result.scalars().all()
@@ -1529,7 +1549,7 @@ async def update_model_config(
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(select(ModelConfig).where(ModelConfig.id == config_id).limit(1))
result = await db.execute(select(ModelConfig).where(ModelConfig.id == config_id, ModelConfig.deleted_at.is_(None)).limit(1))
config = result.scalar_one_or_none()
if not config:
raise HTTPException(status_code=404, detail="配置不存在")
@@ -1560,23 +1580,25 @@ async def delete_model_config(
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(select(ModelConfig).where(ModelConfig.id == config_id).limit(1))
result = await db.execute(select(ModelConfig).where(ModelConfig.id == config_id, ModelConfig.deleted_at.is_(None)).limit(1))
config = result.scalar_one_or_none()
if not config:
raise HTTPException(status_code=404, detail="配置不存在")
await db.delete(config)
config_name = config.name
config.deleted_at = datetime.now(timezone.utc)
config.is_active = False
await db.flush()
await log_operation(
db,
admin.id,
admin.username,
f"删除模型配置: {config.name}",
f"删除模型配置: {config_name}",
"DELETE",
f"/admin/model-configs/{config_id}",
detail=json.dumps(
{
"config_id": config_id,
"name": config.name,
"name": config_name,
},
ensure_ascii=False,
),
@@ -1982,87 +2004,156 @@ async def admin_update_generation_status(
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""Admin update generation record status (e.g., confirm/reject)."""
result = await db.execute(
select(GenerationRecord).where(
GenerationRecord.id == record_id,
GenerationRecord.deleted_at.is_(None),
"""管理员只能终止正在执行或待生成的记录,禁止绕过流水线裸改生成/完成状态。"""
try:
result = await execute_with_lock_timeout(
db,
select(GenerationRecord).where(
GenerationRecord.id == record_id,
GenerationRecord.deleted_at.is_(None),
)
.with_for_update()
.limit(1),
)
.with_for_update()
.limit(1)
)
except DatabaseRowLockBusy as exc:
raise HTTPException(status_code=409, detail=exc.detail) from exc
record = result.scalar_one_or_none()
if not record:
raise HTTPException(status_code=404, detail="记录不存在")
new_status = body.get("status")
if new_status not in ("prompt_optimized", "generating", "completed", "failed"):
raise HTTPException(status_code=400, detail="无效状态")
if new_status == "failed":
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=body.get("error_message") or record.error_message or "管理员设置为失败",
new_status = str(body.get("status") or "").strip()
if new_status in {"generating", "completed", "prompt_optimized"}:
raise HTTPException(
status_code=409,
detail="禁止直接修改为该状态;生成请调用生成接口,完成必须由下载/超分流水线落库",
)
else:
record.status = new_status
if new_status != "failed":
raise HTTPException(status_code=400, detail="该接口仅允许管理员终止任务")
if record.status == "completed":
raise HTTPException(status_code=409, detail="已完成记录不能直接改为失败")
error_message = body.get("error_message") or record.error_message or "管理员终止生成任务"
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=error_message,
generation_attempt_no=int(record.generation_attempt_no or 1),
)
record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
record.provider_create_claim_token = None
record.provider_create_lease_until = None
record.poll_claim_token = None
record.poll_lease_until = None
record.next_poll_at = None
record.download_claim_token = None
record.download_lease_until = None
record.download_next_retry_at = None
# 若任务已进入超分,必须同时撤销超分数据库租约;执行中的超分 Worker
# 在回填前校验 lease_token,发现 token 被清除后会中止,不得覆盖管理员终止状态。
from app.enums.video_upscale import VideoUpscaleStage, VideoUpscaleTaskStatus
from app.models.video_upscale_task import VideoUpscaleTask
try:
upscale_result = await execute_with_lock_timeout(
db,
select(VideoUpscaleTask)
.where(VideoUpscaleTask.generation_record_id == record.id)
.with_for_update()
.limit(1),
)
except DatabaseRowLockBusy as exc:
raise HTTPException(status_code=409, detail=exc.detail) from exc
upscale = upscale_result.scalar_one_or_none()
if upscale and upscale.status not in {
VideoUpscaleTaskStatus.COMPLETED.value,
VideoUpscaleTaskStatus.FAILED.value,
}:
upscale.status = VideoUpscaleTaskStatus.FAILED.value
upscale.stage = VideoUpscaleStage.FAILED.value
upscale.last_error = error_message
upscale.failed_at = datetime.now(CST)
upscale.next_retry_at = None
upscale.lease_token = None
upscale.lease_until = None
if body.get("video_url"):
record.video_url = body["video_url"]
if body.get("video_cover_url"):
record.video_cover_url = body["video_cover_url"]
if body.get("image_url"):
record.image_url = body["image_url"]
if new_status == "completed":
record.generated_at = datetime.now(CST)
await db.flush()
await log_operation(
db,
admin.id,
admin.username,
f"更新生成记录状态: {new_status}",
"管理员终止生成记录",
"PUT",
f"/admin/generation-records/{record_id}/status",
detail=json.dumps(
{
"record_id": record_id,
"new_status": new_status,
"generation_attempt_no": int(record.generation_attempt_no or 1),
"error_message": error_message,
},
ensure_ascii=False,
),
)
await db.commit()
# Redis 注册表只做调度加速;删除失败不回滚已提交的业务终止状态。
try:
from app.services.celery_download_recovery_service import remove_download_active
from app.services.generation.pipeline.owner_service import redis_owner_item_id
from app.services.redis_registry_service import redis_remove_registry_item
from app.config import settings
registry_id = redis_owner_item_id(
"generation_record",
record_id,
int(record.generation_attempt_no or 1),
)
await remove_download_active(registry_id)
await redis_remove_registry_item(
hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY,
zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY,
item_id=registry_id,
log_context="admin_generation_record_terminate",
)
except Exception:
pass
return {"message": "ok"}
@router.post("/generation-records/{record_id}/generate")
async def admin_generate_video(
async def admin_generate_record_resource(
record_id: str,
body: dict,
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""Admin trigger video/image generation for a record with specified params."""
from app.models.project import Project
from app.services.video_queue import task_queue
result = await db.execute(
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(
GenerationRecord.id == record_id,
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
)
.with_for_update()
"""管理员触发 GenerationRecord 图片或视频资源生成。"""
from app.services.generation.pipeline.generation_record_service import (
commit_and_enqueue_generation_record,
prepare_generation_record_execution,
)
try:
result = await execute_with_lock_timeout(
db,
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(
GenerationRecord.id == record_id,
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
)
.with_for_update(),
)
except DatabaseRowLockBusy as exc:
raise HTTPException(status_code=409, detail=exc.detail) from exc
row = result.first()
if not row:
raise HTTPException(status_code=404, detail="记录不存在")
record, project_name = row
type_str = "视频" if record.gen_type == GenerationType.video else "图片"
if record.status not in ("prompt_optimized", "failed"):
raise HTTPException(status_code=400, detail=f"当前状态不允许生成{type_str}")
if record.pipeline_stage == GenerationRecordPipelineStage.UPSCALE_FAILED.value:
@@ -2075,7 +2166,6 @@ async def admin_generate_video(
)
if record.gen_type == GenerationType.video:
# Video Generation
aspect_ratio = body.get("aspect_ratio", "16:9")
resolution = body.get("resolution", "720p")
if aspect_ratio not in ASPECT_RATIOS:
@@ -2097,97 +2187,34 @@ async def admin_generate_video(
aspect_ratio=aspect_ratio,
supported_provider_resolutions=supported_provider_resolutions,
)
duration = record.duration or 5
media_billing = await charge_generation_media_by_params(
db,
user_id=record.user_id,
record_id=record.id,
gen_type="video",
duration=duration,
resolution=resolution,
project_name=project_name,
description_prefix="视频生成(管理后台)",
owner_type=OWNER_GENERATION_RECORD,
attempt_no=attempt_no,
)
record.aspect_ratio = aspect_ratio
record.resolution = resolution
record.provider_generation_resolution = provider_resolution
record.video_upscale_enabled_snapshot = upscale_enabled
record.video_upscale_snapshot_json = upscale_snapshot_json
record.pipeline_stage = GenerationRecordPipelineStage.CREATING_PROVIDER_TASK.value
record.credits_cost = round(float(record.credits_cost or 0) + media_billing.total_charged, 2)
record.status = "generating"
record.error_message = None
record.video_url = None
record.video_cover_url = None
record.image_url = None
record.seedance_task_id = None
await db.flush()
try:
from app.services.video_gen import submit_video_task
task_id = await submit_video_task(
db,
engine,
record,
include_media_references=False,
)
record.seedance_task_id = task_id
record.pipeline_stage = GenerationRecordPipelineStage.WAITING_REMOTE.value
await db.flush()
await task_queue.enqueue(record_id)
except Exception as e:
record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=str(e),
)
await db.flush()
elif record.gen_type == GenerationType.image:
# Image generation
from app.services.image_gen import get_active_image_engine
post_image_size = body.get("image_size", "")
image_size = post_image_size or record.image_size or "2K"
media_billing = await charge_generation_media_by_params(
db,
user_id=record.user_id,
record_id=record.id,
gen_type="image",
image_size=image_size,
project_name=project_name,
description_prefix="图片生成(管理后台)",
owner_type=OWNER_GENERATION_RECORD,
attempt_no=attempt_no,
)
record.image_size = image_size
record.credits_cost = round(float(record.credits_cost or 0) + media_billing.total_charged, 2)
record.status = "generating"
record.error_message = None
record.image_url = None
record.video_url = None
record.video_cover_url = None
record.seedance_task_id = None
engine = await get_active_image_engine(db)
record.image_size = body.get("image_size") or record.image_size or "2K"
record.provider_generation_resolution = None
record.video_upscale_enabled_snapshot = False
record.video_upscale_snapshot_json = None
record.pipeline_stage = None
await db.flush()
else:
raise HTTPException(status_code=400, detail="不支持的生成类型")
try:
await task_queue.enqueue(record_id)
except Exception as e:
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=str(e),
)
await db.flush()
media_billing = await charge_generation_media_for_record(
db,
record=record,
project_name=project_name,
description_prefix=f"{type_str}生成(管理后台)-",
attempt_no=attempt_no,
engine_id=engine.id,
)
record.credits_cost = round(float(record.credits_cost or 0) + float(media_billing.total_charged or 0), 2)
prepare_generation_record_execution(record, engine=engine, attempt_no=attempt_no)
await db.flush()
await commit_and_enqueue_generation_record(db, record, reason="generation_record_admin_generate")
await log_operation(
db,
@@ -2201,6 +2228,7 @@ async def admin_generate_video(
"record_id": record_id,
"gen_type": record.gen_type,
"project_name": project_name,
"generation_attempt_no": record.generation_attempt_no,
},
ensure_ascii=False,
),
+2 -2
View File
@@ -62,14 +62,14 @@ async def get_credit_ratios(
video_engines_result = await db.execute(
select(VideoEngine.id)
.where(VideoEngine.is_active == True)
.where(VideoEngine.is_active == True, VideoEngine.deleted_at.is_(None))
.order_by(VideoEngine.priority.desc())
)
video_engine_ids = video_engines_result.scalars().all()
image_engines_result = await db.execute(
select(ImageEngine.id)
.where(ImageEngine.is_active == True)
.where(ImageEngine.is_active == True, ImageEngine.deleted_at.is_(None))
.order_by(ImageEngine.priority.desc())
)
image_engine_ids = image_engines_result.scalars().all()
+97 -264
View File
@@ -28,6 +28,10 @@ from app.schemas.generation import (
RESOLUTIONS,
IMAGE_SIZES,
)
from app.services.generation.pipeline.db_lock_service import (
DatabaseRowLockBusy,
execute_with_lock_timeout,
)
from app.services.credits import deduct_credits, calc_text_credits
from app.services.llm import optimize_prompt
from app.services.video_url import generate_temp_url, validate_and_get_record_id, get_video_stream_url
@@ -288,7 +292,7 @@ async def optimize(
references=req.references,
gen_type=req.gen_type,
)
# Create record BEFORE LLM call so it's visible if user refreshes
# LLM 成功后再创建记录;LLM 失败不写 GenerationRecord。
record = GenerationRecord(
id=generate_id(),
user_id=current_user.id,
@@ -356,16 +360,22 @@ async def optimize(
# 提示词积分不足时,之前已落库的 optimizing 记录必须改为 failed,避免前端长期显示生成中。
# 此阶段没有媒体生成扣费,不调用生成失败退款逻辑。
await db.rollback()
result = await db.execute(
select(GenerationRecord)
.where(
GenerationRecord.id == failed_record_id,
GenerationRecord.user_id == failed_user_id,
GenerationRecord.deleted_at.is_(None),
try:
result = await execute_with_lock_timeout(
db,
select(GenerationRecord)
.where(
GenerationRecord.id == failed_record_id,
GenerationRecord.user_id == failed_user_id,
GenerationRecord.deleted_at.is_(None),
)
.with_for_update()
.limit(1),
)
.with_for_update()
.limit(1)
)
except DatabaseRowLockBusy:
# Preserve the original 402 response; a later admin/manual check can
# reconcile the rare record-state update lock conflict.
raise e
failed_record = result.scalar_one_or_none()
if failed_record:
failed_record.status = "failed"
@@ -395,27 +405,30 @@ async def optimize(
@router.post("/{record_id}/generate")
async def generate(
async def generate_record_resource(
record_id: str,
req: GenerateParams,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(
GenerationRecord.id == record_id,
GenerationRecord.user_id == current_user.id,
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
try:
result = await execute_with_lock_timeout(
db,
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(
GenerationRecord.id == record_id,
GenerationRecord.user_id == current_user.id,
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
)
.with_for_update(),
)
.with_for_update()
)
except DatabaseRowLockBusy as exc:
raise HTTPException(status_code=409, detail=exc.detail) from exc
row = result.first()
if not row:
raise RecordNotFoundError()
record, project_name = row
if record.status not in ("prompt_optimized", "failed"):
raise InvalidStatusError("当前状态不允许生成")
@@ -423,20 +436,18 @@ async def generate(
raise InvalidStatusError("该任务生成失败,请联系客服进行修复")
await assert_user_resource_capacity_available(db, current_user.id)
attempt_no = await get_next_credit_attempt_no(db, owner_type=OWNER_GENERATION_RECORD, owner_id=record.id)
attempt_no = await get_next_credit_attempt_no(
db,
owner_type=OWNER_GENERATION_RECORD,
owner_id=record.id,
from app.services.generation.pipeline.generation_record_service import (
commit_and_enqueue_generation_record,
prepare_generation_record_execution,
)
if record.gen_type == GenerationType.video:
# Video generation
if req.aspect_ratio not in ASPECT_RATIOS:
raise HTTPException(status_code=400, detail="不支持的画面比例")
if req.resolution not in RESOLUTIONS:
raise HTTPException(status_code=400, detail="不支持的分辨率")
from app.services.video_gen import get_active_engine
from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot
@@ -451,100 +462,39 @@ async def generate(
aspect_ratio=req.aspect_ratio,
supported_provider_resolutions=supported_provider_resolutions,
)
duration = record.duration or 5
media_billing = await charge_generation_media_by_params(
db,
user_id=current_user.id,
record_id=record.id,
gen_type="video",
duration=duration,
resolution=req.resolution,
project_name=project_name,
description_prefix=project_name+"-",
owner_type=OWNER_GENERATION_RECORD,
attempt_no=attempt_no,
billing = await charge_generation_media_by_params(
db, user_id=current_user.id, record_id=record.id, gen_type="video",
duration=record.duration or 5, resolution=req.resolution, engine_id=engine.id,
project_name=project_name, description_prefix=project_name + "-",
owner_type=OWNER_GENERATION_RECORD, attempt_no=attempt_no,
)
record.aspect_ratio = req.aspect_ratio
record.resolution = req.resolution
record.provider_generation_resolution = provider_resolution
record.video_upscale_enabled_snapshot = upscale_enabled
record.video_upscale_snapshot_json = upscale_snapshot_json
record.pipeline_stage = GenerationRecordPipelineStage.CREATING_PROVIDER_TASK.value
record.credits_cost = round(float(record.credits_cost or 0) + media_billing.total_charged, 2)
record.status = "generating"
record.error_message = None
record.video_url = None
record.video_cover_url = None
record.image_url = None
record.seedance_task_id = None
await db.flush()
try:
from app.services.video_gen import submit_video_task
from app.services.error_codes import extract_error_message
from app.services.video_queue import task_queue
task_id = await submit_video_task(
db,
engine,
record,
include_media_references=False,
)
record.seedance_task_id = task_id
record.pipeline_stage = GenerationRecordPipelineStage.WAITING_REMOTE.value
await db.flush()
await task_queue.enqueue(record_id)
except Exception as e:
record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=extract_error_message(e, "视频"),
)
await db.flush()
elif record.gen_type == GenerationType.image:
image_size = req.image_size or record.image_size or "2K"
media_billing = await charge_generation_media_by_params(
db,
user_id=current_user.id,
record_id=record.id,
gen_type="image",
image_size=image_size,
project_name=project_name,
description_prefix=project_name+"-",
owner_type=OWNER_GENERATION_RECORD,
attempt_no=attempt_no,
else:
from app.services.image_gen import get_active_image_engine
engine = await get_active_image_engine(db)
image_size = req.image_size or record.image_size or engine.default_size or "2K"
billing = await charge_generation_media_by_params(
db, user_id=current_user.id, record_id=record.id, gen_type="image",
image_size=image_size, engine_id=engine.id, project_name=project_name,
description_prefix=project_name + "-", owner_type=OWNER_GENERATION_RECORD, attempt_no=attempt_no,
)
record.image_size = image_size
record.credits_cost = round(float(record.credits_cost or 0) + media_billing.total_charged, 2)
record.status = "generating"
record.error_message = None
record.image_url = None
record.video_url = None
record.video_cover_url = None
record.seedance_task_id = None
record.provider_generation_resolution = None
record.video_upscale_enabled_snapshot = False
record.video_upscale_snapshot_json = None
record.pipeline_stage = None
await db.flush()
try:
from app.services.video_queue import task_queue
await task_queue.enqueue(record_id)
except Exception as e:
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=f"图片任务队列投递失败: {e}",
)
await db.flush()
record.credits_cost = round(float(record.credits_cost or 0) + float(billing.total_charged or 0), 2)
prepare_generation_record_execution(record, engine=engine, attempt_no=attempt_no)
await db.flush()
await commit_and_enqueue_generation_record(db, record, reason="generation_record_api_generate")
refs = await resolve_private_portrait_reference_display_urls(db, json.loads(record.media_references) if record.media_references else None, user_id=current_user.id)
refs = await resolve_private_portrait_reference_display_urls(
db, json.loads(record.media_references) if record.media_references else None, user_id=current_user.id
)
return _record_to_out(record, project_name, refs_override=refs)
@@ -554,21 +504,24 @@ async def retry_generation(
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(
GenerationRecord.id == record_id,
GenerationRecord.user_id == current_user.id,
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
try:
result = await execute_with_lock_timeout(
db,
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(
GenerationRecord.id == record_id,
GenerationRecord.user_id == current_user.id,
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
)
.with_for_update(),
)
.with_for_update()
)
except DatabaseRowLockBusy as exc:
raise HTTPException(status_code=409, detail=exc.detail) from exc
row = result.first()
if not row:
raise RecordNotFoundError()
record, project_name = row
if record.status != "failed":
raise InvalidStatusError("只有失败的记录可以重试")
@@ -576,78 +529,44 @@ async def retry_generation(
raise InvalidStatusError("该任务生成失败,请联系客服进行修复")
await assert_user_resource_capacity_available(db, current_user.id)
attempt_no = await get_next_credit_attempt_no(
db,
owner_type=OWNER_GENERATION_RECORD,
owner_id=record.id,
attempt_no = await get_next_credit_attempt_no(db, owner_type=OWNER_GENERATION_RECORD, owner_id=record.id)
from app.services.generation.pipeline.generation_record_service import (
commit_and_enqueue_generation_record,
prepare_generation_record_execution,
)
engine = None
if record.gen_type == GenerationType.video:
from app.services.video_gen import get_active_engine
from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot
engine = await get_active_engine(db)
try:
supported_provider_resolutions = json.loads(engine.supported_resolutions or "[]")
except (TypeError, json.JSONDecodeError):
supported_provider_resolutions = []
provider_resolution, upscale_enabled, upscale_snapshot_json = await build_video_upscale_snapshot(
db,
target_resolution=record.resolution or "",
aspect_ratio=record.aspect_ratio or "",
db, target_resolution=record.resolution or "480p", aspect_ratio=record.aspect_ratio or "16:9",
supported_provider_resolutions=supported_provider_resolutions,
)
record.provider_generation_resolution = provider_resolution
record.video_upscale_enabled_snapshot = upscale_enabled
record.video_upscale_snapshot_json = upscale_snapshot_json
record.pipeline_stage = GenerationRecordPipelineStage.CREATING_PROVIDER_TASK.value
else:
from app.services.image_gen import get_active_image_engine
engine = await get_active_image_engine(db)
media_billing = await charge_generation_media_for_record(
db,
record=record,
project_name=project_name,
description_prefix="视频重试",
attempt_no=attempt_no,
billing = await charge_generation_media_for_record(
db, record=record, project_name=project_name, description_prefix="资源生成重试-", attempt_no=attempt_no, engine_id=engine.id
)
record.status = "generating"
record.error_message = None
record.video_url = None
record.video_cover_url = None
record.image_url = None
record.seedance_task_id = None
record.generated_at = None
record.credits_cost = round(float(record.credits_cost or 0) + media_billing.total_charged, 2)
record.credits_cost = round(float(record.credits_cost or 0) + float(billing.total_charged or 0), 2)
record.manual_retry_count = int(record.manual_retry_count or 0) + 1
record.retry_count = int(record.manual_retry_count or 0)
prepare_generation_record_execution(record, engine=engine, attempt_no=attempt_no)
await db.flush()
await commit_and_enqueue_generation_record(db, record, reason="generation_record_api_retry")
try:
from app.services.video_queue import task_queue
if record.gen_type == GenerationType.video:
from app.services.video_gen import submit_video_task
assert engine is not None
task_id = await submit_video_task(
db,
engine,
record,
include_media_references=False,
)
record.seedance_task_id = task_id
record.pipeline_stage = GenerationRecordPipelineStage.WAITING_REMOTE.value
await db.flush()
await task_queue.enqueue(record_id)
except Exception as e:
from app.services.error_codes import extract_error_message
if record.gen_type == GenerationType.video:
record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=extract_error_message(e, "重试"),
)
await db.flush()
refs = await resolve_private_portrait_reference_display_urls(db, json.loads(record.media_references) if record.media_references else None, user_id=current_user.id)
refs = await resolve_private_portrait_reference_display_urls(
db, json.loads(record.media_references) if record.media_references else None, user_id=current_user.id
)
return _record_to_out(record, project_name, refs_override=refs)
@@ -720,11 +639,15 @@ async def get_queue_status(
estimated_wait_seconds = None
if record.status == "generating":
resource_started_at = record.resource_generation_started_at or record.created_at
ahead_result = await db.execute(
select(func.count(GenerationRecord.id)).where(
GenerationRecord.status == "generating",
GenerationRecord.deleted_at.is_(None),
GenerationRecord.created_at < record.created_at,
func.coalesce(
GenerationRecord.resource_generation_started_at,
GenerationRecord.created_at,
) < resource_started_at,
)
)
ahead = ahead_result.scalar() or 0
@@ -741,96 +664,6 @@ async def get_queue_status(
}
@router.post("/callbacks/seedance")
async def seedance_callback(request: Request, db: AsyncSession = Depends(get_db)):
"""Receive async callback from Seedance API."""
data = await request.json()
task_id = data.get("id")
task_status = data.get("status")
if not task_id:
return {"message": "ignored"}
result = await db.execute(
select(GenerationRecord).where(
GenerationRecord.seedance_task_id == task_id,
GenerationRecord.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
record = result.scalar_one_or_none()
if not record:
return {"message": "record not found"}
if record.status == "completed":
return {"message": "already completed"}
if str(record.pipeline_stage or "").startswith("upscale_"):
return {"message": "upscale already started"}
if task_status == "succeeded":
remote_url = str(data.get("content", {}).get("video_url", "") or "").strip()
if not remote_url:
record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
await mark_generation_record_failed_and_refund_once(
db, record=record, error_message="供应商回调成功但未返回视频地址"
)
else:
usage = data.get("usage", {}) if isinstance(data.get("usage"), dict) else {}
from app.services.video_queue import handle_generation_record_video_succeeded
try:
entered_upscale = await handle_generation_record_video_succeeded(
db,
record,
remote_url=remote_url,
provider_response=data,
video_tokens=usage.get("total_tokens", 0),
)
except Exception as exc:
await db.rollback()
result = await db.execute(
select(GenerationRecord).where(GenerationRecord.id == record.id).with_for_update().limit(1)
)
failed_record = result.scalar_one_or_none()
if failed_record:
failed_record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
await mark_generation_record_failed_and_refund_once(
db, record=failed_record, error_message=f"视频结果下载失败: {exc}"
)
entered_upscale = False
from app.services.video_gen import _log_video_response
_log_video_response(record.id, data)
if not entered_upscale and record.status == "completed":
from app.services.notification import create_notification
from app.api.v1.notifications import push_notification_to_user
notif = await create_notification(
db, record.user_id, "视频生成完成",
"您的视频已生成完成,可以查看了。", "video", record.id,
)
await push_notification_to_user(record.user_id, notif)
elif task_status == "failed":
error_message = data.get("error", "视频生成失败")
record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=error_message,
)
# Log callback response
from app.services.video_gen import _log_video_response
_log_video_response(record.id, data, error=record.error_message)
# Notify user
from app.services.notification import create_notification
from app.api.v1.notifications import push_notification_to_user
notif = await create_notification(
db, record.user_id, "视频生成失败",
f"视频生成失败:{record.error_message}", "video", record.id,
)
await push_notification_to_user(record.user_id, notif)
await db.flush()
return {"message": "ok"}
@router.post(
"/upload-image",
summary="上传 AI 创作普通参考图片",
+46 -10
View File
@@ -1,10 +1,11 @@
from datetime import datetime
from datetime import datetime, timedelta, timezone
from fastapi import APIRouter, Body, Depends, HTTPException, Path, Query
from sqlalchemy import select
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.dependencies import get_current_user, get_db
from app.models.chat_generation_task import ChatGenerationTask
from app.models.user import User
@@ -20,6 +21,10 @@ from app.schemas.generation_ai import (
GenerationAITaskListOut,
GenerationAITaskOut,
)
from app.services.generation.pipeline.db_lock_service import (
DatabaseRowLockBusy,
execute_with_lock_timeout,
)
from app.services.generation.ai.service import (
build_task_out_list,
list_generation_ai_engine_options,
@@ -715,13 +720,17 @@ async def retry_task(
if celery_app is None:
raise HTTPException(status_code=503, detail="Celery未启用:请配置 REDIS_URL 或 CELERY_BROKER_URL 后启动 worker")
result = await db.execute(
select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.user_id == current_user.id,
ChatGenerationTask.deleted_at.is_(None),
).with_for_update().limit(1)
)
try:
result = await execute_with_lock_timeout(
db,
select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.user_id == current_user.id,
ChatGenerationTask.deleted_at.is_(None),
).with_for_update().limit(1),
)
except DatabaseRowLockBusy as exc:
raise HTTPException(status_code=409, detail=exc.detail) from exc
task = result.scalar_one_or_none()
if not task:
raise HTTPException(status_code=404, detail="任务不存在")
@@ -779,7 +788,7 @@ async def retry_task(
enqueue_ids: list[str] = []
download_retry_ids: list[str] = []
for target in retry_targets:
if int(target.retry_count or 0) >= 3:
if int(target.manual_retry_count or 0) >= 3:
raise HTTPException(status_code=400, detail=f"任务 {target.id} 已超过最大重试次数")
is_download_retry = bool(
@@ -811,6 +820,17 @@ async def retry_task(
quantity=quantity,
)
target.credits_cost = round(float(target.credits_cost or 0) + media_billing.total_charged, 2)
resource_started_at = datetime.now(timezone.utc)
target.generation_attempt_no = int(attempt_no)
target.resource_generation_started_at = resource_started_at
if target.gen_type == "image":
target.deadline_at = resource_started_at + timedelta(
minutes=int(settings.CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES or 30)
)
else:
target.deadline_at = resource_started_at + timedelta(
hours=int(settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS or 24)
)
target.provider_task_id = None
target.seedance_task_id = None
target.remote_result_url = None
@@ -818,6 +838,21 @@ async def retry_task(
target.provider_create_claim_token = None
target.provider_create_lease_until = None
target.provider_create_started_at = None
target.poll_started_at = None
target.poll_claim_token = None
target.poll_lease_until = None
target.poll_error_count = 0
target.next_poll_at = None
target.poll_interval_seconds = 0
target.download_celery_task_id = None
target.download_enqueued_at = None
target.download_started_at = None
target.download_claim_token = None
target.download_lease_until = None
target.download_next_retry_at = None
target.download_attempt_count = 0
target.download_last_error = None
target.download_storage_date_dir = None
target.image_url = None
target.video_url = None
target.video_cover_url = None
@@ -832,7 +867,8 @@ async def retry_task(
target.poll_count = 0
target.last_poll_at = None
target.generated_at = None
target.retry_count = int(target.retry_count or 0) + 1
target.manual_retry_count = int(target.manual_retry_count or 0) + 1
target.retry_count = int(target.manual_retry_count or 0)
if retrying_group_children:
await db.flush()
@@ -10,6 +10,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_current_user, get_db
from app.models.user import User
from app.enums.common import ModuleProjectStatusEnum
from app.enums.generation_task import GenerationOwnerType
from app.enums.hot_opening_replicate import HotOpeningLogEventEnum, HotOpeningStepCodeEnum, ModuleCodeEnum
from app.schemas.hot_opening_replicate import (
HotOpeningActionOut,
@@ -610,14 +611,21 @@ async def generate_image(
from app.tasks.generation_create_tasks import chatapi_create_generation_task
try:
chatapi_create_generation_task.delay(chat_task_id_value)
chatapi_create_generation_task.apply_async(
args=[chat_task_id_value],
kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": 1},
queue="gen_chatapi_create",
countdown=0,
)
except Exception as exc:
await _mark_dispatch_failed_and_raise(
db,
_log_api_error(
event_type=HotOpeningLogEventEnum.CELERY_DISPATCH_FAILED.value,
current_user=current_user,
project_id=project_id_value,
step_id=step_id_value,
message=f"图片生成任务投递失败: {exc}",
message=f"图片生成任务投递失败,等待生成恢复任务补投: {exc}",
detail={"recoverable": True, "chat_task_id": chat_task_id_value},
exc=exc,
)
return HotOpeningActionOut(
@@ -747,14 +755,21 @@ async def generate_video(
from app.tasks.generation_create_tasks import chatapi_create_generation_task
try:
chatapi_create_generation_task.delay(chat_task_id_value)
chatapi_create_generation_task.apply_async(
args=[chat_task_id_value],
kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": 1},
queue="gen_chatapi_create",
countdown=0,
)
except Exception as exc:
await _mark_dispatch_failed_and_raise(
db,
_log_api_error(
event_type=HotOpeningLogEventEnum.CELERY_DISPATCH_FAILED.value,
current_user=current_user,
project_id=project_id_value,
step_id=step_id_value,
message=f"视频生成任务投递失败: {exc}",
message=f"视频生成任务投递失败,等待生成恢复任务补投: {exc}",
detail={"recoverable": True, "chat_task_id": chat_task_id_value},
exc=exc,
)
return HotOpeningActionOut(
+1 -1
View File
@@ -21,7 +21,7 @@ async def list_active_engines(
"""Public endpoint returning active image engine capabilities."""
result = await db.execute(
select(ImageEngine)
.where(ImageEngine.is_active == True)
.where(ImageEngine.is_active == True, ImageEngine.deleted_at.is_(None))
.order_by(ImageEngine.priority.desc())
)
engines = result.scalars().all()
+31 -4
View File
@@ -10,6 +10,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.dependencies import get_current_user, get_db
from app.models.user import User
from app.enums.generation_task import GenerationOwnerType
from app.enums.shot_replicate import (
ModuleCodeEnum,
ShotAnalysisStatusEnum,
@@ -1061,9 +1062,22 @@ async def generate_image(
try:
from app.tasks.generation_create_tasks import chatapi_create_generation_task
chatapi_create_generation_task.apply_async(args=[chat_task_id_value], queue="gen_chatapi_create", countdown=0)
chatapi_create_generation_task.apply_async(
args=[chat_task_id_value],
kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": 1},
queue="gen_chatapi_create",
countdown=0,
)
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}")
_log_api_error(
event_type=ShotReplicateLogEventEnum.CELERY_DISPATCH_FAILED.value,
current_user=current_user,
project_id=project_id_value,
step_id=step_id_value,
message=f"图片生成任务投递失败,等待生成恢复任务补投: {exc}",
detail={"recoverable": True, "chat_task_id": chat_task_id_value},
exc=exc,
)
return ShotReplicateActionOut(message="图片生成任务已提交", project_id=project_id_value, step_id=step_id_value, detail=await _reload_project_detail(db, current_user, project_id_value))
@@ -1148,9 +1162,22 @@ async def generate_video(
try:
from app.tasks.generation_create_tasks import chatapi_create_generation_task
chatapi_create_generation_task.apply_async(args=[chat_task_id_value], queue="gen_chatapi_create", countdown=0)
chatapi_create_generation_task.apply_async(
args=[chat_task_id_value],
kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": 1},
queue="gen_chatapi_create",
countdown=0,
)
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}")
_log_api_error(
event_type=ShotReplicateLogEventEnum.CELERY_DISPATCH_FAILED.value,
current_user=current_user,
project_id=project_id_value,
step_id=step_id_value,
message=f"视频生成任务投递失败,等待生成恢复任务补投: {exc}",
detail={"recoverable": True, "chat_task_id": chat_task_id_value},
exc=exc,
)
return ShotReplicateActionOut(message="视频生成任务已提交", project_id=project_id_value, step_id=step_id_value, detail=await _reload_project_detail(db, current_user, project_id_value))
+1 -1
View File
@@ -21,7 +21,7 @@ async def list_active_engines(
"""Public endpoint returning active video engine capabilities."""
result = await db.execute(
select(VideoEngine)
.where(VideoEngine.is_active == True)
.where(VideoEngine.is_active == True, VideoEngine.deleted_at.is_(None))
.order_by(VideoEngine.priority.desc())
)
engines = result.scalars().all()
+20 -1
View File
@@ -117,6 +117,9 @@ class Settings(BaseSettings):
VIDEO_UPSCALE_RETRY_BACKOFF_SECONDS: int = 60
VIDEO_UPSCALE_RECOVERY_BATCH_SIZE: int = 50
VIDEO_UPSCALE_RECOVERY_LOCK_KEY: str = "vg:celery:video_upscale_recovery_lock"
VIDEO_UPSCALE_EXECUTION_LOCK_KEY_PREFIX: str = "vg:lock:upscale:execute"
VIDEO_UPSCALE_EXECUTION_LOCK_TTL_SECONDS: int = 30 * 60
VIDEO_UPSCALE_STAGE_HANDOFF_DELAY_SECONDS: int = 2
VIDEO_COVER_SEEK_TIME: str = "00:00:01"
VIDEO_COVER_FALLBACK_SEEK_TIME: str = "00:00:00"
@@ -149,7 +152,7 @@ class Settings(BaseSettings):
CHATAPI_ASYNC_MAX_RETRIES: int = 3
CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS: int = 30
CHATAPI_ASYNC_POLL_INTERVAL_SECONDS: int = 30
CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES: int = 10
CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES: int = 30
# 视频异步生成不再使用 30 分钟最终超时;前 10 分钟高频轮询,之后降频,24 小时最后判定失败才退款。
CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS: int = 24
CHATAPI_ASYNC_VIDEO_HIGH_FREQ_MINUTES: int = 10
@@ -228,6 +231,22 @@ class Settings(BaseSettings):
POLL_DUE_DISPATCH_LOCK_KEY: str = "vg:celery:poll_due_dispatch_lock"
POLL_DUE_DISPATCH_LOCK_TTL_SECONDS: int = 55
# Redis execution locks. These locks are fail-closed: when Redis is
# unavailable, the current Celery task retries and does not fall back to an
# unlocked database-only execution path.
GENERATION_CREATE_LOCK_KEY_PREFIX: str = "vg:lock:generation:create"
GENERATION_POLL_LOCK_KEY_PREFIX: str = "vg:lock:generation:poll"
GENERATION_DOWNLOAD_LOCK_KEY_PREFIX: str = "vg:lock:generation:download"
GENERATION_CREATE_LOCK_TTL_SECONDS: int = 10 * 60
GENERATION_POLL_LOCK_TTL_SECONDS: int = 5 * 60
GENERATION_DOWNLOAD_LOCK_TTL_SECONDS: int = 10 * 60
REDIS_EXECUTION_LOCK_RENEW_INTERVAL_SECONDS: int = 30
# PostgreSQL short row-lock wait and whole ordinary download timeout.
# The lock timeout is transaction-local (SET LOCAL), not a global DB setting.
GENERATION_DB_LOCK_TIMEOUT_SECONDS: int = 5
GENERATION_DOWNLOAD_TOTAL_TIMEOUT_SECONDS: int = 8 * 60
# 模块异步任务容灾配置。
# 覆盖 ModuleGenerationStep 提词任务、shot 原视频/片段分析、shot ffmpeg 切割 active 注册。
# 恢复扫描走 CELERY_RECOVERY_QUEUE,真实业务任务回到原始队列。
+7 -1
View File
@@ -10,13 +10,17 @@ class GenerationStatus(str, Enum):
class GenerationRecordPipelineStage(str, Enum):
"""GenerationRecord 视频生成与超分流水线阶段。"""
"""GenerationRecord 图片/视频生成、下载与超分流水线阶段。"""
QUEUED = "queued"
PREPARING = "preparing"
CREATING_PROVIDER_TASK = "creating_provider_task"
WAITING_REMOTE = "waiting_remote"
POLLING = "polling"
RESULT_READY = "result_ready"
DOWNLOAD_QUEUED = "download_queued"
DOWNLOADING = "downloading"
RETRY_WAITING = "retry_waiting"
UPSCALE_QUEUED = "upscale_queued"
UPSCALE_PROCESSING = "upscale_processing"
UPSCALE_POLLING = "upscale_polling"
@@ -26,6 +30,8 @@ class GenerationRecordPipelineStage(str, Enum):
UPSCALE_FAILED = "upscale_failed"
DONE = "done"
FAILED = "failed"
TIMEOUT = "timeout"
DOWNLOAD_FAILED = "download_failed"
class GenerationType(str, Enum):
@@ -1,12 +1,18 @@
from enum import Enum
class GenerationOwnerType(str, Enum):
CHAT_GENERATION_TASK = "chat_generation_task"
GENERATION_RECORD = "generation_record"
class GenerationMode(str, Enum):
CHATAPI_ASYNC = "chatapi_async"
CHATAPI_MAIN = "chatapi_main"
CHATAPI_CHILD = "chatapi_child"
HOT_OPENING_REPLICATE = "hot_opening_replicate"
SHOT_REPLICATE = "shot_replicate"
GENERATION_RECORD = "generation_record"
class GenerationType(str, Enum):
@@ -138,6 +144,9 @@ class ChatGenerationTaskEventType(str, Enum):
DOWNLOAD_SKIP_RETRY_NOT_DUE = "DOWNLOAD_SKIP_RETRY_NOT_DUE"
DOWNLOAD_SKIP_FINAL_STATE = "DOWNLOAD_SKIP_FINAL_STATE"
DOWNLOAD_SKIP_DISABLED = "DOWNLOAD_SKIP_DISABLED"
STALE_ATTEMPT_MESSAGE_SKIPPED = "STALE_ATTEMPT_MESSAGE_SKIPPED"
GENERATION_RECORD_ENQUEUE_SUCCESS = "GENERATION_RECORD_ENQUEUE_SUCCESS"
GENERATION_RECORD_ENQUEUE_FAILED = "GENERATION_RECORD_ENQUEUE_FAILED"
TASK_TIMEOUT = "TASK_TIMEOUT"
TASK_FAILED = "TASK_FAILED"
+2 -9
View File
@@ -31,11 +31,6 @@ async def lifespan(app: FastAPI):
await init_redis()
# await _seed_data()
# Start task queue (handles both video and image generation)
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
@@ -104,8 +99,6 @@ async def lifespan(app: FastAPI):
yield
task_queue.stop()
await queue_task
upload_queue.stop()
await upload_queue_task
pre_test_queue.stop()
@@ -315,7 +308,7 @@ async def _seed_data():
default_video_engine_result = await db.execute(
select(VideoEngine)
.where(VideoEngine.is_active == True)
.where(VideoEngine.is_active == True, VideoEngine.deleted_at.is_(None))
.order_by(VideoEngine.priority.desc(), VideoEngine.id.desc())
.limit(1)
)
@@ -348,7 +341,7 @@ async def _seed_data():
default_image_engine_result = await db.execute(
select(ImageEngine)
.where(ImageEngine.is_active == True)
.where(ImageEngine.is_active == True, ImageEngine.deleted_at.is_(None))
.order_by(ImageEngine.priority.desc(), ImageEngine.id.desc())
.limit(1)
)
@@ -91,6 +91,8 @@ class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin):
)
generation_count: Mapped[int] = mapped_column(Integer, default=1, server_default="1", nullable=False)
generation_index: Mapped[int | None] = mapped_column(Integer, nullable=True)
generation_attempt_no: Mapped[int] = mapped_column(Integer, nullable=False, default=1, server_default="1")
resource_generation_started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True, index=True)
# 图片主任务同步调用供应商时的分布式执行租约。
# 防止重复 Celery 消息或恢复任务同时触发多次组图请求。
@@ -118,7 +120,10 @@ class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin):
video_tokens_used: Mapped[int] = mapped_column(Integer, default=0)
image_tokens_used: Mapped[int] = mapped_column(Integer, default=0)
# retry_count is retained as a compatibility mirror of manual retries.
retry_count: Mapped[int] = mapped_column(Integer, default=0)
manual_retry_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default="0")
poll_error_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default="0")
poll_count: Mapped[int] = mapped_column(Integer, default=0)
last_poll_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
# 视频降频轮询调度字段。
@@ -126,6 +131,8 @@ class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin):
poll_started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
next_poll_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
poll_interval_seconds: Mapped[int] = mapped_column(Integer, default=0)
poll_claim_token: Mapped[str | None] = mapped_column(String(64), nullable=True, index=True)
poll_lease_until: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True, index=True)
deadline_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
generated_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
error_message: Mapped[str | None] = mapped_column(Text, nullable=True)
@@ -135,6 +142,7 @@ class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin):
download_celery_task_id: Mapped[str | None] = mapped_column(String(160), nullable=True, index=True)
download_enqueued_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
download_started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
download_claim_token: Mapped[str | None] = mapped_column(String(64), nullable=True, index=True)
download_lease_until: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True, index=True)
download_next_retry_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True, index=True)
download_attempt_count: Mapped[int] = mapped_column(Integer, default=0)
@@ -1,4 +1,4 @@
from sqlalchemy import DateTime, ForeignKey, String, Text, func
from sqlalchemy import CheckConstraint, DateTime, ForeignKey, Index, Integer, String, Text, func
from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base
@@ -6,20 +6,31 @@ from app.utils.id_gen import generate_id
class ChatGenerationTaskEvent(Base):
"""Append-only event log for project-independent chat generation tasks."""
"""Append-only event log shared by ChatGenerationTask and GenerationRecord."""
__tablename__ = "chat_generation_task_events"
__table_args__ = (
CheckConstraint(
"(owner_type = 'chat_generation_task' AND task_id IS NOT NULL AND generation_record_id IS NULL) "
"OR (owner_type = 'generation_record' AND task_id IS NULL AND generation_record_id IS NOT NULL)",
name="ck_chat_generation_task_events_owner",
),
Index("idx_chat_generation_task_events_task_created", "owner_type", "task_id", "created_at"),
Index("idx_chat_generation_task_events_record_created", "owner_type", "generation_record_id", "created_at"),
Index("idx_chat_generation_task_events_attempt_created", "owner_type", "generation_attempt_no", "created_at"),
)
id: Mapped[str] = mapped_column(String(32), primary_key=True, default=generate_id)
task_id: Mapped[str] = mapped_column(
String(32), ForeignKey("chat_generation_tasks.id", ondelete="CASCADE"), index=True
)
owner_type: Mapped[str] = mapped_column(String(32), nullable=False, default="chat_generation_task", server_default="chat_generation_task", index=True)
task_id: Mapped[str | None] = mapped_column(String(32), ForeignKey("chat_generation_tasks.id", ondelete="CASCADE"), nullable=True, index=True)
generation_record_id: Mapped[str | None] = mapped_column(String(32), ForeignKey("generation_records.id", ondelete="CASCADE"), nullable=True, index=True)
generation_attempt_no: Mapped[int] = mapped_column(Integer, nullable=False, default=1, server_default="1", index=True)
generation_mode: Mapped[str | None] = mapped_column(String(32), nullable=True, index=True)
event_type: Mapped[str] = mapped_column(String(64), index=True)
from_status: Mapped[str | None] = mapped_column(String(32), nullable=True)
to_status: Mapped[str | None] = mapped_column(String(32), nullable=True)
from_stage: Mapped[str | None] = mapped_column(String(32), nullable=True)
to_stage: Mapped[str | None] = mapped_column(String(32), nullable=True)
message: Mapped[str | None] = mapped_column(String(512), nullable=True)
from_stage: Mapped[str | None] = mapped_column(String(48), nullable=True)
to_stage: Mapped[str | None] = mapped_column(String(48), nullable=True)
message: Mapped[str | None] = mapped_column(Text, nullable=True)
detail_json: Mapped[str | None] = mapped_column(Text, nullable=True)
created_at = mapped_column(DateTime(timezone=True), server_default=func.now(), index=True)
@@ -1,4 +1,4 @@
from sqlalchemy import DateTime, ForeignKey, Integer, String, Text, func
from sqlalchemy import CheckConstraint, DateTime, ForeignKey, Index, Integer, String, Text, func
from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base
@@ -6,14 +6,25 @@ from app.utils.id_gen import generate_id
class ChatProviderCallLog(Base):
"""Provider call audit log for chat_generation_tasks."""
"""Provider call audit log shared by both generation owner models."""
__tablename__ = "chat_provider_call_logs"
__table_args__ = (
CheckConstraint(
"(owner_type = 'chat_generation_task' AND task_id IS NOT NULL AND generation_record_id IS NULL) "
"OR (owner_type = 'generation_record' AND task_id IS NULL AND generation_record_id IS NOT NULL)",
name="ck_chat_provider_call_logs_owner",
),
Index("idx_chat_provider_call_logs_task_created", "owner_type", "task_id", "created_at"),
Index("idx_chat_provider_call_logs_record_created", "owner_type", "generation_record_id", "created_at"),
Index("idx_chat_provider_call_logs_attempt_created", "owner_type", "generation_attempt_no", "created_at"),
)
id: Mapped[str] = mapped_column(String(32), primary_key=True, default=generate_id)
task_id: Mapped[str] = mapped_column(
String(32), ForeignKey("chat_generation_tasks.id", ondelete="CASCADE"), index=True
)
owner_type: Mapped[str] = mapped_column(String(32), nullable=False, default="chat_generation_task", server_default="chat_generation_task", index=True)
task_id: Mapped[str | None] = mapped_column(String(32), ForeignKey("chat_generation_tasks.id", ondelete="CASCADE"), nullable=True, index=True)
generation_record_id: Mapped[str | None] = mapped_column(String(32), ForeignKey("generation_records.id", ondelete="CASCADE"), nullable=True, index=True)
generation_attempt_no: Mapped[int] = mapped_column(Integer, nullable=False, default=1, server_default="1", index=True)
generation_mode: Mapped[str | None] = mapped_column(String(32), nullable=True, index=True)
provider: Mapped[str | None] = mapped_column(String(64), nullable=True, index=True)
api_type: Mapped[str] = mapped_column(String(64), index=True)
@@ -1,6 +1,6 @@
from datetime import date, datetime
from sqlalchemy import BigInteger, Date, DateTime, ForeignKey, Index, String, Text
from sqlalchemy import BigInteger, Date, DateTime, ForeignKey, Index, String, Text, text
from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base, TimestampMixin, SoftDeleteMixin
@@ -44,4 +44,12 @@ class GeneratedResource(Base, TimestampMixin, SoftDeleteMixin):
Index("ix_generated_resources_user_month", GeneratedResource.user_id, GeneratedResource.resource_month)
Index("ix_generated_resources_source", GeneratedResource.source_model, GeneratedResource.source_id)
Index(
"uq_generated_resources_active_source_type",
GeneratedResource.source_model,
GeneratedResource.source_id,
GeneratedResource.resource_type,
unique=True,
postgresql_where=text("deleted_at IS NULL"),
)
Index("ix_generated_resources_active_user", GeneratedResource.user_id, GeneratedResource.deleted_at)
+43 -1
View File
@@ -1,6 +1,6 @@
from datetime import datetime
from sqlalchemy import Boolean, DateTime, ForeignKey, Integer, String, Text, Float, Index
from sqlalchemy import Boolean, DateTime, ForeignKey, Integer, String, Text, Float, Index, text
from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base, TimestampMixin, SoftDeleteMixin
@@ -52,7 +52,49 @@ class GenerationRecord(Base, TimestampMixin, SoftDeleteMixin):
error_message: Mapped[str | None] = mapped_column(Text, nullable=True)
idempotency_key: Mapped[str | None] = mapped_column(String(64), nullable=True, index=True)
# Unified Celery generation pipeline state.
generation_attempt_no: Mapped[int] = mapped_column(Integer, nullable=False, default=1, server_default="1")
resource_generation_started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True, index=True)
deadline_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True, index=True)
engine_id: Mapped[str | None] = mapped_column(String(32), nullable=True, index=True)
engine_snapshot_json: Mapped[str | None] = mapped_column(Text, nullable=True)
provider_response_json: Mapped[str | None] = mapped_column(Text, nullable=True)
remote_result_url: Mapped[str | None] = mapped_column(Text, nullable=True)
provider_create_claim_token: Mapped[str | None] = mapped_column(String(64), nullable=True, index=True)
provider_create_lease_until: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True, index=True)
provider_create_started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
# retry_count is retained for API/backward compatibility. New business
# logic uses manual_retry_count and poll_error_count separately.
retry_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default="0")
manual_retry_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default="0")
poll_error_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default="0")
poll_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default="0")
last_poll_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
poll_started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
next_poll_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
poll_interval_seconds: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default="0")
poll_claim_token: Mapped[str | None] = mapped_column(String(64), nullable=True, index=True)
poll_lease_until: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True, index=True)
download_celery_task_id: Mapped[str | None] = mapped_column(String(160), nullable=True, index=True)
download_enqueued_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
download_started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
download_claim_token: Mapped[str | None] = mapped_column(String(64), nullable=True, index=True)
download_lease_until: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True, index=True)
download_next_retry_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True, index=True)
download_attempt_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default="0")
download_last_error: Mapped[str | None] = mapped_column(Text, nullable=True)
download_storage_date_dir: Mapped[str | None] = mapped_column(String(16), nullable=True)
__table_args__ = (
Index('idx_genrec_user_status_created', 'user_id', 'status', 'created_at'),
Index('idx_genrec_project_status', 'project_id', 'status'),
Index(
'idx_genrec_next_poll_at',
'next_poll_at',
postgresql_where=text("deleted_at IS NULL AND status = 'generating' AND gen_type = 'video' AND next_poll_at IS NOT NULL"),
),
)
+2 -2
View File
@@ -1,10 +1,10 @@
from sqlalchemy import Boolean, CheckConstraint, Integer, String, Text
from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base, TimestampMixin
from app.models.base import Base, TimestampMixin, SoftDeleteMixin
class ImageEngine(Base, TimestampMixin):
class ImageEngine(Base, TimestampMixin, SoftDeleteMixin):
__tablename__ = "image_engines"
__table_args__ = (
CheckConstraint("max_generation_count BETWEEN 1 AND 5", name="ck_image_engines_max_generation_count"),
+2 -2
View File
@@ -1,10 +1,10 @@
from sqlalchemy import Boolean, Float, Integer, String
from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base, TimestampMixin
from app.models.base import Base, TimestampMixin, SoftDeleteMixin
class ModelConfig(Base, TimestampMixin):
class ModelConfig(Base, TimestampMixin, SoftDeleteMixin):
__tablename__ = "model_configs"
id: Mapped[str] = mapped_column(String(32), primary_key=True)
+2 -2
View File
@@ -1,10 +1,10 @@
from sqlalchemy import Boolean, CheckConstraint, Integer, String
from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base, TimestampMixin
from app.models.base import Base, TimestampMixin, SoftDeleteMixin
class VideoEngine(Base, TimestampMixin):
class VideoEngine(Base, TimestampMixin, SoftDeleteMixin):
__tablename__ = "video_engines"
__table_args__ = (
CheckConstraint("max_generation_count BETWEEN 1 AND 5", name="ck_video_engines_max_generation_count"),
+1
View File
@@ -25,6 +25,7 @@ class ModelConfigCreate(BaseModel):
class ModelConfigOut(ModelConfigCreate):
id: str
created_at: NaiveDatetime
deleted_at: NaiveDatetimeOptional = None
model_config = {"from_attributes": True}
+2 -1
View File
@@ -1,7 +1,7 @@
from pydantic import BaseModel, Field, model_validator
from app.enums.generation_provider import IMAGE_MULTI_OUTPUT_MAX, IMAGE_MULTI_REFERENCE_MAX
from app.schemas.common import NaiveDatetime
from app.schemas.common import NaiveDatetime, NaiveDatetimeOptional
class ImageEngineCreate(BaseModel):
@@ -60,6 +60,7 @@ class ImageEngineCreate(BaseModel):
class ImageEngineOut(ImageEngineCreate):
id: str
created_at: NaiveDatetime
deleted_at: NaiveDatetimeOptional = None
model_config = {"from_attributes": True}
+2 -1
View File
@@ -1,6 +1,6 @@
from pydantic import BaseModel, Field
from app.schemas.common import NaiveDatetime
from app.schemas.common import NaiveDatetime, NaiveDatetimeOptional
class VideoEngineCreate(BaseModel):
@@ -37,6 +37,7 @@ class VideoEngineCreate(BaseModel):
class VideoEngineOut(VideoEngineCreate):
id: str
created_at: NaiveDatetime
deleted_at: NaiveDatetimeOptional = None
model_config = {"from_attributes": True}
+2 -2
View File
@@ -81,7 +81,7 @@ async def calc_video_credits(
if not engine_id:
video_engines_result = await db.execute(
select(VideoEngine.id)
.where(VideoEngine.is_active == True)
.where(VideoEngine.is_active == True, VideoEngine.deleted_at.is_(None))
.order_by(VideoEngine.priority.desc())
.limit(1)
)
@@ -144,7 +144,7 @@ async def calc_image_credits(
if not engine_id:
image_engines_result = await db.execute(
select(ImageEngine.id)
.where(ImageEngine.is_active == True)
.where(ImageEngine.is_active == True, ImageEngine.deleted_at.is_(None))
.order_by(ImageEngine.priority.desc())
.limit(1)
)
@@ -50,7 +50,7 @@ def normalize_generation_count(value: int | None) -> int:
async def get_image_engine(db: AsyncSession, engine_id: str | None) -> ImageEngine:
query = select(ImageEngine).where(ImageEngine.is_active == True)
query = select(ImageEngine).where(ImageEngine.is_active == True, ImageEngine.deleted_at.is_(None))
if engine_id:
query = query.where(ImageEngine.id == engine_id)
else:
@@ -63,7 +63,7 @@ async def get_image_engine(db: AsyncSession, engine_id: str | None) -> ImageEngi
async def get_video_engine(db: AsyncSession, engine_id: str | None) -> VideoEngine:
query = select(VideoEngine).where(VideoEngine.is_active == True)
query = select(VideoEngine).where(VideoEngine.is_active == True, VideoEngine.deleted_at.is_(None))
if engine_id:
query = query.where(VideoEngine.id == engine_id)
else:
@@ -4,7 +4,7 @@ import json
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from uuid import uuid4
from typing import Awaitable, Callable
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
@@ -18,6 +18,7 @@ from app.enums.generation_task import (
GenerationType,
)
from app.models.chat_generation_task import ChatGenerationTask
from app.services.generation.pipeline.db_lock_service import execute_with_lock_timeout
from app.services.generation.ai.task_group_service import aggregate_main_task_status, load_children_map
from app.services.generation.log_service import log_task_event
from app.services.generation.provider_service import (
@@ -27,6 +28,7 @@ from app.services.generation.provider_service import (
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.image_gen import ImageProviderError
from app.services.operation_log_service import build_exception_detail, log_operation_event
from app.services.redis_registry_service import RedisExecutionLockError
from app.utils.id_gen import generate_id
@@ -84,8 +86,15 @@ def _task_snapshot(main: ChatGenerationTask) -> SimpleNamespace:
)
async def _claim_image_main_batch(db: AsyncSession, main_task_id: str) -> ImageBatchClaim:
result = await db.execute(
async def _claim_image_main_batch(
db: AsyncSession,
main_task_id: str,
*,
execution_token: str,
) -> ImageBatchClaim:
result = await execute_with_lock_timeout(
db,
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == main_task_id,
@@ -146,7 +155,7 @@ async def _claim_image_main_batch(db: AsyncSession, main_task_id: str) -> ImageB
await db.commit()
return ImageBatchClaim(False, main_task_id, reason="deadline_expired")
claim_token = uuid4().hex
claim_token = execution_token
main.provider_create_claim_token = claim_token
main.provider_create_started_at = now
main.provider_create_lease_until = now + timedelta(seconds=IMAGE_PROVIDER_CLAIM_LEASE_SECONDS)
@@ -232,7 +241,9 @@ async def _fail_claimed_main(
await db.rollback()
except Exception:
pass
result = await db.execute(
result = await execute_with_lock_timeout(
db,
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == main_task_id,
@@ -296,7 +307,9 @@ async def _split_children(
provider_result: dict,
provider_items: list[dict],
) -> list[str]:
result = await db.execute(
result = await execute_with_lock_timeout(
db,
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == main_task_id,
@@ -343,8 +356,12 @@ async def _split_children(
index = int(item.get("generation_index") or 0)
if index < 1 or index > expected_count:
raise RuntimeError(f"无效的图片生成序号: {index}")
child_created_at = datetime.now(timezone.utc)
child = ChatGenerationTask(
id=generate_id(),
created_at=child_created_at,
resource_generation_started_at=child_created_at,
generation_attempt_no=1,
user_id=main.user_id,
original_prompt=main.original_prompt,
optimized_prompt=main.optimized_prompt,
@@ -433,13 +450,23 @@ async def _enqueue_child_downloads(db: AsyncSession, child_ids: list[str]) -> di
return {"enqueued": enqueued, "failed": failed}
async def run_image_main_batch(db: AsyncSession, main_task: ChatGenerationTask) -> list[str]:
async def run_image_main_batch(
db: AsyncSession,
main_task: ChatGenerationTask,
*,
execution_token: str,
execution_guard: Callable[[], Awaitable[None]],
) -> list[str]:
"""单次同步组图,全部成功后原子拆分 child。
绝不在组图 API 失败后退化为 N 次单图请求。
"""
main_task_id = str(main_task.id)
claim = await _claim_image_main_batch(db, main_task_id)
claim = await _claim_image_main_batch(
db,
main_task_id,
execution_token=execution_token,
)
if claim.existing_child_ids is not None:
await _enqueue_child_downloads(db, claim.existing_child_ids)
return claim.existing_child_ids
@@ -463,6 +490,7 @@ async def run_image_main_batch(db: AsyncSession, main_task: ChatGenerationTask)
claim.runtime_engine,
generation_count=generation_count,
)
await execution_guard()
provider_items = _validate_provider_batch(provider_result, generation_count)
log_operation_event(
domain="generation_ai_batch",
@@ -480,7 +508,10 @@ async def run_image_main_batch(db: AsyncSession, main_task: ChatGenerationTask)
"fallback_to_single_requests": False,
},
)
except RedisExecutionLockError:
raise
except Exception as exc:
await execution_guard()
message = exc.safe_message if isinstance(exc, ImageProviderError) else str(exc)
await _fail_claimed_main(
db,
@@ -493,6 +524,7 @@ async def run_image_main_batch(db: AsyncSession, main_task: ChatGenerationTask)
return []
try:
await execution_guard()
child_ids = await _split_children(
db,
main_task_id=main_task_id,
@@ -500,7 +532,10 @@ async def run_image_main_batch(db: AsyncSession, main_task: ChatGenerationTask)
provider_result=provider_result,
provider_items=provider_items,
)
except RedisExecutionLockError:
raise
except Exception as exc:
await execution_guard()
await _fail_claimed_main(
db,
main_task_id=main_task_id,
@@ -95,12 +95,12 @@ async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEng
"""获取当前启用的图片/视频生成引擎,供前端创建任务时选择 engine_id。"""
image_result = await db.execute(
select(ImageEngine)
.where(ImageEngine.is_active == True)
.where(ImageEngine.is_active == True, ImageEngine.deleted_at.is_(None))
.order_by(ImageEngine.priority.desc())
)
video_result = await db.execute(
select(VideoEngine)
.where(VideoEngine.is_active == True)
.where(VideoEngine.is_active == True, VideoEngine.deleted_at.is_(None))
.order_by(VideoEngine.priority.desc())
)
@@ -274,7 +274,7 @@ def record_to_out(
text_tokens_used=task.text_tokens_used or 0,
image_tokens_used=task.image_tokens_used or 0,
video_tokens_used=task.video_tokens_used or 0,
retry_count=task.retry_count or 0,
retry_count=task.manual_retry_count or 0,
poll_count=task.poll_count or 0,
error_message=task.error_message if is_main else _resolve_error_message(task.error_message),
created_at=task.created_at,
@@ -142,8 +142,16 @@ def _base_task_kwargs(
credits_cost: float = 0.0,
idempotency_key: str | None = None,
) -> dict[str, Any]:
resource_started_at = (
deadline_at - timedelta(minutes=int(settings.CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES or 30))
if gen_type == GenerationType.IMAGE.value
else deadline_at - timedelta(hours=int(settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS or 24))
)
return {
"id": task_id,
"created_at": resource_started_at,
"resource_generation_started_at": resource_started_at,
"generation_attempt_no": 1,
"user_id": user_id,
"original_prompt": req.original_prompt,
"gen_type": gen_type,
@@ -515,12 +523,10 @@ async def enqueue_created_generation_tasks(
) -> list[str]:
"""在业务事务提交后投递任务;返回投递失败的任务ID。
投递失败会在补偿事务中将对应任务置为失败并幂等退款,视频子任务
同时触发主任务状态汇总。调用方不应在初始事务提交前调用本函数
投递失败时保留 queued 状态和已提交的计费/快照,交由 generation recovery
补投,避免 broker 短暂故障被误判为业务生成失败并提前退款
"""
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
from app.services.generation.log_service import log_task_event
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
from app.tasks.generation_create_tasks import chatapi_create_generation_task
normalized_ids = list(dict.fromkeys(str(item) for item in task_ids if item))
@@ -530,6 +536,7 @@ async def enqueue_created_generation_tasks(
ChatGenerationTask.user_id,
ChatGenerationTask.parent_task_id,
ChatGenerationTask.generation_index,
ChatGenerationTask.generation_attempt_no,
).where(ChatGenerationTask.id.in_(normalized_ids))
) if normalized_ids else None
task_meta = {
@@ -537,6 +544,7 @@ async def enqueue_created_generation_tasks(
"user_id": str(row.user_id),
"parent_task_id": str(row.parent_task_id) if row.parent_task_id else None,
"generation_index": row.generation_index,
"generation_attempt_no": int(row.generation_attempt_no or 1),
}
for row in (meta_result.all() if meta_result is not None else [])
}
@@ -555,7 +563,11 @@ async def enqueue_created_generation_tasks(
detail={"generation_index": meta.get("generation_index")},
)
try:
chatapi_create_generation_task.delay(task_id)
chatapi_create_generation_task.apply_async(
args=[task_id],
kwargs={"owner_type": "chat_generation_task", "generation_attempt_no": int(meta.get("generation_attempt_no") or 1)},
queue="gen_chatapi_create",
)
await log_task_event(
task_id=task_id,
event_type="CHILD_ENQUEUE_SUCCESS",
@@ -575,33 +587,25 @@ async def enqueue_created_generation_tasks(
)
except Exception as exc:
failed_ids.append(task_id)
await db.rollback()
failed_task = await mark_chat_generation_task_failed_and_refund_once(
db,
task_id=task_id,
error_message=f"任务队列投递失败: {exc}",
pipeline_stage="failed",
)
await aggregate_parent_for_child(db, failed_task)
await db.commit()
await log_task_event(
task_id=task_id,
generation_attempt_no=int(meta.get("generation_attempt_no") or 1),
event_type="CHILD_ENQUEUE_FAILED",
to_status="failed",
to_stage="failed",
to_status="generating",
to_stage="queued",
message=str(exc),
detail={"task_id": task_id},
detail={"task_id": task_id, "recoverable": True},
)
log_operation_event(
domain="generation_ai_batch",
event_type="CHILD_ENQUEUE_FAILED",
event_status="failed",
source="api",
user_id=getattr(failed_task, "user_id", None),
group_id=getattr(failed_task, "parent_task_id", None) or task_id,
user_id=meta.get("user_id"),
group_id=meta.get("parent_task_id") or task_id,
task_id=task_id,
message=str(exc),
detail={"physical_files_deleted": False},
detail={"recoverable": True, "physical_files_deleted": False},
error=str(exc),
)
return failed_ids
@@ -14,6 +14,7 @@ from app.enums.generation_task import (
GenerationMode,
)
from app.models.chat_generation_task import ChatGenerationTask
from app.services.generation.pipeline.db_lock_service import execute_with_lock_timeout
from app.services.operation_log_service import log_operation_event
from app.services.video_upscale.guard_service import assert_no_recoverable_failed_upscale_tasks
from app.services.resource_accounting_service import (
@@ -157,7 +158,9 @@ async def aggregate_main_task_status(
*,
parent_task_id: str,
) -> ChatGenerationTask | None:
result = await db.execute(
result = await execute_with_lock_timeout(
db,
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == parent_task_id,
@@ -229,8 +232,10 @@ async def aggregate_main_task_status(
main.text_tokens_used = sum(int(child.text_tokens_used or 0) for child in children)
main.image_tokens_used = sum(int(child.image_tokens_used or 0) for child in children)
main.video_tokens_used = sum(int(child.video_tokens_used or 0) for child in children)
main.retry_count = sum(int(child.retry_count or 0) for child in children)
# main 的手动重试次数只代表 main 自身,不能累加 child 的轮询/重试次数。
main.retry_count = int(main.manual_retry_count or 0)
main.poll_count = sum(int(child.poll_count or 0) for child in children)
main.poll_error_count = sum(int(child.poll_error_count or 0) for child in children)
await db.flush()
log_operation_event(
@@ -272,7 +277,9 @@ async def soft_delete_child_tasks_batch(
if not ids:
return 0
deleted_at = deleted_at or datetime.now(timezone.utc)
result = await db.execute(
result = await execute_with_lock_timeout(
db,
select(ChatGenerationTask)
.where(
ChatGenerationTask.id.in_(ids),
@@ -374,7 +381,9 @@ async def soft_delete_top_level_task_group(
deleted_at: datetime | None = None,
) -> int:
deleted_at = deleted_at or datetime.now(timezone.utc)
result = await db.execute(
result = await execute_with_lock_timeout(
db,
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == task_id,
@@ -558,6 +558,7 @@ async def charge_generation_media_for_record(
project_name: str | None = None,
description_prefix: str = "AI创作-",
attempt_no: int | None = None,
engine_id: str | None = None,
) -> BillingSummary:
return await charge_generation_media_by_params(
db,
@@ -567,6 +568,7 @@ async def charge_generation_media_for_record(
image_size=record.image_size,
duration=record.duration,
resolution=record.resolution,
engine_id=engine_id or getattr(record, "engine_id", None),
project_name=project_name,
description_prefix=description_prefix,
owner_type=OWNER_GENERATION_RECORD,
@@ -1,12 +1,15 @@
from __future__ import annotations
import json
import os
import uuid
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Awaitable, Callable
from urllib.parse import urlparse
from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask
from app.services.generation.pipeline.owner_service import GenerationOwner
from app.services.image_gen import download_image
from app.services.provider_limit import provider_limit
from app.services.resource_accounting_service import safe_file_size
@@ -34,7 +37,7 @@ def _to_aware_utc(value: datetime | None) -> datetime | None:
return value.astimezone(timezone.utc)
def _build_storage_date_dir(record: ChatGenerationTask) -> str:
def _build_storage_date_dir(record: GenerationOwner) -> str:
fixed = (getattr(record, "download_storage_date_dir", None) or "").strip().strip("/")
if fixed:
return fixed
@@ -42,6 +45,32 @@ def _build_storage_date_dir(record: ChatGenerationTask) -> str:
return created_at.strftime("%Y/%m/%d")
def _normalize_image_extension(record: GenerationOwner) -> str:
output_format = ""
try:
snapshot = json.loads(getattr(record, "engine_snapshot_json", None) or "{}")
if isinstance(snapshot, dict):
output_format = str(snapshot.get("output_format") or "").strip().lower()
except Exception:
output_format = ""
if output_format in {"jpg", "jpeg"}:
return "jpg"
if output_format in {"png", "webp"}:
return output_format
remote_url = str(getattr(record, "remote_result_url", None) or "")
try:
suffix = os.path.splitext(urlparse(remote_url).path or "")[1].lower().lstrip(".")
except Exception:
suffix = ""
if suffix in {"jpg", "jpeg"}:
return "jpg"
if suffix in {"png", "webp"}:
return suffix
return "jpg"
def _make_part_path(final_path: str) -> str:
return f"{final_path}.{uuid.uuid4().hex}.part"
@@ -65,16 +94,27 @@ def _safe_remove(path: str | None) -> None:
pass
async def _download_image_atomically(remote_url: str, final_path: str) -> str:
async def _download_image_atomically(
remote_url: str,
final_path: str,
*,
execution_guard: Callable[[], Awaitable[None]] | None = None,
) -> str:
if _is_valid_file(final_path):
return final_path
os.makedirs(os.path.dirname(final_path), exist_ok=True)
part_path = _make_part_path(final_path)
try:
await download_image(remote_url, part_path)
await download_image(
remote_url,
part_path,
execution_guard=execution_guard,
)
if not _is_valid_file(part_path):
raise RuntimeError("图片下载完成但临时文件为空")
if execution_guard is not None:
await execution_guard()
os.replace(part_path, final_path)
return final_path
except Exception:
@@ -82,16 +122,27 @@ async def _download_image_atomically(remote_url: str, final_path: str) -> str:
raise
async def _download_video_atomically(remote_url: str, final_path: str) -> str:
async def _download_video_atomically(
remote_url: str,
final_path: str,
*,
execution_guard: Callable[[], Awaitable[None]] | None = None,
) -> str:
if _is_valid_file(final_path):
return final_path
os.makedirs(os.path.dirname(final_path), exist_ok=True)
part_path = _make_part_path(final_path)
try:
await download_video(remote_url, part_path)
await download_video(
remote_url,
part_path,
execution_guard=execution_guard,
)
if not _is_valid_file(part_path):
raise RuntimeError("视频下载完成但临时文件为空")
if execution_guard is not None:
await execution_guard()
os.replace(part_path, final_path)
return final_path
except Exception:
@@ -101,7 +152,11 @@ async def _download_video_atomically(remote_url: str, final_path: str) -> str:
async def download_video_upscale_source(record: ChatGenerationTask) -> DownloadedGenerationResult:
async def download_video_upscale_source(
record: GenerationOwner,
*,
execution_guard: Callable[[], Awaitable[None]] | None = None,
) -> DownloadedGenerationResult:
"""下载超分源视频。
源视频只供后处理使用,不生成封面,也不作为用户 GeneratedResource。
@@ -119,10 +174,16 @@ async def download_video_upscale_source(record: ChatGenerationTask) -> Downloade
part_path = build_part_mp4_path(dest)
try:
async with provider_limit("result_download", settings.RESULT_DOWNLOAD_MAX_CONCURRENCY):
await download_video(record.remote_result_url, part_path)
await download_video(
record.remote_result_url,
part_path,
execution_guard=execution_guard,
)
if not _is_valid_file(part_path):
raise RuntimeError("超分源视频下载完成但临时文件为空")
await probe_video(part_path)
if execution_guard is not None:
await execution_guard()
os.replace(part_path, dest)
except Exception:
_safe_remove(part_path)
@@ -138,7 +199,11 @@ async def download_video_upscale_source(record: ChatGenerationTask) -> Downloade
)
async def download_generation_result(record: ChatGenerationTask) -> DownloadedGenerationResult:
async def download_generation_result(
record: GenerationOwner,
*,
execution_guard: Callable[[], Awaitable[None]] | None = None,
) -> DownloadedGenerationResult:
if not record.remote_result_url:
raise ValueError("缺少远程结果URL")
@@ -147,13 +212,18 @@ async def download_generation_result(record: ChatGenerationTask) -> DownloadedGe
if record.gen_type == "image":
dest_dir = os.path.join(settings.STORAGE_IMAGE_LOCAL_PATH, date_dir)
os.makedirs(dest_dir, exist_ok=True)
dest = os.path.join(dest_dir, f"{record.id}.png")
extension = _normalize_image_extension(record)
dest = os.path.join(dest_dir, f"{record.id}.{extension}")
async with provider_limit("result_download", settings.RESULT_DOWNLOAD_MAX_CONCURRENCY):
await _download_image_atomically(record.remote_result_url if record.remote_result_url else "", dest)
await _download_image_atomically(
record.remote_result_url if record.remote_result_url else "",
dest,
execution_guard=execution_guard,
)
return DownloadedGenerationResult(
url=f"/generate/images/{date_dir}/{record.id}.png",
url=f"/generate/images/{date_dir}/{record.id}.{extension}",
storage_path=dest,
file_size_bytes=safe_file_size(dest),
resource_type="image",
@@ -164,14 +234,20 @@ async def download_generation_result(record: ChatGenerationTask) -> DownloadedGe
dest = os.path.join(dest_dir, f"{record.id}.mp4")
async with provider_limit("result_download", settings.RESULT_DOWNLOAD_MAX_CONCURRENCY):
await _download_video_atomically(record.remote_result_url if record.remote_result_url else "", dest)
await _download_video_atomically(
record.remote_result_url if record.remote_result_url else "",
dest,
execution_guard=execution_guard,
)
cover_url, cover_storage_path = create_video_cover_for_local_video(
record_id=record.id,
video_path=dest,
date_dir=date_dir,
log_prefix=f"ChatGenerationTask视频封面生成 task_id={record.id}",
log_prefix=f"生成资源视频封面 task_id={record.id}",
)
if execution_guard is not None:
await execution_guard()
return DownloadedGenerationResult(
url=f"/generate/videos/{date_dir}/{record.id}.mp4",
@@ -23,6 +23,7 @@ from app.models.module_generation_step import ModuleGenerationStep
from app.models.shot_replicate_segment import ShotReplicateSegment
from app.models.user import User
from app.schemas.generation_ai import GenerationAIHistoryBatchDeleteOut
from app.services.generation.pipeline.db_lock_service import execute_with_lock_timeout
from app.services.generation.ai.task_group_service import soft_delete_child_tasks_batch
from app.services.module_generation_flow_base_service import is_active_chat_generation_task
from app.services.module_generation_log_service import log_module_event_file
@@ -105,7 +106,9 @@ async def _load_chat_tasks_for_steps(
task_ids = _task_id_list(steps)
if not task_ids:
return {}
result = await db.execute(
result = await execute_with_lock_timeout(
db,
select(ChatGenerationTask)
.where(
ChatGenerationTask.id.in_(task_ids),
@@ -194,7 +197,9 @@ async def _delete_generation_records(
ids: list[str],
deleted_at: datetime,
) -> GenerationAIHistoryBatchDeleteOut:
result = await db.execute(
result = await execute_with_lock_timeout(
db,
select(GenerationRecord)
.where(
GenerationRecord.id.in_(ids),
@@ -238,7 +243,9 @@ async def _delete_chat_tasks(
ids: list[str],
deleted_at: datetime,
) -> GenerationAIHistoryBatchDeleteOut:
result = await db.execute(
result = await execute_with_lock_timeout(
db,
select(ChatGenerationTask)
.where(
ChatGenerationTask.id.in_(ids),
@@ -325,7 +332,9 @@ async def _load_module_projects_by_ids(
source: GenerationHistorySourceEnum,
project_ids: list[str],
) -> list[ModuleGenerationProject]:
result = await db.execute(
result = await execute_with_lock_timeout(
db,
select(ModuleGenerationProject)
.where(
ModuleGenerationProject.id.in_(project_ids),
@@ -351,7 +360,8 @@ async def _soft_delete_module_projects(
if not project_ids:
return [], [], 0
step_result = await db.execute(
step_result = await execute_with_lock_timeout(
db,
select(ModuleGenerationStep)
.where(
ModuleGenerationStep.project_id.in_(project_ids),
@@ -452,7 +462,9 @@ async def _delete_shot_segments(
ids: list[str],
deleted_at: datetime,
) -> GenerationAIHistoryBatchDeleteOut:
result = await db.execute(
result = await execute_with_lock_timeout(
db,
select(ShotReplicateSegment)
.where(
ShotReplicateSegment.id.in_(ids),
@@ -4,9 +4,13 @@ import hashlib
import json
from typing import Any
from app.enums.generation_task import GenerationMode, GenerationOwnerType
from app.models.base import async_session
from app.models.chat_generation_task import ChatGenerationTask
from app.models.chat_generation_task_event import ChatGenerationTaskEvent
from app.models.chat_provider_call_log import ChatProviderCallLog
from app.models.generation_record import GenerationRecord
from app.services.operation_log_service import build_exception_detail, log_operation_event, sanitize_log_value
from app.utils.id_gen import generate_id
MAX_EXCERPT_CHARS = 2000
@@ -22,12 +26,9 @@ def _safe_json(data: Any) -> str | None:
def _excerpt(data: Any, limit: int = MAX_EXCERPT_CHARS) -> str | None:
text = _safe_json(data)
text = _safe_json(sanitize_log_value(data))
if text is None:
return None
# Avoid storing secrets in logs.
text = text.replace("Authorization", "Authorization-REDACTED")
text = text.replace("api_key", "api_key_REDACTED")
if len(text) > limit:
return text[:limit] + "...[truncated]"
return text
@@ -40,12 +41,74 @@ def _hash(data: Any) -> str | None:
return hashlib.sha256(text.encode("utf-8")).hexdigest()
def _owner_fields(
obj: Any | None,
*,
owner_type: str | None,
owner_id: str | None,
task_id: str | None,
record_id: str | None,
generation_attempt_no: int | None,
generation_mode: str | None,
) -> dict[str, Any] | None:
if isinstance(obj, GenerationRecord) or record_id:
resolved_owner_type = GenerationOwnerType.GENERATION_RECORD.value
resolved_owner_id = record_id or owner_id or getattr(obj, "id", None)
resolved_task_id = None
resolved_record_id = resolved_owner_id
resolved_mode = generation_mode or GenerationMode.GENERATION_RECORD.value
elif isinstance(obj, ChatGenerationTask) or task_id:
resolved_owner_type = GenerationOwnerType.CHAT_GENERATION_TASK.value
resolved_owner_id = task_id or owner_id or getattr(obj, "id", None)
resolved_task_id = resolved_owner_id
resolved_record_id = None
resolved_mode = generation_mode or getattr(obj, "generation_mode", GenerationMode.CHATAPI_ASYNC.value)
else:
resolved_owner_type = owner_type or GenerationOwnerType.CHAT_GENERATION_TASK.value
resolved_owner_id = owner_id
if resolved_owner_type == GenerationOwnerType.GENERATION_RECORD.value:
resolved_task_id = None
resolved_record_id = resolved_owner_id
resolved_mode = generation_mode or GenerationMode.GENERATION_RECORD.value
else:
resolved_task_id = resolved_owner_id
resolved_record_id = None
resolved_mode = generation_mode or GenerationMode.CHATAPI_ASYNC.value
if not resolved_owner_id:
return None
return {
"owner_type": resolved_owner_type,
"owner_id": str(resolved_owner_id),
"task_id": str(resolved_task_id) if resolved_task_id else None,
"generation_record_id": str(resolved_record_id) if resolved_record_id else None,
"generation_attempt_no": int(generation_attempt_no or getattr(obj, "generation_attempt_no", 1) or 1),
"generation_mode": str(resolved_mode or ""),
}
def _fallback_log(event_type: str, fields: dict[str, Any] | None, exc: Exception) -> None:
fields = fields or {}
log_operation_event(
domain="generation_pipeline",
event_type="PIPELINE_DB_LOG_FAILED",
event_status="failed",
task_id=fields.get("owner_id"),
message=f"数据库生成日志写入失败: {event_type}",
detail=build_exception_detail(exc, fields),
error=str(exc),
)
async def log_task_event(
task: Any | None = None,
*,
record: Any | None = None,
owner_type: str | None = None,
owner_id: str | None = None,
task_id: str | None = None,
record_id: str | None = None,
generation_attempt_no: int | None = None,
generation_mode: str | None = None,
event_type: str,
from_status: str | None = None,
to_status: str | None = None,
@@ -54,17 +117,52 @@ async def log_task_event(
message: str | None = None,
detail: Any = None,
) -> None:
"""Write task event in a separate transaction; failure must not affect main flow."""
"""Write an owner-scoped event in a separate transaction."""
obj = task or record
fields = _owner_fields(
obj,
owner_type=owner_type,
owner_id=owner_id,
task_id=task_id,
record_id=record_id,
generation_attempt_no=generation_attempt_no,
generation_mode=generation_mode,
)
if not fields:
return
upper_event = str(event_type or "").upper()
failed_event = any(marker in upper_event for marker in ("FAILED", "TIMEOUT", "ERROR"))
log_operation_event(
domain="generation_pipeline",
event_type=event_type,
event_status="failed" if failed_event else "success",
source="pipeline",
user_id=str(getattr(obj, "user_id", "") or "") or None,
project_id=str(getattr(obj, "project_id", "") or "") or None,
task_id=fields["owner_id"],
message=message,
detail={
"owner_type": fields["owner_type"],
"owner_id": fields["owner_id"],
"generation_attempt_no": fields["generation_attempt_no"],
"generation_mode": fields["generation_mode"],
"from_status": from_status,
"to_status": to_status,
"from_stage": from_stage,
"to_stage": to_stage,
"detail_excerpt": _excerpt(detail),
},
error=message if failed_event else None,
)
try:
obj = task or record
tid = task_id or record_id or (obj.id if obj else None)
if not tid:
return
async with async_session() as db:
db.add(ChatGenerationTaskEvent(
id=generate_id(),
task_id=tid,
generation_mode=getattr(obj, "generation_mode", "chatapi_async"),
owner_type=fields["owner_type"],
task_id=fields["task_id"],
generation_record_id=fields["generation_record_id"],
generation_attempt_no=fields["generation_attempt_no"],
generation_mode=fields["generation_mode"],
event_type=event_type,
from_status=from_status,
to_status=to_status,
@@ -74,16 +172,20 @@ async def log_task_event(
detail_json=_excerpt(detail),
))
await db.commit()
except Exception:
return
except Exception as exc:
_fallback_log(event_type, fields, exc)
async def log_provider_call(
task: Any | None = None,
*,
record: Any | None = None,
owner_type: str | None = None,
owner_id: str | None = None,
task_id: str | None = None,
record_id: str | None = None,
generation_attempt_no: int | None = None,
generation_mode: str | None = None,
provider: str | None,
api_type: str,
model: str | None = None,
@@ -100,17 +202,28 @@ async def log_provider_call(
error_code: str | None = None,
error_message: str | None = None,
) -> None:
"""Write provider call log in a separate transaction; failure must not affect main flow."""
"""Write an owner-scoped provider call log in a separate transaction."""
obj = task or record
fields = _owner_fields(
obj,
owner_type=owner_type,
owner_id=owner_id,
task_id=task_id,
record_id=record_id,
generation_attempt_no=generation_attempt_no,
generation_mode=generation_mode,
)
if not fields:
return
try:
obj = task or record
tid = task_id or record_id or (obj.id if obj else None)
if not tid:
return
async with async_session() as db:
db.add(ChatProviderCallLog(
id=generate_id(),
task_id=tid,
generation_mode=getattr(obj, "generation_mode", "chatapi_async"),
owner_type=fields["owner_type"],
task_id=fields["task_id"],
generation_record_id=fields["generation_record_id"],
generation_attempt_no=fields["generation_attempt_no"],
generation_mode=fields["generation_mode"],
provider=provider,
api_type=api_type,
model=model,
@@ -130,5 +243,5 @@ async def log_provider_call(
error_message=error_message,
))
await db.commit()
except Exception:
return
except Exception as exc:
_fallback_log(api_type, fields, exc)
@@ -0,0 +1 @@
"""Shared Celery generation pipeline for ChatGenerationTask and GenerationRecord."""
@@ -0,0 +1,84 @@
from __future__ import annotations
from typing import Any
from fastapi import HTTPException
from sqlalchemy import text
from sqlalchemy.exc import DBAPIError, OperationalError
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
class DatabaseRowLockBusy(HTTPException, RuntimeError):
"""A short-lived PostgreSQL row/table lock could not be acquired in time."""
def __init__(self, message: str = "当前任务正在被其他流程处理,请稍后重试") -> None:
super().__init__(status_code=409, detail=message)
_LOCK_NOT_AVAILABLE_SQLSTATE = "55P03"
def _sqlstate_from_exception(exc: BaseException | None) -> str | None:
current: BaseException | None = exc
seen: set[int] = set()
while current is not None and id(current) not in seen:
seen.add(id(current))
for attr in ("sqlstate", "pgcode"):
value = getattr(current, attr, None)
if value:
return str(value)
current = getattr(current, "orig", None) or getattr(current, "__cause__", None)
return None
def is_postgres_lock_timeout(exc: BaseException) -> bool:
if _sqlstate_from_exception(exc) == _LOCK_NOT_AVAILABLE_SQLSTATE:
return True
message = str(exc).lower()
return "lock timeout" in message or "could not obtain lock" in message
def raise_if_database_lock_busy(exc: BaseException) -> None:
if is_postgres_lock_timeout(exc):
raise DatabaseRowLockBusy("数据库任务行正在被其他事务处理,请稍后重试") from exc
async def apply_short_lock_timeout(
db: AsyncSession,
*,
seconds: int | None = None,
) -> None:
"""Apply a transaction-local PostgreSQL lock wait limit.
It deliberately does not change the global database configuration and is a
no-op on non-PostgreSQL test/development databases.
"""
bind = db.get_bind()
if bind is None or bind.dialect.name != "postgresql":
return
timeout_seconds = max(
1,
int(
seconds
if seconds is not None
else getattr(settings, "GENERATION_DB_LOCK_TIMEOUT_SECONDS", 5)
or 5
),
)
await db.execute(text(f"SET LOCAL lock_timeout = '{timeout_seconds}s'"))
async def execute_with_lock_timeout(
db: AsyncSession,
statement: Any,
*,
seconds: int | None = None,
):
await apply_short_lock_timeout(db, seconds=seconds)
try:
return await db.execute(statement)
except (OperationalError, DBAPIError) as exc:
raise_if_database_lock_busy(exc)
raise
@@ -0,0 +1,39 @@
from __future__ import annotations
from app.enums.celery_queue import CeleryQueue
from app.enums.generation_task import ChatGenerationTaskEventType
from app.services.generation.log_service import log_task_event
from app.services.generation.pipeline.owner_service import GenerationOwner, owner_type_of
async def enqueue_generation_create(
owner: GenerationOwner,
*,
reason: str,
) -> None:
"""Commit caller-owned state before invoking this function."""
from app.tasks.generation_create_tasks import chatapi_create_generation_task
owner_type = owner_type_of(owner)
attempt_no = int(getattr(owner, "generation_attempt_no", 1) or 1)
try:
chatapi_create_generation_task.apply_async(
args=[str(owner.id)],
kwargs={"owner_type": owner_type, "generation_attempt_no": attempt_no},
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
task_id=f"generation-create:{owner_type}:{owner.id}:attempt:{attempt_no}",
)
await log_task_event(
owner,
event_type=ChatGenerationTaskEventType.GENERATION_RECORD_ENQUEUE_SUCCESS.value,
message="资源生成创建任务已投递",
detail={"reason": reason, "queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
)
except Exception as exc:
await log_task_event(
owner,
event_type=ChatGenerationTaskEventType.GENERATION_RECORD_ENQUEUE_FAILED.value,
message=str(exc),
detail={"reason": reason, "queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
)
raise
@@ -0,0 +1,59 @@
from __future__ import annotations
import json
from datetime import datetime, timezone
from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.generation_status import GenerationRecordPipelineStage, GenerationStatus
from app.models.generation_record import GenerationRecord
from app.models.image_engine import ImageEngine
from app.models.video_engine import VideoEngine
from app.services.generation.ai.engine_service import build_image_snapshot, build_video_snapshot
from app.services.generation.pipeline.enqueue_service import enqueue_generation_create
from app.services.generation.pipeline.lifecycle_service import reset_execution_fields
def _json(data: dict) -> str:
return json.dumps(data, ensure_ascii=False, default=str)
def prepare_generation_record_execution(
record: GenerationRecord,
*,
engine: ImageEngine | VideoEngine,
attempt_no: int,
) -> None:
now = datetime.now(timezone.utc)
reset_execution_fields(record, started_at=now, attempt_no=attempt_no)
record.engine_id = engine.id
if record.gen_type == "image":
record.engine_snapshot_json = _json(build_image_snapshot(
engine,
record.image_size or getattr(engine, "default_size", "2K") or "2K",
record.image_proportion or "1:1",
record.image_px or "2048x2048",
))
else:
record.engine_snapshot_json = _json(build_video_snapshot(
engine,
record.aspect_ratio or "16:9",
record.resolution or "480p",
int(record.duration or 4),
))
record.status = GenerationStatus.generating.value
record.pipeline_stage = GenerationRecordPipelineStage.QUEUED.value
async def commit_and_enqueue_generation_record(
db: AsyncSession,
record: GenerationRecord,
*,
reason: str,
) -> None:
await db.commit()
try:
await enqueue_generation_create(record, reason=reason)
except Exception:
# queued stage and all execution metadata are already committed; recovery will retry.
return
@@ -0,0 +1,149 @@
from __future__ import annotations
import asyncio
from datetime import datetime, timedelta, timezone
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.enums.generation_status import GenerationRecordPipelineStage
from app.enums.generation_task import ChatGenerationPipelineStage, GenerationType
from app.models.chat_generation_task import ChatGenerationTask
from app.models.generation_record import GenerationRecord
from app.models.base import async_session
from app.services.generation.pipeline.db_lock_service import DatabaseRowLockBusy
from app.services.generation.pipeline.owner_service import GenerationOwner
from app.services.generation.refund_service import (
mark_chat_generation_task_failed_and_refund_once,
mark_generation_record_failed_and_refund_once,
)
def utc_now() -> datetime:
return datetime.now(timezone.utc)
def generation_deadline(*, gen_type: str, started_at: datetime) -> datetime:
if str(gen_type or "").lower() == GenerationType.IMAGE.value:
return started_at + timedelta(minutes=max(1, int(settings.CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES or 30)))
return started_at + timedelta(hours=max(1, int(settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS or 24)))
def reset_execution_fields(owner: GenerationOwner, *, started_at: datetime, attempt_no: int) -> None:
owner.generation_attempt_no = max(1, int(attempt_no or 1))
owner.resource_generation_started_at = started_at
owner.deadline_at = generation_deadline(gen_type=owner.gen_type, started_at=started_at)
owner.error_message = None
owner.seedance_task_id = None
if hasattr(owner, "provider_task_id"):
owner.provider_task_id = None
owner.remote_result_url = None
owner.provider_response_json = None
owner.provider_create_claim_token = None
owner.provider_create_lease_until = None
owner.provider_create_started_at = None
owner.retry_count = int(getattr(owner, "manual_retry_count", 0) or 0)
owner.poll_error_count = 0
owner.poll_count = 0
owner.last_poll_at = None
owner.poll_started_at = None
owner.next_poll_at = None
owner.poll_interval_seconds = 0
owner.poll_claim_token = None
owner.poll_lease_until = None
owner.download_celery_task_id = None
owner.download_enqueued_at = None
owner.download_started_at = None
owner.download_claim_token = None
owner.download_lease_until = None
owner.download_next_retry_at = None
owner.download_attempt_count = 0
owner.download_last_error = None
owner.download_storage_date_dir = None
owner.generated_at = None
owner.image_url = None
owner.video_url = None
owner.video_cover_url = None
async def notify_owner_finished(db: AsyncSession, owner: GenerationOwner) -> None:
"""Run module hooks after the generation owner transaction has committed.
Module project/step rows have their own short locks. A transient lock conflict
must not turn an already completed generation into a download/provider retry,
so hooks use a fresh short transaction with a small local retry window.
"""
_ = db
if not isinstance(owner, ChatGenerationTask):
return
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
from app.services.operation_log_service import log_operation_event
task_id = str(owner.id)
attempt_no = int(owner.generation_attempt_no or 1)
last_error: Exception | None = None
for retry_index in range(3):
try:
async with async_session() as hook_db:
result = await hook_db.execute(
select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
).limit(1)
)
fresh_task = result.scalar_one_or_none()
if fresh_task is None:
return
await notify_chat_generation_task_finished(hook_db, fresh_task)
await aggregate_parent_for_child(hook_db, fresh_task)
await hook_db.commit()
return
except DatabaseRowLockBusy as exc:
last_error = exc
await asyncio.sleep(1 + retry_index)
except Exception as exc:
last_error = exc
break
log_operation_event(
domain="generation_pipeline",
event_type="MODULE_HOOK_DEFERRED_MANUAL_CHECK",
event_status="failed",
source="pipeline",
task_id=task_id,
message="生成任务已进入终态,但模块状态回填失败,需要人工排查",
detail={
"generation_attempt_no": attempt_no,
"retry_count": 3,
},
error=str(last_error or "unknown module hook error"),
)
async def mark_owner_failed_and_refund_once(
db: AsyncSession,
owner: GenerationOwner,
*,
error_message: str,
pipeline_stage: str,
) -> GenerationOwner:
if isinstance(owner, ChatGenerationTask):
await mark_chat_generation_task_failed_and_refund_once(
db,
task=owner,
error_message=error_message,
pipeline_stage=pipeline_stage,
generation_attempt_no=int(owner.generation_attempt_no or 1),
)
return owner
owner.pipeline_stage = pipeline_stage if pipeline_stage in {item.value for item in GenerationRecordPipelineStage} else GenerationRecordPipelineStage.FAILED.value
await mark_generation_record_failed_and_refund_once(
db,
record=owner,
error_message=error_message,
generation_attempt_no=int(owner.generation_attempt_no or 1),
)
return owner
@@ -0,0 +1,192 @@
from __future__ import annotations
import asyncio
from dataclasses import dataclass
from typing import TypeAlias
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.generation_status import GenerationStatus
from app.enums.generation_task import (
ChatGenerationTaskStatus,
GenerationMode,
GenerationOwnerType,
)
from app.models.chat_generation_task import ChatGenerationTask
from app.models.generation_record import GenerationRecord
from app.services.generation.pipeline.db_lock_service import (
DatabaseRowLockBusy,
apply_short_lock_timeout,
raise_if_database_lock_busy,
)
GenerationOwner: TypeAlias = ChatGenerationTask | GenerationRecord
@dataclass(frozen=True, slots=True)
class GenerationOwnerRef:
owner_type: str
owner_id: str
generation_attempt_no: int | None = None
@classmethod
def from_owner(cls, owner: GenerationOwner) -> "GenerationOwnerRef":
return cls(
owner_type=owner_type_of(owner),
owner_id=str(owner.id),
generation_attempt_no=int(getattr(owner, "generation_attempt_no", 1) or 1),
)
def normalize_owner_type(value: str | GenerationOwnerType | None) -> str:
if isinstance(value, GenerationOwnerType):
return value.value
text = str(value or "").strip().lower()
if not text:
return GenerationOwnerType.CHAT_GENERATION_TASK.value
if text not in {item.value for item in GenerationOwnerType}:
raise ValueError(f"不支持的生成任务所有者类型: {value}")
return text
def owner_type_of(owner: GenerationOwner) -> str:
if isinstance(owner, ChatGenerationTask):
return GenerationOwnerType.CHAT_GENERATION_TASK.value
if isinstance(owner, GenerationRecord):
return GenerationOwnerType.GENERATION_RECORD.value
raise TypeError(f"不支持的生成任务对象: {type(owner)!r}")
def owner_mode(owner: GenerationOwner) -> str:
if isinstance(owner, ChatGenerationTask):
return str(owner.generation_mode or GenerationMode.CHATAPI_ASYNC.value)
return GenerationMode.GENERATION_RECORD.value
def owner_provider_task_id(owner: GenerationOwner) -> str | None:
return str(getattr(owner, "provider_task_id", None) or getattr(owner, "seedance_task_id", None) or "") or None
def set_owner_provider_task_id(owner: GenerationOwner, value: str | None) -> None:
if hasattr(owner, "provider_task_id"):
owner.provider_task_id = value
owner.seedance_task_id = value
def owner_is_generating(owner: GenerationOwner) -> bool:
if isinstance(owner, ChatGenerationTask):
return owner.status == ChatGenerationTaskStatus.GENERATING.value
return owner.status == GenerationStatus.generating.value
def owner_is_completed(owner: GenerationOwner) -> bool:
if isinstance(owner, ChatGenerationTask):
return owner.status == ChatGenerationTaskStatus.COMPLETED.value
return owner.status == GenerationStatus.completed.value
def owner_is_failed(owner: GenerationOwner) -> bool:
if isinstance(owner, ChatGenerationTask):
return owner.status == ChatGenerationTaskStatus.FAILED.value
return owner.status == GenerationStatus.failed.value
def set_owner_generating(owner: GenerationOwner) -> None:
owner.status = ChatGenerationTaskStatus.GENERATING.value if isinstance(owner, ChatGenerationTask) else GenerationStatus.generating.value
def set_owner_completed(owner: GenerationOwner) -> None:
owner.status = ChatGenerationTaskStatus.COMPLETED.value if isinstance(owner, ChatGenerationTask) else GenerationStatus.completed.value
def set_owner_failed(owner: GenerationOwner) -> None:
owner.status = ChatGenerationTaskStatus.FAILED.value if isinstance(owner, ChatGenerationTask) else GenerationStatus.failed.value
def is_attempt_current(owner: GenerationOwner, attempt_no: int | None) -> bool:
if attempt_no is None:
return True
return int(getattr(owner, "generation_attempt_no", 1) or 1) == int(attempt_no)
async def load_generation_owner(
db: AsyncSession,
*,
owner_type: str | GenerationOwnerType | None,
owner_id: str,
for_update: bool = False,
include_deleted: bool = False,
) -> GenerationOwner | None:
normalized = normalize_owner_type(owner_type)
if normalized == GenerationOwnerType.CHAT_GENERATION_TASK.value:
query = select(ChatGenerationTask).where(ChatGenerationTask.id == owner_id)
if not include_deleted:
query = query.where(ChatGenerationTask.deleted_at.is_(None))
else:
query = select(GenerationRecord).where(GenerationRecord.id == owner_id)
if not include_deleted:
query = query.where(GenerationRecord.deleted_at.is_(None))
if for_update:
await apply_short_lock_timeout(db)
query = query.with_for_update().execution_options(populate_existing=True)
try:
result = await db.execute(query.limit(1))
except Exception as exc:
raise_if_database_lock_busy(exc)
raise
return result.scalar_one_or_none()
async def load_generation_owner_for_update_retry(
db: AsyncSession,
*,
owner_type: str | GenerationOwnerType | None,
owner_id: str,
attempts: int = 3,
) -> GenerationOwner | None:
"""Retry a short owner row lock inside the same Worker.
This is intended after an external call/download has already completed so a
transient row lock does not force the Worker to repeat that external side effect.
"""
last_error: DatabaseRowLockBusy | None = None
for retry_index in range(max(1, int(attempts or 1))):
try:
return await load_generation_owner(
db,
owner_type=owner_type,
owner_id=owner_id,
for_update=True,
)
except DatabaseRowLockBusy as exc:
last_error = exc
await db.rollback()
if retry_index + 1 < max(1, int(attempts or 1)):
await asyncio.sleep(1 + retry_index)
raise last_error or DatabaseRowLockBusy()
def redis_owner_item_id(owner_type: str | GenerationOwnerType | None, owner_id: str, attempt_no: int | None = None) -> str:
normalized = normalize_owner_type(owner_type)
if attempt_no is None:
return f"{normalized}:{owner_id}"
return f"{normalized}:{owner_id}:attempt:{int(attempt_no)}"
def parse_redis_owner_item_id(value: str) -> GenerationOwnerRef:
text = str(value or "").strip()
for owner_type in (GenerationOwnerType.CHAT_GENERATION_TASK.value, GenerationOwnerType.GENERATION_RECORD.value):
prefix = f"{owner_type}:"
if text.startswith(prefix):
rest = text[len(prefix):]
marker = ":attempt:"
if marker in rest:
owner_id, attempt = rest.rsplit(marker, 1)
try:
return GenerationOwnerRef(owner_type, owner_id, int(attempt))
except ValueError:
return GenerationOwnerRef(owner_type, owner_id, None)
return GenerationOwnerRef(owner_type, rest, None)
# Historical Redis/Celery identifiers always belonged to ChatGenerationTask.
return GenerationOwnerRef(GenerationOwnerType.CHAT_GENERATION_TASK.value, text, None)
@@ -0,0 +1,148 @@
from __future__ import annotations
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.enums.generation_status import GenerationRecordPipelineStage, GenerationStatus
from app.enums.generation_task import GenerationOwnerType
from app.models.generation_record import GenerationRecord
from app.services.generation.pipeline.owner_service import GenerationOwnerRef
from app.services.redis_registry_service import ensure_aware_utc
@dataclass(frozen=True, slots=True)
class GenerationRecordRecoveryCursor:
owner_id: str
@dataclass(slots=True)
class GenerationRecordRecoveryBatch:
create: list[GenerationOwnerRef]
poll: list[GenerationOwnerRef]
download: list[GenerationOwnerRef]
next_cursor: GenerationRecordRecoveryCursor | None
async def find_generation_record_recovery_batch(
db: AsyncSession,
*,
limit: int,
cursor: GenerationRecordRecoveryCursor | None = None,
) -> GenerationRecordRecoveryBatch:
"""按稳定游标读取恢复分流所需列,避免大字段加载和 offset 扫描。"""
stages = {
GenerationRecordPipelineStage.QUEUED.value,
GenerationRecordPipelineStage.PREPARING.value,
GenerationRecordPipelineStage.CREATING_PROVIDER_TASK.value,
GenerationRecordPipelineStage.WAITING_REMOTE.value,
GenerationRecordPipelineStage.POLLING.value,
GenerationRecordPipelineStage.RESULT_READY.value,
GenerationRecordPipelineStage.DOWNLOAD_QUEUED.value,
GenerationRecordPipelineStage.DOWNLOADING.value,
GenerationRecordPipelineStage.RETRY_WAITING.value,
}
page_size = max(1, int(limit))
now = datetime.now(timezone.utc)
queue_timeout = timedelta(
seconds=max(1, int(settings.DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS or 300))
)
query = select(
GenerationRecord.id,
GenerationRecord.generation_attempt_no,
GenerationRecord.seedance_task_id,
GenerationRecord.remote_result_url,
GenerationRecord.pipeline_stage,
GenerationRecord.provider_create_lease_until,
GenerationRecord.next_poll_at,
GenerationRecord.poll_lease_until,
GenerationRecord.download_enqueued_at,
GenerationRecord.download_lease_until,
GenerationRecord.download_next_retry_at,
).where(
GenerationRecord.deleted_at.is_(None),
GenerationRecord.status == GenerationStatus.generating.value,
GenerationRecord.pipeline_stage.in_(stages),
)
if cursor is not None:
query = query.where(GenerationRecord.id > cursor.owner_id)
result = await db.execute(
query.order_by(GenerationRecord.id.asc()).limit(page_size)
)
rows = list(result.all())
create: list[GenerationOwnerRef] = []
poll: list[GenerationOwnerRef] = []
download: list[GenerationOwnerRef] = []
for (
owner_id,
attempt_no,
provider_task_id,
remote_result_url,
pipeline_stage,
provider_create_lease_until,
next_poll_at,
poll_lease_until,
download_enqueued_at,
download_lease_until,
download_next_retry_at,
) in rows:
ref = GenerationOwnerRef(
GenerationOwnerType.GENERATION_RECORD.value,
str(owner_id),
int(attempt_no or 1),
)
stage = str(pipeline_stage or "")
if str(remote_result_url or "").strip():
if stage == GenerationRecordPipelineStage.RESULT_READY.value:
download.append(ref)
elif stage == GenerationRecordPipelineStage.DOWNLOAD_QUEUED.value:
checked_enqueued_at = ensure_aware_utc(download_enqueued_at)
if (
checked_enqueued_at is None
or checked_enqueued_at + queue_timeout <= now
):
download.append(ref)
elif stage == GenerationRecordPipelineStage.DOWNLOADING.value:
if (
ensure_aware_utc(download_lease_until) is None
or ensure_aware_utc(download_lease_until) <= now
):
download.append(ref)
elif stage == GenerationRecordPipelineStage.RETRY_WAITING.value:
if (
ensure_aware_utc(download_next_retry_at) is None
or ensure_aware_utc(download_next_retry_at) <= now
):
download.append(ref)
elif str(provider_task_id or "").strip():
checked_next_poll_at = ensure_aware_utc(next_poll_at)
checked_poll_lease_until = ensure_aware_utc(poll_lease_until)
if (
(checked_next_poll_at is None or checked_next_poll_at <= now)
and (
checked_poll_lease_until is None
or checked_poll_lease_until <= now
)
):
poll.append(ref)
else:
if (
ensure_aware_utc(provider_create_lease_until) is None
or ensure_aware_utc(provider_create_lease_until) <= now
):
create.append(ref)
next_cursor = None
if len(rows) == page_size:
last = rows[-1]
next_cursor = GenerationRecordRecoveryCursor(owner_id=str(last.id))
return GenerationRecordRecoveryBatch(
create=create,
poll=poll,
download=download,
next_cursor=next_cursor,
)
@@ -5,7 +5,7 @@ from datetime import datetime, timedelta, timezone
from app.config import settings
from app.enums.generation_task import GenerationType
from app.models.chat_generation_task import ChatGenerationTask
from app.services.generation.pipeline.owner_service import GenerationOwner
from app.services.redis_registry_service import ensure_aware_utc
@@ -23,7 +23,7 @@ def utc_now() -> datetime:
return datetime.now(timezone.utc)
def is_video_generation_task(task: ChatGenerationTask) -> bool:
def is_video_generation_task(task: GenerationOwner) -> bool:
return str(getattr(task, "gen_type", "") or "").lower() == GenerationType.VIDEO.value
@@ -33,7 +33,7 @@ def video_final_deadline_from(now: datetime | None = None) -> datetime:
return current_time + timedelta(hours=hours)
def ensure_video_poll_fields(task: ChatGenerationTask, *, now: datetime | None = None) -> None:
def ensure_video_poll_fields(task: GenerationOwner, *, now: datetime | None = None) -> None:
"""补齐视频轮询调度字段,兼容历史任务。"""
if not is_video_generation_task(task):
return
@@ -42,17 +42,22 @@ def ensure_video_poll_fields(task: ChatGenerationTask, *, now: datetime | None =
if ensure_aware_utc(getattr(task, "poll_started_at", None)) is None:
task.poll_started_at = current_time
if ensure_aware_utc(getattr(task, "deadline_at", None)) is None:
task.deadline_at = video_final_deadline_from(current_time)
resource_started_at = (
ensure_aware_utc(getattr(task, "resource_generation_started_at", None))
or ensure_aware_utc(getattr(task, "created_at", None))
or current_time
)
task.deadline_at = video_final_deadline_from(resource_started_at)
if getattr(task, "poll_interval_seconds", None) is None:
task.poll_interval_seconds = 0
def is_final_poll_due(task: ChatGenerationTask, *, now: datetime | None = None) -> bool:
def is_final_poll_due(task: GenerationOwner, *, now: datetime | None = None) -> bool:
deadline_at = ensure_aware_utc(getattr(task, "deadline_at", None))
return bool(deadline_at and deadline_at <= (now or utc_now()))
def is_poll_not_due(task: ChatGenerationTask, *, now: datetime | None = None, tolerance_seconds: int = 1) -> bool:
def is_poll_not_due(task: GenerationOwner, *, now: datetime | None = None, tolerance_seconds: int = 1) -> bool:
"""判断当前 poll 任务是否早于 next_poll_at。只对视频降频轮询生效。"""
if not is_video_generation_task(task):
return False
@@ -73,7 +78,7 @@ def _clamp_positive_seconds(value: int | float | None, default: int) -> int:
def build_video_pending_poll_schedule(
task: ChatGenerationTask,
task: GenerationOwner,
*,
now: datetime | None = None,
) -> PollScheduleDecision:
@@ -133,7 +138,7 @@ def build_video_pending_poll_schedule(
def build_default_poll_schedule(
task: ChatGenerationTask,
task: GenerationOwner,
*,
now: datetime | None = None,
delay_seconds: int | None = None,
@@ -75,7 +75,7 @@ async def _build_user_content(record: ChatGenerationTask, db: AsyncSession | Non
async def _get_model_config(db: AsyncSession) -> ModelConfig:
result = await db.execute(
select(ModelConfig)
.where(ModelConfig.is_active == True)
.where(ModelConfig.is_active == True, ModelConfig.deleted_at.is_(None))
.order_by(ModelConfig.priority.desc())
.limit(1)
)
@@ -11,6 +11,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask
from app.services.generation.pipeline.owner_service import GenerationOwner, owner_provider_task_id
from app.models.image_engine import ImageEngine
from app.models.video_engine import VideoEngine
from app.services.generation.log_service import log_provider_call
@@ -39,7 +40,7 @@ def _try_json(value: Any) -> Any:
return None
async def get_runtime_engine(db: AsyncSession, task: ChatGenerationTask) -> Any:
async def get_runtime_engine(db: AsyncSession, task: GenerationOwner) -> Any:
"""使用任务快照冻结历史参数,只从当前引擎记录读取密钥。"""
snapshot = _loads(task.engine_snapshot_json)
if not task.engine_id:
@@ -50,7 +51,7 @@ async def get_runtime_engine(db: AsyncSession, task: ChatGenerationTask) -> Any:
result = await db.execute(select(VideoEngine).where(VideoEngine.id == task.engine_id).limit(1))
engine = result.scalar_one_or_none()
if not engine:
raise ValueError("引擎不存在或已删除")
raise ValueError("任务绑定的引擎不存在")
return SimpleNamespace(
id=task.engine_id,
name=snapshot.get("name") or engine.name,
@@ -89,7 +90,7 @@ async def get_runtime_engine(db: AsyncSession, task: ChatGenerationTask) -> Any:
)
async def create_provider_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
async def create_provider_task(db: AsyncSession, task: GenerationOwner) -> dict:
if task.gen_type == "video":
return await _create_video_task(db, task)
if task.gen_type == "image":
@@ -97,12 +98,14 @@ async def create_provider_task(db: AsyncSession, task: ChatGenerationTask) -> di
raise ValueError(f"不支持的生成类型: {task.gen_type}")
async def _create_video_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
async def _create_video_task(db: AsyncSession, task: GenerationOwner) -> dict:
engine = await get_runtime_engine(db, task)
# Close the engine lookup transaction before the long provider HTTP call.
await db.commit()
started = time.perf_counter()
async with provider_limit("ark_video_create", settings.ARK_VIDEO_CREATE_MAX_CONCURRENCY):
try:
provider_task_id = await submit_video_task(None, engine, task, include_media_references=True)
provider_task_id = await submit_video_task(None, engine, task, include_media_references=isinstance(task, ChatGenerationTask))
response = {"task_id": provider_task_id}
await log_provider_call(
task,
@@ -132,11 +135,13 @@ async def _create_video_task(db: AsyncSession, task: ChatGenerationTask) -> dict
async def create_image_sync_batch_result(
db: AsyncSession,
task: ChatGenerationTask,
task: GenerationOwner,
*,
generation_count: int,
) -> ImageProviderBatchResult:
engine = await get_runtime_engine(db, task)
# Do not keep a database transaction open while the synchronous provider call runs.
await db.commit()
return await create_image_sync_batch_result_with_engine(
task,
engine,
@@ -145,7 +150,7 @@ async def create_image_sync_batch_result(
async def create_image_sync_batch_result_with_engine(
task: ChatGenerationTask,
task: GenerationOwner,
engine: Any,
*,
generation_count: int,
@@ -164,7 +169,7 @@ async def create_image_sync_batch_result_with_engine(
None,
engine,
task,
include_media_references=True,
include_media_references=isinstance(task, ChatGenerationTask),
generation_count=count,
)
response_data = result.get("response_data") or result
@@ -197,7 +202,7 @@ async def create_image_sync_batch_result_with_engine(
raise
async def create_image_sync_result(db: AsyncSession, task: ChatGenerationTask) -> dict:
async def create_image_sync_result(db: AsyncSession, task: GenerationOwner) -> dict:
result = await create_image_sync_batch_result(db, task, generation_count=1)
items = result.get("items") or []
if len(items) != 1:
@@ -216,9 +221,13 @@ async def create_image_sync_result(db: AsyncSession, task: ChatGenerationTask) -
}
async def poll_provider_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
async def poll_provider_task(db: AsyncSession, task: GenerationOwner) -> dict:
engine = await get_runtime_engine(db, task)
task_id = task.seedance_task_id or task.provider_task_id
# Polling may block on the remote provider; release the lookup transaction first.
await db.commit()
task_id = owner_provider_task_id(task)
if not task_id:
raise ValueError("缺少供应商任务ID")
if task.gen_type == "video":
async with provider_limit("ark_video_poll", settings.ARK_VIDEO_POLL_MAX_CONCURRENCY):
return await poll_task_status(engine, task_id)
@@ -16,6 +16,7 @@ from app.enums.generation_task import (
ChatGenerationTaskEventType,
ChatGenerationTaskStatus,
GenerationMode,
GenerationOwnerType,
GenerationType,
)
from app.models.chat_generation_task import ChatGenerationTask
@@ -27,9 +28,15 @@ 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.pipeline.db_lock_service import apply_short_lock_timeout
from app.services.generation.pipeline.lifecycle_service import notify_owner_finished
from app.services.generation.poll_schedule_service import ensure_video_poll_fields, is_poll_not_due, is_video_generation_task
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.generation.pipeline.owner_service import (
load_generation_owner,
parse_redis_owner_item_id,
redis_owner_item_id,
)
from app.services.redis_registry_service import (
redis_get_due_registry_ids,
redis_get_registry_payloads,
@@ -42,6 +49,26 @@ logger = logging.getLogger("video_gen")
POLL_QUEUE = CeleryQueue.GEN_PROVIDER_POLL.value
def _chat_registry_id(task: ChatGenerationTask) -> str:
return redis_owner_item_id(
GenerationOwnerType.CHAT_GENERATION_TASK.value,
str(task.id),
int(getattr(task, "generation_attempt_no", 1) or 1),
)
async def _load_chat_task_for_update(
db: AsyncSession, task_id: str
) -> ChatGenerationTask | None:
owner = await load_generation_owner(
db,
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
owner_id=str(task_id),
for_update=True,
)
return owner if isinstance(owner, ChatGenerationTask) else None
def _now() -> datetime:
return datetime.now(timezone.utc)
@@ -140,26 +167,47 @@ async def recover_one_download_task(
if not task:
return "skip_missing_task"
if task.generation_mode not in ALLOWED_GENERATION_MODES:
await remove_download_active(task.id)
await remove_download_active(_chat_registry_id(task))
return "clean_invalid_mode"
if _is_final_task_state(task):
await remove_download_active(task.id)
await remove_download_active(_chat_registry_id(task))
return "clean_final_state"
if task.status != ChatGenerationTaskStatus.GENERATING.value:
await remove_download_active(task.id)
task_id = str(task.id)
generation_attempt_no = int(task.generation_attempt_no or 1)
generation_mode = str(task.generation_mode or "")
status = task.status
stage = task.pipeline_stage
registry_id = _chat_registry_id(task)
await db.rollback()
await remove_download_active(registry_id)
await log_task_event(
task,
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
owner_id=task_id,
task_id=task_id,
generation_attempt_no=generation_attempt_no,
generation_mode=generation_mode,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NOT_GENERATING.value,
message=f"{source} 下载恢复跳过:任务不是 generating",
detail={"status": task.status, "stage": task.pipeline_stage},
detail={"status": status, "stage": stage},
)
return "clean_not_generating"
if not task.remote_result_url:
task_id = str(task.id)
generation_attempt_no = int(task.generation_attempt_no or 1)
generation_mode = str(task.generation_mode or "")
status = task.status
stage = task.pipeline_stage
await db.rollback()
await log_task_event(
task,
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
owner_id=task_id,
task_id=task_id,
generation_attempt_no=generation_attempt_no,
generation_mode=generation_mode,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL.value,
message=f"{source} 下载恢复跳过:缺少 remote_result_url",
detail={"status": task.status, "stage": task.pipeline_stage},
detail={"status": status, "stage": stage},
)
return "skip_no_remote_result_url"
@@ -167,38 +215,38 @@ async def recover_one_download_task(
redis_payload = payload or {}
if stage == ChatGenerationPipelineStage.RESULT_READY.value:
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 result_ready 未完成下载,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_result_ready",
)
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 result_ready 未完成下载,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
return "recover_result_ready"
if stage == DOWNLOAD_STAGE_QUEUED:
if _is_queue_timeout(task, current_time):
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 download_queued 长时间未消费,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_download_queued_timeout",
)
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 download_queued 长时间未消费,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
return "recover_queued_timeout"
await postpone_download_active_check(
record_id=task.id,
record_id=_chat_registry_id(task),
payload=payload,
check_at=_queue_timeout_at(task, current_time),
)
@@ -206,22 +254,22 @@ async def recover_one_download_task(
if stage == DOWNLOAD_STAGE_DOWNLOADING:
if _is_expired(task.download_lease_until, current_time):
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 downloading lease 过期,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_downloading_lease_expired",
)
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 downloading lease 过期,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
return "recover_downloading_expired"
await postpone_download_active_check(
record_id=task.id,
record_id=_chat_registry_id(task),
payload=payload,
check_at=task.download_lease_until,
)
@@ -229,22 +277,22 @@ async def recover_one_download_task(
if stage == DOWNLOAD_STAGE_RETRY_WAITING:
if _is_expired(task.download_next_retry_at, current_time):
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 retry_waiting 到期,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_retry_waiting_due",
)
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 retry_waiting 到期,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
return "recover_retry_due"
await postpone_download_active_check(
record_id=task.id,
record_id=_chat_registry_id(task),
payload=payload,
check_at=task.download_next_retry_at,
)
@@ -267,33 +315,42 @@ async def recover_download_tasks_once(db: AsyncSession) -> dict[str, Any]:
)
payloads = await get_download_active_payloads(due_ids)
for task_id in due_ids:
result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
task = result.scalar_one_or_none()
download_refs = {
registry_item_id: parse_redis_owner_item_id(registry_item_id)
for registry_item_id in due_ids
}
for registry_item_id, ref in download_refs.items():
if ref.owner_type != GenerationOwnerType.CHAT_GENERATION_TASK.value:
continue
task = await _load_chat_task_for_update(db, ref.owner_id)
if task is None:
await remove_download_active(task_id)
await db.rollback()
await remove_download_active(registry_item_id)
action = "clean_missing_task"
elif (
ref.generation_attempt_no is not None
and int(task.generation_attempt_no or 1)
!= int(ref.generation_attempt_no)
):
await db.rollback()
await remove_download_active(registry_item_id)
action = "clean_stale_download_attempt"
else:
current_registry_id = _chat_registry_id(task)
if registry_item_id != current_registry_id:
await remove_download_active(registry_item_id)
checked_ids.add(task.id)
action = await recover_one_download_task(
db,
task,
payload=payloads.get(task_id),
payload=payloads.get(registry_item_id),
source="startup_redis",
)
results[action] = results.get(action, 0) + 1
# DB fallback:不依赖 Redis active 注册表。
fallback_result = await db.execute(
select(ChatGenerationTask)
select(ChatGenerationTask.id)
.where(
ChatGenerationTask.deleted_at.is_(None),
ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)),
@@ -310,12 +367,15 @@ async def recover_download_tasks_once(db: AsyncSession) -> dict[str, Any]:
)
.order_by(ChatGenerationTask.updated_at.asc())
.limit(int(settings.DOWNLOAD_RECOVERY_BATCH_SIZE or 100))
.with_for_update(skip_locked=True)
)
fallback_tasks = fallback_result.scalars().all()
fallback_ids = [str(value) for value in fallback_result.scalars().all()]
for task in fallback_tasks:
if task.id in checked_ids:
for task_id in fallback_ids:
if task_id in checked_ids:
continue
task = await _load_chat_task_for_update(db, task_id)
if task is None:
await db.rollback()
continue
action = await recover_one_download_task(
db,
@@ -324,7 +384,7 @@ async def recover_download_tasks_once(db: AsyncSession) -> dict[str, Any]:
source="startup_db",
)
results[action] = results.get(action, 0) + 1
checked_ids.add(task.id)
checked_ids.add(task_id)
return {"checked": len(checked_ids), "results": results}
@@ -341,11 +401,9 @@ async def _mark_timeout(
error_message=error_message,
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
)
await notify_chat_generation_task_finished(db, task)
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
await aggregate_parent_for_child(db, task)
await db.commit()
await _remove_poll_active(task.id)
await notify_owner_finished(db, task)
await _remove_poll_active(_chat_registry_id(task))
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
@@ -369,11 +427,9 @@ async def _mark_failed(
error_message=error_message,
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
)
await notify_chat_generation_task_finished(db, task)
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
await aggregate_parent_for_child(db, task)
await db.commit()
await _remove_poll_active(task.id)
await notify_owner_finished(db, task)
await _remove_poll_active(_chat_registry_id(task))
await log_task_event(task, event_type=event_type, message=task.error_message, detail=detail)
return "mark_failed"
@@ -403,19 +459,19 @@ async def recover_one_generation_task(
if not task:
return "skip_missing_task"
if task.generation_mode not in ALLOWED_GENERATION_MODES:
await _remove_poll_active(task.id)
await _remove_poll_active(_chat_registry_id(task))
return "clean_invalid_mode"
if _is_final_task_state(task):
await _remove_poll_active(task.id)
await _remove_poll_active(_chat_registry_id(task))
return "clean_final_state"
if task.status != ChatGenerationTaskStatus.GENERATING.value:
await _remove_poll_active(task.id)
await _remove_poll_active(_chat_registry_id(task))
return "clean_not_generating"
if bool(getattr(task, "video_upscale_enabled_snapshot", False)) and str(task.pipeline_stage or "").startswith("upscale_"):
# 原视频已经进入超分流水线,后续由 video_upscale 恢复扫描处理。
# 这里禁止再次投递原结果下载,避免覆盖保留的 source.mp4 或提前生成用户资源。
await _remove_poll_active(task.id)
await _remove_poll_active(_chat_registry_id(task))
return "delegate_video_upscale_recovery"
has_remote_result = bool(str(task.remote_result_url or "").strip())
@@ -425,7 +481,13 @@ async def recover_one_generation_task(
# 最高优先级:只要远程结果 URL 已经落库,说明生成侧已经成功。
# 不管当前 pipeline_stage 是 queued/creating/waiting/result_ready/download_*,恢复时都不能重复 create 或 poll。
if has_remote_result:
await _remove_poll_active(task.id)
await _remove_poll_active(_chat_registry_id(task))
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_has_remote_result_url",
)
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
@@ -436,12 +498,6 @@ async def recover_one_generation_task(
"deadline_expired": is_deadline_expired,
},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_has_remote_result_url",
)
return "recover_download_has_remote_result"
# 已经过 deadline 且没有结果 URL
@@ -459,7 +515,7 @@ async def recover_one_generation_task(
)
poll_generation_task.apply_async(
args=[task.id],
kwargs={"force_due": True},
kwargs={"force_due": True, "owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
queue=POLL_QUEUE,
countdown=0,
)
@@ -470,13 +526,11 @@ async def recover_one_generation_task(
)
return "recover_deadline_final_poll"
await log_task_event(
return await _mark_timeout(
db,
task,
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_TIMEOUT.value,
message=f"{source} 发现任务已到 deadline,且没有 remote_result_url/供应商任务ID,按超时失败处理",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
error_message=f"{source} 发现任务已到 deadline,且没有远程结果或供应商任务ID",
)
return await _mark_timeout(db, task)
# 未过 deadline:有供应商任务 ID 才允许恢复到 poll 队列。
# 视频任务如果 next_poll_at 未到期,不提前 poll,只刷新 active 注册表等待 Beat dispatcher 到期投递。
@@ -523,7 +577,7 @@ async def recover_one_generation_task(
)
poll_generation_task.apply_async(
args=[task.id],
kwargs={"force_due": True},
kwargs={"force_due": True, "owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
queue=POLL_QUEUE,
countdown=0,
)
@@ -552,9 +606,11 @@ async def recover_one_generation_task(
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
):
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
await db.commit()
# Release the recovery row lock before writing an event through the
# independent logging session or talking to the broker.
await db.commit()
await _remove_poll_active(task.id)
await _remove_poll_active(_chat_registry_id(task))
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
@@ -563,6 +619,7 @@ async def recover_one_generation_task(
)
chatapi_create_generation_task.apply_async(
args=[task.id],
kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
countdown=0,
)
@@ -572,7 +629,7 @@ async def recover_one_generation_task(
if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value:
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
await db.commit()
await _remove_poll_active(task.id)
await _remove_poll_active(_chat_registry_id(task))
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
@@ -581,6 +638,7 @@ async def recover_one_generation_task(
)
chatapi_create_generation_task.apply_async(
args=[task.id],
kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
countdown=0,
)
@@ -621,28 +679,33 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
if image_main_cursor:
image_main_query = image_main_query.where(ChatGenerationTask.id > image_main_cursor)
image_main_result = await db.execute(
image_main_query.order_by(ChatGenerationTask.id.asc())
image_main_query.with_only_columns(ChatGenerationTask.id)
.order_by(ChatGenerationTask.id.asc())
.limit(image_main_batch_size)
.with_for_update()
)
image_mains = list(image_main_result.scalars().all())
if not image_mains:
image_main_ids = [str(value) for value in image_main_result.scalars().all()]
if not image_main_ids:
break
for main in image_mains:
main_id = str(main.id)
child_parent_result = await db.execute(
select(ChatGenerationTask.parent_task_id)
.where(
ChatGenerationTask.parent_task_id.in_(image_main_ids),
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
)
.distinct()
)
split_parent_ids = {str(value) for value in child_parent_result.scalars().all() if value}
for main_id in image_main_ids:
image_main_cursor = main_id
checked_ids.add(main_id)
main = await _load_chat_task_for_update(db, main_id)
if main is None:
await db.rollback()
continue
child_result = await db.execute(
select(ChatGenerationTask.id)
.where(
ChatGenerationTask.parent_task_id == main_id,
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
)
.limit(1)
)
if child_result.scalar_one_or_none() is not None:
if main_id in split_parent_ids:
main.provider_create_claim_token = None
main.provider_create_lease_until = None
await db.commit()
@@ -653,7 +716,7 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
lease_until = ensure_aware_utc(main.provider_create_lease_until)
lease_alive = bool(main.provider_create_claim_token and lease_until and lease_until > now)
if lease_alive:
await db.commit()
await db.rollback()
results["image_main_claim_alive"] = results.get("image_main_claim_alive", 0) + 1
continue
@@ -670,19 +733,24 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
results["image_main_timeout"] = results.get("image_main_timeout", 0) + 1
continue
claim_expired = False
if main.provider_create_claim_token or main.provider_create_lease_until:
main.provider_create_claim_token = None
main.provider_create_lease_until = None
main.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
claim_expired = True
attempt_no = int(main.generation_attempt_no or 1)
await db.commit()
if claim_expired:
await log_task_event(
main,
event_type=ChatGenerationTaskEventType.IMAGE_MAIN_CLAIM_EXPIRED.value,
message="图片主任务供应商执行租约已过期,恢复重新投递",
)
await db.commit()
try:
chatapi_create_generation_task.apply_async(
args=[main_id],
kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": attempt_no},
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
countdown=0,
)
@@ -691,7 +759,7 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
logger.exception("恢复投递图片主任务失败 task_id=%s: %s", main_id, exc)
results["recover_image_main_enqueue_failed"] = results.get("recover_image_main_enqueue_failed", 0) + 1
if len(image_mains) < image_main_batch_size:
if len(image_main_ids) < image_main_batch_size:
break
due_poll_ids = await redis_get_due_registry_ids(
@@ -705,26 +773,34 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
log_context="poll_active",
)
for task_id in due_poll_ids:
result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
task = result.scalar_one_or_none()
poll_refs = {
registry_item_id: parse_redis_owner_item_id(registry_item_id)
for registry_item_id in due_poll_ids
}
for registry_item_id, ref in poll_refs.items():
if ref.owner_type != GenerationOwnerType.CHAT_GENERATION_TASK.value:
continue
task = await _load_chat_task_for_update(db, ref.owner_id)
if task is None:
await _remove_poll_active(task_id)
await db.rollback()
await _remove_poll_active(registry_item_id)
action = "clean_missing_poll_task"
elif (
ref.generation_attempt_no is not None
and int(task.generation_attempt_no or 1) != int(ref.generation_attempt_no)
):
await db.rollback()
await _remove_poll_active(registry_item_id)
action = "clean_stale_poll_attempt"
else:
current_registry_id = _chat_registry_id(task)
if registry_item_id != current_registry_id:
await _remove_poll_active(registry_item_id)
checked_ids.add(task.id)
action = await recover_one_generation_task(
db,
task,
payload=poll_payloads.get(task_id),
payload=poll_payloads.get(registry_item_id),
source="startup_poll_redis",
)
results[action] = results.get(action, 0) + 1
@@ -735,7 +811,7 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
for _round in range(max_rounds):
query_result = await db.execute(
select(ChatGenerationTask)
select(ChatGenerationTask.id)
.where(
ChatGenerationTask.deleted_at.is_(None),
ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)),
@@ -753,15 +829,18 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
)
.order_by(ChatGenerationTask.updated_at.asc())
.limit(batch_size)
.with_for_update(skip_locked=True)
)
tasks = query_result.scalars().all()
if not tasks:
task_ids = [str(value) for value in query_result.scalars().all()]
if not task_ids:
break
progressed_this_round = 0
for task in tasks:
if task.id in checked_ids:
for task_id in task_ids:
if task_id in checked_ids:
continue
task = await _load_chat_task_for_update(db, task_id)
if task is None:
await db.rollback()
continue
action = await recover_one_generation_task(
db,
@@ -770,11 +849,11 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
source="startup_db",
)
results[action] = results.get(action, 0) + 1
checked_ids.add(task.id)
checked_ids.add(task_id)
total_db_checked += 1
progressed_this_round += 1
if len(tasks) < batch_size or progressed_this_round <= 0:
if len(task_ids) < batch_size or progressed_this_round <= 0:
break
# 子任务可能在 worker 中断前已进入终态但主任务尚未汇总,按稳定游标完整重算全部主任务。
@@ -820,6 +899,7 @@ async def dispatch_due_poll_tasks_once(db: AsyncSession) -> dict[str, Any]:
batch_size = max(1, int(settings.POLL_DUE_DISPATCH_BATCH_SIZE or 100))
poll_lease_expired_at = current_time - timedelta(seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300))
await apply_short_lock_timeout(db)
query_result = await db.execute(
select(ChatGenerationTask)
.where(
@@ -846,6 +926,7 @@ async def dispatch_due_poll_tasks_once(db: AsyncSession) -> dict[str, Any]:
dispatched_task_ids: list[str] = []
dispatched_due_next_poll_at_by_id: dict[str, datetime | None] = {}
dispatched_queue_hold_until_by_id: dict[str, datetime] = {}
post_commit_logs: list[dict[str, Any]] = []
for task in tasks:
action = "skip_unknown"
@@ -863,13 +944,14 @@ async def dispatch_due_poll_tasks_once(db: AsyncSession) -> dict[str, Any]:
continue
if not (task.seedance_task_id or task.provider_task_id):
# dispatcher 不负责重新 create;没有 provider id 的异常状态交给启动容灾或 create 任务处理。
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.POLL_DISPATCH_SKIP.value,
message="视频到期轮询调度跳过:缺少外部任务ID",
detail={"pipeline_stage": task.pipeline_stage, "next_poll_at": task.next_poll_at},
)
# Defer FK-backed event logging until the batch row locks are released.
post_commit_logs.append({
"task_id": str(task.id),
"generation_attempt_no": int(task.generation_attempt_no or 1),
"event_type": ChatGenerationTaskEventType.POLL_DISPATCH_SKIP.value,
"message": "视频到期轮询调度跳过:缺少外部任务ID",
"detail": {"pipeline_stage": task.pipeline_stage, "next_poll_at": task.next_poll_at},
})
action = "skip_no_provider_task_id"
continue
@@ -886,16 +968,28 @@ async def dispatch_due_poll_tasks_once(db: AsyncSession) -> dict[str, Any]:
except Exception as exc:
logger.exception("视频到期轮询调度单条处理失败。task_id=%s", getattr(task, "id", None))
action = "error"
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.POLL_DISPATCH_SKIP.value,
message=f"视频到期轮询调度单条处理失败:{exc}",
)
post_commit_logs.append({
"task_id": str(getattr(task, "id", "") or ""),
"generation_attempt_no": int(getattr(task, "generation_attempt_no", 1) or 1),
"event_type": ChatGenerationTaskEventType.POLL_DISPATCH_SKIP.value,
"message": f"视频到期轮询调度单条处理失败:{exc}",
"detail": None,
})
finally:
results[action] = results.get(action, 0) + 1
await db.commit()
for item in post_commit_logs:
if item["task_id"]:
await log_task_event(
task_id=item["task_id"],
generation_attempt_no=item["generation_attempt_no"],
event_type=item["event_type"],
message=item["message"],
detail=item["detail"],
)
fresh_tasks = []
if dispatched_task_ids:
fresh_result = await db.execute(
@@ -931,7 +1025,7 @@ async def dispatch_due_poll_tasks_once(db: AsyncSession) -> dict[str, Any]:
)
poll_generation_task.apply_async(
args=[task.id],
kwargs={"force_due": True},
kwargs={"force_due": True, "owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
queue=POLL_QUEUE,
countdown=0,
)
@@ -7,6 +7,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from app.models.chat_generation_task import ChatGenerationTask
from app.models.credit_record import CreditRecord
from app.models.generation_record import GenerationRecord
from app.services.generation.pipeline.db_lock_service import execute_with_lock_timeout
from app.services.credits import refund_credits
from app.services.credit_record_meta_service import build_refund_meta_from_charge
from app.services.generation.billing_service import (
@@ -25,46 +26,52 @@ def _round2(value: float | int | None) -> float:
return round(float(value or 0), 2)
async def _has_refund_for_biz_key(db: AsyncSession, *, user_id: str, charge_biz_key: str) -> bool:
result = await db.execute(
select(CreditRecord.id)
.where(
CreditRecord.user_id == user_id,
CreditRecord.type == "refund",
CreditRecord.refund_for_biz_key == charge_biz_key,
)
.limit(1)
)
return result.scalar_one_or_none() is not None
async def _find_unrefunded_media_charges(
db: AsyncSession,
*,
user_id: str,
owner_type: str,
owner_id: str,
generation_attempt_no: int | None = None,
) -> list[CreditRecord]:
"""查找当前任务下所有未退款的媒体生成扣费流水"""
pattern = f"{owner_type}:{owner_id}:attempt:%:{CHARGE_MEDIA}:charge"
"""批量查找指定任务、指定生成轮次下尚未退款的媒体扣费"""
result = await db.execute(
select(CreditRecord)
.where(
CreditRecord.user_id == user_id,
CreditRecord.related_id == owner_id,
CreditRecord.type == "consume",
CreditRecord.biz_key.like(pattern),
CreditRecord.biz_key.is_not(None),
)
.order_by(CreditRecord.created_at.asc())
)
charges = list(result.scalars().all())
unrefunded: list[CreditRecord] = []
for charge in charges:
if not charge.biz_key:
charges: list[CreditRecord] = []
for charge in result.scalars().all():
parsed = parse_credit_biz_key(charge.biz_key)
if not parsed:
continue
if not await _has_refund_for_biz_key(db, user_id=user_id, charge_biz_key=charge.biz_key):
unrefunded.append(charge)
return unrefunded
if parsed.get("owner_type") != owner_type or parsed.get("owner_id") != owner_id:
continue
if parsed.get("charge_kind") != CHARGE_MEDIA or parsed.get("action") != "charge":
continue
if generation_attempt_no is not None and int(parsed.get("attempt_no") or 0) != int(generation_attempt_no):
continue
charges.append(charge)
charge_keys = [str(charge.biz_key) for charge in charges if charge.biz_key]
if not charge_keys:
return []
refund_result = await db.execute(
select(CreditRecord.refund_for_biz_key).where(
CreditRecord.user_id == user_id,
CreditRecord.type == "refund",
CreditRecord.refund_for_biz_key.in_(charge_keys),
)
)
refunded_keys = {str(value) for (value,) in refund_result.all() if value}
return [charge for charge in charges if charge.biz_key not in refunded_keys]
async def refund_unrefunded_media_charges(
@@ -74,8 +81,9 @@ async def refund_unrefunded_media_charges(
owner_type: str,
owner_id: str,
description_prefix: str,
generation_attempt_no: int | None = None,
) -> float:
"""回退当前任务所有未退款媒体扣费流水。
"""回退当前生成轮次尚未退款媒体扣费流水。
容灾考虑
- 不依赖 retry_count 推断当前轮次
@@ -88,6 +96,7 @@ async def refund_unrefunded_media_charges(
user_id=user_id,
owner_type=owner_type,
owner_id=owner_id,
generation_attempt_no=generation_attempt_no,
)
for charge in charges:
parsed = parse_credit_biz_key(charge.biz_key)
@@ -124,6 +133,7 @@ async def mark_generation_record_failed_and_refund_once(
record_id: str | None = None,
record: GenerationRecord | None = None,
error_message: str | None = None,
generation_attempt_no: int | None = None,
) -> GenerationRecord | None:
"""把 GenerationRecord 标记为最终失败并幂等退回媒体生成积分。
@@ -132,7 +142,9 @@ async def mark_generation_record_failed_and_refund_once(
if record is None:
if not record_id:
return None
result = await db.execute(
result = await execute_with_lock_timeout(
db,
select(GenerationRecord)
.where(GenerationRecord.id == record_id, GenerationRecord.deleted_at.is_(None))
.with_for_update()
@@ -155,6 +167,7 @@ async def mark_generation_record_failed_and_refund_once(
owner_type=OWNER_GENERATION_RECORD,
owner_id=record.id,
description_prefix="生成记录",
generation_attempt_no=generation_attempt_no or int(getattr(record, "generation_attempt_no", 1) or 1),
)
if refunded_amount > 0:
# GenerationRecord.credits_cost 只代表视频/图片生成媒体积分。
@@ -171,12 +184,15 @@ async def mark_chat_generation_task_failed_and_refund_once(
task: ChatGenerationTask | None = None,
error_message: str | None = None,
pipeline_stage: str = "failed",
generation_attempt_no: int | None = None,
) -> ChatGenerationTask | None:
"""把 ChatGenerationTask 标记为最终失败并幂等退回媒体生成积分。"""
if task is None:
if not task_id:
return None
result = await db.execute(
result = await execute_with_lock_timeout(
db,
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == task_id,
@@ -204,6 +220,7 @@ async def mark_chat_generation_task_failed_and_refund_once(
owner_type=OWNER_CHAT_GENERATION_TASK,
owner_id=task.id,
description_prefix="任务生成",
generation_attempt_no=generation_attempt_no or int(getattr(task, "generation_attempt_no", 1) or 1),
)
if refunded_amount > 0:
@@ -132,6 +132,9 @@ async def create_chat_generation_task_for_module(
snapshot["generation_count"] = 1
task = ChatGenerationTask(
id=task_id,
created_at=now,
resource_generation_started_at=now,
generation_attempt_no=1,
user_id=current_user.id,
original_prompt=original_prompt,
optimized_prompt=optimized_prompt,
@@ -195,6 +198,9 @@ async def create_chat_generation_task_for_module(
snapshot["video_upscale_enabled_snapshot"] = upscale_enabled_snapshot
task = ChatGenerationTask(
id=task_id,
created_at=now,
resource_generation_started_at=now,
generation_attempt_no=1,
user_id=current_user.id,
original_prompt=original_prompt,
optimized_prompt=optimized_prompt,
@@ -1,5 +1,6 @@
from __future__ import annotations
import asyncio
import json
from datetime import datetime
from typing import Any
@@ -42,6 +43,11 @@ from app.services.generation.ai.engine_service import (
)
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.pipeline.db_lock_service import (
DatabaseRowLockBusy,
apply_short_lock_timeout,
execute_with_lock_timeout,
)
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_error, log_module_event_file, log_module_prompt_event
@@ -328,6 +334,18 @@ async def _soft_delete_steps_from_index(
start_index: int,
deleted_at: datetime | None = None,
) -> None:
processing_result = await db.execute(
select(ModuleGenerationStep.id).where(
ModuleGenerationStep.project_id == project.id,
ModuleGenerationStep.module == MODULE,
ModuleGenerationStep.deleted_at.is_(None),
ModuleGenerationStep.is_current == True,
ModuleGenerationStep.status == ModuleStepStatusEnum.PROCESSING.value,
ModuleGenerationStep.step_index >= start_index,
).limit(1)
)
if processing_result.scalar_one_or_none() is not None:
raise HTTPException(status_code=409, detail="当前步骤正在处理中,请等待完成后再操作")
await _base_soft_delete_steps_from_index(
db,
project=project,
@@ -750,6 +768,69 @@ async def update_hot_opening_video_prompt_schema(
)
async def _reload_prompt_context_for_update(
db: AsyncSession,
*,
project_id: str,
step_id: str,
step_code: str,
) -> tuple[ModuleGenerationProject | None, ModuleGenerationStep | None]:
last_error: DatabaseRowLockBusy | None = None
for retry_index in range(3):
try:
project_result = await execute_with_lock_timeout(
db,
select(ModuleGenerationProject)
.where(
ModuleGenerationProject.id == project_id,
ModuleGenerationProject.module == MODULE,
ModuleGenerationProject.deleted_at.is_(None),
)
.with_for_update()
.execution_options(populate_existing=True)
.limit(1)
)
project = project_result.scalar_one_or_none()
if project is None:
return None, None
step_result = await execute_with_lock_timeout(
db,
select(ModuleGenerationStep)
.where(
ModuleGenerationStep.id == step_id,
ModuleGenerationStep.project_id == project_id,
ModuleGenerationStep.module == MODULE,
ModuleGenerationStep.step_code == step_code,
ModuleGenerationStep.deleted_at.is_(None),
ModuleGenerationStep.is_current == True,
)
.with_for_update()
.execution_options(populate_existing=True)
.limit(1)
)
return project, step_result.scalar_one_or_none()
except DatabaseRowLockBusy as exc:
last_error = exc
await db.rollback()
if retry_index < 2:
await asyncio.sleep(1 + retry_index)
raise last_error or DatabaseRowLockBusy()
def _prompt_context_matches(
step: ModuleGenerationStep | None,
*,
expected_version: int,
expected_input_json: str,
) -> bool:
if step is None or step.status != ModuleStepStatusEnum.PROCESSING.value:
return False
if int(step.version or 1) != int(expected_version):
return False
current_input = json.dumps(step.input_json, ensure_ascii=False, sort_keys=True, default=str)
return current_input == expected_input_json
async def submit_image_prompt_optimize(
db: AsyncSession,
*,
@@ -782,6 +863,7 @@ async def submit_image_prompt_optimize(
async def run_image_prompt_optimize(db: AsyncSession, *, project_id: str, step_id: str | None = None) -> ModuleGenerationStep | None:
await apply_short_lock_timeout(db)
project_result = await db.execute(
select(ModuleGenerationProject)
.where(ModuleGenerationProject.id == project_id, ModuleGenerationProject.module == MODULE, ModuleGenerationProject.deleted_at.is_(None))
@@ -799,6 +881,7 @@ async def run_image_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
return None
if step_id:
await apply_short_lock_timeout(db)
result = await db.execute(
select(ModuleGenerationStep)
.where(
@@ -845,24 +928,45 @@ async def run_image_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
{"type": "video", "url": material.get("material_video_url"), "name": "参考素材视频"},
{"type": "image", "url": material.get("material_image_url"), "name": "新产品图片"},
]
project_id_value = str(project.id)
step_id_value = str(step.id)
user_id_value = str(project.user_id)
module_value = str(project.module)
expected_step_version = int(step.version or 1)
expected_input_json = json.dumps(step.input_json, ensure_ascii=False, sort_keys=True, default=str)
await db.commit()
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,
project_id=project_id_value,
step_id=step_id_value,
user_id=user_id_value,
module=module_value,
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,
user_id=user_id_value,
references=references,
gen_type="image",
)
project, step = await _reload_prompt_context_for_update(
db,
project_id=project_id_value,
step_id=step_id_value,
step_code=HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value,
)
if not _prompt_context_matches(
step,
expected_version=expected_step_version,
expected_input_json=expected_input_json,
):
await db.rollback()
return None
billing = await charge_module_prompt_usage(
db,
user_id=project.user_id,
@@ -907,7 +1011,25 @@ async def run_image_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
token_usage=usage,
)
await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.IMAGE_PROMPT_SUCCESS.value, message="图片 AI 提词生成成功")
await db.commit()
except DatabaseRowLockBusy:
await db.rollback()
raise
except Exception as exc:
await db.rollback()
project, step = await _reload_prompt_context_for_update(
db,
project_id=project_id_value,
step_id=step_id_value,
step_code=HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value,
)
if not _prompt_context_matches(
step,
expected_version=expected_step_version,
expected_input_json=expected_input_json,
):
await db.rollback()
return None
step.status = ModuleStepStatusEnum.FAILED.value
step.error_message = str(exc)
step.completed_at = _now()
@@ -925,6 +1047,7 @@ async def run_image_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
)
_log_project_error(project=project, step=step, event_type="IMAGE_PROMPT_FAILED", message=project.error_message, exc=exc)
await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.IMAGE_PROMPT_FAILED.value, message=project.error_message)
await db.commit()
return step
@@ -1093,6 +1216,7 @@ async def submit_video_prompt_optimize(
async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_id: str | None = None) -> ModuleGenerationStep | None:
await apply_short_lock_timeout(db)
project_result = await db.execute(
select(ModuleGenerationProject)
.where(ModuleGenerationProject.id == project_id, ModuleGenerationProject.module == MODULE, ModuleGenerationProject.deleted_at.is_(None))
@@ -1112,6 +1236,7 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
return None
if step_id:
await apply_short_lock_timeout(db)
result = await db.execute(
select(ModuleGenerationStep)
.where(
@@ -1151,8 +1276,17 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
step.error_message = "缺少新项目图片结果,不能生成视频提词"
project.status = ModuleProjectStatusEnum.FAILED.value
project.error_message = step.error_message
await db.commit()
return step
project_id_value = str(project.id)
step_id_value = str(step.id)
user_id_value = str(project.user_id)
module_value = str(project.module)
expected_step_version = int(step.version or 1)
expected_input_json = json.dumps(step.input_json, ensure_ascii=False, sort_keys=True, default=str)
await db.commit()
try:
request_log = {
"source_project_name": material.get("source_project_name") or "",
@@ -1165,10 +1299,10 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
}
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,
project_id=project_id_value,
step_id=step_id_value,
user_id=user_id_value,
module=module_value,
prompt_type=ModulePromptTypeEnum.VIDEO_PROMPT.value,
request=request_log,
)
@@ -1176,7 +1310,7 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
request_log["schema_config_source"] = schema_config_snapshot.get("source")
prompt_schema, final_prompt, token_usage = await optimize_hot_opening_video_prompt(
db,
user_id=project.user_id,
user_id=user_id_value,
source_project_name=request_log["source_project_name"],
target_project_name=request_log["target_project_name"],
core_content_point=request_log["core_content_point"],
@@ -1185,11 +1319,24 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
video_config=video_config,
target_platform=target_platform,
schema_config_snapshot=schema_config_snapshot,
module=project.module,
project_id=project.id,
step_id=step.id,
trace_id=f"hot-video-prompt:{step.id}",
module=module_value,
project_id=project_id_value,
step_id=step_id_value,
trace_id=f"hot-video-prompt:{step_id_value}",
)
project, step = await _reload_prompt_context_for_update(
db,
project_id=project_id_value,
step_id=step_id_value,
step_code=HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value,
)
if not _prompt_context_matches(
step,
expected_version=expected_step_version,
expected_input_json=expected_input_json,
):
await db.rollback()
return None
billing = await charge_module_prompt_usage(
db,
user_id=project.user_id,
@@ -1237,7 +1384,25 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
token_usage=usage,
)
await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.VIDEO_PROMPT_SUCCESS.value, message="视频 AI 提词生成成功")
await db.commit()
except DatabaseRowLockBusy:
await db.rollback()
raise
except Exception as exc:
await db.rollback()
project, step = await _reload_prompt_context_for_update(
db,
project_id=project_id_value,
step_id=step_id_value,
step_code=HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value,
)
if not _prompt_context_matches(
step,
expected_version=expected_step_version,
expected_input_json=expected_input_json,
):
await db.rollback()
return None
step.status = ModuleStepStatusEnum.FAILED.value
step.error_message = str(exc)
step.completed_at = _now()
@@ -1255,6 +1420,7 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
)
_log_project_error(project=project, step=step, event_type="VIDEO_PROMPT_FAILED", message=project.error_message, exc=exc)
await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.VIDEO_PROMPT_FAILED.value, message=project.error_message)
await db.commit()
return step
@@ -1366,9 +1532,37 @@ async def generate_video_from_prompt(
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(
meta_result = await db.execute(
select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where(
ModuleGenerationStep.chat_task_id == task.id,
ModuleGenerationStep.module == MODULE,
ModuleGenerationStep.is_current == True,
ModuleGenerationStep.deleted_at.is_(None),
).limit(1)
)
meta = meta_result.first()
if not meta:
return
step_id_value, project_id_value = str(meta.id), str(meta.project_id)
project_result = await execute_with_lock_timeout(
db,
select(ModuleGenerationProject)
.where(
ModuleGenerationProject.id == project_id_value,
ModuleGenerationProject.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
project = project_result.scalar_one_or_none()
if not project:
return
step_result = await execute_with_lock_timeout(
db,
select(ModuleGenerationStep)
.where(
ModuleGenerationStep.id == step_id_value,
ModuleGenerationStep.project_id == project_id_value,
ModuleGenerationStep.chat_task_id == task.id,
ModuleGenerationStep.module == MODULE,
ModuleGenerationStep.is_current == True,
@@ -1377,18 +1571,9 @@ async def handle_chat_generation_task_completed(db: AsyncSession, task: ChatGene
.with_for_update()
.limit(1)
)
step = result.scalar_one_or_none()
step = 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
@@ -1429,19 +1614,48 @@ async def handle_chat_generation_task_completed(db: AsyncSession, task: ChatGene
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))
meta_result = await db.execute(
select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where(
ModuleGenerationStep.chat_task_id == task.id,
ModuleGenerationStep.module == MODULE,
ModuleGenerationStep.is_current == True,
ModuleGenerationStep.deleted_at.is_(None),
).limit(1)
)
meta = meta_result.first()
if not meta:
return
step_id_value, project_id_value = str(meta.id), str(meta.project_id)
project_result = await execute_with_lock_timeout(
db,
select(ModuleGenerationProject)
.where(
ModuleGenerationProject.id == project_id_value,
ModuleGenerationProject.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_result = await execute_with_lock_timeout(
db,
select(ModuleGenerationStep)
.where(
ModuleGenerationStep.id == step_id_value,
ModuleGenerationStep.project_id == project_id_value,
ModuleGenerationStep.chat_task_id == task.id,
ModuleGenerationStep.module == MODULE,
ModuleGenerationStep.is_current == True,
ModuleGenerationStep.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
step = step_result.scalar_one_or_none()
if not step:
return
step.status = ModuleStepStatusEnum.FAILED.value
step.error_message = task.error_message
step.completed_at = _now()
@@ -1459,7 +1673,8 @@ async def mark_hot_opening_step_dispatch_failed(
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(
result = await execute_with_lock_timeout(
db,
select(ModuleGenerationStep)
.where(
ModuleGenerationStep.id == step_id,
@@ -1515,7 +1515,7 @@ def _log_video_prompt_ai_event(
)
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))
result = await db.execute(select(ModelConfig).where(ModelConfig.is_active == True, ModelConfig.deleted_at.is_(None)).order_by(ModelConfig.priority.desc()).limit(1))
return result.scalar_one_or_none()
@@ -1557,6 +1557,9 @@ async def optimize_hot_opening_video_prompt(
# return result, build_final_video_prompt(result), {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}
config = await _select_model_config(db)
# All module/project claims are committed by the caller. Release this
# configuration read transaction before the remote model request.
await db.commit()
if not config:
result = normalize_video_prompt_schema_from_ai(_mock_result(video_config, target_platform), video_config, schema_config_snapshot)
return result, build_final_video_prompt(result), {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}
+13 -2
View File
@@ -107,7 +107,7 @@ def _log_image_response(record_id: str, response_data: dict, error: str | None =
async def get_active_image_engine(db: AsyncSession) -> ImageEngine:
result = await db.execute(
select(ImageEngine)
.where(ImageEngine.is_active == True)
.where(ImageEngine.is_active == True, ImageEngine.deleted_at.is_(None))
.order_by(ImageEngine.priority.desc())
.limit(1)
)
@@ -455,12 +455,23 @@ async def poll_image_task_status(engine: ImageEngine, task_id: str) -> dict:
}
async def download_image(image_url: str, dest_path: str) -> str:
async def download_image(
image_url: str,
dest_path: str,
*,
execution_guard=None,
) -> str:
os.makedirs(os.path.dirname(dest_path), exist_ok=True)
async with httpx.AsyncClient(timeout=300) as client:
async with client.stream("GET", image_url) as response:
response.raise_for_status()
with open(dest_path, "wb") as file:
chunk_no = 0
async for chunk in response.aiter_bytes(chunk_size=8192):
file.write(chunk)
chunk_no += 1
if execution_guard is not None and chunk_no % 32 == 0:
await execution_guard()
if execution_guard is not None:
await execution_guard()
return dest_path
+5 -1
View File
@@ -102,10 +102,13 @@ async def optimize_prompt(
result = await db.execute(
select(ModelConfig)
.where(ModelConfig.is_active == True)
.where(ModelConfig.is_active == True, ModelConfig.deleted_at.is_(None))
.order_by(ModelConfig.priority.desc())
)
configs = list(result.scalars().all())
# Release the read transaction before the external LLM request. Callers
# must commit their business claim before invoking optimize_prompt.
await db.commit()
if configs:
# 按 priority 从大到小依次尝试,跳过 mock,失败则用下一个
@@ -178,6 +181,7 @@ async def _call_openai_compatible(
break
except (json.JSONDecodeError, TypeError):
pass
await db.commit()
if not system_prompt:
if gen_type == "image":
@@ -85,6 +85,7 @@ async def _find_latest_media_charge(
owner_type: str,
owner_id: str,
media_type: str | None,
attempt_no: int | None = None,
) -> CreditRecord | None:
query = (
select(CreditRecord)
@@ -97,6 +98,8 @@ async def _find_latest_media_charge(
)
if media_type:
query = query.where(CreditRecord.media_type == media_type)
if attempt_no is not None:
query = query.where(CreditRecord.attempt_no == int(attempt_no))
query = query.order_by(CreditRecord.attempt_no.desc().nullslast(), CreditRecord.created_at.desc()).limit(1)
result = await db.execute(query)
return result.scalar_one_or_none()
@@ -213,6 +216,7 @@ async def sync_chat_generation_task_media_token_snapshot(
owner_type=CreditRecordOwnerType.CHAT_GENERATION_TASK.value,
owner_id=task.id,
media_type=gen_type or None,
attempt_no=int(getattr(task, "generation_attempt_no", 1) or 1),
)
return await _sync_charge_snapshot(
db,
@@ -243,6 +247,7 @@ async def sync_generation_record_media_token_snapshot(
owner_type=CreditRecordOwnerType.GENERATION_RECORD.value,
owner_id=record.id,
media_type=gen_type or None,
attempt_no=int(getattr(record, "generation_attempt_no", 1) or 1),
)
return await _sync_charge_snapshot(
db,
@@ -16,6 +16,11 @@ from app.models.module_generation_step import ModuleGenerationStep
from app.models.user import User
from app.enums.module_generation_flow import ModuleGenerationFlowConfig
from app.services.module_generation_step_common_service import build_step_input, build_step_output, utc_now
from app.services.generation.pipeline.db_lock_service import (
apply_short_lock_timeout,
execute_with_lock_timeout,
raise_if_database_lock_busy,
)
from app.services.resource_accounting_service import (
SOURCE_MODEL_CHAT_TASK,
soft_delete_resources_by_source,
@@ -44,8 +49,13 @@ async def get_project_for_user(
if populate_existing:
query = query.execution_options(populate_existing=True)
if for_update:
await apply_short_lock_timeout(db)
query = query.with_for_update()
result = await db.execute(query.limit(1))
try:
result = await db.execute(query.limit(1))
except Exception as exc:
raise_if_database_lock_busy(exc)
raise
project = result.scalar_one_or_none()
if not project:
raise HTTPException(status_code=404, detail=config.project_not_found_message)
@@ -72,8 +82,13 @@ async def get_step_for_user(
if not user.is_admin:
query = query.where(ModuleGenerationStep.user_id == user.id)
if for_update:
await apply_short_lock_timeout(db)
query = query.with_for_update()
result = await db.execute(query.limit(1))
try:
result = await db.execute(query.limit(1))
except Exception as exc:
raise_if_database_lock_busy(exc)
raise
step = result.scalar_one_or_none()
if not step:
raise HTTPException(status_code=404, detail=config.step_not_found_message)
@@ -278,7 +293,9 @@ async def load_chat_tasks_for_steps(
)
if for_update:
stmt = stmt.with_for_update()
result = await db.execute(stmt)
result = await execute_with_lock_timeout(db, stmt)
else:
result = await db.execute(stmt)
return {task.id: task for task in result.scalars().all()}
@@ -314,7 +331,8 @@ async def assert_project_has_no_active_chat_tasks(
config: ModuleGenerationFlowConfig,
detail_message: str = "当前存在生成中任务,请等待生成完成或失败后再操作",
) -> dict[str, ChatGenerationTask]:
result = await db.execute(
result = await execute_with_lock_timeout(
db,
select(ModuleGenerationStep)
.where(
ModuleGenerationStep.project_id == project.id,
@@ -353,7 +371,8 @@ async def soft_delete_steps_from_index(
- 已完成任务只做软删任务与 generated_resources释放容量统计失败任务只软删任务
"""
deleted_at = deleted_at or utc_now()
result = await db.execute(
result = await execute_with_lock_timeout(
db,
select(ModuleGenerationStep)
.where(
ModuleGenerationStep.project_id == project.id,
@@ -49,7 +49,7 @@ def _mask_string(value: str) -> str:
def _is_sensitive_key(key: str) -> bool:
lower = str(key).replace("-", "_").lower()
return any(pattern in lower for pattern in SENSITIVE_KEY_PATTERNS)
return lower == "sign" or any(pattern in lower for pattern in SENSITIVE_KEY_PATTERNS)
def _sanitize_url(value: str) -> str:
@@ -8,6 +8,7 @@ import logging
import os
import threading
import uuid
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Any, Dict, Iterable, List, Optional, Union
@@ -24,6 +25,40 @@ logger = logging.getLogger("video_gen")
_redis_clients: Dict[tuple[int, int, int], Any] = {}
class RedisExecutionLockError(RuntimeError):
"""Redis execution-lock infrastructure error.
Execution locks are fail-closed: callers must stop the current Celery task
and retry later instead of falling back to an unlocked database path.
"""
class RedisExecutionLockUnavailable(RedisExecutionLockError):
"""Redis is unavailable, so execution ownership cannot be established."""
class RedisExecutionLockLost(RedisExecutionLockError):
"""The current worker no longer owns the execution lock."""
_RELEASE_LOCK_SCRIPT = """
if redis.call('get', KEYS[1]) == ARGV[1] then
return redis.call('del', KEYS[1])
else
return 0
end
"""
_RENEW_LOCK_SCRIPT = """
if redis.call('get', KEYS[1]) == ARGV[1] then
return redis.call('pexpire', KEYS[1], ARGV[2])
else
return 0
end
"""
def utc_now() -> datetime:
return datetime.now(timezone.utc)
@@ -73,8 +108,9 @@ async def get_registry_redis() -> Optional[Any]:
Redis 客户端会触发 got Future attached to a different loop
因此这里按 pid + thread_id + event_loop_id 缓存客户端确保同一个客户端
只在创建它的事件循环里使用Redis 不可用时返回 None调用方降级为
DB fallback不能影响生成主链路
只在创建它的事件循环里使用普通 active 注册表在 Redis 不可用时返回
None执行锁封装会把 None 转为 RedisExecutionLockUnavailable严格中止
当前任务不允许无锁执行外部副作用
"""
redis_url = registry_redis_url()
if not _is_supported_redis_url(redis_url):
@@ -328,6 +364,212 @@ async def redis_acquire_lock(
return None
async def redis_acquire_execution_lock(
*,
lock_key: str,
ttl_seconds: int,
token: Optional[str] = None,
log_context: str = "execution_lock",
) -> Optional[str]:
"""Fail-closed execution lock.
Returns a token when acquired and ``None`` when another worker owns the
lock. Redis connection/command failures raise
:class:`RedisExecutionLockUnavailable`; callers must retry the Celery task
and must not execute the external side effect without a lock.
"""
redis = await get_registry_redis()
if redis is None:
raise RedisExecutionLockUnavailable(
f"Redis execution lock unavailable: context={log_context}, key={lock_key}"
)
lock_token = token or uuid.uuid4().hex
ttl_ms = max(1000, int(ttl_seconds or 60) * 1000)
try:
acquired = await redis.set(lock_key, lock_token, nx=True, px=ttl_ms)
return lock_token if acquired else None
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.error(
"获取 Redis 执行锁失败。context=%s, lock_key=%s, error=%s",
log_context,
lock_key,
exc,
)
raise RedisExecutionLockUnavailable(
f"Redis execution lock acquire failed: {lock_key}: {exc}"
) from exc
async def redis_check_lock_owner(
*,
lock_key: str,
token: str,
log_context: str = "execution_lock",
) -> bool:
redis = await get_registry_redis()
if redis is None:
raise RedisExecutionLockUnavailable(
f"Redis execution lock unavailable: context={log_context}, key={lock_key}"
)
try:
value = await redis.get(lock_key)
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.error(
"检查 Redis 执行锁失败。context=%s, lock_key=%s, error=%s",
log_context,
lock_key,
exc,
)
raise RedisExecutionLockUnavailable(
f"Redis execution lock check failed: {lock_key}: {exc}"
) from exc
return bool(value and str(value) == str(token))
async def redis_renew_lock(
*,
lock_key: str,
token: str,
ttl_seconds: int,
log_context: str = "execution_lock",
) -> bool:
redis = await get_registry_redis()
if redis is None:
raise RedisExecutionLockUnavailable(
f"Redis execution lock unavailable: context={log_context}, key={lock_key}"
)
ttl_ms = max(1000, int(ttl_seconds or 60) * 1000)
try:
renewed = await redis.eval(_RENEW_LOCK_SCRIPT, 1, lock_key, token, ttl_ms)
return bool(renewed)
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.error(
"续期 Redis 执行锁失败。context=%s, lock_key=%s, error=%s",
log_context,
lock_key,
exc,
)
raise RedisExecutionLockUnavailable(
f"Redis execution lock renew failed: {lock_key}: {exc}"
) from exc
@dataclass(slots=True)
class RedisExecutionLockLease:
"""Owned Redis execution lock with compare-and-expire heartbeat."""
lock_key: str
token: str
ttl_seconds: int
log_context: str = "execution_lock"
renew_interval_seconds: int | None = None
_stop_event: asyncio.Event = field(default_factory=asyncio.Event, init=False, repr=False)
_heartbeat_task: asyncio.Task[Any] | None = field(default=None, init=False, repr=False)
_lost_error: RedisExecutionLockError | None = field(default=None, init=False, repr=False)
@classmethod
async def acquire(
cls,
*,
lock_key: str,
ttl_seconds: int,
token: str | None = None,
log_context: str = "execution_lock",
renew_interval_seconds: int | None = None,
) -> "RedisExecutionLockLease | None":
acquired_token = await redis_acquire_execution_lock(
lock_key=lock_key,
ttl_seconds=ttl_seconds,
token=token,
log_context=log_context,
)
if not acquired_token:
return None
lease = cls(
lock_key=lock_key,
token=acquired_token,
ttl_seconds=max(1, int(ttl_seconds or 60)),
log_context=log_context,
renew_interval_seconds=renew_interval_seconds,
)
lease.start_heartbeat()
return lease
def start_heartbeat(self) -> None:
if self._heartbeat_task is not None:
return
self._heartbeat_task = asyncio.create_task(self._heartbeat())
async def _heartbeat(self) -> None:
interval = int(
self.renew_interval_seconds
or max(1, min(60, self.ttl_seconds // 3))
)
while not self._stop_event.is_set():
try:
await asyncio.wait_for(self._stop_event.wait(), timeout=interval)
return
except asyncio.TimeoutError:
pass
try:
renewed = await redis_renew_lock(
lock_key=self.lock_key,
token=self.token,
ttl_seconds=self.ttl_seconds,
log_context=self.log_context,
)
if not renewed:
self._lost_error = RedisExecutionLockLost(
f"Redis execution lock ownership lost: {self.lock_key}"
)
return
except RedisExecutionLockError as exc:
self._lost_error = exc
return
async def ensure_owned(self) -> None:
if self._lost_error is not None:
raise self._lost_error
owned = await redis_check_lock_owner(
lock_key=self.lock_key,
token=self.token,
log_context=self.log_context,
)
if not owned:
self._lost_error = RedisExecutionLockLost(
f"Redis execution lock ownership lost: {self.lock_key}"
)
raise self._lost_error
async def close(self) -> None:
self._stop_event.set()
heartbeat = self._heartbeat_task
if heartbeat is not None:
try:
await heartbeat
except Exception:
logger.debug(
"Redis execution lock heartbeat close failed. context=%s key=%s",
self.log_context,
self.lock_key,
exc_info=True,
)
try:
await redis_release_lock(
lock_key=self.lock_key,
token=self.token,
log_context=self.log_context,
)
except Exception:
logger.debug(
"Redis execution lock release failed. context=%s key=%s",
self.log_context,
self.lock_key,
exc_info=True,
)
async def redis_release_lock(
*,
lock_key: str,
@@ -338,16 +580,8 @@ async def redis_release_lock(
redis = await get_registry_redis()
if redis is None:
return False
script = """
if redis.call('get', KEYS[1]) == ARGV[1] then
return redis.call('del', KEYS[1])
else
return 0
end
"""
try:
released = await redis.eval(script, 1, lock_key, token)
released = await redis.eval(_RELEASE_LOCK_SCRIPT, 1, lock_key, token)
return bool(released)
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning("释放 Redis 锁失败。context=%s, lock_key=%s, error=%s", log_context, lock_key, exc)
@@ -317,20 +317,29 @@ async def record_generation_record_generated_resource(
remote_url: str | None = None,
generated_at: datetime | None = None,
) -> ResourceAccountingResult:
snapshot = _parse_json(getattr(record, "engine_snapshot_json", None))
return await record_generated_resource(
db,
user_id=record.user_id,
resource_type=record.gen_type,
resource_url=resource_url,
remote_url=remote_url,
remote_url=remote_url or getattr(record, "remote_result_url", None),
storage_type="local" if storage_path else "remote",
storage_path=storage_path,
file_size_bytes=file_size_bytes,
source_model=SOURCE_MODEL_GENERATION_RECORD,
source_model_module="app.models.generation_record",
source_id=record.id,
engine_id=getattr(record, "engine_id", None),
engine_type=snapshot.get("engine_type") or record.gen_type,
provider=snapshot.get("provider"),
model_name=snapshot.get("model_name"),
generated_at=generated_at or record.generated_at or datetime.now(timezone.utc),
extra={"project_id": record.project_id},
extra={
"project_id": record.project_id,
"generation_attempt_no": int(getattr(record, "generation_attempt_no", 1) or 1),
"pipeline_stage": record.pipeline_stage,
},
)
@@ -1,5 +1,6 @@
from __future__ import annotations
import asyncio
import json
from datetime import datetime
from typing import Any
@@ -42,6 +43,11 @@ from app.services.generation.ai.engine_service import (
)
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.pipeline.db_lock_service import (
DatabaseRowLockBusy,
apply_short_lock_timeout,
execute_with_lock_timeout,
)
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,
@@ -338,6 +344,18 @@ async def _soft_delete_steps_from_index(
refund_unfinished: bool = False,
release_stats: dict[str, int] | None = None,
) -> None:
processing_result = await db.execute(
select(ModuleGenerationStep.id).where(
ModuleGenerationStep.project_id == project.id,
ModuleGenerationStep.module == MODULE,
ModuleGenerationStep.deleted_at.is_(None),
ModuleGenerationStep.is_current == True,
ModuleGenerationStep.status == ModuleStepStatusEnum.PROCESSING.value,
ModuleGenerationStep.step_index >= start_index,
).limit(1)
)
if processing_result.scalar_one_or_none() is not None:
raise HTTPException(status_code=409, detail="当前步骤正在处理中,请等待完成后再操作")
await _base_soft_delete_steps_from_index(
db,
project=project,
@@ -693,6 +711,69 @@ async def update_shot_replicate_video_prompt_schema(
)
async def _reload_prompt_context_for_update(
db: AsyncSession,
*,
project_id: str,
step_id: str,
step_code: str,
) -> tuple[ModuleGenerationProject | None, ModuleGenerationStep | None]:
last_error: DatabaseRowLockBusy | None = None
for retry_index in range(3):
try:
project_result = await execute_with_lock_timeout(
db,
select(ModuleGenerationProject)
.where(
ModuleGenerationProject.id == project_id,
ModuleGenerationProject.module == MODULE,
ModuleGenerationProject.deleted_at.is_(None),
)
.with_for_update()
.execution_options(populate_existing=True)
.limit(1)
)
project = project_result.scalar_one_or_none()
if project is None:
return None, None
step_result = await execute_with_lock_timeout(
db,
select(ModuleGenerationStep)
.where(
ModuleGenerationStep.id == step_id,
ModuleGenerationStep.project_id == project_id,
ModuleGenerationStep.module == MODULE,
ModuleGenerationStep.step_code == step_code,
ModuleGenerationStep.deleted_at.is_(None),
ModuleGenerationStep.is_current == True,
)
.with_for_update()
.execution_options(populate_existing=True)
.limit(1)
)
return project, step_result.scalar_one_or_none()
except DatabaseRowLockBusy as exc:
last_error = exc
await db.rollback()
if retry_index < 2:
await asyncio.sleep(1 + retry_index)
raise last_error or DatabaseRowLockBusy()
def _prompt_context_matches(
step: ModuleGenerationStep | None,
*,
expected_version: int,
expected_input_json: str,
) -> bool:
if step is None or step.status != ModuleStepStatusEnum.PROCESSING.value:
return False
if int(step.version or 1) != int(expected_version):
return False
current_input = json.dumps(step.input_json, ensure_ascii=False, sort_keys=True, default=str)
return current_input == expected_input_json
async def submit_image_prompt_optimize(
db: AsyncSession,
*,
@@ -732,6 +813,7 @@ async def submit_image_prompt_optimize(
async def run_image_prompt_optimize(db: AsyncSession, *, project_id: str, step_id: str | None = None) -> ModuleGenerationStep | None:
await apply_short_lock_timeout(db)
project_result = await db.execute(
select(ModuleGenerationProject)
.where(ModuleGenerationProject.id == project_id, ModuleGenerationProject.module == MODULE, ModuleGenerationProject.deleted_at.is_(None))
@@ -749,6 +831,7 @@ async def run_image_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
return None
if step_id:
await apply_short_lock_timeout(db)
result = await db.execute(
select(ModuleGenerationStep)
.where(
@@ -795,24 +878,45 @@ async def run_image_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
{"type": "video", "url": material.get("material_video_url"), "name": "参考素材视频"},
{"type": "image", "url": material.get("material_image_url"), "name": "新产品图片"},
]
project_id_value = str(project.id)
step_id_value = str(step.id)
user_id_value = str(project.user_id)
module_value = str(project.module)
expected_step_version = int(step.version or 1)
expected_input_json = json.dumps(step.input_json, ensure_ascii=False, sort_keys=True, default=str)
await db.commit()
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,
project_id=project_id_value,
step_id=step_id_value,
user_id=user_id_value,
module=module_value,
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,
user_id=user_id_value,
references=references,
gen_type="image",
)
project, step = await _reload_prompt_context_for_update(
db,
project_id=project_id_value,
step_id=step_id_value,
step_code=ShotReplicateStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value,
)
if not _prompt_context_matches(
step,
expected_version=expected_step_version,
expected_input_json=expected_input_json,
):
await db.rollback()
return None
billing = await charge_module_prompt_usage(
db,
user_id=project.user_id,
@@ -857,7 +961,25 @@ async def run_image_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
token_usage=usage,
)
await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.IMAGE_PROMPT_SUCCESS.value, message="图片 AI 提词生成成功")
await db.commit()
except DatabaseRowLockBusy:
await db.rollback()
raise
except Exception as exc:
await db.rollback()
project, step = await _reload_prompt_context_for_update(
db,
project_id=project_id_value,
step_id=step_id_value,
step_code=ShotReplicateStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value,
)
if not _prompt_context_matches(
step,
expected_version=expected_step_version,
expected_input_json=expected_input_json,
):
await db.rollback()
return None
step.status = ModuleStepStatusEnum.FAILED.value
step.error_message = str(exc)
step.completed_at = _now()
@@ -875,6 +997,7 @@ async def run_image_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
)
_log_project_error(project=project, step=step, event_type="IMAGE_PROMPT_FAILED", message=project.error_message, exc=exc)
await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.IMAGE_PROMPT_FAILED.value, message=project.error_message)
await db.commit()
return step
@@ -1053,6 +1176,7 @@ async def submit_video_prompt_optimize(
async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_id: str | None = None) -> ModuleGenerationStep | None:
await apply_short_lock_timeout(db)
project_result = await db.execute(
select(ModuleGenerationProject)
.where(ModuleGenerationProject.id == project_id, ModuleGenerationProject.module == MODULE, ModuleGenerationProject.deleted_at.is_(None))
@@ -1072,6 +1196,7 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
return None
if step_id:
await apply_short_lock_timeout(db)
result = await db.execute(
select(ModuleGenerationStep)
.where(
@@ -1111,8 +1236,17 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
step.error_message = "缺少新项目图片结果,不能生成视频提词"
project.status = ModuleProjectStatusEnum.FAILED.value
project.error_message = step.error_message
await db.commit()
return step
project_id_value = str(project.id)
step_id_value = str(step.id)
user_id_value = str(project.user_id)
module_value = str(project.module)
expected_step_version = int(step.version or 1)
expected_input_json = json.dumps(step.input_json, ensure_ascii=False, sort_keys=True, default=str)
await db.commit()
try:
request_log = {
"source_project_name": material.get("source_project_name") or "",
@@ -1125,10 +1259,10 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
}
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,
project_id=project_id_value,
step_id=step_id_value,
user_id=user_id_value,
module=module_value,
prompt_type=ModulePromptTypeEnum.VIDEO_PROMPT.value,
request=request_log,
)
@@ -1136,7 +1270,7 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
request_log["schema_config_source"] = schema_config_snapshot.get("source")
prompt_schema, final_prompt, token_usage = await optimize_shot_replicate_video_prompt(
db,
user_id=project.user_id,
user_id=user_id_value,
source_project_name=request_log["source_project_name"],
target_project_name=request_log["target_project_name"],
core_content_point=request_log["core_content_point"],
@@ -1145,11 +1279,24 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
video_config=video_config,
target_platform=target_platform,
schema_config_snapshot=schema_config_snapshot,
module=project.module,
project_id=project.id,
step_id=step.id,
trace_id=f"shot-video-prompt:{step.id}",
module=module_value,
project_id=project_id_value,
step_id=step_id_value,
trace_id=f"shot-video-prompt:{step_id_value}",
)
project, step = await _reload_prompt_context_for_update(
db,
project_id=project_id_value,
step_id=step_id_value,
step_code=ShotReplicateStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value,
)
if not _prompt_context_matches(
step,
expected_version=expected_step_version,
expected_input_json=expected_input_json,
):
await db.rollback()
return None
billing = await charge_module_prompt_usage(
db,
user_id=project.user_id,
@@ -1197,7 +1344,25 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
token_usage=usage,
)
await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.VIDEO_PROMPT_SUCCESS.value, message="视频 AI 提词生成成功")
await db.commit()
except DatabaseRowLockBusy:
await db.rollback()
raise
except Exception as exc:
await db.rollback()
project, step = await _reload_prompt_context_for_update(
db,
project_id=project_id_value,
step_id=step_id_value,
step_code=ShotReplicateStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value,
)
if not _prompt_context_matches(
step,
expected_version=expected_step_version,
expected_input_json=expected_input_json,
):
await db.rollback()
return None
step.status = ModuleStepStatusEnum.FAILED.value
step.error_message = str(exc)
step.completed_at = _now()
@@ -1215,6 +1380,7 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
)
_log_project_error(project=project, step=step, event_type="VIDEO_PROMPT_FAILED", message=project.error_message, exc=exc)
await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.VIDEO_PROMPT_FAILED.value, message=project.error_message)
await db.commit()
return step
@@ -1331,9 +1497,37 @@ async def generate_video_from_prompt(
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(
meta_result = await db.execute(
select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where(
ModuleGenerationStep.chat_task_id == task.id,
ModuleGenerationStep.module == MODULE,
ModuleGenerationStep.is_current == True,
ModuleGenerationStep.deleted_at.is_(None),
).limit(1)
)
meta = meta_result.first()
if not meta:
return
step_id_value, project_id_value = str(meta.id), str(meta.project_id)
project_result = await execute_with_lock_timeout(
db,
select(ModuleGenerationProject)
.where(
ModuleGenerationProject.id == project_id_value,
ModuleGenerationProject.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
project = project_result.scalar_one_or_none()
if not project:
return
step_result = await execute_with_lock_timeout(
db,
select(ModuleGenerationStep)
.where(
ModuleGenerationStep.id == step_id_value,
ModuleGenerationStep.project_id == project_id_value,
ModuleGenerationStep.chat_task_id == task.id,
ModuleGenerationStep.module == MODULE,
ModuleGenerationStep.is_current == True,
@@ -1342,18 +1536,9 @@ async def handle_chat_generation_task_completed(db: AsyncSession, task: ChatGene
.with_for_update()
.limit(1)
)
step = result.scalar_one_or_none()
step = 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 == ShotReplicateStepCodeEnum.IMAGE_GENERATE.value:
step.status = ModuleStepStatusEnum.COMPLETED.value
@@ -1394,19 +1579,48 @@ async def handle_chat_generation_task_completed(db: AsyncSession, task: ChatGene
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))
meta_result = await db.execute(
select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where(
ModuleGenerationStep.chat_task_id == task.id,
ModuleGenerationStep.module == MODULE,
ModuleGenerationStep.is_current == True,
ModuleGenerationStep.deleted_at.is_(None),
).limit(1)
)
meta = meta_result.first()
if not meta:
return
step_id_value, project_id_value = str(meta.id), str(meta.project_id)
project_result = await execute_with_lock_timeout(
db,
select(ModuleGenerationProject)
.where(
ModuleGenerationProject.id == project_id_value,
ModuleGenerationProject.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_result = await execute_with_lock_timeout(
db,
select(ModuleGenerationStep)
.where(
ModuleGenerationStep.id == step_id_value,
ModuleGenerationStep.project_id == project_id_value,
ModuleGenerationStep.chat_task_id == task.id,
ModuleGenerationStep.module == MODULE,
ModuleGenerationStep.is_current == True,
ModuleGenerationStep.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
step = step_result.scalar_one_or_none()
if not step:
return
step.status = ModuleStepStatusEnum.FAILED.value
step.error_message = task.error_message
step.completed_at = _now()
@@ -1439,7 +1653,8 @@ async def mark_shot_replicate_step_dispatch_failed(
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(
result = await execute_with_lock_timeout(
db,
select(ModuleGenerationStep)
.where(
ModuleGenerationStep.id == step_id,
@@ -414,7 +414,7 @@ def filter_and_normalize_breakdown(result: dict[str, Any], *, mode: AnalysisMode
async def _select_model_config(db: AsyncSession) -> ModelConfig | None:
result = await db.execute(
select(ModelConfig)
.where(ModelConfig.is_active == True)
.where(ModelConfig.is_active == True, ModelConfig.deleted_at.is_(None))
.order_by(ModelConfig.priority.desc())
.limit(1)
)
+13 -2
View File
@@ -77,7 +77,7 @@ async def get_active_engine(db: AsyncSession) -> VideoEngine:
"""Get the active video engine with highest priority."""
result = await db.execute(
select(VideoEngine)
.where(VideoEngine.is_active == True)
.where(VideoEngine.is_active == True, VideoEngine.deleted_at.is_(None))
.order_by(VideoEngine.priority.desc())
.limit(1)
)
@@ -221,7 +221,12 @@ async def poll_task_status(engine: VideoEngine, task_id: str) -> dict:
}
async def download_video(video_url: str, dest_path: str) -> str:
async def download_video(
video_url: str,
dest_path: str,
*,
execution_guard=None,
) -> str:
"""Download video to local storage."""
import os
@@ -231,6 +236,12 @@ async def download_video(video_url: str, dest_path: str) -> str:
async with client.stream("GET", video_url) as response:
response.raise_for_status()
with open(dest_path, "wb") as f:
chunk_no = 0
async for chunk in response.aiter_bytes(chunk_size=8192):
f.write(chunk)
chunk_no += 1
if execution_guard is not None and chunk_no % 32 == 0:
await execution_guard()
if execution_guard is not None:
await execution_guard()
return dest_path
-433
View File
@@ -1,433 +0,0 @@
import asyncio
import json
import logging
import os
from datetime import datetime, timezone
from urllib.parse import urlparse
from sqlalchemy import or_, select
from app.config import settings
from app.enums.generation_status import GenerationRecordPipelineStage
from app.models.base import async_session
from app.models.generation_record import GenerationRecord
from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once
from app.services.image_gen import download_image, get_active_image_engine
from app.services.media_token_usage_snapshot_service import sync_generation_record_media_token_snapshot
from app.services.provider_limit import provider_limit
from app.services.resource_accounting_service import (
record_generation_record_generated_resource,
safe_file_size,
)
from app.services.video_cover_service import create_video_cover_for_local_video
from app.services.video_gen import _log_video_response, download_video, get_active_engine, poll_task_status
from app.services.video_upscale.media_service import build_part_mp4_path, probe_video, safe_remove
logger = logging.getLogger("videogen")
POLL_INTERVAL = 30
MAX_POLLS = 60
_PROVIDER_RECOVERABLE_STAGES = {
None,
"",
GenerationRecordPipelineStage.CREATING_PROVIDER_TASK.value,
GenerationRecordPipelineStage.WAITING_REMOTE.value,
GenerationRecordPipelineStage.POLLING.value,
GenerationRecordPipelineStage.RESULT_READY.value,
GenerationRecordPipelineStage.DOWNLOADING.value,
}
def _is_provider_stage(record: GenerationRecord) -> bool:
return (record.pipeline_stage or "") in _PROVIDER_RECOVERABLE_STAGES
def _source_date_dir(record: GenerationRecord) -> str:
created = record.created_at
if created is None:
created = datetime.now(timezone.utc)
return created.strftime("%Y/%m/%d")
def _normalize_image_extension(output_format: str | None, remote_url: str | None = None) -> str:
value = str(output_format or "").strip().lower()
if value in {"jpg", "jpeg"}:
return "jpg"
if value == "png":
return "png"
if value == "webp":
return "webp"
if remote_url:
try:
path = urlparse(remote_url).path or ""
except Exception:
path = str(remote_url)
suffix = os.path.splitext(path)[1].lower().lstrip(".")
if suffix in {"jpg", "jpeg"}:
return "jpg"
if suffix == "png":
return "png"
if suffix == "webp":
return "webp"
return "jpg"
async def _download_generation_record_upscale_source(
record: GenerationRecord,
remote_url: str,
) -> tuple[str, int]:
date_dir = _source_date_dir(record)
dest_dir = os.path.join(settings.STORAGE_LOCAL_PATH, "_upscale_source", date_dir)
os.makedirs(dest_dir, exist_ok=True)
final_path = os.path.join(dest_dir, f"{record.id}.source.mp4")
if os.path.isfile(final_path) and os.path.getsize(final_path) > 0:
await probe_video(final_path)
return final_path, safe_file_size(final_path)
part_path = build_part_mp4_path(final_path)
try:
async with provider_limit("result_download", settings.RESULT_DOWNLOAD_MAX_CONCURRENCY):
await download_video(remote_url, part_path)
if not os.path.isfile(part_path) or os.path.getsize(part_path) <= 0:
raise RuntimeError("超分源视频下载完成但临时文件为空")
await probe_video(part_path)
os.replace(part_path, final_path)
return final_path, safe_file_size(final_path)
except Exception:
safe_remove(part_path)
raise
async def handle_generation_record_video_succeeded(
db,
record: GenerationRecord,
*,
remote_url: str,
provider_response: dict | None,
video_tokens: int = 0,
) -> bool:
"""处理 GenerationRecord 原视频生成成功。
返回 True 表示已进入超分队列False 表示按原流程直接完成
"""
record.video_tokens_used = int(video_tokens or 0)
await sync_generation_record_media_token_snapshot(db, record, provider_response=provider_response or {})
if bool(record.video_upscale_enabled_snapshot) and record.video_upscale_snapshot_json:
record.pipeline_stage = GenerationRecordPipelineStage.DOWNLOADING.value
await db.flush()
source_path, source_size = await _download_generation_record_upscale_source(record, remote_url)
from app.services.video_upscale.task_service import enqueue_upscale_task, prepare_video_upscale_task
upscale = await prepare_video_upscale_task(
db,
generation_record=record,
source_local_path=source_path,
source_file_size_bytes=source_size,
source_remote_url=remote_url,
)
# enqueue_upscale_task 会先提交数据库,再投递 Celery;投递失败由 gen_recovery 补投。
await enqueue_upscale_task(db, upscale=upscale, reason="generation_record_source_ready")
logger.info("GenerationRecord 已进入视频超分队列: record_id=%s upscale_task_id=%s", record.id, upscale.id)
return True
storage_path = None
file_size_bytes = 0
if settings.STORAGE_TYPE == "local" and remote_url:
try:
date_dir = _source_date_dir(record)
dest_dir = os.path.join(settings.STORAGE_LOCAL_PATH, date_dir)
os.makedirs(dest_dir, exist_ok=True)
dest = os.path.join(dest_dir, f"{record.id}.mp4")
await download_video(remote_url, dest)
record.video_url = f"/generate/videos/{date_dir}/{record.id}.mp4"
cover_url, _cover_storage_path = create_video_cover_for_local_video(
record_id=record.id,
video_path=dest,
date_dir=date_dir,
log_prefix=f"GenerationRecord视频封面生成 record_id={record.id}",
)
record.video_cover_url = cover_url
storage_path = dest
file_size_bytes = safe_file_size(dest)
except Exception as exc:
logger.warning("GenerationRecord 最终视频本地保存失败,回退远程地址: record_id=%s error=%s", record.id, exc)
record.video_url = remote_url
else:
record.video_url = remote_url
record.status = "completed"
record.pipeline_stage = GenerationRecordPipelineStage.DONE.value
record.generated_at = datetime.now(timezone.utc)
record.error_message = None
if record.video_url:
await record_generation_record_generated_resource(
db,
record,
resource_url=record.video_url,
storage_path=storage_path,
file_size_bytes=file_size_bytes,
remote_url=remote_url,
generated_at=record.generated_at,
)
await db.commit()
return False
class TaskQueue:
def __init__(self):
self.queue: asyncio.Queue[str] = asyncio.Queue()
self.running = False
self._active: dict[str, int] = {}
async def enqueue(self, record_id: str):
await self.queue.put(record_id)
async def recover(self):
async with async_session() as db:
result = await db.execute(
select(GenerationRecord).where(
GenerationRecord.status == "generating",
GenerationRecord.seedance_task_id.isnot(None),
GenerationRecord.deleted_at.is_(None),
or_(
GenerationRecord.pipeline_stage.is_(None),
GenerationRecord.pipeline_stage.in_(
[stage for stage in _PROVIDER_RECOVERABLE_STAGES if stage]
),
),
)
)
records = result.scalars().all()
for record in records:
await self.queue.put(record.id)
logger.info("Recovered task: %s (seedance: %s stage=%s)", record.id, record.seedance_task_id, record.pipeline_stage)
async def run(self):
self.running = True
logger.info("Video queue started")
while self.running:
try:
record_id = await asyncio.wait_for(self.queue.get(), timeout=5.0)
except asyncio.TimeoutError:
continue
try:
await self._process(record_id)
except Exception as exc:
logger.exception("Error processing %s: %s", record_id, exc)
finally:
self.queue.task_done()
logger.info("Video queue stopped")
async def _process(self, record_id: str):
async with async_session() as db:
result = await db.execute(
select(GenerationRecord).where(
GenerationRecord.id == record_id,
GenerationRecord.deleted_at.is_(None),
).with_for_update().limit(1)
)
record = result.scalar_one_or_none()
if not record or record.status != "generating":
return
if record.gen_type == "video":
if not _is_provider_stage(record):
return
await self._process_video(db, record)
else:
await self._process_image(db, record)
async def _process_video(self, db, record: GenerationRecord):
record_id = record.id
if not record.seedance_task_id:
record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
await mark_generation_record_failed_and_refund_once(db, record=record, error_message="缺少外部任务ID")
await db.commit()
return
record.pipeline_stage = GenerationRecordPipelineStage.POLLING.value
try:
engine = await get_active_engine(db)
poll_result = await poll_task_status(engine, record.seedance_task_id)
except Exception as exc:
logger.error("Poll error for %s: %s", record_id, exc)
count = self._active.get(record_id, 0) + 1
self._active[record_id] = count
if count >= MAX_POLLS:
record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
await mark_generation_record_failed_and_refund_once(db, record=record, error_message=f"轮询超时: {exc}")
await db.commit()
self._active.pop(record_id, None)
else:
await db.commit()
await asyncio.sleep(POLL_INTERVAL)
await self.queue.put(record_id)
return
status = poll_result["status"]
try:
resp_data = json.loads(poll_result.get("response_data", "{}"))
except (json.JSONDecodeError, TypeError):
resp_data = {}
_log_video_response(record_id, resp_data, poll_result.get("error"))
if status == "succeeded":
file_url = str(poll_result.get("video_url") or "").strip()
if not file_url:
record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
await mark_generation_record_failed_and_refund_once(db, record=record, error_message="供应商成功但未返回视频地址")
await db.commit()
return
record.pipeline_stage = GenerationRecordPipelineStage.RESULT_READY.value
await db.flush()
try:
await handle_generation_record_video_succeeded(
db,
record,
remote_url=file_url,
provider_response=resp_data,
video_tokens=poll_result.get("video_tokens", 0),
)
except Exception as exc:
await db.rollback()
result = await db.execute(
select(GenerationRecord).where(GenerationRecord.id == record_id).with_for_update().limit(1)
)
failed_record = result.scalar_one_or_none()
if failed_record:
failed_record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
await mark_generation_record_failed_and_refund_once(
db,
record=failed_record,
error_message=f"视频结果下载失败: {exc}",
)
await db.commit()
logger.exception("GenerationRecord 视频成功结果处理失败: %s", record_id)
self._active.pop(record_id, None)
return
if status == "failed":
record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=poll_result.get("error", "视频生成失败"),
)
self._active.pop(record_id, None)
await db.commit()
logger.info("Video task failed: %s", record_id)
return
count = self._active.get(record_id, 0) + 1
self._active[record_id] = count
if count >= MAX_POLLS:
record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
await mark_generation_record_failed_and_refund_once(db, record=record, error_message="视频生成超时")
self._active.pop(record_id, None)
await db.commit()
logger.info("Video task timed out: %s", record_id)
else:
await db.commit()
await asyncio.sleep(POLL_INTERVAL)
await self.queue.put(record_id)
async def _process_image(self, db, record: GenerationRecord):
"""Process image generation task - calls API directly."""
record_id = record.id
from app.services.image_gen import submit_image_task
try:
engine = await get_active_image_engine(db)
poll_result = await asyncio.to_thread(
submit_image_task,
db,
engine,
record,
include_media_references=False,
)
if not isinstance(poll_result, dict):
raise RuntimeError("图片供应商返回结构异常")
items = poll_result.get("items") or []
if not isinstance(items, list):
raise RuntimeError("图片供应商返回结果列表异常")
if not items:
raise RuntimeError("图片供应商未返回图片结果")
if len(items) != 1:
raise RuntimeError(f"图片供应商单图返回数量异常,期望 1,实际 {len(items)}")
item = items[0] or {}
if not isinstance(item, dict):
raise RuntimeError("图片供应商返回单项结果结构异常")
item_error = item.get("error_message") or item.get("error_code")
if item_error:
raise RuntimeError(str(item_error))
remote_url = str(item.get("remote_result_url") or "").strip()
if not remote_url:
raise RuntimeError("图片供应商成功响应但没有图片地址")
storage_path = None
file_size_bytes = 0
if settings.STORAGE_TYPE == "local":
try:
date_dir = _source_date_dir(record)
dest_dir = os.path.join(settings.STORAGE_IMAGE_LOCAL_PATH, date_dir)
os.makedirs(dest_dir, exist_ok=True)
extension = _normalize_image_extension(item.get("output_format"), remote_url)
dest = os.path.join(dest_dir, f"{record_id}.{extension}")
await download_image(remote_url, dest)
record.image_url = f"/generate/images/{date_dir}/{record_id}.{extension}"
storage_path = dest
file_size_bytes = safe_file_size(dest)
except Exception as exc:
logger.warning("GenerationRecord 图片本地保存失败,回退远程地址: record_id=%s error=%s", record_id, exc)
record.image_url = remote_url
else:
record.image_url = remote_url
record.image_tokens_used = int(poll_result.get("image_tokens", 0) or 0)
provider_response = poll_result.get("response_data") or {}
await sync_generation_record_media_token_snapshot(
db,
record,
provider_response=provider_response if isinstance(provider_response, dict) else {},
)
record.status = "completed"
record.pipeline_stage = GenerationRecordPipelineStage.DONE.value
record.generated_at = datetime.now(timezone.utc)
record.error_message = None
if record.image_url:
await record_generation_record_generated_resource(
db,
record,
resource_url=record.image_url,
storage_path=storage_path,
file_size_bytes=file_size_bytes,
remote_url=remote_url,
generated_at=record.generated_at,
)
await db.commit()
logger.info("Image task completed: %s", record_id)
except Exception as exc:
record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=(getattr(exc, "safe_message", None) or str(exc) or "图片生成失败"),
)
await db.commit()
logger.error("Image task failed: %s, error: %s", record_id, exc, exc_info=True)
def stop(self):
"""Signal the queue to stop."""
self.running = False
task_queue = TaskQueue()
@@ -2,12 +2,16 @@ from __future__ import annotations
import asyncio
import os
import subprocess
from pathlib import Path
from collections.abc import Awaitable, Callable
from app.config import settings
from app.services.video_cover_service import get_ffmpeg_bin
from app.services.video_upscale.media_service import build_part_mp4_path, is_valid_file, probe_video, safe_remove
from app.services.video_upscale.media_service import (
build_part_mp4_path,
is_valid_file,
probe_video,
safe_remove,
)
class LocalVideoUpscaleError(RuntimeError):
@@ -28,17 +32,15 @@ def _build_filter(target_width: int, target_height: int) -> str:
)
def _run_ffmpeg_sync(
def _build_command(
*,
source_path: str,
part_path: str,
target_width: int,
target_height: int,
timeout_seconds: int,
) -> None:
ffmpeg_bin = get_ffmpeg_bin() # 明确复用 config.py 的 FFMPEG_BIN。
cmd = [
ffmpeg_bin,
) -> list[str]:
return [
get_ffmpeg_bin(),
"-hide_banner",
"-nostdin",
"-y",
@@ -68,23 +70,69 @@ def _run_ffmpeg_sync(
"192k",
part_path,
]
async def _terminate_process(process: asyncio.subprocess.Process) -> None:
if process.returncode is not None:
return
process.terminate()
try:
result = subprocess.run(
cmd,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
encoding="utf-8",
errors="replace",
timeout=max(30, int(timeout_seconds)),
shell=False,
await asyncio.wait_for(process.wait(), timeout=5)
except asyncio.TimeoutError:
process.kill()
await process.wait()
async def _run_ffmpeg(
*,
source_path: str,
part_path: str,
target_width: int,
target_height: int,
timeout_seconds: int,
execution_guard: Callable[[], Awaitable[None]] | None,
) -> None:
try:
process = await asyncio.create_subprocess_exec(
*_build_command(
source_path=source_path,
part_path=part_path,
target_width=target_width,
target_height=target_height,
),
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
)
except subprocess.TimeoutExpired as exc:
raise LocalVideoUpscaleError(f"本地 FFmpeg 超分超时: {timeout_seconds}") from exc
except OSError as exc:
raise LocalVideoUpscaleError(f"本地 FFmpeg 无法启动: {exc}") from exc
if result.returncode != 0:
raise LocalVideoUpscaleError(f"本地 FFmpeg 超分失败: {(result.stderr or '').strip()[-4000:]}")
communicate_task = asyncio.create_task(process.communicate())
deadline = asyncio.get_running_loop().time() + max(30, int(timeout_seconds))
try:
while not communicate_task.done():
if execution_guard is not None:
await execution_guard()
remaining = deadline - asyncio.get_running_loop().time()
if remaining <= 0:
raise LocalVideoUpscaleError(f"本地 FFmpeg 超分超时: {timeout_seconds}")
try:
await asyncio.wait_for(asyncio.shield(communicate_task), timeout=min(2.0, remaining))
except asyncio.TimeoutError:
continue
stdout, stderr = await communicate_task
except BaseException:
await _terminate_process(process)
if not communicate_task.done():
communicate_task.cancel()
try:
await communicate_task
except BaseException:
pass
raise
if process.returncode != 0:
error_text = (stderr or b"").decode("utf-8", errors="replace").strip()
raise LocalVideoUpscaleError(f"本地 FFmpeg 超分失败: {error_text[-4000:]}")
async def execute_local_ffmpeg_crop(
@@ -94,6 +142,7 @@ async def execute_local_ffmpeg_crop(
target_width: int,
target_height: int,
timeout_seconds: int | None = None,
execution_guard: Callable[[], Awaitable[None]] | None = None,
) -> str:
if not is_valid_file(source_path):
raise LocalVideoUpscaleError(f"超分源视频不存在或为空: {source_path}")
@@ -106,13 +155,13 @@ async def execute_local_ffmpeg_crop(
part_path = build_part_mp4_path(final_path)
safe_remove(part_path)
try:
await asyncio.to_thread(
_run_ffmpeg_sync,
await _run_ffmpeg(
source_path=source_path,
part_path=part_path,
target_width=target_width,
target_height=target_height,
timeout_seconds=int(timeout_seconds or settings.VIDEO_UPSCALE_LOCAL_TIMEOUT_SECONDS),
execution_guard=execution_guard,
)
if not is_valid_file(part_path):
raise LocalVideoUpscaleError("本地 FFmpeg 输出文件为空")
@@ -121,6 +170,8 @@ async def execute_local_ffmpeg_crop(
raise LocalVideoUpscaleError(
f"本地 FFmpeg 输出尺寸不正确: {info.width}x{info.height},预期 {target_width}x{target_height}"
)
if execution_guard is not None:
await execution_guard()
os.replace(part_path, final_path)
return final_path
except Exception:
@@ -9,6 +9,7 @@ import uuid
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from pathlib import Path
from collections.abc import Awaitable, Callable
from typing import Any
from urllib.parse import parse_qsl, urlsplit
@@ -180,7 +181,13 @@ def build_local_source_signed_url(source_local_path: str, expire_seconds: int) -
)
async def download_video_to_path(url: str, final_path: str, timeout_seconds: int) -> str:
async def download_video_to_path(
url: str,
final_path: str,
timeout_seconds: int,
*,
execution_guard: Callable[[], Awaitable[None]] | None = None,
) -> str:
if is_valid_file(final_path):
try:
await probe_video(final_path)
@@ -196,10 +203,14 @@ async def download_video_to_path(url: str, final_path: str, timeout_seconds: int
response.raise_for_status()
with open(part_path, "wb") as file_obj:
async for chunk in response.aiter_bytes(chunk_size=1024 * 1024):
if execution_guard is not None:
await execution_guard()
file_obj.write(chunk)
if not is_valid_file(part_path):
raise RuntimeError("视频下载完成但临时文件为空")
await probe_video(part_path)
if execution_guard is not None:
await execution_guard()
os.replace(part_path, final_path)
return final_path
except Exception:
@@ -10,8 +10,7 @@ from app.enums.generation_task import ChatGenerationPipelineStage, ChatGeneratio
from app.models.chat_generation_task import ChatGenerationTask
from app.models.generation_record import GenerationRecord
from app.models.video_upscale_task import VideoUpscaleTask
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
from app.services.generation.pipeline.db_lock_service import apply_short_lock_timeout
VideoUpscaleOwner: TypeAlias = ChatGenerationTask | GenerationRecord
@@ -73,6 +72,7 @@ async def load_upscale_owner(
else:
return None
if for_update:
await apply_short_lock_timeout(db)
query = query.with_for_update()
result = await db.execute(query.limit(1))
return result.scalar_one_or_none()
@@ -88,8 +88,6 @@ async def mark_owner_upscale_failed(
owner.status = ChatGenerationTaskStatus.FAILED.value
owner.pipeline_stage = ChatGenerationPipelineStage.UPSCALE_FAILED.value
owner.error_message = error_message
await notify_chat_generation_task_finished(db, owner)
await aggregate_parent_for_child(db, owner)
return
owner.status = GenerationStatus.failed.value
@@ -5,6 +5,7 @@ import math
import os
import uuid
from datetime import datetime, timedelta, timezone
from collections.abc import Awaitable, Callable
from typing import Any
from urllib.parse import urlsplit, urlunsplit
@@ -30,8 +31,9 @@ from app.enums.video_upscale import (
from app.models.chat_generation_task import ChatGenerationTask
from app.models.generation_record import GenerationRecord
from app.models.video_upscale_task import VideoUpscaleTask
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
from app.services.generation.pipeline.db_lock_service import apply_short_lock_timeout
from app.services.generation.pipeline.lifecycle_service import notify_owner_finished
from app.services.redis_registry_service import RedisExecutionLockError, RedisExecutionLockLost
from app.services.media_token_usage_snapshot_service import (
sync_chat_generation_task_media_token_snapshot,
sync_generation_record_media_token_snapshot,
@@ -73,6 +75,14 @@ from app.services.video_upscale.volc_service import (
from app.utils.id_gen import generate_id
ExecutionGuard = Callable[[], Awaitable[None]]
async def _ensure_guard(execution_guard: ExecutionGuard | None) -> None:
if execution_guard is not None:
await execution_guard()
def _now() -> datetime:
return datetime.now(timezone.utc)
@@ -185,6 +195,7 @@ async def _load_pair(
) -> tuple[VideoUpscaleTask | None, VideoUpscaleOwner | None]:
query = select(VideoUpscaleTask).where(VideoUpscaleTask.id == upscale_task_id)
if for_update:
await apply_short_lock_timeout(db)
query = query.with_for_update()
result = await db.execute(query.limit(1))
upscale = result.scalar_one_or_none()
@@ -217,6 +228,7 @@ async def prepare_video_upscale_task(
source_local_path: str,
source_file_size_bytes: int,
source_remote_url: str | None = None,
source_info: Any | None = None,
) -> VideoUpscaleTask:
owner: VideoUpscaleOwner | None = task or generation_record
if owner is None:
@@ -224,7 +236,9 @@ async def prepare_video_upscale_task(
if owner.gen_type != GenerationType.VIDEO.value or not bool(owner.video_upscale_enabled_snapshot):
raise RuntimeError("当前任务未启用视频超分快照")
snapshot = _snapshot(owner)
source_info = await probe_video(source_local_path)
# ffprobe should run before the owner row is locked by the caller.
# Keep this fallback for non-generation callers that do not pass probe data.
source_info = source_info or await probe_video(source_local_path)
remote_url = str(source_remote_url or getattr(owner, "remote_result_url", None) or "").strip() or None
signed_at, expires_at = parse_tos_signed_url_expiry(remote_url)
@@ -233,6 +247,7 @@ async def prepare_video_upscale_task(
if isinstance(owner, ChatGenerationTask)
else VideoUpscaleTask.generation_record_id == owner.id
)
await apply_short_lock_timeout(db)
result = await db.execute(
select(VideoUpscaleTask).where(owner_filter).with_for_update().limit(1)
)
@@ -333,6 +348,7 @@ async def _claim(
stage: str,
chat_stage: str,
increment_attempt: bool,
execution_token: str,
lease_seconds: int | None = None,
) -> tuple[VideoUpscaleTask, VideoUpscaleOwner, dict[str, Any]] | None:
upscale, task = await _load_pair(db, upscale_task_id, for_update=True)
@@ -349,7 +365,7 @@ async def _claim(
snapshot = _snapshot(task)
upscale.status = VideoUpscaleTaskStatus.PROCESSING.value
upscale.stage = stage
upscale.lease_token = uuid.uuid4().hex
upscale.lease_token = execution_token
upscale.lease_until = _lease_until(lease_seconds)
upscale.started_at = upscale.started_at or _now()
upscale.next_retry_at = None
@@ -360,6 +376,21 @@ async def _claim(
return upscale, task, snapshot
async def _load_owned_pair(
db: AsyncSession,
*,
upscale_task_id: str,
execution_token: str,
for_update: bool = True,
) -> tuple[VideoUpscaleTask, VideoUpscaleOwner]:
upscale, task = await _load_pair(db, upscale_task_id, for_update=for_update)
if not upscale or not task:
raise RedisExecutionLockLost(f"超分任务或所属任务不存在: {upscale_task_id}")
if str(upscale.lease_token or "") != str(execution_token):
raise RedisExecutionLockLost(f"超分数据库执行租约已失效: {upscale_task_id}")
return upscale, task
async def _final_fail(
db: AsyncSession,
*,
@@ -381,6 +412,7 @@ async def _final_fail(
upscale.lease_token = None
await mark_owner_upscale_failed(db, task, error_message=error_message)
await db.commit()
await notify_owner_finished(db, task)
log_video_upscale_event(
event_type="upscale_retry_exhausted",
event_status="failed",
@@ -464,8 +496,13 @@ async def _queue_finalize(
upscale_task_id: str,
final_path: str,
reason: str,
execution_token: str,
execution_guard: ExecutionGuard | None = None,
) -> None:
upscale, task = await _load_pair(db, upscale_task_id, for_update=True)
await _ensure_guard(execution_guard)
upscale, task = await _load_owned_pair(
db, upscale_task_id=upscale_task_id, execution_token=execution_token, for_update=True
)
if not upscale or not task:
return
if not is_valid_file(final_path):
@@ -488,6 +525,7 @@ async def _queue_finalize(
task=task,
upscale_task=upscale,
action="finalize",
countdown=max(1, int(settings.VIDEO_UPSCALE_STAGE_HANDOFF_DELAY_SECONDS or 2)),
)
log_video_upscale_event(
event_type="upscale_finalize_enqueued",
@@ -498,8 +536,18 @@ async def _queue_finalize(
)
async def _finalize_success(db: AsyncSession, *, upscale_task_id: str, final_path: str) -> None:
upscale, task = await _load_pair(db, upscale_task_id, for_update=False)
async def _finalize_success(
db: AsyncSession,
*,
upscale_task_id: str,
final_path: str,
execution_token: str,
execution_guard: ExecutionGuard | None = None,
) -> None:
await _ensure_guard(execution_guard)
upscale, task = await _load_owned_pair(
db, upscale_task_id=upscale_task_id, execution_token=execution_token, for_update=False
)
if not upscale or not task:
return
snapshot = _snapshot(task)
@@ -550,7 +598,10 @@ async def _finalize_success(db: AsyncSession, *, upscale_task_id: str, final_pat
if not cover_url or not cover_path:
raise RuntimeError("超分最终视频封面生成失败")
upscale, task = await _load_pair(db, upscale_task_id, for_update=True)
await _ensure_guard(execution_guard)
upscale, task = await _load_owned_pair(
db, upscale_task_id=upscale_task_id, execution_token=execution_token, for_update=True
)
if not upscale or not task:
return
if owner_is_completed(task) and task.video_url:
@@ -569,7 +620,7 @@ async def _finalize_success(db: AsyncSession, *, upscale_task_id: str, final_pat
task.generated_at = now
task.error_message = None
if isinstance(task, ChatGenerationTask):
task.retry_count = 0
task.retry_count = int(task.manual_retry_count or 0)
upscale.status = VideoUpscaleTaskStatus.COMPLETED.value
upscale.stage = VideoUpscaleStage.COMPLETED.value
@@ -595,8 +646,6 @@ async def _finalize_success(db: AsyncSession, *, upscale_task_id: str, final_pat
generated_at=now,
)
await sync_chat_generation_task_media_token_snapshot(db, task)
await notify_chat_generation_task_finished(db, task)
await aggregate_parent_for_child(db, task)
else:
await record_generation_record_generated_resource(
db,
@@ -614,9 +663,12 @@ async def _finalize_success(db: AsyncSession, *, upscale_task_id: str, final_pat
if not delete_source_after_success:
upscale.source_delete_error = VIDEO_UPSCALE_SOURCE_RETAINED_MARKER
await db.commit()
if isinstance(task, ChatGenerationTask):
await notify_owner_finished(db, task)
if delete_source_after_success and source_path and os.path.abspath(source_path) != os.path.abspath(final_path):
removed = safe_remove(source_path)
cleanup_error = None if removed else f"源视频删除失败: {source_path}"
await apply_short_lock_timeout(db)
cleanup_result = await db.execute(
select(VideoUpscaleTask).where(VideoUpscaleTask.id == upscale_task_id).with_for_update().limit(1)
)
@@ -668,13 +720,16 @@ async def _finalize_success(db: AsyncSession, *, upscale_task_id: str, final_pat
)
async def run_local_upscale(db: AsyncSession, upscale_task_id: str) -> None:
async def run_local_upscale(
db: AsyncSession, upscale_task_id: str, *, execution_token: str, execution_guard: ExecutionGuard | None = None
) -> None:
claimed = await _claim(
db,
upscale_task_id=upscale_task_id,
stage=VideoUpscaleStage.LOCAL_PROCESSING.value,
chat_stage=ChatGenerationPipelineStage.UPSCALE_PROCESSING.value,
increment_attempt=True,
execution_token=execution_token,
lease_seconds=int(settings.VIDEO_UPSCALE_LOCAL_TIMEOUT_SECONDS or 3600) + 300,
)
if not claimed:
@@ -690,15 +745,21 @@ async def run_local_upscale(db: AsyncSession, upscale_task_id: str) -> None:
target_width=int(snapshot.get("target_width") or upscale.target_width),
target_height=int(snapshot.get("target_height") or upscale.target_height),
timeout_seconds=int(processor.get("timeout_seconds") or settings.VIDEO_UPSCALE_LOCAL_TIMEOUT_SECONDS),
execution_guard=execution_guard,
)
await _ensure_guard(execution_guard)
log_video_upscale_event(
event_type="upscale_local_success",
task=task,
upscale_task=upscale,
detail={"final_local_path": final_path},
)
except RedisExecutionLockError:
await db.rollback()
raise
except Exception as exc:
await db.rollback()
await _ensure_guard(execution_guard)
await _schedule_retry(
db,
upscale_task_id=upscale_task_id,
@@ -714,9 +775,15 @@ async def run_local_upscale(db: AsyncSession, upscale_task_id: str) -> None:
upscale_task_id=upscale_task_id,
final_path=final_path,
reason="local_upscale_completed",
execution_token=execution_token,
execution_guard=execution_guard,
)
except RedisExecutionLockError:
await db.rollback()
raise
except Exception as exc:
await db.rollback()
await _ensure_guard(execution_guard)
await _schedule_retry(
db,
upscale_task_id=upscale_task_id,
@@ -770,13 +837,17 @@ async def _select_remote_input(db: AsyncSession, upscale: VideoUpscaleTask, proc
return signed_url, VideoUpscaleInputSourceType.LOCAL_SIGNED.value
async def run_remote_submit(db: AsyncSession, upscale_task_id: str, *, count_attempt: bool = True) -> None:
async def run_remote_submit(
db: AsyncSession, upscale_task_id: str, *, count_attempt: bool = True, execution_token: str,
execution_guard: ExecutionGuard | None = None,
) -> None:
claimed = await _claim(
db,
upscale_task_id=upscale_task_id,
stage=VideoUpscaleStage.REMOTE_SUBMITTING.value,
chat_stage=ChatGenerationPipelineStage.UPSCALE_PROCESSING.value,
increment_attempt=count_attempt,
execution_token=execution_token,
lease_seconds=300,
)
if not claimed:
@@ -831,9 +902,10 @@ async def run_remote_submit(db: AsyncSession, upscale_task_id: str, *, count_att
processor=processor,
client_token=f"{upscale.id}-{int(upscale.attempt_count or 0)}-{int(upscale.input_source_fallback_count or 0)}",
)
upscale, task = await _load_pair(db, upscale_task_id, for_update=True)
if not upscale or not task:
return
await _ensure_guard(execution_guard)
upscale, task = await _load_owned_pair(
db, upscale_task_id=upscale_task_id, execution_token=execution_token, for_update=True
)
upscale.provider_task_id = submit_result.task_id
upscale.provider_submitted_at = _now()
upscale.provider_request_json = json.dumps(_sanitize_provider_payload(submit_result.request_payload), ensure_ascii=False, default=str)
@@ -863,6 +935,9 @@ async def run_remote_submit(db: AsyncSession, upscale_task_id: str, *, count_att
remote_request_id=submit_result.request_id,
detail={"provider_task_id": submit_result.task_id, "input_source_type": source_type},
)
except RedisExecutionLockError:
await db.rollback()
raise
except VolcMediaKitError as exc:
log_video_upscale_event(
event_type="upscale_provider_submit_failed",
@@ -875,6 +950,7 @@ async def run_remote_submit(db: AsyncSession, upscale_task_id: str, *, count_att
error=str(exc),
)
await db.rollback()
await _ensure_guard(execution_guard)
await _persist_provider_error_payload(
db,
upscale_task_id=upscale_task_id,
@@ -897,6 +973,7 @@ async def run_remote_submit(db: AsyncSession, upscale_task_id: str, *, count_att
error=str(exc),
)
await db.rollback()
await _ensure_guard(execution_guard)
await _schedule_retry(
db,
upscale_task_id=upscale_task_id,
@@ -905,13 +982,16 @@ async def run_remote_submit(db: AsyncSession, upscale_task_id: str, *, count_att
)
async def run_remote_poll(db: AsyncSession, upscale_task_id: str) -> None:
async def run_remote_poll(
db: AsyncSession, upscale_task_id: str, *, execution_token: str, execution_guard: ExecutionGuard | None = None
) -> None:
claimed = await _claim(
db,
upscale_task_id=upscale_task_id,
stage=VideoUpscaleStage.REMOTE_POLLING.value,
chat_stage=ChatGenerationPipelineStage.UPSCALE_POLLING.value,
increment_attempt=False,
execution_token=execution_token,
lease_seconds=300,
)
if not claimed:
@@ -939,9 +1019,10 @@ async def run_remote_poll(db: AsyncSession, upscale_task_id: str) -> None:
provider_task_id,
request_timeout_seconds=int(processor.get("request_timeout_seconds") or 30),
)
upscale, task = await _load_pair(db, upscale_task_id, for_update=True)
if not upscale or not task:
return
await _ensure_guard(execution_guard)
upscale, task = await _load_owned_pair(
db, upscale_task_id=upscale_task_id, execution_token=execution_token, for_update=True
)
upscale.provider_response_json = json.dumps(query_result.response_payload, ensure_ascii=False, default=str)
upscale.lease_until = None
upscale.lease_token = None
@@ -1001,6 +1082,7 @@ async def run_remote_poll(db: AsyncSession, upscale_task_id: str) -> None:
task=task,
upscale_task=upscale,
action="submit_local_source_fallback",
countdown=max(1, int(settings.VIDEO_UPSCALE_STAGE_HANDOFF_DELAY_SECONDS or 2)),
)
log_video_upscale_event(
event_type="upscale_source_fallback_local",
@@ -1071,6 +1153,7 @@ async def run_remote_poll(db: AsyncSession, upscale_task_id: str) -> None:
task=task,
upscale_task=upscale,
action="download_remote_result",
countdown=max(1, int(settings.VIDEO_UPSCALE_STAGE_HANDOFF_DELAY_SECONDS or 2)),
)
log_video_upscale_event(
event_type="upscale_provider_poll_success",
@@ -1079,6 +1162,9 @@ async def run_remote_poll(db: AsyncSession, upscale_task_id: str) -> None:
remote_request_id=query_result.request_id,
detail={"provider_output_url_expires_at": upscale.provider_output_url_expires_at},
)
except RedisExecutionLockError:
await db.rollback()
raise
except VolcMediaKitError as exc:
log_video_upscale_event(
event_type="upscale_provider_poll_failed",
@@ -1091,6 +1177,7 @@ async def run_remote_poll(db: AsyncSession, upscale_task_id: str) -> None:
error=str(exc),
)
await db.rollback()
await _ensure_guard(execution_guard)
await _persist_provider_error_payload(
db,
upscale_task_id=upscale_task_id,
@@ -1113,6 +1200,7 @@ async def run_remote_poll(db: AsyncSession, upscale_task_id: str) -> None:
error=str(exc),
)
await db.rollback()
await _ensure_guard(execution_guard)
await _schedule_retry(
db,
upscale_task_id=upscale_task_id,
@@ -1121,13 +1209,16 @@ async def run_remote_poll(db: AsyncSession, upscale_task_id: str) -> None:
)
async def run_finalize_upscale(db: AsyncSession, upscale_task_id: str) -> None:
async def run_finalize_upscale(
db: AsyncSession, upscale_task_id: str, *, execution_token: str, execution_guard: ExecutionGuard | None = None
) -> None:
claimed = await _claim(
db,
upscale_task_id=upscale_task_id,
stage=VideoUpscaleStage.FINALIZING.value,
chat_stage=ChatGenerationPipelineStage.UPSCALE_FINALIZING.value,
increment_attempt=False,
execution_token=execution_token,
lease_seconds=max(300, int(settings.VIDEO_COVER_TIMEOUT_SECONDS or 15) + 300),
)
if not claimed:
@@ -1143,13 +1234,19 @@ async def run_finalize_upscale(db: AsyncSession, upscale_task_id: str) -> None:
detail={"final_local_path": final_path},
)
try:
await _finalize_success(db, upscale_task_id=upscale_task_id, final_path=final_path)
await _finalize_success(
db, upscale_task_id=upscale_task_id, final_path=final_path,
execution_token=execution_token, execution_guard=execution_guard,
)
log_video_upscale_event(
event_type="upscale_finalize_success",
task=task,
upscale_task=upscale,
detail={"final_local_path": final_path},
)
except RedisExecutionLockError:
await db.rollback()
raise
except Exception as exc:
log_video_upscale_event(
event_type="upscale_finalize_failed",
@@ -1161,6 +1258,7 @@ async def run_finalize_upscale(db: AsyncSession, upscale_task_id: str) -> None:
error=str(exc),
)
await db.rollback()
await _ensure_guard(execution_guard)
await _schedule_retry(
db,
upscale_task_id=upscale_task_id,
@@ -1170,13 +1268,16 @@ async def run_finalize_upscale(db: AsyncSession, upscale_task_id: str) -> None:
)
async def run_remote_result_download(db: AsyncSession, upscale_task_id: str) -> None:
async def run_remote_result_download(
db: AsyncSession, upscale_task_id: str, *, execution_token: str, execution_guard: ExecutionGuard | None = None
) -> None:
claimed = await _claim(
db,
upscale_task_id=upscale_task_id,
stage=VideoUpscaleStage.RESULT_DOWNLOADING.value,
chat_stage=ChatGenerationPipelineStage.UPSCALE_DOWNLOADING.value,
increment_attempt=False,
execution_token=execution_token,
lease_seconds=int(settings.VIDEO_UPSCALE_REMOTE_RESULT_DOWNLOAD_TIMEOUT_SECONDS or 600) + 120,
)
if not claimed:
@@ -1198,6 +1299,7 @@ async def run_remote_result_download(db: AsyncSession, upscale_task_id: str) ->
task=task,
upscale_task=upscale,
action="renew_remote_result",
countdown=max(1, int(settings.VIDEO_UPSCALE_STAGE_HANDOFF_DELAY_SECONDS or 2)),
)
log_video_upscale_event(
event_type="upscale_provider_result_renew_query",
@@ -1228,13 +1330,18 @@ async def run_remote_result_download(db: AsyncSession, upscale_task_id: str) ->
output_url,
final_path,
int(settings.VIDEO_UPSCALE_REMOTE_RESULT_DOWNLOAD_TIMEOUT_SECONDS or 600),
execution_guard=execution_guard,
)
await _ensure_guard(execution_guard)
log_video_upscale_event(
event_type="upscale_output_download_success",
task=task,
upscale_task=upscale,
detail={"final_local_path": final_path, "file_size_bytes": safe_file_size(final_path)},
)
except RedisExecutionLockError:
await db.rollback()
raise
except Exception as exc:
log_video_upscale_event(
event_type="upscale_output_download_failed",
@@ -1246,6 +1353,7 @@ async def run_remote_result_download(db: AsyncSession, upscale_task_id: str) ->
error=str(exc),
)
await db.rollback()
await _ensure_guard(execution_guard)
expires_at = _aware(upscale.provider_output_url_expires_at)
action = "submit" if expires_at and expires_at <= _now() else "download"
await _schedule_retry(
@@ -1263,9 +1371,15 @@ async def run_remote_result_download(db: AsyncSession, upscale_task_id: str) ->
upscale_task_id=upscale_task_id,
final_path=final_path,
reason="remote_result_download_completed",
execution_token=execution_token,
execution_guard=execution_guard,
)
except RedisExecutionLockError:
await db.rollback()
raise
except Exception as exc:
await db.rollback()
await _ensure_guard(execution_guard)
await _schedule_retry(
db,
upscale_task_id=upscale_task_id,
@@ -1278,6 +1392,7 @@ async def run_remote_result_download(db: AsyncSession, upscale_task_id: str) ->
async def recover_video_upscale_tasks_once(db: AsyncSession) -> dict[str, Any]:
now = _now()
batch_size = max(1, int(settings.VIDEO_UPSCALE_RECOVERY_BATCH_SIZE or 50))
await apply_short_lock_timeout(db)
result = await db.execute(
select(VideoUpscaleTask)
.where(
@@ -1339,6 +1454,7 @@ async def recover_video_upscale_tasks_once(db: AsyncSession) -> dict[str, Any]:
continue
source_path = str(cleanup_upscale.source_local_path or "")
removed = safe_remove(source_path)
await apply_short_lock_timeout(db)
cleanup_result = await db.execute(
select(VideoUpscaleTask).where(VideoUpscaleTask.id == item.id).with_for_update().limit(1)
)
+483 -300
View File
@@ -1,29 +1,51 @@
from app.tasks.async_runner import run_async
from __future__ import annotations
import asyncio
import json
from datetime import datetime, timedelta, timezone
from typing import Any, Optional
from sqlalchemy import select
from app.config import settings
from app.enums.celery_queue import CeleryQueue
from app.enums.generation_status import GenerationRecordPipelineStage
from app.enums.generation_task import (
ALLOWED_GENERATION_MODES,
ChatGenerationPipelineStage,
ChatGenerationTaskEventType,
ChatGenerationTaskStatus,
GenerationMode,
GenerationOwnerType,
GenerationType,
)
from app.models.base import async_session
from app.models.chat_generation_task import ChatGenerationTask
from app.services.error_codes import extract_error_message
from app.services.generation.log_service import log_task_event
from app.services.generation.pipeline.db_lock_service import DatabaseRowLockBusy
from app.services.generation.pipeline.lifecycle_service import (
mark_owner_failed_and_refund_once,
notify_owner_finished,
)
from app.services.generation.pipeline.owner_service import (
GenerationOwner,
is_attempt_current,
load_generation_owner,
normalize_owner_type,
owner_is_generating,
owner_mode,
owner_provider_task_id,
set_owner_provider_task_id,
)
from app.services.generation.poll_schedule_service import ensure_video_poll_fields
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.generation.provider_service import create_provider_task
from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
from app.services.redis_registry_service import ensure_aware_utc
from app.services.media_token_usage_snapshot_service import (
sync_chat_generation_task_media_token_snapshot,
sync_generation_record_media_token_snapshot,
)
from app.services.redis_registry_service import (
RedisExecutionLockError,
RedisExecutionLockLease,
)
from app.tasks.async_runner import run_async
from app.tasks.celery_app import celery_app
@@ -32,341 +54,502 @@ def _now() -> datetime:
def _get_first_value(obj: Any, *field_names: str) -> Optional[Any]:
"""
兼容不同版本字段名避免字段调整后 Celery 任务直接报错
"""
for field_name in field_names:
if hasattr(obj, field_name):
value = getattr(obj, field_name)
if value is not None and str(value).strip() != "":
return value
value = getattr(obj, field_name, None)
if value is not None and str(value).strip() != "":
return value
return None
def _to_clean_str(value: Any) -> Optional[str]:
if value is None:
return None
text = str(value).strip()
return text if text else None
def _clean(value: Any) -> str | None:
text = str(value).strip() if value is not None else ""
return text or None
def _format_duration(value: Any) -> Optional[str]:
"""
duration=4 -> 4
duration="4秒" -> 4
"""
text = _to_clean_str(value)
if not text:
return None
lower_text = text.lower()
if text.endswith("") or lower_text.endswith("s") or lower_text.endswith("sec") or lower_text.endswith("seconds"):
return text
return f"{text}"
def _build_optimized_prompt_by_params(task: ChatGenerationTask) -> str:
"""
不调用提词优化 API直接将 original_prompt 拼接上对应类型的生成参数
视频示例
original_prompt时长4画面比例16:9分辨率480p
图片示例
original_prompt分辨率2K画布比例1:1像素尺寸2048×2048
"""
original_prompt = _to_clean_str(getattr(task, "original_prompt", None)) or ""
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 in {GenerationMode.HOT_OPENING_REPLICATE.value, GenerationMode.SHOT_REPLICATE.value} and gen_type == GenerationType.VIDEO.value:
def _build_optimized_prompt_by_params(owner: GenerationOwner) -> str:
base_prompt = (_clean(owner.original_prompt) or "").rstrip(",。;; \n\t")
gen_type = (_clean(owner.gen_type) or "").lower()
generation_mode = owner_mode(owner)
if (
generation_mode
in {
GenerationMode.HOT_OPENING_REPLICATE.value,
GenerationMode.SHOT_REPLICATE.value,
}
and gen_type == GenerationType.VIDEO.value
):
stripped = base_prompt.strip()
if stripped.startswith("{") or stripped.startswith("["):
if stripped.startswith(("{", "[")):
return base_prompt
duration = _get_first_value(task, "duration")
aspect_ratio = _get_first_value(task, "aspect_ratio")
resolution = _get_first_value(task, "provider_generation_resolution", "resolution")
image_size = _get_first_value(task, "image_size")
image_px = _get_first_value(task, "image_px")
image_proportion = _get_first_value(task, "image_proportion")
parts = []
parts: list[str] = []
if gen_type == GenerationType.VIDEO.value:
# 时长:4秒,画面比例:16:9,分辨率:480p
if duration:
parts.append(f"时长:{duration}")
parts.append(f"画面比例:{aspect_ratio}")
parts.append(f"分辨率:{resolution}")
else:
parts.append("时长:4秒")
parts.append("画面比例:16:9")
parts.append("分辨率:480p")
elif gen_type == GenerationType.IMAGE.value:
if image_size:
parts.append(f"分辨率:{image_size}")
parts.append(f"画布比例:{image_proportion}")
parts.append(f"像素尺寸:{image_px}")
else:
parts.append("分辨率:2K")
parts.append("画布比例:1:1")
parts.append("像素尺寸:2048x2048")
else:
# 未知类型时返回原始字符
return base_prompt
suffix = "".join(parts)
if base_prompt and suffix:
return f"{base_prompt}{suffix}"
if base_prompt:
return base_prompt
return suffix
async def _run(task_id: str):
async with async_session() as db:
result = await db.execute(select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
).with_for_update().limit(1))
task = result.scalar_one_or_none()
is_image_main = bool(
task
and task.generation_mode == GenerationMode.CHATAPI_MAIN.value
and task.gen_type == GenerationType.IMAGE.value
and int(task.generation_count or 1) > 1
parts.extend(
[
f"时长:{_get_first_value(owner, 'duration') or 4}",
f"画面比例:{_get_first_value(owner, 'aspect_ratio') or '16:9'}",
f"分辨率:{_get_first_value(owner, 'provider_generation_resolution', 'resolution') or '480p'}",
]
)
if not task or (task.generation_mode not in ALLOWED_GENERATION_MODES and not is_image_main):
return
elif gen_type == GenerationType.IMAGE.value:
parts.extend(
[
f"分辨率:{_get_first_value(owner, 'image_size') or '2K'}",
f"画布比例:{_get_first_value(owner, 'image_proportion') or '1:1'}",
f"像素尺寸:{_get_first_value(owner, 'image_px') or '2048x2048'}",
]
)
suffix = "".join(parts)
return f"{base_prompt}{suffix}" if base_prompt and suffix else base_prompt or suffix
if task.status != ChatGenerationTaskStatus.GENERATING.value:
return
deadline_at = ensure_aware_utc(task.deadline_at)
# 图片 main 的 deadline 与 provider claim 由 image_batch_service 原子处理,
# 避免重复 Celery 消息在有效租约期间把正在执行的批次错误退款。
if not is_image_main and deadline_at and datetime.now(timezone.utc) > deadline_at:
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message="任务超时",
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
)
await db.commit()
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
to_status=ChatGenerationTaskStatus.FAILED.value,
to_stage=ChatGenerationPipelineStage.TIMEOUT.value,
)
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
await notify_chat_generation_task_finished(db, task)
await aggregate_parent_for_child(db, task)
await db.commit()
return
def _stage(owner: GenerationOwner, chat_stage: ChatGenerationPipelineStage) -> str:
if isinstance(owner, ChatGenerationTask):
return chat_stage.value
try:
return GenerationRecordPipelineStage(chat_stage.value).value
except ValueError:
return chat_stage.value
if task.pipeline_stage not in (
ChatGenerationPipelineStage.QUEUED.value,
ChatGenerationPipelineStage.PREPARING.value,
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
):
return
def _lock_key(owner_type: str, owner_id: str, attempt_no: int) -> str:
return (
f"{settings.GENERATION_CREATE_LOCK_KEY_PREFIX}:"
f"{owner_type}:{owner_id}:attempt:{attempt_no}"
)
async def _sync_media_snapshot(
db, owner: GenerationOwner, provider_response: Any = None
) -> None:
if isinstance(owner, ChatGenerationTask):
await sync_chat_generation_task_media_token_snapshot(
db, owner, provider_response=provider_response
)
else:
await sync_generation_record_media_token_snapshot(
db, owner, provider_response=provider_response
)
async def _stale(owner: GenerationOwner, attempt_no: int | None) -> bool:
if is_attempt_current(owner, attempt_no):
return False
await log_task_event(
owner,
event_type=ChatGenerationTaskEventType.STALE_ATTEMPT_MESSAGE_SKIPPED.value,
message="创建任务消息属于旧生成轮次,已跳过",
detail={
"message_attempt": attempt_no,
"current_attempt": owner.generation_attempt_no,
},
)
return True
async def _reload_owner_after_external_call(
db,
*,
owner_type: str,
owner_id: str,
) -> GenerationOwner | None:
"""Keep the provider result in the current Worker while briefly retrying a busy row lock."""
last_error: DatabaseRowLockBusy | None = None
for retry_index in range(3):
try:
if not task.optimized_prompt:
old_stage = task.pipeline_stage
task.pipeline_stage = ChatGenerationPipelineStage.PREPARING.value
await db.commit()
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.PROMPT_CONCAT_START.value,
from_stage=old_stage,
to_stage=ChatGenerationPipelineStage.PREPARING.value,
message="开始本地拼接提示词,不调用提词优化API",
)
return await load_generation_owner(
db,
owner_type=owner_type,
owner_id=owner_id,
for_update=True,
)
except DatabaseRowLockBusy as exc:
last_error = exc
await db.rollback()
if retry_index < 2:
await asyncio.sleep(1 + retry_index)
raise last_error or DatabaseRowLockBusy()
optimized_prompt = _build_optimized_prompt_by_params(task)
task.optimized_prompt = optimized_prompt
# 不调用提词优化 API,因此不产生模型 token 消耗。
task.text_tokens_used = task.text_tokens_used or 0
async def _resolve_attempt(
task_id: str, *, owner_type: str, message_attempt: int | None
) -> int | None:
async with async_session() as db:
owner = await load_generation_owner(
db, owner_type=owner_type, owner_id=task_id, for_update=False
)
if not owner:
return None
if await _stale(owner, message_attempt):
return None
return int(owner.generation_attempt_no or 1)
await db.commit()
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.PROMPT_CONCAT_SUCCESS.value,
to_stage=ChatGenerationPipelineStage.PREPARING.value,
detail={
"optimized_prompt": optimized_prompt,
"gen_type": task.gen_type,
"message": "已完成本地提示词拼接,未调用提词优化API",
},
)
if is_image_main:
from app.services.generation.ai.image_batch_service import run_image_main_batch
async def _dispatch_next_stage(
db,
owner: GenerationOwner,
*,
normalized_owner_type: str,
) -> None:
if owner.pipeline_stage == _stage(
owner, ChatGenerationPipelineStage.RESULT_READY
):
from app.tasks.generation_download_tasks import enqueue_download_task
await run_image_main_batch(db, task)
await enqueue_download_task(db, owner, reason="create_result_ready")
return
from app.tasks.generation_poll_tasks import (
poll_generation_task,
register_poll_active,
)
try:
check_at = _now() + timedelta(
seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300)
)
await register_poll_active(
owner,
reason="create_provider_success",
check_at=check_at,
next_poll_at=owner.next_poll_at,
)
poll_generation_task.apply_async(
args=[owner.id],
kwargs={
"owner_type": normalized_owner_type,
"generation_attempt_no": int(owner.generation_attempt_no or 1),
"force_due": False,
},
queue=CeleryQueue.GEN_PROVIDER_POLL.value,
)
except Exception as exc:
# Provider creation is already committed. A broker/registry failure is
# an infrastructure enqueue failure, not a generation failure. Keep the
# remote task ID and let due-poll/startup recovery enqueue it again.
owner.next_poll_at = _now()
await db.commit()
await log_task_event(
owner,
event_type=ChatGenerationTaskEventType.POLL_SCHEDULED.value,
message="供应商任务已创建,但轮询任务投递失败,等待恢复扫描",
detail={
"error": str(exc),
"provider_task_id": owner_provider_task_id(owner),
"generation_refunded": False,
},
to_stage=owner.pipeline_stage,
)
async def _run(
task_id: str,
*,
owner_type: str = GenerationOwnerType.CHAT_GENERATION_TASK.value,
generation_attempt_no: int | None = None,
):
normalized_owner_type = normalize_owner_type(owner_type)
effective_attempt = await _resolve_attempt(
task_id,
owner_type=normalized_owner_type,
message_attempt=generation_attempt_no,
)
if effective_attempt is None:
return
lease = await RedisExecutionLockLease.acquire(
lock_key=_lock_key(normalized_owner_type, task_id, effective_attempt),
ttl_seconds=int(settings.GENERATION_CREATE_LOCK_TTL_SECONDS or 600),
log_context="generation_create",
renew_interval_seconds=int(
settings.REDIS_EXECUTION_LOCK_RENEW_INTERVAL_SECONDS or 30
),
)
if lease is None:
return
try:
async with async_session() as db:
owner = await load_generation_owner(
db,
owner_type=normalized_owner_type,
owner_id=task_id,
for_update=True,
)
if not owner:
return
if task.seedance_task_id or task.provider_task_id:
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
if task.gen_type == GenerationType.VIDEO.value:
ensure_video_poll_fields(task, now=_now())
task.next_poll_at = _now()
await db.commit()
elif task.remote_result_url:
task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value
await db.commit()
from app.tasks.generation_download_tasks import enqueue_download_task
await enqueue_download_task(db, task, reason="create_remote_result_ready")
return
else:
old_stage = task.pipeline_stage
task.pipeline_stage = ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value
await db.commit()
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_START.value,
from_stage=old_stage,
to_stage=ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
)
created = await create_provider_task(db, task)
provider_task_id = created.get("task_id")
if provider_task_id:
task.provider_task_id = provider_task_id
task.seedance_task_id = provider_task_id
task.remote_result_url = created.get("remote_result_url") or task.remote_result_url
if task.gen_type == GenerationType.IMAGE.value:
task.image_tokens_used = created.get("image_tokens", task.image_tokens_used or 0) or 0
task.provider_response_json = json.dumps(
created.get("response_data") or {},
ensure_ascii=False,
default=str,
)
await sync_chat_generation_task_media_token_snapshot(db, task, provider_response=task.provider_response_json)
if task.remote_result_url and not task.seedance_task_id:
# 同步图片路径:原 SDK 已经返回最终 URL。
task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value
else:
# 视频路径:provider 返回 task id,后续轮询。
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
if task.gen_type == GenerationType.VIDEO.value:
current_time = _now()
ensure_video_poll_fields(task, now=current_time)
task.next_poll_at = current_time
task.poll_interval_seconds = int(task.poll_interval_seconds or 0)
task.status = ChatGenerationTaskStatus.GENERATING.value
await db.commit()
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_SUCCESS.value,
to_stage=task.pipeline_stage,
detail=created,
)
if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value:
from app.tasks.generation_download_tasks import enqueue_download_task
await enqueue_download_task(db, task, reason="create_result_ready")
else:
from app.tasks.generation_poll_tasks import poll_generation_task, register_poll_active
await register_poll_active(
task,
reason="create_provider_success",
check_at=_now() + timedelta(seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300)),
next_poll_at=getattr(task, "next_poll_at", None),
)
poll_generation_task.apply_async(
args=[task.id],
queue=CeleryQueue.GEN_PROVIDER_POLL.value,
countdown=0,
)
except Exception as exc:
try:
if not is_attempt_current(owner, effective_attempt):
await db.rollback()
except Exception:
pass
return
result = await db.execute(select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
).with_for_update().limit(1))
task = result.scalar_one_or_none()
is_image_main = bool(
isinstance(owner, ChatGenerationTask)
and owner.generation_mode == GenerationMode.CHATAPI_MAIN.value
and owner.gen_type == GenerationType.IMAGE.value
and int(owner.generation_count or 1) > 1
)
if (
isinstance(owner, ChatGenerationTask)
and owner.generation_mode not in ALLOWED_GENERATION_MODES
and not is_image_main
):
return
if not owner_is_generating(owner):
return
if owner.deadline_at and _now() > owner.deadline_at and not is_image_main:
await mark_owner_failed_and_refund_once(
db,
owner,
error_message="任务超时",
pipeline_stage=_stage(
owner, ChatGenerationPipelineStage.TIMEOUT
),
)
await db.commit()
await notify_owner_finished(db, owner)
await db.commit()
await log_task_event(
owner,
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
to_stage=owner.pipeline_stage,
)
return
if task:
error_message = extract_error_message(exc, "生成任务") if callable(extract_error_message) else str(exc)
if is_image_main:
# image_batch_service 负责供应商/拆分失败退款。若 child 已落库,
# 顶层兜底绝不能再把 main 退款。
child_result = await db.execute(
select(ChatGenerationTask.id).where(
ChatGenerationTask.parent_task_id == task.id,
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
).limit(1)
allowed_stages = {
_stage(owner, ChatGenerationPipelineStage.QUEUED),
_stage(owner, ChatGenerationPipelineStage.PREPARING),
_stage(owner, ChatGenerationPipelineStage.CREATING_PROVIDER_TASK),
}
if owner.pipeline_stage not in allowed_stages:
return
try:
if not owner.optimized_prompt:
owner.pipeline_stage = _stage(
owner, ChatGenerationPipelineStage.PREPARING
)
has_children = child_result.scalar_one_or_none() is not None
if not has_children:
task.provider_create_claim_token = None
task.provider_create_lease_until = None
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=error_message,
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
owner.optimized_prompt = _build_optimized_prompt_by_params(owner)
owner.text_tokens_used = int(owner.text_tokens_used or 0)
await db.commit()
await log_task_event(
owner,
event_type=ChatGenerationTaskEventType.PROMPT_CONCAT_SUCCESS.value,
to_stage=owner.pipeline_stage,
)
if is_image_main:
from app.services.generation.ai.image_batch_service import (
run_image_main_batch,
)
await run_image_main_batch(
db,
owner,
execution_token=lease.token,
execution_guard=lease.ensure_owned,
)
return
if owner_provider_task_id(owner):
owner.pipeline_stage = _stage(
owner, ChatGenerationPipelineStage.WAITING_REMOTE
)
if owner.gen_type == GenerationType.VIDEO.value:
ensure_video_poll_fields(owner, now=_now())
owner.next_poll_at = _now()
await db.commit()
elif owner.remote_result_url:
owner.pipeline_stage = _stage(
owner, ChatGenerationPipelineStage.RESULT_READY
)
await db.commit()
else:
current_time = _now()
owner.pipeline_stage = _stage(
owner, ChatGenerationPipelineStage.CREATING_PROVIDER_TASK
)
owner.provider_create_claim_token = lease.token
owner.provider_create_started_at = current_time
owner.provider_create_lease_until = current_time + timedelta(
seconds=int(
settings.GENERATION_CREATE_LOCK_TTL_SECONDS or 600
)
)
await db.commit()
await log_task_event(
owner,
event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_START.value,
to_stage=owner.pipeline_stage,
)
created = await create_provider_task(db, owner)
await lease.ensure_owned()
owner = await _reload_owner_after_external_call(
db,
owner_type=normalized_owner_type,
owner_id=task_id,
)
if not owner:
return
if not is_attempt_current(owner, effective_attempt):
await db.rollback()
return
if owner.provider_create_claim_token != lease.token:
return
provider_task_id = created.get("task_id")
if provider_task_id:
set_owner_provider_task_id(owner, str(provider_task_id))
owner.remote_result_url = (
created.get("remote_result_url") or owner.remote_result_url
)
owner.provider_response_json = json.dumps(
created.get("response_data") or {},
ensure_ascii=False,
default=str,
)
owner.provider_create_claim_token = None
owner.provider_create_lease_until = None
if owner.gen_type == GenerationType.IMAGE.value:
owner.image_tokens_used = int(
created.get(
"image_tokens", owner.image_tokens_used or 0
)
or 0
)
await _sync_media_snapshot(
db, owner, owner.provider_response_json
)
if owner.remote_result_url and not owner_provider_task_id(owner):
owner.pipeline_stage = _stage(
owner, ChatGenerationPipelineStage.RESULT_READY
)
else:
from app.services.generation.ai.task_group_service import aggregate_main_task_status
await aggregate_main_task_status(db, parent_task_id=str(task.id))
owner.pipeline_stage = _stage(
owner, ChatGenerationPipelineStage.WAITING_REMOTE
)
if owner.gen_type == GenerationType.VIDEO.value:
ensure_video_poll_fields(owner, now=_now())
owner.next_poll_at = _now()
await db.commit()
await log_task_event(
owner,
event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_SUCCESS.value,
to_stage=owner.pipeline_stage,
detail=created,
)
except (RedisExecutionLockError, DatabaseRowLockBusy):
# Redis ownership is mandatory. Do not convert infrastructure
# lock loss into a business failure/refund.
try:
await db.rollback()
except Exception:
pass
raise
except Exception as exc:
try:
await db.rollback()
except Exception:
pass
# 只有仍持有 Redis 执行权时,才能把供应商异常收敛为业务失败。
await lease.ensure_owned()
owner = await load_generation_owner(
db,
owner_type=normalized_owner_type,
owner_id=task_id,
for_update=True,
)
if not owner or not is_attempt_current(owner, effective_attempt):
return
error_message = (
extract_error_message(exc, "生成任务")
if callable(extract_error_message)
else str(exc)
)
if is_image_main and isinstance(owner, ChatGenerationTask):
from sqlalchemy import select
child_result = await db.execute(
select(ChatGenerationTask.id)
.where(ChatGenerationTask.parent_task_id == owner.id)
.limit(1)
)
if child_result.scalar_one_or_none() is not None:
from app.services.generation.ai.task_group_service import (
aggregate_main_task_status,
)
await aggregate_main_task_status(
db, parent_task_id=str(owner.id)
)
else:
await mark_owner_failed_and_refund_once(
db,
owner,
error_message=error_message,
pipeline_stage=_stage(
owner, ChatGenerationPipelineStage.FAILED
),
)
else:
await mark_chat_generation_task_failed_and_refund_once(
await mark_owner_failed_and_refund_once(
db,
task=task,
owner,
error_message=error_message,
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
pipeline_stage=_stage(
owner, ChatGenerationPipelineStage.FAILED
),
)
await db.commit()
await log_task_event(task, event_type=ChatGenerationTaskEventType.TASK_FAILED.value, message=error_message)
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
await notify_chat_generation_task_finished(db, task)
await aggregate_parent_for_child(db, task)
await notify_owner_finished(db, owner)
await db.commit()
await log_task_event(
owner,
event_type=ChatGenerationTaskEventType.TASK_FAILED.value,
message=error_message,
)
return
# The provider state is committed. Enqueue failures below are
# recoverable infrastructure failures and must not trigger refunds.
await _dispatch_next_stage(
db, owner, normalized_owner_type=normalized_owner_type
)
finally:
await lease.close()
if celery_app:
@celery_app.task(name="generation.chatapi_create_generation_task", bind=True, max_retries=3, default_retry_delay=30)
def chatapi_create_generation_task(self, task_id: str):
@celery_app.task(
name="generation.chatapi_create_generation_task",
bind=True,
max_retries=3,
default_retry_delay=30,
)
def chatapi_create_generation_task(
self,
task_id: str,
owner_type: str = GenerationOwnerType.CHAT_GENERATION_TASK.value,
generation_attempt_no: int | None = None,
):
try:
return run_async(_run(task_id))
return run_async(
_run(
task_id,
owner_type=owner_type,
generation_attempt_no=generation_attempt_no,
)
)
except Exception as exc:
# 只处理 run_async/连接池/worker 中断等基础设施异常;业务异常已在 _run 内落库并退款。
retries = int(getattr(self.request, "retries", 0) or 0) + 1
countdown = int(settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30) * max(1, retries)
countdown = int(
settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30
) * max(1, retries)
raise self.retry(exc=exc, countdown=countdown)
else:
class _DisabledTask:
def delay(self, *args, **kwargs):
raise RuntimeError("Celery is disabled")
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -7,7 +7,11 @@ from typing import Any, Awaitable, Callable, Dict
from app.config import settings
from app.enums.celery_queue import CeleryQueue
from app.models.base import async_session
from app.services.redis_registry_service import get_registry_redis, redis_acquire_lock, redis_release_lock
from app.services.redis_registry_service import (
RedisExecutionLockLease,
get_registry_redis,
redis_acquire_execution_lock,
)
from app.tasks.async_runner import run_async
from app.tasks.celery_app import celery_app
@@ -17,25 +21,160 @@ RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or CeleryQueue.GEN_RECOVERY.valu
RecoveryRunner = Callable[[], Awaitable[Dict[str, Any]]]
async def _recover_generation_records_once(*, include_create: bool, include_poll: bool, include_download: bool) -> Dict[str, Any]:
from app.services.generation.pipeline.recovery_repository import find_generation_record_recovery_batch
from app.tasks.generation_create_tasks import chatapi_create_generation_task
from app.tasks.generation_poll_tasks import poll_generation_task
from app.tasks.generation_download_tasks import enqueue_download_task
counts: dict[str, Any] = {"create": 0, "poll": 0, "download": 0, "errors": []}
batch_size = max(1, int(settings.GENERATION_RECOVERY_BATCH_SIZE or 20))
cursor = None
async with async_session() as db:
while True:
batch = await find_generation_record_recovery_batch(db, limit=batch_size, cursor=cursor)
if include_create:
for ref in batch.create:
try:
chatapi_create_generation_task.apply_async(
args=[ref.owner_id],
kwargs={"owner_type": ref.owner_type, "generation_attempt_no": ref.generation_attempt_no},
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
)
counts["create"] += 1
except Exception as exc:
counts["errors"].append({"owner_id": ref.owner_id, "stage": "create", "error": str(exc)})
if include_poll:
for ref in batch.poll:
try:
poll_generation_task.apply_async(
args=[ref.owner_id],
kwargs={"owner_type": ref.owner_type, "generation_attempt_no": ref.generation_attempt_no, "force_due": False},
queue=CeleryQueue.GEN_PROVIDER_POLL.value,
)
counts["poll"] += 1
except Exception as exc:
counts["errors"].append({"owner_id": ref.owner_id, "stage": "poll", "error": str(exc)})
if include_download and batch.download:
from app.services.generation.pipeline.owner_service import load_generation_owner
for ref in batch.download:
try:
# 每条候选重新读取最新状态并只锁当前一行;前一条 commit 后
# 不继续使用批量查询得到的旧 ORM 对象。
owner = await load_generation_owner(
db,
owner_type=ref.owner_type,
owner_id=ref.owner_id,
for_update=True,
)
if (
owner is None
or int(owner.generation_attempt_no or 1)
!= int(ref.generation_attempt_no or 1)
):
await db.rollback()
continue
task_id = await enqueue_download_task(
db,
owner,
recover=True,
reason="generation_record_recovery",
)
if task_id:
counts["download"] += 1
except Exception as exc:
await db.rollback()
counts["errors"].append(
{
"owner_id": ref.owner_id,
"stage": "download",
"error": str(exc),
}
)
if batch.next_cursor is None:
break
cursor = batch.next_cursor
return counts
async def _run_download_once() -> Dict[str, Any]:
from app.services.generation.recovery_service import recover_download_tasks_once
async with async_session() as db:
return await recover_download_tasks_once(db)
chat_result = await recover_download_tasks_once(db)
record_result = await _recover_generation_records_once(include_create=False, include_poll=False, include_download=True)
return {"chat_generation_task": chat_result, "generation_record": record_result}
async def _run_generation_once() -> Dict[str, Any]:
from app.services.generation.recovery_service import recover_generation_tasks_once
async with async_session() as db:
return await recover_generation_tasks_once(db)
chat_result = await recover_generation_tasks_once(db)
record_result = await _recover_generation_records_once(include_create=True, include_poll=True, include_download=False)
return {"chat_generation_task": chat_result, "generation_record": record_result}
async def _run_due_poll_dispatch_once() -> Dict[str, Any]:
from datetime import datetime, timedelta, timezone
from sqlalchemy import select
from app.enums.generation_status import GenerationRecordPipelineStage, GenerationStatus
from app.enums.generation_task import GenerationOwnerType
from app.models.generation_record import GenerationRecord
from app.services.generation.pipeline.db_lock_service import apply_short_lock_timeout
from app.services.generation.recovery_service import dispatch_due_poll_tasks_once
from app.tasks.generation_poll_tasks import poll_generation_task
async with async_session() as db:
return await dispatch_due_poll_tasks_once(db)
chat_result = await dispatch_due_poll_tasks_once(db)
current_time = datetime.now(timezone.utc)
queue_hold_until = current_time + timedelta(seconds=int(settings.POLL_TASK_QUEUE_TIMEOUT_SECONDS or 120))
await apply_short_lock_timeout(db)
due_result = await db.execute(
select(GenerationRecord)
.where(
GenerationRecord.deleted_at.is_(None),
GenerationRecord.status == GenerationStatus.generating.value,
GenerationRecord.gen_type == "video",
GenerationRecord.pipeline_stage.in_([
GenerationRecordPipelineStage.WAITING_REMOTE.value,
GenerationRecordPipelineStage.POLLING.value,
]),
GenerationRecord.next_poll_at.is_not(None),
GenerationRecord.next_poll_at <= current_time,
)
.order_by(GenerationRecord.next_poll_at.asc(), GenerationRecord.id.asc())
.limit(int(settings.POLL_DUE_DISPATCH_BATCH_SIZE or 100))
.with_for_update(skip_locked=True)
)
owners = list(due_result.scalars().all())
dispatch_refs = []
for owner in owners:
dispatch_refs.append((str(owner.id), int(owner.generation_attempt_no or 1)))
owner.pipeline_stage = GenerationRecordPipelineStage.WAITING_REMOTE.value
owner.next_poll_at = queue_hold_until
await db.commit()
dispatched = 0
errors = []
for owner_id, attempt_no in dispatch_refs:
try:
poll_generation_task.apply_async(
args=[owner_id],
kwargs={
"owner_type": GenerationOwnerType.GENERATION_RECORD.value,
"generation_attempt_no": attempt_no,
"force_due": True,
},
queue=CeleryQueue.GEN_PROVIDER_POLL.value,
)
dispatched += 1
except Exception as exc:
errors.append({"owner_id": owner_id, "error": str(exc)})
return {"chat_generation_task": chat_result, "generation_record": {"dispatched": dispatched, "errors": errors}}
async def _run_module_async_once() -> Dict[str, Any]:
@@ -66,33 +205,22 @@ async def _run_with_execution_lock(
runner: RecoveryRunner,
ttl_seconds: int | None = None,
) -> Dict[str, Any]:
"""恢复任务执行锁。
worker_ready 的启动锁只保证只投递一次如果 broker 中残留旧消息
或者人工手动触发恢复任务仍可能并发执行这里再加执行锁避免多个
恢复扫描同时扫库抢行锁抢连接池
"""
redis = await get_registry_redis()
token: str | None = None
if redis is not None:
token = await redis_acquire_lock(
lock_key=lock_key,
ttl_seconds=int(ttl_seconds or settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600),
log_context=log_context,
)
if not token:
return {"skipped": "lock_held", "lock_key": lock_key}
else:
# Redis 不可用时仍允许 DB fallback 执行一次,避免恢复能力彻底失效。
logger.warning("恢复任务执行锁不可用,降级直接执行。context=%s", log_context)
"""恢复协调器严格依赖 Redis 执行锁,不允许 Redis 故障时无锁扫库。"""
lease = await RedisExecutionLockLease.acquire(
lock_key=lock_key,
ttl_seconds=int(ttl_seconds or settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600),
log_context=log_context,
renew_interval_seconds=max(1, int(settings.REDIS_EXECUTION_LOCK_RENEW_INTERVAL_SECONDS or 30)),
)
if lease is None:
return {"skipped": "lock_held", "lock_key": lock_key}
try:
result = await runner()
result["execution_lock"] = "lock_acquired" if token else "redis_unavailable_run_db_fallback"
await lease.ensure_owned()
result["execution_lock"] = "lock_acquired"
return result
finally:
if token:
await redis_release_lock(lock_key=lock_key, token=token, log_context=log_context)
await lease.close()
async def _is_lock_held(lock_key: str) -> bool:
@@ -126,15 +254,8 @@ async def _run_due_poll_dispatch_with_guard() -> Dict[str, Any]:
async def _acquire_download_recovery_loop_lock() -> tuple[bool, str]:
"""下载恢复循环锁。
Redis 不可用时降级为直接执行 DB fallback避免恢复能力彻底失效
Redis 可用但锁被其他 worker 持有时本轮跳过不再重复投递下一轮
"""
redis = await get_registry_redis()
if redis is None:
return True, "redis_unavailable_run_db_fallback"
token = await redis_acquire_lock(
"""下载恢复循环调度锁;Redis 不可用直接抛错,不做 DB 降级。"""
token = await redis_acquire_execution_lock(
lock_key=settings.DOWNLOAD_RECOVERY_LOOP_LOCK_KEY,
ttl_seconds=int(settings.DOWNLOAD_RECOVERY_LOOP_LOCK_TTL_SECONDS or 55),
log_context="download_recovery_loop",
@@ -1,2 +0,0 @@
# Video generation is now handled by app.services.video_queue
# This file is kept as an empty shell to avoid import errors.
+65 -14
View File
@@ -1,9 +1,11 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable
from typing import Any
from app.config import settings
from app.models.base import async_session
from app.services.redis_registry_service import RedisExecutionLockLease
from app.services.video_upscale.task_service import (
recover_video_upscale_tasks_once,
run_finalize_upscale,
@@ -16,34 +18,80 @@ from app.tasks.async_runner import run_async
from app.tasks.celery_app import celery_app
async def _run_with_execution_lock(
upscale_task_id: str,
callback: Callable[[str, str, Callable[[], Awaitable[None]]], Awaitable[None]],
) -> None:
lock_key = f"{settings.VIDEO_UPSCALE_EXECUTION_LOCK_KEY_PREFIX}:{upscale_task_id}"
lease = await RedisExecutionLockLease.acquire(
lock_key=lock_key,
ttl_seconds=max(30, int(settings.VIDEO_UPSCALE_EXECUTION_LOCK_TTL_SECONDS or 900)),
log_context="video_upscale_execution",
renew_interval_seconds=max(1, int(settings.REDIS_EXECUTION_LOCK_RENEW_INTERVAL_SECONDS or 20)),
)
if lease is None:
# 重复消息已有其他 Worker 推进,不属于业务失败。
return
try:
await callback(upscale_task_id, lease.token, lease.ensure_owned)
await lease.ensure_owned()
finally:
await lease.close()
async def _run_local(upscale_task_id: str) -> None:
async with async_session() as db:
await run_local_upscale(db, upscale_task_id)
async def _execute(task_id: str, token: str, guard: Callable[[], Awaitable[None]]) -> None:
async with async_session() as db:
await run_local_upscale(db, task_id, execution_token=token, execution_guard=guard)
await _run_with_execution_lock(upscale_task_id, _execute)
async def _run_submit(upscale_task_id: str, *, count_attempt: bool = True) -> None:
async with async_session() as db:
await run_remote_submit(db, upscale_task_id, count_attempt=count_attempt)
async def _execute(task_id: str, token: str, guard: Callable[[], Awaitable[None]]) -> None:
async with async_session() as db:
await run_remote_submit(
db, task_id, count_attempt=count_attempt, execution_token=token, execution_guard=guard
)
await _run_with_execution_lock(upscale_task_id, _execute)
async def _run_poll(upscale_task_id: str) -> None:
async with async_session() as db:
await run_remote_poll(db, upscale_task_id)
async def _execute(task_id: str, token: str, guard: Callable[[], Awaitable[None]]) -> None:
async with async_session() as db:
await run_remote_poll(db, task_id, execution_token=token, execution_guard=guard)
await _run_with_execution_lock(upscale_task_id, _execute)
async def _run_download(upscale_task_id: str) -> None:
async with async_session() as db:
await run_remote_result_download(db, upscale_task_id)
async def _execute(task_id: str, token: str, guard: Callable[[], Awaitable[None]]) -> None:
async with async_session() as db:
await run_remote_result_download(db, task_id, execution_token=token, execution_guard=guard)
await _run_with_execution_lock(upscale_task_id, _execute)
async def _run_finalize(upscale_task_id: str) -> None:
async with async_session() as db:
await run_finalize_upscale(db, upscale_task_id)
async def _execute(task_id: str, token: str, guard: Callable[[], Awaitable[None]]) -> None:
async with async_session() as db:
await run_finalize_upscale(db, task_id, execution_token=token, execution_guard=guard)
await _run_with_execution_lock(upscale_task_id, _execute)
async def _run_recovery() -> dict[str, Any]:
async with async_session() as db:
return await recover_video_upscale_tasks_once(db)
lease = await RedisExecutionLockLease.acquire(
lock_key=settings.VIDEO_UPSCALE_RECOVERY_LOCK_KEY,
ttl_seconds=max(30, int(settings.VIDEO_UPSCALE_EXECUTION_LOCK_TTL_SECONDS or 900)),
log_context="video_upscale_recovery",
renew_interval_seconds=max(1, int(settings.REDIS_EXECUTION_LOCK_RENEW_INTERVAL_SECONDS or 20)),
)
if lease is None:
return {"checked": 0, "results": {"lock_busy": 1}}
try:
async with async_session() as db:
result = await recover_video_upscale_tasks_once(db)
await lease.ensure_owned()
return result
finally:
await lease.close()
if celery_app:
@@ -88,9 +136,12 @@ if celery_app:
raise self.retry(exc=exc, countdown=max(5, int(settings.VIDEO_UPSCALE_RETRY_BACKOFF_SECONDS or 60)))
@celery_app.task(name="video_upscale.recover_once", bind=True)
@celery_app.task(name="video_upscale.recover_once", bind=True, max_retries=2)
def recover_once(self) -> dict[str, Any]:
return run_async(_run_recovery())
try:
return run_async(_run_recovery())
except Exception as exc:
raise self.retry(exc=exc, countdown=max(5, int(settings.VIDEO_UPSCALE_RETRY_BACKOFF_SECONDS or 60)))
else: