火山引擎SMS API|celery容灾优化|生成模型引擎积分列表API

This commit is contained in:
2026-06-04 16:28:59 +08:00
parent 450e509b1a
commit 8d43d5871c
25 changed files with 1936 additions and 302 deletions
+8 -1
View File
@@ -54,4 +54,11 @@ VIDEO_COVER_SEEK_TIME=00:00:01
VIDEO_COVER_FALLBACK_SEEK_TIME=00:00:00
VIDEO_COVER_WIDTH=600
VIDEO_COVER_TIMEOUT_SECONDS=15
VIDEO_COVER_FORMAT=png
VIDEO_COVER_FORMAT=png
# VOLC_SMS
VOLC_SMS_ACCESS_KEY_ID=AKLTYWY5Yjc5YjM3N2IwNDc3M2I3NTU2YjlmNTczYzQzMmM
VOLC_SMS_SECRET_ACCESS_KEY=TXpjM01HUTFZMlV5TUdKbE5Ea3lNRGhqTUdSak16UTFOV0ptTW1SaE5XRQ==
VOLC_SMS_ACCOUNT=8b3f6ca8
VOLC_SMS_TEMPLATE_ID=S1T_1y2pb9ej2g9vm
VOLC_SMS_SIGN=纬佳网络科技
@@ -0,0 +1,159 @@
"""add sms password setup and celery startup recovery
Revision ID: 476b259992de
Revises: 516b84e3bf17
Create Date: 2026-06-04 14:25:03.284184
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = '476b259992de'
down_revision: Union[str, None] = '516b84e3bf17'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
# ### chat_generation_tasks 下载容灾字段 ###
op.add_column(
"chat_generation_tasks",
sa.Column("download_celery_task_id", sa.String(length=160), nullable=True),
)
op.add_column(
"chat_generation_tasks",
sa.Column("download_enqueued_at", sa.DateTime(timezone=True), nullable=True),
)
op.add_column(
"chat_generation_tasks",
sa.Column("download_started_at", sa.DateTime(timezone=True), nullable=True),
)
op.add_column(
"chat_generation_tasks",
sa.Column("download_lease_until", sa.DateTime(timezone=True), nullable=True),
)
op.add_column(
"chat_generation_tasks",
sa.Column("download_next_retry_at", sa.DateTime(timezone=True), nullable=True),
)
# 重点:已有数据时,新增 NOT NULL 字段必须给默认值,否则 PostgreSQL 可能直接失败。
op.add_column(
"chat_generation_tasks",
sa.Column(
"download_attempt_count",
sa.Integer(),
nullable=False,
server_default=sa.text("0"),
),
)
op.add_column(
"chat_generation_tasks",
sa.Column("download_last_error", sa.Text(), nullable=True),
)
op.add_column(
"chat_generation_tasks",
sa.Column("download_storage_date_dir", sa.String(length=16), nullable=True),
)
op.create_index(
op.f("ix_chat_generation_tasks_download_celery_task_id"),
"chat_generation_tasks",
["download_celery_task_id"],
unique=False,
)
op.create_index(
op.f("ix_chat_generation_tasks_download_lease_until"),
"chat_generation_tasks",
["download_lease_until"],
unique=False,
)
op.create_index(
op.f("ix_chat_generation_tasks_download_next_retry_at"),
"chat_generation_tasks",
["download_next_retry_at"],
unique=False,
)
# 可选:如果你不想 DB 层长期保留 server_default,可以取消下面注释。
# 但我建议保留 DB 默认值,避免后续手写 SQL 插入时 download_attempt_count 为空。
#
# op.alter_column(
# "chat_generation_tasks",
# "download_attempt_count",
# server_default=None,
# existing_type=sa.Integer(),
# existing_nullable=False,
# )
# ### users 短信注册 / 设置密码字段 ###
op.add_column(
"users",
sa.Column("password_set_at", sa.DateTime(timezone=True), nullable=True),
)
# 老用户已有密码,回填 password_set_at,避免被误判成未设置密码。
op.execute(
"""
UPDATE users
SET password_set_at = CURRENT_TIMESTAMP
WHERE hashed_password IS NOT NULL
AND hashed_password <> ''
"""
)
# 短信注册用户允许暂时没有密码。
op.alter_column(
"users",
"hashed_password",
existing_type=sa.VARCHAR(length=128),
nullable=True,
)
def downgrade() -> None:
# 如果已经有短信注册但未设置密码的用户,直接改回 NOT NULL 会失败。
# 这里写一个不可登录的占位值,只是保证 downgrade 能执行。
op.execute(
"""
UPDATE users
SET hashed_password = 'PASSWORD_NOT_SET_AFTER_DOWNGRADE'
WHERE hashed_password IS NULL
OR hashed_password = ''
"""
)
op.alter_column(
"users",
"hashed_password",
existing_type=sa.VARCHAR(length=128),
nullable=False,
)
op.drop_column("users", "password_set_at")
op.drop_index(
op.f("ix_chat_generation_tasks_download_next_retry_at"),
table_name="chat_generation_tasks",
)
op.drop_index(
op.f("ix_chat_generation_tasks_download_lease_until"),
table_name="chat_generation_tasks",
)
op.drop_index(
op.f("ix_chat_generation_tasks_download_celery_task_id"),
table_name="chat_generation_tasks",
)
op.drop_column("chat_generation_tasks", "download_storage_date_dir")
op.drop_column("chat_generation_tasks", "download_last_error")
op.drop_column("chat_generation_tasks", "download_attempt_count")
op.drop_column("chat_generation_tasks", "download_next_retry_at")
op.drop_column("chat_generation_tasks", "download_lease_until")
op.drop_column("chat_generation_tasks", "download_started_at")
op.drop_column("chat_generation_tasks", "download_enqueued_at")
op.drop_column("chat_generation_tasks", "download_celery_task_id")
+171 -89
View File
@@ -1,112 +1,70 @@
from datetime import datetime, timezone
from datetime import datetime
from fastapi import APIRouter, Depends
from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.dependencies import get_db, get_current_user
from app.dependencies import (
get_current_user,
get_current_user_allow_password_pending,
get_db,
)
from app.models.system_config import SystemConfig
from app.models.user import User
from app.schemas.auth import LoginRequest, ChangePasswordRequest, RegisterRequest
from app.schemas.auth import (
ChangePasswordRequest,
LoginRequest,
RegisterRequest,
SetPasswordRequest,
SmsLoginRequest,
)
from app.schemas.user import UserOut
from app.services.auth import (
authenticate_user,
create_access_token,
decode_access_token,
get_user_by_phone,
hash_password,
verify_password,
)
from app.services.sms import verify_sms_code
from app.utils.id_gen import generate_id
router = APIRouter(prefix="/auth", tags=["auth"])
@router.post("/login")
async def login(req: LoginRequest, db: AsyncSession = Depends(get_db)):
# Validate captcha in production (when SMS_MOCK is false)
if not settings.SMS_MOCK:
if not req.captcha_token:
from fastapi import HTTPException, status
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="需要验证码",
)
token_sub = decode_access_token(req.captcha_token)
if not token_sub or not token_sub.startswith("captcha:"):
from fastapi import HTTPException, status
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="验证码无效或已过期",
)
user = await authenticate_user(db, req.username, req.password)
if not user:
from fastapi import HTTPException, status
def _validate_captcha_if_needed(captcha_token: str | None) -> None:
# 保留原密码登录的图形验证码逻辑,不改成短信验证码。
if settings.SMS_MOCK:
return
if not captcha_token:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="用户名或密码错误",
status_code=status.HTTP_400_BAD_REQUEST,
detail="需要验证码",
)
token_sub = decode_access_token(captcha_token)
if not token_sub or not token_sub.startswith("captcha:"):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="验证码无效或已过期",
)
# Only allow frontend users to login via this endpoint
if user.user_type != "frontend":
from fastapi import HTTPException, status
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="该账号不允许在此登录",
)
user.last_login_at = datetime.now()
def _token_response(user: User, remember_me: bool = False) -> dict:
user.credits = round(user.credits, 2)
await db.flush()
token = create_access_token(user.id, req.remember_me)
return {"access_token": token, "token_type": "bearer", "user": UserOut.model_validate(user)}
token = create_access_token(user.id, remember_me)
return {
"access_token": token,
"token_type": "bearer",
"user": UserOut.model_validate(user),
"must_set_password": bool(user.user_type == "frontend" and not user.hashed_password),
}
@router.post("/register")
async def register(req: RegisterRequest, db: AsyncSession = Depends(get_db)):
"""Register a new user with phone + SMS code + password."""
# Verify SMS code
from app.services.sms import verify_sms_code
ok = await verify_sms_code(req.phone, req.code)
if not ok:
from fastapi import HTTPException, status
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="验证码错误或已过期",
)
# Check if phone already registered
existing = await db.execute(select(User).where(User.phone == req.phone).limit(1))
if existing.scalar_one_or_none():
from fastapi import HTTPException, status
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="该手机号已注册",
)
# Create user with default username: 用户 + last 4 digits of phone
import random
username = f"用户{req.phone[-4:]}"
existing_name = await db.execute(select(User).where(User.username == username).limit(1))
if existing_name.scalar_one_or_none():
username = f"用户{req.phone[-4:]}{random.randint(10, 99)}"
user = User(
id=generate_id(),
username=username,
phone=req.phone,
hashed_password=hash_password(req.password),
credits=100,
is_admin=False,
user_type="frontend",
)
db.add(user)
await db.flush()
# Assign default menus to new user
async def _assign_default_frontend_menus(db: AsyncSession, user: User) -> None:
from app.models.menu_config import MenuConfig
result = await db.execute(
select(MenuConfig).where(
MenuConfig.is_default == True,
@@ -119,22 +77,149 @@ async def register(req: RegisterRequest, db: AsyncSession = Depends(get_db)):
if default_menus:
user.allowed_menus = [m.path for m in default_menus if m.path]
@router.post(
"/login",
summary="客户端密码登录",
description="保留原有用户名/手机号 + 密码登录。仅允许 frontend 用户登录;管理员仍使用 /auth/admin-login。",
)
async def login(req: LoginRequest, db: AsyncSession = Depends(get_db)):
_validate_captcha_if_needed(req.captcha_token)
user = await authenticate_user(db, req.username, req.password)
if not user:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="用户名或密码错误",
)
# Only allow frontend users to login via this endpoint
if user.user_type != "frontend":
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="该账号不允许在此登录",
)
user.last_login_at = datetime.now()
await db.flush()
return _token_response(user, req.remember_me)
@router.post(
"/sms-login",
summary="客户端短信验证码登录",
description="新增兼容登录方式:手机号 + 短信验证码登录。不覆盖 /auth/login 密码登录。仅允许 frontend 用户登录。",
)
async def sms_login(req: SmsLoginRequest, db: AsyncSession = Depends(get_db)):
ok = await verify_sms_code(req.phone, req.code, "login")
if not ok:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="验证码错误或已过期",
)
user = await get_user_by_phone(db, req.phone)
if not user or not user.is_active:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="账号不存在或已禁用",
)
if user.user_type != "frontend":
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="该账号不允许在此登录",
)
user.last_login_at = datetime.now()
await db.flush()
return _token_response(user, req.remember_me)
@router.post(
"/register",
summary="客户端手机号短信注册",
description="手机号 + 注册短信验证码注册。注册成功后 username 默认等于手机号,不生成密码;前端需根据 must_set_password 引导用户设置密码。",
)
async def register(req: RegisterRequest, db: AsyncSession = Depends(get_db)):
ok = await verify_sms_code(req.phone, req.code, "register")
if not ok:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="验证码错误或已过期",
)
existing_phone = await db.execute(select(User).where(User.phone == req.phone).limit(1))
if existing_phone.scalar_one_or_none():
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="该手机号已注册",
)
existing_username = await db.execute(select(User).where(User.username == req.phone).limit(1))
if existing_username.scalar_one_or_none():
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="该手机号已注册",
)
user = User(
id=generate_id(),
username=req.phone,
phone=req.phone,
hashed_password=None,
password_set_at=None,
credits=100,
is_admin=False,
user_type="frontend",
)
db.add(user)
await db.flush()
await _assign_default_frontend_menus(db, user)
user.credits = round(user.credits, 2)
token = create_access_token(user.id)
return {"access_token": token, "token_type": "bearer", "user": UserOut.model_validate(user)}
return {
"access_token": token,
"token_type": "bearer",
"user": UserOut.model_validate(user),
"must_set_password": True,
}
@router.post("/logout")
async def logout(current_user: User = Depends(get_current_user)):
async def logout(current_user: User = Depends(get_current_user_allow_password_pending)):
return {"message": "ok"}
@router.get("/me", response_model=UserOut)
async def get_me(current_user: User = Depends(get_current_user)):
async def get_me(current_user: User = Depends(get_current_user_allow_password_pending)):
current_user.credits = round(current_user.credits, 2)
return current_user
@router.post(
"/set-password",
summary="设置登录密码",
description="短信注册或短信登录后,用户没有密码时调用该接口设置密码。该接口允许未设置密码用户访问。",
)
async def set_password(
req: SetPasswordRequest,
current_user: User = Depends(get_current_user_allow_password_pending),
db: AsyncSession = Depends(get_db),
):
if len(req.new_password) < 6:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="密码至少6位",
)
current_user.hashed_password = hash_password(req.new_password)
current_user.password_set_at = datetime.now()
await db.flush()
return {"message": "密码设置成功", "must_set_password": False}
@router.post("/change-password")
async def change_password(
req: ChangePasswordRequest,
@@ -142,20 +227,19 @@ async def change_password(
db: AsyncSession = Depends(get_db),
):
if not verify_password(req.old_password, current_user.hashed_password):
from fastapi import HTTPException, status
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="原密码错误",
)
if len(req.new_password) < 6:
from fastapi import HTTPException, status
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="密码至少6位",
)
current_user.hashed_password = hash_password(req.new_password)
current_user.password_set_at = datetime.now()
await db.flush()
return {"message": "密码修改成功"}
@@ -183,22 +267,20 @@ async def admin_login(req: LoginRequest, db: AsyncSession = Depends(get_db)):
"""Admin-only login endpoint."""
user = await authenticate_user(db, req.username, req.password)
if not user:
from fastapi import HTTPException, status
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="用户名或密码错误",
)
if user.user_type != "admin":
from fastapi import HTTPException, status
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="该账号不是管理员账号",
)
user.last_login_at = datetime.now()
user.credits = round(user.credits, 2)
await db.flush()
token = create_access_token(user.id, req.remember_me)
user.credits = round(user.credits, 2)
return {"access_token": token, "token_type": "bearer", "user": UserOut.model_validate(user)}
+14
View File
@@ -9,6 +9,7 @@ from app.models.video_engine import VideoEngine
from app.models.image_engine import ImageEngine
from app.schemas.credit import CreditBalanceOut, CreditRecordOut
from app.schemas.credit_ratio import CreditRatioOut
from app.services.credit_ratio_service import list_all_credit_ratios
from app.services.credits import get_records
router = APIRouter(prefix="/credits", tags=["credits"])
@@ -26,6 +27,19 @@ async def get_credits(
)
@router.get(
"/credit-ratios",
response_model=list[CreditRatioOut],
summary="获取积分比例列表",
description="客户端获取当前系统配置的积分计费规则列表。普通登录用户可访问,只读返回 credit_ratios 表中的图片/视频积分比例配置。",
)
async def list_client_credit_ratios(
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
return await list_all_credit_ratios(db)
@router.get("/ratios", response_model=dict)
async def get_credit_ratios(
current_user: User = Depends(get_current_user),
+42 -20
View File
@@ -4,27 +4,45 @@ from app.config import settings
from app.schemas.sms import SmsSendRequest, SmsVerifyRequest, SmsResponse
from app.services.sms import generate_and_send_sms, verify_sms_code
router = APIRouter(prefix="/sms", tags=["sms"])
router = APIRouter(prefix="/sms", tags=["短信验证码"])
@router.post("/send", response_model=SmsResponse)
def _validate_captcha_token(captcha_token: str | None) -> None:
if settings.SMS_MOCK:
return
if not (req_captcha := captcha_token):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="需要验证码",
)
from app.services.auth import decode_access_token
token_sub = decode_access_token(req_captcha)
if not token_sub or not token_sub.startswith("captcha:"):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="验证码无效",
)
@router.post(
"/send",
response_model=SmsResponse,
summary="发送短信验证码",
description="客户端发送短信验证码。scene=register 用于注册,scene=login 用于短信登录,scene=set_password 用于设置密码。正式环境会按配置校验图形验证码。",
)
async def send_sms_code(req: SmsSendRequest):
"""Send SMS verification code. Requires captcha_token in production."""
if not settings.SMS_MOCK:
if not req.captcha_token:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="需要验证码",
)
from app.services.auth import decode_access_token
token_sub = decode_access_token(req.captcha_token)
if not token_sub or not token_sub.startswith("captcha:"):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="验证码无效",
)
_validate_captcha_token(req.captcha_token)
try:
ok = await generate_and_send_sms(req.phone, req.scene.value)
except ValueError as exc:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=str(exc),
) from exc
ok = await generate_and_send_sms(req.phone)
if not ok:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
@@ -33,10 +51,14 @@ async def send_sms_code(req: SmsSendRequest):
return SmsResponse(message="验证码已发送", success=True)
@router.post("/verify", response_model=SmsResponse)
@router.post(
"/verify",
response_model=SmsResponse,
summary="校验短信验证码",
description="校验指定手机号、场景下的短信验证码。业务接口一般会内部校验,本接口主要用于前端调试或单独校验。",
)
async def verify_sms(req: SmsVerifyRequest):
"""Verify SMS code."""
ok = await verify_sms_code(req.phone, req.code)
ok = await verify_sms_code(req.phone, req.code, req.scene.value)
if not ok:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
+27
View File
@@ -29,12 +29,26 @@ class Settings(BaseSettings):
ENCRYPTION_KEY: str = "changeme-32bytes-base64-key-here!!"
# SMS settings
# 兼容旧字段:SMS_API_URL/SMS_API_KEY/SMS_SIGN_NAME/SMS_TEMPLATE_CODE 保留,
# 新接入默认使用火山引擎短信 SDK。
SMS_API_URL: str = ""
SMS_API_KEY: str = ""
SMS_SIGN_NAME: str = "VideoGen"
SMS_TEMPLATE_CODE: str = "SMS_001"
SMS_MOCK: bool = True
# 火山引擎短信配置。
# SmsAccount=消息组IDTemplateID=模板IDSign=短信签名内容。
VOLC_SMS_ACCESS_KEY_ID: str = ""
VOLC_SMS_SECRET_ACCESS_KEY: str = ""
VOLC_SMS_ACCOUNT: str = ""
VOLC_SMS_TEMPLATE_ID: str = ""
VOLC_SMS_SIGN: str = ""
SMS_CODE_LENGTH: int = 4
SMS_CODE_TTL_SECONDS: int = 300
SMS_SEND_INTERVAL_SECONDS: int = 60
SMS_DAILY_LIMIT: int = 20
# Payment settings
WECHAT_MCH_ID: str = ""
WECHAT_API_KEY: str = ""
@@ -95,6 +109,19 @@ class Settings(BaseSettings):
CELERY_DB_POOL_TIMEOUT: int = 30
CELERY_DB_POOL_RECYCLE: int = 1800
# Celery 图片/视频下载容灾配置。
# Redis broker 下 priority=0 最高,priority=9 最低。
DOWNLOAD_TASK_PRIORITY_NORMAL: int = 5
DOWNLOAD_TASK_PRIORITY_RECOVER: int = 0
DOWNLOAD_TASK_MAX_ATTEMPTS: int = 3
DOWNLOAD_TASK_RETRY_BACKOFF_SECONDS: int = 30
DOWNLOAD_TASK_LEASE_SECONDS: int = 10 * 60
DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS: int = 5 * 60
DOWNLOAD_RECOVERY_BATCH_SIZE: int = 100
DOWNLOAD_RECOVERY_STARTUP_DELAY_SECONDS: int = 3
DOWNLOAD_ACTIVE_REDIS_HASH_KEY: str = "vg:celery:download:active"
DOWNLOAD_ACTIVE_REDIS_ZSET_KEY: str = "vg:celery:download:active_index"
RESOURCE_SIGN_SECRET: str = "resource-signature-secret-key-for-API-authentication"
RESOURCE_SIGN_EXPIRE_SECONDS: int = 60
RESOURCE_SIGN_ARG_EXPIRE: str = "exp"
+18 -4
View File
@@ -1,11 +1,11 @@
from fastapi import Depends, HTTPException, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.base import async_session
from app.models.user import User
from app.services.auth import decode_access_token
from sqlalchemy import select
from app.services.auth import decode_access_token, user_must_set_password
security = HTTPBearer(auto_error=False)
@@ -22,7 +22,7 @@ async def get_db():
await session.close()
async def get_current_user(
async def get_current_user_allow_password_pending(
credentials: HTTPAuthorizationCredentials | None = Depends(security),
db: AsyncSession = Depends(get_db),
) -> User:
@@ -56,8 +56,22 @@ async def get_current_user(
return user
async def get_current_user(
current_user: User = Depends(get_current_user_allow_password_pending),
) -> User:
if user_must_set_password(current_user):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={
"code": "PASSWORD_REQUIRED",
"message": "请先设置登录密码",
},
)
return current_user
async def get_admin_user(
current_user: User = Depends(get_current_user),
current_user: User = Depends(get_current_user_allow_password_pending),
) -> User:
if not current_user.is_admin or current_user.user_type != "admin":
raise HTTPException(
@@ -69,3 +69,13 @@ class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin):
generated_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
error_message: Mapped[str | None] = mapped_column(Text, nullable=True)
idempotency_key: Mapped[str | None] = mapped_column(String(64), nullable=True, index=True)
# Celery 下载容灾字段。
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_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)
download_last_error: Mapped[str | None] = mapped_column(Text, nullable=True)
download_storage_date_dir: Mapped[str | None] = mapped_column(String(16), nullable=True)
+9 -1
View File
@@ -13,7 +13,8 @@ class User(Base, TimestampMixin):
username: Mapped[str] = mapped_column(String(64), unique=True, index=True)
email: Mapped[str | None] = mapped_column(String(255), unique=True, nullable=True)
phone: Mapped[str | None] = mapped_column(String(20), unique=True, nullable=True)
hashed_password: Mapped[str] = mapped_column(String(128))
# 短信注册用户允许先没有密码,后续通过 /auth/set-password 设置。
hashed_password: Mapped[str | None] = mapped_column(String(128), nullable=True)
avatar: Mapped[str | None] = mapped_column(String(512), nullable=True)
credits: Mapped[float] = mapped_column(Float, default=0.0)
is_active: Mapped[bool] = mapped_column(Boolean, default=True)
@@ -22,4 +23,11 @@ class User(Base, TimestampMixin):
last_login_at: Mapped[datetime | None] = mapped_column(
DateTime(timezone=True), nullable=True
)
password_set_at: Mapped[datetime | None] = mapped_column(
DateTime(timezone=True), nullable=True
)
allowed_menus: Mapped[list | None] = mapped_column(JSON, nullable=True)
@property
def must_set_password(self) -> bool:
return self.user_type == "frontend" and not self.hashed_password
+19 -10
View File
@@ -1,22 +1,31 @@
from pydantic import BaseModel
from pydantic import BaseModel, Field
class LoginRequest(BaseModel):
username: str
password: str
captcha_token: str | None = None
remember_me: bool = False
username: str = Field(..., description="用户名或手机号")
password: str = Field(..., description="登录密码")
captcha_token: str | None = Field(None, description="图形验证码令牌,按现有配置需要时传入")
remember_me: bool = Field(False, description="是否记住登录")
class SmsLoginRequest(BaseModel):
phone: str = Field(..., pattern=r"^1[3-9]\d{9}$", description="手机号")
code: str = Field(..., min_length=4, max_length=8, description="短信验证码")
remember_me: bool = Field(False, description="是否记住登录")
class RegisterRequest(BaseModel):
phone: str
code: str
password: str
phone: str = Field(..., pattern=r"^1[3-9]\d{9}$", description="手机号,注册成功后 username 默认等于手机号")
code: str = Field(..., min_length=4, max_length=8, description="注册短信验证码")
class SetPasswordRequest(BaseModel):
new_password: str = Field(..., min_length=6, description="新密码,至少6位")
class ChangePasswordRequest(BaseModel):
old_password: str
new_password: str
old_password: str = Field(..., description="旧密码")
new_password: str = Field(..., min_length=6, description="新密码,至少6位")
class TokenResponse(BaseModel):
+11 -1
View File
@@ -1,13 +1,23 @@
from enum import Enum
from pydantic import BaseModel, Field
class SmsScene(str, Enum):
register = "register"
login = "login"
set_password = "set_password"
class SmsSendRequest(BaseModel):
phone: str = Field(..., pattern=r"^1[3-9]\d{9}$", description="手机号")
captcha_token: str | None = None
scene: SmsScene = Field(..., description="短信场景:register=注册,login=短信登录,set_password=设置密码")
captcha_token: str | None = Field(None, description="图形验证码令牌,正式环境按配置要求传入")
class SmsVerifyRequest(BaseModel):
phone: str = Field(..., pattern=r"^1[3-9]\d{9}$", description="手机号")
scene: SmsScene = Field(..., description="短信场景")
code: str = Field(..., min_length=4, max_length=8, description="验证码")
+2
View File
@@ -5,10 +5,12 @@ class UserOut(BaseModel):
id: str
username: str
email: str | None = None
phone: str | None = None
avatar: str | None = None
credits: float
is_admin: bool = False
user_type: str = "frontend"
allowed_menus: list | None = None
must_set_password: bool = False
model_config = {"from_attributes": True}
+33 -9
View File
@@ -2,7 +2,7 @@ from datetime import datetime, timedelta, timezone
import bcrypt
import jwt
from sqlalchemy import select
from sqlalchemy import or_, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
@@ -13,8 +13,13 @@ def hash_password(plain: str) -> str:
return bcrypt.hashpw(plain.encode(), bcrypt.gensalt()).decode()
def verify_password(plain: str, hashed: str) -> bool:
return bcrypt.checkpw(plain.encode(), hashed.encode())
def verify_password(plain: str, hashed: str | None) -> bool:
if not hashed:
return False
try:
return bcrypt.checkpw(plain.encode(), hashed.encode())
except Exception:
return False
def create_access_token(user_id: str, remember_me: bool = False) -> str:
@@ -35,17 +40,36 @@ def decode_access_token(token: str) -> str | None:
return None
async def get_user_by_username_or_phone(
db: AsyncSession, username_or_phone: str
) -> User | None:
value = (username_or_phone or "").strip()
if not value:
return None
result = await db.execute(
select(User)
.where(or_(User.username == value, User.phone == value))
.limit(1)
)
return result.scalar_one_or_none()
async def get_user_by_phone(db: AsyncSession, phone: str) -> User | None:
result = await db.execute(select(User).where(User.phone == phone).limit(1))
return result.scalar_one_or_none()
async def authenticate_user(
db: AsyncSession, username: str, password: str
) -> User | None:
result = await db.execute(select(User).where(User.username == username).limit(1))
user = result.scalar_one_or_none()
if not user:
# Try phone number lookup for frontend users
result = await db.execute(select(User).where(User.phone == username).limit(1))
user = result.scalar_one_or_none()
user = await get_user_by_username_or_phone(db, username)
if not user or not verify_password(password, user.hashed_password):
return None
if not user.is_active:
return None
return user
def user_must_set_password(user: User | None) -> bool:
return bool(user and user.user_type == "frontend" and not user.hashed_password)
@@ -0,0 +1,338 @@
# app/services/celery_download_recovery_service.py
from __future__ import annotations
import inspect
import json
import logging
from datetime import datetime, timezone
from typing import Any, Dict, Iterable, List, Optional, Union
from app.config import settings
try:
from redis.exceptions import RedisError
except ImportError:
RedisError = RuntimeError # type: ignore[assignment]
logger = logging.getLogger("video_gen")
_redis_client: Optional[Any] = None
def utc_now() -> datetime:
return datetime.now(timezone.utc)
def ensure_aware_utc(value: Optional[datetime]) -> Optional[datetime]:
if value is None:
return None
if value.tzinfo is None:
return value.replace(tzinfo=timezone.utc)
return value.astimezone(timezone.utc)
def datetime_to_epoch(value: Optional[datetime]) -> int:
checked_value = ensure_aware_utc(value) or utc_now()
return int(checked_value.timestamp())
def _registry_redis_url() -> str:
return settings.CELERY_BROKER_URL or settings.REDIS_URL or ""
async def get_registry_redis() -> Optional[Any]:
global _redis_client
if _redis_client is not None:
return _redis_client
redis_url = _registry_redis_url()
if not redis_url:
return None
try:
from redis.asyncio import Redis
except ImportError as exc:
logger.warning(
"下载容灾 Redis 注册表不可用,redis 依赖未安装。error=%s",
exc,
)
return None
try:
redis_client = Redis.from_url(redis_url, decode_responses=True)
await redis_client.ping()
_redis_client = redis_client
return _redis_client
except (RedisError, OSError, RuntimeError) as exc:
logger.warning(
"下载容灾 Redis 注册表不可用,降级为仅 DB 容灾。error=%s",
exc,
)
_redis_client = None
return None
async def close_registry_redis() -> None:
global _redis_client
client = _redis_client
_redis_client = None
if client is None:
return
try:
close_method = getattr(client, "close", None)
if close_method is None:
return
close_result = close_method()
if inspect.isawaitable(close_result):
await close_result
except (RedisError, OSError, RuntimeError) as exc:
logger.debug(
"关闭下载容灾 Redis 注册表连接失败。error=%s",
exc,
)
def build_download_active_payload(
*,
record_id: str,
celery_task_id: Optional[str],
stage: str,
attempt: Optional[int] = None,
queue: str = "gen_result_download",
priority: Optional[int] = None,
enqueue_at: Optional[datetime] = None,
started_at: Optional[datetime] = None,
updated_at: Optional[datetime] = None,
lease_until: Optional[datetime] = None,
next_retry_at: Optional[datetime] = None,
check_at: Optional[datetime] = None,
reason: Optional[str] = None,
) -> Dict[str, Any]:
now = utc_now()
checked_updated_at = ensure_aware_utc(updated_at) or now
checked_enqueue_at = ensure_aware_utc(enqueue_at)
checked_started_at = ensure_aware_utc(started_at)
checked_lease_until = ensure_aware_utc(lease_until)
checked_next_retry_at = ensure_aware_utc(next_retry_at)
checked_check_at = ensure_aware_utc(check_at)
return {
"record_id": record_id,
"celery_task_id": celery_task_id,
"stage": stage,
"attempt": int(attempt or 0),
"queue": queue,
"priority": priority,
"enqueue_at": (
datetime_to_epoch(checked_enqueue_at)
if checked_enqueue_at
else None
),
"started_at": (
datetime_to_epoch(checked_started_at)
if checked_started_at
else None
),
"updated_at": datetime_to_epoch(checked_updated_at),
"lease_until": (
datetime_to_epoch(checked_lease_until)
if checked_lease_until
else None
),
"next_retry_at": (
datetime_to_epoch(checked_next_retry_at)
if checked_next_retry_at
else None
),
"check_at": (
datetime_to_epoch(checked_check_at)
if checked_check_at
else None
),
"reason": reason,
}
async def upsert_download_active(
*,
record_id: str,
payload: Dict[str, Any],
check_at: Optional[Union[datetime, int, float]],
) -> None:
redis = await get_registry_redis()
if redis is None:
return
if isinstance(check_at, datetime):
score = datetime_to_epoch(check_at)
elif check_at is None:
score = datetime_to_epoch(utc_now())
else:
score = int(float(check_at))
updated_payload = dict(payload)
updated_payload["check_at"] = score
try:
pipe: Any = redis.pipeline(transaction=True)
pipe.hset(
settings.DOWNLOAD_ACTIVE_REDIS_HASH_KEY,
record_id,
json.dumps(updated_payload, ensure_ascii=False, default=str),
)
pipe.zadd(
settings.DOWNLOAD_ACTIVE_REDIS_ZSET_KEY,
{record_id: score},
)
await pipe.execute()
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning(
"写入下载容灾 Redis 注册表失败。record_id=%s, error=%s",
record_id,
exc,
)
async def remove_download_active(record_id: str) -> None:
redis = await get_registry_redis()
if redis is None:
return
try:
pipe: Any = redis.pipeline(transaction=True)
pipe.hdel(settings.DOWNLOAD_ACTIVE_REDIS_HASH_KEY, record_id)
pipe.zrem(settings.DOWNLOAD_ACTIVE_REDIS_ZSET_KEY, record_id)
await pipe.execute()
except (RedisError, OSError, RuntimeError) as exc:
logger.warning(
"删除下载容灾 Redis 注册表失败。record_id=%s, error=%s",
record_id,
exc,
)
async def get_due_download_record_ids(
*,
limit: Optional[int] = None,
now: Optional[datetime] = None,
) -> List[str]:
redis = await get_registry_redis()
if redis is None:
return []
batch_limit = int(limit or settings.DOWNLOAD_RECOVERY_BATCH_SIZE or 100)
score = datetime_to_epoch(now or utc_now())
try:
result = await redis.zrangebyscore(
settings.DOWNLOAD_ACTIVE_REDIS_ZSET_KEY,
min="-inf",
max=score,
start=0,
num=batch_limit,
)
return [str(item) for item in result]
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning(
"扫描下载容灾 Redis ZSet 失败。error=%s",
exc,
)
return []
async def get_download_active_payloads(
record_ids: Iterable[str],
) -> Dict[str, Dict[str, Any]]:
cleaned_record_ids = [str(item) for item in record_ids if item]
if not cleaned_record_ids:
return {}
redis = await get_registry_redis()
if redis is None:
return {}
try:
raw_values = await redis.hmget(
settings.DOWNLOAD_ACTIVE_REDIS_HASH_KEY,
cleaned_record_ids,
)
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning(
"读取下载容灾 Redis Hash 失败。error=%s",
exc,
)
return {}
result: Dict[str, Dict[str, Any]] = {}
for record_id, raw in zip(cleaned_record_ids, raw_values):
if not raw:
continue
try:
value = json.loads(raw)
except (TypeError, ValueError, json.JSONDecodeError):
continue
if isinstance(value, dict):
result[record_id] = value
return result
async def postpone_download_active_check(
*,
record_id: str,
payload: Optional[Dict[str, Any]] = None,
check_at: Optional[Union[datetime, int, float]] = None,
) -> None:
redis = await get_registry_redis()
if redis is None:
return
if isinstance(check_at, datetime):
score = datetime_to_epoch(check_at)
elif check_at is None:
score = datetime_to_epoch(utc_now()) + int(
settings.DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS or 300
)
else:
score = int(float(check_at))
try:
pipe: Any = redis.pipeline(transaction=True)
pipe.zadd(
settings.DOWNLOAD_ACTIVE_REDIS_ZSET_KEY,
{record_id: score},
)
if payload is not None:
updated_payload = dict(payload)
updated_payload["check_at"] = score
updated_payload["updated_at"] = datetime_to_epoch(utc_now())
pipe.hset(
settings.DOWNLOAD_ACTIVE_REDIS_HASH_KEY,
record_id,
json.dumps(
updated_payload,
ensure_ascii=False,
default=str,
),
)
await pipe.execute()
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning(
"刷新下载容灾 Redis 检查时间失败。record_id=%s, error=%s",
record_id,
exc,
)
@@ -0,0 +1,17 @@
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.credit_ratio import CreditRatio
async def list_all_credit_ratios(db: AsyncSession) -> list[CreditRatio]:
"""获取全部积分比例规则,供客户端只读展示和后台复用。"""
result = await db.execute(
select(CreditRatio).order_by(
CreditRatio.gen_type.asc(),
CreditRatio.model_config_id.asc(),
CreditRatio.resolution.asc(),
CreditRatio.created_at.desc(),
)
)
return list(result.scalars().all())
@@ -1,8 +1,9 @@
from __future__ import annotations
import os
import uuid
from dataclasses import dataclass
from datetime import datetime
from datetime import datetime, timezone
from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask
@@ -24,17 +25,93 @@ class DownloadedGenerationResult:
cover_storage_path: str | None = None
def _to_aware_utc(value: datetime | None) -> datetime | None:
if value is None:
return None
if value.tzinfo is None:
return value.replace(tzinfo=timezone.utc)
return value.astimezone(timezone.utc)
def _build_storage_date_dir(record: ChatGenerationTask) -> str:
fixed = (getattr(record, "download_storage_date_dir", None) or "").strip().strip("/")
if fixed:
return fixed
created_at = _to_aware_utc(getattr(record, "created_at", None)) or datetime.now(timezone.utc)
return created_at.strftime("%Y/%m/%d")
def _make_part_path(final_path: str) -> str:
return f"{final_path}.{uuid.uuid4().hex}.part"
def _is_valid_file(path: str | None) -> bool:
if not path:
return False
try:
return os.path.isfile(path) and os.path.getsize(path) > 0
except OSError:
return False
def _safe_remove(path: str | None) -> None:
if not path:
return
try:
if os.path.exists(path):
os.remove(path)
except OSError:
pass
async def _download_image_atomically(remote_url: str, final_path: str) -> 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)
if not _is_valid_file(part_path):
raise RuntimeError("图片下载完成但临时文件为空")
os.replace(part_path, final_path)
return final_path
except Exception:
_safe_remove(part_path)
raise
async def _download_video_atomically(remote_url: str, final_path: str) -> 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)
if not _is_valid_file(part_path):
raise RuntimeError("视频下载完成但临时文件为空")
os.replace(part_path, final_path)
return final_path
except Exception:
_safe_remove(part_path)
raise
async def download_generation_result(record: ChatGenerationTask) -> DownloadedGenerationResult:
if not record.remote_result_url:
raise ValueError("缺少远程结果URL")
date_dir = datetime.now().strftime("%Y/%m/%d")
date_dir = _build_storage_date_dir(record)
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")
async with provider_limit("result_download", settings.RESULT_DOWNLOAD_MAX_CONCURRENCY):
await download_image(record.remote_result_url, dest)
await _download_image_atomically(record.remote_result_url if record.remote_result_url else "", dest)
return DownloadedGenerationResult(
url=f"/generate/images/{date_dir}/{record.id}.png",
storage_path=dest,
@@ -45,8 +122,9 @@ async def download_generation_result(record: ChatGenerationTask) -> DownloadedGe
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)
await _download_video_atomically(record.remote_result_url if record.remote_result_url else "", dest)
cover_url, cover_storage_path = create_video_cover_for_local_video(
record_id=record.id,
@@ -0,0 +1,378 @@
from __future__ import annotations
import logging
from datetime import datetime, timedelta, timezone
from typing import Any, Dict
from sqlalchemy import or_, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask
from app.services.celery_download_recovery_service import (
ensure_aware_utc,
get_download_active_payloads,
get_due_download_record_ids,
postpone_download_active_check,
remove_download_active,
)
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
logger = logging.getLogger("video_gen")
def _now() -> datetime:
return datetime.now(timezone.utc)
def _is_expired(value: datetime | None, now: datetime | None = None) -> bool:
checked = ensure_aware_utc(value)
if checked is None:
return True
return checked <= (now or _now())
def _queue_timeout_at(task: ChatGenerationTask, now: datetime | None = None) -> datetime:
current_time = now or _now()
enqueued_at = ensure_aware_utc(task.download_enqueued_at)
if enqueued_at is None:
return current_time
return enqueued_at + timedelta(seconds=int(settings.DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS or 300))
def _is_queue_timeout(task: ChatGenerationTask, now: datetime | None = None) -> bool:
current_time = now or _now()
return _queue_timeout_at(task, current_time) <= current_time
def _is_final_task_state(task: ChatGenerationTask) -> bool:
return task.status in ("completed", "failed") or task.pipeline_stage in (
"done",
"failed",
"timeout",
"download_failed",
)
async def recover_one_download_task(
db: AsyncSession,
task: ChatGenerationTask,
*,
payload: dict[str, Any] | None = None,
source: str = "startup_db",
) -> str:
from app.tasks.generation_download_tasks import (
DOWNLOAD_STAGE_DOWNLOADING,
DOWNLOAD_STAGE_QUEUED,
DOWNLOAD_STAGE_RETRY_WAITING,
enqueue_download_task,
)
current_time = _now()
if not task:
return "skip_missing_task"
if task.generation_mode != "chatapi_async":
await remove_download_active(task.id)
return "clean_invalid_mode"
if _is_final_task_state(task):
await remove_download_active(task.id)
return "clean_final_state"
if task.status != "generating":
await remove_download_active(task.id)
return "clean_not_generating"
if not task.remote_result_url:
return "skip_no_remote_result_url"
stage = task.pipeline_stage
redis_payload = payload or {}
if stage == "result_ready":
await log_task_event(
task,
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
message=f"{source} 发现 result_ready 未完成下载,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_result_ready",
)
return "recover_result_ready"
if stage == DOWNLOAD_STAGE_QUEUED:
if _is_queue_timeout(task, current_time):
await log_task_event(
task,
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
message=f"{source} 发现 download_queued 长时间未消费,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_download_queued_timeout",
)
return "recover_queued_timeout"
await postpone_download_active_check(
record_id=task.id,
payload=payload,
check_at=_queue_timeout_at(task, current_time),
)
return "skip_queued_not_timeout"
if stage == DOWNLOAD_STAGE_DOWNLOADING:
if _is_expired(task.download_lease_until, current_time):
await log_task_event(
task,
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
message=f"{source} 发现 downloading lease 过期,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_downloading_lease_expired",
)
return "recover_downloading_expired"
await postpone_download_active_check(
record_id=task.id,
payload=payload,
check_at=task.download_lease_until,
)
return "skip_downloading_alive"
if stage == DOWNLOAD_STAGE_RETRY_WAITING:
if _is_expired(task.download_next_retry_at, current_time):
await log_task_event(
task,
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
message=f"{source} 发现 retry_waiting 到期,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_retry_waiting_due",
)
return "recover_retry_due"
await postpone_download_active_check(
record_id=task.id,
payload=payload,
check_at=task.download_next_retry_at,
)
return "skip_retry_waiting_not_due"
return f"skip_stage_{stage}"
async def recover_download_tasks_once(db: AsyncSession) -> dict[str, Any]:
"""启动时下载容灾扫描。
先按 Redis active_index 找到到期下载任务;Redis 不可用或索引丢失时,
再通过 DB fallback 扫描 result_ready/download_* 状态,避免任务永久卡住。
"""
checked_ids: set[str] = set()
results: dict[str, int] = {}
due_ids = await get_due_download_record_ids(
limit=settings.DOWNLOAD_RECOVERY_BATCH_SIZE,
)
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()
if task is None:
await remove_download_active(task_id)
action = "clean_missing_task"
else:
checked_ids.add(task.id)
action = await recover_one_download_task(
db,
task,
payload=payloads.get(task_id),
source="startup_redis",
)
results[action] = results.get(action, 0) + 1
# DB fallback:不依赖 Redis active 注册表。
fallback_result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.deleted_at.is_(None),
ChatGenerationTask.generation_mode == "chatapi_async",
ChatGenerationTask.status == "generating",
ChatGenerationTask.remote_result_url.is_not(None),
ChatGenerationTask.pipeline_stage.in_(
["result_ready", "download_queued", "downloading", "retry_waiting"]
),
)
.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()
for task in fallback_tasks:
if task.id in checked_ids:
continue
action = await recover_one_download_task(
db,
task,
payload=None,
source="startup_db",
)
results[action] = results.get(action, 0) + 1
checked_ids.add(task.id)
return {"checked": len(checked_ids), "results": results}
async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
"""启动时生成链路容灾扫描。
只在 Celery worker 启动时跑一次,不引入 beat,不新增第四条启动命令。
用于把 queued/creating/waiting_remote/polling/result_ready 等中间态重新投递到现有三个队列。
"""
from app.tasks.generation_create_tasks import chatapi_create_generation_task
from app.tasks.generation_download_tasks import enqueue_download_task
from app.tasks.generation_poll_tasks import poll_generation_task
current_time = _now()
results: dict[str, int] = {}
query_result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.deleted_at.is_(None),
ChatGenerationTask.generation_mode == "chatapi_async",
ChatGenerationTask.status == "generating",
ChatGenerationTask.pipeline_stage.in_(
[
"queued",
"preparing",
"creating_provider_task",
"waiting_remote",
"polling",
"result_ready",
]
),
)
.order_by(ChatGenerationTask.updated_at.asc())
.limit(int(settings.DOWNLOAD_RECOVERY_BATCH_SIZE or 100))
.with_for_update(skip_locked=True)
)
tasks = query_result.scalars().all()
for task in tasks:
if task.deadline_at and _is_expired(task.deadline_at, current_time):
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message="任务超时",
pipeline_stage="timeout",
)
await db.commit()
await log_task_event(
task,
event_type="TASK_TIMEOUT",
to_status="failed",
to_stage="timeout",
)
action = "mark_timeout"
elif task.pipeline_stage in ("queued", "preparing", "creating_provider_task"):
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message="启动时发现创建阶段任务未完成,恢复投递创建队列",
detail={"pipeline_stage": task.pipeline_stage},
)
chatapi_create_generation_task.apply_async(
args=[task.id],
queue="gen_chatapi_create",
countdown=0,
)
action = "recover_create"
elif task.pipeline_stage in ("waiting_remote", "polling"):
if task.remote_result_url:
await enqueue_download_task(
db,
task,
recover=True,
reason="startup_waiting_remote_has_result",
)
action = "recover_waiting_has_result"
elif task.provider_task_id or task.seedance_task_id:
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message="启动时发现远程等待/轮询阶段任务未完成,恢复投递轮询队列",
detail={"pipeline_stage": task.pipeline_stage},
)
task.pipeline_stage = "waiting_remote"
await db.commit()
poll_generation_task.apply_async(
args=[task.id],
queue="gen_provider_poll",
countdown=0,
)
action = "recover_poll"
else:
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message="启动时发现任务缺少供应商任务ID,恢复投递创建队列",
detail={"pipeline_stage": task.pipeline_stage},
)
task.pipeline_stage = "queued"
await db.commit()
chatapi_create_generation_task.apply_async(
args=[task.id],
queue="gen_chatapi_create",
countdown=0,
)
action = "recover_create_missing_provider_id"
elif task.pipeline_stage == "result_ready":
if task.remote_result_url:
await enqueue_download_task(
db,
task,
recover=True,
reason="startup_generation_result_ready",
)
action = "recover_result_ready"
else:
action = "skip_result_ready_no_url"
else:
action = f"skip_stage_{task.pipeline_stage}"
results[action] = results.get(action, 0) + 1
# 下载阶段单独跑 DB fallback。
download_result = await recover_download_tasks_once(db)
return {
"checked": len(tasks),
"results": results,
"download_recovery": download_result,
}
+152 -45
View File
@@ -1,79 +1,186 @@
import asyncio
import json
import logging
import random
import time
import httpx
from datetime import datetime
from typing import Any
from app.config import settings
from app.utils.redis import get_redis
logger = logging.getLogger("videogen")
# In-memory fallback for verification codes
# In-memory fallback for local development when Redis is disabled.
_sms_code_store: dict[str, tuple[str, float]] = {}
_sms_send_interval_store: dict[str, float] = {}
_sms_daily_count_store: dict[str, tuple[str, int]] = {}
def _generate_code(length: int = 6) -> str:
return "".join(random.choices("0123456789", k=length))
def _generate_code(length: int | None = None) -> str:
code_length = int(length or settings.SMS_CODE_LENGTH or 4)
code_length = max(4, min(code_length, 8))
return "".join(random.choices("0123456789", k=code_length))
async def send_sms(phone: str, code: str) -> bool:
"""Send SMS verification code. Supports mock mode and real HTTP gateway."""
if settings.SMS_MOCK or not settings.SMS_API_URL:
logger.info(f"[SMS MOCK] To={phone}, Code={code}")
return True
def _normalize_scene(scene: str | None) -> str:
value = (scene or "").strip().lower()
if value not in {"register", "login", "set_password"}:
value = "login"
return value
def _code_key(phone: str, scene: str | None) -> str:
return f"sms_code:{_normalize_scene(scene)}:{phone}"
def _interval_key(phone: str, scene: str | None) -> str:
return f"sms_interval:{_normalize_scene(scene)}:{phone}"
def _daily_key(phone: str, scene: str | None) -> str:
today = datetime.now().strftime("%Y%m%d")
return f"sms_daily:{_normalize_scene(scene)}:{phone}:{today}"
def _today() -> str:
return datetime.now().strftime("%Y%m%d")
async def _check_send_limit(phone: str, scene: str) -> None:
"""发送频控:同场景同手机号间隔限制 + 每日次数限制。"""
redis = get_redis()
interval_seconds = int(settings.SMS_SEND_INTERVAL_SECONDS or 60)
daily_limit = int(settings.SMS_DAILY_LIMIT or 20)
if redis:
interval_key = _interval_key(phone, scene)
if await redis.get(interval_key):
raise ValueError(f"短信发送过于频繁,请{interval_seconds}秒后再试")
daily_key = _daily_key(phone, scene)
count = await redis.incr(daily_key)
if count == 1:
await redis.expire(daily_key, 24 * 60 * 60)
if count > daily_limit:
raise ValueError("今日短信发送次数已达上限,请明天再试")
await redis.setex(interval_key, interval_seconds, "1")
return
now_ts = time.time()
interval_key = _interval_key(phone, scene)
last_send_at = _sms_send_interval_store.get(interval_key)
if last_send_at and now_ts - last_send_at < interval_seconds:
raise ValueError(f"短信发送过于频繁,请{interval_seconds}秒后再试")
_sms_send_interval_store[interval_key] = now_ts
daily_key = f"{_normalize_scene(scene)}:{phone}"
current_day = _today()
stored_day, count = _sms_daily_count_store.get(daily_key, (current_day, 0))
if stored_day != current_day:
stored_day, count = current_day, 0
count += 1
_sms_daily_count_store[daily_key] = (stored_day, count)
if count > daily_limit:
raise ValueError("今日短信发送次数已达上限,请明天再试")
def _send_volc_sms_sync(phone: str, code: str) -> dict[str, Any]:
from volcengine.sms.SmsService import SmsService
if not settings.VOLC_SMS_ACCESS_KEY_ID or not settings.VOLC_SMS_SECRET_ACCESS_KEY:
raise RuntimeError("火山短信 AK/SK 未配置")
if not settings.VOLC_SMS_ACCOUNT:
raise RuntimeError("火山短信消息组ID VOLC_SMS_ACCOUNT 未配置")
if not settings.VOLC_SMS_TEMPLATE_ID:
raise RuntimeError("火山短信模板ID VOLC_SMS_TEMPLATE_ID 未配置")
if not settings.VOLC_SMS_SIGN:
raise RuntimeError("火山短信签名 VOLC_SMS_SIGN 未配置")
sms_service = SmsService()
sms_service.set_ak(settings.VOLC_SMS_ACCESS_KEY_ID)
sms_service.set_sk(settings.VOLC_SMS_SECRET_ACCESS_KEY)
body = {
"SmsAccount": settings.VOLC_SMS_ACCOUNT,
"Sign": settings.VOLC_SMS_SIGN,
"TemplateID": settings.VOLC_SMS_TEMPLATE_ID,
"TemplateParam": json.dumps({"xxxx": str(code)}, ensure_ascii=False, separators=(",", ":")),
"Tag": f"{phone}:{int(time.time())}",
"PhoneNumbers": phone,
}
raw_resp = sms_service.send_sms(json.dumps(body, ensure_ascii=False, separators=(",", ":")))
if isinstance(raw_resp, str):
try:
resp: dict[str, Any] = json.loads(raw_resp)
except json.JSONDecodeError:
resp = {"raw": raw_resp}
elif isinstance(raw_resp, dict):
resp = raw_resp
else:
resp = {"raw": raw_resp}
error = (resp.get("ResponseMetadata") or {}).get("Error") if isinstance(resp, dict) else None
if error:
raise RuntimeError(f"火山短信发送失败:{error.get('Code')} {error.get('Message')}")
return resp
async def send_sms(phone: str, code: str, scene: str = "login") -> bool:
"""发送短信验证码。
SMS_MOCK=true 时只写日志,方便本地调试;否则使用火山引擎短信 SDK。
火山 SDK 为同步调用,这里放到线程中执行,避免阻塞 FastAPI event loop。
"""
scene = _normalize_scene(scene)
try:
payload = {
"phone": phone,
"code": code,
"sign_name": settings.SMS_SIGN_NAME,
"template_code": settings.SMS_TEMPLATE_CODE,
}
headers = {
"Authorization": f"Bearer {settings.SMS_API_KEY}",
"Content-Type": "application/json",
}
async with httpx.AsyncClient(timeout=10) as client:
resp = await client.post(settings.SMS_API_URL, json=payload, headers=headers)
resp.raise_for_status()
return True
logger.info("[SMS] scene=%s, To=%s, Code=%s", scene, phone, code)
await asyncio.to_thread(_send_volc_sms_sync, phone, code)
return True
except Exception:
logger.exception(f"SMS send failed for {phone}")
logger.exception("SMS send failed. scene=%s phone=%s", scene, phone)
return False
async def store_sms_code(phone: str, code: str, ttl: int = 300) -> None:
"""Store SMS verification code with TTL (default 5 minutes)."""
async def store_sms_code(phone: str, code: str, scene: str = "login", ttl: int | None = None) -> None:
ttl_seconds = int(ttl or settings.SMS_CODE_TTL_SECONDS or 300)
key = _code_key(phone, scene)
redis = get_redis()
if redis:
await redis.setex(f"sms_code:{phone}", ttl, code)
await redis.setex(key, ttl_seconds, code)
else:
_sms_code_store[phone] = (code, time.time() + ttl)
_sms_code_store[key] = (code, time.time() + ttl_seconds)
async def generate_and_send_sms(phone: str) -> bool:
"""Generate a code, store it, and send it via SMS."""
async def generate_and_send_sms(phone: str, scene: str = "login") -> bool:
scene = _normalize_scene(scene)
await _check_send_limit(phone, scene)
code = _generate_code()
ok = await send_sms(phone, code)
ok = await send_sms(phone, code, scene)
if ok:
await store_sms_code(phone, code)
await store_sms_code(phone, code, scene)
return ok
async def verify_sms_code(phone: str, code: str) -> bool:
"""Verify an SMS verification code."""
async def verify_sms_code(phone: str, code: str, scene: str = "login") -> bool:
key = _code_key(phone, scene)
redis = get_redis()
if redis:
stored = await redis.get(f"sms_code:{phone}")
if stored and stored == code:
await redis.delete(f"sms_code:{phone}")
stored = await redis.get(key)
if isinstance(stored, bytes):
stored = stored.decode()
if stored and str(stored) == str(code):
await redis.delete(key)
return True
return False
else:
entry = _sms_code_store.pop(phone, None)
if entry:
stored_code, expires = entry
if time.time() < expires and stored_code == code:
return True
return False
entry = _sms_code_store.pop(key, None)
if entry:
stored_code, expires = entry
if time.time() < expires and stored_code == code:
return True
return False
@@ -1,13 +1,13 @@
import asyncio
import logging
import os
import shutil
import subprocess
import uuid
from pathlib import Path
from app.config import settings
import asyncio
logger = logging.getLogger("videogen")
@@ -20,6 +20,29 @@ def _clean_cover_format(value: str | None) -> str:
return ext or "jpg"
def _is_valid_file(path: str | None) -> bool:
if not path:
return False
try:
return os.path.isfile(path) and os.path.getsize(path) > 0
except OSError:
return False
def _safe_remove(path: str | None) -> None:
if not path:
return
try:
if os.path.exists(path):
os.remove(path)
except OSError:
pass
def _make_part_path(final_path: str) -> str:
return f"{final_path}.{uuid.uuid4().hex}.part"
def get_ffmpeg_bin() -> str:
"""
获取 ffmpeg 可执行文件路径。
@@ -119,6 +142,34 @@ def generate_video_cover(
return str(output_file)
def generate_video_cover_atomically(
video_path: str,
output_path: str,
seek_time: str = "00:00:01",
width: int = 720,
timeout: int = 15,
) -> str:
if _is_valid_file(output_path):
return output_path
part_path = _make_part_path(output_path)
try:
generate_video_cover(
video_path=video_path,
output_path=part_path,
seek_time=seek_time,
width=width,
timeout=timeout,
)
if not _is_valid_file(part_path):
raise VideoCoverError(f"封面临时文件为空: {part_path}")
os.replace(part_path, output_path)
return output_path
except Exception:
_safe_remove(part_path)
raise
def build_video_cover_path_and_url(record_id: str, date_dir: str) -> tuple[str, str]:
"""
根据记录ID和日期目录生成本地封面文件路径与对外URL。
@@ -158,7 +209,7 @@ def try_generate_video_cover(
cover_timeout = timeout or settings.VIDEO_COVER_TIMEOUT_SECONDS
try:
return generate_video_cover(
return generate_video_cover_atomically(
video_path=video_path,
output_path=output_path,
seek_time=first_seek,
@@ -168,7 +219,7 @@ def try_generate_video_cover(
except Exception as first_exc:
if second_seek and second_seek != first_seek:
try:
return generate_video_cover(
return generate_video_cover_atomically(
video_path=video_path,
output_path=output_path,
seek_time=second_seek,
@@ -222,6 +273,7 @@ def create_video_cover_for_local_video(
return None, None
return cover_url, generated_path
async def async_create_video_cover_for_local_video(
*,
record_id: str,
+6 -2
View File
@@ -5,7 +5,11 @@ custom named tasks are registered when workers start.
"""
try:
from app.tasks import generation_create_tasks, generation_poll_tasks, generation_download_tasks # noqa: F401
from app.tasks import ( # noqa: F401
generation_create_tasks,
generation_poll_tasks,
generation_download_tasks,
generation_recovery_tasks,
)
except Exception:
# Keep application importable even when optional Celery dependencies/config are absent.
pass
+61 -15
View File
@@ -1,16 +1,20 @@
import logging
from celery import Celery
from celery.signals import worker_process_init, worker_process_shutdown, worker_ready
from app.config import settings
from celery.signals import worker_process_init, worker_process_shutdown
from app.tasks.async_runner import run_async, close_loop
from app.models.base import engine
from app.tasks.async_runner import close_loop, run_async
logger = logging.getLogger("video_gen")
def _derive_redis_db(url: str, db_no: int) -> str:
if not url:
return url
import re
if re.search(r"/\d+$", url):
return re.sub(r"/\d+$", f"/{db_no}", url)
return url.rstrip("/") + f"/{db_no}"
@@ -37,11 +41,16 @@ if broker_url:
worker_prefetch_multiplier=1,
broker_transport_options={
"visibility_timeout": 3600,
"queue_order_strategy": "priority",
"priority_steps": list(range(10)),
"sep": ":",
},
task_routes={
"generation.chatapi_create_generation_task": {"queue": "gen_chatapi_create"},
"generation.poll_generation_task": {"queue": "gen_provider_poll"},
"generation.download_generation_result_task": {"queue": "gen_result_download"},
"generation.recover_download_tasks_once": {"queue": "gen_result_download"},
"generation.recover_generation_tasks_once": {"queue": "gen_result_download"},
"app.tasks.cleanup.*": {"queue": "default"},
},
)
@@ -50,15 +59,47 @@ else:
celery_app = None
@worker_ready.connect
def on_worker_ready(sender=None, **kwargs):
"""Celery worker 启动时做一次容灾恢复。
注意:
- 不启用 Celery beat。
- 不要求新增第四条启动命令。
- 只让 gen_result_download worker 投递恢复任务,避免三个 worker 同时重复扫描。
"""
if celery_app is None:
return
hostname = str(getattr(sender, "hostname", "") or "")
if "gen_result_download" not in hostname:
return
try:
from app.tasks.generation_recovery_tasks import (
recover_download_tasks_once,
recover_generation_tasks_once,
)
countdown = max(0, int(settings.DOWNLOAD_RECOVERY_STARTUP_DELAY_SECONDS or 0))
recover_generation_tasks_once.apply_async(
countdown=countdown,
queue="gen_result_download",
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
)
recover_download_tasks_once.apply_async(
countdown=countdown + 5,
queue="gen_result_download",
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
)
except Exception:
logger.exception("启动容灾恢复任务投递失败")
@worker_process_init.connect
def on_worker_process_init(**kwargs):
"""
Linux prefork 子进程启动后执行。
目的:
1. 丢弃 fork 前可能继承的连接池状态。
2. 后续任务会在当前子进程自己的长期 event loop 上重新建连接池。
"""
"""Linux prefork 子进程启动后丢弃 fork 前可能继承的连接池状态。"""
try:
run_async(engine.dispose())
except Exception:
@@ -67,12 +108,17 @@ def on_worker_process_init(**kwargs):
@worker_process_shutdown.connect
def on_worker_process_shutdown(**kwargs):
"""
子进程退出前关闭连接池和 event loop。
"""
"""子进程退出前关闭连接池和 event loop。"""
try:
run_async(engine.dispose())
except Exception:
pass
try:
from app.services.celery_download_recovery_service import close_registry_redis
run_async(close_registry_redis())
except Exception:
pass
finally:
close_loop()
close_loop()
@@ -172,9 +172,9 @@ async def _run(task_id: str):
task.pipeline_stage = "result_ready"
await db.commit()
from app.tasks.generation_download_tasks import download_generation_result_task
from app.tasks.generation_download_tasks import enqueue_download_task
download_generation_result_task.delay(task.id)
await enqueue_download_task(db, task, reason="create_remote_result_ready")
return
else:
@@ -223,9 +223,9 @@ async def _run(task_id: str):
)
if task.pipeline_stage == "result_ready":
from app.tasks.generation_download_tasks import download_generation_result_task
from app.tasks.generation_download_tasks import enqueue_download_task
download_generation_result_task.delay(task.id)
await enqueue_download_task(db, task, reason="create_result_ready")
else:
from app.tasks.generation_poll_tasks import poll_generation_task
@@ -1,10 +1,19 @@
from app.tasks.async_runner import run_async
from datetime import datetime, timezone, timedelta
import uuid
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.models.base import async_session
from app.models.chat_generation_task import ChatGenerationTask
from app.services.celery_download_recovery_service import (
build_download_active_payload,
ensure_aware_utc,
remove_download_active,
upsert_download_active,
)
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
@@ -12,52 +21,132 @@ from app.services.generation_refund_service import mark_chat_generation_task_fai
from app.services.resource_accounting_service import record_chat_task_generated_resource
from app.tasks.celery_app import celery_app
# downloading 卡住多久后允许自动恢复。
# 说明:
# - worker 在 pipeline_stage 改成 downloading 后,如果被 kill,任务可能永远停在 downloading。
# - 这里允许超过该时间的 downloading 任务重新进入下载流程。
# - 如果你的视频文件特别大,可以把这个时间调大,比如 20 * 60。
DOWNLOAD_STUCK_SECONDS = 10 * 60
DOWNLOAD_QUEUE = "gen_result_download"
DOWNLOAD_STAGE_QUEUED = "download_queued"
DOWNLOAD_STAGE_DOWNLOADING = "downloading"
DOWNLOAD_STAGE_RETRY_WAITING = "retry_waiting"
DOWNLOAD_STAGE_DONE = "done"
DOWNLOAD_STAGE_FAILED = "download_failed"
def _to_aware_utc(dt):
"""
把 datetime 统一转成 timezone-aware UTC,避免 offset-naive 和 offset-aware 比较报错。
PostgreSQL / SQLite / 不同驱动返回的 updated_at 可能有时区,也可能没有。
"""
if not dt:
def _now() -> datetime:
return datetime.now(timezone.utc)
def _task_created_date_dir(task: ChatGenerationTask) -> str:
created_at = ensure_aware_utc(getattr(task, "created_at", None)) or _now()
return created_at.strftime("%Y/%m/%d")
def _queue_timeout_at(now: datetime | None = None) -> datetime:
now = now or _now()
return now + timedelta(seconds=int(settings.DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS or 300))
def _lease_until(now: datetime | None = None) -> datetime:
now = now or _now()
return now + timedelta(seconds=int(settings.DOWNLOAD_TASK_LEASE_SECONDS or 600))
def _retry_at(attempt: int, now: datetime | None = None) -> datetime:
now = now or _now()
base = int(settings.DOWNLOAD_TASK_RETRY_BACKOFF_SECONDS or 30)
return now + timedelta(seconds=max(1, base * max(1, attempt)))
def _is_expired(value: datetime | None, now: datetime | None = None) -> bool:
value = ensure_aware_utc(value)
if value is None:
return True
return value <= (now or _now())
def _is_already_completed(task: ChatGenerationTask) -> bool:
if task.status == "completed" or task.pipeline_stage == DOWNLOAD_STAGE_DONE:
if task.gen_type == "image" and task.image_url:
return True
if task.gen_type == "video" and task.video_url:
return True
return False
def _build_celery_task_id(task_id: str, attempt: int | None = None, reason: str | None = None) -> str:
safe_reason = (reason or "download").replace(" ", "_")[:32]
return f"download:{task_id}:{int(attempt or 0)}:{safe_reason}:{uuid.uuid4().hex[:12]}"
async def _register_active_from_task(
task: ChatGenerationTask,
*,
check_at: datetime,
priority: int,
reason: str | None = None,
) -> None:
payload = build_download_active_payload(
record_id=task.id,
celery_task_id=task.download_celery_task_id,
stage=task.pipeline_stage or "",
attempt=task.download_attempt_count or 0,
queue=DOWNLOAD_QUEUE,
priority=priority,
enqueue_at=task.download_enqueued_at,
started_at=task.download_started_at,
lease_until=task.download_lease_until,
next_retry_at=task.download_next_retry_at,
check_at=check_at,
reason=reason,
)
await upsert_download_active(record_id=task.id, payload=payload, check_at=check_at)
async def enqueue_download_task(
db: AsyncSession,
task: ChatGenerationTask,
*,
recover: bool = False,
reason: str | None = None,
countdown: int | None = None,
) -> str | None:
"""统一投递图片/视频下载任务,并同步 DB + Redis active 注册表。"""
if not task or task.generation_mode != "chatapi_async":
return None
if dt.tzinfo is None:
return dt.replace(tzinfo=timezone.utc)
return dt.astimezone(timezone.utc)
if task.status != "generating":
return None
if not task.remote_result_url:
return None
now = _now()
priority = settings.DOWNLOAD_TASK_PRIORITY_RECOVER if recover else settings.DOWNLOAD_TASK_PRIORITY_NORMAL
celery_task_id = _build_celery_task_id(
task.id,
attempt=task.download_attempt_count or task.retry_count or 0,
reason=reason or ("recover" if recover else "normal"),
)
task.pipeline_stage = DOWNLOAD_STAGE_QUEUED
task.download_celery_task_id = celery_task_id
task.download_enqueued_at = now
task.download_next_retry_at = None
if not task.download_storage_date_dir:
task.download_storage_date_dir = _task_created_date_dir(task)
await db.commit()
check_at = _queue_timeout_at(now)
await _register_active_from_task(task, check_at=check_at, priority=priority, reason=reason)
if celery_app:
download_generation_result_task.apply_async(
args=[task.id],
queue=DOWNLOAD_QUEUE,
priority=priority,
countdown=countdown,
task_id=celery_task_id,
)
return celery_task_id
def _is_recent_downloading(task: ChatGenerationTask) -> bool:
"""
判断 downloading 是否仍然是较新的下载任务。
返回 True
- 说明可能有另一个 worker 刚进入下载,不要重复下载。
返回 False
- 说明 downloading 已经超过 DOWNLOAD_STUCK_SECONDS,认为可能卡死,可以恢复。
"""
updated_at = _to_aware_utc(getattr(task, "updated_at", None))
if not updated_at:
return False
return datetime.now(timezone.utc) - updated_at < timedelta(seconds=DOWNLOAD_STUCK_SECONDS)
async def _reload_task(db, task_id: str) -> ChatGenerationTask | None:
"""
rollback 后重新查询任务对象。
说明:
- SQLAlchemy rollback 后,当前 ORM 对象可能过期。
- 继续访问旧 task 有概率触发异步懒加载异常。
"""
async def _reload_task(db: AsyncSession, task_id: str) -> ChatGenerationTask | None:
result = await db.execute(
select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
@@ -67,6 +156,110 @@ async def _reload_task(db, task_id: str) -> ChatGenerationTask | None:
return result.scalar_one_or_none()
async def _claim_download_lease(db: AsyncSession, task: ChatGenerationTask) -> bool:
now = _now()
if not task or task.generation_mode != "chatapi_async":
return False
if task.status != "generating":
return False
if _is_already_completed(task):
return False
if not task.remote_result_url:
return False
stage = task.pipeline_stage
if stage == DOWNLOAD_STAGE_DOWNLOADING:
if not _is_expired(task.download_lease_until, now):
return False
await log_task_event(
task,
event_type="DOWNLOAD_STUCK_RECOVER",
message=f"downloading lease 已过期,重新抢占下载。lease_until={task.download_lease_until}",
)
elif stage == DOWNLOAD_STAGE_RETRY_WAITING:
if not _is_expired(task.download_next_retry_at, now):
return False
elif stage in (DOWNLOAD_STAGE_QUEUED, "result_ready"):
pass
else:
return False
old_stage = stage
task.pipeline_stage = DOWNLOAD_STAGE_DOWNLOADING
task.download_started_at = now
task.download_lease_until = _lease_until(now)
task.download_next_retry_at = None
task.download_attempt_count = int(task.download_attempt_count or 0) + 1
task.retry_count = task.download_attempt_count
if not task.download_storage_date_dir:
task.download_storage_date_dir = _task_created_date_dir(task)
await db.commit()
await _register_active_from_task(
task,
check_at=task.download_lease_until,
priority=settings.DOWNLOAD_TASK_PRIORITY_NORMAL,
reason="claim_download_lease",
)
await log_task_event(
task,
event_type="DOWNLOAD_START",
from_stage=old_stage,
to_stage=DOWNLOAD_STAGE_DOWNLOADING,
detail={
"attempt": task.download_attempt_count,
"lease_until": task.download_lease_until,
"download_celery_task_id": task.download_celery_task_id,
},
)
return True
async def _mark_retry_waiting(db: AsyncSession, task: ChatGenerationTask, exc: Exception) -> datetime:
now = _now()
attempt = int(task.download_attempt_count or task.retry_count or 0)
next_retry_at = _retry_at(attempt, now)
error_message = extract_error_message(exc, "下载") if callable(extract_error_message) else str(exc)
retry_celery_task_id = _build_celery_task_id(task.id, attempt=attempt, reason="retry_waiting")
task.pipeline_stage = DOWNLOAD_STAGE_RETRY_WAITING
task.download_celery_task_id = retry_celery_task_id
task.download_next_retry_at = next_retry_at
task.download_lease_until = None
task.download_last_error = error_message
task.retry_count = attempt
await db.commit()
await _register_active_from_task(
task,
check_at=next_retry_at,
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
reason="download_retry_waiting",
)
await log_task_event(
task,
event_type="DOWNLOAD_RETRY_WAITING",
message=error_message,
to_stage=DOWNLOAD_STAGE_RETRY_WAITING,
detail={
"attempt": attempt,
"next_retry_at": next_retry_at,
"download_celery_task_id": retry_celery_task_id,
},
)
return next_retry_at
def _should_final_fail(task: ChatGenerationTask) -> bool:
return int(task.download_attempt_count or task.retry_count or 0) >= int(settings.DOWNLOAD_TASK_MAX_ATTEMPTS or 3)
async def _run(task_id: str):
async with async_session() as db:
result = await db.execute(
@@ -76,46 +269,20 @@ async def _run(task_id: str):
).with_for_update().limit(1)
)
task = result.scalar_one_or_none()
if not task or task.generation_mode != "chatapi_async":
if not task:
return
if task.status != "generating":
return
# 关键修改 3
# 原来只允许 result_ready 进入下载。
# 现在允许 downloading 恢复,但只有“卡住超过 DOWNLOAD_STUCK_SECONDS”的 downloading 才继续。
if task.pipeline_stage == "downloading":
if _is_recent_downloading(task):
# downloading 很新,说明可能有 worker 正在下载,直接跳过,避免并发重复下载。
return
# downloading 已经很久没更新,认为 worker 可能挂了,允许恢复下载。
await log_task_event(
task,
event_type="DOWNLOAD_STUCK_RECOVER",
message=f"downloading 超过 {DOWNLOAD_STUCK_SECONDS} 秒,重新进入下载流程",
)
elif task.pipeline_stage != "result_ready":
claimed = await _claim_download_lease(db, task)
if not claimed:
return
try:
old_stage = task.pipeline_stage
# 无论从 result_ready 进入,还是从 stuck downloading 恢复,都重新标记为 downloading。
task.pipeline_stage = "downloading"
await db.commit()
await log_task_event(
task,
event_type="DOWNLOAD_START",
from_stage=old_stage,
to_stage="downloading",
)
downloaded = await download_generation_result(task)
task = await _reload_task(db, task_id)
if not task:
return
if task.gen_type == "image":
task.image_url = downloaded.url
else:
@@ -123,9 +290,12 @@ async def _run(task_id: str):
task.video_cover_url = downloaded.cover_url
task.status = "completed"
task.pipeline_stage = "done"
task.generated_at = datetime.now(timezone.utc)
task.pipeline_stage = DOWNLOAD_STAGE_DONE
task.generated_at = _now()
task.retry_count = 0
task.download_lease_until = None
task.download_next_retry_at = None
task.download_last_error = None
await record_chat_task_generated_resource(
db,
@@ -138,22 +308,22 @@ async def _run(task_id: str):
)
await db.commit()
await remove_download_active(task.id)
await log_task_event(
task,
event_type="DOWNLOAD_SUCCESS",
to_status="completed",
to_stage="done",
to_stage=DOWNLOAD_STAGE_DONE,
detail={
"resource_url": downloaded.url,
"video_cover_url": downloaded.cover_url,
"file_size_bytes": downloaded.file_size_bytes,
"download_attempt_count": task.download_attempt_count,
},
)
except Exception as exc:
# 关键修改 2
# 异常后先 rollback,再重新查询 task,不继续使用 rollback 前的旧 ORM 对象。
try:
await db.rollback()
except Exception:
@@ -163,33 +333,42 @@ async def _run(task_id: str):
if not task:
return
task.retry_count = (task.retry_count or 0) + 1
if task.retry_count > 3:
if _should_final_fail(task):
error_message = extract_error_message(exc, "下载") if callable(extract_error_message) else str(exc)
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=error_message,
pipeline_stage="download_failed",
pipeline_stage=DOWNLOAD_STAGE_FAILED,
)
task.download_last_error = error_message
task.download_lease_until = None
task.download_next_retry_at = None
await db.commit()
await remove_download_active(task.id)
await log_task_event(
task,
event_type="DOWNLOAD_FAILED",
message=task.error_message,
detail={
"download_attempt_count": task.download_attempt_count,
"max_attempts": settings.DOWNLOAD_TASK_MAX_ATTEMPTS,
},
)
else:
# 下载失败但未超过重试次数,改回 result_ready,等待下一次下载。
# 这样不会卡死在 downloading。
task.pipeline_stage = "result_ready"
await db.commit()
next_retry_at = await _mark_retry_waiting(db, task, exc)
download_generation_result_task.apply_async(
args=[task.id],
countdown=30 * task.retry_count,
)
if celery_app:
delay_seconds = max(1, int((next_retry_at - _now()).total_seconds()))
download_generation_result_task.apply_async(
args=[task.id],
queue=DOWNLOAD_QUEUE,
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
countdown=delay_seconds,
task_id=task.download_celery_task_id,
)
if celery_app:
@@ -140,8 +140,9 @@ async def _run(task_id: str):
await log_task_event(task, event_type="POLL_SUCCESS", to_stage="result_ready")
from app.tasks.generation_download_tasks import download_generation_result_task
download_generation_result_task.delay(task.id)
from app.tasks.generation_download_tasks import enqueue_download_task
await enqueue_download_task(db, task, reason="poll_success_result_ready")
return
if _is_failed(status):
@@ -0,0 +1,46 @@
# app/tasks/generation_recovery_tasks.py
from __future__ import annotations
from typing import Any, Dict
from app.models.base import async_session
from app.tasks.async_runner import run_async
from app.tasks.celery_app import celery_app
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)
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)
if celery_app:
@celery_app.task(name="generation.recover_download_tasks_once")
def recover_download_tasks_once() -> Dict[str, Any]:
return run_async(_run_download_once())
@celery_app.task(name="generation.recover_generation_tasks_once")
def recover_generation_tasks_once() -> Dict[str, Any]:
return run_async(_run_generation_once())
else:
class _DisabledTask:
def delay(self, *args: Any, **kwargs: Any) -> None:
raise RuntimeError("Celery is disabled")
def apply_async(self, *args: Any, **kwargs: Any) -> None:
raise RuntimeError("Celery is disabled")
recover_download_tasks_once = _DisabledTask()
recover_generation_tasks_once = _DisabledTask()