火山引擎SMS API|celery容灾优化|生成模型引擎积分列表API
This commit is contained in:
+8
-1
@@ -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")
|
||||
@@ -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)}
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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=消息组ID,TemplateID=模板ID,Sign=短信签名内容。
|
||||
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"
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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="验证码")
|
||||
|
||||
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user