#!/usr/bin/env python
# -*- coding: utf-8 -*-
r"""
phpbb_mysql2pg.py — phpBB 3.3 ACP 数据库备份转换器（MySQL/MariaDB → PostgreSQL）

将 phpBB 后台（ACP → 维护 → 数据库 → 备份）在 MySQL/MariaDB 下导出的备份文件，
转换为可在 PostgreSQL 环境下通过 ACP「恢复」功能直接导入的 .sql 单文件。

输出格式与 phpBB 3.3 自带 phpbb\\db\\extractor\\postgres_extractor 生成的备份逐句一致：
  * 文件头 "--" 注释（MySQL 备份的 "#" 注释在 PostgreSQL 下非法，必须改写）+ BEGIN TRANSACTION;
  * 每表: DROP TABLE IF EXISTS / [DROP|CREATE SEQUENCE {表}_seq] / CREATE TABLE / CREATE INDEX
  * 数据: COPY {表} (列) FROM stdin; + 制表符分隔数据行 + "\." 结束符
  * 自增列: 数据导入后 SELECT SETVAL('{表}_seq', ...) 同步序列
  * 文件尾 COMMIT;

文件名默认生成为 backup_<10位时间戳>_<16位hex>.sql，符合 phpBB ACP 恢复页的
文件名正则校验（acp_database.php: get_file_list()），可直接放入 phpBB 的
store/ 目录后在 ACP → 维护 → 数据库 → 恢复 中选择。

用法:
  python phpbb_mysql2pg.py <输入备份.sql|.sql.gz|.sql.bz2> [-o 输出文件.sql] [--quiet]

说明:
  * 输入为 phpBB ACP 导出的「完整」或「仅数据」备份；完整备份含建表语句时效果最佳。
  * MySQL 反斜杠转义（\\' \\\\ \\n \\r \\0 \\Z 等）会被完整解码后按 PostgreSQL
    COPY 文本格式重新编码，中文、引号、换行、制表符等均可无损往返。
  * 全文索引（FULLTEXT KEY）在 PostgreSQL 下语法不同，会被跳过；恢复完成后请在
    ACP → 维护 → 搜索索引 处重建（phpBB 默认使用自建搜索索引表，通常无影响）。

实现基于 phpBB 3.3.17 源码（acp_database.php / postgres_extractor.php /
mysql_extractor.php）的逐行比对。
"""

import argparse
import bz2
import gzip
import io
import os
import re
import secrets
import sys
import time

# phpBB ACP 恢复页对备份文件名的校验正则（见 acp_database.php get_file_list()）
ACP_FILENAME_RE = re.compile(r'^backup_(\d{10,})_(?:[a-z\d]{16}|[a-z\d]{32})\.(sql(?:\.(?:gz|bz2))?)$', re.I)


# ----------------------------------------------------------------------------
# 输入读取（自动识别 gzip / bzip2 压缩备份）
# ----------------------------------------------------------------------------

def read_input(path):
    """读取备份文件为文本（UTF-8，保留无法解码的原始字节）。"""
    with open(path, 'rb') as f:
        magic = f.read(4)
    if magic[:2] == b'\x1f\x8b':
        wrapper = gzip.open(path, 'rb')
    elif magic[:3] == b'BZh':
        wrapper = bz2.open(path, 'rb')
    else:
        wrapper = open(path, 'rb')
    with wrapper:
        with io.TextIOWrapper(wrapper, encoding='utf-8', errors='surrogateescape', newline='') as tw:
            text = tw.read()
    # 备份由服务端以 "\n" 写出；若被 Windows 工具转成 CRLF 则归一化。
    # MySQL 转储内不会出现字面 \r（\r 总是被写成两字符 "\r"），此替换安全。
    return text.replace('\r\n', '\n')


# ----------------------------------------------------------------------------
# 语句切分：按 ";\n" 终止符（跳过字符串字面量与反引号标识符内部）
# ----------------------------------------------------------------------------

def split_statements(text):
    stmts = []
    buf = []
    in_str = False
    in_bt = False
    i = 0
    n = len(text)
    while i < n:
        c = text[i]
        buf.append(c)
        if in_str:
            if c == '\\':
                if i + 1 < n:
                    buf.append(text[i + 1])
                    i += 1
            elif c == "'":
                if i + 1 < n and text[i + 1] == "'":
                    buf.append("'")          # '' 双写转义为一个单引号
                    i += 1
                else:
                    in_str = False
        elif in_bt:
            if c == '`':
                in_bt = False
        else:
            if c == "'":
                in_str = True
            elif c == '`':
                in_bt = True
            elif c == ';' and i + 1 < n and text[i + 1] == '\n':
                stmts.append(''.join(buf))
                buf = []
                i += 1                      # 跳过换行符本身
        i += 1
    if buf:
        tail = ''.join(buf).strip()
        if tail:
            stmts.append(''.join(buf))
    return stmts


def split_comments(stmt):
    """将语句块拆为（注释行列表, 代码部分）。'#' 注释在 PostgreSQL 下非法。"""
    comments, code_lines = [], []
    for ln in stmt.split('\n'):
        if ln.lstrip().startswith('#') or ln.lstrip().startswith('--'):
            comments.append(ln)
        else:
            code_lines.append(ln)
    return comments, '\n'.join(code_lines).strip()


# ----------------------------------------------------------------------------
# MySQL 字符串解码
# ----------------------------------------------------------------------------

_MYSQL_ESCAPES = {
    'n': '\n', 'r': '\r', 't': '\t', '0': '\x00', 'Z': '\x1a',
    'b': '\b', '"': '"', "'": "'", '\\': '\\', '%': '\\%', '_': '\\_',
}


def decode_mysql_string(s):
    """解码 MySQL 反斜杠转义 + '' 双写转义。"""
    out = []
    i = 0
    n = len(s)
    while i < n:
        c = s[i]
        if c == '\\' and i + 1 < n:
            d = s[i + 1]
            out.append(_MYSQL_ESCAPES.get(d, d))
            i += 2
        elif c == "'" and i + 1 < n and s[i + 1] == "'":
            out.append("'")
            i += 2
        else:
            out.append(c)
            i += 1
    return ''.join(out)


# ----------------------------------------------------------------------------
# PostgreSQL COPY 文本格式编码
# ----------------------------------------------------------------------------

_COPY_SPECIAL = {
    '\\': '\\\\', '\n': '\\n', '\r': '\\r', '\t': '\\t',
    '\b': '\\b', '\f': '\\f', '\v': '\\v',
}


def copy_escape_text(val):
    """按 PostgreSQL COPY 文本格式转义（与 phpBB postgres_extractor 行为一致）。"""
    out = []
    for ch in val:
        if ch in _COPY_SPECIAL:
            out.append(_COPY_SPECIAL[ch])
        else:
            o = ord(ch)
            if o < 0x20 or o == 0x7f:       # 其余控制字符用八进制转义
                out.append('\\%03o' % o)
            else:
                out.append(ch)
    return ''.join(out)


def copy_escape_bytea(val):
    """bytea 列：恢复原始字节后逐字节八进制转义。"""
    data = val.encode('utf-8', 'surrogateescape')
    return ''.join('\\%03o' % b for b in data)


# ----------------------------------------------------------------------------
# INSERT 语句解析（phpBB/mysql 转储的多行 VALUES）
# ----------------------------------------------------------------------------

_INSERT_RE = re.compile(
    r'^\s*INSERT\s+(?:IGNORE\s+)?INTO\s+`?([\w$]+)`?\s*\(([^)]*)\)\s+VALUES\s*(.*)$',
    re.I | re.S)


def parse_values(vs):
    """解析 VALUES 后的 (...),(...) 数据，返回 [(值, ...), ...]。

    每个值为：None（NULL）或解码后的字符串（数字保持原样文本）。
    """
    rows = []
    i = 0
    n = len(vs)
    while i < n:
        # 定位行起始 '('
        while i < n and vs[i] in ' \t\r\n,':
            i += 1
        if i >= n or vs[i] != '(':
            break
        i += 1
        row = []
        while True:
            while i < n and vs[i] in ' \t\r\n':
                i += 1
            if i < n and vs[i] == ')':
                rows.append(row)
                i += 1
                break
            if i < n and vs[i] == "'":
                i += 1
                buf = []
                while i < n:
                    c = vs[i]
                    if c == '\\':
                        buf.append(c)
                        if i + 1 < n:
                            buf.append(vs[i + 1])
                            i += 1
                        i += 1
                        continue
                    if c == "'":
                        if i + 1 < n and vs[i + 1] == "'":
                            buf.append("'")
                            i += 2
                            continue
                        break
                    buf.append(c)
                    i += 1
                i += 1                     # 越过收尾引号
                row.append(decode_mysql_string(''.join(buf)))
            else:
                j = i
                while j < n and vs[j] not in ',)':
                    j += 1
                tok = vs[i:j].strip()
                if tok.upper() == 'NULL':
                    row.append(None)
                else:
                    row.append(tok)
                i = j
            while i < n and vs[i] in ' \t\r\n':
                i += 1
            if i < n and vs[i] == ',':
                i += 1
    return rows


# ----------------------------------------------------------------------------
# CREATE TABLE 解析与类型映射
# ----------------------------------------------------------------------------

_CREATE_RE = re.compile(r'\s*CREATE\s+TABLE\s+(?:IF\s+NOT\s+EXISTS\s+)?`?([^\s`(]+)`?\s*\(', re.I)


def split_create_table(stmt):
    """返回 (表名, 列/索引定义体, 尾部表选项)。"""
    m = _CREATE_RE.match(stmt)
    if not m:
        raise ValueError('无法解析 CREATE TABLE: %.80s...' % stmt.strip())
    table = m.group(1)
    start = m.end() - 1                      # 指向 '('
    depth = 0
    in_str = False
    i = start
    n = len(stmt)
    while i < n:
        c = stmt[i]
        if in_str:
            if c == '\\':
                i += 2
                continue
            if c == "'":
                in_str = False
            i += 1
            continue
        if c == "'":
            in_str = True
        elif c == '(':
            depth += 1
        elif c == ')':
            depth -= 1
            if depth == 0:
                return table, stmt[start + 1:i], stmt[i + 1:]
        i += 1
    raise ValueError('CREATE TABLE 括号不配对: %s' % table)


def split_top_level(s, sep=','):
    """按顶层分隔符切分（忽略引号与括号内部）。"""
    parts, buf = [], []
    depth = 0
    in_str = False
    in_bt = False
    i = 0
    n = len(s)
    while i < n:
        c = s[i]
        if in_str:
            buf.append(c)
            if c == '\\':
                if i + 1 < n:
                    buf.append(s[i + 1])
                    i += 1
            elif c == "'":
                in_str = False
        elif in_bt:
            buf.append(c)
            if c == '`':
                in_bt = False
        else:
            if c == "'":
                in_str = True
            elif c == '`':
                in_bt = True
            elif c == '(':
                depth += 1
            elif c == ')':
                depth -= 1
            if c == sep and depth == 0:
                parts.append(''.join(buf))
                buf = []
                i += 1
                continue
            buf.append(c)
        i += 1
    parts.append(''.join(buf))
    return parts


def map_type(base, argstr):
    """MySQL 类型 → PostgreSQL (类型, COPY 编码类别)。"""
    b = base.lower()
    if b in ('tinyint', 'smallint', 'bool', 'boolean'):
        return 'smallint', 'num'
    if b in ('mediumint', 'int', 'integer'):
        return 'integer', 'num'
    if b == 'bigint':
        return 'bigint', 'num'
    if b in ('float', 'real'):
        return 'real', 'num'
    if b in ('double', 'double precision', 'float8'):
        return 'double precision', 'num'
    if b in ('decimal', 'numeric'):
        return ('numeric(%s)' % argstr.replace(' ', ''), 'num') if argstr else ('numeric', 'num')
    if b in ('char', 'character'):
        return 'char(%s)' % (argstr or '1'), 'text'
    if b in ('varchar', 'character varying'):
        return 'varchar(%s)' % (argstr or '255'), 'text'
    if b in ('tinytext', 'mediumtext', 'longtext', 'text'):
        return 'text', 'text'
    if b in ('tinyblob', 'mediumblob', 'longblob', 'blob', 'binary', 'varbinary'):
        return 'bytea', 'bytea'
    if b == 'date':
        return 'date', 'text'
    if b in ('datetime', 'timestamp'):
        return 'timestamp', 'text'
    if b == 'time':
        return 'time', 'text'
    if b == 'year':
        return 'smallint', 'num'
    if b == 'bit':
        return 'smallint', 'num'
    if b == 'enum':
        return ('varchar(%s)' % enum_maxlen(argstr)) if argstr else 'varchar(255)', 'text'
    if b == 'set':
        return 'varchar(255)', 'text'
    return None, 'text'                     # 未知类型：按文本透传


def enum_maxlen(argstr):
    items = split_top_level(argstr)
    maxlen = 1
    for it in items:
        it = it.strip()
        if it.startswith("'") and it.endswith("'"):
            v = decode_mysql_string(it[1:-1])
            maxlen = max(maxlen, len(v))
    return maxlen


def extract_default(r):
    """从列定义剩余文本中抽出 DEFAULT 子句。返回 (原始默认值, 剩余文本)。"""
    m = re.search(r'\bDEFAULT\b', r, re.I)
    if not m:
        return None, r
    i = m.end()
    n = len(r)
    while i < n and r[i] in ' \t':
        i += 1
    j = i
    if i < n and r[i] == "'":
        j = i + 1
        while j < n:
            if r[j] == '\\':
                j += 2
                continue
            if r[j] == "'":
                if j + 1 < n and r[j + 1] == "'":
                    j += 2
                    continue
                break
            j += 1
        j += 1
        raw = r[i:j]
    elif r[i:i + 17].upper() == 'CURRENT_TIMESTAMP':
        j = i + 17
        if j < n and r[j] == '(':
            while j < n and r[j] != ')':
                j += 1
            j += 1
        raw = r[i:j]
    else:
        while j < n and (r[j].isalnum() or r[j] in '._-+'):
            j += 1
        raw = r[i:j]
    return raw, r[:m.start()] + ' ' + r[j:]


def convert_default(raw):
    """把 MySQL 默认值转换为 PostgreSQL 字面量；零日期返回 ('ZERO_DATE', 原值)。"""
    if raw is None:
        return None
    raw = raw.strip()
    if not raw or raw.upper() == 'NULL':
        return None
    if raw.upper().startswith('CURRENT_TIMESTAMP'):
        return 'CURRENT_TIMESTAMP'
    if raw.upper().startswith("B'"):
        return raw[2:-1] or '0'
    if raw.startswith("'"):
        inner = raw[1:-1] if raw.endswith("'") else raw[1:]
        val = decode_mysql_string(inner)
        if re.match(r'^0{4}-0{2}-0{2}', val or ''):
            return ('ZERO_DATE', val)
        return "'%s'" % val.replace("'", "''")
    return raw                              # 数字等


_COL_RE = re.compile(r'^`?([\w$]+)`?\s+(.+)$', re.S)


def parse_column(item, table, warn):
    """解析单个列定义 → (列名, PG行, 类别, 是否自增, 附加CHECK或None)。"""
    m = _COL_RE.match(item.strip())
    if not m:
        warn('无法解析列定义: %.60s' % item.strip())
        return None
    name, rest = m.group(1), m.group(2).strip()
    tm = re.match(r'^([A-Za-z][A-Za-z0-9_]*)\s*(?:\(\s*(.*?)\s*\))?', rest, re.S)
    base = tm.group(1)
    argstr = tm.group(2)
    r = rest[tm.end():].strip()

    pgtype, cat = map_type(base, argstr)
    if pgtype is None:
        warn('%s.%s: 未知类型 %s，按 text 处理' % (table, name, base))
        pgtype = 'text'

    extra_check = None
    if base and base.lower() == 'enum' and argstr:
        items = [it.strip() for it in split_top_level(argstr)]
        extra_check = '  CHECK (%s IN (%s))' % (name, ', '.join(items))
    if base and base.lower() == 'set':
        warn('%s.%s: SET 类型转为 varchar(255)，多值语义需人工核对' % (table, name))

    # 修饰符清洗
    auto = bool(re.search(r'\bAUTO_INCREMENT\b', r, re.I))
    r = re.sub(r'\bUNSIGNED\b|\bZEROFILL\b|\bAUTO_INCREMENT\b', ' ', r, flags=re.I)
    r = re.sub(r'\bCHARACTER\s+SET\s+\w+|\bCHARSET\s+\w+', ' ', r, flags=re.I)
    r = re.sub(r'\bCOLLATE\s+\w+', ' ', r, flags=re.I)
    r = re.sub(r'\bON\s+UPDATE\s+CURRENT_TIMESTAMP(?:\s*\(\d*\))?', ' ', r, flags=re.I)
    r = re.sub(r"\bCOMMENT\s+'(?:[^'\\]|\\.)*'", ' ', r, flags=re.I)

    default_raw, r = extract_default(r)
    default = convert_default(default_raw)
    zero_date = isinstance(default, tuple) and default[0] == 'ZERO_DATE'
    if zero_date:
        default = None
        warn('%s.%s: 零日期默认值 %s 已改为 NULL（PostgreSQL 不接受零日期）'
             % (table, name, default_raw))

    notnull = bool(re.search(r'\bNOT\s+NULL\b', r, re.I))
    r = re.sub(r'\bNOT\s+NULL\b|\bNULL\b', ' ', r, flags=re.I)
    leftover = re.sub(r'[\s,]+', '', r)
    if leftover:
        warn('%s.%s: 未识别的列修饰符 "%s" 已忽略' % (table, name, leftover))

    line = '  %s %s' % (name, pgtype)
    if auto:
        line += " DEFAULT nextval('%s_seq')" % table
    elif default is not None:
        line += ' DEFAULT %s' % default
    if notnull and not zero_date:
        line += ' NOT NULL'
    return name, line, cat, auto, extra_check


_IDX_COL_RE = re.compile(r'^`?([\w$]+)`?(?:\((\d+)\))?(?:\s+(ASC|DESC))?\s*$', re.I)


def parse_index_cols(paren_body):
    """解析索引列清单 → [(列名, 排序或None)]，长度前缀被剥离。"""
    cols = []
    for c in split_top_level(paren_body):
        c = c.strip()
        m = _IDX_COL_RE.match(c)
        if m:
            cols.append((m.group(1), m.group(3).upper() if m.group(3) else None))
        else:
            cols.append((c, None))          # 表达式索引等，交由上层告警
    return cols


# ----------------------------------------------------------------------------
# 主转换器
# ----------------------------------------------------------------------------

class Mysql2Pg(object):

    def __init__(self, output_path, quiet=False):
        self.path = output_path
        self.quiet = quiet
        self.f = open(output_path, 'w', encoding='utf-8',
                      errors='surrogateescape', newline='')
        self.warnings = []
        self.tables = {}                    # 表名 → {'cats': {列:类别}, 'auto': 列名}
        self.table_rows = {}
        self.used_index_names = set()
        self._prefix = 'phpbb_'

    def warn(self, msg):
        self.warnings.append(msg)

    def w(self, s):
        self.f.write(s)

    # ---------------- 主流程 ----------------

    def convert(self, text):
        m = re.search(r'^#\s*Dump of tables for\s*(\S*)\s*$', text, re.M)
        if m:
            self._prefix = m.group(1) or 'phpbb_'
        if not re.search(r'^#\s*phpBB Backup Script', text, re.M):
            self.warn('输入文件头部未找到 "phpBB Backup Script" 标识，请确认这是 phpBB ACP 导出的备份')

        self.w('--\n')
        self.w('-- phpBB Backup Script\n')
        self.w('-- Dump of tables for %s\n' % self._prefix)
        self.w('-- DATE : %s GMT\n' % time.strftime('%d-%m-%Y %H:%M:%S', time.gmtime()))
        self.w('--\n')
        self.w('BEGIN TRANSACTION;\n')

        stmts = split_statements(text)
        n_stmt = 0
        for stmt in stmts:
            _comments, code = split_comments(stmt)
            if not code:
                continue
            head = code.split(None, 1)[0].upper()
            if head in ('SET', 'LOCK', 'UNLOCK', 'USE', 'FLUSH', 'DELIMITER') \
                    or code.startswith('/*'):
                continue                    # mysqldump 环境语句，丢弃
            if re.match(r'^DROP\s+TABLE', code, re.I):
                continue                    # DROP 由转换器自行生成
            if re.match(r'^CREATE\s+TABLE', code, re.I):
                self.handle_create_table(code)
                n_stmt += 1
            elif re.match(r'^INSERT\s', code, re.I):
                self.handle_insert(code)
                n_stmt += 1
            else:
                self.warn('跳过不支持的语句: %.60s' % code.replace('\n', ' '))

        self.w('COMMIT;\n')
        self.f.close()
        return n_stmt

    # ---------------- DDL ----------------

    def handle_create_table(self, code):
        table, body, _tail = split_create_table(code)
        items = split_top_level(body)
        lines = []
        pk_cols = None
        index_stmts = []
        cats = {}
        auto_col = None

        for it in items:
            s = it.strip()
            if not s:
                continue
            u = s.upper()
            if u.startswith('PRIMARY KEY'):
                pk_cols = parse_index_cols(s[s.index('(') + 1:s.rindex(')')])
            elif re.match(r'^(FULLTEXT|SPATIAL)\b', u):
                self.warn('%s: 跳过 %s 索引（PostgreSQL 语法不同）；'
                          '全文检索请在恢复后于 ACP 重建搜索索引'
                          % (table, u.split()[0]))
            elif re.match(r'^UNIQUE\s+(KEY|INDEX)\b', u) or \
                    (u.startswith('CONSTRAINT') and 'UNIQUE' in u):
                self.handle_index(s, table, index_stmts, unique=True)
            elif re.match(r'^(KEY|INDEX)\b', u):
                self.handle_index(s, table, index_stmts, unique=False)
            else:
                parsed = parse_column(s, table, self.warn)
                if not parsed:
                    continue
                name, line, cat, auto, extra_check = parsed
                lines.append(line)
                if extra_check:
                    lines.append(extra_check)
                cats[name] = cat
                if auto:
                    auto_col = name

        if pk_cols:
            lines.append('  PRIMARY KEY (%s)' %
                         ', '.join(c for c, _o in pk_cols))

        self.w('-- Table: %s\n' % table)
        self.w('DROP TABLE IF EXISTS %s;\n' % table)
        if auto_col:
            self.w('DROP SEQUENCE IF EXISTS %s_seq;\n' % table)
            self.w('CREATE SEQUENCE %s_seq;\n' % table)
        self.w('CREATE TABLE %s(\n%s\n);\n' % (table, ', \n'.join(lines)))
        for s in index_stmts:
            self.w(s + ';\n')
        self.w('\n')

        self.tables[table] = {'cats': cats, 'auto': auto_col}

    def handle_index(self, s, table, index_stmts, unique):
        s = s.strip()
        # 形如: [UNIQUE] KEY `name` (`col`(10), `col2` DESC) 或 CONSTRAINT `name` UNIQUE (...)
        m = re.match(r'^(?:UNIQUE\s+)?(?:KEY|INDEX)\s+(?:`?([\w$]+)`?\s+)?(\(.*\))\s*$',
                     s, re.S | re.I) or \
            re.match(r'^CONSTRAINT\s+`?([\w$]+)`?\s+UNIQUE\s+(\(.*\))\s*$', s, re.S | re.I)
        if not m:
            self.warn('%s: 跳过无法解析的索引 %.60s' % (table, s))
            return
        keyname, paren = m.group(1), m.group(2)
        cols = parse_index_cols(paren[1:-1])
        if not cols or any(not re.match(r'^[\w$]+$', c) for c, _o in cols):
            self.warn('%s: 跳过无法解析的索引 %.60s' % (table, s))
            return
        if keyname is None:
            keyname = 'k%d' % len(index_stmts)
        idx_name = '%s_%s' % (table, keyname)
        while idx_name in self.used_index_names:
            idx_name += '_2'
        self.used_index_names.add(idx_name)
        col_txt = ', '.join(c + (' ' + o if o else '') for c, o in cols)
        index_stmts.append('CREATE %sINDEX %s ON %s (%s)' %
                           ('UNIQUE ' if unique else '', idx_name, table, col_txt))

    # ---------------- 数据 ----------------

    def col_category(self, table, col):
        info = self.tables.get(table)
        if not info:
            return 'text'
        return info['cats'].get(col, 'text')

    def handle_insert(self, code):
        m = _INSERT_RE.match(code)
        if not m:
            self.warn('跳过无法解析的 INSERT: %.60s' % code.replace('\n', ' '))
            return
        table = m.group(1)
        cols = [c.strip().strip('`') for c in m.group(2).split(',')]
        rows = parse_values(m.group(3))
        if not rows:
            return
        cats = [self.col_category(table, c) for c in cols]

        self.w('COPY %s (%s) FROM stdin;\n' % (table, ', '.join(cols)))
        for row in rows:
            fields = []
            for val, cat in zip(row, cats):
                if val is None:
                    fields.append('\\N')
                elif cat == 'bytea':
                    fields.append(copy_escape_bytea(val))
                elif cat == 'text':
                    fields.append(copy_escape_text(val))
                else:                       # 数值等
                    sv = str(val).strip()
                    if sv == '':
                        fields.append('\\N')
                    else:
                        fields.append(sv)
            self.w('\t'.join(fields) + '\n')
        self.w('\\.\n')

        info = self.tables.get(table)
        if info and info['auto']:
            col = info['auto']
            self.w("SELECT SETVAL('%s_seq',(select case when max(%s)>0 "
                   "then max(%s)+1 else 1 end FROM %s));\n" % (table, col, col, table))
        self.table_rows[table] = self.table_rows.get(table, 0) + len(rows)


# ----------------------------------------------------------------------------
# CLI
# ----------------------------------------------------------------------------

def main(argv=None):
    ap = argparse.ArgumentParser(
        description='phpBB 3.3 ACP 备份转换器：MySQL/MariaDB → PostgreSQL（ACP 可直接恢复）')
    ap.add_argument('input', help='phpBB ACP 导出的备份文件（.sql / .sql.gz / .sql.bz2）')
    ap.add_argument('-o', '--output', help='输出文件名（默认按 ACP 命名规则自动生成）')
    ap.add_argument('--quiet', action='store_true', help='不打印统计信息')
    args = ap.parse_args(argv)

    if not os.path.isfile(args.input):
        sys.stderr.write('错误: 找不到输入文件 %s\n' % args.input)
        return 1

    output = args.output or 'backup_%d_%s.sql' % (int(time.time()), secrets.token_hex(8))
    if not ACP_FILENAME_RE.match(os.path.basename(output)):
        sys.stderr.write('警告: 输出文件名不符合 phpBB ACP 恢复页校验规则\n'
                         '（backup_<10位时间戳>_<16位hex>.sql），恢复时可能无法识别\n')

    text = read_input(args.input)
    conv = Mysql2Pg(output, quiet=args.quiet)
    n_stmt = conv.convert(text)

    if not args.quiet:
        sys.stdout.write('输入: %s\n输出: %s\n' % (args.input, output))
        sys.stdout.write('已处理语句: %d；表: %d；数据行: %d\n'
                          % (n_stmt, len(conv.tables), sum(conv.table_rows.values())))
        if conv.warnings:
            sys.stdout.write('警告 (%d 条):\n' % len(conv.warnings))
            for wmsg in conv.warnings:
                sys.stdout.write('  - %s\n' % wmsg)
        else:
            sys.stdout.write('警告: 无\n')
    return 0


if __name__ == '__main__':
    sys.exit(main())
