diff --git a/.github/codeowners b/.github/codeowners new file mode 100644 index 0000000..20a7ed5 --- /dev/null +++ b/.github/codeowners @@ -0,0 +1,2 @@ +* @loanpal-engineering/business-intelligence-bi +.github/codeowners @loanpal-engineering/security @loanpal-engineering/DevOps diff --git a/.gitignore b/.gitignore index 07a31f7..0254d49 100644 --- a/.gitignore +++ b/.gitignore @@ -10,3 +10,6 @@ target/ project/target dependency-reduced-pom.xml /bin/ +.DS_Store +.databricks +.vscode/ \ No newline at end of file diff --git a/README.md b/README.md index 9d24698..bf65d19 100644 --- a/README.md +++ b/README.md @@ -62,6 +62,8 @@ $ bin/spark-shell --packages com.springml:spark-salesforce_2.11:1.1.3 * `timeout`: (Optional) The maximum time spent polling for the completion of bulk query job. This option can only be used when `bulk` is `true`. * `externalIdFieldName`: (Optional) The name of the field used as the external ID for Salesforce Object. This value is only used when doing an update or upsert. Default "Id". * `queryAll`: (Optional) Toggle to retrieve deleted and archived records for SOQL queries. Default value is `false`. +### Options only supported for fetching Salesforce Objects. +* `batchSize`: (Optional) maximum number of records per batch when performing updates. Defaults to 5000 (note that batches greater than 10000 will result in a error) ### Scala API diff --git a/build.sbt b/build.sbt deleted file mode 100644 index e6f5f62..0000000 --- a/build.sbt +++ /dev/null @@ -1,89 +0,0 @@ -name := "spark-salesforce" - -version := "1.1.3" - -organization := "com.springml" - -scalaVersion := "2.11.8" - -resolvers += "sonatype-snapshots" at "https://oss.sonatype.org/content/repositories/snapshots/" - -libraryDependencies ++= Seq( - "com.force.api" % "force-wsc" % "40.0.0", - "com.force.api" % "force-partner-api" % "40.0.0", - "com.springml" % "salesforce-wave-api" % "1.0.10", - "org.mockito" % "mockito-core" % "2.0.31-beta" -) - -parallelExecution in Test := false - -resolvers += Resolver.url("artifactory", url("http://scalasbt.artifactoryonline.com/scalasbt/sbt-plugin-releases"))(Resolver.ivyStylePatterns) - -resolvers += "Typesafe Repository" at "http://repo.typesafe.com/typesafe/releases/" - -resolvers += "sonatype-releases" at "https://oss.sonatype.org/content/repositories/releases/" - -resolvers += "sonatype-snapshots" at "https://oss.sonatype.org/content/repositories/snapshots/" - -resolvers += "Spark Package Main Repo" at "https://dl.bintray.com/spark-packages/maven" - -libraryDependencies += "org.scalatest" %% "scalatest" % "2.2.1" % "test" -libraryDependencies += "com.madhukaraphatak" %% "java-sizeof" % "0.1" -libraryDependencies += "com.fasterxml.jackson.dataformat" % "jackson-dataformat-xml" % "2.4.4" -libraryDependencies += "org.codehaus.woodstox" % "woodstox-core-asl" % "4.4.0" - -// Spark Package Details (sbt-spark-package) -spName := "springml/spark-salesforce" - -spAppendScalaVersion := true - -sparkVersion := "2.2.0" - -sparkComponents += "sql" - -publishMavenStyle := true - -spIncludeMaven := true - -spShortDescription := "Spark Salesforce Wave Connector" - -spDescription := """Spark Salesforce Wave Connector - | - Creates Salesforce Wave Datasets using dataframe - | - Constructs Salesforce Wave dataset's metadata using schema present in dataframe - | - Can use custom metadata for constructing Salesforce Wave dataset's metadata""".stripMargin - -// licenses += "Apache-2.0" -> url("http://opensource.org/licenses/Apache-2.0") - -credentials += Credentials(Path.userHome / ".ivy2" / ".credentials") - -publishTo := { - val nexus = "https://oss.sonatype.org/" - if (version.value.endsWith("SNAPSHOT")) - Some("snapshots" at nexus + "content/repositories/snapshots") - else - Some("releases" at nexus + "service/local/staging/deploy/maven2") -} - -pomExtra := ( - https://github.com/springml/spark-salesforce - - - Apache License, Verision 2.0 - http://www.apache.org/licenses/LICENSE-2.0.html - repo - - - - scm:git:github.com/springml/spark-salesforce - scm:git:git@github.com:springml/spark-salesforce - github.com/springml/spark-salesforce - - - - springml - Springml - http://www.springml.com - - ) - - diff --git a/pom.xml b/pom.xml new file mode 100644 index 0000000..75e413c --- /dev/null +++ b/pom.xml @@ -0,0 +1,202 @@ + + + 4.0.0 + com.goodleap + spark-salesforce + jar + spark-salesforce + 1.1.6 + spark-salesforce + + com.springml + + https://github.com/springml/spark-salesforce + + + Apache License, Verision 2.0 + http://www.apache.org/licenses/LICENSE-2.0.html + repo + + + + scm:git:github.com/springml/spark-salesforce + scm:git:git@github.com:springml/spark-salesforce + github.com/springml/spark-salesforce + + + 1.8 + 1.8 + UTF-8 + 3.3.0 + 2.12.11 + 2.12 + + + + springml + Springml + http://www.springml.com + + + + + org.scala-lang + scala-library + ${scala.version} + + + org.apache.spark + spark-sql_${scala.compat.version} + ${spark.version} + provided + + + + com.force.api + force-wsc + 53.0.0 + + + com.force.api + force-partner-api + 53.0.0 + + + com.springml + salesforce-wave-api + 1.0.8-loanpal + + + + + + + + + + + + + + org.mockito + mockito-core + 2.0.31-beta + + + org.scalatest + scalatest_2.12 + 3.0.1 + test + + + endolabs.salesforce + bulkv2 + 1.0.0 + + + com.squareup.okhttp3 + okhttp + 3.14.2 + + + com.squareup.okhttp3 + logging-interceptor + 3.14.2 + + + com.frejo + force-rest-api + 0.0.42 + + + org.codehaus.woodstox + woodstox-core-asl + 4.4.1 + + + org.codehaus.woodstox + stax2-api + + + + + + + + + + + src/main/scala + + + + net.alchim31.maven + scala-maven-plugin + 3.4.6 + + + + compile + testCompile + + + + -Xss16m + -Xms1028m + -Xmx4096m + + + + + + + org.apache.maven.plugins + maven-shade-plugin + 3.4.1 + + + package + + shade + + + + + com.fasterxml.jackson.dataformat + com.shaded.fasterxml.jackson.dataformat + + + + + + + + org.apache.maven.plugins + maven-shade-plugin + 3.2.1 + + + package + + shade + + + + + + + *:* + + META-INF/*.SF + META-INF/*.DSA + META-INF/*.RSA + + + + uber-${project.artifactId}-${project.version} + + + + + \ No newline at end of file diff --git a/project/plugins.sbt b/project/plugins.sbt deleted file mode 100644 index f996987..0000000 --- a/project/plugins.sbt +++ /dev/null @@ -1,7 +0,0 @@ -resolvers += "bintray-spark-packages" at "https://dl.bintray.com/spark-packages/maven/" - -addSbtPlugin("org.spark-packages" % "sbt-spark-package" % "0.2.6") -addSbtPlugin("com.eed3si9n" % "sbt-assembly" % "0.14.3") -addSbtPlugin("com.typesafe.sbteclipse" % "sbteclipse-plugin" % "4.0.0") -addSbtPlugin("org.xerial.sbt" % "sbt-sonatype" % "0.5.0") -addSbtPlugin("com.jsuereth" % "sbt-pgp" % "1.1.0") diff --git a/src/main/scala/com/springml/spark/salesforce/BulkRelation.scala b/src/main/scala/com/springml/spark/salesforce/BulkRelation.scala index e7c5333..694fca3 100644 --- a/src/main/scala/com/springml/spark/salesforce/BulkRelation.scala +++ b/src/main/scala/com/springml/spark/salesforce/BulkRelation.scala @@ -34,7 +34,9 @@ case class BulkRelation( userSchema: StructType, sqlContext: SQLContext, inferSchema: Boolean, - timeout: Long) extends BaseRelation with TableScan { + timeout: Long, + maxCharsPerColumn: Int, + maxColumns: Int) extends BaseRelation with TableScan { import sqlContext.sparkSession.implicits._ @@ -71,12 +73,15 @@ case class BulkRelation( // Use Csv parser to split CSV by rows to cover edge cases (ex. escaped characters, new line within string, etc) def splitCsvByRows(csvString: String): Seq[String] = { + if (csvString == "Records not found for this query") Seq.empty // The CsvParser interface only interacts with IO, so StringReader and StringWriter val inputReader = new StringReader(csvString) val parserSettings = new CsvParserSettings() parserSettings.setLineSeparatorDetectionEnabled(true) parserSettings.getFormat.setNormalizedNewline(' ') + parserSettings.setMaxCharsPerColumn(maxCharsPerColumn) + parserSettings.setMaxColumns(maxColumns) val readerParser = new CsvParser(parserSettings) val parsedInput = readerParser.parseAll(inputReader).asScala @@ -90,7 +95,7 @@ case class BulkRelation( val writer = new CsvWriter(outputWriter, writerSettings) parsedInput.foreach { writer.writeRow(_) } - outputWriter.toString.lines.toList + outputWriter.toString.split("\n").toList } splitCsvByRows(result) @@ -109,6 +114,7 @@ case class BulkRelation( .option("quote", "\"") .option("escape", "\"") .option("multiLine", true) + .option("maxColumns", maxColumns) .csv(csvData) } else { bulkAPI.closeJob(jobId) diff --git a/src/main/scala/com/springml/spark/salesforce/DataWriter.scala b/src/main/scala/com/springml/spark/salesforce/DataWriter.scala index 2565dad..c6ce2d5 100644 --- a/src/main/scala/com/springml/spark/salesforce/DataWriter.scala +++ b/src/main/scala/com/springml/spark/salesforce/DataWriter.scala @@ -28,18 +28,18 @@ import org.apache.spark.sql.SaveMode * It uses Partner External Metadata SOAP API to write the dataset */ class DataWriter ( - val userName: String, - val password: String, - val login: String, - val version: String, - val datasetName: String, - val appName: String - ) extends Serializable { + val userName: String, + val password: String, + val login: String, + val version: String, + val datasetName: String, + val appName: String + ) extends Serializable { @transient val logger = Logger.getLogger(classOf[DataWriter]) def writeMetadata(metaDataJson: String, - mode: SaveMode, - upsert: Boolean): Option[String] = { + mode: SaveMode, + upsert: Boolean): Option[String] = { val partnerConnection = createConnection(userName, password, login, version) val oper = operation(mode, upsert) @@ -62,19 +62,25 @@ class DataWriter ( Some(saveResult.getId) } else { logger.error("failed to write metadata") - println("******************************************************************") - println("failed to write metadata") + logger.error("******************************************************************") + logger.error("failed to write metadata") logSaveResultError(saveResult) - println("******************************************************************") - println(metaDataJson) - println("******************************************************************") + logger.error("******************************************************************") + logger.error(metaDataJson) + logger.error("******************************************************************") None } }).head } def writeData(rdd: RDD[Row], metadataId: String): Boolean = { - val csvRDD = rdd.map(row => row.toSeq.map(value => Utils.rowValue(value)).mkString(",")) + val csvRDD = rdd.map{row => + val schema = row.schema.fields + row.toSeq.indices.map( + index => Utils.cast(row, schema(index).dataType, index) + ).mkString(",") + } + csvRDD.mapPartitionsWithIndex { case (index, iterator) => { @transient val logger = Logger.getLogger(classOf[DataWriter]) @@ -128,7 +134,6 @@ class DataWriter ( saved } - private def operation(mode: SaveMode, upsert: Boolean): String = { if (upsert) { logger.warn("Ignoring SaveMode as upsert set to true") diff --git a/src/main/scala/com/springml/spark/salesforce/DatasetRelation.scala b/src/main/scala/com/springml/spark/salesforce/DatasetRelation.scala index 4dc46d0..bf01854 100644 --- a/src/main/scala/com/springml/spark/salesforce/DatasetRelation.scala +++ b/src/main/scala/com/springml/spark/salesforce/DatasetRelation.scala @@ -21,26 +21,26 @@ import scala.collection.JavaConversions.{asScalaBuffer, mapAsScalaMap} * Relation class for reading data from Salesforce and construct RDD */ case class DatasetRelation( - waveAPI: WaveAPI, - forceAPI: ForceAPI, - query: String, - userSchema: StructType, - sqlContext: SQLContext, - resultVariable: Option[String], - pageSize: Int, - sampleSize: Int, - encodeFields: Option[String], - inferSchema: Boolean, - replaceDatasetNameWithId: Boolean, - sdf: SimpleDateFormat, - queryAll: Boolean) extends BaseRelation with TableScan { + waveAPI: WaveAPI, + forceAPI: ForceAPI, + query: String, + userSchema: StructType, + sqlContext: SQLContext, + resultVariable: Option[String], + pageSize: Int, + sampleSize: Int, + encodeFields: Option[String], + inferSchema: Boolean, + replaceDatasetNameWithId: Boolean, + sdf: SimpleDateFormat, + queryAll: Boolean) extends BaseRelation with TableScan { private val logger = Logger.getLogger(classOf[DatasetRelation]) val records = read() def read(): java.util.List[java.util.Map[String, String]] = { - var records: java.util.List[java.util.Map[String, String]]= null + var records: java.util.List[java.util.Map[String, String]] = null // Query getting executed here if (waveAPI != null) { records = queryWave() @@ -52,7 +52,7 @@ case class DatasetRelation( } private def queryWave(): java.util.List[java.util.Map[String, String]] = { - var records: java.util.List[java.util.Map[String, String]]= null + var records: java.util.List[java.util.Map[String, String]] = null var saql = query if (replaceDatasetNameWithId) { @@ -77,7 +77,7 @@ case class DatasetRelation( records } - def replaceDatasetNameWithId(query : String, startIndex : Integer) : String = { + def replaceDatasetNameWithId(query: String, startIndex: Integer): String = { var modQuery = query logger.debug("start Index : " + startIndex) @@ -101,22 +101,25 @@ case class DatasetRelation( } private def querySF(): java.util.List[java.util.Map[String, String]] = { - var records: java.util.List[java.util.Map[String, String]]= null + var records: java.util.List[java.util.Map[String, String]] = null - var resultSet = forceAPI.query(query, queryAll) - records = resultSet.filterRecords() + var resultSet = forceAPI.query(query, queryAll) + records = resultSet.filterRecords() - while (!resultSet.isDone()) { - resultSet = forceAPI.queryMore(resultSet) - records.addAll(resultSet.filterRecords()) - } + while (!resultSet.isDone()) { + resultSet = forceAPI.queryMore(resultSet) + records.addAll(resultSet.filterRecords()) + } - return records + return records } + // NOTE: this function is only called when reading the data NOT when writing the data private def cast(fieldValue: String, toType: DataType, - nullable: Boolean = true, fieldName: String): Any = { - if (fieldValue == "" && nullable && !toType.isInstanceOf[StringType]) { + nullable: Boolean = true, fieldName: String): Any = { + if (fieldValue == null || fieldValue.isEmpty) { + null + } else if (fieldValue == "" && nullable && !toType.isInstanceOf[StringType]) { null } else { toType match { @@ -150,7 +153,7 @@ case class DatasetRelation( } } - private def shouldEncode(fieldName: String) : Boolean = { + private def shouldEncode(fieldName: String): Boolean = { if (encodeFields != null && encodeFields.isDefined) { val toBeEncodedField = encodeFields.get.split(",") return toBeEncodedField.contains(fieldName) @@ -181,14 +184,14 @@ case class DatasetRelation( sqlContext.sparkContext.parallelize(sampleRowArray) } - private def getSampleSize : Integer = { + private def getSampleSize: Integer = { // If the record is less than sampleSize, then the whole data is used as sample val totalRecordsSize = records.size() logger.debug("Total Record Size: " + totalRecordsSize) if (totalRecordsSize < sampleSize) { logger.debug("Total Record Size " + totalRecordsSize - + " is Smaller than Sample Size " - + sampleSize + ". So total records are used for sampling") + + " is Smaller than Sample Size " + + sampleSize + ". So total records are used for sampling") totalRecordsSize } else { sampleSize @@ -198,7 +201,7 @@ case class DatasetRelation( private def header: Array[String] = { val sampleList = sample - var header : Array[String] = null + var header: Array[String] = null for (currentRecord <- sampleList) { logger.debug("record size " + currentRecord.size()) val recordHeader = new Array[String](currentRecord.size()) @@ -273,7 +276,7 @@ case class DatasetRelation( sqlContext.sparkContext.parallelize(rowArray) } - private def fieldValue(row: java.util.Map[String, String], name: String) : String = { + private def fieldValue(row: java.util.Map[String, String], name: String): String = { if (row.contains(name)) { row(name) } else { diff --git a/src/main/scala/com/springml/spark/salesforce/DefaultSource.scala b/src/main/scala/com/springml/spark/salesforce/DefaultSource.scala index fc077a6..060cc08 100644 --- a/src/main/scala/com/springml/spark/salesforce/DefaultSource.scala +++ b/src/main/scala/com/springml/spark/salesforce/DefaultSource.scala @@ -26,6 +26,7 @@ import org.apache.spark.sql.types.StructType import org.apache.spark.sql.{DataFrame, SQLContext, SaveMode} import scala.collection.mutable.ListBuffer +import scala.util.{Failure, Success, Try} /** * Default source for Salesforce wave data source. @@ -75,6 +76,7 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr val bulkFlag = flag(bulkStr, "bulk") val queryAllStr = parameters.getOrElse("queryAll", "false") + // val maxColumns = parameters.getOrElse("maxColumns", "512") val queryAllFlag = flag(queryAllStr, "queryAll") validateMutualExclusive(saql, soql, "saql", "soql") @@ -123,7 +125,25 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr val encodeFields = parameters.get("encodeFields") val monitorJob = parameters.getOrElse("monitorJob", "false") val externalIdFieldName = parameters.getOrElse("externalIdFieldName", "Id") + val batchSizeStr = parameters.getOrElse("batchSize", "5000") + val bulkApiV2Str = parameters.getOrElse("bulkApiV2", "false") + val batchSize = Try(batchSizeStr.toInt) match { + case Success(v)=> v + case Failure(e)=> { + val errorMsg = "batchSize parameter not an integer." + logger.error(errorMsg) + throw new Exception(errorMsg) + } + } + val bulkAPIV2 = Try(bulkApiV2Str.toBoolean) match { + case Success(v) => v + case Failure(e) => { + val errorMsg = "bulkAPIV2 parameter not a boolean." + logger.error(errorMsg) + throw new Exception(errorMsg) + } + } validateMutualExclusive(datasetName, sfObject, "datasetName", "sfObject") if (datasetName.isDefined) { @@ -140,8 +160,14 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr flag(upsert, "upsert"), flag(monitorJob, "monitorJob"), data, metadataFile) } else { logger.info("Updating Salesforce Object") - updateSalesforceObject(username, password, login, version, sfObject.get, mode, - flag(upsert, "upsert"), externalIdFieldName, data) + if(bulkAPIV2) { + updateSalesforceObjectV2(username, password, login, version, sfObject.get, mode, + flag(upsert, "upsert"), externalIdFieldName, batchSize, data) + } else { + updateSalesforceObject(username, password, login, version, sfObject.get, mode, + flag(upsert, "upsert"), externalIdFieldName, batchSize, data) + } + } return createReturnRelation(data) @@ -156,6 +182,7 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr mode: SaveMode, upsert: Boolean, externalIdFieldName: String, + batchSize: Integer, data: DataFrame) { val csvHeader = Utils.csvHeadder(data.schema) @@ -164,7 +191,7 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr val repartitionedRDD = Utils.repartition(data.rdd) logger.info("no of partitions after repartitioning is " + repartitionedRDD.partitions.length) - val writer = new SFObjectWriter(username, password, login, version, sfObject, mode, upsert, externalIdFieldName, csvHeader) + val writer = new SFObjectWriter(username, password, login, version, sfObject, mode, upsert, externalIdFieldName, csvHeader, batchSize) logger.info("Writing data") val successfulWrite = writer.writeData(repartitionedRDD) logger.info(s"Writing data was successful was $successfulWrite") @@ -174,6 +201,30 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr } + private def updateSalesforceObjectV2( + username: String, + password: String, + login: String, + version: String, + sfObject: String, + mode: SaveMode, + upsert: Boolean, + externalIdFieldName: String, + batchSize: Integer, + data: DataFrame) { + + val csvHeader = Utils.csvHeadder(data.schema) + + val writer = new SFObjectWriter2(username, password, login, version, sfObject, mode, upsert, externalIdFieldName, csvHeader, batchSize) + logger.info("Writing data") + val successfulWrite = writer.writeData(data.rdd) + logger.info(s"Writing data was successful was $successfulWrite") + if (!successfulWrite) { + sys.error("Unable to update salesforce object") + } + + } + private def createBulkRelation( sqlContext: SQLContext, username: String, @@ -190,6 +241,13 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr throw new Exception("sfObject must not be empty when performing bulk query") } + val maxCharsPerColumnStr = parameters.getOrElse("maxCharsPerColumn", "4096") + val maxCharsPerColumn = try { + maxCharsPerColumnStr.toInt + } catch { + case e: Exception => throw new Exception("maxCharsPerColumn must be a valid integer") + } + val timeoutStr = parameters.getOrElse("timeout", "600000") val timeout = try { timeoutStr.toLong @@ -216,6 +274,11 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr customHeaders += new BasicHeader("Sforce-Enable-PKChunking", "true") } } + val maxColumns = try { + parameters.getOrElse("maxColumns", "512").toInt + } catch { + case e: Exception => throw new Exception("max columns needs to be a valid integer") + } BulkRelation( username, @@ -228,7 +291,9 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr schema, sqlContext, inferSchemaFlag, - timeout + timeout, + maxCharsPerColumn, + maxColumns ) } diff --git a/src/main/scala/com/springml/spark/salesforce/SFObjectWriter.scala b/src/main/scala/com/springml/spark/salesforce/SFObjectWriter.scala index 1856bbb..47e84bf 100644 --- a/src/main/scala/com/springml/spark/salesforce/SFObjectWriter.scala +++ b/src/main/scala/com/springml/spark/salesforce/SFObjectWriter.scala @@ -8,73 +8,92 @@ import com.springml.salesforce.wave.api.BulkAPI import com.springml.salesforce.wave.util.WaveAPIConstants import com.springml.salesforce.wave.model.JobInfo +import scala.collection.JavaConverters.asScalaBufferConverter +import scala.util.Try + /** * Write class responsible for update Salesforce object using data provided in dataframe * First column of dataframe contains Salesforce Object * Next subsequent columns are fields to be updated */ -class SFObjectWriter ( - val username: String, - val password: String, - val login: String, - val version: String, - val sfObject: String, - val mode: SaveMode, - val upsert: Boolean, - val externalIdFieldName: String, - val csvHeader: String - ) extends Serializable { +class SFObjectWriter( + val username: String, + val password: String, + val login: String, + val version: String, + val sfObject: String, + val mode: SaveMode, + val upsert: Boolean, + val externalIdFieldName: String, + val csvHeader: String, + val batchSize: Integer + ) extends Serializable { @transient val logger = Logger.getLogger(classOf[SFObjectWriter]) + // Single BulkAPI instance reused for all driver-side operations (createJob, + // closeJob, isCompleted). SFConfig.getPartnerConnection() is lazy, so this + // performs exactly ONE SOAP login for the entire driver-side lifecycle. + // Marked @transient because BulkAPI is not serializable — executors must + // create their own instances via newBulkAPI(). + @transient private lazy val driverBulkAPI: BulkAPI = newBulkAPI() + + // Fresh instance for executor-side operations (addBatch inside + // mapPartitionsWithIndex). Each partition authenticates once. + private def newBulkAPI(): BulkAPI = { + APIFactory.getInstance().bulkAPI(username, password, login, version) + } + def writeData(rdd: RDD[Row]): Boolean = { - val csvRDD = rdd.map(row => row.toSeq.map(value => Utils.rowValue(value)).mkString(",")) + + val csvRDD = rdd.map { row => + val schema = row.schema.fields + row.toSeq.indices.map( + index => Utils.cast(row, schema(index).dataType, index) + ).mkString(",") + } + + val partitionCnt = (1 + csvRDD.count() / batchSize).toInt + val partitionedRDD = csvRDD.repartition(partitionCnt) val jobInfo = new JobInfo(WaveAPIConstants.STR_CSV, sfObject, operation(mode, upsert)) jobInfo.setExternalIdFieldName(externalIdFieldName) - val jobId = bulkAPI.createJob(jobInfo).getId + val jobId = driverBulkAPI.createJob(jobInfo).getId - csvRDD.mapPartitionsWithIndex { + partitionedRDD.mapPartitionsWithIndex { case (index, iterator) => { val records = iterator.toArray.mkString("\n") - var batchInfoId : String = null + var batchInfoId: String = null if (records != null && !records.isEmpty()) { val data = csvHeader + "\n" + records - val batchInfo = bulkAPI.addBatch(jobId, data) + val batchInfo = newBulkAPI().addBatch(jobId, data) batchInfoId = batchInfo.getId } val success = (batchInfoId != null) - // Job status will be checked after completing all batches List(success).iterator } }.reduce((a, b) => a & b) - bulkAPI.closeJob(jobId) + driverBulkAPI.closeJob(jobId) var i = 1 while (i < 999999) { - if (bulkAPI.isCompleted(jobId)) { + if (driverBulkAPI.isCompleted(jobId)) { logger.info("Job completed") return true } - logger.info("Job not completed, waiting...") Thread.sleep(200) i = i + 1 } print("Returning false...") - logger.info("Job not completed. Timeout..." ) + logger.info("Job not completed. Timeout...") false } - // Create new instance of BulkAPI every time because Spark workers cannot serialize the object - private def bulkAPI(): BulkAPI = { - APIFactory.getInstance().bulkAPI(username, password, login, version) - } - private def operation(mode: SaveMode, upsert: Boolean): String = { if (upsert) { "upsert" diff --git a/src/main/scala/com/springml/spark/salesforce/SFObjectWriter2.scala b/src/main/scala/com/springml/spark/salesforce/SFObjectWriter2.scala new file mode 100644 index 0000000..f3a26a8 --- /dev/null +++ b/src/main/scala/com/springml/spark/salesforce/SFObjectWriter2.scala @@ -0,0 +1,149 @@ +package com.springml.spark.salesforce + +import java.io.BufferedReader + +import com.springml.salesforce.wave.api.{APIFactory, BulkAPI} +import com.springml.salesforce.wave.model.JobInfo +import com.springml.salesforce.wave.util.WaveAPIConstants +import org.apache.log4j.Logger +import org.apache.spark.rdd.RDD +import org.apache.spark.sql.{Row, SaveMode} +import endolabs.salesforce.bulkv2.{AccessToken, Bulk2Client, Bulk2ClientBuilder} +import com.force.api._ +import endolabs.salesforce.bulkv2.`type`.OperationEnum +import endolabs.salesforce.bulkv2.response.GetJobInfoResponse + +import scala.util.Try +import scala.collection.JavaConverters.{asScalaBufferConverter, asScalaIteratorConverter} + +class SFObjectWriter2(val username: String, + val password: String, + val login: String, + val version: String, + val sfObject: String, + val mode: SaveMode, + val upsert: Boolean, + val externalIdFieldName: String, + val csvHeader: String, + val batchSize: Integer) extends Serializable { + + @transient val logger = Logger.getLogger(classOf[SFObjectWriter]) + + def writeData(rdd: RDD[Row]): Boolean = { + + val csvRDD = rdd.map { row => + val schema = row.schema.fields + row.toSeq.indices.map( + index => Utils.cast(row, schema(index).dataType, index) + ).mkString(",") + } + + val OperationEnum = operation(mode, upsert) + var bulkJobIDs = csvRDD.mapPartitionsWithIndex { + case (index, iterator) => { + + val records = iterator.toArray.mkString("\n") + var batchInfoId: String = null + var id = "" + if (records != null && !records.isEmpty()) { + val partitionClient = newClient() + val createResponse = partitionClient.createJob(sfObject, OperationEnum) + id = createResponse.getId + val data = csvHeader + "\n" + records + partitionClient.uploadJobData(id, data) + partitionClient.closeJob(createResponse.getId) + } + + List(id).iterator + } + }.collect().filter(_ != null) + val allIds = bulkJobIDs + + // Single client instance for all driver-side polling (1 SOAP login total) + val pollingClient = newClient() + + var i = 1 + var failedRecords = 0 + val TIMEOUT_MAX = 7000 + val BREAK_LOOP = 9000 + var isEmpty = false + while (i < TIMEOUT_MAX) { + var data: Array[GetJobInfoResponse] = Array() + try { + data = bulkJobIDs.map{ + id => + pollingClient.getJobInfo(id) + } + failedRecords += data.map(x => x.getNumberRecordsFailed.toInt).sum + + data.foreach{ + x => if (x.getNumberRecordsFailed != 0) { + logger.info(s"${x.getRetries} number of retries") + val results = new BufferedReader(pollingClient.getJobFailedRecordResults(x.getId)) + results.lines().iterator().asScala.foreach{ + line => + println(line) + logger.info(line) + } + } + } + println(s"${data.count(x=> x.isFinished)} finished jobs ${data.length} remanining jobs") + println(s"${data.take(10).map(x => (x.getId, x.getState.toJsonValue)).mkString(",")} sample states") + } catch { + case e: Exception => println(e) + bulkJobIDs = allIds + } + if (data.count(x=> x.isFinished) == data.length) { + i = BREAK_LOOP + isEmpty = data.isEmpty + } else { + bulkJobIDs = data.filter(x => !x.isFinished).map(x => x.getId) + logger.info("Job not completed, waiting...") + Thread.sleep(10000) + i = i + 1 + } + } + if(isEmpty) { + return true + } + + if (i == BREAK_LOOP && failedRecords == 0){ + return true + } else if(failedRecords != 0) { + logger.info("Job failed. Timeout...") + return true + } + print("Returning false...") + logger.info("Job not completed. Timeout...") + false + + } + + // Authenticate once, reuse session for building Bulk2Client instances. + // ForceApi is not serializable, so executors call newClient() to get their own. + private def newClient(): Bulk2Client = { + val api = new ForceApi(new ApiConfig() + .setUsername(username) + .setPassword(password) + .setLoginEndpoint(login) + .setApiVersion(ApiVersion.V48)) + val session = api.getSession + new Bulk2ClientBuilder() + .withSessionId(session.getAccessToken, session.getApiEndpoint) + .build() + } + + private def operation(mode: SaveMode, upsert: Boolean): OperationEnum = { + if (upsert) { + OperationEnum.UPSERT + } else if (mode != null && SaveMode.Overwrite.name().equalsIgnoreCase(mode.name())) { + OperationEnum.UPDATE + } else if (mode != null && SaveMode.Append.name().equalsIgnoreCase(mode.name())) { + OperationEnum.INSERT + } else { + logger.warn("SaveMode " + mode + " Not supported. Using 'insert' operation") + OperationEnum.INSERT + } + } + +} diff --git a/src/main/scala/com/springml/spark/salesforce/Utils.scala b/src/main/scala/com/springml/spark/salesforce/Utils.scala index e3f0363..a337ebc 100644 --- a/src/main/scala/com/springml/spark/salesforce/Utils.scala +++ b/src/main/scala/com/springml/spark/salesforce/Utils.scala @@ -16,24 +16,23 @@ package com.springml.spark.salesforce -import scala.io.Source -import scala.util.parsing.json._ +import java.text.SimpleDateFormat + +//import com.madhukaraphatak.sizeof.SizeEstimator +import org.apache.spark.util.SizeEstimator +import com.sforce.soap.partner.fault.UnexpectedErrorFault import com.sforce.soap.partner.{Connector, PartnerConnection, SaveResult} import com.sforce.ws.ConnectorConfig -import com.madhukaraphatak.sizeof.SizeEstimator +import com.springml.spark.salesforce.metadata.MetadataConstructor import org.apache.log4j.Logger import org.apache.spark.rdd.RDD import org.apache.spark.sql.Row -import org.apache.spark.sql.types.{DoubleType, IntegerType, StructType} - -import scala.collection.immutable.HashMap -import com.springml.spark.salesforce.metadata.MetadataConstructor -import com.sforce.soap.partner.sobject.SObject -import scala.concurrent.duration._ -import com.sforce.soap.partner.fault.UnexpectedErrorFault +import org.apache.spark.sql.types.{BooleanType, DataType, StringType, StructType} import scala.concurrent.duration.FiniteDuration +import scala.io.Source import scala.util.Try +import scala.util.parsing.json._ /** * Utility to construct metadata and repartition RDD @@ -42,7 +41,7 @@ object Utils extends Serializable { @transient val logger = Logger.getLogger("Utils") def createConnection(username: String, password: String, - login: String, version: String):PartnerConnection = { + login: String, version: String): PartnerConnection = { val config = new ConnectorConfig() config.setUsername(username) config.setPassword(password) @@ -52,13 +51,17 @@ object Utils extends Serializable { Connector.newConnection(config) } + @transient val formatter = new SimpleDateFormat("yyyy-MM-dd'T'HH:mm:ss.SSS") + def logSaveResultError(result: SaveResult): Unit = { result.getErrors.map(error => { logger.error(error.getMessage) println(error.getMessage) error.getFields.map(logger.error(_)) - error.getFields.map { println } + error.getFields.map { + println + } }) } @@ -92,24 +95,24 @@ object Utils extends Serializable { totalSize } - def rddSize(rdd: RDD[Row]) : Long = { + def rddSize(rdd: RDD[Row]): Long = { rowSize(rdd.collect()) } - def rowSize(rows: Array[Row]) : Long = { - var sizeOfRows = 0l - for (row <- rows) { - // Converting to bytes - val rowSize = SizeEstimator.estimate(row.toSeq.map { value => rowValue(value) }.mkString(",")) - sizeOfRows += rowSize - } + def rowSize(rows: Array[Row]): Long = { + var sizeOfRows = 0l + for (row <- rows) { + // Converting to bytes + val rowSize = SizeEstimator.estimate(row.toSeq.map { value => rowValue(value) }.mkString(",")) + sizeOfRows += rowSize + } - sizeOfRows + sizeOfRows } - def rowValue(rowVal: Any) : String = { + def rowValue(rowVal: Any): String = { if (rowVal == null) { - "" + "#N/A" } else { var value = rowVal.toString() if (value.contains("\"")) { @@ -122,6 +125,24 @@ object Utils extends Serializable { } } + def cast(row: Row, toType: DataType, index: Int): String = { + toType match { + case _: BooleanType => { + // salesforce doesn't allow null booleans + Option(row.getAs[Boolean](index)).getOrElse(false).toString + } + case _: StringType => { + val fieldValue = row.getAs[String](index) + if (fieldValue == "") { + rowValue(null) + } else { + rowValue(fieldValue) + } + } + case _ => rowValue(row.get(index)) + } + } + def metadataConfig(usersMetadataConfig: Option[String]) = { var systemMetadataConfig = readMetadataConfig() if (usersMetadataConfig != null && usersMetadataConfig.isDefined) { @@ -132,15 +153,15 @@ object Utils extends Serializable { systemMetadataConfig } - def csvHeadder(schema: StructType) : String = { + def csvHeadder(schema: StructType): String = { schema.fields.map(field => field.name).mkString(",") } def metadata( - metadataFile: Option[String], - usersMetadataConfig: Option[String], - schema: StructType, - datasetName: String) : String = { + metadataFile: Option[String], + usersMetadataConfig: Option[String], + schema: StructType, + datasetName: String): String = { if (metadataFile != null && metadataFile.isDefined) { logger.info("Using provided Metadata Configuration") @@ -155,8 +176,8 @@ object Utils extends Serializable { } def monitorJob(objId: String, username: String, password: - String, login: String, version: String) : Boolean = { - var partnerConnection = Utils.createConnection(username, password, login, version) + String, login: String, version: String): Boolean = { + val partnerConnection = Utils.createConnection(username, password, login, version) try { monitorJob(partnerConnection, objId, 500) } catch { @@ -165,7 +186,7 @@ object Utils extends Serializable { logger.info("Error Message from Salesforce Wave " + exMsg) if (exMsg contains "Invalid Session") { logger.info("Session expired. Monitoring Job using new connection") - return monitorJob(objId, username, password, login, version) + monitorJob(objId, username, password, login, version) } else { throw uefault } @@ -180,10 +201,10 @@ object Utils extends Serializable { } def retryWithExponentialBackoff( - func:() => Boolean, - timeoutDuration: FiniteDuration, - initSleepInterval: FiniteDuration, - maxSleepInterval: FiniteDuration): Boolean = { + func: () => Boolean, + timeoutDuration: FiniteDuration, + initSleepInterval: FiniteDuration, + maxSleepInterval: FiniteDuration): Boolean = { val timeout = timeoutDuration.toMillis var waited = 0L @@ -208,7 +229,7 @@ object Utils extends Serializable { } private def monitorJob(connection: PartnerConnection, - objId: String, waitDuration: Long) : Boolean = { + objId: String, waitDuration: Long): Boolean = { val sobjects = connection.retrieve("Status", "InsightsExternalData", Array(objId)) if (sobjects != null && sobjects.length > 0) { val status = sobjects(0).getField("Status") @@ -255,20 +276,20 @@ object Utils extends Serializable { } } - private def maxWaitSeconds(waitDuration: Long) : Long = { + private def maxWaitSeconds(waitDuration: Long): Long = { // 2 Minutes val maxWaitDuration = 120000 if (waitDuration >= maxWaitDuration) maxWaitDuration else waitDuration * 2 } - private def readMetadataConfig() : Map[String, Map[String, String]] = { + private def readMetadataConfig(): Map[String, Map[String, String]] = { val source = Source.fromURL(getClass.getResource("/metadata_config.json")) val jsonContent = try source.mkString finally source.close() readJSON(jsonContent) } - private def readJSON(jsonContent : String) : Map[String, Map[String, String]]= { + private def readJSON(jsonContent: String): Map[String, Map[String, String]] = { val result = JSON.parseFull(jsonContent) val resMap: Map[String, Map[String, String]] = result.get.asInstanceOf[Map[String, Map[String, String]]] resMap