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)})