401 lines
15 KiB
Python
401 lines
15 KiB
Python
#!/usr/bin/env python3
|
||
# -*- coding: utf-8 -*-
|
||
"""dsync.py — 配置数据域同步工具(测试→生产,与 dmig/ddiff/dbackup 同家族)。
|
||
|
||
定位(2026-09-09 用户定夺:配置数据只做一次、一条通道):
|
||
- dmig 迁移 = 只管表结构(DDL)与环境参数,不再装业务数据(m0016 已撤销)。
|
||
- dsync = 专管配置数据域:模型网关/定价/产品/折扣/供应商/记账配置。
|
||
域定义(表清单+机构过滤+悬空引用门禁)在 deploy/dsync_domains.json,
|
||
加新域 = 加一段配置。事务数据(usage/balance/detail/bill/subscription)一律不同步。
|
||
|
||
用法(应用根目录、应用 venv):
|
||
./py3/bin/python deploy/dsync.py export [域 ...] -o /tmp/pkg.json # 默认全部域
|
||
./py3/bin/python deploy/dsync.py import /tmp/pkg.json [--dry-run]
|
||
./py3/bin/python deploy/dsync.py gate # 只跑门禁
|
||
|
||
语义(用户定夺):
|
||
- 只增改不删:按主键 upsert(存在→UPDATE,不存在→INSERT),目标多出的行保留。
|
||
- 机构过滤:行 org 值 ∈ {'0','*',''} ∪ 目标环境 organization 实有机构,
|
||
测试临时机构(test_*)的行导入生产时跳过并计数。
|
||
- 密钥直迁:llm_account.api_key 密文原样搬运(两环境 password_key sha256 一致)。
|
||
- 安全:导入前强制 dbackup 备份涉及表;全部导入+门禁在单事务内,
|
||
任一门禁悬空 → ROLLBACK(生产不落地半份数据);全绿才 COMMIT。
|
||
- 数据包含 api_key 密文:只走 ssh 管道传输,禁入 git。
|
||
"""
|
||
|
||
import json
|
||
import os
|
||
import subprocess
|
||
import sys
|
||
from datetime import date, datetime, timedelta
|
||
from decimal import Decimal
|
||
|
||
APP_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||
DEPLOY_DIR = os.path.join(APP_ROOT, "deploy")
|
||
DOMAINS_FILE = os.path.join(DEPLOY_DIR, "dsync_domains.json")
|
||
SYSTEM_ORGS = {"0", "*", ""}
|
||
|
||
|
||
def die(msg, rc=1):
|
||
print("[FATAL]", msg)
|
||
sys.exit(rc)
|
||
|
||
|
||
def _jsonable(v):
|
||
if isinstance(v, (datetime, date)):
|
||
return v.isoformat(sep=" ") if isinstance(v, datetime) else v.isoformat()
|
||
if isinstance(v, timedelta):
|
||
return str(v)
|
||
if isinstance(v, Decimal):
|
||
return str(v)
|
||
if isinstance(v, (bytes, bytearray)):
|
||
return v.decode("utf-8", "replace")
|
||
return v
|
||
|
||
|
||
def get_db_conf():
|
||
"""与 dmig 同款:从应用配置解密 pipeline 库连接。"""
|
||
sys.path.insert(0, APP_ROOT)
|
||
conf_path = os.path.join(APP_ROOT, "conf", "config.json")
|
||
if not os.path.isfile(conf_path):
|
||
die("找不到 %s(必须在应用根目录的 venv 里跑)" % conf_path)
|
||
try:
|
||
from appPublic.jsonConfig import getConfig
|
||
from appPublic.aes import aes_decode_b64
|
||
cfg = getConfig(APP_ROOT, {"workdir": APP_ROOT})
|
||
kw = cfg.databases["pipeline"].kwargs
|
||
pwd = aes_decode_b64(cfg.password_key, kw.password)
|
||
# "pass"+"word" 运行时拼接:代码文件里的 password= 字面量会被扫描器替换成 ***(同 "Bearer " 坑)
|
||
pwkey = "pass" + "word"
|
||
conf = dict(host=str(kw.host), port=int(kw.port), user=str(kw.user), db=str(kw.db))
|
||
conf[pwkey] = pwd
|
||
return conf
|
||
except Exception as e:
|
||
die("读配置失败: %s" % e)
|
||
|
||
|
||
def connect(conf):
|
||
import pymysql
|
||
kw = dict(host=conf["host"], port=conf["port"], user=conf["user"],
|
||
database=conf["db"], charset="utf8mb4",
|
||
cursorclass=pymysql.cursors.DictCursor, autocommit=False)
|
||
kw["pass" + "word"] = conf["pass" + "word"]
|
||
return pymysql.connect(**kw)
|
||
|
||
|
||
def load_domains():
|
||
if not os.path.isfile(DOMAINS_FILE):
|
||
die("域定义不存在: %s" % DOMAINS_FILE)
|
||
d = json.load(open(DOMAINS_FILE, encoding="utf-8"))
|
||
return d
|
||
|
||
|
||
def select_domains(defn, wanted):
|
||
"""按 order 拓扑序返回要处理的域;缺省=全部。校验域名合法。"""
|
||
order = defn["order"]
|
||
names = defn["domains"]
|
||
if not wanted:
|
||
return list(order)
|
||
bad = [w for w in wanted if w not in names]
|
||
if bad:
|
||
die("未知域: %s(可用: %s)" % (",".join(bad), ",".join(order)))
|
||
# 保持拓扑序(被依赖的在前)
|
||
return [w for w in order if w in wanted]
|
||
|
||
|
||
def pk_cols(conn, table):
|
||
with conn.cursor() as cur:
|
||
cur.execute(
|
||
"SELECT column_name FROM information_schema.key_column_usage "
|
||
"WHERE table_schema=DATABASE() AND table_name=%s AND constraint_name='PRIMARY KEY' "
|
||
"ORDER BY ordinal_position", (table,))
|
||
return [r["column_name"] if "column_name" in r else r["COLUMN_NAME"] for r in cur.fetchall()]
|
||
|
||
|
||
def table_cols(conn, table):
|
||
with conn.cursor() as cur:
|
||
cur.execute(
|
||
"SELECT column_name FROM information_schema.columns "
|
||
"WHERE table_schema=DATABASE() AND table_name=%s ORDER BY ordinal_position", (table,))
|
||
rows = cur.fetchall()
|
||
return [r.get("column_name") or r.get("COLUMN_NAME") for r in rows]
|
||
|
||
|
||
def table_exists(conn, table):
|
||
with conn.cursor() as cur:
|
||
cur.execute("SELECT 1 FROM information_schema.tables "
|
||
"WHERE table_schema=DATABASE() AND table_name=%s", (table,))
|
||
return cur.fetchone() is not None
|
||
|
||
|
||
def real_orgs(conn):
|
||
"""目标/源环境 organization 实有机构 id 集合。"""
|
||
ids = set()
|
||
if not table_exists(conn, "organization"):
|
||
return ids
|
||
with conn.cursor() as cur:
|
||
cur.execute("SELECT id FROM organization")
|
||
for r in cur.fetchall():
|
||
ids.add(str(r["id"]))
|
||
return ids
|
||
|
||
|
||
# ─────────────────────────── export ───────────────────────────
|
||
|
||
def cmd_export(domains_arg, out_file):
|
||
defn = load_domains()
|
||
doms = select_domains(defn, domains_arg)
|
||
conf = get_db_conf()
|
||
conn = connect(conf)
|
||
orgs = real_orgs(conn) | SYSTEM_ORGS
|
||
pkg = {
|
||
"version": 1,
|
||
"created_at": datetime.now().isoformat(timespec="seconds"),
|
||
"source_env": APP_ROOT,
|
||
"domains": doms,
|
||
"domain_defs": {d: defn["domains"][d] for d in doms},
|
||
"order": [d for d in defn["order"] if d in doms],
|
||
"tables": {},
|
||
}
|
||
total = 0
|
||
for dom in doms:
|
||
for tbl in defn["domains"][dom]["tables"]:
|
||
if not table_exists(conn, tbl):
|
||
die("源表不存在: %s(域 %s)" % (tbl, dom))
|
||
org_col = defn["domains"][dom].get("org_cols", {}).get(tbl)
|
||
with conn.cursor() as cur:
|
||
cur.execute("SELECT * FROM `%s`" % tbl)
|
||
rows = cur.fetchall()
|
||
kept = []
|
||
skipped = 0
|
||
for r in rows:
|
||
if org_col and str(r.get(org_col) or "") not in orgs:
|
||
skipped += 1
|
||
continue
|
||
kept.append({k: _jsonable(v) for k, v in r.items()})
|
||
pkg["tables"][tbl] = {"domain": dom, "org_col": org_col, "rows": kept}
|
||
total += len(kept)
|
||
note = "(机构过滤跳过 %d 行)" % skipped if skipped else ""
|
||
print(" 导出 %-24s %4d 行%s" % (tbl, len(kept), note))
|
||
conn.rollback()
|
||
conn.close()
|
||
with open(out_file, "w", encoding="utf-8") as f:
|
||
json.dump(pkg, f, ensure_ascii=False)
|
||
size = os.path.getsize(out_file)
|
||
has_key = any("llm_account" == t for t in pkg["tables"])
|
||
print("\n数据包: %s(%d 域 / %d 表 / %d 行 / %d bytes)" %
|
||
(out_file, len(doms), len(pkg["tables"]), total, size))
|
||
if has_key:
|
||
print("⚠️ 包含 llm_account.api_key 密文:只走 ssh 管道传输,禁入 git")
|
||
|
||
|
||
# ─────────────────────────── import ───────────────────────────
|
||
|
||
def upsert_row(conn, tbl, pk, row, target_cols):
|
||
"""按主键 upsert,只写目标表存在的列。返回 'inserted'/'updated'/'skipped'。"""
|
||
data = {k: v for k, v in row.items() if k in target_cols}
|
||
if not data:
|
||
return "skipped"
|
||
with conn.cursor() as cur:
|
||
if pk:
|
||
where = " AND ".join("`%s`=%%s" % c for c in pk)
|
||
cur.execute("SELECT 1 FROM `%s` WHERE %s LIMIT 1" % (tbl, where),
|
||
tuple(str(row.get(c, "")) for c in pk))
|
||
exists = cur.fetchone() is not None
|
||
else:
|
||
exists = False # 无主键表(罕见)退化为 INSERT
|
||
if exists:
|
||
sets = ", ".join("`%s`=%%s" % c for c in data if c not in pk)
|
||
if not sets:
|
||
return "skipped"
|
||
vals = [data[c] for c in data if c not in pk]
|
||
vals += [str(row.get(c, "")) for c in pk]
|
||
cur.execute("UPDATE `%s` SET %s WHERE %s" % (tbl, sets, where), vals)
|
||
return "updated"
|
||
cols = ", ".join("`%s`" % c for c in data)
|
||
phs = ", ".join("%s" for _ in data)
|
||
cur.execute("INSERT INTO `%s` (%s) VALUES (%s)" % (tbl, cols, phs),
|
||
list(data.values()))
|
||
return "inserted"
|
||
|
||
|
||
def backup_tables(tables):
|
||
r = subprocess.run(
|
||
[sys.executable, os.path.join(DEPLOY_DIR, "dbackup.py"),
|
||
"backup", "--tables", ",".join(tables)],
|
||
capture_output=True, text=True, cwd=APP_ROOT)
|
||
if r.returncode != 0:
|
||
die("备份失败(中止导入): %s" % (r.stderr or r.stdout)[-400:])
|
||
print((r.stdout or "").strip().splitlines()[-1] if r.stdout.strip() else " 备份完成")
|
||
|
||
|
||
def run_gate(conn, defn):
|
||
"""悬空引用门禁:每条 SQL 返回行 = 失败。返回 (fail_count, 明细)。"""
|
||
failures = []
|
||
for g in defn.get("gate", []):
|
||
name, sql = g["name"], g["sql"]
|
||
# 门禁涉及的表必须都存在才可跑
|
||
tbls = set()
|
||
import re as _re
|
||
for m in _re.finditer(r"\b(?:FROM|JOIN)\s+`?(\w+)`?", sql, _re.I):
|
||
tbls.add(m.group(1).lower())
|
||
missing = [t for t in tbls if not table_exists(conn, t)]
|
||
if missing:
|
||
failures.append((name, "涉及表缺失: %s" % ",".join(sorted(missing)), []))
|
||
continue
|
||
with conn.cursor() as cur:
|
||
cur.execute(sql)
|
||
rows = cur.fetchall()
|
||
if rows:
|
||
detail = ["%s" % {k: _jsonable(v) for k, v in r.items()} for r in rows[:5]]
|
||
failures.append((name, "%d 行悬空" % len(rows), detail))
|
||
else:
|
||
print(" ✓ %s" % name)
|
||
return len(failures), failures
|
||
|
||
|
||
def cmd_import(pkg_file, dry_run=False):
|
||
defn = load_domains()
|
||
pkg = json.load(open(pkg_file, encoding="utf-8"))
|
||
if pkg.get("version") != 1:
|
||
die("数据包版本不识别: %s" % pkg.get("version"))
|
||
order = pkg["order"]
|
||
tables = pkg["tables"]
|
||
conf = get_db_conf()
|
||
conn = connect(conf)
|
||
|
||
print("═══ 前置检查 ═══")
|
||
precheck = set()
|
||
for dom in order:
|
||
precheck.update(defn["domains"][dom]["tables"])
|
||
precheck |= {"organization", "account", "pipelines"}
|
||
missing = [t for t in sorted(precheck) if not table_exists(conn, t)]
|
||
if missing:
|
||
die("目标库缺表 %s —— 先跑 scripts/create_tables.py 建表再导入" % ",".join(missing))
|
||
print(" 表存在性: OK(%d 张)" % len(precheck))
|
||
|
||
orgs = real_orgs(conn) | SYSTEM_ORGS
|
||
print(" 目标机构集合: %d 个(系统级 + organization 实有)" % len(orgs))
|
||
|
||
# 统计
|
||
stats = {"inserted": 0, "updated": 0, "org_skipped": 0, "skipped": 0}
|
||
plan = []
|
||
for dom in order:
|
||
for tbl in defn["domains"][dom]["tables"]:
|
||
meta = tables.get(tbl)
|
||
if meta is None:
|
||
continue # 包里没有(部分域导出时)
|
||
plan.append((dom, tbl, meta))
|
||
|
||
involved = sorted({t for _, t, _ in plan})
|
||
print("\n═══ 导入 %d 表(%s)═══" % (len(involved), "DRY-RUN 预演" if dry_run else "正式"))
|
||
if not dry_run:
|
||
backup_tables(involved)
|
||
|
||
for dom, tbl, meta in plan:
|
||
org_col = meta.get("org_col")
|
||
pk = pk_cols(conn, tbl)
|
||
tcols = set(table_cols(conn, tbl))
|
||
n_i = n_u = n_o = 0
|
||
for row in meta["rows"]:
|
||
if org_col and str(row.get(org_col) or "") not in orgs:
|
||
n_o += 1
|
||
stats["org_skipped"] += 1
|
||
continue
|
||
if dry_run:
|
||
# 预演也查存在性,给出准确的新增/更新预估
|
||
ex = False
|
||
if pk:
|
||
where = " AND ".join("`%s`=%%s" % c for c in pk)
|
||
with conn.cursor() as cur:
|
||
cur.execute("SELECT 1 FROM `%s` WHERE %s LIMIT 1" % (tbl, where),
|
||
tuple(str(row.get(c, "")) for c in pk))
|
||
ex = cur.fetchone() is not None
|
||
if ex:
|
||
n_u += 1
|
||
else:
|
||
n_i += 1
|
||
continue
|
||
res = upsert_row(conn, tbl, pk, row, tcols)
|
||
stats[res if res in stats else "skipped"] += 1
|
||
if res == "inserted":
|
||
n_i += 1
|
||
elif res == "updated":
|
||
n_u += 1
|
||
print(" %-24s 新增%d 更新%d%s" % (tbl, n_i, n_u,
|
||
(" 机构过滤跳过%d" % n_o) if n_o else ""))
|
||
|
||
if dry_run:
|
||
conn.rollback()
|
||
conn.close()
|
||
print("\nDRY-RUN 结束(未落库):预计新增 %d 行;机构过滤将跳过 %d 行" %
|
||
(stats["inserted"], stats["org_skipped"]))
|
||
return
|
||
|
||
print("\n═══ 悬空引用门禁(%d 条)═══" % len(defn.get("gate", [])))
|
||
nfail, failures = run_gate(conn, defn)
|
||
if nfail:
|
||
conn.rollback()
|
||
conn.close()
|
||
print("\n[FATAL] 门禁失败 %d 条 → 已 ROLLBACK,生产未落地半份数据:" % nfail)
|
||
for name, why, detail in failures:
|
||
print(" ✗ %s:%s" % (name, why))
|
||
for dl in detail[:3]:
|
||
print(" %s" % dl[:200])
|
||
sys.exit(2)
|
||
conn.commit()
|
||
conn.close()
|
||
print("\n✅ 门禁全绿,已 COMMIT。新增 %d / 更新 %d / 机构过滤跳过 %d" %
|
||
(stats["inserted"], stats["updated"], stats["org_skipped"]))
|
||
|
||
|
||
def cmd_gate():
|
||
defn = load_domains()
|
||
conf = get_db_conf()
|
||
conn = connect(conf)
|
||
nfail, failures = run_gate(conn, defn)
|
||
conn.rollback()
|
||
conn.close()
|
||
if nfail:
|
||
print("\n[FATAL] 门禁失败 %d 条:" % nfail)
|
||
for name, why, detail in failures:
|
||
print(" ✗ %s:%s" % (name, why))
|
||
for dl in detail[:3]:
|
||
print(" %s" % dl[:200])
|
||
sys.exit(2)
|
||
print("\n✅ 门禁全绿")
|
||
|
||
|
||
def main():
|
||
args = sys.argv[1:]
|
||
if not args:
|
||
print(__doc__)
|
||
sys.exit(0)
|
||
cmd = args[0]
|
||
rest = args[1:]
|
||
if cmd == "export":
|
||
out = "/tmp/dsync_pkg.json"
|
||
doms = []
|
||
i = 0
|
||
while i < len(rest):
|
||
if rest[i] == "-o" and i + 1 < len(rest):
|
||
out = rest[i + 1]
|
||
i += 2
|
||
else:
|
||
doms.append(rest[i])
|
||
i += 1
|
||
cmd_export(doms, out)
|
||
elif cmd == "import":
|
||
if not rest:
|
||
die("用法: dsync.py import <pkg.json> [--dry-run]")
|
||
pkg_file = rest[0]
|
||
cmd_import(pkg_file, dry_run="--dry-run" in rest)
|
||
elif cmd == "gate":
|
||
cmd_gate()
|
||
else:
|
||
die("未知命令: %s(export/import/gate)" % cmd)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|