From e4ffbf49354d8a6bfc4e723a200f101e1c8f2dfb Mon Sep 17 00:00:00 2001 From: Matt Aitken Date: Thu, 11 Dec 2025 11:41:29 +0000 Subject: [PATCH] Adding stronger types --- internal-packages/tsql/src/query/ast.ts | 185 ++++++++---- internal-packages/tsql/src/query/parser.ts | 279 +++++++++++------- .../tsql/src/query/property_types.ts | 5 + 3 files changed, 311 insertions(+), 158 deletions(-) diff --git a/internal-packages/tsql/src/query/ast.ts b/internal-packages/tsql/src/query/ast.ts index a78b76edf..4efedc9cc 100644 --- a/internal-packages/tsql/src/query/ast.ts +++ b/internal-packages/tsql/src/query/ast.ts @@ -45,9 +45,43 @@ export interface UnknownType extends ConstantType { data_type: "unknown"; } +export type Expression = + | CTE + | Alias + | ArithmeticOperation + | And + | Or + | CompareOperation + | Not + | BetweenExpr + | OrderExpr + | ArrayAccess + | Array + | Dict + | TupleAccess + | Tuple + | Lambda + | Constant + | Field + | Placeholder + | Call + | ExprCall + | JoinConstraint + | JoinExpr + | WindowFrameExpr + | WindowExpr + | WindowFunction + | LimitByExpr + | SelectQuery + | SelectSetQuery + | RatioExpr + | SampleExpr + | HogQLXTag; + export interface CTE extends Expr { + expression_type: "cte"; name: string; - expr: Expr; + expr: Expression; cte_type: "column" | "subquery"; } @@ -186,7 +220,7 @@ export interface FieldTraverserType extends Type { export interface ExpressionFieldType extends Type { name: string; - expr: Expr; + expr: Expression; table_type: TableOrSelectType; isolate_scope?: boolean; } @@ -265,27 +299,27 @@ export type SetOperator = export interface Declaration extends AST {} export interface VariableAssignment extends Declaration { - left: Expr; - right: Expr; + left: Expression; + right: Expression; } export interface VariableDeclaration extends Declaration { name: string; - expr?: Expr; + expr?: Expression; } export interface Statement extends Declaration {} export interface ExprStatement extends Statement { - expr?: Expr; + expr?: Expression; } export interface ReturnStatement extends Statement { - expr?: Expr; + expr?: Expression; } export interface ThrowStatement extends Statement { - expr: Expr; + expr: Expression; } export interface TryCatchStatement extends Statement { @@ -295,27 +329,27 @@ export interface TryCatchStatement extends Statement { } export interface IfStatement extends Statement { - expr: Expr; + expr: Expression; then: Statement; else_?: Statement; } export interface WhileStatement extends Statement { - expr: Expr; + expr: Expression; body: Statement; } export interface ForStatement extends Statement { - initializer?: VariableDeclaration | VariableAssignment | Expr; - condition?: Expr; - increment?: Expr; + initializer?: VariableDeclaration | VariableAssignment | Expression; + condition?: Expression; + increment?: Expression; body: Statement; } export interface ForInStatement extends Statement { keyVar?: string; valueVar: string; - expr: Expr; + expr: Expression; body: Statement; } @@ -335,120 +369,141 @@ export interface Program extends AST { // Expression types export interface Alias extends Expr { + expression_type: "alias"; alias: string; - expr: Expr; + expr: Expression; hidden?: boolean; from_asterisk?: boolean; } export interface ArithmeticOperation extends Expr { - left: Expr; - right: Expr; + expression_type: "arithmetic_operation"; + left: Expression; + right: Expression; op: ArithmeticOperationOp; } export interface And extends Expr { + expression_type: "and"; type?: ConstantType; - exprs: Expr[]; + exprs: Expression[]; } export interface Or extends Expr { - exprs: Expr[]; + expression_type: "or"; + exprs: Expression[]; type?: ConstantType; } export interface CompareOperation extends Expr { - left: Expr; - right: Expr; + expression_type: "compare_operation"; + left: Expression; + right: Expression; op: CompareOperationOp; type?: ConstantType; } export interface Not extends Expr { - expr: Expr; + expression_type: "not"; + expr: Expression; type?: ConstantType; } export interface BetweenExpr extends Expr { - expr: Expr; - low: Expr; - high: Expr; + expression_type: "between_expr"; + expr: Expression; + low: Expression; + high: Expression; negated?: boolean; type?: ConstantType; } export interface OrderExpr extends Expr { - expr: Expr; + expression_type: "order_expr"; + expr: Expression; order?: "ASC" | "DESC"; } export interface ArrayAccess extends Expr { - array: Expr; - property: Expr; + expression_type: "array_access"; + array: Expression; + property: Expression; nullish?: boolean; } export interface Array extends Expr { - exprs: Expr[]; + expression_type: "array"; + exprs: Expression[]; } export interface Dict extends Expr { - items: [Expr, Expr][]; + expression_type: "dict"; + items: [Expression, Expression][]; } export interface TupleAccess extends Expr { - tuple: Expr; + expression_type: "tuple_access"; + tuple: Expression; index: number; nullish?: boolean; } export interface Tuple extends Expr { - exprs: Expr[]; + expression_type: "tuple"; + exprs: Expression[]; } export interface Lambda extends Expr { + expression_type: "lambda"; args: string[]; - expr: Expr | Block; + expr: Expression | Block; } export interface Constant extends Expr { + expression_type: "constant"; value: any; } export interface Field extends Expr { + expression_type: "field"; chain: (string | number)[]; from_asterisk?: boolean; } export interface Placeholder extends Expr { - expr: Expr; + expression_type: "placeholder"; + expr: Expression; // Computed properties chain?: (string | number)[] | null; field?: string | null; } export interface Call extends Expr { + expression_type: "call"; name: string; - args: Expr[]; - params?: Expr[]; + args: Expression[]; + params?: Expression[]; distinct?: boolean; } export interface ExprCall extends Expr { - expr: Expr; - args: Expr[]; + expression_type: "expr_call"; + expr: Expression; + args: Expression[]; } export interface JoinConstraint extends Expr { - expr: Expr; + expression_type: "join_constraint"; + expr: Expression; constraint_type: "ON" | "USING"; } export interface JoinExpr extends Expr { + expression_type: "join_expr"; type?: TableOrSelectType; join_type?: string; table?: SelectQuery | SelectSetQuery | Placeholder | HogQLXTag | Field; - table_args?: Expr[]; + table_args?: Expression[]; alias?: string; table_final?: boolean; constraint?: JoinConstraint; @@ -457,12 +512,14 @@ export interface JoinExpr extends Expr { } export interface WindowFrameExpr extends Expr { + expression_type: "window_frame_expr"; frame_type?: "CURRENT ROW" | "PRECEDING" | "FOLLOWING"; frame_value?: number; } export interface WindowExpr extends Expr { - partition_by?: Expr[]; + expression_type: "window_expr"; + partition_by?: Expression[]; order_by?: OrderExpr[]; frame_method?: "ROWS" | "RANGE"; frame_start?: WindowFrameExpr; @@ -470,37 +527,40 @@ export interface WindowExpr extends Expr { } export interface WindowFunction extends Expr { + expression_type: "window_function"; name: string; - args?: Expr[]; - exprs?: Expr[]; + args?: Expression[]; + exprs?: Expression[]; over_expr?: WindowExpr; over_identifier?: string; } export interface LimitByExpr extends Expr { - n: Expr; - exprs: Expr[]; - offset_value?: Expr; + expression_type: "limit_by_expr"; + n: Expression; + exprs: Expression[]; + offset_value?: Expression; } export interface SelectQuery extends Expr { + expression_type: "select_query"; type?: SelectQueryType; ctes?: Record; - select: Expr[]; + select: Expression[]; distinct?: boolean; select_from?: JoinExpr; array_join_op?: string; - array_join_list?: Expr[]; + array_join_list?: Expression[]; window_exprs?: Record; - where?: Expr; - prewhere?: Expr; - having?: Expr; - group_by?: Expr[]; + where?: Expression; + prewhere?: Expression; + having?: Expression; + group_by?: Expression[]; order_by?: OrderExpr[]; - limit?: Expr; + limit?: Expression; limit_by?: LimitByExpr; limit_with_ties?: boolean; - offset?: Expr; + offset?: Expression; settings?: HogQLQuerySettings; view_name?: string; } @@ -511,6 +571,7 @@ export interface SelectSetNode extends AST { } export interface SelectSetQuery extends Expr { + expression_type: "select_set_query"; type?: SelectSetQueryType; initial_select_query: SelectQuery | SelectSetQuery; subsequent_select_queries: SelectSetNode[]; @@ -529,11 +590,13 @@ export namespace SelectSetQuery { } export interface RatioExpr extends Expr { + expression_type: "ratio_expr"; left: Constant; right?: Constant; } export interface SampleExpr extends Expr { + expression_type: "sample_expr"; sample_value: RatioExpr; offset_value?: RatioExpr; } @@ -544,6 +607,7 @@ export interface HogQLXAttribute extends AST { } export interface HogQLXTag extends Expr { + expression_type: "hogqlx_tag"; kind: string; attributes: HogQLXAttribute[]; // Equivalent to to_dict() method @@ -557,12 +621,17 @@ export function createEmptySelectQuery(columns?: Record): } return { + expression_type: "select_query", select: Object.entries(columns).map(([column, field]) => ({ + expression_type: "alias" as const, alias: column, - expr: { value: (field as DatabaseField).default_value?.() ?? null } as Constant, - })) as Alias[], - where: { value: false } as Constant, - } as SelectQuery; + expr: { + expression_type: "constant", + value: (field as DatabaseField).default_value?.() ?? null, + } as Constant, + })), + where: { expression_type: "constant", value: false } as Constant, + }; } // Add static method equivalent for SelectQuery.empty() diff --git a/internal-packages/tsql/src/query/parser.ts b/internal-packages/tsql/src/query/parser.ts index c15b1e239..b83d09654 100644 --- a/internal-packages/tsql/src/query/parser.ts +++ b/internal-packages/tsql/src/query/parser.ts @@ -57,11 +57,35 @@ import { HogQLXAttribute, HogQLXTag, Declaration, + Expression, + AST, } from "./ast"; import { RESERVED_KEYWORDS } from "./constants"; import { SyntaxError, BaseHogQLError, NotImplementedError } from "./errors"; import type { HogQLTimings } from "./timings"; import { parseStringLiteralCtx, parseStringLiteralText, parseStringTextCtx } from "./parse_string"; +import { + CatchBlockContext, + DeclarationContext, + ExprContext, + ExpressionContext, + ExprStmtContext, + ForInStmtContext, + ForStmtContext, + FuncStmtContext, + IdentifierListContext, + IfStmtContext, + KvPairContext, + KvPairListContext, + ProgramContext, + ReturnStmtContext, + StatementContext, + ThrowStmtContext, + TryCatchStmtContext, + VarAssignmentContext, + VarDeclContext, + WhileStmtContext, +} from "../grammar/TSQLParser.js"; /** * Token with position information. @@ -182,73 +206,82 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { } // Program and declarations - visitProgram(ctx: any): Program { + visitProgram(ctx: ProgramContext): Program { const declarations: Declaration[] = []; // Implement based on your parser context structure throw new NotImplementedError("visitProgram not implemented"); } - visitDeclaration(ctx: any): Declaration { + visitDeclaration(ctx: DeclarationContext): Declaration { return this.visitChildren(ctx); } - visitExpression(ctx: any): Expr { + visitExpression(ctx: ExpressionContext): Expr { return this.visitChildren(ctx); } - visitVarDecl(ctx: any): VariableDeclaration { + visitVarDecl(ctx: VarDeclContext): VariableDeclaration { + const expr = ctx.expression(); return { name: this.visitIdentifier(ctx.identifier()), - expr: ctx.expression() ? this.visit(ctx.expression()) : undefined, + expr: expr ? this.visit(expr) : undefined, }; } - visitVarAssignment(ctx: any): VariableAssignment { + visitVarAssignment(ctx: VarAssignmentContext): VariableAssignment { return { left: this.visit(ctx.expression(0)), right: this.visit(ctx.expression(1)), }; } - visitStatement(ctx: any): Statement { + visitStatement(ctx: StatementContext): Statement { return this.visitChildren(ctx); } - visitExprStmt(ctx: any): ExprStatement { + visitExprStmt(ctx: ExprStmtContext): ExprStatement { return { expr: this.visit(ctx.expression()), }; } - visitReturnStmt(ctx: any): ReturnStatement { + visitReturnStmt(ctx: ReturnStmtContext): ReturnStatement { + const expr = ctx.expression(); return { - expr: ctx.expression() ? this.visit(ctx.expression()) : undefined, + expr: expr ? this.visit(expr) : undefined, }; } - visitThrowStmt(ctx: any): ThrowStatement { + visitThrowStmt(ctx: ThrowStmtContext): ThrowStatement { + const expr = ctx.expression(); return { - expr: ctx.expression() ? this.visit(ctx.expression()) : undefined, + expr: expr ? this.visit(expr) : undefined, }; } - visitCatchBlock(ctx: any): [string | null, string | null, Statement] { + visitCatchBlock(ctx: CatchBlockContext): [string | null, string | null, Statement] { + const catchVar = ctx._catchVar; + const catchType = ctx._catchType; + const catchStmt = ctx._catchStmt; return [ - ctx.catchVar ? this.visit(ctx.catchVar) : null, - ctx.catchType ? this.visit(ctx.catchType) : null, - this.visit(ctx.catchStmt), + catchVar ? this.visit(catchVar) : null, + catchType ? this.visit(catchType) : null, + catchStmt ? this.visit(catchStmt) : undefined, ]; } - visitTryCatchStmt(ctx: any): TryCatchStatement { + visitTryCatchStmt(ctx: TryCatchStmtContext): TryCatchStatement { + const tryStmt = ctx._tryStmt; + const catchBlocks = ctx.catchBlock(); + const finallyStmt = ctx._finallyStmt; return { - try_stmt: this.visit(ctx.tryStmt), - catches: ctx.catchBlock().map((c: any) => this.visit(c)), - finally_stmt: ctx.finallyStmt ? this.visit(ctx.finallyStmt) : undefined, + try_stmt: tryStmt ? this.visit(tryStmt) : undefined, + catches: catchBlocks.map((c: CatchBlockContext) => this.visit(c)), + finally_stmt: finallyStmt ? this.visit(finallyStmt) : undefined, }; } - visitIfStmt(ctx: any): IfStatement { + visitIfStmt(ctx: IfStmtContext): IfStatement { return { expr: this.visit(ctx.expression()), then: this.visit(ctx.statement(0)), @@ -256,14 +289,14 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { }; } - visitWhileStmt(ctx: any): WhileStatement { + visitWhileStmt(ctx: WhileStmtContext): WhileStatement { return { expr: this.visit(ctx.expression()), body: ctx.statement() ? this.visit(ctx.statement()) : undefined, }; } - visitForInStmt(ctx: any): ForInStatement { + visitForInStmt(ctx: ForInStmtContext): ForInStatement { const firstIdentifier = this.visitIdentifier(ctx.identifier(0)); const secondIdentifier = ctx.identifier(1) ? this.visitIdentifier(ctx.identifier(1)) : null; return { @@ -274,38 +307,39 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { }; } - visitForStmt(ctx: any): ForStatement { + visitForStmt(ctx: ForStmtContext): ForStatement { const initializer = - ctx.initializerVarDeclr || ctx.initializerVarAssignment || ctx.initializerExpression; + ctx._initializerVarDeclr || ctx._initializerVarAssignment || ctx._initializerExpression; const increment = - ctx.incrementVarDeclr || ctx.incrementVarAssignment || ctx.incrementExpression; + ctx._incrementVarDeclr || ctx._incrementVarAssignment || ctx._incrementExpression; return { initializer: initializer ? this.visit(initializer) : undefined, - condition: ctx.condition ? this.visit(ctx.condition) : undefined, + condition: ctx._condition ? this.visit(ctx._condition) : undefined, increment: increment ? this.visit(increment) : undefined, body: this.visit(ctx.statement()), }; } - visitFuncStmt(ctx: any): Function { + visitFuncStmt(ctx: FuncStmtContext): Function { + const params = ctx.identifierList(); return { name: this.visitIdentifier(ctx.identifier()), - params: ctx.identifierList() ? this.visit(ctx.identifierList()) : [], + params: params ? this.visit(params) : [], body: this.visit(ctx.block()), }; } - visitKvPairList(ctx: any): [Expr, Expr][] { + visitKvPairList(ctx: KvPairListContext): [Expr, Expr][] { return ctx.kvPair().map((kv: any) => this.visit(kv)); } - visitKvPair(ctx: any): [Expr, Expr] { + visitKvPair(ctx: KvPairContext): [Expr, Expr] { const exprs = ctx.expression(); return [this.visit(exprs[0]), this.visit(exprs[1])]; } - visitIdentifierList(ctx: any): string[] { + visitIdentifierList(ctx: IdentifierListContext): string[] { return ctx.identifier().map((ident: any) => this.visitIdentifier(ident)); } @@ -356,6 +390,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { return initialQuery; } return { + expression_type: "select_set_query", initial_select_query: initialQuery, subsequent_select_queries: selectQueries, }; @@ -367,6 +402,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitSelectStmt(ctx: any): SelectQuery { const selectQuery: SelectQuery = { + expression_type: "select_query", ctes: ctx.withClause() ? this.visit(ctx.withClause()) : undefined, select: ctx.columnExprList() ? this.visit(ctx.columnExprList()) : [], distinct: ctx.DISTINCT() ? true : undefined, @@ -471,6 +507,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { if (Array.isArray(limitExpr) && limitExpr.length === 2) { const [n, offsetValue] = limitExpr; return { + expression_type: "limit_by_expr", n, offset_value: offsetValue, exprs: this.visit(ctx.columnExprList()), @@ -479,6 +516,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { // If no offset, just use limitExpr as n return { + expression_type: "limit_by_expr", n: limitExpr, offset_value: undefined, exprs: this.visit(ctx.columnExprList()), @@ -535,6 +573,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { return table; } return { + expression_type: "join_expr", table, table_final: tableFinal, sample, @@ -594,6 +633,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { throw new NotImplementedError("Unsupported: JOIN ... ON with multiple expressions"); } return { + expression_type: "join_constraint", expr: columnExprList[0], constraint_type: ctx.USING() ? "USING" : "ON", }; @@ -606,6 +646,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { ratioExpressions.length > 1 && ctx.OFFSET() ? this.visit(ratioExpressions[1]) : undefined; return { + expression_type: "sample_expr", sample_value: sampleRatioExpr, offset_value: offsetRatioExpr, }; @@ -618,6 +659,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitOrderExpr(ctx: any): OrderExpr { const order = ctx.DESC() || ctx.DESCENDING() ? "DESC" : "ASC"; return { + expression_type: "order_expr", expr: this.visit(ctx.columnExpr()), order: order as "ASC" | "DESC", }; @@ -633,6 +675,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { const right = ctx.SLASH() && numberLiterals.length > 1 ? numberLiterals[1] : null; return { + expression_type: "ratio_expr", left: this.visitNumberLiteral(left), right: right ? this.visitNumberLiteral(right) : undefined, }; @@ -642,6 +685,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { const frame = ctx.winFrameClause(); const visitedFrame = frame ? this.visit(frame) : undefined; return { + expression_type: "window_expr", partition_by: ctx.winPartitionByClause() ? this.visit(ctx.winPartitionByClause()) : undefined, order_by: ctx.winOrderByClause() ? this.visit(ctx.winOrderByClause()) : undefined, frame_method: frame && frame.RANGE() ? "RANGE" : frame && frame.ROWS() ? "ROWS" : undefined, @@ -673,6 +717,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitWinFrameBound(ctx: any): WindowFrameExpr { if (ctx.PRECEDING()) { return { + expression_type: "window_frame_expr", frame_type: "PRECEDING", frame_value: ctx.numberLiteral() ? (this.visit(ctx.numberLiteral()) as Constant).value @@ -681,13 +726,14 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { } if (ctx.FOLLOWING()) { return { + expression_type: "window_frame_expr", frame_type: "FOLLOWING", frame_value: ctx.numberLiteral() ? (this.visit(ctx.numberLiteral()) as Constant).value : undefined, }; } - return { frame_type: "CURRENT ROW" }; + return { expression_type: "window_frame_expr", frame_type: "CURRENT ROW" }; } // Column expressions @@ -697,6 +743,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitColumnExprTernaryOp(ctx: any): Call { return { + expression_type: "call", name: "if", args: [ this.visit(ctx.columnExpr(0)), @@ -723,11 +770,12 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { ); } - return { expr, alias }; + return { expression_type: "alias", expr, alias }; } visitColumnExprNegate(ctx: any): ArithmeticOperation { return { + expression_type: "arithmetic_operation", op: ArithmeticOperationOp.Sub, left: { value: 0 } as Constant, right: this.visit(ctx.columnExpr()), @@ -736,6 +784,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitColumnExprDict(ctx: any): Dict { return { + expression_type: "dict", items: ctx.kvPairList() ? this.visit(ctx.kvPairList()) : [], }; } @@ -750,6 +799,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitColumnExprArray(ctx: any): ArrayExpression { return { + expression_type: "array", exprs: ctx.columnExprList() ? this.visit(ctx.columnExprList()) : [], }; } @@ -768,7 +818,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { // Use columnExpr() method to get left and right operands const left = this.visit(ctx.columnExpr(0)); const right = this.visit(ctx.columnExpr(1)); - return { left, right, op }; + return { expression_type: "arithmetic_operation", left, right, op }; } visitColumnExprPrecedence2(ctx: any): ArithmeticOperation | Call { @@ -777,11 +827,21 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { const right = this.visit(ctx.columnExpr(1)); if (ctx.PLUS()) { - return { left, right, op: ArithmeticOperationOp.Add }; + return { + expression_type: "arithmetic_operation", + left, + right, + op: ArithmeticOperationOp.Add, + }; } else if (ctx.DASH()) { - return { left, right, op: ArithmeticOperationOp.Sub }; + return { + expression_type: "arithmetic_operation", + left, + right, + op: ArithmeticOperationOp.Sub, + }; } else if (ctx.CONCAT()) { - const args: Expr[] = []; + const args: Expression[] = []; if ("name" in left && left.name === "concat" && "args" in left) { args.push(...left.args); } else { @@ -794,7 +854,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { args.push(right); } - return { name: "concat", args }; + return { expression_type: "call", name: "concat", args }; } else { throw new NotImplementedError(`Unsupported ColumnExprPrecedence2: ${ctx.getText()}`); } @@ -840,7 +900,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { throw new NotImplementedError(`Unsupported ColumnExprPrecedence3: ${ctx.getText()}`); } - return { left, right, op }; + return { expression_type: "compare_operation", left, right, op }; } visitColumnExprInterval(ctx: any): Call { @@ -866,11 +926,12 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { throw new NotImplementedError(`Unsupported interval type: ${interval.getText()}`); } - return { name, args: [this.visit(ctx.columnExpr())] }; + return { expression_type: "call", name, args: [this.visit(ctx.columnExpr())] }; } visitColumnExprIsNull(ctx: any): CompareOperation { return { + expression_type: "compare_operation", left: this.visit(ctx.columnExpr()), right: { value: null } as Constant, op: ctx.NOT() ? CompareOperationOp.NotEq : CompareOperationOp.Eq, @@ -879,36 +940,38 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitColumnExprTuple(ctx: any): Tuple { return { + expression_type: "tuple", exprs: ctx.columnExprList() ? this.visit(ctx.columnExprList()) : [], }; } visitColumnExprArrayAccess(ctx: any): ArrayAccess { - const object: Expr = this.visit(ctx.columnExpr(0)); - const property: Expr = this.visit(ctx.columnExpr(1)); - return { array: object, property }; + const object: Expression = this.visit(ctx.columnExpr(0)); + const property: Expression = this.visit(ctx.columnExpr(1)); + return { expression_type: "array_access", array: object, property }; } visitColumnExprNullArrayAccess(ctx: any): ArrayAccess { - const object: Expr = this.visit(ctx.columnExpr(0)); - const property: Expr = this.visit(ctx.columnExpr(1)); - return { array: object, property, nullish: true }; + const object: Expression = this.visit(ctx.columnExpr(0)); + const property: Expression = this.visit(ctx.columnExpr(1)); + return { expression_type: "array_access", array: object, property, nullish: true }; } visitColumnExprPropertyAccess(ctx: any): ArrayAccess { const object = this.visit(ctx.columnExpr()); const property = { value: this.visitIdentifier(ctx.identifier()) } as Constant; - return { array: object, property }; + return { expression_type: "array_access", array: object, property }; } visitColumnExprNullPropertyAccess(ctx: any): ArrayAccess { const object = this.visit(ctx.columnExpr()); const property = { value: this.visitIdentifier(ctx.identifier()) } as Constant; - return { array: object, property, nullish: true }; + return { expression_type: "array_access", array: object, property, nullish: true }; } visitColumnExprBetween(ctx: any): BetweenExpr { return { + expression_type: "between_expr", expr: this.visit(ctx.columnExpr(0)), low: this.visit(ctx.columnExpr(1)), high: this.visit(ctx.columnExpr(2)), @@ -927,7 +990,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { let right = this.visit(ctx.columnExpr(1)); const rightArray = "exprs" in right ? right.exprs : [right]; - return { exprs: [...leftArray, ...rightArray] }; + return { expression_type: "and", exprs: [...leftArray, ...rightArray] }; } visitColumnExprOr(ctx: any): Or { @@ -937,28 +1000,28 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { let right = this.visit(ctx.columnExpr(1)); const rightArray = "exprs" in right ? right.exprs : [right]; - return { exprs: [...leftArray, ...rightArray] }; + return { expression_type: "or", exprs: [...leftArray, ...rightArray] }; } visitColumnExprTupleAccess(ctx: any): TupleAccess { const tuple = this.visit(ctx.columnExpr()); const index = parseInt(ctx.DECIMAL_LITERAL().getText()); - return { tuple, index }; + return { expression_type: "tuple_access", tuple, index }; } visitColumnExprNullTupleAccess(ctx: any): TupleAccess { const tuple = this.visit(ctx.columnExpr()); const index = parseInt(ctx.DECIMAL_LITERAL().getText()); - return { tuple, index, nullish: true }; + return { expression_type: "tuple_access", tuple, index, nullish: true }; } visitColumnExprCase(ctx: any): Call { const columns = ctx.columnExpr().map((column: any) => this.visit(column)); if (ctx.caseExpr) { - const args: Expr[] = [ + const args: Expression[] = [ columns[0], - { exprs: [] } as ArrayExpression, - { exprs: [] } as ArrayExpression, + { expression_type: "array", exprs: [] }, + { expression_type: "array", exprs: [] }, , columns[columns.length - 1], ]; @@ -966,20 +1029,21 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { const arrayIndex = ((index - 1) % 2) + 1; (args[arrayIndex] as ArrayExpression).exprs.push(columns[index]); } - return { name: "transform", args }; + return { expression_type: "call", name: "transform", args }; } else if (columns.length === 3) { - return { name: "if", args: columns }; + return { expression_type: "call", name: "if", args: columns }; } else { - return { name: "multiIf", args: columns }; + return { expression_type: "call", name: "multiIf", args: columns }; } } visitColumnExprNot(ctx: any): Not { - return { expr: this.visit(ctx.columnExpr()) }; + return { expression_type: "not", expr: this.visit(ctx.columnExpr()) }; } visitColumnExprWinFunctionTarget(ctx: any): 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) : [], @@ -989,6 +1053,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitColumnExprWinFunction(ctx: any): WindowFunction { return { + expression_type: "window_function", name: this.visitIdentifier(ctx.identifier()), exprs: ctx.columnExprs ? this.visit(ctx.columnExprs) : [], args: ctx.columnArgList ? this.visit(ctx.columnArgList) : [], @@ -1003,27 +1068,30 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitColumnExprFunction(ctx: any): Call { const name = this.visitIdentifier(ctx.identifier()); - let parameters: Expr[] | undefined = ctx.columnExprs ? this.visit(ctx.columnExprs) : undefined; + let parameters: Expression[] | undefined = ctx.columnExprs + ? this.visit(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: Expr[] = ctx.columnArgList ? this.visit(ctx.columnArgList) : []; + const args: Expression[] = ctx.columnArgList ? this.visit(ctx.columnArgList) : []; const distinct = ctx.DISTINCT() ? true : false; - return { name, params: parameters, args, distinct }; + return { expression_type: "call", name, params: parameters, args, distinct }; } visitColumnExprAsterisk(ctx: any): Field { if (ctx.tableIdentifier()) { const table = this.visit(ctx.tableIdentifier()); - return { chain: [...table, "*"] }; + return { expression_type: "field", chain: [...table, "*"] }; } - return { chain: ["*"] }; + return { expression_type: "field", chain: ["*"] }; } visitColumnLambdaExpr(ctx: any): Lambda { return { + expression_type: "lambda", args: ctx.identifier().map((identifier: any) => this.visitIdentifier(identifier)), expr: ctx.columnExpr() ? this.visit(ctx.columnExpr()) : this.visit(ctx.block()), }; @@ -1041,16 +1109,16 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitWithExprSubquery(ctx: any): CTE { const subquery = this.visit(ctx.selectSetStmt()); const name = this.visitIdentifier(ctx.identifier()); - return { name, expr: subquery, cte_type: "subquery" }; + return { expression_type: "cte", name, expr: subquery, cte_type: "subquery" }; } visitWithExprColumn(ctx: any): CTE { const expr = this.visit(ctx.columnExpr()); const name = this.visitIdentifier(ctx.identifier()); - return { name, expr, cte_type: "column" }; + return { expression_type: "cte", name, expr, cte_type: "column" }; } - visitColumnIdentifier(ctx: any): Expr { + visitColumnIdentifier(ctx: any): Expression { if (ctx.placeholder()) { return this.visit(ctx.placeholder()); } @@ -1061,15 +1129,15 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { if (table.length === 0 && nested.length > 0) { const text = ctx.getText().toLowerCase(); if (text === "true") { - return { value: true } as Constant; + return { expression_type: "constant", value: true }; } if (text === "false") { - return { value: false } as Constant; + return { expression_type: "constant", value: false }; } - return { chain: nested } as Field; + return { expression_type: "field", chain: nested }; } - return { chain: [...table, ...nested] } as Field; + return { expression_type: "field", chain: [...table, ...nested] }; } visitNestedIdentifier(ctx: any): string[] { @@ -1078,7 +1146,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitTableExprIdentifier(ctx: any): Field { const chain = this.visit(ctx.tableIdentifier()); - return { chain }; + return { expression_type: "field", chain }; } visitTableExprSubquery(ctx: any): SelectQuery | SelectSetQuery { @@ -1101,7 +1169,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { table.alias = alias; return table; } - return { table, alias }; + return { expression_type: "join_expr", table, alias }; } visitTableExprFunction(ctx: any): JoinExpr { @@ -1115,7 +1183,11 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitTableFunctionExpr(ctx: any): JoinExpr { const name = this.visitIdentifier(ctx.identifier()); const args = ctx.tableArgList() ? this.visit(ctx.tableArgList()) : []; - return { table: { chain: [name] } as Field, table_args: args }; + return { + expression_type: "join_expr", + table: { expression_type: "field", chain: [name] }, + table_args: args, + }; } visitTableIdentifier(ctx: any): string[] { @@ -1131,7 +1203,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { return nestedArray; } - visitTableArgList(ctx: any): Expr[] { + visitTableArgList(ctx: any): Expression[] { return ctx.columnExpr().map((arg: any) => this.visit(arg)); } @@ -1154,20 +1226,20 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { text === "inf" || text === "nan" ) { - return { value: parseFloat(text) }; + return { expression_type: "constant", value: parseFloat(text) }; } - return { value: parseInt(text) }; + return { expression_type: "constant", value: parseInt(text) }; } visitLiteral(ctx: any): Constant { if (ctx.NULL_SQL()) { - return { value: null }; + return { expression_type: "constant", value: null }; } if (ctx.STRING_LITERAL()) { // STRING_LITERAL() returns a TerminalNode, which has getText() const stringLiteral = ctx.STRING_LITERAL(); const text = parseStringLiteralCtx(stringLiteral); - return { value: text }; + return { expression_type: "constant", value: text }; } if (ctx.numberLiteral()) { return this.visitNumberLiteral(ctx.numberLiteral()); @@ -1230,6 +1302,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitColumnExprNullish(ctx: any): Call { return { + expression_type: "call", name: "ifNull", args: [this.visit(ctx.columnExpr(0)), this.visit(ctx.columnExpr(1))], }; @@ -1237,6 +1310,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitColumnExprCall(ctx: any): ExprCall { return { + expression_type: "expr_call", expr: this.visit(ctx.columnExpr()), args: ctx.columnExprList() ? this.visit(ctx.columnExprList()) : [], }; @@ -1246,11 +1320,13 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { const expr = this.visit(ctx.columnExpr()); if ("chain" in expr && expr.chain.length === 1) { return { + expression_type: "call", name: String(expr.chain[0]), args: [this.visit(ctx.selectSetStmt())], }; } return { + expression_type: "expr_call", expr, args: [this.visit(ctx.selectSetStmt())], }; @@ -1267,7 +1343,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { } visitHogqlxText(ctx: any): Constant { - return { value: ctx.HOGQLX_TEXT_TEXT().getText() }; + return { expression_type: "constant", value: ctx.HOGQLX_TEXT_TEXT().getText() }; } visitHogqlxTagElementClosed(ctx: any): HogQLXTag { @@ -1275,7 +1351,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { const attributes = ctx.hogqlxTagAttribute() ? ctx.hogqlxTagAttribute().map((a: any) => this.visit(a)) : []; - return { kind, attributes }; + return { expression_type: "hogqlx_tag", kind, attributes }; } visitHogqlxTagElementNested(ctx: any): HogQLXTag { @@ -1292,7 +1368,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { : []; // ── collect child nodes, discarding pure-indentation whitespace ── - const keptChildren: Expr[] = []; + const keptChildren: Expression[] = []; for (const element of ctx.hogqlxChildElement()) { const child = this.visit(element); @@ -1317,7 +1393,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { attributes.push({ name: "children", value: keptChildren }); } - return { kind: opening, attributes }; + return { expression_type: "hogqlx_tag", kind: opening, attributes }; } visitHogqlxTagAttribute(ctx: any): HogQLXAttribute { @@ -1332,7 +1408,7 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { } visitPlaceholder(ctx: any): Placeholder { - return { expr: this.visit(ctx.columnExpr()) }; + return { expression_type: "placeholder", expr: this.visit(ctx.columnExpr()) }; } visitColumnExprTemplateString(ctx: any): Expr { @@ -1341,66 +1417,69 @@ export class TSQLParseTreeConverter implements TSQLParserVisitor { visitString(ctx: any): Constant | Expr { if (ctx.STRING_LITERAL()) { - return { value: parseStringLiteralCtx(ctx.STRING_LITERAL()) }; + return { expression_type: "constant", value: parseStringLiteralCtx(ctx.STRING_LITERAL()) }; } return this.visit(ctx.templateString()); } visitTemplateString(ctx: any): Constant | Call { - const pieces: Expr[] = []; + const pieces: Expression[] = []; for (const chunk of ctx.stringContents()) { pieces.push(this.visit(chunk)); } if (pieces.length === 0) { - return { value: "" }; + return { expression_type: "constant", value: "" }; } else if (pieces.length === 1) { const first = pieces[0]; // If it's already a Constant or Call, return as-is, otherwise wrap in Call if ("value" in first || "name" in first) { return first as Constant | Call; } - return { name: "concat", args: [first] }; + return { expression_type: "call", name: "concat", args: [first] }; } - return { name: "concat", args: pieces }; + return { expression_type: "call", name: "concat", args: pieces }; } visitFullTemplateString(ctx: any): Constant | Call { - const pieces: Expr[] = []; + const pieces: Expression[] = []; for (const chunk of ctx.stringContentsFull()) { pieces.push(this.visit(chunk)); } if (pieces.length === 0) { - return { value: "" }; + return { expression_type: "constant", value: "" }; } else if (pieces.length === 1) { const first = pieces[0]; // If it's already a Constant or Call, return as-is, otherwise wrap in Call if ("value" in first || "name" in first) { return first as Constant | Call; } - return { name: "concat", args: [first] }; + return { expression_type: "call", name: "concat", args: [first] }; } - return { name: "concat", args: pieces }; + return { expression_type: "call", name: "concat", args: pieces }; } - visitStringContents(ctx: any): Constant | Expr { + visitStringContents(ctx: any): Constant | Expression { if (ctx.STRING_TEXT()) { - return { value: parseStringTextCtx(ctx.STRING_TEXT(), true) }; + return { expression_type: "constant", value: parseStringTextCtx(ctx.STRING_TEXT(), true) }; } else if (ctx.columnExpr()) { return this.visit(ctx.columnExpr()); } - return { value: "" }; + return { expression_type: "constant", value: "" }; } - visitStringContentsFull(ctx: any): Constant | Expr { + visitStringContentsFull(ctx: any): Constant | Expression { if (ctx.FULL_STRING_TEXT()) { - return { value: parseStringTextCtx(ctx.FULL_STRING_TEXT(), false) }; + return { + expression_type: "constant", + value: parseStringTextCtx(ctx.FULL_STRING_TEXT(), false), + }; } else if (ctx.columnExpr()) { return this.visit(ctx.columnExpr()); } - return { value: "" }; + return { expression_type: "constant", value: "" }; } } diff --git a/internal-packages/tsql/src/query/property_types.ts b/internal-packages/tsql/src/query/property_types.ts index 9fefa579e..13f2a1dcf 100644 --- a/internal-packages/tsql/src/query/property_types.ts +++ b/internal-packages/tsql/src/query/property_types.ts @@ -520,6 +520,7 @@ export class PropertySwapper extends CloningVisitor { private createToTimeZoneCall(node: Field): Call { return { + expression_type: "call", name: "toTimeZone", args: [node, this.createConstant(this.timezone)], type: { @@ -534,6 +535,7 @@ export class PropertySwapper extends CloningVisitor { private createToDateTimeCall(node: Field): Call { return { + expression_type: "call", name: "toDateTime", args: [node], start: node.start, @@ -543,6 +545,7 @@ export class PropertySwapper extends CloningVisitor { private createToFloatCall(node: Field): Call { return { + expression_type: "call", name: "toFloat", args: [node], start: node.start, @@ -552,6 +555,7 @@ export class PropertySwapper extends CloningVisitor { private createToBoolCall(node: Field): Call { return { + expression_type: "call", name: "toBool", args: [ { @@ -578,6 +582,7 @@ export class PropertySwapper extends CloningVisitor { private createConstant(value: any): Constant { return { + expression_type: "constant", value, }; }