火山引擎SMS API|celery容灾优化|生成模型引擎积分列表API
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user