diff --git a/config.nims b/config.nims index 9eb46ad..0377d94 100644 --- a/config.nims +++ b/config.nims @@ -20,6 +20,8 @@ task test, "Run all test suite": exec "nim c -f -r tests/tsqlite" exec "nim c -f -r tests/tdb_utils" exec "nim c -f -r tests/timportstatic" + exec "nim c -f -r tests/tpostgres_schema_import" + exec "nim c -f -r tests/tqualified_schema_queries" task setup_postgres, "Ensure local Postgres has test DB/user": # Use a simple script to avoid Nim/psql quoting pitfalls diff --git a/ormin.nimble b/ormin.nimble index 1de2b6e..dcb6a15 100644 --- a/ormin.nimble +++ b/ormin.nimble @@ -1,6 +1,6 @@ # Package -version = "0.9.0" +version = "0.10.0" author = "Araq" description = "Prepared SQL statement generator. A lightweight ORM." license = "MIT" diff --git a/ormin/db_types.nim b/ormin/db_types.nim index 29cc10f..df4b5f8 100644 --- a/ormin/db_types.nim +++ b/ormin/db_types.nim @@ -5,7 +5,7 @@ proc dbTypFromName*(name: string): DbTypeKind = var k = dbUnknown case name.toLowerAscii - of "int", "integer", "int8", "smallint", "int16", + of "int", "integer", "int2", "int4", "int8", "smallint", "bigint", "int16", "longint", "int32", "int64", "tinyint", "hugeint": k = dbInt of "uint", "uint8", "uint16", "uint32", "uint64": k = dbUInt of "serial": k = dbSerial @@ -14,7 +14,7 @@ proc dbTypFromName*(name: string): DbTypeKind = of "blob": k = dbBlob of "fixedchar": k = dbFixedChar of "varchar", "text", "string": k = dbVarchar - of "json": k = dbJson + of "json", "jsonb": k = dbJson of "xml": k = dbXml of "decimal": k = dbDecimal of "float", "double", "longdouble", "real": k = dbFloat diff --git a/ormin/db_utils.nim b/ormin/db_utils.nim index 0d0398b..e0fae18 100644 --- a/ormin/db_utils.nim +++ b/ormin/db_utils.nim @@ -113,12 +113,12 @@ iterator tableDefs(sql: DbSql): tuple[name, tableName, model: string] = for i in 0 ..< ast.len: let node = ast[i] if node.kind in {nkCreateTable, nkCreateTableIfNotExists}: - yield (node[0].strVal.toLowerAscii(), $node[0], $node) + yield (sqlIdentBaseName(node[0]).toLowerAscii(), sqlIdentName(node[0]), $node) else: # Fallback: ast might be a single statement (not a list) let node = ast if node.kind in {nkCreateTable, nkCreateTableIfNotExists}: - yield (node[0].strVal.toLowerAscii(), $node[0], $node) + yield (sqlIdentBaseName(node[0]).toLowerAscii(), sqlIdentName(node[0]), $node) iterator tablePairs*(sql: string): tuple[name, model: string] = for name, _, model in tableDefs(DbSql(sql)): diff --git a/ormin/importer_core.nim b/ormin/importer_core.nim index ae91a11..901349a 100644 --- a/ormin/importer_core.nim +++ b/ormin/importer_core.nim @@ -13,6 +13,8 @@ type name: string tabIndex: int typ: DbTypekind + typeName: string + validValues: seq[string] key: int # 0 nothing special, # +1 -- primary key # -N -- references attribute N @@ -28,6 +30,7 @@ type DbColumns* = seq[DbColumn] KnownTables* = OrderedTable[string, DbColumns] + KnownEnums = Table[string, seq[string]] ImportTarget* = enum postgre, sqlite, mysql @@ -40,13 +43,30 @@ proc hasRefs(colDesc: SqlNode): (string, string) = for i in 2 ..< colDesc.len: let c = colDesc[i] if c.kind == nkReferences: - if c[0].kind == nkCall: - return (c[0][0].strVal, c[0][1].strVal) - elif c[0].kind == nkIdent: - return ($c[0], "id") + if c[0].kind == nkColumnReference: + return (sqlIdentName(c[0][0]), sqlIdentBaseName(c[0][1])) + elif c[0].kind == nkCall: + return (sqlIdentName(c[0][0]), sqlIdentBaseName(c[0][1])) + elif c[0].kind in {nkIdent, nkQuotedIdent, nkDot}: + return (sqlIdentName(c[0]), "id") ("", "") -proc getType(n: SqlNode): DbType = +proc collectEnumTypes(schemaSql: string): KnownEnums = + result = initTable[string, seq[string]]() + let definitions = parseEnumTypeDefs(schemaSql) + for i in 0 ..< definitions.len: + let definition = definitions[i] + let typeName = sqlIdentName(definition[0]).toLowerAscii() + var values: seq[string] = @[] + for value in definition[1].sons: + values.add(value.strVal) + result[typeName] = values + + let dotPos = typeName.rfind('.') + if dotPos >= 0 and dotPos + 1 < typeName.len: + result[typeName[dotPos + 1 .. ^1]] = values + +proc getType(n: SqlNode; enums: KnownEnums): DbType = var it = n if it.kind == nkCall: it = it[0] @@ -56,21 +76,28 @@ proc getType(n: SqlNode): DbType = for i in 0 ..< it.len: assert it[i].kind == nkStringLit result.validValues.add it[i].strVal - elif it.kind in {nkIdent, nkStringLit}: - result.kind = dbTypFromName(it.strVal) - result.name = it.strVal + elif it.kind in {nkIdent, nkQuotedIdent, nkStringLit, nkDot}: + let typeName = sqlIdentName(it) + let normalized = typeName.toLowerAscii + if enums.hasKey(normalized): + result.kind = dbEnum + result.validValues = enums[normalized] + else: + let baseName = sqlIdentBaseName(it) + result.kind = dbTypFromName(baseName) + result.name = typeName -proc collectTables*(n: SqlNode; t: var KnownTables) = +proc collectTables*(n: SqlNode; t: var KnownTables; enums: KnownEnums) = if n.isNil: return case n.kind of nkCreateTable, nkCreateTableIfNotExists: - let tableName = n[0].strVal + let tableName = sqlIdentName(n[0]) var cols: DbColumns = @[] for i in 1 ..< n.len: let it = n[i] if it.kind == nkColumnDef: - var typ = getType(it[1]) + var typ = getType(it[1], enums) if hasAttribute(it, {nkNotNull}): typ.notNull = true cols.add DbColumn( @@ -97,9 +124,9 @@ proc collectTables*(n: SqlNode; t: var KnownTables) = var refTable = "" var refCols: seq[string] = @[] if r.kind == nkColumnReference or r.kind == nkCall: - refTable = r[0].strVal + refTable = sqlIdentName(r[0]) for k in 1 ..< r.len: - refCols.add(r[k].strVal) + refCols.add(sqlIdentBaseName(r[k])) let pairCount = min(localCols.len, refCols.len) for k in 0 ..< pairCount: let localName = localCols[k] @@ -111,25 +138,54 @@ proc collectTables*(n: SqlNode; t: var KnownTables) = t[tableName] = cols else: for i in 0 ..< n.len: - collectTables(n[i], t) + collectTables(n[i], t, enums) proc attrToKey(a: DbColumn; t: KnownTables): int = if a.primaryKey: return 1 if a.refs[0].len > 0: + var referencedTable = "" + for tableName in keys(t): + if cmpIgnoreCase(tableName, a.refs[0]) == 0: + referencedTable = tableName + break + + if referencedTable.len == 0 and '.' notin a.refs[0]: + for tableName in keys(t): + let dotPos = tableName.rfind('.') + let baseName = + if dotPos >= 0: tableName[dotPos + 1 .. ^1] + else: tableName + if cmpIgnoreCase(baseName, a.refs[0]) == 0: + if referencedTable.len > 0: + return 0 + referencedTable = tableName + + if referencedTable.len == 0: + return 0 + var i = 0 for k, v in pairs(t): for b in v: - if cmpIgnoreCase(k, a.refs[0]) == 0 and cmpIgnoreCase(b.name, a.refs[1]) == 0: + if cmpIgnoreCase(k, referencedTable) == 0 and cmpIgnoreCase(b.name, a.refs[1]) == 0: return -i - 1 inc i 0 +proc addStringSeq(dest: var string; values: openArray[string]) = + dest.add "@[" + for i, value in values: + if i > 0: + dest.add ", " + dest.add escape(value) + dest.add "]" + proc renderModelCode(schemaSql, schemaPath: string; target: ImportTarget; includeStatic = false): string = discard target let sql = parseSql(schemaSql, schemaPath) + let enums = collectEnumTypes(schemaSql) var knownTables = initOrderedTable[string, DbColumns]() - collectTables(sql, knownTables) + collectTables(sql, knownTables, enums) result.add FileHeader result.add "const tableNames = [" @@ -158,6 +214,10 @@ proc renderModelCode(schemaSql, schemaPath: string; target: ImportTarget; includ result.add $i result.add ", typ: " result.add $a.typ.kind + result.add ", typeName: " + result.add escape(a.typ.name) + result.add ", validValues: " + result.addStringSeq(a.typ.validValues) result.add ", key: " result.add $attrToKey(a, knownTables) result.add ")" diff --git a/ormin/parsesql_tmp.nim b/ormin/parsesql_tmp.nim index b2b4eb9..8563f3e 100644 --- a/ormin/parsesql_tmp.nim +++ b/ormin/parsesql_tmp.nim @@ -477,6 +477,7 @@ type nkHexStringLit, nkIntegerLit, nkNumericLit, + nkRaw, nkPrimaryKey, nkForeignKey, nkNotNull, @@ -540,7 +541,7 @@ type const LiteralNodes = { nkIdent, nkQuotedIdent, nkStringLit, nkBitStringLit, nkHexStringLit, - nkIntegerLit, nkNumericLit + nkIntegerLit, nkNumericLit, nkRaw } type @@ -588,6 +589,26 @@ proc `[]`*(n: SqlNode; i: BackwardsIndex): SqlNode = n.sons[n.len - int(i)] proc add*(father, n: SqlNode) = add(father.sons, n) +proc sqlIdentName*(n: SqlNode): string = + ## Return a SQL identifier name from an identifier or dotted identifier node. + case n.kind + of nkIdent, nkQuotedIdent: + result = n.strVal + of nkDot: + result = sqlIdentName(n[0]) & "." & sqlIdentName(n[1]) + else: + result = "" + +proc sqlIdentBaseName*(n: SqlNode): string = + ## Return the unqualified final identifier from an identifier node. + case n.kind + of nkIdent, nkQuotedIdent: + result = n.strVal + of nkDot: + result = sqlIdentBaseName(n[1]) + else: + result = "" + proc getTok(p: var SqlParser) = getTok(p, p.tok) @@ -632,6 +653,145 @@ proc eat(p: var SqlParser, keyw: string) = proc opt(p: var SqlParser, kind: TokKind) = if p.tok.kind == kind: getTok(p) +proc skipToSemicolon(p: var SqlParser) = + while p.tok.kind notin {tkSemicolon, tkEof}: + getTok(p) + +proc skipBalancedParens(p: var SqlParser) = + if p.tok.kind != tkParLe: + return + + var depth = 0 + while p.tok.kind != tkEof: + var shouldReadNext = true + case p.tok.kind + of tkParLe: + inc depth + of tkParRi: + dec depth + getTok(p) + if depth == 0: + return + else: + shouldReadNext = false + else: + discard + if shouldReadNext: + getTok(p) + +proc stripTrailingSqlSpace(sql: var string) = + while sql.len > 0 and sql[^1] in Whitespace: + sql.setLen(sql.len - 1) + +proc addSqlStringLiteral(sql: var string; value: string; prefix = "") = + sql.add(prefix) + sql.add('\'') + sql.add(value.replace("'", "''")) + sql.add('\'') + +proc addSqlToken(sql: var string; tok: Token) = + case tok.kind + of tkParLe: + sql.add('(') + of tkParRi: + sql.stripTrailingSqlSpace() + sql.add(')') + of tkBracketLe: + sql.add('[') + of tkBracketRi: + sql.stripTrailingSqlSpace() + sql.add(']') + of tkComma: + sql.stripTrailingSqlSpace() + sql.add(", ") + of tkDot, tkColon: + sql.stripTrailingSqlSpace() + sql.add(tok.literal) + of tkOperator: + if sql.len > 0 and sql[^1] notin Whitespace + {'(', '[', '.', ':'}: + sql.add(' ') + sql.add(tok.literal) + sql.add(' ') + of tkQuotedIdentifier: + if sql.len > 0 and sql[^1] notin Whitespace + {'(', '[', '.', ':'}: + sql.add(' ') + sql.add('"') + sql.add(tok.literal.replace("\"", "\"\"")) + sql.add('"') + of tkStringConstant: + if sql.len > 0 and sql[^1] notin Whitespace + {'(', '[', '.', ':'}: + sql.add(' ') + sql.addSqlStringLiteral(tok.literal) + of tkEscapeConstant: + if sql.len > 0 and sql[^1] notin Whitespace + {'(', '[', '.', ':'}: + sql.add(' ') + var escaped = tok.literal.replace("\\", "\\\\") + escaped = escaped.replace("'", "''") + sql.add("E'") + sql.add(escaped) + sql.add('\'') + of tkDollarQuotedConstant: + if sql.len > 0 and sql[^1] notin Whitespace + {'(', '[', '.', ':'}: + sql.add(' ') + var tag = "$ormin$" + while tag in tok.literal: + tag = tag[0 ..< ^1] & "_or$" + sql.add(tag) + sql.add(tok.literal) + sql.add(tag) + of tkBitStringConstant: + if sql.len > 0 and sql[^1] notin Whitespace + {'(', '[', '.', ':'}: + sql.add(' ') + sql.addSqlStringLiteral(tok.literal, "B") + of tkHexStringConstant: + if sql.len > 0 and sql[^1] notin Whitespace + {'(', '[', '.', ':'}: + sql.add(' ') + sql.addSqlStringLiteral(tok.literal, "X") + of tkEof: + discard + else: + if sql.len > 0 and sql[^1] notin Whitespace + {'(', '[', '.', ':'}: + sql.add(' ') + sql.add(tok.literal) + +proc readBalancedSql(p: var SqlParser): string = + if p.tok.kind != tkParLe: + return + + var depth = 0 + while p.tok.kind != tkEof: + let kind = p.tok.kind + case kind + of tkParLe: + inc depth + of tkParRi: + dec depth + else: + discard + result.addSqlToken(p.tok) + getTok(p) + if kind == tkParRi and depth == 0: + return + + sqlError(p, "closing parenthesis expected") + +proc parseIdentNode(p: var SqlParser): SqlNode = + expectIdent(p) + if p.tok.kind == tkQuotedIdentifier: + result = newNode(nkQuotedIdent, p.tok.literal) + else: + result = newNode(nkIdent, p.tok.literal) + getTok(p) + +proc parseQualifiedIdentifier(p: var SqlParser): SqlNode = + result = parseIdentNode(p) + while p.tok.kind == tkDot: + getTok(p) + let left = result + result = newNode(nkDot) + result.add(left) + result.add(parseIdentNode(p)) + proc parseDataType(p: var SqlParser): SqlNode = if isKeyw(p, "enum"): result = newNode(nkEnumDef) @@ -646,9 +806,7 @@ proc parseDataType(p: var SqlParser): SqlNode = getTok(p) eat(p, tkParRi) else: - expectIdent(p) - result = newNode(nkIdent, p.tok.literal) - getTok(p) + result = parseQualifiedIdentifier(p) if p.tok.kind == tkParLe: var complexType = newNode(nkCall) complexType.add(result) @@ -766,6 +924,10 @@ proc primary(p: var SqlParser): SqlNode = else: sqlError(p, "identifier expected") getTok(p) + of tkColon: + getTok(p) + eat(p, tkColon) + discard parseDataType(p) else: break proc lowestExprAux(p: var SqlParser, v: out SqlNode, limit: int): int = @@ -789,8 +951,7 @@ proc parseExpr(p: var SqlParser): SqlNode = discard lowestExprAux(p, result, - 1) proc parseTableName(p: var SqlParser): SqlNode = - expectIdent(p) - result = primary(p) + result = parseQualifiedIdentifier(p) proc parseColumnReference(p: var SqlParser): SqlNode = result = parseTableName(p) @@ -805,19 +966,37 @@ proc parseColumnReference(p: var SqlParser): SqlNode = result.add(parseTableName(p)) eat(p, tkParRi) +proc parseTableConstraint(p: var SqlParser): SqlNode + proc parseCheck(p: var SqlParser): SqlNode = getTok(p) result = newNode(nkCheck) - result.add(parseExpr(p)) + if p.tok.kind == tkParLe: + result.add(newNode(nkRaw, readBalancedSql(p))) + else: + result.add(parseExpr(p)) proc parseConstraint(p: var SqlParser): SqlNode = getTok(p) - result = newNode(nkConstraint) expectIdent(p) - result.add(newNode(nkIdent, p.tok.literal)) + let constraintName = newNode(nkIdent, p.tok.literal) getTok(p) - optKeyw(p, "check") - result.add(parseExpr(p)) + if isKeyw(p, "foreign") or isKeyw(p, "primary") or isKeyw(p, "unique"): + result = parseTableConstraint(p) + elif isKeyw(p, "check"): + result = newNode(nkConstraint) + result.add(constraintName) + let checkNode = parseCheck(p) + if checkNode.len > 0: + result.add(checkNode[0]) + else: + result.add(newNode(nkIdent, "true")) + else: + result = newNode(nkConstraint) + result.add(constraintName) + result.add(newNode(nkIdent, "true")) + while p.tok.kind notin {tkComma, tkParRi, tkEof}: + getTok(p) proc parseParIdentList(p: var SqlParser, father: SqlNode) = eat(p, tkParLe) @@ -937,6 +1116,22 @@ proc parseColumnConstraints(p: var SqlParser, result: SqlNode) = elif isKeyw(p, "identity"): getTok(p) result.add(newNode(nkIdentity)) + elif isKeyw(p, "generated"): + getTok(p) + optKeyw(p, "always") + if isKeyw(p, "by"): + getTok(p) + optKeyw(p, "default") + optKeyw(p, "as") + if isKeyw(p, "identity"): + getTok(p) + result.add(newNode(nkIdentity)) + if p.tok.kind == tkParLe: + skipBalancedParens(p) + elif p.tok.kind == tkParLe: + skipBalancedParens(p) + optKeyw(p, "stored") + optKeyw(p, "virtual") elif isKeyw(p, "primary"): getTok(p) eat(p, "key") @@ -1076,7 +1271,7 @@ proc parseTableConstraint(p: var SqlParser): SqlNode = result.add(m) elif isKeyw(p, "unique"): getTok(p) - eat(p, "key") + optKeyw(p, "key") result = newNode(nkUnique) parseParIdentList(p, result) elif isKeyw(p, "check"): @@ -1092,12 +1287,7 @@ proc parseUnique(p: var SqlParser): SqlNode = proc parseTableDef(p: var SqlParser): SqlNode = result = parseIfNotExists(p, nkCreateTable) - expectIdent(p) - if p.tok.kind == tkQuotedIdentifier: - result.add(newNode(nkQuotedIdent, p.tok.literal)) - else: - result.add(newNode(nkIdent, p.tok.literal)) - getTok(p) + result.add(parseQualifiedIdentifier(p)) if p.tok.kind == tkParLe: getTok(p) while p.tok.kind != tkParRi: @@ -1120,9 +1310,7 @@ proc parseTableDef(p: var SqlParser): SqlNode = proc parseTypeDef(p: var SqlParser): SqlNode = result = parseIfNotExists(p, nkCreateType) - expectIdent(p) - result.add(newNode(nkIdent, p.tok.literal)) - getTok(p) + result.add(parseQualifiedIdentifier(p)) eat(p, "as") result.add(parseDataType(p)) @@ -1212,27 +1400,16 @@ proc parseIndexDef(p: var SqlParser): SqlNode = result.add(newNode(nkIdent, p.tok.literal)) getTok(p) eat(p, "on") - expectIdent(p) - result.add(newNode(nkIdent, p.tok.literal)) - getTok(p) + result.add(parseQualifiedIdentifier(p)) eat(p, tkParLe) - expectIdent(p) - result.add(newNode(nkIdent, p.tok.literal)) - getTok(p) - while p.tok.kind == tkComma: - getTok(p) - expectIdent(p) - result.add(newNode(nkIdent, p.tok.literal)) - getTok(p) - eat(p, tkParRi) + skipBalancedParens(p) + skipToSemicolon(p) proc parseInsert(p: var SqlParser): SqlNode = getTok(p) eat(p, "into") - expectIdent(p) result = newNode(nkInsert) - result.add(newNode(nkIdent, p.tok.literal)) - getTok(p) + result.add(parseQualifiedIdentifier(p)) if p.tok.kind == tkParLe: var n = newNode(nkColumnList) parseParIdentList(p, n) @@ -1253,6 +1430,7 @@ proc parseInsert(p: var SqlParser): SqlNode = getTok(p) result.add(n) eat(p, tkParRi) + skipToSemicolon(p) proc parseUpdate(p: var SqlParser): SqlNode = getTok(p) @@ -1401,6 +1579,9 @@ proc parseSelect(p: var SqlParser): SqlNode = proc parseStmt(p: var SqlParser; parent: SqlNode) = if isKeyw(p, "create"): getTok(p) + if isKeyw(p, "or"): + getTok(p) + optKeyw(p, "replace") optKeyw(p, "cached") optKeyw(p, "memory") optKeyw(p, "temp") @@ -1416,7 +1597,7 @@ proc parseStmt(p: var SqlParser; parent: SqlNode) = elif isKeyw(p, "index"): parent.add parseIndexDef(p) else: - sqlError(p, "TABLE expected") + skipToSemicolon(p) elif isKeyw(p, "insert"): parent.add parseInsert(p) elif isKeyw(p, "update"): @@ -1430,8 +1611,12 @@ proc parseStmt(p: var SqlParser; parent: SqlNode) = parent.add parsePragma(p) elif isKeyw(p, "begin"): getTok(p) + elif isKeyw(p, "do") or isKeyw(p, "drop") or isKeyw(p, "alter") or + isKeyw(p, "grant") or isKeyw(p, "revoke") or isKeyw(p, "comment") or + isKeyw(p, "notify") or isKeyw(p, "listen"): + skipToSemicolon(p) else: - sqlError(p, "SELECT, CREATE, UPDATE or DELETE expected") + skipToSemicolon(p) proc parse(p: var SqlParser): SqlNode = ## parses the content of `p`'s input stream and returns the SQL AST. @@ -1534,6 +1719,8 @@ proc ra(n: SqlNode, s: var SqlWriter) = s.add("x'" & n.strVal & "'") of nkIntegerLit, nkNumericLit: s.add(n.strVal) + of nkRaw: + s.add(n.strVal) of nkPrimaryKey: s.addKeyw("primary key") rs(n, s) @@ -1856,3 +2043,32 @@ proc parseSql*(input: string, filename = "", considerTypeParams = false): SqlNod ## `filename` is only used for error messages. ## Syntax errors raise an `SqlParseError` exception. parseSql(newStringStream(input), "", considerTypeParams) + +proc scanEnumTypeDefs(input, filename: string; definitions: SqlNode) = + var p: SqlParser + open(p, newStringStream(input), filename) + try: + while p.tok.kind != tkEof: + if p.tok.kind == tkDollarQuotedConstant: + let body = p.tok.literal + getTok(p) + scanEnumTypeDefs(body, filename, definitions) + elif isKeyw(p, "create"): + getTok(p) + if isKeyw(p, "type"): + let definition = parseIfNotExists(p, nkCreateType) + definition.add(parseQualifiedIdentifier(p)) + if isKeyw(p, "as"): + getTok(p) + if isKeyw(p, "enum"): + definition.add(parseDataType(p)) + definitions.add(definition) + else: + getTok(p) + finally: + close(p) + +proc parseEnumTypeDefs*(input: string; filename = ""): SqlNode = + ## Finds enum type definitions, including definitions inside dollar-quoted blocks. + result = newNode(nkStmtList) + scanEnumTypeDefs(input, filename, result) diff --git a/ormin/queries.nim b/ormin/queries.nim index 8848126..85a7fc6 100644 --- a/ormin/queries.nim +++ b/ormin/queries.nim @@ -217,6 +217,13 @@ proc lookupCte(ctes: openArray[CteDef]; name: string): int {.compileTime.} = if cmpIgnoreCase(cte.name, name) == 0: return i +proc baseTableName(name: string): string {.compileTime.} = + let dotPos = name.rfind('.') + if dotPos >= 0: + result = name[dotPos + 1 .. ^1] + else: + result = name + proc sourceName(q: QueryBuilder; source: int): string {.compileTime.} = if isCteEnvIndex(source): result = q.ctes[fromCteEnvIndex(source)].name @@ -229,17 +236,39 @@ proc sourceColumns(q: QueryBuilder; source: int): seq[SourceColumn] {.compileTim else: for a in attributes: if a.tabIndex == source: - result.add SourceColumn(name: a.name, typ: DbType(kind: a.typ)) + var typ = DbType(kind: a.typ) + when compiles(a.typeName): + typ.name = a.typeName + when compiles(a.validValues): + typ.validValues = a.validValues + result.add SourceColumn(name: a.name, typ: typ) proc sourceLookup(q: QueryBuilder; table: string): int {.compileTime.} = for i, t in tableNames: if cmpIgnoreCase(t, table) == 0: return i + let cteIdx = lookupCte(q.ctes, table) if cteIdx >= 0: return cteEnvIndex(cteIdx) + + var baseMatch = -1 + if '.' notin table: + for i, t in tableNames: + if cmpIgnoreCase(baseTableName(t), table) == 0: + if baseMatch >= 0: + return -1 + baseMatch = i + if baseMatch >= 0: + return baseMatch result = -1 +proc sourceMatches(q: QueryBuilder; source: int; table: string): bool {.compileTime.} = + let name = sourceName(q, source) + result = cmpIgnoreCase(name, table) == 0 + if not result and not isCteEnvIndex(source) and '.' notin table: + result = cmpIgnoreCase(baseTableName(name), table) == 0 + proc sourceAlias(q: QueryBuilder; source: int; sourceName: string): string {.compileTime.} = if q.kind == qkJoin and q.env.len > 0 and q.env[^1][0] == source: result = q.env[^1][1] @@ -270,7 +299,7 @@ proc lookup(table, attr: string; qb: QueryBuilder; alias: var string): DbType = var found = false var foundSource = -1 for e in qb.env: - if table.len == 0 or cmpIgnoreCase(sourceName(qb, e[0]), table) == 0: + if table.len == 0 or sourceMatches(qb, e[0], table): for col in sourceColumns(qb, e[0]): if cmpIgnoreCase(col.name, attr) == 0: if found: @@ -359,6 +388,12 @@ proc nodeName(n: NimNode): string {.compileTime.} = result = nodeName(n[0]) else: result = "" + of nnkDotExpr: + if n.len == 2: + let left = nodeName(n[0]) + let right = nodeName(n[1]) + if left.len > 0 and right.len > 0: + result = left & "." & right else: result = "" @@ -506,8 +541,8 @@ proc cond(n: NimNode; q: var string; params: var Params; else: result = lookupColumnInEnv(n, q, params, expected, qb) of nnkDotExpr: - let t = $n[0] - let a = $n[1] + let t = nodeName(n[0]) + let a = nodeName(n[1]) escIdent(q, t) q.add '.' escIdent(q, a) @@ -922,7 +957,7 @@ proc selectAll(q: QueryBuilder; tabIndex: int; arg, lineInfo: NimNode) = proc tableSel(n: NimNode; q: QueryBuilder) = if n.kind == nnkCall and q.kind != qkDelete: let call = n - let tab = $call[0] + let tab = nodeName(call[0]) let tabIndex = sourceLookup(q, tab) if tabIndex < 0: macros.error "unknown table name: " & tab & " from: " & fmtTableList(tableNames), n @@ -1009,8 +1044,8 @@ proc tableSel(n: NimNode; q: QueryBuilder) = else: macros.error "unknown selector: " & repr(n), n if q.kind notin {qkUpdate, qkSelect, qkJoin}: q.head.add ")" - elif n.kind in {nnkIdent, nnkAccQuoted, nnkSym} and q.kind == qkDelete: - let tab = $n + elif n.kind in {nnkIdent, nnkAccQuoted, nnkSym, nnkDotExpr} and q.kind == qkDelete: + let tab = nodeName(n) let tabIndex = sourceLookup(q, tab) if tabIndex < 0: macros.error "unknown table name: " & tab & " from: " & fmtTableList(tableNames), n @@ -1137,7 +1172,7 @@ proc queryh(n: NimNode; q: QueryBuilder) = if joinClause.kind == nnkCommand and joinClause.len == 2 and joinClause[1].kind == nnkCommand and joinClause[1].len == 2 and $joinClause[1][0] == "on" and joinClause[0].kind == nnkCall: - let tab = $joinClause[0][0] + let tab = nodeName(joinClause[0][0]) let tabIndex = sourceLookup(q, tab) if tabIndex < 0: macros.error "unknown table name: " & tab & " from: " & fmtTableList(tableNames), n @@ -1158,7 +1193,7 @@ proc queryh(n: NimNode; q: QueryBuilder) = swap q.env, oldEnv checkBool(t, onn) elif joinClause.kind == nnkCall: - let tab = $joinClause[0] + let tab = nodeName(joinClause[0]) let tabIndex = sourceLookup(q, tab) if tabIndex < 0: macros.error "unknown table name: " & tab & " from: " & fmtTableList(tableNames), n[1][0] diff --git a/tests/qualified_schema_model.sql b/tests/qualified_schema_model.sql new file mode 100644 index 0000000..ef90b85 --- /dev/null +++ b/tests/qualified_schema_model.sql @@ -0,0 +1,11 @@ +create table public.events ( + public_value text +); + +create table audit.events ( + audit_value text +); + +create table public.users ( + username text +); diff --git a/tests/tdb_utils.nim b/tests/tdb_utils.nim index 525c1ef..45c037f 100644 --- a/tests/tdb_utils.nim +++ b/tests/tdb_utils.nim @@ -1,4 +1,4 @@ -import unittest, os, sequtils +import std/[assertions, os, sequtils, strutils, unittest] import db_connector/db_common from db_connector/db_sqlite import open, exec, getValue import ormin/db_utils @@ -41,6 +41,22 @@ let sqlContent = """ const staticSqlContent = staticLoad("db_utils_case_quoted.sql") +block checkConstraintRoundTrip: + const schema = """ +create table accounts ( + balance integer check (balance >= 0), + code text constraint normalized_code check ( + code ~ '^[A-Z]+$' and code = 'X'::text + ) +); +""" + let pairs = tablePairs(schema).toSeq() + doAssert pairs.len == 1 + doAssert pairs[0].model.contains("balance >= 0") + doAssert pairs[0].model.contains("code ~ '^[A-Z]+$'") + doAssert pairs[0].model.contains("'X'::text") + doAssert not pairs[0].model.contains("check true") + writeFile($sqlFile, sqlContent) suite "db_utils: case and quoted names": diff --git a/tests/tpostgres_schema_import.nim b/tests/tpostgres_schema_import.nim new file mode 100644 index 0000000..569e23e --- /dev/null +++ b/tests/tpostgres_schema_import.nim @@ -0,0 +1,134 @@ +import std/[assertions, strutils] + +import ormin/importer_core + +const postgresSchema = """ +do $$ +begin + create type public.client_kind as enum ( + 'device', + 'internal', + 'partner' + ); +exception + when duplicate_object then null; +end $$; + +do $$ +begin + create type public.resource_kind as enum ( + 'platform', + 'organization', + 'device', + 'integration' + ); +exception + when duplicate_object then null; +end $$; + +create table if not exists public.clients ( + id uuid primary key default gen_random_uuid(), + client_id text not null unique + check (client_id ~ '^[A-Za-z0-9._:-]{8,128}$'), + kind public.client_kind not null, + enabled boolean not null default true, + metadata jsonb not null default '{}'::jsonb, + created_at timestamptz not null default now() +); + +insert into public.clients (client_id, kind) +values + ('device-001', 'device'), + ('partner-001', 'partner') +on conflict (client_id) do nothing; + +create table if not exists public.client_resource_grants ( + id bigint generated always as identity primary key, + service_client_id uuid not null + references public.clients(id) on delete cascade, + resource_kind public.resource_kind not null, + organization_id uuid, + device_id text, + integration text, + resource_key text generated always as ( + case resource_kind + when 'platform'::public.resource_kind then 'platform' + when 'organization'::public.resource_kind then organization_id::text + when 'device'::public.resource_kind then organization_id::text || ':' || device_id + when 'integration'::public.resource_kind then lower(integration) + end + ) stored, + constraint client_resource_grants_shape check ( + resource_kind = 'platform'::public.resource_kind + or resource_kind = 'device'::public.resource_kind + ) +); + +create index if not exists idx_client_resource_grants_device + on public.client_resource_grants(organization_id, device_id) + where device_id is not null; + +drop trigger if exists set_clients_updated_at on public.clients; +create trigger set_clients_updated_at +before update on public.clients +for each row execute function public.set_updated_at(); +""" + +let schema = postgresSchema +let model = generateModelCode(schema, "postgres_schema.sql", postgre) + +doAssert model.contains("\"public.clients\"") +doAssert model.contains("\"public.client_resource_grants\"") +doAssert model.contains("Attr(name: \"id\", tabIndex: 0, typ: dbUuid") +doAssert model.contains("typeName: \"uuid\", validValues: @[], key: 1") +doAssert model.contains("Attr(name: \"kind\", tabIndex: 0, typ: dbEnum") +doAssert model.contains( + "typeName: \"public.client_kind\", validValues: @[" & + "\"device\", \"internal\", \"partner\"]" +) +doAssert model.contains("Attr(name: \"metadata\", tabIndex: 0, typ: dbJson") +doAssert model.contains("typeName: \"jsonb\", validValues: @[], key: 0") +doAssert model.contains("Attr(name: \"id\", tabIndex: 1, typ: dbInt") +doAssert model.contains("typeName: \"bigint\", validValues: @[], key: 1") +doAssert model.contains("Attr(name: \"resource_kind\", tabIndex: 1, typ: dbEnum") +doAssert model.contains( + "typeName: \"public.resource_kind\", validValues: @[" & + "\"platform\", \"organization\", \"device\", \"integration\"]" +) +doAssert model.contains("Attr(name: \"resource_key\", tabIndex: 1, typ: dbVarchar") + +block qualifiedTableNamesRemainDistinct: + const schemaText = """ +create table public.events (public_value text); +create table audit.events (audit_value text); +""" + let schema = schemaText + let generated = generateModelCode(schema, "qualified.sql", postgre) + doAssert generated.contains("\"public.events\"") + doAssert generated.contains("\"audit.events\"") + doAssert generated.contains("Attr(name: \"public_value\", tabIndex: 0") + doAssert generated.contains("Attr(name: \"audit_value\", tabIndex: 1") + +block enumKeywordsMayBeSeparatedByWhitespace: + const schemaText = """ +create +type public.mood as enum ('happy', 'sad'); +create table public.people (mood public.mood); +""" + let schema = schemaText + let generated = generateModelCode(schema, "enum_whitespace.sql", postgre) + doAssert generated.contains("Attr(name: \"mood\", tabIndex: 0, typ: dbEnum") + doAssert generated.contains( + "typeName: \"public.mood\", validValues: @[\"happy\", \"sad\"]" + ) + +block namedUniqueConstraintImports: + const schemaText = """ +create table public.people ( + email text, + constraint people_email_key unique (email) +); +""" + let schema = schemaText + let generated = generateModelCode(schema, "named_unique.sql", postgre) + doAssert generated.contains("Attr(name: \"email\", tabIndex: 0, typ: dbVarchar") diff --git a/tests/tqualified_schema_queries.nim b/tests/tqualified_schema_queries.nim new file mode 100644 index 0000000..663a18a --- /dev/null +++ b/tests/tqualified_schema_queries.nim @@ -0,0 +1,31 @@ +import std/assertions + +import ormin + +importModel(DbBackend.postgre, "qualified_schema_model", includeStatic = true) + +var db {.global.}: DbConn + +proc selectPublicEvents() = + discard query: + select public.events(public_value) + where public.events.public_value == "visible" + +proc selectAuditEvents() = + discard query: + select audit.events(audit_value) + +proc selectUsersByUnambiguousBaseName() = + discard query: + select users(username) + +proc joinQualifiedEvents() = + discard query: + select public.events(public_value) + join audit.events(audit_value) on public.events.public_value == audit.events.audit_value + +static: + doAssert not compiles(block: + discard query: + select events(public_value) + )