From 8d43d5871c7dbbdfb6552f80aafe2712afd975fa Mon Sep 17 00:00:00 2001 From: GinHa <15201596918@163.com> Date: Thu, 4 Jun 2026 16:28:59 +0800 Subject: [PATCH] =?UTF-8?q?=E7=81=AB=E5=B1=B1=E5=BC=95=E6=93=8ESMS=20API|c?= =?UTF-8?q?elery=E5=AE=B9=E7=81=BE=E4=BC=98=E5=8C=96|=E7=94=9F=E6=88=90?= =?UTF-8?q?=E6=A8=A1=E5=9E=8B=E5=BC=95=E6=93=8E=E7=A7=AF=E5=88=86=E5=88=97?= =?UTF-8?q?=E8=A1=A8API?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- video-gen-api/.env | 9 +- ...92de_add_sms_password_setup_and_celery_.py | 159 ++++++++ video-gen-api/app/api/v1/auth.py | 260 +++++++----- video-gen-api/app/api/v1/credits.py | 14 + video-gen-api/app/api/v1/sms.py | 62 ++- video-gen-api/app/config.py | 27 ++ video-gen-api/app/dependencies.py | 22 +- .../app/models/chat_generation_task.py | 10 + video-gen-api/app/models/user.py | 10 +- video-gen-api/app/schemas/auth.py | 29 +- video-gen-api/app/schemas/sms.py | 12 +- video-gen-api/app/schemas/user.py | 2 + video-gen-api/app/services/auth.py | 42 +- .../celery_download_recovery_service.py | 338 ++++++++++++++++ .../app/services/credit_ratio_service.py | 17 + .../services/generation_download_service.py | 86 +++- .../services/generation_recovery_service.py | 378 ++++++++++++++++++ video-gen-api/app/services/sms.py | 197 ++++++--- .../app/services/video_cover_service.py | 60 ++- video-gen-api/app/tasks/__init__.py | 8 +- video-gen-api/app/tasks/celery_app.py | 76 +++- .../app/tasks/generation_create_tasks.py | 8 +- .../app/tasks/generation_download_tasks.py | 361 ++++++++++++----- .../app/tasks/generation_poll_tasks.py | 5 +- .../app/tasks/generation_recovery_tasks.py | 46 +++ 25 files changed, 1936 insertions(+), 302 deletions(-) create mode 100644 video-gen-api/alembic/versions/476b259992de_add_sms_password_setup_and_celery_.py create mode 100644 video-gen-api/app/services/celery_download_recovery_service.py create mode 100644 video-gen-api/app/services/credit_ratio_service.py create mode 100644 video-gen-api/app/services/generation_recovery_service.py create mode 100644 video-gen-api/app/tasks/generation_recovery_tasks.py diff --git a/video-gen-api/.env b/video-gen-api/.env index d74dafef..04decd64 100644 --- a/video-gen-api/.env +++ b/video-gen-api/.env @@ -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 \ No newline at end of file +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=纬佳网络科技 \ No newline at end of file diff --git a/video-gen-api/alembic/versions/476b259992de_add_sms_password_setup_and_celery_.py b/video-gen-api/alembic/versions/476b259992de_add_sms_password_setup_and_celery_.py new file mode 100644 index 00000000..170db685 --- /dev/null +++ b/video-gen-api/alembic/versions/476b259992de_add_sms_password_setup_and_celery_.py @@ -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") \ No newline at end of file diff --git a/video-gen-api/app/api/v1/auth.py b/video-gen-api/app/api/v1/auth.py index 1f309877..86751082 100644 --- a/video-gen-api/app/api/v1/auth.py +++ b/video-gen-api/app/api/v1/auth.py @@ -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)} diff --git a/video-gen-api/app/api/v1/credits.py b/video-gen-api/app/api/v1/credits.py index d1fe383a..243fceb1 100644 --- a/video-gen-api/app/api/v1/credits.py +++ b/video-gen-api/app/api/v1/credits.py @@ -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), diff --git a/video-gen-api/app/api/v1/sms.py b/video-gen-api/app/api/v1/sms.py index 7f794abb..2aad1b0a 100644 --- a/video-gen-api/app/api/v1/sms.py +++ b/video-gen-api/app/api/v1/sms.py @@ -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, diff --git a/video-gen-api/app/config.py b/video-gen-api/app/config.py index 00c50746..5b3bf258 100644 --- a/video-gen-api/app/config.py +++ b/video-gen-api/app/config.py @@ -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" diff --git a/video-gen-api/app/dependencies.py b/video-gen-api/app/dependencies.py index 00a8eb2f..76d12cc0 100644 --- a/video-gen-api/app/dependencies.py +++ b/video-gen-api/app/dependencies.py @@ -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( diff --git a/video-gen-api/app/models/chat_generation_task.py b/video-gen-api/app/models/chat_generation_task.py index 3283f589..2bd74af6 100644 --- a/video-gen-api/app/models/chat_generation_task.py +++ b/video-gen-api/app/models/chat_generation_task.py @@ -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) diff --git a/video-gen-api/app/models/user.py b/video-gen-api/app/models/user.py index 6be791c8..573d6d87 100644 --- a/video-gen-api/app/models/user.py +++ b/video-gen-api/app/models/user.py @@ -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 diff --git a/video-gen-api/app/schemas/auth.py b/video-gen-api/app/schemas/auth.py index 922f01d4..a562ba2e 100644 --- a/video-gen-api/app/schemas/auth.py +++ b/video-gen-api/app/schemas/auth.py @@ -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): diff --git a/video-gen-api/app/schemas/sms.py b/video-gen-api/app/schemas/sms.py index 6681427c..d070b1f9 100644 --- a/video-gen-api/app/schemas/sms.py +++ b/video-gen-api/app/schemas/sms.py @@ -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="验证码") diff --git a/video-gen-api/app/schemas/user.py b/video-gen-api/app/schemas/user.py index 6ba00398..63321728 100644 --- a/video-gen-api/app/schemas/user.py +++ b/video-gen-api/app/schemas/user.py @@ -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} diff --git a/video-gen-api/app/services/auth.py b/video-gen-api/app/services/auth.py index 9f7d0b0b..89e26056 100644 --- a/video-gen-api/app/services/auth.py +++ b/video-gen-api/app/services/auth.py @@ -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) diff --git a/video-gen-api/app/services/celery_download_recovery_service.py b/video-gen-api/app/services/celery_download_recovery_service.py new file mode 100644 index 00000000..a74b5f28 --- /dev/null +++ b/video-gen-api/app/services/celery_download_recovery_service.py @@ -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, + ) diff --git a/video-gen-api/app/services/credit_ratio_service.py b/video-gen-api/app/services/credit_ratio_service.py new file mode 100644 index 00000000..a5ab5e07 --- /dev/null +++ b/video-gen-api/app/services/credit_ratio_service.py @@ -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()) diff --git a/video-gen-api/app/services/generation_download_service.py b/video-gen-api/app/services/generation_download_service.py index 74e8beb1..bf3182ef 100644 --- a/video-gen-api/app/services/generation_download_service.py +++ b/video-gen-api/app/services/generation_download_service.py @@ -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, diff --git a/video-gen-api/app/services/generation_recovery_service.py b/video-gen-api/app/services/generation_recovery_service.py new file mode 100644 index 00000000..cb54a932 --- /dev/null +++ b/video-gen-api/app/services/generation_recovery_service.py @@ -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, + } diff --git a/video-gen-api/app/services/sms.py b/video-gen-api/app/services/sms.py index 7190a900..57f22f78 100644 --- a/video-gen-api/app/services/sms.py +++ b/video-gen-api/app/services/sms.py @@ -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 diff --git a/video-gen-api/app/services/video_cover_service.py b/video-gen-api/app/services/video_cover_service.py index 79951f5c..ae17709d 100644 --- a/video-gen-api/app/services/video_cover_service.py +++ b/video-gen-api/app/services/video_cover_service.py @@ -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, diff --git a/video-gen-api/app/tasks/__init__.py b/video-gen-api/app/tasks/__init__.py index f86ca85d..8ff0cbff 100644 --- a/video-gen-api/app/tasks/__init__.py +++ b/video-gen-api/app/tasks/__init__.py @@ -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 diff --git a/video-gen-api/app/tasks/celery_app.py b/video-gen-api/app/tasks/celery_app.py index 546921f4..4c6c59be 100644 --- a/video-gen-api/app/tasks/celery_app.py +++ b/video-gen-api/app/tasks/celery_app.py @@ -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() \ No newline at end of file + close_loop() diff --git a/video-gen-api/app/tasks/generation_create_tasks.py b/video-gen-api/app/tasks/generation_create_tasks.py index 46e30c6d..ea83e2f7 100644 --- a/video-gen-api/app/tasks/generation_create_tasks.py +++ b/video-gen-api/app/tasks/generation_create_tasks.py @@ -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 diff --git a/video-gen-api/app/tasks/generation_download_tasks.py b/video-gen-api/app/tasks/generation_download_tasks.py index 9c974f25..61543671 100644 --- a/video-gen-api/app/tasks/generation_download_tasks.py +++ b/video-gen-api/app/tasks/generation_download_tasks.py @@ -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: diff --git a/video-gen-api/app/tasks/generation_poll_tasks.py b/video-gen-api/app/tasks/generation_poll_tasks.py index 8999d625..a6797593 100644 --- a/video-gen-api/app/tasks/generation_poll_tasks.py +++ b/video-gen-api/app/tasks/generation_poll_tasks.py @@ -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): diff --git a/video-gen-api/app/tasks/generation_recovery_tasks.py b/video-gen-api/app/tasks/generation_recovery_tasks.py new file mode 100644 index 00000000..f27aa64d --- /dev/null +++ b/video-gen-api/app/tasks/generation_recovery_tasks.py @@ -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()