Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 6 additions & 5 deletions src/main/scala/com/springml/spark/salesforce/DataWriter.scala
Original file line number Diff line number Diff line change
Expand Up @@ -28,9 +28,10 @@ import org.apache.spark.sql.SaveMode
* It uses Partner External Metadata SOAP API to write the dataset
*/
class DataWriter (
val userName: String,
val userName: Option[String],
val password: String,
val login: String,
val authToken: Option[String],
val serverUrl: String,
val version: String,
val datasetName: String,
val appName: String
Expand All @@ -40,7 +41,7 @@ class DataWriter (
def writeMetadata(metaDataJson: String,
mode: SaveMode,
upsert: Boolean): Option[String] = {
val partnerConnection = createConnection(userName, password, login, version)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We need to support existing login as well. i.e with username and password

val partnerConnection = createConnection(userName, password, authToken, serverUrl, version)
val oper = operation(mode, upsert)

val sobj = new SObject()
Expand Down Expand Up @@ -87,7 +88,7 @@ class DataWriter (
sobj.setField("InsightsExternalDataId", metadataId)
sobj.setField("PartNumber", partNumber)

val partnerConnection = Utils.createConnection(userName, password, login, version)
val partnerConnection = Utils.createConnection(userName, password, authToken, serverUrl, version)
val results = partnerConnection.create(Array(sobj))

val resultSuccess = results.map(saveResult => {
Expand All @@ -108,7 +109,7 @@ class DataWriter (

def commit(id: String): Boolean = {

val partnerConnection = Utils.createConnection(userName, password, login, version)
val partnerConnection = Utils.createConnection(userName, password, authToken, serverUrl, version)

val sobj = new SObject()
sobj.setType("InsightsExternalData")
Expand Down
92 changes: 72 additions & 20 deletions src/main/scala/com/springml/spark/salesforce/DefaultSource.scala
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@ package com.springml.spark.salesforce

import java.text.SimpleDateFormat

import com.springml.salesforce.wave.api.APIFactory
import com.springml.salesforce.wave.api.{APIFactory, BulkAPI, ForceAPI, WaveAPI}
import org.apache.log4j.Logger
import org.apache.spark.sql.sources.{BaseRelation, CreatableRelationProvider, RelationProvider, SchemaRelationProvider}
import org.apache.spark.sql.types.StructType
Expand Down Expand Up @@ -50,9 +50,19 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr
*
*/
override def createRelation(sqlContext: SQLContext, parameters: Map[String, String], schema: StructType) = {
val username = param(parameters, "SF_USERNAME", "username")
val password = param(parameters, "SF_PASSWORD", "password")
val login = parameters.getOrElse("login", "https://login.salesforce.com")
val username = optionalParam(parameters, "SF_USERNAME", "username")
val authToken = optionalParam(parameters, "SF_AUTHTOKEN", "authToken")
validateMutualExclusive(username, authToken, "username", "authToken")
var password = ""
var serverUrl = ""
if (username.isDefined) {
password = param(parameters, "SF_PASSWORD", "password")
serverUrl = parameters.getOrElse("login", "https://login.salesforce.com")
} else {
serverUrl = parameters.getOrElse("instanceUrl", "https://login.salesforce.com")

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Why don't we use the existing property for instanceUrl?

}


val version = parameters.getOrElse("version", "36.0")
val saql = parameters.get("saql")
val soql = parameters.get("soql")
Expand All @@ -71,17 +81,30 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr
val inferSchemaFlag = flag(inferSchema, "inferSchema")

if (saql.isDefined) {
val waveAPI = APIFactory.getInstance.waveAPI(username, password, login, version)

var waveAPI:WaveAPI = null

if (username.isDefined) {
waveAPI = APIFactory.getInstance.waveAPI(username.get, password, serverUrl, version)
} else {
waveAPI = APIFactory.getInstance.waveAPIwAuthToken(authToken.get, serverUrl, version)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We need PR for salesforce-wave-api as well

}

DatasetRelation(waveAPI, null, saql.get, schema, sqlContext,
resultVariable, pageSize.toInt, sampleSize.toInt,
encodeFields, inferSchemaFlag, replaceDatasetNameWithId.toBoolean, sdf(dateFormat))
} else {
if (replaceDatasetNameWithId.toBoolean) {
logger.warn("Ignoring 'replaceDatasetNameWithId' option as it is not applicable to soql")
}

val forceAPI = APIFactory.getInstance.forceAPI(username, password, login,
var forceAPI: ForceAPI = null
if (username.isDefined) {
forceAPI = APIFactory.getInstance.forceAPI(username.get, password, serverUrl,
version, Integer.getInteger(pageSize), Integer.getInteger(maxRetry))
} else {
forceAPI = APIFactory.getInstance.forceAPIwAuthToken(authToken.get, serverUrl,
version, Integer.getInteger(pageSize), Integer.getInteger(maxRetry))
}
DatasetRelation(null, forceAPI, soql.get, schema, sqlContext,
null, 0, sampleSize.toInt, encodeFields, inferSchemaFlag,
replaceDatasetNameWithId.toBoolean, sdf(dateFormat))
Expand All @@ -91,12 +114,23 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr

override def createRelation(sqlContext: SQLContext, mode: SaveMode, parameters: Map[String, String], data: DataFrame): BaseRelation = {

val username = param(parameters, "SF_USERNAME", "username")
val password = param(parameters, "SF_PASSWORD", "password")
val username = optionalParam(parameters, "SF_USERNAME", "username")
val authToken = optionalParam(parameters, "SF_AUTHTOKEN", "authToken")

validateMutualExclusive(username, authToken, "username", "authToken")

var password = ""
var serverUrl = ""
if (username.isDefined) {
password = param(parameters, "SF_PASSWORD", "password")
serverUrl = parameters.getOrElse("login", "https://login.salesforce.com")
} else {
serverUrl = parameters.getOrElse("instanceUrl", "https://login.salesforce.com")
}

val datasetName = parameters.get("datasetName")
val sfObject = parameters.get("sfObject")
val appName = parameters.getOrElse("appName", null)
val login = parameters.getOrElse("login", "https://login.salesforce.com")
val version = parameters.getOrElse("version", "36.0")
val usersMetadataConfig = parameters.get("metadataConfig")
val upsert = parameters.getOrElse("upsert", "false")
Expand All @@ -115,21 +149,22 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr
}

logger.info("Writing dataframe into Salesforce Wave")
writeInSalesforceWave(username, password, login, version,
writeInSalesforceWave(username, password, authToken, serverUrl, version,
datasetName.get, appName, usersMetadataConfig, mode,
flag(upsert, "upsert"), flag(monitorJob, "monitorJob"), data, metadataFile)
} else {
logger.info("Updating Salesforce Object")
updateSalesforceObject(username, password, login, version, sfObject.get, mode, data)
updateSalesforceObject(username, password, authToken, serverUrl, version, sfObject.get, mode, data)
}

return createReturnRelation(data)
}

private def updateSalesforceObject(
username: String,
username: Option[String],
password: String,
login: String,
authToken: Option[String],
serverUrl: String,
version: String,
sfObject: String,
mode: SaveMode,
Expand All @@ -141,8 +176,13 @@ 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 bulkAPI = APIFactory.getInstance.bulkAPI(username, password, login, version)
val writer = new SFObjectWriter(username, password, login, version, sfObject, mode, csvHeader)
var bulkAPI:BulkAPI = null
if (username.isDefined) {
bulkAPI = APIFactory.getInstance.bulkAPI(username.get, password, serverUrl, version)
} else {
bulkAPI = APIFactory.getInstance.bulkAPIwAuthToken(authToken.get, serverUrl, version)
}
val writer = new SFObjectWriter(username, password, authToken, serverUrl, version, sfObject, mode, csvHeader)
logger.info("Writing data")
val successfulWrite = writer.writeData(repartitionedRDD)
logger.info(s"Writing data was successful was $successfulWrite")
Expand All @@ -153,9 +193,10 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr
}

private def writeInSalesforceWave(
username: String,
username: Option[String],
password: String,
login: String,
authToken: Option[String],
serverUrl: String,
version: String,
datasetName: String,
appName: String,
Expand All @@ -165,7 +206,7 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr
monitorJob: Boolean,
data: DataFrame,
metadata: Option[String]) {
val dataWriter = new DataWriter(username, password, login, version, datasetName, appName)
val dataWriter = new DataWriter(username, password, authToken, serverUrl, version, datasetName, appName)

val metaDataJson = Utils.metadata(metadata, usersMetadataConfig, data.schema, datasetName)

Expand Down Expand Up @@ -202,7 +243,7 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr

if (monitorJob) {
logger.info("Monitoring Job status in Salesforce wave")
if (Utils.monitorJob(writtenId.get, username, password, login, version)) {
if (Utils.monitorJob(writtenId.get, username, password, authToken, serverUrl, version)) {

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

How is the authToken expiry should be handled?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Great question, would you like automatic logic built into the spark layer or would you expect that to take place outside of spark?

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@wgorman I would like to handle in spark layer. There might be a chance that the access token got expired while fetching millions of rows

logger.info(s"Successfully dataset $datasetName has been processed in Salesforce Wave")
} else {
sys.error(s"Upload Job for dataset $datasetName failed in Salesforce Wave. Check Monitor Job in Salesforce Wave for more details")
Expand All @@ -223,6 +264,17 @@ class DefaultSource extends RelationProvider with SchemaRelationProvider with Cr
}
}


private def optionalParam(parameters: Map[String, String], envName: String, paramName: String) : Option[String] = {
val envProp = sys.env.get(envName)
if (envProp != null && envProp.isDefined) {
return Option(envProp.get)
}

val param = parameters.getOrElse(paramName, null)
return Option(param)
}

private def param(parameters: Map[String, String], envName: String, paramName: String) : String = {
val envProp = sys.env.get(envName)
if (envProp != null && envProp.isDefined) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -13,9 +13,10 @@ import com.springml.salesforce.wave.util.WaveAPIConstants
* Next subsequent columns are fields to be updated
*/
class SFObjectWriter (
val username: String,
val username: Option[String],
val password: String,
val login: String,
val authToken: Option[String],
val serverUrl: String,
val apiVersion: String,
val sfObject: String,
val mode: SaveMode,
Expand Down Expand Up @@ -63,7 +64,11 @@ class SFObjectWriter (
}

def bulkAPI() : BulkAPI = {
APIFactory.getInstance.bulkAPI(username, password, login, apiVersion)
if (username.isDefined) {
return APIFactory.getInstance.bulkAPI(username.get, password, serverUrl, apiVersion)
} else {
return APIFactory.getInstance.bulkAPIwAuthToken(authToken.get, serverUrl, apiVersion)
}
}

private def operation(mode: SaveMode): String = {
Expand Down
23 changes: 14 additions & 9 deletions src/main/scala/com/springml/spark/salesforce/Utils.scala
Original file line number Diff line number Diff line change
Expand Up @@ -37,12 +37,17 @@ import scala.util.Try
object Utils extends Serializable {
@transient val logger = Logger.getLogger("Utils")

def createConnection(username: String, password: String,
login: String, version: String):PartnerConnection = {
def createConnection(username: Option[String], password: String, authToken: Option[String],
serverUrl: String, version: String):PartnerConnection = {
val config = new ConnectorConfig()
config.setUsername(username)
config.setPassword(password)
val endpoint = if (login.endsWith("/")) (login + "services/Soap/u/" + version) else (login + "/services/Soap/u/" + version)
if (username.isDefined) {
config.setUsername(username.get)
config.setPassword(password)
} else {
config.setManualLogin(true)
config.setSessionId(authToken.get)
}
val endpoint = if (serverUrl.endsWith("/")) (serverUrl + "services/Soap/u/" + version) else (serverUrl + "/services/Soap/u/" + version)
config.setAuthEndpoint(endpoint)
config.setServiceEndpoint(endpoint)
Connector.newConnection(config)
Expand Down Expand Up @@ -150,9 +155,9 @@ 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)
def monitorJob(objId: String, username: Option[String], password:
String, authToken:Option[String], login: String, version: String) : Boolean = {
var partnerConnection = Utils.createConnection(username, password, authToken, login, version)
try {
monitorJob(partnerConnection, objId, 500)
} catch {
Expand All @@ -161,7 +166,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)
return monitorJob(objId, username, password, authToken, login, version)
} else {
throw uefault
}
Expand Down