forked from erp-dev/erp
feat: big version, added tasks for backup_database and stock change, added health check api, approve sse (support channel via merchant)
This commit is contained in:
126
api_v1/tasks.py
Normal file
126
api_v1/tasks.py
Normal file
@@ -0,0 +1,126 @@
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
|
||||
from celery import shared_task
|
||||
from django.conf import settings
|
||||
from django.utils import timezone
|
||||
|
||||
from basic_info import models as basic_models
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@shared_task(bind=True)
|
||||
def ping_task(self, message: str = 'ping'):
|
||||
"""
|
||||
最简单的心跳任务,用于验证 Celery worker 是否能够
|
||||
正确消费队列并返回结果。
|
||||
"""
|
||||
payload = {
|
||||
'task_id': self.request.id,
|
||||
'message': message,
|
||||
'timestamp': timezone.now().isoformat(),
|
||||
}
|
||||
logger.info('Celery ping_task 执行成功: %s', payload)
|
||||
return payload
|
||||
|
||||
|
||||
@shared_task(bind=True)
|
||||
def merchant_product_count(self, merchant_id: int):
|
||||
"""
|
||||
计算指定商户下的产品数量,用于演示如何在任务中访问数据库。
|
||||
"""
|
||||
count = basic_models.Product.objects.filter(merchant_id=merchant_id).count()
|
||||
payload = {
|
||||
'task_id': self.request.id,
|
||||
'merchant_id': merchant_id,
|
||||
'product_count': count,
|
||||
'calculated_at': timezone.now().isoformat(),
|
||||
}
|
||||
logger.info(
|
||||
'Celery merchant_product_count 统计完成: merchant=%s count=%s task=%s',
|
||||
merchant_id,
|
||||
count,
|
||||
self.request.id,
|
||||
)
|
||||
return payload
|
||||
|
||||
|
||||
def _ensure_backup_dir(output_dir: str | None) -> Path:
|
||||
base_dir = Path(settings.BASE_DIR)
|
||||
backup_dir = Path(output_dir) if output_dir else (base_dir / 'data-bak')
|
||||
backup_dir.mkdir(parents=True, exist_ok=True)
|
||||
return backup_dir
|
||||
|
||||
|
||||
def _build_backup_path(backup_dir: Path, filename_prefix: str) -> Path:
|
||||
timestamp = timezone.now().strftime('%Y%m%d-%H%M%S')
|
||||
return backup_dir / f'{filename_prefix}-{timestamp}.sql'
|
||||
|
||||
|
||||
def _run_pg_dump(backup_path: Path):
|
||||
db_settings = settings.DATABASES['default']
|
||||
pg_dump = shutil.which('pg_dump')
|
||||
if not pg_dump:
|
||||
raise RuntimeError('pg_dump 不存在,请确认 PostgreSQL 客户端工具已安装')
|
||||
|
||||
host = db_settings.get('HOST') or 'localhost'
|
||||
port = db_settings.get('PORT') or '5432'
|
||||
user = db_settings.get('USER') or ''
|
||||
name = db_settings['NAME']
|
||||
password = db_settings.get('PASSWORD') or ''
|
||||
|
||||
cmd = [
|
||||
pg_dump,
|
||||
'-h',
|
||||
host,
|
||||
'-p',
|
||||
str(port),
|
||||
'-U',
|
||||
user,
|
||||
'-F',
|
||||
'p',
|
||||
'-d',
|
||||
name,
|
||||
]
|
||||
|
||||
env = os.environ.copy()
|
||||
if password:
|
||||
env['PGPASSWORD'] = password
|
||||
|
||||
with backup_path.open('wb') as stream:
|
||||
subprocess.run(cmd, check=True, stdout=stream, env=env)
|
||||
|
||||
|
||||
def _dump_database_to_sql(backup_path: Path):
|
||||
engine = settings.DATABASES['default']['ENGINE']
|
||||
if 'postgresql' not in engine:
|
||||
raise NotImplementedError('当前项目只支持 PostgreSQL 数据库备份,请检查 DATABASES 配置')
|
||||
_run_pg_dump(backup_path)
|
||||
|
||||
|
||||
@shared_task(bind=True)
|
||||
def backup_database(self, output_dir: str | None = None, filename_prefix: str = 'db-backup'):
|
||||
"""
|
||||
备份当前数据库为 .sql 文件(仅数据,不包含表结构 DDL),存放在项目根目录 data-bak 下。
|
||||
|
||||
参数:
|
||||
output_dir: 可选,指定备份目录(默认 BASE_DIR/data-bak)
|
||||
filename_prefix: 备份文件名前缀
|
||||
"""
|
||||
backup_dir = _ensure_backup_dir(output_dir)
|
||||
backup_path = _build_backup_path(backup_dir, filename_prefix)
|
||||
_dump_database_to_sql(backup_path)
|
||||
|
||||
payload = {
|
||||
'task_id': self.request.id,
|
||||
'backup_path': str(backup_path),
|
||||
'created_at': timezone.now().isoformat(),
|
||||
}
|
||||
logger.info('数据库备份完成: %s', payload)
|
||||
return payload
|
||||
|
||||
Reference in New Issue
Block a user