diff --git a/src/main/java/net/sf/jsqlparser/util/TablesNamesFinder.java b/src/main/java/net/sf/jsqlparser/util/TablesNamesFinder.java index 94b7e1de8..a2600f2de 100644 --- a/src/main/java/net/sf/jsqlparser/util/TablesNamesFinder.java +++ b/src/main/java/net/sf/jsqlparser/util/TablesNamesFinder.java @@ -9,52 +9,6 @@ */ package net.sf.jsqlparser.util; -import net.sf.jsqlparser.statement.create.accessmethod.CreateAccessMethod; - -import net.sf.jsqlparser.statement.create.fdw.CreateForeignDataWrapper; -import net.sf.jsqlparser.statement.alter.AlterForeignDataWrapper; -import net.sf.jsqlparser.statement.create.server.CreateServer; -import net.sf.jsqlparser.statement.alter.AlterServer; -import net.sf.jsqlparser.statement.create.usermapping.CreateUserMapping; -import net.sf.jsqlparser.statement.alter.AlterUserMapping; -import net.sf.jsqlparser.statement.create.textsearch.CreateTextSearchConfiguration; -import net.sf.jsqlparser.statement.alter.AlterTextSearchConfiguration; -import net.sf.jsqlparser.statement.create.rule.CreateRule; -import net.sf.jsqlparser.statement.notify.NotifyStatement; -import net.sf.jsqlparser.statement.create.collation.CreateCollation; -import net.sf.jsqlparser.statement.alter.AlterCollation; -import net.sf.jsqlparser.statement.alter.AlterPolicy; -import net.sf.jsqlparser.statement.drop.DropPolicy; -import net.sf.jsqlparser.statement.create.statistics.CreateStatistics; -import net.sf.jsqlparser.statement.alter.AlterStatistics; -import net.sf.jsqlparser.statement.alter.AlterRelation; -import net.sf.jsqlparser.statement.alter.AlterTablespaceMove; -import net.sf.jsqlparser.statement.alter.database.AlterDatabase; -import net.sf.jsqlparser.statement.alter.schema.AlterSchema; -import net.sf.jsqlparser.statement.select.MatchRecognize; -import net.sf.jsqlparser.expression.RowPatternFunction; - -import net.sf.jsqlparser.expression.AliasedExpression; - -import net.sf.jsqlparser.statement.oracle.OracleBlock; -import net.sf.jsqlparser.statement.oracle.OracleAssignment; -import net.sf.jsqlparser.statement.oracle.OracleNullStatement; - -import net.sf.jsqlparser.statement.role.CreateRole; -import net.sf.jsqlparser.statement.role.AlterRole; -import net.sf.jsqlparser.statement.grant.Revoke; -import net.sf.jsqlparser.statement.grant.AlterDefaultPrivileges; -import net.sf.jsqlparser.statement.create.type.CreateType; -import net.sf.jsqlparser.statement.alter.AlterType; -import net.sf.jsqlparser.statement.create.domain.CreateDomain; -import net.sf.jsqlparser.statement.alter.AlterDomain; -import net.sf.jsqlparser.statement.create.extension.CreateExtension; -import net.sf.jsqlparser.statement.alter.AlterExtension; -import net.sf.jsqlparser.statement.create.publication.CreatePublication; -import net.sf.jsqlparser.statement.alter.AlterPublication; -import net.sf.jsqlparser.statement.create.subscription.CreateSubscription; -import net.sf.jsqlparser.statement.alter.AlterSubscription; - import java.util.ArrayList; import java.util.HashSet; import java.util.List; @@ -62,6 +16,8 @@ import java.util.Set; import net.sf.jsqlparser.JSQLParserException; import net.sf.jsqlparser.expression.*; +import net.sf.jsqlparser.expression.AliasedExpression; +import net.sf.jsqlparser.expression.RowPatternFunction; import net.sf.jsqlparser.expression.operators.arithmetic.Addition; import net.sf.jsqlparser.expression.operators.arithmetic.BitwiseAnd; import net.sf.jsqlparser.expression.operators.arithmetic.BitwiseLeftShift; @@ -95,9 +51,9 @@ import net.sf.jsqlparser.expression.operators.relational.Intersects; import net.sf.jsqlparser.expression.operators.relational.IsBooleanExpression; import net.sf.jsqlparser.expression.operators.relational.IsDistinctExpression; +import net.sf.jsqlparser.expression.operators.relational.IsJsonExpression; import net.sf.jsqlparser.expression.operators.relational.IsNullExpression; import net.sf.jsqlparser.expression.operators.relational.IsUnknownExpression; -import net.sf.jsqlparser.expression.operators.relational.IsJsonExpression; import net.sf.jsqlparser.expression.operators.relational.JsonOperator; import net.sf.jsqlparser.expression.operators.relational.LikeExpression; import net.sf.jsqlparser.expression.operators.relational.Matches; @@ -114,60 +70,105 @@ import net.sf.jsqlparser.parser.CCJSqlParserUtil; import net.sf.jsqlparser.schema.Column; import net.sf.jsqlparser.schema.Table; +import net.sf.jsqlparser.statement.AssertStatement; +import net.sf.jsqlparser.statement.AttachStatement; import net.sf.jsqlparser.statement.Block; import net.sf.jsqlparser.statement.Commit; -import net.sf.jsqlparser.statement.StartTransaction; -import net.sf.jsqlparser.statement.ReleaseSavepointStatement; +import net.sf.jsqlparser.statement.ConnectStatement; +import net.sf.jsqlparser.statement.CopyStatement; import net.sf.jsqlparser.statement.CreateFunctionalStatement; +import net.sf.jsqlparser.statement.DeallocateStatement; import net.sf.jsqlparser.statement.DeclareStatement; import net.sf.jsqlparser.statement.DescribeStatement; +import net.sf.jsqlparser.statement.DetachStatement; +import net.sf.jsqlparser.statement.DisconnectStatement; import net.sf.jsqlparser.statement.DoStatement; import net.sf.jsqlparser.statement.ExplainStatement; +import net.sf.jsqlparser.statement.ExtensionStatement; import net.sf.jsqlparser.statement.IfElseStatement; import net.sf.jsqlparser.statement.OutputClause; +import net.sf.jsqlparser.statement.PragmaStatement; +import net.sf.jsqlparser.statement.PrepareStatement; import net.sf.jsqlparser.statement.PurgeObjectType; import net.sf.jsqlparser.statement.PurgeStatement; +import net.sf.jsqlparser.statement.ReleaseSavepointStatement; import net.sf.jsqlparser.statement.ResetStatement; import net.sf.jsqlparser.statement.ReturningClause; import net.sf.jsqlparser.statement.RollbackStatement; import net.sf.jsqlparser.statement.SavepointStatement; import net.sf.jsqlparser.statement.SessionStatement; +import net.sf.jsqlparser.statement.SetIdentityInsertStatement; import net.sf.jsqlparser.statement.SetStatement; import net.sf.jsqlparser.statement.ShowColumnsStatement; import net.sf.jsqlparser.statement.ShowStatement; +import net.sf.jsqlparser.statement.StartTransaction; import net.sf.jsqlparser.statement.Statement; import net.sf.jsqlparser.statement.StatementVisitor; import net.sf.jsqlparser.statement.Statements; import net.sf.jsqlparser.statement.UnsupportedStatement; import net.sf.jsqlparser.statement.UseStatement; -import net.sf.jsqlparser.statement.SetIdentityInsertStatement; import net.sf.jsqlparser.statement.alter.Alter; +import net.sf.jsqlparser.statement.alter.AlterCollation; +import net.sf.jsqlparser.statement.alter.AlterDomain; +import net.sf.jsqlparser.statement.alter.AlterExtension; +import net.sf.jsqlparser.statement.alter.AlterForeignDataWrapper; +import net.sf.jsqlparser.statement.alter.AlterPolicy; +import net.sf.jsqlparser.statement.alter.AlterPublication; +import net.sf.jsqlparser.statement.alter.AlterRelation; +import net.sf.jsqlparser.statement.alter.AlterServer; import net.sf.jsqlparser.statement.alter.AlterSession; +import net.sf.jsqlparser.statement.alter.AlterStatistics; +import net.sf.jsqlparser.statement.alter.AlterSubscription; import net.sf.jsqlparser.statement.alter.AlterSystemStatement; +import net.sf.jsqlparser.statement.alter.AlterTablespaceMove; +import net.sf.jsqlparser.statement.alter.AlterTextSearchConfiguration; +import net.sf.jsqlparser.statement.alter.AlterType; +import net.sf.jsqlparser.statement.alter.AlterUserMapping; import net.sf.jsqlparser.statement.alter.RenameTableStatement; +import net.sf.jsqlparser.statement.alter.database.AlterDatabase; +import net.sf.jsqlparser.statement.alter.schema.AlterSchema; import net.sf.jsqlparser.statement.alter.sequence.AlterSequence; import net.sf.jsqlparser.statement.analyze.Analyze; import net.sf.jsqlparser.statement.comment.Comment; +import net.sf.jsqlparser.statement.create.accessmethod.CreateAccessMethod; +import net.sf.jsqlparser.statement.create.collation.CreateCollation; import net.sf.jsqlparser.statement.create.database.CreateDatabase; +import net.sf.jsqlparser.statement.create.domain.CreateDomain; import net.sf.jsqlparser.statement.create.event.AlterEvent; import net.sf.jsqlparser.statement.create.event.CreateEvent; +import net.sf.jsqlparser.statement.create.extension.CreateExtension; +import net.sf.jsqlparser.statement.create.extension.CreateExtensionRepository; +import net.sf.jsqlparser.statement.create.fdw.CreateForeignDataWrapper; import net.sf.jsqlparser.statement.create.index.CreateIndex; +import net.sf.jsqlparser.statement.create.macro.CreateMacro; import net.sf.jsqlparser.statement.create.policy.CreatePolicy; +import net.sf.jsqlparser.statement.create.publication.CreatePublication; +import net.sf.jsqlparser.statement.create.rule.CreateRule; import net.sf.jsqlparser.statement.create.schema.CreateSchema; import net.sf.jsqlparser.statement.create.sequence.CreateSequence; +import net.sf.jsqlparser.statement.create.server.CreateServer; +import net.sf.jsqlparser.statement.create.statistics.CreateStatistics; +import net.sf.jsqlparser.statement.create.subscription.CreateSubscription; import net.sf.jsqlparser.statement.create.synonym.CreateSynonym; import net.sf.jsqlparser.statement.create.table.CreateTable; +import net.sf.jsqlparser.statement.create.textsearch.CreateTextSearchConfiguration; import net.sf.jsqlparser.statement.create.trigger.CreateTrigger; +import net.sf.jsqlparser.statement.create.type.CreateType; import net.sf.jsqlparser.statement.create.user.CreateUser; +import net.sf.jsqlparser.statement.create.usermapping.CreateUserMapping; import net.sf.jsqlparser.statement.create.view.AlterView; import net.sf.jsqlparser.statement.create.view.CreateView; import net.sf.jsqlparser.statement.delete.Delete; import net.sf.jsqlparser.statement.delete.ParenthesedDelete; import net.sf.jsqlparser.statement.drop.Drop; +import net.sf.jsqlparser.statement.drop.DropPolicy; import net.sf.jsqlparser.statement.execute.Execute; import net.sf.jsqlparser.statement.execute.ExecuteArgument; import net.sf.jsqlparser.statement.export.Export; +import net.sf.jsqlparser.statement.export.ExportDataStatement; +import net.sf.jsqlparser.statement.grant.AlterDefaultPrivileges; import net.sf.jsqlparser.statement.grant.Grant; +import net.sf.jsqlparser.statement.grant.Revoke; import net.sf.jsqlparser.statement.imprt.Import; import net.sf.jsqlparser.statement.insert.Insert; import net.sf.jsqlparser.statement.insert.InsertBulk; @@ -176,6 +177,7 @@ import net.sf.jsqlparser.statement.insert.OracleMultiInsertBranch; import net.sf.jsqlparser.statement.insert.OracleMultiInsertClause; import net.sf.jsqlparser.statement.insert.ParenthesedInsert; +import net.sf.jsqlparser.statement.load.LoadDataStatement; import net.sf.jsqlparser.statement.lock.LockStatement; import net.sf.jsqlparser.statement.merge.Merge; import net.sf.jsqlparser.statement.merge.MergeDelete; @@ -183,6 +185,10 @@ import net.sf.jsqlparser.statement.merge.MergeOperation; import net.sf.jsqlparser.statement.merge.MergeOperationVisitor; import net.sf.jsqlparser.statement.merge.MergeUpdate; +import net.sf.jsqlparser.statement.notify.NotifyStatement; +import net.sf.jsqlparser.statement.oracle.OracleAssignment; +import net.sf.jsqlparser.statement.oracle.OracleBlock; +import net.sf.jsqlparser.statement.oracle.OracleNullStatement; import net.sf.jsqlparser.statement.piped.AggregatePipeOperator; import net.sf.jsqlparser.statement.piped.AsPipeOperator; import net.sf.jsqlparser.statement.piped.CallPipeOperator; @@ -204,6 +210,8 @@ import net.sf.jsqlparser.statement.piped.WherePipeOperator; import net.sf.jsqlparser.statement.piped.WindowPipeOperator; import net.sf.jsqlparser.statement.refresh.RefreshMaterializedViewStatement; +import net.sf.jsqlparser.statement.role.AlterRole; +import net.sf.jsqlparser.statement.role.CreateRole; import net.sf.jsqlparser.statement.select.AllColumns; import net.sf.jsqlparser.statement.select.AllTableColumns; import net.sf.jsqlparser.statement.select.FromItem; @@ -212,12 +220,12 @@ import net.sf.jsqlparser.statement.select.Join; import net.sf.jsqlparser.statement.select.LateralSubSelect; import net.sf.jsqlparser.statement.select.LateralView; +import net.sf.jsqlparser.statement.select.MatchRecognize; import net.sf.jsqlparser.statement.select.OrderByElement; import net.sf.jsqlparser.statement.select.ParenthesedFromItem; import net.sf.jsqlparser.statement.select.ParenthesedSelect; import net.sf.jsqlparser.statement.select.Pivot; import net.sf.jsqlparser.statement.select.PivotQuery; -import net.sf.jsqlparser.statement.select.UnPivotQuery; import net.sf.jsqlparser.statement.select.PivotVisitor; import net.sf.jsqlparser.statement.select.PivotXml; import net.sf.jsqlparser.statement.select.PlainSelect; @@ -227,8 +235,9 @@ import net.sf.jsqlparser.statement.select.SelectVisitor; import net.sf.jsqlparser.statement.select.SetOperationList; import net.sf.jsqlparser.statement.select.TableFunction; -import net.sf.jsqlparser.statement.select.UnPivot; import net.sf.jsqlparser.statement.select.TableStatement; +import net.sf.jsqlparser.statement.select.UnPivot; +import net.sf.jsqlparser.statement.select.UnPivotQuery; import net.sf.jsqlparser.statement.select.Values; import net.sf.jsqlparser.statement.select.WithItem; import net.sf.jsqlparser.statement.show.ShowIndexStatement; @@ -238,20 +247,6 @@ import net.sf.jsqlparser.statement.update.Update; import net.sf.jsqlparser.statement.update.UpdateSet; import net.sf.jsqlparser.statement.upsert.Upsert; -import net.sf.jsqlparser.statement.PragmaStatement; -import net.sf.jsqlparser.statement.ExtensionStatement; -import net.sf.jsqlparser.statement.AttachStatement; -import net.sf.jsqlparser.statement.DetachStatement; -import net.sf.jsqlparser.statement.ConnectStatement; -import net.sf.jsqlparser.statement.DisconnectStatement; -import net.sf.jsqlparser.statement.PrepareStatement; -import net.sf.jsqlparser.statement.DeallocateStatement; -import net.sf.jsqlparser.statement.CopyStatement; -import net.sf.jsqlparser.statement.create.macro.CreateMacro; -import net.sf.jsqlparser.statement.create.extension.CreateExtensionRepository; -import net.sf.jsqlparser.statement.AssertStatement; -import net.sf.jsqlparser.statement.export.ExportDataStatement; -import net.sf.jsqlparser.statement.load.LoadDataStatement; /** @@ -649,6 +644,9 @@ public Void visit(Column tableColumn, S context) { MatchRecognize.normalizeVariableName(tableColumn.getTable().getName()))) { visit(tableColumn.getTable(), context); } + if (tableColumn.getArrayConstructor() != null) { + tableColumn.getArrayConstructor().accept(this, context); + } return null; } @@ -686,6 +684,10 @@ public Void visit(Function function, S context) { if (exprList != null) { visit(exprList, context); } + exprList = function.getNamedParameters(); + if (exprList != null) { + visit(exprList, context); + } return null; } @@ -1020,6 +1022,9 @@ public Void visit(AnalyticExpression analytic, S context) { if (analytic.getFilterExpression() != null) { analytic.getFilterExpression().accept(this, context); } + if (analytic.getPartitionExpressionList() != null) { + visit(analytic.getPartitionExpressionList(), context); + } if (analytic.getFuncOrderBy() != null) { for (OrderByElement element : analytic.getFuncOrderBy()) { element.getExpression().accept(this, context); @@ -1126,6 +1131,12 @@ public Void visit(FromQuery fromQuery, S context) { if (fromQuery.getFromItem() != null) { fromQuery.getFromItem().accept(this, context); } + if (fromQuery.getLateralViews() != null) { + for (LateralView lateralView : fromQuery.getLateralViews()) { + lateralView.getGeneratorFunction().accept(this, context); + } + } + visitJoins(fromQuery.getJoins(), context); for (PipeOperator pipeOperator : fromQuery.getPipeOperators()) { pipeOperator.accept(this, null); } @@ -1578,9 +1589,17 @@ private void visitInsertAction(InsertConflictAction action, S context) { @Override public Void visitOutputClause(OutputClause outputClause, S context) { - if (outputClause != null && outputClause.getSelectItemList() != null) { - for (SelectItem selectItem : outputClause.getSelectItemList()) { - selectItem.accept(this, context); + if (outputClause != null) { + if (outputClause.getSelectItemList() != null) { + for (SelectItem selectItem : outputClause.getSelectItemList()) { + selectItem.accept(this, context); + } + } + if (outputClause.getOutputTable() != null) { + visit(outputClause.getOutputTable(), context); + } + if (outputClause.getTableVariable() != null) { + outputClause.getTableVariable().accept(this, context); } } return null; @@ -1902,6 +1921,7 @@ public Void visit(Merge merge, S context) { operation.accept(this, context); } } + visitOutputClause(merge.getOutputClause(), context); visitReturningClause(merge.getReturningClause(), context); return null; } @@ -2032,6 +2052,9 @@ public Void visit(Upsert upsert, S context) { if (upsert.getSelect() != null) { visit(upsert.getSelect(), context); } + if (upsert.getDuplicateAction() != null) { + visitInsertAction(upsert.getDuplicateAction(), context); + } return null; } @@ -2207,7 +2230,7 @@ public void visit(Grant grant) { @Override public Void visit(ArrayExpression array, S context) { array.getObjExpression().accept(this, context); - if (array.getStartIndexExpression() != null) { + if (array.getIndexExpression() != null) { array.getIndexExpression().accept(this, context); } if (array.getStartIndexExpression() != null) { @@ -2520,6 +2543,7 @@ public Void visit(KeyExpression keyExpression, S context) { @Override public Void visit(IfElseStatement ifElseStatement, S context) { + ifElseStatement.getCondition().accept(this, context); ifElseStatement.getIfStatement().accept(this, context); if (ifElseStatement.getElseStatement() != null) { ifElseStatement.getElseStatement().accept(this, context); diff --git a/src/test/java/net/sf/jsqlparser/util/TablesNamesFinderTraversalTest.java b/src/test/java/net/sf/jsqlparser/util/TablesNamesFinderTraversalTest.java new file mode 100644 index 000000000..b96ede8d6 --- /dev/null +++ b/src/test/java/net/sf/jsqlparser/util/TablesNamesFinderTraversalTest.java @@ -0,0 +1,241 @@ +/*- + * #%L + * JSQLParser library + * %% + * Copyright (C) 2004 - 2026 JSQLParser + * %% + * Dual licensed under GNU LGPL 2.1 or Apache License 2.0 + * #L% + */ +package net.sf.jsqlparser.util; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.stream.Stream; +import net.sf.jsqlparser.JSQLParserException; +import net.sf.jsqlparser.expression.ArrayExpression; +import net.sf.jsqlparser.expression.Expression; +import net.sf.jsqlparser.expression.JsonExpression; +import net.sf.jsqlparser.parser.CCJSqlParserUtil; +import net.sf.jsqlparser.schema.Table; +import net.sf.jsqlparser.statement.Statement; +import net.sf.jsqlparser.statement.upsert.Upsert; +import net.sf.jsqlparser.test.TestUtils; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; +import org.junit.jupiter.params.provider.ValueSource; + +class TablesNamesFinderTraversalTest { + @ParameterizedTest + @MethodSource("nestedTables") + void findsTablesInNestedChildren(String sql, List expected) throws JSQLParserException { + Statement statement = CCJSqlParserUtil.parse(sql); + assertThat(TablesNamesFinder.findTables(sql)).containsExactlyInAnyOrderElementsOf(expected); + assertThat(new TablesNamesFinder<>().getTables(statement)) + .containsExactlyInAnyOrderElementsOf(expected); + } + + private static Stream nestedTables() { + return Stream.of( + tables("SELECT a[(SELECT max(x) FROM t2)] FROM t1", "t1", "t2"), + tables("SELECT a[1][(SELECT max(x) FROM t2)] FROM t1", "t1", "t2"), + tables("SELECT f()[(SELECT max(x) FROM t2)] FROM t1", "t1", "t2"), + tables("SELECT f((SELECT x FROM t2))[(SELECT y FROM t3)] FROM t1", "t1", "t2", + "t3"), + tables("SELECT a[OFFSET(1):2] FROM t1", "t1"), + tables("SELECT f()[OFFSET(1):OFFSET(2)] FROM t1", "t1"), + tables("SELECT f()[1:] FROM t1", "t1"), + tables("SELECT f()[:2] FROM t1", "t1"), + tables("SELECT f()[:] FROM t1", "t1"), + tables("SELECT f()[OFFSET((SELECT x FROM t2)):OFFSET((SELECT y FROM t3))] FROM t1", + "t1", "t2", + "t3"), + tables("FROM t1 JOIN t2 ON t1.id=t2.id |> SELECT t1.id", "t1", "t2"), + tables("FROM t1 JOIN t2 ON t1.id=(SELECT id FROM t3) |> SELECT t1.id", "t1", "t2", + "t3"), + tables("FROM t1 LEFT JOIN (SELECT id FROM t2) q ON t1.id=q.id |> SELECT t1.id", + "t1", "t2"), + tables("FROM t1 LATERAL VIEW explode((SELECT a FROM t2)) e AS x |> SELECT x", "t1", + "t2"), + tables("WITH c AS (SELECT a FROM t2) FROM t1 JOIN c ON t1.a=c.a |> SELECT t1.a", + "t1", "t2"), + tables("UPDATE t1 SET a=1 OUTPUT inserted.* INTO t2", "t1", "t2"), + tables("INSERT INTO t1 OUTPUT inserted.* INTO t2 VALUES (1)", "t1", "t2"), + tables("DELETE t1 OUTPUT deleted.* INTO t2 FROM t1", "t1", "t2"), + tables("MERGE INTO t1 USING t2 ON t1.a=t2.a WHEN MATCHED THEN DELETE OUTPUT deleted.* INTO t3", + "t1", "t2", "t3"), + tables("UPDATE t1 SET a=1 OUTPUT (SELECT x FROM t2) INTO t3", "t1", "t2", "t3"), + tables("MERGE INTO t1 USING t2 ON t1.a=t2.a WHEN MATCHED THEN DELETE OUTPUT (SELECT x FROM t3) INTO t4", + "t1", "t2", "t3", "t4"), + tables("MERGE INTO t1 USING t2 ON t1.a=t2.a WHEN MATCHED THEN DELETE OUTPUT deleted.* INTO @tv", + "t1", "t2"), + tables("UPSERT INTO t1 VALUES (1) ON DUPLICATE KEY UPDATE (a)=(SELECT max(x) FROM t2)", + "t1", "t2"), + tables("UPSERT INTO t1 SELECT a FROM t2 ON DUPLICATE KEY UPDATE b=(SELECT x FROM t3)", + "t1", "t2", "t3"), + tables("SELECT POSITION('x' IN (SELECT max(a) FROM t2)) FROM t1", "t1", "t2"), + tables("SELECT SUBSTRING((SELECT a FROM t2) FROM (SELECT n FROM t3) FOR (SELECT n FROM t4)) FROM t1", + "t1", "t2", "t3", "t4"), + tables("IF (SELECT count(*) FROM t2)>1 DELETE FROM t3", "t2", "t3"), + tables("IF EXISTS (SELECT 1 FROM t1) UPDATE t2 SET a=1 ELSE DELETE FROM t3", "t1", + "t2", "t3"), + tables("SELECT sum(x) OVER (PARTITION BY (SELECT a FROM t2)) FROM t1", "t1", "t2"), + tables("SELECT sum(x) OVER (PARTITION BY (SELECT a FROM t2) ORDER BY (SELECT b FROM t3)) FROM t1", + "t1", "t2", "t3"), + tables("SELECT a[1], f()[1] FROM t1", "t1"), + tables("SELECT a FROM t1 JOIN t2 ON t1.a=t2.a", "t1", "t2"), + tables("SELECT * FROM t1 LATERAL VIEW explode((SELECT a FROM t2)) e AS x", "t1", + "t2"), + tables("FROM t1 |> JOIN t2 ON t1.a=t2.a |> SELECT t1.a", "t1", "t2"), + tables("UPDATE t1 SET a=1 OUTPUT inserted.*, deleted.a INTO @tv", "t1"), + tables("UPDATE t1 SET a=1 OUTPUT inserted.*", "t1"), + tables("INSERT INTO t1 VALUES (1) ON DUPLICATE KEY UPDATE a=(SELECT x FROM t2)", + "t1", "t2"), + tables("UPSERT INTO t1 VALUES (1)", "t1"), + tables("SELECT POSITION('x' IN 'xx') FROM t1", "t1"), + tables("IF a>1 DELETE FROM t3", "t3"), + tables("SELECT sum(x) OVER (ORDER BY (SELECT a FROM t2)) FROM t1", "t1", "t2")); + } + + private static Arguments tables(String sql, String... names) { + return Arguments.of(sql, List.of(names)); + } + + @ParameterizedTest + @MethodSource("expressionTables") + void findsTablesThroughExpressionEntryPoints(String sql, List expected) + throws JSQLParserException { + Expression expression = CCJSqlParserUtil.parseExpression(sql); + assertThat(TablesNamesFinder.findTablesInExpression(sql)) + .containsExactlyInAnyOrderElementsOf(expected); + assertThat(new TablesNamesFinder<>().getTables(expression)) + .containsExactlyInAnyOrderElementsOf(expected); + } + + private static Stream expressionTables() { + return Stream.of(tables("a[(SELECT x FROM t2)]", "t2"), + tables("f()[(SELECT x FROM t2)]", "t2"), + tables("f()[OFFSET((SELECT x FROM t2)):OFFSET((SELECT y FROM t3))]", "t2", "t3"), + tables("POSITION('x' IN (SELECT a FROM t2))", "t2"), + tables("sum(x) OVER (PARTITION BY (SELECT a FROM t2))", "t2"), + tables("a[1]"), tables("f()[1]"), tables("f()[:2]"), tables("f()[:]")); + } + + @ParameterizedTest + @ValueSource(ints = {0, 1, 2, 3, 4, 5, 6, 7}) + void visitsArrayOperandsOnce(int fields) throws JSQLParserException { + Expression index = (fields & 1) == 0 ? null + : CCJSqlParserUtil.parseExpression("(SELECT x FROM index_table)"); + Expression start = (fields & 2) == 0 ? null + : CCJSqlParserUtil.parseExpression("(SELECT x FROM start_table)"); + Expression stop = (fields & 4) == 0 ? null + : CCJSqlParserUtil.parseExpression("(SELECT x FROM stop_table)"); + ArrayExpression expression = new ArrayExpression( + CCJSqlParserUtil.parseExpression("f((SELECT x FROM object_table))"), index, start, + stop); + Map expected = new HashMap<>(); + expected.put("object_table", 1); + if (index != null) { + expected.put("index_table", 1); + } + if (start != null) { + expected.put("start_table", 1); + } + if (stop != null) { + expected.put("stop_table", 1); + } + CountingFinder finder = new CountingFinder(); + assertThat(finder.getTables(expression)) + .containsExactlyInAnyOrderElementsOf(expected.keySet()); + assertThat(finder.visits).isEqualTo(expected); + } + + @Test + void visitsNewChildrenWithoutRepeatingExistingChildren() throws JSQLParserException { + CountingFinder finder = new CountingFinder(); + Statement statement = CCJSqlParserUtil.parse( + "FROM t1 JOIN t2 ON t1.a=(SELECT a FROM t3) |> SELECT sum(a) OVER (PARTITION BY (SELECT a FROM t4) ORDER BY (SELECT a FROM t5))"); + assertThat(finder.getTables(statement)).containsExactlyInAnyOrder("t1", "t2", "t3", "t4", + "t5"); + assertThat(finder.visits).isEqualTo(Map.of("t1", 1, "t2", 1, "t3", 1, "t4", 1, "t5", 1)); + } + + @Test + void resetsStateBetweenStatementAndExpressionVisits() throws JSQLParserException { + TablesNamesFinder finder = new TablesNamesFinder<>(); + assertThat(finder.getTables(CCJSqlParserUtil.parse("SELECT a FROM t1"))) + .containsExactlyInAnyOrder("t1"); + assertThat(finder.getTables(CCJSqlParserUtil.parseExpression("qualified.a"))) + .containsExactlyInAnyOrder("qualified"); + assertThat(finder + .getTables(CCJSqlParserUtil.parse("UPDATE t3 SET a=1 OUTPUT inserted.* INTO @tv"))) + .containsExactlyInAnyOrder("t3"); + assertThat(finder.getTables(CCJSqlParserUtil.parse("SELECT a FROM t4"))) + .containsExactlyInAnyOrder("t4"); + } + + @Test + void visitsParsedSliceBounds() throws JSQLParserException { + String sql = "f()[OFFSET((SELECT x FROM t1)):OFFSET((SELECT y FROM t2))]"; + Expression expression = CCJSqlParserUtil.parseExpression(sql); + assertThat(expression).isInstanceOf(ArrayExpression.class); + ArrayExpression array = (ArrayExpression) expression; + assertThat(array.getIndexExpression()).isNull(); + assertThat(array.getStartIndexExpression()).isNotNull(); + assertThat(array.getStopIndexExpression()).isNotNull(); + assertThat(new TablesNamesFinder<>().getTables(array)).containsExactlyInAnyOrder("t1", + "t2"); + } + + @Test + void preservesJsonArrayIndexTraversal() throws JSQLParserException { + Expression expression = CCJSqlParserUtil.parseExpression("f()[1:2]"); + assertThat(expression).isInstanceOf(ArrayExpression.class); + assertThat(((ArrayExpression) expression).getIndexExpression()) + .isInstanceOf(JsonExpression.class); + assertThat(new TablesNamesFinder<>().getTables(expression)).isEmpty(); + } + + @Test + void findsTablesInUpsertActionCondition() throws JSQLParserException { + Upsert upsert = (Upsert) CCJSqlParserUtil + .parse("UPSERT INTO t1 VALUES (1) ON DUPLICATE KEY UPDATE a=(SELECT x FROM t2)"); + upsert.getDuplicateAction().setWhereExpression( + CCJSqlParserUtil.parseCondExpression("EXISTS (SELECT 1 FROM t3)")); + assertThat(new TablesNamesFinder<>().getTables(upsert)).containsExactlyInAnyOrder("t1", + "t2", "t3"); + } + + @Test + void preservesCteNamesInOtherSourcesEntryPoints() throws JSQLParserException { + String sql = "WITH c AS (SELECT a FROM t2) FROM t1 JOIN c ON t1.a=c.a |> SELECT t1.a"; + Statement statement = CCJSqlParserUtil.parse(sql); + assertThat(TablesNamesFinder.findTablesOrOtherSources(sql)).containsExactlyInAnyOrder("t1", + "t2", "c"); + assertThat(new TablesNamesFinder<>().getTablesOrOtherSources(statement)) + .containsExactlyInAnyOrder("t1", "t2", "c"); + } + + @Test + void preservesParseAndDeparse() throws JSQLParserException { + String sql = "SELECT a[(SELECT max(x) FROM t2)] FROM t1"; + Statement statement = TestUtils.assertSqlCanBeParsedAndDeparsed(sql); + assertThat(new TablesNamesFinder<>().getTables(statement)).containsExactlyInAnyOrder("t1", + "t2"); + } + + private static class CountingFinder extends TablesNamesFinder { + private final Map visits = new HashMap<>(); + + @Override + public Void visit(Table table, S context) { + visits.merge(table.getFullyQualifiedName(), 1, Integer::sum); + return super.visit(table, context); + } + } +}