11package io .substrait .isthmus ;
22
3+ import static org .junit .jupiter .api .Assertions .assertEquals ;
4+ import static org .junit .jupiter .api .Assertions .assertSame ;
35import static org .junit .jupiter .api .Assertions .assertTrue ;
46
57import com .google .protobuf .util .JsonFormat ;
68import io .substrait .proto .Expression ;
9+ import io .substrait .proto .Expression .Subquery .SetPredicate .PredicateOp ;
710import io .substrait .proto .FilterRel ;
811import io .substrait .proto .Plan ;
912import java .io .IOException ;
1013import org .apache .calcite .sql .parser .SqlParseException ;
11- import org .junit .jupiter .api .Assertions ;
1214import org .junit .jupiter .api .Test ;
1315
1416public class SubqueryPlanTest extends PlanTestBase {
1517 // TODO: Add a roundtrip test once the ProtoRelConverter is committed and updated to support
1618 // subqueries
19+
1720 @ Test
1821 public void existsCorrelatedSubquery () throws SqlParseException {
1922 SqlToSubstrait s = new SqlToSubstrait ();
@@ -34,9 +37,7 @@ public void existsCorrelatedSubquery() throws SqlParseException {
3437 .getSubquery ();
3538
3639 assertTrue (subquery .hasSetPredicate ());
37- assertTrue (
38- subquery .getSetPredicate ().getPredicateOp ()
39- == Expression .Subquery .SetPredicate .PredicateOp .PREDICATE_OP_EXISTS );
40+ assertSame (PredicateOp .PREDICATE_OP_EXISTS , subquery .getSetPredicate ().getPredicateOp ());
4041
4142 FilterRel setPredicateFilter =
4243 subquery
@@ -52,8 +53,8 @@ public void existsCorrelatedSubquery() throws SqlParseException {
5253 .getValue ()
5354 .getSelection (); // l_orderkey
5455
55- assertTrue ( correlatedCol .getDirectReference ().getStructField ().getField () == 0 );
56- assertTrue ( correlatedCol .getOuterReference ().getStepsOut () == 1 );
56+ assertEquals ( 0 , correlatedCol .getDirectReference ().getStructField ().getField ());
57+ assertEquals ( 1 , correlatedCol .getOuterReference ().getStepsOut ());
5758 }
5859
5960 @ Test
@@ -85,9 +86,7 @@ public void uniqueCorrelatedSubquery() throws IOException, SqlParseException {
8586 .getFilter (); // unique (select ... from orders where o_orderkey = l_orderkey)
8687
8788 assertTrue (subquery .hasSetPredicate ());
88- assertTrue (
89- subquery .getSetPredicate ().getPredicateOp ()
90- == Expression .Subquery .SetPredicate .PredicateOp .PREDICATE_OP_UNIQUE );
89+ assertSame (PredicateOp .PREDICATE_OP_UNIQUE , subquery .getSetPredicate ().getPredicateOp ());
9190
9291 Expression .FieldReference correlatedCol =
9392 setPredicateFilter
@@ -97,8 +96,8 @@ public void uniqueCorrelatedSubquery() throws IOException, SqlParseException {
9796 .getValue ()
9897 .getSelection (); // l_orderkey
9998
100- assertTrue ( correlatedCol .getDirectReference ().getStructField ().getField () == 0 );
101- assertTrue ( correlatedCol .getOuterReference ().getStepsOut () == 1 );
99+ assertEquals ( 0 , correlatedCol .getDirectReference ().getStructField ().getField ());
100+ assertEquals ( 1 , correlatedCol .getOuterReference ().getStepsOut ());
102101 }
103102
104103 @ Test
@@ -135,8 +134,8 @@ public void inPredicateCorrelatedSubQuery() throws IOException, SqlParseExceptio
135134 .getValue ()
136135 .getSelection (); // l_partkey
137136
138- assertTrue ( correlatedCol .getDirectReference ().getStructField ().getField () == 1 );
139- assertTrue ( correlatedCol .getOuterReference ().getStepsOut () == 1 );
137+ assertEquals ( 1 , correlatedCol .getDirectReference ().getStructField ().getField ());
138+ assertEquals ( 1 , correlatedCol .getOuterReference ().getStepsOut ());
140139 }
141140
142141 @ Test
@@ -175,8 +174,8 @@ public void notInPredicateCorrelatedSubquery() throws IOException, SqlParseExcep
175174 .getValue ()
176175 .getSelection (); // l_partkey
177176
178- assertTrue ( correlatedCol .getDirectReference ().getStructField ().getField () == 1 );
179- assertTrue ( correlatedCol .getOuterReference ().getStepsOut () == 1 );
177+ assertEquals ( 1 , correlatedCol .getDirectReference ().getStructField ().getField ());
178+ assertEquals ( 1 , correlatedCol .getOuterReference ().getStepsOut ());
180179 }
181180
182181 @ Test
@@ -207,9 +206,7 @@ public void existsNestedCorrelatedSubquery() throws IOException, SqlParseExcepti
207206 .getSubquery ();
208207
209208 assertTrue (outer_subquery .hasSetPredicate ());
210- assertTrue (
211- outer_subquery .getSetPredicate ().getPredicateOp ()
212- == Expression .Subquery .SetPredicate .PredicateOp .PREDICATE_OP_EXISTS );
209+ assertSame (PredicateOp .PREDICATE_OP_EXISTS , outer_subquery .getSetPredicate ().getPredicateOp ());
213210
214211 FilterRel exists_filter =
215212 outer_subquery
@@ -221,9 +218,7 @@ public void existsNestedCorrelatedSubquery() throws IOException, SqlParseExcepti
221218 exists_filter .getCondition ().getScalarFunction ().getArguments (1 ).getValue ().getSubquery ();
222219 assertTrue (inner_subquery .hasSetPredicate ());
223220
224- assertTrue (
225- inner_subquery .getSetPredicate ().getPredicateOp ()
226- == Expression .Subquery .SetPredicate .PredicateOp .PREDICATE_OP_UNIQUE );
221+ assertSame (PredicateOp .PREDICATE_OP_UNIQUE , inner_subquery .getSetPredicate ().getPredicateOp ());
227222
228223 Expression inner_subquery_condition =
229224 inner_subquery
@@ -251,17 +246,17 @@ public void existsNestedCorrelatedSubquery() throws IOException, SqlParseExcepti
251246 .getArguments (1 )
252247 .getValue ()
253248 .getSelection (); // p.p_partkey
254- assertTrue ( correlatedCol1 .getDirectReference ().getStructField ().getField () == 0 );
255- assertTrue ( correlatedCol1 .getOuterReference ().getStepsOut () == 2 );
249+ assertEquals ( 0 , correlatedCol1 .getDirectReference ().getStructField ().getField ());
250+ assertEquals ( 2 , correlatedCol1 .getOuterReference ().getStepsOut ());
256251
257252 Expression .FieldReference correlatedCol2 =
258253 inner_subquery_cond2
259254 .getScalarFunction ()
260255 .getArguments (1 )
261256 .getValue ()
262257 .getSelection (); // l.l_suppkey
263- assertTrue ( correlatedCol2 .getDirectReference ().getStructField ().getField () == 2 );
264- assertTrue ( correlatedCol2 .getOuterReference ().getStepsOut () == 1 );
258+ assertEquals ( 2 , correlatedCol2 .getDirectReference ().getStructField ().getField ());
259+ assertEquals ( 1 , correlatedCol2 .getOuterReference ().getStepsOut ());
265260 }
266261
267262 @ Test
@@ -326,27 +321,22 @@ public void nestedScalarCorrelatedSubquery() throws IOException, SqlParseExcepti
326321 .getArguments (1 )
327322 .getValue ()
328323 .getSelection (); // p.p_partkey
329- assertTrue ( correlatedCol1 .getDirectReference ().getStructField ().getField () == 0 );
330- assertTrue ( correlatedCol1 .getOuterReference ().getStepsOut () == 2 );
324+ assertEquals ( 0 , correlatedCol1 .getDirectReference ().getStructField ().getField ());
325+ assertEquals ( 2 , correlatedCol1 .getOuterReference ().getStepsOut ());
331326
332327 Expression .FieldReference correlatedCol2 =
333328 inner_subquery_cond2
334329 .getScalarFunction ()
335330 .getArguments (1 )
336331 .getValue ()
337332 .getSelection (); // l.l_suppkey
338- assertTrue ( correlatedCol2 .getDirectReference ().getStructField ().getField () == 2 );
339- assertTrue ( correlatedCol2 .getOuterReference ().getStepsOut () == 1 );
333+ assertEquals ( 2 , correlatedCol2 .getDirectReference ().getStructField ().getField ());
334+ assertEquals ( 1 , correlatedCol2 .getOuterReference ().getStepsOut ());
340335 }
341336
342337 @ Test
343- public void correlatedScalarSubQInSelect () throws IOException {
344- SqlToSubstrait s = new SqlToSubstrait ();
338+ public void correlatedScalarSubQueryInSelect () throws Exception {
345339 String sql = asString ("subquery/nested_scalar_subquery_in_select.sql" );
346- Assertions .assertThrows (
347- UnsupportedOperationException .class ,
348- () -> {
349- s .convert (sql , TPCH_CATALOG );
350- });
340+ assertSqlSubstraitRelRoundTrip (sql );
351341 }
352342}
0 commit comments