Cai
4 days ago bf157a136d9b08c14b4da2997dc5a03b1a1af33d
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
from __future__ import annotations
 
import argparse
import getpass
import json
import os
import re
from pathlib import Path
 
import pymysql
 
 
ROOT = Path(__file__).resolve().parents[1]
DEFAULT_SEED_DIR = ROOT / "data" / "mysql_seed_20260624"
 
 
def split_sql(sql: str) -> list[str]:
    parts: list[str] = []
    buff: list[str] = []
    for line in sql.splitlines():
        buff.append(line)
        if line.strip().endswith(";"):
            stmt = "\n".join(buff).strip()
            if stmt:
                parts.append(stmt[:-1] if stmt.endswith(";") else stmt)
            buff = []
    tail = "\n".join(buff).strip()
    if tail:
        parts.append(tail)
    return parts
 
 
def get_password(args: argparse.Namespace) -> str:
    if args.password:
        return args.password
    if os.environ.get("DARKLINE_MYSQL_PASSWORD"):
        return os.environ["DARKLINE_MYSQL_PASSWORD"]
    return getpass.getpass("MySQL password: ")
 
 
def connect_without_db(args: argparse.Namespace, password: str):
    return pymysql.connect(
        host=args.host,
        port=args.port,
        user=args.user,
        password=password,
        charset="utf8mb4",
        autocommit=True,
    )
 
 
def connect_with_db(args: argparse.Namespace, password: str):
    return pymysql.connect(
        host=args.host,
        port=args.port,
        user=args.user,
        password=password,
        database=args.database,
        charset="utf8mb4",
        autocommit=True,
    )
 
 
def main() -> None:
    parser = argparse.ArgumentParser(description="Import darkline seed data into local MySQL.")
    parser.add_argument("--host", default=os.environ.get("DARKLINE_MYSQL_HOST", "127.0.0.1"))
    parser.add_argument("--port", type=int, default=int(os.environ.get("DARKLINE_MYSQL_PORT", "3306")))
    parser.add_argument("--user", default=os.environ.get("DARKLINE_MYSQL_USER", "root"))
    parser.add_argument("--password", default=None)
    parser.add_argument("--database", default=os.environ.get("DARKLINE_MYSQL_DATABASE", "tianxia"))
    parser.add_argument("--seed-dir", default=str(DEFAULT_SEED_DIR))
    parser.add_argument("--create-database", action="store_true")
    args = parser.parse_args()
 
    seed_dir = Path(args.seed_dir)
    password = get_password(args)
 
    if args.create_database:
        conn = connect_without_db(args, password)
        with conn.cursor() as cur:
            cur.execute(
                f"CREATE DATABASE IF NOT EXISTS `{args.database}` "
                "CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci"
            )
        conn.close()
 
    conn = connect_with_db(args, password)
    with conn.cursor() as cur:
        schema_sql = (seed_dir / "schema.sql").read_text(encoding="utf-8")
        for stmt in split_sql(schema_sql):
            cur.execute(stmt)
 
        table_order = json.loads((seed_dir / "table_order.json").read_text(encoding="utf-8"))
        imported_counts = {}
        for table in table_order:
            jsonl_path = seed_dir / "tables" / f"{table}.jsonl"
            rows = [
                json.loads(line)
                for line in jsonl_path.read_text(encoding="utf-8").splitlines()
                if line.strip()
            ]
            if not rows:
                imported_counts[table] = 0
                continue
            columns = list(rows[0].keys())
            col_sql = ", ".join(f"`{c}`" for c in columns)
            placeholders = ", ".join(["%s"] * len(columns))
            sql = f"REPLACE INTO `{table}` ({col_sql}) VALUES ({placeholders})"
            cur.executemany(sql, [[row.get(c) for c in columns] for row in rows])
            imported_counts[table] = len(rows)
 
        verify_counts = {}
        for table in table_order:
            cur.execute(f"SELECT COUNT(*) FROM `{table}`")
            verify_counts[table] = cur.fetchone()[0]
 
    conn.close()
    print(
        json.dumps(
            {
                "import_status": "PASS",
                "database": args.database,
                "seed_dir": str(seed_dir),
                "imported_counts": imported_counts,
                "verify_counts": verify_counts,
            },
            ensure_ascii=False,
            indent=2,
        )
    )
 
 
if __name__ == "__main__":
    main()