diff --git a/.cursor/rules/branch/333-prepared-statement.mdc b/.cursor/rules/branch/333-prepared-statement.mdc new file mode 100644 index 00000000..c84330b2 --- /dev/null +++ b/.cursor/rules/branch/333-prepared-statement.mdc @@ -0,0 +1,167 @@ +--- +description: 333-prepared-statementブランチでの開発時に読み込む +alwaysApply: false +--- +333-prepared-statement ブランチでの高速化内容 +=== + +このブランチで実装することは以下の通りです。 +- prepared statement を、高頻度な同一 SQL 実行で通常経路より速くなるように改善する +- とくに PostgreSQL で、prepare/deallocate の無駄と接続拡散を減らす +- benchmark が prepared statement の warm 状態を正しく測れるようにする + +## 進捗 + +- [x] `example/benchmark.nim` の PostgreSQL 計測で prepared update が遅く見える理由を調査した +- [x] prepared statement 高速化の設計方針を整理した +- [ ] PostgreSQL にプール単位の prepared statement cache を追加する +- [ ] `close()` を論理 close 化し、物理破棄は別 API に分離する +- [ ] 接続固定で prepared statement を連続実行できる API を追加する +- [ ] `example/benchmark.nim` を cold/warm 計測に分離し、warm 経路で同一 statement を再利用する +- [ ] prepared cache と接続固定 API のテストを追加する + +## 参考資料 + +- PostgreSQL libpq prepared statements: https://www.postgresql.jp/document/12/html/libpq-exec.html + - `PQsendPrepare` + - `PQsendQueryPrepared` + - `DEALLOCATE` +- PostgreSQL 非同期 libpq: https://www.postgresql.jp/document/12/html/libpq-async.html +- 対象実装 + - `src/allographer/query_builder/models/postgres/postgres_exec.nim` + - `src/allographer/query_builder/models/postgres/postgres_types.nim` + - `src/allographer/query_builder/libs/postgres/postgres_impl.nim` + - `example/benchmark.nim` + +## 調査結果・設計まとめ + +### 現状の問題 + +PostgreSQL で prepared statement が遅く見える主因は、prepared statement 自体の性質ではなく、現在の実装と benchmark の計測条件にある。 + +- `benchmark.nim` は計測のたびに `prepare` と `close` を行っており、steady-state の再利用性能ではなくライフサイクル全体を測っている +- PostgreSQL 実装は prepared statement をコネクションごとに lazy prepare するため、プールが大きいと `PREPARE` が広く拡散する +- `first()` と `exec()` が毎回 `getFreeConn()` するため、1 つの業務操作でも接続が変わり、prepared state の局所性が弱い +- 通常経路も `PQsendQueryParams` による bind 実行なので、prepared の利得は parse/plan 再利用ぶんに限られる + +### 高速化の基本方針 + +prepared statement を速くするには、単発 API を追加するだけでは足りない。以下をまとめて入れる必要がある。 + +- prepare/deallocate の回数を減らす +- 同じ logical operation を同じ接続に寄せる +- warm 済み prepared state をプール内で継続利用する + +### 設計方針 1: プール単位の prepared cache + +`PreparedStatement` ハンドルごとに物理 state を持つのではなく、`Connections` が SQL ごとの prepared state を保持する。 + +概念設計: + +```nim +type PostgresPreparedEntry = ref object + sql*: string + nArgs*: int + stmtBaseName*: string + stmtNames*: seq[string] ## connI ごとの物理 prepared 名 + refCount*: int + lastUsedAt*: int64 + +type Connections* = ref object + conns*: seq[Connection] + timeout*: int + waiters*: Deque[Future[void]] + columnTypeCache*: Table[string, seq[Row]] + preparedCache*: Table[string, PostgresPreparedEntry] +``` + +期待効果: + +- 同じ SQL に対する `prepare()` の重複を避けられる +- 同一プール内では physical prepare を接続ごとに 1 回へ寄せられる +- 短命な prepared handle を繰り返し作っても warm state を失いにくくなる + +### 設計方針 2: `close()` の論理 close 化 + +現在の `close()` は、その handle が使った接続すべてに `DEALLOCATE` を送る。高速化の観点ではこれは不利なので、意味論を分ける。 + +- `close()`: handle の参照を閉じるだけ +- `clearPreparedCache()` または `flushPrepared(sql)`: 物理 `DEALLOCATE` +- 接続 close 時: その接続に紐づく prepared state を一括破棄 + +期待効果: + +- 利用者が毎回 `prepare` / `close` しても、プール側は hot SQL を保持できる +- benchmark の warm 計測でも、毎回の物理破棄コストを避けられる + +### 設計方針 3: 接続固定 API の追加 + +prepared statement の locality を高めるため、同じ接続上で複数回の prepared 実行を束ねる API を追加する。 + +候補 API: + +```nim +proc withPreparedConnection*( + self: PostgresConnections, + body: proc (ctx: PostgresPreparedContext): Future[void] +): Future[void] +``` + +```nim +type PostgresPreparedContext* = ref object + owner*: PostgresConnections + connI*: int +``` + +prepared statement 側には `ctx` 付き overload を足す。 + +```nim +proc first*(self: PostgresPreparedStatement, ctx: PostgresPreparedContext, args: seq[string]): Future[Option[JsonNode]] +proc exec*(self: PostgresPreparedStatement, ctx: PostgresPreparedContext, args: seq[string]): Future[void] +``` + +期待効果: + +- `SELECT` と `UPDATE` が別接続に飛ぶのを防げる +- lazy prepare の発生先が必要な接続に限定される +- transaction 中の `transactionConn` と整合しやすい + +### benchmark の変更方針 + +`example/benchmark.nim` は prepared statement の高速化効果を正しく測れるように変更する。 + +- `selectStmt` / `updateStmt` を `timeProcess` の外で 1 回だけ作る +- 繰り返し計測では同じ prepared statement を再利用する +- `close()` は全計測の最後に 1 回だけ呼ぶ +- benchmark 名を `cold` と `warm` に分ける +- warm 計測では `withPreparedConnection` を使って接続固定経路を測る + +### 実装順序 + +#### Phase 1 + +- `Connections.preparedCache` を追加する +- `prepare()` を cache lookup ベースに変更する +- `close()` を論理 close 化する +- `clearPreparedCache()` を追加する + +#### Phase 2 + +- `PostgresPreparedContext` を追加する +- `withPreparedConnection` を追加する +- prepared `get/first/exec` に `ctx` overload を追加する +- `benchmark.nim` を warm/cold 計測へ更新する + +#### Phase 3 + +- `preparedCacheMaxEntries` を追加する +- `refCount == 0` の entry を対象に LRU eviction を入れる +- eviction 時に `DEALLOCATE` を発行する + +### テスト方針 + +- 同じ SQL を複数回 `prepare()` しても、同一接続で physical prepare が 1 回に収まること +- `close()` 後に再 `prepare()` しても、cache が残る設計なら physical prepare が増えないこと +- `withPreparedConnection` の block 内で同じ `connI` が使われること +- transaction 中は `transactionConn` が優先されること +- cache eviction 時に `DEALLOCATE` が正しく走ること diff --git a/example/benchmark.nim b/example/benchmark.nim index e2ef1c1a..cc67b7b8 100644 --- a/example/benchmark.nim +++ b/example/benchmark.nim @@ -27,6 +27,7 @@ let sqlitePath = getEnv("SQLITE_PATH", "db.sqlite3") + mysqlUrl = getEnv("MYSQL_URL", "mysql://user:pass@mysql:3306/database") database = getEnv("DB_DATABASE", "database") user = getEnv("DB_USER", "user") password = getEnv("DB_PASSWORD", "pass") @@ -43,7 +44,7 @@ let surrealPort = getEnv("SURREAL_PORT", "8000").parseInt -template benchmarkScenario(rdb: untyped): untyped = +template benchmarkScenario(rdb: untyped, useBackticks: static[bool]): untyped = proc migrate() {.async.} = rdb.create( table("World", [ @@ -78,6 +79,31 @@ template benchmarkScenario(rdb: untyped): untyped = await all(futures) return response + proc benchUpdatePrepared(): Future[seq[JsonNode]] {.async.} = + when isExistsMariaDB or isExistsMySQL: + let selectSql = """SELECT `index` as id, `randomNumber` FROM `World` WHERE `index` = ?""" + let updateSql = """UPDATE `World` SET `randomNumber` = ? WHERE `index` = ?""" + else: + let selectSql = """SELECT "index" as id, "randomNumber" FROM "World" WHERE "index" = ?""" + let updateSql = """UPDATE "World" SET "randomNumber" = ? WHERE "index" = ?""" + + let selectStmt = rdb.prepare(selectSql) + let updateStmt = rdb.prepare(updateSql) + var response = newSeq[JsonNode](countNum) + var futures = newSeq[Future[void]](countNum) + for i in 1..countNum: + let index = rand(range1_10000) + let number = rand(range1_10000) + futures[i - 1] = (proc(): Future[void] {.async.} = + discard await selectStmt.first(@[$index]) + await updateStmt.exec(@[$number, $index]) + )() + response[i - 1] = %*{"id": index, "randomNumber": number} + await all(futures) + await selectStmt.close() + await updateStmt.close() + return response + proc timeProcess[T](name: system.string, cb: proc(): Future[T]) {.async.} = var eachTime = 0.0 var sumTime = 0.0 @@ -102,36 +128,46 @@ template benchmarkScenario(rdb: untyped): untyped = migrate().waitFor waitFor timeProcess("update", benchUpdate) + when compiles(rdb.prepare("SELECT 1")): + waitFor timeProcess("update prepared", benchUpdatePrepared) when isExistsSqlite: proc runSqlite() = echo "=== sqlite" let rdb = dbOpen(SQLite3, sqlitePath, maxConnections, timeout, shouldDisplayLog=shouldDisplayLog) - benchmarkScenario(rdb) + benchmarkScenario(rdb, false) + +when isExistsMysql: + proc runMysql() = + echo "=== mysql" + let rdb = dbOpen(MySQL, mysqlUrl, maxConnections, timeout, shouldDisplayLog=shouldDisplayLog) + benchmarkScenario(rdb, true) when isExistsMariadb: proc runMariadb() = echo "=== mariadb" - let rdb = dbOpen(MariaDB, database, user, password, mariaHost, mariaPort, maxConnections, timeout, shouldDisplayLog=shouldDisplayLog) - benchmarkScenario(rdb) + let rdb = dbOpen(MariaDB, "mariadb://user:pass@mariadb:3306/database", maxConnections, timeout, shouldDisplayLog=shouldDisplayLog) + benchmarkScenario(rdb, true) when isExistsPostgres: proc runPostgres() = echo "=== postgres" - let rdb = dbOpen(PostgreSQL, database, user, password, pgHost, pgPort, maxConnections, timeout, shouldDisplayLog=shouldDisplayLog) - benchmarkScenario(rdb) + let rdb = dbOpen(PostgreSQL, "postgresql://user:pass@postgres:5432/database", maxConnections, timeout, shouldDisplayLog=shouldDisplayLog) + benchmarkScenario(rdb, false) when isExistsSurrealdb: proc runSurreal() = echo "=== surrealdb" let rdb = waitFor dbOpen(SurrealDB, surrealNamespace, surrealDatabase, surrealUser, surrealPassword, surrealHost, surrealPort, maxConnections, timeout, shouldDisplayLog=shouldDisplayLog) - benchmarkScenario(rdb) + benchmarkScenario(rdb, false) proc main() = when isExistsSqlite: runSqlite() + when isExistsMysql: + runMysql() when isExistsMariadb: runMariadb() when isExistsPostgres: diff --git a/src/allographer/query_builder/libs/mariadb/mariadb_impl.nim b/src/allographer/query_builder/libs/mariadb/mariadb_impl.nim index 4d42d1d6..c20af556 100644 --- a/src/allographer/query_builder/libs/mariadb/mariadb_impl.nim +++ b/src/allographer/query_builder/libs/mariadb/mariadb_impl.nim @@ -4,6 +4,8 @@ import std/times import std/json import ../../error import ../../models/database_types +import ../../models/mariadb/mariadb_types +import ../../prepared_param import ./mariadb_rdb import ./mariadb_lib @@ -136,6 +138,243 @@ proc rawExec(conn: PMySQL, query: string, args: MariadbParams, timeout: int) {.a await runRealQuery(conn, q, deadline) +proc runStmtPrepare(conn: PMySQL, stmt: PSTMT, sql: string, deadline: MonoTime): Future[void] {.async.} = + var ret = 0.cint + var waitStatus = stmt_prepare_start(addr ret, stmt, sql.cstring, culong(sql.len)) + while waitStatus != 0: + let ready = await waitMariadb(conn, waitStatus, deadline) + waitStatus = stmt_prepare_cont(addr ret, stmt, ready) + if ret != 0: + raise newException(DbError, $stmt_error(stmt)) + + +proc runStmtReset(conn: PMySQL, stmt: PSTMT, deadline: MonoTime): Future[void] {.async.} = + var ret = false + var waitStatus = stmt_reset_start(addr ret, stmt) + while waitStatus != 0: + let ready = await waitMariadb(conn, waitStatus, deadline) + waitStatus = stmt_reset_cont(addr ret, stmt, ready) + if ret: + raise newException(DbError, $stmt_error(stmt)) + + +proc runStmtFreeResult(conn: PMySQL, stmt: PSTMT, deadline: MonoTime): Future[void] {.async.} = + var ret = false + var waitStatus = stmt_free_result_start(addr ret, stmt) + while waitStatus != 0: + let ready = await waitMariadb(conn, waitStatus, deadline) + waitStatus = stmt_free_result_cont(addr ret, stmt, ready) + if ret: + raise newException(DbError, $stmt_error(stmt)) + + +proc runStmtExecute(conn: PMySQL, stmt: PSTMT, deadline: MonoTime): Future[void] {.async.} = + var ret = 0.cint + var waitStatus = stmt_execute_start(addr ret, stmt) + while waitStatus != 0: + let ready = await waitMariadb(conn, waitStatus, deadline) + waitStatus = stmt_execute_cont(addr ret, stmt, ready) + if ret != 0: + raise newException(DbError, $stmt_error(stmt)) + + +proc runStmtStoreResult(conn: PMySQL, stmt: PSTMT, deadline: MonoTime): Future[void] {.async.} = + var ret = 0.cint + var waitStatus = stmt_store_result_start(addr ret, stmt) + while waitStatus != 0: + let ready = await waitMariadb(conn, waitStatus, deadline) + waitStatus = stmt_store_result_cont(addr ret, stmt, ready) + if ret != 0: + raise newException(DbError, $stmt_error(stmt)) + + +proc runStmtFetch(conn: PMySQL, stmt: PSTMT, deadline: MonoTime): Future[cint] {.async.} = + var ret = 0.cint + var waitStatus = stmt_fetch_start(addr ret, stmt) + while waitStatus != 0: + let ready = await waitMariadb(conn, waitStatus, deadline) + waitStatus = stmt_fetch_cont(addr ret, stmt, ready) + return ret + + +proc prepareStmt*(conn: PMySQL, sql: string, timeout: int): Future[PSTMT] {.async.} = + assert(not conn.isNil, "Database not connected.") + result = stmt_init(conn) + if result.isNil: + dbError(conn) + let deadline = makeDeadline(timeout) + try: + await runStmtPrepare(conn, result, sql, deadline) + except CatchableError: + discard stmt_close(result) + raise + + +proc bindStmtParams(stmt: PSTMT, args: seq[PreparedParam]) = + if stmt_param_count(stmt) != args.len: + raise newException(DbError, "Prepared statement parameter count mismatch.") + + if args.len == 0: + return + + var binds = newSeq[BIND](args.len) + var values = newSeq[string](args.len) + var lengths = newSeq[culong](args.len) + var nullFlags = newSeq[my_bool](args.len) + var errorFlags = newSeq[my_bool](args.len) + + for i, arg in args: + if arg.isNull: + nullFlags[i] = true + binds[i].buffer_type = TYPE_NULL + binds[i].is_null = addr nullFlags[i] + binds[i].error = addr errorFlags[i] + continue + + values[i] = arg.value + lengths[i] = values[i].len.culong + binds[i].buffer_type = TYPE_STRING + if values[i].len > 0: + binds[i].buffer = cast[pointer](values[i].cstring) + else: + binds[i].buffer = nil + binds[i].buffer_length = values[i].len.culong + binds[i].length = addr lengths[i] + binds[i].is_null = addr nullFlags[i] + binds[i].error = addr errorFlags[i] + + if stmt_bind_param(stmt, binds[0].addr): + raise newException(DbError, $stmt_error(stmt)) + + +proc bindStmtResults( + stmt: PSTMT, + metadata: PRES, + resultBinds: MariadbResultBindCache +) = + let cols = int(num_fields(metadata)) + if resultBinds.binds.len != cols: + resultBinds.binds = newSeq[BIND](cols) + resultBinds.buffers = newSeq[string](cols) + resultBinds.lengths = newSeq[culong](cols) + resultBinds.nullFlags = newSeq[my_bool](cols) + resultBinds.errorFlags = newSeq[my_bool](cols) + for i in 0 ..< cols: + let field = fetch_field_direct(metadata, cast[mariadb_rdb.cuint](i)) + var bufferLen = int(field.len) + if bufferLen < 4096: + bufferLen = 4096 + resultBinds.buffers[i] = newString(bufferLen) + else: + for i in 0 ..< cols: + resultBinds.lengths[i] = 0 + resultBinds.nullFlags[i] = false + resultBinds.errorFlags[i] = false + + for i in 0 ..< cols: + resultBinds.binds[i].buffer_type = TYPE_STRING + resultBinds.binds[i].buffer = if resultBinds.buffers[i].len > 0: cast[pointer](resultBinds.buffers[i].cstring) else: nil + resultBinds.binds[i].buffer_length = resultBinds.buffers[i].len.culong + resultBinds.binds[i].length = addr resultBinds.lengths[i] + resultBinds.binds[i].is_null = addr resultBinds.nullFlags[i] + resultBinds.binds[i].error = addr resultBinds.errorFlags[i] + + if cols > 0 and stmt_bind_result(stmt, resultBinds.binds[0].addr): + raise newException(DbError, $stmt_error(stmt)) + + +proc refetchTruncatedColumns( + stmt: PSTMT, + resultBinds: MariadbResultBindCache +) = + for i in 0 ..< resultBinds.binds.len: + if not resultBinds.errorFlags[i]: + continue + let needed = max(int(resultBinds.lengths[i]), resultBinds.buffers[i].len) + if needed <= 0: + continue + resultBinds.buffers[i] = newString(needed) + resultBinds.binds[i].buffer = cast[pointer](resultBinds.buffers[i].cstring) + resultBinds.binds[i].buffer_length = needed.culong + if stmt_fetch_column(stmt, resultBinds.binds[i].addr, cast[mariadb_rdb.cuint](i), 0) != 0: + raise newException(DbError, $stmt_error(stmt)) + + +proc execPreparedStmt*(conn: PMySQL, stmt: PSTMT, args: seq[PreparedParam], timeout: int) {.async.} = + assert(not conn.isNil, "Database not connected.") + let deadline = makeDeadline(timeout) + await runStmtReset(conn, stmt, deadline) + await runStmtFreeResult(conn, stmt, deadline) + bindStmtParams(stmt, args) + await runStmtExecute(conn, stmt, deadline) + await runStmtFreeResult(conn, stmt, deadline) + + +proc queryPreparedStmt*( + conn: PMySQL, + stmt: PSTMT, + args: seq[PreparedParam], + timeout: int, + resultBinds: MariadbResultBindCache +): Future[(seq[database_types.Row], DbRows)] {.async.} = + assert(not conn.isNil, "Database not connected.") + let deadline = makeDeadline(timeout) + await runStmtReset(conn, stmt, deadline) + await runStmtFreeResult(conn, stmt, deadline) + bindStmtParams(stmt, args) + await runStmtExecute(conn, stmt, deadline) + + var dbRows: DbRows + var rows = newSeq[seq[string]]() + let metadata = stmt_result_metadata(stmt) + if metadata.isNil: + await runStmtFreeResult(conn, stmt, deadline) + return (rows, dbRows) + + defer: + free_result(metadata) + await runStmtFreeResult(conn, stmt, deadline) + + await runStmtStoreResult(conn, stmt, deadline) + + let cols = int(num_fields(metadata)) + var baseColumns: DbColumns + setColumnInfo(baseColumns, metadata, cols) + bindStmtResults(stmt, metadata, resultBinds) + + while true: + let fetchRes = await runStmtFetch(conn, stmt, deadline) + if fetchRes == 100: + break + if fetchRes notin {0, 101}: + raise newException(DbError, $stmt_error(stmt)) + if fetchRes == 101: + refetchTruncatedColumns(stmt, resultBinds) + + var rowColumns = baseColumns + var row = newSeq[string](cols) + for i in 0 ..< cols: + if resultBinds.nullFlags[i]: + rowColumns[i].typ.kind = dbNull + row[i] = "" + else: + let length = min(int(resultBinds.lengths[i]), resultBinds.buffers[i].len) + if length <= 0: + row[i] = "" + else: + row[i] = resultBinds.buffers[i][0 ..< length] + rows.add(row) + dbRows.add(rowColumns) + + return (rows, dbRows) + + +proc closePreparedStmt*(stmt: PSTMT): void = + if stmt.isNil: + return + discard stmt_close(stmt) + + proc jsonObjValuesToStrSeq(args: JsonNode): seq[string] = result = newSeq[string](args.len) var i = 0 diff --git a/src/allographer/query_builder/libs/mariadb/mariadb_rdb.nim b/src/allographer/query_builder/libs/mariadb/mariadb_rdb.nim index 01d2cc17..ee2e845e 100644 --- a/src/allographer/query_builder/libs/mariadb/mariadb_rdb.nim +++ b/src/allographer/query_builder/libs/mariadb/mariadb_rdb.nim @@ -833,8 +833,8 @@ type STMT* = St_mysql_stmt -# Enum_stmt_attr_type* = enum -# STMT_ATTR_UPDATE_MAX_LENGTH, STMT_ATTR_CURSOR_TYPE, STMT_ATTR_PREFETCH_ROWS + Enum_stmt_attr_type* = enum + STMT_ATTR_UPDATE_MAX_LENGTH, STMT_ATTR_CURSOR_TYPE, STMT_ATTR_PREFETCH_ROWS # {.deprecated: [Tst_dynamic_array: St_dynamic_array, Tst_mysql_options: St_mysql_options, # TDYNAMIC_ARRAY: DYNAMIC_ARRAY, Tprotocol_type: Protocol_type, # Trpl_type: Rpl_type, Tcharset_info_st: Charset_info_st, @@ -1095,45 +1095,65 @@ proc store_result_start*(ret: ptr PRES, MySQL: PMySQL): cint{.stdcall, dynlib: l proc store_result_cont*(ret: ptr PRES, MySQL: PMySQL, status: cint): cint{.stdcall, dynlib: lib, importc: "mysql_store_result_cont".} proc stmt_init*(MySQL: PMySQL): PSTMT{.stdcall, dynlib: lib, importc: "mysql_stmt_init".} -# proc stmt_prepare*(stmt: PSTMT, query: cstring, len: int): cint{.stdcall, -# dynlib: lib, importc: "mysql_stmt_prepare".} -# proc stmt_execute*(stmt: PSTMT): cint{.stdcall, dynlib: lib, -# importc: "mysql_stmt_execute".} -# proc stmt_fetch*(stmt: PSTMT): cint{.stdcall, dynlib: lib, -# importc: "mysql_stmt_fetch".} -# proc stmt_fetch_column*(stmt: PSTMT, `bind`: PBIND, column: cuint, offset: int): cint{. -# stdcall, dynlib: lib, importc: "mysql_stmt_fetch_column".} -# proc stmt_store_result*(stmt: PSTMT): cint{.stdcall, dynlib: lib, -# importc: "mysql_stmt_store_result".} -# proc stmt_param_count*(stmt: PSTMT): int{.stdcall, dynlib: lib, -# importc: "mysql_stmt_param_count".} -# proc stmt_attr_set*(stmt: PSTMT, attr_type: Enum_stmt_attr_type, attr: pointer): my_bool{. -# stdcall, dynlib: lib, importc: "mysql_stmt_attr_set".} -# proc stmt_attr_get*(stmt: PSTMT, attr_type: Enum_stmt_attr_type, attr: pointer): my_bool{. -# stdcall, dynlib: lib, importc: "mysql_stmt_attr_get".} -# proc stmt_bind_param*(stmt: PSTMT, bnd: PBIND): my_bool{.stdcall, dynlib: lib, -# importc: "mysql_stmt_bind_param".} -# proc stmt_bind_result*(stmt: PSTMT, bnd: PBIND): my_bool{.stdcall, dynlib: lib, -# importc: "mysql_stmt_bind_result".} -# proc stmt_close*(stmt: PSTMT): my_bool{.stdcall, dynlib: lib, -# importc: "mysql_stmt_close".} -# proc stmt_reset*(stmt: PSTMT): my_bool{.stdcall, dynlib: lib, -# importc: "mysql_stmt_reset".} -# proc stmt_free_result*(stmt: PSTMT): my_bool{.stdcall, dynlib: lib, -# importc: "mysql_stmt_free_result".} +proc stmt_prepare_start*(ret: ptr cint, stmt: PSTMT, query: cstring, len: culong): cint{.stdcall, + dynlib: lib, importc: "mysql_stmt_prepare_start".} +proc stmt_prepare_cont*(ret: ptr cint, stmt: PSTMT, status: cint): cint{.stdcall, + dynlib: lib, importc: "mysql_stmt_prepare_cont".} +proc stmt_execute_start*(ret: ptr cint, stmt: PSTMT): cint{.stdcall, dynlib: lib, + importc: "mysql_stmt_execute_start".} +proc stmt_execute_cont*(ret: ptr cint, stmt: PSTMT, status: cint): cint{.stdcall, dynlib: lib, + importc: "mysql_stmt_execute_cont".} +proc stmt_fetch_start*(ret: ptr cint, stmt: PSTMT): cint{.stdcall, dynlib: lib, + importc: "mysql_stmt_fetch_start".} +proc stmt_fetch_cont*(ret: ptr cint, stmt: PSTMT, status: cint): cint{.stdcall, dynlib: lib, + importc: "mysql_stmt_fetch_cont".} +proc stmt_fetch_column*(stmt: PSTMT, `bind`: PBIND, column: cuint, offset: int): cint{. + stdcall, dynlib: lib, importc: "mysql_stmt_fetch_column".} +proc stmt_store_result_start*(ret: ptr cint, stmt: PSTMT): cint{.stdcall, dynlib: lib, + importc: "mysql_stmt_store_result_start".} +proc stmt_store_result_cont*(ret: ptr cint, stmt: PSTMT, status: cint): cint{.stdcall, dynlib: lib, + importc: "mysql_stmt_store_result_cont".} +proc stmt_param_count*(stmt: PSTMT): int{.stdcall, dynlib: lib, + importc: "mysql_stmt_param_count".} +proc stmt_attr_set*(stmt: PSTMT, attr_type: Enum_stmt_attr_type, attr: pointer): my_bool{. + stdcall, dynlib: lib, importc: "mysql_stmt_attr_set".} +proc stmt_attr_get*(stmt: PSTMT, attr_type: Enum_stmt_attr_type, attr: pointer): my_bool{. + stdcall, dynlib: lib, importc: "mysql_stmt_attr_get".} +proc stmt_bind_param*(stmt: PSTMT, bnd: PBIND): my_bool{.stdcall, dynlib: lib, + importc: "mysql_stmt_bind_param".} +proc stmt_bind_result*(stmt: PSTMT, bnd: PBIND): my_bool{.stdcall, dynlib: lib, + importc: "mysql_stmt_bind_result".} +proc stmt_close*(stmt: PSTMT): my_bool{.stdcall, dynlib: lib, + importc: "mysql_stmt_close".} +proc stmt_reset*(stmt: PSTMT): my_bool{.stdcall, dynlib: lib, + importc: "mysql_stmt_reset".} +proc stmt_free_result*(stmt: PSTMT): my_bool{.stdcall, dynlib: lib, + importc: "mysql_stmt_free_result".} +proc stmt_reset_start*(ret: ptr my_bool, stmt: PSTMT): cint{.stdcall, dynlib: lib, + importc: "mysql_stmt_reset_start".} +proc stmt_reset_cont*(ret: ptr my_bool, stmt: PSTMT, status: cint): cint{.stdcall, dynlib: lib, + importc: "mysql_stmt_reset_cont".} +proc stmt_free_result_start*(ret: ptr my_bool, stmt: PSTMT): cint{.stdcall, dynlib: lib, + importc: "mysql_stmt_free_result_start".} +proc stmt_free_result_cont*(ret: ptr my_bool, stmt: PSTMT, status: cint): cint{.stdcall, dynlib: lib, + importc: "mysql_stmt_free_result_cont".} +proc stmt_close_start*(ret: ptr my_bool, stmt: PSTMT): cint{.stdcall, dynlib: lib, + importc: "mysql_stmt_close_start".} +proc stmt_close_cont*(ret: ptr my_bool, stmt: PSTMT, status: cint): cint{.stdcall, dynlib: lib, + importc: "mysql_stmt_close_cont".} +proc stmt_result_metadata*(stmt: PSTMT): PRES{.stdcall, dynlib: lib, + importc: "mysql_stmt_result_metadata".} +proc stmt_errno*(stmt: PSTMT): cuint{.stdcall, dynlib: lib, + importc: "mysql_stmt_errno".} +proc stmt_error*(stmt: PSTMT): cstring{.stdcall, dynlib: lib, + importc: "mysql_stmt_error".} +proc stmt_sqlstate*(stmt: PSTMT): cstring{.stdcall, dynlib: lib, + importc: "mysql_stmt_sqlstate".} # proc stmt_send_long_data*(stmt: PSTMT, param_number: cuint, data: cstring, # len: int): my_bool{.stdcall, dynlib: lib, # importc: "mysql_stmt_send_long_data".} -# proc stmt_result_metadata*(stmt: PSTMT): PRES{.stdcall, dynlib: lib, -# importc: "mysql_stmt_result_metadata".} # proc stmt_param_metadata*(stmt: PSTMT): PRES{.stdcall, dynlib: lib, # importc: "mysql_stmt_param_metadata".} -# proc stmt_errno*(stmt: PSTMT): cuint{.stdcall, dynlib: lib, -# importc: "mysql_stmt_errno".} -# proc stmt_error*(stmt: PSTMT): cstring{.stdcall, dynlib: lib, -# importc: "mysql_stmt_error".} -# proc stmt_sqlstate*(stmt: PSTMT): cstring{.stdcall, dynlib: lib, -# importc: "mysql_stmt_sqlstate".} # proc stmt_row_seek*(stmt: PSTMT, offset: ROW_OFFSET): ROW_OFFSET{.stdcall, # dynlib: lib, importc: "mysql_stmt_row_seek".} # proc stmt_row_tell*(stmt: PSTMT): ROW_OFFSET{.stdcall, dynlib: lib, diff --git a/src/allographer/query_builder/libs/mysql/mysql_impl.nim b/src/allographer/query_builder/libs/mysql/mysql_impl.nim index f9f3c743..f4d8637e 100644 --- a/src/allographer/query_builder/libs/mysql/mysql_impl.nim +++ b/src/allographer/query_builder/libs/mysql/mysql_impl.nim @@ -5,6 +5,8 @@ import std/strformat import std/json import ../../error import ../../models/database_types +import ../../models/mysql/mysql_types +import ../../prepared_param import ./mysql_rdb import ./mysql_lib @@ -33,6 +35,193 @@ proc rawExec(conn:PMySQL, query: string, args: MysqlParams) = if realQuery(conn, q.cstring, q.len) != 0'i32: dbError(conn) +proc prepareStmt*(conn: PMySQL, sql: string, timeout: int): Future[PSTMT] {.async.} = + assert(not conn.isNil, "Database not connected.") + await sleepAsync(0) + result = mysql_rdb.stmt_init(conn) + if result.isNil: + dbError(conn) + if mysql_rdb.stmt_prepare(result, sql.cstring, sql.len) != 0: + let errmsg = $mysql_rdb.stmt_error(result) + discard mysql_rdb.stmt_close(result) + raise newException(DbError, errmsg) + + +proc bindStmtParams(stmt: PSTMT, args: seq[PreparedParam]) = + if mysql_rdb.stmt_param_count(stmt) != args.len: + raise newException(DbError, "Prepared statement parameter count mismatch.") + + if args.len == 0: + return + + var binds = newSeq[BIND](args.len) + var values = newSeq[string](args.len) + var lengths = newSeq[culong](args.len) + var nullFlags = newSeq[my_bool](args.len) + var errorFlags = newSeq[my_bool](args.len) + + for i, arg in args: + if arg.isNull: + nullFlags[i] = true + binds[i].buffer_type = TYPE_NULL + binds[i].is_null = addr nullFlags[i] + binds[i].error = addr errorFlags[i] + continue + + values[i] = arg.value + lengths[i] = values[i].len.culong + binds[i].buffer_type = TYPE_STRING + if values[i].len > 0: + binds[i].buffer = cast[pointer](values[i].cstring) + else: + binds[i].buffer = nil + binds[i].buffer_length = values[i].len.culong + binds[i].length = addr lengths[i] + binds[i].is_null = addr nullFlags[i] + binds[i].error = addr errorFlags[i] + + if mysql_rdb.stmt_bind_param(stmt, binds[0].addr): + raise newException(DbError, $mysql_rdb.stmt_error(stmt)) + + +proc bindStmtResults( + stmt: PSTMT, + metadata: PRES, + resultBinds: MysqlResultBindCache +) = + let cols = int(mysql_rdb.num_fields(metadata)) + if resultBinds.binds.len != cols: + resultBinds.binds = newSeq[BIND](cols) + resultBinds.buffers = newSeq[string](cols) + resultBinds.lengths = newSeq[culong](cols) + resultBinds.nullFlags = newSeq[my_bool](cols) + resultBinds.errorFlags = newSeq[my_bool](cols) + for i in 0 ..< cols: + let field = mysql_rdb.fetch_field_direct(metadata, cast[mysql_rdb.cuint](i)) + var bufferLen = int(field.len) + if bufferLen < 4096: + bufferLen = 4096 + resultBinds.buffers[i] = newString(bufferLen) + else: + for i in 0 ..< cols: + resultBinds.lengths[i] = 0 + resultBinds.nullFlags[i] = false + resultBinds.errorFlags[i] = false + + for i in 0 ..< cols: + resultBinds.binds[i].buffer_type = TYPE_STRING + resultBinds.binds[i].buffer = if resultBinds.buffers[i].len > 0: cast[pointer](resultBinds.buffers[i].cstring) else: nil + resultBinds.binds[i].buffer_length = resultBinds.buffers[i].len.culong + resultBinds.binds[i].length = addr resultBinds.lengths[i] + resultBinds.binds[i].is_null = addr resultBinds.nullFlags[i] + resultBinds.binds[i].error = addr resultBinds.errorFlags[i] + + if cols > 0 and mysql_rdb.stmt_bind_result(stmt, resultBinds.binds[0].addr): + raise newException(DbError, $mysql_rdb.stmt_error(stmt)) + + +proc refetchTruncatedColumns( + stmt: PSTMT, + resultBinds: MysqlResultBindCache +) = + for i in 0 ..< resultBinds.binds.len: + if not resultBinds.errorFlags[i]: + continue + let needed = max(int(resultBinds.lengths[i]), resultBinds.buffers[i].len) + if needed <= 0: + continue + resultBinds.buffers[i] = newString(needed) + resultBinds.binds[i].buffer = cast[pointer](resultBinds.buffers[i].cstring) + resultBinds.binds[i].buffer_length = needed.culong + if mysql_rdb.stmt_fetch_column(stmt, resultBinds.binds[i].addr, cast[mysql_rdb.cuint](i), 0) != 0: + raise newException(DbError, $mysql_rdb.stmt_error(stmt)) + + +proc execPreparedStmt*(conn: PMySQL, stmt: PSTMT, args: seq[PreparedParam], timeout: int) {.async.} = + assert(not conn.isNil, "Database not connected.") + await sleepAsync(0) + if mysql_rdb.stmt_reset(stmt): + raise newException(DbError, $mysql_rdb.stmt_error(stmt)) + if mysql_rdb.stmt_free_result(stmt): + raise newException(DbError, $mysql_rdb.stmt_error(stmt)) + bindStmtParams(stmt, args) + if mysql_rdb.stmt_execute(stmt) != 0: + raise newException(DbError, $mysql_rdb.stmt_error(stmt)) + if mysql_rdb.stmt_free_result(stmt): + raise newException(DbError, $mysql_rdb.stmt_error(stmt)) + + +proc queryPreparedStmt*( + conn: PMySQL, + stmt: PSTMT, + args: seq[PreparedParam], + timeout: int, + resultBinds: MysqlResultBindCache +): Future[(seq[database_types.Row], DbRows)] {.async.} = + assert(not conn.isNil, "Database not connected.") + await sleepAsync(0) + if mysql_rdb.stmt_reset(stmt): + raise newException(DbError, $mysql_rdb.stmt_error(stmt)) + if mysql_rdb.stmt_free_result(stmt): + raise newException(DbError, $mysql_rdb.stmt_error(stmt)) + bindStmtParams(stmt, args) + if mysql_rdb.stmt_execute(stmt) != 0: + raise newException(DbError, $mysql_rdb.stmt_error(stmt)) + + var dbRows: DbRows + var rows = newSeq[seq[string]]() + let metadata = mysql_rdb.stmt_result_metadata(stmt) + if metadata.isNil: + if mysql_rdb.stmt_free_result(stmt): + raise newException(DbError, $mysql_rdb.stmt_error(stmt)) + return (rows, dbRows) + + defer: + mysql_rdb.free_result(metadata) + if mysql_rdb.stmt_free_result(stmt): + raise newException(DbError, $mysql_rdb.stmt_error(stmt)) + + if mysql_rdb.stmt_store_result(stmt) != 0: + raise newException(DbError, $mysql_rdb.stmt_error(stmt)) + + let cols = int(mysql_rdb.num_fields(metadata)) + var baseColumns: DbColumns + setColumnInfo(baseColumns, metadata, cols) + bindStmtResults(stmt, metadata, resultBinds) + + while true: + let fetchRes = mysql_rdb.stmt_fetch(stmt) + if fetchRes == 100: + break + if fetchRes notin {0, 101}: + raise newException(DbError, $mysql_rdb.stmt_error(stmt)) + if fetchRes == 101: + refetchTruncatedColumns(stmt, resultBinds) + + var rowColumns = baseColumns + var row = newSeq[string](cols) + for i in 0 ..< cols: + if resultBinds.nullFlags[i]: + rowColumns[i].typ.kind = dbNull + row[i] = "" + else: + let length = min(int(resultBinds.lengths[i]), resultBinds.buffers[i].len) + if length <= 0: + row[i] = "" + else: + row[i] = resultBinds.buffers[i][0 ..< length] + rows.add(row) + dbRows.add(rowColumns) + + return (rows, dbRows) + + +proc closePreparedStmt*(stmt: PSTMT) = + if stmt.isNil: + return + discard mysql_rdb.stmt_close(stmt) + + proc query*(db:PMySQL, query: string, args: seq[string], timeout:int):Future[(seq[database_types.Row], DbRows)] {.async.} = assert db.ping == 0 var dbRows: DbRows diff --git a/src/allographer/query_builder/libs/mysql/mysql_rdb.nim b/src/allographer/query_builder/libs/mysql/mysql_rdb.nim index 39379123..1fbfe561 100644 --- a/src/allographer/query_builder/libs/mysql/mysql_rdb.nim +++ b/src/allographer/query_builder/libs/mysql/mysql_rdb.nim @@ -1053,43 +1053,39 @@ proc real_escape_string*(MySQL: PMySQL, fto: cstring, `from`: cstring, len: int) # proc read_query_result*(MySQL: PMySQL): my_bool{.stdcall, dynlib: lib, # importc: "mysql_read_query_result".} proc stmt_init*(MySQL: PMySQL): PSTMT{.stdcall, dynlib: lib, importc: "mysql_stmt_init".} -# proc stmt_prepare*(stmt: PSTMT, query: cstring, len: int): cint{.stdcall, -# dynlib: lib, importc: "mysql_stmt_prepare".} -# proc stmt_execute*(stmt: PSTMT): cint{.stdcall, dynlib: lib, -# importc: "mysql_stmt_execute".} -# proc stmt_fetch*(stmt: PSTMT): cint{.stdcall, dynlib: lib, -# importc: "mysql_stmt_fetch".} -# proc stmt_fetch_column*(stmt: PSTMT, `bind`: PBIND, column: cuint, offset: int): cint{. -# stdcall, dynlib: lib, importc: "mysql_stmt_fetch_column".} -# proc stmt_store_result*(stmt: PSTMT): cint{.stdcall, dynlib: lib, -# importc: "mysql_stmt_store_result".} -# proc stmt_param_count*(stmt: PSTMT): int{.stdcall, dynlib: lib, -# importc: "mysql_stmt_param_count".} -# proc stmt_attr_set*(stmt: PSTMT, attr_type: Enum_stmt_attr_type, attr: pointer): my_bool{. -# stdcall, dynlib: lib, importc: "mysql_stmt_attr_set".} -# proc stmt_attr_get*(stmt: PSTMT, attr_type: Enum_stmt_attr_type, attr: pointer): my_bool{. -# stdcall, dynlib: lib, importc: "mysql_stmt_attr_get".} -# proc stmt_bind_param*(stmt: PSTMT, bnd: PBIND): my_bool{.stdcall, dynlib: lib, -# importc: "mysql_stmt_bind_param".} -# proc stmt_bind_result*(stmt: PSTMT, bnd: PBIND): my_bool{.stdcall, dynlib: lib, -# importc: "mysql_stmt_bind_result".} -# proc stmt_close*(stmt: PSTMT): my_bool{.stdcall, dynlib: lib, -# importc: "mysql_stmt_close".} -# proc stmt_reset*(stmt: PSTMT): my_bool{.stdcall, dynlib: lib, -# importc: "mysql_stmt_reset".} -# proc stmt_free_result*(stmt: PSTMT): my_bool{.stdcall, dynlib: lib, -# importc: "mysql_stmt_free_result".} +proc stmt_prepare*(stmt: PSTMT, query: cstring, len: int): cint{.stdcall, + dynlib: lib, importc: "mysql_stmt_prepare".} +proc stmt_execute*(stmt: PSTMT): cint{.stdcall, dynlib: lib, + importc: "mysql_stmt_execute".} +proc stmt_fetch*(stmt: PSTMT): cint{.stdcall, dynlib: lib, + importc: "mysql_stmt_fetch".} +proc stmt_fetch_column*(stmt: PSTMT, `bind`: PBIND, column: cuint, offset: int): cint{. + stdcall, dynlib: lib, importc: "mysql_stmt_fetch_column".} +proc stmt_store_result*(stmt: PSTMT): cint{.stdcall, dynlib: lib, + importc: "mysql_stmt_store_result".} +proc stmt_param_count*(stmt: PSTMT): int{.stdcall, dynlib: lib, + importc: "mysql_stmt_param_count".} +proc stmt_bind_param*(stmt: PSTMT, bnd: PBIND): my_bool{.stdcall, dynlib: lib, + importc: "mysql_stmt_bind_param".} +proc stmt_bind_result*(stmt: PSTMT, bnd: PBIND): my_bool{.stdcall, dynlib: lib, + importc: "mysql_stmt_bind_result".} +proc stmt_close*(stmt: PSTMT): my_bool{.stdcall, dynlib: lib, + importc: "mysql_stmt_close".} +proc stmt_reset*(stmt: PSTMT): my_bool{.stdcall, dynlib: lib, + importc: "mysql_stmt_reset".} +proc stmt_free_result*(stmt: PSTMT): my_bool{.stdcall, dynlib: lib, + importc: "mysql_stmt_free_result".} # proc stmt_send_long_data*(stmt: PSTMT, param_number: cuint, data: cstring, # len: int): my_bool{.stdcall, dynlib: lib, # importc: "mysql_stmt_send_long_data".} -# proc stmt_result_metadata*(stmt: PSTMT): PRES{.stdcall, dynlib: lib, -# importc: "mysql_stmt_result_metadata".} +proc stmt_result_metadata*(stmt: PSTMT): PRES{.stdcall, dynlib: lib, + importc: "mysql_stmt_result_metadata".} # proc stmt_param_metadata*(stmt: PSTMT): PRES{.stdcall, dynlib: lib, # importc: "mysql_stmt_param_metadata".} # proc stmt_errno*(stmt: PSTMT): cuint{.stdcall, dynlib: lib, # importc: "mysql_stmt_errno".} -# proc stmt_error*(stmt: PSTMT): cstring{.stdcall, dynlib: lib, -# importc: "mysql_stmt_error".} +proc stmt_error*(stmt: PSTMT): cstring{.stdcall, dynlib: lib, + importc: "mysql_stmt_error".} # proc stmt_sqlstate*(stmt: PSTMT): cstring{.stdcall, dynlib: lib, # importc: "mysql_stmt_sqlstate".} # proc stmt_row_seek*(stmt: PSTMT, offset: ROW_OFFSET): ROW_OFFSET{.stdcall, diff --git a/src/allographer/query_builder/libs/postgres/postgres_impl.nim b/src/allographer/query_builder/libs/postgres/postgres_impl.nim index 269483b4..bc20c12a 100644 --- a/src/allographer/query_builder/libs/postgres/postgres_impl.nim +++ b/src/allographer/query_builder/libs/postgres/postgres_impl.nim @@ -7,6 +7,7 @@ import std/strutils import std/times import ../../error import ../../models/database_types +import ../../prepared_param import ./postgres_rdb import ./postgres_lib @@ -333,10 +334,9 @@ proc getColumns*(db: PPGconn, query: string, args: seq[string], timeout: int): F result.add(column.name) -proc prepare*(db: PPGconn, query: string, timeout: int, stmtName: string): Future[int] {.async.} = +proc prepare*(db: PPGconn, query: string, timeout: int, stmtName: string, nArgs: int): Future[void] {.async.} = assert db.status == CONNECTION_OK - let nArgs = query.count('$') - let success = pqsendPrepare(db, stmtName, dbFormat(query).cstring, int32(nArgs), nil) + let success = pqsendPrepare(db, stmtName, questionToDaller(query).cstring, int32(nArgs), nil) if success != 1: dbError(db) let deadline = makePgDeadline(timeout) await pgFlushOutgoing(db, deadline) @@ -346,15 +346,52 @@ proc prepare*(db: PPGconn, query: string, timeout: int, stmtName: string): Futur db.checkError() break pqclear(pqresult) - return nArgs -proc preparedQuery*(db: PPGconn, args: seq[string], nArgs: int, timeout: int, stmtName: string): Future[(seq[Row], DbRows)] {.async.} = + +proc deallocate*(db: PPGconn, stmtName: string, timeout: int): Future[void] {.async.} = + assert db.status == CONNECTION_OK + if stmtName.len == 0: + return + let success = pqsendQuery(db, ("DEALLOCATE " & stmtName).cstring) + if success != 1: + dbError(db) + let deadline = makePgDeadline(timeout) + await pgFlushOutgoing(db, deadline) + while true: + let pqresult = await pgNextResult(db, deadline) + if pqresult == nil: + db.checkError() + break + pqclear(pqresult) + +proc allocPreparedCStringArray(args: seq[PreparedParam]): cstringArray = + result = cast[cstringArray](alloc0((args.len + 1) * sizeof(cstring))) + for i, arg in args: + if arg.isNull: + continue + let cstrLen = arg.value.len + 1 + let cstr = cast[cstring](alloc0(cstrLen)) + copyMem(cstr, arg.value.cstring, arg.value.len) + result[i] = cstr + + +proc freePreparedCStringArray(values: cstringArray, n: int) = + if values.isNil: + return + for i in 0 ..< n: + if values[i] != nil: + dealloc(values[i]) + dealloc(values) + + +proc preparedQuery*(db: PPGconn, args: seq[PreparedParam], nArgs: int, timeout: int, stmtName: string): Future[(seq[Row], DbRows)] {.async.} = assert db.status == CONNECTION_OK let deadline = makePgDeadline(timeout) await pgEnsureIdle(db, deadline) - let arr = allocCStringArray(args) - let status = pqsendQueryPrepared(db, stmtName, int32(nArgs), arr, nil, nil, 0) - deallocCStringArray(arr) + let values = allocPreparedCStringArray(args) + defer: + freePreparedCStringArray(values, args.len) + let status = pqsendQueryPrepared(db, stmtName, int32(nArgs), values, nil, nil, 0) if status != 1: dbError(db) var dbRows: DbRows var rows = newSeq[Row]() @@ -376,13 +413,14 @@ proc preparedQuery*(db: PPGconn, args: seq[string], nArgs: int, timeout: int, st return (rows, dbRows) -proc preparedExec*(db: PPGconn, args: seq[string], nArgs: int, timeout: int, stmtName: string) {.async.} = +proc preparedExec*(db: PPGconn, args: seq[PreparedParam], nArgs: int, timeout: int, stmtName: string) {.async.} = assert db.status == CONNECTION_OK let deadline = makePgDeadline(timeout) await pgEnsureIdle(db, deadline) - let arr = allocCStringArray(args) - let status = pqsendQueryPrepared(db, stmtName, int32(nArgs), arr, nil, nil, 0) - deallocCStringArray(arr) + let values = allocPreparedCStringArray(args) + defer: + freePreparedCStringArray(values, args.len) + let status = pqsendQueryPrepared(db, stmtName, int32(nArgs), values, nil, nil, 0) if status != 1: dbError(db) await pgFlushOutgoing(db, deadline) while true: diff --git a/src/allographer/query_builder/libs/sqlite/sqlite_impl.nim b/src/allographer/query_builder/libs/sqlite/sqlite_impl.nim index 1d468941..14f154ec 100644 --- a/src/allographer/query_builder/libs/sqlite/sqlite_impl.nim +++ b/src/allographer/query_builder/libs/sqlite/sqlite_impl.nim @@ -2,6 +2,7 @@ import std/asyncdispatch import std/strutils import std/json import ../../models/database_types +import ../../prepared_param import ./sqlite_rdb import ./sqlite_lib @@ -164,6 +165,69 @@ proc prepare*(db:PSqlite3, query:string, timeout:int):Future[PStmt] {.async.} = dbError(db) +proc bindPreparedParams(db: PSqlite3, stmt: PStmt, args: seq[PreparedParam]) = + if reset(stmt) != SQLITE_OK: + dbError(db) + if clear_bindings(stmt) != SQLITE_OK: + dbError(db) + + for i, arg in args: + let paramIdx = i.int32 + 1 + if arg.isNull: + if bind_null(stmt, paramIdx) != SQLITE_OK: + dbError(db) + else: + if bind_text(stmt, paramIdx, arg.value.cstring, arg.value.len.int32, SQLITE_TRANSIENT) != SQLITE_OK: + dbError(db) + + +proc preparedQueryReuse*(db: PSqlite3, stmt: PStmt, args: seq[PreparedParam], timeout: int, + cachedColumns: DbColumns): Future[(seq[Row], DbRows)] {.async.} = + assert(not db.isNil, "Database not connected.") + sleepAsync(0).await + bindPreparedParams(db, stmt, args) + defer: + discard clear_bindings(stmt) + + var dbRows: DbRows + var rows = newSeq[seq[string]]() + + while true: + let stepRes = step(stmt) + if stepRes == SQLITE_ROW: + var columns = cachedColumns + setColumnsRuntimeTypes(columns, stmt) + dbRows.add(columns) + var row = newSeq[string](int(column_count(stmt))) + for i in 0 ..< row.len: + let text = column_text(stmt, i.int32) + if text.isNil: + row[i] = "" + else: + row[i] = $text + rows.add(row) + continue + if stepRes == SQLITE_DONE: + break + dbError(db) + + return (rows, dbRows) + + +proc preparedExecReuse*(db: PSqlite3, stmt: PStmt, args: seq[PreparedParam], timeout: int) {.async.} = + assert(not db.isNil, "Database not connected.") + sleepAsync(0).await + bindPreparedParams(db, stmt, args) + defer: + discard clear_bindings(stmt) + + var stepRes = step(stmt) + while stepRes == SQLITE_ROW: + stepRes = step(stmt) + if stepRes != SQLITE_DONE: + dbError(db) + + proc preparedQuery*(db:PSqlite3, args:seq[string] = @[], sqliteStmt:PStmt):Future[(seq[Row], DbRows)] {.async.} = # bind params for i, row in args: diff --git a/src/allographer/query_builder/libs/sqlite/sqlite_lib.nim b/src/allographer/query_builder/libs/sqlite/sqlite_lib.nim index 39bdd9db..07238bf0 100644 --- a/src/allographer/query_builder/libs/sqlite/sqlite_lib.nim +++ b/src/allographer/query_builder/libs/sqlite/sqlite_lib.nim @@ -79,7 +79,7 @@ proc toTypeKind(t: var DbType; x: int32) = of SQLITE_TEXT: t.kind = dbVarchar else: t.kind = dbUnknown -proc setColumnsStaticMeta(columns: var DbColumns; x: PStmt) = +proc setColumnsStaticMeta*(columns: var DbColumns; x: PStmt) = ## ステップ前でも列名・宣言型・テーブル名は取得できる(行に依存しない)。 let L = column_count(x) setLen(columns, L.int) @@ -88,7 +88,7 @@ proc setColumnsStaticMeta(columns: var DbColumns; x: PStmt) = columns[i].typ.name = $column_decltype(x, i) columns[i].tableName = $column_table_name(x, i) -proc setColumnsRuntimeTypes(columns: var DbColumns; x: PStmt) = +proc setColumnsRuntimeTypes*(columns: var DbColumns; x: PStmt) = ## 行ごとに変わりうるのは `column_type` のみ。 let L = column_count(x) for i in 0'i32 ..< L: diff --git a/src/allographer/query_builder/models/mariadb/mariadb_exec.nim b/src/allographer/query_builder/models/mariadb/mariadb_exec.nim index 6ab226ad..244a5062 100644 --- a/src/allographer/query_builder/models/mariadb/mariadb_exec.nim +++ b/src/allographer/query_builder/models/mariadb/mariadb_exec.nim @@ -8,8 +8,10 @@ import std/sequtils import std/tables import std/times import ../../libs/mariadb/mariadb_impl +import ../../libs/mariadb/mariadb_rdb except Option, cuint import ../../log import ../database_types +import ../../prepared_param import ./query/mariadb_builder import ./mariadb_types @@ -56,6 +58,22 @@ proc returnConn(self:MariadbConnections | MariadbQuery | RawMariadbQuery, i: int wakeOnePoolWaiter(self.pools) +proc prepare*(self: MariadbConnections, sql: string): MariadbPreparedStatement = + new(result) + result.owner = self + result.info = self.info + result.sql = sql + result.stmts = newSeq[PSTMT](self.pools.conns.len) + result.nArgs = countQuestionMarks(sql) + result.resultBindCache = newSeq[MariadbResultBindCache](self.pools.conns.len) + + +proc ensurePreparedStmt(self: MariadbPreparedStatement, connI: int): Future[PSTMT] {.async.} = + if self.stmts[connI].isNil: + self.stmts[connI] = await mariadb_impl.prepareStmt(self.owner.pools.conns[connI].conn, self.sql, self.owner.pools.timeout) + return self.stmts[connI] + + # ================================================================================ # toJson # ================================================================================ @@ -364,7 +382,7 @@ proc transactionStart(self:MariadbConnections) {.async.} = self.isInTransaction = true self.transactionConn = connI - mariadb_impl.exec(self.pools.conns[connI].conn, "BEGIN", newJArray(), newSeq[Row](), self.pools.timeout).await + mariadb_impl.exec(self.pools.conns[connI].conn, "BEGIN", newJArray(), newSeq[seq[string]](), self.pools.timeout).await proc transactionEnd(self:MariadbConnections, query:string) {.async.} = @@ -373,7 +391,169 @@ proc transactionEnd(self:MariadbConnections, query:string) {.async.} = self.transactionConn = 0 self.isInTransaction = false - mariadb_impl.exec(self.pools.conns[self.transactionConn].conn, query, newJArray(), newSeq[Row](), self.pools.timeout).await + mariadb_impl.exec(self.pools.conns[self.transactionConn].conn, query, newJArray(), newSeq[seq[string]](), self.pools.timeout).await + + +proc getPreparedRows(self: MariadbPreparedStatement, args: seq[PreparedParam]): Future[(seq[seq[string]], DbRows)] {.async.} = + var connI = self.owner.transactionConn + if not self.owner.isInTransaction: + connI = getFreeConn(self.owner).await + defer: + if not self.owner.isInTransaction: + self.owner.returnConn(connI).await + if connI == errorConnectionNum: + return + + let stmt = await self.ensurePreparedStmt(connI) + if connI >= self.resultBindCache.len: + self.resultBindCache.setLen(connI + 1) + if self.resultBindCache[connI].isNil: + new(self.resultBindCache[connI]) + return mariadb_impl.queryPreparedStmt( + self.owner.pools.conns[connI].conn, + stmt, + args, + self.owner.pools.timeout, + self.resultBindCache[connI] + ).await + + +proc getPreparedAllRows(self: MariadbPreparedStatement, args: seq[PreparedParam]): Future[seq[JsonNode]] {.async.} = + let (rows, dbRows) = await self.getPreparedRows(args) + if rows.len == 0: + self.owner.log.echoErrorMsg(self.sql) + return newSeq[JsonNode](0) + return toJson(rows, dbRows) + + +proc getPreparedRow(self: MariadbPreparedStatement, args: seq[PreparedParam]): Future[Option[JsonNode]] {.async.} = + let (rows, dbRows) = await self.getPreparedRows(args) + if rows.len == 0: + self.owner.log.echoErrorMsg(self.sql) + return none(JsonNode) + return toJson(rows, dbRows)[0].some() + + +proc getPreparedAllRowsPlain(self: MariadbPreparedStatement, args: seq[PreparedParam]): Future[seq[seq[string]]] {.async.} = + let (rows, _) = await self.getPreparedRows(args) + return rows + + +proc getPreparedRowPlain(self: MariadbPreparedStatement, args: seq[PreparedParam]): Future[seq[string]] {.async.} = + let (rows, _) = await self.getPreparedRows(args) + if rows.len == 0: + self.owner.log.echoErrorMsg(self.sql) + return newSeq[string](0) + return rows[0] + + +proc execPrepared(self: MariadbPreparedStatement, args: seq[PreparedParam]) {.async.} = + var connI = self.owner.transactionConn + if not self.owner.isInTransaction: + connI = getFreeConn(self.owner).await + defer: + if not self.owner.isInTransaction: + self.owner.returnConn(connI).await + if connI == errorConnectionNum: + return + + let stmt = await self.ensurePreparedStmt(connI) + await mariadb_impl.execPreparedStmt( + self.owner.pools.conns[connI].conn, + stmt, + args, + self.owner.pools.timeout + ) + + +proc preparedGet(self: MariadbPreparedStatement, args: seq[PreparedParam]): Future[seq[JsonNode]] {.async.} = + try: + self.owner.log.logger(self.sql) + return await self.getPreparedAllRows(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc preparedFirst(self: MariadbPreparedStatement, args: seq[PreparedParam]): Future[Option[JsonNode]] {.async.} = + try: + self.owner.log.logger(self.sql) + return await self.getPreparedRow(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc preparedGetPlain(self: MariadbPreparedStatement, args: seq[PreparedParam]): Future[seq[seq[string]]] {.async.} = + try: + self.owner.log.logger(self.sql) + return await self.getPreparedAllRowsPlain(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc preparedFirstPlain(self: MariadbPreparedStatement, args: seq[PreparedParam]): Future[seq[string]] {.async.} = + try: + self.owner.log.logger(self.sql) + return await self.getPreparedRowPlain(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc preparedExec(self: MariadbPreparedStatement, args: seq[PreparedParam]) {.async.} = + try: + self.owner.log.logger(self.sql) + await self.execPrepared(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc get*(self: MariadbPreparedStatement, args: seq[string]): Future[seq[JsonNode]] {.async.} = + return await self.preparedGet(args.toPreparedParams) + + +proc get*(self: MariadbPreparedStatement, args: JsonNode): Future[seq[JsonNode]] {.async.} = + return await self.preparedGet(args.toPreparedParams) + + +proc first*(self: MariadbPreparedStatement, args: seq[string]): Future[Option[JsonNode]] {.async.} = + return await self.preparedFirst(args.toPreparedParams) + + +proc first*(self: MariadbPreparedStatement, args: JsonNode): Future[Option[JsonNode]] {.async.} = + return await self.preparedFirst(args.toPreparedParams) + + +proc getPlain*(self: MariadbPreparedStatement, args: seq[string]): Future[seq[seq[string]]] {.async.} = + return await self.preparedGetPlain(args.toPreparedParams) + + +proc getPlain*(self: MariadbPreparedStatement, args: JsonNode): Future[seq[seq[string]]] {.async.} = + return await self.preparedGetPlain(args.toPreparedParams) + + +proc firstPlain*(self: MariadbPreparedStatement, args: seq[string]): Future[seq[string]] {.async.} = + return await self.preparedFirstPlain(args.toPreparedParams) + + +proc firstPlain*(self: MariadbPreparedStatement, args: JsonNode): Future[seq[string]] {.async.} = + return await self.preparedFirstPlain(args.toPreparedParams) + + +proc exec*(self: MariadbPreparedStatement, args: seq[string]) {.async.} = + await self.preparedExec(args.toPreparedParams) + + +proc exec*(self: MariadbPreparedStatement, args: JsonNode) {.async.} = + await self.preparedExec(args.toPreparedParams) # ================================================================================ @@ -663,6 +843,17 @@ proc firstPlain*(self: RawMariadbQuery):Future[seq[string]] {.async.} = return self.getRowPlain(self.queryString, self.placeHolder).await +proc close*(self: MariadbPreparedStatement) {.async.} = + for i, stmt in self.stmts: + if stmt.isNil: + continue + try: + mariadb_impl.closePreparedStmt(stmt) + except CatchableError: + self.owner.log.echoErrorMsg("close failed for prepared stmt: " & getCurrentExceptionMsg()) + self.stmts[i] = nil + + template seeder*(rdb:MariadbConnections, tableName:string, body:untyped):untyped = ## The `seeder` block allows the code in the block to work only when the table is empty. block: diff --git a/src/allographer/query_builder/models/mariadb/mariadb_types.nim b/src/allographer/query_builder/models/mariadb/mariadb_types.nim index ec729c12..c79f5623 100644 --- a/src/allographer/query_builder/models/mariadb/mariadb_types.nim +++ b/src/allographer/query_builder/models/mariadb/mariadb_types.nim @@ -65,6 +65,25 @@ type RawMariadbQuery* = ref object transactionConn*: int +type MariadbResultBindCache* = ref object + binds*: seq[BIND] + buffers*: seq[string] + lengths*: seq[culong] + nullFlags*: seq[my_bool] + errorFlags*: seq[my_bool] + + +type MariadbPreparedStatement* = ref object + owner*: MariadbConnections + info*: ConnectionInfo + sql*: string + stmts*: seq[PSTMT] + nArgs*: int + resultBindCache*: seq[MariadbResultBindCache] + + + + proc `$`*(self:MariadbConnections|MariadbQuery|RawMariadbQuery):string = return "MariaDB" diff --git a/src/allographer/query_builder/models/mysql/mysql_exec.nim b/src/allographer/query_builder/models/mysql/mysql_exec.nim index a85960cb..3aacc972 100644 --- a/src/allographer/query_builder/models/mysql/mysql_exec.nim +++ b/src/allographer/query_builder/models/mysql/mysql_exec.nim @@ -6,8 +6,10 @@ import std/strutils import std/sequtils import std/times import ../../libs/mysql/mysql_impl +import ../../libs/mysql/mysql_rdb except Option import ../../log import ../database_types +import ../../prepared_param import ./query/mysql_builder import ./mysql_types @@ -338,7 +340,7 @@ proc transactionStart(self:MysqlConnections) {.async.} = self.isInTransaction = true self.transactionConn = connI - mysql_impl.exec(self.pools.conns[connI].conn, "BEGIN", newJArray(), newSeq[Row](), self.pools.timeout).await + mysql_impl.exec(self.pools.conns[connI].conn, "BEGIN", newJArray(), newSeq[seq[string]](), self.pools.timeout).await proc transactionEnd(self:MysqlConnections, query:string) {.async.} = @@ -347,7 +349,7 @@ proc transactionEnd(self:MysqlConnections, query:string) {.async.} = self.transactionConn = 0 self.isInTransaction = false - mysql_impl.exec(self.pools.conns[self.transactionConn].conn, query, newJArray(), newSeq[Row](), self.pools.timeout).await + mysql_impl.exec(self.pools.conns[self.transactionConn].conn, query, newJArray(), newSeq[seq[string]](), self.pools.timeout).await # ================================================================================ @@ -606,6 +608,189 @@ proc commit*(self:MysqlConnections) {.async.} = self.transactionEnd("COMMIT").await +proc prepare*(self: MysqlConnections, sql: string): MysqlPreparedStatement = + new(result) + result.owner = self + result.info = self.info + result.sql = sql + result.stmts = newSeq[PSTMT](self.pools.conns.len) + result.nArgs = countQuestionMarks(sql) + result.resultBindCache = newSeq[MysqlResultBindCache](self.pools.conns.len) + + +proc ensurePreparedStmt(self: MysqlPreparedStatement, connI: int): Future[PSTMT] {.async.} = + if self.stmts[connI].isNil: + self.stmts[connI] = await mysql_impl.prepareStmt(self.owner.pools.conns[connI].conn, self.sql, self.owner.pools.timeout) + return self.stmts[connI] + + +proc getPreparedRows(self: MysqlPreparedStatement, args: seq[PreparedParam]): Future[(seq[seq[string]], DbRows)] {.async.} = + var connI = self.owner.transactionConn + if not self.owner.isInTransaction: + connI = getFreeConn(self.owner).await + defer: + if not self.owner.isInTransaction: + self.owner.returnConn(connI).await + if connI == errorConnectionNum: + return + + let stmt = await self.ensurePreparedStmt(connI) + if connI >= self.resultBindCache.len: + self.resultBindCache.setLen(connI + 1) + if self.resultBindCache[connI].isNil: + new(self.resultBindCache[connI]) + return mysql_impl.queryPreparedStmt( + self.owner.pools.conns[connI].conn, + stmt, + args, + self.owner.pools.timeout, + self.resultBindCache[connI] + ).await + + +proc getPreparedAllRows(self: MysqlPreparedStatement, args: seq[PreparedParam]): Future[seq[JsonNode]] {.async.} = + let (rows, dbRows) = await self.getPreparedRows(args) + if rows.len == 0: + self.owner.log.echoErrorMsg(self.sql) + return newSeq[JsonNode](0) + return toJson(rows, dbRows) + + +proc getPreparedRow(self: MysqlPreparedStatement, args: seq[PreparedParam]): Future[Option[JsonNode]] {.async.} = + let (rows, dbRows) = await self.getPreparedRows(args) + if rows.len == 0: + self.owner.log.echoErrorMsg(self.sql) + return none(JsonNode) + return toJson(rows, dbRows)[0].some() + + +proc getPreparedAllRowsPlain(self: MysqlPreparedStatement, args: seq[PreparedParam]): Future[seq[seq[string]]] {.async.} = + let (rows, _) = await self.getPreparedRows(args) + return rows + + +proc getPreparedRowPlain(self: MysqlPreparedStatement, args: seq[PreparedParam]): Future[seq[string]] {.async.} = + let (rows, _) = await self.getPreparedRows(args) + if rows.len == 0: + self.owner.log.echoErrorMsg(self.sql) + return newSeq[string](0) + return rows[0] + + +proc execPrepared(self: MysqlPreparedStatement, args: seq[PreparedParam]) {.async.} = + var connI = self.owner.transactionConn + if not self.owner.isInTransaction: + connI = getFreeConn(self.owner).await + defer: + if not self.owner.isInTransaction: + self.owner.returnConn(connI).await + if connI == errorConnectionNum: + return + + let stmt = await self.ensurePreparedStmt(connI) + await mysql_impl.execPreparedStmt( + self.owner.pools.conns[connI].conn, + stmt, + args, + self.owner.pools.timeout + ) + + +proc preparedGet(self: MysqlPreparedStatement, args: seq[PreparedParam]): Future[seq[JsonNode]] {.async.} = + try: + self.owner.log.logger(self.sql) + return await self.getPreparedAllRows(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc preparedFirst(self: MysqlPreparedStatement, args: seq[PreparedParam]): Future[Option[JsonNode]] {.async.} = + try: + self.owner.log.logger(self.sql) + return await self.getPreparedRow(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc preparedGetPlain(self: MysqlPreparedStatement, args: seq[PreparedParam]): Future[seq[seq[string]]] {.async.} = + try: + self.owner.log.logger(self.sql) + return await self.getPreparedAllRowsPlain(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc preparedFirstPlain(self: MysqlPreparedStatement, args: seq[PreparedParam]): Future[seq[string]] {.async.} = + try: + self.owner.log.logger(self.sql) + return await self.getPreparedRowPlain(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc preparedExec(self: MysqlPreparedStatement, args: seq[PreparedParam]) {.async.} = + try: + self.owner.log.logger(self.sql) + await self.execPrepared(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc get*(self: MysqlPreparedStatement, args: seq[string]): Future[seq[JsonNode]] {.async.} = + return await self.preparedGet(args.toPreparedParams) + + +proc get*(self: MysqlPreparedStatement, args: JsonNode): Future[seq[JsonNode]] {.async.} = + return await self.preparedGet(args.toPreparedParams) + + +proc first*(self: MysqlPreparedStatement, args: seq[string]): Future[Option[JsonNode]] {.async.} = + return await self.preparedFirst(args.toPreparedParams) + + +proc first*(self: MysqlPreparedStatement, args: JsonNode): Future[Option[JsonNode]] {.async.} = + return await self.preparedFirst(args.toPreparedParams) + + +proc getPlain*(self: MysqlPreparedStatement, args: seq[string]): Future[seq[seq[string]]] {.async.} = + return await self.preparedGetPlain(args.toPreparedParams) + + +proc getPlain*(self: MysqlPreparedStatement, args: JsonNode): Future[seq[seq[string]]] {.async.} = + return await self.preparedGetPlain(args.toPreparedParams) + + +proc firstPlain*(self: MysqlPreparedStatement, args: seq[string]): Future[seq[string]] {.async.} = + return await self.preparedFirstPlain(args.toPreparedParams) + + +proc firstPlain*(self: MysqlPreparedStatement, args: JsonNode): Future[seq[string]] {.async.} = + return await self.preparedFirstPlain(args.toPreparedParams) + + +proc exec*(self: MysqlPreparedStatement, args: seq[string]) {.async.} = + await self.preparedExec(args.toPreparedParams) + + +proc exec*(self: MysqlPreparedStatement, args: JsonNode) {.async.} = + await self.preparedExec(args.toPreparedParams) + + +proc close*(self: MysqlPreparedStatement) {.async.} = + for stmt in self.stmts: + mysql_impl.closePreparedStmt(stmt) + + proc get*(self: RawMysqlQuery):Future[seq[JsonNode]] {.async.} = ## It is only used with raw() self.log.logger(self.queryString) diff --git a/src/allographer/query_builder/models/mysql/mysql_types.nim b/src/allographer/query_builder/models/mysql/mysql_types.nim index 6b698ab9..f4e826fa 100644 --- a/src/allographer/query_builder/models/mysql/mysql_types.nim +++ b/src/allographer/query_builder/models/mysql/mysql_types.nim @@ -60,6 +60,23 @@ type RawMysqlQuery* = ref object transactionConn*: int +type MysqlResultBindCache* = ref object + binds*: seq[BIND] + buffers*: seq[string] + lengths*: seq[culong] + nullFlags*: seq[my_bool] + errorFlags*: seq[my_bool] + + +type MysqlPreparedStatement* = ref object + owner*: MysqlConnections + info*: ConnectionInfo + sql*: string + stmts*: seq[PSTMT] + nArgs*: int + resultBindCache*: seq[MysqlResultBindCache] + + proc `$`*(self:MysqlConnections|MysqlQuery|RawMysqlQuery):string = return "MySQL" diff --git a/src/allographer/query_builder/models/postgres/postgres_exec.nim b/src/allographer/query_builder/models/postgres/postgres_exec.nim index 66f833c5..70ed803c 100644 --- a/src/allographer/query_builder/models/postgres/postgres_exec.nim +++ b/src/allographer/query_builder/models/postgres/postgres_exec.nim @@ -1,5 +1,6 @@ import std/asyncdispatch import std/deques +import std/atomics import std/json import std/monotimes import std/options @@ -12,9 +13,12 @@ import ../../libs/postgres/postgres_lib import ../../libs/postgres/postgres_impl import ../../log import ../database_types +import ../../prepared_param import ./query/postgres_builder import ./postgres_types +var gPreparedStmtCounter: Atomic[int] + # ================================================================================ # connection @@ -74,6 +78,29 @@ proc returnConn(self: PostgresConnections | PostgresQuery | RawPostgresQuery, i: wakeOnePoolWaiter(self.pools) +proc prepare*(self: PostgresConnections, sql: string): PostgresPreparedStatement = + new(result) + result.owner = self + result.sql = sql + result.stmtBaseName = &"allographer_stmt_{gPreparedStmtCounter.fetchAdd(1)}" + result.stmtNames = newSeq[string](self.pools.conns.len) + result.nArgs = countQuestionMarks(sql) + + +proc ensurePreparedStmt(self: PostgresPreparedStatement, connI: int): Future[string] {.async.} = + if self.stmtNames[connI].len == 0: + let stmtName = &"{self.stmtBaseName}_{connI}" + await postgres_impl.prepare( + self.owner.pools.conns[connI].conn, + self.sql, + self.owner.pools.timeout, + stmtName, + self.nArgs + ) + self.stmtNames[connI] = stmtName + return self.stmtNames[connI] + + # ================================================================================ # toJson # ================================================================================ @@ -401,16 +428,175 @@ proc transactionStart(self:PostgresConnections|PostgresQuery) {.async.} = self.isInTransaction = true self.transactionConn = connI - postgres_impl.exec(self.pools.conns[connI].conn, "BEGIN", newJArray(), newSeq[Row](), self.pools.timeout).await + postgres_impl.exec(self.pools.conns[connI].conn, "BEGIN", newJArray(), newSeq[seq[string]](), self.pools.timeout).await proc transactionEnd(self:PostgresConnections|PostgresQuery, query:string) {.async.} = - postgres_impl.exec(self.pools.conns[self.transactionConn].conn, query, newJArray(), newSeq[Row](), self.pools.timeout).await + postgres_impl.exec(self.pools.conns[self.transactionConn].conn, query, newJArray(), newSeq[seq[string]](), self.pools.timeout).await self.returnConn(self.transactionConn).await self.transactionConn = 0 self.isInTransaction = false +proc getPreparedRows(self: PostgresPreparedStatement, args: seq[PreparedParam]): Future[(seq[seq[string]], DbRows)] {.async.} = + var connI = self.owner.transactionConn + if not self.owner.isInTransaction: + connI = getFreeConn(self.owner).await + defer: + if not self.owner.isInTransaction: + self.owner.returnConn(connI).await + if connI == errorConnectionNum: + return + + let stmtName = await self.ensurePreparedStmt(connI) + return postgres_impl.preparedQuery( + self.owner.pools.conns[connI].conn, + args, + self.nArgs, + self.owner.pools.timeout, + stmtName + ).await + + +proc getPreparedAllRows(self: PostgresPreparedStatement, args: seq[PreparedParam]): Future[seq[JsonNode]] {.async.} = + let (rows, dbRows) = await self.getPreparedRows(args) + if rows.len == 0: + self.owner.log.echoErrorMsg(self.sql) + return newSeq[JsonNode](0) + return toJson(rows, dbRows) + + +proc getPreparedRow(self: PostgresPreparedStatement, args: seq[PreparedParam]): Future[Option[JsonNode]] {.async.} = + let (rows, dbRows) = await self.getPreparedRows(args) + if rows.len == 0: + self.owner.log.echoErrorMsg(self.sql) + return none(JsonNode) + return toJson(rows, dbRows)[0].some() + + +proc getPreparedAllRowsPlain(self: PostgresPreparedStatement, args: seq[PreparedParam]): Future[seq[seq[string]]] {.async.} = + let (rows, _) = await self.getPreparedRows(args) + return rows + + +proc getPreparedRowPlain(self: PostgresPreparedStatement, args: seq[PreparedParam]): Future[seq[string]] {.async.} = + let (rows, _) = await self.getPreparedRows(args) + if rows.len == 0: + self.owner.log.echoErrorMsg(self.sql) + return newSeq[string](0) + return rows[0] + + +proc execPrepared(self: PostgresPreparedStatement, args: seq[PreparedParam]) {.async.} = + var connI = self.owner.transactionConn + if not self.owner.isInTransaction: + connI = getFreeConn(self.owner).await + defer: + if not self.owner.isInTransaction: + self.owner.returnConn(connI).await + if connI == errorConnectionNum: + return + + let stmtName = await self.ensurePreparedStmt(connI) + await postgres_impl.preparedExec( + self.owner.pools.conns[connI].conn, + args, + self.nArgs, + self.owner.pools.timeout, + stmtName + ) + + +proc preparedGet(self: PostgresPreparedStatement, args: seq[PreparedParam]): Future[seq[JsonNode]] {.async.} = + try: + self.owner.log.logger(self.sql) + return await self.getPreparedAllRows(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc preparedFirst(self: PostgresPreparedStatement, args: seq[PreparedParam]): Future[Option[JsonNode]] {.async.} = + try: + self.owner.log.logger(self.sql) + return await self.getPreparedRow(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc preparedGetPlain(self: PostgresPreparedStatement, args: seq[PreparedParam]): Future[seq[seq[string]]] {.async.} = + try: + self.owner.log.logger(self.sql) + return await self.getPreparedAllRowsPlain(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc preparedFirstPlain(self: PostgresPreparedStatement, args: seq[PreparedParam]): Future[seq[string]] {.async.} = + try: + self.owner.log.logger(self.sql) + return await self.getPreparedRowPlain(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc preparedExec(self: PostgresPreparedStatement, args: seq[PreparedParam]) {.async.} = + try: + self.owner.log.logger(self.sql) + await self.execPrepared(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc get*(self: PostgresPreparedStatement, args: seq[string]): Future[seq[JsonNode]] {.async.} = + return await self.preparedGet(args.toPreparedParams) + + +proc get*(self: PostgresPreparedStatement, args: JsonNode): Future[seq[JsonNode]] {.async.} = + return await self.preparedGet(args.toPreparedParams) + + +proc first*(self: PostgresPreparedStatement, args: seq[string]): Future[Option[JsonNode]] {.async.} = + return await self.preparedFirst(args.toPreparedParams) + + +proc first*(self: PostgresPreparedStatement, args: JsonNode): Future[Option[JsonNode]] {.async.} = + return await self.preparedFirst(args.toPreparedParams) + + +proc getPlain*(self: PostgresPreparedStatement, args: seq[string]): Future[seq[seq[string]]] {.async.} = + return await self.preparedGetPlain(args.toPreparedParams) + + +proc getPlain*(self: PostgresPreparedStatement, args: JsonNode): Future[seq[seq[string]]] {.async.} = + return await self.preparedGetPlain(args.toPreparedParams) + + +proc firstPlain*(self: PostgresPreparedStatement, args: seq[string]): Future[seq[string]] {.async.} = + return await self.preparedFirstPlain(args.toPreparedParams) + + +proc firstPlain*(self: PostgresPreparedStatement, args: JsonNode): Future[seq[string]] {.async.} = + return await self.preparedFirstPlain(args.toPreparedParams) + + +proc exec*(self: PostgresPreparedStatement, args: seq[string]) {.async.} = + await self.preparedExec(args.toPreparedParams) + + +proc exec*(self: PostgresPreparedStatement, args: JsonNode) {.async.} = + await self.preparedExec(args.toPreparedParams) + + # ================================================================================ # public exec # ================================================================================ @@ -725,6 +911,25 @@ proc firstPlain*(self: RawPostgresQuery):Future[seq[string]] {.async.} = return self.getRowPlain(self.queryString, self.placeHolder).await +proc deallocatePreparedStmtSafely(self: PostgresPreparedStatement, connI: int, stmtName: string): Future[void] {.async.} = + try: + await postgres_impl.deallocate(self.owner.pools.conns[connI].conn, stmtName, self.owner.pools.timeout) + except CatchableError: + self.owner.log.echoErrorMsg("deallocate failed for " & stmtName & ": " & getCurrentExceptionMsg()) + + +proc close*(self: PostgresPreparedStatement) {.async.} = + var futs: seq[Future[void]] + for i, stmtName in self.stmtNames: + if stmtName.len == 0: + continue + futs.add(self.deallocatePreparedStmtSafely(i, stmtName)) + if futs.len > 0: + await all(futs) + for i in 0 ..< self.stmtNames.len: + self.stmtNames[i] = "" + + template seeder*(rdb:PostgresConnections, tableName:string, body:untyped):untyped = ## The `seeder` block allows the code in the block to work only when the table is empty. block: diff --git a/src/allographer/query_builder/models/postgres/postgres_types.nim b/src/allographer/query_builder/models/postgres/postgres_types.nim index 9d668d9c..f7024aec 100644 --- a/src/allographer/query_builder/models/postgres/postgres_types.nim +++ b/src/allographer/query_builder/models/postgres/postgres_types.nim @@ -57,6 +57,14 @@ type RawPostgresQuery* = ref object transactionConn*: int +type PostgresPreparedStatement* = ref object + owner*: PostgresConnections + sql*: string + stmtBaseName*: string + stmtNames*: seq[string] + nArgs*: int + + proc `$`*(self:PostgresConnections|PostgresQuery|RawPostgresQuery):string = return "PostgreSQL" diff --git a/src/allographer/query_builder/models/sqlite/sqlite_exec.nim b/src/allographer/query_builder/models/sqlite/sqlite_exec.nim index d658aa4e..b7e6a999 100644 --- a/src/allographer/query_builder/models/sqlite/sqlite_exec.nim +++ b/src/allographer/query_builder/models/sqlite/sqlite_exec.nim @@ -9,8 +9,10 @@ import std/tables import std/times import ../../libs/sqlite/sqlite_impl import ../../libs/sqlite/sqlite_lib +import ../../libs/sqlite/sqlite_rdb import ../../log import ../database_types +import ../../prepared_param import ./query/sqlite_builder import ./sqlite_types @@ -73,6 +75,22 @@ proc returnConn(self: SqliteConnections | SqliteQuery | RawSqliteQuery, i: int) wakeOnePoolWaiter(self.pools) +proc prepare*(self: SqliteConnections, sql: string): SqlitePreparedStatement = + SqlitePreparedStatement( + owner: self, + sql: sql, + stmts: newSeq[PStmt](self.pools.conns.len), + nArgs: countQuestionMarks(sql) + ) + + +proc ensurePreparedStmt(self: SqlitePreparedStatement, connI: int): Future[PStmt] {.async.} = + if self.stmts[connI].isNil: + self.stmts[connI] = sqlite_impl.prepare(self.owner.pools.conns[connI].conn, self.sql, self.owner.pools.timeout).await + return self.stmts[connI] + + + # ================================================================================ # toJson # ================================================================================ @@ -509,6 +527,128 @@ proc transactionEnd(self:SqliteConnections, query:string) {.async.} = sqlite_impl.exec(self.pools.conns[self.transactionConn].conn, query, newJArray(), self.pools.timeout).await +proc getPreparedRows(self: SqlitePreparedStatement, args: seq[PreparedParam]): Future[(seq[seq[string]], DbRows)] {.async.} = + var connI = self.owner.transactionConn + if not self.owner.isInTransaction: + connI = getFreeConn(self.owner).await + defer: + if not self.owner.isInTransaction: + self.owner.returnConn(connI).await + if connI == errorConnectionNum: + return + + let stmt = await self.ensurePreparedStmt(connI) + if not self.hasCachedColumns: + setColumnsStaticMeta(self.cachedColumns, stmt) + self.hasCachedColumns = true + + return sqlite_impl.preparedQueryReuse( + self.owner.pools.conns[connI].conn, + stmt, + args, + self.owner.pools.timeout, + self.cachedColumns + ).await + + +proc getPreparedAllRows(self: SqlitePreparedStatement, args: seq[PreparedParam]): Future[seq[JsonNode]] {.async.} = + let (rows, dbRows) = await self.getPreparedRows(args) + if rows.len == 0: + self.owner.log.echoErrorMsg(self.sql) + return newSeq[JsonNode](0) + return toJson(rows, dbRows) + + +proc getPreparedRow(self: SqlitePreparedStatement, args: seq[PreparedParam]): Future[Option[JsonNode]] {.async.} = + let (rows, dbRows) = await self.getPreparedRows(args) + if rows.len == 0: + self.owner.log.echoErrorMsg(self.sql) + return none(JsonNode) + return toJson(rows, dbRows)[0].some() + + +proc getPreparedAllRowsPlain(self: SqlitePreparedStatement, args: seq[PreparedParam]): Future[seq[seq[string]]] {.async.} = + let (rows, _) = await self.getPreparedRows(args) + return rows + + +proc getPreparedRowPlain(self: SqlitePreparedStatement, args: seq[PreparedParam]): Future[seq[string]] {.async.} = + let (rows, _) = await self.getPreparedRows(args) + if rows.len == 0: + self.owner.log.echoErrorMsg(self.sql) + return newSeq[string](0) + return rows[0] + + +proc execPrepared(self: SqlitePreparedStatement, args: seq[PreparedParam]) {.async.} = + var connI = self.owner.transactionConn + if not self.owner.isInTransaction: + connI = getFreeConn(self.owner).await + defer: + if not self.owner.isInTransaction: + self.owner.returnConn(connI).await + if connI == errorConnectionNum: + return + + let stmt = await self.ensurePreparedStmt(connI) + await sqlite_impl.preparedExecReuse( + self.owner.pools.conns[connI].conn, + stmt, + args, + self.owner.pools.timeout + ) + + +proc preparedGet(self: SqlitePreparedStatement, args: seq[PreparedParam]): Future[seq[JsonNode]] {.async.} = + try: + self.owner.log.logger(self.sql) + return await self.getPreparedAllRows(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc preparedFirst(self: SqlitePreparedStatement, args: seq[PreparedParam]): Future[Option[JsonNode]] {.async.} = + try: + self.owner.log.logger(self.sql) + return await self.getPreparedRow(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc preparedGetPlain(self: SqlitePreparedStatement, args: seq[PreparedParam]): Future[seq[seq[string]]] {.async.} = + try: + self.owner.log.logger(self.sql) + return await self.getPreparedAllRowsPlain(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc preparedFirstPlain(self: SqlitePreparedStatement, args: seq[PreparedParam]): Future[seq[string]] {.async.} = + try: + self.owner.log.logger(self.sql) + return await self.getPreparedRowPlain(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + +proc preparedExec(self: SqlitePreparedStatement, args: seq[PreparedParam]) {.async.} = + try: + self.owner.log.logger(self.sql) + await self.execPrepared(args) + except CatchableError: + self.owner.log.echoErrorMsg(self.sql) + self.owner.log.echoErrorMsg(getCurrentExceptionMsg()) + raise getCurrentException() + + # ================================================================================ # public exec # ================================================================================ @@ -591,55 +731,6 @@ proc findPlain*(self:SqliteQuery, id: int, key="id"):Future[seq[string]] {.async return self.findPlain($id, key).await -# ==================== return Object ==================== -# proc get*[T](self: SqliteQuery, typ:typedesc[T]):Future[seq[T]] {.async.} = -# var sql = self.selectBuilder() -# try: -# self.log.logger(sql) -# let rows = self.getAllRows(sql).await -# for row in rows: -# result.add(row.to(typ)) -# except CatchableError: -# self.log.echoErrorMsg(sql) -# self.log.echoErrorMsg( getCurrentExceptionMsg() ) -# raise getCurrentException() - - -# proc first*[T](self: SqliteQuery, typ:typedesc[T]):Future[Option[T]] {.async.} = -# var sql = self.selectFirstBuilder() -# try: -# self.log.logger(sql) -# let row = self.getRow(sql).await -# if row.isSome(): -# return row.get().to(typ).some() -# else: -# return none(typ) -# except CatchableError: -# self.log.echoErrorMsg(sql) -# self.log.echoErrorMsg( getCurrentExceptionMsg() ) -# raise getCurrentException() - - -# proc find*[T](self: SqliteQuery, id:string, typ:typedesc[T], key="id"):Future[Option[T]] {.async.} = -# self.placeHolder.add(%*{"key":key, "value": id}) -# var sql = self.selectFindBuilder(key) -# try: -# self.log.logger(sql) -# let row = self.getRow(sql).await -# if row.isSome(): -# return row.get().to(typ).some() -# else: -# return none(typ) -# except CatchableError: -# self.log.echoErrorMsg(sql) -# self.log.echoErrorMsg( getCurrentExceptionMsg() ) -# raise getCurrentException() - - -# proc find*[T](self: SqliteQuery, id:int, typ:typedesc[T], key="id"):Future[Option[T]] {.async.} = -# return self.find($id, typ, key).await - - # ==================== insert JsonNode ==================== proc insert*(self:SqliteQuery, items:JsonNode) {.async.} = let sql = self.insertValueBuilder(items) @@ -836,6 +927,62 @@ proc firstPlain*(self: RawSqliteQuery):Future[seq[string]] {.async.} = return self.getRowPlain(self.queryString, self.placeHolder).await +proc close*(self: SqlitePreparedStatement) {.async.} = + for i, stmt in self.stmts: + if stmt.isNil: + continue + try: + discard finalize(stmt) + except CatchableError: + self.owner.log.echoErrorMsg("finalize failed for prepared stmt: " & getCurrentExceptionMsg()) + self.stmts[i] = nil + + +# ================================================================================ +# public prepared exec +# ================================================================================ + +proc get*(self: SqlitePreparedStatement, args: seq[string]): Future[seq[JsonNode]] {.async.} = + return await self.preparedGet(args.toPreparedParams) + + +proc get*(self: SqlitePreparedStatement, args: JsonNode): Future[seq[JsonNode]] {.async.} = + return await self.preparedGet(args.toPreparedParams) + + +proc first*(self: SqlitePreparedStatement, args: seq[string]): Future[Option[JsonNode]] {.async.} = + return await self.preparedFirst(args.toPreparedParams) + + +proc first*(self: SqlitePreparedStatement, args: JsonNode): Future[Option[JsonNode]] {.async.} = + return await self.preparedFirst(args.toPreparedParams) + + +proc getPlain*(self: SqlitePreparedStatement, args: seq[string]): Future[seq[seq[string]]] {.async.} = + return await self.preparedGetPlain(args.toPreparedParams) + + +proc getPlain*(self: SqlitePreparedStatement, args: JsonNode): Future[seq[seq[string]]] {.async.} = + return await self.preparedGetPlain(args.toPreparedParams) + + +proc firstPlain*(self: SqlitePreparedStatement, args: seq[string]): Future[seq[string]] {.async.} = + return await self.preparedFirstPlain(args.toPreparedParams) + + +proc firstPlain*(self: SqlitePreparedStatement, args: JsonNode): Future[seq[string]] {.async.} = + return await self.preparedFirstPlain(args.toPreparedParams) + + +proc exec*(self: SqlitePreparedStatement, args: seq[string]) {.async.} = + await self.preparedExec(args.toPreparedParams) + + +proc exec*(self: SqlitePreparedStatement, args: JsonNode) {.async.} = + await self.preparedExec(args.toPreparedParams) + + + template seeder*(rdb:SqliteConnections, tableName:string, body:untyped):untyped = ## The `seeder` block allows the code in the block to work only when the table is empty. block: diff --git a/src/allographer/query_builder/models/sqlite/sqlite_types.nim b/src/allographer/query_builder/models/sqlite/sqlite_types.nim index 214cea69..91a30b24 100644 --- a/src/allographer/query_builder/models/sqlite/sqlite_types.nim +++ b/src/allographer/query_builder/models/sqlite/sqlite_types.nim @@ -2,6 +2,7 @@ import std/asyncdispatch import std/deques import std/json import std/tables +import ../database_types import ../../log import ../../libs/sqlite/sqlite_rdb @@ -56,5 +57,15 @@ type RawSqliteQuery* = ref object transactionConn*: int +type SqlitePreparedStatement* = ref object + owner*: SqliteConnections + sql*: string + stmts*: seq[PStmt] + nArgs*: int + cachedColumns*: DbColumns + hasCachedColumns*: bool + + + proc `$`*(self:SqliteConnections|SqliteQuery|RawSqliteQuery):string = return "SQLite" diff --git a/src/allographer/query_builder/prepared_param.nim b/src/allographer/query_builder/prepared_param.nim new file mode 100644 index 00000000..452bb278 --- /dev/null +++ b/src/allographer/query_builder/prepared_param.nim @@ -0,0 +1,115 @@ +import std/json +import std/options +import std/macros + + +type PreparedParam* = object + value*: string + isNull*: bool + + +proc nullPreparedParam*(): PreparedParam = + result.isNull = true + + +proc toPreparedParam*(v: string): PreparedParam = + result.value = v + + +proc toPreparedParam*(v: cstring): PreparedParam = + if v.isNil: + return nullPreparedParam() + result.value = $v + + +proc toPreparedParam*(v: bool): PreparedParam = + result.value = if v: "1" else: "0" + + +proc toPreparedParam*[T: SomeInteger](v: T): PreparedParam = + result.value = $v + + +proc toPreparedParam*[T: SomeFloat](v: T): PreparedParam = + result.value = $v + + +proc toPreparedParam*(v: JsonNode): PreparedParam = + if v.isNil or v.kind == JNull: + return nullPreparedParam() + + case v.kind + of JBool: + result.value = if v.getBool: "1" else: "0" + of JInt: + result.value = $v.getInt + of JFloat: + result.value = $v.getFloat + of JString: + result.value = v.getStr + of JArray, JObject: + result.value = v.pretty + of JNull: + discard + + +proc toPreparedParam*[T](v: Option[T]): PreparedParam = + if v.isSome: + return toPreparedParam(v.get) + return nullPreparedParam() + + +proc toPreparedParam*[T](v: T): PreparedParam = + result.value = $v + + +proc toPreparedParams*(args: seq[string]): seq[PreparedParam] = + result = newSeq[PreparedParam](args.len) + for i, arg in args: + if arg == "NULL" or arg == "null": + result[i] = nullPreparedParam() + else: + result[i] = toPreparedParam(arg) + + +proc toPreparedParams*(args: JsonNode): seq[PreparedParam] = + if args.isNil or args.kind == JNull: + return + + if args.kind == JArray: + result = newSeq[PreparedParam](args.len) + for i in 0 ..< args.len: + result[i] = toPreparedParam(args[i]) + return + + result = @[toPreparedParam(args)] + + +proc preparedText*(param: PreparedParam): string = + if param.isNull: + return "NULL" + return param.value + + +proc preparedTextSeq*(args: openArray[PreparedParam]): seq[string] = + result = newSeq[string](args.len) + for i, arg in args: + result[i] = arg.preparedText + + +proc countQuestionMarks*(s: string): int = + for ch in s: + if ch == '?': + inc result + + +proc buildPreparedArgsExpr*(args: NimNode): NimNode = + let toPreparedParamSym = bindSym("toPreparedParam") + let nullPreparedParamSym = bindSym("nullPreparedParam") + var arr = nnkBracket.newTree() + for arg in args: + if arg.kind == nnkNilLit: + arr.add(newCall(nullPreparedParamSym)) + else: + arr.add(newCall(toPreparedParamSym, arg)) + result = newTree(nnkPrefix, ident("@"), arr) diff --git a/tests/config.nims b/tests/config.nims index ef7ea0c4..66addc57 100644 --- a/tests/config.nims +++ b/tests/config.nims @@ -1,8 +1,8 @@ import os switch("path", "$projectDir/../src") -# putEnv("DB_SQLITE", $true) -# putEnv("DB_POSTGRES", $true) -# putEnv("DB_MYSQL", $true) -# putEnv("DB_MARIADB", $true) -# putEnv("DB_SURREAL", $true) +putEnv("DB_SQLITE", $true) +putEnv("DB_POSTGRES", $true) +putEnv("DB_MYSQL", $true) +putEnv("DB_MARIADB", $true) +putEnv("DB_SURREAL", $true) diff --git a/tests/mariadb/test_prepared_statement.nim b/tests/mariadb/test_prepared_statement.nim new file mode 100644 index 00000000..c49edb93 --- /dev/null +++ b/tests/mariadb/test_prepared_statement.nim @@ -0,0 +1,89 @@ +discard """ + cmd: "nim c -d:reset -d:ssl -r $file" +""" + +import std/unittest +import std/asyncdispatch +import std/json +import std/options +import std/strformat +import ../../src/allographer/schema_builder +import ../../src/allographer/query_builder +import ./connections + + +let rdb = mariadb + + +proc setup(rdb: MariadbConnections) = + rdb.create([ + table("auth", [ + Column.increments("id"), + Column.string("auth") + ]), + table("user", [ + Column.increments("id"), + Column.string("name").nullable(), + Column.string("email").nullable(), + Column.string("address").nullable(), + Column.date("submit_on").nullable(), + Column.datetime("submit_at").nullable(), + Column.foreign("auth_id").reference("id").onTable("auth").onDelete(SET_NULL).nullable() + ]) + ]) + + seeder(rdb, "auth"): + rdb.table("auth").insert(@[ + %*{"auth": "admin"}, + %*{"auth": "user"} + ]).waitFor + + seeder(rdb, "user"): + var users: seq[JsonNode] + for i in 1..10: + let authId = if i mod 2 == 0: 2 else: 1 + let month = if i > 9: $i else: &"0{i}" + users.add( + %*{ + "name": &"user{i}", + "email": &"user{i}@example.com", + "auth_id": authId, + "submit_on": &"2020-{month}-01", + "submit_at": &"2020-{month}-01 00:00:00", + } + ) + + rdb.table("user").insert(users).waitFor + + +setup(rdb) + + +suite($rdb & " prepared statement"): + test("select"): + let stmt = rdb.prepare("""SELECT `id`, `name`, `email`, `address` FROM `user` WHERE `id` = ?""") + defer: + waitFor stmt.close() + + let args = newJArray() + args.add(newJInt(1)) + + let rows = stmt.get(args).waitFor + check rows.len == 1 + check rows[0] == %*{"id": 1, "name": "user1", "email": "user1@example.com", "address": newJNull()} + let rowOpt = stmt.first(args).waitFor + let row = options.get(rowOpt) + check row["name"].getStr == "user1" + check stmt.getPlain(args).waitFor[0][1] == "user1" + check stmt.firstPlain(args).waitFor[1] == "user1" + + + test("update null"): + let stmt = rdb.prepare("""UPDATE `user` SET `address` = ? WHERE `id` = ?""") + defer: + waitFor stmt.close() + + waitFor stmt.exec(@["NULL", "1"]) + let rowOpt = rdb.table("user").find(1).waitFor + let row = options.get(rowOpt) + check row["address"].kind == JNull diff --git a/tests/mysql/test_prepared_statement.nim b/tests/mysql/test_prepared_statement.nim new file mode 100644 index 00000000..4a6b424e --- /dev/null +++ b/tests/mysql/test_prepared_statement.nim @@ -0,0 +1,89 @@ +discard """ + cmd: "nim c -d:reset -d:ssl -r $file" +""" + +import std/unittest +import std/asyncdispatch +import std/json +import std/options +import std/strformat +import ../../src/allographer/schema_builder +import ../../src/allographer/query_builder +import ./connections + + +let rdb = mysql + + +proc setup(rdb: MysqlConnections) = + rdb.create([ + table("auth", [ + Column.increments("id"), + Column.string("auth") + ]), + table("user", [ + Column.increments("id"), + Column.string("name").nullable(), + Column.string("email").nullable(), + Column.string("address").nullable(), + Column.date("submit_on").nullable(), + Column.datetime("submit_at").nullable(), + Column.foreign("auth_id").reference("id").onTable("auth").onDelete(SET_NULL).nullable() + ]) + ]) + + seeder(rdb, "auth"): + rdb.table("auth").insert(@[ + %*{"auth": "admin"}, + %*{"auth": "user"} + ]).waitFor + + seeder(rdb, "user"): + var users: seq[JsonNode] + for i in 1..10: + let authId = if i mod 2 == 0: 2 else: 1 + let month = if i > 9: $i else: &"0{i}" + users.add( + %*{ + "name": &"user{i}", + "email": &"user{i}@example.com", + "auth_id": authId, + "submit_on": &"2020-{month}-01", + "submit_at": &"2020-{month}-01 00:00:00", + } + ) + + rdb.table("user").insert(users).waitFor + + +setup(rdb) + + +suite($rdb & " prepared statement"): + test("select"): + let stmt = rdb.prepare("""SELECT `id`, `name`, `email`, `address` FROM `user` WHERE `id` = ?""") + defer: + waitFor stmt.close() + + let args = newJArray() + args.add(newJInt(1)) + + let rows = stmt.get(args).waitFor + check rows.len == 1 + check rows[0] == %*{"id": 1, "name": "user1", "email": "user1@example.com", "address": newJNull()} + let rowOpt = stmt.first(args).waitFor + let row = options.get(rowOpt) + check row["name"].getStr == "user1" + check stmt.getPlain(args).waitFor[0][1] == "user1" + check stmt.firstPlain(args).waitFor[1] == "user1" + + + test("update null"): + let stmt = rdb.prepare("""UPDATE `user` SET `address` = ? WHERE `id` = ?""") + defer: + waitFor stmt.close() + + waitFor stmt.exec(@["NULL", "1"]) + let rowOpt = rdb.table("user").find(1).waitFor + let row = options.get(rowOpt) + check row["address"].kind == JNull diff --git a/tests/postgres/test_prepared_statement.nim b/tests/postgres/test_prepared_statement.nim new file mode 100644 index 00000000..81926baa --- /dev/null +++ b/tests/postgres/test_prepared_statement.nim @@ -0,0 +1,89 @@ +discard """ + cmd: "nim c -d:reset -d:ssl -r $file" +""" + +import std/unittest +import std/asyncdispatch +import std/json +import std/options +import std/strformat +import ../../src/allographer/schema_builder +import ../../src/allographer/query_builder +import ./connections + + +let rdb = postgres + + +proc setup(rdb: PostgresConnections) = + rdb.create([ + table("auth", [ + Column.increments("id"), + Column.string("auth") + ]), + table("user", [ + Column.increments("id"), + Column.string("name").nullable(), + Column.string("email").nullable(), + Column.string("address").nullable(), + Column.date("submit_on").nullable(), + Column.datetime("submit_at").nullable(), + Column.foreign("auth_id").reference("id").onTable("auth").onDelete(SET_NULL).nullable() + ]) + ]) + + seeder(rdb, "auth"): + rdb.table("auth").insert(@[ + %*{"auth": "admin"}, + %*{"auth": "user"} + ]).waitFor + + seeder(rdb, "user"): + var users: seq[JsonNode] + for i in 1..10: + let authId = if i mod 2 == 0: 2 else: 1 + let month = if i > 9: $i else: &"0{i}" + users.add( + %*{ + "name": &"user{i}", + "email": &"user{i}@example.com", + "auth_id": authId, + "submit_on": &"2020-{month}-01", + "submit_at": &"2020-{month}-01 00:00:00", + } + ) + + rdb.table("user").insert(users).waitFor + + +setup(rdb) + + +suite($rdb & " prepared statement"): + test("select"): + let stmt = rdb.prepare("""SELECT "id", "name", "email", "address" FROM "user" WHERE "id" = ?""") + defer: + waitFor stmt.close() + + let args = newJArray() + args.add(newJInt(1)) + + let rows = stmt.get(args).waitFor + check rows.len == 1 + check rows[0] == %*{"id": 1, "name": "user1", "email": "user1@example.com", "address": newJNull()} + let rowOpt = stmt.first(args).waitFor + let row = options.get(rowOpt) + check row["name"].getStr == "user1" + check stmt.getPlain(args).waitFor[0][1] == "user1" + check stmt.firstPlain(args).waitFor[1] == "user1" + + + test("update null"): + let stmt = rdb.prepare("""UPDATE "user" SET "address" = ? WHERE "id" = ?""") + defer: + waitFor stmt.close() + + waitFor stmt.exec(@["NULL", "1"]) + let rowOpt = rdb.table("user").find(1).waitFor + let row = options.get(rowOpt) + check row["address"].kind == JNull diff --git a/tests/sqlite/test_prepared_statement.nim b/tests/sqlite/test_prepared_statement.nim new file mode 100644 index 00000000..b7d65c84 --- /dev/null +++ b/tests/sqlite/test_prepared_statement.nim @@ -0,0 +1,89 @@ +discard """ + cmd: "nim c -d:reset -d:ssl -r $file" +""" + +import std/unittest +import std/asyncdispatch +import std/json +import std/options +import std/strformat +import ../../src/allographer/schema_builder +import ../../src/allographer/query_builder +import ./connections + + +let rdb = sqlite + + +proc setup(rdb: SqliteConnections) = + rdb.create([ + table("auth", [ + Column.increments("id"), + Column.string("auth") + ]), + table("user", [ + Column.increments("id"), + Column.string("name").nullable(), + Column.string("email").nullable(), + Column.string("address").nullable(), + Column.date("submit_on").nullable(), + Column.datetime("submit_at").nullable(), + Column.foreign("auth_id").reference("id").onTable("auth").onDelete(SET_NULL).nullable() + ]) + ]) + + seeder(rdb, "auth"): + rdb.table("auth").insert(@[ + %*{"auth": "admin"}, + %*{"auth": "user"} + ]).waitFor + + seeder(rdb, "user"): + var users: seq[JsonNode] + for i in 1..10: + let authId = if i mod 2 == 0: 2 else: 1 + let month = if i > 9: $i else: &"0{i}" + users.add( + %*{ + "name": &"user{i}", + "email": &"user{i}@example.com", + "auth_id": authId, + "submit_on": &"2020-{month}-01", + "submit_at": &"2020-{month}-01 00:00:00", + } + ) + + rdb.table("user").insert(users).waitFor + + +setup(rdb) + + +suite($rdb & " prepared statement"): + test("select"): + let stmt = rdb.prepare("""SELECT "id", "name", "email", "address" FROM "user" WHERE "id" = ?""") + defer: + waitFor stmt.close() + + let args = newJArray() + args.add(newJInt(1)) + + let rows = stmt.get(args).waitFor + check rows.len == 1 + check rows[0] == %*{"id": 1, "name": "user1", "email": "user1@example.com", "address": newJNull()} + let rowOpt = stmt.first(args).waitFor + let row = options.get(rowOpt) + check row["name"].getStr == "user1" + check stmt.getPlain(args).waitFor[0][1] == "user1" + check stmt.firstPlain(args).waitFor[1] == "user1" + + + test("update null"): + let stmt = rdb.prepare("""UPDATE "user" SET "address" = ? WHERE "id" = ?""") + defer: + waitFor stmt.close() + + waitFor stmt.exec(@["NULL", "1"]) + let rowOpt = rdb.table("user").find(1).waitFor + let row = options.get(rowOpt) + check row["address"].kind == JNull