diff --git a/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/CaseExpression.java b/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/CaseExpression.java new file mode 100644 index 00000000..81c42322 --- /dev/null +++ b/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/CaseExpression.java @@ -0,0 +1,94 @@ +package org.springframework.data.relational.core.sql; + +import java.util.ArrayList; +import java.util.List; + +import static java.util.stream.Collectors.joining; + +/** + * Case with one or more conditions expression. + *

+ * Results in a rendered condition: + *

+ *   CASE
+ *     WHEN condition1 THEN result1
+ *     WHEN condition2 THEN result2
+ *     ELSE result
+ *   END
+ * 
+ *

+ * + * @author Sven Rienstra + * @since 3.4 + */ +public class CaseExpression extends AbstractSegment implements Expression { + private final List whenList; + private final Expression elseExpression; + + private CaseExpression(List whenList, Expression elseExpression) { + + super(children(whenList, elseExpression)); + this.whenList = whenList; + this.elseExpression = elseExpression; + } + + /** + * Create CASE {@link Expression} with initial {@link When} condition. + * @param condition initial {@link When} condition + * @return the {@link CaseExpression} + */ + public static CaseExpression create(When condition) { + return new CaseExpression(List.of(condition), null); + } + + /** + * Add additional {@link When} condition + * @param condition the {@link When} condition + * @return the {@link CaseExpression} + */ + public CaseExpression when(When condition) { + List conditions = new ArrayList<>(this.whenList); + conditions.add(condition); + return new CaseExpression(conditions, elseExpression); + } + + /** + * Add ELSE clause + * @param elseExpression the {@link Expression} else value + * @return the {@link CaseExpression} + */ + public CaseExpression elseExpression(Literal elseExpression) { + return new CaseExpression(whenList, elseExpression); + } + + /** + * @return the {@link When} conditions + */ + public List getWhenList() { + return whenList; + } + + /** + * @return the ELSE {@link Literal} value + */ + public Expression getElseExpression() { + return elseExpression; + } + + @Override + public String toString() { + return "CASE " + whenList.stream().map(When::toString).collect(joining(" ")) + (elseExpression != null ? " ELSE " + elseExpression : "") + " END"; + } + + private static Segment[] children(List whenList, Expression elseExpression) { + + List segments = new ArrayList<>(); + segments.addAll(whenList); + + if (elseExpression != null) { + segments.add(elseExpression); + } + + return segments.toArray(new Segment[segments.size()]); + } +} diff --git a/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/When.java b/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/When.java new file mode 100644 index 00000000..e90338a9 --- /dev/null +++ b/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/When.java @@ -0,0 +1,53 @@ +package org.springframework.data.relational.core.sql; + +/** + * When segment for Case statement. + *

+ * Results in a rendered condition: {@code WHEN THEN }. + *

+ * + * @author Sven Rienstra + * @since 3.4 + */ +public class When extends AbstractSegment { + private final Condition condition; + private final Expression value; + + private When(Condition condition, Expression value) { + + super(condition, value); + + this.condition = condition; + this.value = value; + } + + /** + * Creates a new {@link When} given two {@link Expression} condition and {@link Literal} value. + * + * @param condition the condition {@link Expression}. + * @param value the {@link Literal} value. + * @return the {@link When}. + */ + public static When when(Condition condition, Expression value) { + return new When(condition, value); + } + + /** + * @return the condition + */ + public Condition getCondition() { + return condition; + } + + /** + * @return the value + */ + public Expression getValue() { + return value; + } + + @Override + public String toString() { + return "WHEN " + condition + " THEN " + value; + } +} diff --git a/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/CaseExpressionVisitor.java b/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/CaseExpressionVisitor.java new file mode 100644 index 00000000..8e66ccb7 --- /dev/null +++ b/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/CaseExpressionVisitor.java @@ -0,0 +1,59 @@ +package org.springframework.data.relational.core.sql.render; + +import org.springframework.data.relational.core.sql.CaseExpression; +import org.springframework.data.relational.core.sql.Literal; +import org.springframework.data.relational.core.sql.Visitable; +import org.springframework.data.relational.core.sql.When; + +/** + * Renderer for {@link CaseExpression}. + * + * @author Sven Rienstra + * @since 3.4 + */ +public class CaseExpressionVisitor extends TypedSingleConditionRenderSupport implements PartRenderer { + private final StringBuilder part = new StringBuilder(); + + CaseExpressionVisitor(RenderContext context) { + super(context); + } + + @Override + Delegation leaveNested(Visitable segment) { + + if (hasDelegatedRendering()) { + CharSequence renderedPart = consumeRenderedPart(); + + if (segment instanceof When) { + part.append(" "); + part.append(renderedPart); + } else if (segment instanceof Literal) { + part.append(" ELSE "); + part.append(renderedPart); + } + } + + return super.leaveNested(segment); + } + + @Override + Delegation enterMatched(CaseExpression segment) { + + part.append("CASE"); + + return super.enterMatched(segment); + } + + @Override + Delegation leaveMatched(CaseExpression segment) { + + part.append(" END"); + + return super.leaveMatched(segment); + } + + @Override + public CharSequence getRenderedPart() { + return part; + } +} diff --git a/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/ExpressionVisitor.java b/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/ExpressionVisitor.java index 67e648fb..c3bdcebc 100644 --- a/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/ExpressionVisitor.java +++ b/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/ExpressionVisitor.java @@ -108,6 +108,10 @@ class ExpressionVisitor extends TypedSubtreeVisitor implements PartR CastVisitor visitor = new CastVisitor(context); partRenderer = visitor; return Delegation.delegateTo(visitor); + } else if (segment instanceof CaseExpression) { + CaseExpressionVisitor visitor = new CaseExpressionVisitor(context); + partRenderer = visitor; + return Delegation.delegateTo(visitor); } else { // works for literals and just and possibly more value = segment.toString(); diff --git a/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/OrderByClauseVisitor.java b/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/OrderByClauseVisitor.java index 81a850bb..b536b899 100644 --- a/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/OrderByClauseVisitor.java +++ b/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/OrderByClauseVisitor.java @@ -15,7 +15,9 @@ */ package org.springframework.data.relational.core.sql.render; + import org.springframework.data.relational.core.sql.Column; +import org.springframework.data.relational.core.sql.CaseExpression; import org.springframework.data.relational.core.sql.Expressions; import org.springframework.data.relational.core.sql.OrderByField; import org.springframework.data.relational.core.sql.SimpleFunction; @@ -83,7 +85,7 @@ class OrderByClauseVisitor extends TypedSubtreeVisitor implements return Delegation.delegateTo((SimpleFunctionVisitor)delegate); } - if (segment instanceof Expressions.SimpleExpression) { + if (segment instanceof Expressions.SimpleExpression || segment instanceof CaseExpression) { delegate = new ExpressionVisitor(context); return Delegation.delegateTo((ExpressionVisitor)delegate); } diff --git a/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/TypedSingleConditionRenderSupport.java b/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/TypedSingleConditionRenderSupport.java index 932ecef1..04b87f08 100644 --- a/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/TypedSingleConditionRenderSupport.java +++ b/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/TypedSingleConditionRenderSupport.java @@ -18,6 +18,7 @@ package org.springframework.data.relational.core.sql.render; import org.springframework.data.relational.core.sql.Condition; import org.springframework.data.relational.core.sql.Expression; import org.springframework.data.relational.core.sql.Visitable; +import org.springframework.data.relational.core.sql.When; import org.springframework.lang.Nullable; import org.springframework.util.Assert; @@ -26,6 +27,7 @@ import org.springframework.util.Assert; * delegate nested {@link Expression} and {@link Condition} rendering. * * @author Mark Paluch + * @author Sven Rienstra * @since 1.1 */ abstract class TypedSingleConditionRenderSupport extends TypedSubtreeVisitor { @@ -40,8 +42,8 @@ abstract class TypedSingleConditionRenderSupport extends Ty @Override Delegation enterNested(Visitable segment) { - if (segment instanceof Expression) { - ExpressionVisitor visitor = new ExpressionVisitor(context); + if (segment instanceof When) { + WhenVisitor visitor = new WhenVisitor(context); current = visitor; return Delegation.delegateTo(visitor); } @@ -52,6 +54,12 @@ abstract class TypedSingleConditionRenderSupport extends Ty return Delegation.delegateTo(visitor); } + if (segment instanceof Expression) { + ExpressionVisitor visitor = new ExpressionVisitor(context); + current = visitor; + return Delegation.delegateTo(visitor); + } + throw new IllegalStateException("Cannot provide visitor for " + segment); } diff --git a/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/WhenVisitor.java b/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/WhenVisitor.java new file mode 100644 index 00000000..70e8c1da --- /dev/null +++ b/spring-data-relational/src/main/java/org/springframework/data/relational/core/sql/render/WhenVisitor.java @@ -0,0 +1,48 @@ +package org.springframework.data.relational.core.sql.render; + +import org.springframework.data.relational.core.sql.Visitable; +import org.springframework.data.relational.core.sql.When; + +/** + * Renderer for {@link When} segments. + * + * @author Sven Rienstra + * @since 3.4 + */ +public class WhenVisitor extends TypedSingleConditionRenderSupport implements PartRenderer { + private final StringBuilder part = new StringBuilder(); + private boolean conditionRendered; + + WhenVisitor(RenderContext context) { + super(context); + } + + @Override + Delegation leaveNested(Visitable segment) { + + if (hasDelegatedRendering()) { + + if (conditionRendered) { + part.append(" THEN "); + } + + part.append(consumeRenderedPart()); + conditionRendered = true; + } + + return super.leaveNested(segment); + } + + @Override + Delegation enterMatched(When segment) { + + part.append("WHEN "); + + return super.enterMatched(segment); + } + + @Override + public CharSequence getRenderedPart() { + return part; + } +} diff --git a/spring-data-relational/src/test/java/org/springframework/data/relational/core/sql/render/OrderByClauseVisitorUnitTests.java b/spring-data-relational/src/test/java/org/springframework/data/relational/core/sql/render/OrderByClauseVisitorUnitTests.java index e04fea49..baadc8b0 100644 --- a/spring-data-relational/src/test/java/org/springframework/data/relational/core/sql/render/OrderByClauseVisitorUnitTests.java +++ b/spring-data-relational/src/test/java/org/springframework/data/relational/core/sql/render/OrderByClauseVisitorUnitTests.java @@ -18,14 +18,7 @@ package org.springframework.data.relational.core.sql.render; import static org.assertj.core.api.Assertions.*; import org.junit.jupiter.api.Test; -import org.springframework.data.relational.core.sql.Column; -import org.springframework.data.relational.core.sql.Expression; -import org.springframework.data.relational.core.sql.Expressions; -import org.springframework.data.relational.core.sql.OrderByField; -import org.springframework.data.relational.core.sql.SQL; -import org.springframework.data.relational.core.sql.Select; -import org.springframework.data.relational.core.sql.SimpleFunction; -import org.springframework.data.relational.core.sql.Table; +import org.springframework.data.relational.core.sql.*; import java.util.Arrays; import java.util.List; @@ -129,4 +122,18 @@ class OrderByClauseVisitorUnitTests { assertThat(visitor.getRenderedPart().toString()).isEqualTo("1 ASC"); } + + @Test + void shouldRenderOrderByCase() { + Table employee = SQL.table("employee").as("emp"); + Column column = employee.column("name"); + + CaseExpression caseExpression = CaseExpression.create(When.when(column.isNull(), SQL.literalOf(1))).elseExpression(SQL.literalOf(2)); + Select select = Select.builder().select(column).from(employee).orderBy(OrderByField.from(caseExpression).asc()).build(); + + OrderByClauseVisitor visitor = new OrderByClauseVisitor(new SimpleRenderContext(NamingStrategies.asIs())); + select.visit(visitor); + + assertThat(visitor.getRenderedPart().toString()).isEqualTo("CASE WHEN emp.name IS NULL THEN 1 ELSE 2 END ASC"); + } } diff --git a/spring-data-relational/src/test/java/org/springframework/data/relational/core/sql/render/SelectRendererUnitTests.java b/spring-data-relational/src/test/java/org/springframework/data/relational/core/sql/render/SelectRendererUnitTests.java index 68842f50..1251aefd 100644 --- a/spring-data-relational/src/test/java/org/springframework/data/relational/core/sql/render/SelectRendererUnitTests.java +++ b/spring-data-relational/src/test/java/org/springframework/data/relational/core/sql/render/SelectRendererUnitTests.java @@ -688,6 +688,23 @@ class SelectRendererUnitTests { assertThat(rendered).isEqualTo("SELECT e.*, e.id FROM employee e"); } + @Test + void rendersCaseExpression() { + Table table = SQL.table("table"); + Column column = table.column("name"); + + CaseExpression caseExpression = CaseExpression.create(When.when(column.isNull(), SQL.literalOf(1))) // + .when(When.when(column.isNotNull(), SQL.literalOf(2))) // + .elseExpression(SQL.literalOf(3)); + + Select select = StatementBuilder.select(caseExpression) // + .from(table) // + .build(); + + String rendered = SqlRenderer.toString(select); + assertThat(rendered).isEqualTo("SELECT CASE WHEN table.name IS NULL THEN 1 WHEN table.name IS NOT NULL THEN 2 ELSE 3 END FROM table"); + } + /** * Tests the rendering of analytic functions. */