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

This commit is contained in:
2026-06-04 16:28:59 +08:00
parent 450e509b1a
commit 8d43d5871c
25 changed files with 1936 additions and 302 deletions
+33 -9
View File
@@ -2,7 +2,7 @@ from datetime import datetime, timedelta, timezone
import bcrypt
import jwt
from sqlalchemy import select
from sqlalchemy import or_, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
@@ -13,8 +13,13 @@ def hash_password(plain: str) -> str:
return bcrypt.hashpw(plain.encode(), bcrypt.gensalt()).decode()
def verify_password(plain: str, hashed: str) -> bool:
return bcrypt.checkpw(plain.encode(), hashed.encode())
def verify_password(plain: str, hashed: str | None) -> bool:
if not hashed:
return False
try:
return bcrypt.checkpw(plain.encode(), hashed.encode())
except Exception:
return False
def create_access_token(user_id: str, remember_me: bool = False) -> str:
@@ -35,17 +40,36 @@ def decode_access_token(token: str) -> str | None:
return None
async def get_user_by_username_or_phone(
db: AsyncSession, username_or_phone: str
) -> User | None:
value = (username_or_phone or "").strip()
if not value:
return None
result = await db.execute(
select(User)
.where(or_(User.username == value, User.phone == value))
.limit(1)
)
return result.scalar_one_or_none()
async def get_user_by_phone(db: AsyncSession, phone: str) -> User | None:
result = await db.execute(select(User).where(User.phone == phone).limit(1))
return result.scalar_one_or_none()
async def authenticate_user(
db: AsyncSession, username: str, password: str
) -> User | None:
result = await db.execute(select(User).where(User.username == username).limit(1))
user = result.scalar_one_or_none()
if not user:
# Try phone number lookup for frontend users
result = await db.execute(select(User).where(User.phone == username).limit(1))
user = result.scalar_one_or_none()
user = await get_user_by_username_or_phone(db, username)
if not user or not verify_password(password, user.hashed_password):
return None
if not user.is_active:
return None
return user
def user_must_set_password(user: User | None) -> bool:
return bool(user and user.user_type == "frontend" and not user.hashed_password)
@@ -0,0 +1,338 @@
# app/services/celery_download_recovery_service.py
from __future__ import annotations
import inspect
import json
import logging
from datetime import datetime, timezone
from typing import Any, Dict, Iterable, List, Optional, Union
from app.config import settings
try:
from redis.exceptions import RedisError
except ImportError:
RedisError = RuntimeError # type: ignore[assignment]
logger = logging.getLogger("video_gen")
_redis_client: Optional[Any] = None
def utc_now() -> datetime:
return datetime.now(timezone.utc)
def ensure_aware_utc(value: Optional[datetime]) -> Optional[datetime]:
if value is None:
return None
if value.tzinfo is None:
return value.replace(tzinfo=timezone.utc)
return value.astimezone(timezone.utc)
def datetime_to_epoch(value: Optional[datetime]) -> int:
checked_value = ensure_aware_utc(value) or utc_now()
return int(checked_value.timestamp())
def _registry_redis_url() -> str:
return settings.CELERY_BROKER_URL or settings.REDIS_URL or ""
async def get_registry_redis() -> Optional[Any]:
global _redis_client
if _redis_client is not None:
return _redis_client
redis_url = _registry_redis_url()
if not redis_url:
return None
try:
from redis.asyncio import Redis
except ImportError as exc:
logger.warning(
"下载容灾 Redis 注册表不可用,redis 依赖未安装。error=%s",
exc,
)
return None
try:
redis_client = Redis.from_url(redis_url, decode_responses=True)
await redis_client.ping()
_redis_client = redis_client
return _redis_client
except (RedisError, OSError, RuntimeError) as exc:
logger.warning(
"下载容灾 Redis 注册表不可用,降级为仅 DB 容灾。error=%s",
exc,
)
_redis_client = None
return None
async def close_registry_redis() -> None:
global _redis_client
client = _redis_client
_redis_client = None
if client is None:
return
try:
close_method = getattr(client, "close", None)
if close_method is None:
return
close_result = close_method()
if inspect.isawaitable(close_result):
await close_result
except (RedisError, OSError, RuntimeError) as exc:
logger.debug(
"关闭下载容灾 Redis 注册表连接失败。error=%s",
exc,
)
def build_download_active_payload(
*,
record_id: str,
celery_task_id: Optional[str],
stage: str,
attempt: Optional[int] = None,
queue: str = "gen_result_download",
priority: Optional[int] = None,
enqueue_at: Optional[datetime] = None,
started_at: Optional[datetime] = None,
updated_at: Optional[datetime] = None,
lease_until: Optional[datetime] = None,
next_retry_at: Optional[datetime] = None,
check_at: Optional[datetime] = None,
reason: Optional[str] = None,
) -> Dict[str, Any]:
now = utc_now()
checked_updated_at = ensure_aware_utc(updated_at) or now
checked_enqueue_at = ensure_aware_utc(enqueue_at)
checked_started_at = ensure_aware_utc(started_at)
checked_lease_until = ensure_aware_utc(lease_until)
checked_next_retry_at = ensure_aware_utc(next_retry_at)
checked_check_at = ensure_aware_utc(check_at)
return {
"record_id": record_id,
"celery_task_id": celery_task_id,
"stage": stage,
"attempt": int(attempt or 0),
"queue": queue,
"priority": priority,
"enqueue_at": (
datetime_to_epoch(checked_enqueue_at)
if checked_enqueue_at
else None
),
"started_at": (
datetime_to_epoch(checked_started_at)
if checked_started_at
else None
),
"updated_at": datetime_to_epoch(checked_updated_at),
"lease_until": (
datetime_to_epoch(checked_lease_until)
if checked_lease_until
else None
),
"next_retry_at": (
datetime_to_epoch(checked_next_retry_at)
if checked_next_retry_at
else None
),
"check_at": (
datetime_to_epoch(checked_check_at)
if checked_check_at
else None
),
"reason": reason,
}
async def upsert_download_active(
*,
record_id: str,
payload: Dict[str, Any],
check_at: Optional[Union[datetime, int, float]],
) -> None:
redis = await get_registry_redis()
if redis is None:
return
if isinstance(check_at, datetime):
score = datetime_to_epoch(check_at)
elif check_at is None:
score = datetime_to_epoch(utc_now())
else:
score = int(float(check_at))
updated_payload = dict(payload)
updated_payload["check_at"] = score
try:
pipe: Any = redis.pipeline(transaction=True)
pipe.hset(
settings.DOWNLOAD_ACTIVE_REDIS_HASH_KEY,
record_id,
json.dumps(updated_payload, ensure_ascii=False, default=str),
)
pipe.zadd(
settings.DOWNLOAD_ACTIVE_REDIS_ZSET_KEY,
{record_id: score},
)
await pipe.execute()
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning(
"写入下载容灾 Redis 注册表失败。record_id=%s, error=%s",
record_id,
exc,
)
async def remove_download_active(record_id: str) -> None:
redis = await get_registry_redis()
if redis is None:
return
try:
pipe: Any = redis.pipeline(transaction=True)
pipe.hdel(settings.DOWNLOAD_ACTIVE_REDIS_HASH_KEY, record_id)
pipe.zrem(settings.DOWNLOAD_ACTIVE_REDIS_ZSET_KEY, record_id)
await pipe.execute()
except (RedisError, OSError, RuntimeError) as exc:
logger.warning(
"删除下载容灾 Redis 注册表失败。record_id=%s, error=%s",
record_id,
exc,
)
async def get_due_download_record_ids(
*,
limit: Optional[int] = None,
now: Optional[datetime] = None,
) -> List[str]:
redis = await get_registry_redis()
if redis is None:
return []
batch_limit = int(limit or settings.DOWNLOAD_RECOVERY_BATCH_SIZE or 100)
score = datetime_to_epoch(now or utc_now())
try:
result = await redis.zrangebyscore(
settings.DOWNLOAD_ACTIVE_REDIS_ZSET_KEY,
min="-inf",
max=score,
start=0,
num=batch_limit,
)
return [str(item) for item in result]
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning(
"扫描下载容灾 Redis ZSet 失败。error=%s",
exc,
)
return []
async def get_download_active_payloads(
record_ids: Iterable[str],
) -> Dict[str, Dict[str, Any]]:
cleaned_record_ids = [str(item) for item in record_ids if item]
if not cleaned_record_ids:
return {}
redis = await get_registry_redis()
if redis is None:
return {}
try:
raw_values = await redis.hmget(
settings.DOWNLOAD_ACTIVE_REDIS_HASH_KEY,
cleaned_record_ids,
)
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning(
"读取下载容灾 Redis Hash 失败。error=%s",
exc,
)
return {}
result: Dict[str, Dict[str, Any]] = {}
for record_id, raw in zip(cleaned_record_ids, raw_values):
if not raw:
continue
try:
value = json.loads(raw)
except (TypeError, ValueError, json.JSONDecodeError):
continue
if isinstance(value, dict):
result[record_id] = value
return result
async def postpone_download_active_check(
*,
record_id: str,
payload: Optional[Dict[str, Any]] = None,
check_at: Optional[Union[datetime, int, float]] = None,
) -> None:
redis = await get_registry_redis()
if redis is None:
return
if isinstance(check_at, datetime):
score = datetime_to_epoch(check_at)
elif check_at is None:
score = datetime_to_epoch(utc_now()) + int(
settings.DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS or 300
)
else:
score = int(float(check_at))
try:
pipe: Any = redis.pipeline(transaction=True)
pipe.zadd(
settings.DOWNLOAD_ACTIVE_REDIS_ZSET_KEY,
{record_id: score},
)
if payload is not None:
updated_payload = dict(payload)
updated_payload["check_at"] = score
updated_payload["updated_at"] = datetime_to_epoch(utc_now())
pipe.hset(
settings.DOWNLOAD_ACTIVE_REDIS_HASH_KEY,
record_id,
json.dumps(
updated_payload,
ensure_ascii=False,
default=str,
),
)
await pipe.execute()
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning(
"刷新下载容灾 Redis 检查时间失败。record_id=%s, error=%s",
record_id,
exc,
)
@@ -0,0 +1,17 @@
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.credit_ratio import CreditRatio
async def list_all_credit_ratios(db: AsyncSession) -> list[CreditRatio]:
"""获取全部积分比例规则,供客户端只读展示和后台复用。"""
result = await db.execute(
select(CreditRatio).order_by(
CreditRatio.gen_type.asc(),
CreditRatio.model_config_id.asc(),
CreditRatio.resolution.asc(),
CreditRatio.created_at.desc(),
)
)
return list(result.scalars().all())
@@ -1,8 +1,9 @@
from __future__ import annotations
import os
import uuid
from dataclasses import dataclass
from datetime import datetime
from datetime import datetime, timezone
from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask
@@ -24,17 +25,93 @@ class DownloadedGenerationResult:
cover_storage_path: str | None = None
def _to_aware_utc(value: datetime | None) -> datetime | None:
if value is None:
return None
if value.tzinfo is None:
return value.replace(tzinfo=timezone.utc)
return value.astimezone(timezone.utc)
def _build_storage_date_dir(record: ChatGenerationTask) -> str:
fixed = (getattr(record, "download_storage_date_dir", None) or "").strip().strip("/")
if fixed:
return fixed
created_at = _to_aware_utc(getattr(record, "created_at", None)) or datetime.now(timezone.utc)
return created_at.strftime("%Y/%m/%d")
def _make_part_path(final_path: str) -> str:
return f"{final_path}.{uuid.uuid4().hex}.part"
def _is_valid_file(path: str | None) -> bool:
if not path:
return False
try:
return os.path.isfile(path) and os.path.getsize(path) > 0
except OSError:
return False
def _safe_remove(path: str | None) -> None:
if not path:
return
try:
if os.path.exists(path):
os.remove(path)
except OSError:
pass
async def _download_image_atomically(remote_url: str, final_path: str) -> str:
if _is_valid_file(final_path):
return final_path
os.makedirs(os.path.dirname(final_path), exist_ok=True)
part_path = _make_part_path(final_path)
try:
await download_image(remote_url, part_path)
if not _is_valid_file(part_path):
raise RuntimeError("图片下载完成但临时文件为空")
os.replace(part_path, final_path)
return final_path
except Exception:
_safe_remove(part_path)
raise
async def _download_video_atomically(remote_url: str, final_path: str) -> str:
if _is_valid_file(final_path):
return final_path
os.makedirs(os.path.dirname(final_path), exist_ok=True)
part_path = _make_part_path(final_path)
try:
await download_video(remote_url, part_path)
if not _is_valid_file(part_path):
raise RuntimeError("视频下载完成但临时文件为空")
os.replace(part_path, final_path)
return final_path
except Exception:
_safe_remove(part_path)
raise
async def download_generation_result(record: ChatGenerationTask) -> DownloadedGenerationResult:
if not record.remote_result_url:
raise ValueError("缺少远程结果URL")
date_dir = datetime.now().strftime("%Y/%m/%d")
date_dir = _build_storage_date_dir(record)
if record.gen_type == "image":
dest_dir = os.path.join(settings.STORAGE_IMAGE_LOCAL_PATH, date_dir)
os.makedirs(dest_dir, exist_ok=True)
dest = os.path.join(dest_dir, f"{record.id}.png")
async with provider_limit("result_download", settings.RESULT_DOWNLOAD_MAX_CONCURRENCY):
await download_image(record.remote_result_url, dest)
await _download_image_atomically(record.remote_result_url if record.remote_result_url else "", dest)
return DownloadedGenerationResult(
url=f"/generate/images/{date_dir}/{record.id}.png",
storage_path=dest,
@@ -45,8 +122,9 @@ async def download_generation_result(record: ChatGenerationTask) -> DownloadedGe
dest_dir = os.path.join(settings.STORAGE_LOCAL_PATH, date_dir)
os.makedirs(dest_dir, exist_ok=True)
dest = os.path.join(dest_dir, f"{record.id}.mp4")
async with provider_limit("result_download", settings.RESULT_DOWNLOAD_MAX_CONCURRENCY):
await download_video(record.remote_result_url, dest)
await _download_video_atomically(record.remote_result_url if record.remote_result_url else "", dest)
cover_url, cover_storage_path = create_video_cover_for_local_video(
record_id=record.id,
@@ -0,0 +1,378 @@
from __future__ import annotations
import logging
from datetime import datetime, timedelta, timezone
from typing import Any, Dict
from sqlalchemy import or_, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask
from app.services.celery_download_recovery_service import (
ensure_aware_utc,
get_download_active_payloads,
get_due_download_record_ids,
postpone_download_active_check,
remove_download_active,
)
from app.services.generation_log_service import log_task_event
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
logger = logging.getLogger("video_gen")
def _now() -> datetime:
return datetime.now(timezone.utc)
def _is_expired(value: datetime | None, now: datetime | None = None) -> bool:
checked = ensure_aware_utc(value)
if checked is None:
return True
return checked <= (now or _now())
def _queue_timeout_at(task: ChatGenerationTask, now: datetime | None = None) -> datetime:
current_time = now or _now()
enqueued_at = ensure_aware_utc(task.download_enqueued_at)
if enqueued_at is None:
return current_time
return enqueued_at + timedelta(seconds=int(settings.DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS or 300))
def _is_queue_timeout(task: ChatGenerationTask, now: datetime | None = None) -> bool:
current_time = now or _now()
return _queue_timeout_at(task, current_time) <= current_time
def _is_final_task_state(task: ChatGenerationTask) -> bool:
return task.status in ("completed", "failed") or task.pipeline_stage in (
"done",
"failed",
"timeout",
"download_failed",
)
async def recover_one_download_task(
db: AsyncSession,
task: ChatGenerationTask,
*,
payload: dict[str, Any] | None = None,
source: str = "startup_db",
) -> str:
from app.tasks.generation_download_tasks import (
DOWNLOAD_STAGE_DOWNLOADING,
DOWNLOAD_STAGE_QUEUED,
DOWNLOAD_STAGE_RETRY_WAITING,
enqueue_download_task,
)
current_time = _now()
if not task:
return "skip_missing_task"
if task.generation_mode != "chatapi_async":
await remove_download_active(task.id)
return "clean_invalid_mode"
if _is_final_task_state(task):
await remove_download_active(task.id)
return "clean_final_state"
if task.status != "generating":
await remove_download_active(task.id)
return "clean_not_generating"
if not task.remote_result_url:
return "skip_no_remote_result_url"
stage = task.pipeline_stage
redis_payload = payload or {}
if stage == "result_ready":
await log_task_event(
task,
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
message=f"{source} 发现 result_ready 未完成下载,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_result_ready",
)
return "recover_result_ready"
if stage == DOWNLOAD_STAGE_QUEUED:
if _is_queue_timeout(task, current_time):
await log_task_event(
task,
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
message=f"{source} 发现 download_queued 长时间未消费,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_download_queued_timeout",
)
return "recover_queued_timeout"
await postpone_download_active_check(
record_id=task.id,
payload=payload,
check_at=_queue_timeout_at(task, current_time),
)
return "skip_queued_not_timeout"
if stage == DOWNLOAD_STAGE_DOWNLOADING:
if _is_expired(task.download_lease_until, current_time):
await log_task_event(
task,
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
message=f"{source} 发现 downloading lease 过期,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_downloading_lease_expired",
)
return "recover_downloading_expired"
await postpone_download_active_check(
record_id=task.id,
payload=payload,
check_at=task.download_lease_until,
)
return "skip_downloading_alive"
if stage == DOWNLOAD_STAGE_RETRY_WAITING:
if _is_expired(task.download_next_retry_at, current_time):
await log_task_event(
task,
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
message=f"{source} 发现 retry_waiting 到期,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_retry_waiting_due",
)
return "recover_retry_due"
await postpone_download_active_check(
record_id=task.id,
payload=payload,
check_at=task.download_next_retry_at,
)
return "skip_retry_waiting_not_due"
return f"skip_stage_{stage}"
async def recover_download_tasks_once(db: AsyncSession) -> dict[str, Any]:
"""启动时下载容灾扫描。
先按 Redis active_index 找到到期下载任务;Redis 不可用或索引丢失时,
再通过 DB fallback 扫描 result_ready/download_* 状态,避免任务永久卡住。
"""
checked_ids: set[str] = set()
results: dict[str, int] = {}
due_ids = await get_due_download_record_ids(
limit=settings.DOWNLOAD_RECOVERY_BATCH_SIZE,
)
payloads = await get_download_active_payloads(due_ids)
for task_id in due_ids:
result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
task = result.scalar_one_or_none()
if task is None:
await remove_download_active(task_id)
action = "clean_missing_task"
else:
checked_ids.add(task.id)
action = await recover_one_download_task(
db,
task,
payload=payloads.get(task_id),
source="startup_redis",
)
results[action] = results.get(action, 0) + 1
# DB fallback:不依赖 Redis active 注册表。
fallback_result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.deleted_at.is_(None),
ChatGenerationTask.generation_mode == "chatapi_async",
ChatGenerationTask.status == "generating",
ChatGenerationTask.remote_result_url.is_not(None),
ChatGenerationTask.pipeline_stage.in_(
["result_ready", "download_queued", "downloading", "retry_waiting"]
),
)
.order_by(ChatGenerationTask.updated_at.asc())
.limit(int(settings.DOWNLOAD_RECOVERY_BATCH_SIZE or 100))
.with_for_update(skip_locked=True)
)
fallback_tasks = fallback_result.scalars().all()
for task in fallback_tasks:
if task.id in checked_ids:
continue
action = await recover_one_download_task(
db,
task,
payload=None,
source="startup_db",
)
results[action] = results.get(action, 0) + 1
checked_ids.add(task.id)
return {"checked": len(checked_ids), "results": results}
async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
"""启动时生成链路容灾扫描。
只在 Celery worker 启动时跑一次,不引入 beat,不新增第四条启动命令。
用于把 queued/creating/waiting_remote/polling/result_ready 等中间态重新投递到现有三个队列。
"""
from app.tasks.generation_create_tasks import chatapi_create_generation_task
from app.tasks.generation_download_tasks import enqueue_download_task
from app.tasks.generation_poll_tasks import poll_generation_task
current_time = _now()
results: dict[str, int] = {}
query_result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.deleted_at.is_(None),
ChatGenerationTask.generation_mode == "chatapi_async",
ChatGenerationTask.status == "generating",
ChatGenerationTask.pipeline_stage.in_(
[
"queued",
"preparing",
"creating_provider_task",
"waiting_remote",
"polling",
"result_ready",
]
),
)
.order_by(ChatGenerationTask.updated_at.asc())
.limit(int(settings.DOWNLOAD_RECOVERY_BATCH_SIZE or 100))
.with_for_update(skip_locked=True)
)
tasks = query_result.scalars().all()
for task in tasks:
if task.deadline_at and _is_expired(task.deadline_at, current_time):
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message="任务超时",
pipeline_stage="timeout",
)
await db.commit()
await log_task_event(
task,
event_type="TASK_TIMEOUT",
to_status="failed",
to_stage="timeout",
)
action = "mark_timeout"
elif task.pipeline_stage in ("queued", "preparing", "creating_provider_task"):
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message="启动时发现创建阶段任务未完成,恢复投递创建队列",
detail={"pipeline_stage": task.pipeline_stage},
)
chatapi_create_generation_task.apply_async(
args=[task.id],
queue="gen_chatapi_create",
countdown=0,
)
action = "recover_create"
elif task.pipeline_stage in ("waiting_remote", "polling"):
if task.remote_result_url:
await enqueue_download_task(
db,
task,
recover=True,
reason="startup_waiting_remote_has_result",
)
action = "recover_waiting_has_result"
elif task.provider_task_id or task.seedance_task_id:
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message="启动时发现远程等待/轮询阶段任务未完成,恢复投递轮询队列",
detail={"pipeline_stage": task.pipeline_stage},
)
task.pipeline_stage = "waiting_remote"
await db.commit()
poll_generation_task.apply_async(
args=[task.id],
queue="gen_provider_poll",
countdown=0,
)
action = "recover_poll"
else:
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message="启动时发现任务缺少供应商任务ID,恢复投递创建队列",
detail={"pipeline_stage": task.pipeline_stage},
)
task.pipeline_stage = "queued"
await db.commit()
chatapi_create_generation_task.apply_async(
args=[task.id],
queue="gen_chatapi_create",
countdown=0,
)
action = "recover_create_missing_provider_id"
elif task.pipeline_stage == "result_ready":
if task.remote_result_url:
await enqueue_download_task(
db,
task,
recover=True,
reason="startup_generation_result_ready",
)
action = "recover_result_ready"
else:
action = "skip_result_ready_no_url"
else:
action = f"skip_stage_{task.pipeline_stage}"
results[action] = results.get(action, 0) + 1
# 下载阶段单独跑 DB fallback。
download_result = await recover_download_tasks_once(db)
return {
"checked": len(tasks),
"results": results,
"download_recovery": download_result,
}
+152 -45
View File
@@ -1,79 +1,186 @@
import asyncio
import json
import logging
import random
import time
import httpx
from datetime import datetime
from typing import Any
from app.config import settings
from app.utils.redis import get_redis
logger = logging.getLogger("videogen")
# In-memory fallback for verification codes
# In-memory fallback for local development when Redis is disabled.
_sms_code_store: dict[str, tuple[str, float]] = {}
_sms_send_interval_store: dict[str, float] = {}
_sms_daily_count_store: dict[str, tuple[str, int]] = {}
def _generate_code(length: int = 6) -> str:
return "".join(random.choices("0123456789", k=length))
def _generate_code(length: int | None = None) -> str:
code_length = int(length or settings.SMS_CODE_LENGTH or 4)
code_length = max(4, min(code_length, 8))
return "".join(random.choices("0123456789", k=code_length))
async def send_sms(phone: str, code: str) -> bool:
"""Send SMS verification code. Supports mock mode and real HTTP gateway."""
if settings.SMS_MOCK or not settings.SMS_API_URL:
logger.info(f"[SMS MOCK] To={phone}, Code={code}")
return True
def _normalize_scene(scene: str | None) -> str:
value = (scene or "").strip().lower()
if value not in {"register", "login", "set_password"}:
value = "login"
return value
def _code_key(phone: str, scene: str | None) -> str:
return f"sms_code:{_normalize_scene(scene)}:{phone}"
def _interval_key(phone: str, scene: str | None) -> str:
return f"sms_interval:{_normalize_scene(scene)}:{phone}"
def _daily_key(phone: str, scene: str | None) -> str:
today = datetime.now().strftime("%Y%m%d")
return f"sms_daily:{_normalize_scene(scene)}:{phone}:{today}"
def _today() -> str:
return datetime.now().strftime("%Y%m%d")
async def _check_send_limit(phone: str, scene: str) -> None:
"""发送频控:同场景同手机号间隔限制 + 每日次数限制。"""
redis = get_redis()
interval_seconds = int(settings.SMS_SEND_INTERVAL_SECONDS or 60)
daily_limit = int(settings.SMS_DAILY_LIMIT or 20)
if redis:
interval_key = _interval_key(phone, scene)
if await redis.get(interval_key):
raise ValueError(f"短信发送过于频繁,请{interval_seconds}秒后再试")
daily_key = _daily_key(phone, scene)
count = await redis.incr(daily_key)
if count == 1:
await redis.expire(daily_key, 24 * 60 * 60)
if count > daily_limit:
raise ValueError("今日短信发送次数已达上限,请明天再试")
await redis.setex(interval_key, interval_seconds, "1")
return
now_ts = time.time()
interval_key = _interval_key(phone, scene)
last_send_at = _sms_send_interval_store.get(interval_key)
if last_send_at and now_ts - last_send_at < interval_seconds:
raise ValueError(f"短信发送过于频繁,请{interval_seconds}秒后再试")
_sms_send_interval_store[interval_key] = now_ts
daily_key = f"{_normalize_scene(scene)}:{phone}"
current_day = _today()
stored_day, count = _sms_daily_count_store.get(daily_key, (current_day, 0))
if stored_day != current_day:
stored_day, count = current_day, 0
count += 1
_sms_daily_count_store[daily_key] = (stored_day, count)
if count > daily_limit:
raise ValueError("今日短信发送次数已达上限,请明天再试")
def _send_volc_sms_sync(phone: str, code: str) -> dict[str, Any]:
from volcengine.sms.SmsService import SmsService
if not settings.VOLC_SMS_ACCESS_KEY_ID or not settings.VOLC_SMS_SECRET_ACCESS_KEY:
raise RuntimeError("火山短信 AK/SK 未配置")
if not settings.VOLC_SMS_ACCOUNT:
raise RuntimeError("火山短信消息组ID VOLC_SMS_ACCOUNT 未配置")
if not settings.VOLC_SMS_TEMPLATE_ID:
raise RuntimeError("火山短信模板ID VOLC_SMS_TEMPLATE_ID 未配置")
if not settings.VOLC_SMS_SIGN:
raise RuntimeError("火山短信签名 VOLC_SMS_SIGN 未配置")
sms_service = SmsService()
sms_service.set_ak(settings.VOLC_SMS_ACCESS_KEY_ID)
sms_service.set_sk(settings.VOLC_SMS_SECRET_ACCESS_KEY)
body = {
"SmsAccount": settings.VOLC_SMS_ACCOUNT,
"Sign": settings.VOLC_SMS_SIGN,
"TemplateID": settings.VOLC_SMS_TEMPLATE_ID,
"TemplateParam": json.dumps({"xxxx": str(code)}, ensure_ascii=False, separators=(",", ":")),
"Tag": f"{phone}:{int(time.time())}",
"PhoneNumbers": phone,
}
raw_resp = sms_service.send_sms(json.dumps(body, ensure_ascii=False, separators=(",", ":")))
if isinstance(raw_resp, str):
try:
resp: dict[str, Any] = json.loads(raw_resp)
except json.JSONDecodeError:
resp = {"raw": raw_resp}
elif isinstance(raw_resp, dict):
resp = raw_resp
else:
resp = {"raw": raw_resp}
error = (resp.get("ResponseMetadata") or {}).get("Error") if isinstance(resp, dict) else None
if error:
raise RuntimeError(f"火山短信发送失败:{error.get('Code')} {error.get('Message')}")
return resp
async def send_sms(phone: str, code: str, scene: str = "login") -> bool:
"""发送短信验证码。
SMS_MOCK=true 时只写日志,方便本地调试;否则使用火山引擎短信 SDK。
火山 SDK 为同步调用,这里放到线程中执行,避免阻塞 FastAPI event loop。
"""
scene = _normalize_scene(scene)
try:
payload = {
"phone": phone,
"code": code,
"sign_name": settings.SMS_SIGN_NAME,
"template_code": settings.SMS_TEMPLATE_CODE,
}
headers = {
"Authorization": f"Bearer {settings.SMS_API_KEY}",
"Content-Type": "application/json",
}
async with httpx.AsyncClient(timeout=10) as client:
resp = await client.post(settings.SMS_API_URL, json=payload, headers=headers)
resp.raise_for_status()
return True
logger.info("[SMS] scene=%s, To=%s, Code=%s", scene, phone, code)
await asyncio.to_thread(_send_volc_sms_sync, phone, code)
return True
except Exception:
logger.exception(f"SMS send failed for {phone}")
logger.exception("SMS send failed. scene=%s phone=%s", scene, phone)
return False
async def store_sms_code(phone: str, code: str, ttl: int = 300) -> None:
"""Store SMS verification code with TTL (default 5 minutes)."""
async def store_sms_code(phone: str, code: str, scene: str = "login", ttl: int | None = None) -> None:
ttl_seconds = int(ttl or settings.SMS_CODE_TTL_SECONDS or 300)
key = _code_key(phone, scene)
redis = get_redis()
if redis:
await redis.setex(f"sms_code:{phone}", ttl, code)
await redis.setex(key, ttl_seconds, code)
else:
_sms_code_store[phone] = (code, time.time() + ttl)
_sms_code_store[key] = (code, time.time() + ttl_seconds)
async def generate_and_send_sms(phone: str) -> bool:
"""Generate a code, store it, and send it via SMS."""
async def generate_and_send_sms(phone: str, scene: str = "login") -> bool:
scene = _normalize_scene(scene)
await _check_send_limit(phone, scene)
code = _generate_code()
ok = await send_sms(phone, code)
ok = await send_sms(phone, code, scene)
if ok:
await store_sms_code(phone, code)
await store_sms_code(phone, code, scene)
return ok
async def verify_sms_code(phone: str, code: str) -> bool:
"""Verify an SMS verification code."""
async def verify_sms_code(phone: str, code: str, scene: str = "login") -> bool:
key = _code_key(phone, scene)
redis = get_redis()
if redis:
stored = await redis.get(f"sms_code:{phone}")
if stored and stored == code:
await redis.delete(f"sms_code:{phone}")
stored = await redis.get(key)
if isinstance(stored, bytes):
stored = stored.decode()
if stored and str(stored) == str(code):
await redis.delete(key)
return True
return False
else:
entry = _sms_code_store.pop(phone, None)
if entry:
stored_code, expires = entry
if time.time() < expires and stored_code == code:
return True
return False
entry = _sms_code_store.pop(key, None)
if entry:
stored_code, expires = entry
if time.time() < expires and stored_code == code:
return True
return False
@@ -1,13 +1,13 @@
import asyncio
import logging
import os
import shutil
import subprocess
import uuid
from pathlib import Path
from app.config import settings
import asyncio
logger = logging.getLogger("videogen")
@@ -20,6 +20,29 @@ def _clean_cover_format(value: str | None) -> str:
return ext or "jpg"
def _is_valid_file(path: str | None) -> bool:
if not path:
return False
try:
return os.path.isfile(path) and os.path.getsize(path) > 0
except OSError:
return False
def _safe_remove(path: str | None) -> None:
if not path:
return
try:
if os.path.exists(path):
os.remove(path)
except OSError:
pass
def _make_part_path(final_path: str) -> str:
return f"{final_path}.{uuid.uuid4().hex}.part"
def get_ffmpeg_bin() -> str:
"""
获取 ffmpeg 可执行文件路径。
@@ -119,6 +142,34 @@ def generate_video_cover(
return str(output_file)
def generate_video_cover_atomically(
video_path: str,
output_path: str,
seek_time: str = "00:00:01",
width: int = 720,
timeout: int = 15,
) -> str:
if _is_valid_file(output_path):
return output_path
part_path = _make_part_path(output_path)
try:
generate_video_cover(
video_path=video_path,
output_path=part_path,
seek_time=seek_time,
width=width,
timeout=timeout,
)
if not _is_valid_file(part_path):
raise VideoCoverError(f"封面临时文件为空: {part_path}")
os.replace(part_path, output_path)
return output_path
except Exception:
_safe_remove(part_path)
raise
def build_video_cover_path_and_url(record_id: str, date_dir: str) -> tuple[str, str]:
"""
根据记录ID和日期目录生成本地封面文件路径与对外URL。
@@ -158,7 +209,7 @@ def try_generate_video_cover(
cover_timeout = timeout or settings.VIDEO_COVER_TIMEOUT_SECONDS
try:
return generate_video_cover(
return generate_video_cover_atomically(
video_path=video_path,
output_path=output_path,
seek_time=first_seek,
@@ -168,7 +219,7 @@ def try_generate_video_cover(
except Exception as first_exc:
if second_seek and second_seek != first_seek:
try:
return generate_video_cover(
return generate_video_cover_atomically(
video_path=video_path,
output_path=output_path,
seek_time=second_seek,
@@ -222,6 +273,7 @@ def create_video_cover_for_local_video(
return None, None
return cover_url, generated_path
async def async_create_video_cover_for_local_video(
*,
record_id: str,