From ca45a6e359bd655e73fa52358a6b6c77d676d15c Mon Sep 17 00:00:00 2001 From: Matt Aitken Date: Fri, 12 Dec 2025 13:38:46 +0000 Subject: [PATCH] Casting --- .../tsql/src/query/parser.test.ts | 2 +- internal-packages/tsql/src/query/parser.ts | 497 ++++++++++++------ 2 files changed, 334 insertions(+), 165 deletions(-) diff --git a/internal-packages/tsql/src/query/parser.test.ts b/internal-packages/tsql/src/query/parser.test.ts index f8aa547e1..72cd54f8a 100644 --- a/internal-packages/tsql/src/query/parser.test.ts +++ b/internal-packages/tsql/src/query/parser.test.ts @@ -33,7 +33,7 @@ function parseAndConvert(input: string) { describe("TSQLParseTreeConverter", () => { describe("SELECT statements", () => { - it.only("should convert a simple SELECT statement", () => { + it("should convert a simple SELECT statement", () => { const ast = parseAndConvert("SELECT * FROM users"); expect(ast).toBeDefined(); diff --git a/internal-packages/tsql/src/query/parser.ts b/internal-packages/tsql/src/query/parser.ts index 9425d077d..80762083e 100644 --- a/internal-packages/tsql/src/query/parser.ts +++ b/internal-packages/tsql/src/query/parser.ts @@ -225,12 +225,107 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { this.start = start; } + // ───────────────────────────────────────────────────────────────────────────── + // Typed visit helpers - these provide type-safe wrappers around the generic visit() + // ───────────────────────────────────────────────────────────────────────────── + + /** Visit a column expression and return an Expression */ + private visitAsExpr(ctx: ParserRuleContext): Expression { + return this.visit(ctx) as Expression; + } + + /** Visit a join expression context */ + private visitJoin(ctx: ParserRuleContext): JoinExpr { + return this.visit(ctx) as JoinExpr; + } + + /** Visit an order expression list */ + private visitOrderList(ctx: ParserRuleContext): OrderExpr[] { + return this.visit(ctx) as OrderExpr[]; + } + + /** Visit a window expression */ + private visitWindow(ctx: ParserRuleContext): WindowExpr { + return this.visit(ctx) as WindowExpr; + } + + /** Visit a ratio expression */ + private visitRatio(ctx: ParserRuleContext): RatioExpr { + return this.visit(ctx) as RatioExpr; + } + + /** Visit a CTE list */ + private visitCTEs(ctx: ParserRuleContext): Record { + return this.visit(ctx) as Record; + } + + /** Visit an expression list */ + private visitExprList(ctx: ParserRuleContext): Expression[] { + return this.visit(ctx) as Expression[]; + } + + /** Visit a select statement */ + private visitSelectQuery(ctx: ParserRuleContext): SelectQuery | SelectSetQuery { + return this.visit(ctx) as SelectQuery | SelectSetQuery; + } + + /** Visit a placeholder */ + private visitPlaceholderExpr(ctx: ParserRuleContext): Placeholder { + return this.visit(ctx) as Placeholder; + } + + /** Visit and return a string (for identifiers, aliases) */ + private visitAsString(ctx: ParserRuleContext): string { + return this.visit(ctx) as string; + } + + /** Visit and return string array (for table identifiers, nested identifiers) */ + private visitStringArray(ctx: ParserRuleContext): string[] { + return this.visit(ctx) as string[]; + } + + /** Visit a limit by expression */ + private visitLimitBy(ctx: ParserRuleContext): LimitByExpr { + return this.visit(ctx) as LimitByExpr; + } + + /** Visit a join operator and return the type string */ + private visitJoinOpString(ctx: ParserRuleContext): string { + return this.visit(ctx) as string; + } + + /** Visit a join constraint */ + private visitConstraint(ctx: ParserRuleContext): JoinConstraint { + return this.visit(ctx) as JoinConstraint; + } + + /** Visit a sample expression */ + private visitSample(ctx: ParserRuleContext): SampleExpr { + return this.visit(ctx) as SampleExpr; + } + + /** Visit a window frame expression */ + private visitFrame(ctx: ParserRuleContext): WindowFrameExpr | [WindowFrameExpr, WindowFrameExpr] { + return this.visit(ctx) as WindowFrameExpr | [WindowFrameExpr, WindowFrameExpr]; + } + + /** Visit and return a Constant */ + private visitConst(ctx: ParserRuleContext): Constant { + return this.visit(ctx) as Constant; + } + + // ───────────────────────────────────────────────────────────────────────────── + // Core visitor methods + // ───────────────────────────────────────────────────────────────────────────── + visit(ctx: ParserRuleContext): ParseResult { const start = getTokenStart(ctx.start); const stop = getTokenStop(ctx.stop); const end = stop !== undefined ? stop + 1 : undefined; try { - const node = this.visitChildren(ctx); + // Use accept() for proper visitor dispatch - this calls the correct visitXxx method + // based on the runtime type of the context + const node = ctx.accept(this); // Only set position if node is a valid object and we have position info if ( node && @@ -513,18 +608,21 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { ); } const selectQuery = this.visitSelectStmtWithParens(subsequent.selectStmtWithParens()); + // SelectSetNode expects SelectQuery | SelectSetQuery, but visitSelectStmtWithParens may return Placeholder + // In practice, a Placeholder in this position would be substituted at runtime selectQueries.push({ - select_query: selectQuery, + select_query: selectQuery as SelectQuery | SelectSetQuery, set_operator: unionType, }); } if (selectQueries.length === 0) { - return initialQuery; + // initialQuery may be Placeholder but we cast since SelectSetQuery expects SelectQuery | SelectSetQuery + return initialQuery as SelectQuery | SelectSetQuery; } return { expression_type: "select_set_query", - initial_select_query: initialQuery, + initial_select_query: initialQuery as SelectQuery | SelectSetQuery, subsequent_select_queries: selectQueries, }; } @@ -562,36 +660,36 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { ctes: withClause ? this.visitWithClause(withClause) : undefined, select: columnExprList ? this.visitColumnExprList(columnExprList) : [], distinct: ctx.DISTINCT() ? true : undefined, - select_from: fromClause ? this.visit(fromClause) : undefined, - where: whereClause ? this.visit(whereClause) : undefined, - prewhere: prewhereClause ? this.visit(prewhereClause) : undefined, - having: havingClause ? this.visit(havingClause) : undefined, - group_by: groupByClause ? this.visit(groupByClause) : undefined, - order_by: orderByClause ? this.visit(orderByClause) : undefined, - limit_by: limitByClause ? this.visit(limitByClause) : undefined, + select_from: fromClause ? this.visitJoin(fromClause) : undefined, + where: whereClause ? this.visitAsExpr(whereClause) : undefined, + prewhere: prewhereClause ? this.visitAsExpr(prewhereClause) : undefined, + having: havingClause ? this.visitAsExpr(havingClause) : undefined, + group_by: groupByClause ? this.visitExprList(groupByClause) : undefined, + order_by: orderByClause ? this.visitOrderList(orderByClause) : undefined, + limit_by: limitByClause ? this.visitLimitBy(limitByClause) : undefined, }; const windowClause = ctx.windowClause(); if (windowClause) { selectQuery.window_exprs = {}; for (let index = 0; index < windowClause.windowExpr().length; index++) { - const name = this.visit(windowClause.identifier()[index]); - selectQuery.window_exprs![name] = this.visit(windowClause.windowExpr()[index]); + const name = this.visitAsString(windowClause.identifier()[index]); + selectQuery.window_exprs![name] = this.visitWindow(windowClause.windowExpr()[index]); } } const limitAndOffsetClause = ctx.limitAndOffsetClause(); const offsetOnlyClause = ctx.offsetOnlyClause(); if (limitAndOffsetClause) { - selectQuery.limit = this.visit(limitAndOffsetClause.columnExpr(0)); + selectQuery.limit = this.visitAsExpr(limitAndOffsetClause.columnExpr(0)); if (limitAndOffsetClause.columnExpr(1)) { - selectQuery.offset = this.visit(limitAndOffsetClause.columnExpr(1)); + selectQuery.offset = this.visitAsExpr(limitAndOffsetClause.columnExpr(1)); } if (limitAndOffsetClause.WITH() && limitAndOffsetClause.TIES()) { selectQuery.limit_with_ties = true; } } else if (offsetOnlyClause) { - selectQuery.offset = this.visit(offsetOnlyClause.columnExpr()); + selectQuery.offset = this.visitAsExpr(offsetOnlyClause.columnExpr()); } const arrayJoinClause = ctx.arrayJoinClause(); @@ -606,7 +704,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { } else { selectQuery.array_join_op = "ARRAY JOIN"; } - selectQuery.array_join_list = this.visit(arrayJoinClause.columnExprList()); + selectQuery.array_join_list = this.visitExprList(arrayJoinClause.columnExprList()); if (selectQuery.array_join_list) { for (const expr of selectQuery.array_join_list) { if (!("alias" in expr)) { @@ -630,19 +728,19 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { } visitWithClause(ctx: WithClauseContext): Record { - return this.visit(ctx.withExprList()); + return this.visitCTEs(ctx.withExprList()); } visitFromClause(ctx: FromClauseContext): JoinExpr { - return this.visit(ctx.joinExpr()); + return this.visitJoin(ctx.joinExpr()); } visitPrewhereClause(ctx: PrewhereClauseContext): Expr { - return this.visit(ctx.columnExpr()); + return this.visitAsExpr(ctx.columnExpr()); } visitWhereClause(ctx: WhereClauseContext): Expr { - return this.visit(ctx.columnExpr()); + return this.visitAsExpr(ctx.columnExpr()); } visitGroupByClause(ctx: GroupByClauseContext): Expr[] { @@ -650,46 +748,51 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { if (!columnExprList) { throw new SyntaxError("GROUP BY clause must have a column expression list"); } - return this.visit(columnExprList); + return this.visitExprList(columnExprList); } visitHavingClause(ctx: HavingClauseContext): Expr { - return this.visit(ctx.columnExpr()); + return this.visitAsExpr(ctx.columnExpr()); } visitOrderByClause(ctx: OrderByClauseContext): OrderExpr[] { - return this.visit(ctx.orderExprList()); + return this.visitOrderList(ctx.orderExprList()); } visitLimitByClause(ctx: LimitByClauseContext): LimitByExpr { - const limitExpr = this.visit(ctx.limitExpr()); + const limitExpr = this.visitLimitExprResult(ctx.limitExpr()); // If limitExpr is a tuple (n, offset), split it - if (Array.isArray(limitExpr) && limitExpr.length === 2) { + if (Array.isArray(limitExpr)) { const [n, offsetValue] = limitExpr; return { expression_type: "limit_by_expr", n, offset_value: offsetValue, - exprs: this.visit(ctx.columnExprList()), + exprs: this.visitExprList(ctx.columnExprList()), }; } - // If no offset, just use limitExpr as n + // If no offset, just use limitExpr as n (TypeScript now knows it's Expression) return { expression_type: "limit_by_expr", n: limitExpr, offset_value: undefined, - exprs: this.visit(ctx.columnExprList()), + exprs: this.visitExprList(ctx.columnExprList()), }; } - visitLimitExpr(ctx: LimitExprContext): Expr | [Expr, Expr] { - const n = this.visit(ctx.columnExpr(0)); + /** Helper for visitLimitExpr which returns Expression | [Expression, Expression] */ + private visitLimitExprResult(ctx: LimitExprContext): Expression | [Expression, Expression] { + return this.visitLimitExpr(ctx); + } + + visitLimitExpr(ctx: LimitExprContext): Expression | [Expression, Expression] { + const n = this.visitAsExpr(ctx.columnExpr(0)); // Check if we have an offset (second expression) if (ctx.columnExpr(1)) { - const offsetValue = this.visit(ctx.columnExpr(1)); + const offsetValue = this.visitAsExpr(ctx.columnExpr(1)); // For "LIMIT a, b" syntax: a is offset, b is limit if (ctx.COMMA()) { return [offsetValue, n]; // Return tuple as (offset, limit) @@ -703,16 +806,16 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { // JOIN expressions visitJoinExprOp(ctx: JoinExprOpContext): JoinExpr { - const join1: JoinExpr = this.visit(ctx.joinExpr(0)); - const join2: JoinExpr = this.visit(ctx.joinExpr(1)); + const join1: JoinExpr = this.visitJoin(ctx.joinExpr(0)); + const join2: JoinExpr = this.visitJoin(ctx.joinExpr(1)); const joinOp = ctx.joinOp(); if (joinOp) { - join2.join_type = `${this.visit(joinOp)} JOIN`; + join2.join_type = `${this.visitJoinOpString(joinOp)} JOIN`; } else { join2.join_type = "JOIN"; } - join2.constraint = this.visit(ctx.joinConstraintClause()); + join2.constraint = this.visitConstraint(ctx.joinConstraintClause()); let lastJoin = join1; while (lastJoin.next_join) { @@ -725,16 +828,19 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitJoinExprTable(ctx: JoinExprTableContext): JoinExpr { const sampleClause = ctx.sampleClause(); - const sample = sampleClause ? this.visit(sampleClause) : undefined; - const table = this.visit(ctx.tableExpr()); + const sample = sampleClause ? this.visitSample(sampleClause) : undefined; + const tableResult = this.visitTableExprResult(ctx.tableExpr()); const tableFinal = ctx.FINAL() ? true : undefined; - if ("table" in table) { + // Check if result is already a JoinExpr (from visitTableExprAlias or visitTableExprFunction) + if ("expression_type" in tableResult && tableResult.expression_type === "join_expr") { // visitTableExprAlias returns a JoinExpr to pass the alias // visitTableExprFunction returns a JoinExpr to pass the args - table.table_final = tableFinal; - table.sample = sample; - return table; + tableResult.table_final = tableFinal; + tableResult.sample = sample; + return tableResult; } + // Otherwise, wrap the table expression in a JoinExpr + const table = tableResult as SelectQuery | SelectSetQuery | Placeholder | HogQLXTag | Field; return { expression_type: "join_expr", table, @@ -743,13 +849,26 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { }; } + /** Helper for visiting table expressions that may return JoinExpr or table types */ + private visitTableExprResult( + ctx: ParserRuleContext + ): JoinExpr | SelectQuery | SelectSetQuery | Placeholder | HogQLXTag | Field { + return this.visit(ctx) as + | JoinExpr + | SelectQuery + | SelectSetQuery + | Placeholder + | HogQLXTag + | Field; + } + visitJoinExprParens(ctx: JoinExprParensContext): JoinExpr { - return this.visit(ctx.joinExpr()); + return this.visitJoin(ctx.joinExpr()); } visitJoinExprCrossOp(ctx: JoinExprCrossOpContext): JoinExpr { - const join1: JoinExpr = this.visit(ctx.joinExpr(0)); - const join2: JoinExpr = this.visit(ctx.joinExpr(1)); + const join1: JoinExpr = this.visitJoin(ctx.joinExpr(0)); + const join2: JoinExpr = this.visitJoin(ctx.joinExpr(1)); join2.join_type = "CROSS JOIN"; let lastJoin = join1; while (lastJoin.next_join) { @@ -791,7 +910,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { } visitJoinConstraintClause(ctx: JoinConstraintClauseContext): JoinConstraint { - const columnExprList = this.visit(ctx.columnExprList()); + const columnExprList = this.visitExprList(ctx.columnExprList()); if (columnExprList.length !== 1) { throw new NotImplementedError("Unsupported: JOIN ... ON with multiple expressions"); } @@ -804,9 +923,11 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitSampleClause(ctx: SampleClauseContext): SampleExpr { const ratioExpressions = ctx.ratioExpr(); - const sampleRatioExpr = this.visit(ratioExpressions[0]); + const sampleRatioExpr = this.visitRatio(ratioExpressions[0]); const offsetRatioExpr = - ratioExpressions.length > 1 && ctx.OFFSET() ? this.visit(ratioExpressions[1]) : undefined; + ratioExpressions.length > 1 && ctx.OFFSET() + ? this.visitRatio(ratioExpressions[1]) + : undefined; return { expression_type: "sample_expr", @@ -816,14 +937,14 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { } visitOrderExprList(ctx: OrderExprListContext): OrderExpr[] { - return ctx.orderExpr().map((expr: any) => this.visit(expr)); + return ctx.orderExpr().map((expr: OrderExprContext) => this.visitOrderExpr(expr)); } visitOrderExpr(ctx: OrderExprContext): OrderExpr { const order = ctx.DESC() || ctx.DESCENDING() ? "DESC" : "ASC"; return { expression_type: "order_expr", - expr: this.visit(ctx.columnExpr()), + expr: this.visitAsExpr(ctx.columnExpr()), order: order as "ASC" | "DESC", }; } @@ -831,7 +952,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitRatioExpr(ctx: RatioExprContext): RatioExpr { const placeholder = ctx.placeholder(); if (placeholder) { - return this.visit(placeholder); + return this.visitPlaceholderExpr(placeholder) as unknown as RatioExpr; } const numberLiterals = ctx.numberLiteral(); @@ -847,13 +968,13 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitWindowExpr(ctx: WindowExprContext): WindowExpr { const frame = ctx.winFrameClause(); - const visitedFrame = frame ? this.visit(frame) : undefined; + const visitedFrame = frame ? this.visitFrame(frame) : undefined; const partitionByClause = ctx.winPartitionByClause(); const orderByClause = ctx.winOrderByClause(); return { expression_type: "window_expr", - partition_by: partitionByClause ? this.visit(partitionByClause) : undefined, - order_by: orderByClause ? this.visit(orderByClause) : undefined, + partition_by: partitionByClause ? this.visitExprList(partitionByClause) : undefined, + order_by: orderByClause ? this.visitOrderList(orderByClause) : undefined, frame_method: frame && frame.RANGE() ? "RANGE" : frame && frame.ROWS() ? "ROWS" : undefined, frame_start: Array.isArray(visitedFrame) ? visitedFrame[0] : visitedFrame, frame_end: Array.isArray(visitedFrame) ? visitedFrame[1] : undefined, @@ -861,25 +982,30 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { } visitWinPartitionByClause(ctx: WinPartitionByClauseContext): Expr[] { - return this.visit(ctx.columnExprList()); + return this.visitExprList(ctx.columnExprList()); } visitWinOrderByClause(ctx: WinOrderByClauseContext): OrderExpr[] { - return this.visit(ctx.orderExprList()); + return this.visitOrderList(ctx.orderExprList()); } visitWinFrameClause( ctx: WinFrameClauseContext ): WindowFrameExpr | [WindowFrameExpr, WindowFrameExpr] { - return this.visit(ctx.winFrameExtend()); + return this.visitFrame(ctx.winFrameExtend()); } visitFrameStart(ctx: FrameStartContext): WindowFrameExpr { - return this.visit(ctx.winFrameBound()); + return this.visitFrameBound(ctx.winFrameBound()); } visitFrameBetween(ctx: FrameBetweenContext): [WindowFrameExpr, WindowFrameExpr] { - return [this.visit(ctx.winFrameBound(0)), this.visit(ctx.winFrameBound(1))]; + return [this.visitFrameBound(ctx.winFrameBound(0)), this.visitFrameBound(ctx.winFrameBound(1))]; + } + + /** Helper for visiting a single frame bound */ + private visitFrameBound(ctx: WinFrameBoundContext): WindowFrameExpr { + return this.visitWinFrameBound(ctx); } visitWinFrameBound(ctx: WinFrameBoundContext): WindowFrameExpr { @@ -888,7 +1014,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { return { expression_type: "window_frame_expr", frame_type: "PRECEDING", - frame_value: numberLiteral ? (this.visit(numberLiteral) as Constant).value : undefined, + frame_value: numberLiteral ? this.visitNumberLiteral(numberLiteral).value : undefined, }; } if (ctx.FOLLOWING()) { @@ -896,7 +1022,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { return { expression_type: "window_frame_expr", frame_type: "FOLLOWING", - frame_value: numberLiteral ? (this.visit(numberLiteral) as Constant).value : undefined, + frame_value: numberLiteral ? this.visitNumberLiteral(numberLiteral).value : undefined, }; } return { expression_type: "window_frame_expr", frame_type: "CURRENT ROW" }; @@ -904,7 +1030,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { // Column expressions visitColumnExprList(ctx: ColumnExprListContext): Expression[] { - return ctx.columnExpr().map((c: any) => this.visitWithExprColumn(c)); + return ctx.columnExpr().map((c: ParserRuleContext) => this.visitAsExpr(c)); } visitColumnExprTernaryOp(ctx: ColumnExprTernaryOpContext): Call { @@ -912,9 +1038,9 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { expression_type: "call", name: "if", args: [ - this.visit(ctx.columnExpr(0)), - this.visit(ctx.columnExpr(1)), - this.visit(ctx.columnExpr(2)), + this.visitAsExpr(ctx.columnExpr(0)), + this.visitAsExpr(ctx.columnExpr(1)), + this.visitAsExpr(ctx.columnExpr(2)), ], }; } @@ -930,7 +1056,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { } else { throw new SyntaxError("Must specify an alias"); } - const expr = this.visit(ctx.columnExpr()); + const expr = this.visitAsExpr(ctx.columnExpr()); if (RESERVED_KEYWORDS.includes(alias.toLowerCase() as any)) { throw new SyntaxError( @@ -946,7 +1072,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { expression_type: "arithmetic_operation", op: ArithmeticOperationOp.Sub, left: { value: 0 } as Constant, - right: this.visit(ctx.columnExpr()), + right: this.visitAsExpr(ctx.columnExpr()), }; } @@ -954,12 +1080,12 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { const kvPairList = ctx.kvPairList(); return { expression_type: "dict", - items: kvPairList ? this.visit(kvPairList) : [], + items: kvPairList ? this.visitKvPairList(kvPairList) : [], }; } visitColumnExprSubquery(ctx: ColumnExprSubqueryContext): SelectQuery | SelectSetQuery { - return this.visit(ctx.selectSetStmt()); + return this.visitSelectQuery(ctx.selectSetStmt()); } visitColumnExprLiteral(ctx: ColumnExprLiteralContext): Expr { @@ -970,7 +1096,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { const columnExprList = ctx.columnExprList(); return { expression_type: "array", - exprs: columnExprList ? this.visit(columnExprList) : [], + exprs: columnExprList ? this.visitExprList(columnExprList) : [], }; } @@ -986,15 +1112,15 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { throw new NotImplementedError(`Unsupported ColumnExprPrecedence1: ${ctx.text}`); } // Use columnExpr() method to get left and right operands - const left = this.visit(ctx.columnExpr(0)); - const right = this.visit(ctx.columnExpr(1)); + const left = this.visitAsExpr(ctx.columnExpr(0)); + const right = this.visitAsExpr(ctx.columnExpr(1)); return { expression_type: "arithmetic_operation", left, right, op }; } visitColumnExprPrecedence2(ctx: ColumnExprPrecedence2Context): ArithmeticOperation | Call { // Use columnExpr() method to get left and right operands - const left = this.visit(ctx.columnExpr(0)); - const right = this.visit(ctx.columnExpr(1)); + const left = this.visitAsExpr(ctx.columnExpr(0)); + const right = this.visitAsExpr(ctx.columnExpr(1)); if (ctx.PLUS()) { return { @@ -1012,13 +1138,13 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { }; } else if (ctx.CONCAT()) { const args: Expression[] = []; - if ("name" in left && left.name === "concat" && "args" in left) { + if ("name" in left && left.name === "concat" && "args" in left && left.args) { args.push(...left.args); } else { args.push(left); } - if ("name" in right && right.name === "concat" && "args" in right) { + if ("name" in right && right.name === "concat" && "args" in right && right.args) { args.push(...right.args); } else { args.push(right); @@ -1032,8 +1158,8 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitColumnExprPrecedence3(ctx: ColumnExprPrecedence3Context): CompareOperation { // Use columnExpr() method to get left and right operands - const left = this.visit(ctx.columnExpr(0)); - const right = this.visit(ctx.columnExpr(1)); + const left = this.visitAsExpr(ctx.columnExpr(0)); + const right = this.visitAsExpr(ctx.columnExpr(1)); let op: CompareOperationOp; if (ctx.EQ_SINGLE() || ctx.EQ_DOUBLE()) { @@ -1096,110 +1222,119 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { throw new NotImplementedError(`Unsupported interval type: ${interval.text}`); } - return { expression_type: "call", name, args: [this.visit(ctx.columnExpr())] }; + return { expression_type: "call", name, args: [this.visitAsExpr(ctx.columnExpr())] }; } visitColumnExprIsNull(ctx: ColumnExprIsNullContext): CompareOperation { return { expression_type: "compare_operation", - left: this.visit(ctx.columnExpr()), + left: this.visitAsExpr(ctx.columnExpr()), right: { value: null } as Constant, op: ctx.NOT() ? CompareOperationOp.NotEq : CompareOperationOp.Eq, }; } visitColumnExprTuple(ctx: ColumnExprTupleContext): Tuple { + const columnExprList = ctx.columnExprList(); return { expression_type: "tuple", - exprs: ctx.columnExprList() ? this.visit(ctx.columnExprList()) : [], + exprs: columnExprList ? this.visitExprList(columnExprList) : [], }; } visitColumnExprArrayAccess(ctx: ColumnExprArrayAccessContext): ArrayAccess { - const object: Expression = this.visit(ctx.columnExpr(0)); - const property: Expression = this.visit(ctx.columnExpr(1)); + const object: Expression = this.visitAsExpr(ctx.columnExpr(0)); + const property: Expression = this.visitAsExpr(ctx.columnExpr(1)); return { expression_type: "array_access", array: object, property }; } visitColumnExprNullArrayAccess(ctx: ColumnExprNullArrayAccessContext): ArrayAccess { - const object: Expression = this.visit(ctx.columnExpr(0)); - const property: Expression = this.visit(ctx.columnExpr(1)); + const object: Expression = this.visitAsExpr(ctx.columnExpr(0)); + const property: Expression = this.visitAsExpr(ctx.columnExpr(1)); return { expression_type: "array_access", array: object, property, nullish: true }; } visitColumnExprPropertyAccess(ctx: ColumnExprPropertyAccessContext): ArrayAccess { - const object = this.visit(ctx.columnExpr()); - const property = { value: this.visitIdentifier(ctx.identifier()) } as Constant; + const object = this.visitAsExpr(ctx.columnExpr()); + const property: Constant = { + expression_type: "constant", + value: this.visitIdentifier(ctx.identifier()), + }; return { expression_type: "array_access", array: object, property }; } visitColumnExprNullPropertyAccess(ctx: ColumnExprNullPropertyAccessContext): ArrayAccess { - const object = this.visit(ctx.columnExpr()); - const property = { value: this.visitIdentifier(ctx.identifier()) } as Constant; + const object = this.visitAsExpr(ctx.columnExpr()); + const property: Constant = { + expression_type: "constant", + value: this.visitIdentifier(ctx.identifier()), + }; return { expression_type: "array_access", array: object, property, nullish: true }; } visitColumnExprBetween(ctx: ColumnExprBetweenContext): BetweenExpr { return { expression_type: "between_expr", - expr: this.visit(ctx.columnExpr(0)), - low: this.visit(ctx.columnExpr(1)), - high: this.visit(ctx.columnExpr(2)), + expr: this.visitAsExpr(ctx.columnExpr(0)), + low: this.visitAsExpr(ctx.columnExpr(1)), + high: this.visitAsExpr(ctx.columnExpr(2)), negated: !!ctx.NOT(), }; } visitColumnExprParens(ctx: ColumnExprParensContext): Expr { - return this.visit(ctx.columnExpr()); + return this.visitAsExpr(ctx.columnExpr()); } visitColumnExprAnd(ctx: ColumnExprAndContext): And { - let left = this.visit(ctx.columnExpr(0)); - const leftArray = "exprs" in left ? left.exprs : [left]; + const left = this.visitAsExpr(ctx.columnExpr(0)); + const leftArray = "exprs" in left && left.exprs ? left.exprs : [left]; - let right = this.visit(ctx.columnExpr(1)); - const rightArray = "exprs" in right ? right.exprs : [right]; + const right = this.visitAsExpr(ctx.columnExpr(1)); + const rightArray = "exprs" in right && right.exprs ? right.exprs : [right]; return { expression_type: "and", exprs: [...leftArray, ...rightArray] }; } visitColumnExprOr(ctx: ColumnExprOrContext): Or { - let left = this.visit(ctx.columnExpr(0)); - const leftArray = "exprs" in left ? left.exprs : [left]; + const left = this.visitAsExpr(ctx.columnExpr(0)); + const leftArray = "exprs" in left && left.exprs ? left.exprs : [left]; - let right = this.visit(ctx.columnExpr(1)); - const rightArray = "exprs" in right ? right.exprs : [right]; + const right = this.visitAsExpr(ctx.columnExpr(1)); + const rightArray = "exprs" in right && right.exprs ? right.exprs : [right]; return { expression_type: "or", exprs: [...leftArray, ...rightArray] }; } visitColumnExprTupleAccess(ctx: ColumnExprTupleAccessContext): TupleAccess { - const tuple = this.visit(ctx.columnExpr()); + const tuple = this.visitAsExpr(ctx.columnExpr()); const index = parseInt(ctx.DECIMAL_LITERAL().text); return { expression_type: "tuple_access", tuple, index }; } visitColumnExprNullTupleAccess(ctx: ColumnExprNullTupleAccessContext): TupleAccess { - const tuple = this.visit(ctx.columnExpr()); + const tuple = this.visitAsExpr(ctx.columnExpr()); const index = parseInt(ctx.DECIMAL_LITERAL().text); return { expression_type: "tuple_access", tuple, index, nullish: true }; } visitColumnExprCase(ctx: ColumnExprCaseContext): Call { - const columns = ctx.columnExpr().map((column: any) => this.visit(column)); + const columns = ctx + .columnExpr() + .map((column: ParserRuleContext) => this.visitAsExpr(column) as Expression); if (ctx._caseExpr) { - const args: Expression[] = [ + const args: (Expression | undefined)[] = [ columns[0], { expression_type: "array", exprs: [] }, { expression_type: "array", exprs: [] }, - , + undefined, columns[columns.length - 1], ]; for (let index = 1; index < columns.length - 1; index++) { const arrayIndex = ((index - 1) % 2) + 1; (args[arrayIndex] as ArrayExpression).exprs.push(columns[index]); } - return { expression_type: "call", name: "transform", args }; + return { expression_type: "call", name: "transform", args: args as Expression[] }; } else if (columns.length === 3) { return { expression_type: "call", name: "if", args: columns }; } else { @@ -1208,45 +1343,46 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { } visitColumnExprNot(ctx: ColumnExprNotContext): Not { - return { expression_type: "not", expr: this.visit(ctx.columnExpr()) }; + return { expression_type: "not", expr: this.visitAsExpr(ctx.columnExpr()) }; } visitColumnExprWinFunctionTarget(ctx: ColumnExprWinFunctionTargetContext): WindowFunction { return { expression_type: "window_function", name: this.visitIdentifier(ctx.identifier(0)), - exprs: ctx._columnExprs ? this.visit(ctx._columnExprs) : [], - args: ctx._columnArgList ? this.visit(ctx._columnArgList) : [], + exprs: ctx._columnExprs ? this.visitExprList(ctx._columnExprs) : [], + args: ctx._columnArgList ? this.visitExprList(ctx._columnArgList) : [], over_identifier: this.visitIdentifier(ctx.identifier(1)), }; } visitColumnExprWinFunction(ctx: ColumnExprWinFunctionContext): WindowFunction { + const windowExpr = ctx.windowExpr(); return { expression_type: "window_function", name: this.visitIdentifier(ctx.identifier()), - exprs: ctx._columnExprs ? this.visit(ctx._columnExprs) : [], - args: ctx._columnArgList ? this.visit(ctx._columnArgList) : [], - over_expr: ctx.windowExpr() ? this.visit(ctx.windowExpr()) : undefined, + exprs: ctx._columnExprs ? this.visitExprList(ctx._columnExprs) : [], + args: ctx._columnArgList ? this.visitExprList(ctx._columnArgList) : [], + over_expr: windowExpr ? this.visitWindowExpr(windowExpr) : undefined, }; } visitColumnExprIdentifier(ctx: ColumnExprIdentifierContext): Expr { - return this.visit(ctx.columnIdentifier()); + return this.visitColumnIdentifier(ctx.columnIdentifier()); } visitColumnExprFunction(ctx: ColumnExprFunctionContext): Call { const name = this.visitIdentifier(ctx.identifier()); let parameters: Expression[] | undefined = ctx._columnExprs - ? this.visit(ctx._columnExprs) + ? this.visitExprList(ctx._columnExprs) : undefined; // two sets of parameters fn()(), return an empty list for the first even if no parameters if (ctx.LPAREN && ctx.LPAREN().length > 1 && parameters === undefined) { parameters = []; } - const args: Expression[] = ctx._columnArgList ? this.visit(ctx._columnArgList) : []; + const args: Expression[] = ctx._columnArgList ? this.visitExprList(ctx._columnArgList) : []; const distinct = ctx.DISTINCT() ? true : false; return { expression_type: "call", name, params: parameters, args, distinct }; } @@ -1254,7 +1390,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitColumnExprAsterisk(ctx: ColumnExprAsteriskContext): Field { const tableIdentifier = ctx.tableIdentifier(); if (tableIdentifier) { - const table = this.visit(tableIdentifier); + const table = this.visitStringArray(tableIdentifier); return { expression_type: "field", chain: [...table, "*"] }; } return { expression_type: "field", chain: ["*"] }; @@ -1263,30 +1399,45 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitColumnLambdaExpr(ctx: ColumnLambdaExprContext): Lambda { const columnExpr = ctx.columnExpr(); const block = ctx.block(); + let expr: Expression | Block; + if (columnExpr) { + expr = this.visitAsExpr(columnExpr); + } else if (block) { + expr = this.visitBlock(block) as Block; + } else { + throw new SyntaxError("Lambda expression must have either an expression or a block"); + } return { expression_type: "lambda", - args: ctx.identifier().map((identifier: any) => this.visitIdentifier(identifier)), - expr: columnExpr ? this.visit(columnExpr) : block ? this.visit(block) : undefined, + args: ctx + .identifier() + .map((identifier: IdentifierContext) => this.visitIdentifier(identifier)), + expr, }; } visitWithExprList(ctx: WithExprListContext): Record { const ctes: Record = {}; for (const expr of ctx.withExpr()) { - const cte = this.visit(expr); + const cte = this.visitCTE(expr); ctes[cte.name] = cte; } return ctes; } + /** Helper to visit a CTE expression */ + private visitCTE(ctx: ParserRuleContext): CTE { + return this.visit(ctx) as CTE; + } + visitWithExprSubquery(ctx: WithExprSubqueryContext): CTE { - const subquery = this.visit(ctx.selectSetStmt()); + const subquery = this.visitSelectQuery(ctx.selectSetStmt()); const name = this.visitIdentifier(ctx.identifier()); return { expression_type: "cte", name, expr: subquery, cte_type: "subquery" }; } visitWithExprColumn(ctx: WithExprColumnContext): CTE { - const expr = this.visitColumnExpr(ctx.columnExpr()); + const expr = this.visitAsExpr(ctx.columnExpr()); const name = this.visitIdentifier(ctx.identifier()); return { expression_type: "cte", name, expr, cte_type: "column" }; } @@ -1294,14 +1445,14 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitColumnIdentifier(ctx: ColumnIdentifierContext): Expression { const placeholder = ctx.placeholder(); if (placeholder) { - return this.visit(placeholder); + return this.visitPlaceholder(placeholder); } const tableIdentifier = ctx.tableIdentifier(); - const table = tableIdentifier ? this.visit(tableIdentifier) : []; + const table = tableIdentifier ? this.visitTableIdentifier(tableIdentifier) : []; const nestedIdentifier = ctx.nestedIdentifier(); - const nested = nestedIdentifier ? this.visit(nestedIdentifier) : []; + const nested = nestedIdentifier ? this.visitNestedIdentifier(nestedIdentifier) : []; if (table.length === 0 && nested.length > 0) { const text = ctx.text.toLowerCase(); @@ -1318,20 +1469,22 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { } visitNestedIdentifier(ctx: NestedIdentifierContext): string[] { - return ctx.identifier().map((identifier: any) => this.visitIdentifier(identifier)); + return ctx + .identifier() + .map((identifier: IdentifierContext) => this.visitIdentifier(identifier)); } visitTableExprIdentifier(ctx: TableExprIdentifierContext): Field { - const chain = this.visit(ctx.tableIdentifier()); + const chain = this.visitTableIdentifier(ctx.tableIdentifier()); return { expression_type: "field", chain }; } visitTableExprSubquery(ctx: TableExprSubqueryContext): SelectQuery | SelectSetQuery { - return this.visit(ctx.selectSetStmt()); + return this.visitSelectQuery(ctx.selectSetStmt()); } visitTableExprPlaceholder(ctx: TableExprPlaceholderContext): Placeholder { - return this.visit(ctx.placeholder()); + return this.visitPlaceholder(ctx.placeholder()); } visitTableExprAlias(ctx: TableExprAliasContext): JoinExpr { @@ -1339,32 +1492,38 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { if (!exp) { throw new SyntaxError("Must specify an alias"); } - const alias: string = this.visit(exp); + const alias: string = + exp instanceof AliasContext + ? this.visitAlias(exp) + : this.visitIdentifier(exp as IdentifierContext); if (RESERVED_KEYWORDS.includes(alias.toLowerCase() as any)) { throw new SyntaxError( `"${alias}" cannot be an alias or identifier, as it's a reserved keyword` ); } - const table = this.visit(ctx.tableExpr()); - if ("table" in table) { - table.alias = alias; - return table; + const tableResult = this.visitTableExprResult(ctx.tableExpr()); + // Check if result is already a JoinExpr + if ("expression_type" in tableResult && tableResult.expression_type === "join_expr") { + tableResult.alias = alias; + return tableResult; } + // Otherwise, wrap in a JoinExpr + const table = tableResult as SelectQuery | SelectSetQuery | Placeholder | HogQLXTag | Field; return { expression_type: "join_expr", table, alias }; } visitTableExprFunction(ctx: TableExprFunctionContext): JoinExpr { - return this.visit(ctx.tableFunctionExpr()); + return this.visitTableFunctionExpr(ctx.tableFunctionExpr()); } visitTableExprTag(ctx: TableExprTagContext): HogQLXTag { - return this.visit(ctx.tSQLxTagElement()); + return this.visitHogqlxTagElementNested(ctx.tSQLxTagElement()); } visitTableFunctionExpr(ctx: TableFunctionExprContext): JoinExpr { const name = this.visitIdentifier(ctx.identifier()); const tableArgList = ctx.tableArgList(); - const args = tableArgList ? this.visit(tableArgList) : []; + const args = tableArgList ? this.visitTableArgList(tableArgList) : []; return { expression_type: "join_expr", table: { expression_type: "field", chain: [name] }, @@ -1373,13 +1532,14 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { } visitTableIdentifier(ctx: TableIdentifierContext): string[] { - const nested = ctx.nestedIdentifier() ? this.visit(ctx.nestedIdentifier()) : []; + const nestedIdentifier = ctx.nestedIdentifier(); + const nested = nestedIdentifier ? this.visitNestedIdentifier(nestedIdentifier) : []; // Ensure nested is always an array const nestedArray = Array.isArray(nested) ? nested : nested ? [nested] : []; const databaseIdentifier = ctx.databaseIdentifier(); if (databaseIdentifier) { - const dbId = this.visit(databaseIdentifier); + const dbId = this.visitDatabaseIdentifier(databaseIdentifier); return [dbId, ...nestedArray]; } @@ -1387,7 +1547,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { } visitTableArgList(ctx: TableArgListContext): Expression[] { - return ctx.columnExpr().map((arg: any) => this.visit(arg)); + return ctx.columnExpr().map((arg: ParserRuleContext) => this.visitAsExpr(arg)); } visitDatabaseIdentifier(ctx: DatabaseIdentifierContext): string { @@ -1477,7 +1637,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { return { expression_type: "call", name: "ifNull", - args: [this.visit(ctx.columnExpr(0)), this.visit(ctx.columnExpr(1))], + args: [this.visitAsExpr(ctx.columnExpr(0)), this.visitAsExpr(ctx.columnExpr(1))], }; } @@ -1485,36 +1645,36 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { const columnExprList = ctx.columnExprList(); return { expression_type: "expr_call", - expr: this.visit(ctx.columnExpr()), - args: columnExprList ? this.visit(columnExprList) : [], + expr: this.visitAsExpr(ctx.columnExpr()), + args: columnExprList ? this.visitExprList(columnExprList) : [], }; } visitColumnExprCallSelect(ctx: ColumnExprCallSelectContext): Call | ExprCall { - const expr = this.visit(ctx.columnExpr()); - if ("chain" in expr && expr.chain.length === 1) { + const expr = this.visitAsExpr(ctx.columnExpr()); + if ("chain" in expr && expr.chain && expr.chain.length === 1) { return { expression_type: "call", name: String(expr.chain[0]), - args: [this.visit(ctx.selectSetStmt())], + args: [this.visitSelectQuery(ctx.selectSetStmt())], }; } return { expression_type: "expr_call", expr, - args: [this.visit(ctx.selectSetStmt())], + args: [this.visitSelectQuery(ctx.selectSetStmt())], }; } - visitHogqlxChildElement(ctx: TSQLxChildElementContext): Expr | HogQLXTag { + visitHogqlxChildElement(ctx: TSQLxChildElementContext): Expression { const tSQLxTagElement = ctx.tSQLxTagElement(); if (tSQLxTagElement) { - return this.visit(tSQLxTagElement); + return this.visitHogqlxTagElementNested(tSQLxTagElement); } if (ctx.TSQLX_TEXT_TEXT()) { return this.visitHogqlxText(ctx); } - return this.visit(ctx.columnExpr()!); + return this.visitAsExpr(ctx.columnExpr()!); } visitHogqlxText(ctx: TSQLxChildElementContext): Constant { @@ -1525,7 +1685,9 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitHogqlxTagElementClosed(ctx: TSQLxTagElementContext): HogQLXTag { const kind = this.visitIdentifier(ctx.identifier()[0]); const attributes = ctx.tSQLxTagAttribute() - ? ctx.tSQLxTagAttribute().map((a: any) => this.visit(a)) + ? ctx + .tSQLxTagAttribute() + .map((a: TSQLxTagAttributeContext) => this.visitHogqlxTagAttribute(a)) : []; return { expression_type: "hogqlx_tag", kind, attributes }; } @@ -1540,13 +1702,15 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { } const attributes = ctx.tSQLxTagAttribute() - ? ctx.tSQLxTagAttribute().map((a: any) => this.visit(a)) + ? ctx + .tSQLxTagAttribute() + .map((a: TSQLxTagAttributeContext) => this.visitHogqlxTagAttribute(a)) : []; // ── collect child nodes, discarding pure-indentation whitespace ── const keptChildren: Expression[] = []; for (const element of ctx.tSQLxChildElement()) { - const child = this.visit(element); + const child = this.visitHogqlxChildElement(element); if ("value" in child && typeof child.value === "string") { const v = child.value; @@ -1577,30 +1741,35 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { const columnExpr = ctx.columnExpr(); const string = ctx.string(); if (columnExpr) { - return { name, value: this.visit(columnExpr) }; + return { name, value: this.visitAsExpr(columnExpr) }; } else if (string) { - return { name, value: this.visit(string) }; + return { name, value: this.visitStringExpr(string) }; } else { - return { name, value: { value: true } as Constant }; + return { name, value: { expression_type: "constant", value: true } as Constant }; } } visitPlaceholder(ctx: PlaceholderContext): Placeholder { - return { expression_type: "placeholder", expr: this.visitCol(ctx.columnExpr()) }; + return { expression_type: "placeholder", expr: this.visitAsExpr(ctx.columnExpr()) }; } visitColumnExprTemplateString(ctx: ColumnExprTemplateStringContext): Expr { - return this.visit(ctx.templateString()); + return this.visitTemplateString(ctx.templateString()); } - visitString(ctx: StringContext): Constant | Expr { + /** Helper to visit a string context that returns Constant | Expr */ + private visitStringExpr(ctx: StringContext): Constant | Expr { + return this.visitStringLiteral(ctx); + } + + visitStringLiteral(ctx: StringContext): Constant | Expr { const stringLiteral = ctx.STRING_LITERAL(); if (stringLiteral) { return { expression_type: "constant", value: parseStringLiteralText(stringLiteral.text) }; } const templateString = ctx.templateString(); if (templateString) { - return this.visit(templateString); + return this.visitTemplateString(templateString); } return { expression_type: "constant", value: "" }; } @@ -1608,7 +1777,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitTemplateString(ctx: TemplateStringContext): Constant | Call { const pieces: Expression[] = []; for (const chunk of ctx.stringContents()) { - pieces.push(this.visit(chunk)); + pieces.push(this.visitStringContents(chunk)); } if (pieces.length === 0) { @@ -1628,7 +1797,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitFullTemplateString(ctx: FullTemplateStringContext): Constant | Call { const pieces: Expression[] = []; for (const chunk of ctx.stringContentsFull()) { - pieces.push(this.visit(chunk)); + pieces.push(this.visitStringContentsFull(chunk)); } if (pieces.length === 0) { @@ -1651,7 +1820,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { if (stringText) { return { expression_type: "constant", value: parseStringLiteralText(stringText.text) }; } else if (columnExpr) { - return this.visit(columnExpr); + return this.visitAsExpr(columnExpr); } return { expression_type: "constant", value: "" }; } @@ -1665,7 +1834,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { value: parseStringLiteralText(fullStringText.text), }; } else if (columnExpr) { - return this.visit(columnExpr); + return this.visitAsExpr(columnExpr); } return { expression_type: "constant", value: "" }; }