生成项目任务/chat任务软删|生成资源管控回收|生成资源token验签API预处理
This commit is contained in:
@@ -42,6 +42,7 @@ from app.services.credits import add_credits, deduct_credits
|
||||
from app.services.notification import create_notification
|
||||
from app.services.auth import hash_password, verify_password
|
||||
from app.services.operation_log import log_operation
|
||||
from app.services.resource_signed_url_service import build_resource_signed_url
|
||||
from app.utils.id_gen import generate_id
|
||||
from app.schemas.generation import GenerationType, ASPECT_RATIOS, RESOLUTIONS
|
||||
|
||||
@@ -884,9 +885,9 @@ async def get_stats(
|
||||
total_users = (await db.execute(
|
||||
select(func.count(User.id)).where(User.user_type == "frontend")
|
||||
)).scalar() or 0
|
||||
total_projects = (await db.execute(select(func.count(Project.id)))).scalar() or 0
|
||||
total_projects = (await db.execute(select(func.count(Project.id)).where(Project.deleted_at.is_(None)))).scalar() or 0
|
||||
total_generations = (
|
||||
await db.execute(select(func.count(GenerationRecord.id)))
|
||||
await db.execute(select(func.count(GenerationRecord.id)).where(GenerationRecord.deleted_at.is_(None)))
|
||||
).scalar() or 0
|
||||
total_revenue = (
|
||||
await db.execute(
|
||||
@@ -973,6 +974,7 @@ async def admin_list_generation_records(
|
||||
select(GenerationRecord, User.username, Project.name)
|
||||
.join(User, GenerationRecord.user_id == User.id)
|
||||
.join(Project, GenerationRecord.project_id == Project.id)
|
||||
.where(GenerationRecord.deleted_at.is_(None), Project.deleted_at.is_(None))
|
||||
.order_by(GenerationRecord.created_at.desc())
|
||||
)
|
||||
if user_id:
|
||||
@@ -981,7 +983,7 @@ async def admin_list_generation_records(
|
||||
query = query.where(GenerationRecord.status == status)
|
||||
|
||||
# Count total
|
||||
count_query = select(func.count(GenerationRecord.id))
|
||||
count_query = select(func.count(GenerationRecord.id)).where(GenerationRecord.deleted_at.is_(None))
|
||||
if user_id:
|
||||
count_query = count_query.where(GenerationRecord.user_id == user_id)
|
||||
if status:
|
||||
@@ -1015,7 +1017,7 @@ async def admin_list_generation_records(
|
||||
"aspect_ratio": record.aspect_ratio,
|
||||
"resolution": record.resolution,
|
||||
"status": record.status,
|
||||
"video_url": record.video_url,
|
||||
"video_url": build_resource_signed_url(record.video_url) if record.video_url else '',
|
||||
"references": refs,
|
||||
"credits_cost": record.credits_cost or 0,
|
||||
"text_credits_cost": record.text_credits_cost or 0,
|
||||
@@ -1028,7 +1030,7 @@ async def admin_list_generation_records(
|
||||
# append img param
|
||||
"gen_type": record.gen_type,
|
||||
"image_size": record.image_size or '',
|
||||
"image_url": record.image_url or '',
|
||||
"image_url": build_resource_signed_url(record.image_url) if record.image_url else '',
|
||||
"image_tokens_used": record.image_tokens_used or 0,
|
||||
"image_proportion": record.image_proportion or '',
|
||||
"image_px": record.image_px or '',
|
||||
@@ -1046,7 +1048,10 @@ async def admin_update_generation_status(
|
||||
):
|
||||
"""Admin update generation record status (e.g., confirm/reject)."""
|
||||
result = await db.execute(
|
||||
select(GenerationRecord).where(GenerationRecord.id == record_id)
|
||||
select(GenerationRecord).where(
|
||||
GenerationRecord.id == record_id,
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
record = result.scalar_one_or_none()
|
||||
if not record:
|
||||
@@ -1080,7 +1085,11 @@ async def admin_generate_video(
|
||||
result = await db.execute(
|
||||
select(GenerationRecord, Project.name)
|
||||
.join(Project, GenerationRecord.project_id == Project.id)
|
||||
.where(GenerationRecord.id == record_id)
|
||||
.where(
|
||||
GenerationRecord.id == record_id,
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
Project.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
row = result.first()
|
||||
if not row:
|
||||
|
||||
@@ -28,6 +28,11 @@ from app.schemas.generation import (
|
||||
from app.services.credits import deduct_credits, calc_text_credits, calc_video_credits, calc_image_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
|
||||
from app.services.resource_accounting_service import (
|
||||
record_generation_record_generated_resource,
|
||||
safe_file_size,
|
||||
)
|
||||
from app.services.resource_signed_url_service import build_resource_signed_url
|
||||
from app.utils.id_gen import generate_id
|
||||
from app.utils.exceptions import InsufficientCreditsError, RecordNotFoundError, InvalidStatusError
|
||||
|
||||
@@ -71,8 +76,8 @@ def _record_to_out(record: GenerationRecord, project_name: str) -> GenerationRec
|
||||
image_proportion=record.image_proportion,
|
||||
image_px=record.image_px,
|
||||
status=record.status,
|
||||
video_url=record.video_url,
|
||||
image_url=record.image_url,
|
||||
video_url=build_resource_signed_url(record.video_url) if record.video_url else '',
|
||||
image_url=build_resource_signed_url(record.image_url) if record.image_url else '',
|
||||
references=refs,
|
||||
text_credits_cost=round(record.text_credits_cost or 0.00, 2),
|
||||
# text_tokens_used=record.text_tokens_used or 0,
|
||||
@@ -94,7 +99,11 @@ async def list_records(
|
||||
query = (
|
||||
select(GenerationRecord, Project.name)
|
||||
.join(Project, GenerationRecord.project_id == Project.id)
|
||||
.where(GenerationRecord.user_id == current_user.id)
|
||||
.where(
|
||||
GenerationRecord.user_id == current_user.id,
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
Project.deleted_at.is_(None),
|
||||
)
|
||||
.order_by(GenerationRecord.created_at.desc())
|
||||
)
|
||||
if project_id:
|
||||
@@ -133,6 +142,8 @@ async def optimize(
|
||||
.join(Project, GenerationRecord.project_id == Project.id)
|
||||
.where(
|
||||
GenerationRecord.user_id == current_user.id,
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
Project.deleted_at.is_(None),
|
||||
GenerationRecord.idempotency_key == req.idempotency_key,
|
||||
GenerationRecord.gen_type == req.gen_type,
|
||||
GenerationRecord.status == "prompt_optimized",
|
||||
@@ -155,6 +166,7 @@ async def optimize(
|
||||
select(Project).where(
|
||||
Project.id == req.project_id,
|
||||
Project.user_id == current_user.id,
|
||||
Project.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
project = proj_result.scalar_one_or_none()
|
||||
@@ -243,6 +255,8 @@ async def generate(
|
||||
.where(
|
||||
GenerationRecord.id == record_id,
|
||||
GenerationRecord.user_id == current_user.id,
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
Project.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
row = result.first()
|
||||
@@ -326,6 +340,8 @@ async def retry_generation(
|
||||
.where(
|
||||
GenerationRecord.id == record_id,
|
||||
GenerationRecord.user_id == current_user.id,
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
Project.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
row = result.first()
|
||||
@@ -378,6 +394,7 @@ async def update_prompt(
|
||||
select(GenerationRecord).where(
|
||||
GenerationRecord.id == record_id,
|
||||
GenerationRecord.user_id == current_user.id,
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
record = result.scalar_one_or_none()
|
||||
@@ -422,6 +439,7 @@ async def get_queue_status(
|
||||
select(GenerationRecord).where(
|
||||
GenerationRecord.id == record_id,
|
||||
GenerationRecord.user_id == current_user.id,
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
record = result.scalar_one_or_none()
|
||||
@@ -435,6 +453,7 @@ async def get_queue_status(
|
||||
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,
|
||||
)
|
||||
)
|
||||
@@ -461,7 +480,10 @@ async def seedance_callback(request: Request, db: AsyncSession = Depends(get_db)
|
||||
return {"message": "ignored"}
|
||||
|
||||
result = await db.execute(
|
||||
select(GenerationRecord).where(GenerationRecord.seedance_task_id == task_id)
|
||||
select(GenerationRecord).where(
|
||||
GenerationRecord.seedance_task_id == task_id,
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
record = result.scalar_one_or_none()
|
||||
if not record:
|
||||
@@ -470,19 +492,33 @@ async def seedance_callback(request: Request, db: AsyncSession = Depends(get_db)
|
||||
if task_status == "succeeded":
|
||||
remote_url = data.get("content", {}).get("video_url", "")
|
||||
record.status = "completed"
|
||||
storage_path = None
|
||||
file_size_bytes = 0
|
||||
# Download video to local storage
|
||||
if settings.STORAGE_TYPE == "local" and remote_url:
|
||||
try:
|
||||
from app.services.video_gen import download_video
|
||||
dest = os.path.join(settings.STORAGE_LOCAL_PATH, f"{record.id}.mp4")
|
||||
await download_video(remote_url, dest)
|
||||
record.video_url = f"/videos/{record.id}.mp4"
|
||||
record.video_url = f"/generate/videos/{record.id}.mp4"
|
||||
storage_path = dest
|
||||
file_size_bytes = safe_file_size(dest)
|
||||
except Exception as e:
|
||||
logger.warning(f"Callback download failed, using remote URL: {e}")
|
||||
record.video_url = remote_url
|
||||
else:
|
||||
record.video_url = remote_url
|
||||
record.generated_at = datetime.now()
|
||||
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,
|
||||
)
|
||||
# Extract video token usage from callback
|
||||
usage = data.get("usage", {})
|
||||
if usage:
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException, Path, Query
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
@@ -10,6 +12,7 @@ from app.schemas.generation_ai import (
|
||||
GenerationAIHistoryDayItemsOut,
|
||||
GenerationAIHistoryGroupedOut,
|
||||
GenerationAIRetryOut,
|
||||
GenerationAITaskDeleteOut,
|
||||
GenerationAITaskCreate,
|
||||
GenerationAITaskListOut,
|
||||
GenerationAITaskOut,
|
||||
@@ -21,6 +24,7 @@ from app.services.generation_ai_service import (
|
||||
list_generation_history_day_items,
|
||||
list_generation_history_grouped_days,
|
||||
record_to_out,
|
||||
soft_delete_chat_generation_task,
|
||||
)
|
||||
from app.services.generation_log_service import log_task_event
|
||||
from app.tasks.celery_app import celery_app
|
||||
@@ -389,6 +393,7 @@ async def get_task(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.user_id == current_user.id,
|
||||
ChatGenerationTask.generation_mode == "chatapi_async",
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
task = result.scalar_one_or_none()
|
||||
@@ -397,6 +402,71 @@ async def get_task(
|
||||
return record_to_out(task)
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/tasks/{task_id}",
|
||||
response_model=GenerationAITaskDeleteOut,
|
||||
summary="删除AI生成任务",
|
||||
description=(
|
||||
"软删除当前登录用户自己的AI生成任务。"
|
||||
"该接口不会物理删除数据库记录和本地文件,只会设置 deleted_at,后续列表、详情、历史统计默认不再返回。"
|
||||
"删除已完成任务时会联动软删 generated_resources 资源账本,并重新扣减用户有效资源空间统计。"
|
||||
"如果任务仍处于 generating 生成中状态,接口会直接拦截,不允许删除。"
|
||||
),
|
||||
responses={
|
||||
200: {
|
||||
"description": "软删除成功,返回任务ID和本次释放的资源空间字节数",
|
||||
},
|
||||
400: {
|
||||
"description": "任务正在生成中,暂不能删除",
|
||||
},
|
||||
401: {
|
||||
"description": "未登录或 Token 无效",
|
||||
},
|
||||
404: {
|
||||
"description": "任务不存在,或任务不属于当前用户,或任务已经被删除",
|
||||
},
|
||||
},
|
||||
)
|
||||
async def delete_task(
|
||||
task_id: str = Path(
|
||||
...,
|
||||
description="需要删除的AI生成任务ID",
|
||||
examples=["0019e0a44895b6d837d"],
|
||||
),
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.user_id == current_user.id,
|
||||
ChatGenerationTask.generation_mode == "chatapi_async",
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
task = result.scalar_one_or_none()
|
||||
if not task:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
|
||||
if task.status == "generating":
|
||||
raise HTTPException(status_code=400, detail="当前任务正在生成中,暂不能删除")
|
||||
|
||||
deleted_at = datetime.now(timezone.utc)
|
||||
freed_size_bytes = await soft_delete_chat_generation_task(
|
||||
db,
|
||||
task=task,
|
||||
deleted_at=deleted_at,
|
||||
)
|
||||
await db.flush()
|
||||
|
||||
return GenerationAITaskDeleteOut(
|
||||
message="任务已删除",
|
||||
task_id=task.id,
|
||||
deleted=True,
|
||||
freed_size_bytes=freed_size_bytes,
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/tasks/{task_id}/retry",
|
||||
response_model=GenerationAIRetryOut,
|
||||
@@ -443,6 +513,7 @@ async def retry_task(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.user_id == current_user.id,
|
||||
ChatGenerationTask.generation_mode == "chatapi_async",
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
task = result.scalar_one_or_none()
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.dependencies import get_db, get_current_user
|
||||
@@ -7,6 +9,7 @@ from app.models.user import User
|
||||
from app.models.project import Project
|
||||
from app.models.generation_record import GenerationRecord
|
||||
from app.schemas.project import ProjectCreate, ProjectOut
|
||||
from app.services.resource_accounting_service import soft_delete_generation_record_resources
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
router = APIRouter(prefix="/projects", tags=["projects"])
|
||||
@@ -19,7 +22,10 @@ async def list_projects(
|
||||
):
|
||||
result = await db.execute(
|
||||
select(Project)
|
||||
.where(Project.user_id == current_user.id)
|
||||
.where(
|
||||
Project.user_id == current_user.id,
|
||||
Project.deleted_at.is_(None),
|
||||
)
|
||||
.order_by(Project.created_at.desc())
|
||||
)
|
||||
return result.scalars().all()
|
||||
@@ -52,17 +58,52 @@ async def delete_project(
|
||||
select(Project).where(
|
||||
Project.id == project_id,
|
||||
Project.user_id == current_user.id,
|
||||
Project.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
project = result.scalar_one_or_none()
|
||||
if not project:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="项目不存在")
|
||||
|
||||
# Cascade delete generation records
|
||||
from sqlalchemy import delete
|
||||
await db.execute(
|
||||
delete(GenerationRecord).where(GenerationRecord.project_id == project_id)
|
||||
generating_count = (
|
||||
await db.execute(
|
||||
select(func.count(GenerationRecord.id)).where(
|
||||
GenerationRecord.project_id == project_id,
|
||||
GenerationRecord.user_id == current_user.id,
|
||||
GenerationRecord.status == "generating",
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
).scalar() or 0
|
||||
if generating_count > 0:
|
||||
raise HTTPException(status_code=400, detail="当前项目下存在生成中任务,暂不能删除")
|
||||
|
||||
records_result = await db.execute(
|
||||
select(GenerationRecord).where(
|
||||
GenerationRecord.project_id == project_id,
|
||||
GenerationRecord.user_id == current_user.id,
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
await db.delete(project)
|
||||
records = list(records_result.scalars().all())
|
||||
record_ids = [record.id for record in records]
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
project.deleted_at = now
|
||||
for record in records:
|
||||
record.deleted_at = now
|
||||
|
||||
freed_size_bytes = await soft_delete_generation_record_resources(
|
||||
db,
|
||||
record_ids,
|
||||
deleted_at=now,
|
||||
)
|
||||
|
||||
await db.flush()
|
||||
return {"message": "ok"}
|
||||
return {
|
||||
"message": "ok",
|
||||
"project_id": project_id,
|
||||
"deleted": True,
|
||||
"deleted_records": len(record_ids),
|
||||
"freed_size_bytes": freed_size_bytes,
|
||||
}
|
||||
|
||||
@@ -45,8 +45,8 @@ class Settings(BaseSettings):
|
||||
PAYMENT_MOCK: bool = True
|
||||
|
||||
STORAGE_TYPE: str = "local"
|
||||
STORAGE_LOCAL_PATH: str = "./storage/videos"
|
||||
STORAGE_IMAGE_LOCAL_PATH: str = "./storage/images"
|
||||
STORAGE_LOCAL_PATH: str = "./storage/generate/videos"
|
||||
STORAGE_IMAGE_LOCAL_PATH: str = "./storage/generate/images"
|
||||
UPLOAD_LOCAL_PATH: str = "./storage/uploads"
|
||||
|
||||
|
||||
@@ -84,5 +84,10 @@ class Settings(BaseSettings):
|
||||
CELERY_DB_POOL_TIMEOUT: int = 30
|
||||
CELERY_DB_POOL_RECYCLE: int = 1800
|
||||
|
||||
RESOURCE_SIGN_SECRET: str = "resource-signature-secret-key-for-API-authentication"
|
||||
RESOURCE_SIGN_EXPIRE_SECONDS: int = 60
|
||||
RESOURCE_SIGN_ARG_EXPIRE: str = "exp"
|
||||
RESOURCE_SIGN_ARG_SIGNATURE: str = "sign"
|
||||
|
||||
|
||||
settings = Settings()
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from app.models.base import Base, TimestampMixin, engine, async_session, init_database, close_database
|
||||
from app.models.base import Base, TimestampMixin, SoftDeleteMixin, engine, async_session, init_database, close_database
|
||||
from app.models.user import User
|
||||
from app.models.project import Project
|
||||
from app.models.generation_record import GenerationRecord
|
||||
@@ -18,13 +18,17 @@ from app.models.operation_log import OperationLog
|
||||
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.generated_resource import GeneratedResource
|
||||
from app.models.user_resource_month_stat import UserResourceMonthStat
|
||||
from app.models.user_resource_total_stat import UserResourceTotalStat
|
||||
|
||||
__all__ = [
|
||||
"Base", "TimestampMixin", "engine", "async_session",
|
||||
"Base", "TimestampMixin", "SoftDeleteMixin", "engine", "async_session",
|
||||
"init_database", "close_database",
|
||||
"User", "Project", "GenerationRecord", "CreditRecord",
|
||||
"ModelConfig", "SystemConfig", "Notification", "PaymentOrder",
|
||||
"TokenUsage", "IndustryConfig", "VideoEngine", "CreditRatio",
|
||||
"MenuConfig", "RechargePackage", "OperationLog",
|
||||
"ChatGenerationTask", "ChatGenerationTaskEvent", "ChatProviderCallLog",
|
||||
"GeneratedResource", "UserResourceMonthStat", "UserResourceTotalStat",
|
||||
]
|
||||
|
||||
@@ -51,6 +51,12 @@ class TimestampMixin:
|
||||
)
|
||||
|
||||
|
||||
class SoftDeleteMixin:
|
||||
deleted_at: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True), nullable=True, index=True
|
||||
)
|
||||
|
||||
|
||||
async def init_database() -> None:
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
|
||||
@@ -3,10 +3,10 @@ from datetime import datetime
|
||||
from sqlalchemy import DateTime, Float, ForeignKey, 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 ChatGenerationTask(Base, TimestampMixin):
|
||||
class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin):
|
||||
"""Project-independent AI chat/image/video generation task.
|
||||
|
||||
This table is intentionally NOT linked to projects. It is used by the
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
from datetime import date, datetime
|
||||
|
||||
from sqlalchemy import BigInteger, Date, DateTime, ForeignKey, Index, String, Text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.models.base import Base, TimestampMixin, SoftDeleteMixin
|
||||
|
||||
|
||||
class GeneratedResource(Base, TimestampMixin, SoftDeleteMixin):
|
||||
"""统一生成资源账本。
|
||||
|
||||
只记录生成成功后的图片/视频资源,不直接绑定具体业务外键,
|
||||
通过 source_model + source_id 兼容 ChatGenerationTask、GenerationRecord 以及后续新模块。
|
||||
"""
|
||||
|
||||
__tablename__ = "generated_resources"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
user_id: Mapped[str] = mapped_column(
|
||||
String(32), ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||
)
|
||||
|
||||
resource_type: Mapped[str] = mapped_column(String(16), index=True, nullable=False) # image / video
|
||||
resource_url: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
remote_url: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
storage_type: Mapped[str] = mapped_column(String(32), default="local", nullable=False)
|
||||
storage_path: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
file_size_bytes: Mapped[int] = mapped_column(BigInteger, default=0, nullable=False)
|
||||
|
||||
source_model: Mapped[str] = mapped_column(String(64), index=True, nullable=False)
|
||||
source_model_module: Mapped[str | None] = mapped_column(String(255), nullable=True)
|
||||
source_id: Mapped[str] = mapped_column(String(32), index=True, nullable=False)
|
||||
|
||||
engine_id: Mapped[str | None] = mapped_column(String(32), nullable=True, index=True)
|
||||
engine_type: Mapped[str | None] = mapped_column(String(32), nullable=True)
|
||||
provider: Mapped[str | None] = mapped_column(String(64), nullable=True)
|
||||
model_name: Mapped[str | None] = mapped_column(String(128), nullable=True)
|
||||
|
||||
generated_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True, index=True)
|
||||
resource_month: Mapped[date] = mapped_column(Date, index=True, nullable=False)
|
||||
extra_json: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
|
||||
|
||||
Index("ix_generated_resources_user_month", GeneratedResource.user_id, GeneratedResource.resource_month)
|
||||
Index("ix_generated_resources_source", GeneratedResource.source_model, GeneratedResource.source_id)
|
||||
Index("ix_generated_resources_active_user", GeneratedResource.user_id, GeneratedResource.deleted_at)
|
||||
@@ -3,10 +3,10 @@ from datetime import datetime
|
||||
from sqlalchemy import DateTime, ForeignKey, Integer, String, Text, Float
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.models.base import Base, TimestampMixin
|
||||
from app.models.base import Base, TimestampMixin, SoftDeleteMixin
|
||||
|
||||
|
||||
class GenerationRecord(Base, TimestampMixin):
|
||||
class GenerationRecord(Base, TimestampMixin, SoftDeleteMixin):
|
||||
__tablename__ = "generation_records"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
from sqlalchemy import ForeignKey, String
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.models.base import Base, TimestampMixin
|
||||
from app.models.base import Base, TimestampMixin, SoftDeleteMixin
|
||||
|
||||
|
||||
class Project(Base, TimestampMixin):
|
||||
class Project(Base, TimestampMixin, SoftDeleteMixin):
|
||||
__tablename__ = "projects"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
from datetime import date, datetime
|
||||
|
||||
from sqlalchemy import BigInteger, Date, DateTime, ForeignKey, Integer, String, UniqueConstraint
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.models.base import Base, TimestampMixin
|
||||
|
||||
|
||||
class UserResourceMonthStat(Base, TimestampMixin):
|
||||
"""用户月份资源空间聚合表。"""
|
||||
|
||||
__tablename__ = "user_resource_month_stats"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("user_id", "stat_month", name="uq_user_resource_month_stats_user_month"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
user_id: Mapped[str] = mapped_column(
|
||||
String(32), ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||
)
|
||||
stat_month: Mapped[date] = mapped_column(Date, index=True, nullable=False)
|
||||
|
||||
active_size_bytes: Mapped[int] = mapped_column(BigInteger, default=0, nullable=False)
|
||||
deleted_size_bytes: Mapped[int] = mapped_column(BigInteger, default=0, nullable=False)
|
||||
total_generated_size_bytes: Mapped[int] = mapped_column(BigInteger, default=0, nullable=False)
|
||||
|
||||
image_size_bytes: Mapped[int] = mapped_column(BigInteger, default=0, nullable=False)
|
||||
video_size_bytes: Mapped[int] = mapped_column(BigInteger, default=0, nullable=False)
|
||||
|
||||
active_count: Mapped[int] = mapped_column(Integer, default=0, nullable=False)
|
||||
deleted_count: Mapped[int] = mapped_column(Integer, default=0, nullable=False)
|
||||
image_count: Mapped[int] = mapped_column(Integer, default=0, nullable=False)
|
||||
video_count: Mapped[int] = mapped_column(Integer, default=0, nullable=False)
|
||||
|
||||
last_recalculated_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
@@ -0,0 +1,34 @@
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import BigInteger, DateTime, ForeignKey, Integer, String, UniqueConstraint
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.models.base import Base, TimestampMixin
|
||||
|
||||
|
||||
class UserResourceTotalStat(Base, TimestampMixin):
|
||||
"""用户全局资源空间聚合表。"""
|
||||
|
||||
__tablename__ = "user_resource_total_stats"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("user_id", name="uq_user_resource_total_stats_user"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
user_id: Mapped[str] = mapped_column(
|
||||
String(32), ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||
)
|
||||
|
||||
active_size_bytes: Mapped[int] = mapped_column(BigInteger, default=0, nullable=False)
|
||||
deleted_size_bytes: Mapped[int] = mapped_column(BigInteger, default=0, nullable=False)
|
||||
total_generated_size_bytes: Mapped[int] = mapped_column(BigInteger, default=0, nullable=False)
|
||||
|
||||
image_size_bytes: Mapped[int] = mapped_column(BigInteger, default=0, nullable=False)
|
||||
video_size_bytes: Mapped[int] = mapped_column(BigInteger, default=0, nullable=False)
|
||||
|
||||
active_count: Mapped[int] = mapped_column(Integer, default=0, nullable=False)
|
||||
deleted_count: Mapped[int] = mapped_column(Integer, default=0, nullable=False)
|
||||
image_count: Mapped[int] = mapped_column(Integer, default=0, nullable=False)
|
||||
video_count: Mapped[int] = mapped_column(Integer, default=0, nullable=False)
|
||||
|
||||
last_recalculated_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
@@ -416,6 +416,28 @@ class GenerationAITaskListOut(BaseModel):
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
class GenerationAITaskDeleteOut(BaseModel):
|
||||
"""AI生成任务删除响应体。"""
|
||||
|
||||
model_config = ConfigDict(
|
||||
json_schema_extra={
|
||||
"example": {
|
||||
"message": "任务已删除",
|
||||
"task_id": "0019e0a44895b6d837d",
|
||||
"deleted": True,
|
||||
"freed_size_bytes": 123456,
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
message: str = Field(..., description="操作结果提示信息")
|
||||
task_id: str = Field(..., description="被软删除的AI生成任务ID")
|
||||
deleted: bool = Field(..., description="是否已完成软删除")
|
||||
freed_size_bytes: int = Field(0, description="本次软删联动释放的有效资源空间字节数")
|
||||
|
||||
|
||||
class GenerationAIRetryOut(BaseModel):
|
||||
"""AI生成任务重试响应体。"""
|
||||
|
||||
|
||||
@@ -25,6 +25,8 @@ from app.schemas.generation_ai import (
|
||||
GenerationAIVideoEngineOptionOut,
|
||||
)
|
||||
from app.services.generation_billing_service import charge_generation_media_by_params
|
||||
from app.services.resource_accounting_service import soft_delete_chat_task_resources
|
||||
from app.services.resource_signed_url_service import build_resource_signed_url
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
IMAGE_DEFAULT_SIZE = "2K"
|
||||
@@ -201,6 +203,7 @@ async def create_async_generation_task(db: AsyncSession, current_user: User, req
|
||||
ChatGenerationTask.user_id == current_user.id,
|
||||
ChatGenerationTask.idempotency_key == req.idempotency_key,
|
||||
ChatGenerationTask.generation_mode == "chatapi_async",
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
).order_by(ChatGenerationTask.created_at.desc()).limit(1)
|
||||
)
|
||||
existing = result.scalar_one_or_none()
|
||||
@@ -330,8 +333,8 @@ def record_to_out(task: ChatGenerationTask) -> GenerationAITaskOut:
|
||||
provider_task_id=task.provider_task_id,
|
||||
seedance_task_id=task.seedance_task_id,
|
||||
# remote_result_url=task.remote_result_url,
|
||||
image_url=task.image_url,
|
||||
video_url=task.video_url,
|
||||
image_url=build_resource_signed_url(task.image_url) if task.image_url else "",
|
||||
video_url=build_resource_signed_url(task.video_url) if task.video_url else "",
|
||||
engine_id=task.engine_id,
|
||||
engine_snapshot=snapshot,
|
||||
credits_cost=task.credits_cost or 0.0,
|
||||
@@ -381,6 +384,7 @@ async def list_async_generation_tasks(
|
||||
query = select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.generation_mode == "chatapi_async",
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
if gen_type:
|
||||
query = query.where(ChatGenerationTask.gen_type == gen_type)
|
||||
@@ -436,6 +440,7 @@ def _history_base_filters(user_id: str, gen_type: str):
|
||||
return [
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.generation_mode == "chatapi_async",
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
ChatGenerationTask.status == "completed",
|
||||
ChatGenerationTask.gen_type == gen_type,
|
||||
ChatGenerationTask.generated_at.is_not(None),
|
||||
@@ -445,6 +450,7 @@ def _history_base_filters(user_id: str, gen_type: str):
|
||||
def _generation_record_history_base_filters(user_id: str, gen_type: str):
|
||||
return [
|
||||
GenerationRecord.user_id == user_id,
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
GenerationRecord.status == "completed",
|
||||
GenerationRecord.gen_type == gen_type,
|
||||
GenerationRecord.generated_at.is_not(None),
|
||||
@@ -478,8 +484,8 @@ def generation_record_to_history_out(
|
||||
provider_task_id=record.seedance_task_id,
|
||||
seedance_task_id=record.seedance_task_id,
|
||||
remote_result_url=None,
|
||||
image_url=record.image_url,
|
||||
video_url=record.video_url,
|
||||
image_url=build_resource_signed_url(record.image_url) if record.image_url else '',
|
||||
video_url=build_resource_signed_url(record.video_url) if record.video_url else '',
|
||||
engine_id=None,
|
||||
engine_snapshot=None,
|
||||
credits_cost=record.credits_cost or 0.0,
|
||||
@@ -544,7 +550,7 @@ async def list_generation_record_history_grouped_days(
|
||||
for generated_day, day_total in day_rows:
|
||||
item_result = await db.execute(
|
||||
select(GenerationRecord, Project.name.label("project_name"))
|
||||
.outerjoin(Project, GenerationRecord.project_id == Project.id)
|
||||
.outerjoin(Project, (GenerationRecord.project_id == Project.id) & (Project.deleted_at.is_(None)))
|
||||
.where(
|
||||
*filters,
|
||||
func.date(GenerationRecord.generated_at) == generated_day,
|
||||
@@ -603,7 +609,7 @@ async def list_generation_record_history_day_items(
|
||||
|
||||
result = await db.execute(
|
||||
select(GenerationRecord, Project.name.label("project_name"))
|
||||
.outerjoin(Project, GenerationRecord.project_id == Project.id)
|
||||
.outerjoin(Project, (GenerationRecord.project_id == Project.id) & (Project.deleted_at.is_(None)))
|
||||
.where(
|
||||
*filters,
|
||||
day_expr == target_day,
|
||||
@@ -773,4 +779,15 @@ async def list_generation_history_day_items(
|
||||
"page": page,
|
||||
"page_size": page_size,
|
||||
"items": [record_to_out(task) for task in tasks],
|
||||
}
|
||||
}
|
||||
|
||||
async def soft_delete_chat_generation_task(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
task: ChatGenerationTask,
|
||||
deleted_at: datetime | None = None,
|
||||
) -> int:
|
||||
"""软删 ChatGenerationTask 并联动软删资源账本,返回释放的 active 空间字节数。"""
|
||||
deleted_at = deleted_at or datetime.now(timezone.utc)
|
||||
task.deleted_at = deleted_at
|
||||
return await soft_delete_chat_task_resources(db, task.id, deleted_at=deleted_at)
|
||||
|
||||
@@ -1,16 +1,27 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
|
||||
from app.config import settings
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
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
|
||||
from app.services.video_gen import download_video
|
||||
|
||||
|
||||
async def download_generation_result(record: ChatGenerationTask) -> str:
|
||||
@dataclass(slots=True)
|
||||
class DownloadedGenerationResult:
|
||||
url: str
|
||||
storage_path: str | None
|
||||
file_size_bytes: int
|
||||
resource_type: str
|
||||
storage_type: str = "local"
|
||||
|
||||
|
||||
async def download_generation_result(record: ChatGenerationTask) -> DownloadedGenerationResult:
|
||||
if not record.remote_result_url:
|
||||
raise ValueError("缺少远程结果URL")
|
||||
|
||||
@@ -21,11 +32,21 @@ async def download_generation_result(record: ChatGenerationTask) -> str:
|
||||
dest = os.path.join(dest_dir, f"{record.id}.png")
|
||||
async with provider_limit("result_download", settings.RESULT_DOWNLOAD_MAX_CONCURRENCY):
|
||||
await download_image(record.remote_result_url, dest)
|
||||
return f"/images/{date_dir}/{record.id}.png"
|
||||
return DownloadedGenerationResult(
|
||||
url=f"/generate/images/{date_dir}/{record.id}.png",
|
||||
storage_path=dest,
|
||||
file_size_bytes=safe_file_size(dest),
|
||||
resource_type="image",
|
||||
)
|
||||
|
||||
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")
|
||||
async with provider_limit("result_download", settings.RESULT_DOWNLOAD_MAX_CONCURRENCY):
|
||||
await download_video(record.remote_result_url, dest)
|
||||
return f"/videos/{date_dir}/{record.id}.mp4"
|
||||
return DownloadedGenerationResult(
|
||||
url=f"/generate/videos/{date_dir}/{record.id}.mp4",
|
||||
storage_path=dest,
|
||||
file_size_bytes=safe_file_size(dest),
|
||||
resource_type="video",
|
||||
)
|
||||
|
||||
@@ -0,0 +1,401 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from datetime import date, datetime, timezone
|
||||
from typing import Any, Iterable, Sequence
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.models.generated_resource import GeneratedResource
|
||||
from app.models.generation_record import GenerationRecord
|
||||
from app.models.user_resource_month_stat import UserResourceMonthStat
|
||||
from app.models.user_resource_total_stat import UserResourceTotalStat
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
SOURCE_MODEL_CHAT_TASK = "ChatGenerationTask"
|
||||
SOURCE_MODEL_GENERATION_RECORD = "GenerationRecord"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ResourceAccountingResult:
|
||||
resource_id: str
|
||||
resource_type: str
|
||||
file_size_bytes: int
|
||||
active_size_delta: int
|
||||
|
||||
|
||||
def resource_month_from_datetime(value: datetime | None = None) -> date:
|
||||
value = value or datetime.now(timezone.utc)
|
||||
return date(value.year, value.month, 1)
|
||||
|
||||
|
||||
def safe_file_size(path: str | None) -> int:
|
||||
if not path:
|
||||
return 0
|
||||
try:
|
||||
return int(os.path.getsize(path))
|
||||
except OSError:
|
||||
return 0
|
||||
|
||||
|
||||
def _json(data: Any) -> str | None:
|
||||
if data is None:
|
||||
return None
|
||||
if isinstance(data, str):
|
||||
return data
|
||||
return json.dumps(data, ensure_ascii=False, default=str)
|
||||
|
||||
|
||||
def _parse_json(text: str | None) -> dict:
|
||||
if not text:
|
||||
return {}
|
||||
try:
|
||||
data = json.loads(text)
|
||||
return data if isinstance(data, dict) else {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def _int(value: int | None) -> int:
|
||||
return int(value or 0)
|
||||
|
||||
|
||||
def _add_non_negative(obj: Any, field: str, delta: int) -> None:
|
||||
current = _int(getattr(obj, field, 0))
|
||||
setattr(obj, field, max(0, current + int(delta or 0)))
|
||||
|
||||
|
||||
def _add_raw(obj: Any, field: str, delta: int) -> None:
|
||||
current = _int(getattr(obj, field, 0))
|
||||
setattr(obj, field, current + int(delta or 0))
|
||||
|
||||
|
||||
async def _get_or_create_month_stat(
|
||||
db: AsyncSession,
|
||||
user_id: str,
|
||||
stat_month: date,
|
||||
) -> UserResourceMonthStat:
|
||||
result = await db.execute(
|
||||
select(UserResourceMonthStat).where(
|
||||
UserResourceMonthStat.user_id == user_id,
|
||||
UserResourceMonthStat.stat_month == stat_month,
|
||||
)
|
||||
)
|
||||
stat = result.scalar_one_or_none()
|
||||
if stat:
|
||||
return stat
|
||||
|
||||
stat = UserResourceMonthStat(
|
||||
id=generate_id(),
|
||||
user_id=user_id,
|
||||
stat_month=stat_month,
|
||||
)
|
||||
db.add(stat)
|
||||
await db.flush()
|
||||
return stat
|
||||
|
||||
|
||||
async def _get_or_create_total_stat(
|
||||
db: AsyncSession,
|
||||
user_id: str,
|
||||
) -> UserResourceTotalStat:
|
||||
result = await db.execute(
|
||||
select(UserResourceTotalStat).where(UserResourceTotalStat.user_id == user_id)
|
||||
)
|
||||
stat = result.scalar_one_or_none()
|
||||
if stat:
|
||||
return stat
|
||||
|
||||
stat = UserResourceTotalStat(
|
||||
id=generate_id(),
|
||||
user_id=user_id,
|
||||
)
|
||||
db.add(stat)
|
||||
await db.flush()
|
||||
return stat
|
||||
|
||||
|
||||
async def apply_resource_stat_delta(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
user_id: str,
|
||||
stat_month: date,
|
||||
resource_type: str,
|
||||
active_size_delta: int = 0,
|
||||
active_count_delta: int = 0,
|
||||
deleted_size_delta: int = 0,
|
||||
deleted_count_delta: int = 0,
|
||||
total_generated_size_delta: int = 0,
|
||||
) -> None:
|
||||
month_stat = await _get_or_create_month_stat(db, user_id, stat_month)
|
||||
total_stat = await _get_or_create_total_stat(db, user_id)
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
for stat in (month_stat, total_stat):
|
||||
_add_non_negative(stat, "active_size_bytes", active_size_delta)
|
||||
_add_non_negative(stat, "active_count", active_count_delta)
|
||||
_add_non_negative(stat, "deleted_size_bytes", deleted_size_delta)
|
||||
_add_non_negative(stat, "deleted_count", deleted_count_delta)
|
||||
_add_raw(stat, "total_generated_size_bytes", total_generated_size_delta)
|
||||
|
||||
if resource_type == "image":
|
||||
_add_non_negative(stat, "image_size_bytes", active_size_delta)
|
||||
_add_non_negative(stat, "image_count", active_count_delta)
|
||||
elif resource_type == "video":
|
||||
_add_non_negative(stat, "video_size_bytes", active_size_delta)
|
||||
_add_non_negative(stat, "video_count", active_count_delta)
|
||||
|
||||
stat.last_recalculated_at = now
|
||||
|
||||
|
||||
async def record_generated_resource(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
user_id: str,
|
||||
resource_type: str,
|
||||
resource_url: str,
|
||||
source_model: str,
|
||||
source_id: str,
|
||||
source_model_module: str | None = None,
|
||||
remote_url: str | None = None,
|
||||
storage_type: str = "local",
|
||||
storage_path: str | None = None,
|
||||
file_size_bytes: int | None = None,
|
||||
engine_id: str | None = None,
|
||||
engine_type: str | None = None,
|
||||
provider: str | None = None,
|
||||
model_name: str | None = None,
|
||||
generated_at: datetime | None = None,
|
||||
extra: Any = None,
|
||||
) -> ResourceAccountingResult:
|
||||
"""记录生成成功资源,并增量维护用户月份/全局空间统计。
|
||||
|
||||
file_size_bytes 获取不到时允许为 0,符合当前确认方案。
|
||||
同一个 source_model + source_id + resource_type 已存在未软删账本时,更新账本并按大小差值修正统计。
|
||||
"""
|
||||
resource_type = (resource_type or "").lower().strip()
|
||||
if resource_type not in ("image", "video"):
|
||||
raise ValueError("resource_type 仅支持 image 或 video")
|
||||
if not resource_url:
|
||||
raise ValueError("resource_url 不能为空")
|
||||
|
||||
generated_at = generated_at or datetime.now(timezone.utc)
|
||||
stat_month = resource_month_from_datetime(generated_at)
|
||||
size = int(file_size_bytes if file_size_bytes is not None else safe_file_size(storage_path))
|
||||
if size < 0:
|
||||
size = 0
|
||||
|
||||
result = await db.execute(
|
||||
select(GeneratedResource).where(
|
||||
GeneratedResource.source_model == source_model,
|
||||
GeneratedResource.source_id == source_id,
|
||||
GeneratedResource.resource_type == resource_type,
|
||||
GeneratedResource.deleted_at.is_(None),
|
||||
).order_by(GeneratedResource.created_at.desc()).limit(1)
|
||||
)
|
||||
existing = result.scalar_one_or_none()
|
||||
|
||||
if existing:
|
||||
old_size = _int(existing.file_size_bytes)
|
||||
active_size_delta = size - old_size
|
||||
existing.user_id = user_id
|
||||
existing.resource_url = resource_url
|
||||
existing.remote_url = remote_url
|
||||
existing.storage_type = storage_type or existing.storage_type or "local"
|
||||
existing.storage_path = storage_path
|
||||
existing.file_size_bytes = size
|
||||
existing.engine_id = engine_id
|
||||
existing.engine_type = engine_type
|
||||
existing.provider = provider
|
||||
existing.model_name = model_name
|
||||
existing.generated_at = generated_at
|
||||
existing.resource_month = stat_month
|
||||
existing.extra_json = _json(extra)
|
||||
|
||||
if active_size_delta:
|
||||
await apply_resource_stat_delta(
|
||||
db,
|
||||
user_id=user_id,
|
||||
stat_month=stat_month,
|
||||
resource_type=resource_type,
|
||||
active_size_delta=active_size_delta,
|
||||
)
|
||||
|
||||
return ResourceAccountingResult(
|
||||
resource_id=existing.id,
|
||||
resource_type=resource_type,
|
||||
file_size_bytes=size,
|
||||
active_size_delta=active_size_delta,
|
||||
)
|
||||
|
||||
resource = GeneratedResource(
|
||||
id=generate_id(),
|
||||
user_id=user_id,
|
||||
resource_type=resource_type,
|
||||
resource_url=resource_url,
|
||||
remote_url=remote_url,
|
||||
storage_type=storage_type or "local",
|
||||
storage_path=storage_path,
|
||||
file_size_bytes=size,
|
||||
source_model=source_model,
|
||||
source_model_module=source_model_module,
|
||||
source_id=source_id,
|
||||
engine_id=engine_id,
|
||||
engine_type=engine_type,
|
||||
provider=provider,
|
||||
model_name=model_name,
|
||||
generated_at=generated_at,
|
||||
resource_month=stat_month,
|
||||
extra_json=_json(extra),
|
||||
)
|
||||
db.add(resource)
|
||||
await db.flush()
|
||||
|
||||
await apply_resource_stat_delta(
|
||||
db,
|
||||
user_id=user_id,
|
||||
stat_month=stat_month,
|
||||
resource_type=resource_type,
|
||||
active_size_delta=size,
|
||||
active_count_delta=1,
|
||||
total_generated_size_delta=size,
|
||||
)
|
||||
|
||||
return ResourceAccountingResult(
|
||||
resource_id=resource.id,
|
||||
resource_type=resource_type,
|
||||
file_size_bytes=size,
|
||||
active_size_delta=size,
|
||||
)
|
||||
|
||||
|
||||
async def record_chat_task_generated_resource(
|
||||
db: AsyncSession,
|
||||
task: ChatGenerationTask,
|
||||
*,
|
||||
resource_url: str,
|
||||
storage_path: str | None = None,
|
||||
file_size_bytes: int | None = None,
|
||||
remote_url: str | None = None,
|
||||
generated_at: datetime | None = None,
|
||||
) -> ResourceAccountingResult:
|
||||
snapshot = _parse_json(task.engine_snapshot_json)
|
||||
return await record_generated_resource(
|
||||
db,
|
||||
user_id=task.user_id,
|
||||
resource_type=task.gen_type,
|
||||
resource_url=resource_url,
|
||||
remote_url=remote_url or task.remote_result_url,
|
||||
storage_type="local" if storage_path else "remote",
|
||||
storage_path=storage_path,
|
||||
file_size_bytes=file_size_bytes,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_model_module="app.models.chat_generation_task",
|
||||
source_id=task.id,
|
||||
engine_id=task.engine_id,
|
||||
engine_type=snapshot.get("engine_type") or task.gen_type,
|
||||
provider=snapshot.get("provider"),
|
||||
model_name=snapshot.get("model_name"),
|
||||
generated_at=generated_at or task.generated_at or datetime.now(timezone.utc),
|
||||
extra={"pipeline_stage": task.pipeline_stage},
|
||||
)
|
||||
|
||||
|
||||
async def record_generation_record_generated_resource(
|
||||
db: AsyncSession,
|
||||
record: GenerationRecord,
|
||||
*,
|
||||
resource_url: str,
|
||||
storage_path: str | None = None,
|
||||
file_size_bytes: int | None = None,
|
||||
remote_url: str | None = None,
|
||||
generated_at: datetime | None = None,
|
||||
) -> ResourceAccountingResult:
|
||||
return await record_generated_resource(
|
||||
db,
|
||||
user_id=record.user_id,
|
||||
resource_type=record.gen_type,
|
||||
resource_url=resource_url,
|
||||
remote_url=remote_url,
|
||||
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,
|
||||
generated_at=generated_at or record.generated_at or datetime.now(timezone.utc),
|
||||
extra={"project_id": record.project_id},
|
||||
)
|
||||
|
||||
|
||||
async def soft_delete_resources_by_source(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
source_model: str,
|
||||
source_ids: Sequence[str] | Iterable[str],
|
||||
deleted_at: datetime | None = None,
|
||||
) -> int:
|
||||
"""按业务来源软删资源账本,并返回本次释放的 active 空间字节数。"""
|
||||
ids = [item for item in source_ids if item]
|
||||
if not ids:
|
||||
return 0
|
||||
|
||||
deleted_at = deleted_at or datetime.now(timezone.utc)
|
||||
result = await db.execute(
|
||||
select(GeneratedResource).where(
|
||||
GeneratedResource.source_model == source_model,
|
||||
GeneratedResource.source_id.in_(ids),
|
||||
GeneratedResource.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
resources = list(result.scalars().all())
|
||||
freed_size = 0
|
||||
|
||||
for resource in resources:
|
||||
size = _int(resource.file_size_bytes)
|
||||
resource.deleted_at = deleted_at
|
||||
freed_size += size
|
||||
await apply_resource_stat_delta(
|
||||
db,
|
||||
user_id=resource.user_id,
|
||||
stat_month=resource.resource_month,
|
||||
resource_type=resource.resource_type,
|
||||
active_size_delta=-size,
|
||||
active_count_delta=-1,
|
||||
deleted_size_delta=size,
|
||||
deleted_count_delta=1,
|
||||
)
|
||||
|
||||
return freed_size
|
||||
|
||||
|
||||
async def soft_delete_chat_task_resources(
|
||||
db: AsyncSession,
|
||||
task_id: str,
|
||||
*,
|
||||
deleted_at: datetime | None = None,
|
||||
) -> int:
|
||||
return await soft_delete_resources_by_source(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=[task_id],
|
||||
deleted_at=deleted_at,
|
||||
)
|
||||
|
||||
|
||||
async def soft_delete_generation_record_resources(
|
||||
db: AsyncSession,
|
||||
record_ids: Sequence[str] | Iterable[str],
|
||||
*,
|
||||
deleted_at: datetime | None = None,
|
||||
) -> int:
|
||||
return await soft_delete_resources_by_source(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_GENERATION_RECORD,
|
||||
source_ids=record_ids,
|
||||
deleted_at=deleted_at,
|
||||
)
|
||||
@@ -0,0 +1,265 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import hmac
|
||||
import time
|
||||
from typing import Any, Optional
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
from app.config import settings
|
||||
|
||||
|
||||
class ResourceSignedUrlError(RuntimeError):
|
||||
"""
|
||||
资源签名 URL 生成异常。
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
def _to_int(value: Any, default: int) -> int:
|
||||
"""
|
||||
安全转换整数。
|
||||
"""
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
|
||||
|
||||
def _get_sign_secret(secret: Optional[str] = None) -> str:
|
||||
"""
|
||||
获取资源签名密钥。
|
||||
|
||||
优先级:
|
||||
1. 函数传入 secret
|
||||
2. app.config.settings.RESOURCE_SIGN_SECRET
|
||||
"""
|
||||
sign_secret = secret or getattr(settings, "RESOURCE_SIGN_SECRET", "")
|
||||
|
||||
if not sign_secret or not str(sign_secret).strip():
|
||||
raise ResourceSignedUrlError("RESOURCE_SIGN_SECRET 未配置")
|
||||
|
||||
return str(sign_secret)
|
||||
|
||||
|
||||
def _get_sign_expire_seconds(expire_seconds: Optional[int] = None) -> int:
|
||||
"""
|
||||
获取资源签名有效期秒数。
|
||||
"""
|
||||
if expire_seconds is not None:
|
||||
seconds = _to_int(expire_seconds, 3600)
|
||||
else:
|
||||
seconds = _to_int(getattr(settings, "RESOURCE_SIGN_EXPIRE_SECONDS", 3600), 3600)
|
||||
|
||||
if seconds <= 0:
|
||||
seconds = 3600
|
||||
|
||||
return seconds
|
||||
|
||||
|
||||
def _get_expire_arg_name() -> str:
|
||||
"""
|
||||
获取过期时间参数名。
|
||||
默认 exp。
|
||||
"""
|
||||
name = getattr(settings, "RESOURCE_SIGN_ARG_EXPIRE", "exp")
|
||||
name = str(name or "exp").strip()
|
||||
return name or "exp"
|
||||
|
||||
|
||||
def _get_signature_arg_name() -> str:
|
||||
"""
|
||||
获取签名参数名。
|
||||
默认 sign。
|
||||
"""
|
||||
name = getattr(settings, "RESOURCE_SIGN_ARG_SIGNATURE", "sign")
|
||||
name = str(name or "sign").strip()
|
||||
return name or "sign"
|
||||
|
||||
|
||||
def _extract_sign_uri(resource_url: str) -> str:
|
||||
"""
|
||||
提取用于签名的 URI path。
|
||||
|
||||
例如:
|
||||
http://www.test6.com/generation/video/a.mp4?x=1
|
||||
|
||||
用于签名的是:
|
||||
/generation/video/a.mp4
|
||||
|
||||
注意:
|
||||
OpenResty/Lua 侧建议使用 ngx.var.uri 或 r.uri 参与签名,
|
||||
不要使用完整 URL,也不要包含 query string。
|
||||
"""
|
||||
if not resource_url or not str(resource_url).strip():
|
||||
raise ResourceSignedUrlError("resource_url 不能为空")
|
||||
|
||||
url = str(resource_url).strip()
|
||||
parsed = urlsplit(url)
|
||||
|
||||
sign_uri = parsed.path or url
|
||||
|
||||
if not sign_uri.startswith("/"):
|
||||
sign_uri = "/" + sign_uri
|
||||
|
||||
return sign_uri
|
||||
|
||||
|
||||
def generate_resource_signature(
|
||||
resource_url: str,
|
||||
expires_at: int,
|
||||
secret: Optional[str] = None,
|
||||
) -> str:
|
||||
"""
|
||||
生成资源 URL 签名。
|
||||
|
||||
签名规则:
|
||||
|
||||
message = "{uri}:{exp}"
|
||||
sign = hmac_sha256(secret, message).hexdigest()
|
||||
|
||||
例如:
|
||||
|
||||
uri = "/generation/video/a.mp4"
|
||||
exp = 1780000000
|
||||
message = "/generation/video/a.mp4:1780000000"
|
||||
|
||||
OpenResty/Lua 侧必须使用完全一致的 message 规则。
|
||||
"""
|
||||
sign_secret = _get_sign_secret(secret)
|
||||
sign_uri = _extract_sign_uri(resource_url)
|
||||
|
||||
expires_at = _to_int(expires_at, 0)
|
||||
if expires_at <= 0:
|
||||
raise ResourceSignedUrlError("expires_at 必须是有效的 Unix 时间戳")
|
||||
|
||||
message = f"generation_resource_controller:{sign_uri}:{expires_at}"
|
||||
|
||||
return hmac.new(
|
||||
sign_secret.encode("utf-8"),
|
||||
message.encode("utf-8"),
|
||||
hashlib.sha256,
|
||||
).hexdigest()
|
||||
|
||||
|
||||
def _append_query_params(resource_url: str, params: dict[str, Any]) -> str:
|
||||
"""
|
||||
按用户要求追加 URL 参数:
|
||||
|
||||
- 原 URL 有 ?,用 ¶m=value 追加
|
||||
- 原 URL 没有 ?,用 ?param=value 追加
|
||||
|
||||
同时兼容 #fragment,参数会追加到 # 前面。
|
||||
"""
|
||||
url = str(resource_url).strip()
|
||||
|
||||
if not url:
|
||||
raise ResourceSignedUrlError("resource_url 不能为空")
|
||||
|
||||
base_url = url
|
||||
fragment = ""
|
||||
|
||||
if "#" in url:
|
||||
base_url, fragment_part = url.split("#", 1)
|
||||
fragment = "#" + fragment_part
|
||||
|
||||
separator = "&" if "?" in base_url else "?"
|
||||
|
||||
query_string = "&".join(
|
||||
f"{key}={value}"
|
||||
for key, value in params.items()
|
||||
if key and value is not None
|
||||
)
|
||||
|
||||
if not query_string:
|
||||
return url
|
||||
|
||||
return f"{base_url}{separator}{query_string}{fragment}"
|
||||
|
||||
|
||||
def build_resource_signed_url(
|
||||
resource_url: str | None,
|
||||
expire_seconds: Optional[int] = None,
|
||||
secret: Optional[str] = None,
|
||||
now_ts: Optional[int] = None,
|
||||
) -> str:
|
||||
"""
|
||||
生成带时效签名的资源 URL。
|
||||
|
||||
参数:
|
||||
resource_url:
|
||||
外部传入的资源 URL,可以是完整 URL,也可以是 path。
|
||||
|
||||
例如:
|
||||
http://www.test6.com/generation/video/a.mp4
|
||||
http://www.test6.com/generation/video/a.mp4?from=history
|
||||
/generation/video/a.mp4
|
||||
|
||||
expire_seconds:
|
||||
有效期秒数,不传则使用 settings.RESOURCE_SIGN_EXPIRE_SECONDS。
|
||||
|
||||
secret:
|
||||
可选,自定义签名密钥。不传则使用 settings.RESOURCE_SIGN_SECRET。
|
||||
|
||||
now_ts:
|
||||
可选,当前时间戳。主要用于单元测试,正常业务不需要传。
|
||||
|
||||
返回:
|
||||
带 exp 和 sign 参数的 URL。
|
||||
|
||||
示例:
|
||||
http://www.test6.com/generation/video/a.mp4?exp=1780000000&sign=xxxx
|
||||
|
||||
http://www.test6.com/generation/video/a.mp4?from=history&exp=1780000000&sign=xxxx
|
||||
"""
|
||||
if not resource_url:
|
||||
return resource_url
|
||||
|
||||
seconds = _get_sign_expire_seconds(expire_seconds)
|
||||
current_ts = _to_int(now_ts, int(time.time())) if now_ts is not None else int(time.time())
|
||||
|
||||
expires_at = current_ts + seconds
|
||||
|
||||
expire_arg_name = _get_expire_arg_name()
|
||||
signature_arg_name = _get_signature_arg_name()
|
||||
|
||||
signature = generate_resource_signature(
|
||||
resource_url=resource_url,
|
||||
expires_at=expires_at,
|
||||
secret=secret,
|
||||
)
|
||||
|
||||
return _append_query_params(
|
||||
resource_url=resource_url,
|
||||
params={
|
||||
expire_arg_name: expires_at,
|
||||
signature_arg_name: signature,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def build_resource_signed_urls(
|
||||
resource_urls: list[str],
|
||||
expire_seconds: Optional[int] = None,
|
||||
secret: Optional[str] = None,
|
||||
) -> list[str]:
|
||||
"""
|
||||
批量生成资源签名 URL。
|
||||
|
||||
用于历史记录列表、资源列表等场景。
|
||||
"""
|
||||
if not resource_urls:
|
||||
return []
|
||||
|
||||
now_ts = int(time.time())
|
||||
|
||||
return [
|
||||
build_resource_signed_url(
|
||||
resource_url=url,
|
||||
expire_seconds=expire_seconds,
|
||||
secret=secret,
|
||||
now_ts=now_ts,
|
||||
)
|
||||
for url in resource_urls
|
||||
if url
|
||||
]
|
||||
@@ -9,7 +9,11 @@ from sqlalchemy import select
|
||||
from app.models.base import async_session
|
||||
from app.models.generation_record import GenerationRecord
|
||||
from app.services.video_gen import get_active_engine, poll_task_status, download_video, _log_video_response
|
||||
from app.services.image_gen import get_active_image_engine, poll_image_task_status, download_image
|
||||
from app.services.image_gen import get_active_image_engine, download_image
|
||||
from app.services.resource_accounting_service import (
|
||||
record_generation_record_generated_resource,
|
||||
safe_file_size,
|
||||
)
|
||||
from app.config import settings
|
||||
|
||||
logger = logging.getLogger("videogen")
|
||||
@@ -35,6 +39,7 @@ class TaskQueue:
|
||||
select(GenerationRecord).where(
|
||||
GenerationRecord.status == "generating",
|
||||
GenerationRecord.seedance_task_id.isnot(None),
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
records = result.scalars().all()
|
||||
@@ -66,7 +71,10 @@ class TaskQueue:
|
||||
"""Process a single record: poll status and update DB."""
|
||||
async with async_session() as db:
|
||||
result = await db.execute(
|
||||
select(GenerationRecord).where(GenerationRecord.id == record_id)
|
||||
select(GenerationRecord).where(
|
||||
GenerationRecord.id == record_id,
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
record = result.scalar_one_or_none()
|
||||
if not record or record.status != "generating":
|
||||
@@ -114,6 +122,8 @@ class TaskQueue:
|
||||
|
||||
if status == "succeeded":
|
||||
file_url = poll_result.get("video_url", "")
|
||||
storage_path = None
|
||||
file_size_bytes = 0
|
||||
if settings.STORAGE_TYPE == "local" and file_url:
|
||||
try:
|
||||
date_dir = datetime.now().strftime("%Y/%m/%d")
|
||||
@@ -121,7 +131,9 @@ class TaskQueue:
|
||||
os.makedirs(dest_dir, exist_ok=True)
|
||||
dest = os.path.join(dest_dir, f"{record_id}.mp4")
|
||||
await download_video(file_url, dest)
|
||||
record.video_url = f"/videos/{date_dir}/{record_id}.mp4"
|
||||
record.video_url = f"/generate/videos/{date_dir}/{record_id}.mp4"
|
||||
storage_path = dest
|
||||
file_size_bytes = safe_file_size(dest)
|
||||
except Exception as e:
|
||||
logger.warning(f"Download failed, using remote URL: {e}")
|
||||
record.video_url = file_url
|
||||
@@ -130,6 +142,16 @@ class TaskQueue:
|
||||
record.video_tokens_used = poll_result.get("video_tokens", 0)
|
||||
record.status = "completed"
|
||||
record.generated_at = datetime.now()
|
||||
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=file_url,
|
||||
generated_at=record.generated_at,
|
||||
)
|
||||
self._active.pop(record_id, None)
|
||||
await db.commit()
|
||||
logger.info(f"Video task completed: {record_id}")
|
||||
@@ -158,29 +180,44 @@ class TaskQueue:
|
||||
async def _process_image(self, db, record):
|
||||
"""Process image generation task - calls API directly."""
|
||||
record_id = record.id
|
||||
from app.services.image_gen import submit_image_task, download_image, _log_image_response
|
||||
from app.services.image_gen import submit_image_task, _log_image_response
|
||||
|
||||
try:
|
||||
engine = await get_active_image_engine(db)
|
||||
poll_result = await asyncio.to_thread(submit_image_task, db, engine, record)
|
||||
|
||||
if poll_result["error"] == "":
|
||||
if settings.STORAGE_TYPE == "local" and poll_result.get("image_url"):
|
||||
remote_url = poll_result.get("image_url")
|
||||
storage_path = None
|
||||
file_size_bytes = 0
|
||||
if settings.STORAGE_TYPE == "local" and remote_url:
|
||||
try:
|
||||
date_dir = datetime.now().strftime("%Y/%m/%d")
|
||||
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")
|
||||
await download_image(poll_result.get("image_url"), dest)
|
||||
record.image_url = f"/images/{date_dir}/{record_id}.png"
|
||||
await download_image(remote_url, dest)
|
||||
record.image_url = f"/generate/images/{date_dir}/{record_id}.png"
|
||||
storage_path = dest
|
||||
file_size_bytes = safe_file_size(dest)
|
||||
except Exception as e:
|
||||
logger.warning(f"Download failed, using remote URL: {e}")
|
||||
record.image_url = poll_result.get("image_url")
|
||||
record.image_url = remote_url
|
||||
else:
|
||||
record.image_url = poll_result.get("image_url")
|
||||
record.image_url = remote_url
|
||||
record.image_tokens_used = poll_result.get("image_tokens", 0)
|
||||
record.status = "completed"
|
||||
record.generated_at = datetime.now()
|
||||
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(f"Image task completed: {record_id}")
|
||||
else:
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import time
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import select
|
||||
@@ -10,6 +9,8 @@ from app.utils.security import encrypt_temp_token, decrypt_temp_token
|
||||
|
||||
async def generate_temp_url(db: AsyncSession, record: GenerationRecord) -> str:
|
||||
"""Generate a temporary encrypted URL for video access (1 hour expiry)."""
|
||||
if getattr(record, "deleted_at", None) is not None:
|
||||
return ""
|
||||
token = encrypt_temp_token(record.id, expires_in=3600)
|
||||
record.video_url_expires_at = datetime.now().replace(second=0, microsecond=0)
|
||||
# We store just the token, the full URL is constructed by the frontend
|
||||
@@ -24,6 +25,9 @@ async def validate_and_get_record_id(token: str) -> str | None:
|
||||
async def get_video_stream_url(db: AsyncSession, record_id: str) -> str | None:
|
||||
"""Get the actual video URL for a record (for proxying/redirecting)."""
|
||||
result = await db.execute(
|
||||
select(GenerationRecord.video_url).where(GenerationRecord.id == record_id)
|
||||
select(GenerationRecord.video_url).where(
|
||||
GenerationRecord.id == record_id,
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
@@ -22,6 +22,7 @@ async def _cleanup_urls():
|
||||
.where(
|
||||
GenerationRecord.video_url_expires_at.isnot(None),
|
||||
GenerationRecord.video_url_expires_at < now,
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
)
|
||||
.values(video_url_expires_at=None)
|
||||
)
|
||||
|
||||
@@ -106,7 +106,10 @@ def _build_optimized_prompt_by_params(task: ChatGenerationTask) -> str:
|
||||
|
||||
async def _run(task_id: str):
|
||||
async with async_session() as db:
|
||||
result = await db.execute(select(ChatGenerationTask).where(ChatGenerationTask.id == task_id))
|
||||
result = await db.execute(select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
))
|
||||
task = result.scalar_one_or_none()
|
||||
|
||||
if not task or task.generation_mode != "chatapi_async":
|
||||
@@ -230,7 +233,10 @@ async def _run(task_id: str):
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
result = await db.execute(select(ChatGenerationTask).where(ChatGenerationTask.id == task_id))
|
||||
result = await db.execute(select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
))
|
||||
task = result.scalar_one_or_none()
|
||||
|
||||
if task:
|
||||
|
||||
@@ -8,6 +8,7 @@ from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.services.error_codes import extract_error_message
|
||||
from app.services.generation_download_service import download_generation_result
|
||||
from app.services.generation_log_service import log_task_event
|
||||
from app.services.resource_accounting_service import record_chat_task_generated_resource
|
||||
from app.tasks.celery_app import celery_app
|
||||
|
||||
|
||||
@@ -57,14 +58,22 @@ async def _reload_task(db, task_id: str) -> ChatGenerationTask | None:
|
||||
- 继续访问旧 task 有概率触发异步懒加载异常。
|
||||
"""
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask).where(ChatGenerationTask.id == task_id)
|
||||
select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
|
||||
async def _run(task_id: str):
|
||||
async with async_session() as db:
|
||||
result = await db.execute(select(ChatGenerationTask).where(ChatGenerationTask.id == task_id))
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
task = result.scalar_one_or_none()
|
||||
if not task or task.generation_mode != "chatapi_async":
|
||||
return
|
||||
@@ -104,17 +113,28 @@ async def _run(task_id: str):
|
||||
to_stage="downloading",
|
||||
)
|
||||
|
||||
local_url = await download_generation_result(task)
|
||||
downloaded = await download_generation_result(task)
|
||||
|
||||
if task.gen_type == "image":
|
||||
task.image_url = local_url
|
||||
task.image_url = downloaded.url
|
||||
else:
|
||||
task.video_url = local_url
|
||||
task.video_url = downloaded.url
|
||||
|
||||
task.status = "completed"
|
||||
task.pipeline_stage = "done"
|
||||
task.generated_at = datetime.now(timezone.utc)
|
||||
task.retry_count = 0
|
||||
|
||||
await record_chat_task_generated_resource(
|
||||
db,
|
||||
task,
|
||||
resource_url=downloaded.url,
|
||||
storage_path=downloaded.storage_path,
|
||||
file_size_bytes=downloaded.file_size_bytes,
|
||||
remote_url=task.remote_result_url,
|
||||
generated_at=task.generated_at,
|
||||
)
|
||||
|
||||
await db.commit()
|
||||
|
||||
await log_task_event(
|
||||
@@ -122,6 +142,10 @@ async def _run(task_id: str):
|
||||
event_type="DOWNLOAD_SUCCESS",
|
||||
to_status="completed",
|
||||
to_stage="done",
|
||||
detail={
|
||||
"resource_url": downloaded.url,
|
||||
"file_size_bytes": downloaded.file_size_bytes,
|
||||
},
|
||||
)
|
||||
|
||||
except Exception as exc:
|
||||
@@ -173,4 +197,4 @@ else:
|
||||
def apply_async(self, *args, **kwargs):
|
||||
raise RuntimeError("Celery is disabled")
|
||||
|
||||
download_generation_result_task = _DisabledTask()
|
||||
download_generation_result_task = _DisabledTask()
|
||||
|
||||
@@ -38,14 +38,20 @@ async def _reload_task(db, task_id: str) -> ChatGenerationTask | None:
|
||||
- 所以 poll/download 的异常分支统一 rollback 后重新 select。
|
||||
"""
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask).where(ChatGenerationTask.id == task_id)
|
||||
select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
|
||||
async def _run(task_id: str):
|
||||
async with async_session() as db:
|
||||
result = await db.execute(select(ChatGenerationTask).where(ChatGenerationTask.id == task_id))
|
||||
result = await db.execute(select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
))
|
||||
task = result.scalar_one_or_none()
|
||||
if not task or task.generation_mode != "chatapi_async":
|
||||
return
|
||||
|
||||
Reference in New Issue
Block a user