Query.utils.ts457 lines · main
| 1 | import { |
| 2 | format, |
| 3 | ident, |
| 4 | joinSqlFragments, |
| 5 | literal, |
| 6 | safeSql, |
| 7 | type SafeSqlFragment, |
| 8 | } from '../pg-format' |
| 9 | import type { Dictionary, Filter, QueryPagination, QueryTable, Sort } from './types' |
| 10 | |
| 11 | export function countQuery( |
| 12 | table: QueryTable, |
| 13 | options?: { |
| 14 | filters?: Filter[] |
| 15 | } |
| 16 | ) { |
| 17 | let query = safeSql`select count(*) from ${queryTable(table)}` |
| 18 | const { filters } = options ?? {} |
| 19 | if (filters) { |
| 20 | query = applyFilters(query, filters) |
| 21 | } |
| 22 | return safeSql`${query};` |
| 23 | } |
| 24 | |
| 25 | export function truncateQuery( |
| 26 | table: QueryTable, |
| 27 | options?: { |
| 28 | // [Joshen] yet to implement cascade from UI, just adding first |
| 29 | cascade?: boolean |
| 30 | } |
| 31 | ) { |
| 32 | let query = safeSql`truncate ${queryTable(table)}` |
| 33 | const { cascade } = options ?? {} |
| 34 | if (cascade) { |
| 35 | query = safeSql`${query} cascade` |
| 36 | } |
| 37 | return safeSql`${query};` |
| 38 | } |
| 39 | |
| 40 | export function deleteQuery( |
| 41 | table: QueryTable, |
| 42 | filters?: Filter[], |
| 43 | options?: { |
| 44 | returning?: boolean |
| 45 | enumArrayColumns?: string[] |
| 46 | } |
| 47 | ) { |
| 48 | if (!filters || filters.length === 0) { |
| 49 | throw new Error('no filters for this delete query') |
| 50 | } |
| 51 | let query = safeSql`delete from ${queryTable(table)}` |
| 52 | const { returning, enumArrayColumns } = options ?? {} |
| 53 | if (filters) { |
| 54 | query = applyFilters(query, filters) |
| 55 | } |
| 56 | if (returning) { |
| 57 | const returningFragment = |
| 58 | enumArrayColumns === undefined || enumArrayColumns.length === 0 |
| 59 | ? safeSql` returning *` |
| 60 | : safeSql` returning *, ${joinSqlFragments( |
| 61 | enumArrayColumns.map((x) => safeSql`${ident(x)}::text[]`), |
| 62 | ',' |
| 63 | )}` |
| 64 | query = safeSql`${query}${returningFragment}` |
| 65 | } |
| 66 | return safeSql`${query};` |
| 67 | } |
| 68 | |
| 69 | export function insertQuery( |
| 70 | table: QueryTable, |
| 71 | values: Dictionary<any>[], |
| 72 | options?: { |
| 73 | returning?: boolean |
| 74 | enumArrayColumns?: string[] |
| 75 | } |
| 76 | ) { |
| 77 | if (!values || values.length === 0) { |
| 78 | throw new Error('no value to insert') |
| 79 | } |
| 80 | const { returning, enumArrayColumns } = options ?? {} |
| 81 | const queryColumns = joinSqlFragments( |
| 82 | Object.keys(values[0]).map((x) => ident(x)), |
| 83 | ',' |
| 84 | ) |
| 85 | let query = safeSql`` |
| 86 | if (queryColumns.length == 0) { |
| 87 | query = format( |
| 88 | safeSql`insert into %1$s select from jsonb_populate_recordset(null::%1$s, %2$s)`, |
| 89 | queryTable(table), |
| 90 | literal(JSON.stringify(values)) |
| 91 | ) |
| 92 | } else { |
| 93 | query = format( |
| 94 | safeSql`insert into %1$s (%2$s) select %2$s from jsonb_populate_recordset(null::%1$s, %3$s)`, |
| 95 | queryTable(table), |
| 96 | queryColumns, |
| 97 | literal(JSON.stringify(values)) |
| 98 | ) |
| 99 | } |
| 100 | if (returning) { |
| 101 | const returningStatement = |
| 102 | enumArrayColumns === undefined || enumArrayColumns.length === 0 |
| 103 | ? safeSql` returning *` |
| 104 | : safeSql` returning *, ${joinSqlFragments( |
| 105 | enumArrayColumns.map((x) => safeSql`${ident(x)}::text[]`), |
| 106 | ',' |
| 107 | )}` |
| 108 | query = safeSql`${query}${returningStatement}` |
| 109 | } |
| 110 | return safeSql`${query};` |
| 111 | } |
| 112 | |
| 113 | export function selectQuery( |
| 114 | table: QueryTable, |
| 115 | columns?: SafeSqlFragment, |
| 116 | options?: { |
| 117 | filters?: Filter[] |
| 118 | pagination?: QueryPagination |
| 119 | sorts?: Sort[] |
| 120 | }, |
| 121 | isFinal = true, |
| 122 | isCTE = false |
| 123 | ) { |
| 124 | let query = safeSql`` |
| 125 | const queryColumn = columns ?? safeSql`*` |
| 126 | query = safeSql`select ${queryColumn} from ${isCTE ? queryCTE(table) : queryTable(table)}` |
| 127 | |
| 128 | const { filters, pagination, sorts } = options ?? {} |
| 129 | if (filters) { |
| 130 | query = applyFilters(query, filters) |
| 131 | } |
| 132 | if (sorts) { |
| 133 | query = applySorts(query, sorts) |
| 134 | } |
| 135 | if (pagination) { |
| 136 | const { limit, offset } = pagination ?? {} |
| 137 | query = safeSql`${query} limit ${literal(limit)} offset ${literal(offset)}` |
| 138 | } |
| 139 | return safeSql`${query}${isFinal ? safeSql`;` : safeSql``}` |
| 140 | } |
| 141 | |
| 142 | export function updateQuery( |
| 143 | table: QueryTable, |
| 144 | value: Dictionary<any>, |
| 145 | options?: { |
| 146 | filters?: Filter[] |
| 147 | returning?: boolean |
| 148 | enumArrayColumns?: string[] |
| 149 | } |
| 150 | ) { |
| 151 | const { filters, returning, enumArrayColumns } = options ?? {} |
| 152 | if (!filters || filters.length === 0) { |
| 153 | throw new Error('no filters for this update query') |
| 154 | } |
| 155 | const queryColumns = joinSqlFragments( |
| 156 | Object.keys(value).map((x) => ident(x)), |
| 157 | ',' |
| 158 | ) |
| 159 | let query = format( |
| 160 | safeSql`update %1$s set (%2$s) = (select %2$s from json_populate_record(null::%1$s, %3$s))`, |
| 161 | queryTable(table), |
| 162 | queryColumns, |
| 163 | literal(JSON.stringify(value)) |
| 164 | ) |
| 165 | if (filters) { |
| 166 | query = applyFilters(query, filters) |
| 167 | } |
| 168 | if (returning) { |
| 169 | const returning = |
| 170 | enumArrayColumns === undefined || enumArrayColumns.length === 0 |
| 171 | ? safeSql` returning *` |
| 172 | : safeSql` returning *, ${joinSqlFragments( |
| 173 | enumArrayColumns.map((x) => safeSql`${ident(x)}::text[]`), |
| 174 | ',' |
| 175 | )}` |
| 176 | query = safeSql`${query}${returning}` |
| 177 | } |
| 178 | |
| 179 | return safeSql`${query};` |
| 180 | } |
| 181 | |
| 182 | //============================================================ |
| 183 | // Filter Utils |
| 184 | //============================================================ |
| 185 | |
| 186 | function applyFilters(query: SafeSqlFragment, filters: Filter[]) { |
| 187 | if (filters.length === 0) return query |
| 188 | query = safeSql`${query} where ${joinSqlFragments( |
| 189 | filters.map((filter) => { |
| 190 | // Handle composite values |
| 191 | if (Array.isArray(filter.column)) { |
| 192 | switch (filter.operator) { |
| 193 | case 'in': |
| 194 | return inTupleFilterSql(filter) |
| 195 | case '=': |
| 196 | case '<>': |
| 197 | case '>': |
| 198 | case '<': |
| 199 | case '>=': |
| 200 | case '<=': |
| 201 | return defaultTupleFilterSql(filter) |
| 202 | default: |
| 203 | throw new Error(`Cannot use ${filter.operator} operator in a tuple filter`) |
| 204 | } |
| 205 | } |
| 206 | |
| 207 | switch (filter.operator) { |
| 208 | case 'in': |
| 209 | return inFilterSql(filter) |
| 210 | case 'is': |
| 211 | return isFilterSql(filter) |
| 212 | case '~~': |
| 213 | case '~~*': |
| 214 | case '!~~': |
| 215 | case '!~~*': |
| 216 | return castColumnToText(filter) |
| 217 | default: |
| 218 | return safeSql`${ident(filter.column)} ${filter.operator as SafeSqlFragment} ${filterLiteral(filter.value)}` |
| 219 | } |
| 220 | }), |
| 221 | ' and ' |
| 222 | )}` |
| 223 | return query |
| 224 | } |
| 225 | |
| 226 | function inFilterSql(filter: Filter) { |
| 227 | let values: Array<SafeSqlFragment> |
| 228 | if (Array.isArray(filter.value)) { |
| 229 | values = filter.value.map((x) => filterLiteral(x)) |
| 230 | } else { |
| 231 | const filterValueTxt = String(filter.value) |
| 232 | values = filterValueTxt.split(',').map((x) => filterLiteral(x)) |
| 233 | } |
| 234 | return safeSql`${ident(filter.column)} ${filter.operator as SafeSqlFragment} (${joinSqlFragments(values, ',')})` |
| 235 | } |
| 236 | |
| 237 | function defaultTupleFilterSql(filter: Filter) { |
| 238 | if (!Array.isArray(filter.column)) { |
| 239 | throw new Error('Use standard applyFilters for single column') |
| 240 | } |
| 241 | if (!Array.isArray(filter.value)) { |
| 242 | throw new Error('Tuple filter value must be an array') |
| 243 | } |
| 244 | if (filter.value.length !== filter.column.length) { |
| 245 | throw new Error('Tuple filter value must have the same length as the column array') |
| 246 | } |
| 247 | |
| 248 | const columns = safeSql`(${joinSqlFragments( |
| 249 | filter.column.map((c) => ident(c)), |
| 250 | ', ' |
| 251 | )})` |
| 252 | const values = safeSql`(${joinSqlFragments( |
| 253 | filter.value.map((v) => filterLiteral(v)), |
| 254 | ', ' |
| 255 | )})` |
| 256 | return safeSql`${columns} ${filter.operator as SafeSqlFragment} ${values}` |
| 257 | } |
| 258 | |
| 259 | function inTupleFilterSql(filter: Filter) { |
| 260 | if (!Array.isArray(filter.column)) { |
| 261 | throw new Error('Use inFilterSql for single columns') |
| 262 | } |
| 263 | if (!Array.isArray(filter.value)) { |
| 264 | throw new Error(`Values for a tuple 'in' filter must be an array`) |
| 265 | } |
| 266 | |
| 267 | const columns = safeSql`(${joinSqlFragments( |
| 268 | filter.column.map((c) => ident(c)), |
| 269 | ', ' |
| 270 | )})` |
| 271 | |
| 272 | const values = filter.value.map((v) => { |
| 273 | if (Array.isArray(v)) { |
| 274 | if (v.length !== filter.column.length) { |
| 275 | throw new Error(`Tuple value length must match column length`) |
| 276 | } |
| 277 | return safeSql`(${joinSqlFragments( |
| 278 | v.map((x) => filterLiteral(x)), |
| 279 | ', ' |
| 280 | )})` |
| 281 | } else { |
| 282 | const filterValueTxt = String(v) |
| 283 | const currValues = filterValueTxt.split(',') |
| 284 | if (currValues.length !== filter.column.length) { |
| 285 | throw new Error(`Tuple value length must match column length`) |
| 286 | } |
| 287 | return safeSql`(${joinSqlFragments( |
| 288 | currValues.map((x) => filterLiteral(x)), |
| 289 | ', ' |
| 290 | )})` |
| 291 | } |
| 292 | }) |
| 293 | |
| 294 | return safeSql`${columns} ${filter.operator as SafeSqlFragment} (${joinSqlFragments(values, ', ')})` |
| 295 | } |
| 296 | |
| 297 | function isFilterSql(filter: Filter) { |
| 298 | const filterValueTxt = String(filter.value) |
| 299 | switch (filterValueTxt) { |
| 300 | case 'null': |
| 301 | case 'false': |
| 302 | case 'true': |
| 303 | case 'not null': |
| 304 | return safeSql`${ident(filter.column)} ${filter.operator as SafeSqlFragment} ${filterValueTxt as SafeSqlFragment}` |
| 305 | default: |
| 306 | return safeSql`${ident(filter.column)} ${filter.operator as SafeSqlFragment} ${filterLiteral(filter.value)}` |
| 307 | } |
| 308 | } |
| 309 | |
| 310 | function castColumnToText(filter: Filter) { |
| 311 | return safeSql`${ident(filter.column)}::text ${filter.operator as SafeSqlFragment} ${filterLiteral(filter.value)}` |
| 312 | } |
| 313 | |
| 314 | function parseArrayLiteral(value: string): SafeSqlFragment | null { |
| 315 | if (!value.startsWith('ARRAY[')) return null |
| 316 | |
| 317 | // Find the closing ] of the ARRAY, tracking quoted strings |
| 318 | const afterPrefix = value.slice(6) |
| 319 | let inString = false |
| 320 | let arrayCloseIdx = -1 |
| 321 | for (let i = 0; i < afterPrefix.length; i++) { |
| 322 | const ch = afterPrefix[i] |
| 323 | if (!inString) { |
| 324 | if (ch === ']') { |
| 325 | arrayCloseIdx = i |
| 326 | break |
| 327 | } else if (ch === "'") { |
| 328 | inString = true |
| 329 | } |
| 330 | } else { |
| 331 | if (ch === "'" && afterPrefix[i + 1] === "'") { |
| 332 | i++ // escaped '' |
| 333 | } else if (ch === "'") { |
| 334 | inString = false |
| 335 | } |
| 336 | } |
| 337 | } |
| 338 | if (arrayCloseIdx === -1) return null |
| 339 | |
| 340 | const contents = afterPrefix.slice(0, arrayCloseIdx) |
| 341 | const suffix = afterPrefix.slice(arrayCloseIdx + 1) // e.g. "::status_type[]" or "" |
| 342 | |
| 343 | // Validate type cast suffix: only allow ::word_chars[]? or empty |
| 344 | let typeCast: SafeSqlFragment = safeSql`` |
| 345 | if (suffix !== '') { |
| 346 | const match = suffix.match(/^::([A-Za-z_][A-Za-z0-9_]*)(\[\])?$/) |
| 347 | if (!match) return null |
| 348 | typeCast = safeSql`::${match[1] as SafeSqlFragment}${match[2] ? safeSql`[]` : safeSql``}` |
| 349 | } |
| 350 | |
| 351 | // Parse comma-separated, single-quoted items |
| 352 | const rawItems: Array<string> = [] |
| 353 | let current = '' |
| 354 | let inStr = false |
| 355 | for (let i = 0; i < contents.length; i++) { |
| 356 | const ch = contents[i] |
| 357 | if (!inStr) { |
| 358 | if (ch === "'") { |
| 359 | inStr = true |
| 360 | current += ch |
| 361 | } else if (ch === ',') { |
| 362 | rawItems.push(current.trim()) |
| 363 | current = '' |
| 364 | } else { |
| 365 | current += ch |
| 366 | } |
| 367 | } else { |
| 368 | if (ch === "'" && contents[i + 1] === "'") { |
| 369 | current += "''" |
| 370 | i++ |
| 371 | } else if (ch === "'") { |
| 372 | current += ch |
| 373 | inStr = false |
| 374 | } else { |
| 375 | current += ch |
| 376 | } |
| 377 | } |
| 378 | } |
| 379 | if (current.trim()) rawItems.push(current.trim()) |
| 380 | |
| 381 | const unquoted = rawItems.map((item) => { |
| 382 | if (item.startsWith("'") && item.endsWith("'")) { |
| 383 | return item.slice(1, -1).replace(/''/g, "'") |
| 384 | } |
| 385 | return item |
| 386 | }) |
| 387 | |
| 388 | const formattedItems = joinSqlFragments( |
| 389 | unquoted.map((x) => literal(x)), |
| 390 | ',' |
| 391 | ) |
| 392 | return safeSql`ARRAY[${formattedItems}]${typeCast}` |
| 393 | } |
| 394 | |
| 395 | function filterLiteral(value: any): SafeSqlFragment { |
| 396 | if (typeof value === 'boolean') { |
| 397 | return (value ? 'true' : 'false') as SafeSqlFragment |
| 398 | } |
| 399 | if (typeof value === 'string') { |
| 400 | if (value.startsWith('ARRAY[')) { |
| 401 | const parsed = parseArrayLiteral(value) |
| 402 | if (parsed !== null) return parsed |
| 403 | } |
| 404 | return literal(value) |
| 405 | } |
| 406 | return literal(value) |
| 407 | } |
| 408 | |
| 409 | //============================================================ |
| 410 | // Sort Utils |
| 411 | //============================================================ |
| 412 | |
| 413 | function applySorts(query: SafeSqlFragment, sorts: Sort[]): SafeSqlFragment { |
| 414 | const validSorts = sorts.filter((sort) => sort.column) |
| 415 | if (validSorts.length === 0) return query |
| 416 | query = safeSql`${query} order by ${joinSqlFragments( |
| 417 | validSorts.map((x) => { |
| 418 | const order = x.ascending ? safeSql`asc` : safeSql`desc` |
| 419 | const nullOrder = x.nullsFirst ? safeSql`nulls first` : safeSql`nulls last` |
| 420 | return safeSql`${ident(x.table)}.${ident(x.column)} ${order} ${nullOrder}` |
| 421 | }), |
| 422 | ', ' |
| 423 | )}` |
| 424 | return query |
| 425 | } |
| 426 | |
| 427 | //============================================================ |
| 428 | // Misc |
| 429 | //============================================================ |
| 430 | |
| 431 | function queryTable(table: QueryTable) { |
| 432 | return safeSql`${ident(table.schema)}.${ident(table.name)}` |
| 433 | } |
| 434 | |
| 435 | function queryCTE(table: QueryTable) { |
| 436 | return safeSql`${ident(table.name)}` |
| 437 | } |
| 438 | |
| 439 | export function wrapWithTransaction(sql: SafeSqlFragment) { |
| 440 | return safeSql` |
| 441 | begin; |
| 442 | |
| 443 | ${sql} |
| 444 | |
| 445 | commit; |
| 446 | ` |
| 447 | } |
| 448 | |
| 449 | export function wrapWithRollback(sql: SafeSqlFragment) { |
| 450 | return safeSql` |
| 451 | begin; |
| 452 | |
| 453 | ${sql} |
| 454 | |
| 455 | rollback; |
| 456 | ` |
| 457 | } |