package aceboot::orm.postgres
import aceboot::orm.common.*

import std.collection.*
import std.convert.*

// ===== libpq FFI(可执行/测试模块需 link-option "-lpq")=====
foreign func PQconnectdb(conninfo: CString): CPointer<Unit>

foreign func PQstatus(conn: CPointer<Unit>): Int32

foreign func PQerrorMessage(conn: CPointer<Unit>): CString

foreign func PQexecParams(conn: CPointer<Unit>, command: CString, nParams: Int32, paramTypes: CPointer<UInt32>,
    paramValues: CPointer<CString>, paramLengths: CPointer<Int32>, paramFormats: CPointer<Int32>,
    resultFormat: Int32): CPointer<Unit>

foreign func PQresultStatus(res: CPointer<Unit>): Int32

foreign func PQresultErrorMessage(res: CPointer<Unit>): CString

foreign func PQntuples(res: CPointer<Unit>): Int32

// 最近一条命令影响的行数(UPDATE/DELETE/INSERT),以字符串返回;非 DML 返回空串。
foreign func PQcmdTuples(res: CPointer<Unit>): CString

foreign func PQnfields(res: CPointer<Unit>): Int32

foreign func PQfname(res: CPointer<Unit>, col: Int32): CString

// 列类型 OID(pg_type.oid),用于区分数组列与 JSONB/TEXT 等标量列。
foreign func PQftype(res: CPointer<Unit>, col: Int32): UInt32

foreign func PQgetvalue(res: CPointer<Unit>, row: Int32, col: Int32): CString

foreign func PQgetisnull(res: CPointer<Unit>, row: Int32, col: Int32): Int32

foreign func PQclear(res: CPointer<Unit>): Unit

foreign func PQfinish(conn: CPointer<Unit>): Unit

foreign func PQresultErrorField(res: CPointer<Unit>, fieldcode: Int32): CString

const PG_CONNECTION_OK: Int32 = 0
let PG_DIAG_SQLSTATE: Int32 = 67  // ASCII 'C'
const PGRES_COMMAND_OK: Int32 = 1
const PGRES_TUPLES_OK: Int32 = 2

/**
 * PostgreSQL 驱动(FFI libpq)。conninfo 形如
 * "host=127.0.0.1 port=5432 user=postgres password=xxx dbname=app",可附 `pool.max`/`pool.min`。
 * 参数化走 PQexecParams(text 协议,$n 占位),杜绝注入;自增 id 由 Repository 经 RETURNING 取。
 *
 * 连接池化:不再持有单连接 + Mutex 串行化,而是每次操作从 `ConnectionPool` 借一条连接、
 * 用完归还,故并发请求不再相互串行(libpq 单连接非并发安全,借出的连接同一时刻只属一个线程)。
 * 事务期间(begin..commit/rollback)用 ThreadLocal 把同一条连接钉在当前线程,保证
 * BEGIN/中间语句/COMMIT 走同一连接、互不干扰——天然支持多线程各自独立事务。
 */
public class PostgresDriver <: Driver {
    let pool: ConnectionPool
    // 当前线程绑定的事务连接(无事务时为 None);保证事务内语句走同一连接。
    let txn = ThreadLocal<CPointer<Unit>>()
    // 最近一条 exec 影响的行数(乐观锁判定用);-1 表示未知/非 DML。
    var lastChanges: Int64 = -1

    public init(conninfo: String) {
        let (clean, cfg) = extractPoolConfig(conninfo)
        this.pool = ConnectionPool(cfg, {=> PostgresDriver.openConn(clean)}, {c => unsafe { PQfinish(c) }})
        // 立刻借一条验证连通性(连不上则抛,沿用「无 PG 环境则跳过」的测试约定)。
        let c = pool.borrow()
        pool.giveBack(c)
    }

    // 打开并校验一条 libpq 连接;失败抛 OrmException(并销毁半成品连接)。
    static func openConn(conninfo: String): CPointer<Unit> {
        let ci = unsafe { LibC.mallocCString(conninfo) }
        let c = unsafe { PQconnectdb(ci) }
        unsafe { LibC.free(ci) }
        if (unsafe { PQstatus(c) } != PG_CONNECTION_OK) {
            let msg = unsafe { PQerrorMessage(c) }.toString()
            unsafe { PQfinish(c) }
            throw OrmException("postgres connect failed: ${msg}")
        }
        return c
    }

    public func dialect(): Dialect {
        PostgresDialect()
    }

    // DbValue -> 文本(text 协议绑定)。
    static func toText(v: DbValue): String {
        match (v) {
            case DbInt(x) => x.toString()
            case DbText(s) => s
            case DbReal(x) => x.toString()
            case DbBool(b) => if (b) {
                "true"
            } else {
                "false"
            }
            case DbBytes(b) => "\\x" + bytesToHex(b) // BYTEA 文本输入格式 \xHEX
            case DbNull => "" // 防御兜底:绑定路径不经此(run() 对 DbNull 直接绑 null 指针即 SQL NULL)
            case DbArray(a) => pgArrayLiteral(a) // PG 数组字面量文本(如 {1,2,NULL}),可绑给 $n::type[]
        }
    }

    // DbValue 列表 -> PG 数组字面量(text 协议输入格式,如 {1,2,NULL} / {"a","b,c"}),嵌套数组递归。
    static func pgArrayLiteral(items: Array<DbValue>): String {
        var s = "{"
        var first = true
        for (it in items) {
            if (!first) {
                s += ","
            }
            first = false
            s += pgArrayElem(it)
        }
        return s + "}"
    }

    // 单个数组元素:数值/布尔裸写,NULL 关键字,文本/字节加引号转义(防逗号/引号/花括号破坏结构)。
    static func pgArrayElem(v: DbValue): String {
        match (v) {
            case DbInt(x) => x.toString()
            case DbReal(x) => x.toString()
            case DbBool(b) => if (b) {
                "t"
            } else {
                "f"
            }
            case DbNull => "NULL"
            case DbText(s) => quotePgArrayElem(s)
            case DbBytes(b) => quotePgArrayElem("\\x" + bytesToHex(b))
            case DbArray(a) => pgArrayLiteral(a)
        }
    }

    // 元素加双引号,内部 \ 与 " 前置反斜杠(PG 数组字面量元素转义规则)。
    static func quotePgArrayElem(s: String): String {
        let out = ArrayList<UInt8>()
        out.add(34u8) // "
        for (b in s.toArray()) {
            if (b == 34u8 || b == 92u8) {
                out.add(92u8)
            }
            out.add(b)
        }
        out.add(34u8)
        return String.fromUtf8(out.toArray())
    }

    // 字节数组 -> 小写十六进制串(无分隔)。
    static func bytesToHex(b: Array<UInt8>): String {
        let digits = "0123456789abcdef".toArray()
        let out = ArrayList<UInt8>()
        for (x in b) {
            out.add(digits[Int64(x >> 4)])
            out.add(digits[Int64(x & 15u8)])
        }
        return String.fromUtf8(out.toArray())
    }

    // 常见一维数组类型 OID(pg_type):_json/_bool/_bytea/_int2/_int4/_text/_bpchar/_varchar/_int8/
    // _float4/_float8/_timestamp/_date/_timestamptz/_numeric/_uuid/_jsonb。JSONB(3802)/JSON(114) 不在其列。
    static func isPgArrayOid(oid: UInt32): Bool {
        match (oid) {
            case 199 | 1000 | 1001 | 1005 | 1007 | 1009 | 1014 | 1015 | 1016 | 1021 | 1022 |
                1115 | 1182 | 1185 | 1231 | 2951 | 3807 => true
            case _ => false
        }
    }

    // 解析 PG 文本协议一维数组(如 {1,2,NULL})为 DbValue 列表。
    static func parsePgArray(s: String): Array<DbValue> {
        let bytes = s.toArray()
        let n = bytes.size
        if (n < 2) {
            return Array<DbValue>()
        }
        // 去掉首尾 { }
        let inner = bytes[1..n - 1]
        let tokens = ArrayList<String>()
        var cur = ArrayList<UInt8>()
        var inQ = false
        var i = 0
        let m = inner.size
        while (i < m) {
            let b = inner[i]
            if (inQ && b == 92u8 && i + 1 < m) { // 引号内反斜杠转义:下一字节取字面值
                cur.add(inner[i + 1])
                i += 2
                continue
            }
            if (b == 34u8) {
                inQ = !inQ
                i += 1
                continue
            }
            if (b == 44u8 && !inQ) {
                tokens.add(String.fromUtf8(cur.toArray()))
                cur = ArrayList<UInt8>()
                i += 1
                continue
            }
            cur.add(b)
            i += 1
        }
        tokens.add(String.fromUtf8(cur.toArray()))
        let result = ArrayList<DbValue>()
        for (tok in tokens) {
            let t = tok.trimAscii()
            if (t == "NULL" || t.isEmpty()) {
                result.add(DbNull)
                continue
            }
            let vi: ?Int64 = try { Some(Int64.parse(t)) } catch (_: Exception) { None }
            match (vi) {
                case Some(i) => result.add(DbInt(i)); continue
                case None => ()
            }
            let vr: ?Float64 = try { Some(Float64.parse(t)) } catch (_: Exception) { None }
            match (vr) {
                case Some(f) => result.add(DbReal(f)); continue
                case None => ()
            }
            result.add(DbText(t))
        }
        return result.toArray()
    }

    // 在指定连接上执行 PQexecParams 并校验状态;返回 result 指针(调用方负责 PQclear)。
    static func run(conn: CPointer<Unit>, sql: String, params: Array<DbValue>): CPointer<Unit> {
        let n = params.size
        let csql = unsafe { LibC.mallocCString(sql) }
        // DbNull 须绑成 null 指针(libpq 语义:值指针为 NULL 即 SQL NULL);
        // 绑空串会对 timestamp/integer 等类型列报 invalid input syntax。
        let cstrs = Array<CString>(n, {i => match (params[i]) {
            case DbNull => unsafe { CString(CPointer<UInt8>()) }
            case _ => unsafe { LibC.mallocCString(toText(params[i])) }
        }})
        let res = if (n == 0) {
            unsafe {
                PQexecParams(conn, csql, 0, CPointer<UInt32>(), CPointer<CString>(), CPointer<Int32>(),
                    CPointer<Int32>(), 0)
            }
        } else {
            let handle = unsafe { acquireArrayRawData(cstrs) }
            let r = unsafe {
                PQexecParams(conn, csql, Int32(n), CPointer<UInt32>(), handle.pointer, CPointer<Int32>(),
                    CPointer<Int32>(), 0)
            }
            unsafe { releaseArrayRawData(handle) }
            r
        }
        unsafe { LibC.free(csql) }
        for (cs in cstrs) {
            if (!cs.isNull()) { // DbNull 占位是 null 指针,无需释放
                unsafe { LibC.free(cs) }
            }
        }
        let status = unsafe { PQresultStatus(res) }
        if (status != PGRES_COMMAND_OK && status != PGRES_TUPLES_OK) {
            let msg = unsafe { PQresultErrorMessage(res) }.toString()
            let statePtr = unsafe { PQresultErrorField(res, PG_DIAG_SQLSTATE) }
            let sqlstate = if (statePtr.isNull()) { "" } else { statePtr.toString() }
            unsafe { PQclear(res) }
            throw OrmException("postgres exec failed: ${msg} | sql=${sql}", code: sqlstate)
        }
        return res
    }

    // 在指定连接上执行写语句,返回受影响行数(自增 id 经 RETURNING(query)取,故此处不返 id)。
    static func execOn(conn: CPointer<Unit>, sql: String, params: Array<DbValue>): Int64 {
        let res = run(conn, sql, params)
        let aff = parseAffected(res)
        unsafe { PQclear(res) }
        return aff
    }

    // 解析 PQcmdTuples(受影响行数文本);空串(非 DML)或无法解析时返回 -1。
    static func parseAffected(res: CPointer<Unit>): Int64 {
        let s = unsafe { PQcmdTuples(res) }.toString()
        if (s.isEmpty()) {
            return -1
        }
        return try {
            Int64.parse(s)
        } catch (_: Exception) {
            -1
        }
    }

    // 在指定连接上执行查询,物化为 Row 集。
    static func queryOn(conn: CPointer<Unit>, sql: String, params: Array<DbValue>): Array<Row> {
        let res = run(conn, sql, params)
        let nrows = unsafe { PQntuples(res) }
        let ncols = unsafe { PQnfields(res) }
        let rows = ArrayList<Row>()
        for (i in 0..nrows) {
            let cells = HashMap<String, DbValue>()
            for (j in 0..ncols) {
                let name = unsafe { PQfname(res, j) }.toString()
                let v = if (unsafe { PQgetisnull(res, i, j) } == 1) {
                    DbNull
                } else {
                    let textVal = unsafe { PQgetvalue(res, i, j) }.toString()
                    // 按列类型 OID 判定数组,不嗅探内容——JSONB `{"k":"v"}`/存 JSON 的 TEXT 列也以 { 开头
                    if (isPgArrayOid(unsafe { PQftype(res, j) })) {
                        DbArray(parsePgArray(textVal))
                    } else {
                        DbText(textVal) // text 协议,Row 取值器按需解析
                    }
                }
                cells.add(name, v)
            }
            rows.add(Row(cells))
        }
        unsafe { PQclear(res) }
        return rows.toArray()
    }

    public func exec(sql: String, params: Array<DbValue>): Int64 {
        let aff = match (txn.get()) {
            case Some(c) => PostgresDriver.execOn(c, sql, params) // 事务内:走钉住的连接
            case None =>
                let c = pool.borrow()
                try {
                    PostgresDriver.execOn(c, sql, params)
                } finally {
                    pool.giveBack(c)
                }
        }
        lastChanges = aff
        return aff // 受影响行数(对标 JDBC executeUpdate;非 DML 返回 -1)
    }

    /** PG 无 last-insert-id 概念,恒返 0;自增 id 经 INSERT ... RETURNING(query)获取。 */
    public func lastInsertId(): Int64 {
        0
    }

    /** 上一条 exec 影响的行数(供 @VersionColumn 乐观锁判定;非 DML 返回 -1)。 */
    public func affectedRows(): Int64 {
        lastChanges
    }

    public func query(sql: String, params: Array<DbValue>): Array<Row> {
        match (txn.get()) {
            case Some(c) => PostgresDriver.queryOn(c, sql, params)
            case None =>
                let c = pool.borrow()
                try {
                    PostgresDriver.queryOn(c, sql, params)
                } finally {
                    pool.giveBack(c)
                }
        }
    }

    public func begin(): Unit {
        // 借一条连接钉到当前线程,直到 commit/rollback 释放。BEGIN 失败即归还。
        let c = pool.borrow()
        try {
            let res = PostgresDriver.run(c, "BEGIN", [])
            unsafe { PQclear(res) }
            txn.set(Some(c))
        } catch (e: Exception) {
            pool.giveBack(c)
            throw e
        }
    }

    public func commit(): Unit {
        endTxn("COMMIT")
    }

    public func rollback(): Unit {
        endTxn("ROLLBACK")
    }

    // 收尾事务:在钉住的连接上执行 COMMIT/ROLLBACK,无论成败都解钉并归还连接。
    func endTxn(verb: String): Unit {
        match (txn.get()) {
            case Some(c) =>
                try {
                    let res = PostgresDriver.run(c, verb, [])
                    unsafe { PQclear(res) }
                } finally {
                    txn.set(None)
                    pool.giveBack(c)
                }
            case None => () // 无活动事务:静默忽略(对齐旧版「裸 COMMIT 无害」语义)
        }
    }

    public func close(): Unit {
        pool.closeAll()
    }

    /**
     * 覆盖 Driver 默认实现:使用 PG 原生游标(DECLARE/FETCH)实现流式读取。
     * - 若当前协程已在 ds.transaction{} 内(txn 已钉连接),直接用该连接声明游标,跳过 BEGIN/COMMIT。
     * - 否则从连接池借一条新连接并自建事务包裹游标生命周期,异常时 ROLLBACK。
     * 在外层事务内调用时,游标名加 nowMillis() 后缀以避免同连接上多游标同名冲突。
     */
    public func streamQuery(sql: String, params: Array<DbValue>, chunkSize: Int64,
                            cb: (Array<Row>) -> Unit): Unit {
        match (txn.get()) {
            case Some(c) =>
                // 外层事务:用钉住的连接,游标名带时间戳后缀防同连接上的嵌套冲突
                let cursorName = "ace_stream_" + nowMillis().toString()
                let declRes = PostgresDriver.run(c, "DECLARE " + cursorName + " CURSOR FOR " + sql, params)
                unsafe { PQclear(declRes) }
                let fetchSql = "FETCH " + chunkSize.toString() + " FROM " + cursorName
                while (true) {
                    let rows = PostgresDriver.queryOn(c, fetchSql, [])
                    if (rows.size == 0) {
                        break
                    }
                    cb(rows)
                    if (rows.size < chunkSize) {
                        break
                    }
                }
                PostgresDriver.execOn(c, "CLOSE " + cursorName, [])
            case None =>
                // 无活动事务:借新连接,自建事务包裹游标生命周期
                let conn = pool.borrow()
                try {
                    PostgresDriver.execOn(conn, "BEGIN", [])
                    let declRes = PostgresDriver.run(conn, "DECLARE ace_stream_cur CURSOR FOR " + sql, params)
                    unsafe { PQclear(declRes) }
                    let fetchSql = "FETCH " + chunkSize.toString() + " FROM ace_stream_cur"
                    while (true) {
                        let rows = PostgresDriver.queryOn(conn, fetchSql, [])
                        if (rows.size == 0) {
                            break
                        }
                        cb(rows)
                        if (rows.size < chunkSize) {
                            break
                        }
                    }
                    PostgresDriver.execOn(conn, "CLOSE ace_stream_cur", [])
                    PostgresDriver.execOn(conn, "COMMIT", [])
                    pool.giveBack(conn)
                } catch (e: Exception) {
                    try {
                        PostgresDriver.execOn(conn, "ROLLBACK", [])
                    } catch (_: Exception) {}
                    pool.giveBack(conn)
                    throw e
                }
        }
    }
}

// import 本模块即登记 "postgres" 驱动工厂(供 OrmComponent 据 [datasource].driver 选用)。
let _ = registerDriverFactory("postgres", {url => PostgresDriver(url)})