1
0
forked from erp-dev/erp
Files
erpnew/api_v1/tasks.py

444 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import asyncio
import json
import logging
import os
import shutil
import subprocess
from datetime import datetime
from decimal import Decimal, InvalidOperation
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
from api_v1 import models as api_models
from flower.utils import (
fetch_products_from_mingdaoyun,
fetch_customers_from_mingdaoyun,
)
logger = logging.getLogger(__name__)
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
def _get_mdy_merchant():
merchant_id = getattr(settings, 'MDY_MERCHANT_ID', None)
qs = basic_models.Merchant.objects.all()
if merchant_id:
qs = qs.filter(id=merchant_id)
merchant = qs.order_by('id').first()
if not merchant:
raise RuntimeError('未找到用于明道云同步的商户,请先创建商户或配置 MDY_MERCHANT_ID')
return merchant
def _get_mdy_category(merchant):
category_id = getattr(settings, 'MDY_PRODUCT_CATEGORY_ID', None)
qs = basic_models.ProductCategory.objects.filter(merchant=merchant)
if category_id:
qs = qs.filter(id=category_id)
category = qs.order_by('id').first()
if not category:
raise RuntimeError('未找到用于明道云同步的产品类别,请先创建类别或配置 MDY_PRODUCT_CATEGORY_ID')
return category
def _parse_mdy_datetime(value: str | None):
if not value:
return None
try:
dt = datetime.strptime(value, '%Y-%m-%d %H:%M:%S')
except ValueError:
return None
if timezone.is_naive(dt):
dt = timezone.make_aware(dt, timezone.get_current_timezone())
return dt
def _map_unit(unit_label: str | None) -> int:
if not unit_label:
return basic_models.ProductUnitEnum.METER
mapping = {
'': basic_models.ProductUnitEnum.METER,
'': basic_models.ProductUnitEnum.YARD,
'公斤': basic_models.ProductUnitEnum.KG,
'': basic_models.ProductUnitEnum.SEGMENT,
}
return mapping.get(unit_label, basic_models.ProductUnitEnum.METER)
def _ensure_decimal(value: str | None):
if not value:
return None
try:
return Decimal(str(value))
except (InvalidOperation, TypeError, ValueError):
return None
def _ensure_int(value):
if value is None:
return None
try:
return int(value)
except (TypeError, ValueError, InvalidOperation):
# 有些值可能是 Decimal 或字符串数字
try:
return int(Decimal(str(value)))
except Exception:
return None
def _run_fetch(page: int, page_size: int):
return asyncio.run(fetch_products_from_mingdaoyun(page=page, page_size=page_size))
def _run_fetch_customers(page: int, page_size: int):
return asyncio.run(fetch_customers_from_mingdaoyun(page=page, page_size=page_size))
def _upsert_product(product_data, merchant, category):
description = ''
if product_data.detail:
description = json.dumps(product_data.detail, ensure_ascii=False)
defaults = {
'category': category,
'name': product_data.name or product_data.uid,
'color': product_data.color or '',
'description': description,
'unit': _map_unit(product_data.unit),
'from_mdy': True,
}
width_decimal = _ensure_decimal(product_data.width)
if width_decimal is not None:
defaults['width_size'] = width_decimal
pieces_int = _ensure_int(product_data.pieces)
if pieces_int is not None:
defaults['pieces'] = pieces_int
segment_int = _ensure_int(product_data.segment_size)
if segment_int is not None:
defaults['segment_size'] = segment_int
product_obj, created = basic_models.Product.objects.get_or_create(
merchant=merchant,
human_id=product_data.uid,
defaults=defaults,
)
if created:
return True
if not product_obj.from_mdy:
return False
updated = False
for field, value in defaults.items():
if getattr(product_obj, field) != value:
setattr(product_obj, field, value)
updated = True
if updated:
product_obj.save(update_fields=list(defaults.keys()))
return updated
@shared_task(bind=True)
def sync_mdy_products(self, page_size: int = 300, max_pages: int | None = None, max_records: int | None = None):
"""
从明道云同步产品数据(按 ctime 升序分页扫描)。
- 默认从上次同步记录的 page_index 继续翻页
- last_ctime/last_rowid 用于页内游标(避免重复处理)
- max_pages 表示“单次任务最多处理多少页”(不是最大页码)
"""
merchant = _get_mdy_merchant()
category = _get_mdy_category(merchant)
max_records = max_records or 0 # 0 表示不限制
last_sync = api_models.DataSync.objects.filter(
table_name=api_models.DataSync.TableName.PRODUCT
).order_by('-created_at').first()
last_ctime = last_sync.last_ctime if last_sync else None
last_rowid = last_sync.last_rowid if last_sync else ''
start_page_index = last_sync.page_index if last_sync else 1
synced_rows = 0
page_index = max(1, start_page_index)
pages_processed = 0
total_count = 0
latest_ctime = last_ctime
latest_rowid = last_rowid
while True:
# max_pages: 单次任务最多处理多少页(不是“最大页码”)
if max_pages is not None and pages_processed >= max_pages:
break
if max_records and synced_rows >= max_records:
break
products, total = _run_fetch(page_index, page_size)
total_count = total
if not products:
break
hit_max_records = False
for item in products:
if max_records and synced_rows >= max_records:
hit_max_records = True
break
product_ctime = _parse_mdy_datetime(item.created_at)
# 跳过已同步到的游标
if last_ctime and product_ctime:
if product_ctime < last_ctime:
continue
if product_ctime == last_ctime and last_rowid and item.rowid == last_rowid:
continue
changed = _upsert_product(item, merchant, category)
if changed:
synced_rows += 1
if product_ctime:
if latest_ctime is None or product_ctime > latest_ctime:
latest_ctime = product_ctime
latest_rowid = item.rowid
elif product_ctime == latest_ctime:
# 同一秒内可能有多条记录,尽量把 rowid 推进到最后处理的那条
latest_rowid = item.rowid
pages_processed += 1
if hit_max_records:
# 达到单次任务的记录上限:下次从同一页继续(依赖 last_ctime/last_rowid 跳过已处理部分)
break
# 若返回不足一页,说明到尾部,可结束
if len(products) < page_size:
break
page_index += 1
record_last_ctime = latest_ctime or last_ctime
record_last_rowid = latest_rowid or last_rowid
api_models.DataSync.objects.create(
table_name=api_models.DataSync.TableName.PRODUCT,
page_index=page_index,
page_size=page_size,
synced_rows=synced_rows,
total_count=total_count,
last_ctime=record_last_ctime,
last_rowid=record_last_rowid,
note='asc scan',
)
payload = {
'task_id': self.request.id,
'synced_rows': synced_rows,
'page_index': page_index,
'page_size': page_size,
'total_count': total_count,
'last_ctime': record_last_ctime.isoformat() if record_last_ctime else None,
}
logger.info('明道云产品同步完成: %s', payload)
return payload
def _upsert_customer(customer_data, merchant):
mdy_uid = customer_data.uid or customer_data.rowid
if not mdy_uid:
return False
defaults = {
'merchant': merchant,
'name': customer_data.name or customer_data.uid,
'area': customer_data.area or '',
'from_mdy': True,
}
customer_obj, created = basic_models.Customer.objects.get_or_create(
mdy_uid=mdy_uid,
defaults=defaults,
)
if created:
return True
if not customer_obj.from_mdy:
return False
updated_fields = {}
for field, value in defaults.items():
if getattr(customer_obj, field) != value:
updated_fields[field] = value
if updated_fields:
for field, value in updated_fields.items():
setattr(customer_obj, field, value)
customer_obj.save(update_fields=list(updated_fields.keys()))
return True
return False
@shared_task(bind=True)
def sync_mdy_customers(self, page_size: int = 300, max_pages: int | None = None, max_records: int | None = None):
"""
从明道云同步客户数据(按 ctime 升序分页扫描)。
- 默认从上次同步记录的 page_index 继续翻页
- last_ctime/last_rowid 用于页内游标(避免重复处理)
- max_pages 表示“单次任务最多处理多少页”(不是最大页码)
"""
merchant = _get_mdy_merchant()
max_records = max_records or 0 # 0 表示不限制
last_sync = api_models.DataSync.objects.filter(
table_name=api_models.DataSync.TableName.CUSTOMER
).order_by('-created_at').first()
last_ctime = last_sync.last_ctime if last_sync else None
last_rowid = last_sync.last_rowid if last_sync else ''
start_page_index = last_sync.page_index if last_sync else 1
synced_rows = 0
page_index = max(1, start_page_index)
pages_processed = 0
total_count = 0
latest_ctime = last_ctime
latest_rowid = last_rowid
while True:
# max_pages: 单次任务最多处理多少页(不是“最大页码”)
if max_pages is not None and pages_processed >= max_pages:
break
if max_records and synced_rows >= max_records:
break
customers, total = _run_fetch_customers(page_index, page_size)
total_count = total
if not customers:
break
hit_max_records = False
for item in customers:
if max_records and synced_rows >= max_records:
hit_max_records = True
break
record_ctime = _parse_mdy_datetime(item.created_at)
if last_ctime and record_ctime:
if record_ctime < last_ctime:
continue
if record_ctime == last_ctime and last_rowid and item.rowid == last_rowid:
continue
changed = _upsert_customer(item, merchant)
if changed:
synced_rows += 1
if record_ctime:
if latest_ctime is None or record_ctime > latest_ctime:
latest_ctime = record_ctime
latest_rowid = item.rowid
elif record_ctime == latest_ctime:
latest_rowid = item.rowid
pages_processed += 1
if hit_max_records:
break
if len(customers) < page_size:
break
page_index += 1
record_last_ctime = latest_ctime or last_ctime
record_last_rowid = latest_rowid or last_rowid
api_models.DataSync.objects.create(
table_name=api_models.DataSync.TableName.CUSTOMER,
page_index=page_index,
page_size=page_size,
synced_rows=synced_rows,
total_count=total_count,
last_ctime=record_last_ctime,
last_rowid=record_last_rowid,
note='asc scan',
)
payload = {
'task_id': self.request.id,
'synced_rows': synced_rows,
'page_index': page_index,
'page_size': page_size,
'total_count': total_count,
'last_ctime': record_last_ctime.isoformat() if record_last_ctime else None,
}
logger.info('明道云客户同步完成: %s', payload)
return payload