diff --git a/build.sbt b/build.sbt
index 40ae2d2..12ab0ef 100644
--- a/build.sbt
+++ b/build.sbt
@@ -17,7 +17,7 @@
name := "spark-bigquery"
organization := "com.spotify"
-scalaVersion := "2.10.6"
+scalaVersion := "2.11.8"
crossScalaVersions := Seq("2.10.6", "2.11.8")
spName := "spotify/spark-bigquery"
@@ -27,7 +27,7 @@ spAppendScalaVersion := true
spIncludeMaven := true
libraryDependencies ++= Seq(
- "com.databricks" %% "spark-avro" % "3.0.0",
+// "com.databricks" %% "spark-avro" % "3.2.0",
"com.google.cloud.bigdataoss" % "bigquery-connector" % "0.7.5-hadoop2"
exclude ("com.google.guava", "guava-jdk5"),
"org.slf4j" % "slf4j-simple" % "1.7.21",
diff --git a/project/plugins.sbt b/project/plugins.sbt
index 46d1cb0..c6a83eb 100644
--- a/project/plugins.sbt
+++ b/project/plugins.sbt
@@ -5,3 +5,4 @@ addSbtPlugin("org.scalastyle" % "scalastyle-sbt-plugin" % "0.8.0")
addSbtPlugin("org.scoverage" % "sbt-scoverage" % "1.3.5")
addSbtPlugin("org.spark-packages" % "sbt-spark-package" % "0.2.4")
addSbtPlugin("org.xerial.sbt" % "sbt-sonatype" % "1.1")
+addSbtPlugin("net.virtual-void" % "sbt-dependency-graph" % "0.8.2")
diff --git a/scalastyle-config.xml b/scalastyle-config.xml
index 2b10c34..1806655 100644
--- a/scalastyle-config.xml
+++ b/scalastyle-config.xml
@@ -32,7 +32,7 @@
-
+
diff --git a/src/main/scala/com/databricks/spark/avro/DefaultSource.scala b/src/main/scala/com/databricks/spark/avro/DefaultSource.scala
new file mode 100644
index 0000000..0d481e0
--- /dev/null
+++ b/src/main/scala/com/databricks/spark/avro/DefaultSource.scala
@@ -0,0 +1,289 @@
+/*
+ * Copyright 2014 Databricks
+ *
+ * Licensed 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 com.databricks.spark.avro
+
+import java.io._
+import java.net.URI
+import java.util.zip.Deflater
+
+import com.databricks.spark.avro.DefaultSource.{IgnoreFilesWithoutExtensionProperty, SerializableConfiguration}
+import com.databricks.spark.avro.clone.{AvroOutputWriterFactory, SchemaConverters}
+import com.esotericsoftware.kryo.io.{Input, Output}
+import com.esotericsoftware.kryo.{Kryo, KryoSerializable}
+import org.apache.avro.SchemaBuilder
+import org.apache.avro.file.{DataFileConstants, DataFileReader}
+import org.apache.avro.generic.{GenericDatumReader, GenericRecord}
+import org.apache.avro.mapred.{AvroOutputFormat, FsInput}
+import org.apache.avro.mapreduce.AvroJob
+import org.apache.hadoop.conf.Configuration
+import org.apache.hadoop.fs.{FileStatus, Path}
+import org.apache.hadoop.mapreduce.Job
+import org.apache.spark.TaskContext
+import org.apache.spark.sql.SparkSession
+import org.apache.spark.sql.catalyst.InternalRow
+import org.apache.spark.sql.catalyst.encoders.RowEncoder
+import org.apache.spark.sql.catalyst.expressions.GenericRow
+import org.apache.spark.sql.execution.datasources.{FileFormat, OutputWriterFactory, PartitionedFile}
+import org.apache.spark.sql.sources.{DataSourceRegister, Filter}
+import org.apache.spark.sql.types.StructType
+import org.slf4j.LoggerFactory
+
+import scala.util.control.NonFatal
+
+class DefaultSource extends FileFormat with DataSourceRegister {
+ private val log = LoggerFactory.getLogger(getClass)
+
+ override def equals(other: Any): Boolean = other match {
+ case _: DefaultSource => true
+ case _ => false
+ }
+
+ override def inferSchema(
+ spark: SparkSession,
+ options: Map[String, String],
+ files: Seq[FileStatus]): Option[StructType] = {
+ val conf = spark.sparkContext.hadoopConfiguration
+
+ // Schema evolution is not supported yet. Here we only pick a single random sample file to
+ // figure out the schema of the whole dataset.
+ val sampleFile = if (conf.getBoolean(IgnoreFilesWithoutExtensionProperty, true)) {
+ files.find(_.getPath.getName.endsWith(".avro")).getOrElse {
+ throw new FileNotFoundException(
+ "No Avro files found. Hadoop option \"avro.mapred.ignore.inputs.without.extension\" is " +
+ "set to true. Do all input files have \".avro\" extension?"
+ )
+ }
+ } else {
+ files.headOption.getOrElse {
+ throw new FileNotFoundException("No Avro files found.")
+ }
+ }
+
+ val avroSchema = {
+ val in = new FsInput(sampleFile.getPath, conf)
+ try {
+ val reader = DataFileReader.openReader(in, new GenericDatumReader[GenericRecord]())
+ try {
+ reader.getSchema
+ } finally {
+ reader.close()
+ }
+ } finally {
+ in.close()
+ }
+ }
+
+ SchemaConverters.toSqlType(avroSchema).dataType match {
+ case t: StructType => Some(t)
+ case _ => throw new RuntimeException(
+ s"""Avro schema cannot be converted to a Spark SQL StructType:
+ |
+ |${avroSchema.toString(true)}
+ |""".stripMargin)
+ }
+ }
+
+ override def shortName(): String = "avro"
+
+ override def prepareWrite(
+ spark: SparkSession,
+ job: Job,
+ options: Map[String, String],
+ dataSchema: StructType): OutputWriterFactory = {
+ val recordName = options.getOrElse("recordName", "topLevelRecord")
+ val recordNamespace = options.getOrElse("recordNamespace", "")
+ val build = SchemaBuilder.record(recordName).namespace(recordNamespace)
+ val outputAvroSchema = SchemaConverters.convertStructToAvro(dataSchema, build, recordNamespace)
+
+ AvroJob.setOutputKeySchema(job, outputAvroSchema)
+ val AVRO_COMPRESSION_CODEC = "spark.sql.avro.compression.codec"
+ val AVRO_DEFLATE_LEVEL = "spark.sql.avro.deflate.level"
+ val COMPRESS_KEY = "mapred.output.compress"
+
+ spark.conf.get(AVRO_COMPRESSION_CODEC, "snappy") match {
+ case "uncompressed" =>
+ log.info("writing uncompressed Avro records")
+ job.getConfiguration.setBoolean(COMPRESS_KEY, false)
+
+ case "snappy" =>
+ log.info("compressing Avro output using Snappy")
+ job.getConfiguration.setBoolean(COMPRESS_KEY, true)
+ job.getConfiguration.set(AvroJob.CONF_OUTPUT_CODEC, DataFileConstants.SNAPPY_CODEC)
+
+ case "deflate" =>
+ val deflateLevel = spark.conf.get(
+ AVRO_DEFLATE_LEVEL, Deflater.DEFAULT_COMPRESSION.toString).toInt
+ log.info(s"compressing Avro output using deflate (level=$deflateLevel)")
+ job.getConfiguration.setBoolean(COMPRESS_KEY, true)
+ job.getConfiguration.set(AvroJob.CONF_OUTPUT_CODEC, DataFileConstants.DEFLATE_CODEC)
+ job.getConfiguration.setInt(AvroOutputFormat.DEFLATE_LEVEL_KEY, deflateLevel)
+
+ case unknown: String =>
+ log.error(s"unsupported compression codec $unknown")
+ }
+
+ new AvroOutputWriterFactory(dataSchema, recordName, recordNamespace)
+ }
+
+ override def buildReader(
+ spark: SparkSession,
+ dataSchema: StructType,
+ partitionSchema: StructType,
+ requiredSchema: StructType,
+ filters: Seq[Filter],
+ options: Map[String, String],
+ hadoopConf: Configuration): (PartitionedFile) => Iterator[InternalRow] = {
+
+ val broadcastedConf =
+ spark.sparkContext.broadcast(new SerializableConfiguration(hadoopConf))
+
+ (file: PartitionedFile) => {
+ val log = LoggerFactory.getLogger(classOf[DefaultSource])
+ val conf = broadcastedConf.value.value
+
+ // TODO Removes this check once `FileFormat` gets a general file filtering interface method.
+ // Doing input file filtering is improper because we may generate empty tasks that process no
+ // input files but stress the scheduler. We should probably add a more general input file
+ // filtering mechanism for `FileFormat` data sources. See SPARK-16317.
+ if (
+ conf.getBoolean(IgnoreFilesWithoutExtensionProperty, true) &&
+ !file.filePath.endsWith(".avro")
+ ) {
+ Iterator.empty
+ } else {
+ val reader = {
+ val in = new FsInput(new Path(new URI(file.filePath)), conf)
+ try {
+ DataFileReader.openReader(in, new GenericDatumReader[GenericRecord]())
+ } catch {
+ case NonFatal(e) =>
+ log.error("Exception while opening DataFileReader", e)
+ in.close()
+ throw e
+ }
+ }
+
+ // Ensure that the reader is closed even if the task fails or doesn't consume the entire
+ // iterator of records.
+ Option(TaskContext.get()).foreach { taskContext =>
+ taskContext.addTaskCompletionListener { _ =>
+ reader.close()
+ }
+ }
+
+ val fieldExtractors = {
+ val avroSchema = reader.getSchema
+ requiredSchema.zipWithIndex.map { case (field, index) =>
+ val avroField = Option(avroSchema.getField(field.name)).getOrElse {
+ throw new IllegalArgumentException(
+ s"""Cannot find required column ${field.name} in Avro schema:"
+ |
+ |${avroSchema.toString(true)}
+ """.stripMargin
+ )
+ }
+
+ val converter = SchemaConverters.createConverterToSQL(avroField.schema())
+
+ (record: GenericRecord, buffer: Array[Any]) => {
+ buffer(index) = converter(record.get(avroField.pos()))
+ }
+ }
+ }
+
+ new Iterator[InternalRow] {
+ private val rowBuffer = Array.fill[Any](requiredSchema.length)(null)
+
+ private val safeDataRow = new GenericRow(rowBuffer)
+
+ // Used to convert `Row`s containing data columns into `InternalRow`s.
+ private val encoderForDataColumns = RowEncoder(requiredSchema)
+
+ private[this] var completed = false
+
+ override def hasNext: Boolean = {
+ if (completed) {
+ false
+ } else {
+ val r = reader.hasNext
+ if (!r) {
+ reader.close()
+ completed = true
+ }
+ r
+ }
+ }
+
+ override def next(): InternalRow = {
+ val record = reader.next()
+
+ var i = 0
+ while (i < requiredSchema.length) {
+ fieldExtractors(i)(record, rowBuffer)
+ i += 1
+ }
+
+ encoderForDataColumns.toRow(safeDataRow)
+ }
+ }
+ }
+ }
+ }
+}
+
+object DefaultSource {
+ val IgnoreFilesWithoutExtensionProperty = "avro.mapred.ignore.inputs.without.extension"
+
+ class SerializableConfiguration(@transient var value: Configuration)
+ extends Serializable with KryoSerializable {
+ @transient private[avro] lazy val log = LoggerFactory.getLogger(getClass)
+
+ private def writeObject(out: ObjectOutputStream): Unit = tryOrIOException {
+ out.defaultWriteObject()
+ value.write(out)
+ }
+
+ private def readObject(in: ObjectInputStream): Unit = tryOrIOException {
+ value = new Configuration(false)
+ value.readFields(in)
+ }
+
+ private def tryOrIOException[T](block: => T): T = {
+ try {
+ block
+ } catch {
+ case e: IOException =>
+ log.error("Exception encountered", e)
+ throw e
+ case NonFatal(e) =>
+ log.error("Exception encountered", e)
+ throw new IOException(e)
+ }
+ }
+
+ def write(kryo: Kryo, out: Output): Unit = {
+ val dos = new DataOutputStream(out)
+ value.write(dos)
+ dos.flush()
+ }
+
+ def read(kryo: Kryo, in: Input): Unit = {
+ value = new Configuration(false)
+ value.readFields(new DataInputStream(in))
+ }
+ }
+}
diff --git a/src/main/scala/com/databricks/spark/avro/clone/AvroOutputWriter.scala b/src/main/scala/com/databricks/spark/avro/clone/AvroOutputWriter.scala
new file mode 100644
index 0000000..698c929
--- /dev/null
+++ b/src/main/scala/com/databricks/spark/avro/clone/AvroOutputWriter.scala
@@ -0,0 +1,167 @@
+/*
+ * Copyright 2014 Databricks
+ *
+ * Licensed 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 com.databricks.spark.avro.clone
+
+import java.io.{IOException, OutputStream}
+import java.nio.ByteBuffer
+import java.sql.{Date, Timestamp}
+import java.util.HashMap
+
+import org.apache.hadoop.fs.Path
+
+import scala.collection.immutable.Map
+import org.apache.avro.generic.GenericData.Record
+import org.apache.avro.generic.GenericRecord
+import org.apache.avro.{Schema, SchemaBuilder}
+import org.apache.avro.mapred.AvroKey
+import org.apache.avro.mapreduce.AvroKeyOutputFormat
+import org.apache.hadoop.conf.Configuration
+import org.apache.hadoop.io.NullWritable
+import org.apache.hadoop.mapreduce.{RecordWriter, TaskAttemptContext, TaskAttemptID}
+import org.apache.spark.sql.Row
+import org.apache.spark.sql.execution.datasources.OutputWriter
+import org.apache.spark.sql.types._
+
+// NOTE: This class is instantiated and used on executor side only, no need to be serializable.
+class AvroOutputWriter(
+ path: String,
+ context: TaskAttemptContext,
+ schema: StructType,
+ recordName: String,
+ recordNamespace: String) extends OutputWriter {
+
+ private lazy val converter = createConverterToAvro(schema, recordName, recordNamespace)
+
+ /**
+ * Overrides the couple of methods responsible for generating the output streams / files so
+ * that the data can be correctly partitioned
+ */
+ private val recordWriter: RecordWriter[AvroKey[GenericRecord], NullWritable] =
+ new AvroKeyOutputFormat[GenericRecord]() {
+
+ private def getConfigurationFromContext(context: TaskAttemptContext): Configuration = {
+ // Use reflection to get the Configuration. This is necessary because TaskAttemptContext
+ // is a class in Hadoop 1.x and an interface in Hadoop 2.x.
+ val method = context.getClass.getMethod("getConfiguration")
+ method.invoke(context).asInstanceOf[Configuration]
+ }
+
+ override def getDefaultWorkFile(context: TaskAttemptContext, extension: String): Path = {
+ val uniqueWriteJobId =
+ getConfigurationFromContext(context).get("spark.sql.sources.writeJobUUID")
+ val taskAttemptId: TaskAttemptID = {
+ // Use reflection to get the TaskAttemptID. This is necessary because TaskAttemptContext
+ // is a class in Hadoop 1.x and an interface in Hadoop 2.x.
+ val method = context.getClass.getMethod("getTaskAttemptID")
+ method.invoke(context).asInstanceOf[TaskAttemptID]
+ }
+ val split = taskAttemptId.getTaskID.getId
+ new Path(path, f"part-r-$split%05d-$uniqueWriteJobId$extension")
+ }
+
+ @throws(classOf[IOException])
+ override def getAvroFileOutputStream(c: TaskAttemptContext): OutputStream = {
+ val path = getDefaultWorkFile(context, ".avro")
+ path.getFileSystem(getConfigurationFromContext(context)).create(path)
+ }
+
+ }.getRecordWriter(context)
+
+ override def write(row: Row): Unit = {
+ val key = new AvroKey(converter(row).asInstanceOf[GenericRecord])
+ recordWriter.write(key, NullWritable.get())
+ }
+
+ override def close(): Unit = recordWriter.close(context)
+
+ // scalastyle:off
+
+ /**
+ * This function constructs converter function for a given sparkSQL datatype. This is used in
+ * writing Avro records out to disk
+ */
+ private def createConverterToAvro(
+ dataType: DataType,
+ structName: String,
+ recordNamespace: String): (Any) => Any = {
+ dataType match {
+ case BinaryType => (item: Any) => item match {
+ case null => null
+ case bytes: Array[Byte] => ByteBuffer.wrap(bytes)
+ }
+ case ByteType | ShortType | IntegerType | LongType |
+ FloatType | DoubleType | StringType | BooleanType => identity
+ case _: DecimalType => (item: Any) => if (item == null) null else item.toString
+ case TimestampType => (item: Any) =>
+ if (item == null) null else item.asInstanceOf[Timestamp].getTime * 1000l // MNC: mills -> micros hack.
+ case DateType => (item: Any) =>
+ if (item == null) null else item.asInstanceOf[Date].toString
+ case ArrayType(elementType, _) =>
+ val elementConverter = createConverterToAvro(elementType, structName, recordNamespace)
+ (item: Any) => {
+ if (item == null) {
+ null
+ } else {
+ val sourceArray = item.asInstanceOf[Seq[Any]]
+ val sourceArraySize = sourceArray.size
+ val targetArray = new Array[Any](sourceArraySize)
+ var idx = 0
+ while (idx < sourceArraySize) {
+ targetArray(idx) = elementConverter(sourceArray(idx))
+ idx += 1
+ }
+ targetArray
+ }
+ }
+ case MapType(StringType, valueType, _) =>
+ val valueConverter = createConverterToAvro(valueType, structName, recordNamespace)
+ (item: Any) => {
+ if (item == null) {
+ null
+ } else {
+ val javaMap = new HashMap[String, Any]()
+ item.asInstanceOf[Map[String, Any]].foreach { case (key, value) =>
+ javaMap.put(key, valueConverter(value))
+ }
+ javaMap
+ }
+ }
+ case structType: StructType =>
+ val builder = SchemaBuilder.record(structName).namespace(recordNamespace)
+ val schema: Schema = SchemaConverters.convertStructToAvro(
+ structType, builder, recordNamespace)
+ val fieldConverters = structType.fields.map(field =>
+ createConverterToAvro(field.dataType, field.name, recordNamespace))
+ (item: Any) => {
+ if (item == null) {
+ null
+ } else {
+ val record = new Record(schema)
+ val convertersIterator = fieldConverters.iterator
+ val fieldNamesIterator = dataType.asInstanceOf[StructType].fieldNames.iterator
+ val rowIterator = item.asInstanceOf[Row].toSeq.iterator
+
+ while (convertersIterator.hasNext) {
+ val converter = convertersIterator.next()
+ record.put(fieldNamesIterator.next(), converter(rowIterator.next()))
+ }
+ record
+ }
+ }
+ }
+ }
+}
diff --git a/src/main/scala/com/databricks/spark/avro/clone/AvroOutputWriterFactory.scala b/src/main/scala/com/databricks/spark/avro/clone/AvroOutputWriterFactory.scala
new file mode 100644
index 0000000..07313bd
--- /dev/null
+++ b/src/main/scala/com/databricks/spark/avro/clone/AvroOutputWriterFactory.scala
@@ -0,0 +1,36 @@
+/*
+ * Copyright 2014 Databricks
+ *
+ * Licensed 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 com.databricks.spark.avro.clone
+
+import org.apache.hadoop.mapreduce.TaskAttemptContext
+
+import org.apache.spark.sql.execution.datasources.{OutputWriter, OutputWriterFactory}
+import org.apache.spark.sql.types.StructType
+
+class AvroOutputWriterFactory(
+ schema: StructType,
+ recordName: String,
+ recordNamespace: String) extends OutputWriterFactory {
+
+ def newInstance(
+ path: String,
+ bucketId: Option[Int],
+ dataSchema: StructType,
+ context: TaskAttemptContext): OutputWriter = {
+ new AvroOutputWriter(path, context, schema, recordName, recordNamespace)
+ }
+}
diff --git a/src/main/scala/com/databricks/spark/avro/SchemaConverters.scala b/src/main/scala/com/databricks/spark/avro/clone/SchemaConverters.scala
similarity index 98%
rename from src/main/scala/com/databricks/spark/avro/SchemaConverters.scala
rename to src/main/scala/com/databricks/spark/avro/clone/SchemaConverters.scala
index 350d117..45bef66 100644
--- a/src/main/scala/com/databricks/spark/avro/SchemaConverters.scala
+++ b/src/main/scala/com/databricks/spark/avro/clone/SchemaConverters.scala
@@ -13,8 +13,7 @@
* See the License for the specific language governing permissions and
* limitations under the License.
*/
-// scalastyle:off
-package com.databricks.spark.avro
+package com.databricks.spark.avro.clone
import java.nio.ByteBuffer
import java.util.HashMap
@@ -29,6 +28,8 @@ import org.apache.avro.Schema.Type._
import org.apache.spark.sql.Row
import org.apache.spark.sql.types._
+// scalastyle:off
+
/**
* This object contains method that are used to convert sparkSQL schemas to avro schemas and vice
* versa.
@@ -225,6 +226,7 @@ object SchemaConverters {
case BinaryType => schemaBuilder.bytesType()
case BooleanType => schemaBuilder.booleanType()
case TimestampType => schemaBuilder.longType()
+ case DateType => schemaBuilder.stringType()
case ArrayType(elementType, _) =>
val builder = getSchemaBuilder(dataType.asInstanceOf[ArrayType].containsNull)
@@ -268,6 +270,7 @@ object SchemaConverters {
case BinaryType => newFieldBuilder.bytesType()
case BooleanType => newFieldBuilder.booleanType()
case TimestampType => newFieldBuilder.longType()
+ case DateType => newFieldBuilder.stringType()
case ArrayType(elementType, _) =>
val builder = getSchemaBuilder(dataType.asInstanceOf[ArrayType].containsNull)
@@ -297,4 +300,3 @@ object SchemaConverters {
}
}
}
-// scalastyle:on
\ No newline at end of file
diff --git a/src/main/scala/com/databricks/spark/avro/clone/package.scala b/src/main/scala/com/databricks/spark/avro/clone/package.scala
new file mode 100644
index 0000000..c249610
--- /dev/null
+++ b/src/main/scala/com/databricks/spark/avro/clone/package.scala
@@ -0,0 +1,36 @@
+/*
+ * Copyright 2014 Databricks
+ *
+ * Licensed 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 com.databricks.spark.avro
+
+import org.apache.spark.sql.{DataFrame, DataFrameReader, DataFrameWriter}
+
+package object clone {
+ /**
+ * Adds a method, `avro`, to DataFrameWriter that allows you to write avro files using
+ * the DataFileWriter
+ */
+ implicit class AvroDataFrameWriter[T](writer: DataFrameWriter[T]) {
+ def avro: String => Unit = writer.format("com.databricks.spark.avro").save
+ }
+
+ /**
+ * Adds a method, `avro`, to DataFrameReader that allows you to read avro files using
+ * the DataFileReade
+ */
+ implicit class AvroDataFrameReader(reader: DataFrameReader) {
+ def avro: String => DataFrame = reader.format("com.databricks.spark.avro").load
+ }
+}
diff --git a/src/main/scala/com/spotify/spark/bigquery/BigQueryClient.scala b/src/main/scala/com/spotify/spark/bigquery/BigQueryClient.scala
index e14120e..c993b7c 100644
--- a/src/main/scala/com/spotify/spark/bigquery/BigQueryClient.scala
+++ b/src/main/scala/com/spotify/spark/bigquery/BigQueryClient.scala
@@ -104,18 +104,27 @@ private[bigquery] class BigQueryClient(conf: Configuration) {
*/
def query(sqlQuery: String): TableReference = queryCache.get(sqlQuery)
+ /**
+ * Check if the table exists - useful for pre-write logic.
+ */
+ def tableExists(table: TableReference): Boolean = {
+ new BigQueryHelper(bigquery).tableExists(table)
+ }
+
/**
* Load an Avro data set on GCS to a BigQuery table.
*/
def load(gcsPath: String, destinationTable: TableReference,
writeDisposition: WriteDisposition.Value = null,
- createDisposition: CreateDisposition.Value = null): Unit = {
+ createDisposition: CreateDisposition.Value = null,
+ schema: TableSchema = null): Unit = {
val tableName = BigQueryStrings.toString(destinationTable)
logger.info(s"Loading $gcsPath into $tableName")
var loadConfig = new JobConfigurationLoad()
.setDestinationTable(destinationTable)
.setSourceFormat("AVRO")
.setSourceUris(List(gcsPath + "/*.avro").asJava)
+ .setSchema(schema)
if (writeDisposition != null) {
loadConfig = loadConfig.setWriteDisposition(writeDisposition.toString)
}
diff --git a/src/main/scala/com/spotify/spark/bigquery/package.scala b/src/main/scala/com/spotify/spark/bigquery/package.scala
index d295c4e..a5d34c0 100644
--- a/src/main/scala/com/spotify/spark/bigquery/package.scala
+++ b/src/main/scala/com/spotify/spark/bigquery/package.scala
@@ -17,8 +17,10 @@
package com.spotify.spark
-import com.databricks.spark.avro._
-import com.google.api.services.bigquery.model.TableReference
+import com.appsflyer.spark.bigquery.BigQuerySchema
+import com.databricks.spark.avro.clone._
+import com.databricks.spark.avro.clone.SchemaConverters
+import com.google.api.services.bigquery.model.{TableReference, TableSchema}
import com.google.cloud.hadoop.fs.gcs.GoogleHadoopFileSystem
import com.google.cloud.hadoop.io.bigquery._
import org.apache.avro.Schema
@@ -147,17 +149,27 @@ package object bigquery {
sqlContext.setConf("spark.sql.avro.compression.codec", "deflate")
+ /**
+ * Check for existence of a table.
+ * @param tableRef
+ * @return
+ */
+ def bigQueryTableExists(tableRef: TableReference): Boolean = {
+ bq.tableExists(tableRef)
+ }
+
/**
* Save a [[DataFrame]] to a BigQuery table.
*/
def saveAsBigQueryTable(tableRef: TableReference,
writeDisposition: WriteDisposition.Value,
- createDisposition: CreateDisposition.Value): Unit = {
+ createDisposition: CreateDisposition.Value,
+ schema: TableSchema): Unit = {
val bucket = conf.get(BigQueryConfiguration.GCS_BUCKET_KEY)
val temp = s"spark-bigquery-${System.currentTimeMillis()}=${Random.nextInt(Int.MaxValue)}"
val gcsPath = s"gs://$bucket/hadoop/tmp/spark-bigquery/$temp"
self.write.avro(gcsPath)
- val df = bq.load(gcsPath, tableRef, writeDisposition, createDisposition)
+ val df = bq.load(gcsPath, tableRef, writeDisposition, createDisposition, schema)
delete(new Path(gcsPath))
df
}
@@ -167,11 +179,33 @@ package object bigquery {
*/
def saveAsBigQueryTable(tableSpec: String,
writeDisposition: WriteDisposition.Value = null,
- createDisposition: CreateDisposition.Value = null): Unit =
+ createDisposition: CreateDisposition.Value = null,
+ schema: TableSchema = null): Unit =
+ saveAsBigQueryTable(
+ BigQueryStrings.parseTableReference(tableSpec),
+ writeDisposition,
+ createDisposition,
+ schema)
+
+ /**
+ * Save a [[DataFrame]] to a BigQuery table. Extract the schema definition from
+ * the Dataframe and force BigQuery to use this schema when importing the Avro data.
+ * This has the benefit that Timestamp and Date fields don't get coerced to longs but
+ * instead use the types available in BigQuery.
+ */
+ def saveAsBigQueryTableWithRichSchema(tableSpec: String,
+ writeDisposition: WriteDisposition.Value = null,
+ createDisposition: CreateDisposition.Value = null): Unit = {
+
+ val schemaJson = BigQuerySchema(self)
+ val tableSchema = new TableSchema().setFields(BigQueryUtils.getSchemaFromString(schemaJson))
+
saveAsBigQueryTable(
BigQueryStrings.parseTableReference(tableSpec),
writeDisposition,
- createDisposition)
+ createDisposition,
+ tableSchema)
+ }
private def delete(path: Path): Unit = {
val fs = FileSystem.get(path.toUri, conf)
diff --git a/src/main/scala/spark/bigquery/BigQuerySchema.scala b/src/main/scala/spark/bigquery/BigQuerySchema.scala
new file mode 100644
index 0000000..9a03dcf
--- /dev/null
+++ b/src/main/scala/spark/bigquery/BigQuerySchema.scala
@@ -0,0 +1,80 @@
+/*
+ * Copyright 2016 Appsflyer.
+ *
+ * Licensed 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 com.appsflyer.spark.bigquery
+
+import org.apache.spark.sql.DataFrame
+import org.apache.spark.sql.types._
+import org.json4s.JsonAST.{JArray, JValue}
+import org.json4s.JsonDSL._
+import org.json4s.jackson.JsonMethods.{pretty, render}
+
+/**
+ * Builds BigQuery input JSON schema based on DataFrame.
+ * Example schema can be found here: https://cloud.google.com/bigquery/docs/personsDataSchema.json
+ */
+object BigQuerySchema {
+
+ private def getMode(field: StructField) = {
+ field.dataType match {
+ case ArrayType(_, _) => "REPEATED"
+ case _ => if (field.nullable) "NULLABLE" else "REQUIRED"
+ }
+ }
+
+ private def getTypeName(dataType: DataType) = {
+ dataType match {
+ case ByteType | ShortType | IntegerType | LongType => "INTEGER"
+ case FloatType | DoubleType => "FLOAT"
+ case _: DecimalType | StringType => "STRING"
+ case BinaryType => "BYTES"
+ case BooleanType => "BOOLEAN"
+ case TimestampType => "TIMESTAMP"
+ case ArrayType(_, _) | MapType(_, _, _) | _: StructType => "RECORD"
+ case DateType => "STRING"
+ }
+ }
+
+ private def typeToJson(field: StructField, dataType: DataType): JValue = {
+ dataType match {
+ case structType: StructType =>
+ ("type" -> getTypeName(dataType)) ~
+ ("fields" -> structType.fields.map(fieldToJson(_)).toList)
+ case arrayType: ArrayType =>
+ arrayType.elementType match {
+ case _: ArrayType =>
+ throw new IllegalArgumentException(s"Multidimensional arrays are not supported: ${field.name}")
+ case other =>
+ typeToJson(field, other)
+ }
+ case mapType: MapType =>
+ throw new IllegalArgumentException(s"Unsupported type: ${dataType}")
+ case other =>
+ ("type" -> getTypeName(dataType))
+ }
+ }
+
+ private def fieldToJson(field: StructField): JValue = {
+ ("name" -> field.name) ~
+ ("mode" -> getMode(field)) merge
+ typeToJson(field, field.dataType)
+ }
+
+ def apply(df: DataFrame): String = {
+ pretty(render(JArray(df.schema.fields.map(fieldToJson(_)).toList)))
+ }
+}
diff --git a/version.sbt b/version.sbt
index cfeb331..74410b2 100644
--- a/version.sbt
+++ b/version.sbt
@@ -1 +1 @@
-version in ThisBuild := "0.2.1-SNAPSHOT"
+version in ThisBuild := "0.2.1-BRIGADE02"