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