feat(platform): close AI expense value loop
Add tenant-safe value, telemetry, connector, commercial, and production-readiness foundations.
This commit is contained in:
398
server/scripts/backfill_standard_adjustment_savings.py
Normal file
398
server/scripts/backfill_standard_adjustment_savings.py
Normal file
@@ -0,0 +1,398 @@
|
||||
#!/usr/bin/env python3
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import uuid
|
||||
from collections import Counter
|
||||
from dataclasses import asdict
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from alembic.config import Config
|
||||
from alembic.script import ScriptDirectory
|
||||
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.cli.savings_standard_adjustment_backfill import ( # noqa: E402
|
||||
DEFAULT_BATCH_SIZE,
|
||||
MAX_BATCH_SIZE,
|
||||
StandardAdjustmentBackfillCursor,
|
||||
StandardAdjustmentSavingsBackfillService,
|
||||
)
|
||||
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,
|
||||
)
|
||||
|
||||
REQUIRED_ALEMBIC_REVISION = "20260716_0015"
|
||||
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 revision_contains_required(
|
||||
current_revision: str | None,
|
||||
required_revision: str = REQUIRED_ALEMBIC_REVISION,
|
||||
) -> bool:
|
||||
"""确认当前迁移沿 down_revision 链包含回填所需的数据契约。"""
|
||||
|
||||
current = str(current_revision or "").strip()
|
||||
required = str(required_revision or "").strip()
|
||||
if not current or not required:
|
||||
return False
|
||||
|
||||
config = Config(str(SERVER_DIR / "alembic.ini"))
|
||||
script_directory = ScriptDirectory.from_config(config)
|
||||
pending = [current]
|
||||
visited: set[str] = set()
|
||||
while pending:
|
||||
revision_id = pending.pop()
|
||||
if revision_id in visited:
|
||||
continue
|
||||
if revision_id == required:
|
||||
return True
|
||||
visited.add(revision_id)
|
||||
revision = script_directory.get_revision(revision_id)
|
||||
if revision is None:
|
||||
return False
|
||||
down_revision = revision.down_revision
|
||||
if isinstance(down_revision, str):
|
||||
pending.append(down_revision)
|
||||
elif down_revision:
|
||||
pending.extend(str(item) for item in down_revision)
|
||||
return False
|
||||
|
||||
|
||||
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_BATCH_SIZE:
|
||||
raise argparse.ArgumentTypeError(f"不能超过 {MAX_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_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: StandardAdjustmentBackfillCursor | 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
|
||||
for field in ("original_amount", "target_amount", "saving_amount"):
|
||||
if payload[field] is not None:
|
||||
payload[field] = str(payload[field])
|
||||
return payload
|
||||
|
||||
|
||||
def _base_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),
|
||||
"claims_inspected": 0,
|
||||
"flags_inspected": 0,
|
||||
"eligible": 0,
|
||||
"created": 0,
|
||||
"replayed": 0,
|
||||
"skipped": 0,
|
||||
"reasons": {},
|
||||
"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 _merge_page(summary: dict[str, Any], page: Any, *, sample_limit: int) -> None:
|
||||
summary["batches"] += 1
|
||||
summary["claims_inspected"] += page.claims_inspected
|
||||
summary["flags_inspected"] += page.flags_inspected
|
||||
for field in ("eligible", "created", "replayed", "skipped"):
|
||||
if hasattr(page, field):
|
||||
summary[field] += int(getattr(page, field))
|
||||
reasons = Counter(summary["reasons"])
|
||||
reasons.update(page.reasons)
|
||||
summary["reasons"] = dict(sorted(reasons.items()))
|
||||
available = max(0, sample_limit - len(summary["samples"]))
|
||||
summary["samples"].extend(_item_payload(item) for item in page.items[:available])
|
||||
summary["last_cursor"] = _cursor_payload(page.next_cursor)
|
||||
|
||||
|
||||
def _preview(session: Session, args: argparse.Namespace, summary: dict[str, Any]) -> None:
|
||||
service = StandardAdjustmentSavingsBackfillService(
|
||||
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 page.claims_inspected == 0:
|
||||
break
|
||||
_merge_page(summary, page, sample_limit=args.sample_limit)
|
||||
cursor = page.next_cursor
|
||||
last_has_more = page.has_more
|
||||
if remaining is not None:
|
||||
remaining -= page.claims_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"savings-standard-adjustment-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(
|
||||
"同一租户已有标准调整 Savings 回填正在运行。",
|
||||
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"savings-adjustment-{uuid.uuid4().hex}"
|
||||
summary["run_id"] = run_id
|
||||
cursor = None
|
||||
remaining = args.max_claims
|
||||
last_has_more = False
|
||||
try:
|
||||
with Session(bind=connection, autoflush=False, expire_on_commit=False) as session:
|
||||
service = StandardAdjustmentSavingsBackfillService(
|
||||
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,
|
||||
)
|
||||
if result.claims_inspected == 0:
|
||||
session.rollback()
|
||||
break
|
||||
session.commit()
|
||||
_merge_page(summary, result, sample_limit=args.sample_limit)
|
||||
cursor = result.next_cursor
|
||||
last_has_more = result.has_more
|
||||
if remaining is not None:
|
||||
remaining -= result.claims_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 not revision_contains_required(state.revision):
|
||||
raise BackfillCommandError(
|
||||
f"数据库迁移链必须包含 {REQUIRED_ALEMBIC_REVISION},"
|
||||
f"实际为 {state.revision or 'unversioned/base'}。",
|
||||
code="migration_revision_mismatch",
|
||||
exit_code=EXIT_SAFETY,
|
||||
)
|
||||
connection.rollback()
|
||||
summary = _base_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(
|
||||
claims_inspected=0,
|
||||
flags_inspected=0,
|
||||
eligible=0,
|
||||
created=0,
|
||||
replayed=0,
|
||||
skipped=0,
|
||||
reasons={},
|
||||
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, ValueError, RuntimeError) as exc:
|
||||
print(
|
||||
json.dumps(_error_payload(exc, code="backfill_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())
|
||||
Reference in New Issue
Block a user