diff --git a/src/main/java/com/avaje/ebeaninternal/server/expression/InExpression.java b/src/main/java/com/avaje/ebeaninternal/server/expression/InExpression.java index b9ae17902..89ad86c40 100644 --- a/src/main/java/com/avaje/ebeaninternal/server/expression/InExpression.java +++ b/src/main/java/com/avaje/ebeaninternal/server/expression/InExpression.java @@ -1,35 +1,54 @@ package com.avaje.ebeaninternal.server.expression; import com.avaje.ebean.bean.EntityBean; +import com.avaje.ebean.event.BeanQueryRequest; import com.avaje.ebeaninternal.api.HashQueryPlanBuilder; import com.avaje.ebeaninternal.api.SpiExpression; import com.avaje.ebeaninternal.api.SpiExpressionRequest; import com.avaje.ebeaninternal.server.el.ElPropertyValue; import java.io.IOException; +import java.util.ArrayList; +import java.util.Arrays; import java.util.Collection; +import java.util.List; class InExpression extends AbstractExpression { private final boolean not; - private final Object[] values; + private final Collection sourceValues; - InExpression(String propertyName, Collection coll, boolean not) { + private Object[] bindValues; + + InExpression(String propertyName, Collection sourceValues, boolean not) { super(propertyName); - this.values = coll.toArray(new Object[coll.size()]); + this.sourceValues = sourceValues; this.not = not; } InExpression(String propertyName, Object[] array, boolean not) { super(propertyName); - this.values = array; + this.sourceValues = Arrays.asList(array); this.not = not; } + private Object[] values() { + List vals = new ArrayList(); + for (Object sourceValue : sourceValues) { + NamedParamHelp.valueAdd(vals, sourceValue); + } + return vals.toArray(); + } + + @Override + public void prepareExpression(BeanQueryRequest request) { + bindValues = values(); + } + @Override public void writeDocQuery(DocQueryContext context) throws IOException { - context.writeIn(propName, values, not); + context.writeIn(propName, values(), not); } @Override @@ -40,13 +59,13 @@ class InExpression extends AbstractExpression { prop = null; } - for (int i = 0; i < values.length; i++) { + for (int i = 0; i < bindValues.length; i++) { if (prop == null) { - request.addBindValue(values[i]); + request.addBindValue(bindValues[i]); } else { // extract the id values from the bean - Object[] ids = prop.getAssocIdValues((EntityBean) values[i]); + Object[] ids = prop.getAssocIdValues((EntityBean) bindValues[i]); if (ids != null) { for (int j = 0; j < ids.length; j++) { request.addBindValue(ids[j]); @@ -59,7 +78,7 @@ class InExpression extends AbstractExpression { @Override public void addSql(SpiExpressionRequest request) { - if (values.length == 0) { + if (bindValues.length == 0) { String expr = not ? "1=1" : "1=0"; request.append(expr); return; @@ -72,7 +91,7 @@ class InExpression extends AbstractExpression { if (prop != null) { request.append(prop.getAssocIdInExpr(propName)); - String inClause = prop.getAssocIdInValueExpr(values.length); + String inClause = prop.getAssocIdInValueExpr(bindValues.length); request.append(inClause); } else { @@ -81,7 +100,7 @@ class InExpression extends AbstractExpression { request.append(" not"); } request.append(" in (?"); - for (int i = 1; i < values.length; i++) { + for (int i = 1; i < bindValues.length; i++) { request.append(", ").append("?"); } @@ -94,15 +113,15 @@ class InExpression extends AbstractExpression { */ @Override public void queryPlanHash(HashQueryPlanBuilder builder) { - builder.add(InExpression.class).add(propName).add(values.length).add(not); - builder.bind(values.length); + builder.add(InExpression.class).add(propName).add(bindValues.length).add(not); + builder.bind(bindValues.length); } @Override public int queryBindHash() { int hc = 31; - for (int i = 0; i < values.length; i++) { - hc = 31 * hc + values[i].hashCode(); + for (int i = 0; i < bindValues.length; i++) { + hc = 31 * hc + bindValues[i].hashCode(); } return hc; } @@ -116,17 +135,17 @@ class InExpression extends AbstractExpression { InExpression that = (InExpression) other; return propName.equals(that.propName) && not == that.not - && values.length == that.values.length; + && bindValues.length == that.bindValues.length; } @Override public boolean isSameByBind(SpiExpression other) { InExpression that = (InExpression) other; - if (this.values.length != that.values.length) { + if (this.bindValues.length != that.bindValues.length) { return false; } - for (int i = 0; i < values.length; i++) { - if (!values[i].equals(that.values[i])) { + for (int i = 0; i < bindValues.length; i++) { + if (!bindValues[i].equals(that.bindValues[i])) { return false; } } diff --git a/src/main/java/com/avaje/ebeaninternal/server/expression/NamedParamHelp.java b/src/main/java/com/avaje/ebeaninternal/server/expression/NamedParamHelp.java index 97df6e828..1defe77e5 100644 --- a/src/main/java/com/avaje/ebeaninternal/server/expression/NamedParamHelp.java +++ b/src/main/java/com/avaje/ebeaninternal/server/expression/NamedParamHelp.java @@ -2,6 +2,9 @@ package com.avaje.ebeaninternal.server.expression; import com.avaje.ebeaninternal.api.SpiNamedParam; +import java.util.Collection; +import java.util.List; + /** * Helper for evaluating named parameters. */ @@ -25,4 +28,16 @@ class NamedParamHelp { return (value == null) ? null : value.toString(); } + /** + * Add the potentially named parameter(s) to the values. + */ + public static void valueAdd(List values, Object sourceValue) { + + Object value = value(sourceValue); + if (value instanceof Collection) { + values.addAll((Collection)value); + } else { + values.add(value); + } + } } diff --git a/src/main/java/com/avaje/ebeaninternal/server/grammer/EqlAdapter.java b/src/main/java/com/avaje/ebeaninternal/server/grammer/EqlAdapter.java index fe04c9073..7c02fbe99 100644 --- a/src/main/java/com/avaje/ebeaninternal/server/grammer/EqlAdapter.java +++ b/src/main/java/com/avaje/ebeaninternal/server/grammer/EqlAdapter.java @@ -3,7 +3,6 @@ package com.avaje.ebeaninternal.server.grammer; import com.avaje.ebean.Expression; import com.avaje.ebean.ExpressionList; import com.avaje.ebean.LikeType; -import com.avaje.ebean.Query; import com.avaje.ebeaninternal.api.SpiQuery; import com.avaje.ebeaninternal.server.grammer.antlr.EQLBaseListener; import com.avaje.ebeaninternal.server.grammer.antlr.EQLLexer; @@ -13,6 +12,9 @@ import org.antlr.v4.runtime.ParserRuleContext; import org.antlr.v4.runtime.tree.ParseTree; import org.antlr.v4.runtime.tree.TerminalNode; +import java.util.ArrayList; +import java.util.List; + class EqlAdapter extends EQLBaseListener { private static final OperatorMapping operatorMapping = new OperatorMapping(); @@ -27,6 +29,10 @@ class EqlAdapter extends EQLBaseListener { private boolean textMode; + private List inValues; + + private String inPropertyName; + public EqlAdapter(SpiQuery query) { this.query = query; this.helper = new EqlAdapterHelper(this); @@ -100,7 +106,25 @@ class EqlAdapter extends EQLBaseListener { @Override public void enterIn_expression(EQLParser.In_expressionContext ctx) { + this.inValues = new ArrayList(); + this.inPropertyName = getLeftHandSidePath(ctx); + } + @Override + public void enterIn_value(EQLParser.In_valueContext ctx) { + int childCount = ctx.getChildCount(); + for (int i = 0; i < childCount; i++) { + ParseTree child = ctx.getChild(i); + String text = child.getText(); + if (!text.equals("(") && !text.equals(")")) { + inValues.add(helper.bind(text)); + } + } + } + + @Override + public void exitIn_expression(EQLParser.In_expressionContext ctx) { + helper.addIn(inPropertyName, inValues); } @Override diff --git a/src/main/java/com/avaje/ebeaninternal/server/grammer/EqlAdapterHelper.java b/src/main/java/com/avaje/ebeaninternal/server/grammer/EqlAdapterHelper.java index 203233b3f..ce49f5dfe 100644 --- a/src/main/java/com/avaje/ebeaninternal/server/grammer/EqlAdapterHelper.java +++ b/src/main/java/com/avaje/ebeaninternal/server/grammer/EqlAdapterHelper.java @@ -4,6 +4,7 @@ import com.avaje.ebean.ExpressionList; import com.avaje.ebean.LikeType; import java.math.BigDecimal; +import java.util.List; class EqlAdapterHelper { @@ -40,6 +41,10 @@ class EqlAdapterHelper { } } + @SuppressWarnings("unchecked") + protected void addIn(String path, List inValues) { + peekExprList().in(path, inValues); + } protected void addExpression(String path, EqlOperator op, String value) { @@ -99,7 +104,7 @@ class EqlAdapterHelper { peekExprList().add(owner.like(caseInsensitive, likeType, path, bindValue)); } - private Object bind(String value) { + protected Object bind(String value) { ValueType valueType = getValueType(value); return getBindValue(valueType, value); } diff --git a/src/test/java/com/avaje/ebeaninternal/server/grammer/EqlParserTest.java b/src/test/java/com/avaje/ebeaninternal/server/grammer/EqlParserTest.java index 5a246c4d0..85d0dcc62 100644 --- a/src/test/java/com/avaje/ebeaninternal/server/grammer/EqlParserTest.java +++ b/src/test/java/com/avaje/ebeaninternal/server/grammer/EqlParserTest.java @@ -6,6 +6,8 @@ import com.avaje.ebeaninternal.api.SpiQuery; import com.avaje.tests.model.basic.Customer; import org.junit.Test; +import java.util.Arrays; + import static org.assertj.core.api.Assertions.assertThat; public class EqlParserTest { @@ -92,6 +94,36 @@ public class EqlParserTest { } + @Test + public void where_in() throws Exception { + + Query query = parse("where name in ('Rob','Jim')"); + query.findList(); + + assertThat(query.getGeneratedSql()).contains("where t0.name in (?, ? )"); + } + + @Test + public void where_in_when_namedParams() throws Exception { + + Query query = parse("where name in (:one, :two)"); + query.setParameter("one", "Foo"); + query.setParameter("two", "Bar"); + query.findList(); + + assertThat(query.getGeneratedSql()).contains("where t0.name in (?, ? )"); + } + + @Test + public void where_in_when_namedParamAsList() throws Exception { + + Query query = parse("where name in (:names)"); + query.setParameter("names", Arrays.asList("Baz","Maz","Jim")); + query.findList(); + + assertThat(query.getGeneratedSql()).contains("where t0.name in (?, ?, ? )"); + } + private Query parse(String raw) { Query query = Ebean.find(Customer.class);