Files
X-Financial/server/scripts/backfill_approval_tasks.py

362 lines
13 KiB
Python
Raw Normal View History

#!/usr/bin/env python3
from __future__ import annotations
import argparse
import json
import os
import sys
import uuid
from dataclasses import asdict
from datetime import UTC, datetime
from pathlib import Path
from typing import Any
from sqlalchemy import create_engine, text
from sqlalchemy.engine import Connection
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.orm import Session
from sqlalchemy.pool import NullPool
SERVER_DIR = Path(__file__).resolve().parents[1]
SRC_DIR = SERVER_DIR / "src"
if str(SRC_DIR) not in sys.path:
sys.path.insert(0, str(SRC_DIR))
from app.db.maintenance_database_target import ( # noqa: E402
MaintenanceDatabaseTargetError,
validate_maintenance_database_target,
)
from app.db.migration_preflight import ( # noqa: E402
MigrationPreflightError,
validate_migration_state,
)
from app.services.approval_task_backfill import ( # noqa: E402
DEFAULT_BACKFILL_BATCH_SIZE,
MAX_BACKFILL_BATCH_SIZE,
ApprovalTaskBackfillCursor,
ApprovalTaskBackfillService,
)
REQUIRED_ALEMBIC_REVISION = "20260716_0014"
EXIT_SAFETY = 3
EXIT_LOCKED = 4
EXIT_RUNTIME = 6
class BackfillCommandError(RuntimeError):
def __init__(self, message: str, *, code: str, exit_code: int) -> None:
super().__init__(message)
self.code = code
self.exit_code = exit_code
def parse_timestamp(value: str) -> datetime:
normalized = str(value or "").strip()
if normalized.endswith("Z"):
normalized = f"{normalized[:-1]}+00:00"
try:
parsed = datetime.fromisoformat(normalized)
except ValueError as exc:
raise argparse.ArgumentTypeError("必须是带时区的 ISO 8601 时间") from exc
if parsed.tzinfo is None or parsed.utcoffset() is None:
raise argparse.ArgumentTypeError("时间必须显式包含时区")
return parsed.astimezone(UTC)
def positive_int(value: str) -> int:
try:
parsed = int(value)
except ValueError as exc:
raise argparse.ArgumentTypeError("必须是正整数") from exc
if parsed < 1:
raise argparse.ArgumentTypeError("必须是正整数")
return parsed
def batch_size(value: str) -> int:
parsed = positive_int(value)
if parsed > MAX_BACKFILL_BATCH_SIZE:
raise argparse.ArgumentTypeError(f"不能超过 {MAX_BACKFILL_BATCH_SIZE}")
return parsed
def non_empty_text(value: str) -> str:
normalized = str(value or "").strip()
if not normalized:
raise argparse.ArgumentTypeError("不能为空")
return normalized
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description="为历史待审批费用单生成审批任务;默认只读预览。",
)
mode = parser.add_mutually_exclusive_group()
mode.add_argument("--dry-run", action="store_true", help="只读预览(默认)。")
mode.add_argument("--apply", action="store_true", help="显式写入审批任务和审计事件。")
parser.add_argument("--tenant-id", required=True, type=non_empty_text)
parser.add_argument("--created-before", required=True, type=parse_timestamp)
parser.add_argument("--batch-size", type=batch_size, default=DEFAULT_BACKFILL_BATCH_SIZE)
parser.add_argument("--max-claims", type=positive_int)
parser.add_argument("--sample-limit", type=positive_int, default=20)
parser.add_argument("--expected-host", required=True)
parser.add_argument("--expected-database", required=True)
parser.add_argument(
"--confirm-target",
help="apply 时必须精确等于解析后的 host:port/database。",
)
parser.add_argument("--allow-non-disposable-target", action="store_true")
return parser
def _cursor_payload(cursor: ApprovalTaskBackfillCursor | None) -> dict[str, str] | None:
if cursor is None:
return None
return {
"created_at": _isoformat(cursor.created_at),
"claim_id": cursor.claim_id,
}
def _item_payload(item: Any) -> dict[str, Any]:
payload = asdict(item)
payload["disposition"] = item.disposition.value
if item.entered_at is not None:
payload["entered_at"] = _isoformat(item.entered_at)
return payload
def _summary(args: argparse.Namespace, *, target: Any, revision: str) -> dict[str, Any]:
return {
"mode": "apply" if args.apply else "dry-run",
"database": {
"target": target.exact_target,
"url": target.sanitized_url,
"revision": revision,
},
"tenant_id": args.tenant_id,
"created_before": _isoformat(args.created_before),
"inspected": 0,
"eligible": 0,
"existing": 0,
"skipped": 0,
"created": 0,
"batches": 0,
"limited": False,
"last_cursor": None,
"samples": [],
}
def _page_size(configured: int, remaining: int | None) -> int:
return configured if remaining is None else min(configured, remaining)
def _preview(session: Session, args: argparse.Namespace, summary: dict[str, Any]) -> None:
service = ApprovalTaskBackfillService(
session,
tenant_id=args.tenant_id,
created_before=args.created_before,
)
cursor = None
remaining = args.max_claims
last_has_more = False
while remaining is None or remaining > 0:
page = service.preview(
batch_size=_page_size(args.batch_size, remaining),
after=cursor,
)
if not page.items:
break
summary["batches"] += 1
summary["inspected"] += page.inspected
summary["eligible"] += page.eligible
summary["existing"] += page.existing
summary["skipped"] += page.skipped
available = max(0, args.sample_limit - len(summary["samples"]))
summary["samples"].extend(_item_payload(item) for item in page.items[:available])
cursor = page.next_cursor
summary["last_cursor"] = _cursor_payload(cursor)
last_has_more = page.has_more
if remaining is not None:
remaining -= page.inspected
if not page.has_more:
break
summary["limited"] = bool(remaining == 0 and last_has_more)
def _acquire_lock(connection: Connection, tenant_id: str) -> str:
lock_name = f"approval-task-backfill:{tenant_id}"
acquired = connection.scalar(
text("SELECT pg_try_advisory_lock(hashtextextended(:name, 0))"),
{"name": lock_name},
)
connection.commit()
if not acquired:
raise BackfillCommandError(
"同一租户已有审批任务回填正在运行。",
code="advisory_lock_unavailable",
exit_code=EXIT_LOCKED,
)
return lock_name
def _release_lock(connection: Connection, lock_name: str) -> None:
if connection.in_transaction():
connection.rollback()
connection.execute(
text("SELECT pg_advisory_unlock(hashtextextended(:name, 0))"),
{"name": lock_name},
)
connection.commit()
def _apply(connection: Connection, args: argparse.Namespace, summary: dict[str, Any]) -> None:
lock_name = _acquire_lock(connection, args.tenant_id)
run_id = f"approval-task-backfill-{uuid.uuid4().hex}"
backfilled_at = datetime.now(UTC)
summary["run_id"] = run_id
summary["backfilled_at"] = _isoformat(backfilled_at)
cursor = None
remaining = args.max_claims
last_has_more = False
try:
with Session(bind=connection, autoflush=False, expire_on_commit=False) as session:
service = ApprovalTaskBackfillService(
session,
tenant_id=args.tenant_id,
created_before=args.created_before,
)
while remaining is None or remaining > 0:
session.execute(text("SET LOCAL lock_timeout = '5s'"))
result = service.apply_batch(
run_id=run_id,
batch_size=_page_size(args.batch_size, remaining),
after=cursor,
backfilled_at=backfilled_at,
)
if not result.items:
session.rollback()
break
session.commit()
summary["batches"] += 1
summary["inspected"] += result.inspected
summary["created"] += result.created
summary["existing"] += result.existing
summary["skipped"] += result.skipped
available = max(0, args.sample_limit - len(summary["samples"]))
summary["samples"].extend(_item_payload(item) for item in result.items[:available])
cursor = result.next_cursor
summary["last_cursor"] = _cursor_payload(cursor)
last_has_more = result.has_more
if remaining is not None:
remaining -= result.inspected
if not result.has_more:
break
summary["limited"] = bool(remaining == 0 and last_has_more)
finally:
_release_lock(connection, lock_name)
def run(args: argparse.Namespace) -> dict[str, Any]:
if args.apply and not str(args.confirm_target or "").strip():
raise BackfillCommandError(
"--apply 必须提供 --confirm-target 精确确认数据库目标。",
code="confirm_target_required",
exit_code=EXIT_SAFETY,
)
database_url = os.environ.get("DATABASE_URL", "")
target = validate_maintenance_database_target(
database_url,
expected_host=args.expected_host,
expected_database=args.expected_database,
apply=args.apply,
allow_non_disposable=args.allow_non_disposable_target,
confirm_target=args.confirm_target,
)
engine = create_engine(database_url, pool_pre_ping=True, poolclass=NullPool)
try:
with engine.connect() as connection:
database = str(connection.scalar(text("SELECT current_database()")) or "")
if database != target.database:
raise BackfillCommandError(
"连接后的数据库名与 DATABASE_URL 不一致。",
code="connected_database_mismatch",
exit_code=EXIT_SAFETY,
)
state = validate_migration_state(connection)
if state.revision != REQUIRED_ALEMBIC_REVISION:
raise BackfillCommandError(
f"数据库迁移版本必须为 {REQUIRED_ALEMBIC_REVISION}"
f"实际为 {state.revision or 'unversioned/base'}",
code="migration_revision_mismatch",
exit_code=EXIT_SAFETY,
)
connection.rollback()
summary = _summary(args, target=target, revision=state.revision)
with Session(bind=connection, autoflush=False) as session:
_preview(session, args, summary)
session.rollback()
if not args.apply:
summary["would_create"] = summary["eligible"]
return summary
preview = dict(summary)
summary.update(
inspected=0,
existing=0,
skipped=0,
created=0,
batches=0,
limited=False,
last_cursor=None,
samples=[],
preview=preview,
)
_apply(connection, args, summary)
return summary
finally:
engine.dispose()
def _isoformat(value: datetime) -> str:
normalized = value.replace(tzinfo=UTC) if value.tzinfo is None else value.astimezone(UTC)
return normalized.isoformat().replace("+00:00", "Z")
def _error_payload(exc: Exception, *, code: str) -> dict[str, str]:
return {"status": "error", "code": code, "message": str(exc)}
def main(argv: list[str] | None = None) -> int:
args = build_parser().parse_args(argv)
try:
payload = run(args)
except BackfillCommandError as exc:
print(json.dumps(_error_payload(exc, code=exc.code), ensure_ascii=False), file=sys.stderr)
return exc.exit_code
except MaintenanceDatabaseTargetError as exc:
print(json.dumps(_error_payload(exc, code=exc.code), ensure_ascii=False), file=sys.stderr)
return EXIT_SAFETY
except MigrationPreflightError as exc:
print(
json.dumps(
_error_payload(exc, code="migration_preflight_failed"),
ensure_ascii=False,
),
file=sys.stderr,
)
return EXIT_SAFETY
except (OSError, SQLAlchemyError) as exc:
print(
json.dumps(_error_payload(exc, code="database_runtime_error"), ensure_ascii=False),
file=sys.stderr,
)
return EXIT_RUNTIME
print(json.dumps(payload, ensure_ascii=False, indent=2, default=str))
return 0
if __name__ == "__main__":
raise SystemExit(main())