Repository: spark
Updated Branches:
  refs/heads/master 119f45d61 -> d74373225


[SPARK-5097][SQL] Test cases for DataFrame expressions.

Author: Reynold Xin <[email protected]>

Closes #4235 from rxin/df-tests1 and squashes the following commits:

f341db6 [Reynold Xin] [SPARK-5097][SQL] Test cases for DataFrame expressions.


Project: http://git-wip-us.apache.org/repos/asf/spark/repo
Commit: http://git-wip-us.apache.org/repos/asf/spark/commit/d7437322
Tree: http://git-wip-us.apache.org/repos/asf/spark/tree/d7437322
Diff: http://git-wip-us.apache.org/repos/asf/spark/diff/d7437322

Branch: refs/heads/master
Commit: d74373225ef78cabd6b76830439d6b4936b0c4a6
Parents: 119f45d
Author: Reynold Xin <[email protected]>
Authored: Tue Jan 27 18:10:49 2015 -0800
Committer: Reynold Xin <[email protected]>
Committed: Tue Jan 27 18:10:49 2015 -0800

----------------------------------------------------------------------
 .../spark/sql/catalyst/expressions/rows.scala   |   1 -
 .../scala/org/apache/spark/sql/Column.scala     |   5 +-
 .../scala/org/apache/spark/sql/DataFrame.scala  |   9 +-
 .../org/apache/spark/sql/dsl/package.scala      |   1 +
 .../spark/sql/ColumnExpressionSuite.scala       | 302 ++++++++++++++++
 .../org/apache/spark/sql/DataFrameSuite.scala   | 280 +++++++++++++++
 .../org/apache/spark/sql/DslQuerySuite.scala    | 346 -------------------
 .../scala/org/apache/spark/sql/TestData.scala   |   2 +-
 8 files changed, 594 insertions(+), 352 deletions(-)
----------------------------------------------------------------------


http://git-wip-us.apache.org/repos/asf/spark/blob/d7437322/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/rows.scala
----------------------------------------------------------------------
diff --git 
a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/rows.scala
 
b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/rows.scala
index 8df150e..73ec7a6 100644
--- 
a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/rows.scala
+++ 
b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/rows.scala
@@ -114,7 +114,6 @@ class GenericRow(protected[sql] val values: Array[Any]) 
extends Row {
   }
 
   override def getString(i: Int): String = {
-    if (values(i) == null) sys.error("Failed to check null bit for primitive 
String value.")
     values(i).asInstanceOf[String]
   }
 

http://git-wip-us.apache.org/repos/asf/spark/blob/d7437322/sql/core/src/main/scala/org/apache/spark/sql/Column.scala
----------------------------------------------------------------------
diff --git a/sql/core/src/main/scala/org/apache/spark/sql/Column.scala 
b/sql/core/src/main/scala/org/apache/spark/sql/Column.scala
index 7fc8347..7f20cf8 100644
--- a/sql/core/src/main/scala/org/apache/spark/sql/Column.scala
+++ b/sql/core/src/main/scala/org/apache/spark/sql/Column.scala
@@ -252,7 +252,10 @@ class Column(
   /**
    * Equality test with an expression that is safe for null values.
    */
-  override def <=> (other: Column): Column = EqualNullSafe(expr, other.expr)
+  override def <=> (other: Column): Column = other match {
+    case null => EqualNullSafe(expr, Literal.anyToLiteral(null).expr)
+    case _ => EqualNullSafe(expr, other.expr)
+  }
 
   /**
    * Equality test with a literal value that is safe for null values.

http://git-wip-us.apache.org/repos/asf/spark/blob/d7437322/sql/core/src/main/scala/org/apache/spark/sql/DataFrame.scala
----------------------------------------------------------------------
diff --git a/sql/core/src/main/scala/org/apache/spark/sql/DataFrame.scala 
b/sql/core/src/main/scala/org/apache/spark/sql/DataFrame.scala
index d0bb364..3198215 100644
--- a/sql/core/src/main/scala/org/apache/spark/sql/DataFrame.scala
+++ b/sql/core/src/main/scala/org/apache/spark/sql/DataFrame.scala
@@ -230,9 +230,12 @@ class DataFrame protected[sql](
   /**
    * Selecting a single column and return it as a [[Column]].
    */
-  override def apply(colName: String): Column = {
-    val expr = resolve(colName)
-    new Column(Some(sqlContext), Some(Project(Seq(expr), logicalPlan)), expr)
+  override def apply(colName: String): Column = colName match {
+    case "*" =>
+      Column("*")
+    case _ =>
+      val expr = resolve(colName)
+      new Column(Some(sqlContext), Some(Project(Seq(expr), logicalPlan)), expr)
   }
 
   /**

http://git-wip-us.apache.org/repos/asf/spark/blob/d7437322/sql/core/src/main/scala/org/apache/spark/sql/dsl/package.scala
----------------------------------------------------------------------
diff --git a/sql/core/src/main/scala/org/apache/spark/sql/dsl/package.scala 
b/sql/core/src/main/scala/org/apache/spark/sql/dsl/package.scala
index 29c3d26..4c44e17 100644
--- a/sql/core/src/main/scala/org/apache/spark/sql/dsl/package.scala
+++ b/sql/core/src/main/scala/org/apache/spark/sql/dsl/package.scala
@@ -53,6 +53,7 @@ package object dsl {
   def last(e: Column): Column = Last(e.expr)
   def min(e: Column): Column = Min(e.expr)
   def max(e: Column): Column = Max(e.expr)
+
   def upper(e: Column): Column = Upper(e.expr)
   def lower(e: Column): Column = Lower(e.expr)
   def sqrt(e: Column): Column = Sqrt(e.expr)

http://git-wip-us.apache.org/repos/asf/spark/blob/d7437322/sql/core/src/test/scala/org/apache/spark/sql/ColumnExpressionSuite.scala
----------------------------------------------------------------------
diff --git 
a/sql/core/src/test/scala/org/apache/spark/sql/ColumnExpressionSuite.scala 
b/sql/core/src/test/scala/org/apache/spark/sql/ColumnExpressionSuite.scala
new file mode 100644
index 0000000..825a186
--- /dev/null
+++ b/sql/core/src/test/scala/org/apache/spark/sql/ColumnExpressionSuite.scala
@@ -0,0 +1,302 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *    http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import org.apache.spark.sql.dsl._
+import org.apache.spark.sql.test.TestSQLContext
+import org.apache.spark.sql.types.{BooleanType, IntegerType, StructField, 
StructType}
+
+
+class ColumnExpressionSuite extends QueryTest {
+  import org.apache.spark.sql.TestData._
+
+  // TODO: Add test cases for bitwise operations.
+
+  test("star") {
+    checkAnswer(testData.select($"*"), testData.collect().toSeq)
+  }
+
+  ignore("star qualified by data frame object") {
+    // This is not yet supported.
+    val df = testData.toDF
+    checkAnswer(df.select(df("*")), df.collect().toSeq)
+  }
+
+  test("star qualified by table name") {
+    checkAnswer(testData.as("testData").select($"testData.*"), 
testData.collect().toSeq)
+  }
+
+  test("+") {
+    checkAnswer(
+      testData2.select($"a" + 1),
+      testData2.collect().toSeq.map(r => Row(r.getInt(0) + 1)))
+
+    checkAnswer(
+      testData2.select($"a" + $"b" + 2),
+      testData2.collect().toSeq.map(r => Row(r.getInt(0) + r.getInt(1) + 2)))
+  }
+
+  test("-") {
+    checkAnswer(
+      testData2.select($"a" - 1),
+      testData2.collect().toSeq.map(r => Row(r.getInt(0) - 1)))
+
+    checkAnswer(
+      testData2.select($"a" - $"b" - 2),
+      testData2.collect().toSeq.map(r => Row(r.getInt(0) - r.getInt(1) - 2)))
+  }
+
+  test("*") {
+    checkAnswer(
+      testData2.select($"a" * 10),
+      testData2.collect().toSeq.map(r => Row(r.getInt(0) * 10)))
+
+    checkAnswer(
+      testData2.select($"a" * $"b"),
+      testData2.collect().toSeq.map(r => Row(r.getInt(0) * r.getInt(1))))
+  }
+
+  test("/") {
+    checkAnswer(
+      testData2.select($"a" / 2),
+      testData2.collect().toSeq.map(r => Row(r.getInt(0).toDouble / 2)))
+
+    checkAnswer(
+      testData2.select($"a" / $"b"),
+      testData2.collect().toSeq.map(r => Row(r.getInt(0).toDouble / 
r.getInt(1))))
+  }
+
+
+  test("%") {
+    checkAnswer(
+      testData2.select($"a" % 2),
+      testData2.collect().toSeq.map(r => Row(r.getInt(0) % 2)))
+
+    checkAnswer(
+      testData2.select($"a" % $"b"),
+      testData2.collect().toSeq.map(r => Row(r.getInt(0) % r.getInt(1))))
+  }
+
+  test("unary -") {
+    checkAnswer(
+      testData2.select(-$"a"),
+      testData2.collect().toSeq.map(r => Row(-r.getInt(0))))
+  }
+
+  test("unary !") {
+    checkAnswer(
+      complexData.select(!$"b"),
+      complexData.collect().toSeq.map(r => Row(!r.getBoolean(3))))
+  }
+
+  test("isNull") {
+    checkAnswer(
+      nullStrings.toDF.where($"s".isNull),
+      nullStrings.collect().toSeq.filter(r => r.getString(1) eq null))
+  }
+
+  test("isNotNull") {
+    checkAnswer(
+      nullStrings.toDF.where($"s".isNotNull),
+      nullStrings.collect().toSeq.filter(r => r.getString(1) ne null))
+  }
+
+  test("===") {
+    checkAnswer(
+      testData2.filter($"a" === 1),
+      testData2.collect().toSeq.filter(r => r.getInt(0) == 1))
+
+    checkAnswer(
+      testData2.filter($"a" === $"b"),
+      testData2.collect().toSeq.filter(r => r.getInt(0) == r.getInt(1)))
+  }
+
+  test("<=>") {
+    checkAnswer(
+      testData2.filter($"a" === 1),
+      testData2.collect().toSeq.filter(r => r.getInt(0) == 1))
+
+    checkAnswer(
+      testData2.filter($"a" === $"b"),
+      testData2.collect().toSeq.filter(r => r.getInt(0) == r.getInt(1)))
+  }
+
+  test("!==") {
+    val nullData = 
TestSQLContext.applySchema(TestSQLContext.sparkContext.parallelize(
+      Row(1, 1) ::
+      Row(1, 2) ::
+      Row(1, null) ::
+      Row(null, null) :: Nil),
+      StructType(Seq(StructField("a", IntegerType), StructField("b", 
IntegerType))))
+
+    checkAnswer(
+      nullData.filter($"b" <=> 1),
+      Row(1, 1) :: Nil)
+
+    checkAnswer(
+      nullData.filter($"b" <=> null),
+      Row(1, null) :: Row(null, null) :: Nil)
+
+    checkAnswer(
+      nullData.filter($"a" <=> $"b"),
+      Row(1, 1) :: Row(null, null) :: Nil)
+  }
+
+  test(">") {
+    checkAnswer(
+      testData2.filter($"a" > 1),
+      testData2.collect().toSeq.filter(r => r.getInt(0) > 1))
+
+    checkAnswer(
+      testData2.filter($"a" > $"b"),
+      testData2.collect().toSeq.filter(r => r.getInt(0) > r.getInt(1)))
+  }
+
+  test(">=") {
+    checkAnswer(
+      testData2.filter($"a" >= 1),
+      testData2.collect().toSeq.filter(r => r.getInt(0) >= 1))
+
+    checkAnswer(
+      testData2.filter($"a" >= $"b"),
+      testData2.collect().toSeq.filter(r => r.getInt(0) >= r.getInt(1)))
+  }
+
+  test("<") {
+    checkAnswer(
+      testData2.filter($"a" < 2),
+      testData2.collect().toSeq.filter(r => r.getInt(0) < 2))
+
+    checkAnswer(
+      testData2.filter($"a" < $"b"),
+      testData2.collect().toSeq.filter(r => r.getInt(0) < r.getInt(1)))
+  }
+
+  test("<=") {
+    checkAnswer(
+      testData2.filter($"a" <= 2),
+      testData2.collect().toSeq.filter(r => r.getInt(0) <= 2))
+
+    checkAnswer(
+      testData2.filter($"a" <= $"b"),
+      testData2.collect().toSeq.filter(r => r.getInt(0) <= r.getInt(1)))
+  }
+
+  val booleanData = 
TestSQLContext.applySchema(TestSQLContext.sparkContext.parallelize(
+    Row(false, false) ::
+      Row(false, true) ::
+      Row(true, false) ::
+      Row(true, true) :: Nil),
+    StructType(Seq(StructField("a", BooleanType), StructField("b", 
BooleanType))))
+
+  test("&&") {
+    checkAnswer(
+      booleanData.filter($"a" && true),
+      Row(true, false) :: Row(true, true) :: Nil)
+
+    checkAnswer(
+      booleanData.filter($"a" && false),
+      Nil)
+
+    checkAnswer(
+      booleanData.filter($"a" && $"b"),
+      Row(true, true) :: Nil)
+  }
+
+  test("||") {
+    checkAnswer(
+      booleanData.filter($"a" || true),
+      booleanData.collect())
+
+    checkAnswer(
+      booleanData.filter($"a" || false),
+      Row(true, false) :: Row(true, true) :: Nil)
+
+    checkAnswer(
+      booleanData.filter($"a" || $"b"),
+      Row(false, true) :: Row(true, false) :: Row(true, true) :: Nil)
+  }
+
+  test("sqrt") {
+    checkAnswer(
+      testData.select(sqrt('key)).orderBy('key.asc),
+      (1 to 100).map(n => Row(math.sqrt(n)))
+    )
+
+    checkAnswer(
+      testData.select(sqrt('value), 'key).orderBy('key.asc, 'value.asc),
+      (1 to 100).map(n => Row(math.sqrt(n), n))
+    )
+
+    checkAnswer(
+      testData.select(sqrt(Literal(null))),
+      (1 to 100).map(_ => Row(null))
+    )
+  }
+
+  test("abs") {
+    checkAnswer(
+      testData.select(abs('key)).orderBy('key.asc),
+      (1 to 100).map(n => Row(n))
+    )
+
+    checkAnswer(
+      negativeData.select(abs('key)).orderBy('key.desc),
+      (1 to 100).map(n => Row(n))
+    )
+
+    checkAnswer(
+      testData.select(abs(Literal(null))),
+      (1 to 100).map(_ => Row(null))
+    )
+  }
+
+  test("upper") {
+    checkAnswer(
+      lowerCaseData.select(upper('l)),
+      ('a' to 'd').map(c => Row(c.toString.toUpperCase))
+    )
+
+    checkAnswer(
+      testData.select(upper('value), 'key),
+      (1 to 100).map(n => Row(n.toString, n))
+    )
+
+    checkAnswer(
+      testData.select(upper(Literal(null))),
+      (1 to 100).map(n => Row(null))
+    )
+  }
+
+  test("lower") {
+    checkAnswer(
+      upperCaseData.select(lower('L)),
+      ('A' to 'F').map(c => Row(c.toString.toLowerCase))
+    )
+
+    checkAnswer(
+      testData.select(lower('value), 'key),
+      (1 to 100).map(n => Row(n.toString, n))
+    )
+
+    checkAnswer(
+      testData.select(lower(Literal(null))),
+      (1 to 100).map(n => Row(null))
+    )
+  }
+}

http://git-wip-us.apache.org/repos/asf/spark/blob/d7437322/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSuite.scala
----------------------------------------------------------------------
diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSuite.scala 
b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSuite.scala
new file mode 100644
index 0000000..6d7d5aa
--- /dev/null
+++ b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSuite.scala
@@ -0,0 +1,280 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *    http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import org.apache.spark.sql.dsl._
+import org.apache.spark.sql.types._
+
+/* Implicits */
+import org.apache.spark.sql.test.TestSQLContext._
+
+import scala.language.postfixOps
+
+class DataFrameSuite extends QueryTest {
+  import org.apache.spark.sql.TestData._
+
+  test("table scan") {
+    checkAnswer(
+      testData,
+      testData.collect().toSeq)
+  }
+
+  test("repartition") {
+    checkAnswer(
+      testData.select('key).repartition(10).select('key),
+      testData.select('key).collect().toSeq)
+  }
+
+  test("agg") {
+    checkAnswer(
+      testData2.groupBy("a").agg($"a", sum($"b")),
+      Seq(Row(1,3), Row(2,3), Row(3,3))
+    )
+    checkAnswer(
+      testData2.groupBy("a").agg($"a", sum($"b").as("totB")).agg(sum('totB)),
+      Row(9)
+    )
+    checkAnswer(
+      testData2.agg(sum('b)),
+      Row(9)
+    )
+  }
+
+  test("convert $\"attribute name\" into unresolved attribute") {
+    checkAnswer(
+      testData.where($"key" === Literal(1)).select($"value"),
+      Row("1"))
+  }
+
+  test("convert Scala Symbol 'attrname into unresolved attribute") {
+    checkAnswer(
+      testData.where('key === Literal(1)).select('value),
+      Row("1"))
+  }
+
+  test("select *") {
+    checkAnswer(
+      testData.select($"*"),
+      testData.collect().toSeq)
+  }
+
+  test("simple select") {
+    checkAnswer(
+      testData.where('key === Literal(1)).select('value),
+      Row("1"))
+  }
+
+  test("select with functions") {
+    checkAnswer(
+      testData.select(sum('value), avg('value), count(Literal(1))),
+      Row(5050.0, 50.5, 100))
+
+    checkAnswer(
+      testData2.select('a + 'b, 'a < 'b),
+      Seq(
+        Row(2, false),
+        Row(3, true),
+        Row(3, false),
+        Row(4, false),
+        Row(4, false),
+        Row(5, false)))
+
+    checkAnswer(
+      testData2.select(sumDistinct('a)),
+      Row(6))
+  }
+
+  test("global sorting") {
+    checkAnswer(
+      testData2.orderBy('a.asc, 'b.asc),
+      Seq(Row(1,1), Row(1,2), Row(2,1), Row(2,2), Row(3,1), Row(3,2)))
+
+    checkAnswer(
+      testData2.orderBy('a.asc, 'b.desc),
+      Seq(Row(1,2), Row(1,1), Row(2,2), Row(2,1), Row(3,2), Row(3,1)))
+
+    checkAnswer(
+      testData2.orderBy('a.desc, 'b.desc),
+      Seq(Row(3,2), Row(3,1), Row(2,2), Row(2,1), Row(1,2), Row(1,1)))
+
+    checkAnswer(
+      testData2.orderBy('a.desc, 'b.asc),
+      Seq(Row(3,1), Row(3,2), Row(2,1), Row(2,2), Row(1,1), Row(1,2)))
+
+    checkAnswer(
+      arrayData.orderBy('data.getItem(0).asc),
+      arrayData.toDF.collect().sortBy(_.getAs[Seq[Int]](0)(0)).toSeq)
+
+    checkAnswer(
+      arrayData.orderBy('data.getItem(0).desc),
+      arrayData.toDF.collect().sortBy(_.getAs[Seq[Int]](0)(0)).reverse.toSeq)
+
+    checkAnswer(
+      arrayData.orderBy('data.getItem(1).asc),
+      arrayData.toDF.collect().sortBy(_.getAs[Seq[Int]](0)(1)).toSeq)
+
+    checkAnswer(
+      arrayData.orderBy('data.getItem(1).desc),
+      arrayData.toDF.collect().sortBy(_.getAs[Seq[Int]](0)(1)).reverse.toSeq)
+  }
+
+  test("limit") {
+    checkAnswer(
+      testData.limit(10),
+      testData.take(10).toSeq)
+
+    checkAnswer(
+      arrayData.limit(1),
+      arrayData.take(1).map(r => Row.fromSeq(r.productIterator.toSeq)))
+
+    checkAnswer(
+      mapData.limit(1),
+      mapData.take(1).map(r => Row.fromSeq(r.productIterator.toSeq)))
+  }
+
+  test("average") {
+    checkAnswer(
+      testData2.agg(avg('a)),
+      Row(2.0))
+
+    checkAnswer(
+      testData2.agg(avg('a), sumDistinct('a)), // non-partial
+      Row(2.0, 6.0) :: Nil)
+
+    checkAnswer(
+      decimalData.agg(avg('a)),
+      Row(new java.math.BigDecimal(2.0)))
+    checkAnswer(
+      decimalData.agg(avg('a), sumDistinct('a)), // non-partial
+      Row(new java.math.BigDecimal(2.0), new java.math.BigDecimal(6)) :: Nil)
+
+    checkAnswer(
+      decimalData.agg(avg('a cast DecimalType(10, 2))),
+      Row(new java.math.BigDecimal(2.0)))
+    checkAnswer(
+      decimalData.agg(avg('a cast DecimalType(10, 2)), sumDistinct('a cast 
DecimalType(10, 2))), // non-partial
+      Row(new java.math.BigDecimal(2.0), new java.math.BigDecimal(6)) :: Nil)
+  }
+
+  test("null average") {
+    checkAnswer(
+      testData3.agg(avg('b)),
+      Row(2.0))
+
+    checkAnswer(
+      testData3.agg(avg('b), countDistinct('b)),
+      Row(2.0, 1))
+
+    checkAnswer(
+      testData3.agg(avg('b), sumDistinct('b)), // non-partial
+      Row(2.0, 2.0))
+  }
+
+  test("zero average") {
+    checkAnswer(
+      emptyTableData.agg(avg('a)),
+      Row(null))
+
+    checkAnswer(
+      emptyTableData.agg(avg('a), sumDistinct('b)), // non-partial
+      Row(null, null))
+  }
+
+  test("count") {
+    assert(testData2.count() === testData2.map(_ => 1).count())
+
+    checkAnswer(
+      testData2.agg(count('a), sumDistinct('a)), // non-partial
+      Row(6, 6.0))
+  }
+
+  test("null count") {
+    checkAnswer(
+      testData3.groupBy('a).agg('a, count('b)),
+      Seq(Row(1,0), Row(2, 1))
+    )
+
+    checkAnswer(
+      testData3.groupBy('a).agg('a, count('a + 'b)),
+      Seq(Row(1,0), Row(2, 1))
+    )
+
+    checkAnswer(
+      testData3.agg(count('a), count('b), count(Literal(1)), 
countDistinct('a), countDistinct('b)),
+      Row(2, 1, 2, 2, 1)
+    )
+
+    checkAnswer(
+      testData3.agg(count('b), countDistinct('b), sumDistinct('b)), // 
non-partial
+      Row(1, 1, 2)
+    )
+  }
+
+  test("zero count") {
+    assert(emptyTableData.count() === 0)
+
+    checkAnswer(
+      emptyTableData.agg(count('a), sumDistinct('a)), // non-partial
+      Row(0, null))
+  }
+
+  test("zero sum") {
+    checkAnswer(
+      emptyTableData.agg(sum('a)),
+      Row(null))
+  }
+
+  test("zero sum distinct") {
+    checkAnswer(
+      emptyTableData.agg(sumDistinct('a)),
+      Row(null))
+  }
+
+  test("except") {
+    checkAnswer(
+      lowerCaseData.except(upperCaseData),
+      Row(1, "a") ::
+      Row(2, "b") ::
+      Row(3, "c") ::
+      Row(4, "d") :: Nil)
+    checkAnswer(lowerCaseData.except(lowerCaseData), Nil)
+    checkAnswer(upperCaseData.except(upperCaseData), Nil)
+  }
+
+  test("intersect") {
+    checkAnswer(
+      lowerCaseData.intersect(lowerCaseData),
+      Row(1, "a") ::
+      Row(2, "b") ::
+      Row(3, "c") ::
+      Row(4, "d") :: Nil)
+    checkAnswer(lowerCaseData.intersect(upperCaseData), Nil)
+  }
+
+  test("udf") {
+    val foo = (a: Int, b: String) => a.toString + b
+
+    checkAnswer(
+      // SELECT *, foo(key, value) FROM testData
+      testData.select($"*", callUDF(foo, 'key, 'value)).limit(3),
+      Row(1, "1", "11") :: Row(2, "2", "22") :: Row(3, "3", "33") :: Nil
+    )
+  }
+
+
+}

http://git-wip-us.apache.org/repos/asf/spark/blob/d7437322/sql/core/src/test/scala/org/apache/spark/sql/DslQuerySuite.scala
----------------------------------------------------------------------
diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DslQuerySuite.scala 
b/sql/core/src/test/scala/org/apache/spark/sql/DslQuerySuite.scala
deleted file mode 100644
index a5848f2..0000000
--- a/sql/core/src/test/scala/org/apache/spark/sql/DslQuerySuite.scala
+++ /dev/null
@@ -1,346 +0,0 @@
-/*
- * Licensed to the Apache Software Foundation (ASF) under one or more
- * contributor license agreements.  See the NOTICE file distributed with
- * this work for additional information regarding copyright ownership.
- * The ASF licenses this file to You under the Apache License, Version 2.0
- * (the "License"); you may not use this file except in compliance with
- * the License.  You may obtain a copy of the License at
- *
- *    http://www.apache.org/licenses/LICENSE-2.0
- *
- * Unless required by applicable law or agreed to in writing, software
- * distributed under the License is distributed on an "AS IS" BASIS,
- * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
- * See the License for the specific language governing permissions and
- * limitations under the License.
- */
-
-package org.apache.spark.sql
-
-import org.apache.spark.sql.dsl._
-import org.apache.spark.sql.types._
-
-/* Implicits */
-import org.apache.spark.sql.test.TestSQLContext._
-
-import scala.language.postfixOps
-
-class DslQuerySuite extends QueryTest {
-  import org.apache.spark.sql.TestData._
-
-  test("table scan") {
-    checkAnswer(
-      testData,
-      testData.collect().toSeq)
-  }
-
-  test("repartition") {
-    checkAnswer(
-      testData.select('key).repartition(10).select('key),
-      testData.select('key).collect().toSeq)
-  }
-
-  test("agg") {
-    checkAnswer(
-      testData2.groupBy("a").agg($"a", sum($"b")),
-      Seq(Row(1,3), Row(2,3), Row(3,3))
-    )
-    checkAnswer(
-      testData2.groupBy("a").agg($"a", sum($"b").as("totB")).agg(sum('totB)),
-      Row(9)
-    )
-    checkAnswer(
-      testData2.agg(sum('b)),
-      Row(9)
-    )
-  }
-
-  test("convert $\"attribute name\" into unresolved attribute") {
-    checkAnswer(
-      testData.where($"key" === Literal(1)).select($"value"),
-      Row("1"))
-  }
-
-  test("convert Scala Symbol 'attrname into unresolved attribute") {
-    checkAnswer(
-      testData.where('key === Literal(1)).select('value),
-      Row("1"))
-  }
-
-  test("select *") {
-    checkAnswer(
-      testData.select($"*"),
-      testData.collect().toSeq)
-  }
-
-  test("simple select") {
-    checkAnswer(
-      testData.where('key === Literal(1)).select('value),
-      Row("1"))
-  }
-
-  test("select with functions") {
-    checkAnswer(
-      testData.select(sum('value), avg('value), count(Literal(1))),
-      Row(5050.0, 50.5, 100))
-
-    checkAnswer(
-      testData2.select('a + 'b, 'a < 'b),
-      Seq(
-        Row(2, false),
-        Row(3, true),
-        Row(3, false),
-        Row(4, false),
-        Row(4, false),
-        Row(5, false)))
-
-    checkAnswer(
-      testData2.select(sumDistinct('a)),
-      Row(6))
-  }
-
-  test("global sorting") {
-    checkAnswer(
-      testData2.orderBy('a.asc, 'b.asc),
-      Seq(Row(1,1), Row(1,2), Row(2,1), Row(2,2), Row(3,1), Row(3,2)))
-
-    checkAnswer(
-      testData2.orderBy('a.asc, 'b.desc),
-      Seq(Row(1,2), Row(1,1), Row(2,2), Row(2,1), Row(3,2), Row(3,1)))
-
-    checkAnswer(
-      testData2.orderBy('a.desc, 'b.desc),
-      Seq(Row(3,2), Row(3,1), Row(2,2), Row(2,1), Row(1,2), Row(1,1)))
-
-    checkAnswer(
-      testData2.orderBy('a.desc, 'b.asc),
-      Seq(Row(3,1), Row(3,2), Row(2,1), Row(2,2), Row(1,1), Row(1,2)))
-
-    checkAnswer(
-      arrayData.orderBy('data.getItem(0).asc),
-      arrayData.toDF.collect().sortBy(_.getAs[Seq[Int]](0)(0)).toSeq)
-
-    checkAnswer(
-      arrayData.orderBy('data.getItem(0).desc),
-      arrayData.toDF.collect().sortBy(_.getAs[Seq[Int]](0)(0)).reverse.toSeq)
-
-    checkAnswer(
-      arrayData.orderBy('data.getItem(1).asc),
-      arrayData.toDF.collect().sortBy(_.getAs[Seq[Int]](0)(1)).toSeq)
-
-    checkAnswer(
-      arrayData.orderBy('data.getItem(1).desc),
-      arrayData.toDF.collect().sortBy(_.getAs[Seq[Int]](0)(1)).reverse.toSeq)
-  }
-
-  test("limit") {
-    checkAnswer(
-      testData.limit(10),
-      testData.take(10).toSeq)
-
-    checkAnswer(
-      arrayData.limit(1),
-      arrayData.take(1).map(r => Row.fromSeq(r.productIterator.toSeq)))
-
-    checkAnswer(
-      mapData.limit(1),
-      mapData.take(1).map(r => Row.fromSeq(r.productIterator.toSeq)))
-  }
-
-  test("average") {
-    checkAnswer(
-      testData2.agg(avg('a)),
-      Row(2.0))
-
-    checkAnswer(
-      testData2.agg(avg('a), sumDistinct('a)), // non-partial
-      Row(2.0, 6.0) :: Nil)
-
-    checkAnswer(
-      decimalData.agg(avg('a)),
-      Row(new java.math.BigDecimal(2.0)))
-    checkAnswer(
-      decimalData.agg(avg('a), sumDistinct('a)), // non-partial
-      Row(new java.math.BigDecimal(2.0), new java.math.BigDecimal(6)) :: Nil)
-
-    checkAnswer(
-      decimalData.agg(avg('a cast DecimalType(10, 2))),
-      Row(new java.math.BigDecimal(2.0)))
-    checkAnswer(
-      decimalData.agg(avg('a cast DecimalType(10, 2)), sumDistinct('a cast 
DecimalType(10, 2))), // non-partial
-      Row(new java.math.BigDecimal(2.0), new java.math.BigDecimal(6)) :: Nil)
-  }
-
-  test("null average") {
-    checkAnswer(
-      testData3.agg(avg('b)),
-      Row(2.0))
-
-    checkAnswer(
-      testData3.agg(avg('b), countDistinct('b)),
-      Row(2.0, 1))
-
-    checkAnswer(
-      testData3.agg(avg('b), sumDistinct('b)), // non-partial
-      Row(2.0, 2.0))
-  }
-
-  test("zero average") {
-    checkAnswer(
-      emptyTableData.agg(avg('a)),
-      Row(null))
-
-    checkAnswer(
-      emptyTableData.agg(avg('a), sumDistinct('b)), // non-partial
-      Row(null, null))
-  }
-
-  test("count") {
-    assert(testData2.count() === testData2.map(_ => 1).count())
-
-    checkAnswer(
-      testData2.agg(count('a), sumDistinct('a)), // non-partial
-      Row(6, 6.0))
-  }
-
-  test("null count") {
-    checkAnswer(
-      testData3.groupBy('a).agg('a, count('b)),
-      Seq(Row(1,0), Row(2, 1))
-    )
-
-    checkAnswer(
-      testData3.groupBy('a).agg('a, count('a + 'b)),
-      Seq(Row(1,0), Row(2, 1))
-    )
-
-    checkAnswer(
-      testData3.agg(count('a), count('b), count(Literal(1)), 
countDistinct('a), countDistinct('b)),
-      Row(2, 1, 2, 2, 1)
-    )
-
-    checkAnswer(
-      testData3.agg(count('b), countDistinct('b), sumDistinct('b)), // 
non-partial
-      Row(1, 1, 2)
-    )
-  }
-
-  test("zero count") {
-    assert(emptyTableData.count() === 0)
-
-    checkAnswer(
-      emptyTableData.agg(count('a), sumDistinct('a)), // non-partial
-      Row(0, null))
-  }
-
-  test("zero sum") {
-    checkAnswer(
-      emptyTableData.agg(sum('a)),
-      Row(null))
-  }
-
-  test("zero sum distinct") {
-    checkAnswer(
-      emptyTableData.agg(sumDistinct('a)),
-      Row(null))
-  }
-
-  test("except") {
-    checkAnswer(
-      lowerCaseData.except(upperCaseData),
-      Row(1, "a") ::
-      Row(2, "b") ::
-      Row(3, "c") ::
-      Row(4, "d") :: Nil)
-    checkAnswer(lowerCaseData.except(lowerCaseData), Nil)
-    checkAnswer(upperCaseData.except(upperCaseData), Nil)
-  }
-
-  test("intersect") {
-    checkAnswer(
-      lowerCaseData.intersect(lowerCaseData),
-      Row(1, "a") ::
-      Row(2, "b") ::
-      Row(3, "c") ::
-      Row(4, "d") :: Nil)
-    checkAnswer(lowerCaseData.intersect(upperCaseData), Nil)
-  }
-
-  test("udf") {
-    val foo = (a: Int, b: String) => a.toString + b
-
-    checkAnswer(
-      // SELECT *, foo(key, value) FROM testData
-      testData.select($"*", callUDF(foo, 'key, 'value)).limit(3),
-      Row(1, "1", "11") :: Row(2, "2", "22") :: Row(3, "3", "33") :: Nil
-    )
-  }
-
-  test("sqrt") {
-    checkAnswer(
-      testData.select(sqrt('key)).orderBy('key asc),
-      (1 to 100).map(n => Row(math.sqrt(n)))
-    )
-
-    checkAnswer(
-      testData.select(sqrt('value), 'key).orderBy('key asc, 'value asc),
-      (1 to 100).map(n => Row(math.sqrt(n), n))
-    )
-
-    checkAnswer(
-      testData.select(sqrt(Literal(null))),
-      (1 to 100).map(_ => Row(null))
-    )
-  }
-
-  test("abs") {
-    checkAnswer(
-      testData.select(abs('key)).orderBy('key asc),
-      (1 to 100).map(n => Row(n))
-    )
-
-    checkAnswer(
-      negativeData.select(abs('key)).orderBy('key desc),
-      (1 to 100).map(n => Row(n))
-    )
-
-    checkAnswer(
-      testData.select(abs(Literal(null))),
-      (1 to 100).map(_ => Row(null))
-    )
-  }
-
-  test("upper") {
-    checkAnswer(
-      lowerCaseData.select(upper('l)),
-      ('a' to 'd').map(c => Row(c.toString.toUpperCase))
-    )
-
-    checkAnswer(
-      testData.select(upper('value), 'key),
-      (1 to 100).map(n => Row(n.toString, n))
-    )
-
-    checkAnswer(
-      testData.select(upper(Literal(null))),
-      (1 to 100).map(n => Row(null))
-    )
-  }
-
-  test("lower") {
-    checkAnswer(
-      upperCaseData.select(lower('L)),
-      ('A' to 'F').map(c => Row(c.toString.toLowerCase))
-    )
-
-    checkAnswer(
-      testData.select(lower('value), 'key),
-      (1 to 100).map(n => Row(n.toString, n))
-    )
-
-    checkAnswer(
-      testData.select(lower(Literal(null))),
-      (1 to 100).map(n => Row(null))
-    )
-  }
-}

http://git-wip-us.apache.org/repos/asf/spark/blob/d7437322/sql/core/src/test/scala/org/apache/spark/sql/TestData.scala
----------------------------------------------------------------------
diff --git a/sql/core/src/test/scala/org/apache/spark/sql/TestData.scala 
b/sql/core/src/test/scala/org/apache/spark/sql/TestData.scala
index fffa2b7..9eefe67 100644
--- a/sql/core/src/test/scala/org/apache/spark/sql/TestData.scala
+++ b/sql/core/src/test/scala/org/apache/spark/sql/TestData.scala
@@ -161,7 +161,7 @@ object TestData {
     TestSQLContext.sparkContext.parallelize(
       NullStrings(1, "abc") ::
       NullStrings(2, "ABC") ::
-      NullStrings(3, null) :: Nil)
+      NullStrings(3, null) :: Nil).toDF
   nullStrings.registerTempTable("nullStrings")
 
   case class TableName(tableName: String)


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to